From 444963a91ff16f9b90b98b69b6e4f694d3864eb8 Mon Sep 17 00:00:00 2001 From: leilei3167 Date: Sat, 12 Sep 2026 11:06:00 +0800 Subject: [PATCH] chore: rebase onto litellm_internal_staging and resolve utils.py conflicts --- .../litellm_core_utils/streaming_handler.py | 2657 +-- litellm/utils.py | 9885 +----------- .../test_streaming_handler.py | 5015 +----- tests/test_litellm/test_router.py | 13428 +--------------- 4 files changed, 4 insertions(+), 30981 deletions(-) diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 7da843b4d5e..eaa7e580556 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -1,2656 +1 @@ -import asyncio -import collections.abc -import datetime -import json -import logging -import threading -import time -import traceback -from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping, Sequence -from dataclasses import dataclass -from types import MappingProxyType -from typing import Any, Final, NoReturn, Protocol, TypeVar, cast - -import anyio -import httpx -from pydantic import BaseModel, ValidationError -from typing_extensions import NotRequired, TypedDict - -import litellm -from litellm import verbose_logger -from litellm._uuid import uuid -from litellm.litellm_core_utils.model_response_utils import ( - is_model_response_stream_empty, -) -from litellm.litellm_core_utils.redact_messages import LiteLLMLoggingObject -from litellm.litellm_core_utils.thread_pool_executor import executor -from litellm.types.llms.openai import OpenAIChatCompletionChunk -from litellm.types.router import GenericLiteLLMParams -from litellm.types.utils import ( - CacheCreationTokenDetails, - CompletionTokensDetailsWrapper, - Delta, - LlmProviders, - ModelResponse, - ModelResponseStream, - PromptTokensDetailsWrapper, - StreamingChoices, - Usage, -) -from litellm.types.utils import GenericStreamingChunk as GChunk - -from ..exceptions import OpenAIError -from .core_helpers import map_finish_reason, process_response_headers -from .exception_mapping_utils import exception_type -from .llm_response_utils.get_api_base import get_api_base -from .rules import Rules - -# Constants for special delta attribute names -AUDIO_ATTRIBUTE: Final = "audio" -IMAGE_ATTRIBUTE: Final = "images" -TOOL_CALLS_ATTRIBUTE: Final = "tool_calls" -FUNCTION_CALL_ATTRIBUTE: Final = "function_call" - -_SYNC_ITER_EXHAUSTED: Final = object() - -_GCHUNK_FIELDS: Final[frozenset] = frozenset(GChunk.__annotations__) -_USAGE_COST_HEADER_PROVIDERS: Final[frozenset[str]] = frozenset({LlmProviders.OPENROUTER.value}) - - -def _next_sync_or_exhausted(it: Any) -> object: - """ - Call next(it) from a thread and return _SYNC_ITER_EXHAUSTED on StopIteration. - - asyncio.to_thread re-raises thread exceptions inside a coroutine, where PEP 479 - converts StopIteration to RuntimeError before any except clause can catch it. - Returning a sentinel instead keeps StopIteration out of the coroutine boundary. - """ - try: - return next(it) - except StopIteration: - return _SYNC_ITER_EXHAUSTED - - -def is_async_iterable(obj: object) -> bool: - """ - Check if an object is an async iterable (can be used with 'async for'). - - Args: - obj: Any Python object to check - - Returns: - bool: True if the object is async iterable, False otherwise - """ - return isinstance(obj, collections.abc.AsyncIterable) - - -def print_verbose(print_statement: object): - try: - if litellm.set_verbose: - print(print_statement) # noqa: T201 - except Exception: - pass - - -@dataclass(frozen=True, slots=True) -class _ProviderChunkParsed: - response_obj: dict[str, object] - - -@dataclass(frozen=True, slots=True) -class _ProviderChunkEarlyReturn: - value: "ModelResponseStream | None" - - -_ProviderChunkResult = _ProviderChunkParsed | _ProviderChunkEarlyReturn - - -class _PredibaseStreamData(TypedDict): - token: NotRequired[Mapping[str, str]] - details: Mapping[str, str] - generated_text: str | None - error: str | None - - -class _Ai21StreamData(TypedDict): - completions: Sequence[Mapping[str, Mapping[str, str]]] - - -class _MaritalkStreamData(TypedDict): - answer: str - - -class _NlpCloudStreamData(TypedDict): - generated_text: str - - -class _AlephAlphaStreamData(TypedDict): - completions: Sequence[Mapping[str, str]] - - -class _AzureStreamChoice(TypedDict): - delta: Mapping[str, str] | None - finish_reason: str | None - - -class _AzureStreamData(TypedDict): - choices: Sequence[_AzureStreamChoice] - - -class _BasetenModelOutput(TypedDict): - data: NotRequired[Sequence[str]] - - -class _BasetenStreamData(TypedDict): - token: NotRequired[Mapping[str, str]] - model_output: NotRequired["_BasetenModelOutput | str"] - completion: NotRequired[object] - - -class _DeltaDumpDict(TypedDict): - role: NotRequired[str | None] - tool_calls: NotRequired[Sequence[Mapping[str, object]]] - - -class _TextCompletionChoiceLike(Protocol): - text: str - finish_reason: str | None - - -class _VertexFunctionCallLike(Protocol): - name: str - args: Mapping[str, Iterable[object]] - - -class _VertexPartLike(Protocol): - function_call: _VertexFunctionCallLike - - -class _VertexContentLike(Protocol): - parts: Sequence[_VertexPartLike] - - -class _VertexFinishReasonLike(Protocol): - name: str - - -class _VertexCandidateLike(Protocol): - content: _VertexContentLike - finish_reason: _VertexFinishReasonLike - - -class _VertexChunkLike(Protocol): - text: str - candidates: Sequence[_VertexCandidateLike] - - -class _ParsedChunkHiddenParams(BaseModel): - provider_specific_fields: Mapping[str, object] | None = None - - -def _provider_response_model(chunk: object) -> str | None: - model: Final[object] = chunk.get("model") if isinstance(chunk, Mapping) else getattr(chunk, "model", None) - return model if isinstance(model, str) and model else None - - -def _parsed_provider_hidden_params(hidden: object) -> _ParsedChunkHiddenParams | None: - if not isinstance(hidden, dict): - return None - try: - return _ParsedChunkHiddenParams.model_validate(hidden) - except ValidationError: - return None - - -def _provider_hidden_params( - chunk: object, - provider_response_model: str | None, -) -> Mapping[str, object] | None: - hidden: Final[object] = getattr(chunk, "_hidden_params", None) - parsed: Final = _parsed_provider_hidden_params(hidden) - provider_specific_fields: Final[object | None] = ( - dict(parsed.provider_specific_fields) # mutable-ok: stream assembly merges provider metadata into this dict - if parsed is not None and parsed.provider_specific_fields - else None - ) - params: Final[Mapping[str, object]] = MappingProxyType( - { - key: value - for key, value in ( - ("provider_response_model", provider_response_model), - ("provider_specific_fields", provider_specific_fields), - ) - if value is not None - } - ) - return params or None - - -class CustomStreamWrapper: - def __init__( - self, - completion_stream, - model, - logging_obj: LiteLLMLoggingObject, - custom_llm_provider: str | None = None, - stream_options=None, - make_call: Callable | None = None, - _response_headers: dict | httpx.Headers | None = None, - ): - self.model = model - self.make_call = make_call - self.custom_llm_provider = custom_llm_provider - self.logging_obj: LiteLLMLoggingObject = logging_obj - self.completion_stream = completion_stream - self.sent_first_chunk = False - self.sent_last_chunk = False - self._stream_created_time: float = time.time() - - litellm_params: Final[GenericLiteLLMParams] = GenericLiteLLMParams.model_validate( - dict(**self.logging_obj.model_call_details.get("litellm_params", {})) - ) - self.merge_reasoning_content_in_choices: bool = litellm_params.merge_reasoning_content_in_choices or False - self.sent_first_thinking_block = False - self.sent_last_thinking_block = False - self.thinking_content = "" - - self.system_fingerprint: str | None = None - self._provider_response_model: str | None = None - self.received_finish_reason: str | None = None - self.intermittent_finish_reason: str | None = None # finish reasons that show up mid-stream - self.special_tokens = [ - "<|assistant|>", - "<|system|>", - "<|user|>", - "", - "", - "<|im_end|>", - "<|im_start|>", - ] - self.holding_chunk = "" - self.complete_response = "" - self.response_uptil_now = "" - _model_info: Final[dict] = litellm_params.model_info or {} - - _api_base: Final = get_api_base( - model=model or "", - optional_params=self.logging_obj.model_call_details.get("litellm_params", {}), - ) - - self._hidden_params = { - "model_id": (_model_info.get("id", None)), - "api_base": _api_base, - } # returned as x-litellm-model-id response header in proxy - - self._hidden_params["additional_headers"] = process_response_headers( - _response_headers or {} - ) # GUARANTEE OPENAI HEADERS IN RESPONSE - - self._response_headers = _response_headers - self.response_id: str | None = None - self.logging_loop = None - self.rules = Rules() - self.stream_options = stream_options or getattr(logging_obj, "stream_options", None) - self.messages = getattr(logging_obj, "messages", None) - self.sent_stream_usage = False - self.send_stream_usage = True if self.check_send_stream_usage(self.stream_options) else False - self.tool_call = False - self.chunks: list = [] # keep track of the returned chunks - used for calculating the input/output tokens for stream options - self._repeated_messages_count = 1 - self.is_function_call = self.check_is_function_call(logging_obj=logging_obj) - self.created: int | None = None - self._last_returned_hidden_params: dict | None = None - - _cached_logging_provider: Final = self.logging_obj.model_call_details.get("custom_llm_provider", None) - self._cached_logging_llm_provider: str | None = _cached_logging_provider - _effective_model = model or "" - if custom_llm_provider == "openai" and custom_llm_provider != _cached_logging_provider: - _effective_model = f"{_cached_logging_provider}/{_effective_model}" - self._cached_model_name: str = _effective_model - - # Snapshot assumes self._hidden_params is populated from litellm_params - # at init and never mutated during the stream. If that ever changes, - # this cache must be removed. - self._base_hidden_params: dict[str, object] = { - **self._hidden_params, - "response_cost": None, - } - - self._post_streaming_hooks: list | None = None - - def _check_max_streaming_duration(self) -> None: - """Raise litellm.Timeout if the stream has exceeded LITELLM_MAX_STREAMING_DURATION_SECONDS.""" - from litellm.constants import LITELLM_MAX_STREAMING_DURATION_SECONDS - - if LITELLM_MAX_STREAMING_DURATION_SECONDS is None: - return - elapsed: Final = time.time() - self._stream_created_time - if elapsed > LITELLM_MAX_STREAMING_DURATION_SECONDS: - raise litellm.Timeout( - message=f"Stream exceeded max streaming duration of {LITELLM_MAX_STREAMING_DURATION_SECONDS}s (elapsed {elapsed:.1f}s)", - model=self.model or "", - llm_provider=self.custom_llm_provider or "", - ) - - def __iter__(self) -> Iterator["ModelResponseStream"]: - return self - - def __aiter__(self) -> AsyncIterator["ModelResponseStream"]: - return self - - def _restore_consumer_correlation_context(self, *, guarded: bool = False) -> None: - """Restore trace_id/session_id in the *consuming* thread/task/context. - - wrapper_async() deliberately skips restoring correlation context when - it returns a stream, so log lines emitted while the caller iterates it - still carry this call's ids (see request_correlation_in_logs). - wrapper() (the sync path) never stamps anything in the first place - - see Logging.__init__'s supports_correlation_logging - so this method - is an inert no-op for sync-created streams, harmless to call anyway - since the class is shared between __next__ and __anext__. - But the terminal success/failure handlers this stream dispatches to - finish the job run on a *different* Task/thread (asyncio.create_task, - threading.Thread, or the shared executor) - restoring there fixes up - that detached context, not the one actually running the caller's - `for`/`async for` loop. Call this at every point control genuinely - returns to that consuming context: natural exhaustion (StopIteration/ - StopAsyncIteration), a raised failure, or explicit aclose(). Never let - this raise - it must not break the caller's actual stream handling. - - guarded=True (only __del__ uses this) skips the restore unless the - contextvars still hold the ids this stream's own call set, so a - delayed finalizer never overwrites a different, still-active call - that has since taken over the same Task/thread's context. - """ - try: - logging_obj: Final[object | None] = getattr(self, "logging_obj", None) - if logging_obj is None: - return - method_name: Final = ( - "_restore_correlation_context_if_unclaimed" if guarded else "_restore_correlation_context" - ) - restore: Final[Callable[[], object] | None] = getattr(logging_obj, method_name, None) - if restore is not None: - restore() - except Exception as restore_error: # noqa: BLE001 # best-effort cleanup; must not raise into the caller - verbose_logger.debug("could not restore correlation context: %s", restore_error) - - def __del__(self) -> None: - """Best-effort correlation-context cleanup for an abandoned async stream. - - Only meaningfully applies to streams created by wrapper_async(): it - leaves contextvars "open" across the caller's iteration, so if the - caller never fully consumes the stream - stops early, drops the - reference, cancels it - none of the exit points - _restore_consumer_correlation_context() is called from ever run. For a - sync stream (wrapper()), this is a no-op in practice: wrapper() never - stamps trace_id/session_id for sync calls in the first place (see - Logging.__init__'s supports_correlation_logging), so there is nothing - for this to clean up. - - This is a best-effort fallback, not a guarantee: __del__ timing is - unpredictable (delayed by cyclic GC, not guaranteed at interpreter - shutdown, and may run on a different thread), so this can only reduce - how long the leak persists, not eliminate it. That's an acceptable - trade specifically because its blast radius is bounded to the one - asyncio Task this stream's own call ran in - each async call has its - own copy of the contextvars, and Tasks (unlike a thread pool's worker - threads) are never recycled across requests, so a delayed or missed - cleanup here can never misattribute a *different* request's logs. - guarded=True additionally ensures it never clobbers a different, - still-active call's context within that same Task if this fires late. - """ - self._restore_consumer_correlation_context(guarded=True) - - async def aclose(self): - # Restore the consumer's outer context only after the underlying - # provider stream's own close (and its diagnostic logging below, if - # closing fails) completes - not before - so those log lines still - # carry this closing stream's own trace_id/session_id. - if self.completion_stream is not None: - stream_to_close: Final = self.completion_stream - self.completion_stream = None - # Shield from anyio cancellation so cleanup awaits can complete. - # Without this, CancelledError is thrown into every await during - # task group cancellation, preventing HTTP connection release. - with anyio.CancelScope(shield=True): - try: - if hasattr(stream_to_close, "aclose"): - await stream_to_close.aclose() - elif hasattr(stream_to_close, "close"): - result: Final = stream_to_close.close() - if result is not None: - await result - except BaseException as e: - verbose_logger.debug( - "CustomStreamWrapper.aclose: error closing completion_stream: %s", - e, - ) - self._restore_consumer_correlation_context() - - def check_send_stream_usage(self, stream_options: dict | None): - return stream_options is not None and stream_options.get("include_usage", False) is True - - def check_is_function_call(self, logging_obj) -> bool: - from litellm.litellm_core_utils.prompt_templates.common_utils import ( - is_function_call, - ) - - if hasattr(logging_obj, "optional_params") and isinstance(logging_obj.optional_params, dict): - if is_function_call(logging_obj.optional_params): - return True - - return False - - def process_chunk(self, chunk: str): - """ - NLP Cloud streaming returns the entire response, for each chunk. Process this, to only return the delta. - """ - try: - chunk = chunk.strip() - self.complete_response = self.complete_response.strip() - - chunk = chunk.removeprefix(self.complete_response) - - self.complete_response += chunk - return chunk - except Exception as e: - raise e - - def raise_on_model_repetition(self) -> None: - """ - Fixes - https://github.com/BerriAI/litellm/issues/5158 - - if the model enters a loop and starts repeating the same chunk again, break out of loop and raise an internalservererror - allows for retries. - - Raises - InternalServerError, if LLM enters infinite loop while streaming - """ - if len(self.chunks) < 2: - return - - # Providers like Vertex Gemini (Flash / Flash Lite with web search) emit - # metadata-only / usage-only chunks with no choices. These get stored in - # self.chunks but carry no comparable content, so skip repetition detection. - if not self.chunks[-1].choices or not self.chunks[-2].choices: - return - - last_content: Final = self.chunks[-1].choices[0].delta.content - - if ( - last_content is None or not isinstance(last_content, str) or len(last_content) <= 2 - ): # ignore empty content - https://github.com/BerriAI/litellm/issues/5158#issuecomment-2287156946 - self._repeated_messages_count = 1 - return - - second_to_last_content: Final = self.chunks[-2].choices[0].delta.content - - if last_content == second_to_last_content: - self._repeated_messages_count += 1 - else: - self._repeated_messages_count = 1 - - if self._repeated_messages_count >= litellm.REPEATED_STREAMING_CHUNK_LIMIT: - # All last n chunks are identical - raise litellm.InternalServerError( - message=f"The model is repeating the same chunk = {last_content}.", - model="", - llm_provider="", - ) - - def check_special_tokens(self, chunk: str, finish_reason: str | None): - """ - Output parse / special tokens for sagemaker + hf streaming. - """ - hold = False - if self.custom_llm_provider != "sagemaker": - return hold, chunk - - if finish_reason: - for token in self.special_tokens: - if token in chunk: - chunk = chunk.replace(token, "") - return hold, chunk - - if self.sent_first_chunk is True: - return hold, chunk - - curr_chunk = self.holding_chunk + chunk - curr_chunk = curr_chunk.strip() - - for token in self.special_tokens: - if len(curr_chunk) < len(token) and curr_chunk in token: - hold = True - self.holding_chunk = curr_chunk - elif len(curr_chunk) >= len(token): - if token in curr_chunk: - self.holding_chunk = curr_chunk.replace(token, "") - hold = True - else: - pass - - if hold is False: # reset - self.holding_chunk = "" - return hold, curr_chunk - - def handle_predibase_chunk(self, chunk): - try: - if not isinstance(chunk, str): - chunk = chunk.decode("utf-8") # DO NOT REMOVE this: This is required for HF inference API + Streaming - text = "" - is_finished = False - finish_reason = "" - print_verbose(f"chunk: {chunk}") - if chunk.startswith("data:"): - data_json: Final[_PredibaseStreamData] = json.loads(chunk[5:]) - print_verbose(f"data json: {data_json}") - if "token" in data_json and "text" in data_json["token"]: - text = data_json["token"]["text"] - if data_json.get("details", False) and data_json["details"].get("finish_reason", False): - is_finished = True - finish_reason = data_json["details"]["finish_reason"] - elif data_json.get("generated_text", False): # if full generated text exists, then stream is complete - text = "" # don't return the final bos token - is_finished = True - finish_reason = "stop" - elif data_json.get("error", False): - raise Exception(data_json.get("error")) - return { - "text": text, - "is_finished": is_finished, - "finish_reason": finish_reason, - } - elif "error" in chunk: - raise ValueError(chunk) - return { - "text": text, - "is_finished": is_finished, - "finish_reason": finish_reason, - } - except Exception as e: - raise e - - def handle_ai21_chunk(self, chunk): # fake streaming - chunk = chunk.decode("utf-8") - data_json: Final[_Ai21StreamData] = json.loads(chunk) - try: - text: Final = data_json["completions"][0]["data"]["text"] - is_finished: Final = True - finish_reason: Final = "stop" - return { - "text": text, - "is_finished": is_finished, - "finish_reason": finish_reason, - } - except Exception: - raise ValueError(f"Unable to parse response. Original response: {chunk}") - - def handle_maritalk_chunk(self, chunk): # fake streaming - chunk = chunk.decode("utf-8") - data_json: Final[_MaritalkStreamData] = json.loads(chunk) - try: - text: Final = data_json["answer"] - is_finished: Final = True - finish_reason: Final = "stop" - return { - "text": text, - "is_finished": is_finished, - "finish_reason": finish_reason, - } - except Exception: - raise ValueError(f"Unable to parse response. Original response: {chunk}") - - def handle_nlp_cloud_chunk(self, chunk): - text = "" - is_finished = False - finish_reason = "" - try: - if self.model and "dolphin" in self.model: - chunk = self.process_chunk(chunk=chunk) - else: - data_json: Final[_NlpCloudStreamData] = json.loads(chunk) - chunk = data_json["generated_text"] - text = chunk - if "[DONE]" in text: - text = text.replace("[DONE]", "") - is_finished = True - finish_reason = "stop" - return { - "text": text, - "is_finished": is_finished, - "finish_reason": finish_reason, - } - except Exception: - raise ValueError(f"Unable to parse response. Original response: {chunk}") - - def handle_aleph_alpha_chunk(self, chunk): - chunk = chunk.decode("utf-8") - data_json: Final[_AlephAlphaStreamData] = json.loads(chunk) - try: - text: Final = data_json["completions"][0]["completion"] - is_finished: Final = True - finish_reason: Final = "stop" - return { - "text": text, - "is_finished": is_finished, - "finish_reason": finish_reason, - } - except Exception: - raise ValueError(f"Unable to parse response. Original response: {chunk}") - - def handle_azure_chunk(self, chunk): - is_finished = False - finish_reason = "" - text = "" - print_verbose(f"chunk: {chunk}") - if "data: [DONE]" in chunk: - text = "" - is_finished = True - finish_reason = "stop" - return { - "text": text, - "is_finished": is_finished, - "finish_reason": finish_reason, - } - elif chunk.startswith("data:"): - data_json: Final[_AzureStreamData] = json.loads(chunk[5:]) # chunk.startswith("data:"): - try: - if len(data_json["choices"]) > 0: - delta: Final = data_json["choices"][0]["delta"] - text = "" if delta is None else delta.get("content", "") - if data_json["choices"][0].get("finish_reason", None): - is_finished = True - finish_reason = data_json["choices"][0]["finish_reason"] - print_verbose(f"text: {text}; is_finished: {is_finished}; finish_reason: {finish_reason}") - return { - "text": text, - "is_finished": is_finished, - "finish_reason": finish_reason, - } - except Exception: - raise ValueError(f"Unable to parse response. Original response: {chunk}") - elif "error" in chunk: - raise ValueError(f"Unable to parse response. Original response: {chunk}") - else: - return { - "text": text, - "is_finished": is_finished, - "finish_reason": finish_reason, - } - - def handle_replicate_chunk(self, chunk): - try: - text = "" - is_finished = False - finish_reason = "" - if "output" in chunk: - text = chunk["output"] - if "status" in chunk: - if chunk["status"] == "succeeded": - is_finished = True - finish_reason = "stop" - elif chunk.get("error", None): - raise Exception(chunk["error"]) - return { - "text": text, - "is_finished": is_finished, - "finish_reason": finish_reason, - } - except Exception: - raise ValueError(f"Unable to parse response. Original response: {chunk}") - - def handle_openai_chat_completion_chunk(self, chunk): - try: - str_line: Final = chunk - text = "" - is_finished = False - finish_reason = None - logprobs = None - usage = None - if str_line and str_line.choices and len(str_line.choices) > 0: - if str_line.choices[0].delta is not None and str_line.choices[0].delta.content is not None: - text = str_line.choices[0].delta.content - else: # function/tool calling chunk - when content is None. in this case we just return the original chunk from openai - pass - if str_line.choices[0].finish_reason: - is_finished = True # check if str_line._hidden_params["is_finished"] is True - if hasattr(str_line, "_hidden_params") and str_line._hidden_params.get("is_finished") is not None: - is_finished = str_line._hidden_params.get("is_finished") - finish_reason = str_line.choices[0].finish_reason - - # checking for logprobs - if hasattr(str_line.choices[0], "logprobs") and str_line.choices[0].logprobs is not None: - logprobs = str_line.choices[0].logprobs - else: - logprobs = None - - usage = getattr(str_line, "usage", None) - - return { - "text": text, - "is_finished": is_finished, - "finish_reason": finish_reason, - "logprobs": logprobs, - "original_chunk": str_line, - "usage": usage, - } - except Exception as e: - raise e - - def handle_azure_text_completion_chunk(self, chunk): - try: - text = "" - is_finished = False - finish_reason = None - choices: Final[Sequence[_TextCompletionChoiceLike]] = getattr(chunk, "choices", []) - if len(choices) > 0: - text = choices[0].text - if choices[0].finish_reason is not None: - is_finished = True - finish_reason = choices[0].finish_reason - return { - "text": text, - "is_finished": is_finished, - "finish_reason": finish_reason, - } - - except Exception as e: - raise e - - def handle_openai_text_completion_chunk(self, chunk): - try: - text = "" - is_finished = False - finish_reason = None - usage = None - choices: Final[Sequence[_TextCompletionChoiceLike]] = getattr(chunk, "choices", []) - if len(choices) > 0: - text = choices[0].text - if choices[0].finish_reason is not None: - is_finished = True - finish_reason = choices[0].finish_reason - usage = getattr(chunk, "usage", None) - return { - "text": text, - "is_finished": is_finished, - "finish_reason": finish_reason, - "usage": usage, - } - - except Exception as e: - raise e - - def handle_baseten_chunk(self, chunk) -> str: - try: - chunk = chunk.decode("utf-8") - if len(chunk) > 0: - if chunk.startswith("data:"): - data_json: _BasetenStreamData = json.loads(chunk[5:]) - if "token" in data_json and "text" in data_json["token"]: - return data_json["token"]["text"] - else: - return "" - data_json = json.loads(chunk) - if "model_output" in data_json: - if ( - isinstance(data_json["model_output"], dict) - and "data" in data_json["model_output"] - and isinstance(data_json["model_output"]["data"], list) - ): - return data_json["model_output"]["data"][0] - elif isinstance(data_json["model_output"], str): - return data_json["model_output"] - elif "completion" in data_json and isinstance(data_json["completion"], str): - return data_json["completion"] - else: - raise ValueError(f"Unable to parse response. Original response: {chunk}") - else: - return "" - else: - return "" - except Exception as e: - verbose_logger.exception("litellm.CustomStreamWrapper.handle_baseten_chunk(): Exception occured - %s", e) - return "" - - def handle_triton_stream(self, chunk): - try: - if isinstance(chunk, dict): - parsed_response = chunk - elif isinstance(chunk, (str, bytes)): - if isinstance(chunk, bytes): - chunk = chunk.decode("utf-8") - if "text_output" in chunk: - response = CustomStreamWrapper._strip_sse_data_from_chunk(chunk) or "" - response = response.strip() - parsed_response = json.loads(response) - else: - return { - "text": "", - "is_finished": False, - "prompt_tokens": 0, - "completion_tokens": 0, - } - else: - print_verbose(f"chunk: {chunk} (Type: {type(chunk)})") - raise ValueError(f"Unable to parse response. Original response: {chunk}") - text: Final = parsed_response.get("text_output", "") - finish_reason: Final = parsed_response.get("stop_reason") - is_finished: Final = parsed_response.get("is_finished", False) - return { - "text": text, - "is_finished": is_finished, - "finish_reason": finish_reason, - "prompt_tokens": parsed_response.get("input_token_count", 0), - "completion_tokens": parsed_response.get("generated_token_count", 0), - } - return {"text": "", "is_finished": False} - except Exception as e: - raise e - - def model_response_creator( - self, chunk: dict | None = None, hidden_params: Mapping[str, object] | None = None - ) -> ModelResponseStream: - _model: Final = self._cached_model_name - _logging_obj_llm_provider: Final = self._cached_logging_llm_provider - - if chunk is None: - args: dict[str, Any] = {"model": _model} - else: - chunk.pop("model", None) - args = {"model": _model} - if chunk: - args.update({k: v for k, v in chunk.items() if k != "stream"}) - - model_response: Final = ModelResponseStream(**args) - if self.response_id is not None: - model_response.id = self.response_id - elif model_response.id: - self.response_id = model_response.id - if self.system_fingerprint is not None: - model_response.system_fingerprint = self.system_fingerprint - - if ( - self.created is not None - ): # maintain same 'created' across all chunks - https://github.com/BerriAI/litellm/issues/11437 - model_response.created = self.created - else: - self.created = model_response.created - - # Spread order is load-bearing: _base_hidden_params (model_id, api_base, ...) - # must win over both caller-supplied hidden_params and the computed - # custom_llm_provider/created_at values, so it comes last. - if hidden_params is not None: - model_response._hidden_params = { - **hidden_params, - "custom_llm_provider": _logging_obj_llm_provider, - "created_at": time.time(), - **self._base_hidden_params, - } - else: - model_response._hidden_params = { - "custom_llm_provider": _logging_obj_llm_provider, - "created_at": time.time(), - **self._base_hidden_params, - } - - if len(model_response.choices) > 0 and getattr(model_response.choices[0], "delta") is not None: - # do nothing, if object instantiated - pass - else: - model_response.choices = [StreamingChoices(finish_reason=None)] - return model_response - - def is_delta_empty(self, delta: Delta) -> bool: - is_empty = True - if delta.content or delta.tool_calls is not None or delta.function_call is not None: - is_empty = False - return is_empty - - def set_model_id(self, id: str, model_response: ModelResponseStream) -> ModelResponseStream: - """ - Set the model id and response id to the given id. - - Ensure model id is always the same across all chunks. - - If a valid ID is received in any chunk, use it for the response. - """ - if self.response_id is None and id and isinstance(id, str) and id.strip(): - self.response_id = id - - if id and isinstance(id, str) and id.strip(): - model_response._hidden_params["received_model_id"] = id - - if self.response_id is not None and isinstance(self.response_id, str): - model_response.id = self.response_id - return model_response - - def copy_model_response_level_provider_specific_fields( - self, - original_chunk: ModelResponseStream | OpenAIChatCompletionChunk, - model_response: ModelResponseStream, - ) -> ModelResponseStream: - """ - Copy provider_specific_fields from original_chunk to model_response. - """ - provider_specific_fields: Final = getattr(original_chunk, "provider_specific_fields", None) - if provider_specific_fields is not None: - model_response.provider_specific_fields = provider_specific_fields - for k, v in provider_specific_fields.items(): - setattr(model_response, k, v) - return model_response - - def is_chunk_non_empty( - self, - completion_obj: dict[str, Any], - model_response: ModelResponseStream, - response_obj: dict[str, Any], - ) -> bool: - if ( - "content" in completion_obj - and (isinstance(completion_obj["content"], str) and len(completion_obj["content"]) > 0) - or ( - "tool_calls" in completion_obj - and completion_obj["tool_calls"] is not None - and len(completion_obj["tool_calls"]) > 0 - ) - or ("function_call" in completion_obj and completion_obj["function_call"] is not None) - or ( - "tool_calls" in model_response.choices[0].delta - and model_response.choices[0].delta["tool_calls"] is not None - and len(model_response.choices[0].delta["tool_calls"]) > 0 - ) - or ( - "function_call" in model_response.choices[0].delta - and model_response.choices[0].delta["function_call"] is not None - ) - or ( - "reasoning_content" in model_response.choices[0].delta - and model_response.choices[0].delta.reasoning_content is not None - ) - or (model_response.choices[0].delta.provider_specific_fields is not None) - or ( - "provider_specific_fields" in model_response - and model_response.choices[0].delta.provider_specific_fields is not None - ) - or ("provider_specific_fields" in response_obj and response_obj["provider_specific_fields"] is not None) - or ( - "annotations" in model_response.choices[0].delta - and model_response.choices[0].delta.annotations is not None - ) - or ( - not self.sent_first_chunk - and hasattr(model_response.choices[0].delta, "role") - and model_response.choices[0].delta.role is not None - ) - or (getattr(model_response.choices[0].delta, "reasoning_items", None) is not None) - ): - return True - else: - return False - - def strip_role_from_delta(self, model_response: ModelResponseStream) -> ModelResponseStream: - """ - Strip the role from the delta. - """ - if self.sent_first_chunk is False: - model_response.choices[0].delta["role"] = "assistant" - self.sent_first_chunk = True - elif self.sent_first_chunk is True and hasattr(model_response.choices[0].delta, "role"): - _initial_delta: Final = model_response.choices[0].delta.model_dump() - - _initial_delta.pop("role", None) - model_response.choices[0].delta = Delta(**_initial_delta) - return model_response - - def _has_special_delta_content(self, model_response: ModelResponseStream) -> bool: - """ - Check if the delta contains special content types (tool_calls, function_call, audio, or image). - """ - if len(model_response.choices) == 0: - return False - - delta: Final = model_response.choices[0].delta - - # Check for tool_calls or function_call - if ( - getattr(delta, TOOL_CALLS_ATTRIBUTE, None) is not None - or getattr(delta, FUNCTION_CALL_ATTRIBUTE, None) is not None - ): - return True - - # Check for audio - if hasattr(delta, AUDIO_ATTRIBUTE) and getattr(delta, AUDIO_ATTRIBUTE, None) is not None: - return True - - # Check for image - if hasattr(delta, IMAGE_ATTRIBUTE) and getattr(delta, IMAGE_ATTRIBUTE, None) is not None: - return True - - return False - - def _handle_special_delta_content(self, model_response: ModelResponseStream) -> ModelResponseStream: - """ - Handle special delta content types by stripping role and returning the response. - """ - return self.strip_role_from_delta(model_response) - - def _has_special_delta_attribute(self, delta, attribute_name: str) -> bool: - """ - Check if delta has a specific attribute and it's not None. - """ - return delta is not None and getattr(delta, attribute_name, None) is not None - - def _copy_delta_attribute(self, source_delta, target_delta, attribute_name: str) -> None: - """ - Copy a specific attribute from source delta to target delta. - """ - setattr(target_delta, attribute_name, getattr(source_delta, attribute_name)) - - def _has_any_special_delta_attributes(self, delta) -> bool: - """ - Check if delta has any special attributes (audio, image). - """ - special_attributes: Final = [AUDIO_ATTRIBUTE, IMAGE_ATTRIBUTE] - for attribute in special_attributes: - if self._has_special_delta_attribute(delta, attribute): - return True - return False - - def _handle_special_delta_attributes(self, delta, model_response: "ModelResponseStream") -> None: - """ - Handle special delta attributes (audio, image) by copying them to model_response. - """ - special_attributes: Final = [AUDIO_ATTRIBUTE, IMAGE_ATTRIBUTE] - for attribute in special_attributes: - if self._has_special_delta_attribute(delta, attribute): - self._copy_delta_attribute(delta, model_response.choices[0].delta, attribute) - - def return_processed_chunk_logic( # noqa: C901 - self, - completion_obj: dict[str, Any], - model_response: ModelResponseStream, - response_obj: dict[str, Any], - ): - from litellm.litellm_core_utils.core_helpers import ( - preserve_upstream_non_openai_attributes, - ) - - is_chunk_non_empty: Final = self.is_chunk_non_empty(completion_obj, model_response, response_obj) - - if is_chunk_non_empty: # cannot set content of an OpenAI Object to be an empty string - self.raise_on_model_repetition() - hold, model_response_str = self.check_special_tokens( - chunk=completion_obj["content"], - finish_reason=model_response.choices[0].finish_reason, - ) # filter out bos/eos tokens from openai-compatible hf endpoints - - if hold is False: - ## check if openai/azure chunk - original_chunk: Final = response_obj.get("original_chunk", None) - if original_chunk: - if len(original_chunk.choices) > 0: - choices: Final = [] - for choice in original_chunk.choices: - try: - if isinstance(choice, BaseModel): - choice_json = choice.model_dump() - choice_json.pop( - "finish_reason", None - ) # for mistral etc. which return a value in their last chunk (not-openai compatible). - choices.append(StreamingChoices(**choice_json)) - except Exception: - choices.append(StreamingChoices()) - setattr(model_response, "choices", choices) - else: - return - model_response.system_fingerprint = original_chunk.system_fingerprint - setattr( - model_response, - "citations", - getattr(original_chunk, "citations", None), - ) - preserve_upstream_non_openai_attributes( - model_response=model_response, - original_chunk=original_chunk, - ) - - model_response = self.strip_role_from_delta(model_response) - if verbose_logger.isEnabledFor(logging.DEBUG): - verbose_logger.debug( - "model_response.choices[0].delta: %s", - model_response.choices[0].delta, - ) - else: - ## else - completion_obj["content"] = model_response_str - if self.sent_first_chunk is False: - completion_obj["role"] = "assistant" - self.sent_first_chunk = True - if response_obj.get("provider_specific_fields") is not None: - completion_obj["provider_specific_fields"] = response_obj["provider_specific_fields"] - model_response.choices[0].delta = Delta(**completion_obj) - _index: Final[int | None] = completion_obj.get("index") - if _index is not None: - model_response.choices[0].index = _index - - self._optional_combine_thinking_block_in_choices(model_response=model_response) - - return model_response - else: - return - elif self.received_finish_reason is not None: - if self.sent_last_chunk is True: - # Bedrock returns the guardrail trace in the last chunk - we want to return this here - if self.custom_llm_provider == "bedrock" and "trace" in model_response: - return model_response - - # Don't raise StopIteration here - some providers (like OpenRouter) - # send usage/cost data in chunks after the finish_reason chunk - if hasattr(model_response, "usage") and model_response.usage is not None: - return model_response - return - # flush any remaining holding chunk - if len(self.holding_chunk) > 0: - if model_response.choices[0].delta.content is None: - model_response.choices[0].delta.content = self.holding_chunk - else: - model_response.choices[0].delta.content = ( - self.holding_chunk + model_response.choices[0].delta.content - ) - self.holding_chunk = "" - # if delta is None - _is_delta_empty: Final = self.is_delta_empty(delta=model_response.choices[0].delta) - - # Preserve custom attributes from original chunk (applies to both - # empty and non-empty delta final chunks). - _original_chunk: Final = response_obj.get("original_chunk", None) - if _original_chunk is not None: - preserve_upstream_non_openai_attributes( - model_response=model_response, - original_chunk=_original_chunk, - ) - - if _is_delta_empty: - model_response.choices[0].delta = Delta(content=None) # ensure empty delta chunk returned - # get any function call arguments - model_response.choices[0].finish_reason = map_finish_reason( - finish_reason=self.received_finish_reason - ) # ensure consistent output to openai - - self.sent_last_chunk = True - - return model_response - elif self._has_special_delta_content(model_response): - return self._handle_special_delta_content(model_response) - else: - if hasattr(model_response, "usage"): - self.chunks.append(model_response) - return - - def _optional_combine_thinking_block_in_choices(self, model_response: ModelResponseStream) -> None: - """ - UI's Like OpenWebUI expect to get 1 chunk with ... tags in the chunk content - - In place updates the model_response object with reasoning_content in content with ... tags - - Enabled when `merge_reasoning_content_in_choices=True` passed in request params - - - """ - if self.merge_reasoning_content_in_choices is True: - reasoning_content: Final = getattr(model_response.choices[0].delta, "reasoning_content", None) - if reasoning_content: - if self.sent_first_thinking_block is False: - # Ensure content is not None before concatenation - if model_response.choices[0].delta.content is None: - model_response.choices[0].delta.content = "" - model_response.choices[0].delta.content += "" + reasoning_content - self.sent_first_thinking_block = True - elif ( - self.sent_first_thinking_block is True - and hasattr(model_response.choices[0].delta, "reasoning_content") - and model_response.choices[0].delta.reasoning_content - ): - model_response.choices[0].delta.content = reasoning_content - elif ( - self.sent_first_thinking_block is True - and not self.sent_last_thinking_block - and model_response.choices[0].delta.content - ): - model_response.choices[0].delta.content = "" + (model_response.choices[0].delta.content or "") - self.sent_last_thinking_block = True - - if hasattr(model_response.choices[0].delta, "reasoning_content"): - del model_response.choices[0].delta.reasoning_content - - def _dispatch_provider_chunk( - self, - chunk: Any, - model_response: ModelResponseStream, - completion_obj: dict[str, Any], - ) -> _ProviderChunkResult: - response_obj: dict[str, Any] = {} - if ( - isinstance(chunk, ModelResponseStream) - and self.custom_llm_provider is not None - and self.custom_llm_provider in litellm._custom_providers - ): - _has_content: Final = bool( - chunk.choices - and chunk.choices[0].delta is not None - and (chunk.choices[0].delta.content or chunk.choices[0].delta.tool_calls) - ) - if self.received_finish_reason is not None: - if not _has_content: - raise StopIteration - if chunk.choices and chunk.choices[0].finish_reason: - self.received_finish_reason = chunk.choices[0].finish_reason - if not _has_content: - return _ProviderChunkEarlyReturn(None) - # Strip finish_reason from the content chunk so it appears - # only on the trailing empty-delta chunk (OpenAI spec). - # finish_reason_handler() will emit the proper terminal chunk. - chunk.choices[0].finish_reason = None - return _ProviderChunkEarlyReturn(chunk) - - if ( - isinstance(chunk, dict) - and generic_chunk_has_all_required_fields(chunk=chunk) # check if chunk is a generic streaming chunk - ) or (self.custom_llm_provider and self.custom_llm_provider in litellm._custom_providers): - if self.received_finish_reason is not None: - _chunk_has_content: Final = isinstance(chunk, dict) and ( - bool(chunk.get("text", "")) - or chunk.get("tool_use") is not None - # Usage-only final chunks are valid and needed to surface - # finish_reason/usage to downstream translators. - or chunk.get("usage") is not None - ) - if not _chunk_has_content and (not isinstance(chunk, dict) or "provider_specific_fields" not in chunk): - raise StopIteration - anthropic_response_obj: Final[GChunk] = cast(GChunk, chunk) - completion_obj["content"] = anthropic_response_obj["text"] - if anthropic_response_obj["is_finished"]: - self.received_finish_reason = anthropic_response_obj["finish_reason"] - - if anthropic_response_obj["finish_reason"]: - self.intermittent_finish_reason = anthropic_response_obj["finish_reason"] - - if anthropic_response_obj["usage"] is not None: - setattr( - model_response, - "usage", - litellm.Usage(**anthropic_response_obj["usage"]), - ) - - if "tool_use" in anthropic_response_obj and anthropic_response_obj["tool_use"] is not None: - completion_obj["tool_calls"] = [anthropic_response_obj["tool_use"]] - - if ( - "provider_specific_fields" in anthropic_response_obj - and anthropic_response_obj["provider_specific_fields"] is not None - ): - for key, value in anthropic_response_obj["provider_specific_fields"].items(): - setattr(model_response, key, value) - - response_obj = cast(dict[str, object], anthropic_response_obj) - elif self.model == "replicate" or self.custom_llm_provider == "replicate": - response_obj = self.handle_replicate_chunk(chunk) - completion_obj["content"] = response_obj["text"] - if response_obj["is_finished"]: - self.received_finish_reason = response_obj["finish_reason"] - elif self.custom_llm_provider and self.custom_llm_provider == "predibase": - response_obj = self.handle_predibase_chunk(chunk) - completion_obj["content"] = response_obj["text"] - if response_obj["is_finished"]: - self.received_finish_reason = response_obj["finish_reason"] - elif self.custom_llm_provider and self.custom_llm_provider == "baseten": # baseten doesn't provide streaming - completion_obj["content"] = self.handle_baseten_chunk(chunk) - elif self.custom_llm_provider and self.custom_llm_provider == "ai21": # ai21 doesn't provide streaming - response_obj = self.handle_ai21_chunk(chunk) - completion_obj["content"] = response_obj["text"] - if response_obj["is_finished"]: - self.received_finish_reason = response_obj["finish_reason"] - elif self.custom_llm_provider and self.custom_llm_provider == "maritalk": - response_obj = self.handle_maritalk_chunk(chunk) - completion_obj["content"] = response_obj["text"] - if response_obj["is_finished"]: - self.received_finish_reason = response_obj["finish_reason"] - elif self.custom_llm_provider and self.custom_llm_provider == "vllm": - completion_obj["content"] = chunk[0].outputs[0].text - elif ( - self.custom_llm_provider and self.custom_llm_provider == "aleph_alpha" - ): # aleph alpha doesn't provide streaming - response_obj = self.handle_aleph_alpha_chunk(chunk) - completion_obj["content"] = response_obj["text"] - if response_obj["is_finished"]: - self.received_finish_reason = response_obj["finish_reason"] - elif self.custom_llm_provider == "nlp_cloud": - try: - response_obj = self.handle_nlp_cloud_chunk(chunk) - completion_obj["content"] = response_obj["text"] - if response_obj["is_finished"]: - self.received_finish_reason = response_obj["finish_reason"] - except Exception as e: - if self.received_finish_reason: - raise e - else: - if self.sent_first_chunk is False: - raise Exception("An unknown error occurred with the stream") - self.received_finish_reason = "stop" - elif self.custom_llm_provider == "vertex_ai" and not isinstance(chunk, ModelResponseStream): - vertex_chunk: Final = cast(_VertexChunkLike, chunk) - import proto - - if hasattr(vertex_chunk, "candidates") is True: - try: - try: - completion_obj["content"] = vertex_chunk.text - except Exception as e: - original_exception: Final = e - if "Part has no text." in str(e): - ## check for function calling - function_call: Final = vertex_chunk.candidates[0].content.parts[0].function_call - - args_dict: Final = {} - - # Check if it's a RepeatedComposite instance - for key, val in function_call.args.items(): - if isinstance( - val, - proto.marshal.collections.repeated.RepeatedComposite, - ): - # If so, convert to list - args_dict[key] = [v for v in val] - else: - args_dict[key] = val - - try: - args_str: Final = json.dumps(args_dict) - except Exception as e: - raise e - _delta_obj: Final = litellm.utils.Delta( - content=None, - tool_calls=[ - { - "id": f"call_{uuid.uuid4()}", - "function": { - "arguments": args_str, - "name": function_call.name, - }, - "type": "function", - } - ], - ) - _streaming_response: Final = StreamingChoices(delta=_delta_obj) - _model_response: Final = ModelResponseStream() - _model_response.choices = [_streaming_response] - response_obj = {"original_chunk": _model_response} - else: - raise original_exception - if ( - hasattr(vertex_chunk.candidates[0], "finish_reason") - and vertex_chunk.candidates[0].finish_reason.name != "FINISH_REASON_UNSPECIFIED" - ): # every non-final chunk in vertex ai has this - self.received_finish_reason = map_finish_reason(vertex_chunk.candidates[0].finish_reason.name) - except Exception: - if vertex_chunk.candidates[0].finish_reason.name == "SAFETY": - raise Exception(f"The response was blocked by VertexAI. {vertex_chunk}") - else: - completion_obj["content"] = str(vertex_chunk) - elif self.custom_llm_provider == "petals": - if self.completion_stream is None or len(self.completion_stream) == 0: - if self.received_finish_reason is not None: - raise StopIteration - else: - self.received_finish_reason = "stop" - chunk_size = 30 - stream = cast(Any, self.completion_stream) - new_chunk = stream[:chunk_size] - completion_obj["content"] = new_chunk - self.completion_stream = stream[chunk_size:] - elif self.custom_llm_provider == "palm": - # fake streaming - response_obj = {} - if self.completion_stream is None or len(self.completion_stream) == 0: - if self.received_finish_reason is not None: - raise StopIteration - else: - self.received_finish_reason = "stop" - chunk_size = 30 - stream = cast(Any, self.completion_stream) - new_chunk = stream[:chunk_size] - completion_obj["content"] = new_chunk - self.completion_stream = stream[chunk_size:] - elif self.custom_llm_provider == "triton": - response_obj = self.handle_triton_stream(chunk) - completion_obj["content"] = response_obj["text"] - print_verbose(f"completion obj content: {completion_obj['content']}") - if response_obj["is_finished"]: - self.received_finish_reason = response_obj["finish_reason"] - elif self.custom_llm_provider == "text-completion-openai": - response_obj = self.handle_openai_text_completion_chunk(chunk) - completion_obj["content"] = response_obj["text"] - print_verbose(f"completion obj content: {completion_obj['content']}") - if response_obj["is_finished"]: - self.received_finish_reason = response_obj["finish_reason"] - if response_obj["usage"] is not None: - _text_completion_usage: Final[Usage] = response_obj["usage"] - setattr( - model_response, - "usage", - litellm.Usage( - prompt_tokens=_text_completion_usage.prompt_tokens, - completion_tokens=_text_completion_usage.completion_tokens, - total_tokens=_text_completion_usage.total_tokens, - ), - ) - elif self.custom_llm_provider == "text-completion-codestral": - if not isinstance(chunk, str): - raise ValueError(f"chunk is not a string: {chunk}") - response_obj = cast( - dict[str, object], - litellm.CodestralTextCompletionConfig()._chunk_parser(chunk), - ) - completion_obj["content"] = response_obj["text"] - print_verbose(f"completion obj content: {completion_obj['content']}") - if response_obj["is_finished"]: - self.received_finish_reason = response_obj["finish_reason"] - if "usage" in response_obj is not None: - _codestral_usage: Final[Usage] = response_obj["usage"] - setattr( - model_response, - "usage", - litellm.Usage( - prompt_tokens=_codestral_usage.prompt_tokens, - completion_tokens=_codestral_usage.completion_tokens, - total_tokens=_codestral_usage.total_tokens, - ), - ) - elif self.custom_llm_provider == "azure_text": - response_obj = self.handle_azure_text_completion_chunk(chunk) - completion_obj["content"] = response_obj["text"] - print_verbose(f"completion obj content: {completion_obj['content']}") - if response_obj["is_finished"]: - self.received_finish_reason = response_obj["finish_reason"] - elif self.custom_llm_provider == "cached_response": - cached_chunk: Final = cast(ModelResponseStream, chunk) - chunk_finish_reason: Final = cached_chunk.choices[0].finish_reason - response_obj = { - "text": cached_chunk.choices[0].delta.content, - "is_finished": chunk_finish_reason is not None, - "finish_reason": chunk_finish_reason, - "original_chunk": cached_chunk, - "tool_calls": ( - cached_chunk.choices[0].delta.tool_calls - if hasattr(cached_chunk.choices[0].delta, "tool_calls") - else None - ), - } - - completion_obj["content"] = response_obj["text"] - if response_obj["tool_calls"] is not None: - completion_obj["tool_calls"] = response_obj["tool_calls"] - print_verbose(f"completion obj content: {completion_obj['content']}") - if hasattr(cached_chunk, "id"): - model_response.id = cached_chunk.id - self.response_id = cached_chunk.id - if hasattr(cached_chunk, "system_fingerprint"): - self.system_fingerprint = cached_chunk.system_fingerprint - if response_obj["is_finished"]: - self.received_finish_reason = response_obj["finish_reason"] - else: # openai / azure chat model - if self.custom_llm_provider in [ - LlmProviders.AZURE.value, - LlmProviders.AZURE_AI.value, - ]: - if isinstance(chunk, BaseModel) and hasattr(chunk, "model"): - # for azure, we need to pass the model from the original chunk - self.model = getattr(chunk, "model", self.model) - response_obj = self.handle_openai_chat_completion_chunk(chunk) - if response_obj is None: - return _ProviderChunkEarlyReturn(None) - completion_obj["content"] = response_obj["text"] - self.intermittent_finish_reason = response_obj.get("finish_reason", None) - if response_obj["is_finished"]: - if response_obj["finish_reason"] == "error": - raise Exception( - f"{self.custom_llm_provider} raised a streaming error - finish_reason: error, no content string given. Received Chunk={response_obj}" - ) - self.received_finish_reason = response_obj["finish_reason"] - if response_obj.get("original_chunk", None) is not None: - if hasattr(response_obj["original_chunk"], "id"): - model_response = self.set_model_id(response_obj["original_chunk"].id, model_response) - if hasattr(response_obj["original_chunk"], "system_fingerprint"): - model_response.system_fingerprint = response_obj["original_chunk"].system_fingerprint - self.system_fingerprint = response_obj["original_chunk"].system_fingerprint - if response_obj["logprobs"] is not None: - model_response.choices[0].logprobs = response_obj["logprobs"] - - if response_obj["usage"] is not None: - if isinstance(response_obj["usage"], dict): - setattr( - model_response, - "usage", - litellm.Usage( - prompt_tokens=response_obj["usage"].get("prompt_tokens", None) or None, - completion_tokens=response_obj["usage"].get("completion_tokens", None) or None, - total_tokens=response_obj["usage"].get("total_tokens", None) or None, - ), - ) - elif isinstance(response_obj["usage"], Usage): - setattr( - model_response, - "usage", - response_obj["usage"], - ) - elif isinstance(response_obj["usage"], BaseModel): - setattr( - model_response, - "usage", - litellm.Usage(**response_obj["usage"].model_dump()), - ) - return _ProviderChunkParsed(response_obj) - - def chunk_creator(self, chunk: Any): - if hasattr(chunk, "id"): - self.response_id = chunk.id - provider_response_model: Final = _provider_response_model(chunk) - if provider_response_model is not None: - self._provider_response_model = provider_response_model - model_response = self.model_response_creator( - hidden_params=_provider_hidden_params(chunk, self._provider_response_model) - ) - response_obj: dict[str, Any] = {} - try: - # return this for all models - completion_obj: Final[dict[str, Any]] = {"content": ""} - dispatch_result: Final = self._dispatch_provider_chunk( - chunk=chunk, - model_response=model_response, - completion_obj=completion_obj, - ) - if isinstance(dispatch_result, _ProviderChunkEarlyReturn): - return dispatch_result.value - response_obj = dispatch_result.response_obj - - model_response.model = self.model - ## FUNCTION CALL PARSING - original_chunk: Final = response_obj.get("original_chunk") if response_obj is not None else None - if ( - original_chunk is not None - ): # function / tool calling branch - only set for openai/azure compatible endpoints - # enter this branch when no content has been passed in response - if hasattr(original_chunk, "id"): - model_response = self.set_model_id(original_chunk.id, model_response) - if hasattr(original_chunk, "provider_specific_fields"): - model_response = self.copy_model_response_level_provider_specific_fields( - original_chunk, model_response - ) - if original_chunk.choices and len(original_chunk.choices) > 0: - delta = original_chunk.choices[0].delta - if delta is not None and (delta.function_call is not None or delta.tool_calls is not None): - try: - model_response.system_fingerprint = original_chunk.system_fingerprint - ## AZURE - check if arguments is not None - if original_chunk.choices[0].delta.function_call is not None: - if ( - getattr( - original_chunk.choices[0].delta.function_call, - "arguments", - ) - is None - ): - original_chunk.choices[0].delta.function_call.arguments = "" - elif original_chunk.choices[0].delta.tool_calls is not None: - if isinstance(original_chunk.choices[0].delta.tool_calls, list): - for t in original_chunk.choices[0].delta.tool_calls: - if hasattr(t, "functions") and hasattr(t.functions, "arguments"): - if ( - getattr( - t.function, - "arguments", - ) - is None - ): - t.function.arguments = "" - _json_delta: Final[_DeltaDumpDict] = delta.model_dump() - if "role" not in _json_delta or _json_delta["role"] is None: - _json_delta["role"] = "assistant" # mistral's api returns role as None - if "tool_calls" in _json_delta and isinstance(_json_delta["tool_calls"], list): - for tool in _json_delta["tool_calls"]: - if ( - isinstance(tool, dict) - and "function" in tool - and isinstance(tool["function"], dict) - and ("type" not in tool or tool["type"] is None) - ): - # if function returned but type set to None - mistral's api returns type: None - tool["type"] = "function" - model_response.choices[0].delta = Delta(**_json_delta) - except Exception as e: - verbose_logger.exception( - "litellm.CustomStreamWrapper.chunk_creator(): Exception occured - %s", e - ) - model_response.choices[0].delta = Delta() - elif self._has_any_special_delta_attributes(delta): - self._handle_special_delta_attributes(delta, model_response) - else: - try: - delta = ( - dict() - if original_chunk.choices[0].delta is None - else dict(original_chunk.choices[0].delta) - ) - model_response.choices[0].delta = Delta(**delta) - except Exception: - model_response.choices[0].delta = Delta() - else: - if self.stream_options is not None and self.stream_options["include_usage"] is True: - model_response.choices = [] - return model_response - self._record_usage_only_chunk(model_response=model_response) - return - ## CHECK FOR TOOL USE - - if "tool_calls" in completion_obj and len(completion_obj["tool_calls"]) > 0: - if self.is_function_call is True: # user passed in 'functions' param - completion_obj["function_call"] = completion_obj["tool_calls"][0]["function"] - completion_obj["tool_calls"] = None - - self.tool_call = True - - if hasattr(chunk, "usage") and chunk.usage is not None: - model_response.usage = chunk.usage - - ## RETURN ARG - result: Final = self.return_processed_chunk_logic( - completion_obj=completion_obj, - model_response=model_response, - response_obj=response_obj, - ) - return result - - except StopIteration: - raise StopIteration - except Exception as e: - traceback.format_exc() - setattr(e, "message", str(e)) - raise exception_type( - model=self.model, - custom_llm_provider=self.custom_llm_provider, - original_exception=e, - ) - - def set_logging_event_loop(self, loop): - """ - import litellm, asyncio - - loop = asyncio.get_event_loop() # 👈 gets the current event loop - - response = litellm.completion(.., stream=True) - - response.set_logging_event_loop(loop=loop) # 👈 enables async_success callbacks for sync logging - - for chunk in response: - ... - """ - self.logging_loop = loop - - async def _call_post_streaming_deployment_hook(self, chunk): - """ - Call the post-call streaming deployment hook for callbacks. - - This allows callbacks to modify streaming chunks before they're returned. - """ - try: - import litellm - from litellm.integrations.custom_logger import CustomLogger - from litellm.types.utils import CallTypes - - if self._post_streaming_hooks is None: - self._post_streaming_hooks = [ - cb - for cb in litellm.callbacks - if isinstance(cb, CustomLogger) and hasattr(cb, "async_post_call_streaming_deployment_hook") - ] - - if not self._post_streaming_hooks: - return chunk - - request_data: Final = self.logging_obj.model_call_details - call_type_str: Final = self.logging_obj.call_type - - try: - typed_call_type = CallTypes(call_type_str) - except ValueError: - typed_call_type = None - - for callback in self._post_streaming_hooks: - result = await callback.async_post_call_streaming_deployment_hook( - request_data=request_data, - response_chunk=chunk, - call_type=typed_call_type, - ) - if result is not None: - chunk = result - - return chunk - except Exception as e: - from litellm._logging import verbose_logger - - verbose_logger.exception("Error in post-call streaming deployment hook: %s", e) - return chunk - - def _add_mcp_list_tools_to_first_chunk(self, chunk: ModelResponseStream) -> ModelResponseStream: - """ - Add mcp_list_tools from _hidden_params to the first chunk's delta.provider_specific_fields. - - This method checks if MCP metadata with mcp_list_tools is stored in _hidden_params - and adds it to the first chunk's delta.provider_specific_fields. - """ - try: - # Check if MCP metadata should be added to first chunk - if not hasattr(self, "_hidden_params") or not self._hidden_params: - return chunk - - mcp_metadata: Final = self._hidden_params.get("mcp_metadata") - if not mcp_metadata or not isinstance(mcp_metadata, dict): - return chunk - - # Only add mcp_list_tools to first chunk (not tool_calls or tool_results) - mcp_list_tools: Final = mcp_metadata.get("mcp_list_tools") - if not mcp_list_tools: - return chunk - - # Add mcp_list_tools to delta.provider_specific_fields - if hasattr(chunk, "choices") and chunk.choices: - for choice in chunk.choices: - if isinstance(choice, StreamingChoices) and hasattr(choice, "delta") and choice.delta: - # Get existing provider_specific_fields or create new dict - provider_fields = getattr(choice.delta, "provider_specific_fields", None) or {} - - # Add only mcp_list_tools to first chunk - provider_fields["mcp_list_tools"] = mcp_list_tools - - # Set the provider_specific_fields - setattr(choice.delta, "provider_specific_fields", provider_fields) - - except Exception as e: - from litellm._logging import verbose_logger - - verbose_logger.exception("Error adding MCP list tools to first chunk: %s", e) - - return chunk - - def _add_mcp_metadata_to_final_chunk(self, chunk: ModelResponseStream) -> ModelResponseStream: - """ - Add MCP metadata from _hidden_params to the final chunk's delta.provider_specific_fields. - - This method checks if MCP metadata is stored in _hidden_params and adds it to - the chunk's delta.provider_specific_fields, similar to how RAG adds search results. - """ - try: - # Check if MCP metadata should be added to final chunk - if not hasattr(self, "_hidden_params") or not self._hidden_params: - return chunk - - mcp_metadata: Final = self._hidden_params.get("mcp_metadata") - if not mcp_metadata: - return chunk - - # Add MCP metadata to delta.provider_specific_fields - if hasattr(chunk, "choices") and chunk.choices: - for choice in chunk.choices: - if isinstance(choice, StreamingChoices) and hasattr(choice, "delta") and choice.delta: - # Get existing provider_specific_fields or create new dict - provider_fields = getattr(choice.delta, "provider_specific_fields", None) or {} - - # Add MCP metadata - if isinstance(mcp_metadata, dict): - provider_fields.update(mcp_metadata) - - # Set the provider_specific_fields - setattr(choice.delta, "provider_specific_fields", provider_fields) - - except Exception as e: - from litellm._logging import verbose_logger - - verbose_logger.exception("Error adding MCP metadata to final chunk: %s", e) - - return chunk - - def cache_streaming_response(self, processed_chunk, cache_hit: bool): - """ - Caches the streaming response - """ - if not cache_hit and self.logging_obj._llm_caching_handler is not None: - self.logging_obj._llm_caching_handler._sync_add_streaming_response_to_cache(processed_chunk) - - async def async_cache_streaming_response(self, processed_chunk, cache_hit: bool): - """ - Caches the streaming response - """ - if not cache_hit and self.logging_obj._llm_caching_handler is not None: - await self.logging_obj._llm_caching_handler._add_streaming_response_to_cache(processed_chunk) - - def run_success_logging_and_cache_storage(self, processed_chunk, cache_hit: bool): - """ - Runs success logging in a thread and adds the response to the cache - """ - if litellm.disable_streaming_logging is True: - """ - [NOT RECOMMENDED] - Set this via `litellm.disable_streaming_logging = True`. - - Disables streaming logging. - """ - return - ## ASYNC LOGGING - # Create an event loop for the new thread - if self.logging_loop is not None: - future: Final = asyncio.run_coroutine_threadsafe( - self.logging_obj.async_success_handler(processed_chunk, None, None, cache_hit), - loop=self.logging_loop, - ) - future.result() - else: - asyncio.run(self.logging_obj.async_success_handler(processed_chunk, None, None, cache_hit)) - ## SYNC LOGGING — only for sync SDK entrypoints; async proxy paths export via async_success_handler - litellm_params: Final = self.logging_obj.model_call_details.get("litellm_params", {}) - if self.logging_obj._is_sync_litellm_request(litellm_params): - self.logging_obj.success_handler(processed_chunk, None, None, cache_hit) - - _PROVIDERS_WITHOUT_STREAM_FINISH_REASON = frozenset({"baseten", "vllm"}) - - def _has_provider_finish_reason(self) -> bool: - return self.received_finish_reason is not None or self.intermittent_finish_reason is not None - - def _should_require_provider_finish_reason(self) -> bool: - return self.custom_llm_provider not in self._PROVIDERS_WITHOUT_STREAM_FINISH_REASON - - def _raise_incomplete_stream_without_finish_reason(self) -> "NoReturn": - message = ( - "Stream ended without a finish_reason from the provider. " - "Partial content was received but the response was not successfully completed." - ) - self._record_partial_usage_for_failure() - self._handle_stream_fallback_error(RuntimeError(message)) - - def finish_reason_handler(self): - model_response: Final = self.model_response_creator() - _finish_reason: Final = self.received_finish_reason or self.intermittent_finish_reason - if _finish_reason is not None: - model_response.choices[0].finish_reason = _finish_reason - else: - model_response.choices[0].finish_reason = "stop" - - ## if tool use - if ( - model_response.choices[0].finish_reason == "stop" and self.tool_call - ): # don't overwrite for other - potential error finish reasons - model_response.choices[0].finish_reason = "tool_calls" - return model_response - - def _record_usage_only_chunk(self, model_response: "ModelResponseStream") -> None: - """ - Keep provider usage-only chunks (e.g. OpenRouter's post-finish chunk, which carries a - provider-reported cost) available to cost tracking. They are never returned to the - caller; ``stream_options.include_usage`` only controls what the caller sees. - """ - if getattr(model_response, "usage", None) is None: - return - self.chunks.append(model_response.model_copy(update={"choices": []})) - - @staticmethod - def _resolve_provider_reported_cost(usage_cost: object) -> float | None: - """ - Providers report usage.cost either as a number or as a breakdown object - whose total lives under ``total_cost``. - """ - if isinstance(usage_cost, bool): - return None - if isinstance(usage_cost, (int, float)): - return float(usage_cost) - if isinstance(usage_cost, dict): - return CustomStreamWrapper._resolve_provider_reported_cost(usage_cost.get("total_cost")) - return None - - @staticmethod - def _propagate_usage_cost_to_hidden_params( - response: "ModelResponse", - custom_llm_provider: str | None, - ) -> None: - if custom_llm_provider not in _USAGE_COST_HEADER_PROVIDERS: - return - _usage: Final[Usage | None] = getattr(response, "usage", None) - _cost: Final = CustomStreamWrapper._resolve_provider_reported_cost(getattr(_usage, "cost", None)) - if _cost is not None: - if "additional_headers" not in response._hidden_params: - response._hidden_params["additional_headers"] = {} - response._hidden_params["additional_headers"]["llm_provider-x-litellm-response-cost"] = _cost - - def __next__(self) -> "ModelResponseStream": - cache_hit = False - if self.custom_llm_provider is not None and self.custom_llm_provider == "cached_response": - cache_hit = True - self._check_max_streaming_duration() - try: - if self.completion_stream is None: - self.fetch_sync_stream() - - while True: - if ( - isinstance(self.completion_stream, str) - or isinstance(self.completion_stream, bytes) - or isinstance(self.completion_stream, ModelResponse) - ): - chunk = self.completion_stream - else: - chunk = next(self.completion_stream) - if chunk is not None and chunk != b"": - print_verbose( - f"PROCESSED CHUNK PRE CHUNK CREATOR: {chunk.decode('utf-8', errors='replace') if isinstance(chunk, bytes) else chunk}; custom_llm_provider: {self.custom_llm_provider}" - ) - response: ModelResponseStream | None = self.chunk_creator(chunk=chunk) - print_verbose(f"PROCESSED CHUNK POST CHUNK CREATOR: {response}") - - if response is None: - continue - if self.logging_obj.completion_start_time is None: - self.logging_obj._update_completion_start_time(completion_start_time=datetime.datetime.now()) - ## LOGGING - if not litellm.disable_streaming_logging: - executor.submit( - self.run_success_logging_and_cache_storage, - response, - cache_hit, - ) # log response - if response.choices: - choice = response.choices[0] - if isinstance(choice, StreamingChoices): - self.response_uptil_now += choice.delta.get("content", "") or "" - else: - self.response_uptil_now += "" - self.rules.post_call_rules(input=self.response_uptil_now, model=self.model) - # HANDLE STREAM OPTIONS - self.chunks.append(response) - - # Add mcp_list_tools to first chunk if present - if not self.sent_first_chunk and response.choices: - response = self._add_mcp_list_tools_to_first_chunk(response) - self.sent_first_chunk = True - - # ModelResponseStream declares `usage` as a field, so - # hasattr(response, "usage") is always True — must check - # `is not None` to avoid running this path on every chunk. - if getattr(response, "usage", None) is not None: - usage_to_preserve = response.usage - if usage_to_preserve: - response._hidden_params["usage"] = usage_to_preserve - - obj_dict = response.model_dump() - - if "usage" in obj_dict: - del obj_dict["usage"] - - response = self.model_response_creator(chunk=obj_dict, hidden_params=response._hidden_params) - ## check if empty - is_empty = is_model_response_stream_empty(model_response=cast(ModelResponseStream, response)) - - if is_empty: - continue - # add usage as hidden param - if self.sent_last_chunk is True and self.stream_options is None: - usage = calculate_total_usage(chunks=self.chunks) - response._hidden_params["usage"] = usage - self._last_returned_hidden_params = response._hidden_params - # Add MCP metadata to final chunk if present - response = self._add_mcp_metadata_to_final_chunk(response) - # RETURN RESULT - return response - - except StopIteration: - if self.sent_last_chunk is True: - try: - complete_streaming_response = litellm.stream_chunk_builder( - chunks=self.chunks, - messages=self.messages, - logging_obj=self.logging_obj, - ) - except Exception as e: - # stream_chunk_builder can re-raise (as APIError) on large agentic - # streams. The raise originates inside this except-StopIteration block, - # so the sibling `except Exception` below does not catch it; it would - # escape __next__ and drop the request from SpendLogs. Recover - # best-effort usage from the raw chunks so cost is still tracked - verbose_logger.warning( - "stream_chunk_builder raised at end-of-stream (%s); logging best-effort usage from chunks.", - str(e), - ) - try: - complete_streaming_response = self.model_response_creator( - chunk={"usage": calculate_total_usage(chunks=self.chunks)} - ) - except Exception: - complete_streaming_response = None - - response = self.model_response_creator() - if complete_streaming_response is not None: - self._propagate_usage_cost_to_hidden_params(complete_streaming_response, self.custom_llm_provider) - - setattr( - response, - "usage", - getattr(complete_streaming_response, "usage"), - ) - try: - _cache_copy = complete_streaming_response.model_copy(deep=True) - _log_copy = complete_streaming_response.model_copy(deep=True) - except RuntimeError: - _cache_copy = complete_streaming_response.model_copy() - _log_copy = complete_streaming_response.model_copy() - self.cache_streaming_response( - processed_chunk=_cache_copy, - cache_hit=cache_hit, - ) - executor.submit( - self.logging_obj.success_handler, - _log_copy, - None, - None, - cache_hit, - ) - else: - executor.submit( - self.logging_obj.success_handler, - response, - None, - None, - cache_hit, - ) - # Update hidden_params with final usage from - # stream_chunk_builder. Some providers (e.g. OpenRouter) - # send usage in a chunk after finish_reason, which arrives - # after _hidden_params["usage"] was initially set. The - # _hidden_params dict is the same object the user received - # (shared by reference), so mutating it here also corrects - # the user's copy. - if ( - self.stream_options is None - and complete_streaming_response is not None - and self._last_returned_hidden_params is not None - ): - final_usage: Final = getattr(complete_streaming_response, "usage", None) - if final_usage is not None: - self._last_returned_hidden_params["usage"] = final_usage - - if self.sent_stream_usage is False and self.send_stream_usage is True: - self.sent_stream_usage = True - return response - self._restore_consumer_correlation_context() - raise # Re-raise StopIteration - else: - if self._should_require_provider_finish_reason() and not self._has_provider_finish_reason(): - self._raise_incomplete_stream_without_finish_reason() - self.sent_last_chunk = True - processed_chunk: Final = self.finish_reason_handler() - if self.stream_options is None: # add usage as hidden param - usage = calculate_total_usage(chunks=self.chunks) - processed_chunk._hidden_params["usage"] = usage - ## LOGGING - executor.submit( - self.run_success_logging_and_cache_storage, - processed_chunk, - cache_hit, - ) # log response - # Deliberately do NOT restore context here even though - # completion_stream is already exhausted: this chunk is still - # real data belonging to this call, and the caller's own - # (application-level) log statements processing it run - # immediately after this return, in this same synchronous - # frame - restoring first would make those lines carry the - # wrong ids, which is exactly what leaving context open during - # iteration is meant to prevent (see - # _restore_consumer_correlation_context's docstring). A caller - # that keeps iterating gets cleaned up on its next __next__() - # call (immediate StopIteration, handled above); one that - # stops right here relies on aclose() or the best-effort - # __del__ guard instead. - return processed_chunk - except Exception as e: - traceback_exception: Final = traceback.format_exc() - # LOG FAILURE - handle streaming failure logging in the _next_ object, remove `handle_failure` once it's deprecated - threading.Thread(target=self.logging_obj.failure_handler, args=(e, traceback_exception)).start() - self._handle_stream_fallback_error(e) - - def fetch_sync_stream(self): - if self.completion_stream is None and self.make_call is not None: - # Call make_call to get the completion stream - self.completion_stream = self.make_call(client=litellm.module_level_client) - self._stream_iter = self.completion_stream.__iter__() - - return self.completion_stream - - async def fetch_stream(self): - if self.completion_stream is None and self.make_call is not None: - # Call make_call to get the completion stream - self.completion_stream = await self.make_call(client=litellm.module_level_aclient) - self._stream_iter = self.completion_stream.__aiter__() - - return self.completion_stream - - async def __anext__(self) -> "ModelResponseStream": - cache_hit = False - if self.custom_llm_provider is not None and self.custom_llm_provider == "cached_response": - cache_hit = True - try: - # Inside the try (not before it) so a raised litellm.Timeout flows - # through the same except Exception -> _handle_stream_fallback_error - # path as every other failure, restoring the consumer's correlation - # context - a check before the try would bypass that entirely. - self._check_max_streaming_duration() - if self.completion_stream is None: - await self.fetch_stream() - - if is_async_iterable(self.completion_stream): - async for chunk in self.completion_stream: # pyright: ignore[reportOptionalIterable] # is_async_iterable guard proves __aiter__ - if chunk == "None" or chunk is None: - continue # skip None chunks - - elif self.custom_llm_provider == "gemini" and hasattr(chunk, "parts") and len(chunk.parts) == 0: - continue - processed_chunk: ModelResponseStream | None = self.chunk_creator(chunk=chunk) - if processed_chunk is None: - continue - - if self.logging_obj.completion_start_time is None: - self.logging_obj._update_completion_start_time(completion_start_time=datetime.datetime.now()) - - if processed_chunk.choices: - choice = processed_chunk.choices[0] - if isinstance(choice, StreamingChoices): - self.response_uptil_now += choice.delta.get("content", "") or "" - else: - self.response_uptil_now += "" - self.rules.post_call_rules(input=self.response_uptil_now, model=self.model) - # Add mcp_list_tools to first chunk if present - if not self.sent_first_chunk and processed_chunk.choices: - processed_chunk = self._add_mcp_list_tools_to_first_chunk(processed_chunk) - self.sent_first_chunk = True - - _has_usage = ( - hasattr(processed_chunk, "usage") and getattr(processed_chunk, "usage", None) is not None - ) - - if _has_usage: - # Store a copy ONLY when usage stripping below will mutate - # the chunk. For non-usage chunks (vast majority), store - # directly to avoid expensive model_copy() per chunk. - self.chunks.append(processed_chunk.model_copy()) - - # Strip usage from the outgoing chunk so it's not sent twice - # (once in the chunk, once in _hidden_params). - obj_dict = processed_chunk.model_dump() - if "usage" in obj_dict: - del obj_dict["usage"] - processed_chunk = self.model_response_creator( - chunk=obj_dict, hidden_params=processed_chunk._hidden_params - ) - is_empty = is_model_response_stream_empty( - model_response=cast(ModelResponseStream, processed_chunk) - ) - if is_empty: - continue - else: - # No usage data — safe to store directly without copying - self.chunks.append(processed_chunk) - - # add usage as hidden param - if self.sent_last_chunk is True and self.stream_options is None: - usage = calculate_total_usage(chunks=self.chunks) - processed_chunk._hidden_params["usage"] = usage - self._last_returned_hidden_params = processed_chunk._hidden_params - - # Call post-call streaming deployment hook for final chunk - if self.sent_last_chunk is True: - processed_chunk = await self._call_post_streaming_deployment_hook(processed_chunk) - # Add MCP metadata to final chunk if present (after hooks) - processed_chunk = self._add_mcp_metadata_to_final_chunk(processed_chunk) - - return processed_chunk - raise StopAsyncIteration - else: # temporary patch for non-aiohttp async calls - # example - boto3 bedrock llms - while True: - if isinstance(self.completion_stream, str) or isinstance(self.completion_stream, bytes): - chunk = self.completion_stream - else: - chunk = await asyncio.to_thread(_next_sync_or_exhausted, self.completion_stream) - if chunk is _SYNC_ITER_EXHAUSTED: - raise StopAsyncIteration - if chunk is not None and chunk != b"": - processed_chunk = self.chunk_creator(chunk=chunk) - if processed_chunk is None: - continue - - choice = processed_chunk.choices[0] - if isinstance(choice, StreamingChoices): - self.response_uptil_now += choice.delta.get("content", "") or "" - else: - self.response_uptil_now += "" - self.rules.post_call_rules(input=self.response_uptil_now, model=self.model) - # RETURN RESULT - self.chunks.append(processed_chunk) - return processed_chunk - except (StopAsyncIteration, StopIteration): - return await self._finalize_completed_stream(cache_hit=cache_hit) - except httpx.TimeoutException as e: # if httpx read timeout error occues - traceback_exception = traceback.format_exc() - ## ADD DEBUG INFORMATION - E.G. LITELLM REQUEST TIMEOUT - traceback_exception += f"\nLiteLLM Default Request Timeout - {litellm.request_timeout}" - if self.logging_obj is not None: - self._record_partial_usage_for_failure() - ## LOGGING - asyncio.create_task( - self.logging_obj.dispatch_failure_handlers(e, traceback_exception, prefer_async_handlers=True) - ) - self._handle_stream_fallback_error(e) - except (httpx.ReadError, httpx.RemoteProtocolError) as e: - if self.received_finish_reason is None: - self._log_stream_failure_and_raise(e) - return await self._finalize_completed_stream(cache_hit=cache_hit) - except Exception as e: - self._log_stream_failure_and_raise(e) - - async def _finalize_completed_stream(self, cache_hit: bool) -> "ModelResponseStream": - if self.sent_last_chunk is True: - # log the final chunk with accurate streaming values - try: - complete_streaming_response = litellm.stream_chunk_builder( - chunks=self.chunks, - messages=self.messages, - logging_obj=self.logging_obj, - ) - except Exception as e: - # see sync __next__: a raise from stream_chunk_builder inside this - # except handler escapes __anext__ and drops the request from SpendLogs. - # Recover best-effort usage from the raw chunks so cost is still tracked - verbose_logger.warning( - "stream_chunk_builder raised at end-of-stream (%s); logging best-effort usage from chunks.", - str(e), - ) - try: - complete_streaming_response = self.model_response_creator( - chunk={"usage": calculate_total_usage(chunks=self.chunks)} - ) - except Exception: - complete_streaming_response = None - - response: Final = self.model_response_creator() - if complete_streaming_response is not None: - self._propagate_usage_cost_to_hidden_params(complete_streaming_response, self.custom_llm_provider) - - setattr( - response, - "usage", - getattr(complete_streaming_response, "usage"), - ) - try: - _copy = complete_streaming_response.model_copy(deep=True) - except RuntimeError: - _copy = complete_streaming_response.model_copy() - asyncio.create_task( - self.async_cache_streaming_response( - processed_chunk=_copy, - cache_hit=cache_hit, - ) - ) - # Update hidden_params with final usage from - # stream_chunk_builder (see sync __next__ for full comment). - if ( - self.stream_options is None - and complete_streaming_response is not None - and self._last_returned_hidden_params is not None - ): - final_usage: Final = getattr(complete_streaming_response, "usage", None) - if final_usage is not None: - self._last_returned_hidden_params["usage"] = final_usage - - if self.sent_stream_usage is False and self.send_stream_usage is True: - self.sent_stream_usage = True - return response - - _deferred_cb: Final = getattr( - self.logging_obj, - "_on_deferred_stream_complete", - None, - ) - if _deferred_cb is not None: - # Proxy has post-call guardrails. Store the assembled - # response so the outer streaming consumer - # (ProxyLogging.async_post_call_streaming_iterator_hook) - # can fire the deferred callback AFTER all guardrail - # end-of-stream blocks complete. Scheduling here via - # create_task would race with unified_guardrail's - # end-of-stream block for short-stream providers. - self.logging_obj._deferred_stream_complete_args = ( - complete_streaming_response, - cache_hit, - ) - else: - # prefer_async_handlers routes CustomLogger to async_success_handler - # when consumers use ``async for`` on sync-SDK streams. Legacy string - # callbacks still run via executor.submit inside dispatch_success_handlers. - asyncio.create_task( - self.logging_obj.dispatch_success_handlers( - complete_streaming_response, - cache_hit=cache_hit, - start_time=None, - end_time=None, - prefer_async_handlers=True, - ) - ) - - self._restore_consumer_correlation_context() - raise StopAsyncIteration # Re-raise StopIteration - else: - if self._should_require_provider_finish_reason() and not self._has_provider_finish_reason(): - self._raise_incomplete_stream_without_finish_reason() - self.sent_last_chunk = True - processed_chunk: Final = self.finish_reason_handler() - if self.stream_options is None: - usage: Final = calculate_total_usage(chunks=self.chunks) - processed_chunk._hidden_params["usage"] = usage # pyright: ignore[reportPrivateUsage] # sync parity - # see sync __next__'s sibling branch: deliberately do NOT restore - # here - this chunk is still this call's own data, and restoring - # before returning it would corrupt the caller's own log - # statements processing it. A caller that keeps iterating gets - # cleaned up on the next __anext__() call; one that stops here - # relies on aclose() or the best-effort __del__ guard. - return processed_chunk - - def _log_stream_failure_and_raise(self, e: Exception) -> NoReturn: - traceback_exception: Final = traceback.format_exc() - if self.logging_obj is not None: - self._record_partial_usage_for_failure() - ## LOGGING - asyncio.create_task( - self.logging_obj.dispatch_failure_handlers(e, traceback_exception, prefer_async_handlers=True) - ) - self._handle_stream_fallback_error(e) - - def _record_partial_usage_for_failure(self) -> None: - """ - A stream that breaks mid-flight still billed the provider for the chunks - already delivered. Recover that partial usage from the chunks seen so - far and stash it, with its cost, on the logging object so the failure - handler records the real partial spend instead of zero. A request that - later recovers via a router fallback overwrites this with the combined - success log on the same request id, so this never double counts. - """ - if self.logging_obj is None or not self.chunks: - return - try: - partial_response: Final = litellm.stream_chunk_builder( - chunks=self.chunks, - messages=self.messages if isinstance(self.messages, list) else None, - logging_obj=self.logging_obj, - ) - if partial_response is None: - return - usage: Final = cast(Usage | None, getattr(partial_response, "usage", None)) - if usage is None: - return - if self.model: - partial_response.model = self.model - backfill_missing_cache_usage_fields(usage) - self.logging_obj.model_call_details["combined_usage_object"] = usage - self.logging_obj.model_call_details["response_cost"] = ( - self.logging_obj._response_cost_calculator(result=partial_response) or 0.0 - ) - except Exception as recover_error: - verbose_logger.debug( - "could not recover partial usage for interrupted stream: %s", - recover_error, - ) - - def _handle_stream_fallback_error(self, e: Exception) -> "NoReturn": - """ - Common error handling for both __next__ and __anext__. - - Maps the raw exception to an OpenAI-compatible type, then decides - whether to raise it directly (non-retriable 4xx) or wrap it in - MidStreamFallbackError so the Router can trigger a fallback. - - 429 (rate-limit) is explicitly exempted from the 4xx filter because - it is transient and the Router should switch to another model group. - """ - from litellm.exceptions import MidStreamFallbackError - - # Map to OpenAI exception format. Some providers' mappers (e.g. - # _map_anthropic_exception, _map_aleph_alpha_exception) synchronously - # log a debug diagnostic (the raw status code) as part of mapping - - # restore the consumer's outer context only after this completes, so - # that diagnostic log line still carries the failing stream's own - # trace_id/session_id instead of the consumer's (or an empty one). - if isinstance(e, OpenAIError): - mapped_exception: Exception = e - else: - try: - mapped_exception = exception_type( - model=self.model, - custom_llm_provider=self.custom_llm_provider, - original_exception=e, - completion_kwargs={}, - extra_kwargs={}, - ) - except Exception as mapping_error: - mapped_exception = mapping_error - self._restore_consumer_correlation_context() - - def _normalize_status_code(exc: Exception) -> int | None: - """Best-effort status_code extraction.""" - try: - code: Final[int | str | None] = getattr(exc, "status_code", None) - if code is not None: - return int(code) - except Exception: - pass - - response: Final[object | None] = getattr(exc, "response", None) - if response is not None: - try: - status_code: Final[int | str | None] = getattr(response, "status_code", None) - if status_code is not None: - return int(status_code) - except Exception: - pass - return None - - mapped_status_code: Final = _normalize_status_code(mapped_exception) - original_status_code: Final = _normalize_status_code(e) - - # Raise non-retriable client errors directly (skip fallback). - # Exception: 429 (rate-limit) IS retriable/transient — allow it - # through so the Router can switch to a different model group. - if mapped_status_code is not None and 400 <= mapped_status_code < 500 and mapped_status_code != 429: - raise mapped_exception - if original_status_code is not None and 400 <= original_status_code < 500 and original_status_code != 429: - raise mapped_exception - - raise MidStreamFallbackError( - message=str(mapped_exception), - model=self.model, - llm_provider=self.custom_llm_provider or "anthropic", - original_exception=mapped_exception, - generated_content=self.response_uptil_now, - is_pre_first_chunk=not self.sent_first_chunk, - ) - - @staticmethod - def _strip_sse_data_from_chunk(chunk: str | None) -> str | None: - """ - Strips the 'data: ' prefix from Server-Sent Events (SSE) chunks. - - Some providers like sagemaker send it as `data:`, need to handle both - - SSE messages are prefixed with 'data: ' which is part of the protocol, - not the actual content from the LLM. This method removes that prefix - and returns the actual content. - - Args: - chunk: The SSE chunk that may contain the 'data: ' prefix (string or bytes) - - Returns: - The chunk with the 'data: ' prefix removed, or the original chunk - if no prefix was found. Returns None if input is None. - - See OpenAI Python Ref for this: https://github.com/openai/openai-python/blob/041bf5a8ec54da19aad0169671793c2078bd6173/openai/api_requestor.py#L100 - """ - if chunk is None: - return None - - if isinstance(chunk, str): - # OpenAI sends `data: ` - if chunk.startswith("data: "): - # Strip the prefix and any leading whitespace that might follow it - _length_of_sse_data_prefix = len("data: ") - return chunk[_length_of_sse_data_prefix:] - elif chunk.startswith("data:"): - # Sagemaker sends `data:`, no trailing whitespace - _length_of_sse_data_prefix = len("data:") - return chunk[_length_of_sse_data_prefix:] - - return chunk - - -def _cache_token_count(details: PromptTokensDetailsWrapper | None, keys: tuple[str, ...]) -> int: - for key in keys: - value = getattr(details, key, None) - if isinstance(value, int) and not isinstance(value, bool) and value: - return value - return 0 - - -def backfill_missing_cache_usage_fields(usage: Usage) -> None: - """Give partial-stream usage the same cache fields a complete stream reports. - - Carries OpenAI-style ``prompt_tokens_details`` counts up to the Anthropic-style - top-level keys, defaulting to zero. It must carry the real count rather than a - flat zero: downstream readers treat these keys as authoritative once present and - skip their own normalization, so a zero here would overwrite a real cache read. - """ - details: Final = usage.prompt_tokens_details - if getattr(usage, "cache_read_input_tokens", None) is None: - usage.cache_read_input_tokens = _cache_token_count( # rebind-ok: in-place backfill is the contract - details, ("cached_tokens",) - ) - if getattr(usage, "cache_creation_input_tokens", None) is None: - usage.cache_creation_input_tokens = _cache_token_count( # rebind-ok: in-place backfill is the contract - details, ("cache_write_tokens", "cache_creation_tokens") - ) - if usage.prompt_tokens_details is None: - usage.prompt_tokens_details = PromptTokensDetailsWrapper(cached_tokens=0) # rebind-ok: backfill in place - - -_TokenDetails = TypeVar("_TokenDetails", PromptTokensDetailsWrapper, CompletionTokensDetailsWrapper) - - -def _coerce_token_details( - usage: dict | BaseModel, field: str, details_type: type[_TokenDetails] -) -> _TokenDetails | None: - raw = usage.get(field) if isinstance(usage, dict) else getattr(usage, field, None) - if raw is None: - return None - if isinstance(raw, details_type): - return raw - return details_type(**(raw if isinstance(raw, dict) else raw.model_dump())) - - -def calculate_total_usage(chunks: list[ModelResponse]) -> Usage: - """Assume most recent usage chunk has total usage uptil then.""" - from litellm.litellm_core_utils.streaming_chunk_builder_utils import ( - attach_cache_creation_token_details, - capture_cache_creation_token_details, - ) - - prompt_tokens: int = 0 - completion_tokens: int = 0 - latest_usage_chunk: Usage | Mapping[str, int] | None = None - prompt_tokens_details: PromptTokensDetailsWrapper | None = None - completion_tokens_details: CompletionTokensDetailsWrapper | None = None - cache_creation_token_details: CacheCreationTokenDetails | None = None - - for chunk in chunks: - if "usage" in chunk and chunk["usage"] is not None: - usage = chunk["usage"] - latest_usage_chunk = usage - if "prompt_tokens" in usage: - prompt_tokens = usage.get("prompt_tokens", 0) or 0 - if "completion_tokens" in usage: - completion_tokens = usage.get("completion_tokens", 0) or 0 - incoming_prompt_tokens_details = _coerce_token_details( - usage, "prompt_tokens_details", PromptTokensDetailsWrapper - ) - cache_creation_token_details = capture_cache_creation_token_details( - incoming_prompt_tokens_details, cache_creation_token_details - ) - prompt_tokens_details = incoming_prompt_tokens_details or prompt_tokens_details - completion_tokens_details = ( - _coerce_token_details(usage, "completion_tokens_details", CompletionTokensDetailsWrapper) - or completion_tokens_details - ) - - returned_usage_chunk: Final = Usage( - prompt_tokens=prompt_tokens, - completion_tokens=completion_tokens, - total_tokens=prompt_tokens + completion_tokens, - prompt_tokens_details=attach_cache_creation_token_details(prompt_tokens_details, cache_creation_token_details), - completion_tokens_details=completion_tokens_details, - ) - - if latest_usage_chunk is not None: - latest_cost: Final = ( - latest_usage_chunk.get("cost") - if isinstance(latest_usage_chunk, dict) - else getattr(latest_usage_chunk, "cost", None) - ) - if latest_cost is not None: - returned_usage_chunk.cost = latest_cost - - return returned_usage_chunk - - -def generic_chunk_has_all_required_fields(chunk: dict) -> bool: - """ - Checks if the provided chunk dictionary contains all required fields for GenericStreamingChunk. - - :param chunk: The dictionary to check. - :return: True if all required fields are present, False otherwise. - """ - return all(key in _GCHUNK_FIELDS for key in chunk) - - -def convert_generic_chunk_to_model_response_stream( - chunk: GChunk, -) -> ModelResponseStream: - from litellm.types.utils import Delta - - model_response_stream: Final = ModelResponseStream( - id=str(uuid.uuid4()), - model="", - choices=[ - StreamingChoices( - index=chunk.get("index", 0), - delta=Delta( - content=chunk["text"], - tool_calls=chunk.get("tool_use", None), - ), - ) - ], - finish_reason=chunk["finish_reason"] if chunk["is_finished"] else None, - ) - - if "usage" in chunk and chunk["usage"] is not None: - setattr(model_response_stream, "usage", chunk["usage"]) - - return model_response_stream +@file:/tmp/litellm-work/litellm/litellm/litellm_core_utils/streaming_handler.py \ No newline at end of file diff --git a/litellm/utils.py b/litellm/utils.py index e3497ad71a6..905c2cfb6bb 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -1,9884 +1 @@ -"""Utility helpers for LiteLLM core request handling and provider support.""" - -# from __future__ import annotations must be the first non-comment statement -from __future__ import annotations - -import ast -import asyncio -import base64 -import binascii -import contextvars -import copy -import datetime -import hashlib -import inspect -import io -import itertools -import json -import logging -import os -import random -import re -import struct -import subprocess - -# What is this? -## Generic utils.py file. Problem-specific utils (e.g. 'cost calculation), should all be in `litellm_core_utils/`. -import sys -import textwrap -import threading -import time -import traceback -from dataclasses import dataclass, field -from functools import lru_cache, wraps -from importlib import resources -from inspect import iscoroutine -from io import StringIO -from os.path import abspath, dirname, join -from types import MappingProxyType - -import dotenv -import httpx -import openai -import tiktoken -from httpx import Proxy -from httpx._utils import get_environment_proxies -from openai.lib import _parsing, _pydantic -from openai.types.chat.completion_create_params import ResponseFormat -from pydantic import BaseModel -from tiktoken import Encoding -from tokenizers import Tokenizer - -import litellm -import litellm.litellm_core_utils - -# audio_utils.utils is lazy-loaded - only imported when needed for transcription calls -import litellm.litellm_core_utils.json_validation_rule -from litellm._internal_context import is_internal_call -from litellm._lazy_imports import ( - _get_default_encoding, - _get_modified_max_tokens, - _get_token_counter_new, -) -from litellm._uuid import uuid -from litellm.constants import ( - DEFAULT_CHAT_COMPLETION_PARAM_VALUES, - DEFAULT_EMBEDDING_PARAM_VALUES, - DEFAULT_MAX_LRU_CACHE_SIZE, - DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT, - DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET, - DEFAULT_TRIM_RATIO, - FUNCTION_DEFINITION_TOKEN_COUNT, - HF_CONFIG_FETCH_TIMEOUT_SECONDS, - INITIAL_RETRY_DELAY, - JITTER, - MAX_RETRY_DELAY, - MAX_TOKEN_TRIMMING_ATTEMPTS, - MINIMUM_PROMPT_CACHE_TOKEN_COUNT_OVERRIDE, - NON_INFERENCE_CALL_TYPES, - OPENAI_EMBEDDING_PARAMS, - PROVIDERS_THAT_AUTHENTICATE_ON_PROVIDER_INFO, - TOOL_CHOICE_OBJECT_TOKEN_COUNT, -) -from litellm.litellm_core_utils.fallback_generalizations import ( - match_capability_generalizations, -) -from litellm.litellm_core_utils.sensitive_data_masker import redact_credentials_in_payload - -_CachingHandlerResponse = None -_LLMCachingHandler = None -_CustomGuardrail = None -_CustomLogger = None - - -def _get_cached_custom_logger(): - """ - Get cached CustomLogger class. - Lazy imports on first call to avoid loading custom_logger at import time. - Subsequent calls use cached class for better performance. - """ - global _CustomLogger - if _CustomLogger is None: - from litellm.integrations.custom_logger import CustomLogger - - _CustomLogger = CustomLogger - return _CustomLogger - - -@lru_cache(maxsize=None) -def _accepts_fallback_depth_kwarg_for_class(cls: type) -> bool: - """ - Whether cls's async_post_call_failure_deployment_hook override accepts a - fallback_depth keyword, cached per class so a signature the base class added after a - subscriber's override was written (e.g. the PR's own earlier 3-arg proof-of-fix - example) doesn't raise TypeError - swallowed at debug level - on every call. - """ - params: Final = inspect.signature(cls.async_post_call_failure_deployment_hook).parameters - return "fallback_depth" in params or any(p.kind is inspect.Parameter.VAR_KEYWORD for p in params.values()) - - -def _snapshot_exception_for_hook(exception: Exception) -> Exception: - """ - A same-class copy of exception that skips __init__ (many litellm exceptions require - constructor args beyond what .args carries, so copy.copy's pickle-based reconstruction - fails on them). Handed to failure-hook callbacks instead of the live object so a - callback setting e.g. exception.status_code cannot change the status code the real - caller actually receives. Falls back to the live object if snapshotting fails for a - type this doesn't anticipate, since the real exception must still reach the callback. - """ - try: - cls: Final = type(exception) - snapshot: Final = cls.__new__(cls) - snapshot.__dict__.update(exception.__dict__) - snapshot.args = exception.args - snapshot.__traceback__ = exception.__traceback__ - snapshot.__cause__ = exception.__cause__ - snapshot.__context__ = exception.__context__ - # Setting __cause__ implicitly forces __suppress_context__ to True (CPython - # behavior for `raise ... from ...`), so this must be set after, not before. - snapshot.__suppress_context__ = exception.__suppress_context__ - return snapshot - except Exception: # noqa: BLE001 # any snapshot failure must fall back to the live object, not break the failure path - return exception - - -def _get_cached_custom_guardrail(): - """ - Get cached CustomGuardrail class. - Lazy imports on first call to avoid loading custom_guardrail at import time. - Subsequent calls use cached class for better performance. - """ - global _CustomGuardrail - if _CustomGuardrail is None: - from litellm.integrations.custom_guardrail import CustomGuardrail - - _CustomGuardrail = CustomGuardrail - return _CustomGuardrail - - -def _get_cached_caching_handler_response(): - """ - Get cached CachingHandlerResponse class. - Lazy imports on first call to avoid loading caching_handler at import time. - Subsequent calls use cached class for better performance. - """ - global _CachingHandlerResponse - if _CachingHandlerResponse is None: - from litellm.caching.caching_handler import CachingHandlerResponse - - _CachingHandlerResponse = CachingHandlerResponse - return _CachingHandlerResponse - - -def _get_cached_llm_caching_handler(): - """ - Get cached LLMCachingHandler class. - Lazy imports on first call to avoid loading caching_handler at import time. - Subsequent calls use cached class for better performance. - """ - global _LLMCachingHandler - if _LLMCachingHandler is None: - from litellm.caching.caching_handler import LLMCachingHandler - - _LLMCachingHandler = LLMCachingHandler - return _LLMCachingHandler - - -# Cached lazy import for audio_utils.utils -# Module-level cache to avoid repeated imports while preserving memory benefits -_audio_utils_module = None - - -def _get_cached_audio_utils(): - """ - Get cached audio_utils.utils module. - Lazy imports on first call to avoid loading audio_utils.utils at import time. - Subsequent calls use cached module for better performance. - """ - global _audio_utils_module - if _audio_utils_module is None: - import litellm.litellm_core_utils.audio_utils.utils - - _audio_utils_module = litellm.litellm_core_utils.audio_utils.utils - return _audio_utils_module - - -from litellm.types.llms.openai import ( - AllMessageValues, - AllPromptValues, - ChatCompletionAssistantToolCall, - ChatCompletionNamedToolChoiceParam, - ChatCompletionToolParam, - ChatCompletionToolParamFunctionChunk, - OpenAITextCompletionUserMessage, - OpenAIWebSearchOptions, -) -from litellm.types.utils import ( - OPENAI_RESPONSE_HEADERS, - CallTypes, - ChatCompletionDeltaToolCall, - ChatCompletionMessageToolCall, - Choices, - CostPerToken, - CredentialItem, - CustomHuggingfaceTokenizer, - Delta, - Embedding, - EmbeddingResponse, - FileTypes, - Function, - ImageResponse, - LlmProviders, - LlmProvidersSet, - LLMResponseTypes, - Message, - ModelInfo, - ModelInfoBase, - ModelResponse, - ModelResponseStream, - ProviderField, - ProviderSpecificModelInfo, - RawRequestTypedDict, - SandboxProviders, - SearchProviders, - SelectTokenizerResponse, - StreamingChoices, - TextChoices, - TextCompletionResponse, - TranscriptionResponse, - Usage, - all_litellm_params, -) - -_CALL_TYPE_ENUM_MAP: Final[dict] = {ct.value: ct for ct in CallTypes} - -# +-----------------------------------------------+ -# | | -# | Give Feedback / Get Help | -# | https://github.com/BerriAI/litellm/issues/new | -# | | -# +-----------------------------------------------+ -# -# Thank you users! We ❤️ you! - Krrish & Ishaan - - -try: - # Python 3.9+ - with ( - resources.files("litellm.litellm_core_utils.tokenizers") - .joinpath("anthropic_tokenizer.json") - .open("r", encoding="utf-8") as f - ): - json_data = json.load(f) -except (ImportError, AttributeError, TypeError): - with resources.open_text("litellm.litellm_core_utils.tokenizers", "anthropic_tokenizer.json") as f: - json_data = json.load(f) - -# Convert to str (if necessary) -claude_json_str = json.dumps(json_data) -import importlib.metadata -from collections.abc import Callable, Iterable, Mapping, Sequence -from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast, get_args - -from litellm import utils as litellm_utils - -# These are lazy loaded via __getattr__ -from litellm.llms.base_llm.base_utils import ( - BaseLLMModelInfo, - type_to_response_format_param, -) - -if TYPE_CHECKING: - # Heavy types that are only needed for type checking; avoid importing - # their modules at runtime during `litellm` import. - from litellm.caching.caching_handler import ( - CachingHandlerResponse, - LLMCachingHandler, - ) - from litellm.integrations.custom_logger import CustomLogger - - # Type stubs for lazy-loaded functions and classes - from litellm.litellm_core_utils.cached_imports import ( - get_coroutine_checker, - get_litellm_logging_class, - get_set_callbacks, - ) - from litellm.litellm_core_utils.core_helpers import ( - get_litellm_metadata_from_kwargs, - map_finish_reason, - process_response_headers, - ) - from litellm.litellm_core_utils.credential_accessor import CredentialAccessor - from litellm.litellm_core_utils.dot_notation_indexing import ( - delete_nested_value, - is_nested_path, - ) - - # Type stubs for lazy-loaded functions to help mypy understand their types - # These imports allow mypy to understand the types when these are accessed via __getattr__ - from litellm.litellm_core_utils.exception_mapping_utils import exception_type - from litellm.litellm_core_utils.get_litellm_params import ( - _get_base_model_from_litellm_call_metadata, - get_litellm_params, - ) - from litellm.litellm_core_utils.get_llm_provider_logic import ( - _is_non_openai_azure_model, - get_llm_provider, - ) - from litellm.litellm_core_utils.get_supported_openai_params import ( - get_supported_openai_params, - ) - from litellm.litellm_core_utils.llm_request_utils import _ensure_extra_body_is_safe - from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - LiteLLMResponseObjectHandler, - _handle_invalid_parallel_tool_calls, - convert_to_model_response_object, - convert_to_streaming_response, - convert_to_streaming_response_async, - ) - from litellm.litellm_core_utils.llm_response_utils.get_api_base import get_api_base - from litellm.litellm_core_utils.llm_response_utils.get_formatted_prompt import ( - get_formatted_prompt, - ) - from litellm.litellm_core_utils.llm_response_utils.get_headers import ( - get_response_headers, - ) - from litellm.litellm_core_utils.llm_response_utils.response_metadata import ( - ResponseMetadata, - update_response_metadata, - ) - from litellm.litellm_core_utils.prompt_templates.common_utils import ( - _parse_content_for_reasoning, - ) - from litellm.litellm_core_utils.redact_messages import ( - LiteLLMLoggingObject, - redact_message_input_output_from_logging, - ) - from litellm.litellm_core_utils.rules import Rules - from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper - from litellm.litellm_core_utils.thread_pool_executor import executor - from litellm.llms.base_llm.anthropic_messages.transformation import ( - BaseAnthropicMessagesConfig, - ) - from litellm.llms.base_llm.audio_transcription.transformation import ( - BaseAudioTranscriptionConfig, - ) - - # Type stubs for lazy-loaded config classes and types - from litellm.llms.base_llm.batches.transformation import BaseBatchesConfig - from litellm.llms.base_llm.containers.transformation import BaseContainerConfig - from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig - from litellm.llms.base_llm.files.transformation import BaseFilesConfig - from litellm.llms.base_llm.google_genai.transformation import ( - BaseGoogleGenAIGenerateContentConfig, - ) - from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig - from litellm.llms.base_llm.image_generation.transformation import ( - BaseImageGenerationConfig, - ) - from litellm.llms.base_llm.image_variations.transformation import ( - BaseImageVariationConfig, - ) - from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig - from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig - from litellm.llms.base_llm.realtime.http_transformation import ( - BaseRealtimeHTTPConfig, - ) - from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig - from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig - from litellm.llms.base_llm.sandbox.transformation import BaseSandboxConfig - from litellm.llms.base_llm.search.transformation import BaseSearchConfig - from litellm.llms.base_llm.text_to_speech.transformation import ( - BaseTextToSpeechConfig, - ) - from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig - from litellm.llms.base_llm.vector_store_files.transformation import ( - BaseVectorStoreFilesConfig, - ) - from litellm.llms.base_llm.videos.transformation import BaseVideoConfig - from litellm.llms.bedrock.common_utils import BedrockModelInfo - from litellm.llms.bedrock.embed.amazon_nova_transformation import ( - AmazonNovaEmbeddingConfig, - ) - from litellm.llms.bedrock.embed.amazon_titan_g1_transformation import ( - AmazonTitanG1Config, - ) - from litellm.llms.bedrock.embed.amazon_titan_multimodal_transformation import ( - AmazonTitanMultimodalEmbeddingG1Config, - ) - from litellm.llms.bedrock.embed.amazon_titan_v2_transformation import ( - AmazonTitanV2Config, - ) - from litellm.llms.bedrock.embed.cohere_transformation import ( - BedrockCohereEmbeddingConfig, - ) - from litellm.llms.bedrock.embed.twelvelabs_marengo_transformation import ( - TwelveLabsMarengoEmbeddingConfig, - ) - from litellm.llms.cohere.common_utils import CohereModelInfo - from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler - from litellm.llms.mistral.ocr.transformation import MistralOCRConfig - from litellm.proxy._types import AllowedModelRegion - from litellm.router_utils.get_retry_from_policy import ( - get_num_retries_from_retry_policy, - reset_retry_policy, - ) - from litellm.types.llms.anthropic import ( - ANTHROPIC_API_ONLY_HEADERS, - AnthropicThinkingParam, - ) - from litellm.types.llms.openai import ( - ChatCompletionDeltaToolCallChunk, - ChatCompletionToolCallChunk, - ChatCompletionToolCallFunctionChunk, - ) - from litellm.types.rerank import RerankResponse - from litellm.types.router import LiteLLM_Params - -from litellm.llms.base_llm.chat.transformation import BaseConfig -from litellm.llms.base_llm.completion.transformation import BaseTextCompletionConfig -from litellm.llms.base_llm.evals.transformation import BaseEvalsAPIConfig -from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig -from litellm.llms.base_llm.skills.transformation import BaseSkillsAPIConfig -from litellm.secret_managers.main import get_secret - -from ._logging import _is_debugging_on, verbose_logger -from .caching.caching import ( - AzureBlobCache, - Cache, - QdrantSemanticCache, - RedisCache, - RedisSemanticCache, - S3Cache, -) -from .exceptions import ( - APIConnectionError, - APIError, - AuthenticationError, - BadRequestError, - BudgetExceededError, - ContentPolicyViolationError, - ContextWindowExceededError, - NotFoundError, - OpenAIError, - PermissionDeniedError, - RateLimitError, - ServiceUnavailableError, - Timeout, - UnprocessableEntityError, - UnsupportedParamsError, -) - -if TYPE_CHECKING: - from litellm import MockException - -####### ENVIRONMENT VARIABLES #################### -# Adjust to your specific application needs / system capabilities. -sentry_sdk_instance: Final = None -capture_exception: Final = None -add_breadcrumb = None -posthog: Final = None -slack_app: Final = None -alerts_channel: Final = None -heliconeLogger: Final = None -athinaLogger: Final = None -promptLayerLogger: Final = None -langsmithLogger: Final = None -logfireLogger: Final = None -weightsBiasesLogger: Final = None -customLogger: Final = None -langFuseLogger: Final = None -openMeterLogger: Final = None -lagoLogger: Final = None -dataDogLogger: Final = None -prometheusLogger: Final = None -dynamoLogger: Final = None -s3Logger: Final = None -greenscaleLogger: Final = None -lunaryLogger: Final = None -aispendLogger: Final = None -supabaseClient: Final = None -callback_list: list[str] | None = [] -user_logger_fn = None -additional_details: Final[dict[str, str] | None] = {} -local_cache: Final[dict[str, str] | None] = {} -last_fetched_at: Final = None -last_fetched_at_keys: Final = None -######## Model Response ######################### - -# All liteLLM Model responses will be in this format, Follows the OpenAI Format -# https://docs.litellm.ai/docs/completion/output -# { -# 'choices': [ -# { -# 'finish_reason': 'stop', -# 'index': 0, -# 'message': { -# 'role': 'assistant', -# 'content': " I'm doing well, thank you for asking. I am Claude, an AI assistant created by Anthropic." -# } -# } -# ], -# 'created': 1691429984.3852863, -# 'model': 'claude-instant-1', -# 'usage': {'prompt_tokens': 18, 'completion_tokens': 23, 'total_tokens': 41} -# } - - -############################################################ -def print_verbose( - print_statement, - logger_only: bool = False, - log_level: Literal["DEBUG", "INFO", "ERROR"] = "DEBUG", -): - try: - if log_level == "DEBUG": - verbose_logger.debug(print_statement) - elif log_level == "INFO": - verbose_logger.info(print_statement) - elif log_level == "ERROR": - verbose_logger.error(print_statement) - if litellm.set_verbose is True and logger_only is False: - print(print_statement) # noqa: T201 - except Exception: - pass - - -def _print_verbose_is_active() -> bool: - """Whether print_verbose would reach either of its two consumers, so a call site can skip - building a payload nothing would read. _is_debugging_on() is not the same predicate: it reads - litellm._logging.set_verbose, while print_verbose's print reads litellm.set_verbose, and - assigning the documented litellm.set_verbose = True rebinds only the latter.""" - return litellm.set_verbose is True or verbose_logger.isEnabledFor(logging.DEBUG) - - -####### CLIENT ################### -# make it easy to log if completion/embedding runs succeeded or failed + see what happened | Non-Blocking -def custom_llm_setup(): - """ - Add custom_llm provider to provider list - """ - for custom_llm in litellm.custom_provider_map: - if custom_llm["provider"] not in litellm.provider_list: - litellm.provider_list.append(custom_llm["provider"]) - - if custom_llm["provider"] not in litellm._custom_providers: - litellm._custom_providers.append(custom_llm["provider"]) - - -def _add_custom_logger_callback_to_specific_event(callback: str, logging_event: Literal["success", "failure"]) -> None: - """ - Add a custom logger callback to the specific event - """ - from litellm import _custom_logger_compatible_callbacks_literal - from litellm.litellm_core_utils.litellm_logging import ( - _init_custom_logger_compatible_class, - ) - - if callback not in litellm._known_custom_logger_compatible_callbacks: - verbose_logger.debug( - "Callback %s is not a valid custom logger compatible callback. Known list - %s", - callback, - litellm._known_custom_logger_compatible_callbacks, - ) - return - - callback_class: Final = _init_custom_logger_compatible_class( - cast(_custom_logger_compatible_callbacks_literal, callback), - internal_usage_cache=None, - llm_router=None, - ) - - if callback_class: - if logging_event == "success" and _custom_logger_class_exists_in_success_callbacks(callback_class) is False: - litellm.logging_callback_manager.add_litellm_success_callback(callback_class) - litellm.logging_callback_manager.add_litellm_async_success_callback(callback_class) - if callback in litellm.success_callback: - litellm.success_callback.remove(callback) # remove the string from the callback list - if callback in litellm._async_success_callback: - litellm._async_success_callback.remove(callback) # remove the string from the callback list - elif logging_event == "failure" and _custom_logger_class_exists_in_failure_callbacks(callback_class) is False: - litellm.logging_callback_manager.add_litellm_failure_callback(callback_class) - litellm.logging_callback_manager.add_litellm_async_failure_callback(callback_class) - if callback in litellm.failure_callback: - litellm.failure_callback.remove(callback) # remove the string from the callback list - if callback in litellm._async_failure_callback: - litellm._async_failure_callback.remove(callback) # remove the string from the callback list - - -def _custom_logger_class_exists_in_success_callbacks( - callback_class: CustomLogger, -) -> bool: - """ - Returns True if an instance of the custom logger exists in litellm.success_callback or litellm._async_success_callback - - e.g if `LangfusePromptManagement` is passed in, it will return True if an instance of `LangfusePromptManagement` exists in litellm.success_callback or litellm._async_success_callback - - Prevents double adding a custom logger callback to the litellm callbacks - - Matches on the exact class; an instance of a subclass does not count as registered - """ - return any(type(cb) is type(callback_class) for cb in litellm.success_callback + litellm._async_success_callback) - - -def _custom_logger_class_exists_in_failure_callbacks( - callback_class: CustomLogger, -) -> bool: - """ - Returns True if an instance of the custom logger exists in litellm.failure_callback or litellm._async_failure_callback - - e.g if `LangfusePromptManagement` is passed in, it will return True if an instance of `LangfusePromptManagement` exists in litellm.failure_callback or litellm._async_failure_callback - - Prevents double adding a custom logger callback to the litellm callbacks - - Matches on the exact class; an instance of a subclass does not count as registered - """ - return any(type(cb) is type(callback_class) for cb in litellm.failure_callback + litellm._async_failure_callback) - - -def get_request_guardrails(kwargs: dict[str, Any]) -> list[str]: - """ - Get the request guardrails from the kwargs - """ - metadata: Final = kwargs.get("metadata") or {} - requester_metadata: Final = metadata.get("requester_metadata") or {} - applied_guardrails: Final = requester_metadata.get("guardrails") or [] - return applied_guardrails - - -def get_applied_guardrails(kwargs: dict[str, object]) -> list[str]: - """ - - Add 'default_on' guardrails to the list - - Add request guardrails to the list - """ - - request_guardrails: Final = get_request_guardrails(kwargs) - applied_guardrails: Final = [] - CustomGuardrail: Final = _get_cached_custom_guardrail() - for callback in litellm.callbacks: - if callback is not None and isinstance(callback, CustomGuardrail): - if callback.guardrail_name is not None: - if callback.default_on is True or callback.guardrail_name in request_guardrails: - applied_guardrails.append(callback.guardrail_name) - - return applied_guardrails - - -def load_credentials_from_list(kwargs: dict): - """ - Updates kwargs with the credentials if credential_name in kwarg - """ - # Access CredentialAccessor via module to trigger lazy loading if needed - CredentialAccessor: Final = getattr(sys.modules[__name__], "CredentialAccessor") - - credential_name: Final = kwargs.get("litellm_credential_name") - if credential_name and litellm.credential_list: - credential_accessor: Final[Mapping[str, object]] = CredentialAccessor.get_credential_values(credential_name) - for key, value in credential_accessor.items(): - if key not in kwargs: - kwargs[key] = value - - -def get_dynamic_callbacks( - dynamic_callbacks: list[str | Callable | CustomLogger] | None, -) -> list: - returned_callbacks: Final = litellm.callbacks.copy() - if dynamic_callbacks: - returned_callbacks.extend(dynamic_callbacks) - return returned_callbacks - - -def _is_gemini_model(model: str | None, custom_llm_provider: str | None) -> bool: - """ - Check if the target model is a Gemini or Vertex AI Gemini model. - """ - if custom_llm_provider in ["gemini", "vertex_ai", "vertex_ai_beta"]: - # For vertex_ai, check if it's actually a Gemini model - if custom_llm_provider in ["vertex_ai", "vertex_ai_beta"]: - return model is not None and "gemini" in model.lower() - return True - - # Check if model name contains gemini - return model is not None and "gemini" in model.lower() - - -def _remove_thought_signature_from_id(tool_call_id: str, separator: str) -> str: - """ - Remove thought signature from a tool call ID. - """ - if separator in tool_call_id: - return tool_call_id.split(separator, 1)[0] - return tool_call_id - - -def _process_assistant_message_tool_calls(msg_copy: dict, thought_signature_separator: str) -> dict: - """ - Process assistant message to remove thought signatures from tool call IDs. - """ - role: Final = msg_copy.get("role") - tool_calls: Final = msg_copy.get("tool_calls") - - if role == "assistant" and isinstance(tool_calls, list): - new_tool_calls: Final = [] - for tc in tool_calls: - # Handle both dict and Pydantic model tool calls - if hasattr(tc, "model_dump"): - # It's a Pydantic model, convert to dict - tc_dict = tc.model_dump() - elif isinstance(tc, dict): - tc_dict = tc.copy() - else: - new_tool_calls.append(tc) - continue - - # Remove thought signature from ID if present - if isinstance(tc_dict.get("id"), str): - if thought_signature_separator in tc_dict["id"]: - tc_dict["id"] = _remove_thought_signature_from_id(tc_dict["id"], thought_signature_separator) - - new_tool_calls.append(tc_dict) - msg_copy["tool_calls"] = new_tool_calls - - return msg_copy - - -def _process_tool_message_id(msg_copy: dict, thought_signature_separator: str) -> dict: - """ - Process tool message to remove thought signature from tool_call_id. - """ - if msg_copy.get("role") == "tool" and isinstance(msg_copy.get("tool_call_id"), str): - if thought_signature_separator in msg_copy["tool_call_id"]: - msg_copy["tool_call_id"] = _remove_thought_signature_from_id( - msg_copy["tool_call_id"], thought_signature_separator - ) - - return msg_copy - - -def _remove_thought_signatures_from_messages(messages: list, thought_signature_separator: str) -> list: - """ - Remove thought signatures from tool call IDs in all messages. - """ - processed_messages: Final = [] - - for msg in messages: - # Handle Pydantic models (convert to dict) - if hasattr(msg, "model_dump"): - msg_dict = msg.model_dump() - elif isinstance(msg, dict): - msg_dict = msg.copy() - else: - # Unknown type, keep as is - processed_messages.append(msg) - continue - - # Process assistant messages with tool_calls - msg_dict = _process_assistant_message_tool_calls(msg_dict, thought_signature_separator) - - # Process tool messages with tool_call_id - msg_dict = _process_tool_message_id(msg_dict, thought_signature_separator) - - processed_messages.append(msg_dict) - - return processed_messages - - -def _restore_correlation_context_if_supported(logging_obj: object) -> None: - """Call logging_obj._restore_correlation_context() if it's actually there. - - Some call sites (tests, narrow unit paths) inject a minimal stand-in - object as litellm_logging_obj instead of a real Logging instance - this - method is new plumbing specific to request_correlation_in_logs, not part - of any pre-existing stand-in's expected interface. `object` (not `Any`) - is deliberate: the getattr() below is exactly how this stays type-safe - while still tolerating a stand-in that lacks the method. - """ - restore: Final = getattr(logging_obj, "_restore_correlation_context", None) - if restore is not None: - restore() - - -def _is_streaming_response_for_correlation(result: object) -> bool: - """True if `result` is a lazy stream wrapper rather than an already-complete response. - - Only wrapper_async() consults this - it must NOT restore the originating - Task's trace_id/session_id as soon as a streaming call returns this: the - caller is about to iterate it over however many subsequent lines of their - own code, and those log lines should still show this call's ids, not the - pre-call ones. This is safe specifically because each async call already - runs in its own asyncio Task with its own copy of the contextvars, so - leaving it "open" can only affect that one Task, never a different, - unrelated future request - Tasks, unlike a thread pool's worker threads, - are never recycled across requests. The corresponding terminal handler - (async_success_handler, dispatched once the full stream is actually - assembled) is what restores it once streaming genuinely finishes. - - wrapper() (the sync path) does NOT consult this at all: sync calls pass - supports_correlation_logging=False into function_setup()/Logging(), so - they never stamp trace_id/session_id in the first place - a plain OS - thread has no per-call isolation the way an asyncio Task does, and a - thread pool's worker threads *are* recycled across unrelated requests, so - stamping ids there without a safe restore mechanism could permanently - misattribute a later, unrelated request's logs. Full sync support is - deferred to a follow-up PR with its own restore mechanism; see - Logging.__init__'s supports_correlation_logging parameter. - - Genuinely circular otherwise: utils.py -> streaming_handler.py -> - redact_messages.py -> llms/vertex_ai/common_utils.py -> utils.py, which - needs names (supports_response_schema, etc.) this module hasn't finished - defining yet at that point in its own top-to-bottom execution. - """ - from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper - - return isinstance(result, CustomStreamWrapper) - - -# Runs once per call to check if the user wants to send their data anywhere - PostHog/Sentry/Slack/etc. -def function_setup( - original_function: str, - rules_obj: Rules, - start_time: datetime.datetime, - *args: Any, # positional passthrough to the wrapped LLM call (ANN401 ignored, see ruff-strict.toml) - is_async_call: bool = True, - **kwargs: Any, # kwargs-ok: forwarded to Logging()/callbacks, varies per call_type -) -> tuple[LiteLLMLoggingObject, dict[str, Any]]: - ### NOTICES ### - if litellm.set_verbose is True: - verbose_logger.warning( - "`litellm.set_verbose` is deprecated. Please set `os.environ['LITELLM_LOG'] = 'DEBUG'` for debug logs." - ) - logging_obj: LiteLLMLoggingObject | None = None # rebind-ok: set to the real object further down on success - try: - global callback_list, add_breadcrumb, user_logger_fn, Logging - - ## CUSTOM LLM SETUP ## - custom_llm_setup() - - ## GET APPLIED GUARDRAILS - applied_guardrails: Final = get_applied_guardrails(kwargs) - - ## LOGGING SETUP - function_id: Final[str | None] = kwargs["id"] if "id" in kwargs else None - - ## LAZY LOAD COROUTINE CHECKER ## - get_coroutine_checker_fn: Final = litellm_utils.get_coroutine_checker - coroutine_checker: Final = get_coroutine_checker_fn() - - ## DYNAMIC CALLBACKS ## - dynamic_callbacks: Final[list[str | Callable | CustomLogger] | None] = kwargs.pop("callbacks", None) - all_callbacks: Final = get_dynamic_callbacks(dynamic_callbacks=dynamic_callbacks) - - if len(all_callbacks) > 0: - for callback in all_callbacks: - # check if callback is a string - e.g. "lago", "openmeter" - if isinstance(callback, str): - callback = litellm.litellm_core_utils.litellm_logging._init_custom_logger_compatible_class( - callback, - internal_usage_cache=None, - llm_router=None, - ) - if callback is None or any( - type(cb) is type(callback) for cb in litellm._async_success_callback - ): # don't double add a callback - continue - if callback not in litellm.input_callback: - litellm.input_callback.append(callback) - if callback not in litellm.success_callback: - litellm.logging_callback_manager.add_litellm_success_callback(callback) - if callback not in litellm.failure_callback: - litellm.logging_callback_manager.add_litellm_failure_callback(callback) - if callback not in litellm._async_success_callback: - litellm.logging_callback_manager.add_litellm_async_success_callback(callback) - if callback not in litellm._async_failure_callback: - litellm.logging_callback_manager.add_litellm_async_failure_callback(callback) - print_verbose(f"Initialized litellm callbacks, Async Success Callbacks: {litellm._async_success_callback}") - - if ( - len(litellm.input_callback) > 0 or len(litellm.success_callback) > 0 or len(litellm.failure_callback) > 0 - ) and len(callback_list) == 0: - callback_list = list(set(litellm.input_callback + litellm.success_callback + litellm.failure_callback)) - get_set_callbacks: Final = getattr(sys.modules[__name__], "get_set_callbacks") - get_set_callbacks()(callback_list=callback_list, function_id=function_id) - ## ASYNC CALLBACKS - safety net for callbacks added via direct append - if len(litellm.input_callback) > 0: - removed_async_items = [] - for index, callback in enumerate(litellm.input_callback): - if coroutine_checker.is_async_callable(callback): - litellm._async_input_callback.append(callback) - removed_async_items.append(index) - - # Pop the async items from input_callback in reverse order to avoid index issues - for index in reversed(removed_async_items): - litellm.input_callback.pop(index) - if len(litellm.success_callback) > 0: - removed_async_items = [] - for index, callback in enumerate(litellm.success_callback): - if coroutine_checker.is_async_callable(callback): - litellm.logging_callback_manager.add_litellm_async_success_callback(callback) - removed_async_items.append(index) - elif callback == "dynamodb" or callback == "openmeter": - # dynamo is an async callback, it's used for the proxy and needs to be async - # we only support async dynamo db logging for acompletion/aembedding since that's used on proxy - litellm.logging_callback_manager.add_litellm_async_success_callback(callback) - removed_async_items.append(index) - elif callback in litellm._known_custom_logger_compatible_callbacks and isinstance(callback, str): - _add_custom_logger_callback_to_specific_event(callback, "success") - - # Pop the async items from success_callback in reverse order to avoid index issues - for index in reversed(removed_async_items): - litellm.success_callback.pop(index) - - if len(litellm.failure_callback) > 0: - removed_async_items = [] - for index, callback in enumerate(litellm.failure_callback): - if coroutine_checker.is_async_callable(callback): - litellm.logging_callback_manager.add_litellm_async_failure_callback(callback) - removed_async_items.append(index) - elif callback in litellm._known_custom_logger_compatible_callbacks and isinstance(callback, str): - _add_custom_logger_callback_to_specific_event(callback, "failure") - - # Pop the async items from failure_callback in reverse order to avoid index issues - for index in reversed(removed_async_items): - litellm.failure_callback.pop(index) - ### DYNAMIC CALLBACKS ### - dynamic_success_callbacks: list[str | Callable | CustomLogger] | None = None - dynamic_async_success_callbacks: list[str | Callable | CustomLogger] | None = None - dynamic_failure_callbacks: list[str | Callable | CustomLogger] | None = None - dynamic_async_failure_callbacks: Final[list[str | Callable | CustomLogger] | None] = None - if kwargs.get("success_callback", None) is not None and isinstance(kwargs["success_callback"], list): - removed_async_items = [] - for index, callback in enumerate(kwargs["success_callback"]): - if coroutine_checker.is_async_callable(callback) or callback == "dynamodb" or callback == "s3": - if dynamic_async_success_callbacks is not None and isinstance( - dynamic_async_success_callbacks, list - ): - dynamic_async_success_callbacks.append(callback) - else: - dynamic_async_success_callbacks = [callback] - removed_async_items.append(index) - # Pop the async items from success_callback in reverse order to avoid index issues - for index in reversed(removed_async_items): - kwargs["success_callback"].pop(index) - dynamic_success_callbacks = kwargs.pop("success_callback") - if kwargs.get("failure_callback", None) is not None and isinstance(kwargs["failure_callback"], list): - dynamic_failure_callbacks = kwargs.pop("failure_callback") - - if add_breadcrumb: - try: - from litellm.litellm_core_utils.core_helpers import safe_deep_copy - - details_to_log = safe_deep_copy(kwargs) - except Exception: - details_to_log = kwargs - - if litellm.turn_off_message_logging: - # make a copy of the _model_Call_details and log it - details_to_log.pop("messages", None) - details_to_log.pop("input", None) - details_to_log.pop("prompt", None) - add_breadcrumb( - category="litellm.llm_call", - message=f"Keyword Args: {details_to_log}", - level="info", - ) - if "logger_fn" in kwargs: - user_logger_fn = kwargs["logger_fn"] - # INIT LOGGER - for user-specified integrations - model: Final = args[0] if len(args) > 0 else kwargs.get("model", None) - call_type: Final = original_function - if ( - call_type == CallTypes.completion.value - or call_type == CallTypes.acompletion.value - or call_type == CallTypes.anthropic_messages.value - ): - messages = None - if len(args) > 1: - messages = args[1] - elif kwargs.get("messages", None): - messages = kwargs["messages"] - ### PRE-CALL RULES ### - Rules: Final = litellm_utils.Rules - if ( - Rules.has_pre_call_rules() - and isinstance(messages, list) - and len(messages) > 0 - and isinstance(messages[0], dict) - and "content" in messages[0] - ): - buffer: Final = StringIO() - for m in messages: - content = m.get("content", "") - if content is not None and isinstance(content, str): - buffer.write(content) - - rules_obj.pre_call_rules( - input=buffer.getvalue(), - model=model, - ) - - ### REMOVE THOUGHT SIGNATURES FROM TOOL CALL IDS FOR NON-GEMINI MODELS ### - # Gemini models embed thought signatures in tool call IDs. When sending - # messages with tool calls to non-Gemini providers, we need to remove these - # signatures to ensure compatibility. - if isinstance(messages, list) and len(messages) > 0: - try: - from litellm.litellm_core_utils.get_llm_provider_logic import ( - get_llm_provider, - ) - from litellm.litellm_core_utils.prompt_templates.factory import ( - THOUGHT_SIGNATURE_SEPARATOR, - ) - - # Get custom_llm_provider to determine target provider - custom_llm_provider = kwargs.get("custom_llm_provider") - - # If custom_llm_provider not in kwargs, try to determine it from the model - if not custom_llm_provider and model: - try: - _, custom_llm_provider, _, _ = get_llm_provider( - model=model, - custom_llm_provider=custom_llm_provider, - ) - except Exception: - # If we can't determine the provider, skip this processing - pass - - # Only process if target is NOT a Gemini model - if not _is_gemini_model(model, custom_llm_provider): - verbose_logger.debug("Removing thought signatures from tool call IDs for non-Gemini model") - - # Process messages to remove thought signatures - processed_messages: Final = _remove_thought_signatures_from_messages( - messages, THOUGHT_SIGNATURE_SEPARATOR - ) - - # Update messages in kwargs or args - if "messages" in kwargs: - kwargs["messages"] = processed_messages - elif len(args) > 1: - args_list: Final = list(args) - args_list[1] = processed_messages - args = tuple(args_list) - - except Exception as e: - # Log the error but don't fail the request - verbose_logger.warning("Error removing thought signatures from tool call IDs: %s", e) - elif call_type == CallTypes.embedding.value or call_type == CallTypes.aembedding.value: - messages = args[1] if len(args) > 1 else kwargs.get("input", None) - elif call_type == CallTypes.image_generation.value or call_type == CallTypes.aimage_generation.value: - messages = args[0] if len(args) > 0 else kwargs["prompt"] - elif call_type == CallTypes.moderation.value or call_type == CallTypes.amoderation.value: - messages = args[1] if len(args) > 1 else kwargs["input"] - elif call_type == CallTypes.atext_completion.value or call_type == CallTypes.text_completion.value: - messages = args[0] if len(args) > 0 else kwargs["prompt"] - elif call_type == CallTypes.rerank.value or call_type == CallTypes.arerank.value: - messages = kwargs.get("query") - elif call_type == CallTypes.atranscription.value or call_type == CallTypes.transcription.value: - _file_obj: Final[FileTypes] = args[1] if len(args) > 1 else kwargs["file"] - # Lazy import audio_utils.utils only when needed for transcription calls - audio_utils: Final = _get_cached_audio_utils() - file_checksum: Final = audio_utils.get_audio_file_content_hash(file_obj=_file_obj) - if "metadata" in kwargs: - kwargs["metadata"]["file_checksum"] = file_checksum - else: - kwargs["metadata"] = {"file_checksum": file_checksum} - messages = file_checksum - elif call_type == CallTypes.aspeech.value or call_type == CallTypes.speech.value: - messages = kwargs.get("input", "speech") - elif call_type == CallTypes.aresponses.value or call_type == CallTypes.responses.value: - # Handle both 'input' (standard Responses API) and 'messages' (Cursor chat format) - messages = ( - args[0] if len(args) > 0 else kwargs.get("input") or kwargs.get("messages", "default-message-value") - ) - elif ( - call_type == CallTypes.generate_content.value - or call_type == CallTypes.agenerate_content.value - or call_type == CallTypes.generate_content_stream.value - or call_type == CallTypes.agenerate_content_stream.value - ): - try: - from litellm.google_genai.adapters.transformation import ( - GoogleGenAIAdapter, - ) - from litellm.litellm_core_utils.prompt_templates.common_utils import ( - get_last_user_message, - ) - - contents_param: Final = args[1] if len(args) > 1 else kwargs.get("contents") - model_param: Final[str] = args[0] if len(args) > 0 else kwargs.get("model", "") - - if contents_param: - adapter: Final = GoogleGenAIAdapter() - transformed: Final = adapter.translate_generate_content_to_completion( - model=model_param, - contents=contents_param, - config=kwargs.get("config"), - ) - transformed_messages: Final = transformed.get("messages", []) - messages = get_last_user_message(transformed_messages) or "default-message-value" - else: - messages = "default-message-value" - except Exception as e: - verbose_logger.debug("Error extracting messages from Google contents: %s", e) - messages = "default-message-value" - elif call_type in NON_INFERENCE_CALL_TYPES: - messages = [] # mutable-ok: loggers require a list here and Logging copies it - else: - messages = "default-message-value" - stream = False - if _is_streaming_request( - kwargs=kwargs, - call_type=call_type, - ): - stream = True - get_litellm_logging_class: Final = getattr(sys.modules[__name__], "get_litellm_logging_class") - # Victim for object pool - logging_obj = get_litellm_logging_class()( # rebind-ok: 2nd assignment to logging_obj (see initial None above) - model=model, - messages=messages, - stream=stream, - litellm_call_id=kwargs["litellm_call_id"], - litellm_trace_id=kwargs.get("litellm_trace_id"), - function_id=function_id or "", - call_type=call_type, - start_time=start_time, - dynamic_success_callbacks=dynamic_success_callbacks, - dynamic_failure_callbacks=dynamic_failure_callbacks, - dynamic_async_success_callbacks=dynamic_async_success_callbacks, - dynamic_async_failure_callbacks=dynamic_async_failure_callbacks, - kwargs=kwargs, - applied_guardrails=applied_guardrails, - supports_correlation_logging=is_async_call, - ) - - ## check if metadata is passed in - litellm_params: Final[dict[str, object]] = {"api_base": ""} - if "metadata" in kwargs: - litellm_params["metadata"] = kwargs["metadata"] - if "litellm_metadata" in kwargs and isinstance(kwargs["litellm_metadata"], dict): - litellm_params["litellm_metadata"] = kwargs["litellm_metadata"] - # For endpoints like /v1/messages that use "litellm_metadata" instead - # of "metadata" (to avoid conflicting with provider API metadata fields), - # populate litellm_params["metadata"] so callbacks (e.g. Langfuse) that - # read API key info from litellm_params["metadata"] see the fields. - if not litellm_params.get("metadata"): - litellm_params["metadata"] = kwargs["litellm_metadata"].copy() - - logging_obj.update_environment_variables( - model=model, - user="", - optional_params={}, - litellm_params=litellm_params, - stream_options=kwargs.get("stream_options", None), - ) - return logging_obj, kwargs - except Exception as e: - # If Logging() was constructed above before this failed, its __init__ already - # mutated trace_id_var/session_id_var - restore them *before* logging the - # exception below, since we're about to raise without ever returning - # logging_obj to the caller's wrapper()/wrapper_async() (which would - # otherwise be the one doing this restore). Restoring first means this - # diagnostic log line itself doesn't get stamped with a call's ids when - # that call never actually produced a usable logging object. - if logging_obj is not None: - _restore_correlation_context_if_supported(logging_obj) - verbose_logger.exception("litellm.utils.py::function_setup() - [Non-Blocking] Error in function_setup") - raise e - - -async def _client_async_logging_helper( - logging_obj: LiteLLMLoggingObject, - result, - start_time, - end_time, - is_completion_with_fallbacks: bool, -): - if ( - is_completion_with_fallbacks is False - ): # don't log the parent event litellm.completion_with_fallbacks as a 'log_success_event', this will lead to double logging the same call - https://github.com/BerriAI/litellm/issues/7477 - print_verbose( - f"Async Wrapper: Completed Call, calling async_success_handler: {logging_obj.async_success_handler}" - ) - ################################################ - # Async Logging Worker - ################################################ - from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER - - GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue( - async_coroutine=logging_obj.async_success_handler(result=result, start_time=start_time, end_time=end_time) - ) - - ################################################ - # Sync Logging Worker - ################################################ - logging_obj.handle_sync_success_callbacks_for_async_calls( - result=result, - start_time=start_time, - end_time=end_time, - ) - - -def _get_wrapper_num_retries(kwargs: dict[str, Any], exception: Exception) -> tuple[int | None, dict[str, Any]]: - """ - Get the number of retries from the kwargs and the retry policy. - Used for the wrapper functions. - """ - - num_retries = kwargs.get("num_retries", None) - if num_retries is None: - num_retries = litellm.num_retries - if kwargs.get("retry_policy", None): - get_num_retries_from_retry_policy: Final[Callable[..., int | None]] = getattr( - sys.modules[__name__], "get_num_retries_from_retry_policy" - ) - reset_retry_policy: Final = litellm_utils.reset_retry_policy - retry_policy_num_retries: Final[int | None] = get_num_retries_from_retry_policy( - exception=exception, - retry_policy=kwargs.get("retry_policy"), - ) - kwargs["retry_policy"] = reset_retry_policy() - if retry_policy_num_retries is not None: - num_retries = retry_policy_num_retries - - return num_retries, kwargs - - -def _get_wrapper_timeout(kwargs: dict[str, object], exception: Exception) -> float | int | httpx.Timeout | None: - """ - Get the timeout from the kwargs - Used for the wrapper functions. - """ - - timeout: Final = cast(float | int | httpx.Timeout | None, kwargs.get("timeout", None)) - - return timeout - - -def check_coroutine(value) -> bool: - get_coroutine_checker: Final = litellm_utils.get_coroutine_checker - return get_coroutine_checker().is_async_callable(value) - - -async def async_pre_call_deployment_hook(kwargs: dict[str, Any], call_type: str): - """ - Allow modifying the request just before it's sent to the deployment. - - Use this instead of 'async_pre_call_hook' when you need to modify the request AFTER a deployment is selected, but BEFORE the request is sent. - """ - try: - typed_call_type = CallTypes(call_type) - except ValueError: - typed_call_type = None # unknown call type - - modified_kwargs = kwargs.copy() - - CustomLogger: Final = _get_cached_custom_logger() - for callback in litellm.callbacks: - if isinstance(callback, CustomLogger): - result = await callback.async_pre_call_deployment_hook(modified_kwargs, typed_call_type) - if result is not None: - modified_kwargs = result - - return modified_kwargs - - -async def async_post_call_success_deployment_hook( - request_data: dict, response: object, call_type: CallTypes | None -) -> Any | None: - """ - Allow modifying / reviewing the response just after it's received from the deployment. - """ - try: - typed_call_type = CallTypes(call_type) - except ValueError: - typed_call_type = None # unknown call type - - modified_response = response - - CustomLogger: Final = _get_cached_custom_logger() - for callback in litellm.callbacks: - if isinstance(callback, CustomLogger): - result = await callback.async_post_call_success_deployment_hook( - request_data, cast(LLMResponseTypes, modified_response), typed_call_type - ) - if result is not None: - modified_response = result - - return modified_response - - -async def async_post_call_failure_deployment_hook( - request_data: Mapping[str, object], exception: Exception, call_type: str -) -> None: - """ - Notify CustomLogger callbacks that a deployment attempt failed. - - Unlike its pre-call/post-success siblings, this wraps each callback call - in its own try/except: it runs on the wrapper's exception path, so a - broken callback must never replace the real exception that's about to be - re-raised to the caller. - - Reads ``fallback_depth`` off ``request_data`` (set by ``Router`` on each - fallback hop) and passes it through to the callback; ``None`` when - missing or not an int, since a bare SDK call has no fallback chain. - - Callbacks receive a same-class snapshot of ``exception``, not the live - object that's about to be re-raised, so a callback setting an attribute - on it (e.g. ``status_code``) cannot change what the real caller sees. - ``request_data`` omits ``attempted_targets``: unlike the rest of this - attempt's own kwargs, it's the *same* object shared by reference across - every hop of the live fallback walk, so a callback calling ``.record()`` - on it would make the router skip a deployment it hasn't actually tried. - """ - try: - typed_call_type = CallTypes(call_type) - except ValueError: - typed_call_type = None # unknown call type - - _raw_fallback_depth: Final = request_data.get("fallback_depth") - fallback_depth: Final = _raw_fallback_depth if isinstance(_raw_fallback_depth, int) else None - safe_request_data: Final = MappingProxyType({k: v for k, v in request_data.items() if k != "attempted_targets"}) - safe_exception: Final = _snapshot_exception_for_hook(exception) - - CustomLogger: Final = _get_cached_custom_logger() - for callback in litellm.callbacks: - if isinstance(callback, CustomLogger): - try: - if _accepts_fallback_depth_kwarg_for_class(type(callback)): - await callback.async_post_call_failure_deployment_hook( - safe_request_data, safe_exception, typed_call_type, fallback_depth=fallback_depth - ) - else: - await callback.async_post_call_failure_deployment_hook( - safe_request_data, safe_exception, typed_call_type - ) - except Exception as callback_error: # noqa: BLE001 # a broken callback must not mask the real failure - verbose_logger.debug( - "async_post_call_failure_deployment_hook error in %s: %s", - type(callback).__name__, - callback_error, - ) - - -def post_call_processing( - original_response, - model, - optional_params: dict | None, - original_function, - rules_obj, -): - try: - if original_response is None: - pass - else: - call_type: Final = original_function.__name__ - if call_type == CallTypes.completion.value or call_type == CallTypes.acompletion.value: - is_coroutine: Final = check_coroutine(original_response) - if is_coroutine is True: - pass - else: - if isinstance(original_response, ModelResponse) and len(original_response.choices) > 0: - model_response: Final[str | None] = original_response.choices[0].message.content - if model_response is not None: - ### POST-CALL RULES ### - rules_obj.post_call_rules(input=model_response, model=model) - ### JSON SCHEMA VALIDATION ### - # Per-request flag takes priority over global flag - _per_request_validation: Final = ( - optional_params.get("enable_json_schema_validation") - if optional_params is not None - else None - ) - _enable_json_schema_validation: Final = ( - _per_request_validation - if _per_request_validation is not None - else litellm.enable_json_schema_validation - ) - if _enable_json_schema_validation is True: - try: - if ( - optional_params is not None - and "response_format" in optional_params - and optional_params["response_format"] is not None - ): - json_response_format: dict | None = None - if ( - isinstance( - optional_params["response_format"], - dict, - ) - and optional_params["response_format"].get("json_schema") is not None - ): - json_response_format = optional_params["response_format"] - elif _parsing._completions.is_basemodel_type( - optional_params["response_format"] - ): - json_response_format = type_to_response_format_param( - response_format=optional_params["response_format"] - ) - if json_response_format is not None: - litellm.litellm_core_utils.json_validation_rule.validate_schema( - schema=json_response_format["json_schema"]["schema"], - response=model_response, - ) - except TypeError: - pass - if ( - optional_params is not None - and "response_format" in optional_params - and isinstance(optional_params["response_format"], dict) - and "type" in optional_params["response_format"] - and optional_params["response_format"]["type"] == "json_object" - and "response_schema" in optional_params["response_format"] - and isinstance( - optional_params["response_format"]["response_schema"], - dict, - ) - and "enforce_validation" in optional_params["response_format"] - and optional_params["response_format"]["enforce_validation"] is True - ): - # schema given, json response expected, and validation enforced - litellm.litellm_core_utils.json_validation_rule.validate_schema( - schema=optional_params["response_format"]["response_schema"], - response=model_response, - ) - - except Exception as e: - raise e - - -def client(original_function): - Rules: Final = litellm_utils.Rules - rules_obj: Final = Rules() - - @wraps(original_function) - def wrapper(*args, **kwargs): - # DO NOT MOVE THIS. It always needs to run first - # Check if this is an async function. If so only execute the async function - call_type = original_function.__name__ - if _is_async_request(kwargs): - # [OPTIONAL] CHECK MAX RETRIES / REQUEST - if litellm.num_retries_per_request is not None: - # check if previous_models passed in as ['litellm_params']['metadata]['previous_models'] - previous_models = (kwargs.get("metadata") or {}).get("previous_models", None) - if previous_models is not None: - if litellm.num_retries_per_request <= len(previous_models): - raise Exception("Max retries per request hit!") - - # MODEL CALL - result = original_function(*args, **kwargs) - if _is_streaming_request( - kwargs=kwargs, - call_type=call_type, - ): - if "complete_response" in kwargs and kwargs["complete_response"] is True: - chunks = [] - for idx, chunk in enumerate(result): - chunks.append(chunk) - return litellm.stream_chunk_builder(chunks, messages=kwargs.get("messages", None)) - else: - return result - - return result - - # Prints Exactly what was passed to litellm function - don't execute any logic here - it should just print - print_args_passed_to_litellm(original_function, args, kwargs) - start_time: Final = datetime.datetime.now() - result = None - logging_obj: LiteLLMLoggingObject | None = kwargs.get("litellm_logging_obj", None) - - # only set litellm_call_id if its not in kwargs - if "litellm_call_id" not in kwargs: - kwargs["litellm_call_id"] = str(uuid.uuid4()) - - model: Final[str | None] = args[0] if len(args) > 0 else kwargs.get("model", None) - - try: - if logging_obj is None: - logging_obj, kwargs = function_setup( - original_function.__name__, rules_obj, start_time, *args, is_async_call=False, **kwargs - ) - - # Type assertion: logging_obj is guaranteed to be non-None after function_setup - assert logging_obj is not None, "logging_obj should not be None after function_setup" - - ## LOAD CREDENTIALS - load_credentials_from_list(kwargs) - kwargs["litellm_logging_obj"] = logging_obj - LLMCachingHandler: Final = _get_cached_llm_caching_handler() - _llm_caching_handler: Final[LLMCachingHandler] = LLMCachingHandler( - original_function=original_function, - request_kwargs=kwargs, - start_time=start_time, - ) - logging_obj._llm_caching_handler = _llm_caching_handler - - # [OPTIONAL] CHECK BUDGET - if litellm.max_budget: - if litellm._current_cost > litellm.max_budget: - raise BudgetExceededError( - current_cost=litellm._current_cost, - max_budget=litellm.max_budget, - ) - - # [OPTIONAL] CHECK MAX RETRIES / REQUEST - if litellm.num_retries_per_request is not None: - # check if previous_models passed in as ['litellm_params']['metadata]['previous_models'] - previous_models = (kwargs.get("metadata") or {}).get("previous_models", None) - if previous_models is not None: - if litellm.num_retries_per_request <= len(previous_models): - raise Exception("Max retries per request hit!") - - # [OPTIONAL] CHECK CACHE - print_verbose( - f"SYNC kwargs[caching]: {kwargs.get('caching', False)}; litellm.cache: {litellm.cache}; kwargs.get('cache')['no-cache']: {kwargs.get('cache', {}).get('no-cache', False)}" - ) - # if caching is false or cache["no-cache"]==True, don't run this - if ( - ( - ( - (kwargs.get("caching", None) is None and litellm.cache is not None) - or kwargs.get("caching", False) is True - ) - and kwargs.get("cache", {}).get("no-cache", False) is not True - ) - and kwargs.get("aembedding", False) is not True - and kwargs.get("atext_completion", False) is not True - and kwargs.get("acompletion", False) is not True - and kwargs.get("aimg_generation", False) is not True - and kwargs.get("atranscription", False) is not True - and kwargs.get("arerank", False) is not True - and kwargs.get("_arealtime", False) is not True - ): # allow users to control returning cached responses from the completion function - # checking cache - verbose_logger.debug("INSIDE CHECKING SYNC CACHE") - caching_handler_response: Final[CachingHandlerResponse] = _llm_caching_handler._sync_get_cache( - model=model or "", - original_function=original_function, - logging_obj=logging_obj, - start_time=start_time, - call_type=call_type, - kwargs=kwargs, - args=args, - ) - - if caching_handler_response.cached_result is not None: - verbose_logger.debug("Cache hit!") - return caching_handler_response.cached_result - - # CHECK MAX TOKENS - if ( - kwargs.get("max_tokens", None) is not None - and model is not None - and litellm.modify_params is True # user is okay with params being modified - and ( - call_type == CallTypes.acompletion.value - or call_type == CallTypes.completion.value - or call_type == CallTypes.anthropic_messages.value - ) - ): - try: - base_model = model - if kwargs.get("hf_model_name", None) is not None: - base_model = f"huggingface/{kwargs.get('hf_model_name')}" - messages = None - if len(args) > 1: - messages = args[1] - elif kwargs.get("messages", None): - messages = kwargs["messages"] - user_max_tokens: Final = kwargs.get("max_tokens") - modified_max_tokens: Final = _get_modified_max_tokens()( - model=model, - base_model=base_model, - messages=messages, - user_max_tokens=user_max_tokens, - buffer_num=None, - buffer_perc=None, - ) - kwargs["max_tokens"] = modified_max_tokens - except Exception as e: - print_verbose(f"Error while checking max token limit: {e}") - # MODEL CALL - result = original_function(*args, **kwargs) - end_time = datetime.datetime.now() - if _is_streaming_request( - kwargs=kwargs, - call_type=call_type, - ): - if "complete_response" in kwargs and kwargs["complete_response"] is True: - chunks = [] - for idx, chunk in enumerate(result): - chunks.append(chunk) - return litellm.stream_chunk_builder(chunks, messages=kwargs.get("messages", None)) - else: - # RETURN RESULT - update_response_metadata = getattr(sys.modules[__name__], "update_response_metadata") - update_response_metadata( - result=result, - logging_obj=logging_obj, - model=model, - kwargs=kwargs, - start_time=start_time, - end_time=end_time, - ) - return result - elif ( - "acompletion" in kwargs - and kwargs["acompletion"] is True - or "aembedding" in kwargs - and kwargs["aembedding"] is True - or "aimg_generation" in kwargs - and kwargs["aimg_generation"] is True - or "atranscription" in kwargs - and kwargs["atranscription"] is True - or "aspeech" in kwargs - and kwargs["aspeech"] is True - or asyncio.iscoroutine(result) - ): - return result - - ### POST-CALL RULES ### - post_call_processing( - original_response=result, - model=model or None, - optional_params=kwargs, - original_function=original_function, - rules_obj=rules_obj, - ) - - # [OPTIONAL] ADD TO CACHE - _llm_caching_handler.sync_set_cache( - result=result, - args=args, - kwargs=kwargs, - ) - - # LOG SUCCESS - handle streaming success logging in the _next_ object, remove `handle_success` once it's deprecated - verbose_logger.info("Wrapper: Completed Call, calling success_handler") - # Copy the current context to propagate it to the background thread - # This is essential for OpenTelemetry span context propagation - ctx: Final = contextvars.copy_context() - executor: Final = getattr(sys.modules[__name__], "executor") - executor.submit( - ctx.run, - logging_obj.success_handler, - result, - start_time, - end_time, - ) - # RETURN RESULT - update_response_metadata = getattr(sys.modules[__name__], "update_response_metadata") - update_response_metadata( - result=result, - logging_obj=logging_obj, - model=model, - kwargs=kwargs, - start_time=start_time, - end_time=end_time, - ) - return result - except Exception as e: - call_type = original_function.__name__ - if call_type == CallTypes.completion.value: - num_retries = kwargs.get("num_retries", None) or litellm.num_retries or None - if kwargs.get("retry_policy", None): - get_num_retries_from_retry_policy: Callable[..., int | None] = getattr( - sys.modules[__name__], "get_num_retries_from_retry_policy" - ) - reset_retry_policy = litellm_utils.reset_retry_policy - num_retries = get_num_retries_from_retry_policy( - exception=e, - retry_policy=kwargs.get("retry_policy"), - ) - kwargs["retry_policy"] = reset_retry_policy() # prevent infinite loops - litellm.num_retries = None # set retries to None to prevent infinite loops - context_window_fallback_dict: Final = kwargs.get("context_window_fallback_dict", {}) - - _is_litellm_router_call = "model_group" in ( - kwargs.get("metadata") or {} - ) # check if call from litellm.router/proxy - if ( - num_retries and not _is_litellm_router_call - ): # only enter this if call is not from litellm router/proxy. router has it's own logic for retrying - if ( - isinstance(e, openai.APIError) - or isinstance(e, openai.Timeout) - or isinstance(e, openai.APIConnectionError) - ): - kwargs["num_retries"] = num_retries - return litellm.completion_with_retries(*args, **kwargs) - elif ( - isinstance(e, litellm.exceptions.ContextWindowExceededError) - and context_window_fallback_dict - and model in context_window_fallback_dict - and not _is_litellm_router_call - ): - if len(args) > 0: - args[0] = context_window_fallback_dict[model] - else: - kwargs["model"] = context_window_fallback_dict[model] - return original_function(*args, **kwargs) - elif call_type == CallTypes.responses.value: - num_retries = kwargs.get("num_retries", None) or litellm.num_retries or None - if kwargs.get("retry_policy", None): - get_num_retries_from_retry_policy = getattr( - sys.modules[__name__], "get_num_retries_from_retry_policy" - ) - reset_retry_policy = litellm_utils.reset_retry_policy - num_retries = get_num_retries_from_retry_policy( - exception=e, - retry_policy=kwargs.get("retry_policy"), - ) - kwargs["retry_policy"] = reset_retry_policy() # prevent infinite loops - litellm.num_retries = None # set retries to None to prevent infinite loops - - _is_litellm_router_call = "model_group" in ( - kwargs.get("metadata") or {} - ) # check if call from litellm.router/proxy - if ( - num_retries and not _is_litellm_router_call - ): # only enter this if call is not from litellm router/proxy. router has it's own logic for retrying - if ( - isinstance(e, openai.APIError) - or isinstance(e, openai.Timeout) - or isinstance(e, openai.APIConnectionError) - ): - kwargs["num_retries"] = num_retries - return litellm.responses_with_retries(*args, **kwargs) - traceback_exception: Final = traceback.format_exc() - end_time = datetime.datetime.now() - - # LOG FAILURE - handle streaming failure logging in the _next_ object, remove `handle_failure` once it's deprecated - if logging_obj: - logging_obj.failure_handler( - e, traceback_exception, start_time, end_time - ) # DO NOT MAKE THREADED - router retry fallback relies on this! - raise e - - @wraps(original_function) - async def wrapper_async(*args, **kwargs): - print_args_passed_to_litellm(original_function, args, kwargs) - start_time: Final = datetime.datetime.now() - result = None - _update_response_metadata: Final = getattr(sys.modules[__name__], "update_response_metadata") - logging_obj: LiteLLMLoggingObject | None = kwargs.get("litellm_logging_obj", None) - LLMCachingHandler: Final = _get_cached_llm_caching_handler() - _llm_caching_handler: Final[LLMCachingHandler] = LLMCachingHandler( - original_function=original_function, - request_kwargs=kwargs, - start_time=start_time, - ) - # only set litellm_call_id if its not in kwargs - call_type = original_function.__name__ - if "litellm_call_id" not in kwargs: - kwargs["litellm_call_id"] = str(uuid.uuid4()) - - model: Final[str | None] = args[0] if len(args) > 0 else kwargs.get("model", None) - is_completion_with_fallbacks: Final = kwargs.get("fallbacks") is not None - kwargs.pop("_is_litellm_internal_call", None) # discard if injected - _is_litellm_internal_call: Final = is_internal_call.get() - _deployment_call_end_time: datetime.datetime | None = ( - None # rebind-ok: set once, from inside the except below, only if the model call itself fails - ) - - try: - if logging_obj is None: - logging_obj, kwargs = function_setup(original_function.__name__, rules_obj, start_time, *args, **kwargs) - - # Type assertion: logging_obj is guaranteed to be non-None after function_setup - assert logging_obj is not None, "logging_obj should not be None after function_setup" - - modified_kwargs: Final = await async_pre_call_deployment_hook(kwargs, call_type) - if modified_kwargs is not None: - kwargs = modified_kwargs - - # Sync logging_obj.stream after deployment hooks (they may convert it). - _hook_stream: Final = kwargs.get("stream") - if _hook_stream is not None and logging_obj.stream != _hook_stream: - logging_obj.stream = _hook_stream - - kwargs["litellm_logging_obj"] = logging_obj - ## LOAD CREDENTIALS - load_credentials_from_list(kwargs) - logging_obj._llm_caching_handler = _llm_caching_handler - # [OPTIONAL] CHECK BUDGET - if litellm.max_budget: - if litellm._current_cost > litellm.max_budget: - raise BudgetExceededError( - current_cost=litellm._current_cost, - max_budget=litellm.max_budget, - ) - - # [OPTIONAL] CHECK CACHE - if _is_debugging_on(): - print_verbose( - f"ASYNC kwargs[caching]: {kwargs.get('caching', False)}; litellm.cache: {litellm.cache}; kwargs.get('cache'): {kwargs.get('cache', None)}" - ) - _caching_handler_response: CachingHandlerResponse | None = await _llm_caching_handler._async_get_cache( - model=model or "", - original_function=original_function, - logging_obj=logging_obj, - start_time=start_time, - call_type=call_type, - kwargs=kwargs, - args=args, - ) - - if _caching_handler_response is not None: - if ( - _caching_handler_response.cached_result is not None - and _caching_handler_response.final_embedding_cached_response is None - ): - return _caching_handler_response.cached_result - - elif _caching_handler_response.embedding_all_elements_cache_hit is True: - return _caching_handler_response.final_embedding_cached_response - - # CHECK MAX TOKENS - if ( - kwargs.get("max_tokens", None) is not None - and model is not None - and litellm.modify_params is True # user is okay with params being modified - and ( - call_type == CallTypes.acompletion.value - or call_type == CallTypes.completion.value - or call_type == CallTypes.anthropic_messages.value - ) - ): - try: - base_model = model - if kwargs.get("hf_model_name", None) is not None: - base_model = f"huggingface/{kwargs.get('hf_model_name')}" - messages = None - if len(args) > 1: - messages = args[1] - elif kwargs.get("messages", None): - messages = kwargs["messages"] - user_max_tokens: Final = kwargs.get("max_tokens") - modified_max_tokens: Final = _get_modified_max_tokens()( - model=model, - base_model=base_model, - messages=messages, - user_max_tokens=user_max_tokens, - buffer_num=None, - buffer_perc=None, - ) - kwargs["max_tokens"] = modified_max_tokens - except Exception as e: - print_verbose(f"Error while checking max token limit: {e}") - - # MODEL CALL - try: - result = await original_function(*args, **kwargs) - except Exception as deployment_error: - _deployment_call_end_time = datetime.datetime.now() # noqa: DTZ005 # matches the naive datetimes this whole function already times start_time/end_time with - try: - await async_post_call_failure_deployment_hook( - request_data=kwargs, - exception=deployment_error, - call_type=call_type, - ) - except BaseException: # noqa: S110, BLE001 # hook dispatch - including cancellation mid-await - must never replace the real deployment failure, so there is nothing to do with what it raises - pass - raise - end_time = datetime.datetime.now() - - if _is_streaming_request( - kwargs=kwargs, - call_type=call_type, - ): - if "complete_response" in kwargs and kwargs["complete_response"] is True: - chunks: Final = [] - for idx, chunk in enumerate(result): - chunks.append(chunk) - return litellm.stream_chunk_builder(chunks, messages=kwargs.get("messages", None)) - else: - _update_response_metadata( - result=result, - logging_obj=logging_obj, - model=model, - kwargs=kwargs, - start_time=start_time, - end_time=end_time, - ) - return _llm_caching_handler.wrap_streaming_result_for_cache( - result=result, - call_type=call_type, - ) - elif call_type == CallTypes.arealtime.value: - return result - ### POST-CALL RULES ### - post_call_processing( - original_response=result, - model=model, - optional_params=kwargs, - original_function=original_function, - rules_obj=rules_obj, - ) - # Only run if call_type is a valid value in CallTypes - _call_type_enum: Final = _CALL_TYPE_ENUM_MAP.get(call_type) - if _call_type_enum is not None: - result = await async_post_call_success_deployment_hook( - request_data=kwargs, - response=result, - call_type=_call_type_enum, - ) - - ## Add response to cache - await _llm_caching_handler.async_set_cache( - result=result, - original_function=original_function, - kwargs=kwargs, - args=args, - ) - - # LOG SUCCESS - handle streaming success logging in the _next_ object - # Internal sub-calls (e.g. emulated file-search steps) share the - # parent's logging obj; skip async logging here so only the outer call bills once. - # NOTE: streaming requests return early (before this point) via - # CustomStreamWrapper, so this block is non-streaming only. - if not _is_litellm_internal_call: - if getattr(logging_obj, "_defer_async_logging", False): - - def _enqueue_deferred_logging() -> None: - asyncio.create_task( - _client_async_logging_helper( - logging_obj=logging_obj, - result=result, - start_time=start_time, - end_time=end_time, - is_completion_with_fallbacks=is_completion_with_fallbacks, - ) - ) - - logging_obj._enqueue_deferred_logging = _enqueue_deferred_logging - else: - asyncio.create_task( - _client_async_logging_helper( - logging_obj=logging_obj, - result=result, - start_time=start_time, - end_time=end_time, - is_completion_with_fallbacks=is_completion_with_fallbacks, - ) - ) - - logging_obj.handle_sync_success_callbacks_for_async_calls( - result=result, - start_time=start_time, - end_time=end_time, - ) - # REBUILD EMBEDDING CACHING - if ( - isinstance(result, EmbeddingResponse) - and _caching_handler_response is not None - and _caching_handler_response.final_embedding_cached_response is not None - ): - return _llm_caching_handler._combine_cached_embedding_response_with_api_result( - _caching_handler_response=_caching_handler_response, - embedding_response=result, - start_time=start_time, - end_time=end_time, - ) - - _update_response_metadata( - result=result, - logging_obj=logging_obj, - model=model, - kwargs=kwargs, - start_time=start_time, - end_time=end_time, - ) - - return result - except Exception as e: - traceback_exception: Final = traceback.format_exc() - # Reuse the timestamp taken right when the deployment call itself failed, before - # the failure hook ran, so a slow callback doesn't inflate the reported duration. - end_time = _deployment_call_end_time if _deployment_call_end_time is not None else datetime.datetime.now() # noqa: DTZ005 # matches the naive datetimes this whole function already times start_time/end_time with - if logging_obj and not _is_litellm_internal_call: - try: - logging_obj.failure_handler( - e, traceback_exception, start_time, end_time - ) # DO NOT MAKE THREADED - router retry fallback relies on this! - except Exception as e: - raise e - try: - await logging_obj.async_failure_handler(e, traceback_exception, start_time, end_time) - except Exception as e: - raise e - - call_type = original_function.__name__ - num_retries, kwargs = _get_wrapper_num_retries(kwargs=kwargs, exception=e) - if call_type == CallTypes.acompletion.value: - context_window_fallback_dict: Final = kwargs.get("context_window_fallback_dict", {}) - - _is_litellm_router_call = "model_group" in ( - kwargs.get("metadata") or {} - ) # check if call from litellm.router/proxy - - if ( - num_retries and not _is_litellm_router_call - ): # only enter this if call is not from litellm router/proxy. router has it's own logic for retrying - try: - litellm.num_retries = None # set retries to None to prevent infinite loops - kwargs["num_retries"] = num_retries - kwargs["original_function"] = original_function - if isinstance(e, openai.RateLimitError): # rate limiting specific error - kwargs["retry_strategy"] = "exponential_backoff_retry" - elif isinstance(e, openai.APIError): # generic api error - kwargs["retry_strategy"] = "constant_retry" - result = await litellm.acompletion_with_retries(*args, **kwargs) - except Exception: - pass - else: - return result - elif ( - isinstance(e, litellm.exceptions.ContextWindowExceededError) - and context_window_fallback_dict - and model in context_window_fallback_dict - and not _is_litellm_router_call - ): - if len(args) > 0: - args[0] = context_window_fallback_dict[model] - else: - kwargs["model"] = context_window_fallback_dict[model] - result = await original_function(*args, **kwargs) - return result - elif call_type == CallTypes.aresponses.value: - _is_litellm_router_call = "model_group" in ( - kwargs.get("metadata") or {} - ) # check if call from litellm.router/proxy - - if ( - num_retries and not _is_litellm_router_call - ): # only enter this if call is not from litellm router/proxy. router has it's own logic for retrying - try: - litellm.num_retries = None # set retries to None to prevent infinite loops - kwargs["num_retries"] = num_retries - kwargs["original_function"] = original_function - if isinstance(e, openai.RateLimitError): # rate limiting specific error - kwargs["retry_strategy"] = "exponential_backoff_retry" - elif isinstance(e, openai.APIError): # generic api error - kwargs["retry_strategy"] = "constant_retry" - result = await litellm.aresponses_with_retries(*args, **kwargs) - except Exception: - pass - else: - return result - - deployment_num_retries: Final = kwargs.get("num_retries") - if deployment_num_retries is not None: - setattr(e, "num_retries", deployment_num_retries) - - timeout: Final = _get_wrapper_timeout(kwargs=kwargs, exception=e) - setattr(e, "timeout", timeout) - raise e - - finally: - # Restore trace_id/session_id contextvars to their pre-call value once - # this call (in this asyncio Task) is fully done - see - # request_correlation_in_logs. Unlike wrapper()'s sync path, it's safe to - # skip restoring when returning a stream: each async call already runs in - # its own Task with its own copy of the contextvars (asyncio.create_task - # copies context at creation), so leaving this Task's own view "open" - # while the caller iterates the stream can only affect that one Task - - # never a different, unrelated future request, since Tasks (unlike a - # thread pool's worker threads) are never recycled across requests. The - # corresponding terminal handler (async_success_handler) restores it once - # streaming genuinely finishes; aclose()/__del__ cover early termination. - if not _is_streaming_response_for_correlation(result): - _restore_correlation_context_if_supported(logging_obj) - - get_coroutine_checker: Final = litellm_utils.get_coroutine_checker - is_coroutine: Final = get_coroutine_checker().is_async_callable(original_function) - - # Return the appropriate wrapper based on the original function type - if is_coroutine: - return wrapper_async - else: - return wrapper - - -def _is_async_request( - kwargs: dict | None, - is_pass_through: bool = False, -) -> bool: - """ - Returns True if the call type is an internal async request. - - eg. litellm.acompletion, litellm.aimage_generation, litellm.acreate_batch, litellm._arealtime - - Args: - kwargs (dict): The kwargs passed to the litellm function - is_pass_through (bool): Whether the call is a pass-through call. By default all pass through calls are async. - """ - if kwargs is None: - return False - if ( - kwargs.get("acompletion", False) is True - or kwargs.get("aembedding", False) is True - or kwargs.get("aimg_generation", False) is True - or kwargs.get("amoderation", False) is True - or kwargs.get("atext_completion", False) is True - or kwargs.get("atranscription", False) is True - or kwargs.get("arerank", False) is True - or kwargs.get("_arealtime", False) is True - or kwargs.get("acreate_batch", False) is True - or kwargs.get("acreate_fine_tuning_job", False) is True - or is_pass_through is True - ): - return True - return False - - -_STREAMING_CALL_TYPES: Final = frozenset( - { - CallTypes.generate_content_stream, - CallTypes.agenerate_content_stream, - CallTypes.generate_content_stream.value, - CallTypes.agenerate_content_stream.value, - } -) - - -def _is_streaming_request( - kwargs: dict[str, object], - call_type: CallTypes | str, -) -> bool: - """ - Returns True if the call type is a streaming request. - Returns True if: - - if "stream=True" in kwargs (litellm chat completion, litellm text completion, litellm messages) - - if call_type is generate_content_stream or agenerate_content_stream (litellm google genai) - """ - if "stream" in kwargs and kwargs["stream"] is True: - return True - return call_type in _STREAMING_CALL_TYPES - - -def _select_tokenizer(model: str, custom_tokenizer: CustomHuggingfaceTokenizer | None = None): - if custom_tokenizer is not None: - _tokenizer: Final = create_pretrained_tokenizer( - identifier=custom_tokenizer["identifier"], - revision=custom_tokenizer["revision"], - auth_token=custom_tokenizer["auth_token"], - ) - return _tokenizer - return _select_tokenizer_helper(model=model) - - -@lru_cache(maxsize=DEFAULT_MAX_LRU_CACHE_SIZE) -def _select_tokenizer_helper(model: str) -> SelectTokenizerResponse: - if litellm.disable_hf_tokenizer_download is True: - return _return_openai_tokenizer(model) - - try: - result: Final = _return_huggingface_tokenizer(model) - if result is not None: - return result - except Exception as e: - verbose_logger.debug("Error selecting tokenizer: %s", e) - - # default - tiktoken - return _return_openai_tokenizer(model) - - -def _return_openai_tokenizer(model: str) -> SelectTokenizerResponse: - return {"type": "openai_tokenizer", "tokenizer": _get_default_encoding()} - - -def _return_huggingface_tokenizer(model: str) -> SelectTokenizerResponse | None: - if model in litellm.cohere_models and "command-r" in model: - # cohere - cohere_tokenizer: Final = Tokenizer.from_pretrained("Xenova/c4ai-command-r-v01-tokenizer") - return {"type": "huggingface_tokenizer", "tokenizer": cohere_tokenizer} - # anthropic - elif model in litellm.anthropic_models and "claude-3" not in model: - claude_tokenizer: Final = Tokenizer.from_str(claude_json_str) - return {"type": "huggingface_tokenizer", "tokenizer": claude_tokenizer} - # llama2 - elif "llama-2" in model.lower() or "replicate" in model.lower(): - tokenizer = Tokenizer.from_pretrained("hf-internal-testing/llama-tokenizer") - return {"type": "huggingface_tokenizer", "tokenizer": tokenizer} - # llama3 - elif "llama-3" in model.lower(): - tokenizer = Tokenizer.from_pretrained("Xenova/llama-3-tokenizer") - return {"type": "huggingface_tokenizer", "tokenizer": tokenizer} - else: - return None - - -def encode(model="", text="", custom_tokenizer: dict | None = None): - """ - Encodes the given text using the specified model. - - Args: - model (str): The name of the model to use for tokenization. - custom_tokenizer (Optional[dict]): A custom tokenizer created with the `create_pretrained_tokenizer` or `create_tokenizer` method. Must be a dictionary with a string value for `type` and Tokenizer for `tokenizer`. Default is None. - text (str): The text to be encoded. - - Returns: - enc: The encoded text. - """ - tokenizer_json: Final = custom_tokenizer or _select_tokenizer(model=model) - if isinstance(tokenizer_json["tokenizer"], Encoding): - enc = tokenizer_json["tokenizer"].encode(text, disallowed_special=()) - else: - enc = tokenizer_json["tokenizer"].encode(text) - # Normalize: HuggingFace Tokenizer.encode() returns an Encoding object; - # extract .ids so the return type is always List[int]. - if hasattr(enc, "ids"): - return enc.ids - return enc - - -def decode( - model="", - tokens: Sequence[int] = (), - custom_tokenizer: dict | None = None, - skip_special_tokens: bool = True, -): - """ - Decodes token ids using the selected tokenizer. - - Args: - skip_special_tokens: For HuggingFace tokenizers, keep the historical - LiteLLM round-trip behavior by omitting special tokens by default. - Set to False to inspect decoded BOS/EOS tokens. - """ - tokenizer_json: Final = custom_tokenizer or _select_tokenizer(model=model) - if tokenizer_json["type"] == "huggingface_tokenizer": - if skip_special_tokens: - tokens = _strip_huggingface_special_token_ids(tokenizer_json["tokenizer"], tokens) - dec = tokenizer_json["tokenizer"].decode(tokens, skip_special_tokens=skip_special_tokens) - return dec - dec = tokenizer_json["tokenizer"].decode(tokens) - return dec - - -def _strip_huggingface_special_token_ids(tokenizer: Tokenizer, tokens: Sequence[int]) -> Sequence[int]: - try: - added_tokens_decoder: Final = tokenizer.get_added_tokens_decoder() - except Exception: - return tokens - - special_token_ids: Final = { - token_id for token_id, added_token in added_tokens_decoder.items() if getattr(added_token, "special", False) - } - if not special_token_ids: - return tokens - return [token for token in tokens if token not in special_token_ids] - - -def create_pretrained_tokenizer(identifier: str, revision="main", auth_token: str | None = None): - """ - Creates a tokenizer from an existing file on a HuggingFace repository to be used with `token_counter`. - - Args: - identifier (str): The identifier of a Model on the Hugging Face Hub, that contains a tokenizer.json file - revision (str, defaults to main): A branch or commit id - auth_token (str, optional, defaults to None): An optional auth token used to access private repositories on the Hugging Face Hub - - Returns: - dict: A dictionary with the tokenizer and its type. - """ - - try: - tokenizer = Tokenizer.from_pretrained( - identifier, - revision=revision, - auth_token=auth_token, - ) - except Exception as e: - verbose_logger.error("Error creating pretrained tokenizer: %s. Defaulting to version without 'auth_token'.", e) - tokenizer = Tokenizer.from_pretrained(identifier, revision=revision) - return {"type": "huggingface_tokenizer", "tokenizer": tokenizer} - - -def create_tokenizer(json: str): - """ - Creates a tokenizer from a valid JSON string for use with `token_counter`. - - Args: - json (str): A valid JSON string representing a previously serialized tokenizer - - Returns: - dict: A dictionary with the tokenizer and its type. - """ - - tokenizer: Final = Tokenizer.from_str(json) - return {"type": "huggingface_tokenizer", "tokenizer": tokenizer} - - -def token_counter( - model="", - custom_tokenizer: dict | SelectTokenizerResponse | None = None, - text: str | list[str] | None = None, - messages: Sequence | None = None, - count_response_tokens: bool | None = False, - tools: list[ChatCompletionToolParam] | None = None, - tool_choice: ChatCompletionNamedToolChoiceParam | None = None, - use_default_image_token_count: bool | None = False, - default_token_count: int | None = None, -) -> int: - """ - The same as `litellm.litellm_core_utils.token_counter`. - - Kept for backwards compatibility. - """ - - ######################################################### - # Flag to disable token counter - # We've gotten reports of this consuming CPU cycles, - # exposing this flag to allow users to disable - # it to confirm if this is indeed the issue - ######################################################### - if litellm.disable_token_counter is True: - return 0 - - return _get_token_counter_new()( - model, - custom_tokenizer, - text, - messages, - count_response_tokens, - tools, - tool_choice, - use_default_image_token_count, - default_token_count, - ) - - -def supports_httpx_timeout(custom_llm_provider: str) -> bool: - """ - Helper function to know if a provider implementation supports httpx timeout - """ - supported_providers: Final = ["openai", "azure", "bedrock"] - - if custom_llm_provider in supported_providers: - return True - - return False - - -def supports_system_messages(model: str, custom_llm_provider: str | None) -> bool: - """ - Check if the given model supports system messages and return a boolean value. - - Parameters: - model (str): The model name to be checked. - custom_llm_provider (str): The provider to be checked. - - Returns: - bool: True if the model supports system messages, False otherwise. - - Raises: - Exception: If the given model is not found in model_prices_and_context_window.json. - """ - return _supports_factory( - model=model, - custom_llm_provider=custom_llm_provider, - key="supports_system_messages", - ) - - -def supports_web_search(model: str, custom_llm_provider: str | None = None) -> bool: - """ - Check if the given model supports web search and return a boolean value. - - Parameters: - model (str): The model name to be checked. - custom_llm_provider (str): The provider to be checked. - - Returns: - bool: True if the model supports web search, False otherwise. - - Raises: - Exception: If the given model is not found in model_prices_and_context_window.json. - """ - return _supports_factory( - model=model, - custom_llm_provider=custom_llm_provider, - key="supports_web_search", - ) - - -def supports_url_context(model: str, custom_llm_provider: str | None = None) -> bool: - """ - Check if the given model supports URL context and return a boolean value. - - Parameters: - model (str): The model name to be checked. - custom_llm_provider (str): The provider to be checked. - - Returns: - bool: True if the model supports URL context, False otherwise. - - Raises: - Exception: If the given model is not found in model_prices_and_context_window.json. - """ - return _supports_factory( - model=model, - custom_llm_provider=custom_llm_provider, - key="supports_url_context", - ) - - -def supports_native_streaming(model: str, custom_llm_provider: str | None) -> bool: - """ - Check if the given model supports native streaming and return a boolean value. - - Parameters: - model (str): The model name to be checked. - custom_llm_provider (str): The provider to be checked. - - Returns: - bool: True if the model supports native streaming, False otherwise. - - Raises: - Exception: If the given model is not found in model_prices_and_context_window.json. - """ - try: - model, custom_llm_provider, _, _ = litellm.get_llm_provider( - model=model, custom_llm_provider=custom_llm_provider - ) - - model_info: Final = _get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider) - supports_native_streaming = model_info.get("supports_native_streaming", True) - if supports_native_streaming is None: - supports_native_streaming = True - return supports_native_streaming - except Exception as e: - verbose_logger.debug( - "Model not found or error in checking supports_native_streaming support. You passed model=%s, custom_llm_provider=%s. Error: %s", - model, - custom_llm_provider, - e, - ) - return False - - -def supports_response_schema(model: str, custom_llm_provider: str | None = None) -> bool: - """ - Check if the given model + provider supports 'response_schema' as a param. - - Parameters: - model (str): The model name to be checked. - custom_llm_provider (str): The provider to be checked. - - Returns: - bool: True if the model supports response_schema, False otherwise. - - Does not raise error. Defaults to 'False'. Outputs logging.error. - """ - ## GET LLM PROVIDER ## - try: - get_llm_provider: Final = litellm_utils.get_llm_provider - model, custom_llm_provider, _, _ = get_llm_provider(model=model, custom_llm_provider=custom_llm_provider) - except Exception as e: - verbose_logger.debug( - "Model not found or error in checking response schema support. You passed model=%s, custom_llm_provider=%s. Error: %s", - model, - custom_llm_provider, - e, - ) - return False - - # providers that globally support response schema - PROVIDERS_GLOBALLY_SUPPORT_RESPONSE_SCHEMA: Final = [ - litellm.LlmProviders.PREDIBASE, - litellm.LlmProviders.FIREWORKS_AI, - litellm.LlmProviders.LM_STUDIO, - litellm.LlmProviders.NEBIUS, - litellm.LlmProviders.DATABRICKS, - ] - - if custom_llm_provider in PROVIDERS_GLOBALLY_SUPPORT_RESPONSE_SCHEMA: - return True - return _supports_factory( - model=model, - custom_llm_provider=custom_llm_provider, - key="supports_response_schema", - ) - - -def supports_parallel_function_calling(model: str, custom_llm_provider: str | None = None) -> bool: - """ - Check if the given model supports parallel tool calls and return a boolean value. - """ - return _supports_factory( - model=model, - custom_llm_provider=custom_llm_provider, - key="supports_parallel_function_calling", - ) - - -def supports_function_calling(model: str, custom_llm_provider: str | None = None) -> bool: - """ - Check if the given model supports function calling and return a boolean value. - - Parameters: - model (str): The model name to be checked. - custom_llm_provider (Optional[str]): The provider to be checked. - - Returns: - bool: True if the model supports function calling, False otherwise. - - Raises: - Exception: If the given model is not found or there's an error in retrieval. - """ - return _supports_factory( - model=model, - custom_llm_provider=custom_llm_provider, - key="supports_function_calling", - ) - - -def supports_tool_choice(model: str, custom_llm_provider: str | None = None) -> bool: - """ - Check if the given model supports `tool_choice` and return a boolean value. - """ - return _supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_tool_choice") - - -def _supports_provider_info_factory(model: str, custom_llm_provider: str | None, key: str) -> Literal[True] | None: - """ - Check if the given model supports a provider specific model info and return a boolean value. - """ - - provider_info: Final = get_provider_info(model=model, custom_llm_provider=custom_llm_provider) - - if provider_info is not None and provider_info.get(key, False) is True: - return True - return None - - -def _supports_factory(model: str, custom_llm_provider: str | None, key: str) -> bool: - """ - Check if the given model supports function calling and return a boolean value. - - Parameters: - model (str): The model name to be checked. - custom_llm_provider (Optional[str]): The provider to be checked. - - Returns: - bool: True if the model supports function calling, False otherwise. - - Raises: - Exception: If the given model is not found or there's an error in retrieval. - """ - from litellm.litellm_core_utils.get_llm_provider_logic import declared_authenticating_provider - - try: - declared: Final = declared_authenticating_provider(model, custom_llm_provider) - if declared is not None: - model = model.removeprefix( - f"{declared}/" - ) # rebind-ok: mirrors get_llm_provider's split without its OAuth flow - custom_llm_provider = declared # rebind-ok: same - else: - model, custom_llm_provider, _, _ = litellm.get_llm_provider( - model=model, custom_llm_provider=custom_llm_provider - ) - - model_info: Final = _get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider) - - if model_info.get(key, False) is True: - return True - elif model_info.get(key) is None: # don't check if 'False' explicitly set - # Fallback: when the provider-prefixed entry (e.g. - # "deepseek/deepseek-chat") exists but is missing a capability - # field, check the bare model-name entry (e.g. "deepseek-chat") - # which may carry the complete metadata. See #20885. - bare_model_key: Final = _get_model_cost_key(model) - if bare_model_key is not None: - bare_entry: Final = litellm.model_cost.get(bare_model_key) or {} - if bare_entry.get(key, False) is True: - return True - - supported_by_provider = _supports_provider_info_factory(model, custom_llm_provider, key) - if supported_by_provider is not None: - return supported_by_provider - - return False - except Exception as e: - verbose_logger.debug( - "Model not found or error in checking %s support. You passed model=%s, custom_llm_provider=%s. Error: %s", - key, - model, - custom_llm_provider, - e, - ) - - supported_by_provider = _supports_provider_info_factory(model, custom_llm_provider, key) - if supported_by_provider is not None: - return supported_by_provider - - return False - - -def declared_value_factory(model: str, custom_llm_provider: str | None, key: str) -> str | None: - """Return a string value the model map declares for *key*, or ``None`` when it says nothing. - - The string-valued sibling of :func:`_supports_factory` and - :func:`_is_explicitly_disabled_factory`, public where those two are not because it is read - from the provider configs rather than from this module, sharing their - ``get_llm_provider`` -> ``_get_model_info_helper`` chain and their unprefixed-twin - fallback (#20885), so a provider-prefixed entry that omits the key still answers - from the bare entry that carries it. - - ``None`` means "the map does not say", never "the map says no" - callers decide what - an unknown declaration implies, and for a capability gate that decision must be the - conservative one. - """ - try: - resolved: Final = litellm.get_llm_provider(model=model, custom_llm_provider=custom_llm_provider) - resolved_model: Final = resolved[0] - resolved_provider: Final = resolved[1] - model_info: Final = _get_model_info_helper(model=resolved_model, custom_llm_provider=resolved_provider) - declared: Final = model_info.get(key) - if isinstance(declared, str): - return declared - bare_model_key: Final = _get_model_cost_key(resolved_model) - bare_entry: Final = litellm.model_cost.get(bare_model_key) if bare_model_key is not None else None - if isinstance(bare_entry, dict): - bare_declared: Final = bare_entry.get(key) - if isinstance(bare_declared, str): - return bare_declared - return None - except Exception as e: # noqa: BLE001 # an unreadable map entry means "not declared", never a failed call - verbose_logger.debug( - "Model not found or error in reading %s. You passed model=%s, custom_llm_provider=%s. Error: %s", - key, - model, - custom_llm_provider, - e, - ) - return None - - -def _is_explicitly_disabled_factory(model: str, custom_llm_provider: str | None, key: str) -> bool: - """Return True only when the model map explicitly sets *key* to ``False``. - - This is the opt-out mirror of :func:`_supports_factory`. Where - ``_supports_factory`` requires an explicit ``True`` to return ``True``, - this function requires an explicit ``False``. A missing key (``None``) - is treated as *not* disabled so that unknown or newly-added models are - allowed through without any model-map entry. - - Uses the same ``get_llm_provider`` → ``_get_model_info_helper`` chain as - ``_supports_factory`` so caching, fallback, and normalisation improvements - apply here automatically. - """ - from litellm.litellm_core_utils.get_llm_provider_logic import declared_authenticating_provider - - try: - declared: Final = declared_authenticating_provider(model, custom_llm_provider) - if declared is not None: - model = model.removeprefix( - f"{declared}/" - ) # rebind-ok: mirrors get_llm_provider's split without its OAuth flow - custom_llm_provider = declared # rebind-ok: same - else: - model, custom_llm_provider, _, _ = litellm.get_llm_provider( - model=model, custom_llm_provider=custom_llm_provider - ) - model_info: Final = _get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider) - val: Final = model_info.get(key) - if val is False: - return True - if val is None: - bare_model_key: Final = _get_model_cost_key(model) - if bare_model_key is not None: - bare_entry: Final = litellm.model_cost.get(bare_model_key) or {} - if bare_entry.get(key) is False: - return True - return False - except Exception as e: - verbose_logger.debug( - "Model not found or error in checking %s disabled state. You passed model=%s, custom_llm_provider=%s. Error: %s", - key, - model, - custom_llm_provider, - e, - ) - return False - - -def supports_audio_input(model: str, custom_llm_provider: str | None = None) -> bool: - """Check if a given model supports audio input in a chat completion call""" - return _supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_audio_input") - - -def supports_pdf_input(model: str, custom_llm_provider: str | None = None) -> bool: - """Check if a given model supports pdf input in a chat completion call""" - return _supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_pdf_input") - - -def supports_audio_output(model: str, custom_llm_provider: str | None = None) -> bool: - """Check if a given model supports audio output in a chat completion call""" - return _supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_audio_input") - - -def supports_prompt_caching(model: str, custom_llm_provider: str | None = None) -> bool: - """ - Check if the given model supports prompt caching and return a boolean value. - - Parameters: - model (str): The model name to be checked. - custom_llm_provider (Optional[str]): The provider to be checked. - - Returns: - bool: True if the model supports prompt caching, False otherwise. - - Raises: - Exception: If the given model is not found or there's an error in retrieval. - """ - return _supports_factory( - model=model, - custom_llm_provider=custom_llm_provider, - key="supports_prompt_caching", - ) - - -def supports_prompt_cache_breakpoint(model: str, custom_llm_provider: str | None = None) -> bool: - return _supports_factory( - model=model, - custom_llm_provider=custom_llm_provider, - key="supports_prompt_cache_breakpoint", - ) - - -def supports_computer_use(model: str, custom_llm_provider: str | None = None) -> bool: - """ - Check if the given model supports computer use and return a boolean value. - - Parameters: - model (str): The model name to be checked. - custom_llm_provider (Optional[str]): The provider to be checked. - - Returns: - bool: True if the model supports computer use, False otherwise. - - Raises: - Exception: If the given model is not found or there's an error in retrieval. - """ - return _supports_factory( - model=model, - custom_llm_provider=custom_llm_provider, - key="supports_computer_use", - ) - - -def is_vision_explicitly_disabled(model: str, custom_llm_provider: str | None = None) -> bool: - """True only when supports_vision is explicitly declared false for the model. - - The opt-out mirror of :func:`supports_vision`: a missing declaration reads as not - disabled, so unknown or newly added models stay eligible for image routing. - """ - return _is_explicitly_disabled_factory(model, custom_llm_provider, "supports_vision") - - -def supports_vision(model: str, custom_llm_provider: str | None = None) -> bool: - """ - Check if the given model supports vision and return a boolean value. - - Parameters: - model (str): The model name to be checked. - custom_llm_provider (Optional[str]): The provider to be checked. - - Returns: - bool: True if the model supports vision, False otherwise. - """ - return _supports_factory( - model=model, - custom_llm_provider=custom_llm_provider, - key="supports_vision", - ) - - -def supports_reasoning(model: str, custom_llm_provider: str | None = None) -> bool: - """ - Check if the given model supports reasoning and return a boolean value. - """ - return _supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_reasoning") - - -def supports_native_structured_output(model: str, custom_llm_provider: str | None = None) -> bool: - """ - Check if the given model supports native structured outputs and return a boolean value. - """ - return _supports_factory( - model=model, - custom_llm_provider=custom_llm_provider, - key="supports_native_structured_output", - ) - - -def get_supported_regions(model: str, custom_llm_provider: str | None = None) -> list[str] | None: - """ - Get a list of supported regions for a given model and provider. - - Parameters: - model (str): The model name to be checked. - custom_llm_provider (Optional[str]): The provider to be checked. - """ - try: - model, custom_llm_provider, _, _ = litellm.get_llm_provider( - model=model, custom_llm_provider=custom_llm_provider - ) - - model_info: Final = _get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider) - - # Get the key used in model_cost to look up supported_regions - # since ModelInfoBase doesn't include this field - model_key: Final = model_info.get("key") - if model_key is None: - return None - - model_cost_data: Final = litellm.model_cost.get(model_key, {}) - supported_regions: Final = model_cost_data.get("supported_regions", None) - if supported_regions is None: - return None - - ######################################################### - # Ensure only list supported regions are returned - ######################################################### - if isinstance(supported_regions, list): - return supported_regions - else: - return None - except Exception as e: - verbose_logger.debug( - "Model not found or error in checking supported_regions support. You passed model=%s, custom_llm_provider=%s. Error: %s", - model, - custom_llm_provider, - e, - ) - return None - - -def supports_embedding_image_input(model: str, custom_llm_provider: str | None = None) -> bool: - """ - Check if the given model supports embedding image input and return a boolean value. - """ - return _supports_factory( - model=model, - custom_llm_provider=custom_llm_provider, - key="supports_embedding_image_input", - ) - - -####### HELPER FUNCTIONS ################ -def _update_dictionary(existing_dict: dict, new_dict: dict) -> dict: - for k, v in new_dict.items(): - if v is not None: - # Convert stringified numbers to appropriate numeric types - if isinstance(v, str): - existing_dict[k] = _convert_stringified_numbers(v) - elif isinstance(v, dict): - existing_nested_dict = existing_dict.get(k) - if isinstance(existing_nested_dict, dict): - existing_dict[k] = {**existing_nested_dict, **v} # mutable-ok: copy-on-write merge - else: - existing_dict[k] = dict(v) # mutable-ok: detached copy, never the caller's dict by reference - else: - existing_dict[k] = v - - return existing_dict - - -def _convert_stringified_numbers(value): - """Convert stringified numbers (including scientific notation) to appropriate numeric types.""" - if isinstance(value, str): - try: - # Try to convert to float first to handle scientific notation like "3e-07" - if "e" in value.lower() or "." in value: - return float(value) - # Try to convert to int for whole numbers like "8192" - else: - return int(value) - except (ValueError, TypeError): - # If conversion fails, return the original string - return value - return value - - -_BEDROCK_REGION_PREFIXES: Final = ( - "us.", - "eu.", - "apac.", - "jp.", - "au.", - "us-gov.", - "global.", - "ap-northeast-1.", -) - -_CACHE_PRICING_FIELDS: Final = ( - "cache_creation_input_token_cost", - "cache_creation_input_token_cost_above_1hr", - "cache_creation_input_token_cost_above_200k_tokens", - "cache_read_input_token_cost", - "cache_read_input_token_cost_above_200k_tokens", -) - - -def _resolve_builtin_model_cost_entry(key: str, provider: str) -> dict[str, object] | None: - """Best-effort lookup of a built-in ``model_cost`` entry for a custom key - whose shape ``get_model_info`` cannot resolve (repeated provider prefixes - like ``bedrock/bedrock/bedrock/us.anthropic.claude-sonnet-4-6`` or region - aliases). - - Returns a copy of the matching entry so the caller can inherit its defaults - (most importantly cache pricing) without mutating the shared built-in. - Returns ``None`` when no safe match exists. - """ - candidates: Final[list[str]] = [] - segments: Final = key.split("/") - idx = 0 - while idx < len(segments) - 1 and segments[idx] in LlmProvidersSet: - idx += 1 - candidates.append("/".join(segments[idx:])) - - base: Final = candidates[-1] if candidates else key - for region_prefix in _BEDROCK_REGION_PREFIXES: - if base.startswith(region_prefix): - candidates.append(base[len(region_prefix) :]) - - if provider: - stripped: Final = _strip_model_name(model=base, custom_llm_provider=provider) - if stripped != base: - candidates.append(stripped) - - for candidate in candidates: - entry = litellm.model_cost.get(candidate) - if entry is not None and entry.get("litellm_provider") is not None: - return dict(entry) - return None - - -def _get_builtin_model_info_for_registration(model: str) -> ModelInfo | None: - """Resolve ``model`` to its built-in cost-map entry for registration merging. - - Returns ``None`` when the lookup raises or when it resolved via a - fallback-generalization capability rule, detected as the resolved key missing - ``litellm.model_cost`` while matching a capability rule. A rule-derived entry - carries no pricing, so treating it as a hit would skip the built-in - cache-pricing inheritance for prefix-mangled keys. - """ - try: - info: Final = get_model_info(model=model) - except Exception: - return None - if info["key"] in litellm.model_cost: - return info - if match_capability_generalizations(info["key"]) is None: - return info - return None - - -_runtime_registered_model_cost: Final[dict[str, dict[str, object]]] = {} # mutable-ok: replayed on reload - - -class _LiveDeploymentReplay: - """Single-slot holder for the callback that rebuilds live router deployments. - - A class attribute rather than a module global so there is one writer and one - reader, and neither needs a ``global`` statement. - """ - - callback: Callable[[], None] | None = None - - -def set_live_deployment_replay(replay: Callable[[], None]) -> None: - """Install the callback that re-asserts live router deployments after a refresh. - - ``litellm.router`` installs this at import time. The seam exists because the - deployment metadata a refresh has to restore belongs to whichever Router - objects are alive at that moment, which this module cannot see, and importing - the router here would be circular. - """ - _LiveDeploymentReplay.callback = replay - - -def reapply_runtime_model_cost_registrations() -> None: - """Re-apply runtime model metadata on top of a freshly adopted cost map. - - Adopting a new catalog replaces ``litellm.model_cost`` wholesale, which on - its own discards everything registered at runtime: the deployment - ``model_info`` the Router registers from ``model_list``, and pricing - overrides passed to ``register_model``. Both are re-applied here so a price - data reload only updates pricing rather than erasing operator-supplied model - metadata. - - The two are restored differently, and the difference is what keeps this - bounded. Deployment metadata is re-derived from the routers that are alive - right now, so a deployment that has been deleted or repointed, and a router - that has been discarded, are simply not part of the rebuild; nothing has to - withdraw them and nothing accumulates. Only ``register_model`` calls that - have no such owner are recorded and replayed, and a registration describing - a single request opts out of even that. - """ - if _LiveDeploymentReplay.callback is not None: - _LiveDeploymentReplay.callback() - if _runtime_registered_model_cost: - register_model(model_cost=dict(_runtime_registered_model_cost)) # mutable-ok: snapshot, replay rewrites it - - -def register_model( - model_cost: str | dict, - *, - persist_across_reloads: bool = True, - warning_display_name: str | None = None, -): - """ - Register new / Override existing models (and their pricing) to specific providers. - Provide EITHER a model cost dictionary or a url to a hosted json blob - Example usage: - model_cost_dict = { - "gpt-4": { - "max_tokens": 8192, - "input_cost_per_token": 0.00003, - "output_cost_per_token": 0.00006, - "litellm_provider": "openai", - "mode": "chat" - }, - } - - ``persist_across_reloads`` controls whether the registration is replayed - when the cost map is refreshed. It defaults to True because a caller - registering a model is declaring durable intent. Pass False for a - registration that only describes one request, so it is dropped rather than - re-asserted over every future catalog. - - ``warning_display_name`` names the model in the missing-cache-pricing - warning instead of the registered key, for callers that register under an - opaque key (e.g. the router's hashed deployment ids). - """ - - loaded_model_cost = {} - if isinstance(model_cost, dict): - # Convert stringified numbers to appropriate numeric types - loaded_model_cost = model_cost - elif isinstance(model_cost, str): - loaded_model_cost = litellm.get_model_cost_map(url=model_cost) - - if persist_across_reloads: - _registrations: Final[Mapping[str, Mapping[str, object]]] = loaded_model_cost - for _registered_key, _registered_value in _registrations.items(): - _runtime_registered_model_cost[_registered_key] = dict(_registered_value) # mutable-ok: caller-owned - - _skip_get_model_info_providers: Final = PROVIDERS_THAT_AUTHENTICATE_ON_PROVIDER_INFO - - for key, value in loaded_model_cost.items(): - ## get model info ## - provider = value.get("litellm_provider", "") - _key_str = str(key) - if provider in _skip_get_model_info_providers or any( - _key_str.startswith(f"{p}/") for p in _skip_get_model_info_providers - ): - existing_model = litellm.model_cost.get(key, {}) - model_cost_key = key - else: - builtin_model_info = _get_builtin_model_info_for_registration(model=_key_str) - if builtin_model_info is not None: - existing_model = cast(dict, builtin_model_info) - model_cost_key = existing_model["key"] - else: - existing_model = {} - model_cost_key = key - builtin_entry = _resolve_builtin_model_cost_entry(key=_key_str, provider=provider) - if builtin_entry is not None: - for field in _CACHE_PRICING_FIELDS: - if value.get(field) is None and builtin_entry.get(field) is not None: - existing_model[field] = builtin_entry[field] - elif ( - value.get("cache_creation_input_token_cost") is None - and value.get("cache_read_input_token_cost") is None - and value.get("tiered_pricing") is None - and ( - value.get("input_cost_per_token") is not None or value.get("output_cost_per_token") is not None - ) - ): - verbose_logger.warning( - "register_model: model=%s has custom pricing but not in built-in cost map and no prefix/region variant matched; cache_creation_input_token_cost and cache_read_input_token_cost will default to 0 for this model (input/output cost tracking is unaffected). To track cache cost, add them to model_info", - warning_display_name or key, - ) - # ``get_model_info`` returns ``litellm_provider: None`` when the - # provider is unknown (e.g. custom deployments registered via - # ``Router.add_deployment``). Persisting that None into - # ``litellm.model_cost`` causes ``_check_provider_match`` to drop - # custom pricing on subsequent cost lookups. - if existing_model.get("litellm_provider") is None: - existing_model.pop("litellm_provider", None) - # Same pattern for cost fields (#30198): ``_get_model_info_helper`` - # synthesizes ``input_cost_per_token`` / ``output_cost_per_token`` - # = 0 when they are absent from the raw entry. Writing those zeros - # back flips a sparse entry from "no cost keys" (priced via name) - # to "cost keys = 0" (free), which makes - # ``_is_cost_explicitly_configured`` return True and silently - # disables budget enforcement on the next re-registration. - _raw_entry = litellm.model_cost.get(model_cost_key) - if _raw_entry is None: - _raw_entry = litellm.model_cost.get(key) - if _raw_entry is None: - _raw_entry = {} - for _cost_field in ("input_cost_per_token", "output_cost_per_token"): - if _cost_field not in _raw_entry and _cost_field not in value: - existing_model.pop(_cost_field, None) - ## override / add new keys to the existing model cost dictionary - updated_dictionary = _update_dictionary(existing_model, value) - litellm.model_cost.setdefault(model_cost_key, {}).update(updated_dictionary) - - # Invalidate case-insensitive lookup map since model_cost was modified - _invalidate_model_cost_lowercase_map() - - verbose_logger.debug("added/updated model=%s in litellm.model_cost: %s", model_cost_key, model_cost_key) - # add new model names to provider lists - if value.get("litellm_provider") == "openai": - if key not in litellm.open_ai_chat_completion_models: - litellm.open_ai_chat_completion_models.add(key) - elif value.get("litellm_provider") == "text-completion-openai": - if key not in litellm.open_ai_text_completion_models: - litellm.open_ai_text_completion_models.add(key) - elif value.get("litellm_provider") == "cohere": - if key not in litellm.cohere_models: - litellm.cohere_models.add(key) - elif value.get("litellm_provider") == "anthropic": - if key not in litellm.anthropic_models: - litellm.anthropic_models.add(key) - elif value.get("litellm_provider") == "openrouter": - split_string = key.split("/", 1) - if split_string[-1] not in litellm.openrouter_models: - litellm.openrouter_models.add(split_string[-1]) - elif value.get("litellm_provider") == "vercel_ai_gateway": - if key not in litellm.vercel_ai_gateway_models: - litellm.vercel_ai_gateway_models.add(key) - elif value.get("litellm_provider") == "vertex_ai-text-models": - if key not in litellm.vertex_text_models: - litellm.vertex_text_models.add(key) - elif value.get("litellm_provider") == "vertex_ai-code-text-models": - if key not in litellm.vertex_code_text_models: - litellm.vertex_code_text_models.add(key) - elif value.get("litellm_provider") == "vertex_ai-chat-models": - if key not in litellm.vertex_chat_models: - litellm.vertex_chat_models.add(key) - elif value.get("litellm_provider") == "vertex_ai-code-chat-models": - if key not in litellm.vertex_code_chat_models: - litellm.vertex_code_chat_models.add(key) - elif value.get("litellm_provider") == "ai21": - if key not in litellm.ai21_models: - litellm.ai21_models.add(key) - elif value.get("litellm_provider") == "nlp_cloud": - if key not in litellm.nlp_cloud_models: - litellm.nlp_cloud_models.add(key) - elif value.get("litellm_provider") == "aleph_alpha": - if key not in litellm.aleph_alpha_models: - litellm.aleph_alpha_models.add(key) - elif value.get("litellm_provider") == "bedrock": - if key not in litellm.bedrock_models: - litellm.bedrock_models.add(key) - elif value.get("litellm_provider") == "novita": - if key not in litellm.novita_models: - litellm.novita_models.add(key) - return model_cost - - -def _should_drop_param(k, additional_drop_params) -> bool: - if additional_drop_params is not None and isinstance(additional_drop_params, list) and k in additional_drop_params: - return True # allow user to drop specific params for a model - e.g. vllm - logit bias - - return False - - -def _get_non_default_params(passed_params: dict, default_params: dict, additional_drop_params: list | None) -> dict: - non_default_params: Final = {} - for k, v in passed_params.items(): - if ( - k in default_params - and v != default_params[k] - and _should_drop_param(k=k, additional_drop_params=additional_drop_params) is False - ): - non_default_params[k] = v - - return non_default_params - - -def get_optional_params_transcription( - model: str, - custom_llm_provider: str, - language: str | None = None, - prompt: str | None = None, - response_format: str | None = None, - temperature: int | None = None, - timestamp_granularities: list[Literal["word", "segment"]] | None = None, - drop_params: bool | None = None, - **kwargs, -): - from litellm.constants import OPENAI_TRANSCRIPTION_PARAMS - - # retrieve all parameters passed to the function - passed_params: Final = locals() - - passed_params.pop("OPENAI_TRANSCRIPTION_PARAMS") - custom_llm_provider = passed_params.pop("custom_llm_provider") - drop_params = passed_params.pop("drop_params") - special_params: Final[Mapping[str, object]] = passed_params.pop("kwargs") - for k, v in special_params.items(): - passed_params[k] = v - - default_params: Final = { - "language": None, - "prompt": None, - "response_format": None, - "temperature": None, # openai defaults this to 0 - "timestamp_granularities": None, - } - - non_default_params = {k: v for k, v in passed_params.items() if (k in default_params and v != default_params[k])} - optional_params = {} - - ## raise exception if non-default value passed for non-openai/azure embedding calls - def _check_valid_arg(supported_params): - if len(non_default_params.keys()) > 0: - keys: Final = list(non_default_params.keys()) - for k in keys: - if ( - drop_params is True or litellm.drop_params is True - ) and k not in supported_params: # drop the unsupported non-default values - non_default_params.pop(k, None) - elif k not in supported_params: - raise UnsupportedParamsError( - status_code=500, - message=f"Setting user/encoding format is not supported by {custom_llm_provider}. To drop it from the call, set `litellm.drop_params = True`.", - ) - return non_default_params - - provider_config: BaseAudioTranscriptionConfig | None = None - if custom_llm_provider is not None: - provider_config = ProviderConfigManager.get_provider_audio_transcription_config( - model=model, - provider=LlmProviders(custom_llm_provider), - ) - - if custom_llm_provider == "openai" or custom_llm_provider == "azure": - optional_params = non_default_params - elif custom_llm_provider == "groq": - supported_params = litellm.GroqSTTConfig().get_supported_openai_params_stt() - _check_valid_arg(supported_params=supported_params) - optional_params = litellm.GroqSTTConfig().map_openai_params_stt( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=drop_params if drop_params is not None else False, - ) - elif provider_config is not None: # custom audio transcription config - supported_params = provider_config.get_supported_openai_params(model=model) - _check_valid_arg(supported_params=supported_params) - optional_params = provider_config.map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=drop_params if drop_params is not None else False, - ) - - optional_params = add_provider_specific_params_to_optional_params( - optional_params=optional_params, - passed_params=passed_params, - custom_llm_provider=custom_llm_provider, - openai_params=OPENAI_TRANSCRIPTION_PARAMS, - additional_drop_params=kwargs.get("additional_drop_params", None), - ) - - return optional_params - - -def _map_openai_size_to_vertex_ai_aspect_ratio(size: str | None) -> str: - """Map OpenAI size parameter to Vertex AI aspectRatio.""" - if size is None: - return "1:1" - - # Map OpenAI size strings to Vertex AI aspect ratio strings - # Vertex AI accepts: "1:1", "9:16", "16:9", "4:3", "3:4" - size_to_aspect_ratio: Final = { - "256x256": "1:1", # Square - "512x512": "1:1", # Square - "1024x1024": "1:1", # Square (default) - "1792x1024": "16:9", # Landscape - "1024x1792": "9:16", # Portrait - } - return size_to_aspect_ratio.get(size, "1:1") # Default to square if size not recognized - - -def get_optional_params_image_gen( - model: str | None = None, - n: int | None = None, - quality: str | None = None, - response_format: str | None = None, - size: str | None = None, - style: str | None = None, - user: str | None = None, - imageConfig: dict | None = None, - custom_llm_provider: str | None = None, - additional_drop_params: list | None = None, - provider_config: BaseImageGenerationConfig | None = None, - drop_params: bool | None = None, - **kwargs, -): - # retrieve all parameters passed to the function - passed_params: Final = locals() - model = passed_params.pop("model", None) - custom_llm_provider = passed_params.pop("custom_llm_provider") - provider_config = passed_params.pop("provider_config", None) - drop_params = passed_params.pop("drop_params", None) - additional_drop_params = passed_params.pop("additional_drop_params", None) - special_params: Final[Mapping[str, object]] = passed_params.pop("kwargs") - for k, v in special_params.items(): - if ( - k.startswith("aws_") - and (custom_llm_provider != "bedrock" and custom_llm_provider != "sagemaker") - or k == "hf_model_name" - and custom_llm_provider != "sagemaker" - ): # allow dynamically setting boto3 init logic - continue - elif ( - k.startswith("vertex_") and custom_llm_provider != "vertex_ai" and custom_llm_provider != "vertex_ai_beta" - ): # allow dynamically setting vertex ai init logic - continue - passed_params[k] = v - - provider_supported_params: Final[tuple[str, ...]] = ( - tuple(provider_config.get_supported_openai_params(model=model or "")) if provider_config is not None else () - ) - default_params: Final = { - "n": None, - "quality": None, - "response_format": None, - "size": None, - "style": None, - "user": None, - "imageConfig": None, - "tools": None, - "web_search_options": None, - **{k: None for k in provider_supported_params}, - } - - non_default_params: Final = _get_non_default_params( - passed_params=passed_params, - default_params=default_params, - additional_drop_params=additional_drop_params, - ) - optional_params: dict[str, object] = {} - - ## raise exception if non-default value passed for non-openai/azure embedding calls - def _check_valid_arg(supported_params): - if len(non_default_params.keys()) > 0: - keys: Final = list(non_default_params.keys()) - for k in keys: - if ( - litellm.drop_params is True or drop_params is True - ) and k not in supported_params: # drop the unsupported non-default values - non_default_params.pop(k, None) - passed_params.pop(k, None) - elif k not in supported_params: - raise UnsupportedParamsError( - status_code=500, - message=f"Setting `{k}` is not supported by {custom_llm_provider}, {model}. To drop it from the call, set `litellm.drop_params = True`.", - ) - return non_default_params - - if provider_config is not None: - supported_params = provider_config.get_supported_openai_params(model=model or "") - _check_valid_arg(supported_params=supported_params) - optional_params = provider_config.map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model or "", - drop_params=drop_params if drop_params is not None else False, - ) - elif ( - custom_llm_provider == "openai" - or custom_llm_provider == "azure" - or custom_llm_provider in litellm.openai_compatible_providers - ): - optional_params = non_default_params - elif custom_llm_provider == "bedrock": - config_class: Final = litellm.BedrockImageGeneration.get_config_class(model=model) - supported_params = config_class.get_supported_openai_params(model=model) - _check_valid_arg(supported_params=supported_params) - optional_params = config_class.map_openai_params(non_default_params=non_default_params, optional_params={}) - elif custom_llm_provider == "vertex_ai": - supported_params = ["n", "size"] - """ - All params here: https://console.cloud.google.com/vertex-ai/publishers/google/model-garden/imagegeneration?project=adroit-crow-413218 - """ - _check_valid_arg(supported_params=supported_params) - if n is not None: - optional_params["sampleCount"] = int(n) - - # Map OpenAI size parameter to Vertex AI aspectRatio - if size is not None: - optional_params["aspectRatio"] = _map_openai_size_to_vertex_ai_aspect_ratio(size) - - openai_params: Final[list[str]] = ( - list(provider_supported_params) if provider_config is not None else list(default_params.keys()) - ) - - optional_params = add_provider_specific_params_to_optional_params( - optional_params=optional_params, - passed_params=passed_params, - custom_llm_provider=custom_llm_provider or "", - openai_params=openai_params, - additional_drop_params=additional_drop_params, - ) - # remove keys with None or empty dict/list values to avoid sending empty payloads - optional_params = { - k: v for k, v in optional_params.items() if v is not None and (not isinstance(v, (dict, list)) or len(v) > 0) - } - return optional_params - - -def get_optional_params_embeddings( - # 2 optional params - model: str, - user: str | None = None, - encoding_format: str | None = None, - dimensions: int | None = None, - custom_llm_provider="", - drop_params: bool | None = None, - additional_drop_params: list[str] | None = None, - allowed_openai_params: list[str] | None = None, - **kwargs, -): - # Lazy load get_supported_openai_params - get_supported_openai_params: Final = getattr(sys.modules[__name__], "get_supported_openai_params") - - # retrieve all parameters passed to the function - passed_params: Final = locals() - custom_llm_provider = passed_params.pop("custom_llm_provider", None) - special_params: Final = passed_params.pop("kwargs") - - drop_params = passed_params.pop("drop_params", None) - additional_drop_params = passed_params.pop("additional_drop_params", None) - allowed_openai_params = passed_params.pop("allowed_openai_params", None) or [] - # Remove function objects from passed_params to avoid JSON serialization errors - passed_params.pop("get_supported_openai_params", None) - - def _check_valid_arg(supported_params: list | None): - if supported_params is None: - return - unsupported_params: Final = {} - for k in non_default_params: - if k not in supported_params: - unsupported_params[k] = non_default_params[k] - if unsupported_params: - if litellm.drop_params is True or (drop_params is not None and drop_params is True): - pass - else: - raise UnsupportedParamsError( - status_code=500, - message=f"{custom_llm_provider} does not support parameters: {unsupported_params}, for model={model}. To drop these, set `litellm.drop_params=True` or for proxy:\n\n`litellm_settings:\n drop_params: true`\n", - ) - - non_default_params: Final = PreProcessNonDefaultParams.embedding_pre_process_non_default_params( - passed_params=passed_params, - special_params=special_params, - custom_llm_provider=custom_llm_provider, - additional_drop_params=additional_drop_params, - model=model, - ) - - provider_config: BaseEmbeddingConfig | None = None - - optional_params = {} - if custom_llm_provider is not None and custom_llm_provider in LlmProviders._member_map_.values(): - provider_config = ProviderConfigManager.get_provider_embedding_config( - model=model, - provider=LlmProviders(custom_llm_provider), - ) - - if provider_config is not None: - supported_params: list | None = provider_config.get_supported_openai_params(model=model) - _check_valid_arg(supported_params=supported_params) - optional_params = provider_config.map_openai_params( - non_default_params=non_default_params, - optional_params={}, - model=model, - drop_params=drop_params if drop_params is not None else False, - ) - # Provider-only params (e.g. Cohere input_type) are not in - # OPENAI_EMBEDDING_PARAMS, so embedding_pre_process drops them from - # non_default_params before map_openai_params. Restore only those extras - # from passed_params — skip OPENAI_EMBEDDING_PARAMS to avoid duplicating - # values already mapped (e.g. dimensions -> output_dimension). - if supported_params: - for param in supported_params: - if param in OPENAI_EMBEDDING_PARAMS: - continue - if param in passed_params and passed_params[param] is not None and param not in optional_params: - optional_params[param] = passed_params[param] - ## raise exception if non-default value passed for non-openai/azure embedding calls - elif custom_llm_provider == "openai": - # 'dimensions` is only supported in `text-embedding-3` and later models - if ( - model is not None - and "text-embedding-3" not in model - and "dimensions" in non_default_params - and "dimensions" not in (allowed_openai_params or []) - ): - # Honor drop_params (per-call) and litellm.drop_params (global) the same - # way `_check_valid_arg` does above. The raised error message itself - # tells users to set `drop_params=True`, so respect it here. - if litellm.drop_params is True or drop_params is True: - non_default_params.pop("dimensions", None) - else: - raise UnsupportedParamsError( - status_code=500, - message="Setting dimensions is not supported for OpenAI `text-embedding-3` and later models. To drop it from the call, set `litellm.drop_params = True`.", - ) - optional_params = non_default_params - elif custom_llm_provider == "triton": - supported_params = get_supported_openai_params( - model=model, - custom_llm_provider=custom_llm_provider, - request_type="embeddings", - ) - _check_valid_arg(supported_params=supported_params) - optional_params = litellm.TritonEmbeddingConfig().map_openai_params( - non_default_params=non_default_params, - optional_params={}, - model=model, - drop_params=drop_params if drop_params is not None else False, - ) - elif custom_llm_provider == "databricks": - supported_params = get_supported_openai_params( - model=model or "", - custom_llm_provider="databricks", - request_type="embeddings", - ) - _check_valid_arg(supported_params=supported_params) - optional_params = litellm.DatabricksEmbeddingConfig().map_openai_params( - non_default_params=non_default_params, optional_params={} - ) - - elif custom_llm_provider == "nvidia_nim": - supported_params = get_supported_openai_params( - model=model or "", - custom_llm_provider="nvidia_nim", - request_type="embeddings", - ) - _check_valid_arg(supported_params=supported_params) - optional_params = litellm.nvidiaNimEmbeddingConfig.map_openai_params( - non_default_params=non_default_params, optional_params={}, kwargs=kwargs - ) - elif custom_llm_provider == "vertex_ai" or custom_llm_provider == "gemini": - # OpenAI SDKs send encoding_format="float" by default; float lists are - # exactly what the vertex API returns, so the param is a no-op and the - # provider default is not rejected. Other values (e.g. "base64") stay - # on the unsupported-param path below. - if non_default_params.get("encoding_format") == "float": - non_default_params.pop("encoding_format") - supported_params = get_supported_openai_params( - model=model, - custom_llm_provider="vertex_ai", - request_type="embeddings", - ) - _check_valid_arg(supported_params=supported_params) - ( - optional_params, - kwargs, - ) = litellm.VertexAITextEmbeddingConfig().map_openai_params( - non_default_params=non_default_params, optional_params={}, kwargs=kwargs - ) - elif custom_llm_provider == "lm_studio": - supported_params = litellm.LmStudioEmbeddingConfig().get_supported_openai_params() - _check_valid_arg(supported_params=supported_params) - optional_params = litellm.LmStudioEmbeddingConfig().map_openai_params( - non_default_params=non_default_params, optional_params={} - ) - elif custom_llm_provider == "bedrock": - # if dimensions is in non_default_params -> pass it for model=bedrock/amazon.titan-embed-text-v2 - if "amazon.titan-embed-text-v1" in model: - object: ( - AmazonTitanG1Config - | AmazonTitanMultimodalEmbeddingG1Config - | AmazonTitanV2Config - | BedrockCohereEmbeddingConfig - | TwelveLabsMarengoEmbeddingConfig - | AmazonNovaEmbeddingConfig - ) = litellm.AmazonTitanG1Config() - elif "amazon.titan-embed-image-v1" in model: - object = litellm.AmazonTitanMultimodalEmbeddingG1Config() - elif "amazon.titan-embed-text-v2:0" in model: - object = litellm.AmazonTitanV2Config() - elif "cohere.embed" in model: - object = litellm.BedrockCohereEmbeddingConfig() - elif "twelvelabs" in model or "marengo" in model: - object = litellm.TwelveLabsMarengoEmbeddingConfig() - elif "nova" in model.lower(): - object = litellm.AmazonNovaEmbeddingConfig() - else: # unmapped model - supported_params = [] - _check_valid_arg(supported_params=supported_params) - final_params = {**kwargs} - return final_params - - supported_params = object.get_supported_openai_params() - _check_valid_arg(supported_params=supported_params) - optional_params = object.map_openai_params(non_default_params=non_default_params, optional_params={}) - elif custom_llm_provider == "mistral": - supported_params = get_supported_openai_params( - model=model, - custom_llm_provider="mistral", - request_type="embeddings", - ) - _check_valid_arg(supported_params=supported_params) - optional_params = litellm.MistralEmbeddingConfig().map_openai_params( - non_default_params=non_default_params, optional_params={} - ) - elif custom_llm_provider == "jina_ai": - supported_params = get_supported_openai_params( - model=model, - custom_llm_provider="jina_ai", - request_type="embeddings", - ) - _check_valid_arg(supported_params=supported_params) - optional_params = litellm.JinaAIEmbeddingConfig().map_openai_params( - non_default_params=non_default_params, - optional_params={}, - model=model, - drop_params=drop_params if drop_params is not None else False, - ) - elif custom_llm_provider == "voyage": - supported_params = get_supported_openai_params( - model=model, - custom_llm_provider="voyage", - request_type="embeddings", - ) - _check_valid_arg(supported_params=supported_params) - if litellm.VoyageContextualEmbeddingConfig.is_contextualized_embeddings(model): - optional_params = litellm.VoyageContextualEmbeddingConfig().map_openai_params( - non_default_params=non_default_params, - optional_params={}, - model=model, - drop_params=drop_params if drop_params is not None else False, - ) - elif litellm.VoyageMultimodalEmbeddingConfig.is_multimodal_embeddings(model): - optional_params = litellm.VoyageMultimodalEmbeddingConfig().map_openai_params( - non_default_params=non_default_params, - optional_params={}, - model=model, - drop_params=drop_params if drop_params is not None else False, - ) - else: - optional_params = litellm.VoyageEmbeddingConfig().map_openai_params( - non_default_params=non_default_params, - optional_params={}, - model=model, - drop_params=drop_params if drop_params is not None else False, - ) - final_params = {**optional_params, **kwargs} - return final_params - elif custom_llm_provider == "sap": - supported_params = get_supported_openai_params( - model=model, - custom_llm_provider="sap", - request_type="embeddings", - ) - _check_valid_arg(supported_params=supported_params) - optional_params = litellm.GenAIHubEmbeddingConfig().map_openai_params( - non_default_params=non_default_params, - optional_params={}, - model=model, - drop_params=drop_params if drop_params is not None else False, - ) - elif custom_llm_provider == "infinity": - supported_params = get_supported_openai_params( - model=model, - custom_llm_provider="infinity", - request_type="embeddings", - ) - _check_valid_arg(supported_params=supported_params) - optional_params = litellm.InfinityEmbeddingConfig().map_openai_params( - non_default_params=non_default_params, - optional_params={}, - model=model, - drop_params=drop_params if drop_params is not None else False, - ) - - final_params = {**optional_params, **kwargs} - return final_params - - elif custom_llm_provider == "fireworks_ai": - supported_params = get_supported_openai_params( - model=model, - custom_llm_provider="fireworks_ai", - request_type="embeddings", - ) - _check_valid_arg(supported_params=supported_params) - optional_params = litellm.FireworksAIEmbeddingConfig().map_openai_params( - non_default_params=non_default_params, optional_params={}, model=model - ) - elif custom_llm_provider == "sambanova": - supported_params = get_supported_openai_params( - model=model, - custom_llm_provider="sambanova", - request_type="embeddings", - ) - _check_valid_arg(supported_params=supported_params) - optional_params = litellm.SambaNovaEmbeddingConfig().map_openai_params( - non_default_params=non_default_params, - optional_params={}, - model=model, - drop_params=drop_params if drop_params is not None else False, - ) - elif custom_llm_provider == "ovhcloud": - supported_params = get_supported_openai_params( - model=model, - custom_llm_provider="ovhcloud", - request_type="embeddings", - ) - _check_valid_arg(supported_params=supported_params) - optional_params = litellm.OVHCloudEmbeddingConfig().map_openai_params( - non_default_params=non_default_params, - optional_params={}, - model=model, - drop_params=drop_params if drop_params is not None else False, - ) - - elif custom_llm_provider == "ollama": - if "dimensions" in non_default_params: - optional_params["dimensions"] = non_default_params.pop("dimensions") - if len(non_default_params.keys()) > 0: - if litellm.drop_params is True or drop_params is True: # drop the unsupported non-default values - keys = list(non_default_params.keys()) - for k in keys: - non_default_params.pop(k, None) - else: - raise UnsupportedParamsError( - status_code=500, - message=f"Setting {non_default_params} is not supported by {custom_llm_provider}. To drop it from the call, set `litellm.drop_params = True`.", - ) - elif ( - custom_llm_provider != "openai" - and custom_llm_provider != "azure" - and custom_llm_provider not in litellm.openai_compatible_providers - ): - if len(non_default_params.keys()) > 0: - if litellm.drop_params is True or drop_params is True: # drop the unsupported non-default values - keys = list(non_default_params.keys()) - for k in keys: - non_default_params.pop(k, None) - else: - raise UnsupportedParamsError( - status_code=500, - message=f"Setting {non_default_params} is not supported by {custom_llm_provider}. To drop it from the call, set `litellm.drop_params = True`.", - ) - else: - optional_params = non_default_params - else: - optional_params = non_default_params - - final_params = add_provider_specific_params_to_optional_params( - optional_params=optional_params, - passed_params=passed_params, - custom_llm_provider=custom_llm_provider, - openai_params=list(DEFAULT_EMBEDDING_PARAM_VALUES.keys()), - additional_drop_params=kwargs.get("additional_drop_params", None), - ) - - if "extra_body" in final_params and len(final_params["extra_body"]) == 0: - final_params.pop("extra_body", None) - - return final_params - - -def _remove_additional_properties(schema): - """ - clean out 'additionalProperties = False'. Causes vertexai/gemini OpenAI API Schema errors - https://github.com/langchain-ai/langchainjs/issues/5240 - - Relevant Issues: https://github.com/BerriAI/litellm/issues/6136, https://github.com/BerriAI/litellm/issues/6088 - """ - if isinstance(schema, dict): - # Remove the 'additionalProperties' key if it exists and is set to False - if "additionalProperties" in schema and schema["additionalProperties"] is False: - del schema["additionalProperties"] - - # Recursively process all dictionary values - for key, value in schema.items(): - _remove_additional_properties(value) - - elif isinstance(schema, list): - # Recursively process all items in the list - for item in schema: - _remove_additional_properties(item) - - return schema - - -def _remove_strict_from_schema(schema): - """ - Relevant Issues: https://github.com/BerriAI/litellm/issues/6136, https://github.com/BerriAI/litellm/issues/6088 - """ - if isinstance(schema, dict): - # Remove the 'additionalProperties' key if it exists and is set to False - if "strict" in schema: - del schema["strict"] - - # Recursively process all dictionary values - for key, value in schema.items(): - _remove_strict_from_schema(value) - - elif isinstance(schema, list): - # Recursively process all items in the list - for item in schema: - _remove_strict_from_schema(item) - - return schema - - -def _remove_json_schema_refs(schema, max_depth=10): - """ - Remove JSON schema reference fields like '$id' and '$schema' that can cause issues with some providers. - - These fields are used for schema validation but can cause problems when the schema references - are not accessible to the provider's validation system. - - Args: - schema: The schema object to clean (dict, list, or other) - max_depth: Maximum recursion depth to prevent infinite loops (default: 10) - - Relevant Issues: Mistral API grammar validation fails when schema contains $id and $schema references - """ - if max_depth <= 0: - return schema - - if isinstance(schema, dict): - # Remove JSON schema reference fields - schema.pop("$id", None) - schema.pop("$schema", None) - - # Recursively process all dictionary values - for key, value in schema.items(): - _remove_json_schema_refs(value, max_depth - 1) - - elif isinstance(schema, list): - # Recursively process all items in the list - for item in schema: - _remove_json_schema_refs(item, max_depth - 1) - - return schema - - -def _remove_unsupported_params(non_default_params: dict, supported_openai_params: list[str] | None) -> dict: - """ - Remove unsupported params from non_default_params - """ - remove_keys: Final = [] - if supported_openai_params is None: - return {} # no supported params, so no optional openai params to send - for param in non_default_params: - if param not in supported_openai_params: - remove_keys.append(param) - for key in remove_keys: - non_default_params.pop(key, None) - return non_default_params - - -def filter_out_litellm_params(kwargs: dict) -> dict: - """ - Filter out LiteLLM internal parameters from kwargs dict. - - Returns a new dict containing only non-LiteLLM parameters that should be - passed to external provider APIs. - - Args: - kwargs: Dictionary that may contain LiteLLM internal parameters - - Returns: - Dictionary with LiteLLM internal parameters filtered out - - Example: - >>> kwargs = {"query": "test", "shared_session": session_obj, "metadata": {}} - >>> filtered = filter_out_litellm_params(kwargs) - >>> # filtered = {"query": "test"} - """ - - return {key: value for key, value in kwargs.items() if key not in all_litellm_params} - - -def _provider_supports_vertex_params(custom_llm_provider: str) -> bool: - if custom_llm_provider in ("vertex_ai", "vertex_ai_beta"): - return True - try: - provider: Final = LlmProviders(custom_llm_provider) - except ValueError: - return False - provider_config: Final = ProviderConfigManager.get_provider_chat_config(model="", provider=provider) - return bool(getattr(provider_config, "supports_vertex_params", False)) - - -class PreProcessNonDefaultParams: - @staticmethod - def base_pre_process_non_default_params( - passed_params: dict, - special_params: dict, - custom_llm_provider: str, - additional_drop_params: list[str] | None, - default_param_values: dict, - additional_endpoint_specific_params: list[str], - ) -> dict: - for k, v in special_params.items(): - if k == "aws_bedrock_project_id": - # sent as a request header (read from litellm_params by the - # bedrock-mantle configs), never as a request body field - continue - if ( - k.startswith("aws_") - and (custom_llm_provider != "bedrock" and not custom_llm_provider.startswith("sagemaker")) - or k == "hf_model_name" - and custom_llm_provider != "sagemaker" - or k.startswith("vertex_") - and not _provider_supports_vertex_params(custom_llm_provider) - ): # allow dynamically setting boto3 init logic - continue - passed_params[k] = v - - # filter out those parameters that were passed with non-default values - non_default_params: Final = { - k: v - for k, v in passed_params.items() - if ( - k != "model" - and k != "custom_llm_provider" - and k != "api_version" - and k != "drop_params" - and k != "allowed_openai_params" - and k != "additional_drop_params" - and k not in additional_endpoint_specific_params - and k in default_param_values - and v != default_param_values[k] - and _should_drop_param(k=k, additional_drop_params=additional_drop_params) is False - ) - } - - return non_default_params - - @staticmethod - def embedding_pre_process_non_default_params( - passed_params: dict, - special_params: dict, - custom_llm_provider: str, - additional_drop_params: list[str] | None, - model: str, - remove_sensitive_keys: bool = False, - add_provider_specific_params: bool = False, - ) -> dict: - non_default_params: Final = PreProcessNonDefaultParams.base_pre_process_non_default_params( - passed_params=passed_params, - special_params=special_params, - custom_llm_provider=custom_llm_provider, - additional_drop_params=additional_drop_params, - default_param_values={k: None for k in OPENAI_EMBEDDING_PARAMS}, - additional_endpoint_specific_params=["input"], - ) - - return non_default_params - - -def pre_process_non_default_params( - passed_params: dict, - special_params: dict, - custom_llm_provider: str, - additional_drop_params: list[str] | None, - model: str, - remove_sensitive_keys: bool = False, - add_provider_specific_params: bool = False, - provider_config: BaseConfig | None = None, -) -> dict: - """ - Pre-process non-default params to a standardized format - """ - # retrieve all parameters passed to the function - - non_default_params = PreProcessNonDefaultParams.base_pre_process_non_default_params( - passed_params=passed_params, - special_params=special_params, - custom_llm_provider=custom_llm_provider, - additional_drop_params=additional_drop_params, - default_param_values=DEFAULT_CHAT_COMPLETION_PARAM_VALUES, - additional_endpoint_specific_params=["messages"], - ) - - if "response_format" in non_default_params: - if provider_config is not None: - non_default_params["response_format"] = provider_config.get_json_schema_from_pydantic_object( - response_format=non_default_params["response_format"] - ) - else: - non_default_params["response_format"] = type_to_response_format_param( - response_format=non_default_params["response_format"] - ) - - if "tools" in non_default_params and isinstance( - non_default_params, list - ): # fixes https://github.com/BerriAI/litellm/issues/4933 - tools: Final = non_default_params["tools"] - for tool in tools: # clean out 'additionalProperties = False'. Causes vertexai/gemini OpenAI API Schema errors - https://github.com/langchain-ai/langchainjs/issues/5240 - tool_function = tool.get("function", {}) - parameters = tool_function.get("parameters", None) - if parameters is not None: - new_parameters = copy.deepcopy(parameters) - if "additionalProperties" in new_parameters and new_parameters["additionalProperties"] is False: - new_parameters.pop("additionalProperties", None) - tool_function["parameters"] = new_parameters - - if add_provider_specific_params: - non_default_params = add_provider_specific_params_to_optional_params( - optional_params=non_default_params, - passed_params=passed_params, - custom_llm_provider=custom_llm_provider, - openai_params=list(DEFAULT_CHAT_COMPLETION_PARAM_VALUES.keys()), - additional_drop_params=additional_drop_params, - ) - - if remove_sensitive_keys: - non_default_params = remove_sensitive_keys_from_dict(non_default_params) - return non_default_params - - -def remove_sensitive_keys_from_dict(d: dict) -> dict: - """ - Remove sensitive keys from a dictionary - """ - sensitive_key_phrases: Final = ["key", "secret", "access", "credential"] - remove_keys: Final = [] - for key in d: - if any(phrase in key.lower() for phrase in sensitive_key_phrases): - remove_keys.append(key) - for key in remove_keys: - d.pop(key) - return d - - -def pre_process_optional_params(passed_params: dict, non_default_params: dict, custom_llm_provider: str) -> dict: - """For .completion(), preprocess optional params""" - optional_params: dict = {} - - common_auth_dict: Final = litellm.common_cloud_provider_auth_params - if custom_llm_provider in common_auth_dict["providers"]: - """ - Check if params = ["project", "region_name", "token"] - and correctly translate for = ["azure", "vertex_ai", "watsonx", "aws"] - """ - if custom_llm_provider == "azure": - optional_params = litellm.AzureOpenAIConfig().map_special_auth_params( - non_default_params=passed_params, optional_params=optional_params - ) - elif custom_llm_provider == "bedrock": - optional_params = litellm.AmazonBedrockGlobalConfig().map_special_auth_params( - non_default_params=passed_params, optional_params=optional_params - ) - elif custom_llm_provider == "vertex_ai" or custom_llm_provider == "vertex_ai_beta": - optional_params = litellm.VertexAIConfig().map_special_auth_params( - non_default_params=passed_params, optional_params=optional_params - ) - elif custom_llm_provider == "watsonx": - optional_params = litellm.IBMWatsonXAIConfig().map_special_auth_params( - non_default_params=passed_params, optional_params=optional_params - ) - - ## raise exception if function calling passed in for a provider that doesn't support it - if "functions" in non_default_params or "function_call" in non_default_params or "tools" in non_default_params: - if ( - custom_llm_provider == "ollama" - and custom_llm_provider != "text-completion-openai" - and custom_llm_provider != "azure" - and custom_llm_provider != "vertex_ai" - and custom_llm_provider != "anyscale" - and custom_llm_provider != "together_ai" - and custom_llm_provider != "groq" - and custom_llm_provider != "nvidia_nim" - and custom_llm_provider != "cerebras" - and custom_llm_provider != "xai" - and custom_llm_provider != "ai21_chat" - and custom_llm_provider != "volcengine" - and custom_llm_provider != "deepseek" - and custom_llm_provider != "codestral" - and custom_llm_provider != "mistral" - and custom_llm_provider != "anthropic" - and custom_llm_provider != "cohere_chat" - and custom_llm_provider != "cohere" - and custom_llm_provider != "bedrock" - and custom_llm_provider != "ollama_chat" - and custom_llm_provider != "openrouter" - and custom_llm_provider != "vercel_ai_gateway" - and custom_llm_provider != "nebius" - and custom_llm_provider != "wandb" - and custom_llm_provider not in litellm.openai_compatible_providers - ): - if custom_llm_provider == "ollama": - # ollama actually supports json output - optional_params["format"] = "json" - litellm.add_function_to_prompt = True # so that main.py adds the function call to the prompt - if "tools" in non_default_params: - optional_params["functions_unsupported_model"] = non_default_params.pop("tools") - non_default_params.pop("tool_choice", None) # causes ollama requests to hang - elif "functions" in non_default_params: - optional_params["functions_unsupported_model"] = non_default_params.pop("functions") - elif litellm.add_function_to_prompt: # if user opts to add it to prompt instead - optional_params["functions_unsupported_model"] = non_default_params.pop( - "tools", non_default_params.pop("functions", None) - ) - else: - raise UnsupportedParamsError( - status_code=500, - message=f"Function calling is not supported by {custom_llm_provider}.", - ) - - return optional_params - - -def get_optional_params( - # use the openai defaults - # https://platform.openai.com/docs/api-reference/chat/create - model: str, - functions=None, - function_call=None, - temperature=None, - top_p=None, - n=None, - stream=False, - stream_options=None, - stop=None, - max_tokens=None, - max_completion_tokens=None, - modalities=None, - prediction=None, - audio=None, - presence_penalty=None, - frequency_penalty=None, - logit_bias=None, - user=None, - custom_llm_provider="", - response_format=None, - seed=None, - tools=None, - tool_choice=None, - max_retries=None, - logprobs=None, - top_logprobs=None, - extra_headers=None, - api_version=None, - parallel_tool_calls=None, - drop_params=None, - allowed_openai_params: list[str] | None = None, - reasoning_effort=None, - verbosity=None, - additional_drop_params=None, - messages: list[AllMessageValues] | None = None, - thinking: AnthropicThinkingParam | None = None, - web_search_options: OpenAIWebSearchOptions | None = None, - safety_identifier: str | None = None, - store: bool | None = None, - prompt_cache_key: str | None = None, - base_model: str | None = None, - **kwargs, -): - passed_params: Final = locals().copy() - special_params: Final = passed_params.pop("kwargs") - # Remove base_model from passed_params so it doesn't interfere with - # non_default_params / _check_valid_arg — it's a routing hint, not an - # OpenAI param. - passed_params.pop("base_model", None) - provider_config: BaseConfig | None = None - if custom_llm_provider is not None and custom_llm_provider in [provider.value for provider in LlmProviders]: - provider_config = ProviderConfigManager.get_provider_chat_config( - model=model, - provider=LlmProviders(custom_llm_provider), - base_model=base_model, - ) - non_default_params: Final = pre_process_non_default_params( - passed_params=passed_params, - special_params=special_params, - custom_llm_provider=custom_llm_provider, - additional_drop_params=additional_drop_params, - model=model, - provider_config=provider_config, - ) - optional_params = pre_process_optional_params( - passed_params=passed_params, - non_default_params=non_default_params, - custom_llm_provider=custom_llm_provider, - ) - - def _check_valid_arg(supported_params: list[str]): - """ - Check if the params passed to completion() are supported by the provider - - Args: - supported_params: List[str] - supported params from the litellm config - """ - verbose_logger.info("\nLiteLLM completion() model= %s; provider = %s", model, custom_llm_provider) - verbose_logger.debug("\nLiteLLM: Params passed to completion() %s", passed_params) - verbose_logger.debug("\nLiteLLM: Non-Default params passed to completion() %s", non_default_params) - unsupported_params: Final = {} - for k in non_default_params: - if k not in supported_params: - if k in PROVIDER_UNVALIDATED_PARAMS: - continue - if k == "n" and n == 1: # langchain sends n=1 as a default value - continue # skip this param - unsupported_params[k] = non_default_params[k] - - if unsupported_params: - if litellm.drop_params is True or (drop_params is not None and drop_params is True): - for k in unsupported_params: - non_default_params.pop(k, None) - else: - raise UnsupportedParamsError( - status_code=500, - message=f"{custom_llm_provider} does not support parameters: {list(unsupported_params.keys())}, for model={model}. To drop these, set `litellm.drop_params=True` or for proxy:\n\n`litellm_settings:\n drop_params: true`\n. \n If you want to use these params dynamically send allowed_openai_params={list(unsupported_params.keys())} in your request.", - ) - - get_supported_openai_params: Final = getattr(sys.modules[__name__], "get_supported_openai_params") - supported_params = get_supported_openai_params( - model=model, custom_llm_provider=custom_llm_provider, base_model=base_model - ) - if supported_params is None: - supported_params = get_supported_openai_params(model=model, custom_llm_provider="openai") - - supported_params = supported_params or [] - allowed_openai_params = allowed_openai_params or [] - supported_params.extend(allowed_openai_params) - - _check_valid_arg( - supported_params=supported_params or [], - ) - ## raise exception if provider doesn't support passed in param - if custom_llm_provider == "anthropic": - ## check if unsupported param passed in - optional_params = litellm.AnthropicConfig().map_openai_params( - model=model, - non_default_params=non_default_params, - optional_params=optional_params, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "anthropic_text": - optional_params = litellm.AnthropicTextConfig().map_openai_params( - model=model, - non_default_params=non_default_params, - optional_params=optional_params, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - optional_params = litellm.AnthropicTextConfig().map_openai_params( - model=model, - non_default_params=non_default_params, - optional_params=optional_params, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - - elif custom_llm_provider == "cohere_chat" or custom_llm_provider == "cohere": - # handle cohere params - optional_params = litellm.CohereChatConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "triton": - optional_params = litellm.TritonConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=drop_params if drop_params is not None else False, - ) - - elif custom_llm_provider == "maritalk": - optional_params = litellm.MaritalkConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "replicate": - optional_params = litellm.ReplicateConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "predibase": - optional_params = litellm.PredibaseConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "huggingface": - optional_params = litellm.HuggingFaceChatConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "together_ai": - optional_params = litellm.TogetherAIChatConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "vertex_ai" and ( - model in litellm.vertex_chat_models - or model in litellm.vertex_code_chat_models - or model in litellm.vertex_text_models - or model in litellm.vertex_code_text_models - or model in litellm.vertex_language_models - or model in litellm.vertex_vision_models - ): - optional_params = litellm.VertexGeminiConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - - elif custom_llm_provider == "gemini": - optional_params = litellm.GoogleAIStudioGeminiConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "vertex_ai_beta" or (custom_llm_provider == "vertex_ai" and "gemini" in model): - optional_params = litellm.VertexGeminiConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif litellm.VertexAIAnthropicConfig.is_supported_model(model=model, custom_llm_provider=custom_llm_provider): - optional_params = litellm.VertexAIAnthropicConfig().map_openai_params( - model=model, - non_default_params=non_default_params, - optional_params=optional_params, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "vertex_ai": - if model in litellm.vertex_mistral_models: - if "codestral" in model: - optional_params = litellm.CodestralTextCompletionConfig().map_openai_params( - model=model, - non_default_params=non_default_params, - optional_params=optional_params, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - else: - optional_params = litellm.MistralConfig().map_openai_params( - model=model, - non_default_params=non_default_params, - optional_params=optional_params, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif model in litellm.vertex_ai_ai21_models: - optional_params = litellm.VertexAIAi21Config().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif provider_config is not None: - optional_params = provider_config.map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - else: # use generic openai-like param mapping - optional_params = litellm.VertexAILlama3Config().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - - elif custom_llm_provider == "sagemaker": - # temperature, top_p, n, stream, stop, max_tokens, n, presence_penalty default to None - optional_params = litellm.SagemakerConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "bedrock": - BedrockModelInfo: Final = getattr(sys.modules[__name__], "BedrockModelInfo") - bedrock_route: Final = BedrockModelInfo.get_bedrock_route(model) - bedrock_base_model: Final = BedrockModelInfo.get_base_model(model) - if bedrock_route == "converse" or bedrock_route == "converse_like": - optional_params = litellm.AmazonConverseConfig().map_openai_params( - model=model, - non_default_params=non_default_params, - optional_params=optional_params, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif bedrock_route == "openai": - optional_params = litellm.AmazonBedrockOpenAIConfig().map_openai_params( - model=model, - non_default_params=non_default_params, - optional_params=optional_params, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif "anthropic" in bedrock_base_model and bedrock_route == "invoke": - if bedrock_base_model in litellm.AmazonAnthropicConfig.get_legacy_anthropic_model_names(): - optional_params = litellm.AmazonAnthropicConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - else: - optional_params = litellm.AmazonAnthropicClaudeConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif provider_config is not None: - optional_params = provider_config.map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - if bedrock_route == "claude_platform": - optional_params = BedrockModelInfo.map_claude_platform_auth_params( - passed_params=passed_params, optional_params=optional_params - ) - elif custom_llm_provider == "cloudflare": - optional_params = litellm.CloudflareChatConfig().map_openai_params( - model=model, - non_default_params=non_default_params, - optional_params=optional_params, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "ollama": - optional_params = litellm.OllamaConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "ollama_chat": - optional_params = litellm.OllamaChatConfig().map_openai_params( - model=model, - non_default_params=non_default_params, - optional_params=optional_params, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "nlp_cloud": - optional_params = litellm.NLPCloudConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - - elif custom_llm_provider == "petals": - optional_params = litellm.PetalsConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "deepinfra": - optional_params = litellm.DeepInfraConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "perplexity" and provider_config is not None: - optional_params = provider_config.map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "mistral" or custom_llm_provider == "codestral": - optional_params = litellm.MistralConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "text-completion-codestral": - optional_params = litellm.CodestralTextCompletionConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - - elif custom_llm_provider == "text-completion-inception": - optional_params = litellm.InceptionTextCompletionConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - - elif custom_llm_provider == "databricks": - optional_params = litellm.DatabricksConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "nvidia_nim": - optional_params = litellm.NvidiaNimConfig().map_openai_params( - model=model, - non_default_params=non_default_params, - optional_params=optional_params, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "cerebras": - optional_params = litellm.CerebrasConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "xai": - optional_params = litellm.XAIChatConfig().map_openai_params( - model=model, - non_default_params=non_default_params, - optional_params=optional_params, - ) - elif custom_llm_provider == "ai21_chat" or custom_llm_provider == "ai21": - optional_params = litellm.AI21ChatConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "fireworks_ai": - optional_params = litellm.FireworksAIConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "volcengine": - optional_params = litellm.VolcEngineConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "hosted_vllm": - optional_params = litellm.HostedVLLMChatConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "vllm": - optional_params = litellm.VLLMConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "groq": - optional_params = litellm.GroqChatConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "bedrock_mantle": - optional_params = litellm.BedrockMantleChatConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "deepseek": - optional_params = litellm.DeepSeekChatConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "tencent": - optional_params = litellm.TencentChatConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "openrouter": - optional_params = litellm.OpenrouterConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "watsonx": - optional_params = litellm.IBMWatsonXChatConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - # WatsonX-text param check - for param in passed_params: - if litellm.IBMWatsonXAIConfig().is_watsonx_text_param(param): - raise ValueError( - f"LiteLLM now defaults to Watsonx's `/text/chat` endpoint. Please use the `watsonx_text` provider instead, to call the `/text/generation` endpoint. Param: {param}" - ) - elif custom_llm_provider == "watsonx_text": - optional_params = litellm.IBMWatsonXAIConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "openai": - optional_params = litellm.OpenAIConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "nebius": - optional_params = litellm.NebiusConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif custom_llm_provider == "azure": - _azure_detection_model: Final = base_model or model - if litellm.AzureOpenAIO1Config().is_o_series_model(model=_azure_detection_model): - optional_params = litellm.AzureOpenAIO1Config().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=_azure_detection_model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model(model=_azure_detection_model): - optional_params = litellm.AzureOpenAIGPT5Config().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=_azure_detection_model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - else: - verbose_logger.debug( - "Azure optional params - api_version: api_version={}, litellm.api_version={}, os.environ['AZURE_API_VERSION']={}".format( - api_version, litellm.api_version, get_secret("AZURE_API_VERSION") - ) - ) - api_version = ( - api_version - or litellm.api_version - or get_secret("AZURE_API_VERSION") - or litellm.AZURE_DEFAULT_API_VERSION - ) - optional_params = litellm.AzureOpenAIConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=_azure_detection_model, - api_version=api_version, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - elif provider_config is not None: - optional_params = provider_config.map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - else: # assume passing in params for openai-like api - optional_params = litellm.OpenAILikeChatConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), - ) - # if user passed in non-default kwargs for specific providers/models, pass them along - optional_params = add_provider_specific_params_to_optional_params( - optional_params=optional_params, - passed_params=passed_params, - custom_llm_provider=custom_llm_provider, - openai_params=list(DEFAULT_CHAT_COMPLETION_PARAM_VALUES.keys()), - additional_drop_params=additional_drop_params, - ) - if _print_verbose_is_active(): - print_verbose(f"Final returned optional params: {redact_credentials_in_payload(optional_params)}") - optional_params = _apply_openai_param_overrides( - optional_params=optional_params, - non_default_params=non_default_params, - allowed_openai_params=allowed_openai_params, - ) - - # Apply nested drops from additional_drop_params - if additional_drop_params: - is_nested_path: Final = getattr(sys.modules[__name__], "is_nested_path") - delete_nested_value: Final = getattr(sys.modules[__name__], "delete_nested_value") - nested_paths: Final = [p for p in additional_drop_params if is_nested_path(p)] - for path in nested_paths: - optional_params = delete_nested_value(optional_params, path) - - return optional_params - - -def add_provider_specific_params_to_optional_params( - optional_params: dict, - passed_params: dict, - custom_llm_provider: str, - openai_params: list[str], - additional_drop_params: list | None = None, -) -> dict: - """ - Add provider specific params to optional_params - """ - - if custom_llm_provider in ["openai", "azure", "text-completion-openai"] + litellm.openai_compatible_providers: - # for openai, azure we should pass the extra/passed params within `extra_body` https://github.com/openai/openai-python/blob/ac33853ba10d13ac149b1fa3ca6dba7d613065c9/src/openai/resources/models.py#L46 - if _should_drop_param(k="extra_body", additional_drop_params=additional_drop_params) is False: - extra_body: Final = dict(passed_params.pop("extra_body", None) or {}) - for k in passed_params: - if k not in openai_params and passed_params[k] is not None: - extra_body[k] = passed_params[k] - if not isinstance(optional_params.get("extra_body"), dict): - optional_params["extra_body"] = {} - initial_extra_body: Final = { - **optional_params["extra_body"], - **extra_body, - } - - if additional_drop_params is not None: - processed_extra_body = {k: v for k, v in initial_extra_body.items() if k not in additional_drop_params} - else: - processed_extra_body = initial_extra_body - - _ensure_extra_body_is_safe: Final = getattr(sys.modules[__name__], "_ensure_extra_body_is_safe") - optional_params["extra_body"] = _ensure_extra_body_is_safe(extra_body=processed_extra_body) - else: - for k in passed_params: - if k not in openai_params and passed_params[k] is not None: - if _should_drop_param(k=k, additional_drop_params=additional_drop_params): - continue - optional_params[k] = passed_params[k] - return optional_params - - -def _apply_openai_param_overrides(optional_params: dict, non_default_params: dict, allowed_openai_params: list): - """ - If user passes in allowed_openai_params, apply them to optional_params - - These params will get passed as is to the LLM API since the user opted in to passing them in the request - - Only params the caller actually sent are forwarded. Previously this - function unconditionally wrote `None` for any allowed param missing from - the request, which then reached the provider SDK as a top-level kwarg it - did not recognize (e.g. openai SDK raising - `AsyncCompletions.create() got an unexpected keyword argument 'enable_thinking'`). - See https://github.com/BerriAI/litellm/issues/25697 - """ - if allowed_openai_params: - for param in allowed_openai_params: - if param in optional_params: - continue - if param not in non_default_params: - continue - optional_params[param] = non_default_params.pop(param) - return optional_params - - -PROVIDER_UNVALIDATED_PARAMS: Final = frozenset({"user", "stream_options", "stream", "max_retries"}) - - -def provider_rejectable_params(passed_params: Mapping[str, object]) -> frozenset[str]: - """The params a provider can actually be rejected for, i.e. the ones _check_valid_arg compares - against its supported list. - - Anything outside this set never reaches that comparison. Endpoint and transport controls such as - base_url, timeout, default_headers, organization and deployment_id are not chat completion - params at all, so a caller filtering on "is this an OpenAI param" would discard configuration the - request needs while never touching what the provider would have rejected. - """ - params: Final = dict(passed_params) # mutable-ok: get_non_default_params takes a dict - return frozenset(get_non_default_params(params)) - PROVIDER_UNVALIDATED_PARAMS - - -def get_non_default_params(passed_params: dict) -> dict: - # filter out those parameters that were passed with non-default values - non_default_params: Final = { - k: v - for k, v in passed_params.items() - if ( - k != "model" - and k != "custom_llm_provider" - and k in DEFAULT_CHAT_COMPLETION_PARAM_VALUES - and v != DEFAULT_CHAT_COMPLETION_PARAM_VALUES[k] - ) - } - - return non_default_params - - -def calculate_max_parallel_requests( - max_parallel_requests: int | None, - rpm: int | None, - tpm: int | None, - default_max_parallel_requests: int | None, -) -> int | None: - """ - Returns the max parallel requests to send to a deployment. - - Used in semaphore for async requests on router. - - Parameters: - - max_parallel_requests - Optional[int] - max_parallel_requests allowed for that deployment - - rpm - Optional[int] - requests per minute allowed for that deployment - - tpm - Optional[int] - tokens per minute allowed for that deployment - - default_max_parallel_requests - Optional[int] - default_max_parallel_requests allowed for any deployment - - Returns: - - int or None (if all params are None) - - Order: - max_parallel_requests > rpm > tpm / 6 (azure formula) > default max_parallel_requests - - Azure RPM formula: - 6 rpm per 1000 TPM - https://learn.microsoft.com/en-us/azure/ai-services/openai/quotas-limits - - - """ - if max_parallel_requests is not None: - return max_parallel_requests - elif rpm is not None: - return rpm - elif tpm is not None: - calculated_rpm = int(tpm / 1000 * 6) - if calculated_rpm == 0: - calculated_rpm = 1 - return calculated_rpm - elif default_max_parallel_requests is not None: - return default_max_parallel_requests - return None - - -def _get_deployment_order(deployment: dict | Any) -> int | None: - """ - Returns the routing order for a deployment. - - Checks litellm_params first (static config), then model_info (dynamic/team - models added via API where order lives in model_info, not litellm_params). - """ - order = deployment.get("litellm_params", {}).get("order") - if order is None: - order = deployment.get("model_info", {}).get("order") - return order - - -def _get_order_filtered_deployments(healthy_deployments: list[dict], target_order: int | None = None) -> list: - if target_order is not None: - return [d for d in healthy_deployments if _get_deployment_order(d) == target_order] - - # Default: pick min order group - _valid_orders: Final[list[int]] = [ - o for deployment in healthy_deployments for o in [_get_deployment_order(deployment)] if o is not None - ] - min_order: Final[int | None] = min(_valid_orders) if _valid_orders else None - - if min_order is not None: - filtered_deployments: Final = [ - deployment for deployment in healthy_deployments if _get_deployment_order(deployment) == min_order - ] - - return filtered_deployments - return healthy_deployments - - -def _get_excluded_filtered_deployments( - healthy_deployments: list[dict], - excluded_deployment_ids: Iterable[str] | None = None, -) -> list: - """ - Filter out deployments whose `model_info.id` appears in `excluded_deployment_ids`. - - Used by weighted-routing failover so a single logical request can re-pick - across the remaining deployments in the same model group after one of them - has failed. - - If the filter would leave no deployments, an empty list is returned so the - caller raises its usual no-deployments error and the weighted-failover - helper falls through to the cross-group fallback path. Returning the - original unfiltered list here would re-include the just-failed deployment. - """ - if not excluded_deployment_ids: - return healthy_deployments - - excluded_set: Final = set(excluded_deployment_ids) - return [d for d in healthy_deployments if (d.get("model_info") or {}).get("id") not in excluded_set] - - -def _get_model_region(custom_llm_provider: str, litellm_params: LiteLLM_Params) -> str | None: - """ - Return the region for a model, for a given provider - """ - if custom_llm_provider == "vertex_ai": - # check 'vertex_location' - vertex_ai_location: Final = ( - litellm_params.vertex_location - or litellm.vertex_location - or get_secret("VERTEXAI_LOCATION") - or get_secret("VERTEX_LOCATION") - ) - if vertex_ai_location is not None and isinstance(vertex_ai_location, str): - return vertex_ai_location - elif custom_llm_provider == "bedrock": - aws_region_name: Final = litellm_params.aws_region_name - if aws_region_name is not None: - return aws_region_name - elif custom_llm_provider == "watsonx": - watsonx_region_name: Final = litellm_params.watsonx_region_name - if watsonx_region_name is not None: - return watsonx_region_name - return litellm_params.region_name - - -def _infer_model_region(litellm_params: LiteLLM_Params) -> AllowedModelRegion | None: - """ - Infer if a model is in the EU or US region - - Returns: - - str (region) - "eu" or "us" - - None (if region not found) - """ - model, custom_llm_provider, _, _ = litellm.get_llm_provider( - model=litellm_params.model, litellm_params=litellm_params - ) - - model_region: Final = _get_model_region(custom_llm_provider=custom_llm_provider, litellm_params=litellm_params) - - if model_region is None: - verbose_logger.debug("Cannot infer model region for model: %s", litellm_params.model) - return None - - if custom_llm_provider == "azure": - eu_regions = litellm.AzureOpenAIConfig().get_eu_regions() - us_regions = litellm.AzureOpenAIConfig().get_us_regions() - elif custom_llm_provider == "vertex_ai": - eu_regions = litellm.VertexAIConfig().get_eu_regions() - us_regions = litellm.VertexAIConfig().get_us_regions() - elif custom_llm_provider == "bedrock": - eu_regions = litellm.AmazonBedrockGlobalConfig().get_eu_regions() - us_regions = litellm.AmazonBedrockGlobalConfig().get_us_regions() - elif custom_llm_provider == "watsonx": - eu_regions = litellm.IBMWatsonXAIConfig().get_eu_regions() - us_regions = litellm.IBMWatsonXAIConfig().get_us_regions() - else: - eu_regions = [] - us_regions = [] - - for region in eu_regions: - if region in model_region.lower(): - return "eu" - for region in us_regions: - if region in model_region.lower(): - return "us" - return None - - -def _is_region_eu(litellm_params: LiteLLM_Params) -> bool: - """ - Return true/false if a deployment is in the EU - """ - if litellm_params.region_name == "eu": - return True - - ## Else - try and infer from model region - model_region: Final = _infer_model_region(litellm_params=litellm_params) - if model_region is not None and model_region == "eu": - return True - return False - - -def _is_region_us(litellm_params: LiteLLM_Params) -> bool: - """ - Return true/false if a deployment is in the US - """ - if litellm_params.region_name == "us": - return True - - ## Else - try and infer from model region - model_region: Final = _infer_model_region(litellm_params=litellm_params) - if model_region is not None and model_region == "us": - return True - return False - - -def is_region_allowed(litellm_params: LiteLLM_Params, allowed_model_region: str) -> bool: - """ - Return true/false if a deployment is in the EU - """ - if litellm_params.region_name == allowed_model_region: - return True - return False - - -def get_model_region(litellm_params: LiteLLM_Params, mode: str | None) -> str | None: - """ - Pass the litellm params for an azure model, and get back the region - """ - if ( - "azure" in litellm_params.model - and isinstance(litellm_params.api_key, str) - and isinstance(litellm_params.api_base, str) - ): - _model: Final = litellm_params.model.replace("azure/", "") - response: Final[dict] = litellm.AzureChatCompletion().get_headers( - model=_model, - api_key=litellm_params.api_key, - api_base=litellm_params.api_base, - api_version=litellm_params.api_version or litellm.AZURE_DEFAULT_API_VERSION, - timeout=10, - mode=mode or "chat", - ) - - region: Final[str | None] = response.get("x-ms-region", None) - return region - return None - - -def get_first_chars_messages(kwargs: dict) -> str: - try: - _messages = kwargs.get("messages") - _messages = str(_messages)[:100] - return _messages - except Exception: - return "" - - -def _count_characters(text: str) -> int: - # Remove white spaces and count characters - filtered_text: Final = "".join(char for char in text if not char.isspace()) - return len(filtered_text) - - -def get_response_string(response_obj: ModelResponse | ModelResponseStream) -> str: - # Handle Responses API streaming events - if hasattr(response_obj, "type") and hasattr(response_obj, "response"): - # This is a Responses API streaming event (e.g., ResponseCreatedEvent, ResponseCompletedEvent) - # Extract text from the response object's output if available - responses_api_response: Final = getattr(response_obj, "response", None) - if responses_api_response and hasattr(responses_api_response, "output"): - output_list: Final = responses_api_response.output - # Use list accumulation to avoid O(n^2) string concatenation: - # repeatedly doing `response_str += part` copies the full string each time - # because Python strings are immutable, so total work grows with n^2. - response_output_parts: Final[list[str]] = [] - for output_item in output_list: - # Handle output items with content array - if hasattr(output_item, "content"): - for content_part in output_item.content: - if hasattr(content_part, "text"): - response_output_parts.append(content_part.text) - # Handle output items with direct text field - elif hasattr(output_item, "text"): - response_output_parts.append(output_item.text) - return "".join(response_output_parts) - - # Handle Responses API text delta events - if hasattr(response_obj, "type") and hasattr(response_obj, "delta"): - event_type: Final = getattr(response_obj, "type", "") - if "text.delta" in event_type or "output_text.delta" in event_type: - delta: Final = getattr(response_obj, "delta", "") - return delta if isinstance(delta, str) else "" - - # Handle standard ModelResponse and ModelResponseStream - _choices: Final[list[Choices] | list[StreamingChoices]] = response_obj.choices - - # Use list accumulation to avoid O(n^2) string concatenation across choices - response_parts: Final[list[str]] = [] - for choice in _choices: - if isinstance(choice, Choices): - if choice.message.content is not None: - response_parts.append(str(choice.message.content)) - elif isinstance(choice, StreamingChoices): - if choice.delta.content is not None: - response_parts.append(str(choice.delta.content)) - - return "".join(response_parts) - - -def get_utc_datetime() -> datetime.datetime: - return datetime.datetime.now(datetime.timezone.utc) - - -def get_max_tokens(model: str) -> int | None: - """ - Get the maximum number of output tokens allowed for a given model. - - Parameters: - model (str): The name of the model. - - Returns: - int: The maximum number of tokens allowed for the given model. - - Raises: - Exception: If the model is not mapped yet. - - Example: - >>> get_max_tokens("gpt-4") - 8192 - """ - - def _get_max_position_embeddings(model_name): - # Construct the URL for the config.json file - config_url: Final = f"https://huggingface.co/{model_name}/raw/main/config.json" - try: - # Make the HTTP request to get the raw JSON file - response: Final = litellm.module_level_client.get(config_url, timeout=HF_CONFIG_FETCH_TIMEOUT_SECONDS) - response.raise_for_status() # Raise an exception for bad responses (4xx or 5xx) - - # Parse the JSON response - config_json: Final[Mapping[str, int]] = response.json() - # Extract and return the max_position_embeddings - max_position_embeddings: Final = config_json.get("max_position_embeddings") - if max_position_embeddings is not None: - return max_position_embeddings - else: - return None - except Exception: - return None - - try: - if model in litellm.model_cost: - if "max_output_tokens" in litellm.model_cost[model]: - return litellm.model_cost[model]["max_output_tokens"] - elif "max_tokens" in litellm.model_cost[model]: - return litellm.model_cost[model]["max_tokens"] - get_llm_provider: Final = litellm_utils.get_llm_provider - model, custom_llm_provider, _, _ = get_llm_provider(model=model) - if custom_llm_provider == "huggingface": - max_tokens: Final = _get_max_position_embeddings(model_name=model) - return max_tokens - if model in litellm.model_cost: # check if extracted model is in model_list - if "max_output_tokens" in litellm.model_cost[model]: - return litellm.model_cost[model]["max_output_tokens"] - elif "max_tokens" in litellm.model_cost[model]: - return litellm.model_cost[model]["max_tokens"] - else: - raise Exception() - return None - except Exception: - raise Exception( - f"Model {model} isn't mapped yet. Add it here - https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json" - ) - - -def _strip_stable_vertex_version(model_name) -> str: - return re.sub(r"-\d+$", "", model_name) - - -def _get_base_bedrock_model(model_name) -> str: - """ - Get the base model from the given model name. - - Handle model names like - "us.meta.llama3-2-11b-instruct-v1:0" -> "meta.llama3-2-11b-instruct-v1" - AND "meta.llama3-2-11b-instruct-v1:0" -> "meta.llama3-2-11b-instruct-v1" - """ - from litellm.llms.bedrock.common_utils import BedrockModelInfo - - return BedrockModelInfo.get_base_model(model_name) - - -def _strip_openai_finetune_model_name(model_name: str) -> str: - """ - Strips the organization, custom suffix, and ID from an OpenAI fine-tuned model name. - - input: ft:gpt-3.5-turbo:my-org:custom_suffix:id - output: ft:gpt-3.5-turbo - - Args: - model_name (str): The full model name - - Returns: - str: The stripped model name - """ - return re.sub(r"(:[^:]+){3}$", "", model_name) - - -def _strip_model_name(model: str, custom_llm_provider: str | None) -> str: - if custom_llm_provider and custom_llm_provider in ["bedrock", "bedrock_converse"]: - stripped_bedrock_model: Final = _get_base_bedrock_model(model_name=model) - return stripped_bedrock_model - elif ( - custom_llm_provider - and (custom_llm_provider == "vertex_ai" or custom_llm_provider == "gemini") - or custom_llm_provider - and (custom_llm_provider == "databricks") - ): - strip_version: Final = _strip_stable_vertex_version(model_name=model) - return strip_version - elif "ft:" in model: - strip_finetune: Final = _strip_openai_finetune_model_name(model_name=model) - return strip_finetune - else: - return model - - -# Global case-insensitive lookup map for model_cost (built eagerly at module import) -_model_cost_lowercase_map: dict[str, str] | None = None - -# Monotonic counter bumped on every model_cost mutation. Consumers that -# memoize derived state (e.g. provider-specific indices) can include this -# value in their cache key so they invalidate even when key add+remove or -# in-place value replacement leaves len/id unchanged. -_model_cost_mutation_generation: int = 0 - - -def get_model_cost_mutation_generation() -> int: - return _model_cost_mutation_generation - - -def _invalidate_model_cost_lowercase_map() -> None: - """Invalidate the case-insensitive lookup map for model_cost. - - Call this whenever litellm.model_cost is modified to ensure the map is rebuilt. - Also clears related LRU caches that depend on model_cost data. - """ - global _model_cost_lowercase_map, _model_cost_mutation_generation - _model_cost_lowercase_map = None - _model_cost_mutation_generation += 1 - - # Clear LRU caches that depend on model_cost data - _cached_get_model_info.cache_clear() - _cached_get_model_info_helper.cache_clear() - - -def _rebuild_model_cost_lowercase_map() -> dict[str, str]: - """Rebuild the case-insensitive lookup map from the current model_cost. - - Returns: - The rebuilt map (guaranteed to be not None). - """ - global _model_cost_lowercase_map - _model_cost_lowercase_map = {k.lower(): k for k in litellm.model_cost} - return _model_cost_lowercase_map - - -def _handle_stale_map_entry_rebuild( - potential_key_lower: str, -) -> str | None: - """ - Handle stale _model_cost_lowercase_map entry (key was popped). - - Rebuilds the map and retries the lookup. - - Returns: - The matched key if found after rebuild, None otherwise. - """ - global _model_cost_lowercase_map - _model_cost_lowercase_map = _rebuild_model_cost_lowercase_map() - matched_key: Final = _model_cost_lowercase_map.get(potential_key_lower) - if matched_key is not None and matched_key in litellm.model_cost: - return matched_key - return None - - -def _handle_new_key_with_scan( - potential_key_lower: str, -) -> str | None: - """ - Handle new key added to model_cost without invalidating _model_cost_lowercase_map. - - Scans model_cost for case-insensitive match and rebuilds the map if found. - - Returns: - The matched key if found, None otherwise. - """ - global _model_cost_lowercase_map - for key in litellm.model_cost: - if key.lower() == potential_key_lower: - _model_cost_lowercase_map = _rebuild_model_cost_lowercase_map() - return key - return None - - -def _get_model_cost_key(potential_key: str) -> str | None: - """ - Get the actual key from model_cost, with case-insensitive fallback. - - WARNING: Only O(1) lookup operations are acceptable. O(n) lookups will cause severe - CPU overhead. This function is called frequently during router operations. - - ALLOWED HELPER FUNCTIONS (conditionally called, O(n) operations are acceptable): - - _rebuild_model_cost_lowercase_map: Rebuilds the lookup map (only when map is None) - - _handle_stale_map_entry_rebuild: Rebuilds map when stale entry detected (rare case) - - If you need to add a new helper function with O(n) operations that is conditionally - called and confirmed not to cause performance issues, add it to the allowed_helpers - list in: tests/code_coverage_tests/check_get_model_cost_key_performance.py - """ - global _model_cost_lowercase_map - - # Exact match (O(1)) - if potential_key in litellm.model_cost: - return potential_key - - # Case-insensitive lookup via map (O(1)) - if _model_cost_lowercase_map is None: - _model_cost_lowercase_map = _rebuild_model_cost_lowercase_map() - - potential_key_lower: Final = potential_key.lower() - matched_key = _model_cost_lowercase_map.get(potential_key_lower) - - # Verify key exists (O(1) - handles model_cost.pop() case) - if matched_key is not None and matched_key in litellm.model_cost: - return matched_key - - # Rebuild map if stale entry detected (O(n) rebuild, but only when stale entry found) - if matched_key is not None: - matched_key = _handle_stale_map_entry_rebuild(potential_key_lower) - if matched_key is not None: - return matched_key - - return None - - -def _get_model_info_from_model_cost(key: str) -> dict: - return litellm.model_cost[key] - - -def _check_provider_match(model_info: dict, custom_llm_provider: str | None) -> bool: - """ - Check if the model info provider matches the custom provider. - - A missing ``litellm_provider`` key and a ``litellm_provider`` set to - ``None`` both mean "no specific provider constraint" and are treated - as a wildcard match. ``register_model`` may persist ``None`` here via - ``get_model_info`` when a deployment is registered without a provider, - so normalising the two cases keeps custom pricing applied consistently. - """ - if custom_llm_provider and ( - model_info.get("litellm_provider") is not None and model_info["litellm_provider"] != custom_llm_provider - ): - if ( - custom_llm_provider == "vertex_ai" - and model_info["litellm_provider"].startswith("vertex_ai") - or custom_llm_provider == "fireworks_ai" - and model_info["litellm_provider"].startswith("fireworks_ai") - or custom_llm_provider.startswith("bedrock") - and model_info["litellm_provider"].startswith("bedrock") - ): - return True - elif ( - custom_llm_provider == "litellm_proxy" - ): # litellm_proxy is a special case, it's not a provider, it's a proxy for the provider - return True - elif custom_llm_provider == "azure_ai" and model_info["litellm_provider"] in ( - "azure", - "openai", - ): - # Azure AI also works with azure models - # as a last attempt if the model is not on Azure AI, Azure then fallback to OpenAI cost - # tracking the cost is better than attributing 0 cost to it. - return True - elif custom_llm_provider == "github": - # Allow github/ aliases to reuse existing provider metadata. - return True - else: - return False - - return True - - -from typing_extensions import ReadOnly, TypedDict - - -class PotentialModelNamesAndCustomLLMProvider(TypedDict): - split_model: str - combined_model_name: str - stripped_model_name: str - combined_stripped_model_name: str - provider_prefixed_model_name: ReadOnly[str] - custom_llm_provider: str - - -def _get_model_info_from_generalization( - model: str, - potential_model_names: PotentialModelNamesAndCustomLLMProvider, - custom_llm_provider: str | None, -) -> tuple[str, dict] | None: - """Resolve an unmapped model via the declarative capability generalization rules. - - Tries the same name candidates as the exact lookups, in the same order, and - returns ``(matched_name, model_info)`` for the first candidate matched by at - least one capability rule, with ``litellm_provider`` backfilled from the - provider the caller requested. Rules lose to exact entries: if ANY candidate is - an exact ``litellm.model_cost`` key (necessarily provider-mismatched, or the - exact lookups would have returned it), the model is known rather than unmapped, - and resolving it from rules would hand an unpriced rule-derived entry to - callers whose fallback ladder (e.g. the cost calculator's model-name variants) - still had a priced exact name to try. O(number of rules); only call after the - exact lookups have missed. - """ - candidates: Final = ( - potential_model_names["combined_model_name"], - model, - potential_model_names["split_model"], - potential_model_names["combined_stripped_model_name"], - potential_model_names["stripped_model_name"], - potential_model_names["provider_prefixed_model_name"], - ) - if any(_get_model_cost_key(candidate) is not None for candidate in candidates): - return None - for candidate in candidates: - generalized_info = match_capability_generalizations(candidate) - if generalized_info is None: - continue - if custom_llm_provider is None: - return candidate, generalized_info - return candidate, {**generalized_info, "litellm_provider": custom_llm_provider} - return None - - -def _get_potential_model_names(model: str, custom_llm_provider: str | None) -> PotentialModelNamesAndCustomLLMProvider: - if custom_llm_provider is None: - # Get custom_llm_provider - try: - get_llm_provider: Final = litellm_utils.get_llm_provider - split_model, custom_llm_provider, _, _ = get_llm_provider(model=model) - except Exception: - split_model = model - combined_model_name = model - stripped_model_name = _strip_model_name(model=model, custom_llm_provider=custom_llm_provider) - combined_stripped_model_name = stripped_model_name - provider_prefixed_model_name = model - elif custom_llm_provider and model.startswith( - custom_llm_provider + "/" - ): # handle case where custom_llm_provider is provided and model starts with custom_llm_provider - split_model = model.split("/", 1)[1] - combined_model_name = model - stripped_model_name = _strip_model_name(model=split_model, custom_llm_provider=custom_llm_provider) - combined_stripped_model_name = f"{custom_llm_provider}/{stripped_model_name}" - provider_prefixed_model_name = f"{custom_llm_provider}/{model}" - else: - split_model = model - combined_model_name = f"{custom_llm_provider}/{model}" - stripped_model_name = _strip_model_name(model=model, custom_llm_provider=custom_llm_provider) - combined_stripped_model_name = f"{custom_llm_provider}/{stripped_model_name}" - provider_prefixed_model_name = combined_model_name - - if custom_llm_provider in ("bedrock", "bedrock_converse"): - from litellm.llms.bedrock.common_utils import strip_bedrock_routing_prefix - - split_model = strip_bedrock_routing_prefix(split_model) - - return PotentialModelNamesAndCustomLLMProvider( - split_model=split_model, - combined_model_name=combined_model_name, - stripped_model_name=stripped_model_name, - combined_stripped_model_name=combined_stripped_model_name, - provider_prefixed_model_name=provider_prefixed_model_name, - custom_llm_provider=cast(str, custom_llm_provider), - ) - - -def _get_max_position_embeddings(model_name: str) -> int | None: - # Construct the URL for the config.json file - config_url: Final = f"https://huggingface.co/{model_name}/raw/main/config.json" - - try: - # Make the HTTP request to get the raw JSON file - response: Final = litellm.module_level_client.get(config_url, timeout=HF_CONFIG_FETCH_TIMEOUT_SECONDS) - response.raise_for_status() # Raise an exception for bad responses (4xx or 5xx) - - # Parse the JSON response - config_json: Final[Mapping[str, int]] = response.json() - - # Extract and return the max_position_embeddings - max_position_embeddings: Final = config_json.get("max_position_embeddings") - - if max_position_embeddings is not None: - return max_position_embeddings - else: - return None - except Exception: - return None - - -@lru_cache(maxsize=DEFAULT_MAX_LRU_CACHE_SIZE) -def _cached_get_model_info_helper( - model: str, - custom_llm_provider: str | None, - api_base: str | None = None, -) -> ModelInfoBase: - """ - _get_model_info_helper wrapped with lru_cache - - Speed Optimization to hit high RPS - """ - return _get_model_info_helper( - model=model, - custom_llm_provider=custom_llm_provider, - api_base=api_base, - ) - - -def get_provider_info(model: str, custom_llm_provider: str | None) -> ProviderSpecificModelInfo | None: - ## PROVIDER-SPECIFIC INFORMATION - # if custom_llm_provider == "predibase": - # _model_info["supports_response_schema"] = True - provider_config: BaseLLMModelInfo | None = None - if custom_llm_provider and custom_llm_provider in LlmProvidersSet: - # Check if the provider string exists in LlmProviders enum - provider_config = ProviderConfigManager.get_provider_model_info( - model=model, provider=LlmProviders(custom_llm_provider) - ) - - model_info: ProviderSpecificModelInfo | None = None - if provider_config: - model_info = provider_config.get_provider_info(model=model) - - return model_info - - -def _is_potential_model_name_in_model_cost( - potential_model_names: PotentialModelNamesAndCustomLLMProvider, -) -> bool: - """ - Check if the potential model name is in the model cost (case-insensitive). - """ - return any( - _get_model_cost_key(str(potential_model_name)) is not None - for potential_model_name in potential_model_names.values() - ) - - -_ABOVE_THRESHOLD_COST_KEY: Final = re.compile(r"_above_\d+k?_tokens$") - - -def _get_model_info_helper( - model: str, - custom_llm_provider: str | None = None, - api_base: str | None = None, - api_key: str | None = None, -) -> ModelInfoBase: - """ - Helper for 'get_model_info'. Separated out to avoid infinite loop caused by returning 'supported_openai_param's - """ - from litellm.litellm_core_utils.get_llm_provider_logic import declared_authenticating_provider - - try: - azure_llms: Final = {**litellm.azure_llms, **litellm.azure_embedding_models} - if model in azure_llms: - model = azure_llms[model] - if custom_llm_provider is not None and custom_llm_provider == "vertex_ai_beta": - custom_llm_provider = "vertex_ai" - if custom_llm_provider is not None and custom_llm_provider == "vertex_ai": - if "meta/" + model in litellm.vertex_llama3_models: - model = "meta/" + model - elif ( - model + "@latest" in litellm.vertex_mistral_models or model + "@latest" in litellm.vertex_ai_ai21_models - ): - model = model + "@latest" - ########################## - potential_model_names: Final = _get_potential_model_names( - model=model, custom_llm_provider=custom_llm_provider or declared_authenticating_provider(model) - ) - - verbose_logger.debug("checking potential_model_names in litellm.model_cost: %s", potential_model_names) - - combined_model_name: Final = potential_model_names["combined_model_name"] - stripped_model_name: Final = potential_model_names["stripped_model_name"] - combined_stripped_model_name: Final = potential_model_names["combined_stripped_model_name"] - provider_prefixed_model_name: Final = potential_model_names["provider_prefixed_model_name"] - split_model: Final = potential_model_names["split_model"] - custom_llm_provider = potential_model_names["custom_llm_provider"] - model_cost_custom_llm_provider: Final = custom_llm_provider - ######################### - provider_config: BaseLLMModelInfo | None = None - if custom_llm_provider and custom_llm_provider in LlmProvidersSet: - provider_config = ProviderConfigManager.get_provider_model_info( - model=model, provider=LlmProviders(custom_llm_provider) - ) - if provider_config is not None: - provider_get_model_info: Final = getattr(provider_config, "get_model_info", None) - if callable(provider_get_model_info): - try: - provider_model_info: Final = provider_get_model_info( - model=model, - api_base=api_base, - api_key=api_key, - ) - if provider_model_info is not None: - return provider_model_info - except Exception as e: - verbose_logger.warning( - "Could not get dynamic model info for model=%s, provider=%s; " - "falling back to the static cost map: %s", - model, - custom_llm_provider, - e, - ) - - if custom_llm_provider == "huggingface": - max_tokens: Final = _get_max_position_embeddings(model_name=model) - return ModelInfoBase( - key=model, - max_tokens=max_tokens, - max_input_tokens=None, - max_output_tokens=None, - input_cost_per_token=0, - output_cost_per_token=0, - litellm_provider="huggingface", - mode="chat", - supports_system_messages=None, - supports_response_schema=None, - supports_function_calling=None, - supports_tool_choice=None, - supports_assistant_prefill=None, - supports_prompt_caching=None, - supports_prompt_cache_breakpoint=None, - supports_computer_use=None, - supports_pdf_input=None, - ) - else: - """ - Check if: (in order of specificity) - 1. 'custom_llm_provider/model' in litellm.model_cost. Checks "groq/llama3-8b-8192" if model="llama3-8b-8192" and custom_llm_provider="groq" - 2. 'model' in litellm.model_cost. Checks "gemini-1.5-pro-002" in litellm.model_cost if model="gemini-1.5-pro-002" and custom_llm_provider=None - 3. 'split_model' in litellm.model_cost. Checks "au.anthropic.claude-opus-4-8" in litellm.model_cost if model="bedrock/au.anthropic.claude-opus-4-8" - 4. 'combined_stripped_model_name' in litellm.model_cost. Checks if 'gemini/gemini-1.5-flash' in model map, if 'gemini/gemini-1.5-flash-001' given. - 5. 'stripped_model_name' in litellm.model_cost. Checks if 'ft:gpt-3.5-turbo' in model map, if 'ft:gpt-3.5-turbo:my-org:custom_suffix:id' given. - 6. 'provider_prefixed_model_name' in litellm.model_cost, for providers whose own model ids repeat the - litellm provider name. Checks "perplexity/perplexity/glm-5.2" if model="perplexity/glm-5.2" and - custom_llm_provider="perplexity", where 1-5 all read the leading "perplexity/" as the litellm prefix - and strip it. Tried last so no model that already resolves through 1-5 can change. - """ - - _model_info: dict[str, Any] | None = None - key: str | None = None - - # Use case-insensitive lookup for all model name checks - _matched_key = _get_model_cost_key(combined_model_name) - if _matched_key is not None: - key = _matched_key - _model_info = _get_model_info_from_model_cost(key=cast(str, key)) - if not _check_provider_match( - model_info=_model_info, - custom_llm_provider=model_cost_custom_llm_provider, - ): - _model_info = None - if _model_info is None: - _matched_key = _get_model_cost_key(model) - if _matched_key is not None: - key = _matched_key - _model_info = _get_model_info_from_model_cost(key=cast(str, key)) - if not _check_provider_match( - model_info=_model_info, - custom_llm_provider=model_cost_custom_llm_provider, - ): - _model_info = None - if _model_info is None: - _matched_key = _get_model_cost_key(split_model) - if _matched_key is not None: - key = _matched_key - _model_info = _get_model_info_from_model_cost(key=cast(str, key)) - if not _check_provider_match( - model_info=_model_info, - custom_llm_provider=model_cost_custom_llm_provider, - ): - _model_info = None - if _model_info is None: - _matched_key = _get_model_cost_key(combined_stripped_model_name) - if _matched_key is not None: - key = _matched_key - _model_info = _get_model_info_from_model_cost(key=cast(str, key)) - if not _check_provider_match( - model_info=_model_info, - custom_llm_provider=model_cost_custom_llm_provider, - ): - _model_info = None - if _model_info is None: - _matched_key = _get_model_cost_key(stripped_model_name) - if _matched_key is not None: - key = _matched_key - _model_info = _get_model_info_from_model_cost(key=cast(str, key)) - if not _check_provider_match( - model_info=_model_info, - custom_llm_provider=model_cost_custom_llm_provider, - ): - _model_info = None - if _model_info is None: - _matched_key = _get_model_cost_key(provider_prefixed_model_name) - if _matched_key is not None: - key = _matched_key - _model_info = _get_model_info_from_model_cost(key=cast(str, key)) - if not _check_provider_match( - model_info=_model_info, - custom_llm_provider=model_cost_custom_llm_provider, - ): - _model_info = None - - if _model_info is None: - generalization: Final = _get_model_info_from_generalization( - model=model, - potential_model_names=potential_model_names, - custom_llm_provider=custom_llm_provider, - ) - if generalization is not None: - key, _model_info = generalization - - if _model_info is None or key is None: - raise ValueError( - "This model isn't mapped yet. Add it here - https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json" - ) - _input_cost_per_token: float | None = _model_info.get("input_cost_per_token") - if _input_cost_per_token is None: - # default value to 0, be noisy about this - verbose_logger.debug( - "model=%s, custom_llm_provider=%s has no input_cost_per_token in model_cost_map. Defaulting to 0.", - model, - custom_llm_provider, - ) - _input_cost_per_token = 0 - - _output_cost_per_token: float | None = _model_info.get("output_cost_per_token") - if _output_cost_per_token is None: - # default value to 0, be noisy about this - verbose_logger.debug( - "model=%s, custom_llm_provider=%s has no output_cost_per_token in model_cost_map. Defaulting to 0.", - model, - custom_llm_provider, - ) - _output_cost_per_token = 0 - - returned_model_info: Final = ModelInfoBase( - key=key, - max_tokens=_model_info.get("max_tokens", None), - max_input_tokens=_model_info.get("max_input_tokens", None), - max_output_tokens=_model_info.get("max_output_tokens", None), - input_cost_per_token=_input_cost_per_token, - input_cost_per_token_flex=_model_info.get("input_cost_per_token_flex", None), - input_cost_per_token_priority=_model_info.get("input_cost_per_token_priority", None), - input_cost_per_token_ultrafast=_model_info.get("input_cost_per_token_ultrafast", None), - cache_creation_input_token_cost=_model_info.get("cache_creation_input_token_cost", None), - cache_creation_input_token_cost_above_200k_tokens=_model_info.get( - "cache_creation_input_token_cost_above_200k_tokens", None - ), - cache_creation_input_token_cost_above_272k_tokens=_model_info.get( - "cache_creation_input_token_cost_above_272k_tokens", None - ), - cache_creation_input_token_cost_above_272k_tokens_priority=_model_info.get( - "cache_creation_input_token_cost_above_272k_tokens_priority", None - ), - cache_creation_input_token_cost_above_272k_tokens_flex=_model_info.get( - "cache_creation_input_token_cost_above_272k_tokens_flex", None - ), - cache_creation_input_token_cost_flex=_model_info.get("cache_creation_input_token_cost_flex", None), - cache_creation_input_token_cost_priority=_model_info.get( - "cache_creation_input_token_cost_priority", None - ), - cache_creation_input_token_cost_ultrafast=_model_info.get( - "cache_creation_input_token_cost_ultrafast", None - ), - cache_read_input_token_cost=_model_info.get("cache_read_input_token_cost", None), - prompt_cache_min_tokens=_model_info.get("prompt_cache_min_tokens", None), - cache_read_input_token_cost_above_200k_tokens=_model_info.get( - "cache_read_input_token_cost_above_200k_tokens", None - ), - cache_read_input_token_cost_above_200k_tokens_priority=_model_info.get( - "cache_read_input_token_cost_above_200k_tokens_priority", None - ), - cache_read_input_token_cost_above_272k_tokens=_model_info.get( - "cache_read_input_token_cost_above_272k_tokens", None - ), - cache_read_input_token_cost_above_272k_tokens_priority=_model_info.get( - "cache_read_input_token_cost_above_272k_tokens_priority", None - ), - cache_read_input_token_cost_above_272k_tokens_flex=_model_info.get( - "cache_read_input_token_cost_above_272k_tokens_flex", None - ), - cache_read_input_token_cost_above_512k_tokens=_model_info.get( - "cache_read_input_token_cost_above_512k_tokens", None - ), - cache_read_input_token_cost_flex=_model_info.get("cache_read_input_token_cost_flex", None), - cache_read_input_token_cost_priority=_model_info.get("cache_read_input_token_cost_priority", None), - cache_read_input_token_cost_ultrafast=_model_info.get("cache_read_input_token_cost_ultrafast", None), - cache_creation_input_token_cost_above_1hr=_model_info.get( - "cache_creation_input_token_cost_above_1hr", None - ), - off_peak_pricing=_model_info.get("off_peak_pricing", None), - input_cost_per_character=_model_info.get("input_cost_per_character", None), - input_cost_per_token_above_128k_tokens=_model_info.get("input_cost_per_token_above_128k_tokens", None), - input_cost_per_token_above_200k_tokens=_model_info.get("input_cost_per_token_above_200k_tokens", None), - input_cost_per_token_above_200k_tokens_priority=_model_info.get( - "input_cost_per_token_above_200k_tokens_priority", None - ), - input_cost_per_token_above_272k_tokens=_model_info.get("input_cost_per_token_above_272k_tokens", None), - input_cost_per_token_above_272k_tokens_priority=_model_info.get( - "input_cost_per_token_above_272k_tokens_priority", None - ), - input_cost_per_token_above_272k_tokens_flex=_model_info.get( - "input_cost_per_token_above_272k_tokens_flex", None - ), - input_cost_per_token_above_512k_tokens=_model_info.get("input_cost_per_token_above_512k_tokens", None), - input_cost_per_query=_model_info.get("input_cost_per_query", None), - input_cost_per_second=_model_info.get("input_cost_per_second", None), - input_cost_per_audio_token=_model_info.get("input_cost_per_audio_token", None), - input_cost_per_image_token=_model_info.get("input_cost_per_image_token", None), - input_cost_per_video_token=_model_info.get("input_cost_per_video_token", None), - input_cost_per_image=_model_info.get("input_cost_per_image", None), - input_cost_per_audio_per_second=_model_info.get("input_cost_per_audio_per_second", None), - input_cost_per_video_per_second=_model_info.get("input_cost_per_video_per_second", None), - input_cost_per_token_batches=_model_info.get("input_cost_per_token_batches"), - output_cost_per_token_batches=_model_info.get("output_cost_per_token_batches"), - output_cost_per_token=_output_cost_per_token, - output_cost_per_token_flex=_model_info.get("output_cost_per_token_flex", None), - output_cost_per_token_priority=_model_info.get("output_cost_per_token_priority", None), - output_cost_per_token_ultrafast=_model_info.get("output_cost_per_token_ultrafast", None), - regional_processing_uplift_multiplier_eu=_model_info.get( - "regional_processing_uplift_multiplier_eu", None - ), - regional_processing_uplift_multiplier_us=_model_info.get( - "regional_processing_uplift_multiplier_us", None - ), - regional_endpoint_uplift_multiplier=_model_info.get("regional_endpoint_uplift_multiplier", None), - output_cost_per_audio_token=_model_info.get("output_cost_per_audio_token", None), - output_cost_per_character=_model_info.get("output_cost_per_character", None), - output_cost_per_reasoning_token=_model_info.get("output_cost_per_reasoning_token", None), - output_cost_per_reasoning_token_flex=_model_info.get("output_cost_per_reasoning_token_flex", None), - output_cost_per_reasoning_token_priority=_model_info.get( - "output_cost_per_reasoning_token_priority", None - ), - output_cost_per_token_above_128k_tokens=_model_info.get( - "output_cost_per_token_above_128k_tokens", None - ), - output_cost_per_character_above_128k_tokens=_model_info.get( - "output_cost_per_character_above_128k_tokens", None - ), - output_cost_per_token_above_200k_tokens=_model_info.get( - "output_cost_per_token_above_200k_tokens", None - ), - output_cost_per_token_above_200k_tokens_priority=_model_info.get( - "output_cost_per_token_above_200k_tokens_priority", None - ), - output_cost_per_token_above_272k_tokens=_model_info.get( - "output_cost_per_token_above_272k_tokens", None - ), - output_cost_per_token_above_272k_tokens_priority=_model_info.get( - "output_cost_per_token_above_272k_tokens_priority", None - ), - output_cost_per_token_above_272k_tokens_flex=_model_info.get( - "output_cost_per_token_above_272k_tokens_flex", None - ), - output_cost_per_token_above_512k_tokens=_model_info.get( - "output_cost_per_token_above_512k_tokens", None - ), - output_cost_per_second=_model_info.get("output_cost_per_second", None), - output_cost_per_second_1080p=_model_info.get("output_cost_per_second_1080p", None), - output_cost_per_second_480p=_model_info.get("output_cost_per_second_480p", None), - output_cost_per_second_4k=_model_info.get("output_cost_per_second_4k", None), - output_cost_per_video_per_second=_model_info.get("output_cost_per_video_per_second", None), - output_cost_per_image=_model_info.get("output_cost_per_image", None), - output_cost_per_image_token=_model_info.get("output_cost_per_image_token", None), - output_cost_per_video_token=_model_info.get("output_cost_per_video_token", None), - output_vector_size=_model_info.get("output_vector_size", None), - citation_cost_per_token=_model_info.get("citation_cost_per_token", None), - tiered_pricing=_model_info.get("tiered_pricing", None), - litellm_provider=_model_info.get("litellm_provider", custom_llm_provider), - mode=_model_info.get("mode"), - supported_endpoints=_model_info.get("supported_endpoints", None), - supports_system_messages=_model_info.get("supports_system_messages", None), - supports_response_schema=_model_info.get("supports_response_schema", None), - supports_vision=_model_info.get("supports_vision", None), - supports_function_calling=_model_info.get("supports_function_calling", None), - supports_parallel_function_calling=_model_info.get("supports_parallel_function_calling", None), - supports_tool_choice=_model_info.get("supports_tool_choice", None), - supports_assistant_prefill=_model_info.get("supports_assistant_prefill", None), - supports_prompt_caching=_model_info.get("supports_prompt_caching", None), - supports_prompt_cache_breakpoint=_model_info.get("supports_prompt_cache_breakpoint", None), - supports_audio_input=_model_info.get("supports_audio_input", None), - supports_audio_output=_model_info.get("supports_audio_output", None), - supports_pdf_input=_model_info.get("supports_pdf_input", None), - supports_embedding_image_input=_model_info.get("supports_embedding_image_input", None), - supports_native_streaming=_model_info.get("supports_native_streaming", None), - supports_native_structured_output=_model_info.get("supports_native_structured_output", None), - supports_web_search=_model_info.get("supports_web_search", None), - supports_url_context=_model_info.get("supports_url_context", None), - supports_reasoning=_model_info.get("supports_reasoning", None), - supports_adaptive_thinking=_model_info.get("supports_adaptive_thinking", None), - supports_legacy_thinking=_model_info.get("supports_legacy_thinking", None), - thinking_always_on=_model_info.get("thinking_always_on", None), - supports_tool_search=_model_info.get("supports_tool_search", None), - supports_mid_conversation_system=_model_info.get("supports_mid_conversation_system", None), - supports_none_reasoning_effort=_model_info.get("supports_none_reasoning_effort", None), - supports_minimal_reasoning_effort=_model_info.get("supports_minimal_reasoning_effort", None), - supports_low_reasoning_effort=_model_info.get("supports_low_reasoning_effort", None), - supports_xhigh_reasoning_effort=_model_info.get("supports_xhigh_reasoning_effort", None), - supports_max_reasoning_effort=_model_info.get("supports_max_reasoning_effort", None), - reasoning_effort_levels=_model_info.get("reasoning_effort_levels", None), - default_reasoning_effort=_model_info.get("default_reasoning_effort", None), - bedrock_output_config_effort_ceiling=_model_info.get("bedrock_output_config_effort_ceiling", None), - bedrock_converse_supports_strict_tools=_model_info.get("bedrock_converse_supports_strict_tools", None), - supports_computer_use=_model_info.get("supports_computer_use", None), - search_context_cost_per_query=_model_info.get("search_context_cost_per_query", None), - web_search_billing_unit=_model_info.get("web_search_billing_unit", None), - google_maps_grounding_cost_per_query=_model_info.get("google_maps_grounding_cost_per_query", None), - tpm=_model_info.get("tpm", None), - rpm=_model_info.get("rpm", None), - ocr_cost_per_page=_model_info.get("ocr_cost_per_page", None), - ocr_cost_per_credit=_model_info.get("ocr_cost_per_credit", None), - annotation_cost_per_page=_model_info.get("annotation_cost_per_page", None), - provider_specific_entry=_model_info.get("provider_specific_entry", None), - uses_embed_content=_model_info.get("uses_embed_content", None), - supports_image_size=_model_info.get("supports_image_size", None), - ) - for cost_key, cost_value in _model_info.items(): - if cost_key not in returned_model_info and _ABOVE_THRESHOLD_COST_KEY.search(cost_key) is not None: - returned_model_info[cost_key] = cost_value - return returned_model_info - except Exception as e: - verbose_logger.debug("Error getting model info: %s", e) - raise Exception( - f"This model isn't mapped yet. model={model}, custom_llm_provider={custom_llm_provider}. Add it here - https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json." - ) - - -def _build_model_info( - model: str, - custom_llm_provider: str | None = None, - api_base: str | None = None, - api_key: str | None = None, -) -> ModelInfo: - supported_openai_params = litellm.get_supported_openai_params(model=model, custom_llm_provider=custom_llm_provider) - - _model_info: Final = _get_model_info_helper( - model=model, - custom_llm_provider=custom_llm_provider, - api_base=api_base, - api_key=api_key, - ) - - provider_info: Final = get_provider_info(model=model, custom_llm_provider=custom_llm_provider) - if provider_info: - for key, value in provider_info.items(): - if value is not None: - _model_info[key] = value - - # if verbose_logger.isEnabledFor(logging.DEBUG): - # verbose_logger.debug(f"model_info: {_model_info}") - - return ModelInfo(**_model_info, supported_openai_params=supported_openai_params) - - -@lru_cache(maxsize=DEFAULT_MAX_LRU_CACHE_SIZE) -def _cached_get_model_info( - model: str, - custom_llm_provider: str | None = None, - api_base: str | None = None, -) -> ModelInfo: - return _build_model_info(model=model, custom_llm_provider=custom_llm_provider, api_base=api_base) - - -def get_model_info( - model: str, - custom_llm_provider: str | None = None, - api_base: str | None = None, - api_key: str | None = None, -) -> ModelInfo: - """ - Get a dict for the maximum tokens (context window), input_cost_per_token, output_cost_per_token for a given model. - - Parameters: - - model (str): The name of the model. - - custom_llm_provider (str | null): the provider used for the model. If provided, used to check if the litellm model info is for that provider. - - Returns: - dict: A dictionary containing the following information: - key: Required[str] # the key in litellm.model_cost which is returned - max_tokens: Required[Optional[int]] - max_input_tokens: Required[Optional[int]] - max_output_tokens: Required[Optional[int]] - input_cost_per_token: Required[float] - input_cost_per_character: Optional[float] # only for vertex ai models - input_cost_per_token_above_128k_tokens: Optional[float] # only for vertex ai models - input_cost_per_character_above_128k_tokens: Optional[ - float - ] # only for vertex ai models - input_cost_per_query: Optional[float] # only for rerank models - input_cost_per_image: Optional[float] # only for vertex ai models - input_cost_per_audio_token: Optional[float] - input_cost_per_audio_per_second: Optional[float] # only for vertex ai models - input_cost_per_video_per_second: Optional[float] # only for vertex ai models - output_cost_per_token: Required[float] - output_cost_per_audio_token: Optional[float] - output_cost_per_character: Optional[float] # only for vertex ai models - output_cost_per_token_above_128k_tokens: Optional[ - float - ] # only for vertex ai models - output_cost_per_character_above_128k_tokens: Optional[ - float - ] # only for vertex ai models - output_cost_per_image: Optional[float] - output_vector_size: Optional[int] - output_cost_per_video_per_second: Optional[float] # only for vertex ai models - output_cost_per_audio_per_second: Optional[float] # only for vertex ai models - litellm_provider: Required[str] - mode: Required[ - Literal[ - "completion", "embedding", "image_generation", "chat", "audio_transcription" - ] - ] - supported_openai_params: Required[Optional[List[str]]] - supports_system_messages: Optional[bool] - supports_response_schema: Optional[bool] - supports_vision: Optional[bool] - supports_function_calling: Optional[bool] - supports_tool_choice: Optional[bool] - supports_prompt_caching: Optional[bool] - supports_prompt_cache_breakpoint: Optional[bool] - supports_audio_input: Optional[bool] - supports_audio_output: Optional[bool] - supports_pdf_input: Optional[bool] - supports_web_search: Optional[bool] - supports_url_context: Optional[bool] - supports_reasoning: Optional[bool] - Raises: - Exception: If the model is not mapped yet. - - Example: - >>> get_model_info("gpt-4") - { - "max_tokens": 8192, - "input_cost_per_token": 0.00003, - "output_cost_per_token": 0.00006, - "litellm_provider": "openai", - "mode": "chat", - "supported_openai_params": ["temperature", "max_tokens", "top_p", "frequency_penalty", "presence_penalty"] - } - """ - # api_key is a per-caller credential, not part of the model identity, so it is - # kept out of the cache key; explicit keys are resolved without the cache. - if api_key is not None: - return _build_model_info(model, custom_llm_provider, api_base, api_key) - return _cached_get_model_info(model, custom_llm_provider, api_base) - - -get_model_info.cache_clear = _cached_get_model_info.cache_clear -get_model_info.cache_info = _cached_get_model_info.cache_info - - -def json_schema_type(python_type_name: str): - """Converts standard python types to json schema types - - Parameters - ---------- - python_type_name : str - __name__ of type - - Returns - ------- - str - a standard JSON schema type, "string" if not recognized. - """ - python_to_json_schema_types: Final = { - str.__name__: "string", - int.__name__: "integer", - float.__name__: "number", - bool.__name__: "boolean", - list.__name__: "array", - dict.__name__: "object", - "NoneType": "null", - } - - return python_to_json_schema_types.get(python_type_name, "string") - - -def function_to_dict(input_function) -> dict: - """Using type hints and numpy-styled docstring, - produce a dictionary usable for OpenAI function calling - - Parameters - ---------- - input_function : function - A function with a numpy-style docstring - - Returns - ------- - dictionnary - A dictionnary to add to the list passed to `functions` parameter of `litellm.completion` - """ - # Get function name and docstring - try: - import inspect - from ast import literal_eval - - from numpydoc.docscrape import NumpyDocString - except Exception as e: - raise e - - name: Final = input_function.__name__ - docstring: Final = inspect.getdoc(input_function) - numpydoc: Final = NumpyDocString(docstring) - description: Final = "\n".join([s.strip() for s in numpydoc["Summary"]]) - - # Get function parameters and their types from annotations and docstring - parameters: Final = {} - required_params: Final = [] - param_info: Final = inspect.signature(input_function).parameters - - for param_name, param in param_info.items(): - if hasattr(param, "annotation"): - param_type = json_schema_type(param.annotation.__name__) - else: - param_type = None - param_description = None - param_enum = None - - # Try to extract param description from docstring using numpydoc - for param_data in numpydoc["Parameters"]: - if param_data.name == param_name: - if hasattr(param_data, "type"): - # replace type from docstring rather than annotation - param_type = param_data.type - if "optional" in param_type: - param_type = param_type.split(",")[0] - elif "{" in param_type: - # may represent a set of acceptable values - # translating as enum for function calling - try: - param_enum = str(list(literal_eval(param_type))) - param_type = "string" - except Exception: - pass - param_type = json_schema_type(param_type) - param_description = "\n".join([s.strip() for s in param_data.desc]) - - param_dict = { - "type": param_type, - "description": param_description, - "enum": param_enum, - } - - parameters[param_name] = dict([(k, v) for k, v in param_dict.items() if isinstance(v, str)]) - - # Check if the parameter has no default value (i.e., it's required) - if param.default == param.empty: - required_params.append(param_name) - - # Create the dictionary - result: Final = { - "name": name, - "description": description, - "parameters": { - "type": "object", - "properties": parameters, - }, - } - - # Add "required" key if there are required parameters - if required_params: - result["parameters"]["required"] = required_params - - return result - - -def modify_url(original_url, new_path): - url: Final = httpx.URL(original_url) - modified_url: Final = url.copy_with(path=new_path) - return str(modified_url) - - -def load_test_model( - model: str, - custom_llm_provider: str = "", - api_base: str = "", - prompt: str = "", - num_calls: int = 0, - force_timeout: int = 0, -): - test_prompt = "Hey, how's it going" - test_calls = 100 - if prompt: - test_prompt = prompt - if num_calls: - test_calls = num_calls - messages: Final = [[{"role": "user", "content": test_prompt}] for _ in range(test_calls)] - start_time: Final = time.time() - try: - litellm.batch_completion( - model=model, - messages=messages, - custom_llm_provider=custom_llm_provider, - api_base=api_base, - force_timeout=force_timeout, - ) - end_time = time.time() - response_time = end_time - start_time - return { - "total_response_time": response_time, - "calls_made": 100, - "status": "success", - "exception": None, - } - except Exception as e: - end_time = time.time() - response_time = end_time - start_time - return { - "total_response_time": response_time, - "calls_made": 100, - "status": "failed", - "exception": e, - } - - -def get_provider_fields(custom_llm_provider: str) -> list[ProviderField]: - """Return the fields required for each provider""" - - if custom_llm_provider == "databricks": - return litellm.DatabricksConfig().get_required_params() - - elif custom_llm_provider == "ollama": - return litellm.OllamaConfig().get_required_params() - - elif custom_llm_provider == "azure_ai": - return litellm.AzureAIStudioConfig().get_required_params() - - else: - return [] - - -def create_proxy_transport_and_mounts(): - proxies: Final = {key: None if url is None else Proxy(url=url) for key, url in get_environment_proxies().items()} - - sync_proxy_mounts: Final = {} - async_proxy_mounts: Final = {} - - # Retrieve NO_PROXY environment variable - no_proxy: Final = os.getenv("NO_PROXY", None) - no_proxy_urls: Final = no_proxy.split(",") if no_proxy else [] - - for key, proxy in proxies.items(): - if proxy is None: - sync_proxy_mounts[key] = httpx.HTTPTransport() - async_proxy_mounts[key] = httpx.AsyncHTTPTransport() - else: - sync_proxy_mounts[key] = httpx.HTTPTransport(proxy=proxy) - async_proxy_mounts[key] = httpx.AsyncHTTPTransport(proxy=proxy) - - for url in no_proxy_urls: - sync_proxy_mounts[url] = httpx.HTTPTransport() - async_proxy_mounts[url] = httpx.AsyncHTTPTransport() - - return sync_proxy_mounts, async_proxy_mounts - - -def validate_environment( - model: str | None = None, - api_key: str | None = None, - api_base: str | None = None, - api_version: str | None = None, -) -> dict: - """ - Checks if the environment variables are valid for the given model. - - Args: - model (Optional[str]): The name of the model. Defaults to None. - api_key (Optional[str]): If the user passed in an api key, of their own. - - Returns: - dict: A dictionary containing the following keys: - - keys_in_environment (bool): True if all the required keys are present in the environment, False otherwise. - - missing_keys (List[str]): A list of missing keys in the environment. - """ - keys_in_environment = False - missing_keys: list[str] = [] - - if model is None: - return { - "keys_in_environment": keys_in_environment, - "missing_keys": missing_keys, - } - ## EXTRACT LLM PROVIDER - if model name provided - try: - get_llm_provider: Final = litellm_utils.get_llm_provider - _, custom_llm_provider, _, _ = get_llm_provider(model=model) - except Exception: - custom_llm_provider = None - - if custom_llm_provider: - if custom_llm_provider == "openai": - if "OPENAI_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("OPENAI_API_KEY") - elif custom_llm_provider == "azure": - if "AZURE_API_BASE" in os.environ and "AZURE_API_VERSION" in os.environ and "AZURE_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.extend(["AZURE_API_BASE", "AZURE_API_VERSION", "AZURE_API_KEY"]) - elif custom_llm_provider == "anthropic": - if "ANTHROPIC_API_KEY" in os.environ or "ANTHROPIC_AUTH_TOKEN" in os.environ: - keys_in_environment = True - else: - missing_keys.append("ANTHROPIC_API_KEY") - elif custom_llm_provider == "cohere": - if "COHERE_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("COHERE_API_KEY") - elif custom_llm_provider == "replicate": - if "REPLICATE_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("REPLICATE_API_KEY") - elif custom_llm_provider == "openrouter": - if "OPENROUTER_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("OPENROUTER_API_KEY") - elif custom_llm_provider == "vercel_ai_gateway": - if "VERCEL_AI_GATEWAY_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("VERCEL_AI_GATEWAY_API_KEY") - elif custom_llm_provider == "datarobot": - if "DATAROBOT_API_TOKEN" in os.environ: - keys_in_environment = True - else: - missing_keys.append("DATAROBOT_API_TOKEN") - elif custom_llm_provider == "vertex_ai": - if "VERTEXAI_PROJECT" in os.environ and "VERTEXAI_LOCATION" in os.environ: - keys_in_environment = True - else: - missing_keys.extend(["VERTEXAI_PROJECT", "VERTEXAI_LOCATION"]) - elif custom_llm_provider == "huggingface": - if "HUGGINGFACE_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("HUGGINGFACE_API_KEY") - elif custom_llm_provider == "ai21": - if "AI21_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("AI21_API_KEY") - elif custom_llm_provider == "together_ai": - if "TOGETHERAI_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("TOGETHERAI_API_KEY") - elif custom_llm_provider == "aleph_alpha": - if "ALEPH_ALPHA_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("ALEPH_ALPHA_API_KEY") - elif custom_llm_provider == "baseten": - if "BASETEN_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("BASETEN_API_KEY") - elif custom_llm_provider == "nlp_cloud": - if "NLP_CLOUD_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("NLP_CLOUD_API_KEY") - elif custom_llm_provider == "bedrock" or custom_llm_provider == "sagemaker": - if ("AWS_ACCESS_KEY_ID" in os.environ and "AWS_SECRET_ACCESS_KEY" in os.environ) or ( - # IAM role, profile, or web identity auth don't require access keys - "AWS_ROLE_ARN" in os.environ - or "AWS_PROFILE" in os.environ - or "AWS_WEB_IDENTITY_TOKEN_FILE" in os.environ - or "AWS_CONTAINER_CREDENTIALS_RELATIVE_URI" in os.environ # ECS task role - or "AWS_CONTAINER_CREDENTIALS_FULL_URI" in os.environ # ECS/Fargate full URI credential delivery - ): - keys_in_environment = True - else: - missing_keys.append("AWS_ACCESS_KEY_ID") - missing_keys.append("AWS_SECRET_ACCESS_KEY") - elif custom_llm_provider in ["ollama", "ollama_chat"]: - if "OLLAMA_API_BASE" in os.environ: - keys_in_environment = True - else: - missing_keys.append("OLLAMA_API_BASE") - elif custom_llm_provider == "anyscale": - if "ANYSCALE_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("ANYSCALE_API_KEY") - elif custom_llm_provider == "deepinfra": - if "DEEPINFRA_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("DEEPINFRA_API_KEY") - elif custom_llm_provider == "featherless_ai": - if "FEATHERLESS_AI_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("FEATHERLESS_AI_API_KEY") - elif custom_llm_provider == "gemini": - if ("GOOGLE_API_KEY" in os.environ) or ("GEMINI_API_KEY" in os.environ): - keys_in_environment = True - else: - missing_keys.append("GOOGLE_API_KEY") - missing_keys.append("GEMINI_API_KEY") - elif custom_llm_provider == "groq": - if "GROQ_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("GROQ_API_KEY") - elif custom_llm_provider == "nvidia_nim": - if "NVIDIA_NIM_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("NVIDIA_NIM_API_KEY") - elif custom_llm_provider == "cerebras": - if "CEREBRAS_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("CEREBRAS_API_KEY") - elif custom_llm_provider == "baseten": - if "BASETEN_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("BASETEN_API_KEY") - elif custom_llm_provider == "xai": - if "XAI_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("XAI_API_KEY") - elif custom_llm_provider == "ai21_chat": - if "AI21_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("AI21_API_KEY") - elif custom_llm_provider == "volcengine": - if "VOLCENGINE_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("VOLCENGINE_API_KEY") - elif custom_llm_provider == "codestral" or custom_llm_provider == "text-completion-codestral": - if "CODESTRAL_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("CODESTRAL_API_KEY") - elif custom_llm_provider == "inception" or custom_llm_provider == "text-completion-inception": - if "INCEPTION_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("INCEPTION_API_KEY") - elif custom_llm_provider == "deepseek": - if "DEEPSEEK_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("DEEPSEEK_API_KEY") - elif custom_llm_provider == "tencent": - if "TENCENT_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("TENCENT_API_KEY") - elif custom_llm_provider == "mistral": - if "MISTRAL_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("MISTRAL_API_KEY") - elif custom_llm_provider == "palm": - if "PALM_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("PALM_API_KEY") - elif custom_llm_provider == "perplexity": - if "PERPLEXITYAI_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("PERPLEXITYAI_API_KEY") - elif custom_llm_provider == "voyage": - if "VOYAGE_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("VOYAGE_API_KEY") - elif custom_llm_provider == "infinity": - if "INFINITY_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("INFINITY_API_KEY") - elif custom_llm_provider == "fireworks_ai": - if ( - "FIREWORKS_AI_API_KEY" in os.environ - or "FIREWORKS_API_KEY" in os.environ - or "FIREWORKSAI_API_KEY" in os.environ - or "FIREWORKS_AI_TOKEN" in os.environ - ): - keys_in_environment = True - else: - missing_keys.append("FIREWORKS_AI_API_KEY") - elif custom_llm_provider == "cloudflare": - if "CLOUDFLARE_API_KEY" in os.environ and ( - "CLOUDFLARE_ACCOUNT_ID" in os.environ or "CLOUDFLARE_API_BASE" in os.environ - ): - keys_in_environment = True - else: - missing_keys.append("CLOUDFLARE_API_KEY") - missing_keys.append("CLOUDFLARE_API_BASE") - elif custom_llm_provider == "novita": - if "NOVITA_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("NOVITA_API_KEY") - elif custom_llm_provider == "nebius": - if "NEBIUS_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("NEBIUS_API_KEY") - elif custom_llm_provider == "wandb": - if "WANDB_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("WANDB_API_KEY") - elif custom_llm_provider in ("dashscope", "qwencloud", "qwen_ai_platform"): - if f"{custom_llm_provider.upper()}_API_KEY" in os.environ or "DASHSCOPE_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append(f"{custom_llm_provider.upper()}_API_KEY") - elif custom_llm_provider == "modelscope": - if "MODELSCOPE_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("MODELSCOPE_API_KEY") - elif custom_llm_provider == "moonshot": - if "MOONSHOT_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("MOONSHOT_API_KEY") - else: - ## openai - chatcompletion + text completion - if ( - model in litellm.open_ai_chat_completion_models - or model in litellm.open_ai_text_completion_models - or model in litellm.open_ai_embedding_models - or model in litellm.openai_image_generation_models - or model.startswith("gpt-image") - ): - if "OPENAI_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("OPENAI_API_KEY") - ## anthropic - elif model in litellm.anthropic_models: - if "ANTHROPIC_API_KEY" in os.environ or "ANTHROPIC_AUTH_TOKEN" in os.environ: - keys_in_environment = True - else: - missing_keys.append("ANTHROPIC_API_KEY") - ## cohere - elif model in litellm.cohere_models: - if "COHERE_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("COHERE_API_KEY") - ## replicate - elif model in litellm.replicate_models: - if "REPLICATE_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("REPLICATE_API_KEY") - ## openrouter - elif model in litellm.openrouter_models: - if "OPENROUTER_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("OPENROUTER_API_KEY") - ## vercel_ai_gateway - elif model in litellm.vercel_ai_gateway_models: - if "VERCEL_AI_GATEWAY_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("VERCEL_AI_GATEWAY_API_KEY") - ## datarobot - elif model in litellm.datarobot_models: - if "DATAROBOT_API_TOKEN" in os.environ: - keys_in_environment = True - else: - missing_keys.append("DATAROBOT_API_TOKEN") - ## vertex - text + chat models - elif ( - model in litellm.vertex_chat_models - or model in litellm.vertex_text_models - or model in litellm.models_by_provider["vertex_ai"] - ): - if "VERTEXAI_PROJECT" in os.environ and "VERTEXAI_LOCATION" in os.environ: - keys_in_environment = True - else: - missing_keys.extend(["VERTEXAI_PROJECT", "VERTEXAI_LOCATION"]) - ## huggingface - elif model in litellm.huggingface_models: - if "HUGGINGFACE_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("HUGGINGFACE_API_KEY") - ## ai21 - elif model in litellm.ai21_models: - if "AI21_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("AI21_API_KEY") - ## together_ai - elif model in litellm.together_ai_models: - if "TOGETHERAI_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("TOGETHERAI_API_KEY") - ## aleph_alpha - elif model in litellm.aleph_alpha_models: - if "ALEPH_ALPHA_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("ALEPH_ALPHA_API_KEY") - ## baseten - elif model in litellm.baseten_models: - if "BASETEN_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("BASETEN_API_KEY") - ## nlp_cloud - elif model in litellm.nlp_cloud_models: - if "NLP_CLOUD_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("NLP_CLOUD_API_KEY") - elif model in litellm.novita_models: - if "NOVITA_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("NOVITA_API_KEY") - elif model in litellm.nebius_models: - if "NEBIUS_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("NEBIUS_API_KEY") - elif model in litellm.wandb_models: - if "WANDB_API_KEY" in os.environ: - keys_in_environment = True - else: - missing_keys.append("WANDB_API_KEY") - - def filter_missing_keys(keys: list[str], exclude_pattern: str) -> list[str]: - """Filter out keys that contain the exclude_pattern (case insensitive).""" - return [key for key in keys if exclude_pattern not in key.lower()] - - if api_key is not None: - missing_keys = filter_missing_keys(missing_keys, "api_key") - - if api_base is not None: - missing_keys = filter_missing_keys(missing_keys, "api_base") - - if api_version is not None: - missing_keys = filter_missing_keys(missing_keys, "api_version") - - if len(missing_keys) == 0: # no missing keys - keys_in_environment = True - - return {"keys_in_environment": keys_in_environment, "missing_keys": missing_keys} - - -def acreate(*args, **kwargs): ## Thin client to handle the acreate langchain call - return litellm.acompletion(*args, **kwargs) - - -def valid_model(model): - try: - # for a given model name, check if the user has the right permissions to access the model - if model in litellm.open_ai_chat_completion_models or model in litellm.open_ai_text_completion_models: - openai.models.retrieve(model) - else: - messages: Final = [{"role": "user", "content": "Hello World"}] - litellm.completion(model=model, messages=messages) - except Exception: - raise BadRequestError(message="", model=model, llm_provider="") - - -def check_valid_key(model: str, api_key: str): - """ - Checks if a given API key is valid for a specific model by making a litellm.completion call with max_tokens=10 - - Args: - model (str): The name of the model to check the API key against. - api_key (str): The API key to be checked. - - Returns: - bool: True if the API key is valid for the model, False otherwise. - """ - messages: Final = [{"role": "user", "content": "Hey, how's it going?"}] - try: - litellm.completion(model=model, messages=messages, api_key=api_key, max_tokens=10) - return True - except AuthenticationError: - return False - except Exception: - return False - - -def _should_retry(status_code: int): - """ - Retries on 408, 409, 429 and 500 errors. - - Any client error in the 400-499 range that isn't explicitly handled (such as 400 Bad Request, 401 Unauthorized, 403 Forbidden, 404 Not Found, etc.) would not trigger a retry. - - Reimplementation of openai's should retry logic, since that one can't be imported. - https://github.com/openai/openai-python/blob/af67cfab4210d8e497c05390ce14f39105c77519/src/openai/_base_client.py#L639 - """ - # If the server explicitly says whether or not to retry, obey. - # Retry on request timeouts. - if status_code == 408: - return True - - # Retry on lock timeouts. - if status_code == 409: - return True - - # Retry on rate limits. - if status_code == 429: - return True - - # Retry internal errors. - if status_code >= 500: - return True - - return False - - -def _get_retry_after_from_exception_header( - response_headers: httpx.Headers | None = None, -): - """ - Reimplementation of openai's calculate retry after, since that one can't be imported. - https://github.com/openai/openai-python/blob/af67cfab4210d8e497c05390ce14f39105c77519/src/openai/_base_client.py#L631 - """ - try: - import email # openai import - - # About the Retry-After header: https://developer.mozilla.org/en-US/docs/Web/HTTP/Headers/Retry-After - # - # ". See https://developer.mozilla.org/en-US/docs/Web/HTTP/Headers/Retry-After#syntax for - # details. - if response_headers is not None: - retry_header: Final[str] = response_headers.get("retry-after") - try: - retry_after = int(retry_header) - except Exception: - retry_date_tuple: Final = email.utils.parsedate_tz(retry_header) - if retry_date_tuple is None: - retry_after = -1 - else: - retry_date: Final = email.utils.mktime_tz(retry_date_tuple) - retry_after = int(retry_date - time.time()) - else: - retry_after = -1 - - return retry_after - - except Exception: - retry_after = -1 - - -def _calculate_retry_after( - remaining_retries: int, - max_retries: int, - response_headers: httpx.Headers | None = None, - min_timeout: int = 0, -) -> float | int: - retry_after: Final = _get_retry_after_from_exception_header(response_headers) - - # Add some jitter (default JITTER is 0.75 - so upto 0.75s) - jitter: Final = JITTER * random.random() - - # If the API asks us to wait a certain amount of time (and it's a reasonable amount), just do what it says. - if retry_after is not None and 0 < retry_after <= 60: - return retry_after + jitter - - # Calculate exponential backoff - num_retries: Final = max_retries - remaining_retries - sleep_seconds = INITIAL_RETRY_DELAY * pow(2.0, num_retries) - - # Make sure sleep_seconds is boxed between min_timeout and MAX_RETRY_DELAY - sleep_seconds = max(sleep_seconds, min_timeout) - sleep_seconds = min(sleep_seconds, MAX_RETRY_DELAY) - - return sleep_seconds + jitter - - -# custom prompt helper function -def register_prompt_template( - model: str, - roles: dict = {}, - initial_prompt_value: str = "", - final_prompt_value: str = "", - tokenizer_config: dict = {}, -): - """ - Register a prompt template to follow your custom format for a given model - - Args: - model (str): The name of the model. - roles (dict): A dictionary mapping roles to their respective prompt values. - initial_prompt_value (str, optional): The initial prompt value. Defaults to "". - final_prompt_value (str, optional): The final prompt value. Defaults to "". - - Returns: - dict: The updated custom prompt dictionary. - Example usage: - ``` - import litellm - litellm.register_prompt_template( - model="llama-2", - initial_prompt_value="You are a good assistant" # [OPTIONAL] - roles={ - "system": { - "pre_message": "[INST] <>\n", # [OPTIONAL] - "post_message": "\n<>\n [/INST]\n" # [OPTIONAL] - }, - "user": { - "pre_message": "[INST] ", # [OPTIONAL] - "post_message": " [/INST]" # [OPTIONAL] - }, - "assistant": { - "pre_message": "\n" # [OPTIONAL] - "post_message": "\n" # [OPTIONAL] - } - } - final_prompt_value="Now answer as best you can:" # [OPTIONAL] - ) - ``` - """ - complete_model: Final = model - potential_models: Final = [complete_model] - try: - get_llm_provider: Final = litellm_utils.get_llm_provider - model = get_llm_provider(model=model)[0] - potential_models.append(model) - except Exception: - pass - if tokenizer_config: - for m in potential_models: - litellm.known_tokenizer_config[m] = { - "tokenizer": tokenizer_config, - "status": "success", - } - else: - for m in potential_models: - litellm.custom_prompt_dict[m] = { - "roles": roles, - "initial_prompt_value": initial_prompt_value, - "final_prompt_value": final_prompt_value, - } - - return litellm.custom_prompt_dict - - -class TextCompletionStreamWrapper: - def __init__( - self, - completion_stream, - model, - stream_options: dict | None = None, - custom_llm_provider: str | None = None, - ): - self.completion_stream = completion_stream - self.model = model - self.stream_options = stream_options - self.custom_llm_provider = custom_llm_provider - - def __iter__(self): - return self - - def __aiter__(self): - return self - - def convert_to_text_completion_object(self, chunk: ModelResponse): - try: - response: Final = TextCompletionResponse() - response["id"] = chunk.get("id", None) - response["object"] = "text_completion" - response["created"] = chunk.get("created", None) - response["model"] = chunk.get("model", None) - text_choices: Final = TextChoices() - if isinstance(chunk, Choices): # chunk should always be of type StreamingChoices - raise Exception - delta: Final = chunk["choices"][0]["delta"] - text_choices["text"] = delta["content"] - text_choices["reasoning_content"] = delta.get("reasoning_content") - text_choices["index"] = chunk["choices"][0]["index"] - text_choices["finish_reason"] = chunk["choices"][0]["finish_reason"] - response["choices"] = [text_choices] - - # only pass usage when stream_options["include_usage"] is True - if self.stream_options and self.stream_options.get("include_usage", False) is True: - response["usage"] = chunk.get("usage", None) - - return response - except Exception as e: - raise Exception(f"Error occurred converting to text completion object - chunk: {chunk}; Error: {e}") - - def __next__(self): - # model_response = ModelResponse(stream=True, model=self.model) - TextCompletionResponse() - try: - for chunk in self.completion_stream: - if chunk == "None" or chunk is None: - raise Exception - processed_chunk = self.convert_to_text_completion_object(chunk=chunk) - return processed_chunk - raise StopIteration - except StopIteration: - raise StopIteration - except Exception as e: - exception_type: Final = getattr(sys.modules[__name__], "exception_type") - raise exception_type( - model=self.model, - custom_llm_provider=self.custom_llm_provider or "", - original_exception=e, - completion_kwargs={}, - extra_kwargs={}, - ) - - async def __anext__(self): - try: - async for chunk in self.completion_stream: - if chunk == "None" or chunk is None: - raise Exception - processed_chunk = self.convert_to_text_completion_object(chunk=chunk) - return processed_chunk - raise StopIteration - except StopIteration: - raise StopAsyncIteration - - -def mock_completion_streaming_obj(model_response, mock_response, model, n: int | None = None): - if isinstance(mock_response, litellm.MockException): - raise mock_response - if isinstance(mock_response, ModelResponseStream): - yield mock_response - return - for i in range(0, len(mock_response), 3): - completion_obj = Delta(role="assistant", content=mock_response[i : i + 3]) - if n is None: - model_response.choices[0].delta = completion_obj - else: - _all_choices = [] - for j in range(n): - _streaming_choice = litellm.utils.StreamingChoices( - index=j, - delta=litellm.utils.Delta(role="assistant", content=mock_response[i : i + 3]), - ) - _all_choices.append(_streaming_choice) - model_response.choices = _all_choices - yield model_response - # Separate terminal object so content chunks keep finish_reason unset. - terminal: Final = model_response.model_copy(deep=True) - if n is None: - terminal.choices[0].delta = Delta(role="assistant", content=None) - terminal.choices[0].finish_reason = "stop" - else: - for j in range(n): - terminal.choices[j].delta = litellm.utils.Delta(role="assistant", content=None) - terminal.choices[j].finish_reason = "stop" - yield terminal - - -async def async_mock_completion_streaming_obj( - model_response, - mock_response: str | MockException | ModelResponseStream, - model, - n: int | None = None, -): - if isinstance(mock_response, litellm.MockException): - raise mock_response - if isinstance(mock_response, ModelResponseStream): - yield mock_response - return - for i in range(0, len(mock_response), 3): - completion_obj = Delta(role="assistant", content=mock_response[i : i + 3]) - if n is None: - model_response.choices[0].delta = completion_obj - else: - _all_choices = [] - for j in range(n): - _streaming_choice = litellm.utils.StreamingChoices( - index=j, - delta=litellm.utils.Delta(role="assistant", content=mock_response[i : i + 3]), - ) - _all_choices.append(_streaming_choice) - model_response.choices = _all_choices - yield model_response - # Separate terminal object so content chunks keep finish_reason unset. - terminal: Final = model_response.model_copy(deep=True) - if n is None: - terminal.choices[0].delta = Delta(role="assistant", content=None) - terminal.choices[0].finish_reason = "stop" - else: - for j in range(n): - terminal.choices[j].delta = litellm.utils.Delta(role="assistant", content=None) - terminal.choices[j].finish_reason = "stop" - yield terminal - - -########## Reading Config File ############################ -def read_config_args(config_path) -> dict: - try: - import os - - os.getcwd() - with open(config_path, "r") as config_file: - config: Final = json.load(config_file) - - # read keys/ values from config file and return them - return config - except Exception as e: - raise e - - -########## experimental completion variants ############################ - - -def process_system_message(system_message, max_tokens, model): - system_message_event: Final = {"role": "system", "content": system_message} - system_message_tokens = get_token_count([system_message_event], model) - - if system_message_tokens > max_tokens: - print_verbose("`tokentrimmer`: Warning, system message exceeds token limit. Trimming...") - # shorten system message to fit within max_tokens - new_system_message: Final = shorten_message_to_fit_limit(system_message_event, max_tokens, model) - system_message_tokens = get_token_count([new_system_message], model) - - return system_message_event, max_tokens - system_message_tokens - - -def process_messages(messages, max_tokens, model): - # Process messages from older to more recent - messages = messages[::-1] - final_messages = [] - verbose_logger.debug( - "calling process_messages with messages: %s, max_tokens: %s, model: %s", messages, max_tokens, model - ) - for message in messages: - verbose_logger.debug("processing final_messages: %s", final_messages) - used_tokens = get_token_count(final_messages, model) - available_tokens = max_tokens - used_tokens - verbose_logger.debug("used_tokens: %s, available_tokens: %s", used_tokens, available_tokens) - if available_tokens <= 3: - break - - final_messages = attempt_message_addition( - final_messages=final_messages, - message=message, - available_tokens=available_tokens, - max_tokens=max_tokens, - model=model, - ) - verbose_logger.debug("final_messages after attempt_message_addition: %s", final_messages) - verbose_logger.debug("Final messages: %s", final_messages) - return final_messages - - -def attempt_message_addition(final_messages, message, available_tokens, max_tokens, model): - temp_messages: Final = [message] + final_messages - temp_message_tokens: Final = get_token_count(messages=temp_messages, model=model) - verbose_logger.debug("temp_message_tokens: %s, max_tokens: %s", temp_message_tokens, max_tokens) - if temp_message_tokens <= max_tokens: - return temp_messages - - # if temp_message_tokens > max_tokens, try shortening temp_messages - elif "function_call" not in message: - verbose_logger.debug("attempting to shorten message to fit limit") - # fit updated_message to be within temp_message_tokens - max_tokens (aka the amount temp_message_tokens is greate than max_tokens) - updated_message: Final = shorten_message_to_fit_limit(message, available_tokens, model) - if can_add_message(updated_message, final_messages, max_tokens, model): - verbose_logger.debug("can add message, returning [updated_message] + final_messages") - return [updated_message] + final_messages - else: - verbose_logger.debug("cannot add message, returning final_messages") - return final_messages - - -def can_add_message(message, messages, max_tokens, model): - if get_token_count(messages + [message], model) <= max_tokens: - return True - return False - - -def get_token_count(messages, model): - return token_counter(model=model, messages=messages) - - -def shorten_message_to_fit_limit(message, tokens_needed, model: str | None, raise_error_on_max_limit: bool = False): - """ - Shorten a message to fit within a token limit by removing characters from the middle. - - Args: - message: The message to shorten - tokens_needed: The maximum number of tokens allowed - model: The model being used (optional) - raise_error_on_max_limit: If True, raises an error when max attempts reached. If False, returns final trimmed content. - """ - - # For OpenAI models, even blank messages cost 7 token, - # and if the buffer is less than 3, the while loop will never end, - # hence the value 10. - if model is not None and "gpt" in model and tokens_needed <= 10: - return message - - content = message["content"] - attempts = 0 - - verbose_logger.debug("content: %s", content) - - while attempts < MAX_TOKEN_TRIMMING_ATTEMPTS: - verbose_logger.debug("getting token count for message: %s", message) - total_tokens = get_token_count([message], model) - verbose_logger.debug("total_tokens: %s, tokens_needed: %s", total_tokens, tokens_needed) - - if total_tokens <= tokens_needed: - break - - ratio = (tokens_needed) / total_tokens - - new_length = int(len(content) * ratio) - 1 - new_length = max(0, new_length) - - half_length = new_length // 2 - left_half = content[:half_length] - right_half = content[-half_length:] - - trimmed_content = left_half + ".." + right_half - message["content"] = trimmed_content - verbose_logger.debug("trimmed_content: %s", trimmed_content) - content = trimmed_content - attempts += 1 - - if attempts >= MAX_TOKEN_TRIMMING_ATTEMPTS and raise_error_on_max_limit: - raise Exception( - f"Failed to trim message to fit within {tokens_needed} tokens after {MAX_TOKEN_TRIMMING_ATTEMPTS} attempts" - ) - - return message - - -# LiteLLM token trimmer -# this code is borrowed from https://github.com/KillianLucas/tokentrim/blob/main/tokentrim/tokentrim.py -# Credits for this code go to Killian Lucas -def trim_messages( - messages, - model: str | None = None, - trim_ratio: float = DEFAULT_TRIM_RATIO, - return_response_tokens: bool = False, - max_tokens=None, -): - """ - Trim a list of messages to fit within a model's token limit. - - Args: - messages: Input messages to be trimmed. Each message is a dictionary with 'role' and 'content'. - model: The LiteLLM model being used (determines the token limit). - trim_ratio: Target ratio of tokens to use after trimming. Default is 0.75, meaning it will trim messages so they use about 75% of the model's token limit. - return_response_tokens: If True, also return the number of tokens left available for the response after trimming. - max_tokens: Instead of specifying a model or trim_ratio, you can specify this directly. - - Returns: - Trimmed messages and optionally the number of tokens available for response. - """ - # Initialize max_tokens - # if users pass in max tokens, trim to this amount - original_messages: Final = messages - messages = copy.deepcopy(messages) - try: - if max_tokens is None: - # Check if model is valid - if model in litellm.model_cost: - max_tokens_for_model: Final = litellm.model_cost[model].get( - "max_input_tokens", litellm.model_cost[model]["max_tokens"] - ) - max_tokens = int(max_tokens_for_model * trim_ratio) - else: - # if user did not specify max (input) tokens - # or passed an llm litellm does not know - # do nothing, just return messages - return messages - - system_message = "" - for message in messages: - if message["role"] == "system": - system_message += "\n" if system_message else "" - system_message += message["content"] - - ## Handle Tool Call ## - check if last message is a tool response, return as is - https://github.com/BerriAI/litellm/issues/4931 - tool_messages: Final = [] - - for message in reversed(messages): - if message["role"] != "tool": - break - tool_messages.append(message) - tool_messages.reverse() - # # Remove the collected tool messages from the original list - if len(tool_messages): - messages = messages[: -len(tool_messages)] - - current_tokens: Final = token_counter(model=model or "", messages=messages) - print_verbose(f"Current tokens: {current_tokens}, max tokens: {max_tokens}") - - # Do nothing if current tokens under messages - if current_tokens < max_tokens: - return messages + tool_messages - - #### Trimming messages if current_tokens > max_tokens - print_verbose( - f"Need to trim input messages: {messages}, current_tokens{current_tokens}, max_tokens: {max_tokens}" - ) - system_message_event: dict | None = None - if system_message: - system_message_event, max_tokens = process_system_message( - system_message=system_message, max_tokens=max_tokens, model=model - ) - - if max_tokens == 0: # the system messages are too long - return [system_message_event] - - # Since all system messages are combined and trimmed to fit the max_tokens, - # we remove all system messages from the messages list - messages = [message for message in messages if message["role"] != "system"] - - verbose_logger.debug("Processed system message: %s", system_message_event) - final_messages = process_messages(messages=messages, max_tokens=max_tokens, model=model) - verbose_logger.debug("Processed messages: %s", final_messages) - - # Add system message to the beginning of the final messages - if system_message_event: - final_messages = [system_message_event] + final_messages - - if len(tool_messages) > 0: - final_messages.extend(tool_messages) - - verbose_logger.debug("Final messages: %s, return_response_tokens: %s", final_messages, return_response_tokens) - if return_response_tokens: # if user wants token count with new trimmed messages - response_tokens: Final = max_tokens - get_token_count(final_messages, model) - return final_messages, response_tokens - return final_messages - except Exception as e: # [NON-Blocking, if error occurs just return final_messages - verbose_logger.exception("Got exception while token trimming - %s", e) - return original_messages - - -from litellm.caching.in_memory_cache import InMemoryCache - - -class AvailableModelsCache(InMemoryCache): - def __init__(self, ttl_seconds: int = 300, max_size: int = 1000): - super().__init__(ttl_seconds, max_size) - self._env_hash: str | None = None - - def _get_env_hash(self) -> str: - """Create a hash of relevant environment variables""" - env_vars: Final = {k: v for k, v in os.environ.items() if k.startswith(("OPENAI", "ANTHROPIC", "AZURE", "AWS"))} - return str(hash(frozenset(env_vars.items()))) - - def _check_env_changed(self) -> bool: - """Check if environment variables have changed""" - current_hash: Final = self._get_env_hash() - if self._env_hash is None: - self._env_hash = current_hash - return True - return current_hash != self._env_hash - - def _get_cache_key( - self, - custom_llm_provider: str | None, - litellm_params: LiteLLM_Params | None, - ) -> str: - valid_str = "" - - if litellm_params is not None: - valid_str = litellm_params.model_dump_json() - if custom_llm_provider is not None: - valid_str = f"{custom_llm_provider}:{valid_str}" - return hashlib.sha256(valid_str.encode()).hexdigest() - - def get_cached_model_info( - self, - custom_llm_provider: str | None = None, - litellm_params: LiteLLM_Params | None = None, - ) -> list[str] | None: - """Get cached model info""" - # Check if environment has changed - if litellm_params is None and self._check_env_changed(): - self.cache_dict.clear() - return None - - cache_key: Final = self._get_cache_key(custom_llm_provider, litellm_params) - - result: Final = cast(list[str] | None, self.get_cache(cache_key)) - - if result is not None: - return copy.deepcopy(result) - return result - - def set_cached_model_info( - self, - custom_llm_provider: str, - litellm_params: LiteLLM_Params | None, - available_models: list[str], - ): - """Set cached model info""" - cache_key: Final = self._get_cache_key(custom_llm_provider, litellm_params) - self.set_cache(cache_key, copy.deepcopy(available_models)) - - -# Global cache instance -_model_cache: Final = AvailableModelsCache() - - -def _infer_valid_provider_from_env_vars( - custom_llm_provider: str | None = None, -) -> list[str]: - valid_providers: Final[list[str]] = [] - environ_keys: Final = os.environ.keys() - for provider in litellm.provider_list: - if custom_llm_provider and provider != custom_llm_provider: - continue - - # edge case litellm has together_ai as a provider, it should be togetherai - env_provider_1 = provider.replace("_", "") - env_provider_2 = provider - - # litellm standardizes expected provider keys to - # PROVIDER_API_KEY. Example: OPENAI_API_KEY, COHERE_API_KEY - expected_provider_key_1 = f"{env_provider_1.upper()}_API_KEY" - expected_provider_key_2 = f"{env_provider_2.upper()}_API_KEY" - if expected_provider_key_1 in environ_keys or expected_provider_key_2 in environ_keys: - # key is set - valid_providers.append(provider) - - return valid_providers - - -def _get_valid_models_from_provider_api( - provider_config: BaseLLMModelInfo, - custom_llm_provider: str, - litellm_params: LiteLLM_Params | None = None, -) -> list[str]: - try: - cached_result: Final = _model_cache.get_cached_model_info(custom_llm_provider, litellm_params) - - if cached_result is not None: - return cached_result - models: Final = provider_config.get_models( - api_key=litellm_params.api_key if litellm_params is not None else None, - api_base=litellm_params.api_base if litellm_params is not None else None, - ) - - _model_cache.set_cached_model_info(custom_llm_provider, litellm_params, models) - return models - except Exception as e: - verbose_logger.warning("Error getting valid models: %s", e) - return [] - - -def get_valid_models( - check_provider_endpoint: bool | None = None, - custom_llm_provider: str | None = None, - litellm_params: LiteLLM_Params | None = None, - api_key: str | None = None, - api_base: str | None = None, -) -> list[str]: - """ - Returns a list of valid LLMs based on the set environment variables - - Args: - check_provider_endpoint: If True, will check the provider's endpoint for valid models. - custom_llm_provider: If provided, will only check the provider's endpoint for valid models. - api_key: If provided, will use the API key to get valid models. - api_base: If provided, will use the API base to get valid models. - Returns: - A list of valid LLMs - """ - - try: - ################################ - # init litellm_params - ################################# - from litellm.types.router import LiteLLM_Params - - if litellm_params is None: - litellm_params = LiteLLM_Params(model="") - if api_key is not None: - litellm_params.api_key = api_key - if api_base is not None: - litellm_params.api_base = api_base - ################################# - - check_provider_endpoint = check_provider_endpoint or litellm.check_provider_endpoint - # get keys set in .env - - valid_providers: list[str] = [] - valid_models: Final[list[str]] = [] - # for all valid providers, make a list of supported llms - - if custom_llm_provider: - valid_providers = [custom_llm_provider] - else: - valid_providers = _infer_valid_provider_from_env_vars(custom_llm_provider) - - for provider in valid_providers: - provider_config = ProviderConfigManager.get_provider_model_info( - model=None, - provider=LlmProviders(provider), - ) - - if custom_llm_provider and provider != custom_llm_provider: - continue - - if provider == "azure": - valid_models.append("Azure-LLM") - elif provider_config is not None and check_provider_endpoint and provider is not None: - valid_models.extend( - _get_valid_models_from_provider_api( - provider_config, - provider, - litellm_params, - ) - ) - else: - models_for_provider = copy.deepcopy(litellm.models_by_provider.get(provider, [])) - valid_models.extend(models_for_provider) - - return valid_models - except Exception as e: - verbose_logger.warning("Error getting valid models: %s", e) - return [] # NON-Blocking - - -def print_args_passed_to_litellm(original_function, args, kwargs): - if not _is_debugging_on(): - return - try: - # we've already printed this for acompletion, don't print for completion - if ( - "acompletion" in kwargs - and kwargs["acompletion"] is True - and original_function.__name__ == "completion" - or "aembedding" in kwargs - and kwargs["aembedding"] is True - and original_function.__name__ == "embedding" - or ( - "aimg_generation" in kwargs - and kwargs["aimg_generation"] is True - and original_function.__name__ == "img_generation" - ) - ): - return - - args_str: Final = ", ".join(map(repr, args)) - redacted_kwargs: Final = redact_credentials_in_payload(kwargs) - kwargs_str: Final = ", ".join(f"{key}={value!r}" for key, value in redacted_kwargs.items()) - print_verbose( - "\n", - ) # new line before - print_verbose( - "\033[92mRequest to litellm:\033[0m", - ) - if args and kwargs: - print_verbose(f"\033[92mlitellm.{original_function.__name__}({args_str}, {kwargs_str})\033[0m") - elif args: - print_verbose(f"\033[92mlitellm.{original_function.__name__}({args_str})\033[0m") - elif kwargs: - print_verbose(f"\033[92mlitellm.{original_function.__name__}({kwargs_str})\033[0m") - else: - print_verbose(f"\033[92mlitellm.{original_function.__name__}()\033[0m") - print_verbose("\n") # new line after - except Exception: - # This should always be non blocking - pass - - -def get_logging_id(start_time, response_obj): - try: - response_id: Final = "time-" + start_time.strftime("%H-%M-%S-%f") + "_" + response_obj.get("id") - return response_id - except Exception: - return None - - -def _get_base_model_from_metadata(model_call_details=None): - if model_call_details is None: - return None - litellm_params: Final = model_call_details.get("litellm_params", {}) - if litellm_params is not None: - _base_model: Final = litellm_params.get("base_model", None) - if _base_model is not None: - return _base_model - metadata: Final = litellm_params.get("metadata") or {} - - _get_base_model_from_litellm_call_metadata: Callable[..., str | None] = getattr( - sys.modules[__name__], "_get_base_model_from_litellm_call_metadata" - ) - base_model_from_metadata: Final = _get_base_model_from_litellm_call_metadata(metadata=metadata) - if base_model_from_metadata is not None: - return base_model_from_metadata - - # Also check litellm_metadata (used by Responses API and other generic API calls) - litellm_metadata: Final = litellm_params.get("litellm_metadata", {}) - _get_base_model_from_litellm_call_metadata = getattr( - sys.modules[__name__], "_get_base_model_from_litellm_call_metadata" - ) - return _get_base_model_from_litellm_call_metadata(metadata=litellm_metadata) - return None - - -class ModelResponseIterator: - def __init__(self, model_response: ModelResponse, convert_to_delta: bool = False): - if convert_to_delta is True: - _stream_response: Final = ModelResponseStream() - _stream_response.choices[0].delta.content = model_response.choices[0].message.content - self.model_response: ModelResponse | ModelResponseStream = _stream_response - else: - self.model_response = model_response - self.is_done = False - - # Sync iterator - def __iter__(self): - return self - - def __next__(self): - if self.is_done: - raise StopIteration - self.is_done = True - return self.model_response - - # Async iterator - def __aiter__(self): - return self - - async def __anext__(self): - if self.is_done: - raise StopAsyncIteration - self.is_done = True - return self.model_response - - -class ModelResponseListIterator: - def __init__(self, model_responses, delay: float | None = None): - self.model_responses = model_responses - self.index = 0 - self.delay = delay - - # Sync iterator - def __iter__(self): - return self - - def __next__(self): - if self.index >= len(self.model_responses): - raise StopIteration - model_response: Final = self.model_responses[self.index] - self.index += 1 - if self.delay: - time.sleep(self.delay) - return model_response - - # Async iterator - def __aiter__(self): - return self - - async def __anext__(self): - if self.index >= len(self.model_responses): - raise StopAsyncIteration - model_response: Final = self.model_responses[self.index] - self.index += 1 - if self.delay: - await asyncio.sleep(self.delay) - return model_response - - -class CustomModelResponseIterator(Iterable): - def __init__(self) -> None: - super().__init__() - - -def is_cached_message(message: AllMessageValues) -> bool: - """ - Returns true, if message is marked as needing to be cached. - - Used for anthropic/gemini context caching. - - Follows the anthropic format {"cache_control": {"type": "ephemeral"}} - - Can be disabled globally by setting litellm.disable_anthropic_gemini_context_caching_transform = True - """ - # Check if context caching is disabled globally - if litellm.disable_anthropic_gemini_context_caching_transform is True: - return False - - # Check message-level cache_control (set by cache_control_injection_points hook for string content) - message_level_cache_control: Final = message.get("cache_control") - if ( - message_level_cache_control is not None - and isinstance(message_level_cache_control, dict) - and message_level_cache_control.get("type") == "ephemeral" - ): - return True - - if "content" not in message: - return False - - content: Final = message["content"] - - # Handle non-list content types (None, str, etc.) - if not isinstance(content, list): - return False - - for content_item in content: - # Ensure content_item is a dictionary before accessing keys - if not isinstance(content_item, dict): - continue - - cache_control = content_item.get("cache_control") - if ( - content_item.get("type") == "text" - and cache_control is not None - and isinstance(cache_control, dict) - and cache_control.get("type") == "ephemeral" - ): - return True - - return False - - -def is_base64_encoded(s: str) -> bool: - try: - # Strip out the prefix if it exists - if not s.startswith( - "data:" - ): # require `data:` for base64 str, like openai. Prevents false positives like s='Dog' - return False - - s = s.split(",")[1] - - # Try to decode the string - decoded_bytes: Final = base64.b64decode(s, validate=True) - - # Check if the original string can be re-encoded to the same string - return base64.b64encode(decoded_bytes).decode("utf-8") == s - except Exception: - return False - - -def get_base64_str(s: str) -> str: - """ - s: b64str OR data:image/png;base64,b64str - """ - if "," in s: - return s.split(",")[1] - return s - - -def has_tool_call_blocks(messages: list[AllMessageValues]) -> bool: - """ - Returns true, if messages has tool call blocks. - - Used for anthropic/bedrock message validation. - """ - for message in messages: - if message.get("tool_calls") is not None: - return True - return False - - -def any_assistant_message_has_thinking_blocks( - messages: list[AllMessageValues], -) -> bool: - """ - Returns true if ANY assistant message has thinking_blocks. - - This is used to prevent dropping the thinking param when some messages - in the conversation already contain thinking blocks. Dropping thinking - when thinking blocks exist causes Anthropic error: - "When thinking is disabled, an assistant message cannot contain thinking" - - Related issue: https://github.com/BerriAI/litellm/issues/18926 - """ - for message in messages: - if message.get("role") == "assistant": - thinking_blocks = message.get("thinking_blocks") - if thinking_blocks is not None and (not hasattr(thinking_blocks, "__len__") or len(thinking_blocks) > 0): - return True - return False - - -def last_assistant_with_tool_calls_has_no_thinking_blocks( - messages: list[AllMessageValues], -) -> bool: - """ - Returns true if the last assistant message with tool_calls has no thinking_blocks. - - This is used to detect when thinking param should be dropped to avoid - Anthropic error: "Expected thinking or redacted_thinking, but found tool_use" - - When thinking is enabled, assistant messages with tool_calls must include thinking_blocks. - If the client didn't preserve thinking_blocks, we need to drop the thinking param. - - IMPORTANT: This should only be used in conjunction with - any_assistant_message_has_thinking_blocks() to ensure we don't drop thinking - when other messages in the conversation contain thinking blocks. - - Related issues: https://github.com/BerriAI/litellm/issues/14194, https://github.com/BerriAI/litellm/issues/9020 - """ - # Find the last assistant message with tool_calls - last_assistant_with_tools = None - for message in messages: - if message.get("role") == "assistant" and message.get("tool_calls") is not None: - last_assistant_with_tools = message - - if last_assistant_with_tools is None: - return False - - # Check if it has thinking_blocks - thinking_blocks: Final = last_assistant_with_tools.get("thinking_blocks") - return thinking_blocks is None or (hasattr(thinking_blocks, "__len__") and len(thinking_blocks) == 0) - - -def add_dummy_tool(custom_llm_provider: str) -> list[ChatCompletionToolParam]: - """ - Prevent Anthropic from raising error when tool_use block exists but no tools are provided. - - Relevent Issues: https://github.com/BerriAI/litellm/issues/5388, https://github.com/BerriAI/litellm/issues/5747 - """ - return [ - ChatCompletionToolParam( - type="function", - function=ChatCompletionToolParamFunctionChunk( - name="dummy_tool", - description="This is a dummy tool call", # provided to satisfy bedrock constraint. - parameters={ - "type": "object", - "properties": {}, - }, - ), - ) - ] - - -from litellm.types.llms.openai import ( - ChatCompletionAudioObject, - ChatCompletionImageObject, - ChatCompletionTextObject, - ChatCompletionUserMessage, - OpenAIMessageContent, - ValidUserMessageContentTypes, -) - - -def convert_to_dict(message: BaseModel | dict) -> dict: - """ - Converts a message to a dictionary if it's a Pydantic model. - - Args: - message: The message, which may be a Pydantic model or a dictionary. - - Returns: - dict: The converted message. - """ - if isinstance(message, BaseModel): - return message.model_dump(exclude_none=True) - elif isinstance(message, dict): - return message - else: - raise TypeError(f"Invalid message type: {type(message)}. Expected dict or Pydantic model.") - - -def convert_list_message_to_dict(messages: Sequence): - new_messages: Final = [] - for message in messages: - convert_msg_to_dict = cast(AllMessageValues, convert_to_dict(message)) - cleaned_message = cleanup_none_field_in_message(message=convert_msg_to_dict) - new_messages.append(cleaned_message) - return new_messages - - -def validate_and_fix_openai_messages(messages: list): - """ - Ensures all messages are valid OpenAI chat completion messages. - - Handles missing role for assistant messages. - """ - new_messages: Final = [] - for message in messages: - if not message.get("role"): - message["role"] = "assistant" - if message.get("tool_calls"): - message["tool_calls"] = jsonify_tools(tools=message["tool_calls"]) - - convert_msg_to_dict = cast(AllMessageValues, convert_to_dict(message)) - cleaned_message = cleanup_none_field_in_message(message=convert_msg_to_dict) - new_messages.append(cleaned_message) - return validate_chat_completion_user_messages(messages=new_messages) - - -def validate_and_fix_openai_tools(tools: list | None) -> list[dict] | None: - """ - Ensure tools is List[dict] and not List[BaseModel] - """ - new_tools: Final = [] - if tools is None: - return tools - for tool in tools: - if isinstance(tool, BaseModel): - new_tools.append(tool.model_dump()) - elif isinstance(tool, dict): - new_tools.append(tool) - return new_tools - - -def validate_and_fix_thinking_param( - thinking: AnthropicThinkingParam | bool | None, -) -> AnthropicThinkingParam | None: - """ - Coerces bool thinking values (True becomes enabled with the default medium budget, False becomes None) - and normalizes camelCase keys in the thinking param to snake_case. - Handles clients that send budgetTokens instead of budget_tokens. - """ - if thinking is True: - return cast( - "AnthropicThinkingParam", - {"type": "enabled", "budget_tokens": DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET}, - ) - if thinking is False: - return None - if thinking is None or not isinstance(thinking, dict): - return thinking - normalized: Final = dict(thinking) - if "budgetTokens" in normalized and "budget_tokens" not in normalized: - normalized["budget_tokens"] = normalized.pop("budgetTokens") - elif "budgetTokens" in normalized and "budget_tokens" in normalized: - normalized.pop("budgetTokens") - return cast("AnthropicThinkingParam", normalized) - - -def cleanup_none_field_in_message(message: AllMessageValues): - """ - Cleans up the message by removing the none field. - - remove None fields in the message - e.g. {"function": None} - some providers raise validation errors - """ - new_message: Final = message.copy() - return {k: v for k, v in new_message.items() if v is not None} - - -def validate_chat_completion_user_messages(messages: list[AllMessageValues]): - """ - Ensures all user messages are valid OpenAI chat completion messages. - - Args: - messages: List of message dictionaries - message_content_type: Type to validate content against - - Returns: - List[dict]: The validated messages - - Raises: - ValueError: If any message is invalid - """ - for idx, m in enumerate(messages): - try: - if m["role"] == "user": - user_content = m.get("content") - if user_content is not None: - if isinstance(user_content, str): - continue - elif isinstance(user_content, list): - for item in user_content: - if isinstance(item, dict): - if item.get("type") not in ValidUserMessageContentTypes: - raise Exception(f"invalid content type={item.get('type')}") - except Exception as e: - if isinstance(e, KeyError): - raise Exception( - f"Invalid message at index {idx}. Please ensure all messages are valid OpenAI chat completion messages." - ) - if "invalid content type" in str(e): - raise Exception( - f"Invalid user message at index {idx}. Please ensure all user messages are valid OpenAI chat completion messages." - ) - else: - raise e - - return messages - - -def validate_chat_completion_tool_choice( - tool_choice: dict | str | None, -) -> dict | str | None: - """ - Confirm the tool choice is passed in the OpenAI format. - - Prevents user errors like: https://github.com/BerriAI/litellm/issues/7483 - """ - if tool_choice is None or isinstance(tool_choice, str): - return tool_choice - elif isinstance(tool_choice, dict): - tool_choice_type = tool_choice.get("type") - if tool_choice_type in ("auto", "none", "required") and "function" not in tool_choice: - return tool_choice_type - - # Standard OpenAI format: {"type": "function", "function": {...}} - if tool_choice.get("type") is None or tool_choice.get("function") is None: - raise Exception( - f"Invalid tool choice, tool_choice={tool_choice}. Please ensure tool_choice follows the OpenAI spec" - ) - return tool_choice - raise Exception( - f"Invalid tool choice, tool_choice={tool_choice}. Got={type(tool_choice)}. Expecting str, or dict. Please ensure tool_choice follows the OpenAI tool_choice spec" - ) - - -def validate_openai_optional_params(stop: str | list[str] | None = None, **kwargs) -> str | list[str] | None: - """ - Validates and fixes OpenAI optional parameters. - - Args: - stop: Stop sequences (string or list of strings) - **kwargs: Additional optional parameters - - Returns: - Validated stop parameter (truncated to 4 elements if needed) - """ - if stop is not None and isinstance(stop, list) and not litellm.disable_stop_sequence_limit: - # Truncate to 4 elements if more are provided as openai only supports up to 4 stop sequences - if len(stop) > 4: - stop = stop[:4] - - return stop - - -@lru_cache(maxsize=1) -def _get_bundled_model_cost_map() -> dict[str, Any]: - try: - model_cost_path: Final = resources.files("litellm").joinpath("model_prices_and_context_window_backup.json") - return json.loads(model_cost_path.read_text()) - except Exception: - return {} - - -def _get_model_cost_entry_for_provider_config( - model: str, - provider: LlmProviders, -) -> dict[str, Any]: - candidate_keys: Final = (model, f"{provider.value}/{model}") - for model_key in candidate_keys: - model_info = litellm.model_cost.get(model_key) - if model_info is not None: - return model_info - - bundled_model_cost: Final = _get_bundled_model_cost_map() - for model_key in candidate_keys: - model_info = bundled_model_cost.get(model_key) - if model_info is not None: - return model_info - return {} - - -class ProviderConfigManager: - # Dictionary mapping for O(1) provider lookup - # Stores tuples of (factory_function, needs_model_parameter) - # This is initialized lazily on first access to avoid circular imports - _PROVIDER_CONFIG_MAP: dict[LlmProviders, tuple[Callable, bool]] | None = None - - @staticmethod - def _build_provider_config_map() -> dict[LlmProviders, tuple[Callable, bool]]: - """Build the provider-to-config mapping dictionary. - - Returns a dict mapping provider to (factory_function, needs_model_parameter). - This avoids expensive inspect.signature() calls at runtime. - """ - return { - # Most common providers first for readability - # Format: (factory_function, needs_model_parameter: bool) - LlmProviders.OPENAI: (lambda: litellm.OpenAIGPTConfig(), False), - LlmProviders.ANTHROPIC: (lambda: litellm.AnthropicConfig(), False), - # AZURE is handled as a special case in get_provider_chat_config() - # so that base_model can be threaded through for model-type detection. - LlmProviders.AZURE_AI: ( - lambda model: ProviderConfigManager._get_azure_ai_config(model), - True, - ), - LlmProviders.VERTEX_AI: ( - lambda model: ProviderConfigManager._get_vertex_ai_config(model), - True, - ), - LlmProviders.BEDROCK: ( - lambda model: ProviderConfigManager._get_bedrock_config(model), - True, - ), - LlmProviders.COHERE: ( - lambda model: ProviderConfigManager._get_cohere_config(model), - True, - ), - LlmProviders.COHERE_CHAT: ( - lambda model: ProviderConfigManager._get_cohere_config(model), - True, - ), - # Simple provider mappings (no model parameter needed) - LlmProviders.DEEPSEEK: (lambda: litellm.DeepSeekChatConfig(), False), - LlmProviders.TENCENT: (lambda: litellm.TencentChatConfig(), False), - LlmProviders.GROQ: (lambda: litellm.GroqChatConfig(), False), - LlmProviders.BEDROCK_MANTLE: ( - lambda: litellm.BedrockMantleChatConfig(), - False, - ), - LlmProviders.A2A: (lambda: litellm.A2AConfig(), False), - LlmProviders.BYTEZ: (lambda: litellm.BytezChatConfig(), False), - LlmProviders.DATABRICKS: (lambda: litellm.DatabricksConfig(), False), - LlmProviders.XAI: (lambda: litellm.XAIChatConfig(), False), - LlmProviders.ZAI: (lambda: litellm.ZAIChatConfig(), False), - LlmProviders.LAMBDA_AI: (lambda: litellm.LambdaAIChatConfig(), False), - LlmProviders.INCEPTION: (lambda: litellm.InceptionChatConfig(), False), - LlmProviders.LLAMA: (lambda: litellm.LlamaAPIConfig(), False), - LlmProviders.TEXT_COMPLETION_OPENAI: ( - lambda: litellm.OpenAITextCompletionConfig(), - False, - ), - LlmProviders.SNOWFLAKE: (lambda: litellm.SnowflakeConfig(), False), - LlmProviders.CLARIFAI: (lambda: litellm.ClarifaiConfig(), False), - LlmProviders.ANTHROPIC_TEXT: (lambda: litellm.AnthropicTextConfig(), False), - LlmProviders.VERTEX_AI_BETA: (lambda: litellm.VertexGeminiConfig(), False), - LlmProviders.CLOUDFLARE: (lambda: litellm.CloudflareChatConfig(), False), - LlmProviders.SAGEMAKER_CHAT: (lambda: litellm.SagemakerChatConfig(), False), - LlmProviders.SAGEMAKER_NOVA: (lambda: litellm.SagemakerNovaConfig(), False), - LlmProviders.SAGEMAKER: (lambda: litellm.SagemakerConfig(), False), - LlmProviders.FIREWORKS_AI: (lambda: litellm.FireworksAIConfig(), False), - LlmProviders.FRIENDLIAI: (lambda: litellm.FriendliaiChatConfig(), False), - LlmProviders.WATSONX: (lambda: litellm.IBMWatsonXChatConfig(), False), - LlmProviders.WATSONX_TEXT: (lambda: litellm.IBMWatsonXAIConfig(), False), - LlmProviders.EMPOWER: (lambda: litellm.EmpowerChatConfig(), False), - LlmProviders.MINIMAX: (lambda: litellm.MinimaxChatConfig(), False), - LlmProviders.GITHUB: (lambda: litellm.GithubChatConfig(), False), - LlmProviders.COMPACTIFAI: (lambda: litellm.CompactifAIChatConfig(), False), - LlmProviders.GITHUB_COPILOT: (lambda: litellm.GithubCopilotConfig(), False), - LlmProviders.CHATGPT: (lambda: litellm.ChatGPTConfig(), False), - LlmProviders.GIGACHAT: (lambda: litellm.GigaChatConfig(), False), - LlmProviders.RAGFLOW: (lambda: litellm.RAGFlowConfig(), False), - LlmProviders.CUSTOM: (lambda: litellm.OpenAILikeChatConfig(), False), - LlmProviders.CUSTOM_OPENAI: (lambda: litellm.OpenAILikeChatConfig(), False), - LlmProviders.OPENAI_LIKE: (lambda: litellm.OpenAILikeChatConfig(), False), - LlmProviders.AIOHTTP_OPENAI: ( - lambda: litellm.AiohttpOpenAIChatConfig(), - False, - ), - LlmProviders.HOSTED_VLLM: (lambda: litellm.HostedVLLMChatConfig(), False), - LlmProviders.LLAMAFILE: (lambda: litellm.LlamafileChatConfig(), False), - LlmProviders.LM_STUDIO: (lambda: litellm.LMStudioChatConfig(), False), - LlmProviders.GALADRIEL: (lambda: litellm.GaladrielChatConfig(), False), - LlmProviders.REPLICATE: (lambda: litellm.ReplicateConfig(), False), - LlmProviders.HUGGINGFACE: (lambda: litellm.HuggingFaceChatConfig(), False), - LlmProviders.TOGETHER_AI: (lambda: litellm.TogetherAIChatConfig(), False), - LlmProviders.OPENROUTER: (lambda: litellm.OpenrouterConfig(), False), - LlmProviders.VERCEL_AI_GATEWAY: ( - lambda: litellm.VercelAIGatewayConfig(), - False, - ), - LlmProviders.COMETAPI: (lambda: litellm.CometAPIConfig(), False), - LlmProviders.DATAROBOT: (lambda: litellm.DataRobotConfig(), False), - LlmProviders.GEMINI: (lambda: litellm.GoogleAIStudioGeminiConfig(), False), - LlmProviders.AI21: (lambda: litellm.AI21ChatConfig(), False), - LlmProviders.AI21_CHAT: (lambda: litellm.AI21ChatConfig(), False), - LlmProviders.AZURE_TEXT: (lambda: litellm.AzureOpenAITextConfig(), False), - LlmProviders.NLP_CLOUD: (lambda: litellm.NLPCloudConfig(), False), - LlmProviders.OOBABOOGA: (lambda: litellm.OobaboogaConfig(), False), - LlmProviders.OLLAMA_CHAT: (lambda: litellm.OllamaChatConfig(), False), - LlmProviders.DEEPINFRA: (lambda: litellm.DeepInfraConfig(), False), - LlmProviders.PERPLEXITY: (lambda: litellm.PerplexityChatConfig(), False), - LlmProviders.MISTRAL: (lambda: litellm.MistralConfig(), False), - LlmProviders.CODESTRAL: (lambda: litellm.MistralConfig(), False), - LlmProviders.NVIDIA_NIM: (lambda: litellm.NvidiaNimConfig(), False), - LlmProviders.CEREBRAS: (lambda: litellm.CerebrasConfig(), False), - LlmProviders.BASETEN: (lambda: litellm.BasetenConfig(), False), - LlmProviders.VOLCENGINE: (lambda: litellm.VolcEngineConfig(), False), - LlmProviders.TEXT_COMPLETION_CODESTRAL: ( - lambda: litellm.CodestralTextCompletionConfig(), - False, - ), - LlmProviders.TEXT_COMPLETION_INCEPTION: ( - lambda: litellm.InceptionTextCompletionConfig(), - False, - ), - LlmProviders.SAMBANOVA: (lambda: litellm.SambanovaConfig(), False), - LlmProviders.MARITALK: (lambda: litellm.MaritalkConfig(), False), - LlmProviders.VLLM: (lambda: litellm.VLLMConfig(), False), - LlmProviders.OLLAMA: (lambda: litellm.OllamaConfig(), False), - LlmProviders.PREDIBASE: (lambda: litellm.PredibaseConfig(), False), - LlmProviders.TRITON: (lambda: litellm.TritonConfig(), False), - LlmProviders.PETALS: (lambda: litellm.PetalsConfig(), False), - LlmProviders.SAP_GENERATIVE_AI_HUB: ( - lambda: litellm.GenAIHubOrchestrationConfig(), - False, - ), - LlmProviders.FEATHERLESS_AI: (lambda: litellm.FeatherlessAIConfig(), False), - LlmProviders.NOVITA: (lambda: litellm.NovitaConfig(), False), - LlmProviders.NEBIUS: (lambda: litellm.NebiusConfig(), False), - LlmProviders.WANDB: (lambda: litellm.WandbConfig(), False), - LlmProviders.DASHSCOPE: (lambda: litellm.DashScopeChatConfig(), False), - LlmProviders.QWENCLOUD: (lambda: litellm.QwenCloudChatConfig(), False), - LlmProviders.QWEN_AI_PLATFORM: ( - lambda: litellm.QwenAIPlatformChatConfig(), - False, - ), - LlmProviders.MODELSCOPE: (lambda: litellm.ModelScopeChatConfig(), False), - LlmProviders.MOONSHOT: (lambda: litellm.MoonshotChatConfig(), False), - LlmProviders.DOCKER_MODEL_RUNNER: ( - lambda: litellm.DockerModelRunnerChatConfig(), - False, - ), - LlmProviders.V0: (lambda: litellm.V0ChatConfig(), False), - LlmProviders.MORPH: (lambda: litellm.MorphChatConfig(), False), - LlmProviders.LITELLM_PROXY: ( - lambda: litellm.LiteLLMProxyChatConfig(), - False, - ), - LlmProviders.GRADIENT_AI: (lambda: litellm.GradientAIConfig(), False), - LlmProviders.NSCALE: (lambda: litellm.NscaleConfig(), False), - LlmProviders.HEROKU: (lambda: litellm.HerokuChatConfig(), False), - LlmProviders.OCI: (lambda: litellm.OCIChatConfig(), False), - LlmProviders.HYPERBOLIC: (lambda: litellm.HyperbolicChatConfig(), False), - LlmProviders.OVHCLOUD: (lambda: litellm.OVHCloudChatConfig(), False), - LlmProviders.AMAZON_NOVA: (lambda: litellm.AmazonNovaChatConfig(), False), - LlmProviders.LANGGRAPH: ( - lambda: ProviderConfigManager._get_langgraph_config(), - False, - ), - LlmProviders.LANGFLOW: ( - lambda: ProviderConfigManager._get_langflow_config(), - False, - ), - LlmProviders.GDC: ( - lambda: litellm.GDCGeminiConfig(), - False, - ), - } - - @staticmethod - def _get_azure_config(model: str, base_model: str | None = None) -> BaseConfig: - """Get Azure config based on model type. - - When *base_model* is provided (e.g. ``"azure/gpt-5.2"``), it is used - for model-type detection instead of *model* (the deployment name). - This allows non-standard deployment names like ``"azure/foo"`` to be - routed through the correct config when the user specifies the true - underlying model via ``base_model``. - """ - detection_model: Final = base_model or model - if litellm.AzureOpenAIO1Config().is_o_series_model(model=detection_model): - return litellm.AzureOpenAIO1Config() - if litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model(model=detection_model): - return litellm.AzureOpenAIGPT5Config() - return litellm.AzureOpenAIConfig() - - @staticmethod - def _get_azure_ai_config(model: str) -> BaseConfig: - """Get Azure AI config based on model type.""" - from litellm.llms.azure_ai.common_utils import AzureFoundryModelInfo - - return AzureFoundryModelInfo.get_azure_ai_config_for_model(model) - - @staticmethod - def _get_vertex_ai_config(model: str) -> BaseConfig: - """Get Vertex AI config based on model type.""" - if "gemini" in model: - return litellm.VertexGeminiConfig() - elif "claude" in model: - return litellm.VertexAIAnthropicConfig() - elif "gpt-oss" in model: - from litellm.llms.vertex_ai.vertex_ai_partner_models.gpt_oss.transformation import ( - VertexAIGPTOSSTransformation, - ) - - return VertexAIGPTOSSTransformation() - elif model in litellm.vertex_mistral_models: - if "codestral" in model: - return litellm.CodestralTextCompletionConfig() - return litellm.MistralConfig() - elif model in litellm.vertex_ai_ai21_models: - return litellm.VertexAIAi21Config() - else: - return litellm.VertexAILlama3Config() - - @staticmethod - def _get_bedrock_config(model: str) -> BaseConfig: - """Get Bedrock config based on model.""" - from litellm.llms.bedrock.common_utils import get_bedrock_chat_config - - return get_bedrock_chat_config(model=model) - - @staticmethod - def _get_cohere_config(model: str) -> BaseConfig: - """Get Cohere config based on route.""" - CohereModelInfo: Final = litellm_utils.CohereModelInfo - route: Final = CohereModelInfo.get_cohere_route(model) - if route == "v2": - return litellm.CohereV2ChatConfig() - return litellm.CohereChatConfig() - - @staticmethod - def _get_langgraph_config() -> BaseConfig: - """Get LangGraph config.""" - from litellm.llms.langgraph.chat.transformation import LangGraphConfig - - return LangGraphConfig() - - @staticmethod - def _get_langflow_config() -> BaseConfig: - """Get LangFlow config.""" - from litellm.llms.langflow.chat.transformation import LangFlowConfig - - return LangFlowConfig() - - @staticmethod - def get_provider_chat_config( - model: str, - provider: LlmProviders, - base_model: str | None = None, - ) -> BaseConfig | None: - """ - Returns the provider config for a given provider. - - Uses O(1) dictionary lookup for fast provider resolution. - Python classes take priority over JSON (they have custom overrides). - - For Azure, *base_model* (when set) drives model-type detection so that - non-standard deployment names still route to the correct config. - """ - # Handle OpenAI special cases (O-series and GPT-5 models) - if provider == LlmProviders.OPENAI: - from litellm.llms.openai.chat.gpt_transformation import ( - OpenAIGPTConfig, - OpenAIUnknownModelConfig, - ) - - if litellm.openaiOSeriesConfig.is_model_o_series_model(model=model): - return litellm.openaiOSeriesConfig - if litellm.OpenAIGPT5Config.is_model_gpt_5_model(model=model): - return litellm.OpenAIGPT5Config() - if not OpenAIGPTConfig.is_openai_catalog_model(model): - return OpenAIUnknownModelConfig() - - # Handle Azure before the generic map so base_model can be threaded through - if provider == LlmProviders.AZURE: - return ProviderConfigManager._get_azure_config(model=model, base_model=base_model) - - # Initialize provider config map lazily (avoids circular imports) - if ProviderConfigManager._PROVIDER_CONFIG_MAP is None: - ProviderConfigManager._PROVIDER_CONFIG_MAP = ProviderConfigManager._build_provider_config_map() - - # O(1) dictionary lookup — Python classes first (custom overrides take priority) - config_entry: Final = ProviderConfigManager._PROVIDER_CONFIG_MAP.get(provider) - if config_entry is not None: - config_factory, needs_model = config_entry - if needs_model: - return config_factory(model) - else: - return config_factory() - - # Fall back to JSON providers (generic OpenAI-compatible) - from litellm.llms.openai_like.dynamic_config import create_config_class - from litellm.llms.openai_like.json_loader import JSONProviderRegistry - - if JSONProviderRegistry.exists(provider.value): - provider_config: Final = JSONProviderRegistry.get(provider.value) - if provider_config is None: - raise ValueError(f"Provider {provider.value} not found") - return create_config_class(provider_config)() - - return None - - @staticmethod - def get_provider_embedding_config( - model: str, - provider: LlmProviders, - ) -> BaseEmbeddingConfig | None: - if ( - litellm.LlmProviders.VOYAGE == provider - and litellm.VoyageContextualEmbeddingConfig.is_contextualized_embeddings(model) - ): - return litellm.VoyageContextualEmbeddingConfig() - elif ( - litellm.LlmProviders.VOYAGE == provider - and litellm.VoyageMultimodalEmbeddingConfig.is_multimodal_embeddings(model) - ): - return litellm.VoyageMultimodalEmbeddingConfig() - elif litellm.LlmProviders.VOYAGE == provider: - return litellm.VoyageEmbeddingConfig() - elif litellm.LlmProviders.TRITON == provider: - return litellm.TritonEmbeddingConfig() - elif litellm.LlmProviders.WATSONX == provider: - return litellm.IBMWatsonXEmbeddingConfig() - elif litellm.LlmProviders.SAP_GENERATIVE_AI_HUB == provider: - return litellm.GenAIHubEmbeddingConfig() - elif litellm.LlmProviders.INFINITY == provider: - return litellm.InfinityEmbeddingConfig() - elif litellm.LlmProviders.SAMBANOVA == provider: - return litellm.SambaNovaEmbeddingConfig() - elif litellm.LlmProviders.OCI == provider: - from litellm.llms.oci.embed.transformation import OCIEmbedConfig - - return OCIEmbedConfig() - elif litellm.LlmProviders.COHERE == provider or litellm.LlmProviders.COHERE_CHAT == provider: - from litellm.llms.cohere.embed.transformation import CohereEmbeddingConfig - - return CohereEmbeddingConfig() - elif litellm.LlmProviders.JINA_AI == provider: - from litellm.llms.jina_ai.embedding.transformation import ( - JinaAIEmbeddingConfig, - ) - - return JinaAIEmbeddingConfig() - elif litellm.LlmProviders.VOLCENGINE == provider: - from litellm.llms.volcengine.embedding.transformation import ( - VolcEngineEmbeddingConfig, - ) - - return VolcEngineEmbeddingConfig() - elif provider in ( - litellm.LlmProviders.DASHSCOPE, - litellm.LlmProviders.QWENCLOUD, - litellm.LlmProviders.QWEN_AI_PLATFORM, - ): - from litellm.llms.dashscope.common_utils import ( - get_dashscope_family_embedding_config, - ) - - return get_dashscope_family_embedding_config(provider.value) - elif litellm.LlmProviders.OVHCLOUD == provider: - return litellm.OVHCloudEmbeddingConfig() - elif litellm.LlmProviders.SNOWFLAKE == provider: - return litellm.SnowflakeEmbeddingConfig() - elif litellm.LlmProviders.COMETAPI == provider: - return litellm.CometAPIEmbeddingConfig() - elif litellm.LlmProviders.GITHUB_COPILOT == provider: - return litellm.GithubCopilotEmbeddingConfig() - elif litellm.LlmProviders.OPENROUTER == provider: - from litellm.llms.openrouter.embedding.transformation import ( - OpenrouterEmbeddingConfig, - ) - - return OpenrouterEmbeddingConfig() - elif litellm.LlmProviders.VERCEL_AI_GATEWAY == provider: - from litellm.llms.vercel_ai_gateway.embedding.transformation import ( - VercelAIGatewayEmbeddingConfig, - ) - - return VercelAIGatewayEmbeddingConfig() - elif litellm.LlmProviders.GIGACHAT == provider: - return litellm.GigaChatEmbeddingConfig() - elif litellm.LlmProviders.HOSTED_VLLM == provider: - return litellm.HostedVLLMEmbeddingConfig() - elif litellm.LlmProviders.SAGEMAKER == provider: - from litellm.llms.sagemaker.embedding.transformation import ( - SagemakerEmbeddingConfig, - ) - - return SagemakerEmbeddingConfig.get_model_config(model) - elif litellm.LlmProviders.PERPLEXITY == provider: - return litellm.PerplexityEmbeddingConfig() - return None - - @staticmethod - def get_provider_rerank_config( - model: str, - provider: LlmProviders, - api_base: str | None, - present_version_params: list[str], - ) -> BaseRerankConfig: - if litellm.LlmProviders.COHERE == provider or litellm.LlmProviders.COHERE_CHAT == provider: - if should_use_cohere_v1_client(api_base, present_version_params): - return litellm.CohereRerankConfig() - else: - return litellm.CohereRerankV2Config() - elif litellm.LlmProviders.AZURE_AI == provider: - return litellm.AzureAIRerankConfig() - elif litellm.LlmProviders.INFINITY == provider: - return litellm.InfinityRerankConfig() - elif litellm.LlmProviders.JINA_AI == provider: - return litellm.JinaAIRerankConfig() - elif litellm.LlmProviders.HOSTED_VLLM == provider: - return litellm.HostedVLLMRerankConfig() - elif litellm.LlmProviders.HUGGINGFACE == provider: - return litellm.HuggingFaceRerankConfig() - elif litellm.LlmProviders.DEEPINFRA == provider: - return litellm.DeepinfraRerankConfig() - elif litellm.LlmProviders.NVIDIA_NIM == provider: - from litellm.llms.nvidia_nim.rerank.common_utils import ( - get_nvidia_nim_rerank_config, - ) - - return get_nvidia_nim_rerank_config(model) - elif litellm.LlmProviders.VERTEX_AI == provider: - return litellm.VertexAIRerankConfig() - elif litellm.LlmProviders.FIREWORKS_AI == provider: - return litellm.FireworksAIRerankConfig() - elif litellm.LlmProviders.VOYAGE == provider: - return litellm.VoyageRerankConfig() - elif litellm.LlmProviders.WATSONX == provider: - return litellm.IBMWatsonXRerankConfig() - elif provider in ( - litellm.LlmProviders.DASHSCOPE, - litellm.LlmProviders.QWENCLOUD, - litellm.LlmProviders.QWEN_AI_PLATFORM, - ): - from litellm.llms.dashscope.common_utils import ( - get_dashscope_family_rerank_config, - ) - - return get_dashscope_family_rerank_config(provider.value) - return litellm.CohereRerankConfig() - - @staticmethod - def get_provider_anthropic_messages_config( - model: str, - provider: LlmProviders, - ) -> BaseAnthropicMessagesConfig | None: - return ProviderConfigManager._get_provider_anthropic_messages_config_cached(model=model, provider=provider) - - @staticmethod - @lru_cache(maxsize=DEFAULT_MAX_LRU_CACHE_SIZE) - def _get_provider_anthropic_messages_config_cached( - model: str, - provider: LlmProviders, - ) -> BaseAnthropicMessagesConfig | None: - model_lower: Final = model.lower() - if litellm.LlmProviders.ANTHROPIC == provider: - return litellm.AnthropicMessagesConfig() - # The 'BEDROCK' provider corresponds to Amazon's implementation of Anthropic Claude v3. - # This mapping ensures that the correct configuration is returned for BEDROCK. - elif litellm.LlmProviders.BEDROCK == provider: - from litellm.llms.bedrock.common_utils import BedrockModelInfo - - return BedrockModelInfo.get_bedrock_provider_config_for_messages_api(model) - elif litellm.LlmProviders.VERTEX_AI == provider: - if "claude" in model_lower: - from litellm.llms.vertex_ai.vertex_ai_partner_models.anthropic.experimental_pass_through.transformation import ( - VertexAIPartnerModelsAnthropicMessagesConfig, - ) - - return VertexAIPartnerModelsAnthropicMessagesConfig() - elif litellm.LlmProviders.AZURE_AI == provider: - if "claude" in model_lower: - from litellm.llms.azure_ai.anthropic.messages_transformation import ( - AzureAnthropicMessagesConfig, - ) - - return AzureAnthropicMessagesConfig() - elif litellm.LlmProviders.MINIMAX == provider: - from litellm.llms.minimax.messages.transformation import ( - MinimaxMessagesConfig, - ) - - return MinimaxMessagesConfig() - elif litellm.LlmProviders.DEEPSEEK == provider: - from litellm.llms.deepseek.messages.transformation import ( - DeepSeekAnthropicMessagesConfig, - ) - - return DeepSeekAnthropicMessagesConfig() - elif litellm.LlmProviders.TENCENT == provider: - from litellm.llms.tencent.messages.transformation import ( - TencentAnthropicMessagesConfig, - ) - - return TencentAnthropicMessagesConfig() - elif litellm.LlmProviders.GITHUB_COPILOT == provider: - if "claude" in model_lower: - from litellm.llms.github_copilot.messages.transformation import ( - GithubCopilotAnthropicMessagesConfig, - ) - - return GithubCopilotAnthropicMessagesConfig() - - from litellm.llms.openai_like.json_loader import JSONProviderRegistry - - json_provider: Final = JSONProviderRegistry.get(provider.value) - if json_provider is not None and "/v1/messages" in json_provider.supported_endpoints: - from litellm.llms.openai_like.messages.transformation import ( - JSONProviderAnthropicMessagesConfig, - ) - - return JSONProviderAnthropicMessagesConfig(json_provider) - return None - - @staticmethod - def get_provider_audio_transcription_config( - model: str, - provider: LlmProviders, - ) -> BaseAudioTranscriptionConfig | None: - model_cost_entry: Final = _get_model_cost_entry_for_provider_config( - model=model, - provider=provider, - ) - if ( - litellm.LlmProviders.AZURE == provider - and model_cost_entry.get("audio_transcription_config") == "azure_speech" - ): - from litellm.llms.azure.audio_transcription.transformation import ( - AzureSpeechAudioTranscriptionConfig, - ) - - return AzureSpeechAudioTranscriptionConfig() - elif litellm.LlmProviders.DEEPGRAM == provider: - return litellm.DeepgramAudioTranscriptionConfig() - elif litellm.LlmProviders.ELEVENLABS == provider: - from litellm.llms.elevenlabs.audio_transcription.transformation import ( - ElevenLabsAudioTranscriptionConfig, - ) - - return ElevenLabsAudioTranscriptionConfig() - elif litellm.LlmProviders.OPENAI == provider: - if "gpt-4o" in model: - return litellm.OpenAIGPTAudioTranscriptionConfig() - else: - return litellm.OpenAIWhisperAudioTranscriptionConfig() - elif litellm.LlmProviders.HOSTED_VLLM == provider: - from litellm.llms.hosted_vllm.transcriptions.transformation import ( - HostedVLLMAudioTranscriptionConfig, - ) - - return HostedVLLMAudioTranscriptionConfig() - elif litellm.LlmProviders.WATSONX == provider: - from litellm.llms.watsonx.audio_transcription.transformation import ( - IBMWatsonXAudioTranscriptionConfig, - ) - - return IBMWatsonXAudioTranscriptionConfig() - elif litellm.LlmProviders.OVHCLOUD == provider: - from litellm.llms.ovhcloud.audio_transcription.transformation import ( - OVHCloudAudioTranscriptionConfig, - ) - - return OVHCloudAudioTranscriptionConfig() - elif litellm.LlmProviders.SCALEWAY == provider: - from litellm.llms.scaleway.audio_transcription.transformation import ( - ScalewayAudioTranscriptionConfig, - ) - - return ScalewayAudioTranscriptionConfig() - elif litellm.LlmProviders.MISTRAL == provider: - from litellm.llms.mistral.audio_transcription.transformation import ( - MistralAudioTranscriptionConfig, - ) - - return MistralAudioTranscriptionConfig() - elif litellm.LlmProviders.NVIDIA_RIVA == provider: - from litellm.llms.nvidia_riva.audio_transcription.transformation import ( - NvidiaRivaAudioTranscriptionConfig, - ) - - return NvidiaRivaAudioTranscriptionConfig() - elif litellm.LlmProviders.SONIOX == provider: - from litellm.llms.soniox.audio_transcription.transformation import ( - SonioxAudioTranscriptionConfig, - ) - - return SonioxAudioTranscriptionConfig() - elif litellm.LlmProviders.VERTEX_AI == provider: - bare_vertex_model: Final = model.removeprefix("vertex_ai/") - if bare_vertex_model.startswith("gemini") and "transcribe" in bare_vertex_model: - from litellm.llms.vertex_ai.audio_transcription.gemini_transcribe_transformation import ( - VertexGeminiAudioTranscriptionConfig, - ) - - return VertexGeminiAudioTranscriptionConfig() - from litellm.llms.vertex_ai.audio_transcription.transformation import ( - VertexAIAudioTranscriptionConfig, - ) - - return VertexAIAudioTranscriptionConfig() - elif litellm.LlmProviders.GEMINI == provider: - from litellm.llms.gemini.audio_transcription.transformation import ( - GeminiAudioTranscriptionConfig, - ) - - return GeminiAudioTranscriptionConfig() - return None - - @staticmethod - def get_provider_responses_api_config( - provider: LlmProviders | str, - model: str | None = None, - ) -> BaseResponsesAPIConfig | None: - from litellm.llms.openai_like.dynamic_config import ( - create_responses_config_class, - ) - from litellm.llms.openai_like.json_loader import JSONProviderRegistry - - # Resolve provider string for JSON lookup - provider_str: Final = provider.value if isinstance(provider, LlmProviders) else str(provider) - - # Try to convert to enum for Python class lookup first. - # Python classes take priority over JSON (they have custom overrides). - provider_enum: LlmProviders | None = None - if isinstance(provider, LlmProviders): - provider_enum = provider - else: - try: - provider_enum = LlmProviders(provider) - except ValueError: - pass - - # Check Python classes first (custom overrides take priority) - result: Final = ProviderConfigManager._get_python_responses_api_config(provider_enum, model) - if result is not None: - return result - - # Fall back to JSON providers (generic OpenAI-compatible) - if JSONProviderRegistry.exists(provider_str) and JSONProviderRegistry.supports_responses_api(provider_str): - provider_config: Final = JSONProviderRegistry.get(provider_str) - if provider_config is not None: - return create_responses_config_class(provider_config)() - - return None - - @staticmethod - def _get_python_responses_api_config( - provider: LlmProviders | None, - model: str | None = None, - ) -> BaseResponsesAPIConfig | None: - """Check for Python-class-based responses API configs (custom overrides).""" - if provider is None: - return None - - if litellm.LlmProviders.OPENAI == provider: - return litellm.OpenAIResponsesAPIConfig() - elif litellm.LlmProviders.AZURE == provider: - # Check if it's an O-series model - # Note: GPT models (gpt-3.5, gpt-4, gpt-5, etc.) support temperature parameter - # O-series models (o1, o3) do not contain "gpt" and have different parameter restrictions - is_gpt_model: Final = model and "gpt" in model.lower() - is_o_series = model and ("o_series" in model.lower() or (supports_reasoning(model) and not is_gpt_model)) - - if is_o_series: - return litellm.AzureOpenAIOSeriesResponsesAPIConfig() - else: - return litellm.AzureOpenAIResponsesAPIConfig() - elif litellm.LlmProviders.XAI == provider: - return litellm.XAIResponsesAPIConfig() - elif litellm.LlmProviders.GITHUB_COPILOT == provider: - from litellm.llms.github_copilot.responses.transformation import ( - github_copilot_supports_responses_api, - ) - - if model is None or github_copilot_supports_responses_api(model=model): - return litellm.GithubCopilotResponsesAPIConfig() - return None - elif litellm.LlmProviders.CHATGPT == provider: - return litellm.ChatGPTResponsesAPIConfig() - elif litellm.LlmProviders.LITELLM_PROXY == provider: - return litellm.LiteLLMProxyResponsesAPIConfig() - elif litellm.LlmProviders.VOLCENGINE == provider: - return litellm.VolcEngineResponsesAPIConfig() - elif litellm.LlmProviders.MANUS == provider: - return litellm.ManusResponsesAPIConfig() - elif litellm.LlmProviders.PERPLEXITY == provider: - return litellm.PerplexityResponsesConfig() - elif litellm.LlmProviders.DATABRICKS == provider: - # Databricks Responses API is only compatible with OpenAI GPT models - if model and "gpt" in model.lower(): - return litellm.DatabricksResponsesAPIConfig() - return None - elif litellm.LlmProviders.OPENROUTER == provider: - return litellm.OpenRouterResponsesAPIConfig() - elif litellm.LlmProviders.HOSTED_VLLM == provider: - return litellm.HostedVLLMResponsesAPIConfig() - elif litellm.LlmProviders.BEDROCK_MANTLE == provider: - # Both decisions are data-driven from the model's price-map entry, with - # no model-name logic. Capability (can it serve Responses?) comes from - # mantle_supports_responses (supported_endpoints / mode); - # chat-only models (gpt-oss safeguard, nvidia, ...) return None and keep - # the chat-completions emulation (responses/main.py "config is None"). - # The wire path comes from mantle_base_segment, which reads the - # use_openai_responses_path flag: gpt-5.x and gemma-4-* on - # /openai/v1/responses, everything else (incl. gpt-oss) on - # /v1/responses. - from litellm.llms.bedrock_mantle.common_utils import ( - mantle_base_segment, - mantle_supports_responses, - ) - - if not model or not mantle_supports_responses(model, litellm.model_cost): - return None - return litellm.BedrockMantleResponsesAPIConfig( - use_openai_path=mantle_base_segment(model, litellm.model_cost) == "openai/v1" - ) - return None - - @staticmethod - def get_provider_skills_api_config( - provider: LlmProviders, - ) -> BaseSkillsAPIConfig | None: - """ - Get provider-specific Skills API configuration - - Args: - provider: The LLM provider - - Returns: - Provider-specific Skills API config or None - """ - if litellm.LlmProviders.ANTHROPIC == provider: - return litellm.AnthropicSkillsConfig() - return None - - @staticmethod - def get_provider_evals_api_config( - provider: LlmProviders, - ) -> BaseEvalsAPIConfig | None: - """ - Get provider-specific Evals API configuration - - Args: - provider: The LLM provider - - Returns: - Provider-specific Evals API config or None - """ - if litellm.LlmProviders.OPENAI == provider: - from litellm.llms.openai.evals.transformation import OpenAIEvalsConfig - - return OpenAIEvalsConfig() - return None - - @staticmethod - def get_provider_text_completion_config( - model: str, - provider: LlmProviders, - ) -> BaseTextCompletionConfig: - if LlmProviders.FIREWORKS_AI == provider: - return litellm.FireworksAITextCompletionConfig() - elif LlmProviders.TOGETHER_AI == provider: - return litellm.TogetherAITextCompletionConfig() - elif LlmProviders.TEXT_COMPLETION_INCEPTION == provider: - return litellm.InceptionTextCompletionConfig() - return litellm.OpenAITextCompletionConfig() - - @staticmethod - def get_provider_model_info( - model: str | None, - provider: LlmProviders, - ) -> BaseLLMModelInfo | None: - if LlmProviders.FIREWORKS_AI == provider: - return litellm.FireworksAIConfig() - elif LlmProviders.OPENAI == provider: - return litellm.OpenAIGPTConfig() - elif LlmProviders.GEMINI == provider: - return litellm.GeminiModelInfo() - elif LlmProviders.VERTEX_AI == provider: - from litellm.llms.vertex_ai.common_utils import VertexAIModelInfo - - return VertexAIModelInfo() - elif LlmProviders.LITELLM_PROXY == provider: - return litellm.LiteLLMProxyChatConfig() - elif LlmProviders.TOPAZ == provider: - return litellm.TopazModelInfo() - elif LlmProviders.ANTHROPIC == provider: - return litellm.AnthropicModelInfo() - elif LlmProviders.XAI == provider: - return litellm.XAIModelInfo() - elif LlmProviders.OLLAMA == provider or LlmProviders.OLLAMA_CHAT == provider: - # Dynamic model listing for Ollama server - from litellm.llms.ollama.common_utils import OllamaModelInfo - - return OllamaModelInfo() - elif LlmProviders.VLLM == provider or LlmProviders.HOSTED_VLLM == provider: - from litellm.llms.vllm.common_utils import ( - VLLMModelInfo, # experimental approach, to reduce bloat on __init__.py - ) - - return VLLMModelInfo() - elif LlmProviders.LEMONADE == provider: - return litellm.LemonadeChatConfig() - elif LlmProviders.CLARIFAI == provider: - return litellm.ClarifaiConfig() - elif LlmProviders.BEDROCK == provider: - from litellm.llms.bedrock.common_utils import BedrockModelInfo - - return BedrockModelInfo() - elif LlmProviders.AZURE_AI == provider: - from litellm.llms.azure_ai.common_utils import AzureFoundryModelInfo - - return AzureFoundryModelInfo(model=model) - return None - - @staticmethod - def get_provider_passthrough_config( - model: str, - provider: LlmProviders, - ) -> BasePassthroughConfig | None: - if LlmProviders.BEDROCK == provider: - from litellm.llms.bedrock.passthrough.transformation import ( - BedrockPassthroughConfig, - ) - - return BedrockPassthroughConfig() - elif LlmProviders.BEDROCK_MANTLE == provider: - from litellm.llms.bedrock_mantle.passthrough.transformation import ( - BedrockMantlePassthroughConfig, - ) - - return BedrockMantlePassthroughConfig() - elif LlmProviders.VLLM == provider or LlmProviders.HOSTED_VLLM == provider: - from litellm.llms.vllm.passthrough.transformation import ( - VLLMPassthroughConfig, - ) - - return VLLMPassthroughConfig() - elif LlmProviders.AZURE == provider: - from litellm.llms.azure.passthrough.transformation import ( - AzurePassthroughConfig, - ) - - return AzurePassthroughConfig() - elif LlmProviders.GIGACHAT == provider: - from litellm.llms.gigachat.passthrough.transformation import ( - GigaChatPassthroughConfig, - ) - - return GigaChatPassthroughConfig() - elif LlmProviders.WATSONX == provider: - from litellm.llms.watsonx.passthrough.transformation import ( - WatsonxPassthroughConfig, - ) - - return WatsonxPassthroughConfig() - return None - - @staticmethod - def get_provider_image_variation_config( - model: str, - provider: LlmProviders, - ) -> BaseImageVariationConfig | None: - if LlmProviders.OPENAI == provider: - return litellm.OpenAIImageVariationConfig() - elif LlmProviders.TOPAZ == provider: - return litellm.TopazImageVariationConfig() - return None - - @staticmethod - def get_provider_files_config( - model: str, - provider: LlmProviders, - ) -> BaseFilesConfig | None: - if LlmProviders.GEMINI == provider: - from litellm.llms.gemini.files.transformation import ( - GoogleAIStudioFilesHandler, # experimental approach, to reduce bloat on __init__.py - ) - - return GoogleAIStudioFilesHandler() - elif LlmProviders.VERTEX_AI == provider: - from litellm.llms.vertex_ai.files.transformation import VertexAIFilesConfig - - return VertexAIFilesConfig() - elif LlmProviders.BEDROCK == provider: - from litellm.llms.bedrock.files.transformation import BedrockFilesConfig - - return BedrockFilesConfig() - elif LlmProviders.MANUS == provider: - from litellm.llms.manus.files.transformation import ManusFilesConfig - - return ManusFilesConfig() - elif LlmProviders.ANTHROPIC == provider: - from litellm.llms.anthropic.files.transformation import AnthropicFilesConfig - - return AnthropicFilesConfig() - return None - - @staticmethod - def get_provider_batches_config( - model: str, - provider: LlmProviders, - ) -> BaseBatchesConfig | None: - if LlmProviders.BEDROCK == provider: - from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig - - return BedrockBatchesConfig() - return None - - @staticmethod - def get_provider_vector_store_config( - provider: LlmProviders, - ) -> CustomLogger | None: - from litellm.integrations.vector_store_integrations.bedrock_vector_store import ( - BedrockVectorStore, - ) - - if LlmProviders.BEDROCK == provider: - return BedrockVectorStore.get_initialized_custom_logger() - return None - - @staticmethod - def get_provider_vector_stores_config( - provider: LlmProviders, - api_type: str | None = None, - ) -> BaseVectorStoreConfig | None: - """ - v2 vector store config, use this for new vector store integrations - """ - if litellm.LlmProviders.OPENAI == provider: - from litellm.llms.openai.vector_stores.transformation import ( - OpenAIVectorStoreConfig, - ) - - return OpenAIVectorStoreConfig() - elif litellm.LlmProviders.AZURE == provider: - from litellm.llms.azure.vector_stores.transformation import ( - AzureOpenAIVectorStoreConfig, - ) - - return AzureOpenAIVectorStoreConfig() - elif litellm.LlmProviders.VERTEX_AI == provider: - if api_type == "rag_api" or api_type is None: # default to rag_api - from litellm.llms.vertex_ai.vector_stores.rag_api.transformation import ( - VertexVectorStoreConfig, - ) - - return VertexVectorStoreConfig() - elif api_type == "search_api": - from litellm.llms.vertex_ai.vector_stores.search_api.transformation import ( - VertexSearchAPIVectorStoreConfig, - ) - - return VertexSearchAPIVectorStoreConfig() - elif litellm.LlmProviders.BEDROCK == provider: - from litellm.llms.bedrock.vector_stores.transformation import ( - BedrockVectorStoreConfig, - ) - - return BedrockVectorStoreConfig() - elif litellm.LlmProviders.PG_VECTOR == provider: - from litellm.llms.pg_vector.vector_stores.transformation import ( - PGVectorStoreConfig, - ) - - return PGVectorStoreConfig() - elif litellm.LlmProviders.AZURE_AI == provider: - from litellm.llms.azure_ai.vector_stores.transformation import ( - AzureAIVectorStoreConfig, - ) - - return AzureAIVectorStoreConfig() - elif litellm.LlmProviders.MILVUS == provider: - from litellm.llms.milvus.vector_stores.transformation import ( - MilvusVectorStoreConfig, - ) - - return MilvusVectorStoreConfig() - elif litellm.LlmProviders.GEMINI == provider: - from litellm.llms.gemini.vector_stores.transformation import ( - GeminiVectorStoreConfig, - ) - - return GeminiVectorStoreConfig() - elif litellm.LlmProviders.RAGFLOW == provider: - from litellm.llms.ragflow.vector_stores.transformation import ( - RAGFlowVectorStoreConfig, - ) - - return RAGFlowVectorStoreConfig() - elif litellm.LlmProviders.S3_VECTORS == provider: - from litellm.llms.s3_vectors.vector_stores.transformation import ( - S3VectorsVectorStoreConfig, - ) - - return S3VectorsVectorStoreConfig() - elif litellm.LlmProviders.VALKEY == provider: - from litellm.llms.valkey.vector_stores.transformation import ( - ValkeyVectorStoreConfig, - ) - - return ValkeyVectorStoreConfig() - elif litellm.LlmProviders.MONGODB == provider: - from litellm.llms.mongodb.vector_stores.transformation import ( - MongoDBVectorStoreConfig, - ) - - return MongoDBVectorStoreConfig() - return None - - @staticmethod - def get_provider_vector_store_files_config( - provider: LlmProviders, - ) -> BaseVectorStoreFilesConfig | None: - if litellm.LlmProviders.OPENAI == provider: - from litellm.llms.openai.vector_store_files.transformation import ( - OpenAIVectorStoreFilesConfig, - ) - - return OpenAIVectorStoreFilesConfig() - return None - - @staticmethod - def get_provider_image_generation_config( - model: str, - provider: LlmProviders, - ) -> BaseImageGenerationConfig | None: - if LlmProviders.OPENAI == provider: - from litellm.llms.openai.image_generation import ( - get_openai_image_generation_config, - ) - - return get_openai_image_generation_config(model) - elif LlmProviders.AZURE == provider: - from litellm.llms.azure.image_generation import ( - get_azure_image_generation_config, - ) - - return get_azure_image_generation_config(model) - elif LlmProviders.AZURE_AI == provider: - from litellm.llms.azure_ai.image_generation import ( - get_azure_ai_image_generation_config, - ) - - return get_azure_ai_image_generation_config(model) - elif LlmProviders.XINFERENCE == provider: - from litellm.llms.xinference.image_generation import ( - get_xinference_image_generation_config, - ) - - return get_xinference_image_generation_config(model) - elif LlmProviders.RECRAFT == provider: - from litellm.llms.recraft.image_generation import ( - get_recraft_image_generation_config, - ) - - return get_recraft_image_generation_config(model) - elif LlmProviders.AIML == provider: - from litellm.llms.aiml.image_generation import ( - get_aiml_image_generation_config, - ) - - return get_aiml_image_generation_config(model) - elif LlmProviders.COMETAPI == provider: - from litellm.llms.cometapi.image_generation import ( - get_cometapi_image_generation_config, - ) - - return get_cometapi_image_generation_config(model) - elif LlmProviders.GEMINI == provider: - from litellm.llms.gemini.image_generation import ( - get_gemini_image_generation_config, - ) - - return get_gemini_image_generation_config(model) - elif LlmProviders.LITELLM_PROXY == provider: - from litellm.llms.litellm_proxy.image_generation.transformation import ( - LiteLLMProxyImageGenerationConfig, - ) - - return LiteLLMProxyImageGenerationConfig() - elif LlmProviders.FAL_AI == provider: - from litellm.llms.fal_ai.image_generation import ( - get_fal_ai_image_generation_config, - ) - - return get_fal_ai_image_generation_config(model) - elif LlmProviders.STABILITY == provider: - from litellm.llms.stability.image_generation import ( - get_stability_image_generation_config, - ) - - return get_stability_image_generation_config(model) - elif LlmProviders.RUNWAYML == provider: - from litellm.llms.runwayml.image_generation import ( - get_runwayml_image_generation_config, - ) - - return get_runwayml_image_generation_config(model) - elif LlmProviders.BLACK_FOREST_LABS == provider: - from litellm.llms.black_forest_labs.image_generation import ( - get_black_forest_labs_image_generation_config, - ) - - return get_black_forest_labs_image_generation_config(model) - elif LlmProviders.VERTEX_AI == provider: - from litellm.llms.vertex_ai.image_generation import ( - get_vertex_ai_image_generation_config, - ) - - return get_vertex_ai_image_generation_config(model) - elif LlmProviders.OPENROUTER == provider: - from litellm.llms.openrouter.image_generation import ( - get_openrouter_image_generation_config, - ) - - return get_openrouter_image_generation_config(model) - elif provider in ( - LlmProviders.DASHSCOPE, - LlmProviders.QWENCLOUD, - LlmProviders.QWEN_AI_PLATFORM, - ): - from litellm.llms.dashscope.common_utils import ( - get_dashscope_family_image_generation_config, - ) - - return get_dashscope_family_image_generation_config(provider.value) - elif LlmProviders.MODELSCOPE == provider: - from litellm.llms.modelscope.image_generation import ( - get_modelscope_image_generation_config, - ) - - return get_modelscope_image_generation_config(model) - return None - - @staticmethod - def get_provider_video_config( - model: str | None, - provider: LlmProviders, - ) -> BaseVideoConfig | None: - if LlmProviders.OPENAI == provider: - from litellm.llms.openai.videos.transformation import OpenAIVideoConfig - - return OpenAIVideoConfig() - elif LlmProviders.AZURE == provider: - from litellm.llms.azure.videos.transformation import AzureVideoConfig - - return AzureVideoConfig() - elif LlmProviders.GEMINI == provider: - from litellm.llms.gemini.videos.transformation import GeminiVideoConfig - - return GeminiVideoConfig() - elif LlmProviders.VERTEX_AI == provider: - from litellm.llms.vertex_ai.videos.transformation import VertexAIVideoConfig - - return VertexAIVideoConfig() - elif LlmProviders.RUNWAYML == provider: - from litellm.llms.runwayml.videos.transformation import RunwayMLVideoConfig - - return RunwayMLVideoConfig() - elif LlmProviders.HOSTED_VLLM == provider: - from litellm.llms.hosted_vllm.videos import get_hosted_vllm_video_config - - return get_hosted_vllm_video_config(model) - return None - - @staticmethod - def get_provider_container_config( - provider: LlmProviders, - ) -> BaseContainerConfig | None: - if LlmProviders.OPENAI == provider: - from litellm.llms.openai.containers.transformation import ( - OpenAIContainerConfig, - ) - - return OpenAIContainerConfig() - if provider in (LlmProviders.AZURE, LlmProviders.AZURE_TEXT): - from litellm.llms.azure.containers.transformation import ( - AzureContainerConfig, - ) - - return AzureContainerConfig() - return None - - @staticmethod - def get_provider_realtime_config( - model: str, - provider: LlmProviders, - ) -> BaseRealtimeConfig | None: - if LlmProviders.GEMINI == provider: - from litellm.llms.gemini.realtime.transformation import GeminiRealtimeConfig - - return GeminiRealtimeConfig() - return None - - @staticmethod - def get_provider_realtime_http_config( - model: str, - provider: LlmProviders, - ) -> BaseRealtimeHTTPConfig | None: - """ - Return the HTTP transformation config for realtime HTTP endpoints - (POST /realtime/client_secrets and POST /realtime/calls). - """ - - if LlmProviders.OPENAI == provider: - from litellm.llms.openai.realtime.http_transformation import ( - OpenAIRealtimeHTTPConfig, - ) - - return OpenAIRealtimeHTTPConfig() - if LlmProviders.AZURE == provider: - from litellm.llms.azure.realtime.http_transformation import ( - AzureRealtimeHTTPConfig, - ) - - return AzureRealtimeHTTPConfig() - return None - - @staticmethod - def get_provider_image_edit_config( - model: str, - provider: LlmProviders, - ) -> BaseImageEditConfig | None: - if LlmProviders.OPENAI == provider: - from litellm.llms.openai.image_edit import get_openai_image_edit_config - - return get_openai_image_edit_config(model=model) - elif LlmProviders.AZURE == provider: - from litellm.llms.azure.image_edit.transformation import ( - AzureImageEditConfig, - ) - - return AzureImageEditConfig() - elif LlmProviders.RECRAFT == provider: - from litellm.llms.recraft.image_edit.transformation import ( - RecraftImageEditConfig, - ) - - return RecraftImageEditConfig() - elif LlmProviders.BLACK_FOREST_LABS == provider: - from litellm.llms.black_forest_labs.image_edit.transformation import ( - BlackForestLabsImageEditConfig, - ) - - return BlackForestLabsImageEditConfig() - elif LlmProviders.AZURE_AI == provider: - from litellm.llms.azure_ai.image_edit import get_azure_ai_image_edit_config - - return get_azure_ai_image_edit_config(model) - elif LlmProviders.GEMINI == provider: - from litellm.llms.gemini.image_edit import get_gemini_image_edit_config - - return get_gemini_image_edit_config(model) - elif LlmProviders.LITELLM_PROXY == provider: - from litellm.llms.litellm_proxy.image_edit.transformation import ( - LiteLLMProxyImageEditConfig, - ) - - return LiteLLMProxyImageEditConfig() - elif LlmProviders.VERTEX_AI == provider: - from litellm.llms.vertex_ai.image_edit import ( - get_vertex_ai_image_edit_config, - ) - - return get_vertex_ai_image_edit_config(model) - elif LlmProviders.STABILITY == provider: - from litellm.llms.stability.image_edit import ( - get_stability_image_edit_config, - ) - - return get_stability_image_edit_config(model) - elif LlmProviders.BEDROCK == provider: - from litellm.llms.bedrock.image_edit.amazon_nova_canvas_image_edit_transformation import ( - get_bedrock_image_edit_config_for_model, - ) - - return get_bedrock_image_edit_config_for_model(model) - elif LlmProviders.OPENROUTER == provider: - from litellm.llms.openrouter.image_edit import ( - get_openrouter_image_edit_config, - ) - - return get_openrouter_image_edit_config(model) - return None - - @staticmethod - def get_provider_ocr_config( - model: str, - provider: LlmProviders, - ) -> BaseOCRConfig | None: - """ - Get OCR configuration for a given provider. - """ - from litellm.llms.vertex_ai.ocr.transformation import VertexAIOCRConfig - - # Special handling for Azure AI - distinguish between Mistral OCR and Document Intelligence - if provider == litellm.LlmProviders.AZURE_AI: - from litellm.llms.azure_ai.ocr.common_utils import get_azure_ai_ocr_config - - return get_azure_ai_ocr_config(model=model) - - if provider == litellm.LlmProviders.VERTEX_AI: - from litellm.llms.vertex_ai.ocr.common_utils import get_vertex_ai_ocr_config - - return get_vertex_ai_ocr_config(model=model) - - if provider == litellm.LlmProviders.COHERE: - from litellm.llms.cohere.ocr.transformation import CohereParseConfig - - return CohereParseConfig() - - if provider == litellm.LlmProviders.REDUCTO: - from litellm.llms.reducto.ocr.transformation import ( - ReductoParseLegacyConfig, - ReductoParseV3Config, - ) - - if model == "parse-v3": - return ReductoParseV3Config() - if model == "parse-legacy": - return ReductoParseLegacyConfig() - return None - - MistralOCRConfig: Final = litellm_utils.MistralOCRConfig - PROVIDER_TO_CONFIG_MAP: Final = { - litellm.LlmProviders.MISTRAL: MistralOCRConfig, - } - config_class: Final = PROVIDER_TO_CONFIG_MAP.get(provider, None) - if config_class is None: - return None - return config_class() - - @staticmethod - def get_provider_search_config( - provider: SearchProviders, - ) -> BaseSearchConfig | None: - """ - Get Search configuration for a given provider. - """ - from litellm.llms.apiserpent.search.transformation import ( - APISerpentSearchConfig, - ) - from litellm.llms.azure.search.transformation import BingGroundingSearchConfig - from litellm.llms.bedrock.search.transformation import AgentCoreSearchConfig - from litellm.llms.brave.search.transformation import BraveSearchConfig - from litellm.llms.dataforseo.search.transformation import DataForSEOSearchConfig - from litellm.llms.duckduckgo.search.transformation import DuckDuckGoSearchConfig - from litellm.llms.exa_ai.search.transformation import ExaAISearchConfig - from litellm.llms.fastcrw.search.transformation import FastCRWSearchConfig - from litellm.llms.firecrawl.search.transformation import FirecrawlSearchConfig - from litellm.llms.google_pse.search.transformation import GooglePSESearchConfig - from litellm.llms.linkup.search.transformation import LinkupSearchConfig - from litellm.llms.nimble.search.transformation import NimbleSearchConfig - from litellm.llms.parallel_ai.search.transformation import ( - ParallelAISearchConfig, - ) - from litellm.llms.perplexity.search.transformation import PerplexitySearchConfig - from litellm.llms.searchapi.search.transformation import SearchAPIConfig - from litellm.llms.searxng.search.transformation import SearXNGSearchConfig - from litellm.llms.serper.search.transformation import SerperSearchConfig - from litellm.llms.tavily.search.transformation import TavilySearchConfig - from litellm.llms.tinyfish.search.transformation import TinyfishSearchConfig - from litellm.llms.you_com.search.transformation import YouComSearchConfig - - PROVIDER_TO_CONFIG_MAP: Final = { - SearchProviders.PERPLEXITY: PerplexitySearchConfig, - SearchProviders.TAVILY: TavilySearchConfig, - SearchProviders.PARALLEL_AI: ParallelAISearchConfig, - SearchProviders.EXA_AI: ExaAISearchConfig, - SearchProviders.BRAVE: BraveSearchConfig, - SearchProviders.GOOGLE_PSE: GooglePSESearchConfig, - SearchProviders.DATAFORSEO: DataForSEOSearchConfig, - SearchProviders.FIRECRAWL: FirecrawlSearchConfig, - SearchProviders.FASTCRW: FastCRWSearchConfig, - SearchProviders.SEARXNG: SearXNGSearchConfig, - SearchProviders.LINKUP: LinkupSearchConfig, - SearchProviders.DUCKDUCKGO: DuckDuckGoSearchConfig, - SearchProviders.SEARCHAPI: SearchAPIConfig, - SearchProviders.SERPER: SerperSearchConfig, - SearchProviders.YOU_COM: YouComSearchConfig, - SearchProviders.APISERPENT: APISerpentSearchConfig, - SearchProviders.TINYFISH: TinyfishSearchConfig, - SearchProviders.AGENTCORE: AgentCoreSearchConfig, - SearchProviders.NIMBLE: NimbleSearchConfig, - SearchProviders.BING_GROUNDING: BingGroundingSearchConfig, - } - config_class: Final = PROVIDER_TO_CONFIG_MAP.get(provider, None) - if config_class is None: - return None - return config_class() - - @staticmethod - def get_provider_sandbox_config( - provider: SandboxProviders, - ) -> BaseSandboxConfig | None: - """ - Get sandbox (code execution) configuration for a given provider. - """ - from litellm.llms.e2b.sandbox.transformation import E2BSandboxConfig - from litellm.llms.opensandbox.sandbox.transformation import ( - OpenSandboxSandboxConfig, - ) - - if provider == SandboxProviders.E2B: - return E2BSandboxConfig() - if provider == SandboxProviders.OPENSANDBOX: - return OpenSandboxSandboxConfig() - return None - - @staticmethod - def get_provider_text_to_speech_config( - model: str, - provider: LlmProviders, - ) -> BaseTextToSpeechConfig | None: - """ - Get text-to-speech configuration for a given provider. - """ - from litellm.llms.base_llm.text_to_speech.transformation import ( - BaseTextToSpeechConfig, - ) - - if litellm.LlmProviders.AZURE == provider: - # Only return Azure AVA config for Azure Speech Service models (speech/) - # Azure OpenAI TTS models (azure/azure-tts) should not use this config - if model.startswith("speech/"): - from litellm.llms.azure.text_to_speech.transformation import ( - AzureAVATextToSpeechConfig, - ) - - return AzureAVATextToSpeechConfig() - elif litellm.LlmProviders.ELEVENLABS == provider: - from litellm.llms.elevenlabs.text_to_speech.transformation import ( - ElevenLabsTextToSpeechConfig, - ) - - return ElevenLabsTextToSpeechConfig() - elif litellm.LlmProviders.RUNWAYML == provider: - from litellm.llms.runwayml.text_to_speech.transformation import ( - RunwayMLTextToSpeechConfig, - ) - - return RunwayMLTextToSpeechConfig() - elif litellm.LlmProviders.VERTEX_AI == provider: - if "gemini" in model: - # Gemini TTS uses the speech_to_completion bridge, and Google Cloud TTS param - # mapping would drop response_format before the bridge sees it (LIT-6501) - return None - from litellm.llms.vertex_ai.text_to_speech.transformation import ( - VertexAITextToSpeechConfig, - ) - - return VertexAITextToSpeechConfig() - elif litellm.LlmProviders.MINIMAX == provider: - from litellm.llms.minimax.text_to_speech.transformation import ( - MinimaxTextToSpeechConfig, - ) - - return MinimaxTextToSpeechConfig() - elif litellm.LlmProviders.AWS_POLLY == provider: - from litellm.llms.aws_polly.text_to_speech.transformation import ( - AWSPollyTextToSpeechConfig, - ) - - return AWSPollyTextToSpeechConfig() - return None - - @staticmethod - def get_provider_google_genai_generate_content_config( - model: str, - provider: LlmProviders, - ) -> BaseGoogleGenAIGenerateContentConfig | None: - if litellm.LlmProviders.GEMINI == provider: - from litellm.llms.gemini.google_genai.transformation import ( - GoogleGenAIConfig, - ) - - return GoogleGenAIConfig() - elif litellm.LlmProviders.VERTEX_AI == provider: - from litellm.llms.vertex_ai.google_genai.transformation import ( - VertexAIGoogleGenAIConfig, - ) - from litellm.llms.vertex_ai.vertex_ai_partner_models.main import ( - VertexAIPartnerModels, - ) - - ######################################################### - # If Vertex Partner models like Anthropic, Mistral, etc. are used, - # return None as we want this to go through the litellm.completion() adapter - # and not the Google Gen AI adapter - ######################################################### - if VertexAIPartnerModels.is_vertex_partner_model(model): - return None - - ######################################################### - # If the model is not a Vertex Partner model, return the Vertex AI Google Gen AI Config - # This is for Vertex `gemini` models - ######################################################### - return VertexAIGoogleGenAIConfig() - return None - - -def get_end_user_id_for_cost_tracking( - litellm_params: dict, - service_type: Literal["litellm_logging", "prometheus"] = "litellm_logging", -) -> str | None: - """ - Used for enforcing `disable_end_user_cost_tracking` param. - - service_type: "litellm_logging" or "prometheus" - used to allow prometheus only disable cost tracking. - """ - get_litellm_metadata_from_kwargs: Final = getattr(sys.modules[__name__], "get_litellm_metadata_from_kwargs") - _metadata: Final = cast(dict, get_litellm_metadata_from_kwargs(dict(litellm_params=litellm_params))) - - end_user_id: Final = cast( - str | None, - litellm_params.get("user_api_key_end_user_id") or _metadata.get("user_api_key_end_user_id"), - ) - if litellm.disable_end_user_cost_tracking: - return None - - ####################################### - # By default we don't track end_user on prometheus since we don't want to increase cardinality - # by default litellm.enable_end_user_cost_tracking_prometheus_only is None, so we don't track end_user on prometheus - ####################################### - if service_type == "prometheus": - if litellm.enable_end_user_cost_tracking_prometheus_only is not True: - return None - return end_user_id - - -def should_use_cohere_v1_client(api_base: str | None, present_version_params: list[str]): - if not api_base: - return False - uses_v1_params: Final = ("max_chunks_per_doc" in present_version_params) and ( - "max_tokens_per_doc" not in present_version_params - ) - return api_base.endswith("/v1/rerank") or (uses_v1_params and not api_base.endswith("/v2/rerank")) - - -def get_prompt_cache_min_tokens(model: str) -> int: - """ - Returns the smallest prefix `model` will actually cache. - - Resolution order is an explicitly configured `MINIMUM_PROMPT_CACHE_TOKEN_COUNT`, then the - model's `prompt_cache_min_tokens` in the cost map, then the provider-agnostic default. The - cost map is the source of truth because the real minimum is per-model and per-platform: - Anthropic's ranges from 512 to 4096 and moves in both directions across releases, and the - same model can differ by platform. - - Never raises. An unresolvable model falls back to the default rather than propagating, so a - caller cannot mistake "no entry for this model" for "this prompt is not cacheable". - """ - if MINIMUM_PROMPT_CACHE_TOKEN_COUNT_OVERRIDE is not None: - return MINIMUM_PROMPT_CACHE_TOKEN_COUNT_OVERRIDE - try: - min_tokens: Final = get_model_info(model=model).get("prompt_cache_min_tokens") - except Exception: - return DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT - if min_tokens is None: - return DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT - return min_tokens - - -def is_prompt_caching_valid_prompt( - model: str, - messages: list[AllMessageValues] | None, - tools: list[ChatCompletionToolParam] | None = None, - custom_llm_provider: str | None = None, - min_token_count: int | None = None, -) -> bool: - """ - Returns true if the prompt is valid for prompt caching. - - The minimum cacheable prefix is per-model, so it is resolved from `model` unless the caller - passes `min_token_count`. Callers that only hold a model-group alias (the router's deployment - checks) must resolve the threshold themselves and pass it, because an alias resolves to - nothing here and would silently fall back to the default. - - OpenAI's minimum is a flat 1024 across models, which the default already covers. - """ - try: - if messages is None and tools is None: - return False - if custom_llm_provider is not None and not model.startswith(custom_llm_provider): - model = custom_llm_provider + "/" + model - token_count: Final = token_counter( - messages=messages, - tools=tools, - model=model, - use_default_image_token_count=True, - ) - if min_token_count is None: - min_token_count = get_prompt_cache_min_tokens(model=model) - return token_count >= min_token_count - except Exception as e: - verbose_logger.error("Error in is_prompt_caching_valid_prompt: %s", e) - return False - - -def extract_duration_from_srt_or_vtt(srt_or_vtt_content: str) -> float | None: - """ - Extracts the total duration (in seconds) from SRT or VTT content. - - Args: - srt_or_vtt_content (str): The content of an SRT or VTT file as a string. - - Returns: - Optional[float]: The total duration in seconds, or None if no timestamps are found. - """ - # Regular expression to match timestamps in the format "hh:mm:ss,ms" or "hh:mm:ss.ms" - timestamp_pattern: Final = r"(\d{2}):(\d{2}):(\d{2})[.,](\d{3})" - - timestamps: Final[Sequence[tuple[str, str, str, str]]] = re.findall(timestamp_pattern, srt_or_vtt_content) - - if not timestamps: - return None - - # Convert timestamps to seconds and find the max (end time) - durations: Final = [] - match: tuple[str, str, str, str] - for match in timestamps: - hours, minutes, seconds, milliseconds = map(int, match) - total_seconds = hours * 3600 + minutes * 60 + seconds + milliseconds / 1000.0 - durations.append(total_seconds) - - return max(durations) if durations else None - - -def _add_path_to_api_base(api_base: str, ending_path: str) -> str: - """ - Adds an ending path to an API base URL while preventing duplicate path segments. - - Args: - api_base: Base URL string - ending_path: Path to append to the base URL - - Returns: - Modified URL string with proper path handling - """ - original_url: Final = httpx.URL(api_base) - base_url: Final = original_url.copy_with(params={}) # Removes query params - base_path: Final = original_url.path.rstrip("/") - end_path: Final = ending_path.lstrip("/") - - # Split paths into segments - base_segments: Final = [s for s in base_path.split("/") if s] - end_segments: Final = [s for s in end_path.split("/") if s] - - # Find overlapping segments from the end of base_path and start of ending_path - final_segments = [] - for i in range(len(base_segments)): - if base_segments[i:] == end_segments[: len(base_segments) - i]: - final_segments = base_segments[:i] + end_segments - break - else: - # No overlap found, just combine all segments - final_segments = base_segments + end_segments - - # Construct the new path - modified_path: Final = "/" + "/".join(final_segments) - modified_url: Final = base_url.copy_with(path=modified_path) - - # Re-add the original query parameters - return str(modified_url.copy_with(params=original_url.params)) - - -def get_standard_openai_params(params: Mapping[str, object]) -> dict: - return {k: v for k, v in params.items() if k in litellm.OPENAI_CHAT_COMPLETION_PARAMS and v is not None} - - -def get_non_default_completion_params(kwargs: Mapping[str, object]) -> dict: - openai_params: Final = litellm.OPENAI_CHAT_COMPLETION_PARAMS - default_params: Final = openai_params + all_litellm_params - non_default_params: Final = { - k: v for k, v in kwargs.items() if k not in default_params - } # model-specific params - pass them straight to the model/provider - - return non_default_params - - -def peek_reasoning_summary_aliases(optional_params: dict) -> object | None: - """Read AI-SDK-style reasoning summary from optional_params or nested extra_body. - - Uses key membership (not ``or`` chains) so falsy values like ``""`` are not skipped. - """ - if "reasoningSummary" in optional_params: - return optional_params["reasoningSummary"] - if "reasoning_summary" in optional_params: - return optional_params["reasoning_summary"] - extra_body: Final = optional_params.get("extra_body") - if isinstance(extra_body, dict): - if "reasoningSummary" in extra_body: - return extra_body["reasoningSummary"] - if "reasoning_summary" in extra_body: - return extra_body["reasoning_summary"] - return None - - -def strip_reasoning_summary_aliases_from_optional_params( - optional_params: dict, -) -> tuple[dict, object | None]: - """Copy optional_params; remove reasoningSummary aliases from top-level and extra_body.""" - op: Final = dict(optional_params) - rs_val = op.pop("reasoningSummary", None) - snake_rs_val: Final = op.pop("reasoning_summary", None) - if rs_val is None: - rs_val = snake_rs_val - eb = op.get("extra_body") - if isinstance(eb, dict): - eb = dict(eb) - eb_rs_val: Final = eb.pop("reasoningSummary", None) - eb_snake_rs_val: Final = eb.pop("reasoning_summary", None) - if rs_val is None: - rs_val = eb_rs_val - if rs_val is None: - rs_val = eb_snake_rs_val - if eb: - op["extra_body"] = eb - else: - op.pop("extra_body", None) - return op, rs_val - - -def get_non_default_transcription_params(kwargs: dict) -> dict: - from litellm.constants import OPENAI_TRANSCRIPTION_PARAMS - - default_params: Final = OPENAI_TRANSCRIPTION_PARAMS + all_litellm_params - non_default_params: Final = {k: v for k, v in kwargs.items() if k not in default_params} - return non_default_params - - -def add_openai_metadata( - metadata: Mapping[str, object] | None, -) -> dict[str, str] | None: - """ - Add metadata to openai optional parameters, excluding hidden params. - - OpenAI 'metadata' only supports string values. - - Args: - params (dict): Dictionary of API parameters - metadata (dict, optional): Metadata to include in the request - - Returns: - dict: Updated parameters dictionary with visible metadata only - """ - if metadata is None: - return None - # Only include non-hidden parameters - visible_metadata: dict[str, str] = { - str(k): v for k, v in metadata.items() if k != "hidden_params" and isinstance(v, str) - } - - # max 16 keys allowed by openai - trim down to 16 - if len(visible_metadata) > 16: - filtered_metadata: Final = {} - idx = 0 - for k, v in visible_metadata.items(): - if idx < 16: - filtered_metadata[k] = v - idx += 1 - visible_metadata = filtered_metadata - - return visible_metadata.copy() - - -def get_requester_metadata(metadata: Mapping[str, object]): - if not metadata: - return None - - requester_metadata: Final = metadata.get("requester_metadata") - if isinstance(requester_metadata, dict): - cleaned_metadata = add_openai_metadata(requester_metadata) - if cleaned_metadata: - return cleaned_metadata - - cleaned_metadata = add_openai_metadata(metadata) - if cleaned_metadata: - return cleaned_metadata - - return None - - -def return_raw_request(endpoint: CallTypes, kwargs: dict) -> RawRequestTypedDict: - """ - Return the json str of the request - - This is currently in BETA, and tested for `/chat/completions` -> `litellm.completion` calls. - """ - from datetime import datetime - - from litellm.litellm_core_utils.litellm_logging import Logging - - litellm_logging_obj: Final = Logging( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hi"}], - stream=False, - call_type="acompletion", - litellm_call_id="1234", - start_time=datetime.now(), - function_id="1234", - log_raw_request_response=True, - ) - - llm_api_endpoint: Final = getattr(litellm, endpoint.value) - - received_exception = "" - - try: - llm_api_endpoint( - **kwargs, - litellm_logging_obj=litellm_logging_obj, - api_key="my-fake-api-key", # 👈 ensure the request fails - ) - except Exception as e: - received_exception = str(e) - - raw_request_typed_dict: Final = litellm_logging_obj.model_call_details.get("raw_request_typed_dict") - if raw_request_typed_dict: - return cast(RawRequestTypedDict, raw_request_typed_dict) - else: - return RawRequestTypedDict( - error=received_exception, - ) - - -def jsonify_tools(tools: Sequence[object]) -> list[dict]: - """ - Fixes https://github.com/BerriAI/litellm/issues/9321 - - Where user passes in a pydantic base model - """ - new_tools: Final[list[dict]] = [] - for tool in tools: - if isinstance(tool, BaseModel): - tool = tool.model_dump(exclude_none=True) - elif isinstance(tool, dict): - tool = tool.copy() - if isinstance(tool, dict): - new_tools.append(tool) - return new_tools - - -def get_empty_usage() -> Usage: - return Usage( - prompt_tokens=0, - completion_tokens=0, - total_tokens=0, - ) - - -def should_run_mock_completion( - mock_response: object | None, - mock_tool_calls: object | None, - mock_timeout: object | None, -) -> bool: - if mock_response or mock_tool_calls or mock_timeout: - return True - return False - - -def __getattr__(name: str) -> Any: - """Lazy import handler for utils module with cached registry for improved performance.""" - # Use cached registry from _lazy_imports instead of importing tuples every time - from litellm._lazy_imports import _get_lazy_import_registry - - registry: Final = _get_lazy_import_registry() - - # Check if name is in registry and call the cached handler function - if name in registry: - handler_func: Final = registry[name] - return handler_func(name) - - raise AttributeError(f"module {__name__!r} has no attribute {name!r}") +@file:/tmp/litellm-work/litellm/litellm/utils.py \ No newline at end of file diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py index 0a96d2d466b..12a40131089 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py @@ -1,5014 +1 @@ -import json -import time -from unittest.mock import AsyncMock, MagicMock, Mock, patch - -import pytest - -import asyncio -import traceback -from typing import Optional - -import litellm -from litellm import verbose_logger -from litellm._logging import session_id_var, trace_id_var -from litellm.litellm_core_utils.litellm_logging import Logging -from litellm.litellm_core_utils.streaming_handler import ( - AUDIO_ATTRIBUTE, - CustomStreamWrapper, - _ProviderChunkEarlyReturn, - _ProviderChunkParsed, -) -from litellm.types.utils import ( - CompletionTokensDetailsWrapper, - Delta, - ModelResponse, - ModelResponseStream, - PromptTokensDetailsWrapper, - StandardLoggingPayload, - StreamingChoices, - Usage, -) -from litellm.utils import ModelResponseListIterator - - -@pytest.fixture -def initialized_custom_stream_wrapper() -> CustomStreamWrapper: - streaming_handler = CustomStreamWrapper( - completion_stream=None, - model=None, - logging_obj=MagicMock(), - custom_llm_provider=None, - ) - return streaming_handler - - -@pytest.fixture -def logging_obj() -> Logging: - import time - - logging_obj = Logging( - model="my-random-model", - messages=[{"role": "user", "content": "Hey"}], - stream=True, - call_type="completion", - start_time=time.time(), - litellm_call_id="12345", - function_id="1245", - ) - return logging_obj - - -bedrock_chunks = [ - ModelResponseStream( - id="chatcmpl-d249def8-a78b-464c-87b5-3a6f43565292", - created=1742056047, - model=None, - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta( - provider_specific_fields=None, - content="I'm Claude", - role="assistant", - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields={}, - usage=None, - ), - ModelResponseStream( - id="chatcmpl-fe559823-b383-4249-ab87-52f6ad9d08c2", - created=1742056047, - model=None, - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta( - provider_specific_fields=None, - content=", an AI", - role="assistant", - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields={}, - usage=None, - ), - ModelResponseStream( - id="chatcmpl-c1c6cc2f-75b9-4a24-88b9-4e5aacd0268b", - created=1742056047, - model=None, - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason="stop", - index=0, - delta=Delta( - provider_specific_fields=None, - content="", - role="assistant", - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields={}, - usage=None, - ), -] - - -def test_is_chunk_non_empty(initialized_custom_stream_wrapper: CustomStreamWrapper): - """Unit test if non-empty when reasoning_content is present""" - chunk = { - "id": "e89b6501-8ac2-464c-9550-7cd3daf94350", - "object": "chat.completion.chunk", - "created": 1741037890, - "model": "deepseek-reasoner", - "system_fingerprint": "fp_5417b77867_prod0225", - "choices": [ - { - "index": 0, - "delta": {"content": None, "reasoning_content": "."}, - "logprobs": None, - "finish_reason": None, - } - ], - } - assert initialized_custom_stream_wrapper.is_chunk_non_empty( - completion_obj=MagicMock(), - model_response=ModelResponseStream(**chunk), - response_obj=MagicMock(), - ) - - -def test_is_chunk_non_empty_with_annotations( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """Unit test if non-empty when annotations are present""" - chunk = { - "id": "e89b6501-8ac2-464c-9550-7cd3daf94350", - "object": "chat.completion.chunk", - "created": 1741037890, - "model": "deepseek-reasoner", - "system_fingerprint": "fp_5417b77867_prod0225", - "choices": [ - { - "index": 0, - "delta": { - "content": None, - "annotations": [ - {"type": "url_citation", "url": "https://www.google.com"} - ], - }, - "logprobs": None, - "finish_reason": None, - } - ], - } - assert ( - initialized_custom_stream_wrapper.is_chunk_non_empty( - completion_obj=MagicMock(), - model_response=ModelResponseStream(**chunk), - response_obj=MagicMock(), - ) - is True - ) - - -def test_optional_combine_thinking_block_in_choices( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """Test that reasoning_content is properly combined with content using tags""" - # Setup the wrapper to use the merge feature - initialized_custom_stream_wrapper.merge_reasoning_content_in_choices = True - - # First chunk with reasoning_content - should add tag - first_chunk = { - "id": "chunk1", - "object": "chat.completion.chunk", - "created": 1741037890, - "model": "deepseek-reasoner", - "choices": [ - { - "index": 0, - "delta": { - "content": "", - "reasoning_content": "Let me think about this", - }, - "finish_reason": None, - } - ], - } - - # Middle chunk with more reasoning_content - middle_chunk = { - "id": "chunk2", - "object": "chat.completion.chunk", - "created": 1741037891, - "model": "deepseek-reasoner", - "choices": [ - { - "index": 0, - "delta": {"content": "", "reasoning_content": " step by step"}, - "finish_reason": None, - } - ], - } - - # Final chunk with actual content - should add tag - final_chunk = { - "id": "chunk3", - "object": "chat.completion.chunk", - "created": 1741037892, - "model": "deepseek-reasoner", - "choices": [ - { - "index": 0, - "delta": {"content": "The answer is 42", "reasoning_content": None}, - "finish_reason": None, - } - ], - } - - # Process first chunk - first_response = ModelResponseStream(**first_chunk) - initialized_custom_stream_wrapper._optional_combine_thinking_block_in_choices( - first_response - ) - print("first_response", json.dumps(first_response, indent=4, default=str)) - assert first_response.choices[0].delta.content == "Let me think about this" - # assert the response does not have attribute reasoning_content - assert not hasattr(first_response.choices[0].delta, "reasoning_content") - - assert initialized_custom_stream_wrapper.sent_first_thinking_block is True - - # Process middle chunk - middle_response = ModelResponseStream(**middle_chunk) - initialized_custom_stream_wrapper._optional_combine_thinking_block_in_choices( - middle_response - ) - print("middle_response", json.dumps(middle_response, indent=4, default=str)) - assert middle_response.choices[0].delta.content == " step by step" - assert not hasattr(middle_response.choices[0].delta, "reasoning_content") - - # Process final chunk - final_response = ModelResponseStream(**final_chunk) - initialized_custom_stream_wrapper._optional_combine_thinking_block_in_choices( - final_response - ) - print("final_response", json.dumps(final_response, indent=4, default=str)) - assert final_response.choices[0].delta.content == "The answer is 42" - assert initialized_custom_stream_wrapper.sent_last_thinking_block is True - assert not hasattr(final_response.choices[0].delta, "reasoning_content") - - -def test_multi_chunk_reasoning_and_content( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """Test handling of multiple reasoning chunks followed by multiple content chunks""" - # Setup the wrapper to use the merge feature - initialized_custom_stream_wrapper.merge_reasoning_content_in_choices = True - initialized_custom_stream_wrapper.sent_first_thinking_block = False - initialized_custom_stream_wrapper.sent_last_thinking_block = False - - # Create test chunks - chunks = [ - # Chunk 1: First reasoning chunk - { - "id": "chunk1", - "object": "chat.completion.chunk", - "created": 1741037890, - "model": "deepseek-reasoner", - "choices": [ - { - "index": 0, - "delta": { - "content": "", - "reasoning_content": "To solve this problem", - }, - "finish_reason": None, - } - ], - }, - # Chunk 2: Second reasoning chunk - { - "id": "chunk2", - "object": "chat.completion.chunk", - "created": 1741037891, - "model": "deepseek-reasoner", - "choices": [ - { - "index": 0, - "delta": { - "content": "", - "reasoning_content": ", I need to calculate 6 * 7", - }, - "finish_reason": None, - } - ], - }, - # Chunk 3: Third reasoning chunk - { - "id": "chunk3", - "object": "chat.completion.chunk", - "created": 1741037892, - "model": "deepseek-reasoner", - "choices": [ - { - "index": 0, - "delta": {"content": "", "reasoning_content": " which equals 42"}, - "finish_reason": None, - } - ], - }, - # Chunk 4: First content chunk (transition from reasoning to content) - { - "id": "chunk4", - "object": "chat.completion.chunk", - "created": 1741037893, - "model": "deepseek-reasoner", - "choices": [ - { - "index": 0, - "delta": { - "content": "The answer to your question", - "reasoning_content": None, - }, - "finish_reason": None, - } - ], - }, - # Chunk 5: Second content chunk - { - "id": "chunk5", - "object": "chat.completion.chunk", - "created": 1741037894, - "model": "deepseek-reasoner", - "choices": [ - { - "index": 0, - "delta": {"content": " is 42.", "reasoning_content": None}, - "finish_reason": None, - } - ], - }, - ] - - # Expected content after processing each chunk - expected_contents = [ - "To solve this problem", - ", I need to calculate 6 * 7", - " which equals 42", - "The answer to your question", - " is 42.", - ] - - # Process each chunk and verify results - for i, (chunk, expected_content) in enumerate(zip(chunks, expected_contents)): - response = ModelResponseStream(**chunk) - initialized_custom_stream_wrapper._optional_combine_thinking_block_in_choices( - response - ) - - # Check content - assert ( - response.choices[0].delta.content == expected_content - ), f"Chunk {i+1}: content mismatch" - - # Check reasoning_content was removed - assert not hasattr( - response.choices[0].delta, "reasoning_content" - ), f"Chunk {i+1}: reasoning_content should be removed" - - # Verify final state - assert initialized_custom_stream_wrapper.sent_first_thinking_block is True - assert initialized_custom_stream_wrapper.sent_last_thinking_block is True - - -def test_strip_sse_data_from_chunk(): - """Test the static method that strips 'data: ' prefix from SSE chunks""" - # Test with string inputs - assert CustomStreamWrapper._strip_sse_data_from_chunk("data: content") == "content" - assert ( - CustomStreamWrapper._strip_sse_data_from_chunk("data: spaced content") - == " spaced content" - ) - assert ( - CustomStreamWrapper._strip_sse_data_from_chunk("regular content") - == "regular content" - ) - assert ( - CustomStreamWrapper._strip_sse_data_from_chunk("regular content with data:") - == "regular content with data:" - ) - - # Test with None input - assert CustomStreamWrapper._strip_sse_data_from_chunk(None) is None - - -@pytest.mark.parametrize("sync_mode", [True, False]) -@pytest.mark.asyncio -async def test_streaming_handler_with_usage( - sync_mode: bool, final_usage_block: Optional[Usage] = None -): - import time - - final_usage_block = final_usage_block or Usage( - completion_tokens=392, - prompt_tokens=1799, - total_tokens=2191, - completion_tokens_details=CompletionTokensDetailsWrapper( # <-- This has a value - accepted_prediction_tokens=None, - audio_tokens=None, - reasoning_tokens=0, - rejected_prediction_tokens=None, - text_tokens=None, - ), - prompt_tokens_details=PromptTokensDetailsWrapper( - audio_tokens=None, cached_tokens=1796, text_tokens=None, image_tokens=None - ), - ) - - final_chunk = ModelResponseStream( - id="chatcmpl-87291500-d8c5-428e-b187-36fe5a4c97ab", - created=1742056047, - model=None, - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta( - provider_specific_fields=None, - content="", - role="assistant", - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields={}, - usage=final_usage_block, - ) - test_chunks = bedrock_chunks + [final_chunk] - completion_stream = ModelResponseListIterator(model_responses=test_chunks) - - response = CustomStreamWrapper( - completion_stream=completion_stream, - model="bedrock/claude-haiku-4-5-20251001-v1:0", - custom_llm_provider="bedrock", - logging_obj=Logging( - model="bedrock/claude-haiku-4-5-20251001-v1:0", - messages=[{"role": "user", "content": "Hey"}], - stream=True, - call_type="completion", - start_time=time.time(), - litellm_call_id="12345", - function_id="1245", - ), - stream_options={"include_usage": True}, - ) - - chunk_has_usage = False - if sync_mode: - for chunk in response: - if hasattr(chunk, "usage"): - assert chunk.usage == final_usage_block - chunk_has_usage = True - else: - async for chunk in response: - if hasattr(chunk, "usage"): - assert chunk.usage == final_usage_block - chunk_has_usage = True - assert chunk_has_usage - - -@pytest.mark.parametrize("sync_mode", [False]) -@pytest.mark.asyncio -@pytest.mark.flaky(reruns=3) -async def test_streaming_with_usage_and_logging(sync_mode: bool): - import time - - from litellm.integrations.custom_logger import CustomLogger - - class MockCallback(CustomLogger): - pass - - mock_callback = MockCallback() - litellm.success_callback = [mock_callback] - litellm._async_success_callback = [mock_callback] - - final_usage_block = Usage( - completion_tokens=392, - prompt_tokens=1799, - total_tokens=2191, - completion_tokens_details=CompletionTokensDetailsWrapper( - accepted_prediction_tokens=None, - audio_tokens=None, - reasoning_tokens=0, - rejected_prediction_tokens=None, - text_tokens=None, - ), - prompt_tokens_details=PromptTokensDetailsWrapper( - audio_tokens=None, - cached_tokens=1796, - text_tokens=None, - image_tokens=None, - ), - cache_creation_input_tokens=0, - cache_read_input_tokens=1796, - ) - - with ( - patch.object(mock_callback, "log_success_event") as mock_log_success_event, - patch.object(mock_callback, "log_stream_event") as mock_log_stream_event, - patch.object( - mock_callback, "async_log_success_event" - ) as mock_async_log_success_event, - patch.object( - mock_callback, "async_log_stream_event" - ) as mock_async_log_stream_event, - ): - await test_streaming_handler_with_usage( - sync_mode=sync_mode, final_usage_block=final_usage_block - ) - if sync_mode: - time.sleep(1) - mock_log_success_event.assert_called_once() - # mock_log_stream_event.assert_called() - assert ( - mock_log_success_event.call_args.kwargs["response_obj"].usage - == final_usage_block - ) - else: - await asyncio.sleep(1) - mock_async_log_success_event.assert_called_once() - # mock_async_log_stream_event.assert_called() - assert ( - mock_async_log_success_event.call_args.kwargs["response_obj"].usage - == final_usage_block - ) - - -def test_streaming_handler_with_stop_chunk( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - args = { - "completion_obj": {"content": ""}, - "response_obj": { - "text": "", - "is_finished": True, - "finish_reason": "length", - "logprobs": None, - "original_chunk": ModelResponseStream( - id="chatcmpl-ad517c2e-c197-48de-a2e6-a559cca48124", - created=1742093326, - model=None, - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason="length", - index=0, - delta=Delta( - provider_specific_fields=None, - content="", - role="assistant", - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields={}, - usage=None, - ), - "usage": None, - }, - } - - returned_chunk = initialized_custom_stream_wrapper.return_processed_chunk_logic( - **args, model_response=ModelResponseStream() - ) - assert returned_chunk is None - - -def test_finish_reason_chunk_preserves_non_openai_attributes( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """ - Regression test for #23444: - Preserve upstream non-OpenAI attributes on final finish_reason chunk. - """ - initialized_custom_stream_wrapper.received_finish_reason = "stop" - - original_chunk = ModelResponseStream( - id="chatcmpl-test", - created=1742093326, - model=None, - object="chat.completion.chunk", - choices=[ - StreamingChoices( - finish_reason="stop", - index=0, - delta=Delta(content=""), - logprobs=None, - ) - ], - ) - setattr(original_chunk, "custom_field", {"key": "value"}) - - returned_chunk = initialized_custom_stream_wrapper.return_processed_chunk_logic( - completion_obj={"content": ""}, - response_obj={"original_chunk": original_chunk}, - model_response=ModelResponseStream(), - ) - - assert returned_chunk is not None - assert getattr(returned_chunk, "custom_field", None) == {"key": "value"} - - -def test_finish_reason_with_holding_chunk_preserves_non_openai_attributes( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """ - Regression test for #23444 holding-chunk path: - preserve custom attributes when _is_delta_empty is False after flushing - holding_chunk. - """ - initialized_custom_stream_wrapper.received_finish_reason = "stop" - initialized_custom_stream_wrapper.holding_chunk = "filtered text" - - original_chunk = ModelResponseStream( - id="chatcmpl-test-2", - created=1742093327, - model=None, - object="chat.completion.chunk", - choices=[ - StreamingChoices( - finish_reason="stop", - index=0, - delta=Delta(content=""), - logprobs=None, - ) - ], - ) - setattr(original_chunk, "custom_field", {"key": "value"}) - - returned_chunk = initialized_custom_stream_wrapper.return_processed_chunk_logic( - completion_obj={"content": ""}, - response_obj={"original_chunk": original_chunk}, - model_response=ModelResponseStream(), - ) - - assert returned_chunk is not None - assert returned_chunk.choices[0].delta.content == "filtered text" - assert getattr(returned_chunk, "custom_field", None) == {"key": "value"} - - -def test_set_response_id_propagation_empty_to_valid( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """Test that response_id is properly set when first chunk has empty ID and second chunk has valid ID""" - - model_response1 = ModelResponseStream(id="", created=1742056047, model=None) - model_response1 = initialized_custom_stream_wrapper.set_model_id( - model_response1.id, model_response1 - ) - assert model_response1.id == "" - - model_response2 = ModelResponseStream( - id="valid-id-123", created=1742056048, model=None - ) - model_response2 = initialized_custom_stream_wrapper.set_model_id( - "valid-id-123", model_response2 - ) - assert model_response2.id == "valid-id-123" - assert initialized_custom_stream_wrapper.response_id == "valid-id-123" - - -def test_set_response_id_propagation_valid_to_invalid( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """Test that response_id is maintained when first chunk has valid ID and second chunk has invalid ID""" - - model_response1 = ModelResponseStream( - id="first-valid-id", created=1742056049, model=None - ) - model_response1 = initialized_custom_stream_wrapper.set_model_id( - "first-valid-id", model_response1 - ) - assert model_response1.id == "first-valid-id" - assert initialized_custom_stream_wrapper.response_id == "first-valid-id" - - model_response2 = ModelResponseStream(id="", created=1742056050, model=None) - model_response2 = initialized_custom_stream_wrapper.set_model_id( - "", model_response2 - ) - assert model_response2.id == "first-valid-id" - assert initialized_custom_stream_wrapper.response_id == "first-valid-id" - - -@pytest.mark.asyncio -async def test_streaming_completion_start_time(logging_obj: Logging): - """Test that the start time is set correctly""" - from litellm.integrations.custom_logger import CustomLogger - - class MockCallback(CustomLogger): - pass - - mock_callback = MockCallback() - litellm.success_callback = [mock_callback, "langfuse"] - - completion_stream = ModelResponseListIterator( - model_responses=bedrock_chunks, delay=0.1 - ) - - response = CustomStreamWrapper( - completion_stream=completion_stream, - model="bedrock/claude-haiku-4-5-20251001-v1:0", - logging_obj=logging_obj, - ) - - async for chunk in response: - print(chunk) - - await asyncio.sleep(2) - - assert logging_obj.model_call_details["completion_start_time"] is not None - assert ( - logging_obj.model_call_details["completion_start_time"] - < logging_obj.model_call_details["end_time"] - ) - - -@pytest.mark.asyncio -async def test_vertex_streaming_bad_request_not_midstream(logging_obj: Logging): - """Ensure Vertex bad request errors surface as 400, not mid-stream fallbacks.""" - from litellm.llms.vertex_ai.common_utils import VertexAIError - - async def _raise_bad_request(**kwargs): - raise VertexAIError( - status_code=400, message="invalid maxOutputTokens", headers=None - ) - - response = CustomStreamWrapper( - completion_stream=None, - model="gemini-3-pro-preview", - logging_obj=logging_obj, - custom_llm_provider="vertex_ai_beta", - make_call=_raise_bad_request, - ) - - with pytest.raises(litellm.BadRequestError) as excinfo: - await response.__anext__() - - assert getattr(excinfo.value, "status_code", None) == 400 - assert "invalid maxOutputTokens" in str(excinfo.value) - - -@pytest.mark.asyncio -async def test_vertex_streaming_rate_limit_triggers_midstream_fallback( - logging_obj: Logging, -): - """Ensure Vertex 429 rate-limit errors raise MidStreamFallbackError, not RateLimitError. - - Regression test for https://github.com/BerriAI/litellm/issues/20870 - """ - from litellm.exceptions import MidStreamFallbackError - from litellm.llms.vertex_ai.common_utils import VertexAIError - - async def _raise_rate_limit(**kwargs): - raise VertexAIError( - status_code=429, message="Resource exhausted.", headers=None - ) - - response = CustomStreamWrapper( - completion_stream=None, - model="gemini-3-flash-preview", - logging_obj=logging_obj, - custom_llm_provider="vertex_ai_beta", - make_call=_raise_rate_limit, - ) - - with pytest.raises(MidStreamFallbackError) as excinfo: - await response.__anext__() - - assert excinfo.value.is_pre_first_chunk is True - assert excinfo.value.generated_content == "" - - -def test_sync_streaming_rate_limit_triggers_midstream_fallback(logging_obj: Logging): - """Ensure __next__ raises MidStreamFallbackError on 429, not RateLimitError. - - This is the sync-streaming equivalent of the async test above. Before - this fix, __next__ would raise RateLimitError directly, bypassing the - Router's fallback chain entirely. - """ - from litellm.exceptions import MidStreamFallbackError - from litellm.llms.vertex_ai.common_utils import VertexAIError - - def _raise_rate_limit(**kwargs): - raise VertexAIError( - status_code=429, message="Resource exhausted.", headers=None - ) - - response = CustomStreamWrapper( - completion_stream=None, - model="gemini-3-flash-preview", - logging_obj=logging_obj, - custom_llm_provider="vertex_ai_beta", - make_call=_raise_rate_limit, - ) - - with pytest.raises(MidStreamFallbackError) as excinfo: - next(response) - - assert excinfo.value.is_pre_first_chunk is True - assert excinfo.value.generated_content == "" - - -def test_sync_streaming_bad_request_not_midstream(logging_obj: Logging): - """Ensure __next__ raises BadRequestError (400) directly, not MidStreamFallbackError. - - Non-retriable 4xx errors should surface immediately to the caller. - """ - from litellm.llms.vertex_ai.common_utils import VertexAIError - - def _raise_bad_request(**kwargs): - raise VertexAIError( - status_code=400, message="invalid maxOutputTokens", headers=None - ) - - response = CustomStreamWrapper( - completion_stream=None, - model="gemini-3-pro-preview", - logging_obj=logging_obj, - custom_llm_provider="vertex_ai_beta", - make_call=_raise_bad_request, - ) - - with pytest.raises(litellm.BadRequestError) as excinfo: - next(response) - - assert getattr(excinfo.value, "status_code", None) == 400 - assert "invalid maxOutputTokens" in str(excinfo.value) - - -def _bedrock_error_event(exception_type: str): - """A mocked botocore event-stream error event: status_code is botocore's - hard-coded 400, with the real type in the :exception-type header.""" - event = Mock() - event.to_response_dict = Mock( - return_value={ - "status_code": 400, - "headers": { - ":exception-type": exception_type, - ":content-type": "application/json", - ":message-type": "exception", - }, - "body": b'{"message":"Bedrock had an internal error."}', - } - ) - return event - - -@pytest.mark.asyncio -async def test_bedrock_midstream_internal_server_error_wraps_for_fallback( - logging_obj: Logging, -): - """End-to-end regression for https://github.com/BerriAI/litellm/issues/24608: - a Bedrock mid-stream internalServerException event (botocore stamps it 400) - must flow through the real decoder, gain its modeled 500 status, and wrap - into MidStreamFallbackError so the Router can run streaming fallback. - - Calls the real AWSEventStreamDecoder, so reverting the decoder status fix - makes the decoder raise BedrockError(400) and the gate raises BadRequestError - directly -> this test fails without the fix.""" - from litellm.exceptions import MidStreamFallbackError - from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder - - decoder = AWSEventStreamDecoder(model="anthropic.claude-3-sonnet-20240229-v1:0") - - async def _bedrock_stream(): - decoder._parse_message_from_event( - _bedrock_error_event("internalServerException") - ) - yield # unreachable; the line above raises - - async def _make_call(**kwargs): - return _bedrock_stream() - - response = CustomStreamWrapper( - completion_stream=None, - model="anthropic.claude-3-sonnet-20240229-v1:0", - logging_obj=logging_obj, - custom_llm_provider="bedrock", - make_call=_make_call, - ) - - with pytest.raises(MidStreamFallbackError): - await response.__anext__() - - -@pytest.mark.asyncio -async def test_bedrock_5xx_wraps_for_midstream_fallback(logging_obj: Logging): - """Gate contract: a Bedrock 5xx (here 503 serviceUnavailableException) wraps - into MidStreamFallbackError so the Router can run streaming fallback.""" - from litellm.exceptions import MidStreamFallbackError - from litellm.llms.bedrock.chat.invoke_handler import BedrockError - - async def _raise_503(**kwargs): - raise BedrockError( - status_code=503, - message="serviceUnavailableException Bedrock is unavailable.", - ) - - response = CustomStreamWrapper( - completion_stream=None, - model="anthropic.claude-3-sonnet-20240229-v1:0", - logging_obj=logging_obj, - custom_llm_provider="bedrock", - make_call=_raise_503, - ) - - with pytest.raises(MidStreamFallbackError): - await response.__anext__() - - -@pytest.mark.asyncio -async def test_bedrock_validation_error_raises_directly(logging_obj: Logging): - """Gate contract: a Bedrock validationException (400) is a client error and - must surface directly, never wrapped into MidStreamFallbackError.""" - from litellm.exceptions import MidStreamFallbackError - from litellm.llms.bedrock.chat.invoke_handler import BedrockError - - async def _raise_400(**kwargs): - raise BedrockError( - status_code=400, - message="validationException malformed input.", - ) - - response = CustomStreamWrapper( - completion_stream=None, - model="anthropic.claude-3-sonnet-20240229-v1:0", - logging_obj=logging_obj, - custom_llm_provider="bedrock", - make_call=_raise_400, - ) - - with pytest.raises(Exception, match='litellm\\.BadRequestError: BedrockException') as excinfo: - await response.__anext__() - assert not isinstance(excinfo.value, MidStreamFallbackError) - assert getattr(excinfo.value, "status_code", None) == 400 - - -def _hosted_vllm_stream_wrapper(logging_obj: Logging, error_payload: dict) -> CustomStreamWrapper: - """A CustomStreamWrapper over the real OpenAI-compatible line iterator, - fed an HTTP 200 SSE body that carries an in-body error payload the way - vLLM/sglang emit it.""" - from litellm.llms.openai.chat.gpt_transformation import ( - OpenAIChatCompletionStreamingHandler, - ) - - async def _stream(): - yield f"data: {json.dumps(error_payload)}" - yield "data: [DONE]" - - completion_stream = OpenAIChatCompletionStreamingHandler( - streaming_response=_stream(), sync_stream=False - ) - return CustomStreamWrapper( - completion_stream=completion_stream, - model="qwen-vl", - logging_obj=logging_obj, - custom_llm_provider="hosted_vllm", - ) - - -@pytest.mark.asyncio -async def test_in_body_stream_error_400_raises_bad_request(logging_obj: Logging): - """Regression for https://github.com/BerriAI/litellm/issues/25492: a 400 - error returned inside a 200 SSE body must surface as BadRequestError with - the provider's message, not be parsed as an empty chunk that silently - ends the stream (and never as an internal MidStreamFallbackError).""" - from litellm.exceptions import MidStreamFallbackError - - response = _hosted_vllm_stream_wrapper( - logging_obj, - { - "error": { - "object": "error", - "message": "The model is not multimodal. Please remove image inputs.", - "type": "BadRequestError", - "param": None, - "code": 400, - } - }, - ) - - with pytest.raises(litellm.BadRequestError) as excinfo: - await response.__anext__() - - assert not isinstance(excinfo.value, MidStreamFallbackError) - assert excinfo.value.status_code == 400 - assert "not multimodal" in str(excinfo.value) - - -@pytest.mark.asyncio -async def test_in_body_stream_error_500_wraps_for_midstream_fallback( - logging_obj: Logging, -): - """An in-body 5xx error wraps into MidStreamFallbackError so the Router's - FallbackStreamWrapper can switch to a configured fallback deployment.""" - from litellm.exceptions import MidStreamFallbackError - - response = _hosted_vllm_stream_wrapper( - logging_obj, - { - "error": { - "object": "error", - "message": "internal engine crash", - "type": "InternalServerError", - "param": None, - "code": 500, - } - }, - ) - - with pytest.raises(MidStreamFallbackError) as excinfo: - await response.__anext__() - - assert excinfo.value.is_pre_first_chunk is True - assert "internal engine crash" in str(excinfo.value) - - -@pytest.mark.asyncio -async def test_async_streaming_read_timeout_triggers_midstream_fallback( - logging_obj: Logging, -): - """A mid-stream httpx.ReadTimeout must wrap into MidStreamFallbackError so - the Router's FallbackStreamWrapper can switch to a fallback model. - - Previously __anext__ caught httpx.TimeoutException and re-raised it raw, - which bypassed _handle_stream_fallback_error and prevented stream_timeout - from triggering fallbacks the way connection-phase timeout does. - """ - import httpx - - from litellm.exceptions import MidStreamFallbackError - - async def _raise_read_timeout(**kwargs): - raise httpx.ReadTimeout("Timeout on reading data from socket") - - response = CustomStreamWrapper( - completion_stream=None, - model="gpt-4", - logging_obj=logging_obj, - custom_llm_provider="openai", - make_call=_raise_read_timeout, - ) - - with pytest.raises(MidStreamFallbackError) as excinfo: - await response.__anext__() - - assert excinfo.value.is_pre_first_chunk is True - assert isinstance(excinfo.value.original_exception, Exception) - - -def test_streaming_handler_with_created_time_propagation( - initialized_custom_stream_wrapper: CustomStreamWrapper, logging_obj: Logging -): - """Test that the created time is consistent across chunks""" - import time - - bad_chunk = ModelResponseStream( - choices=[], created=int(time.time()) - ) # chunk with different created time - - completion_stream = ModelResponseListIterator( - model_responses=bedrock_chunks + [bad_chunk] - ) - - response = CustomStreamWrapper( - completion_stream=completion_stream, - model="bedrock/claude-haiku-4-5-20251001-v1:0", - logging_obj=logging_obj, - ) - - created: Optional[int] = None - for chunk in response: - if created is None: - created = chunk.created - else: - assert created == chunk.created - - -def test_streaming_handler_with_stream_options( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """Test that the stream options are propagated to the response""" - - mr = initialized_custom_stream_wrapper.model_response_creator() - mr_dict = mr.model_dump() - print(mr_dict) - assert "stream_options" not in mr_dict - - -def test_optional_combine_thinking_block_with_none_content( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """Test that reasoning_content is properly combined when delta.content is None""" - # Setup the wrapper to use the merge feature - initialized_custom_stream_wrapper.merge_reasoning_content_in_choices = True - - # First chunk with reasoning_content and None content - should handle None gracefully - first_chunk = { - "id": "chunk1", - "object": "chat.completion.chunk", - "created": 1741037890, - "model": "deepseek-reasoner", - "choices": [ - { - "index": 0, - "delta": { - "content": None, # This is None, not empty string - "reasoning_content": "Let me think about this problem", - }, - "finish_reason": None, - } - ], - } - - # Second chunk with reasoning_content and None content - second_chunk = { - "id": "chunk2", - "object": "chat.completion.chunk", - "created": 1741037891, - "model": "deepseek-reasoner", - "choices": [ - { - "index": 0, - "delta": { - "content": None, # This is None, not empty string - "reasoning_content": " step by step", - }, - "finish_reason": None, - } - ], - } - - # Final chunk with actual content - should add tag - final_chunk = { - "id": "chunk3", - "object": "chat.completion.chunk", - "created": 1741037892, - "model": "deepseek-reasoner", - "choices": [ - { - "index": 0, - "delta": {"content": "The answer is 42", "reasoning_content": None}, - "finish_reason": None, - } - ], - } - - # Process first chunk - should not raise TypeError - first_response = ModelResponseStream(**first_chunk) - initialized_custom_stream_wrapper._optional_combine_thinking_block_in_choices( - first_response - ) - assert ( - first_response.choices[0].delta.content - == "Let me think about this problem" - ) - assert not hasattr(first_response.choices[0].delta, "reasoning_content") - assert initialized_custom_stream_wrapper.sent_first_thinking_block is True - - # Process second chunk - should work with continued reasoning - second_response = ModelResponseStream(**second_chunk) - initialized_custom_stream_wrapper._optional_combine_thinking_block_in_choices( - second_response - ) - assert second_response.choices[0].delta.content == " step by step" - assert not hasattr(second_response.choices[0].delta, "reasoning_content") - - # Process final chunk - should add tag - final_response = ModelResponseStream(**final_chunk) - initialized_custom_stream_wrapper._optional_combine_thinking_block_in_choices( - final_response - ) - assert final_response.choices[0].delta.content == "The answer is 42" - assert initialized_custom_stream_wrapper.sent_last_thinking_block is True - assert not hasattr(final_response.choices[0].delta, "reasoning_content") - - -def test_has_special_delta_content( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """Test the _has_special_delta_content helper method""" - - # Test empty choices - empty_response = ModelResponseStream( - id="test", created=1742056047, model=None, choices=[] - ) - assert not initialized_custom_stream_wrapper._has_special_delta_content( - empty_response - ) - - # Test with tool_calls (simulate with mock object) - tool_call_response = ModelResponseStream( - id="test", - created=1742056047, - model=None, - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta( - content=None, - tool_calls=[ - { - "id": "test", - "function": {"arguments": "{}", "name": "test_func"}, - } - ], - ), - ) - ], - ) - assert initialized_custom_stream_wrapper._has_special_delta_content( - tool_call_response - ) - - # Test with function_call (simulate with mock object) - function_call_response = ModelResponseStream( - id="test", - created=1742056047, - model=None, - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta( - content=None, function_call={"name": "test_func", "arguments": "{}"} - ), - ) - ], - ) - assert initialized_custom_stream_wrapper._has_special_delta_content( - function_call_response - ) - - # Test with audio (simulate by adding audio attribute) - audio_response = ModelResponseStream( - id="test", - created=1742056047, - model=None, - choices=[ - StreamingChoices(finish_reason=None, index=0, delta=Delta(content=None)) - ], - ) - # Manually add audio attribute to delta - audio_response.choices[0].delta.audio = {"transcript": "test"} - assert initialized_custom_stream_wrapper._has_special_delta_content(audio_response) - - # Test with image (simulate by adding image attribute) - image_response = ModelResponseStream( - id="test", - created=1742056047, - model=None, - choices=[ - StreamingChoices(finish_reason=None, index=0, delta=Delta(content=None)) - ], - ) - # Manually add image attribute to delta - image_response.choices[0].delta.images = [{"url": "test.jpg"}] - assert initialized_custom_stream_wrapper._has_special_delta_content(image_response) - - # Test with regular content (should return False) - regular_response = ModelResponseStream( - id="test", - created=1742056047, - model=None, - choices=[ - StreamingChoices( - finish_reason=None, index=0, delta=Delta(content="Hello world") - ) - ], - ) - assert not initialized_custom_stream_wrapper._has_special_delta_content( - regular_response - ) - - -def test_handle_special_delta_content( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """Test the _handle_special_delta_content helper method""" - test_response = ModelResponseStream( - id="test", - created=1742056047, - model=None, - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta(content="test", role="assistant"), - ) - ], - ) - - # The method should call strip_role_from_delta - result = initialized_custom_stream_wrapper._handle_special_delta_content( - test_response - ) - - # Should return the same response object (modified) - assert result is test_response - - # Should have set sent_first_chunk to True - assert initialized_custom_stream_wrapper.sent_first_chunk is True - - -def test_has_any_special_delta_attributes( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """Test the _has_any_special_delta_attributes helper method""" - - # Test with delta that has audio attribute - class MockDelta: - def __init__(self): - self.audio = {"transcript": "Hello world"} - - audio_delta = MockDelta() - result = initialized_custom_stream_wrapper._has_any_special_delta_attributes( - audio_delta - ) - assert result is True - - # Test with delta that has image attribute - class MockDeltaImage: - def __init__(self): - self.images = [{"url": "test.jpg"}] - - image_delta = MockDeltaImage() - result = initialized_custom_stream_wrapper._has_any_special_delta_attributes( - image_delta - ) - assert result is True - - # Test with delta that has no special attributes - class MockDeltaRegular: - def __init__(self): - self.content = "regular content" - - regular_delta = MockDeltaRegular() - result = initialized_custom_stream_wrapper._has_any_special_delta_attributes( - regular_delta - ) - assert result is False - - -def test_calculate_total_usage_with_cost(): - from litellm.litellm_core_utils.streaming_handler import calculate_total_usage - - chunk1_usage = Usage(completion_tokens=1, prompt_tokens=10, total_tokens=11) - chunk1 = ModelResponseStream( - id="test-1", - created=1745513206, - model="openrouter/test", - choices=[ - StreamingChoices(finish_reason=None, index=0, delta=Delta(content="Hi")) - ], - usage=chunk1_usage, - ) - - chunk2_usage = Usage( - completion_tokens=5, prompt_tokens=10, total_tokens=15, cost=0.00025 - ) - chunk2 = ModelResponseStream( - id="test-1", - created=1745513207, - model="openrouter/test", - choices=[ - StreamingChoices(finish_reason="stop", index=0, delta=Delta(content="")) - ], - usage=chunk2_usage, - ) - - usage = calculate_total_usage([chunk1, chunk2]) - - assert hasattr(usage, "cost") - assert usage.cost == 0.00025 - assert usage.prompt_tokens == 10 - assert usage.completion_tokens == 5 - - -def test_calculate_total_usage_with_dict_usage_cost(): - """Regression: dict-shaped `usage` with a `cost` key must still surface - provider cost even though `hasattr` on a dict does not consult its keys.""" - from litellm.litellm_core_utils.streaming_handler import calculate_total_usage - - chunk = { - "usage": { - "prompt_tokens": 10, - "completion_tokens": 5, - "total_tokens": 15, - "cost": 0.00025, - } - } - - usage = calculate_total_usage([chunk]) - - assert usage.prompt_tokens == 10 - assert usage.completion_tokens == 5 - assert getattr(usage, "cost", None) == 0.00025 - - -def test_calculate_total_usage_preserves_prompt_cache_token_details(): - """Regression for #34801: dropping `prompt_tokens_details` here re-prices OpenAI - cache-read tokens at the uncached input rate, overstating spend.""" - from litellm.litellm_core_utils.streaming_handler import calculate_total_usage - - usage_with_details = Usage( - prompt_tokens=6017, - completion_tokens=4, - total_tokens=6021, - prompt_tokens_details=PromptTokensDetailsWrapper( - cached_tokens=6004, cache_write_tokens=10 - ), - completion_tokens_details=CompletionTokensDetailsWrapper(reasoning_tokens=2), - ) - chunk_with_details = ModelResponseStream( - id="chatcmpl-1", - created=1745513206, - model="openai/gpt-5.6-sol", - choices=[ - StreamingChoices(finish_reason=None, index=0, delta=Delta(content="Hi")) - ], - usage=usage_with_details, - ) - chunk_without_details = ModelResponseStream( - id="chatcmpl-1", - created=1745513207, - model="openai/gpt-5.6-sol", - choices=[ - StreamingChoices(finish_reason="stop", index=0, delta=Delta(content="")) - ], - usage=Usage(prompt_tokens=6017, completion_tokens=4, total_tokens=6021), - ) - - usage = calculate_total_usage([chunk_with_details, chunk_without_details]) - - assert usage.prompt_tokens == 6017 - assert usage.prompt_tokens_details is not None - assert usage.prompt_tokens_details.cached_tokens == 6004 - assert usage.prompt_tokens_details.cache_write_tokens == 10 - assert usage.completion_tokens_details is not None - assert usage.completion_tokens_details.reasoning_tokens == 2 - - -def test_calculate_total_usage_preserves_anthropic_cache_creation_ttl_breakdown(): - """Anthropic sends the 5m/1h cache-write split only on `message_start`; the later - `message_delta` repeats the flat count without the split. Losing it here bills 1h - cache writes at the cheaper 5m rate.""" - from litellm.litellm_core_utils.streaming_handler import calculate_total_usage - from litellm.types.utils import CacheCreationTokenDetails - - message_start_chunk = ModelResponseStream( - id="chatcmpl-1", - created=1745513206, - model="claude-sonnet-5", - choices=[ - StreamingChoices(finish_reason=None, index=0, delta=Delta(content="Hi")) - ], - usage=Usage( - prompt_tokens=120, - completion_tokens=1, - total_tokens=121, - prompt_tokens_details=PromptTokensDetailsWrapper( - cached_tokens=0, - cache_creation_tokens=100, - cache_creation_token_details=CacheCreationTokenDetails( - ephemeral_5m_input_tokens=20, ephemeral_1h_input_tokens=80 - ), - ), - ), - ) - message_delta_chunk = ModelResponseStream( - id="chatcmpl-1", - created=1745513207, - model="claude-sonnet-5", - choices=[ - StreamingChoices(finish_reason="stop", index=0, delta=Delta(content="")) - ], - usage=Usage( - prompt_tokens=120, - completion_tokens=4, - total_tokens=124, - prompt_tokens_details=PromptTokensDetailsWrapper( - cached_tokens=0, cache_creation_tokens=100 - ), - ), - ) - - usage = calculate_total_usage([message_start_chunk, message_delta_chunk]) - - assert usage.prompt_tokens_details is not None - assert usage.prompt_tokens_details.cache_creation_tokens == 100 - ttl_breakdown = usage.prompt_tokens_details.cache_creation_token_details - assert ttl_breakdown is not None - assert ttl_breakdown.ephemeral_5m_input_tokens == 20 - assert ttl_breakdown.ephemeral_1h_input_tokens == 80 - - -@pytest.mark.asyncio -async def test_openrouter_streaming_cost_after_finish_reason(logging_obj: Logging): - from litellm.utils import ModelResponseListIterator - - chunk1 = ModelResponseStream( - id="chatcmpl-or", - created=1742056047, - model="openrouter/claude", - choices=[ - StreamingChoices( - finish_reason=None, index=0, delta=Delta(content="Hi", role="assistant") - ) - ], - usage=None, - ) - chunk2 = ModelResponseStream( - id="chatcmpl-or", - created=1742056048, - model="openrouter/claude", - choices=[ - StreamingChoices(finish_reason="stop", index=0, delta=Delta(content="")) - ], - usage=None, - ) - chunk3_usage = Usage( - completion_tokens=5, prompt_tokens=10, total_tokens=15, cost=0.00025 - ) - chunk3 = ModelResponseStream( - id="chatcmpl-or", - created=1742056049, - model="openrouter/claude", - choices=[ - StreamingChoices(finish_reason=None, index=0, delta=Delta(content="")) - ], - usage=chunk3_usage, - ) - - completion_stream = ModelResponseListIterator( - model_responses=[chunk1, chunk2, chunk3] - ) - response = CustomStreamWrapper( - completion_stream=completion_stream, - model="openrouter/claude", - custom_llm_provider="openrouter", - logging_obj=logging_obj, - stream_options={"include_usage": True}, - ) - - collected_chunks = [] - async for chunk in response: - collected_chunks.append(chunk) - - usage_chunks = [c for c in collected_chunks if hasattr(c, "usage") and c.usage] - assert len(usage_chunks) > 0 - assert hasattr(usage_chunks[-1].usage, "cost") - assert usage_chunks[-1].usage.cost == 0.00025 - - -@pytest.mark.asyncio -async def test_openrouter_streaming_usage_only_chunk_without_stream_options(): - """ - Regression: OpenRouter's post-finish chunk has `choices: []`. When the caller did not - pass stream_options.include_usage it was dropped before cost tracking, so the - provider-reported cost never reached the assembled response. - """ - import time - - from litellm.integrations.custom_logger import CustomLogger - from litellm.utils import ModelResponseListIterator - - chunk1 = ModelResponseStream( - id="chatcmpl-or", - created=1742056047, - model="openrouter/claude", - choices=[ - StreamingChoices( - finish_reason=None, index=0, delta=Delta(content="Hi", role="assistant") - ) - ], - usage=None, - ) - chunk2 = ModelResponseStream( - id="chatcmpl-or", - created=1742056048, - model="openrouter/claude", - choices=[ - StreamingChoices(finish_reason="stop", index=0, delta=Delta(content="")) - ], - usage=None, - ) - usage_only_chunk = ModelResponseStream( - id="chatcmpl-or", - created=1742056049, - model="openrouter/claude", - choices=[], - usage=Usage( - completion_tokens=5, prompt_tokens=10, total_tokens=15, cost=0.00025 - ), - ) - - class MockCallback(CustomLogger): - pass - - mock_callback = MockCallback() - previous_success_callback = litellm.success_callback - previous_async_success_callback = litellm._async_success_callback - litellm.success_callback = [mock_callback] - litellm._async_success_callback = [mock_callback] - - stream_logging_obj = Logging( - model="openrouter/claude", - messages=[{"role": "user", "content": "Hey"}], - stream=True, - call_type="completion", - start_time=time.time(), - litellm_call_id="12345", - function_id="1245", - ) - stream_logging_obj.update_environment_variables( - model="openrouter/claude", - optional_params={}, - litellm_params={}, - custom_llm_provider="openrouter", - ) - - response = CustomStreamWrapper( - completion_stream=ModelResponseListIterator( - model_responses=[chunk1, chunk2, usage_only_chunk] - ), - model="openrouter/claude", - custom_llm_provider="openrouter", - logging_obj=stream_logging_obj, - ) - - success_logged = asyncio.Event() - try: - with patch.object( - mock_callback, - "async_log_success_event", - new_callable=AsyncMock, - side_effect=lambda *args, **kwargs: success_logged.set(), - ) as mock_success_event: - collected_chunks = [chunk async for chunk in response] - await asyncio.wait_for(success_logged.wait(), timeout=30) - finally: - litellm.success_callback = previous_success_callback - litellm._async_success_callback = previous_async_success_callback - - assert all(getattr(chunk, "usage", None) is None for chunk in collected_chunks) - - mock_success_event.assert_called_once() - logged_kwargs = mock_success_event.call_args.kwargs["kwargs"] - assert logged_kwargs["response_cost"] == 0.00025 - assert logged_kwargs["standard_logging_object"]["response_cost"] == 0.00025 - - -def test_openrouter_streaming_cost_propagates_to_hidden_params(): - """ - Verify that provider-reported cost from usage.cost flows into - _hidden_params["additional_headers"]["llm_provider-x-litellm-response-cost"] - on the complete streaming response, so litellm's cost calculator uses it. - """ - import litellm - - chunk1 = ModelResponseStream( - id="chatcmpl-or", - created=1742056047, - model="openrouter/claude", - choices=[ - StreamingChoices( - finish_reason=None, index=0, delta=Delta(content="Hi", role="assistant") - ) - ], - usage=None, - ) - chunk2 = ModelResponseStream( - id="chatcmpl-or", - created=1742056048, - model="openrouter/claude", - choices=[ - StreamingChoices(finish_reason="stop", index=0, delta=Delta(content="")) - ], - usage=None, - ) - chunk3 = ModelResponseStream( - id="chatcmpl-or", - created=1742056049, - model="openrouter/claude", - choices=[ - StreamingChoices(finish_reason=None, index=0, delta=Delta(content="")) - ], - usage=Usage( - completion_tokens=5, prompt_tokens=10, total_tokens=15, cost=0.00025 - ), - ) - - # Build the complete response as stream_chunk_builder does - complete_response = litellm.stream_chunk_builder( - chunks=[chunk1, chunk2, chunk3], - messages=[{"role": "user", "content": "test"}], - ) - - assert complete_response is not None - assert hasattr(complete_response.usage, "cost") - assert complete_response.usage.cost == 0.00025 - - # Use the real propagation method from CustomStreamWrapper - CustomStreamWrapper._propagate_usage_cost_to_hidden_params(complete_response, "openrouter") - - assert "additional_headers" in complete_response._hidden_params - assert ( - complete_response._hidden_params["additional_headers"][ - "llm_provider-x-litellm-response-cost" - ] - == 0.00025 - ) - - # Verify the cost calculator would pick this up - from litellm.cost_calculator import get_response_cost_from_hidden_params - - provider_cost = get_response_cost_from_hidden_params( - complete_response._hidden_params - ) - assert provider_cost == 0.00025 - - -def test_perplexity_streaming_dict_cost_bills_through_its_own_calculator(): - import litellm - from litellm.cost_calculator import ( - get_response_cost_from_hidden_params, - response_cost_calculator, - ) - - chunks = [ - ModelResponseStream( - id="chatcmpl-pplx", - created=1742056047, - model="perplexity/sonar", - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta(content="Hi", role="assistant"), - ) - ], - usage=None, - ), - ModelResponseStream( - id="chatcmpl-pplx", - created=1742056048, - model="perplexity/sonar", - choices=[ - StreamingChoices(finish_reason="stop", index=0, delta=Delta(content="")) - ], - usage=None, - ), - ModelResponseStream( - id="chatcmpl-pplx", - created=1742056049, - model="perplexity/sonar", - choices=[ - StreamingChoices(finish_reason=None, index=0, delta=Delta(content="")) - ], - usage=Usage( - completion_tokens=18, - prompt_tokens=12, - total_tokens=30, - cost={ - "input_tokens_cost": 0.000012, - "output_tokens_cost": 0.000018, - "request_cost": 0.005, - "total_cost": 0.00503, - }, - ), - ), - ] - - complete_response = litellm.stream_chunk_builder( - chunks=chunks, messages=[{"role": "user", "content": "test"}] - ) - - assert complete_response is not None - - CustomStreamWrapper._propagate_usage_cost_to_hidden_params(complete_response, "perplexity") - - assert get_response_cost_from_hidden_params(complete_response._hidden_params) is None - assert response_cost_calculator( - response_object=complete_response, - model="perplexity/sonar", - custom_llm_provider="perplexity", - call_type="completion", - optional_params={}, - ) == pytest.approx(0.00503) - - -def test_openai_compatible_streaming_cost_is_priced_from_the_cost_map(): - import litellm - from litellm.cost_calculator import ( - get_response_cost_from_hidden_params, - response_cost_calculator, - ) - - model = "openai/streams-cost-in-nanodollars" - litellm.register_model( - { - model: { - "input_cost_per_token": 1e-6, - "output_cost_per_token": 2e-6, - "litellm_provider": "openai", - "mode": "chat", - } - } - ) - complete_response = ModelResponse( - id="chatcmpl-openai-compatible", - model=model, - choices=[], - usage=Usage(completion_tokens=5, prompt_tokens=10, total_tokens=15, cost=3_144_000), - ) - - CustomStreamWrapper._propagate_usage_cost_to_hidden_params(complete_response, "openai") - - assert get_response_cost_from_hidden_params(complete_response._hidden_params) is None - assert response_cost_calculator( - response_object=complete_response, - model=model, - custom_llm_provider="openai", - call_type="completion", - optional_params={}, - ) == pytest.approx(2e-5) - - -def test_xai_streaming_reported_cost_still_takes_the_margin(monkeypatch): - import litellm - from litellm.cost_calculator import ( - get_response_cost_from_hidden_params, - response_cost_calculator, - ) - - complete_response = ModelResponse( - id="chatcmpl-xai", - model="grok-4-latest", - choices=[], - usage=Usage(completion_tokens=353, prompt_tokens=198, total_tokens=551, cost=0.0009956), - ) - - CustomStreamWrapper._propagate_usage_cost_to_hidden_params(complete_response, "xai") - - assert get_response_cost_from_hidden_params(complete_response._hidden_params) is None - monkeypatch.setattr(litellm, "cost_margin_config", {"xai": 0.5}) - assert response_cost_calculator( - response_object=complete_response, - model="xai/grok-4-latest", - custom_llm_provider="xai", - call_type="completion", - optional_params={}, - ) == pytest.approx(0.0009956 * 1.5) - - -def test_provider_reported_cost_ignores_unusable_shapes(): - assert CustomStreamWrapper._resolve_provider_reported_cost(None) is None - assert CustomStreamWrapper._resolve_provider_reported_cost({}) is None - assert CustomStreamWrapper._resolve_provider_reported_cost({"total_cost": None}) is None - assert CustomStreamWrapper._resolve_provider_reported_cost(0.5) == 0.5 - - -def test_handle_special_delta_attributes( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """Test the _handle_special_delta_attributes helper method""" - - # Create a model response - model_response = ModelResponseStream( - id="test", - created=1742056047, - model=None, - choices=[ - StreamingChoices(finish_reason=None, index=0, delta=Delta(content="test")) - ], - ) - - # Test with delta that has audio attribute - class MockDelta: - def __init__(self): - self.audio = {"transcript": "Hello world"} - - audio_delta = MockDelta() - initialized_custom_stream_wrapper._handle_special_delta_attributes( - audio_delta, model_response - ) - - # Should copy the audio attribute - assert hasattr(model_response.choices[0].delta, "audio") - assert model_response.choices[0].delta.audio == {"transcript": "Hello world"} - - # Test with delta that has image attribute - class MockDeltaImage: - def __init__(self): - self.images = [{"url": "test.jpg"}] - - image_delta = MockDeltaImage() - model_response2 = ModelResponseStream( - id="test", - created=1742056047, - model=None, - choices=[ - StreamingChoices(finish_reason=None, index=0, delta=Delta(content="test")) - ], - ) - - initialized_custom_stream_wrapper._handle_special_delta_attributes( - image_delta, model_response2 - ) - - # Should copy the image attribute - print(f"delta: {model_response2.choices[0].delta}") - assert hasattr(model_response2.choices[0].delta, "images") - print(f"images: {model_response2.choices[0].delta.images}") - assert model_response2.choices[0].delta.images[0] == {"url": "test.jpg"} - - -def test_has_special_delta_attribute( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """Test the _has_special_delta_attribute helper method""" - - # Test with None delta - assert not initialized_custom_stream_wrapper._has_special_delta_attribute( - None, "audio" - ) - - # Test with delta that has the attribute - class MockDelta: - def __init__(self): - self.audio = {"transcript": "test"} - - delta_with_audio = MockDelta() - assert initialized_custom_stream_wrapper._has_special_delta_attribute( - delta_with_audio, "audio" - ) - - # Test with delta that doesn't have the attribute - class MockDeltaNoAudio: - def __init__(self): - self.content = "test" - - delta_without_audio = MockDeltaNoAudio() - assert not initialized_custom_stream_wrapper._has_special_delta_attribute( - delta_without_audio, "audio" - ) - - # Test with delta that has the attribute but it's None - class MockDeltaNone: - def __init__(self): - self.audio = None - - delta_with_none = MockDeltaNone() - assert not initialized_custom_stream_wrapper._has_special_delta_attribute( - delta_with_none, "audio" - ) - - -def test_is_chunk_non_empty_with_empty_tool_calls( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """ - Test that is_chunk_non_empty returns False when tool_calls is an empty list. - - Regression test for https://github.com/BerriAI/litellm/issues/17425 - Empty tool_calls in delta should not be considered non-empty chunks. - """ - chunk = { - "id": "test-chunk-id", - "object": "chat.completion.chunk", - "created": 1741037890, - "model": "claude-sonnet-4-20250514", - "choices": [ - { - "index": 0, - "delta": { - "content": None, - "tool_calls": [], # Empty tool_calls list - }, - "logprobs": None, - "finish_reason": None, - } - ], - } - # Empty tool_calls should return False - assert ( - initialized_custom_stream_wrapper.is_chunk_non_empty( - completion_obj={}, # completion_obj has no tool_calls - model_response=ModelResponseStream(**chunk), - response_obj={}, - ) - is False - ) - - -def test_is_chunk_non_empty_with_valid_tool_calls( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """ - Test that is_chunk_non_empty returns True when tool_calls has valid entries. - - Companion test for https://github.com/BerriAI/litellm/issues/17425 - Non-empty tool_calls in delta should be considered non-empty chunks. - """ - chunk = { - "id": "test-chunk-id", - "object": "chat.completion.chunk", - "created": 1741037890, - "model": "claude-sonnet-4-20250514", - "choices": [ - { - "index": 0, - "delta": { - "content": None, - "tool_calls": [ - { - "id": "call_123", - "type": "function", - "function": { - "name": "get_weather", - "arguments": '{"location": "NYC"}', - }, - } - ], - }, - "logprobs": None, - "finish_reason": None, - } - ], - } - # Non-empty tool_calls should return True - assert ( - initialized_custom_stream_wrapper.is_chunk_non_empty( - completion_obj={}, - model_response=ModelResponseStream(**chunk), - response_obj={}, - ) - is True - ) - - -def _make_chunk(content: Optional[str]) -> ModelResponseStream: - return ModelResponseStream( - id="test", - created=1741037890, - model="test-model", - choices=[StreamingChoices(index=0, delta=Delta(content=content))], - ) - - -def _build_chunks(pattern: list[str], N: int) -> list[ModelResponseStream]: - """ - Build a list of chunks based on a pattern specification. - """ - chunks = [] - for i, p in enumerate(pattern): - if p == "same": - chunks.append(_make_chunk("same_chunk")) - elif p == "diff": - chunks.append(_make_chunk(f"chunk_{i}")) - else: - chunks.append(_make_chunk(p)) - return chunks - - -_REPETITION_TEST_CASES = [ - # Basic cases - pytest.param( - ["same"] * litellm.REPEATED_STREAMING_CHUNK_LIMIT, - True, - id="all_identical_raises", - ), - pytest.param( - ["same"] * (litellm.REPEATED_STREAMING_CHUNK_LIMIT - 1), - False, - id="below_threshold_no_raise", - ), - pytest.param( - [None] * litellm.REPEATED_STREAMING_CHUNK_LIMIT, - False, - id="none_content_no_raise", - ), - pytest.param( - [""] * litellm.REPEATED_STREAMING_CHUNK_LIMIT, - False, - id="empty_content_no_raise", - ), - # Short content (len <= 2) should not raise - pytest.param( - ["##"] * litellm.REPEATED_STREAMING_CHUNK_LIMIT, - False, - id="short_content_2chars_no_raise", - ), - pytest.param( - ["{"] * litellm.REPEATED_STREAMING_CHUNK_LIMIT, - False, - id="short_content_1char_no_raise", - ), - pytest.param( - ["ab"] * litellm.REPEATED_STREAMING_CHUNK_LIMIT, - False, - id="short_content_2chars_ab_no_raise", - ), - # All different chunks - pytest.param( - ["diff"] * litellm.REPEATED_STREAMING_CHUNK_LIMIT, - False, - id="all_different_no_raise", - ), - # One chunk different at various positions - pytest.param( - ["different_first"] + ["same"] * (litellm.REPEATED_STREAMING_CHUNK_LIMIT - 1), - False, - id="first_chunk_different_no_raise", - ), - pytest.param( - ["same"] * (litellm.REPEATED_STREAMING_CHUNK_LIMIT - 1) + ["different_last"], - False, - id="last_chunk_different_no_raise", - ), - pytest.param( - ["same"] * (litellm.REPEATED_STREAMING_CHUNK_LIMIT // 2 + 1) - + ["different_mid"] - + ["same"] - * ( - litellm.REPEATED_STREAMING_CHUNK_LIMIT - - litellm.REPEATED_STREAMING_CHUNK_LIMIT // 2 - + 1 - ), - False, - id="middle_chunk_different_no_raise", - ), - pytest.param( - ["same"] * (litellm.REPEATED_STREAMING_CHUNK_LIMIT - 2) + ["diff", "diff"], - False, - id="last_two_different_no_raise", - ), - pytest.param( - ["diff"] * litellm.REPEATED_STREAMING_CHUNK_LIMIT - + ["same"] * litellm.REPEATED_STREAMING_CHUNK_LIMIT - + ["diff"], - True, - id="in_between_same_and_diff_raise", - ), -] - - -@pytest.mark.parametrize("chunks_pattern,should_raise", _REPETITION_TEST_CASES) -def test_raise_on_model_repetition( - initialized_custom_stream_wrapper: CustomStreamWrapper, - chunks_pattern: list, - should_raise: bool, -): - wrapper = initialized_custom_stream_wrapper - chunks = _build_chunks(chunks_pattern, len(chunks_pattern)) - - if should_raise: - def _feed(): - for chunk in chunks: - wrapper.chunks.append(chunk) - wrapper.raise_on_model_repetition() - - with pytest.raises(litellm.InternalServerError) as exc_info: - _feed() - assert "repeating the same chunk" in str(exc_info.value) - else: - for chunk in chunks: - wrapper.chunks.append(chunk) - wrapper.raise_on_model_repetition() - - -@pytest.mark.parametrize( - "empty_chunk_index", - [-1, -2], - ids=["last_chunk_empty", "second_to_last_chunk_empty"], -) -def test_raise_on_model_repetition_tolerates_empty_choices( - initialized_custom_stream_wrapper: CustomStreamWrapper, - empty_chunk_index: int, -): - """ - Regression test for https://github.com/BerriAI/litellm/issues/28884 - - Vertex Gemini Flash / Flash Lite with web search streaming emits - metadata-only and usage-only chunks that carry no choices. These are - appended to self.chunks, and raise_on_model_repetition() previously - accessed choices[0] unconditionally, raising IndexError mid-stream - (surfaced to users as MidStreamFallbackError -> APIConnectionError). - """ - wrapper = initialized_custom_stream_wrapper - - chunks = [ - _make_chunk("hello world"), - ModelResponseStream( - id="usage-only", - created=1741037890, - model="vertex_ai/gemini-3.1-flash-lite", - choices=[], - usage=Usage(prompt_tokens=10, completion_tokens=0, total_tokens=10), - ), - ] - if empty_chunk_index == -2: - chunks.append(_make_chunk("hello world again")) - - for chunk in chunks: - wrapper.chunks.append(chunk) - wrapper.raise_on_model_repetition() - - -def test_usage_chunk_after_finish_reason_updates_hidden_params(logging_obj): - """ - Test that provider-reported usage from a post-finish_reason chunk - is surfaced in _hidden_params even when stream_options is NOT set. - - Reproduces issue #20760: OpenRouter sends a final chunk with usage data - after the finish_reason chunk. The hidden_params["usage"] on the last - user-visible chunk was being calculated before this usage chunk arrived, - resulting in zeros. The fix recalculates it in the StopIteration handler - after stream_chunk_builder processes all chunks. - """ - # Simulate OpenRouter's actual streaming pattern: - # 1) content chunk - # 2) finish_reason chunk (content="") - # 3) usage chunk (content="", finish_reason=None, usage={...}) - chunks = [ - ModelResponseStream( - id="gen-abc", - object="chat.completion.chunk", - created=1000000, - model="openrouter/openai/gpt-4o-mini", - choices=[ - StreamingChoices( - index=0, - delta=Delta(role="assistant", content="Hello"), - finish_reason=None, - ) - ], - ), - ModelResponseStream( - id="gen-abc", - object="chat.completion.chunk", - created=1000000, - model="openrouter/openai/gpt-4o-mini", - choices=[ - StreamingChoices( - index=0, - delta=Delta(content=""), - finish_reason="stop", - ) - ], - ), - ModelResponseStream( - id="gen-abc", - object="chat.completion.chunk", - created=1000000, - model="openrouter/openai/gpt-4o-mini", - choices=[ - StreamingChoices( - index=0, - delta=Delta(role="assistant", content=""), - finish_reason=None, - ) - ], - usage=Usage( - prompt_tokens=20, - completion_tokens=135, - total_tokens=155, - ), - ), - ] - - # Create a CustomStreamWrapper with NO stream_options - wrapper = CustomStreamWrapper( - completion_stream=ModelResponseListIterator(model_responses=chunks), - model="openrouter/openai/gpt-4o-mini", - logging_obj=logging_obj, - custom_llm_provider="openrouter", - stream_options=None, - ) - - # Consume the stream - collected = [] - for chunk in wrapper: - collected.append(chunk) - - # The last user-visible chunk's _hidden_params["usage"] should - # contain the provider-reported values, not zeros. - last_chunk = collected[-1] - hidden_usage = last_chunk._hidden_params.get("usage") - assert hidden_usage is not None, "Expected usage in _hidden_params" - assert ( - hidden_usage.prompt_tokens == 20 - ), f"Expected prompt_tokens=20 from provider, got {hidden_usage.prompt_tokens}" - assert ( - hidden_usage.completion_tokens == 135 - ), f"Expected completion_tokens=135 from provider, got {hidden_usage.completion_tokens}" - - -@pytest.mark.asyncio -async def test_custom_stream_wrapper_aclose(): - """Test that aclose() delegates to the underlying completion_stream's aclose()""" - mock_stream = AsyncMock() - mock_stream.aclose = AsyncMock() - - wrapper = CustomStreamWrapper( - completion_stream=mock_stream, - model=None, - logging_obj=MagicMock(), - custom_llm_provider=None, - ) - - await wrapper.aclose() - mock_stream.aclose.assert_awaited_once() - - -@pytest.mark.asyncio -async def test_custom_stream_wrapper_aclose_no_underlying(): - """Test that aclose() is safe when completion_stream has no aclose method""" - mock_stream = MagicMock(spec=[]) # No aclose attribute - - wrapper = CustomStreamWrapper( - completion_stream=mock_stream, - model=None, - logging_obj=MagicMock(), - custom_llm_provider=None, - ) - - # Should not raise - await wrapper.aclose() - - -@pytest.mark.asyncio -async def test_custom_stream_wrapper_aclose_none_stream(): - """Test that aclose() is safe when completion_stream is None""" - wrapper = CustomStreamWrapper( - completion_stream=None, - model=None, - logging_obj=MagicMock(), - custom_llm_provider=None, - ) - - # Should not raise - await wrapper.aclose() - - -def test_content_not_dropped_when_finish_reason_already_set( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """ - Regression test for #22098: Vertex AI Claude streaming truncation. - - When content_block_delta and message_delta arrive in rapid succession, - received_finish_reason can be set BEFORE the last content chunk is - processed. The old code raised StopIteration unconditionally, dropping - content. The fix checks for text/tool_use content before stopping. - """ - initialized_custom_stream_wrapper.received_finish_reason = "stop" - initialized_custom_stream_wrapper.custom_llm_provider = "anthropic" - - content_chunk = { - "text": "world!", - "tool_use": None, - "is_finished": False, - "finish_reason": "", - "usage": None, - "index": 0, - } - - result = initialized_custom_stream_wrapper.chunk_creator(chunk=content_chunk) - - assert ( - result is not None - ), "chunk_creator() returned None — content was dropped (issue #22098)" - assert result.choices[0].delta.content == "world!" - - -def test_empty_chunk_still_stops_after_finish_reason_set( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """ - Companion test for #22098: an empty GenericStreamingChunk must still - raise StopIteration when received_finish_reason is already set. - """ - initialized_custom_stream_wrapper.received_finish_reason = "stop" - initialized_custom_stream_wrapper.custom_llm_provider = "anthropic" - - empty_chunk = { - "text": "", - "tool_use": None, - "is_finished": False, - "finish_reason": "", - "usage": None, - "index": 0, - } - - with pytest.raises(StopIteration): - initialized_custom_stream_wrapper.chunk_creator(chunk=empty_chunk) - - -def test_tool_use_not_dropped_when_finish_reason_already_set( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """ - Regression test for #22098: tool_use-only chunks must not be dropped - when received_finish_reason is already set. - """ - initialized_custom_stream_wrapper.received_finish_reason = "stop" - initialized_custom_stream_wrapper.custom_llm_provider = "anthropic" - - tool_chunk = { - "text": "", - "tool_use": { - "id": "call_1", - "type": "function", - "function": {"name": "get_weather", "arguments": "{}"}, - }, - "is_finished": False, - "finish_reason": "", - "usage": None, - "index": 0, - } - - result = initialized_custom_stream_wrapper.chunk_creator(chunk=tool_chunk) - - assert ( - result is not None - ), "chunk_creator() returned None — tool_use data was dropped" - - tool_calls = result.choices[0].delta.tool_calls - assert ( - tool_calls is not None and len(tool_calls) > 0 - ), "tool_calls should contain at least one tool call" - assert tool_calls[0].id == "call_1" - assert tool_calls[0].function.name == "get_weather" - - -def test_usage_only_chunk_not_dropped_when_finish_reason_already_set( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """ - Regression test: usage-only chunks must not be dropped once finish_reason - is already set. Dropping these chunks can lose terminal finish_reason in - downstream Responses API streaming translation. - """ - initialized_custom_stream_wrapper.received_finish_reason = "content_filter" - initialized_custom_stream_wrapper.custom_llm_provider = "anthropic" - - usage_only_chunk = { - "text": "", - "tool_use": None, - "is_finished": False, - "finish_reason": "", - "usage": {"prompt_tokens": 10, "completion_tokens": 1, "total_tokens": 11}, - "index": 0, - } - - result = initialized_custom_stream_wrapper.chunk_creator(chunk=usage_only_chunk) - - assert result is not None, "usage-only chunk should not be dropped" - assert result.choices[0].finish_reason == "content_filter" - assert result.usage is not None - - -def _run_dispatch(wrapper: CustomStreamWrapper, chunk): - model_response = wrapper.model_response_creator() - completion_obj = {"content": ""} - result = wrapper._dispatch_provider_chunk( - chunk=chunk, - model_response=model_response, - completion_obj=completion_obj, - ) - return result, model_response, completion_obj - - -def test_dispatch_vllm_extracts_output_text( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """vllm chunks expose text at chunk[0].outputs[0].text; the dispatch must - surface that as the content and report a parsed result.""" - initialized_custom_stream_wrapper.custom_llm_provider = "vllm" - - class _Output: - text = "hello from vllm" - - class _VLLMChunk: - outputs = [_Output()] - - result, _, completion_obj = _run_dispatch( - initialized_custom_stream_wrapper, [_VLLMChunk()] - ) - - assert isinstance(result, _ProviderChunkParsed) - assert completion_obj["content"] == "hello from vllm" - - -def test_dispatch_petals_slices_completion_stream( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """petals fakes streaming by slicing 30 chars off the buffered completion - stream each call, leaving the remainder for the next chunk.""" - initialized_custom_stream_wrapper.custom_llm_provider = "petals" - initialized_custom_stream_wrapper.completion_stream = "A" * 50 - - result, _, completion_obj = _run_dispatch( - initialized_custom_stream_wrapper, chunk=None - ) - - assert isinstance(result, _ProviderChunkParsed) - assert completion_obj["content"] == "A" * 30 - assert initialized_custom_stream_wrapper.completion_stream == "A" * 20 - assert initialized_custom_stream_wrapper.received_finish_reason is None - - -def test_dispatch_petals_empty_stream_sets_stop( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """An exhausted petals stream marks the turn finished with a stop reason.""" - initialized_custom_stream_wrapper.custom_llm_provider = "petals" - initialized_custom_stream_wrapper.completion_stream = "" - - result, _, completion_obj = _run_dispatch( - initialized_custom_stream_wrapper, chunk=None - ) - - assert isinstance(result, _ProviderChunkParsed) - assert completion_obj["content"] == "" - assert initialized_custom_stream_wrapper.received_finish_reason == "stop" - - -def test_dispatch_petals_empty_stream_after_finish_raises( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """Once petals has already finished, an empty stream signals end-of-iteration.""" - initialized_custom_stream_wrapper.custom_llm_provider = "petals" - initialized_custom_stream_wrapper.completion_stream = "" - initialized_custom_stream_wrapper.received_finish_reason = "stop" - - with pytest.raises(StopIteration): - _run_dispatch(initialized_custom_stream_wrapper, chunk=None) - - -def test_dispatch_palm_slices_completion_stream( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """palm uses the same fake-streaming slice strategy as petals.""" - initialized_custom_stream_wrapper.custom_llm_provider = "palm" - initialized_custom_stream_wrapper.completion_stream = "B" * 40 - - result, _, completion_obj = _run_dispatch( - initialized_custom_stream_wrapper, chunk=None - ) - - assert isinstance(result, _ProviderChunkParsed) - assert completion_obj["content"] == "B" * 30 - assert initialized_custom_stream_wrapper.completion_stream == "B" * 10 - - -def test_dispatch_cached_response_extracts_delta( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """cached_response replays a stored ModelResponseStream; the dispatch lifts - its delta content, finish_reason and id back onto the live response.""" - initialized_custom_stream_wrapper.custom_llm_provider = "cached_response" - chunk = ModelResponseStream( - id="chatcmpl-cache-1", - choices=[ - StreamingChoices( - index=0, - delta=Delta(content="cached text"), - finish_reason="stop", - ) - ], - ) - - result, model_response, completion_obj = _run_dispatch( - initialized_custom_stream_wrapper, chunk - ) - - assert isinstance(result, _ProviderChunkParsed) - assert completion_obj["content"] == "cached text" - assert initialized_custom_stream_wrapper.received_finish_reason == "stop" - assert model_response.id == "chatcmpl-cache-1" - assert initialized_custom_stream_wrapper.response_id == "chatcmpl-cache-1" - - -def test_dispatch_vertex_ai_legacy_text_and_finish_reason( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """Legacy vertex_ai chunks (non-ModelResponseStream) expose .text and a - candidate finish_reason enum that must be normalised to an OpenAI reason.""" - initialized_custom_stream_wrapper.custom_llm_provider = "vertex_ai" - - class _FinishReason: - name = "STOP" - - class _Candidate: - finish_reason = _FinishReason() - - class _VertexChunk: - candidates = [_Candidate()] - text = "vertex content" - - result, _, completion_obj = _run_dispatch( - initialized_custom_stream_wrapper, _VertexChunk() - ) - - assert isinstance(result, _ProviderChunkParsed) - assert completion_obj["content"] == "vertex content" - assert initialized_custom_stream_wrapper.received_finish_reason == "stop" - - -def test_dispatch_vertex_ai_legacy_without_candidates_stringifies_chunk( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """A legacy vertex_ai chunk with no candidates falls back to str(chunk).""" - initialized_custom_stream_wrapper.custom_llm_provider = "vertex_ai" - - class _RawChunk: - def __str__(self) -> str: - return "raw vertex blob" - - result, _, completion_obj = _run_dispatch( - initialized_custom_stream_wrapper, _RawChunk() - ) - - assert isinstance(result, _ProviderChunkParsed) - assert completion_obj["content"] == "raw vertex blob" - - -def test_dispatch_vertex_ai_legacy_function_call( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """A legacy vertex_ai chunk whose part has no text but carries a - function_call is converted into an OpenAI tool-call delta.""" - initialized_custom_stream_wrapper.custom_llm_provider = "vertex_ai" - - class _FunctionCall: - name = "get_weather" - args = {"location": "SF"} - - class _Part: - function_call = _FunctionCall() - - class _Content: - parts = [_Part()] - - class _FinishReason: - name = "STOP" - - class _Candidate: - content = _Content() - finish_reason = _FinishReason() - - class _VertexFunctionChunk: - candidates = [_Candidate()] - - @property - def text(self): - raise RuntimeError("Part has no text.") - - result, _, _ = _run_dispatch( - initialized_custom_stream_wrapper, _VertexFunctionChunk() - ) - - assert isinstance(result, _ProviderChunkParsed) - tool_calls = result.response_obj["original_chunk"].choices[0].delta.tool_calls - assert tool_calls[0].function.name == "get_weather" - assert json.loads(tool_calls[0].function.arguments) == {"location": "SF"} - assert initialized_custom_stream_wrapper.received_finish_reason == "stop" - - -def test_dispatch_custom_provider_returns_chunk_early( - monkeypatch: pytest.MonkeyPatch, - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """A registered custom provider passes its already-OpenAI-shaped chunk - straight through as an early return rather than re-parsing it.""" - monkeypatch.setattr(litellm, "_custom_providers", ["my-custom-llm"]) - initialized_custom_stream_wrapper.custom_llm_provider = "my-custom-llm" - chunk = ModelResponseStream( - choices=[ - StreamingChoices(index=0, delta=Delta(content="hi"), finish_reason=None) - ] - ) - - result, _, _ = _run_dispatch(initialized_custom_stream_wrapper, chunk) - - assert isinstance(result, _ProviderChunkEarlyReturn) - assert result.value is chunk - - -def test_dispatch_custom_provider_finish_only_returns_none_early( - monkeypatch: pytest.MonkeyPatch, - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """A custom-provider chunk that carries only a finish_reason (no content) - records the reason and returns None so no empty delta is emitted.""" - monkeypatch.setattr(litellm, "_custom_providers", ["my-custom-llm"]) - initialized_custom_stream_wrapper.custom_llm_provider = "my-custom-llm" - chunk = ModelResponseStream( - choices=[ - StreamingChoices(index=0, delta=Delta(content=None), finish_reason="stop") - ] - ) - - result, _, _ = _run_dispatch(initialized_custom_stream_wrapper, chunk) - - assert isinstance(result, _ProviderChunkEarlyReturn) - assert result.value is None - assert initialized_custom_stream_wrapper.received_finish_reason == "stop" - - -def test_dispatch_text_completion_codestral_parses_chunk( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """text-completion-codestral streams raw SSE JSON strings that the dispatch - routes through CodestralTextCompletionConfig to extract content/finish.""" - initialized_custom_stream_wrapper.custom_llm_provider = "text-completion-codestral" - chunk = json.dumps( - {"choices": [{"delta": {"content": "codestral text"}, "finish_reason": "stop"}]} - ) - - result, _, completion_obj = _run_dispatch(initialized_custom_stream_wrapper, chunk) - - assert isinstance(result, _ProviderChunkParsed) - assert completion_obj["content"] == "codestral text" - assert initialized_custom_stream_wrapper.received_finish_reason == "stop" - - -def test_dispatch_text_completion_codestral_requires_string( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """The codestral branch only knows how to parse raw strings; anything else - is a programming error and must surface loudly.""" - initialized_custom_stream_wrapper.custom_llm_provider = "text-completion-codestral" - - with pytest.raises(ValueError, match="chunk is not a string: \\{'not': 'a string'\\}"): - _run_dispatch(initialized_custom_stream_wrapper, {"not": "a string"}) - - -def test_dispatch_triton_stream( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """triton stream chunks arrive as dicts keyed by text_output/stop_reason.""" - initialized_custom_stream_wrapper.custom_llm_provider = "triton" - chunk = {"text_output": "triton text", "is_finished": True, "stop_reason": "stop"} - - result, _, completion_obj = _run_dispatch(initialized_custom_stream_wrapper, chunk) - - assert isinstance(result, _ProviderChunkParsed) - assert completion_obj["content"] == "triton text" - assert initialized_custom_stream_wrapper.received_finish_reason == "stop" - - -def test_dispatch_ai21_decodes_completion( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """ai21 does fake streaming over a single byte-encoded JSON completion.""" - initialized_custom_stream_wrapper.custom_llm_provider = "ai21" - chunk = json.dumps({"completions": [{"data": {"text": "ai21 text"}}]}).encode( - "utf-8" - ) - - result, _, completion_obj = _run_dispatch(initialized_custom_stream_wrapper, chunk) - - assert isinstance(result, _ProviderChunkParsed) - assert completion_obj["content"] == "ai21 text" - assert initialized_custom_stream_wrapper.received_finish_reason == "stop" - - -def test_dispatch_text_completion_openai_with_usage( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """text-completion-openai chunks expose choices[].text and an optional usage - block that the dispatch lifts onto the model response.""" - initialized_custom_stream_wrapper.custom_llm_provider = "text-completion-openai" - - class _Choice: - text = "oai text" - finish_reason = "stop" - - class _Usage: - prompt_tokens = 5 - completion_tokens = 3 - total_tokens = 8 - - class _TextChunk: - choices = [_Choice()] - usage = _Usage() - - result, model_response, completion_obj = _run_dispatch( - initialized_custom_stream_wrapper, _TextChunk() - ) - - assert isinstance(result, _ProviderChunkParsed) - assert completion_obj["content"] == "oai text" - assert initialized_custom_stream_wrapper.received_finish_reason == "stop" - assert model_response.usage.prompt_tokens == 5 - assert model_response.usage.total_tokens == 8 - - -@pytest.mark.asyncio -async def test_custom_stream_wrapper_anext_does_not_block_event_loop_for_sync_iterators( - logging_obj: Logging, -): - """ - Regression test: __anext__ must not call blocking next() on a sync iterator on the - event loop thread. This happens for some provider streams which are sync iterators - but used in async contexts (e.g. boto3-style streaming). - """ - - class BlockingIterator: - def __init__(self, chunks, delay_s: float): - self._it = iter(chunks) - self._delay_s = delay_s - - def __iter__(self): - return self - - def __next__(self): - time.sleep(self._delay_s) # simulate blocking I/O - return next(self._it) - - test_chunk = ModelResponseStream( - id="chatcmpl-test", - created=int(time.time()), - model="test-model", - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason="stop", - index=0, - delta=Delta( - provider_specific_fields=None, - content="hello", - role="assistant", - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields={}, - usage=None, - ) - - # Delay is intentionally > the wait_for timeout used to detect event loop blocking. - wrapper = CustomStreamWrapper( - completion_stream=BlockingIterator([test_chunk], delay_s=0.3), - model="test-model", - logging_obj=logging_obj, - custom_llm_provider="cached_response", - ) - - tick_event = asyncio.Event() - - async def background_tick(): - await asyncio.sleep(0.05) - tick_event.set() - - # Run the two coroutines concurrently and measure wall time. - # If __anext__ blocks the event loop, background_tick can't run and the gather - # takes the full 0.3 s delay; if non-blocking both finish within ~0.35 s total. - start = asyncio.get_event_loop().time() - - out, _ = await asyncio.gather( - wrapper.__anext__(), - background_tick(), - ) - - elapsed = asyncio.get_event_loop().time() - start - assert isinstance(out, ModelResponseStream) - # background_tick sleeps 0.05 s; total must finish well under 2 × 0.3 s - assert elapsed < 0.5, f"Event loop was likely blocked (elapsed={elapsed:.2f}s)" - - -@pytest.mark.asyncio -async def test_custom_stream_wrapper_anext_exhaustion_raises_stop_async_iteration( - logging_obj: Logging, -): - """ - PEP 479 regression: when a sync iterator is exhausted, asyncio.to_thread(next, it) - raises StopIteration inside a coroutine, which Python converts to RuntimeError. - The wrapper must catch StopIteration in the thread and raise StopAsyncIteration - in the coroutine instead, so callers get clean stream termination. - """ - - class SingleChunkIterator: - def __init__(self, chunk: ModelResponseStream): - self._chunk = chunk - self._done = False - - def __iter__(self): - return self - - def __next__(self): - if self._done: - raise StopIteration - self._done = True - return self._chunk - - test_chunk = ModelResponseStream( - id="chatcmpl-exhaustion-test", - created=int(time.time()), - model="test-model", - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason="stop", - index=0, - delta=Delta( - provider_specific_fields=None, - content="done", - role="assistant", - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields={}, - usage=None, - ) - - wrapper = CustomStreamWrapper( - completion_stream=SingleChunkIterator(test_chunk), - model="test-model", - logging_obj=logging_obj, - custom_llm_provider="cached_response", - ) - - # Drain the wrapper fully. The wrapper's except-handler calls finish_reason_handler() - # on the first StopAsyncIteration (sent_last_chunk=False→True), then re-raises on the - # next call. What must NOT happen is a RuntimeError from PEP 479 converting - # StopIteration (raised inside the thread) to RuntimeError inside the coroutine. - try: - while True: - await wrapper.__anext__() - except StopAsyncIteration: - pass # expected clean termination - except RuntimeError as e: - pytest.fail(f"PEP 479 regression: StopIteration leaked as RuntimeError: {e}") - - -# Azure streaming chunks that reproduce issue #24221: -# Azure sends an initial chunk with prompt_filter_results and choices=[], -# then a chunk with role='assistant' and content='', then content chunks. -# With stream_options.include_usage=True, the empty-choices chunk was -# forwarded with an inflated default choice, consuming the sent_first_chunk -# flag and causing strip_role_from_delta to strip the role from the real -# first chunk. -_AZURE_CHUNKS_WITH_PROMPT_FILTER = [ - # Chunk 1: prompt_filter_results, no choices (Azure-specific) - ModelResponseStream( - id="chatcmpl-abc123", - created=1742056047, - model=None, - object="chat.completion.chunk", - choices=[], - usage=None, - ), - # Chunk 2: first real chunk with role='assistant' and empty content - ModelResponseStream( - id="chatcmpl-abc123", - created=1742056047, - model=None, - object="chat.completion.chunk", - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta(content="", role="assistant"), - ) - ], - usage=None, - ), - # Chunk 3: content - ModelResponseStream( - id="chatcmpl-abc123", - created=1742056047, - model=None, - object="chat.completion.chunk", - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta(content="Hello!"), - ) - ], - usage=None, - ), - # Chunk 4: finish_reason - ModelResponseStream( - id="chatcmpl-abc123", - created=1742056047, - model=None, - object="chat.completion.chunk", - choices=[ - StreamingChoices( - finish_reason="stop", - index=0, - delta=Delta(), - ) - ], - usage=None, - ), - # Chunk 5: final usage chunk, no choices - ModelResponseStream( - id="chatcmpl-abc123", - created=1742056047, - model=None, - object="chat.completion.chunk", - choices=[], - usage=Usage( - completion_tokens=10, - prompt_tokens=20, - total_tokens=30, - ), - ), -] - - -@pytest.mark.parametrize("sync_mode", [True, False], ids=["sync", "async"]) -@pytest.mark.asyncio -async def test_azure_streaming_role_preserved_with_include_usage(sync_mode: bool): - """ - Regression test for https://github.com/BerriAI/litellm/issues/24221 - - Azure sends an initial chunk with choices=[] (prompt_filter_results) - before the first content chunk. With stream_options.include_usage=True, - this chunk was forwarded with an inflated default choice, which: - 1. Consumed the sent_first_chunk flag - 2. Caused strip_role_from_delta to strip role from the real first chunk - - The fix ensures: - - Chunks with choices=[] are forwarded faithfully (no inflated choices) - - sent_first_chunk is only marked for chunks with real choices - - Chunks with role in delta are not discarded as empty - """ - completion_stream = ModelResponseListIterator( - model_responses=_AZURE_CHUNKS_WITH_PROMPT_FILTER - ) - - response = CustomStreamWrapper( - completion_stream=completion_stream, - model="azure/gpt-5-nano", - custom_llm_provider="azure", - logging_obj=Logging( - model="azure/gpt-5-nano", - messages=[{"role": "user", "content": "Hey"}], - stream=True, - call_type="completion", - start_time=time.time(), - litellm_call_id="12345", - function_id="1245", - ), - stream_options={"include_usage": True}, - ) - - chunks = [] - if sync_mode: - for chunk in response: - chunks.append(chunk) - else: - async for chunk in response: - chunks.append(chunk) - - # The prompt_filter chunk should be forwarded with choices=[] - assert ( - len(chunks[0].choices) == 0 - ), f"Expected prompt_filter chunk with choices=[], got {len(chunks[0].choices)} choices" - - # At least one chunk must have role='assistant' in its delta - has_role = any( - len(c.choices) > 0 and getattr(c.choices[0].delta, "role", None) == "assistant" - for c in chunks - ) - assert has_role, ( - "No chunk contained role='assistant' in delta (issue #24221). " - "Chunk deltas: " - + str([c.choices[0].delta if c.choices else "no choices" for c in chunks]) - ) - - -def test_gemini_legacy_vertex_stop_finish_reason_normalised(): - """ - The legacy vertex_ai SDK streaming path sets finish_reason from a proto enum - whose .name attribute is an uppercase string (e.g. "STOP", "MAX_TOKENS"). - Before the fix, received_finish_reason was stored as "STOP" which never - matched "stop" in finish_reason_handler, silently breaking the tool_calls - override. After the fix, map_finish_reason() is applied so the value is - always an OpenAI-normalised lowercase string. - """ - wrapper = CustomStreamWrapper( - completion_stream=None, - model="gemini-1.5-pro", - logging_obj=MagicMock(), - custom_llm_provider="vertex_ai", - ) - - # Simulate a proto-like chunk: .candidates[0].finish_reason.name == "STOP" - mock_finish_reason = MagicMock() - mock_finish_reason.name = "STOP" - mock_candidate = MagicMock() - mock_candidate.finish_reason = mock_finish_reason - mock_chunk = MagicMock() - mock_chunk.candidates = [mock_candidate] - # Ensure the chunk is not treated as a ModelResponseStream - mock_chunk.__class__ = type("FakeProtoChunk", (), {}) - - with patch("litellm.litellm_core_utils.streaming_handler.proto", create=True): - wrapper.chunk_creator(chunk=mock_chunk) - - assert wrapper.received_finish_reason == "stop", ( - f"Expected 'stop' but got {wrapper.received_finish_reason!r}. " - "map_finish_reason() was not applied to the Gemini enum name." - ) - - -def test_gemini_legacy_vertex_tool_calls_finish_reason_with_stop_enum(): - """ - When Gemini emits finish_reason STOP alongside tool-call content, the final - chunk must report finish_reason='tool_calls'. This requires that the raw - "STOP" enum name is first normalised to lowercase "stop" by map_finish_reason() - so that finish_reason_handler's equality check fires correctly. - """ - wrapper = CustomStreamWrapper( - completion_stream=None, - model="gemini-1.5-pro", - logging_obj=MagicMock(), - custom_llm_provider="vertex_ai", - ) - - mock_finish_reason = MagicMock() - mock_finish_reason.name = "STOP" - mock_candidate = MagicMock() - mock_candidate.finish_reason = mock_finish_reason - mock_chunk = MagicMock() - mock_chunk.candidates = [mock_candidate] - mock_chunk.__class__ = type("FakeProtoChunk", (), {}) - - with patch("litellm.litellm_core_utils.streaming_handler.proto", create=True): - wrapper.chunk_creator(chunk=mock_chunk) - - # Signal that tool_calls were present in the stream - wrapper.tool_call = True - - final = wrapper.finish_reason_handler() - assert final.choices[0].finish_reason == "tool_calls", ( - f"Expected 'tool_calls' but got {final.choices[0].finish_reason!r}. " - "STOP enum was not normalised through map_finish_reason()." - ) - - -@pytest.mark.parametrize( - "finish_reason", ["stop", "tool_calls", "length", "content_filter"] -) -def test_chunk_creator_passes_through_model_response_stream( - initialized_custom_stream_wrapper: CustomStreamWrapper, - finish_reason: str, -): - """ - chunk_creator must pass ModelResponseStream chunks from custom providers - straight through and preserve finish_reason exactly — not force-cast to GChunk. - Regression test for issue #27389. - """ - initialized_custom_stream_wrapper.custom_llm_provider = "my-custom-provider" - litellm._custom_providers.append("my-custom-provider") - - chunk = ModelResponseStream( - id="test-id", - choices=[ - StreamingChoices( - index=0, - delta=Delta(content="Hello", role="assistant"), - finish_reason=finish_reason, - ) - ], - ) - - result = initialized_custom_stream_wrapper.chunk_creator(chunk=chunk) - - litellm._custom_providers.remove("my-custom-provider") - - assert result is not None - assert initialized_custom_stream_wrapper.received_finish_reason == finish_reason - - -def test_chunk_creator_drops_empty_finish_chunk( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """ - A ModelResponseStream chunk with finish_reason but no content should return - None so finish_reason_handler() synthesises the final chunk — mirrors GChunk - behaviour via is_chunk_non_empty. - """ - initialized_custom_stream_wrapper.custom_llm_provider = "my-custom-provider" - litellm._custom_providers.append("my-custom-provider") - - chunk = ModelResponseStream( - id="test-id", - choices=[ - StreamingChoices( - index=0, - delta=Delta(content=""), - finish_reason="stop", - ) - ], - ) - - result = initialized_custom_stream_wrapper.chunk_creator(chunk=chunk) - - litellm._custom_providers.remove("my-custom-provider") - - assert result is None - assert initialized_custom_stream_wrapper.received_finish_reason == "stop" - - -def test_chunk_creator_stops_iteration_on_trailing_chunk( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """ - After received_finish_reason is set, any empty trailing chunk (e.g. provider - metadata flush) must raise StopIteration to end the stream cleanly. - """ - initialized_custom_stream_wrapper.custom_llm_provider = "my-custom-provider" - initialized_custom_stream_wrapper.received_finish_reason = "stop" - litellm._custom_providers.append("my-custom-provider") - - trailing_chunk = ModelResponseStream( - id="test-id", - choices=[ - StreamingChoices( - index=0, - delta=Delta(content=None), - finish_reason="stop", - ) - ], - ) - - with pytest.raises(StopIteration): - initialized_custom_stream_wrapper.chunk_creator(chunk=trailing_chunk) - - litellm._custom_providers.remove("my-custom-provider") - - -def test_chunk_creator_strips_finish_reason_from_content_chunk( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """ - When content and finish_reason arrive in the same chunk, finish_reason must be - stripped so finish_reason_handler() emits it on the synthetic terminal chunk — - preventing two terminal chunks (double finish_reason bug). - """ - initialized_custom_stream_wrapper.custom_llm_provider = "my-custom-provider" - litellm._custom_providers.append("my-custom-provider") - - chunk = ModelResponseStream( - id="test-id", - choices=[ - StreamingChoices( - index=0, - delta=Delta(content="Hello"), - finish_reason="stop", - ) - ], - ) - - result = initialized_custom_stream_wrapper.chunk_creator(chunk=chunk) - - litellm._custom_providers.remove("my-custom-provider") - - assert result is not None - assert ( - result.choices[0].finish_reason is None - ), "finish_reason must be stripped from content chunks to avoid double terminal chunks" - assert initialized_custom_stream_wrapper.received_finish_reason == "stop" - - -def test_chunk_creator_tool_calls_not_dropped_on_finish( - initialized_custom_stream_wrapper: CustomStreamWrapper, -): - """ - A terminal chunk with finish_reason="tool_calls" and delta.tool_calls must NOT - be silently dropped — tool_calls counts as content so the chunk is passed through - (with finish_reason stripped) rather than returning None. - """ - from litellm.types.utils import ChatCompletionDeltaToolCall, Function - - initialized_custom_stream_wrapper.custom_llm_provider = "my-custom-provider" - litellm._custom_providers.append("my-custom-provider") - - chunk = ModelResponseStream( - id="test-id", - choices=[ - StreamingChoices( - index=0, - delta=Delta( - content=None, - tool_calls=[ - ChatCompletionDeltaToolCall( - id="call_abc", - function=Function( - name="get_weather", arguments='{"city":"NYC"}' - ), - type="function", - index=0, - ) - ], - ), - finish_reason="tool_calls", - ) - ], - ) - - result = initialized_custom_stream_wrapper.chunk_creator(chunk=chunk) - - litellm._custom_providers.remove("my-custom-provider") - - assert result is not None, "tool_calls chunk must not be dropped" - assert result.choices[0].delta.tool_calls is not None - assert result.choices[0].finish_reason is None - assert initialized_custom_stream_wrapper.received_finish_reason == "tool_calls" - - -def test_record_partial_usage_for_failure_stashes_usage_and_cost(): - """A stream that breaks mid-flight must surface the usage assembled from the - chunks already delivered, plus its cost, on the logging object so the - failure handler records the real partial spend instead of zero. - """ - logging_obj = Logging( - model="gpt-4o-mini", - messages=[{"role": "user", "content": "Hey"}], - stream=True, - call_type="completion", - start_time=time.time(), - litellm_call_id="partial-usage-1", - function_id="1245", - ) - logging_obj.model_call_details["custom_llm_provider"] = "openai" - - wrapper = CustomStreamWrapper( - completion_stream=None, - model="gpt-4o-mini", - logging_obj=logging_obj, - custom_llm_provider="openai", - ) - wrapper.chunks = [ - ModelResponseStream( - id="chatcmpl-partial-1", - created=1742056047, - model="gpt-4o-mini", - object="chat.completion.chunk", - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta( - content="The Roman Empire began when", role="assistant" - ), - ) - ], - usage=Usage(prompt_tokens=30, completion_tokens=1, total_tokens=31), - ) - ] - - wrapper._record_partial_usage_for_failure() - - stashed = logging_obj.model_call_details["combined_usage_object"] - assert stashed.prompt_tokens == 30 - assert stashed.completion_tokens == 1 - assert stashed.total_tokens == 31 - assert isinstance(logging_obj.model_call_details["response_cost"], float) - - -def test_record_partial_usage_for_failure_noop_without_chunks(): - """With no chunks delivered there is nothing billed to recover, so the - failure stash must stay absent and not force a zero-usage row. - """ - logging_obj = Logging( - model="gpt-4o-mini", - messages=[{"role": "user", "content": "Hey"}], - stream=True, - call_type="completion", - start_time=time.time(), - litellm_call_id="partial-usage-2", - function_id="1245", - ) - wrapper = CustomStreamWrapper( - completion_stream=None, - model="gpt-4o-mini", - logging_obj=logging_obj, - custom_llm_provider="openai", - ) - wrapper.chunks = [] - - wrapper._record_partial_usage_for_failure() - - assert "combined_usage_object" not in logging_obj.model_call_details - - -def _wrapper_with_partial_chunks( - chunk_model: str, - usage: Optional[Usage] = None, - model: str = "gpt-4o-mini", - custom_llm_provider: str = "openai", -) -> tuple: - logging_obj = Logging( - model=model, - messages=[{"role": "user", "content": "Tell me a long story"}], - stream=True, - call_type="completion", - start_time=time.time(), - litellm_call_id="partial-usage-alias", - function_id="1245", - ) - logging_obj.model_call_details["custom_llm_provider"] = custom_llm_provider - logging_obj.optional_params = {} - wrapper = CustomStreamWrapper( - completion_stream=None, - model=model, - logging_obj=logging_obj, - custom_llm_provider=custom_llm_provider, - ) - wrapper.chunks = [ - ModelResponseStream( - id="chatcmpl-partial-alias-1", - created=1742056047, - model=chunk_model, - object="chat.completion.chunk", - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta( - content="The Roman Empire began when", role="assistant" - ), - ) - ], - usage=usage, - ) - ] - return wrapper, logging_obj - - -def test_record_partial_usage_for_failure_prices_alias_restamped_chunks_at_real_model(): - wrapper, logging_obj = _wrapper_with_partial_chunks( - chunk_model="bedrock-claude-opus-5", - usage=Usage(prompt_tokens=40, completion_tokens=5, total_tokens=45), - model="us.anthropic.claude-opus-5", - custom_llm_provider="bedrock", - ) - assert "bedrock/bedrock-claude-opus-5" not in litellm.model_cost - - wrapper._record_partial_usage_for_failure() - - stashed = logging_obj.model_call_details["combined_usage_object"] - assert stashed.completion_tokens == 5 - rates = litellm.model_cost["us.anthropic.claude-opus-5"] - expected = 40 * rates["input_cost_per_token"] + 5 * rates["output_cost_per_token"] - assert logging_obj.model_call_details["response_cost"] == pytest.approx(expected) - - -def test_record_partial_usage_for_failure_counts_prompt_tokens_from_request_messages(): - wrapper, logging_obj = _wrapper_with_partial_chunks(chunk_model="my-public-alias") - - wrapper._record_partial_usage_for_failure() - - stashed = logging_obj.model_call_details["combined_usage_object"] - assert stashed.prompt_tokens > 0 - - -def test_record_partial_usage_for_failure_backfills_missing_cache_fields(): - wrapper, logging_obj = _wrapper_with_partial_chunks(chunk_model="gpt-4o-mini") - - wrapper._record_partial_usage_for_failure() - - stashed = logging_obj.model_call_details["combined_usage_object"] - assert stashed.cache_creation_input_tokens == 0 - assert stashed.cache_read_input_tokens == 0 - assert stashed.prompt_tokens_details is not None - assert stashed.prompt_tokens_details.cached_tokens == 0 - - -def test_record_partial_usage_for_failure_prices_corrected_model_not_chunk_model(): - wrapper, logging_obj = _wrapper_with_partial_chunks( - chunk_model="claude-opus-5", - usage=Usage(prompt_tokens=40, completion_tokens=5, total_tokens=45), - model="gpt-4o-mini", - custom_llm_provider="openai", - ) - - wrapper._record_partial_usage_for_failure() - - rates = litellm.model_cost["gpt-4o-mini"] - expected = 40 * rates["input_cost_per_token"] + 5 * rates["output_cost_per_token"] - assert logging_obj.model_call_details["response_cost"] == pytest.approx(expected) - - -def test_record_partial_usage_for_failure_carries_up_openai_style_cached_tokens(): - recovered = Usage( - prompt_tokens=1000, - completion_tokens=10, - total_tokens=1010, - prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=500), - ) - wrapper, logging_obj = _wrapper_with_partial_chunks( - chunk_model="gpt-4o-mini", usage=recovered - ) - - wrapper._record_partial_usage_for_failure() - - stashed = logging_obj.model_call_details["combined_usage_object"] - assert stashed.cache_read_input_tokens == 500 - assert stashed.cache_creation_input_tokens == 0 - - -def test_record_partial_usage_for_failure_keeps_cache_values_recovered_from_chunks(): - recovered = Usage( - prompt_tokens=40, - completion_tokens=5, - total_tokens=45, - cache_read_input_tokens=7, - cache_creation_input_tokens=3, - ) - wrapper, logging_obj = _wrapper_with_partial_chunks( - chunk_model="gpt-4o-mini", usage=recovered - ) - - wrapper._record_partial_usage_for_failure() - - stashed = logging_obj.model_call_details["combined_usage_object"] - assert stashed.cache_read_input_tokens == 7 - assert stashed.cache_creation_input_tokens == 3 - assert stashed.prompt_tokens_details is not None - assert stashed.prompt_tokens_details.cached_tokens == 7 - - -@pytest.mark.parametrize("sync_mode", [True, False]) -@pytest.mark.asyncio -async def test_stream_chunk_builder_raise_at_end_of_stream_still_recovers_usage( - sync_mode, -): - """stream_chunk_builder re-raises (as APIError) on large agentic tool-use - streams. That raise originates inside the except-StopIteration handler, so - before the fix it escaped __next__/__anext__ and the request was dropped from - SpendLogs while the provider billed the tokens. The wrapper must catch it and - recover usage from the raw chunks so cost is still tracked.""" - final_usage_block = Usage( - completion_tokens=392, prompt_tokens=1799, total_tokens=2191 - ) - final_chunk = ModelResponseStream( - id="chatcmpl-raise-test", - created=1742056047, - model=None, - object="chat.completion.chunk", - choices=[ - StreamingChoices( - finish_reason="stop", - index=0, - delta=Delta(content="", role="assistant"), - ) - ], - usage=final_usage_block, - ) - test_chunks = bedrock_chunks + [final_chunk] - - logging_obj = Logging( - model="bedrock/claude-haiku-4-5-20251001-v1:0", - messages=[{"role": "user", "content": "Hey"}], - stream=True, - call_type="completion", - start_time=time.time(), - litellm_call_id="raise-test", - function_id="1245", - ) - - response = CustomStreamWrapper( - completion_stream=ModelResponseListIterator(model_responses=test_chunks), - model="bedrock/claude-haiku-4-5-20251001-v1:0", - custom_llm_provider="bedrock", - logging_obj=logging_obj, - stream_options={"include_usage": True}, - ) - - seen_usage = [] - with patch.object( - litellm, - "stream_chunk_builder", - side_effect=Exception("simulated assembly failure"), - ): - # before the fix this raised and dropped the request; it must not raise now - if sync_mode: - for chunk in response: - if getattr(chunk, "usage", None) is not None: - seen_usage.append(chunk.usage) - else: - async for chunk in response: - if getattr(chunk, "usage", None) is not None: - seen_usage.append(chunk.usage) - - assert any( - u.total_tokens == final_usage_block.total_tokens for u in seen_usage - ), "usage recovered from raw chunks was not emitted after stream_chunk_builder raised" - - -@pytest.mark.parametrize("sync_mode", [True, False]) -@pytest.mark.asyncio -async def test_stream_chunk_builder_raise_and_usage_recovery_failure_does_not_crash( - sync_mode, -): - """If end-of-stream assembly raises AND best-effort usage recovery from the raw - chunks also fails, the stream must still complete cleanly rather than propagate - the exception to the consumer.""" - from litellm.litellm_core_utils import streaming_handler as sh_module - - final_chunk = ModelResponseStream( - id="chatcmpl-raise-recover-fail", - created=1742056047, - model=None, - object="chat.completion.chunk", - choices=[ - StreamingChoices( - finish_reason="stop", - index=0, - delta=Delta(content="", role="assistant"), - ) - ], - usage=Usage(completion_tokens=1, prompt_tokens=1, total_tokens=2), - ) - - response = CustomStreamWrapper( - completion_stream=ModelResponseListIterator( - model_responses=bedrock_chunks + [final_chunk] - ), - model="bedrock/claude-haiku-4-5-20251001-v1:0", - custom_llm_provider="bedrock", - logging_obj=Logging( - model="bedrock/claude-haiku-4-5-20251001-v1:0", - messages=[{"role": "user", "content": "Hey"}], - stream=True, - call_type="completion", - start_time=time.time(), - litellm_call_id="raise-recover-fail", - function_id="1245", - ), - stream_options={"include_usage": True}, - ) - - with ( - patch.object( - litellm, "stream_chunk_builder", side_effect=Exception("assembly failed") - ), - patch.object( - sh_module, "calculate_total_usage", side_effect=Exception("recovery failed") - ), - ): - # must not raise even though both assembly and recovery fail - if sync_mode: - chunks = [c for c in response] - else: - chunks = [c async for c in response] - - assert len(chunks) > 0 - - -class TransportErrorAfterChunksIterator: - """Yields the given chunks, then raises the given exception once, then StopAsyncIteration.""" - - def __init__(self, model_responses, exception): - self.model_responses = model_responses - self.exception = exception - self.index = 0 - self.raised = False - - def __aiter__(self): - return self - - async def __anext__(self): - if self.index < len(self.model_responses): - chunk = self.model_responses[self.index] - self.index += 1 - return chunk - if not self.raised: - self.raised = True - raise self.exception - raise StopAsyncIteration - - -def _reset_test_chunk(content: Optional[str] = None, finish_reason: Optional[str] = None) -> ModelResponseStream: - return ModelResponseStream( - id="chatcmpl-reset-test", - created=1783458104, - model="stub-model", - object="chat.completion.chunk", - choices=[ - StreamingChoices( - index=0, - delta=Delta(content=content), - finish_reason=finish_reason, - ) - ], - ) - - -@pytest.mark.asyncio -async def test_transport_read_error_after_finish_reason_ends_stream_gracefully( - logging_obj: Logging, -): - """A trailing connection reset after the provider's finish chunk must not fail the stream.""" - import httpx - - completion_stream = TransportErrorAfterChunksIterator( - model_responses=[ - _reset_test_chunk(content="Hello"), - _reset_test_chunk(finish_reason="stop"), - ], - exception=httpx.ReadError("Response payload is not completed"), - ) - response = CustomStreamWrapper( - completion_stream=completion_stream, - model="hosted_vllm/stub-model", - custom_llm_provider="hosted_vllm", - logging_obj=logging_obj, - ) - - chunks = [chunk async for chunk in response] - - finish_reasons = [ - chunk.choices[0].finish_reason - for chunk in chunks - if chunk.choices and chunk.choices[0].finish_reason - ] - contents = [ - chunk.choices[0].delta.content - for chunk in chunks - if chunk.choices and chunk.choices[0].delta and chunk.choices[0].delta.content - ] - assert finish_reasons == ["stop"] - assert contents == ["Hello"] - - -@pytest.mark.asyncio -async def test_transport_read_error_before_finish_reason_raises(logging_obj: Logging): - """A connection reset before any finish chunk must surface, never end as a clean stop. - - Regression test for silent empty/truncated HTTP 200 streams: the aiohttp - transport used to swallow mid-stream connection resets, so the wrapper saw a - clean end-of-stream and fabricated finish_reason "stop". - """ - import httpx - - from litellm.exceptions import MidStreamFallbackError - - completion_stream = TransportErrorAfterChunksIterator( - model_responses=[_reset_test_chunk(content="Hel")], - exception=httpx.ReadError("Response payload is not completed"), - ) - response = CustomStreamWrapper( - completion_stream=completion_stream, - model="hosted_vllm/stub-model", - custom_llm_provider="hosted_vllm", - logging_obj=logging_obj, - ) - - received = [] - async def _drain(): - async for chunk in response: - received.append(chunk) - - with pytest.raises(MidStreamFallbackError): - await _drain() - - fabricated_finish_reasons = [ - chunk.choices[0].finish_reason - for chunk in received - if chunk.choices and chunk.choices[0].finish_reason - ] - assert fabricated_finish_reasons == [] - - -def test_openai_custom_tool_call_stream_deltas_survive_conversion(logging_obj: Logging): - """ - Regression test: OpenAI chat completions custom tool calls stream as - delta.tool_calls entries with a `custom` payload and NO `function` key. - Delta() used to raise on those dicts and chunk_creator's except branch - replaced the choice with an empty Delta, silently dropping the entire - tool call from the client stream. - """ - from openai.types.chat.chat_completion_chunk import ChatCompletionChunk - - from litellm.types.utils import ChatCompletionDeltaCustomToolCall - - raw_chunks = [ - { - "id": "chatcmpl-custom", - "object": "chat.completion.chunk", - "created": 1784657671, - "model": "gpt-5.6", - "choices": [ - { - "index": 0, - "delta": { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "index": 0, - "id": "call_TBs", - "type": "custom", - "custom": {"name": "ApplyPatch", "input": ""}, - } - ], - }, - "finish_reason": None, - } - ], - }, - { - "id": "chatcmpl-custom", - "object": "chat.completion.chunk", - "created": 1784657671, - "model": "gpt-5.6", - "choices": [ - { - "index": 0, - "delta": {"tool_calls": [{"index": 0, "custom": {"input": "*** Begin Patch\n"}}]}, - "finish_reason": None, - } - ], - }, - { - "id": "chatcmpl-custom", - "object": "chat.completion.chunk", - "created": 1784657671, - "model": "gpt-5.6", - "choices": [ - { - "index": 0, - "delta": {"tool_calls": [{"index": 0, "custom": {"input": "*** End Patch\n"}}]}, - "finish_reason": None, - } - ], - }, - { - "id": "chatcmpl-custom", - "object": "chat.completion.chunk", - "created": 1784657671, - "model": "gpt-5.6", - "choices": [{"index": 0, "delta": {}, "finish_reason": "tool_calls"}], - }, - ] - sdk_chunks = [ChatCompletionChunk.construct(**raw) for raw in raw_chunks] - first_dumped = sdk_chunks[0].choices[0].model_dump() - assert first_dumped["delta"]["tool_calls"][0]["custom"] == {"name": "ApplyPatch", "input": ""} - - wrapper = CustomStreamWrapper( - completion_stream=iter(sdk_chunks), - model="gpt-5.6", - custom_llm_provider="openai", - logging_obj=logging_obj, - ) - - emitted = list(wrapper) - tool_call_deltas = [ - chunk.choices[0].delta.tool_calls[0] - for chunk in emitted - if chunk.choices and chunk.choices[0].delta and chunk.choices[0].delta.tool_calls - ] - assert len(tool_call_deltas) == 3 - assert isinstance(tool_call_deltas[0], ChatCompletionDeltaCustomToolCall) - assert tool_call_deltas[0].id == "call_TBs" - assert tool_call_deltas[0].type == "custom" - assert tool_call_deltas[0].custom.name == "ApplyPatch" - combined_input = "".join(tc.custom.input or "" for tc in tool_call_deltas) - assert combined_input == "*** Begin Patch\n*** End Patch\n" - finish_reasons = [chunk.choices[0].finish_reason for chunk in emitted if chunk.choices] - assert "tool_calls" in finish_reasons - - -def test_sync_completion_never_stamps_correlation_context(monkeypatch): - """wrapper() (the sync entry point) does not participate in - request_correlation_in_logs at all: Logging.__init__() is called with - supports_correlation_logging=False for every sync call, so - trace_id_var/session_id_var are never touched, regardless of whether the - caller passes litellm_trace_id/litellm_session_id or the call streams. - - This is a deliberate scoping decision, not an oversight: a plain OS - thread has no per-call isolation the way an asyncio Task does, and a - thread pool's worker threads are recycled across unrelated requests, so - safely supporting this for the sync path needs its own restore mechanism - with its own tests - tracked as a separate, follow-up piece of work. - Async (acompletion/wrapper_async, the only path the proxy uses) is - unaffected - see test_async_streaming_completion_does_not_reset_context_before_iteration.""" - monkeypatch.setattr(litellm, "request_correlation_in_logs", True) - # Reset explicitly rather than asserting a clean slate - this must hold - # regardless of what any other test left behind in these module-level - # contextvars. - trace_id_var.set("") - session_id_var.set("") - try: - litellm.completion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hi"}], - mock_response="Hello there!", - litellm_trace_id="should-never-appear", - litellm_session_id="should-never-appear-either", - num_retries=0, - ) - assert trace_id_var.get() == "" - assert session_id_var.get() == "" - - response = litellm.completion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hi"}], - mock_response="Hello there!", - stream=True, - litellm_trace_id="should-never-appear-stream", - litellm_session_id="should-never-appear-stream-either", - num_retries=0, - ) - for _ in response: - pass - assert trace_id_var.get() == "" - assert session_id_var.get() == "" - finally: - trace_id_var.set("") - session_id_var.set("") - - -def test_abandoned_sync_stream_cannot_contaminate_a_later_call_on_the_same_thread(monkeypatch): - """The maintainer-reported blocking bug reproduced live in this session - - request A starts a sync stream, consumes one chunk, abandons it; request - B runs next on the same forced-reuse ThreadPoolExecutor worker - is now - structurally impossible rather than merely restored-after-the-fact: since - sync calls never stamp trace_id_var/session_id_var at all - (supports_correlation_logging=False), there is nothing for request A to - leave behind for request B to inherit.""" - monkeypatch.setattr(litellm, "request_correlation_in_logs", True) - - from concurrent.futures import ThreadPoolExecutor - - pool = ThreadPoolExecutor(max_workers=1) - try: - - def call_a_abandon_stream(): - response = litellm.completion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "call A"}], - mock_response="call A response", - stream=True, - litellm_session_id="SESSION-AAA", - litellm_trace_id="TRACE-AAA", - num_retries=0, - ) - next(response) # consume exactly one chunk, then abandon it - - def call_b_non_streaming(): - litellm.completion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "call B"}], - mock_response="call B response", - litellm_session_id="SESSION-BBB", - litellm_trace_id="TRACE-BBB", - num_retries=0, - ) - return trace_id_var.get(), session_id_var.get() - - pool.submit(call_a_abandon_stream).result() - ids_after_b = pool.submit(call_b_non_streaming).result() - - assert ids_after_b == ("", "") - finally: - pool.shutdown(wait=True) - - -@pytest.mark.asyncio -async def test_async_streaming_completion_does_not_reset_context_before_iteration(monkeypatch): - """Same as above for wrapper_async()/acompletion().""" - monkeypatch.setattr(litellm, "request_correlation_in_logs", True) - trace_id_var.set("outer-trace-async-stream") - session_id_var.set("outer-session-async-stream") - try: - response = await litellm.acompletion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hi"}], - mock_response="Hello there!", - stream=True, - litellm_session_id="async-streaming-call-session", - num_retries=0, - ) - assert session_id_var.get() == "async-streaming-call-session" - - async for _ in response: - pass - - # Once the stream is genuinely exhausted, the *consuming* task's own - # context must be restored - async_success_handler's own dispatch (via - # asyncio.create_task) only fixes up its own detached task, not this one. - assert session_id_var.get() == "outer-session-async-stream" - assert trace_id_var.get() == "outer-trace-async-stream" - finally: - trace_id_var.set("") - session_id_var.set("") - - -def test_stream_wrapper_del_restores_correlation_context(): - """CustomStreamWrapper.__del__ is the best-effort fallback for an abandoned - stream (caller never exhausts it, so the normal terminal-handler restore - never fires). Testing this via real garbage collection is unreliable in - practice - CPython's per-chunk logging submits work to a thread pool - executor whose worker thread transiently holds its own reference to the - wrapper (a bound method argument) until that task completes, so refcount - doesn't reliably hit zero on a deterministic schedule even with polling. - Call __del__ directly instead: it's a plain method, calling it early - doesn't run actual finalization, and this exercises exactly the logic that - real garbage collection would eventually trigger. - """ - trace_id_var.set("outer-trace-abandoned") - session_id_var.set("outer-session-abandoned") - try: - log_obj = Logging( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hi"}], - stream=True, - call_type="completion", - start_time=None, - litellm_call_id="abandoned-stream-call", - function_id="fn-abandoned-stream", - kwargs={"litellm_session_id": "abandoned-stream-session"}, - ) - wrapper = CustomStreamWrapper( - completion_stream=iter([]), - model="gpt-3.5-turbo", - logging_obj=log_obj, - ) - wrapper.__del__() - - assert trace_id_var.get() == "outer-trace-abandoned" - assert session_id_var.get() == "outer-session-abandoned" - finally: - trace_id_var.set("") - session_id_var.set("") - - -def test_stream_wrapper_del_never_raises_with_broken_logging_obj(): - """__del__ runs during garbage collection, possibly at interpreter - shutdown - it must never raise regardless of what's wrong with logging_obj, - or Python prints an ignored "exception in __del__" warning and, worse, - could mask the real error a caller is in the middle of handling.""" - - class ExplodingLogging: - model_call_details: dict = {} - - def _restore_correlation_context(self): - raise RuntimeError("logging_obj is in a bad state") - - wrapper = CustomStreamWrapper( - completion_stream=iter([]), - model="gpt-3.5-turbo", - logging_obj=ExplodingLogging(), - ) - wrapper.__del__() # must not raise - - -def test_stream_wrapper_del_does_not_clobber_a_newer_active_call(): - """A delayed finalizer must never stomp a different, still-active call's - context. If an abandoned stream's __del__ fires late - after a new call - has already started in the same Task/thread and claimed the contextvars - - unconditionally restoring the abandoned stream's own pre-call snapshot - would corrupt the active call's subsequent log lines with stale ids.""" - trace_id_var.set("outer-trace-before-abandoned-call") - session_id_var.set("outer-session-before-abandoned-call") - try: - abandoned_log_obj = Logging( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hi"}], - stream=True, - call_type="completion", - start_time=None, - litellm_call_id="abandoned-stream-call", - function_id="fn-abandoned-stream", - kwargs={"litellm_session_id": "abandoned-stream-session"}, - ) - wrapper = CustomStreamWrapper( - completion_stream=iter([]), - model="gpt-3.5-turbo", - logging_obj=abandoned_log_obj, - ) - - # A new, unrelated call starts in this same Task/thread before the - # abandoned stream's __del__ ever fires, and claims the contextvars. - Logging( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hi"}], - stream=False, - call_type="completion", - start_time=None, - litellm_call_id="newer-active-call", - function_id="fn-newer-active-call", - kwargs={"litellm_session_id": "newer-active-session"}, - ) - assert trace_id_var.get() != abandoned_log_obj.litellm_trace_id - assert session_id_var.get() == "newer-active-session" - - # The delayed finalizer for the abandoned stream must not clobber - # the newer call's still-active ids. - wrapper.__del__() - - assert trace_id_var.get() != abandoned_log_obj.litellm_trace_id - assert session_id_var.get() == "newer-active-session" - finally: - trace_id_var.set("") - session_id_var.set("") - - -def test_stream_wrapper_del_restores_when_own_session_id_needed_sanitizing(): - """The __del__ guard must compare against the *sanitized* id actually - stored in the contextvar, not the raw litellm_session_id/litellm_trace_id - - set_session_id()/set_trace_id() strip control characters before - storing, so a caller-supplied id containing e.g. a newline would never - equal the raw attribute, and the guard would wrongly conclude some other - call has claimed the context and skip cleanup forever.""" - trace_id_var.set("outer-trace-needs-sanitizing") - session_id_var.set("outer-session-needs-sanitizing") - try: - raw_session_id = "abandoned\nsession\rwith-control-chars" - log_obj = Logging( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hi"}], - stream=True, - call_type="completion", - start_time=None, - litellm_call_id="abandoned-stream-needs-sanitizing", - function_id="fn-abandoned-stream-needs-sanitizing", - kwargs={"litellm_session_id": raw_session_id}, - ) - # Sanity: the contextvar holds the sanitized value, which differs - # from the raw litellm_session_id this test constructed it with. - assert session_id_var.get() != raw_session_id - assert log_obj.litellm_session_id == raw_session_id - - wrapper = CustomStreamWrapper( - completion_stream=iter([]), - model="gpt-3.5-turbo", - logging_obj=log_obj, - ) - wrapper.__del__() - - assert trace_id_var.get() == "outer-trace-needs-sanitizing" - assert session_id_var.get() == "outer-session-needs-sanitizing" - finally: - trace_id_var.set("") - session_id_var.set("") - - -def test_stream_wrapper_next_keeps_context_active_through_synthesized_finish_reason_chunk(): - """When the provider already supplied a finish_reason (e.g. stripped from a - content chunk) and the underlying stream then ends, __next__ synthesizes the - terminal chunk via finish_reason_handler() and returns it. That chunk is still - this call's own data - the caller's own (application-level) log statements - processing it run immediately after this return, in the same synchronous - frame, so context must NOT be restored yet or those log lines would carry the - wrong ids. A caller that keeps iterating (the common, non-early-break pattern) - still gets a correct, deterministic restore on the very next __next__() call, - since completion_stream is already exhausted and immediately re-raises - StopIteration.""" - trace_id_var.set("outer-trace-finish-reason") - session_id_var.set("outer-session-finish-reason") - try: - log_obj = Logging( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hi"}], - stream=True, - call_type="completion", - start_time=None, - litellm_call_id="finish-reason-call", - function_id="fn-finish-reason", - kwargs={"litellm_session_id": "finish-reason-session"}, - ) - wrapper = CustomStreamWrapper( - completion_stream=iter([]), - model="gpt-3.5-turbo", - logging_obj=log_obj, - ) - wrapper.received_finish_reason = "stop" - assert trace_id_var.get() == log_obj.litellm_trace_id - assert session_id_var.get() == "finish-reason-session" - - chunk = next(wrapper) - - assert chunk.choices[0].finish_reason is not None - # Still this call's own ids - not restored yet. - assert trace_id_var.get() == log_obj.litellm_trace_id - assert session_id_var.get() == "finish-reason-session" - - # A caller that keeps iterating (doesn't break early) still gets a - # deterministic restore right here, on the next real StopIteration. - with pytest.raises(StopIteration): - next(wrapper) - assert trace_id_var.get() == "outer-trace-finish-reason" - assert session_id_var.get() == "outer-session-finish-reason" - finally: - trace_id_var.set("") - session_id_var.set("") - - -def test_stream_wrapper_del_cleans_up_after_synthesized_finish_reason_chunk(): - """A caller that breaks immediately after seeing finish_reason (the - early-break pattern) never triggers the next()-driven restore above - it - relies on the best-effort __del__ guard instead, same as any other - abandoned stream. The guard must still recognize this call's own - (unrestored) ids as unclaimed and clean them up.""" - trace_id_var.set("outer-trace-finish-reason-del") - session_id_var.set("outer-session-finish-reason-del") - try: - log_obj = Logging( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hi"}], - stream=True, - call_type="completion", - start_time=None, - litellm_call_id="finish-reason-del-call", - function_id="fn-finish-reason-del", - kwargs={"litellm_session_id": "finish-reason-del-session"}, - ) - wrapper = CustomStreamWrapper( - completion_stream=iter([]), - model="gpt-3.5-turbo", - logging_obj=log_obj, - ) - wrapper.received_finish_reason = "stop" - - chunk = next(wrapper) - assert chunk.choices[0].finish_reason is not None - - wrapper.__del__() - - assert trace_id_var.get() == "outer-trace-finish-reason-del" - assert session_id_var.get() == "outer-session-finish-reason-del" - finally: - trace_id_var.set("") - session_id_var.set("") - - -@pytest.mark.asyncio -async def test_stream_wrapper_anext_keeps_context_active_through_synthesized_finish_reason_chunk(): - """Async sibling of test_stream_wrapper_next_keeps_context_active_through_synthesized_finish_reason_chunk - - _finalize_completed_stream()'s else branch must not restore before - returning the synthesized chunk either.""" - trace_id_var.set("outer-trace-anext-finish-reason") - session_id_var.set("outer-session-anext-finish-reason") - try: - log_obj = Logging( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hi"}], - stream=True, - call_type="completion", - start_time=None, - litellm_call_id="anext-finish-reason-call", - function_id="fn-anext-finish-reason", - kwargs={"litellm_session_id": "anext-finish-reason-session"}, - ) - - async def _empty_aiter(): - return - yield # pragma: no cover - makes this an async generator - - wrapper = CustomStreamWrapper( - completion_stream=_empty_aiter(), - model="gpt-3.5-turbo", - logging_obj=log_obj, - ) - wrapper.received_finish_reason = "stop" - assert trace_id_var.get() == log_obj.litellm_trace_id - assert session_id_var.get() == "anext-finish-reason-session" - - chunk = await wrapper.__anext__() - - assert chunk.choices[0].finish_reason is not None - # Still this call's own ids - not restored yet. - assert trace_id_var.get() == log_obj.litellm_trace_id - assert session_id_var.get() == "anext-finish-reason-session" - - # A caller that keeps iterating still gets a deterministic restore - # right here, on the next real StopAsyncIteration. - with pytest.raises(StopAsyncIteration): - await wrapper.__anext__() - assert trace_id_var.get() == "outer-trace-anext-finish-reason" - assert session_id_var.get() == "outer-session-anext-finish-reason" - finally: - trace_id_var.set("") - session_id_var.set("") - - -@pytest.mark.asyncio -async def test_stream_wrapper_anext_max_duration_timeout_restores_consumer_correlation_context(monkeypatch): - """_check_max_streaming_duration() raises litellm.Timeout when a client keeps - an async stream open past LITELLM_MAX_STREAMING_DURATION_SECONDS. That raise - must flow through the same except Exception -> _handle_stream_fallback_error - path as every other failure so the consumer's outer correlation context gets - restored - calling the check before entering __anext__()'s try block would - let the Timeout bypass that restoration entirely.""" - monkeypatch.setattr(litellm.constants, "LITELLM_MAX_STREAMING_DURATION_SECONDS", 1) - trace_id_var.set("outer-trace-max-duration") - session_id_var.set("outer-session-max-duration") - try: - log_obj = Logging( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hi"}], - stream=True, - call_type="completion", - start_time=None, - litellm_call_id="max-duration-call", - function_id="fn-max-duration", - kwargs={"litellm_session_id": "max-duration-session"}, - ) - - async def _empty_aiter(): - return - yield # pragma: no cover - makes this an async generator - - wrapper = CustomStreamWrapper( - completion_stream=_empty_aiter(), - model="gpt-3.5-turbo", - logging_obj=log_obj, - ) - assert trace_id_var.get() == log_obj.litellm_trace_id - assert session_id_var.get() == "max-duration-session" - - wrapper._stream_created_time = time.time() - 10 - - with pytest.raises(litellm.Timeout): - await wrapper.__anext__() - - assert trace_id_var.get() == "outer-trace-max-duration" - assert session_id_var.get() == "outer-session-max-duration" - finally: - trace_id_var.set("") - session_id_var.set("") - - -@pytest.mark.asyncio -async def test_stream_wrapper_aclose_restores_consumer_correlation_context(): - """Explicit early termination (aclose(), e.g. on client disconnect or a - router fallback aborting an in-progress stream) must restore the caller's - correlation context too - not just __del__'s best-effort GC-timed fallback, - since aclose() is normally called deterministically by the consumer/ - framework, unlike __del__.""" - trace_id_var.set("outer-trace-aclose") - session_id_var.set("outer-session-aclose") - try: - log_obj = Logging( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hi"}], - stream=True, - call_type="completion", - start_time=None, - litellm_call_id="aclose-call", - function_id="fn-aclose", - kwargs={"litellm_session_id": "aclose-session"}, - ) - - async def _empty_aiter(): - return - yield # pragma: no cover - makes this an async generator - - wrapper = CustomStreamWrapper( - completion_stream=_empty_aiter(), - model="gpt-3.5-turbo", - logging_obj=log_obj, - ) - assert trace_id_var.get() == log_obj.litellm_trace_id - assert session_id_var.get() == "aclose-session" - - await wrapper.aclose() - - assert trace_id_var.get() == "outer-trace-aclose" - assert session_id_var.get() == "outer-session-aclose" - finally: - trace_id_var.set("") - session_id_var.set("") - - -@pytest.mark.asyncio -async def test_stream_wrapper_aclose_keeps_context_active_through_close_failure_diagnostic(monkeypatch): - """If closing the underlying provider stream raises, aclose()'s except - branch logs a debug diagnostic. That log line must still carry the - closing stream's own trace_id/session_id - the outer context must not be - restored until after the close attempt (and its diagnostic) completes.""" - trace_id_var.set("outer-trace-close-fail") - session_id_var.set("outer-session-close-fail") - try: - log_obj = Logging( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hi"}], - stream=True, - call_type="completion", - start_time=None, - litellm_call_id="close-fail-call", - function_id="fn-close-fail", - kwargs={"litellm_session_id": "close-fail-session"}, - ) - - class _RaisingAsyncCloseStream: - async def aclose(self): - raise RuntimeError("boom closing stream") - - def __aiter__(self): - return self - - async def __anext__(self): - raise StopAsyncIteration - - wrapper = CustomStreamWrapper( - completion_stream=_RaisingAsyncCloseStream(), - model="gpt-3.5-turbo", - logging_obj=log_obj, - ) - assert trace_id_var.get() == log_obj.litellm_trace_id - assert session_id_var.get() == "close-fail-session" - - captured_ids = {} - real_debug = verbose_logger.debug - - def fake_debug(msg, *args, **kwargs): - if "error closing completion_stream" in msg: - captured_ids["trace_id"] = trace_id_var.get() - captured_ids["session_id"] = session_id_var.get() - return real_debug(msg, *args, **kwargs) - - monkeypatch.setattr(verbose_logger, "debug", fake_debug) - - await wrapper.aclose() - - assert captured_ids["trace_id"] == log_obj.litellm_trace_id - assert captured_ids["session_id"] == "close-fail-session" - assert trace_id_var.get() == "outer-trace-close-fail" - assert session_id_var.get() == "outer-session-close-fail" - finally: - trace_id_var.set("") - session_id_var.set("") - - -def test_handle_stream_fallback_error_restores_context_only_after_exception_mapping(monkeypatch): - """_map_anthropic_exception/_map_aleph_alpha_exception synchronously log a - debug diagnostic (the raw status code) as part of exception_type()'s - mapping. The consumer's outer context must not be restored until that - mapping call returns, or the diagnostic log line would carry the outer - (or empty) trace_id/session_id instead of the failing stream's own.""" - trace_id_var.set("outer-trace-fallback") - session_id_var.set("outer-session-fallback") - try: - log_obj = Logging( - model="claude-3-opus", - messages=[{"role": "user", "content": "hi"}], - stream=True, - call_type="completion", - start_time=None, - litellm_call_id="fallback-error-call", - function_id="fn-fallback-error", - kwargs={"litellm_session_id": "fallback-error-session"}, - ) - wrapper = CustomStreamWrapper( - completion_stream=iter([]), - model="claude-3-opus", - custom_llm_provider="anthropic", - logging_obj=log_obj, - ) - - captured_ids = {} - - def fake_exception_type(**kwargs): - captured_ids["trace_id"] = trace_id_var.get() - captured_ids["session_id"] = session_id_var.get() - return ValueError("mapped boom") - - monkeypatch.setattr("litellm.litellm_core_utils.streaming_handler.exception_type", fake_exception_type) - - from litellm.exceptions import MidStreamFallbackError - - with pytest.raises(MidStreamFallbackError): - wrapper._handle_stream_fallback_error(RuntimeError("boom")) - - # The mapper ran while the stream's own ids were still active. - assert captured_ids["trace_id"] == log_obj.litellm_trace_id - assert captured_ids["session_id"] == "fallback-error-session" - # Restored to the consumer's outer context once mapping/raise completes. - assert trace_id_var.get() == "outer-trace-fallback" - assert session_id_var.get() == "outer-session-fallback" - finally: - trace_id_var.set("") - session_id_var.set("") - - -def test_chunk_creator_preserves_hidden_provider_specific_fields_from_parsed_chunk(): - wrapper = CustomStreamWrapper( - completion_stream=None, - model="gemini-3.5-flash", - logging_obj=MagicMock(), - custom_llm_provider="vertex_ai", - ) - parsed_chunk = ModelResponseStream( - choices=[StreamingChoices(index=0, delta=Delta(content="hello", role="assistant"), finish_reason=None)], - ) - parsed_chunk._hidden_params["provider_specific_fields"] = {"traffic_type": "ON_DEMAND_FLEX"} - - result = wrapper.chunk_creator(chunk=parsed_chunk) - - assert result is not None - assert result._hidden_params["provider_specific_fields"] == {"traffic_type": "ON_DEMAND_FLEX"} - assembled = litellm.stream_chunk_builder(chunks=[result]) - assert assembled is not None - assert assembled._hidden_params["provider_specific_fields"] == {"traffic_type": "ON_DEMAND_FLEX"} - - -def test_chunk_creator_keeps_provider_model_private_across_stream(): - from litellm.router_utils.add_retry_fallback_headers import ( - get_hidden_params_dict, - ) - - wrapper = CustomStreamWrapper( - completion_stream=None, - model="requested-route", - logging_obj=MagicMock(), - custom_llm_provider="openai", - ) - selected_chunk = ModelResponseStream( - id="chunk-1", - model="selected-model", - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta(content="hello"), - ) - ], - ) - terminal_chunk = ModelResponseStream( - id="chunk-1", - model=None, - choices=[ - StreamingChoices( - finish_reason="stop", - index=0, - delta=Delta(), - ) - ], - ) - - first_result = wrapper.chunk_creator(chunk=selected_chunk) - terminal_result = wrapper.chunk_creator(chunk=terminal_chunk) - - assert first_result is not None - assert terminal_result is not None - assert first_result.model == "requested-route" - assert terminal_result.model == "requested-route" - assert ( - get_hidden_params_dict(first_result)["provider_response_model"] - == "selected-model" - ) - assert ( - get_hidden_params_dict(terminal_result)["provider_response_model"] - == "selected-model" - ) - - assembled = litellm.stream_chunk_builder(chunks=[first_result, terminal_result]) - assert assembled is not None - assert assembled.model == "requested-route" - assert ( - get_hidden_params_dict(assembled)["provider_response_model"] - == "selected-model" - ) - - -def test_assembled_stream_uses_later_provider_model_for_cost( - monkeypatch: pytest.MonkeyPatch, -): - from litellm.router_utils.add_retry_fallback_headers import ( - get_hidden_params_dict, - ) - - selected_model_info = { - "input_cost_per_token": 0.000002, - "output_cost_per_token": 0.000004, - "litellm_provider": "azure", - } - monkeypatch.setitem( - litellm.model_cost, - "azure/gpt-4.1-nano-2025-04-14", - selected_model_info, - ) - monkeypatch.setitem( - litellm.model_cost, - "azure/azure-model-router", - { - "input_cost_per_token": 0.00002, - "output_cost_per_token": 0.00004, - "litellm_provider": "azure", - }, - ) - logging_obj = MagicMock() - logging_obj.model_call_details = {"custom_llm_provider": "azure"} - wrapper = CustomStreamWrapper( - completion_stream=None, - model="azure-model-router", - logging_obj=logging_obj, - custom_llm_provider="azure", - ) - router_chunk = ModelResponseStream( - id="chunk-1", - model="azure-model-router", - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta(content="hello "), - ) - ], - ) - selected_chunk = ModelResponseStream( - id="chunk-1", - model="gpt-4.1-nano-2025-04-14", - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta(content="world"), - ) - ], - ) - terminal_chunk = ModelResponseStream( - id="chunk-1", - model="azure-model-router", - choices=[ - StreamingChoices( - finish_reason="stop", - index=0, - delta=Delta(), - ) - ], - ) - - router_result = wrapper.chunk_creator(chunk=router_chunk) - selected_result = wrapper.chunk_creator(chunk=selected_chunk) - terminal_result = wrapper.chunk_creator(chunk=terminal_chunk) - - assert router_result is not None - assert selected_result is not None - assert terminal_result is not None - assert ( - get_hidden_params_dict(router_result)["provider_response_model"] - == "azure-model-router" - ) - assert ( - get_hidden_params_dict(selected_result)["provider_response_model"] - == "gpt-4.1-nano-2025-04-14" - ) - assert ( - get_hidden_params_dict(terminal_result)["provider_response_model"] - == "azure-model-router" - ) - - assembled = litellm.stream_chunk_builder( - chunks=[router_result, selected_result, terminal_result] - ) - assert assembled is not None - assert assembled.model == "gpt-4.1-nano-2025-04-14" - assert ( - get_hidden_params_dict(assembled)["provider_response_model"] - == "gpt-4.1-nano-2025-04-14" - ) - assembled.usage = Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15) - assert litellm.completion_cost( - completion_response=assembled, - custom_llm_provider="azure", - ) == pytest.approx( - 10 * selected_model_info["input_cost_per_token"] - + 5 * selected_model_info["output_cost_per_token"] - ) - - -@pytest.mark.asyncio -async def test_async_stream_assembled_response_keeps_vertex_traffic_type(logging_obj: Logging): - content_chunk = ModelResponseStream( - choices=[StreamingChoices(index=0, delta=Delta(content="hello", role="assistant"), finish_reason=None)], - ) - final_chunk = ModelResponseStream( - choices=[StreamingChoices(index=0, delta=Delta(content=""), finish_reason="stop")], - ) - setattr(final_chunk, "usage", Usage(prompt_tokens=7, completion_tokens=5, total_tokens=12)) - final_chunk._hidden_params["provider_specific_fields"] = {"traffic_type": "ON_DEMAND_FLEX"} - - async def _stream(): - yield content_chunk - yield final_chunk - - wrapper = CustomStreamWrapper( - completion_stream=_stream(), - model="gemini-3.5-flash", - logging_obj=logging_obj, - custom_llm_provider="vertex_ai", - stream_options={"include_usage": True}, - ) - - received = [chunk async for chunk in wrapper] - - assembled = litellm.stream_chunk_builder(chunks=received, messages=[{"role": "user", "content": "hi"}]) - assert assembled is not None - assert assembled._hidden_params["provider_specific_fields"]["traffic_type"] == "ON_DEMAND_FLEX" - - -@pytest.mark.asyncio -async def test_async_fake_stream_final_chunk_carries_hidden_usage(logging_obj: Logging): - from litellm.llms.base_llm.base_model_iterator import MockResponseIterator - from litellm.types.utils import ModelResponse - - model_response = ModelResponse( - id="chatcmpl-fake-stream", - model="my-random-model", - choices=[ - { - "index": 0, - "message": {"role": "assistant", "content": "hello world"}, - "finish_reason": "stop", - } - ], - ) - model_response.usage = Usage(prompt_tokens=1234, completion_tokens=7, total_tokens=1241) - - wrapper = CustomStreamWrapper( - completion_stream=MockResponseIterator(model_response=model_response), - model="my-random-model", - custom_llm_provider="anthropic", - logging_obj=logging_obj, - ) - - final_chunk = None - async for chunk in wrapper: - final_chunk = chunk - - assert final_chunk is not None - hidden_usage = final_chunk._hidden_params.get("usage") - assert hidden_usage is not None - assert hidden_usage.prompt_tokens == 1234 - assert hidden_usage.completion_tokens == 7 - assert hidden_usage.total_tokens == 1241 - - -class TestStableStreamingResponseId: - """ - All chunks of one streamed response must share the same top-level id - (OpenAI streaming contract). Providers streaming via GenericStreamingChunk - (e.g. GigaChat) do not propagate an upstream response id, so - CustomStreamWrapper must pin the id from the first chunk it creates, - mirroring the existing `created` pinning (issue #11437). - - Clients such as goose merge streamed deltas into one assistant message by - chunk id; per-chunk ids split a single reply into many messages. - """ - - def test_generic_chunks_share_one_id(self): - def _generic_chunks(): - return iter( - [ - { - "text": "Hello", - "tool_use": None, - "is_finished": False, - "finish_reason": "", - "usage": None, - "index": 0, - }, - { - "text": " world", - "tool_use": None, - "is_finished": False, - "finish_reason": "", - "usage": None, - "index": 0, - }, - { - "text": "", - "tool_use": None, - "is_finished": True, - "finish_reason": "stop", - "usage": { - "prompt_tokens": 1, - "completion_tokens": 2, - "total_tokens": 3, - }, - "index": 0, - }, - ] - ) - - wrapper = CustomStreamWrapper( - completion_stream=_generic_chunks(), - model="gigachat/GigaChat-2-Max", - logging_obj=MagicMock(), - custom_llm_provider="gigachat", - ) - ids = [chunk.id for chunk in wrapper if chunk.id] - assert ids, "no chunks emitted" - assert len(set(ids)) == 1, f"chunk ids differ across one stream: {ids}" - - def test_creator_pins_id_from_first_chunk(self): - wrapper = CustomStreamWrapper( - completion_stream=iter([]), - model="gigachat/GigaChat-2-Max", - logging_obj=MagicMock(), - custom_llm_provider="gigachat", - ) - first = wrapper.model_response_creator() - assert wrapper.response_id == first.id - assert wrapper.model_response_creator().id == first.id - - def test_provider_supplied_id_still_wins(self): - wrapper = CustomStreamWrapper( - completion_stream=iter([]), - model="gigachat/GigaChat-2-Max", - logging_obj=MagicMock(), - custom_llm_provider="gigachat", - ) - wrapper.response_id = "chatcmpl-from-provider" - assert wrapper.model_response_creator().id == "chatcmpl-from-provider" - -@pytest.mark.asyncio -async def test_clean_eof_without_finish_reason_raises_midstream_error_issue_40260(): - from litellm.exceptions import MidStreamFallbackError - - async def source(): - yield ModelResponseStream( - id="chatcmpl-offline-repro", - created=0, - model="gpt-5.6", - choices=[ - StreamingChoices( - index=0, - delta=Delta(content='{"findings":[{"title":"unfinished'), - finish_reason=None, - ) - ], - ) - - log = Logging( - model="gpt-5.6", - messages=[{"role": "user", "content": "Return JSON"}], - stream=True, - call_type="acompletion", - start_time=time.time(), - litellm_call_id="offline-eof-repro", - function_id="offline-eof-repro", - ) - wrapper = CustomStreamWrapper( - completion_stream=source(), - model="gpt-5.6", - custom_llm_provider="azure", - logging_obj=log, - stream_options={"include_usage": True}, - ) - - chunks = [] - - async def _consume(): - async for c in wrapper: - chunks.append(c) - - with pytest.raises(MidStreamFallbackError) as excinfo: - await _consume() - - assert chunks, "partial content should have been yielded before the error" - assert all(getattr(c.choices[0], "finish_reason", None) is None for c in chunks) - assert wrapper.received_finish_reason is None - assert "without a finish_reason" in str(excinfo.value).lower() or "finish_reason" in str(excinfo.value) - assert '{"findings"' in (excinfo.value.generated_content or "") - - -def test_clean_eof_without_finish_reason_raises_sync_issue_40260(): - from litellm.exceptions import MidStreamFallbackError - - log = Logging( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hi"}], - stream=True, - call_type="completion", - start_time=time.time(), - litellm_call_id="sync-eof-repro", - function_id="sync-eof-repro", - ) - wrapper = CustomStreamWrapper( - completion_stream=iter([]), - model="gpt-3.5-turbo", - custom_llm_provider="openai", - logging_obj=log, - ) - with pytest.raises(MidStreamFallbackError): - next(wrapper) - - -def test_provider_finish_reason_still_synthesizes_terminal_chunk_issue_40260(): - log = Logging( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hi"}], - stream=True, - call_type="completion", - start_time=time.time(), - litellm_call_id="length-finish-repro", - function_id="length-finish-repro", - ) - wrapper = CustomStreamWrapper( - completion_stream=iter([]), - model="gpt-3.5-turbo", - custom_llm_provider="openai", - logging_obj=log, - ) - wrapper.received_finish_reason = "length" - chunk = next(wrapper) - assert chunk.choices[0].finish_reason == "length" - - -def test_baseten_eof_without_finish_reason_still_succeeds_issue_40260(): - log = Logging( - model="baseten/model", - messages=[{"role": "user", "content": "hi"}], - stream=True, - call_type="completion", - start_time=time.time(), - litellm_call_id="baseten-eof", - function_id="baseten-eof", - ) - wrapper = CustomStreamWrapper( - completion_stream=iter([]), - model="baseten/model", - custom_llm_provider="baseten", - logging_obj=log, - ) - chunk = next(wrapper) - assert chunk.choices[0].finish_reason == "stop" - - -def test_vllm_eof_without_finish_reason_still_succeeds_issue_40260(): - log = Logging( - model="vllm/model", - messages=[{"role": "user", "content": "hi"}], - stream=True, - call_type="completion", - start_time=time.time(), - litellm_call_id="vllm-eof", - function_id="vllm-eof", - ) - wrapper = CustomStreamWrapper( - completion_stream=iter([]), - model="vllm/model", - custom_llm_provider="vllm", - logging_obj=log, - ) - chunk = next(wrapper) - assert chunk.choices[0].finish_reason == "stop" - +@file:/tmp/litellm-work/litellm/tests/test_litellm/litellm_core_utils/test_streaming_handler.py \ No newline at end of file diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index e3694ad5cdc..f6b2d8d50fc 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -1,13427 +1 @@ -import asyncio -import copy -import functools -import json -import logging -import os -import threading -from datetime import datetime -from types import SimpleNamespace -from typing import Final -from unittest.mock import AsyncMock, MagicMock, patch - -import httpx -import openai -import pytest -import respx - - - -import litellm -from litellm import Router -from litellm.exceptions import MidStreamFallbackError -from litellm.integrations.custom_guardrail import CustomGuardrail -from litellm.integrations.custom_logger import CustomLogger -from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging -from litellm.llms.bedrock.common_utils import BedrockError -from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import ( - SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES, -) -from litellm.router import ( - MAX_BUFFERED_PRE_CONTENT_ANTHROPIC_CHUNKS, - FallbackAwareAnthropicMessagesStream, - _anthropic_stream_commits_now, - _anthropic_stream_fallback_error_for_raised, - _anthropic_stream_raised_error_status, - _anthropic_stream_should_decline_fallback, - _anthropic_stream_error_is_gateway_verdict, - _anthropic_stream_forwards_ping_live, - _anthropic_stream_should_drop_pre_content_ping, - _is_retriable_anthropic_status, -) -from litellm.types.router import DeploymentTypedDict - - -def test_update_kwargs_does_not_mutate_defaults_and_merges_metadata(): - # initialize a real Router (env‑vars can be empty) - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": { - "model": "azure/gpt-4.1-mini", - "api_key": os.getenv("AZURE_AI_API_KEY"), - "api_version": os.getenv("AZURE_API_VERSION"), - "api_base": os.getenv("AZURE_AI_API_BASE"), - }, - } - ], - ) - - # override to known defaults for the test - router.default_litellm_params = { - "foo": "bar", - "metadata": {"baz": 123}, - } - original = copy.deepcopy(router.default_litellm_params) - kwargs: dict = {} - - # invoke the helper - router._update_kwargs_with_default_litellm_params( - kwargs=kwargs, - metadata_variable_name="litellm_metadata", - ) - - # 1) router.defaults must be unchanged - assert router.default_litellm_params == original - - # 2) non‑metadata keys get merged - assert kwargs["foo"] == "bar" - - # 3) metadata lands under "metadata" - assert kwargs["litellm_metadata"] == {"baz": 123} - - -def test_router_with_model_info_and_model_group(): - """ - Test edge case where user specifies model_group in model_info - """ - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": { - "model": "gpt-3.5-turbo", - }, - "model_info": { - "tpm": 1000, - "rpm": 1000, - "model_group": "gpt-3.5-turbo", - }, - } - ], - ) - - router._set_model_group_info( - model_group="gpt-3.5-turbo", - user_facing_model_group_name="gpt-3.5-turbo", - ) - - -def test_router_model_group_encrypted_content_affinity_callback_registration(): - from litellm.router_utils.pre_call_checks.deployment_affinity_check import ( - DeploymentAffinityCheck, - ) - from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import ( - EncryptedContentAffinityCheck, - ) - - model_group = "openai.gpt-5.1-codex" - model_group_affinity_config = { - model_group: ["encrypted_content_affinity"], - } - original_callbacks = list(litellm.callbacks) - litellm.callbacks = [] - router = None - - try: - router = litellm.Router( - model_list=[ - { - "model_name": model_group, - "litellm_params": { - "model": "openai/gpt-5.1-codex", - "api_key": "mock-api-key", - }, - } - ], - model_group_affinity_config=model_group_affinity_config, - num_retries=0, - ) - callbacks = router.optional_callbacks or [] - encrypted_content_callbacks = [ - cb for cb in callbacks if isinstance(cb, EncryptedContentAffinityCheck) - ] - deployment_callback = next( - cb for cb in callbacks if isinstance(cb, DeploymentAffinityCheck) - ) - assert len(encrypted_content_callbacks) == 1 - assert encrypted_content_callbacks[0].enable_global_affinity is False - assert ( - encrypted_content_callbacks[0].model_group_affinity_config - == model_group_affinity_config - ) - assert callbacks.index(encrypted_content_callbacks[0]) < callbacks.index( - deployment_callback - ) - assert litellm.callbacks.index(encrypted_content_callbacks[0]) < ( - litellm.callbacks.index(deployment_callback) - ) - - router._add_encrypted_content_affinity_check(enable_global_affinity=True) - - callbacks = router.optional_callbacks or [] - encrypted_content_callbacks = [ - cb for cb in callbacks if isinstance(cb, EncryptedContentAffinityCheck) - ] - assert len(encrypted_content_callbacks) == 1 - assert encrypted_content_callbacks[0].enable_global_affinity is True - assert encrypted_content_callbacks[0].router is router - finally: - if router is not None: - router.discard() - litellm.callbacks = original_callbacks - - -@pytest.mark.asyncio -async def test_encrypted_content_affinity_model_group_config_is_additive(): - from litellm.responses.utils import ResponsesAPIRequestUtils - from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import ( - EncryptedContentAffinityCheck, - ) - - model_group = "openai.gpt-5.1-codex" - target_deployment = { - "model_name": model_group, - "litellm_params": {"model": "openai/gpt-5.1-codex"}, - "model_info": {"id": "deployment-b"}, - } - healthy_deployments = [ - { - "model_name": model_group, - "litellm_params": {"model": "openai/gpt-5.1-codex"}, - "model_info": {"id": "deployment-a"}, - }, - target_deployment, - ] - encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id( - "deployment-b", "rs_test" - ) - - assert EncryptedContentAffinityCheck.has_model_group_affinity_enabled( - {model_group: ["encrypted_content_affinity"]} - ) - assert not EncryptedContentAffinityCheck.has_model_group_affinity_enabled(None) - - per_group_check = EncryptedContentAffinityCheck( - enable_global_affinity=False, - model_group_affinity_config={ - model_group: ["encrypted_content_affinity"], - }, - ) - request_kwargs = { - "input": [{"type": "reasoning", "id": encoded_id}], - "litellm_metadata": {}, - } - filtered = await per_group_check.async_filter_deployments( - model=model_group, - healthy_deployments=healthy_deployments, - messages=None, - request_kwargs=request_kwargs, - ) - - assert filtered == [target_deployment] - assert request_kwargs["litellm_metadata"]["encrypted_content_affinity_enabled"] - - disabled_check = EncryptedContentAffinityCheck( - enable_global_affinity=False, - model_group_affinity_config={ - "other-model-group": ["encrypted_content_affinity"], - }, - ) - disabled_request_kwargs = { - "input": [{"type": "reasoning", "id": encoded_id}], - "litellm_metadata": {}, - } - unfiltered = await disabled_check.async_filter_deployments( - model=model_group, - healthy_deployments=healthy_deployments, - messages=None, - request_kwargs=disabled_request_kwargs, - ) - - assert unfiltered == healthy_deployments - assert ( - "encrypted_content_affinity_enabled" - not in disabled_request_kwargs["litellm_metadata"] - ) - - global_check = EncryptedContentAffinityCheck( - enable_global_affinity=True, - model_group_affinity_config={ - model_group: ["deployment_affinity"], - }, - ) - global_request_kwargs = { - "input": [{"type": "reasoning", "id": encoded_id}], - "litellm_metadata": {}, - } - globally_filtered = await global_check.async_filter_deployments( - model=model_group, - healthy_deployments=healthy_deployments, - messages=None, - request_kwargs=global_request_kwargs, - ) - - assert globally_filtered == [target_deployment] - assert global_request_kwargs["litellm_metadata"][ - "encrypted_content_affinity_enabled" - ] - - -@pytest.mark.asyncio -async def test_encrypted_content_affinity_takes_priority_over_user_key_affinity(): - from litellm.responses.utils import ResponsesAPIRequestUtils - from litellm.router_utils.pre_call_checks.deployment_affinity_check import ( - DeploymentAffinityCheck, - ) - from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import ( - EncryptedContentAffinityCheck, - ) - - model_group = "openai.gpt-5.1-codex" - user_api_key_hash = "test-user-key" - deployment_a = { - "model_name": model_group, - "litellm_params": { - "model": "openai/gpt-5.1-codex", - "api_key": "mock-api-key-a", - }, - "model_info": {"id": "deployment-a"}, - } - deployment_b = { - "model_name": model_group, - "litellm_params": { - "model": "openai/gpt-5.1-codex", - "api_key": "mock-api-key-b", - }, - "model_info": {"id": "deployment-b"}, - } - original_callbacks = list(litellm.callbacks) - litellm.callbacks = [] - router = None - - try: - router = litellm.Router( - model_list=[deployment_a, deployment_b], - model_group_affinity_config={ - model_group: [ - "deployment_affinity", - "encrypted_content_affinity", - ], - }, - num_retries=0, - ) - callbacks = router.optional_callbacks or [] - deployment_callback = next( - cb for cb in callbacks if isinstance(cb, DeploymentAffinityCheck) - ) - encrypted_content_callback = next( - cb for cb in callbacks if isinstance(cb, EncryptedContentAffinityCheck) - ) - assert callbacks.index(encrypted_content_callback) < callbacks.index( - deployment_callback - ) - assert litellm.callbacks.index(encrypted_content_callback) < ( - litellm.callbacks.index(deployment_callback) - ) - - cache_key = DeploymentAffinityCheck.get_affinity_cache_key( - model_group=model_group, - user_key=user_api_key_hash, - ) - await deployment_callback.cache.async_set_cache( - key=cache_key, - value={"model_id": "deployment-a"}, - ttl=60, - ) - encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id( - "deployment-b", "rs_test" - ) - request_kwargs = { - "input": [{"type": "reasoning", "id": encoded_id}], - "litellm_metadata": {"user_api_key_hash": user_api_key_hash}, - } - - filtered = await router.async_callback_filter_deployments( - model=model_group, - healthy_deployments=[deployment_a, deployment_b], - messages=None, - parent_otel_span=None, - request_kwargs=request_kwargs, - ) - - assert filtered == [deployment_b] - assert request_kwargs.get("_encrypted_content_affinity_pinned") is True - finally: - if router is not None: - router.discard() - litellm.callbacks = original_callbacks - - -@pytest.mark.asyncio -async def test_arouter_with_tags_and_fallbacks(): - """ - If fallback model missing tag, raise error - """ - from litellm import Router - - router = Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": { - "model": "gpt-3.5-turbo", - "mock_response": "Hello, world!", - "tags": ["test"], - }, - }, - { - "model_name": "anthropic-claude-3-5-sonnet", - "litellm_params": { - "model": "claude-sonnet-4-5-20250929", - "mock_response": "Hello, world 2!", - }, - }, - ], - fallbacks=[ - {"gpt-3.5-turbo": ["anthropic-claude-3-5-sonnet"]}, - ], - enable_tag_filtering=True, - ) - - with pytest.raises(litellm.InternalServerError): - response = await router.acompletion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "Hello, world!"}], - mock_testing_fallbacks=True, - metadata={"tags": ["test"]}, - ) - - -@pytest.mark.asyncio -async def test_async_router_acreate_file(): - """ - Write to all deployments of a model - """ - from unittest.mock import MagicMock, patch - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo"}, - }, - {"model_name": "gpt-3.5-turbo", "litellm_params": {"model": "gpt-4o-mini"}}, - ], - ) - - with patch("litellm.acreate_file", return_value=MagicMock()) as mock_acreate_file: - mock_acreate_file.return_value = MagicMock() - response = await router.acreate_file( - model="gpt-3.5-turbo", - purpose="test", - file=MagicMock(), - ) - - # assert that the mock_acreate_file was called twice - assert mock_acreate_file.call_count == 2 - - -@pytest.mark.asyncio -async def test_async_router_acreate_file_with_jsonl(): - """ - Test router.acreate_file with both JSONL and non-JSONL files - """ - import json - from io import BytesIO - from unittest.mock import MagicMock, patch - - # Create test JSONL content - jsonl_data = [ - { - "body": { - "model": "gpt-3.5-turbo-router", - "messages": [{"role": "user", "content": "test"}], - } - }, - { - "body": { - "model": "gpt-3.5-turbo-router", - "messages": [{"role": "user", "content": "test2"}], - } - }, - ] - jsonl_content = "\n".join(json.dumps(item) for item in jsonl_data) - jsonl_file = BytesIO(jsonl_content.encode("utf-8")) - jsonl_file.name = "test.jsonl" - - # Create test non-JSONL content - non_jsonl_content = "This is not a JSONL file" - non_jsonl_file = BytesIO(non_jsonl_content.encode("utf-8")) - non_jsonl_file.name = "test.txt" - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo-router", - "litellm_params": {"model": "gpt-3.5-turbo"}, - }, - { - "model_name": "gpt-3.5-turbo-router", - "litellm_params": {"model": "gpt-4o-mini"}, - }, - ], - ) - - with patch("litellm.acreate_file", return_value=MagicMock()) as mock_acreate_file: - # Test with JSONL file - response = await router.acreate_file( - model="gpt-3.5-turbo-router", - purpose="batch", - file=jsonl_file, - ) - - # Verify mock was called twice (once for each deployment) - print(f"mock_acreate_file.call_count: {mock_acreate_file.call_count}") - print(f"mock_acreate_file.call_args_list: {mock_acreate_file.call_args_list}") - assert mock_acreate_file.call_count == 2 - - # Get the file content passed to the first call - first_call_file = mock_acreate_file.call_args_list[0][1]["file"] - first_call_content = first_call_file.read().decode("utf-8") - - # Verify the model name was replaced in the JSONL content - first_line = json.loads(first_call_content.split("\n")[0]) - assert first_line["body"]["model"] == "gpt-3.5-turbo" - - # Reset mock for next test - mock_acreate_file.reset_mock() - - # Test with non-JSONL file - response = await router.acreate_file( - model="gpt-3.5-turbo-router", - purpose="user_data", - file=non_jsonl_file, - ) - - # Verify mock was called twice - assert mock_acreate_file.call_count == 2 - - # Get the file content passed to the first call - first_call_file = mock_acreate_file.call_args_list[0][1]["file"] - first_call_content = first_call_file.read().decode("utf-8") - - # Verify the non-JSONL content was not modified - assert first_call_content == non_jsonl_content - - -@pytest.mark.asyncio -async def test_async_router_acreate_file_does_not_fall_back_across_model_groups(): - """A file created for batches only exists under the credentials of the model group - the caller named. A cross-group fallback silently stores it with the wrong provider - and the later batch create against the named group permanently fails.""" - from unittest.mock import MagicMock, patch - - router = litellm.Router( - model_list=[ - { - "model_name": "azure-gpt", - "litellm_params": { - "model": "azure/my-azure-deployment", - "api_base": "http://127.0.0.1:9", - "api_key": "dummy-key", - "api_version": "2024-06-01", - }, - }, - { - "model_name": "openai-gpt", - "litellm_params": {"model": "gpt-4o-mini"}, - }, - ], - fallbacks=[{"azure-gpt": ["openai-gpt"]}], - ) - - def fail_azure(*args: object, **kwargs: object) -> MagicMock: - if kwargs.get("model") == "azure/my-azure-deployment": - raise litellm.APIConnectionError( - message="Connection error.", - llm_provider="azure", - model="azure/my-azure-deployment", - ) - return MagicMock() - - with patch("litellm.acreate_file", side_effect=fail_azure) as mock_acreate_file: - with pytest.raises(litellm.APIConnectionError): - await router.acreate_file( - model="azure-gpt", - purpose="batch", - file=MagicMock(), - ) - - called_models = [call.kwargs.get("model") for call in mock_acreate_file.call_args_list] - assert "azure/my-azure-deployment" in called_models - assert "gpt-4o-mini" not in called_models - - -@pytest.mark.asyncio -async def test_async_router_acancel_batch_does_not_fall_back_across_model_groups(monkeypatch: pytest.MonkeyPatch): - """The proxy cancels a managed batch by handing the router the deployment id decoded - from the unified batch id. A default (``*``) fallback matches that id like any other - model string, and the fallback provider is then asked to cancel a batch it never - issued, which can only answer not-found. The router re-raises the owner's error after - that wasted round trip, so the pin's observable is the foreign call never happening.""" - - monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) - router = litellm.Router( - model_list=[ - { - "model_name": "azure-gpt", - "litellm_params": { - "model": "azure/my-azure-deployment", - "api_base": "http://127.0.0.1:9", - "api_key": "dummy-key", - "api_version": "2024-06-01", - }, - "model_info": {"id": "azure-batch-dep"}, - }, - { - "model_name": "openai-gpt", - "litellm_params": {"model": "gpt-4o-mini", "api_key": "dummy-key"}, - }, - ], - default_fallbacks=["openai-gpt"], - ) - - with respx.mock(assert_all_called=False) as respx_mock: - azure_route = respx_mock.post(host="127.0.0.1").mock( - return_value=httpx.Response(401, json={"error": {"code": "401", "message": "invalid subscription key"}}) - ) - openai_route = respx_mock.post("https://api.openai.com/v1/batches/batch_owned_by_azure/cancel").mock( - return_value=httpx.Response( - 404, - json={ - "error": { - "message": "No batch found with id 'batch_owned_by_azure'.", - "type": "invalid_request_error", - "code": "batch_not_found", - } - }, - ) - ) - with pytest.raises(openai.AuthenticationError, match="invalid subscription key"): - await router.acancel_batch(model="azure-batch-dep", batch_id="batch_owned_by_azure") - - assert azure_route.called - assert not openai_route.called - - -@pytest.mark.asyncio -async def test_async_router_acreate_file_uses_deployment_custom_llm_provider(): - """ - Ensure file routing preserves deployment custom_llm_provider instead of - inferring provider from model string alone. - """ - from unittest.mock import MagicMock, patch - - router = litellm.Router( - model_list=[ - { - "model_name": "team-azure-batch", - "litellm_params": { - "model": "gpt-4.1-mini", - "custom_llm_provider": "azure", - "api_base": "https://example-resource.openai.azure.com", - }, - }, - ], - ) - - with patch("litellm.acreate_file", return_value=MagicMock()) as mock_acreate_file: - await router.acreate_file( - model="team-azure-batch", - purpose="batch", - file=MagicMock(), - ) - - assert mock_acreate_file.call_count == 1 - assert mock_acreate_file.call_args.kwargs["custom_llm_provider"] == "azure" - - -@pytest.mark.asyncio -async def test_async_router_acreate_file_forwards_target_model_names_to_litellm_proxy(): - import json - from io import BytesIO - from unittest.mock import MagicMock, patch - - jsonl_file = BytesIO( - json.dumps({"body": {"model": "chained-batch", "messages": [{"role": "user", "content": "hi"}]}}).encode( - "utf-8" - ) - ) - jsonl_file.name = "test.jsonl" - - router = litellm.Router( - model_list=[ - { - "model_name": "chained-batch", - "litellm_params": { - "model": "litellm_proxy/gpt-4.1-batch", - "api_base": "http://localhost:4001/v1", - "api_key": "sk-proxy-b", - }, - }, - ], - ) - - with patch("litellm.acreate_file", return_value=MagicMock()) as mock_acreate_file: - await router.acreate_file( - model="chained-batch", - purpose="batch", - file=jsonl_file, - ) - - assert mock_acreate_file.call_count == 1 - call_kwargs = mock_acreate_file.call_args.kwargs - assert call_kwargs["custom_llm_provider"] == "litellm_proxy" - assert call_kwargs["extra_body"] == {"target_model_names": "gpt-4.1-batch"} - uploaded_line = json.loads(call_kwargs["file"].read().decode("utf-8").split("\n")[0]) - assert uploaded_line["body"]["model"] == "gpt-4.1-batch" - - -@pytest.mark.asyncio -async def test_async_router_acreate_file_does_not_inject_target_model_names_for_other_providers(): - from unittest.mock import MagicMock, patch - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4.1-batch", - "litellm_params": {"model": "gpt-4.1"}, - }, - ], - ) - - with patch("litellm.acreate_file", return_value=MagicMock()) as mock_acreate_file: - await router.acreate_file( - model="gpt-4.1-batch", - purpose="batch", - file=MagicMock(), - ) - - assert mock_acreate_file.call_count == 1 - assert mock_acreate_file.call_args.kwargs.get("extra_body") is None - - -@pytest.mark.asyncio -async def test_async_router_acreate_file_litellm_proxy_sends_target_model_names_in_multipart_form(): - import json - from io import BytesIO - - import httpx - - jsonl_file = BytesIO( - json.dumps({"body": {"model": "chained-batch", "messages": [{"role": "user", "content": "hi"}]}}).encode( - "utf-8" - ) - ) - jsonl_file.name = "test.jsonl" - - router = litellm.Router( - model_list=[ - { - "model_name": "chained-batch", - "litellm_params": { - "model": "litellm_proxy/gpt-4.1-batch", - "api_base": "http://localhost:4001/v1", - "api_key": "sk-proxy-b", - }, - }, - ], - ) - - file_object_json = { - "id": "file-abc123", - "object": "file", - "bytes": 100, - "created_at": 1700000000, - "filename": "test.jsonl", - "purpose": "batch", - "status": "processed", - } - - with respx.mock(assert_all_called=True) as respx_mock: - create_route = respx_mock.post("http://localhost:4001/v1/files").mock( - return_value=httpx.Response(200, json=file_object_json) - ) - response = await router.acreate_file( - model="chained-batch", - purpose="batch", - file=jsonl_file, - ) - - assert response.id == "file-abc123" - request_body = create_route.calls.last.request.content - assert b'name="target_model_names"' in request_body - assert b"gpt-4.1-batch" in request_body - assert b'name="purpose"' in request_body - - -@pytest.mark.asyncio -async def test_async_router_afile_content_uses_deployment_custom_llm_provider(): - """ - Regression test: Ensure afile_content preserves deployment custom_llm_provider - when model name lacks provider prefix (e.g., "gpt-4.1-mini" instead of "azure/gpt-4.1-mini"). - - This prevents "None is not a valid LlmProviders" errors when calling file content operations. - """ - from unittest.mock import AsyncMock, MagicMock, patch - from litellm.types.llms.openai import HttpxBinaryResponseContent - - router = litellm.Router( - model_list=[ - { - "model_name": "team-azure-batch", - "litellm_params": { - "model": "gpt-4.1-mini", # No provider prefix - "custom_llm_provider": "azure", - "api_base": "https://example-resource.openai.azure.com", - "api_key": "test-key", - }, - }, - ], - ) - - # Mock the Azure file handler's afile_content method - mock_response = MagicMock(spec=HttpxBinaryResponseContent) - mock_response.response = MagicMock() - - with patch( - "litellm.llms.azure.files.handler.AzureOpenAIFilesAPI.afile_content", - return_value=mock_response, - ) as mock_afile_content: - result = await router.afile_content( - model="team-azure-batch", - file_id="file-123", - ) - - # Verify the call was made (proves custom_llm_provider was correctly passed) - assert mock_afile_content.call_count == 1 - assert result == mock_response - - -@pytest.mark.asyncio -async def test_arouter_async_get_healthy_deployments(): - """ - Test that afile_content returns the correct file content - """ - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo"}, - }, - ], - ) - - result = await router.async_get_healthy_deployments( - model="gpt-3.5-turbo", - request_kwargs={}, - messages=None, - input=None, - specific_deployment=False, - parent_otel_span=None, - ) - - assert len(result) == 1 - assert result[0]["model_name"] == "gpt-3.5-turbo" - assert result[0]["litellm_params"]["model"] == "gpt-3.5-turbo" - - -@pytest.mark.asyncio -@patch("litellm.amoderation") -async def test_arouter_amoderation_with_credential_name(mock_amoderation): - """ - Test that router.amoderation passes litellm_credential_name to the underlying litellm.amoderation call - """ - mock_amoderation.return_value = AsyncMock() - - router = litellm.Router( - model_list=[ - { - "model_name": "text-moderation-stable", - "litellm_params": { - "model": "text-moderation-stable", - "litellm_credential_name": "my-custom-auth", - }, - }, - ], - ) - - await router.amoderation(input="I love everyone!", model="text-moderation-stable") - - mock_amoderation.assert_called_once() - call_kwargs = mock_amoderation.call_args[1] # Get the kwargs of the call - print( - "call kwargs for router.amoderation=", - json.dumps(call_kwargs, indent=4, default=str), - ) - assert call_kwargs["litellm_credential_name"] == "my-custom-auth" - assert call_kwargs["model"] == "text-moderation-stable" - - -def test_arouter_test_team_model(): - """ - Test that router.test_team_model returns the correct model - """ - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo"}, - "model_info": { - "team_id": "test-team", - "team_public_model_name": "test-model", - }, - }, - ], - ) - - result = router.map_team_model(team_model_name="test-model", team_id="test-team") - assert result is not None - - -def test_arouter_ignore_invalid_deployments(): - """ - Test that router.ignore_invalid_deployments is set to True - """ - from litellm.types.router import Deployment - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "my-bad-model"}, - }, - ], - ignore_invalid_deployments=True, - ) - - assert router.ignore_invalid_deployments is True - assert router.get_model_list() == [] - - ## check upsert deployment - router.upsert_deployment( - Deployment( - model_name="gpt-3.5-turbo", - litellm_params={"model": "my-bad-model"}, # type: ignore - model_info={"tpm": 1000, "rpm": 1000}, - ) - ) - - assert router.get_model_list() == [] - - -@pytest.mark.asyncio -async def test_arouter_aretrieve_batch(): - """ - Test that router.aretrieve_batch returns the correct response - """ - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": { - "model": "gpt-3.5-turbo", - "custom_llm_provider": "azure", - "api_key": "my-custom-key", - "api_base": "my-custom-base", - }, - } - ], - ) - - with patch.object( - litellm, "aretrieve_batch", return_value=AsyncMock() - ) as mock_aretrieve_batch: - try: - response = await router.aretrieve_batch( - model="gpt-3.5-turbo", - ) - except Exception as e: - print(f"Error: {e}") - - mock_aretrieve_batch.assert_called_once() - - print(mock_aretrieve_batch.call_args.kwargs) - assert mock_aretrieve_batch.call_args.kwargs["api_key"] == "my-custom-key" - assert mock_aretrieve_batch.call_args.kwargs["api_base"] == "my-custom-base" - - -@pytest.mark.asyncio -async def test_arouter_aretrieve_file_content(): - """ - Test that router.acreate_file with JSONL file returns the correct response - """ - - with patch.object( - litellm, "afile_content", return_value=AsyncMock() - ) as mock_afile_content: - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": { - "model": "gpt-3.5-turbo", - "custom_llm_provider": "azure", - "api_key": "my-custom-key", - "api_base": "my-custom-base", - }, - } - ], - ) - try: - response = await router.afile_content( - **{ - "model": "gpt-3.5-turbo", - "file_id": "my-unique-file-id", - } - ) # type: ignore - except Exception as e: - print(f"Error: {e}") - - mock_afile_content.assert_called_once() - - print(mock_afile_content.call_args.kwargs) - assert mock_afile_content.call_args.kwargs["api_key"] == "my-custom-key" - assert mock_afile_content.call_args.kwargs["api_base"] == "my-custom-base" - - -@pytest.mark.asyncio -async def test_arouter_filter_team_based_models(): - """ - Test that router.filter_team_based_models filters out models that are not in the team - """ - from litellm.types.router import Deployment - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo"}, - "model_info": { - "team_id": "test-team", - }, - }, - ], - ) - - # WORKS - result = await router.acompletion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "Hello, world!"}], - metadata={"user_api_key_team_id": "test-team"}, - mock_response="Hello, world!", - ) - - assert result is not None - - # FAILS - with pytest.raises(Exception, match='No deployments available for selected model, Try again in') as e: - result = await router.acompletion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "Hello, world!"}], - metadata={"user_api_key_team_id": "test-team-2"}, - mock_response="Hello, world!", - ) - assert "No deployments available" in str(e.value) - - ## ADD A MODEL THAT IS NOT IN THE TEAM - router.add_deployment( - Deployment( - model_name="gpt-3.5-turbo", - litellm_params={"model": "gpt-3.5-turbo"}, # type: ignore - model_info={"tpm": 1000, "rpm": 1000}, - ) - ) - - result = await router.acompletion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "Hello, world!"}], - metadata={"user_api_key_team_id": "test-team-2"}, - mock_response="Hello, world!", - ) - - assert result is not None - - -def test_arouter_should_include_deployment(): - """ - Test the should_include_deployment method with various scenarios - - The method logic: - 1. Returns True if: team_id matches AND model_name matches team_public_model_name - 2. Returns True if: model_name matches AND deployment has no team_id - 3. Otherwise returns False - """ - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo"}, - "model_info": { - "team_id": "test-team", - }, - }, - ], - ) - - # Test deployment structures - deployment_with_team_and_public_name = { - "model_name": "gpt-3.5-turbo", - "model_info": { - "team_id": "test-team", - "team_public_model_name": "team-gpt-model", - }, - } - - deployment_with_team_no_public_name = { - "model_name": "gpt-3.5-turbo", - "model_info": { - "team_id": "test-team", - }, - } - - deployment_without_team = { - "model_name": "gpt-4", - "model_info": {}, - } - - deployment_different_team = { - "model_name": "claude-3", - "model_info": { - "team_id": "other-team", - "team_public_model_name": "team-claude-model", - }, - } - - # Test Case 1: Team-specific deployment - team_id and team_public_model_name match - result = router.should_include_deployment( - model_name="team-gpt-model", - model=deployment_with_team_and_public_name, - team_id="test-team", - ) - assert ( - result is True - ), "Should return True when team_id and team_public_model_name match" - - # Test Case 2: Team-specific deployment - team_id matches but model_name doesn't match team_public_model_name - result = router.should_include_deployment( - model_name="different-model", - model=deployment_with_team_and_public_name, - team_id="test-team", - ) - assert ( - result is False - ), "Should return False when team_id matches but model_name doesn't match team_public_model_name" - - # Test Case 3: Team-specific deployment - team_id doesn't match - result = router.should_include_deployment( - model_name="team-gpt-model", - model=deployment_with_team_and_public_name, - team_id="different-team", - ) - assert result is False, "Should return False when team_id doesn't match" - - # Test Case 4: Team-specific deployment with no team_public_model_name - should fail - result = router.should_include_deployment( - model_name="gpt-3.5-turbo", - model=deployment_with_team_no_public_name, - team_id="test-team", - ) - assert ( - result is True - ), "Should return True when team deployment has no team_public_model_name to match" - - # Test Case 5: Non-team deployment - model_name matches and no team_id - result = router.should_include_deployment( - model_name="gpt-4", model=deployment_without_team, team_id=None - ) - assert ( - result is True - ), "Should return True when model_name matches and deployment has no team_id" - - # Test Case 6: Non-team deployment - model_name matches but team_id provided (should still work) - result = router.should_include_deployment( - model_name="gpt-4", model=deployment_without_team, team_id="any-team" - ) - assert ( - result is True - ), "Should return True when model_name matches non-team deployment, regardless of team_id param" - - # Test Case 7: Non-team deployment - model_name doesn't match - result = router.should_include_deployment( - model_name="different-model", model=deployment_without_team, team_id=None - ) - assert result is False, "Should return False when model_name doesn't match" - - # Test Case 8: Team deployment accessed without matching team_id - result = router.should_include_deployment( - model_name="gpt-3.5-turbo", - model=deployment_with_team_and_public_name, - team_id=None, - ) - assert ( - result is True - ), "Should return True when matching model with exact model_name" - - -def test_arouter_responses_api_bridge(): - """ - Test that router.responses_api_bridge returns the correct response - """ - from unittest.mock import MagicMock, patch - - from litellm.llms.custom_httpx.http_handler import HTTPHandler - - router = litellm.Router( - model_list=[ - { - "model_name": "[IP-approved] o3-pro", - "litellm_params": { - "model": "azure/responses/o_series/webinterface-o3-pro", - "api_base": "https://webhook.site/fba79dae-220a-4bb7-9a3a-8caa49604e55", - "api_key": "sk-1234567890", - "api_version": "preview", - "stream": True, - }, - "model_info": { - "input_cost_per_token": 0.00002, - "output_cost_per_token": 0.00008, - }, - } - ], - ) - - ## CONFIRM BRIDGE IS CALLED - with patch.object(litellm, "responses", return_value=AsyncMock()) as mock_responses: - result = router.completion( - model="[IP-approved] o3-pro", - messages=[{"role": "user", "content": "Hello, world!"}], - ) - assert mock_responses.call_count == 1 - - ## CONFIRM MODEL NAME IS STRIPPED - client = HTTPHandler() - - mock_response = MagicMock() - mock_response.status_code = 200 - mock_response.headers = {"content-type": "application/json"} - mock_response.json.return_value = { - "id": "resp_test", - "object": "response", - "status": "completed", - "output": [], - } - mock_response.text = ( - '{"id": "resp_test", "object": "response", "status": "completed", "output": []}' - ) - - with patch.object(client, "post", return_value=mock_response) as mock_post: - try: - result = router.completion( - model="[IP-approved] o3-pro", - messages=[{"role": "user", "content": "Hello, world!"}], - client=client, - num_retries=0, - ) - except Exception as e: - print(f"Error: {e}") - - assert mock_post.call_count == 1 - assert ( - mock_post.call_args.kwargs["url"] - == "https://webhook.site/fba79dae-220a-4bb7-9a3a-8caa49604e55/openai/v1/responses?api-version=preview" - ) - assert mock_post.call_args.kwargs["json"]["model"] == "webinterface-o3-pro" - - -@pytest.mark.asyncio -async def test_router_v1_messages_fallbacks(): - """ - Test that router.v1_messages_fallbacks returns the correct response - """ - router = litellm.Router( - model_list=[ - { - "model_name": "claude-sonnet-4-5-20250929", - "litellm_params": { - "model": "anthropic/claude-sonnet-4-5-20250929", - "mock_response": "litellm.InternalServerError", - }, - }, - { - "model_name": "bedrock-claude", - "litellm_params": { - "model": "anthropic.claude-haiku-4-5-20251001-v1:0", - "mock_response": "Hello, world I am a fallback!", - }, - }, - ], - fallbacks=[ - {"claude-sonnet-4-5-20250929": ["bedrock-claude"]}, - ], - ) - - result = await router.aanthropic_messages( - model="claude-sonnet-4-5-20250929", - messages=[{"role": "user", "content": "Hello, world!"}], - max_tokens=256, - ) - assert result is not None - - print(result) - assert result["content"][0]["text"] == "Hello, world I am a fallback!" - - -def test_add_invalid_provider_to_router(): - """ - Test that router.add_deployment raises an error if the provider is invalid - """ - from litellm.types.router import Deployment - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo"}, - } - ], - ) - - with pytest.raises(Exception, match='Unsupported provider - vertex_ai_eu') as e: - router.add_deployment( - Deployment( - model_name="vertex_ai/*", - litellm_params={ - "model": "vertex_ai/*", - "custom_llm_provider": "vertex_ai_eu", - }, - ) - ) - - assert router.pattern_router.patterns == {} - - -@pytest.mark.asyncio -async def test_router_ageneric_api_call_with_fallbacks_helper(): - """ - Test the _ageneric_api_call_with_fallbacks_helper method with various scenarios - """ - from unittest.mock import patch - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": { - "model": "gpt-3.5-turbo", - "api_key": "test-key", - "api_base": "https://api.openai.com/v1", - }, - "model_info": { - "tpm": 1000, - "rpm": 1000, - }, - }, - ], - ) - - # Test 1: Successful call - async def mock_generic_function(**kwargs): - return {"result": "success", "model": kwargs.get("model")} - - with patch.object(router, "async_get_available_deployment") as mock_get_deployment: - mock_get_deployment.return_value = { - "model_name": "gpt-3.5-turbo", - "litellm_params": { - "model": "gpt-3.5-turbo", - "api_key": "test-key", - "api_base": "https://api.openai.com/v1", - }, - } - - with patch.object( - router, "_update_kwargs_with_deployment" - ) as mock_update_kwargs: - with patch.object( - router, "async_routing_strategy_pre_call_checks" - ) as mock_pre_call_checks: - with patch.object( - router, "_get_client", return_value=None - ) as mock_get_client: - result = await router._ageneric_api_call_with_fallbacks_helper( - model="gpt-3.5-turbo", - original_generic_function=mock_generic_function, - messages=[{"role": "user", "content": "test"}], - ) - - assert result is not None - assert result["result"] == "success" - mock_get_deployment.assert_called_once() - mock_update_kwargs.assert_called_once() - mock_pre_call_checks.assert_called_once() - - # Test 2: Passthrough on no deployment (success case) - async def mock_passthrough_function(**kwargs): - return {"result": "passthrough", "model": kwargs.get("model")} - - with patch.object(router, "async_get_available_deployment") as mock_get_deployment: - mock_get_deployment.side_effect = Exception("No deployment available") - - result = await router._ageneric_api_call_with_fallbacks_helper( - model="gpt-3.5-turbo", - original_generic_function=mock_passthrough_function, - passthrough_on_no_deployment=True, - messages=[{"role": "user", "content": "test"}], - ) - - assert result is not None - assert result["result"] == "passthrough" - assert result["model"] == "gpt-3.5-turbo" - - # Test 3: No deployment available and passthrough=False (should raise exception) - with patch.object(router, "async_get_available_deployment") as mock_get_deployment: - mock_get_deployment.side_effect = Exception("No deployment available") - - with pytest.raises(Exception, match='No deployment available') as exc_info: - await router._ageneric_api_call_with_fallbacks_helper( - model="gpt-3.5-turbo", - original_generic_function=mock_generic_function, - passthrough_on_no_deployment=False, - messages=[{"role": "user", "content": "test"}], - ) - - assert "No deployment available" in str(exc_info.value) - - # Test 4: Test with semaphore (rate limiting) - import asyncio - - async def mock_semaphore_function(**kwargs): - return {"result": "semaphore_success", "model": kwargs.get("model")} - - with patch.object(router, "async_get_available_deployment") as mock_get_deployment: - mock_get_deployment.return_value = { - "model_name": "gpt-3.5-turbo", - "litellm_params": { - "model": "gpt-3.5-turbo", - "api_key": "test-key", - "api_base": "https://api.openai.com/v1", - }, - } - - mock_semaphore = asyncio.Semaphore(1) - - with patch.object( - router, "_update_kwargs_with_deployment" - ) as mock_update_kwargs: - with patch.object( - router, "_get_client", return_value=mock_semaphore - ) as mock_get_client: - with patch.object( - router, "async_routing_strategy_pre_call_checks" - ) as mock_pre_call_checks: - result = await router._ageneric_api_call_with_fallbacks_helper( - model="gpt-3.5-turbo", - original_generic_function=mock_semaphore_function, - messages=[{"role": "user", "content": "test"}], - ) - - assert result is not None - assert result["result"] == "semaphore_success" - mock_get_client.assert_called_once() - mock_pre_call_checks.assert_called_once() - - # Test 5: Test call tracking (success and failure counts) - initial_success_count = router.success_calls.get("gpt-3.5-turbo", 0) - initial_fail_count = router.fail_calls.get("gpt-3.5-turbo", 0) - - async def mock_failing_function(**kwargs): - raise Exception("Mock failure") - - with patch.object(router, "async_get_available_deployment") as mock_get_deployment: - mock_get_deployment.return_value = { - "model_name": "gpt-3.5-turbo", - "litellm_params": { - "model": "gpt-3.5-turbo", - "api_key": "test-key", - "api_base": "https://api.openai.com/v1", - }, - } - - with patch.object( - router, "_update_kwargs_with_deployment" - ) as mock_update_kwargs: - with patch.object( - router, "_get_client", return_value=None - ) as mock_get_client: - with patch.object( - router, "async_routing_strategy_pre_call_checks" - ) as mock_pre_call_checks: - with pytest.raises(Exception, match='Mock failure') as exc_info: - await router._ageneric_api_call_with_fallbacks_helper( - model="gpt-3.5-turbo", - original_generic_function=mock_failing_function, - messages=[{"role": "user", "content": "test"}], - ) - - assert "Mock failure" in str(exc_info.value) - # Check that fail_calls was incremented - assert router.fail_calls["gpt-3.5-turbo"] == initial_fail_count + 1 - - -@pytest.mark.asyncio -async def test_ageneric_api_call_deployment_model_overrides_alias(): - """ - Regression: when a model alias (e.g. "not-gemini-2.5-flash") maps to a deployment - with model="vertex_ai/gemini-2.5-flash", the underlying litellm function must receive - the deployment model, not the alias. Before the fix, **kwargs overwrote data["model"]. - """ - from unittest.mock import patch - - captured: dict = {} - - async def capture_model(**kwargs): - captured["model"] = kwargs.get("model") - return {"result": "ok"} - - router = litellm.Router( - model_list=[ - { - "model_name": "not-gemini-2.5-flash", - "litellm_params": { - "model": "vertex_ai/gemini-2.5-flash", - "api_key": "fake-key", - }, - } - ] - ) - - def inject_alias_into_kwargs(deployment, kwargs, function_name=None): - # Simulate the alias leaking into kwargs (as happens when - # _ageneric_api_call_with_fallbacks sets kwargs["model"] = alias before - # calling the helper through async_function_with_fallbacks). - kwargs["model"] = "not-gemini-2.5-flash" - - with ( - patch.object(router, "async_get_available_deployment") as mock_dep, - patch.object( - router, - "_update_kwargs_with_deployment", - side_effect=inject_alias_into_kwargs, - ), - patch.object(router, "async_routing_strategy_pre_call_checks"), - patch.object(router, "_get_client", return_value=None), - ): - mock_dep.return_value = { - "model_name": "not-gemini-2.5-flash", - "litellm_params": { - "model": "vertex_ai/gemini-2.5-flash", - "api_key": "fake-key", - }, - } - - await router._ageneric_api_call_with_fallbacks_helper( - model="not-gemini-2.5-flash", - original_generic_function=capture_model, - ) - - assert ( - captured["model"] == "vertex_ai/gemini-2.5-flash" - ), f"Expected deployment model 'vertex_ai/gemini-2.5-flash', got '{captured['model']}'" - - -@pytest.mark.asyncio -async def test_ageneric_api_call_resolves_realtime_session_model(): - """ - Regression for #36742: realtime client secret requests carry the model inside `session` too, and the proxy - fills it with the pre-routing model group name. The underlying litellm function reads session.model first, - so it must see the resolved deployment, while a caller's nested transcription model stays untouched. - """ - routed: Final = AsyncMock(return_value={"result": "ok"}) - - router = litellm.Router( - model_list=[ - { - "model_name": "my-realtime-group", - "litellm_params": { - "model": "openai/gpt-realtime-2.1-mini", - "api_key": "fake-key", - }, - "model_info": {"mode": "realtime"}, - } - ] - ) - - await router._ageneric_api_call_with_fallbacks( - model="my-realtime-group", - original_function=routed, - session={ - "type": "realtime", - "model": "my-realtime-group", - "audio": {"input": {"transcription": {"model": "gpt-4o-transcribe"}}}, - }, - ) - - sent: Final = routed.call_args.kwargs - assert sent["model"] == "openai/gpt-realtime-2.1-mini" - assert sent["session"]["model"] == "openai/gpt-realtime-2.1-mini" - assert sent["session"]["audio"]["input"]["transcription"]["model"] == "gpt-4o-transcribe" - - -@pytest.mark.asyncio -async def test_ageneric_api_call_does_not_add_session_model(): - """ - A session that never carried a model must not gain one from routing: the underlying function then falls back - to the resolved `model` kwarg itself, and the outgoing session body keeps the caller's shape. - """ - routed: Final = AsyncMock(return_value={"result": "ok"}) - - router = litellm.Router( - model_list=[ - { - "model_name": "my-realtime-group", - "litellm_params": { - "model": "openai/gpt-realtime-2.1-mini", - "api_key": "fake-key", - }, - "model_info": {"mode": "realtime"}, - } - ] - ) - - await router._ageneric_api_call_with_fallbacks( - model="my-realtime-group", - original_function=routed, - session={"type": "realtime"}, - ) - - sent: Final = routed.call_args.kwargs - assert sent["model"] == "openai/gpt-realtime-2.1-mini" - assert sent["session"] == {"type": "realtime"} - - -@pytest.mark.parametrize( - "session, expected", - [ - ({"type": "realtime", "model": "my-realtime-group"}, {"session": {"type": "realtime", "model": "resolved"}}), - ({"type": "realtime"}, {}), - (None, {}), - ("not-a-session", {}), - ], -) -def test_with_router_resolved_session_model(session, expected): - from litellm.router import _with_router_resolved_session_model - - assert dict(_with_router_resolved_session_model(session, "resolved")) == expected - - -def test_router_get_model_access_groups_team_only_models(): - """ - Test that Router.get_model_access_groups returns the correct response for team-only models - """ - router = litellm.Router( - model_list=[ - { - "model_name": "my-custom-model-name", - "litellm_params": {"model": "gpt-3.5-turbo"}, - "model_info": { - "team_id": "team_1", - "access_groups": ["default-models"], - "team_public_model_name": "gpt-3.5-turbo", - }, - }, - ] - ) - - access_groups = router.get_model_access_groups( - model_name="gpt-3.5-turbo", team_id=None - ) - assert len(access_groups) == 0 - - access_groups = router.get_model_access_groups( - model_name="gpt-3.5-turbo", team_id="team_1" - ) - assert list(access_groups.keys()) == ["default-models"] - - -def test_cached_get_model_group_info(): - """ - Test that _cached_get_model_group_info caches results and - invalidates on deployment changes. - """ - from litellm.types.router import Deployment, LiteLLM_Params - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4", - "litellm_params": {"model": "gpt-4", "api_key": "fake"}, - "model_info": {"tpm": 1000, "rpm": 100}, - }, - ] - ) - - # First call should compute and cache - result1 = router._cached_get_model_group_info("gpt-4") - assert result1 is not None - assert result1.tpm == 1000 - - # Second call should hit cache (same object) - result2 = router._cached_get_model_group_info("gpt-4") - assert result1 is result2 - - # Add a deployment — cache should be invalidated - router.add_deployment( - Deployment( - model_name="gpt-4", - litellm_params=LiteLLM_Params(model="gpt-4", api_key="fake2"), - model_info={"tpm": 2000, "rpm": 200}, - ) - ) - result3 = router._cached_get_model_group_info("gpt-4") - assert result3 is not result2 - assert result3 is not None - assert result3.tpm == 3000 # 1000 + 2000 - - # Delete a deployment — cache should be invalidated - deployment_id = router.model_list[-1]["model_info"]["id"] - router.delete_deployment(id=deployment_id) - result4 = router._cached_get_model_group_info("gpt-4") - assert result4 is not result3 - assert result4 is not None - assert result4.tpm == 1000 - - # set_model_list — cache should be invalidated - router.set_model_list( - [ - { - "model_name": "gpt-4", - "litellm_params": {"model": "gpt-4", "api_key": "fake"}, - "model_info": {"tpm": 5000}, - }, - ] - ) - result5 = router._cached_get_model_group_info("gpt-4") - assert result5 is not result4 - assert result5 is not None - assert result5.tpm == 5000 - - # Verify cache still works after invalidation - result6 = router._cached_get_model_group_info("gpt-4") - assert result5 is result6 - - -def test_model_group_info_cost_from_db_model_info(): - """ - When get_deployment_model_info fails (model_info is None fallback), - input_cost_per_token and output_cost_per_token should be read from db model_info. - """ - from unittest.mock import patch - - router = litellm.Router( - model_list=[ - { - "model_name": "my-custom-model", - "litellm_params": { - "model": "openai/my-custom-model", - "api_key": "fake", - "api_base": "https://my-custom-endpoint.com", - }, - "model_info": { - "input_cost_per_token": 0.0001, - "output_cost_per_token": 0.0002, - }, - }, - ] - ) - - with patch.object( - router, "get_deployment_model_info", side_effect=Exception("not found") - ): - result = router._cached_get_model_group_info("my-custom-model") - assert result is not None - assert result.input_cost_per_token == 0.0001 - assert result.output_cost_per_token == 0.0002 - - -def test_model_group_info_cost_none_when_db_model_info_has_no_cost(): - """ - When get_deployment_model_info fails and db model_info has no cost fields, - input/output_cost_per_token should be None. - """ - from unittest.mock import patch - - router = litellm.Router( - model_list=[ - { - "model_name": "my-custom-model-no-cost", - "litellm_params": { - "model": "openai/my-custom-model-no-cost", - "api_key": "fake", - "api_base": "https://my-custom-endpoint.com", - }, - "model_info": {}, - }, - ] - ) - - with patch.object( - router, "get_deployment_model_info", side_effect=Exception("not found") - ): - result = router._cached_get_model_group_info("my-custom-model-no-cost") - assert result is not None - assert result.input_cost_per_token is None - assert result.output_cost_per_token is None - - -@pytest.mark.parametrize( - "value,expected", - [ - ("1e-05", 1e-05), - ("0.00001", 1e-05), - (1e-05, 1e-05), - (5, 5.0), - (None, None), - ("not-a-number", None), - ], -) -def test_cost_value_as_float(value, expected): - from litellm.router import _cost_value_as_float - - assert _cost_value_as_float(value) == expected - - -def test_model_group_info_with_stringified_cost_values(): - """ - YAML 1.2 parsers emit '1e-05' (integer mantissa) as a string, so cost - values in deployment model_info can arrive as str. Aggregating the model - group must not raise TypeError('>' between str and float) and must return - float costs. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "my-custom-model", - "litellm_params": { - "model": "openai/my-custom-backend-1", - "api_key": "fake", - }, - "model_info": { - "input_cost_per_token": "1e-05", - "output_cost_per_token": "1e-05", - }, - }, - { - "model_name": "my-custom-model", - "litellm_params": { - "model": "openai/my-custom-backend-2", - "api_key": "fake", - }, - "model_info": { - "input_cost_per_token": "2e-05", - "output_cost_per_token": "2e-05", - }, - }, - ] - ) - - def _model_info_with_str_costs(model_id: str, model_name: str): - for model in router.model_list: - if model["model_info"]["id"] == model_id: - return { - "key": model_name, - "input_cost_per_token": model["model_info"]["input_cost_per_token"], - "output_cost_per_token": model["model_info"]["output_cost_per_token"], - "litellm_provider": "openai", - "mode": "chat", - } - return None - - with patch.object( - router, "get_deployment_model_info", side_effect=_model_info_with_str_costs - ): - result = router._set_model_group_info( - model_group="my-custom-model", - user_facing_model_group_name="my-custom-model", - ) - - assert result is not None - assert result.input_cost_per_token == 2e-05 - assert result.output_cost_per_token == 2e-05 - assert isinstance(result.input_cost_per_token, float) - assert isinstance(result.output_cost_per_token, float) - - -def test_model_group_info_db_fallback_with_stringified_cost_values(): - """ - Fallback path: when get_deployment_model_info returns nothing, costs are - read straight from the deployment's model_info dict, which can hold - stringified floats parsed from YAML. They must be coerced to float. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "my-custom-model", - "litellm_params": { - "model": "openai/my-custom-backend-1", - "api_key": "fake", - }, - "model_info": { - "input_cost_per_token": "1e-05", - "output_cost_per_token": "3e-05", - }, - }, - { - "model_name": "my-custom-model", - "litellm_params": { - "model": "openai/my-custom-backend-2", - "api_key": "fake", - }, - "model_info": { - "input_cost_per_token": "2e-05", - "output_cost_per_token": "2e-05", - }, - }, - ] - ) - - with patch.object( - router, "get_deployment_model_info", side_effect=Exception("not found") - ): - result = router._set_model_group_info( - model_group="my-custom-model", - user_facing_model_group_name="my-custom-model", - ) - - assert result is not None - assert result.input_cost_per_token == 2e-05 - assert result.output_cost_per_token == 3e-05 - assert isinstance(result.input_cost_per_token, float) - assert isinstance(result.output_cost_per_token, float) - - -def test_get_model_access_groups_caching(): - """ - Test that get_model_access_groups caches the no-args result - and invalidates on deployment changes. - """ - from litellm.types.router import Deployment, LiteLLM_Params - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4", - "litellm_params": {"model": "gpt-4"}, - "model_info": {"access_groups": ["premium"]}, - }, - ] - ) - - # First call computes and populates cache - result1 = router.get_model_access_groups() - assert "premium" in result1 - - # All subsequent calls should return the same cached object (including first) - result2 = router.get_model_access_groups() - assert result1 is result2 - - # Calls with args should bypass cache - result_with_args = router.get_model_access_groups(model_name="gpt-4") - assert result_with_args is not result2 - - # Add a deployment — cache should be invalidated - router.add_deployment( - Deployment( - model_name="gpt-3.5", - litellm_params=LiteLLM_Params(model="gpt-3.5-turbo"), - model_info={"access_groups": ["default"]}, - ) - ) - result3 = router.get_model_access_groups() - assert result3 is not result2 - assert "premium" in result3 - assert "default" in result3 - - # Delete the deployment — cache should be invalidated again - deployment_id = None - for m in router.model_list: - if m.get("model_name") == "gpt-3.5": - deployment_id = m.get("model_info", {}).get("id") - break - assert deployment_id is not None - router.delete_deployment(id=deployment_id) - result4 = router.get_model_access_groups() - assert result4 is not result3 - assert "default" not in result4 - assert "premium" in result4 - - -def test_get_model_access_groups_cache_invalidation_set_model_list(): - """ - Test that set_model_list invalidates the access groups cache. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4", - "litellm_params": {"model": "gpt-4"}, - "model_info": {"access_groups": ["premium"]}, - }, - ] - ) - - # Populate cache - result1 = router.get_model_access_groups() - assert "premium" in result1 - - # set_model_list should invalidate cache - router.set_model_list( - [ - { - "model_name": "claude-3", - "litellm_params": {"model": "anthropic/claude-3-opus-20240229"}, - "model_info": {"access_groups": ["research"]}, - }, - ] - ) - result2 = router.get_model_access_groups() - assert result2 is not result1 - assert "research" in result2 - assert "premium" not in result2 - - -def test_get_model_access_groups_cache_invalidation_upsert_deployment(): - """ - Test that upsert_deployment invalidates the access groups cache. - """ - from litellm.types.router import Deployment, LiteLLM_Params - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4", - "litellm_params": {"model": "gpt-4"}, - "model_info": {"access_groups": ["premium"]}, - }, - ] - ) - - # Populate cache - result1 = router.get_model_access_groups() - assert "premium" in result1 - - # Get the existing deployment's ID - existing_id = router.model_list[0]["model_info"]["id"] - - # Upsert with the same ID but different params — triggers pop + re-add - router.upsert_deployment( - Deployment( - model_name="gpt-4-updated", - litellm_params=LiteLLM_Params(model="gpt-4-turbo"), - model_info={"id": existing_id, "access_groups": ["updated-group"]}, - ) - ) - result2 = router.get_model_access_groups() - assert result2 is not result1 - assert "updated-group" in result2 - - -@pytest.mark.asyncio -async def test_acompletion_streaming_iterator(): - """Test _acompletion_streaming_iterator for normal streaming and fallback behavior.""" - from unittest.mock import MagicMock - - from litellm.exceptions import MidStreamFallbackError - - # Helper class for creating async iterators - class AsyncIterator: - def __init__(self, items, error_after=None): - self.items = items - self.index = 0 - self.error_after = error_after - - def __aiter__(self): - return self - - async def __anext__(self): - if self.error_after is not None and self.index >= self.error_after: - raise self.error_after - if self.index >= len(self.items): - raise StopAsyncIteration - item = self.items[self.index] - self.index += 1 - return item - - # Set up router with fallback configuration - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4", - "litellm_params": {"model": "gpt-4", "api_key": "fake-key-1"}, - }, - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo", "api_key": "fake-key-2"}, - }, - ], - fallbacks=[{"gpt-4": ["gpt-3.5-turbo"]}], - set_verbose=True, - ) - - # Test data - messages = [{"role": "user", "content": "Hello"}] - initial_kwargs = {"model": "gpt-4", "stream": True, "temperature": 0.7} - - # Test 1: Successful streaming (no errors) - print("\n=== Test 1: Successful streaming ===") - - # Mock successful streaming response - mock_chunks = [ - MagicMock(choices=[MagicMock(delta=MagicMock(content="Hello"))]), - MagicMock(choices=[MagicMock(delta=MagicMock(content=" there"))]), - MagicMock(choices=[MagicMock(delta=MagicMock(content="!"))]), - ] - - mock_response = AsyncIterator(mock_chunks) - - setattr(mock_response, "model", "gpt-4") - setattr(mock_response, "custom_llm_provider", "openai") - setattr(mock_response, "logging_obj", MagicMock()) - - result = await router._acompletion_streaming_iterator( - model_response=mock_response, messages=messages, initial_kwargs=initial_kwargs - ) - - # Collect streamed chunks - collected_chunks = [] - async for chunk in result: - collected_chunks.append(chunk) - - assert len(collected_chunks) == 3 - assert all(chunk in mock_chunks for chunk in collected_chunks) - print("✓ Successfully streamed all chunks") - - # Test 2: MidStreamFallbackError with generated content is re-raised, not silently continued - print("\n=== Test 2: MidStreamFallbackError re-raises when content already generated ===") - - # Error with generated content and is_pre_first_chunk=False (the default): - # the router must re-raise instead of attempting a continuation-prompt fallback, - # because partial content has already been sent to the client. - error = MidStreamFallbackError( - message="Connection lost", - model="gpt-4", - llm_provider="openai", - generated_content="Hello", - ) - - class AsyncIteratorWithError: - def __init__(self, items, error_after_index): - self.items = items - self.index = 0 - self.error_after_index = error_after_index - self.chunks = [] - - def __aiter__(self): - return self - - async def __anext__(self): - if self.index >= len(self.items): - raise StopAsyncIteration - if self.index == self.error_after_index: - raise error - item = self.items[self.index] - self.index += 1 - return item - - mock_error_response = AsyncIteratorWithError(mock_chunks, 1) # Error after first chunk - - setattr(mock_error_response, "model", "gpt-4") - setattr(mock_error_response, "custom_llm_provider", "openai") - setattr(mock_error_response, "logging_obj", MagicMock()) - - result = await router._acompletion_streaming_iterator( - model_response=mock_error_response, - messages=messages, - initial_kwargs=initial_kwargs, - ) - - # Collect streamed chunks — the first chunk succeeds, then the error re-raises - collected_chunks = [] - async def _drain(): - async for chunk in result: - collected_chunks.append(chunk) - - with pytest.raises(MidStreamFallbackError): - await _drain() - - assert len(collected_chunks) == 1, "one chunk yielded before the error" - print("✓ MidStreamFallbackError re-raised correctly when content was already generated") - - print("\n=== All tests passed! ===") - - -@pytest.mark.asyncio -async def test_acompletion_streaming_iterator_reraises_original_exception_when_available(): - """Async: when the mid-stream MidStreamFallbackError wraps a real provider - exception (original_exception), the router must re-raise that original - exception instead of the internal wrapper, so the client sees the - specific error type/code (e.g. RateLimitError) rather than a generic - MidStreamFallbackError.""" - from unittest.mock import MagicMock - - from litellm.exceptions import MidStreamFallbackError, RateLimitError - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4", - "litellm_params": {"model": "gpt-4", "api_key": "fake-key"}, - } - ], - set_verbose=True, - ) - - messages = [{"role": "user", "content": "Test"}] - initial_kwargs = {"model": "gpt-4", "stream": True} - - original_exception = RateLimitError( - message="rate limited", - llm_provider="vertex_ai", - model="gpt-4", - ) - error = MidStreamFallbackError( - message="rate limited", - model="gpt-4", - llm_provider="openai", - original_exception=original_exception, - generated_content="Hello", - ) - - mock_chunks = [ - MagicMock(choices=[MagicMock(delta=MagicMock(content="Hello"))]), - MagicMock(choices=[MagicMock(delta=MagicMock(content=" there"))]), - ] - - class AsyncIteratorWithError: - def __init__(self, items, error_after_index): - self.items = items - self.index = 0 - self.error_after_index = error_after_index - - def __aiter__(self): - return self - - async def __anext__(self): - if self.index >= len(self.items): - raise StopAsyncIteration - if self.index == self.error_after_index: - raise error - item = self.items[self.index] - self.index += 1 - return item - - mock_error_response = AsyncIteratorWithError(mock_chunks, 1) - setattr(mock_error_response, "model", "gpt-4") - setattr(mock_error_response, "custom_llm_provider", "openai") - setattr(mock_error_response, "logging_obj", MagicMock()) - - result = await router._acompletion_streaming_iterator( - model_response=mock_error_response, - messages=messages, - initial_kwargs=initial_kwargs, - ) - - with pytest.raises(RateLimitError) as exc_info: - async for _ in result: - pass - assert exc_info.value is original_exception - assert exc_info.value.type == "throttling_error" - assert exc_info.value.code == "429" - - -@pytest.mark.asyncio -async def test_acompletion_streaming_iterator_edge_cases(): - """Test edge cases for _acompletion_streaming_iterator.""" - from unittest.mock import MagicMock - - from litellm.exceptions import MidStreamFallbackError - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4", - "litellm_params": {"model": "gpt-4", "api_key": "fake-key"}, - } - ], - set_verbose=True, - ) - - messages = [{"role": "user", "content": "Test"}] - initial_kwargs = {"model": "gpt-4", "stream": True} - - # Test: Empty generated content - empty_error = MidStreamFallbackError( - message="Error", - model="gpt-4", - llm_provider="openai", - generated_content="", # Empty content - ) - - class AsyncIteratorImmediateError: - def __init__(self): - self.model = "gpt-4" - self.custom_llm_provider = "openai" - self.logging_obj = MagicMock() - self.chunks = [] - - def __aiter__(self): - return self - - async def __anext__(self): - raise empty_error - - mock_response = AsyncIteratorImmediateError() - - # Mock empty fallback response using AsyncIterator - class EmptyAsyncIterator: - def __aiter__(self): - return self - - async def __anext__(self): - raise StopAsyncIteration - - mock_fallback_response = EmptyAsyncIterator() - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - return_value=mock_fallback_response, - ) as mock_fallback_utils: - collected_chunks = [] - iterator = await router._acompletion_streaming_iterator( - model_response=mock_response, - messages=messages, - initial_kwargs=initial_kwargs, - ) - - async for chunk in iterator: - collected_chunks.append(chunk) - - # Should still call fallback even with empty content - assert mock_fallback_utils.called - fallback_kwargs = mock_fallback_utils.call_args.kwargs["kwargs"] - modified_messages = fallback_kwargs["messages"] - - # Empty content → pre-first-chunk path uses original messages - # (no continuation prompt added) - assert modified_messages == messages - print("✓ Handles empty generated content correctly") - - print("✓ Edge case tests passed!") - - -@pytest.mark.asyncio -async def test_acompletion_streaming_iterator_preserves_hidden_params(): - """ - Regression test: FallbackStreamWrapper must copy _hidden_params from the - original CustomStreamWrapper so that x-litellm-overhead-duration-ms (and - other hidden params) are present in the proxy response headers for streaming. - """ - from unittest.mock import MagicMock - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4", - "litellm_params": {"model": "gpt-4", "api_key": "fake-key"}, - } - ], - ) - - # Simulate a CustomStreamWrapper that already has timing metadata set by - # update_response_metadata (litellm_overhead_time_ms, _response_ms, etc.) - mock_response = MagicMock() - mock_response.model = "gpt-4" - mock_response.custom_llm_provider = "openai" - mock_response.logging_obj = MagicMock() - mock_response._hidden_params = { - "litellm_overhead_time_ms": 12.34, - "_response_ms": 500.0, - "litellm_call_id": "test-call-id", - "api_base": "https://api.openai.com", - "additional_headers": {}, - } - - # Make the mock iterable (yields nothing — we only care about hidden_params) - async def _empty(): - return - yield # make it an async generator - - mock_response.__aiter__ = lambda self: _empty().__aiter__() - - result = await router._acompletion_streaming_iterator( - model_response=mock_response, - messages=[{"role": "user", "content": "hi"}], - initial_kwargs={"model": "gpt-4", "stream": True}, - ) - - # The returned FallbackStreamWrapper must carry the original _hidden_params - assert hasattr(result, "_hidden_params"), "result must have _hidden_params" - assert result._hidden_params.get("litellm_overhead_time_ms") == 12.34, ( - "litellm_overhead_time_ms must be preserved — " - "this is what drives x-litellm-overhead-duration-ms in streaming responses" - ) - assert result._hidden_params.get("litellm_call_id") == "test-call-id" - assert result._hidden_params.get("_response_ms") == 500.0 - - -def test_completion_streaming_iterator_fallback_on_429(): - """Sync streaming: MidStreamFallbackError (429 pre-first-chunk) triggers fallback. - - This is the sync counterpart of test_acompletion_streaming_iterator. - Before this fix, __next__ raised RateLimitError directly and the Router - never got a chance to fall back. - """ - from unittest.mock import MagicMock - - from litellm.exceptions import MidStreamFallbackError - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4", - "litellm_params": {"model": "gpt-4", "api_key": "fake-key"}, - } - ], - ) - - messages = [{"role": "user", "content": "Test"}] - initial_kwargs = {"model": "gpt-4", "stream": True} - - rate_limit_error = MidStreamFallbackError( - message="Resource exhausted", - model="gpt-4", - llm_provider="vertex_ai", - generated_content="", - is_pre_first_chunk=True, - ) - - class SyncIteratorImmediateError: - def __init__(self): - self.model = "gpt-4" - self.custom_llm_provider = "openai" - self.logging_obj = MagicMock() - self.chunks = [] - - def __iter__(self): - return self - - def __next__(self): - raise rate_limit_error - - mock_response = SyncIteratorImmediateError() - - # Fallback returns a simple non-streaming response (fallback may not stream) - mock_fallback_response = MagicMock() - mock_fallback_response.__iter__ = MagicMock(return_value=iter([])) - - with patch.object( - router, - "function_with_fallbacks", - return_value=mock_fallback_response, - ) as mock_fallback: - result = router._completion_streaming_iterator( - model_response=mock_response, - messages=messages, - initial_kwargs=initial_kwargs, - ) - - collected_chunks = list(result) - - assert mock_fallback.called - call_kwargs = mock_fallback.call_args - # Pre-first-chunk: should use original messages, no continuation prompt - assert call_kwargs.kwargs.get("messages") == messages - # Verify original_function is _completion (sync) - assert call_kwargs.kwargs.get("original_function") == router._completion - - -def test_completion_streaming_iterator_preserves_hidden_params(): - """SyncFallbackStreamWrapper must copy _hidden_params from original response.""" - from unittest.mock import MagicMock - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4", - "litellm_params": {"model": "gpt-4", "api_key": "fake-key"}, - } - ], - ) - - mock_response = MagicMock() - mock_response.model = "gpt-4" - mock_response.custom_llm_provider = "openai" - mock_response.logging_obj = MagicMock() - mock_response._hidden_params = { - "litellm_overhead_time_ms": 42.0, - "litellm_call_id": "test-sync-call", - } - mock_response.__iter__ = MagicMock(return_value=iter([])) - - result = router._completion_streaming_iterator( - model_response=mock_response, - messages=[{"role": "user", "content": "hi"}], - initial_kwargs={"model": "gpt-4", "stream": True}, - ) - - assert hasattr(result, "_hidden_params") - assert result._hidden_params.get("litellm_overhead_time_ms") == 42.0 - assert result._hidden_params.get("litellm_call_id") == "test-sync-call" - - -def test_completion_streaming_iterator_reraises_mid_chunk_error(): - """Sync: MidStreamFallbackError with generated_content and is_pre_first_chunk=False - must be re-raised immediately; the router cannot recover after partial content - has already been sent to the client.""" - from unittest.mock import MagicMock - - from litellm.exceptions import MidStreamFallbackError - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4", - "litellm_params": {"model": "gpt-4", "api_key": "fake-key"}, - } - ], - ) - - messages = [{"role": "user", "content": "Test"}] - initial_kwargs = {"model": "gpt-4", "stream": True} - - mid_chunk_error = MidStreamFallbackError( - message="Connection reset", - model="gpt-4", - llm_provider="openai", - generated_content="Hello, I am", - is_pre_first_chunk=False, - ) - - class SyncIteratorMidChunkError: - def __init__(self): - self.model = "gpt-4" - self.custom_llm_provider = "openai" - self.logging_obj = MagicMock() - self.chunks = [] - - def __iter__(self): - return self - - def __next__(self): - raise mid_chunk_error - - mock_response = SyncIteratorMidChunkError() - - result = router._completion_streaming_iterator( - model_response=mock_response, - messages=messages, - initial_kwargs=initial_kwargs, - ) - - with pytest.raises(MidStreamFallbackError): - list(result) - - -def test_completion_streaming_iterator_reraises_original_exception_when_available(): - """Sync: when the mid-chunk MidStreamFallbackError wraps a real provider - exception (original_exception), the router must re-raise that original - exception instead of the internal wrapper, so the client sees the - specific error type/code (e.g. RateLimitError) rather than a generic - MidStreamFallbackError.""" - from unittest.mock import MagicMock - - from litellm.exceptions import MidStreamFallbackError, RateLimitError - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4", - "litellm_params": {"model": "gpt-4", "api_key": "fake-key"}, - } - ], - ) - - messages = [{"role": "user", "content": "Test"}] - initial_kwargs = {"model": "gpt-4", "stream": True} - - original_exception = RateLimitError( - message="rate limited", - llm_provider="vertex_ai", - model="gpt-4", - ) - mid_chunk_error = MidStreamFallbackError( - message="rate limited", - model="gpt-4", - llm_provider="openai", - original_exception=original_exception, - generated_content="Hello, I am", - is_pre_first_chunk=False, - ) - - class SyncIteratorMidChunkError: - def __init__(self): - self.model = "gpt-4" - self.custom_llm_provider = "openai" - self.logging_obj = MagicMock() - self.chunks = [] - - def __iter__(self): - return self - - def __next__(self): - raise mid_chunk_error - - mock_response = SyncIteratorMidChunkError() - - result = router._completion_streaming_iterator( - model_response=mock_response, - messages=messages, - initial_kwargs=initial_kwargs, - ) - - with pytest.raises(RateLimitError) as exc_info: - list(result) - assert exc_info.value is original_exception - assert exc_info.value.type == "throttling_error" - assert exc_info.value.code == "429" - - -def test_completion_streaming_iterator_reraises_mid_chunk_error_with_no_text_content(): - """Sync: a reasoning-only chunk sets is_pre_first_chunk=False without populating - generated_content (which only tracks text deltas). The re-raise guard must still - detect this via the raw chunks on the wrapper, or the router silently retries and - the client receives duplicated/inconsistent output.""" - from unittest.mock import MagicMock - - from litellm.exceptions import MidStreamFallbackError - from litellm.types.utils import Delta, StreamingChoices - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4", - "litellm_params": {"model": "gpt-4", "api_key": "fake-key"}, - } - ], - ) - - messages = [{"role": "user", "content": "Test"}] - initial_kwargs = {"model": "gpt-4", "stream": True} - - mid_chunk_error = MidStreamFallbackError( - message="Connection reset", - model="gpt-4", - llm_provider="openai", - generated_content="", - is_pre_first_chunk=False, - ) - - reasoning_chunk = litellm.ModelResponseStream( - id="chatcmpl-partial-1", - model="gpt-4", - object="chat.completion.chunk", - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta(reasoning_content="Thinking about the answer", role="assistant"), - ) - ], - ) - - class SyncIteratorNoTextChunkError: - def __init__(self): - self.model = "gpt-4" - self.custom_llm_provider = "openai" - self.logging_obj = MagicMock() - self.chunks = [reasoning_chunk] - - def __iter__(self): - return self - - def __next__(self): - raise mid_chunk_error - - mock_response = SyncIteratorNoTextChunkError() - - with patch.object(router, "function_with_fallbacks") as mock_fallback: - result = router._completion_streaming_iterator( - model_response=mock_response, - messages=messages, - initial_kwargs=initial_kwargs, - ) - - with pytest.raises(MidStreamFallbackError): - list(result) - - assert not mock_fallback.called, ( - "fallback must not be attempted once any content, text or non-text, has already streamed" - ) - - -@pytest.mark.asyncio -async def test_acompletion_streaming_iterator_pre_first_chunk_skips_continuation(): - """When MidStreamFallbackError has is_pre_first_chunk=True, use original messages.""" - from unittest.mock import MagicMock - - from litellm.exceptions import MidStreamFallbackError - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4", - "litellm_params": {"model": "gpt-4", "api_key": "fake-key"}, - } - ], - ) - - messages = [{"role": "user", "content": "Hello"}] - initial_kwargs = {"model": "gpt-4", "stream": True} - - pre_first_chunk_error = MidStreamFallbackError( - message="429 Resource exhausted", - model="gpt-4", - llm_provider="vertex_ai", - generated_content="", - is_pre_first_chunk=True, - ) - - class AsyncIteratorPreFirstChunkError: - def __init__(self): - self.model = "gpt-4" - self.custom_llm_provider = "openai" - self.logging_obj = MagicMock() - self.chunks = [] - - def __aiter__(self): - return self - - async def __anext__(self): - raise pre_first_chunk_error - - mock_response = AsyncIteratorPreFirstChunkError() - - class EmptyAsyncIterator: - def __aiter__(self): - return self - - async def __anext__(self): - raise StopAsyncIteration - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - return_value=EmptyAsyncIterator(), - ) as mock_fallback_utils: - iterator = await router._acompletion_streaming_iterator( - model_response=mock_response, - messages=messages, - initial_kwargs=initial_kwargs, - ) - async for _ in iterator: - pass - - assert mock_fallback_utils.called - fallback_kwargs = mock_fallback_utils.call_args.kwargs["kwargs"] - # Pre-first-chunk: should use original messages, no continuation prompt - assert fallback_kwargs["messages"] == messages - - -@pytest.mark.asyncio -async def test_acompletion_streaming_iterator_reraises_mid_chunk_error_with_no_text_content(): - """Async: a reasoning-only chunk sets is_pre_first_chunk=False without populating - generated_content (which only tracks text deltas). The re-raise guard must still - detect this via the raw chunks on the wrapper, or the router silently retries and - the client receives duplicated/inconsistent output.""" - from unittest.mock import MagicMock - - from litellm.exceptions import MidStreamFallbackError - from litellm.types.utils import Delta, StreamingChoices - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4", - "litellm_params": {"model": "gpt-4", "api_key": "fake-key"}, - } - ], - ) - - messages = [{"role": "user", "content": "Test"}] - initial_kwargs = {"model": "gpt-4", "stream": True} - - mid_chunk_error = MidStreamFallbackError( - message="Connection reset", - model="gpt-4", - llm_provider="openai", - generated_content="", - is_pre_first_chunk=False, - ) - - reasoning_chunk = litellm.ModelResponseStream( - id="chatcmpl-partial-1", - model="gpt-4", - object="chat.completion.chunk", - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta(reasoning_content="Thinking about the answer", role="assistant"), - ) - ], - ) - - class AsyncIteratorNoTextChunkError: - def __init__(self): - self.model = "gpt-4" - self.custom_llm_provider = "openai" - self.logging_obj = MagicMock() - self.chunks = [reasoning_chunk] - - def __aiter__(self): - return self - - async def __anext__(self): - raise mid_chunk_error - - mock_response = AsyncIteratorNoTextChunkError() - - with patch.object(router, "async_function_with_fallbacks_common_utils") as mock_fallback_utils: - iterator = await router._acompletion_streaming_iterator( - model_response=mock_response, - messages=messages, - initial_kwargs=initial_kwargs, - ) - - with pytest.raises(MidStreamFallbackError): - async for _ in iterator: - pass - - assert not mock_fallback_utils.called, ( - "fallback must not be attempted once any content, text or non-text, has already streamed" - ) - - -# --------------------------------------------------------------------------- -# Shared helpers for the _aresponses_streaming_iterator test suite. -# --------------------------------------------------------------------------- -def _make_responses_iterator( - *, - chunks=(), - error=None, - bridge=False, - model="gpt-4", - hidden_params=None, - chat_chunks=None, -): - """Build a minimal mock Responses-API streaming iterator. - - Bypasses BaseResponsesAPIStreamingIterator.__init__ but mirrors every - attribute production code reads. Yields *chunks*, then raises *error* - (or StopAsyncIteration). Set bridge=True to inherit from - LiteLLMCompletionStreamingIterator so the wrapper's bridge-path - isinstance check (used by usage extraction) matches. - """ - from litellm.responses.litellm_completion_transformation.streaming_iterator import ( - LiteLLMCompletionStreamingIterator, - ) - from litellm.responses.streaming_iterator import ( - BaseResponsesAPIStreamingIterator, - ) - - base = ( - LiteLLMCompletionStreamingIterator - if bridge - else BaseResponsesAPIStreamingIterator - ) - - class _Iter(base): - def __init__(self): - self._chunks = list(chunks) - self._idx = 0 - self._hidden_params = hidden_params or {} - self.model = model - self.custom_llm_provider = "anthropic" - self.logging_obj = MagicMock() - self.litellm_metadata = None - self.responses_api_provider_config = None - self.finished = False - self.completed_response = None - self.response = None - self.start_time = None - self.request_data = {} - self.call_type = None - if chat_chunks is not None: - self.collected_chat_completion_chunks = chat_chunks - - def __aiter__(self): - return self - - async def __anext__(self): - if self._idx < len(self._chunks): - self._idx += 1 - return self._chunks[self._idx - 1] - if error is not None: - raise error - raise StopAsyncIteration - - return _Iter() - - -class _AsyncList: - """Generic async iterator over a list — used as the fallback response.""" - - def __init__(self, items=()): - self._items = list(items) - self._idx = 0 - - def __aiter__(self): - return self - - async def __anext__(self): - if self._idx >= len(self._items): - raise StopAsyncIteration - item = self._items[self._idx] - self._idx += 1 - return item - - -def _make_router_with_fallback(primary="gpt-4", secondary="gpt-3.5-turbo"): - return litellm.Router( - model_list=[ - { - "model_name": primary, - "litellm_params": {"model": primary, "api_key": "k1"}, - }, - { - "model_name": secondary, - "litellm_params": {"model": secondary, "api_key": "k2"}, - }, - ], - fallbacks=[{primary: [secondary]}], - ) - - -@pytest.mark.asyncio -async def test_aresponses_streaming_iterator_fallback(): - """Catches MidStreamFallbackError, re-enters the fallback chain via - async_function_with_fallbacks_common_utils with the per-attempt helper - and original_generic_function preserved. Mirrors - test_acompletion_streaming_iterator for the aresponses path.""" - from litellm.exceptions import MidStreamFallbackError - from litellm.responses.streaming_iterator import ( - BaseResponsesAPIStreamingIterator, - ) - - router = _make_router_with_fallback( - "anthropic/claude-sonnet-4-6", "vertex_ai/claude-sonnet-4-6" - ) - src = _make_responses_iterator( - chunks=[MagicMock(type="response.created")], - error=MidStreamFallbackError( - message="anthropic socket timeout", - model="anthropic/claude-sonnet-4-6", - llm_provider="anthropic", - is_pre_first_chunk=False, - generated_content="", - ), - model="anthropic/claude-sonnet-4-6", - hidden_params={"model_id": "src-deployment-1"}, - ) - fallback_chunks = [ - MagicMock(type="response.output_text.delta"), - MagicMock(type="response.completed"), - ] - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - return_value=_AsyncList(fallback_chunks), - ) as mock_fallback_utils: - wrapped = await router._aresponses_streaming_iterator( - response=src, - initial_kwargs={ - "model": "anthropic/claude-sonnet-4-6", - "stream": True, - "input": "Hi", - "original_generic_function": litellm.aresponses, - }, - ) - assert isinstance(wrapped, BaseResponsesAPIStreamingIterator) - assert wrapped._hidden_params.get("model_id") == "src-deployment-1" - collected = [c async for c in wrapped] - - assert len(collected) == 3 # 1 primary chunk + 2 fallback chunks - call_kwargs = mock_fallback_utils.call_args.kwargs - fbk = call_kwargs["kwargs"] - # Bound methods compare equal when they share the same instance + __func__. - assert fbk["original_function"] == router._ageneric_api_call_with_fallbacks_helper - assert fbk["original_generic_function"] is litellm.aresponses - assert call_kwargs["model_group"] == "anthropic/claude-sonnet-4-6" - assert call_kwargs["disable_fallbacks"] is False - - -@pytest.mark.asyncio -async def test_aresponses_streaming_iterator_writes_litellm_metadata_on_fallback(): - """Regression: model_group must land under "litellm_metadata" (the key - litellm.aresponses reads), not the default "metadata".""" - from litellm.exceptions import MidStreamFallbackError - - router = _make_router_with_fallback() - src = _make_responses_iterator( - error=MidStreamFallbackError( - message="boom", - model="gpt-4", - llm_provider="anthropic", - is_pre_first_chunk=True, - generated_content="", - ) - ) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - return_value=_AsyncList(), - ) as mock_fallback_utils: - wrapped = await router._aresponses_streaming_iterator( - response=src, - initial_kwargs={ - "model": "gpt-4", - "stream": True, - "input": "Hello", - "original_generic_function": litellm.aresponses, - }, - ) - async for _ in wrapped: - pass - - fbk = mock_fallback_utils.call_args.kwargs["kwargs"] - assert "litellm_metadata" in fbk, "wrong metadata_variable_name" - assert fbk["litellm_metadata"]["model_group"] == "gpt-4" - assert "model_group" not in fbk.get( - "metadata", {} - ), "model_group leaked into 'metadata' instead of 'litellm_metadata'" - - -@pytest.mark.asyncio -async def test_aresponses_streaming_iterator_pre_first_chunk_skips_continuation(): - """Pre-first-chunk error: original input is preserved unchanged.""" - from litellm.exceptions import MidStreamFallbackError - - router = _make_router_with_fallback() - src = _make_responses_iterator( - error=MidStreamFallbackError( - message="socket timeout before first chunk", - model="gpt-4", - llm_provider="anthropic", - is_pre_first_chunk=True, - generated_content="", - ) - ) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - return_value=_AsyncList(), - ) as mock_fallback_utils: - wrapped = await router._aresponses_streaming_iterator( - response=src, - initial_kwargs={ - "model": "gpt-4", - "stream": True, - "input": "Hello", - "original_generic_function": litellm.aresponses, - }, - ) - async for _ in wrapped: - pass - - fbk = mock_fallback_utils.call_args.kwargs["kwargs"] - assert fbk["input"] == "Hello" # original input, no continuation messages - - -@pytest.mark.asyncio -async def test_aresponses_streaming_iterator_partial_content_injects_continuation(): - """Mid-stream error: input is rewritten to include user prompt + - developer instruction + prior assistant message with partial output.""" - from litellm.exceptions import MidStreamFallbackError - - router = _make_router_with_fallback() - src = _make_responses_iterator( - chunks=[MagicMock(type="response.output_text.delta")], - error=MidStreamFallbackError( - message="socket reset mid-stream", - model="gpt-4", - llm_provider="anthropic", - is_pre_first_chunk=False, - generated_content="The capital of France is", - ), - ) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - return_value=_AsyncList(), - ) as mock_fallback_utils: - wrapped = await router._aresponses_streaming_iterator( - response=src, - initial_kwargs={ - "model": "gpt-4", - "stream": True, - "input": "What's the capital of France?", - "original_generic_function": litellm.aresponses, - }, - ) - async for _ in wrapped: - pass - - new_input = mock_fallback_utils.call_args.kwargs["kwargs"]["input"] - assert isinstance(new_input, list) - assert new_input[0]["role"] == "user" - assert new_input[0]["content"][0]["text"] == "What's the capital of France?" - assert new_input[1]["role"] == "developer" - assert "do not repeat" in new_input[1]["content"][0]["text"].lower() - assert new_input[2]["role"] == "assistant" - assert new_input[2]["content"][0]["type"] == "output_text" - assert new_input[2]["content"][0]["text"] == "The capital of France is" - - -@pytest.mark.asyncio -async def test_aresponses_streaming_iterator_combines_partial_usage(): - """Partial usage from the bridge path is normalized to ResponseAPIUsage - and summed onto the fallback's response.completed event — no token-name - split, clean ResponseAPIUsage on output.""" - from types import SimpleNamespace - - from litellm.exceptions import MidStreamFallbackError - from litellm.types.llms.openai import ( - ResponseAPIUsage, - ResponseCompletedEvent, - ResponsesAPIResponse, - ResponsesAPIStreamEvents, - ) - - router = _make_router_with_fallback() - src = _make_responses_iterator( - bridge=True, - chat_chunks=[MagicMock()], - chunks=[MagicMock(type="response.output_text.delta")], - error=MidStreamFallbackError( - message="boom", - model="gpt-4", - llm_provider="anthropic", - is_pre_first_chunk=False, - generated_content="hello", - ), - ) - - fallback_response_object = ResponsesAPIResponse( - id="resp_test", created_at=0, model="gpt-4", object="response", output=[] - ) - fallback_response_object.usage = ResponseAPIUsage( - input_tokens=20, output_tokens=15, total_tokens=35 - ) - fallback_event = ResponseCompletedEvent( - type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED, - response=fallback_response_object, - ) - - with ( - patch( - "litellm.main.stream_chunk_builder", - return_value=SimpleNamespace( - usage=SimpleNamespace(prompt_tokens=10, completion_tokens=4) - ), - ), - patch.object( - router, - "async_function_with_fallbacks_common_utils", - return_value=_AsyncList([fallback_event]), - ), - ): - wrapped = await router._aresponses_streaming_iterator( - response=src, - initial_kwargs={ - "model": "gpt-4", - "stream": True, - "input": "hi", - "original_generic_function": litellm.aresponses, - }, - ) - async for _ in wrapped: - pass - - merged = fallback_response_object.usage - assert isinstance(merged, ResponseAPIUsage) - assert merged.input_tokens == 30 # 10 (translated from prompt_tokens) + 20 - assert merged.output_tokens == 19 # 4 (translated from completion_tokens) + 15 - assert merged.total_tokens == 49 - - -def _midstream_rate_limit_error(): - rate_limit_error = litellm.RateLimitError( - message="vertex_ai_betaException - Resource exhausted.", - model="gemini", - llm_provider="vertex_ai_beta", - ) - midstream_error = MidStreamFallbackError( - message=str(rate_limit_error), - model="gemini", - llm_provider="vertex_ai_beta", - original_exception=rate_limit_error, - is_pre_first_chunk=True, - ) - return rate_limit_error, midstream_error - - -@pytest.mark.asyncio -async def test_acompletion_streaming_iterator_surfaces_rate_limit_without_fallbacks(): - """Regression for #26015: a mid-stream 429 with no fallbacks configured must - surface a clean RateLimitError, not leak the internal MidStreamFallbackError - wrapper to the client, and must terminate instead of hanging.""" - rate_limit_error, midstream_error = _midstream_rate_limit_error() - - router = litellm.Router( - model_list=[ - { - "model_name": "gemini", - "litellm_params": { - "model": "vertex_ai/gemini-2.0-flash", - "api_key": "fake-key", - }, - }, - ], - num_retries=0, - ) - - class _RaisingStream: - def __init__(self): - self.chunks = [] - - def __aiter__(self): - return self - - async def __anext__(self): - raise midstream_error - - stream = _RaisingStream() - setattr(stream, "model", "gemini") - setattr(stream, "custom_llm_provider", "vertex_ai_beta") - setattr(stream, "logging_obj", MagicMock()) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(side_effect=midstream_error), - ): - result = await router._acompletion_streaming_iterator( - model_response=stream, - messages=[{"role": "user", "content": "Hello"}], - initial_kwargs={"model": "gemini", "stream": True}, - ) - - async def _consume(): - async for _ in result: - pass - - with pytest.raises(litellm.RateLimitError) as exc_info: - await asyncio.wait_for(_consume(), timeout=10) - - assert not isinstance(exc_info.value, MidStreamFallbackError) - assert exc_info.value.status_code == 429 - assert exc_info.value is rate_limit_error - - -def test_completion_streaming_iterator_surfaces_rate_limit_without_fallbacks(): - """Sync counterpart of - test_acompletion_streaming_iterator_surfaces_rate_limit_without_fallbacks.""" - rate_limit_error, midstream_error = _midstream_rate_limit_error() - - router = litellm.Router( - model_list=[ - { - "model_name": "gemini", - "litellm_params": { - "model": "vertex_ai/gemini-2.0-flash", - "api_key": "fake-key", - }, - }, - ], - num_retries=0, - ) - - class _RaisingSyncStream: - def __init__(self): - self.model = "gemini" - self.custom_llm_provider = "vertex_ai_beta" - self.logging_obj = MagicMock() - self.chunks = [] - - def __iter__(self): - return self - - def __next__(self): - raise midstream_error - - with patch.object( - router, - "function_with_fallbacks", - side_effect=midstream_error, - ): - result = router._completion_streaming_iterator( - model_response=_RaisingSyncStream(), - messages=[{"role": "user", "content": "Hello"}], - initial_kwargs={"model": "gemini", "stream": True}, - ) - - with pytest.raises(litellm.RateLimitError) as exc_info: - list(result) - - assert not isinstance(exc_info.value, MidStreamFallbackError) - assert exc_info.value.status_code == 429 - assert exc_info.value is rate_limit_error - - -@pytest.mark.asyncio -async def test_aresponses_streaming_iterator_surfaces_rate_limit_without_fallbacks(): - """Responses-API counterpart of - test_acompletion_streaming_iterator_surfaces_rate_limit_without_fallbacks.""" - rate_limit_error, midstream_error = _midstream_rate_limit_error() - - router = litellm.Router( - model_list=[ - { - "model_name": "gemini", - "litellm_params": { - "model": "vertex_ai/gemini-2.0-flash", - "api_key": "fake-key", - }, - }, - ], - num_retries=0, - ) - src = _make_responses_iterator(error=midstream_error, model="gemini") - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(side_effect=midstream_error), - ): - wrapped = await router._aresponses_streaming_iterator( - response=src, - initial_kwargs={ - "model": "gemini", - "stream": True, - "input": "Hello", - "original_generic_function": litellm.aresponses, - }, - ) - - async def _consume(): - async for _ in wrapped: - pass - - with pytest.raises(litellm.RateLimitError) as exc_info: - await asyncio.wait_for(_consume(), timeout=10) - - assert not isinstance(exc_info.value, MidStreamFallbackError) - assert exc_info.value.status_code == 429 - assert exc_info.value is rate_limit_error - - -@pytest.mark.asyncio -async def test_async_function_with_fallbacks_common_utils(): - """Test the async_function_with_fallbacks_common_utils method""" - # Create a basic router for testing - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": { - "model": "gpt-3.5-turbo", - }, - } - ], - max_fallbacks=5, - ) - - # Test case 1: disable_fallbacks=True should raise original exception - test_exception = Exception("Test error") - with pytest.raises(Exception, match="Test error"): - await router.async_function_with_fallbacks_common_utils( - e=test_exception, - disable_fallbacks=True, - fallbacks=None, - context_window_fallbacks=None, - content_policy_fallbacks=None, - model_group="gpt-3.5-turbo", - args=(), - kwargs=MagicMock(), - ) - - # Test case 2: original_model_group=None should raise original exception - with pytest.raises(Exception, match="Test error"): - await router.async_function_with_fallbacks_common_utils( - e=test_exception, - disable_fallbacks=False, - fallbacks=None, - context_window_fallbacks=None, - content_policy_fallbacks=None, - model_group="gpt-3.5-turbo", - args=(), - kwargs={}, # No model key - ) - - -def test_should_include_deployment(): - """Test that Router.should_include_deployment returns the correct response""" - router = litellm.Router( - model_list=[ - { - "model_name": "model_name_a28a12f9-3e44-4861-bd4f-325f2d309ce8_cd5dc6fb-b046-4e05-ae1d-32ba4d936266", - "litellm_params": {"model": "openai/*"}, - "model_info": { - "team_id": "a28a12f9-3e44-4861-bd4f-325f2d309ce8", - "team_public_model_name": "openai/*", - }, - } - ], - ) - - model = { - "model_name": "model_name_a28a12f9-3e44-4861-bd4f-325f2d309ce8_cd5dc6fb-b046-4e05-ae1d-32ba4d936266", - "litellm_params": { - "api_key": "sk-proj-1234567890", - "custom_llm_provider": "openai", - "use_in_pass_through": False, - "use_litellm_proxy": False, - "merge_reasoning_content_in_choices": False, - "model": "openai/*", - }, - "model_info": { - "id": "95f58039-d54a-4d1c-b700-5e32e99a1120", - "db_model": True, - "updated_by": "64a2f787-0863-4d76-9516-2dc49c1598e8", - "created_by": "64a2f787-0863-4d76-9516-2dc49c1598e8", - "team_id": "a28a12f9-3e44-4861-bd4f-325f2d309ce8", - "team_public_model_name": "openai/*", - "mode": "completion", - "access_groups": ["restricted-models-openai"], - }, - } - model_name = "openai/o4-mini-deep-research" - team_id = "a28a12f9-3e44-4861-bd4f-325f2d309ce8" - assert router.get_model_list( - model_name=model_name, - team_id=team_id, - ) - - -def test_pre_call_checks_skips_token_count_without_max_input_tokens(monkeypatch): - """ - tiktoken token counting is the dominant on-loop cost for large prompts. When no - deployment in the group declares max_input_tokens, the count is never consumed, so - _pre_call_checks must not run it at all. - """ - router = litellm.Router( - model_list=[ - {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, - ], - enable_pre_call_checks=True, - ) - monkeypatch.setattr(router, "get_router_model_info", lambda **kwargs: {}) - - calls = [] - monkeypatch.setattr( - litellm, "token_counter", lambda *a, **k: calls.append(1) or 1000 - ) - - deployments = [ - {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, - {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d2"}}, - ] - result = router._pre_call_checks( - model="m", - healthy_deployments=deployments, - messages=[{"role": "user", "content": "hi"}], - ) - - assert calls == [] - assert len(result) == 2 - - -def test_pre_call_checks_counts_once_and_filters_on_max_input_tokens(monkeypatch): - """ - When a deployment declares max_input_tokens the count must still run, be performed - at most once across the group (memoized), and filter deployments whose limit is - exceeded. - """ - router = litellm.Router( - model_list=[ - {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, - ], - enable_pre_call_checks=True, - ) - monkeypatch.setattr( - router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5} - ) - - calls = [] - monkeypatch.setattr( - litellm, "token_counter", lambda *a, **k: calls.append(1) or 1000 - ) - - deployments = [ - {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, - {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d2"}}, - ] - with pytest.raises(litellm.ContextWindowExceededError): - router._pre_call_checks( - model="m", - healthy_deployments=deployments, - messages=[{"role": "user", "content": "hi"}], - ) - - assert calls == [1] - - -def test_pre_call_checks_uses_precounted_tokens(monkeypatch): - """ - An async caller counts off the event loop and passes the result in. _pre_call_checks - must filter on that count instead of re-counting on the loop. - """ - router = litellm.Router( - model_list=[ - {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, - ], - enable_pre_call_checks=True, - ) - monkeypatch.setattr( - router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5} - ) - - calls = [] - monkeypatch.setattr( - litellm, "token_counter", lambda *a, **k: calls.append(1) or 1 - ) - - deployments = [ - {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, - ] - with pytest.raises(litellm.ContextWindowExceededError): - router._pre_call_checks( - model="m", - healthy_deployments=deployments, - messages=[{"role": "user", "content": "hi"}], - input_token_count=1000, - ) - - assert calls == [] - - -async def test_async_get_healthy_deployments_counts_tokens_off_the_event_loop(monkeypatch): - """ - The async deployment path must hand _pre_call_checks a count taken in a worker thread, - so a multi-MB prompt never blocks the proxy during deployment selection. - """ - router = litellm.Router( - model_list=[ - {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, - ], - enable_pre_call_checks=True, - ) - monkeypatch.setattr( - router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 1_000_000} - ) - - counting_threads = [] - monkeypatch.setattr( - litellm, - "token_counter", - lambda *a, **k: counting_threads.append(threading.current_thread()) or 42, - ) - - counts_passed_in = [] - original_pre_call_checks = router._pre_call_checks - - def spy(**kwargs): - counts_passed_in.append(kwargs.get("input_token_count")) - return original_pre_call_checks(**kwargs) - - monkeypatch.setattr(router, "_pre_call_checks", spy) - - result = await router.async_get_healthy_deployments( - model="m", - request_kwargs={}, - messages=[{"role": "user", "content": "hi"}], - input=None, - specific_deployment=False, - parent_otel_span=None, - ) - - assert len(result) == 1 - assert counts_passed_in == [42] - assert len(counting_threads) == 1 - assert counting_threads[0] is not threading.current_thread() - - -@pytest.mark.parametrize( - "model_info,expected", - [ - ({"max_input_tokens": 100}, True), - ({"max_input_tokens": None}, False), - ({}, False), - ], -) -def test_pre_call_checks_need_token_count(monkeypatch, model_info, expected): - """Only a deployment that declares an integer context window makes a token count worth taking.""" - router = litellm.Router( - model_list=[ - {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, - ], - enable_pre_call_checks=True, - ) - monkeypatch.setattr(router, "get_router_model_info", lambda **kwargs: model_info) - - deployments = [ - {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, - ] - assert router._pre_call_checks_need_token_count("m", deployments) is expected - - -def test_deployment_max_input_tokens_survives_an_unmappable_deployment(monkeypatch): - """ - _pre_call_checks skips a deployment it cannot resolve and carries on. The off-loop - pre-count must do the same, or an unmapped first deployment hides the limit declared by - a later one and the count lands back on the event loop. - """ - router = litellm.Router( - model_list=[ - {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, - ], - enable_pre_call_checks=True, - ) - - def flaky_model_info(deployment, received_model_name, id=None): - if deployment["model_info"]["id"] == "unmapped": - raise ValueError("This model isn't mapped yet.") - return {"max_input_tokens": 100} - - monkeypatch.setattr(router, "get_router_model_info", flaky_model_info) - - unmapped = {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "unmapped"}} - mapped = {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "mapped"}} - - assert router._deployment_max_input_tokens("m", unmapped) is None - assert router._deployment_max_input_tokens("m", mapped) == 100 - assert router._pre_call_checks_need_token_count("m", [unmapped, mapped]) is True - - -def test_pre_call_checks_does_not_recount_inline_after_an_off_loop_failure(monkeypatch): - """ - When the off-loop count failed there is nothing left to filter on, so _pre_call_checks must - return the deployments unfiltered rather than repeating the count on the event loop. - """ - router = litellm.Router( - model_list=[ - {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, - ], - enable_pre_call_checks=True, - ) - monkeypatch.setattr( - router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5} - ) - - calls = [] - monkeypatch.setattr( - litellm, "token_counter", lambda *a, **k: calls.append(1) or 1000 - ) - - deployments = [ - {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, - ] - result = router._pre_call_checks( - model="m", - healthy_deployments=deployments, - messages=[{"role": "user", "content": "hi"}], - input_token_count=None, - skip_inline_token_count=True, - ) - - assert calls == [] - assert len(result) == 1 - - -async def test_async_get_healthy_deployments_never_recounts_on_the_loop(monkeypatch): - """ - An off-loop count that raises must not send the same work back onto the event loop through - _pre_call_checks' inline fallback. - """ - router = litellm.Router( - model_list=[ - {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, - ], - enable_pre_call_checks=True, - ) - monkeypatch.setattr( - router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5} - ) - - counting_threads = [] - - def exploding_counter(*args, **kwargs): - counting_threads.append(threading.current_thread()) - raise ValueError("Invalid content item type: image") - - monkeypatch.setattr(litellm, "token_counter", exploding_counter) - - result = await router.async_get_healthy_deployments( - model="m", - request_kwargs={}, - messages=[{"role": "user", "content": "hi"}], - input=None, - specific_deployment=False, - parent_otel_span=None, - ) - - assert len(result) == 1 - assert len(counting_threads) == 1 - assert counting_threads[0] is not threading.current_thread() - - -async def test_acount_pre_call_check_tokens_leaves_the_event_loop_free(monkeypatch): - """ - A multi-MB prompt must not stall the proxy: a competing coroutine has to get - scheduled while the router's context-window count is in flight. - """ - router = litellm.Router( - model_list=[ - {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, - ], - enable_pre_call_checks=True, - ) - monkeypatch.setattr( - router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5} - ) - - deployments = [ - {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, - ] - ran = [] - - async def competitor(): - ran.append("competitor") - - task = asyncio.create_task(competitor()) - count = await router._acount_pre_call_check_tokens( - model="m", - healthy_deployments=deployments, - messages=[{"role": "user", "content": "A" * 512 * 1024}], - input=None, - request_kwargs=None, - ) - ran.append("count") - await task - - assert count is not None and count > 0 - assert ran == ["competitor", "count"] - - -async def test_acount_pre_call_check_tokens_skips_without_max_input_tokens(monkeypatch): - """No deployment limits its context window, so there is nothing to count.""" - router = litellm.Router( - model_list=[ - {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, - ], - enable_pre_call_checks=True, - ) - monkeypatch.setattr(router, "get_router_model_info", lambda **kwargs: {}) - - calls = [] - monkeypatch.setattr( - litellm, "token_counter", lambda *a, **k: calls.append(1) or 1000 - ) - - count = await router._acount_pre_call_check_tokens( - model="m", - healthy_deployments=[ - {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, - ], - messages=[{"role": "user", "content": "hi"}], - input=None, - request_kwargs=None, - ) - - assert count is None - assert calls == [] - - -def test_pre_call_checks_counts_tokens_from_responses_input_string(monkeypatch): - """ - Responses API calls pass `input` (str) instead of `messages`. Context-window - checks must count tokens from `input` and filter deployments over the limit. Uses - the real token_counter so the transform + counting path is a true regression guard. - """ - router = litellm.Router( - model_list=[ - {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, - ], - enable_pre_call_checks=True, - ) - monkeypatch.setattr( - router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 1} - ) - - deployments = [ - {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, - ] - with pytest.raises(litellm.ContextWindowExceededError): - router._pre_call_checks( - model="m", - healthy_deployments=deployments, - input="a very long prompt that exceeds the tiny context window", - ) - - -def test_pre_call_checks_counts_tokens_from_responses_input_list(monkeypatch): - """ - Responses API `input` can be a list of input items. It must be normalized to - chat messages and counted so oversized requests are filtered out. Uses the real - token_counter (no mock) so the transform + counting path is a true regression guard. - """ - router = litellm.Router( - model_list=[ - {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, - ], - enable_pre_call_checks=True, - ) - monkeypatch.setattr( - router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 1} - ) - - deployments = [ - {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, - ] - with pytest.raises(litellm.ContextWindowExceededError): - router._pre_call_checks( - model="m", - healthy_deployments=deployments, - input=[ - {"role": "user", "content": "count these tokens against the one token limit please"}, - ], - ) - - -def test_pre_call_checks_counts_responses_instructions_tokens(monkeypatch): - """ - Responses API `instructions` become a system message the model receives, so their - tokens must be counted too. A request whose `input` alone fits under the limit but - whose `input` + `instructions` exceeds it must be filtered (regression for the - context-window check under-filtering when instructions were ignored). - """ - router = litellm.Router( - model_list=[ - {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, - ], - enable_pre_call_checks=True, - ) - - deployments = [ - {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, - ] - - short_input = "hi" - long_instructions = "you are a helpful assistant. " * 20 - - input_only_tokens = router._count_pre_call_check_tokens(messages=None, input=short_input) - with_instructions_tokens = router._count_pre_call_check_tokens( - messages=None, input=short_input, request_kwargs={"instructions": long_instructions} - ) - assert with_instructions_tokens > input_only_tokens - - monkeypatch.setattr( - router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": input_only_tokens} - ) - with pytest.raises(litellm.ContextWindowExceededError): - router._pre_call_checks( - model="m", - healthy_deployments=deployments, - input=short_input, - request_kwargs={"instructions": long_instructions}, - ) - - -_OVERSIZED_TOOL_DESCRIPTION = "look up the answer in the knowledge base. " * 40 - - -@pytest.mark.parametrize( - "prompt_kwargs, tool", - [ - pytest.param( - {"messages": [{"role": "user", "content": "hi"}]}, - { - "type": "function", - "function": { - "name": "lookup", - "description": _OVERSIZED_TOOL_DESCRIPTION, - "parameters": {"type": "object", "properties": {"q": {"type": "string"}}}, - }, - }, - id="chat_completions_tool", - ), - pytest.param( - {"input": "hi"}, - { - "type": "function", - "name": "lookup", - "description": _OVERSIZED_TOOL_DESCRIPTION, - "parameters": {"type": "object", "properties": {"q": {"type": "string"}}}, - }, - id="responses_tool", - ), - pytest.param( - {"messages": [{"role": "user", "content": "hi"}]}, - { - "name": "lookup", - "description": _OVERSIZED_TOOL_DESCRIPTION, - "input_schema": {"type": "object", "properties": {"q": {"type": "string"}}}, - }, - id="anthropic_messages_tool", - ), - ], -) -def test_pre_call_checks_counts_tool_definition_tokens(monkeypatch, prompt_kwargs, tool): - """ - Tool definitions are sent to the model as prompt tokens but never appear in - `messages` or `input`. A request whose prompt alone fits the context window but - whose prompt plus `tools` exceeds it must be rejected before dispatch, for the - Chat Completions, Responses and Anthropic Messages tool shapes alike. - """ - router = litellm.Router( - model_list=[ - {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, - ], - enable_pre_call_checks=True, - ) - deployments = [ - {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, - ] - - prompt_only_tokens = router._count_pre_call_check_tokens( - messages=prompt_kwargs.get("messages"), input=prompt_kwargs.get("input") - ) - monkeypatch.setattr( - router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": prompt_only_tokens} - ) - - assert len(router._pre_call_checks(model="m", healthy_deployments=deployments, **prompt_kwargs)) == 1 - with pytest.raises(litellm.ContextWindowExceededError): - router._pre_call_checks( - model="m", - healthy_deployments=deployments, - request_kwargs={"tools": [tool]}, - **prompt_kwargs, - ) - - -@pytest.mark.parametrize( - "system", - [ - pytest.param("You are a meticulous assistant. " * 40, id="system_string"), - pytest.param( - [{"type": "text", "text": "You are a meticulous assistant. " * 40}], - id="system_blocks", - ), - ], -) -def test_pre_call_checks_counts_anthropic_system_tokens(monkeypatch, system): - """ - The Anthropic Messages API carries the system prompt as a top-level `system` field, - not as a message. Its tokens reach the model, so a request whose `messages` fit but - whose `messages` plus `system` exceed the context window must be rejected. - """ - router = litellm.Router( - model_list=[ - {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, - ], - enable_pre_call_checks=True, - ) - deployments = [ - {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, - ] - messages = [{"role": "user", "content": "hi"}] - - messages_only_tokens = router._count_pre_call_check_tokens(messages=messages, input=None) - monkeypatch.setattr(router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": messages_only_tokens}) - - assert len(router._pre_call_checks(model="m", healthy_deployments=deployments, messages=messages)) == 1 - with pytest.raises(litellm.ContextWindowExceededError): - router._pre_call_checks( - model="m", - healthy_deployments=deployments, - messages=messages, - request_kwargs={"system": system}, - ) - - -@pytest.mark.asyncio -async def test_aanthropic_messages_enforces_context_window_with_system_and_tools(): - """ - End-to-end router regression for /v1/messages: a request whose only oversized - content lives in the top-level `system` field or in `tools` must trip the pre-call - context-window check instead of being dispatched (the deployment uses mock_response, - so reaching the provider handler would return a response rather than raise). - """ - router = litellm.Router( - model_list=[ - { - "model_name": "small-ctx", - "litellm_params": {"model": "anthropic/claude-3-5-haiku-20241022", "mock_response": "hi"}, - "model_info": {"max_input_tokens": 20}, - } - ], - enable_pre_call_checks=True, - ) - messages = [{"role": "user", "content": "hi"}] - - response = await router.aanthropic_messages(model="small-ctx", messages=messages, max_tokens=5) - assert response is not None - - with pytest.raises(litellm.ContextWindowExceededError): - await router.aanthropic_messages( - model="small-ctx", - messages=messages, - max_tokens=5, - system="You are a meticulous assistant. " * 40, - ) - with pytest.raises(litellm.ContextWindowExceededError): - await router.aanthropic_messages( - model="small-ctx", - messages=messages, - max_tokens=5, - tools=[ - { - "name": "lookup", - "description": _OVERSIZED_TOOL_DESCRIPTION, - "input_schema": {"type": "object", "properties": {"q": {"type": "string"}}}, - } - ], - ) - - -def test_count_pre_call_check_tokens_across_api_surfaces(): - """ - _count_pre_call_check_tokens must count tokens from chat `messages`, a Responses - API string `input`, and a Responses API list `input`, and raise when given neither. - """ - router = litellm.Router( - model_list=[ - {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, - ], - ) - - messages_tokens = router._count_pre_call_check_tokens( - messages=[{"role": "user", "content": "hello world"}], input=None - ) - string_input_tokens = router._count_pre_call_check_tokens(messages=None, input="hello world") - list_input_tokens = router._count_pre_call_check_tokens( - messages=None, input=[{"role": "user", "content": "hello world"}] - ) - - assert messages_tokens > 0 - assert string_input_tokens > 0 - assert list_input_tokens > 0 - - with pytest.raises(ValueError, match='Either messages or input must be provided to count tokens'): - router._count_pre_call_check_tokens(messages=None, input=None) - - -def test_pre_call_checks_no_messages_or_input_does_not_crash(monkeypatch): - """ - When neither messages nor input is provided (e.g. endpoints without prompt text), - token counting is skipped gracefully and all deployments are returned. - """ - router = litellm.Router( - model_list=[ - {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, - ], - enable_pre_call_checks=True, - ) - monkeypatch.setattr( - router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5} - ) - - counted: list[dict] = [] - original = router._count_pre_call_check_tokens - monkeypatch.setattr( - router, - "_count_pre_call_check_tokens", - lambda **kwargs: counted.append(kwargs) or original(**kwargs), - ) - - deployments = [ - {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, - ] - result = router._pre_call_checks(model="m", healthy_deployments=deployments) - assert len(result) == 1 - assert counted == [] # token counting skipped entirely, so no misleading error is logged - - -@pytest.mark.asyncio -async def test_aresponses_enforces_context_window_pre_call_check(): - """ - End-to-end router regression: a Responses API call whose `input` exceeds the - deployment's max_input_tokens must be filtered by the pre-call check, raising - ContextWindowExceededError instead of being silently routed. This guards the - wiring that forwards `input` from the generic-call path into deployment selection - (the deployment uses mock_response, so the check must trip before any real call). - """ - router = litellm.Router( - model_list=[ - { - "model_name": "small-ctx", - "litellm_params": {"model": "gpt-3.5-turbo", "mock_response": "hi"}, - "model_info": {"max_input_tokens": 5}, - } - ], - enable_pre_call_checks=True, - ) - with pytest.raises(litellm.ContextWindowExceededError): - await router.aresponses( - model="small-ctx", - input="this responses input is definitely much longer than five tokens for sure", - ) - - -def test_get_deployment_model_info_base_model_flow(): - """Test that get_deployment_model_info correctly handles the base model flow""" - from unittest.mock import patch - - router = litellm.Router( - model_list=[ - { - "model_name": "test-model", - "litellm_params": {"model": "gpt-3.5-turbo"}, - } - ], - ) - - # Mock data for the test - mock_custom_model_info = { - "base_model": "gpt-3.5-turbo", - "input_cost_per_token": 0.001, - "output_cost_per_token": 0.002, - "custom_field": "custom_value", - } - - mock_base_model_info = { - "key": "gpt-3.5-turbo", - "max_tokens": 4096, - "max_input_tokens": 4096, - "max_output_tokens": 4096, - "input_cost_per_token": 0.0015, # This should be overridden by custom model info - "output_cost_per_token": 0.002, - "litellm_provider": "openai", - "mode": "chat", - "supported_openai_params": ["temperature", "max_tokens"], - } - - mock_litellm_model_name_info = { - "key": "test-model", - "max_tokens": 2048, - "max_input_tokens": 2048, - "max_output_tokens": 2048, - "input_cost_per_token": 0.0005, - "output_cost_per_token": 0.001, - "litellm_provider": "test_provider", - "mode": "completion", - "supported_openai_params": ["temperature"], - } - - # Test Case 1: Base model flow with custom model info that has base_model - with patch.object( - litellm, "model_cost", {"test-custom-model": mock_custom_model_info} - ): - with patch.object(litellm, "get_model_info") as mock_get_model_info: - # Configure mock returns - mock_get_model_info.side_effect = lambda model: { - "gpt-3.5-turbo": mock_base_model_info, - "test-model": mock_litellm_model_name_info, - }.get(model) - - result = router.get_deployment_model_info( - model_id="test-custom-model", model_name="test-model" - ) - - # Verify that get_model_info was called for both base model and model name - assert mock_get_model_info.call_count == 2 - mock_get_model_info.assert_any_call( - model="gpt-3.5-turbo" - ) # base model call - mock_get_model_info.assert_any_call(model="test-model") # model name call - - # Verify the result contains merged information - assert result is not None - - # Test the correct merging behavior after fix: - # 1. base_model_info provides defaults, custom_model_info overrides (correct priority) - # 2. The result of step 1 gets merged into litellm_model_name_info (custom+base override litellm) - - # Fields from custom model (should override base model values) - assert ( - result["input_cost_per_token"] == 0.001 - ) # From custom model (overrides base 0.0015) - assert ( - result["output_cost_per_token"] == 0.002 - ) # From custom model (same as base) - assert result["custom_field"] == "custom_value" # From custom model - - # Fields from base model that weren't overridden by custom - assert result["max_tokens"] == 4096 # From base model - assert result["litellm_provider"] == "openai" # From base model - assert ( - result["mode"] == "chat" - ) # From base model (overrides litellm "completion") - - # The key field comes from base model since both base and litellm have it - # and base model info overrides litellm model name info in final merge - assert ( - result["key"] == "gpt-3.5-turbo" - ) # From base model (overrides litellm key) - - # Test Case 2: Custom model info without base_model - mock_custom_model_info_no_base = { - "input_cost_per_token": 0.001, - "output_cost_per_token": 0.002, - "custom_field": "custom_value", - } - - with patch.object( - litellm, - "model_cost", - {"test-custom-model-no-base": mock_custom_model_info_no_base}, - ): - with patch.object(litellm, "get_model_info") as mock_get_model_info: - mock_get_model_info.side_effect = lambda model: { - "test-model": mock_litellm_model_name_info, - }.get(model) - - result = router.get_deployment_model_info( - model_id="test-custom-model-no-base", model_name="test-model" - ) - - # Should only call get_model_info once for model name (no base model) - assert mock_get_model_info.call_count == 1 - mock_get_model_info.assert_called_with(model="test-model") - - # Verify the result contains merged information - assert result is not None - assert result["input_cost_per_token"] == 0.001 # From custom model - assert result["max_tokens"] == 2048 # From litellm model name info - assert result["custom_field"] == "custom_value" # From custom model - assert result["mode"] == "completion" # From litellm model name info - - # Test Case 3: No custom model info, only litellm model name info - with patch.object(litellm, "model_cost", {}): # Empty model cost - with patch.object(litellm, "get_model_info") as mock_get_model_info: - mock_get_model_info.side_effect = lambda model: { - "test-model": mock_litellm_model_name_info, - }.get(model) - - result = router.get_deployment_model_info( - model_id="non-existent-model", model_name="test-model" - ) - - # Should only call get_model_info once for model name - assert mock_get_model_info.call_count == 1 - mock_get_model_info.assert_called_with(model="test-model") - - # Result should be just the litellm model name info - assert result is not None - assert result == mock_litellm_model_name_info - - # Test Case 4: Base model info retrieval fails (exception handling) - mock_custom_model_info_invalid_base = { - "base_model": "invalid-base-model", - "input_cost_per_token": 0.001, - "output_cost_per_token": 0.002, - } - - with patch.object( - litellm, - "model_cost", - {"test-custom-model-invalid": mock_custom_model_info_invalid_base}, - ): - with patch.object(litellm, "get_model_info") as mock_get_model_info: - # Mock get_model_info to raise exception for invalid base model - def mock_get_model_info_side_effect(model): - if model == "invalid-base-model": - raise Exception("Model not found") - elif model == "test-model": - return mock_litellm_model_name_info - return None - - mock_get_model_info.side_effect = mock_get_model_info_side_effect - - result = router.get_deployment_model_info( - model_id="test-custom-model-invalid", model_name="test-model" - ) - - # Should handle exception gracefully and still return merged result - assert result is not None - assert result["input_cost_per_token"] == 0.001 # From custom model - assert result["mode"] == "completion" # From litellm model name info - - # Test Case 5: Both model_cost.get() and get_model_info() return None - with patch.object(litellm, "model_cost", {}): - with patch.object( - litellm, "get_model_info", side_effect=Exception("Not found") - ): - result = router.get_deployment_model_info( - model_id="non-existent", model_name="non-existent" - ) - - # Should return None when no model info is found - assert result is None - - # Test Case 6: custom_model_info present but litellm_model_name_model_info is None - # (model has custom pricing in config but is not in built-in model_prices_and_context_window.json) - mock_custom_pricing_only = { - "input_cost_per_token": 1.74e-06, - "output_cost_per_token": 3.48e-06, - "cache_read_input_token_cost": 1.45e-08, - "mode": "chat", - } - - with patch.object( - litellm, - "model_cost", - {"custom-model-id": mock_custom_pricing_only}, - ): - with patch.object(litellm, "get_model_info") as mock_get_model_info: - # Model NOT in built-in cost map — raise exception - mock_get_model_info.side_effect = Exception("Model not in cost map") - - result = router.get_deployment_model_info( - model_id="custom-model-id", model_name="unknown-model" - ) - - # Should return custom_model_info even when litellm_model_name_model_info is None - assert result is not None - assert result["input_cost_per_token"] == 1.74e-06 - assert result["output_cost_per_token"] == 3.48e-06 - assert result["cache_read_input_token_cost"] == 1.45e-08 - assert result["mode"] == "chat" - - # Test Case 7: custom_model_info with base_model but litellm_model_name_model_info None - mock_custom_with_base = { - "base_model": "some-base-model", - "input_cost_per_token": 0.01, - "output_cost_per_token": 0.02, - } - mock_base_info = { - "key": "some-base-model", - "max_tokens": 8192, - "mode": "chat", - "litellm_provider": "openai", - } - - with patch.object( - litellm, - "model_cost", - {"custom-with-base": mock_custom_with_base}, - ): - with patch.object(litellm, "get_model_info") as mock_get_model_info: - - def get_info_side_effect(model): - if model == "some-base-model": - return mock_base_info - raise Exception("Model not in cost map") - - mock_get_model_info.side_effect = get_info_side_effect - - result = router.get_deployment_model_info( - model_id="custom-with-base", model_name="unknown-model" - ) - - # Should return custom_model_info merged with base model info - assert result is not None - assert ( - result["input_cost_per_token"] == 0.01 - ) # From custom (overrides base) - assert result["max_tokens"] == 8192 # From base model - assert result["litellm_provider"] == "openai" # From base model - - print("✓ All base model flow test cases passed!") - - -@patch("litellm.model_cost", {}) -def test_get_deployment_model_info_base_model_merge_priority(): - """Test that base model info merging respects the correct priority order""" - from unittest.mock import patch - - router = litellm.Router( - model_list=[ - { - "model_name": "test-model", - "litellm_params": {"model": "gpt-3.5-turbo"}, - } - ], - ) - - # Test data with overlapping fields to test merge priority - mock_custom_model_info = { - "base_model": "gpt-4", - "input_cost_per_token": 0.01, # Should override base model value - "max_tokens": 8000, # Should override base model value - "custom_only_field": "custom_value", - } - - mock_base_model_info = { - "key": "gpt-4", - "max_tokens": 4096, # Should be overridden by custom model - "input_cost_per_token": 0.03, # Should be overridden by custom model - "output_cost_per_token": 0.06, # Should be preserved (not in custom) - "litellm_provider": "openai", - "base_only_field": "base_value", - } - - mock_litellm_model_name_info = { - "key": "test-model", - "max_tokens": 2048, # Should be overridden by final custom model info - "input_cost_per_token": 0.005, # Should be overridden by final custom model info - "output_cost_per_token": 0.01, # Should be overridden by final custom model info - "mode": "completion", - "litellm_only_field": "litellm_value", - } - - with patch.object( - litellm, "model_cost", {"custom-model-id": mock_custom_model_info} - ): - with patch.object(litellm, "get_model_info") as mock_get_model_info: - mock_get_model_info.side_effect = lambda model: { - "gpt-4": mock_base_model_info, - "test-model": mock_litellm_model_name_info, - }.get(model) - - result = router.get_deployment_model_info( - model_id="custom-model-id", model_name="test-model" - ) - - assert result is not None - - # Test correct merge priority after fix: - # 1. base_model_info provides defaults - # 2. custom_model_info overrides base_model_info - # 3. Result from steps 1-2 overrides litellm_model_name_info - - # Fields that should come from custom model info (highest priority) - assert ( - result["input_cost_per_token"] == 0.01 - ) # From custom model (overrides base 0.03) - assert ( - result["max_tokens"] == 8000 - ) # From custom model (overrides base 4096) - assert result["custom_only_field"] == "custom_value" # From custom model - - # Fields that should come from base model (not overridden by custom) - assert ( - result["output_cost_per_token"] == 0.06 - ) # From base model (not in custom) - assert ( - result["litellm_provider"] == "openai" - ) # From base model (not in custom) - assert ( - result["base_only_field"] == "base_value" - ) # From base model (not in custom) - - # Fields that should come from litellm model name info (not overridden by custom+base) - assert ( - result["mode"] == "completion" - ) # From litellm model name info (not in custom or base) - assert ( - result["litellm_only_field"] == "litellm_value" - ) # From litellm model name info (not in custom or base) - - # Key comes from base model since both base and litellm have key fields - # and the merged custom+base overrides litellm in the final merge - assert result["key"] == "gpt-4" - - print("✓ Base model merge priority test passed!") - - -def test_add_deployment_model_to_endpoint_for_llm_passthrough_route(): - """ - Test that _add_deployment_model_to_endpoint_for_llm_passthrough_route correctly strips bedrock provider prefix - """ - router = litellm.Router( - model_list=[ - { - "model_name": "special-bedrock-model", - "litellm_params": { - "model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", - }, - } - ], - ) - - # Test Case 1: Bedrock model with provider prefix - should strip "bedrock/" prefix - kwargs = { - "endpoint": "/model/special-bedrock-model/invoke", - "custom_llm_provider": "bedrock", - } - result = router._add_deployment_model_to_endpoint_for_llm_passthrough_route( - kwargs=kwargs, - model="special-bedrock-model", - model_name="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", - ) - assert ( - result["endpoint"] - == "/model/us.anthropic.claude-haiku-4-5-20251001-v1:0/invoke" - ), f"Expected '/model/us.anthropic.claude-haiku-4-5-20251001-v1:0/invoke', got '{result['endpoint']}'" - - # Test Case 2: Bedrock invoke-with-response-stream endpoint - kwargs = { - "endpoint": "/model/special-bedrock-model/invoke-with-response-stream", - "custom_llm_provider": "bedrock", - } - result = router._add_deployment_model_to_endpoint_for_llm_passthrough_route( - kwargs=kwargs, - model="special-bedrock-model", - model_name="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", - ) - assert ( - result["endpoint"] - == "/model/us.anthropic.claude-haiku-4-5-20251001-v1:0/invoke-with-response-stream" - ), f"Expected streaming endpoint with stripped prefix, got '{result['endpoint']}'" - - # Test Case 3: Bedrock converse endpoint - kwargs = { - "endpoint": "/model/bedrock-model/converse", - "custom_llm_provider": "bedrock", - } - result = router._add_deployment_model_to_endpoint_for_llm_passthrough_route( - kwargs=kwargs, - model="bedrock-model", - model_name="bedrock/us.meta.llama3-8b-instruct-v1:0", - ) - assert ( - result["endpoint"] == "/model/us.meta.llama3-8b-instruct-v1:0/converse" - ), f"Expected '/model/us.meta.llama3-8b-instruct-v1:0/converse', got '{result['endpoint']}'" - - # Test Case 4: Bedrock provider prefix auto-detected from model_name - kwargs = { - "endpoint": "/model/router-model/invoke", - } - result = router._add_deployment_model_to_endpoint_for_llm_passthrough_route( - kwargs=kwargs, - model="router-model", - model_name="bedrock/us.meta.llama3-8b-instruct-v1:0", - ) - assert ( - result["endpoint"] == "/model/us.meta.llama3-8b-instruct-v1:0/invoke" - ), f"Expected '/model/us.meta.llama3-8b-instruct-v1:0/invoke', got '{result['endpoint']}'" - - -def test_update_kwargs_with_deployment_uses_pass_through_request_timeout(): - router = litellm.Router( - model_list=[ - { - "model_name": "my-bedrock-model", - "litellm_params": { - "model": "bedrock/us.anthropic.claude-opus-4-5-20251101-v1:0", - }, - } - ], - ) - deployment = router.model_list[0] - kwargs: dict = {} - - with patch( - "litellm.proxy.proxy_server.general_settings", - {"pass_through_request_timeout": 6}, - ): - router._update_kwargs_with_deployment( - deployment=deployment, - kwargs=kwargs, - function_name="_ageneric_api_call_with_fallbacks", - ) - - assert kwargs["timeout"] == 6.0 - - -@pytest.mark.asyncio -async def test_router_acompletion_with_unknown_model_and_default_fallback(): - """ - Test that the router successfully uses a default fallback when a completely - unknown model is requested. It should not raise a BadRequestError. - This test verifies the fix for issue #15114. - """ - model_list = [ - { - "model_name": "gpt-4o", # This is the fallback model - "litellm_params": { - "model": "azure/gpt-4o-real", # The actual underlying model name - "api_key": "fake-key", - "api_base": "https://fake-endpoint.openai.azure.com/", - "mock_response": "this is the fallback response", # Mocked response to prevent real API calls - }, - } - ] - - # Initialize the router with a default fallback - router = litellm.Router(model_list=model_list, default_fallbacks=["gpt-4o"]) - - messages = [ - {"role": "user", "content": "This call should succeed by falling back."} - ] - - # Call completion with a model name that is NOT in the model_list - response = await router.acompletion( - model="completely-unknown-model", messages=messages - ) - - # Check that the call did not fail and we received a valid response object. - assert response is not None - - # Check that the content of the response is from the MOCKED fallback model. - assert response.choices[0].message.content == "this is the fallback response" - - # Check that the response object reports the model that was *actually* called. - assert response.model == "gpt-4o-real" - - -@pytest.mark.asyncio -async def test_router_acompletion_with_unknown_model_and_no_fallback(): - """ - Test that the router still raises a BadRequestError for an unknown model - when no default fallbacks are configured. This ensures we don't break - the original behavior. - """ - model_list = [ - { - "model_name": "gpt-4o", - "litellm_params": { - "model": "azure/gpt-4o-real", - "api_key": "fake-key", - "mock_response": "this should not be called", - }, - } - ] - - # Initialize the router WITHOUT any default fallbacks - router = litellm.Router(model_list=model_list) - - messages = [{"role": "user", "content": "This call should fail."}] - - # Use pytest.raises to assert that a BadRequestError is thrown. - with pytest.raises(litellm.BadRequestError) as excinfo: - await router.acompletion(model="completely-unknown-model", messages=messages) - - # Check that the error message is correct. - # The router returns 'no healthy deployments' because get_model_list returns [] not None. - assert "no healthy deployments for this model" in str(excinfo.value) - - -@pytest.mark.asyncio -async def test_router_unknown_model_error_message_renders_model_name_literally(): - """ - The unknown-model error message renders the caller-supplied model name - verbatim. A name containing Python format-field syntax must be treated as - literal text, not re-interpreted as a format template, which would distort - the message and balloon its length. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4o", - "litellm_params": {"model": "azure/gpt-4o-real", "api_key": "fake-key"}, - } - ] - ) - - weird_model = "ghost{:>200}model" - messages = [{"role": "user", "content": "hi"}] - - with pytest.raises(litellm.BadRequestError) as excinfo: - await router.acompletion(model=weird_model, messages=messages) - - message = str(excinfo.value) - assert weird_model in message - assert " " not in message # no padding run from an expanded format field - - -def test_get_deployment_credentials_with_provider_aws_bedrock_runtime_endpoint(): - """ - Test that get_deployment_credentials_with_provider correctly copies - aws_bedrock_runtime_endpoint from deployment litellm_params to credentials. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "bedrock-claude-model", - "litellm_params": { - "model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", - "aws_access_key_id": "test-access-key", - "aws_secret_access_key": "test-secret-key", - "aws_region_name": "us-east-1", - "aws_bedrock_runtime_endpoint": "https://bedrock-runtime.us-east-1.amazonaws.com", - }, - } - ], - ) - - credentials = router.get_deployment_credentials_with_provider( - model_id="bedrock-claude-model" - ) - - assert credentials is not None - assert ( - credentials["aws_bedrock_runtime_endpoint"] - == "https://bedrock-runtime.us-east-1.amazonaws.com" - ) - assert credentials["aws_access_key_id"] == "test-access-key" - assert credentials["aws_secret_access_key"] == "test-secret-key" - assert credentials["aws_region_name"] == "us-east-1" - assert credentials["custom_llm_provider"] == "bedrock" - - -def test_get_deployment_credentials_with_provider_includes_bucket_name(): - """ - Regression: bucket_name must survive the CredentialLiteLLMParams filter so - managed-files batch retrieval can resolve the GCS/S3 bucket. Previously it was - dropped, causing "GCS bucket_name is required" when fetching batch output files. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "vertex-gemini", - "litellm_params": { - "model": "vertex_ai/gemini-3.5-flash", - "vertex_project": "my-project", - "vertex_location": "global", - "gcs_bucket_name": "my-batch-bucket", - }, - } - ], - ) - - credentials = router.get_deployment_credentials_with_provider( - model_id="vertex-gemini" - ) - - assert credentials is not None - assert credentials["gcs_bucket_name"] == "my-batch-bucket" - assert credentials["vertex_project"] == "my-project" - assert credentials["custom_llm_provider"] == "vertex_ai" - - -def test_get_deployment_credentials_with_provider_resolves_credential_name(): - """ - Test that get_deployment_credentials_with_provider correctly resolves - litellm_credential_name to actual credential values (for UI-created models). - """ - from litellm.types.utils import CredentialItem - - # Setup credential list with a test credential - litellm.credential_list = [ - CredentialItem( - credential_name="test-azure-cred", - credential_info={"custom_llm_provider": "azure"}, - credential_values={ - "api_key": "resolved-api-key", - "api_base": "https://resolved.openai.azure.com", - "api_version": "2024-02-01", - }, - ) - ] - - router = litellm.Router( - model_list=[ - { - "model_name": "azure-gpt-4", - "litellm_params": { - "model": "azure/gpt-4", - "litellm_credential_name": "test-azure-cred", - }, - } - ], - ) - - credentials = router.get_deployment_credentials_with_provider( - model_id="azure-gpt-4" - ) - - assert credentials is not None - assert credentials["api_key"] == "resolved-api-key" - assert credentials["api_base"] == "https://resolved.openai.azure.com" - assert credentials["api_version"] == "2024-02-01" - assert credentials["custom_llm_provider"] == "azure" - # Ensure credential name is removed after resolution - assert "litellm_credential_name" not in credentials - - # Cleanup - litellm.credential_list = [] - - -def test_get_deployment_credentials_with_provider_bedrock_batch_fields(): - """ - Test that get_deployment_credentials_with_provider returns the deployment's - model and the Bedrock batch/S3 fields (s3_region_name, s3_encryption_key_id, - aws_batch_role_arn) instead of silently dropping them (#25104). - """ - router = litellm.Router( - model_list=[ - { - "model_name": "bedrock-batch-model", - "litellm_params": { - "model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", - "aws_region_name": "us-west-2", - "s3_bucket_name": "my-batch-bucket", - "s3_region_name": "us-east-1", - "s3_encryption_key_id": "arn:aws:kms:us-west-2:123:key/abc", - "aws_batch_role_arn": "arn:aws:iam::123:role/batch-role", - }, - } - ], - ) - - credentials = router.get_deployment_credentials_with_provider( - model_id="bedrock-batch-model" - ) - - assert credentials is not None - assert credentials["custom_llm_provider"] == "bedrock" - assert credentials["model"] == "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0" - assert credentials["aws_region_name"] == "us-west-2" - assert credentials["s3_bucket_name"] == "my-batch-bucket" - assert credentials["s3_region_name"] == "us-east-1" - assert credentials["s3_encryption_key_id"] == "arn:aws:kms:us-west-2:123:key/abc" - assert credentials["aws_batch_role_arn"] == "arn:aws:iam::123:role/batch-role" - - -def test_get_deployment_credentials_with_provider_preserves_aws_auth_params(): - """ - Test that get_deployment_credentials_with_provider preserves every AWS auth - selector (session token, assume-role, web identity, profile) so bedrock - files/batches deployments using temporary or role-based credentials do not - silently fall back to the server's ambient identity (#36155). - """ - aws_auth_params = { - "aws_access_key_id": "deployment-access-key", - "aws_secret_access_key": "deployment-secret", - "aws_session_token": "deployment-session-token", - "aws_region_name": "us-west-2", - "aws_session_name": "deployment-session", - "aws_profile_name": "deployment-profile", - "aws_role_name": "arn:aws:iam::123:role/deployment-role", - "aws_web_identity_token": "deployment-web-identity", - "aws_sts_endpoint": "https://sts.us-west-2.amazonaws.com", - "aws_external_id": "deployment-external-id", - } - router = litellm.Router( - model_list=[ - { - "model_name": "bedrock-batch-model", - "litellm_params": { - "model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", - **aws_auth_params, - }, - } - ], - ) - - credentials = router.get_deployment_credentials_with_provider( - model_id="bedrock-batch-model" - ) - - assert credentials is not None - for key, value in aws_auth_params.items(): - assert credentials.get(key) == value, key - - -def _team_wildcard_model(api_key: str, model_id: str = "team-wildcard-id") -> dict: - return { - "model_name": f"model_name_team-1_{model_id}", - "litellm_params": {"model": "openai/*", "api_key": api_key}, - "model_info": { - "id": model_id, - "team_id": "team-1", - "team_public_model_name": "openai/*", - }, - } - - -def test_get_deployment_credentials_with_provider_team_wildcard_priority(): - """ - Regression: a global wildcard pattern (e.g. "openai/*") must not shadow a - team's own wildcard entry. When team_id is provided, the team wildcard - deployment's credentials win; without team_id the global one is used. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "openai/*", - "litellm_params": {"model": "openai/*", "api_key": "global-key"}, - }, - _team_wildcard_model(api_key="team-key"), - ], - ) - - team_credentials = router.get_deployment_credentials_with_provider( - model_id="openai/gpt-5.2", team_id="team-1" - ) - assert team_credentials is not None - assert team_credentials["api_key"] == "team-key" - - global_credentials = router.get_deployment_credentials_with_provider( - model_id="openai/gpt-5.2" - ) - assert global_credentials is not None - assert global_credentials["api_key"] == "global-key" - - -def test_get_deployment_credentials_with_provider_skips_other_team_deployment(): - """ - Regression: a team-scoped deployment sharing a model_name with a global - deployment must never resolve for another team's (or an unscoped) caller, - even when it is indexed first; the shared global deployment wins instead. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "gemini-2.5-pro", - "litellm_params": { - "model": "vertex_ai/gemini-2.5-pro", - "vertex_project": "team-b-project", - }, - "model_info": { - "id": "team-b-vertex", - "team_id": "team-b", - "team_public_model_name": "gemini-2.5-pro", - }, - }, - { - "model_name": "gemini-2.5-pro", - "litellm_params": { - "model": "vertex_ai/gemini-2.5-pro", - "vertex_project": "shared-project", - }, - }, - ], - ) - - other_team_credentials = router.get_deployment_credentials_with_provider( - model_id="gemini-2.5-pro", team_id="team-a" - ) - assert other_team_credentials is not None - assert other_team_credentials["vertex_project"] == "shared-project" - - unscoped_credentials = router.get_deployment_credentials_with_provider( - model_id="gemini-2.5-pro" - ) - assert unscoped_credentials is not None - assert unscoped_credentials["vertex_project"] == "shared-project" - - owner_credentials = router.get_deployment_credentials_with_provider( - model_id="gemini-2.5-pro", team_id="team-b" - ) - assert owner_credentials is not None - assert owner_credentials["vertex_project"] == "team-b-project" - - -def test_get_deployment_credentials_with_provider_no_fallback_to_other_team_only_name(): - """ - When the only deployments under a model name belong to another team, other - callers must get None (env fallback) instead of that team's credentials. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "gemini-2.5-pro", - "litellm_params": { - "model": "vertex_ai/gemini-2.5-pro", - "vertex_project": "team-b-project", - }, - "model_info": { - "id": "team-b-vertex", - "team_id": "team-b", - "team_public_model_name": "gemini-2.5-pro", - }, - }, - ], - ) - - assert ( - router.get_deployment_credentials_with_provider( - model_id="gemini-2.5-pro", team_id="team-a" - ) - is None - ) - assert ( - router.get_deployment_credentials_with_provider(model_id="gemini-2.5-pro") - is None - ) - - -def test_deployment_usable_by_team_helpers(): - """ - Direct coverage of the team-ownership filter: a team-scoped deployment is - usable only by its owning team, shared deployments by anyone, and the - model-group picker returns the first usable deployment or None. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "gemini-2.5-pro", - "litellm_params": { - "model": "vertex_ai/gemini-2.5-pro", - "vertex_project": "team-b-project", - }, - "model_info": { - "id": "team-b-vertex", - "team_id": "team-b", - "team_public_model_name": "gemini-2.5-pro", - }, - }, - { - "model_name": "gemini-2.5-pro", - "litellm_params": { - "model": "vertex_ai/gemini-2.5-pro", - "vertex_project": "shared-project", - }, - }, - ], - ) - - team_owned, shared = router.model_list - assert router._deployment_usable_by_team(team_owned, "team-b") is True - assert router._deployment_usable_by_team(team_owned, "team-a") is False - assert router._deployment_usable_by_team(team_owned, None) is False - assert router._deployment_usable_by_team(shared, "team-a") is True - assert router._deployment_usable_by_team(shared, None) is True - - picked = router._get_model_group_deployment_usable_by_team( - model_group_name="gemini-2.5-pro", team_id="team-a" - ) - assert picked is not None - assert picked.litellm_params.vertex_project == "shared-project" - - owner_picked = router._get_model_group_deployment_usable_by_team( - model_group_name="gemini-2.5-pro", team_id="team-b" - ) - assert owner_picked is not None - assert owner_picked.litellm_params.vertex_project == "team-b-project" - - assert ( - router._get_model_group_deployment_usable_by_team( - model_group_name="unknown-model", team_id="team-a" - ) - is None - ) - - -def test_get_deployment_credentials_with_provider_skips_other_team_wildcard(): - """ - Global wildcard resolution must skip a team-scoped wildcard deployment for - callers outside that team, falling through to the shared wildcard entry. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "openai/*", - "litellm_params": {"model": "openai/*", "api_key": "team-b-key"}, - "model_info": { - "id": "team-b-wildcard", - "team_id": "team-b", - "team_public_model_name": "openai/*", - }, - }, - { - "model_name": "openai/*", - "litellm_params": {"model": "openai/*", "api_key": "global-key"}, - }, - ], - ) - - other_team_credentials = router.get_deployment_credentials_with_provider( - model_id="openai/gpt-5.2", team_id="team-a" - ) - assert other_team_credentials is not None - assert other_team_credentials["api_key"] == "global-key" - - owner_credentials = router.get_deployment_credentials_with_provider( - model_id="openai/gpt-5.2", team_id="team-b" - ) - assert owner_credentials is not None - assert owner_credentials["api_key"] == "team-b-key" - - -def test_team_wildcard_credentials_not_usable_after_delete_deployment(): - """ - Regression: team_pattern_routers retained deleted deployments, so a team - user could keep resolving credentials of a deleted wildcard deployment. - """ - router = litellm.Router(model_list=[_team_wildcard_model(api_key="old-key")]) - - assert ( - router.get_deployment_credentials_with_provider( - model_id="openai/gpt-5.2", team_id="team-1" - ) - is not None - ) - - router.delete_deployment(id="team-wildcard-id") - - assert ( - router.get_deployment_credentials_with_provider( - model_id="openai/gpt-5.2", team_id="team-1" - ) - is None - ) - - -def test_global_wildcard_pattern_router_evicts_stale_entry_on_upsert_and_delete(): - """ - Regression for #29064: upsert_deployment removed the old deployment from - model_list but left it in the global pattern_router, so wildcard requests - round-robined between the stale and the corrected deployment. - """ - from litellm.types.router import Deployment, LiteLLM_Params - - router = litellm.Router( - model_list=[ - { - "model_name": "openai/*", - "litellm_params": {"model": "openai/openai/*", "api_key": "sk-old"}, - "model_info": {"id": "global-wildcard"}, - } - ] - ) - - router.upsert_deployment( - Deployment( - model_name="openai/*", - litellm_params=LiteLLM_Params(model="openai/*", api_key="sk-new"), - model_info={"id": "global-wildcard"}, - ) - ) - - matches = router.pattern_router.route("openai/gpt-5.2") - assert matches is not None - assert [m["litellm_params"]["api_key"] for m in matches] == ["sk-new"] - - router.delete_deployment(id="global-wildcard") - assert router.pattern_router.patterns == {} - - -def test_pattern_match_router_remove_deployment(): - """ - remove_deployment must drop only the deployment with the given model id and - delete patterns whose deployment list becomes empty. - """ - from litellm.router_utils.pattern_match_deployments import PatternMatchRouter - - pattern_router = PatternMatchRouter() - pattern_router.add_pattern( - "openai/*", - {"litellm_params": {"model": "openai/*", "api_key": "key-a"}, "model_info": {"id": "dep-a"}}, - ) - pattern_router.add_pattern( - "openai/*", - {"litellm_params": {"model": "openai/*", "api_key": "key-b"}, "model_info": {"id": "dep-b"}}, - ) - - pattern_router.remove_deployment(model_id="dep-a") - matches = pattern_router.route("openai/gpt-5.2") - assert matches is not None - assert [m["model_info"]["id"] for m in matches] == ["dep-b"] - - pattern_router.remove_deployment(model_id="dep-b") - assert pattern_router.patterns == {} - assert pattern_router.route("openai/gpt-5.2") is None - - -def test_team_wildcard_credentials_refreshed_on_upsert_and_set_model_list(): - """ - Regression: replacing a team wildcard deployment (upsert or model list - reload) must serve the new credentials, not the stale cached ones. - """ - from litellm.types.router import Deployment - - router = litellm.Router(model_list=[_team_wildcard_model(api_key="old-key")]) - - router.upsert_deployment( - deployment=Deployment(**_team_wildcard_model(api_key="new-key")) - ) - credentials = router.get_deployment_credentials_with_provider( - model_id="openai/gpt-5.2", team_id="team-1" - ) - assert credentials is not None - assert credentials["api_key"] == "new-key" - - router.set_model_list(model_list=[]) - assert ( - router.get_deployment_credentials_with_provider( - model_id="openai/gpt-5.2", team_id="team-1" - ) - is None - ) - - -def test_get_available_guardrail_single_deployment(): - """ - Test get_available_guardrail returns the single guardrail when only one exists. - """ - guardrail_config = { - "guardrail_name": "content-filter", - "litellm_params": {"guardrail": "custom", "mode": "pre_call"}, - "id": "guardrail-1", - } - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo"}, - } - ], - guardrail_list=[guardrail_config], - ) - - result = router.get_available_guardrail(guardrail_name="content-filter") - assert result == guardrail_config - - -def test_get_available_guardrail_multiple_deployments(): - """ - Test get_available_guardrail load balances across multiple guardrails. - """ - guardrail_1 = { - "guardrail_name": "content-filter", - "litellm_params": {"guardrail": "custom", "mode": "pre_call"}, - "id": "guardrail-1", - } - guardrail_2 = { - "guardrail_name": "content-filter", - "litellm_params": {"guardrail": "custom", "mode": "pre_call"}, - "id": "guardrail-2", - } - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo"}, - } - ], - guardrail_list=[guardrail_1, guardrail_2], - ) - - # Call multiple times to verify load balancing - results = set() - for _ in range(20): - result = router.get_available_guardrail(guardrail_name="content-filter") - results.add(result["id"]) - - # Both guardrails should be selected at least once - assert "guardrail-1" in results or "guardrail-2" in results - - -def test_get_available_guardrail_not_found(): - """ - Test get_available_guardrail raises ValueError when guardrail not found. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo"}, - } - ], - guardrail_list=[], - ) - - with pytest.raises(ValueError, match="No guardrail found with name"): - router.get_available_guardrail(guardrail_name="non-existent") - - -@pytest.mark.asyncio -async def test_aguardrail_helper(): - """ - Test _aguardrail_helper selects a guardrail and executes the original function. - """ - guardrail_config = { - "guardrail_name": "content-filter", - "litellm_params": {"guardrail": "custom", "mode": "pre_call"}, - "id": "guardrail-1", - } - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo"}, - } - ], - guardrail_list=[guardrail_config], - ) - - # Mock the original function - async def mock_original_function(**kwargs): - return { - "result": "success", - "selected_guardrail": kwargs.get("selected_guardrail"), - } - - result = await router._aguardrail_helper( - model="content-filter", - original_generic_function=mock_original_function, - ) - - assert result["result"] == "success" - assert result["selected_guardrail"] == guardrail_config - - -@pytest.mark.asyncio -async def test_aguardrail(): - """ - Test aguardrail executes a guardrail with load balancing and fallbacks. - """ - guardrail_config = { - "guardrail_name": "content-filter", - "litellm_params": {"guardrail": "custom", "mode": "pre_call"}, - "id": "guardrail-1", - } - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo"}, - } - ], - guardrail_list=[guardrail_config], - ) - - # Mock the original function - async def mock_original_function(**kwargs): - return { - "result": "success", - "selected_guardrail": kwargs.get("selected_guardrail"), - } - - result = await router.aguardrail( - guardrail_name="content-filter", - original_function=mock_original_function, - ) - - assert result["result"] == "success" - assert result["selected_guardrail"]["id"] == "guardrail-1" - - -@pytest.mark.asyncio -async def test_anthropic_messages_call_type_is_cached(): - """ - Regression test: Verify that anthropic_messages call type is allowed - in PromptCachingDeploymentCheck.async_log_success_event. - """ - import asyncio - - from litellm.caching.dual_cache import DualCache - from litellm.router_utils.pre_call_checks.prompt_caching_deployment_check import ( - PromptCachingDeploymentCheck, - ) - from litellm.router_utils.prompt_caching_cache import PromptCachingCache - from litellm.types.utils import ( - CallTypes, - StandardLoggingHiddenParams, - StandardLoggingMetadata, - StandardLoggingModelInformation, - StandardLoggingPayload, - ) - - # Create mock standard logging payload inline - def create_standard_logging_payload() -> StandardLoggingPayload: - return StandardLoggingPayload( - id="test_id", - call_type="completion", - response_cost=0.1, - response_cost_failure_debug_info=None, - status="success", - total_tokens=30, - prompt_tokens=20, - completion_tokens=10, - startTime=1234567890.0, - endTime=1234567891.0, - completionStartTime=1234567890.5, - model_map_information=StandardLoggingModelInformation( - model_map_key="gpt-3.5-turbo", model_map_value=None - ), - model="gpt-3.5-turbo", - model_id="model-123", - model_group="openai-gpt", - api_base="https://api.openai.com", - metadata=StandardLoggingMetadata( - user_api_key_hash="test_hash", - user_api_key_org_id=None, - user_api_key_alias="test_alias", - user_api_key_team_id="test_team", - user_api_key_user_id="test_user", - user_api_key_team_alias="test_team_alias", - spend_logs_metadata=None, - requester_ip_address="127.0.0.1", - requester_metadata=None, - ), - cache_hit=False, - cache_key=None, - saved_cache_cost=0.0, - request_tags=[], - end_user=None, - requester_ip_address="127.0.0.1", - messages=[{"role": "user", "content": "Hello, world!"}], - response={"choices": [{"message": {"content": "Hi there!"}}]}, - error_str=None, - model_parameters={"stream": True}, - hidden_params=StandardLoggingHiddenParams( - model_id="model-123", - cache_key=None, - api_base="https://api.openai.com", - response_cost="0.1", - additional_headers=None, - ), - ) - - cache = DualCache() - deployment_check = PromptCachingDeploymentCheck(cache=cache) - prompt_cache = PromptCachingCache(cache=cache) - - # Create messages with enough tokens to pass the caching threshold - test_messages = [ - { - "role": "user", - "content": [ - { - "type": "text", - "text": "test long message here" * 1024, - "cache_control": {"type": "ephemeral", "ttl": "5m"}, - } - ], - } - ] - test_model_id = "test-model-id-123" - - # Create a payload with anthropic_messages call type - payload = create_standard_logging_payload() - payload["call_type"] = CallTypes.anthropic_messages.value - payload["messages"] = test_messages - payload["model"] = "anthropic/claude-3-5-sonnet-20240620" - payload["model_id"] = test_model_id - - # Log the success event (should cache the model_id) - await deployment_check.async_log_success_event( - kwargs={"standard_logging_object": payload}, - response_obj={}, - start_time=1234567890.0, - end_time=1234567891.0, - ) - - # Small delay to ensure cache write completes - await asyncio.sleep(0.1) - - # Verify that the model_id was actually cached - cached_result = await prompt_cache.async_get_model_id( - messages=test_messages, - tools=None, - ) - - # This assertion will FAIL if anthropic_messages is filtered out - assert ( - cached_result is not None - ), "Model ID should be cached for anthropic_messages call type" - assert ( - cached_result["model_id"] == test_model_id - ), f"Expected {test_model_id}, got {cached_result['model_id']}" - - -def test_update_kwargs_with_deployment_propagates_model_tags(): - """ - Test that deployment-level tags from litellm_params are merged into - kwargs metadata when _update_kwargs_with_deployment is called. - - This ensures model-level tags defined in config.yaml appear in SpendLogs. - See: https://github.com/BerriAI/litellm/issues/XXXX - """ - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4o-mini", - "litellm_params": { - "model": "openai/gpt-4o-mini", - "api_key": "fake-key", - "tags": ["openai-account", "production"], - }, - }, - ], - ) - - kwargs: dict = {"metadata": {}} - deployment = router.get_deployment_by_model_group_name( - model_group_name="gpt-4o-mini" - ) - router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) - - # Deployment tags should be propagated to kwargs metadata - assert "tags" in kwargs["metadata"] - assert "openai-account" in kwargs["metadata"]["tags"] - assert "production" in kwargs["metadata"]["tags"] - - -def test_update_kwargs_with_deployment_merges_tags_without_duplicates(): - """ - Test that when both request-level and deployment-level tags exist, - they are merged without duplicates. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4o-mini", - "litellm_params": { - "model": "openai/gpt-4o-mini", - "api_key": "fake-key", - "tags": ["openai-account", "shared-tag"], - }, - }, - ], - ) - - # Simulate request that already has tags (from request body or key/team level) - kwargs: dict = {"metadata": {"tags": ["user-tag", "shared-tag"]}} - deployment = router.get_deployment_by_model_group_name( - model_group_name="gpt-4o-mini" - ) - router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) - - # Both sources should be merged, no duplicates - assert "user-tag" in kwargs["metadata"]["tags"] - assert "openai-account" in kwargs["metadata"]["tags"] - assert "shared-tag" in kwargs["metadata"]["tags"] - assert kwargs["metadata"]["tags"].count("shared-tag") == 1 - - -def test_update_kwargs_with_deployment_no_tags(): - """ - Test that when deployment has no tags, kwargs metadata is not affected. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4o-mini", - "litellm_params": { - "model": "openai/gpt-4o-mini", - "api_key": "fake-key", - }, - }, - ], - ) - - kwargs: dict = {"metadata": {}} - deployment = router.get_deployment_by_model_group_name( - model_group_name="gpt-4o-mini" - ) - router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) - - # No tags key should be added if deployment has no tags - assert "tags" not in kwargs["metadata"] - - -def test_update_kwargs_with_deployment_merges_tools(): - """ - Test that when both deployment litellm_params and request have tools, - they are merged (deployment tools first, then request tools). - - Supports proxy-configured tools (e.g. for o3 deep research) merged with - client-provided tools. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "o3-deep-research", - "litellm_params": { - "model": "openai/o3-deep-research", - "api_key": "fake-key", - "tools": [{"type": "web_search"}], - "tool_choice": "auto", - }, - }, - ], - ) - - kwargs: dict = { - "metadata": {}, - "tools": [ - { - "type": "function", - "function": {"name": "get_weather", "description": "Get weather"}, - }, - ], - } - deployment = router.get_deployment_by_model_group_name( - model_group_name="o3-deep-research" - ) - router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) - - # Tools should be merged: deployment first, then request - assert "tools" in kwargs - assert len(kwargs["tools"]) == 2 - assert kwargs["tools"][0] == {"type": "web_search"} - assert kwargs["tools"][1]["function"]["name"] == "get_weather" - # tool_choice from request (none) - deployment's should be used - assert kwargs["tool_choice"] == "auto" - - -def test_update_kwargs_with_deployment_merge_tools_deployment_only(): - """ - Test that when only deployment has tools, they are applied to kwargs. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "o3-deep-research", - "litellm_params": { - "model": "openai/o3-deep-research", - "api_key": "fake-key", - "tools": [{"type": "web_search"}], - "tool_choice": "required", - }, - }, - ], - ) - - kwargs: dict = {"metadata": {}} - deployment = router.get_deployment_by_model_group_name( - model_group_name="o3-deep-research" - ) - router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) - - assert kwargs["tools"] == [{"type": "web_search"}] - assert kwargs["tool_choice"] == "required" - - -def test_update_kwargs_with_deployment_merge_tools_request_overrides_tool_choice(): - """ - Test that when request has tool_choice, it overrides deployment's. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "o3-deep-research", - "litellm_params": { - "model": "openai/o3-deep-research", - "api_key": "fake-key", - "tools": [{"type": "web_search"}], - "tool_choice": "auto", - }, - }, - ], - ) - - kwargs: dict = { - "metadata": {}, - "tool_choice": "none", - } - deployment = router.get_deployment_by_model_group_name( - model_group_name="o3-deep-research" - ) - router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) - - # Request tool_choice should be preserved (merged tools still applied) - assert kwargs["tool_choice"] == "none" - - -def test_credential_name_injected_as_tag(): - """ - Test that litellm_credential_name from deployment litellm_params - is injected as a tag into metadata during _update_kwargs_with_deployment. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "xai-model", - "litellm_params": { - "model": "xai/grok-4-1-fast", - "litellm_credential_name": "xAI", - }, - } - ], - ) - - kwargs: dict = {"metadata": {"tags": ["A.101"]}} - deployment = router.get_deployment_by_model_group_name(model_group_name="xai-model") - router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) - - assert "Credential: xAI" in kwargs["metadata"]["tags"] - assert "A.101" in kwargs["metadata"]["tags"] - - -def test_credential_name_not_duplicated_in_tags(): - """ - Test that if the credential tag already exists in the tags list, - it is not duplicated. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "xai-model", - "litellm_params": { - "model": "xai/grok-4-1-fast", - "litellm_credential_name": "xAI", - }, - } - ], - ) - - kwargs: dict = {"metadata": {"tags": ["Credential: xAI", "A.101"]}} - deployment = router.get_deployment_by_model_group_name(model_group_name="xai-model") - router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) - - assert kwargs["metadata"]["tags"].count("Credential: xAI") == 1 - - -def test_credential_name_not_injected_when_absent(): - """ - Test that when no litellm_credential_name is set, tags are unchanged. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-model", - "litellm_params": { - "model": "gpt-4o", - }, - } - ], - ) - - kwargs: dict = {"metadata": {"tags": ["A.101"]}} - deployment = router.get_deployment_by_model_group_name(model_group_name="gpt-model") - router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) - - assert kwargs["metadata"]["tags"] == ["A.101"] - - -def test_update_kwargs_with_deployment_model_info_in_litellm_metadata(): - """For generic_api_call, model_info with pricing must go to litellm_metadata. - - Routes like /messages and /responses use generic_api_call which stores - model_info under litellm_metadata. Regression test for #23185. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "claude-sonnet-4", - "litellm_params": { - "model": "anthropic/claude-sonnet-4-20250514", - "api_key": "fake-key", - }, - "model_info": { - "id": "custom-pricing-id", - "input_cost_per_token": 0.0003, - "output_cost_per_token": 0.0015, - }, - }, - ], - ) - - kwargs: dict = {} - deployment = router.get_deployment_by_model_group_name( - model_group_name="claude-sonnet-4" - ) - router._update_kwargs_with_deployment( - deployment=deployment, kwargs=kwargs, function_name="generic_api_call" - ) - - assert "litellm_metadata" in kwargs - model_info = kwargs["litellm_metadata"]["model_info"] - assert model_info["id"] == "custom-pricing-id" - assert model_info["input_cost_per_token"] == 0.0003 - assert model_info["output_cost_per_token"] == 0.0015 - - -def test_update_kwargs_with_deployment_model_info_in_metadata(): - """For acompletion (function_name=None), model_info goes to metadata. - - /chat/completions uses acompletion which stores model_info under metadata. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "claude-sonnet-4", - "litellm_params": { - "model": "anthropic/claude-sonnet-4-20250514", - "api_key": "fake-key", - }, - "model_info": { - "id": "custom-pricing-id", - "input_cost_per_token": 0.0003, - "output_cost_per_token": 0.0015, - }, - }, - ], - ) - - kwargs: dict = {} - deployment = router.get_deployment_by_model_group_name( - model_group_name="claude-sonnet-4" - ) - router._update_kwargs_with_deployment( - deployment=deployment, kwargs=kwargs, function_name=None - ) - - assert "metadata" in kwargs - model_info = kwargs["metadata"]["model_info"] - assert model_info["id"] == "custom-pricing-id" - assert model_info["input_cost_per_token"] == 0.0003 - assert model_info["output_cost_per_token"] == 0.0015 - - -def test_combine_fallback_usage(): - """Test that _combine_fallback_usage merges partial and fallback usage.""" - from litellm.router import Router - from litellm.types.utils import Usage - - # Create a stream chunk with usage - chunk = litellm.ModelResponseStream( - id="test", - model="gpt-4o", - choices=[], - usage=Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15), - ) - - # Call _combine_fallback_usage with no extra usage - Router._combine_fallback_usage(chunk, None) - assert chunk.usage is not None - assert chunk.usage.prompt_tokens == 10 - assert chunk.usage.completion_tokens == 5 - assert chunk.usage.total_tokens == 15 - - -@pytest.mark.asyncio -async def test_acompletion_streaming_iterator_does_not_log_success_on_terminal_failure(): - """A mid-stream failure with no successful fallback raises and is logged as - a failure, so the router must never dispatch it as a success. Partial-spend - recovery for the failure row happens in the streaming handler, not here, so - this guards only against reintroducing a success log for a failed stream. - """ - from litellm.exceptions import MidStreamFallbackError - from litellm.types.utils import Delta, StreamingChoices, Usage - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4", - "litellm_params": {"model": "gpt-4", "api_key": "fake-key-1"}, - }, - ], - set_verbose=True, - ) - - error = MidStreamFallbackError( - message="Connection lost", - model="gpt-4", - llm_provider="openai", - generated_content="The Roman Empire began when", - ) - - def _make_interrupted_model_response(): - partial_chunk = litellm.ModelResponseStream( - id="chatcmpl-partial-1", - created=1742056047, - model="gpt-4", - object="chat.completion.chunk", - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta(content="The Roman Empire began when", role="assistant"), - ) - ], - usage=Usage(prompt_tokens=17, completion_tokens=9, total_tokens=26), - ) - - class _RaisingStream: - def __init__(self): - self.index = 0 - self.chunks = [partial_chunk] - - def __aiter__(self): - return self - - async def __anext__(self): - if self.index == 0: - self.index += 1 - return partial_chunk - raise error - - stream = _RaisingStream() - logging_obj = MagicMock() - logging_obj.dispatch_success_handlers = AsyncMock() - logging_obj.model_call_details = {} - setattr(stream, "model", "gpt-4") - setattr(stream, "custom_llm_provider", "openai") - setattr(stream, "logging_obj", logging_obj) - return stream, logging_obj - - messages = [{"role": "user", "content": "Hello"}] - initial_kwargs = {"model": "gpt-4", "stream": True} - - # Terminal path: no successful fallback -> the error propagates and the - # router never dispatches a success for the failed stream. - model_response, logging_obj = _make_interrupted_model_response() - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(side_effect=error), - ): - result = await router._acompletion_streaming_iterator( - model_response=model_response, - messages=messages, - initial_kwargs=dict(initial_kwargs), - ) - collected = [] - async def _drain(): - async for chunk in result: - collected.append(chunk) - - with pytest.raises(MidStreamFallbackError): - await _drain() - - assert len(collected) == 1 - logging_obj.dispatch_success_handlers.assert_not_called() - - # Mid-stream errors with generated content are now re-raised immediately; - # no continuation-prompt fallback is attempted. Success handlers must - # still not be dispatched in this path. - model_response, logging_obj = _make_interrupted_model_response() - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(), - ) as mock_fallback: - result = await router._acompletion_streaming_iterator( - model_response=model_response, - messages=messages, - initial_kwargs=dict(initial_kwargs), - ) - collected = [] - async def _drain(): - async for chunk in result: - collected.append(chunk) - - with pytest.raises(MidStreamFallbackError): - await _drain() - - assert len(collected) == 1, "only the partial chunk before the error" - mock_fallback.assert_not_called() - logging_obj.dispatch_success_handlers.assert_not_called() - - -@pytest.mark.asyncio -async def test_team_scoped_model_fallback(): - """ - Test that fallback works correctly for team-scoped models. - - When a team-scoped model fails and the fallback model is also team-scoped, - the router should find the fallback deployment by matching team_public_model_name. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "team-a-primary-internal", - "litellm_params": {"model": "gpt-3.5-turbo", "api_key": "fake"}, - "model_info": { - "team_id": "team-a", - "team_public_model_name": "primary-model", - }, - }, - { - "model_name": "team-a-fallback-internal", - "litellm_params": { - "model": "gpt-4", - "api_key": "fake", - "mock_response": "fallback success from team-a", - }, - "model_info": { - "team_id": "team-a", - "team_public_model_name": "fallback-model", - }, - }, - ], - fallbacks=[{"primary-model": ["fallback-model"]}], - ) - - response = await router.acompletion( - model="primary-model", - messages=[{"role": "user", "content": "Hello"}], - metadata={"user_api_key_team_id": "team-a"}, - mock_testing_fallbacks=True, - ) - assert response is not None - assert response.choices[0].message.content == "fallback success from team-a" - - -@pytest.mark.asyncio -async def test_team_scoped_model_fallback_to_global(): - """ - Test that a team-scoped model can fall back to a global (non-team) model. - - Global models (no team_id on deployment) should be accessible as fallback - targets for team-scoped requests. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "team-a-primary-internal", - "litellm_params": {"model": "gpt-3.5-turbo", "api_key": "fake"}, - "model_info": { - "team_id": "team-a", - "team_public_model_name": "primary-model", - }, - }, - { - "model_name": "global-fallback", - "litellm_params": { - "model": "gpt-4", - "api_key": "fake", - "mock_response": "global fallback success", - }, - }, - ], - fallbacks=[{"primary-model": ["global-fallback"]}], - ) - - response = await router.acompletion( - model="primary-model", - messages=[{"role": "user", "content": "Hello"}], - metadata={"user_api_key_team_id": "team-a"}, - mock_testing_fallbacks=True, - ) - assert response is not None - assert response.choices[0].message.content == "global fallback success" - - -@pytest.mark.asyncio -async def test_team_scoped_model_fallback_cross_team_blocked(): - """ - Test that cross-team fallback is correctly blocked. - - When team-a's model fails and the fallback target is scoped to team-b, - the router should NOT use it (team isolation). - """ - router = litellm.Router( - model_list=[ - { - "model_name": "team-a-primary-internal", - "litellm_params": {"model": "gpt-3.5-turbo", "api_key": "fake"}, - "model_info": { - "team_id": "team-a", - "team_public_model_name": "primary-model", - }, - }, - { - "model_name": "team-b-fallback-internal", - "litellm_params": { - "model": "gpt-4", - "api_key": "fake", - "mock_response": "team-b response - should not reach here", - }, - "model_info": { - "team_id": "team-b", - "team_public_model_name": "fallback-model", - }, - }, - ], - fallbacks=[{"primary-model": ["fallback-model"]}], - ) - - with pytest.raises(litellm.InternalServerError): - await router.acompletion( - model="primary-model", - messages=[{"role": "user", "content": "Hello"}], - metadata={"user_api_key_team_id": "team-a"}, - mock_testing_fallbacks=True, - ) - - -def test_get_all_deployments_with_team_id(): - """ - Test that _get_all_deployments with team_id can find deployments - by team_public_model_name when the model_name is not in the index. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "internal-team-deployment", - "litellm_params": {"model": "gpt-4", "api_key": "fake"}, - "model_info": { - "team_id": "team-x", - "team_public_model_name": "gpt-4", - }, - }, - ], - ) - - # Without team_id: "gpt-4" is not in the model_name index (internal name is different) - deployments = router._get_all_deployments(model_name="gpt-4") - assert len(deployments) == 0 - - # With correct team_id: should find via O(n) scan matching team_public_model_name - deployments = router._get_all_deployments(model_name="gpt-4", team_id="team-x") - assert len(deployments) == 1 - assert deployments[0]["model_name"] == "internal-team-deployment" - - # With wrong team_id: should find nothing - deployments = router._get_all_deployments(model_name="gpt-4", team_id="team-y") - assert len(deployments) == 0 - - -def test_multiregion_team_deployments_unique_model_names(): - """ - Simulates athenahealth's exact setup: unique model_names per deployment, - same team_public_model_name, multiple regions. - - Verifies that _get_all_deployments returns ALL regional deployments - for a team when queried by team_public_model_name. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "metis-claude-us-east-1", - "litellm_params": { - "model": "bedrock/anthropic.claude-3-sonnet", - "aws_region_name": "us-east-1", - "api_key": "fake", - }, - "model_info": { - "team_id": "metis-team", - "team_public_model_name": "claude-sonnet", - }, - }, - { - "model_name": "metis-claude-us-west-2", - "litellm_params": { - "model": "bedrock/anthropic.claude-3-sonnet", - "aws_region_name": "us-west-2", - "api_key": "fake", - }, - "model_info": { - "team_id": "metis-team", - "team_public_model_name": "claude-sonnet", - }, - }, - ], - ) - - # "claude-sonnet" is NOT in the model_name index - assert "claude-sonnet" not in router.model_names - - # Without team_id: returns nothing (no model_name="claude-sonnet" in index, no O(n) scan) - deployments = router._get_all_deployments(model_name="claude-sonnet") - assert len(deployments) == 0 - - # With team_id: O(n) scan finds BOTH regional deployments - deployments = router._get_all_deployments( - model_name="claude-sonnet", team_id="metis-team" - ) - assert len(deployments) == 2 - deployment_names = {d["model_name"] for d in deployments} - assert deployment_names == {"metis-claude-us-east-1", "metis-claude-us-west-2"} - - # Each deployment has a unique ID (critical for cooldown/retry to work) - deployment_ids = {d["model_info"]["id"] for d in deployments} - assert ( - len(deployment_ids) == 2 - ), "Each deployment must have a unique ID for cooldown tracking" - - # Wrong team: returns nothing - deployments = router._get_all_deployments( - model_name="claude-sonnet", team_id="other-team" - ) - assert len(deployments) == 0 - - -@pytest.mark.asyncio -async def test_multiregion_team_failover_between_regions(): - """ - Simulates athenahealth's multiregion failover scenario: - - Two Bedrock deployments (us-east-1 and us-west-2) with unique model_names - - Same team_public_model_name ("claude-sonnet") - - Primary region fails → router should failover to second region - - This is the exact scenario Sean Glover from athenahealth will demonstrate. - """ - router = litellm.Router( - model_list=[ - { - "model_name": "metis-claude-us-east-1", - "litellm_params": { - "model": "bedrock/anthropic.claude-3-sonnet", - "api_key": "fake", - "mock_response": "response from us-east-1", - }, - "model_info": { - "team_id": "metis-team", - "team_public_model_name": "claude-sonnet", - }, - }, - { - "model_name": "metis-claude-us-west-2", - "litellm_params": { - "model": "bedrock/anthropic.claude-3-sonnet", - "api_key": "fake", - "mock_response": "response from us-west-2", - }, - "model_info": { - "team_id": "metis-team", - "team_public_model_name": "claude-sonnet", - }, - }, - ], - num_retries=1, - ) - - # Verify the router finds both deployments for the team - deployments = router._get_all_deployments( - model_name="claude-sonnet", team_id="metis-team" - ) - assert ( - len(deployments) == 2 - ), "Router must find both regional deployments by team_public_model_name" - - # Make a normal request — should succeed from one of the regions - response = await router.acompletion( - model="claude-sonnet", - messages=[{"role": "user", "content": "Hello"}], - metadata={"user_api_key_team_id": "metis-team"}, - ) - assert response is not None - assert response.choices[0].message.content in [ - "response from us-east-1", - "response from us-west-2", - ] - - -def test_access_group_scoped_key_filters_deployments_with_same_public_model(): - """ - If a key can access a model only via access group membership, - router candidate deployments for that public model should be constrained - to deployments in the allowed access group. - """ - from litellm.proxy._types import UserAPIKeyAuth - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-5", - "litellm_params": { - "model": "openai/gpt-5.1", - "api_key": "key1", - "mock_response": "response-via-AG1", - }, - "model_info": {"access_groups": ["AG1"]}, - }, - { - "model_name": "gpt-5", - "litellm_params": { - "model": "openai/gpt-4o", - "api_key": "key2", - "mock_response": "response-via-AG2", - }, - "model_info": {"access_groups": ["AG2"]}, - }, - ] - ) - - scoped_key = UserAPIKeyAuth( - api_key="hashed-key", - team_id="team2", - models=["AG2"], - team_models=["AG2"], - ) - - _model, deployments = router._common_checks_available_deployment( - model="gpt-5", - request_kwargs={ - "metadata": { - "user_api_key_team_id": "team2", - "user_api_key_auth": scoped_key, - } - }, - ) - - assert len(deployments) == 1 - assert deployments[0].get("model_info", {}).get("access_groups") == ["AG2"] - - seen = set() - for _ in range(20): - response = router.completion( - model="gpt-5", - messages=[{"role": "user", "content": "hello"}], - metadata={"user_api_key_team_id": "team2", "user_api_key_auth": scoped_key}, - ) - seen.add(response.choices[0].message.content) - - assert seen == {"response-via-AG2"} - - -def test_explicit_model_access_does_not_force_access_group_filtering(): - """ - If a key has explicit model access in addition to access group entries, - do not force access-group-only filtering for deployment selection. - """ - from litellm.proxy._types import UserAPIKeyAuth - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-5", - "litellm_params": { - "model": "openai/gpt-5.1", - "api_key": "key1", - "mock_response": "response-via-AG1", - }, - "model_info": {"access_groups": ["AG1"]}, - }, - { - "model_name": "gpt-5", - "litellm_params": { - "model": "openai/gpt-4o", - "api_key": "key2", - "mock_response": "response-via-AG2", - }, - "model_info": {"access_groups": ["AG2"]}, - }, - ] - ) - - explicit_key = UserAPIKeyAuth( - api_key="hashed-key", - team_id="team2", - models=["AG2", "gpt-5"], - team_models=["AG2", "gpt-5"], - ) - - _model, deployments = router._common_checks_available_deployment( - model="gpt-5", - request_kwargs={ - "metadata": { - "user_api_key_team_id": "team2", - "user_api_key_auth": explicit_key, - } - }, - ) - - deployment_groups = [ - d.get("model_info", {}).get("access_groups") for d in deployments - ] - assert ["AG1"] in deployment_groups - assert ["AG2"] in deployment_groups - - -def test_access_group_filter_empty_does_not_bypass_via_litellm_model_fallback( - monkeypatch: pytest.MonkeyPatch, -): - """ - When access-group filtering removes all candidates, _get_deployment_by_litellm_model - must not run: it does not re-apply access groups and could return blocked deployments - that share the same litellm_params.model as the request model string. - - ``get_model_access_groups`` is patched to expose AG1 for the public model (so the - access-group filter runs with a non-empty allowed set) while every deployment - returned for that name is AG2-only — filtered to empty. Without the guard, the - litellm-model fallback would return both rows because ``litellm_params.model`` matches. - """ - from litellm.proxy._types import UserAPIKeyAuth - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-5", - "litellm_params": { - "model": "gpt-5", - "api_key": "key1", - "mock_response": "blocked-dep-1", - }, - "model_info": {"access_groups": ["AG2"]}, - }, - { - "model_name": "gpt-5", - "litellm_params": { - "model": "gpt-5", - "api_key": "key2", - "mock_response": "blocked-dep-2", - }, - "model_info": {"access_groups": ["AG2"]}, - }, - ] - ) - - orig_groups = router.get_model_access_groups - - def fake_get_model_access_groups( - model_name=None, model_access_group=None, team_id=None - ): - if model_name == "gpt-5" and model_access_group is None: - return {"AG1": ["gpt-5"], "AG2": ["gpt-5"]} - return orig_groups( - model_name=model_name, - model_access_group=model_access_group, - team_id=team_id, - ) - - monkeypatch.setattr(router, "get_model_access_groups", fake_get_model_access_groups) - - scoped_key = UserAPIKeyAuth( - api_key="hashed-key", - team_id="team2", - models=["AG1"], - team_models=["AG1"], - ) - - with pytest.raises(litellm.BadRequestError): - router._common_checks_available_deployment( - model="gpt-5", - request_kwargs={ - "metadata": { - "user_api_key_team_id": "team2", - "user_api_key_auth": scoped_key, - } - }, - ) - - -def test_access_group_block_does_not_silently_use_default_fallback_model( - monkeypatch: pytest.MonkeyPatch, -): - """ - When access-group filtering empties candidates for model X, the router must not use - ``fallbacks`` default ``*`` routing to model Y: Y may have no ``access_groups``, so - ``_filter_deployments_by_model_access_groups`` would not constrain Y and the caller - would be served despite being blocked from X. - """ - from litellm.proxy._types import UserAPIKeyAuth - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-5", - "litellm_params": { - "model": "gpt-5", - "api_key": "key1", - "mock_response": "blocked-dep-1", - }, - "model_info": {"access_groups": ["AG2"]}, - }, - { - "model_name": "gpt-5", - "litellm_params": { - "model": "gpt-5", - "api_key": "key2", - "mock_response": "blocked-dep-2", - }, - "model_info": {"access_groups": ["AG2"]}, - }, - { - "model_name": "gpt-4-fallback", - "litellm_params": { - "model": "gpt-4", - "api_key": "fallback-key", - "mock_response": "should-not-reach", - }, - }, - ], - fallbacks=[{"*": ["gpt-4-fallback"]}], - ) - - orig_groups = router.get_model_access_groups - - def fake_get_model_access_groups( - model_name=None, model_access_group=None, team_id=None - ): - if model_name == "gpt-5" and model_access_group is None: - return {"AG1": ["gpt-5"], "AG2": ["gpt-5"]} - return orig_groups( - model_name=model_name, - model_access_group=model_access_group, - team_id=team_id, - ) - - monkeypatch.setattr(router, "get_model_access_groups", fake_get_model_access_groups) - - scoped_key = UserAPIKeyAuth( - api_key="hashed-key", - team_id="team2", - models=["AG1"], - team_models=["AG1"], - ) - - with pytest.raises(litellm.BadRequestError): - router._common_checks_available_deployment( - model="gpt-5", - request_kwargs={ - "metadata": { - "user_api_key_team_id": "team2", - "user_api_key_auth": scoped_key, - } - }, - ) - - -def test_access_group_block_via_litellm_model_branch_does_not_use_default_fallback( - monkeypatch: pytest.MonkeyPatch, -): - """ - When the by-name lookup returns no deployments and the litellm-model fallback - branch finds candidates that access-group filtering then empties, the router - must not fall through to default ``fallbacks`` routing — the default fallback - model may have no ``access_groups`` and would short-circuit the filter, - silently serving a caller blocked by access-group restrictions. - """ - from litellm.proxy._types import UserAPIKeyAuth - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-5-alias", - "litellm_params": { - "model": "gpt-5", - "api_key": "key1", - "mock_response": "blocked-dep-1", - }, - "model_info": {"access_groups": ["AG2"]}, - }, - { - "model_name": "gpt-4-fallback", - "litellm_params": { - "model": "gpt-4", - "api_key": "fallback-key", - "mock_response": "should-not-reach", - }, - }, - ], - fallbacks=[{"*": ["gpt-4-fallback"]}], - ) - - orig_groups = router.get_model_access_groups - - def fake_get_model_access_groups( - model_name=None, model_access_group=None, team_id=None - ): - if model_name == "gpt-5" and model_access_group is None: - return {"AG1": ["gpt-5"], "AG2": ["gpt-5"]} - return orig_groups( - model_name=model_name, - model_access_group=model_access_group, - team_id=team_id, - ) - - monkeypatch.setattr(router, "get_model_access_groups", fake_get_model_access_groups) - - scoped_key = UserAPIKeyAuth( - api_key="hashed-key", - team_id="team2", - models=["AG1"], - team_models=["AG1"], - ) - - with pytest.raises(litellm.BadRequestError): - router._common_checks_available_deployment( - model="gpt-5", - request_kwargs={ - "metadata": { - "user_api_key_team_id": "team2", - "user_api_key_auth": scoped_key, - } - }, - ) - - -def test_try_early_resolve_deployments_for_model_not_in_names(): - """ - Direct coverage for ``_try_early_resolve_deployments_for_model_not_in_names``: - - - Returns ``None`` when the requested model is already in ``self.model_names`` - (the by-name lookup path will handle it). - - Returns ``None`` when there are no team deployments, no pattern matches, and - no default deployment to fall back to. - - Returns the pattern-router match when the model matches a wildcard route. - - Returns the default deployment with the request model substituted in when one - is configured, without mutating the stored default. - """ - router_in_names = litellm.Router( - model_list=[ - { - "model_name": "gpt-5", - "litellm_params": { - "model": "openai/gpt-5", - "api_key": "key1", - }, - }, - ] - ) - - assert ( - router_in_names._try_early_resolve_deployments_for_model_not_in_names( - model="gpt-5", request_team_id=None - ) - is None - ) - assert ( - router_in_names._try_early_resolve_deployments_for_model_not_in_names( - model="some-unknown-model", request_team_id=None - ) - is None - ) - - pattern_router = litellm.Router( - model_list=[ - { - "model_name": "openai/*", - "litellm_params": { - "model": "openai/*", - "api_key": "key-pattern", - }, - }, - ] - ) - - pattern_result = ( - pattern_router._try_early_resolve_deployments_for_model_not_in_names( - model="openai/gpt-4o-mini", request_team_id=None - ) - ) - assert pattern_result is not None - resolved_model, pattern_deployments = pattern_result - assert resolved_model == "openai/gpt-4o-mini" - assert isinstance(pattern_deployments, list) and len(pattern_deployments) == 1 - - default_router = litellm.Router( - model_list=[ - { - "model_name": "named-model", - "litellm_params": { - "model": "openai/gpt-4o", - "api_key": "key-named", - }, - }, - ] - ) - default_router.default_deployment = { - "model_name": "default", - "litellm_params": { - "model": "openai/will-be-overridden", - "api_key": "key-default", - }, - } - - default_result = ( - default_router._try_early_resolve_deployments_for_model_not_in_names( - model="brand-new-model", request_team_id=None - ) - ) - assert default_result is not None - resolved_model, default_deployment = default_result - assert resolved_model == "brand-new-model" - assert isinstance(default_deployment, dict) - assert default_deployment["litellm_params"]["model"] == "brand-new-model" - # The original default_deployment must not be mutated. - assert ( - default_router.default_deployment["litellm_params"]["model"] - == "openai/will-be-overridden" - ) - - -def _router_with_two_deployments(blocked_flags): - import litellm - - model_list = [] - for idx, blocked in enumerate(blocked_flags): - model_list.append( - { - "model_name": "gpt-4o", - "litellm_params": {"model": f"openai/gpt-4o-{idx}"}, - "model_info": {"id": f"dep-{idx}", "blocked": blocked}, - } - ) - return litellm.Router(model_list=model_list) - - -def test_get_fully_blocked_model_names_marks_name_when_all_deployments_blocked(): - router = _router_with_two_deployments([True, True]) - assert router.get_fully_blocked_model_names() == {"gpt-4o"} - - -def test_get_fully_blocked_model_names_keeps_name_when_partial_blocked(): - router = _router_with_two_deployments([True, False]) - assert router.get_fully_blocked_model_names() == set() - - -def test_get_fully_blocked_model_names_treats_missing_key_as_unblocked(): - import litellm - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4o", - "litellm_params": {"model": "openai/gpt-4o"}, - "model_info": {"id": "dep-0"}, - } - ] - ) - assert router.get_fully_blocked_model_names() == set() - - -def _seed_unhealthy_states(router, unhealthy_ids, timestamp=None): - import time - - ts = timestamp if timestamp is not None else time.time() - router.health_state_cache.set_deployment_health_states( - { - uid: {"is_healthy": False, "timestamp": ts, "reason": "test_unhealthy"} - for uid in unhealthy_ids - } - ) - - -@pytest.mark.asyncio -async def test_async_get_fully_unhealthy_model_names_marks_name_when_all_unhealthy(): - router = _router_with_two_deployments([False, False]) - _seed_unhealthy_states(router, {"dep-0", "dep-1"}) - assert await router.async_get_fully_unhealthy_model_names() == {"gpt-4o"} - - -@pytest.mark.asyncio -async def test_async_get_fully_unhealthy_model_names_keeps_name_when_partial(): - router = _router_with_two_deployments([False, False]) - _seed_unhealthy_states(router, {"dep-0"}) - assert await router.async_get_fully_unhealthy_model_names() == set() - - -@pytest.mark.asyncio -async def test_async_get_fully_unhealthy_model_names_empty_without_health_state(): - router = _router_with_two_deployments([False, False]) - assert await router.async_get_fully_unhealthy_model_names() == set() - - -@pytest.mark.asyncio -async def test_async_get_fully_unhealthy_model_names_ignores_stale_state(): - import time - - router = _router_with_two_deployments([False, False]) - stale_ts = time.time() - (router.health_state_cache.staleness_threshold + 10) - _seed_unhealthy_states(router, {"dep-0", "dep-1"}, timestamp=stale_ts) - assert await router.async_get_fully_unhealthy_model_names() == set() - - -@pytest.mark.asyncio -async def test_async_get_fully_unhealthy_model_names_includes_team_alias(): - import litellm - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-4o", - "litellm_params": {"model": "openai/gpt-4o"}, - "model_info": { - "id": "dep-0", - "team_id": "team-1", - "team_public_model_name": "team-gpt", - }, - } - ] - ) - _seed_unhealthy_states(router, {"dep-0"}) - assert await router.async_get_fully_unhealthy_model_names() == { - "gpt-4o", - "team-gpt", - } - - -@pytest.mark.asyncio -async def test_async_get_fully_unhealthy_model_names_noop_with_allowed_fails_policy(): - from litellm.types.router import AllowedFailsPolicy - - router = _router_with_two_deployments([False, False]) - router.allowed_fails_policy = AllowedFailsPolicy(BadRequestErrorAllowedFails=1) - _seed_unhealthy_states(router, {"dep-0", "dep-1"}) - assert await router.async_get_fully_unhealthy_model_names() == set() - - -@pytest.mark.asyncio -async def test_async_get_healthy_deployments_skips_blocked_deployment(): - router = _router_with_two_deployments([True, False]) - healthy, all_dep = await router._async_get_healthy_deployments( - model="gpt-4o", parent_otel_span=None - ) - healthy_ids = [d["model_info"]["id"] for d in healthy] - assert "dep-0" not in healthy_ids - assert "dep-1" in healthy_ids - assert len(all_dep) == 2 - - -def test_get_healthy_deployments_sync_skips_blocked_deployment(): - router = _router_with_two_deployments([False, True]) - healthy, all_dep = router._get_healthy_deployments( - model="gpt-4o", parent_otel_span=None - ) - healthy_ids = [d["model_info"]["id"] for d in healthy] - assert "dep-0" in healthy_ids - assert "dep-1" not in healthy_ids - assert len(all_dep) == 2 - - -def test_filter_blocked_deployments_drops_blocked_keeps_unblocked(): - router = _router_with_two_deployments([True, False]) - filtered = router._filter_blocked_deployments(router.get_model_list() or []) - ids = [d["model_info"]["id"] for d in filtered] - assert ids == ["dep-1"] - - -@pytest.mark.asyncio -async def test_public_async_get_healthy_deployments_skips_blocked_on_primary_path(): - router = _router_with_two_deployments([True, False]) - deployments = await router.async_get_healthy_deployments( - model="gpt-4o", request_kwargs={} - ) - assert isinstance(deployments, list) - ids = [d["model_info"]["id"] for d in deployments] - assert "dep-0" not in ids - assert "dep-1" in ids - - -def test_public_get_available_deployment_skips_blocked_on_primary_path(): - router = _router_with_two_deployments([True, False]) - deployment = router.get_available_deployment(model="gpt-4o", request_kwargs={}) - assert deployment["model_info"]["id"] == "dep-1" - - -def test_get_available_deployment_raises_when_addressed_dict_is_blocked(): - import litellm - - router = _router_with_two_deployments([True, True]) - with pytest.raises(litellm.ServiceUnavailableError): - router.get_available_deployment(model="dep-0", request_kwargs={}) - - -def _router_with_two_pass_through_deployments(blocked_flags): - import litellm - - model_list = [] - for idx, blocked in enumerate(blocked_flags): - model_list.append( - { - "model_name": "gpt-4o", - "litellm_params": { - "model": f"openai/gpt-4o-{idx}", - "api_key": "sk-fake-for-tests", - "use_in_pass_through": True, - }, - "model_info": {"id": f"pt-{idx}", "blocked": blocked}, - } - ) - return litellm.Router(model_list=model_list) - - -def test_get_available_deployment_for_pass_through_skips_blocked(): - router = _router_with_two_pass_through_deployments([True, False]) - deployment = router.get_available_deployment_for_pass_through( - model="gpt-4o", request_kwargs={} - ) - assert deployment["model_info"]["id"] == "pt-1" - - -def test_get_available_deployment_for_pass_through_raises_when_dict_blocked(): - import litellm - - router = _router_with_two_pass_through_deployments([True, True]) - with pytest.raises(litellm.ServiceUnavailableError): - router.get_available_deployment_for_pass_through( - model="pt-0", request_kwargs={} - ) - - -def test_initialize_deployment_for_pass_through_keeps_bedrock_iam_deployment(): - """ - Bedrock deployments using IAM/OIDC auth have no api_key; pass-through - init must not raise and drop them from routing (#27728). - """ - import litellm - - router = litellm.Router( - model_list=[ - { - "model_name": "bedrock-claude", - "litellm_params": { - "model": "bedrock/anthropic.claude-3-5-sonnet-20241022-v2:0", - "aws_role_name": "arn:aws:iam::123456789012:role/my-role", - "aws_session_name": "my-session", - "use_in_pass_through": True, - }, - "model_info": {"id": "bedrock-iam-pt"}, - } - ] - ) - assert [m["model_info"]["id"] for m in router.get_model_list()] == [ - "bedrock-iam-pt" - ] - - -def test_pass_through_deployment_api_key_resolves_via_get_credentials(): - from litellm.proxy.pass_through_endpoints.passthrough_endpoint_router import ( - PassthroughEndpointRouter, - ) - - router = _router_with_two_pass_through_deployments([False, False]) - passthrough_router = PassthroughEndpointRouter(llm_router_getter=lambda: router) - assert len(router.get_model_list()) == 2 - assert ( - passthrough_router.get_credentials( - custom_llm_provider="openai", region_name=None - ) - == "sk-fake-for-tests" - ) - - -def test_get_deployment_credentials_returns_none_for_blocked_deployment(): - router = _router_with_two_deployments([True, False]) - assert router.get_deployment_credentials(model_id="dep-0") is None - assert router.get_deployment_credentials(model_id="dep-1") is not None - - -def test_get_deployment_credentials_with_provider_returns_none_for_blocked_deployment(): - router = _router_with_two_deployments([True, False]) - assert router.get_deployment_credentials_with_provider(model_id="dep-0") is None - assert router.get_deployment_credentials_with_provider(model_id="dep-1") is not None - - -def test_is_deployment_blocked_static_helper_reflects_blocked_flag(): - """ - Exercises Router._is_deployment_blocked so router_code_coverage.py (AST call graph) - marks the helper as covered by router-named tests. - """ - import types - - import litellm - - router = _router_with_two_deployments([True, False]) - blocked_dep = router.get_deployment("dep-0") - unblocked_dep = router.get_deployment("dep-1") - assert blocked_dep is not None and unblocked_dep is not None - assert litellm.Router._is_deployment_blocked(blocked_dep) is True - assert litellm.Router._is_deployment_blocked(unblocked_dep) is False - - # No model_info on deployment object → treated as not blocked - assert litellm.Router._is_deployment_blocked(object()) is False - missing_blocked = types.SimpleNamespace() - assert ( - litellm.Router._is_deployment_blocked( - types.SimpleNamespace(model_info=missing_blocked) - ) - is False - ) - assert ( - litellm.Router._is_deployment_blocked( - types.SimpleNamespace(model_info=types.SimpleNamespace(blocked=True)) - ) - is True - ) - - -class TestRouterRequestTimeoutPropagation: - """litellm_settings.request_timeout must act as an independent per-attempt timeout. - - Regression for LIT-2369: request_timeout was shadowed by router_settings.timeout, - so Bedrock (and other provider) calls fell back to the hardcoded 600s httpx - default instead of the configured value. - """ - - def _make_router(self, timeout=None, stream_timeout=None): - return litellm.Router( - model_list=[ - { - "model_name": "test-model", - "litellm_params": { - "model": "openai/gpt-4", - "api_key": "sk-test", - }, - } - ], - timeout=timeout, - stream_timeout=stream_timeout, - ) - - @pytest.fixture - def explicit_request_timeout(self): - original_value = litellm.request_timeout - original_flag = litellm.request_timeout_explicitly_set - litellm.request_timeout = 300 - litellm.request_timeout_explicitly_set = True - try: - yield 300 - finally: - litellm.request_timeout = original_value - litellm.request_timeout_explicitly_set = original_flag - - def test_request_timeout_stored_independently_when_both_set( - self, explicit_request_timeout - ): - router = self._make_router(timeout=330) - assert router.timeout == 330 - assert router.request_timeout == 300 - - def test_request_timeout_none_when_not_explicitly_configured(self): - original_value = litellm.request_timeout - original_flag = litellm.request_timeout_explicitly_set - litellm.request_timeout = litellm.constants.DEFAULT_REQUEST_TIMEOUT_SECONDS - litellm.request_timeout_explicitly_set = False - try: - router = self._make_router(timeout=330) - assert router.timeout == 330 - assert router.request_timeout is None - finally: - litellm.request_timeout = original_value - litellm.request_timeout_explicitly_set = original_flag - - def test_non_stream_prefers_request_timeout_over_router_timeout( - self, explicit_request_timeout - ): - router = self._make_router(timeout=330) - assert router._get_non_stream_timeout(kwargs={}, data={}) == 300 - - def test_stream_prefers_request_timeout_over_router_timeout( - self, explicit_request_timeout - ): - router = self._make_router(timeout=330) - # stream=True resolves through _get_stream_timeout; request_timeout must win. - assert router._get_timeout(kwargs={"stream": True}, data={}) == 300 - - def test_explicit_stream_timeout_still_wins_over_request_timeout( - self, explicit_request_timeout - ): - router = self._make_router(timeout=330, stream_timeout=45) - assert router._get_stream_timeout(kwargs={}, data={}) == 45 - - def test_non_stream_falls_through_to_router_timeout_without_request_timeout(self): - original_value = litellm.request_timeout - original_flag = litellm.request_timeout_explicitly_set - litellm.request_timeout = litellm.constants.DEFAULT_REQUEST_TIMEOUT_SECONDS - litellm.request_timeout_explicitly_set = False - try: - router = self._make_router(timeout=330) - assert router._get_non_stream_timeout(kwargs={}, data={}) == 330 - finally: - litellm.request_timeout = original_value - litellm.request_timeout_explicitly_set = original_flag - - def test_per_deployment_timeout_overrides_request_timeout( - self, explicit_request_timeout - ): - router = self._make_router(timeout=330) - assert router._get_non_stream_timeout(kwargs={}, data={"timeout": 120}) == 120 - - def test_per_request_timeout_overrides_request_timeout( - self, explicit_request_timeout - ): - router = self._make_router(timeout=330) - assert ( - router._get_non_stream_timeout( - kwargs={"timeout": 60}, data={"timeout": 120} - ) - == 60 - ) - - -# --------------------------------------------------------------------------- -# Deferred-stream eager-fetch tests -# --------------------------------------------------------------------------- - - -def _make_deferred_stream_wrapper(make_call_fn): - """Return a CustomStreamWrapper with completion_stream=None and the given make_call.""" - from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper - - logging_obj = MagicMock() - logging_obj.model_call_details = {"litellm_params": {}} - return CustomStreamWrapper( - completion_stream=None, - model="vertex_ai/gemini-2.0-flash", - logging_obj=logging_obj, - custom_llm_provider="vertex_ai_beta", - make_call=make_call_fn, - ) - - -def _make_router_with_vertex_and_fallback(): - return litellm.Router( - model_list=[ - { - "model_name": "my-gemini", - "litellm_params": { - "model": "vertex_ai/gemini-2.0-flash", - "vertex_project": "test-project", - "vertex_location": "us-central1", - }, - }, - { - "model_name": "my-fallback", - "litellm_params": { - "model": "openai/gpt-4o-mini", - "api_key": "sk-fake", - }, - }, - ], - fallbacks=[{"my-gemini": ["my-fallback"]}], - num_retries=0, - ) - - -@pytest.mark.asyncio -async def test_acompletion_deferred_stream_error_propagates_through_acompletion(): - """Regression: a deferred-stream CustomStreamWrapper whose make_call raises a 429 - must propagate the exception from within _acompletion's except block so that - fail_calls is incremented (i.e., deployment cooldown fires) and the standard - router fallback chain can handle it. - - Before the fix, the HTTP call happened inside __anext__ (outside the except block), - so fail_calls was never incremented. - """ - import litellm as _litellm - - rate_limit_err = _litellm.RateLimitError( - message="Resource exhausted", - llm_provider="vertex_ai", - model="gemini-2.0-flash", - ) - - async def failing_make_call(**kwargs): - raise rate_limit_err - - router = _make_router_with_vertex_and_fallback() - deferred_wrapper = _make_deferred_stream_wrapper(failing_make_call) - - with patch( - "litellm.acompletion", - new_callable=AsyncMock, - return_value=deferred_wrapper, - ): - with pytest.raises(_litellm.RateLimitError): - await router._acompletion( - model="vertex_ai/gemini-2.0-flash", - messages=[{"role": "user", "content": "Hello"}], - stream=True, - specific_deployment=router.model_list[0], - ) - - model_name = router.model_list[0]["litellm_params"]["model"] - assert router.fail_calls[model_name] == 1, ( - "fail_calls must be incremented when the deferred HTTP call fails; " - "without the eager fetch_stream() fix this stays at 0" - ) - - -@pytest.mark.asyncio -async def test_acompletion_deferred_stream_preserves_original_headers_on_error(): - """Router is used both by the proxy and directly as an SDK. HTTP-framing headers - (Content-Length, Transfer-Encoding, ...) must NOT be stripped at this layer, or - direct SDK callers lose legitimate provider metadata (e.g. content-type, - proxy-authenticate) that only the proxy's own response construction needs to - worry about. Stripping happens in the proxy layer instead - (_handle_llm_api_exception).""" - import litellm as _litellm - - err = _litellm.RateLimitError( - message="Resource exhausted", - llm_provider="vertex_ai", - model="gemini-2.0-flash", - ) - err.headers = { - "content-length": "42", - "transfer-encoding": "chunked", - "content-encoding": "gzip", - "content-type": "application/json", - "x-request-id": "abc-123", - } - - async def failing_make_call(**kwargs): - raise err - - router = _make_router_with_vertex_and_fallback() - deferred_wrapper = _make_deferred_stream_wrapper(failing_make_call) - - with patch( - "litellm.acompletion", - new_callable=AsyncMock, - return_value=deferred_wrapper, - ): - with pytest.raises(_litellm.RateLimitError) as exc_info: - await router._acompletion( - model="vertex_ai/gemini-2.0-flash", - messages=[{"role": "user", "content": "Hello"}], - stream=True, - specific_deployment=router.model_list[0], - ) - - raised = exc_info.value - headers = getattr(raised, "headers", {}) - assert headers.get("content-length") == "42" - assert headers.get("transfer-encoding") == "chunked" - assert headers.get("content-encoding") == "gzip" - assert headers.get("content-type") == "application/json" - assert headers.get("x-request-id") == "abc-123" - - -@pytest.mark.asyncio -async def test_acompletion_deferred_stream_skipped_when_stream_already_set(): - """When completion_stream is already populated (non-deferred provider), the eager - fetch_stream() call must be skipped entirely; no exception should be raised even - if make_call would fail. - """ - from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper - - async def would_fail(**kwargs): - raise RuntimeError("should not be called") - - logging_obj = MagicMock() - logging_obj.model_call_details = {"litellm_params": {}} - - async def noop_aiter(): - return - yield - - noop_stream = noop_aiter() - already_set_wrapper = CustomStreamWrapper( - completion_stream=noop_stream, - model="openai/gpt-4o", - logging_obj=logging_obj, - custom_llm_provider="openai", - make_call=would_fail, - ) - - router = litellm.Router( - model_list=[ - { - "model_name": "my-model", - "litellm_params": { - "model": "openai/gpt-4o", - "api_key": "sk-fake", - }, - } - ], - ) - - with patch( - "litellm.acompletion", - new_callable=AsyncMock, - return_value=already_set_wrapper, - ): - result = await router._acompletion( - model="openai/gpt-4o", - messages=[{"role": "user", "content": "Hello"}], - stream=True, - specific_deployment=router.model_list[0], - ) - - assert result is not None, "should return a streaming wrapper without errors" - assert already_set_wrapper.completion_stream is noop_stream, "completion_stream must not be re-fetched" - await noop_stream.aclose() - - -def test_completion_deferred_stream_error_propagates_through_completion(): - """Regression: the sync router path needs the same eager fetch as the async one. - - A deferred-stream CustomStreamWrapper hands back a wrapper whose HTTP call has - not happened yet, so without fetch_sync_stream() the provider error surfaces on - first iteration, outside _completion's except block. The deployment is then never - marked failed and function_with_fallbacks never sees the error. - """ - import litellm as _litellm - - rate_limit_err = _litellm.RateLimitError( - message="Resource exhausted", - llm_provider="vertex_ai", - model="gemini-2.0-flash", - ) - make_call_invocations = [] - - def failing_make_call(**kwargs): - make_call_invocations.append(kwargs) - raise rate_limit_err - - router = _make_router_with_vertex_and_fallback() - deferred_wrapper = _make_deferred_stream_wrapper(failing_make_call) - - with patch("litellm.completion", return_value=deferred_wrapper): - with pytest.raises(_litellm.RateLimitError): - router._completion( - model="vertex_ai/gemini-2.0-flash", - messages=[{"role": "user", "content": "Hello"}], - stream=True, - specific_deployment=router.model_list[0], - ) - - assert len(make_call_invocations) == 1, ( - "the deferred HTTP call must run inside _completion's try block; " - "without the eager fetch_sync_stream() fix it is deferred to first iteration" - ) - - -def test_completion_deferred_stream_skipped_when_stream_already_set(): - """A non-deferred sync provider already has completion_stream populated, so the - eager fetch must be skipped and make_call left untouched. - """ - from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper - - def would_fail(**kwargs): - raise RuntimeError("should not be called") - - logging_obj = MagicMock() - logging_obj.model_call_details = {"litellm_params": {}} - already_set_stream = iter([]) - - already_set_wrapper = CustomStreamWrapper( - completion_stream=already_set_stream, - model="openai/gpt-4o", - logging_obj=logging_obj, - custom_llm_provider="openai", - make_call=would_fail, - ) - - router = litellm.Router( - model_list=[ - { - "model_name": "my-model", - "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-fake"}, - } - ], - ) - - with patch("litellm.completion", return_value=already_set_wrapper): - result = router._completion( - model="openai/gpt-4o", - messages=[{"role": "user", "content": "Hello"}], - stream=True, - specific_deployment=router.model_list[0], - ) - - assert result is not None, "should return a streaming wrapper without errors" - assert already_set_wrapper.completion_stream is already_set_stream, "completion_stream must not be re-fetched" - - -class TestAdvisorSubCallCooldown: - """Regression for LIT-4565: an advisor orchestration failure must not cool - down the selected (healthy) deployment, which would reject unrelated - callers to the same model group.""" - - def _router(self): - return litellm.Router( - model_list=[ - { - "model_name": "claude-sonnet-5", - "litellm_params": {"model": "bedrock/us.anthropic.claude-opus-4-8"}, - "model_info": {"id": "dep-1"}, - } - ], - ) - - def _kwargs(self, exception): - return { - "exception": exception, - "litellm_params": {"model_info": {"id": "dep-1"}, "metadata": {}}, - } - - def _auth_error(self): - return litellm.AuthenticationError( - message="x-api-key header is required", - llm_provider="anthropic", - model="claude-opus-4-8", - ) - - def _cooled_down_ids(self, router): - active = router.cooldown_cache.get_active_cooldowns( - model_ids=["dep-1"], parent_otel_span=None - ) - return [entry[0] for entry in active] - - @pytest.mark.asyncio - async def test_untagged_auth_error_cools_down_deployment(self): - from datetime import datetime - - router = self._router() - now = datetime.now() - assert ( - router.deployment_callback_on_failure( - self._kwargs(self._auth_error()), None, now, now - ) - is True - ) - assert "dep-1" in self._cooled_down_ids(router) - - def test_advisor_orchestration_failure_does_not_cool_down_deployment(self): - from datetime import datetime - - from litellm.router_utils.cooldown_handlers import ( - mark_advisor_orchestration_failure, - ) - - router = self._router() - exception = self._auth_error() - mark_advisor_orchestration_failure(exception) - - now = datetime.now() - assert ( - router.deployment_callback_on_failure( - self._kwargs(exception), None, now, now - ) - is False - ) - assert "dep-1" not in self._cooled_down_ids(router) - - -def test_stream_chunks_have_generated_content_detects_text_and_non_text(): - from litellm.router import _stream_chunks_have_generated_content - from litellm.types.utils import ( - ChatCompletionDeltaToolCall, - Delta, - Function, - StreamingChoices, - ) - - def _chunk(delta): - return litellm.ModelResponseStream( - id="chatcmpl-1", - model="gpt-4", - object="chat.completion.chunk", - choices=[StreamingChoices(finish_reason=None, index=0, delta=delta)], - ) - - assert _stream_chunks_have_generated_content([]) is False - - empty_chunk = _chunk(Delta(role="assistant")) - assert _stream_chunks_have_generated_content([empty_chunk]) is False - - text_chunk = _chunk(Delta(content="Hello")) - assert _stream_chunks_have_generated_content([text_chunk]) is True - - reasoning_chunk = _chunk(Delta(reasoning_content="Thinking")) - assert _stream_chunks_have_generated_content([reasoning_chunk]) is True - - tool_call_delta = Delta( - tool_calls=[ - ChatCompletionDeltaToolCall( - id="call_1", - function=Function(name="get_weather", arguments="{}"), - type="function", - index=0, - ) - ] - ) - tool_call_chunk = _chunk(tool_call_delta) - assert _stream_chunks_have_generated_content([tool_call_chunk]) is True - - thinking_delta = Delta(thinking_blocks=[{"type": "thinking", "thinking": "Let me think..."}]) - thinking_chunk = _chunk(thinking_delta) - assert _stream_chunks_have_generated_content([thinking_chunk]) is True - - reasoning_items_delta = Delta(reasoning_items=[{"type": "reasoning", "id": "rs_1"}]) - reasoning_items_chunk = _chunk(reasoning_items_delta) - assert _stream_chunks_have_generated_content([reasoning_items_chunk]) is True - - audio_delta = Delta(audio={"data": "abc123", "expires_at": 1234567890, "transcript": "hello"}) - audio_chunk = _chunk(audio_delta) - assert _stream_chunks_have_generated_content([audio_chunk]) is True - - images_delta = Delta(images=[{"image_url": {"url": "https://example.com/img.png"}, "index": 0, "type": "image_url"}]) - images_chunk = _chunk(images_delta) - assert _stream_chunks_have_generated_content([images_chunk]) is True - - annotations_delta = Delta( - annotations=[{"type": "url_citation", "url_citation": {"url": "https://example.com"}}] - ) - annotations_chunk = _chunk(annotations_delta) - assert _stream_chunks_have_generated_content([annotations_chunk]) is True - - -def test_get_configured_token_limits_reads_deployment_model_info(): - router = litellm.Router( - model_list=[ - { - "model_name": "my-custom-model", - "litellm_params": {"model": "openai/some-unmapped-model"}, - "model_info": {"max_input_tokens": 32000, "max_output_tokens": 8000}, - } - ] - ) - - assert router.get_configured_token_limits("my-custom-model") == (32000, 8000) - - -def test_get_configured_token_limits_returns_none_for_unset_or_unknown(): - router = litellm.Router( - model_list=[ - { - "model_name": "no-limits-model", - "litellm_params": {"model": "openai/some-unmapped-model"}, - } - ] - ) - - assert router.get_configured_token_limits("no-limits-model") == (None, None) - assert router.get_configured_token_limits("not-a-real-model") == (None, None) - - -def test_get_configured_token_limits_skips_wildcard_pattern_matching(): - router = litellm.Router( - model_list=[ - { - "model_name": "bedrock/*", - "litellm_params": {"model": "bedrock/*"}, - "model_info": {"max_input_tokens": 12345}, - } - ] - ) - - with patch.object( - router.pattern_router, "route", side_effect=AssertionError("pattern route called") - ): - assert router.get_configured_token_limits( - "bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0" - ) == (None, None) - - -def test_get_configured_token_limits_treats_malformed_values_as_absent(): - malformed = ["", "unlimited", "128,000", [128000], {"max": 128000}, True] - router = litellm.Router( - model_list=[ - { - "model_name": f"bad-limit-{i}", - "litellm_params": {"model": "openai/some-unmapped-model"}, - "model_info": {"max_input_tokens": bad, "max_output_tokens": bad}, - } - for i, bad in enumerate(malformed) - ] - ) - - for i in range(len(malformed)): - assert router.get_configured_token_limits(f"bad-limit-{i}") == (None, None) - - -def test_get_configured_token_limits_coerces_numeric_strings(): - router = litellm.Router( - model_list=[ - { - "model_name": "quoted-limits-model", - "litellm_params": {"model": "openai/some-unmapped-model"}, - "model_info": {"max_input_tokens": "32000", "max_output_tokens": "8000"}, - } - ] - ) - - assert router.get_configured_token_limits("quoted-limits-model") == (32000, 8000) - - -def test_get_configured_mode_reads_deployment_model_info(): - router = litellm.Router( - model_list=[ - { - "model_name": "tts-model", - "litellm_params": {"model": "openai/some-unmapped-tts-model"}, - "model_info": {"mode": "audio_speech"}, - } - ] - ) - - assert router.get_configured_mode("tts-model") == "audio_speech" - - -def test_get_configured_mode_returns_none_for_unset_or_unknown(): - router = litellm.Router( - model_list=[ - { - "model_name": "no-mode-model", - "litellm_params": {"model": "openai/some-unmapped-model"}, - } - ] - ) - - assert router.get_configured_mode("no-mode-model") is None - assert router.get_configured_mode("not-a-real-model") is None - - -def test_get_configured_mode_skips_wildcard_pattern_matching(): - router = litellm.Router( - model_list=[ - { - "model_name": "bedrock/*", - "litellm_params": {"model": "bedrock/*"}, - "model_info": {"mode": "chat"}, - } - ] - ) - - with patch.object( - router.pattern_router, "route", side_effect=AssertionError("pattern route called") - ): - assert ( - router.get_configured_mode("bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0") - is None - ) - - -def test_get_configured_mode_treats_malformed_values_as_absent(): - malformed = ["", " ", 12345, ["chat"], {"mode": "chat"}, True] - router = litellm.Router( - model_list=[ - { - "model_name": f"bad-mode-{i}", - "litellm_params": {"model": "openai/some-unmapped-model"}, - "model_info": {"mode": bad}, - } - for i, bad in enumerate(malformed) - ] - ) - - for i in range(len(malformed)): - assert router.get_configured_mode(f"bad-mode-{i}") is None - - -def test_get_configured_display_name_reads_deployment_model_info(): - router = litellm.Router( - model_list=[ - { - "model_name": "Kimi K3-claude-compatible", - "litellm_params": {"model": "openai/some-unmapped-model"}, - "model_info": {"display_name": "Kimi K3"}, - } - ] - ) - - assert router.get_configured_display_name("Kimi K3-claude-compatible") == "Kimi K3" - - -def test_get_configured_display_name_returns_none_for_unset_or_unknown(): - router = litellm.Router( - model_list=[ - { - "model_name": "no-display-model", - "litellm_params": {"model": "openai/some-unmapped-model"}, - } - ] - ) - - assert router.get_configured_display_name("no-display-model") is None - assert router.get_configured_display_name("not-a-real-model") is None - - -def test_get_configured_display_name_skips_wildcard_pattern_matching(): - router = litellm.Router( - model_list=[ - { - "model_name": "bedrock/*", - "litellm_params": {"model": "bedrock/*"}, - "model_info": {"display_name": "Bedrock"}, - } - ] - ) - - with patch.object( - router.pattern_router, "route", side_effect=AssertionError("pattern route called") - ): - assert ( - router.get_configured_display_name("bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0") - is None - ) - - -def test_get_configured_display_name_treats_malformed_values_as_absent(): - malformed = ["", " ", 12345, ["Kimi K3"], {"name": "Kimi K3"}, True] - router = litellm.Router( - model_list=[ - { - "model_name": f"bad-display-{i}", - "litellm_params": {"model": "openai/some-unmapped-model"}, - "model_info": {"display_name": bad}, - } - for i, bad in enumerate(malformed) - ] - ) - - for i in range(len(malformed)): - assert router.get_configured_display_name(f"bad-display-{i}") is None - - -@pytest.mark.asyncio -async def test_acreate_batch_disable_fallbacks_surfaces_owning_provider_error(): - router = litellm.Router( - model_list=[ - { - "model_name": "owning-model", - "litellm_params": { - "model": "openai/gpt-4o-mini", - "api_key": "sk-owning", - }, - }, - { - "model_name": "fallback-model", - "litellm_params": { - "model": "azure/gpt-4o-mini", - "api_key": "sk-fallback", - "api_base": "https://fallback.openai.azure.com", - "api_version": "2024-08-01-preview", - }, - }, - ], - fallbacks=[{"owning-model": ["fallback-model"]}], - num_retries=0, - ) - owning_provider_error = litellm.BadRequestError( - message="completion_window must be one of: 24h", - model="openai/gpt-4o-mini", - llm_provider="openai", - ) - mock_create = AsyncMock(side_effect=owning_provider_error) - - with patch.object(router, "_acreate_batch", mock_create): - with pytest.raises(litellm.BadRequestError, match="24h"): - await router.acreate_batch( - model="owning-model", - input_file_id="file-owned-by-openai", - endpoint="/v1/chat/completions", - completion_window="5m", - disable_fallbacks=True, - ) - - mock_create.assert_awaited_once() - assert mock_create.call_args.kwargs["model"] == "owning-model" - - -@pytest.mark.asyncio -async def test_acreate_batch_surfaces_owning_provider_error_without_disable_fallbacks(): - """The router itself has to keep a batch inside the group that owns the input file: - the proxy only sets disable_fallbacks on the managed-files route, so the caller - otherwise gets the fallback provider's error for a file it never received.""" - from litellm.types.utils import LiteLLMBatch - - router = litellm.Router( - model_list=[ - { - "model_name": "owning-model", - "litellm_params": { - "model": "openai/gpt-4o-mini", - "api_key": "sk-owning", - }, - }, - { - "model_name": "fallback-model", - "litellm_params": { - "model": "azure/gpt-4o-mini", - "api_key": "sk-fallback", - "api_base": "https://fallback.openai.azure.com", - "api_version": "2024-08-01-preview", - }, - }, - ], - fallbacks=[{"owning-model": ["fallback-model"]}], - num_retries=0, - ) - attempted_models = [] - - async def _acreate_batch(model, **kwargs): - attempted_models.append(model) - if model == "owning-model": - raise litellm.APIConnectionError( - message="Connection error - openai is unreachable", - model="openai/gpt-4o-mini", - llm_provider="openai", - ) - return LiteLLMBatch( - id="batch-created-on-the-wrong-provider", - completion_window="24h", - created_at=0, - endpoint="/v1/chat/completions", - input_file_id="file-owned-by-openai", - object="batch", - status="validating", - ) - - with patch.object(router, "_acreate_batch", _acreate_batch): - with pytest.raises(litellm.APIConnectionError, match="openai is unreachable"): - await router.acreate_batch( - model="owning-model", - input_file_id="file-owned-by-openai", - endpoint="/v1/chat/completions", - completion_window="24h", - metadata={"team": "batch-jobs"}, - ) - - assert attempted_models == ["owning-model"] - - -@pytest.mark.asyncio -async def test_acreate_batch_still_falls_back_within_the_owning_model_group(): - """Holding a batch inside the model group that owns its input file must not - disable fallbacks outright (#35359): the owning group's second deployment is - still tried in `order`, and only the cross-group target is skipped.""" - completion_window_error = "Invalid value: '5m'. Supported values are: '24h'." - attempted_models = [] - - async def _acreate_batch(**kwargs): - model = kwargs["model"] - attempted_models.append(model) - if model.startswith("azure/"): - raise litellm.BadRequestError( - message="Error code: 400 - {'error': {'code': 'quotaExceeded'}}", - model=model, - llm_provider="azure", - ) - raise litellm.BadRequestError( - message=completion_window_error, - model=model, - llm_provider="openai", - ) - - router = litellm.Router( - model_list=[ - { - "model_name": "my-gpt", - "litellm_params": { - "model": "openai/gpt-4o-mini", - "api_key": "sk-owning", - "order": 1, - }, - "model_info": {"id": "my-gpt-1"}, - }, - { - "model_name": "my-gpt", - "litellm_params": { - "model": "openai/gpt-4o-mini-backup", - "api_key": "sk-owning", - "order": 2, - }, - "model_info": {"id": "my-gpt-2"}, - }, - { - "model_name": "my-azure-gpt", - "litellm_params": { - "model": "azure/gpt-4o-mini", - "api_key": "sk-fallback", - "api_base": "https://fallback.openai.azure.com", - "api_version": "2024-08-01-preview", - }, - "model_info": {"id": "my-azure-gpt-1"}, - }, - ], - fallbacks=[{"my-gpt": ["my-azure-gpt"]}], - num_retries=0, - ) - - with patch.object(litellm, "acreate_batch", new=_acreate_batch): - with pytest.raises(litellm.BadRequestError) as raised: - await router.acreate_batch( - model="my-gpt", - input_file_id="file-owned-by-my-gpt", - endpoint="/v1/chat/completions", - completion_window="5m", - ) - - assert "24h" in str(raised.value) - assert "quotaExceeded" not in str(raised.value) - assert attempted_models == ["openai/gpt-4o-mini", "openai/gpt-4o-mini-backup"] - - -@pytest.mark.asyncio -async def test_acreate_batch_request_bedrock_tags_override_deployment_tags(): - import httpx - - from litellm.llms.bedrock.common_utils import CommonBatchFilesUtils - - deployment_tags = [{"key": "application", "value": "config-level"}] - request_tags = [{"key": "application", "value": "request-level"}] - router = litellm.Router( - model_list=[ - { - "model_name": "bedrock-batch-model", - "litellm_params": { - "model": "bedrock/us.anthropic.claude-sonnet-5", - "aws_batch_role_arn": "arn:aws:iam::123:role/batch-role", - "aws_region_name": "us-west-2", - "bedrock_tags": deployment_tags, - }, - } - ] - ) - - def fake_response(): - return httpx.Response( - status_code=200, - json={ - "jobArn": "arn:aws:bedrock:us-west-2:123:model-invocation-job/abc1234567", - "status": "Submitted", - }, - ) - - mock_client = MagicMock() - mock_client.post = AsyncMock(side_effect=lambda *args, **kwargs: fake_response()) - - with patch.object( - CommonBatchFilesUtils, - "sign_aws_request", - return_value=({"Authorization": "signed"}, b"{}"), - ) as mock_sign, patch( - "litellm.llms.custom_httpx.llm_http_handler.get_async_httpx_client", - return_value=mock_client, - ): - await router.acreate_batch( - model="bedrock-batch-model", - input_file_id="s3://bucket/input.jsonl", - endpoint="/v1/chat/completions", - completion_window="24h", - ) - assert mock_sign.call_args.kwargs["data"]["tags"] == deployment_tags - - await router.acreate_batch( - model="bedrock-batch-model", - input_file_id="s3://bucket/input.jsonl", - endpoint="/v1/chat/completions", - completion_window="24h", - bedrock_tags=request_tags, - ) - assert mock_sign.call_args.kwargs["data"]["tags"] == request_tags - - -@pytest.mark.asyncio -async def test_avector_store_search_injects_router(): - """ - Regression: router.avector_store_search must pass the router down to the - SDK search call so provider transforms can resolve router-managed - embedding models (e.g. S3 Vectors query embeddings). - """ - from litellm.types.vector_stores import VectorStoreSearchResponse - - expected_response = VectorStoreSearchResponse( - object="vector_store.search_results.page", search_query="q", data=[] - ) - mock_asearch = AsyncMock(return_value=expected_response) - # Router.__init__ binds asearch via a local import, so patch the module - # attribute before constructing the Router. - with patch("litellm.vector_stores.main.asearch", new=mock_asearch): # test-quality-ok: the SDK call is the only place the injected router kwarg is observable - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "openai/gpt-3.5-turbo", "api_key": "test-key"}, - } - ] - ) - search_response = await router.avector_store_search( - vector_store_id="v", query="q", custom_llm_provider="s3_vectors" - ) - - assert search_response is expected_response - mock_asearch.assert_awaited_once() - assert mock_asearch.await_args.kwargs["router"] is router - - -@pytest.mark.asyncio -async def test_avector_store_create_does_not_inject_router(): - """The router injection is gated on the search call type: the create path - must keep calling the SDK without a router kwarg.""" - expected_response = {"id": "vs_1", "object": "vector_store"} - mock_acreate = AsyncMock(return_value=expected_response) - # avector_store_create(model=None) resolves acreate via a local import at - # call time, so patching after Router construction works here. - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "openai/gpt-3.5-turbo", "api_key": "test-key"}, - } - ] - ) - with patch("litellm.vector_stores.main.acreate", new=mock_acreate): # test-quality-ok: the SDK call is the only place a leaked router kwarg would surface - create_response = await router.avector_store_create(model=None, custom_llm_provider="openai") - - assert create_response is expected_response - mock_acreate.assert_awaited_once() - assert "router" not in mock_acreate.await_args.kwargs - - -def test_vector_store_search_injects_router(): - """ - Sync parity for the router injection: router.vector_store_search must pass - the router down to the SDK search call so provider transforms can resolve - router-managed embedding models, same as avector_store_search. - """ - from litellm.types.vector_stores import VectorStoreSearchResponse - - expected_response = VectorStoreSearchResponse( - object="vector_store.search_results.page", search_query="q", data=[] - ) - mock_search = MagicMock(return_value=expected_response) - # Router.__init__ binds search via a local import, so patch the module - # attribute before constructing the Router. - with patch("litellm.vector_stores.main.search", new=mock_search): # test-quality-ok: the SDK call is the only place the injected router kwarg is observable - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "openai/gpt-3.5-turbo", "api_key": "test-key"}, - } - ] - ) - search_response = router.vector_store_search( - vector_store_id="v", query="q", custom_llm_provider="s3_vectors" - ) - - assert search_response is expected_response - mock_search.assert_called_once() - assert mock_search.call_args.kwargs["router"] is router - assert mock_search.call_args.kwargs["custom_llm_provider"] == "s3_vectors" - - -def test_vector_store_create_does_not_inject_router(): - """The sync create path must keep calling the SDK without a router kwarg.""" - expected_response = {"id": "vs_1", "object": "vector_store"} - mock_create = MagicMock(return_value=expected_response) - # Router.__init__ binds create via a local import, so patch the module - # attribute before constructing the Router. - with patch("litellm.vector_stores.main.create", new=mock_create): # test-quality-ok: the SDK call is the only place a leaked router kwarg would surface - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "openai/gpt-3.5-turbo", "api_key": "test-key"}, - } - ] - ) - create_response = router.vector_store_create(custom_llm_provider="openai") - - assert create_response is expected_response - mock_create.assert_called_once() - assert "router" not in mock_create.call_args.kwargs - - -class TestPreRoutingStrategyRegistryLifecycle: - """ - Regression tests: a deployment leaving the model_list must release the - pre-routing strategy slot it holds in `auto_routers` / `complexity_routers` / - `adaptive_routers` / `quality_routers`. - - Before this fix, editing an auto-router-family model (a UI save, which reaches - every other pod as an `upsert_deployment` from the periodic DB reload) popped - the deployment out of the model_list and then failed to re-add it: registration - raised "already exists" against the stale registry entry, and - `ignore_invalid_deployments=True` swallowed the error. The router vanished from - the Models page and stayed gone until a proxy restart, while the DB row and the - "saved successfully" response both looked fine. - """ - - @staticmethod - def _complexity_router_params(default_model: str, tags=None) -> dict: - return { - "model": "auto_router/complexity_router", - "complexity_router_config": { - "tiers": {"SIMPLE": "gpt-4o-mini", "MEDIUM": "gpt-4o", "COMPLEX": "gpt-4o"} - }, - "complexity_router_default_model": default_model, - **({"tags": tags} if tags else {}), - } - - @classmethod - def _router_with_complexity_router(cls, default_model: str = "gpt-4o") -> "litellm.Router": - return litellm.Router( - model_list=[ - {"model_name": "gpt-4o", "litellm_params": {"model": "gpt-4o"}}, - {"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}}, - { - "model_name": "smart-router", - "litellm_params": cls._complexity_router_params(default_model), - "model_info": {"id": "router-1", "db_model": True}, - }, - ], - ignore_invalid_deployments=True, - ) - - @staticmethod - def _model_names(router: "litellm.Router") -> list: - return [model["model_name"] for model in router.model_list] - - def test_upsert_of_edited_router_keeps_it_routable(self): - from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo - - router = self._router_with_complexity_router() - - router.upsert_deployment( - deployment=Deployment( - model_name="smart-router", - litellm_params=LiteLLM_Params(**self._complexity_router_params("gpt-4o-mini")), - model_info=ModelInfo(id="router-1", db_model=True), - ) - ) - - assert "smart-router" in self._model_names(router) - registered = router.complexity_routers["smart-router"] - assert len(registered) == 1 - # the surviving strategy is the edited one, not the pre-edit leftover - assert registered[0].strategy.config.default_model == "gpt-4o-mini" - - def test_unchanged_upsert_leaves_router_untouched(self): - from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo - - router = self._router_with_complexity_router() - strategy_before = router.complexity_routers["smart-router"][0].strategy - - for _ in range(3): - router.upsert_deployment( - deployment=Deployment( - model_name="smart-router", - litellm_params=LiteLLM_Params(**self._complexity_router_params("gpt-4o")), - model_info=ModelInfo(id="router-1", db_model=True), - ) - ) - - assert "smart-router" in self._model_names(router) - assert router.complexity_routers["smart-router"][0].strategy is strategy_before - - def test_delete_frees_the_name_for_a_new_router(self): - from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo - - router = self._router_with_complexity_router() - - router.delete_deployment(id="router-1") - assert "smart-router" not in router.complexity_routers - - router.add_deployment( - deployment=Deployment( - model_name="smart-router", - litellm_params=LiteLLM_Params(**self._complexity_router_params("gpt-4o-mini")), - model_info=ModelInfo(id="router-2", db_model=True), - ) - ) - - assert "smart-router" in self._model_names(router) - assert router.complexity_routers["smart-router"][0].strategy.config.default_model == "gpt-4o-mini" - - def test_delete_only_frees_the_matching_tag_slot(self): - router = litellm.Router( - model_list=[ - {"model_name": "gpt-4o", "litellm_params": {"model": "gpt-4o"}}, - {"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}}, - { - "model_name": "shared-router", - "litellm_params": self._complexity_router_params("gpt-4o", tags=["team-a"]), - "model_info": {"id": "router-a"}, - }, - { - "model_name": "shared-router", - "litellm_params": self._complexity_router_params("gpt-4o-mini", tags=["team-b"]), - "model_info": {"id": "router-b"}, - }, - ], - ignore_invalid_deployments=True, - ) - assert len(router.complexity_routers["shared-router"]) == 2 - - router.delete_deployment(id="router-a") - - remaining = router.complexity_routers["shared-router"] - assert len(remaining) == 1 - assert remaining[0].tags == ("team-b",) - - def test_delete_of_regular_model_preserves_router_sharing_its_name(self): - router = litellm.Router( - model_list=[ - {"model_name": "gpt-4o", "litellm_params": {"model": "gpt-4o"}}, - {"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}}, - { - "model_name": "shared-name", - "litellm_params": self._complexity_router_params("gpt-4o"), - "model_info": {"id": "router-1"}, - }, - { - "model_name": "shared-name", - "litellm_params": {"model": "openai/gpt-4o"}, - "model_info": {"id": "regular-1"}, - }, - ], - ignore_invalid_deployments=True, - ) - strategy = router.complexity_routers["shared-name"][0].strategy - - router.delete_deployment(id="regular-1") - - assert router.complexity_routers["shared-name"][0].strategy is strategy - - def test_upsert_of_edited_adaptive_router_rebuilds_it(self): - """Adaptive routers are built by set_model_list()'s deferred pass, not by - add_deployment(), so releasing the slot on edit must be paired with a rebuild - - otherwise the edit silently turns adaptive routing off.""" - from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo - - def adaptive_params(available_models: list) -> dict: - return { - "model": "auto_router/adaptive_router", - "adaptive_router_config": {"available_models": available_models}, - } - - router = litellm.Router( - model_list=[ - {"model_name": "gpt-4o", "litellm_params": {"model": "openai/gpt-4o"}}, - {"model_name": "gpt-4o-mini", "litellm_params": {"model": "openai/gpt-4o-mini"}}, - { - "model_name": "adaptive-router", - "litellm_params": adaptive_params(["gpt-4o-mini"]), - "model_info": {"id": "router-1", "db_model": True}, - }, - ], - ignore_invalid_deployments=True, - ) - assert "adaptive-router" in router.adaptive_routers - - router.upsert_deployment( - deployment=Deployment( - model_name="adaptive-router", - litellm_params=LiteLLM_Params(**adaptive_params(["gpt-4o", "gpt-4o-mini"])), - model_info=ModelInfo(id="router-1", db_model=True), - ) - ) - - assert "adaptive-router" in self._model_names(router) - registered = router.adaptive_routers["adaptive-router"] - assert len(registered) == 1 - assert set(registered[0].strategy.config.available_models) == {"gpt-4o", "gpt-4o-mini"} - - def test_delete_repairs_indices_even_when_strategy_release_fails(self): - """Structural removal and strategy release are not equally critical. Once the entry - leaves model_list the index maps must be repaired no matter what, so releasing the - registry slot runs after that repair and cannot abandon the router half-updated.""" - router = self._router_with_complexity_router() - idx = router.model_id_to_deployment_index_map["router-1"] - router.model_list[idx] = {"model_name": "smart-router", "litellm_params": None} - - returned = router.delete_deployment(id="router-1") - - assert returned is not None - assert "router-1" not in router.model_id_to_deployment_index_map - assert all(entry.get("model_info", {}).get("id") != "router-1" for entry in router.model_list) - assert router.get_deployment(model_id="router-1") is None - assert "gpt-4o" in self._model_names(router) - - def test_delete_of_adaptive_enabled_complexity_router_frees_both_registries(self): - """A complexity router with adaptive set is registered in BOTH complexity_routers - and adaptive_routers under the same (model_name, tags). Releasing only the first - match leaves the adaptive strategy live, so a deleted alias stays routable and its - post-call hook keeps recording.""" - import litellm as litellm_module - from litellm.router_strategy.adaptive_router.hooks import AdaptiveRouterPostCallHook - - params = { - "model": "auto_router/complexity_router", - "complexity_router_config": { - "tiers": {"SIMPLE": "gpt-4o-mini", "MEDIUM": "gpt-4o"}, - "adaptive": True, - }, - "complexity_router_default_model": "gpt-4o", - } - router = litellm.Router( - model_list=[ - {"model_name": "gpt-4o", "litellm_params": {"model": "openai/gpt-4o"}}, - {"model_name": "gpt-4o-mini", "litellm_params": {"model": "openai/gpt-4o-mini"}}, - { - "model_name": "hybrid-router", - "litellm_params": params, - "model_info": {"id": "router-1", "db_model": True}, - }, - ], - ignore_invalid_deployments=True, - ) - assert "hybrid-router" in router.complexity_routers - assert "hybrid-router" in router.adaptive_routers - hooks = litellm_module.logging_callback_manager.get_custom_loggers_for_type(AdaptiveRouterPostCallHook) - assert len(hooks) == 1 - - router.delete_deployment(id="router-1") - - assert "hybrid-router" not in router.complexity_routers - assert "hybrid-router" not in router.adaptive_routers - remaining_hooks = litellm_module.logging_callback_manager.get_custom_loggers_for_type( - AdaptiveRouterPostCallHook - ) - assert remaining_hooks == [] - - def test_upsert_of_edited_quality_router_keeps_it_routable(self): - """_unregister_pre_routing_strategy_for_deployment dispatches on four prefixes; - quality_router is one of them and would otherwise go unexercised.""" - from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo - - def quality_params(default_model: str) -> dict: - return { - "model": "auto_router/quality_router", - "quality_router_default_model": default_model, - } - - router = litellm.Router( - model_list=[ - {"model_name": "gpt-4o", "litellm_params": {"model": "openai/gpt-4o"}}, - {"model_name": "gpt-4o-mini", "litellm_params": {"model": "openai/gpt-4o-mini"}}, - { - "model_name": "quality-router", - "litellm_params": quality_params("gpt-4o"), - "model_info": {"id": "router-1", "db_model": True}, - }, - ], - ignore_invalid_deployments=True, - ) - assert "quality-router" in router.quality_routers - - router.upsert_deployment( - deployment=Deployment( - model_name="quality-router", - litellm_params=LiteLLM_Params(**quality_params("gpt-4o-mini")), - model_info=ModelInfo(id="router-1", db_model=True), - ) - ) - - assert "quality-router" in self._model_names(router) - registered = router.quality_routers["quality-router"] - assert len(registered) == 1 - assert registered[0].strategy.config.default_model == "gpt-4o-mini" - - @staticmethod - def _hybrid_router_params(tiers: dict) -> dict: - return { - "model": "auto_router/complexity_router", - "complexity_router_config": {"tiers": tiers, "adaptive": True}, - "complexity_router_default_model": "gpt-4o", - } - - @classmethod - def _router_with_hybrid_router(cls) -> "litellm.Router": - return litellm.Router( - model_list=[ - {"model_name": "gpt-4o", "litellm_params": {"model": "openai/gpt-4o"}}, - {"model_name": "gpt-4o-mini", "litellm_params": {"model": "openai/gpt-4o-mini"}}, - { - "model_name": "hybrid-router", - "litellm_params": cls._hybrid_router_params({"SIMPLE": "gpt-4o-mini", "MEDIUM": "gpt-4o"}), - "model_info": {"id": "router-1", "db_model": True}, - }, - ], - ignore_invalid_deployments=True, - ) - - def test_upsert_of_edited_hybrid_complexity_router_relinks_adaptive(self): - """Editing an adaptive-enabled complexity router releases its adaptive companion - along with the complexity slot; the finalize re-run must fire for it (not just for - `auto_router/adaptive_router` deployments) or the rebuilt complexity router keeps - routing while bandit recording, DB persistence and /adaptive_router/state all - silently stop until the next full reload.""" - import litellm as litellm_module - from litellm.router_strategy.adaptive_router.hooks import AdaptiveRouterPostCallHook - from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo - - router = self._router_with_hybrid_router() - assert "hybrid-router" in router.adaptive_routers - - router.upsert_deployment( - deployment=Deployment( - model_name="hybrid-router", - litellm_params=LiteLLM_Params( - **self._hybrid_router_params( - {"SIMPLE": "gpt-4o-mini", "MEDIUM": "gpt-4o", "COMPLEX": "gpt-4o"} - ) - ), - model_info=ModelInfo(id="router-1", db_model=True), - ) - ) - - assert "hybrid-router" in self._model_names(router) - assert "hybrid-router" in router.complexity_routers - assert "hybrid-router" in router.adaptive_routers - rebuilt = router.complexity_routers["hybrid-router"][0].strategy - assert router.adaptive_routers["hybrid-router"][0].strategy is rebuilt.adaptive_router - hooks = litellm_module.logging_callback_manager.get_custom_loggers_for_type(AdaptiveRouterPostCallHook) - assert len(hooks) == 1 - - def test_upsert_turning_adaptive_on_builds_the_companion(self): - """An edit that flips `adaptive: true` on an existing complexity router must - register the companion immediately; neither side of the old prefix-only gate - matches a complexity deployment, so the flip was a silent no-op until restart.""" - from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo - - router = self._router_with_complexity_router() - assert "smart-router" not in router.adaptive_routers - - params = self._complexity_router_params("gpt-4o") - params["complexity_router_config"] = {**params["complexity_router_config"], "adaptive": True} - router.upsert_deployment( - deployment=Deployment( - model_name="smart-router", - litellm_params=LiteLLM_Params(**params), - model_info=ModelInfo(id="router-1", db_model=True), - ) - ) - - assert "smart-router" in router.adaptive_routers - - def test_unregister_pre_routing_strategy_scopes_the_drop_by_tags(self): - """The bool return drives the hook re-sync; a tag mismatch must report False and - leave the registry untouched, and dropping the last entry must free the key.""" - from litellm.types.router import TaggedPreRoutingStrategy - - registry = { - "m": [ - TaggedPreRoutingStrategy(tags=("team-a",), strategy=object()), - TaggedPreRoutingStrategy(tags=(), strategy=object()), - ] - } - - assert litellm.Router._unregister_pre_routing_strategy(registry, "m", ("team-b",)) is False - assert len(registry["m"]) == 2 - - assert litellm.Router._unregister_pre_routing_strategy(registry, "m", ("team-a",)) is True - assert [entry.tags for entry in registry["m"]] == [()] - - assert litellm.Router._unregister_pre_routing_strategy(registry, "m", ()) is True - assert "m" not in registry - - def test_unregister_for_deployment_ignores_non_router_deployments(self): - """Direct twin of the endpoint-level test: a regular deployment that shares a - router's model_name must not evict the router's registry slot.""" - from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo - - router = self._router_with_complexity_router() - - router._unregister_pre_routing_strategy_for_deployment( - deployment=Deployment( - model_name="smart-router", - litellm_params=LiteLLM_Params(model="openai/gpt-4o"), - model_info=ModelInfo(id="plain-1", db_model=True), - ) - ) - - assert "smart-router" in router.complexity_routers - - def test_sync_adaptive_router_hooks_keeps_one_hook_per_registered_router(self): - """Re-syncing must replace, not accumulate: a duplicated hook double-fires - bandit signal recording for every request.""" - import litellm as litellm_module - from litellm.router_strategy.adaptive_router.hooks import AdaptiveRouterPostCallHook - - router = self._router_with_hybrid_router() - - router._sync_adaptive_router_hooks() - router._sync_adaptive_router_hooks() - - hooks = litellm_module.logging_callback_manager.get_custom_loggers_for_type(AdaptiveRouterPostCallHook) - assert len(hooks) == 1 - - def test_deployment_participates_in_adaptive_routing_matrix(self): - """The upsert finalize re-run keys off this predicate for both the incoming and - outgoing deployment; a false negative silently strands the adaptive companion.""" - from litellm.types.router import LiteLLM_Params - - router = self._router_with_complexity_router() - - cases = [ - ({"model": "auto_router/adaptive_router", "adaptive_router_config": {}}, True), - (self._hybrid_router_params({"SIMPLE": "gpt-4o-mini"}), True), - (self._complexity_router_params("gpt-4o"), False), - ( - { - "model": "auto_router/complexity_router", - "complexity_router_config": {"tiers": {"SIMPLE": "gpt-4o-mini"}, "adaptive": False}, - "complexity_router_default_model": "gpt-4o", - }, - False, - ), - ({"model": "openai/gpt-4o"}, False), - ] - for params, expected in cases: - actual = router._deployment_participates_in_adaptive_routing( - litellm_params=LiteLLM_Params(**params) - ) - assert actual is expected, params["model"] - - -def test_model_info_is_active_for_environment_matrix(monkeypatch): - """The model-write endpoints consult this predicate to tell a deliberately - environment-inactive model from one dropped by a failed reload; the Router's own - deployment gate delegates to it, so the two can never diverge.""" - from litellm.router import model_info_is_active_for_environment - - assert model_info_is_active_for_environment(model_info=None) is True - assert model_info_is_active_for_environment(model_info={"id": "m1"}) is True - assert model_info_is_active_for_environment(model_info={"supported_environments": None}) is True - - monkeypatch.setenv("LITELLM_ENVIRONMENT", "development") - assert model_info_is_active_for_environment(model_info={"supported_environments": ["development"]}) is True - assert model_info_is_active_for_environment(model_info={"supported_environments": ["production"]}) is False - - monkeypatch.delenv("LITELLM_ENVIRONMENT") - with pytest.raises(ValueError, match="LITELLM_ENVIRONMENT"): - model_info_is_active_for_environment(model_info={"supported_environments": ["production"]}) - - -def test_pre_call_checks_uses_deployment_model_when_model_info_lookup_raises(monkeypatch): - """ - The supported-params check must run against the deployment's own - provider-qualified model. Resolving the per-deployment model only after the - model-info lookup leaves it unset whenever that lookup raises (an - unregistered custom model), so the check falls back to the bare model group - name and the request dies with 'LLM Provider NOT provided'. - """ - monkeypatch.setattr(litellm, "drop_params", False) - - router = litellm.Router( - model_list=[ - { - "model_name": "custom-alias", - "litellm_params": {"model": "hosted_vllm/not-in-the-catalog"}, - } - ], - enable_pre_call_checks=True, - ) - - def _raise_unmapped(**kwargs): - raise ValueError("This model isn't mapped yet") - - monkeypatch.setattr(router, "get_router_model_info", _raise_unmapped) - - seen: list[tuple] = [] - original_get_supported_openai_params = litellm.get_supported_openai_params - - def _record(model, custom_llm_provider=None, **kwargs): - seen.append((model, custom_llm_provider)) - return original_get_supported_openai_params(model=model, custom_llm_provider=custom_llm_provider, **kwargs) - - monkeypatch.setattr(litellm, "get_supported_openai_params", _record) - - deployments = [ - { - "litellm_params": {"model": "hosted_vllm/not-in-the-catalog"}, - "model_info": {"id": "d1"}, - } - ] - result = router._pre_call_checks( - model="custom-alias", - healthy_deployments=deployments, - messages=[{"role": "user", "content": "hi"}], - request_kwargs={}, - ) - - assert len(result) == 1 - assert seen == [("not-in-the-catalog", "hosted_vllm")] - - -def test_pre_call_checks_keeps_deployment_when_provider_is_unresolvable(monkeypatch): - """ - Pre-call checks filter deployments; they must never be the thing that fails - a request. A deployment whose provider cannot be resolved simply skips the - supported-params check instead of raising out of deployment selection. - """ - monkeypatch.setattr(litellm, "drop_params", False) - - router = litellm.Router( - model_list=[ - { - "model_name": "custom-alias", - "litellm_params": {"model": "gpt-3.5-turbo"}, - } - ], - enable_pre_call_checks=True, - ) - - def _raise_no_provider(**kwargs): - raise litellm.BadRequestError( - message="LLM Provider NOT provided.", - model="custom-alias", - llm_provider="", - ) - - monkeypatch.setattr(litellm, "get_llm_provider", _raise_no_provider) - - deployments = [ - { - "litellm_params": {"model": "some-unresolvable-model"}, - "model_info": {"id": "d1"}, - } - ] - result = router._pre_call_checks( - model="custom-alias", - healthy_deployments=deployments, - messages=[{"role": "user", "content": "hi"}], - request_kwargs={}, - ) - - assert len(result) == 1 - - -class TestUpsertDeploymentRollback: - """ - Regression tests: `upsert_deployment` pops the previous deployment before - re-adding the edited one. When the re-add raises under - `ignore_invalid_deployments=True`, the pop must be rolled back so this pod - keeps serving the previous configuration instead of silently dropping a live - deployment (the "Error upserting deployment" drop behind the access-group - reload 500 in the 2-replica e2e suite). - """ - - def test_failed_upsert_keeps_previous_deployment_serving(self): - from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo - - router = litellm.Router( - model_list=[ - { - "model_name": "prod-model", - "litellm_params": {"model": "gpt-4o", "api_key": "sk-old"}, - "model_info": {"id": "prod-1", "db_model": True}, - } - ], - ignore_invalid_deployments=True, - ) - - result = router.upsert_deployment( - deployment=Deployment( - model_name="prod-model", - litellm_params=LiteLLM_Params(model="auto_router/broken"), - model_info=ModelInfo(id="prod-1", db_model=True), - ) - ) - - assert result is None - restored = router.get_deployment(model_id="prod-1") - assert restored is not None - assert restored.litellm_params.model == "gpt-4o" - assert [model["model_name"] for model in router.model_list] == ["prod-model"] - - def test_failed_fresh_add_returns_none_without_restore(self): - from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo - - router = litellm.Router(model_list=[], ignore_invalid_deployments=True) - - result = router.upsert_deployment( - deployment=Deployment( - model_name="fresh-router", - litellm_params=LiteLLM_Params(model="auto_router/broken"), - model_info=ModelInfo(id="fresh-1", db_model=True), - ) - ) - - assert result is None - assert router.get_deployment(model_id="fresh-1") is None - assert router.model_list == [] - - def test_restore_re_adds_popped_deployment(self): - router = litellm.Router( - model_list=[ - { - "model_name": "prod-model", - "litellm_params": {"model": "gpt-4o", "api_key": "sk-old"}, - "model_info": {"id": "prod-1", "db_model": True}, - } - ], - ignore_invalid_deployments=True, - ) - previous = router.get_deployment(model_id="prod-1") - router.delete_deployment(id="prod-1") - assert router.has_model_id("prod-1") is False - - router._restore_deployment_after_failed_upsert( - previous_deployment=previous, model_id="prod-1" - ) - - restored = router.get_deployment(model_id="prod-1") - assert restored is not None - assert restored.litellm_params.model == "gpt-4o" - - router._restore_deployment_after_failed_upsert( - previous_deployment=previous, model_id="prod-1" - ) - assert len(router.model_list) == 1 - - router._restore_deployment_after_failed_upsert( - previous_deployment=None, model_id="prod-1" - ) - assert len(router.model_list) == 1 - - -class TestUpsertDeploymentRename: - """ - Issue #38360: renaming a model wrote the new `model_name` to the db, but the reload's - `upsert_deployment` compared only `litellm_params` and `model_info`. A rename with no - other edit therefore compared equal and the router kept the old name until a restart, - so `/model/info` and `/v1/models` served the stale name and the new one was unroutable. - """ - - @staticmethod - def _router() -> "litellm.Router": - return litellm.Router( - model_list=[ - { - "model_name": "old-name", - "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"}, - "model_info": {"id": "rename-1", "db_model": True}, - } - ] - ) - - @staticmethod - def _deployment(model_name: str, tpm: int | None = None): - from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo - - return Deployment( - model_name=model_name, - litellm_params=LiteLLM_Params(model="openai/gpt-4o", api_key="sk-test", tpm=tpm), - model_info=ModelInfo(id="rename-1", db_model=True), - ) - - def test_rename_only_updates_the_router(self): - router = self._router() - - assert router.upsert_deployment(deployment=self._deployment("new-name")) is not None - - assert [model["model_name"] for model in router.model_list] == ["new-name"] - renamed = router.get_deployment(model_id="rename-1") - assert renamed is not None - assert renamed.model_name == "new-name" - - def test_rename_only_makes_the_new_name_routable(self): - router = self._router() - - router.upsert_deployment(deployment=self._deployment("new-name")) - - assert router.get_model_ids(model_name="new-name") == ["rename-1"] - assert router.get_model_ids(model_name="old-name") == [] - - def test_rename_alongside_another_edit_still_updates(self): - router = self._router() - - router.upsert_deployment(deployment=self._deployment("new-name", tpm=1234)) - - assert router.get_model_ids(model_name="new-name") == ["rename-1"] - renamed = router.get_deployment(model_id="rename-1") - assert renamed is not None - assert renamed.litellm_params.tpm == 1234 - - def test_unchanged_deployment_is_still_a_no_op(self): - router = self._router() - - assert router.upsert_deployment(deployment=self._deployment("old-name")) is None - assert [model["model_name"] for model in router.model_list] == ["old-name"] - - -class TestConsumedRequestTagsStamp: - """Issue #36621: when a request's tags select a tagged pre-routing strategy, those - tags are consumed by the selection; the hook must stamp the rewritten model group so - tag filtering skips request-body tags there, and must clear the stamp on every - re-entry (fallbacks reuse the same request_kwargs) so it cannot leak elsewhere.""" - - class _RewriteStrategy: - def __init__(self, rewrite_to: str): - self.rewrite_to = rewrite_to - - async def async_pre_routing_hook( - self, model, request_kwargs, messages=None, input=None, specific_deployment=False - ): - from litellm.types.router import PreRoutingHookResponse - - return PreRoutingHookResponse(model=self.rewrite_to, messages=messages) - - @classmethod - def _router(cls, marker_tags=("route",)) -> "litellm.Router": - from litellm.types.router import TaggedPreRoutingStrategy - - router = litellm.Router( - model_list=[ - {"model_name": "gpt4o", "litellm_params": {"model": "openai/gpt-4o"}}, - {"model_name": "gemini-flash", "litellm_params": {"model": "gemini/gemini-3.6-flash"}}, - ], - enable_tag_filtering=True, - ) - router.auto_routers = { - "gpt4o": [TaggedPreRoutingStrategy(tags=marker_tags, strategy=cls._RewriteStrategy("gemini-flash"))] - } - return router - - @pytest.mark.asyncio - async def test_stamps_the_rewritten_group_when_request_tags_selected_the_router(self): - from litellm.constants import CONSUMED_REQUEST_TAGS_METADATA_KEY - from litellm.types.router import ConsumedRequestTagsStamp - - router = self._router() - request_kwargs = {"metadata": {"tags": ["route"]}} - - await router.async_pre_routing_hook(model="gpt4o", request_kwargs=request_kwargs) - - assert request_kwargs["metadata"][CONSUMED_REQUEST_TAGS_METADATA_KEY] == ConsumedRequestTagsStamp( - model_group="gemini-flash", tags=("route",) - ) - - @pytest.mark.asyncio - async def test_stamps_into_litellm_metadata_when_the_request_uses_that_bucket(self): - from litellm.constants import CONSUMED_REQUEST_TAGS_METADATA_KEY - from litellm.types.router import ConsumedRequestTagsStamp - - router = self._router() - request_kwargs = {"litellm_metadata": {"tags": ["route"]}} - - await router.async_pre_routing_hook(model="gpt4o", request_kwargs=request_kwargs) - - assert request_kwargs["litellm_metadata"][CONSUMED_REQUEST_TAGS_METADATA_KEY] == ConsumedRequestTagsStamp( - model_group="gemini-flash", tags=("route",) - ) - - @pytest.mark.asyncio - async def test_fallback_reentry_with_a_plain_group_clears_the_stale_stamp(self): - from litellm.constants import CONSUMED_REQUEST_TAGS_METADATA_KEY - - router = self._router() - request_kwargs = {"metadata": {"tags": ["route"]}} - - await router.async_pre_routing_hook(model="gpt4o", request_kwargs=request_kwargs) - await router.async_pre_routing_hook(model="gemini-flash", request_kwargs=request_kwargs) - - assert CONSUMED_REQUEST_TAGS_METADATA_KEY not in request_kwargs["metadata"] - - @pytest.mark.asyncio - async def test_no_stamp_when_the_request_is_untagged(self): - from litellm.constants import CONSUMED_REQUEST_TAGS_METADATA_KEY - - router = self._router() - request_kwargs = {"metadata": {}} - - await router.async_pre_routing_hook(model="gpt4o", request_kwargs=request_kwargs) - - assert CONSUMED_REQUEST_TAGS_METADATA_KEY not in request_kwargs["metadata"] - - @pytest.mark.asyncio - async def test_no_stamp_when_the_selected_strategy_carries_no_tags(self): - from litellm.constants import CONSUMED_REQUEST_TAGS_METADATA_KEY - - router = self._router(marker_tags=()) - request_kwargs = {"metadata": {"tags": ["route"]}} - - await router.async_pre_routing_hook(model="gpt4o", request_kwargs=request_kwargs) - - assert CONSUMED_REQUEST_TAGS_METADATA_KEY not in request_kwargs["metadata"] - - -class TestClaudeCodeSubagentSessionRouterBinding: - class _RewriteStrategy: - def __init__(self, routed_model: str = "cheap-model") -> None: - self.routed_model = routed_model - - async def async_pre_routing_hook( - self, model, request_kwargs, messages=None, input=None, specific_deployment=False - ): - from litellm.types.router import PreRoutingHookResponse - - return PreRoutingHookResponse( - model=self.routed_model, - messages=messages, - routing_decision={ - "router_model_name": "smart-router", - "router_type": "complexity", - "routed_model": self.routed_model, - "cause": "heuristic_scorer", - }, - ) - - @classmethod - def _router( - cls, - cheap_response: str = "cheap response", - fallbacks: list[dict[str, list[str]]] | None = None, - ) -> "litellm.Router": - from litellm.types.router import TaggedPreRoutingStrategy - - router = litellm.Router( - model_list=[ - { - "model_name": "cheap-model", - "litellm_params": {"model": "openai/gpt-4o-mini", "mock_response": cheap_response}, - }, - { - "model_name": "expensive-model", - "litellm_params": {"model": "openai/gpt-4o", "mock_response": "expensive response"}, - }, - ], - fallbacks=fallbacks, - num_retries=0, - ) - router.complexity_routers = { - "smart-router": (TaggedPreRoutingStrategy(tags=(), strategy=cls._RewriteStrategy()),), - "premium-router": (TaggedPreRoutingStrategy(tags=(), strategy=cls._RewriteStrategy("expensive-model")),), - } - return router - - @staticmethod - def _request_kwargs( - *, - key_hash: str = "key-hash-a", - app: str = "cli", - agent_id: str | None = None, - fallback_depth: int | None = None, - ) -> dict: - headers = { - "X-Claude-Code-Session-Id": "session-1234", - "x-app": app, - **({"x-claude-code-agent-id": agent_id} if agent_id is not None else {}), - } - return { - "metadata": {"user_api_key_hash": key_hash}, - "proxy_server_request": {"headers": headers}, - **({"fallback_depth": fallback_depth} if fallback_depth is not None else {}), - } - - @pytest.mark.asyncio - async def test_subagent_concrete_model_uses_the_main_sessions_router(self): - router = self._router() - - await router.acompletion( - model="smart-router", - messages=[{"role": "user", "content": "main turn"}], - **self._request_kwargs(), - ) - subagent_kwargs = self._request_kwargs(agent_id="agent-1234") - - response = await router.acompletion( - model="expensive-model", - messages=[{"role": "user", "content": "subagent turn"}], - **subagent_kwargs, - ) - - assert response.choices[0].message.content == "cheap response" - assert subagent_kwargs["metadata"]["model_group"] == "smart-router" - assert subagent_kwargs["metadata"]["routing_decision"]["router_model_name"] == "smart-router" - - @pytest.mark.asyncio - async def test_main_thread_side_calls_to_a_plain_model_keep_the_session_router(self): - router = self._router() - - await router.async_pre_routing_hook(model="smart-router", request_kwargs=self._request_kwargs()) - await router.async_pre_routing_hook(model="expensive-model", request_kwargs=self._request_kwargs()) - - response = await router.async_pre_routing_hook( - model="expensive-model", - request_kwargs=self._request_kwargs(agent_id="agent-1234"), - ) - - assert response is not None - assert response.model == "cheap-model" - - @pytest.mark.asyncio - async def test_redis_cleanup_failure_does_not_reject_a_subagent_request(self): - from litellm.caching.caching import RedisCache - - router = self._router() - del router.complexity_routers["smart-router"] - redis_cache = MagicMock(spec=RedisCache) - redis_cache.async_get_cache = AsyncMock(return_value="smart-router") - redis_cache.async_delete_cache = AsyncMock(side_effect=ConnectionError("redis unavailable")) - router._update_redis_cache(cache=redis_cache) - - response = await router.async_pre_routing_hook( - model="expensive-model", - request_kwargs=self._request_kwargs(agent_id="agent-1234"), - ) - - assert response is None - redis_cache.async_delete_cache.assert_awaited_once() - - @pytest.mark.asyncio - async def test_redis_read_failure_does_not_reject_a_subagent_request(self): - from litellm.caching.caching import RedisCache - - router = self._router() - request_kwargs = self._request_kwargs(agent_id="agent-1234") - cache_key = router._claude_code_session_router_cache_key(request_kwargs) - assert cache_key is not None - await router._claude_code_session_router_cache.in_memory_cache.async_set_cache( - cache_key, - "smart-router", - ) - redis_cache = MagicMock(spec=RedisCache) - redis_cache.async_get_cache = AsyncMock(side_effect=Exception("Redis circuit breaker is open")) - router._update_redis_cache(cache=redis_cache) - - response = await router.async_pre_routing_hook( - model="expensive-model", - request_kwargs=request_kwargs, - ) - - assert response is None - assert "model_group" not in request_kwargs["metadata"] - redis_cache.async_get_cache.assert_awaited_once() - - @pytest.mark.asyncio - async def test_redis_write_failures_do_not_reject_main_or_subagent_requests(self): - from litellm.caching.caching import RedisCache - - router = self._router() - redis_cache = MagicMock(spec=RedisCache) - redis_cache.async_get_cache = AsyncMock(return_value="smart-router") - redis_cache.async_set_cache = AsyncMock(side_effect=Exception("redis unavailable")) - router._update_redis_cache(cache=redis_cache) - - main_response = await router.async_pre_routing_hook( - model="smart-router", - request_kwargs=self._request_kwargs(), - ) - subagent_response = await router.async_pre_routing_hook( - model="expensive-model", - request_kwargs=self._request_kwargs(agent_id="agent-1234"), - ) - - assert main_response is not None - assert main_response.model == "cheap-model" - assert subagent_response is not None - assert subagent_response.model == "cheap-model" - assert redis_cache.async_set_cache.await_count == 2 - - @pytest.mark.asyncio - async def test_subagents_follow_the_main_threads_latest_router_across_workers(self): - from types import SimpleNamespace - - from litellm.caching.caching import RedisCache - - shared_binding = SimpleNamespace(value=None) - shared_redis = MagicMock(spec=RedisCache) - shared_redis.async_get_cache = AsyncMock(side_effect=lambda key, **_: shared_binding.value) - shared_redis.async_set_cache = AsyncMock( - side_effect=lambda key, value, **_: setattr(shared_binding, "value", value) - ) - main_worker, subagent_worker = self._router(), self._router() - main_worker._update_redis_cache(cache=shared_redis) - subagent_worker._update_redis_cache(cache=shared_redis) - - await main_worker.async_pre_routing_hook(model="smart-router", request_kwargs=self._request_kwargs()) - first = await subagent_worker.async_pre_routing_hook( - model="expensive-model", - request_kwargs=self._request_kwargs(agent_id="agent-1234"), - ) - await main_worker.async_pre_routing_hook(model="premium-router", request_kwargs=self._request_kwargs()) - second = await subagent_worker.async_pre_routing_hook( - model="expensive-model", - request_kwargs=self._request_kwargs(agent_id="agent-1234"), - ) - - assert first is not None - assert first.model == "cheap-model" - assert second is not None - assert second.model == "expensive-model" - assert shared_binding.value == "premium-router" - - @pytest.mark.asyncio - async def test_no_pre_routing_strategies_means_no_session_cache_traffic(self): - from litellm.caching.caching import RedisCache - - router = self._router() - router.complexity_routers.clear() - redis_cache = MagicMock(spec=RedisCache) - redis_cache.async_get_cache = AsyncMock(return_value=None) - redis_cache.async_set_cache = AsyncMock() - redis_cache.async_delete_cache = AsyncMock() - router._update_redis_cache(cache=redis_cache) - - for request_kwargs in (self._request_kwargs(), self._request_kwargs(agent_id="agent-1234")): - response = await router.async_pre_routing_hook(model="expensive-model", request_kwargs=request_kwargs) - assert response is None - - redis_cache.async_get_cache.assert_not_awaited() - redis_cache.async_set_cache.assert_not_awaited() - redis_cache.async_delete_cache.assert_not_awaited() - - @pytest.mark.asyncio - async def test_session_bindings_do_not_evict_router_rate_limit_state(self): - router = self._router() - assert router._update_usage(deployment_id="deployment-id", parent_otel_span=None) == 1 - - for session_index in range(201): - request_kwargs = self._request_kwargs() - request_kwargs["proxy_server_request"]["headers"]["X-Claude-Code-Session-Id"] = ( - f"session-{session_index:04d}" - ) - await router.async_pre_routing_hook(model="smart-router", request_kwargs=request_kwargs) - - assert router._update_usage(deployment_id="deployment-id", parent_otel_span=None) == 2 - - @pytest.mark.asyncio - async def test_background_and_fallback_requests_do_not_clear_the_session_router(self): - router = self._router() - - await router.async_pre_routing_hook(model="smart-router", request_kwargs=self._request_kwargs()) - await router.async_pre_routing_hook( - model="expensive-model", - request_kwargs=self._request_kwargs(app="cli-bg"), - ) - await router.async_pre_routing_hook( - model="expensive-model", - request_kwargs=self._request_kwargs(fallback_depth=1), - ) - - response = await router.async_pre_routing_hook( - model="expensive-model", - request_kwargs=self._request_kwargs(agent_id="agent-1234"), - ) - - assert response is not None - assert response.model == "cheap-model" - - @pytest.mark.asyncio - async def test_subagent_fallback_does_not_reapply_the_session_router(self): - router = self._router() - - await router.async_pre_routing_hook(model="smart-router", request_kwargs=self._request_kwargs()) - - response = await router.async_pre_routing_hook( - model="expensive-model", - request_kwargs=self._request_kwargs(agent_id="agent-1234", fallback_depth=1), - ) - - assert response is None - - @pytest.mark.asyncio - async def test_subagent_can_fallback_to_its_original_requested_model(self): - router = self._router( - cheap_response="litellm.RateLimitError", - fallbacks=[{"cheap-model": ["expensive-model"]}], - ) - - await router.acompletion( - model="smart-router", - messages=[{"role": "user", "content": "main turn"}], - **self._request_kwargs(), - ) - subagent_kwargs = self._request_kwargs(agent_id="agent-1234") - - response = await router.acompletion( - model="expensive-model", - messages=[{"role": "user", "content": "subagent turn"}], - **subagent_kwargs, - ) - - assert response.choices[0].message.content == "expensive response" - assert subagent_kwargs["metadata"]["routing_decision"]["routed_model"] == "cheap-model" - - @pytest.mark.asyncio - async def test_subagent_can_use_the_bound_router_name_fallback(self): - router = self._router( - cheap_response="litellm.RateLimitError", - fallbacks=[{"smart-router": ["expensive-model"]}], - ) - - await router.acompletion( - model="smart-router", - messages=[{"role": "user", "content": "main turn"}], - **self._request_kwargs(), - ) - - response = await router.acompletion( - model="expensive-model", - messages=[{"role": "user", "content": "subagent turn"}], - **self._request_kwargs(agent_id="agent-1234"), - ) - - assert response.choices[0].message.content == "expensive response" - - @pytest.mark.asyncio - async def test_anthropic_subagent_four_fallback_hops_use_each_current_model_chain(self): - from litellm.types.router import TaggedPreRoutingStrategy - - failing_groups = ("cheap-model", "fallback-1", "fallback-2", "fallback-3") - router = litellm.Router( - model_list=[ - *( - { - "model_name": group, - "litellm_params": { - "model": "anthropic/claude-3-haiku-20240307", - "mock_response": "litellm.RateLimitError", - }, - } - for group in failing_groups - ), - { - "model_name": "requested-model", - "litellm_params": { - "model": "anthropic/claude-3-haiku-20240307", - "mock_response": "requested response", - }, - }, - { - "model_name": "fallback-4", - "litellm_params": { - "model": "anthropic/claude-3-haiku-20240307", - "mock_response": "fourth fallback response", - }, - }, - ], - fallbacks=[ - {"smart-router": ["fallback-1"]}, - {"fallback-1": ["fallback-2"]}, - {"fallback-2": ["fallback-3"]}, - {"fallback-3": ["fallback-4"]}, - ], - num_retries=0, - max_fallbacks=4, - ) - router.complexity_routers = { - "smart-router": (TaggedPreRoutingStrategy(tags=(), strategy=self._RewriteStrategy()),) - } - main_kwargs = self._request_kwargs() - main_kwargs["litellm_metadata"] = main_kwargs.pop("metadata") - await router.async_pre_routing_hook(model="smart-router", request_kwargs=main_kwargs) - subagent_kwargs = self._request_kwargs(agent_id="agent-1234") - subagent_kwargs["litellm_metadata"] = subagent_kwargs.pop("metadata") - - response = await router.aanthropic_messages( - model="requested-model", - messages=[{"role": "user", "content": "subagent turn"}], - max_tokens=64, - **subagent_kwargs, - ) - - assert response["content"][0]["text"] == "fourth fallback response" - - @pytest.mark.asyncio - async def test_session_router_binding_is_scoped_to_the_authenticated_key(self): - router = self._router() - - await router.async_pre_routing_hook(model="smart-router", request_kwargs=self._request_kwargs()) - - response = await router.async_pre_routing_hook( - model="expensive-model", - request_kwargs=self._request_kwargs(key_hash="key-hash-b", agent_id="agent-1234"), - ) - - assert response is None - - -class TestAutoRouterMaxInputCharsWiring: - """`auto_router_max_input_chars` on the deployment has to reach the AutoRouter that embeds prompts. - - Without it the cap silently reverts to the default, so an operator whose embedding model has a - 512-token window cannot lower it and every long prompt falls back to the default model instead - of being routed. - """ - - @staticmethod - def _router(**extra_params) -> "litellm.Router": - pytest.importorskip("semantic_router", reason="auto-router needs the semantic-router extra") - return litellm.Router( - model_list=[ - {"model_name": "gpt-4o", "litellm_params": {"model": "gpt-4o"}}, - { - "model_name": "my-auto-router", - "litellm_params": { - "model": "auto_router/my-auto-router", - "auto_router_config": json.dumps( - {"routes": [{"name": "gpt-4o", "utterances": ["write me code"]}]} - ), - "auto_router_default_model": "gpt-4o", - "auto_router_embedding_model": "text-embedding-3-small", - **extra_params, - }, - }, - ] - ) - - @staticmethod - def _registered_auto_router(router: "litellm.Router"): - return router.auto_routers["my-auto-router"][0].strategy - - def test_should_pass_the_configured_cap_to_the_auto_router(self): - router = self._router(auto_router_max_input_chars=512) - - assert self._registered_auto_router(router).max_input_chars == 512 - - def test_should_fall_back_to_the_shared_default_when_the_deployment_omits_it(self): - from litellm.constants import DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS - - router = self._router() - - assert self._registered_auto_router(router).max_input_chars == DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS - - -class TestTaggedAutoRouterOnSharedModelName: - """A tagged auto-router marker sharing its model_name with a plain deployment must not - capture requests whose tags don't match it when tag filtering is enabled (#36620).""" - - class _FixedRouteLayer: - def __call__(self, text: str): - from semantic_router.schema import RouteChoice - - return RouteChoice(name="gemini-flash") - - @classmethod - def _router(cls, marker_tags, include_plain_sibling: bool, enable_tag_filtering: bool) -> "litellm.Router": - pytest.importorskip("semantic_router", reason="auto-router needs the semantic-router extra") - marker = { - "model_name": "gpt4o", - "litellm_params": { - "model": "auto_router/gpt4o-router", - "auto_router_config": json.dumps( - {"routes": [{"name": "gemini-flash", "utterances": ["capital city questions"]}]} - ), - "auto_router_default_model": "gemini-flash", - "auto_router_embedding_model": "text-embedding-3-small", - **({"tags": marker_tags} if marker_tags else {}), - }, - } - plain = {"model_name": "gpt4o", "litellm_params": {"model": "openai/gpt-4o"}} - tier = {"model_name": "gemini-flash", "litellm_params": {"model": "gemini/gemini-3.6-flash"}} - router = litellm.Router( - model_list=[plain, marker, tier] if include_plain_sibling else [marker, tier], - enable_tag_filtering=enable_tag_filtering, - ) - router.auto_routers["gpt4o"][0].strategy.routelayer = cls._FixedRouteLayer() - return router - - @staticmethod - async def _hook_response(router: "litellm.Router", request_kwargs: dict): - return await router.async_pre_routing_hook( - model="gpt4o", - request_kwargs=request_kwargs, - messages=[{"role": "user", "content": "What is the capital of France?"}], - ) - - @pytest.mark.asyncio - async def test_untagged_request_bypasses_the_tagged_marker_when_a_plain_deployment_shares_the_name(self): - router = self._router(marker_tags=["route"], include_plain_sibling=True, enable_tag_filtering=True) - - assert await self._hook_response(router, {}) is None - - @pytest.mark.asyncio - async def test_request_tagged_for_the_marker_is_still_semantically_routed(self): - router = self._router(marker_tags=["route"], include_plain_sibling=True, enable_tag_filtering=True) - - response = await self._hook_response(router, {"metadata": {"tags": ["route"]}}) - - assert response is not None - assert response.model == "gemini-flash" - - @pytest.mark.asyncio - async def test_request_level_tag_filtering_from_key_settings_bypasses_the_marker(self): - router = self._router(marker_tags=["route"], include_plain_sibling=True, enable_tag_filtering=False) - - assert await self._hook_response(router, {"enable_tag_filtering": True}) is None - - @pytest.mark.asyncio - async def test_globally_disabled_filtering_still_lets_the_marker_capture_untagged_requests(self): - router = self._router(marker_tags=["route"], include_plain_sibling=True, enable_tag_filtering=False) - - response = await self._hook_response(router, {}) - - assert response is not None - assert response.model == "gemini-flash" - - @pytest.mark.asyncio - async def test_marker_only_alias_still_captures_untagged_requests(self): - router = self._router(marker_tags=["route"], include_plain_sibling=False, enable_tag_filtering=True) - - response = await self._hook_response(router, {}) - - assert response is not None - assert response.model == "gemini-flash" - - @pytest.mark.asyncio - async def test_untagged_marker_sharing_the_name_still_captures_untagged_requests(self): - router = self._router(marker_tags=None, include_plain_sibling=True, enable_tag_filtering=True) - - response = await self._hook_response(router, {}) - - assert response is not None - assert response.model == "gemini-flash" - - @pytest.mark.asyncio - async def test_untagged_selection_never_lands_on_the_marker_deployment(self): - router = self._router(marker_tags=["route"], include_plain_sibling=True, enable_tag_filtering=True) - - for _ in range(20): - deployment = await router.async_get_available_deployment( - model="gpt4o", - request_kwargs={}, - messages=[{"role": "user", "content": "What is the capital of France?"}], - ) - assert deployment["litellm_params"]["model"] == "openai/gpt-4o" - - def test_deployment_without_litellm_params_mapping_is_not_a_marker(self): - assert litellm.Router._is_strategy_marker_deployment({"model_name": "gpt4o"}) is False - - def test_model_name_has_plain_deployments_reflects_the_pool(self): - mixed = self._router(marker_tags=["route"], include_plain_sibling=True, enable_tag_filtering=True) - marker_only = self._router(marker_tags=["route"], include_plain_sibling=False, enable_tag_filtering=True) - - assert mixed._model_name_has_plain_deployments("gpt4o") is True - assert marker_only._model_name_has_plain_deployments("gpt4o") is False - - -class TestAutoRouterSharedModelNameConnectionParams: - """A plain deployment sharing its model_name with an `auto_router/` marker must not have - its api_base and api_key grafted onto the routed tier's outbound call (#36619).""" - - PLAIN_API_BASE = "https://plain-sibling.openai.example/v1" - PLAIN_API_KEY = "sk-plain-sibling-secret" - - class _FixedRouteLayer: - def __call__(self, text: str): - from semantic_router.schema import RouteChoice - - return RouteChoice(name="gemini-flash") - - @classmethod - def _router(cls, plain_entry_first: bool) -> "litellm.Router": - pytest.importorskip("semantic_router", reason="auto-router needs the semantic-router extra") - plain = { - "model_name": "gpt4o", - "litellm_params": { - "model": "openai/gpt-4o", - "api_key": cls.PLAIN_API_KEY, - "api_base": cls.PLAIN_API_BASE, - }, - } - marker = { - "model_name": "gpt4o", - "litellm_params": { - "model": "auto_router/gpt4o-router", - "auto_router_config": json.dumps( - {"routes": [{"name": "gemini-flash", "utterances": ["capital city questions"]}]} - ), - "auto_router_default_model": "gemini-flash", - "auto_router_embedding_model": "text-embedding-3-small", - "drop_params": True, - }, - } - tier = { - "model_name": "gemini-flash", - "litellm_params": {"model": "gemini/gemini-3.6-flash", "api_key": "sk-tier-key"}, - } - shared_name_entries = [plain, marker] if plain_entry_first else [marker, plain] - router = litellm.Router(model_list=[*shared_name_entries, tier]) - router.auto_routers["gpt4o"][0].strategy.routelayer = cls._FixedRouteLayer() - return router - - @staticmethod - def _gemini_response() -> httpx.Response: - return httpx.Response( - status_code=200, - json={ - "candidates": [ - {"content": {"parts": [{"text": "Paris"}], "role": "model"}, "finishReason": "STOP"} - ], - "usageMetadata": {"promptTokenCount": 5, "candidatesTokenCount": 1, "totalTokenCount": 6}, - "modelVersion": "gemini-3.6-flash", - }, - request=httpx.Request("POST", "https://generativelanguage.googleapis.com"), - ) - - @pytest.mark.parametrize( - "plain_entry_first", [True, False], ids=["plain_entry_first", "marker_entry_first"] - ) - async def test_routed_tier_call_goes_out_on_its_own_endpoint_and_credentials(self, plain_entry_first): - """The outbound provider request for the routed tier hits the tier's own Gemini host - with the tier's own key, never the plain sibling's api_base or api_key.""" - from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler - - router = self._router(plain_entry_first) - - with patch.object( - AsyncHTTPHandler, "post", new_callable=AsyncMock, return_value=self._gemini_response() - ) as mock_post: - await router.acompletion( - model="gpt4o", - messages=[{"role": "user", "content": "What is the capital of France?"}], - ) - - call = mock_post.call_args - outbound_url = str(call.kwargs["url"] if "url" in call.kwargs else call.args[0]) - outbound_headers = dict(call.kwargs.get("headers") or {}) - - assert "generativelanguage.googleapis.com" in outbound_url - assert "gemini-3.6-flash" in outbound_url - assert self.PLAIN_API_BASE not in outbound_url - assert self.PLAIN_API_KEY not in outbound_url - assert self.PLAIN_API_KEY not in json.dumps(outbound_headers) - - -class TestGetAllowedFailsFromPolicy: - def _make_router(self, **policy_kwargs) -> litellm.Router: - from litellm.types.router import AllowedFailsPolicy - - return litellm.Router( - model_list=[{"model_name": "gpt-4", "litellm_params": {"model": "gpt-4", "api_key": "fake"}}], - allowed_fails_policy=AllowedFailsPolicy(**policy_kwargs), - ) - - def test_no_policy_returns_none(self): - router = litellm.Router( - model_list=[{"model_name": "gpt-4", "litellm_params": {"model": "gpt-4", "api_key": "fake"}}], - ) - assert router.get_allowed_fails_from_policy(litellm.RateLimitError("429", "openai", "gpt-4")) is None - - def test_internal_server_error_allowed_fails(self): - router = self._make_router(InternalServerErrorAllowedFails=7) - exc = litellm.InternalServerError("500", "openai", "gpt-4") - assert router.get_allowed_fails_from_policy(exc) == 7 - - def test_service_unavailable_error_allowed_fails(self): - router = self._make_router(ServiceUnavailableErrorAllowedFails=4) - exc = litellm.ServiceUnavailableError("503", "openai", "gpt-4") - assert router.get_allowed_fails_from_policy(exc) == 4 - - def test_bad_gateway_error_allowed_fails(self): - router = self._make_router(BadGatewayErrorAllowedFails=2) - exc = litellm.BadGatewayError("502", "openai", "gpt-4") - assert router.get_allowed_fails_from_policy(exc) == 2 - - def test_not_found_error_allowed_fails(self): - router = self._make_router(NotFoundErrorAllowedFails=1) - exc = litellm.NotFoundError("404", "openai", "gpt-4") - assert router.get_allowed_fails_from_policy(exc) == 1 - - def test_unmatched_exception_returns_none(self): - router = self._make_router(InternalServerErrorAllowedFails=5) - exc = litellm.RateLimitError("429", "openai", "gpt-4") - assert router.get_allowed_fails_from_policy(exc) is None - - -class _LogCapture(logging.Handler): - def __init__(self, level): - super().__init__(level=level) - self._level = level - self.messages = [] - - def emit(self, record): - if record.levelno == self._level: - self.messages.append(record.getMessage()) - - -class _FallbackAttemptRecorder(CustomLogger): - def __init__(self): - super().__init__() - self.failed_targets = [] - self.breadcrumbs_per_target = [] - - async def log_failure_fallback_event(self, original_model_group, kwargs, original_exception): - self.failed_targets.append(kwargs.get("model")) - self.breadcrumbs_per_target.append(kwargs.get("metadata", {}).get("previous_models", ())) - - -def _cyclic_fallback_router(num_retries=0): - groups = ["group-a", "group-b", "group-c", "group-d"] - return litellm.Router( - model_list=[ - { - "model_name": group, - "litellm_params": { - "model": "openai/gpt-4o-mini", - "api_key": "sk-fake", - "mock_response": "litellm.InternalServerError", - }, - } - for group in groups - ], - fallbacks=[ - {"group-a": ["group-b", "group-c"]}, - {"group-b": ["group-a", "group-c"]}, - {"group-c": ["group-d"]}, - {"group-d": ["group-b", "group-a"]}, - ], - num_retries=num_retries, - ) - - -async def _drive_cyclic_fallback(router, capture, recorder=None, **request_kwargs): - router_logger = logging.getLogger("LiteLLM Router") - previous_level = router_logger.level - router_logger.setLevel(capture.level) - router_logger.addHandler(capture) - if recorder is not None: - litellm.callbacks.append(recorder) - try: - with pytest.raises(litellm.InternalServerError): - await router.acompletion( - model="group-a", messages=[{"role": "user", "content": "hi"}], **request_kwargs - ) - finally: - router_logger.removeHandler(capture) - router_logger.setLevel(previous_level) - if recorder is not None: - litellm.callbacks.remove(recorder) - - -@pytest.mark.asyncio -async def test_cyclic_fallback_graph_does_not_amplify_one_request(): - """A fallback graph whose entries loop back on each other is easy to build by accident, - and every group in the loop fails identically on a deterministic error, so the walk must - not revisit a group and must not re-emit a growing chained traceback at each level. Left - unbounded, one request blocks the event loop long enough for health probes to fail.""" - recorder = _FallbackAttemptRecorder() - capture = _LogCapture(logging.ERROR) - - await _drive_cyclic_fallback(_cyclic_fallback_router(), capture, recorder) - - assert sorted(set(recorder.failed_targets)) == ["group-b", "group-c", "group-d"] - assert len(recorder.failed_targets) == len(set(recorder.failed_targets)) - assert not any("Traceback (most recent call last)" in message for message in capture.messages) - assert sum(len(message) for message in capture.messages) < 5_000 - - -@pytest.mark.asyncio -async def test_retry_breadcrumbs_do_not_carry_the_walk_state(): - """log_retry copies every kwarg into previous_models, which reaches spend logs and - logging callbacks. The set of already-attempted groups is router-internal walk state - with no diagnostic value there, and it is the one entry that is not a plain scalar. - A retry has to be configured for the walk state to reach log_retry at all.""" - router = _cyclic_fallback_router(num_retries=1) - capture = _LogCapture(logging.ERROR) - recorder = _FallbackAttemptRecorder() - - await _drive_cyclic_fallback(router, capture, recorder) - - breadcrumbs = [breadcrumb for hop in recorder.breadcrumbs_per_target for breadcrumb in hop] - assert breadcrumbs, "no retry breadcrumbs were recorded" - assert any( - "fallback_depth" in breadcrumb for breadcrumb in breadcrumbs - ), "no breadcrumb carried router walk state, so this test cannot see the leak" - for breadcrumb in breadcrumbs: - assert "attempted_targets" not in breadcrumb - - -_BREADCRUMB_CREDENTIAL_CANARY = "Bearer sk-ant-oat01-RETRY-BREADCRUMB-CANARY-doNotShip" - - -@pytest.mark.parametrize( - "container_key, request_kwargs", - [ - ( - "provider_specific_header", - { - "provider_specific_header": { - "custom_llm_provider": "openai", - "extra_headers": {"authorization": _BREADCRUMB_CREDENTIAL_CANARY}, - } - }, - ), - ( - "extra_headers", - {"extra_headers": {"authorization": _BREADCRUMB_CREDENTIAL_CANARY}}, - ), - ( - "api_key", - {"api_key": _BREADCRUMB_CREDENTIAL_CANARY}, - ), - ], -) -@pytest.mark.asyncio -async def test_retry_breadcrumbs_never_carry_a_forwarded_credential(container_key, request_kwargs): - """log_retry copies kwargs into previous_models, which reaches spend logs and logging callbacks. - Any of these kwargs can carry a client's forwarded Authorization token or a provider key, and a - breadcrumb has no diagnostic use for the raw secret. A denylist of key names is always one new - credential kwarg behind, so log_retry scrubs credential-named values by pattern instead: the - container still reaches the breadcrumb, but the raw secret never does, whatever key holds it.""" - router = _cyclic_fallback_router(num_retries=1) - capture = _LogCapture(logging.ERROR) - metadata = {} - - await _drive_cyclic_fallback(router, capture, metadata=metadata, **request_kwargs) - - breadcrumbs = metadata["previous_models"] - assert breadcrumbs, "no retry breadcrumbs were recorded" - dumped = json.dumps(breadcrumbs, default=str) - assert container_key in dumped, "the credential-bearing kwarg never reached the breadcrumb, so this test cannot see the leak" - assert _BREADCRUMB_CREDENTIAL_CANARY not in dumped - - -def _always_failing_router(num_retries): - return litellm.Router( - model_list=[ - { - "model_name": "broken-group", - "litellm_params": { - "model": "openai/gpt-4o-mini", - "api_key": "sk-fake", - "mock_response": "litellm.InternalServerError", - }, - } - ], - num_retries=num_retries, - ) - - -async def _fail_one_proxy_shaped_request(router, request_marker): - """The proxy hands the router a metadata dict and a proxy_server_request whose body is a - shallow copy of the request, so body["metadata"] is the very same dict the router later - stamps previous_models onto.""" - metadata = {"request_marker": request_marker} - with pytest.raises(litellm.InternalServerError): - await router.acompletion( - model="broken-group", - messages=[{"role": "user", "content": "hi"}], - metadata=metadata, - proxy_server_request={ - "url": "http://localhost:4000/v1/chat/completions", - "method": "POST", - "headers": {}, - "body": {"model": "broken-group", "metadata": metadata}, - }, - ) - return metadata["previous_models"] - - -def _nested_breadcrumb_lists(node): - if isinstance(node, dict): - return [v for k, v in node.items() if k == "previous_models"] + [ - found for v in node.values() for found in _nested_breadcrumb_lists(v) - ] - if isinstance(node, (list, tuple)): - return [found for item in node for found in _nested_breadcrumb_lists(item)] - return [] - - -@pytest.mark.asyncio -async def test_retry_breadcrumbs_stay_per_request_and_flat_across_failing_requests(): - """Every failed attempt appends a breadcrumb to metadata["previous_models"], and the proxy's - request snapshot aliases that same metadata dict. Kept on the Router and copied wholesale, - each breadcrumb embedded every earlier one from every earlier request, so the breadcrumb - tree, and with it the debug repr of the kwargs, roughly doubled on each failed attempt until - a single-worker proxy spent minutes in the redaction regex and stopped answering.""" - router = _always_failing_router(num_retries=2) - - breadcrumbs_per_request = [ - await _fail_one_proxy_shaped_request(router, f"request-{request_number}") for request_number in range(1, 7) - ] - - for request_number, breadcrumbs in enumerate(breadcrumbs_per_request, start=1): - assert len(breadcrumbs) == 3, "one initial attempt plus two retries failed, each leaving one breadcrumb" - assert {breadcrumb["metadata"]["request_marker"] for breadcrumb in breadcrumbs} == {f"request-{request_number}"} - for breadcrumb in breadcrumbs: - assert _nested_breadcrumb_lists(breadcrumb) == [] - assert len({len(repr(breadcrumbs)) for breadcrumbs in breadcrumbs_per_request}) == 1 - - -@pytest.mark.asyncio -async def test_retry_breadcrumbs_keep_only_the_last_four_attempts(): - router = _always_failing_router(num_retries=6) - - breadcrumbs = await _fail_one_proxy_shaped_request(router, "request-1") - - assert len(breadcrumbs) == 4 - assert [breadcrumb["metadata"]["attempted_retries"] for breadcrumb in breadcrumbs] == [3, 4, 5, 6] - - -@pytest.mark.asyncio -async def test_fallback_traceback_stays_available_at_debug_level(): - """Dropping the stack from the ERROR line is only safe because the fallback path still - emits it once per level at DEBUG, which is what an operator needs to diagnose why every - fallback failed. This pins that remaining debug traceback.""" - capture = _LogCapture(logging.DEBUG) - - await _drive_cyclic_fallback(_cyclic_fallback_router(), capture) - - assert any("Traceback (most recent call last)" in message for message in capture.messages) - - -@pytest.mark.asyncio -async def test_fallback_failure_detail_from_upstream_is_bounded(): - """The detail each level records about the level below it is attacker-influenced, since - it carries whatever the upstream error said. It has to be bounded on its own, so a walk - over several groups cannot compound one large message into the log or into the message - handed back to the caller.""" - huge_message = "z" * 50_000 - capture = _LogCapture(logging.ERROR) - - await _drive_cyclic_fallback( - _cyclic_fallback_router(), - capture, - mock_response=litellm.InternalServerError( - message=huge_message, llm_provider="openai", model="group-a" - ), - ) - - assert capture.messages, "the fallback failure path did not log at ERROR" - assert huge_message not in "".join(capture.messages) - assert max(len(message) for message in capture.messages) < 5_000 - - -def test_stamp_or_clear_metadata_key_writes_and_clears_both_buckets(): - request_kwargs = {"metadata": {}} - litellm.Router._stamp_or_clear_metadata_key(request_kwargs=request_kwargs, key="probe", value=7) - assert request_kwargs["metadata"]["probe"] == 7 - - stale_kwargs = {"metadata": {"probe": 7}, "litellm_metadata": {"probe": 7}} - litellm.Router._stamp_or_clear_metadata_key(request_kwargs=stale_kwargs, key="probe", value=None) - assert "probe" not in stale_kwargs["metadata"] - assert "probe" not in stale_kwargs["litellm_metadata"] - - -@pytest.mark.parametrize( - "complexity_router_config,expect_callback", - [ - ({"tiers": {"SIMPLE": "gpt-4o"}}, True), - ({"tiers": {"SIMPLE": "gpt-4o"}, "deployment_affinity": False}, False), - ({"tiers": {"SIMPLE": "gpt-4o"}, "deployment_affinity": False, "session_affinity": True}, True), - ], -) -def test_complexity_router_registers_affinity_callback_for_deployment_pin(complexity_router_config, expect_callback): - """The marker the complexity router stamps is inert unless a DeploymentAffinityCheck is - registered to read it, so deployment_affinity has to pull the callback in, and its default-on - means a bare config registers one. Opting out must skip the callback entirely rather than - register a filter that can never fire, including when session_affinity is on, since the two - pins are independent.""" - from litellm.router_utils.pre_call_checks.deployment_affinity_check import ( - DeploymentAffinityCheck, - ) - - router = litellm.Router( - model_list=[ - {"model_name": "gpt-4o", "litellm_params": {"model": "gpt-4o"}}, - { - "model_name": "my-complexity-router", - "litellm_params": { - "model": "auto_router/complexity_router", - "complexity_router_config": complexity_router_config, - }, - }, - ] - ) - try: - registered = any(isinstance(cb, DeploymentAffinityCheck) for cb in router.optional_callbacks or []) - assert registered is expect_callback - finally: - for cb in router.optional_callbacks or []: - litellm.logging_callback_manager.remove_callback_from_all_lists(cb) - - -def test_ensure_deployment_affinity_callback_is_idempotent(): - from litellm.router_utils.pre_call_checks.deployment_affinity_check import ( - DeploymentAffinityCheck, - ) - - router = litellm.Router(model_list=[]) - try: - router._ensure_deployment_affinity_callback() - router._ensure_deployment_affinity_callback() - affinity_callbacks = [ - cb for cb in router.optional_callbacks or [] if isinstance(cb, DeploymentAffinityCheck) - ] - assert len(affinity_callbacks) == 1 - finally: - for cb in router.optional_callbacks or []: - litellm.logging_callback_manager.remove_callback_from_all_lists(cb) - - -def test_get_router_model_info_does_not_wipe_cached_pricing(): - """A Deployment's model_info declares the mirrored pricing fields with None defaults; - merging it must not write those Nones into the lru_cache'd dict get_model_info() owns, - or /model/info loses built-in prices for every model a worker serves.""" - from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo - - litellm.get_model_info.cache_clear() - expected = copy.deepcopy(litellm.get_model_info(model="anthropic/claude-sonnet-4-5")) - - router = litellm.Router(model_list=[]) - merged = router.get_router_model_info( - deployment=Deployment( - model_name="sonnet", - litellm_params=LiteLLM_Params(model="claude-sonnet-4-5", custom_llm_provider="anthropic"), - model_info=ModelInfo(id="sonnet-1"), - ), - received_model_name="sonnet", - ) - - assert litellm.get_model_info(model="anthropic/claude-sonnet-4-5") == expected - for field in ("input_cost_per_token", "output_cost_per_token", "cache_read_input_token_cost"): - assert merged[field] == expected[field] - - -def test_get_router_model_info_keeps_explicit_pricing_overrides(): - from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo - - litellm.get_model_info.cache_clear() - router = litellm.Router(model_list=[]) - merged = router.get_router_model_info( - deployment=Deployment( - model_name="sonnet", - litellm_params=LiteLLM_Params(model="claude-sonnet-4-5", custom_llm_provider="anthropic"), - model_info=ModelInfo(id="sonnet-1", input_cost_per_token=1e-08), - ), - received_model_name="sonnet", - ) - - assert merged["input_cost_per_token"] == 1e-08 - assert litellm.get_model_info(model="anthropic/claude-sonnet-4-5")["input_cost_per_token"] != 1e-08 - - -class TestModelGroupAliasReachesPreRoutingStrategies: - """A `model_group_alias` whose target is a strategy router must dispatch exactly like the - router's own model_name. The four strategy registries are keyed by the marker deployment's - model_name, so the alias has to be resolved before the pre-routing hook looks anything up, - and a group that resolves only to markers is not callable at all (LIT-4664).""" - - MARKER_TIMEOUT = 42.0 - REGISTRY_NAMES = ("auto_routers", "complexity_routers", "adaptive_routers", "quality_routers") - - class _RewriteStrategy: - async def async_pre_routing_hook( - self, model, request_kwargs, messages=None, input=None, specific_deployment=False - ): - from litellm.types.router import PreRoutingHookResponse - - return PreRoutingHookResponse(model="gemini-flash", messages=messages) - - @classmethod - def _router(cls, registry_name: str | None) -> "litellm.Router": - from litellm.types.router import TaggedPreRoutingStrategy - - tiers = dict.fromkeys(("SIMPLE", "MEDIUM", "COMPLEX", "REASONING"), "gemini-flash") - router = litellm.Router( - model_list=[ - { - "model_name": "smart-route", - "litellm_params": { - "model": "auto_router/complexity_router", - "complexity_router_config": {"tiers": tiers}, - "complexity_router_default_model": "gemini-flash", - "timeout": cls.MARKER_TIMEOUT, - }, - }, - { - "model_name": "gemini-flash", - "litellm_params": {"model": "gemini/gemini-3.6-flash", "mock_response": "routed by the tier"}, - }, - ], - model_group_alias={"smart-alias": "smart-route"}, - ) - for name in cls.REGISTRY_NAMES: - setattr(router, name, {}) - if registry_name is not None: - setattr( - router, - registry_name, - {"smart-route": [TaggedPreRoutingStrategy(tags=(), strategy=cls._RewriteStrategy())]}, - ) - return router - - @staticmethod - def _messages() -> list[dict[str, str]]: - return [{"role": "user", "content": "What is the capital of France?"}] - - @pytest.mark.parametrize("registry_name", REGISTRY_NAMES) - @pytest.mark.asyncio - async def test_alias_dispatches_to_the_strategy_registered_under_the_target(self, registry_name): - router = self._router(registry_name) - request_kwargs = {"metadata": {}} - - response = await router.async_pre_routing_hook( - model="smart-alias", request_kwargs=request_kwargs, messages=self._messages() - ) - - assert response is not None - assert response.model == "gemini-flash" - - @pytest.mark.asyncio - async def test_alias_call_still_forwards_the_marker_own_params_to_the_routed_tier(self): - router = self._router("auto_routers") - request_kwargs = {"metadata": {}} - - await router.async_pre_routing_hook( - model="smart-alias", request_kwargs=request_kwargs, messages=self._messages() - ) - - assert request_kwargs["timeout"] == self.MARKER_TIMEOUT - - @pytest.mark.asyncio - async def test_alias_deployment_selection_lands_on_the_tier_never_the_marker(self): - router = self._router("auto_routers") - - deployment = await router.async_get_available_deployment( - model="smart-alias", request_kwargs={"metadata": {}}, messages=self._messages() - ) - - assert deployment["litellm_params"]["model"] == "gemini/gemini-3.6-flash" - - @pytest.mark.asyncio - async def test_alias_call_completes_and_still_bills_the_name_the_caller_sent(self): - router = self._router("auto_routers") - metadata: dict = {} - - response = await router.acompletion( - model="smart-alias", messages=self._messages(), metadata=metadata - ) - - assert response.choices[0].message.content == "routed by the tier" - assert metadata["model_group"] == "smart-alias" - assert metadata["model_group_alias"] == "smart-alias" - - def test_a_group_of_only_markers_is_not_a_callable_model(self): - router = self._router(None) - - with pytest.raises(litellm.BadRequestError, match="strategy router marker"): - router.get_available_deployment( - model="smart-route", messages=self._messages(), request_kwargs={"metadata": {}} - ) - - -class TestAutoRouterCompressionDecoupling: - """An auto router's `auto_router_routing_compression` / `auto_router_model_compression` - decouple what the routing decision sees from what the model call sees. The one - assertion that must hold under any mutation: the strategy can be routed on - compressed text while the caller's own `messages` list - the one that would reach - the model - is never touched.""" - - class _RecordingStrategy: - """Echoes back whatever `messages` it was handed, like every real strategy does.""" - - def __init__(self): - self.received_messages: list[dict] | None = None - - async def async_pre_routing_hook( - self, model, request_kwargs, messages=None, input=None, specific_deployment=False - ): - from litellm.types.router import PreRoutingHookResponse - - self.received_messages = messages - return PreRoutingHookResponse(model="gemini-flash", messages=messages) - - class _CompressingGuardrail(CustomGuardrail): - def __init__(self, guardrail_name: str): - super().__init__(guardrail_name=guardrail_name) - self.call_count = 0 - - async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): - self.call_count += 1 - structured_messages = inputs.get("structured_messages") or [] - compressed = [{**m, "content": f"[COMPRESSED] {m.get('content')}"} for m in structured_messages] - return {**inputs, "structured_messages": compressed} - - @staticmethod - def _messages() -> list[dict[str, str]]: - return [{"role": "user", "content": "What is the capital of France?"}] - - def _router(self, marker_litellm_params: dict) -> tuple[litellm.Router, "_RecordingStrategy"]: - from litellm.types.router import TaggedPreRoutingStrategy - - tiers = dict.fromkeys(("SIMPLE", "MEDIUM", "COMPLEX", "REASONING"), "gemini-flash") - router = litellm.Router( - model_list=[ - { - "model_name": "smart-router", - "litellm_params": { - "model": "auto_router/complexity_router", - "complexity_router_config": {"tiers": tiers}, - "complexity_router_default_model": "gemini-flash", - **marker_litellm_params, - }, - }, - { - "model_name": "gemini-flash", - "litellm_params": {"model": "gemini/gemini-3.6-flash", "mock_response": "routed by the tier"}, - }, - ], - ) - for name in ("auto_routers", "complexity_routers", "adaptive_routers", "quality_routers"): - setattr(router, name, {}) - strategy = self._RecordingStrategy() - router.complexity_routers = {"smart-router": [TaggedPreRoutingStrategy(tags=(), strategy=strategy)]} - return router, strategy - - @pytest.fixture - def registered_guardrail(self, monkeypatch): - from litellm.proxy.guardrails import guardrail_registry - - # Registered under a compression provider name: both hops refuse a name that - # does not resolve to one, so a bare callback would never be used. - monkeypatch.setitem(guardrail_registry.guardrail_class_registry, "headroom", self._CompressingGuardrail) - guardrail = self._CompressingGuardrail(guardrail_name="fake-compress") - litellm.logging_callback_manager.add_litellm_callback(guardrail) - yield guardrail - litellm.logging_callback_manager.remove_callback_from_all_lists(guardrail) - - @pytest.mark.asyncio - async def test_routing_side_compression_never_reaches_the_caller_messages(self, registered_guardrail): - router, strategy = self._router( - { - "auto_router_routing_compression": "fake-compress", - "auto_router_model_compression": "none", - } - ) - original_messages = self._messages() - - response = await router.async_pre_routing_hook( - model="smart-router", request_kwargs={"metadata": {}}, messages=original_messages - ) - - assert strategy.received_messages == [ - {"role": "user", "content": "[COMPRESSED] What is the capital of France?"} - ] - assert response.messages == original_messages - - @pytest.mark.asyncio - async def test_model_side_compression_alone_leaves_routing_uncompressed(self, registered_guardrail): - router, strategy = self._router( - { - "auto_router_routing_compression": "none", - "auto_router_model_compression": "fake-compress", - } - ) - original_messages = self._messages() - - response = await router.async_pre_routing_hook( - model="smart-router", request_kwargs={"metadata": {}}, messages=original_messages - ) - - assert strategy.received_messages == original_messages - assert response.messages == original_messages - assert registered_guardrail.call_count == 0 - - @pytest.mark.asyncio - async def test_routing_none_classifies_on_the_live_messages_not_a_pre_guardrail_copy(self, registered_guardrail): - """Routing asked for no compression while the model hop compressed, so the only - messages left are that guardrail's output and the strategy classifies on them. - - Keeping a pre-compression copy to classify on instead is what this deliberately - gives up: that copy is taken before the pre-call guardrails run, so it still - holds whatever a masking guardrail exists to strip, and routing-side compression - POSTs its input to an external service.""" - router, strategy = self._router( - { - "auto_router_routing_compression": "none", - "auto_router_model_compression": "fake-compress", - } - ) - model_compressed = [{"role": "user", "content": "[COMPRESSED] What is the capital of France?"}] - - await router.async_pre_routing_hook( - model="smart-router", request_kwargs={"metadata": {}}, messages=model_compressed - ) - - assert strategy.received_messages == model_compressed - assert registered_guardrail.call_count == 0 - - @pytest.mark.asyncio - @pytest.mark.asyncio - async def test_same_compression_still_compresses_routing_when_nothing_armed_it(self, registered_guardrail): - """Regression: only the proxy calls arm_pre_call. Used through the SDK, nothing - arms the model-side guardrail and nothing has compressed anything, so reusing a - model-hop result that was never produced would serve the request with no - compression on either hop, silently ignoring the configuration.""" - from litellm.proxy.guardrails import auto_router_compression - - router, strategy = self._router( - { - "auto_router_routing_compression": "fake-compress", - "auto_router_model_compression": "fake-compress", - } - ) - uncompressed = self._messages() - assert auto_router_compression.model_hop_compression_armed() is False - - await router.async_pre_routing_hook( - model="smart-router", request_kwargs={"metadata": {}}, messages=uncompressed - ) - - assert strategy.received_messages != uncompressed - assert registered_guardrail.call_count == 1 - - async def test_same_compression_on_both_hops_compresses_once(self, registered_guardrail): - """The same/different distinction exists so a shared choice does not pay for - compression twice: by the time the router runs, `messages` already reflects - whatever the ordinary pre-call guardrail pipeline did for the model call, so - the routing decision must reuse it rather than calling the guardrail again.""" - from litellm.proxy.guardrails import auto_router_compression - - router, strategy = self._router( - { - "auto_router_routing_compression": "fake-compress", - "auto_router_model_compression": "fake-compress", - } - ) - # Stands in for what the proxy's ordinary pre-call guardrail pipeline would - # have already produced for the model call, since `auto_router_model_compression` - # names a guardrail: the router never triggers that pipeline itself. - already_compressed_messages = [{"role": "user", "content": "[COMPRESSED] What is the capital of France?"}] - # arm_pre_call is what would have armed that guardrail, and only the proxy calls - # it; the reuse below is conditional on it having run. - armed = auto_router_compression._model_hop_armed.set(True) - - try: - response = await router.async_pre_routing_hook( - model="smart-router", request_kwargs={"metadata": {}}, messages=already_compressed_messages - ) - finally: - auto_router_compression._model_hop_armed.reset(armed) - - assert strategy.received_messages == already_compressed_messages - assert response.messages == already_compressed_messages - assert registered_guardrail.call_count == 0 - - @pytest.mark.asyncio - async def test_no_policy_is_fully_unaffected(self, registered_guardrail): - router, strategy = self._router({}) - original_messages = self._messages() - - response = await router.async_pre_routing_hook( - model="smart-router", request_kwargs={"metadata": {}}, messages=original_messages - ) - - assert strategy.received_messages is original_messages - assert response.messages == original_messages - assert registered_guardrail.call_count == 0 - - -@pytest.mark.usefixtures("local_model_cost_map") - -@pytest.mark.usefixtures("local_model_cost_map") -class TestAzureBaseModelFallbackLogging: - """When an azure deployment has no base_model but its model name is a known - azure key in the cost map, get_router_model_info resolves it via the - fallback, so it must not log the per-request 'Could not identify azure - model' ERROR. The ERROR must remain for genuinely unmappable deployment - names. Issue #33172.""" - - def _router_with_azure_deployment(self, deployment_model: str): - return litellm.Router( - model_list=[ - { - "model_name": "my-group", - "litellm_params": { - "model": deployment_model, - "api_key": "fake-key", - "api_base": "https://fake.openai.azure.com", - }, - "model_info": {"id": "azure-base-model-test-id"}, - } - ] - ) - - def test_map_known_deployment_name_resolves_without_error_log(self): - router = self._router_with_azure_deployment("azure/gpt-4o") - - with patch( - "litellm.router.verbose_router_logger.error" - ) as mock_error: - model_info = router.get_router_model_info( - deployment=None, received_model_name="my-group", id="azure-base-model-test-id" - ) - - assert not any( - "Could not identify azure model" in str(call) - for call in mock_error.call_args_list - ), f"unexpected error log: {mock_error.call_args_list}" - # the fallback resolution must actually surface the map values - assert model_info["max_input_tokens"] == litellm.model_cost["azure/gpt-4o"]["max_input_tokens"] - assert model_info["input_cost_per_token"] == litellm.model_cost["azure/gpt-4o"]["input_cost_per_token"] - - def test_unmappable_deployment_name_still_logs_error(self): - router = self._router_with_azure_deployment("azure/my-custom-deployment-name") - - with patch( - "litellm.router.verbose_router_logger.error" - ) as mock_error: - model_info = router.get_router_model_info( - deployment=None, received_model_name="my-group", id="azure-base-model-test-id" - ) - - assert any( - "Could not identify azure model" in str(call) - for call in mock_error.call_args_list - ), "expected the error log for an unmappable azure deployment name" - # unmappable names resolve to a zeroed stub — unchanged behavior - assert model_info.get("max_input_tokens") is None - - def test_explicit_base_model_still_wins(self): - router = litellm.Router( - model_list=[ - { - "model_name": "my-group", - "litellm_params": { - "model": "azure/some-deployment", - "api_key": "fake-key", - "api_base": "https://fake.openai.azure.com", - }, - "model_info": { - "id": "azure-base-model-test-id", - "base_model": "azure/gpt-4o-mini", - }, - } - ] - ) - - model_info = router.get_router_model_info( - deployment=None, received_model_name="my-group", id="azure-base-model-test-id" - ) - assert model_info["max_input_tokens"] == litellm.model_cost["azure/gpt-4o-mini"]["max_input_tokens"] - -def test_model_group_info_intersects_supported_reasoning_efforts(): - router = litellm.Router( - model_list=[ - { - "model_name": "smart-group", - "litellm_params": {"model": "anthropic/opus-like"}, - "model_info": {"id": "opus-like-deployment"}, - }, - { - "model_name": "smart-group", - "litellm_params": {"model": "openai/mini-like"}, - "model_info": {"id": "mini-like-deployment"}, - }, - ] - ) - - def _model_info(model_id: str, model_name: str): - if model_id == "opus-like-deployment": - return { - "key": model_name, - "litellm_provider": "anthropic", - "mode": "chat", - "supports_reasoning": True, - "supports_xhigh_reasoning_effort": True, - "supports_max_reasoning_effort": True, - } - return { - "key": model_name, - "litellm_provider": "openai", - "mode": "chat", - "supports_reasoning": True, - "supports_none_reasoning_effort": False, - "supports_minimal_reasoning_effort": True, - "supports_xhigh_reasoning_effort": False, - } - - with patch.object(router, "get_deployment_model_info", side_effect=_model_info): - result = router._set_model_group_info( - model_group="smart-group", - user_facing_model_group_name="smart-group", - ) - - assert result is not None - # opus-like offers all seven levels, mini-like lacks none/xhigh/max; only the common set survives, - # so the group never advertises an effort routing could hand to a deployment that rejects it. - assert result.supported_reasoning_efforts == ("minimal", "low", "medium", "high") - - -def test_model_group_info_reasoning_efforts_are_unknown_when_any_deployment_is_off_the_map(): - """The router fills every ModelInfo key, so a deployment absent from the model map arrives with - supports_reasoning None rather than with the key missing. Its synthesized entry carries no mode, - which is what separates it from a mapped non-reasoning model, and nothing being known about it is - no evidence that the unknown deployment accepts levels its mapped sibling supports. The group - therefore reports unknown instead of advertising a value routing might send to either one.""" - router = litellm.Router( - model_list=[ - { - "model_name": "smart-group", - "litellm_params": {"model": "anthropic/opus-like"}, - "model_info": {"id": "opus-like-deployment"}, - }, - { - "model_name": "smart-group", - "litellm_params": {"model": "openai/unmapped-model"}, - "model_info": {"id": "unmapped-deployment"}, - }, - ] - ) - - def _model_info(model_id: str, model_name: str): - if model_id == "opus-like-deployment": - return { - "key": model_name, - "litellm_provider": "anthropic", - "mode": "chat", - "supports_reasoning": True, - "supports_max_reasoning_effort": True, - } - return {"key": model_name, "litellm_provider": "openai", "mode": None, "supports_reasoning": None} - - with patch.object(router, "get_deployment_model_info", side_effect=_model_info): - result = router._set_model_group_info( - model_group="smart-group", - user_facing_model_group_name="smart-group", - ) - - assert result is not None - assert result.supported_reasoning_efforts is None - - - -def test_model_group_info_surfaces_supports_parallel_function_calling(local_model_cost_map): - """``/model_group/info`` folds each deployment's registry flags into the group; a deployment whose - registry entry declares parallel function calling must flip the group to True instead of False.""" - router = litellm.Router( - model_list=[ - { - "model_name": "glm-group", - "litellm_params": {"model": "together_ai/zai-org/GLM-5.3-Flash", "api_key": "fake-key"}, - } - ] - ) - - result = router._set_model_group_info(model_group="glm-group", user_facing_model_group_name="glm-group") - - assert result is not None - assert result.supports_parallel_function_calling is True - - -def test_model_group_info_reasoning_efforts_empty_on_a_mapped_non_reasoning_deployment(): - """A group mixing a reasoning model with one the map knows is not a reasoning model shares no - level, so it advertises none and the picker offers nothing rather than a level routing would - hand to a deployment that rejects it.""" - router = litellm.Router( - model_list=[ - { - "model_name": "mixed-group", - "litellm_params": {"model": "anthropic/opus-like"}, - "model_info": {"id": "opus-like-deployment"}, - }, - { - "model_name": "mixed-group", - "litellm_params": {"model": "openai/plain-chat"}, - "model_info": {"id": "plain-chat-deployment"}, - }, - ] - ) - - def _model_info(model_id: str, model_name: str): - if model_id == "opus-like-deployment": - return { - "key": model_name, - "litellm_provider": "anthropic", - "mode": "chat", - "supports_reasoning": True, - "supports_max_reasoning_effort": True, - } - return {"key": model_name, "litellm_provider": "openai", "mode": "chat", "supports_reasoning": None} - - with patch.object(router, "get_deployment_model_info", side_effect=_model_info): - result = router._set_model_group_info( - model_group="mixed-group", - user_facing_model_group_name="mixed-group", - ) - - assert result is not None - assert result.supported_reasoning_efforts == () - - -def test_model_group_info_reasoning_efforts_ignore_a_value_declared_in_model_info(): - """The group's levels are computed from its deployments, so a value an operator left in one - deployment's model_info must not seed them. Seeding let the first deployment read narrow the - whole group while the same value on any other deployment was silently ignored.""" - router = litellm.Router( - model_list=[ - { - "model_name": "declared-group", - "litellm_params": {"model": "openai/first-reasoner"}, - "model_info": {"id": "first-deployment"}, - }, - { - "model_name": "declared-group", - "litellm_params": {"model": "openai/second-reasoner"}, - "model_info": {"id": "second-deployment"}, - }, - ] - ) - - def _model_info(model_id: str, model_name: str): - info = { - "key": model_name, - "litellm_provider": "openai", - "mode": "chat", - "supports_reasoning": True, - "supports_none_reasoning_effort": True, - } - if model_id == "first-deployment": - info["supported_reasoning_efforts"] = ("high",) - return info - - with patch.object(router, "get_deployment_model_info", side_effect=_model_info): - result = router._set_model_group_info( - model_group="declared-group", - user_facing_model_group_name="declared-group", - ) - - assert result is not None - assert result.supported_reasoning_efforts == ("none", "minimal", "low", "medium", "high") - - -def test_model_group_info_survives_a_junk_typed_operator_effort_value(): - """A deployment's registered model_info reads back with whatever the operator wrote under any - key, so a wrong-typed supported_reasoning_efforts must not fail the group's info. Only the - constructor's trailing override keeps the junk away from ModelGroupInfo validation.""" - router = litellm.Router( - model_list=[ - { - "model_name": "junk-declared-group", - "litellm_params": {"model": "openai/lone-reasoner"}, - "model_info": {"id": "junk-deployment"}, - }, - ] - ) - - def _model_info(model_id: str, model_name: str): - return { - "key": model_name, - "litellm_provider": "openai", - "mode": "chat", - "supports_reasoning": True, - "supports_none_reasoning_effort": True, - "supported_reasoning_efforts": "high", - } - - with patch.object(router, "get_deployment_model_info", side_effect=_model_info): - result = router._set_model_group_info( - model_group="junk-declared-group", - user_facing_model_group_name="junk-declared-group", - ) - - assert result is not None - assert result.supported_reasoning_efforts == ("none", "minimal", "low", "medium", "high") - - -def test_model_group_info_reasoning_efforts_are_unknown_for_an_operator_declared_mode(): - """A deployment is registered in the cost map under its own id with whatever model_info the - operator wrote, so a mode they set themselves reads back exactly like one the map supplied. Only - a mode the map supplied marks the deployment as known. An off-map deployment carrying an - operator mode remains unknown and must keep the whole group's level support unknown.""" - from litellm.router_utils.reasoning_effort_capability import resolve_supported_reasoning_efforts - - mapped_model = "openai/gpt-5.6-sol" - expected = resolve_supported_reasoning_efforts( - litellm.get_model_info(model=mapped_model), - deployment_is_mapped=True, - ) - assert expected - - router = litellm.Router( - model_list=[ - { - "model_name": "smart-group", - "litellm_params": {"model": mapped_model, "api_key": "sk-fake"}, - "model_info": {"id": "mapped-deployment"}, - }, - { - "model_name": "smart-group", - "litellm_params": {"model": "openai/a-model-the-map-never-heard-of", "api_key": "sk-fake"}, - "model_info": {"id": "off-map-deployment", "mode": "chat"}, - }, - ] - ) - - result = router._set_model_group_info( - model_group="smart-group", - user_facing_model_group_name="smart-group", - ) - - assert result is not None - assert result.supported_reasoning_efforts is None - - -class TestAddDeploymentApiBaseProviderResolution: - def test_bare_model_with_known_api_base_initializes(self): - router = litellm.Router( - model_list=[ - { - "model_name": "groq-pinned", - "litellm_params": { - "model": "llama-3.3-70b-versatile", - "api_base": "https://api.groq.com/openai/v1", - "api_key": "fake-key", - }, - }, - { - "model_name": "deepseek-pinned", - "litellm_params": { - "model": "deepseek-chat", - "api_base": "https://api.deepseek.com/v1", - "api_key": "fake-key", - }, - }, - ] - ) - - model_list = router.get_model_list() - assert model_list is not None - assert {m["model_name"] for m in model_list} == {"groq-pinned", "deepseek-pinned"} - - def test_bare_model_with_unknown_api_base_still_raises(self): - with pytest.raises(litellm.BadRequestError, match="LLM Provider NOT provided"): - litellm.Router( - model_list=[ - { - "model_name": "mystery", - "litellm_params": { - "model": "some-unknown-model", - "api_base": "https://llm.internal.example.com/v1", - "api_key": "fake-key", - }, - } - ] - ) - - def test_explicit_custom_llm_provider_beats_api_base_endpoint_match(self): - router = litellm.Router( - model_list=[ - { - "model_name": "openai-via-gateway", - "litellm_params": { - "model": "gpt-3.5-turbo", - "custom_llm_provider": "openai", - "api_base": "https://api.groq.com/openai/v1", - "api_key": "fake-key", - }, - } - ] - ) - - deployment = router.get_deployment_by_model_group_name("openai-via-gateway") - assert deployment is not None - assert deployment.litellm_params.custom_llm_provider == "openai" - -# ===================================================================== -# anthropic_messages mid-stream-fallback helpers, added for #24004 -# (mid-stream fallback not supported for anthropic_messages route type). -# -# anthropic_messages goes through _ageneric_api_call_with_fallbacks rather -# than _acompletion, so its returned iterator was never wrapped by the chat -# completions fallback handler: an SSE `event: error` frame from a native -# Anthropic/Bedrock passthrough passed through to the client silently, and a -# MidStreamFallbackError raised by the completion-bridge path's -# CustomStreamWrapper (e.g. a Vertex AI transport drop) propagated -# unhandled. -# -# Targets the helpers introduced on Router: -# - _aanthropic_messages_streaming_iterator -# - _aanthropic_messages_fallback_attempt -# - _aanthropic_messages_with_streaming_fallbacks -# - _dispatch_generic_call_type -# ===================================================================== - - -async def _anthropic_messages_empty_generator(): - return - yield # pragma: no cover - makes this an async generator - - -def _anthropic_messages_make_wrapper() -> FallbackAwareAnthropicMessagesStream: - """A minimal wrapper for tests that call _aanthropic_messages_fallback_attempt - directly, bypassing _aanthropic_messages_streaming_iterator.""" - return FallbackAwareAnthropicMessagesStream(_anthropic_messages_empty_generator(), object()) - - -def _anthropic_messages_make_router() -> Router: - return Router( - model_list=[ - { - "model_name": "primary", - "litellm_params": { - "model": "anthropic/claude-sonnet-4-5", - "api_key": "sk-test", - }, - }, - { - "model_name": "fallback", - "litellm_params": { - "model": "bedrock/anthropic.claude-sonnet-4-5", - }, - }, - ] - ) - - -class _AnthropicMessagesFakeByteStream: - """Minimal AsyncIterator[bytes], carrying _hidden_params like - AnthropicMessagesStreamingResponse does.""" - - def __init__(self, chunks: list) -> None: - self._chunks = list(chunks) - self._hidden_params = {"additional_headers": {"x-amzn-requestid": "req-1"}} - self.closed = False - - def __aiter__(self): - return self - - async def __anext__(self) -> bytes: - if not self._chunks: - raise StopAsyncIteration - return self._chunks.pop(0) - - async def aclose(self) -> None: - self.closed = True - - -class _AnthropicMessagesRaisingByteStream: - """Simulates the completion-bridge path: no error SSE chunk is ever - yielded, the underlying CustomStreamWrapper raises MidStreamFallbackError - directly out of the iterator instead (a Vertex AI transport drop).""" - - def __init__(self, chunks: list, error: Exception) -> None: - self._chunks = list(chunks) - self._error = error - self._hidden_params: dict = {} - self.closed = False - - def __aiter__(self): - return self - - async def __anext__(self) -> bytes: - if self._chunks: - return self._chunks.pop(0) - raise self._error - - async def aclose(self) -> None: - self.closed = True - - -class _AnthropicMessagesFallbackByteStream: - def __init__(self, chunks: list, hidden_params: dict | None = None) -> None: - self._chunks = list(chunks) - self._hidden_params = hidden_params if hidden_params is not None else {} - - def __aiter__(self): - return self - - async def __anext__(self) -> bytes: - if not self._chunks: - raise StopAsyncIteration - return self._chunks.pop(0) - - -def _anthropic_messages_overloaded_error_chunk() -> bytes: - return ( - b"event: error\n" - b'data: {"type": "error", "error": {"type": "overloaded_error", "message": "Overloaded"}}\n\n' - ) - - -def _anthropic_messages_invalid_request_error_chunk() -> bytes: - return ( - b"event: error\n" - b'data: {"type": "error", "error": {"type": "invalid_request_error", "message": "bad request"}}\n\n' - ) - - -def _anthropic_messages_rate_limit_error_chunk() -> bytes: - return ( - b"event: error\n" - b'data: {"type": "error", "error": {"type": "rate_limit_error", "message": "Too many requests"}}\n\n' - ) - - -def _anthropic_messages_content_chunk(text: str = "hi") -> bytes: - payload = f'{{"type": "content_block_delta", "delta": {{"type": "text_delta", "text": "{text}"}}}}' - return f"event: content_block_delta\ndata: {payload}\n\n".encode() - - -def _anthropic_messages_message_start_chunk() -> bytes: - """A lifecycle/bookkeeping frame Anthropic sends before any real content - - routinely the very first event before an overload error.""" - return b'event: message_start\ndata: {"type": "message_start", "message": {"id": "msg_1"}}\n\n' - - -def _anthropic_messages_ping_chunk() -> bytes: - return b'event: ping\ndata: {"type": "ping"}\n\n' - - -# -------- _aanthropic_messages_streaming_iterator (passthrough) -------- - - -@pytest.mark.asyncio -async def test_anthropic_messages_streaming_iterator_passthrough(): - """Without any error chunk, the wrapper forwards every chunk unchanged - and carries the source iterator's _hidden_params through (so response - headers like Bedrock's request-id keep flowing to the client).""" - router = _anthropic_messages_make_router() - source = _AnthropicMessagesFakeByteStream( - [_anthropic_messages_content_chunk("hi"), _anthropic_messages_content_chunk(" there")] - ) - - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, initial_kwargs={"model": "primary"} - ) - - collected = [chunk async for chunk in wrapped] - assert collected == [_anthropic_messages_content_chunk("hi"), _anthropic_messages_content_chunk(" there")] - assert wrapped._hidden_params["additional_headers"]["x-amzn-requestid"] == "req-1" - - -@pytest.mark.asyncio -async def test_anthropic_messages_streaming_iterator_flushes_buffered_lifecycle_frames_in_order(): - """Regression: lifecycle frames held back to guard against a mid-stream - fallback must still reach the client, in order, once real content - arrives - buffering them for the fallback-safety check must not silently - drop them on the happy path.""" - router = _anthropic_messages_make_router() - message_stop = b'event: message_stop\ndata: {"type": "message_stop"}\n\n' - source = _AnthropicMessagesFakeByteStream( - [_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("hi"), message_stop] - ) - - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, initial_kwargs={"model": "primary"} - ) - - collected = [chunk async for chunk in wrapped] - assert collected == [_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("hi"), message_stop] - - -@pytest.mark.asyncio -async def test_anthropic_messages_streaming_iterator_flushes_buffered_frames_on_stream_end(): - """Regression: if the primary stream ends with only lifecycle frames and - no content and no error, the buffered frames must still reach the - client rather than being silently swallowed.""" - router = _anthropic_messages_make_router() - message_stop = b'event: message_stop\ndata: {"type": "message_stop"}\n\n' - source = _AnthropicMessagesFakeByteStream([_anthropic_messages_message_start_chunk(), message_stop]) - - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, initial_kwargs={"model": "primary"} - ) - - collected = [chunk async for chunk in wrapped] - assert collected == [_anthropic_messages_message_start_chunk(), message_stop] - - with pytest.raises(StopAsyncIteration): - await wrapped.__anext__() - - -@pytest.mark.asyncio -async def test_anthropic_messages_content_coalesced_with_error_in_one_physical_chunk_skips_fallback(): - """Greptile review round: transport-level buffering can coalesce a real - content_block_delta and a following retriable error into ONE physical - read from the source iterator. Since the whole chunk (content and error - together) is forwarded to the client atomically, the client genuinely - receives the content - so no fallback must be attempted, exactly as if - the two events had arrived as separate reads.""" - router = _anthropic_messages_make_router() - coalesced_chunk = _anthropic_messages_content_chunk("partial") + _anthropic_messages_overloaded_error_chunk() - source = _AnthropicMessagesFakeByteStream([coalesced_chunk]) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(return_value=_AnthropicMessagesFallbackByteStream([])), - ) as mock_fallback: - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, - initial_kwargs={"model": "primary"}, - ) - collected = [chunk async for chunk in wrapped] - - assert collected == [coalesced_chunk] - mock_fallback.assert_not_awaited() - - -@pytest.mark.asyncio -async def test_anthropic_messages_ping_behind_buffered_lifecycle_frame_is_dropped(): - """Bugbot regression: a `ping` keepalive behind buffered lifecycle frames - carries no content and is dropped outright rather than buffered - - otherwise a slow-starting connection sending many pings could grow the - pre-content buffer without bound.""" - router = _anthropic_messages_make_router() - source = _AnthropicMessagesFakeByteStream( - [ - _anthropic_messages_message_start_chunk(), - _anthropic_messages_ping_chunk(), - _anthropic_messages_content_chunk("hi"), - ] - ) - - wrapped = await router._aanthropic_messages_streaming_iterator(response=source, initial_kwargs={"model": "primary"}) - collected = [chunk async for chunk in wrapped] - - assert collected == [_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("hi")] - - -@pytest.mark.asyncio -async def test_anthropic_messages_leading_ping_keepalive_is_forwarded_live(): - """A `ping` that no lifecycle frame precedes is how a hold-back turn keeps - its connection alive (AgenticAnthropicStreamingIterator), so it must reach - the client at once rather than wait behind the pre-content buffer.""" - router = _anthropic_messages_make_router() - content_released = asyncio.Event() - - async def source(): - yield _anthropic_messages_ping_chunk() - await content_released.wait() - yield _anthropic_messages_message_start_chunk() - yield _anthropic_messages_content_chunk("hi") - - wrapped = await router._aanthropic_messages_streaming_iterator(response=source(), initial_kwargs={"model": "primary"}) - - assert await asyncio.wait_for(wrapped.__anext__(), timeout=1) == _anthropic_messages_ping_chunk() - content_released.set() - assert [chunk async for chunk in wrapped] == [ - _anthropic_messages_message_start_chunk(), - _anthropic_messages_content_chunk("hi"), - ] - - -@pytest.mark.asyncio -async def test_anthropic_messages_leading_ping_does_not_disqualify_fallback(): - """A live-forwarded leading `ping` commits nothing: a retriable error after - it still falls back, and the fallback's own lifecycle follows the ping cleanly.""" - router = _anthropic_messages_make_router() - source = _AnthropicMessagesFakeByteStream( - [_anthropic_messages_ping_chunk(), _anthropic_messages_overloaded_error_chunk()] - ) - fallback_message_start = b'event: message_start\ndata: {"type": "message_start", "message": {"id": "msg_2"}}\n\n' - fallback_stream = _AnthropicMessagesFallbackByteStream( - [fallback_message_start, _anthropic_messages_content_chunk("fallback answer")] - ) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(return_value=fallback_stream), - ): - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, - initial_kwargs={"model": "primary"}, - ) - collected = [chunk async for chunk in wrapped] - - assert collected == [ - _anthropic_messages_ping_chunk(), - fallback_message_start, - _anthropic_messages_content_chunk("fallback answer"), - ] - - -@pytest.mark.asyncio -async def test_anthropic_messages_hold_back_retrieval_failure_reaches_client_without_fallback(): - """The hold-back iterator's own retrieval-failure frame is the gateway's verdict, not a - provider failure: a configured fallback stays untouched and the client reads the error - right after the live keepalive.""" - router = _anthropic_messages_make_router() - source = _AnthropicMessagesFakeByteStream( - [_anthropic_messages_ping_chunk(), SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES] - ) - fallback = AsyncMock( - return_value=_AnthropicMessagesFallbackByteStream([_anthropic_messages_content_chunk("fallback answer")]) - ) - - with patch.object(router, "async_function_with_fallbacks_common_utils", new=fallback): - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, - initial_kwargs={"model": "primary"}, - ) - collected = [chunk async for chunk in wrapped] - - fallback.assert_not_called() - assert collected == [_anthropic_messages_ping_chunk(), SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES] - - -@pytest.mark.asyncio -async def test_anthropic_messages_pre_content_buffer_cap_forces_commit(): - """Bugbot regression: a hostile or pathological upstream that never emits - real content or an error must not grow the pre-content lifecycle buffer - without bound - hitting MAX_BUFFERED_PRE_CONTENT_ANTHROPIC_CHUNKS commits - to the primary stream early, exactly as real content arriving would.""" - router = _anthropic_messages_make_router() - lifecycle_chunk = _anthropic_messages_message_start_chunk() - error_chunk = _anthropic_messages_overloaded_error_chunk() - chunks = [lifecycle_chunk] * (MAX_BUFFERED_PRE_CONTENT_ANTHROPIC_CHUNKS + 5) + [error_chunk] - source = _AnthropicMessagesFakeByteStream(chunks) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(return_value=_AnthropicMessagesFallbackByteStream([])), - ) as mock_fallback: - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, initial_kwargs={"model": "primary"} - ) - collected = [chunk async for chunk in wrapped] - - mock_fallback.assert_not_awaited() - assert collected.count(lifecycle_chunk) == MAX_BUFFERED_PRE_CONTENT_ANTHROPIC_CHUNKS + 5 - assert collected[-1] == error_chunk - - -@pytest.mark.asyncio -async def test_anthropic_messages_ping_coalesced_with_content_in_one_physical_chunk_is_forwarded(): - """Greptile/Bugbot regression: transport-level buffering can coalesce a - `ping` keepalive and a real content_block_delta into ONE physical read. - The pre-content ping-drop must only discard PURE ping frames - dropping - the whole coalesced chunk would silently lose generated content.""" - router = _anthropic_messages_make_router() - coalesced_chunk = _anthropic_messages_ping_chunk() + _anthropic_messages_content_chunk("hi") - source = _AnthropicMessagesFakeByteStream([coalesced_chunk]) - - wrapped = await router._aanthropic_messages_streaming_iterator(response=source, initial_kwargs={"model": "primary"}) - collected = [chunk async for chunk in wrapped] - - assert collected == [coalesced_chunk] - - -@pytest.mark.asyncio -async def test_anthropic_messages_ping_coalesced_with_retriable_error_still_falls_back(): - """Greptile/Bugbot regression: a physical chunk coalescing a `ping` with a - retriable `event: error` must not be discarded as a keepalive - the error - inside it must still trigger the mid-stream fallback.""" - router = _anthropic_messages_make_router() - coalesced_chunk = _anthropic_messages_ping_chunk() + _anthropic_messages_overloaded_error_chunk() - source = _AnthropicMessagesFakeByteStream([coalesced_chunk]) - fallback_stream = _AnthropicMessagesFallbackByteStream([_anthropic_messages_content_chunk("fallback answer")]) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(return_value=fallback_stream), - ) as mock_fallback: - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, initial_kwargs={"model": "primary"} - ) - collected = [chunk async for chunk in wrapped] - - mock_fallback.assert_awaited_once() - assert collected == [_anthropic_messages_content_chunk("fallback answer")] - - -# -------- _aanthropic_messages_fallback_attempt -------- - - -@pytest.mark.asyncio -async def test_aanthropic_messages_fallback_attempt_yields_fallback_stream(): - """Direct-call regression: the fallback-attempt helper re-enters the - Router's fallback chain and forwards whatever the fallback produces.""" - router = _anthropic_messages_make_router() - fallback_stream = _AnthropicMessagesFallbackByteStream([_anthropic_messages_content_chunk("fallback answer")]) - error = MidStreamFallbackError(message="overloaded", model="primary", llm_provider="anthropic") - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(return_value=fallback_stream), - ) as mock_fallback: - collected = [ - chunk - async for chunk in router._aanthropic_messages_fallback_attempt( - error, - {"model": "primary", "messages": [{"role": "user", "content": "hi"}]}, - _anthropic_messages_make_wrapper(), - ) - ] - - assert collected == [_anthropic_messages_content_chunk("fallback answer")] - mock_fallback.assert_awaited_once() - assert mock_fallback.await_args.kwargs["e"] is error - - -@pytest.mark.asyncio -async def test_aanthropic_messages_fallback_attempt_raises_original_exception_on_double_failure(): - """Direct-call regression: when the fallback attempt itself fails with a - MidStreamFallbackError wrapping a real provider exception, that real - exception must surface rather than the internal wrapper exception.""" - router = _anthropic_messages_make_router() - error = MidStreamFallbackError(message="overloaded", model="primary", llm_provider="anthropic") - original_exception = litellm.APIError( - status_code=503, message="fallback also overloaded", llm_provider="bedrock", model="fallback" - ) - fallback_failure = MidStreamFallbackError( - message="fallback failed", model="fallback", llm_provider="bedrock", original_exception=original_exception - ) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(side_effect=fallback_failure), - ): - with pytest.raises(litellm.APIError) as exc_info: - async for _ in router._aanthropic_messages_fallback_attempt( - error, {"model": "primary"}, _anthropic_messages_make_wrapper() - ): - pass - - assert exc_info.value is original_exception - - -@pytest.mark.asyncio -async def test_aanthropic_messages_fallback_attempt_yields_non_streaming_fallback_response(): - """Bugbot regression: a fallback that resolves to a non-streaming - response (no __aiter__, e.g. an agentic tool-use interception loop) must - be synthesized into a valid SSE byte sequence, not yielded as a raw dict - into a byte stream - the generator is typed AsyncGenerator[bytes, None] - and every item reaching the client must be a real SSE frame.""" - router = _anthropic_messages_make_router() - error = MidStreamFallbackError(message="overloaded", model="primary", llm_provider="anthropic") - non_streaming_response = {"id": "msg_1", "type": "message", "content": [{"type": "text", "text": "hi"}]} - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(return_value=non_streaming_response), - ): - collected = [ - item - async for item in router._aanthropic_messages_fallback_attempt( - error, {"model": "primary"}, _anthropic_messages_make_wrapper() - ) - ] - - assert all(isinstance(item, bytes) for item in collected) - event_types = [item.split(b"\n")[0].removeprefix(b"event: ") for item in collected] - assert event_types == [ - b"message_start", - b"content_block_start", - b"content_block_delta", - b"content_block_stop", - b"message_delta", - b"message_stop", - ] - assert b'"text": "hi"' in collected[2] - - -@pytest.mark.asyncio -async def test_aanthropic_messages_fallback_attempt_reraises_plain_exception_on_double_failure(): - """Direct-call regression: when the fallback attempt fails with a plain - exception (not a MidStreamFallbackError), that exception itself must - propagate unchanged.""" - router = _anthropic_messages_make_router() - error = MidStreamFallbackError(message="overloaded", model="primary", llm_provider="anthropic") - fallback_failure = ValueError("no healthy deployments") - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(side_effect=fallback_failure), - ): - with pytest.raises(ValueError, match="no healthy deployments") as exc_info: - async for _ in router._aanthropic_messages_fallback_attempt( - error, {"model": "primary"}, _anthropic_messages_make_wrapper() - ): - pass - - assert exc_info.value is fallback_failure - - -# -------- _aanthropic_messages_with_streaming_fallbacks -------- - - -@pytest.mark.asyncio -async def test_aanthropic_messages_with_streaming_fallbacks_non_streaming_passthrough(): - """A non-streaming response (plain dict) is returned unchanged, never wrapped.""" - router = _anthropic_messages_make_router() - plain_response = {"id": "msg_1", "type": "message"} - - async def fake_original(**_kwargs): - return plain_response - - with patch.object( - router, - "_ageneric_api_call_with_fallbacks", - new=AsyncMock(return_value=plain_response), - ): - out = await router._aanthropic_messages_with_streaming_fallbacks( - original_function=fake_original, - model="primary", - stream=False, - ) - assert out is plain_response - - -@pytest.mark.asyncio -async def test_aanthropic_messages_with_streaming_fallbacks_wraps_streaming_iterator(): - """A streaming response is wrapped via _aanthropic_messages_streaming_iterator.""" - router = _anthropic_messages_make_router() - streaming_iter = _AnthropicMessagesFakeByteStream([_anthropic_messages_content_chunk()]) - wrapped_marker = object() - - async def fake_original(**_kwargs): - return streaming_iter - - with ( - patch.object( - router, - "_ageneric_api_call_with_fallbacks", - new=AsyncMock(return_value=streaming_iter), - ), - patch.object( - router, - "_aanthropic_messages_streaming_iterator", - new=AsyncMock(return_value=wrapped_marker), - ) as mock_wrap, - ): - out = await router._aanthropic_messages_with_streaming_fallbacks( - original_function=fake_original, - model="primary", - stream=True, - ) - assert out is wrapped_marker - mock_wrap.assert_awaited_once() - - -# -------- mid-stream error handling -------- - - -@pytest.mark.asyncio -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.""" - router = _anthropic_messages_make_router() - source = _AnthropicMessagesFakeByteStream([_anthropic_messages_overloaded_error_chunk()]) - fallback_stream = _AnthropicMessagesFallbackByteStream([_anthropic_messages_content_chunk("fallback answer")]) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(return_value=fallback_stream), - ) as mock_fallback: - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, - initial_kwargs={"model": "primary", "messages": [{"role": "user", "content": "hi"}]}, - ) - collected = [chunk async for chunk in wrapped] - - assert collected == [_anthropic_messages_content_chunk("fallback answer")] - mock_fallback.assert_awaited_once() - raised = mock_fallback.await_args.kwargs["e"] - assert isinstance(raised, MidStreamFallbackError) - assert raised.status_code == 503 - assert raised.is_pre_first_chunk is True - assert source.closed is True - - -@pytest.mark.asyncio -async def test_anthropic_messages_mid_stream_error_preserves_real_status_code(): - """Bugbot regression: the MidStreamFallbackError raised for a detected SSE - `event: error` frame must carry the error's REAL parsed status code - (via original_exception), not silently default to 503 for every error - type - a rate_limit_error (429) must surface as 429, not 503.""" - router = _anthropic_messages_make_router() - source = _AnthropicMessagesFakeByteStream([_anthropic_messages_rate_limit_error_chunk()]) - fallback_stream = _AnthropicMessagesFallbackByteStream([_anthropic_messages_content_chunk("fallback answer")]) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(return_value=fallback_stream), - ) as mock_fallback: - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, - initial_kwargs={"model": "primary", "messages": [{"role": "user", "content": "hi"}]}, - ) - [chunk async for chunk in wrapped] - - raised = mock_fallback.await_args.kwargs["e"] - assert isinstance(raised, MidStreamFallbackError) - assert raised.status_code == 429 - assert raised.original_exception is not None - assert raised.original_exception.status_code == 429 - assert raised.original_exception.llm_provider == "anthropic" - - -def test_merge_fallback_hidden_params_direct_call(): - """Direct-call regression: merge_fallback_hidden_params combines the - fallback's hidden params/headers with whatever was already present, - with the fallback's values winning on key collisions.""" - wrapper = FallbackAwareAnthropicMessagesStream( - _anthropic_messages_empty_generator(), - _AnthropicMessagesFakeByteStream([]), # carries {"additional_headers": {"x-amzn-requestid": "req-1"}} - ) - wrapper.merge_fallback_hidden_params( - {"model_id": "fallback-deployment"}, - {"x-amzn-requestid": "req-2", "x-fallback-only": "yes"}, - ) - assert wrapper._hidden_params["model_id"] == "fallback-deployment" - assert wrapper._hidden_params["additional_headers"] == { - "x-amzn-requestid": "req-2", - "x-fallback-only": "yes", - } - - -def test_anthropic_stream_should_drop_pre_content_ping_direct_call(): - ping = _anthropic_messages_ping_chunk() - content = _anthropic_messages_content_chunk("hi") - assert _anthropic_stream_should_drop_pre_content_ping(ping, has_generated_content=False) is True - assert _anthropic_stream_should_drop_pre_content_ping(ping, has_generated_content=True) is False - assert _anthropic_stream_should_drop_pre_content_ping(content, has_generated_content=False) is False - - -def test_anthropic_stream_forwards_ping_live_direct_call(): - ping = _anthropic_messages_ping_chunk() - content = _anthropic_messages_content_chunk("hi") - assert _anthropic_stream_forwards_ping_live(ping, has_generated_content=False, buffered_chunk_count=0) is True - assert _anthropic_stream_forwards_ping_live(ping, has_generated_content=False, buffered_chunk_count=1) is False - assert _anthropic_stream_forwards_ping_live(ping, has_generated_content=True, buffered_chunk_count=0) is False - assert _anthropic_stream_forwards_ping_live(content, has_generated_content=False, buffered_chunk_count=0) is False - - -def test_anthropic_stream_error_is_gateway_verdict_direct_call(): - assert _anthropic_stream_error_is_gateway_verdict(SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES) is True - assert _anthropic_stream_error_is_gateway_verdict(_anthropic_messages_overloaded_error_chunk()) is False - assert _anthropic_stream_error_is_gateway_verdict(_anthropic_messages_ping_chunk()) is False - - -def test_fallback_aware_stream_reports_withheld_output_of_its_current_source(): - """The proxy's cancel-refund guard reads this flag off the router wrapper, so it - must reflect the stream actually being drained: the primary, then the fallback.""" - - class _HoldingBack: - _hidden_params = {"additional_headers": {}} - has_buffered_provider_output = True - - wrapper = FallbackAwareAnthropicMessagesStream(_anthropic_messages_empty_generator(), _HoldingBack()) - assert wrapper.has_buffered_provider_output is True - - wrapper.adopt_fallback_source(_AnthropicMessagesFakeByteStream([])) - assert wrapper.has_buffered_provider_output is False - - -def test_is_retriable_anthropic_status_direct_call(): - assert _is_retriable_anthropic_status(429) is True - assert _is_retriable_anthropic_status(503) is True - assert _is_retriable_anthropic_status(500) is True - assert _is_retriable_anthropic_status(400) is False - assert _is_retriable_anthropic_status(404) is False - - -def test_anthropic_stream_should_decline_fallback_direct_call(): - pre_first_chunk_error = MidStreamFallbackError( - message="overloaded", model="primary", llm_provider="anthropic", is_pre_first_chunk=True - ) - post_first_chunk_error = MidStreamFallbackError( - message="overloaded", model="primary", llm_provider="anthropic", is_pre_first_chunk=False - ) - assert _anthropic_stream_should_decline_fallback(False, pre_first_chunk_error) is False - assert _anthropic_stream_should_decline_fallback(True, pre_first_chunk_error) is True - assert _anthropic_stream_should_decline_fallback(False, post_first_chunk_error) is True - - -def test_anthropic_stream_commits_now_direct_call(): - content = _anthropic_messages_content_chunk("hi") - lifecycle_chunk = _anthropic_messages_message_start_chunk() - assert _anthropic_stream_commits_now(content, has_generated_content=False, buffered_chunk_count=0) is True - assert _anthropic_stream_commits_now(content, has_generated_content=True, buffered_chunk_count=0) is False - assert ( - _anthropic_stream_commits_now( - lifecycle_chunk, - has_generated_content=False, - buffered_chunk_count=MAX_BUFFERED_PRE_CONTENT_ANTHROPIC_CHUNKS, - ) - is True - ) - assert ( - _anthropic_stream_commits_now( - lifecycle_chunk, - has_generated_content=False, - buffered_chunk_count=MAX_BUFFERED_PRE_CONTENT_ANTHROPIC_CHUNKS - 1, - ) - is False - ) - - -@pytest.mark.asyncio -async def test_anthropic_messages_fallback_merges_fallback_hidden_params(): - """Bugbot regression: after a successful mid-stream fallback, the - wrapper's _hidden_params must reflect the FALLBACK deployment's own - provider headers (e.g. a different Bedrock request-id), not stay - frozen on the primary's - raw bytes can't carry per-item _hidden_params - the way a ModelResponseStream/ResponsesAPI event can, so the wrapper - itself is the only place left to expose them.""" - router = _anthropic_messages_make_router() - source = _AnthropicMessagesFakeByteStream( - [_anthropic_messages_overloaded_error_chunk()] - ) # carries x-amzn-requestid: req-1 - fallback_stream = _AnthropicMessagesFallbackByteStream( - [_anthropic_messages_content_chunk("fallback answer")], - hidden_params={"additional_headers": {"x-amzn-requestid": "req-2", "x-fallback-only": "yes"}}, - ) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(return_value=fallback_stream), - ): - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, - initial_kwargs={"model": "primary"}, - ) - _ = [chunk async for chunk in wrapped] - - headers = wrapped._hidden_params["additional_headers"] - assert headers["x-amzn-requestid"] == "req-2" - assert headers["x-fallback-only"] == "yes" - - -@pytest.mark.asyncio -async def test_aanthropic_messages_with_streaming_fallbacks_deep_copies_nested_metadata(): - """Bugbot regression: a shallow .copy() of kwargs still shares the - nested litellm_metadata/metadata dict objects with the primary attempt. - _update_kwargs_with_deployment mutates that dict in place with - deployment-specific fields, which must not leak into the fallback - request's metadata.""" - router = _anthropic_messages_make_router() - primary_metadata = {"model_group": "primary"} - streaming_iter_kwargs = {} - - async def fake_original(**_kwargs): - # Simulate _update_kwargs_with_deployment mutating the primary's - # litellm_metadata in place, as the real helper does. - primary_metadata["deployment"] = "primary-deployment-object" - return _AnthropicMessagesFakeByteStream([_anthropic_messages_content_chunk("hi")]) - - with patch.object( - router, - "_aanthropic_messages_streaming_iterator", - new=AsyncMock(side_effect=lambda **kwargs: streaming_iter_kwargs.update(kwargs) or "wrapped"), - ): - with patch.object( - router, - "_ageneric_api_call_with_fallbacks", - new=AsyncMock(side_effect=fake_original), - ): - await router._aanthropic_messages_with_streaming_fallbacks( - original_function=fake_original, - model="primary", - stream=True, - litellm_metadata=primary_metadata, - ) - - fallback_kwargs = streaming_iter_kwargs["initial_kwargs"] - assert fallback_kwargs["litellm_metadata"] is not primary_metadata - assert "deployment" not in fallback_kwargs["litellm_metadata"] - - -@pytest.mark.asyncio -async def test_aanthropic_messages_with_streaming_fallbacks_deep_copies_metadata_field(): - """Same regression as above for the (separate) `metadata` kwarg some - call sites use instead of `litellm_metadata`.""" - router = _anthropic_messages_make_router() - primary_metadata = {"tag": "primary"} - streaming_iter_kwargs = {} - - async def fake_original(**_kwargs): - primary_metadata["deployment"] = "primary-deployment-object" - return _AnthropicMessagesFakeByteStream([_anthropic_messages_content_chunk("hi")]) - - with patch.object( - router, - "_aanthropic_messages_streaming_iterator", - new=AsyncMock(side_effect=lambda **kwargs: streaming_iter_kwargs.update(kwargs) or "wrapped"), - ): - with patch.object( - router, - "_ageneric_api_call_with_fallbacks", - new=AsyncMock(side_effect=fake_original), - ): - await router._aanthropic_messages_with_streaming_fallbacks( - original_function=fake_original, - model="primary", - stream=True, - metadata=primary_metadata, - ) - - fallback_kwargs = streaming_iter_kwargs["initial_kwargs"] - assert fallback_kwargs["metadata"] is not primary_metadata - assert "deployment" not in fallback_kwargs["metadata"] - - -@pytest.mark.asyncio -async def test_anthropic_messages_fallback_triggers_after_lifecycle_only_frame(): - """Regression: Anthropic routinely sends a message_start lifecycle frame - before an overload error even fires. A lifecycle-only frame (no real - content) must not disqualify the fallback attempt, and must not reach - the client either - forwarding it and then appending the fallback's own - message_start would produce two overlapping message lifecycles on one - SSE stream. The primary's buffered lifecycle frame is discarded and the - client sees only the fallback's own, single, clean lifecycle.""" - router = _anthropic_messages_make_router() - source = _AnthropicMessagesFakeByteStream( - [_anthropic_messages_message_start_chunk(), _anthropic_messages_overloaded_error_chunk()] - ) - fallback_message_start = b'event: message_start\ndata: {"type": "message_start", "message": {"id": "msg_2"}}\n\n' - fallback_stream = _AnthropicMessagesFallbackByteStream( - [fallback_message_start, _anthropic_messages_content_chunk("fallback answer")] - ) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(return_value=fallback_stream), - ) as mock_fallback: - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, - initial_kwargs={"model": "primary"}, - ) - collected = [chunk async for chunk in wrapped] - - assert collected == [fallback_message_start, _anthropic_messages_content_chunk("fallback answer")] - assert collected.count(_anthropic_messages_message_start_chunk()) == 0, ( - "the primary's message_start must never reach the client" - ) - assert sum(1 for c in collected if c.startswith(b"event: message_start")) == 1, ( - "exactly one message_start must reach the client" - ) - mock_fallback.assert_awaited_once() - raised = mock_fallback.await_args.kwargs["e"] - assert raised.is_pre_first_chunk is True - - -@pytest.mark.asyncio -async def test_anthropic_messages_raised_error_after_real_content_does_not_restart_stream(): - """Regression: a MidStreamFallbackError raised directly by the source - iterator (the completion-bridge path's CustomStreamWrapper, e.g. a - transport drop) must not trigger a fallback once real content already - reached the client - that would append a second, overlapping message - lifecycle onto the same SSE stream. The original exception must - propagate to the caller instead.""" - router = _anthropic_messages_make_router() - content = _anthropic_messages_content_chunk("partial answer") - original_exception = litellm.APIError( - status_code=503, - message="stream reset", - llm_provider="vertex_ai", - model="primary", - ) - raised_error = MidStreamFallbackError( - message="stream reset", - model="primary", - llm_provider="vertex_ai", - original_exception=original_exception, - is_pre_first_chunk=False, - ) - source = _AnthropicMessagesRaisingByteStream([content], raised_error) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(), - ) as mock_fallback: - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, - initial_kwargs={"model": "primary"}, - ) - collected = [] - - async def _consume(): - async for chunk in wrapped: - collected.append(chunk) - - with pytest.raises(litellm.APIError) as exc_info: - await _consume() - - assert collected == [content] - assert exc_info.value is original_exception - mock_fallback.assert_not_awaited() - - -@pytest.mark.asyncio -async def test_anthropic_messages_fallback_also_catches_raised_midstream_error(): - """Regression for the completion-bridge path (deployments with no native - /v1/messages endpoint): its CustomStreamWrapper raises - MidStreamFallbackError directly (e.g. on a Vertex AI transport drop) - instead of yielding an SSE error chunk - the wrapper must catch that too.""" - router = _anthropic_messages_make_router() - raised_error = MidStreamFallbackError( - message="stream reset", - model="primary", - llm_provider="vertex_ai", - is_pre_first_chunk=True, - ) - source = _AnthropicMessagesRaisingByteStream([], raised_error) - fallback_stream = _AnthropicMessagesFallbackByteStream([_anthropic_messages_content_chunk("fallback answer")]) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(return_value=fallback_stream), - ) as mock_fallback: - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, - initial_kwargs={"model": "primary"}, - ) - collected = [chunk async for chunk in wrapped] - - assert collected == [_anthropic_messages_content_chunk("fallback answer")] - mock_fallback.assert_awaited_once() - assert mock_fallback.await_args.kwargs["e"] is raised_error - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - "raised_error", - [ - BedrockError(status_code=503, message='serviceUnavailableException {"message": "Service unavailable"}'), - BedrockError(status_code=500, message='internalServerException {"message": "Internal error"}'), - BedrockError(status_code=429, message='throttlingException {"message": "Too many requests"}'), - httpx.ReadError("connection reset by upstream"), - ], - ids=["503", "500", "429", "transport-drop"], -) -async def test_anthropic_messages_raised_provider_error_before_content_triggers_fallback(raised_error): - """A retriable raise before content falls over exactly like a detected SSE error event.""" - router = _anthropic_messages_make_router() - source = _AnthropicMessagesRaisingByteStream([_anthropic_messages_message_start_chunk()], raised_error) - fallback_stream = _AnthropicMessagesFallbackByteStream([_anthropic_messages_content_chunk("fallback answer")]) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(return_value=fallback_stream), - ) as mock_fallback: - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, - initial_kwargs={"model": "primary"}, - ) - collected = [chunk async for chunk in wrapped] - - assert collected == [_anthropic_messages_content_chunk("fallback answer")] - mock_fallback.assert_awaited_once() - converted = mock_fallback.await_args.kwargs["e"] - assert isinstance(converted, MidStreamFallbackError) - assert converted.original_exception is raised_error - assert converted.is_pre_first_chunk is True - assert source.closed is True - - -class _AnthropicMessagesStringStatusError(Exception): - def __init__(self): - super().__init__("bad request") - self.status_code = "400" - - -class _AnthropicMessagesResponseOnlyStatusError(Exception): - def __init__(self): - super().__init__("bad request") - self.response = SimpleNamespace(status_code=400) - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - "raised_error", - [ - BedrockError(status_code=400, message='validationException {"message": "Malformed input"}'), - BedrockError(status_code=424, message='modelStreamErrorException {"message": "Model stream error"}'), - _AnthropicMessagesStringStatusError(), - _AnthropicMessagesResponseOnlyStatusError(), - ], - ids=["400", "424", "str-400", "response-only-400"], -) -async def test_anthropic_messages_raised_non_retriable_provider_error_propagates_unchanged(raised_error): - """A raised client error reaches the caller as the same exception, nothing flushed, no fallback.""" - router = _anthropic_messages_make_router() - source = _AnthropicMessagesRaisingByteStream([_anthropic_messages_message_start_chunk()], raised_error) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(), - ) as mock_fallback: - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, - initial_kwargs={"model": "primary"}, - ) - collected = [] - - async def _consume(): - async for chunk in wrapped: - collected.append(chunk) - - with pytest.raises(type(raised_error)) as exc_info: - await _consume() - - assert collected == [] - assert exc_info.value is raised_error - mock_fallback.assert_not_awaited() - - -@pytest.mark.asyncio -async def test_anthropic_messages_raised_provider_error_after_content_propagates_unchanged(): - """A raise after content propagates unchanged even when its status is retriable.""" - router = _anthropic_messages_make_router() - content = _anthropic_messages_content_chunk("partial answer") - raised_error = BedrockError( - status_code=503, message='serviceUnavailableException {"message": "Service unavailable"}' - ) - source = _AnthropicMessagesRaisingByteStream([content], raised_error) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(), - ) as mock_fallback: - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, - initial_kwargs={"model": "primary"}, - ) - collected = [] - - async def _consume(): - async for chunk in wrapped: - collected.append(chunk) - - with pytest.raises(BedrockError) as exc_info: - await _consume() - - assert collected == [content] - assert exc_info.value is raised_error - mock_fallback.assert_not_awaited() - - -@pytest.mark.parametrize( - "error, expected_status", - [ - (BedrockError(status_code=503, message="unavailable"), 503), - (_AnthropicMessagesStringStatusError(), 400), - (_AnthropicMessagesResponseOnlyStatusError(), 400), - (httpx.ReadError("connection reset by upstream"), None), - ], - ids=["int", "digit-str", "response-only", "none"], -) -def test_anthropic_stream_raised_error_status_reads_every_status_shape(error, expected_status): - assert _anthropic_stream_raised_error_status(error) == expected_status - - -@pytest.mark.parametrize( - "error, has_generated_content, converts", - [ - (BedrockError(status_code=503, message="unavailable"), False, True), - (httpx.ReadError("connection reset by upstream"), False, True), - (BedrockError(status_code=400, message="malformed"), False, False), - (BedrockError(status_code=503, message="unavailable"), True, False), - ], - ids=["retriable", "no-status", "client-error", "after-content"], -) -def test_anthropic_stream_fallback_error_for_raised_gates_like_a_detected_error_event( - error, has_generated_content, converts -): - converted = _anthropic_stream_fallback_error_for_raised(error, "primary", has_generated_content) - if not converts: - assert converted is None - return - assert isinstance(converted, MidStreamFallbackError) - assert converted.original_exception is error - assert converted.is_pre_first_chunk is True - assert converted.llm_provider == "anthropic" - - -@pytest.mark.asyncio -async def test_aanthropic_messages_recover_stream_error_flushes_buffered_frames_before_declining(): - router = _anthropic_messages_make_router() - original = BedrockError(status_code=503, message="unavailable") - declined = MidStreamFallbackError( - message="unavailable", - model="primary", - llm_provider="anthropic", - original_exception=original, - is_pre_first_chunk=False, - ) - buffered = (_anthropic_messages_message_start_chunk(),) - flushed = [] - - async def drain(recovery) -> None: - async for chunk in recovery: - flushed.append(chunk) - - with patch.object(router, "_aanthropic_messages_fallback_attempt") as mock_attempt: - recovery = router._aanthropic_messages_recover_stream_error( - declined, True, buffered, "primary", {"model": "primary"}, _anthropic_messages_make_wrapper() - ) - with pytest.raises(BedrockError) as exc_info: - await drain(recovery) - assert flushed == list(buffered) - assert exc_info.value is original - mock_attempt.assert_not_called() - - -@pytest.mark.asyncio -async def test_aanthropic_messages_recover_stream_error_hands_converted_raise_to_fallback_attempt(): - router = _anthropic_messages_make_router() - raised = BedrockError(status_code=503, message="unavailable") - handed_over = [] - - async def fake_attempt(fallback_error, initial_kwargs, wrapper): - handed_over.append(fallback_error) - yield b"fallback" - - with patch.object(router, "_aanthropic_messages_fallback_attempt", new=fake_attempt): - recovery = router._aanthropic_messages_recover_stream_error( - raised, False, (), "primary", {"model": "primary"}, _anthropic_messages_make_wrapper() - ) - collected = [chunk async for chunk in recovery] - assert collected == [b"fallback"] - assert len(handed_over) == 1 - assert isinstance(handed_over[0], MidStreamFallbackError) - assert handed_over[0].original_exception is raised - - -@pytest.mark.asyncio -async def test_anthropic_messages_non_retriable_client_error_skips_fallback(): - """A 4xx (non-429) error type (e.g. invalid_request_error) is a client - error a fallback attempt cannot fix, so it must be forwarded to the - client as-is rather than burning a fallback attempt.""" - router = _anthropic_messages_make_router() - error_chunk = _anthropic_messages_invalid_request_error_chunk() - source = _AnthropicMessagesFakeByteStream([error_chunk]) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(), - ) as mock_fallback: - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, - initial_kwargs={"model": "primary"}, - ) - collected = [chunk async for chunk in wrapped] - - assert collected == [error_chunk] - mock_fallback.assert_not_awaited() - - -@pytest.mark.asyncio -async def test_anthropic_messages_post_first_chunk_error_skips_fallback(): - """Once content has already reached the caller, retrying would start a - second, overlapping Anthropic message lifecycle on the same SSE stream - - the error must be forwarded instead of triggering an invisible retry.""" - router = _anthropic_messages_make_router() - content = _anthropic_messages_content_chunk("partial answer") - error_chunk = _anthropic_messages_overloaded_error_chunk() - source = _AnthropicMessagesFakeByteStream([content, error_chunk]) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(), - ) as mock_fallback: - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, - initial_kwargs={"model": "primary"}, - ) - collected = [chunk async for chunk in wrapped] - - assert collected == [content, error_chunk] - mock_fallback.assert_not_awaited() - - -@pytest.mark.asyncio -async def test_anthropic_messages_non_retriable_error_flushes_buffered_lifecycle_frames(): - """A non-retriable error arriving while lifecycle frames are still - buffered (no content seen yet) must flush those buffered frames before - forwarding the error, so the client still sees the whole primary - attempt rather than losing the buffered message_start silently.""" - router = _anthropic_messages_make_router() - error_chunk = _anthropic_messages_invalid_request_error_chunk() - source = _AnthropicMessagesFakeByteStream([_anthropic_messages_message_start_chunk(), error_chunk]) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(), - ) as mock_fallback: - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, - initial_kwargs={"model": "primary"}, - ) - collected = [chunk async for chunk in wrapped] - - assert collected == [_anthropic_messages_message_start_chunk(), error_chunk] - mock_fallback.assert_not_awaited() - - -@pytest.mark.asyncio -async def test_anthropic_messages_raised_error_declined_flushes_buffered_lifecycle_frames(): - """When a raised MidStreamFallbackError is declined (source says content - was not pre-first-chunk) while lifecycle frames are still buffered, they - must be flushed to the client before the exception propagates.""" - router = _anthropic_messages_make_router() - raised_error = MidStreamFallbackError( - message="stream reset", - model="primary", - llm_provider="vertex_ai", - is_pre_first_chunk=False, - ) - source = _AnthropicMessagesRaisingByteStream([_anthropic_messages_message_start_chunk()], raised_error) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(), - ) as mock_fallback: - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, - initial_kwargs={"model": "primary"}, - ) - collected = [] - - async def _consume(): - async for chunk in wrapped: - collected.append(chunk) - - with pytest.raises(MidStreamFallbackError) as exc_info: - await _consume() - - assert collected == [_anthropic_messages_message_start_chunk()] - assert exc_info.value is raised_error - mock_fallback.assert_not_awaited() - - -@pytest.mark.asyncio -async def test_anthropic_messages_raised_error_without_original_exception_reraises_itself(): - """When a declined MidStreamFallbackError carries no original_exception, - the bare exception itself must propagate rather than being swallowed.""" - router = _anthropic_messages_make_router() - content = _anthropic_messages_content_chunk("partial answer") - raised_error = MidStreamFallbackError( - message="stream reset", - model="primary", - llm_provider="vertex_ai", - is_pre_first_chunk=False, - ) - source = _AnthropicMessagesRaisingByteStream([content], raised_error) - - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, - initial_kwargs={"model": "primary"}, - ) - collected = [] - - async def _consume(): - async for chunk in wrapped: - collected.append(chunk) - - with pytest.raises(MidStreamFallbackError) as exc_info: - await _consume() - - assert collected == [content] - assert exc_info.value is raised_error - - -@pytest.mark.asyncio -async def test_anthropic_messages_fallback_also_failing_raises_original_exception(): - """If the fallback attempt itself fails with a MidStreamFallbackError - wrapping a real provider exception, the client must see that real - exception, not the internal MidStreamFallbackError.""" - router = _anthropic_messages_make_router() - source = _AnthropicMessagesFakeByteStream([_anthropic_messages_overloaded_error_chunk()]) - original_exception = litellm.APIError( - status_code=503, - message="fallback also overloaded", - llm_provider="bedrock", - model="fallback", - ) - fallback_failure = MidStreamFallbackError( - message="fallback failed", - model="fallback", - llm_provider="bedrock", - original_exception=original_exception, - ) - - with patch.object( - router, - "async_function_with_fallbacks_common_utils", - new=AsyncMock(side_effect=fallback_failure), - ): - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, - initial_kwargs={"model": "primary"}, - ) - with pytest.raises(litellm.APIError) as exc_info: - async for _ in wrapped: - pass - - assert exc_info.value is original_exception - - -# -------- _dispatch_generic_call_type -------- - - -@pytest.mark.asyncio -async def test_dispatch_generic_call_type_routes_anthropic_messages_through_streaming_fallbacks(): - router = _anthropic_messages_make_router() - - async def fake_original(**_kwargs): - return {"id": "msg_1"} - - with patch.object( - router, - "_aanthropic_messages_with_streaming_fallbacks", - new=AsyncMock(return_value="anthropic-result"), - ) as mock_anthropic: - out = await router._dispatch_generic_call_type( - call_type="anthropic_messages", - original_function=fake_original, - model="primary", - ) - assert out == "anthropic-result" - mock_anthropic.assert_awaited_once() - - -@pytest.mark.asyncio -async def test_dispatch_generic_call_type_other_call_types_use_generic_fallback(): - router = _anthropic_messages_make_router() - - async def fake_original(**_kwargs): - return {"id": "file_1"} - - with patch.object( - router, - "_ageneric_api_call_with_fallbacks", - new=AsyncMock(return_value="generic-result"), - ) as mock_generic: - out = await router._dispatch_generic_call_type( - call_type="afile_delete", - original_function=fake_original, - model="primary", - ) - assert out == "generic-result" - mock_generic.assert_awaited_once() - - -@pytest.mark.asyncio -async def test_factory_function_anthropic_messages_uses_streaming_fallback_dispatch(): - """anthropic_messages must be wired through the mid-stream-fallback-aware - path rather than the bare generic dispatch every other call type without - special handling uses.""" - router = _anthropic_messages_make_router() - wrapped = router.factory_function(litellm.anthropic_messages, call_type="anthropic_messages") - assert callable(wrapped) - - with patch.object( - router, - "_aanthropic_messages_with_streaming_fallbacks", - new=AsyncMock(return_value="ok"), - ) as mock_anthropic: - result = await wrapped(model="primary") - assert result == "ok" - mock_anthropic.assert_awaited_once() - - -@pytest.mark.asyncio -async def test_async_function_with_fallbacks_stamps_zero_attempted_fallbacks(): - """A request served by the primary model group records attempted_fallbacks=0 and - the requested model group in metadata, mirroring the x-litellm-attempted-fallbacks header.""" - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo", "mock_response": "hi"}, - } - ] - ) - metadata = {} - - await router.acompletion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hey"}], - metadata=metadata, - ) - - assert metadata["attempted_fallbacks"] == 0 - assert metadata["original_model_group"] == "gpt-3.5-turbo" - - -@pytest.mark.asyncio -async def test_async_function_with_fallbacks_stamps_route_bucket_not_litellm_metadata(): - """A chat completion carrying both metadata buckets gets stamped in the route's bucket - (metadata), matching where run_async_fallback rewrites, so the two never diverge.""" - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo", "mock_response": "hi"}, - } - ] - ) - metadata = {} - litellm_metadata = {"client_key": "client_value"} - - await router.acompletion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hey"}], - metadata=metadata, - litellm_metadata=litellm_metadata, - ) - - assert metadata["attempted_fallbacks"] == 0 - assert metadata["original_model_group"] == "gpt-3.5-turbo" - assert litellm_metadata["client_key"] == "client_value" - assert "attempted_fallbacks" not in litellm_metadata - assert "original_model_group" not in litellm_metadata - - -@pytest.mark.asyncio -async def test_async_function_with_fallbacks_overrides_client_supplied_stamp_values(): - """Client-supplied attempted_fallbacks and original_model_group are replaced on entry, - so a reused metadata dict or a spoofed value cannot leak stale attribution into logs.""" - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo", "mock_response": "hi"}, - } - ] - ) - metadata = {"attempted_fallbacks": 99, "original_model_group": "stale-group"} - - await router.acompletion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hey"}], - metadata=metadata, - ) - - assert metadata["attempted_fallbacks"] == 0 - assert metadata["original_model_group"] == "gpt-3.5-turbo" - - -@pytest.mark.asyncio -async def test_async_function_with_fallbacks_stamps_despite_forged_reentry_params(): - """A client injecting fallback_depth or a JSON-shaped attempted_targets via request - litellm params cannot skip the entry stamp; only the router's own in-process - AttemptedFallbackTargets instance marks a genuine re-entrant hop.""" - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo", "mock_response": "hi"}, - } - ] - ) - metadata = {"attempted_fallbacks": 99, "original_model_group": "spoofed-group"} - - await router.acompletion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hey"}], - metadata=metadata, - fallback_depth=3, - attempted_targets={"keys": ["spoofed-group"]}, - ) - - assert metadata["attempted_fallbacks"] == 0 - assert metadata["original_model_group"] == "gpt-3.5-turbo" - - -@pytest.mark.asyncio -async def test_async_function_with_fallbacks_skips_stamp_on_genuine_reentrant_hop(): - """A re-entrant hop carrying the router's own AttemptedFallbackTargets instance keeps - the per-hop metadata that run_async_fallback wrote instead of resetting it to zero.""" - from litellm.router_utils.fallback_event_handlers import AttemptedFallbackTargets - - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo", "mock_response": "hi"}, - } - ] - ) - metadata = {"attempted_fallbacks": 1, "original_model_group": "prod-chat"} - - await router.acompletion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hey"}], - metadata=metadata, - attempted_targets=AttemptedFallbackTargets(keys=frozenset(("prod-chat",))), - ) - - assert metadata["attempted_fallbacks"] == 1 - assert metadata["original_model_group"] == "prod-chat" - - -def _record_router_acompletion_kwargs(router: litellm.Router) -> list: - """Spy on router._acompletion, recording each call's kwargs while delegating through.""" - records = [] - original_acompletion = router._acompletion - - @functools.wraps(original_acompletion) - async def _spy(*args, **spy_kwargs): - records.append(spy_kwargs) - return await original_acompletion(*args, **spy_kwargs) - - router._acompletion = _spy - return records - - -@pytest.mark.asyncio -async def test_async_function_with_fallbacks_scrubs_spoofed_values_from_sibling_bucket(): - """Spend logs read a truthy litellm_metadata dict in preference to metadata, so spoofed - stamp keys planted in the bucket the route does not own are removed on entry, in place, - before they can flow into the spend log row.""" - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo", "mock_response": "hi"}, - } - ] - ) - metadata = {} - litellm_metadata = { - "attempted_fallbacks": 99, - "original_model_group": "spoofed-group", - "client_key": "client_value", - } - downstream_calls = _record_router_acompletion_kwargs(router) - - await router.acompletion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hey"}], - metadata=metadata, - litellm_metadata=litellm_metadata, - ) - - assert len(downstream_calls) == 1 - downstream_sibling = downstream_calls[0]["litellm_metadata"] - assert "attempted_fallbacks" not in downstream_sibling - assert "original_model_group" not in downstream_sibling - assert downstream_sibling["client_key"] == "client_value" - assert "attempted_fallbacks" not in litellm_metadata - assert "original_model_group" not in litellm_metadata - assert litellm_metadata["client_key"] == "client_value" - assert metadata["attempted_fallbacks"] == 0 - assert metadata["original_model_group"] == "gpt-3.5-turbo" - - -@pytest.mark.asyncio -async def test_async_function_with_fallbacks_scrubs_sibling_bucket_in_place(): - """Everything below the router resolves the bucket by key presence, so the scrub edits - the caller's dict object like every other router bucket write. Rebinding kwargs to a - scrubbed copy detaches the proxy's request_data write-backs (guardrail telemetry, retry - accounting) from the object the spend row is built from.""" - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo", "mock_response": "hi"}, - } - ] - ) - litellm_metadata = { - "attempted_fallbacks": 7, - "original_model_group": "planted-group", - "client_key": "client_value", - } - caller_snapshot = copy.deepcopy(litellm_metadata) - downstream_calls = _record_router_acompletion_kwargs(router) - - await router.acompletion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hey"}], - metadata={}, - litellm_metadata=litellm_metadata, - ) - - assert len(downstream_calls) == 1 - assert downstream_calls[0]["litellm_metadata"] is litellm_metadata - assert "attempted_fallbacks" not in litellm_metadata - assert "original_model_group" not in litellm_metadata - assert litellm_metadata["client_key"] == caller_snapshot["client_key"] - - -@pytest.mark.asyncio -async def test_async_function_with_fallbacks_stamps_aliased_buckets_on_every_call(): - """One dict object passed as both metadata and litellm_metadata: the first call's own - stamp puts the reserved keys into the shared object, so the second call enters the - scrub with them present. Scrubbing in place keeps the stamp and the bucket on the same - object; a scrubbed copy would leave the spend reader's preferred bucket unstamped.""" - router = litellm.Router( - model_list=[ - { - "model_name": "chat-group", - "litellm_params": {"model": "gpt-3.5-turbo", "mock_response": "hi"}, - } - ] - ) - shared_metadata = {"team": "alpha"} - downstream_calls = _record_router_acompletion_kwargs(router) - - for _ in range(3): - await router.acompletion( - model="chat-group", - messages=[{"role": "user", "content": "hey"}], - metadata=shared_metadata, - litellm_metadata=shared_metadata, - ) - - assert len(downstream_calls) == 3 - for call_kwargs in downstream_calls: - assert call_kwargs["litellm_metadata"] is shared_metadata - assert call_kwargs["metadata"] is shared_metadata - assert call_kwargs["litellm_metadata"]["attempted_fallbacks"] == 0 - assert call_kwargs["litellm_metadata"]["original_model_group"] == "chat-group" - - -@pytest.mark.asyncio -async def test_async_function_with_fallbacks_passes_clean_sibling_bucket_through_unchanged(): - """A sibling bucket carrying no reserved stamp keys is forwarded downstream as the - caller's own object with no copy made, matching pre-scrub behavior. Retry accounting - stamped into that bucket downstream predates the scrub and is out of its scope.""" - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo", "mock_response": "hi"}, - } - ] - ) - litellm_metadata = {"client_key": "client_value"} - downstream_calls = _record_router_acompletion_kwargs(router) - - await router.acompletion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hey"}], - metadata={}, - litellm_metadata=litellm_metadata, - ) - - assert len(downstream_calls) == 1 - assert downstream_calls[0]["litellm_metadata"] is litellm_metadata - assert litellm_metadata["client_key"] == "client_value" - assert "attempted_fallbacks" not in litellm_metadata - assert "original_model_group" not in litellm_metadata - - -@pytest.mark.asyncio -async def test_run_async_fallback_keeps_caller_metadata_keys_on_the_wire(monkeypatch): - """Under enable_preview_features, add_openai_metadata forwards only the first 16 - string pairs of request metadata to the provider body, so the fallback hop must - spread caller keys before the router's own stamps: a stamp inserted first evicts - the caller's 16th key from the wire while the internal stamp rides in its place.""" - monkeypatch.setattr(litellm, "enable_preview_features", True) - caller_metadata = {f"user_key_{i}": f"value_{i}" for i in range(16)} - router = litellm.Router( - model_list=[ - { - "model_name": "primary-group", - "litellm_params": {"model": "gpt-3.5-turbo", "api_key": "sk-test"}, - }, - { - "model_name": "fallback-group", - "litellm_params": {"model": "gpt-3.5-turbo", "api_key": "sk-test"}, - }, - ], - fallbacks=[{"primary-group": ["fallback-group"]}], - num_retries=0, - ) - - wire_bodies = [] - - def _respond(request: httpx.Request) -> httpx.Response: - wire_bodies.append(json.loads(request.content)) - return httpx.Response( - 200, - json={ - "id": "chatcmpl-wire", - "object": "chat.completion", - "created": 1, - "model": "gpt-3.5-turbo", - "choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}], - "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, - }, - ) - - client = openai.AsyncOpenAI( - api_key="sk-test", - http_client=httpx.AsyncClient(transport=httpx.MockTransport(_respond)), - ) - - await router.acompletion( - model="primary-group", - messages=[{"role": "user", "content": "hey"}], - metadata=dict(caller_metadata), - mock_testing_fallbacks=True, - client=client, - ) - - assert len(wire_bodies) == 1 - assert wire_bodies[0]["metadata"] == caller_metadata - - wire_bodies.clear() - small_metadata = {"team": "alpha", "env": "prod"} - await router.acompletion( - model="primary-group", - messages=[{"role": "user", "content": "hey again"}], - metadata=dict(small_metadata), - mock_testing_fallbacks=True, - client=client, - ) - - assert len(wire_bodies) == 1 - small_wire = wire_bodies[0]["metadata"] - assert {k: small_wire[k] for k in small_metadata} == small_metadata - assert small_wire["original_model_group"] == "primary-group" - assert small_wire["model_group"] == "fallback-group" - - -@pytest.mark.asyncio -async def test_run_async_fallback_two_hop_chain_reports_entry_group_and_hop_count(): - """A two-hop fallback chain stamps attempted_fallbacks=2 on the final leg and keeps - original_model_group at the group requested on entry: a later hop's stamp appends - after caller keys without overriding the value stamped by an earlier hop.""" - router = litellm.Router( - model_list=[ - { - "model_name": "group-a", - "litellm_params": {"model": "gpt-3.5-turbo", "mock_response": "litellm.InternalServerError"}, - }, - { - "model_name": "group-b", - "litellm_params": {"model": "gpt-3.5-turbo", "mock_response": "litellm.InternalServerError"}, - }, - { - "model_name": "group-c", - "litellm_params": {"model": "gpt-3.5-turbo", "mock_response": "ok"}, - }, - ], - fallbacks=[{"group-a": ["group-b"]}, {"group-b": ["group-c"]}], - num_retries=0, - ) - metadata = {} - leg_records = [] - original_acompletion = router._acompletion - - @functools.wraps(original_acompletion) - async def _spy(*args, **spy_kwargs): - leg_records.append((spy_kwargs.get("model"), copy.deepcopy(spy_kwargs.get("metadata")))) - return await original_acompletion(*args, **spy_kwargs) - - router._acompletion = _spy - - await router.acompletion( - model="group-a", - messages=[{"role": "user", "content": "hey"}], - metadata=metadata, - ) - - assert [model for model, _ in leg_records] == ["group-a", "group-b", "group-c"] - hop_one_metadata = leg_records[1][1] - assert hop_one_metadata["attempted_fallbacks"] == 1 - assert hop_one_metadata["original_model_group"] == "group-a" - assert hop_one_metadata["model_group"] == "group-b" - hop_two_metadata = leg_records[2][1] - assert hop_two_metadata["attempted_fallbacks"] == 2 - assert hop_two_metadata["original_model_group"] == "group-a" - assert hop_two_metadata["model_group"] == "group-c" - assert metadata["attempted_fallbacks"] == 0 - assert metadata["original_model_group"] == "group-a" - - -def _permission_denied_error() -> litellm.PermissionDeniedError: - return litellm.PermissionDeniedError( - message="OpenrouterException - this key has no access to the model", - llm_provider="openrouter", - model="openrouter/openai/gpt-4o", - response=httpx.Response(status_code=403, request=httpx.Request(method="POST", url="https://openrouter.ai")), - ) - - -def test_permission_denied_error_is_not_retried_against_a_single_deployment(): - router = litellm.Router( - model_list=[ - {"model_name": "gpt-4o", "litellm_params": {"model": "openrouter/openai/gpt-4o", "api_key": "sk-test"}}, - ] - ) - - with pytest.raises(litellm.PermissionDeniedError): - router.should_retry_this_error( - error=_permission_denied_error(), - healthy_deployments=router.model_list, - all_deployments=router.model_list, - ) - - -def test_permission_denied_error_is_retried_when_other_deployments_exist(): - router = litellm.Router( - model_list=[ - {"model_name": "gpt-4o", "litellm_params": {"model": "openrouter/openai/gpt-4o", "api_key": "sk-test"}}, - {"model_name": "gpt-4o", "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"}}, - ] - ) - - assert ( - router.should_retry_this_error( - error=_permission_denied_error(), - healthy_deployments=router.model_list, - all_deployments=router.model_list, - ) - is True - ) - - -class _AllowlistFallbackAccessCheck: - def __init__(self, allowed_models: frozenset[str]): - self.allowed_models = allowed_models - self.checked_models = [] - - async def __call__(self, *, model, request_kwargs, llm_router): - self.checked_models.append(model) - return model in self.allowed_models - - -def _router_with_failing_primary(fallback_access_check) -> Router: - return Router( - model_list=[ - { - "model_name": "primary", - "litellm_params": { - "model": "openai/primary", - "api_key": "k", - "mock_response": Exception("primary is down"), - }, - }, - { - "model_name": "secret-fallback", - "litellm_params": { - "model": "openai/secret", - "api_key": "k", - "mock_response": "served by secret-fallback", - }, - }, - ], - fallbacks=[{"primary": ["secret-fallback"]}], - num_retries=0, - fallback_access_check=fallback_access_check, - ) - - -@pytest.mark.asyncio -async def test_fallback_access_check_blocks_config_fallback_the_caller_cannot_use(): - access_check = _AllowlistFallbackAccessCheck(allowed_models=frozenset()) - router = _router_with_failing_primary(access_check) - - with pytest.raises(Exception, match="primary is down"): - await router.acompletion(model="primary", messages=[{"role": "user", "content": "hi"}]) - - assert access_check.checked_models == ["secret-fallback"] - - -@pytest.mark.asyncio -async def test_fallback_access_check_lets_an_authorized_config_fallback_through(): - router = _router_with_failing_primary(_AllowlistFallbackAccessCheck(allowed_models=frozenset({"secret-fallback"}))) - - response = await router.acompletion(model="primary", messages=[{"role": "user", "content": "hi"}]) - - assert response.choices[0].message.content == "served by secret-fallback" - - -@pytest.mark.asyncio -async def test_router_without_fallback_access_check_attempts_every_config_fallback(): - router = _router_with_failing_primary(None) - - response = await router.acompletion(model="primary", messages=[{"role": "user", "content": "hi"}]) - - assert response.choices[0].message.content == "served by secret-fallback" - - -def _resolution_router() -> Router: - return Router( - model_list=[ - {"model_name": "pinned", "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"}}, - {"model_name": "pooled", "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-test"}}, - {"model_name": "pooled", "litellm_params": {"model": "anthropic/claude-haiku-4-5", "api_key": "sk-test"}}, - {"model_name": "bedrock/*", "litellm_params": {"model": "bedrock/*", "api_key": "sk-test"}}, - ], - model_group_alias={"nickname": "pinned"}, - ) - - -@pytest.mark.parametrize( - "model_name,expected", - [ - ("pinned", ("openai/gpt-4o",)), - ("nickname", ("openai/gpt-4o",)), - ("pooled", ("openai/gpt-4o-mini", "anthropic/claude-haiku-4-5")), - ("bedrock/anthropic.claude-3-5-sonnet", ("bedrock/anthropic.claude-3-5-sonnet",)), - ("never-configured", ()), - ], - ids=["exact-name", "model-group-alias", "every-member-of-a-pool", "wildcard-expands", "resolves-to-nothing"], -) -def test_resolved_litellm_models_answers_through_every_channel_a_request_uses( - model_name: str, expected: tuple[str, ...] -) -> None: - """A caller comparing two names by what serves them needs each channel the request path - composes, since the deployment name an admin picked carries no information on its own. - - `resolves-to-nothing` is the contract that keeps the fallback out of here: an empty - result is not "the call fails", so what to do about it stays each caller's policy. - """ - assert set(_resolution_router().resolved_litellm_models(model_name)) == set(expected) - - -class TestTierParamsTheTargetAccepts: - """A tier's litellm_params are applied to every request that tier routes, so one the target - cannot take raised UnsupportedParamsError before the request left the proxy, turning the whole - tier into a 400.""" - - @pytest.fixture(autouse=True) - def force_local_model_cost(self, monkeypatch): - from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap - - monkeypatch.setattr(litellm, "model_cost", GetModelCostMap.load_local_model_cost_map()) - - @staticmethod - def _router(model: str) -> litellm.Router: - return litellm.Router( - model_list=[{"model_name": "tiered", "litellm_params": {"model": model, "api_key": "sk-x"}}] - ) - - def test_drops_a_param_no_deployment_declares(self): - router = self._router("novita/moonshotai/kimi-k3") - - accepted = router._tier_params_the_target_accepts("tiered", {"reasoning_effort": "max"}, {}) - - assert accepted == {} - - def test_keeps_a_param_the_deployment_declares(self): - router = self._router("fireworks_ai/kimi-k3") - - accepted = router._tier_params_the_target_accepts("tiered", {"reasoning_effort": "max"}, {}) - - assert accepted == {"reasoning_effort": "max"} - - @pytest.mark.parametrize( - "control, value", - [ - ("api_base", "https://example.invalid"), - ("api_key", "sk-tier"), - ("base_url", "https://example.invalid"), - ("timeout", 30), - ("default_headers", {"x-tier": "1"}), - ("organization", "org-tier"), - ("deployment_id", "dep-tier"), - ], - ) - def test_keeps_credentials_and_transport_controls(self, control, value): - """These are not chat completion params, so get_optional_params never compares them against - a provider's supported list. Filtering on "is this an OpenAI param" would discard the - configuration the request needs while never touching what the provider would reject.""" - router = self._router("novita/moonshotai/kimi-k3") - - accepted = router._tier_params_the_target_accepts("tiered", {control: value, "reasoning_effort": "max"}, {}) - - assert accepted == {control: value} - - @pytest.mark.parametrize( - "control, value", - [ - ("additional_drop_params", ["seed"]), - ("drop_params", True), - ("allowed_openai_params", ["seed"]), - ("api_version", "2024-02-01"), - ("metadata", {"tier": "complex"}), - ], - ) - def test_keeps_litellm_controls_the_provider_never_lists(self, control, value): - """No provider lists a litellm control among its supported params, so "no deployment - declares it" means litellm consumes it, not that the target refuses it. Dropping - drop_params or additional_drop_params would silently disable the operator's sanitization.""" - router = self._router("novita/moonshotai/kimi-k3") - - accepted = router._tier_params_the_target_accepts("tiered", {control: value, "reasoning_effort": "max"}, {}) - - assert accepted == {control: value} - - def test_tier_allowlist_protects_the_param_it_names(self): - """allowed_openai_params is the documented escape hatch for an incomplete supported-params - list, and request-time validation extends the supported list with it, so a param the tier - both sets and allowlists would never 400 and must not be dropped.""" - router = self._router("novita/moonshotai/kimi-k3") - - accepted = router._tier_params_the_target_accepts( - "tiered", {"reasoning_effort": "max", "allowed_openai_params": ["reasoning_effort"]}, {} - ) - - assert accepted == {"reasoning_effort": "max", "allowed_openai_params": ["reasoning_effort"]} - - def test_request_allowlist_protects_the_param_it_names(self): - router = self._router("novita/moonshotai/kimi-k3") - - accepted = router._tier_params_the_target_accepts( - "tiered", {"reasoning_effort": "max"}, {"allowed_openai_params": ["reasoning_effort"]} - ) - - assert accepted == {"reasoning_effort": "max"} - - def test_allowlist_protects_only_the_params_it_names(self): - router = self._router("novita/moonshotai/kimi-k3") - - accepted = router._tier_params_the_target_accepts( - "tiered", {"reasoning_effort": "max", "allowed_openai_params": ["seed"]}, {} - ) - - assert accepted == {"allowed_openai_params": ["seed"]} - - def test_declared_param_allowlist_ignores_malformed_declarations(self): - """A str is iterable, so without the type guard a YAML scalar mistake like - allowed_openai_params: reasoning_effort would allowlist single characters.""" - assert litellm.Router._declared_param_allowlist({"allowed_openai_params": ["reasoning_effort", 3]}) == frozenset( - {"reasoning_effort"} - ) - assert litellm.Router._declared_param_allowlist({"allowed_openai_params": "reasoning_effort"}) == frozenset() - assert litellm.Router._declared_param_allowlist({}) == frozenset() - - def test_deployment_accepts_param_honors_deployment_allowlist(self): - deployment = { - "model_name": "x", - "litellm_params": {"model": "novita/moonshotai/kimi-k3", "allowed_openai_params": ["reasoning_effort"]}, - } - - assert litellm.Router._deployment_accepts_param(deployment, "x", "reasoning_effort") is True - - def test_keeps_a_token_ceiling_the_provider_spells_differently(self): - """petals lists max_tokens but not max_completion_tokens. A tier ceiling in the unsupported - spelling is a cost bound: dropping it would let a caller's larger max_tokens through where - today the mismatch fails loudly.""" - router = self._router("petals/petals-team/StableBeluga2") - - accepted = router._tier_params_the_target_accepts( - "tiered", {"max_completion_tokens": 100, "reasoning_effort": "max"}, {} - ) - - assert accepted == {"max_completion_tokens": 100} - - def test_keeps_extra_headers_even_when_the_provider_omits_it(self): - """Several providers leave extra_headers out of their supported params, so the filter would - drop it. Headers carry auth and tenancy, so sending fewer than the operator configured is - worse than the error they already get.""" - router = self._router("ai21/jamba-1.5-mini") - - accepted = router._tier_params_the_target_accepts( - "tiered", {"extra_headers": {"x-tenant": "acme"}, "reasoning_effort": "max"}, {} - ) - - assert accepted == {"extra_headers": {"x-tenant": "acme"}} - - def test_keeps_a_param_any_deployment_in_the_group_declares(self): - """Routing has not picked a deployment yet, so one capable member keeps the param alive.""" - router = litellm.Router( - model_list=[ - {"model_name": "tiered", "litellm_params": {"model": "novita/moonshotai/kimi-k3", "api_key": "k"}}, - {"model_name": "tiered", "litellm_params": {"model": "fireworks_ai/kimi-k3", "api_key": "k"}}, - ] - ) - - accepted = router._tier_params_the_target_accepts("tiered", {"reasoning_effort": "max"}, {}) - - assert accepted == {"reasoning_effort": "max"} - - def test_deployment_accepts_param_honors_base_model(self): - """An azure deployment named after the deployment rather than the model carries the real - model in base_model, and request-time mapping resolves capability through it, so the filter - has to ask the same question or it drops a param the deployment accepts.""" - by_model_info = { - "model_name": "x", - "litellm_params": {"model": "azure/my-gpt5-deploy"}, - "model_info": {"base_model": "azure/gpt-5"}, - } - by_litellm_params = { - "model_name": "x", - "litellm_params": {"model": "azure/my-gpt5-deploy", "base_model": "azure/gpt-5"}, - } - without_hint = {"model_name": "x", "litellm_params": {"model": "azure/my-gpt5-deploy"}} - - assert litellm.Router._deployment_accepts_param(by_model_info, "x", "reasoning_effort") is True - assert litellm.Router._deployment_accepts_param(by_litellm_params, "x", "reasoning_effort") is True - assert litellm.Router._deployment_accepts_param(without_hint, "x", "reasoning_effort") is False - - def test_deployment_accepts_param_reads_the_provider(self): - deployment = {"model_name": "x", "litellm_params": {"model": "fireworks_ai/kimi-k3"}} - - assert litellm.Router._deployment_accepts_param(deployment, "x", "reasoning_effort") is True - - def test_deployment_accepts_param_is_false_when_the_provider_omits_it(self): - deployment = {"model_name": "x", "litellm_params": {"model": "novita/moonshotai/kimi-k3"}} - - assert litellm.Router._deployment_accepts_param(deployment, "x", "reasoning_effort") is False - - @pytest.mark.parametrize( - "deployment", - [{"model_name": "x"}, {"model_name": "x", "litellm_params": {}}, {"model_name": "x", "litellm_params": {"model": "not-a-real-provider/nope"}}], - ) - def test_deployment_accepts_param_fails_open(self, deployment): - """An unresolvable deployment must not be the reason a param is dropped.""" - assert litellm.Router._deployment_accepts_param(deployment, "x", "reasoning_effort") is True - - @pytest.mark.parametrize( - "litellm_params", - [ - {"model": "github_copilot/gpt-4o"}, - {"model": "chatgpt/gpt-5"}, - {"model": "gpt-4o", "custom_llm_provider": "github_copilot"}, - ], - ) - def test_deployment_accepts_param_never_asks_a_provider_whose_lookup_authenticates( - self, litellm_params, monkeypatch - ): - """Resolving github_copilot or chatgpt runs their OAuth device flow, so a capability - question asked from the routing path can freeze the event loop for minutes waiting on a - human. The deployment counts as accepting everything, and the lookup is never made: an - exception-based sentinel cannot prove that, because the filter swallows exceptions into - the same keep answer.""" - lookups: list = [] - - def _record(*args, **kwargs): - lookups.append((args, kwargs)) - raise RuntimeError("provider resolution must not run for an authenticating provider") - - monkeypatch.setattr(litellm, "get_llm_provider", _record) - deployment = {"model_name": "x", "litellm_params": litellm_params} - - assert litellm.Router._deployment_accepts_param(deployment, "x", "reasoning_effort") is True - assert lookups == [] - - def test_keeps_everything_for_an_unknown_group(self): - """An unresolvable target must never narrow what the request already did.""" - router = self._router("fireworks_ai/kimi-k3") - - accepted = router._tier_params_the_target_accepts("no-such-group", {"reasoning_effort": "max"}, {}) - - assert accepted == {"reasoning_effort": "max"} - - -class TestRequestReasoningEffortOverride: - def test_drop_effort_from_nested_carrier_preserves_other_nested_values(self): - params: dict[str, object] = {"output_config": {"effort": "high", "format": "json"}} - - litellm.Router._pop_effort_from_nested_carrier(params, "output_config") - - assert params == {"output_config": {"format": "json"}} - - @pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"]) - def test_is_classifier_internal_call_recognizes_both_metadata_carriers(self, metadata_key): - kwargs = {metadata_key: {"internal_call_origin": "autorouter_classifier"}} - - assert litellm.Router._is_classifier_internal_call(kwargs) is True - assert litellm.Router._is_classifier_internal_call({metadata_key: {}}) is False - - def test_removes_every_deployment_native_effort_carrier_without_mutating_shared_config(self): - extra_body: dict[str, object] = { - "reasoning_effort": "high", - "thinking": {"type": "enabled"}, - "output_config": {"effort": "high", "format": "json"}, - "reasoning": {"effort": "high", "summary": "detailed"}, - "provider_option": True, - } - deployment_params: dict[str, object] = { - "model": "bedrock/converse/anthropic.claude-3-7-sonnet", - "thinking": {"type": "enabled", "budget_tokens": 2048}, - "output_config": {"effort": "high", "format": {"type": "json_schema"}}, - "reasoning": {"effort": "high", "summary": "auto"}, - "extra_body": extra_body, - } - - sanitized = litellm.Router._deployment_params_with_request_reasoning_override( - deployment_params, {"reasoning_effort": "low"} - ) - - assert sanitized == { - "model": "bedrock/converse/anthropic.claude-3-7-sonnet", - "output_config": {"format": {"type": "json_schema"}}, - "reasoning": {"summary": "auto"}, - "extra_body": { - "output_config": {"format": "json"}, - "reasoning": {"summary": "detailed"}, - "provider_option": True, - }, - } - assert deployment_params["thinking"] == {"type": "enabled", "budget_tokens": 2048} - assert deployment_params["output_config"] == {"effort": "high", "format": {"type": "json_schema"}} - assert extra_body["reasoning_effort"] == "high" - - @pytest.mark.parametrize("request_kwargs", [{}, {"reasoning_effort": None}]) - def test_omitted_override_preserves_deployment_defaults(self, request_kwargs): - deployment_params = { - "model": "deepseek/deepseek-reasoner", - "thinking": {"type": "enabled"}, - "output_config": {"effort": "high"}, - } - - assert ( - litellm.Router._deployment_params_with_request_reasoning_override(deployment_params, request_kwargs) - == deployment_params - ) - - @pytest.mark.asyncio - async def test_280_concurrent_overrides_never_mutate_or_leak_through_shared_deployment_params(self): - deployment_params = { - "model": "fireworks_ai/accounts/fireworks/models/kimi-k2-thinking", - "thinking": {"type": "enabled"}, - "output_config": {"effort": "high", "format": "json"}, - "extra_body": {"reasoning_effort": "high", "tenant": "shared"}, - } - efforts = ("none", "minimal", "low", "medium", "high", "xhigh", "max") - - results = await asyncio.gather( - *( - asyncio.to_thread( - litellm.Router._deployment_params_with_request_reasoning_override, - deployment_params, - {"reasoning_effort": efforts[index % len(efforts)]}, - ) - for index in range(280) - ) - ) - - assert all("thinking" not in result for result in results) - assert all(result["output_config"] == {"format": "json"} for result in results) - assert all(result["extra_body"] == {"tenant": "shared"} for result in results) - assert deployment_params["thinking"] == {"type": "enabled"} - assert deployment_params["output_config"] == {"effort": "high", "format": "json"} - assert deployment_params["extra_body"] == {"reasoning_effort": "high", "tenant": "shared"} - - @pytest.mark.parametrize( - ("metadata", "should_drop"), - [({"internal_call_origin": "autorouter_classifier"}, True), ({}, False)], - ids=["classifier", "ordinary-request"], - ) - def test_only_classifier_calls_drop_effort_for_an_unsupported_fallback(self, metadata, should_drop): - router = litellm.Router(model_list=[]) - body: dict[str, object] = {"model": "classifier", "reasoning_effort": "low"} - kwargs: dict[str, object] = { - "reasoning_effort": "low", - "metadata": metadata, - "proxy_server_request": {"body": body}, - } - deployment: DeploymentTypedDict = { - "model_name": "fallback", - "litellm_params": {"model": "openai/gpt-4o-mini"}, - } - - router._drop_unsupported_classifier_reasoning_effort(deployment, "fallback", kwargs) - - assert ("reasoning_effort" not in kwargs) is should_drop - assert ("reasoning_effort" not in body) is should_drop - - -class TestPreRoutingTierDrivesFallbacks: - """#38832: a complexity/auto router picks a tier behind the router name, but fallback - lookup stayed on the router name, so the tier's configured chain never ran and a - provider failure on the tier's first hop was returned to the client.""" - - class _TierRouter(litellm.Router): - async def async_pre_routing_hook( - self, model, request_kwargs, messages=None, input=None, specific_deployment=False - ): - from litellm.types.router import PreRoutingHookResponse - - if model == "smart-router": - return PreRoutingHookResponse(model="tier1", messages=messages) - return None - - @classmethod - def _router(cls, fallbacks) -> "litellm.Router": - return cls._TierRouter( - model_list=[ - { - "model_name": "smart-router", - "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-x"}, - }, - { - "model_name": "tier1", - "litellm_params": { - "model": "openai/gpt-4o-mini", - "api_key": "sk-x", - "mock_response": "litellm.RateLimitError", - }, - }, - { - "model_name": "backup-a", - "litellm_params": { - "model": "openai/gpt-4o-mini", - "api_key": "sk-x", - "mock_response": "from backup-a", - }, - }, - { - "model_name": "backup-b", - "litellm_params": { - "model": "openai/gpt-4o-mini", - "api_key": "sk-x", - "mock_response": "from backup-b", - }, - }, - { - "model_name": "failing-backup", - "litellm_params": { - "model": "openai/gpt-4o-mini", - "api_key": "sk-x", - "mock_response": "litellm.RateLimitError", - }, - }, - { - "model_name": "plain", - "litellm_params": { - "model": "openai/gpt-4o-mini", - "api_key": "sk-x", - "mock_response": "litellm.RateLimitError", - }, - }, - ], - fallbacks=fallbacks, - num_retries=0, - ) - - @pytest.mark.asyncio - async def test_the_selected_tier_fallback_chain_runs(self): - router = self._router([{"tier1": ["backup-a"]}]) - - response = await router.acompletion( - model="smart-router", messages=[{"role": "user", "content": "hi"}] - ) - - assert response.choices[0].message.content == "from backup-a" - - @pytest.mark.asyncio - async def test_a_chain_keyed_on_the_router_name_is_not_used(self): - """The router name has no chain of its own, so nothing should rescue this call.""" - router = self._router([{"tier2": ["backup-a"]}]) - - with pytest.raises(litellm.RateLimitError): - await router.acompletion(model="smart-router", messages=[{"role": "user", "content": "hi"}]) - - @pytest.mark.asyncio - async def test_a_chain_keyed_on_the_router_name_rescues_when_no_tier_chain_exists(self): - """The documented contract: configs keyed on the requested name keep working behind auto-routers.""" - router = self._router([{"smart-router": ["backup-a"]}]) - - response = await router.acompletion(model="smart-router", messages=[{"role": "user", "content": "hi"}]) - - assert response.choices[0].message.content == "from backup-a" - - @pytest.mark.asyncio - async def test_the_tier_chain_wins_over_the_router_name_chain(self): - router = self._router([{"tier1": ["backup-a"]}, {"smart-router": ["backup-b"]}]) - - response = await router.acompletion(model="smart-router", messages=[{"role": "user", "content": "hi"}]) - - assert response.choices[0].message.content == "from backup-a" - - @pytest.mark.asyncio - async def test_a_request_without_a_pre_routing_hook_still_uses_its_own_group(self): - router = self._router([{"tier1": ["backup-a"]}]) - - response = await router.acompletion( - model="tier1", messages=[{"role": "user", "content": "hi"}] - ) - - assert response.choices[0].message.content == "from backup-a" - - @pytest.mark.asyncio - async def test_a_caller_cannot_pick_the_chain_by_sending_the_selection(self): - """The metadata bucket carries caller-supplied keys, so only the hook may set the tier.""" - router = self._router([{"tier1": ["backup-a"]}]) - - with pytest.raises(litellm.RateLimitError): - await router.acompletion( - model="plain", - messages=[{"role": "user", "content": "hi"}], - metadata={"pre_routing_selected_model": "tier1"}, - ) - - @pytest.mark.asyncio - async def test_each_fallback_hop_resolves_its_own_chain(self): - """The second hop must key off the group it is running, not the tier that failed.""" - router = self._router([{"tier1": ["failing-backup"]}, {"failing-backup": ["backup-b"]}]) - - response = await router.acompletion(model="smart-router", messages=[{"role": "user", "content": "hi"}]) - - assert response.choices[0].message.content == "from backup-b" - - -@pytest.mark.asyncio -async def test_prompt_management_factory_marks_injection_for_every_deployment(monkeypatch): - """The factory stamps a provisional deployment's model_info into kwargs before the - prompt pass runs, then routes on the returned model, so any deployment can end up - billed. An injection recorded there must carry the every-deployment sentinel, never - the provisional deployment's id, or a differently-billed deployment loses the credit.""" - import time - - from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging - - router = litellm.Router( - model_list=[ - { - "model_name": "cached-claude", - "litellm_params": { - "model": "anthropic_cache_control_hook/claude-sonnet-5", - "prompt_id": "cache-points", - }, - "model_info": {"id": "provisional-dep"}, - } - ] - ) - captured: dict = {} - - async def _capture_acompletion(**kwargs): - captured.update(kwargs) - return litellm.ModelResponse() - - monkeypatch.setattr(litellm, "acompletion", _capture_acompletion) - logging_obj = LiteLLMLogging( - model="cached-claude", - messages=[{"role": "user", "content": "hi"}], - stream=False, - call_type="acompletion", - start_time=time.time(), - litellm_call_id="lit-6445", - function_id="f", - ) - await router.acompletion( - model="cached-claude", - messages=[ - {"role": "system", "content": "a static system prompt"}, - {"role": "user", "content": "hi"}, - ], - cache_control_injection_points=[{"location": "message", "role": "system"}], - litellm_logging_obj=logging_obj, - ) - bucket = captured.get("litellm_metadata") or captured["metadata"] - assert captured["model_info"]["id"] == "provisional-dep" - assert bucket["litellm_gateway_injected_cache"] == "" - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - "retry_policy,upstream_status,error_type,expected_upstream_calls", - [ - ({"ServiceUnavailableErrorRetries": 0}, 503, litellm.ServiceUnavailableError, 1), - ({"ServiceUnavailableErrorRetries": 1}, 503, litellm.ServiceUnavailableError, 2), - ({"InternalServerErrorRetries": 0}, 500, litellm.InternalServerError, 1), - ({"DefaultRetries": 0}, 502, litellm.BadGatewayError, 1), - ({"DefaultRetries": 0, "ServiceUnavailableErrorRetries": 1}, 503, litellm.ServiceUnavailableError, 2), - ({"ServiceUnavailableErrorRetries": 0}, 502, litellm.BadGatewayError, 3), - ], -) -async def test_router_retry_policy_controls_upstream_attempt_count( - monkeypatch: pytest.MonkeyPatch, retry_policy, upstream_status, error_type, expected_upstream_calls -): - monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-5.6", - "litellm_params": { - "model": "openai/gpt-5.6", - "api_key": "sk-fake", - "api_base": "https://retry-policy.local/v1", - }, - } - ], - num_retries=2, - retry_policy=retry_policy, - disable_cooldowns=True, - ) - - with respx.mock(assert_all_called=True) as respx_mock: - upstream = respx_mock.post("https://retry-policy.local/v1/chat/completions").mock( - return_value=httpx.Response( - upstream_status, - headers={"retry-after": "0"}, - json={"error": {"message": "model is down", "type": "server_error"}}, - ) - ) - with pytest.raises(error_type): - await router.acompletion(model="gpt-5.6", messages=[{"role": "user", "content": "hi"}]) - - assert upstream.call_count == expected_upstream_calls - - -def _make_failure_logging_obj(): - return LiteLLMLogging( - model="gpt-5.6", - messages=[{"role": "user", "content": "hi"}], - stream=False, - call_type="acompletion", - start_time=datetime.now(), - litellm_call_id="lit-6960", - function_id="f", - ) - - -async def _assert_router_failure_logging_is_coordinated(logging_obj, trigger, expected_exception): - """The sync failure_handler must not start until async_failure_handler has finished on the shared logging_obj.""" - events: list[str] = [] - sync_done = threading.Event() - - async def _async_failure(*args, **kwargs): - events.append("async_start") - await asyncio.sleep(0.05) - events.append("async_end") - - def _sync_failure(*args, **kwargs): - events.append("sync_start") - sync_done.set() - - with ( - patch.object(logging_obj, "async_failure_handler", side_effect=_async_failure), - patch.object(logging_obj, "failure_handler", side_effect=_sync_failure), - patch.object(logging_obj, "_should_run_sync_failure_callbacks_for_async_calls", return_value=True), - ): - with pytest.raises(expected_exception): - await trigger() - pending = [t for t in asyncio.all_tasks() if t is not asyncio.current_task()] - await asyncio.gather(*pending) - assert await asyncio.to_thread(sync_done.wait, 5), "failure_handler never ran" - - assert events == ["async_start", "async_end", "sync_start"] - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - "hook_error", - [ - litellm.RateLimitError(message="rpm exceeded", llm_provider="openai", model="gpt-5.6"), - RuntimeError("pre call check blew up"), - ], -) -async def test_async_routing_strategy_pre_call_checks_failure_logging_is_coordinated(hook_error): - class _RaisingPreCallCheck(CustomLogger): - async def async_pre_call_check(self, deployment, parent_otel_span): - raise hook_error - - router = litellm.Router( - model_list=[{"model_name": "gpt-5.6", "litellm_params": {"model": "openai/gpt-5.6", "api_key": "sk-fake"}}] - ) - deployment = router.model_list[0] - logging_obj = _make_failure_logging_obj() - - with patch.object(litellm, "callbacks", [_RaisingPreCallCheck()]): # test-quality-ok: router reads this global - await _assert_router_failure_logging_is_coordinated( - logging_obj, - lambda: router.async_routing_strategy_pre_call_checks( - deployment=deployment, parent_otel_span=None, logging_obj=logging_obj - ), - type(hook_error), - ) - - -@pytest.mark.asyncio -async def test_async_callback_filter_deployments_failure_logging_is_coordinated(): - class _RaisingFilter(CustomLogger): - async def async_filter_deployments(self, *args, **kwargs): - raise RuntimeError("filter blew up") - - router = litellm.Router( - model_list=[{"model_name": "gpt-5.6", "litellm_params": {"model": "openai/gpt-5.6", "api_key": "sk-fake"}}] - ) - logging_obj = _make_failure_logging_obj() - - with patch.object(litellm, "callbacks", [_RaisingFilter()]): # test-quality-ok: router reads this global - await _assert_router_failure_logging_is_coordinated( - logging_obj, - lambda: router.async_callback_filter_deployments( - model="gpt-5.6", - healthy_deployments=router.model_list, - messages=None, - parent_otel_span=None, - request_kwargs={}, - logging_obj=logging_obj, - ), - RuntimeError, - ) - - -@pytest.mark.asyncio -async def test_async_get_available_deployment_failure_logging_is_coordinated(): - router = litellm.Router( - model_list=[{"model_name": "gpt-5.6", "litellm_params": {"model": "openai/gpt-5.6", "api_key": "sk-fake"}}] - ) - logging_obj = _make_failure_logging_obj() - - await _assert_router_failure_logging_is_coordinated( - logging_obj, - lambda: router.async_get_available_deployment( - model="model-that-is-not-configured", - request_kwargs={"litellm_logging_obj": logging_obj}, - messages=[{"role": "user", "content": "hi"}], - ), - litellm.BadRequestError, - ) - - -@pytest.mark.asyncio -async def test_async_get_available_deployment_for_pass_through_failure_logging_is_coordinated(): - router = litellm.Router( - model_list=[{"model_name": "gpt-5.6", "litellm_params": {"model": "openai/gpt-5.6", "api_key": "sk-fake"}}] - ) - logging_obj = _make_failure_logging_obj() - - await _assert_router_failure_logging_is_coordinated( - logging_obj, - lambda: router.async_get_available_deployment_for_pass_through( - model="gpt-5.6", - request_kwargs={"litellm_logging_obj": logging_obj}, - ), - litellm.BadRequestError, - ) - - -class _InFlightTracker: - def __init__(self) -> None: - self.current = 0 - self.peak = 0 - - def enter(self) -> None: - self.current += 1 - self.peak = max(self.peak, self.current) - - def exit(self) -> None: - self.current -= 1 - - -_SSE_CHUNKS: Final[tuple[bytes, ...]] = ( - *( - b'data: {"id":"c","object":"chat.completion.chunk","created":1,"model":"gpt-5.6",' - b'"choices":[{"index":0,"delta":{"content":"x"},"finish_reason":null}]}\n\n' - for _ in range(5) - ), - b'data: {"id":"c","object":"chat.completion.chunk","created":1,"model":"gpt-5.6",' - b'"choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}\n\n', -) - - -class _CountingSSEStream(httpx.AsyncByteStream): - def __init__(self, tracker: _InFlightTracker) -> None: - self._tracker = tracker - self._in_flight = False - - def _finish(self) -> None: - if self._in_flight: - self._in_flight = False - self._tracker.exit() - - async def __aiter__(self): - self._in_flight = True - self._tracker.enter() - try: - for chunk in _SSE_CHUNKS: - await asyncio.sleep(0.02) - yield chunk - finally: - await self.aclose() - yield b"data: [DONE]\n\n" - - async def aclose(self) -> None: - await asyncio.sleep(0.02) - self._finish() - - -def _max_parallel_router(max_parallel_requests: int) -> Router: - return Router( - model_list=[ - { - "model_name": "gpt-5.6", - "litellm_params": { - "model": "openai/gpt-5.6", - "api_key": "sk-fake", - "api_base": "https://max-parallel.local/v1", - "max_parallel_requests": max_parallel_requests, - }, - } - ], - num_retries=0, - ) - - -@pytest.mark.asyncio -@pytest.mark.parametrize("stream", [False, True]) -async def test_router_max_parallel_requests_bounds_in_flight_upstream_calls( - monkeypatch: pytest.MonkeyPatch, stream: bool -): - monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) - tracker: Final = _InFlightTracker() - router: Final = _max_parallel_router(max_parallel_requests=2) - - async def upstream(request: httpx.Request) -> httpx.Response: - if stream: - return httpx.Response( - 200, headers={"content-type": "text/event-stream"}, stream=_CountingSSEStream(tracker) - ) - tracker.enter() - await asyncio.sleep(0.05) - tracker.exit() - return httpx.Response( - 200, - json={ - "id": "c", - "object": "chat.completion", - "created": 1, - "model": "gpt-5.6", - "choices": [{"index": 0, "message": {"role": "assistant", "content": "x"}, "finish_reason": "stop"}], - }, - ) - - async def one_call() -> None: - response = await router.acompletion( - model="gpt-5.6", messages=[{"role": "user", "content": "hi"}], stream=stream - ) - if stream: - async for _ in response: - pass - - with respx.mock(assert_all_called=True) as respx_mock: - respx_mock.post("https://max-parallel.local/v1/chat/completions").mock(side_effect=upstream) - await asyncio.wait_for(asyncio.gather(*(one_call() for _ in range(10))), timeout=10) - - assert tracker.peak <= 2 - assert tracker.current == 0 - - -@pytest.mark.asyncio -async def test_router_max_parallel_requests_slot_released_when_stream_closed_early(monkeypatch: pytest.MonkeyPatch): - monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) - tracker: Final = _InFlightTracker() - router: Final = _max_parallel_router(max_parallel_requests=1) - - with respx.mock() as respx_mock: - respx_mock.post("https://max-parallel.local/v1/chat/completions").mock( - side_effect=lambda request: httpx.Response( - 200, headers={"content-type": "text/event-stream"}, stream=_CountingSSEStream(tracker) - ) - ) - first: Final = await router.acompletion( - model="gpt-5.6", messages=[{"role": "user", "content": "hi"}], stream=True - ) - await first.__anext__() - - async def second_call() -> None: - second = await router.acompletion( - model="gpt-5.6", messages=[{"role": "user", "content": "hi"}], stream=True - ) - async for _ in second: - pass - - second_task: Final = asyncio.create_task(second_call()) - await asyncio.sleep(0.05) - assert tracker.current == 1 - await first.aclose() - await asyncio.wait_for(second_task, timeout=2) - - assert tracker.peak == 1 - assert tracker.current == 0 +@file:/tmp/litellm-work/litellm/tests/test_litellm/test_router.py \ No newline at end of file