Author SHA1 Message Date
James VeitchandClaude Sonnet 5 f16da579bc 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>
2026-07-29 23:43:13 +01:00
James VeitchandClaude Sonnet 5 7a40d22230 feat(llm): escalate retry timeouts, expose seed_used, add retry UI feedback
Three follow-ups to the refusal-retry feature:

- Each retry's request timeout now escalates (next_timeout_secs: base *
  attempt) across every chat()/chat_structured() retry loop, instead of
  reusing the same budget that just ran out.
- LLMProvider.chat()/chat_structured() gain an optional attempt_info
  out-param (record_attempt_info), populated with the seed/timeout
  actually used and the attempt/refusal counts — backward compatible,
  no return-type change.
- ChatCompletion exposes a new seed_used output and prepends a status
  line to its existing text preview when a retry/refusal happened, so
  both are visible without any frontend/JS changes.

seed_used is appended as the LAST output (after any dynamic
structured-output fields), not inserted after model_name — dynamic
fields must keep starting at their existing index (3) so already-wired
workflows (workflows/ltx-i2v-pipeline.json, and the still-open PR #35
canvas-format file) that link a FormatString input to a ChatCompletion
field by output index aren't silently repointed at the wrong socket.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
2026-07-29 22:19:22 +01:00
James VeitchandClaude Sonnet 5 42c614d5f8 feat(llm): add user-configurable custom_phrases to refusal detection
Ship a comma-separated custom_phrases field on OllamaOptionRefusalRetry
so a new deflection phrasing a specific model uses can be added at
runtime, without waiting on a code change/release. Each phrase is
checked as a free, case-insensitive substring match (no embedding
model required) and, when embedding_model is set, folded in as extra
exemplars for the similarity check too.

Also ship two more default patterns spotted this session: bare
"cannot generate" (no leading "I") and "I am/I'm restricted from" —
both verified against the existing ordinary-content regression cases
to confirm no new false positives.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
2026-07-29 17:44:32 +01:00
James VeitchandClaude Sonnet 5 a0f391f212 fix(llm): catch bare "cannot generate" phrasing without a leading "I"
The lexical refusal pattern required an "I" before cannot/can't/won't
so a refusal fragment like "Cannot generate this image due to the
content" (no leading pronoun — the shape it takes as a JSON string
value or a mid-sentence fragment) went undetected. Made the "I"
prefix optional; the cannot/can't/won't-plus-nearby-generate/create/
etc. shape is specific enough that dropping the pronoun requirement
doesn't pick up ordinary content (verified against the existing
"ordinary content" regression cases).

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
2026-07-29 17:34:41 +01:00
James VeitchandClaude Sonnet 5 3c5136af70 fix(llm): stop gating the embedding refusal check on is_ambiguous
An explicitly configured embedding_model is an opt-in to pay for the
similarity check, but is_refusal still gated it on is_ambiguous — a
length/hedge-keyword heuristic that skips anything long or plainly
worded. For chat_structured() this checks the full serialized JSON
output, whose braces/field-name overhead alone often pushed responses
past the 400-char cutoff, so the embedding call silently never fired
for exactly the subtle, on-topic-looking deflections this feature was
built to catch.

