fix(router): retry a /v1/messages stream the provider drops before the first content chunk (#44276)

* fix(router): retry a /v1/messages stream the provider drops before the first content chunk

A /v1/messages stream that the upstream closed before any content reached the
client answered an error event after a single attempt, so the router's
num_retries never applied to that drop. The pre-content failure is now retried
within the model group before the fallback chain runs, with the budget resolved
the way a failure raised before the stream opened resolves it: a retry policy
that names the error class, then the request's num_retries, then the
deployment's, then the router's. A drop after content reached the client keeps
surfacing the provider's error after one attempt.

Fixes #44238

* fix(router): hand a retry's non-retriable error to the fallback chain and type the retry helpers

A retry that failed before its stream opened with an error no retry covers raised straight to the
client, skipping a fallback the first attempt would have used. assert_never now comes from
typing_extensions so the router imports on Python 3.10, and the retry helpers read their kwargs
through typed narrowing instead of Mapping[str, Any]

* fix(router): cast the untyped router fallback defaults the stream retry gate reads

The retry gate passed the router's fallback attributes, declared without element types, to the
typed request override helper, which basedpyright counted as new unknown-argument errors

* fix(router): consult context_window_fallbacks when a retried /v1/messages stream overflows

A retry attempt raising ContextWindowExceededError reached the fallback chain inside its
mid-stream envelope, so only the regular fallbacks list matched. The fallback attempt now
unwraps it the way it unwraps a content policy error. The new router helpers are covered for
the router code coverage check with two direct-call tests and named covering tests

* fix(router): retry a 408 raised by a /v1/messages retry and honor deployment num_retries before the stream opens

* fix(router): attribute a retried /v1/messages stream to the deployment that served it and bound the retry-policy hold

* fix(router): retry /v1/messages error frames under their retry-policy class and keep the first drop's committed budget

An `event: error` frame that arrives before the first content delta now raises the exception class the pre-stream mapping gives an HTTP answer with the same status (429 RateLimitError, 500 and 529 InternalServerError, 503 ServiceUnavailableError, 504 Timeout), so a retry policy's per-class budget governs it the way it governs the error before the stream opened. The status the client sees is unchanged

A retry that lands on a sibling deployment keeps the budget the first drop committed to, read back from the request's attempted_retries and max_retries, instead of recomputing it from the new deployment's num_retries, matching the pre-stream retry loop

* refactor(anthropic): keep the error-frame exception mapping under llms and type the retry test helper

The status-to-exception mapping an `event: error` frame gets before the retry policy is consulted now lives next to the Anthropic error status map in llms/anthropic/common_utils.py, with its own unit test, and the two-deployment retry test helper takes explicit typed parameters instead of a bare dict and untyped kwargs

* refactor(anthropic): map an error frame's status with explicit returns on every path

* fix(router): map stream error frames through the pre-stream exception mapping

An overloaded `event: error` frame on a /v1/messages stream now raises the InternalServerError a 529 answer maps to, built by exception_type from the frame's own body, so one retry policy class governs the error before and after the first byte; a failed fallback after such a frame answers 500 like every other litellm path instead of the frame map's 503

A model_group_retry_policy that does not parse (a non-integer budget, an entry that is not a mapping) no longer fails every healthy stream of that group before its first attempt: the stream runs with no policy and the plain num_retries budget, with a warning naming the group

* fix(router): forward an error frame nothing can take over for as the provider sent it

A pre-content error frame whose class the retry policy grants no retry, with no fallback configured, raised an HTTP error only on the first attempt while the same frame after exhausted retries reached the client verbatim. Both now pass through as sent, the way the merge base forwarded every frame.

* test(integration): audit /v1/messages pre-content retry across routes and budgets

Adds the /audit cells for the pre-content stream retry: the native Anthropic route
(drops and error frames before content, HTTP rejections before the stream opens, SDK
sync and async, after-content and non-retriable controls, budget exhaustion, cache
twin, spend row and headers), the chat and responses bridges, the generic routes
(responses, chat, vllm pass-through, Gemini generateContent, fine-tuning jobs list),
owned two-worker proxies for router-level budgets, retry policies and fallbacks, and
two chaos cells (a worker killed mid burst, an outage on every first attempt). Shared
helpers for scripted Anthropic SSE upstreams and OpenAI-compatible wire replies live
in tests/integration/_support

---------

Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-05 22:53:39 +00:00 • committed by GitHub
parent 7bad8de067
commit 68d9b8bbb8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
15 changed files with 3212 additions and 61 deletions

View file

@ -3,6 +3,7 @@ This file contains common utils for anthropic calls.
"""
import copy
import json
import re
from collections.abc import Mapping, MutableMapping, Sequence
from datetime import datetime, timezone
@ -82,6 +83,23 @@ ANTHROPIC_ERROR_STATUS_CODE_MAP: Final = MappingProxyType(
}
)
def anthropic_error_frame_exception(error_type: str, message: str, status_code: int, model: str) -> Exception:
"""The exception the pre-stream mapping raises for an HTTP answer carrying this frame's body and status, so a
retry policy's per-class budget governs an `event: error` frame the way it governs the same error before the
stream opened: an overloaded frame is the InternalServerError a real 529 answer is, whatever status the frame
map gives it."""
from litellm.litellm_core_utils.exception_mapping_utils import exception_type
frame_body: Final = json.dumps({"type": "error", "error": {"type": error_type, "message": message}})
frame_error: Final = AnthropicError(status_code=status_code, message=frame_body)
try:
exception_type(model=model, original_exception=frame_error, custom_llm_provider="anthropic")
except Exception as raised: # noqa: BLE001 # exception_type hands the mapped error back by raising it
return raised
return frame_error
_BEDROCK_VERSION_SUFFIX_RE: Final = re.compile(r"-v\d+(?::\d+)?$")
_INFERENCE_PROFILE_MINOR_RE: Final = re.compile(r":\d+$")
_DATED_RELEASE_SUFFIX_RE: Final = re.compile(r"-\d{8}$")

View file

@ -31,6 +31,7 @@ from collections.abc import (
MutableMapping,
Sequence,
)
from dataclasses import dataclass
from datetime import datetime, timezone
from functools import lru_cache, partial
from types import MappingProxyType
@ -41,7 +42,7 @@ import httpx
import openai
from openai import AsyncOpenAI
from pydantic import BaseModel, TypeAdapter, ValidationError
from typing_extensions import overload
from typing_extensions import assert_never, overload
import litellm
import litellm.litellm_core_utils.exception_mapping_utils
@ -108,7 +109,7 @@ from litellm.litellm_core_utils.sensitive_data_masker import (
mask_sensitive_structure,
)
from litellm.litellm_core_utils.token_counter import offload_token_count
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
from litellm.llms.anthropic.common_utils import AnthropicModelInfo, anthropic_error_frame_exception
from litellm.llms.base_llm.passthrough.transformation import replace_path_segment
from litellm.llms.base_llm.vector_store.transformation import (
RouterVectorStoreEmbeddingExecutor,
@ -206,22 +207,29 @@ from litellm.router_utils.fallback_event_handlers import (
MID_STREAM_FALLBACK_CONTROLS_KEY,
AttemptedFallbackTargets,
_check_non_standard_fallback_format,
attempted_retries_for_request,
carry_over_pre_routing_selection,
carry_over_routed_deployment,
clear_pre_routing_selection,
committed_retry_budget_for_request,
fallback_lookup_groups,
fallbacks_disabled_for_request,
get_fallback_model_group_for_lookup_groups,
get_pre_routing_selection,
has_unattempted_fallback_target,
mid_stream_fallback_hop_kwargs,
mid_stream_retry_kwargs,
per_request_fallback_controls,
record_disable_fallbacks,
record_pre_routing_selection,
record_retry_attempt,
routed_deployment_id,
run_async_fallback,
)
from litellm.router_utils.get_retry_from_policy import (
get_num_retries_from_retry_policy as _get_num_retries_from_retry_policy,
)
from litellm.router_utils.get_retry_from_policy import resolve_retry_policy
from litellm.router_utils.handle_error import (
async_raise_no_deployment_exception,
send_llm_exception_alert,
@ -479,6 +487,7 @@ _NO_SESSION_KWARGS: Final[Mapping[str, Mapping[str, object]]] = MappingProxyType
_SESSION_ADAPTER: Final = TypeAdapter(Mapping[str, object])
_MODEL_INFO_ADAPTER: Final = TypeAdapter(Mapping[str, object])
_SILENT_MODEL_ADAPTER: Final = TypeAdapter(str | list[str])
_RESOLVED_RETRY_POLICY_ADAPTER: Final = TypeAdapter(RetryPolicy | None)
_ROUTING_KWARGS_ADAPTER: Final[TypeAdapter[Mapping[str, object] | None]] = TypeAdapter(Mapping[str, object] | None)
_DEPLOYMENT_SELECTED_EVENT: Final = "litellm.request.deployment_selected"
@ -606,14 +615,20 @@ def _anthropic_stream_raised_error_status(error: Exception) -> int | None:
def _anthropic_stream_fallback_error_for_raised(
error: Exception, model: str, has_generated_content: bool
) -> "MidStreamFallbackError | None":
"""Same gate as a detected SSE error event; None means the raise propagates unchanged."""
from litellm.exceptions import MidStreamFallbackError
"""The pre-stream retry rule (408, 409, 429, 5xx); None means the raise propagates unchanged."""
if has_generated_content:
return None
status_code: Final = _anthropic_stream_raised_error_status(error)
if status_code is not None and not _is_retriable_anthropic_status(status_code):
return None
if status_code is None:
return _anthropic_stream_pre_content_error(error, model)
retriable: Final = litellm._should_retry(status_code) # pyright: ignore[reportPrivateUsage] # shared retry rule
return _anthropic_stream_pre_content_error(error, model) if retriable else None
def _anthropic_stream_pre_content_error(error: Exception, model: str) -> "MidStreamFallbackError":
"""The envelope the fallback chain judges a failure by when the client has received no content yet."""
from litellm.exceptions import MidStreamFallbackError
return MidStreamFallbackError(
message=str(error),
model=model,
@ -623,6 +638,30 @@ def _anthropic_stream_fallback_error_for_raised(
)
def _deployment_num_retries(deployment: "Deployment | None") -> int | None:
"""The deployment's own num_retries litellm_param, an int or a digit string the way the config loader leaves it."""
configured: Final = getattr(deployment.litellm_params, "num_retries", None) if deployment is not None else None
if isinstance(configured, bool) or not isinstance(configured, (int, str)):
return None
return int(configured) if str(configured).isdigit() else None
def _request_fallback_list(
kwargs: Mapping[str, object], key: str, router_default: "list[object] | None"
) -> "list[object] | None": # mutable-ok: should_retry_this_error's own parameter type
return cast("list[object] | None", kwargs.get(key, router_default)) # cast-ok: same type as the router attribute
def _request_model_group(kwargs: Mapping[str, object]) -> str | None:
model_group: Final = kwargs.get("model")
return model_group if isinstance(model_group, str) else None
def _mid_stream_retry_trigger(error: "MidStreamFallbackError") -> Exception:
"""The provider's own error, which is what the retry policy and should_retry_this_error classify."""
return error.original_exception if error.original_exception is not None else error
def _anthropic_stream_commits_now(chunk: object, has_generated_content: bool, buffered_chunk_count: int) -> bool:
"""
Whether `chunk` should make Router._aanthropic_messages_streaming_iterator
@ -639,9 +678,29 @@ def _anthropic_stream_commits_now(chunk: object, has_generated_content: bool, bu
return is_anthropic_content_delta_chunk(chunk) or buffered_chunk_count >= MAX_BUFFERED_PRE_CONTENT_ANTHROPIC_CHUNKS
@dataclass(frozen=True, slots=True)
class _AnthropicStreamRetryOpened:
response: object
attempted_retries: int
max_retries: int
@dataclass(frozen=True, slots=True)
class _AnthropicStreamRetriesExhausted:
error: "MidStreamFallbackError"
_AnthropicStreamRetryOutcome: TypeAlias = _AnthropicStreamRetryOpened | _AnthropicStreamRetriesExhausted
MAX_HELD_PRE_OUTPUT_RESPONSES_EVENTS: Final = 200
def _retry_policy_ceiling(policy: RetryPolicy) -> int:
"""The most retries any error class under this policy can be granted."""
return max((retries for retries in policy.model_dump().values() if isinstance(retries, int)), default=0)
def _responses_stream_holds_event(item: object, held_event_count: int) -> bool:
from litellm.responses.streaming_iterator import PRE_OUTPUT_LIFECYCLE_EVENT_TYPES
@ -665,6 +724,7 @@ class FallbackAwareAnthropicMessagesStream:
self._source_iterator = source_iterator
self.fallback_headers_adopted = False
self._hidden_params = dict(getattr(source_iterator, "_hidden_params", None) or {})
self._followed_source_params: object = None
@property
def has_buffered_provider_output(self) -> bool:
@ -692,6 +752,22 @@ class FallbackAwareAnthropicMessagesStream:
self._source_iterator = fallback_response
self.fallback_headers_adopted = True
def follow_source_attribution(self) -> None:
"""
A retry's or a fallback's stream carries a wrapper of its own, so a hop it makes before its
first byte lands on that wrapper while the proxy reads the headers off this one. Mirrors the
source's current attribution onto this wrapper whenever the source has adopted a new one.
"""
source: Final = self._source_iterator
if not getattr(source, "fallback_headers_adopted", False):
return
source_params: Final = getattr(source, "_hidden_params", None)
if source_params is None or source_params is self._followed_source_params:
return
self._followed_source_params = source_params
hidden_params, headers = Router._prepare_fallback_hidden_params(source) # pyright: ignore[reportPrivateUsage] # this wrapper is the Router's own stream type
self.merge_fallback_hidden_params(hidden_params, headers)
def __aiter__(self) -> "FallbackAwareAnthropicMessagesStream":
return self
@ -5435,6 +5511,7 @@ class Router:
if model is not None:
self.fail_calls[model] += 1
if deployment is not None:
self._set_deployment_num_retries_on_exception(e, deployment)
self._stamp_failed_deployment_id_with_effective_model_info(e, deployment, kwargs)
raise e
@ -5565,8 +5642,9 @@ class Router:
# to take over there is nothing to buffer for, so every frame,
# including pings and provider error frames, is forwarded live.
model: Final = cast(str, initial_kwargs.get("model")) # cast-ok: kwargs always carries the model group
has_generated_content = not self._anthropic_messages_stream_can_fall_back( # rebind-ok: set once real content is seen, the buffer cap is hit, or no fallback can take over
model, initial_kwargs
has_generated_content = not ( # rebind-ok: set once real content is seen, the buffer cap is hit, or neither a retry nor a fallback can take over
self._anthropic_messages_stream_can_retry(initial_kwargs)
or self._anthropic_messages_stream_can_fall_back(model, initial_kwargs)
)
buffered_lifecycle_chunks: tuple[bytes, ...] = () # rebind-ok: flushed once committed or on decline
try:
@ -5585,11 +5663,8 @@ class Router:
else chunk
)
error_event = parse_anthropic_error_event(parse_window)
retriable_pending_error = (
not has_generated_content
and error_event is not None
and _is_retriable_anthropic_status(error_event[2])
and not _anthropic_stream_error_is_gateway_verdict(chunk)
recoverable_frame_error = self._anthropic_messages_recoverable_frame_error(
error_event, chunk, has_generated_content, model, initial_kwargs
)
refusal_stop_details = (
parse_anthropic_refusal_stop_details(parse_window)
@ -5605,22 +5680,16 @@ class Router:
original_exception=refusal_error,
is_pre_first_chunk=True,
)
if not has_generated_content and not retriable_pending_error and error_event is None:
if not has_generated_content and error_event is None:
buffered_lifecycle_chunks = (*buffered_lifecycle_chunks, chunk)
continue
if retriable_pending_error:
if recoverable_frame_error is not None:
assert error_event is not None
_error_type, message, status_code = error_event
raise MidStreamFallbackError(
message=message,
message=error_event[1],
model=model,
llm_provider="anthropic",
original_exception=litellm.exceptions.APIError(
status_code=status_code,
message=message,
llm_provider="anthropic",
model=model,
),
original_exception=recoverable_frame_error,
is_pre_first_chunk=True,
)
for buffered_chunk in buffered_lifecycle_chunks:
@ -5676,8 +5745,244 @@ class Router:
)
if fallback_error is None:
raise stream_error
async for item in self._aanthropic_messages_fallback_attempt(fallback_error, initial_kwargs, wrapper):
yield item
outcome: Final = await self._aanthropic_messages_retry_same_group(fallback_error, initial_kwargs)
match outcome:
case _AnthropicStreamRetryOpened(response=retried, attempted_retries=attempted, max_retries=budget):
async for item in self._aanthropic_messages_yield_recovered(retried, wrapper, (attempted, budget)):
yield item
case _AnthropicStreamRetriesExhausted(error=last_error):
async for item in self._aanthropic_messages_fallback_attempt(last_error, initial_kwargs, wrapper):
yield item
case _:
assert_never(outcome)
def _anthropic_messages_group_retry_policy(self, kwargs: Mapping[str, object]) -> dict[str, RetryPolicy] | None:
configured: Final = kwargs.get("model_group_retry_policy", self.model_group_retry_policy)
return cast("dict[str, RetryPolicy] | None", configured) # cast-ok: same type as the router attribute
def _anthropic_messages_resolved_retry_policy(self, kwargs: Mapping[str, object]) -> RetryPolicy | None:
"""
The retry policy for this request's model group, unless the request opted out with num_retries=0.
A policy that does not resolve to a RetryPolicy governs nothing here, so the stream runs as it would
with none; the pre-stream retry loop still reports the malformed policy when an attempt fails.
"""
if kwargs.get("num_retries") == 0:
return None
model_group: Final = _request_model_group(kwargs)
try:
return _RESOLVED_RETRY_POLICY_ADAPTER.validate_python(
resolve_retry_policy(
retry_policy=self.retry_policy,
model_group=model_group,
model_group_retry_policy=self._anthropic_messages_group_retry_policy(kwargs),
)
)
except (TypeError, ValidationError) as malformed:
verbose_router_logger.warning(
"The retry policy for %s is not a RetryPolicy, streaming without one: %s", model_group, malformed
)
return None
def _anthropic_messages_plain_retry_budget(self, kwargs: Mapping[str, object]) -> int:
"""
The same precedence async_function_with_retries resolves for a failure raised before the
stream opened: the request's num_retries, then the routed deployment's, then the router's.
"""
request_num_retries: Final = kwargs.get("num_retries")
if isinstance(request_num_retries, int):
return request_num_retries
deployment_id: Final = routed_deployment_id(kwargs)
deployment: Final = self.get_deployment(deployment_id) if deployment_id is not None else None
deployment_num_retries: Final = _deployment_num_retries(deployment)
if deployment_num_retries is not None:
return deployment_num_retries
return self.num_retries if self.num_retries is not None else 0
def _anthropic_messages_policy_retries(self, trigger: Exception, kwargs: Mapping[str, object]) -> int | None:
policy: Final = self._anthropic_messages_resolved_retry_policy(kwargs)
if policy is None:
return None
return _get_num_retries_from_retry_policy(exception=trigger, retry_policy=policy)
def _anthropic_messages_retry_budget(self, trigger: Exception, kwargs: Mapping[str, object]) -> tuple[int, bool]:
"""
The budget an earlier retry of this request committed to, else the retry policy's grant when one
names this error, else the plain budget, with whether a policy governs the retry: a committed budget
is kept whichever deployment the retry lands on, as async_function_with_retries keeps its own.
"""
policy_retries: Final = self._anthropic_messages_policy_retries(trigger, kwargs)
committed_budget: Final = committed_retry_budget_for_request(kwargs)
if committed_budget is not None:
return committed_budget, policy_retries is not None
if policy_retries is None:
return self._anthropic_messages_plain_retry_budget(kwargs), False
return policy_retries, True
def _anthropic_messages_stream_can_retry(self, kwargs: Mapping[str, object]) -> bool:
"""
Whether a pre-content failure of this stream would be retried within its own model group,
the other case where holding lifecycle frames back from the client buys a clean restart.
A retry policy names its budget per error class, so the largest budget it names bounds the
hold: holding frames one attempt too long is safe, forwarding them before a retry is not.
"""
attempted: Final = attempted_retries_for_request(kwargs)
committed_budget: Final = committed_retry_budget_for_request(kwargs)
if committed_budget is not None:
return committed_budget > attempted
plain_budget: Final = self._anthropic_messages_plain_retry_budget(kwargs)
policy: Final = self._anthropic_messages_resolved_retry_policy(kwargs)
ceiling: Final = plain_budget if policy is None else max(plain_budget, _retry_policy_ceiling(policy))
return ceiling > attempted
def _anthropic_messages_recoverable_frame_error(
self,
error_event: tuple[str, str, int] | None,
chunk: object,
has_generated_content: bool,
model_group: str,
kwargs: Mapping[str, object],
) -> Exception | None:
"""
The exception a provider `event: error` frame before content recovers through when a retry of its
class or a fallback can still take over. A frame nothing can take over for (content already out, a
gateway verdict, or a class granted no retry with no fallback) reaches the client as the provider
sent it, the way the last exhausted attempt's does.
"""
if has_generated_content or error_event is None:
return None
error_type, message, status_code = error_event
if not _is_retriable_anthropic_status(status_code) or _anthropic_stream_error_is_gateway_verdict(chunk):
return None
frame_error: Final = anthropic_error_frame_exception(error_type, message, status_code, model_group)
budget, _ = self._anthropic_messages_retry_budget(frame_error, kwargs)
if budget > attempted_retries_for_request(kwargs):
return frame_error
if self._anthropic_messages_stream_can_fall_back(model_group, kwargs):
return frame_error
return None
def _anthropic_messages_should_retry(
self,
trigger: Exception,
healthy_deployments: list[dict], # mutable-ok: should_retry_this_error's own parameter type
all_deployments: list[dict], # mutable-ok: should_retry_this_error's own parameter type
kwargs: Mapping[str, object],
) -> bool:
try:
self.should_retry_this_error(
error=trigger,
healthy_deployments=healthy_deployments,
all_deployments=all_deployments,
context_window_fallbacks=_request_fallback_list(
kwargs,
"context_window_fallbacks",
cast("list[object] | None", self.context_window_fallbacks), # cast-ok: untyped router attribute
),
content_policy_fallbacks=_request_fallback_list(
kwargs,
"content_policy_fallbacks",
cast("list[object] | None", self.content_policy_fallbacks), # cast-ok: untyped router attribute
),
regular_fallbacks=_request_fallback_list(
kwargs,
"fallbacks",
cast("list[object] | None", self.fallbacks), # cast-ok: untyped router attribute
),
)
except Exception: # noqa: BLE001 # should_retry_this_error declines by raising the error it was given
return False
return True
async def _aanthropic_messages_retry_same_group(
self, e: "MidStreamFallbackError", initial_kwargs: Mapping[str, object]
) -> _AnthropicStreamRetryOutcome:
"""
Re-runs the attempt within the request's own model group, the way async_function_with_retries
would have for a failure raised before the stream opened, until a retry opens a stream or the
budget runs out. Each retry's stream carries its own wrapper with the remaining budget, so a
retry that drops before content again continues the same count instead of starting over.
"""
model_group: Final = cast(str, initial_kwargs.get("model")) # cast-ok: kwargs always carries the model group
retry_kwargs: Final = mid_stream_retry_kwargs(initial_kwargs)
healthy_deployments, all_deployments = await self._async_get_healthy_deployments(
model=model_group, parent_otel_span=_get_parent_otel_span_from_kwargs(retry_kwargs)
)
budget, policy_applies = self._anthropic_messages_retry_budget(_mid_stream_retry_trigger(e), initial_kwargs)
last_error = e # rebind-ok: the newest failure is what the fallback chain and the caller see
for attempt in range(attempted_retries_for_request(initial_kwargs), budget):
trigger = _mid_stream_retry_trigger(last_error)
if not policy_applies and not self._anthropic_messages_should_retry(
trigger, healthy_deployments, all_deployments, initial_kwargs
):
return _AnthropicStreamRetriesExhausted(last_error)
self.log_retry(kwargs=retry_kwargs, e=trigger)
await asyncio.sleep(
self._time_to_sleep_before_retry(
e=trigger,
remaining_retries=budget - attempt,
num_retries=budget,
healthy_deployments=healthy_deployments,
all_deployments=all_deployments,
)
)
record_retry_attempt(retry_kwargs, attempted_retries=attempt + 1, max_retries=budget)
verbose_router_logger.debug(
"Retrying anthropic_messages stream dropped before content, attempt %s of %s", attempt + 1, budget
)
try:
response = await self._ageneric_api_call_with_fallbacks_anthropic_messages_attempt(**retry_kwargs)
except Exception as retry_error: # noqa: BLE001 # every failure of a retry before its stream opens is the fallback chain's to judge
wrapped = _anthropic_stream_fallback_error_for_raised(retry_error, model_group, False)
if wrapped is None:
return _AnthropicStreamRetriesExhausted(
_anthropic_stream_pre_content_error(retry_error, model_group)
)
last_error = wrapped
continue
return _AnthropicStreamRetryOpened(response, attempted_retries=attempt + 1, max_retries=budget)
return _AnthropicStreamRetriesExhausted(last_error)
async def _aanthropic_messages_yield_recovered(
self,
recovered: object,
wrapper: "FallbackAwareAnthropicMessagesStream",
retry_counters: tuple[int, int] | None = None,
) -> AsyncGenerator[bytes, None]:
"""
Hands a retry's or a fallback's response to the client through the wrapper, closing it afterwards.
A retry stamps the retry headers async_function_with_retries would have for a pre-stream retry;
a later hop the recovered stream makes replaces them with its own, the way a fallback's do.
"""
from litellm.llms.anthropic.pass_through.messages.streaming_iterator import (
aclose_if_supported,
anthropic_messages_response_as_sse_events,
)
hidden_params, headers = Router._prepare_fallback_hidden_params(recovered)
wrapper.merge_fallback_hidden_params(hidden_params, headers)
wrapper.adopt_fallback_source(recovered)
if retry_counters is not None:
add_retry_headers_to_response(
response=wrapper, attempted_retries=retry_counters[0], max_retries=retry_counters[1]
)
try:
if hasattr(recovered, "__aiter__"):
async for item in cast("AsyncIterator[bytes]", recovered): # cast-ok: __aiter__ checked above
wrapper.follow_source_attribution()
yield item
return
# A recovery can resolve to a complete AnthropicMessagesResponse
# dict even for a streaming request (e.g. an agentic tool-use
# interception loop) - yielding it as-is would put a raw dict
# into a byte stream, so it's synthesized into the SSE
# lifecycle a real stream would have sent instead.
for event in anthropic_messages_response_as_sse_events(
cast("AnthropicMessagesResponse", recovered) # cast-ok: non-streaming shape by elimination
):
yield event
finally:
with anyio.CancelScope(shield=True), contextlib.suppress(BaseException):
await aclose_if_supported(recovered)
async def _aanthropic_messages_fallback_attempt(
self,
@ -5693,12 +5998,7 @@ class Router:
budget.
"""
from litellm.exceptions import MidStreamFallbackError
from litellm.llms.anthropic.pass_through.messages.streaming_iterator import (
aclose_if_supported,
anthropic_messages_response_as_sse_events,
)
fallback_response = None # rebind-ok: pre-init so finally can close it if a fallback was actually attempted
try:
model_group: Final = cast(str, initial_kwargs.get("model")) # cast-ok: model group
fallbacks: Final[list | None] = initial_kwargs.get( # mutable-ok: matches the common_utils list|None param
@ -5716,12 +6016,16 @@ class Router:
kwargs=initial_kwargs,
metadata_variable_name="litellm_metadata",
)
# The content-policy dispatch branch matches on the trigger's own type, so a refusal's
# MidStreamFallbackError envelope is unwrapped here or the wrong fallback list is consulted.
# The content-policy and context-window dispatch branches match on the trigger's own type, so
# such an error's MidStreamFallbackError envelope is unwrapped here or the wrong fallback list is consulted.
fallback_trigger: Final[Exception] = (
e.original_exception if isinstance(e.original_exception, litellm.ContentPolicyViolationError) else e
e.original_exception
if isinstance(
e.original_exception, (litellm.ContentPolicyViolationError, litellm.ContextWindowExceededError)
)
else e
)
fallback_response = await self.async_function_with_fallbacks_common_utils( # rebind-ok: set on success
fallback_response: Final = await self.async_function_with_fallbacks_common_utils(
e=fallback_trigger,
disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs),
fallbacks=fallbacks,
@ -5732,31 +6036,13 @@ class Router:
kwargs=initial_kwargs,
include_fallback_errors=initial_kwargs.get("include_fallback_errors", False) is True,
)
fallback_hidden_params, fallback_headers = Router._prepare_fallback_hidden_params(fallback_response)
wrapper.merge_fallback_hidden_params(fallback_hidden_params, fallback_headers)
wrapper.adopt_fallback_source(fallback_response)
if hasattr(fallback_response, "__aiter__"):
async for fallback_item in fallback_response:
yield fallback_item
else:
# A fallback can resolve to a complete AnthropicMessagesResponse
# dict even for a streaming request (e.g. an agentic tool-use
# interception loop) - yielding it as-is would put a raw dict
# into a byte stream, so it's synthesized into the SSE
# lifecycle a real stream would have sent instead.
for event in anthropic_messages_response_as_sse_events(
cast("AnthropicMessagesResponse", fallback_response) # cast-ok: non-streaming shape by elimination
):
yield event
async for fallback_item in self._aanthropic_messages_yield_recovered(fallback_response, wrapper):
yield fallback_item
except Exception as fallback_error:
verbose_router_logger.error("Anthropic messages streaming fallback also failed: %s", fallback_error)
if isinstance(fallback_error, MidStreamFallbackError) and fallback_error.original_exception is not None:
raise fallback_error.original_exception from fallback_error
raise
finally:
if fallback_response is not None:
with anyio.CancelScope(shield=True), contextlib.suppress(BaseException):
await aclose_if_supported(fallback_response)
async def _aanthropic_messages_with_streaming_fallbacks(
self,
@ -5795,6 +6081,7 @@ class Router:
model=model, original_generic_function=original_generic_function, **kwargs
)
carry_over_pre_routing_selection(live_kwargs=kwargs, snapshot=hop_kwargs)
carry_over_routed_deployment(live_kwargs=kwargs, snapshot=hop_kwargs)
if kwargs.get("stream") and hasattr(response, "__aiter__"):
return await self._aanthropic_messages_streaming_iterator(
response=cast("AsyncIterator[bytes]", response), # cast-ok: stream=True always returns a byte iterator

View file

@ -362,6 +362,78 @@ def mid_stream_fallback_hop_kwargs(
}
_MID_STREAM_RETRY_STRIPPED_KEYS: Final = (*_PER_REQUEST_FALLBACK_CONTROL_KEYS, "original_function")
_MID_STREAM_RETRY_ATTEMPTED_KEY: Final = "attempted_retries"
_MID_STREAM_RETRY_BUDGET_KEY: Final = "max_retries"
def mid_stream_retry_kwargs(
hop_kwargs: Mapping[str, object],
) -> dict[str, object]: # mutable-ok: unpacked as **kwargs into the attempt function, which pops its controls carrier
"""
The kwargs a same-group retry re-enters the attempt function with. async_function_with_retries
pops the per-request controls and the chain's original_function before any attempt runs, and
the controls carrier the snapshot still holds restores the overrides into the retry's own hop.
"""
return {key: value for key, value in hop_kwargs.items() if key not in _MID_STREAM_RETRY_STRIPPED_KEYS}
def _request_metadata_bucket(kwargs: Mapping[str, object]) -> Mapping[str, object] | None:
bucket: Final = kwargs.get(get_metadata_variable_name_from_kwargs(kwargs))
return bucket if isinstance(bucket, Mapping) else None
def attempted_retries_for_request(kwargs: Mapping[str, object]) -> int:
"""How many same-group retries async_function_with_retries, or a mid-stream retry, already spent on this request."""
bucket: Final = _request_metadata_bucket(kwargs)
attempted: Final = bucket.get(_MID_STREAM_RETRY_ATTEMPTED_KEY) if bucket is not None else None
return attempted if type(attempted) is int and attempted > 0 else 0
def committed_retry_budget_for_request(kwargs: Mapping[str, object]) -> int | None:
"""The budget the first retry of this request committed to, kept by every later attempt the way the
pre-stream retry loop keeps its own; None until a retry has run."""
if attempted_retries_for_request(kwargs) == 0:
return None
bucket: Final = _request_metadata_bucket(kwargs)
budget: Final = bucket.get(_MID_STREAM_RETRY_BUDGET_KEY) if bucket is not None else None
return budget if type(budget) is int else None
def record_retry_attempt(kwargs: Mapping[str, object], attempted_retries: int, max_retries: int) -> None:
"""Stamp the attempt about to run the way async_function_with_retries does before each of its retries."""
bucket: Final = kwargs.get(get_metadata_variable_name_from_kwargs(kwargs))
if not isinstance(bucket, dict):
return
bucket[_MID_STREAM_RETRY_ATTEMPTED_KEY] = attempted_retries
bucket[_MID_STREAM_RETRY_BUDGET_KEY] = max_retries
def _routed_model_info(kwargs: Mapping[str, object]) -> Mapping[str, object] | None:
bucket: Final = _request_metadata_bucket(kwargs)
model_info: Final = bucket.get("model_info") if bucket is not None else None
return model_info if isinstance(model_info, Mapping) else None
def routed_deployment_id(kwargs: Mapping[str, object]) -> str | None:
model_info: Final = _routed_model_info(kwargs)
deployment_id: Final = model_info.get("id") if model_info is not None else None
return deployment_id if isinstance(deployment_id, str) else None
def carry_over_routed_deployment(live_kwargs: Mapping[str, object], snapshot: Mapping[str, object]) -> None:
"""
Copy the deployment this attempt routed to into the snapshot's metadata bucket, which was
taken before routing: a same-group retry reads the deployment's own num_retries off it and
records which deployment failed.
"""
snapshot_bucket: Final = snapshot.get(get_metadata_variable_name_from_kwargs(snapshot))
model_info: Final = _routed_model_info(live_kwargs)
if not isinstance(snapshot_bucket, dict) or model_info is None:
return
snapshot_bucket["model_info"] = dict(model_info)
DISABLE_FALLBACKS_METADATA_KEY: Final = "_disable_fallbacks"

View file

@ -34,7 +34,7 @@ def _retries_for_a_404_answer(exception: Exception, policy: RetryPolicy) -> int
return policy.NotFoundErrorRetries if status_code == 404 else None
def _resolve_policy(
def resolve_retry_policy(
retry_policy: RetryPolicy | Mapping[str, int | None] | None,
model_group: str | None,
model_group_retry_policy: Mapping[str, RetryPolicy | Mapping[str, int | None]] | None,
@ -56,7 +56,7 @@ def get_num_retries_from_retry_policy(
model_group_retry_policy: Mapping[str, RetryPolicy | Mapping[str, int | None]] | None = None,
) -> int | None:
"""Prefer NotFoundErrorRetries for any 404 answer, then walk the exception's MRO most specific class first."""
policy: Final = _resolve_policy(retry_policy, model_group, model_group_retry_policy)
policy: Final = resolve_retry_policy(retry_policy, model_group, model_group_retry_policy)
if policy is None:
return None
by_class: Final = (

View file

@ -93,6 +93,18 @@ ignored_function_names = [
"_async_get_available_deployment_for_pass_through", # Same, through async_get_available_deployment_for_pass_through in test_router.py
"_embedding",
"_aembedding",
"_anthropic_stream_pre_content_error", # Tested through the non-retriable retry error tests in test_router.py
"_deployment_num_retries", # Tested through the deployment num_retries mid-stream budget test in test_router.py
"_request_fallback_list", # Tested through every mid-stream retry test in test_router.py
"_request_model_group", # Tested through test_anthropic_messages_retry_budget_precedence_direct_call
"_mid_stream_retry_trigger", # Tested through the retry policy mid-stream budget test in test_router.py
"_anthropic_messages_group_retry_policy", # Tested through the retry budget precedence test in test_router.py
"_anthropic_messages_resolved_retry_policy", # Tested through the malformed retry policy tests in test_router.py
"_anthropic_messages_plain_retry_budget", # Tested through the retry budget precedence test in test_router.py
"_anthropic_messages_should_retry", # Tested through every mid-stream retry test in test_router.py
"_aanthropic_messages_retry_same_group", # Tested through the dropped-before-content retry tests in test_router.py
"_aanthropic_messages_yield_recovered", # Tested through every mid-stream retry and fallback test in test_router.py
"_anthropic_messages_policy_retries", # Tested through the retry budget precedence test in test_router.py
]

View file

@ -0,0 +1,197 @@
from __future__ import annotations
import json
import threading
from collections import Counter
from collections.abc import Mapping
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final
from integration._support.wire import Reply
from pydantic import JsonValue, TypeAdapter
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
_MESSAGES: Final = TypeAdapter(list[dict[str, JsonValue]])
_EMPTY: Final[Mapping[str, JsonValue]] = MappingProxyType({})
ANTHROPIC_ERROR_TYPES: Final = MappingProxyType(
{
400: "invalid_request_error",
401: "authentication_error",
408: "api_error",
409: "api_error",
429: "rate_limit_error",
500: "api_error",
503: "api_error",
529: "overloaded_error",
}
)
LIFECYCLE: Final = (
"message_start",
"content_block_start",
"content_block_delta",
"content_block_stop",
"message_delta",
"message_stop",
)
def sse(event: str, payload: Mapping[str, JsonValue]) -> bytes:
return f"event: {event}\ndata: {json.dumps(payload)}\n\n".encode()
def message_start(message_id: str, model: str) -> bytes:
return sse(
"message_start",
{
"type": "message_start",
"message": {
"id": message_id,
"type": "message",
"role": "assistant",
"model": model,
"content": [],
"stop_reason": None,
"stop_sequence": None,
"usage": {"input_tokens": 5, "output_tokens": 1},
},
},
)
def text_delta(text: str) -> bytes:
return sse(
"content_block_delta",
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": text}},
)
PING: Final = sse("ping", {"type": "ping"})
CONTENT_BLOCK_START: Final = sse(
"content_block_start",
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
)
CONTENT_TAIL: Final = (
sse("content_block_stop", {"type": "content_block_stop", "index": 0})
+ sse(
"message_delta",
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
"usage": {"output_tokens": 3},
},
)
+ sse("message_stop", {"type": "message_stop"})
)
def message_stream(message_id: str, model: str, text: str) -> tuple[bytes, bytes, bytes, bytes]:
return (message_start(message_id, model), CONTENT_BLOCK_START, text_delta(text), CONTENT_TAIL)
def message_json(message_id: str, model: str, text: str) -> bytes:
return json.dumps(
{
"id": message_id,
"type": "message",
"role": "assistant",
"model": model,
"content": [{"type": "text", "text": text}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 5, "output_tokens": 3},
}
).encode()
def error_frame(status: int, message: str) -> bytes:
return sse("error", {"type": "error", "error": {"type": ANTHROPIC_ERROR_TYPES[status], "message": message}})
def error_body(status: int, message: str) -> bytes:
return json.dumps({"type": "error", "error": {"type": ANTHROPIC_ERROR_TYPES[status], "message": message}}).encode()
DROP_PAUSE: Final = 0.2
def stream_reply(chunks: tuple[bytes, ...], *, abort_after: int | None = None, pause: float = 0) -> Reply:
return Reply(content_type="text/event-stream", chunks=chunks, abort_after=abort_after, pause_between_chunks=pause)
def dropping_reply(chunks: tuple[bytes, ...], *, abort_after: int) -> Reply:
return stream_reply(chunks, abort_after=abort_after, pause=DROP_PAUSE if abort_after else 0)
def status_reply(status: int) -> Reply:
return Reply(status=status, body=error_body(status, f"scripted {status}"))
@dataclass(frozen=True, slots=True)
class SseEvent:
event: str
data: Mapping[str, JsonValue]
def _parse_block(block: str) -> SseEvent:
lines: Final = block.splitlines()
event: Final = next((line.removeprefix("event:").strip() for line in lines if line.startswith("event:")), "")
data: Final = "".join(line.removeprefix("data:").strip() for line in lines if line.startswith("data:"))
return SseEvent(event, _JSON_OBJECT.validate_json(data) if data else _EMPTY)
def _is_event_block(block: str) -> bool:
return bool(block.strip()) and block.strip() != "data: [DONE]"
def parse_sse(text: str) -> tuple[SseEvent, ...]:
return tuple(_parse_block(block) for block in text.replace("\r\n", "\n").split("\n\n") if _is_event_block(block))
def event_type(event: SseEvent) -> str:
return event.event or str(event.data.get("type", ""))
def event_types(events: tuple[SseEvent, ...]) -> tuple[str, ...]:
return tuple(event_type(event) for event in events)
def message_id(events: tuple[SseEvent, ...]) -> str:
start: Final = next(event for event in events if event.event == "message_start")
return str(_JSON_OBJECT.validate_python(start.data["message"])["id"])
def delta_text(events: tuple[SseEvent, ...]) -> str:
deltas: Final = tuple(
_JSON_OBJECT.validate_python(event.data["delta"]) for event in events if event.event == "content_block_delta"
)
return "".join(str(delta.get("text", "")) for delta in deltas)
def error_type(events: tuple[SseEvent, ...]) -> str | None:
error: Final = next((event for event in events if event.event == "error"), None)
if error is None:
return None
return str(_JSON_OBJECT.validate_python(error.data["error"])["type"])
def user_prompt(body: Mapping[str, JsonValue]) -> str:
content: Final = _MESSAGES.validate_python(body["messages"])[0]["content"]
assert isinstance(content, str), content
return content
class Attempts:
def __init__(self) -> None:
self._lock: Final = threading.Lock()
self._seen: Final = Counter[str]()
def record(self, marker: str) -> int:
with self._lock:
self._seen[marker] += 1
return self._seen[marker]
def count(self, marker: str) -> int:
with self._lock:
return self._seen[marker]

View file

@ -0,0 +1,125 @@
from __future__ import annotations
import json
from collections.abc import Callable
from typing import Final
from integration._support.wire import Reply, Request, Wire
from pydantic import JsonValue
_USAGE: Final = {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8}
MODEL_DISCOVERY: Final = ("GET", "/v1/models")
def answering_model_discovery(respond: Callable[[Request], Reply]) -> Callable[[Request], Reply]:
def guarded(request: Request) -> Reply:
if (request.method, request.target) == MODEL_DISCOVERY:
return Reply(body=json.dumps({"object": "list", "data": []}).encode())
return respond(request)
return guarded
def posted_targets(wire: Wire) -> tuple[str, ...]:
return tuple(request.target for request in wire.drain() if request.method == "POST")
def openai_error(status: int) -> Reply:
return Reply(
status=status,
body=json.dumps({"error": {"message": f"scripted {status}", "type": "server_error", "code": None}}).encode(),
)
def _data_frame(frame: dict[str, JsonValue]) -> bytes:
return b"data: " + json.dumps(frame).encode() + b"\n\n"
def _typed_frame(event: dict[str, JsonValue]) -> bytes:
return f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode()
def chat_stream(identity: str, model: str, text: str) -> tuple[bytes, bytes, bytes]:
chunk: Final = {"id": identity, "object": "chat.completion.chunk", "created": 1, "model": model}
role_only: Final = _data_frame({**chunk, "choices": [{"index": 0, "delta": {"role": "assistant", "content": ""}}]})
content: Final = _data_frame({**chunk, "choices": [{"index": 0, "delta": {"content": text}}]})
finish: Final = _data_frame(
{**chunk, "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}], "usage": _USAGE}
)
return (role_only, content, finish + b"data: [DONE]\n\n")
def chat_reply(identity: str, model: str, text: str, *, stream: bool) -> Reply:
if stream:
return Reply(content_type="text/event-stream", chunks=chat_stream(identity, model, text))
return Reply(
body=json.dumps(
{
"id": identity,
"object": "chat.completion",
"created": 1,
"model": model,
"choices": [{"index": 0, "message": {"role": "assistant", "content": text}, "finish_reason": "stop"}],
"usage": _USAGE,
}
).encode()
)
def _response_object(identity: str, model: str, text: str) -> dict[str, JsonValue]:
return {
"id": identity,
"object": "response",
"created_at": 1,
"status": "completed",
"model": model,
"output": [_message_item(identity, text, "completed")],
"usage": {"input_tokens": 5, "output_tokens": 3, "total_tokens": 8},
}
def _message_item(identity: str, text: str, status: str) -> dict[str, JsonValue]:
return {
"id": f"msg_{identity}",
"type": "message",
"role": "assistant",
"status": status,
"content": [{"type": "output_text", "text": text, "annotations": []}] if status == "completed" else [],
}
def responses_stream(identity: str, model: str, text: str) -> tuple[bytes, bytes, bytes]:
response: Final = _response_object(identity, model, text)
item: Final = _message_item(identity, text, "in_progress")
opened: Final = _typed_frame(
{
"type": "response.created",
"sequence_number": 0,
"response": {**response, "status": "in_progress", "output": []},
}
) + _typed_frame({"type": "response.output_item.added", "sequence_number": 1, "output_index": 0, "item": item})
delta: Final = _typed_frame(
{
"type": "response.output_text.delta",
"sequence_number": 2,
"item_id": f"msg_{identity}",
"output_index": 0,
"content_index": 0,
"delta": text,
}
)
closed: Final = _typed_frame(
{
"type": "response.output_item.done",
"sequence_number": 3,
"output_index": 0,
"item": _message_item(identity, text, "completed"),
}
) + _typed_frame({"type": "response.completed", "sequence_number": 4, "response": response})
return (opened, delta, closed)
def responses_reply(identity: str, model: str, text: str, *, stream: bool) -> Reply:
if stream:
return Reply(content_type="text/event-stream", chunks=responses_stream(identity, model, text))
return Reply(body=json.dumps(_response_object(identity, model, text)).encode())

View file

@ -0,0 +1,179 @@
import json
import uuid
from collections.abc import Callable
from dataclasses import dataclass
from typing import Final, Literal
import anthropic
import pytest
from integration._support.anthropic_sse import (
Attempts,
SseEvent,
delta_text,
dropping_reply,
event_types,
parse_sse,
stream_reply,
)
from integration._support.client import Gateway, eventually, object_value
from integration._support.database import read_rows
from integration._support.openai_wire import (
answering_model_discovery,
chat_stream,
openai_error,
posted_targets,
responses_stream,
)
from integration._support.wire import Reply, Request, Wire, wire_server
from pydantic import JsonValue
_BACKEND: Final = "gpt-4o-mini"
_PROVIDER_KEY: Final = "integration-provider-key"
_TEXT: Final = "Hello"
FirstAttempt = Literal["drop_after_headers", "drop_after_pre_content_frame", "http_500", "drop_after_content"]
@dataclass(frozen=True, slots=True)
class _Bridge:
name: str
provider_model: str
target: str
stream: Callable[[str, str, str], tuple[bytes, bytes, bytes]]
_CHAT_COMPLETIONS: Final = _Bridge("chat", f"hosted_vllm/{_BACKEND}", "/v1/chat/completions", chat_stream)
_RESPONSES_API: Final = _Bridge("responses", f"openai/{_BACKEND}", "/v1/responses", responses_stream)
_BRIDGES: Final = (_CHAT_COMPLETIONS, _RESPONSES_API)
def _bridge_id(bridge: _Bridge) -> str:
return bridge.name
def _first_attempt_reply(kind: FirstAttempt, chunks: tuple[bytes, bytes, bytes]) -> Reply:
match kind:
case "drop_after_headers":
return dropping_reply(chunks, abort_after=0)
case "drop_after_pre_content_frame":
return dropping_reply(chunks, abort_after=1)
case "http_500":
return openai_error(500)
case "drop_after_content":
return dropping_reply(chunks, abort_after=2)
def _upstream(bridge: _Bridge, marker: str, kind: FirstAttempt, attempts: Attempts) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert (request.method, request.target) == ("POST", bridge.target), request
assert request.headers["authorization"] == f"Bearer {_PROVIDER_KEY}", request.headers
body: Final = object_value(json.loads(request.body))
assert body["model"] == _BACKEND, body
assert body["stream"] is True, body
assert marker in request.body.decode(), body
assert "num_retries" not in body, body
attempt: Final = attempts.record(marker)
chunks: Final = bridge.stream(f"{bridge.name}-{marker}-a{attempt}", _BACKEND, _TEXT)
if attempt > 1:
return stream_reply(chunks)
return _first_attempt_reply(kind, chunks)
return answering_model_discovery(respond)
def _marker() -> str:
return "bridge-pre-content-" + uuid.uuid4().hex
def _body(model: str, marker: str, **extra: JsonValue) -> dict[str, JsonValue]:
return {
"model": model,
"max_tokens": 16,
"stream": True,
"messages": [{"role": "user", "content": marker}],
**extra,
}
def _stream(gateway: Gateway, body: dict[str, JsonValue]) -> tuple[int, tuple[SseEvent, ...]]:
response: Final = gateway.request("POST", "/v1/messages", body)
return response.status_code, parse_sse(response.text)
def _success_rows(model: str) -> list[dict[str, JsonValue]]:
return eventually(
lambda: read_rows('SELECT request_id, status FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)),
lambda rows: len(rows) >= 1,
seconds=70,
)
def _assert_completed_by_retry(
status: int, events: tuple[SseEvent, ...], model: str, wire: Wire, bridge: _Bridge
) -> None:
assert status == 200, events
types: Final = event_types(events)
assert types[0] == "message_start", events
assert "content_block_delta" in types and "error" not in types, events
assert types[-1] == "message_stop", events
assert delta_text(events) == _TEXT, events
assert posted_targets(wire) == (bridge.target,) * 2
rows: Final = _success_rows(model)
assert [row["status"] for row in rows] == ["success"], rows
@pytest.mark.parametrize("bridge", _BRIDGES, ids=_bridge_id)
@pytest.mark.parametrize("kind", ["drop_after_headers", "drop_after_pre_content_frame", "http_500"])
def test_bridge_stream_failing_before_content_is_retried_per_the_deployment_budget(
gateway: Gateway, bridge: _Bridge, kind: FirstAttempt
) -> None:
marker: Final = _marker()
attempts: Final = Attempts()
with wire_server(_upstream(bridge, marker, kind, attempts)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=bridge.provider_model, api_base=wire.url + "/v1", num_retries=1)
status, events = _stream(gateway, _body(model, marker))
_assert_completed_by_retry(status, events, model, wire, bridge)
@pytest.mark.parametrize("bridge", _BRIDGES, ids=_bridge_id)
def test_bridge_stream_failing_before_content_is_retried_per_the_request_budget(
gateway: Gateway, bridge: _Bridge
) -> None:
marker: Final = _marker()
attempts: Final = Attempts()
with wire_server(_upstream(bridge, marker, "drop_after_headers", attempts)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=bridge.provider_model, api_base=wire.url + "/v1")
status, events = _stream(gateway, _body(model, marker, num_retries=1))
_assert_completed_by_retry(status, events, model, wire, bridge)
@pytest.mark.parametrize("bridge", _BRIDGES, ids=_bridge_id)
async def test_anthropic_sdk_async_stream_over_the_bridge_completes_after_a_drop(
gateway: Gateway, bridge: _Bridge
) -> None:
marker: Final = _marker()
attempts: Final = Attempts()
with wire_server(_upstream(bridge, marker, "drop_after_headers", attempts)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=bridge.provider_model, api_base=wire.url + "/v1", num_retries=1)
client: Final = anthropic.AsyncAnthropic(
base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0, timeout=60
)
async with client.messages.stream(
model=model, max_tokens=16, messages=[{"role": "user", "content": marker}]
) as stream:
final: Final = await stream.get_final_message()
assert [block.text for block in final.content if block.type == "text"] == [_TEXT], final
assert posted_targets(wire) == (bridge.target,) * 2
@pytest.mark.parametrize("bridge", _BRIDGES, ids=_bridge_id)
def test_bridge_stream_dropping_after_content_is_not_retried(gateway: Gateway, bridge: _Bridge) -> None:
marker: Final = _marker()
attempts: Final = Attempts()
with wire_server(_upstream(bridge, marker, "drop_after_content", attempts)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=bridge.provider_model, api_base=wire.url + "/v1", num_retries=1)
status, events = _stream(gateway, _body(model, marker))
assert status == 200, events
assert delta_text(events) == _TEXT, events
assert event_types(events)[-1] == "error", events
assert posted_targets(wire) == (bridge.target,)

View file

@ -0,0 +1,411 @@
import asyncio
import base64
import json
import re
import signal
import threading
import uuid
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from pathlib import Path
from queue import SimpleQueue
from types import MappingProxyType
from typing import Final, Literal
from urllib.parse import urlsplit
import httpx
import psutil
import pytest
import yaml
from integration._support.anthropic_sse import (
Attempts,
delta_text,
dropping_reply,
event_type,
event_types,
message_id,
message_json,
message_stream,
parse_sse,
status_reply,
stream_reply,
user_prompt,
)
from integration._support.client import Gateway, eventually, object_value, string_value
from integration._support.database import read_rows
from integration._support.openai_wire import chat_reply, openai_error, responses_reply
from integration._support.process import owned_proxy_process
from integration._support.wire import Reply, Request, Wire, wire_server
from pydantic import JsonValue
_ANTHROPIC_BACKEND: Final = "claude-under-test"
_ANTHROPIC_KEY: Final = "synthetic-anthropic-key"
_OPENAI_BACKEND: Final = "gpt-4o-mini"
_OPENAI_KEY: Final = "integration-provider-key"
_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]")
_OUTAGE_STATUSES: Final = (529, 429, 503)
_ROUTING_ENCODED_ID: Final = re.compile(r"resp_([A-Za-z0-9+/]+=*)")
pytestmark = pytest.mark.timeout(240)
Endpoint = Literal["messages", "chat", "responses"]
@dataclass(frozen=True, slots=True)
class _Models:
messages: str
openai: str
def for_endpoint(self, endpoint: Endpoint) -> str:
return self.messages if endpoint == "messages" else self.openai
@dataclass(frozen=True, slots=True)
class _Call:
endpoint: Endpoint
stream: bool
marker: str
@dataclass(frozen=True, slots=True)
class _Served:
call: _Call
status: int
text: str
def _answer(marker: str) -> str:
return f"answer-{marker}"
def _path(endpoint: Endpoint) -> str:
match endpoint:
case "messages":
return "/v1/messages"
case "chat":
return "/v1/chat/completions"
case "responses":
return "/v1/responses"
def _body(models: _Models, call: _Call) -> dict[str, JsonValue]:
prompt: Final = f"chaos:{call.marker}"
model: Final = models.for_endpoint(call.endpoint)
match call.endpoint:
case "messages":
return {
"model": model,
"max_tokens": 16,
"stream": call.stream,
"messages": [{"role": "user", "content": prompt}],
}
case "chat":
return {"model": model, "stream": call.stream, "messages": [{"role": "user", "content": prompt}]}
case "responses":
return {"model": model, "stream": call.stream, "input": prompt}
def _marker_of(request: Request) -> str:
body: Final = object_value(json.loads(request.body))
prompt: Final = string_value(body["input"]) if "input" in body else user_prompt(body)
return prompt.removeprefix("chaos:")
def _streaming(request: Request) -> bool:
return object_value(json.loads(request.body)).get("stream") is True
def _served(request: Request, marker: str, attempt: int) -> Reply:
text: Final = _answer(marker)
match request.target:
case "/v1/messages":
served_id: Final = f"msg_{marker}_a{attempt}"
if _streaming(request):
return stream_reply(message_stream(served_id, _ANTHROPIC_BACKEND, text))
return Reply(body=message_json(served_id, _ANTHROPIC_BACKEND, text))
case "/v1/chat/completions":
return chat_reply(f"chatcmpl-{marker}-a{attempt}", _OPENAI_BACKEND, text, stream=_streaming(request))
case "/v1/responses":
return responses_reply(f"resp_{marker}_a{attempt}", _OPENAI_BACKEND, text, stream=_streaming(request))
raise AssertionError(request.target)
def _drop_or_500(request: Request, marker: str) -> Reply:
if request.target == "/v1/messages":
return dropping_reply(message_stream(f"msg_{marker}_a1", _ANTHROPIC_BACKEND, _answer(marker)), abort_after=1)
return openai_error(500)
def _outage(statuses: Mapping[str, int]) -> Callable[[Request, str], Reply]:
def first_attempt(request: Request, marker: str) -> Reply:
if request.target == "/v1/messages":
return status_reply(statuses[marker])
return openai_error(statuses[marker])
return first_attempt
@dataclass(frozen=True, slots=True)
class _Upstream:
attempts: Attempts
first_attempt: Callable[[Request, str], Reply]
held: SimpleQueue[str] | None = None
release: threading.Event | None = None
def __call__(self, request: Request) -> Reply:
if request.method == "GET":
return Reply(body=json.dumps({"object": "list", "data": []}).encode())
assert request.method == "POST", request
marker: Final = _marker_of(request)
attempt: Final = self.attempts.record(marker)
if attempt > 1:
return _served(request, marker, attempt)
if self.held is not None and self.release is not None:
self.held.put(marker)
assert self.release.wait(timeout=120), "The burst was never released"
return self.first_attempt(request, marker)
def _config(wire: Wire, directory: Path, models: _Models) -> Path:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["model_list"] = [
{
"model_name": models.messages,
"litellm_params": {
"model": f"anthropic/{_ANTHROPIC_BACKEND}",
"api_base": wire.url,
"api_key": _ANTHROPIC_KEY,
"num_retries": 1,
},
},
{
"model_name": models.openai,
"litellm_params": {
"model": f"openai/{_OPENAI_BACKEND}",
"api_base": wire.url + "/v1",
"api_key": _OPENAI_KEY,
"num_retries": 1,
},
},
]
config["router_settings"] = {"num_retries": 0, "disable_cooldowns": True}
path: Final = directory / "pre-content-retry-chaos.yaml"
path.write_text(yaml.safe_dump(config))
return path
def _models() -> _Models:
suffix: Final = uuid.uuid4().hex
return _Models(messages=f"audit-chaos-messages-{suffix}", openai=f"audit-chaos-openai-{suffix}")
async def _send(client: httpx.AsyncClient, key: str, models: _Models, call: _Call) -> _Served:
async with client.stream(
"POST",
_path(call.endpoint),
json=_body(models, call),
headers={"Authorization": f"Bearer {key}", "anthropic-version": "2023-06-01"},
) as response:
raw: Final = await response.aread()
return _Served(call=call, status=response.status_code, text=raw.decode())
async def _burst(
base_url: str, key: str, models: _Models, calls: tuple[_Call, ...], *, tolerate_transport_errors: bool = False
) -> tuple[_Served, ...]:
async with httpx.AsyncClient(base_url=base_url, timeout=120, trust_env=False) as client:
results: Final = await asyncio.gather(
*(_send(client, key, models, call) for call in calls), return_exceptions=tolerate_transport_errors
)
for result in results:
assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result)
return tuple(result for result in results if isinstance(result, _Served))
def _calls(plan: tuple[tuple[Endpoint, bool, int], ...]) -> tuple[_Call, ...]:
return tuple(
_Call(endpoint=endpoint, stream=stream, marker=uuid.uuid4().hex)
for endpoint, stream, count in plan
for _ in range(count)
) # comprehension-ok: a flat plan expansion, one marker per planned call
def _first_data_frame(text: str) -> dict[str, JsonValue]:
line: Final = next(line for line in text.splitlines() if line.startswith("data: ") and "[DONE]" not in line)
return object_value(json.loads(line.removeprefix("data: ")))
def _served_id(served: _Served) -> str:
match served.call.endpoint, served.call.stream:
case "messages", True:
return message_id(parse_sse(served.text))
case "responses", True:
completed: Final = next(
event for event in parse_sse(served.text) if event_type(event) == "response.completed"
)
return string_value(object_value(completed.data["response"])["id"])
case "chat", True:
return string_value(_first_data_frame(served.text)["id"])
case _:
return string_value(object_value(json.loads(served.text))["id"])
def _assert_completed_on_the_second_attempt(served: _Served) -> None:
assert served.status == 200, (served.call, served.text)
assert _answer(served.call.marker) in served.text, (served.call, served.text)
assert "error" not in served.text.lower() or served.call.endpoint == "responses", (served.call, served.text)
match served.call.endpoint, served.call.stream:
case "messages", True:
events: Final = parse_sse(served.text)
assert event_types(events)[-1] == "message_stop", events
assert delta_text(events) == _answer(served.call.marker), events
assert message_id(events) == f"msg_{served.call.marker}_a2", events
case "messages", False:
assert _served_id(served) == f"msg_{served.call.marker}_a2", served.text
case "chat", _:
assert _served_id(served) == f"chatcmpl-{served.call.marker}-a2", served.text
case "responses", _:
assert f"msg_resp_{served.call.marker}_a2" in served.text, served.text
def _upstream_served_id(served: _Served) -> str:
match served.call.endpoint:
case "messages":
return f"msg_{served.call.marker}_a2"
case "chat":
return f"chatcmpl-{served.call.marker}-a2"
case "responses":
return f"resp_{served.call.marker}_a2"
def _routing_decoded_upstream_id(request_id: str) -> str | None:
encoded: Final = _ROUTING_ENCODED_ID.fullmatch(request_id)
if encoded is None:
return None
decoded: Final = base64.b64decode(encoded.group(1)).decode(errors="replace")
if not decoded.startswith("litellm:"):
return None
return decoded.rpartition("response_id:")[2]
def _names_of_row(request_id: str) -> frozenset[str]:
return frozenset(name for name in (request_id, _routing_decoded_upstream_id(request_id)) if name is not None)
def _names_of_served(served: _Served) -> frozenset[str]:
return frozenset({_served_id(served), _upstream_served_id(served)})
def _spend_ids(model: str, expected: int) -> tuple[str, ...]:
rows: Final = eventually(
lambda: read_rows('SELECT request_id, status FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)),
lambda found: len(found) >= expected,
seconds=90,
)
assert [row["status"] for row in rows] == ["success"] * len(rows), rows
return tuple(string_value(row["request_id"]) for row in rows)
def _assert_rows_name_each_served_response_once(model: str, served: tuple[_Served, ...]) -> None:
rows: Final = _spend_ids(model, len(served))
named: Final = tuple(
tuple(index for index, item in enumerate(served) if _names_of_served(item) & _names_of_row(request_id))
for request_id in rows
)
assert sorted(named) == [(index,) for index in range(len(served))], (model, named, rows)
def _assert_each_served_id_landed_exactly_once(models: _Models, served: tuple[_Served, ...]) -> None:
_assert_rows_name_each_served_response_once(
models.messages, tuple(item for item in served if item.call.endpoint == "messages")
)
_assert_rows_name_each_served_response_once(
models.openai, tuple(item for item in served if item.call.endpoint != "messages")
)
def _open_upstream_connections(pid: int, upstream: str) -> int:
port: Final = urlsplit(upstream).port
return sum(
1
for connection in psutil.Process(pid).net_connections(kind="tcp")
if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port
)
@pytest.mark.timeout(300)
async def test_worker_sigkill_mid_burst_leaves_the_sibling_retrying_pre_content_failures(
gateway: Gateway, tmp_path: Path
) -> None:
models: Final = _models()
calls: Final = _calls(
(("messages", True, 24), ("chat", False, 3), ("chat", True, 3), ("responses", False, 3), ("responses", True, 3))
)
release: Final = threading.Event()
held: Final[SimpleQueue[str]] = SimpleQueue()
upstream: Final = _Upstream(Attempts(), _drop_or_500, held, release)
with wire_server(upstream) as wire:
config: Final = _config(wire, tmp_path, models)
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned:
candidate: Final = owned.gateway
base_url: Final = str(candidate.client.base_url)
workers: Final = eventually(
lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())),
lambda pids: len(pids) == 2,
seconds=30,
)
burst: Final = asyncio.create_task(
_burst(base_url, candidate.key, models, calls, tolerate_transport_errors=True)
)
await asyncio.to_thread(eventually, held.qsize, lambda size: size == len(calls), 90)
async with httpx.AsyncClient(base_url=base_url, timeout=15, trust_env=False) as probe:
alive: Final = await probe.get("/health/liveliness")
assert alive.status_code == 200, alive.text
held_by: Final = MappingProxyType({pid: _open_upstream_connections(pid, wire.url) for pid in workers})
assert sum(held_by.values()) == len(calls), held_by
victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__)
victim: Final = psutil.Process(victim_pid)
victim.suspend()
victim.send_signal(signal.SIGKILL)
release.set()
served: Final = await burst
assert held_by[survivor_pid] >= 10, held_by
assert len(served) == held_by[survivor_pid], (held_by, len(served))
for item in served:
_assert_completed_on_the_second_attempt(item)
follow_up: Final = _Call(endpoint="messages", stream=True, marker=uuid.uuid4().hex)
(answered,) = await _burst(base_url, candidate.key, models, (follow_up,))
_assert_completed_on_the_second_attempt(answered)
_assert_each_served_id_landed_exactly_once(models, (*served, answered))
assert all(upstream.attempts.count(item.call.marker) == 2 for item in (*served, answered))
@pytest.mark.timeout(300)
async def test_outage_on_every_first_attempt_is_absorbed_by_the_deployment_budget(
gateway: Gateway, tmp_path: Path
) -> None:
models: Final = _models()
calls: Final = _calls(
(
("messages", True, 6),
("messages", False, 6),
("chat", True, 6),
("chat", False, 6),
("responses", True, 6),
("responses", False, 6),
)
)
statuses: Final = MappingProxyType(
{call.marker: _OUTAGE_STATUSES[index % len(_OUTAGE_STATUSES)] for index, call in enumerate(calls)}
)
upstream: Final = _Upstream(Attempts(), _outage(statuses))
with wire_server(upstream) as wire:
config: Final = _config(wire, tmp_path, models)
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned:
candidate: Final = owned.gateway
served: Final = await _burst(str(candidate.client.base_url), candidate.key, models, calls)
assert len(served) == len(calls)
for item in served:
_assert_completed_on_the_second_attempt(item)
_assert_each_served_id_landed_exactly_once(models, served)
assert all(upstream.attempts.count(call.marker) == 2 for call in calls)

View file

@ -0,0 +1,473 @@
import json
import os
import uuid
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from typing import Final, Literal
import anthropic
import httpx
import pytest
from integration._support.anthropic_sse import (
LIFECYCLE,
Attempts,
SseEvent,
delta_text,
dropping_reply,
error_frame,
error_type,
event_types,
message_id,
message_json,
message_stream,
parse_sse,
status_reply,
stream_reply,
user_prompt,
)
from integration._support.client import Gateway, eventually, object_value
from integration._support.database import read_rows
from integration._support.wire import Reply, Request, Wire, wire_server
from pydantic import JsonValue
from redis import Redis
_MODEL: Final = "claude-sonnet-4-5-20250929"
_API_KEY: Final = "synthetic-anthropic-key"
_TEXT: Final = "Hello"
FirstAttempt = Literal[
"drop_after_headers",
"drop_after_message_start",
"error_frame",
"error_frame_after_message_start",
"http_status",
"drop_after_content",
]
BudgetSource = Literal["deployment", "request"]
@dataclass(frozen=True, slots=True)
class _Failure:
kind: FirstAttempt
status: int = 500
def reply(self, served_id: str) -> Reply:
chunks: Final = message_stream(served_id, _MODEL, _TEXT)
match self.kind:
case "drop_after_headers":
return dropping_reply(chunks, abort_after=0)
case "drop_after_message_start":
return dropping_reply(chunks, abort_after=1)
case "error_frame":
return stream_reply((error_frame(self.status, f"scripted {self.status}"),))
case "error_frame_after_message_start":
return stream_reply((chunks[0], error_frame(self.status, f"scripted {self.status}")))
case "http_status":
return status_reply(self.status)
case "drop_after_content":
return dropping_reply(chunks, abort_after=3)
_MID_STREAM_FAILURES: Final = (
_Failure("drop_after_headers"),
_Failure("drop_after_message_start"),
_Failure("error_frame", 529),
_Failure("error_frame", 429),
_Failure("error_frame_after_message_start", 500),
)
_PRE_STREAM_STATUSES: Final = (529, 500, 429, 408, 409)
def _served_id(prompt: str, attempt: int) -> str:
return f"msg_{prompt}_a{attempt}"
def _prompt() -> str:
return "pre-content-" + uuid.uuid4().hex
def _upstream(
prompt: str, failure: _Failure, attempts: Attempts, *, failing_attempts: int = 1, stream: bool = True
) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert (request.method, request.target) == ("POST", "/v1/messages"), request
assert request.headers["x-api-key"] == _API_KEY, request.headers
body: Final = object_value(json.loads(request.body))
assert body["model"] == _MODEL, body
assert body.get("stream", False) is stream, body
assert user_prompt(body) == prompt, body
assert "num_retries" not in body, body
attempt: Final = attempts.record(prompt)
if attempt <= failing_attempts:
return failure.reply(_served_id(prompt, attempt))
if stream:
return stream_reply(message_stream(_served_id(prompt, attempt), _MODEL, _TEXT))
return Reply(body=message_json(_served_id(prompt, attempt), _MODEL, _TEXT))
return respond
def _body(
model: str, prompt: str, source: BudgetSource, *, stream: bool = True, budget: int = 1
) -> dict[str, JsonValue]:
return {
"model": model,
"max_tokens": 16,
"stream": stream,
"messages": [{"role": "user", "content": prompt}],
**({"num_retries": budget} if source == "request" else {}),
}
def _stream(gateway: Gateway, body: Mapping[str, JsonValue]) -> tuple[int, tuple[SseEvent, ...]]:
response: Final = gateway.request("POST", "/v1/messages", body)
return response.status_code, parse_sse(response.text)
def _success_rows(request_id: str) -> list[dict[str, JsonValue]]:
return eventually(
lambda: read_rows('SELECT status, model_group FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (request_id,)),
lambda rows: len(rows) >= 1,
seconds=70,
)
def _assert_completed_by_retry(
status: int, events: tuple[SseEvent, ...], prompt: str, model: str, wire: Wire, *, attempts: int = 2
) -> None:
assert status == 200, events
assert event_types(events) == LIFECYCLE, events
assert message_id(events) == _served_id(prompt, attempts), events
assert delta_text(events) == _TEXT, events
assert [request.target for request in wire.drain()] == ["/v1/messages"] * attempts
assert _success_rows(_served_id(prompt, attempts)) == [{"status": "success", "model_group": model}]
def _assert_rejected_before_content(
response: httpx.Response, wire: Wire, *, status: int, attempts: int, error: str
) -> None:
assert response.status_code == status, response.text
assert error in response.text, response.text
assert "content_block_delta" not in response.text, response.text
assert [request.target for request in wire.drain()] == ["/v1/messages"] * attempts
def _assert_stream_failed_after_message_start(
status: int, events: tuple[SseEvent, ...], wire: Wire, *, attempts: int
) -> None:
assert status == 200, events
types: Final = event_types(events)
assert types[0] == "message_start", events
assert types[-1] == "error", events
assert "content_block_delta" not in types, events
assert [request.target for request in wire.drain()] == ["/v1/messages"] * attempts
@pytest.mark.parametrize("source", ["deployment", "request"])
@pytest.mark.parametrize("failure", _MID_STREAM_FAILURES, ids=lambda failure: f"{failure.kind}-{failure.status}")
def test_stream_failing_before_content_is_retried_on_the_same_group_and_completes(
gateway: Gateway, failure: _Failure, source: BudgetSource
) -> None:
prompt: Final = _prompt()
attempts: Final = Attempts()
with wire_server(_upstream(prompt, failure, attempts)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model=f"anthropic/{_MODEL}",
api_base=wire.url,
api_key=_API_KEY,
**({"num_retries": 1} if source == "deployment" else {}),
)
status, events = _stream(gateway, _body(model, prompt, source))
_assert_completed_by_retry(status, events, prompt, model, wire)
@pytest.mark.parametrize("source", ["deployment", "request"])
@pytest.mark.parametrize("http_status", _PRE_STREAM_STATUSES)
def test_stream_rejected_before_it_opens_is_retried_and_completes(
gateway: Gateway, http_status: int, source: BudgetSource
) -> None:
prompt: Final = _prompt()
attempts: Final = Attempts()
with (
wire_server(_upstream(prompt, _Failure("http_status", http_status), attempts)) as wire,
gateway.scenario() as scenario,
):
model: Final = scenario.model(
model=f"anthropic/{_MODEL}",
api_base=wire.url,
api_key=_API_KEY,
**({"num_retries": 1} if source == "deployment" else {}),
)
status, events = _stream(gateway, _body(model, prompt, source))
_assert_completed_by_retry(status, events, prompt, model, wire)
def _sdk(gateway: Gateway) -> anthropic.Anthropic:
return anthropic.Anthropic(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0, timeout=60)
def _async_sdk(gateway: Gateway) -> anthropic.AsyncAnthropic:
return anthropic.AsyncAnthropic(
base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0, timeout=60
)
def test_anthropic_sdk_sync_stream_completes_after_a_drop_following_message_start(gateway: Gateway) -> None:
prompt: Final = _prompt()
attempts: Final = Attempts()
with (
wire_server(_upstream(prompt, _Failure("drop_after_message_start"), attempts)) as wire,
gateway.scenario() as scenario,
):
model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_API_KEY, num_retries=1)
with _sdk(gateway).messages.stream(
model=model, max_tokens=16, messages=[{"role": "user", "content": prompt}]
) as stream:
final: Final = stream.get_final_message()
assert final.id == _served_id(prompt, 2), final
assert [block.text for block in final.content if block.type == "text"] == [_TEXT], final
assert [request.target for request in wire.drain()] == ["/v1/messages"] * 2
assert _success_rows(final.id) == [{"status": "success", "model_group": model}]
async def test_anthropic_sdk_async_stream_completes_after_a_drop_following_message_start(gateway: Gateway) -> None:
prompt: Final = _prompt()
attempts: Final = Attempts()
with (
wire_server(_upstream(prompt, _Failure("drop_after_message_start"), attempts)) as wire,
gateway.scenario() as scenario,
):
model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_API_KEY)
async with _async_sdk(gateway).messages.stream(
model=model, max_tokens=16, messages=[{"role": "user", "content": prompt}], extra_body={"num_retries": 1}
) as stream:
final: Final = await stream.get_final_message()
assert final.id == _served_id(prompt, 2), final
assert [block.text for block in final.content if block.type == "text"] == [_TEXT], final
assert [request.target for request in wire.drain()] == ["/v1/messages"] * 2
def test_anthropic_sdk_sync_stream_completes_after_an_overloaded_error_frame(gateway: Gateway) -> None:
prompt: Final = _prompt()
attempts: Final = Attempts()
with (
wire_server(_upstream(prompt, _Failure("error_frame", 529), attempts)) as wire,
gateway.scenario() as scenario,
):
model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_API_KEY, num_retries=1)
with _sdk(gateway).messages.create(
model=model, max_tokens=16, messages=[{"role": "user", "content": prompt}], stream=True
) as stream:
events: Final = tuple(stream)
starts: Final = [event.message.id for event in events if event.type == "message_start"]
assert starts == [_served_id(prompt, 2)], events
assert [
event.delta.text
for event in events
if event.type == "content_block_delta" and event.delta.type == "text_delta"
] == [_TEXT]
assert [request.target for request in wire.drain()] == ["/v1/messages"] * 2
async def test_anthropic_sdk_async_stream_completes_after_a_529_before_the_stream_opens(gateway: Gateway) -> None:
prompt: Final = _prompt()
attempts: Final = Attempts()
with (
wire_server(_upstream(prompt, _Failure("http_status", 529), attempts)) as wire,
gateway.scenario() as scenario,
):
model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_API_KEY, num_retries=1)
async with _async_sdk(gateway).messages.stream(
model=model, max_tokens=16, messages=[{"role": "user", "content": prompt}]
) as stream:
final: Final = await stream.get_final_message()
assert final.id == _served_id(prompt, 2), final
assert [request.target for request in wire.drain()] == ["/v1/messages"] * 2
def test_stream_dropping_after_content_is_not_retried(gateway: Gateway) -> None:
prompt: Final = _prompt()
attempts: Final = Attempts()
with (
wire_server(_upstream(prompt, _Failure("drop_after_content"), attempts)) as wire,
gateway.scenario() as scenario,
):
model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_API_KEY, num_retries=1)
status, events = _stream(gateway, _body(model, prompt, "deployment"))
assert status == 200, events
assert event_types(events)[:3] == LIFECYCLE[:3], events
assert event_types(events)[-1] == "error", events
assert message_id(events) == _served_id(prompt, 1), events
assert delta_text(events) == _TEXT, events
assert [request.target for request in wire.drain()] == ["/v1/messages"]
def test_stream_rejected_with_401_before_it_opens_is_not_retried(gateway: Gateway) -> None:
prompt: Final = _prompt()
attempts: Final = Attempts()
with (
wire_server(_upstream(prompt, _Failure("http_status", 401), attempts)) as wire,
gateway.scenario() as scenario,
):
model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_API_KEY, num_retries=1)
response: Final = gateway.request("POST", "/v1/messages", _body(model, prompt, "deployment"))
assert response.status_code == 401, response.text
assert [request.target for request in wire.drain()] == ["/v1/messages"]
def test_stream_invalid_request_error_frame_is_not_retried(gateway: Gateway) -> None:
prompt: Final = _prompt()
attempts: Final = Attempts()
with (
wire_server(_upstream(prompt, _Failure("error_frame", 400), attempts)) as wire,
gateway.scenario() as scenario,
):
model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_API_KEY, num_retries=1)
status, events = _stream(gateway, _body(model, prompt, "deployment"))
assert status == 200, events
assert event_types(events) == ("error",), events
assert error_type(events) == "invalid_request_error", events
assert [request.target for request in wire.drain()] == ["/v1/messages"]
def test_request_num_retries_zero_turns_the_retry_off_for_a_deployment_with_a_budget(gateway: Gateway) -> None:
prompt: Final = _prompt()
attempts: Final = Attempts()
with (
wire_server(_upstream(prompt, _Failure("drop_after_message_start"), attempts)) as wire,
gateway.scenario() as scenario,
):
model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_API_KEY, num_retries=1)
status, events = _stream(gateway, {**_body(model, prompt, "deployment"), "num_retries": 0})
_assert_stream_failed_after_message_start(status, events, wire, attempts=1)
def test_always_dropping_upstream_is_attempted_once_per_budget_unit_plus_the_first_call(gateway: Gateway) -> None:
prompt: Final = _prompt()
attempts: Final = Attempts()
with (
wire_server(_upstream(prompt, _Failure("drop_after_message_start"), attempts, failing_attempts=99)) as wire,
gateway.scenario() as scenario,
):
model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_API_KEY, num_retries=2)
status, events = _stream(gateway, _body(model, prompt, "deployment"))
_assert_stream_failed_after_message_start(status, events, wire, attempts=3)
assert message_id(events) == _served_id(prompt, 3), events
def test_retry_rejected_before_it_opens_counts_against_the_same_budget(gateway: Gateway) -> None:
prompt: Final = _prompt()
attempts: Final = Attempts()
def respond(request: Request) -> Reply:
body: Final = object_value(json.loads(request.body))
assert user_prompt(body) == prompt, body
attempt: Final = attempts.record(prompt)
if attempt == 1:
return dropping_reply(message_stream(_served_id(prompt, attempt), _MODEL, _TEXT), abort_after=1)
return status_reply(529)
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_API_KEY, num_retries=1)
response: Final = gateway.request("POST", "/v1/messages", _body(model, prompt, "deployment"))
_assert_rejected_before_content(response, wire, status=500, attempts=2, error="error")
def test_non_stream_messages_rejected_with_500_is_retried_per_the_deployment_budget(gateway: Gateway) -> None:
prompt: Final = _prompt()
attempts: Final = Attempts()
with (
wire_server(_upstream(prompt, _Failure("http_status", 500), attempts, stream=False)) as wire,
gateway.scenario() as scenario,
):
model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_API_KEY, num_retries=1)
response: Final = gateway.request("POST", "/v1/messages", _body(model, prompt, "deployment", stream=False))
assert response.status_code == 200, response.text
payload: Final = object_value(json.loads(response.content))
assert payload["id"] == _served_id(prompt, 2), response.text
assert payload["content"] == [{"type": "text", "text": _TEXT}], response.text
assert [request.target for request in wire.drain()] == ["/v1/messages"] * 2
assert _success_rows(_served_id(prompt, 2)) == [{"status": "success", "model_group": model}]
def test_retried_stream_response_headers_name_the_attempt(gateway: Gateway) -> None:
prompt: Final = _prompt()
attempts: Final = Attempts()
with (
wire_server(_upstream(prompt, _Failure("drop_after_message_start"), attempts)) as wire,
gateway.scenario() as scenario,
):
model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_API_KEY, num_retries=1)
response: Final = gateway.request("POST", "/v1/messages", _body(model, prompt, "deployment"))
events: Final = parse_sse(response.text)
_assert_completed_by_retry(response.status_code, events, prompt, model, wire)
assert response.headers.get("x-litellm-attempted-retries") == "1", dict(response.headers)
assert response.headers.get("x-litellm-max-retries") == "1", dict(response.headers)
def test_two_drops_stamp_two_attempted_retries_on_the_response(gateway: Gateway) -> None:
prompt: Final = _prompt()
attempts: Final = Attempts()
with (
wire_server(_upstream(prompt, _Failure("drop_after_message_start"), attempts, failing_attempts=2)) as wire,
gateway.scenario() as scenario,
):
model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_API_KEY, num_retries=2)
response: Final = gateway.request("POST", "/v1/messages", _body(model, prompt, "deployment"))
events: Final = parse_sse(response.text)
assert response.status_code == 200, response.text
assert event_types(events) == LIFECYCLE, events
assert message_id(events) == _served_id(prompt, 3), events
assert [request.target for request in wire.drain()] == ["/v1/messages"] * 3
assert response.headers.get("x-litellm-attempted-retries") == "2", dict(response.headers)
assert response.headers.get("x-litellm-max-retries") == "2", dict(response.headers)
def test_retried_stream_spend_row_records_the_attempt_count(gateway: Gateway) -> None:
prompt: Final = _prompt()
attempts: Final = Attempts()
with (
wire_server(_upstream(prompt, _Failure("drop_after_message_start"), attempts)) as wire,
gateway.scenario() as scenario,
):
model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_API_KEY, num_retries=1)
status, events = _stream(gateway, _body(model, prompt, "deployment"))
_assert_completed_by_retry(status, events, prompt, model, wire)
rows: Final = read_rows(
"SELECT metadata->>'attempted_retries' AS attempted, metadata->>'max_retries' AS budget "
'FROM "LiteLLM_SpendLogs" WHERE request_id=%s',
(_served_id(prompt, 2),),
)
assert rows == [{"attempted": "1", "budget": "1"}], rows
def _string_values(cache: Redis) -> tuple[bytes, ...]:
keys: Final = tuple(key for key in cache.scan_iter(count=1000) if cache.type(key) == b"string")
return tuple(value for value in cache.mget(keys) if value is not None) if keys else ()
def _cached_somewhere(served_id: str) -> bool:
with Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"])) as cache:
return any(served_id.encode() in value for value in _string_values(cache))
def test_retried_stream_is_cached_and_the_identical_request_is_served_without_the_upstream(
gateway: Gateway,
) -> None:
prompt: Final = _prompt()
attempts: Final = Attempts()
with (
wire_server(_upstream(prompt, _Failure("drop_after_message_start"), attempts)) as wire,
gateway.scenario() as scenario,
):
model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_API_KEY, num_retries=1)
body: Final = _body(model, prompt, "deployment")
status, events = _stream(gateway, body)
_assert_completed_by_retry(status, events, prompt, model, wire)
eventually(lambda: _cached_somewhere(_served_id(prompt, 2)), bool, seconds=30)
replay_status, replay = _stream(gateway, body)
assert replay_status == 200, replay
assert message_id(replay) == _served_id(prompt, 2), replay
assert delta_text(replay) == _TEXT, replay
assert event_types(replay)[-1] == "message_stop", replay
assert wire.drain() == ()

