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
|
||||
from collections.abc import Callable
|
||||
from typing import cast
|
||||
|
||||
from pydantic import BaseModel, ValidationError
|
||||
@@ -49,6 +50,8 @@ from .provider import Message
|
||||
from .retry import (
|
||||
RETRY_BACKOFF_SECS,
|
||||
EmbedFn,
|
||||
format_recovered_status,
|
||||
format_retry_status,
|
||||
is_refusal,
|
||||
next_seed,
|
||||
next_timeout_secs,
|
||||
@@ -135,6 +138,7 @@ async def chat_structured(
|
||||
timeout_secs: float = 300.0,
|
||||
embed_fn: EmbedFn | None = None,
|
||||
attempt_info: dict | None = None,
|
||||
on_status: Callable[[str], None] | None = None,
|
||||
) -> BaseModel:
|
||||
"""Call ``model`` at ``base_url`` (an OpenAI-compatible ``/v1`` root) and
|
||||
return a validated instance of ``schema``.
|
||||
@@ -208,6 +212,18 @@ async def chat_structured(
|
||||
refusal_count = 0
|
||||
attempt_seed = (options or {}).get("seed", 0) if isinstance(options, dict) else 0
|
||||
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):
|
||||
attempt_timeout = next_timeout_secs(timeout_secs, attempt)
|
||||
# Rebuilt each attempt so the escalated timeout actually takes
|
||||
@@ -282,6 +298,7 @@ async def chat_structured(
|
||||
refusal_count += 1
|
||||
last_error = RuntimeError("refusal/deflection detected")
|
||||
last_invalid_text = content
|
||||
_emit_retry_status("Refusal/deflection detected", attempt)
|
||||
if attempt < total_attempts:
|
||||
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
||||
continue
|
||||
@@ -292,10 +309,15 @@ async def chat_structured(
|
||||
timeout_secs=attempt_timeout,
|
||||
refusals=refusal_count,
|
||||
)
|
||||
if on_status is not None and attempt > 1:
|
||||
on_status(
|
||||
format_recovered_status(attempt, total_attempts, attempt_seed)
|
||||
)
|
||||
return output
|
||||
except _STRUCTURED_OUTPUT_FAILURE_EXCEPTIONS as exc:
|
||||
last_error = exc
|
||||
last_invalid_text = str(exc)
|
||||
_emit_retry_status("Structured output failed", attempt)
|
||||
if attempt < total_attempts:
|
||||
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
||||
|
||||
|
||||
@@ -15,6 +15,7 @@ exist otherwise (spec.md FR-006).
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from collections.abc import Callable
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
@@ -29,6 +30,8 @@ from .ollama_provider import (
|
||||
from .provider import Message, ModelInfo, ModelStatus
|
||||
from .retry import (
|
||||
RETRY_BACKOFF_SECS,
|
||||
format_recovered_status,
|
||||
format_retry_status,
|
||||
is_refusal,
|
||||
next_seed,
|
||||
next_timeout_secs,
|
||||
@@ -197,6 +200,7 @@ class LlamaCppProvider:
|
||||
timeout_secs: float = 300.0,
|
||||
max_retries: int = 2,
|
||||
attempt_info: dict | None = None,
|
||||
on_status: Callable[[str], None] | None = None,
|
||||
) -> str:
|
||||
payload_messages = [_to_openai_message(m) for m in messages]
|
||||
options, think = _pop_think(options)
|
||||
@@ -285,6 +289,7 @@ class LlamaCppProvider:
|
||||
if choices
|
||||
else ""
|
||||
)
|
||||
retry_reason: str | None = None
|
||||
if response_text.strip():
|
||||
refused = False
|
||||
if refusal_cfg and refusal_cfg.get("enabled"):
|
||||
@@ -304,10 +309,31 @@ class LlamaCppProvider:
|
||||
timeout_secs=attempt_timeout,
|
||||
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
|
||||
refusal_count += 1
|
||||
retry_reason = "Refusal/deflection detected"
|
||||
else:
|
||||
retry_reason = "Blank response"
|
||||
|
||||
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)
|
||||
|
||||
record_attempt_info(
|
||||
@@ -331,6 +357,7 @@ class LlamaCppProvider:
|
||||
timeout_secs: float = 300.0,
|
||||
max_retries: int = 2,
|
||||
attempt_info: dict | None = None,
|
||||
on_status: Callable[[str], None] | None = None,
|
||||
) -> BaseModel:
|
||||
from .chat import chat_structured as _chat_structured_impl
|
||||
|
||||
@@ -376,6 +403,7 @@ class LlamaCppProvider:
|
||||
timeout_secs=timeout_secs,
|
||||
embed_fn=embed_fn,
|
||||
attempt_info=attempt_info,
|
||||
on_status=on_status,
|
||||
)
|
||||
_CHAT_RESPONSE_CACHE.set(cache_key, result.model_dump())
|
||||
return result
|
||||
|
||||
@@ -15,12 +15,15 @@ import json
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
from .provider import Message, ModelInfo, ModelStatus
|
||||
from .retry import (
|
||||
RETRY_BACKOFF_SECS,
|
||||
format_recovered_status,
|
||||
format_retry_status,
|
||||
is_refusal,
|
||||
next_seed,
|
||||
next_timeout_secs,
|
||||
@@ -366,6 +369,7 @@ class OllamaProvider:
|
||||
timeout_secs: float = 300.0,
|
||||
max_retries: int = 2,
|
||||
attempt_info: dict | None = None,
|
||||
on_status: Callable[[str], None] | None = None,
|
||||
) -> str:
|
||||
if any(m.images for m in messages):
|
||||
await _require_vision_capability(self.host, model, self.headers)
|
||||
@@ -437,6 +441,7 @@ class OllamaProvider:
|
||||
headers=self.headers,
|
||||
)
|
||||
response_text = result.get("message", {}).get("content", "")
|
||||
retry_reason: str | None = None
|
||||
if response_text.strip():
|
||||
refused = False
|
||||
if refusal_cfg and refusal_cfg.get("enabled"):
|
||||
@@ -456,12 +461,19 @@ class OllamaProvider:
|
||||
timeout_secs=attempt_timeout,
|
||||
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
|
||||
# A detected refusal is handled exactly like a blank
|
||||
# response below: fall through to the backoff/retry with a
|
||||
# bumped seed (next_seed), rather than returning the refusal
|
||||
# text to the caller.
|
||||
refusal_count += 1
|
||||
retry_reason = "Refusal/deflection detected"
|
||||
|
||||
# done: false alongside blank content is a distinct signal from
|
||||
# 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
|
||||
# blank generation would be.
|
||||
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 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)
|
||||
|
||||
record_attempt_info(
|
||||
@@ -506,6 +534,7 @@ class OllamaProvider:
|
||||
timeout_secs: float = 300.0,
|
||||
max_retries: int = 2,
|
||||
attempt_info: dict | None = None,
|
||||
on_status: Callable[[str], None] | None = None,
|
||||
) -> BaseModel:
|
||||
"""Native ``/api/chat`` + ``"format"`` (grammar-constrained JSON
|
||||
decoding), not the shared pydantic-ai ``chat.py`` helper.
|
||||
@@ -567,6 +596,17 @@ class OllamaProvider:
|
||||
attempt_seed = 0
|
||||
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):
|
||||
attempt_options = dict(options) if options else {}
|
||||
if attempt > 1:
|
||||
@@ -597,6 +637,7 @@ class OllamaProvider:
|
||||
except RuntimeError as exc:
|
||||
last_error = exc
|
||||
last_invalid_text = str(exc)
|
||||
_emit_retry_status("Request failed", attempt)
|
||||
if attempt < total_attempts:
|
||||
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
||||
continue
|
||||
@@ -607,6 +648,7 @@ class OllamaProvider:
|
||||
except ValidationError as exc:
|
||||
last_error = exc
|
||||
last_invalid_text = content
|
||||
_emit_retry_status("Schema validation failed", attempt)
|
||||
if attempt < total_attempts:
|
||||
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
||||
continue
|
||||
@@ -628,6 +670,7 @@ class OllamaProvider:
|
||||
refusal_count += 1
|
||||
last_error = RuntimeError("refusal/deflection detected")
|
||||
last_invalid_text = content
|
||||
_emit_retry_status("Refusal/deflection detected", attempt)
|
||||
if attempt < total_attempts:
|
||||
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
||||
continue
|
||||
@@ -640,6 +683,10 @@ class OllamaProvider:
|
||||
timeout_secs=attempt_timeout,
|
||||
refusals=refusal_count,
|
||||
)
|
||||
if on_status is not None and attempt > 1:
|
||||
on_status(
|
||||
format_recovered_status(attempt, total_attempts, attempt_seed)
|
||||
)
|
||||
return parsed
|
||||
|
||||
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.
|
||||
"""
|
||||
|
||||
from collections.abc import Callable
|
||||
from enum import Enum
|
||||
from typing import Literal, Protocol
|
||||
|
||||
@@ -84,6 +85,7 @@ class LLMProvider(Protocol):
|
||||
timeout_secs: float = 300.0,
|
||||
max_retries: int = 2,
|
||||
attempt_info: dict | None = None,
|
||||
on_status: Callable[[str], None] | None = None,
|
||||
) -> str:
|
||||
"""Free-text chat response.
|
||||
|
||||
@@ -110,6 +112,12 @@ class LLMProvider(Protocol):
|
||||
count) via ``_llm/retry.py``'s ``record_attempt_info`` — an optional
|
||||
out-param, not a return-type change, so existing callers that don't
|
||||
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,
|
||||
max_retries: int = 2,
|
||||
attempt_info: dict | None = None,
|
||||
on_status: Callable[[str], None] | None = None,
|
||||
) -> BaseModel:
|
||||
"""Schema-validated chat response.
|
||||
|
||||
@@ -132,7 +141,7 @@ class LLMProvider(Protocol):
|
||||
|
||||
ADR-010: see ``chat()`` — same ``options["think"]`` convention,
|
||||
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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -649,6 +649,27 @@ class ChatCompletion:
|
||||
# concrete provider ``client`` is.
|
||||
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
|
||||
# tool-calling/retry/validation mechanism (pydantic-ai, ADR-007).
|
||||
# ChatCompletion never branches on which concrete provider it got.
|
||||
@@ -662,6 +683,7 @@ class ChatCompletion:
|
||||
timeout_secs=float(timeout_secs),
|
||||
max_retries=max_retries,
|
||||
attempt_info=attempt_info,
|
||||
on_status=on_status,
|
||||
)
|
||||
)
|
||||
else:
|
||||
@@ -677,6 +699,7 @@ class ChatCompletion:
|
||||
timeout_secs=float(timeout_secs),
|
||||
max_retries=max_retries,
|
||||
attempt_info=attempt_info,
|
||||
on_status=on_status,
|
||||
)
|
||||
)
|
||||
response_text = parsed.model_dump_json()
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import json
|
||||
import logging
|
||||
import random
|
||||
import sys
|
||||
@@ -7,6 +8,28 @@ from .utils import any_type
|
||||
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:
|
||||
def __init__(self):
|
||||
pass
|
||||
@@ -25,15 +48,20 @@ class RandomChoice:
|
||||
|
||||
FUNCTION = "random_choice"
|
||||
|
||||
OUTPUT_NODE = False
|
||||
OUTPUT_NODE = True
|
||||
|
||||
CATEGORY = "dv/utils"
|
||||
|
||||
@classmethod
|
||||
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"))
|
||||
if kwargs.get("seed")
|
||||
@@ -41,10 +69,13 @@ class RandomChoice:
|
||||
)
|
||||
input = [i for i in kwargs.items() if i[0] != "seed"]
|
||||
logger.debug("RandomChoice inputs: %s", input)
|
||||
return random.choice(input)[1]
|
||||
|
||||
def random_choice(self, **kwargs):
|
||||
try:
|
||||
choice = random.choice(input)[1]
|
||||
choice = self._pick(**kwargs)
|
||||
logger.debug("RandomChoice chose: %s", choice)
|
||||
return (choice,)
|
||||
return {"ui": {"text": [_preview_text(choice)]}, "result": (choice,)}
|
||||
except Exception as e:
|
||||
logger.error("RandomChoice: unexpected error: %s", e)
|
||||
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,
|
||||
max_retries=2,
|
||||
attempt_info=None,
|
||||
on_status=None,
|
||||
):
|
||||
self.calls.append(("chat", model))
|
||||
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
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -538,6 +565,35 @@ def test_chat_structured_forwards_attempt_info(monkeypatch):
|
||||
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=
|
||||
# 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
|
||||
|
||||
|
||||
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):
|
||||
bad = ValidationError.from_exception_data("Widget", [])
|
||||
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
|
||||
|
||||
|
||||
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(
|
||||
monkeypatch,
|
||||
):
|
||||
|
||||
@@ -15,6 +15,8 @@ from comfydv._llm.ollama_provider import _run_async
|
||||
from comfydv._llm.retry import (
|
||||
REFUSAL_EXEMPLARS,
|
||||
cosine_similarity,
|
||||
format_recovered_status,
|
||||
format_retry_status,
|
||||
is_ambiguous,
|
||||
is_lexical_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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -99,6 +99,7 @@ class _FakeProvider:
|
||||
timeout_secs=300.0,
|
||||
max_retries=2,
|
||||
attempt_info=None,
|
||||
on_status=None,
|
||||
):
|
||||
self.calls.append(("chat", model, messages, options, timeout_secs, max_retries))
|
||||
if attempt_info is not None:
|
||||
@@ -121,6 +122,7 @@ class _FakeProvider:
|
||||
timeout_secs=300.0,
|
||||
max_retries=2,
|
||||
attempt_info=None,
|
||||
on_status=None,
|
||||
):
|
||||
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
|
||||
|
||||
|
||||
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
|
||||
# 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