fix(metadata): keep identity selectors from leaking unselected prompts

This commit is contained in:
Will Miao
2026-08-12 19:44:43 +08:00
parent 5bc6d8286c
commit c2f16784b3
2 changed files with 205 additions and 1 deletions
+23 -1
View File
@@ -104,7 +104,11 @@ class GenericNodeExtractor(NodeMetadataExtractor):
if not output_tuple or len(output_tuple) < 1: if not output_tuple or len(output_tuple) < 1:
return return
output_conditioning = output_tuple[0] conditioning_index = _first_conditioning_index(return_types)
if conditioning_index is None or len(output_tuple) <= conditioning_index:
return
output_conditioning = output_tuple[conditioning_index]
if output_conditioning is None: if output_conditioning is None:
return return
@@ -430,6 +434,16 @@ def _first_output_tuple(outputs):
return None return None
def _first_conditioning_index(return_types):
"""Return the index of the first CONDITIONING output slot, or None."""
if not return_types:
return None
for index, return_type in enumerate(return_types):
if "CONDITIONING" in str(return_type):
return index
return None
def _collect_conditioning_inputs(inputs): def _collect_conditioning_inputs(inputs):
"""Collect conditioning object inputs (``conditioning*`` keys). """Collect conditioning object inputs (``conditioning*`` keys).
@@ -460,6 +474,14 @@ def _record_conditioning_source(
if not sources: if not sources:
return return
# Identity-preserving selectors return one of their inputs unchanged:
# only that input contributed to the output, so record it alone instead
# of treating every input as a combination source.
for conditioning in sources:
if id(conditioning) == id(output_conditioning):
sources = [conditioning]
break
prompt_metadata = _ensure_prompt_metadata(metadata, node_id) prompt_metadata = _ensure_prompt_metadata(metadata, node_id)
prompt_metadata.setdefault("conditioning_sources", []).append( prompt_metadata.setdefault("conditioning_sources", []).append(
{ {
@@ -765,6 +765,188 @@ def test_conditioning_provenance_ignores_scalar_conditioning_fields(
] ]
def test_conditioning_provenance_selector_with_conditioning_named_inputs(
metadata_registry, monkeypatch
):
"""An identity selector whose inputs use ``conditioning*`` names must not
leak the unselected branch's prompt."""
prompt_graph = {
"encode_a": {
"class_type": "CLIPTextEncode",
"inputs": {"text": "AAA", "clip": ["clip", 0]},
},
"encode_b": {
"class_type": "CLIPTextEncode",
"inputs": {"text": "BBB", "clip": ["clip", 0]},
},
"selector": {
"class_type": "ConditioningSelector",
"inputs": {
"conditioning_a": ["encode_a", 0],
"conditioning_b": ["encode_b", 0],
},
},
"sampler": {
"class_type": "ClownsharKSampler_Beta",
"inputs": {
"seed": 123,
"steps": 8,
"cfg": 1.0,
"sampler_name": "linear/euler",
"scheduler": "beta57",
"denoise": 1.0,
"positive": ["selector", 0],
"negative": ["encode_b", 0],
"latent_image": {
"samples": types.SimpleNamespace(shape=(1, 4, 16, 16))
},
},
},
}
prompt = SimpleNamespace(original_prompt=prompt_graph)
conditioning_a = object()
conditioning_b = object()
monkeypatch.setattr(metadata_processor, "standalone_mode", False)
metadata_registry.start_collection("prompt-selector")
metadata_registry.set_current_prompt(prompt)
metadata_registry.record_node_execution(
"encode_a", "CLIPTextEncode", {"text": "AAA"}, None
)
metadata_registry.update_node_execution(
"encode_a", "CLIPTextEncode", [(conditioning_a,)]
)
metadata_registry.record_node_execution(
"encode_b", "CLIPTextEncode", {"text": "BBB"}, None
)
metadata_registry.update_node_execution(
"encode_b", "CLIPTextEncode", [(conditioning_b,)]
)
metadata_registry.record_node_execution(
"selector",
"ConditioningSelector",
{"conditioning_a": conditioning_a, "conditioning_b": conditioning_b},
None,
return_types=("CONDITIONING",),
)
metadata_registry.update_node_execution(
"selector", "ConditioningSelector", [(conditioning_a,)],
return_types=("CONDITIONING",),
)
metadata_registry.record_node_execution(
"sampler",
"ClownsharKSampler_Beta",
{
"seed": 123,
"steps": 8,
"cfg": 1.0,
"sampler_name": "linear/euler",
"scheduler": "beta57",
"denoise": 1.0,
"positive": conditioning_a,
"negative": conditioning_b,
"latent_image": {
"samples": types.SimpleNamespace(shape=(1, 4, 16, 16))
},
},
None,
)
metadata = metadata_registry.get_metadata("prompt-selector")
params = MetadataProcessor.extract_generation_params(metadata)
assert params["prompt"] == "AAA"
assert params["negative_prompt"] == "BBB"
def test_conditioning_provenance_uses_conditioning_output_slot(
metadata_registry, monkeypatch
):
"""Unregistered nodes whose CONDITIONING output is not the first slot
must still be tracked through the correct output position.
The graph's conditioning chain ends at an unexecuted phantom node so the
topology fallback in extract_generation_params cannot mask a runtime
provenance failure.
"""
prompt_graph = {
"diag_node": {
"class_type": "DiagThenCond",
"inputs": {"conditioning": ["phantom_source", 0]},
},
"sampler": {
"class_type": "ClownsharKSampler_Beta",
"inputs": {
"seed": 123,
"steps": 8,
"cfg": 1.0,
"sampler_name": "linear/euler",
"scheduler": "beta57",
"denoise": 1.0,
"positive": ["diag_node", 1],
"latent_image": {
"samples": types.SimpleNamespace(shape=(1, 4, 16, 16))
},
},
},
}
prompt = SimpleNamespace(original_prompt=prompt_graph)
input_conditioning = object()
transformed_conditioning = object()
monkeypatch.setattr(metadata_processor, "standalone_mode", False)
metadata_registry.start_collection("prompt-output-slot")
metadata_registry.set_current_prompt(prompt)
metadata_registry.record_node_execution(
"encode_pos", "CLIPTextEncode", {"text": "AAA"}, None
)
metadata_registry.update_node_execution(
"encode_pos", "CLIPTextEncode", [(input_conditioning,)]
)
metadata_registry.record_node_execution(
"diag_node",
"DiagThenCond",
{"conditioning": input_conditioning},
None,
return_types=("STRING", "CONDITIONING"),
)
metadata_registry.update_node_execution(
"diag_node",
"DiagThenCond",
[("diagnostics", transformed_conditioning)],
return_types=("STRING", "CONDITIONING"),
)
metadata_registry.record_node_execution(
"sampler",
"ClownsharKSampler_Beta",
{
"seed": 123,
"steps": 8,
"cfg": 1.0,
"sampler_name": "linear/euler",
"scheduler": "beta57",
"denoise": 1.0,
"positive": transformed_conditioning,
"negative": input_conditioning,
"latent_image": {
"samples": types.SimpleNamespace(shape=(1, 4, 16, 16))
},
},
None,
)
metadata = metadata_registry.get_metadata("prompt-output-slot")
params = MetadataProcessor.extract_generation_params(metadata)
assert params["prompt"] == "AAA"
def test_conditioning_provenance_recovers_kj_set_get_prompts( def test_conditioning_provenance_recovers_kj_set_get_prompts(
metadata_registry, monkeypatch metadata_registry, monkeypatch
): ):