View file

@ -0,0 +1,255 @@
import json
import uuid
from collections.abc import Callable
from typing import Final
import pytest
from integration._support.anthropic_sse import (
Attempts,
event_type,
event_types,
parse_sse,
user_prompt,
)
from integration._support.client import Gateway, eventually, object_value
from integration._support.database import read_rows
from integration._support.openai_wire import (
answering_model_discovery,
chat_reply,
openai_error,
posted_targets,
responses_reply,
)
from integration._support.wire import Reply, Request, Wire, wire_server
from pydantic import JsonValue, TypeAdapter
_TEXT: Final = "Hello"
_OPENAI_MODEL: Final = "gpt-4o-mini"
_PROVIDER_KEY: Final = "integration-provider-key"
_GEMINI_MODEL: Final = "gemini-2.5-flash"
_GEMINI_KEY: Final = "synthetic-gemini-key"
_OBJECTS: Final = TypeAdapter(list[dict[str, JsonValue]])
def _marker() -> str:
return "generic-retry-" + uuid.uuid4().hex
def _openai_upstream(
marker: str, target: str, attempts: Attempts, served: Callable[[int, bool], Reply]
) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert (request.method, request.target) == ("POST", target), request
assert request.headers["authorization"] == f"Bearer {_PROVIDER_KEY}", request.headers
body: Final = object_value(json.loads(request.body))
assert body["model"] == _OPENAI_MODEL, body
assert "num_retries" not in body, body
attempt: Final = attempts.record(marker)
if attempt == 1:
return openai_error(500)
return served(attempt, body.get("stream") is True)
return answering_model_discovery(respond)
def _responses_served(marker: str) -> Callable[[int, bool], Reply]:
def served(attempt: int, streamed: bool) -> Reply:
return responses_reply(f"resp_{marker}_a{attempt}", _OPENAI_MODEL, _TEXT, stream=streamed)
return served
def _chat_served(marker: str) -> Callable[[int, bool], Reply]:
def served(attempt: int, streamed: bool) -> Reply:
return chat_reply(f"chatcmpl-{marker}-a{attempt}", _OPENAI_MODEL, _TEXT, stream=streamed)
return served
def _success_rows(model: str) -> list[dict[str, JsonValue]]:
return eventually(
lambda: read_rows('SELECT request_id, status FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)),
lambda rows: len(rows) >= 1,
seconds=70,
)
def _assert_two_attempts_one_success(wire: Wire, target: str, model: str) -> None:
assert posted_targets(wire) == (target,) * 2
assert [row["status"] for row in _success_rows(model)] == ["success"]
def _responses_text(response_text: str, stream: bool) -> str:
if not stream:
payload: Final = object_value(json.loads(response_text))
content: Final = _OBJECTS.validate_python(_OBJECTS.validate_python(payload["output"])[0]["content"])
return str(content[0]["text"])
events: Final = parse_sse(response_text)
assert event_types(events)[-1] == "response.completed", events
return "".join(str(event.data["delta"]) for event in events if event_type(event) == "response.output_text.delta")
@pytest.mark.parametrize("stream", [False, True], ids=["non_stream", "stream"])
def test_responses_rejected_before_the_stream_opens_is_retried_per_the_deployment_budget(
gateway: Gateway, stream: bool
) -> None:
marker: Final = _marker()
attempts: Final = Attempts()
served: Final = _responses_served(marker)
with (
wire_server(_openai_upstream(marker, "/v1/responses", attempts, served)) as wire,
gateway.scenario() as scenario,
):
model: Final = scenario.model(api_base=wire.url + "/v1", num_retries=1)
response: Final = gateway.request("POST", "/v1/responses", {"model": model, "input": marker, "stream": stream})
assert response.status_code == 200, response.text
assert _responses_text(response.text, stream) == _TEXT, response.text
_assert_two_attempts_one_success(wire, "/v1/responses", model)
@pytest.mark.parametrize("stream", [False, True], ids=["non_stream", "stream"])
def test_chat_completions_control_keeps_retrying_per_the_deployment_budget(gateway: Gateway, stream: bool) -> None:
marker: Final = _marker()
attempts: Final = Attempts()
served: Final = _chat_served(marker)
with (
wire_server(_openai_upstream(marker, "/v1/chat/completions", attempts, served)) as wire,
gateway.scenario() as scenario,
):
model: Final = scenario.model(api_base=wire.url + "/v1", num_retries=1)
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": marker}], "stream": stream},
)
assert response.status_code == 200, response.text
assert f"chatcmpl-{marker}-a2" in response.text, response.text
assert _TEXT in response.text, response.text
_assert_two_attempts_one_success(wire, "/v1/chat/completions", model)
def test_responses_request_budget_still_wins_over_a_zero_deployment_budget(gateway: Gateway) -> None:
marker: Final = _marker()
attempts: Final = Attempts()
served: Final = _responses_served(marker)
with (
wire_server(_openai_upstream(marker, "/v1/responses", attempts, served)) as wire,
gateway.scenario() as scenario,
):
model: Final = scenario.model(api_base=wire.url + "/v1", num_retries=0)
response: Final = gateway.request("POST", "/v1/responses", {"model": model, "input": marker, "num_retries": 1})
assert response.status_code == 200, response.text
assert _responses_text(response.text, False) == _TEXT, response.text
_assert_two_attempts_one_success(wire, "/v1/responses", model)
def _vllm_passthrough_upstream(marker: str, attempts: Attempts) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert (request.method, request.target) == ("POST", "/v1/chat/completions"), request
body: Final = object_value(json.loads(request.body))
assert user_prompt(body) == marker, body
attempt: Final = attempts.record(marker)
if attempt == 1:
return openai_error(500)
return chat_reply(f"chatcmpl-{marker}-a{attempt}", _OPENAI_MODEL, _TEXT, stream=body.get("stream") is True)
return answering_model_discovery(respond)
@pytest.mark.parametrize("stream", [False, True], ids=["non_stream", "stream"])
def test_vllm_passthrough_rejected_before_it_opens_is_retried_per_the_deployment_budget(
gateway: Gateway, stream: bool
) -> None:
marker: Final = _marker()
attempts: Final = Attempts()
with wire_server(_vllm_passthrough_upstream(marker, attempts)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"hosted_vllm/{_OPENAI_MODEL}", api_base=wire.url + "/v1", num_retries=1)
response: Final = gateway.request(
"POST",
"/vllm/chat/completions",
{"model": model, "messages": [{"role": "user", "content": marker}], **({"stream": True} if stream else {})},
)
assert response.status_code == 200, response.text
assert f"chatcmpl-{marker}-a2" in response.text, response.text
assert _TEXT in response.text, response.text
assert posted_targets(wire) == ("/v1/chat/completions",) * 2
def _gemini_upstream(marker: str, attempts: Attempts) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert request.method == "POST", request
assert request.target.split("?")[0] == f"/models/{_GEMINI_MODEL}:generateContent", request.target
assert request.headers["x-goog-api-key"] == _GEMINI_KEY, request.headers
body: Final = object_value(json.loads(request.body))
assert body["contents"] == [{"role": "user", "parts": [{"text": marker}]}], body
attempt: Final = attempts.record(marker)
if attempt == 1:
return Reply(
status=500,
body=json.dumps({"error": {"code": 500, "message": "scripted", "status": "INTERNAL"}}).encode(),
)
return Reply(
body=json.dumps(
{
"candidates": [
{
"content": {"parts": [{"text": f"{_TEXT} a{attempt}"}], "role": "model"},
"finishReason": "STOP",
"index": 0,
}
],
"usageMetadata": {"promptTokenCount": 5, "candidatesTokenCount": 3, "totalTokenCount": 8},
"modelVersion": _GEMINI_MODEL,
}
).encode()
)
return respond
def test_gemini_generate_content_rejected_is_retried_per_the_deployment_budget(gateway: Gateway) -> None:
marker: Final = _marker()
attempts: Final = Attempts()
with wire_server(_gemini_upstream(marker, attempts)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model=f"gemini/{_GEMINI_MODEL}", api_base=wire.url, api_key=_GEMINI_KEY, num_retries=1
)
response: Final = gateway.request(
"POST",
f"/v1beta/models/{model}:generateContent",
{"contents": [{"role": "user", "parts": [{"text": marker}]}]},
)
assert response.status_code == 200, response.text
assert f"{_TEXT} a2" in response.text, response.text
assert [request.target.split("?")[0] for request in wire.drain()] == [
f"/models/{_GEMINI_MODEL}:generateContent"
] * 2
assert [row["status"] for row in _success_rows(model)] == ["success"]
def _fine_tuning_list_upstream(marker: str, attempts: Attempts) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert request.method == "GET", request
assert request.target.split("?")[0] == "/v1/fine_tuning/jobs", request.target
assert request.headers["authorization"] == f"Bearer {_PROVIDER_KEY}", request.headers
attempt: Final = attempts.record(marker)
if attempt == 1:
return openai_error(500)
return Reply(body=json.dumps({"object": "list", "data": [], "has_more": False}).encode())
return answering_model_discovery(respond)
def test_fine_tuning_jobs_list_rejected_is_retried_per_the_deployment_budget(gateway: Gateway) -> None:
marker: Final = _marker()
attempts: Final = Attempts()
with wire_server(_fine_tuning_list_upstream(marker, attempts)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(api_base=wire.url + "/v1", num_retries=1)
response: Final = gateway.request(
"GET", "/v1/fine_tuning/jobs", params={"target_model_names": model, "limit": "5"}
)
assert response.status_code == 200, response.text
assert object_value(json.loads(response.text))["data"] == [], response.text
assert [request.target.split("?")[0] for request in wire.drain() if request.method == "GET"] == [
"/v1/fine_tuning/jobs"
] * 2

View file

@ -0,0 +1,263 @@
import json
import uuid
from collections.abc import Callable, Iterator
from dataclasses import dataclass
from pathlib import Path
from typing import Final, Literal
import httpx
import pytest
import yaml
from integration._support.anthropic_sse import (
LIFECYCLE,
PING,
Attempts,
delta_text,
dropping_reply,
error_body,
error_frame,
event_types,
message_id,
message_stream,
parse_sse,
status_reply,
stream_reply,
user_prompt,
)
from integration._support.client import Gateway, gateway_from_environment, object_value
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, Wire, wire_server
from pydantic import JsonValue
_PRIMARY: Final = "claude-under-test"
_FALLBACK: Final = "claude-fallback"
_CONTEXT_WINDOW: Final = "claude-context-window"
_API_KEY: Final = "synthetic-anthropic-key"
_TEXT: Final = "Hello"
_ROUTER_BUDGET: Final = "audit-router-budget"
_STRING_BUDGET: Final = "audit-string-budget"
_WITH_FALLBACKS: Final = "audit-primary"
_FALLBACK_GROUP: Final = "audit-fallback"
_CW_FALLBACK_GROUP: Final = "audit-cw-fallback"
_LONELY: Final = "audit-lonely"
_POLICY: Final = "audit-policy"
_UPSTREAM_URL_PLACEHOLDER: Final = "upstream-url"
Behavior = Literal[
"drop-once", "drop-always", "hold-ping", "drop-then-too-long", "drop-then-401", "overloaded-frames-fallback-503"
]
pytestmark = pytest.mark.timeout(240)
def _served_id(backend: str, marker: str, attempt: int) -> str:
return f"msg_{backend}_{marker}_a{attempt}"
def _primary_reply(behavior: str, attempt: int, full: tuple[bytes, bytes, bytes, bytes]) -> Reply:
dropped: Final = dropping_reply(full, abort_after=1)
match behavior:
case "drop-once":
return dropped if attempt == 1 else stream_reply(full)
case "drop-always":
return dropped
case "hold-ping":
return stream_reply((full[0], PING, full[1] + full[2] + full[3]), pause=0.25)
case "drop-then-too-long":
if attempt == 1:
return dropped
return Reply(status=400, body=error_body(400, "prompt is too long: 250000 tokens > 200000 maximum"))
case "drop-then-401":
return dropped if attempt == 1 else status_reply(401)
case "overloaded-frames-fallback-503":
return stream_reply((error_frame(529, "scripted overloaded"),))
raise AssertionError(behavior)
def _fallback_reply(behavior: str, full: tuple[bytes, bytes, bytes, bytes]) -> Reply:
if behavior == "overloaded-frames-fallback-503":
return status_reply(503)
return stream_reply(full)
def _respond(attempts: Attempts) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert (request.method, request.target) == ("POST", "/v1/messages"), request
assert request.headers["x-api-key"] == _API_KEY, request.headers
body: Final = object_value(json.loads(request.body))
assert "num_retries" not in body, body
backend: Final = str(body["model"])
behavior, marker = user_prompt(body).split(":", 1)
attempt: Final = attempts.record(f"{backend}:{marker}")
full: Final = message_stream(_served_id(backend, marker, attempt), backend, _TEXT)
if backend != _PRIMARY:
return _fallback_reply(behavior, full)
return _primary_reply(behavior, attempt, full)
return respond
def _deployment(name: str, backend: str, **extra: JsonValue) -> dict[str, JsonValue]:
return {
"model_name": name,
"litellm_params": {
"model": f"anthropic/{backend}",
"api_base": _UPSTREAM_URL_PLACEHOLDER,
"api_key": _API_KEY,
**extra,
},
}
def _config(wire: Wire, directory: Path) -> Path:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["model_list"] = [
_deployment(_ROUTER_BUDGET, _PRIMARY),
_deployment(_STRING_BUDGET, _PRIMARY, num_retries="2"),
_deployment(_WITH_FALLBACKS, _PRIMARY, num_retries=1),
_deployment(_FALLBACK_GROUP, _FALLBACK),
_deployment(_CW_FALLBACK_GROUP, _CONTEXT_WINDOW),
_deployment(_LONELY, _PRIMARY, num_retries=1),
_deployment(_POLICY, _PRIMARY),
]
config["router_settings"] = {
"num_retries": 1,
"disable_cooldowns": True,
"fallbacks": [{_WITH_FALLBACKS: [_FALLBACK_GROUP]}],
"context_window_fallbacks": [{_WITH_FALLBACKS: [_CW_FALLBACK_GROUP]}],
"model_group_retry_policy": {_POLICY: {"DefaultRetries": 2}},
}
path: Final = directory / "messages-retry-budget-sources.yaml"
path.write_text(yaml.safe_dump(config).replace(_UPSTREAM_URL_PLACEHOLDER, wire.url))
return path
@dataclass(frozen=True, slots=True)
class _Rig:
proxy: Gateway
attempts: Attempts
def stream(self, model: str, behavior: Behavior, marker: str, **extra: JsonValue) -> httpx.Response:
return self.proxy.request(
"POST",
"/v1/messages",
{
"model": model,
"max_tokens": 16,
"stream": True,
"messages": [{"role": "user", "content": f"{behavior}:{marker}"}],
**extra,
},
)
def attempts_on(self, backend: str, marker: str) -> int:
return self.attempts.count(f"{backend}:{marker}")
@pytest.fixture(scope="module")
def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Rig]:
attempts: Final = Attempts()
directory: Final = tmp_path_factory.mktemp("messages-retry-budget-sources")
with gateway_from_environment() as gateway, wire_server(_respond(attempts)) as wire:
with owned_proxy(gateway, directory, {}, config=_config(wire, directory), workers=2) as proxy:
yield _Rig(proxy, attempts)
def _assert_completed(response: httpx.Response, served_id: str) -> None:
assert response.status_code == 200, response.text
events: Final = parse_sse(response.text)
assert event_types(events) == LIFECYCLE, events
assert message_id(events) == served_id, events
assert delta_text(events) == _TEXT, events
def _assert_failed_after_message_start(response: httpx.Response, served_id: str) -> None:
assert response.status_code == 200, response.text
events: Final = parse_sse(response.text)
types: Final = event_types(events)
assert types[0] == "message_start", events
assert types[-1] == "error", events
assert "content_block_delta" not in types, events
assert message_id(events) == served_id, events
def test_router_num_retries_governs_a_group_without_its_own_budget(rig: _Rig) -> None:
marker: Final = uuid.uuid4().hex
_assert_completed(rig.stream(_ROUTER_BUDGET, "drop-once", marker), _served_id(_PRIMARY, marker, 2))
assert rig.attempts_on(_PRIMARY, marker) == 2
def test_a_digit_string_deployment_budget_is_honored_as_a_number(rig: _Rig) -> None:
marker: Final = uuid.uuid4().hex
response: Final = rig.stream(_STRING_BUDGET, "drop-always", marker)
_assert_failed_after_message_start(response, _served_id(_PRIMARY, marker, 3))
assert rig.attempts_on(_PRIMARY, marker) == 3
def test_request_num_retries_zero_turns_the_router_budget_off(rig: _Rig) -> None:
marker: Final = uuid.uuid4().hex
response: Final = rig.stream(_ROUTER_BUDGET, "drop-once", marker, num_retries=0)
_assert_failed_after_message_start(response, _served_id(_PRIMARY, marker, 1))
assert rig.attempts_on(_PRIMARY, marker) == 1
def test_lifecycle_frames_are_held_until_content_while_pings_go_out_live(rig: _Rig) -> None:
marker: Final = uuid.uuid4().hex
response: Final = rig.stream(_ROUTER_BUDGET, "hold-ping", marker)
assert response.status_code == 200, response.text
events: Final = parse_sse(response.text)
assert event_types(events) == ("ping", *LIFECYCLE), events
assert message_id(events) == _served_id(_PRIMARY, marker, 1), events
assert delta_text(events) == _TEXT, events
assert rig.attempts_on(_PRIMARY, marker) == 1
def test_a_retry_policy_default_retries_sets_the_budget_for_a_pre_content_drop(rig: _Rig) -> None:
marker: Final = uuid.uuid4().hex
response: Final = rig.stream(_POLICY, "drop-always", marker)
_assert_failed_after_message_start(response, _served_id(_PRIMARY, marker, 3))
assert rig.attempts_on(_PRIMARY, marker) == 3
def test_request_num_retries_zero_turns_a_retry_policy_off(rig: _Rig) -> None:
marker: Final = uuid.uuid4().hex
response: Final = rig.stream(_POLICY, "drop-once", marker, num_retries=0)
_assert_failed_after_message_start(response, _served_id(_PRIMARY, marker, 1))
assert rig.attempts_on(_PRIMARY, marker) == 1
def test_fallbacks_run_only_after_the_same_group_budget_is_spent(rig: _Rig) -> None:
marker: Final = uuid.uuid4().hex
response: Final = rig.stream(_WITH_FALLBACKS, "drop-always", marker)
_assert_completed(response, _served_id(_FALLBACK, marker, 1))
assert response.headers.get("x-litellm-attempted-fallbacks") == "1", dict(response.headers)
assert response.headers.get("x-litellm-model-group") == _FALLBACK_GROUP, dict(response.headers)
assert rig.attempts_on(_PRIMARY, marker) == 2
assert rig.attempts_on(_FALLBACK, marker) == 1
def test_a_retry_raising_a_context_window_error_reaches_the_context_window_fallback(rig: _Rig) -> None:
marker: Final = uuid.uuid4().hex
_assert_completed(rig.stream(_WITH_FALLBACKS, "drop-then-too-long", marker), _served_id(_CONTEXT_WINDOW, marker, 1))
assert rig.attempts_on(_PRIMARY, marker) == 2
assert rig.attempts_on(_FALLBACK, marker) == 0
assert rig.attempts_on(_CONTEXT_WINDOW, marker) == 1
def test_a_retry_rejected_with_401_ends_the_retries_and_reaches_the_client_unchanged(rig: _Rig) -> None:
marker: Final = uuid.uuid4().hex
response: Final = rig.stream(_LONELY, "drop-then-401", marker)
assert response.status_code == 401, response.text
assert "authentication_error" in response.text, response.text
assert "content_block_delta" not in response.text, response.text
assert rig.attempts_on(_PRIMARY, marker) == 2
def test_overloaded_frames_whose_fallback_fails_answer_the_mapped_internal_server_error(rig: _Rig) -> None:
marker: Final = uuid.uuid4().hex
response: Final = rig.stream(_WITH_FALLBACKS, "overloaded-frames-fallback-503", marker)
assert response.status_code == 500, response.text
assert "content_block_delta" not in response.text, response.text
assert rig.attempts_on(_PRIMARY, marker) == 2
assert rig.attempts_on(_FALLBACK, marker) == 2

