refactor(types): declare litellm-owned kwargs as typed objects and derive the lists from their fields (#42843)

* refactor(types): declare litellm-owned params in one registry

* refactor(types): re-export registry constants without redundant aliases

* refactor(types): satisfy type-discipline rules in registry projections and tests

* refactor(types): classify every registry entry and check groups against typed config models

* test(types): pin load-bearing names and exact projections in registry tests

* style(types): keep agentic projection comment within ruff format

* refactor(types): declare litellm-owned params as typed objects and derive the lists from their fields

* refactor(types): fields of the typed objects become the registry; tests use a hand-written inventory

* refactor(types): split traversal into wire_names and owned_wire_names, move rust to kwarg artifacts

rust is a module-level switch (litellm.rust) that nothing reads from a call's kwargs, so it
joins self, use_client and model_config as a registered artifact instead of a DispatchOptions
field. The field constants now import from litellm.types.litellm_params directly instead of
through a re-export in litellm.types.utils. metadata and litellm_metadata are MutableMapping
because their readers mutate them in place, and client accepts raw httpx clients

* refactor(types): own max_agentic_loops as an option and walk only nested leaves

Move max_agentic_loops from AgenticLoopState to a new AgenticLoopOptions leaf under
LiteLLMOptions, since the interception handlers read it as a deployment ceiling rather
than stamping it. Drop the owned_wire_names fallback that treated an unresolved annotation
as a direct field, which under postponed annotations silently shrank the registry. Re-export
TRUSTED_CALLBACK_VARS_FIELD and ADDRESSED_RESPONSE_ID_FIELD from types.utils so that import
path keeps working. Tests use hand-written inventories for the callback and pricing names

* refactor(types): move data_residency to call state and drop aliased re-exports

data_residency is stamped by get_litellm_params and responses.main during the
call, so it lives on CallState, not CostOptions. mock_response also accepts a
float sequence, which main.py reads for mock embeddings. The types/utils.py
re-exports become one plain import with an exact F401 suppression instead of
two X as X aliases that pushed PLC0414 over its strict-gate ceiling. Redundant
leaf docstrings and the structural artifact test are gone; the re-exported
FIELD constants are checked by identity instead

* refactor(types): project owned kwarg names once and keep pass-through extraction in request order

* refactor(types): type caching_groups from its cache reader and hoist the pass-through ownership set

caching_groups is a sequence of flat model-group sequences, which is what
Cache._get_caching_group iterates. A regression test drives the public
cache key path so two groups in one caching group share a key and a third
does not. The pass-through endpoint builds its frozenset of owned names
once at import instead of per request, reads the two metadata carriers
from the extracted mapping instead of popping them, and its extraction
mappings are read-only. Concatenation tests assert the whole derived list
and tuple, docstrings drop reader claims that nothing in the module backs

* refactor(types): read owned names live in pass-through and pin tests to literal inventories

The pass-through endpoint checks body keys against the public all_litellm_params
list at request time again, as the base does, instead of a frozenset taken at
import, so a name registered after import is still extracted. A test drives
that path with a name added after import, and another sends both metadata
carriers interleaved with provider keys and asserts the whole merged result.

retry_policy accepts the mapping form its router reader builds a RetryPolicy
from. The pricing inventory in the typed tests is a literal tuple checked
against the model's fields, the agentic compatibility test asserts type, length
and set instead of declaration order, and the typed-model overlap tests assert
the exact intersection.

* refactor(types): move model_alias_map to CallState and read the owned registry in registry order in pass-through

* refactor(types): drop restating docstrings, keep FIELD importers on types.utils, pin pass-through registry order

* fix(types): satisfy strict lint for public FIELD re-exports

* fix(types): restore clean parameter re-exports

* fix(tests): compare pass-through extraction order to registry body keys

* refactor(types): type owned request parameter leaves

* refactor(types): share routing strategy literal and tighten leaf tests

* fix(proxy): drop client-supplied proxy-stamped names from pass-through litellm_params

* refactor(proxy): name pass-through litellm key split for what it holds

* refactor(types): drop TODO markers on the kept readerless fields

* fix(types): keep deployment tag_regex and max_file_size_mb out of provider requests

* fix(types): include every routing strategy the router accepts

---------

Co-authored-by: shrey kharbanda <shreshth@berri.ai>
This commit is contained in:
devin-ai-integration[bot] 2026-09-24 20:18:41 -07:00 • committed by GitHub
parent f3cf1cdfef
commit f61b3c3f38
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 1236 additions and 233 deletions

View file

@ -126,6 +126,7 @@ from litellm.types.completion import (
_CompletionDispatchContext,
_CompletionDispatchResult,
)
from litellm.types.litellm_params import RetryStrategy
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import (
CustomPricingLiteLLMParams,
@ -6026,9 +6027,7 @@ def completion_with_retries(*args, **kwargs):
# reset retries in .completion()
kwargs["max_retries"] = 0
kwargs["num_retries"] = 0
retry_strategy: Final[Literal["exponential_backoff_retry", "constant_retry"]] = kwargs.pop(
"retry_strategy", "constant_retry"
)
retry_strategy: Final[RetryStrategy] = kwargs.pop("retry_strategy", "constant_retry")
original_function: Final = kwargs.pop("original_function", completion)
if retry_strategy == "exponential_backoff_retry":
retryer = tenacity.Retrying(
@ -6054,7 +6053,7 @@ async def acompletion_with_retries(*args, **kwargs):
num_retries: Final = kwargs.pop("num_retries", 3)
kwargs["max_retries"] = 0
kwargs["num_retries"] = 0
retry_strategy: Final = kwargs.pop("retry_strategy", "constant_retry")
retry_strategy: Final[RetryStrategy] = kwargs.pop("retry_strategy", "constant_retry")
original_function: Final = kwargs.pop("original_function", completion)
if retry_strategy == "exponential_backoff_retry":
retryer = tenacity.AsyncRetrying(
@ -6082,9 +6081,7 @@ def responses_with_retries(*args, **kwargs):
# reset retries in .responses()
kwargs["max_retries"] = 0
kwargs["num_retries"] = 0
retry_strategy: Final[Literal["exponential_backoff_retry", "constant_retry"]] = kwargs.pop(
"retry_strategy", "constant_retry"
)
retry_strategy: Final[RetryStrategy] = kwargs.pop("retry_strategy", "constant_retry")
original_function: Final = kwargs.pop("original_function", responses)
if retry_strategy == "exponential_backoff_retry":
retryer = tenacity.Retrying(
@ -6111,7 +6108,7 @@ async def aresponses_with_retries(*args, **kwargs):
num_retries: Final = kwargs.pop("num_retries", 3)
kwargs["max_retries"] = 0
kwargs["num_retries"] = 0
retry_strategy: Final = kwargs.pop("retry_strategy", "constant_retry")
retry_strategy: Final[RetryStrategy] = kwargs.pop("retry_strategy", "constant_retry")
original_function: Final = kwargs.pop("original_function", aresponses)
if retry_strategy == "exponential_backoff_retry":
retryer = tenacity.AsyncRetrying(

View file

@ -105,6 +105,8 @@ from litellm.proxy.route_llm_request import ProxyModelNotFoundError
from litellm.proxy.utils import normalize_route_for_root_path
from litellm.repositories.team_repository import TeamRepository
from litellm.secret_managers.main import get_secret_str
from litellm.types import utils as types_utils
from litellm.types.litellm_params import ProxyRequestState, wire_names
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
@ -133,6 +135,9 @@ router: Final = APIRouter()
pass_through_endpoint_logging: Final = PassThroughEndpointLogging()
_METADATA_KEYS: Final = frozenset(("litellm_metadata", "metadata"))
_KEPT_OUT_OF_LITELLM_PARAMS: Final = _METADATA_KEYS | frozenset(wire_names(ProxyRequestState))
# Global registry to track registered pass-through routes and prevent memory leaks
_registered_pass_through_routes: Final[dict[str, dict[str, str | bool | list[str] | Mapping[str, object]]]] = {}
@ -578,21 +583,21 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
"""
Filter out litellm params from the request body
"""
from litellm.types.utils import all_litellm_params
_parsed_body = _parsed_body or {}
litellm_params_in_body: Final = {}
for k in all_litellm_params:
if k in _parsed_body:
litellm_params_in_body[k] = _parsed_body.pop(k, None)
litellm_keys_in_body: Final = MappingProxyType(
{k: _parsed_body.pop(k) for k in types_utils.all_litellm_params if k in _parsed_body}
)
litellm_params_in_body: Final = MappingProxyType(
{k: v for k, v in litellm_keys_in_body.items() if k not in _KEPT_OUT_OF_LITELLM_PARAMS}
)
_metadata = dict(
LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_dict)
)
litellm_metadata: Final = litellm_params_in_body.pop("litellm_metadata", None)
metadata: Final = litellm_params_in_body.pop("metadata", None)
litellm_metadata: Final = litellm_keys_in_body.get("litellm_metadata")
metadata: Final = litellm_keys_in_body.get("metadata")
if litellm_metadata:
_metadata.update(litellm_metadata)
if metadata:

View file

@ -259,6 +259,7 @@ from litellm.router_utils.routing_groups import (
validate_routing_strategy,
)
from litellm.scheduler import FlowItem, Scheduler
from litellm.types.litellm_params import RoutingStrategyName
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionToolParam,
@ -796,15 +797,7 @@ class Router:
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: Literal[
"simple-shuffle",
"least-busy",
"usage-based-routing",
"latency-based-routing",
"cost-based-routing",
"usage-based-routing-v2",
"lar1",
] = "simple-shuffle",
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,

View file

@ -3,8 +3,10 @@ from typing import Any, Final
from pydantic import BaseModel, Field
CHAT_COMPLETION_AGENTIC_SURFACE: Final = "chat_completions"
RESPONSES_AGENTIC_SURFACE: Final = "responses"
from litellm.types.litellm_params import AgenticSurface
CHAT_COMPLETION_AGENTIC_SURFACE: Final[AgenticSurface] = "chat_completions"
RESPONSES_AGENTIC_SURFACE: Final[AgenticSurface] = "responses"
CODE_INTERPRETER_INTERCEPTION_PREFIX: Final = "_code_interpreter_interception"
HEADROOM_INTERCEPTION_PREFIX: Final = "_headroom_interception"
HEADROOM_CONVERTED_STREAM_KEY: Final = f"{HEADROOM_INTERCEPTION_PREFIX}_converted_stream"

View file

@ -0,0 +1,364 @@
"""LiteLLM-owned request kwargs declared as typed fields; types/utils.py splices these with the callback and pricing
models and KWARG_ARTIFACTS into all_litellm_params."""
from collections.abc import Callable, Iterator, Mapping, MutableMapping, Sequence
from dataclasses import dataclass, field, fields, is_dataclass
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Literal, TypeAlias
if TYPE_CHECKING:
import httpx
from aiohttp import ClientSession
from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.router_strategy.complexity_router.context_compaction import CompactionState
from litellm.router_utils.fallback_event_handlers import AttemptedFallbackTargets
from litellm.types.caching import DynamicCacheControl
from litellm.types.llms.openai import ChatCompletionAssistantMessage, ChatCompletionUserMessage
from litellm.types.proxy.litellm_pre_call_utils import SecretFields
from litellm.types.router import ConfigurableClientsideParamsCustomAuth, DeploymentTypedDict, RetryPolicy
from litellm.types.router_weights import RouterWeights
from litellm.types.utils import ModelResponse, ModelResponseStream, ProviderSpecificHeader
ProviderClient: TypeAlias = (
OpenAI
| AsyncOpenAI
| AzureOpenAI
| AsyncAzureOpenAI
| HTTPHandler
| AsyncHTTPHandler
| httpx.Client
| httpx.AsyncClient
)
MockResponse: TypeAlias = (
str | Exception | Mapping[str, object] | Sequence[float] | ModelResponse | ModelResponseStream
)
RetryStrategy: TypeAlias = Literal["constant_retry", "exponential_backoff_retry"]
AgenticSurface: TypeAlias = Literal["chat_completions", "responses"]
RoutingStrategyName: TypeAlias = Literal[
"simple-shuffle",
"least-busy",
"usage-based-routing",
"latency-based-routing",
"cost-based-routing",
"usage-based-routing-v2",
"lar1",
]
TRUSTED_CALLBACK_VARS_FIELD: Final = "litellm_trusted_callback_vars"
ADDRESSED_RESPONSE_ID_FIELD: Final = "_litellm_addressed_response_id"
WIRE_NAME: Final = "wire_name"
def wire(name: str) -> Mapping[str, str]:
return MappingProxyType({WIRE_NAME: name})
@dataclass(frozen=True, slots=True, kw_only=True)
class ProviderConnection:
api_key: str | None = None
api_base: str | None = None
api_version: str | None = None
region_name: str | None = None
headers: Mapping[str, str] | None = None
provider_specific_header: "ProviderSpecificHeader | Sequence[ProviderSpecificHeader] | None" = None
client: "ProviderClient | None" = None
shared_session: "ClientSession | None" = None
ssl_verify: bool | str | None = None
request_timeout: float | None = None
force_timeout: float | None = None
stream_timeout: float | str | None = None
max_retries: int | None = None
tenant_id: str | None = None
client_id: str | None = None
client_secret: str | None = None
azure_username: str | None = None
azure_password: str | None = None
azure_scope: str | None = None
azure_ad_token_provider: Callable[[], str] | None = None
litellm_credential_name: str | None = None
configurable_clientside_auth_params: "Sequence[str | ConfigurableClientsideParamsCustomAuth] | None" = None
use_xai_oauth: bool | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class BedrockBatchConnection:
# Bedrock rejects these names in request bodies, so register them as LiteLLM-owned
aws_batch_role_arn: str | None = None
s3_bucket_name: str | None = None
s3_region_name: str | None = None
s3_endpoint_url: str | None = None
s3_output_bucket_name: str | None = None
s3_bucket_owner: str | None = None
s3_access_key_id: str | None = None
s3_secret_access_key: str | None = None
s3_encryption_key_id: str | None = None
bedrock_tags: Sequence[Mapping[str, str]] | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class ConnectionSettings:
provider: ProviderConnection
bedrock_batch: BedrockBatchConnection
@dataclass(frozen=True, slots=True, kw_only=True)
class DispatchOptions:
custom_llm_provider: str | None = None
azure: bool | None = None
use_litellm_proxy: bool | None = None
use_chat_completions_api: bool | None = None
use_in_pass_through: bool | None = None
allowed_openai_params: Sequence[str] | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class RoutingOptions:
fallbacks: Sequence[str | Mapping[str, object]] | None = None
context_window_fallback_dict: Mapping[str, str] | None = None
num_retries: int | None = None
retry_policy: "RetryPolicy | Mapping[str, object] | None" = None
retry_strategy: RetryStrategy | None = None
routing_strategy: RoutingStrategyName | None = None
cooldown_time: float | None = None
allowed_model_region: str | None = None
enable_tag_filtering: bool | None = None
fastest_response: bool | None = None
provider_affinity_header: str | None = None
search_tool_name: str | None = None
model_list: "Sequence[DeploymentTypedDict] | None" = None
@dataclass(frozen=True, slots=True, kw_only=True)
class DeploymentOptions:
model_info: Mapping[str, object] | None = None
rpm: int | None = None
tpm: int | None = None
itpm: int | None = None
otpm: int | None = None
default_api_key_rpm_limit: int | None = None
default_api_key_tpm_limit: int | None = None
max_parallel_requests: int | None = None
weight: int | None = None
order: int | None = None
tag_regex: Sequence[str] | None = None
max_file_size_mb: float | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class SpecializedRouterOptions:
auto_router_config_path: str | None = None
auto_router_config: str | None = None
auto_router_default_model: str | None = None
auto_router_embedding_model: str | None = None
auto_router_max_input_chars: int | None = None
auto_router_routing_compression: str | None = None
auto_router_model_compression: str | None = None
complexity_router_config: Mapping[str, object] | None = None
complexity_router_default_model: str | None = None
adaptive_router_config: Mapping[str, object] | None = None
adaptive_router_default_model: str | None = None
quality_router_config: Mapping[str, object] | None = None
quality_router_default_model: str | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class CachingOptions:
caching: bool | None = None
cache: "DynamicCacheControl | None" = None
ttl: float | None = None
enable_prompt_caching: bool | None = None
caching_groups: Sequence[Sequence[str]] | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class CostOptions:
cost_per_query: float | None = None
base_model: str | None = None
max_budget: float | None = None
budget_duration: str | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class ObservabilityOptions:
id: str | None = None
metadata: MutableMapping[str, object] | None = None # mutable-ok: the router and logging write keys into it
litellm_metadata: MutableMapping[str, object] | None = None # mutable-ok: the proxy writes keys into it
tags: Sequence[str] | None = None
litellm_trace_id: str | None = None
litellm_session_id: str | None = None
litellm_request_debug: bool | None = None
logger_fn: Callable[[Mapping[str, object]], None] | None = None
verbose: bool | None = None
no_log: bool | None = field(default=None, metadata=wire("no-log"))
@dataclass(frozen=True, slots=True, kw_only=True)
class AgenticLoopOptions:
max_agentic_loops: int | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class GuardrailOptions:
guardrails: Sequence[str] | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class PromptOptions:
prompt_id: str | None = None
prompt_variables: Mapping[str, object] | None = None
prompt_version: str | None = None
prompt_environment: str | None = None
prompt_label: str | None = None
litellm_system_prompt: str | None = None
custom_prompt_dict: Mapping[str, object] | None = None
roles: Mapping[str, object] | None = None
final_prompt_value: str | None = None
bos_token: str | None = None
eos_token: str | None = None
hf_model_name: str | None = None
supports_system_message: bool | None = None
ensure_alternating_roles: bool | None = None
user_continue_message: "ChatCompletionUserMessage | None" = None
assistant_continue_message: "ChatCompletionAssistantMessage | None" = None
disable_add_transform_inline_image_block: bool | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class ResponseOptions:
merge_reasoning_content_in_choices: bool | None = None
enable_json_schema_validation: bool | None = None
complete_response: bool | None = None
stream_chunk_size: int | None = None
keepalive_seconds: float | None = None
allow_client_keepalive_override: bool | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class MockOptions:
mock_response: "MockResponse | None" = None
mock_timeout: bool | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class LiteLLMOptions:
dispatch: DispatchOptions
routing: RoutingOptions
deployment: DeploymentOptions
specialized_routers: SpecializedRouterOptions
caching: CachingOptions
cost: CostOptions
observability: ObservabilityOptions
agentic_loop: AgenticLoopOptions
guardrails: GuardrailOptions
prompt: PromptOptions
response: ResponseOptions
mock: MockOptions
@dataclass(frozen=True, slots=True, kw_only=True)
class CallState:
litellm_call_id: str | None = None
completion_call_id: str | None = None
model_alias_map: Mapping[str, str] | None = None
data_residency: str | None = None
litellm_logging_obj: "Logging | None" = None
preset_cache_key: str | None = None
cache_key: str | None = None
stream_response: "Mapping[str, ModelResponse] | None" = None
context_compaction_state: "CompactionState | None" = field(default=None, metadata=wire("_context_compaction_state"))
@dataclass(frozen=True, slots=True, kw_only=True)
class AgenticLoopState:
depth: int | None = field(default=None, metadata=wire("_agentic_loop_depth"))
fingerprints: Sequence[str] | None = field(default=None, metadata=wire("_agentic_loop_fingerprints"))
api_surface: Literal["chat_completions", "responses"] | None = field(
default=None, metadata=wire("_agentic_loop_api_surface")
)
code_interpreter_active: bool | None = field(default=None, metadata=wire("_code_interpreter_interception_active"))
code_interpreter_sandbox_key: str | None = field(
default=None, metadata=wire("_code_interpreter_interception_sandbox_key")
)
code_interpreter_session_scoped: bool | None = field(
default=None, metadata=wire("_code_interpreter_interception_session_scoped")
)
code_interpreter_converted_stream: bool | None = field(
default=None, metadata=wire("_code_interpreter_interception_converted_stream")
)
websearch_emit_native_blocks: bool | None = field(
default=None, metadata=wire("_websearch_interception_emit_native_blocks")
)
websearch_converted_stream: bool | None = field(
default=None, metadata=wire("_websearch_interception_converted_stream")
)
headroom_converted_stream: bool | None = field(
default=None, metadata=wire("_headroom_interception_converted_stream")
)
@dataclass(frozen=True, slots=True, kw_only=True)
class RouterState:
weights: "RouterWeights | None" = field(default=None, metadata=wire("_router_weights"))
fallback_depth: int | None = None
max_fallbacks: int | None = None
attempted_targets: "AttemptedFallbackTargets | None" = None
@dataclass(frozen=True, slots=True, kw_only=True)
class ProxyRequestState:
proxy_server_request: Mapping[str, object] | None = None
secret_fields: "SecretFields | None" = None
trusted_callback_vars: Mapping[str, str] | None = field(default=None, metadata=wire(TRUSTED_CALLBACK_VARS_FIELD))
addressed_response_id: str | None = field(default=None, metadata=wire(ADDRESSED_RESPONSE_ID_FIELD))
strip_stream_usage: bool | None = field(default=None, metadata=wire("_litellm_strip_stream_usage"))
client_side_timeout: bool | None = None
model_file_id_mapping: Mapping[str, Mapping[str, str]] | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class EntrypointState:
acompletion: bool | None = None
aembedding: bool | None = None
aimg_generation: bool | None = None
atext_completion: bool | None = None
text_completion: bool | None = None
allm_passthrough_route: bool | None = None
async_call: bool | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class InternalState:
call: CallState
agentic_loop: AgenticLoopState
router: RouterState
proxy: ProxyRequestState
entrypoint: EntrypointState
KWARG_ARTIFACTS: Final[tuple[str, ...]] = ("self", "use_client", "model_config", "rust")
LITELLM_OWNED_ROOTS: Final = (ConnectionSettings, LiteLLMOptions, InternalState)
def wire_names(owner: type) -> tuple[str, ...]:
return tuple(owned.metadata.get(WIRE_NAME, owned.name) for owned in fields(owner))
def owned_wire_names(root: type) -> tuple[str, ...]:
def names() -> Iterator[str]:
for leaf in fields(root):
if not is_dataclass(leaf.type):
raise TypeError(f"{root.__name__}.{leaf.name} is not a dataclass leaf")
yield from wire_names(leaf.type) # pyright: ignore[reportArgumentType] # Field.type admits str
return tuple(names())
OWNED_KWARG_NAMES: Final = tuple(name for root in LITELLM_OWNED_ROOTS for name in owned_wire_names(root))
AGENTIC_LOOP_KWARG_NAMES: Final = (*wire_names(AgenticLoopState), *wire_names(AgenticLoopOptions))
BEDROCK_BATCH_KWARG_NAMES: Final = wire_names(BedrockBatchConnection)

View file

@ -24,6 +24,7 @@ if TYPE_CHECKING:
from .completion import CompletionRequest
from .embedding import EmbeddingRequest
from .litellm_params import RoutingStrategyName
from .llms.bedrock import AwsSessionTag
from .llms.openai import OpenAIFileObject
from .search import SearchProvider
@ -104,12 +105,7 @@ class RouterConfig(BaseModel):
context_window_fallbacks: list | None = []
model_group_alias: dict[str, list[str]] | None = {}
retry_after: int | None = 0
routing_strategy: Literal[
"simple-shuffle",
"least-busy",
"usage-based-routing",
"latency-based-routing",
] = "simple-shuffle"
routing_strategy: RoutingStrategyName = "simple-shuffle"
routing_groups: list[RoutingGroup] | None = None
model_config = ConfigDict(protected_namespaces=())

View file

@ -56,8 +56,15 @@ from litellm.types.llms.base import (
from litellm.types.mcp import MCPServerCostInfo
from ..litellm_core_utils.core_helpers import map_finish_reason, process_response_headers
from . import litellm_params as _litellm_params
from .agents import LiteLLMSendMessageResponse
from .guardrails import GuardrailEventHooks
from .litellm_params import (
AGENTIC_LOOP_KWARG_NAMES,
BEDROCK_BATCH_KWARG_NAMES,
KWARG_ARTIFACTS,
OWNED_KWARG_NAMES,
)
from .llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse
from .llms.base import HiddenParams
from .llms.openai import (
@ -3901,205 +3908,20 @@ def pricing_override_fields(*sources: Mapping[str, object]) -> tuple[str, ...]:
)
# Server-controlled fields that bound or drive an interceptor's agentic loop
# (depth, cycle fingerprints, ceiling, code-interpreter sandbox state). Listed
# in all_litellm_params so they are treated as LiteLLM-level and excluded from
# get_non_default_completion_params; otherwise the OpenAI param builder sweeps
# any unrecognized top-level key into extra_body and leaks them to the provider.
# This is what lets the loop carry state across rerun calls without a provider
# scrubber.
agentic_loop_internal_litellm_params: Final = [
"_agentic_loop_depth",
"_agentic_loop_fingerprints",
"_agentic_loop_api_surface",
"max_agentic_loops",
"_code_interpreter_interception_active",
"_code_interpreter_interception_sandbox_key",
"_code_interpreter_interception_session_scoped",
"_code_interpreter_interception_converted_stream",
"_websearch_interception_emit_native_blocks",
"_websearch_interception_converted_stream",
"_headroom_interception_converted_stream",
agentic_loop_internal_litellm_params: Final = list(AGENTIC_LOOP_KWARG_NAMES) # mutable-ok: public type stays a list
bedrock_batch_litellm_params: Final = BEDROCK_BATCH_KWARG_NAMES
TRUSTED_CALLBACK_VARS_FIELD: Final = _litellm_params.TRUSTED_CALLBACK_VARS_FIELD
ADDRESSED_RESPONSE_ID_FIELD: Final = _litellm_params.ADDRESSED_RESPONSE_ID_FIELD
all_litellm_params = [ # rebind-ok: two star imports in litellm/__init__.py re-bind it # mutable-ok: callers concat
*OWNED_KWARG_NAMES,
*KWARG_ARTIFACTS,
*StandardCallbackDynamicParams.__annotations__,
*CustomPricingLiteLLMParams.model_fields,
]
# Proxy-owned callback credentials, stamped from admin-configured team/key callback
# settings. Listed in all_litellm_params for the same reason as the agentic-loop
# fields above: an unrecognized top-level key is swept into extra_body and sent to
# the provider.
TRUSTED_CALLBACK_VARS_FIELD: Final = "litellm_trusted_callback_vars"
ADDRESSED_RESPONSE_ID_FIELD: Final = "_litellm_addressed_response_id"
# Bedrock managed-batch deployment config, read from litellm_params by the batch and
# files transformations. Listed for the same reason as the fields above: these sit on
# a deployment that also serves chat, so leaking them into extra_body makes Bedrock
# reject every non-batch request to that deployment.
bedrock_batch_litellm_params: Final = (
"aws_batch_role_arn",
"s3_bucket_name",
"s3_region_name",
"s3_endpoint_url",
"s3_output_bucket_name",
"s3_bucket_owner",
"s3_access_key_id",
"s3_secret_access_key",
"s3_encryption_key_id",
"bedrock_tags",
)
all_litellm_params = (
agentic_loop_internal_litellm_params
+ [TRUSTED_CALLBACK_VARS_FIELD, ADDRESSED_RESPONSE_ID_FIELD, *bedrock_batch_litellm_params]
+ [
"_context_compaction_state",
"metadata",
"litellm_metadata",
"keepalive_seconds",
"allow_client_keepalive_override",
"litellm_trace_id",
"litellm_request_debug",
"guardrails",
"tags",
"acompletion",
"aimg_generation",
"atext_completion",
"text_completion",
"caching",
"mock_response",
"mock_timeout",
"disable_add_transform_inline_image_block",
"api_key",
"api_version",
"prompt_id",
"prompt_variables",
"litellm_system_prompt",
"provider_specific_header",
"prompt_version",
"prompt_environment",
"api_base",
"force_timeout",
"logger_fn",
"verbose",
"custom_llm_provider",
"model_file_id_mapping",
"litellm_logging_obj",
"litellm_call_id",
"completion_call_id",
"model_alias_map",
"custom_prompt_dict",
"stream_response",
"cost_per_query",
"ssl_verify",
"data_residency",
"async_call",
"aembedding",
"allm_passthrough_route",
"_litellm_strip_stream_usage",
"use_client",
"id",
"fallbacks",
"routing_strategy",
"_router_weights",
"azure",
"headers",
"model_list",
"num_retries",
"context_window_fallback_dict",
"retry_policy",
"retry_strategy",
"roles",
"final_prompt_value",
"bos_token",
"eos_token",
"request_timeout",
"client_side_timeout",
"complete_response",
"self",
"client",
"rpm",
"tpm",
"default_api_key_rpm_limit",
"default_api_key_tpm_limit",
"itpm",
"otpm",
"max_parallel_requests",
"input_cost_per_token",
"output_cost_per_token",
"input_cost_per_second",
"output_cost_per_second",
"hf_model_name",
"model_info",
"proxy_server_request",
"secret_fields",
"preset_cache_key",
"caching_groups",
"ttl",
"cache",
"enable_prompt_caching",
"no-log",
"base_model",
"stream_timeout",
"stream_chunk_size",
"supports_system_message",
"region_name",
"allowed_model_region",
"model_config",
"fastest_response",
"cooldown_time",
"cache_key",
"max_retries",
"azure_ad_token_provider",
"tenant_id",
"client_id",
"azure_username",
"azure_password",
"azure_scope",
"client_secret",
"user_continue_message",
"configurable_clientside_auth_params",
"weight",
"ensure_alternating_roles",
"assistant_continue_message",
"user_continue_message",
"fallback_depth",
"max_fallbacks",
"attempted_targets",
"max_budget",
"budget_duration",
"use_in_pass_through",
"merge_reasoning_content_in_choices",
"litellm_credential_name",
"allowed_openai_params",
"litellm_session_id",
"provider_affinity_header",
"use_litellm_proxy",
"use_chat_completions_api",
"rust",
"prompt_label",
"shared_session",
"search_tool_name",
"order",
"enable_tag_filtering",
"enable_json_schema_validation",
"use_xai_oauth",
"auto_router_config_path",
"auto_router_config",
"auto_router_default_model",
"auto_router_embedding_model",
"auto_router_max_input_chars",
"auto_router_routing_compression",
"auto_router_model_compression",
"complexity_router_config",
"complexity_router_default_model",
"adaptive_router_config",
"adaptive_router_default_model",
"quality_router_config",
"quality_router_default_model",
]
+ list(StandardCallbackDynamicParams.__annotations__.keys())
+ list(CustomPricingLiteLLMParams.model_fields.keys())
)
class KeyGenerationConfig(TypedDict, total=False):
required_params: list[str] # specify params that must be present in the key generation request

View file

@ -5,10 +5,11 @@ import logging
import os
import sys
import zlib
from collections.abc import Callable
from collections.abc import Callable, Mapping
from contextlib import ExitStack, contextmanager
from dataclasses import dataclass
from io import BytesIO
from types import SimpleNamespace
from types import MappingProxyType, SimpleNamespace
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
@ -16,7 +17,7 @@ import httpx
import pytest
from fastapi import HTTPException, Request, Response, UploadFile
from fastapi.responses import StreamingResponse
from pydantic import ValidationError
from pydantic import TypeAdapter, ValidationError
from starlette.datastructures import FormData, Headers, QueryParams
from starlette.datastructures import UploadFile as StarletteUploadFile
@ -45,6 +46,7 @@ from litellm.proxy.pass_through_endpoints.success_handler import (
PassThroughEndpointLogging,
)
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
from litellm.types import utils as types_utils
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY,
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
@ -7305,6 +7307,156 @@ def test_passthrough_logs_the_resolved_deployment_model_info_over_the_request_bo
assert kwargs["litellm_params"]["metadata"]["model_info"] == {"id": "vertex-gemini-38-flash-dep"}
@dataclass(frozen=True, slots=True, kw_only=True)
class _PassThroughSplit:
litellm_params: Mapping[str, object]
forwarded_body: Mapping[str, object]
_LITELLM_PARAMS: Final = TypeAdapter(dict[str, object])
_PROXY_SERVER_REQUEST: Final = TypeAdapter(dict[str, object])
def _split_pass_through_body(body: str) -> _PassThroughSplit:
mock_request: Final = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent"
mock_request.headers = Headers()
mock_request.scope = MappingProxyType({})
init_kwargs_for_pass_through_endpoint: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint # pyright: ignore[reportUnknownVariableType, reportUnknownMemberType] # untyped legacy helper
kwargs: Final = init_kwargs_for_pass_through_endpoint(
request=mock_request,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
passthrough_logging_payload=MagicMock(),
logging_obj=MagicMock(),
_parsed_body=json.loads(body),
litellm_call_id="lit-owned-keys-call-id",
)
validate_litellm_params: Final = _LITELLM_PARAMS.validate_python # pyright: ignore[reportUnknownArgumentType] # untyped legacy helper
litellm_params: Final = validate_litellm_params(kwargs["litellm_params"])
return _PassThroughSplit(
litellm_params=MappingProxyType(litellm_params),
forwarded_body=MappingProxyType(
_LITELLM_PARAMS.validate_python(
_PROXY_SERVER_REQUEST.validate_python(litellm_params["proxy_server_request"])["body"]
)
),
)
GEMINI_BODY: Final = '{"contents": [{"parts": [{"text": "hi"}]}], "generationConfig": {"temperature": 0}}'
def _metadata_of(split: _PassThroughSplit) -> Mapping[str, object]:
return MappingProxyType(_LITELLM_PARAMS.validate_python(split.litellm_params["metadata"]))
def test_passthrough_moves_every_litellm_owned_key_from_the_forwarded_body_into_litellm_params() -> None:
split: Final = _split_pass_through_body(
'{"ttl": 30, "contents": [{"parts": [{"text": "hi"}]}], "num_retries": 2,'
' "generationConfig": {"temperature": 0}, "litellm_trace_id": "trace-a"}'
)
assert frozenset(split.litellm_params) == frozenset(
("ttl", "num_retries", "litellm_trace_id", "metadata", "proxy_server_request")
)
assert tuple(split.litellm_params[k] for k in ("ttl", "num_retries", "litellm_trace_id")) == (30, 2, "trace-a")
assert split.forwarded_body == json.loads(GEMINI_BODY)
PROXY_STAMPED_NAMES: Final = frozenset(
(
"proxy_server_request",
"secret_fields",
"litellm_trusted_callback_vars",
"_litellm_addressed_response_id",
"_litellm_strip_stream_usage",
"client_side_timeout",
"model_file_id_mapping",
)
)
@pytest.mark.parametrize(
"name",
sorted(frozenset(litellm.all_litellm_params) - frozenset(("metadata", "litellm_metadata")) - PROXY_STAMPED_NAMES),
)
def test_passthrough_keeps_each_registered_litellm_owned_name_out_of_the_forwarded_body(name: str) -> None:
split: Final = _split_pass_through_body(json.dumps({name: "owned", **json.loads(GEMINI_BODY)}))
assert frozenset(split.litellm_params) == frozenset((name, "metadata", "proxy_server_request"))
assert split.litellm_params[name] == "owned"
assert split.forwarded_body == json.loads(GEMINI_BODY)
@pytest.mark.parametrize("name", sorted(PROXY_STAMPED_NAMES))
def test_passthrough_drops_a_client_supplied_proxy_stamped_name(name: str) -> None:
split: Final = _split_pass_through_body(json.dumps({name: {"forged": "by-client"}, **json.loads(GEMINI_BODY)}))
assert frozenset(split.litellm_params) == frozenset(("metadata", "proxy_server_request"))
assert split.litellm_params["proxy_server_request"] != {"forged": "by-client"}, split.litellm_params
assert split.forwarded_body == json.loads(GEMINI_BODY)
def test_passthrough_merges_both_metadata_carriers_from_the_body_into_one_metadata_key() -> None:
split: Final = _split_pass_through_body(
'{"metadata": {"client_tag": "a"}, "contents": [{"parts": [{"text": "hi"}]}], "ttl": 30,'
' "litellm_metadata": {"lm": "b"}, "generationConfig": {"temperature": 0}}'
)
assert frozenset(split.litellm_params) == frozenset(("ttl", "metadata", "proxy_server_request"))
assert _metadata_of(split) == {**_metadata_of(_split_pass_through_body(GEMINI_BODY)), "client_tag": "a", "lm": "b"}
assert split.forwarded_body == json.loads(GEMINI_BODY)
def test_passthrough_lets_metadata_win_over_litellm_metadata_on_a_shared_key() -> None:
split: Final = _split_pass_through_body(
'{"litellm_metadata": {"shared": "from-litellm-metadata", "lm": "b"},'
' "metadata": {"shared": "from-metadata", "client_tag": "a"}, "contents": []}'
)
assert _metadata_of(split) == {
**_metadata_of(_split_pass_through_body('{"contents": []}')),
"shared": "from-metadata",
"lm": "b",
"client_tag": "a",
}
def test_passthrough_orders_extracted_litellm_params_by_the_registry() -> None:
body: Final = json.dumps({"ttl": 30, "tags": ["team-a"], "num_retries": 2, "contents": []})
split: Final = _split_pass_through_body(body)
body_keys: Final = frozenset(json.loads(body))
assert tuple(k for k in split.litellm_params if k in body_keys) == tuple(
k for k in types_utils.all_litellm_params if k in body_keys
)
LATE_REGISTERED_BODY: Final = '{"registered_later": 1, "contents": [{"parts": [{"text": "hi"}]}]}'
def test_passthrough_sees_a_name_appended_to_the_public_list_after_import() -> None:
litellm.all_litellm_params.append("registered_later")
try:
split: Final = _split_pass_through_body(LATE_REGISTERED_BODY)
finally:
litellm.all_litellm_params.remove("registered_later")
assert frozenset(split.litellm_params) == frozenset(("registered_later", "metadata", "proxy_server_request"))
assert split.forwarded_body == {"contents": [{"parts": [{"text": "hi"}]}]}
def test_passthrough_sees_the_public_list_rebound_after_import(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(types_utils, "all_litellm_params", (*litellm.all_litellm_params, "registered_later"))
split: Final = _split_pass_through_body(LATE_REGISTERED_BODY)
assert frozenset(split.litellm_params) == frozenset(("registered_later", "metadata", "proxy_server_request"))
assert split.forwarded_body == {"contents": [{"parts": [{"text": "hi"}]}]}
@pytest.mark.asyncio
async def test_chat_completion_pass_through_endpoint_answers_an_openai_typed_error_for_an_unknown_model(
monkeypatch: pytest.MonkeyPatch,

View file

@ -3730,6 +3730,23 @@ def test_scoped_weights_are_excluded_from_provider_params(filter_name: str) -> N
assert filtered == {"provider_option": "kept"}
@pytest.mark.parametrize(
"provider_filter",
[
litellm.utils.get_non_default_completion_params,
litellm.utils.get_non_default_transcription_params,
litellm.utils.filter_out_litellm_params,
],
)
@pytest.mark.parametrize("setting", [("tag_regex", ["^team-a$"]), ("max_file_size_mb", 5)])
def test_deployment_only_settings_copied_by_the_router_stay_out_of_provider_params(
provider_filter: Callable[[dict[str, object]], Mapping[str, object]], setting: tuple[str, object]
) -> None:
name, value = setting
filtered: Final = provider_filter({"provider_option": "kept", name: value})
assert filtered == {"provider_option": "kept"}, filtered
class TestGetOptionalParamsTencent:
"""Tests that tencent provider uses TencentChatConfig for parameter mapping."""

View file

@ -0,0 +1,655 @@
import inspect
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass, field, fields
from operator import attrgetter
from types import MappingProxyType
from typing import Final, TypeAlias, cast, get_type_hints
import httpx
import pytest
from aiohttp import ClientSession
from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
import litellm
from litellm.caching.caching import Cache
from litellm.litellm_core_utils.get_litellm_params import (
get_litellm_params, # pyright: ignore[reportUnknownVariableType] # untyped legacy carrier
)
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.router_strategy.complexity_router.context_compaction import CompactionState
from litellm.router_utils.fallback_event_handlers import AttemptedFallbackTargets
from litellm.types import litellm_params
from litellm.types import utils as types_utils
from litellm.types.caching import DynamicCacheControl
from litellm.types.litellm_params import (
ADDRESSED_RESPONSE_ID_FIELD,
LITELLM_OWNED_ROOTS,
TRUSTED_CALLBACK_VARS_FIELD,
CachingOptions,
owned_wire_names,
wire,
wire_names,
)
from litellm.types.llms.openai import ChatCompletionAssistantMessage, ChatCompletionUserMessage
from litellm.types.proxy.litellm_pre_call_utils import SecretFields
from litellm.types.router import (
ConfigurableClientsideParamsCustomAuth,
CredentialLiteLLMParams,
DeploymentTypedDict,
RetryPolicy,
RouterConfig,
UpdateRouterConfig,
)
from litellm.types.router_weights import RouterWeights
from litellm.types.utils import (
CustomPricingLiteLLMParams,
ModelResponse,
ModelResponseStream,
ProviderSpecificHeader,
StandardCallbackDynamicParams,
agentic_loop_internal_litellm_params,
all_litellm_params,
bedrock_batch_litellm_params,
)
from litellm.utils import (
filter_out_litellm_params, # pyright: ignore[reportUnknownVariableType] # untyped legacy classifier
get_non_default_completion_params, # pyright: ignore[reportUnknownVariableType] # untyped legacy classifier
get_non_default_transcription_params, # pyright: ignore[reportUnknownVariableType] # untyped legacy classifier
)
PROVIDER_KNOB: Final = "registry_test_provider_only_knob"
CONNECTION_NAMES: Final = (
"api_key",
"api_base",
"api_version",
"region_name",
"headers",
"provider_specific_header",
"client",
"shared_session",
"ssl_verify",
"request_timeout",
"force_timeout",
"stream_timeout",
"max_retries",
"tenant_id",
"client_id",
"client_secret",
"azure_username",
"azure_password",
"azure_scope",
"azure_ad_token_provider",
"litellm_credential_name",
"configurable_clientside_auth_params",
"use_xai_oauth",
"aws_batch_role_arn",
"s3_bucket_name",
"s3_region_name",
"s3_endpoint_url",
"s3_output_bucket_name",
"s3_bucket_owner",
"s3_access_key_id",
"s3_secret_access_key",
"s3_encryption_key_id",
"bedrock_tags",
)
OPTION_NAMES: Final = (
"custom_llm_provider",
"azure",
"use_litellm_proxy",
"use_chat_completions_api",
"use_in_pass_through",
"allowed_openai_params",
"fallbacks",
"context_window_fallback_dict",
"num_retries",
"retry_policy",
"retry_strategy",
"routing_strategy",
"cooldown_time",
"allowed_model_region",
"enable_tag_filtering",
"fastest_response",
"provider_affinity_header",
"search_tool_name",
"model_list",
"model_info",
"rpm",
"tpm",
"itpm",
"otpm",
"default_api_key_rpm_limit",
"default_api_key_tpm_limit",
"max_parallel_requests",
"weight",
"order",
"tag_regex",
"max_file_size_mb",
"auto_router_config_path",
"auto_router_config",
"auto_router_default_model",
"auto_router_embedding_model",
"auto_router_max_input_chars",
"auto_router_routing_compression",
"auto_router_model_compression",
"complexity_router_config",
"complexity_router_default_model",
"adaptive_router_config",
"adaptive_router_default_model",
"quality_router_config",
"quality_router_default_model",
"caching",
"cache",
"ttl",
"enable_prompt_caching",
"caching_groups",
"cost_per_query",
"base_model",
"max_budget",
"budget_duration",
"id",
"metadata",
"litellm_metadata",
"tags",
"litellm_trace_id",
"litellm_session_id",
"litellm_request_debug",
"logger_fn",
"verbose",
"no-log",
"max_agentic_loops",
"guardrails",
"prompt_id",
"prompt_variables",
"prompt_version",
"prompt_environment",
"prompt_label",
"litellm_system_prompt",
"custom_prompt_dict",
"roles",
"final_prompt_value",
"bos_token",
"eos_token",
"hf_model_name",
"supports_system_message",
"ensure_alternating_roles",
"user_continue_message",
"assistant_continue_message",
"disable_add_transform_inline_image_block",
"merge_reasoning_content_in_choices",
"enable_json_schema_validation",
"complete_response",
"stream_chunk_size",
"keepalive_seconds",
"allow_client_keepalive_override",
"mock_response",
"mock_timeout",
)
AGENTIC_LOOP_STATE_NAMES: Final = (
"_agentic_loop_depth",
"_agentic_loop_fingerprints",
"_agentic_loop_api_surface",
"_code_interpreter_interception_active",
"_code_interpreter_interception_sandbox_key",
"_code_interpreter_interception_session_scoped",
"_code_interpreter_interception_converted_stream",
"_websearch_interception_emit_native_blocks",
"_websearch_interception_converted_stream",
"_headroom_interception_converted_stream",
)
INTERNAL_STATE_NAMES: Final = (
"litellm_call_id",
"completion_call_id",
"model_alias_map",
"data_residency",
"litellm_logging_obj",
"preset_cache_key",
"cache_key",
"stream_response",
"_context_compaction_state",
*AGENTIC_LOOP_STATE_NAMES,
"_router_weights",
"fallback_depth",
"max_fallbacks",
"attempted_targets",
"proxy_server_request",
"secret_fields",
"litellm_trusted_callback_vars",
"_litellm_addressed_response_id",
"_litellm_strip_stream_usage",
"client_side_timeout",
"model_file_id_mapping",
"acompletion",
"aembedding",
"aimg_generation",
"atext_completion",
"text_completion",
"allm_passthrough_route",
"async_call",
)
BEDROCK_BATCH_NAMES: Final = (
"aws_batch_role_arn",
"s3_bucket_name",
"s3_region_name",
"s3_endpoint_url",
"s3_output_bucket_name",
"s3_bucket_owner",
"s3_access_key_id",
"s3_secret_access_key",
"s3_encryption_key_id",
"bedrock_tags",
)
ARTIFACT_NAMES: Final = ("self", "use_client", "model_config", "rust")
CALLBACK_VAR_NAMES: Final = tuple(StandardCallbackDynamicParams.__annotations__)
PRICING_NAMES: Final = tuple(CustomPricingLiteLLMParams.model_fields)
OWNED_NAMES: Final = (
*CONNECTION_NAMES,
*OPTION_NAMES,
*INTERNAL_STATE_NAMES,
*ARTIFACT_NAMES,
*CALLBACK_VAR_NAMES,
*PRICING_NAMES,
)
Classifier: TypeAlias = Callable[[dict[str, object]], dict[str, object]] # mutable-ok: classifiers use dict
CLASSIFIERS: Final[Mapping[str, Classifier]] = MappingProxyType(
{ # pyright: ignore[reportUnknownArgumentType] # untyped legacy classifiers
"completion": get_non_default_completion_params,
"transcription": get_non_default_transcription_params,
"filter_out": filter_out_litellm_params,
}
)
@pytest.mark.parametrize("classifier_name", CLASSIFIERS)
@pytest.mark.parametrize("name", OWNED_NAMES)
def test_owned_name_is_kept_out_of_provider_params(name: str, classifier_name: str) -> None:
provider_value: Final = object()
classify: Final = CLASSIFIERS[classifier_name]
result: Final = classify({name: object(), PROVIDER_KNOB: provider_value}) # mutable-ok: classifiers take a dict
assert result == MappingProxyType({PROVIDER_KNOB: provider_value})
assert result[PROVIDER_KNOB] is provider_value
def test_a_name_no_object_declares_reaches_the_provider() -> None:
result: Final = CLASSIFIERS["completion"]({PROVIDER_KNOB: 1}) # mutable-ok: classifier input type
assert result == MappingProxyType({PROVIDER_KNOB: 1})
def _cache_key_for_model_group(cache: Cache, model_group: str, options: CachingOptions) -> str:
return cache.get_cache_key( # pyright: ignore[reportUnknownMemberType] # untyped legacy key builder
model=model_group,
messages=(MappingProxyType({"role": "user", "content": "shared prompt"}),),
metadata=MappingProxyType({"caching_groups": options.caching_groups, "model_group": model_group}),
)
def test_caching_groups_is_a_flat_sequence_of_model_groups_that_share_one_cache_key(
monkeypatch: pytest.MonkeyPatch,
) -> None:
for callback_list in ("input_callback", "success_callback", "_async_success_callback"):
monkeypatch.setattr(litellm, callback_list, []) # mutable-ok: Cache() appends "cache" to these lists
options: Final = CachingOptions(caching_groups=(("gpt-4", "gpt-4o"), ("claude-3",)))
cache: Final = Cache()
keys: Final = tuple(_cache_key_for_model_group(cache, group, options) for group in ("gpt-4", "gpt-4o", "claude-3"))
assert (keys[0] == keys[1], keys[0] == keys[2]) == (True, False)
def test_all_litellm_params_is_exactly_the_owned_inventory() -> None:
assert frozenset(all_litellm_params) == frozenset(OWNED_NAMES)
assert frozenset(ARTIFACT_NAMES).isdisjoint(DECLARED_NAMES)
def test_every_owned_name_has_exactly_one_owner() -> None:
duplicated: Final = tuple(name for name in dict.fromkeys(all_litellm_params) if all_litellm_params.count(name) > 1)
assert duplicated == ()
@pytest.mark.parametrize(
("exported", "declared"),
(
pytest.param(
types_utils.TRUSTED_CALLBACK_VARS_FIELD,
litellm_params.TRUSTED_CALLBACK_VARS_FIELD,
id="TRUSTED_CALLBACK_VARS_FIELD",
),
pytest.param(
types_utils.ADDRESSED_RESPONSE_ID_FIELD,
litellm_params.ADDRESSED_RESPONSE_ID_FIELD,
id="ADDRESSED_RESPONSE_ID_FIELD",
),
),
)
def test_types_utils_still_exports_the_field_constant(exported: str, declared: str) -> None:
assert exported == declared
@dataclass(frozen=True, slots=True, kw_only=True)
class _Leaf:
plain: int | None = None
renamed: int | None = field(default=None, metadata=wire("wire-name"))
@dataclass(frozen=True, slots=True, kw_only=True)
class _OtherLeaf:
plain: int | None = None
trailing: int | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class _Root:
first: _Leaf
second: _OtherLeaf
@dataclass(frozen=True, slots=True, kw_only=True)
class _RootDeclaringAKwargDirectly:
first: _Leaf
stray: int | None = None
def test_wire_names_are_the_field_names_in_declaration_order_unless_wire_renames_them() -> None:
assert wire_names(_Leaf) == ("plain", "wire-name")
def test_owned_wire_names_walk_leaves_in_declaration_order_and_keep_every_occurrence() -> None:
assert owned_wire_names(_Root) == ("plain", "wire-name", "plain", "trailing")
def test_owned_wire_names_refuse_a_root_that_declares_a_kwarg_outside_a_leaf() -> None:
with pytest.raises(TypeError):
owned_wire_names(_RootDeclaringAKwargDirectly)
def test_agentic_loop_names_concatenate_as_a_list() -> None:
extended: Final = agentic_loop_internal_litellm_params + ["caller_added"] # mutable-ok: list contract under test
assert (type(extended), len(extended), frozenset(extended)) == (
list,
len(AGENTIC_LOOP_STATE_NAMES) + 2,
frozenset((*AGENTIC_LOOP_STATE_NAMES, "max_agentic_loops", "caller_added")),
)
def test_bedrock_batch_names_concatenate_as_a_tuple() -> None:
extended: Final = bedrock_batch_litellm_params + ("caller_added",)
assert extended == (*BEDROCK_BATCH_NAMES, "caller_added")
def test_proxy_stamped_fields_keep_their_wire_names() -> None:
assert (TRUSTED_CALLBACK_VARS_FIELD, ADDRESSED_RESPONSE_ID_FIELD) == (
"litellm_trusted_callback_vars",
"_litellm_addressed_response_id",
)
def test_all_litellm_params_concatenates_with_a_list_like_the_completion_entrypoint_does() -> None:
extended: Final = ["aembedding", "extra_headers"] + all_litellm_params # mutable-ok: list contract under test
assert (type(extended), frozenset(extended)) == (list, frozenset(("aembedding", "extra_headers", *OWNED_NAMES)))
CARRIED_AND_FORWARDED: Final = frozenset(("drop_params", "hugging_face", "no_log", "replicate", "together_ai"))
CARRIER_SIGNATURE: Final = inspect.signature(get_litellm_params) # pyright: ignore[reportUnknownArgumentType] # legacy
CARRIED_PARAMS: Final = tuple(
name for name in CARRIER_SIGNATURE.parameters if name != "kwargs" and name not in CARRIED_AND_FORWARDED
)
@pytest.mark.parametrize("name", CARRIED_PARAMS)
def test_every_param_get_litellm_params_carries_is_kept_out_of_provider_params(name: str) -> None:
provider_value: Final = object()
result: Final = CLASSIFIERS["completion"](
{name: object(), PROVIDER_KNOB: provider_value} # mutable-ok: classifier input type
)
assert result == MappingProxyType({PROVIDER_KNOB: provider_value})
TYPED_CONFIG_MODELS: Final[Mapping[str, tuple[type[BaseModel], ...]]] = MappingProxyType(
{
"credentials": (CredentialLiteLLMParams,),
"router": (RouterConfig, UpdateRouterConfig),
}
)
DECLARED_NAMES: Final = frozenset(name for root in LITELLM_OWNED_ROOTS for name in owned_wire_names(root))
ProviderClient: TypeAlias = (
OpenAI
| AsyncOpenAI
| AzureOpenAI
| AsyncAzureOpenAI
| HTTPHandler
| AsyncHTTPHandler
| httpx.Client
| httpx.AsyncClient
)
MockResponse: TypeAlias = str | Exception | Mapping[str, object] | Sequence[float] | ModelResponse | ModelResponseStream
TYPE_HINT_NAMESPACE: Final[Mapping[str, object]] = {
"ProviderClient": ProviderClient,
"ProviderSpecificHeader": ProviderSpecificHeader,
"ClientSession": ClientSession,
"AsyncAzureOpenAI": AsyncAzureOpenAI,
"AsyncOpenAI": AsyncOpenAI,
"AzureOpenAI": AzureOpenAI,
"OpenAI": OpenAI,
"AsyncHTTPHandler": AsyncHTTPHandler,
"HTTPHandler": HTTPHandler,
"ConfigurableClientsideParamsCustomAuth": ConfigurableClientsideParamsCustomAuth,
"RetryPolicy": RetryPolicy,
"DeploymentTypedDict": DeploymentTypedDict,
"DynamicCacheControl": DynamicCacheControl,
"ChatCompletionUserMessage": ChatCompletionUserMessage,
"ChatCompletionAssistantMessage": ChatCompletionAssistantMessage,
"MockResponse": MockResponse,
"ModelResponse": ModelResponse,
"ModelResponseStream": ModelResponseStream,
"Logging": Logging,
"SecretFields": SecretFields,
"CompactionState": CompactionState,
"RouterWeights": RouterWeights,
"AttemptedFallbackTargets": AttemptedFallbackTargets,
}
LEAF_SAMPLES: Final[Mapping[type, Mapping[str, object]]] = {
litellm_params.ProviderConnection: {"api_key": "k", "request_timeout": 1.5},
litellm_params.BedrockBatchConnection: {"aws_batch_role_arn": "arn", "bedrock_tags": ({"k": "v"},)},
litellm_params.DispatchOptions: {"custom_llm_provider": "openai"},
litellm_params.RoutingOptions: {
"fallbacks": [{"model": "gpt-4o", "api_key": "k", "temperature": 0}],
"num_retries": 2,
"retry_strategy": "constant_retry",
"routing_strategy": "simple-shuffle",
},
litellm_params.DeploymentOptions: {"model_info": {"region": "us"}, "rpm": 2},
litellm_params.SpecializedRouterOptions: {"adaptive_router_default_model": "gpt-4o"},
litellm_params.CachingOptions: {"ttl": 30.0, "caching_groups": (("gpt-4o", "gpt-4o-mini"),)},
litellm_params.CostOptions: {"max_budget": 10.0},
litellm_params.ObservabilityOptions: {"metadata": {"request": "test"}, "no_log": True},
litellm_params.AgenticLoopOptions: {"max_agentic_loops": 2},
litellm_params.GuardrailOptions: {"guardrails": ("default",)},
litellm_params.PromptOptions: {"prompt_id": "prompt", "prompt_variables": {"name": "value"}},
litellm_params.ResponseOptions: {"stream_chunk_size": 64},
litellm_params.MockOptions: {"mock_timeout": True},
litellm_params.CallState: {
"completion_call_id": "call",
"model_alias_map": {"alias": "gpt-4o"},
"data_residency": "us",
},
litellm_params.AgenticLoopState: {"api_surface": "chat_completions", "depth": 1},
litellm_params.RouterState: {"fallback_depth": 1},
litellm_params.ProxyRequestState: {
"proxy_server_request": {"path": "/chat/completions"},
"trusted_callback_vars": {"dd_api_key": "k"},
},
litellm_params.EntrypointState: {"acompletion": True},
}
LEAF_BAD_SAMPLES: Final[Mapping[type, Mapping[str, object]]] = {
litellm_params.ProviderConnection: {"api_key": 1},
litellm_params.BedrockBatchConnection: {"aws_batch_role_arn": 1},
litellm_params.DispatchOptions: {"custom_llm_provider": 1},
litellm_params.RoutingOptions: {"num_retries": "2"},
litellm_params.DeploymentOptions: {"rpm": "2"},
litellm_params.SpecializedRouterOptions: {"auto_router_max_input_chars": "2"},
litellm_params.CachingOptions: {"ttl": "30"},
litellm_params.CostOptions: {"max_budget": "10"},
litellm_params.ObservabilityOptions: {"verbose": "true"},
litellm_params.AgenticLoopOptions: {"max_agentic_loops": "2"},
litellm_params.GuardrailOptions: {"guardrails": (1,)},
litellm_params.PromptOptions: {"prompt_id": 1},
litellm_params.ResponseOptions: {"stream_chunk_size": "64"},
litellm_params.MockOptions: {"mock_timeout": "true"},
litellm_params.CallState: {"completion_call_id": 1},
litellm_params.AgenticLoopState: {"depth": "1"},
litellm_params.RouterState: {"fallback_depth": "1"},
litellm_params.ProxyRequestState: {"proxy_server_request": "request"},
litellm_params.EntrypointState: {"acompletion": "true"},
}
INVALID_LITERAL_SAMPLES: Final[tuple[tuple[type, Mapping[str, object]], ...]] = (
(litellm_params.RoutingOptions, {"retry_strategy": "linear"}),
(litellm_params.RoutingOptions, {"routing_strategy": "random"}),
(litellm_params.AgenticLoopState, {"api_surface": "batches"}),
)
def _leaf_id(value: object) -> str:
return value.__name__ if isinstance(value, type) else ""
def _leaf_instance(leaf: type, sample: Mapping[str, object]) -> object:
constructor: Final = cast(Callable[..., object], leaf)
return constructor(**sample)
def _strict_leaf_validation(leaf: type, instance: object) -> object:
hints: Final[Mapping[str, object]] = cast(
Mapping[str, object], get_type_hints(type(instance), localns=TYPE_HINT_NAMESPACE)
)
for field_info in fields(leaf):
value = cast(Callable[[object], object], attrgetter(field_info.name))(instance)
field_adapter: TypeAdapter[object] = TypeAdapter[object](
hints[field_info.name],
config=ConfigDict(arbitrary_types_allowed=True),
)
field_adapter.validate_python(value, strict=True)
return instance
@pytest.mark.parametrize("leaf,sample", LEAF_SAMPLES.items(), ids=_leaf_id)
def test_every_owned_leaf_accepts_a_strict_reader_shaped_sample(leaf: type, sample: Mapping[str, object]) -> None:
instance: Final = _leaf_instance(leaf, sample)
result: Final = _strict_leaf_validation(leaf, instance)
assert result == instance
assert frozenset(sample) <= frozenset(field.name for field in fields(leaf))
@pytest.mark.parametrize("leaf,sample", LEAF_BAD_SAMPLES.items(), ids=_leaf_id)
def test_every_owned_leaf_rejects_a_strict_wrong_typed_sample(leaf: type, sample: Mapping[str, object]) -> None:
instance: Final = _leaf_instance(leaf, sample)
with pytest.raises(ValidationError):
_strict_leaf_validation(leaf, instance)
@pytest.mark.parametrize("leaf,sample", INVALID_LITERAL_SAMPLES, ids=_leaf_id)
def test_owned_leaf_literals_reject_unknown_values(leaf: type, sample: Mapping[str, object]) -> None:
instance: Final = _leaf_instance(leaf, sample)
with pytest.raises(ValidationError):
_strict_leaf_validation(leaf, instance)
@pytest.mark.parametrize(
"strategy",
[
"simple-shuffle",
"least-busy",
"usage-based-routing",
"latency-based-routing",
"cost-based-routing",
"usage-based-routing-v2",
"lar1",
],
)
def test_routing_options_accept_every_strategy_the_router_accepts(strategy: str) -> None:
instance: Final = _leaf_instance(litellm_params.RoutingOptions, {"routing_strategy": strategy})
assert _strict_leaf_validation(litellm_params.RoutingOptions, instance) is instance
NAMES_SHARED_WITH_TYPED_MODELS: Final[Mapping[str, tuple[str, ...]]] = MappingProxyType(
{
"credentials": (
"api_base",
"api_key",
"api_version",
"aws_batch_role_arn",
"azure_password",
"azure_scope",
"azure_username",
"bedrock_tags",
"client_id",
"client_secret",
"region_name",
"s3_access_key_id",
"s3_bucket_name",
"s3_bucket_owner",
"s3_encryption_key_id",
"s3_endpoint_url",
"s3_output_bucket_name",
"s3_region_name",
"s3_secret_access_key",
"tenant_id",
),
"router": (
"caching_groups",
"cooldown_time",
"enable_tag_filtering",
"fallbacks",
"max_retries",
"model_list",
"num_retries",
"retry_policy",
"routing_strategy",
),
}
)
@pytest.mark.parametrize("source", TYPED_CONFIG_MODELS)
def test_names_a_typed_config_model_shares_with_the_owned_inventory_are_exactly_these(source: str) -> None:
model_names: Final = frozenset(name for model in TYPED_CONFIG_MODELS[source] for name in model.model_fields)
assert DECLARED_NAMES & model_names == frozenset(NAMES_SHARED_WITH_TYPED_MODELS[source])
@pytest.mark.parametrize("name", PRICING_NAMES)
def test_pricing_name_is_owned_by_the_pricing_model_alone(name: str) -> None:
assert name not in DECLARED_NAMES