fix: retry blank/failed LLM responses with a new seed and backoff
Reported live on a fresh runpod: structured_output extraction on ChatCompletion would fail with RuntimeError: chat_structured: response failed validation against schema after 3 attempt(s) ... Last error: Exceeded maximum output retries (0) ...but succeed if a plain (non-structured) chat call with the same prompt was run first. Root cause: pydantic-ai's Agent is built with retries=0 (comfydv drives its own outer retry loop, not pydantic-ai's internal one), so a single bad/empty tool-call response instantly exhausts one attempt. comfydv's own retry loop then retried with the *exact same request* every time — if the model is genuinely stuck (e.g. still warming up on a freshly-started backend) or produces a degenerate response for a given seed, every attempt fails identically. Confirmed the seed was unset in the reporting user's workflow (Ollama picks its own per call) and it still failed 3/3 in a row, which rules out "stuck on one bad seed" and points at request-level determinism plus zero backoff between attempts as the real gap. Plain chat() had zero retry/validation logic at all before this change — it silently returned whatever content came back, even blank. Given the underlying failure mode (a fresh model's first response sometimes being blank) is identical for both chat modes, this fix applies to both, not just structured output. Adds src/comfydv/_llm/retry.py: next_seed() (deterministic, starting from any user-pinned options["seed"], incrementing per retry — attempt 1 is always the caller's original, untouched request) and RETRY_BACKOFF_SECS (a flat delay between retries, giving a still- loading model time to finish rather than hammering it with identical requests back-to-back). Both OllamaProvider.chat() and LlamaCppProvider.chat() now retry (per the node's existing max_retries widget) when a response comes back blank, injecting the incremented seed each retry — Ollama via the native options.seed field, llama.cpp via the OpenAI spec's top-level "seed" field (nesting it under "options" like the rest of the Ollama-native passthrough would silently not work — llama-server's OpenAI-compatible endpoint doesn't read seed from there, a pre-existing, documented non-goal for general options passthrough that doesn't apply to this internal retry mechanism). Neither raises if every retry stays blank — matches chat()'s existing never-validates contract. chat_structured() in _llm/chat.py now injects the same incrementing seed via pydantic-ai's ModelSettings["seed"] (which correctly maps to the OpenAI API's top-level seed param for both backends) and adds the same backoff between attempts, on top of its existing validation-retry loop. Live-verified against real Ollama: the happy path (model responds normally first try) still makes exactly one network call for chat() and completes chat_structured() normally — no regression from the new retry scaffolding. The specific cold/stuck-model failure itself is runpod-cold-start-dependent and wasn't independently reproducible in this session's environment (a second live attempt succeeded on the very first try) — unit tests directly exercise the retry/seed/backoff logic instead (test_llm_retry.py, and new cases in test_ollama_provider.py, test_llamacpp_provider.py, and test_llm_chat_structured.py), verified against both providers. Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_0132ojafeazQ3ephcBejEWFj
This commit is contained in:
co-authored by
Claude Sonnet 5
parent
92da4fc117
commit
09a173e552
@@ -13,6 +13,7 @@ drives its own retry loop so the error contract is comfydv's, not
|
||||
pydantic-ai's internal one.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import cast
|
||||
|
||||
from pydantic import BaseModel, ValidationError
|
||||
@@ -30,6 +31,7 @@ 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
|
||||
|
||||
_STRUCTURED_OUTPUT_FAILURE_EXCEPTIONS = (
|
||||
UnexpectedModelBehavior,
|
||||
@@ -122,10 +124,26 @@ async def chat_structured(
|
||||
total_attempts = max(0, min(int(max_retries), 5)) + 1
|
||||
last_error: Exception | None = None
|
||||
last_invalid_text = ""
|
||||
for _attempt in range(1, total_attempts + 1):
|
||||
for attempt in range(1, total_attempts + 1):
|
||||
attempt_settings = dict(model_settings) if model_settings else {}
|
||||
if attempt > 1:
|
||||
# Confirmed live: a freshly-loaded model's first structured-output
|
||||
# attempt can fail outright (no valid tool call at all) and then
|
||||
# behave normally on the very next call. Retrying with the exact
|
||||
# same request reproduces the same failure if the model is
|
||||
# genuinely stuck rather than just unlucky, so force a new seed
|
||||
# (pydantic-ai maps ModelSettings["seed"] to the OpenAI API's
|
||||
# top-level "seed" param, which works against both Ollama's and
|
||||
# llama-server's OpenAI-compatible endpoints) and give it a beat
|
||||
# via RETRY_BACKOFF_SECS in case it's still finishing loading.
|
||||
attempt_settings["seed"] = next_seed(options, attempt)
|
||||
try:
|
||||
result = await agent.run(
|
||||
prompt, message_history=history, model_settings=model_settings
|
||||
prompt,
|
||||
message_history=history,
|
||||
model_settings=cast(ModelSettings, attempt_settings)
|
||||
if attempt_settings
|
||||
else None,
|
||||
)
|
||||
# agent's output_type is the caller's `schema` (a runtime value,
|
||||
# not a static type parameter), so the checker can't narrow
|
||||
@@ -135,6 +153,8 @@ async def chat_structured(
|
||||
except _STRUCTURED_OUTPUT_FAILURE_EXCEPTIONS as exc:
|
||||
last_error = exc
|
||||
last_invalid_text = str(exc)
|
||||
if attempt < total_attempts:
|
||||
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
||||
|
||||
raise RuntimeError(
|
||||
f"chat_structured: response failed validation against schema after "
|
||||
|
||||
@@ -13,12 +13,14 @@ Deployment prerequisite: llama-server must be launched with --models-dir or
|
||||
exist otherwise (spec.md FR-006).
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from .ollama_provider import _TTLLRUCache, _cache_key, _get_json, _post_json
|
||||
from .provider import Message, ModelInfo, ModelStatus
|
||||
from .retry import RETRY_BACKOFF_SECS, next_seed
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -156,43 +158,70 @@ class LlamaCppProvider:
|
||||
messages: list[Message],
|
||||
options: dict | None = None,
|
||||
timeout_secs: float = 300.0,
|
||||
max_retries: int = 2,
|
||||
) -> str:
|
||||
payload_messages = [m.model_dump() for m in messages]
|
||||
payload: dict = {"model": model, "messages": payload_messages, "stream": False}
|
||||
if options:
|
||||
# Passed through verbatim, same nesting OllamaProvider.chat() uses
|
||||
# (payload["options"] = options) — the OllamaOption* nodes emit
|
||||
# Ollama-native parameter names (num_predict, repeat_penalty,
|
||||
# ...), which llama-server's OpenAI-compatible endpoint won't
|
||||
# recognize either way; translating them is out of scope for
|
||||
# this epic (plan.md Non-goals — no changes to the generic
|
||||
# nodes). This keeps the two providers' handling consistent
|
||||
# rather than silently special-casing one of them.
|
||||
payload["options"] = options
|
||||
total_attempts = max(0, min(int(max_retries), 5)) + 1
|
||||
response_text = ""
|
||||
|
||||
cache_key = _cache_key(
|
||||
"llamacpp_chat",
|
||||
self.host,
|
||||
self.headers or {},
|
||||
model,
|
||||
payload_messages,
|
||||
options or {},
|
||||
)
|
||||
cached, hit = _CHAT_RESPONSE_CACHE.get(cache_key)
|
||||
if hit:
|
||||
return cached
|
||||
for attempt in range(1, total_attempts + 1):
|
||||
payload: dict = {
|
||||
"model": model,
|
||||
"messages": payload_messages,
|
||||
"stream": False,
|
||||
}
|
||||
if options:
|
||||
# Passed through verbatim, same nesting OllamaProvider.chat()
|
||||
# uses (payload["options"] = options) — the OllamaOption*
|
||||
# nodes emit Ollama-native parameter names (num_predict,
|
||||
# repeat_penalty, ...), which llama-server's OpenAI-compatible
|
||||
# endpoint won't recognize either way; translating them is
|
||||
# out of scope for this epic (plan.md Non-goals — no changes
|
||||
# to the generic nodes). This keeps the two providers'
|
||||
# handling consistent rather than silently special-casing
|
||||
# one of them.
|
||||
payload["options"] = options
|
||||
if attempt > 1:
|
||||
# Unlike the options-passthrough above, this IS the OpenAI
|
||||
# spec's actual top-level "seed" field, so it takes effect
|
||||
# against llama-server's /v1/chat/completions.
|
||||
payload["seed"] = next_seed(options, attempt)
|
||||
|
||||
result = await _post_json(
|
||||
f"{self.host}/v1/chat/completions",
|
||||
payload,
|
||||
timeout=timeout_secs,
|
||||
headers=self.headers,
|
||||
)
|
||||
choices = result.get("choices") or []
|
||||
response_text = (
|
||||
choices[0].get("message", {}).get("content", "") or "" if choices else ""
|
||||
)
|
||||
_CHAT_RESPONSE_CACHE.set(cache_key, response_text)
|
||||
cache_key = _cache_key(
|
||||
"llamacpp_chat",
|
||||
self.host,
|
||||
self.headers or {},
|
||||
model,
|
||||
payload_messages,
|
||||
options or {},
|
||||
payload.get("seed"),
|
||||
)
|
||||
cached, hit = _CHAT_RESPONSE_CACHE.get(cache_key)
|
||||
if hit:
|
||||
return cached
|
||||
|
||||
result = await _post_json(
|
||||
f"{self.host}/v1/chat/completions",
|
||||
payload,
|
||||
timeout=timeout_secs,
|
||||
headers=self.headers,
|
||||
)
|
||||
choices = result.get("choices") or []
|
||||
response_text = (
|
||||
choices[0].get("message", {}).get("content", "") or ""
|
||||
if choices
|
||||
else ""
|
||||
)
|
||||
if response_text.strip():
|
||||
_CHAT_RESPONSE_CACHE.set(cache_key, response_text)
|
||||
return response_text
|
||||
|
||||
if attempt < total_attempts:
|
||||
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
||||
|
||||
# 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.
|
||||
return response_text
|
||||
|
||||
async def chat_structured(
|
||||
|
||||
@@ -19,6 +19,7 @@ import time
|
||||
from pydantic import BaseModel
|
||||
|
||||
from .provider import Message, ModelInfo, ModelStatus
|
||||
from .retry import RETRY_BACKOFF_SECS, next_seed
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -274,29 +275,54 @@ class OllamaProvider:
|
||||
messages: list[Message],
|
||||
options: dict | None = None,
|
||||
timeout_secs: float = 300.0,
|
||||
max_retries: int = 2,
|
||||
) -> str:
|
||||
payload_messages = [m.model_dump() for m in messages]
|
||||
payload: dict = {"model": model, "messages": payload_messages, "stream": False}
|
||||
if options:
|
||||
payload["options"] = options
|
||||
total_attempts = max(0, min(int(max_retries), 5)) + 1
|
||||
response_text = ""
|
||||
|
||||
cache_key = _cache_key(
|
||||
"chat",
|
||||
self.host,
|
||||
self.headers or {},
|
||||
model,
|
||||
payload_messages,
|
||||
options or {},
|
||||
)
|
||||
cached, hit = _CHAT_RESPONSE_CACHE.get(cache_key)
|
||||
if hit:
|
||||
return cached
|
||||
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)
|
||||
|
||||
result = await _post_json(
|
||||
f"{self.host}/api/chat", payload, timeout=timeout_secs, headers=self.headers
|
||||
)
|
||||
response_text = result.get("message", {}).get("content", "")
|
||||
_CHAT_RESPONSE_CACHE.set(cache_key, response_text)
|
||||
payload: dict = {
|
||||
"model": model,
|
||||
"messages": payload_messages,
|
||||
"stream": False,
|
||||
}
|
||||
if attempt_options:
|
||||
payload["options"] = attempt_options
|
||||
|
||||
cache_key = _cache_key(
|
||||
"chat",
|
||||
self.host,
|
||||
self.headers or {},
|
||||
model,
|
||||
payload_messages,
|
||||
attempt_options,
|
||||
)
|
||||
cached, hit = _CHAT_RESPONSE_CACHE.get(cache_key)
|
||||
if hit:
|
||||
return cached
|
||||
|
||||
result = await _post_json(
|
||||
f"{self.host}/api/chat",
|
||||
payload,
|
||||
timeout=timeout_secs,
|
||||
headers=self.headers,
|
||||
)
|
||||
response_text = result.get("message", {}).get("content", "")
|
||||
if response_text.strip():
|
||||
_CHAT_RESPONSE_CACHE.set(cache_key, response_text)
|
||||
return response_text
|
||||
|
||||
if attempt < total_attempts:
|
||||
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
||||
|
||||
# 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.
|
||||
return response_text
|
||||
|
||||
async def chat_structured(
|
||||
|
||||
@@ -71,8 +71,18 @@ class LLMProvider(Protocol):
|
||||
messages: list[Message],
|
||||
options: dict | None = None,
|
||||
timeout_secs: float = 300.0,
|
||||
max_retries: int = 2,
|
||||
) -> str:
|
||||
"""Free-text chat response."""
|
||||
"""Free-text chat response.
|
||||
|
||||
Retries up to ``max_retries`` times (clamped 0-5) with a new seed if
|
||||
the response comes back blank — confirmed live on a freshly-loaded
|
||||
model, whose first response is sometimes empty before it settles
|
||||
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()``.
|
||||
"""
|
||||
...
|
||||
|
||||
async def chat_structured(
|
||||
|
||||
@@ -0,0 +1,34 @@
|
||||
"""Shared retry-on-empty-output helpers for chat()/chat_structured().
|
||||
|
||||
Both providers' chat() calls (ADR-007) and the shared chat_structured()
|
||||
helper (_llm/chat.py) hit the same class of failure, confirmed live against
|
||||
a freshly-started Ollama instance on a fresh runpod: the model's first
|
||||
response after loading is sometimes blank or fails structured-output
|
||||
validation outright, then behaves normally on the very next call. Centralized
|
||||
here so both providers and both chat modes retry the same way rather than
|
||||
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").
|
||||
"""
|
||||
|
||||
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
|
||||
back-to-back."""
|
||||
|
||||
|
||||
def next_seed(options: dict | None, attempt: int) -> int:
|
||||
"""Deterministic seed for retry ``attempt`` (1-indexed).
|
||||
|
||||
Attempt 1 is the caller's original request and is never touched by this
|
||||
function — callers only call it for attempt >= 2. Starts from
|
||||
``options["seed"]`` if the caller pinned one, else 0, and increments by
|
||||
``attempt - 1`` so each retry is a new, reproducible value instead of
|
||||
repeating the exact same request that just failed.
|
||||
"""
|
||||
base = 0
|
||||
if options and isinstance(options.get("seed"), int):
|
||||
base = options["seed"]
|
||||
return base + (attempt - 1)
|
||||
@@ -548,6 +548,7 @@ class ChatCompletion:
|
||||
messages,
|
||||
llm_options,
|
||||
timeout_secs=float(timeout_secs),
|
||||
max_retries=max_retries,
|
||||
)
|
||||
)
|
||||
else:
|
||||
|
||||
@@ -39,7 +39,9 @@ class _FakeProvider:
|
||||
async def unload_model(self, model):
|
||||
self.calls.append(("unload_model", model))
|
||||
|
||||
async def chat(self, model, messages, options=None, timeout_secs=300.0):
|
||||
async def chat(
|
||||
self, model, messages, options=None, timeout_secs=300.0, max_retries=2
|
||||
):
|
||||
self.calls.append(("chat", model))
|
||||
return self.chat_response
|
||||
|
||||
|
||||
@@ -276,7 +276,7 @@ def test_chat_no_choices_returns_empty_string(monkeypatch):
|
||||
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
|
||||
result = _run_async(
|
||||
LlamaCppProvider("http://localhost:8080").chat(
|
||||
"gemma-3-4b", [Message(role="user", content="hi")]
|
||||
"gemma-3-4b", [Message(role="user", content="hi")], max_retries=0
|
||||
)
|
||||
)
|
||||
|
||||
@@ -301,6 +301,83 @@ def test_chat_second_identical_call_is_cached(monkeypatch):
|
||||
assert calls["n"] == 1
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# chat — retry-on-blank-output. Mirrors test_ollama_provider.py's coverage;
|
||||
# the one llama.cpp-specific detail is that the retry seed must land in the
|
||||
# top-level OpenAI "seed" field, not nested under "options" (see chat()'s
|
||||
# comment on why the options passthrough doesn't reach llama-server at all).
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def _fake_sleep(_secs):
|
||||
"""No-op stand-in for asyncio.sleep — keeps retry tests instant."""
|
||||
|
||||
|
||||
def test_chat_retries_on_blank_response_and_returns_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": ""}}]}
|
||||
return {"choices": [{"message": {"content": "real 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")]
|
||||
)
|
||||
)
|
||||
|
||||
assert result == "real answer"
|
||||
assert len(calls) == 2
|
||||
|
||||
|
||||
def test_chat_retry_seed_is_top_level_not_nested_in_options(monkeypatch):
|
||||
calls = []
|
||||
|
||||
async def fake_post(url, payload, *, timeout=120.0, headers=None):
|
||||
calls.append(payload)
|
||||
if len(calls) < 3:
|
||||
return {"choices": [{"message": {"content": ""}}]}
|
||||
return {"choices": [{"message": {"content": "ok"}}]}
|
||||
|
||||
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")], max_retries=2
|
||||
)
|
||||
)
|
||||
|
||||
assert "seed" not in calls[0]
|
||||
assert calls[1]["seed"] == 1
|
||||
assert calls[2]["seed"] == 2
|
||||
|
||||
|
||||
def test_chat_exhausted_retries_returns_blank_without_raising(monkeypatch):
|
||||
calls = {"n": 0}
|
||||
|
||||
async def fake_post(url, payload, *, timeout=120.0, headers=None):
|
||||
calls["n"] += 1
|
||||
return {"choices": [{"message": {"content": ""}}]}
|
||||
|
||||
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")], max_retries=2
|
||||
)
|
||||
)
|
||||
|
||||
assert result == ""
|
||||
assert calls["n"] == 3 # original + 2 retries, per max_retries=2
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# chat_structured — zero new logic, delegates to the shared helper unchanged
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -23,6 +23,18 @@ from comfydv._llm.ollama_provider import _run_async
|
||||
from comfydv._llm.provider import Message
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _no_retry_backoff(monkeypatch):
|
||||
"""Keep the real RETRY_BACKOFF_SECS delay out of this suite's wall-clock
|
||||
time for every test except the ones that specifically assert on it
|
||||
(which re-monkeypatch locally, overriding this)."""
|
||||
|
||||
async def _instant_sleep(_secs):
|
||||
pass
|
||||
|
||||
monkeypatch.setattr(chat_mod.asyncio, "sleep", _instant_sleep)
|
||||
|
||||
|
||||
class _Widget(BaseModel):
|
||||
name: str
|
||||
count: int
|
||||
@@ -187,6 +199,79 @@ def test_chat_structured_requires_last_message_user_role():
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# retry-on-failure seed/backoff — live-verified: a freshly-loaded model's
|
||||
# first structured-output attempt can fail outright, then behave normally on
|
||||
# the very next call.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_chat_structured_retry_injects_incrementing_seed(monkeypatch):
|
||||
bad = ValidationError.from_exception_data("Widget", [])
|
||||
fake = _FakeAgent([bad, bad, _Widget(name="c", count=3)])
|
||||
monkeypatch.setattr(chat_mod, "_build_agent", lambda **kw: fake)
|
||||
|
||||
_run_async(
|
||||
chat_mod.chat_structured(
|
||||
base_url="http://localhost:11434/v1",
|
||||
model="llama3",
|
||||
messages=_messages(),
|
||||
schema=_Widget,
|
||||
max_retries=2,
|
||||
)
|
||||
)
|
||||
|
||||
assert fake.calls[0][2] is None # attempt 1 untouched — no options set
|
||||
assert fake.calls[1][2]["seed"] == 1
|
||||
assert fake.calls[2][2]["seed"] == 2
|
||||
|
||||
|
||||
def test_chat_structured_retry_seed_starts_from_pinned_base(monkeypatch):
|
||||
bad = ValidationError.from_exception_data("Widget", [])
|
||||
fake = _FakeAgent([bad, _Widget(name="c", count=3)])
|
||||
monkeypatch.setattr(chat_mod, "_build_agent", lambda **kw: fake)
|
||||
|
||||
_run_async(
|
||||
chat_mod.chat_structured(
|
||||
base_url="http://localhost:11434/v1",
|
||||
model="llama3",
|
||||
messages=_messages(),
|
||||
schema=_Widget,
|
||||
options={"seed": 42},
|
||||
max_retries=2,
|
||||
)
|
||||
)
|
||||
|
||||
assert fake.calls[0][2] == {"extra_body": {"options": {"seed": 42}}}
|
||||
assert fake.calls[1][2]["seed"] == 43 # base(42) + (attempt 2 - 1)
|
||||
assert fake.calls[1][2]["extra_body"] == {"options": {"seed": 42}}
|
||||
|
||||
|
||||
def test_chat_structured_retry_sleeps_between_attempts(monkeypatch):
|
||||
sleep_calls = []
|
||||
|
||||
async def fake_sleep(secs):
|
||||
sleep_calls.append(secs)
|
||||
|
||||
monkeypatch.setattr(chat_mod.asyncio, "sleep", fake_sleep)
|
||||
|
||||
bad = ValidationError.from_exception_data("Widget", [])
|
||||
fake = _FakeAgent([bad, _Widget(name="c", count=3)])
|
||||
monkeypatch.setattr(chat_mod, "_build_agent", lambda **kw: fake)
|
||||
|
||||
_run_async(
|
||||
chat_mod.chat_structured(
|
||||
base_url="http://localhost:11434/v1",
|
||||
model="llama3",
|
||||
messages=_messages(),
|
||||
schema=_Widget,
|
||||
max_retries=2,
|
||||
)
|
||||
)
|
||||
|
||||
assert sleep_calls == [chat_mod.RETRY_BACKOFF_SECS]
|
||||
|
||||
|
||||
def test_history_to_messages_preserves_order_and_roles():
|
||||
from pydantic_ai.messages import ModelRequest, ModelResponse
|
||||
|
||||
|
||||
@@ -0,0 +1,24 @@
|
||||
"""Tests for comfydv._llm.retry — shared retry-on-blank-output helpers used
|
||||
by both providers' chat() and the shared chat_structured() helper.
|
||||
"""
|
||||
|
||||
from comfydv._llm.retry import next_seed
|
||||
|
||||
|
||||
def test_next_seed_attempt_one_is_zero_by_default():
|
||||
assert next_seed(None, 1) == 0
|
||||
|
||||
|
||||
def test_next_seed_increments_from_zero_when_unset():
|
||||
assert next_seed(None, 2) == 1
|
||||
assert next_seed({}, 3) == 2
|
||||
|
||||
|
||||
def test_next_seed_starts_from_pinned_base():
|
||||
assert next_seed({"seed": 42}, 1) == 42
|
||||
assert next_seed({"seed": 42}, 2) == 43
|
||||
assert next_seed({"seed": 42}, 3) == 44
|
||||
|
||||
|
||||
def test_next_seed_ignores_non_int_seed():
|
||||
assert next_seed({"seed": "not-an-int"}, 2) == 1
|
||||
@@ -89,8 +89,10 @@ class _FakeProvider:
|
||||
async def unload_model(self, model):
|
||||
self.calls.append(("unload_model", model))
|
||||
|
||||
async def chat(self, model, messages, options=None, timeout_secs=300.0):
|
||||
self.calls.append(("chat", model, messages, options, timeout_secs))
|
||||
async def chat(
|
||||
self, model, messages, options=None, timeout_secs=300.0, max_retries=2
|
||||
):
|
||||
self.calls.append(("chat", model, messages, options, timeout_secs, max_retries))
|
||||
return self.chat_response
|
||||
|
||||
async def chat_structured(
|
||||
|
||||
@@ -244,6 +244,148 @@ def test_chat_timeout_forwarded(monkeypatch):
|
||||
assert captured["timeout"] == 600.0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# chat — retry-on-blank-output (live-verified: a freshly-loaded model's first
|
||||
# response is sometimes blank on a fresh runpod, then normal afterwards)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def _fake_sleep(_secs):
|
||||
"""No-op stand-in for asyncio.sleep — keeps retry tests instant."""
|
||||
|
||||
|
||||
def test_chat_retries_on_blank_response_and_returns_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": ""}}
|
||||
return {"message": {"content": "real 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")]
|
||||
)
|
||||
)
|
||||
|
||||
assert result == "real answer"
|
||||
assert len(calls) == 2
|
||||
|
||||
|
||||
def test_chat_retry_injects_incrementing_seed_when_none_pinned(monkeypatch):
|
||||
calls = []
|
||||
|
||||
async def fake_post(url, payload, *, timeout=120.0, headers=None):
|
||||
calls.append(payload)
|
||||
if len(calls) < 3:
|
||||
return {"message": {"content": ""}}
|
||||
return {"message": {"content": "ok"}}
|
||||
|
||||
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")], max_retries=2
|
||||
)
|
||||
)
|
||||
|
||||
assert "seed" not in calls[0].get("options", {})
|
||||
assert calls[1]["options"]["seed"] == 1
|
||||
assert calls[2]["options"]["seed"] == 2
|
||||
|
||||
|
||||
def test_chat_retry_seed_starts_from_pinned_base(monkeypatch):
|
||||
calls = []
|
||||
|
||||
async def fake_post(url, payload, *, timeout=120.0, headers=None):
|
||||
calls.append(payload)
|
||||
if len(calls) == 1:
|
||||
return {"message": {"content": ""}}
|
||||
return {"message": {"content": "ok"}}
|
||||
|
||||
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")],
|
||||
options={"seed": 42},
|
||||
)
|
||||
)
|
||||
|
||||
assert calls[0]["options"]["seed"] == 42 # attempt 1 untouched
|
||||
assert calls[1]["options"]["seed"] == 43 # attempt 2 = base + 1
|
||||
|
||||
|
||||
def test_chat_exhausted_retries_returns_blank_without_raising(monkeypatch):
|
||||
calls = []
|
||||
|
||||
async def fake_post(url, payload, *, timeout=120.0, headers=None):
|
||||
calls.append(payload)
|
||||
return {"message": {"content": ""}}
|
||||
|
||||
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")], max_retries=2
|
||||
)
|
||||
)
|
||||
|
||||
assert result == ""
|
||||
assert len(calls) == 3 # original + 2 retries, per max_retries=2
|
||||
|
||||
|
||||
def test_chat_blank_response_is_not_cached(monkeypatch):
|
||||
"""A blank final result must not poison the cache — the next queue run
|
||||
should try the real backend again, not replay the blank forever (the
|
||||
cache has no TTL)."""
|
||||
calls = {"n": 0}
|
||||
|
||||
async def fake_post(url, payload, *, timeout=120.0, headers=None):
|
||||
calls["n"] += 1
|
||||
return {"message": {"content": ""}}
|
||||
|
||||
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
|
||||
monkeypatch.setattr(provider_mod.asyncio, "sleep", _fake_sleep)
|
||||
|
||||
provider = OllamaProvider("http://localhost:11434")
|
||||
messages = [Message(role="user", content="hi")]
|
||||
_run_async(provider.chat(model="llama3", messages=messages, max_retries=0))
|
||||
calls_after_first_run = calls["n"]
|
||||
_run_async(provider.chat(model="llama3", messages=messages, max_retries=0))
|
||||
|
||||
assert calls["n"] > calls_after_first_run # second run hit the network again
|
||||
|
||||
|
||||
def test_chat_no_retry_needed_does_not_sleep(monkeypatch):
|
||||
sleep_calls = []
|
||||
|
||||
async def fake_sleep(secs):
|
||||
sleep_calls.append(secs)
|
||||
|
||||
async def fake_post(url, payload, *, timeout=120.0, headers=None):
|
||||
return {"message": {"content": "first try works"}}
|
||||
|
||||
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")]
|
||||
)
|
||||
)
|
||||
|
||||
assert sleep_calls == []
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# chat_structured
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
Reference in New Issue
Block a user