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>
This commit is contained in:
James Veitch
2026-07-29 22:19:22 +01:00
co-authored by Claude Sonnet 5
parent 42c614d5f8
commit 7a40d22230
12 changed files with 625 additions and 42 deletions
+39 -8
View File
@@ -46,7 +46,14 @@ from pydantic_ai.providers.openai import OpenAIProvider
from pydantic_ai.settings import ModelSettings
from .provider import Message
from .retry import RETRY_BACKOFF_SECS, EmbedFn, is_refusal, next_seed
from .retry import (
RETRY_BACKOFF_SECS,
EmbedFn,
is_refusal,
next_seed,
next_timeout_secs,
record_attempt_info,
)
_STRUCTURED_OUTPUT_FAILURE_EXCEPTIONS = (
UnexpectedModelBehavior,
@@ -127,6 +134,7 @@ async def chat_structured(
max_retries: int = 2,
timeout_secs: float = 300.0,
embed_fn: EmbedFn | None = None,
attempt_info: dict | None = None,
) -> BaseModel:
"""Call ``model`` at ``base_url`` (an OpenAI-compatible ``/v1`` root) and
return a validated instance of ``schema``.
@@ -171,13 +179,6 @@ async def chat_structured(
"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
@@ -204,7 +205,21 @@ 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
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
@@ -217,6 +232,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"],
@@ -263,11 +279,19 @@ async def chat_structured(
custom_phrases=tuple(refusal_cfg.get("custom_phrases") or ()),
)
if refused:
refusal_count += 1
last_error = RuntimeError("refusal/deflection detected")
last_invalid_text = content
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,
)
return output
except _STRUCTURED_OUTPUT_FAILURE_EXCEPTIONS as exc:
last_error = exc
@@ -275,6 +299,13 @@ async def chat_structured(
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: "
+51 -2
View File
@@ -27,7 +27,13 @@ from .ollama_provider import (
_post_json,
)
from .provider import Message, ModelInfo, ModelStatus
from .retry import RETRY_BACKOFF_SECS, is_refusal, next_seed
from .retry import (
RETRY_BACKOFF_SECS,
is_refusal,
next_seed,
next_timeout_secs,
record_attempt_info,
)
logger = logging.getLogger(__name__)
@@ -190,6 +196,7 @@ class LlamaCppProvider:
options: dict | None = None,
timeout_secs: float = 300.0,
max_retries: int = 2,
attempt_info: dict | None = None,
) -> str:
payload_messages = [_to_openai_message(m) for m in messages]
options, think = _pop_think(options)
@@ -203,8 +210,12 @@ class LlamaCppProvider:
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,
@@ -233,6 +244,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",
@@ -246,12 +264,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 []
@@ -272,11 +297,26 @@ class LlamaCppProvider:
)
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,
)
return response_text
refusal_count += 1
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,
)
# 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.
@@ -290,6 +330,7 @@ class LlamaCppProvider:
options: dict | None = None,
timeout_secs: float = 300.0,
max_retries: int = 2,
attempt_info: dict | None = None,
) -> BaseModel:
from .chat import chat_structured as _chat_structured_impl
@@ -305,6 +346,13 @@ 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
@@ -327,6 +375,7 @@ class LlamaCppProvider:
max_retries=max_retries,
timeout_secs=timeout_secs,
embed_fn=embed_fn,
attempt_info=attempt_info,
)
_CHAT_RESPONSE_CACHE.set(cache_key, result.model_dump())
return result
+66 -3
View File
@@ -19,7 +19,13 @@ import time
from pydantic import BaseModel, ValidationError
from .provider import Message, ModelInfo, ModelStatus
from .retry import RETRY_BACKOFF_SECS, is_refusal, next_seed
from .retry import (
RETRY_BACKOFF_SECS,
is_refusal,
next_seed,
next_timeout_secs,
record_attempt_info,
)
logger = logging.getLogger(__name__)
@@ -359,6 +365,7 @@ class OllamaProvider:
options: dict | None = None,
timeout_secs: float = 300.0,
max_retries: int = 2,
attempt_info: dict | None = None,
) -> str:
if any(m.images for m in messages):
await _require_vision_capability(self.host, model, self.headers)
@@ -380,11 +387,16 @@ class OllamaProvider:
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,
@@ -409,12 +421,19 @@ 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", "")
@@ -430,11 +449,19 @@ class OllamaProvider:
)
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,
)
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
# done: false alongside blank content is a distinct signal from
# an ordinary blank generation — it's Ollama answering before
@@ -448,6 +475,14 @@ class OllamaProvider:
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,
)
if incomplete:
raise RuntimeError(
f"Ollama returned an incomplete response after "
@@ -470,6 +505,7 @@ class OllamaProvider:
options: dict | None = None,
timeout_secs: float = 300.0,
max_retries: int = 2,
attempt_info: dict | None = None,
) -> BaseModel:
"""Native ``/api/chat`` + ``"format"`` (grammar-constrained JSON
decoding), not the shared pydantic-ai ``chat.py`` helper.
@@ -515,16 +551,28 @@ 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
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,
@@ -543,7 +591,7 @@ 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:
@@ -577,6 +625,7 @@ class OllamaProvider:
custom_phrases=custom_phrases,
)
if refused:
refusal_count += 1
last_error = RuntimeError("refusal/deflection detected")
last_invalid_text = content
if attempt < total_attempts:
@@ -584,8 +633,22 @@ class OllamaProvider:
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,
)
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: "
+13 -2
View File
@@ -83,6 +83,7 @@ class LLMProvider(Protocol):
options: dict | None = None,
timeout_secs: float = 300.0,
max_retries: int = 2,
attempt_info: dict | None = None,
) -> str:
"""Free-text chat response.
@@ -92,7 +93,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 +104,12 @@ 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.
"""
...
@@ -112,6 +121,7 @@ class LLMProvider(Protocol):
options: dict | None = None,
timeout_secs: float = 300.0,
max_retries: int = 2,
attempt_info: dict | None = None,
) -> BaseModel:
"""Schema-validated chat response.
@@ -121,7 +131,8 @@ 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`` out-param convention.
"""
...
+42
View File
@@ -47,6 +47,48 @@ def next_seed(options: dict | None, attempt: int) -> int:
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,
}
)
# ---------------------------------------------------------------------------
# Refusal/deflection detection
# ---------------------------------------------------------------------------
+55 -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,12 @@ 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 = {}
# Provider owns transport, caching, and — for structured_output — the
# tool-calling/retry/validation mechanism (pydantic-ai, ADR-007).
@@ -637,6 +661,7 @@ class ChatCompletion:
llm_options,
timeout_secs=float(timeout_secs),
max_retries=max_retries,
attempt_info=attempt_info,
)
)
else:
@@ -651,20 +676,43 @@ class ChatCompletion:
llm_options,
timeout_secs=float(timeout_secs),
max_retries=max_retries,
attempt_info=attempt_info,
)
)
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 +722,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]},
+16 -1
View File
@@ -40,9 +40,24 @@ 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,
):
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
+77
View File
@@ -405,6 +405,54 @@ 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,
}
# ---------------------------------------------------------------------------
# chat_structured — zero new logic, delegates to the shared helper unchanged
# ---------------------------------------------------------------------------
@@ -461,6 +509,35 @@ 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
# ---------------------------------------------------------------------------
# _fetch_models — name-only view used by ComfyUI's /dv/ollama/models?backend=
# llamacpp route (the JS refresh button / node-creation auto-populate).
+70
View File
@@ -142,6 +142,76 @@ 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_exhausted_retries_raises_runtime_error(monkeypatch):
bad = ValidationError.from_exception_data("Widget", [])
fake = _FakeAgent([bad, bad, bad]) # max_retries=2 -> 3 total attempts
+40
View File
@@ -19,6 +19,8 @@ from comfydv._llm.retry import (
is_lexical_refusal,
is_refusal,
next_seed,
next_timeout_secs,
record_attempt_info,
)
@@ -41,6 +43,44 @@ 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,
}
# ---------------------------------------------------------------------------
# Refusal/deflection detection
# ---------------------------------------------------------------------------
+76 -20
View File
@@ -92,13 +92,35 @@ 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,
):
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,
):
self.calls.append(
(
@@ -113,6 +135,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)
@@ -272,7 +303,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"
@@ -280,11 +311,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):
@@ -302,15 +339,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()
@@ -337,7 +383,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",
@@ -359,14 +405,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?",
@@ -380,11 +426,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
@@ -428,7 +474,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()
@@ -504,8 +550,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"}
@@ -532,6 +579,7 @@ class TestStructuredOutput:
"STRING",
"INT",
"BOOLEAN",
"INT",
)
assert ChatCompletion.RETURN_NAMES == (
"response",
@@ -540,8 +588,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)
@@ -569,7 +618,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 ""
@@ -598,7 +647,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(
@@ -818,8 +867,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
@@ -1058,6 +1107,7 @@ class TestUpdateStructuredOutputsRoute:
"response",
"updated_history",
"model_name",
"seed_used",
]
def test_valid_schema_returns_dynamic_outputs(self):
@@ -1077,6 +1127,7 @@ class TestUpdateStructuredOutputsRoute:
"summary",
"score",
"is_positive",
"seed_used",
]
assert types == [
"STRING",
@@ -1085,6 +1136,7 @@ class TestUpdateStructuredOutputsRoute:
"STRING",
"INT",
"BOOLEAN",
"INT",
]
def test_invalid_json_while_typing_falls_back_to_base_outputs(self):
@@ -1102,6 +1154,7 @@ class TestUpdateStructuredOutputsRoute:
"response",
"updated_history",
"model_name",
"seed_used",
]
def test_incomplete_schema_missing_properties_falls_back_to_base_outputs(self):
@@ -1116,6 +1169,7 @@ class TestUpdateStructuredOutputsRoute:
"response",
"updated_history",
"model_name",
"seed_used",
]
def test_response_reflects_class_state_not_just_this_call(self):
@@ -1129,7 +1183,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": "{}"}
@@ -1138,6 +1193,7 @@ class TestUpdateStructuredOutputsRoute:
"response",
"updated_history",
"model_name",
"seed_used",
]
assert ChatCompletion.RETURN_NAMES == ChatCompletion._BASE_RETURN_NAMES
+80
View File
@@ -927,6 +927,86 @@ 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
# ---------------------------------------------------------------------------
# refusal-retry — comfydv-level "refusal_retry" options convention (see
# OllamaOptionRefusalRetry / comfydv._llm.retry.is_refusal)