From 09a173e552a769060178d2d58273c363dca24b8d Mon Sep 17 00:00:00 2001 From: James Veitch <1722315+darth-veitcher@users.noreply.github.com> Date: Mon, 20 Jul 2026 08:49:35 +0100 Subject: [PATCH] fix: retry blank/failed LLM responses with a new seed and backoff MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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 Claude-Session: https://claude.ai/code/session_0132ojafeazQ3ephcBejEWFj --- src/comfydv/_llm/chat.py | 24 ++++- src/comfydv/_llm/llamacpp_provider.py | 95 +++++++++++------ src/comfydv/_llm/ollama_provider.py | 64 ++++++++---- src/comfydv/_llm/provider.py | 12 ++- src/comfydv/_llm/retry.py | 34 ++++++ src/comfydv/ollama.py | 1 + tests/test_llamacpp.py | 4 +- tests/test_llamacpp_provider.py | 79 +++++++++++++- tests/test_llm_chat_structured.py | 85 +++++++++++++++ tests/test_llm_retry.py | 24 +++++ tests/test_ollama.py | 6 +- tests/test_ollama_provider.py | 142 ++++++++++++++++++++++++++ 12 files changed, 511 insertions(+), 59 deletions(-) create mode 100644 src/comfydv/_llm/retry.py create mode 100644 tests/test_llm_retry.py diff --git a/src/comfydv/_llm/chat.py b/src/comfydv/_llm/chat.py index 367f00c..b45283c 100644 --- a/src/comfydv/_llm/chat.py +++ b/src/comfydv/_llm/chat.py @@ -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 " diff --git a/src/comfydv/_llm/llamacpp_provider.py b/src/comfydv/_llm/llamacpp_provider.py index ee90594..dc55e4a 100644 --- a/src/comfydv/_llm/llamacpp_provider.py +++ b/src/comfydv/_llm/llamacpp_provider.py @@ -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( diff --git a/src/comfydv/_llm/ollama_provider.py b/src/comfydv/_llm/ollama_provider.py index aaa5efd..09cb38e 100644 --- a/src/comfydv/_llm/ollama_provider.py +++ b/src/comfydv/_llm/ollama_provider.py @@ -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( diff --git a/src/comfydv/_llm/provider.py b/src/comfydv/_llm/provider.py index 02b3e99..353e594 100644 --- a/src/comfydv/_llm/provider.py +++ b/src/comfydv/_llm/provider.py @@ -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( diff --git a/src/comfydv/_llm/retry.py b/src/comfydv/_llm/retry.py new file mode 100644 index 0000000..afab55b --- /dev/null +++ b/src/comfydv/_llm/retry.py @@ -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) diff --git a/src/comfydv/ollama.py b/src/comfydv/ollama.py index 3fa7008..f010553 100644 --- a/src/comfydv/ollama.py +++ b/src/comfydv/ollama.py @@ -548,6 +548,7 @@ class ChatCompletion: messages, llm_options, timeout_secs=float(timeout_secs), + max_retries=max_retries, ) ) else: diff --git a/tests/test_llamacpp.py b/tests/test_llamacpp.py index 8761bd4..f6507dc 100644 --- a/tests/test_llamacpp.py +++ b/tests/test_llamacpp.py @@ -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 diff --git a/tests/test_llamacpp_provider.py b/tests/test_llamacpp_provider.py index 3030d04..89e57ae 100644 --- a/tests/test_llamacpp_provider.py +++ b/tests/test_llamacpp_provider.py @@ -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 # --------------------------------------------------------------------------- diff --git a/tests/test_llm_chat_structured.py b/tests/test_llm_chat_structured.py index 85a5ceb..b9f2ce2 100644 --- a/tests/test_llm_chat_structured.py +++ b/tests/test_llm_chat_structured.py @@ -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 diff --git a/tests/test_llm_retry.py b/tests/test_llm_retry.py new file mode 100644 index 0000000..7b3ecf0 --- /dev/null +++ b/tests/test_llm_retry.py @@ -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 diff --git a/tests/test_ollama.py b/tests/test_ollama.py index ef03e62..8577f4f 100644 --- a/tests/test_ollama.py +++ b/tests/test_ollama.py @@ -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( diff --git a/tests/test_ollama_provider.py b/tests/test_ollama_provider.py index f14ee46..03a07b4 100644 --- a/tests/test_ollama_provider.py +++ b/tests/test_ollama_provider.py @@ -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 # ---------------------------------------------------------------------------