fix(metadata): record the actual checkpoint on Set/Get and stale-cache workflows

trace_model_path dead-ended on KJNodes GetNode virtual links and fell back
to the first entry in MODELS, which could be a stale checkpoint name
resurrected from the process-lifetime node_cache — images ended up with a
wrong or missing Hashes.model.

- SetNodeExtractor now records MODEL passthrough as a model_variable entry
  (variable name -> source node id), cached per node like other metadata
- trace_model_path resolves GetNode variable references through those
  entries instead of giving up (last SetNode wins, per KJNodes semantics)
- _fill_missing_metadata validates cached checkpoint names against the
  node's current widget values and rebuilds stale entries from the current
  inputs instead of resurrecting the previous model
This commit is contained in:
Will Miao
2026-10-09 20:44:39 +08:00
parent 1a6b0f78a3
commit 406096c0e1
4 changed files with 370 additions and 7 deletions
+90 -4
View File
@@ -1,8 +1,47 @@
import os
import re
import time
from typing import Any
from nodes import NODE_CLASS_MAPPINGS # pyright: ignore[reportMissingImports, reportAttributeAccessIssue]
from .node_extractors import NODE_EXTRACTORS, GenericNodeExtractor
from .constants import METADATA_CATEGORIES, IMAGES, OVERWRITE
from .constants import METADATA_CATEGORIES, IMAGES, MODELS, OVERWRITE
# Input fields that carry a model filename in loader-style nodes. Mirrors
# GenericNodeExtractor._MODEL_NAME_FIELDS, plus the TensorRT engine extension.
_MODEL_NAME_FIELDS = ("ckpt_name", "unet_name", "model_path", "model_name", "gguf_name")
_MODEL_FILE_EXTENSIONS = (
".ckpt", ".pt", ".pt2", ".bin", ".pth",
".safetensors", ".pkl", ".sft", ".gguf", ".engine",
)
def _checkpoint_name_candidates(inputs):
"""Derive the checkpoint names a loader node's current inputs could record.
Used to keep stale ``node_cache`` entries (from before the user switched
models) out of prompts they no longer belong to. Returns an empty set when
the inputs carry no recognizable model field.
"""
candidates = set()
for field in _MODEL_NAME_FIELDS:
value = inputs.get(field)
if not isinstance(value, str) or not value.strip():
continue
name = value.strip()
if not name.lower().endswith(_MODEL_FILE_EXTENSIONS):
continue
candidates.add(name)
base = os.path.splitext(os.path.basename(name))[0]
candidates.add(base)
# TensorRTLoaderExtractor derivation: drop the "_$profile" part and
# any trailing save counter (e.g. "_00001_").
derived = base
if "_$" in derived:
derived = derived[: derived.index("_$")]
derived = re.sub(r"_\d+_?$", "", derived)
candidates.add(derived)
return candidates
class MetadataRegistry:
@@ -155,9 +194,56 @@ class MetadataRegistry:
continue
if category in cached_data and node_id in cached_data[category]:
if node_id not in metadata[category]:
metadata[category][node_id] = cached_data[category][
node_id
]
cached_entry = cached_data[category][node_id]
if category == MODELS:
cached_entry = self._validated_checkpoint_fill(
cached_entry, node_data
)
if cached_entry is None:
continue
metadata[category][node_id] = cached_entry
@staticmethod
def _validated_checkpoint_fill(cached_entry, node_data):
"""Guard a MODELS cache fill against stale checkpoint names.
A loader that has not executed since the user switched models still
holds the previous model in its cache entry. When the node's current
inputs name a different file, the cached name is replaced with the one
derived from the current inputs; when no current name can be derived,
the stale entry is dropped entirely.
The cache is trusted as-is when:
* the entry is not a checkpoint record,
* the cached name carries no model-file extension (e.g. names set
from runtime attachments rather than a widget value), or
* the node inputs carry no recognizable model field (unknown loader).
"""
if not isinstance(cached_entry, dict):
return cached_entry
if cached_entry.get("type") != "checkpoint":
return cached_entry
cached_name = cached_entry.get("name")
if not cached_name or not cached_name.lower().endswith(_MODEL_FILE_EXTENSIONS):
return cached_entry
inputs = node_data.get("inputs", {})
candidates = _checkpoint_name_candidates(inputs)
if not candidates:
return cached_entry
if cached_name in candidates:
return cached_entry
# Stale — rebuild from the current inputs when possible.
for field in _MODEL_NAME_FIELDS:
value = inputs.get(field)
if isinstance(value, str) and value.strip().lower().endswith(
_MODEL_FILE_EXTENSIONS
):
rebuilt = dict(cached_entry)
rebuilt["name"] = value.strip()
return rebuilt
return None
def record_node_execution(self, node_id, class_type, inputs, outputs, return_types=None):
"""Record information about a node's execution"""