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