Now embed_fn, once supplied, always runs on any non-blank response that
didn't already match the lexical pass. Also widened is_ambiguous's own
threshold/keyword list, since it remains a standalone tested heuristic
for other callers even though is_refusal no longer wires it in.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
2026-07-29 17:30:19 +01:00
James VeitchandClaude Sonnet 5 97249ee281 fix(llm): catch "I am unable to" phrasing in refusal detection
The lexical pattern only matched the "I'm unable to" contraction, so
spelled-out refusals like "I am unable to comply with your instruction"
fell through undetected without an embedding model configured.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
2026-07-29 17:16:38 +01:00
James VeitchandClaude Sonnet 5 1f138e5991 feat(llm): detect refusal/deflection responses and retry with a new seed (#36)
Some models (observed with an abliterated Qwen variant) answer a request
they judge "politically sensitive" with hedging refusal language instead of
erroring or returning blank — neither of which the existing blank-response
or schema-validation retry triggers catch, so the refusal just passed
through as a normal result.

Adds a hybrid detector to _llm/retry.py: a free lexical/regex pass catches
blatant refusals ("I cannot generate...") without any network call; a
response that's short and/or hedge-y enough to be ambiguous additionally
gets an embedding-similarity check against canonical refusal exemplars, via
a new optional LLMProvider.embed() (Ollama's /api/embed, llama-server's
/v1/embeddings) — deliberately model-agnostic to the backend, since this is
a model-behavior concern, not a backend one. A detected refusal is treated
exactly like a blank response or validation failure: retried with a bumped
seed (existing next_seed()/RETRY_BACKOFF_SECS machinery), never returned to
the caller.

Wired into all four retry loops: OllamaProvider.chat()/chat_structured(),
LlamaCppProvider.chat(), and the shared _llm/chat.py chat_structured() used
by LlamaCppProvider.chat_structured(). New OllamaOptionRefusalRetry node
(same composable OLLAMA_OPTIONS chain as OllamaOptionDisableThinking)
configures it: enabled toggle, optional embedding_model (blank = lexical-
only), similarity threshold.


Claude-Session: https://claude.ai/code/session_01YArD9ZjBWKsAvazmS48amA

Co-authored-by: Claude Sonnet 5 <noreply@anthropic.com>
2026-07-29 17:11:08 +01:00
16 changed files with 2552 additions and 55 deletions
+3
View File
@@ -17,6 +17,7 @@ from .ollama import (
OllamaOptionDisableThinking,
OllamaOptionExtraBody,
OllamaOptionMaxTokens,
OllamaOptionRefusalRetry,
OllamaOptionRepeatPenalty,
OllamaOptionSeed,
OllamaOptionTemperature,
@@ -48,6 +49,7 @@ NODE_CLASS_MAPPINGS = {
"OllamaOptionTopK": OllamaOptionTopK,
"OllamaOptionRepeatPenalty": OllamaOptionRepeatPenalty,
"OllamaOptionDisableThinking": OllamaOptionDisableThinking,
"OllamaOptionRefusalRetry": OllamaOptionRefusalRetry,
"OllamaOptionExtraBody": OllamaOptionExtraBody,
"OllamaDebugHistory": OllamaDebugHistory,
"OllamaHistoryLength": OllamaHistoryLength,
@@ -75,6 +77,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"OllamaOptionTopK": "Ollama Option — Top K",
"OllamaOptionRepeatPenalty": "Ollama Option — Repeat Penalty",
"OllamaOptionDisableThinking": "Ollama Option — Disable Thinking",
"OllamaOptionRefusalRetry": "Ollama Option — Refusal Retry",
"OllamaOptionExtraBody": "Ollama Option — Extra Body",
"OllamaDebugHistory": "Ollama Debug History",
"OllamaHistoryLength": "Ollama History Length",
+97 -9
View File
@@ -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
@@ -46,7 +47,16 @@ from pydantic_ai.providers.openai import OpenAIProvider
from pydantic_ai.settings import ModelSettings
from .provider import Message
from .retry import RETRY_BACKOFF_SECS, next_seed
from .retry import (
RETRY_BACKOFF_SECS,
EmbedFn,
format_recovered_status,
format_retry_status,
is_refusal,
next_seed,
next_timeout_secs,
record_attempt_info,
)
_STRUCTURED_OUTPUT_FAILURE_EXCEPTIONS = (
UnexpectedModelBehavior,
@@ -126,6 +136,9 @@ async def chat_structured(
options: dict | None = None,
max_retries: int = 2,
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``.
@@ -155,19 +168,21 @@ async def chat_structured(
Retries up to ``max_retries`` times (clamped 0-5) on validation failure
before raising ``RuntimeError``. Never returns a value that failed
validation against ``schema``.
``options`` may also carry a ``"refusal_retry"`` config dict (same
comfydv-level convention as ``"think"``, emitted by
``OllamaOptionRefusalRetry``) — a detected refusal/deflection (see
``_llm/retry.py``) is treated exactly like a validation failure: retried
with a bumped seed rather than returned to the caller. ``embed_fn`` is
``LlamaCppProvider``'s own ``embed()``, bound to whatever embedding
model the config names — passed in rather than looked up here since
this module has no provider instance of its own to call.
"""
if not messages or messages[-1].role != "user":
raise ValueError(
"chat_structured requires the last message to have role='user'"
)
agent = _build_agent(
base_url=base_url,
model=model,
schema=schema,
headers=headers,
timeout_secs=timeout_secs,
)
history = _history_to_messages(messages)
prompt = _user_prompt_content(messages[-1])
think = None
@@ -175,6 +190,11 @@ async def chat_structured(
options = dict(options)
think = options.pop("think")
options = options or None
refusal_cfg = None
if options and "refusal_retry" in options:
options = dict(options)
refusal_cfg = options.pop("refusal_retry")
options = options or None
extra_body: dict = {}
if options:
extra_body["options"] = options
@@ -189,7 +209,33 @@ async def chat_structured(
total_attempts = max(0, min(int(max_retries), 5)) + 1
last_error: Exception | None = None
last_invalid_text = ""
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
# effect — httpx.AsyncClient's timeout is fixed at construction,
# not mutable per-request.
agent = _build_agent(
base_url=base_url,
model=model,
schema=schema,
headers=headers,
timeout_secs=attempt_timeout,
)
attempt_settings = dict(model_settings) if model_settings else {}
if attempt > 1:
# Confirmed live: a freshly-loaded model's first structured-output
@@ -202,6 +248,7 @@ async def chat_structured(
# llama-server's OpenAI-compatible endpoints) and give it a beat
# via RETRY_BACKOFF_SECS in case it's still finishing loading.
seed = next_seed(options, attempt)
attempt_seed = seed
attempt_settings["seed"] = seed
if "extra_body" in attempt_settings:
# beacon-reviewer caught this: if a caller pinned options["seed"],
@@ -233,13 +280,54 @@ async def chat_structured(
# not a static type parameter), so the checker can't narrow
# result.output past Agent's default `str` — cast to the
# function's declared return type, which schema is a subtype of.
return cast(BaseModel, result.output)
output = cast(BaseModel, result.output)
if refusal_cfg and refusal_cfg.get("enabled"):
# Re-serialized, not the original wire text — pydantic-ai's
# NativeOutput doesn't expose that separately, and the
# regex/embedding check works the same either way (same
# textual content, just re-encoded).
content = output.model_dump_json()
refused = await is_refusal(
content,
embed_fn=embed_fn,
embed_cache_key=refusal_cfg.get("embedding_model", ""),
threshold=refusal_cfg.get("threshold", 0.82),
custom_phrases=tuple(refusal_cfg.get("custom_phrases") or ()),
)
if refused:
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
record_attempt_info(
attempt_info,
seed=attempt_seed,
attempts=attempt,
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)
record_attempt_info(
attempt_info,
seed=attempt_seed,
attempts=total_attempts,
timeout_secs=attempt_timeout,
refusals=refusal_count,
)
raise RuntimeError(
f"chat_structured: response failed validation against schema after "
f"{total_attempts} attempt(s) (model={model!r}). Last error: "
+149 -5
View File
@@ -15,12 +15,28 @@ exist otherwise (spec.md FR-006).
import asyncio
import logging
from collections.abc import Callable
from pydantic import BaseModel
from .ollama_provider import _TTLLRUCache, _cache_key, _get_json, _pop_think, _post_json
from .ollama_provider import (
_TTLLRUCache,
_cache_key,
_get_json,
_pop_refusal_retry,
_pop_think,
_post_json,
)
from .provider import Message, ModelInfo, ModelStatus
from .retry import RETRY_BACKOFF_SECS, next_seed
from .retry import (
RETRY_BACKOFF_SECS,
format_recovered_status,
format_retry_status,
is_refusal,
next_seed,
next_timeout_secs,
record_attempt_info,
)
logger = logging.getLogger(__name__)
@@ -183,13 +199,27 @@ class LlamaCppProvider:
options: dict | None = None,
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)
options, refusal_cfg = _pop_refusal_retry(options)
embed_fn = None
custom_phrases: tuple[str, ...] = ()
if refusal_cfg and refusal_cfg.get("enabled"):
custom_phrases = tuple(refusal_cfg.get("custom_phrases") or ())
if refusal_cfg.get("embedding_model"):
embedding_model = refusal_cfg["embedding_model"]
embed_fn = lambda t: self.embed(embedding_model, t) # noqa: E731
total_attempts = max(0, min(int(max_retries), 5)) + 1
response_text = ""
refusal_count = 0
attempt_seed = 0
attempt_timeout = timeout_secs
for attempt in range(1, total_attempts + 1):
attempt_timeout = next_timeout_secs(timeout_secs, attempt)
payload: dict = {
"model": model,
"messages": payload_messages,
@@ -218,6 +248,13 @@ class LlamaCppProvider:
# spec's actual top-level "seed" field, so it takes effect
# against llama-server's /v1/chat/completions.
payload["seed"] = next_seed(options, attempt)
attempt_seed = payload["seed"]
else:
# attempt 1 never sets the top-level "seed" field above (only
# retries do) — fall back to whatever the caller pinned in
# options, so attempt_info/seed_used reports the real seed in
# play even on a first-attempt success, not a stale 0.
attempt_seed = (options or {}).get("seed", 0)
cache_key = _cache_key(
"llamacpp_chat",
@@ -231,12 +268,19 @@ class LlamaCppProvider:
)
cached, hit = _CHAT_RESPONSE_CACHE.get(cache_key)
if hit:
record_attempt_info(
attempt_info,
seed=attempt_seed,
attempts=attempt,
timeout_secs=attempt_timeout,
refusals=refusal_count,
)
return cached
result = await _post_json(
f"{self.host}/v1/chat/completions",
payload,
timeout=timeout_secs,
timeout=attempt_timeout,
headers=self.headers,
)
choices = result.get("choices") or []
@@ -245,13 +289,60 @@ class LlamaCppProvider:
if choices
else ""
)
retry_reason: str | None = None
if response_text.strip():
_CHAT_RESPONSE_CACHE.set(cache_key, response_text)
return response_text
refused = False
if refusal_cfg and refusal_cfg.get("enabled"):
refused = await is_refusal(
response_text,
embed_fn=embed_fn,
embed_cache_key=refusal_cfg.get("embedding_model", ""),
threshold=refusal_cfg.get("threshold", 0.82),
custom_phrases=custom_phrases,
)
if not refused:
_CHAT_RESPONSE_CACHE.set(cache_key, response_text)
record_attempt_info(
attempt_info,
seed=attempt_seed,
attempts=attempt,
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(
attempt_info,
seed=attempt_seed,
attempts=total_attempts,
timeout_secs=attempt_timeout,
refusals=refusal_count,
)
# Every attempt came back blank — never raises here (chat() has
# never validated its output, unlike chat_structured()); return the
# last (blank) attempt uncached so the next queue run tries fresh.
@@ -265,6 +356,8 @@ class LlamaCppProvider:
options: dict | None = None,
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
@@ -280,8 +373,25 @@ class LlamaCppProvider:
)
cached, hit = _CHAT_RESPONSE_CACHE.get(cache_key)
if hit:
record_attempt_info(
attempt_info,
seed=(options or {}).get("seed", 0),
attempts=1,
timeout_secs=timeout_secs,
refusals=0,
)
return schema.model_validate(cached)
embed_fn = None
refusal_cfg = (options or {}).get("refusal_retry")
if (
refusal_cfg
and refusal_cfg.get("enabled")
and refusal_cfg.get("embedding_model")
):
embedding_model = refusal_cfg["embedding_model"]
embed_fn = lambda t: self.embed(embedding_model, t) # noqa: E731
result = await _chat_structured_impl(
base_url=f"{self.host}/v1",
model=model,
@@ -291,6 +401,40 @@ class LlamaCppProvider:
options=options,
max_retries=max_retries,
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
async def embed(self, model: str, text: str) -> list[float] | None:
"""POST {host}/v1/embeddings — llama-server's OpenAI-compatible
embeddings endpoint.
Requires an embedding-capable model to be loaded in the router
(typically a *different* model from whatever's answering chat
requests) — not live-verified against a running llama-server (no
instance available at implementation time), mirroring this
provider's other sourced-from-docs-not-verified caveats. Returns
``None`` rather than raising on any failure, same contract as
``OllamaProvider.embed()``.
"""
if not model.strip() or not text.strip():
return None
try:
result = await _post_json(
f"{self.host}/v1/embeddings",
{"model": model, "input": text},
timeout=30.0,
headers=self.headers,
)
except Exception:
return None
data = result.get("data")
if not isinstance(data, list) or not data:
return None
vec = data[0].get("embedding")
if not isinstance(vec, list) or not vec:
return None
return vec
+209 -5
View File
@@ -15,11 +15,20 @@ 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, next_seed
from .retry import (
RETRY_BACKOFF_SECS,
format_recovered_status,
format_retry_status,
is_refusal,
next_seed,
next_timeout_secs,
record_attempt_info,
)
logger = logging.getLogger(__name__)
@@ -102,6 +111,22 @@ def _pop_think(options: dict | None) -> tuple[dict | None, bool | None]:
return (remaining or None), think
def _pop_refusal_retry(options: dict | None) -> tuple[dict | None, dict | None]:
"""Split a ``"refusal_retry"`` config dict out of a generic ``options``
dict — same convention as ``_pop_think``: ``OllamaOptionRefusalRetry``
merges ``{"refusal_retry": {"enabled", "embedding_model", "threshold"}}``
into the same composable ``OLLAMA_OPTIONS`` chain every other
``OllamaOption*`` node feeds into ``ChatCompletion``'s ``options``
input, and neither Ollama's nor llama.cpp's own API recognizes this key,
so every provider pops it out here before building its request.
"""
if not options or "refusal_retry" not in options:
return options, None
remaining = dict(options)
cfg = remaining.pop("refusal_retry")
return (remaining or None), cfg
_MODEL_LIST_CACHE = _TTLLRUCache(maxsize=32, ttl_seconds=20.0)
_CHAT_RESPONSE_CACHE = _TTLLRUCache(maxsize=64, ttl_seconds=None)
_CAPABILITY_CACHE = _TTLLRUCache(maxsize=32, ttl_seconds=300.0)
@@ -343,6 +368,8 @@ class OllamaProvider:
options: dict | None = None,
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)
@@ -353,14 +380,27 @@ class OllamaProvider:
# array (ADR-008 — no transform needed for /api/chat).
payload_messages = [m.model_dump(exclude_none=True) for m in messages]
options, think = _pop_think(options)
options, refusal_cfg = _pop_refusal_retry(options)
embed_fn = None
custom_phrases: tuple[str, ...] = ()
if refusal_cfg and refusal_cfg.get("enabled"):
custom_phrases = tuple(refusal_cfg.get("custom_phrases") or ())
if refusal_cfg.get("embedding_model"):
embedding_model = refusal_cfg["embedding_model"]
embed_fn = lambda t: self.embed(embedding_model, t) # noqa: E731
total_attempts = max(0, min(int(max_retries), 5)) + 1
response_text = ""
incomplete = False
refusal_count = 0
attempt_seed = 0
attempt_timeout = timeout_secs
for attempt in range(1, total_attempts + 1):
attempt_options = dict(options) if options else {}
if attempt > 1:
attempt_options["seed"] = next_seed(options, attempt)
attempt_seed = attempt_options.get("seed", 0)
attempt_timeout = next_timeout_secs(timeout_secs, attempt)
payload: dict = {
"model": model,
@@ -385,18 +425,55 @@ class OllamaProvider:
)
cached, hit = _CHAT_RESPONSE_CACHE.get(cache_key)
if hit:
record_attempt_info(
attempt_info,
seed=attempt_seed,
attempts=attempt,
timeout_secs=attempt_timeout,
refusals=refusal_count,
)
return cached
result = await _post_json(
f"{self.host}/api/chat",
payload,
timeout=timeout_secs,
timeout=attempt_timeout,
headers=self.headers,
)
response_text = result.get("message", {}).get("content", "")
retry_reason: str | None = None
if response_text.strip():
_CHAT_RESPONSE_CACHE.set(cache_key, response_text)
return response_text
refused = False
if refusal_cfg and refusal_cfg.get("enabled"):
refused = await is_refusal(
response_text,
embed_fn=embed_fn,
embed_cache_key=refusal_cfg.get("embedding_model", ""),
threshold=refusal_cfg.get("threshold", 0.82),
custom_phrases=custom_phrases,
)
if not refused:
_CHAT_RESPONSE_CACHE.set(cache_key, response_text)
record_attempt_info(
attempt_info,
seed=attempt_seed,
attempts=attempt,
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
@@ -406,10 +483,34 @@ 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(
attempt_info,
seed=attempt_seed,
attempts=total_attempts,
timeout_secs=attempt_timeout,
refusals=refusal_count,
)
if incomplete:
raise RuntimeError(
f"Ollama returned an incomplete response after "
@@ -432,6 +533,8 @@ class OllamaProvider:
options: dict | None = None,
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.
@@ -457,6 +560,14 @@ class OllamaProvider:
payload_messages = [m.model_dump(exclude_none=True) for m in messages]
json_schema = schema.model_json_schema()
options, think = _pop_think(options)
options, refusal_cfg = _pop_refusal_retry(options)
embed_fn = None
custom_phrases: tuple[str, ...] = ()
if refusal_cfg and refusal_cfg.get("enabled"):
custom_phrases = tuple(refusal_cfg.get("custom_phrases") or ())
if refusal_cfg.get("embedding_model"):
embedding_model = refusal_cfg["embedding_model"]
embed_fn = lambda t: self.embed(embedding_model, t) # noqa: E731
cache_key = _cache_key(
"chat_structured",
self.host,
@@ -469,16 +580,39 @@ class OllamaProvider:
)
cached, hit = _CHAT_RESPONSE_CACHE.get(cache_key)
if hit:
record_attempt_info(
attempt_info,
seed=(options or {}).get("seed", 0),
attempts=1,
timeout_secs=timeout_secs,
refusals=0,
)
return schema.model_validate(cached)
total_attempts = max(0, min(int(max_retries), 5)) + 1
last_error: Exception | None = None
last_invalid_text = ""
refusal_count = 0
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:
attempt_options["seed"] = next_seed(options, attempt)
attempt_seed = attempt_options.get("seed", 0)
attempt_timeout = next_timeout_secs(timeout_secs, attempt)
payload: dict = {
"model": model,
@@ -497,12 +631,13 @@ class OllamaProvider:
result = await _post_json(
f"{self.host}/api/chat",
payload,
timeout=timeout_secs,
timeout=attempt_timeout,
headers=self.headers,
)
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
@@ -513,15 +648,84 @@ 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
if refusal_cfg and refusal_cfg.get("enabled"):
# Checked against the raw JSON text, not a specific parsed
# field: ChatCompletion's schema is caller-defined and this
# provider has no idea which field would carry refusal
# language — the regex/embedding check still matches text
# sitting inside a JSON string value either way.
refused = await is_refusal(
content,
embed_fn=embed_fn,
embed_cache_key=refusal_cfg.get("embedding_model", ""),
threshold=refusal_cfg.get("threshold", 0.82),
custom_phrases=custom_phrases,
)
if refused:
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
_CHAT_RESPONSE_CACHE.set(cache_key, parsed.model_dump())
record_attempt_info(
attempt_info,
seed=attempt_seed,
attempts=attempt,
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(
attempt_info,
seed=attempt_seed,
attempts=total_attempts,
timeout_secs=attempt_timeout,
refusals=refusal_count,
)
raise RuntimeError(
f"chat_structured: response failed validation against schema after "
f"{total_attempts} attempt(s) (model={model!r}). Last error: "
f"{last_error}. Last response (truncated): {last_invalid_text[:300]!r}"
)
async def embed(self, model: str, text: str) -> list[float] | None:
"""POST {host}/api/embed — Ollama's native embeddings endpoint.
Returns ``None`` rather than raising on any failure (wrong/missing
embedding model, unreachable server, malformed response) — this is
a best-effort capability per the ``LLMProvider`` protocol, and its
one current caller (refusal-retry detection) already treats
``None`` as "skip the embedding check", not an error.
"""
if not model.strip() or not text.strip():
return None
try:
result = await _post_json(
f"{self.host}/api/embed",
{"model": model, "input": text},
timeout=30.0,
headers=self.headers,
)
except Exception:
return None
embeddings = result.get("embeddings")
if not isinstance(embeddings, list) or not embeddings:
return None
vec = embeddings[0]
if not isinstance(vec, list) or not vec:
return None
return vec
+37 -2
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.
"""
from collections.abc import Callable
from enum import Enum
from typing import Literal, Protocol
@@ -83,6 +84,8 @@ class LLMProvider(Protocol):
options: dict | None = None,
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.
@@ -92,7 +95,9 @@ class LLMProvider(Protocol):
into normal behavior. Still returns the (possibly blank) last
attempt's text rather than raising if every retry comes back blank —
this method has never validated its output, unlike
``chat_structured()``.
``chat_structured()``. Each retry's request timeout also escalates
(``_llm/retry.py``'s ``next_timeout_secs``) rather than reusing the
same budget that just ran out.
ADR-010: ``options`` may carry a ``"think"`` key (bool) to disable a
"thinking"-capable model's chain-of-thought reasoning — every
@@ -101,6 +106,18 @@ class LLMProvider(Protocol):
``chat_template_kwargs``/``reasoning_effort`` in the request body),
since neither backend recognizes a literal ``"think"`` key nested
inside a generic options object.
``attempt_info``, if given, is populated in place with the retry
loop's final outcome (seed/timeout used, attempt count, refusal
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.
"""
...
@@ -112,6 +129,8 @@ class LLMProvider(Protocol):
options: dict | None = None,
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.
@@ -121,6 +140,22 @@ class LLMProvider(Protocol):
field.
ADR-010: see ``chat()`` — same ``options["think"]`` convention,
same per-provider translation.
same per-provider translation, same escalating per-attempt timeout,
and the same ``attempt_info``/``on_status`` conventions.
"""
...
async def embed(self, model: str, text: str) -> list[float] | None:
"""Embedding vector for ``text``, or ``None`` if unavailable.
Best-effort, not a core capability every deployment has configured:
``model`` must itself be embedding-capable, which is typically a
*different* model from whatever's answering chat requests (e.g.
``nomic-embed-text``, not the model passed to ``chat()``). Returns
``None`` rather than raising when embeddings aren't usable right now
(wrong/missing model, unreachable server) — the one current caller,
refusal-retry detection (see ``_llm/retry.py``), degrades gracefully
to lexical-only detection when this returns ``None``, so a provider
with no embedding model configured is never a hard failure.
"""
...
+299
View File
@@ -11,8 +11,21 @@ each re-deriving the policy.
Only blank/whitespace-only responses trigger a retry for plain chat() —
not merely "short" ones — because a fixed length threshold would misfire on
legitimately short, valid answers (single-word replies, labels, "yes"/"no").
Refusal/deflection detection (below) is a separate, opt-in trigger for the
same retry-with-a-new-seed mechanism: some models (observed with an
abliterated Qwen variant) answer with a soft refusal on a topic they judge
"sensitive" instead of erroring or returning blank, so neither of the above
checks catches it. This is deliberately a model-behavior concern, not a
backend one — every ``LLMProvider`` implementation (Ollama, llama.cpp, and
whatever comes next) wires the same detector into its own retry loop via its
own ``embed()``, rather than each backend inventing its own heuristic.
"""
import math
import re
from collections.abc import Awaitable, Callable
RETRY_BACKOFF_SECS = 1.5
"""Flat delay between retries — gives a still-loading model time to finish
before the next attempt, rather than hammering it with identical requests
@@ -32,3 +45,289 @@ def next_seed(options: dict | None, attempt: int) -> int:
if options and isinstance(options.get("seed"), int):
base = options["seed"]
return base + (attempt - 1)
def next_timeout_secs(base_timeout: float, attempt: int) -> float:
"""Escalating per-attempt timeout for retries (1-indexed ``attempt``).
Attempt 1 gets the caller's own ``timeout_secs`` unchanged; each retry
multiplies it by the attempt number. A request that timed out may
genuinely need more time — a slow-to-load or heavily-loaded model, a
large prompt — not just an identical retry under the same budget it
just failed to meet.
"""
return base_timeout * attempt
def record_attempt_info(
attempt_info: dict | None,
*,
seed: int,
attempts: int,
timeout_secs: float,
refusals: int,
) -> None:
"""Populate an optional caller-supplied dict with the retry loop's
final outcome — the seed/timeout actually used, how many attempts it
took, and how many were refusal-triggered.
A plain out-param rather than a return-type change, so it's fully
backward compatible: a caller that doesn't pass ``attempt_info`` sees
no change in behavior at all. ``ChatCompletion`` uses this to expose
the seed actually used as a node output and to build a UI status line
when a retry/refusal happened.
"""
if attempt_info is None:
return
attempt_info.update(
{
"seed": seed,
"attempts": attempts,
"timeout_secs": timeout_secs,
"refusals": refusals,
}
)
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
# ---------------------------------------------------------------------------
#
# Hybrid, cheapest-check-first: a fast, free lexical pass catches the blatant
# majority ("I cannot generate...") without ever touching the network; only
# a response that's short and/or hedge-y enough to be genuinely ambiguous
# pays for an embedding call. A long, on-topic response never reaches the
# embedding step at all.
REFUSAL_LEXICAL_PATTERNS: tuple[re.Pattern, ...] = tuple(
re.compile(p, re.IGNORECASE)
for p in (
r"\b(?:I\s*(?:'m|\s+am)?\s*)?(?:cannot|can't|won't|will not)\b[^.]{0,60}?\b"
r"(?:generate|create|produce|write|provide|help|assist|describe|depict|continue)\b",
r"\bI(?:'m|\s+am) (?:(?:not able|unable) to|restricted from)\b",
r"\bI don't feel comfortable\b",
r"\bI'm sorry,?\s*(?:but\s+)?I\s*(?:can't|cannot)\b",
r"\bas an AI\b[^.]{0,60}?\b(?:cannot|can't|unable|not able)\b",
r"\bI must (?:decline|refuse)\b",
r"\bthis (?:request|content|image|prompt) (?:violates|goes against)\b",
r"\bnot something I(?:'m| am)? (?:able to|going to) (?:help with|generate|create|provide)\b",
r"\bI (?:won't|will not) be (?:able to )?(?:generat|creat|provid)",
)
)
_SOFT_HEDGE_KEYWORDS: tuple[str, ...] = (
"sensitive",
"unfortunately",
"i must",
"policy",
"guideline",
"responsible ai",
"not appropriate",
"inappropriate",
"instead, i",
"i'd rather",
"i would rather",
"controversial",
"harmful",
"offensive",
"disturbing",
"explicit content",
"nsfw",
"consider an alternative",
)
_AMBIGUOUS_LENGTH_THRESHOLD = 600
"""Below this many characters, a response is short enough that a soft
refusal is plausible and worth the extra embedding check — chosen well
under this pipeline's normal structured-JSON response sizes (typically
1000+ characters), not a tuned/validated threshold.
Note: ``is_refusal`` no longer applies this gate when the caller supplies
``embed_fn`` (see below) — an explicitly configured embedding model always
runs the check regardless of length/keywords. ``is_ambiguous`` is kept as a
standalone, independently tested heuristic for other callers (e.g. logging
"why was this worth embedding-checking") rather than wired into the hybrid
detector itself."""
REFUSAL_EXEMPLARS: tuple[str, ...] = (
"I cannot generate an image or description involving this topic.",
"I'm not able to help with this request as it involves a politically sensitive subject.",
"I don't feel comfortable creating content related to this figure or event.",
"This request touches on a sensitive political topic that I must decline to engage with.",
"As an AI, I'm unable to produce content depicting this scenario.",
"I'd rather not generate this — it involves content I'm not comfortable with.",
)
def is_lexical_refusal(text: str) -> bool:
"""Cheap, free regex pass — catches the blatant majority of refusals."""
return any(p.search(text) for p in REFUSAL_LEXICAL_PATTERNS)
def is_ambiguous(text: str) -> bool:
"""Whether ``text`` is short/hedge-y enough to be worth the pricier
embedding check, having already failed the free lexical pass.
Deliberately cheap and approximate — false positives here only cost one
extra embedding call, false negatives skip a refusal that a real
similarity check might have caught. Not meant to be a precise signal on
its own, just a gate on when the more expensive check runs at all.
"""
stripped = text.strip()
if not stripped:
return False
if len(stripped) < _AMBIGUOUS_LENGTH_THRESHOLD:
return True
lowered = stripped.lower()
return any(keyword in lowered for keyword in _SOFT_HEDGE_KEYWORDS)
def cosine_similarity(a: list[float], b: list[float]) -> float:
"""Standard cosine similarity, no numpy dependency (comfydv has none)."""
if not a or not b or len(a) != len(b):
return 0.0
dot = sum(x * y for x, y in zip(a, b))
norm_a = math.sqrt(sum(x * x for x in a))
norm_b = math.sqrt(sum(y * y for y in b))
if norm_a == 0.0 or norm_b == 0.0:
return 0.0
return dot / (norm_a * norm_b)
EmbedFn = Callable[[str], Awaitable[list[float] | None]]
_exemplar_embedding_cache: dict[str, list[list[float]]] = {}
async def _exemplar_embeddings(
embed_fn: EmbedFn, cache_key: str, exemplars: tuple[str, ...] = REFUSAL_EXEMPLARS
) -> list[list[float]]:
"""Embed ``exemplars`` once per ``cache_key`` and reuse — the exemplar
set only changes if the caller's custom phrases change (folded into
``cache_key`` by the caller), or the embedding space (i.e. which model
produced the vectors) does.
"""
cached = _exemplar_embedding_cache.get(cache_key)
if cached is not None:
return cached
embeddings = []
for exemplar in exemplars:
vec = await embed_fn(exemplar)
if not vec:
# An embedding call failing for one exemplar almost certainly
# means embeddings aren't usable at all right now (wrong/missing
# embedding model, unreachable server) — bail out rather than
# caching a partial, unusable exemplar set.
return []
embeddings.append(vec)
_exemplar_embedding_cache[cache_key] = embeddings
return embeddings
def _matches_custom_phrase(text: str, custom_phrases: tuple[str, ...]) -> bool:
"""Case-insensitive substring match against user-supplied phrases."""
if not custom_phrases:
return False
lowered = text.lower()
return any(phrase.lower() in lowered for phrase in custom_phrases)
async def is_refusal(
text: str,
*,
embed_fn: EmbedFn | None = None,
embed_cache_key: str = "",
threshold: float = 0.82,
custom_phrases: tuple[str, ...] = (),
) -> bool:
"""Hybrid refusal/deflection detector: free lexical pass first, then an
embedding-similarity fallback whenever the caller has configured one.
``embed_fn`` is supplied by the caller's own ``LLMProvider.embed()`` —
this function has no idea which backend or model produced ``text``, by
design (ADR: refusal detection is a model-behavior concern, not a
backend one). ``embed_fn=None`` (no embedding model configured) degrades
to lexical-only detection rather than erroring; ``embed_fn`` present
means the caller already opted in to the extra cost, so every non-blank,
non-lexically-caught response gets checked — no further length/keyword
gating. Any failure while embedding (unreachable server, no
embedding-capable model loaded) is swallowed the same way — an optional
enhancement failing shouldn't take down the retry loop it's assisting.
``custom_phrases`` lets a caller extend detection at runtime — e.g. a
ComfyUI node field the user edits directly — without touching the
shipped patterns/exemplars. Each phrase is checked two ways: a free
case-insensitive substring match (same cost tier as the lexical pass,
so it runs even with no ``embed_fn`` configured), and, when ``embed_fn``
is present, folded in as additional exemplars for the similarity check
so near-matches (not just exact substrings) of the user's phrases count
too.
"""
if not text or not text.strip():
return False # blank responses are the *other* retry trigger, not this one
if is_lexical_refusal(text):
return True
custom_phrases = tuple(p.strip() for p in custom_phrases if p and p.strip())
if _matches_custom_phrase(text, custom_phrases):
return True
if embed_fn is None:
return False
# embed_fn only exists when the caller explicitly configured an
# embedding_model — that's an opt-in to pay for the check, so run it on
# every non-blank, non-lexically-caught response rather than gating
# further on is_ambiguous. The length/keyword heuristic exists to avoid
# *unwanted* embedding calls when no embedding model is configured (see
# the embed_fn is None branch above); it has no reason to also suppress
# calls once the caller has already asked for them, and doing so was
# exactly what let the subtle/on-topic-looking deflections this feature
# targets slip through undetected.
exemplars = (
REFUSAL_EXEMPLARS + custom_phrases if custom_phrases else REFUSAL_EXEMPLARS
)
exemplar_cache_key = (
f"{embed_cache_key}|custom:{','.join(custom_phrases)}"
if custom_phrases
else embed_cache_key
)
try:
exemplar_vecs = await _exemplar_embeddings(
embed_fn, exemplar_cache_key, exemplars
)
if not exemplar_vecs:
return False
text_vec = await embed_fn(text)
if not text_vec:
return False
except Exception:
return False
return max(cosine_similarity(text_vec, vec) for vec in exemplar_vecs) >= threshold
+184 -6
View File
@@ -509,8 +509,20 @@ def _coerce_structured_value(value, comfy_type: str):
class ChatCompletion:
OUTPUT_NODE = True
_BASE_RETURN_TYPES = ("STRING", "OLLAMA_HISTORY", "STRING")
_BASE_RETURN_NAMES = ("response", "updated_history", "model_name")
# "seed_used" is always the LAST output, after any dynamic structured-
# output fields — not inserted right after model_name — so a schema's
# fields keep starting at the same fixed index (3) they occupied before
# this output existed. Already-wired workflows (e.g.
# workflows/ltx-i2v-pipeline.json) link FormatString inputs to a
# ChatCompletion node's dynamic field by output *index*; inserting a
# new fixed output ahead of those fields would silently repoint every
# such link at the wrong socket.
_FIXED_RETURN_TYPES = ("STRING", "OLLAMA_HISTORY", "STRING")
_FIXED_RETURN_NAMES = ("response", "updated_history", "model_name")
_SEED_RETURN_TYPE = ("INT",)
_SEED_RETURN_NAME = ("seed_used",)
_BASE_RETURN_TYPES = _FIXED_RETURN_TYPES + _SEED_RETURN_TYPE
_BASE_RETURN_NAMES = _FIXED_RETURN_NAMES + _SEED_RETURN_NAME
# Per-node-instance structured-output config, keyed by unique_id — same
# pattern as FormatString.node_configs.
@@ -579,8 +591,14 @@ class ChatCompletion:
cls.RETURN_NAMES = cls._BASE_RETURN_NAMES
return
names = tuple(schema["properties"].keys())
cls.RETURN_TYPES = cls._BASE_RETURN_TYPES + _comfy_types_for_schema(schema)
cls.RETURN_NAMES = cls._BASE_RETURN_NAMES + names
# Dynamic fields land between the fixed base outputs and the
# trailing seed_used — see _SEED_RETURN_TYPE's comment above.
cls.RETURN_TYPES = (
cls._FIXED_RETURN_TYPES
+ _comfy_types_for_schema(schema)
+ cls._SEED_RETURN_TYPE
)
cls.RETURN_NAMES = cls._FIXED_RETURN_NAMES + names + cls._SEED_RETURN_NAME
def chat(
self,
@@ -624,6 +642,33 @@ class ChatCompletion:
if user_images:
messages[-1].images = user_images
llm_options = dict(options) if options else None
# Populated in place by the provider (see _llm/retry.py's
# record_attempt_info) with the seed/timeout actually used and how
# many attempts/refusals it took — an out-param rather than a
# return-type change, so it works the same regardless of which
# 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).
@@ -637,6 +682,8 @@ class ChatCompletion:
llm_options,
timeout_secs=float(timeout_secs),
max_retries=max_retries,
attempt_info=attempt_info,
on_status=on_status,
)
)
else:
@@ -651,20 +698,44 @@ class ChatCompletion:
llm_options,
timeout_secs=float(timeout_secs),
max_retries=max_retries,
attempt_info=attempt_info,
on_status=on_status,
)
)
response_text = parsed.model_dump_json()
seed_used = attempt_info.get("seed", 0)
attempts_made = attempt_info.get("attempts", 1)
refusals = attempt_info.get("refusals", 0)
updated = list(history)
updated.append({"role": "user", "content": prompt})
updated.append({"role": "assistant", "content": response_text})
n = len(updated)
# Visual feedback (US: "show me when a retry/refusal happened") —
# surfaced in the node's existing text preview rather than a new UI
# surface, so it's visible without any frontend/JS changes.
status_line = ""
if attempts_made > 1:
status_line = (
f"⚠️ Refusal/deflection detected — retried {refusals} "
f"time(s), succeeded on attempt {attempts_made} with "
f"seed={seed_used}.\n\n"
if refusals
else f"⚠️ Retried (blank/invalid response) — succeeded on "
f"attempt {attempts_made} with seed={seed_used}.\n\n"
)
ui_text = (
f"{response_text}\n\n── History: {n} message(s) ──\n{_history_preview(updated)}"
f"{status_line}{response_text}\n\n"
f"── History: {n} message(s) ──\n{_history_preview(updated)}"
if n > 2
else response_text
else f"{status_line}{response_text}"
)
# seed_used is appended last, after any dynamic structured fields —
# see ChatCompletion's _SEED_RETURN_TYPE comment for why.
result_tuple = (response_text, updated, effective_model)
if structured_output:
assert schema is not None # structured_output implies this was parsed
@@ -674,6 +745,7 @@ class ChatCompletion:
for name, ctype in zip(schema["properties"].keys(), comfy_types)
)
result_tuple += extra
result_tuple += (seed_used,)
return {
"ui": {"text": [ui_text]},
@@ -914,6 +986,112 @@ class OllamaOptionDisableThinking:
return (_merge_option(options, "think", not disable_thinking),)
class OllamaOptionRefusalRetry:
"""Retry with a bumped seed when a response looks like a soft
refusal/deflection rather than an actual error (see
``_llm/retry.py``'s ``is_refusal()``).
Observed with an abliterated Qwen variant: it sometimes answers a
request it judges "politically sensitive" with hedging refusal
language instead of erroring or returning blank — neither of which the
existing blank-response or schema-validation retry triggers catch, so
without this the response just passes through as-is.
Rides the same composable ``OLLAMA_OPTIONS`` chain as every other
``OllamaOption*`` node, but like ``OllamaOptionDisableThinking``, the
``"refusal_retry"`` key this node emits is a comfydv-level convention,
not an Ollama-native sampling param: every ``LLMProvider``
implementation pops it out of ``options`` and drives its own retry
loop with it (ADR: refusal detection is a model-behavior concern, not
a backend one — see ``LLMProvider.embed()`` in ``_llm/provider.py``).
Works for both backends from the same node.
``custom_phrases`` (comma-separated) lets you add your own trigger
phrases at runtime, without a code change/release — useful for a new
deflection phrasing a specific model uses that the shipped patterns in
``REFUSAL_LEXICAL_PATTERNS``/``REFUSAL_EXEMPLARS`` don't cover yet.
"""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"enabled": ("BOOLEAN", {"default": True}),
"embedding_model": (
"STRING",
{
"default": "",
"tooltip": (
"Name of a separate embedding-capable model "
'(e.g. "nomic-embed-text") — NOT the chat '
"model itself; most chat models can't produce "
"usable embeddings. Leave blank to skip the "
"embedding-similarity check and detect only "
"blatant, literal refusal phrases (still "
"useful, cheaper, catches less)."
),
},
),
"threshold": (
"FLOAT",
{
"default": 0.82,
"min": 0.0,
"max": 1.0,
"step": 0.01,
"tooltip": (
"Cosine-similarity threshold against canonical "
"refusal exemplars, above which an ambiguous "
"(short/hedge-y) response is treated as a "
"refusal. Only consulted when embedding_model "
"is set and the cheap lexical pass didn't "
"already catch it."
),
},
),
"custom_phrases": (
"STRING",
{
"default": "",
"tooltip": (
"Comma-separated phrases you want treated as "
'refusals too, e.g. "I am restricted from, as an '
'AI model, I must avoid". Checked as free, exact '
"case-insensitive substrings (no embedding model "
"needed) and, when embedding_model is set, also "
"folded in as extra exemplars for the similarity "
"check — lets you extend detection at runtime "
"without waiting on a shipped pattern update."
),
},
),
},
"optional": {"options": ("OLLAMA_OPTIONS",)},
}
RETURN_TYPES = ("OLLAMA_OPTIONS",)
RETURN_NAMES = ("options",)
FUNCTION = "set_refusal_retry"
CATEGORY = "dv/ollama/options"
def set_refusal_retry(
self, enabled, embedding_model, threshold, custom_phrases="", options=None
):
phrases = tuple(p.strip() for p in custom_phrases.split(",") if p and p.strip())
return (
_merge_option(
options,
"refusal_retry",
{
"enabled": enabled,
"embedding_model": embedding_model.strip(),
"threshold": threshold,
"custom_phrases": phrases,
},
),
)
class OllamaOptionExtraBody:
@classmethod
def INPUT_TYPES(s):
+36 -5
View File
@@ -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
+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);
};
},
});
+17 -1
View File
@@ -40,9 +40,25 @@ class _FakeProvider:
self.calls.append(("unload_model", model))
async def chat(
self, model, messages, options=None, timeout_secs=300.0, max_retries=2
self,
model,
messages,
options=None,
timeout_secs=300.0,
max_retries=2,
attempt_info=None,
on_status=None,
):
self.calls.append(("chat", model))
if attempt_info is not None:
attempt_info.update(
{
"seed": (options or {}).get("seed", 0),
"attempts": 1,
"timeout_secs": timeout_secs,
"refusals": 0,
}
)
return self.chat_response
+217
View File
@@ -405,6 +405,81 @@ def test_chat_exhausted_retries_returns_blank_without_raising(monkeypatch):
assert calls["n"] == 3 # original + 2 retries, per max_retries=2
def test_chat_timeout_escalates_per_retry_attempt(monkeypatch):
timeouts = []
async def fake_post(url, payload, *, timeout=120.0, headers=None):
timeouts.append(timeout)
if len(timeouts) < 3:
return {"choices": [{"message": {"content": ""}}]}
return {"choices": [{"message": {"content": "done"}}]}
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
monkeypatch.setattr(provider_mod.asyncio, "sleep", _fake_sleep)
_run_async(
LlamaCppProvider("http://localhost:8080").chat(
"gemma-3-4b",
[Message(role="user", content="hi")],
timeout_secs=50.0,
max_retries=2,
)
)
assert timeouts == [50.0, 100.0, 150.0]
def test_chat_attempt_info_populated_on_success(monkeypatch):
async def fake_post(url, payload, *, timeout=120.0, headers=None):
return {"choices": [{"message": {"content": "ok"}}]}
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
attempt_info: dict = {}
_run_async(
LlamaCppProvider("http://localhost:8080").chat(
"gemma-3-4b",
[Message(role="user", content="hi")],
options={"seed": 9},
attempt_info=attempt_info,
)
)
assert attempt_info == {
"seed": 9,
"attempts": 1,
"timeout_secs": 300.0,
"refusals": 0,
}
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
# ---------------------------------------------------------------------------
@@ -461,6 +536,64 @@ def test_chat_structured_forwards_options(monkeypatch):
assert captured["options"] == {"temperature": 0.0}
def test_chat_structured_forwards_attempt_info(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)
attempt_info: dict = {}
_run_async(
LlamaCppProvider("http://localhost:8080").chat_structured(
"gemma-3-4b",
[Message(role="user", content="hi")],
Widget,
attempt_info=attempt_info,
)
)
# Same object passed straight through — the shared helper populates it,
# this provider doesn't need to know its shape.
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).
@@ -558,3 +691,87 @@ def test_chat_text_only_content_stays_plain_string(monkeypatch):
)
assert captured["payload"]["messages"] == [{"role": "user", "content": "hi"}]
# ---------------------------------------------------------------------------
# refusal-retry — parity with test_ollama_provider.py's coverage (ADR: this
# is a model-behavior concern, not a backend one — both providers wire the
# same comfydv._llm.retry.is_refusal() into their own retry loop).
# ---------------------------------------------------------------------------
def test_chat_retries_on_lexical_refusal_and_returns_clean_second_attempt(
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)
result = _run_async(
LlamaCppProvider("http://localhost:8080").chat(
"gemma-3-4b",
[Message(role="user", content="hi")],
options={"refusal_retry": {"enabled": True, "embedding_model": ""}},
)
)
assert result == "a real, on-topic answer"
assert len(calls) == 2
def test_chat_refusal_retry_disabled_returns_refusal_text_unchanged(monkeypatch):
calls = []
async def fake_post(url, payload, *, timeout=120.0, headers=None):
calls.append(payload)
return {"choices": [{"message": {"content": "I cannot generate that."}}]}
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
result = _run_async(
LlamaCppProvider("http://localhost:8080").chat(
"gemma-3-4b", [Message(role="user", content="hi")]
)
)
assert result == "I cannot generate that."
assert len(calls) == 1
def test_embed_returns_vector_from_v1_embeddings(monkeypatch):
captured = {}
async def fake_post(url, payload, *, timeout=120.0, headers=None):
captured["url"] = url
captured["payload"] = payload
return {"data": [{"embedding": [0.4, 0.5, 0.6], "index": 0}]}
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
result = _run_async(
LlamaCppProvider("http://localhost:8080").embed("nomic-embed-text", "hello")
)
assert result == [0.4, 0.5, 0.6]
assert captured["url"] == "http://localhost:8080/v1/embeddings"
assert captured["payload"] == {"model": "nomic-embed-text", "input": "hello"}
def test_embed_returns_none_when_no_embedding_model_configured(monkeypatch):
async def fake_post(url, payload, *, timeout=120.0, headers=None):
raise RuntimeError("llama-server returned HTTP 404 for /v1/embeddings")
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
result = _run_async(
LlamaCppProvider("http://localhost:8080").embed("nomic-embed-text", "hello")
)
assert result is None
+224
View File
@@ -142,6 +142,100 @@ def test_chat_structured_retries_on_validation_failure(monkeypatch):
assert len(fake.calls) == 2
def test_chat_structured_timeout_escalates_per_retry_attempt(monkeypatch):
bad = ValidationError.from_exception_data("Widget", [])
fake = _FakeAgent([bad, bad, _Widget(name="b", count=2)])
build_calls = []
def fake_build_agent(**kw):
build_calls.append(kw["timeout_secs"])
return fake
monkeypatch.setattr(chat_mod, "_build_agent", fake_build_agent)
_run_async(
chat_mod.chat_structured(
base_url="http://localhost:11434/v1",
model="llama3",
messages=_messages(),
schema=_Widget,
timeout_secs=100.0,
max_retries=2,
)
)
assert build_calls == [100.0, 200.0, 300.0]
def test_chat_structured_attempt_info_populated_on_success(monkeypatch):
fake = _FakeAgent([_Widget(name="a", count=1)])
monkeypatch.setattr(chat_mod, "_build_agent", lambda **kw: fake)
attempt_info: dict = {}
_run_async(
chat_mod.chat_structured(
base_url="http://localhost:11434/v1",
model="llama3",
messages=_messages(),
schema=_Widget,
options={"seed": 5},
attempt_info=attempt_info,
)
)
assert attempt_info == {
"seed": 5,
"attempts": 1,
"timeout_secs": 300.0,
"refusals": 0,
}
def test_chat_structured_attempt_info_populated_on_exhaustion(monkeypatch):
bad = ValidationError.from_exception_data("Widget", [])
fake = _FakeAgent([bad, bad])
monkeypatch.setattr(chat_mod, "_build_agent", lambda **kw: fake)
attempt_info: dict = {}
with pytest.raises(RuntimeError):
_run_async(
chat_mod.chat_structured(
base_url="http://localhost:11434/v1",
model="llama3",
messages=_messages(),
schema=_Widget,
max_retries=1,
attempt_info=attempt_info,
)
)
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
@@ -381,6 +475,136 @@ def test_chat_structured_retry_sleeps_between_attempts(monkeypatch):
assert sleep_calls == [chat_mod.RETRY_BACKOFF_SECS]
# ---------------------------------------------------------------------------
# refusal-retry — parity with test_ollama_provider.py's coverage (this
# module serves LlamaCppProvider.chat_structured() — see
# comfydv._llm.retry.is_refusal and OllamaOptionRefusalRetry).
# ---------------------------------------------------------------------------
def test_chat_structured_retries_on_refusal_and_returns_clean_second_attempt(
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)
result = _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,
)
)
assert result == clean
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,
):
refused = _Widget(name="I cannot generate that content.", count=1)
fake = _FakeAgent([refused])
monkeypatch.setattr(chat_mod, "_build_agent", lambda **kw: fake)
result = _run_async(
chat_mod.chat_structured(
base_url="http://localhost:11434/v1",
model="llama3",
messages=_messages(),
schema=_Widget,
)
)
assert result == refused
assert len(fake.calls) == 1
def test_chat_structured_refusal_retry_exhausted_raises(monkeypatch):
refused = _Widget(name="I cannot help with this.", count=1)
fake = _FakeAgent([refused, refused])
monkeypatch.setattr(chat_mod, "_build_agent", lambda **kw: fake)
with pytest.raises(RuntimeError, match="failed validation"):
_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=1,
)
)
def test_chat_structured_refusal_retry_uses_embed_fn(monkeypatch):
from comfydv._llm.retry import REFUSAL_EXEMPLARS
refused = _Widget(name="not today, sorry", count=1)
clean = _Widget(name="a clean value", count=2)
fake = _FakeAgent([refused, clean])
monkeypatch.setattr(chat_mod, "_build_agent", lambda **kw: fake)
embed_calls = []
async def fake_embed(text):
embed_calls.append(text)
if "not today" in text or text in REFUSAL_EXEMPLARS:
return [1.0, 0.0]
return [0.0, 1.0]
result = _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": "nomic-embed-text",
"threshold": 0.5,
}
},
max_retries=2,
embed_fn=fake_embed,
)
)
assert result == clean
assert embed_calls # embedding path was actually exercised
def test_history_to_messages_preserves_order_and_roles():
from pydantic_ai.messages import ModelRequest, ModelResponse
+404 -2
View File
@@ -1,8 +1,29 @@
"""Tests for comfydv._llm.retry — shared retry-on-blank-output helpers used
by both providers' chat() and the shared chat_structured() helper.
by both providers' chat() and the shared chat_structured() helper, plus the
refusal/deflection detector that rides the same retry-with-a-new-seed
mechanism.
Refusal-detection tests here are pure-logic only — no provider/HTTP
involved. See test_ollama_provider.py, test_llamacpp_provider.py, and
test_llm_chat_structured.py for the retry-loop integration (does a detected
refusal actually trigger a reseeded retry).
"""
from comfydv._llm.retry import next_seed
import pytest
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,
next_seed,
next_timeout_secs,
record_attempt_info,
)
def test_next_seed_attempt_one_is_zero_by_default():
@@ -22,3 +43,384 @@ def test_next_seed_starts_from_pinned_base():
def test_next_seed_ignores_non_int_seed():
assert next_seed({"seed": "not-an-int"}, 2) == 1
class TestNextTimeoutSecs:
def test_attempt_one_returns_base_timeout_unchanged(self):
assert next_timeout_secs(300.0, 1) == 300.0
def test_escalates_multiplicatively_per_attempt(self):
assert next_timeout_secs(300.0, 2) == 600.0
assert next_timeout_secs(300.0, 3) == 900.0
class TestRecordAttemptInfo:
def test_none_attempt_info_is_a_no_op(self):
# Must not raise — callers that don't care about this metadata pass
# None and should see no behavior change at all.
record_attempt_info(None, seed=1, attempts=2, timeout_secs=600.0, refusals=1)
def test_populates_dict_in_place(self):
info: dict = {}
record_attempt_info(info, seed=7, attempts=3, timeout_secs=900.0, refusals=2)
assert info == {
"seed": 7,
"attempts": 3,
"timeout_secs": 900.0,
"refusals": 2,
}
def test_overwrites_previous_values(self):
# Callers call this once per attempt (or once per return path) —
# a later call must replace, not merge with, an earlier one.
info: dict = {"seed": 1, "attempts": 1, "timeout_secs": 300.0, "refusals": 0}
record_attempt_info(info, seed=2, attempts=2, timeout_secs=600.0, refusals=1)
assert info == {
"seed": 2,
"attempts": 2,
"timeout_secs": 600.0,
"refusals": 1,
}
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
# ---------------------------------------------------------------------------
class TestLexicalDetection:
def test_blatant_refusal_openers_match(self):
blatant = [
"I cannot generate that image for you.",
"I'm not able to help with this request.",
"I am unable to comply with your instruction.",
"Cannot generate this image due to the sensitive content.",
'{"error": "cannot generate this content"}',
"I don't feel comfortable creating this content.",
"I'm sorry, but I can't produce that description.",
"As an AI, I'm unable to depict this scenario.",
"I must decline this request.",
"This request violates my usage policies.",
"That's not something I'm able to help with right now.",
]
for text in blatant:
assert is_lexical_refusal(text), f"expected refusal match: {text!r}"
def test_ordinary_content_does_not_match(self):
ordinary = [
"The subject turns to face the camera and smiles warmly.",
"A person cannot simply walk into Mordor, the guide joked.",
"I can help you plan a birthday party for your dog.",
"",
]
for text in ordinary:
assert not is_lexical_refusal(text), f"unexpected match: {text!r}"
class TestAmbiguityHeuristic:
def test_short_response_is_ambiguous(self):
assert is_ambiguous("Sorry, can't do that one.")
def test_long_response_without_hedge_keywords_is_not_ambiguous(self):
long_text = "The subject rotates smoothly toward the lens. " * 20
assert len(long_text) >= 400
assert not is_ambiguous(long_text)
def test_long_response_with_hedge_keyword_is_ambiguous(self):
long_text = "Unfortunately, " + "this touches on a sensitive area. " * 20
assert len(long_text) >= 400
assert is_ambiguous(long_text)
def test_blank_text_is_not_ambiguous(self):
assert not is_ambiguous(" ")
class TestCosineSimilarity:
def test_identical_vectors_score_one(self):
assert cosine_similarity([1.0, 0.0], [1.0, 0.0]) == pytest.approx(1.0)
def test_orthogonal_vectors_score_zero(self):
assert cosine_similarity([1.0, 0.0], [0.0, 1.0]) == pytest.approx(0.0)
def test_opposite_vectors_score_negative_one(self):
assert cosine_similarity([1.0, 0.0], [-1.0, 0.0]) == pytest.approx(-1.0)
def test_mismatched_lengths_return_zero(self):
assert cosine_similarity([1.0, 0.0], [1.0, 0.0, 0.0]) == 0.0
def test_empty_vectors_return_zero(self):
assert cosine_similarity([], []) == 0.0
class TestIsRefusalHybrid:
def test_blank_text_is_never_a_refusal(self):
assert not _run_async(is_refusal(""))
assert not _run_async(is_refusal(" "))
def test_lexical_match_short_circuits_without_embedding(self):
calls = []
async def embed_fn(text):
calls.append(text)
return [1.0, 0.0]
result = _run_async(
is_refusal("I cannot generate that image for you.", embed_fn=embed_fn)
)
assert result is True
assert calls == [] # never reached the embedding step
def test_long_clean_response_is_still_embedding_checked_but_not_a_refusal(self):
# embed_fn present means the caller opted in to the embedding check
# regardless of length/hedge-keywords (is_ambiguous no longer gates
# this) — a long, on-topic response should still be embedding-
# checked, it just shouldn't score as similar to the refusal
# exemplars.
calls = []
async def embed_fn(text):
calls.append(text)
if text == long_text:
return [0.0, 1.0] # orthogonal to the exemplar vector below
return [1.0, 0.0] # exemplars
long_text = "The subject rotates smoothly toward the lens. " * 20
result = _run_async(is_refusal(long_text, embed_fn=embed_fn))
assert result is False
assert long_text in calls # embedding check DID run, just scored low
def test_no_embed_fn_degrades_to_lexical_only(self):
# Ambiguous (short), no lexical match, no embed_fn -> can't check further
assert not _run_async(is_refusal("Not today, sorry.", embed_fn=None))
def test_ambiguous_response_above_threshold_is_refusal(self):
async def embed_fn(text):
# Exemplars and a near-identical short "refusal-ish" probe get a
# high similarity score; distinguish by a marker substring.
if "PROBE" in text:
return [1.0, 0.0]
return [0.99, 0.14] # cos-sim with [1,0] is ~0.99
result = _run_async(
is_refusal(
"PROBE: not comfortable with this one",
embed_fn=embed_fn,
embed_cache_key="test-model",
threshold=0.8,
)
)
assert result is True
def test_ambiguous_response_below_threshold_is_not_refusal(self):
async def embed_fn(text):
if "PROBE" in text:
return [1.0, 0.0]
return [0.0, 1.0] # orthogonal -> cos-sim 0.0
result = _run_async(
is_refusal(
"PROBE: a short reply",
embed_fn=embed_fn,
embed_cache_key="test-model-2",
threshold=0.8,
)
)
assert result is False
def test_long_json_shaped_soft_refusal_without_hedge_keywords_is_caught(self):
# Regression case: a structured-output-shaped response (>600 chars
# once you count JSON braces/field names) whose deflection doesn't
# use any of the canned hedge keywords used to be invisible to the
# embedding check entirely, because is_ambiguous gated on length
# and keywords. embed_fn now runs unconditionally once configured.
soft_refusal = (
'{"prompt": "'
+ "Let's take this in a different creative direction that everyone can enjoy. "
* 8
+ '"}'
)
assert len(soft_refusal) >= 600
assert not is_lexical_refusal(soft_refusal)
async def embed_fn(text):
if text == soft_refusal:
return [1.0, 0.0]
return [0.97, 0.24] # exemplars: cos-sim with [1,0] is ~0.97
result = _run_async(
is_refusal(
soft_refusal,
embed_fn=embed_fn,
embed_cache_key="regression-model",
threshold=0.8,
)
)
assert result is True
def test_exemplar_embeddings_cached_across_calls(self):
exemplar_calls = {"n": 0}
async def embed_fn(text):
if "PROBE" in text:
return [1.0, 0.0]
exemplar_calls["n"] += 1
return [1.0, 0.0]
_run_async(
is_refusal(
"PROBE: first ambiguous call",
embed_fn=embed_fn,
embed_cache_key="cache-key-shared",
threshold=0.5,
)
)
first_count = exemplar_calls["n"]
assert first_count > 0
_run_async(
is_refusal(
"PROBE: second ambiguous call",
embed_fn=embed_fn,
embed_cache_key="cache-key-shared",
threshold=0.5,
)
)
# Exemplar embeddings reused from cache -> no additional exemplar calls
assert exemplar_calls["n"] == first_count
def test_embed_fn_failure_degrades_to_not_refused(self):
async def failing_embed_fn(text):
raise RuntimeError("server unreachable")
result = _run_async(
is_refusal(
"Not comfortable with this one, sorry.",
embed_fn=failing_embed_fn,
embed_cache_key="unreachable-model",
)
)
assert result is False
def test_embed_fn_returning_none_degrades_to_not_refused(self):
async def none_embed_fn(text):
return None
result = _run_async(
is_refusal(
"Not comfortable with this one, sorry.",
embed_fn=none_embed_fn,
embed_cache_key="no-embeddings-model",
)
)
assert result is False
class TestCustomPhrases:
def test_custom_phrase_substring_match_needs_no_embed_fn(self):
# A phrase the user added at runtime that the shipped lexical
# patterns don't cover — should be caught for free, no embedding
# model required.
result = _run_async(
is_refusal(
"I am restricted from producing that kind of content.",
custom_phrases=("restricted from",),
)
)
assert result is True
def test_custom_phrase_match_is_case_insensitive(self):
result = _run_async(
is_refusal(
"SORRY, THAT'S OFF LIMITS FOR ME.",
custom_phrases=("off limits",),
)
)
assert result is True
def test_unrelated_custom_phrase_does_not_match(self):
result = _run_async(
is_refusal(
"The subject walks calmly toward the horizon.",
custom_phrases=("restricted from", "off limits"),
)
)
assert result is False
def test_blank_and_whitespace_custom_phrases_are_ignored(self):
# A stray empty entry must never become a universal substring match.
result = _run_async(
is_refusal(
"The subject walks calmly toward the horizon.",
custom_phrases=("", " "),
)
)
assert result is False
def test_custom_phrase_folded_into_embedding_exemplars(self):
# No exact substring match, but embed_fn scores the response as
# similar to the custom phrase (not one of the shipped exemplars).
custom = "my creators have limited what I can show you"
async def embed_fn(text):
if text == custom:
return [1.0, 0.0]
if text in REFUSAL_EXEMPLARS:
return [0.0, 1.0] # shipped exemplars score orthogonal
return [0.99, 0.14] # the probe response is near the custom one
result = _run_async(
is_refusal(
"There are limits my creators placed on what I can show.",
embed_fn=embed_fn,
embed_cache_key="custom-exemplar-model",
threshold=0.8,
custom_phrases=(custom,),
)
)
assert result is True
def test_different_custom_phrase_sets_do_not_share_exemplar_cache(self):
# Regression guard: if the exemplar cache key ignored custom_phrases,
# a second call with a different custom phrase set would incorrectly
# reuse the first call's cached (and now stale) exemplar vectors.
calls = []
async def embed_fn(text):
calls.append(text)
return [1.0, 0.0]
_run_async(
is_refusal(
"short reply",
embed_fn=embed_fn,
embed_cache_key="shared-model",
custom_phrases=("phrase one",),
)
)
first_call_count = len(calls)
_run_async(
is_refusal(
"short reply",
embed_fn=embed_fn,
embed_cache_key="shared-model",
custom_phrases=("phrase two",),
)
)
# A fresh custom phrase set re-embeds the exemplars (including the
# new phrase) rather than reusing the first set's cached vectors.
assert len(calls) > first_call_count
+128 -20
View File
@@ -50,6 +50,7 @@ from comfydv.ollama import (
OllamaOptionDisableThinking,
OllamaOptionExtraBody,
OllamaOptionMaxTokens,
OllamaOptionRefusalRetry,
OllamaOptionRepeatPenalty,
OllamaOptionSeed,
OllamaOptionTemperature,
@@ -91,13 +92,37 @@ class _FakeProvider:
self.calls.append(("unload_model", model))
async def chat(
self, model, messages, options=None, timeout_secs=300.0, max_retries=2
self,
model,
messages,
options=None,
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:
attempt_info.update(
{
"seed": (options or {}).get("seed", 0),
"attempts": 1,
"timeout_secs": timeout_secs,
"refusals": 0,
}
)
return self.chat_response
async def chat_structured(
self, model, messages, schema, options=None, timeout_secs=300.0, max_retries=2
self,
model,
messages,
schema,
options=None,
timeout_secs=300.0,
max_retries=2,
attempt_info=None,
on_status=None,
):
self.calls.append(
(
@@ -112,6 +137,15 @@ class _FakeProvider:
)
if self.raise_on_chat_structured is not None:
raise self.raise_on_chat_structured
if attempt_info is not None:
attempt_info.update(
{
"seed": (options or {}).get("seed", 0),
"attempts": 1,
"timeout_secs": timeout_secs,
"refusals": 0,
}
)
return schema(**self.structured_field_values)
@@ -271,7 +305,7 @@ class TestUS4ChatCompletion:
def test_chat_model_receives_wired_string(self):
"""Wiring LLMLoadModel.model_name → ChatCompletion.model works."""
fake = _FakeProvider(chat_response="ok")
_, _, used_model = ChatCompletion().chat(
_, _, used_model, _ = ChatCompletion().chat(
client=fake, model="llama3:latest", prompt="hi"
)["result"]
assert used_model == "llama3:latest"
@@ -279,11 +313,17 @@ class TestUS4ChatCompletion:
assert fake.calls[0][1] == "llama3:latest"
def test_chat_completion_returns_model_name(self):
assert ChatCompletion.RETURN_TYPES == ("STRING", "OLLAMA_HISTORY", "STRING")
assert ChatCompletion.RETURN_TYPES == (
"STRING",
"OLLAMA_HISTORY",
"STRING",
"INT",
)
assert ChatCompletion.RETURN_NAMES == (
"response",
"updated_history",
"model_name",
"seed_used",
)
def test_chat_is_output_node(self):
@@ -301,15 +341,24 @@ class TestUS4ChatCompletion:
ret = ChatCompletion().chat(client=fake, model="m", prompt="hi")
assert "hello world" in ret["ui"]["text"][0]
def test_chat_result_is_3_tuple(self):
def test_chat_result_is_4_tuple(self):
fake = _FakeProvider(chat_response="hello")
ret = ChatCompletion().chat(client=fake, model="m", prompt="hi")
assert isinstance(ret["result"], tuple)
assert len(ret["result"]) == 3
response, history, model_name = ret["result"]
assert len(ret["result"]) == 4
response, history, model_name, seed_used = ret["result"]
assert response == "hello"
assert isinstance(history, list)
assert model_name == "m"
assert seed_used == 0
def test_chat_result_seed_used_reflects_pinned_seed(self):
fake = _FakeProvider(chat_response="hello")
ret = ChatCompletion().chat(
client=fake, model="m", prompt="hi", options={"seed": 42}
)
_, _, _, seed_used = ret["result"]
assert seed_used == 42
def test_chat_has_timeout_secs_input(self):
input_types = ChatCompletion.INPUT_TYPES()
@@ -336,7 +385,7 @@ class TestUS4ChatCompletion:
):
"""Scenario: Single-turn completion returns non-empty response."""
(client,) = OllamaClient().create_client(ollama_host)
response, updated_history, model_name = ChatCompletion().chat(
response, updated_history, model_name, _seed = ChatCompletion().chat(
client=client,
model=_CHAT_MODEL,
prompt="Say exactly the word: pong",
@@ -358,14 +407,14 @@ class TestUS4ChatCompletion:
"""
(client,) = OllamaClient().create_client(ollama_host)
no_think = {"think": False}
_, history, _ = ChatCompletion().chat(
_, history, _, _ = ChatCompletion().chat(
client=client,
model=_CHAT_MODEL,
prompt="My name is Alice. Remember it.",
history=[],
options=no_think,
)["result"]
response, updated, _ = ChatCompletion().chat(
response, updated, _, _ = ChatCompletion().chat(
client=client,
model=_CHAT_MODEL,
prompt="What is my name?",
@@ -379,11 +428,11 @@ class TestUS4ChatCompletion:
def test_history_accumulated_correctly(self, ollama_host, skip_if_no_ollama):
"""History list grows by 2 entries per turn."""
(client,) = OllamaClient().create_client(ollama_host)
_, h1, _ = ChatCompletion().chat(
_, h1, _, _ = ChatCompletion().chat(
client=client, model=_CHAT_MODEL, prompt="Turn 1", history=[]
)["result"]
assert len(h1) == 2
_, h2, _ = ChatCompletion().chat(
_, h2, _, _ = ChatCompletion().chat(
client=client, model=_CHAT_MODEL, prompt="Turn 2", history=h1
)["result"]
assert len(h2) == 4
@@ -427,7 +476,7 @@ class TestUS4ChatCompletion:
max_retries=1,
unique_id="smoke-test",
)
response_text, _history, model_name, output = result["result"]
response_text, _history, model_name, output, _seed = result["result"]
assert model_name == _CHAT_MODEL
assert json.loads(response_text) == {"output": output}
assert isinstance(output, str) and output.strip()
@@ -503,8 +552,9 @@ class TestStructuredOutput:
"updated_history",
"model_name",
"output",
"seed_used",
)
assert len(ret["result"]) == 4
assert len(ret["result"]) == 5
assert ret["result"][3] == "clean text"
assert json.loads(ret["result"][0]) == {"output": "clean text"}
@@ -531,6 +581,7 @@ class TestStructuredOutput:
"STRING",
"INT",
"BOOLEAN",
"INT",
)
assert ChatCompletion.RETURN_NAMES == (
"response",
@@ -539,8 +590,9 @@ class TestStructuredOutput:
"summary",
"score",
"is_positive",
"seed_used",
)
_, _, _, summary, score, is_positive = ret["result"]
_, _, _, summary, score, is_positive, _seed = ret["result"]
assert summary == "great"
assert score == 9
assert isinstance(score, int)
@@ -568,7 +620,7 @@ class TestStructuredOutput:
output_schema=_SCHEMA_WITH_OPTIONAL_NULLABLE_FIELD,
unique_id="n3b",
)
_, _, _, output, duration_seconds = ret["result"]
_, _, _, output, duration_seconds, _seed = ret["result"]
assert output == "hi"
assert duration_seconds is None # FLOAT socket; only STRING coerces None to ""
@@ -597,7 +649,7 @@ class TestStructuredOutput:
structured_output=True,
unique_id="n5",
)
assert len(ChatCompletion.RETURN_TYPES) == 4
assert len(ChatCompletion.RETURN_TYPES) == 5
fake2 = _FakeProvider(chat_response="x")
ChatCompletion().chat(
@@ -732,6 +784,55 @@ class TestUS5ComposableOptions:
)
assert opts == {"temperature": 0.5, "think": False}
def test_refusal_retry_default_merges_config_dict(self):
(opts,) = OllamaOptionRefusalRetry().set_refusal_retry(
enabled=True, embedding_model="", threshold=0.82
)
assert opts == {
"refusal_retry": {
"enabled": True,
"embedding_model": "",
"threshold": 0.82,
"custom_phrases": (),
}
}
def test_refusal_retry_strips_embedding_model_whitespace(self):
(opts,) = OllamaOptionRefusalRetry().set_refusal_retry(
enabled=True, embedding_model=" nomic-embed-text ", threshold=0.9
)
assert opts["refusal_retry"]["embedding_model"] == "nomic-embed-text"
def test_refusal_retry_merges_existing_options(self):
(opts,) = OllamaOptionRefusalRetry().set_refusal_retry(
enabled=False,
embedding_model="",
threshold=0.82,
options={"temperature": 0.5},
)
assert opts["temperature"] == 0.5
assert opts["refusal_retry"]["enabled"] is False
def test_refusal_retry_parses_custom_phrases_csv(self):
# Whitespace around each phrase is trimmed and empty entries (from
# blank items / trailing commas) are dropped.
(opts,) = OllamaOptionRefusalRetry().set_refusal_retry(
enabled=True,
embedding_model="",
threshold=0.82,
custom_phrases=" I am restricted from , not permitted to help,,",
)
assert opts["refusal_retry"]["custom_phrases"] == (
"I am restricted from",
"not permitted to help",
)
def test_refusal_retry_custom_phrases_defaults_empty(self):
(opts,) = OllamaOptionRefusalRetry().set_refusal_retry(
enabled=True, embedding_model="", threshold=0.82
)
assert opts["refusal_retry"]["custom_phrases"] == ()
def test_extra_body_merges_json(self):
(opts,) = OllamaOptionExtraBody().set_extra_body(
extra_body_json='{"stop": ["</s>"]}', options={"temperature": 0.0}
@@ -768,8 +869,8 @@ class TestUS5ComposableOptions:
history=[],
options=opts2,
)
r1, _, _model = ChatCompletion().chat(**kwargs)["result"]
r2, _, _model = ChatCompletion().chat(**kwargs)["result"]
r1, _, _model, _seed = ChatCompletion().chat(**kwargs)["result"]
r2, _, _model, _seed = ChatCompletion().chat(**kwargs)["result"]
assert r1 == r2
@@ -1008,6 +1109,7 @@ class TestUpdateStructuredOutputsRoute:
"response",
"updated_history",
"model_name",
"seed_used",
]
def test_valid_schema_returns_dynamic_outputs(self):
@@ -1027,6 +1129,7 @@ class TestUpdateStructuredOutputsRoute:
"summary",
"score",
"is_positive",
"seed_used",
]
assert types == [
"STRING",
@@ -1035,6 +1138,7 @@ class TestUpdateStructuredOutputsRoute:
"STRING",
"INT",
"BOOLEAN",
"INT",
]
def test_invalid_json_while_typing_falls_back_to_base_outputs(self):
@@ -1052,6 +1156,7 @@ class TestUpdateStructuredOutputsRoute:
"response",
"updated_history",
"model_name",
"seed_used",
]
def test_incomplete_schema_missing_properties_falls_back_to_base_outputs(self):
@@ -1066,6 +1171,7 @@ class TestUpdateStructuredOutputsRoute:
"response",
"updated_history",
"model_name",
"seed_used",
]
def test_response_reflects_class_state_not_just_this_call(self):
@@ -1079,7 +1185,8 @@ class TestUpdateStructuredOutputsRoute:
"output_schema": _SINGLE_FIELD_SCHEMA,
}
)
assert ChatCompletion.RETURN_NAMES[-1] == "output"
assert ChatCompletion.RETURN_NAMES[-2] == "output"
assert ChatCompletion.RETURN_NAMES[-1] == "seed_used"
data = self._call(
{"unique_id": "r5", "structured_output": False, "output_schema": "{}"}
@@ -1088,6 +1195,7 @@ class TestUpdateStructuredOutputsRoute:
"response",
"updated_history",
"model_name",
"seed_used",
]
assert ChatCompletion.RETURN_NAMES == ChatCompletion._BASE_RETURN_NAMES
+417
View File
@@ -19,6 +19,7 @@ import pytest
import comfydv._llm.ollama_provider as provider_mod
from comfydv._llm.ollama_provider import OllamaProvider, _run_async
from comfydv._llm.provider import Message, ModelStatus
from comfydv._llm.retry import REFUSAL_EXEMPLARS
@pytest.fixture(autouse=True)
@@ -924,3 +925,419 @@ def test_chat_text_only_payload_omits_images_key(monkeypatch):
)
assert captured["payload"]["messages"] == [{"role": "user", "content": "hi"}]
# ---------------------------------------------------------------------------
# attempt_info / escalating per-attempt timeout
# ---------------------------------------------------------------------------
def test_chat_timeout_escalates_per_retry_attempt(monkeypatch):
timeouts = []
async def fake_post(url, payload, *, timeout=120.0, headers=None):
timeouts.append(timeout)
if len(timeouts) < 3:
return {"message": {"content": ""}} # blank -> retry
return {"message": {"content": "done"}}
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
monkeypatch.setattr(provider_mod.asyncio, "sleep", _fake_sleep)
_run_async(
OllamaProvider("http://localhost:11434").chat(
"llama3",
[Message(role="user", content="hi")],
timeout_secs=100.0,
max_retries=2,
)
)
assert timeouts == [100.0, 200.0, 300.0]
def test_chat_attempt_info_populated_on_success(monkeypatch):
async def fake_post(url, payload, *, timeout=120.0, headers=None):
return {"message": {"content": "ok"}}
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
attempt_info: dict = {}
_run_async(
OllamaProvider("http://localhost:11434").chat(
"llama3",
[Message(role="user", content="hi")],
options={"seed": 5},
attempt_info=attempt_info,
)
)
assert attempt_info == {
"seed": 5,
"attempts": 1,
"timeout_secs": 300.0,
"refusals": 0,
}
def test_chat_attempt_info_reflects_refusal_retry(monkeypatch):
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)
attempt_info: dict = {}
_run_async(
OllamaProvider("http://localhost:11434").chat(
"llama3",
[Message(role="user", content="hi")],
options={"refusal_retry": {"enabled": True, "embedding_model": ""}},
attempt_info=attempt_info,
)
)
assert attempt_info["attempts"] == 2
assert attempt_info["refusals"] == 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
# OllamaOptionRefusalRetry / comfydv._llm.retry.is_refusal)
# ---------------------------------------------------------------------------
def test_chat_retries_on_lexical_refusal_and_returns_clean_second_attempt(
monkeypatch,
):
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)
result = _run_async(
OllamaProvider("http://localhost:11434").chat(
"llama3",
[Message(role="user", content="hi")],
options={"refusal_retry": {"enabled": True, "embedding_model": ""}},
)
)
assert result == "a real, on-topic answer"
assert len(calls) == 2
# next_seed(): attempt 2 gets a bumped seed, not a repeat of attempt 1
assert calls[1]["options"]["seed"] == 1
def test_chat_retries_on_custom_phrase_with_no_embedding_model(monkeypatch):
"""Regression coverage for OllamaOptionRefusalRetry's custom_phrases
field: a phrase the user added at runtime, not one of the shipped
lexical patterns, still triggers a retry — with no embedding_model
configured at all."""
calls = []
async def fake_post(url, payload, *, timeout=120.0, headers=None):
calls.append(payload)
if len(calls) == 1:
# deliberately not covered by any shipped lexical pattern —
# only the custom_phrases entry below should catch this
return {"message": {"content": "That's outside my operating parameters."}}
return {"message": {"content": "a real, on-topic answer"}}
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
monkeypatch.setattr(provider_mod.asyncio, "sleep", _fake_sleep)
result = _run_async(
OllamaProvider("http://localhost:11434").chat(
"llama3",
[Message(role="user", content="hi")],
options={
"refusal_retry": {
"enabled": True,
"embedding_model": "",
"custom_phrases": ("outside my operating parameters",),
}
},
)
)
assert result == "a real, on-topic answer"
assert len(calls) == 2
def test_chat_refusal_retry_disabled_returns_refusal_text_unchanged(monkeypatch):
calls = []
async def fake_post(url, payload, *, timeout=120.0, headers=None):
calls.append(payload)
return {"message": {"content": "I cannot generate that for you."}}
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
result = _run_async(
OllamaProvider("http://localhost:11434").chat(
"llama3", [Message(role="user", content="hi")]
)
)
assert result == "I cannot generate that for you."
assert len(calls) == 1
def test_chat_refused_response_is_never_cached(monkeypatch):
"""A refused attempt must not poison the cache the way a genuine success
does — otherwise every subsequent identical request would replay the
refusal forever (the exact bug the disable-thinking fix worked around
for a different failure mode)."""
calls = []
async def fake_post(url, payload, *, timeout=120.0, headers=None):
calls.append(payload)
return {"message": {"content": "I cannot generate that for you."}}
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
monkeypatch.setattr(provider_mod.asyncio, "sleep", _fake_sleep)
messages = [Message(role="user", content="hi")]
options = {"refusal_retry": {"enabled": True, "embedding_model": ""}}
_run_async(
OllamaProvider("http://localhost:11434").chat(
"llama3", messages, options=options, max_retries=0
)
)
_run_async(
OllamaProvider("http://localhost:11434").chat(
"llama3", messages, options=options, max_retries=0
)
)
# Every attempt actually hit the network — nothing was ever cached.
assert len(calls) == 2
def test_chat_structured_retries_on_refusal_embedded_in_json_field(monkeypatch):
from pydantic import BaseModel
class Widget(BaseModel):
name: str
responses = [
{"message": {"content": '{"name": "I cannot generate that content."}'}},
{"message": {"content": '{"name": "a clean value"}'}},
]
calls = []
async def fake_post(url, payload, *, timeout=120.0, headers=None):
calls.append(payload)
return responses[len(calls) - 1]
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
monkeypatch.setattr(provider_mod.asyncio, "sleep", _fake_sleep)
result = _run_async(
OllamaProvider("http://localhost:11434").chat_structured(
"llama3",
[Message(role="user", content="hi")],
Widget,
options={"refusal_retry": {"enabled": True, "embedding_model": ""}},
max_retries=2,
)
)
assert result == Widget(name="a clean value")
assert len(calls) == 2
def test_chat_structured_refusal_retry_exhausted_raises(monkeypatch):
from pydantic import BaseModel
class Widget(BaseModel):
name: str
async def fake_post(url, payload, *, timeout=120.0, headers=None):
return {"message": {"content": '{"name": "I cannot help with this."}'}}
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
monkeypatch.setattr(provider_mod.asyncio, "sleep", _fake_sleep)
with pytest.raises(RuntimeError, match="failed validation"):
_run_async(
OllamaProvider("http://localhost:11434").chat_structured(
"llama3",
[Message(role="user", content="hi")],
Widget,
options={"refusal_retry": {"enabled": True, "embedding_model": ""}},
max_retries=1,
)
)
def test_chat_structured_refusal_retry_uses_embedding_fallback(monkeypatch):
"""embedding_model set -> embed() is actually consulted for an ambiguous
(non-lexical-match) response, not just the free regex pass."""
from pydantic import BaseModel
class Widget(BaseModel):
name: str
responses = [
{"message": {"content": '{"name": "not today, sorry"}'}},
{"message": {"content": '{"name": "a clean value"}'}},
]
calls = []
embed_calls = []
async def fake_post(url, payload, *, timeout=120.0, headers=None):
calls.append(payload)
return responses[len(calls) - 1]
async def fake_embed(self, model, text):
embed_calls.append((model, text))
# Exemplars and the refusal-flavored response land in the same
# direction (high cosine sim); the clean response is orthogonal.
if "not today" in text or text in REFUSAL_EXEMPLARS:
return [1.0, 0.0]
return [0.0, 1.0]
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
monkeypatch.setattr(provider_mod.asyncio, "sleep", _fake_sleep)
monkeypatch.setattr(OllamaProvider, "embed", fake_embed)
result = _run_async(
OllamaProvider("http://localhost:11434").chat_structured(
"llama3",
[Message(role="user", content="hi")],
Widget,
options={
"refusal_retry": {
"enabled": True,
"embedding_model": "nomic-embed-text",
"threshold": 0.5,
}
},
max_retries=2,
)
)
assert result == Widget(name="a clean value")
assert len(calls) == 2
assert embed_calls # embedding path was actually exercised
assert all(model == "nomic-embed-text" for model, _ in embed_calls)
# ---------------------------------------------------------------------------
# embed()
# ---------------------------------------------------------------------------
def test_embed_returns_vector_from_api_embed(monkeypatch):
captured = {}
async def fake_post(url, payload, *, timeout=120.0, headers=None):
captured["url"] = url
captured["payload"] = payload
return {"embeddings": [[0.1, 0.2, 0.3]]}
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
result = _run_async(
OllamaProvider("http://localhost:11434").embed("nomic-embed-text", "hello")
)
assert result == [0.1, 0.2, 0.3]
assert captured["url"] == "http://localhost:11434/api/embed"
assert captured["payload"] == {"model": "nomic-embed-text", "input": "hello"}
def test_embed_returns_none_on_unreachable_server(monkeypatch):
async def fake_post(url, payload, *, timeout=120.0, headers=None):
raise RuntimeError("Cannot reach Ollama at ...")
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
result = _run_async(
OllamaProvider("http://localhost:11434").embed("nomic-embed-text", "hello")
)
assert result is None
def test_embed_returns_none_for_malformed_response(monkeypatch):
async def fake_post(url, payload, *, timeout=120.0, headers=None):
return {"unexpected": "shape"}
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
result = _run_async(
OllamaProvider("http://localhost:11434").embed("nomic-embed-text", "hello")
)
assert result is None
def test_embed_returns_none_for_empty_model_or_text():
provider = OllamaProvider("http://localhost:11434")
assert _run_async(provider.embed("", "hello")) is None
assert _run_async(provider.embed("nomic-embed-text", "")) is None
+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())