diff --git a/addons/official/storycore_asset_creator/COMFYUI_CLIENT.md b/addons/official/storycore_asset_creator/COMFYUI_CLIENT.md new file mode 100644 index 00000000..805868cd --- /dev/null +++ b/addons/official/storycore_asset_creator/COMFYUI_CLIENT.md @@ -0,0 +1,107 @@ +# Suivi des tâches ComfyUI + +Le client Python existant reste le point d'accès HTTP de l'Asset Creator. +Cette correction n'ajoute aucun service externe, modèle ou abonnement. + +## Deux défauts reproduits + +Sur le commit `5b6c83bd26f60f7eb5e4da53f958f2d3b06edb4a`, avec des réponses +simulées et les sockets interdites : + +- Une entrée d'historique contenant `status_str="error"`, `completed=false` + et des sorties partielles était renvoyée comme un résultat réussi. +- L'appel HTTP à `/history/{prompt_id}` ne recevait aucun `timeout`. + La lecture du client montre la même omission pour les autres appels, + sauf `/system_stats`. + +## Comportement du client + +| Situation | Résultat | +|---|---| +| `status.status_str == "success"` et `status.completed is True` | Retourne le dictionnaire `outputs`, même vide. | +| `status.status_str == "error"` | Lève `RuntimeError`, même avec des fichiers intermédiaires. | +| Ancien champ `error` au premier niveau | Lève `RuntimeError`. | +| Historique absent, statut inconnu ou incomplet | Continue le suivi jusqu'à son expiration. | +| Budget de suivi écoulé | `ComfyUIWaitTimeout`, avec `reason="wait_deadline"`. | +| `requests.exceptions.Timeout` pendant un GET de suivi | `ComfyUIWaitTimeout`, avec `reason="http_timeout"` et cause conservée. | +| Autre erreur HTTP ou réseau | Exception Requests transmise à l'appelant. | + +`ComfyUIWaitTimeout` hérite de `TimeoutError` et expose `prompt_id`. +Le client ne supprime pas la tâche, ne l'annule pas et ne la soumet pas à +nouveau quand le suivi expire. Une expiration ne prouve pas que la tâche a +échoué ou qu'elle tourne encore ; son état doit être relu sur le serveur. + +## Deux délais distincts + +- `request_timeout=30.0` : délai Requests appliqué à tous les appels HTTP. + Le contrôle de disponibilité garde un plafond de 5 secondes. +- `wait_for_result(..., timeout=300.0, poll_interval=2.0)` : budget du suivi + et intervalle entre lectures. Le budget utilise une horloge monotone. + Avant chaque GET, le délai HTTP est limité au budget restant ; les pauses + sont également limitées. Le budget est revérifié au retour des GET et + après le callback avant de dormir. + +Ces trois paramètres doivent être finis et strictement positifs. +`get_history` et `get_queue_status` acceptent aussi un `timeout` nommé pour +réduire leur délai HTTP, sans dépasser celui du client. + +Le délai Requests concerne la connexion et l'inactivité de lecture. Il +ne garantit **pas** une durée totale stricte : plusieurs adresses réseau, +une réponse qui arrive lentement ou un callback bloquant peuvent dépasser +le budget. Aucun thread n'est abandonné pour simuler une telle garantie. +Voir la [documentation officielle Requests sur les délais](https://requests.readthedocs.io/en/latest/user/advanced/#timeouts). + +## Reprendre le suivi + +```python +from addons.official.storycore_asset_creator.src.comfyui_client import ( + ComfyUIClient, + ComfyUIWaitTimeout, +) + +client = ComfyUIClient.from_project_config(request_timeout=30.0) +# prompt_id est l'identifiant déjà renvoyé par queue_workflow et sauvegardé. +try: + outputs = client.wait_for_result(prompt_id, timeout=300.0) +except ComfyUIWaitTimeout as error: + prompt_id_a_reprendre = error.prompt_id + # Plus tard : client.wait_for_result(prompt_id_a_reprendre, timeout=300.0) +``` + +L'appelant doit sauvegarder l'identifiant dès l'acceptation. Cette correction +ne crée pas de stockage persistant de tâches et ne raccorde pas encore la +reprise à une interface MCP, CLI ou Blender. + +Si le POST `/prompt` expire avant de renvoyer un identifiant, l'acceptation +reste incertaine. Le client laisse remonter l'exception Requests et n'ajoute +aucune relance automatique. Il faut vérifier la file du serveur avant de +décider d'une nouvelle soumission. + +Les réponses des téléchargements sont fermées même en cas d'erreur. +Un téléchargement interrompu peut toujours laisser un fichier partiel : +la correction ne rend pas l'écriture atomique. + +## Contrat et vérification + +Contrat lu dans ComfyUI au commit +[`387f98aa2822f684b8597959a52a467d88cc4806`](https://github.com/Comfy-Org/ComfyUI/commit/387f98aa2822f684b8597959a52a467d88cc4806) : +[`execution.py`](https://github.com/Comfy-Org/ComfyUI/blob/387f98aa2822f684b8597959a52a467d88cc4806/execution.py#L1316-L1340) +décrit le statut enregistré ; +[`main.py`](https://github.com/Comfy-Org/ComfyUI/blob/387f98aa2822f684b8597959a52a467d88cc4806/main.py#L367-L372) +le renseigne à la fin du traitement. Le client ne déduit pas la réussite +de la seule présence d'un fichier. + +Depuis la racine du dépôt, avec `requests` installé : + +```bash +python -m unittest discover -s tests/asset_creator -p test_comfyui_client.py -v +``` + +Les 18 tests utilisent des réponses HTTP simulées et une horloge contrôlée. +La création de sockets y est interdite. Ils couvrent les statuts terminaux, +les sorties partielles, les délais, la reprise et les ressources HTTP. +Vérification locale effectuée avec Python 3.12.14 et Requests 2.34.2. + +Il ne s'agit pas d'un test sur un serveur ComfyUI, Blender ou un GPU. +Les recettes Trellis2 locales, leurs extensions et le format API réellement +soumis restent à vérifier dans l'environnement d'exécution. diff --git a/addons/official/storycore_asset_creator/README.md b/addons/official/storycore_asset_creator/README.md index 2cd61736..33744eb0 100644 --- a/addons/official/storycore_asset_creator/README.md +++ b/addons/official/storycore_asset_creator/README.md @@ -51,6 +51,12 @@ Le port ComfyUI **n'est jamais supposé** — il varie selon l'édition install 3. Variables d'environnement ← STORYCORE_COMFYUI_HOST / STORYCORE_COMFYUI_PORT ``` +### Délais et reprise du suivi + +Le client distingue le délai HTTP du budget d'attente de la génération. +Une expiration du suivi conserve l'identifiant du prompt ; elle ne l'annule +pas. Voir [le contrat du client et les tests](COMFYUI_CLIENT.md). + ### Override via variables d'environnement ```bash diff --git a/addons/official/storycore_asset_creator/src/comfyui_client.py b/addons/official/storycore_asset_creator/src/comfyui_client.py index 6e5cf545..4a6fafe6 100644 --- a/addons/official/storycore_asset_creator/src/comfyui_client.py +++ b/addons/official/storycore_asset_creator/src/comfyui_client.py @@ -10,10 +10,11 @@ from __future__ import annotations +import math import time import uuid from pathlib import Path -from typing import Any, Dict, Optional +from typing import Any try: import requests @@ -21,6 +22,29 @@ requests = None # Blender embeds its own Python; requests peut manquer +class ComfyUIWaitTimeout(TimeoutError): + """Suivi interrompu; prompt_id permet de reprendre sans soumettre de nouveau.""" + + def __init__(self, prompt_id: str, reason: str = "wait_deadline"): + self.prompt_id = prompt_id + self.reason = reason + detail = ( + "delai HTTP depasse" + if reason == "http_timeout" + else "delai d'attente depasse" + ) + super().__init__( + f"ComfyUI {prompt_id}: {detail}. " + "Le suivi est interrompu; le prompt n'a pas ete annule." + ) + + +def _positive_seconds(value: float, name: str) -> float: + if not math.isfinite(value) or value <= 0: + raise ValueError(f"{name} doit etre un nombre fini strictement positif") + return value + + class ComfyUIClient: """ Client HTTP pour ComfyUI (localhost ou remote). @@ -35,7 +59,14 @@ class ComfyUIClient: NE PAS hardcoder le port 8188 — lire depuis config/comfyui_config.json. """ - def __init__(self, host: str = "127.0.0.1", port: Optional[int] = None): + def __init__( + self, + host: str = "127.0.0.1", + port: int | None = None, + *, + request_timeout: float = 30.0, + ): + self.request_timeout = _positive_seconds(request_timeout, "request_timeout") if port is None: # Tenter de charger depuis la config projet try: @@ -51,7 +82,9 @@ def __init__(self, host: str = "127.0.0.1", port: Optional[int] = None): self.client_id = str(uuid.uuid4()) @classmethod - def from_project_config(cls, blender_prefs=None) -> "ComfyUIClient": + def from_project_config( + cls, blender_prefs=None, *, request_timeout: float = 30.0 + ) -> ComfyUIClient: """ Cree un client en lisant la config depuis config/comfyui_config.json (avec surcharge optionnelle depuis les preferences Blender). @@ -66,19 +99,28 @@ def from_project_config(cls, blender_prefs=None) -> "ComfyUIClient": from .config_loader import get_comfyui_connection host, port = get_comfyui_connection(blender_prefs=blender_prefs) - return cls(host=host, port=port) + return cls(host=host, port=port, request_timeout=request_timeout) + + def _http_timeout(self, timeout: float | None) -> float: + if timeout is None: + return self.request_timeout + return min(self.request_timeout, _positive_seconds(timeout, "timeout")) # ── API ────────────────────────────────────────────────────────────────── def is_alive(self) -> bool: """Verifie que ComfyUI repond.""" + if requests is None: + return False try: - r = requests.get(f"{self.base_url}/system_stats", timeout=5) + r = requests.get( + f"{self.base_url}/system_stats", timeout=min(5.0, self.request_timeout) + ) return r.status_code == 200 - except Exception: + except requests.exceptions.RequestException: return False - def upload_image(self, image_path: str, subfolder: str = "") -> Dict[str, Any]: + def upload_image(self, image_path: str, subfolder: str = "") -> dict[str, Any]: """ Upload une image dans ComfyUI input/. @@ -90,12 +132,17 @@ def upload_image(self, image_path: str, subfolder: str = "") -> Dict[str, Any]: data = {"type": "input", "overwrite": "true"} if subfolder: data["subfolder"] = subfolder - r = requests.post(f"{self.base_url}/upload/image", files=files, data=data) + r = requests.post( + f"{self.base_url}/upload/image", + files=files, + data=data, + timeout=self.request_timeout, + ) r.raise_for_status() return r.json() def queue_workflow( - self, workflow: Dict[str, Any], client_id: Optional[str] = None + self, workflow: dict[str, Any], client_id: str | None = None ) -> str: """ Envoie le workflow dans la queue ComfyUI. @@ -106,19 +153,25 @@ def queue_workflow( "prompt": workflow, "client_id": client_id or self.client_id, } - r = requests.post(f"{self.base_url}/prompt", json=payload) + r = requests.post( + f"{self.base_url}/prompt", json=payload, timeout=self.request_timeout + ) r.raise_for_status() return r.json()["prompt_id"] - def get_queue_status(self) -> Dict[str, Any]: + def get_queue_status(self, *, timeout: float | None = None) -> dict[str, Any]: """Retourne le statut de la queue.""" - r = requests.get(f"{self.base_url}/queue") + r = requests.get(f"{self.base_url}/queue", timeout=self._http_timeout(timeout)) r.raise_for_status() return r.json() - def get_history(self, prompt_id: str) -> Optional[Dict[str, Any]]: + def get_history( + self, prompt_id: str, *, timeout: float | None = None + ) -> dict[str, Any] | None: """Retourne l'historique d'un prompt execute.""" - r = requests.get(f"{self.base_url}/history/{prompt_id}") + r = requests.get( + f"{self.base_url}/history/{prompt_id}", timeout=self._http_timeout(timeout) + ) r.raise_for_status() data = r.json() return data.get(prompt_id) @@ -129,44 +182,79 @@ def wait_for_result( timeout: float = 300.0, poll_interval: float = 2.0, progress_callback=None, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Attend la fin d'un prompt en polling. Args: prompt_id : ID retourne par queue_workflow - timeout : secondes max avant abandon + timeout : budget de suivi (hors garanties de temps reel) poll_interval : intervalle de polling en secondes progress_callback: callable(status_str) optionnel Returns: outputs dict du prompt - Raises: TimeoutError si depasse le timeout + Raises: ComfyUIWaitTimeout si budget ecoule ou delai HTTP depasse RuntimeError si erreur dans le workflow + + Un timeout arrete le suivi, pas le prompt. Les delais Requests bornent + connexion/inactivite de lecture, pas la duree totale d'une requete. """ - start = time.time() - while True: - elapsed = time.time() - start - if elapsed > timeout: - raise TimeoutError(f"Trellis2: timeout apres {timeout}s") + _positive_seconds(timeout, "timeout") + _positive_seconds(poll_interval, "poll_interval") + if requests is None: + raise ImportError("Le client ComfyUI requiert le module requests") + start = time.monotonic() + deadline = start + timeout + + def remaining() -> float: + budget = deadline - time.monotonic() + if budget <= 0: + raise ComfyUIWaitTimeout(prompt_id) + return budget - history = self.get_history(prompt_id) + while True: + try: + history = self.get_history(prompt_id, timeout=remaining()) + except requests.exceptions.Timeout as error: + raise ComfyUIWaitTimeout(prompt_id, "http_timeout") from error + remaining() if history: if "error" in history: - raise RuntimeError(f"ComfyUI erreur: {history['error']}") - outputs = history.get("outputs", {}) - if outputs: - return outputs + raise RuntimeError( + f"ComfyUI {prompt_id} erreur: {history['error']}" + ) + status = history.get("status") + if isinstance(status, dict): + if status.get("status_str") == "error": + raise RuntimeError( + f"ComfyUI {prompt_id} erreur: {status.get('messages', [])}" + ) + if ( + status.get("status_str") == "success" + and status.get("completed") is True + ): + outputs = history.get("outputs", {}) + if not isinstance(outputs, dict): + raise RuntimeError( + f"ComfyUI {prompt_id}: outputs invalides" + ) + return outputs if progress_callback: - queue = self.get_queue_status() + try: + queue = self.get_queue_status(timeout=remaining()) + except requests.exceptions.Timeout as error: + raise ComfyUIWaitTimeout(prompt_id, "http_timeout") from error + remaining() running = len(queue.get("queue_running", [])) pending = len(queue.get("queue_pending", [])) + elapsed = time.monotonic() - start progress_callback( f"Running: {running} | Pending: {pending} | {elapsed:.0f}s" ) - time.sleep(poll_interval) + time.sleep(min(poll_interval, remaining())) def download_output(self, filename: str, dest_dir: str, subfolder: str = "") -> str: """ @@ -177,36 +265,37 @@ def download_output(self, filename: str, dest_dir: str, subfolder: str = "") -> params = {"filename": filename, "type": "output"} if subfolder: params["subfolder"] = subfolder - r = requests.get(f"{self.base_url}/view", params=params, stream=True) - r.raise_for_status() - - dest = Path(dest_dir) - dest.mkdir(parents=True, exist_ok=True) - out_path = dest / filename + with requests.get( + f"{self.base_url}/view", + params=params, + stream=True, + timeout=self.request_timeout, + ) as r: + r.raise_for_status() + dest = Path(dest_dir) + dest.mkdir(parents=True, exist_ok=True) + out_path = dest / filename - with open(out_path, "wb") as f: - for chunk in r.iter_content(chunk_size=8192): - f.write(chunk) + with open(out_path, "wb") as f: + f.writelines(r.iter_content(chunk_size=8192)) return str(out_path) - def get_output_files(self, outputs: Dict[str, Any]) -> list[str]: + def get_output_files(self, outputs: dict[str, Any]) -> list[str]: """ Extrait la liste des noms de fichiers depuis les outputs d'un prompt. Cherche les nodes de type 'images', 'gltf', 'glb_path' etc. """ files = [] - for node_id, node_outputs in outputs.items(): - for key, values in node_outputs.items(): + for node_outputs in outputs.values(): + for values in node_outputs.values(): if isinstance(values, list): for v in values: if isinstance(v, dict) and "filename" in v: files.append(v["filename"]) - elif isinstance(values, str) and ( - values.endswith(".glb") - or values.endswith(".gltf") - or values.endswith(".png") + elif isinstance(values, str) and values.endswith( + (".glb", ".gltf", ".png") ): files.append(Path(values).name) return files diff --git a/tests/asset_creator/test_comfyui_client.py b/tests/asset_creator/test_comfyui_client.py new file mode 100644 index 00000000..56954fe6 --- /dev/null +++ b/tests/asset_creator/test_comfyui_client.py @@ -0,0 +1,342 @@ +"""HTTP/status regressions without a ComfyUI server, Blender or GPU.""" + +import importlib.util +import tempfile +import unittest +from pathlib import Path +from unittest.mock import MagicMock, Mock, patch + +import requests + +SOURCE = ( + Path(__file__).resolve().parents[2] + / "addons/official/storycore_asset_creator/src/comfyui_client.py" +) +SPEC = importlib.util.spec_from_file_location("asset_creator_client", SOURCE) +client_module = importlib.util.module_from_spec(SPEC) +SPEC.loader.exec_module(client_module) +ComfyUIClient = client_module.ComfyUIClient +PROMPT_ID = "accepted-prompt-id" +OUTPUTS = {"7": {"images": [{"filename": "partial.png"}]}} + + +def history(outputs=None, *, status="success", completed=True): + return { + "outputs": OUTPUTS if outputs is None else outputs, + "status": { + "status_str": status, + "completed": completed, + "messages": [["execution_error", {"exception_message": "out of memory"}]] + if status == "error" + else [], + }, + } + + +def response(payload=None): + result = MagicMock() + result.json.return_value = payload + result.status_code = 200 + result.__enter__.return_value = result + return result + + +class Clock: + def __init__(self): + self.now = 0.0 + self.sleeps = [] + + def monotonic(self): + return self.now + + def sleep(self, seconds): + self.sleeps.append(seconds) + self.now += seconds + + +class ComfyUIClientTests(unittest.TestCase): + def setUp(self): + sockets = patch( + "socket.socket", side_effect=AssertionError("network forbidden") + ) + sockets.start() + self.addCleanup(sockets.stop) + self.clock = Clock() + clock_patch = patch.object(client_module, "time", self.clock) + clock_patch.start() + self.addCleanup(clock_patch.stop) + self.client = ComfyUIClient(port=8188, request_timeout=7.0) + + def test_error_status_wins_over_partial_outputs(self): + for completed in (False, True): + with ( + self.subTest(completed=completed), + patch.object( + self.client, + "get_history", + return_value=history(status="error", completed=completed), + ), + ): + with self.assertRaisesRegex(RuntimeError, "out of memory") as caught: + self.client.wait_for_result(PROMPT_ID) + self.assertIn(PROMPT_ID, str(caught.exception)) + self.assertEqual([], self.clock.sleeps) + + def test_legacy_top_level_error_still_fails(self): + with ( + patch.object( + self.client, + "get_history", + return_value={"error": "failed", "outputs": OUTPUTS}, + ), + self.assertRaisesRegex(RuntimeError, "failed"), + ): + self.client.wait_for_result(PROMPT_ID) + + def test_only_explicit_completed_success_returns_outputs(self): + for outputs in ({}, OUTPUTS): + with ( + self.subTest(outputs=outputs), + patch.object(self.client, "get_history", return_value=history(outputs)), + ): + self.assertEqual(outputs, self.client.wait_for_result(PROMPT_ID)) + self.assertEqual([], self.clock.sleeps) + + def test_unknown_or_incomplete_status_does_not_return_partial_outputs(self): + cases = [ + {"outputs": OUTPUTS}, + {"outputs": OUTPUTS, "status": None}, + {"outputs": OUTPUTS, "status": "success"}, + history(completed=False), + history(completed="true"), + history(status="unknown"), + ] + for item in cases: + with ( + self.subTest(item=item), + patch.object(self.client, "get_history", return_value=item), + self.assertRaises(TimeoutError), + ): + self.client.wait_for_result(PROMPT_ID, timeout=1) + + def test_success_with_malformed_outputs_fails(self): + with ( + patch.object(self.client, "get_history", return_value=history([])), + self.assertRaisesRegex(RuntimeError, "outputs invalides"), + ): + self.client.wait_for_result(PROMPT_ID) + + def test_sleep_is_clipped_and_no_request_starts_after_deadline(self): + with ( + patch.object( + client_module.requests, "get", return_value=response({}) + ) as get, + self.assertRaises(client_module.ComfyUIWaitTimeout) as caught, + ): + self.client.wait_for_result(PROMPT_ID, timeout=5, poll_interval=3) + self.assertEqual([3, 2], self.clock.sleeps) + self.assertEqual( + [5, 2], [call.kwargs["timeout"] for call in get.call_args_list] + ) + self.assertEqual(PROMPT_ID, caught.exception.prompt_id) + self.assertEqual("wait_deadline", caught.exception.reason) + + def test_late_history_response_does_not_start_queue_poll(self): + def late_history(*args, **kwargs): + self.clock.now += 6 + return response({PROMPT_ID: history()}) + + callback = Mock() + with ( + patch.object( + client_module.requests, "get", side_effect=late_history + ) as get, + self.assertRaises(client_module.ComfyUIWaitTimeout), + ): + self.client.wait_for_result( + PROMPT_ID, timeout=5, progress_callback=callback + ) + self.assertEqual(1, get.call_count) + callback.assert_not_called() + self.assertEqual([], self.clock.sleeps) + + def test_queue_poll_uses_remaining_budget(self): + budgets = [] + + def get(url, **kwargs): + budgets.append(kwargs["timeout"]) + if "/history/" in url: + self.clock.now += 3 + return response({}) + self.clock.now += 2 + return response({"queue_running": [], "queue_pending": []}) + + callback = Mock() + with ( + patch.object(client_module.requests, "get", side_effect=get), + self.assertRaises(client_module.ComfyUIWaitTimeout), + ): + self.client.wait_for_result( + PROMPT_ID, timeout=5, progress_callback=callback + ) + self.assertEqual([5, 2], budgets) + callback.assert_not_called() + + def test_callback_time_counts_towards_deadline(self): + def slow_callback(status): + self.clock.now += 5 + + with ( + patch.object( + client_module.requests, + "get", + side_effect=[ + response({}), + response({"queue_running": [], "queue_pending": []}), + ], + ) as get, + self.assertRaises(client_module.ComfyUIWaitTimeout), + ): + self.client.wait_for_result( + PROMPT_ID, timeout=5, progress_callback=slow_callback + ) + self.assertEqual(2, get.call_count) + self.assertEqual([], self.clock.sleeps) + + def test_http_timeout_preserves_prompt_and_cause_for_both_poll_endpoints(self): + for endpoint in ("history", "queue"): + error = requests.exceptions.ReadTimeout("server silent") + effects = [error] if endpoint == "history" else [response({}), error] + with ( + self.subTest(endpoint=endpoint), + patch.object(client_module.requests, "get", side_effect=effects) as get, + patch.object(client_module.requests, "post") as post, + ): + with self.assertRaises(client_module.ComfyUIWaitTimeout) as caught: + self.client.wait_for_result(PROMPT_ID, progress_callback=Mock()) + self.assertEqual(PROMPT_ID, caught.exception.prompt_id) + self.assertEqual("http_timeout", caught.exception.reason) + self.assertIs(error, caught.exception.__cause__) + self.assertEqual(len(effects), get.call_count) + post.assert_not_called() + + def test_resume_wait_after_timeout_does_not_submit_again(self): + with ( + patch.object( + client_module.requests, + "get", + side_effect=[ + requests.exceptions.ReadTimeout(), + response({PROMPT_ID: history()}), + ], + ) as get, + patch.object(client_module.requests, "post") as post, + ): + with self.assertRaises(client_module.ComfyUIWaitTimeout) as caught: + self.client.wait_for_result(PROMPT_ID) + outputs = self.client.wait_for_result(caught.exception.prompt_id) + self.assertEqual(OUTPUTS, outputs) + self.assertTrue( + all(call.args[0].endswith(PROMPT_ID) for call in get.call_args_list) + ) + post.assert_not_called() + + def test_invalid_durations_fail_before_http(self): + for value in (0, -1, float("nan"), float("inf"), -float("inf")): + with ( + self.subTest(value=value), + patch.object(client_module.requests, "get") as get, + ): + with self.assertRaises(ValueError): + ComfyUIClient(port=8188, request_timeout=value) + for argument in ("timeout", "poll_interval"): + with self.assertRaises(ValueError): + self.client.wait_for_result(PROMPT_ID, **{argument: value}) + get.assert_not_called() + + def test_all_http_operations_receive_timeout(self): + for duration in (1.5, 30.0): + with self.subTest(duration=duration), tempfile.TemporaryDirectory() as temp: + client = ComfyUIClient(port=8188, request_timeout=duration) + source = Path(temp) / "source.png" + source.write_bytes(b"synthetic image") + reply = response({"prompt_id": PROMPT_ID}) + with patch.object( + client_module.requests, "post", return_value=reply + ) as post: + client.upload_image(str(source), subfolder="assets") + self.assertEqual(PROMPT_ID, client.queue_workflow({})) + self.assertEqual( + [duration, duration], + [c.kwargs["timeout"] for c in post.call_args_list], + ) + self.assertTrue( + post.call_args_list[0].kwargs["files"]["image"][1].closed + ) + reply = response({}) + reply.iter_content.return_value = [b"mesh"] + with patch.object( + client_module.requests, "get", return_value=reply + ) as get: + self.assertTrue(client.is_alive()) + client.get_queue_status() + client.get_history(PROMPT_ID) + output = client.download_output("mesh.glb", temp) + self.assertEqual( + [min(5, duration), duration, duration, duration], + [c.kwargs["timeout"] for c in get.call_args_list], + ) + self.assertEqual(b"mesh", Path(output).read_bytes()) + reply.__exit__.assert_called_once() + + def test_get_timeouts_are_capped_by_client_configuration(self): + with patch.object( + client_module.requests, "get", return_value=response({}) + ) as get: + self.client.get_history(PROMPT_ID, timeout=99) + self.client.get_queue_status(timeout=1) + self.assertEqual([7, 1], [c.kwargs["timeout"] for c in get.call_args_list]) + + def test_submission_timeout_is_not_retried(self): + error = requests.exceptions.ReadTimeout("response lost") + with ( + patch.object(client_module.requests, "post", side_effect=error) as post, + self.assertRaises(requests.exceptions.ReadTimeout) as caught, + ): + self.client.queue_workflow({}) + self.assertIs(error, caught.exception) + post.assert_called_once() + + def test_http_error_is_preserved(self): + reply = response() + reply.raise_for_status.side_effect = requests.exceptions.HTTPError("503") + with ( + patch.object(client_module.requests, "get", return_value=reply), + self.assertRaises(requests.exceptions.HTTPError), + ): + self.client.wait_for_result(PROMPT_ID) + + def test_download_closes_response_on_stream_error(self): + reply = response() + reply.iter_content.side_effect = requests.exceptions.ConnectionError( + "lost stream" + ) + with ( + tempfile.TemporaryDirectory() as temp, + patch.object(client_module.requests, "get", return_value=reply), + self.assertRaises(requests.exceptions.ConnectionError), + ): + self.client.download_output("mesh.glb", temp) + reply.__exit__.assert_called_once() + + def test_health_check_handles_missing_dependency_and_network_errors(self): + with patch.object(client_module, "requests", None): + self.assertFalse(self.client.is_alive()) + with patch.object( + client_module.requests, "get", side_effect=requests.exceptions.ReadTimeout() + ): + self.assertFalse(self.client.is_alive()) + + +if __name__ == "__main__": + unittest.main()