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:
co-authored by
Claude Sonnet 5
parent
42c614d5f8
commit
7a40d22230
@@ -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: "
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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: "
|
||||
|
||||
@@ -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.
|
||||
"""
|
||||
...
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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).
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user