View file

@ -4148,3 +4148,43 @@ def test_tool_call_is_rebuilt_as_server_tool_use_only_with_a_stored_result(
from litellm.llms.anthropic.common_utils import tool_call_is_rebuilt_as_server_tool_use
assert tool_call_is_rebuilt_as_server_tool_use(tool_call_id, provider_specific_fields) is rebuilt
def _pre_stream_exception_for(error_type: str, message: str, status_code: int, model: str) -> Exception:
from litellm.litellm_core_utils.exception_mapping_utils import exception_type
from litellm.llms.anthropic.common_utils import AnthropicError
body: Final = json.dumps({"type": "error", "error": {"type": error_type, "message": message}})
with pytest.raises(Exception, match=message) as raised:
exception_type(
model=model,
original_exception=AnthropicError(status_code=status_code, message=body),
custom_llm_provider="anthropic",
)
return raised.value
@pytest.mark.parametrize(
"error_type",
["overloaded_error", "api_error", "timeout_error", "rate_limit_error", "invalid_request_error", "never_seen_error"],
)
def test_anthropic_error_frame_exception_matches_the_pre_stream_mapping_for_that_frame(error_type: str) -> None:
from litellm.llms.anthropic.common_utils import ANTHROPIC_ERROR_STATUS_CODE_MAP, anthropic_error_frame_exception
status_code: Final = ANTHROPIC_ERROR_STATUS_CODE_MAP.get(error_type, 500)
pre_stream: Final = _pre_stream_exception_for(error_type, "upstream said no", status_code, "claude-sonnet-4-5")
error: Final = anthropic_error_frame_exception(error_type, "upstream said no", status_code, "claude-sonnet-4-5")
assert type(error) is type(pre_stream)
assert getattr(error, "status_code", None) == getattr(pre_stream, "status_code", None)
assert "upstream said no" in str(error)
def test_anthropic_error_frame_exception_classes_an_overloaded_frame_as_internal_server_error() -> None:
import litellm
from litellm.llms.anthropic.common_utils import anthropic_error_frame_exception
error: Final = anthropic_error_frame_exception("overloaded_error", "Overloaded", 503, "claude-sonnet-4-5")
assert type(error) is litellm.InternalServerError

View file

@ -1,5 +1,6 @@
import json
from datetime import datetime, timedelta
from types import MappingProxyType
from typing import Final, NoReturn
from unittest.mock import MagicMock, patch
@ -10,13 +11,21 @@ import litellm
from litellm.litellm_core_utils import get_llm_provider_logic
from litellm.router_utils.cooldown_handlers import mark_advisor_orchestration_failure
from litellm.router_utils.fallback_event_handlers import (
MID_STREAM_FALLBACK_CONTROLS_KEY,
AttemptedFallbackTargets,
MidStreamFallbackControls,
_trigger_cooldown_for_failed_deployment,
fallback_attempt_key,
attempted_retries_for_request,
committed_retry_budget_for_request,
carry_over_routed_deployment,
clear_pre_routing_selection,
fallback_attempt_key,
get_fallback_model_group,
get_pre_routing_selection,
mid_stream_retry_kwargs,
record_pre_routing_selection,
record_retry_attempt,
routed_deployment_id,
run_async_fallback,
)
@ -1449,3 +1458,98 @@ def test_get_fallback_model_group_never_resolves_a_provider_without_a_prefixed_k
assert get_fallback_model_group(fallbacks=fallbacks, model_group="my-alias") == (["gpt-5.5-mini"], 1)
resolver.assert_not_called()
def test_mid_stream_retry_kwargs_strips_what_the_retry_wrapper_pops_and_keeps_the_controls_carrier():
def generic_function(**kwargs) -> None:
return None
def attempt(**kwargs) -> None:
return None
controls = MidStreamFallbackControls(MappingProxyType({"num_retries": 3}))
litellm_metadata = {"model_group": "glm"}
hop_kwargs = {
"model": "glm",
"original_generic_function": generic_function,
"original_function": attempt,
"fallbacks": [{"glm": ["fb"]}],
"context_window_fallbacks": [],
"content_policy_fallbacks": [],
"num_retries": 3,
"model_group_retry_policy": {},
"stream": True,
"litellm_metadata": litellm_metadata,
MID_STREAM_FALLBACK_CONTROLS_KEY: controls,
}
retry_kwargs = mid_stream_retry_kwargs(hop_kwargs)
assert retry_kwargs == {
"model": "glm",
"original_generic_function": generic_function,
"stream": True,
"litellm_metadata": litellm_metadata,
MID_STREAM_FALLBACK_CONTROLS_KEY: controls,
}
assert retry_kwargs["litellm_metadata"] is litellm_metadata
@pytest.mark.parametrize(
"kwargs,expected",
[
pytest.param({"litellm_metadata": {"attempted_retries": 2}, "metadata": {"attempted_retries": 5}}, 2, id="litellm_metadata-wins"),
pytest.param({"metadata": {"attempted_retries": 1}}, 1, id="metadata-bucket"),
pytest.param({"litellm_metadata": {"attempted_retries": "2"}}, 0, id="string-is-not-a-count"),
pytest.param({"litellm_metadata": {"attempted_retries": -1}}, 0, id="negative-is-not-a-count"),
pytest.param({"litellm_metadata": {}}, 0, id="unstamped"),
pytest.param({}, 0, id="no-bucket"),
],
)
def test_attempted_retries_for_request_reads_the_request_bucket(kwargs, expected):
assert attempted_retries_for_request(kwargs) == expected
def test_record_retry_attempt_stamps_the_bucket_the_retry_wrapper_reads():
kwargs = {"litellm_metadata": {"attempted_retries": 0, "max_retries": 2}, "metadata": {}}
record_retry_attempt(kwargs, attempted_retries=1, max_retries=2)
assert kwargs["litellm_metadata"] == {"attempted_retries": 1, "max_retries": 2}
assert kwargs["metadata"] == {}
assert attempted_retries_for_request(kwargs) == 1
assert committed_retry_budget_for_request(kwargs) == 2
@pytest.mark.parametrize(
"kwargs,expected",
[
pytest.param({"litellm_metadata": {"attempted_retries": 1, "max_retries": 3}}, 3, id="committed-by-a-retry"),
pytest.param({"litellm_metadata": {"attempted_retries": 0, "max_retries": 3}}, None, id="stamped-before-any-retry"),
pytest.param({"litellm_metadata": {"attempted_retries": 1, "max_retries": "3"}}, None, id="string-is-not-a-budget"),
pytest.param({"litellm_metadata": {"attempted_retries": 1}}, None, id="no-budget"),
pytest.param({}, None, id="no-bucket"),
],
)
def test_committed_retry_budget_for_request_is_the_budget_a_retry_stamped(kwargs, expected):
assert committed_retry_budget_for_request(kwargs) == expected
def test_carry_over_routed_deployment_copies_model_info_into_the_snapshot():
live_kwargs = {"litellm_metadata": {"model_info": {"id": "dep-1"}, "deployment": "anthropic/glm-a"}}
snapshot = {"litellm_metadata": {"model_group": "glm"}}
carry_over_routed_deployment(live_kwargs=live_kwargs, snapshot=snapshot)
assert snapshot["litellm_metadata"] == {"model_group": "glm", "model_info": {"id": "dep-1"}}
assert snapshot["litellm_metadata"]["model_info"] is not live_kwargs["litellm_metadata"]["model_info"]
assert routed_deployment_id(snapshot) == "dep-1"
def test_carry_over_routed_deployment_leaves_a_snapshot_without_a_bucket_alone():
snapshot = {"model": "glm"}
carry_over_routed_deployment(live_kwargs={"litellm_metadata": {"model_info": {"id": "dep-1"}}}, snapshot=snapshot)
assert snapshot == {"model": "glm"}
assert routed_deployment_id(snapshot) is None

View file

@ -2,6 +2,7 @@ import asyncio
import copy
import functools
import gc
import itertools
import json
import logging
import os
@ -58,7 +59,15 @@ from litellm.router_utils.cooldown_handlers import _async_get_cooldown_deploymen
from litellm.router_utils.fallback_event_handlers import DISABLE_FALLBACKS_METADATA_KEY
from litellm.router_utils.router_callbacks.track_deployment_metrics import get_deployment_successes_for_current_minute
from litellm.types.llms.openai import ChatCompletionRequest
from litellm.types.router import Deployment, DeploymentTypedDict, LiteLLM_Params, ModelInfo, PreRoutingHookResponse, RetryPolicy
from litellm.types.router import (
CustomRoutingStrategyBase,
Deployment,
DeploymentTypedDict,
LiteLLM_Params,
ModelInfo,
PreRoutingHookResponse,
RetryPolicy,
)
def test_update_kwargs_does_not_mutate_defaults_and_merges_metadata():
@ -13962,7 +13971,9 @@ def _anthropic_messages_make_wrapper() -> FallbackAwareAnthropicMessagesStream:
def _anthropic_messages_make_router(**router_kwargs) -> Router:
"""A fallback-only router: no same-group retries unless a test asks for them."""
router_kwargs.setdefault("fallbacks", [{"primary": ["fallback"]}])
router_kwargs.setdefault("num_retries", 0)
return Router(
model_list=[
{
@ -14776,7 +14787,8 @@ async def test_anthropic_messages_fallback_on_pre_first_chunk_error_event():
"""Regression for #24004: a retriable SSE `event: error` frame
(overloaded_error/internal_server_error) that arrives before any real
content must trigger the router's fallback chain instead of passing
through to the client silently."""
through to the client silently. The frame carries the error a 529 answer maps to, an InternalServerError,
so a failed fallback answers the status every other litellm path gives an overload."""
router = _anthropic_messages_make_router()
source = _AnthropicMessagesFakeByteStream([_anthropic_messages_overloaded_error_chunk()])
fallback_stream = _AnthropicMessagesFallbackByteStream([_anthropic_messages_content_chunk("fallback answer")])
@ -14796,7 +14808,8 @@ async def test_anthropic_messages_fallback_on_pre_first_chunk_error_event():
mock_fallback.assert_awaited_once()
raised = mock_fallback.await_args.kwargs["e"]
assert isinstance(raised, MidStreamFallbackError)
assert raised.status_code == 503
assert isinstance(raised.original_exception, litellm.InternalServerError)
assert raised.status_code == 500
assert raised.is_pre_first_chunk is True
assert source.closed is True
@ -15120,6 +15133,707 @@ async def test_anthropic_messages_hop_stream_failure_reaches_second_fallback_ent
assert b"overloaded_error" not in body
_ANTHROPIC_MESSAGES_RETRY_GROUP: Final = ("anthropic/glm-a", "anthropic/glm-b")
def _anthropic_messages_retry_router(
num_retries: int,
deployment_params: Mapping[str, object] | None = None,
fallbacks: list[dict[str, list[str]]] | None = None,
context_window_fallbacks: list[dict[str, list[str]]] | None = None,
retry_policy: RetryPolicy | None = None,
) -> Router:
"""Two deployments in the group, so a same-group retry waits for no backoff; no fallbacks unless asked."""
group_deployments = [
{"model_name": "glm", "litellm_params": {"model": model, "api_key": "sk-test", **(deployment_params or {})}}
for model in _ANTHROPIC_MESSAGES_RETRY_GROUP
]
return Router(
model_list=[
*group_deployments,
{"model_name": "fb", "litellm_params": {"model": "anthropic/fb-model", "api_key": "sk-test"}},
{"model_name": "cw", "litellm_params": {"model": "anthropic/cw-model", "api_key": "sk-test"}},
],
num_retries=num_retries,
fallbacks=fallbacks or [],
context_window_fallbacks=context_window_fallbacks or [],
retry_policy=retry_policy,
)
class _AnthropicMessagesScriptedProvider:
"""Stands in for litellm.anthropic_messages: answers each call with the next scripted stream and records
the deployment it was routed to plus the retry counters the router stamped for that attempt."""
def __init__(self, *streams) -> None:
self._streams = list(streams)
self.calls: list[tuple[str, object, object]] = []
async def __call__(self, **kwargs):
litellm_metadata = kwargs.get("litellm_metadata") or {}
self.calls.append((kwargs["model"], litellm_metadata.get("attempted_retries"), litellm_metadata.get("max_retries")))
assert self._streams, "provider called more times than scripted"
return self._streams.pop(0)()
def _anthropic_messages_transport_drop(original_exception: Exception | None = None) -> MidStreamFallbackError:
"""What the completion bridge raises when the upstream closes the connection before any content."""
return MidStreamFallbackError(
message="Connection closed",
model="glm",
llm_provider="databricks",
original_exception=original_exception
or litellm.APIConnectionError(message="Connection closed", llm_provider="databricks", model="glm"),
is_pre_first_chunk=True,
)
def _anthropic_messages_dropped_before_content():
return _AnthropicMessagesRaisingByteStream([_anthropic_messages_message_start_chunk()], _anthropic_messages_transport_drop())
def _anthropic_messages_bridge_error_chunk() -> bytes:
from litellm.anthropic_interface.exceptions.exception_mapping_utils import anthropic_error_sse_frame
return anthropic_error_sse_frame(status_code=500, raw_message="Connection closed").encode()
def _anthropic_messages_retried_stream():
return _AnthropicMessagesFakeByteStream(
[_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("pong")]
)
async def _anthropic_messages_drain_into(stream, received: list) -> None:
async for chunk in stream:
received.append(chunk)
async def _anthropic_messages_stream_through_router(router: Router, provider, **request_kwargs):
return await router._aanthropic_messages_with_streaming_fallbacks(
original_function=provider,
model="glm",
stream=True,
messages=[{"role": "user", "content": "ping"}],
max_tokens=16,
**request_kwargs,
)
_ANTHROPIC_MESSAGES_PRE_CONTENT_DROPS: Final = (
pytest.param(_anthropic_messages_dropped_before_content, id="bridge-raises-before-content"),
pytest.param(
lambda: _AnthropicMessagesFakeByteStream(
[_anthropic_messages_message_start_chunk(), _anthropic_messages_bridge_error_chunk()]
),
id="bridge-error-frame",
),
pytest.param(
lambda: _AnthropicMessagesFakeByteStream(
[_anthropic_messages_message_start_chunk(), _anthropic_messages_overloaded_error_chunk()]
),
id="provider-overloaded-frame",
),
)
@pytest.mark.asyncio
@pytest.mark.parametrize("dropped_stream", _ANTHROPIC_MESSAGES_PRE_CONTENT_DROPS)
async def test_anthropic_messages_stream_dropped_before_content_is_retried_within_the_group(dropped_stream):
"""Issue #44238: a /v1/messages stream the provider dropped before any content was answered after a
single upstream attempt, num_retries never applied. The drop is retried within the model group, with
the retry counters continuing the request's count, and the client sees one message lifecycle."""
router = _anthropic_messages_retry_router(num_retries=2)
provider = _AnthropicMessagesScriptedProvider(dropped_stream, _anthropic_messages_retried_stream)
stream = await _anthropic_messages_stream_through_router(router, provider)
body = [chunk async for chunk in stream]
assert body == [_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("pong")]
assert all(model in _ANTHROPIC_MESSAGES_RETRY_GROUP for model, _, _ in provider.calls)
assert [(attempted, budget) for _, attempted, budget in provider.calls] == [(0, 2), (1, 2)]
@pytest.mark.asyncio
async def test_anthropic_messages_stream_dropped_after_content_keeps_the_error_and_is_not_retried():
"""A drop once content reached the client cannot be retried without a second overlapping message
lifecycle, so it keeps surfacing the provider's error after a single attempt."""
router = _anthropic_messages_retry_router(num_retries=2)
drop = _anthropic_messages_transport_drop()
provider = _AnthropicMessagesScriptedProvider(
lambda: _AnthropicMessagesRaisingByteStream(
[_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("par")], drop
)
)
stream = await _anthropic_messages_stream_through_router(router, provider)
received = []
with pytest.raises(litellm.APIConnectionError) as raised:
await _anthropic_messages_drain_into(stream, received)
assert raised.value is drop.original_exception
assert received == [_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("par")]
assert len(provider.calls) == 1
@pytest.mark.asyncio
async def test_anthropic_messages_retries_stop_at_num_retries_and_raise_the_last_drop():
"""Every retry's own stream continues the same count, so a group that keeps dropping is tried
exactly 1 + num_retries times before the provider's error reaches the client."""
router = _anthropic_messages_retry_router(num_retries=2)
provider = _AnthropicMessagesScriptedProvider(
_anthropic_messages_dropped_before_content,
_anthropic_messages_dropped_before_content,
_anthropic_messages_dropped_before_content,
)
stream = await _anthropic_messages_stream_through_router(router, provider)
with pytest.raises(litellm.APIConnectionError):
[chunk async for chunk in stream]
assert [(attempted, budget) for _, attempted, budget in provider.calls] == [(0, 2), (1, 2), (2, 2)]
@pytest.mark.asyncio
async def test_anthropic_messages_retries_run_out_before_the_fallback_chain_is_consulted():
"""Same-group retries come first; the fallback group is reached only once num_retries is spent."""
router = _anthropic_messages_retry_router(num_retries=1, fallbacks=[{"glm": ["fb"]}])
provider = _AnthropicMessagesScriptedProvider(
_anthropic_messages_dropped_before_content,
_anthropic_messages_dropped_before_content,
lambda: _AnthropicMessagesFakeByteStream(
[_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("from fb")]
),
)
stream = await _anthropic_messages_stream_through_router(router, provider)
body = [chunk async for chunk in stream]
assert body == [_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("from fb")]
assert [model in _ANTHROPIC_MESSAGES_RETRY_GROUP for model, _, _ in provider.calls] == [True, True, False]
assert provider.calls[-1][0] == "anthropic/fb-model"
def _anthropic_messages_fb_deployment_hidden_params() -> dict:
return {"model_id": "fb-deployment", "additional_headers": {"x-litellm-model-group": "fb"}}
@pytest.mark.asyncio
async def test_anthropic_messages_fallback_after_exhausted_retries_attributes_the_response_to_the_fallback_deployment():
"""The retry's stream carries a wrapper of its own, so a fallback it makes before its first byte must reach
the wrapper the proxy reads headers off: the response names the deployment that served it, not the primary."""
router = _anthropic_messages_retry_router(num_retries=1, fallbacks=[{"glm": ["fb"]}])
provider = _AnthropicMessagesScriptedProvider(
_anthropic_messages_dropped_before_content,
_anthropic_messages_dropped_before_content,
lambda: _AnthropicMessagesFallbackByteStream(
[_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("from fb")],
hidden_params=_anthropic_messages_fb_deployment_hidden_params(),
),
)
stream = await _anthropic_messages_stream_through_router(router, provider)
body = [chunk async for chunk in stream]
assert body == [_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("from fb")]
assert stream._hidden_params["model_id"] == "fb-deployment"
assert stream._hidden_params["additional_headers"]["x-litellm-model-group"] == "fb"
assert stream._hidden_params["additional_headers"]["x-litellm-attempted-fallbacks"] == 1
def test_anthropic_messages_wrapper_follows_the_attribution_of_a_source_that_fell_back():
inner = FallbackAwareAnthropicMessagesStream(_anthropic_messages_empty_generator(), object())
outer = FallbackAwareAnthropicMessagesStream(_anthropic_messages_empty_generator(), inner)
fallback = _AnthropicMessagesFallbackByteStream([], hidden_params=_anthropic_messages_fb_deployment_hidden_params())
outer.follow_source_attribution()
assert "model_id" not in outer._hidden_params
inner.merge_fallback_hidden_params(*Router._prepare_fallback_hidden_params(fallback))
inner.adopt_fallback_source(fallback)
outer.follow_source_attribution()
assert outer._hidden_params["model_id"] == "fb-deployment"
assert outer._hidden_params["additional_headers"]["x-litellm-model-group"] == "fb"
@pytest.mark.asyncio
@pytest.mark.parametrize("drops", [1, 2])
async def test_anthropic_messages_mid_stream_retries_are_counted_in_the_response_retry_headers(drops: int):
"""A retry made after the stream opened never passes through async_function_with_retries, so the wrapper
stamps the retry headers that path would have and the client reads them along with the first byte."""
router = _anthropic_messages_retry_router(num_retries=2)
provider = _AnthropicMessagesScriptedProvider(
*([_anthropic_messages_dropped_before_content] * drops), _anthropic_messages_retried_stream
)
stream = await _anthropic_messages_stream_through_router(router, provider)
first = await stream.__anext__()
headers = stream._hidden_params["additional_headers"]
assert first == _anthropic_messages_message_start_chunk()
assert (headers["x-litellm-attempted-retries"], headers["x-litellm-max-retries"]) == (drops, 2)
assert [chunk async for chunk in stream] == [_anthropic_messages_content_chunk("pong")]
def _anthropic_messages_raise_authentication_error():
raise litellm.AuthenticationError(message="invalid api key", llm_provider="anthropic", model="glm")
@pytest.mark.asyncio
async def test_anthropic_messages_retry_raising_a_non_retriable_error_is_handed_to_the_fallback_chain():
"""A retry that fails before its stream opens with an error no retry covers ends the retries and reaches
the fallback group the way a pre-stream failure does, instead of surfacing as the client's error."""
router = _anthropic_messages_retry_router(num_retries=2, fallbacks=[{"glm": ["fb"]}])
provider = _AnthropicMessagesScriptedProvider(
_anthropic_messages_dropped_before_content,
_anthropic_messages_raise_authentication_error,
lambda: _AnthropicMessagesFakeByteStream(
[_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("from fb")]
),
)
stream = await _anthropic_messages_stream_through_router(router, provider)
body = [chunk async for chunk in stream]
assert body == [_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("from fb")]
assert [model in _ANTHROPIC_MESSAGES_RETRY_GROUP for model, _, _ in provider.calls] == [True, True, False]
def _anthropic_messages_raise_timeout():
raise litellm.Timeout(message="upstream timed out", model="glm", llm_provider="databricks")
def _anthropic_messages_raise_internal_server_error():
raise litellm.InternalServerError(message="upstream reset", llm_provider="databricks", model="glm")
@pytest.mark.asyncio
async def test_anthropic_messages_retry_raising_a_timeout_is_retried_like_a_pre_stream_timeout():
"""A 408 raised by a retry attempt before its stream opens is retried the way the pre-stream path retries
a 408, instead of ending the retries on the error-frame gate that only knows 429 and 5xx."""
router = _anthropic_messages_retry_router(num_retries=3)
provider = _AnthropicMessagesScriptedProvider(
_anthropic_messages_dropped_before_content,
_anthropic_messages_raise_timeout,
_anthropic_messages_retried_stream,
)
stream = await _anthropic_messages_stream_through_router(router, provider)
body = [chunk async for chunk in stream]
assert body == [_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("pong")]
assert [model in _ANTHROPIC_MESSAGES_RETRY_GROUP for model, _, _ in provider.calls] == [True, True, True]
@pytest.mark.asyncio
async def test_anthropic_messages_deployment_num_retries_also_governs_a_failure_before_the_stream_opens():
"""The deployment's num_retries litellm_param sets the budget for a failure raised before the stream opened
on this route, as it does for a mid-stream drop and for chat completions."""
router = _anthropic_messages_retry_router(num_retries=0, deployment_params={"num_retries": 2})
provider = _AnthropicMessagesScriptedProvider(
_anthropic_messages_raise_internal_server_error,
_anthropic_messages_raise_internal_server_error,
_anthropic_messages_retried_stream,
)
stream = await _anthropic_messages_stream_through_router(router, provider)
body = [chunk async for chunk in stream]
assert body == [_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("pong")]
assert len(provider.calls) == 3
@pytest.mark.asyncio
async def test_anthropic_messages_retry_raising_a_non_retriable_error_reaches_the_client_without_fallbacks():
router = _anthropic_messages_retry_router(num_retries=2)
provider = _AnthropicMessagesScriptedProvider(
_anthropic_messages_dropped_before_content,
_anthropic_messages_raise_authentication_error,
)
stream = await _anthropic_messages_stream_through_router(router, provider)
with pytest.raises(litellm.AuthenticationError):
[chunk async for chunk in stream]
assert [(attempted, budget) for _, attempted, budget in provider.calls] == [(0, 2), (1, 2)]
def _anthropic_messages_raise_context_window_error():
raise litellm.ContextWindowExceededError(message="prompt too long", llm_provider="anthropic", model="glm")
@pytest.mark.asyncio
async def test_anthropic_messages_retry_raising_a_context_window_error_takes_the_context_window_fallback():
"""The fallback chain sees the retry's own error type, so a context window overflow on the retried
deployment reaches context_window_fallbacks rather than the regular fallbacks."""
router = _anthropic_messages_retry_router(
num_retries=2, fallbacks=[{"glm": ["fb"]}], context_window_fallbacks=[{"glm": ["cw"]}]
)
provider = _AnthropicMessagesScriptedProvider(
_anthropic_messages_dropped_before_content,
_anthropic_messages_raise_context_window_error,
lambda: _AnthropicMessagesFakeByteStream(
[_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("from cw")]
),
)
stream = await _anthropic_messages_stream_through_router(router, provider)
body = [chunk async for chunk in stream]
assert body == [_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("from cw")]
assert [model for model, _, _ in provider.calls][-1] == "anthropic/cw-model"
def test_anthropic_messages_retry_budget_precedence_direct_call():
"""A retry policy naming the error class outranks the request's num_retries, which outranks the routed
deployment's, which outranks the router's; num_retries=0 on the request turns a policy off too."""
router = _anthropic_messages_retry_router(num_retries=3, deployment_params={"num_retries": 2})
deployment_id = router.get_model_list(model_name="glm")[0]["model_info"]["id"]
routed = {"model": "glm", "litellm_metadata": {"model_info": {"id": deployment_id}}}
drop = litellm.APIConnectionError(message="closed", llm_provider="databricks", model="glm")
reset = litellm.InternalServerError(message="reset", llm_provider="databricks", model="glm")
policy_router = _anthropic_messages_retry_router(
num_retries=3, retry_policy=RetryPolicy(InternalServerErrorRetries=4)
)
assert router._anthropic_messages_retry_budget(drop, {"model": "glm"}) == (3, False)
assert router._anthropic_messages_retry_budget(drop, routed) == (2, False)
assert router._anthropic_messages_retry_budget(drop, {**routed, "num_retries": 1}) == (1, False)
assert policy_router._anthropic_messages_retry_budget(reset, {"model": "glm", "num_retries": 1}) == (4, True)
assert policy_router._anthropic_messages_retry_budget(drop, {"model": "glm", "num_retries": 1}) == (1, False)
assert policy_router._anthropic_messages_retry_budget(reset, {"model": "glm", "num_retries": 0}) == (0, False)
committed = {**routed, "litellm_metadata": {**routed["litellm_metadata"], "attempted_retries": 1, "max_retries": 5}}
assert router._anthropic_messages_retry_budget(drop, committed) == (5, False)
assert policy_router._anthropic_messages_retry_budget(reset, committed) == (5, True)
def test_anthropic_messages_stream_can_retry_direct_call():
router = _anthropic_messages_retry_router(num_retries=1)
policy_router = _anthropic_messages_retry_router(
num_retries=0, retry_policy=RetryPolicy(InternalServerErrorRetries=1)
)
assert router._anthropic_messages_stream_can_retry({"model": "glm"}) is True
spent = {"model": "glm", "litellm_metadata": {"attempted_retries": 1}}
assert router._anthropic_messages_stream_can_retry(spent) is False
assert router._anthropic_messages_stream_can_retry({"model": "glm", "num_retries": 0}) is False
assert policy_router._anthropic_messages_stream_can_retry({"model": "glm"}) is True
assert policy_router._anthropic_messages_stream_can_retry(spent) is False
assert policy_router._anthropic_messages_stream_can_retry({"model": "glm", "num_retries": 0}) is False
assert policy_router._anthropic_messages_resolved_retry_policy({"model": "glm"}) is not None
assert policy_router._anthropic_messages_resolved_retry_policy({"model": "glm", "num_retries": 0}) is None
def test_retry_policy_ceiling_is_the_largest_budget_any_error_class_is_granted():
from litellm.router import _retry_policy_ceiling
assert _retry_policy_ceiling(RetryPolicy(InternalServerErrorRetries=1, RateLimitErrorRetries=3)) == 3
assert _retry_policy_ceiling(RetryPolicy()) == 0
@pytest.mark.asyncio
async def test_anthropic_messages_last_attempt_under_a_retry_policy_forwards_lifecycle_frames_live():
"""A retry policy bounds the hold the way a plain budget does: once the attempts reach the most retries the
policy grants, the stream is the last one, so its frames reach the client as they arrive and a drop after
them is the provider's error in-band rather than an error raised before any byte."""
router = _anthropic_messages_retry_router(num_retries=0, retry_policy=RetryPolicy(DefaultRetries=1))
drop = _anthropic_messages_transport_drop()
provider = _AnthropicMessagesScriptedProvider(
_anthropic_messages_dropped_before_content,
lambda: _AnthropicMessagesRaisingByteStream([_anthropic_messages_message_start_chunk()], drop),
)
stream = await _anthropic_messages_stream_through_router(router, provider)
received: list = []
with pytest.raises(litellm.APIConnectionError) as raised:
await _anthropic_messages_drain_into(stream, received)
assert raised.value is drop.original_exception
assert received == [_anthropic_messages_message_start_chunk()]
assert [(attempted, budget) for _, attempted, budget in provider.calls] == [(0, 0), (1, 1)]
@pytest.mark.asyncio
async def test_anthropic_messages_request_num_retries_zero_opts_out_of_the_mid_stream_retry():
router = _anthropic_messages_retry_router(num_retries=2)
drop = _anthropic_messages_transport_drop()
provider = _AnthropicMessagesScriptedProvider(
lambda: _AnthropicMessagesRaisingByteStream([_anthropic_messages_message_start_chunk()], drop)
)
stream = await _anthropic_messages_stream_through_router(router, provider, num_retries=0)
with pytest.raises(litellm.APIConnectionError) as raised:
[chunk async for chunk in stream]
assert raised.value is drop.original_exception
assert len(provider.calls) == 1
@pytest.mark.asyncio
@pytest.mark.parametrize("configured", [1, "1"], ids=["int", "config-string"])
async def test_anthropic_messages_deployment_num_retries_sets_the_mid_stream_retry_budget(configured):
"""A deployment's own num_retries litellm_param outranks the router's, as it does for a failure
raised before the stream opened."""
router = _anthropic_messages_retry_router(num_retries=0, deployment_params={"num_retries": configured})
provider = _AnthropicMessagesScriptedProvider(
_anthropic_messages_dropped_before_content, _anthropic_messages_retried_stream
)
stream = await _anthropic_messages_stream_through_router(router, provider)
body = [chunk async for chunk in stream]
assert body == [_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("pong")]
assert [(attempted, budget) for _, attempted, budget in provider.calls] == [(0, 0), (1, 1)]
@pytest.mark.asyncio
async def test_anthropic_messages_retry_policy_sets_the_mid_stream_retry_budget_per_error_class():
router = _anthropic_messages_retry_router(num_retries=0, retry_policy=RetryPolicy(InternalServerErrorRetries=1))
provider = _AnthropicMessagesScriptedProvider(
lambda: _AnthropicMessagesRaisingByteStream(
[_anthropic_messages_message_start_chunk()],
_anthropic_messages_transport_drop(
litellm.InternalServerError(message="upstream reset", llm_provider="databricks", model="glm")
),
),
_anthropic_messages_retried_stream,
)
stream = await _anthropic_messages_stream_through_router(router, provider)
body = [chunk async for chunk in stream]
assert body == [_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("pong")]
assert [(attempted, budget) for _, attempted, budget in provider.calls] == [(0, 0), (1, 1)]
def _anthropic_messages_error_frame(error_type: str) -> bytes:
return f"event: error\ndata: {json.dumps({'type': 'error', 'error': {'type': error_type, 'message': error_type}})}\n\n".encode()
_ANTHROPIC_MESSAGES_ERROR_FRAME_POLICIES: Final = (
pytest.param("api_error", RetryPolicy(InternalServerErrorRetries=1), id="api_error-500-internal-server"),
pytest.param("overloaded_error", RetryPolicy(InternalServerErrorRetries=1), id="overloaded-internal-server"),
pytest.param("rate_limit_error", RetryPolicy(RateLimitErrorRetries=1), id="rate-limit-429"),
pytest.param("timeout_error", RetryPolicy(TimeoutErrorRetries=1), id="timeout-504"),
)
@pytest.mark.asyncio
@pytest.mark.parametrize("error_type,policy", _ANTHROPIC_MESSAGES_ERROR_FRAME_POLICIES)
async def test_anthropic_messages_error_frame_is_retried_under_the_class_the_pre_stream_mapping_gives_it(
error_type, policy
):
"""An `event: error` frame before content carried a generic error, so a policy naming only error classes
granted it no retry while the hold still counted the policy: one attempt, then an HTTP error with no
bytes out. The frame now takes the class the pre-stream mapping raises for an answer carrying its body, so
an overloaded frame counts as the InternalServerError a 529 answer is, not a ServiceUnavailableError."""
router = _anthropic_messages_retry_router(num_retries=0, retry_policy=policy)
provider = _AnthropicMessagesScriptedProvider(
lambda: _AnthropicMessagesFakeByteStream(
[_anthropic_messages_message_start_chunk(), _anthropic_messages_error_frame(error_type)]
),
_anthropic_messages_retried_stream,
)
stream = await _anthropic_messages_stream_through_router(router, provider)
body = [chunk async for chunk in stream]
assert body == [_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("pong")]
assert [(attempted, budget) for _, attempted, budget in provider.calls] == [(0, 0), (1, 1)]
@pytest.mark.asyncio
async def test_anthropic_messages_error_frame_of_a_class_granted_no_retry_reaches_the_client_as_sent():
"""A rate limit frame is the RateLimitError a 429 answer is, so a policy granting only InternalServerError
retries leaves it unretried. With no fallback to take over either, the frame reaches the client as the
provider sent it, behind the lifecycle frames held back for a retry that never opened, the way the last
exhausted attempt's frames do; raising it instead turned a provider error frame into an HTTP error only
on the first attempt."""
router = _anthropic_messages_retry_router(num_retries=0, retry_policy=RetryPolicy(InternalServerErrorRetries=1))
frame = _anthropic_messages_error_frame("rate_limit_error")
provider = _AnthropicMessagesScriptedProvider(
lambda: _AnthropicMessagesFakeByteStream([_anthropic_messages_message_start_chunk(), frame])
)
stream = await _anthropic_messages_stream_through_router(router, provider)
body = [chunk async for chunk in stream]
assert body == [_anthropic_messages_message_start_chunk(), frame]
assert len(provider.calls) == 1
@pytest.mark.asyncio
async def test_anthropic_messages_error_frame_of_a_class_granted_no_retry_still_reaches_a_configured_fallback():
"""The same unretried rate limit frame goes to the fallback group when one is configured, since a
fallback can still take over before any byte reached the client."""
router = _anthropic_messages_retry_router(
num_retries=0, fallbacks=[{"glm": ["fb"]}], retry_policy=RetryPolicy(InternalServerErrorRetries=1)
)
provider = _AnthropicMessagesScriptedProvider(
lambda: _AnthropicMessagesFakeByteStream(
[_anthropic_messages_message_start_chunk(), _anthropic_messages_error_frame("rate_limit_error")]
),
lambda: _AnthropicMessagesFakeByteStream(
[_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("from fb")]
),
)
stream = await _anthropic_messages_stream_through_router(router, provider)
body = [chunk async for chunk in stream]
assert body == [_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("from fb")]
assert [model in _ANTHROPIC_MESSAGES_RETRY_GROUP for model, _, _ in provider.calls] == [True, False]
assert provider.calls[-1][0] == "anthropic/fb-model"
def test_anthropic_messages_recoverable_frame_error_direct_call():
"""Which `event: error` frames are intercepted for a retry or a fallback, and which reach the client as sent."""
policy_router = _anthropic_messages_retry_router(
num_retries=0, retry_policy=RetryPolicy(InternalServerErrorRetries=1)
)
fallback_router = _anthropic_messages_retry_router(num_retries=0, fallbacks=[{"glm": ["fb"]}])
api_error = ("api_error", "reset", 500)
rate_limit = ("rate_limit_error", "slow down", 429)
kwargs = {"model": "glm"}
recovered = policy_router._anthropic_messages_recoverable_frame_error(api_error, b"", False, "glm", kwargs)
assert isinstance(recovered, litellm.InternalServerError)
assert policy_router._anthropic_messages_recoverable_frame_error(rate_limit, b"", False, "glm", kwargs) is None
assert policy_router._anthropic_messages_recoverable_frame_error(api_error, b"", True, "glm", kwargs) is None
assert policy_router._anthropic_messages_recoverable_frame_error(None, b"", False, "glm", kwargs) is None
assert (
policy_router._anthropic_messages_recoverable_frame_error(
("invalid_request_error", "bad", 400), b"", False, "glm", kwargs
)
is None
)
spent = {"model": "glm", "litellm_metadata": {"attempted_retries": 1, "max_retries": 1}}
assert policy_router._anthropic_messages_recoverable_frame_error(api_error, b"", False, "glm", spent) is None
assert isinstance(
fallback_router._anthropic_messages_recoverable_frame_error(rate_limit, b"", False, "glm", kwargs),
litellm.RateLimitError,
)
_ANTHROPIC_MESSAGES_MALFORMED_POLICIES: Final = (
pytest.param({"glm": {"RateLimitErrorRetries": "many"}}, id="string-budget"),
pytest.param({"glm": 5}, id="group-policy-is-an-int"),
pytest.param(5, id="policy-map-is-an-int"),
pytest.param({"glm": [1]}, id="group-policy-is-a-list"),
)
@pytest.mark.asyncio
@pytest.mark.parametrize("policy", _ANTHROPIC_MESSAGES_MALFORMED_POLICIES)
@pytest.mark.parametrize("where", ["request", "router"])
async def test_anthropic_messages_malformed_retry_policy_leaves_a_healthy_stream_alone(policy, where):
"""The hold decision resolves the group's retry policy before the first byte, so a policy that does not
parse used to fail every stream of that group with a 500 before any attempt. It now governs nothing."""
router = _anthropic_messages_retry_router(num_retries=0)
request_kwargs = {"model_group_retry_policy": policy} if where == "request" else {}
if where == "router":
router.model_group_retry_policy = policy
provider = _AnthropicMessagesScriptedProvider(_anthropic_messages_retried_stream)
stream = await _anthropic_messages_stream_through_router(router, provider, **request_kwargs)
body = [chunk async for chunk in stream]
assert body == [_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("pong")]
assert len(provider.calls) == 1
@pytest.mark.asyncio
@pytest.mark.parametrize("policy", _ANTHROPIC_MESSAGES_MALFORMED_POLICIES)
async def test_anthropic_messages_malformed_retry_policy_falls_back_to_the_plain_budget(policy):
router = _anthropic_messages_retry_router(num_retries=1)
provider = _AnthropicMessagesScriptedProvider(
lambda: _AnthropicMessagesFakeByteStream(
[_anthropic_messages_message_start_chunk(), _anthropic_messages_error_frame("overloaded_error")]
),
_anthropic_messages_retried_stream,
)
stream = await _anthropic_messages_stream_through_router(router, provider, model_group_retry_policy=policy)
body = [chunk async for chunk in stream]
assert body == [_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("pong")]
assert [(attempted, budget) for _, attempted, budget in provider.calls] == [(0, 1), (1, 1)]
class _AnthropicMessagesAlternatingDeployments(CustomRoutingStrategyBase):
"""Routes each attempt to the group's next deployment in turn, so which sibling a retry lands on is known."""
def __init__(self, router: Router, model_group: str) -> None:
self._deployments = itertools.cycle(router.get_model_list(model_name=model_group) or ())
async def async_get_available_deployment(
self, model, messages=None, input=None, specific_deployment=False, request_kwargs=None
):
return next(self._deployments)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"first_num_retries,sibling_num_retries,expected_counters",
[
pytest.param(3, 1, [(0, 0), (1, 3), (2, 3), (3, 3)], id="sibling-grants-fewer"),
pytest.param(1, 3, [(0, 0), (1, 1)], id="sibling-grants-more"),
],
)
async def test_anthropic_messages_retries_keep_the_budget_the_first_drop_committed_to_across_deployments(
first_num_retries, sibling_num_retries, expected_counters
):
"""A retry's stream recomputed its budget from the sibling deployment it landed on, so a group whose
deployments grant different num_retries stopped early or overshot the budget the first drop stamped
into the retry headers; later attempts now keep that budget, as the pre-stream retry loop does."""
router = Router(
model_list=[
{
"model_name": "glm",
"litellm_params": {"model": "anthropic/glm-a", "api_key": "sk-test", "num_retries": first_num_retries},
},
{
"model_name": "glm",
"litellm_params": {"model": "anthropic/glm-b", "api_key": "sk-test", "num_retries": sibling_num_retries},
},
],
num_retries=0,
fallbacks=None,
)
router.set_custom_routing_strategy(_AnthropicMessagesAlternatingDeployments(router, "glm"))
provider = _AnthropicMessagesScriptedProvider(*[_anthropic_messages_dropped_before_content] * len(expected_counters))
stream = await _anthropic_messages_stream_through_router(router, provider)
with pytest.raises(litellm.APIConnectionError):
[chunk async for chunk in stream]
assert [model for model, _, _ in provider.calls] == (["anthropic/glm-a", "anthropic/glm-b"] * 2)[: len(expected_counters)]
assert [(attempted, budget) for _, attempted, budget in provider.calls] == expected_counters
@pytest.mark.asyncio
async def test_anthropic_messages_lifecycle_frames_wait_for_content_while_a_retry_remains():
"""A retry can only restart cleanly while nothing reached the client, so with retries left a
fallback-less group holds message_start back until the first content frame, as a fallback does."""
router = _anthropic_messages_retry_router(num_retries=2)
content_released = asyncio.Event()
async def held_stream():
yield _anthropic_messages_message_start_chunk()
await content_released.wait()
yield _anthropic_messages_content_chunk("hi")
provider = _AnthropicMessagesScriptedProvider(held_stream)
stream = await _anthropic_messages_stream_through_router(router, provider)
pending = asyncio.ensure_future(stream.__anext__())
await asyncio.sleep(0.2)
assert not pending.done()
content_released.set()
assert await asyncio.wait_for(pending, timeout=1) == _anthropic_messages_message_start_chunk()
assert [chunk async for chunk in stream] == [_anthropic_messages_content_chunk("hi")]
@pytest.mark.asyncio
async def test_anthropic_messages_attempt_strips_the_controls_carrier_and_wraps_every_hop_stream():
"""Each attempt of the chain, not only the primary's, comes back wrapped for mid-stream
@ -15299,6 +16013,7 @@ def _mid_stream_opt_out_router() -> Router:
{"model_name": "fallback", "litellm_params": {"model": "openai/gpt-5.4-mini", "api_key": "k2"}},
],
fallbacks=[{"primary": ["fallback"]}],
num_retries=0,
)