fix(providers): guard chat_with_retry against explicit None max_tokens (#3102)
This commit is contained in:
@@ -512,10 +512,14 @@ class LLMProvider(ABC):
|
|||||||
on_retry_wait: Callable[[str], Awaitable[None]] | None = None,
|
on_retry_wait: Callable[[str], Awaitable[None]] | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
"""Call chat_stream() with retry on transient provider failures."""
|
"""Call chat_stream() with retry on transient provider failures."""
|
||||||
if max_tokens is self._SENTINEL:
|
if max_tokens is self._SENTINEL or max_tokens is None:
|
||||||
max_tokens = self.generation.max_tokens
|
max_tokens = self.generation.max_tokens
|
||||||
if temperature is self._SENTINEL:
|
if max_tokens is None:
|
||||||
|
max_tokens = GenerationSettings.max_tokens
|
||||||
|
if temperature is self._SENTINEL or temperature is None:
|
||||||
temperature = self.generation.temperature
|
temperature = self.generation.temperature
|
||||||
|
if temperature is None:
|
||||||
|
temperature = GenerationSettings.temperature
|
||||||
if reasoning_effort is self._SENTINEL:
|
if reasoning_effort is self._SENTINEL:
|
||||||
reasoning_effort = self.generation.reasoning_effort
|
reasoning_effort = self.generation.reasoning_effort
|
||||||
|
|
||||||
@@ -549,12 +553,20 @@ class LLMProvider(ABC):
|
|||||||
|
|
||||||
Parameters default to ``self.generation`` when not explicitly passed,
|
Parameters default to ``self.generation`` when not explicitly passed,
|
||||||
so callers no longer need to thread temperature / max_tokens /
|
so callers no longer need to thread temperature / max_tokens /
|
||||||
reasoning_effort through every layer.
|
reasoning_effort through every layer. Explicit ``None`` is also
|
||||||
|
normalized to the provider's generation defaults, with a final fallback
|
||||||
|
to :class:`GenerationSettings` class defaults so that downstream
|
||||||
|
``_build_kwargs`` never sees ``None`` for ``max_tokens`` / ``temperature``
|
||||||
|
(which would crash ``max(1, max_tokens)``).
|
||||||
"""
|
"""
|
||||||
if max_tokens is self._SENTINEL:
|
if max_tokens is self._SENTINEL or max_tokens is None:
|
||||||
max_tokens = self.generation.max_tokens
|
max_tokens = self.generation.max_tokens
|
||||||
if temperature is self._SENTINEL:
|
if max_tokens is None:
|
||||||
|
max_tokens = GenerationSettings.max_tokens
|
||||||
|
if temperature is self._SENTINEL or temperature is None:
|
||||||
temperature = self.generation.temperature
|
temperature = self.generation.temperature
|
||||||
|
if temperature is None:
|
||||||
|
temperature = GenerationSettings.temperature
|
||||||
if reasoning_effort is self._SENTINEL:
|
if reasoning_effort is self._SENTINEL:
|
||||||
reasoning_effort = self.generation.reasoning_effort
|
reasoning_effort = self.generation.reasoning_effort
|
||||||
|
|
||||||
|
|||||||
@@ -521,3 +521,59 @@ async def test_persistent_retry_emits_terminal_progress_on_identical_error_limit
|
|||||||
|
|
||||||
assert response.finish_reason == "error"
|
assert response.finish_reason == "error"
|
||||||
assert progress[-1] == "Persistent retry stopped after 10 identical errors."
|
assert progress[-1] == "Persistent retry stopped after 10 identical errors."
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_chat_with_retry_normalizes_explicit_none_max_tokens() -> None:
|
||||||
|
"""Explicit max_tokens=None must fall back to generation defaults.
|
||||||
|
|
||||||
|
Regression for #3102: callers that construct AgentRunSpec with
|
||||||
|
max_tokens=None propagate None into chat_with_retry, which used to
|
||||||
|
reach ``_build_kwargs`` and crash on ``max(1, None)``.
|
||||||
|
"""
|
||||||
|
provider = ScriptedProvider([LLMResponse(content="ok")])
|
||||||
|
|
||||||
|
response = await provider.chat_with_retry(
|
||||||
|
messages=[{"role": "user", "content": "hi"}],
|
||||||
|
max_tokens=None,
|
||||||
|
temperature=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.content == "ok"
|
||||||
|
# Generation settings default to 4096 / 0.7; explicit None should
|
||||||
|
# have been replaced before reaching chat().
|
||||||
|
assert provider.last_kwargs["max_tokens"] == 4096
|
||||||
|
assert provider.last_kwargs["temperature"] == 0.7
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_chat_with_retry_hard_fallback_when_generation_max_tokens_is_none() -> None:
|
||||||
|
"""Final fallback: even if provider.generation.max_tokens is None,
|
||||||
|
chat() must receive the GenerationSettings class default (not None)."""
|
||||||
|
provider = ScriptedProvider([LLMResponse(content="ok")])
|
||||||
|
# Bypass the dataclass type hint by constructing with None explicitly.
|
||||||
|
provider.generation = GenerationSettings(max_tokens=None, temperature=None) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
response = await provider.chat_with_retry(
|
||||||
|
messages=[{"role": "user", "content": "hi"}],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.content == "ok"
|
||||||
|
assert provider.last_kwargs["max_tokens"] == GenerationSettings.max_tokens
|
||||||
|
assert provider.last_kwargs["temperature"] == GenerationSettings.temperature
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_chat_stream_with_retry_normalizes_explicit_none_max_tokens() -> None:
|
||||||
|
"""chat_stream_with_retry must apply the same None-guard as chat_with_retry."""
|
||||||
|
provider = ScriptedProvider([LLMResponse(content="ok")])
|
||||||
|
|
||||||
|
response = await provider.chat_stream_with_retry(
|
||||||
|
messages=[{"role": "user", "content": "hi"}],
|
||||||
|
max_tokens=None,
|
||||||
|
temperature=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.content == "ok"
|
||||||
|
assert provider.last_kwargs["max_tokens"] == 4096
|
||||||
|
assert provider.last_kwargs["temperature"] == 0.7
|
||||||
|
|||||||
Reference in New Issue
Block a user