From 68d9b8bbb845896625a752c36d26fa41d674f874 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 5 Oct 2026 22:53:39 +0000 Subject: [PATCH] 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> --- litellm/llms/anthropic/common_utils.py | 18 + litellm/router.py | 397 ++++++++-- .../router_utils/fallback_event_handlers.py | 72 ++ litellm/router_utils/get_retry_from_policy.py | 4 +- .../router_code_coverage.py | 12 + tests/integration/_support/anthropic_sse.py | 197 +++++ tests/integration/_support/openai_wire.py | 125 +++ ...test_chat_bridge_pre_content_retry_wire.py | 179 +++++ ...hropic_messages_pre_content_retry_chaos.py | 411 ++++++++++ ...thropic_messages_pre_content_retry_wire.py | 473 ++++++++++++ ...loyment_num_retries_generic_routes_wire.py | 255 +++++++ ...stream_retry_budget_sources_owned_proxy.py | 263 +++++++ .../anthropic/test_anthropic_common_utils.py | 40 + .../test_fallback_event_handlers.py | 106 ++- tests/unit/test_router/test_router.py | 721 +++++++++++++++++- 15 files changed, 3212 insertions(+), 61 deletions(-) create mode 100644 tests/integration/_support/anthropic_sse.py create mode 100644 tests/integration/_support/openai_wire.py create mode 100644 tests/integration/messages_endpoint/chat_bridge/test_chat_bridge_pre_content_retry_wire.py create mode 100644 tests/integration/messages_endpoint/providers/anthropic/test_anthropic_messages_pre_content_retry_chaos.py create mode 100644 tests/integration/messages_endpoint/providers/anthropic/test_anthropic_messages_pre_content_retry_wire.py create mode 100644 tests/integration/routing/test_deployment_num_retries_generic_routes_wire.py create mode 100644 tests/integration/routing/test_messages_stream_retry_budget_sources_owned_proxy.py diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index 3e8d7481133..6da3807ab15 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -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}$") diff --git a/litellm/router.py b/litellm/router.py index afc842a05a0..21709bde1f8 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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 diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py index 74504d76736..432351e58f3 100644 --- a/litellm/router_utils/fallback_event_handlers.py +++ b/litellm/router_utils/fallback_event_handlers.py @@ -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" diff --git a/litellm/router_utils/get_retry_from_policy.py b/litellm/router_utils/get_retry_from_policy.py index 8771d072434..7c412dc759b 100644 --- a/litellm/router_utils/get_retry_from_policy.py +++ b/litellm/router_utils/get_retry_from_policy.py @@ -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 = ( diff --git a/tests/code_coverage_tests/router_code_coverage.py b/tests/code_coverage_tests/router_code_coverage.py index 574336791de..e989c40d095 100644 --- a/tests/code_coverage_tests/router_code_coverage.py +++ b/tests/code_coverage_tests/router_code_coverage.py @@ -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 ] diff --git a/tests/integration/_support/anthropic_sse.py b/tests/integration/_support/anthropic_sse.py new file mode 100644 index 00000000000..4bcfe460c3a --- /dev/null +++ b/tests/integration/_support/anthropic_sse.py @@ -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] diff --git a/tests/integration/_support/openai_wire.py b/tests/integration/_support/openai_wire.py new file mode 100644 index 00000000000..d16cb60c455 --- /dev/null +++ b/tests/integration/_support/openai_wire.py @@ -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()) diff --git a/tests/integration/messages_endpoint/chat_bridge/test_chat_bridge_pre_content_retry_wire.py b/tests/integration/messages_endpoint/chat_bridge/test_chat_bridge_pre_content_retry_wire.py new file mode 100644 index 00000000000..9767e3b0fde --- /dev/null +++ b/tests/integration/messages_endpoint/chat_bridge/test_chat_bridge_pre_content_retry_wire.py @@ -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,) diff --git a/tests/integration/messages_endpoint/providers/anthropic/test_anthropic_messages_pre_content_retry_chaos.py b/tests/integration/messages_endpoint/providers/anthropic/test_anthropic_messages_pre_content_retry_chaos.py new file mode 100644 index 00000000000..b6e848bad75 --- /dev/null +++ b/tests/integration/messages_endpoint/providers/anthropic/test_anthropic_messages_pre_content_retry_chaos.py @@ -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) diff --git a/tests/integration/messages_endpoint/providers/anthropic/test_anthropic_messages_pre_content_retry_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_anthropic_messages_pre_content_retry_wire.py new file mode 100644 index 00000000000..cb7f4c2b9ce --- /dev/null +++ b/tests/integration/messages_endpoint/providers/anthropic/test_anthropic_messages_pre_content_retry_wire.py @@ -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() == () diff --git a/tests/integration/routing/test_deployment_num_retries_generic_routes_wire.py b/tests/integration/routing/test_deployment_num_retries_generic_routes_wire.py new file mode 100644 index 00000000000..9b14d05c1e2 --- /dev/null +++ b/tests/integration/routing/test_deployment_num_retries_generic_routes_wire.py @@ -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 diff --git a/tests/integration/routing/test_messages_stream_retry_budget_sources_owned_proxy.py b/tests/integration/routing/test_messages_stream_retry_budget_sources_owned_proxy.py new file mode 100644 index 00000000000..93e4f97ae1e --- /dev/null +++ b/tests/integration/routing/test_messages_stream_retry_budget_sources_owned_proxy.py @@ -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 diff --git a/tests/unit/llms/anthropic/test_anthropic_common_utils.py b/tests/unit/llms/anthropic/test_anthropic_common_utils.py index bd4d3085533..8b9729fa9f0 100644 --- a/tests/unit/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/unit/llms/anthropic/test_anthropic_common_utils.py @@ -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 diff --git a/tests/unit/router_utils/test_fallback_event_handlers.py b/tests/unit/router_utils/test_fallback_event_handlers.py index 86e283a2f2b..ec05a2c2840 100644 --- a/tests/unit/router_utils/test_fallback_event_handlers.py +++ b/tests/unit/router_utils/test_fallback_event_handlers.py @@ -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 diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index d0115593e46..053dd730dab 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -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, )