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:
James Veitch
2026-07-29 23:43:13 +01:00
co-authored by Claude Sonnet 5
parent 7a40d22230
commit f16da579bc
15 changed files with 500 additions and 6 deletions
+22
View File
@@ -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)
+28
View File
@@ -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
+47
View File
@@ -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(
+10 -1
View File
@@ -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.
""" """
... ...
+30
View File
@@ -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
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
+23
View File
@@ -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()
+36 -5
View File
@@ -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
+57
View File
@@ -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);
};
},
});
+1
View File
@@ -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:
+56
View File
@@ -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).
+48
View File
@@ -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,
): ):
+16
View File
@@ -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
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
+2
View File
@@ -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(
( (
+50
View File
@@ -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)
+74
View File
@@ -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())