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, )