# +-----------------------------------------------+ # | | # | Give Feedback / Get Help | # | https://github.com/BerriAI/litellm/issues/new | # | | # +-----------------------------------------------+ # # Thank you ! We ❤️ you! - Krrish & Ishaan import asyncio import contextlib import copy import enum import hashlib import inspect import json import logging import re import threading import time import traceback import weakref from collections import defaultdict from collections.abc import ( AsyncGenerator, AsyncIterator, Callable, Generator, Iterator, Mapping, MutableMapping, Sequence, ) from datetime import datetime, timezone from functools import lru_cache, partial from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeAlias, TypeVar, Union, cast import anyio import httpx import openai from openai import AsyncOpenAI from pydantic import BaseModel, TypeAdapter, ValidationError from typing_extensions import overload import litellm import litellm.litellm_core_utils.exception_mapping_utils from litellm import get_secret_str from litellm._logging import verbose_router_logger from litellm._uuid import uuid from litellm.caching.caching import ( DualCache, InMemoryCache, RedisCache, RedisClusterCache, ) from litellm.caching.redis_cache import log_redis_failure from litellm.constants import ( CLIENT_OUTPUT_CEILING_METADATA_KEY, CONSUMED_REQUEST_TAGS_METADATA_KEY, DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS, DEFAULT_HEALTH_CHECK_INTERVAL, DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER, DEFAULT_MAX_LRU_CACHE_SIZE, INTERNAL_CALL_ORIGIN_METADATA_KEY, OUTPUT_TOKEN_CEILING_PARAMS, ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY, ROUTING_REQUEST_TAGS_METADATA_KEY, RUNTIME_UPDATABLE_ROUTER_SETTINGS, SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY, ) from litellm.integrations.custom_guardrail import is_guardrail_intervention from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.asyncify import run_async_function from litellm.litellm_core_utils.core_helpers import ( _get_parent_otel_span_from_kwargs, coerce_token_limit, get_litellm_metadata_from_kwargs, get_metadata_variable_name_from_kwargs, get_or_create_metadata_bucket, ) from litellm.litellm_core_utils.coroutine_checker import coroutine_checker from litellm.litellm_core_utils.credential_accessor import CredentialAccessor from litellm.litellm_core_utils.dd_tracing import tracer from litellm.litellm_core_utils.get_llm_provider_logic import ( declared_authenticating_provider, is_registered_custom_provider, ) from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.litellm_core_utils.ptu_pricing import ( PTU_COST_ATTRIBUTION_ENV_VAR, declares_ptu, is_ptu_cost_attribution_enabled, ptu_config_error, ptu_identity_error, ptu_terms, zeroed_ptu_pricing, ) from litellm.litellm_core_utils.request_timeout_resolver import ( get_configured_request_timeout, ) from litellm.litellm_core_utils.secret_redaction import redact_string from litellm.litellm_core_utils.sensitive_data_masker import ( SensitiveDataMasker, mask_sensitive_structure, ) from litellm.litellm_core_utils.token_counter import offload_token_count from litellm.llms.anthropic.common_utils import AnthropicModelInfo from litellm.llms.base_llm.passthrough.transformation import replace_path_segment from litellm.llms.base_llm.vector_store.transformation import ( RouterVectorStoreEmbeddingExecutor, vector_store_request_metadata, ) from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, get_async_httpx_client from litellm.llms.openai_like.json_loader import JSONProviderRegistry from litellm.llms.openai_like.model_info import ( MODEL_INFO_DISCOVERY_PROVIDERS, MODEL_INFO_REFRESH_CONCURRENCY, MODEL_INFO_REFRESH_SECONDS, get_openai_compatible_model_info, ) from litellm.router_strategy.base_routing_strategy import BaseRoutingStrategy from litellm.router_strategy.budget_limiter import RouterBudgetLimiting from litellm.router_strategy.complexity_router.context_compaction import ( arm_compaction, compact_to_fit, compaction_pending, initialize_compaction_state, is_native_compaction_call, reject_recursive_compactor, surface_for_call, ) from litellm.router_strategy.least_busy import LeastBusyLoggingHandler from litellm.router_strategy.lowest_cost import LowestCostLoggingHandler from litellm.router_strategy.lowest_latency import LowestLatencyLoggingHandler from litellm.router_strategy.lowest_tpm_rpm import LowestTPMLoggingHandler from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2, PrefetchedUsage from litellm.router_strategy.simple_shuffle import simple_shuffle from litellm.router_strategy.tag_based_routing import ( _get_tags_from_request_kwargs, get_deployments_for_tag, is_valid_deployment_tag, ) from litellm.router_utils.access_windows import access_windows_config_error, filter_reserved_deployments from litellm.router_utils.add_retry_fallback_headers import ( _HiddenParamsHost, add_fallback_headers_to_response, add_retry_headers_to_response, apply_quality_router_decision_headers, apply_remaining_usage_headers, apply_response_model_id, complexity_router_decision_headers, ensure_response_additional_headers, get_hidden_params_dict, prepare_response_for_header_attachment, replace_complexity_router_headers, response_total_token_count, ) from litellm.router_utils.auto_router_model_naming import ( AUTO_ROUTER_MODEL_PREFIX, GatedAutoRouterCapability, capability_limit_violation, claimed_capability, classify_strategy_router_model, count_capability_routers, ) from litellm.router_utils.batch_utils import ( _get_router_metadata_variable_name, is_batch_retrieve_call_type, replace_model_in_jsonl, should_replace_model_in_jsonl, ) from litellm.router_utils.client_initalization_utils import InitalizeCachedClient, MaxParallelRequestsLimit from litellm.router_utils.clientside_credential_handler import ( get_dynamic_litellm_params, is_clientside_credential, ) from litellm.router_utils.common_utils import ( _is_proxy_admin_request, filter_team_based_models, filter_web_search_deployments, format_fallback_outcome_message, format_no_fallback_group_message, get_request_team_id, provider_for_generic_call, resolve_model_group_alias, truncate_fallback_error_detail, warn_on_provider_credential_mismatch, ) from litellm.router_utils.cooldown_cache import CooldownCache from litellm.router_utils.cooldown_handlers import ( DEFAULT_COOLDOWN_TIME_SECONDS, _async_get_cooldown_deployments, _async_get_cooldown_deployments_with_debug_info, _first_present, # pyright: ignore[reportPrivateUsage] - shared internal helper across router_utils submodules, matching the other cooldown_handlers imports on this line _get_cooldown_deployments, _set_cooldown_deployments, is_advisor_orchestration_failure, is_background_response_cost_poll_not_found, is_caller_timeout_408, ) from litellm.router_utils.fallback_event_handlers import ( MID_STREAM_FALLBACK_CONTROLS_KEY, AttemptedFallbackTargets, _check_non_standard_fallback_format, carry_over_pre_routing_selection, clear_pre_routing_selection, fallback_lookup_groups, fallbacks_disabled_for_request, get_fallback_model_group_for_lookup_groups, get_pre_routing_selection, has_unattempted_fallback_target, mid_stream_fallback_hop_kwargs, per_request_fallback_controls, record_disable_fallbacks, record_pre_routing_selection, run_async_fallback, ) from litellm.router_utils.get_retry_from_policy import ( get_num_retries_from_retry_policy as _get_num_retries_from_retry_policy, ) from litellm.router_utils.handle_error import ( async_raise_no_deployment_exception, send_llm_exception_alert, ) from litellm.router_utils.health_state_cache import DeploymentHealthCache from litellm.router_utils.pre_call_checks.deployment_affinity_check import ( DeploymentAffinityCheck, warn_on_unknown_model_group_affinity_flags, ) from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import ( EncryptedContentAffinityCheck, ) from litellm.router_utils.pre_call_checks.io_token_rate_limit_check import ( build_io_token_rate_limit_headers, deployment_has_io_token_limits, refund_stale_reservation_before_retry, set_io_token_rate_limit_request_kwargs, ) from litellm.router_utils.pre_call_checks.model_rate_limit_check import ( ModelRateLimitingCheck, ) from litellm.router_utils.pre_call_checks.prompt_caching_deployment_check import ( PromptCachingDeploymentCheck, ) from litellm.router_utils.reasoning_effort_capability import ( deployment_is_catalog_mapped, intersect_supported_reasoning_efforts, resolve_supported_reasoning_efforts, ) from litellm.router_utils.router_callbacks.track_deployment_metrics import ( find_deployment_metadata, get_counted_usage_tokens, increment_deployment_failures_for_current_minute, increment_deployment_successes_for_current_minute, ) from litellm.router_utils.routing_groups import ( apply_routing_group_priority, parse_routing_groups, validate_routing_strategy, ) from litellm.router_utils.routing_read_batch import RoutingPrefetch, RoutingReadBatch from litellm.scheduler import FlowItem, Scheduler from litellm.types.litellm_params import RoutingStrategyName from litellm.types.llms.openai import ( AllMessageValues, ChatCompletionToolParam, FileTypes, OpenAIFileObject, OpenAIFilesPurpose, ) from litellm.types.router import ( CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS, VALID_LITELLM_ENVIRONMENTS, AlertingConfig, AllowedFailsPolicy, AssistantsTypedDict, AutoRouterCapabilityLimit, ConsumedRequestTagsStamp, CredentialLiteLLMParams, CustomRoutingStrategyBase, Deployment, DeploymentModelListingInfo, DeploymentTypedDict, DiscoveredDeploymentModelInfo, FallbackAccessCheck, FallbackBudgetCheck, GuardrailTypedDict, LiteLLM_Params, MockRouterTestingParams, ModelGroupInfo, OptionalPreCallChecks, PreRoutingStrategy, RetryPolicy, RouterCacheEnum, RouterErrors, RouterGeneralSettings, RouterModelGroupAliasItem, RouterRateLimitError, RouterRateLimitErrorBasic, RoutingContext, RoutingGroup, RoutingPlugin, RoutingStrategy, SearchToolTypedDict, TaggedPreRoutingStrategy, ) from litellm.types.services import ServiceTypes from litellm.types.utils import ( AUTOROUTER_CLASSIFIER_CALL_ORIGIN, PROMPT_QUOTING_ROUTING_DECISION_FIELDS, CustomPricingLiteLLMParams, GenericBudgetConfigType, LiteLLMBatch, LlmProviders, ModelInfo, ModelResponseStream, StandardLoggingPayload, StandardLoggingRoutingDecision, Usage, all_litellm_params, shared_backend_model_info, ) from litellm.types.utils import ModelInfo as ModelMapInfo from litellm.utils import ( CustomStreamWrapper, EmbeddingResponse, ModelResponse, Rules, function_setup, get_llm_provider, get_non_default_completion_params, get_secret, get_utc_datetime, is_region_allowed, provider_rejectable_params, set_live_deployment_replay, ) from .router_utils.pattern_match_deployments import PatternMatchRouter if TYPE_CHECKING: from opentelemetry.trace import Span as _Span from litellm.exceptions import MidStreamFallbackError from litellm.responses.streaming_iterator import ( BaseResponsesAPIStreamingIterator, ) from litellm.router_strategy.adaptive_router.adaptive_router import ( AdaptiveRouter, ) from litellm.router_strategy.auto_router.auto_router import ( AutoRouter, PreRoutingHookResponse, ) from litellm.router_strategy.complexity_router.complexity_router import ( ComplexityRouter, ) from litellm.router_strategy.quality_router.quality_router import ( QualityRouter, ) from litellm.types.llms.anthropic_messages.anthropic_response import ( AnthropicMessagesResponse, ) from litellm.types.llms.base import BaseLiteLLMOpenAIResponseObject from litellm.types.llms.openai import ( ResponseAPIUsage, ResponseInputParam, ResponsesAPIResponse, ) Span = _Span else: Span = Any AutoRouter = Any ComplexityRouter = Any AdaptiveRouter = Any QualityRouter = Any PreRoutingHookResponse = Any RouterStrategySelector: TypeAlias = ( LeastBusyLoggingHandler | LowestCostLoggingHandler | LowestLatencyLoggingHandler | LowestTPMLoggingHandler | LowestTPMLoggingHandler_v2 ) def _cost_value_as_float(value: str | float | None) -> float | None: if value is None: return None try: return float(value) except (TypeError, ValueError): return None def model_info_is_active_for_environment(model_info: Mapping[str, object] | None) -> bool: """Single owner of the environment-gating rule: a deployment whose model_info names `supported_environments` loads only on pods whose LITELLM_ENVIRONMENT is in that list. `Router.deployment_is_active_for_environment` delegates here, and the model-write endpoints consult the same rule to tell a deliberately inactive model from one that was dropped by a failed reload.""" if model_info is None: return True supported_environments: Final = model_info.get("supported_environments") if supported_environments is None: return True if not isinstance(supported_environments, (list, tuple)): raise ValueError( f"supported_environments must be a list of {VALID_LITELLM_ENVIRONMENTS}. " f"but set as: {supported_environments} for model_info: {model_info}" ) litellm_environment: Final = get_secret_str(secret_name="LITELLM_ENVIRONMENT") if litellm_environment is None: raise ValueError("Set 'supported_environments' for model but not 'LITELLM_ENVIRONMENT' set in .env") if litellm_environment not in VALID_LITELLM_ENVIRONMENTS: raise ValueError( f"LITELLM_ENVIRONMENT must be one of {VALID_LITELLM_ENVIRONMENTS}. but set as: {litellm_environment}" ) for _env in supported_environments: if _env not in VALID_LITELLM_ENVIRONMENTS: raise ValueError( f"supported_environments must be one of {VALID_LITELLM_ENVIRONMENTS}. but set as: {_env} " f"for model_info: {model_info}" ) if litellm_environment in supported_environments: return True return False _PreRoutingStrategyT = TypeVar("_PreRoutingStrategyT") _CallbackT = TypeVar("_CallbackT") _ALIAS_PARAMS_NEVER_FORWARDED: Final = frozenset({"model", "api_base", "api_key", "api_version"}) _ALIAS_MARKER_FORWARDED_PARAMS_KWARG: Final = "_alias_marker_forwarded_params" _CLAUDE_CODE_SESSION_ID_RE: Final = re.compile(r"^[a-zA-Z0-9_\-]{8,}$") _CLAUDE_CODE_SESSION_ROUTER_TTL_SECONDS: Final = 3600 _RUNTIME_TOGGLEABLE_PRE_CALL_CHECKS: Final[Mapping[str, type[CustomLogger]]] = MappingProxyType( { "prompt_caching": PromptCachingDeploymentCheck, "enforce_model_rate_limits": ModelRateLimitingCheck, "encrypted_content_affinity": EncryptedContentAffinityCheck, } ) def _stream_chunks_have_generated_content(chunks: Sequence[ModelResponseStream]) -> bool: for chunk in chunks: if not chunk.choices: continue delta = chunk.choices[0].delta if ( delta.get("content") or delta.get("tool_calls") or delta.get("function_call") or delta.get("reasoning_content") or delta.get("thinking_blocks") or delta.get("reasoning_items") or delta.get("audio") or delta.get("images") or delta.get("annotations") ): return True return False _NO_SESSION_KWARGS: Final[Mapping[str, Mapping[str, object]]] = MappingProxyType({}) _SESSION_ADAPTER: Final = TypeAdapter(Mapping[str, object]) _SILENT_MODEL_ADAPTER: Final = TypeAdapter(str | list[str]) def _as_retry_skipped_deployment_ids(value: object) -> tuple[str, ...]: return tuple(item for item in value if isinstance(item, str)) if isinstance(value, tuple) else () def _silent_experiment_targets(silent_model: object) -> tuple[str, ...]: if silent_model is None: return () try: targets: Final = _SILENT_MODEL_ADAPTER.validate_python(silent_model) except ValidationError: verbose_router_logger.warning( "silent_model must be a model name or a list of model names, got %r; skipping shadow traffic", silent_model, ) return () return (targets,) if isinstance(targets, str) else tuple(targets) def _silent_experiment_kwargs_snapshot(kwargs: Mapping[str, object]) -> Mapping[str, object]: metadata: Final = kwargs.get("metadata") if not isinstance(metadata, Mapping): return MappingProxyType({**kwargs}) return MappingProxyType({**kwargs, "metadata": dict(metadata)}) def _with_router_resolved_session_model(session: object, model_name: str) -> Mapping[str, Mapping[str, object]]: """ Realtime client-secret requests carry the model inside ``session`` as well, and the caller's copy of it still holds the pre-routing model group name, so it has to follow the deployment the router just picked. Returns kwargs to merge into the downstream call, empty when there is no session model to resolve. """ try: typed_session: Final = _SESSION_ADAPTER.validate_python(session) except ValidationError: return _NO_SESSION_KWARGS if "model" not in typed_session: return _NO_SESSION_KWARGS return MappingProxyType( {"session": {**typed_session, "model": model_name}} # mutable-ok: callees deepcopy and JSON-dump session ) # Router._aanthropic_messages_streaming_iterator buffers lifecycle chunks # until real content commits the primary stream, and only while a fallback # can still take over; a hostile or slow-starting upstream that never emits # content or an error could otherwise grow that buffer without bound, so # hitting this cap forces an early commit instead. MAX_BUFFERED_PRE_CONTENT_ANTHROPIC_CHUNKS: Final = 200 def _anthropic_stream_forwards_ping_live(chunk: object, has_generated_content: bool) -> bool: """A `ping` keepalive reaches the client live whenever the stream has not committed: it carries no lifecycle, so it cannot create overlapping lifecycles on the wire, and it keeps the connection alive while lifecycle frames sit buffered for a possible fallback during a long thinking pass.""" from litellm.llms.anthropic.pass_through.messages.streaming_iterator import is_anthropic_ping_chunk return not has_generated_content and is_anthropic_ping_chunk(chunk) def _is_retriable_anthropic_status(status_code: int) -> bool: return status_code == 429 or status_code >= 500 def _without_line_breaks(value: object) -> str: return str(value).replace("\r", "").replace("\n", "") def _anthropic_stream_error_is_gateway_verdict(chunk: object) -> bool: """AgenticAnthropicStreamingIterator's own retrieval-failure frame is the gateway's verdict, not a provider failure: another deployment would rerun the same failed hook, so it reaches the client instead of falling back.""" from litellm.llms.anthropic.pass_through.messages.agentic_streaming_iterator import ( is_server_fulfilled_tool_leak_error, ) return is_server_fulfilled_tool_leak_error(chunk) def _anthropic_stream_should_decline_fallback(has_generated_content: bool, error: "MidStreamFallbackError") -> bool: """ A MidStreamFallbackError raised directly by the source iterator (the completion-bridge path's CustomStreamWrapper, e.g. on a transport drop) carries its own pre_first_chunk bookkeeping - gated the same way a detected SSE error event is, so a fallback is never appended after real content already reached the client on either path. """ return has_generated_content or not error.is_pre_first_chunk def _anthropic_stream_raised_error_status(error: Exception) -> int | None: raw_status: Final = getattr(error, "status_code", None) if isinstance(raw_status, int): return raw_status if isinstance(raw_status, str) and raw_status.isdigit(): return int(raw_status) response_status: Final = getattr(getattr(error, "response", None), "status_code", None) return response_status if isinstance(response_status, int) else None def _anthropic_stream_fallback_error_for_raised( error: Exception, model: str, has_generated_content: bool ) -> "MidStreamFallbackError | None": """Same gate as a detected SSE error event; None means the raise propagates unchanged.""" from litellm.exceptions import MidStreamFallbackError if has_generated_content: return None status_code: Final = _anthropic_stream_raised_error_status(error) if status_code is not None and not _is_retriable_anthropic_status(status_code): return None return MidStreamFallbackError( message=str(error), model=model, llm_provider="anthropic", original_exception=error, is_pre_first_chunk=True, ) def _anthropic_stream_commits_now(chunk: object, has_generated_content: bool, buffered_chunk_count: int) -> bool: """ Whether `chunk` should make Router._aanthropic_messages_streaming_iterator commit to the primary Anthropic stream (real content arrived, or the pre-content buffer cap was hit) rather than keep buffering lifecycle frames toward a possible fallback. """ from litellm.llms.anthropic.pass_through.messages.streaming_iterator import ( is_anthropic_content_delta_chunk, ) if has_generated_content: return False return is_anthropic_content_delta_chunk(chunk) or buffered_chunk_count >= MAX_BUFFERED_PRE_CONTENT_ANTHROPIC_CHUNKS MAX_HELD_PRE_OUTPUT_RESPONSES_EVENTS: Final = 200 def _responses_stream_holds_event(item: object, held_event_count: int) -> bool: from litellm.responses.streaming_iterator import PRE_OUTPUT_LIFECYCLE_EVENT_TYPES if held_event_count >= MAX_HELD_PRE_OUTPUT_RESPONSES_EVENTS: return False return getattr(item, "type", None) in PRE_OUTPUT_LIFECYCLE_EVENT_TYPES class FallbackAwareAnthropicMessagesStream: """ Bare async generators can't carry the `_hidden_params` attribute the proxy reads response headers off of (see router_utils.add_retry_fallback_headers.get_hidden_params_dict), so this thin wrapper carries it through from the source iterator - mirrors AnthropicMessagesStreamingResponse. Used by Router._aanthropic_messages_streaming_iterator. """ def __init__(self, async_generator: AsyncGenerator[bytes, None], source_iterator: object) -> None: self._async_generator = async_generator self._source_iterator = source_iterator self.fallback_headers_adopted = False self._hidden_params = dict( # mutable-ok: mutated in place by merge_fallback_hidden_params getattr(source_iterator, "_hidden_params", None) or {} ) @property def has_buffered_provider_output(self) -> bool: return getattr(self._source_iterator, "has_buffered_provider_output", False) is True @property def chunks(self) -> list[ModelResponseStream] | None: return cast( # cast-ok: chunks is a list of ModelResponseStream on the inner stream "list[ModelResponseStream] | None", getattr(self._source_iterator, "chunks", None) ) @property def messages(self) -> list[AllMessageValues] | None: return cast( # cast-ok: messages is a list of AllMessageValues on the inner stream "list[AllMessageValues] | None", getattr(self._source_iterator, "messages", None) ) @property def model(self) -> str | None: return cast( # cast-ok: model is a str on the inner stream "str | None", getattr(self._source_iterator, "model", None) ) def adopt_fallback_source(self, fallback_response: object) -> None: self._source_iterator = fallback_response self.fallback_headers_adopted = True def __aiter__(self) -> "FallbackAwareAnthropicMessagesStream": return self async def __anext__(self) -> bytes: return await self._async_generator.__anext__() async def aclose(self) -> None: await self._async_generator.aclose() def merge_fallback_hidden_params( self, fallback_hidden_params: Mapping[str, object], fallback_headers: Mapping[str, object], ) -> None: """ Raw bytes can't carry their own _hidden_params the way a ModelResponseStream/ResponsesAPI event can, so a mid-stream fallback's provider headers (e.g. Bedrock's x-amzn-requestid) are merged onto the wrapper itself instead - mirrors Router._apply_fallback_hidden_params_to_item's merge shape. """ existing_headers: Final = cast( # cast-ok: additional_headers is always a dict[str, object] when present "dict[str, object]", self._hidden_params.get("additional_headers") or {} ) self._hidden_params = { # mutable-ok: matches _hidden_params' existing dict[str, object] shape **self._hidden_params, **fallback_hidden_params, "additional_headers": dict( # mutable-ok: hidden params expect a writable header bag replace_complexity_router_headers(existing_headers, fallback_headers) ), } class RoutingArgs(enum.Enum): ttl = 60 # 1min (RPM/TPM expire key) # Routers that are still in use, so a price data reload can rebuild the cost-map # entries their deployments own. Weak so a router nothing references any more, such # as the per-request one built from a caller-supplied user_config, drops out on its # own rather than leaving entries behind that nothing can withdraw. _live_routers: Final["weakref.WeakSet[Router]"] = weakref.WeakSet() def _replay_live_router_model_cost() -> None: """Re-assert every live router's deployments after the cost map is refreshed.""" for router in tuple(_live_routers): router._replay_model_cost_registrations() set_live_deployment_replay(_replay_live_router_model_cost) RETRY_BREADCRUMB_LIMIT: Final = 4 class FallbackAwareStreamWrapper(CustomStreamWrapper): """Base for the Router's chat-completion stream wrappers, which are built around the attempt the Router picked first and have to repoint themselves when a fallback takes over.""" fallback_headers_adopted: bool = False def adopt_fallback_response_headers( self, fallback_response: object, prepared_fallback_hidden_params: tuple[dict[str, object], dict[str, object]], ) -> None: """Repoint this wrapper at the deployment that served the stream. Replaces rather than merges, so the failed attempt's `x-request-id`, rate limit counters, `model_id` and `api_base` cannot reach the proxy's response headers or its callbacks. """ self._response_headers = getattr(fallback_response, "_response_headers", None) fallback_hidden_params, fallback_headers = prepared_fallback_hidden_params if fallback_hidden_params: self._hidden_params = { # mutable-ok: the rest of litellm writes into _hidden_params **fallback_hidden_params, # dict() because add_retry_fallback_headers mutates additional_headers in place "additional_headers": dict(fallback_headers), # mutable-ok: see above } self._base_hidden_params = { # mutable-ok: CustomStreamWrapper keeps this snapshot as a dict **self._hidden_params, "response_cost": None, } self.fallback_headers_adopted = True def as_output_cap(value: object) -> int | None: """A client-sent output cap coerced to an int: ints, floats and numeric strings, never bools or negatives.""" if isinstance(value, bool) or not isinstance(value, (int, float, str)): return None try: cap: Final = int(float(value)) except (ValueError, OverflowError): return None return cap if cap >= 0 else None class Router: model_names: set = set() cache_responses: bool | None = False default_cache_time_seconds: int = 1 * 60 * 60 # 1 hour tenacity = None leastbusy_logger: LeastBusyLoggingHandler | None = None lowesttpm_logger: LowestTPMLoggingHandler | None = None optional_callbacks: list[CustomLogger | Callable | str] | None = None def __init__( self, model_list: list[DeploymentTypedDict] | list[dict[str, Any]] | None = None, ## ASSISTANTS API ## assistants_config: AssistantsTypedDict | None = None, ## SEARCH API ## search_tools: list[SearchToolTypedDict] | None = None, ## GUARDRAIL API ## guardrail_list: list[GuardrailTypedDict] | None = None, ## CACHING ## redis_url: str | None = None, redis_host: str | None = None, redis_port: int | None = None, redis_password: str | None = None, redis_db: int | None = None, cache_responses: bool | None = False, cache_kwargs: dict = {}, # additional kwargs to pass to RedisCache (see caching.py) caching_groups: list[tuple] | None = None, # if you want to cache across model groups client_ttl: int = 3600, # ttl for cached clients - will re-initialize after this time in seconds ## SCHEDULER ## polling_interval: float | None = None, default_priority: int | None = None, ## RELIABILITY ## num_retries: int | None = None, max_fallbacks: int | None = None, # max fallbacks to try before exiting the call. Defaults to 5. timeout: float | None = None, stream_timeout: float | None = None, default_litellm_params: dict | None = None, # default params for Router.chat.completion.create default_max_parallel_requests: int | None = None, set_verbose: bool = False, debug_level: Literal["DEBUG", "INFO"] = "INFO", default_fallbacks: list[str] | None = None, # generic fallbacks, works across all deployments fallbacks: list = [], context_window_fallbacks: list = [], content_policy_fallbacks: list = [], model_group_alias: dict[str, str | RouterModelGroupAliasItem] | None = {}, enable_pre_call_checks: bool = False, enable_tag_filtering: bool = False, tag_filtering_match_any: bool = True, tag_routing_prefix: str = "", plugins: list[RoutingPlugin] | None = None, retry_after: int = 0, # min time to wait before retrying a failed request retry_policy: RetryPolicy | dict | None = None, # set custom retries for different exceptions model_group_retry_policy: dict[str, RetryPolicy] = {}, # set custom retry policies based on model group allowed_fails: int | None = None, # Number of times a deployment can failbefore being added to cooldown allowed_fails_policy: AllowedFailsPolicy | None = None, # set custom allowed fails policy cooldown_time: float | None = None, # (seconds) time to cooldown a deployment after failure disable_cooldowns: bool | None = None, routing_strategy: RoutingStrategyName = "simple-shuffle", optional_pre_call_checks: OptionalPreCallChecks | None = None, routing_strategy_args: dict = {}, # just for latency-based routing_groups: list[RoutingGroup | dict] | None = None, provider_budget_config: GenericBudgetConfigType | None = None, alerting_config: AlertingConfig | None = None, router_general_settings: RouterGeneralSettings | None = RouterGeneralSettings(), deployment_affinity_ttl_seconds: int = 3600, model_group_affinity_config: dict[str, list[str]] | None = None, ignore_invalid_deployments: bool = False, enable_health_check_routing: bool = False, health_check_staleness_threshold: int | None = None, health_check_ignore_transient_errors: bool = False, background_health_check_model_groups: Sequence[str] | None = None, enable_weighted_failover: bool = False, fallback_access_check: FallbackAccessCheck | None = None, fallback_budget_check: FallbackBudgetCheck | None = None, auto_router_capability_limit: AutoRouterCapabilityLimit | None = None, ) -> None: """ Initialize the Router class with the given parameters for caching, reliability, and routing strategy. Args: model_list (Optional[list]): List of models to be used. Defaults to None. redis_url (Optional[str]): URL of the Redis server. Defaults to None. redis_host (Optional[str]): Hostname of the Redis server. Defaults to None. redis_port (Optional[int]): Port of the Redis server. Defaults to None. redis_password (Optional[str]): Password of the Redis server. Defaults to None. cache_responses (Optional[bool]): Flag to enable caching of responses. Defaults to False. cache_kwargs (dict): Additional kwargs to pass to RedisCache. Defaults to {}. caching_groups (Optional[List[tuple]]): List of model groups for caching across model groups. Defaults to None. client_ttl (int): Time-to-live for cached clients in seconds. Defaults to 3600. polling_interval: (Optional[float]): frequency of polling queue. Only for '.scheduler_acompletion()'. Default is 3ms. default_priority: (Optional[int]): the default priority for a request. Only for '.scheduler_acompletion()'. Default is None. num_retries (Optional[int]): Number of retries for failed requests. Defaults to 2. timeout (Optional[float]): Timeout for requests. Defaults to None. default_litellm_params (dict): Default parameters for Router.chat.completion.create. Defaults to {}. set_verbose (bool): Flag to set verbose mode. Defaults to False. debug_level (Literal["DEBUG", "INFO"]): Debug level for logging. Defaults to "INFO". fallbacks (List): List of fallback options. Defaults to []. context_window_fallbacks (List): List of context window fallback options. Defaults to []. enable_pre_call_checks (boolean): Filter out deployments which are outside context window limits for a given prompt model_group_alias (Optional[dict]): Alias for model groups. Defaults to {}. retry_after (int): Minimum time to wait before retrying a failed request. Defaults to 0. allowed_fails (Optional[int]): Number of allowed fails before adding to cooldown. Defaults to None. cooldown_time (float): Time to cooldown a deployment after failure in seconds. Defaults to 1. routing_strategy (Literal["simple-shuffle", "least-busy", "usage-based-routing", "latency-based-routing", "cost-based-routing"]): Routing strategy used for the implicit "default" group (any model not claimed by an entry in `routing_groups`). Defaults to "simple-shuffle". routing_strategy_args (dict): Additional args for the default group's routing strategy (e.g. latency window). Defaults to {}. routing_groups (Optional[List[RoutingGroup]]): Named subsets of `model_name`s with a group routing strategy. Priority groups apply only to group calls and may overlap. Other groups supply their members' default strategy, with at most one such group per model. Unclaimed models use the top-level strategy. alerting_config (AlertingConfig): Slack alerting configuration. Defaults to None. provider_budget_config (ProviderBudgetConfig): Provider budget configuration. Use this to set llm_provider budget limits. example $100/day to OpenAI, $100/day to Azure, etc. Defaults to None. deployment_affinity_ttl_seconds (int): TTL for user-key -> deployment affinity mapping. Defaults to 3600. ignore_invalid_deployments (bool): Ignores invalid deployments, and continues with other deployments. Default is to raise an error. enable_weighted_failover (bool): When True and the routing strategy is "simple-shuffle", a retryable failure on one deployment causes the request to re-pick (weighted) across the other deployments in the same model group before any cross-group fallback runs. Bounded by `max_fallbacks`. Async-only: currently honored by `router.acompletion()` and other async entrypoints. The sync `router.completion()` path falls back to the regular fallback flow. Defaults to False. fallback_access_check (Optional[FallbackAccessCheck]): Awaited before each cross-model-group fallback attempt on the async path; a fallback target it rejects is skipped. Defaults to None (every configured fallback is attempted). fallback_budget_check (Optional[FallbackBudgetCheck]): Awaited before each cross-model-group fallback attempt on the async path; a fallback target it rejects as over budget is skipped. Defaults to None (budget is not re-checked on fallback). Returns: Router: An instance of the litellm.Router class. Example Usage: ```python from litellm import Router model_list = [ { "model_name": "azure-gpt-3.5-turbo", # model alias "litellm_params": { # params for litellm completion/embedding call "model": "azure/", "api_key": , "api_version": , "api_base": }, }, { "model_name": "azure-gpt-3.5-turbo", # model alias "litellm_params": { # params for litellm completion/embedding call "model": "azure/", "api_key": , "api_version": , "api_base": }, }, { "model_name": "openai-gpt-3.5-turbo", # model alias "litellm_params": { # params for litellm completion/embedding call "model": "gpt-3.5-turbo", "api_key": , }, ] router = Router(model_list=model_list, fallbacks=[{"azure-gpt-3.5-turbo": "openai-gpt-3.5-turbo"}]) ``` """ self.set_verbose = set_verbose self.ignore_invalid_deployments = ignore_invalid_deployments self.auto_router_capability_limit = auto_router_capability_limit self.fallback_access_check: Final = fallback_access_check self.fallback_budget_check: Final = fallback_budget_check self.debug_level = debug_level self.enable_pre_call_checks = enable_pre_call_checks self.enable_tag_filtering = enable_tag_filtering self.tag_filtering_match_any = tag_filtering_match_any self.tag_routing_prefix = tag_routing_prefix from litellm._service_logger import ServiceLogging self.service_logger_obj: ServiceLogging = ServiceLogging() litellm.suppress_debug_info = True # prevents 'Give Feedback/Get help' message from being emitted on Router - Relevant Issue: https://github.com/BerriAI/litellm/issues/5942 if self.set_verbose is True: if debug_level == "INFO": verbose_router_logger.setLevel(logging.INFO) elif debug_level == "DEBUG": verbose_router_logger.setLevel(logging.DEBUG) self.router_general_settings: RouterGeneralSettings = router_general_settings or RouterGeneralSettings() self.assistants_config = assistants_config self.search_tools = search_tools or [] self.guardrail_list = guardrail_list or [] self.deployment_names: list = [] # names of models under litellm_params. ex. azure/chatgpt-v-2 self.deployment_latency_map = {} ### CACHING ### cache_type: Literal["local", "redis", "redis-semantic", "s3", "disk"] = "local" # default to an in-memory cache redis_cache = None cache_config: Final[dict[str, Any]] = {} self.client_ttl = client_ttl if redis_url is not None or (redis_host is not None and redis_port is not None): cache_type = "redis" if redis_url is not None: cache_config["url"] = redis_url if redis_host is not None: cache_config["host"] = redis_host if redis_port is not None: cache_config["port"] = str(redis_port) if redis_password is not None: cache_config["password"] = redis_password if redis_db is not None: verbose_router_logger.warning( "Deprecated 'redis_db' argument used. Please remove 'redis_db' from your config/database and use 'cache_kwargs' instead." ) cache_config["db"] = str(redis_db) # Add additional key-value pairs from cache_kwargs cache_config.update(cache_kwargs) redis_cache = self._create_redis_cache(cache_config) if cache_responses: if litellm.cache is None: # the cache can be initialized on the proxy server. We should not overwrite it litellm.cache = litellm.Cache(type=cache_type, **cache_config) self.cache_responses = cache_responses self.cache = DualCache( redis_cache=redis_cache, in_memory_cache=InMemoryCache() ) # use a dual cache (Redis+In-Memory) for tracking cooldowns, usage, etc. self._claude_code_session_router_cache: DualCache = DualCache( redis_cache=redis_cache, in_memory_cache=InMemoryCache(), ) ### SCHEDULER ### self.scheduler = Scheduler(polling_interval=polling_interval, redis_cache=redis_cache) self.default_priority = default_priority self.default_deployment = ( None # use this to track the users default deployment, when they want to use model = * ) self.default_max_parallel_requests = default_max_parallel_requests self.provider_default_deployment_ids: list[str] = [] self.pattern_router = PatternMatchRouter() self.team_pattern_routers: dict[str, PatternMatchRouter] = {} # {"TEAM_ID": PatternMatchRouter} self.auto_routers: dict[str, list[TaggedPreRoutingStrategy[AutoRouter]]] = {} self.complexity_routers: dict[str, list[TaggedPreRoutingStrategy[ComplexityRouter]]] = {} self.adaptive_routers: dict[str, list[TaggedPreRoutingStrategy[AdaptiveRouter]]] = {} self.quality_routers: dict[str, list[TaggedPreRoutingStrategy[QualityRouter]]] = {} self.routing_plugins: list[RoutingPlugin] = list(plugins) if plugins else [] # Initialize model_group_alias early since it's used in set_model_list self.model_group_alias: dict[str, str | RouterModelGroupAliasItem] = ( model_group_alias or {} ) # dict to store aliases for router, ex. {"gpt-4": "gpt-3.5-turbo"}, all requests with gpt-4 -> get routed to gpt-3.5-turbo group # Initialize model ID to deployment index mapping for O(1) lookups self.model_id_to_deployment_index_map: dict[str, int] = {} # Initialize model name to deployment indices mapping for O(1) lookups # Maps model_name -> list of indices in model_list self.model_name_to_deployment_indices: dict[str, list[int]] = {} # Maps (team_id, team_public_model_name) -> list of indices in model_list self.team_model_to_deployment_indices: dict[tuple[str, str], list[int]] = {} self.team_public_model_names: frozenset[str] = frozenset() # Initialize cache attributes that ``_invalidate_model_group_info_cache`` # and ``_invalidate_access_groups_cache`` touch *before* the first # ``set_model_list`` below (which calls those invalidations as part of # building the model index) and before ``_init_routing_groups(None)`` # (which calls them on every group rebuild). self._access_groups_cache: dict[str, list[str]] | None = None # Per-router cache for the proxy auth-layer "is this model explicitly # zero-cost?" check. Lives on the router so it is invalidated alongside # ``_cached_get_model_group_info`` and dies with the router (no # ``id()``-reuse risk after GC). See # ``litellm.proxy.auth.auth_checks._is_model_cost_zero``. self._zero_cost_cache: dict[str, bool] = {} self.cached_deployment_model_info = lru_cache(maxsize=DEFAULT_MAX_LRU_CACHE_SIZE)( self.get_deployment_model_info ) self._discovered_model_info_cache: InMemoryCache = InMemoryCache( max_size_in_memory=max(len(model_list or ()), 1), default_ttl=2 * MODEL_INFO_REFRESH_SECONDS, ) self._routing_group_rows: tuple[DeploymentTypedDict, ...] | None = None self._init_routing_groups(None) self._provider_unresolved_deployments: tuple[Callable[[], Deployment | None], ...] = () self.deployment_affinity_ttl_seconds = deployment_affinity_ttl_seconds self.model_group_affinity_config = model_group_affinity_config warn_on_unknown_model_group_affinity_flags(model_group_affinity_config) if model_list is not None: # set_model_list will build indices automatically self.set_model_list(model_list) # Track this router so a price data reload can rebuild its deployments' # cost-map entries from the list it is serving at that moment. _live_routers.add(self) self.healthy_deployments: list = self.model_list for m in model_list: if "model" in m["litellm_params"]: self.deployment_latency_map[m["litellm_params"]["model"]] = 0 else: self.model_list: list = [] # initialize an empty list - to allow _add_deployment and delete_deployment to work if allowed_fails is not None: self.allowed_fails = allowed_fails else: self.allowed_fails = litellm.allowed_fails self.cooldown_time = cooldown_time or DEFAULT_COOLDOWN_TIME_SECONDS self.cooldown_cache = CooldownCache(cache=self.cache, default_cooldown_time=self.cooldown_time) self.disable_cooldowns = disable_cooldowns self.enable_health_check_routing = enable_health_check_routing self.enable_weighted_failover = enable_weighted_failover self.health_check_ignore_transient_errors = health_check_ignore_transient_errors self.background_health_check_model_groups: frozenset[str] | None = ( frozenset(background_health_check_model_groups) if background_health_check_model_groups is not None else None ) _staleness: Final = health_check_staleness_threshold or ( DEFAULT_HEALTH_CHECK_INTERVAL * DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER ) self.health_state_cache = DeploymentHealthCache(cache=self.cache, staleness_threshold=float(_staleness)) if num_retries is not None: self.num_retries = num_retries elif litellm.num_retries is not None: self.num_retries = litellm.num_retries else: self.num_retries = openai.DEFAULT_MAX_RETRIES if max_fallbacks is not None: self.max_fallbacks = max_fallbacks elif litellm.max_fallbacks is not None: self.max_fallbacks = litellm.max_fallbacks else: self.max_fallbacks = litellm.ROUTER_MAX_FALLBACKS self._explicit_timeout = timeout # None when user did not pass timeout self.timeout = timeout or litellm.request_timeout # Per-attempt request_timeout, independent of router_settings.timeout. # Only stored when a router timeout is also set, since otherwise # request_timeout already flows through self.timeout above. self.request_timeout = get_configured_request_timeout() if timeout is not None else None self.stream_timeout = stream_timeout self.retry_after = retry_after self.routing_strategy = self._normalize_strategy(routing_strategy) self._routing_groups_input: list[RoutingGroup | dict] | None = routing_groups ## SETTING FALLBACKS ## ### validate if it's set + in correct format _fallbacks = fallbacks or litellm.fallbacks self.validate_fallbacks(fallback_param=_fallbacks) ### set fallbacks self.fallbacks = _fallbacks if default_fallbacks is not None or litellm.default_fallbacks is not None: _fallbacks = default_fallbacks or litellm.default_fallbacks if self.fallbacks is not None: self.fallbacks.append({"*": _fallbacks}) else: self.fallbacks = [{"*": _fallbacks}] self.context_window_fallbacks = context_window_fallbacks or litellm.context_window_fallbacks _content_policy_fallbacks: Final = content_policy_fallbacks or litellm.content_policy_fallbacks self.validate_fallbacks(fallback_param=_content_policy_fallbacks) self.content_policy_fallbacks = _content_policy_fallbacks self.total_calls: defaultdict = defaultdict(int) # dict to store total calls made to each model self.fail_calls: defaultdict = defaultdict(int) # dict to store fail_calls made to each model self.success_calls: defaultdict = defaultdict(int) # dict to store success_calls made to each model # make Router.chat.completions.create compatible for openai.chat.completions.create default_litellm_params = default_litellm_params or {} self.chat = litellm.Chat(params=default_litellm_params, router_obj=self) # default litellm args self.default_litellm_params = default_litellm_params self.default_litellm_params.setdefault("timeout", timeout) self.default_litellm_params.setdefault("max_retries", 0) self.default_litellm_params.setdefault("metadata", {}).update({"caching_groups": caching_groups}) self.deployment_stats: dict = {} # used for debugging load balancing """ deployment_stats = { "122999-2828282-277: { "model": "gpt-3", "api_base": "http://localhost:4000", "num_requests": 20, "avg_latency": 0.001, "num_failures": 0, "num_successes": 20 } } """ ### ROUTING SETUP ### if self._normalize_strategy(routing_strategy) == "lar1": from litellm.router_strategy.lar1_routing import apply_lar1_routing_strategy apply_lar1_routing_strategy(self, routing_strategy_args) else: self.routing_strategy_init( routing_strategy=routing_strategy, routing_strategy_args=routing_strategy_args, ) self._init_routing_groups(self._routing_groups_input) self._override_selectors: dict[str, RouterStrategySelector | None] = {} self._override_selectors_lock = threading.Lock() self.access_groups = None ## USAGE TRACKING ## if isinstance(litellm._async_success_callback, list): litellm.logging_callback_manager.add_litellm_async_success_callback(self.deployment_callback_on_success) else: litellm.logging_callback_manager.add_litellm_async_success_callback(self.deployment_callback_on_success) if isinstance(litellm.success_callback, list): litellm.logging_callback_manager.add_litellm_success_callback(self.sync_deployment_callback_on_success) else: litellm.success_callback = [self.sync_deployment_callback_on_success] if isinstance(litellm._async_failure_callback, list): litellm.logging_callback_manager.add_litellm_async_failure_callback( self.async_deployment_callback_on_failure ) else: litellm._async_failure_callback = [self.async_deployment_callback_on_failure] ## COOLDOWNS ## if isinstance(litellm.failure_callback, list): litellm.logging_callback_manager.add_litellm_failure_callback(self.deployment_callback_on_failure) else: litellm.failure_callback = [self.deployment_callback_on_failure] self.routing_strategy_args = routing_strategy_args self.provider_budget_config = provider_budget_config self.router_budget_logger: RouterBudgetLimiting | None = None if RouterBudgetLimiting.should_init_router_budget_limiter( model_list=model_list, provider_budget_config=self.provider_budget_config ): if optional_pre_call_checks is not None: optional_pre_call_checks.append("router_budget_limiting") else: optional_pre_call_checks = ["router_budget_limiting"] self.retry_policy: RetryPolicy | None = None if retry_policy is not None: if isinstance(retry_policy, dict): self.retry_policy = RetryPolicy(**retry_policy) elif isinstance(retry_policy, RetryPolicy): self.retry_policy = retry_policy if self.retry_policy is not None: verbose_router_logger.info( "\x1b[32mRouter Custom Retry Policy Set:\n%s\x1b[0m", self.retry_policy.model_dump(exclude_none=True), ) self.model_group_retry_policy: dict[str, RetryPolicy] | None = model_group_retry_policy self.allowed_fails_policy: AllowedFailsPolicy | None = None if allowed_fails_policy is not None: if isinstance(allowed_fails_policy, dict): self.allowed_fails_policy = AllowedFailsPolicy(**allowed_fails_policy) elif isinstance(allowed_fails_policy, AllowedFailsPolicy): self.allowed_fails_policy = allowed_fails_policy if self.allowed_fails_policy is not None: verbose_router_logger.info( "\x1b[32mRouter Custom Allowed Fails Policy Set:\n%s\x1b[0m", self.allowed_fails_policy.model_dump(exclude_none=True), ) self.alerting_config: AlertingConfig | None = alerting_config if optional_pre_call_checks is not None: self.add_optional_pre_call_checks(optional_pre_call_checks) # If model_group_affinity_config is set but no global affinity checks were # enabled, we still need the DeploymentAffinityCheck callback (with global # flags all False) so per-group config can activate affinity per model group. if self.model_group_affinity_config: self._ensure_deployment_affinity_callback() if self.alerting_config is not None: self._initialize_alerting() self.initialize_assistants_endpoint() self.initialize_router_endpoints() self.apply_default_settings() @staticmethod def get_valid_args() -> list[str]: """ Returns a list of valid arguments for the Router.__init__ method. """ arg_spec: Final = inspect.getfullargspec(Router.__init__) valid_args: Final = arg_spec.args + arg_spec.kwonlyargs if "self" in valid_args: valid_args.remove("self") return valid_args def apply_default_settings(self): """ Apply the default settings to the router. """ default_pre_call_checks: Final[OptionalPreCallChecks] = [] self.add_optional_pre_call_checks(default_pre_call_checks) def discard(self): """ Pseudo-destructor to be invoked to clean up global data structures when router is no longer used. For now, unhook router's callbacks from all lists """ # Stop contributing to cost-map rebuilds straight away rather than waiting # for this router to be collected. _live_routers.discard(self) litellm.logging_callback_manager.remove_callback_from_list_by_object(litellm._async_success_callback, self) litellm.logging_callback_manager.remove_callback_from_list_by_object(litellm.success_callback, self) litellm.logging_callback_manager.remove_callback_from_list_by_object(litellm._async_failure_callback, self) litellm.logging_callback_manager.remove_callback_from_list_by_object(litellm.failure_callback, self) litellm.logging_callback_manager.remove_callback_from_list_by_object(litellm.input_callback, self) litellm.logging_callback_manager.remove_callback_from_list_by_object(litellm.service_callback, self) litellm.logging_callback_manager.remove_callback_from_list_by_object(litellm.callbacks, self) # Remove ForwardClientSideHeadersByModelGroup if it exists if self.optional_callbacks is not None: for callback in self.optional_callbacks: litellm.logging_callback_manager.remove_callback_from_list_by_object( litellm.callbacks, callback, require_self=False ) @staticmethod def _create_redis_cache( cache_config: dict[str, Any], ) -> RedisCache | RedisClusterCache: """ Initializes either a RedisCache or RedisClusterCache based on the cache_config. """ startup_nodes = cache_config.get("startup_nodes") if not startup_nodes: _env_cluster_nodes: Final = get_secret("REDIS_CLUSTER_NODES") if _env_cluster_nodes is not None and isinstance(_env_cluster_nodes, str): startup_nodes = json.loads(_env_cluster_nodes) if startup_nodes: return RedisClusterCache(**{**cache_config, "startup_nodes": startup_nodes}) else: return RedisCache(**cache_config) def _update_redis_cache(self, cache: RedisCache): """ Update the redis cache for the router, if none set. Allows proxy user to just do ```yaml litellm_settings: cache: true ``` and caching to just work. """ self.cache.attach_redis_cache(cache) self._claude_code_session_router_cache.attach_redis_cache(cache) # Maps a routing strategy string to the attribute on `self` that holds # the default group's strategy selector for that strategy. (The selectors # double as `CustomLogger` callbacks, hence the legacy `*_logger` attrs.) _DEFAULT_SELECTOR_ATTR_BY_STRATEGY: dict[str, str] = { "least-busy": "leastbusy_logger", "usage-based-routing": "lowesttpm_logger", "usage-based-routing-v2": "lowesttpm_logger_v2", "latency-based-routing": "lowestlatency_logger", "cost-based-routing": "lowestcost_logger", } @staticmethod def _normalize_strategy( strategy: RoutingStrategy | str | None, ) -> str | None: if strategy is None: return None if isinstance(strategy, RoutingStrategy): return strategy.value return strategy @staticmethod def _validate_routing_strategy(routing_strategy: RoutingStrategy | str | None) -> None: validate_routing_strategy(routing_strategy) def _build_strategy_selector( self, strategy: RoutingStrategy | str, routing_strategy_args: dict, register_callbacks: bool = True, ) -> RouterStrategySelector | None: """ Constructs a strategy selector for a given strategy. Returns None for `simple-shuffle` (no selector needed) and unknown strategies. """ selector: RouterStrategySelector | None = None match self._normalize_strategy(strategy): case RoutingStrategy.LEAST_BUSY.value: selector = LeastBusyLoggingHandler(router_cache=self.cache) case RoutingStrategy.USAGE_BASED_ROUTING.value: selector = LowestTPMLoggingHandler( router_cache=self.cache, routing_args=routing_strategy_args, ) case RoutingStrategy.USAGE_BASED_ROUTING_V2.value: selector = LowestTPMLoggingHandler_v2( router_cache=self.cache, routing_args=routing_strategy_args, ) case RoutingStrategy.LATENCY_BASED.value: selector = LowestLatencyLoggingHandler( router_cache=self.cache, routing_args=routing_strategy_args, ) case RoutingStrategy.COST_BASED.value: selector = LowestCostLoggingHandler( router_cache=self.cache, routing_args={}, ) case _: pass if selector is not None and register_callbacks: self._register_router_selector(selector) return selector @staticmethod def _register_router_selector(selector: RouterStrategySelector) -> None: if isinstance(selector, LeastBusyLoggingHandler): if isinstance(litellm.input_callback, list): litellm.logging_callback_manager.add_litellm_input_callback(selector) else: litellm.input_callback = [selector] if isinstance(litellm.callbacks, list): litellm.logging_callback_manager.add_litellm_callback(selector) def _unregister_router_selectors(self, selectors: Sequence[object]) -> None: """ Drop router-owned strategy selectors from litellm's global callback lists by identity. Used before re-init (`routing_strategy_init` / `_init_routing_groups`) so repeated `update_settings` calls don't accumulate dead selectors that keep receiving callback events. """ for selector in selectors: if isinstance(selector, BaseRoutingStrategy): selector.retire() selector_ids: Final = {id(s) for s in selectors if s is not None} if not selector_ids: return if isinstance(litellm.callbacks, list): litellm.callbacks = [c for c in litellm.callbacks if id(c) not in selector_ids] if isinstance(litellm.input_callback, list): litellm.input_callback = [c for c in litellm.input_callback if id(c) not in selector_ids] def _apply_updated_routing_strategy_args(self) -> None: """ Re-link the default group's selector to the current `routing_strategy_args`. Selectors freeze their `RoutingArgs` at construction, so a runtime args update would otherwise keep serving the boot-time values until restart. Latency/usage state survives the rebuild: it lives in the shared router cache, not on the selector. """ strategy: Final = self._normalize_strategy(self.routing_strategy) if strategy == "lar1": from litellm.router_strategy.lar1_routing import apply_lar1_routing_strategy apply_lar1_routing_strategy(self, self.routing_strategy_args) return attr: Final = self._DEFAULT_SELECTOR_ATTR_BY_STRATEGY.get(strategy or "") current: Final = getattr(self, attr, None) if attr is not None else None if attr is None or current is None: return try: rebuilt: Final = self._build_strategy_selector( strategy=strategy or "", routing_strategy_args=self.routing_strategy_args, ) except (TypeError, ValidationError): verbose_router_logger.exception( "Invalid routing_strategy_args %s for '%s'; keeping the previous ones", self.routing_strategy_args, strategy, ) return self._unregister_router_selectors((current,)) setattr(self, attr, rebuilt) def routing_strategy_init(self, routing_strategy: RoutingStrategy | str, routing_strategy_args: dict): verbose_router_logger.info("Routing strategy: %s", routing_strategy) self._validate_routing_strategy(routing_strategy) self._reset_custom_routing_strategy() self._unregister_router_selectors( [getattr(self, attr, None) for attr in self._DEFAULT_SELECTOR_ATTR_BY_STRATEGY.values()] + list(getattr(self, "_override_selectors", {}).values()) ) self._override_selectors = {} self.leastbusy_logger: LeastBusyLoggingHandler | None = None self.lowesttpm_logger: LowestTPMLoggingHandler | None = None self.lowesttpm_logger_v2: LowestTPMLoggingHandler_v2 | None = None self.lowestlatency_logger: LowestLatencyLoggingHandler | None = None self.lowestcost_logger: LowestCostLoggingHandler | None = None selector: Final = self._build_strategy_selector( strategy=routing_strategy, routing_strategy_args=routing_strategy_args, ) attr: Final = self._DEFAULT_SELECTOR_ATTR_BY_STRATEGY.get(self._normalize_strategy(routing_strategy) or "") # TODO: legacy `self._logger` attributes are read directly by # `get_settings()` and external callers. Fold the default group into # `self._group_selectors["default"]` and drop these attribute writes — # the dual storage is an antipattern preserved only for back-compat. if attr is not None: setattr(self, attr, selector) def _init_routing_groups( self, groups_input: list[RoutingGroup | dict] | None, ) -> None: """ Validates and indexes `routing_groups`. Each `model_name` may belong to at most one non-priority group. Constructs per-group strategy selectors so groups with different `routing_strategy_args` track independent state. Models not claimed by a non-priority group are served by the implicit `"default"` group, whose selectors are the `self._logger` attributes set up in `routing_strategy_init`. """ if not groups_input: self._replace_routing_groups(()) return known_model_names: Final = frozenset(m["model_name"] for m in (self.model_list or ()) if m.get("model_name")) groups: Final = parse_routing_groups(groups_input, known_model_names=known_model_names) alias_names: Final = frozenset(self.model_group_alias or ()) for group in groups: if group.group_name in known_model_names or group.group_name in alias_names: if group.routing_strategy == "priority": raise ValueError("Priority routing group names must not shadow a model or alias") verbose_router_logger.warning( "routing_groups: group_name '%s' is shadowed by an existing model_name or model_group_alias; " "the group's strategy still applies to its members, but the name is not callable until renamed.", group.group_name, ) built: Final = tuple( ( group, self._build_strategy_selector( strategy="simple-shuffle" if group.routing_strategy == "priority" else group.routing_strategy, routing_strategy_args=group.routing_strategy_args or {}, register_callbacks=False, ), ) for group in groups ) self._replace_routing_groups(built) def _replace_routing_groups( self, built: tuple[tuple[RoutingGroup, RouterStrategySelector | None], ...], ) -> None: previous_selectors: Final[Mapping[str, Mapping[str, RouterStrategySelector]]] = getattr( self, "_group_selectors", {} ) self._unregister_router_selectors( tuple(sel for selectors in previous_selectors.values() for sel in selectors.values()) ) for _, selector in built: if selector is not None: self._register_router_selector(selector) self._routing_groups: dict[str, RoutingGroup] = {group.group_name: group for group, _ in built} self._model_to_group: dict[str, str] = { model_name: group.group_name for group, _ in built if group.routing_strategy != "priority" for model_name in group.models } self._group_selectors: dict[str, dict[str, RouterStrategySelector]] = { group.group_name: ( {} if selector is None else {self._normalize_strategy(group.routing_strategy) or "": selector} ) for group, selector in built } self._invalidate_model_group_info_cache() self._invalidate_access_groups_cache() def get_routing_group(self, model_name: str) -> RoutingGroup | None: """ The routing group callable as `model_name`, or None. A real deployment `model_name` added after init shadows a same-named group (mirroring `_try_early_resolve_deployments_for_model_not_in_names`, where concrete models win over indirection); config-time collisions are rejected by `_init_routing_groups`. """ if not self._routing_groups: return None group: Final = self._routing_groups.get(model_name) if ( group is None or model_name in self.model_name_to_deployment_indices or model_name in (self.model_group_alias or {}) ): return None return group def _get_routing_group_deployments( self, model: str, team_id: str | None = None ) -> list[DeploymentTypedDict] | None: # mutable-ok: list matches _get_all_deployments' contract for callers """ The union of member deployments for a routing group called as `model`, or None when `model` is not a callable group. The requested name stays the group name so strategy selectors key their state by it. `_common_checks_available_deployment` consults this BEFORE its early-resolve step so a wildcard `default_deployment` or pattern route cannot hijack a group call. Overall resolution precedence there: specific deployment > model id > model_group_alias > routing group > model_name > team/pattern/default fallbacks. """ if not self._routing_groups: return None routing_group: Final = self.get_routing_group(model) if routing_group is None: return None return [ # mutable-ok: matches _get_all_deployments' list contract expected by downstream filters apply_routing_group_priority(routing_group, member, deployment) for member in routing_group.models for deployment in self._get_all_deployments(model_name=member, team_id=team_id) ] def _is_priority_routing_group(self, model: str) -> bool: resolved: Final = self._get_model_from_alias(model=model) or model group: Final = self.get_routing_group(resolved) return group is not None and group.routing_strategy == "priority" def is_recognized_model(self, model: str) -> bool: """ Whether `model` names something this router serves directly: a deployment model_name, a deployment id, a `model_group_alias`, or a callable routing group. Proxy request gates share this predicate so a new virtual-model kind cannot be forgotten at one of them; wildcard, default-deployment, and deployment-name fallbacks stay caller policy. """ return ( model in self.model_names or self.has_model_id(model) or (self.model_group_alias is not None and model in self.model_group_alias) or self.get_routing_group(model) is not None ) def routing_group_has_alternatives(self, model_group: str | None) -> bool: """ True when `model_group` names a callable routing group whose member union spans more than one deployment. Cooldown handling passes the FAILING REQUEST's model group here: a 429 on a group call cools the member down so selection moves to the group's alternatives, while a direct call to a single-deployment member keeps the single-deployment-model-group cooldown exemption. """ if model_group is None: return False resolved: Final = self._get_model_from_alias(model=model_group) or model_group group: Final = self.get_routing_group(resolved) if group is None: return False return sum(len(self.model_name_to_deployment_indices.get(member) or ()) for member in group.models) > 1 def team_model_has_alternatives(self, deployment_id: str) -> bool: deployment: Final = self.get_deployment(model_id=deployment_id) if deployment is None: return False team_id: Final = deployment.model_info.team_id public_model_name: Final = deployment.model_info.team_public_model_name if team_id is None or public_model_name is None: return False sibling_indices: Final = self.team_model_to_deployment_indices.get((team_id, public_model_name)) or () routable_siblings: Final = self._filter_blocked_deployments([self.model_list[idx] for idx in sibling_indices]) return len(routable_siblings) > 1 _OVERRIDABLE_ROUTING_STRATEGIES: frozenset[str] = frozenset({"simple-shuffle", *_DEFAULT_SELECTOR_ATTR_BY_STRATEGY}) def _get_request_routing_strategy_override(self, request_kwargs: dict | None) -> str | None: """ Reads a per-request `routing_strategy` override (forwarded by the proxy from key/team `router_settings`) out of the request kwargs. Only strategies with a per-request-capable selector are honored; anything else (unknown strings, `lar1`, `provider-budget-routing`) is ignored with a warning so a bad value stored on a key or team can never take down that caller's traffic. """ if not request_kwargs: return None raw_strategy: Final = request_kwargs.get("routing_strategy") if raw_strategy is None: return None strategy = self._normalize_strategy(raw_strategy) if isinstance(raw_strategy, (str, RoutingStrategy)) else None if not isinstance(strategy, str) or strategy not in self._OVERRIDABLE_ROUTING_STRATEGIES: verbose_router_logger.warning( "Ignoring per-request routing_strategy override '%s'; supported overrides: %s.", raw_strategy, sorted(self._OVERRIDABLE_ROUTING_STRATEGIES), ) return None return strategy def _get_override_strategy_selector(self, strategy: str) -> RouterStrategySelector | None: """ Returns the selector for a per-request strategy override. Reuses the default group's selector when the override matches the router's configured strategy (so shared state keeps accumulating in one place); otherwise lazily builds one selector per strategy and caches it for the router's lifetime so its usage/latency state persists across requests. """ if strategy == self._normalize_strategy(self.routing_strategy): attr: Final = self._DEFAULT_SELECTOR_ATTR_BY_STRATEGY.get(strategy) return getattr(self, attr, None) if attr is not None else None with self._override_selectors_lock: if strategy not in self._override_selectors: self._override_selectors[strategy] = self._build_strategy_selector( strategy=strategy, routing_strategy_args={}, register_callbacks=False, ) return self._override_selectors[strategy] def _override_selector_pre_call_check( self, strategy: str | None, selector: RouterStrategySelector | None, deployment: dict ) -> None: """ Override selectors are not in `litellm.callbacks`, so the pre-call check that `routing_strategy_pre_call_checks` runs for the router's own selectors (rpm accounting for `usage-based-routing-v2`) runs here, for the overriding request only. """ if selector is None or strategy is None or selector is not self._override_selectors.get(strategy): return selector.pre_call_check(deployment) async def _async_override_selector_pre_call_check( self, strategy: str | None, selector: RouterStrategySelector | None, deployment: dict, parent_otel_span: Span | None, ) -> None: if selector is None or strategy is None or selector is not self._override_selectors.get(strategy): return await selector.async_pre_call_check(deployment, parent_otel_span) def _bind_override_selector_to_request( self, strategy: str, selector: RouterStrategySelector | None, request_kwargs: Mapping[str, object] | None ) -> None: if selector is None or request_kwargs is None or strategy in self._globally_registered_strategies(): return logging_obj: Final = request_kwargs.get("litellm_logging_obj") if isinstance(logging_obj, LiteLLMLogging): logging_obj.add_dynamic_callback(selector) def _globally_registered_strategies(self) -> frozenset[str]: configured: Final = ( self.routing_strategy, *(group.routing_strategy for group in self._routing_groups.values()), ) return frozenset( normalized for normalized in map(self._normalize_strategy, configured) if normalized is not None ) def arm_routing_read_prefetch(self, model: str, request_kwargs: dict[str, object] | None = None) -> None: """Declare the cooldown read (and, for usage-based routing, the usage read) that `async_get_available_deployment` will make for `model` on the request's Redis batch, so admission's flush carries it. A miss (alias, no batch) costs nothing: routing then reads as it always has.""" try: strategy, selector = self._get_routing_context(model, request_kwargs) usage_selector: Final = ( selector if strategy == "usage-based-routing-v2" and isinstance(selector, LowestTPMLoggingHandler_v2) else None ) deployments: Final = self.get_model_list(model_name=model) if deployments: RoutingPrefetch.arm(self, usage_selector, deployments) except Exception as e: # noqa: BLE001 # a prefetch is an optimisation, never a reason to fail the request verbose_router_logger.debug( "routing read prefetch not armed for %s: %s", _without_line_breaks(model), _without_line_breaks(e) ) def _get_routing_context( self, model: str, request_kwargs: dict | None = None ) -> tuple[str | None, RouterStrategySelector | None]: """ Resolves the routing strategy and selector to use for the given model. A per-request `routing_strategy` in `request_kwargs` (forwarded by the proxy from key/team `router_settings`) takes precedence over both the model's routing group and the router's top-level strategy, since it is the most specific expression of caller intent. Otherwise every model belongs to exactly one group: an explicit entry from `routing_groups` (either because `model` IS a callable group name, or because it is a member of one), or the implicit `"default"` group driven by the router's top-level `routing_strategy` / `routing_strategy_args`. `self.routing_strategy` may be either a string or a `RoutingStrategy` enum member (the constructor accepts both), so it is normalized to a string here. Downstream call sites and `_select_deployment_*` arms compare against string literals. """ override: Final = self._get_request_routing_strategy_override(request_kwargs) if override is not None: verbose_router_logger.debug("routing_group=request-override model=%s strategy=%s", model, override) override_selector: Final = self._get_override_strategy_selector(override) self._bind_override_selector_to_request(override, override_selector, request_kwargs) return override, override_selector resolved_model: Final = self._get_model_from_alias(model=model) or model group_name: Final = ( resolved_model if self.get_routing_group(resolved_model) is not None else self._model_to_group.get(resolved_model) ) if group_name is None: strategy = self._normalize_strategy(self.routing_strategy) attr: Final = self._DEFAULT_SELECTOR_ATTR_BY_STRATEGY.get(strategy or "") selector = getattr(self, attr, None) if attr is not None else None verbose_router_logger.debug("routing_group=default model=%s strategy=%s", model, strategy) return strategy, selector group: Final = self._routing_groups[group_name] if group.routing_strategy == "priority": return "simple-shuffle", None strategy = self._normalize_strategy(group.routing_strategy) selector = self._group_selectors.get(group_name, {}).get(strategy or "") verbose_router_logger.debug("routing_group=%s model=%s strategy=%s", group_name, model, strategy) return strategy, selector async def _select_deployment_async( self, *, strategy: str | None, selector: Any | None, model: str, healthy_deployments: list, messages: list[dict[str, str]] | None, input: str | list | None, request_kwargs: dict | None, ) -> Any | None: """ Asks the strategy selector for a deployment. Caller handles `simple-shuffle` separately (it does not flow through a selector). Returns None for unknown strategies or when the selector is missing. """ if selector is None: return None match strategy: case "least-busy": return await selector.async_get_available_deployments( model_group=model, healthy_deployments=healthy_deployments, ) case "usage-based-routing": # `LowestTPMLoggingHandler` (v1) only exposes the sync # `get_available_deployments`. Mirror the pre-routing-groups # top-level fallback by calling it inline so groups using v1 # still work from async callers. return selector.get_available_deployments( model_group=model, healthy_deployments=healthy_deployments, messages=messages, input=input, ) case "usage-based-routing-v2" | "cost-based-routing": return await selector.async_get_available_deployments( model_group=model, healthy_deployments=healthy_deployments, messages=messages, input=input, ) case "latency-based-routing": return await selector.async_get_available_deployments( model_group=model, healthy_deployments=healthy_deployments, messages=messages, input=input, request_kwargs=request_kwargs, ) case _: return None def _select_deployment_sync( self, *, strategy: str | None, selector: Any | None, model: str, healthy_deployments: list, messages: list[dict[str, str]] | None, input: str | list | None, request_kwargs: dict | None, ) -> Any | None: """ Sync sibling of `_select_deployment_async`. Caller handles `simple-shuffle` separately. """ if selector is None: return None # `cost-based-routing` is intentionally omitted — # `LowestCostLoggingHandler` only implements # `async_get_available_deployments` match strategy: case "least-busy": return selector.get_available_deployments( model_group=model, healthy_deployments=healthy_deployments, ) case "usage-based-routing" | "usage-based-routing-v2": return selector.get_available_deployments( model_group=model, healthy_deployments=healthy_deployments, messages=messages, input=input, ) case "latency-based-routing": return selector.get_available_deployments( model_group=model, healthy_deployments=healthy_deployments, messages=messages, input=input, request_kwargs=request_kwargs, ) case _: return None def initialize_assistants_endpoint(self): ## INITIALIZE PASS THROUGH ASSISTANTS ENDPOINT ## self.acreate_assistants = self.factory_function(litellm.acreate_assistants) self.adelete_assistant = self.factory_function(litellm.adelete_assistant) self.aget_assistants = self.factory_function(litellm.aget_assistants) self.acreate_thread = self.factory_function(litellm.acreate_thread) self.aget_thread = self.factory_function(litellm.aget_thread) self.a_add_message = self.factory_function(litellm.a_add_message) self.aget_messages = self.factory_function(litellm.aget_messages) self.arun_thread = self.factory_function(litellm.arun_thread) def _initialize_core_endpoints(self): """Helper to initialize core router endpoints.""" self.amoderation = self.factory_function(litellm.amoderation, call_type="moderation") self.aanthropic_messages = self.factory_function(litellm.anthropic_messages, call_type="anthropic_messages") self.anthropic_messages = self.factory_function(litellm.anthropic_messages, call_type="anthropic_messages") self.agenerate_content = self.factory_function(litellm.agenerate_content, call_type="agenerate_content") self.aadapter_generate_content = self.factory_function( litellm.aadapter_generate_content, call_type="aadapter_generate_content" ) self.aresponses = self.factory_function(litellm.aresponses, call_type="aresponses") self.afile_delete = self.factory_function(litellm.afile_delete, call_type="afile_delete") self.afile_content = self.factory_function(litellm.afile_content, call_type="afile_content") self.responses = self.factory_function(litellm.responses, call_type="responses") self.aget_responses = self.factory_function(litellm.aget_responses, call_type="aget_responses") self.acancel_responses = self.factory_function(litellm.acancel_responses, call_type="acancel_responses") self.acompact_responses = self.factory_function(litellm.acompact_responses, call_type="acompact_responses") self.adelete_responses = self.factory_function(litellm.adelete_responses, call_type="adelete_responses") self.alist_input_items = self.factory_function(litellm.alist_input_items, call_type="alist_input_items") self._arealtime = self.factory_function(litellm._arealtime, call_type="_arealtime") self.acreate_realtime_client_secret = self.factory_function( litellm.acreate_realtime_client_secret, call_type="acreate_realtime_client_secret" ) self.arealtime_calls = self.factory_function(litellm.arealtime_calls, call_type="arealtime_calls") self.acreate_realtime_transcription_session = self.factory_function( litellm.acreate_realtime_transcription_session, call_type="acreate_realtime_transcription_session" ) self._aresponses_websocket = self.factory_function( litellm._aresponses_websocket, call_type="_aresponses_websocket" ) self.acreate_fine_tuning_job = self.factory_function( litellm.acreate_fine_tuning_job, call_type="acreate_fine_tuning_job" ) self.acancel_fine_tuning_job = self.factory_function( litellm.acancel_fine_tuning_job, call_type="acancel_fine_tuning_job" ) self.alist_fine_tuning_jobs = self.factory_function( litellm.alist_fine_tuning_jobs, call_type="alist_fine_tuning_jobs" ) self.aretrieve_fine_tuning_job = self.factory_function( litellm.aretrieve_fine_tuning_job, call_type="aretrieve_fine_tuning_job" ) self.afile_list = self.factory_function(litellm.afile_list, call_type="alist_files") self.aimage_edit = self.factory_function(litellm.aimage_edit, call_type="aimage_edit") self.allm_passthrough_route = self.factory_function( litellm.allm_passthrough_route, call_type="allm_passthrough_route" ) # Note: acancel_batch is defined as a method on the Router class (not using factory_function) # to properly handle model-to-provider mapping like acreate_batch and aretrieve_batch def _initialize_vector_store_endpoints(self): """Initialize vector store endpoints.""" from litellm.vector_stores.main import ( adelete, alist, aretrieve, asearch, aupdate, create, delete, list, retrieve, search, update, ) self.avector_store_search = self.factory_function(asearch, call_type="avector_store_search") self.vector_store_search = self.factory_function(search, call_type="vector_store_search") self.vector_store_create = self.factory_function(create, call_type="vector_store_create") self.avector_store_retrieve = self.factory_function(aretrieve, call_type="avector_store_retrieve") self.vector_store_retrieve = self.factory_function(retrieve, call_type="vector_store_retrieve") self.avector_store_list = self.factory_function(alist, call_type="avector_store_list") self.vector_store_list = self.factory_function(list, call_type="vector_store_list") self.avector_store_update = self.factory_function(aupdate, call_type="avector_store_update") self.vector_store_update = self.factory_function(update, call_type="vector_store_update") self.avector_store_delete = self.factory_function(adelete, call_type="avector_store_delete") self.vector_store_delete = self.factory_function(delete, call_type="vector_store_delete") def _initialize_vector_store_file_endpoints(self): """Initialize vector store file endpoints.""" from litellm.vector_store_files.main import ( acreate as avector_store_file_create_fn, ) from litellm.vector_store_files.main import ( adelete as avector_store_file_delete_fn, ) from litellm.vector_store_files.main import alist as avector_store_file_list_fn from litellm.vector_store_files.main import ( aretrieve as avector_store_file_retrieve_fn, ) from litellm.vector_store_files.main import ( aretrieve_content as avector_store_file_content_fn, ) from litellm.vector_store_files.main import ( aupdate as avector_store_file_update_fn, ) from litellm.vector_store_files.main import ( create as vector_store_file_create_fn, ) from litellm.vector_store_files.main import ( delete as vector_store_file_delete_fn, ) from litellm.vector_store_files.main import list as vector_store_file_list_fn from litellm.vector_store_files.main import ( retrieve as vector_store_file_retrieve_fn, ) from litellm.vector_store_files.main import ( retrieve_content as vector_store_file_content_fn, ) from litellm.vector_store_files.main import ( update as vector_store_file_update_fn, ) self.avector_store_file_create = self.factory_function( avector_store_file_create_fn, call_type="avector_store_file_create" ) self.vector_store_file_create = self.factory_function( vector_store_file_create_fn, call_type="vector_store_file_create" ) self.avector_store_file_list = self.factory_function( avector_store_file_list_fn, call_type="avector_store_file_list" ) self.vector_store_file_list = self.factory_function( vector_store_file_list_fn, call_type="vector_store_file_list" ) self.avector_store_file_retrieve = self.factory_function( avector_store_file_retrieve_fn, call_type="avector_store_file_retrieve" ) self.vector_store_file_retrieve = self.factory_function( vector_store_file_retrieve_fn, call_type="vector_store_file_retrieve" ) self.avector_store_file_content = self.factory_function( avector_store_file_content_fn, call_type="avector_store_file_content" ) self.vector_store_file_content = self.factory_function( vector_store_file_content_fn, call_type="vector_store_file_content" ) self.avector_store_file_update = self.factory_function( avector_store_file_update_fn, call_type="avector_store_file_update" ) self.vector_store_file_update = self.factory_function( vector_store_file_update_fn, call_type="vector_store_file_update" ) self.avector_store_file_delete = self.factory_function( avector_store_file_delete_fn, call_type="avector_store_file_delete" ) self.vector_store_file_delete = self.factory_function( vector_store_file_delete_fn, call_type="vector_store_file_delete" ) def _initialize_google_genai_endpoints(self): """Initialize Google GenAI endpoints.""" from litellm.google_genai import ( agenerate_content, agenerate_content_stream, generate_content, generate_content_stream, ) self.agenerate_content = self.factory_function(agenerate_content, call_type="agenerate_content") self.generate_content = self.factory_function(generate_content, call_type="generate_content") self.agenerate_content_stream = self.factory_function( agenerate_content_stream, call_type="agenerate_content_stream" ) self.generate_content_stream = self.factory_function( generate_content_stream, call_type="generate_content_stream" ) def _initialize_ocr_search_endpoints(self): """Initialize OCR and search endpoints.""" from litellm.ocr import aocr, ocr self.aocr = self.factory_function(aocr, call_type="aocr") self.ocr = self.factory_function(ocr, call_type="ocr") from litellm.search import asearch, search self.asearch = self.factory_function(asearch, call_type="asearch") self.search = self.factory_function(search, call_type="search") def _initialize_video_endpoints(self): """Initialize video endpoints.""" from litellm.videos import ( avideo_content, avideo_create_character, avideo_edit, avideo_extension, avideo_generation, avideo_get_character, avideo_list, avideo_remix, avideo_status, video_content, video_create_character, video_edit, video_extension, video_generation, video_get_character, video_list, video_remix, video_status, ) self.avideo_generation = self.factory_function(avideo_generation, call_type="avideo_generation") self.video_generation = self.factory_function(video_generation, call_type="video_generation") self.avideo_list = self.factory_function(avideo_list, call_type="avideo_list") self.video_list = self.factory_function(video_list, call_type="video_list") self.avideo_status = self.factory_function(avideo_status, call_type="avideo_status") self.video_status = self.factory_function(video_status, call_type="video_status") self.avideo_content = self.factory_function(avideo_content, call_type="avideo_content") self.video_content = self.factory_function(video_content, call_type="video_content") self.avideo_remix = self.factory_function(avideo_remix, call_type="avideo_remix") self.video_remix = self.factory_function(video_remix, call_type="video_remix") self.avideo_create_character = self.factory_function( avideo_create_character, call_type="avideo_create_character" ) self.video_create_character = self.factory_function(video_create_character, call_type="video_create_character") self.avideo_get_character = self.factory_function(avideo_get_character, call_type="avideo_get_character") self.video_get_character = self.factory_function(video_get_character, call_type="video_get_character") self.avideo_edit = self.factory_function(avideo_edit, call_type="avideo_edit") self.video_edit = self.factory_function(video_edit, call_type="video_edit") self.avideo_extension = self.factory_function(avideo_extension, call_type="avideo_extension") self.video_extension = self.factory_function(video_extension, call_type="video_extension") def _initialize_container_endpoints(self): """Initialize container endpoints.""" from litellm.containers import ( acreate_container, adelete_container, alist_containers, aretrieve_container, create_container, delete_container, list_containers, retrieve_container, ) from litellm.containers.endpoint_factory import ( _generated_endpoints as container_file_endpoints, ) self.acreate_container = self.factory_function(acreate_container, call_type="acreate_container") self.create_container = self.factory_function(create_container, call_type="create_container") self.alist_containers = self.factory_function(alist_containers, call_type="alist_containers") self.list_containers = self.factory_function(list_containers, call_type="list_containers") self.aretrieve_container = self.factory_function(aretrieve_container, call_type="aretrieve_container") self.retrieve_container = self.factory_function(retrieve_container, call_type="retrieve_container") self.adelete_container = self.factory_function(adelete_container, call_type="adelete_container") self.delete_container = self.factory_function(delete_container, call_type="delete_container") # Auto-register JSON-generated container file endpoints for name, func in container_file_endpoints.items(): setattr(self, name, self.factory_function(func, call_type=name)) def _initialize_skills_endpoints(self): """Initialize Anthropic Skills API endpoints.""" self.acreate_skill = self.factory_function(litellm.acreate_skill, call_type="acreate_skill") self.alist_skills = self.factory_function(litellm.alist_skills, call_type="alist_skills") self.aget_skill = self.factory_function(litellm.aget_skill, call_type="aget_skill") self.adelete_skill = self.factory_function(litellm.adelete_skill, call_type="adelete_skill") def _initialize_interactions_endpoints(self): """Initialize Google Interactions API endpoints.""" from litellm.interactions import acancel as acancel_interaction from litellm.interactions import acreate as acreate_interaction from litellm.interactions import adelete as adelete_interaction from litellm.interactions import aget as aget_interaction from litellm.interactions import cancel as cancel_interaction from litellm.interactions import create as create_interaction from litellm.interactions import delete as delete_interaction from litellm.interactions import get as get_interaction self.acreate_interaction = self.factory_function(acreate_interaction, call_type="acreate_interaction") self.create_interaction = self.factory_function(create_interaction, call_type="create_interaction") self.aget_interaction = self.factory_function(aget_interaction, call_type="aget_interaction") self.get_interaction = self.factory_function(get_interaction, call_type="get_interaction") self.adelete_interaction = self.factory_function(adelete_interaction, call_type="adelete_interaction") self.delete_interaction = self.factory_function(delete_interaction, call_type="delete_interaction") self.acancel_interaction = self.factory_function(acancel_interaction, call_type="acancel_interaction") self.cancel_interaction = self.factory_function(cancel_interaction, call_type="cancel_interaction") def _initialize_managed_agents_endpoints(self): """Initialize Google Managed Agents API endpoints (v1beta/agents).""" from litellm.interactions.agents import acreate as acreate_agent from litellm.interactions.agents import adelete as adelete_agent from litellm.interactions.agents import aget as aget_agent from litellm.interactions.agents import alist as alist_agents from litellm.interactions.agents import alist_versions as alist_agent_versions from litellm.interactions.agents import create as create_agent from litellm.interactions.agents import delete as delete_agent from litellm.interactions.agents import get as get_agent from litellm.interactions.agents import list as list_agents from litellm.interactions.agents import list_versions as list_agent_versions self.acreate_agent = self.factory_function(acreate_agent, call_type="acreate_agent") self.create_agent = self.factory_function(create_agent, call_type="create_agent") self.alist_agents = self.factory_function(alist_agents, call_type="alist_agents") self.list_agents = self.factory_function(list_agents, call_type="list_agents") self.aget_agent = self.factory_function(aget_agent, call_type="aget_agent") self.get_agent = self.factory_function(get_agent, call_type="get_agent") self.adelete_agent = self.factory_function(adelete_agent, call_type="adelete_agent") self.delete_agent = self.factory_function(delete_agent, call_type="delete_agent") self.alist_agent_versions = self.factory_function(alist_agent_versions, call_type="alist_agent_versions") self.list_agent_versions = self.factory_function(list_agent_versions, call_type="list_agent_versions") def _initialize_specialized_endpoints(self): """Helper to initialize specialized router endpoints (vector store, OCR, search, video, container, skills, interactions).""" self._initialize_vector_store_endpoints() self._initialize_vector_store_file_endpoints() self._initialize_google_genai_endpoints() self._initialize_ocr_search_endpoints() # Override vector store methods with router-aware implementations self._override_vector_store_methods_for_router() self._initialize_video_endpoints() self._initialize_container_endpoints() self._initialize_skills_endpoints() self._initialize_interactions_endpoints() self._initialize_managed_agents_endpoints() def initialize_router_endpoints(self): self._initialize_core_endpoints() self._initialize_specialized_endpoints() def validate_fallbacks(self, fallback_param: list | None): """ Validate the fallbacks parameter. """ if fallback_param is None: return for fallback_dict in fallback_param: if not isinstance(fallback_dict, dict): raise ValueError(f"Item '{fallback_dict}' is not a dictionary.") if len(fallback_dict) != 1: raise ValueError( f"Dictionary '{fallback_dict}' must have exactly one key, but has {len(fallback_dict)} keys." ) def _add_encrypted_content_affinity_check(self, enable_global_affinity: bool) -> None: def _move_before_deployment_affinity( callback_list: list[_CallbackT], callback_to_move: _CallbackT, ) -> None: if callback_to_move not in callback_list: return callback_list.remove(callback_to_move) insert_index: Final = next( (idx for idx, callback in enumerate(callback_list) if isinstance(callback, DeploymentAffinityCheck)), len(callback_list), ) callback_list.insert(insert_index, callback_to_move) if enable_global_affinity or EncryptedContentAffinityCheck.has_model_group_affinity_enabled( self.model_group_affinity_config ): if self.optional_callbacks is None: self.optional_callbacks = [] existing_ec_callback: EncryptedContentAffinityCheck | None = None for cb in self.optional_callbacks: if isinstance(cb, EncryptedContentAffinityCheck): existing_ec_callback = cb break if existing_ec_callback is not None: existing_ec_callback.router = self existing_ec_callback.enable_global_affinity = ( existing_ec_callback.enable_global_affinity or enable_global_affinity ) existing_ec_callback.model_group_affinity_config = self.model_group_affinity_config or {} ec_callback = existing_ec_callback else: ec_callback = EncryptedContentAffinityCheck( router=self, enable_global_affinity=enable_global_affinity, model_group_affinity_config=self.model_group_affinity_config, ) self.optional_callbacks.append(ec_callback) litellm.logging_callback_manager.add_litellm_callback(ec_callback) _move_before_deployment_affinity(self.optional_callbacks, ec_callback) _move_before_deployment_affinity(litellm.callbacks, ec_callback) def _ensure_deployment_affinity_callback(self) -> None: """Register the DeploymentAffinityCheck callback (global flags all False) if absent. Needed when nothing enabled a global affinity flag but affinity can still activate per request: per-group `model_group_affinity_config` entries, or the session-affinity marker a complexity router stamps at pre-routing time. """ if any(isinstance(cb, DeploymentAffinityCheck) for cb in (self.optional_callbacks or [])): return if self.optional_callbacks is None: self.optional_callbacks = [] affinity_callback: Final = DeploymentAffinityCheck( cache=self.cache, ttl_seconds=self.deployment_affinity_ttl_seconds, enable_user_key_affinity=False, enable_responses_api_affinity=False, enable_session_id_affinity=False, model_group_affinity_config=self.model_group_affinity_config, is_priority_group=self._is_priority_routing_group, ) self.optional_callbacks.append(affinity_callback) litellm.logging_callback_manager.add_litellm_callback(affinity_callback) def add_optional_pre_call_checks(self, optional_pre_call_checks: OptionalPreCallChecks | None): if optional_pre_call_checks is None: return # --------------------------------------------------------------------- # Unified deployment affinity (session stickiness) # --------------------------------------------------------------------- enable_user_key_affinity: Final = "deployment_affinity" in optional_pre_call_checks enable_responses_api_affinity: Final = "responses_api_deployment_check" in optional_pre_call_checks enable_session_id_affinity: Final = "session_affinity" in optional_pre_call_checks if enable_user_key_affinity or enable_responses_api_affinity or enable_session_id_affinity: if self.optional_callbacks is None: self.optional_callbacks = [] existing_affinity_callback: DeploymentAffinityCheck | None = None for cb in self.optional_callbacks: if isinstance(cb, DeploymentAffinityCheck): existing_affinity_callback = cb break if existing_affinity_callback is not None: existing_affinity_callback.enable_user_key_affinity = ( existing_affinity_callback.enable_user_key_affinity or enable_user_key_affinity ) existing_affinity_callback.enable_responses_api_affinity = ( existing_affinity_callback.enable_responses_api_affinity or enable_responses_api_affinity ) existing_affinity_callback.enable_session_id_affinity = ( existing_affinity_callback.enable_session_id_affinity or enable_session_id_affinity ) existing_affinity_callback.ttl_seconds = self.deployment_affinity_ttl_seconds if self.model_group_affinity_config: existing_affinity_callback.model_group_affinity_config = self.model_group_affinity_config else: affinity_callback: Final = DeploymentAffinityCheck( cache=self.cache, ttl_seconds=self.deployment_affinity_ttl_seconds, enable_user_key_affinity=enable_user_key_affinity, enable_responses_api_affinity=enable_responses_api_affinity, enable_session_id_affinity=enable_session_id_affinity, model_group_affinity_config=self.model_group_affinity_config, is_priority_group=self._is_priority_routing_group, ) self.optional_callbacks.append(affinity_callback) litellm.logging_callback_manager.add_litellm_callback(affinity_callback) # --------------------------------------------------------------------- # Encrypted content affinity # --------------------------------------------------------------------- self._add_encrypted_content_affinity_check( enable_global_affinity=("encrypted_content_affinity" in optional_pre_call_checks) ) # --------------------------------------------------------------------- # Remaining optional pre-call checks # --------------------------------------------------------------------- for pre_call_check in optional_pre_call_checks: _callback: CustomLogger | None = None if pre_call_check in ( "deployment_affinity", "responses_api_deployment_check", "session_affinity", "encrypted_content_affinity", ): continue if pre_call_check == "prompt_caching": _callback = PromptCachingDeploymentCheck( cache=self.cache, is_priority_group=self._is_priority_routing_group ) elif pre_call_check == "router_budget_limiting": if self._get_router_deployment_budget_limiter() is not None: continue _callback = RouterBudgetLimiting( dual_cache=self.cache, provider_budget_config=self.provider_budget_config, model_list=self.model_list, ) self.router_budget_logger = _callback elif pre_call_check == "enforce_model_rate_limits": _callback = ModelRateLimitingCheck(dual_cache=self.cache) if _callback is None: continue if self.optional_callbacks is not None and any( isinstance(callback, type(_callback)) for callback in self.optional_callbacks ): continue if self.optional_callbacks is None: self.optional_callbacks = [] self.optional_callbacks.append(_callback) litellm.logging_callback_manager.add_litellm_callback(_callback) def set_optional_pre_call_checks(self, optional_pre_call_checks: OptionalPreCallChecks | None) -> None: if optional_pre_call_checks is None: return requested: Final = frozenset(optional_pre_call_checks) for name, callback_cls in _RUNTIME_TOGGLEABLE_PRE_CALL_CHECKS.items(): if name not in requested: self._remove_optional_callbacks_of_type(callback_cls) self.add_optional_pre_call_checks(optional_pre_call_checks) def _remove_optional_callbacks_of_type(self, callback_cls: type[CustomLogger]) -> None: if self.optional_callbacks is None or not any(type(cb) is callback_cls for cb in self.optional_callbacks): return self.optional_callbacks = [cb for cb in self.optional_callbacks if type(cb) is not callback_cls] if any( router is not self and any(type(cb) is callback_cls for cb in (router.optional_callbacks or [])) for router in tuple(_live_routers) ): return for cb in tuple(litellm.callbacks): if type(cb) is callback_cls: litellm.logging_callback_manager.remove_callback_from_list_by_object( litellm.callbacks, cb, require_self=False ) def print_deployment(self, deployment: dict): """ returns a copy of the deployment with the api key masked Only returns 2 characters of the api key and masks the rest with * (10 *). """ try: _deployment_copy: Final = copy.deepcopy(deployment) litellm_params: Final[dict] = _deployment_copy["litellm_params"] if litellm.redact_user_api_key_info: masker: Final = SensitiveDataMasker(visible_prefix=2, visible_suffix=0) _deployment_copy["litellm_params"] = masker.mask_dict(litellm_params) elif "api_key" in litellm_params: litellm_params["api_key"] = litellm_params["api_key"][:2] + "*" * 10 return _deployment_copy except Exception as e: verbose_router_logger.debug("Error occurred while printing deployment - %s", e) raise e @staticmethod def _deployment_params_with_request_reasoning_override( deployment_params: Mapping[str, object], request_kwargs: Mapping[str, object] ) -> dict[str, object]: # mutable-ok: litellm's request pipeline consumes a mutable kwargs mapping """Return deployment params whose equivalent effort controls cannot outrank a request override. Providers expose the same setting through several native carriers. A request-level ``reasoning_effort`` is the portable override, so a deployment's ``thinking`` or nested ``*.effort`` must not remain beside it and either win or trigger a conflicting-params 400. Every changed mapping is copied so the Router's shared deployment config stays immutable. """ sanitized: Final = dict(deployment_params) # mutable-ok: request-local copy protects shared Router state if request_kwargs.get("reasoning_effort") is None: return sanitized sanitized.pop("thinking", None) Router._pop_effort_from_nested_carrier(sanitized, "output_config") Router._pop_effort_from_nested_carrier(sanitized, "reasoning") extra_body: Final = sanitized.get("extra_body") if isinstance(extra_body, Mapping): sanitized_extra_body: Final = dict(extra_body) # mutable-ok: request-local nested copy sanitized_extra_body.pop("reasoning_effort", None) sanitized_extra_body.pop("thinking", None) Router._pop_effort_from_nested_carrier(sanitized_extra_body, "output_config") Router._pop_effort_from_nested_carrier(sanitized_extra_body, "reasoning") if sanitized_extra_body: sanitized["extra_body"] = sanitized_extra_body else: sanitized.pop("extra_body", None) return sanitized @staticmethod def _is_classifier_internal_call(kwargs: Mapping[str, object]) -> bool: metadata: Final = kwargs.get("metadata") litellm_metadata: Final = kwargs.get("litellm_metadata") return any( isinstance(candidate, Mapping) and candidate.get(INTERNAL_CALL_ORIGIN_METADATA_KEY) == AUTOROUTER_CLASSIFIER_CALL_ORIGIN for candidate in (metadata, litellm_metadata) ) def _drop_unsupported_classifier_reasoning_effort( self, deployment: DeploymentTypedDict, model: str, kwargs: dict[str, object], # mutable-ok: fallback must update the active request and its log body together ) -> None: """Let a classifier fallback without reasoning support remain a usable fallback. The dashboard only offers explicitly advertised levels, but an existing config can outlive a model change and fallbacks can target a different group. Unknown capability fails open; only a provider that explicitly rejects the parameter has it removed. """ if kwargs.get("reasoning_effort") is None or not self._is_classifier_internal_call(kwargs): return if self._deployment_accepts_param(deployment, model, "reasoning_effort"): return verbose_router_logger.warning( "litellm.router.py: dropping classifier reasoning_effort for model=%s because the selected deployment does not support it", model, ) kwargs.pop("reasoning_effort", None) proxy_server_request: Final = kwargs.get("proxy_server_request") if not isinstance(proxy_server_request, dict): return body: Final = proxy_server_request.get("body") if isinstance(body, dict): body.pop("reasoning_effort", None) ### COMPLETION, EMBEDDING, IMG GENERATION FUNCTIONS def completion(self, model: str, messages: list[dict[str, str]], **kwargs) -> ModelResponse | CustomStreamWrapper: """ Example usage: response = router.completion(model="gpt-3.5-turbo", messages=[{"role": "user", "content": "Hey, how's it going?"}] """ try: verbose_router_logger.debug("router.completion(model=%s,..)", model) kwargs["model"] = model kwargs["messages"] = messages kwargs["original_function"] = self._completion self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs) response: Final = self.function_with_fallbacks(**kwargs) return response except Exception as e: raise e def _completion(self, model: str, messages: list[dict[str, str]], **kwargs) -> ModelResponse | CustomStreamWrapper: model_name = None deployment = None try: # Capture kwargs before deployment selection so the streaming # fallback iterator can re-dispatch with the original model group. input_kwargs_for_streaming_fallback: Final = kwargs.copy() input_kwargs_for_streaming_fallback["model"] = model # pick the one that is available (lowest TPM/RPM) deployment = self.get_available_deployment( model=model, messages=messages, specific_deployment=kwargs.pop("specific_deployment", None), request_kwargs=kwargs, ) self._drop_unsupported_classifier_reasoning_effort( deployment=cast(DeploymentTypedDict, deployment), # cast-ok: selection returns a router deployment model=model, kwargs=kwargs, ) # Check for silent model experiment # Make a local copy of litellm_params to avoid mutating the Router's state litellm_params: Final = self._deployment_params_with_request_reasoning_override( deployment["litellm_params"], kwargs ) silent_model: Final = litellm_params.pop("silent_model", None) for silent_target in _silent_experiment_targets(silent_model): # Mirroring traffic to a secondary model # Use threading.Thread (not ThreadPoolExecutor) - executor.submit() # requires pickling args, which fails when kwargs contain unpicklable # objects (e.g. _thread.RLock from OTEL spans, loggers) in deployment. threading.Thread( target=self._silent_experiment_completion, args=(silent_target, messages), kwargs=_silent_experiment_kwargs_snapshot(kwargs), daemon=True, ).start() kwargs.setdefault("messages", messages) self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) kwargs.pop("silent_model", None) # Ensure it's not in kwargs either model_name = litellm_params["model"] potential_model_client: Final = self._get_client(deployment=deployment, kwargs=kwargs) # check if provided keys == client keys # dynamic_api_key: Final = kwargs.get("api_key", None) if ( dynamic_api_key is not None and potential_model_client is not None and dynamic_api_key != potential_model_client.api_key ): model_client = None else: model_client = potential_model_client ### DEPLOYMENT-SPECIFIC PRE-CALL CHECKS ### (e.g. update rpm pre-call. Raise error, if deployment over limit) ## only run if model group given, not model id if model in self.model_names or not self.has_model_id(model): self.routing_strategy_pre_call_checks(deployment=deployment) input_kwargs: Final = { **litellm_params, "messages": messages, "caching": self.cache_responses, "client": model_client, **kwargs, } response: Final = litellm.completion(**input_kwargs) verbose_router_logger.info("litellm.completion(model=%s)\x1b[32m 200 OK\x1b[0m", model_name) ## CHECK CONTENT FILTER ERROR ## if isinstance(response, ModelResponse): _should_raise = self._should_raise_content_policy_error(model=model, response=response, kwargs=kwargs) if _should_raise: raise litellm.ContentPolicyViolationError( message="Response output was blocked.", model=model, llm_provider="", ) if ( isinstance(response, CustomStreamWrapper) and response.completion_stream is None and response.make_call is not None ): response.fetch_sync_stream() # Wrap streaming responses so MidStreamFallbackError (raised # during iteration) triggers the Router's fallback chain. if isinstance(response, CustomStreamWrapper): return self._completion_streaming_iterator( model_response=response, messages=messages, initial_kwargs=input_kwargs_for_streaming_fallback, ) return response except Exception as e: verbose_router_logger.info("litellm.completion(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e) # Set per-deployment num_retries on exception for retry logic if deployment is not None: self._set_deployment_num_retries_on_exception(e, deployment) self._stamp_failed_deployment_id_with_effective_model_info(e, deployment, kwargs) raise e def _get_silent_experiment_kwargs(self, **kwargs) -> dict: """ Prepare kwargs for a silent experiment by ensuring isolation from the primary call. Guarantee metadata isolation: safe_deep_copy falls back to the original reference when deepcopy fails (e.g. metadata contains UserAPIKeyAuth with parent_otel_span — an OTel Span that is not deepcopy-able). Force a shallow copy of the metadata dict so mutations (model_group, is_silent_experiment) never corrupt the main call's metadata. """ from litellm.litellm_core_utils.core_helpers import safe_deep_copy silent_kwargs: Final = safe_deep_copy(kwargs) # safe_deep_copy may fall back to the original metadata reference when # deepcopy fails (UserAPIKeyAuth.parent_otel_span is not deepcopy-able). # Detect this via identity check and force a shallow copy so that setting # model_group / is_silent_experiment on the silent dict doesn't corrupt # the primary call's metadata. original_metadata: Final = kwargs.get("metadata") if original_metadata is not None and silent_kwargs.get("metadata") is original_metadata: silent_kwargs["metadata"] = dict(original_metadata) if "metadata" not in silent_kwargs: silent_kwargs["metadata"] = {} # OTel spans are not safe to use across event loops. The silent # experiment runs in a new event loop, so strip the span to prevent # cross-loop tracing races or span corruption. silent_kwargs["metadata"].pop("litellm_parent_otel_span", None) silent_kwargs["metadata"]["is_silent_experiment"] = True # Pop logging objects and call IDs to ensure a fresh logging context # This prevents collisions in the Proxy's database (spend_logs) silent_kwargs.pop("litellm_call_id", None) silent_kwargs.pop("litellm_logging_obj", None) silent_kwargs.pop("standard_logging_object", None) # DON'T pop proxy_server_request — it's needed for spend log metadata return silent_kwargs async def _run_silent_experiment( self, silent_model: str, messages: Sequence[Mapping[str, str]], silent_kwargs: Mapping[str, object] ) -> None: remaining_kwargs: Final = MappingProxyType( {key: value for key, value in silent_kwargs.items() if key != "stream"} ) response: Final = await self.acompletion( model=silent_model, messages=cast(list[AllMessageValues], messages), stream=bool(silent_kwargs.get("stream", False)), **remaining_kwargs, ) if not isinstance(response, CustomStreamWrapper): return async for _ in response: pass def _silent_experiment_completion(self, silent_model: str, messages: Sequence[Mapping[str, str]], **kwargs): """ Run a silent experiment in the background (thread). """ try: # Prevent infinite recursion if silent model also has a silent model if kwargs.get("metadata", {}).get("is_silent_experiment", False): return messages = copy.deepcopy(messages) verbose_router_logger.info("Starting silent experiment for model %s", silent_model) silent_kwargs: Final = self._get_silent_experiment_kwargs(**kwargs) # Override model_group to correctly attribute metrics to the silent model silent_kwargs["metadata"]["model_group"] = silent_model # Create a new event loop for this thread so that async success # callbacks (e.g. _ProxyDBLogger) can schedule and run DB writes. loop: Final = asyncio.new_event_loop() asyncio.set_event_loop(loop) try: async def _run_silent_completion(): await self._run_silent_experiment(silent_model, messages, silent_kwargs) # Drain any fire-and-forget tasks (e.g. alerting hooks) # scheduled via asyncio.create_task during acompletion. pending: Final = asyncio.all_tasks() current: Final = asyncio.current_task() if current is not None: pending.discard(current) if pending: await asyncio.gather(*pending, return_exceptions=True) loop.run_until_complete(_run_silent_completion()) finally: loop.close() except Exception as e: verbose_router_logger.error("Silent experiment failed for model %s: %s", silent_model, e) # fmt: off @overload async def acompletion( self, model: str, messages: list[AllMessageValues], stream: Literal[True], **kwargs ) -> CustomStreamWrapper: ... @overload async def acompletion( self, model: str, messages: list[AllMessageValues], stream: Literal[False] = False, **kwargs ) -> ModelResponse: ... @overload async def acompletion( self, model: str, messages: list[AllMessageValues], stream: Literal[True, False] = False, **kwargs ) -> CustomStreamWrapper | ModelResponse: ... # fmt: on # The actual implementation of the function async def acompletion( self, model: str, messages: list[AllMessageValues], stream: bool = False, **kwargs, ): try: kwargs["model"] = model kwargs["messages"] = messages kwargs["stream"] = stream kwargs["original_function"] = self._acompletion self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs) request_priority: Final = kwargs.get("priority") or self.default_priority start_time: Final = time.time() _is_prompt_management_model: Final = self._is_prompt_management_model(model) if _is_prompt_management_model: return await self._prompt_management_factory( model=model, messages=messages, kwargs=kwargs, ) if request_priority is not None and isinstance(request_priority, int): response = await self.schedule_acompletion(**kwargs) else: response = await self.async_function_with_fallbacks(**kwargs) end_time: Final = time.time() _duration: Final = end_time - start_time asyncio.create_task( self.service_logger_obj.async_service_success_hook( service=ServiceTypes.ROUTER, duration=_duration, call_type="acompletion", start_time=start_time, end_time=end_time, parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs), ) ) return response except Exception as e: asyncio.create_task( send_llm_exception_alert( litellm_router_instance=self, request_kwargs=kwargs, error_traceback_str=traceback.format_exc(), original_exception=e, ) ) raise e @staticmethod def _combine_fallback_usage( fallback_item: ModelResponseStream, complete_response_object_usage: Usage | None, ) -> None: """Merge partial-stream usage with fallback-stream usage on the chunk.""" from litellm.cost_calculator import BaseTokenUsageProcessor usage: Final = cast(Usage | None, getattr(fallback_item, "usage", None)) usage_objects: Final = [usage] if usage is not None else [] if ( complete_response_object_usage is not None and hasattr(complete_response_object_usage, "usage") and complete_response_object_usage.usage is not None ): usage_objects.append(complete_response_object_usage) combined_usage: Final = BaseTokenUsageProcessor.combine_usage_objects(usage_objects=usage_objects) setattr(fallback_item, "usage", combined_usage) @staticmethod def _prepare_fallback_hidden_params( fallback_response: object, ) -> tuple[dict[str, object], dict[str, object]]: fallback_hidden_params: Final = get_hidden_params_dict(fallback_response) fallback_headers: Final = fallback_hidden_params.get("additional_headers") if not isinstance(fallback_headers, dict): return fallback_hidden_params, {} return fallback_hidden_params, cast("dict[str, object]", fallback_headers) @staticmethod def _adopt_fallback_response_headers( wrapper_ref: "weakref.ref[FallbackAwareStreamWrapper]", fallback_response: object, ) -> tuple[dict[str, object], dict[str, object]]: """Repoint the wrapper at `fallback_response`, returning its prepared hidden params.""" prepared: Final = Router._prepare_fallback_hidden_params(fallback_response) adopting_wrapper: Final = wrapper_ref() if adopting_wrapper is not None: adopting_wrapper.adopt_fallback_response_headers(fallback_response, prepared) return prepared @staticmethod def _apply_fallback_hidden_params_to_item( fallback_item: object, prepared_fallback_hidden_params: tuple[dict[str, object], dict[str, object]], ) -> None: if fallback_item is None or not hasattr(fallback_item, "_hidden_params"): return fallback_hidden_params, fallback_headers = prepared_fallback_hidden_params item_hidden_params: Final = get_hidden_params_dict(fallback_item) item_headers = item_hidden_params.get("additional_headers") if not isinstance(item_headers, dict): item_headers = {} cast(_HiddenParamsHost, fallback_item)._hidden_params = { **item_hidden_params, **fallback_hidden_params, "additional_headers": {**item_headers, **fallback_headers}, } async def _acompletion_streaming_iterator( self, model_response: CustomStreamWrapper, messages: list[dict[str, str]], initial_kwargs: dict, deployment_slot: contextlib.AsyncExitStack | None = None, ) -> CustomStreamWrapper: """ Helper to iterate over a streaming response. Catches errors for fallbacks using the router's fallback system `deployment_slot` holds the deployment's max_parallel_requests semaphore; it is released when the stream is exhausted, closed, or falls back to another deployment """ from litellm.exceptions import MidStreamFallbackError held_slot: Final = deployment_slot if deployment_slot is not None else contextlib.AsyncExitStack() class FallbackStreamWrapper(FallbackAwareStreamWrapper): def __init__(self, async_generator: AsyncGenerator): # Copy attributes from the original model_response super().__init__( completion_stream=async_generator, model=model_response.model, custom_llm_provider=model_response.custom_llm_provider, logging_obj=model_response.logging_obj, _response_headers=getattr(model_response, "_response_headers", None), ) self._async_generator = async_generator inner_chunks: Final[object] = getattr(model_response, "chunks", None) if isinstance(inner_chunks, list): self.chunks = inner_chunks # Preserve hidden params (including litellm_overhead_time_ms) from original response if hasattr(model_response, "_hidden_params"): self._hidden_params = model_response._hidden_params.copy() def __aiter__(self): return self async def __anext__(self): return await self._async_generator.__anext__() async def close_model_response() -> None: if not hasattr(model_response, "aclose"): return try: await model_response.aclose() except BaseException as e: verbose_router_logger.debug( "stream_with_fallbacks: error closing model_response: %s", e, ) async def stream_with_fallbacks(): fallback_response = None # Track for cleanup in finally try: async for item in model_response: yield item except MidStreamFallbackError as e: with anyio.CancelScope(shield=True): await close_model_response() await held_slot.aclose() if not e.is_pre_first_chunk and ( e.generated_content or _stream_chunks_have_generated_content(model_response.chunks) ): if e.original_exception is not None: raise e.original_exception from e raise from litellm.main import stream_chunk_builder complete_response_object: Final = stream_chunk_builder(chunks=model_response.chunks) complete_response_object_usage: Final = cast( Usage | None, getattr(complete_response_object, "usage", None), ) try: # Use the router's fallback system model_group: Final = cast(str, initial_kwargs.get("model")) fallbacks: Final[list | None] = initial_kwargs.get("fallbacks", self.fallbacks) context_window_fallbacks: Final[list | None] = initial_kwargs.get( "context_window_fallbacks", self.context_window_fallbacks ) content_policy_fallbacks: Final[list | None] = initial_kwargs.get( "content_policy_fallbacks", self.content_policy_fallbacks ) initial_kwargs["original_function"] = self._acompletion initial_kwargs["messages"] = messages self._update_kwargs_before_fallbacks(model=model_group, kwargs=initial_kwargs) fallback_response = await self.async_function_with_fallbacks_common_utils( e=e, disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs), fallbacks=fallbacks, context_window_fallbacks=context_window_fallbacks, content_policy_fallbacks=content_policy_fallbacks, model_group=model_group, args=(), kwargs=initial_kwargs, include_fallback_errors=initial_kwargs.get("include_fallback_errors", False) is True, ) # If fallback returns a streaming response, iterate over it if hasattr(fallback_response, "__aiter__"): prepared_fallback_hidden_params = Router._adopt_fallback_response_headers( wrapper_ref, fallback_response ) fallback_headers_are_settled = False async for fallback_item in fallback_response: if not fallback_headers_are_settled: fallback_headers_are_settled = True # a fallback that failed over again only repoints itself once it yields prepared_fallback_hidden_params = Router._adopt_fallback_response_headers( wrapper_ref, fallback_response ) Router._apply_fallback_hidden_params_to_item(fallback_item, prepared_fallback_hidden_params) if ( fallback_item and isinstance(fallback_item, ModelResponseStream) and hasattr(fallback_item, "usage") ): self._combine_fallback_usage(fallback_item, complete_response_object_usage) yield fallback_item else: # If fallback returns a non-streaming response, yield None yield None except Exception as fallback_error: # If fallback also fails, log and re-raise original error verbose_router_logger.error("Fallback also failed: %s", fallback_error) # No fallback handled the mid-stream error, so surface the # real provider exception (e.g. RateLimitError) instead of # leaking the internal MidStreamFallbackError to the client if ( isinstance(fallback_error, MidStreamFallbackError) and fallback_error.original_exception is not None ): raise fallback_error.original_exception from fallback_error raise fallback_error finally: # Close the underlying streams to release HTTP connections # back to the connection pool when the generator is closed # (e.g. on client disconnect). # Shield from anyio cancellation so the awaits can complete. with anyio.CancelScope(shield=True): await close_model_response() await held_slot.aclose() if fallback_response is not None and hasattr(fallback_response, "aclose"): try: await fallback_response.aclose() except BaseException as e: verbose_router_logger.debug( "stream_with_fallbacks: error closing fallback_response: %s", e, ) wrapped_response: Final = FallbackStreamWrapper(stream_with_fallbacks()) # weak, so the generator closing over it does not keep the wrapper out of # refcount teardown and delay the `finally` that releases the deployment slot wrapper_ref: Final = weakref.ref(wrapped_response) return wrapped_response @staticmethod def _extract_partial_responses_usage( source_iterator: "BaseResponsesAPIStreamingIterator", ) -> Optional["ResponseAPIUsage"]: """ Best-effort: pull partial token usage from a Responses-API streaming iterator that errored mid-stream, normalized to ResponseAPIUsage so the caller can combine without crossing token-naming conventions. Two sources, in priority order: 1. The bridge path (LiteLLMCompletionStreamingIterator) accumulates chat-completion chunks while streaming — feed them through stream_chunk_builder to recover chat Usage, then translate (prompt_tokens → input_tokens, completion_tokens → output_tokens). 2. The native path (ResponsesAPIStreamingIterator) only has a completed_response object if the stream reached RESPONSE_COMPLETED before erroring — uncommon mid-stream but worth checking. Already ResponseAPIUsage-shaped. Returns None when no partial usage is recoverable. """ from litellm.responses.litellm_completion_transformation.streaming_iterator import ( LiteLLMCompletionStreamingIterator, ) from litellm.types.llms.openai import ( ResponseAPIUsage, ResponseCompletedEvent, ResponseFailedEvent, ResponseIncompleteEvent, ) # Bridge subclass is the only iterator that accumulates chat-completion # chunks. isinstance narrows the type so we can read the attribute # directly instead of getattr-ing on the base class. if isinstance(source_iterator, LiteLLMCompletionStreamingIterator): chunks: Final = source_iterator.collected_chat_completion_chunks if chunks: try: from litellm.main import stream_chunk_builder built: Final = stream_chunk_builder(chunks=chunks) # stream_chunk_builder returns ModelResponse | # TextCompletionResponse | None. ModelResponse sets .usage # in __init__ rather than declaring it as a class field, so # static narrowing doesn't expose it. Mirror the sync path # (_completion_streaming_iterator) and pull via getattr. chat: Final[object | None] = getattr(built, "usage", None) if built is not None else None if chat is not None: # getattr-with-default because the test path may # substitute a SimpleNamespace lacking some fields; # real Usage instances always have them. prompt: Final = int(getattr(chat, "prompt_tokens", 0) or 0) completion: Final = int(getattr(chat, "completion_tokens", 0) or 0) total: Final = int(getattr(chat, "total_tokens", prompt + completion) or (prompt + completion)) return ResponseAPIUsage( input_tokens=prompt, output_tokens=completion, total_tokens=total, ) except Exception: # Builder is best-effort — fall through to native path. pass # Native path: completed_response is set only if RESPONSE_COMPLETED # arrived before the error (uncommon mid-stream but worth checking). # Already ResponseAPIUsage-shaped — return as-is. completed: Final = source_iterator.completed_response if isinstance( completed, (ResponseCompletedEvent, ResponseFailedEvent, ResponseIncompleteEvent), ): return completed.response.usage return None @staticmethod def _combine_responses_fallback_usage( fallback_item: "BaseLiteLLMOpenAIResponseObject", partial_usage: "ResponseAPIUsage", ) -> None: """ Merge partial-stream usage with fallback-stream usage on a Responses-API streaming event. Only mutates events that carry a `response` with a `usage` field (response.completed / response.failed / response.incomplete). Other events pass through unchanged. Both inputs are ResponseAPIUsage-shaped (see _extract_partial_responses_usage which normalizes the bridge path), so we can sum input_tokens / output_tokens / total_tokens directly and produce a clean ResponseAPIUsage — no token-naming split, no setattr bypass. """ from litellm.types.llms.openai import ( ResponseAPIUsage, ResponseCompletedEvent, ResponseFailedEvent, ResponseIncompleteEvent, ) if not isinstance( fallback_item, (ResponseCompletedEvent, ResponseFailedEvent, ResponseIncompleteEvent), ): return response: Final = fallback_item.response if response.usage is None: return fb: Final = response.usage response.usage = ResponseAPIUsage( input_tokens=(partial_usage.input_tokens or 0) + (fb.input_tokens or 0), output_tokens=(partial_usage.output_tokens or 0) + (fb.output_tokens or 0), total_tokens=(partial_usage.total_tokens or 0) + (fb.total_tokens or 0), ) @staticmethod def _build_responses_continuation_input( input_val: Union[str, "ResponseInputParam"] | None, generated_content: str, ) -> "ResponseInputParam": """ Convert Responses-API input + partial assistant output into a continuation input that asks the fallback model to pick up where the prior assistant message stopped. Best effort across providers. The chat-completions path uses Anthropic's `prefix: True` prefill trick on the assistant message; the Responses-API input schema has no direct equivalent, so we append an instruction (developer role) plus a prior assistant message containing the partial output. Providers without prefill semantics (OpenAI, Vertex) treat this as conversational context and may regenerate — same trade-off as the chat-completions path for non-Anthropic fallbacks. """ # base/continuation are List[Any] because ResponseInputParam items # are a wide Union of TypedDicts (EasyInputMessageParam, Message, # ResponseOutputMessageParam, ...) — annotating as List[Dict[str, Any]] # rejects the list() spread of input_val. We cast the combined list to # ResponseInputParam at the return. base: list[object] if isinstance(input_val, str): base = [ { "type": "message", "role": "user", "content": [{"type": "input_text", "text": input_val}], } ] elif isinstance(input_val, list): base = list(input_val) else: base = [] continuation: Final[list[object]] = [ { "type": "message", "role": "developer", "content": [ { "type": "input_text", "text": ( "The previous assistant response was interrupted " "mid-stream. Continue exactly where it stopped — " "do not repeat any of its content. Your response " "must read as a seamless continuation." ), } ], }, { "type": "message", "role": "assistant", "content": [{"type": "output_text", "text": generated_content}], }, ] return cast("ResponseInputParam", base + continuation) async def _aresponses_streaming_iterator( self, response: "BaseResponsesAPIStreamingIterator", initial_kwargs: dict[str, Any], ) -> "BaseResponsesAPIStreamingIterator": """ Wrap a Responses-API streaming iterator so MidStreamFallbackError triggers the Router's fallback chain (parity with _acompletion_streaming_iterator for the chat-completions path). The Responses-API streaming path goes through _ageneric_api_call_with_fallbacks rather than _acompletion, so the returned iterator is never wrapped by the chat completions fallback handler. Without this wrapper, MidStreamFallbackError raised mid-stream from the underlying CustomStreamWrapper (used by LiteLLMCompletionStreamingIterator when the Responses API is served via the completion bridge) propagates unhandled and the configured cross-provider fallback never fires. Full parity with the chat-completions path: - Pre-first-chunk: retry with the original input unchanged. - Partial content: inject a developer instruction + prior assistant message carrying the generated text so the fallback model continues rather than restarts. - Usage combining: merge partial-stream usage onto the fallback's response.completed event so accounting reflects both attempts. - Stream cleanup: shielded aclose() on both source and fallback iterators on terminate. """ from litellm.exceptions import MidStreamFallbackError from litellm.responses.streaming_iterator import ( BaseResponsesAPIStreamingIterator, _get_openai_response_types, ) source_iterator: Final = response # Pre-resolve the set of terminal stream event types so the # per-chunk type check inside FallbackResponsesStreamWrapper # stays cheap; mirrors the source-iterator filter at # responses/streaming_iterator.py:243-247. _openai_types: Final = _get_openai_response_types() _RESPONSES_TERMINAL_EVENT_TYPES: Final = ( _openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED, _openai_types.ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE, _openai_types.ResponsesAPIStreamEvents.RESPONSE_FAILED, ) class FallbackResponsesStreamWrapper(BaseResponsesAPIStreamingIterator): """ Subclasses BaseResponsesAPIStreamingIterator only for isinstance compatibility (proxy + interactions code paths check the type). Bypasses the parent constructor and delegates iteration to an async generator. """ fallback_headers_adopted: bool = False def __init__(self, async_generator: AsyncGenerator): import time from datetime import datetime self._async_generator = async_generator # Mirror every attribute BaseResponsesAPIStreamingIterator.__init__ # would have set. The wrapper bypasses super().__init__ (it has no # httpx.Response of its own and no provider config to drive), so # we copy from source_iterator where applicable and use safe # defaults elsewhere. This keeps inherited methods (e.g. # _check_max_streaming_duration, _handle_failure) safe to call. # # The bridge path (LiteLLMCompletionStreamingIterator used by # Anthropic/Bedrock/Vertex) does not call super().__init__ and # is missing many of these attributes — use getattr fallbacks # so wrapper construction never raises AttributeError. The # bridge stores the logging object as `litellm_logging_obj`. # base class declares non-Optional types for these # fields but the bridge path (LiteLLMCompletionStreamingIterator) # can legitimately omit them at runtime — keep the None # fallback. Same lines passed mypy on the pre-fix file # because the surrounding function body wasn't fully # type-narrowed; the new typed terminal-event tuple above # is what made these surface. self.response = getattr(source_iterator, "response", None) self.model = getattr(source_iterator, "model", None) self.logging_obj = getattr( source_iterator, "logging_obj", getattr(source_iterator, "litellm_logging_obj", None), ) self.finished = False self.responses_api_provider_config = getattr(source_iterator, "responses_api_provider_config", None) self.completed_response = None self.start_time = getattr(source_iterator, "start_time", datetime.now()) self._failure_handled = False self._yielded_first_chunk = False self._generated_content = "" self._completed_response_cached = False self._completed_response_logged = False self._completed_response_cache_hit = None self._persist_completed_response_before_logging = True self._stream_created_time = time.time() self.litellm_metadata = getattr(source_iterator, "litellm_metadata", None) self.custom_llm_provider = getattr(source_iterator, "custom_llm_provider", None) self.request_data = getattr(source_iterator, "request_data", {}) or {} self.call_type = getattr(source_iterator, "call_type", None) # Preserve hidden params so response headers (model_id, # api_base, additional_headers) keep flowing. self._hidden_params = dict(getattr(source_iterator, "_hidden_params", None) or {}) def adopt_fallback_headers(self, fallback_response: object) -> tuple[dict[str, object], dict[str, object]]: prepared: Final = Router._prepare_fallback_hidden_params(fallback_response) self._hidden_params = {**prepared[0], "additional_headers": prepared[1]} # mutable-ok: stream metadata self.fallback_headers_adopted = True return prepared def __aiter__(self): return self async def __anext__(self): try: chunk: Final = await self._async_generator.__anext__() except StopAsyncIteration: # The inner generator is exhausted. If we never sniffed a # terminal event off a chunk (the bridge path emits the # final response.completed via common_done_event_logic, # which raises StopAsyncIteration after returning it), # fall back to whatever the source iterator latched so # the proxy's container-ownership hook still sees a # completed_response instead of logging a spurious # "no completed_response" warning. if self.completed_response is None: self.completed_response = getattr(source_iterator, "completed_response", None) raise # Sniff the terminal stream event off each forwarded chunk # so ``self.completed_response`` is populated regardless of # which inner iterator produced it (source_iterator, # fallback_iterator, or any future wrapper). Without this # the proxy's container-ownership hook (which reads # ``getattr(stream_response, "completed_response", None)`` # via _extract_completed_responses_response) silently # records nothing on streaming /v1/responses calls — every # follow-up /v1/containers//files call then 403s for # the very key that created the container (#30210). if self.completed_response is None and getattr(chunk, "type", None) in _RESPONSES_TERMINAL_EVENT_TYPES: self.completed_response = chunk return chunk async def aclose(self): # async generators always expose aclose — no defensive check needed. await self._async_generator.aclose() async def stream_with_fallbacks(): held_lifecycle_events: tuple[object, ...] = () # rebind-ok: flushed at first output, dropped on fallback try: async for item in source_iterator: if _responses_stream_holds_event(item, len(held_lifecycle_events)): held_lifecycle_events = (*held_lifecycle_events, item) continue for held_event in held_lifecycle_events: yield held_event held_lifecycle_events = () yield item for held_event in held_lifecycle_events: yield held_event except MidStreamFallbackError as e: async with contextlib.aclosing( self._aresponses_fallback_attempt( e, source_iterator, initial_kwargs, wrapper.adopt_fallback_headers, held_lifecycle_events ) ) as fallback_stream: async for fallback_item in fallback_stream: yield fallback_item except Exception: for held_event in held_lifecycle_events: yield held_event raise finally: with anyio.CancelScope(shield=True): if hasattr(source_iterator, "aclose"): try: await source_iterator.aclose() except Exception as exc: verbose_router_logger.debug( "stream_with_fallbacks(aresponses): error closing source: %s", exc, ) wrapper: Final = FallbackResponsesStreamWrapper(stream_with_fallbacks()) return wrapper async def _aresponses_fallback_attempt( self, e: "MidStreamFallbackError", source_iterator: "BaseResponsesAPIStreamingIterator", initial_kwargs: dict[str, Any], # mutable-ok: mutated in-place before re-entering the fallback chain adopt_headers: Callable[[object], tuple[dict[str, object], dict[str, object]]], # mutable-ok: hidden params held_lifecycle_events: tuple[object, ...], ) -> AsyncGenerator[object, None]: """ Re-enters the Router's fallback chain for a mid-stream Responses API error and yields whatever the fallback attempt produces. The lifecycle events the primary stream held back reach the client only when no fallback lands, so the client sees exactly one response announced, the one whose id completes. Split out of _aresponses_streaming_iterator to keep each function's cyclomatic complexity within the repo's C901 budget. """ from litellm.exceptions import MidStreamFallbackError partial_usage: Final = Router._extract_partial_responses_usage(source_iterator) fallback_response = None # rebind-ok: pre-init so finally can close it if a fallback was actually attempted fallback_yielded = False # rebind-ok: flipped on the first fallback item so a fallback that dies before its first event still replays the primary's held announcement try: model_group: Final = cast(str, initial_kwargs.get("model")) # cast-ok: model group fallbacks: Final[list | None] = initial_kwargs.get( # mutable-ok: matches the common_utils list|None param "fallbacks", self.fallbacks ) context_window_fallbacks: Final[list | None] = initial_kwargs.get( # mutable-ok: matches the param below "context_window_fallbacks", self.context_window_fallbacks ) content_policy_fallbacks: Final[list | None] = initial_kwargs.get( # mutable-ok: matches the param below "content_policy_fallbacks", self.content_policy_fallbacks ) initial_kwargs["original_function"] = ( # rebind-ok: the fallback chain re-enters on the same kwargs self._ageneric_api_call_with_fallbacks_responses_attempt ) if e.generated_content and not e.is_pre_first_chunk: initial_kwargs["input"] = Router._build_responses_continuation_input( # rebind-ok: fallback hop input initial_kwargs.get("input"), e.generated_content, ) # The Responses-API path stores observability metadata # under "litellm_metadata" (not the default "metadata") — # see _ageneric_api_call_with_fallbacks. Mirroring that # here ensures model_group, model_group_alias, and trace # ids land in the same key litellm.aresponses reads from. self._update_kwargs_before_fallbacks( model=model_group, kwargs=initial_kwargs, metadata_variable_name="litellm_metadata", ) # The content-policy dispatch branch matches on the trigger's own type, so a refusal's # MidStreamFallbackError envelope is unwrapped here or the wrong fallback list is consulted. fallback_trigger: Final[Exception] = ( e.original_exception if isinstance(e.original_exception, litellm.ContentPolicyViolationError) else e ) fallback_response = await self.async_function_with_fallbacks_common_utils( # rebind-ok: set on success e=fallback_trigger, disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs), fallbacks=fallbacks, context_window_fallbacks=context_window_fallbacks, content_policy_fallbacks=content_policy_fallbacks, model_group=model_group, args=(), kwargs=initial_kwargs, include_fallback_errors=initial_kwargs.get("include_fallback_errors", False) is True, ) prepared_fallback_hidden_params: Final = adopt_headers(fallback_response) if hasattr(fallback_response, "__aiter__"): async for fallback_item in fallback_response: Router._apply_fallback_hidden_params_to_item(fallback_item, prepared_fallback_hidden_params) if partial_usage is not None: Router._combine_responses_fallback_usage(fallback_item, partial_usage) fallback_yielded = True yield fallback_item else: fallback_yielded = True # rebind-ok: see the pre-init above yield fallback_response except Exception as fallback_error: verbose_router_logger.error("Responses streaming fallback also failed: %s", fallback_error) if not fallback_yielded: for held_event in held_lifecycle_events: yield held_event if isinstance(fallback_error, MidStreamFallbackError) and fallback_error.original_exception is not None: raise fallback_error.original_exception from fallback_error raise finally: if fallback_response is not None and hasattr(fallback_response, "aclose"): with anyio.CancelScope(shield=True): try: await fallback_response.aclose() except Exception as exc: verbose_router_logger.debug( "stream_with_fallbacks(aresponses): error closing fallback: %s", exc, ) def _completion_streaming_iterator( self, model_response: CustomStreamWrapper, messages: list[dict[str, str]], initial_kwargs: dict, ) -> CustomStreamWrapper: """ Sync equivalent of _acompletion_streaming_iterator. Wraps a sync streaming response so that MidStreamFallbackError (raised by CustomStreamWrapper.__next__) triggers the Router's fallback chain instead of surfacing directly to the caller. """ from litellm.exceptions import MidStreamFallbackError class SyncFallbackStreamWrapper(FallbackAwareStreamWrapper): def __init__(self, sync_generator: Generator): super().__init__( completion_stream=sync_generator, model=model_response.model, custom_llm_provider=model_response.custom_llm_provider, logging_obj=model_response.logging_obj, _response_headers=getattr(model_response, "_response_headers", None), ) self._sync_generator = sync_generator if hasattr(model_response, "_hidden_params"): self._hidden_params = model_response._hidden_params.copy() def __iter__(self): return self def __next__(self): return next(self._sync_generator) router_self: Final = self def stream_with_fallbacks(): fallback_response = None try: for item in model_response: yield item except MidStreamFallbackError as e: if fallbacks_disabled_for_request(initial_kwargs) or ( not e.is_pre_first_chunk and (e.generated_content or _stream_chunks_have_generated_content(model_response.chunks)) ): if e.original_exception is not None: raise e.original_exception from e raise from litellm.main import stream_chunk_builder complete_response_object: Final = stream_chunk_builder(chunks=model_response.chunks) complete_response_object_usage: Final = cast( Usage | None, getattr(complete_response_object, "usage", None), ) try: model_group: Final = cast(str, initial_kwargs.get("model")) fallbacks: Final[list | None] = initial_kwargs.get("fallbacks", router_self.fallbacks) context_window_fallbacks: Final[list | None] = initial_kwargs.get( "context_window_fallbacks", router_self.context_window_fallbacks, ) content_policy_fallbacks: Final[list | None] = initial_kwargs.get( "content_policy_fallbacks", router_self.content_policy_fallbacks, ) initial_kwargs["original_function"] = router_self._completion initial_kwargs["messages"] = messages router_self._update_kwargs_before_fallbacks(model=model_group, kwargs=initial_kwargs) fallback_response = router_self.function_with_fallbacks( **initial_kwargs, fallbacks=fallbacks, context_window_fallbacks=context_window_fallbacks, content_policy_fallbacks=content_policy_fallbacks, ) if hasattr(fallback_response, "__iter__"): prepared_fallback_hidden_params = Router._adopt_fallback_response_headers( wrapper_ref, fallback_response ) fallback_headers_are_settled = False for fallback_item in fallback_response: if not fallback_headers_are_settled: fallback_headers_are_settled = True # a fallback that failed over again only repoints itself once it yields prepared_fallback_hidden_params = Router._adopt_fallback_response_headers( wrapper_ref, fallback_response ) Router._apply_fallback_hidden_params_to_item(fallback_item, prepared_fallback_hidden_params) if ( fallback_item and isinstance(fallback_item, ModelResponseStream) and hasattr(fallback_item, "usage") ): router_self._combine_fallback_usage(fallback_item, complete_response_object_usage) yield fallback_item else: yield None except Exception as fallback_error: verbose_router_logger.error("Fallback also failed: %s", fallback_error) if ( isinstance(fallback_error, MidStreamFallbackError) and fallback_error.original_exception is not None ): raise fallback_error.original_exception from fallback_error raise fallback_error finally: if hasattr(model_response, "close"): try: model_response.close() except BaseException as close_err: verbose_router_logger.debug( "stream_with_fallbacks: error closing model_response: %s", close_err, ) if fallback_response is not None and hasattr(fallback_response, "close"): try: fallback_response.close() except BaseException as close_err: verbose_router_logger.debug( "stream_with_fallbacks: error closing fallback_response: %s", close_err, ) wrapped_response: Final = SyncFallbackStreamWrapper(stream_with_fallbacks()) # weak, for the same reason as the async twin wrapper_ref: Final = weakref.ref(wrapped_response) return wrapped_response async def _silent_experiment_acompletion(self, silent_model: str, messages: Sequence[Mapping[str, str]], **kwargs): """ Run a silent experiment in the background. """ try: # Prevent infinite recursion if silent model also has a silent model if kwargs.get("metadata", {}).get("is_silent_experiment", False): return messages = copy.deepcopy(messages) verbose_router_logger.info("Starting silent experiment for model %s", silent_model) silent_kwargs: Final = self._get_silent_experiment_kwargs(**kwargs) # Override model_group to correctly attribute metrics to the silent model silent_kwargs["metadata"]["model_group"] = silent_model # Trigger the silent request await self._run_silent_experiment(silent_model, messages, silent_kwargs) except Exception as e: verbose_router_logger.error("Silent experiment failed for model %s: %s", silent_model, e) async def _acompletion( self, model: str, messages: list[dict[str, str]], **kwargs ) -> ModelResponse | CustomStreamWrapper: """ - Get an available deployment - call it with a semaphore over the call - semaphore specific to it's rpm - in the semaphore, make a check against it's local rpm before running """ model_name = None deployment = None _timeout_debug_deployment_dict = {} # this is a temporary dict to debug timeout issues try: input_kwargs_for_streaming_fallback: Final = kwargs.copy() input_kwargs_for_streaming_fallback["model"] = model parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) start_time: Final = time.time() deployment = await self.async_get_available_deployment( model=model, messages=messages, specific_deployment=kwargs.pop("specific_deployment", None), request_kwargs=kwargs, ) self._drop_unsupported_classifier_reasoning_effort( deployment=cast(DeploymentTypedDict, deployment), # cast-ok: selection returns a router deployment model=model, kwargs=kwargs, ) _timeout_debug_deployment_dict = deployment end_time: Final = time.time() _duration: Final = end_time - start_time asyncio.create_task( self.service_logger_obj.async_service_success_hook( service=ServiceTypes.ROUTER, duration=_duration, call_type="async_get_available_deployment", start_time=start_time, end_time=end_time, parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs), ) ) # debug how often this deployment picked self._track_deployment_metrics(deployment=deployment, parent_otel_span=parent_otel_span) # Check for silent model experiment # Make a local copy of litellm_params to avoid mutating the Router's state litellm_params: Final = self._deployment_params_with_request_reasoning_override( deployment["litellm_params"], kwargs ) silent_model: Final = litellm_params.pop("silent_model", None) for silent_target in _silent_experiment_targets(silent_model): # Mirroring traffic to a secondary model # This is a silent experiment, so we don't want to block the primary request asyncio.create_task( self._silent_experiment_acompletion( silent_model=silent_target, messages=messages, # Use messages instead of *args **_silent_experiment_kwargs_snapshot(kwargs), ) ) kwargs.setdefault("messages", messages) self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) kwargs.pop("silent_model", None) # Ensure it's not in kwargs either model_name = litellm_params["model"] model_client: Final = self._get_async_openai_model_client( deployment=deployment, kwargs=kwargs, ) self.total_calls[model_name] += 1 input_kwargs: Final = { **litellm_params, "messages": messages, "caching": self.cache_responses, "client": model_client, **kwargs, } input_kwargs.pop("silent_model", None) input_kwargs.pop("include_fallback_errors", None) logging_obj: Final[LiteLLMLogging | None] = kwargs.get("litellm_logging_obj", None) max_parallel_requests_limit: Final = self._get_client( deployment=deployment, kwargs=kwargs, client_type="max_parallel_requests", ) compacted_input: Final = await compact_to_fit(self, deployment, input_kwargs, "chat") async with contextlib.AsyncExitStack() as deployment_slot: if isinstance(max_parallel_requests_limit, MaxParallelRequestsLimit): deployment_slot.enter_context(max_parallel_requests_limit) await self.async_routing_strategy_pre_call_checks( deployment=deployment, logging_obj=logging_obj, parent_otel_span=parent_otel_span, ) response = await litellm.acompletion(**compacted_input) ## CHECK CONTENT FILTER ERROR ## if isinstance(response, ModelResponse): _should_raise = self._should_raise_content_policy_error( model=model, response=response, kwargs=kwargs ) if _should_raise: raise litellm.ContentPolicyViolationError( message="Response output was blocked.", model=model, llm_provider="", ) if ( isinstance(response, CustomStreamWrapper) and response.completion_stream is None and response.make_call is not None ): await response.fetch_stream() self.success_calls[model_name] += 1 verbose_router_logger.info("litellm.acompletion(model=%s)\x1b[32m 200 OK\x1b[0m", model_name) # debug how often this deployment picked self._track_deployment_metrics( deployment=deployment, response=response, parent_otel_span=parent_otel_span, ) if isinstance(response, CustomStreamWrapper): return await self._acompletion_streaming_iterator( model_response=response, messages=messages, initial_kwargs=input_kwargs_for_streaming_fallback, deployment_slot=deployment_slot.pop_all(), ) return response except litellm.Timeout as e: deployment_request_timeout_param: Final = _timeout_debug_deployment_dict.get("litellm_params", {}).get( "request_timeout", None ) deployment_timeout_param = _timeout_debug_deployment_dict.get("litellm_params", {}).get("timeout", None) if litellm.expose_router_debug_in_errors: e.message += f"\n\nDeployment Info: request_timeout: {deployment_request_timeout_param}\ntimeout: {deployment_timeout_param}" # Set per-deployment num_retries on exception for retry logic if deployment is not None: self._set_deployment_num_retries_on_exception(e, deployment) self._stamp_failed_deployment_id_with_effective_model_info(e, deployment, kwargs) raise e except Exception as e: verbose_router_logger.info("litellm.acompletion(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e) if model_name is not None: self.fail_calls[model_name] += 1 # Set per-deployment num_retries on exception for retry logic if deployment is not None: self._set_deployment_num_retries_on_exception(e, deployment) self._stamp_failed_deployment_id_with_effective_model_info(e, deployment, kwargs) raise e def _update_kwargs_before_fallbacks( self, model: str, kwargs: dict, metadata_variable_name: str | None = "metadata", ) -> None: """ Adds/updates to kwargs: - litellm_trace_id - metadata num_retries is deliberately left as the caller passed it, with no default filled in. async_function_with_retries resolves the router/global default itself, and it can only honour the precedence request > deployment litellm_params > litellm_settings while an absent num_retries still means "the caller did not ask for one". """ kwargs.setdefault("litellm_trace_id", str(uuid.uuid4())) model_group_alias: str | None = None if self._get_model_from_alias(model=model): model_group_alias = model kwargs.setdefault(metadata_variable_name, {}).update( {"model_group": model, "model_group_alias": model_group_alias} ) def _set_deployment_num_retries_on_exception(self, exception: Exception, deployment: dict) -> None: """ Set num_retries from deployment litellm_params on the exception. This allows the retry logic in async_function_with_retries to use per-deployment retry settings instead of the global setting. """ # Only set if exception doesn't already have num_retries if hasattr(exception, "num_retries") and exception.num_retries is not None: return litellm_params: Final = deployment.get("litellm_params", {}) dep_num_retries: Final = litellm_params.get("num_retries") if dep_num_retries is not None: try: exception.num_retries = int(dep_num_retries) # Handle both int and str except (ValueError, TypeError): pass # Skip if value can't be converted to int def _set_failed_deployment_id_on_exception(self, exception: Exception, deployment: Mapping[str, Any]) -> None: """ Stamp the failed deployment's `model_info.id` on the exception so the fallback layer can exclude it from subsequent re-picks within the same request (used by weighted-routing failover). Idempotent: never overwrites an existing value, so the id of the deployment that *first* failed in a chain is preserved if multiple layers re-raise. """ if getattr(exception, "failed_deployment_id", None): return deployment_id: Final = (deployment.get("model_info") or {}).get("id") if deployment_id: try: exception.failed_deployment_id = deployment_id except Exception: pass def _stamp_failed_deployment_id_with_effective_model_info( self, exception: Exception, deployment: Mapping[str, object], kwargs: Mapping[str, object] ) -> None: # A client-side-credential call gets a dynamic deployment id generated inside # _update_kwargs_with_deployment and stamped into kwargs["model_info"]; stamping # the static shared deployment's id instead would let one tenant's bad credentials # cool down the deployment every other tenant sharing this config relies on. effective_model_info: Final = kwargs.get("model_info") or deployment.get("model_info") or MappingProxyType({}) self._set_failed_deployment_id_on_exception(exception, MappingProxyType({"model_info": effective_model_info})) @staticmethod def _stamp_retry_skip_deployment_id(exception: Exception, kwargs: Mapping[str, object]) -> None: effective_model_info: Final = kwargs.get("model_info") deployment_id: Final = effective_model_info.get("id") if isinstance(effective_model_info, Mapping) else None if isinstance(deployment_id, str) and deployment_id: exception.retry_skip_deployment_id = deployment_id # pyright: ignore[reportAttributeAccessIssue] # dynamic stamp, read by _deployment_ids_to_skip_on_retry def _update_kwargs_with_default_litellm_params( self, kwargs: dict, metadata_variable_name: str | None = "metadata" ) -> None: """ Adds default litellm params to kwargs, if set. Handles inserting this as either "metadata" or "litellm_metadata" depending on the metadata_variable_name """ # 1) copy your defaults and pull out metadata defaults: Final = self.default_litellm_params.copy() metadata_defaults: Final = defaults.pop("metadata", {}) or {} # 2) add any non-metadata defaults that aren't already in kwargs for key, value in defaults.items(): if value is None: continue kwargs.setdefault(key, value) # 3) merge in metadata, this handles inserting this as either "metadata" or "litellm_metadata" kwargs.setdefault(metadata_variable_name, {}).update(metadata_defaults) def _handle_clientside_credential( self, deployment: dict, kwargs: dict, function_name: str | None = None ) -> Deployment: """ Build a per-request Deployment carrying the caller-supplied api_key/api_base, with its own stable id for cooldown, logging, and cost-map identity. This deployment is deliberately never registered with the router (no upsert_deployment/add_deployment call): doing so used to add it to self.model_list under the shared model_name, which made a request-scoped, caller-supplied provider credential a permanent, load-balanced deployment that every other caller of that model group could be routed onto. Its pricing is still registered directly, so a custom price configured on the underlying deployment still applies to this call. """ model_info: Final = deployment.get("model_info", {}).copy() litellm_params: Final = deployment["litellm_params"].copy() dynamic_litellm_params: Final = get_dynamic_litellm_params(litellm_params=litellm_params, request_kwargs=kwargs) # Use deployment model_name as model_group for generating model_id metadata_variable_name: Final = _get_router_metadata_variable_name( function_name=function_name, ) model_group: Final = kwargs.get(metadata_variable_name, {}).get("model_group") _model_id: Final = self.generate_model_id(model_group=model_group, litellm_params=dynamic_litellm_params) original_model_id: Final = model_info.get("id") model_info["id"] = _model_id model_info["original_model_id"] = original_model_id deployment_pydantic_obj: Final = Deployment( model_name=model_group, litellm_params=LiteLLM_Params(**dynamic_litellm_params), model_info=model_info, ) Router._register_deployment_pricing(deployment=deployment_pydantic_obj) return deployment_pydantic_obj @staticmethod def _merge_tools_from_deployment(deployment: dict, kwargs: dict) -> None: """ Merge tools from deployment litellm_params with request kwargs. When both have tools, concatenate them (deployment tools first, then request tools). tool_choice: use request value if provided, else deployment's. """ dep_params_raw: Final = deployment.get("litellm_params", {}) or {} if isinstance(dep_params_raw, dict): dep_params = dep_params_raw else: dep_params = dep_params_raw.model_dump(exclude_none=True) dep_tools: Final = dep_params.get("tools") or [] req_tools: Final = kwargs.get("tools") or [] if dep_tools or req_tools: merged: Final = list(dep_tools) + list(req_tools) kwargs["tools"] = merged if "tool_choice" not in kwargs and dep_params.get("tool_choice") is not None: kwargs["tool_choice"] = dep_params["tool_choice"] def _update_kwargs_with_deployment( self, deployment: dict, kwargs: dict, function_name: str | None = None, ) -> None: """ 3 jobs: - Adds selected deployment, model_info and api_base to kwargs["metadata"] (used for logging) - Adds default litellm params to kwargs, if set. - Merges tools from deployment with request (proxy-configured tools + request tools). """ for key in self._forwarded_alias_marker_keys_the_deployment_sets( deployment=deployment, forwarded_keys=kwargs.pop(_ALIAS_MARKER_FORWARDED_PARAMS_KWARG, ()) ): kwargs.pop(key, None) self._merge_tools_from_deployment(deployment=deployment, kwargs=kwargs) model_info = deployment.get("model_info", {}).copy() deployment_litellm_model_name = deployment["litellm_params"]["model"] deployment_api_base = deployment["litellm_params"].get("api_base") deployment_model_name: Final = deployment["model_name"] if is_clientside_credential(request_kwargs=kwargs): deployment_pydantic_obj: Final = self._handle_clientside_credential( deployment=deployment, kwargs=kwargs, function_name=function_name ) model_info = deployment_pydantic_obj.model_info.model_dump() deployment_litellm_model_name = deployment_pydantic_obj.litellm_params.model deployment_api_base = deployment_pydantic_obj.litellm_params.api_base metadata_variable_name: Final = _get_router_metadata_variable_name( function_name=function_name, ) kwargs.setdefault(metadata_variable_name, {}).update( { "deployment": deployment_litellm_model_name, "model_info": model_info, "api_base": deployment_api_base, "deployment_model_name": deployment_model_name, } ) # A retry/fallback reuses this same kwargs dict for the next deployment. # Refund and clear any reservation the previous deployment attempt left # here before it's wiped below, instead of relying on that attempt's # (possibly still-pending) failure event to do it. refund_stale_reservation_before_retry(self.cache, kwargs) set_io_token_rate_limit_request_kwargs(kwargs, store_in_context=deployment_has_io_token_limits(deployment)) kwargs[metadata_variable_name].setdefault( ROUTING_REQUEST_TAGS_METADATA_KEY, tuple(_get_tags_from_request_kwargs(kwargs, metadata_variable_name=metadata_variable_name)), ) ## DEPLOYMENT-LEVEL TAGS deployment_tags: Final = deployment.get("litellm_params", {}).get("tags") if deployment_tags: existing_tags = kwargs[metadata_variable_name].get("tags") or [] merged_tags: Final = list(existing_tags) for tag in deployment_tags: if tag not in merged_tags: merged_tags.append(tag) kwargs[metadata_variable_name]["tags"] = merged_tags ## CREDENTIAL NAME AS TAG credential_name: Final = deployment.get("litellm_params", {}).get("litellm_credential_name") if credential_name: credential_tag: Final = f"Credential: {credential_name}" existing_tags = kwargs[metadata_variable_name].get("tags") or [] if credential_tag not in existing_tags: existing_tags.append(credential_tag) kwargs[metadata_variable_name]["tags"] = existing_tags kwargs["model_info"] = model_info if function_name == "_ageneric_api_call_with_fallbacks": from litellm.passthrough.timeout_utils import ( resolve_llm_passthrough_timeout, ) _router_timeout: Final = ( self.request_timeout if self.request_timeout is not None else float(self._explicit_timeout) if isinstance(self._explicit_timeout, (int, float)) else None ) _router_stream_timeout: Final = ( self.stream_timeout if self.stream_timeout is not None else self.request_timeout if self.request_timeout is not None else self.default_litellm_params.get("stream_timeout") ) kwargs["timeout"] = resolve_llm_passthrough_timeout( kwargs=kwargs, litellm_params=deployment["litellm_params"], router_timeout=_router_timeout, router_stream_timeout=_router_stream_timeout, ) else: kwargs["timeout"] = self._get_timeout(kwargs=kwargs, data=deployment["litellm_params"]) self._update_kwargs_with_default_litellm_params(kwargs=kwargs, metadata_variable_name=metadata_variable_name) def _get_async_openai_model_client(self, deployment: dict, kwargs: dict): """ Helper to get AsyncOpenAI or AsyncAzureOpenAI client that was created for the deployment The same OpenAI client is re-used to optimize latency / performance in production If dynamic api key is provided: Do not re-use the client. Pass model_client=None. The OpenAI/ AzureOpenAI client will be recreated in the handler for the llm provider """ potential_model_client: Final = self._get_client(deployment=deployment, kwargs=kwargs, client_type="async") # check if provided keys == client keys # dynamic_api_key: Final = kwargs.get("api_key", None) if ( dynamic_api_key is not None and potential_model_client is not None and dynamic_api_key != potential_model_client.api_key ): model_client = None else: model_client = potential_model_client return model_client def _get_stream_timeout(self, kwargs: dict, data: dict) -> float | int | None: """Helper to get stream timeout from kwargs or deployment params""" return ( kwargs.get("stream_timeout", None) # the params dynamically set by user or data.get("stream_timeout", None) # timeout set on litellm_params for this deployment or self.stream_timeout # timeout set on router or self.request_timeout # litellm_settings.request_timeout (per-attempt) or self.default_litellm_params.get("stream_timeout", None) ) def _get_non_stream_timeout(self, kwargs: dict, data: dict) -> float | int | None: """Helper to get non-stream timeout from kwargs or deployment params""" timeout: Final = ( kwargs.get("timeout", None) # the params dynamically set by user or kwargs.get("request_timeout", None) # the params dynamically set by user or data.get("timeout", None) # timeout set on litellm_params for this deployment or data.get("request_timeout", None) # timeout set on litellm_params for this deployment or self.request_timeout # litellm_settings.request_timeout (per-attempt) or self.timeout # timeout set on router (router_settings.timeout) or self.default_litellm_params.get("timeout", None) ) return timeout def _get_timeout(self, kwargs: dict, data: dict) -> float | int | None: """Helper to get timeout from kwargs or deployment params""" timeout: float | int | None = None if kwargs.get("stream", False): timeout = self._get_stream_timeout(kwargs=kwargs, data=data) if timeout is None: timeout = self._get_non_stream_timeout( kwargs=kwargs, data=data ) # default to this if no stream specific timeout set return timeout async def abatch_completion( self, models: list[str], messages: list[dict[str, str]] | list[list[dict[str, str]]], **kwargs, ): """ Async Batch Completion. Used for 2 scenarios: 1. Batch Process 1 request to N models on litellm.Router. Pass messages as List[Dict[str, str]] to use this 2. Batch Process N requests to M models on litellm.Router. Pass messages as List[List[Dict[str, str]]] to use this Example Request for 1 request to N models: ``` response = await router.abatch_completion( models=["gpt-3.5-turbo", "groq-llama"], messages=[ {"role": "user", "content": "is litellm becoming a better product ?"} ], max_tokens=15, ) ``` Example Request for N requests to M models: ``` response = await router.abatch_completion( models=["gpt-3.5-turbo", "groq-llama"], messages=[ [{"role": "user", "content": "is litellm becoming a better product ?"}], [{"role": "user", "content": "who is this"}], ], ) ``` """ ############## Helpers for async completion ################## async def _async_completion_no_exceptions(model: str, messages: list[AllMessageValues], **kwargs): """ Wrapper around self.async_completion that catches exceptions and returns them as a result """ try: return await self.acompletion(model=model, messages=messages, **kwargs) except Exception as e: return e async def _async_completion_no_exceptions_return_idx( model: str, messages: list[AllMessageValues], idx: int, # index of message this response corresponds to **kwargs, ): """ Wrapper around self.async_completion that catches exceptions and returns them as a result """ try: return ( await self.acompletion(model=model, messages=messages, **kwargs), idx, ) except Exception as e: return e, idx ############## Helpers for async completion ################## if isinstance(messages, list) and all(isinstance(m, dict) for m in messages): _tasks = [] for model in models: # add each task but if the task fails _tasks.append(_async_completion_no_exceptions(model=model, messages=messages, **kwargs)) response = await asyncio.gather(*_tasks) return response elif isinstance(messages, list) and all(isinstance(m, list) for m in messages): _tasks = [] for idx, message in enumerate(messages): for model in models: # Request Number X, Model Number Y _tasks.append( _async_completion_no_exceptions_return_idx( model=model, idx=idx, messages=message, **kwargs, ) ) responses: Final = await asyncio.gather(*_tasks) final_responses: Final[list[list[Any]]] = [[] for _ in range(len(messages))] for response in responses: if isinstance(response, tuple): final_responses[response[1]].append(response[0]) else: final_responses[0].append(response) return final_responses async def abatch_completion_one_model_multiple_requests( self, model: str, messages: list[list[AllMessageValues]], **kwargs ): """ Async Batch Completion - Batch Process multiple Messages to one model_group on litellm.Router Use this for sending multiple requests to 1 model Args: model (List[str]): model group messages (List[List[Dict[str, str]]]): list of messages. Each element in the list is one request **kwargs: additional kwargs Usage: response = await self.abatch_completion_one_model_multiple_requests( model="gpt-3.5-turbo", messages=[ [{"role": "user", "content": "hello"}, {"role": "user", "content": "tell me something funny"}], [{"role": "user", "content": "hello good mornign"}], ] ) """ async def _async_completion_no_exceptions(model: str, messages: list[AllMessageValues], **kwargs): """ Wrapper around self.async_completion that catches exceptions and returns them as a result """ try: return await self.acompletion(model=model, messages=messages, **kwargs) except Exception as e: return e _tasks: Final = [] for message_request in messages: # add each task but if the task fails _tasks.append(_async_completion_no_exceptions(model=model, messages=message_request, **kwargs)) response: Final = await asyncio.gather(*_tasks) return response # fmt: off @overload async def abatch_completion_fastest_response( self, model: str, messages: list[dict[str, str]], stream: Literal[True], **kwargs ) -> CustomStreamWrapper: ... @overload async def abatch_completion_fastest_response( self, model: str, messages: list[dict[str, str]], stream: Literal[False] = False, **kwargs ) -> ModelResponse: ... # fmt: on async def abatch_completion_fastest_response( self, model: str, messages: list[dict[str, str]], stream: bool = False, **kwargs, ): """ model - List of comma-separated model names. E.g. model="gpt-4, gpt-3.5-turbo" Returns fastest response from list of model names. OpenAI-compatible endpoint. """ models: Final = [m.strip() for m in model.split(",")] async def _async_completion_no_exceptions( model_name: str, messages: list[dict[str, str]], stream: bool, **kwargs: object ) -> ModelResponse | CustomStreamWrapper | Exception: """ Wrapper around self.acompletion that catches exceptions and returns them as a result """ try: result = await self.acompletion(model=model_name, messages=messages, stream=stream, **kwargs) return result except asyncio.CancelledError: verbose_router_logger.debug("Received 'task.cancel'. Cancelling call w/ model=%s.", model_name) raise except Exception as e: return e pending_tasks = [] async def check_response(task: asyncio.Task): nonlocal pending_tasks try: result: Final = await task if isinstance(result, (ModelResponse, CustomStreamWrapper)): verbose_router_logger.debug("Received successful response. Cancelling other LLM API calls.") # If a desired response is received, cancel all other pending tasks for t in pending_tasks: t.cancel() return result except Exception: # Ignore exceptions, let the loop handle them pass finally: # Remove the task from pending tasks if it finishes try: pending_tasks.remove(task) except KeyError: pass for model_name in models: task = asyncio.create_task( _async_completion_no_exceptions(model_name=model_name, messages=messages, stream=stream, **kwargs) ) pending_tasks.append(task) # Await the first task to complete successfully while pending_tasks: done, pending_tasks = await asyncio.wait(pending_tasks, return_when=asyncio.FIRST_COMPLETED) for completed_task in done: result = await check_response(completed_task) if result is not None: # Return the first successful result result._hidden_params["fastest_response_batch_completion"] = True return result # If we exit the loop without returning, all tasks failed raise Exception("All tasks failed") ### SCHEDULER ### # fmt: off @overload async def schedule_acompletion( self, model: str, messages: list[AllMessageValues], priority: int, stream: Literal[False] = False, **kwargs ) -> ModelResponse: ... @overload async def schedule_acompletion( self, model: str, messages: list[AllMessageValues], priority: int, stream: Literal[True], **kwargs ) -> CustomStreamWrapper: ... # fmt: on async def schedule_acompletion( self, model: str, messages: list[AllMessageValues], priority: int, stream=False, **kwargs, ): parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) ### FLOW ITEM ### _request_id: Final = str(uuid.uuid4()) item: Final = FlowItem( priority=priority, # 👈 SET PRIORITY FOR REQUEST request_id=_request_id, # 👈 SET REQUEST ID model_name=model, # 👈 SAME as 'Router' ) ### [fin] ### ## ADDS REQUEST TO QUEUE ## await self.scheduler.add_request(request=item) ## POLL QUEUE end_time: Final = time.monotonic() + self.timeout curr_time = time.monotonic() poll_interval: Final = self.scheduler.polling_interval # poll every 3ms make_request = False while curr_time < end_time: _healthy_deployments, _ = await self._async_get_healthy_deployments( model=model, parent_otel_span=parent_otel_span ) make_request = await self.scheduler.poll( ## POLL QUEUE ## - returns 'True' if there's healthy deployments OR if request is at top of queue id=item.request_id, model_name=item.model_name, health_deployments=_healthy_deployments, ) if make_request: ## IF TRUE -> MAKE REQUEST break else: ## ELSE -> loop till default_timeout await asyncio.sleep(poll_interval) curr_time = time.monotonic() if make_request: try: _response: Final = await self.acompletion(model=model, messages=messages, stream=stream, **kwargs) _response._hidden_params.setdefault("additional_headers", {}) _response._hidden_params["additional_headers"].update({"x-litellm-request-prioritization-used": True}) return _response except Exception as e: setattr(e, "priority", priority) raise e else: # Clean up the request from the scheduler queue also before raising the timeout exception await self.scheduler.remove_request(request_id=item.request_id, model_name=item.model_name) raise litellm.Timeout( message="Request timed out while polling queue", model=model, llm_provider="openai", ) async def _schedule_factory( self, model: str, priority: int, original_function: Callable, args: tuple[object, ...], kwargs: dict[str, object], ): parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) ### FLOW ITEM ### _request_id: Final = str(uuid.uuid4()) item: Final = FlowItem( priority=priority, # 👈 SET PRIORITY FOR REQUEST request_id=_request_id, # 👈 SET REQUEST ID model_name=model, # 👈 SAME as 'Router' ) ### [fin] ### ## ADDS REQUEST TO QUEUE ## await self.scheduler.add_request(request=item) ## POLL QUEUE end_time: Final = time.monotonic() + self.timeout curr_time = time.monotonic() poll_interval: Final = self.scheduler.polling_interval # poll every 3ms make_request = False while curr_time < end_time: _healthy_deployments, _ = await self._async_get_healthy_deployments( model=model, parent_otel_span=parent_otel_span ) make_request = await self.scheduler.poll( ## POLL QUEUE ## - returns 'True' if there's healthy deployments OR if request is at top of queue id=item.request_id, model_name=item.model_name, health_deployments=_healthy_deployments, ) if make_request: ## IF TRUE -> MAKE REQUEST break else: ## ELSE -> loop till default_timeout await asyncio.sleep(poll_interval) curr_time = time.monotonic() if make_request: try: _response: Final = await original_function(*args, **kwargs) if isinstance(_response._hidden_params, dict): _response._hidden_params.setdefault("additional_headers", {}) _response._hidden_params["additional_headers"].update( {"x-litellm-request-prioritization-used": True} ) return _response except Exception as e: setattr(e, "priority", priority) raise e else: # Clean up the request from the scheduler queue also before raising the timeout exception await self.scheduler.remove_request(request_id=item.request_id, model_name=item.model_name) raise litellm.Timeout( message="Request timed out while polling queue", model=model, llm_provider="openai", ) def _is_prompt_management_model(self, model: str) -> bool: model_list: Final = self.get_model_list(model_name=model) if model_list is None or len(model_list) != 1: return False litellm_model: Final = model_list[0]["litellm_params"].get("model", None) if litellm_model is None or "/" not in litellm_model: return False split_litellm_model: Final = litellm_model.split("/")[0] return split_litellm_model in litellm._known_custom_logger_compatible_callbacks async def _prompt_management_factory( self, model: str, messages: list[AllMessageValues], kwargs: dict[str, Any], ): litellm_logging_object = kwargs.get("litellm_logging_obj", None) if litellm_logging_object is None: litellm_logging_object, kwargs = function_setup( **{ "original_function": "acompletion", "rules_obj": Rules(), "start_time": get_utc_datetime(), **kwargs, } ) litellm_logging_object = cast(LiteLLMLogging, litellm_logging_object) specific_deployment: Final = kwargs.pop("specific_deployment", None) prompt_management_deployment: Final = await self.async_get_available_deployment( model=model, messages=cast(list[dict[str, str]], messages), # cast-ok: selection reads messages structurally specific_deployment=specific_deployment, request_kwargs=kwargs, ) self._update_kwargs_with_deployment(deployment=prompt_management_deployment, kwargs=kwargs) data: Final = prompt_management_deployment["litellm_params"].copy() litellm_model: Final = data.get("model", None) # litellm_agent/ prefix only strips the model name, no prompt_id needed is_litellm_agent_model: Final = isinstance(litellm_model, str) and litellm_model.startswith("litellm_agent/") prompt_id = kwargs.get("prompt_id") or prompt_management_deployment["litellm_params"].get("prompt_id", None) prompt_variables: Final = kwargs.get("prompt_variables") or prompt_management_deployment["litellm_params"].get( "prompt_variables", None ) prompt_label: Final = kwargs.get("prompt_label", None) or prompt_management_deployment["litellm_params"].get( "prompt_label", None ) if not is_litellm_agent_model and (prompt_id is None or not isinstance(prompt_id, str)): raise ValueError(f"Prompt ID is not set or not a string. Got={prompt_id}, type={type(prompt_id)}") if prompt_variables is not None and not isinstance(prompt_variables, dict): raise ValueError( f"Prompt variables is set but not a dictionary. Got={prompt_variables}, type={type(prompt_variables)}" ) ( model, messages, optional_params, ) = litellm_logging_object.get_chat_completion_prompt( model=litellm_model, messages=messages, non_default_params=get_non_default_completion_params(kwargs=kwargs), prompt_id=prompt_id, prompt_variables=prompt_variables, prompt_label=prompt_label, request_kwargs=kwargs, injected_for_every_deployment=True, ) # Filter out prompt management specific parameters from data before merging prompt_management_params: Final = { "bitbucket_config", "dotprompt_config", "prompt_id", "prompt_variables", "prompt_label", "prompt_version", } filtered_data: Final = {k: v for k, v in data.items() if k not in prompt_management_params} kwargs = {**filtered_data, **kwargs, **optional_params} kwargs["model"] = model kwargs["messages"] = messages kwargs["litellm_logging_obj"] = litellm_logging_object kwargs["prompt_id"] = prompt_id kwargs["prompt_variables"] = prompt_variables kwargs["prompt_label"] = prompt_label _model_list: Final = self.get_model_list(model_name=model) if _model_list is None or len(_model_list) == 0: # if direct call to model kwargs.pop("original_function") return await litellm.acompletion(**kwargs) return await self.async_function_with_fallbacks(**kwargs) def image_generation(self, prompt: str, model: str, **kwargs): try: kwargs["model"] = model kwargs["prompt"] = prompt kwargs["original_function"] = self._image_generation kwargs.setdefault("metadata", {}).update({"model_group": model}) response: Final = self.function_with_fallbacks(**kwargs) return response except Exception as e: raise e def _image_generation(self, prompt: str, model: str, **kwargs): model_name: Final = "" try: verbose_router_logger.debug("Inside _image_generation()- model: %s; kwargs: %s", model, kwargs) deployment: Final = self.get_available_deployment( model=model, messages=[{"role": "user", "content": "prompt"}], specific_deployment=kwargs.pop("specific_deployment", None), request_kwargs=kwargs, ) self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) data: Final = deployment["litellm_params"].copy() model_client: Final = self._get_async_openai_model_client( deployment=deployment, kwargs=kwargs, ) self.total_calls[model_name] += 1 ### DEPLOYMENT-SPECIFIC PRE-CALL CHECKS ### (e.g. update rpm pre-call. Raise error, if deployment over limit) self.routing_strategy_pre_call_checks(deployment=deployment) response: Final = litellm.image_generation( **{ **data, "prompt": prompt, "caching": self.cache_responses, "client": model_client, **kwargs, } ) self.success_calls[model_name] += 1 verbose_router_logger.info("litellm.image_generation(model=%s)\x1b[32m 200 OK\x1b[0m", model_name) return response except Exception as e: verbose_router_logger.info("litellm.image_generation(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e) if model_name is not None: self.fail_calls[model_name] += 1 self._stamp_retry_skip_deployment_id(e, kwargs) raise e async def aimage_generation(self, prompt: str, model: str, **kwargs): try: kwargs["model"] = model kwargs["prompt"] = prompt kwargs["original_function"] = self._aimage_generation self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs) response: Final = await self.async_function_with_fallbacks(**kwargs) return response except Exception as e: asyncio.create_task( send_llm_exception_alert( litellm_router_instance=self, request_kwargs=kwargs, error_traceback_str=traceback.format_exc(), original_exception=e, ) ) raise e async def _aimage_generation(self, prompt: str, model: str, **kwargs): model_name = model try: verbose_router_logger.debug("Inside _image_generation()- model: %s; kwargs: %s", model, kwargs) parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) deployment: Final = await self.async_get_available_deployment( model=model, messages=[{"role": "user", "content": "prompt"}], specific_deployment=kwargs.pop("specific_deployment", None), request_kwargs=kwargs, ) self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) data: Final = deployment["litellm_params"].copy() model_name = data["model"] model_client: Final = self._get_async_openai_model_client( deployment=deployment, kwargs=kwargs, ) self.total_calls[model_name] += 1 async with self._deployment_slot(deployment=deployment, kwargs=kwargs, parent_otel_span=parent_otel_span): response = await litellm.aimage_generation( **{ **data, "prompt": prompt, "caching": self.cache_responses, "client": model_client, **kwargs, } ) self.success_calls[model_name] += 1 verbose_router_logger.info("litellm.aimage_generation(model=%s)\x1b[32m 200 OK\x1b[0m", model_name) return response except Exception as e: verbose_router_logger.info("litellm.aimage_generation(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e) if model_name is not None: self.fail_calls[model_name] += 1 self._stamp_retry_skip_deployment_id(e, kwargs) raise e async def atranscription(self, file: FileTypes, model: str, **kwargs): """ Example Usage: ``` from litellm import Router client = Router(model_list = [ { "model_name": "whisper", "litellm_params": { "model": "whisper-1", }, }, ]) audio_file = open("speech.mp3", "rb") transcript = await client.atranscription( model="whisper", file=audio_file ) ``` """ try: kwargs["model"] = model kwargs["file"] = file kwargs["original_function"] = self._atranscription self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs) response: Final = await self.async_function_with_fallbacks(**kwargs) return response except Exception as e: asyncio.create_task( send_llm_exception_alert( litellm_router_instance=self, request_kwargs=kwargs, error_traceback_str=traceback.format_exc(), original_exception=e, ) ) raise e async def _atranscription(self, file: FileTypes, model: str, **kwargs): model_name: Final = model try: verbose_router_logger.debug("Inside _atranscription()- model: %s; kwargs: %s", model, kwargs) parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) deployment: Final = await self.async_get_available_deployment( model=model, messages=[{"role": "user", "content": "prompt"}], specific_deployment=kwargs.pop("specific_deployment", None), request_kwargs=kwargs, ) self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) data: Final = deployment["litellm_params"].copy() model_client: Final = self._get_async_openai_model_client( deployment=deployment, kwargs=kwargs, ) self.total_calls[model_name] += 1 async with self._deployment_slot(deployment=deployment, kwargs=kwargs, parent_otel_span=parent_otel_span): response = await litellm.atranscription( **{ **data, "file": file, "caching": self.cache_responses, "client": model_client, **kwargs, } ) self.success_calls[model_name] += 1 verbose_router_logger.info("litellm.atranscription(model=%s)\x1b[32m 200 OK\x1b[0m", model_name) return response except Exception as e: verbose_router_logger.info("litellm.atranscription(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e) if model_name is not None: self.fail_calls[model_name] += 1 self._stamp_retry_skip_deployment_id(e, kwargs) raise e async def aspeech(self, model: str, input: str, voice: str | None = None, **kwargs): """ Example Usage: ``` from litellm import Router client = Router(model_list = [ { "model_name": "tts", "litellm_params": { "model": "tts-1", }, }, ]) async with client.aspeech( model="tts", voice="alloy", input="the quick brown fox jumped over the lazy dogs", api_base=None, api_key=None, organization=None, project=None, max_retries=1, timeout=600, client=None, optional_params={}, ) as response: response.stream_to_file(speech_file_path) ``` """ try: kwargs["model"] = model kwargs["input"] = input kwargs["voice"] = voice kwargs["original_function"] = self._aspeech self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs) response: Final = await self.async_function_with_fallbacks(**kwargs) return response except Exception as e: asyncio.create_task( send_llm_exception_alert( litellm_router_instance=self, request_kwargs=kwargs, error_traceback_str=traceback.format_exc(), original_exception=e, ) ) raise e async def _aspeech(self, model: str, input: str, voice: str | None = None, **kwargs): model_name: Final = model try: verbose_router_logger.debug("Inside _aspeech()- model: %s; kwargs: %s", model, kwargs) parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) deployment: Final = await self.async_get_available_deployment( model=model, messages=[{"role": "user", "content": "prompt"}], specific_deployment=kwargs.pop("specific_deployment", None), request_kwargs=kwargs, ) self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) data: Final = deployment["litellm_params"].copy() model_client: Final = self._get_async_openai_model_client( deployment=deployment, kwargs=kwargs, ) self.total_calls[model_name] += 1 async with self._deployment_slot(deployment=deployment, kwargs=kwargs, parent_otel_span=parent_otel_span): response = await litellm.aspeech( **{ **data, "input": input, "voice": data.get("voice") if voice is None else voice, "client": model_client, **kwargs, } ) self.success_calls[model_name] += 1 verbose_router_logger.info("litellm.aspeech(model=%s)\x1b[32m 200 OK\x1b[0m", model_name) return response except Exception as e: verbose_router_logger.info("litellm.aspeech(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e) if model_name is not None: self.fail_calls[model_name] += 1 self._stamp_retry_skip_deployment_id(e, kwargs) raise e async def arerank(self, model: str, **kwargs): try: kwargs["model"] = model kwargs["input"] = input kwargs["original_function"] = self._arerank self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs) response: Final = await self.async_function_with_fallbacks(**kwargs) return response except Exception as e: asyncio.create_task( send_llm_exception_alert( litellm_router_instance=self, request_kwargs=kwargs, error_traceback_str=traceback.format_exc(), original_exception=e, ) ) raise e async def _arerank(self, model: str, **kwargs): model_name = None try: verbose_router_logger.debug("Inside _rerank()- model: %s; kwargs: %s", model, kwargs) deployment: Final = await self.async_get_available_deployment( model=model, specific_deployment=kwargs.pop("specific_deployment", None), request_kwargs=kwargs, ) self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) data: Final = deployment["litellm_params"].copy() model_name = data["model"] model_client: Final = self._get_async_openai_model_client( deployment=deployment, kwargs=kwargs, ) self.total_calls[model_name] += 1 response: Final = await litellm.arerank( **{ **data, "caching": self.cache_responses, "client": model_client, **kwargs, } ) self.success_calls[model_name] += 1 verbose_router_logger.info("litellm.arerank(model=%s)\x1b[32m 200 OK\x1b[0m", model_name) return response except Exception as e: verbose_router_logger.info("litellm.arerank(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e) if model_name is not None: self.fail_calls[model_name] += 1 self._stamp_retry_skip_deployment_id(e, kwargs) raise e def text_completion( self, model: str, prompt: str, is_retry: bool | None = False, is_fallback: bool | None = False, is_async: bool | None = False, **kwargs, ): messages: Final = [{"role": "user", "content": prompt}] try: kwargs["model"] = model kwargs["prompt"] = prompt kwargs["num_retries"] = kwargs.get("num_retries", self.num_retries) kwargs.setdefault("metadata", {}).update({"model_group": model}) # pick the one that is available (lowest TPM/RPM) deployment: Final = self.get_available_deployment( model=model, messages=messages, specific_deployment=kwargs.pop("specific_deployment", None), request_kwargs=kwargs, ) data: Final = deployment["litellm_params"].copy() for k, v in self.default_litellm_params.items(): if k not in kwargs: # prioritize model-specific params > default router params kwargs[k] = v elif k == "metadata": kwargs[k].update(v) # call via litellm.completion() return litellm.text_completion(**{**data, "prompt": prompt, "caching": self.cache_responses, **kwargs}) except Exception as e: raise e async def atext_completion( self, model: str, prompt: str, is_retry: bool | None = False, is_fallback: bool | None = False, is_async: bool | None = False, **kwargs, ): if kwargs.get("priority", None) is not None: return await self._schedule_factory( model=model, priority=kwargs.pop("priority"), original_function=self.atext_completion, args=(model, prompt), kwargs=kwargs, ) try: kwargs["model"] = model kwargs["prompt"] = prompt kwargs["original_function"] = self._atext_completion self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs) response: Final = await self.async_function_with_fallbacks(**kwargs) return response except Exception as e: asyncio.create_task( send_llm_exception_alert( litellm_router_instance=self, request_kwargs=kwargs, error_traceback_str=traceback.format_exc(), original_exception=e, ) ) raise e async def _atext_completion(self, model: str, prompt: str, **kwargs): try: verbose_router_logger.debug("Inside _atext_completion()- model: %s; kwargs: %s", model, kwargs) parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) deployment: Final = await self.async_get_available_deployment( model=model, messages=[{"role": "user", "content": prompt}], specific_deployment=kwargs.pop("specific_deployment", None), request_kwargs=kwargs, ) self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) data: Final = deployment["litellm_params"].copy() model_name: Final = data["model"] model_client: Final = self._get_async_openai_model_client( deployment=deployment, kwargs=kwargs, ) self.total_calls[model_name] += 1 async with self._deployment_slot(deployment=deployment, kwargs=kwargs, parent_otel_span=parent_otel_span): response = await litellm.atext_completion( **{ **data, "prompt": prompt, "caching": self.cache_responses, "client": model_client, **kwargs, } ) self.success_calls[model_name] += 1 verbose_router_logger.info("litellm.atext_completion(model=%s)\x1b[32m 200 OK\x1b[0m", model_name) return response except Exception as e: verbose_router_logger.info("litellm.atext_completion(model=%s)\x1b[31m Exception %s\x1b[0m", model, e) if model is not None: self.fail_calls[model] += 1 self._stamp_retry_skip_deployment_id(e, kwargs) raise e async def aadapter_completion( self, adapter_id: str, model: str, is_retry: bool | None = False, is_fallback: bool | None = False, is_async: bool | None = False, **kwargs, ): try: kwargs["model"] = model kwargs["adapter_id"] = adapter_id kwargs["original_function"] = self._aadapter_completion kwargs.setdefault("metadata", {}).update({"model_group": model}) response: Final = await self.async_function_with_fallbacks(**kwargs) return response except Exception as e: asyncio.create_task( send_llm_exception_alert( litellm_router_instance=self, request_kwargs=kwargs, error_traceback_str=traceback.format_exc(), original_exception=e, ) ) raise e async def _aadapter_completion(self, adapter_id: str, model: str, **kwargs): try: verbose_router_logger.debug("Inside _aadapter_completion()- model: %s; kwargs: %s", model, kwargs) parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) deployment: Final = await self.async_get_available_deployment( model=model, messages=[{"role": "user", "content": "default text"}], specific_deployment=kwargs.pop("specific_deployment", None), request_kwargs=kwargs, ) self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) data: Final = deployment["litellm_params"].copy() model_name: Final = data["model"] model_client: Final = self._get_async_openai_model_client( deployment=deployment, kwargs=kwargs, ) self.total_calls[model_name] += 1 async with self._deployment_slot(deployment=deployment, kwargs=kwargs, parent_otel_span=parent_otel_span): response = await litellm.aadapter_completion( **{ **data, "adapter_id": adapter_id, "caching": self.cache_responses, "client": model_client, **kwargs, } ) self.success_calls[model_name] += 1 verbose_router_logger.info("litellm.aadapter_completion(model=%s)\x1b[32m 200 OK\x1b[0m", model_name) return response except Exception as e: verbose_router_logger.info("litellm.aadapter_completion(model=%s)\x1b[31m Exception %s\x1b[0m", model, e) if model is not None: self.fail_calls[model] += 1 self._stamp_retry_skip_deployment_id(e, kwargs) raise e async def _asearch_with_fallbacks(self, original_function: Callable, **kwargs): """ Helper function to make a search API call through the router with load balancing and fallbacks. Reuses the router's retry/fallback infrastructure. """ from litellm.router_utils.search_api_router import SearchAPIRouter return await SearchAPIRouter.async_search_with_fallbacks( router_instance=self, original_function=original_function, **kwargs, ) async def _asearch_with_fallbacks_helper(self, model: str, original_generic_function: Callable, **kwargs): """ Helper function for search API calls - selects a search tool and calls the original function. Called by async_function_with_fallbacks for each retry attempt. """ from litellm.router_utils.search_api_router import SearchAPIRouter return await SearchAPIRouter.async_search_with_fallbacks_helper( router_instance=self, model=model, original_generic_function=original_generic_function, **kwargs, ) async def aguardrail( self, guardrail_name: str, original_function: Callable, **kwargs, ): """ Execute a guardrail with load balancing and fallbacks. Args: guardrail_name: Name of the guardrail to execute original_function: The guardrail's execution function (e.g., async_pre_call_hook) **kwargs: Additional arguments passed to the guardrail Returns: Result from the guardrail execution """ kwargs["model"] = guardrail_name # For fallback system compatibility kwargs["original_generic_function"] = original_function kwargs["original_function"] = self._aguardrail_helper self._update_kwargs_before_fallbacks( model=guardrail_name, kwargs=kwargs, metadata_variable_name="litellm_metadata", ) verbose_router_logger.debug("Inside aguardrail() - guardrail_name: %s; kwargs: %s", guardrail_name, kwargs) response: Final = await self.async_function_with_fallbacks(**kwargs) return response async def _aguardrail_helper( self, model: str, original_generic_function: Callable, **kwargs, ): """ Helper for aguardrail - selects a guardrail deployment and executes it. Called by async_function_with_fallbacks for each retry attempt. Args: model: The guardrail_name (named 'model' for fallback system compatibility) original_generic_function: The guardrail's execution function **kwargs: Additional arguments """ guardrail_name: Final = model selected_guardrail: Final = self.get_available_guardrail( guardrail_name=guardrail_name, ) verbose_router_logger.debug( "Selected guardrail deployment: %s", selected_guardrail.get("litellm_params", {}).get("guardrail") ) # Pass the selected guardrail config to the original function kwargs["selected_guardrail"] = selected_guardrail response: Final = await original_generic_function(**kwargs) return response def get_available_guardrail( self, guardrail_name: str, ) -> "GuardrailTypedDict": """ Select a guardrail deployment using the router's load balancing strategy. Args: guardrail_name: Name of the guardrail to select Returns: Selected guardrail configuration dict """ from litellm.router_strategy.simple_shuffle import simple_shuffle healthy_deployments: Final = [g for g in self.guardrail_list if g.get("guardrail_name") == guardrail_name] if not healthy_deployments: raise ValueError(f"No guardrail found with name: {guardrail_name}") if len(healthy_deployments) == 1: return healthy_deployments[0] # Use simple_shuffle for weighted selection return simple_shuffle( resolve_model_alias=self._get_model_from_alias, healthy_deployments=healthy_deployments, model=guardrail_name, request_kwargs=None, ) async def _ageneric_api_call_with_fallbacks( self, model: str, original_function: Callable, attempt_function: Callable | None = None, **kwargs ): """ Helper function to make a generic LLM API call through the router, this allows you to use retries/fallbacks with litellm router attempt_function runs every attempt of the chain instead of the plain helper, so a streaming endpoint can wrap each attempt's stream with its own mid-stream fallback handling. """ try: kwargs["model"] = model kwargs["original_generic_function"] = original_function kwargs["original_function"] = attempt_function or self._ageneric_api_call_with_fallbacks_helper if attempt_function is not None: controls: Final = per_request_fallback_controls(kwargs) kwargs[MID_STREAM_FALLBACK_CONTROLS_KEY] = controls # rebind-ok: forwarded to every hop self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs, metadata_variable_name="litellm_metadata") verbose_router_logger.debug( "Inside ageneric_api_call_with_fallbacks() - model: %s; kwargs: %s", model, kwargs ) response: Final = await self.async_function_with_fallbacks(**kwargs) return response except Exception as e: asyncio.create_task( send_llm_exception_alert( litellm_router_instance=self, request_kwargs=kwargs, error_traceback_str=traceback.format_exc(), original_exception=e, ) ) raise e def _add_deployment_model_to_endpoint_for_llm_passthrough_route( self, kwargs: dict[str, Any], model: str, model_name: str ) -> dict[str, Any]: """ Add the deployment model to the endpoint for LLM passthrough route. e.g for bedrock invoke users can pass endpoint as /model/special-bedrock-model/invoke it should be actually sent as /model/us.anthropic.claude-3-5-sonnet-20240620-v1:0/invoke """ if "endpoint" in kwargs and kwargs["endpoint"]: # For provider-specific endpoints, strip the provider prefix from model_name # e.g., "bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0" -> "us.anthropic.claude-3-5-sonnet-20240620-v1:0" from litellm import get_llm_provider try: # get_llm_provider returns (model_without_prefix, provider, api_key, api_base) stripped_model_name, _, _, _ = get_llm_provider( model=model_name, custom_llm_provider=kwargs.get("custom_llm_provider"), api_base=kwargs.get("api_base"), ) replacement_model_name = stripped_model_name except Exception: # If get_llm_provider fails, fall back to using model_name as-is replacement_model_name = model_name kwargs["endpoint"] = replace_path_segment(kwargs["endpoint"], model, replacement_model_name) return kwargs async def _ageneric_api_call_with_fallbacks_helper(self, model: str, original_generic_function: Callable, **kwargs): """ Helper function to make a generic LLM API call through the router, this allows you to use retries/fallbacks with litellm router """ passthrough_on_no_deployment: Final = kwargs.pop("passthrough_on_no_deployment", False) function_name: Final = "_ageneric_api_call_with_fallbacks" deployment = None # rebind-ok: pre-init so the except block can stamp a failure with no deployment picked try: parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) try: deployment = await self.async_get_available_deployment( # rebind-ok: set on success, see pre-init above model=model, request_kwargs=kwargs, messages=kwargs.get("messages", None), input=kwargs.get("input", None), specific_deployment=kwargs.pop("specific_deployment", None), ) except Exception as e: if passthrough_on_no_deployment: return await original_generic_function(model=model, **kwargs) raise e self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs, function_name=function_name) data: Final = deployment["litellm_params"].copy() model_name: Final = data["model"] self.total_calls[model_name] += 1 self._add_deployment_model_to_endpoint_for_llm_passthrough_route( kwargs=kwargs, model=model, model_name=model_name ) custom_llm_provider: Final = provider_for_generic_call(data) response_kwargs: Final = { **data, "caching": self.cache_responses, **kwargs, "model": model_name, **_with_router_resolved_session_model(kwargs.get("session"), model_name), } # Only set custom_llm_provider if it's not None if custom_llm_provider is not None: response_kwargs["custom_llm_provider"] = custom_llm_provider compacted_input: Final = await compact_to_fit( self, deployment, response_kwargs, surface_for_call(getattr(original_generic_function, "__name__", "")), ) async with self._deployment_slot(deployment=deployment, kwargs=kwargs, parent_otel_span=parent_otel_span): response = await original_generic_function(**compacted_input) if self._should_raise_anthropic_refusal_error( model=model, original_generic_function=original_generic_function, response=response, kwargs=kwargs, ): from litellm.llms.anthropic.pass_through.messages.utils import ( safeguard_refusal_error, ) refusal_details: Final = cast(dict, response["stop_details"]) # cast-ok: gate verified the shape raise safeguard_refusal_error(model=model, stop_details=refusal_details) self.success_calls[model_name] += 1 verbose_router_logger.info("ageneric_api_call_with_fallbacks(model=%s)\x1b[32m 200 OK\x1b[0m", model_name) return response except Exception as e: verbose_router_logger.info( "ageneric_api_call_with_fallbacks(model=%s)\x1b[31m Exception %s\x1b[0m", model, e ) if model is not None: self.fail_calls[model] += 1 if deployment is not None: self._stamp_failed_deployment_id_with_effective_model_info(e, deployment, kwargs) raise e async def _aresponses_with_streaming_fallbacks( self, original_function: Callable, **kwargs: Any ) -> Union["ResponsesAPIResponse", "BaseResponsesAPIStreamingIterator"]: """ _ageneric_api_call_with_fallbacks for the Responses API, with every attempt's stream carrying its own mid-stream fallback handling (see _ageneric_api_call_with_fallbacks_responses_attempt). """ return await self._ageneric_api_call_with_fallbacks( original_function=original_function, attempt_function=self._ageneric_api_call_with_fallbacks_responses_attempt, **kwargs, ) async def _ageneric_api_call_with_fallbacks_responses_attempt( self, model: str, original_generic_function: Callable, **kwargs: object, # kwargs-ok: forwarded verbatim to the per-attempt helper, shape varies per call site ) -> Union["ResponsesAPIResponse", "BaseResponsesAPIStreamingIterator"]: """ One attempt of the Responses API fallback chain. A streaming result is wrapped with _aresponses_streaming_iterator over this attempt's own kwargs, so a fallback hop that fails mid-stream resumes the original group's chain instead of re-raising; the name keeps _get_router_metadata_variable_name resolving to litellm_metadata for every hop. """ from litellm.responses.streaming_iterator import ( BaseResponsesAPIStreamingIterator, ) controls: Final = kwargs.pop(MID_STREAM_FALLBACK_CONTROLS_KEY, None) hop_kwargs: Final = mid_stream_fallback_hop_kwargs( model=model, original_generic_function=original_generic_function, controls=controls, kwargs=kwargs ) response: Final = await self._ageneric_api_call_with_fallbacks_helper( model=model, original_generic_function=original_generic_function, **kwargs ) carry_over_pre_routing_selection(live_kwargs=kwargs, snapshot=hop_kwargs) if kwargs.get("stream") and isinstance(response, BaseResponsesAPIStreamingIterator): return await self._aresponses_streaming_iterator(response=response, initial_kwargs=hop_kwargs) return response async def _aanthropic_messages_streaming_iterator( self, response: AsyncIterator[bytes], initial_kwargs: dict[str, Any], # mutable-ok: mutated in-place before re-entering the fallback chain ) -> AsyncIterator[bytes]: """ Wrap an anthropic_messages (/v1/messages) streaming response so a mid-stream provider error triggers the Router's fallback chain (parity with _acompletion_streaming_iterator for the chat-completions path). See #24004. anthropic_messages goes through _ageneric_api_call_with_fallbacks rather than _acompletion, so the returned byte iterator is never wrapped by the chat-completions fallback handler. Two failure shapes land here: - the completion-bridge path (deployments with no native /v1/messages endpoint, via LiteLLMMessagesToCompletionTransformationHandler) already raises MidStreamFallbackError out of its underlying CustomStreamWrapper; this wrapper only needs to catch it. - a native Anthropic/Bedrock passthrough never raises anything for a provider SSE `event: error` frame (e.g. `overloaded_error`, `internal_server_error`) - it is forwarded to the client as-is - so this wrapper detects it via parse_anthropic_error_event and raises MidStreamFallbackError itself. Only an error before any real content (a content_block_delta frame) has reached the caller triggers a fallback attempt, mirroring the restriction _acompletion_streaming_iterator applies: once generated output has already reached the caller, retrying would start a second, overlapping Anthropic message lifecycle on the same SSE stream, so the error is left to propagate instead of being retried invisibly. A non-retriable client error (4xx other than 429) is never worth a fallback attempt either, so it is also left to propagate. Lifecycle/bookkeeping frames (message_start, content_block_start, ping, ...) do not by themselves disqualify a fallback attempt - Anthropic routinely sends message_start before an overload error. When a fallback can still take over they are BUFFERED rather than forwarded immediately, since forwarding one and then appending a fallback attempt's own message_start would produce two overlapping message lifecycles on one SSE stream; a `ping` carries no lifecycle, so it is forwarded live even while lifecycle frames sit buffered, keeping the connection alive during a long thinking pass. Buffered frames are flushed, in order, the moment real content arrives (the primary attempt has committed by then anyway) or once the stream ends without ever producing content or an error. When no fallback can take over the request is already committed, so every frame, including pings and provider error frames, is forwarded live and verbatim instead. """ from litellm.llms.anthropic.pass_through.messages.streaming_iterator import ( aclose_if_supported, parse_anthropic_error_event, parse_anthropic_refusal_stop_details, ) from litellm.llms.anthropic.pass_through.messages.utils import ( safeguard_refusal_error, ) source_iterator: Final = response async def stream_with_fallbacks() -> AsyncGenerator[bytes, None]: from litellm.exceptions import MidStreamFallbackError # Lifecycle/bookkeeping frames (message_start, content_block_start, # ...) are held back rather than forwarded immediately, but only # while a fallback can still take over: Anthropic routinely sends # message_start before an overload error, and once a byte reaches # the client a fallback attempt can only append its OWN # message_start, producing two overlapping message lifecycles on # one SSE stream. A `ping` keepalive carries no lifecycle, so it # is forwarded live even behind buffered frames, keeping the # connection alive through a long thinking pass. Buffered frames # are flushed the moment real content (content_block_delta) # arrives - at that point the primary attempt has committed and a # clean retry is no longer possible anyway - or once the primary # stream ends without ever producing content. Hitting # MAX_BUFFERED_PRE_CONTENT_ANTHROPIC_CHUNKS forces the same early # commit as real content arriving, so a hostile or pathological # upstream can't grow the buffer forever. With no fallback able # to take over there is nothing to buffer for, so every frame, # including pings and provider error frames, is forwarded live. model: Final = cast(str, initial_kwargs.get("model")) # cast-ok: kwargs always carries the model group has_generated_content = not self._anthropic_messages_stream_can_fall_back( # rebind-ok: set once real content is seen, the buffer cap is hit, or no fallback can take over model, initial_kwargs ) buffered_lifecycle_chunks: tuple[bytes, ...] = () # rebind-ok: flushed once committed or on decline try: async for chunk in source_iterator: if _anthropic_stream_forwards_ping_live(chunk, has_generated_content): yield chunk continue if _anthropic_stream_commits_now(chunk, has_generated_content, len(buffered_lifecycle_chunks)): has_generated_content = True # A transport can split one SSE data line across byte chunks, so pre-content # detection parses the accumulated buffer plus the current chunk, never the # chunk alone; the buffer is already capped, which bounds this window too. parse_window = ( b"".join(c for c in (*buffered_lifecycle_chunks, chunk) if isinstance(c, (bytes, bytearray))) # pyright: ignore[reportUnnecessaryIsInstance] # bridge-path chunks are not always bytes at runtime if not has_generated_content and isinstance(chunk, (bytes, bytearray)) # pyright: ignore[reportUnnecessaryIsInstance] # bridge-path chunks are not always bytes at runtime else chunk ) error_event = parse_anthropic_error_event(parse_window) retriable_pending_error = ( not has_generated_content and error_event is not None and _is_retriable_anthropic_status(error_event[2]) and not _anthropic_stream_error_is_gateway_verdict(chunk) ) refusal_stop_details = ( parse_anthropic_refusal_stop_details(parse_window) if not has_generated_content and error_event is None else None ) if refusal_stop_details is not None and self._refusal_fallback_available(model, initial_kwargs): refusal_error = safeguard_refusal_error(model=model, stop_details=refusal_stop_details) raise MidStreamFallbackError( message=refusal_error.message, model=model, llm_provider="anthropic", original_exception=refusal_error, is_pre_first_chunk=True, ) if not has_generated_content and not retriable_pending_error and error_event is None: buffered_lifecycle_chunks = (*buffered_lifecycle_chunks, chunk) continue if retriable_pending_error: assert error_event is not None _error_type, message, status_code = error_event raise MidStreamFallbackError( message=message, model=model, llm_provider="anthropic", original_exception=litellm.exceptions.APIError( status_code=status_code, message=message, llm_provider="anthropic", model=model, ), is_pre_first_chunk=True, ) for buffered_chunk in buffered_lifecycle_chunks: yield buffered_chunk buffered_lifecycle_chunks = () yield chunk for buffered_chunk in buffered_lifecycle_chunks: yield buffered_chunk except Exception as stream_error: # noqa: BLE001 # any raised provider error must reach the fallback gate async for item in self._aanthropic_messages_recover_stream_error( stream_error, has_generated_content, buffered_lifecycle_chunks, model, initial_kwargs, wrapper, ): yield item finally: with anyio.CancelScope(shield=True), contextlib.suppress(BaseException): await aclose_if_supported(source_iterator) # Referenced by stream_with_fallbacks via closure - assigned here, before # the generator body ever runs, so the reference resolves fine despite # being defined textually after the function that captures it. wrapper: Final = FallbackAwareAnthropicMessagesStream(stream_with_fallbacks(), source_iterator) return wrapper async def _aanthropic_messages_recover_stream_error( self, stream_error: Exception, has_generated_content: bool, buffered_lifecycle_chunks: tuple[bytes, ...], model: str, initial_kwargs: dict[str, Any], # mutable-ok: handed to _aanthropic_messages_fallback_attempt, which mutates it wrapper: "FallbackAwareAnthropicMessagesStream", ) -> AsyncGenerator[bytes, None]: """Turns a source-iterator failure into a fallback attempt or the error reaching the caller.""" from litellm.exceptions import MidStreamFallbackError if isinstance(stream_error, MidStreamFallbackError) and _anthropic_stream_should_decline_fallback( has_generated_content, stream_error ): for buffered_chunk in buffered_lifecycle_chunks: yield buffered_chunk if stream_error.original_exception is not None: raise stream_error.original_exception from stream_error raise stream_error fallback_error: Final = ( stream_error if isinstance(stream_error, MidStreamFallbackError) else _anthropic_stream_fallback_error_for_raised(stream_error, model, has_generated_content) ) if fallback_error is None: raise stream_error async for item in self._aanthropic_messages_fallback_attempt(fallback_error, initial_kwargs, wrapper): yield item async def _aanthropic_messages_fallback_attempt( self, e: "MidStreamFallbackError", initial_kwargs: dict[str, Any], # mutable-ok: mutated in-place before re-entering the fallback chain wrapper: "FallbackAwareAnthropicMessagesStream", ) -> AsyncGenerator[bytes, None]: """ Re-enters the Router's fallback chain for a mid-stream anthropic_messages error and yields whatever the fallback attempt produces. Split out of _aanthropic_messages_streaming_iterator to keep each function's cyclomatic complexity within the repo's C901 budget. """ from litellm.exceptions import MidStreamFallbackError from litellm.llms.anthropic.pass_through.messages.streaming_iterator import ( aclose_if_supported, anthropic_messages_response_as_sse_events, ) fallback_response = None # rebind-ok: pre-init so finally can close it if a fallback was actually attempted try: model_group: Final = cast(str, initial_kwargs.get("model")) # cast-ok: model group fallbacks: Final[list | None] = initial_kwargs.get( # mutable-ok: matches the common_utils list|None param "fallbacks", self.fallbacks ) context_window_fallbacks: Final[list | None] = initial_kwargs.get( # mutable-ok: matches the param below "context_window_fallbacks", self.context_window_fallbacks ) content_policy_fallbacks: Final[list | None] = initial_kwargs.get( # mutable-ok: matches the param below "content_policy_fallbacks", self.content_policy_fallbacks ) initial_kwargs["original_function"] = self._ageneric_api_call_with_fallbacks_anthropic_messages_attempt self._update_kwargs_before_fallbacks( model=model_group, kwargs=initial_kwargs, metadata_variable_name="litellm_metadata", ) # The content-policy dispatch branch matches on the trigger's own type, so a refusal's # MidStreamFallbackError envelope is unwrapped here or the wrong fallback list is consulted. fallback_trigger: Final[Exception] = ( e.original_exception if isinstance(e.original_exception, litellm.ContentPolicyViolationError) else e ) fallback_response = await self.async_function_with_fallbacks_common_utils( # rebind-ok: set on success e=fallback_trigger, disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs), fallbacks=fallbacks, context_window_fallbacks=context_window_fallbacks, content_policy_fallbacks=content_policy_fallbacks, model_group=model_group, args=(), kwargs=initial_kwargs, include_fallback_errors=initial_kwargs.get("include_fallback_errors", False) is True, ) fallback_hidden_params, fallback_headers = Router._prepare_fallback_hidden_params(fallback_response) wrapper.merge_fallback_hidden_params(fallback_hidden_params, fallback_headers) wrapper.adopt_fallback_source(fallback_response) if hasattr(fallback_response, "__aiter__"): async for fallback_item in fallback_response: yield fallback_item else: # A fallback can resolve to a complete AnthropicMessagesResponse # dict even for a streaming request (e.g. an agentic tool-use # interception loop) - yielding it as-is would put a raw dict # into a byte stream, so it's synthesized into the SSE # lifecycle a real stream would have sent instead. for event in anthropic_messages_response_as_sse_events( cast("AnthropicMessagesResponse", fallback_response) # cast-ok: non-streaming shape by elimination ): yield event except Exception as fallback_error: verbose_router_logger.error("Anthropic messages streaming fallback also failed: %s", fallback_error) if isinstance(fallback_error, MidStreamFallbackError) and fallback_error.original_exception is not None: raise fallback_error.original_exception from fallback_error raise finally: if fallback_response is not None: with anyio.CancelScope(shield=True), contextlib.suppress(BaseException): await aclose_if_supported(fallback_response) async def _aanthropic_messages_with_streaming_fallbacks( self, original_function: Callable, **kwargs: object, # kwargs-ok: forwarded verbatim to original_function, shape varies per call site ) -> Union["AnthropicMessagesResponse", AsyncIterator[bytes]]: """ _ageneric_api_call_with_fallbacks for anthropic_messages, with every attempt's stream carrying its own mid-stream fallback handling (see _ageneric_api_call_with_fallbacks_anthropic_messages_attempt). Parity with _aresponses_with_streaming_fallbacks for the Responses API. """ return await self._ageneric_api_call_with_fallbacks( original_function=original_function, attempt_function=self._ageneric_api_call_with_fallbacks_anthropic_messages_attempt, **kwargs, ) async def _ageneric_api_call_with_fallbacks_anthropic_messages_attempt( self, model: str, original_generic_function: Callable, **kwargs: object, # kwargs-ok: forwarded verbatim to the per-attempt helper, shape varies per call site ) -> Union["AnthropicMessagesResponse", AsyncIterator[bytes]]: """ One attempt of the anthropic_messages fallback chain. A streaming result is wrapped with _aanthropic_messages_streaming_iterator over this attempt's own kwargs, so a fallback hop that fails mid-stream resumes the original group's chain instead of re-raising; the name keeps _get_router_metadata_variable_name resolving to litellm_metadata for every hop. """ controls: Final = kwargs.pop(MID_STREAM_FALLBACK_CONTROLS_KEY, None) hop_kwargs: Final = mid_stream_fallback_hop_kwargs( model=model, original_generic_function=original_generic_function, controls=controls, kwargs=kwargs ) response: Final = await self._ageneric_api_call_with_fallbacks_helper( model=model, original_generic_function=original_generic_function, **kwargs ) carry_over_pre_routing_selection(live_kwargs=kwargs, snapshot=hop_kwargs) if kwargs.get("stream") and hasattr(response, "__aiter__"): return await self._aanthropic_messages_streaming_iterator( response=cast("AsyncIterator[bytes]", response), # cast-ok: stream=True always returns a byte iterator initial_kwargs=hop_kwargs, ) return response async def _dispatch_generic_call_type( self, call_type: str, original_function: Callable, **kwargs: object, # kwargs-ok: forwarded verbatim to the per-call-type helper, shape varies per call site ): """ factory_function's shared dispatch for call types with no call-specific handling, except anthropic_messages: kept out of factory_function's own async_wrapper (already at the repo's C901 complexity ceiling) so routing its mid-stream fallback handling (#24004) doesn't add another branch there. """ if call_type == "anthropic_messages": return await self._aanthropic_messages_with_streaming_fallbacks( original_function=original_function, **kwargs ) return await self._ageneric_api_call_with_fallbacks(original_function=original_function, **kwargs) def _generic_api_call_with_fallbacks(self, model: str, original_function: Callable, **kwargs): """ Make a generic LLM API call through the router, this allows you to use retries/fallbacks with litellm router Args: model: The model to use original_function: The handler function to call (e.g., litellm.completion) **kwargs: Additional arguments to pass to the handler function Returns: The response from the handler function """ handler_name: Final = original_function.__name__ metadata_variable_name: Final = _get_router_metadata_variable_name(function_name="generic_api_call") try: verbose_router_logger.debug( "Inside _generic_api_call() - handler: %s, model: %s; kwargs: %s", handler_name, model, kwargs ) self._update_kwargs_before_fallbacks( model=model, kwargs=kwargs, metadata_variable_name=metadata_variable_name, ) deployment: Final = self.get_available_deployment( model=model, messages=kwargs.get("messages", None), input=kwargs.get("input", None), specific_deployment=kwargs.pop("specific_deployment", None), request_kwargs=kwargs, ) self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs, function_name="generic_api_call") data: Final = deployment["litellm_params"].copy() model_name: Final = data["model"] self.total_calls[model_name] += 1 # For passthrough routes, use the actual model from deployment # and swap model name in endpoint if present if "endpoint" in kwargs and kwargs["endpoint"]: kwargs["endpoint"] = kwargs["endpoint"].replace(model, model_name) kwargs["model"] = model_name # Perform pre-call checks for routing strategy self.routing_strategy_pre_call_checks(deployment=deployment) custom_llm_provider: Final = provider_for_generic_call(data) response: Final = original_function( **{ **data, "custom_llm_provider": custom_llm_provider, "caching": self.cache_responses, **kwargs, } ) self.success_calls[model_name] += 1 verbose_router_logger.info("%s(model=%s)\x1b[32m 200 OK\x1b[0m", handler_name, model_name) return response except Exception as e: verbose_router_logger.info("%s(model=%s)\x1b[31m Exception %s\x1b[0m", handler_name, model, e) if model is not None: self.fail_calls[model] += 1 raise e def embedding( self, model: str, input: str | list, is_async: bool | None = False, **kwargs, ) -> EmbeddingResponse: try: kwargs["model"] = model kwargs["input"] = input kwargs["original_function"] = self._embedding self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs) response: Final = self.function_with_fallbacks(**kwargs) return response except Exception as e: raise e def _embedding(self, input: str | list, model: str, **kwargs): model_name = None try: verbose_router_logger.debug("Inside embedding()- model: %s; kwargs: %s", model, kwargs) deployment: Final = self.get_available_deployment( model=model, input=input, specific_deployment=kwargs.pop("specific_deployment", None), request_kwargs=kwargs, ) self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) data: Final = deployment["litellm_params"].copy() model_name = data["model"] potential_model_client: Final = self._get_client(deployment=deployment, kwargs=kwargs, client_type="sync") # check if provided keys == client keys # dynamic_api_key: Final = kwargs.get("api_key", None) if ( dynamic_api_key is not None and potential_model_client is not None and dynamic_api_key != potential_model_client.api_key ): model_client = None else: model_client = potential_model_client self.total_calls[model_name] += 1 ### DEPLOYMENT-SPECIFIC PRE-CALL CHECKS ### (e.g. update rpm pre-call. Raise error, if deployment over limit) self.routing_strategy_pre_call_checks(deployment=deployment) response: Final = litellm.embedding( **{ **data, "input": input, "caching": self.cache_responses, "client": model_client, **kwargs, } ) self.success_calls[model_name] += 1 verbose_router_logger.info("litellm.embedding(model=%s)\x1b[32m 200 OK\x1b[0m", model_name) return response except Exception as e: verbose_router_logger.info("litellm.embedding(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e) if model_name is not None: self.fail_calls[model_name] += 1 self._stamp_retry_skip_deployment_id(e, kwargs) raise e async def aembedding( self, model: str, input: str | list, is_async: bool | None = True, **kwargs, ) -> EmbeddingResponse: try: kwargs["model"] = model kwargs["input"] = input kwargs["original_function"] = self._aembedding self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs) response: Final = await self.async_function_with_fallbacks(**kwargs) return response except Exception as e: asyncio.create_task( send_llm_exception_alert( litellm_router_instance=self, request_kwargs=kwargs, error_traceback_str=traceback.format_exc(), original_exception=e, ) ) raise e async def _aembedding(self, input: str | list, model: str, **kwargs): model_name = None try: verbose_router_logger.debug("Inside _aembedding()- model: %s; kwargs: %s", model, kwargs) parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) deployment: Final = await self.async_get_available_deployment( model=model, input=input, specific_deployment=kwargs.pop("specific_deployment", None), request_kwargs=kwargs, ) self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) data: Final = deployment["litellm_params"].copy() model_name = data["model"] model_client: Final = self._get_async_openai_model_client( deployment=deployment, kwargs=kwargs, ) self.total_calls[model_name] += 1 async with self._deployment_slot(deployment=deployment, kwargs=kwargs, parent_otel_span=parent_otel_span): response = await litellm.aembedding( **{ **data, "input": input, "caching": self.cache_responses, "client": model_client, **kwargs, } ) self.success_calls[model_name] += 1 verbose_router_logger.info("litellm.aembedding(model=%s)\x1b[32m 200 OK\x1b[0m", model_name) return response except Exception as e: verbose_router_logger.info("litellm.aembedding(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e) if model_name is not None: self.fail_calls[model_name] += 1 self._stamp_retry_skip_deployment_id(e, kwargs) raise e #### FILES API #### async def acreate_file( self, model: str, **kwargs, ) -> OpenAIFileObject: try: kwargs["model"] = model kwargs["original_function"] = self._acreate_file self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs) response: Final = await self.async_function_with_fallbacks(**kwargs) return response except Exception as e: asyncio.create_task( send_llm_exception_alert( litellm_router_instance=self, request_kwargs=kwargs, error_traceback_str=traceback.format_exc(), original_exception=e, ) ) raise e async def _acreate_file( self, model: str, **kwargs, ) -> OpenAIFileObject: try: from litellm.router_utils.common_utils import add_model_file_id_mappings verbose_router_logger.debug("Inside _atext_completion()- model: %s; kwargs: %s", model, kwargs) parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) healthy_deployments: Final = await self.async_get_healthy_deployments( model=model, messages=[{"role": "user", "content": "files-api-fake-text"}], specific_deployment=kwargs.pop("specific_deployment", None), request_kwargs=kwargs, parent_otel_span=parent_otel_span, ) async def create_file_for_deployment(deployment: dict) -> OpenAIFileObject: from litellm.litellm_core_utils.core_helpers import safe_deep_copy kwargs_copy: Final = safe_deep_copy(kwargs) self._update_kwargs_with_deployment( deployment=deployment, kwargs=kwargs_copy, function_name="acreate_file", ) data: Final = deployment["litellm_params"].copy() model_name: Final = data["model"] model_client: Final = self._get_async_openai_model_client( deployment=deployment, kwargs=kwargs_copy, ) self.total_calls[model_name] += 1 ## REPLACE MODEL IN FILE WITH SELECTED DEPLOYMENT ## # For DB/config deployments, use provider from deployment params custom_llm_provider = data.get("custom_llm_provider") stripped_model, inferred_custom_llm_provider, _, _ = get_llm_provider( model=data["model"], custom_llm_provider=custom_llm_provider, ) # Preserve explicitly stored provider, fallback to inferred custom_llm_provider = custom_llm_provider or inferred_custom_llm_provider ## REPLACE MODEL IN FILE WITH SELECTED DEPLOYMENT ## purpose: Final = cast(OpenAIFilesPurpose | None, kwargs.get("purpose")) file = cast(FileTypes | None, kwargs.get("file")) if not file or not purpose: raise Exception("file and file_purpose are required for create_file") replace_model_in_jsonl_bool: Final = should_replace_model_in_jsonl( purpose=purpose, passthrough=kwargs.get("passthrough") is True, ) if replace_model_in_jsonl_bool: file = replace_model_in_jsonl( file_content=file, new_model_name=stripped_model, ) kwargs_copy["file"] = file if custom_llm_provider == LlmProviders.LITELLM_PROXY.value: kwargs_copy["extra_body"] = MappingProxyType( { **(kwargs_copy.get("extra_body") or MappingProxyType({})), "target_model_names": stripped_model, } ) if ( "gcs_bucket_name" in data ): # TODO: Remove this once we have a better way to handle GCS bucket name: Problem is that we need to pass the gcs_bucket_name to the router for the create_file call but it doesn't show up there kwargs_copy.setdefault("litellm_metadata", {})["gcs_bucket_name"] = data["gcs_bucket_name"] async with self._deployment_slot( deployment=deployment, kwargs=kwargs_copy, parent_otel_span=parent_otel_span ): response = await litellm.acreate_file( **{ **data, "custom_llm_provider": custom_llm_provider, "caching": self.cache_responses, "client": model_client, **kwargs_copy, } ) self.success_calls[model_name] += 1 verbose_router_logger.info("litellm.acreate_file(model=%s)\x1b[32m 200 OK\x1b[0m", model_name) return response tasks: Final = [] if isinstance(healthy_deployments, dict): tasks.append(create_file_for_deployment(healthy_deployments)) else: for deployment in healthy_deployments: tasks.append(create_file_for_deployment(deployment)) responses: Final = await asyncio.gather(*tasks) if len(responses) == 0: raise Exception("No healthy deployments found.") model_file_id_mapping: Final = add_model_file_id_mappings( healthy_deployments=healthy_deployments, responses=responses ) returned_response: Final = cast(OpenAIFileObject, responses[0]) returned_response._hidden_params["model_file_id_mapping"] = model_file_id_mapping return returned_response except Exception as e: verbose_router_logger.exception( "litellm.acreate_file(model=%s, %s)\x1b[31m Exception %s\x1b[0m", model, kwargs, e ) if model is not None: self.fail_calls[model] += 1 raise e #### VECTOR STORES API #### async def avector_store_create( self, model: str | None, **kwargs, ): """ Create a vector store for a specific model. Args: model: Model name from router config **kwargs: Vector store creation parameters Returns: VectorStoreCreateResponse """ try: # If model is None, use the factory function approach (direct SDK call) if model is None: from litellm.vector_stores.main import acreate # Use the factory function to handle the call factory_fn: Final = self.factory_function(acreate, call_type="avector_store_create") return await factory_fn(**kwargs) from litellm.vector_stores import acreate as avector_store_create_sdk parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) deployment: Final = await self.async_get_available_deployment( model=model, messages=[{"role": "user", "content": "vector-store-api-fake-text"}], specific_deployment=kwargs.pop("specific_deployment", None), request_kwargs=kwargs, ) data: Final = deployment["litellm_params"].copy() model_name: Final = data["model"] self._update_kwargs_with_deployment( deployment=deployment, kwargs=kwargs, function_name="avector_store_create", ) model_client: Final = self._get_async_openai_model_client( deployment=deployment, kwargs=kwargs, ) self.total_calls[model_name] += 1 # Get custom provider from deployment params custom_llm_provider = data.get("custom_llm_provider") _, inferred_custom_llm_provider, _, _ = get_llm_provider( model=data["model"], custom_llm_provider=custom_llm_provider, ) custom_llm_provider = custom_llm_provider or inferred_custom_llm_provider async with self._deployment_slot(deployment=deployment, kwargs=kwargs, parent_otel_span=parent_otel_span): response = await avector_store_create_sdk( **{ **data, "custom_llm_provider": custom_llm_provider, "caching": self.cache_responses, "client": model_client, **kwargs, } ) self.success_calls[model_name] += 1 verbose_router_logger.info("litellm.avector_store_create(model=%s)\x1b[32m 200 OK\x1b[0m", model_name) return response except Exception as e: verbose_router_logger.exception( "litellm.avector_store_create(model=%s)\x1b[31m Exception %s\x1b[0m", model, e ) if model is not None: self.fail_calls[model] += 1 raise e def _override_vector_store_methods_for_router(self): """ Override factory-generated vector store methods with router-aware implementations. This is called after _initialize_vector_store_endpoints() to ensure our custom methods that handle deployment selection and credential injection are used instead of the generic factory-generated ones. """ # Store references to the custom methods defined above # These methods handle proper routing through deployments # The methods are already defined as instance methods above async def acreate_batch( self, model: str, **kwargs, ) -> LiteLLMBatch: try: kwargs["model"] = model kwargs["original_function"] = self._acreate_batch metadata_variable_name: Final = _get_router_metadata_variable_name(function_name="_acreate_batch") self._update_kwargs_before_fallbacks( model=model, kwargs=kwargs, metadata_variable_name=metadata_variable_name, ) response: Final = await self.async_function_with_fallbacks(**kwargs) return response except Exception as e: asyncio.create_task( send_llm_exception_alert( litellm_router_instance=self, request_kwargs=kwargs, error_traceback_str=traceback.format_exc(), original_exception=e, ) ) raise e async def _acreate_batch( self, model: str, **kwargs, ) -> LiteLLMBatch: try: verbose_router_logger.debug("Inside _acreate_batch()- model: %s; kwargs: %s", model, kwargs) parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) deployment: Final = await self.async_get_available_deployment( model=model, messages=[{"role": "user", "content": "files-api-fake-text"}], specific_deployment=kwargs.pop("specific_deployment", None), request_kwargs=kwargs, ) data: Final = deployment["litellm_params"].copy() model_name: Final = data["model"] self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs, function_name="_acreate_batch") model_client: Final = self._get_async_openai_model_client( deployment=deployment, kwargs=kwargs, ) self.total_calls[model_name] += 1 ## SET CUSTOM PROVIDER TO SELECTED DEPLOYMENT ## custom_llm_provider = data.get("custom_llm_provider") _, inferred_custom_llm_provider, _, _ = get_llm_provider( model=data["model"], custom_llm_provider=custom_llm_provider, ) custom_llm_provider = custom_llm_provider or inferred_custom_llm_provider async with self._deployment_slot(deployment=deployment, kwargs=kwargs, parent_otel_span=parent_otel_span): response = await litellm.acreate_batch( **{ **data, "custom_llm_provider": custom_llm_provider, "caching": self.cache_responses, "client": model_client, **kwargs, } ) self.success_calls[model_name] += 1 verbose_router_logger.info("litellm.acreate_batch(model=%s)\x1b[32m 200 OK\x1b[0m", model_name) return response except Exception as e: verbose_router_logger.exception( "litellm._acreate_batch(model=%s, %s)\x1b[31m Exception %s\x1b[0m", model, kwargs, e ) if model is not None: self.fail_calls[model] += 1 self._stamp_retry_skip_deployment_id(e, kwargs) raise e async def aretrieve_batch( self, model: str | None = None, **kwargs, ) -> LiteLLMBatch: """ Iterate through all models in a model group to check for batch Future Improvement - cache the result. """ try: parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) requested_model_group: Final = model metadata_variable_name: Final = _get_router_metadata_variable_name(function_name="aretrieve_batch") if model is not None: filtered_model_list: ( list[DeploymentTypedDict] | list[dict] | dict | None ) = await self.async_get_healthy_deployments( model=model, messages=[{"role": "user", "content": "retrieve-api-fake-text"}], specific_deployment=kwargs.pop("specific_deployment", None), request_kwargs=kwargs, parent_otel_span=parent_otel_span, ) else: filtered_model_list = self.get_model_list() if filtered_model_list is None: raise Exception("Router not yet initialized.") receieved_exceptions: Final = [] async def try_retrieve_batch(model_name: DeploymentTypedDict): try: from litellm.litellm_core_utils.core_helpers import safe_deep_copy model: Final = model_name["litellm_params"].get("model") data: Final = model_name["litellm_params"].copy() custom_llm_provider = data.get("custom_llm_provider") if model is None: raise Exception(f"Model not found in litellm_params for deployment: {model_name}") # Update kwargs with the current model name or any other model-specific adjustments ## SET CUSTOM PROVIDER TO SELECTED DEPLOYMENT ## if not custom_llm_provider: _, custom_llm_provider, _, _ = get_llm_provider(model=model) new_kwargs: Final = safe_deep_copy(kwargs) self._update_kwargs_with_deployment( deployment=cast(dict, model_name), kwargs=new_kwargs, function_name="aretrieve_batch", ) model_group: Final = requested_model_group or model_name["model_name"] if not new_kwargs[metadata_variable_name].get("model_group"): new_kwargs[metadata_variable_name]["model_group"] = model_group new_kwargs.pop("custom_llm_provider", None) data.pop("custom_llm_provider", None) return await litellm.aretrieve_batch( **{ **data, "custom_llm_provider": custom_llm_provider, **new_kwargs, }, ) except Exception as e: import traceback traceback.print_exc() receieved_exceptions.append(e) return None # Check all models in parallel if ( filtered_model_list is not None and isinstance(filtered_model_list, list) and len(filtered_model_list) > 0 ): results = await asyncio.gather( *[try_retrieve_batch(cast(DeploymentTypedDict, model)) for model in filtered_model_list], return_exceptions=True, ) elif filtered_model_list is not None and isinstance(filtered_model_list, dict): results = await try_retrieve_batch(cast(DeploymentTypedDict, filtered_model_list)) else: raise Exception("No healthy deployments found.") # Check for successful responses and handle exceptions if results is not None: if isinstance(results, LiteLLMBatch): return results elif isinstance(results, list): for result in results: if isinstance(result, LiteLLMBatch): return result # If no valid Batch response was found, raise the first encountered exception if receieved_exceptions: raise receieved_exceptions[0] # Raising the first exception encountered # If no exceptions were encountered, raise a generic exception raise Exception(f"Unable to find batch in any model. Received errors - {receieved_exceptions}") except Exception as e: asyncio.create_task( send_llm_exception_alert( litellm_router_instance=self, request_kwargs=kwargs, error_traceback_str=traceback.format_exc(), original_exception=e, ) ) raise e async def acancel_batch( self, model: str, **kwargs, ) -> LiteLLMBatch: """ Cancel a batch through the router with proper model-to-provider mapping. """ try: kwargs["model"] = model kwargs["original_function"] = self._acancel_batch metadata_variable_name: Final = _get_router_metadata_variable_name(function_name="_acancel_batch") self._update_kwargs_before_fallbacks( model=model, kwargs=kwargs, metadata_variable_name=metadata_variable_name, ) response: Final = await self.async_function_with_fallbacks(**kwargs) return response except Exception as e: asyncio.create_task( send_llm_exception_alert( litellm_router_instance=self, request_kwargs=kwargs, error_traceback_str=traceback.format_exc(), original_exception=e, ) ) raise e async def _acancel_batch( self, model: str, **kwargs, ) -> LiteLLMBatch: try: verbose_router_logger.debug("Inside _acancel_batch()- model: %s; kwargs: %s", model, kwargs) parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) deployment: Final = await self.async_get_available_deployment( model=model, messages=[{"role": "user", "content": "batch-api-fake-text"}], specific_deployment=kwargs.pop("specific_deployment", None), request_kwargs=kwargs, ) selected_deployment_id: Final = (deployment.get("model_info") or {}).get("id") data: Final = deployment["litellm_params"].copy() resolved_credentials: Final = self.get_deployment_credentials_with_provider( model_id=selected_deployment_id or model ) if resolved_credentials is not None: data.update(resolved_credentials) data.pop("litellm_credential_name", None) model_name: Final = data["model"] self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs, function_name="_acancel_batch") model_client: Final = self._get_async_openai_model_client( deployment=deployment, kwargs=kwargs, ) self.total_calls[model_name] += 1 ## SET CUSTOM PROVIDER TO SELECTED DEPLOYMENT ## custom_llm_provider = data.get("custom_llm_provider") _, inferred_custom_llm_provider, _, _ = get_llm_provider( model=data["model"], custom_llm_provider=custom_llm_provider, ) custom_llm_provider = custom_llm_provider or inferred_custom_llm_provider async with self._deployment_slot(deployment=deployment, kwargs=kwargs, parent_otel_span=parent_otel_span): response = await litellm.acancel_batch( **{ **data, "custom_llm_provider": custom_llm_provider, "caching": self.cache_responses, "client": model_client, **kwargs, } ) self.success_calls[model_name] += 1 verbose_router_logger.info("litellm.acancel_batch(model=%s)\x1b[32m 200 OK\x1b[0m", model_name) return response except Exception as e: verbose_router_logger.exception( "litellm._acancel_batch(model=%s, %s)\x1b[31m Exception %s\x1b[0m", model, kwargs, e ) if model is not None: self.fail_calls[model] += 1 self._stamp_retry_skip_deployment_id(e, kwargs) raise e async def alist_batches( self, model: str, **kwargs, ): """ Return all the batches across all deployments of a model group. """ filtered_model_list: Final = self.get_model_list(model_name=model) if filtered_model_list is None: raise Exception("Router not yet initialized.") async def try_retrieve_batch(model: DeploymentTypedDict): try: # Update kwargs with the current model name or any other model-specific adjustments return await litellm.alist_batches(**{**model["litellm_params"], **kwargs}) except Exception: return None # Check all models in parallel results: Final = await asyncio.gather(*[try_retrieve_batch(model) for model in filtered_model_list]) final_results: Final[dict] = { "object": "list", "data": [], "first_id": None, "last_id": None, "has_more": False, } for result in results: if result is not None: ## check batch id if final_results["first_id"] is None and hasattr(result, "first_id"): final_results["first_id"] = getattr(result, "first_id") final_results["last_id"] = getattr(result, "last_id") final_results["data"].extend(result.data) ## check 'has_more' if getattr(result, "has_more", False) is True: final_results["has_more"] = True return final_results #### PASSTHROUGH API #### async def _pass_through_moderation_endpoint_factory( self, original_function: Callable, custom_llm_provider: str | None = None, **kwargs, ): # update kwargs with model_group self._update_kwargs_before_fallbacks( model=kwargs.get("model", ""), kwargs=kwargs, ) if kwargs.get("model") and self.get_model_list(model_name=kwargs["model"]): deployment: Final = await self.async_get_available_deployment( model=kwargs["model"], request_kwargs=kwargs, ) kwargs["model"] = deployment["litellm_params"]["model"] data: Final = deployment["litellm_params"].copy() self._update_kwargs_with_deployment( deployment=deployment, kwargs=kwargs, ) kwargs.update(data) return await original_function(**kwargs) def factory_function( self, original_function: Callable, call_type: Literal[ "assistants", "moderation", "anthropic_messages", "aresponses", "acancel_responses", "acompact_responses", "responses", "aget_responses", "adelete_responses", "afile_delete", "afile_content", "_arealtime", "acreate_realtime_client_secret", "arealtime_calls", "acreate_realtime_transcription_session", "_aresponses_websocket", "acreate_fine_tuning_job", "acancel_fine_tuning_job", "alist_fine_tuning_jobs", "aretrieve_fine_tuning_job", "alist_files", "aimage_edit", "allm_passthrough_route", "alist_input_items", "agenerate_content", "generate_content", "agenerate_content_stream", "generate_content_stream", "avector_store_search", "avector_store_create", "avector_store_retrieve", "avector_store_list", "avector_store_update", "avector_store_delete", "avector_store_file_create", "avector_store_file_list", "avector_store_file_retrieve", "avector_store_file_content", "avector_store_file_update", "avector_store_file_delete", "vector_store_search", "vector_store_create", "vector_store_retrieve", "vector_store_list", "vector_store_update", "vector_store_delete", "vector_store_file_create", "vector_store_file_list", "vector_store_file_retrieve", "vector_store_file_content", "vector_store_file_update", "vector_store_file_delete", "aocr", "ocr", "asearch", "search", "aadapter_generate_content", "avideo_generation", "video_generation", "avideo_list", "video_list", "avideo_status", "video_status", "avideo_content", "video_content", "avideo_remix", "video_remix", "avideo_create_character", "video_create_character", "avideo_get_character", "video_get_character", "avideo_edit", "video_edit", "avideo_extension", "video_extension", "acreate_container", "create_container", "alist_containers", "list_containers", "aretrieve_container", "retrieve_container", "adelete_container", "delete_container", "aupload_container_file", "upload_container_file", "alist_container_files", "list_container_files", "aretrieve_container_file", "retrieve_container_file", "adelete_container_file", "delete_container_file", "acreate_skill", "alist_skills", "aget_skill", "adelete_skill", "acreate_interaction", "create_interaction", "aget_interaction", "get_interaction", "adelete_interaction", "delete_interaction", "acancel_interaction", "cancel_interaction", "acreate_agent", "create_agent", "alist_agents", "list_agents", "aget_agent", "get_agent", "adelete_agent", "delete_agent", "alist_agent_versions", "list_agent_versions", ] = "assistants", ): """ Creates appropriate wrapper functions for different API call types. Returns: - A synchronous function for synchronous call types - An asynchronous function for asynchronous call types """ # Handle synchronous call types if call_type in ( "responses", "generate_content", "generate_content_stream", "ocr", "search", "video_generation", "video_list", "video_status", "video_content", "video_remix", "create_container", "list_containers", "retrieve_container", "delete_container", ): def sync_wrapper( custom_llm_provider: str | None = None, client: object | None = None, **kwargs, ): return self._generic_api_call_with_fallbacks(original_function=original_function, **kwargs) return sync_wrapper if call_type in ( "vector_store_search", "vector_store_create", "vector_store_retrieve", "vector_store_list", "vector_store_update", "vector_store_delete", ): def vector_store_sync_wrapper( custom_llm_provider: str | None = None, client: object | None = None, **kwargs, ): provider_kwargs: Final = ( MappingProxyType({**kwargs, "custom_llm_provider": custom_llm_provider}) if custom_llm_provider and "custom_llm_provider" not in kwargs else MappingProxyType(kwargs) ) search_kwargs: Final = ( MappingProxyType( { **provider_kwargs, "_direct_vector_store_embedding_executor": RouterVectorStoreEmbeddingExecutor( router=self, metadata=self._vector_store_request_metadata(kwargs), ), } ) if call_type == "vector_store_search" else provider_kwargs ) if search_kwargs.get("model"): return self._generic_api_call_with_fallbacks(original_function=original_function, **search_kwargs) if call_type == "vector_store_search": return original_function(**MappingProxyType({**search_kwargs, "router": self})) return original_function(**search_kwargs) return vector_store_sync_wrapper if call_type in ( "vector_store_file_create", "vector_store_file_list", "vector_store_file_retrieve", "vector_store_file_content", "vector_store_file_update", "vector_store_file_delete", ): def vector_store_file_sync_wrapper( custom_llm_provider: str | None = None, client: object | None = None, **kwargs, ): return original_function( custom_llm_provider=custom_llm_provider, client=client, **kwargs, ) return vector_store_file_sync_wrapper if call_type in ( "create_agent", "list_agents", "get_agent", "delete_agent", "list_agent_versions", ): def managed_agents_sync_wrapper( custom_llm_provider: str | None = None, client: object | None = None, **kwargs, ): if custom_llm_provider and "custom_llm_provider" not in kwargs: kwargs["custom_llm_provider"] = custom_llm_provider if "custom_llm_provider" not in kwargs: kwargs["custom_llm_provider"] = "gemini" return original_function(**kwargs) return managed_agents_sync_wrapper # Handle asynchronous call types async def async_wrapper( custom_llm_provider: str | None = None, client: AsyncOpenAI | None = None, **kwargs, ): if call_type == "assistants": return await self._pass_through_assistants_endpoint_factory( original_function=original_function, custom_llm_provider=custom_llm_provider, client=client, **kwargs, ) elif call_type == "moderation": return await self._pass_through_moderation_endpoint_factory( original_function=original_function, **kwargs ) elif call_type in ("asearch", "search"): return await self._asearch_with_fallbacks( original_function=original_function, **kwargs, ) elif call_type in ( "avector_store_file_create", "avector_store_file_list", "avector_store_file_retrieve", "avector_store_file_content", "avector_store_file_update", "avector_store_file_delete", ): return await self._init_vector_store_api_endpoints( original_function=original_function, custom_llm_provider=custom_llm_provider, **kwargs, ) elif call_type == "aresponses": return await self._aresponses_with_streaming_fallbacks( original_function=original_function, **kwargs, ) elif call_type in ( "acreate_realtime_client_secret", "arealtime_calls", "acreate_realtime_transcription_session", ): return await self._ageneric_api_call_with_fallbacks( original_function=original_function, client=client, **kwargs, ) elif call_type in ( "anthropic_messages", "_arealtime", "_aresponses_websocket", "acreate_fine_tuning_job", "acancel_fine_tuning_job", "alist_fine_tuning_jobs", "aretrieve_fine_tuning_job", "alist_files", "aimage_edit", "agenerate_content", "agenerate_content_stream", "aocr", "ocr", "avideo_generation", "avideo_list", "avideo_status", "avideo_content", "avideo_remix", "avideo_create_character", "avideo_get_character", "avideo_edit", "avideo_extension", "acreate_skill", "alist_skills", "aget_skill", "adelete_skill", ): return await self._dispatch_generic_call_type( call_type=call_type, original_function=original_function, **kwargs, ) elif call_type in ( "acreate_container", "alist_containers", "aretrieve_container", "adelete_container", "aupload_container_file", "alist_container_files", "aretrieve_container_file", "adelete_container_file", "aretrieve_container_file_content", ): return await self._init_containers_api_endpoints( original_function=original_function, custom_llm_provider=custom_llm_provider, **kwargs, ) elif call_type == "allm_passthrough_route": if client: kwargs["client"] = client return await self._ageneric_api_call_with_fallbacks( original_function=original_function, passthrough_on_no_deployment=True, **kwargs, ) elif call_type in ( "aget_responses", "acancel_responses", "acompact_responses", "adelete_responses", "alist_input_items", ): return await self._init_responses_api_endpoints( original_function=original_function, **kwargs, ) elif call_type in ( "avector_store_search", "avector_store_create", "avector_store_retrieve", "avector_store_list", "avector_store_update", "avector_store_delete", ): vector_store_kwargs: Final = ( { # mutable-ok: the async routed request requires dynamic keyword arguments **kwargs, "_direct_vector_store_embedding_executor": RouterVectorStoreEmbeddingExecutor( router=self, metadata=self._vector_store_request_metadata(kwargs), ), } if call_type == "avector_store_search" else kwargs ) return await self._init_vector_store_api_endpoints( original_function=original_function, custom_llm_provider=custom_llm_provider, call_type=call_type, **vector_store_kwargs, ) elif call_type in ("afile_delete", "afile_content"): return await self._ageneric_api_call_with_fallbacks( original_function=original_function, custom_llm_provider=custom_llm_provider, client=client, **kwargs, ) elif call_type in ( "acreate_interaction", "create_interaction", "aget_interaction", "adelete_interaction", "acancel_interaction", ): return await self._init_interactions_api_endpoints( original_function=original_function, custom_llm_provider=custom_llm_provider, **kwargs, ) elif call_type in ( "acreate_agent", "alist_agents", "aget_agent", "adelete_agent", "alist_agent_versions", ): return await self._init_managed_agents_api_endpoints( original_function=original_function, custom_llm_provider=custom_llm_provider, **kwargs, ) return async_wrapper @staticmethod def _vector_store_request_metadata(kwargs: Mapping[str, object]) -> Mapping[str, object]: return vector_store_request_metadata(kwargs) async def _init_vector_store_api_endpoints( self, original_function: Callable, custom_llm_provider: str | None = None, call_type: str | None = None, **kwargs, ): """ Initialize the Vector Store API endpoints on the router. If a model is provided in kwargs, use model-based routing to get the deployment credentials. Otherwise, call the original function directly. """ if custom_llm_provider and "custom_llm_provider" not in kwargs: kwargs["custom_llm_provider"] = custom_llm_provider # If model is provided, use generic API call with fallbacks for proper routing if kwargs.get("model"): return await self._ageneric_api_call_with_fallbacks( original_function=original_function, **kwargs, ) # For search, pass the router so provider transforms can resolve # router-managed embedding models (e.g. S3 Vectors query embeddings). # The merge also overrides any client-supplied `router` key. if call_type == "avector_store_search": search_kwargs: Final = MappingProxyType({**kwargs, "router": self}) return await original_function(**search_kwargs) # Otherwise, call the original function directly return await original_function(**kwargs) async def _init_containers_api_endpoints( self, original_function: Callable, custom_llm_provider: str | None = None, **kwargs, ): """ Initialize the Containers API endpoints on the router. LiteLLM-managed container IDs (``cntr_...``) encode ``model_id`` and provider metadata. When present, decode the ID, replace ``container_id`` with the upstream value, and route through ``_ageneric_api_call_with_fallbacks`` so deployment credentials (e.g. regional ``api_base`` for Azure) match :meth:`_init_responses_api_endpoints`. Create/list calls carry no container ID, so they route through the deployment named by ``model`` when the caller passes one, falling back to the direct call when no deployment matches. Otherwise call the handler directly with global provider credentials. """ if custom_llm_provider and "custom_llm_provider" not in kwargs: kwargs["custom_llm_provider"] = custom_llm_provider from litellm.responses.utils import ResponsesAPIRequestUtils container_id: Final = kwargs.get("container_id") _forwarded_model_id: Final = kwargs.get("model_id") if isinstance(container_id, str): decoded: Final = ResponsesAPIRequestUtils._decode_container_id(container_id) original_id: Final = decoded.get("response_id", container_id) if original_id != container_id: kwargs["container_id"] = original_id decoded_provider: Final = decoded.get("custom_llm_provider") if decoded_provider and kwargs.get("custom_llm_provider") == "openai": kwargs["custom_llm_provider"] = decoded_provider # Fall back to the model_id forwarded by the proxy when the container_id # is a native upstream ID (e.g. Azure hex cntr_) that carries no LiteLLM # routing payload, so deployment credentials (api_base, api_key) are applied. model_id: Final = decoded.get("model_id") or ( _forwarded_model_id.strip() if isinstance(_forwarded_model_id, str) and _forwarded_model_id.strip() else None ) if model_id: kwargs["model"] = model_id return await self._ageneric_api_call_with_fallbacks( original_function=original_function, **kwargs, ) requested_model: Final = kwargs.get("model") if isinstance(requested_model, str) and requested_model.strip(): return await self._ageneric_api_call_with_fallbacks( original_function=original_function, passthrough_on_no_deployment=True, **kwargs, ) return await original_function(**kwargs) async def _init_responses_api_endpoints( self, original_function: Callable, **kwargs, ): """ Initialize the Responses API endpoints on the router. GET, DELETE, CANCEL Responses API Requests encode the model_id in the response_id, this function decodes the response_id and sets the model to the model_id. """ from litellm.responses.utils import ResponsesAPIRequestUtils model_id: Final = ResponsesAPIRequestUtils.get_model_id_from_response_id(kwargs.get("response_id")) if model_id is not None: kwargs["model"] = model_id return await self._ageneric_api_call_with_fallbacks( original_function=original_function, **kwargs, ) async def _init_interactions_api_endpoints( self, original_function: Callable, custom_llm_provider: str | None = None, **kwargs, ): """ Initialize the Interactions API endpoints on the router. GET, DELETE, CANCEL Interactions API Requests don't need model-based routing, so we call the original function directly with the custom_llm_provider. """ if custom_llm_provider and "custom_llm_provider" not in kwargs: kwargs["custom_llm_provider"] = custom_llm_provider # Default to gemini for interactions API if "custom_llm_provider" not in kwargs: kwargs["custom_llm_provider"] = "gemini" # If the proxy accidentally passed agent name as model, clear it if kwargs.get("agent") and kwargs.get("model") == kwargs.get("agent"): kwargs["model"] = None # Model-based interactions use deployment routing + fallbacks; agent-only calls # must not enter model-group lookup (agent name is not a LiteLLM deployment). if kwargs.get("model"): return await self._ageneric_api_call_with_fallbacks( original_function=original_function, **kwargs, ) return await original_function(**kwargs) async def _init_managed_agents_api_endpoints( self, original_function: Callable, custom_llm_provider: str | None = None, **kwargs, ): """ Initialize the Managed Agents API endpoints on the router (v1beta/agents). CRUD operations for Gemini managed agents don't need model-based routing, so we call the original function directly with the custom_llm_provider. """ if custom_llm_provider and "custom_llm_provider" not in kwargs: kwargs["custom_llm_provider"] = custom_llm_provider if "custom_llm_provider" not in kwargs: kwargs["custom_llm_provider"] = "gemini" return await original_function(**kwargs) async def _pass_through_assistants_endpoint_factory( self, original_function: Callable, custom_llm_provider: str | None = None, client: AsyncOpenAI | None = None, **kwargs, ): """Internal helper function to pass through the assistants endpoint""" if custom_llm_provider is None: if self.assistants_config is not None: custom_llm_provider = self.assistants_config["custom_llm_provider"] kwargs.update(self.assistants_config["litellm_params"]) else: raise Exception( "'custom_llm_provider' must be set. Either via:\n `Router(assistants_config={'custom_llm_provider': ..})` \nor\n `router.arun_thread(custom_llm_provider=..)`" ) return await original_function(custom_llm_provider=custom_llm_provider, client=client, **kwargs) #### [END] ASSISTANTS API #### async def _maybe_run_weighted_failover( self, exception: Exception, original_model_group: str, all_deployments: Sequence[DeploymentTypedDict], args: tuple, kwargs: dict, input_kwargs: dict, ) -> Any | None: """Same-model-group retry after a failed deployment; returns None if not applicable.""" strategy, _ = self._get_routing_context(original_model_group, kwargs) if strategy != "simple-shuffle": return None failed_id: Final[str | None] = getattr(exception, "failed_deployment_id", None) if not failed_id: return None metadata_variable_name: Final = self._get_metadata_variable_name_from_kwargs(kwargs) meta = kwargs.get(metadata_variable_name) if meta is None: meta = {} kwargs[metadata_variable_name] = meta if not isinstance(meta, dict): return None prev_excluded: Final = set(meta.get("_failover_excluded_ids") or []) excluded: Final = prev_excluded | {failed_id} all_ids: Final = { (d.get("model_info") or {}).get("id") for d in all_deployments if (d.get("model_info") or {}).get("id") is not None } # Only consider deployments that are currently healthy (not in cooldown). # Using all_ids here would cause a wasteful run_async_fallback invocation # that fails with RouterRateLimitError whenever the "remaining" entries # are all in cooldown — the inner async_get_healthy_deployments call # would find an empty list and raise immediately. cooldown_ids = set(await _async_get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=None)) remaining: Final = (all_ids - cooldown_ids) - excluded if not remaining: return None verbose_router_logger.debug( "Weighted failover: exclude=%r, remaining=%s for model_group=%r", excluded, len(remaining), original_model_group, ) meta["_failover_excluded_ids"] = list(excluded) entry: Final = { "model": original_model_group, "_excluded_deployment_ids": list(excluded), } # Build a local copy so the weighted-failover keys do not leak back to # the caller's shared kwargs dict (any downstream fallback path reads # the same dict and must not inherit our `_excluded_deployment_ids` # entry). failover_kwargs: Final = { **input_kwargs, "fallback_model_group": [entry], "original_model_group": original_model_group, } try: return await run_async_fallback(*args, **failover_kwargs) except (openai.APIError, RouterRateLimitError, RouterRateLimitErrorBasic): # Expected model-level failure on the retried deployment. All # litellm provider errors derive from openai.APIError; if every # remaining deployment in the group is in cooldown the router # raises RouterRateLimitError (a ValueError, not an APIError). # In either case defer to the regular fallback path. Programming # errors (AttributeError, KeyError, TypeError, etc.) intentionally # propagate so they remain visible. return None async def async_function_with_fallbacks_common_utils( self, e: Exception, disable_fallbacks: bool | None, fallbacks: list | None, context_window_fallbacks: list | None, content_policy_fallbacks: list | None, model_group: str | None, args: tuple, kwargs: dict, include_fallback_errors: bool = False, ): """ Common utilities for async_function_with_fallbacks """ if verbose_router_logger.isEnabledFor(logging.DEBUG): verbose_router_logger.debug("Traceback%s", redact_string(traceback.format_exc())) original_exception: Final = e fallback_model_group = None original_model_group: Final[str | None] = kwargs.get("model") # A pre-routing hook (complexity / auto / adaptive / quality routers) picks a tier # behind the router name, and fallbacks are configured per tier, not per router. lookup_groups: Final[tuple[str, ...]] = fallback_lookup_groups(kwargs, model_group) fallback_failure_exception_str = "" no_fallback_group_explained = False hop_depth: Final = kwargs.get("fallback_depth") nested_fallback_hop: Final = isinstance(hop_depth, int) and hop_depth > 0 if disable_fallbacks is True or original_model_group is None or is_guardrail_intervention(e): raise e input_kwargs: Final = { "litellm_router": self, "original_exception": original_exception, **kwargs, } if "max_fallbacks" not in input_kwargs: input_kwargs["max_fallbacks"] = self.max_fallbacks if "fallback_depth" not in input_kwargs: input_kwargs["fallback_depth"] = 0 if include_fallback_errors: input_kwargs["include_fallback_errors"] = True # ORDER-BASED FALLBACKS: prepend higher order levels to the fallback list # Skip for error types that have their own dedicated fallback handlers _skip_order_fallback: Final = isinstance( e, (litellm.ContextWindowExceededError, litellm.ContentPolicyViolationError), ) _request_team_id: Final[str | None] = (kwargs.get("metadata", {}) or {}).get("user_api_key_team_id") # Use wildcard-aware lookup so order-based fallback also works for model # groups resolved via pattern routing (e.g. `openai/*` -> `openai/gpt-4.1-mini`). order_model_group: Final = get_pre_routing_selection(kwargs) or original_model_group all_deployments: Final = self.get_model_list(model_name=order_model_group, team_id=_request_team_id) or () _order_set: Final[set] = { litellm.utils._get_deployment_order(d) for d in all_deployments if litellm.utils._get_deployment_order(d) is not None } order_values: Final[list] = sorted(_order_set) if len(order_values) > 1 and not _skip_order_fallback: # Determine which order levels have already been tried current_target: Final = kwargs.get("_target_order") skip_up_to: Final = current_target if current_target is not None else order_values[0] # Build order-based fallback entries (skip already-tried levels) order_fallback_entries: Final[list] = [ {"model": order_model_group, "_target_order": o} for o in order_values if o > skip_up_to ] # Get external fallbacks — handle both standard and non-standard formats external_fallback_group: list | None = None if fallbacks is not None and lookup_groups: if _check_non_standard_fallback_format(fallbacks=fallbacks): # Non-standard formats (e.g. ["claude-3-haiku"] or # [{"model": "...", "messages": [...]}]) are passed through directly external_fallback_group = fallbacks else: external_fallback_group, generic_idx = get_fallback_model_group_for_lookup_groups( fallbacks=fallbacks, lookup_groups=lookup_groups, ) if external_fallback_group is None and generic_idx is not None: external_fallback_group = fallbacks[generic_idx]["*"] # Combined list: order fallbacks first, then external combined_fallbacks: Final = order_fallback_entries + (external_fallback_group or []) if combined_fallbacks: input_kwargs.update( { "fallback_model_group": combined_fallbacks, "original_model_group": original_model_group, } ) response = await run_async_fallback( *args, **input_kwargs, ) return response # Weighted intra-group failover (simple-shuffle only); see _maybe_run_weighted_failover. if self.enable_weighted_failover and not _skip_order_fallback and original_model_group is not None: response = await self._maybe_run_weighted_failover( exception=e, original_model_group=original_model_group, all_deployments=all_deployments, args=args, kwargs=kwargs, input_kwargs=input_kwargs, ) if response is not None: return response try: verbose_router_logger.info("Trying to fallback b/w models") # check if client-side fallbacks are used (e.g. fallbacks = ["gpt-3.5-turbo", "claude-3-haiku"] or fallbacks=[{"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hey, how's it going?"}]}] is_non_standard_fallback_format: Final = _check_non_standard_fallback_format(fallbacks=fallbacks) if is_non_standard_fallback_format: input_kwargs.update( { "fallback_model_group": fallbacks, "original_model_group": original_model_group, } ) response = await run_async_fallback( *args, **input_kwargs, ) return response if isinstance(e, litellm.ContextWindowExceededError): if context_window_fallbacks is not None: context_window_fallback_model_group: Final[list[str] | None] = ( self._get_fallback_model_group_for_lookup_groups( fallbacks=context_window_fallbacks, lookup_groups=lookup_groups, ) ) if context_window_fallback_model_group is None: raise original_exception input_kwargs.update( { "fallback_model_group": context_window_fallback_model_group, "original_model_group": original_model_group, } ) response = await run_async_fallback( *args, **input_kwargs, ) return response else: error_message = f"model={model_group}. context_window_fallbacks={mask_sensitive_structure(context_window_fallbacks)}. fallbacks={mask_sensitive_structure(fallbacks)}.\n\nSet 'context_window_fallback' - https://docs.litellm.ai/docs/routing#fallbacks" verbose_router_logger.info( msg=f"Got 'ContextWindowExceededError'. No context_window_fallback set. Defaulting \ to fallbacks, if available.{error_message}" ) if litellm.expose_router_debug_in_errors: e.message += f"\n{error_message}" elif isinstance(e, litellm.ContentPolicyViolationError): if content_policy_fallbacks is not None: content_policy_fallback_model_group: Final[list[str] | None] = ( self._get_fallback_model_group_for_lookup_groups( fallbacks=content_policy_fallbacks, lookup_groups=lookup_groups, ) ) if content_policy_fallback_model_group is None: raise original_exception input_kwargs.update( { "fallback_model_group": content_policy_fallback_model_group, "original_model_group": original_model_group, } ) response = await run_async_fallback( *args, **input_kwargs, ) return response else: error_message = f"model={model_group}. content_policy_fallback={mask_sensitive_structure(content_policy_fallbacks)}. fallbacks={mask_sensitive_structure(fallbacks)}.\n\nSet 'content_policy_fallback' - https://docs.litellm.ai/docs/routing#fallbacks" verbose_router_logger.info( msg=f"Got 'ContentPolicyViolationError'. No content_policy_fallback set. Defaulting \ to fallbacks, if available.{error_message}" ) if litellm.expose_router_debug_in_errors: e.message += f"\n{error_message}" if fallbacks is not None and lookup_groups: verbose_router_logger.debug("inside model fallbacks: %s", mask_sensitive_structure(fallbacks)) ( fallback_model_group, generic_fallback_idx, ) = get_fallback_model_group_for_lookup_groups( fallbacks=fallbacks, # if fallbacks = [{"gpt-3.5-turbo": ["claude-3-haiku"]}] lookup_groups=lookup_groups, ) ## if none, check for generic fallback if fallback_model_group is None and generic_fallback_idx is not None: fallback_model_group = fallbacks[generic_fallback_idx]["*"] if fallback_model_group is None: masked_fallbacks: Final = mask_sensitive_structure(fallbacks) verbose_router_logger.info( "No fallback model group found for lookup_groups=%s. Fallbacks=%s", " -> ".join(lookup_groups), masked_fallbacks, ) if ( hasattr(original_exception, "message") and litellm.expose_router_debug_in_errors and not nested_fallback_hop ): original_exception.message += format_no_fallback_group_message(lookup_groups, fallbacks) no_fallback_group_explained = True raise original_exception input_kwargs.update( { "fallback_model_group": fallback_model_group, "original_model_group": original_model_group, } ) response = await run_async_fallback( *args, **input_kwargs, ) return response except Exception as new_exception: parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) fallback_failure_exception_str = truncate_fallback_error_detail(redact_string(str(new_exception))) cooldown_info: Final = await _async_get_cooldown_deployments_with_debug_info( litellm_router_instance=self, parent_otel_span=parent_otel_span, ) verbose_router_logger.error( "litellm.router.py::async_function_with_fallbacks() - " "Error occurred while trying to do fallbacks - %s\n" "Debug Information:\nCooldown Deployments=%s", fallback_failure_exception_str, cooldown_info, ) attempted_fallback_group: Final = input_kwargs.get("fallback_model_group") if ( hasattr(original_exception, "message") and litellm.expose_router_debug_in_errors and not no_fallback_group_explained and not nested_fallback_hop ): original_exception.message += format_fallback_outcome_message( model_group, attempted_fallback_group, fallback_failure_exception_str ) raise original_exception @tracer.wrap() async def async_function_with_fallbacks(self, *args, **kwargs): """ Try calling the function_with_retries If it fails after num_retries, fall back to another model group """ model_group: Final[str | None] = kwargs.get("model") compaction_surface: Final = surface_for_call( getattr(kwargs.get("original_generic_function") or kwargs.get("original_function"), "__name__", "") ) if compaction_surface is not None: kwargs["_context_compaction_state"] = initialize_compaction_state(kwargs, compaction_surface) clear_pre_routing_selection(kwargs) # pyright: ignore[reportUnknownArgumentType] # **kwargs is untyped at this boundary if not isinstance(kwargs.get("attempted_targets"), AttemptedFallbackTargets): _fallback_metadata_key: Final = _get_router_metadata_variable_name( function_name=getattr(kwargs.get("original_function"), "__name__", None) ) _sibling_metadata_key: Final = ( "metadata" if _fallback_metadata_key == "litellm_metadata" else "litellm_metadata" ) if isinstance(_sibling_metadata := kwargs.get(_sibling_metadata_key), dict): # In place, like every other router bucket write: downstream resolves the bucket by # key presence, so rebinding kwargs to a copy detaches the proxy's request_data write-backs _sibling_metadata.pop("attempted_fallbacks", None) _sibling_metadata.pop("original_model_group", None) if isinstance(_fallback_metadata := kwargs.get(_fallback_metadata_key), dict): _fallback_metadata["attempted_fallbacks"] = 0 if model_group is not None: _fallback_metadata["original_model_group"] = model_group include_fallback_errors: Final = kwargs.get("include_fallback_errors", False) is True disable_fallbacks: Final[bool | None] = kwargs.pop("disable_fallbacks", False) record_disable_fallbacks(kwargs, disable_fallbacks is True) fallbacks: Final[list | None] = kwargs.get("fallbacks", self.fallbacks) context_window_fallbacks: list | None = kwargs.get("context_window_fallbacks", self.context_window_fallbacks) content_policy_fallbacks: list | None = kwargs.get("content_policy_fallbacks", self.content_policy_fallbacks) mock_timeout: Final = kwargs.pop("mock_timeout", None) try: self._handle_mock_testing_fallbacks( kwargs=kwargs, model_group=model_group, fallbacks=fallbacks, context_window_fallbacks=context_window_fallbacks, content_policy_fallbacks=content_policy_fallbacks, ) if mock_timeout is not None: response = await self.async_function_with_retries(*args, **kwargs, mock_timeout=mock_timeout) else: response = await self.async_function_with_retries(*args, **kwargs) if verbose_router_logger.isEnabledFor(logging.DEBUG): verbose_router_logger.debug("Async Response: %s", response) response = add_fallback_headers_to_response( response=response, attempted_fallbacks=0, ) return response except Exception as e: return await self.async_function_with_fallbacks_common_utils( e, disable_fallbacks, fallbacks, context_window_fallbacks, content_policy_fallbacks, model_group, args, kwargs, include_fallback_errors=include_fallback_errors, ) def _handle_mock_testing_fallbacks( self, kwargs: dict, model_group: str | None = None, fallbacks: list | None = None, context_window_fallbacks: list | None = None, content_policy_fallbacks: list | None = None, ): """ Helper function to raise a litellm Error for mock testing purposes. Raises: litellm.InternalServerError: when `mock_testing_fallbacks=True` passed in request params litellm.ContextWindowExceededError: when `mock_testing_context_fallbacks=True` passed in request params litellm.ContentPolicyViolationError: when `mock_testing_content_policy_fallbacks=True` passed in request params """ mock_testing_params: Final = MockRouterTestingParams.from_kwargs(kwargs) if ( mock_testing_params.mock_testing_fallbacks is not None and mock_testing_params.mock_testing_fallbacks is True ): raise litellm.InternalServerError( model=model_group, llm_provider="", message=f"This is a mock exception for model={model_group}, to trigger a fallback. Fallbacks={fallbacks}", ) elif ( mock_testing_params.mock_testing_context_fallbacks is not None and mock_testing_params.mock_testing_context_fallbacks is True ): raise litellm.ContextWindowExceededError( model=model_group, llm_provider="", message=f"This is a mock exception for model={model_group}, to trigger a fallback. \ Context_Window_Fallbacks={context_window_fallbacks}", ) elif ( mock_testing_params.mock_testing_content_policy_fallbacks is not None and mock_testing_params.mock_testing_content_policy_fallbacks is True ): raise litellm.ContentPolicyViolationError( model=model_group, llm_provider="", message=f"This is a mock exception for model={model_group}, to trigger a fallback. \ Context_Policy_Fallbacks={content_policy_fallbacks}", ) @staticmethod def _deployment_ids_to_skip_on_retry(exception: Exception, already_skipped: object) -> tuple[str, ...]: failed_deployment_id: Final[str | None] = getattr(exception, "retry_skip_deployment_id", None) or getattr( exception, "failed_deployment_id", None ) status_code: Final = getattr(exception, "status_code", None) if not failed_deployment_id or not isinstance(status_code, int): return () if litellm._should_retry(status_code): # pyright: ignore[reportPrivateUsage] # as in should_retry_this_error return () already_skipped_ids: Final = _as_retry_skipped_deployment_ids(already_skipped) skipped: Final = tuple(sorted(frozenset((*already_skipped_ids, failed_deployment_id)))) verbose_router_logger.debug( "Retry skips deployments that already answered %s to this request: %s", status_code, skipped ) return skipped @tracer.wrap() async def async_function_with_retries(self, *args, **kwargs): verbose_router_logger.debug("Inside async function with retries.") original_function: Final = kwargs.pop("original_function") fallbacks: Final = kwargs.pop("fallbacks", self.fallbacks) parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) context_window_fallbacks: Final = kwargs.pop("context_window_fallbacks", self.context_window_fallbacks) content_policy_fallbacks: Final = kwargs.pop("content_policy_fallbacks", self.content_policy_fallbacks) # Support per-request model_group_retry_policy override (from key/team settings) model_group_retry_policy: Final = kwargs.pop("model_group_retry_policy", self.model_group_retry_policy) model_group: Final[str | None] = kwargs.get("model") request_num_retries: Final[int | None] = kwargs.pop("num_retries", None) num_retries = request_num_retries if num_retries is None: # Fall back to the router setting (then 0) so the comparisons below never # hit `None > int`, which would mask the real upstream error with a TypeError. num_retries = self.num_retries if self.num_retries is not None else 0 ## ADD MODEL GROUP SIZE TO METADATA - used for model_group_rate_limit_error tracking _metadata: Final[dict] = kwargs.get("litellm_metadata", kwargs.get("metadata")) or {} if "model_group" in _metadata and isinstance(_metadata["model_group"], str): model_list: Final = self.get_model_list(model_name=_metadata["model_group"]) if model_list is not None: _metadata.update({"model_group_size": len(model_list)}) verbose_router_logger.debug( "async function w/ retries: original_function - %s, num_retries - %s", original_function, num_retries ) ## ADD RETRY TRACKING TO METADATA - used for spend logs retry tracking _metadata["attempted_retries"] = 0 _metadata["max_retries"] = num_retries # Updated after overrides in exception handler try: self._handle_mock_testing_rate_limit_error(model_group=model_group, kwargs=kwargs) # if the function call is successful, no exception will be raised and we'll break out of the loop response = await self.make_call(original_function, *args, **kwargs) response = add_retry_headers_to_response(response=response, attempted_retries=0, max_retries=None) return response except Exception as e: if is_guardrail_intervention(e): raise current_attempt = None original_exception = e deployment_num_retries: Final = getattr(e, "num_retries", None) if ( request_num_retries is None and deployment_num_retries is not None and isinstance(deployment_num_retries, int) ): num_retries = deployment_num_retries """ Retry Logic """ ( _healthy_deployments, _all_deployments, ) = await self._async_get_healthy_deployments( model=kwargs.get("model") or "", parent_otel_span=parent_otel_span, ) # Check retry policy FIRST, before should_retry_this_error # This allows retry policies to override the healthy deployments check _retry_policy_applies = False if request_num_retries != 0 and (self.retry_policy is not None or model_group_retry_policy is not None): # get num_retries from retry policy # Use the model_group captured at the start of the function, or get it from metadata # kwargs.get("model") at this point is the deployment model, not the model_group _model_group_for_retry_policy = model_group or _metadata.get("model_group") or kwargs.get("model") # Use per-request model_group_retry_policy if provided, otherwise use self _retry_policy_retries: Final = _get_num_retries_from_retry_policy( exception=original_exception, model_group=_model_group_for_retry_policy, model_group_retry_policy=model_group_retry_policy, retry_policy=self.retry_policy, ) if _retry_policy_retries is not None: num_retries = _retry_policy_retries _retry_policy_applies = True # raises an exception if this error should not be retries # Skip this check if retry policy applies (retry policy takes precedence) if not _retry_policy_applies: self.should_retry_this_error( error=e, healthy_deployments=_healthy_deployments, all_deployments=_all_deployments, context_window_fallbacks=context_window_fallbacks, regular_fallbacks=fallbacks, content_policy_fallbacks=content_policy_fallbacks, ) # Update max_retries after overrides (deployment_num_retries / retry_policy) _metadata["max_retries"] = num_retries ## LOGGING if num_retries > 0: kwargs = self.log_retry(kwargs=kwargs, e=original_exception) first_skipped_ids: Final = self._deployment_ids_to_skip_on_retry( exception=original_exception, already_skipped=kwargs.get("_retry_skipped_deployment_ids"), ) if first_skipped_ids: kwargs["_retry_skipped_deployment_ids"] = first_skipped_ids # rebind-ok: the next attempt reads it else: raise verbose_router_logger.debug("Retrying request with num_retries: %s", num_retries) # decides how long to sleep before retry retry_after: Final = self._time_to_sleep_before_retry( e=original_exception, remaining_retries=num_retries, num_retries=num_retries, healthy_deployments=_healthy_deployments, all_deployments=_all_deployments, ) await asyncio.sleep(retry_after) for current_attempt in range(num_retries): try: # Update retry tracking metadata before each retry attempt _metadata["attempted_retries"] = current_attempt + 1 _metadata["max_retries"] = num_retries # if the function call is successful, no exception will be raised and we'll break out of the loop response = await self.make_call(original_function, *args, **kwargs) if coroutine_checker.is_async_callable(response): # async errors are often returned as coroutines response = await response response = add_retry_headers_to_response( response=response, attempted_retries=current_attempt + 1, max_retries=num_retries, ) return response except Exception as e: # Always track the latest error so we raise the most # recent exception instead of the first one. original_exception = e ## LOGGING kwargs = self.log_retry(kwargs=kwargs, e=e) remaining_retries = num_retries - current_attempt - 1 _model: str | None = kwargs.get("model") if _model is not None: ( _healthy_deployments, _, ) = await self._async_get_healthy_deployments( model=_model, parent_otel_span=parent_otel_span, ) else: _healthy_deployments = [] # Check if this error is non-retryable (e.g., 400 context # window exceeded). If so, raise immediately instead of # continuing the retry loop. Respect retry policy # precedence - only check when no retry policy applies. if not _retry_policy_applies: try: self.should_retry_this_error( error=e, healthy_deployments=_healthy_deployments, all_deployments=_all_deployments, context_window_fallbacks=context_window_fallbacks, regular_fallbacks=fallbacks, content_policy_fallbacks=content_policy_fallbacks, ) except Exception: raise e skipped_ids = self._deployment_ids_to_skip_on_retry( exception=e, already_skipped=kwargs.get("_retry_skipped_deployment_ids"), ) if skipped_ids: kwargs["_retry_skipped_deployment_ids"] = skipped_ids # rebind-ok: the next attempt reads it _timeout = self._time_to_sleep_before_retry( e=e, remaining_retries=remaining_retries, num_retries=num_retries, healthy_deployments=_healthy_deployments, all_deployments=_all_deployments, ) await asyncio.sleep(_timeout) if type(original_exception) in litellm.LITELLM_EXCEPTION_TYPES: setattr(original_exception, "max_retries", num_retries) # current_attempt is 0-indexed (0 to num_retries-1), so after loop completion # it represents the last attempt index. The actual number of retries attempted # is current_attempt + 1, which equals num_retries when all retries are exhausted. # We've already verified num_retries > 0 before entering the loop, so current_attempt # will always be set (never None) when we reach this point. actual_retries_attempted: Final = current_attempt + 1 if current_attempt is not None else num_retries setattr(original_exception, "num_retries", actual_retries_attempted) raise original_exception async def make_call(self, original_function: Any, *args, **kwargs): """ Handler for making a call to the .completion()/.embeddings()/etc. functions. """ model_group: Final = kwargs.get("model") response = original_function(*args, **kwargs) if coroutine_checker.is_async_callable(response) or inspect.isawaitable(response): response = await response await self.increment_deployment_usage_for_response(response=response, request_kwargs=kwargs) ## PROCESS RESPONSE HEADERS response = await self.set_response_headers(response=response, model_group=model_group, request_kwargs=kwargs) return response def _handle_mock_testing_rate_limit_error(self, kwargs: dict, model_group: str | None = None): """ Helper function to raise a mock litellm.RateLimitError error for testing purposes. Raises: litellm.RateLimitError error when `mock_testing_rate_limit_error=True` passed in request params """ mock_testing_rate_limit_error: Final[bool | None] = kwargs.pop("mock_testing_rate_limit_error", None) available_models: Final = self.get_model_list(model_name=model_group) num_retries: int | None = None if available_models is not None and len(available_models) == 1: num_retries = cast(int | None, available_models[0]["litellm_params"].get("num_retries")) if mock_testing_rate_limit_error is not None and mock_testing_rate_limit_error is True: verbose_router_logger.info( "litellm.router.py::_mock_rate_limit_error() - Raising mock RateLimitError for model=%s", model_group ) raise litellm.RateLimitError( model=model_group, llm_provider="", message=f"This is a mock exception for model={model_group}, to trigger a rate limit error.", num_retries=num_retries, ) def should_retry_this_error( self, error: Exception, healthy_deployments: list | None = None, all_deployments: list | None = None, context_window_fallbacks: list | None = None, content_policy_fallbacks: list | None = None, regular_fallbacks: list | None = None, ): """ 1. raise an exception for ContextWindowExceededError if context_window_fallbacks is not None 2. raise an exception for ContentPolicyViolationError if content_policy_fallbacks is not None 2. raise an exception for RateLimitError if - there are no fallbacks - there are no healthy deployments in the same model group """ _num_healthy_deployments = 0 if healthy_deployments is not None and isinstance(healthy_deployments, list): _num_healthy_deployments = len(healthy_deployments) _num_all_deployments = 0 if all_deployments is not None and isinstance(all_deployments, list): _num_all_deployments = len(all_deployments) ### CHECK IF RATE LIMIT / CONTEXT WINDOW ERROR / CONTENT POLICY VIOLATION ERROR w/ fallbacks available / Bad Request Error if isinstance(error, litellm.ContextWindowExceededError) and context_window_fallbacks is not None: raise error if isinstance(error, litellm.ContentPolicyViolationError) and content_policy_fallbacks is not None: raise error status_code: Final = getattr(error, "status_code", None) if status_code is not None and not litellm._should_retry(status_code): # 401/403 are special cases - allow retry if multiple deployments exist (handled below) if status_code not in (401, 403): raise error if isinstance(error, litellm.NotFoundError): raise error # Error we should only retry if there are other deployments if isinstance(error, openai.RateLimitError): if ( _num_healthy_deployments <= 0 # if no healthy deployments and regular_fallbacks is not None # and fallbacks available and len(regular_fallbacks) > 0 ): raise error # then raise the error if isinstance(error, (openai.AuthenticationError, openai.PermissionDeniedError)): """ - if other deployments available -> retry - else -> raise error """ if _num_all_deployments <= 1: # if there is only 1 deployment for this model group then don't retry raise error # then raise error # Do not retry if there are no healthy deployments # just raise the error if _num_healthy_deployments <= 0: # if no healthy deployments raise error return True def function_with_fallbacks(self, *args, **kwargs): """ Sync wrapper for async_function_with_fallbacks Wrapped to reduce code duplication and prevent bugs. """ return run_async_function(self.async_function_with_fallbacks, *args, **kwargs) def _get_fallback_model_group_for_lookup_groups( self, fallbacks: list[dict[str, list[str]]], # mutable-ok: mirrors the shared resolver's contract lookup_groups: tuple[str, ...], ) -> list[str] | None: # mutable-ok: mirrors the shared resolver's contract fallback_model_group, _ = get_fallback_model_group_for_lookup_groups( fallbacks=fallbacks, lookup_groups=lookup_groups ) return fallback_model_group def _get_first_default_fallback(self) -> str | None: """ Returns the first model from the default_fallbacks list, if it exists. """ if self.fallbacks is None: return None for fallback in self.fallbacks: if isinstance(fallback, dict) and "*" in fallback: default_list = fallback["*"] if isinstance(default_list, list) and len(default_list) > 0: return default_list[0] return None def _time_to_sleep_before_retry( self, e: Exception, remaining_retries: int, num_retries: int, healthy_deployments: list | None = None, all_deployments: list | None = None, ) -> int | float: """ Calculate back-off, then retry It should instantly retry only when: 1. there are healthy deployments in the same model group 2. there are fallbacks for the completion call """ ## base case - single deployment if all_deployments is not None and len(all_deployments) == 1: pass elif healthy_deployments is not None and isinstance(healthy_deployments, list) and len(healthy_deployments) > 0: return 0 response_headers: httpx.Headers | None = None if hasattr(e, "response") and hasattr(e.response, "headers"): response_headers = e.response.headers if hasattr(e, "litellm_response_headers"): response_headers = e.litellm_response_headers if response_headers is not None: timeout = litellm._calculate_retry_after( remaining_retries=remaining_retries, max_retries=num_retries, response_headers=response_headers, min_timeout=self.retry_after, ) else: timeout = litellm._calculate_retry_after( remaining_retries=remaining_retries, max_retries=num_retries, min_timeout=self.retry_after, ) return timeout ### HELPER FUNCTIONS async def deployment_callback_on_success( self, kwargs, # kwargs to completion completion_response, # response from completion start_time, end_time, # start/end time ): """ Track remaining tpm/rpm quota for model in model_list """ try: # WS session wrappers fire with result=None; per-turn costs tracked by inner calls. if kwargs.get("call_type") in ("_aresponses_websocket", "_arealtime"): return if is_batch_retrieve_call_type(kwargs.get("call_type")): return standard_logging_object: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None) if standard_logging_object is None: raise ValueError("standard_logging_object is None") litellm_params: Final = kwargs["litellm_params"] metadata: Final = litellm_params.get("metadata") if metadata is None: return model_group: Final = metadata.get("model_group", None) model_info: Final = litellm_params.get("model_info", {}) or {} deployment_id: Final = model_info.get("id", None) if model_group is None or deployment_id is None or self.get_deployment(model_id=str(deployment_id)) is None: return # Always track deployment successes for cooldown logic, regardless of TPM/RPM limits increment_deployment_successes_for_current_minute( litellm_router_instance=self, deployment_id=str(deployment_id), ) total_tokens: Final[float] = standard_logging_object.get("total_tokens", 0) counted_tokens: Final = get_counted_usage_tokens(litellm_params) deployment_name: Final = metadata.get("deployment", None) return await self._increment_deployment_usage( deployment_id=str(deployment_id), deployment_name=deployment_name if isinstance(deployment_name, str) else None, model_group=model_group, total_tokens=total_tokens if counted_tokens is None else max(0, total_tokens - counted_tokens), rpm_increment=1 if counted_tokens is None else 0, parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs), ) except Exception as e: verbose_router_logger.debug( "litellm.router.Router::deployment_callback_on_success(): Exception occured - %s", e ) async def increment_deployment_usage_for_response( self, response: object, request_kwargs: dict[str, object], ) -> None: if response is None: return try: deployment_metadata: Final = find_deployment_metadata(request_kwargs) model_group: Final = request_kwargs.get("model") if deployment_metadata is None or not isinstance(model_group, str): return model_info: Final = deployment_metadata["model_info"] deployment_id: Final = model_info.get("id") if isinstance(model_info, dict) else None if deployment_id is None: return total_tokens: Final = response_total_token_count(response) deployment_name: Final = deployment_metadata.get("deployment") deployment_metadata[ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY] = total_tokens try: await self._increment_deployment_usage( deployment_id=str(deployment_id), deployment_name=deployment_name if isinstance(deployment_name, str) else None, model_group=model_group, total_tokens=total_tokens, rpm_increment=1, parent_otel_span=_get_parent_otel_span_from_kwargs(request_kwargs), ) except Exception: deployment_metadata.pop(ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY, None) raise except Exception as e: verbose_router_logger.debug( "litellm.router.Router::increment_deployment_usage_for_response(): Exception occured - %s", e ) async def _increment_deployment_usage( self, *, deployment_id: str, deployment_name: str | None, model_group: str, total_tokens: float, rpm_increment: int, parent_otel_span: Span | None, ) -> str | None: from litellm.types.caching import RedisPipelineIncrementOperation deployment_info: Final = self.get_deployment(model_id=deployment_id) if deployment_info is None: return None deployment_model_info: Final = self.get_router_model_info( deployment=deployment_info, received_model_name=model_group, ) configured_limits: Final = ( deployment_info.get("tpm", None), deployment_info.get("rpm", None), deployment_info.litellm_params.tpm, deployment_info.litellm_params.rpm, deployment_model_info.get("tpm", None), deployment_model_info.get("rpm", None), ) ## Nothing to track only when neither tpm/rpm nor itpm/otpm limits are ## set. IO deployments still record TPM/RPM usage here so TPM-aware ## routing strategies see their real load in mixed model groups; their ## itpm/otpm enforcement runs separately in ModelRateLimitingCheck. if all(limit is None for limit in configured_limits) and not deployment_has_io_token_limits( deployment_info.model_dump() ): return None if total_tokens <= 0 and rpm_increment <= 0: return None current_minute: Final = get_utc_datetime().strftime("%H-%M") # use the same timezone regardless of system clock tpm_key: Final = RouterCacheEnum.TPM.value.format( id=deployment_id, current_minute=current_minute, model=deployment_name ) rpm_key: Final = RouterCacheEnum.RPM.value.format( id=deployment_id, current_minute=current_minute, model=deployment_name ) pipeline_operations: Final[list[RedisPipelineIncrementOperation]] = [ RedisPipelineIncrementOperation(key=key, increment_value=increment_value, ttl=RoutingArgs.ttl.value) for key, increment_value in ((tpm_key, total_tokens), (rpm_key, rpm_increment)) ] post_increment_values: Final = await self.cache.async_increment_cache_pipeline( increment_list=pipeline_operations, parent_otel_span=parent_otel_span, ) if post_increment_values is not None and self.cache.redis_cache is not None: for operation, value in zip(pipeline_operations, post_increment_values): await self.cache.async_set_cache( operation["key"], int(value), local_only=True, ttl=RoutingArgs.ttl.value ) return tpm_key def sync_deployment_callback_on_success( self, kwargs, # kwargs to completion completion_response, # response from completion start_time, end_time, # start/end time ) -> str | None: """ Tracks the number of successes for a deployment in the current minute (using in-memory cache) Returns: - key: str - The key used to increment the cache - None: if no key is found """ if is_batch_retrieve_call_type(kwargs.get("call_type")): return None id = None if kwargs["litellm_params"].get("metadata") is None: pass else: model_group: Final = kwargs["litellm_params"]["metadata"].get("model_group", None) model_info: Final = kwargs["litellm_params"].get("model_info", {}) or {} id = model_info.get("id", None) if model_group is None or id is None: return None elif isinstance(id, int): id = str(id) if id is not None: key: Final = increment_deployment_successes_for_current_minute( litellm_router_instance=self, deployment_id=id, ) return key return None def deployment_callback_on_failure( self, kwargs, # kwargs to completion completion_response, # response from completion start_time, end_time, # start/end time ) -> bool: """ 2 jobs: - Tracks the number of failures for a deployment in the current minute (using in-memory cache) - Puts the deployment in cooldown if it exceeds the allowed fails / minute Returns: - True if the deployment should be put in cooldown - False if the deployment should not be put in cooldown """ verbose_router_logger.debug("Router: Entering 'deployment_callback_on_failure'") try: exception: Final = kwargs.get("exception", None) if is_advisor_orchestration_failure(exception): verbose_router_logger.debug( "Router: Exiting 'deployment_callback_on_failure' without cooldown. " "Failure originated from advisor orchestration, not the selected deployment." ) return False # Cache litellm_params to avoid repeated dict lookups litellm_params: Final = kwargs.get("litellm_params", {}) _model_info: Final = litellm_params.get("model_info", {}) if is_background_response_cost_poll_not_found(exception, litellm_params): verbose_router_logger.debug( "Router: Exiting 'deployment_callback_on_failure' without cooldown. " "Provider 404 came from the background response cost poll, not the deployment's health." ) return False exception_status: Final = getattr(exception, "status_code", "") if is_caller_timeout_408(kwargs, exception_status): verbose_router_logger.debug( "Router: Exiting 'deployment_callback_on_failure' without cooldown. " "A timeout the caller set caused this 408, not the deployment's health." ) return False exception_headers: Final = litellm.litellm_core_utils.exception_mapping_utils._get_response_headers( original_exception=exception ) # Determine cooldown time with priority: deployment config > response header > router default deployment_cooldown: Final = _first_present( _model_info if isinstance(_model_info, dict) else None, litellm_params, key="cooldown_time" ) header_cooldown = None if exception_headers is not None: header_cooldown = litellm.utils._get_retry_after_from_exception_header( response_headers=exception_headers ) ############################################## # Logic to determine cooldown time # 1. Check if a cooldown time is set in the deployment config # 2. Check if a cooldown time is set in the response header # 3. If no cooldown time is set, use the router default cooldown time ############################################## if deployment_cooldown is not None and deployment_cooldown >= 0: _time_to_cooldown = deployment_cooldown elif header_cooldown is not None and header_cooldown >= 0: _time_to_cooldown = header_cooldown else: _time_to_cooldown = self.cooldown_time if isinstance(_model_info, dict): deployment_id: Final[str | None] = _model_info.get("id") if deployment_id is None: return False increment_deployment_failures_for_current_minute( litellm_router_instance=self, deployment_id=deployment_id, ) result: Final = _set_cooldown_deployments( litellm_router_instance=self, exception_status=exception_status, original_exception=exception, deployment=deployment_id, time_to_cooldown=_time_to_cooldown, requested_model_group=(get_litellm_metadata_from_kwargs(kwargs) or {}).get("model_group"), ) # setting deployment_id in cooldown deployments return result else: verbose_router_logger.debug( "Router: Exiting 'deployment_callback_on_failure' without cooldown. No model_info found." ) return False except Exception as e: raise e async def async_deployment_callback_on_failure( self, kwargs, completion_response: object | None, start_time, end_time ): """ Update RPM usage for a deployment """ if is_batch_retrieve_call_type(kwargs.get("call_type")): return deployment_name: Final = kwargs["litellm_params"]["metadata"].get( "deployment", None ) # handles wildcard routes - by giving the original name sent to `litellm.completion` model_group: Final = kwargs["litellm_params"]["metadata"].get("model_group", None) model_info: Final = kwargs["litellm_params"].get("model_info", {}) or {} id = model_info.get("id", None) if model_group is None or id is None: return elif isinstance(id, int): id = str(id) parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) dt: Final = get_utc_datetime() current_minute: Final = dt.strftime("%H-%M") # use the same timezone regardless of system clock ## RPM rpm_key: Final = RouterCacheEnum.RPM.value.format(id=id, current_minute=current_minute, model=deployment_name) await self.cache.async_increment_cache( key=rpm_key, value=1, parent_otel_span=parent_otel_span, ttl=RoutingArgs.ttl.value, ) def _get_metadata_variable_name_from_kwargs(self, kwargs: dict) -> Literal["metadata", "litellm_metadata"]: """ Helper to return what the "metadata" field should be called in the request data - New endpoints return `litellm_metadata` - Old endpoints return `metadata` Context: - LiteLLM used `metadata` as an internal field for storing metadata - OpenAI then started using this field for their metadata - LiteLLM is now moving to using `litellm_metadata` for our metadata """ return get_metadata_variable_name_from_kwargs(kwargs) def log_retry(self, kwargs: dict, e: Exception) -> dict: """ When a retry or fallback happens, record which model group, deployment and attempt just failed and why, and count it toward the request-wide num_retries_per_request cap """ from litellm.types.router import RetryAttemptRecord _metadata_var: Final = "litellm_metadata" if "litellm_metadata" in kwargs else "metadata" request_metadata: Final[Mapping[str, object]] = kwargs[_metadata_var] model_group: Final = kwargs.get("model") model_info: Final = request_metadata.get("model_info") deployment_id: Final = model_info.get("id") if isinstance(model_info, Mapping) else None attempted_retries: Final = request_metadata.get("attempted_retries") attempt_record: Final[RetryAttemptRecord] = { "model_group": model_group if isinstance(model_group, str) else None, "deployment_id": deployment_id if isinstance(deployment_id, str) else None, "exception_type": type(e).__name__, "exception_string": str(e), "attempted_retries": attempted_retries if type(attempted_retries) is int else None, } earlier_breadcrumbs: Final = request_metadata.get("previous_models") kept_breadcrumbs: Final[tuple[object, ...]] = ( tuple(earlier_breadcrumbs)[-(RETRY_BREADCRUMB_LIMIT - 1) :] if isinstance(earlier_breadcrumbs, (list, tuple)) else () ) breadcrumbs: Final = (*kept_breadcrumbs, attempt_record) earlier: Final = request_metadata.get("request_retry_count") request_retry_count: Final = (earlier if type(earlier) is int and 0 <= earlier else 0) + 1 kwargs[_metadata_var]["previous_models"] = breadcrumbs # rebind-ok: the logging object already holds this dict kwargs[_metadata_var]["request_retry_count"] = request_retry_count # rebind-ok: same dict, read by the cap return kwargs def _update_usage(self, deployment_id: str, parent_otel_span: Span | None) -> int: """ Update deployment rpm for that minute Returns: - int: request count """ rpm_key: Final = deployment_id request_count = self.cache.get_cache(key=rpm_key, parent_otel_span=parent_otel_span, local_only=True) if request_count is None: request_count = 1 self.cache.set_cache(key=rpm_key, value=request_count, local_only=True, ttl=60) # only store for 60s else: request_count += 1 self.cache.set_cache(key=rpm_key, value=request_count, local_only=True) # don't change existing ttl return request_count def _has_default_fallbacks(self) -> bool: if self.fallbacks is None: return False for fallback in self.fallbacks: if isinstance(fallback, dict): if "*" in fallback: return True return False def _has_content_policy_fallback(self, model_group: str, kwargs: Mapping[str, Any]) -> bool: """ Whether a content-policy fallback would resolve for this request, keyed the same way async_function_with_fallbacks_common_utils resolves it: the tier a pre-routing hook selected wins over the requested group. Raising without this returning True would turn a deliverable response into an error the fallback chain cannot recover from. """ content_policy_fallbacks: Final = kwargs.get("content_policy_fallbacks", self.content_policy_fallbacks) if content_policy_fallbacks is not None: return has_unattempted_fallback_target( self._get_fallback_model_group_for_lookup_groups( fallbacks=content_policy_fallbacks, lookup_groups=fallback_lookup_groups(kwargs, model_group), ), kwargs, ) if self._has_default_fallbacks(): return True verbose_router_logger.debug( "No content-policy fallback available. Returning original response. model=%s, content_policy_fallbacks=%s", model_group, content_policy_fallbacks, ) return False def _refusal_fallback_available(self, model_group: str, kwargs: Mapping[str, Any]) -> bool: """ Whether a safeguard refusal can actually be recovered by the dispatcher. A configured content-policy list is authoritative; with none configured at all, the dispatcher falls through to the generic fallbacks lookup, so the gate mirrors that reachability and arms on a resolving generic chain (tier first, then the requested group, then "*"). """ if fallbacks_disabled_for_request(kwargs): return False content_policy_fallbacks: Final = kwargs.get("content_policy_fallbacks", self.content_policy_fallbacks) if content_policy_fallbacks is not None: return self._has_content_policy_fallback(model_group, kwargs) if self._has_default_fallbacks(): return True fallbacks: Final = kwargs.get("fallbacks", self.fallbacks) if fallbacks is None: return False resolved, _ = get_fallback_model_group_for_lookup_groups( fallbacks=fallbacks, lookup_groups=fallback_lookup_groups(kwargs, model_group), ) return has_unattempted_fallback_target(resolved, kwargs) def _anthropic_messages_order_levels(self, model_group: str, kwargs: Mapping[str, Any]) -> tuple[int, ...]: """ The distinct deployment order levels the fallback dispatcher would see for this request, computed the same way: the tier a pre-routing hook selected wins over the requested group. """ request_team_id: Final[str | None] = (kwargs.get("metadata", {}) or {}).get("user_api_key_team_id") order_model_group: Final = get_pre_routing_selection(kwargs) or model_group all_deployments: Final = self.get_model_list(model_name=order_model_group, team_id=request_team_id) or () return tuple( sorted( { litellm.utils._get_deployment_order(d) for d in all_deployments if litellm.utils._get_deployment_order(d) is not None } ) ) def _anthropic_messages_stream_can_fall_back(self, model_group: str, kwargs: Mapping[str, Any]) -> bool: """ Whether async_function_with_fallbacks_common_utils could still route a MidStreamFallbackError somewhere for this request (order levels, weighted failover, content-policy or generic fallbacks), which is the only case where holding lifecycle frames back from the client buys a clean retry. Errs toward True whenever a dispatcher path might reach a fallback. """ if fallbacks_disabled_for_request(kwargs): return False if self.enable_weighted_failover: return True order_levels: Final = self._anthropic_messages_order_levels(model_group, kwargs) if len(order_levels) > 1: current_target: Final = kwargs.get("_target_order") skip_up_to: Final = current_target if current_target is not None else order_levels[0] if any(o > skip_up_to for o in order_levels): return True content_policy_fallbacks: Final = kwargs.get("content_policy_fallbacks", self.content_policy_fallbacks) if content_policy_fallbacks is not None and self._has_content_policy_fallback(model_group, kwargs): return True fallbacks: Final = kwargs.get("fallbacks", self.fallbacks) if not fallbacks: return False if _check_non_standard_fallback_format(fallbacks=fallbacks): return True resolved, _ = get_fallback_model_group_for_lookup_groups( fallbacks=fallbacks, lookup_groups=fallback_lookup_groups(kwargs, model_group), ) return has_unattempted_fallback_target(resolved, kwargs) def _should_raise_content_policy_error(self, model: str, response: ModelResponse, kwargs: dict) -> bool: """ Determines if a content policy error should be raised. Only raised if a fallback is available. Else, original response is returned. """ if response.choices and len(response.choices) > 0: if response.choices[0].finish_reason != "content_filter": return False return self._has_content_policy_fallback(model, kwargs) def _should_raise_anthropic_refusal_error( self, model: str, original_generic_function: Callable, response: object, kwargs: Mapping[str, object] ) -> bool: """ The /v1/messages twin of _should_raise_content_policy_error: an Anthropic safeguard refusal (stop_reason "refusal" carrying stop_details) re-enters the fallback chain only when a content-policy fallback is configured; a plain refusal without stop_details, or any response with nothing configured, is returned to the client unchanged. """ from litellm.llms.anthropic.pass_through.messages.utils import ( get_safeguard_refusal_stop_details, ) if getattr(original_generic_function, "__name__", "") != "anthropic_messages": return False if get_safeguard_refusal_stop_details(response) is None: return False return self._refusal_fallback_available(model, kwargs) def _get_healthy_deployments(self, model: str, parent_otel_span: Span | None): _all_deployments: list = [] try: _, _all_deployments = self._common_checks_available_deployment( model=model, ) if isinstance(_all_deployments, dict): return [] except Exception: pass unhealthy_deployments: Final = _get_cooldown_deployments( litellm_router_instance=self, parent_otel_span=parent_otel_span ) unhealthy_set: Final = set(unhealthy_deployments) healthy_deployments: list = [d for d in _all_deployments if d["model_info"]["id"] not in unhealthy_set] healthy_deployments = self._filter_blocked_deployments(healthy_deployments) return healthy_deployments, _all_deployments async def _async_get_healthy_deployments( self, model: str, parent_otel_span: Span | None ) -> tuple[list[dict], list[dict]]: """ Returns Tuple of: - Tuple[List[Dict], List[Dict]]: 1. healthy_deployments: list of healthy deployments 2. all_deployments: list of all deployments """ _all_deployments: list = [] try: _, _all_deployments = self._common_checks_available_deployment( model=model, ) if isinstance(_all_deployments, dict): return [], _all_deployments except Exception: pass unhealthy_deployments: Final = await _async_get_cooldown_deployments( litellm_router_instance=self, parent_otel_span=parent_otel_span ) # Convert to set for O(1) lookup instead of O(n) unhealthy_deployments_set: Final = set(unhealthy_deployments) healthy_deployments: list = [ d for d in _all_deployments if d["model_info"]["id"] not in unhealthy_deployments_set ] healthy_deployments = self._filter_blocked_deployments(healthy_deployments) return healthy_deployments, _all_deployments def routing_strategy_pre_call_checks(self, deployment: dict): """ Mimics 'async_routing_strategy_pre_call_checks' Ensures consistent update rpm implementation for 'usage-based-routing-v2' Returns: - None Raises: - Rate Limit Exception - If the deployment is over it's tpm/rpm limits """ for _callback in litellm.callbacks: if isinstance(_callback, CustomLogger): _callback.pre_call_check(deployment) async def async_routing_strategy_pre_call_checks( self, deployment: dict, parent_otel_span: Span | None, logging_obj: LiteLLMLogging | None = None, ): """ For usage-based-routing-v2, enables running rpm checks before the call is made, inside the semaphore. -> makes the calls concurrency-safe, when rpm limits are set for a deployment Returns: - None Raises: - Rate Limit Exception - If the deployment is over it's tpm/rpm limits """ for _callback in litellm.callbacks: if isinstance(_callback, CustomLogger): try: await _callback.async_pre_call_check(deployment, parent_otel_span) except litellm.RateLimitError as e: ## LOG FAILURE EVENT if logging_obj is not None: asyncio.create_task( logging_obj.dispatch_failure_handlers( exception=e, traceback_exception=traceback.format_exc(), prefer_async_handlers=True, ) ) _set_cooldown_deployments( litellm_router_instance=self, exception_status=e.status_code, original_exception=e, deployment=deployment["model_info"]["id"], time_to_cooldown=self.cooldown_time, ) raise e except Exception as e: ## LOG FAILURE EVENT if logging_obj is not None: asyncio.create_task( logging_obj.dispatch_failure_handlers( exception=e, traceback_exception=traceback.format_exc(), prefer_async_handlers=True, ) ) raise e @contextlib.asynccontextmanager async def _deployment_slot( self, deployment: dict, kwargs: Mapping[str, object], parent_otel_span: Span | None ) -> AsyncGenerator[None, None]: """Holds the deployment's max_parallel_requests slot, if it has one, around the provider call. Routing strategy pre-call checks run inside the slot so their rpm accounting stays concurrency-safe.""" max_parallel_requests_limit: Final = self._get_client( deployment=deployment, kwargs=kwargs, client_type="max_parallel_requests", ) async with contextlib.AsyncExitStack() as slot: if isinstance(max_parallel_requests_limit, MaxParallelRequestsLimit): slot.enter_context(max_parallel_requests_limit) await self.async_routing_strategy_pre_call_checks(deployment=deployment, parent_otel_span=parent_otel_span) yield async def async_callback_filter_deployments( self, model: str, healthy_deployments: list[dict], messages: list[AllMessageValues] | None, parent_otel_span: Span | None, request_kwargs: dict | None = None, logging_obj: LiteLLMLogging | None = None, ): """ For usage-based-routing-v2, enables running rpm checks before the call is made, inside the semaphore. -> makes the calls concurrency-safe, when rpm limits are set for a deployment Returns: - None Raises: - Rate Limit Exception - If the deployment is over it's tpm/rpm limits """ returned_healthy_deployments = healthy_deployments for _callback in litellm.callbacks: if isinstance(_callback, CustomLogger): try: returned_healthy_deployments = await _callback.async_filter_deployments( model=model, healthy_deployments=returned_healthy_deployments, messages=messages, request_kwargs=request_kwargs, parent_otel_span=parent_otel_span, ) except Exception as e: ## LOG FAILURE EVENT if logging_obj is not None: asyncio.create_task( logging_obj.dispatch_failure_handlers( exception=e, traceback_exception=traceback.format_exc(), prefer_async_handlers=True, ) ) raise e return returned_healthy_deployments @staticmethod def _json_default_stable_id(value: object) -> str: """json.dumps default= for generate_model_id: plain str() on an arbitrary object (e.g. a RoutingPlugin instance) falls back to object.__repr__'s ``, so the hash -- and deployment id -- would change every restart. Use the class name instead, stable across restarts.""" return f"{type(value).__module__}.{type(value).__qualname__}" @staticmethod def generate_model_id(model_group: str, litellm_params: dict) -> str: # mutable-ok: hashed read-only """ Helper function to consistently generate the same id for a deployment - create a string from all the litellm params - hash - use hash as id """ # Optimized: Use list and join instead of string concatenation in loop # This avoids creating many temporary string objects (O(n) vs O(n²) complexity) parts: Final = [model_group] for k, v in litellm_params.items(): if isinstance(k, str): parts.append(k) elif isinstance(k, dict): parts.append(json.dumps(k, default=Router._json_default_stable_id)) else: parts.append(str(k)) if isinstance(v, str): parts.append(v) elif isinstance(v, dict): parts.append(json.dumps(v, default=Router._json_default_stable_id)) else: parts.append(str(v)) concat_str: Final = "".join(parts) hash_object: Final = hashlib.sha256(concat_str.encode()) return hash_object.hexdigest() @staticmethod def _inherit_builtin_cache_pricing(model_info: dict, backend_model: str, custom_llm_provider: str | None) -> None: """Fill missing cache pricing on a custom-priced deployment entry from the backend model's built-in cost map entry, so a deployment that only spells out ``input_cost_per_token``/``output_cost_per_token`` does not silently bill cache_read/cache_creation at 0. User-specified cache fields always win; only ``None``/missing entries are inherited. No-op when the backend model has no canonical entry. """ cache_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", ) if all(model_info.get(f) is not None for f in cache_fields): return try: backend_info: Final = litellm.get_model_info(model=backend_model, custom_llm_provider=custom_llm_provider) except Exception: return for field in cache_fields: if model_info.get(field) is None: backend_value = backend_info.get(field) if backend_value is not None: model_info[field] = backend_value @staticmethod def _inherit_builtin_base_rates_for_off_peak( model_info: dict, # mutable-ok: cost-map entry filled in place backend_model: str, custom_llm_provider: str | None, ) -> None: """Fill missing pricing fields on a deployment entry that only sets ``off_peak_pricing``, from the backend model's built-in cost map entry. Cost lookup selects the deployment-scoped entry over the shared backend entry only when the deployment entry carries a base pricing field, and ``off_peak_pricing`` is deliberately kept off the shared entry, so a deployment spelling out only its off-peak schedule would otherwise never receive the discount. The backend model's entire canonical cost map entry is copied, field by field, so threshold, tiered, service-tier, cache, character, and per-second rates as well as companion billing fields like ``web_search_billing_unit`` and the regional uplift multipliers all carry over, and peak-hour billing through the deployment entry matches the shared backend entry exactly. The raw ``litellm.model_cost`` entry is the copy source rather than ``get_model_info``'s view of it, since that view synthesizes zero flat token rates for backends without one and storing those would mark a tiered-only backend explicitly priced free. Values are deep-copied to keep the builtin entry isolated. User-specified fields always win; no-op when any base pricing field is already set or the backend model has no canonical entry. """ if not model_info.get("off_peak_pricing"): return if any( model_info.get(field) is not None for field in ("input_cost_per_token", "input_cost_per_second", "cost_per_second", "tiered_pricing") ): return try: backend_info: Final = litellm.get_model_info(model=backend_model, custom_llm_provider=custom_llm_provider) except Exception: # noqa: BLE001 # get_model_info raises plain Exception for an unmapped backend model return backend_entry: Final = litellm.model_cost.get(backend_info.get("key") or "") if not isinstance(backend_entry, dict): return for field, backend_value in backend_entry.items(): if model_info.get(field) is not None or backend_value is None: continue model_info[field] = copy.deepcopy(backend_value) @staticmethod def _inherit_builtin_tiered_output_rate( model_info: dict, backend_model: str, custom_llm_provider: str | None ) -> None: """Fill a missing entry-level output rate on a deployment entry whose tier table omits one, from the backend model's built-in cost map entry. A deployment's custom pricing is registered as its own standalone ``litellm.model_cost`` entry holding only the supplied fields, and the tiered-cost output fallback reads that same entry, so a tier table that spells out only input-side rates would bill every completion at 0. A user-specified ``output_cost_per_token`` always wins. No-op without a tier table, when every tier declares its own output rate, or when the backend model has no canonical entry or no flat output rate: ``get_model_info`` synthesizes a zero for tiered-only backends, and storing that zero would mark the deployment as explicitly priced free. """ tiers: Final = model_info.get("tiered_pricing") if not isinstance(tiers, list) or not tiers: return if model_info.get("output_cost_per_token") is not None: return if all(isinstance(tier, dict) and "output_cost_per_token" in tier for tier in tiers): return try: backend_info: Final = litellm.get_model_info(model=backend_model, custom_llm_provider=custom_llm_provider) except Exception: # noqa: BLE001 # get_model_info raises plain Exception for an unmapped backend model return backend_rate: Final = backend_info.get("output_cost_per_token") if backend_rate: model_info["output_cost_per_token"] = backend_rate def _create_deployment( self, deployment_info: dict, _model_name: str, _litellm_params: dict, _model_info: dict, *, declared_id: str | None = None, duplicate_ids: frozenset[str] = frozenset(), ) -> Deployment | None: """ Create a deployment object and add it to the model list If the deployment is not active for the current environment, it is ignored Returns: - Deployment: The deployment object - None: If the deployment is not active for the current environment (if 'supported_environments' is set in litellm_params) """ try: config_sourced: Final = _model_info.get("db_model") is not True identity_error: Final = ( ptu_identity_error( declared_id=declared_id, taken=declared_id in duplicate_ids, current_id=_model_info.get("id"), model_name=_model_name, ) if config_sourced and ptu_terms(_model_info) is not None else None ) ptu_error: Final = ( (ptu_config_error(_model_info, model_name=_model_name) or identity_error) if config_sourced else None ) if ptu_error is not None and is_ptu_cost_attribution_enabled(): raise ValueError(ptu_error) access_windows_error: Final = access_windows_config_error(_model_info, model_name=_model_name) if access_windows_error is not None: raise ValueError(access_windows_error) zeroed_pricing: Final = zeroed_ptu_pricing(_model_info, _litellm_params) if config_sourced else None litellm_params: Final[LiteLLM_Params] = LiteLLM_Params( **( # pyright: ignore[reportArgumentType] # untyped merged dict; already true for every field here _litellm_params if zeroed_pricing is None else MappingProxyType({**_litellm_params, **zeroed_pricing}) ) ) warn_on_provider_credential_mismatch(model_name=_model_name, litellm_params=_litellm_params) deployment = Deployment( **deployment_info, model_name=_model_name, litellm_params=litellm_params, model_info=_model_info, ) for field in CustomPricingLiteLLMParams.model_fields: if deployment.litellm_params.get(field) is not None: _model_info[field] = deployment.litellm_params[field] Router._inherit_builtin_base_rates_for_off_peak( model_info=_model_info, backend_model=deployment.litellm_params.model, custom_llm_provider=deployment.litellm_params.custom_llm_provider, ) if _model_info.get("input_cost_per_token") is not None: Router._inherit_builtin_cache_pricing( model_info=_model_info, backend_model=deployment.litellm_params.model, custom_llm_provider=deployment.litellm_params.custom_llm_provider, ) Router._inherit_builtin_tiered_output_rate( model_info=_model_info, backend_model=deployment.litellm_params.model, custom_llm_provider=deployment.litellm_params.custom_llm_provider, ) ## REGISTER MODEL INFO IN LITELLM MODEL COST MAP Router._register_deployment_in_model_cost( model_id=deployment.model_info.id, model_info=_model_info, model=deployment.litellm_params.model, custom_llm_provider=deployment.litellm_params.custom_llm_provider, ) ## Check if LLM Deployment is allowed for this deployment if self.deployment_is_active_for_environment(deployment=deployment) is not True: verbose_router_logger.warning( "Ignoring deployment %s as it is not active for environment %s", deployment.model_name, deployment.model_info["supported_environments"], ) return None # Validate tag_regex patterns BEFORE adding the deployment so we never # have partially-initialised router state if a pattern is invalid. _tag_regex: Final = deployment.litellm_params.get("tag_regex") or [] for pattern in _tag_regex: try: re.compile(pattern) except re.error as exc: raise ValueError( f"Invalid regex in tag_regex for model '{deployment.model_name}': {pattern!r} — {exc}" ) from exc deployment = self._add_deployment(deployment=deployment) model: Final = deployment.to_json(exclude_none=True) self._add_model_to_list_and_index_map(model=model, model_id=deployment.model_info.id) return deployment except Exception as e: if self.ignore_invalid_deployments: if isinstance(e, litellm.BadRequestError): self._provider_unresolved_deployments = ( *self._provider_unresolved_deployments, partial( self._create_deployment, deployment_info=deployment_info, _model_name=_model_name, _litellm_params=_litellm_params, _model_info=_model_info, declared_id=declared_id, duplicate_ids=duplicate_ids, ), ) verbose_router_logger.exception( "Error creating deployment: %s, ignoring and continuing with other deployments.", e ) return None else: raise e def _is_auto_router_deployment(self, litellm_params: LiteLLM_Params) -> bool: """ Check if the deployment is an auto-router deployment (semantic router). Returns True if the litellm_params model starts with "auto_router/" but NOT "auto_router/complexity_router" or "auto_router/adaptive_router" (which use the complexity-router and adaptive-router strategies). """ return classify_strategy_router_model(litellm_params.model) == "semantic" @staticmethod def _deployment_tags(deployment: Deployment) -> tuple[str, ...]: """Deployment tags used to disambiguate strategy registries keyed by model_name.""" return tuple(deployment.litellm_params.tags or ()) def init_auto_router_deployment(self, deployment: Deployment): """ Initialize the auto-router deployment. This will initialize the auto-router and add it to the auto-routers dictionary. """ from litellm.router_strategy.auto_router.auto_router import AutoRouter auto_router_config_path: Final[str | None] = deployment.litellm_params.auto_router_config_path auto_router_config: Final[str | None] = deployment.litellm_params.auto_router_config if auto_router_config_path is None and auto_router_config is None: raise ValueError( "auto_router_config_path or auto_router_config is required for auto-router deployments. Please set it in the litellm_params" ) default_model: Final[str | None] = deployment.litellm_params.auto_router_default_model if default_model is None: raise ValueError( "auto_router_default_model is required for auto-router deployments. Please set it in the litellm_params" ) embedding_model: Final[str | None] = deployment.litellm_params.auto_router_embedding_model if embedding_model is None: raise ValueError( "auto_router_embedding_model is required for auto-router deployments. Please set it in the litellm_params" ) max_input_chars: Final[int | None] = deployment.litellm_params.auto_router_max_input_chars autor_router: Final[AutoRouter] = AutoRouter( model_name=deployment.model_name, auto_router_config_path=auto_router_config_path, auto_router_config=auto_router_config, default_model=default_model, embedding_model=embedding_model, litellm_router_instance=self, max_input_chars=(max_input_chars if max_input_chars is not None else DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS), ) self._register_pre_routing_strategy( registry=self.auto_routers, deployment=deployment, strategy=autor_router, strategy_label="Auto-router", ) def _is_complexity_router_deployment(self, litellm_params: LiteLLM_Params) -> bool: """ Check if the deployment is a complexity-router deployment. Returns True if the litellm_params model starts with "auto_router/complexity_router" """ return classify_strategy_router_model(litellm_params.model) == "complexity" def config_deployments(self) -> Iterator[Mapping[str, object]]: """The model_list rows that came from config.yaml rather than the DB (``model_info.db_model`` unset).""" for deployment in self.model_list: if not isinstance(deployment, Mapping): continue model_info = deployment.get("model_info") if not (isinstance(model_info, Mapping) and model_info.get("db_model")): yield deployment def auto_router_capability_violation(self, capability: GatedAutoRouterCapability) -> str | None: """ Why one more router claiming ``capability`` cannot join this router, or None when it can. Judged against every deployment currently on the model_list; an upsert pops the row being edited first, so an edit of an existing gated router keeps its own slot. The limit is resolved on every call through ``auto_router_capability_limit``; unset means unlimited, which is the SDK default, and the proxy injects a resolver backed by its license. """ limit: Final = self.auto_router_capability_limit() if self.auto_router_capability_limit is not None else None others: Final = count_capability_routers( (deployment for deployment in self.model_list if isinstance(deployment, Mapping)), capability=capability, ) return capability_limit_violation(capability=capability, held=others + 1, limit=limit) def init_complexity_router_deployment(self, deployment: Deployment): """ Initialize the complexity-router deployment. This will initialize the complexity-router and add it to the complexity-routers dictionary. """ # Import here to avoid circular imports — ComplexityRouter is a CustomLogger # subclass that imports litellm internals which depend on router.py. # This matches the AutoRouter pattern in init_auto_router_deployment above. from litellm.router_strategy.complexity_router.complexity_router import ( ComplexityRouter, ) from litellm.router_strategy.complexity_router.config import ( ComplexityRouterConfig, ) complexity_router_config: Final[dict | None] = deployment.litellm_params.complexity_router_config capability: Final = claimed_capability(complexity_router_config) if capability is not None: limit_violation: Final = self.auto_router_capability_violation(capability) if limit_violation is not None: raise ValueError(limit_violation) default_model: str | None = deployment.litellm_params.complexity_router_default_model # If no default model specified, try to get from config tiers. Derived from the # validated model, not the raw dict, so normalization (e.g. fallback_tier # whitespace) is applied by its one owner before the tiers lookup. if default_model is None and complexity_router_config: validated: Final = ComplexityRouterConfig.model_validate(complexity_router_config) # Custom tier sets name their fallback tier; built-in sets default to MEDIUM or SIMPLE derived: Final = ( (validated.tiers.get(validated.fallback_tier) if validated.fallback_tier is not None else None) or validated.tiers.get("MEDIUM") or validated.tiers.get("SIMPLE") ) if isinstance(derived, list): default_model = derived[0] if derived else None else: default_model = derived if default_model is None: raise ValueError( "complexity_router_default_model is required for complexity-router deployments, " "or configure tiers in complexity_router_config. Please set it in the litellm_params" ) complexity_router: Final[ComplexityRouter] = ComplexityRouter( model_name=deployment.model_name, default_model=default_model, litellm_router_instance=self, complexity_router_config=complexity_router_config, ) self._register_pre_routing_strategy( registry=self.complexity_routers, deployment=deployment, strategy=complexity_router, strategy_label="Complexity-router", ) if complexity_router._uses_deployment_pin: self._ensure_deployment_affinity_callback() def _is_adaptive_router_deployment(self, litellm_params: LiteLLM_Params) -> bool: """True when this deployment opts in via the `auto_router/adaptive_router` model prefix.""" return classify_strategy_router_model(litellm_params.model) == "adaptive" def _deployment_participates_in_adaptive_routing(self, litellm_params: LiteLLM_Params) -> bool: """True when this deployment owns an `adaptive_routers` entry once finalized: a dedicated adaptive router, or a complexity router whose config enables the adaptive companion. Mirrors the two arms of `_finalize_adaptive_router_if_configured`, which is the registry's only writer.""" if self._is_adaptive_router_deployment(litellm_params=litellm_params): return True if not self._is_complexity_router_deployment(litellm_params=litellm_params): return False config: Final = litellm_params.complexity_router_config if not config: return False adaptive_flag: Final[object] = config.get("adaptive") return bool(adaptive_flag) @staticmethod def _has_registered_strategy( registry: dict[str, list[TaggedPreRoutingStrategy[_PreRoutingStrategyT]]], model_name: str, tags: tuple[str, ...], ) -> bool: """True when a strategy for this (model_name, tags) pair is already registered.""" return any(existing.tags == tags for existing in registry.get(model_name, [])) def _register_pre_routing_strategy( self, registry: dict[str, list[TaggedPreRoutingStrategy[_PreRoutingStrategyT]]], deployment: Deployment, strategy: _PreRoutingStrategyT, strategy_label: str, ) -> None: """ Register `strategy` under `deployment.model_name`, scoped by its tags. Reusing a `model_name` is allowed when tags differ; a repeat of the same (model_name, tags) pair is a misconfiguration and is rejected. """ tags: Final = self._deployment_tags(deployment) if self._has_registered_strategy(registry, deployment.model_name, tags): raise ValueError( f"{strategy_label} deployment {deployment.model_name} with tags {list(tags)} already exists. " "Please use a different model name or set different tags." ) registry[deployment.model_name] = [ *registry.get(deployment.model_name, []), TaggedPreRoutingStrategy(tags=tags, strategy=strategy), ] @staticmethod def _unregister_pre_routing_strategy( registry: dict[str, list[TaggedPreRoutingStrategy[_PreRoutingStrategyT]]], model_name: str, tags: tuple[str, ...], ) -> bool: """Drop the strategy registered for this exact (model_name, tags) pair, leaving strategies registered under the same name with different tags in place. Returns whether anything was actually dropped.""" existing: Final = registry.get(model_name, []) remaining: Final = [entry for entry in existing if entry.tags != tags] if len(remaining) == len(existing): return False if remaining: registry[model_name] = remaining else: registry.pop(model_name, None) return True def _unregister_pre_routing_strategy_for_deployment(self, deployment: Deployment) -> None: """ Release the pre-routing strategy a deployment holds, so removing it from the model_list also frees its (model_name, tags) slot. Without this, re-adding the deployment (an edit arriving via upsert_deployment, or a router recreated under a name that was deleted earlier) hits the "already exists" guard in `_register_pre_routing_strategy`, which `ignore_invalid_deployments` swallows - the deployment then silently never makes it back into the model_list. Released from every registry rather than the first match, because registration is one-to-many: a complexity router configured with `adaptive` is also registered in `adaptive_routers` under the same (model_name, tags) by the deferred finalize pass. Guarded on the auto_router/ prefix so removing a *regular* deployment can't evict a router that merely shares its model_name. """ if not deployment.litellm_params.model.startswith("auto_router/"): return model_name: Final = deployment.model_name tags: Final = self._deployment_tags(deployment) for registry in (self.auto_routers, self.complexity_routers, self.quality_routers): self._unregister_pre_routing_strategy(registry, model_name, tags) if self._unregister_pre_routing_strategy(self.adaptive_routers, model_name, tags): self._sync_adaptive_router_hooks() def _finalize_adaptive_router_if_configured(self) -> None: """Locate every adaptive-router deployment in the finalized model_list and build an AdaptiveRouter for each. Safe no-op when none are configured. Idempotent: skips any deployment whose (model_name, tags) pair is already initialized, so hot-reloads don't rebuild routers that would lose state.""" for entry in self.model_list or []: lp = entry.get("litellm_params") if isinstance(entry, dict) else entry.litellm_params lp_model = (lp.get("model") if isinstance(lp, dict) else lp.model) if lp else None if not (lp_model and lp_model.startswith("auto_router/adaptive_router")): continue model_name = entry.get("model_name") if isinstance(entry, dict) else entry.model_name if not model_name or not lp: continue deployment = Deployment( model_name=model_name, litellm_params=(lp if not isinstance(lp, dict) else LiteLLM_Params(**lp)), model_info=(entry.get("model_info") if isinstance(entry, dict) else entry.model_info), ) if self._has_registered_strategy(self.adaptive_routers, model_name, self._deployment_tags(deployment)): continue self.init_adaptive_router_deployment(deployment=deployment) for model_name, tagged_complexity_routers in self.complexity_routers.items(): for tagged in tagged_complexity_routers: complexity_router = tagged.strategy if not complexity_router.config.adaptive: continue if self._has_registered_strategy(self.adaptive_routers, model_name, tagged.tags): continue adaptive_router: AdaptiveRouter | None = complexity_router._ensure_adaptive_router() if adaptive_router is not None: self.adaptive_routers[model_name] = [ *self.adaptive_routers.get(model_name, []), TaggedPreRoutingStrategy(tags=tagged.tags, strategy=adaptive_router), ] self._sync_adaptive_router_hooks() def _sync_adaptive_router_hooks(self) -> None: """Rebuild the AdaptiveRouterPostCallHook set so it is exactly one hook per currently registered adaptive router. Run at every point the adaptive registry changes, otherwise a released router keeps recording turns through its hook.""" from litellm.router_strategy.adaptive_router.hooks import ( AdaptiveRouterPostCallHook, ) for callback in litellm.logging_callback_manager.get_custom_loggers_for_type(AdaptiveRouterPostCallHook): litellm.logging_callback_manager.remove_callback_from_all_lists(callback) for tagged_adaptive_routers in self.adaptive_routers.values(): for tagged in tagged_adaptive_routers: litellm.logging_callback_manager.add_litellm_callback( AdaptiveRouterPostCallHook(adaptive_router=tagged.strategy) ) def init_adaptive_router_deployment(self, deployment: Deployment) -> None: """ Build an AdaptiveRouter instance for this deployment and register its post-call hook. Multiple adaptive routers can coexist on a single Router, keyed by `deployment.model_name`. `model_to_prefs` and `model_to_cost` are derived from the OTHER models already registered in `self.model_list` whose `model_name` appears in `available_models`. Models not yet registered fall back to defaults. """ # Local import: AdaptiveRouter -> hooks -> classifier all import litellm # internals which transitively import this module. (AGENTS.md exception clause.) from litellm.router_strategy.adaptive_router.adaptive_router import ( AdaptiveRouter, ) from litellm.router_strategy.adaptive_router.hooks import ( AdaptiveRouterPostCallHook, ) from litellm.types.router import ( AdaptiveRouterConfig, AdaptiveRouterPreferences, ) raw_config: Final = deployment.litellm_params.adaptive_router_config if raw_config is None: raise ValueError("adaptive_router_config is required for adaptive-router deployments.") config: Final = AdaptiveRouterConfig(**raw_config) model_to_prefs: Final[dict[str, AdaptiveRouterPreferences]] = {} model_to_cost: Final[dict[str, float]] = {} # O(k) via the name→indices map: only touch deployments whose name # is listed in `available_models`, instead of scanning model_list. for name in config.available_models: indices = self.model_name_to_deployment_indices.get(name, []) if not indices: continue d = (self.model_list or [])[indices[0]] mi = d.get("model_info") if isinstance(d, dict) else d.model_info mi_dict: dict[str, Any] = mi if isinstance(mi, dict) else (mi.model_dump() if mi else {}) prefs_raw = mi_dict.get("adaptive_router_preferences") if prefs_raw is not None: model_to_prefs[name] = AdaptiveRouterPreferences(**prefs_raw) # model_info is the conventional pricing location elsewhere in LiteLLM; litellm_params wins if set. lp = d.get("litellm_params") if isinstance(d, dict) else d.litellm_params lp_dict: dict[str, Any] = lp if isinstance(lp, dict) else (lp.model_dump() if lp else {}) cost = lp_dict.get("input_cost_per_token") if cost is None: cost = mi_dict.get("input_cost_per_token") if cost is not None: model_to_cost[name] = float(cost) adaptive_router: Final = AdaptiveRouter( router_name=deployment.model_name, config=config, model_to_prefs=model_to_prefs, model_to_cost=model_to_cost, ) self._register_pre_routing_strategy( registry=self.adaptive_routers, deployment=deployment, strategy=adaptive_router, strategy_label="Adaptive-router", ) litellm.logging_callback_manager.add_litellm_callback( AdaptiveRouterPostCallHook(adaptive_router=adaptive_router) ) verbose_router_logger.info( "AdaptiveRouter[%s] initialized with %d models", deployment.model_name, len(config.available_models), ) def _is_quality_router_deployment(self, litellm_params: LiteLLM_Params) -> bool: """ Check if the deployment is a quality-router deployment. Returns True if the litellm_params model starts with "auto_router/quality_router". """ return classify_strategy_router_model(litellm_params.model) == "quality" def init_quality_router_deployment(self, deployment: Deployment): """ Initialize the quality-router deployment. Resolves the default model from either `quality_router_default_model` or `quality_router_config["default_model"]`, then instantiates the QualityRouter and stores it in `self.quality_routers`. """ # Import here to mirror the AutoRouter / ComplexityRouter init pattern # and avoid circular imports. from litellm.router_strategy.quality_router.quality_router import ( QualityRouter, ) quality_router_config: Final[dict | None] = deployment.litellm_params.quality_router_config default_model: str | None = deployment.litellm_params.quality_router_default_model if default_model is None and quality_router_config: default_model = quality_router_config.get("default_model") if default_model is None: raise ValueError( "quality_router_default_model is required for quality-router deployments, " "or set default_model in quality_router_config. Please configure it in the litellm_params" ) quality_router: Final[QualityRouter] = QualityRouter( model_name=deployment.model_name, default_model=default_model, litellm_router_instance=self, quality_router_config=quality_router_config, ) self._register_pre_routing_strategy( registry=self.quality_routers, deployment=deployment, strategy=quality_router, strategy_label="Quality-router", ) def deployment_is_active_for_environment(self, deployment: Deployment) -> bool: """ Function to check if a llm deployment is active for a given environment. Allows using the same config.yaml across multople environments Requires `LITELLM_ENVIRONMENT` to be set in .env. Valid values for environment: - development - staging - production Raises: - ValueError: If LITELLM_ENVIRONMENT is not set in .env or not one of the valid values - ValueError: If supported_environments is not set in model_info or not one of the valid values """ return model_info_is_active_for_environment(model_info=deployment.model_info) def set_model_list(self, model_list: list): original_model_list: Final = copy.deepcopy(model_list) self._discovered_model_info_cache.flush_cache() self.model_list = [] self.model_id_to_deployment_index_map = {} # Reset the index self.model_name_to_deployment_indices = {} # Reset the model_name index self.team_model_to_deployment_indices = {} # Reset the team_model index self.team_pattern_routers = {} self.team_public_model_names = frozenset() # Reset per-strategy router registries so hot-reload doesn't leave # stale routers pointing at the old model_list. self.quality_routers = {} self.complexity_routers = {} self.auto_routers = {} self._provider_unresolved_deployments = () self._invalidate_model_group_info_cache() self._invalidate_access_groups_cache() # we add api_base/api_key each model so load balancing between azure/gpt on api_base1 and api_base2 works declared_ids: Final = tuple( str(entry["model_info"]["id"]) for entry in original_model_list if isinstance(entry.get("model_info"), dict) and entry["model_info"].get("id") is not None ) duplicate_ids: Final = frozenset(model_id for model_id in declared_ids if declared_ids.count(model_id) > 1) ptu_declared: Final = tuple( str(entry.get("model_name")) for entry in original_model_list if isinstance(entry.get("model_info"), dict) and entry["model_info"].get("db_model") is not True and declares_ptu(entry["model_info"]) ) if ptu_declared and not is_ptu_cost_attribution_enabled(): verbose_router_logger.warning( "PTU fields are set on config.yaml deployment(s) %s, but PTU cost attribution is disabled, so no " "flat cost accrues and this traffic is billed per token. Set %s=True to enable it", ", ".join(ptu_declared), PTU_COST_ATTRIBUTION_ENV_VAR, ) for model in original_model_list: _model_name = model.pop("model_name") _litellm_params = model.pop("litellm_params") ## check if litellm params in os.environ if isinstance(_litellm_params, dict): for k, v in _litellm_params.items(): if isinstance(v, str) and v.startswith("os.environ/"): _litellm_params[k] = get_secret(v) _model_info: dict = model.pop("model_info", {}) declared_id = None if _model_info.get("id") is None else str(_model_info["id"]) # check if model info has id if "id" not in _model_info: _id = self.generate_model_id(_model_name, _litellm_params) _model_info["id"] = _id if _litellm_params.get("organization", None) is not None and isinstance( _litellm_params["organization"], list ): # Addresses https://github.com/BerriAI/litellm/issues/3949 for org in _litellm_params["organization"]: _litellm_params["organization"] = org self._create_deployment( deployment_info=model, _model_name=_model_name, _litellm_params=_litellm_params, _model_info=_model_info, declared_id=declared_id, duplicate_ids=duplicate_ids, ) else: self._create_deployment( deployment_info=model, _model_name=_model_name, _litellm_params=_litellm_params, _model_info=_model_info, declared_id=declared_id, duplicate_ids=duplicate_ids, ) verbose_router_logger.debug("\nInitialized Model List %s", self.get_model_names()) self.model_names = {m["model_name"] for m in model_list} # Note: model_name_to_deployment_indices is already built incrementally # by _create_deployment -> _add_model_to_list_and_index_map # Deferred: build the AdaptiveRouter strategy now that all underlying # deployments have been registered. self._finalize_adaptive_router_if_configured() def _add_deployment(self, deployment: Deployment) -> Deployment: import os #### VALIDATE MODEL ######## # Check if this is a prompt management model before validating as LLM provider litellm_model: Final = deployment.litellm_params.model if isinstance(deployment.litellm_params.drop_params, str): verbose_router_logger.warning( "model=%s drop_params=%r is not a flag value, treating it as unset", deployment.model_name, deployment.litellm_params.drop_params, ) is_prompt_management_model = False if "/" in litellm_model: split_litellm_model: Final = litellm_model.split("/")[0] if split_litellm_model in litellm._known_custom_logger_compatible_callbacks: is_prompt_management_model = True if is_prompt_management_model: # For prompt management models, skip LLM provider validation # The actual model will be resolved at runtime from the prompt file _model = litellm_model custom_llm_provider = None dynamic_api_key = None api_base = None else: # check if model provider in supported providers ( _model, custom_llm_provider, dynamic_api_key, api_base, ) = litellm.get_llm_provider( model=deployment.litellm_params.model, custom_llm_provider=deployment.litellm_params.get("custom_llm_provider", None), api_base=deployment.litellm_params.api_base, ) # done reading model["litellm_params"] # Check if provider is supported: either in enum or JSON-configured if ( custom_llm_provider not in litellm.provider_list and not JSONProviderRegistry.exists(custom_llm_provider) and not is_registered_custom_provider(custom_llm_provider) ): raise Exception(f"Unsupported provider - {custom_llm_provider}") #### DEPLOYMENT NAMES INIT ######## self.deployment_names.append(deployment.litellm_params.model) ############ Users can either pass tpm/rpm as a litellm_param or a router param ########### # for get_available_deployment, we use the litellm_param["rpm"] # in this snippet we also set rpm to be a litellm_param if deployment.litellm_params.rpm is None and getattr(deployment, "rpm", None) is not None: deployment.litellm_params.rpm = getattr(deployment, "rpm") if deployment.litellm_params.tpm is None and getattr(deployment, "tpm", None) is not None: deployment.litellm_params.tpm = getattr(deployment, "tpm") # Check if user is trying to use model_name == "*" # this is a catch all model for their specific api key # if deployment.model_name == "*": # if deployment.litellm_params.model == "*": # # user wants to pass through all requests to litellm.acompletion for unknown deployments # self.router_general_settings.pass_through_all_models = True # else: # self.default_deployment = deployment.to_json(exclude_none=True) # Check if user is using provider specific wildcard routing # example model_name = "databricks/*" or model_name = "anthropic/*" if "*" in deployment.model_name: # store this as a regex pattern - all deployments matching this pattern will be sent to this deployment # Store deployment.model_name as a regex pattern self.pattern_router.add_pattern(deployment.model_name, deployment.to_json(exclude_none=True)) if deployment.model_info.id: self.provider_default_deployment_ids.append(deployment.model_info.id) _team_id: Final = deployment.model_info.get("team_id") _team_public_model_name: Final[str | None] = deployment.model_info.get("team_public_model_name") if _team_id is not None and _team_public_model_name is not None and "*" in _team_public_model_name: if _team_id not in self.team_pattern_routers: self.team_pattern_routers[_team_id] = PatternMatchRouter() self.team_pattern_routers[_team_id].add_pattern( _team_public_model_name, deployment.to_json(exclude_none=True) ) # Azure GPT-Vision Enhancements, users can pass os.environ/ data_sources: Final = deployment.litellm_params.get("dataSources", []) or [] for data_source in data_sources: params = data_source.get("parameters", {}) for param_key in ["endpoint", "key"]: # if endpoint or key set for Azure GPT Vision Enhancements, check if it's an env var if param_key in params and params[param_key].startswith("os.environ/"): env_name = params[param_key].replace("os.environ/", "") params[param_key] = os.environ.get(env_name, "") # # init OpenAI, Azure clients # InitalizeOpenAISDKClient.set_client( # litellm_router_instance=self, model=deployment.to_json(exclude_none=True) # ) if custom_llm_provider is not None: self._initialize_deployment_for_pass_through( deployment=deployment, custom_llm_provider=custom_llm_provider, ) ######################################################### # Check if this is an auto-router deployment ######################################################### if self._is_auto_router_deployment(litellm_params=deployment.litellm_params): self.init_auto_router_deployment(deployment=deployment) ######################################################### # Check if this is a complexity-router deployment ######################################################### if self._is_complexity_router_deployment(litellm_params=deployment.litellm_params): self.init_complexity_router_deployment(deployment=deployment) # NOTE: adaptive-router deployments are deferred to the end of # set_model_list() because their init needs visibility into the OTHER # deployments listed in `available_models` (which may not yet have # been processed when this one is created). ######################################################### # Check if this is a quality-router deployment ######################################################### if self._is_quality_router_deployment(litellm_params=deployment.litellm_params): self.init_quality_router_deployment(deployment=deployment) return deployment def _initialize_deployment_for_pass_through(self, deployment: Deployment, custom_llm_provider: str): """ Optional: Register vertex credentials for pass-through endpoints if `deployment.litellm_params.use_in_pass_through` is True Other providers need no registration here: PassthroughEndpointRouter.get_credentials resolves their credentials per-request from the live router deployments """ if deployment.litellm_params.use_in_pass_through is not True: return if custom_llm_provider != "vertex_ai": return from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( passthrough_endpoint_router, ) credential_name: Final = deployment.litellm_params.litellm_credential_name credential_values: Final = ( CredentialAccessor.get_credential_values(credential_name) if credential_name is not None else {} ) vertex_project: Final = credential_values.get("vertex_project") or deployment.litellm_params.vertex_project vertex_location: Final = credential_values.get("vertex_location") or deployment.litellm_params.vertex_location vertex_credentials: Final = ( credential_values.get("vertex_credentials") or deployment.litellm_params.vertex_credentials ) if vertex_project is None or vertex_location is None: raise ValueError( "vertex_project, and vertex_location must be set in litellm_params for pass-through endpoints." ) passthrough_endpoint_router.add_vertex_credentials( project_id=vertex_project, location=vertex_location, vertex_credentials=vertex_credentials, ) def add_deployment(self, deployment: Deployment) -> Deployment | None: """ Parameters: - deployment: Deployment - the deployment to be added to the Router Returns: - The added deployment - OR None (if deployment already exists) """ # check if deployment already exists _deployment_model_id: Final = deployment.model_info.id if _deployment_model_id and self.has_model_id(_deployment_model_id): return None warn_on_provider_credential_mismatch( model_name=deployment.model_name, litellm_params=deployment.litellm_params.model_dump(exclude_none=True), ) # add to model list _deployment: Final = deployment.to_json(exclude_none=True) # initialize client self._add_deployment(deployment=deployment) Router._register_deployment_pricing(deployment=deployment) # add to model names self._add_model_to_list_and_index_map(model=_deployment, model_id=deployment.model_info.id) self.model_names.add(deployment.model_name) self._sync_deployment_budget_config(deployment=deployment) return deployment def _update_deployment_indices_after_removal(self, model_id: str, removal_idx: int) -> None: """ Helper method to update deployment indices after a deployment has been removed from model_list. Parameters: - model_id: str - the id of the deployment that was removed - removal_idx: int - the index where the deployment was removed from model_list """ self._discovered_model_info_cache.delete_cache(model_id) # Update indices for all models after the removed one for deployment_id, idx in self.model_id_to_deployment_index_map.items(): if idx > removal_idx: self.model_id_to_deployment_index_map[deployment_id] = idx - 1 # Remove the deleted model from index if model_id in self.model_id_to_deployment_index_map: del self.model_id_to_deployment_index_map[model_id] # Update model_name_to_deployment_indices for model_name, indices in list(self.model_name_to_deployment_indices.items()): # Build new list without mutating the original updated_indices = [] for idx in indices: if idx == removal_idx: # Skip the removed index continue elif idx > removal_idx: # Decrement indices after removal updated_indices.append(idx - 1) else: # Keep indices before removal unchanged updated_indices.append(idx) # Update or remove the entry if len(updated_indices) > 0: self.model_name_to_deployment_indices[model_name] = updated_indices else: del self.model_name_to_deployment_indices[model_name] self.model_names.discard(model_name) # Update team_model_to_deployment_indices for key, indices in list(self.team_model_to_deployment_indices.items()): # Build new list without mutating the original updated_indices = [] for idx in indices: if idx == removal_idx: # Skip the removed index continue elif idx > removal_idx: # Decrement indices after removal updated_indices.append(idx - 1) else: # Keep indices before removal unchanged updated_indices.append(idx) # Update or remove the entry if len(updated_indices) > 0: self.team_model_to_deployment_indices[key] = updated_indices else: del self.team_model_to_deployment_indices[key] self.team_public_model_names = frozenset( public_model_name for _, public_model_name in self.team_model_to_deployment_indices ) self.pattern_router.remove_deployment(model_id) for team_id in list(self.team_pattern_routers.keys()): team_pattern_router = self.team_pattern_routers[team_id] team_pattern_router.remove_deployment(model_id) if not team_pattern_router.patterns: del self.team_pattern_routers[team_id] def _update_team_model_index(self, model: dict, idx: int) -> None: """ Helper to update team_model_to_deployment_indices for a single deployment. Parameters: - model: dict - the deployment to index - idx: int - the index in model_list """ team_id: Final = (model.get("model_info") or {}).get("team_id") team_public_model_name: Final = (model.get("model_info") or {}).get("team_public_model_name") if team_id and team_public_model_name: key: Final = (team_id, team_public_model_name) self.team_public_model_names = self.team_public_model_names | frozenset({team_public_model_name}) if key not in self.team_model_to_deployment_indices: self.team_model_to_deployment_indices[key] = [] if idx not in self.team_model_to_deployment_indices[key]: self.team_model_to_deployment_indices[key].append(idx) def _add_model_to_list_and_index_map(self, model: dict, model_id: str | None = None) -> None: """ Helper method to add a model to the model_list and update both indices. Parameters: - model: dict - the model to add to the list - model_id: Optional[str] - the model ID to use for indexing. If None, will try to get from model["model_info"]["id"] """ idx: Final = len(self.model_list) self.model_list.append(model) self._invalidate_model_group_info_cache() self._invalidate_access_groups_cache() # Update model_id index for O(1) lookup if model_id is not None: self.model_id_to_deployment_index_map[model_id] = idx elif model.get("model_info", {}).get("id") is not None: self.model_id_to_deployment_index_map[model["model_info"]["id"]] = idx # Update model_name index for O(1) lookup model_name: Final = model.get("model_name") if model_name: if model_name not in self.model_name_to_deployment_indices: self.model_name_to_deployment_indices[model_name] = [] self.model_name_to_deployment_indices[model_name].append(idx) # Update team_model index for O(1) team-scoped lookup self._update_team_model_index(model, idx) def upsert_deployment(self, deployment: Deployment) -> Deployment | None: """ Add or update deployment Parameters: - deployment: Deployment - the deployment to be added to the Router Returns: - The added/updated deployment """ _deployment_model_id: Final = deployment.model_info.id or "" _deployment_on_router: Final[Deployment | None] = self.get_deployment(model_id=_deployment_model_id) try: if _deployment_on_router is not None: # deployment with this model_id exists on the router if ( deployment.model_name == _deployment_on_router.model_name and deployment.litellm_params == _deployment_on_router.litellm_params and deployment.model_info == _deployment_on_router.model_info ): # No need to update return None # if there is a new litellm param -> then update the deployment # remove the previous deployment removal_idx: int | None = None deployment_id: Final = deployment.model_info.id deployment_fast_mapping: Final = self.model_id_to_deployment_index_map if deployment_id in deployment_fast_mapping: removal_idx = deployment_fast_mapping[deployment_id] if removal_idx is not None: self.model_list.pop(removal_idx) self._invalidate_model_group_info_cache() self._invalidate_access_groups_cache() self._update_deployment_indices_after_removal(model_id=deployment_id, removal_idx=removal_idx) # Free the outgoing deployment's pre-routing strategy slot (keyed by the # OLD model_name/tags) before the re-add below re-registers it. self._unregister_pre_routing_strategy_for_deployment(deployment=_deployment_on_router) # if the model_id is not in router self.add_deployment(deployment=deployment) # add_deployment() builds every strategy EXCEPT the adaptive one, which # set_model_list() defers until the whole model_list is visible. Re-run that # deferred pass so an adaptive router whose slot was just released above is # rebuilt rather than left unregistered. if self._deployment_participates_in_adaptive_routing(litellm_params=deployment.litellm_params) or ( _deployment_on_router is not None and self._deployment_participates_in_adaptive_routing( litellm_params=_deployment_on_router.litellm_params ) ): self._finalize_adaptive_router_if_configured() return deployment except Exception as e: if self.ignore_invalid_deployments: verbose_router_logger.warning( "Error upserting deployment %s (id=%s): %s. Dropping it and continuing with other deployments.", deployment.model_name, deployment.model_info.id, e, ) self._restore_deployment_after_failed_upsert( previous_deployment=_deployment_on_router, model_id=_deployment_model_id ) return None else: raise e def _restore_deployment_after_failed_upsert(self, previous_deployment: Deployment | None, model_id: str) -> None: """Put a deployment back the way it was before a failed upsert popped it. A rollback re-admits state that was already serving, so it does not go through the capability ceiling a newcomer gets: with the ceiling tightened since the deployment first registered, judging the rollback would drop a serving router over an unrelated failed edit. """ if previous_deployment is None or self.has_model_id(model_id): return limit_resolver: Final = self.auto_router_capability_limit self.auto_router_capability_limit = None try: self.add_deployment(deployment=previous_deployment) verbose_router_logger.info( "Restored deployment %s (id=%s); it keeps serving its previous configuration.", previous_deployment.model_name, model_id, ) except Exception as restore_error: # noqa: BLE001 # best-effort restore: a second failure must not abort the reload verbose_router_logger.warning( "Could not restore previously served deployment %s (id=%s) after the failed upsert: %s", previous_deployment.model_name, model_id, restore_error, ) finally: self.auto_router_capability_limit = limit_resolver @staticmethod def _backend_cost_map_keys(model: str, custom_llm_provider: str | None) -> tuple[str, ...]: """The ``litellm.model_cost`` keys a deployment's shared backend info is registered under.""" backend_key: Final = model if custom_llm_provider is None else f"{custom_llm_provider}/{model}" if "responses/" in backend_key: return (backend_key, backend_key.replace("responses/", "")) return (backend_key,) @staticmethod def _deployment_model_cost_payload(deployment: Deployment) -> dict: # mutable-ok: cost-map entry """The ``model_info`` a deployment contributes to ``litellm.model_cost``. Custom pricing lives on ``litellm_params`` rather than ``model_info``, and the built-in cache-pricing inheritance is derived rather than stored, so both are folded back in here. That keeps this reproducible from a deployment alone, which is what lets a refresh rebuild the same entries. """ model_info: Final[dict] = deployment.model_info.model_dump(exclude_none=True) # mutable-ok: built in place for field in CustomPricingLiteLLMParams.model_fields: field_value = deployment.litellm_params.get(field) if field_value is not None: model_info[field] = field_value Router._inherit_builtin_base_rates_for_off_peak( model_info=model_info, backend_model=deployment.litellm_params.model, custom_llm_provider=deployment.litellm_params.custom_llm_provider, ) if model_info.get("input_cost_per_token") is not None: Router._inherit_builtin_cache_pricing( model_info=model_info, backend_model=deployment.litellm_params.model, custom_llm_provider=deployment.litellm_params.custom_llm_provider, ) Router._inherit_builtin_tiered_output_rate( model_info=model_info, backend_model=deployment.litellm_params.model, custom_llm_provider=deployment.litellm_params.custom_llm_provider, ) return model_info @staticmethod def _register_deployment_pricing(deployment: Deployment) -> None: """Register a deployment's custom/inherited pricing in ``litellm.model_cost``. Takes only a ``Deployment``, so it registers pricing for a deployment that is never added to ``self.model_list`` (a per-request client-side-credential deployment) just as readily as one that is. """ Router._register_deployment_in_model_cost( model_id=deployment.model_info.id, model_info=Router._deployment_model_cost_payload(deployment), model=deployment.litellm_params.model, custom_llm_provider=deployment.litellm_params.custom_llm_provider, ) @staticmethod def _register_deployment_in_model_cost( *, model_id: str | None, model_info: dict, # mutable-ok: cost-map entry model: str, custom_llm_provider: str | None, ) -> None: """Write a deployment's metadata into ``litellm.model_cost``. Runs when a deployment is added and again after a price data reload, so the entries a refresh rebuilds are the ones a fresh boot would produce. Nothing is recorded for replay: a refresh walks the live routers instead, so a deleted, repointed or never-added deployment, and a discarded router, drop out of the rebuild on their own. A strategy-router alias is never the deployment actually called or billed, so custom pricing configured on it must not become a cost-map price: an explicit zero would let ``_is_cost_explicitly_configured`` treat the alias as a genuinely free model and waive budget checks for requests that route to (and bill as) a real deployment. """ if classify_strategy_router_model(model) is not None: model_info = { # mutable-ok: filtered copy of the caller's entry, handed straight to register_model k: v for k, v in model_info.items() if k not in CustomPricingLiteLLMParams.model_fields } if model_id is not None: litellm.register_model( model_cost={model_id: model_info}, persist_across_reloads=False, warning_display_name=model, ) ## OLD MODEL REGISTRATION ## Kept to prevent breaking changes backend_keys: Final = Router._backend_cost_map_keys(model=model, custom_llm_provider=custom_llm_provider) backend_key: Final = backend_keys[0] # For the shared backend key, keep only cost-map schema fields # (minus custom pricing) so that one deployment's pricing overrides # or custom metadata (id, access_via_team_ids, arbitrary keys) # don't pollute another deployment sharing the same backend model # name. Each deployment's full model_info is already stored under # its unique model_id above. shared_model_info: Final = shared_backend_model_info(model_info) existing_shared_mode: Final = (cast(dict | None, litellm.model_cost.get(backend_key, {})) or {}).get("mode") deployment_mode: Final = shared_model_info.get("mode") # Keep the built-in bridge mode stable for shared backend keys. # Multiple aliases can point at the same provider/model backend, # but their deployment-level overrides should not downgrade the # backend from responses -> chat via last-write-wins registration. # Only preserve in that specific direction so legitimate upgrades # (e.g. chat -> responses) and unrelated mode changes still apply, # and so a missing deployment mode does not silently clear the # existing shared backend mode. is_responses_to_chat_downgrade: Final = existing_shared_mode == "responses" and deployment_mode == "chat" would_clear_existing_mode: Final = existing_shared_mode is not None and deployment_mode is None if is_responses_to_chat_downgrade or would_clear_existing_mode: if deployment_mode is not None: verbose_router_logger.warning( "Router: preserving existing mode=%s for shared backend " "key %s instead of the deployment-specified mode=%s " "(prevents alias registration from downgrading the " "shared backend mode).", existing_shared_mode, backend_key, deployment_mode, ) shared_model_info["mode"] = existing_shared_mode # Always register the (possibly mode-preserved) shared backend info. litellm.register_model( model_cost={_key: shared_model_info for _key in backend_keys}, persist_across_reloads=False, ) def _replay_model_cost_registrations(self) -> None: """Re-assert this router's deployments onto a freshly fetched catalog. Reads ``model_list`` at call time, so only deployments the router still serves are restored, plus any config deployment the fresh catalog now resolves. """ provider_unresolved: Final = self._provider_unresolved_deployments self._provider_unresolved_deployments = () for create_deployment in provider_unresolved: create_deployment() for entry in tuple(self.model_list): try: deployment = entry if isinstance(entry, Deployment) else Deployment(**entry) except Exception: # noqa: BLE001 # a malformed entry must not abort the rest of the rebuild verbose_router_logger.exception( "Router: could not rebuild cost-map entry for a deployment during a price data reload" ) continue Router._register_deployment_in_model_cost( model_id=deployment.model_info.id, model_info=Router._deployment_model_cost_payload(deployment), model=deployment.litellm_params.model, custom_llm_provider=deployment.litellm_params.custom_llm_provider, ) self._invalidate_model_group_info_cache() def delete_deployment(self, id: str) -> Deployment | None: """ Parameters: - id: str - the id of the deployment to be deleted Returns: - The deleted deployment - OR None (if deleted deployment not found) """ deployment_idx = None if id in self.model_id_to_deployment_index_map: deployment_idx = self.model_id_to_deployment_index_map[id] try: if deployment_idx is not None: # Pop the item from the list first item: Final = self.model_list.pop(deployment_idx) self._invalidate_model_group_info_cache() self._invalidate_access_groups_cache() self._update_deployment_indices_after_removal(model_id=id, removal_idx=deployment_idx) _budget_limiter: Final = self._get_router_deployment_budget_limiter() if _budget_limiter is not None: _budget_limiter.unregister_deployment_budget(model_id=id) try: self._unregister_pre_routing_strategy_for_deployment( deployment=item if isinstance(item, Deployment) else Deployment(**item) ) except Exception: verbose_router_logger.exception( "delete_deployment: could not release pre-routing strategies for model_id=%s; " "the deployment is out of the model_list and its indices are repaired", id, ) return item else: return None except Exception: return None def _get_router_deployment_budget_limiter( self, ) -> RouterBudgetLimiting | None: """ Return the router's deployment-budget callback. Uses exact-type matching so proxy subclasses (e.g. virtual-key model budgets) registered on litellm.callbacks are not mistaken for router deployment budgets. """ if self.router_budget_logger is not None: return self.router_budget_logger if self.optional_callbacks: for _cb in self.optional_callbacks: if type(_cb) is RouterBudgetLimiting: self.router_budget_logger = _cb return _cb return None def _deployment_has_budget_limits(self, deployment: Deployment) -> bool: return ( deployment.litellm_params.get("max_budget") is not None and deployment.litellm_params.get("budget_duration") is not None and deployment.model_info.id is not None ) def _sync_deployment_budget_config(self, deployment: Deployment) -> None: model_id: Final = deployment.model_info.id if model_id is None: return _budget_limiter = self._get_router_deployment_budget_limiter() if not self._deployment_has_budget_limits(deployment=deployment): if _budget_limiter is not None: _budget_limiter.unregister_deployment_budget(model_id=model_id) return if _budget_limiter is None: self.add_optional_pre_call_checks(optional_pre_call_checks=["router_budget_limiting"]) _budget_limiter = self._get_router_deployment_budget_limiter() if _budget_limiter is not None: _budget_limiter.register_deployment_budget(deployment=deployment.to_json(exclude_none=True)) def get_deployment(self, model_id: str) -> Deployment | None: """ Returns -> Deployment or None Raise Exception -> if model found in invalid format """ # Use O(1) lookup via model_id_to_deployment_index_map only if model_id in self.model_id_to_deployment_index_map: idx: Final = self.model_id_to_deployment_index_map[model_id] model: Final = self.model_list[idx] if isinstance(model, dict): return Deployment(**model) elif isinstance(model, Deployment): return model else: raise Exception(f"Model invalid format - {type(model)}") return None def get_deployment_credentials(self, model_id: str) -> dict | None: """ Returns -> dict of credentials for a given model id. Returns None if the deployment is paused via `LiteLLM_ProxyModelTable.blocked`, so file/batch/passthrough callers that resolve credentials directly cannot keep using a paused deployment. """ deployment: Final = self.get_deployment(model_id=model_id) if deployment is None or self._is_deployment_blocked(deployment): return None return CredentialLiteLLMParams.model_validate( deployment.litellm_params.model_dump(exclude_none=True) ).model_dump(exclude_none=True) def get_deployment_by_model_group_name(self, model_group_name: str) -> Deployment | None: """ Returns -> Deployment or None Raise Exception -> if model found in invalid format Optimized with O(1) index lookup instead of O(n) linear scan. """ # O(1) lookup in model_name index if model_group_name in self.model_name_to_deployment_indices: indices: Final = self.model_name_to_deployment_indices[model_group_name] if indices: # Return first deployment for this model_name model: Final = self.model_list[indices[0]] if isinstance(model, dict): return Deployment(**model) elif isinstance(model, Deployment): return model else: raise Exception(f"Model Name invalid - {type(model)}") return None @staticmethod def _deployment_usable_by_team(model: Mapping | Deployment, team_id: str | None) -> bool: """ A team-scoped deployment (``model_info.team_id`` set) is only usable by callers from that same team; deployments without a team owner are shared. """ model_info: Final = model.get("model_info") if isinstance(model, dict) else model.model_info owner_team_id: Final = model_info.get("team_id") if model_info is not None else None return owner_team_id is None or owner_team_id == team_id def _get_model_group_deployment_usable_by_team( self, model_group_name: str, team_id: str | None ) -> Deployment | None: """ Like ``get_deployment_by_model_group_name``, but skips deployments owned by other teams so a shared model name never resolves another team's credentials. """ indices: Final = self.model_name_to_deployment_indices.get(model_group_name) or () usable: Final = ( self.model_list[idx] for idx in indices if self._deployment_usable_by_team(self.model_list[idx], team_id) ) first_usable: Final = next(usable, None) if first_usable is None: return None return Deployment(**first_usable) if isinstance(first_usable, dict) else first_usable async def arefresh_model_info(self, *, client: AsyncHTTPHandler | None = None) -> None: """Refresh token limits advertised by configured OpenAI-compatible deployments.""" deployments: Final = iter(tuple(self.model_list)) async def refresh_worker() -> None: for raw_deployment in deployments: try: await self._arefresh_deployment_model_info(raw_deployment, client=client) except Exception: # noqa: BLE001 # one invalid deployment must not prevent refreshing the others verbose_router_logger.debug("Could not refresh deployment model info") await asyncio.gather(*(refresh_worker() for _ in range(MODEL_INFO_REFRESH_CONCURRENCY))) self._invalidate_model_group_info_cache() async def _arefresh_deployment_model_info( self, raw_deployment: Mapping[str, object], *, client: AsyncHTTPHandler | None ) -> None: deployment: Final = Deployment.model_validate(raw_deployment) params: Final = LiteLLM_Params.model_validate( MappingProxyType( { **deployment.litellm_params.model_dump(exclude_none=True), **( self.get_deployment_credentials_with_provider(deployment.model_info.id or "") or MappingProxyType({}) ), } ) ) model, provider, dynamic_api_key, api_base = litellm.get_llm_provider(model=params.model, litellm_params=params) if provider not in MODEL_INFO_DISCOVERY_PROVIDERS: return if api_base is None or "*" in model or params.get("use_clientside_credentials"): return api_key: Final = params.api_key or dynamic_api_key headers: Final = TypeAdapter(Mapping[str, str]).validate_python( params.get("extra_headers") or params.get("headers") or MappingProxyType({}) ) auth_headers: Final = ( MappingProxyType({"authorization": f"Bearer {api_key}"}) if api_key else MappingProxyType({}) ) limits: Final = await get_openai_compatible_model_info( model=model, api_base=api_base, headers=MappingProxyType( { **auth_headers, **MappingProxyType({key.lower(): value for key, value in headers.items()}), } ), client=client or get_async_httpx_client(llm_provider=LlmProviders.OPENAI), cache=self.cache.in_memory_cache, ) model_id: Final = deployment.model_info.id if not limits or model_id is None or self.get_model_info(model_id) is not raw_deployment: return self._discovered_model_info_cache.max_size_in_memory = max(len(self.model_list), 1) self._discovered_model_info_cache.delete_cache(model_id) self._discovered_model_info_cache.set_cache( model_id, DiscoveredDeploymentModelInfo(deployment=raw_deployment, limits=limits) ) self._invalidate_model_group_info_cache() def get_discovered_model_info(self, model_id: str | None) -> Mapping[str, int]: cached: Final[object] = self._discovered_model_info_cache.get_cache(model_id) if ( model_id is not None and isinstance(cached, DiscoveredDeploymentModelInfo) and cached.deployment is self.get_model_info(model_id) ): configured: Final = TypeAdapter(Mapping[str, object]).validate_python(cached.deployment["model_info"]) return MappingProxyType({key: value for key, value in cached.limits.items() if configured.get(key) is None}) return MappingProxyType({}) def get_model_listing_info(self, model_name: str) -> DeploymentModelListingInfo | None: """ Return what the concrete deployments behind model_name contribute to its /v1/models entry: the cost-map keys for their underlying models, plus the widest configured or discovered token limits. Resolved via O(1) index lookup. Returns None for wildcard-expanded or unknown names, where the listed name is the real model name and no deployment-specific information exists, and treats a malformed configured limit as absent rather than failing the listing. The whole group is read rather than just its first deployment, so a group that mixes models does not advertise a window that depends on config order; the widest one is reported, which is what get_model_group_info already shows the Admin UI. Keys are deduplicated, so the ordinary group of interchangeable deployments of one model still costs the caller a single cost-map lookup. Unlike get_model_group_info, this never triggers pattern matching or deep copies, so it is safe to call per listed model on the /v1/models hot path. """ indices: Final = self.model_name_to_deployment_indices.get(model_name) if not indices: return None deployments: Final = tuple(self.model_list[index] for index in indices) model_infos: Final = tuple( MappingProxyType( { **self.get_discovered_model_info((deployment.get("model_info") or MappingProxyType({})).get("id")), **MappingProxyType( { k: v for k, v in (deployment.get("model_info") or MappingProxyType({})).items() if v is not None } ), } ) for deployment in deployments ) params: Final = tuple(deployment.get("litellm_params") or MappingProxyType({}) for deployment in deployments) # base_model resolution mirrors get_router_model_info: unset or blank means the # deployment's own model name is the cost-map key. cost_map_keys: Final = tuple( dict.fromkeys( # deduplicates while preserving config order key for key in ( model_info.get("base_model") or litellm_params.get("base_model") or litellm_params.get("model") for model_info, litellm_params in zip(model_infos, params) ) if isinstance(key, str) and key ) ) return DeploymentModelListingInfo( cost_map_keys=cost_map_keys, max_input_tokens=self._widest_configured_limit(model_infos, "max_input_tokens"), max_output_tokens=self._widest_configured_limit(model_infos, "max_output_tokens"), ) @staticmethod def _widest_configured_limit(model_infos: Sequence[Mapping[str, object]], field: str) -> int | None: """The largest usable value of ``field`` across a group's configured model_info blocks.""" limits: Final = tuple( limit for limit in (coerce_token_limit(model_info.get(field)) for model_info in model_infos) if limit is not None ) return max(limits) if limits else None def get_configured_token_limits(self, model_name: str) -> "tuple[int | None, int | None]": """ Return (max_input_tokens, max_output_tokens) configured or discovered for a concrete deployment of model_name, via O(1) index lookup. Returns (None, None) for wildcard-expanded or unknown names, and treats a malformed configured value as absent rather than failing the caller. Deliberately reads one deployment rather than aggregating the group the way get_model_listing_info does: its caller truncates an embedding input to this value, so the widest window in a mixed group would be the wrong answer there. """ deployment: Final = self.get_deployment_by_model_group_name(model_group_name=model_name) if deployment is None: return (None, None) model_info: Final = MappingProxyType( { **self.get_discovered_model_info(deployment.model_info.id), **deployment.model_info.model_dump(exclude_none=True), } ) return ( coerce_token_limit(model_info.get("max_input_tokens")), coerce_token_limit(model_info.get("max_output_tokens")), ) def get_configured_mode(self, model_name: str) -> "str | None": """Return the mode explicitly configured for a concrete deployment.""" deployment: Final = self.get_deployment_by_model_group_name(model_group_name=model_name) if deployment is None: return None mode: Final = deployment.model_info.get("mode") if isinstance(mode, str) and mode.strip(): return mode return None def get_configured_display_name(self, model_name: str) -> "str | None": """ Return the display_name explicitly configured in a concrete deployment's model_info for model_name, via O(1) index lookup. Returns None for wildcard-expanded or unknown names, and treats a non-string or empty configured value as absent rather than failing the listing. Like get_configured_token_limits, this never triggers pattern matching or deep copies, so it is safe to call per listed model on the /v1/models hot path. """ deployment: Final = self.get_deployment_by_model_group_name(model_group_name=model_name) if deployment is None: return None display_name: Final = deployment.model_info.get("display_name") if isinstance(display_name, str) and display_name.strip(): return display_name return None def get_credential_deployment(self, model_id: str, team_id: str | None = None) -> Deployment | None: """ The deployment a passthrough endpoint (files, batches, etc.) resolves for a model id or model name: by deployment id first, then by model_name, then by the team's exact public model name, then by wildcard pattern (team wildcards before global ones, so a global "openai/*" never shadows the team's own entry). Name and wildcard lookups never resolve another team's deployment. Returns None when nothing matches or the match is paused via `LiteLLM_ProxyModelTable.blocked`, so callers cannot bypass an admin pause by resolving the deployment directly. """ deployment: Final = ( self.get_deployment(model_id=model_id) or self._get_model_group_deployment_usable_by_team(model_group_name=model_id, team_id=team_id) or self._get_team_public_name_deployment(model_id=model_id, team_id=team_id) or self._get_wildcard_deployment_usable_by_team(model_id=model_id, team_id=team_id) ) if deployment is None or self._is_deployment_blocked(deployment): return None return deployment def _get_team_public_name_deployment(self, model_id: str, team_id: str | None) -> Deployment | None: if team_id is None: return None team_indices: Final = self.team_model_to_deployment_indices.get((team_id, model_id)) if not team_indices: return None team_model: Final = self.model_list[team_indices[0]] return Deployment(**team_model) if isinstance(team_model, dict) else team_model def _get_wildcard_deployment_usable_by_team(self, model_id: str, team_id: str | None) -> Deployment | None: team_pattern_router: Final = self.team_pattern_routers.get(team_id) if team_id is not None else None team_wildcard_models: Final = team_pattern_router.route(model_id) if team_pattern_router else None global_wildcard_models: Final = tuple( wildcard_model for wildcard_model in (self.pattern_router.route(model_id) or ()) if self._deployment_usable_by_team(wildcard_model, team_id) ) potential_wildcard_models: Final = team_wildcard_models or global_wildcard_models if not potential_wildcard_models: return None wildcard_deployment: Final = potential_wildcard_models[0] if isinstance(wildcard_deployment, dict): return Deployment(**wildcard_deployment) if isinstance(wildcard_deployment, Deployment): return wildcard_deployment return None def get_deployment_credentials_with_provider( self, model_id: str, team_id: str | None = None ) -> dict[str, Any] | None: """ Get API credentials and provider info from a model name in model_list. Useful for passthrough endpoints (files, batches, etc.) that need credentials. Resolves the deployment with `get_credential_deployment` (by deployment id, then model_name, team public model name, and wildcard pattern). Args: model_id: Model ID or model name from model_list (e.g., "gpt-4o-litellm") team_id: Optional team id of the caller. When set, team-scoped deployments (indexed by team public model name, including team wildcard models like "openai/*") are also considered. Name and wildcard lookups never resolve a deployment owned by a different team, so shared model names can't leak another team's credentials. Returns: Dictionary containing api_key, api_base, custom_llm_provider, etc. Returns None if model not found, or if the resolved deployment is paused via `LiteLLM_ProxyModelTable.blocked` (so passthrough callers cannot bypass an admin pause by resolving credentials directly). Example: credentials = router.get_deployment_credentials_with_provider("gpt-4o-litellm") # Returns: {"api_key": "sk-...", "custom_llm_provider": "openai", "model": "gpt-4o", ...} """ deployment: Final = self.get_credential_deployment(model_id=model_id, team_id=team_id) if deployment is None: return None # Get basic credentials credentials: Final = CredentialLiteLLMParams.model_validate( deployment.litellm_params.model_dump(exclude_none=True) ).model_dump(exclude_none=True) # Resolve litellm_credential_name to actual credentials if deployment.litellm_params.litellm_credential_name is not None: credential_values: Final = CredentialAccessor.get_credential_values( deployment.litellm_params.litellm_credential_name ) if not credential_values: verbose_router_logger.warning( "Credential '%s' not found in credential_list", deployment.litellm_params.litellm_credential_name ) credentials.update(credential_values) # Remove the credential name since we've resolved it credentials.pop("litellm_credential_name", None) credentials["model"] = deployment.litellm_params.model # Add custom_llm_provider if deployment.litellm_params.custom_llm_provider: credentials["custom_llm_provider"] = deployment.litellm_params.custom_llm_provider elif "/" in deployment.litellm_params.model: # Extract provider from "provider/model" format credentials["custom_llm_provider"] = deployment.litellm_params.model.split("/")[0] else: credentials["custom_llm_provider"] = "openai" # default return credentials @overload def get_router_model_info( self, deployment: Union[dict, "Deployment"], received_model_name: str, id: None = None, ) -> ModelMapInfo: pass @overload def get_router_model_info(self, deployment: None, received_model_name: str, id: str) -> ModelMapInfo: pass def get_router_model_info( self, deployment: Union[dict, "Deployment"] | None, received_model_name: str, id: str | None = None, ) -> ModelMapInfo: """ For a given model id, return the model info (max tokens, input cost, output cost, etc.). Augment litellm info with additional params set in `model_info`. For azure models, ignore the `model:`. Only set max tokens, cost values if base_model is set. Returns - ModelInfo - If found -> typed dict with max tokens, input cost, etc. Raises: - ValueError -> If model is not mapped yet """ if id is not None: _deployment: Final = self.get_deployment(model_id=id) if _deployment is not None: deployment = _deployment if deployment is None: raise ValueError("Deployment not found") ## GET BASE MODEL base_model = (deployment.get("model_info") or {}).get("base_model", None) if base_model is None: base_model = (deployment.get("litellm_params") or {}).get("base_model", None) model = base_model ## GET PROVIDER - reuse LiteLLM_Params if already constructed litellm_params_data: Final = deployment.get("litellm_params") litellm_params: LiteLLM_Params if isinstance(litellm_params_data, LiteLLM_Params): litellm_params = litellm_params_data elif isinstance(litellm_params_data, dict) and "model" in litellm_params_data: litellm_params = LiteLLM_Params(**litellm_params_data) else: raise ValueError( f"Deployment missing valid litellm_params. " f"Got: {type(litellm_params_data).__name__}, " f"deployment_id: {(deployment.get('model_info') or {}).get('id', 'unknown')}" ) _model, custom_llm_provider, _, _ = litellm.get_llm_provider( model=litellm_params.model, litellm_params=litellm_params, ) ## SET MODEL TO 'model=' - if base_model is None + not azure if custom_llm_provider == "azure" and base_model is None: # Router init auto-registers every deployment name into # litellm.model_cost as a zeroed stub, so membership alone can't # tell a resolvable name apart; require usable limits/costs. _azure_fallback_key = _model if _model.startswith("azure/") else f"azure/{_model}" _fallback_entry = litellm.model_cost.get(_azure_fallback_key) _fallback_resolves = _fallback_entry is not None and ( (_fallback_entry.get("max_input_tokens") or 0) > 0 or (_fallback_entry.get("max_tokens") or 0) > 0 or (_fallback_entry.get("input_cost_per_token") or 0) > 0 ) if _fallback_resolves: verbose_router_logger.debug( "Azure deployment '%s' has no base_model set; using '%s' from the model cost map for max tokens, cost tracking, etc.", _model, _azure_fallback_key, ) else: verbose_router_logger.error( "Could not identify azure model '%s'. Set azure 'base_model' for accurate max tokens, cost tracking, etc.- https://docs.litellm.ai/docs/proxy/cost_tracking#spend-tracking-for-azure-openai-models", _model, ) elif custom_llm_provider != "azure": model = _model if "*" in model: # only call pattern_router for wildcard models potential_models: Final = self.pattern_router.route(received_model_name) if potential_models is not None: for potential_model in potential_models: try: if (potential_model.get("model_info") or {}).get("id") == ( deployment.get("model_info") or {} ).get("id"): model = (potential_model.get("litellm_params") or {}).get("model") break except Exception: pass ## GET LITELLM MODEL INFO - raises exception, if model is not mapped if model is None: # Handle case where base_model is None (e.g., Azure models without base_model set) # Use the original model from litellm_params model = _model if not model.startswith(f"{custom_llm_provider}/"): model_info_name = f"{custom_llm_provider}/{model}" else: model_info_name = model model_info: Final = litellm.get_model_info(model=model_info_name) if model_info is None: return model_info ## CHECK USER SET MODEL INFO raw_user_model_info: Final = deployment.get("model_info") user_model_info: Final = ( raw_user_model_info.model_dump(exclude_none=True) if isinstance(raw_user_model_info, BaseModel) else raw_user_model_info ) # get_model_info() hands back an lru_cache'd dict, so merge into a copy; unset # values are skipped or Deployment's None pricing defaults would erase the map's merged_model_info: Final[ModelMapInfo] = { **copy.deepcopy(model_info), **self.get_discovered_model_info((deployment.get("model_info") or {}).get("id")), **MappingProxyType( {key: value for key, value in (user_model_info or MappingProxyType({})).items() if value is not None} ), } return merged_model_info def get_model_info(self, id: str) -> dict | None: """ For a given model id, return the model info Returns - dict: the model in list with 'model_name', 'litellm_params', Optional['model_info'] - None: could not find deployment in list Optimized with O(1) index lookup instead of O(n) linear scan. """ # O(1) lookup via model_id_to_deployment_index_map if id in self.model_id_to_deployment_index_map: idx: Final = self.model_id_to_deployment_index_map[id] return self.model_list[idx] return None def get_model_group(self, id: str) -> list | None: """ Return list of all models in the same model group as that model id """ model_info: Final = self.get_model_info(id=id) if model_info is None: return None model_name: Final = model_info["model_name"] return self.get_model_list(model_name=model_name) def get_deployment_model_info(self, model_id: str, model_name: str) -> ModelInfo | None: """ For a given model id, return the model info 1. Check if model_id is in model info 2. If not, check if litellm model name is in model info 3. If not, return None """ from litellm.utils import _update_dictionary, cost_map_omits_token_price model_info: ModelInfo | None = None custom_model_info: dict | None = None litellm_model_name_model_info: ModelInfo | None = None base_model_key: str | None = None try: custom_model_info = ( { # mutable-ok: the legacy model-info merge updates this private copy **copy.deepcopy(litellm.model_cost.get(model_id) or MappingProxyType({})), **self.get_discovered_model_info(model_id), } if model_id in litellm.model_cost else None ) except Exception: pass try: litellm_model_name_model_info = litellm.get_model_info(model=model_name) except Exception: pass ## check for base model try: if custom_model_info is not None: base_model: Final = custom_model_info.get("base_model", None) if base_model is not None: ## update litellm model info with base model info base_model_info: Final = copy.deepcopy(litellm.get_model_info(model=base_model)) if base_model_info is not None: base_model_key = base_model_info.get("key") # Base model provides defaults, custom model info overrides custom_model_info = _update_dictionary( cast(dict, base_model_info), custom_model_info, ) except Exception: pass # Three mutually exclusive scenarios for the model's metadata: if custom_model_info is not None and litellm_model_name_model_info is not None: # (1) It has both custom model_info set and exists in the built-in map # merge with custom overriding built-in model_info = cast( ModelInfo, _update_dictionary( copy.deepcopy(cast(dict, litellm_model_name_model_info)), custom_model_info, ), ) elif litellm_model_name_model_info is not None: # (2) Built-in only — no custom pricing to merge model_info = copy.deepcopy(litellm_model_name_model_info) elif custom_model_info is not None: # (3) Custom only — model not in built-in cost map yet # custom_model_info already includes base_model defaults at this point, if applicable model_info = cast(ModelInfo, custom_model_info) if model_info is None: return None builtin_key: Final = ( litellm_model_name_model_info.get("key") if litellm_model_name_model_info is not None else None ) if cost_map_omits_token_price(model_id, builtin_key, base_model_key): return cast( # cast-ok: TypedDict spread with overridden keys loses its type ModelInfo, {**model_info, "input_cost_per_token": None, "output_cost_per_token": None} ) return model_info def _set_model_group_info(self, model_group: str, user_facing_model_group_name: str) -> ModelGroupInfo | None: """ For a given model group name, return the combined model info Returns: - ModelGroupInfo if able to construct a model group - None if error constructing model group info """ model_group_info: ModelGroupInfo | None = None total_tpm: int | None = None total_rpm: int | None = None total_itpm: int | None = None total_otpm: int | None = None configurable_clientside_auth_params: CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS = None reasoning_efforts_initialized = False reasoning_efforts_unknown = False model_list: Final = self.get_model_list(model_name=model_group) if model_list is None: return None for model in model_list: is_match = False if ( "model_name" in model and model["model_name"] == model_group or "model_name" in model and self.pattern_router.route(model_group) is not None ): # exact match is_match = True if not is_match: continue # model in model group found # litellm_params = LiteLLM_Params(**model["litellm_params"]) # get configurable clientside auth params configurable_clientside_auth_params = litellm_params.configurable_clientside_auth_params # Cache nested dict access to avoid repeated temporary dict allocations model_litellm_params = model.get("litellm_params", {}) model_info_dict = model.get("model_info", {}) # get model tpm _deployment_tpm: int | None = None if _deployment_tpm is None: _deployment_tpm = model.get("tpm", None) if _deployment_tpm is None: _deployment_tpm = model_litellm_params.get("tpm", None) if _deployment_tpm is None: _deployment_tpm = model_info_dict.get("tpm", None) # get model rpm _deployment_rpm: int | None = None if _deployment_rpm is None: _deployment_rpm = model.get("rpm", None) if _deployment_rpm is None: _deployment_rpm = model_litellm_params.get("rpm", None) if _deployment_rpm is None: _deployment_rpm = model_info_dict.get("rpm", None) _deployment_itpm: int | None = model.get("itpm") if _deployment_itpm is None: _deployment_itpm = model_litellm_params.get("itpm", None) if _deployment_itpm is None: _deployment_itpm = model_info_dict.get("itpm", None) _deployment_otpm: int | None = model.get("otpm") if _deployment_otpm is None: _deployment_otpm = model_litellm_params.get("otpm", None) if _deployment_otpm is None: _deployment_otpm = model_info_dict.get("otpm", None) # get model info try: model_id = model_info_dict.get("id", None) if model_id is not None: model_info = self.get_deployment_model_info(model_id=model_id, model_name=litellm_params.model) else: model_info = None except Exception: model_info = None deployment_is_mapped = deployment_is_catalog_mapped(model_info, model_info_dict) # get llm provider litellm_model, llm_provider = "", "" try: litellm_model, llm_provider, _, _ = litellm.get_llm_provider( model=litellm_params.model, custom_llm_provider=litellm_params.custom_llm_provider, ) except litellm.exceptions.BadRequestError as e: verbose_router_logger.error("litellm.router.py::get_model_group_info() - %s", e) if model_info is None: supported_openai_params = litellm.get_supported_openai_params( model=litellm_model, custom_llm_provider=llm_provider ) if supported_openai_params is None: supported_openai_params = [] # Get mode from database model_info if available, otherwise default to "chat" db_model_info = model.get("model_info", {}) mode = db_model_info.get("mode", "chat") input_cost_per_token = _cost_value_as_float(db_model_info.get("input_cost_per_token")) output_cost_per_token = _cost_value_as_float(db_model_info.get("output_cost_per_token")) model_info = ModelMapInfo( key=model_group, max_tokens=None, max_input_tokens=None, max_output_tokens=None, input_cost_per_token=input_cost_per_token, output_cost_per_token=output_cost_per_token, litellm_provider=llm_provider, mode=mode, supported_openai_params=supported_openai_params, supports_system_messages=None, ) if model_group_info is None: model_group_info = ModelGroupInfo( **{ "model_group": user_facing_model_group_name, "providers": [llm_provider], **model_info, "supports_fast_mode": True, "supported_reasoning_efforts": None, } ) else: # if max_input_tokens > curr # if max_output_tokens > curr # if input_cost_per_token > curr # if output_cost_per_token > curr # supports_parallel_function_calling == True # supports_vision == True # supports_function_calling == True if llm_provider not in model_group_info.providers: model_group_info.providers.append(llm_provider) if ( model_info.get("max_input_tokens", None) is not None and model_info["max_input_tokens"] is not None and ( model_group_info.max_input_tokens is None or model_info["max_input_tokens"] > model_group_info.max_input_tokens ) ): model_group_info.max_input_tokens = model_info["max_input_tokens"] if ( model_info.get("max_output_tokens", None) is not None and model_info["max_output_tokens"] is not None and ( model_group_info.max_output_tokens is None or model_info["max_output_tokens"] > model_group_info.max_output_tokens ) ): model_group_info.max_output_tokens = model_info["max_output_tokens"] _input_cost_per_token = _cost_value_as_float(model_info.get("input_cost_per_token")) if _input_cost_per_token is not None and ( model_group_info.input_cost_per_token is None or _input_cost_per_token > (model_group_info.input_cost_per_token or 0.0) ): model_group_info.input_cost_per_token = _input_cost_per_token _output_cost_per_token = _cost_value_as_float(model_info.get("output_cost_per_token")) if _output_cost_per_token is not None and ( model_group_info.output_cost_per_token is None or _output_cost_per_token > (model_group_info.output_cost_per_token or 0.0) ): model_group_info.output_cost_per_token = _output_cost_per_token if ( model_info.get("supports_parallel_function_calling", None) is not None and model_info["supports_parallel_function_calling"] is True ): model_group_info.supports_parallel_function_calling = True if model_info.get("supports_vision", None) is not None and model_info["supports_vision"] is True: model_group_info.supports_vision = True if ( model_info.get("supports_function_calling", None) is not None and model_info["supports_function_calling"] is True ): model_group_info.supports_function_calling = True if ( model_info.get("supports_web_search", None) is not None and model_info["supports_web_search"] is True ): model_group_info.supports_web_search = True if ( model_info.get("supports_url_context", None) is not None and model_info["supports_url_context"] is True ): model_group_info.supports_url_context = True if model_info.get("supports_reasoning", None) is not None and model_info["supports_reasoning"] is True: model_group_info.supports_reasoning = True if ( model_info.get("supported_openai_params", None) is not None and model_info["supported_openai_params"] is not None ): model_group_info.supported_openai_params = model_info["supported_openai_params"] if model_info.get("tpm", None) is not None and _deployment_tpm is None: _deployment_tpm = model_info.get("tpm") if model_info.get("rpm", None) is not None and _deployment_rpm is None: _deployment_rpm = model_info.get("rpm") model_group_info.supports_fast_mode = model_group_info.supports_fast_mode and ( AnthropicModelInfo.supports_fast_mode(litellm_model, llm_provider) ) deployment_reasoning_efforts = resolve_supported_reasoning_efforts( model_info, deployment_is_mapped=deployment_is_mapped ) if deployment_reasoning_efforts is None: reasoning_efforts_unknown = True model_group_info.supported_reasoning_efforts = None elif not reasoning_efforts_initialized: reasoning_efforts_initialized = True if not reasoning_efforts_unknown: model_group_info.supported_reasoning_efforts = deployment_reasoning_efforts elif not reasoning_efforts_unknown: model_group_info.supported_reasoning_efforts = intersect_supported_reasoning_efforts( model_group_info.supported_reasoning_efforts, deployment_reasoning_efforts, ) if _deployment_tpm is not None: if total_tpm is None: total_tpm = 0 total_tpm += _deployment_tpm if _deployment_rpm is not None: if total_rpm is None: total_rpm = 0 total_rpm += _deployment_rpm if _deployment_itpm is not None: if total_itpm is None: total_itpm = 0 total_itpm += _deployment_itpm if _deployment_otpm is not None: if total_otpm is None: total_otpm = 0 total_otpm += _deployment_otpm if model_group_info is not None: ## UPDATE WITH TOTAL TPM/RPM FOR MODEL GROUP if total_tpm is not None: model_group_info.tpm = total_tpm if total_rpm is not None: model_group_info.rpm = total_rpm if total_itpm is not None: model_group_info.itpm = total_itpm if total_otpm is not None: model_group_info.otpm = total_otpm ## UPDATE WITH CONFIGURABLE CLIENTSIDE AUTH PARAMS FOR MODEL GROUP if configurable_clientside_auth_params is not None: model_group_info.configurable_clientside_auth_params = configurable_clientside_auth_params return model_group_info def get_model_group_info(self, model_group: str) -> ModelGroupInfo | None: """ For a given model group name, return the combined model info Returns: - ModelGroupInfo if able to construct a model group - None if error constructing model group info or hidden model group """ ## Check if model group alias if model_group in self.model_group_alias: item: Final = self.model_group_alias[model_group] if isinstance(item, str): _router_model_group = item elif isinstance(item, dict): if item["hidden"] is True: return None else: _router_model_group = item["model"] else: return None return self._set_model_group_info( model_group=_router_model_group, user_facing_model_group_name=model_group, ) ## Check if actual model return self._set_model_group_info(model_group=model_group, user_facing_model_group_name=model_group) async def get_model_group_usage(self, model_group: str) -> tuple[int | None, int | None]: """ Returns current tpm/rpm usage for model group Parameters: - model_group: str - the received model name from the user (can be a wildcard route). Returns: - usage: Tuple[tpm, rpm] """ dt: Final = get_utc_datetime() current_minute: Final = dt.strftime("%H-%M") # use the same timezone regardless of system clock tpm_keys: Final[list[str]] = [] rpm_keys: Final[list[str]] = [] model_list: Final = self.get_model_list(model_name=model_group) if model_list is None: # no matching deployments return None, None for model in model_list: id: str | None = model.get("model_info", {}).get("id") litellm_model: str | None = model["litellm_params"].get( "model" ) # USE THE MODEL SENT TO litellm.completion() - consistent with how global_router cache is written. if id is None or litellm_model is None: continue tpm_keys.append( RouterCacheEnum.TPM.value.format( id=id, model=litellm_model, current_minute=current_minute, ) ) rpm_keys.append( RouterCacheEnum.RPM.value.format( id=id, model=litellm_model, current_minute=current_minute, ) ) combined_tpm_rpm_keys: Final = tpm_keys + rpm_keys combined_tpm_rpm_values: Final = await self.cache.async_batch_get_cache(keys=combined_tpm_rpm_keys) if combined_tpm_rpm_values is None: return None, None tpm_usage_list: Final[list | None] = combined_tpm_rpm_values[: len(tpm_keys)] rpm_usage_list: Final[list | None] = combined_tpm_rpm_values[len(tpm_keys) :] ## TPM tpm_usage: int | None = None if tpm_usage_list is not None: for t in tpm_usage_list: if isinstance(t, int): if tpm_usage is None: tpm_usage = 0 tpm_usage += t ## RPM rpm_usage: int | None = None if rpm_usage_list is not None: for t in rpm_usage_list: if isinstance(t, int): if rpm_usage is None: rpm_usage = 0 rpm_usage += t return tpm_usage, rpm_usage async def get_model_group_io_token_usage(self, model_group: str) -> tuple[int | None, int | None]: """ Returns current ITPM/OTPM usage for a model group (sum across deployments). """ dt: Final = get_utc_datetime() current_minute: Final = dt.strftime("%H-%M") itpm_keys: Final[list[str]] = [] otpm_keys: Final[list[str]] = [] model_list: Final = self.get_model_list(model_name=model_group) if model_list is None: return None, None for model in model_list: model_id: str | None = model.get("model_info", {}).get("id") litellm_model: str | None = model["litellm_params"].get("model") if model_id is None or litellm_model is None: continue itpm_keys.append( RouterCacheEnum.ITPM.value.format( id=model_id, model=litellm_model, current_minute=current_minute, ) ) otpm_keys.append( RouterCacheEnum.OTPM.value.format( id=model_id, model=litellm_model, current_minute=current_minute, ) ) combined_values: Final = await self.cache.async_batch_get_cache(keys=itpm_keys + otpm_keys) if combined_values is None: return None, None itpm_values: Final = combined_values[: len(itpm_keys)] otpm_values: Final = combined_values[len(itpm_keys) :] total_itpm: int | None = None for value in itpm_values: if isinstance(value, int): total_itpm = (total_itpm or 0) + value total_otpm: int | None = None for value in otpm_values: if isinstance(value, int): total_otpm = (total_otpm or 0) + value return total_itpm, total_otpm @lru_cache(maxsize=DEFAULT_MAX_LRU_CACHE_SIZE) def _cached_get_model_group_info(self, model_group: str) -> ModelGroupInfo | None: """ Cached version of get_model_group_info, uses @lru_cache wrapper This is a speed optimization, since set_response_headers makes a call to get_model_group_info on every request """ return self.get_model_group_info(model_group) def cached_model_group_info(self, model_group: str) -> ModelGroupInfo | None: return self._cached_get_model_group_info(model_group) async def get_remaining_model_group_usage(self, model_group: str) -> dict[str, int]: model_group_info: Final = self._cached_get_model_group_info(model_group) returned_dict: Final[dict[str, int]] = {} # ITPM/OTPM groups emit input/output token headers, but they may also set # tpm/rpm, so build both sets rather than returning early - clients and # prometheus gauges that read the standard headers still get data. if model_group_info is not None and (model_group_info.itpm is not None or model_group_info.otpm is not None): current_itpm, current_otpm = await self.get_model_group_io_token_usage(model_group) returned_dict.update( build_io_token_rate_limit_headers( itpm_limit=model_group_info.itpm, otpm_limit=model_group_info.otpm, current_itpm=current_itpm, current_otpm=current_otpm, ) ) tpm_limit: Final = model_group_info.tpm if model_group_info is not None else None rpm_limit: Final = model_group_info.rpm if model_group_info is not None else None if tpm_limit is None and rpm_limit is None: return returned_dict current_tpm, current_rpm = await self.get_model_group_usage(model_group) if tpm_limit is not None: returned_dict["x-ratelimit-remaining-tokens"] = tpm_limit - (current_tpm or 0) returned_dict["x-ratelimit-limit-tokens"] = tpm_limit if rpm_limit is not None: returned_dict["x-ratelimit-remaining-requests"] = rpm_limit - (current_rpm or 0) returned_dict["x-ratelimit-limit-requests"] = rpm_limit return returned_dict async def set_response_headers( self, response: object, model_group: str | None = None, request_kwargs: dict[str, object] | None = None, ) -> Any: """ Add the most accurate rate limit headers for a given model response. ## TODO: add model group rate limit headers # - if healthy_deployments > 1, return model group rate limit headers # - else return the model's rate limit headers """ response = prepare_response_for_header_attachment(response) if response is None: return response additional_headers: Final = ensure_response_additional_headers(response) apply_response_model_id( response, find_deployment_metadata(request_kwargs) if request_kwargs is not None else None ) additional_headers["x-litellm-model-group"] = model_group apply_quality_router_decision_headers(additional_headers, request_kwargs) additional_headers.update(complexity_router_decision_headers(request_kwargs)) if model_group is not None: remaining_usage: Final = await self.get_remaining_model_group_usage(model_group) apply_remaining_usage_headers(additional_headers, remaining_usage) return response def _build_model_name_index(self, model_list: list) -> None: """ Build model_name -> deployment indices mapping for O(1) lookups. This index allows us to find all deployments for a given model_name in O(1) time instead of O(n) linear scan through the entire model_list. """ self.model_name_to_deployment_indices.clear() self.team_model_to_deployment_indices.clear() self.team_public_model_names = frozenset() for idx, model in enumerate(model_list): model_name = model.get("model_name") if model_name: if model_name not in self.model_name_to_deployment_indices: self.model_name_to_deployment_indices[model_name] = [] self.model_name_to_deployment_indices[model_name].append(idx) self._update_team_model_index(model, idx) def _build_model_id_to_deployment_index_map(self, model_list: list): """ Build model index from model list to enable O(1) lookups immediately. This is called during initialization to avoid the race condition where requests arrive before model_id_to_deployment_index_map is populated. """ # First populate the model_list self.model_list = [] self._invalidate_model_group_info_cache() self._invalidate_access_groups_cache() for _, model in enumerate(model_list): # Extract model_info from the model dict model_info = model.get("model_info", {}) model_id = model_info.get("id") # If no ID exists, generate one using the same logic as set_model_list if model_id is None: model_name = model.get("model_name", "") litellm_params = model.get("litellm_params", {}) model_id = self.generate_model_id(model_name, litellm_params) # Update the model_info in the original list if "model_info" not in model: model["model_info"] = {} model["model_info"]["id"] = model_id self._add_model_to_list_and_index_map(model=model, model_id=model_id) def get_model_ids(self, model_name: str | None = None, exclude_team_models: bool = False) -> list[str]: """ if 'model_name' is none, returns all. Returns list of model id's. Optimized with O(1) or O(k) index lookup when model_name provided, instead of O(n) linear scan. """ ids: Final = [] if model_name is not None: # O(1) lookup in model_name index, then O(k) iteration where k = deployments for this model_name if model_name in self.model_name_to_deployment_indices: indices: Final = self.model_name_to_deployment_indices[model_name] for idx in indices: model = self.model_list[idx] if "model_info" in model and "id" in model["model_info"]: if exclude_team_models and model["model_info"].get("team_id"): continue ids.append(model["model_info"]["id"]) else: # When model_name is None, return all model IDs # Use the index map keys for O(n) where n = total deployments for model_id in self.model_id_to_deployment_index_map: idx = self.model_id_to_deployment_index_map[model_id] model = self.model_list[idx] if "model_info" in model and "id" in model["model_info"]: if exclude_team_models and model["model_info"].get("team_id"): continue ids.append(model_id) return ids def get_candidate_model_ids_for_route(self, model: str, team_id: str | None = None) -> frozenset[str]: """ Deployment ids that could serve ``model`` for ``team_id``, following the same precedence ``_common_checks_available_deployment`` uses to build a candidate pool: ``model_group_alias``, then a routing group, then the first matching early-resolve path for a name that is not a ``model_name`` (team route, wildcard pattern via ``get_deployments_by_pattern``, team pattern router, default deployment), then the ``model_name`` and team indexes. Delegating to the router's own resolvers keeps this aligned with how a route actually resolves rather than re-deriving it, and unlike ``_common_checks_available_deployment`` it is read-only: it does not apply request fallbacks and (with ``include_team_models`` left off) does not raise. Lets a pre-call check tell a genuine cross-group route from same-group unavailability without leaking deployment ids into request kwargs bound for the provider. """ resolved: Final = self._get_model_from_alias(model=model) or model routing_group_members: Final = self._get_routing_group_deployments(model=resolved, team_id=team_id) if routing_group_members is not None: return self._deployment_ids(routing_group_members) early: Final = self._try_early_resolve_deployments_for_model_not_in_names( model=resolved, request_team_id=team_id ) if early is not None: early_deployments: Final = early[1] return self._deployment_ids( (early_deployments,) if isinstance(early_deployments, Mapping) else early_deployments ) return self._deployment_ids(self._get_all_deployments(model_name=resolved, team_id=team_id)) @staticmethod def _deployment_ids(deployments: Sequence[Mapping[str, object]]) -> frozenset[str]: return frozenset( str(model_info["id"]) for deployment in deployments for model_info in (deployment.get("model_info"),) if isinstance(model_info, Mapping) and model_info.get("id") is not None ) def has_model_id(self, candidate_id: str) -> bool: """ O(1) membership check for a deployment ID without allocating large lists. Note: Call sites may pass a variable named `model` when it actually contains a deployment ID. This helper expects the deployment ID string. Uses the existing `model_id_to_deployment_index_map` which is kept in sync by `_build_model_id_to_deployment_index_map` and model-list mutation helpers. """ return candidate_id in self.model_id_to_deployment_index_map def resolve_model_name_from_model_id(self, model_id: str | None) -> str | None: """ Resolve model_name from model_id. This method attempts to find the correct model_name to use with the router so that litellm_params can be automatically injected from the model config. Strategy: 1. First, check if model_id directly matches a model_name or deployment ID 2. If not, search through router's model_list to find a match by litellm_params.model 3. Return the model_name if found, None otherwise Args: model_id: The model_id extracted from decoded video_id (could be model_name or litellm_params.model value) Returns: model_name if found, None otherwise. If None, the request will fall through to normal flow using environment variables. """ if not model_id: return None # Strategy 1: Check if model_id directly matches a model_name or deployment ID if model_id in self.model_names: return model_id if self.has_model_id(model_id): deployment = self.get_deployment(model_id=model_id) if deployment is not None and deployment.model_name: return deployment.model_name return model_id # Strategy 2: Search through router's model_list to find by litellm_params.model all_models: Final = self.get_model_list(model_name=None) if not all_models: return None for deployment in all_models: litellm_params = deployment.get("litellm_params", {}) actual_model = litellm_params.get("model") # Match by exact match or by checking if actual_model ends with /model_id or :model_id # e.g., model_id="veo-2.0-generate-001" matches actual_model="vertex_ai/veo-2.0-generate-001" matches = ( actual_model == model_id or (actual_model and actual_model.endswith(f"/{model_id}")) or (actual_model and actual_model.endswith(f":{model_id}")) ) if matches: model_name = deployment.get("model_name") if model_name: return model_name # No match found return None def map_team_model(self, team_model_name: str | None, team_id: str) -> str | None: """ Check if team_model_name resolves to team-specific deployments. Returns the public model name (unchanged) so the router can find all sibling deployments via team_id filtering, instead of collapsing to a single internal model_name. When team_model_name is None (e.g. vector store / file endpoints that don't include a model in their request), returns the first matching team deployment's team_public_model_name so the router can inject BYOK credentials from the team-scoped deployment. Returns: - str: the team_model_name if team deployments exist for this team - None: if no team-specific model is found """ models: Final = self.get_model_list(model_name=team_model_name, team_id=team_id) if not models: return None for model in models: if model.get("model_info", {}).get("team_id") == team_id: if team_model_name is None: # No model was specified (e.g. vector store endpoints). # Return the deployment's public model name so the router # can route to it and inject the BYOK API key. return model.get("model_info", {}).get("team_public_model_name") or model.get("model_name") return team_model_name # No team-scoped deployment found; wildcard/pattern routes are # handled downstream by the pattern_router in _common_checks_available_deployment. return None def should_include_deployment(self, model_name: str, model: dict, team_id: str | None = None) -> bool: """ Get the team-specific model name if team_id matches the deployment. """ if ( team_id is not None and (model.get("model_info") or {}).get("team_id") == team_id and model_name == (model.get("model_info") or {}).get("team_public_model_name") ): return True elif model_name is not None and model["model_name"] == model_name: # Fallback: check by internal model_name for non-team deployments # or deployments that haven't been migrated to team_public_model_name yet model_team_id: Final = (model.get("model_info") or {}).get("team_id") if ( team_id is None # requester has no team constraint or model_team_id is None # global deployment - accessible to all teams or model_team_id == team_id # deployment belongs to requester's team ): return True # No match: deployment is for a different team or doesn't match the requested model return False def _get_all_deployments( self, model_name: str, model_alias: str | None = None, team_id: str | None = None, ) -> list[DeploymentTypedDict]: """ Return all deployments of a model name Used for accurate 'get_model_list'. if team_id specified, only return team-specific models Optimized with O(1) index lookup instead of O(n) linear scan. Note: when team_id is provided, O(1) lookup in `team_model_to_deployment_indices` only applies when `model_name` is the team public model name. If a caller passes an internal deployment model name (for example, `model_name__`), this method falls back to the standard model-name index / scan path. """ returned_models: Final[list[DeploymentTypedDict]] = [] # O(1) lookup in team_model index when team_id is provided if team_id is not None: key: Final = (team_id, model_name) if key in self.team_model_to_deployment_indices: indices = self.team_model_to_deployment_indices[key] # O(k) where k = team deployments for this model_name (typically 1-10) for idx in indices: model = self.model_list[idx] if not self.should_include_deployment(model_name=model_name, model=model, team_id=team_id): continue if model_alias is not None: alias_model = model.copy() alias_model["model_name"] = model_alias returned_models.append(alias_model) else: returned_models.append(model) if returned_models: return returned_models # O(1) lookup in model_name index if model_name in self.model_name_to_deployment_indices: indices = self.model_name_to_deployment_indices[model_name] # O(k) where k = deployments for this model_name (typically 1-10) for idx in indices: model = self.model_list[idx] if self.should_include_deployment(model_name=model_name, model=model, team_id=team_id): if model_alias is not None: # Optimized: Use shallow copy since we only modify top-level model_name # This is much faster than deepcopy for nested dict structures alias_model = model.copy() alias_model["model_name"] = model_alias returned_models.append(alias_model) else: returned_models.append(model) elif team_id is not None: # Fallback: if team_id is provided and model_name not in index, # check if model_name matches any team_public_model_name # O(n) scan but only when team_id lookup fails for idx, model in enumerate(self.model_list): if self.should_include_deployment(model_name=model_name, model=model, team_id=team_id): if model_alias is not None: # Optimized: Use shallow copy since we only modify top-level model_name alias_model = model.copy() alias_model["model_name"] = model_alias returned_models.append(alias_model) else: returned_models.append(model) return returned_models def get_model_names(self, team_id: str | None = None) -> list[str]: """ Returns all possible model names for the router, including models defined via model_group_alias. If a team_id is provided, only deployments configured with that team_id (i.e. team‐specific models) will yield their team public name. """ deployments: Final = self.get_model_list() or [] model_names: Final = [] for deployment in deployments: model_info = deployment.get("model_info") if self._is_team_specific_model(model_info): team_model_name = self._get_team_specific_model(deployment=deployment, team_id=team_id) if team_model_name: model_names.append(team_model_name) else: model_names.append(deployment.get("model_name", "")) return model_names def get_fully_blocked_model_names(self) -> set[str]: """ Returns the set of model_names where every backing deployment has `blocked=True`. Used by `/v1/models` to hide paused models from client listings while still surfacing them on admin endpoints (e.g. `/model/info`). A model with at least one non-blocked deployment is still serviceable and remains visible. """ deployments: Final = self.get_model_list() or [] blocked_by_name: Final[dict[str, bool]] = {} for deployment in deployments: name = deployment.get("model_name") or "" if not name: continue is_blocked = (deployment.get("model_info") or {}).get("blocked") is True if name in blocked_by_name: blocked_by_name[name] = blocked_by_name[name] and is_blocked else: blocked_by_name[name] = is_blocked return {name for name, fully_blocked in blocked_by_name.items() if fully_blocked} @staticmethod def _are_all_deployments_blocked( deployments: list[DeploymentTypedDict], ) -> bool: return len(deployments) > 0 and all( (deployment.get("model_info") or {}).get("blocked") is True for deployment in deployments ) def _is_model_fully_blocked(self, model: str) -> bool: deployments: Final = self.get_model_list(model_name=model) or [] return self._are_all_deployments_blocked(deployments=deployments) async def async_get_fully_unhealthy_model_names(self) -> set[str]: """ Returns the set of model names where every backing deployment is currently marked unhealthy by background health checks (and the health state is not stale). Used by `/v1/models?healthy_only=true` to hide models that cannot serve any request. A model with at least one healthy (or unknown-health) deployment remains visible. Returns an empty set when no health state is available, so callers fail open to the unfiltered listing. Notes: - Mirrors `_async_filter_health_check_unhealthy_deployments`: when `allowed_fails_policy` is set, cooldown is the sole routing exclusion mechanism, so nothing is hidden here either. - Team-specific public model names (`team_public_model_name`) are aggregated alongside `model_name`, so team aliases of fully-unhealthy deployments are hidden too (unlike `get_fully_blocked_model_names`, which matches `model_name` only). - Wildcard routes (e.g. `openai/*`) are matched by their literal deployment name only; models expanded from a wildcard route are not hidden (fail open). - Intentionally diverges from the routing-time safety net (which bypasses the health filter when every candidate is unhealthy and still attempts the request): hiding here is presentation-only — it answers "should this model be advertised?", not "should a request for it still be attempted?". A hidden model can still be called directly. """ if self.allowed_fails_policy is not None: return set() unhealthy_ids: Final = await self.health_state_cache.async_get_unhealthy_deployment_ids() if not unhealthy_ids: return set() deployments: Final = self.get_model_list() or [] unhealthy_by_name: Final[dict[str, bool]] = {} for deployment in deployments: model_info = deployment.get("model_info") or {} names = [deployment.get("model_name") or ""] team_public_model_name = model_info.get("team_public_model_name") if team_public_model_name: names.append(team_public_model_name) is_unhealthy = model_info.get("id") in unhealthy_ids for name in names: if not name: continue if name in unhealthy_by_name: unhealthy_by_name[name] = unhealthy_by_name[name] and is_unhealthy else: unhealthy_by_name[name] = is_unhealthy return {name for name, fully_unhealthy in unhealthy_by_name.items() if fully_unhealthy} def _get_team_specific_model(self, deployment: DeploymentTypedDict, team_id: str | None = None) -> str | None: """ Get the team-specific model name if team_id matches the deployment. Args: deployment: DeploymentTypedDict - The model deployment team_id: Optional[str] - If passed, will return router models set with a `team_id` matching the passed `team_id`. Returns: str: The `team_public_model_name` if team_id matches None: If team_id doesn't match or no team info exists """ model_info: Final[dict | None] = deployment.get("model_info") or {} if model_info is None: return None if team_id == model_info.get("team_id"): return model_info.get("team_public_model_name") return None def _is_team_specific_model(self, model_info: dict | None) -> bool: """ Check if model info contains team-specific configuration. Args: model_info: Model information dictionary Returns: bool: True if model has team-specific configuration """ return bool(model_info and model_info.get("team_id")) def get_model_list_from_model_alias(self, model_name: str | None = None) -> list[DeploymentTypedDict]: """ Helper function to get model list from model alias. Used by `.get_model_list` to get model list from model alias. """ returned_models: Final[list[DeploymentTypedDict]] = [] if model_name is not None: # Fast path: direct dict lookup avoids scanning all aliases for non-alias model names. if model_name not in self.model_group_alias: return returned_models alias_items = [(model_name, self.model_group_alias[model_name])] else: alias_items = list(self.model_group_alias.items()) for model_alias, model_value in alias_items: if isinstance(model_value, str): _router_model_name: str = model_value elif isinstance(model_value, dict): _model_value = RouterModelGroupAliasItem(**model_value) if _model_value["hidden"] is True and model_name is None: continue _router_model_name = _model_value["model"] else: continue if (alias_group := self.get_routing_group(_router_model_name)) is not None: returned_models.extend( {**row, "model_name": model_alias} for row in self._materialize_routing_group_rows((alias_group,)) ) else: returned_models.extend( self._get_all_deployments(model_name=_router_model_name, model_alias=model_alias) ) return returned_models def get_model_list_from_routing_groups(self, model_name: str | None = None) -> Sequence[DeploymentTypedDict]: """ Callable routing groups materialized as model-list rows, mirroring `get_model_list_from_model_alias`: each member deployment is emitted under the group's name (via `_get_all_deployments`' `model_alias` rewrite), which is what surfaces groups in `get_model_names`, `/v1/models` discovery, `get_model_group_usage`, and the blocked/unhealthy hiding that all read `get_model_list`. """ if model_name is not None: group: Final = self.get_routing_group(model_name) return self._materialize_routing_group_rows((group,)) if group is not None else () cached: Final = self._routing_group_rows if cached is not None: return cached rows: Final = self._materialize_routing_group_rows( tuple( callable_group for name in self._routing_groups if (callable_group := self.get_routing_group(name)) is not None ) ) self._routing_group_rows = rows return rows def _materialize_routing_group_rows(self, groups: tuple[RoutingGroup, ...]) -> tuple[DeploymentTypedDict, ...]: return tuple( self._as_routing_group_row(apply_routing_group_priority(group, member, deployment)) for group in groups for member in group.models for deployment in self._get_all_deployments(model_name=member, model_alias=group.group_name) ) @staticmethod def _as_routing_group_row(deployment: DeploymentTypedDict) -> DeploymentTypedDict: """ A member deployment re-emitted under its group's name must not carry the member's `access_groups`: access groups grant member names, never the group, so inheriting them here would let a key holding a member's access group list and call the whole group. """ model_info: Final = { # mutable-ok: DeploymentTypedDict rows are plain dicts k: v for k, v in (deployment.get("model_info") or {}).items() if k != "access_groups" } return {**deployment, "model_info": model_info} # mutable-ok: DeploymentTypedDict rows are plain dicts TIER_PARAMS_NEVER_DROPPED: Final = frozenset(all_litellm_params) | frozenset( { "additional_drop_params", "drop_params", "messages", "model", "extra_headers", "max_tokens", "max_completion_tokens", } ) @staticmethod def _declared_param_allowlist(params: Mapping[str, object]) -> frozenset[str]: declared: Final = params.get("allowed_openai_params") if not isinstance(declared, (list, tuple, set, frozenset)): return frozenset() return frozenset(entry for entry in declared if isinstance(entry, str)) @staticmethod def _deployment_accepts_param(deployment: DeploymentTypedDict, group: str, param: str) -> bool: deployment_params: Final = deployment.get("litellm_params") if not deployment_params: return True if param in Router._declared_param_allowlist(deployment_params): return True if declared_authenticating_provider( str(deployment_params.get("model") or ""), deployment_params.get("custom_llm_provider") ): return True deployment_model_info: Final = deployment.get("model_info") base_model: Final = ( deployment_model_info.get("base_model") if deployment_model_info else None ) or deployment_params.get("base_model") try: model, custom_llm_provider, _, _ = litellm.get_llm_provider( model=deployment_params.get("model") or group, custom_llm_provider=deployment_params.get("custom_llm_provider"), ) supported: Final = litellm.get_supported_openai_params( model=model, custom_llm_provider=custom_llm_provider, base_model=base_model if isinstance(base_model, str) else None, ) except Exception as e: # noqa: BLE001 # best-effort filter: an unresolvable provider must not narrow the request verbose_router_logger.debug( "litellm.router.py::_deployment_accepts_param: keeping %s for model=%s. Got - %s", param, group, e ) return True return supported is None or param in supported def _tier_params_the_target_accepts( self, model: str, tier_params: Mapping[str, object], request_kwargs: Mapping[str, object] ) -> Mapping[str, object]: """Drop an OpenAI param that no deployment behind ``model`` declares. A tier's litellm_params are an operator override applied to every request the tier routes, so one the target cannot take turns that whole tier into a 400 raised before the request leaves the proxy. The candidates are exactly what get_optional_params can reject, asked of the module that raises, so credentials and endpoint controls are never at risk. TIER_PARAMS_NEVER_DROPPED is excluded on top of that, for two reasons. No provider lists a litellm control among its supported params, so "no deployment declares it" means litellm consumes it rather than that the target refuses it, and dropping one changes litellm's own behavior: dropping drop_params or additional_drop_params silently disables the sanitization the operator configured. Providers do list extra_headers, but it carries auth, tenancy and routing information, so sending fewer headers than configured is worse than today's error. Token ceilings stay for the same reason: a tier's max_tokens or max_completion_tokens is a cost bound, and dropping it would let a caller's own larger value through where today the mismatch fails loudly. The trade this filter makes is a param for a working request, which is right for one that only shapes how the model answers and wrong for anything else. A param survives if ANY deployment could take it, because routing has not chosen one yet, and it survives both an unresolvable provider and a group with no deployments, because a best-effort filter must never narrow what the request already did. A github_copilot or chatgpt deployment counts as accepting everything, decided before any lookup: resolving either provider runs its OAuth device flow, so a capability question asked from the routing path can freeze the event loop for minutes waiting on a human. allowed_openai_params is the documented escape hatch for an outdated or incomplete supported-params list: request-time validation extends the supported list with it before comparing. The filter asks the same question, so a param named by the allowlist on the tier overlay, the request, or a deployment's own litellm_params is never a drop candidate. """ deployments: Final = self.get_model_list(model_name=model) or () if not deployments: return tier_params allowlisted: Final = self._declared_param_allowlist(tier_params) | self._declared_param_allowlist( request_kwargs ) candidates: Final = provider_rejectable_params(tier_params) - self.TIER_PARAMS_NEVER_DROPPED - allowlisted unsupported: Final = frozenset( param for param in candidates if not any(self._deployment_accepts_param(deployment, model, param) for deployment in deployments) ) if not unsupported: return tier_params verbose_router_logger.warning( "litellm.router.py: dropping tier params %s for model=%s, no deployment behind it declares them", ", ".join(sorted(unsupported)), model, ) return MappingProxyType({key: value for key, value in tier_params.items() if key not in unsupported}) def get_model_list( self, model_name: str | None = None, team_id: str | None = None ) -> list[DeploymentTypedDict] | None: """ Includes router model_group_alias'es as well if team_id specified, returns matching team-specific models """ # Note: model_list and model_group_alias are always initialized in __init__ # so hasattr checks are unnecessary returned_models: list[DeploymentTypedDict] = [] if model_name is not None: returned_models.extend(self._get_all_deployments(model_name=model_name, team_id=team_id)) returned_models.extend(self.get_model_list_from_model_alias(model_name=model_name)) returned_models.extend(self.get_model_list_from_routing_groups(model_name=model_name)) if len(returned_models) == 0: # check if wildcard route potential_wildcard_models: Final = self.pattern_router.get_deployments_by_pattern(model=model_name or "") ## check for team-specific wildcard models if team_id is not None and team_id in self.team_pattern_routers: potential_team_only_wildcard_models: Final = self.team_pattern_routers[ team_id ].get_deployments_by_pattern(model=model_name or "") potential_wildcard_models.extend(potential_team_only_wildcard_models) if model_name is not None and potential_wildcard_models is not None: for m in potential_wildcard_models: deployment_typed_dict = DeploymentTypedDict(**m) deployment_typed_dict["model_name"] = model_name returned_models.append(deployment_typed_dict) if model_name is None: returned_models += self.model_list return returned_models def resolved_litellm_models(self, model_name: str, team_id: str | None = None) -> tuple[str, ...]: """The provider model strings `model_name` can actually be served by on this proxy. `get_model_list` composes every channel the request path itself uses (exact name, model_group_alias, routing groups, wildcards), so this answers "which models will answer a call to this name" rather than "what did the admin call it": the deployment name is admin-arbitrary, and two names over one provider model are one model. Empty when the name resolves to no deployment. That is not the same fact as "the call will fail" - a provider-qualified public name is served by the SDK with no deployment behind it - so the fallback for an empty result is the caller's policy, never this function's. """ return tuple( litellm_model for deployment in self.get_model_list(model_name=model_name, team_id=team_id) or () if isinstance(litellm_model := deployment.get("litellm_params", {}).get("model"), str) and litellm_model ) def _invalidate_model_group_info_cache(self) -> None: """Invalidate the cached model group info. Call this whenever self.model_list is modified to ensure the cache is rebuilt. Also clears the auth-layer zero-cost cache, which depends on the same ``ModelGroupInfo`` data — without this, an in-place pricing update on an existing deployment (same model count) would keep a stale ``True`` result and bypass budget enforcement. """ self._cached_get_model_group_info.cache_clear() self.cached_deployment_model_info.cache_clear() self._zero_cost_cache.clear() self._routing_group_rows = None def _invalidate_access_groups_cache(self) -> None: """Invalidate the cached access groups. Call this whenever self.model_list is modified to ensure the cache is rebuilt. """ self._access_groups_cache = None def get_model_access_groups( self, model_name: str | None = None, model_access_group: str | None = None, team_id: str | None = None, ) -> dict[str, list[str]]: """ If model_name is provided, only return access groups for that model. Parameters: - model_name: Optional[str] - the received model name from the user (can be a wildcard route). If set, will only return access groups for that model. - model_access_group: Optional[str] - the received model access group from the user. If set, will only return models for that access group. - team_id: Optional[str] - the team id, to resolve team-specific models """ # Check if this is the no-args hot path (cacheable) _use_cache: Final = model_name is None and model_access_group is None and team_id is None # Return cached result for the no-args hot path if _use_cache and self._access_groups_cache is not None: return self._access_groups_cache from collections import defaultdict access_groups: Final = defaultdict(list) model_list: Final = self.get_model_list(model_name=model_name, team_id=team_id) if model_list: for m in model_list: _model_info = m.get("model_info") if _model_info: for group in _model_info.get("access_groups", []) or []: if model_access_group is not None: if group == model_access_group: model_name = m["model_name"] access_groups[group].append(model_name) else: model_name = m["model_name"] access_groups[group].append(model_name) # Cache the result for the no-args hot path if _use_cache: self._access_groups_cache = dict(access_groups) return self._access_groups_cache return access_groups def _is_model_access_group_for_wildcard_route(self, model_access_group: str) -> bool: """ Return True if model access group is a wildcard route """ # GET ACCESS GROUPS access_groups: Final = self.get_model_access_groups(model_access_group=model_access_group) if len(access_groups) == 0: return False models: Final = access_groups.get(model_access_group, []) for model in models: # CHECK IF MODEL ACCESS GROUP IS A WILDCARD ROUTE if self.pattern_router.route(request=model) is not None: return True return False def get_settings(self): """ Get router settings method, returns a dictionary of the settings and their values. For example get the set values for routing_strategy_args, routing_strategy, allowed_fails, cooldown_time, num_retries, timeout, max_retries, retry_after """ _all_vars: Final = vars(self) _settings_to_return: Final = {} vars_to_include: Final = [ "routing_strategy_args", "routing_strategy", "allowed_fails", "cooldown_time", "num_retries", "timeout", "max_retries", "retry_after", "fallbacks", "context_window_fallbacks", "model_group_retry_policy", "retry_policy", "model_group_alias", "enable_weighted_failover", "enable_tag_filtering", "tag_routing_prefix", ] for var in vars_to_include: if var in _all_vars: _settings_to_return[var] = _all_vars[var] if ( var == "routing_strategy_args" and self.routing_strategy == "latency-based-routing" and self.lowestlatency_logger is not None ): _settings_to_return[var] = self.lowestlatency_logger.routing_args.json() _settings_to_return["routing_groups"] = [ group.model_dump(exclude=frozenset(("model_priorities",)) if group.model_priorities is None else None) for group in self._routing_groups.values() ] return _settings_to_return def update_settings(self, **kwargs): """ Update the router settings. """ _int_settings: Final = [ "timeout", "num_retries", "retry_after", "allowed_fails", "cooldown_time", ] _existing_router_settings: Final = self.get_settings() rebuild_routing_groups = False routing_args_updated = False for var in kwargs: if var in RUNTIME_UPDATABLE_ROUTER_SETTINGS: if var in _int_settings: _casted_value = int(kwargs[var]) setattr(self, var, _casted_value) elif var == "routing_groups": rebuild_routing_groups = True elif var == "optional_pre_call_checks": self.set_optional_pre_call_checks(kwargs[var]) elif var == "retry_policy": value = kwargs[var] if isinstance(value, dict): value = RetryPolicy(**value) if value is None or isinstance(value, RetryPolicy): setattr(self, var, value) else: value = kwargs[var] # only run routing strategy init if it has changed if var == "routing_strategy": value = self._normalize_strategy(value) if _existing_router_settings["routing_strategy"] != value: if value == "lar1": from litellm.router_strategy.lar1_routing import ( apply_lar1_routing_strategy, ) apply_lar1_routing_strategy( self, kwargs.get("routing_strategy_args"), ) else: self.routing_strategy_init( routing_strategy=value, routing_strategy_args=kwargs.get("routing_strategy_args", {}), ) rebuild_routing_groups = True elif var == "routing_strategy_args": routing_args_updated = value != self.routing_strategy_args setattr(self, var, value) else: verbose_router_logger.debug("Setting %s is not allowed", var) if routing_args_updated: self._apply_updated_routing_strategy_args() if rebuild_routing_groups: routing_groups_input: Final = kwargs.get("routing_groups", self._routing_groups_input) self._init_routing_groups(routing_groups_input) self._routing_groups_input = routing_groups_input verbose_router_logger.debug("Updated Router settings: %s", self.get_settings()) def _get_client(self, deployment, kwargs, client_type=None): """ Returns the appropriate client based on the given deployment, kwargs, and client_type. Parameters: deployment (dict): The deployment dictionary containing the clients. kwargs (dict): The keyword arguments passed to the function. client_type (str): The type of client to return. Returns: The appropriate client based on the given client_type and kwargs. """ model_id: Final = deployment["model_info"]["id"] parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(kwargs) if client_type == "max_parallel_requests": cache_key = f"{model_id}_max_parallel_requests_client" client = self.cache.get_cache(key=cache_key, local_only=True, parent_otel_span=parent_otel_span) if client is None: InitalizeCachedClient.set_max_parallel_requests_client(litellm_router_instance=self, model=deployment) client = self.cache.get_cache(key=cache_key, local_only=True, parent_otel_span=parent_otel_span) return client elif client_type == "async": if kwargs.get("stream") is True: cache_key = f"{model_id}_stream_async_client" client = self.cache.get_cache(key=cache_key, local_only=True, parent_otel_span=parent_otel_span) return client else: cache_key = f"{model_id}_async_client" client = self.cache.get_cache(key=cache_key, local_only=True, parent_otel_span=parent_otel_span) return client else: if kwargs.get("stream") is True: cache_key = f"{model_id}_stream_client" client = self.cache.get_cache(key=cache_key, parent_otel_span=parent_otel_span) return client else: cache_key = f"{model_id}_client" client = self.cache.get_cache(key=cache_key, parent_otel_span=parent_otel_span) return client def _count_pre_call_check_tokens( self, messages: Sequence[Mapping[str, object]] | None, input: str | list[object] | None, request_kwargs: Mapping[str, object] | None = None, ) -> int: """ Count input tokens for context-window pre-call checks. Chat Completions send `messages`; the Responses API sends `input` (a string or a list of Responses input items) plus an optional `instructions` system prompt. The Responses payload is normalized to chat messages via the shared LiteLLMCompletionResponsesConfig transform so the same token_counter path covers both API surfaces and `instructions` tokens are included in the count. Prompt content the message list never carries is read from `request_kwargs`: `tools` (Chat Completions, Responses and Anthropic Messages shapes) and the Anthropic Messages top-level `system` block. """ from litellm.llms.anthropic.pass_through.messages.utils import ( anthropic_system_to_openai_message, ) extras: Final = request_kwargs if request_kwargs is not None else MappingProxyType({}) raw_instructions: Final = extras.get("instructions") instructions: Final = raw_instructions if isinstance(raw_instructions, str) else None raw_tools: Final = extras.get("tools") tools: Final = ( cast(list[ChatCompletionToolParam], raw_tools) # cast-ok: token_counter formats any tool dict shape if isinstance(raw_tools, list) and raw_tools else None ) system_message: Final = anthropic_system_to_openai_message(extras.get("system")) if messages is not None: counted_messages: Final = (system_message, *messages) if system_message is not None else messages return litellm.token_counter(messages=counted_messages, tools=tools) if input is not None: from openai.types.responses.response_create_params import ResponseInputParam from litellm.responses.litellm_completion_transformation.transformation import ( LiteLLMCompletionResponsesConfig, ) typed_input: Final = cast(str | ResponseInputParam, input) # cast-ok: str | list matches transform input input_messages: Final = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages( input=typed_input, responses_api_request={"instructions": instructions} if instructions is not None else {}, ) return litellm.token_counter( messages=cast(list, input_messages), # cast-ok: transformed chat messages tools=tools, ) raise ValueError("Either messages or input must be provided to count tokens") def _deployment_max_input_tokens(self, model: str, deployment: Mapping[str, object]) -> int | None: """The deployment's declared context window, or None when it declares none or cannot be resolved.""" try: model_info: Final = self.get_router_model_info( deployment=cast(dict, deployment), # cast-ok: router deployments are plain dicts received_model_name=model, ) except Exception as e: # noqa: BLE001 # best-effort: an unmappable deployment must not hide the others verbose_router_logger.debug( "litellm.router.py::_deployment_max_input_tokens: skipping deployment. Got - %s", e ) return None max_input_tokens: Final = model_info.get("max_input_tokens") return max_input_tokens if isinstance(max_input_tokens, int) else None def _pre_call_checks_need_token_count( self, model: str, healthy_deployments: Sequence[Mapping[str, object]] ) -> bool: """Whether any healthy deployment declares a context window that a token count could exceed. Resolves each deployment the way ``_pre_call_checks`` does, so one unmappable deployment cannot hide a later one that does declare a limit. """ return any( self._deployment_max_input_tokens(model, deployment) is not None for deployment in healthy_deployments ) async def _acount_pre_call_check_tokens( self, model: str, healthy_deployments: Sequence[Mapping[str, object]], messages: Sequence[Mapping[str, str]] | None, input: str | Sequence[object] | None, request_kwargs: Mapping[str, object] | None, ) -> int | None: """Count input tokens off the event loop, so a multi-MB prompt cannot stall the proxy. Returns None when no deployment limits its context window, and when counting fails. The caller pairs this with ``skip_inline_token_count`` so neither case puts the count back on the loop: a failed count leaves the deployments unfiltered, exactly as before. """ if messages is None and input is None: return None try: if not self._pre_call_checks_need_token_count(model, healthy_deployments): return None return await offload_token_count(self._count_pre_call_check_tokens)( messages=cast(list[dict[str, str]] | None, messages), # cast-ok: forwarded to the sync counter input=cast(str | list | None, input), # cast-ok: forwarded to the sync counter request_kwargs=request_kwargs, ) except Exception as e: # noqa: BLE001 # best-effort: an uncountable prompt must not fail the request verbose_router_logger.error( "litellm.router.py::_acount_pre_call_check_tokens: failed to count tokens. Got - %s", e ) return None def _pre_call_checks( self, model: str, healthy_deployments: list, messages: list[dict[str, str]] | None = None, input: str | list | None = None, request_kwargs: dict | None = None, input_token_count: int | None = None, skip_inline_token_count: bool = False, ): """ Filter out model in model group, if: - model context window < message length. For azure openai models, requires 'base_model' is set. - https://docs.litellm.ai/docs/proxy/cost_tracking#spend-tracking-for-azure-openai-models - filter models above rpm limits - if region given, filter out models not in that region / unknown region - [TODO] function call and model doesn't support function calling """ verbose_router_logger.debug("Starting Pre-call checks for deployments in model=%s", model) # Optimized: Use list() shallow copy instead of deepcopy # We only pop from the list, not modify deployment dicts - 100x+ faster on hot path (every request) _returned_deployments = list(healthy_deployments) invalid_model_indices: Final = set() # Use set for O(1) membership checks # Token counting (tiktoken) is the dominant on-loop cost for large prompts. # Only count when a deployment actually declares max_input_tokens, and count # at most once; for model groups with no context-window limit it is skipped. # Async callers pass the count in, already computed off the event loop, and set # skip_inline_token_count so a failed off-loop count is not retried back on the loop. input_tokens: int | None = input_token_count _context_window_error = False _potential_error_str = "" _rate_limit_error = False parent_otel_span: Final = _get_parent_otel_span_from_kwargs(request_kwargs) has_countable_input: Final = (messages is not None or input is not None) and not compaction_pending( request_kwargs ) ## get model group RPM ## dt: Final = get_utc_datetime() current_minute: Final = dt.strftime("%H-%M") rpm_key: Final = f"{model}:rpm:{current_minute}" model_group_cache: Final = ( self.cache.get_cache(key=rpm_key, local_only=True, parent_otel_span=parent_otel_span) or {} ) # check the in-memory cache used by lowest_latency and usage-based routing. Only check the local cache. for idx, deployment in enumerate(_returned_deployments): # Cache nested dict access to avoid repeated temporary dict allocations _litellm_params = deployment.get("litellm_params", {}) _model_info = deployment.get("model_info", {}) # see if we have the info for this model _deployment_model = None # per-deployment model name (avoids overwriting the outer `model` group name) try: base_model = _model_info.get("base_model", None) if base_model is None: base_model = _litellm_params.get("base_model", None) _deployment_model = base_model or _litellm_params.get("model", None) model_info = self.get_router_model_info(deployment=deployment, received_model_name=model) max_input_tokens = model_info.get("max_input_tokens") if isinstance(model_info, dict) else None if isinstance(max_input_tokens, int) and has_countable_input: if input_tokens is None: if skip_inline_token_count: return _returned_deployments try: input_tokens = self._count_pre_call_check_tokens( messages=messages, input=input, request_kwargs=request_kwargs ) except Exception as e: verbose_router_logger.error( "litellm.router.py::_pre_call_checks: failed to count tokens. Returning initial list of deployments. Got - %s", e, ) return _returned_deployments if input_tokens > max_input_tokens: invalid_model_indices.add(idx) _context_window_error = True _potential_error_str += ( f"Model={_deployment_model}, Max Input Tokens={max_input_tokens}, Got={input_tokens}" ) continue except Exception as e: verbose_router_logger.exception("An error occurs - %s", e) model_id = _model_info.get("id", "") ## RPM CHECK ## ### get local router cache ### current_request_cache_local = ( self.cache.get_cache(key=model_id, local_only=True, parent_otel_span=parent_otel_span) or 0 ) ### get usage based cache ### if isinstance(model_group_cache, dict) and self.routing_strategy != "usage-based-routing-v2": model_group_cache[model_id] = model_group_cache.get(model_id, 0) current_request = max(current_request_cache_local, model_group_cache[model_id]) if isinstance(_litellm_params, dict) and _litellm_params.get("rpm", None) is not None: if isinstance(_litellm_params["rpm"], int) and _litellm_params["rpm"] <= current_request: invalid_model_indices.add(idx) _rate_limit_error = True continue ## REGION CHECK ## if request_kwargs is not None and request_kwargs.get("allowed_model_region") is not None: allowed_model_region = request_kwargs.get("allowed_model_region") if allowed_model_region is not None: if not is_region_allowed( litellm_params=LiteLLM_Params(**_litellm_params), allowed_model_region=allowed_model_region, ): invalid_model_indices.add(idx) continue ## INVALID PARAMS ## -> catch 'gpt-3.5-turbo-16k' not supporting 'response_format' param if request_kwargs is not None and litellm.drop_params is False: # get supported params — use per-deployment model to avoid overwriting the outer model group name _dep_model_for_params: str = _deployment_model or model try: ( _dep_model_for_params, custom_llm_provider, _, _, ) = litellm.get_llm_provider( model=_dep_model_for_params, litellm_params=LiteLLM_Params(**_litellm_params), ) except Exception as e: # noqa: BLE001 # best-effort filter: an unresolvable provider must not fail the request verbose_router_logger.debug( "litellm.router.py::_pre_call_checks: skipping supported-params check for model=%s. Got - %s", _dep_model_for_params, e, ) continue supported_openai_params = litellm.get_supported_openai_params( model=_dep_model_for_params, custom_llm_provider=custom_llm_provider, ) if supported_openai_params is None: continue else: # check the non-default openai params in request kwargs non_default_params = litellm.utils.get_non_default_params(passed_params=request_kwargs) special_params = ["response_format"] # check if all params are supported for k, v in non_default_params.items(): if k not in supported_openai_params and k in special_params: # if not -> invalid model verbose_router_logger.debug("INVALID MODEL INDEX @ REQUEST KWARG FILTERING, k=%s", k) invalid_model_indices.add(idx) if len(invalid_model_indices) == len(_returned_deployments): """ - no healthy deployments available b/c context window checks or rate limit error - First check for rate limit errors (if this is true, it means the model passed the context window check but failed the rate limit check) """ if _rate_limit_error is True: # allow generic fallback logic to take place raise RouterRateLimitErrorBasic( model=model, ) elif _context_window_error is True: raise litellm.ContextWindowExceededError( message=f"litellm._pre_call_checks: Context Window exceeded for given call. No models have context window large enough for this call.\n{_potential_error_str}", model=model, llm_provider="", ) if len(invalid_model_indices) > 0: # Single-pass filter using set for O(1) lookups (avoids O(n^2) from repeated pops) _returned_deployments = [d for i, d in enumerate(_returned_deployments) if i not in invalid_model_indices] return _returned_deployments def _get_model_from_alias(self, model: str) -> str | None: """ Get the model from the alias. Returns: - str, the litellm model name - None, if model is not in model group alias """ return resolve_model_group_alias(self.model_group_alias, model) def _get_deployment_by_litellm_model(self, model: str) -> list: """ Get the deployment by litellm model. """ return [m for m in self.model_list if m["litellm_params"]["model"] == model] def _try_early_resolve_deployments_for_model_not_in_names( self, model: str, request_team_id: str | None, include_team_models: bool = False, ) -> tuple[str, list | dict] | None: """ When ``model`` is not in ``self.model_names``, try team routes, pattern routes, team pattern routers, then default deployment. Returns None if none apply. """ if model in self.model_names: return None # Check for team-specific deployments by team_public_model_name. # This intentionally takes priority over team pattern routers below, # so that named team deployments shadow wildcard/pattern routes. if request_team_id is not None: team_deployments = self._get_all_deployments(model_name=model, team_id=request_team_id) if team_deployments: return model, team_deployments elif include_team_models: team_deployments = self._team_deployments_across_teams(model) if team_deployments: return model, team_deployments pattern_deployments = self.pattern_router.get_deployments_by_pattern( model=model, ) if pattern_deployments: return model, pattern_deployments if request_team_id is not None and request_team_id in self.team_pattern_routers: pattern_deployments = self.team_pattern_routers[request_team_id].get_deployments_by_pattern( model=model, ) if pattern_deployments: return model, pattern_deployments if self.default_deployment is not None: # Shallow copy with nested litellm_params copy (100x+ faster than deepcopy) updated_deployment: Final = self.default_deployment.copy() updated_deployment["litellm_params"] = self.default_deployment["litellm_params"].copy() updated_deployment["litellm_params"]["model"] = model return model, updated_deployment return None def _team_deployments_across_teams(self, model: str) -> list[DeploymentTypedDict]: """Every team's deployments under public name `model`, for a proxy admin calling without a team.""" team_deployments: Final = [ self.model_list[index] for (_, public_model_name), indices in self.team_model_to_deployment_indices.items() if public_model_name == model for index in indices ] team_ids: Final = { team_id for deployment in team_deployments for team_id in [(deployment.get("model_info") or {}).get("team_id")] if team_id is not None } if len(team_ids) > 1: raise litellm.BadRequestError( message=( f"Model name '{model}' matches deployments from multiple teams. " "Specify the deployment ID directly to disambiguate." ), model=model, llm_provider="", ) return team_deployments def deployments_for_request( self, model: str, request_kwargs: Mapping[str, object] ) -> Sequence[DeploymentTypedDict]: """The deployments `model` names for this caller, through the same alias, then team-first, then global, then admin-across-teams resolution `_common_checks_available_deployment` applies, so strategy selection and compression policy can never disagree with deployment selection about which marker a name means.""" registered_name: Final = self._get_model_from_alias(model=model) or model team_id: Final = get_request_team_id(request_kwargs) deployments: Final = self._get_all_deployments(model_name=registered_name, team_id=team_id) if deployments or team_id is not None or not _is_proxy_admin_request(request_kwargs): return deployments return self._team_deployments_across_teams(registered_name) @staticmethod def _is_strategy_marker_deployment(deployment: Mapping[str, object]) -> bool: litellm_params: Final = deployment.get("litellm_params") if not isinstance(litellm_params, Mapping): return False deployment_model: Final = litellm_params.get("model") return isinstance(deployment_model, str) and classify_strategy_router_model(deployment_model) is not None def _common_checks_available_deployment( self, model: str, messages: list[dict[str, str]] | None = None, input: str | list | None = None, specific_deployment: bool | None = False, request_kwargs: dict | None = None, ) -> tuple[str, list | dict]: """ Common checks for 'get_available_deployment' across sync + async call. If 'healthy_deployments' returned is None, this means the user chose a specific deployment Returns - str, the litellm model name - List, if multiple models chosen - Dict, if specific model chosen """ request_team_id: Final = get_request_team_id(request_kwargs) # check if aliases set on litellm model alias map if specific_deployment is True: return model, self._drop_strategy_markers( model, self._filter_reserved_deployments( model=model, healthy_deployments=self._get_deployment_by_litellm_model(model=model), request_team_id=request_team_id, ), ) elif model not in self.model_names and self.has_model_id(model): deployment: Final = self.get_deployment(model_id=model) if deployment is not None: deployment_model: Final = deployment.litellm_params.model return deployment_model, cast( # cast-ok: contract requires a plain dict for a single deployment dict, self._filter_reserved_deployments( model=deployment_model, healthy_deployments=( cast( # cast-ok: model_dump of a router deployment DeploymentTypedDict, deployment.model_dump(exclude_none=True) ), ), request_team_id=request_team_id, )[0], ) raise ValueError( f"LiteLLM Router: Trying to call specific deployment, but Model ID :{model} does not exist in Model ID map" ) _model_from_alias: Final = self._get_model_from_alias(model=model) if _model_from_alias is not None: model = _model_from_alias _routing_group_deployments: Final = self._get_routing_group_deployments(model=model, team_id=request_team_id) if _routing_group_deployments is None: early: Final = self._try_early_resolve_deployments_for_model_not_in_names( model=model, request_team_id=request_team_id, include_team_models=_is_proxy_admin_request(request_kwargs), ) if early is not None: if not isinstance(early[1], list): return early[0], cast( # cast-ok: contract requires a plain dict for a single deployment dict, self._filter_reserved_deployments( model=early[0], healthy_deployments=( cast( # cast-ok: early resolve returns a router deployment DeploymentTypedDict, early[1] ), ), request_team_id=request_team_id, )[0], ) return early[0], self._drop_strategy_markers( early[0], self._filter_reserved_deployments( model=early[0], healthy_deployments=early[1], request_team_id=request_team_id, ), ) ## get healthy deployments ### get all deployments healthy_deployments = ( _routing_group_deployments if _routing_group_deployments is not None else self._get_all_deployments(model_name=model, team_id=request_team_id) ) _pre_model_access_group_filter_len: Final = len(healthy_deployments) healthy_deployments = self._filter_reserved_deployments( model=model, healthy_deployments=self._filter_deployments_by_model_access_groups( model=model, healthy_deployments=healthy_deployments, request_kwargs=request_kwargs, request_team_id=request_team_id, ), request_team_id=request_team_id, ) _access_group_filter_emptied_candidates = ( _pre_model_access_group_filter_len > 0 and len(healthy_deployments) == 0 ) if len(healthy_deployments) == 0: # check if the user sent in a deployment name instead # Do not fall back when access-group filtering removed every candidate; # _get_deployment_by_litellm_model does not re-apply that filter. if _pre_model_access_group_filter_len == 0: _litellm_model_deployments: Final = self._get_deployment_by_litellm_model(model=model) healthy_deployments = self._filter_reserved_deployments( model=model, healthy_deployments=self._filter_deployments_by_model_access_groups( model=model, healthy_deployments=_litellm_model_deployments, request_kwargs=request_kwargs, request_team_id=request_team_id, ), request_team_id=request_team_id, ) # If the litellm-model lookup produced candidates that access-group # filtering then removed, treat this the same as the by-name path # being emptied: prevent default-model fallback from bypassing the # restriction (the fallback model may have no access_groups and # would short-circuit the filter). if len(_litellm_model_deployments) > 0 and len(healthy_deployments) == 0: _access_group_filter_emptied_candidates = True if verbose_router_logger.isEnabledFor(logging.DEBUG): verbose_router_logger.debug("initial list of deployments: %s", healthy_deployments) if len(healthy_deployments) == 0: # Check for default fallbacks if no deployments are found for the requested model # Do not fall back to another model when access-group filtering removed every # candidate for the requested name: re-filtering the fallback model can be a # no-op when it has no access_groups, incorrectly serving a different model. if self._has_default_fallbacks() and not _access_group_filter_emptied_candidates: fallback_model: Final = self._get_first_default_fallback() if fallback_model: verbose_router_logger.info( "Model '%s' not found. Attempting to use default fallback model '%s'.", model, fallback_model ) # Re-assign model to the fallback and try to get deployments again model = fallback_model healthy_deployments = self._get_all_deployments(model_name=model, team_id=request_team_id) healthy_deployments = self._filter_reserved_deployments( model=model, healthy_deployments=self._filter_deployments_by_model_access_groups( model=model, healthy_deployments=healthy_deployments, request_kwargs=request_kwargs, request_team_id=request_team_id, ), request_team_id=request_team_id, ) # If still no deployments after checking for fallbacks, raise an error if len(healthy_deployments) == 0: raise litellm.BadRequestError( message=f"You passed in model={model}. {RouterErrors.no_healthy_deployments.value}", model=model, llm_provider="", ) if litellm.model_alias_map and model in litellm.model_alias_map: model = litellm.model_alias_map[ model ] # update the model to the actual value if an alias has been passed in return model, self._drop_strategy_markers(model, healthy_deployments) def _drop_strategy_markers( self, model: str, deployments: Sequence[DeploymentTypedDict] ) -> list[DeploymentTypedDict]: """A strategy marker is never a callable deployment, whichever resolution arm produced it.""" selectable: Final = [ # mutable-ok: matches _common_checks_available_deployment's list contract d for d in deployments if not self._is_strategy_marker_deployment(d) ] if deployments and not selectable: raise litellm.BadRequestError( message=f"You passed in model={model}. {RouterErrors.only_strategy_marker_deployments.value}", model=model, llm_provider="", ) return selectable def _filter_reserved_deployments( self, model: str, healthy_deployments: Sequence[DeploymentTypedDict], request_team_id: str | None, ) -> tuple[DeploymentTypedDict, ...]: result: Final = filter_reserved_deployments( self._drop_strategy_markers(model, healthy_deployments), request_team_id, now=datetime.now(timezone.utc) ) if result.blocking_window is not None and len(result.deployments) == 0: raise litellm.BadRequestError( message=f"Deployment {model} is reserved for another team until " f"{result.blocking_window.end:%H:%M} {result.blocking_window.timezone}", model=model, llm_provider="", ) return result.deployments def _filter_deployments_by_model_access_groups( self, model: str, healthy_deployments: list, request_kwargs: dict | None, request_team_id: str | None, ) -> list: """ Restrict candidate deployments to caller-authorized model access groups. This is only applied when: - request metadata includes `user_api_key_auth`, and - caller permissions for this model are access-group-only (no explicit model, wildcard, or all-proxy grants). """ if not healthy_deployments or request_kwargs is None: return healthy_deployments metadata: Final = request_kwargs.get("metadata") or {} litellm_metadata: Final = request_kwargs.get("litellm_metadata") or {} user_api_key_auth: Final = metadata.get("user_api_key_auth") or litellm_metadata.get("user_api_key_auth") if user_api_key_auth is None: return healthy_deployments object_models: Final = set(getattr(user_api_key_auth, "models", []) or []) object_team_models: Final = set(getattr(user_api_key_auth, "team_models", []) or []) allowed_models: Final = object_models | object_team_models if not allowed_models: return healthy_deployments # If caller has direct model/wildcard/all-proxy access, do not constrain # deployment choice by access group. if model in allowed_models or "*" in allowed_models or "all-proxy-models" in allowed_models: return healthy_deployments access_groups_for_model: Final = self.get_model_access_groups(model_name=model, team_id=request_team_id) if len(access_groups_for_model) == 0: return healthy_deployments allowed_access_groups: Final = set(access_groups_for_model.keys()) & allowed_models if not allowed_access_groups: # No overlap means this request was not authorized via model access # group membership for this model, so do not force group filtering. return healthy_deployments filtered_deployments: Final = [] for deployment in healthy_deployments: deployment_model_info = deployment.get("model_info") or {} deployment_access_groups = set(deployment_model_info.get("access_groups", []) or []) if deployment_access_groups & allowed_access_groups: filtered_deployments.append(deployment) return filtered_deployments async def async_get_healthy_deployments( self, model: str, request_kwargs: dict, messages: list[dict[str, str]] | None = None, input: str | list | None = None, specific_deployment: bool | None = False, parent_otel_span: Span | None = None, health_check_probe: bool = False, ) -> list[dict] | dict: """ Get the healthy deployments for a model. Returns: - List[Dict], if multiple models chosen *OR* - Dict, if specific model chosen """ model, healthy_deployments = self._common_checks_available_deployment( model=model, messages=messages, input=input, specific_deployment=specific_deployment, request_kwargs=request_kwargs, ) # IF TEAM ID SPECIFIED ON MODEL, AND REQUEST CONTAINS USER_API_KEY_TEAM_ID, FILTER OUT MODELS THAT ARE NOT IN THE TEAM ## THIS PREVENTS WRITING FILES OF OTHER TEAMS TO MODELS THAT ARE TEAM-ONLY MODELS healthy_deployments = filter_team_based_models( healthy_deployments=healthy_deployments, request_kwargs=request_kwargs, ) if verbose_router_logger.isEnabledFor(logging.DEBUG): verbose_router_logger.debug("healthy_deployments after team filter: %s", healthy_deployments) healthy_deployments = filter_web_search_deployments( healthy_deployments=healthy_deployments, request_kwargs=request_kwargs, ) if verbose_router_logger.isEnabledFor(logging.DEBUG): verbose_router_logger.debug("healthy_deployments after web search filter: %s", healthy_deployments) if isinstance(healthy_deployments, dict): if (healthy_deployments.get("model_info") or {}).get("blocked") is True: raise litellm.ServiceUnavailableError( message=f"Model '{model}' is currently paused and cannot accept requests.", model=model, llm_provider="", ) return healthy_deployments # Health-check-based filtering (before cooldown) healthy_deployments = await self._async_filter_health_check_unhealthy_deployments( healthy_deployments=healthy_deployments, parent_otel_span=parent_otel_span, health_check_probe=health_check_probe, ) routing_read_batch: Final = RoutingReadBatch.active() cooldown_deployments: Final = ( await _async_get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=parent_otel_span) if routing_read_batch is None else await routing_read_batch.async_get_cooldown_deployments( litellm_router_instance=self, healthy_deployments=healthy_deployments, parent_otel_span=parent_otel_span, ) ) if verbose_router_logger.isEnabledFor(logging.DEBUG): verbose_router_logger.debug("cooldown deployments: %s", cooldown_deployments) _pre_cooldown_deployments: Final = healthy_deployments healthy_deployments = self._filter_cooldown_deployments( healthy_deployments=healthy_deployments, cooldown_deployments=cooldown_deployments, ) # Safety net: only bypass cooldown filter when health-check routing is # driving cooldown (i.e. allowed_fails_policy is set). Without a policy, # cooldowns are from real request failures and must not be bypassed. if not healthy_deployments and self.enable_health_check_routing and self.allowed_fails_policy is not None: verbose_router_logger.warning( "All deployments in cooldown via health-check routing, bypassing cooldown filter" ) healthy_deployments = _pre_cooldown_deployments healthy_deployments = self._filter_blocked_deployments(healthy_deployments) healthy_deployments = await self.async_callback_filter_deployments( model=model, healthy_deployments=healthy_deployments, messages=(cast(list[AllMessageValues], messages) if messages is not None else None), request_kwargs=request_kwargs, parent_otel_span=parent_otel_span, ) if self.enable_pre_call_checks and (messages is not None or input is not None): deployments_to_check: Final = cast(list[dict], healthy_deployments) healthy_deployments = self._pre_call_checks( model=model, healthy_deployments=deployments_to_check, messages=messages, input=input, request_kwargs=request_kwargs, input_token_count=await self._acount_pre_call_check_tokens( model=model, healthy_deployments=deployments_to_check, messages=messages, input=input, request_kwargs=request_kwargs, ), skip_inline_token_count=True, ) # check if user wants to do tag based routing healthy_deployments = await get_deployments_for_tag( llm_router_instance=self, model=model, request_kwargs=request_kwargs, healthy_deployments=healthy_deployments, metadata_variable_name=self._get_metadata_variable_name_from_kwargs(request_kwargs), ) # narrow to whatever `self.routing_plugins` left in candidate_models healthy_deployments = self._filter_by_routing_plugin_candidates( healthy_deployments=healthy_deployments, request_kwargs=request_kwargs, ) ## ORDER FILTERING ## -> if user set 'order' in deployments, return deployments with lowest order (e.g. order=1 > order=2) _target_order: Final = (request_kwargs or {}).pop("_target_order", None) healthy_deployments = litellm.utils.get_order_filtered_deployments( cast(list[dict], healthy_deployments), target_order=_target_order ) ## WEIGHTED FAILOVER EXCLUSION ## -> drop deployments already tried in ## this request via weighted-failover. Always honored, regardless of the ## router-level flag, so a stale exclusion key on kwargs cannot escape. _excluded_deployment_ids: Final = (request_kwargs or {}).pop("_excluded_deployment_ids", None) healthy_deployments = litellm.utils.get_excluded_filtered_deployments( cast(list[dict], healthy_deployments), excluded_deployment_ids=_excluded_deployment_ids, ) ## RETRY SKIP ## -> drop deployments that already refused this request with a ## non-retryable status, unless that leaves nothing, so the caller still gets ## the provider's own error instead of a no-deployments error. _retry_skipped_deployment_ids: Final = _as_retry_skipped_deployment_ids( request_kwargs.pop("_retry_skipped_deployment_ids", None) if request_kwargs else None ) healthy_deployments = ( litellm.utils.get_excluded_filtered_deployments( healthy_deployments, excluded_deployment_ids=_retry_skipped_deployment_ids ) or healthy_deployments ) if len(healthy_deployments) == 0: exception: Final = await async_raise_no_deployment_exception( litellm_router_instance=self, model=model, parent_otel_span=parent_otel_span, ) raise exception return healthy_deployments @staticmethod def _pop_effort_from_nested_carrier(request_kwargs: dict[str, object], carrier: str) -> None: nested: Final = request_kwargs.get(carrier) if not isinstance(nested, Mapping): return sanitized: Final = {key: value for key, value in nested.items() if key != "effort"} if sanitized: request_kwargs[carrier] = sanitized # rebind-ok: copy-on-write, so a shared nested carrier is never edited else: request_kwargs.pop(carrier, None) @staticmethod def _tier_ceiling_under_the_surface_name( tier_litellm_params: Mapping[str, object], responses_call: bool ) -> Mapping[str, object]: """``max_tokens``, ``max_completion_tokens`` and ``max_output_tokens`` are one ceiling under three names, and each surface reads exactly one of them: the Responses bridge builds its internal ``max_tokens`` from ``max_output_tokens`` and would overwrite the tier's, chat and /v1/messages never read ``max_output_tokens``, and litellm already renames ``max_tokens`` to ``max_completion_tokens`` for the OpenAI models that require it. Collapse whatever the tier carries onto the surface's own name, preferring a value the operator already wrote under that name.""" surface_key: Final = "max_output_tokens" if responses_call else "max_tokens" carried: Final = tuple( key for key in (surface_key, "max_tokens", "max_completion_tokens", "max_output_tokens") if key in tier_litellm_params ) if not carried: return tier_litellm_params return MappingProxyType( { **{k: v for k, v in tier_litellm_params.items() if k not in OUTPUT_TOKEN_CEILING_PARAMS}, surface_key: tier_litellm_params[carried[0]], } ) def _pin_tier_params_onto_request( self, model: str, tier_litellm_params: Mapping[str, object] | None, request_kwargs: dict, responses_call: bool, ) -> bool: """Apply a routing strategy's per-tier litellm_params on top of the request and report whether they pinned an output ceiling, so the caller can hand the request its own ceiling back on a routing pass that pins none.""" if not tier_litellm_params: return False accepted_tier_params: Final = self._tier_params_the_target_accepts(model, tier_litellm_params, request_kwargs) surface_tier_params: Final = self._tier_ceiling_under_the_surface_name( accepted_tier_params, responses_call=responses_call ) self._drop_client_carriers_a_tier_pin_supersedes(request_kwargs, surface_tier_params) request_kwargs.update(surface_tier_params) return not OUTPUT_TOKEN_CEILING_PARAMS.isdisjoint(surface_tier_params) @staticmethod def _restore_client_ceiling_no_tier_pins(request_kwargs: MutableMapping[str, object]) -> None: """A model-group fallback re-enters routing with the kwargs an earlier auto-router pass already rewrote, so a ceiling sized for that pass's tier would ride onto a group no tier chose. When this pass pins none, hand the request back exactly the carriers the caller sent, which the first pinning pass stamped. The stamp lives in a metadata bucket a caller can also write, so the proxy strips the key at ingestion and this read takes nothing but the three ceiling carriers as integers: no other key ever reaches kwargs.""" stamped: Final = next( ( bucket.get(CLIENT_OUTPUT_CEILING_METADATA_KEY) for bucket in (request_kwargs.get("metadata"), request_kwargs.get("litellm_metadata")) if isinstance(bucket, dict) and CLIENT_OUTPUT_CEILING_METADATA_KEY in bucket ), None, ) if not isinstance(stamped, dict): return callers_ceiling: Final = MappingProxyType( { carrier: cap for carrier, value in stamped.items() if carrier in OUTPUT_TOKEN_CEILING_PARAMS and (cap := as_output_cap(value)) is not None } ) for carrier in OUTPUT_TOKEN_CEILING_PARAMS: request_kwargs.pop(carrier, None) request_kwargs.update(callers_ceiling) @staticmethod def _drop_client_carriers_a_tier_pin_supersedes( request_kwargs: dict[str, object], tier_litellm_params: Mapping[str, object], ) -> None: """Tier litellm_params are deliberate operator overrides, but provider translations let a caller-supplied carrier of the same setting (``thinking``, ``output_config.effort``, ``reasoning.effort``) outrank the ``reasoning_effort`` alias, so a pinned effort only reaches the wire if the client's other encodings are removed before the merge. Non-effort fields a carrier also holds (``output_config.format``, ``reasoning.summary``) are kept. An output ceiling has the same shape: ``max_tokens``, ``max_completion_tokens`` and ``max_output_tokens`` are one setting under three names, and a provider handed two of them either rejects the request or picks one by iteration order.""" if not OUTPUT_TOKEN_CEILING_PARAMS.isdisjoint(tier_litellm_params): _, metadata_bucket = get_or_create_metadata_bucket(request_kwargs) metadata_bucket.setdefault( CLIENT_OUTPUT_CEILING_METADATA_KEY, { carrier: request_kwargs[carrier] for carrier in OUTPUT_TOKEN_CEILING_PARAMS if carrier in request_kwargs }, ) for carrier in OUTPUT_TOKEN_CEILING_PARAMS: request_kwargs.pop(carrier, None) if "reasoning_effort" not in tier_litellm_params: return request_kwargs.pop("thinking", None) Router._pop_effort_from_nested_carrier(request_kwargs, "output_config") Router._pop_effort_from_nested_carrier(request_kwargs, "reasoning") async def async_get_available_deployment( self, model: str, request_kwargs: dict, messages: list[dict[str, str]] | None = None, input: str | list | None = None, specific_deployment: bool | None = False, ): """ Async implementation of 'get_available_deployments'. Allows all cache calls to be made async => 10x perf impact (8rps -> 100 rps). """ if ( self.routing_strategy != "usage-based-routing-v2" and self.routing_strategy != "simple-shuffle" and self.routing_strategy != "cost-based-routing" and self.routing_strategy != "latency-based-routing" and self.routing_strategy != "least-busy" ): # prevent regressions for other routing strategies, that don't have async get available deployments implemented. return self.get_available_deployment( model=model, messages=messages, input=input, specific_deployment=specific_deployment, request_kwargs=request_kwargs, ) try: parent_otel_span: Final = _get_parent_otel_span_from_kwargs(request_kwargs) ######################################################### # Execute Pre-Routing Hooks # this hook can modify the model, messages before the routing decision is made ######################################################### responses_call: Final = input is not None and messages is None pre_routing_hook_response: Final = await self.async_pre_routing_hook( model=model, request_kwargs=request_kwargs, messages=messages, input=input, specific_deployment=specific_deployment, ) if pre_routing_hook_response is not None: model = pre_routing_hook_response.model messages = pre_routing_hook_response.messages record_pre_routing_selection(request_kwargs, model) tier_pins_ceiling: Final = self._pin_tier_params_onto_request( model=model, tier_litellm_params=pre_routing_hook_response.litellm_params if pre_routing_hook_response else None, request_kwargs=request_kwargs, responses_call=responses_call, ) if not tier_pins_ceiling: self._restore_client_ceiling_no_tier_pins(request_kwargs) ######################################################### # Resolve the strategy and logger AFTER the pre-routing hook, since # the hook can replace `model` and routing-group lookup must key # off the final model name. strategy, strategy_selector = self._get_routing_context(model, request_kwargs) routing_read_batch: Final = RoutingReadBatch.for_strategy(strategy, strategy_selector) with RoutingReadBatch.scoped(routing_read_batch): healthy_deployments: Final = await self.async_get_healthy_deployments( model=model, request_kwargs=request_kwargs, messages=messages, input=input, specific_deployment=specific_deployment, parent_otel_span=parent_otel_span, ) if isinstance(healthy_deployments, dict): await self._async_override_selector_pre_call_check( strategy, strategy_selector, healthy_deployments, parent_otel_span ) return healthy_deployments # When encrypted content affinity pins to a specific deployment, if request_kwargs.get("_encrypted_content_affinity_pinned") and len(healthy_deployments) == 1: await self._async_override_selector_pre_call_check( strategy, strategy_selector, healthy_deployments[0], parent_otel_span ) return healthy_deployments[0] start_time: Final = time.time() if strategy == "simple-shuffle": return simple_shuffle( resolve_model_alias=self._get_model_from_alias, healthy_deployments=healthy_deployments, model=model, request_kwargs=request_kwargs, ) with PrefetchedUsage.scoped( routing_read_batch.prefetched_usage if routing_read_batch is not None else None ): deployment: Final = await self._select_deployment_async( strategy=strategy, selector=strategy_selector, model=model, healthy_deployments=healthy_deployments, messages=messages, input=input, request_kwargs=request_kwargs, ) if deployment is None: exception: Final = await async_raise_no_deployment_exception( litellm_router_instance=self, model=model, parent_otel_span=parent_otel_span, ) raise exception await self._async_override_selector_pre_call_check( strategy, strategy_selector, deployment, parent_otel_span ) verbose_router_logger.info( "get_available_deployment for model: %s, Selected deployment: %s for model: %s", model, self.print_deployment(deployment), model, ) end_time: Final = time.time() _duration: Final = end_time - start_time asyncio.create_task( self.service_logger_obj.async_service_success_hook( service=ServiceTypes.ROUTER, duration=_duration, call_type=".async_get_available_deployments", parent_otel_span=parent_otel_span, start_time=start_time, end_time=end_time, ) ) return deployment except Exception as e: traceback_exception: Final = traceback.format_exc() # if router rejects call -> log to langfuse/otel/etc. if request_kwargs is not None: logging_obj: Final = request_kwargs.get("litellm_logging_obj", None) if logging_obj is not None: asyncio.create_task( logging_obj.dispatch_failure_handlers( exception=e, traceback_exception=traceback_exception, prefer_async_handlers=True, ) ) raise e async def async_get_available_deployment_for_pass_through( self, model: str, request_kwargs: dict, messages: list[dict[str, str]] | None = None, input: str | list | None = None, specific_deployment: bool | None = False, ): """ Async version of get_available_deployment_for_pass_through Only returns deployments configured with use_in_pass_through=True """ try: parent_otel_span: Final = _get_parent_otel_span_from_kwargs(request_kwargs) # 1. Execute pre-routing hook responses_call: Final = input is not None and messages is None pre_routing_hook_response: Final = await self.async_pre_routing_hook( model=model, request_kwargs=request_kwargs, messages=messages, input=input, specific_deployment=specific_deployment, ) if pre_routing_hook_response is not None: model = pre_routing_hook_response.model messages = pre_routing_hook_response.messages record_pre_routing_selection(request_kwargs, model) tier_pins_ceiling: Final = self._pin_tier_params_onto_request( model=model, tier_litellm_params=pre_routing_hook_response.litellm_params if pre_routing_hook_response else None, request_kwargs=request_kwargs, responses_call=responses_call, ) if not tier_pins_ceiling: self._restore_client_ceiling_no_tier_pins(request_kwargs) # 2. Get healthy deployments healthy_deployments: Final = await self.async_get_healthy_deployments( model=model, request_kwargs=request_kwargs, messages=messages, input=input, specific_deployment=specific_deployment, parent_otel_span=parent_otel_span, ) strategy, strategy_selector = self._get_routing_context(model, request_kwargs) # 3. If specific deployment returned, verify if it supports pass-through if isinstance(healthy_deployments, dict): if (healthy_deployments.get("model_info") or {}).get("blocked") is True: raise litellm.ServiceUnavailableError( message=f"Model '{model}' is currently paused and cannot accept requests.", model=model, llm_provider="", ) litellm_params: Final = healthy_deployments.get("litellm_params", {}) if litellm_params.get("use_in_pass_through"): await self._async_override_selector_pre_call_check( strategy, strategy_selector, healthy_deployments, parent_otel_span ) return healthy_deployments else: raise litellm.BadRequestError( message=f"Deployment {healthy_deployments.get('model_info', {}).get('id')} does not support pass-through endpoint (use_in_pass_through=False)", model=model, llm_provider="", ) # 4. Filter deployments that support pass-through pass_through_deployments = self._filter_pass_through_deployments(healthy_deployments=healthy_deployments) if len(pass_through_deployments) == 0: raise litellm.BadRequestError( message=f"Model {model} has no deployments configured with use_in_pass_through=True. Please add use_in_pass_through: true to the deployment configuration", model=model, llm_provider="", ) # 5. Apply load balancing strategy start_time: Final = time.perf_counter() if strategy == "simple-shuffle": return simple_shuffle( resolve_model_alias=self._get_model_from_alias, healthy_deployments=pass_through_deployments, model=model, request_kwargs=request_kwargs, ) deployment: Final = await self._select_deployment_async( strategy=strategy, selector=strategy_selector, model=model, healthy_deployments=pass_through_deployments, messages=messages, input=input, request_kwargs=request_kwargs, ) if deployment is None: exception: Final = await async_raise_no_deployment_exception( litellm_router_instance=self, model=model, parent_otel_span=parent_otel_span, ) raise exception await self._async_override_selector_pre_call_check( strategy, strategy_selector, deployment, parent_otel_span ) verbose_router_logger.info( "async_get_available_deployment_for_pass_through model: %s, selected deployment: %s", model, self.print_deployment(deployment), ) end_time: Final = time.perf_counter() _duration: Final = end_time - start_time asyncio.create_task( self.service_logger_obj.async_service_success_hook( service=ServiceTypes.ROUTER, duration=_duration, call_type=".async_get_available_deployments", parent_otel_span=parent_otel_span, start_time=start_time, end_time=end_time, ) ) return deployment except Exception as e: traceback_exception: Final = traceback.format_exc() if request_kwargs is not None: logging_obj: Final = request_kwargs.get("litellm_logging_obj", None) if logging_obj is not None: asyncio.create_task( logging_obj.dispatch_failure_handlers( exception=e, traceback_exception=traceback_exception, prefer_async_handlers=True, ) ) raise e async def _run_routing_plugins( self, model: str, request_kwargs: dict, messages: list[dict[str, object]] | None, ) -> RoutingContext: """ Build a RoutingContext for `model`, run it through `self.routing_plugins` in order, then stash the narrowed candidate list and accumulated signals onto `request_kwargs["metadata"]` so `_filter_by_routing_plugin_candidates` (called later, during healthy-deployment filtering) can consume them. """ from litellm.litellm_core_utils.prompt_templates.factory import ( resolve_structured_messages, ) candidate_models: Final = list(self.resolved_litellm_models(model)) metadata_key: Final = self._get_metadata_variable_name_from_kwargs(request_kwargs) metadata: Final = request_kwargs.setdefault(metadata_key, {}) context = RoutingContext( raw_messages=messages or [], structured_messages=resolve_structured_messages(messages=messages, request_kwargs=request_kwargs) or [], candidate_models=candidate_models, metadata=metadata, ) for plugin in self.routing_plugins: context = await plugin.run(context) metadata["routing_plugin_signals"] = context.signals if len(context.candidate_models) < len(candidate_models): metadata["_routing_plugin_candidate_models"] = context.candidate_models return context def _filter_by_routing_plugin_candidates( self, healthy_deployments: list[dict] | dict, request_kwargs: dict, ) -> list[dict] | dict: """ Narrow `healthy_deployments` to whatever `self.routing_plugins` left in `context.candidate_models`. Raises rather than silently falling back to the unfiltered pool -- a plugin narrowing to nothing is a policy decision (e.g. no model this tenant's budget allows), not something to bypass. """ if not self.routing_plugins or not isinstance(healthy_deployments, list): return healthy_deployments metadata_key: Final = self._get_metadata_variable_name_from_kwargs(request_kwargs) candidate_models: Final = (request_kwargs.get(metadata_key) or {}).get("_routing_plugin_candidate_models") # `is None` (not falsy-check): a plugin narrowing to an empty list must # still hit the "no deployments left" raise below, not be treated the # same as "no plugin ever set this key". if candidate_models is None: return healthy_deployments candidate_set: Final = set(candidate_models) filtered: Final = [d for d in healthy_deployments if d.get("litellm_params", {}).get("model") in candidate_set] if not filtered: raise ValueError(f"No deployments left after routing-plugin filtering. candidate_models={candidate_models}") return filtered def _select_pre_routing_strategy( self, model: str, request_kwargs: Mapping[str, object] ) -> "TaggedPreRoutingStrategy[PreRoutingStrategy] | None": """ Resolve the pre-routing strategy for `model`, disambiguating deployments that share a `model_name` by matching the request's tags against each registered strategy's tags before falling back to the first registered. Returns the tagged registry entry so the caller can tell whether the request's tags were what selected it, and can locate the marker deployment the strategy was registered from via its (model_name, tags) pair. The registries are keyed by the marker deployment's own `model_name`, which for a team-scoped router is the internal `model_name_{team}_{uuid}` while the caller sends the team's public name. So the names looked up are the `model_name`s of whatever deployments this caller's request resolves `model` to, and `model` itself when it resolves to none. With tag filtering enabled, router-wide or by the request's enable_tag_filtering (which the proxy sets from key/team router_settings), strategies that all carry real tags matching none of the request's do not capture it when the name also has plain deployments: returning None hands the request to ordinary tag-aware deployment selection. """ registries: Final = (self.auto_routers, self.complexity_routers, self.adaptive_routers, self.quality_routers) if not any(registries): return None deployments: Final = self.deployments_for_request(model, request_kwargs) registered_names: Final = tuple(dict.fromkeys(str(d["model_name"]) for d in deployments)) or (model,) candidates: Final = tuple( tagged for registry in registries for name in registered_names for tagged in registry.get(name, []) ) if not candidates: return None request_tags: Final = _get_tags_from_request_kwargs(request_kwargs) if request_tags: for tagged in candidates: if tagged.tags and is_valid_deployment_tag( list(tagged.tags), request_tags, self.tag_filtering_match_any ): return tagged for tagged in candidates: if "default" in tagged.tags: return tagged request_scoped_filtering: Final = request_kwargs.get("enable_tag_filtering") is True if ( (self.enable_tag_filtering or request_scoped_filtering) and all(tagged.tags for tagged in candidates) and any(not self._is_strategy_marker_deployment(d) for d in deployments) ): return None return candidates[0] @staticmethod def _request_header(request_kwargs: Mapping[str, object], header_name: str) -> str | None: proxy_server_request: Final = request_kwargs.get("proxy_server_request") if not isinstance(proxy_server_request, Mapping): return None headers: Final = proxy_server_request.get("headers") if not isinstance(headers, Mapping): return None return next( ( value for key, value in headers.items() if isinstance(key, str) and key.lower() == header_name and isinstance(value, str) ), None, ) def _claude_code_session_router_cache_key(self, request_kwargs: Mapping[str, object]) -> str | None: session_id: Final = self._request_header(request_kwargs, "x-claude-code-session-id") if session_id is None or _CLAUDE_CODE_SESSION_ID_RE.fullmatch(session_id) is None: return None metadata_name: Final = "litellm_metadata" if "litellm_metadata" in request_kwargs else "metadata" metadata: Final = request_kwargs.get(metadata_name) if not isinstance(metadata, Mapping): return None caller_scope: Final = metadata.get("user_api_key_hash") if not isinstance(caller_scope, str) or not caller_scope: return None return f"claude_code_session_router:v1:{caller_scope}:{session_id}" async def _delete_claude_code_session_router_binding(self, cache_key: str) -> None: try: await self._claude_code_session_router_cache.async_delete_cache(key=cache_key) except Exception as e: # noqa: BLE001 # cache cleanup must not fail an otherwise routable request verbose_router_logger.warning( "Failed to delete Claude Code session router binding; the binding may remain until its TTL expires: %s", e, ) async def _get_claude_code_session_router_binding(self, cache_key: str) -> object: session_cache: Final = self._claude_code_session_router_cache try: if session_cache.redis_cache is None: return await session_cache.async_get_cache(key=cache_key) return await session_cache.redis_cache.async_get_cache(key=cache_key) except Exception as e: # noqa: BLE001 # an optional binding must not make routing depend on Redis log_redis_failure( verbose_router_logger, logging.WARNING, "Failed to read Claude Code session router binding; using the requested model", e, ) return None async def _resolve_claude_code_session_router( self, model: str, registered_model_name: str, request_kwargs: Mapping[str, object], ) -> str: if is_native_compaction_call(): return registered_model_name if not any((self.auto_routers, self.complexity_routers, self.adaptive_routers, self.quality_routers)): return registered_model_name cache_key: Final = self._claude_code_session_router_cache_key(request_kwargs) if cache_key is None or not isinstance(request_kwargs, dict): return registered_model_name if request_kwargs.get("fallback_depth") not in (None, 0): return registered_model_name agent_id: Final = self._request_header(request_kwargs, "x-claude-code-agent-id") if agent_id is not None: bound_model: Final = await self._get_claude_code_session_router_binding(cache_key) if not isinstance(bound_model, str): return registered_model_name bound_registered_model: Final = self._get_model_from_alias(model=bound_model) or bound_model if self._select_pre_routing_strategy(bound_registered_model, request_kwargs) is None: await self._delete_claude_code_session_router_binding(cache_key) return registered_model_name await self._claude_code_session_router_cache.async_set_cache( key=cache_key, value=bound_model, ttl=_CLAUDE_CODE_SESSION_ROUTER_TTL_SECONDS, ) self._stamp_or_clear_metadata_key(request_kwargs, "model_group", bound_model) return bound_registered_model if self._select_pre_routing_strategy(registered_model_name, request_kwargs) is None: return registered_model_name await self._claude_code_session_router_cache.async_set_cache( key=cache_key, value=model, ttl=_CLAUDE_CODE_SESSION_ROUTER_TTL_SECONDS, ) return registered_model_name async def async_pre_routing_hook( self, model: str, request_kwargs: dict[str, object], messages: list[dict[str, Any]] | None = None, input: str | list | None = None, specific_deployment: bool | None = False, ) -> PreRoutingHookResponse | None: """ This hook is called before the routing decision is made. Used for the litellm auto-router to modify the request before the routing decision is made. `model` is whatever the caller asked for, which may be a `model_group_alias` key or a team's public model name, while the strategy registries and the marker deployment are keyed by the marker's own `model_name`, so every lookup below resolves the alias first and the team name through the deployment path. Only the lookups: the caller-facing name stays the alias, since spend metadata is stamped before routing and the response carries the tier group the strategy picked. """ requested_registered_model_name: Final = self._get_model_from_alias(model=model) or model registered_model_name: Final = await self._resolve_claude_code_session_router( model=model, registered_model_name=requested_registered_model_name, request_kwargs=request_kwargs, ) ######################################################### # Run the routing-plugin pipeline, if any plugins are configured. # Plugins narrow the candidate deployment pool (consumed later by # `_filter_by_routing_plugin_candidates`) and may attach signals for # downstream strategies (auto-router, complexity-router, ...) to read. ######################################################### if self.routing_plugins: await self._run_routing_plugins( model=registered_model_name, request_kwargs=request_kwargs, messages=messages ) selected_strategy: Final = self._select_pre_routing_strategy( model=registered_model_name, request_kwargs=request_kwargs ) if selected_strategy is None: await arm_compaction(request_kwargs, None) self._record_routing_decision(request_kwargs=request_kwargs, routing_decision=None) self._stamp_or_clear_metadata_key( request_kwargs=request_kwargs, key=SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY, value=None ) self._stamp_or_clear_metadata_key( request_kwargs=request_kwargs, key=CONSUMED_REQUEST_TAGS_METADATA_KEY, value=None ) return None from litellm.proxy.auth.auto_router_checks import authorize_member_auto_router_inference from litellm.router_strategy.complexity_router.complexity_router import ComplexityRouter reject_recursive_compactor(registered_model_name) await arm_compaction( request_kwargs, selected_strategy.strategy.config.context_compaction if isinstance(selected_strategy.strategy, ComplexityRouter) else None, tuple( dict.fromkeys( member for pool in selected_strategy.strategy.config.tiers.values() for member in ((pool,) if isinstance(pool, str) else pool) ) ) if isinstance(selected_strategy.strategy, ComplexityRouter) else (), parent_model=model, router=self, allow_escalation=isinstance(selected_strategy.strategy, ComplexityRouter) and selected_strategy.strategy.config.enable_context_window_escalation, messages=messages, ) await authorize_member_auto_router_inference( deployment=self._selected_strategy_marker_deployment( model=registered_model_name, strategy_tags=selected_strategy.tags, request_kwargs=request_kwargs, ), request_kwargs=request_kwargs, llm_router=self, ) from litellm.proxy.guardrails.auto_router_compression import ( messages_for_routing, model_hop_compression_armed, policy_for_model, ) # Same tag-aware lookup the proxy's pre-call arming used, so an alias with # several tag-scoped markers cannot suppress under one and route under another. compression_policy: Final = policy_for_model( llm_router=self, model_alias=registered_model_name, request_kwargs=request_kwargs, request_tags=_get_tags_from_request_kwargs(request_kwargs), ) # Shared compression already ran in the pre-call hook, so reuse it rather than # compressing twice. Conditional on arming having actually happened: only the # proxy arms, and on the SDK path the shortcut would skip both hops entirely. needs_independent_routing_compression: Final = compression_policy is not None and not ( compression_policy.is_same and compression_policy.model is not None and model_hop_compression_armed() ) routing_messages: Final = ( await messages_for_routing(policy=compression_policy, messages=messages, request_kwargs=request_kwargs) if needs_independent_routing_compression else None ) routed: Final = await selected_strategy.strategy.async_pre_routing_hook( model=registered_model_name, request_kwargs=request_kwargs, messages=routing_messages if routing_messages is not None else messages, input=input, specific_deployment=specific_deployment, ) # Routing-only compression must not leak into the response: the model call and # deployment-context filtering key off this field. Compared by value, since # pydantic rebuilds the list rather than keeping the object passed in. pre_routing_hook_response: Final = ( routed.model_copy(update={"messages": messages}) # mutable-ok: pydantic's model_copy takes a dict if routed is not None and routing_messages is not None and routed.messages == routing_messages else routed ) self._record_routing_decision( request_kwargs=request_kwargs, routing_decision=(pre_routing_hook_response.routing_decision if pre_routing_hook_response else None), ) self._stamp_or_clear_metadata_key( request_kwargs=request_kwargs, key=SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY, value=(pre_routing_hook_response.session_affinity_ttl_seconds if pre_routing_hook_response else None), ) self._stamp_or_clear_metadata_key( request_kwargs=request_kwargs, key=CONSUMED_REQUEST_TAGS_METADATA_KEY, value=self._consumed_request_tags_stamp( selected_strategy=selected_strategy, pre_routing_hook_response=pre_routing_hook_response, request_tags=_get_tags_from_request_kwargs(request_kwargs), ), ) # `model` (the alias, e.g. "smart-router") is never the deployment actually # called - apply the router marker's own litellm_params to the request, # since the tier/route deployment the hook selected won't have them. The # marker entry is looked up by its `auto_router/` model prefix and the # selected strategy's tags, never by list position: plain deployments may # share the alias `model_name` and must not leak their params (`api_base`, # `api_key`, ...) onto the routed call. Router-only fields (tpm, rpm, # weight, complexity_router_config, ...) are excluded from the actual # outbound LLM call downstream by litellm.types.utils.all_litellm_params, # not here. Custom pricing fields ARE call params, so they must be # excluded here: they price the alias, not the deployment the hook # selected, and forwarding them re-registers the routed deployment at # the alias's price (an explicit 0 makes every alias request bill $0). # Forwarded params only fill gaps: the keys inserted here ride along on the # request (top level, so sibling requests sharing a `metadata` dict never see # them) until `_update_kwargs_with_deployment` drops any the selected # deployment sets itself (its own `aws_region_name` beats the marker's). # Per-tier `litellm_params` on the hook response are deliberate overrides # the caller applies on top, so those keys are never forwarded here. marker_params: Final = ( self._forwardable_alias_marker_params( model=registered_model_name, strategy_tags=selected_strategy.tags, request_kwargs=request_kwargs ) if pre_routing_hook_response is not None else () ) tier_param_keys: Final = ( tuple(pre_routing_hook_response.litellm_params or ()) if pre_routing_hook_response is not None else () ) newly_forwarded: Final = tuple( (key, value) for key, value in marker_params if key not in request_kwargs and key not in tier_param_keys ) request_kwargs.pop(_ALIAS_MARKER_FORWARDED_PARAMS_KWARG, None) request_kwargs.update(newly_forwarded) if newly_forwarded: request_kwargs.update(((_ALIAS_MARKER_FORWARDED_PARAMS_KWARG, tuple(key for key, _ in newly_forwarded)),)) return pre_routing_hook_response def _selected_strategy_marker_deployment( self, model: str, strategy_tags: tuple[str, ...], request_kwargs: Mapping[str, object] ) -> DeploymentTypedDict | None: markers: Final = tuple( deployment for deployment in self.deployments_for_request(model, request_kwargs) if "model" in deployment["litellm_params"] and str(deployment["litellm_params"]["model"]).startswith(AUTO_ROUTER_MODEL_PREFIX) ) tag_matched: Final = tuple( deployment for deployment in markers if (tuple(deployment["litellm_params"]["tags"] or ()) if "tags" in deployment["litellm_params"] else ()) == strategy_tags ) return tag_matched[0] if tag_matched else (markers[0] if markers else None) def _forwardable_alias_marker_params( self, model: str, strategy_tags: tuple[str, ...], request_kwargs: Mapping[str, object] ) -> tuple[tuple[str, object], ...]: marker: Final = self._selected_strategy_marker_deployment( model=model, strategy_tags=strategy_tags, request_kwargs=request_kwargs ) if marker is None: return () return tuple( (key, value) for key, value in marker["litellm_params"].items() if key not in _ALIAS_PARAMS_NEVER_FORWARDED and key not in CustomPricingLiteLLMParams.model_fields and value is not None ) @staticmethod def _forwarded_alias_marker_keys_the_deployment_sets( deployment: Mapping[str, object], forwarded_keys: object ) -> tuple[str, ...]: deployment_litellm_params: Final = deployment.get("litellm_params") if not isinstance(deployment_litellm_params, Mapping) or not isinstance(forwarded_keys, tuple): return () return tuple( key for key in forwarded_keys if isinstance(key, str) and Router._deployment_sets_litellm_param(deployment_litellm_params, key) ) @staticmethod def _deployment_sets_litellm_param(deployment_litellm_params: Mapping[str, object], key: str) -> bool: value: Final = deployment_litellm_params.get(key) if value is None: return False field: Final = LiteLLM_Params.model_fields.get(key) return field is None or value != field.default def _consumed_request_tags_stamp( self, selected_strategy: "TaggedPreRoutingStrategy[PreRoutingStrategy]", pre_routing_hook_response: PreRoutingHookResponse | None, request_tags: Sequence[str], ) -> ConsumedRequestTagsStamp | None: """Record which tags picked the router and which model group it rewrote to, or None. A request whose tags matched the selected strategy's tags has already spent those tags on picking the router; re-applying them to the routed tier's model group would empty the pool unless every tier deployment repeats the marker's tag. Only the strategy's own tags are spent: the request's other tags keep constraining deployment selection inside the routed group, and key/team policy tags are untouched because tag filtering separately re-applies whatever `metadata.inherited_tags` carries for the stamped group. """ if pre_routing_hook_response is None or not selected_strategy.tags or not request_tags: return None if not is_valid_deployment_tag(selected_strategy.tags, request_tags, self.tag_filtering_match_any): return None return ConsumedRequestTagsStamp(model_group=pre_routing_hook_response.model, tags=selected_strategy.tags) @staticmethod def _record_routing_decision( request_kwargs: dict, routing_decision: StandardLoggingRoutingDecision | None, ) -> None: """Make the request's metadata describe THIS routing attempt, and only this one. Fallbacks re-enter the hook with the same `request_kwargs`, so an attempt that picks a plain model group after an auto-router group failed must clear the earlier decision; leaving it would attribute the first router's tier and cause to the deployment that actually served the request. Every attempt therefore writes or clears, never just writes. """ from litellm.types.router import BaselineRouteStamp baseline_model: Final = routing_decision.get("savings_baseline_model") if routing_decision else None baseline_id: Final = routing_decision.get("savings_baseline_deployment_id") if routing_decision else None router_name: Final = routing_decision.get("router_model_name") if routing_decision else None Router._stamp_or_clear_metadata_key( request_kwargs=request_kwargs, key="_autorouter_baseline_route", value=( BaselineRouteStamp(router_name, baseline_model, baseline_id) if router_name and baseline_model and baseline_id else None ), ) Router._stamp_or_clear_metadata_key( request_kwargs=request_kwargs, key="routing_decision", value=( None if routing_decision is None else Router._redact_prompt_text_if_needed( request_kwargs=request_kwargs, routing_decision=routing_decision ) ), ) @staticmethod def _stamp_or_clear_metadata_key(request_kwargs: dict, key: str, value: object | None) -> None: """Write a proxy-internal metadata key for THIS routing attempt, or clear it. Fallbacks and retries re-enter the pre-routing hook with the same `request_kwargs`, so every attempt must write or clear, never just write; a value left behind by an earlier attempt would be attributed to this one. `get_or_create_metadata_bucket` is the single owner of "which dict holds proxy-internal metadata": it picks `litellm_metadata` when present (so the value never lands in the `metadata` dict that routes like /v1/messages forward to the provider) and replaces a non-dict value rather than silently skipping the write. Clearing pops from BOTH buckets so a request whose bucket resolution changed between attempts cannot resurrect a stale value. """ if value is None: for bucket in (request_kwargs.get("metadata"), request_kwargs.get("litellm_metadata")): if isinstance(bucket, dict): bucket.pop(key, None) return _, metadata_bucket = get_or_create_metadata_bucket(request_kwargs) metadata_bucket[key] = value @staticmethod def _redact_prompt_text_if_needed( request_kwargs: Mapping[str, object], routing_decision: StandardLoggingRoutingDecision, ) -> StandardLoggingRoutingDecision: """Drop verbatim prompt text from the record when message logging is redacted. An operator who turns message logging off has said prompt content must not reach the logs, so the fields that quote the prompt (the matched keywords, and the signals that name them) are omitted. Derived values are kept, because a tier, a cause, a score or an escalation flag aggregates the prompt rather than reproducing any of it, and dropping them would leave the row unexplainable for no privacy gain. Applied here rather than in each strategy so a strategy added later cannot bypass it. """ from litellm.litellm_core_utils.redact_messages import ( should_redact_message_logging, ) if not should_redact_message_logging( { "litellm_params": request_kwargs, "standard_callback_dynamic_params": request_kwargs.get("standard_callback_dynamic_params"), } ): return routing_decision kept: Final = { field: value for field, value in routing_decision.items() if field not in PROMPT_QUOTING_ROUTING_DECISION_FIELDS } return cast(StandardLoggingRoutingDecision, kept) # cast-ok: dropping optional keys preserves the type def get_available_deployment( self, model: str, messages: list[dict[str, str]] | None = None, input: str | list | None = None, specific_deployment: bool | None = False, request_kwargs: dict | None = None, ): """ Returns the deployment based on routing strategy """ if self.routing_plugins: raise ValueError( "Router(plugins=[...]) is configured, but this call resolved to the synchronous " "deployment-selection path, which never runs the routing-plugin pipeline. This " "happens for sync Router methods (e.g. Router.completion()) and for async calls " "with a routing_strategy that has no async-native selector (e.g. legacy " "'usage-based-routing', v1). Silently skipping " "configured plugins would let a policy plugin (e.g. a deny-all rule) be bypassed. " "Use an async Router method with a supported routing_strategy (simple-shuffle, " "usage-based-routing-v2, cost-based-routing, latency-based-routing, least-busy), " "or remove `plugins` from the Router config." ) # users need to explicitly call a specific deployment, by setting `specific_deployment = True` as completion()/embedding() kwarg # When this was no explicit we had several issues with fallbacks timing out model, healthy_deployments = self._common_checks_available_deployment( model=model, messages=messages, input=input, specific_deployment=specific_deployment, request_kwargs=request_kwargs, ) strategy, strategy_selector = self._get_routing_context(model, request_kwargs) if isinstance(healthy_deployments, dict): if (healthy_deployments.get("model_info") or {}).get("blocked") is True: raise litellm.ServiceUnavailableError( message=f"Model '{model}' is currently paused and cannot accept requests.", model=model, llm_provider="", ) self._override_selector_pre_call_check(strategy, strategy_selector, healthy_deployments) return healthy_deployments parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(request_kwargs) # Health-check-based filtering (before cooldown) healthy_deployments = self._filter_health_check_unhealthy_deployments( healthy_deployments=healthy_deployments, parent_otel_span=parent_otel_span, ) cooldown_deployments: Final = _get_cooldown_deployments( litellm_router_instance=self, parent_otel_span=parent_otel_span ) _pre_cooldown_deployments: Final = healthy_deployments healthy_deployments = self._filter_cooldown_deployments( healthy_deployments=healthy_deployments, cooldown_deployments=cooldown_deployments, ) if not healthy_deployments and self.enable_health_check_routing and self.allowed_fails_policy is not None: verbose_router_logger.warning( "All deployments in cooldown via health-check routing, bypassing cooldown filter" ) healthy_deployments = _pre_cooldown_deployments healthy_deployments = self._filter_blocked_deployments(healthy_deployments) # filter pre-call checks if self.enable_pre_call_checks and (messages is not None or input is not None): healthy_deployments = self._pre_call_checks( model=model, healthy_deployments=healthy_deployments, messages=messages, input=input, request_kwargs=request_kwargs, ) ## ORDER FILTERING ## -> if user set 'order' in deployments, return deployments with lowest order (e.g. order=1 > order=2) _target_order: Final = (request_kwargs or {}).pop("_target_order", None) healthy_deployments = litellm.utils.get_order_filtered_deployments( healthy_deployments, target_order=_target_order ) ## WEIGHTED FAILOVER EXCLUSION ## -> drop deployments already tried in ## this request via weighted-failover. See async counterpart in ## async_get_healthy_deployments for details. _excluded_deployment_ids: Final = (request_kwargs or {}).pop("_excluded_deployment_ids", None) healthy_deployments = litellm.utils.get_excluded_filtered_deployments( healthy_deployments, excluded_deployment_ids=_excluded_deployment_ids, ) ## RETRY SKIP ## -> see async counterpart in async_get_healthy_deployments. _retry_skipped_deployment_ids: Final = _as_retry_skipped_deployment_ids( request_kwargs.pop("_retry_skipped_deployment_ids", None) if request_kwargs else None ) healthy_deployments = ( litellm.utils.get_excluded_filtered_deployments( healthy_deployments, excluded_deployment_ids=_retry_skipped_deployment_ids ) or healthy_deployments ) if len(healthy_deployments) == 0: model_ids = self.get_model_ids(model_name=model) _cooldown_time = self.cooldown_cache.get_min_cooldown( model_ids=model_ids, parent_otel_span=parent_otel_span ) _cooldown_list = _get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=parent_otel_span) raise RouterRateLimitError( model=model, cooldown_time=_cooldown_time, enable_pre_call_checks=self.enable_pre_call_checks, cooldown_list=_cooldown_list, model_ids=model_ids, ) if strategy == "simple-shuffle": # if users pass rpm or tpm, we do a random weighted pick - based on rpm/tpm ############## Check 'weight' param set for weighted pick ################# return simple_shuffle( resolve_model_alias=self._get_model_from_alias, healthy_deployments=healthy_deployments, model=model, request_kwargs=request_kwargs, ) deployment: Final = self._select_deployment_sync( strategy=strategy, selector=strategy_selector, model=model, healthy_deployments=healthy_deployments, messages=messages, input=input, request_kwargs=request_kwargs, ) if deployment is None: verbose_router_logger.info("get_available_deployment for model: %s, No deployment available", model) model_ids = self.get_model_ids(model_name=model) _cooldown_time = self.cooldown_cache.get_min_cooldown( model_ids=model_ids, parent_otel_span=parent_otel_span ) _cooldown_list = _get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=parent_otel_span) raise RouterRateLimitError( model=model, cooldown_time=_cooldown_time, enable_pre_call_checks=self.enable_pre_call_checks, cooldown_list=_cooldown_list, model_ids=model_ids, ) self._override_selector_pre_call_check(strategy, strategy_selector, deployment) verbose_router_logger.info( "get_available_deployment for model: %s, Selected deployment: %s for model: %s", model, self.print_deployment(deployment), model, ) return deployment def get_available_deployment_for_pass_through( self, model: str, messages: list[dict[str, str]] | None = None, input: str | list | None = None, specific_deployment: bool | None = False, request_kwargs: dict | None = None, ): """ Returns deployments available for pass-through endpoints (based on load balancing strategy) Similar to get_available_deployment, but only returns deployments with use_in_pass_through=True Args: model: Model name messages: Optional list of messages input: Optional input data specific_deployment: Whether to find a specific deployment request_kwargs: Optional request parameters Returns: Dict: Selected deployment configuration Raises: BadRequestError: If no deployment is configured with use_in_pass_through=True RouterRateLimitError: If no pass-through deployments are available """ # 1. Perform common checks to get healthy deployments list model, healthy_deployments = self._common_checks_available_deployment( model=model, messages=messages, input=input, specific_deployment=specific_deployment, request_kwargs=request_kwargs, ) strategy, strategy_selector = self._get_routing_context(model, request_kwargs) # 2. If the returned is a specific deployment (Dict), verify and return directly if isinstance(healthy_deployments, dict): if (healthy_deployments.get("model_info") or {}).get("blocked") is True: raise litellm.ServiceUnavailableError( message=f"Model '{model}' is currently paused and cannot accept requests.", model=model, llm_provider="", ) litellm_params: Final = healthy_deployments.get("litellm_params", {}) if litellm_params.get("use_in_pass_through"): self._override_selector_pre_call_check(strategy, strategy_selector, healthy_deployments) return healthy_deployments else: # Specific deployment does not support pass-through raise litellm.BadRequestError( message=f"Deployment {healthy_deployments.get('model_info', {}).get('id')} does not support pass-through endpoint (use_in_pass_through=False)", model=model, llm_provider="", ) # 3. Filter deployments that support pass-through pass_through_deployments = self._filter_pass_through_deployments(healthy_deployments=healthy_deployments) if len(pass_through_deployments) == 0: # No deployments support pass-through raise litellm.BadRequestError( message=f"Model {model} has no deployment configured with use_in_pass_through=True. Please add use_in_pass_through: true in the deployment configuration", model=model, llm_provider="", ) pass_through_model_ids: Final = tuple( deployment["model_info"]["id"] for deployment in pass_through_deployments if "id" in deployment.get("model_info", {}) ) # 4. Apply health-check and cooldown filtering parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(request_kwargs) pass_through_deployments = self._filter_health_check_unhealthy_deployments( healthy_deployments=pass_through_deployments, parent_otel_span=parent_otel_span, ) cooldown_deployments: Final = _get_cooldown_deployments( litellm_router_instance=self, parent_otel_span=parent_otel_span ) pass_through_deployments = self._filter_cooldown_deployments( healthy_deployments=pass_through_deployments, cooldown_deployments=cooldown_deployments, ) pass_through_deployments = self._filter_blocked_deployments(pass_through_deployments) # 5. Apply pre-call checks (if enabled) if self.enable_pre_call_checks and (messages is not None or input is not None): pass_through_deployments = self._pre_call_checks( model=model, healthy_deployments=pass_through_deployments, messages=messages, input=input, request_kwargs=request_kwargs, ) if self._is_priority_routing_group(model): pass_through_deployments = litellm.utils.get_order_filtered_deployments( pass_through_deployments, target_order=request_kwargs.pop("_target_order", None) if request_kwargs is not None else None, ) if len(pass_through_deployments) == 0: model_ids = self.get_model_ids(model_name=model) _cooldown_time = self.cooldown_cache.get_min_cooldown( model_ids=model_ids, parent_otel_span=parent_otel_span ) _cooldown_list = _get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=parent_otel_span) raise RouterRateLimitError( model=model, cooldown_time=_cooldown_time, enable_pre_call_checks=self.enable_pre_call_checks, cooldown_list=_cooldown_list, model_ids=pass_through_model_ids, ) # 6. Apply load balancing strategy if strategy == "simple-shuffle": return simple_shuffle( resolve_model_alias=self._get_model_from_alias, healthy_deployments=pass_through_deployments, model=model, request_kwargs=request_kwargs, ) deployment: Final = self._select_deployment_sync( strategy=strategy, selector=strategy_selector, model=model, healthy_deployments=pass_through_deployments, messages=messages, input=input, request_kwargs=request_kwargs, ) if deployment is None: verbose_router_logger.info( "get_available_deployment_for_pass_through model: %s, no available deployments", model ) model_ids = self.get_model_ids(model_name=model) _cooldown_time = self.cooldown_cache.get_min_cooldown( model_ids=model_ids, parent_otel_span=parent_otel_span ) _cooldown_list = _get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=parent_otel_span) raise RouterRateLimitError( model=model, cooldown_time=_cooldown_time, enable_pre_call_checks=self.enable_pre_call_checks, cooldown_list=_cooldown_list, model_ids=model_ids, ) self._override_selector_pre_call_check(strategy, strategy_selector, deployment) verbose_router_logger.info( "get_available_deployment_for_pass_through model: %s, selected deployment: %s", model, self.print_deployment(deployment), ) return deployment def _filter_cooldown_deployments( self, healthy_deployments: list[dict], cooldown_deployments: list[str] ) -> list[dict]: """ Filters out the deployments currently cooling down from the list of healthy deployments Args: healthy_deployments: List of healthy deployments cooldown_deployments: List of model_ids cooling down. cooldown_deployments is a list of model_id's cooling down, cooldown_deployments = ["16700539-b3cd-42f4-b426-6a12a1bb706a", "16700539-b3cd-42f4-b426-7899"] Returns: List of healthy deployments """ if verbose_router_logger.isEnabledFor(logging.DEBUG): verbose_router_logger.debug("cooldown deployments: %s", cooldown_deployments) # Convert to set for O(1) lookup and use list comprehension for O(n) filtering cooldown_set: Final = set(cooldown_deployments) return [deployment for deployment in healthy_deployments if deployment["model_info"]["id"] not in cooldown_set] def _filter_blocked_deployments(self, healthy_deployments: list[dict]) -> list[dict]: """ Filters out deployments that an admin has paused via `LiteLLM_ProxyModelTable.blocked`. Applied alongside the cooldown filter on every routing entry point that calls `_common_checks_available_deployment` directly — the primary sync/async path, the sync pass-through path, and the retry / health-check helpers — so paused deployments never serve a request. The async pass-through path inherits this filter through its delegation to `async_get_healthy_deployments`. """ return [ deployment for deployment in healthy_deployments if (deployment.get("model_info") or {}).get("blocked") is not True ] @staticmethod def _is_deployment_blocked(deployment: "Deployment") -> bool: """ Returns True when a `Deployment` Pydantic instance carries the admin-paused flag. Used by credential-lookup helpers so passthrough file / batch endpoints cannot bypass the pause by resolving credentials directly. """ model_info: Final[object | None] = getattr(deployment, "model_info", None) if model_info is None: return False return getattr(model_info, "blocked", None) is True async def _async_filter_health_check_unhealthy_deployments( self, healthy_deployments: list[dict], parent_otel_span: Span | None = None, health_check_probe: bool = False, ) -> list[dict]: """ Filter out deployments marked unhealthy by background health checks. No-op when enable_health_check_routing is False. When background_health_check_model_groups is set, only deployments in the listed model groups are filtered; every other group keeps its configured routing strategy untouched, and a router-level allowed_fails_policy no longer disables the filter for the listed groups. Returns all deployments if health state is unavailable, stale, or would exclude every candidate (safety net). """ if not self.enable_health_check_routing: return healthy_deployments # When allowed_fails_policy is set, cooldown is the sole routing exclusion # mechanism -- skip the binary health check filter so the policy threshold # is respected before any deployment is excluded. With a model-group # allowlist the filter is already scoped, so listed groups keep it. scoped_groups: Final = self.background_health_check_model_groups if self.allowed_fails_policy is not None and scoped_groups is None: return healthy_deployments unhealthy_ids: Final = await self.health_state_cache.async_get_unhealthy_deployment_ids( parent_otel_span=parent_otel_span ) if not unhealthy_ids: return healthy_deployments filtered: Final = [ d for d in healthy_deployments if d["model_info"]["id"] not in unhealthy_ids or (scoped_groups is not None and d["model_name"] not in scoped_groups) ] if not filtered: return [] if health_check_probe else healthy_deployments # mutable-ok: empty list signals unavailable probe return filtered def _filter_health_check_unhealthy_deployments( self, healthy_deployments: list[dict], parent_otel_span: Span | None = None, ) -> list[dict]: """Sync version of _async_filter_health_check_unhealthy_deployments.""" if not self.enable_health_check_routing: return healthy_deployments scoped_groups: Final = self.background_health_check_model_groups if self.allowed_fails_policy is not None and scoped_groups is None: return healthy_deployments unhealthy_ids: Final = self.health_state_cache.get_unhealthy_deployment_ids(parent_otel_span=parent_otel_span) if not unhealthy_ids: return healthy_deployments filtered: Final = [ d for d in healthy_deployments if d["model_info"]["id"] not in unhealthy_ids or (scoped_groups is not None and d["model_name"] not in scoped_groups) ] if not filtered: verbose_router_logger.warning("All deployments marked unhealthy by health checks, bypassing health filter") return healthy_deployments return filtered def _filter_pass_through_deployments(self, healthy_deployments: list[dict]) -> list[dict]: """ Filter out deployments configured with use_in_pass_through=True Args: healthy_deployments: List of healthy deployments Returns: List[Dict]: Only includes a list of deployments that support pass-through """ verbose_router_logger.debug( "Filter pass-through deployments from %s healthy deployments", len(healthy_deployments) ) pass_through_deployments: Final = [ deployment for deployment in healthy_deployments if deployment.get("litellm_params", {}).get("use_in_pass_through", False) ] verbose_router_logger.debug("Found %s deployments with pass-through enabled", len(pass_through_deployments)) return pass_through_deployments def _track_deployment_metrics(self, deployment, parent_otel_span: Span | None, response=None): """ Tracks successful requests rpm usage. """ try: model_id: Final = deployment.get("model_info", {}).get("id", None) if response is None: # update self.deployment_stats if model_id is not None: self._update_usage(model_id, parent_otel_span) # update in-memory cache for tracking except Exception as e: verbose_router_logger.error("Error in _track_deployment_metrics: %s", e) def get_num_retries_from_retry_policy(self, exception: Exception, model_group: str | None = None): return _get_num_retries_from_retry_policy( exception=exception, model_group=model_group, model_group_retry_policy=self.model_group_retry_policy, retry_policy=self.retry_policy, ) def get_allowed_fails_from_policy(self, exception: Exception): """ BadRequestErrorRetries: Optional[int] = None AuthenticationErrorRetries: Optional[int] = None TimeoutErrorRetries: Optional[int] = None RateLimitErrorRetries: Optional[int] = None ContentPolicyViolationErrorRetries: Optional[int] = None """ # if we can find the exception then in the retry policy -> return the number of retries allowed_fails_policy: Final[AllowedFailsPolicy | None] = self.allowed_fails_policy if allowed_fails_policy is None: return None if ( isinstance(exception, litellm.AuthenticationError) and allowed_fails_policy.AuthenticationErrorAllowedFails is not None ): return allowed_fails_policy.AuthenticationErrorAllowedFails if isinstance(exception, litellm.Timeout) and allowed_fails_policy.TimeoutErrorAllowedFails is not None: return allowed_fails_policy.TimeoutErrorAllowedFails if ( isinstance(exception, litellm.RateLimitError) and allowed_fails_policy.RateLimitErrorAllowedFails is not None ): return allowed_fails_policy.RateLimitErrorAllowedFails if ( isinstance(exception, litellm.ContentPolicyViolationError) and allowed_fails_policy.ContentPolicyViolationErrorAllowedFails is not None ): return allowed_fails_policy.ContentPolicyViolationErrorAllowedFails if ( isinstance(exception, litellm.BadRequestError) and allowed_fails_policy.BadRequestErrorAllowedFails is not None ): return allowed_fails_policy.BadRequestErrorAllowedFails if ( isinstance(exception, litellm.InternalServerError) and allowed_fails_policy.InternalServerErrorAllowedFails is not None ): return allowed_fails_policy.InternalServerErrorAllowedFails if ( isinstance(exception, litellm.ServiceUnavailableError) and allowed_fails_policy.ServiceUnavailableErrorAllowedFails is not None ): return allowed_fails_policy.ServiceUnavailableErrorAllowedFails if ( isinstance(exception, litellm.BadGatewayError) and allowed_fails_policy.BadGatewayErrorAllowedFails is not None ): return allowed_fails_policy.BadGatewayErrorAllowedFails if isinstance(exception, litellm.NotFoundError) and allowed_fails_policy.NotFoundErrorAllowedFails is not None: return allowed_fails_policy.NotFoundErrorAllowedFails def _initialize_alerting(self): from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting if self.alerting_config is None: return router_alerting_config: Final[AlertingConfig] = self.alerting_config _slack_alerting_logger: Final = SlackAlerting( alerting_threshold=router_alerting_config.alerting_threshold, alerting=["slack"], default_webhook_url=router_alerting_config.webhook_url, ) self.slack_alerting_logger = _slack_alerting_logger litellm.logging_callback_manager.add_litellm_callback(_slack_alerting_logger) litellm.logging_callback_manager.add_litellm_success_callback( _slack_alerting_logger.response_taking_too_long_callback ) verbose_router_logger.info("\033[94m\nInitialized Alerting for litellm.Router\033[0m\n") def set_custom_routing_strategy(self, CustomRoutingStrategy: CustomRoutingStrategyBase): """ Sets get_available_deployment and async_get_available_deployment on an instanced of litellm.Router Use this to set your custom routing strategy Args: CustomRoutingStrategy: litellm.router.CustomRoutingStrategyBase """ setattr( self, "get_available_deployment", CustomRoutingStrategy.get_available_deployment, ) setattr( self, "async_get_available_deployment", CustomRoutingStrategy.async_get_available_deployment, ) def _reset_custom_routing_strategy(self) -> None: for attr in ("get_available_deployment", "async_get_available_deployment"): if attr in self.__dict__: delattr(self, attr) def flush_cache(self): litellm.cache = None self.cache.flush_cache() session_in_memory_cache: Final = self._claude_code_session_router_cache.in_memory_cache if session_in_memory_cache is not None: session_in_memory_cache.flush_cache() def reset(self): ## clean up on close litellm.success_callback = [] litellm._async_success_callback = [] litellm.failure_callback = [] litellm._async_failure_callback = [] self.retry_policy = None self.flush_cache()