diff --git a/backend/app/core/agent_service.py b/backend/app/core/agent_service.py index 6cb3b42..7f4e3a0 100644 --- a/backend/app/core/agent_service.py +++ b/backend/app/core/agent_service.py @@ -1,7 +1,8 @@ -"""Conversational Agent Service for workflow and engine control (M9). +"""Rule- and regex-based conversational assistant service for workflow and engine control (M9). -Provides natural-language intent parsing, model/engine recommendation, transparent -action plan formulation, and human-in-the-loop proposal execution. +Provides deterministic intent matching, parameter extraction, model/engine recommendation, transparent +action plan formulation, and human-in-the-loop proposal execution. Bounded heuristic capabilities; +does not claim unconstrained natural-language understanding or arbitrary workflow synthesis. """ import re @@ -164,7 +165,8 @@ async def process_chat(self, req: AgentChatRequest) -> AgentChatResponse: "Berry Studio includes an automated ComfyUI DAG validator and repair engine (M10). " "You can submit workflows via `/api/v1/workflow/validate` to inspect topological cycles, " "missing MODEL/CLIP/VAE connections, and missing checkpoint dependencies. " - "The repair engine automatically reconnects broken links and substitutes indexed checkpoints." + "The repair engine reconnects unambiguous broken links and substitutes family-compatible " + "indexed checkpoints while blocking ambiguous or cross-family mutations." ) assistant_msg = AgentChatMessage(role="assistant", content=content) history.append(assistant_msg) diff --git a/backend/app/core/workflow_repair.py b/backend/app/core/workflow_repair.py index c587641..f805b50 100644 --- a/backend/app/core/workflow_repair.py +++ b/backend/app/core/workflow_repair.py @@ -2,6 +2,7 @@ Applies deterministic heuristics to repair disconnected ports, missing model references, and broken graph topology in ComfyUI prompt workflows. +Strictly prevents automatic model or port substitutions when compatibility is uncertain. """ import copy @@ -17,6 +18,22 @@ ) +def infer_model_family(name: str) -> str: + """Infer model architecture family from filename.""" + lower = name.lower() + if any(k in lower for k in ["xl", "sdxl", "base-1.0", "refiner"]): + return "sdxl" + if "flux" in lower: + return "flux" + if "svd" in lower: + return "svd" + if any(k in lower for k in ["v1-5", "sd15", "1.5", "v1.5", "stable-diffusion-v1"]): + return "sd15" + if "sd2" in lower or "v2-1" in lower: + return "sd2" + return "unknown" + + class WorkflowRepairer: def __init__(self, model_catalog=None): self.model_catalog = model_catalog @@ -36,9 +53,9 @@ def _find_nodes_with_output(self, workflow: Dict[str, Any], port_type: str) -> L return matches def repair(self, workflow: Dict[str, Any]) -> WorkflowRepairResult: - """Repair issues in a ComfyUI prompt workflow.""" + """Repair issues in a ComfyUI prompt workflow with strict compatibility guards.""" orig_report = self.validator.validate(workflow) - if orig_report.valid: + if orig_report.valid and not orig_report.missing_models: return WorkflowRepairResult( success=True, original_valid=True, @@ -51,32 +68,87 @@ def repair(self, workflow: Dict[str, Any]) -> WorkflowRepairResult: repaired = copy.deepcopy(workflow) repairs: List[RepairAction] = [] - # 1. Substitute missing checkpoints if available in model catalog + # 1. Substitute missing checkpoints ONLY when architecture compatibility is certain if self.model_catalog: try: indexed_models = self.model_catalog.get_models() - if indexed_models: - default_ckpt = indexed_models[0].name - for nid, node_data in repaired.items(): - if node_data.get("class_type") in ("CheckpointLoaderSimple", "LoadCheckpoint"): - cur_ckpt = node_data.get("inputs", {}).get("ckpt_name", "") - indexed_names = [m.name for m in indexed_models] - if (not cur_ckpt) or (cur_ckpt not in indexed_names): - node_data.setdefault("inputs", {})["ckpt_name"] = default_ckpt + indexed_names = [m.name for m in indexed_models] + + for nid, node_data in repaired.items(): + if node_data.get("class_type") in ("CheckpointLoaderSimple", "LoadCheckpoint"): + cur_ckpt = node_data.get("inputs", {}).get("ckpt_name", "") + + if not cur_ckpt: + # No model specified at all: pick first indexed model if available + if indexed_models: + fallback_ckpt = indexed_models[0].name + node_data.setdefault("inputs", {})["ckpt_name"] = fallback_ckpt repairs.append( RepairAction( node_id=str(nid), node_class=node_data["class_type"], action_type="replace_model", - description=f"Substituted unavailable checkpoint '{cur_ckpt}' with indexed model '{default_ckpt}'.", + description=f"Assigned available indexed model '{fallback_ckpt}' to empty checkpoint slot.", + target_input="ckpt_name", + new_value=fallback_ckpt, + ) + ) + elif cur_ckpt not in indexed_names: + # A specific model was requested but is not in the catalog + cur_family = infer_model_family(cur_ckpt) + if cur_family == "unknown": + # Unknown architecture family: DO NOT guess or substitute! + repairs.append( + RepairAction( + node_id=str(nid), + node_class=node_data["class_type"], + action_type="substitution_blocked", + description=( + f"Model substitution blocked for '{cur_ckpt}': architecture family is unknown. " + "Automatic replacement prevented to avoid graph corruption." + ), target_input="ckpt_name", - new_value=default_ckpt, ) ) + else: + # Find an indexed model of the exact same family + matching_candidates = [ + m.name for m in indexed_models if infer_model_family(m.name) == cur_family + ] + if matching_candidates: + sub_model = matching_candidates[0] + node_data.setdefault("inputs", {})["ckpt_name"] = sub_model + repairs.append( + RepairAction( + node_id=str(nid), + node_class=node_data["class_type"], + action_type="replace_model", + description=( + f"Substituted unavailable {cur_family.upper()} checkpoint '{cur_ckpt}' " + f"with compatible indexed model '{sub_model}'." + ), + target_input="ckpt_name", + new_value=sub_model, + ) + ) + else: + # No model in the matching family exists + repairs.append( + RepairAction( + node_id=str(nid), + node_class=node_data["class_type"], + action_type="substitution_blocked", + description=( + f"Model substitution blocked for '{cur_ckpt}': no indexed model in the " + f"matching '{cur_family.upper()}' family is available." + ), + target_input="ckpt_name", + ) + ) except Exception: pass - # 2. Repair missing required connections + # 2. Repair missing required connections with ambiguity guards vae_sources = self._find_nodes_with_output(repaired, "VAE") clip_sources = self._find_nodes_with_output(repaired, "CLIP") model_sources = self._find_nodes_with_output(repaired, "MODEL") @@ -92,7 +164,7 @@ def repair(self, workflow: Dict[str, Any]) -> WorkflowRepairResult: # Repair VAEDecode / VAEEncode missing VAE if cls_name in ("VAEDecode", "VAEEncode"): if "vae" not in inputs or not inputs["vae"]: - if vae_sources: + if len(vae_sources) == 1: src_id, slot_idx = vae_sources[0] inputs["vae"] = [src_id, slot_idx] repairs.append( @@ -100,17 +172,27 @@ def repair(self, workflow: Dict[str, Any]) -> WorkflowRepairResult: node_id=str(nid), node_class=cls_name, action_type="reconnect_slot", - description=f"Connected missing 'vae' input to node '{src_id}' [slot {slot_idx}].", + description=f"Connected missing 'vae' input to unambiguous source node '{src_id}' [slot {slot_idx}].", target_input="vae", source_node_id=src_id, source_output_slot=slot_idx, ) ) + elif len(vae_sources) > 1: + repairs.append( + RepairAction( + node_id=str(nid), + node_class=cls_name, + action_type="substitution_blocked", + description=f"Automatic VAE connection blocked: {len(vae_sources)} ambiguous VAE sources detected.", + target_input="vae", + ) + ) # Repair CLIPTextEncode missing CLIP if cls_name == "CLIPTextEncode": if "clip" not in inputs or not inputs["clip"]: - if clip_sources: + if len(clip_sources) == 1: src_id, slot_idx = clip_sources[0] inputs["clip"] = [src_id, slot_idx] repairs.append( @@ -118,17 +200,27 @@ def repair(self, workflow: Dict[str, Any]) -> WorkflowRepairResult: node_id=str(nid), node_class=cls_name, action_type="reconnect_slot", - description=f"Connected missing 'clip' input to node '{src_id}' [slot {slot_idx}].", + description=f"Connected missing 'clip' input to unambiguous source node '{src_id}' [slot {slot_idx}].", target_input="clip", source_node_id=src_id, source_output_slot=slot_idx, ) ) + elif len(clip_sources) > 1: + repairs.append( + RepairAction( + node_id=str(nid), + node_class=cls_name, + action_type="substitution_blocked", + description=f"Automatic CLIP connection blocked: {len(clip_sources)} ambiguous CLIP sources detected.", + target_input="clip", + ) + ) # Repair KSampler missing MODEL if cls_name in ("KSampler", "KSamplerAdvanced"): if "model" not in inputs or not inputs["model"]: - if model_sources: + if len(model_sources) == 1: src_id, slot_idx = model_sources[0] inputs["model"] = [src_id, slot_idx] repairs.append( @@ -136,24 +228,32 @@ def repair(self, workflow: Dict[str, Any]) -> WorkflowRepairResult: node_id=str(nid), node_class=cls_name, action_type="reconnect_slot", - description=f"Connected missing 'model' input to node '{src_id}' [slot {slot_idx}].", + description=f"Connected missing 'model' input to unambiguous source node '{src_id}' [slot {slot_idx}].", target_input="model", source_node_id=src_id, source_output_slot=slot_idx, ) ) + elif len(model_sources) > 1: + repairs.append( + RepairAction( + node_id=str(nid), + node_class=cls_name, + action_type="substitution_blocked", + description=f"Automatic MODEL connection blocked: {len(model_sources)} ambiguous MODEL sources detected.", + target_input="model", + ) + ) # Repair KSampler missing latent_image if "latent_image" not in inputs or not inputs["latent_image"]: - # Look for EmptyLatentImage first, or any latent source empty_latents = [ (s_id, s_slot) for s_id, s_slot in latent_sources if repaired.get(s_id, {}).get("class_type") == "EmptyLatentImage" ] - target = empty_latents[0] if empty_latents else (latent_sources[0] if latent_sources else None) - if target: - src_id, slot_idx = target + if len(empty_latents) == 1: + src_id, slot_idx = empty_latents[0] inputs["latent_image"] = [src_id, slot_idx] repairs.append( RepairAction( @@ -166,9 +266,19 @@ def repair(self, workflow: Dict[str, Any]) -> WorkflowRepairResult: source_output_slot=slot_idx, ) ) + elif len(empty_latents) > 1: + repairs.append( + RepairAction( + node_id=str(nid), + node_class=cls_name, + action_type="substitution_blocked", + description=f"Automatic latent connection blocked: {len(empty_latents)} ambiguous EmptyLatentImage sources detected.", + target_input="latent_image", + ) + ) # Repair KSampler missing positive / negative conditioning - if ("positive" not in inputs or not inputs["positive"]) and len(cond_sources) >= 1: + if ("positive" not in inputs or not inputs["positive"]) and len(cond_sources) == 2: src_id, slot_idx = cond_sources[0] inputs["positive"] = [src_id, slot_idx] repairs.append( @@ -183,7 +293,7 @@ def repair(self, workflow: Dict[str, Any]) -> WorkflowRepairResult: ) ) - if ("negative" not in inputs or not inputs["negative"]) and len(cond_sources) >= 2: + if ("negative" not in inputs or not inputs["negative"]) and len(cond_sources) == 2: src_id, slot_idx = cond_sources[1] inputs["negative"] = [src_id, slot_idx] repairs.append( diff --git a/backend/tests/test_m10_workflow_repair.py b/backend/tests/test_m10_workflow_repair.py index 513e0f0..c945a7f 100644 --- a/backend/tests/test_m10_workflow_repair.py +++ b/backend/tests/test_m10_workflow_repair.py @@ -186,3 +186,91 @@ def test_api_workflow_validate_and_repair_endpoints(): assert rep_data["success"] is True assert rep_data["repaired_valid"] is True assert rep_data["repaired_workflow"]["8"]["inputs"]["vae"] == ["4", 2] + + +def test_repairer_blocks_ambiguous_connections(): + """Repairer blocks automatic reconnection when multiple ambiguous sources exist.""" + broken = build_valid_comfyui_graph() + # Add a second VAE loader to create ambiguity + broken["10"] = { + "class_type": "VAELoader", + "inputs": {"vae_name": "vae-ft-mse-840000-ema-pruned.safetensors"}, + } + del broken["8"]["inputs"]["vae"] + + repairer = WorkflowRepairer() + result = repairer.repair(broken) + + # Should have recorded a substitution_blocked action + blocked_repairs = [r for r in result.repairs_applied if r.action_type == "substitution_blocked"] + assert len(blocked_repairs) >= 1 + assert "ambiguous" in blocked_repairs[0].description.lower() + # The slot should not have been blindly wired + assert "vae" not in result.repaired_workflow["8"]["inputs"] + + +def test_repairer_blocks_cross_family_checkpoint_substitution(): + """Repairer blocks model substitution if indexed models belong to an incompatible or unknown family.""" + from unittest.mock import MagicMock + from app.schemas.model import ModelRecord, ModelCategory, ModelArchitecture + + broken = build_valid_comfyui_graph() + # Workflow expects a Flux model + broken["4"]["inputs"]["ckpt_name"] = "flux1-schnell.safetensors" + + # Catalog only has an SD 1.5 model + mock_catalog = MagicMock() + mock_catalog.get_models.return_value = [ + ModelRecord( + id="sd15", + name="v1-5-pruned-emaonly.safetensors", + file_path="models/v1-5-pruned-emaonly.safetensors", + category=ModelCategory.CHECKPOINT, + architecture=ModelArchitecture.SD15, + format="safetensors", + size_bytes=4000000000, + size_mb=4000.0, + ) + ] + + repairer = WorkflowRepairer(model_catalog=mock_catalog) + result = repairer.repair(broken) + + # Cross-family substitution must be blocked + blocked = [r for r in result.repairs_applied if r.action_type == "substitution_blocked"] + assert len(blocked) >= 1 + assert "flux" in blocked[0].description.lower() + # The checkpoint in workflow should remain untouched (not mutated to SD1.5) + assert result.repaired_workflow["4"]["inputs"]["ckpt_name"] == "flux1-schnell.safetensors" + + +def test_repairer_allows_same_family_checkpoint_substitution(): + """Repairer substitutes missing checkpoint when a candidate with matching family is indexed.""" + from unittest.mock import MagicMock + from app.schemas.model import ModelRecord, ModelCategory, ModelArchitecture + + broken = build_valid_comfyui_graph() + # Workflow has missing SD 1.5 checkpoint + broken["4"]["inputs"]["ckpt_name"] = "dreamshaper_8_sd15.safetensors" + + mock_catalog = MagicMock() + mock_catalog.get_models.return_value = [ + ModelRecord( + id="sd15_alt", + name="v1-5-pruned-emaonly.safetensors", + file_path="models/v1-5-pruned-emaonly.safetensors", + category=ModelCategory.CHECKPOINT, + architecture=ModelArchitecture.SD15, + format="safetensors", + size_bytes=4000000000, + size_mb=4000.0, + ) + ] + + repairer = WorkflowRepairer(model_catalog=mock_catalog) + result = repairer.repair(broken) + + substitutions = [r for r in result.repairs_applied if r.action_type == "replace_model"] + assert len(substitutions) == 1 + assert result.repaired_workflow["4"]["inputs"]["ckpt_name"] == "v1-5-pruned-emaonly.safetensors" +