Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 32 additions & 2 deletions agent_core/runtime/loop/_runaway.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@

from agent_core.llm import LLMResponse
from agent_core.messages import Message
from agent_core.model_capabilities import ModelCapabilities
from agent_core.runtime.env import first_configured
from agent_core.runtime.llm_request_overrides import ThinkingRetryOverride
from agent_core.tokens import estimate_message_tokens
Expand Down Expand Up @@ -258,15 +259,18 @@ def _expanded_retry_max_tokens(
active_cap: Any,
messages: list[Message],
context_token_limit_hint: int | None,
max_output_tokens_limit: int | None = None,
) -> int | None:
"""Return a safe 1.5脳 cap, or ``None`` when context cannot hold it."""
"""Return a safe 1.5脳 cap within context and model output limits."""
try:
cap = int(active_cap)
except (TypeError, ValueError):
return None
if cap <= 0:
return None
expanded = int(cap * _RUNAWAY_EXPAND_FACTOR)
if max_output_tokens_limit is not None:
expanded = min(expanded, max_output_tokens_limit)
if context_token_limit_hint:
available = max(
int(context_token_limit_hint)
Expand All @@ -278,19 +282,45 @@ def _expanded_retry_max_tokens(
return expanded if expanded > cap else None


def _model_output_limit(client: Any, _depth: int = 0) -> int | None:
"""Return the largest ``max_tokens`` every model behind ``client`` accepts.

A plain client (or a middleware proxy forwarding to one) reports its own
``capabilities``. A fallback chain binds one ``max_tokens`` for whichever
leg serves the request, so it is limited by its *smallest* known leg;
legs with no declared limit do not constrain it.
"""
if client is None or _depth > 8:
return None
capabilities = getattr(client, "capabilities", None)
if isinstance(capabilities, ModelCapabilities):
return capabilities.max_output_tokens
entries = getattr(client, "entries", None)
if not isinstance(entries, list):
return None
limits = [
limit
for entry in entries
if (limit := _model_output_limit(getattr(entry, "model", None), _depth + 1)) is not None
]
return min(limits) if limits else None


def _bind_expanded_max_tokens(
llm: Any,
*,
active_cap: Any,
messages: list[Message],
context_token_limit_hint: int | None,
) -> Any | None:
bound = _ensure_bound(llm)
expanded = _expanded_retry_max_tokens(
active_cap, messages, context_token_limit_hint,
_model_output_limit(bound.client),
)
if expanded is None:
return None
return replace(_ensure_bound(llm), max_tokens=expanded)
return replace(bound, max_tokens=expanded)


def _phase_reasoning_guard(
Expand Down
1 change: 1 addition & 0 deletions changes/74.fix.md
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Runaway retry no longer expands the output cap past the model's limit when the client is driven through an `LLMFallbackChain`. The chain binds one `max_tokens` for whichever leg serves the request, so expansion is now clamped to the smallest declared output limit among its legs instead of reading a `capabilities` attribute the chain does not expose; a chain whose legs declare no limit keeps the previous context-bounded behaviour.
73 changes: 73 additions & 0 deletions tests/test_runaway_retry_transient.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
from __future__ import annotations

from copy import deepcopy
from dataclasses import replace
from types import SimpleNamespace
from typing import Any
from unittest.mock import AsyncMock
Expand All @@ -11,8 +12,80 @@

from agent_core.messages import Message, assistant_msg, for_wire, system_msg, tool_msg, user_msg
from agent_core.providers.anthropic import AnthropicClient
from agent_core.providers.fallback import FallbackEntry, LLMFallbackChain
from agent_core.runtime.loop import _call
from agent_core.runtime.loop._bind import bind_max_tokens
from agent_core.runtime.loop._runaway import (
_bind_expanded_max_tokens,
_expanded_retry_max_tokens,
_model_output_limit,
)


@pytest.mark.parametrize(("active_cap", "expected"), [
(100_000, 128_000),
(128_000, None),
])
def test_expansion_respects_model_output_limit(active_cap: int, expected: int | None) -> None:
assert _expanded_retry_max_tokens(
active_cap, [user_msg("question")], 1_048_576, 128_000,
) == expected


def _opus(max_tokens: int = 128_000) -> AnthropicClient:
return AnthropicClient("claude-opus-5-5", api_key="test-key", max_tokens=max_tokens)


def test_expansion_through_fallback_chain_respects_smallest_leg_limit() -> None:
small = _opus(64_000)
small.capabilities = replace(small.capabilities, max_output_tokens=64_000)
chain = LLMFallbackChain(entries=[FallbackEntry(model=_opus(48_000)), FallbackEntry(model=small)])
bound = _bind_expanded_max_tokens(
chain, active_cap=48_000, messages=[user_msg("question")],
context_token_limit_hint=1_048_576,
)
assert bound is not None and bound.max_tokens == 64_000


def test_fallback_chain_at_its_output_limit_skips_expansion() -> None:
chain = LLMFallbackChain(entries=[FallbackEntry(model=_opus()), FallbackEntry(model=_opus())])
assert _bind_expanded_max_tokens(
chain, active_cap=128_000, messages=[user_msg("question")],
context_token_limit_hint=1_048_576,
) is None


def test_unknown_leg_limits_leave_expansion_to_context() -> None:
chain = LLMFallbackChain(entries=[FallbackEntry(model=SimpleNamespace(model="custom"))])
assert _model_output_limit(chain) is None


@pytest.mark.asyncio
async def test_runaway_at_model_output_limit_skips_invalid_expansion(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(_call, "_RUNAWAY_BACKOFF_S", 0.0)
client = AnthropicClient("claude-opus-5-5", api_key="test-key", max_tokens=128_000)
wire_caps: list[int] = []

async def create(**kwargs: Any) -> Any:
wire_caps.append(kwargs["max_tokens"])
done = len(wire_caps) == 2
return SimpleNamespace(
content=[SimpleNamespace(type="text", text="done")] if done else [],
stop_reason="end_turn" if done else "max_tokens",
model="claude-opus-5-5", id=f"response-{len(wire_caps)}",
usage=SimpleNamespace(input_tokens=10, output_tokens=2 if done else 128_000),
)

client._client = SimpleNamespace(messages=SimpleNamespace(create=AsyncMock(side_effect=create)))
response = await _call.call_llm(
client, [user_msg("question")], timeout=30, max_retries=3, turn=1,
max_completion_tokens_hint=128_000, context_token_limit_hint=1_048_576,
)

assert response is not None and response.content == "done"
assert wire_caps == [128_000, 8_192]


@pytest.mark.parametrize("retry_path", [
Expand Down
Loading