feat(ui): add output-text previews and live retry status to comfydv nodes
Fixes a real gap discovered this session: ChatCompletion/FormatString's
"ui": {"text": [...]} return value was never rendered anywhere — no JS in
this repo ever implemented an onExecuted handler to read it, so the
retry-status banner from the previous commit was invisible in practice.
RandomChoice never returned a ui payload at all (OUTPUT_NODE was False).
- src/js/preview_text.js: read-only STRING preview widget for
ChatCompletion/FormatString/RandomChoice, mirroring ComfyUI core's own
PreviewAny node/JS extension (confirmed via the actual bundled frontend
source, not guessed) rather than inventing a new convention.
- RandomChoice: OUTPUT_NODE=True + a ui.text preview of its arbitrary-typed
output (str/number passthrough, else JSON, else str() — same fallback
chain as core's PreviewAny). IS_CHANGED still returns the raw picked
value, not the new ui-wrapped dict, so change-detection is unaffected.
- Live retry feedback: chat()/chat_structured() gain an on_status
callback, invoked at each retry boundary with a human-readable status
(_llm/retry.py's format_retry_status/format_recovered_status).
ChatCompletion wires this to ComfyUI's own PromptServer.send_progress_text
(the same first-party mechanism core's gaussian-splat/image nodes use)
rather than a hand-rolled custom event — the frontend already renders a
live "progressText" widget for it, so no additional JS was needed for
the live half of this feature.
Verified live against a real Ollama backend (docker-compose.dev.yml +
Playwright): forced a genuine refusal/retry via custom_phrases, confirmed
both the final preview widget and the live progress widget render
correctly on ChatCompletion, and confirmed FormatString/RandomChoice's
preview widgets render too. 410 tests passing.
Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Sonnet 5
parent
7a40d22230
commit
f16da579bc
@@ -28,6 +28,7 @@ pydantic-ai's internal one.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
from collections.abc import Callable
|
||||||
from typing import cast
|
from typing import cast
|
||||||
|
|
||||||
from pydantic import BaseModel, ValidationError
|
from pydantic import BaseModel, ValidationError
|
||||||
@@ -49,6 +50,8 @@ from .provider import Message
|
|||||||
from .retry import (
|
from .retry import (
|
||||||
RETRY_BACKOFF_SECS,
|
RETRY_BACKOFF_SECS,
|
||||||
EmbedFn,
|
EmbedFn,
|
||||||
|
format_recovered_status,
|
||||||
|
format_retry_status,
|
||||||
is_refusal,
|
is_refusal,
|
||||||
next_seed,
|
next_seed,
|
||||||
next_timeout_secs,
|
next_timeout_secs,
|
||||||
@@ -135,6 +138,7 @@ async def chat_structured(
|
|||||||
timeout_secs: float = 300.0,
|
timeout_secs: float = 300.0,
|
||||||
embed_fn: EmbedFn | None = None,
|
embed_fn: EmbedFn | None = None,
|
||||||
attempt_info: dict | None = None,
|
attempt_info: dict | None = None,
|
||||||
|
on_status: Callable[[str], None] | None = None,
|
||||||
) -> BaseModel:
|
) -> BaseModel:
|
||||||
"""Call ``model`` at ``base_url`` (an OpenAI-compatible ``/v1`` root) and
|
"""Call ``model`` at ``base_url`` (an OpenAI-compatible ``/v1`` root) and
|
||||||
return a validated instance of ``schema``.
|
return a validated instance of ``schema``.
|
||||||
@@ -208,6 +212,18 @@ async def chat_structured(
|
|||||||
refusal_count = 0
|
refusal_count = 0
|
||||||
attempt_seed = (options or {}).get("seed", 0) if isinstance(options, dict) else 0
|
attempt_seed = (options or {}).get("seed", 0) if isinstance(options, dict) else 0
|
||||||
attempt_timeout = timeout_secs
|
attempt_timeout = timeout_secs
|
||||||
|
|
||||||
|
def _emit_retry_status(reason: str, attempt: int) -> None:
|
||||||
|
if on_status is None or attempt >= total_attempts:
|
||||||
|
return
|
||||||
|
upcoming_seed = next_seed(options, attempt + 1)
|
||||||
|
upcoming_timeout = next_timeout_secs(timeout_secs, attempt + 1)
|
||||||
|
on_status(
|
||||||
|
format_retry_status(
|
||||||
|
reason, attempt, total_attempts, upcoming_seed, upcoming_timeout
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
for attempt in range(1, total_attempts + 1):
|
for attempt in range(1, total_attempts + 1):
|
||||||
attempt_timeout = next_timeout_secs(timeout_secs, attempt)
|
attempt_timeout = next_timeout_secs(timeout_secs, attempt)
|
||||||
# Rebuilt each attempt so the escalated timeout actually takes
|
# Rebuilt each attempt so the escalated timeout actually takes
|
||||||
@@ -282,6 +298,7 @@ async def chat_structured(
|
|||||||
refusal_count += 1
|
refusal_count += 1
|
||||||
last_error = RuntimeError("refusal/deflection detected")
|
last_error = RuntimeError("refusal/deflection detected")
|
||||||
last_invalid_text = content
|
last_invalid_text = content
|
||||||
|
_emit_retry_status("Refusal/deflection detected", attempt)
|
||||||
if attempt < total_attempts:
|
if attempt < total_attempts:
|
||||||
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
||||||
continue
|
continue
|
||||||
@@ -292,10 +309,15 @@ async def chat_structured(
|
|||||||
timeout_secs=attempt_timeout,
|
timeout_secs=attempt_timeout,
|
||||||
refusals=refusal_count,
|
refusals=refusal_count,
|
||||||
)
|
)
|
||||||
|
if on_status is not None and attempt > 1:
|
||||||
|
on_status(
|
||||||
|
format_recovered_status(attempt, total_attempts, attempt_seed)
|
||||||
|
)
|
||||||
return output
|
return output
|
||||||
except _STRUCTURED_OUTPUT_FAILURE_EXCEPTIONS as exc:
|
except _STRUCTURED_OUTPUT_FAILURE_EXCEPTIONS as exc:
|
||||||
last_error = exc
|
last_error = exc
|
||||||
last_invalid_text = str(exc)
|
last_invalid_text = str(exc)
|
||||||
|
_emit_retry_status("Structured output failed", attempt)
|
||||||
if attempt < total_attempts:
|
if attempt < total_attempts:
|
||||||
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
||||||
|
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ exist otherwise (spec.md FR-006).
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
|
from collections.abc import Callable
|
||||||
|
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
|
|
||||||
@@ -29,6 +30,8 @@ from .ollama_provider import (
|
|||||||
from .provider import Message, ModelInfo, ModelStatus
|
from .provider import Message, ModelInfo, ModelStatus
|
||||||
from .retry import (
|
from .retry import (
|
||||||
RETRY_BACKOFF_SECS,
|
RETRY_BACKOFF_SECS,
|
||||||
|
format_recovered_status,
|
||||||
|
format_retry_status,
|
||||||
is_refusal,
|
is_refusal,
|
||||||
next_seed,
|
next_seed,
|
||||||
next_timeout_secs,
|
next_timeout_secs,
|
||||||
@@ -197,6 +200,7 @@ class LlamaCppProvider:
|
|||||||
timeout_secs: float = 300.0,
|
timeout_secs: float = 300.0,
|
||||||
max_retries: int = 2,
|
max_retries: int = 2,
|
||||||
attempt_info: dict | None = None,
|
attempt_info: dict | None = None,
|
||||||
|
on_status: Callable[[str], None] | None = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
payload_messages = [_to_openai_message(m) for m in messages]
|
payload_messages = [_to_openai_message(m) for m in messages]
|
||||||
options, think = _pop_think(options)
|
options, think = _pop_think(options)
|
||||||
@@ -285,6 +289,7 @@ class LlamaCppProvider:
|
|||||||
if choices
|
if choices
|
||||||
else ""
|
else ""
|
||||||
)
|
)
|
||||||
|
retry_reason: str | None = None
|
||||||
if response_text.strip():
|
if response_text.strip():
|
||||||
refused = False
|
refused = False
|
||||||
if refusal_cfg and refusal_cfg.get("enabled"):
|
if refusal_cfg and refusal_cfg.get("enabled"):
|
||||||
@@ -304,10 +309,31 @@ class LlamaCppProvider:
|
|||||||
timeout_secs=attempt_timeout,
|
timeout_secs=attempt_timeout,
|
||||||
refusals=refusal_count,
|
refusals=refusal_count,
|
||||||
)
|
)
|
||||||
|
if on_status is not None and attempt > 1:
|
||||||
|
on_status(
|
||||||
|
format_recovered_status(
|
||||||
|
attempt, total_attempts, attempt_seed
|
||||||
|
)
|
||||||
|
)
|
||||||
return response_text
|
return response_text
|
||||||
refusal_count += 1
|
refusal_count += 1
|
||||||
|
retry_reason = "Refusal/deflection detected"
|
||||||
|
else:
|
||||||
|
retry_reason = "Blank response"
|
||||||
|
|
||||||
if attempt < total_attempts:
|
if attempt < total_attempts:
|
||||||
|
if on_status is not None:
|
||||||
|
upcoming_seed = next_seed(options, attempt + 1)
|
||||||
|
upcoming_timeout = next_timeout_secs(timeout_secs, attempt + 1)
|
||||||
|
on_status(
|
||||||
|
format_retry_status(
|
||||||
|
retry_reason,
|
||||||
|
attempt,
|
||||||
|
total_attempts,
|
||||||
|
upcoming_seed,
|
||||||
|
upcoming_timeout,
|
||||||
|
)
|
||||||
|
)
|
||||||
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
||||||
|
|
||||||
record_attempt_info(
|
record_attempt_info(
|
||||||
@@ -331,6 +357,7 @@ class LlamaCppProvider:
|
|||||||
timeout_secs: float = 300.0,
|
timeout_secs: float = 300.0,
|
||||||
max_retries: int = 2,
|
max_retries: int = 2,
|
||||||
attempt_info: dict | None = None,
|
attempt_info: dict | None = None,
|
||||||
|
on_status: Callable[[str], None] | None = None,
|
||||||
) -> BaseModel:
|
) -> BaseModel:
|
||||||
from .chat import chat_structured as _chat_structured_impl
|
from .chat import chat_structured as _chat_structured_impl
|
||||||
|
|
||||||
@@ -376,6 +403,7 @@ class LlamaCppProvider:
|
|||||||
timeout_secs=timeout_secs,
|
timeout_secs=timeout_secs,
|
||||||
embed_fn=embed_fn,
|
embed_fn=embed_fn,
|
||||||
attempt_info=attempt_info,
|
attempt_info=attempt_info,
|
||||||
|
on_status=on_status,
|
||||||
)
|
)
|
||||||
_CHAT_RESPONSE_CACHE.set(cache_key, result.model_dump())
|
_CHAT_RESPONSE_CACHE.set(cache_key, result.model_dump())
|
||||||
return result
|
return result
|
||||||
|
|||||||
@@ -15,12 +15,15 @@ import json
|
|||||||
import logging
|
import logging
|
||||||
import threading
|
import threading
|
||||||
import time
|
import time
|
||||||
|
from collections.abc import Callable
|
||||||
|
|
||||||
from pydantic import BaseModel, ValidationError
|
from pydantic import BaseModel, ValidationError
|
||||||
|
|
||||||
from .provider import Message, ModelInfo, ModelStatus
|
from .provider import Message, ModelInfo, ModelStatus
|
||||||
from .retry import (
|
from .retry import (
|
||||||
RETRY_BACKOFF_SECS,
|
RETRY_BACKOFF_SECS,
|
||||||
|
format_recovered_status,
|
||||||
|
format_retry_status,
|
||||||
is_refusal,
|
is_refusal,
|
||||||
next_seed,
|
next_seed,
|
||||||
next_timeout_secs,
|
next_timeout_secs,
|
||||||
@@ -366,6 +369,7 @@ class OllamaProvider:
|
|||||||
timeout_secs: float = 300.0,
|
timeout_secs: float = 300.0,
|
||||||
max_retries: int = 2,
|
max_retries: int = 2,
|
||||||
attempt_info: dict | None = None,
|
attempt_info: dict | None = None,
|
||||||
|
on_status: Callable[[str], None] | None = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
if any(m.images for m in messages):
|
if any(m.images for m in messages):
|
||||||
await _require_vision_capability(self.host, model, self.headers)
|
await _require_vision_capability(self.host, model, self.headers)
|
||||||
@@ -437,6 +441,7 @@ class OllamaProvider:
|
|||||||
headers=self.headers,
|
headers=self.headers,
|
||||||
)
|
)
|
||||||
response_text = result.get("message", {}).get("content", "")
|
response_text = result.get("message", {}).get("content", "")
|
||||||
|
retry_reason: str | None = None
|
||||||
if response_text.strip():
|
if response_text.strip():
|
||||||
refused = False
|
refused = False
|
||||||
if refusal_cfg and refusal_cfg.get("enabled"):
|
if refusal_cfg and refusal_cfg.get("enabled"):
|
||||||
@@ -456,12 +461,19 @@ class OllamaProvider:
|
|||||||
timeout_secs=attempt_timeout,
|
timeout_secs=attempt_timeout,
|
||||||
refusals=refusal_count,
|
refusals=refusal_count,
|
||||||
)
|
)
|
||||||
|
if on_status is not None and attempt > 1:
|
||||||
|
on_status(
|
||||||
|
format_recovered_status(
|
||||||
|
attempt, total_attempts, attempt_seed
|
||||||
|
)
|
||||||
|
)
|
||||||
return response_text
|
return response_text
|
||||||
# A detected refusal is handled exactly like a blank
|
# A detected refusal is handled exactly like a blank
|
||||||
# response below: fall through to the backoff/retry with a
|
# response below: fall through to the backoff/retry with a
|
||||||
# bumped seed (next_seed), rather than returning the refusal
|
# bumped seed (next_seed), rather than returning the refusal
|
||||||
# text to the caller.
|
# text to the caller.
|
||||||
refusal_count += 1
|
refusal_count += 1
|
||||||
|
retry_reason = "Refusal/deflection detected"
|
||||||
|
|
||||||
# done: false alongside blank content is a distinct signal from
|
# done: false alongside blank content is a distinct signal from
|
||||||
# an ordinary blank generation — it's Ollama answering before
|
# an ordinary blank generation — it's Ollama answering before
|
||||||
@@ -471,8 +483,24 @@ class OllamaProvider:
|
|||||||
# raised on below instead of silently returned like a real
|
# raised on below instead of silently returned like a real
|
||||||
# blank generation would be.
|
# blank generation would be.
|
||||||
incomplete = result.get("done") is False
|
incomplete = result.get("done") is False
|
||||||
|
if retry_reason is None and not response_text.strip():
|
||||||
|
retry_reason = (
|
||||||
|
"Model still loading/swapping" if incomplete else "Blank response"
|
||||||
|
)
|
||||||
|
|
||||||
if attempt < total_attempts:
|
if attempt < total_attempts:
|
||||||
|
if on_status is not None and retry_reason is not None:
|
||||||
|
upcoming_seed = next_seed(options, attempt + 1)
|
||||||
|
upcoming_timeout = next_timeout_secs(timeout_secs, attempt + 1)
|
||||||
|
on_status(
|
||||||
|
format_retry_status(
|
||||||
|
retry_reason,
|
||||||
|
attempt,
|
||||||
|
total_attempts,
|
||||||
|
upcoming_seed,
|
||||||
|
upcoming_timeout,
|
||||||
|
)
|
||||||
|
)
|
||||||
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
||||||
|
|
||||||
record_attempt_info(
|
record_attempt_info(
|
||||||
@@ -506,6 +534,7 @@ class OllamaProvider:
|
|||||||
timeout_secs: float = 300.0,
|
timeout_secs: float = 300.0,
|
||||||
max_retries: int = 2,
|
max_retries: int = 2,
|
||||||
attempt_info: dict | None = None,
|
attempt_info: dict | None = None,
|
||||||
|
on_status: Callable[[str], None] | None = None,
|
||||||
) -> BaseModel:
|
) -> BaseModel:
|
||||||
"""Native ``/api/chat`` + ``"format"`` (grammar-constrained JSON
|
"""Native ``/api/chat`` + ``"format"`` (grammar-constrained JSON
|
||||||
decoding), not the shared pydantic-ai ``chat.py`` helper.
|
decoding), not the shared pydantic-ai ``chat.py`` helper.
|
||||||
@@ -567,6 +596,17 @@ class OllamaProvider:
|
|||||||
attempt_seed = 0
|
attempt_seed = 0
|
||||||
attempt_timeout = timeout_secs
|
attempt_timeout = timeout_secs
|
||||||
|
|
||||||
|
def _emit_retry_status(reason: str, attempt: int) -> None:
|
||||||
|
if on_status is None or attempt >= total_attempts:
|
||||||
|
return
|
||||||
|
upcoming_seed = next_seed(options, attempt + 1)
|
||||||
|
upcoming_timeout = next_timeout_secs(timeout_secs, attempt + 1)
|
||||||
|
on_status(
|
||||||
|
format_retry_status(
|
||||||
|
reason, attempt, total_attempts, upcoming_seed, upcoming_timeout
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
for attempt in range(1, total_attempts + 1):
|
for attempt in range(1, total_attempts + 1):
|
||||||
attempt_options = dict(options) if options else {}
|
attempt_options = dict(options) if options else {}
|
||||||
if attempt > 1:
|
if attempt > 1:
|
||||||
@@ -597,6 +637,7 @@ class OllamaProvider:
|
|||||||
except RuntimeError as exc:
|
except RuntimeError as exc:
|
||||||
last_error = exc
|
last_error = exc
|
||||||
last_invalid_text = str(exc)
|
last_invalid_text = str(exc)
|
||||||
|
_emit_retry_status("Request failed", attempt)
|
||||||
if attempt < total_attempts:
|
if attempt < total_attempts:
|
||||||
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
||||||
continue
|
continue
|
||||||
@@ -607,6 +648,7 @@ class OllamaProvider:
|
|||||||
except ValidationError as exc:
|
except ValidationError as exc:
|
||||||
last_error = exc
|
last_error = exc
|
||||||
last_invalid_text = content
|
last_invalid_text = content
|
||||||
|
_emit_retry_status("Schema validation failed", attempt)
|
||||||
if attempt < total_attempts:
|
if attempt < total_attempts:
|
||||||
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
||||||
continue
|
continue
|
||||||
@@ -628,6 +670,7 @@ class OllamaProvider:
|
|||||||
refusal_count += 1
|
refusal_count += 1
|
||||||
last_error = RuntimeError("refusal/deflection detected")
|
last_error = RuntimeError("refusal/deflection detected")
|
||||||
last_invalid_text = content
|
last_invalid_text = content
|
||||||
|
_emit_retry_status("Refusal/deflection detected", attempt)
|
||||||
if attempt < total_attempts:
|
if attempt < total_attempts:
|
||||||
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
||||||
continue
|
continue
|
||||||
@@ -640,6 +683,10 @@ class OllamaProvider:
|
|||||||
timeout_secs=attempt_timeout,
|
timeout_secs=attempt_timeout,
|
||||||
refusals=refusal_count,
|
refusals=refusal_count,
|
||||||
)
|
)
|
||||||
|
if on_status is not None and attempt > 1:
|
||||||
|
on_status(
|
||||||
|
format_recovered_status(attempt, total_attempts, attempt_seed)
|
||||||
|
)
|
||||||
return parsed
|
return parsed
|
||||||
|
|
||||||
record_attempt_info(
|
record_attempt_info(
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ project-management/ADRs/ADR-007-llm-provider-adapter-pattern.md and
|
|||||||
specs/007-llm-provider-abstraction/contracts/llm_provider_protocol.md.
|
specs/007-llm-provider-abstraction/contracts/llm_provider_protocol.md.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from collections.abc import Callable
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from typing import Literal, Protocol
|
from typing import Literal, Protocol
|
||||||
|
|
||||||
@@ -84,6 +85,7 @@ class LLMProvider(Protocol):
|
|||||||
timeout_secs: float = 300.0,
|
timeout_secs: float = 300.0,
|
||||||
max_retries: int = 2,
|
max_retries: int = 2,
|
||||||
attempt_info: dict | None = None,
|
attempt_info: dict | None = None,
|
||||||
|
on_status: Callable[[str], None] | None = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Free-text chat response.
|
"""Free-text chat response.
|
||||||
|
|
||||||
@@ -110,6 +112,12 @@ class LLMProvider(Protocol):
|
|||||||
count) via ``_llm/retry.py``'s ``record_attempt_info`` — an optional
|
count) via ``_llm/retry.py``'s ``record_attempt_info`` — an optional
|
||||||
out-param, not a return-type change, so existing callers that don't
|
out-param, not a return-type change, so existing callers that don't
|
||||||
pass it see no behavior change.
|
pass it see no behavior change.
|
||||||
|
|
||||||
|
``on_status``, if given, is called synchronously at each retry
|
||||||
|
boundary with a one-line human-readable status (see
|
||||||
|
``_llm/retry.py``'s ``format_retry_status``/``format_recovered_status``)
|
||||||
|
— a live counterpart to ``attempt_info``, which only reports the
|
||||||
|
final outcome after the call returns.
|
||||||
"""
|
"""
|
||||||
...
|
...
|
||||||
|
|
||||||
@@ -122,6 +130,7 @@ class LLMProvider(Protocol):
|
|||||||
timeout_secs: float = 300.0,
|
timeout_secs: float = 300.0,
|
||||||
max_retries: int = 2,
|
max_retries: int = 2,
|
||||||
attempt_info: dict | None = None,
|
attempt_info: dict | None = None,
|
||||||
|
on_status: Callable[[str], None] | None = None,
|
||||||
) -> BaseModel:
|
) -> BaseModel:
|
||||||
"""Schema-validated chat response.
|
"""Schema-validated chat response.
|
||||||
|
|
||||||
@@ -132,7 +141,7 @@ class LLMProvider(Protocol):
|
|||||||
|
|
||||||
ADR-010: see ``chat()`` — same ``options["think"]`` convention,
|
ADR-010: see ``chat()`` — same ``options["think"]`` convention,
|
||||||
same per-provider translation, same escalating per-attempt timeout,
|
same per-provider translation, same escalating per-attempt timeout,
|
||||||
and the same ``attempt_info`` out-param convention.
|
and the same ``attempt_info``/``on_status`` conventions.
|
||||||
"""
|
"""
|
||||||
...
|
...
|
||||||
|
|
||||||
|
|||||||
@@ -89,6 +89,36 @@ def record_attempt_info(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
OnStatus = Callable[[str], None]
|
||||||
|
"""A caller-supplied, synchronous, best-effort progress callback — see
|
||||||
|
``format_retry_status``/``format_recovered_status``. Not async: providers
|
||||||
|
call it inline mid-retry-loop, and the one real implementation
|
||||||
|
(``ChatCompletion``'s closure over ``PromptServer.send_progress_text``) is
|
||||||
|
itself synchronous, so there's nothing to await."""
|
||||||
|
|
||||||
|
|
||||||
|
def format_retry_status(
|
||||||
|
reason: str, attempt: int, total_attempts: int, seed: int, timeout_secs: float
|
||||||
|
) -> str:
|
||||||
|
"""One-line, human-readable status for ``on_status()`` callers — shown
|
||||||
|
live on the node via ComfyUI's ``PromptServer.send_progress_text``
|
||||||
|
(see ``ChatCompletion.chat()``). Centralized so every provider's retry
|
||||||
|
loop describes a retry the same way rather than each inventing its own
|
||||||
|
wording.
|
||||||
|
"""
|
||||||
|
return (
|
||||||
|
f"⚠ {reason} on attempt {attempt}/{total_attempts} — "
|
||||||
|
f"retrying with seed={seed}, timeout={timeout_secs:.0f}s"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def format_recovered_status(attempt: int, total_attempts: int, seed: int) -> str:
|
||||||
|
"""Final status shown once a retry loop succeeds after >1 attempt —
|
||||||
|
lets a live status left over from ``format_retry_status`` resolve to
|
||||||
|
something other than a stale "retrying..." message."""
|
||||||
|
return f"✅ Recovered on attempt {attempt}/{total_attempts} (seed={seed})"
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Refusal/deflection detection
|
# Refusal/deflection detection
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|||||||
@@ -649,6 +649,27 @@ class ChatCompletion:
|
|||||||
# concrete provider ``client`` is.
|
# concrete provider ``client`` is.
|
||||||
attempt_info: dict = {}
|
attempt_info: dict = {}
|
||||||
|
|
||||||
|
# Live counterpart to attempt_info: ComfyUI's own send_progress_text
|
||||||
|
# mechanism (already used by core nodes like PreviewAny/gaussian
|
||||||
|
# splat count) shows this text on the node WHILE it's still
|
||||||
|
# executing, via a "progressText" widget the frontend creates
|
||||||
|
# automatically — no custom JS needed on our side. Best-effort:
|
||||||
|
# a failure here must never take down the actual chat call.
|
||||||
|
on_status = None
|
||||||
|
if unique_id and "comfy" in sys.modules:
|
||||||
|
|
||||||
|
def on_status(message: str) -> None:
|
||||||
|
try:
|
||||||
|
from server import PromptServer
|
||||||
|
|
||||||
|
PromptServer.instance.send_progress_text(message, unique_id)
|
||||||
|
except Exception:
|
||||||
|
logger.debug(
|
||||||
|
"Failed to send live retry status for node %s",
|
||||||
|
unique_id,
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
|
|
||||||
# Provider owns transport, caching, and — for structured_output — the
|
# Provider owns transport, caching, and — for structured_output — the
|
||||||
# tool-calling/retry/validation mechanism (pydantic-ai, ADR-007).
|
# tool-calling/retry/validation mechanism (pydantic-ai, ADR-007).
|
||||||
# ChatCompletion never branches on which concrete provider it got.
|
# ChatCompletion never branches on which concrete provider it got.
|
||||||
@@ -662,6 +683,7 @@ class ChatCompletion:
|
|||||||
timeout_secs=float(timeout_secs),
|
timeout_secs=float(timeout_secs),
|
||||||
max_retries=max_retries,
|
max_retries=max_retries,
|
||||||
attempt_info=attempt_info,
|
attempt_info=attempt_info,
|
||||||
|
on_status=on_status,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -677,6 +699,7 @@ class ChatCompletion:
|
|||||||
timeout_secs=float(timeout_secs),
|
timeout_secs=float(timeout_secs),
|
||||||
max_retries=max_retries,
|
max_retries=max_retries,
|
||||||
attempt_info=attempt_info,
|
attempt_info=attempt_info,
|
||||||
|
on_status=on_status,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
response_text = parsed.model_dump_json()
|
response_text = parsed.model_dump_json()
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
import json
|
||||||
import logging
|
import logging
|
||||||
import random
|
import random
|
||||||
import sys
|
import sys
|
||||||
@@ -7,6 +8,28 @@ from .utils import any_type
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _preview_text(value) -> str:
|
||||||
|
"""Best-effort text preview for RandomChoice's arbitrary-typed output.
|
||||||
|
|
||||||
|
Mirrors ComfyUI core's own ``PreviewAny`` node's value handling (str/
|
||||||
|
number passthrough, else JSON, else ``str()``) rather than inventing a
|
||||||
|
new convention — RandomChoice's output can be anything (an IMAGE
|
||||||
|
tensor, a LATENT, a plain string), so this only needs to be "good
|
||||||
|
enough to glance at," not a faithful repr of every type.
|
||||||
|
"""
|
||||||
|
if isinstance(value, str):
|
||||||
|
return value
|
||||||
|
if isinstance(value, (int, float, bool)):
|
||||||
|
return str(value)
|
||||||
|
try:
|
||||||
|
return json.dumps(value, default=str, indent=2)
|
||||||
|
except Exception:
|
||||||
|
try:
|
||||||
|
return str(value)
|
||||||
|
except Exception:
|
||||||
|
return "<value could not be serialized>"
|
||||||
|
|
||||||
|
|
||||||
class RandomChoice:
|
class RandomChoice:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
pass
|
pass
|
||||||
@@ -25,15 +48,20 @@ class RandomChoice:
|
|||||||
|
|
||||||
FUNCTION = "random_choice"
|
FUNCTION = "random_choice"
|
||||||
|
|
||||||
OUTPUT_NODE = False
|
OUTPUT_NODE = True
|
||||||
|
|
||||||
CATEGORY = "dv/utils"
|
CATEGORY = "dv/utils"
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def IS_CHANGED(s, **kwargs):
|
def IS_CHANGED(s, **kwargs):
|
||||||
return s.random_choice(s, **kwargs)
|
# Unchanged from before the UI-preview addition: returns the raw
|
||||||
|
# picked value (not the ui-wrapped dict random_choice() now returns)
|
||||||
|
# so ComfyUI's change-detection comparison keeps working exactly as
|
||||||
|
# it did previously.
|
||||||
|
return s._pick(**kwargs)
|
||||||
|
|
||||||
def random_choice(self, **kwargs):
|
@staticmethod
|
||||||
|
def _pick(**kwargs):
|
||||||
(
|
(
|
||||||
random.seed(kwargs.get("seed"))
|
random.seed(kwargs.get("seed"))
|
||||||
if kwargs.get("seed")
|
if kwargs.get("seed")
|
||||||
@@ -41,10 +69,13 @@ class RandomChoice:
|
|||||||
)
|
)
|
||||||
input = [i for i in kwargs.items() if i[0] != "seed"]
|
input = [i for i in kwargs.items() if i[0] != "seed"]
|
||||||
logger.debug("RandomChoice inputs: %s", input)
|
logger.debug("RandomChoice inputs: %s", input)
|
||||||
|
return random.choice(input)[1]
|
||||||
|
|
||||||
|
def random_choice(self, **kwargs):
|
||||||
try:
|
try:
|
||||||
choice = random.choice(input)[1]
|
choice = self._pick(**kwargs)
|
||||||
logger.debug("RandomChoice chose: %s", choice)
|
logger.debug("RandomChoice chose: %s", choice)
|
||||||
return (choice,)
|
return {"ui": {"text": [_preview_text(choice)]}, "result": (choice,)}
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error("RandomChoice: unexpected error: %s", e)
|
logger.error("RandomChoice: unexpected error: %s", e)
|
||||||
raise
|
raise
|
||||||
|
|||||||
@@ -0,0 +1,57 @@
|
|||||||
|
/**
|
||||||
|
* preview_text.js — read-only output preview for comfydv's OUTPUT_NODE=True
|
||||||
|
* nodes that return a ComfyUI "ui": {"text": [...]} payload
|
||||||
|
* (ChatCompletion, FormatString, RandomChoice).
|
||||||
|
*
|
||||||
|
* ComfyUI does NOT auto-render an arbitrary node's ui.text — each node type
|
||||||
|
* that wants one implements its own onExecuted handler. This mirrors core's
|
||||||
|
* own ``PreviewAny`` node (comfy_extras/nodes_preview_any.py +
|
||||||
|
* "Comfy.PreviewAny" in the frontend bundle) minus its Markdown/Plaintext
|
||||||
|
* toggle, which none of these three nodes need.
|
||||||
|
*/
|
||||||
|
|
||||||
|
import { app } from "../../scripts/app.js";
|
||||||
|
import { ComfyWidgets } from "../../scripts/widgets.js";
|
||||||
|
|
||||||
|
const PREVIEW_NODES = new Set(["ChatCompletion", "FormatString", "RandomChoice"]);
|
||||||
|
|
||||||
|
app.registerExtension({
|
||||||
|
name: "comfydv.previewText",
|
||||||
|
|
||||||
|
async beforeRegisterNodeDef(nodeType, nodeData) {
|
||||||
|
if (!PREVIEW_NODES.has(nodeData.name)) return;
|
||||||
|
|
||||||
|
const onNodeCreated = nodeType.prototype.onNodeCreated;
|
||||||
|
nodeType.prototype.onNodeCreated = function () {
|
||||||
|
const result = onNodeCreated?.apply(this, arguments);
|
||||||
|
|
||||||
|
const widget = ComfyWidgets.STRING(
|
||||||
|
this,
|
||||||
|
"comfydv_preview_text",
|
||||||
|
["STRING", { multiline: true }],
|
||||||
|
app
|
||||||
|
).widget;
|
||||||
|
widget.label = "Preview";
|
||||||
|
widget.options.read_only = true;
|
||||||
|
// Not a real input — nothing to save/replay in the saved
|
||||||
|
// workflow JSON, and read-only anyway.
|
||||||
|
widget.options.serialize = false;
|
||||||
|
widget.serialize = false;
|
||||||
|
widget.inputEl.readOnly = true;
|
||||||
|
|
||||||
|
return result;
|
||||||
|
};
|
||||||
|
|
||||||
|
const onExecuted = nodeType.prototype.onExecuted;
|
||||||
|
nodeType.prototype.onExecuted = function (message) {
|
||||||
|
onExecuted?.apply(this, arguments);
|
||||||
|
|
||||||
|
const widget = this.widgets?.find(w => w.name === "comfydv_preview_text");
|
||||||
|
if (!widget) return;
|
||||||
|
|
||||||
|
const text = message?.text ?? "";
|
||||||
|
widget.value = Array.isArray(text) ? (text.join("\n\n") ?? "") : text;
|
||||||
|
this.setDirtyCanvas(true, true);
|
||||||
|
};
|
||||||
|
},
|
||||||
|
});
|
||||||
@@ -47,6 +47,7 @@ class _FakeProvider:
|
|||||||
timeout_secs=300.0,
|
timeout_secs=300.0,
|
||||||
max_retries=2,
|
max_retries=2,
|
||||||
attempt_info=None,
|
attempt_info=None,
|
||||||
|
on_status=None,
|
||||||
):
|
):
|
||||||
self.calls.append(("chat", model))
|
self.calls.append(("chat", model))
|
||||||
if attempt_info is not None:
|
if attempt_info is not None:
|
||||||
|
|||||||
@@ -453,6 +453,33 @@ def test_chat_attempt_info_populated_on_success(monkeypatch):
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_chat_on_status_called_on_retry_and_recovery(monkeypatch):
|
||||||
|
calls = []
|
||||||
|
|
||||||
|
async def fake_post(url, payload, *, timeout=120.0, headers=None):
|
||||||
|
calls.append(payload)
|
||||||
|
if len(calls) == 1:
|
||||||
|
return {"choices": [{"message": {"content": "I cannot generate that."}}]}
|
||||||
|
return {"choices": [{"message": {"content": "a real, on-topic answer"}}]}
|
||||||
|
|
||||||
|
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
|
||||||
|
monkeypatch.setattr(provider_mod.asyncio, "sleep", _fake_sleep)
|
||||||
|
|
||||||
|
statuses = []
|
||||||
|
_run_async(
|
||||||
|
LlamaCppProvider("http://localhost:8080").chat(
|
||||||
|
"gemma-3-4b",
|
||||||
|
[Message(role="user", content="hi")],
|
||||||
|
options={"refusal_retry": {"enabled": True, "embedding_model": ""}},
|
||||||
|
on_status=statuses.append,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(statuses) == 2
|
||||||
|
assert "Refusal/deflection detected" in statuses[0]
|
||||||
|
assert "Recovered" in statuses[1]
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# chat_structured — zero new logic, delegates to the shared helper unchanged
|
# chat_structured — zero new logic, delegates to the shared helper unchanged
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -538,6 +565,35 @@ def test_chat_structured_forwards_attempt_info(monkeypatch):
|
|||||||
assert captured["attempt_info"] is attempt_info
|
assert captured["attempt_info"] is attempt_info
|
||||||
|
|
||||||
|
|
||||||
|
def test_chat_structured_forwards_on_status(monkeypatch):
|
||||||
|
from pydantic import BaseModel
|
||||||
|
|
||||||
|
class Widget(BaseModel):
|
||||||
|
name: str
|
||||||
|
|
||||||
|
captured = {}
|
||||||
|
|
||||||
|
async def fake_chat_structured(**kwargs):
|
||||||
|
captured.update(kwargs)
|
||||||
|
return Widget(name="x")
|
||||||
|
|
||||||
|
monkeypatch.setattr("comfydv._llm.chat.chat_structured", fake_chat_structured)
|
||||||
|
|
||||||
|
def on_status(msg):
|
||||||
|
pass
|
||||||
|
|
||||||
|
_run_async(
|
||||||
|
LlamaCppProvider("http://localhost:8080").chat_structured(
|
||||||
|
"gemma-3-4b",
|
||||||
|
[Message(role="user", content="hi")],
|
||||||
|
Widget,
|
||||||
|
on_status=on_status,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert captured["on_status"] is on_status
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# _fetch_models — name-only view used by ComfyUI's /dv/ollama/models?backend=
|
# _fetch_models — name-only view used by ComfyUI's /dv/ollama/models?backend=
|
||||||
# llamacpp route (the JS refresh button / node-creation auto-populate).
|
# llamacpp route (the JS refresh button / node-creation auto-populate).
|
||||||
|
|||||||
@@ -212,6 +212,30 @@ def test_chat_structured_attempt_info_populated_on_exhaustion(monkeypatch):
|
|||||||
assert attempt_info["attempts"] == 2
|
assert attempt_info["attempts"] == 2
|
||||||
|
|
||||||
|
|
||||||
|
def test_chat_structured_on_status_called_on_retry_and_recovery(monkeypatch):
|
||||||
|
bad = ValidationError.from_exception_data("Widget", [])
|
||||||
|
fake = _FakeAgent([bad, _Widget(name="b", count=2)])
|
||||||
|
monkeypatch.setattr(chat_mod, "_build_agent", lambda **kw: fake)
|
||||||
|
|
||||||
|
statuses = []
|
||||||
|
_run_async(
|
||||||
|
chat_mod.chat_structured(
|
||||||
|
base_url="http://localhost:11434/v1",
|
||||||
|
model="llama3",
|
||||||
|
messages=_messages(),
|
||||||
|
schema=_Widget,
|
||||||
|
max_retries=2,
|
||||||
|
on_status=statuses.append,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(statuses) == 2
|
||||||
|
assert "Structured output failed" in statuses[0]
|
||||||
|
assert "attempt 1/3" in statuses[0]
|
||||||
|
assert "Recovered" in statuses[1]
|
||||||
|
assert "attempt 2/3" in statuses[1]
|
||||||
|
|
||||||
|
|
||||||
def test_chat_structured_exhausted_retries_raises_runtime_error(monkeypatch):
|
def test_chat_structured_exhausted_retries_raises_runtime_error(monkeypatch):
|
||||||
bad = ValidationError.from_exception_data("Widget", [])
|
bad = ValidationError.from_exception_data("Widget", [])
|
||||||
fake = _FakeAgent([bad, bad, bad]) # max_retries=2 -> 3 total attempts
|
fake = _FakeAgent([bad, bad, bad]) # max_retries=2 -> 3 total attempts
|
||||||
@@ -481,6 +505,30 @@ def test_chat_structured_retries_on_refusal_and_returns_clean_second_attempt(
|
|||||||
assert len(fake.calls) == 2
|
assert len(fake.calls) == 2
|
||||||
|
|
||||||
|
|
||||||
|
def test_chat_structured_on_status_reports_refusal_reason(monkeypatch):
|
||||||
|
refused = _Widget(name="I cannot generate that content.", count=1)
|
||||||
|
clean = _Widget(name="clean", count=2)
|
||||||
|
fake = _FakeAgent([refused, clean])
|
||||||
|
monkeypatch.setattr(chat_mod, "_build_agent", lambda **kw: fake)
|
||||||
|
|
||||||
|
statuses = []
|
||||||
|
_run_async(
|
||||||
|
chat_mod.chat_structured(
|
||||||
|
base_url="http://localhost:11434/v1",
|
||||||
|
model="llama3",
|
||||||
|
messages=_messages(),
|
||||||
|
schema=_Widget,
|
||||||
|
options={"refusal_retry": {"enabled": True, "embedding_model": ""}},
|
||||||
|
max_retries=2,
|
||||||
|
on_status=statuses.append,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(statuses) == 2
|
||||||
|
assert "Refusal/deflection detected" in statuses[0]
|
||||||
|
assert "Recovered" in statuses[1]
|
||||||
|
|
||||||
|
|
||||||
def test_chat_structured_refusal_retry_disabled_returns_refusal_unchanged(
|
def test_chat_structured_refusal_retry_disabled_returns_refusal_unchanged(
|
||||||
monkeypatch,
|
monkeypatch,
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -15,6 +15,8 @@ from comfydv._llm.ollama_provider import _run_async
|
|||||||
from comfydv._llm.retry import (
|
from comfydv._llm.retry import (
|
||||||
REFUSAL_EXEMPLARS,
|
REFUSAL_EXEMPLARS,
|
||||||
cosine_similarity,
|
cosine_similarity,
|
||||||
|
format_recovered_status,
|
||||||
|
format_retry_status,
|
||||||
is_ambiguous,
|
is_ambiguous,
|
||||||
is_lexical_refusal,
|
is_lexical_refusal,
|
||||||
is_refusal,
|
is_refusal,
|
||||||
@@ -81,6 +83,20 @@ class TestRecordAttemptInfo:
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class TestFormatStatus:
|
||||||
|
def test_retry_status_includes_reason_attempt_seed_and_timeout(self):
|
||||||
|
msg = format_retry_status("Blank response", 1, 3, seed=1, timeout_secs=200.0)
|
||||||
|
assert "Blank response" in msg
|
||||||
|
assert "attempt 1/3" in msg
|
||||||
|
assert "seed=1" in msg
|
||||||
|
assert "timeout=200s" in msg
|
||||||
|
|
||||||
|
def test_recovered_status_includes_attempt_and_seed(self):
|
||||||
|
msg = format_recovered_status(2, 3, seed=1)
|
||||||
|
assert "attempt 2/3" in msg
|
||||||
|
assert "seed=1" in msg
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Refusal/deflection detection
|
# Refusal/deflection detection
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|||||||
@@ -99,6 +99,7 @@ class _FakeProvider:
|
|||||||
timeout_secs=300.0,
|
timeout_secs=300.0,
|
||||||
max_retries=2,
|
max_retries=2,
|
||||||
attempt_info=None,
|
attempt_info=None,
|
||||||
|
on_status=None,
|
||||||
):
|
):
|
||||||
self.calls.append(("chat", model, messages, options, timeout_secs, max_retries))
|
self.calls.append(("chat", model, messages, options, timeout_secs, max_retries))
|
||||||
if attempt_info is not None:
|
if attempt_info is not None:
|
||||||
@@ -121,6 +122,7 @@ class _FakeProvider:
|
|||||||
timeout_secs=300.0,
|
timeout_secs=300.0,
|
||||||
max_retries=2,
|
max_retries=2,
|
||||||
attempt_info=None,
|
attempt_info=None,
|
||||||
|
on_status=None,
|
||||||
):
|
):
|
||||||
self.calls.append(
|
self.calls.append(
|
||||||
(
|
(
|
||||||
|
|||||||
@@ -1007,6 +1007,56 @@ def test_chat_attempt_info_reflects_refusal_retry(monkeypatch):
|
|||||||
assert attempt_info["seed"] == 1 # next_seed(): attempt 2 -> seed 1
|
assert attempt_info["seed"] == 1 # next_seed(): attempt 2 -> seed 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_chat_on_status_called_on_retry_and_recovery(monkeypatch):
|
||||||
|
"""Live counterpart to attempt_info: called mid-loop (not just at the
|
||||||
|
end) so a caller (ChatCompletion's PromptServer.send_progress_text
|
||||||
|
closure) can show retry progress while the node is still executing."""
|
||||||
|
calls = []
|
||||||
|
|
||||||
|
async def fake_post(url, payload, *, timeout=120.0, headers=None):
|
||||||
|
calls.append(payload)
|
||||||
|
if len(calls) == 1:
|
||||||
|
return {"message": {"content": "I cannot generate that for you."}}
|
||||||
|
return {"message": {"content": "a real, on-topic answer"}}
|
||||||
|
|
||||||
|
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
|
||||||
|
monkeypatch.setattr(provider_mod.asyncio, "sleep", _fake_sleep)
|
||||||
|
|
||||||
|
statuses = []
|
||||||
|
_run_async(
|
||||||
|
OllamaProvider("http://localhost:11434").chat(
|
||||||
|
"llama3",
|
||||||
|
[Message(role="user", content="hi")],
|
||||||
|
options={"refusal_retry": {"enabled": True, "embedding_model": ""}},
|
||||||
|
on_status=statuses.append,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(statuses) == 2
|
||||||
|
assert "Refusal/deflection detected" in statuses[0]
|
||||||
|
assert "attempt 1/3" in statuses[0]
|
||||||
|
assert "Recovered" in statuses[1]
|
||||||
|
assert "attempt 2/3" in statuses[1]
|
||||||
|
|
||||||
|
|
||||||
|
def test_chat_on_status_not_called_when_first_attempt_succeeds(monkeypatch):
|
||||||
|
async def fake_post(url, payload, *, timeout=120.0, headers=None):
|
||||||
|
return {"message": {"content": "ok"}}
|
||||||
|
|
||||||
|
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
|
||||||
|
|
||||||
|
statuses = []
|
||||||
|
_run_async(
|
||||||
|
OllamaProvider("http://localhost:11434").chat(
|
||||||
|
"llama3",
|
||||||
|
[Message(role="user", content="hi")],
|
||||||
|
on_status=statuses.append,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert statuses == []
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# refusal-retry — comfydv-level "refusal_retry" options convention (see
|
# refusal-retry — comfydv-level "refusal_retry" options convention (see
|
||||||
# OllamaOptionRefusalRetry / comfydv._llm.retry.is_refusal)
|
# OllamaOptionRefusalRetry / comfydv._llm.retry.is_refusal)
|
||||||
|
|||||||
@@ -0,0 +1,74 @@
|
|||||||
|
"""Tests for comfydv.random_choice.RandomChoice.
|
||||||
|
|
||||||
|
Covers the UI-preview addition (OUTPUT_NODE=True + a "ui": {"text": [...]}
|
||||||
|
return, mirroring ChatCompletion/FormatString so all three get a visible
|
||||||
|
text preview via src/js/preview_text.js) without changing IS_CHANGED's
|
||||||
|
change-detection semantics.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
|
||||||
|
from comfydv.random_choice import RandomChoice, _preview_text
|
||||||
|
|
||||||
|
|
||||||
|
def test_output_node_is_true():
|
||||||
|
assert RandomChoice.OUTPUT_NODE is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_random_choice_returns_ui_result_dict():
|
||||||
|
ret = RandomChoice().random_choice(input1="a", seed=42)
|
||||||
|
assert isinstance(ret, dict)
|
||||||
|
assert "ui" in ret
|
||||||
|
assert "result" in ret
|
||||||
|
assert ret["result"] == ("a",)
|
||||||
|
|
||||||
|
|
||||||
|
def test_ui_text_matches_the_chosen_value_for_a_string():
|
||||||
|
ret = RandomChoice().random_choice(input1="hello", seed=42)
|
||||||
|
assert ret["ui"]["text"] == ["hello"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_ui_text_for_a_number_is_stringified():
|
||||||
|
ret = RandomChoice().random_choice(input1=7, seed=42)
|
||||||
|
assert ret["ui"]["text"] == ["7"]
|
||||||
|
assert ret["result"] == (7,)
|
||||||
|
|
||||||
|
|
||||||
|
def test_seed_pins_the_choice_deterministically():
|
||||||
|
ret1 = RandomChoice().random_choice(input1="a", input2="b", input3="c", seed=42)
|
||||||
|
ret2 = RandomChoice().random_choice(input1="a", input2="b", input3="c", seed=42)
|
||||||
|
assert ret1["result"] == ret2["result"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_is_changed_returns_raw_pick_not_ui_wrapped_dict():
|
||||||
|
"""Regression guard: IS_CHANGED must keep returning the same shape it
|
||||||
|
did before the ui-preview addition (the raw picked value), not the new
|
||||||
|
{"ui": ..., "result": ...} dict random_choice() now returns — otherwise
|
||||||
|
ComfyUI's change-detection comparison would be comparing dicts full of
|
||||||
|
UI-only text noise instead of the actual output value."""
|
||||||
|
result = RandomChoice.IS_CHANGED(input1="only-choice", seed=42)
|
||||||
|
assert result == "only-choice"
|
||||||
|
|
||||||
|
|
||||||
|
class TestPreviewText:
|
||||||
|
def test_string_passthrough(self):
|
||||||
|
assert _preview_text("hello") == "hello"
|
||||||
|
|
||||||
|
def test_number_stringified(self):
|
||||||
|
assert _preview_text(42) == "42"
|
||||||
|
assert _preview_text(3.14) == "3.14"
|
||||||
|
assert _preview_text(True) == "True"
|
||||||
|
|
||||||
|
def test_list_json_dumped(self):
|
||||||
|
assert json.loads(_preview_text([1, 2, 3])) == [1, 2, 3]
|
||||||
|
|
||||||
|
def test_unserializable_falls_back_to_str(self):
|
||||||
|
class Weird:
|
||||||
|
def __repr__(self):
|
||||||
|
return "<Weird thing>"
|
||||||
|
|
||||||
|
# json.dumps(default=str) actually succeeds here (falls back to
|
||||||
|
# str() per-value), so this exercises the "else JSON" branch
|
||||||
|
# rather than the outer except — confirms it never raises either
|
||||||
|
# way.
|
||||||
|
assert "<Weird thing>" in _preview_text(Weird())
|
||||||
Reference in New Issue
Block a user