From e7686e0fe9f478ee0868172013c967c165430680 Mon Sep 17 00:00:00 2001 From: Luke Piette Date: Mon, 21 Sep 2026 11:26:58 -0700 Subject: [PATCH] feat: validate workflow model references before submitting to ComfyUI Workflows referencing checkpoints/LoRAs/VAE/CLIP/UNet weights that are not on the worker previously failed only after the cold start and submit, with ComfyUI's cryptic "Value not in list" / prompt_outputs_failed_validation errors. Some clients also send the literal '__list__' UI placeholder as a model name. Add a pre-flight check in handler() that walks the workflow graph before queue_workflow() and validates every known loader node's model input against ComfyUI's /object_info. All missing files are collected into one clear error naming each file, its model type, and the expected network volume directory; the '__list__' placeholder gets a specific message. Matching is case-sensitive (mirroring ComfyUI) with a did-you-mean hint on case-only mismatches. If /object_info can't be fetched the check is skipped (fail open) and queue_workflow()'s existing 400 enrichment remains as the backstop. get_available_models() is generalized from CheckpointLoaderSimple-only to all known loader node types. Also repairs tests/test_handler.py, which no longer imported (handler.py moved to the repo root) and tested functions that no longer exist; CI tests are currently a placeholder so this had gone unnoticed. Co-Authored-By: Claude Fable 5 --- handler.py | 282 +++++++++++++++++++++++-- tests/test_handler.py | 468 +++++++++++++++++++++++++++++------------- 2 files changed, 593 insertions(+), 157 deletions(-) diff --git a/handler.py b/handler.py index 817cad8a9..ec80043dc 100644 --- a/handler.py +++ b/handler.py @@ -59,6 +59,43 @@ # see https://docs.runpod.io/docs/handler-additional-controls#refresh-worker REFRESH_WORKER = os.environ.get("REFRESH_WORKER", "false").lower() == "true" +# --------------------------------------------------------------------------- +# Model loader nodes — used for pre-flight validation of workflow model refs +# --------------------------------------------------------------------------- +# Maps a loader node's class_type to the model type it loads and the input +# field(s) that name a model file. The node → directory mapping is ComfyUI +# domain knowledge (documented in AGENTS.md), not discoverable from code. +MODEL_LOADER_NODES = { + "CheckpointLoaderSimple": ("checkpoints", ("ckpt_name",)), + "LoraLoader": ("loras", ("lora_name",)), + "VAELoader": ("vae", ("vae_name",)), + "DualCLIPLoader": ("text_encoders", ("clip_name1", "clip_name2")), + "TripleCLIPLoader": ("text_encoders", ("clip_name1", "clip_name2", "clip_name3")), + "UNETLoader": ("diffusion_models", ("unet_name",)), + "UnetLoaderGGUF": ("diffusion_models", ("unet_name",)), + "Hy3DModelLoader": ("diffusion_models", ("model",)), + "UpscaleModelLoader": ("upscale_models", ("model_name",)), +} + +# Where each model type lives on the network volume (see +# src/extra_model_paths.yaml — text encoders mount under clip/ and diffusion +# models under unet/, their legacy ComfyUI directory names). +MODEL_TYPE_VOLUME_DIRS = { + "checkpoints": "/runpod-volume/models/checkpoints/", + "loras": "/runpod-volume/models/loras/", + "vae": "/runpod-volume/models/vae/", + "text_encoders": "/runpod-volume/models/clip/", + "diffusion_models": "/runpod-volume/models/unet/", + "upscale_models": "/runpod-volume/models/upscale_models/", +} + +# Placeholder some clients send when a UI model dropdown was never resolved +# to a real filename. +MODEL_LIST_PLACEHOLDER = "__list__" + +# Cap how many available files an error message lists before truncating +MAX_LISTED_MODELS = 15 + # --------------------------------------------------------------------------- # Helper: quick reachability probe of ComfyUI HTTP endpoint (port 8188) # --------------------------------------------------------------------------- @@ -373,34 +410,234 @@ def upload_images(images): } -def get_available_models(): +def _fetch_object_info(): """ - Get list of available models from ComfyUI + Fetch ComfyUI's /object_info (registered nodes and their input options). Returns: - dict: Dictionary containing available models by type + dict: The parsed /object_info payload, or None if it can't be fetched. """ try: response = requests.get(f"http://{COMFY_HOST}/object_info", timeout=10) response.raise_for_status() - object_info = response.json() - - # Extract available checkpoints from CheckpointLoaderSimple - available_models = {} - if "CheckpointLoaderSimple" in object_info: - checkpoint_info = object_info["CheckpointLoaderSimple"] - if "input" in checkpoint_info and "required" in checkpoint_info["input"]: - ckpt_options = checkpoint_info["input"]["required"].get("ckpt_name") - if ckpt_options and len(ckpt_options) > 0: - available_models["checkpoints"] = ( - ckpt_options[0] if isinstance(ckpt_options[0], list) else [] - ) - - return available_models + return response.json() except Exception as e: - print(f"worker-comfyui - Warning: Could not fetch available models: {e}") + print(f"worker-comfyui - Warning: Could not fetch /object_info: {e}") + return None + + +def _loader_field_options(object_info, class_type, field): + """ + Return the list of valid options for a node input field, or None when the + node/field is not registered or the field is not an option list. + """ + node_info = object_info.get(class_type) + if not isinstance(node_info, dict): + return None + inputs = node_info.get("input", {}) + if not isinstance(inputs, dict): + return None + for section in ("required", "optional"): + section_inputs = inputs.get(section) + if not isinstance(section_inputs, dict): + continue + spec = section_inputs.get(field) + # Option-list inputs look like [["file_a", "file_b"], {...}] + if ( + isinstance(spec, (list, tuple)) + and len(spec) > 0 + and isinstance(spec[0], list) + ): + return spec[0] + return None + + +def get_available_models(): + """ + Get list of available models from ComfyUI, grouped by model type. + + Queries /object_info and extracts the option list of every known model + loader node (see MODEL_LOADER_NODES), so the result covers checkpoints, + loras, vae, text encoders, diffusion models and upscale models. + + Returns: + dict: Mapping of model type (e.g. 'checkpoints') to a sorted list of + available filenames. Empty dict if ComfyUI can't be reached. + """ + object_info = _fetch_object_info() + if object_info is None: return {} + available_models = {} + for class_type, (model_type, fields) in MODEL_LOADER_NODES.items(): + for field in fields: + options = _loader_field_options(object_info, class_type, field) + if options: + available_models.setdefault(model_type, set()).update(options) + return {model_type: sorted(files) for model_type, files in available_models.items()} + + +def _format_options(options): + """Format an option list for an error message, truncating long lists.""" + shown = sorted(options)[:MAX_LISTED_MODELS] + formatted = ", ".join(f"'{o}'" for o in shown) + remaining = len(options) - len(shown) + if remaining > 0: + formatted += f" (and {remaining} more)" + return formatted + + +def _check_model_reference(object_info, node_id, class_type, field, value, model_type): + """ + Validate one model-name input of a loader node. + + Returns a problem description string, or None if the reference is fine + (or cannot be judged from /object_info). + """ + volume_dir = MODEL_TYPE_VOLUME_DIRS.get( + model_type, f"/runpod-volume/models/{model_type}/" + ) + + if value == MODEL_LIST_PLACEHOLDER: + available = _loader_field_options(object_info, class_type, field) or [] + msg = ( + f"Node {node_id} ({class_type}.{field}): got the placeholder " + f"'{MODEL_LIST_PLACEHOLDER}' instead of a real model filename — the client " + f"sent a UI default that was never replaced with an actual file. " + f"Set it to one of the available {model_type} files" + ) + if available: + msg += f": {_format_options(available)}" + return msg + + available = _loader_field_options(object_info, class_type, field) + if available is None: + # Node type or field not registered in this ComfyUI build — let + # ComfyUI's own validation decide instead of guessing. + return None + if value in available: + return None + + msg = ( + f"Node {node_id} ({class_type}.{field}): '{value}' not found in " + f"{model_type}. Expected at {volume_dir}{value}" + ) + # Matching mirrors ComfyUI's own validation and is case-sensitive; when a + # file differs only in case, say so instead of just "not found". + case_match = next((a for a in available if a.lower() == value.lower()), None) + if case_match: + msg += f". Did you mean '{case_match}'? (filenames are case-sensitive)" + elif available: + msg += f". Available {model_type}: {_format_options(available)}" + else: + msg += ( + f". No {model_type} files are available — check that your network " + f"volume contains them under {volume_dir}" + ) + return msg + + +def _check_image_reference(object_info, node_id, value): + """ + Validate a LoadImage 'image' input against ComfyUI's known input images. + + Returns a problem description string, or None if the reference is fine. + """ + if value == MODEL_LIST_PLACEHOLDER: + return ( + f"Node {node_id} (LoadImage.image): got the placeholder " + f"'{MODEL_LIST_PLACEHOLDER}' instead of a real image filename — the client " + f"sent a UI default that was never replaced with an actual file" + ) + available = _loader_field_options(object_info, "LoadImage", "image") + if available is None or value in available: + return None + # ComfyUI accepts annotated names like "example.png [input]" that are not + # part of the option list — don't second-guess those. + if value.endswith("]") and " [" in value: + return None + return ( + f"Node {node_id} (LoadImage.image): '{value}' is not an available input " + f"image. Upload it via the request's 'images' parameter or place it in " + f"ComfyUI's input directory" + ) + + +def validate_workflow_models(workflow): + """ + Pre-flight check that every model file referenced by the workflow exists. + + Walks the workflow graph and, for every known model loader node, verifies + the referenced filename against what ComfyUI actually has registered + (/object_info). This surfaces every missing model in one clear error + *before* the workflow is submitted, instead of ComfyUI's cryptic + 'Value not in list' validation crash after the cold start. + + If /object_info can't be fetched, validation is skipped (fail open): a + transient error must not block a valid workflow, and queue_workflow's + 400 handling remains as the backstop. + + Args: + workflow (dict): The workflow graph ({node_id: {class_type, inputs}}). + + Returns: + str: A user-facing error message listing every problem found, or None + if the workflow looks valid (or validation was skipped). + """ + if not isinstance(workflow, dict): + return None + + object_info = _fetch_object_info() + if object_info is None: + print( + "worker-comfyui - Skipping workflow model pre-flight check " + "(could not fetch /object_info)" + ) + return None + + problems = [] + for node_id, node in workflow.items(): + if not isinstance(node, dict): + continue + class_type = node.get("class_type") + inputs = node.get("inputs") + if not isinstance(inputs, dict): + continue + + if class_type in MODEL_LOADER_NODES: + model_type, fields = MODEL_LOADER_NODES[class_type] + for field in fields: + value = inputs.get(field) + # Non-string values are links to other nodes' outputs + # ([node_id, slot]) — nothing to check. + if not isinstance(value, str): + continue + problem = _check_model_reference( + object_info, node_id, class_type, field, value, model_type + ) + if problem: + problems.append(problem) + elif class_type == "LoadImage": + value = inputs.get("image") + if isinstance(value, str): + problem = _check_image_reference(object_info, node_id, value) + if problem: + problems.append(problem) + + if not problems: + return None + + message = ( + "Workflow validation failed — the following referenced files are not " + "available on this worker:\n" + ) + message += "\n".join(f"• {problem}" for problem in problems) + message += ( + "\n\nUpload the missing file(s) to the matching directory on your " + "network volume, or update the workflow to use one of the available files." + ) + return message + def queue_workflow(workflow, client_id, comfy_org_api_key=None): """ @@ -617,6 +854,15 @@ def handler(job): "details": upload_result["details"], } + # Pre-flight: verify every model file the workflow references actually + # exists before submitting, so a missing model fails fast with a clear + # message instead of ComfyUI's raw validation crash. Runs after image + # upload so freshly uploaded LoadImage inputs are already visible. + preflight_error = validate_workflow_models(workflow) + if preflight_error: + print(f"worker-comfyui - Workflow model pre-flight failed:\n{preflight_error}") + return {"error": preflight_error} + ws = None client_id = str(uuid.uuid4()) prompt_id = None diff --git a/tests/test_handler.py b/tests/test_handler.py index 33efcaa92..fcd0916ca 100644 --- a/tests/test_handler.py +++ b/tests/test_handler.py @@ -1,24 +1,80 @@ import unittest -from unittest.mock import patch, MagicMock, mock_open, Mock +from unittest.mock import patch, MagicMock, Mock import sys import os import json import base64 -# Make sure that "src" is known and can be used to import handler.py -sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "src"))) -from src import handler +# handler.py lives at the repository root; it imports network_volume as a +# sibling module (both are ADDed to / in the Docker image), which lives in +# src/ in the repository — so both directories must be importable. +_REPO_ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "..")) +sys.path.insert(0, os.path.join(_REPO_ROOT, "src")) +sys.path.insert(0, _REPO_ROOT) +import handler + + +def _make_object_info(): + """A minimal /object_info payload covering the loaders under test.""" + return { + "CheckpointLoaderSimple": { + "input": { + "required": { + "ckpt_name": [ + ["sd_xl_base_1.0.safetensors", "subdir/anime_v2.safetensors"], + {}, + ] + } + } + }, + "LoraLoader": { + "input": { + "required": { + "lora_name": [["detail_tweaker.safetensors"], {}], + "model": ["MODEL"], + "clip": ["CLIP"], + } + } + }, + "VAELoader": { + "input": {"required": {"vae_name": [["sdxl_vae.safetensors"], {}]}} + }, + "DualCLIPLoader": { + "input": { + "required": { + "clip_name1": [["clip_l.safetensors"], {}], + "clip_name2": [["t5xxl_fp16.safetensors"], {}], + } + } + }, + "UNETLoader": { + "input": {"required": {"unet_name": [["flux1-dev.safetensors"], {}]}} + }, + "UpscaleModelLoader": { + "input": {"required": {"model_name": [["4x_ultrasharp.pth"], {}]}} + }, + "LoadImage": { + "input": {"required": {"image": [["example.png"], {}]}} + }, + } + -# Local folder for test resources -RUNPOD_WORKER_COMFY_TEST_RESOURCES_IMAGES = "./test_resources/images" +def _mock_object_info_response(object_info): + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = object_info + return mock_response -class TestRunpodWorkerComfy(unittest.TestCase): +class TestValidateInput(unittest.TestCase): def test_valid_input_with_workflow_only(self): input_data = {"workflow": {"key": "value"}} validated_data, error = handler.validate_input(input_data) self.assertIsNone(error) - self.assertEqual(validated_data, {"workflow": {"key": "value"}, "images": None}) + self.assertEqual( + validated_data, + {"workflow": {"key": "value"}, "images": None, "comfy_org_api_key": None}, + ) def test_valid_input_with_workflow_and_images(self): input_data = { @@ -27,7 +83,8 @@ def test_valid_input_with_workflow_and_images(self): } validated_data, error = handler.validate_input(input_data) self.assertIsNone(error) - self.assertEqual(validated_data, input_data) + self.assertEqual(validated_data["workflow"], input_data["workflow"]) + self.assertEqual(validated_data["images"], input_data["images"]) def test_input_missing_workflow(self): input_data = {"images": [{"name": "image1.png", "image": "base64string"}]} @@ -56,7 +113,7 @@ def test_valid_json_string_input(self): input_data = '{"workflow": {"key": "value"}}' validated_data, error = handler.validate_input(input_data) self.assertIsNone(error) - self.assertEqual(validated_data, {"workflow": {"key": "value"}, "images": None}) + self.assertEqual(validated_data["workflow"], {"key": "value"}) def test_empty_input(self): input_data = None @@ -64,6 +121,8 @@ def test_empty_input(self): self.assertIsNotNone(error) self.assertEqual(error, "Please provide input") + +class TestServerAndQueue(unittest.TestCase): @patch("handler.requests.get") def test_check_server_server_up(self, mock_requests): mock_response = MagicMock() @@ -73,166 +132,297 @@ def test_check_server_server_up(self, mock_requests): result = handler.check_server("http://127.0.0.1:8188", 1, 50) self.assertTrue(result) + @patch("handler._is_comfyui_process_alive", return_value=None) @patch("handler.requests.get") - def test_check_server_server_down(self, mock_requests): - mock_requests.get.side_effect = handler.requests.RequestException() + def test_check_server_server_down(self, mock_requests, mock_alive): + mock_requests.side_effect = handler.requests.RequestException() result = handler.check_server("http://127.0.0.1:8188", 1, 50) self.assertFalse(result) - @patch("handler.urllib.request.urlopen") - def test_queue_prompt(self, mock_urlopen): + @patch("handler.requests.post") + def test_queue_workflow(self, mock_post): mock_response = MagicMock() - mock_response.read.return_value = json.dumps({"prompt_id": "123"}).encode() - mock_urlopen.return_value = mock_response - result = handler.queue_workflow({"prompt": "test"}) - self.assertEqual(result, {"prompt_id": "123"}) - - @patch("handler.urllib.request.urlopen") - def test_get_history(self, mock_urlopen): - # Mock response data as a JSON string - mock_response_data = json.dumps({"key": "value"}).encode("utf-8") - - # Define a mock response function for `read` - def mock_read(): - return mock_response_data - - # Create a mock response object - mock_response = Mock() - mock_response.read = mock_read + mock_response.status_code = 200 + mock_response.json.return_value = {"prompt_id": "123"} + mock_post.return_value = mock_response - # Mock the __enter__ and __exit__ methods to support the context manager - mock_response.__enter__ = lambda s: s - mock_response.__exit__ = Mock() + result = handler.queue_workflow({"1": {"class_type": "X"}}, "client-1") + self.assertEqual(result, {"prompt_id": "123"}) - # Set the return value of the urlopen mock - mock_urlopen.return_value = mock_response + @patch("handler.requests.get") + def test_get_history(self, mock_get): + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = {"key": "value"} + mock_get.return_value = mock_response - # Call the function under test result = handler.get_history("123") - - # Assertions self.assertEqual(result, {"key": "value"}) - mock_urlopen.assert_called_with("http://127.0.0.1:8188/history/123") - - @patch("builtins.open", new_callable=mock_open, read_data=b"test") - def test_base64_encode(self, mock_file): - test_data = base64.b64encode(b"test").decode("utf-8") + mock_get.assert_called_with("http://127.0.0.1:8188/history/123", timeout=30) - result = handler.base64_encode("dummy_path") - self.assertEqual(result, test_data) +class TestUploadImages(unittest.TestCase): + @patch("handler.requests.post") + def test_upload_images_successful(self, mock_post): + mock_response = MagicMock() + mock_response.status_code = 200 + mock_post.return_value = mock_response - @patch("handler.os.path.exists") - @patch("handler.rp_upload.upload_image") - @patch.dict( - os.environ, {"COMFY_OUTPUT_PATH": RUNPOD_WORKER_COMFY_TEST_RESOURCES_IMAGES} - ) - def test_bucket_endpoint_not_configured(self, mock_upload_image, mock_exists): - mock_exists.return_value = True - mock_upload_image.return_value = "simulated_uploaded/image.png" + test_image_data = base64.b64encode(b"Test Image Data").decode("utf-8") + images = [{"name": "test_image.png", "image": test_image_data}] - outputs = { - "node_id": {"images": [{"filename": "ComfyUI_00001_.png", "subfolder": ""}]} - } - job_id = "123" + responses = handler.upload_images(images) + self.assertEqual(responses["status"], "success") - result = handler.process_output_images(outputs, job_id) + @patch("handler.requests.post") + def test_upload_images_failed(self, mock_post): + mock_response = MagicMock() + mock_response.status_code = 400 + mock_response.raise_for_status.side_effect = handler.requests.RequestException( + "400 Client Error" + ) + mock_post.return_value = mock_response - self.assertEqual(result["status"], "success") + test_image_data = base64.b64encode(b"Test Image Data").decode("utf-8") + images = [{"name": "test_image.png", "image": test_image_data}] - @patch("handler.os.path.exists") - @patch("handler.rp_upload.upload_image") - @patch.dict( - os.environ, - { - "COMFY_OUTPUT_PATH": RUNPOD_WORKER_COMFY_TEST_RESOURCES_IMAGES, - "BUCKET_ENDPOINT_URL": "http://example.com", - }, - ) - def test_bucket_endpoint_configured(self, mock_upload_image, mock_exists): - # Mock the os.path.exists to return True, simulating that the image exists - mock_exists.return_value = True - - # Mock the rp_upload.upload_image to return a simulated URL - mock_upload_image.return_value = "http://example.com/uploaded/image.png" - - # Define the outputs and job_id for the test - outputs = { - "node_id": { - "images": [{"filename": "ComfyUI_00001_.png", "subfolder": "test"}] - } - } - job_id = "123" + responses = handler.upload_images(images) + self.assertEqual(responses["status"], "error") - # Call the function under test - result = handler.process_output_images(outputs, job_id) - # Assertions - self.assertEqual(result["status"], "success") - self.assertEqual(result["message"], "http://example.com/uploaded/image.png") - mock_upload_image.assert_called_once_with( - job_id, "./test_resources/images/test/ComfyUI_00001_.png" +class TestGetAvailableModels(unittest.TestCase): + @patch("handler.requests.get") + def test_returns_all_model_types(self, mock_get): + mock_get.return_value = _mock_object_info_response(_make_object_info()) + available = handler.get_available_models() + self.assertEqual( + available["checkpoints"], + ["sd_xl_base_1.0.safetensors", "subdir/anime_v2.safetensors"], ) + self.assertEqual(available["loras"], ["detail_tweaker.safetensors"]) + self.assertEqual(available["vae"], ["sdxl_vae.safetensors"]) + self.assertEqual( + available["text_encoders"], + ["clip_l.safetensors", "t5xxl_fp16.safetensors"], + ) + self.assertEqual(available["diffusion_models"], ["flux1-dev.safetensors"]) + self.assertEqual(available["upscale_models"], ["4x_ultrasharp.pth"]) - @patch("handler.os.path.exists") - @patch("handler.rp_upload.upload_image") - @patch.dict( - os.environ, - { - "COMFY_OUTPUT_PATH": RUNPOD_WORKER_COMFY_TEST_RESOURCES_IMAGES, - "BUCKET_ENDPOINT_URL": "http://example.com", - "BUCKET_ACCESS_KEY_ID": "", - "BUCKET_SECRET_ACCESS_KEY": "", - }, - ) - def test_bucket_image_upload_fails_env_vars_wrong_or_missing( - self, mock_upload_image, mock_exists - ): - # Simulate the file existing in the output path - mock_exists.return_value = True - - # When AWS credentials are wrong or missing, upload_image should return 'simulated_uploaded/...' - mock_upload_image.return_value = "simulated_uploaded/image.png" - - outputs = { - "node_id": {"images": [{"filename": "ComfyUI_00001_.png", "subfolder": ""}]} + @patch("handler.requests.get") + def test_returns_empty_dict_when_unreachable(self, mock_get): + mock_get.side_effect = handler.requests.RequestException("boom") + self.assertEqual(handler.get_available_models(), {}) + + +class TestValidateWorkflowModels(unittest.TestCase): + def _validate(self, workflow, object_info=None): + with patch("handler.requests.get") as mock_get: + mock_get.return_value = _mock_object_info_response( + object_info if object_info is not None else _make_object_info() + ) + return handler.validate_workflow_models(workflow) + + def test_all_references_present(self): + workflow = { + "1": { + "class_type": "CheckpointLoaderSimple", + "inputs": {"ckpt_name": "sd_xl_base_1.0.safetensors"}, + }, + "2": { + "class_type": "LoraLoader", + "inputs": {"lora_name": "detail_tweaker.safetensors", "model": ["1", 0]}, + }, + "3": {"class_type": "KSampler", "inputs": {"seed": 42}}, } - job_id = "123" + self.assertIsNone(self._validate(workflow)) - result = handler.process_output_images(outputs, job_id) + def test_subfolder_relative_name_present(self): + workflow = { + "1": { + "class_type": "CheckpointLoaderSimple", + "inputs": {"ckpt_name": "subdir/anime_v2.safetensors"}, + } + } + self.assertIsNone(self._validate(workflow)) - # Check if the image was saved to the 'simulated_uploaded' directory - self.assertIn("simulated_uploaded", result["message"]) - self.assertEqual(result["status"], "success") + def test_missing_checkpoint(self): + workflow = { + "4": { + "class_type": "CheckpointLoaderSimple", + "inputs": {"ckpt_name": "does_not_exist.safetensors"}, + } + } + error = self._validate(workflow) + self.assertIsNotNone(error) + self.assertIn("does_not_exist.safetensors", error) + self.assertIn("checkpoints", error) + self.assertIn("/runpod-volume/models/checkpoints/", error) + self.assertIn("sd_xl_base_1.0.safetensors", error) # lists available + + def test_multiple_missing_models_reported_together(self): + workflow = { + "1": { + "class_type": "LoraLoader", + "inputs": {"lora_name": "missing_lora.safetensors"}, + }, + "2": { + "class_type": "VAELoader", + "inputs": {"vae_name": "missing_vae.safetensors"}, + }, + "3": { + "class_type": "DualCLIPLoader", + "inputs": { + "clip_name1": "missing_clip.safetensors", + "clip_name2": "t5xxl_fp16.safetensors", + }, + }, + } + error = self._validate(workflow) + self.assertIsNotNone(error) + self.assertIn("missing_lora.safetensors", error) + self.assertIn("/runpod-volume/models/loras/", error) + self.assertIn("missing_vae.safetensors", error) + self.assertIn("/runpod-volume/models/vae/", error) + self.assertIn("missing_clip.safetensors", error) + self.assertIn("/runpod-volume/models/clip/", error) + # the valid second clip must not be reported + self.assertNotIn("'t5xxl_fp16.safetensors' not found", error) + + def test_list_placeholder_gets_specific_message(self): + workflow = { + "1": { + "class_type": "CheckpointLoaderSimple", + "inputs": {"ckpt_name": "__list__"}, + } + } + error = self._validate(workflow) + self.assertIsNotNone(error) + self.assertIn("placeholder", error) + self.assertIn("__list__", error) + self.assertIn("UI default", error) + + def test_case_mismatch_gets_hint(self): + workflow = { + "1": { + "class_type": "CheckpointLoaderSimple", + "inputs": {"ckpt_name": "SD_XL_Base_1.0.safetensors"}, + } + } + error = self._validate(workflow) + self.assertIsNotNone(error) + self.assertIn("case-sensitive", error) + self.assertIn("sd_xl_base_1.0.safetensors", error) + + def test_linked_inputs_are_skipped(self): + workflow = { + "1": { + "class_type": "LoraLoader", + "inputs": {"lora_name": ["7", 0], "model": ["1", 0]}, + } + } + self.assertIsNone(self._validate(workflow)) + + def test_unregistered_loader_is_skipped(self): + # UnetLoaderGGUF is a known loader type but not present in this + # ComfyUI build's /object_info — leave it to ComfyUI's validation. + workflow = { + "1": { + "class_type": "UnetLoaderGGUF", + "inputs": {"unet_name": "whatever.gguf"}, + } + } + self.assertIsNone(self._validate(workflow)) - @patch("handler.requests.post") - def test_upload_images_successful(self, mock_post): - mock_response = unittest.mock.Mock() - mock_response.status_code = 200 - mock_response.text = "Successfully uploaded" - mock_post.return_value = mock_response + def test_missing_input_image(self): + workflow = { + "1": { + "class_type": "LoadImage", + "inputs": {"image": "/runpod-volume/not_uploaded.png"}, + } + } + error = self._validate(workflow) + self.assertIsNotNone(error) + self.assertIn("/runpod-volume/not_uploaded.png", error) + self.assertIn("images", error) + + def test_annotated_image_reference_is_skipped(self): + workflow = { + "1": { + "class_type": "LoadImage", + "inputs": {"image": "result.png [output]"}, + } + } + self.assertIsNone(self._validate(workflow)) - test_image_data = base64.b64encode(b"Test Image Data").decode("utf-8") + def test_fails_open_when_object_info_unreachable(self): + workflow = { + "1": { + "class_type": "CheckpointLoaderSimple", + "inputs": {"ckpt_name": "does_not_exist.safetensors"}, + } + } + with patch("handler.requests.get") as mock_get: + mock_get.side_effect = handler.requests.RequestException("network blip") + self.assertIsNone(handler.validate_workflow_models(workflow)) - images = [{"name": "test_image.png", "image": test_image_data}] - responses = handler.upload_images(images) +class TestHandlerPreflightOrdering(unittest.TestCase): + """The pre-flight must run before queue_workflow and not block valid jobs.""" - self.assertEqual(len(responses), 3) - self.assertEqual(responses["status"], "success") + @patch("handler.queue_workflow") + @patch("handler.check_server", return_value=True) + @patch("handler.requests.get") + def test_preflight_failure_short_circuits_queue( + self, mock_get, mock_check_server, mock_queue + ): + mock_get.return_value = _mock_object_info_response(_make_object_info()) + job = { + "id": "job-1", + "input": { + "workflow": { + "1": { + "class_type": "CheckpointLoaderSimple", + "inputs": {"ckpt_name": "missing.safetensors"}, + } + } + }, + } + result = handler.handler(job) + self.assertIn("error", result) + self.assertIn("missing.safetensors", result["error"]) + mock_queue.assert_not_called() + + @patch("handler.get_history") + @patch("handler.queue_workflow") + @patch("handler.websocket.WebSocket") + @patch("handler.check_server", return_value=True) + @patch("handler.requests.get") + def test_valid_workflow_reaches_queue_unchanged( + self, mock_get, mock_check_server, mock_ws_cls, mock_queue, mock_history + ): + mock_get.return_value = _mock_object_info_response(_make_object_info()) + mock_queue.return_value = {"prompt_id": "abc"} + mock_history.return_value = {"abc": {"outputs": {"9": {"images": []}}}} - @patch("handler.requests.post") - def test_upload_images_failed(self, mock_post): - mock_response = unittest.mock.Mock() - mock_response.status_code = 400 - mock_response.text = "Error uploading" - mock_post.return_value = mock_response + mock_ws = MagicMock() + mock_ws.recv.return_value = json.dumps( + {"type": "executing", "data": {"node": None, "prompt_id": "abc"}} + ) + mock_ws_cls.return_value = mock_ws - test_image_data = base64.b64encode(b"Test Image Data").decode("utf-8") + workflow = { + "1": { + "class_type": "CheckpointLoaderSimple", + "inputs": {"ckpt_name": "sd_xl_base_1.0.safetensors"}, + } + } + job = {"id": "job-2", "input": {"workflow": workflow}} + result = handler.handler(job) - images = [{"name": "test_image.png", "image": test_image_data}] + self.assertNotIn("error", result) + mock_queue.assert_called_once() + self.assertEqual(mock_queue.call_args[0][0], workflow) - responses = handler.upload_images(images) - self.assertEqual(len(responses), 3) - self.assertEqual(responses["status"], "error") +if __name__ == "__main__": + unittest.main()