Merge pull request #35926 from BerriAI/litellm_remove_types_ruff_exclusion

chore(lint): remove litellm/types from the ruff lint exclusion
This commit is contained in:
Mateo Wang 2026-08-05 12:35:02 -07:00 • committed by GitHub
commit 332ec6c17a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
270 changed files with 5047 additions and 5348 deletions

View file

@ -1,6 +1,6 @@
{
"reportAny": {
"limit": 29809
"limit": 29806
},
"reportArgumentType": {
"limit": 2645
@ -21,10 +21,10 @@
"limit": 215
},
"reportDuplicateImport": {
"limit": 24
"limit": 19
},
"reportExplicitAny": {
"limit": 9473
"limit": 9469
},
"reportFunctionMemberAccess": {
"limit": 7
@ -105,13 +105,13 @@
"limit": 113
},
"reportUnknownMemberType": {
"limit": 40452
"limit": 40447
},
"reportUnknownParameterType": {
"limit": 20309
},
"reportUnknownVariableType": {
"limit": 31978
"limit": 31880
},
"reportUnnecessaryCast": {
"limit": 124
@ -126,7 +126,7 @@
"limit": 866
},
"reportUntypedBaseClass": {
"limit": 165
"limit": 72
},
"reportUntypedFunctionDecorator": {
"limit": 33
@ -138,9 +138,9 @@
"limit": 139
},
"reportUnusedImport": {
"limit": 588
"limit": 556
},
"reportUnusedVariable": {
"limit": 147
"limit": 146
}
}

View file

@ -1461,32 +1461,30 @@ _UTILS_MODULE_IMPORT_MAP: Final = {
# Export all name tuples and import maps for use in _lazy_imports.py
__all__ = [
# Name tuples
"COST_CALCULATOR_NAMES",
"LITELLM_LOGGING_NAMES",
"UTILS_NAMES",
"TOKEN_COUNTER_NAMES",
"LLM_CLIENT_CACHE_NAMES",
"BEDROCK_TYPES_NAMES",
"TYPES_UTILS_NAMES",
"CACHING_NAMES",
"HTTP_HANDLER_NAMES",
"COST_CALCULATOR_NAMES",
"DOTPROMPT_NAMES",
"HTTP_HANDLER_NAMES",
"LITELLM_LOGGING_NAMES",
"LLM_CLIENT_CACHE_NAMES",
"LLM_CONFIG_NAMES",
"TYPES_NAMES",
"LLM_PROVIDER_LOGIC_NAMES",
"TOKEN_COUNTER_NAMES",
"TYPES_NAMES",
"TYPES_UTILS_NAMES",
"UTILS_MODULE_NAMES",
# Import maps
"_UTILS_IMPORT_MAP",
"_COST_CALCULATOR_IMPORT_MAP",
"_TYPES_UTILS_IMPORT_MAP",
"_TOKEN_COUNTER_IMPORT_MAP",
"UTILS_NAMES",
"_BEDROCK_TYPES_IMPORT_MAP",
"_CACHING_IMPORT_MAP",
"_LITELLM_LOGGING_IMPORT_MAP",
"_COST_CALCULATOR_IMPORT_MAP",
"_DOTPROMPT_IMPORT_MAP",
"_TYPES_IMPORT_MAP",
"_LITELLM_LOGGING_IMPORT_MAP",
"_LLM_CONFIGS_IMPORT_MAP",
"_LLM_PROVIDER_LOGIC_IMPORT_MAP",
"_TOKEN_COUNTER_IMPORT_MAP",
"_TYPES_IMPORT_MAP",
"_TYPES_UTILS_IMPORT_MAP",
"_UTILS_IMPORT_MAP",
"_UTILS_MODULE_IMPORT_MAP",
]

View file

@ -1,6 +1,6 @@
import asyncio
from datetime import datetime, timedelta
from typing import TYPE_CHECKING, Any, Final, Union
from typing import TYPE_CHECKING, Any, Final
import litellm
from litellm._logging import verbose_logger
@ -16,7 +16,7 @@ if TYPE_CHECKING:
from litellm.proxy._types import UserAPIKeyAuth
Span = Union[_Span, Any]
Span = _Span | Any
OTELClass = OpenTelemetry
else:
Span = Any

View file

@ -55,19 +55,15 @@ from litellm.a2a_protocol.main import (
from litellm.types.agents import LiteLLMSendMessageResponse
__all__ = [
# Client
"A2AClient",
# Functions
"asend_message",
"send_message",
"asend_message_streaming",
"aget_agent_card",
"create_a2a_client",
# Response types
"LiteLLMSendMessageResponse",
# Exceptions
"A2AError",
"A2AConnectionError",
"A2AAgentCardError",
"A2AClient",
"A2AConnectionError",
"A2AError",
"A2ALocalhostURLError",
"LiteLLMSendMessageResponse",
"aget_agent_card",
"asend_message",
"asend_message_streaming",
"create_a2a_client",
"send_message",
]

View file

@ -8,8 +8,8 @@ from ..types.llms.openai import *
def get_optional_params_add_message(
role: str | None,
content: str | List[MessageContentTextObject | MessageContentImageFileObject | MessageContentImageURLObject] | None,
attachments: List[Attachment] | None,
content: str | list[MessageContentTextObject | MessageContentImageFileObject | MessageContentImageURLObject] | None,
attachments: list[Attachment] | None,
metadata: dict | None,
custom_llm_provider: str,
**kwargs,
@ -57,7 +57,7 @@ def get_optional_params_add_message(
optional_params = litellm.AzureOpenAIAssistantsAPIConfig().map_openai_params_create_message_params(
non_default_params=non_default_params, optional_params=optional_params
)
for k in passed_params.keys():
for k in passed_params:
if k not in default_params:
optional_params[k] = passed_params[k]
return optional_params
@ -128,7 +128,7 @@ def get_optional_params_image_gen(
if n is not None:
optional_params["sampleCount"] = int(n)
for k in passed_params.keys():
for k in passed_params:
if k not in default_params:
optional_params[k] = passed_params[k]
return optional_params

View file

@ -9,12 +9,12 @@ Has 4 methods:
"""
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any, Final, Union
from typing import TYPE_CHECKING, Any, Final
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
Span = Union[_Span, Any]
Span = _Span | Any
else:
Span = Any

View file

@ -1,12 +1,12 @@
import json
from typing import TYPE_CHECKING, Any, Final, Union
from typing import TYPE_CHECKING, Any, Final
from .base_cache import BaseCache
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
Span = Union[_Span, Any]
Span = _Span | Any
else:
Span = Any

View file

@ -13,7 +13,7 @@ import time
import traceback
from concurrent.futures import ThreadPoolExecutor
from threading import Lock
from typing import TYPE_CHECKING, Any, Final, Union
from typing import TYPE_CHECKING, Any, Final
if TYPE_CHECKING:
from litellm.types.caching import RedisPipelineIncrementOperation
@ -29,7 +29,7 @@ from .redis_cache import RedisCache
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
Span = Union[_Span, Any]
Span = _Span | Any
else:
Span = Any

View file

@ -18,7 +18,7 @@ import time
from collections.abc import Awaitable, Callable, Sequence
from contextvars import ContextVar
from datetime import timedelta
from typing import TYPE_CHECKING, Any, Final, TypeVar, Union, cast
from typing import TYPE_CHECKING, Any, Final, TypeVar, cast
import litellm
from litellm._logging import print_verbose, verbose_logger
@ -49,7 +49,7 @@ if TYPE_CHECKING:
cluster_pipeline = ClusterPipeline
async_redis_client = Redis
async_redis_cluster_client = RedisCluster
Span = Union[_Span, Any]
Span = _Span | Any
else:
pipeline = Any
cluster_pipeline = Any

View file

@ -5,7 +5,7 @@ Key differences:
- RedisClient NEEDs to be re-used across requests, adds 3000ms latency if it's re-created
"""
from typing import TYPE_CHECKING, Any, Final, Union
from typing import TYPE_CHECKING, Any, Final
from litellm.caching.redis_cache import RedisCache
@ -16,7 +16,7 @@ if TYPE_CHECKING:
pipeline = Pipeline
async_redis_client = Redis
Span = Union[_Span, Any]
Span = _Span | Any
else:
pipeline = Any
async_redis_client = Any

View file

@ -367,7 +367,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
stream_options = normalize_responses_api_stream_options(value)
if stream_options is not None:
responses_api_request["stream_options"] = stream_options
elif key in ResponsesAPIOptionalRequestParams.__annotations__.keys():
elif key in ResponsesAPIOptionalRequestParams.__annotations__:
responses_api_request[key] = value
elif key == "previous_response_id":
responses_api_request["previous_response_id"] = value

View file

@ -23,22 +23,20 @@ from .main import (
)
__all__ = [
# Core container operations
"acreate_container",
"adelete_container",
"alist_containers",
"aretrieve_container",
"create_container",
"delete_container",
"list_containers",
"retrieve_container",
# Container file operations (auto-generated from endpoints.json)
"adelete_container_file",
"alist_container_files",
"alist_containers",
"aretrieve_container",
"aretrieve_container_file",
"aretrieve_container_file_content",
"create_container",
"delete_container",
"delete_container_file",
"list_container_files",
"list_containers",
"retrieve_container",
"retrieve_container_file",
"retrieve_container_file_content",
]

View file

@ -80,7 +80,7 @@ async def acreate_fine_tuning_job(
hyperparameters: dict | None = {},
suffix: str | None = None,
validation_file: str | None = None,
integrations: List[str] | None = None,
integrations: list[str] | None = None,
seed: int | None = None,
custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai",
extra_headers: dict[str, str] | None = None,
@ -157,7 +157,7 @@ def create_fine_tuning_job(
hyperparameters: dict | None = {},
suffix: str | None = None,
validation_file: str | None = None,
integrations: List[str] | None = None,
integrations: list[str] | None = None,
seed: int | None = None,
custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai",
extra_headers: dict[str, str] | None = None,

View file

@ -6,7 +6,7 @@ this file has Arize ai specific helper functions
import os
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Union
from typing import TYPE_CHECKING, Any, Final
from litellm.integrations.arize import _utils
from litellm.integrations.arize._utils import ArizeOTELAttributes
@ -21,7 +21,7 @@ if TYPE_CHECKING:
from litellm.types.integrations.arize import Protocol as _Protocol
Protocol = _Protocol
Span = Union[_Span, Any]
Span = _Span | Any
else:
Protocol = Any
Span = Any

View file

@ -1,7 +1,7 @@
import os
import threading
from collections import OrderedDict
from typing import TYPE_CHECKING, Any, Final, Union
from typing import TYPE_CHECKING, Any, Final
from litellm._logging import verbose_logger
from litellm.integrations.arize import _utils
@ -22,7 +22,7 @@ if TYPE_CHECKING:
Protocol = _Protocol
OpenTelemetryConfig = _OpenTelemetryConfig
Span = Union[_Span, Any]
Span = _Span | Any
OpenTelemetry = _OpenTelemetry
LITELLM_TRACER_NAME: str
else:

View file

@ -3,7 +3,7 @@
import re
import traceback
from collections.abc import AsyncGenerator
from typing import TYPE_CHECKING, Any, Final, Optional, Union
from typing import TYPE_CHECKING, Any, Final, Optional
from pydantic import BaseModel
@ -39,7 +39,7 @@ if TYPE_CHECKING:
)
from litellm.types.router import PreRoutingHookResponse
Span = Union[_Span, Any]
Span = _Span | Any
else:
Span = Any
LiteLLMLoggingObj = Any

View file

@ -31,7 +31,7 @@ class PromptTemplate:
self.output_format = self.metadata.get("output", {}).get("format")
self.output_schema = self.metadata.get("output", {}).get("schema", {})
self.optional_params = {}
for key in self.metadata.keys():
for key in self.metadata:
if key not in restricted_keys:
self.optional_params[key] = self.metadata[key]

View file

@ -2,7 +2,7 @@ import base64
import json
import os
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Optional, Union
from typing import TYPE_CHECKING, Any, Final, Optional
from litellm._logging import verbose_logger
from litellm.integrations.arize import _utils
@ -18,7 +18,7 @@ from litellm.types.utils import StandardCallbackDynamicParams
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
Span = Union[_Span, Any]
Span = _Span | Any
else:
Span = Any

View file

@ -4,7 +4,7 @@ Call Hook for LiteLLM Proxy which allows Langfuse prompt management.
import os
from functools import lru_cache
from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias, Union, cast
from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias, cast
from packaging.version import Version
@ -30,7 +30,7 @@ if TYPE_CHECKING:
LangfuseClass: TypeAlias = Langfuse
PROMPT_CLIENT = Union[TextPromptClient, ChatPromptClient]
PROMPT_CLIENT = TextPromptClient | ChatPromptClient
else:
PROMPT_CLIENT = Any
LangfuseClass = Any

View file

@ -1,12 +1,12 @@
import json
from typing import TYPE_CHECKING, Any, Final, Union
from typing import TYPE_CHECKING, Any, Final
from litellm.proxy._types import SpanAttributes
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
Span = Union[_Span, Any]
Span = _Span | Any
else:
Span = Any

View file

@ -1,5 +1,5 @@
import os
from typing import TYPE_CHECKING, Any, Final, Union
from typing import TYPE_CHECKING, Any, Final
from litellm.integrations.opentelemetry import OpenTelemetry
@ -13,7 +13,7 @@ if TYPE_CHECKING:
Protocol = _Protocol
OpenTelemetryConfig = _OpenTelemetryConfig
Span = Union[_Span, Any]
Span = _Span | Any
else:
Protocol = Any
OpenTelemetryConfig = Any

View file

@ -1,7 +1,7 @@
import os
from dataclasses import dataclass, field
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Union, cast
from typing import TYPE_CHECKING, Any, Final, cast
import litellm
from litellm._logging import verbose_logger
@ -47,12 +47,12 @@ if TYPE_CHECKING:
)
from litellm.proxy.proxy_server import UserAPIKeyAuth as _UserAPIKeyAuth
Span = Union[_Span, Any]
Tracer = Union[_Tracer, Any]
Context = Union[_Context, Any]
SpanExporter = Union[_SpanExporter, Any]
UserAPIKeyAuth = Union[_UserAPIKeyAuth, Any]
ManagementEndpointLoggingPayload = Union[_ManagementEndpointLoggingPayload, Any]
Span = _Span | Any
Tracer = _Tracer | Any
Context = _Context | Any
SpanExporter = _SpanExporter | Any
UserAPIKeyAuth = _UserAPIKeyAuth | Any
ManagementEndpointLoggingPayload = _ManagementEndpointLoggingPayload | Any
else:
Span = Any
Tracer = Any
@ -186,16 +186,7 @@ def _normalize_team_metadata_keys(value: Any) -> list[str]:
_FREEZE_MAX_DEPTH: Final = 16
HashableScope = Union[
str,
int,
float,
bool,
bytes,
None,
tuple["HashableScope", ...],
frozenset["HashableScope"],
]
HashableScope = str | int | float | bool | bytes | None | tuple["HashableScope", ...] | frozenset["HashableScope"]
def _freeze_for_dedupe(value: object, _depth: int = 0) -> HashableScope:

View file

@ -31,7 +31,7 @@ Events:
from datetime import datetime
from enum import Enum
from typing import TYPE_CHECKING, Any, Final, Union
from typing import TYPE_CHECKING, Any, Final
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
@ -40,7 +40,7 @@ if TYPE_CHECKING:
from litellm.integrations.opentelemetry import OpenTelemetryConfig
Span = Union[_Span, Any]
Span = _Span | Any
else:
Span = Any

View file

@ -1,7 +1,7 @@
"""Type definitions for Opik payload building."""
from dataclasses import dataclass
from typing import Any, Final, Literal, Union
from typing import Any, Final, Literal
@dataclass
@ -42,5 +42,5 @@ class SpanPayload:
total_cost: float | None = None
PayloadItem = Union[TracePayload, SpanPayload]
PayloadItem = TracePayload | SpanPayload
TraceSpanPayloadTuple: Final = tuple[TracePayload | None, SpanPayload]

View file

@ -72,53 +72,49 @@ from litellm.integrations.otel.model.spans import (
)
__all__ = [
# config
"OTEL_V2_ENV",
"OpenTelemetryV2Config",
"is_otel_v2_enabled",
# semconv
"BAGGAGE_PROMOTED_KEYS",
"DB",
"DEFAULT_BAGGAGE_METADATA_KEYS",
"HTTP",
"MCP",
"OTEL_V2_ENV",
"SPAN_REGISTRY",
"Client",
"Error",
"GenAI",
"GenAIOperation",
"GenAIProvider",
"HTTP",
"JsonRpc",
"LiteLLM",
"LiteLLMError",
"MCP",
"MCPMethod",
"Metric",
"Network",
"NetworkTransport",
"Server",
"resolve_operation",
"resolve_provider",
# spans
"SPAN_REGISTRY",
"LiteLLMSpanKind",
"SpanRole",
"SpanSpec",
"db_system",
"span_role_for_service",
"validate_registry",
# payloads
"GuardrailSpanData",
"JsonRpc",
"LLMCallSpanData",
"LLMRequestParams",
"LLMUsage",
"LiteLLM",
"LiteLLMError",
"LiteLLMSpanKind",
"MCPListToolsSpanData",
"MCPMethod",
"MCPToolCallSpanData",
"Metric",
"Network",
"NetworkTransport",
"OpenTelemetryV2Config",
"ProxyRequestSpanData",
"RequestContext",
"RequestIdentity",
"Server",
"ServerInfo",
"ServiceSpanData",
"SpanError",
"SpanRole",
"SpanSpec",
"db_system",
"is_mcp_list_tools",
"is_mcp_tool_call",
"is_otel_v2_enabled",
"promoted_baggage",
"resolve_operation",
"resolve_provider",
"span_role_for_service",
"validate_registry",
]

View file

@ -66,18 +66,13 @@ from litellm.interactions.main import (
)
__all__ = [
# Create
"create",
"acreate",
# Get
"get",
"aget",
# Delete
"delete",
"adelete",
# Cancel
"cancel",
"acancel",
# Sub-modules
"acreate",
"adelete",
"agents",
"aget",
"cancel",
"create",
"delete",
"get",
]

View file

@ -2,7 +2,7 @@
## Helper utilities
import copy
from collections.abc import Iterable
from typing import TYPE_CHECKING, Any, Final, Literal, Union
from typing import TYPE_CHECKING, Any, Final, Literal
import httpx
@ -14,7 +14,7 @@ if TYPE_CHECKING:
from litellm.types.utils import ModelResponseStream
Span = Union[_Span, Any]
Span = _Span | Any
else:
Span = Any

View file

@ -48,7 +48,7 @@ O(number of rules); callers must only invoke them on a cache miss.
import re
from dataclasses import dataclass
from typing import Final, Union
from typing import Final
from litellm._logging import verbose_logger
@ -100,7 +100,7 @@ class _CapabilityRule:
model_info: dict
_CompiledRule = Union[_RoutingRule, _CapabilityRule]
_CompiledRule = _RoutingRule | _CapabilityRule
def _compile_rule(rule: object) -> tuple[_CompiledRule, ...]:

View file

@ -4827,7 +4827,7 @@ class StandardLoggingPayloadSetup:
# Populate well-known typed fields with int/str coercion where needed
typed_keys: Final[dict] = {}
for key in StandardLoggingAdditionalHeaders.__annotations__.keys():
for key in StandardLoggingAdditionalHeaders.__annotations__:
_key = key.lower().replace("_", "-")
typed_keys[_key] = key
if _key in additiona_headers:
@ -4859,7 +4859,7 @@ class StandardLoggingPayloadSetup:
usage_object=None,
)
if hidden_params is not None:
for key in StandardLoggingHiddenParams.__annotations__.keys():
for key in StandardLoggingHiddenParams.__annotations__:
if key in hidden_params:
if key == "additional_headers":
clean_hidden_params["additional_headers"] = StandardLoggingPayloadSetup.get_additional_headers(
@ -5501,7 +5501,7 @@ def get_standard_logging_metadata(
)
if isinstance(metadata, dict):
# Update the clean_metadata with values from input metadata that match StandardLoggingMetadata fields
for key in StandardLoggingMetadata.__annotations__.keys():
for key in StandardLoggingMetadata.__annotations__:
if key in metadata:
clean_metadata[key] = metadata[key]

View file

@ -4,7 +4,7 @@ import inspect
import re
import time
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Union
from typing import TYPE_CHECKING, Any, Final
from litellm._logging import verbose_logger
from litellm.constants import MAX_BASE64_LENGTH_FOR_LOGGING
@ -23,7 +23,7 @@ if TYPE_CHECKING:
)
LiteLLMModelResponse = _ModelResponse
Span = Union[_Span, Any]
Span = _Span | Any
else:
LiteLLMModelResponse = Any
LiteLLMLoggingObject = Any

View file

@ -47,7 +47,7 @@ def is_model_response_stream_empty(model_response: ModelResponseStream) -> bool:
# Check for any non-base fields that are set
# Access model_fields on the class, not the instance, to avoid Pydantic 2.11+ deprecation warnings
for model_response_field in type(model_response).model_fields.keys():
for model_response_field in type(model_response).model_fields:
# Skip base fields that are always set
if model_response_field in BASE_FIELDS:
continue

View file

@ -8,7 +8,7 @@ import time
import traceback
from collections.abc import AsyncIterator, Callable, Iterator
from dataclasses import dataclass
from typing import Any, Final, NoReturn, TypeVar, Union, cast
from typing import Any, Final, NoReturn, TypeVar, cast
import anyio
import httpx
@ -99,7 +99,7 @@ class _ProviderChunkEarlyReturn:
value: Any
_ProviderChunkResult = Union[_ProviderChunkParsed, _ProviderChunkEarlyReturn]
_ProviderChunkResult = _ProviderChunkParsed | _ProviderChunkEarlyReturn
class CustomStreamWrapper:
@ -256,9 +256,7 @@ class CustomStreamWrapper:
chunk = chunk.strip()
self.complete_response = self.complete_response.strip()
if chunk.startswith(self.complete_response):
# Remove last_sent_chunk only if it appears at the start of the new chunk
chunk = chunk[len(self.complete_response) :]
chunk = chunk.removeprefix(self.complete_response)
self.complete_response += chunk
return chunk

View file

@ -5,7 +5,7 @@
import base64
import json
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any, Final, Generic, TypeVar, Union, cast
from typing import TYPE_CHECKING, Any, Final, Generic, TypeVar, cast
from litellm import verbose_logger
from litellm.llms.base_llm.managed_resources.isolation import (
@ -23,7 +23,7 @@ if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient as _PrismaClient
from litellm.router import Router as _Router
Span = Union[_Span, Any]
Span = _Span | Any
InternalUsageCache = _InternalUsageCache
PrismaClient = _PrismaClient
Router = _Router

View file

@ -57,7 +57,7 @@ class AmazonCohereChatConfig:
Reference - https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-cohere-command-r-plus.html
"""
documents: List[Document] | None = None
documents: list[Document] | None = None
search_queries_only: bool | None = None
preamble: str | None = None
max_tokens: int | None = None
@ -69,12 +69,12 @@ class AmazonCohereChatConfig:
presence_penalty: float | None = None
seed: int | None = None
return_prompt: bool | None = None
stop_sequences: List[str] | None = None
stop_sequences: list[str] | None = None
raw_prompting: bool | None = None
def __init__(
self,
documents: List[Document] | None = None,
documents: list[Document] | None = None,
search_queries_only: bool | None = None,
preamble: str | None = None,
max_tokens: int | None = None,
@ -112,7 +112,7 @@ class AmazonCohereChatConfig:
and v is not None
}
def get_supported_openai_params(self) -> List[str]:
def get_supported_openai_params(self) -> list[str]:
return [
"max_tokens",
"max_completion_tokens",
@ -325,7 +325,7 @@ class AWSEventStreamDecoder:
self.model = model
self.parser = EventStreamJSONParser()
self.content_blocks: List[ContentBlockDeltaEvent] = []
self.content_blocks: list[ContentBlockDeltaEvent] = []
self.tool_calls_index: int | None = None
self.response_id: str | None = None
self.json_mode = json_mode
@ -362,13 +362,13 @@ class AWSEventStreamDecoder:
def translate_thinking_blocks(
self, thinking_block: BedrockConverseReasoningContentBlockDelta
) -> List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None:
) -> list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None:
"""
Translate the thinking blocks to a string
"""
thinking_blocks_list: Final[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]] = []
_thinking_block: Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock] | None = None
thinking_blocks_list: Final[list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock]] = []
_thinking_block: ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock | None = None
if "text" in thinking_block:
_thinking_block = ChatCompletionThinkingBlock(type="thinking")
@ -402,12 +402,12 @@ class AWSEventStreamDecoder:
) -> tuple[
ChatCompletionToolCallChunk | None,
dict,
List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None,
list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None,
]:
"""Handle 'start' event in converse chunk parsing."""
tool_use: ChatCompletionToolCallChunk | None = None
provider_specific_fields: dict = {}
thinking_blocks: List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None = None
thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None = None
self.content_blocks = [] # reset
if start_obj is not None:
@ -450,14 +450,14 @@ class AWSEventStreamDecoder:
ChatCompletionToolCallChunk | None,
dict,
str | None,
List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None,
list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None,
]:
"""Handle 'delta' event in converse chunk parsing."""
text = ""
tool_use: ChatCompletionToolCallChunk | None = None
provider_specific_fields: dict = {}
reasoning_content: str | None = None
thinking_blocks: List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None = None
thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None = None
self.content_blocks.append(delta_obj)
if "text" in delta_obj:
@ -535,7 +535,7 @@ class AWSEventStreamDecoder:
usage: Usage | None = None
provider_specific_fields: dict = {}
reasoning_content: str | None = None
thinking_blocks: List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None = None
thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None = None
content_block_index: Final = int(chunk_data.get("contentBlockIndex", 0))
if "start" in chunk_data:
@ -590,7 +590,7 @@ class AWSEventStreamDecoder:
except Exception as e:
raise Exception(f"Received streaming error - {e}")
def _chunk_parser(self, chunk_data: dict) -> Union[GChunk, ModelResponseStream, dict]:
def _chunk_parser(self, chunk_data: dict) -> GChunk | ModelResponseStream | dict:
text = ""
is_finished = False
finish_reason = ""
@ -645,7 +645,7 @@ class AWSEventStreamDecoder:
tool_use=None,
)
def iter_bytes(self, iterator: Iterator[bytes]) -> Iterator[Union[GChunk, ModelResponseStream, dict]]:
def iter_bytes(self, iterator: Iterator[bytes]) -> Iterator[GChunk | ModelResponseStream | dict]:
"""Given an iterator that yields lines, iterate over it & yield every event encountered"""
from botocore.eventstream import EventStreamBuffer
@ -659,9 +659,7 @@ class AWSEventStreamDecoder:
_data = json.loads(message)
yield self._chunk_parser(chunk_data=_data)
async def aiter_bytes(
self, iterator: AsyncIterator[bytes]
) -> AsyncIterator[Union[GChunk, ModelResponseStream, dict]]:
async def aiter_bytes(self, iterator: AsyncIterator[bytes]) -> AsyncIterator[GChunk | ModelResponseStream | dict]:
"""Given an async iterator that yields lines, iterate over it & yield every event encountered"""
from botocore.eventstream import EventStreamBuffer
@ -741,7 +739,7 @@ class AmazonDeepSeekR1StreamDecoder(AWSEventStreamDecoder):
sync_stream=sync_stream,
)
def _chunk_parser(self, chunk_data: dict) -> Union[GChunk, ModelResponseStream, dict]:
def _chunk_parser(self, chunk_data: dict) -> GChunk | ModelResponseStream | dict:
return self.deepseek_model_response_iterator.chunk_parser(chunk=chunk_data)
@ -756,7 +754,7 @@ class MockResponseIterator: # for returning ai21 streaming responses
return self
def _handle_json_mode_chunk(
self, text: str, tool_calls: List[ChatCompletionToolCallChunk] | None
self, text: str, tool_calls: list[ChatCompletionToolCallChunk] | None
) -> tuple[str, ChatCompletionToolCallChunk | None]:
"""
If JSON mode is enabled, convert the tool call to a message.
@ -789,7 +787,7 @@ class MockResponseIterator: # for returning ai21 streaming responses
text = chunk_data.choices[0].message.content or ""
tool_use = None
_model_response_tool_call: Final = cast(
List[ChatCompletionMessageToolCall] | None,
list[ChatCompletionMessageToolCall] | None,
cast(Choices, chunk_data.choices[0]).message.tool_calls,
)
if self.json_mode is True:

View file

@ -34,7 +34,7 @@ class BedrockCohereEmbeddingConfig:
new_transformed_request: Final = CohereEmbeddingRequest(
input_type=transformed_request["input_type"],
)
for k in CohereEmbeddingRequest.__annotations__.keys():
for k in CohereEmbeddingRequest.__annotations__:
if k in transformed_request:
new_transformed_request[k] = transformed_request[k]

View file

@ -1,7 +1,7 @@
from __future__ import annotations
import json
from typing import TYPE_CHECKING, Any, Final, Union
from typing import TYPE_CHECKING, Any, Final
import httpx
from pydantic import BaseModel
@ -49,12 +49,12 @@ class BedrockImagePreparedRequest(BaseModel):
data: dict
BedrockImageConfigClass = Union[
type[AmazonTitanImageGenerationConfig],
type[AmazonNovaCanvasConfig],
type[AmazonStability3Config],
type[AmazonStabilityConfig],
]
BedrockImageConfigClass = (
type[AmazonTitanImageGenerationConfig]
| type[AmazonNovaCanvasConfig]
| type[AmazonStability3Config]
| type[AmazonStabilityConfig]
)
class BedrockImageGeneration(BaseAWSLLM):

View file

@ -160,10 +160,10 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
aws_filters: dict | None = None
if isinstance(value, dict):
if "operator" in value.keys():
if "operator" in value:
# Single operator - map directly (no wrapping needed)
aws_filters = self._map_operator_filter(value)
elif "and" in value.keys() or "or" in value.keys():
elif "and" in value or "or" in value:
aws_filters = self._map_and_or_filters(value)
else:
# Assume it's already in AWS KB format

View file

@ -345,7 +345,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
client: OpenAI | AsyncOpenAI | None = None,
shared_session: Optional["ClientSession"] = None,
) -> OpenAI | AsyncOpenAI | None:
client_initialization_params: Final[Dict] = locals()
client_initialization_params: Final[dict] = locals()
if client is None:
if not isinstance(max_retries, int):
raise OpenAIError(
@ -408,7 +408,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
data: dict,
timeout: float | httpx.Timeout,
logging_obj: LiteLLMLoggingObj,
) -> Tuple[dict, BaseModel]:
) -> tuple[dict, BaseModel]:
"""
Helper to:
- call chat.completions.create.with_raw_response when litellm.return_response_headers is True
@ -445,7 +445,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
data: dict,
timeout: float | httpx.Timeout,
logging_obj: LiteLLMLoggingObj,
) -> Tuple[dict, BaseModel]:
) -> tuple[dict, BaseModel]:
"""
Helper to:
- call chat.completions.create.with_raw_response when litellm.return_response_headers is True
@ -480,11 +480,11 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
self,
response: Any,
model: str,
messages: list[Dict],
optional_params: Dict,
messages: list[dict],
optional_params: dict,
logging_obj: LiteLLMLoggingObj,
stream: bool,
litellm_params: Dict,
litellm_params: dict,
) -> Any | None:
"""
Call agentic completion hooks for all custom loggers (OpenAI Chat Completions API).
@ -1294,7 +1294,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
)
## embedding CALL
headers: Dict | None = None
headers: dict | None = None
headers, sync_embedding_response = self.make_sync_openai_embedding_request(
openai_client=openai_client,
data=data,
@ -2848,7 +2848,7 @@ class OpenAIAssistantsAPI(BaseLLM):
assistant_id: str,
additional_instructions: str | None,
instructions: str | None,
metadata: Dict | None,
metadata: dict | None,
model: str | None,
stream: bool | None,
tools: Iterable[AssistantToolParam] | None,
@ -2887,12 +2887,12 @@ class OpenAIAssistantsAPI(BaseLLM):
assistant_id: str,
additional_instructions: str | None,
instructions: str | None,
metadata: Dict | None,
metadata: dict | None,
model: str | None,
tools: Iterable[AssistantToolParam] | None,
event_handler: AssistantEventHandler | None,
) -> AsyncAssistantStreamManager[AsyncAssistantEventHandler]:
data: Final[Dict[str, Any]] = {
data: Final[dict[str, Any]] = {
"thread_id": thread_id,
"assistant_id": assistant_id,
"additional_instructions": additional_instructions,
@ -2912,12 +2912,12 @@ class OpenAIAssistantsAPI(BaseLLM):
assistant_id: str,
additional_instructions: str | None,
instructions: str | None,
metadata: Dict | None,
metadata: dict | None,
model: str | None,
tools: Iterable[AssistantToolParam] | None,
event_handler: AssistantEventHandler | None,
) -> AssistantStreamManager[AssistantEventHandler]:
data: Final[Dict[str, Any]] = {
data: Final[dict[str, Any]] = {
"thread_id": thread_id,
"assistant_id": assistant_id,
"additional_instructions": additional_instructions,
@ -2939,7 +2939,7 @@ class OpenAIAssistantsAPI(BaseLLM):
assistant_id: str,
additional_instructions: str | None,
instructions: str | None,
metadata: Dict | None,
metadata: dict | None,
model: str | None,
stream: bool | None,
tools: Iterable[AssistantToolParam] | None,
@ -2961,7 +2961,7 @@ class OpenAIAssistantsAPI(BaseLLM):
assistant_id: str,
additional_instructions: str | None,
instructions: str | None,
metadata: Dict | None,
metadata: dict | None,
model: str | None,
stream: bool | None,
tools: Iterable[AssistantToolParam] | None,
@ -2984,7 +2984,7 @@ class OpenAIAssistantsAPI(BaseLLM):
assistant_id: str,
additional_instructions: str | None,
instructions: str | None,
metadata: Dict | None,
metadata: dict | None,
model: str | None,
stream: bool | None,
tools: Iterable[AssistantToolParam] | None,

View file

@ -165,10 +165,10 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
self,
model: str, # allows overrides to selectively run this
input: str | ResponseInputParam,
tools: List[ALL_RESPONSES_API_TOOL_PARAMS] | None = None,
) -> Tuple[
tools: list[ALL_RESPONSES_API_TOOL_PARAMS] | None = None,
) -> tuple[
str | ResponseInputParam,
List[ALL_RESPONSES_API_TOOL_PARAMS] | None,
list[ALL_RESPONSES_API_TOOL_PARAMS] | None,
]:
"""Sibling of `remove_cache_control_flag_from_messages_and_tools` on
the chat path. Strips Anthropic-only `cache_control` markers from
@ -447,7 +447,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> Tuple[str, dict]:
) -> tuple[str, dict]:
"""
Transform the delete response API request into a URL and data
@ -482,7 +482,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> Tuple[str, dict]:
) -> tuple[str, dict]:
"""
Transform the get response API request into a URL and data
@ -525,10 +525,10 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
headers: dict,
after: str | None = None,
before: str | None = None,
include: List[str] | None = None,
include: list[str] | None = None,
limit: int = 20,
order: Literal["asc", "desc"] = "desc",
) -> Tuple[str, dict]:
) -> tuple[str, dict]:
encoded_response_id: Final = encode_url_path_segment(response_id, field_name="response_id")
url: Final = f"{api_base}/{encoded_response_id}/input_items"
params: Final[dict[str, Any]] = {}
@ -563,7 +563,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> Tuple[str, dict]:
) -> tuple[str, dict]:
"""
Transform the cancel response API request into a URL and data
@ -607,7 +607,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> Tuple[str, dict]:
) -> tuple[str, dict]:
"""
Transform the compact response API request into a URL and data

View file

@ -132,7 +132,7 @@ class OpenAIWhisperAudioTranscriptionConfig(BaseAudioTranscriptionConfig):
raise
return TranscriptionResponse(text=raw_response.text)
if any(key in raw_response_json for key in TranscriptionResponse.model_fields.keys()):
if any(key in raw_response_json for key in TranscriptionResponse.model_fields):
return TranscriptionResponse(**raw_response_json)
else:
raise ValueError(

View file

@ -1,6 +1,6 @@
import warnings
from enum import Enum
from typing import Final, Literal, Union
from typing import Final, Literal
from pydantic import BaseModel, Field, field_validator, model_validator
@ -115,7 +115,7 @@ class SAPToolChatMessage(BaseModel):
_content_validator = field_validator("content", mode="before")(validate_different_content)
ChatMessage = Union[SAPMessage, SAPUserMessage, SAPAssistantMessage, SAPToolChatMessage]
ChatMessage = SAPMessage | SAPUserMessage | SAPAssistantMessage | SAPToolChatMessage
class ResponseFormat(BaseModel):

View file

@ -11,8 +11,6 @@ from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast
import httpx
import litellm
import litellm.litellm_core_utils
import litellm.litellm_core_utils.litellm_logging
from litellm import verbose_logger
from litellm._uuid import uuid
from litellm.constants import (
@ -2429,7 +2427,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
_candidates: Final = completion_response.get("candidates")
if _candidates and len(_candidates) > 0:
content_policy_violations: Final = VertexGeminiConfig().get_flagged_finish_reasons()
if "finishReason" in _candidates[0] and _candidates[0]["finishReason"] in content_policy_violations.keys():
if "finishReason" in _candidates[0] and _candidates[0]["finishReason"] in content_policy_violations:
return self._handle_content_policy_violation(
model_response=model_response,
completion_response=completion_response,

View file

@ -1,4 +1,5 @@
from collections.abc import Callable, Mapping, Sequence
from types import UnionType
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, Union, get_args, get_origin
import httpx
@ -475,14 +476,12 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
return 0
if annotation is list or origin is list:
return []
if origin is Union:
if origin is Union or origin is UnionType:
# Prefer empty list when any option is a list
if any((arg is list or VolcEngineResponsesAPIConfig._annotation_origin(arg) is list) for arg in args):
return []
if type(None) in args:
return None
if origin is Union and type(None) in args:
return None
# Fallback to None when no safer guess exists
return None
@ -514,7 +513,9 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
Choose the best-matching Pydantic model class for a nested dict.
"""
origin: Final = VolcEngineResponsesAPIConfig._annotation_origin(annotation)
union_args: Final = VolcEngineResponsesAPIConfig._annotation_args(annotation) if origin is Union else ()
union_args: Final = (
VolcEngineResponsesAPIConfig._annotation_args(annotation) if origin is Union or origin is UnionType else ()
)
candidates = tuple(candidate for candidate in (annotation, *union_args) if hasattr(candidate, "model_fields"))
if not candidates:

View file

@ -322,7 +322,7 @@ oci_transformation: Final = OCIChatConfig()
ovhcloud_transformation: Final = OVHCloudChatConfig()
lemonade_transformation: Final = LemonadeChatConfig()
MOCK_RESPONSE_TYPE = Union[str, Exception, dict, ModelResponse, ModelResponseStream]
MOCK_RESPONSE_TYPE = str | Exception | dict | ModelResponse | ModelResponseStream
####### COMPLETION ENDPOINTS ################

View file

@ -3375,10 +3375,10 @@ class MCPServerManager:
static_headers: Final = server.static_headers or {}
has_static_authorization: Final = any(
isinstance(k, str) and k.lower() == "authorization" for k in static_headers.keys()
isinstance(k, str) and k.lower() == "authorization" for k in static_headers
)
has_extra_authorization: Final = bool(extra_headers) and any(
isinstance(k, str) and k.lower() == "authorization" for k in (extra_headers or {}).keys()
isinstance(k, str) and k.lower() == "authorization" for k in (extra_headers or {})
)
if (
@ -4419,7 +4419,7 @@ class MCPServerManager:
allowed_params_list: Final = allowed_params[matched]
# Filter arguments to only include allowed parameters
disallowed_params: Final = [param for param in arguments.keys() if param not in allowed_params_list]
disallowed_params: Final = [param for param in arguments if param not in allowed_params_list]
if disallowed_params:
raise HTTPException(

View file

@ -1613,7 +1613,7 @@ if MCP_AVAILABLE:
``mcp_server_auth_headers``). Either form skips the pre-emptive 401.
"""
if oauth2_headers:
for k in oauth2_headers.keys():
for k in oauth2_headers:
if k.lower() == "authorization":
return True
return _client_has_per_server_auth_header(server, mcp_server_auth_headers)

View file

@ -3,7 +3,7 @@ import json
import os
from collections.abc import Callable
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Literal, Union
from typing import TYPE_CHECKING, Any, Final, Literal
import httpx
from pydantic import (
@ -67,7 +67,7 @@ from .types_utils.utils import get_instance_fn, validate_custom_validate_return_
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
Span = Union[_Span, Any]
Span = _Span | Any
else:
Span = Any
@ -4010,7 +4010,7 @@ class JWTKeyItem(TypedDict, total=False):
kid: str
JWKKeyValue = Union[list[JWTKeyItem], JWTKeyItem]
JWKKeyValue = list[JWTKeyItem] | JWTKeyItem
class JWKUrlResponse(TypedDict, total=False):
@ -4053,15 +4053,15 @@ class UserManagementEndpointParamDocStringEnums(str, enum.Enum):
duration_doc_str = """Optional[str] - Duration for the key auto-created on `/user/new`. Default is None."""
PassThroughEndpointLoggingResultValues = Union[
ModelResponse,
TextCompletionResponse,
ImageResponse,
EmbeddingResponse,
VideoObject,
StandardPassThroughResponseObject,
ResponsesAPIResponse,
]
PassThroughEndpointLoggingResultValues = (
ModelResponse
| TextCompletionResponse
| ImageResponse
| EmbeddingResponse
| VideoObject
| StandardPassThroughResponseObject
| ResponsesAPIResponse
)
class PassThroughEndpointLoggingTypedDict(TypedDict):
@ -4162,7 +4162,7 @@ class ClientSideFallbackModel(TypedDict, total=False):
messages: list[AllMessageValues]
ALL_FALLBACK_MODEL_VALUES = Union[str, ClientSideFallbackModel]
ALL_FALLBACK_MODEL_VALUES = str | ClientSideFallbackModel
RBAC_ROLES = Literal[

View file

@ -26,7 +26,7 @@ The two wire shapes:
from collections.abc import Callable
from types import ModuleType
from typing import Final, Literal, Union
from typing import Final, Literal
from pydantic import BaseModel
@ -34,7 +34,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.proxy.a2a.agent_card import normalize_protocol_version
A2AVersion = Literal["0.3", "1.0"]
RequestId = Union[str, int, None]
RequestId = str | int | None
JsonDict = dict[str, object]
_V1_SEND_ENVELOPE_KEYS: Final = frozenset({"message", "task"})

View file

@ -13,7 +13,7 @@ import asyncio
import math
import re
import time
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast
from fastapi import HTTPException, Request, status
from pydantic import BaseModel
@ -109,7 +109,7 @@ from .auth_utils import get_model_from_request, get_request_route_template
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
Span = Union[_Span, Any]
Span = _Span | Any
else:
Span = Any

View file

@ -2,7 +2,7 @@
Handles Authentication Errors
"""
from typing import TYPE_CHECKING, Any, Final, Union
from typing import TYPE_CHECKING, Any, Final
from fastapi import HTTPException, Request, status
@ -28,7 +28,7 @@ DB_UNAVAILABLE_FALLBACK_USER_ID: Final = "__db_unavailable_fallback__"
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
Span = Union[_Span, Any]
Span = _Span | Any
else:
Span = Any

View file

@ -7,7 +7,7 @@ External callers (public IPs) only see servers with available_on_public_internet
import ipaddress
from dataclasses import dataclass
from typing import Any, Final, Union
from typing import Any, Final
from fastapi import Request
from pydantic import TypeAdapter, ValidationError
@ -45,7 +45,7 @@ class _HopCount:
value: int
_HopCountSetting = Union[_HopCountUnset, _HopCountInvalid, _HopCount]
_HopCountSetting = _HopCountUnset | _HopCountInvalid | _HopCount
class IPAddressUtils:

View file

@ -1,14 +1,14 @@
from __future__ import annotations
import ipaddress
from typing import Any, Final, Union
from typing import Any, Final
from fastapi import Request
from pydantic import BaseModel, Field
from litellm._logging import verbose_proxy_logger
TrustedProxyNetwork = Union[ipaddress.IPv4Network, ipaddress.IPv6Network]
TrustedProxyNetwork = ipaddress.IPv4Network | ipaddress.IPv6Network
class NetworkContext(BaseModel):

View file

@ -177,7 +177,7 @@ async def create_batch(
}
input_file_id: Final = _create_batch_data.get("input_file_id", None)
unified_file_id: Union[str, Literal[False]] = False
unified_file_id: str | Literal[False] = False
model_from_file_id = None
if input_file_id:

View file

@ -1,4 +1,4 @@
from typing import Final, Literal, Union
from typing import Final, Literal
from pydantic import BaseModel, ConfigDict, Field, JsonValue, TypeAdapter
@ -56,7 +56,7 @@ class LLMClassifier(BaseModel):
timeout_ms: int = 3000
ClassifierChoice = Union[HeuristicClassifier, LLMClassifier]
ClassifierChoice = HeuristicClassifier | LLMClassifier
class NoSemanticMatching(BaseModel):
@ -88,7 +88,7 @@ class SemanticMatching(BaseModel):
keyword_tier_rules: tuple[KeywordTierRule, ...] = DEFAULT_KEYWORD_TIER_RULES
SemanticMatchingChoice = Union[NoSemanticMatching, SemanticMatching]
SemanticMatchingChoice = NoSemanticMatching | SemanticMatching
class AutorouteConfig(BaseModel):

View file

@ -1230,7 +1230,7 @@ class DBSpendUpdateWriter:
if team_member_list_transactions is not None and len(team_member_list_transactions.keys()) > 0:
# Track which team memberships will be updated for cache invalidation
team_memberships_to_invalidate: Final[list[tuple[str, str]]] = []
for key in team_member_list_transactions.keys():
for key in team_member_list_transactions:
# key is "team_id::<value>::user_id::<value>"
team_id = key.split("::")[1]
user_id = key.split("::")[3]

View file

@ -16,7 +16,7 @@ payload; the secret license key is never sent as an attribute or header.
import os
import tempfile
from dataclasses import dataclass
from typing import TYPE_CHECKING, Final, Optional, Union
from typing import TYPE_CHECKING, Final, Optional
from opentelemetry.exporter.otlp.proto.http.metric_exporter import OTLPMetricExporter
from opentelemetry.metrics import Counter
@ -53,7 +53,7 @@ _CA_CERT_FILENAME: Final = "ca.crt"
METRIC_NAME: Final = "litellm.enterprise.billable_requests"
METER_NAME: Final = "litellm.enterprise.billing"
AttributeValue = Union[str, int]
AttributeValue = str | int
@dataclass(frozen=True, slots=True)

View file

@ -132,7 +132,7 @@ async def create_fine_tuning_job(
)
## CHECK IF MANAGED FILE ID
unified_file_id: Union[str, Literal[False]] = False
unified_file_id: str | Literal[False] = False
training_file: Final = fine_tuning_request.training_file
response: LiteLLMFineTuningJob | None = None
if training_file:
@ -269,7 +269,7 @@ async def retrieve_fine_tuning_job(
custom_llm_provider = request_body.get("custom_llm_provider", None) or custom_llm_provider
## CHECK IF MANAGED FILE ID
unified_finetuning_job_id: Union[str, Literal[False]] = False
unified_finetuning_job_id: str | Literal[False] = False
response: LiteLLMFineTuningJob | None = None
if fine_tuning_job_id:
unified_finetuning_job_id = _is_base64_encoded_unified_file_id(fine_tuning_job_id)
@ -536,7 +536,7 @@ async def cancel_fine_tuning_job(
custom_llm_provider: Final = request_body.get("custom_llm_provider", None)
## CHECK IF MANAGED FILE ID
unified_finetuning_job_id: Union[str, Literal[False]] = False
unified_finetuning_job_id: str | Literal[False] = False
response: LiteLLMFineTuningJob | None = None
if fine_tuning_job_id:
unified_finetuning_job_id = _is_base64_encoded_unified_file_id(fine_tuning_job_id)

View file

@ -8,7 +8,8 @@ import json
import os
from collections.abc import Mapping, Sequence
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Any, Final, Literal, TypeVar, Union, cast
from types import UnionType
from typing import TYPE_CHECKING, Any, Final, Literal, TypeVar, Union, cast, get_args, get_origin
from urllib.parse import urlparse
from fastapi import APIRouter, Depends, HTTPException, Request
@ -1556,13 +1557,9 @@ def _get_field_type_from_annotation(field_annotation: Any) -> str:
Convert a Python type annotation to a UI-friendly type string
"""
# Handle Union types (like Optional[T])
if (
hasattr(field_annotation, "__origin__")
and field_annotation.__origin__ is Union
and hasattr(field_annotation, "__args__")
):
if get_origin(field_annotation) is Union or get_origin(field_annotation) is UnionType:
# For Optional[T], get the non-None type
args: Final = field_annotation.__args__
args: Final = get_args(field_annotation)
non_none_args: Final = [arg for arg in args if arg is not type(None)]
if non_none_args:
field_annotation = non_none_args[0]
@ -1689,13 +1686,9 @@ def _should_skip_optional_params(field_name: str, field_annotation: Any) -> bool
def _unwrap_optional_type(field_annotation: Any) -> Any:
"""Unwrap Optional types to get the actual type."""
if (
hasattr(field_annotation, "__origin__")
and field_annotation.__origin__ is Union
and hasattr(field_annotation, "__args__")
):
if get_origin(field_annotation) is Union or get_origin(field_annotation) is UnionType:
# For Optional[BaseModel], get the non-None type
args: Final = field_annotation.__args__
args: Final = get_args(field_annotation)
non_none_args: Final = [arg for arg in args if arg is not type(None)]
if non_none_args:
return non_none_args[0]

View file

@ -10,7 +10,7 @@ from litellm.types.guardrails import *
sys.path.insert(0, os.path.abspath("../..")) # Adds the parent directory to the system path
def can_modify_guardrails(team_obj: Optional[LiteLLM_TeamTable]) -> bool:
def can_modify_guardrails(team_obj: LiteLLM_TeamTable | None) -> bool:
if team_obj is None:
return True

View file

@ -265,7 +265,7 @@ class GenericGuardrailAPI(CustomGuardrail):
# Dynamically iterate through GenericGuardrailAPIMetadata fields
# and extract matching fields from the source metadata
# Fields in metadata are already prefixed with 'user_api_key_'
for field_name in GenericGuardrailAPIMetadata.__annotations__.keys():
for field_name in GenericGuardrailAPIMetadata.__annotations__:
value = metadata_dict.get(field_name)
if value is not None:
result_metadata[field_name] = value

View file

@ -1,5 +1,5 @@
from collections.abc import AsyncGenerator, Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Literal, Union
from typing import TYPE_CHECKING, Any, Final, Literal
import httpx
from fastapi import HTTPException
@ -57,7 +57,7 @@ class ModelArmorAPIError(Exception):
_SCANNED_CONTENT_KEYS: Final = frozenset({"text", "sanitizedText", "findings", "maliciousUriMatchedItems"})
RedactablePayload = Union[dict, list, str, int, float, bool, None]
RedactablePayload = dict | list | str | int | float | bool | None
def _redact_scanned_content(payload: RedactablePayload, depth: int = 0) -> RedactablePayload:

View file

@ -16,7 +16,6 @@ from typing import (
Any,
Final,
Literal,
Union,
)
from urllib.parse import urljoin
@ -54,7 +53,7 @@ SENSITIVE_DATA_DETECTOR_KEYS: Final[list[str]] = ["sensitiveData", "dataDetector
# Type aliases
MessageRole = Literal["user", "assistant"]
LLMResponse = Union[Any, ModelResponse, EmbeddingResponse, ImageResponse]
LLMResponse = Any | ModelResponse | EmbeddingResponse | ImageResponse
_LEGACY_NOMA_DEPRECATION_WARNED = False
if TYPE_CHECKING:

View file

@ -6,7 +6,7 @@ GET /guardrails/usage/overview, /guardrails/usage/detail/:id, /guardrails/usage/
import json
from collections.abc import Mapping, Sequence
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any, Final, Literal, Union, overload
from typing import TYPE_CHECKING, Any, Final, Literal, overload
from fastapi import APIRouter, Depends, Query
from pydantic import BaseModel
@ -31,8 +31,8 @@ if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient
from litellm.types.guardrails import Guardrail
_DbOrConfigGuardrail = Union[prisma_models.LiteLLM_GuardrailsTable, Guardrail]
_DailyMetricsRow = Union[prisma_models.LiteLLM_DailyGuardrailMetrics, prisma_models.LiteLLM_DailyPolicyMetrics]
_DbOrConfigGuardrail = prisma_models.LiteLLM_GuardrailsTable | Guardrail
_DailyMetricsRow = prisma_models.LiteLLM_DailyGuardrailMetrics | prisma_models.LiteLLM_DailyPolicyMetrics
router: Final = APIRouter()

View file

@ -7,7 +7,7 @@ import time
import traceback
from collections.abc import Iterable
from datetime import datetime, timedelta
from typing import Any, Final, Literal, TypedDict, Union, cast
from typing import Any, Final, Literal, TypedDict, cast
import fastapi
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
@ -110,7 +110,7 @@ def get_callback_identifier(callback):
router: Final = APIRouter()
services = Union[
services = (
Literal[
"slack_budget_alerts",
"langfuse",
@ -127,9 +127,9 @@ services = Union[
"galileo",
"newrelic",
"sqs",
],
str,
]
]
| str
)
@router.get(

View file

@ -19,7 +19,7 @@ Quick summary:
import json
from collections.abc import Iterable
from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn, Union
from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn
from fastapi import HTTPException
from pydantic import BaseModel
@ -61,7 +61,7 @@ if TYPE_CHECKING:
from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
from litellm.router import Router as _Router
Span = Union[_Span, Any]
Span = _Span | Any
InternalUsageCache = _InternalUsageCache
Router = _Router
ParallelRequestLimiter = _ParallelRequestLimiter

View file

@ -53,7 +53,7 @@ class _PROXY_BatchRedisRequests(CustomLogger):
key_value_dict = {}
in_memory_cache_exists = False
for key in cache.in_memory_cache.cache_dict.keys():
for key in cache.in_memory_cache.cache_dict:
if isinstance(key, str) and key.startswith(cache_key_name):
in_memory_cache_exists = True

View file

@ -170,7 +170,7 @@ class SkillsInjectionHook(CustomLogger):
skill_files = self.prompt_handler.extract_all_files(skill)
if skill_files:
all_skill_files[skill.skill_id] = skill_files
for path in skill_files.keys():
for path in skill_files:
if path.endswith(".py"):
all_module_paths.append(path)
@ -238,7 +238,7 @@ class SkillsInjectionHook(CustomLogger):
if skill_files:
all_skill_files[skill.skill_id] = skill_files
# Collect Python module paths
for path in skill_files.keys():
for path in skill_files:
if path.endswith(".py"):
all_module_paths.append(path)

View file

@ -1,7 +1,7 @@
import asyncio
import sys
from datetime import datetime, timedelta
from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn, Union
from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn
from pydantic import BaseModel
from typing_extensions import TypedDict
@ -26,7 +26,7 @@ if TYPE_CHECKING:
from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
Span = Union[_Span, Any]
Span = _Span | Any
InternalUsageCache = _InternalUsageCache
else:
Span = Any

View file

@ -12,7 +12,7 @@ from collections.abc import Callable
from contextvars import ContextVar
from dataclasses import dataclass, field
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, Union, cast
from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, cast
from litellm import DualCache
from litellm._logging import verbose_proxy_logger
@ -49,7 +49,7 @@ if TYPE_CHECKING:
from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
from litellm.types.caching import RedisPipelineIncrementOperation
Span = Union[_Span, Any]
Span = _Span | Any
InternalUsageCache = _InternalUsageCache
else:
Span = Any

View file

@ -249,7 +249,7 @@ def _redact_settings(settings: Mapping[str, object] | None) -> dict[str, object]
"""
if not settings:
return {}
return {k: _REDACTED_VALUE for k in settings.keys()}
return {k: _REDACTED_VALUE for k in settings}
def _log_audit_task_exception(task: "asyncio.Task[None]") -> None:

View file

@ -2,7 +2,7 @@ import asyncio
from collections.abc import Awaitable, Callable, Mapping, Sequence
from datetime import datetime
from types import SimpleNamespace
from typing import TYPE_CHECKING, Final, Protocol, Union
from typing import TYPE_CHECKING, Final, Protocol
from fastapi import HTTPException, status
from typing_extensions import TypedDict
@ -109,7 +109,7 @@ class _KeyMetadataDict(TypedDict, total=False):
team_id: str | None
_WhereValue = Union[str, dict[str, object]]
_WhereValue = str | dict[str, object]
class _AggregatedSpendData(TypedDict):

View file

@ -54,7 +54,7 @@ def _redact_config(config: Mapping[str, Any] | None) -> dict[str, Any]:
"""
if not config:
return {}
return {k: _AUDIT_REDACTED for k in config.keys()}
return {k: _AUDIT_REDACTED for k in config}
def _log_audit_task_exception(task: "asyncio.Task[None]") -> None:

View file

@ -365,7 +365,7 @@ async def new_end_user(
_user_data: Final = data.dict(exclude_none=True)
for k, v in _user_data.items():
if k not in BudgetNewRequest.model_fields.keys():
if k not in BudgetNewRequest.model_fields:
new_end_user_obj[k] = v
## Handle Object Permission - MCP Servers, Vector Stores etc.
@ -573,10 +573,10 @@ async def update_end_user(
# budget_id is for linking to existing budget, not for creating new budget
if k == "budget_id":
update_end_user_table_data[k] = v
elif k in LiteLLM_BudgetTable.model_fields.keys():
elif k in LiteLLM_BudgetTable.model_fields:
budget_table_data[k] = v
elif k in LiteLLM_EndUserTable.model_fields.keys():
elif k in LiteLLM_EndUserTable.model_fields:
update_end_user_table_data[k] = v
## Handle object permission updates (MCP servers, vector stores, etc.)

View file

@ -584,7 +584,7 @@ async def new_user(
special_keys: Final = ["token", "token_id"]
response_dict: Final = {}
for key, value in response.items():
if key in NewUserResponse.model_fields.keys() and key not in special_keys:
if key in NewUserResponse.model_fields and key not in special_keys:
response_dict[key] = value
response_dict["key"] = response.get("token", "")

View file

@ -215,7 +215,7 @@ async def _check_custom_key_allowed(custom_key_value: str | None) -> None:
)
def _is_team_key(data: Union[GenerateKeyRequest, LiteLLM_VerificationToken]):
def _is_team_key(data: GenerateKeyRequest | LiteLLM_VerificationToken):
return data.team_id is not None
@ -498,7 +498,7 @@ def key_generation_check(
def common_key_access_checks(
user_api_key_dict: UserAPIKeyAuth,
data: Union[GenerateKeyRequest, UpdateKeyRequest],
data: GenerateKeyRequest | UpdateKeyRequest,
llm_router: Router | None,
premium_user: bool,
user_id: str | None = None,
@ -752,7 +752,7 @@ _BUDGET_NUMERIC_KEYS = frozenset(["max_budget", "soft_budget", "max_parallel_req
def _enforce_upperbound_key_params(
data: Union[GenerateKeyRequest, UpdateKeyRequest],
data: GenerateKeyRequest | UpdateKeyRequest,
fill_defaults: bool = True,
) -> None:
"""
@ -1161,7 +1161,7 @@ async def _common_key_generation_helper(
def _check_key_model_specific_limits(
keys: list[LiteLLM_VerificationToken],
data: Union[GenerateKeyRequest, UpdateKeyRequest],
data: GenerateKeyRequest | UpdateKeyRequest,
entity_rpm_limit: int | None,
entity_tpm_limit: int | None,
entity_model_rpm_limit_dict: dict[str, int],
@ -1232,7 +1232,7 @@ def _check_key_model_specific_limits(
def _check_key_rpm_tpm_limits(
keys: list[LiteLLM_VerificationToken],
data: Union[GenerateKeyRequest, UpdateKeyRequest],
data: GenerateKeyRequest | UpdateKeyRequest,
entity_rpm_limit: int | None,
entity_tpm_limit: int | None,
entity_type: str, # "team" or "organization"
@ -1271,7 +1271,7 @@ def _check_key_rpm_tpm_limits(
def check_team_key_model_specific_limits(
keys: list[LiteLLM_VerificationToken],
team_table: LiteLLM_TeamTableCachedObj,
data: Union[GenerateKeyRequest, UpdateKeyRequest],
data: GenerateKeyRequest | UpdateKeyRequest,
) -> None:
"""
Check if the team key is allocating model specific limits. If so, raise an error if we're overallocating.
@ -1296,7 +1296,7 @@ def check_team_key_model_specific_limits(
def check_team_key_rpm_tpm_limits(
keys: list[LiteLLM_VerificationToken],
team_table: LiteLLM_TeamTableCachedObj,
data: Union[GenerateKeyRequest, UpdateKeyRequest],
data: GenerateKeyRequest | UpdateKeyRequest,
) -> None:
"""
Check if the team key is allocating rpm/tpm limits. If so, raise an error if we're overallocating.
@ -1312,7 +1312,7 @@ def check_team_key_rpm_tpm_limits(
async def _check_team_key_limits(
team_table: LiteLLM_TeamTableCachedObj,
data: Union[GenerateKeyRequest, UpdateKeyRequest],
data: GenerateKeyRequest | UpdateKeyRequest,
prisma_client: PrismaClient,
) -> None:
"""
@ -1348,7 +1348,7 @@ async def _check_team_key_limits(
async def _check_project_key_limits(
project_id: str,
data: Union[GenerateKeyRequest, UpdateKeyRequest],
data: GenerateKeyRequest | UpdateKeyRequest,
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
) -> None:
@ -1398,7 +1398,7 @@ async def _check_project_key_limits(
def check_org_key_model_specific_limits(
keys: list[LiteLLM_VerificationToken],
org_table: LiteLLM_OrganizationTable,
data: Union[GenerateKeyRequest, UpdateKeyRequest],
data: GenerateKeyRequest | UpdateKeyRequest,
) -> None:
"""
Check if the organization key is allocating model specific limits. If so, raise an error if we're overallocating.
@ -1431,7 +1431,7 @@ def check_org_key_model_specific_limits(
def check_org_key_rpm_tpm_limits(
keys: list[LiteLLM_VerificationToken],
org_table: LiteLLM_OrganizationTable,
data: Union[GenerateKeyRequest, UpdateKeyRequest],
data: GenerateKeyRequest | UpdateKeyRequest,
) -> None:
"""
Check if the organization key is allocating rpm/tpm limits. If so, raise an error if we're overallocating.
@ -1487,7 +1487,7 @@ async def _validate_caller_can_assign_key_org(
async def _check_org_key_limits(
org_table: LiteLLM_OrganizationTable,
data: Union[GenerateKeyRequest, UpdateKeyRequest],
data: GenerateKeyRequest | UpdateKeyRequest,
prisma_client: PrismaClient,
) -> None:
"""
@ -1944,7 +1944,7 @@ def prepare_metadata_fields(data: BaseModel, non_default_values: dict, existing_
async def prepare_key_update_data(
data: Union[UpdateKeyRequest, RegenerateKeyRequest],
data: UpdateKeyRequest | RegenerateKeyRequest,
existing_key_row: LiteLLM_VerificationToken,
):
data_json: Final[dict] = data.model_dump(exclude_unset=True)
@ -5672,7 +5672,7 @@ def _build_key_filter_conditions(
agent_id: str | None = None,
use_substring_matching: bool = False,
expires_filter: str | None = None,
) -> dict[str, Union[str, dict[str, Any], list[dict[str, Any]]]]:
) -> dict[str, str | dict[str, Any] | list[dict[str, Any]]]:
"""Build filter conditions for key listing.
Visibility rules:
@ -5684,7 +5684,7 @@ def _build_key_filter_conditions(
so former members cannot see service accounts they created after leaving.
"""
# Prepare filter conditions
where: dict[str, Union[str, dict[str, Any], list[dict[str, Any]]]] = {}
where: dict[str, str | dict[str, Any] | list[dict[str, Any]]] = {}
where.update(_get_condition_to_filter_out_ui_session_tokens())
# Build the OR conditions for user's keys and admin team keys
@ -5918,7 +5918,7 @@ async def _list_key_helper(
user_map = {user.user_id: user for user in users}
# Prepare response
key_list: Final[list[Union[str, UserAPIKeyAuth, LiteLLM_DeletedVerificationToken]]] = []
key_list: Final[list[str | UserAPIKeyAuth | LiteLLM_DeletedVerificationToken]] = []
for key in keys:
# Convert Prisma model to dict (supports both Pydantic v1 and v2)
try:

View file

@ -696,7 +696,7 @@ async def update_organization(
# Handle budget updates if budget fields are provided
budget_fields: Final = {
k: v for k, v in data.model_dump().items() if k in LiteLLM_BudgetTable.model_fields.keys() and v is not None
k: v for k, v in data.model_dump().items() if k in LiteLLM_BudgetTable.model_fields and v is not None
}
if budget_fields and existing_organization_row.budget_id:
@ -706,7 +706,7 @@ async def update_organization(
)
# Remove budget fields from organization update data
for field in LiteLLM_BudgetTable.model_fields.keys():
for field in LiteLLM_BudgetTable.model_fields:
updated_organization_row.pop(field, None)
response: Final = await _table(OrganizationRepository(prisma_client)).update(

View file

@ -7,7 +7,7 @@ Handles guardrail execution for passthrough endpoints with:
- Automatic inheritance from org/team/key levels when enabled
"""
from typing import Any, Final, Union
from typing import Any, Final
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import (
@ -19,10 +19,10 @@ from litellm.proxy.pass_through_endpoints.jsonpath_extractor import JsonPathExtr
# Type for raw guardrails config input (before normalization)
# Can be a list of names or a dict with settings
PassThroughGuardrailsConfigInput = Union[
list[str], # Simple list: ["guard-1", "guard-2"]
PassThroughGuardrailsConfig, # Dict: {"guard-1": {"request_fields": [...]}}
]
PassThroughGuardrailsConfigInput = (
list[str] # Simple list: ["guard-1", "guard-2"]
| PassThroughGuardrailsConfig # Dict: {"guard-1": {"request_fields": [...]}}
)
class PassthroughGuardrailHandler:
@ -246,7 +246,7 @@ class PassthroughGuardrailHandler:
guardrails_to_run: Final[dict[str, bool]] = {}
# Add passthrough-specific guardrails
for guardrail_name in normalized_config.keys():
for guardrail_name in normalized_config:
guardrails_to_run[guardrail_name] = True
verbose_proxy_logger.debug("Added passthrough-specific guardrail: %s", guardrail_name)

View file

@ -47,14 +47,12 @@ from litellm.proxy.policy_engine.policy_resolver import PolicyResolver
from litellm.proxy.policy_engine.policy_validator import PolicyValidator
__all__ = [
# Registries
"PolicyRegistry",
"get_policy_registry",
"AttachmentRegistry",
"get_attachment_registry",
# Core components
"ConditionEvaluator",
"PolicyMatcher",
"PolicyRegistry",
"PolicyResolver",
"PolicyValidator",
"ConditionEvaluator",
"get_attachment_registry",
"get_policy_registry",
]

View file

@ -180,7 +180,7 @@ class InMemoryPromptRegistry:
from litellm.proxy.prompts.prompt_endpoints import get_base_prompt_id
prompts_to_delete: Final = [
pid for pid in self.IN_MEMORY_PROMPTS.keys() if get_base_prompt_id(prompt_id=pid) == base_prompt_id
pid for pid in self.IN_MEMORY_PROMPTS if get_base_prompt_id(prompt_id=pid) == base_prompt_id
]
for pid in prompts_to_delete:

View file

@ -130,7 +130,7 @@ if TYPE_CHECKING:
from litellm.integrations.opentelemetry import OpenTelemetry
Span = Union[_Span, Any]
Span = _Span | Any
else:
Span = Any
OpenTelemetry = Any
@ -640,7 +640,6 @@ except Exception:
version = "0.0.0"
litellm.suppress_debug_info = True
import json
from typing import Union
from fastapi import (
Depends,
@ -6679,7 +6678,7 @@ class ProxyConfig:
await evict_config_param("anthropic_beta_headers_reload_config")
# Count providers in config
provider_count = sum(1 for k in new_config.keys() if k != "provider_aliases" and k != "description")
provider_count = sum(1 for k in new_config if k != "provider_aliases" and k != "description")
verbose_proxy_logger.info(
"Anthropic beta headers config reloaded successfully. Providers: %s", provider_count
)
@ -15196,7 +15195,7 @@ async def get_config_general_settings(
)
GeneralSettingsUILiteLLMValue = Union[float, bool, str, None]
GeneralSettingsUILiteLLMValue = float | bool | str | None
class GeneralSettingsUILiteLLMFieldSpec(TypedDict):
@ -16130,7 +16129,7 @@ async def reload_anthropic_beta_headers(
)
await invalidate_config_param("anthropic_beta_headers_reload_config")
provider_count: Final = sum(1 for k in new_config.keys() if k not in ["provider_aliases", "description"])
provider_count: Final = sum(1 for k in new_config if k not in ["provider_aliases", "description"])
verbose_proxy_logger.info(
"Anthropic beta headers config reloaded successfully in current pod. Providers: %s", provider_count
)

View file

@ -123,9 +123,7 @@ def _get_spend_logs_metadata(
)
# Filter the metadata dictionary to include only the specified keys
clean_metadata: Final = SpendLogsMetadata(
**{key: metadata.get(key) for key in SpendLogsMetadata.__annotations__.keys()}
)
clean_metadata: Final = SpendLogsMetadata(**{key: metadata.get(key) for key in SpendLogsMetadata.__annotations__})
raw_user_api_key: Final = clean_metadata.get("user_api_key")
if raw_user_api_key is not None and isinstance(raw_user_api_key, str):
clean_metadata["user_api_key"] = _hash_api_key_for_spend_log(raw_user_api_key)

View file

@ -167,7 +167,7 @@ if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy.db.spend_log_tool_index import ToolUsageTransaction
Span = Union[_Span, Any]
Span = _Span | Any
else:
Span = Any

View file

@ -4,7 +4,7 @@ Base repository class with common functionality.
from abc import ABC, abstractmethod
from collections.abc import Iterable, Mapping, Sequence
from typing import Any, Final, Generic, Protocol, TypeVar, Union, runtime_checkable
from typing import Any, Final, Generic, Protocol, TypeVar, runtime_checkable
from pydantic import BaseModel
@ -21,12 +21,7 @@ class SupportsDict(Protocol):
def dict(self) -> dict[str, object]: ...
DbRecord = Union[
Mapping[str, object],
SupportsModelDump,
SupportsDict,
Sequence[tuple[str, object]],
]
DbRecord = Mapping[str, object] | SupportsModelDump | SupportsDict | Sequence[tuple[str, object]]
def record_to_dict(record: DbRecord) -> Mapping[str, object]:

View file

@ -30,7 +30,6 @@ from openai import AsyncOpenAI
from typing_extensions import overload
import litellm
import litellm.litellm_core_utils
import litellm.litellm_core_utils.exception_mapping_utils
from litellm import get_secret_str
from litellm._logging import verbose_router_logger
@ -241,7 +240,7 @@ if TYPE_CHECKING:
ResponsesAPIResponse,
)
Span = Union[_Span, Any]
Span = _Span | Any
else:
Span = Any
AutoRouter = Any
@ -3410,9 +3409,7 @@ class Router:
# Await the first task to complete successfully
while pending_tasks:
done, pending_tasks = await asyncio.wait(
pending_tasks, return_when=asyncio.FIRST_COMPLETED
)
done, pending_tasks = await asyncio.wait(pending_tasks, return_when=asyncio.FIRST_COMPLETED)
for completed_task in done:
result = await check_response(completed_task)
@ -5240,9 +5237,7 @@ class Router:
# 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
)
_, 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),
@ -6029,9 +6024,7 @@ class Router:
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
)
return await original_function(custom_llm_provider=custom_llm_provider, client=client, **kwargs)
#### [END] ASSISTANTS API ####
@ -6359,14 +6352,9 @@ class Router:
if hasattr(original_exception, "message") and litellm.expose_router_debug_in_errors:
# add the available fallbacks to the exception
original_exception.message += ". Received Model Group={}\nAvailable Model Group Fallbacks={}".format(
model_group,
mask_sensitive_structure(fallback_model_group),
)
original_exception.message += f". Received Model Group={model_group}\nAvailable Model Group Fallbacks={mask_sensitive_structure(fallback_model_group)}"
if len(fallback_failure_exception_str) > 0:
original_exception.message += (
f"\nError doing the fallback: {fallback_failure_exception_str}"
)
original_exception.message += f"\nError doing the fallback: {fallback_failure_exception_str}"
raise original_exception
@ -7497,7 +7485,7 @@ class Router:
litellm_params=litellm_params,
model_info=_model_info,
)
for field in CustomPricingLiteLLMParams.model_fields.keys():
for field in CustomPricingLiteLLMParams.model_fields:
if deployment.litellm_params.get(field) is not None:
_model_info[field] = deployment.litellm_params[field]
@ -8246,7 +8234,7 @@ class Router:
self._add_deployment(deployment=deployment)
_model_info_dict: Final[dict] = deployment.model_info.model_dump(exclude_none=True)
for field in CustomPricingLiteLLMParams.model_fields.keys():
for field in CustomPricingLiteLLMParams.model_fields:
field_value = deployment.litellm_params.get(field)
if field_value is not None:
_model_info_dict[field] = field_value
@ -9124,9 +9112,7 @@ class Router:
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
):
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
@ -9144,9 +9130,7 @@ class Router:
):
model_group_info.supports_url_context = True
if (
model_info.get("supports_reasoning", None) is not None and model_info["supports_reasoning"] is 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
@ -9495,7 +9479,7 @@ class Router:
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.keys():
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"]:
@ -10876,9 +10860,7 @@ class Router:
args=(e, traceback_exception),
).start() # log response
# Handle any exceptions that might occur during streaming
asyncio.create_task(
logging_obj.async_failure_handler(e, traceback_exception)
)
asyncio.create_task(logging_obj.async_failure_handler(e, traceback_exception))
raise e
async def async_get_available_deployment_for_pass_through(
@ -11003,9 +10985,7 @@ class Router:
target=logging_obj.failure_handler,
args=(e, traceback_exception),
).start()
asyncio.create_task(
logging_obj.async_failure_handler(e, traceback_exception)
)
asyncio.create_task(logging_obj.async_failure_handler(e, traceback_exception))
raise e
async def _run_routing_plugins(

View file

@ -2,7 +2,7 @@
# picks based on response time (for streaming, this is time to first token)
import random
from datetime import datetime, timedelta
from typing import TYPE_CHECKING, Any, Final, Union
from typing import TYPE_CHECKING, Any, Final
import litellm
from litellm import ModelResponse, token_counter, verbose_logger
@ -14,7 +14,7 @@ from litellm.types.utils import LiteLLMPydanticObjectBase
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
Span = Union[_Span, Any]
Span = _Span | Any
else:
Span = Any

View file

@ -1,7 +1,7 @@
#### What this does ####
# identifies lowest tpm deployment
import random
from typing import TYPE_CHECKING, Any, Final, Union
from typing import TYPE_CHECKING, Any, Final
import httpx
@ -20,7 +20,7 @@ from .base_routing_strategy import BaseRoutingStrategy
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
Span = Union[_Span, Any]
Span = _Span | Any
else:
Span = Any

View file

@ -4,7 +4,7 @@ Wrapper around router cache. Meant to handle model cooldown logic
import functools
import time
from typing import TYPE_CHECKING, Any, Final, Union
from typing import TYPE_CHECKING, Any, Final
from typing_extensions import TypedDict
@ -16,7 +16,7 @@ from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
Span = Union[_Span, Any]
Span = _Span | Any
else:
Span = Any

View file

@ -8,7 +8,7 @@ Router cooldown handlers
import asyncio
import math
from typing import TYPE_CHECKING, Any, Final, Union
from typing import TYPE_CHECKING, Any, Final
import litellm
from litellm._logging import verbose_router_logger
@ -31,7 +31,7 @@ if TYPE_CHECKING:
from litellm.router import Router as _Router
LitellmRouter = _Router
Span = Union[_Span, Any]
Span = _Span | Any
else:
LitellmRouter = Any
Span = Any

View file

@ -237,7 +237,7 @@ def _check_non_standard_fallback_format(fallbacks: list[Any] | None) -> bool:
return True
elif all(isinstance(item, dict) for item in fallbacks):
for item in fallbacks:
for key in LiteLLMParamsTypedDict.__annotations__.keys():
for key in LiteLLMParamsTypedDict.__annotations__:
if key in item:
# If the value is a list, it's likely a standard fallback model group mapping
# (e.g. {"model": ["backup"]}) rather than a parameter override.

View file

@ -1,4 +1,4 @@
from typing import TYPE_CHECKING, Any, Final, Union
from typing import TYPE_CHECKING, Any, Final
from litellm._logging import redact_secrets, verbose_router_logger
from litellm.constants import MAX_EXCEPTION_MESSAGE_LENGTH
@ -14,7 +14,7 @@ if TYPE_CHECKING:
from litellm.router import Router as _Router
LitellmRouter = _Router
Span = Union[_Span, Any]
Span = _Span | Any
else:
LitellmRouter = Any
Span = Any

View file

@ -6,7 +6,7 @@ and exposes it for router candidate filtering.
"""
import time
from typing import TYPE_CHECKING, Any, Final, Union
from typing import TYPE_CHECKING, Any, Final
from typing_extensions import TypedDict
@ -16,7 +16,7 @@ from litellm.caching.caching import DualCache
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
Span = Union[_Span, Any]
Span = _Span | Any
else:
Span = Any

View file

@ -11,7 +11,7 @@ is logged the first time such a deployment is seen.
"""
import contextlib
from typing import TYPE_CHECKING, Any, Final, Union
from typing import TYPE_CHECKING, Any, Final
import httpx
@ -37,7 +37,7 @@ from litellm.utils import get_utc_datetime
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
Span = Union[_Span, Any]
Span = _Span | Any
else:
Span = Any

View file

@ -4,7 +4,7 @@ Wrapper around router cache. Meant to store model id when prompt caching support
import hashlib
import json
from typing import TYPE_CHECKING, Any, Final, Union, cast
from typing import TYPE_CHECKING, Any, Final, cast
from typing_extensions import TypedDict
@ -18,7 +18,7 @@ if TYPE_CHECKING:
from litellm.router import Router
litellm_router = Router
Span = Union[_Span, Any]
Span = _Span | Any
else:
Span = Any
litellm_router = Any

View file

@ -1,39 +1,38 @@
from datetime import datetime
from typing import List, Optional
from pydantic import BaseModel
class AccessGroupCreateRequest(BaseModel):
access_group_name: str
description: Optional[str] = None
access_model_names: Optional[List[str]] = None
access_mcp_server_ids: Optional[List[str]] = None
access_agent_ids: Optional[List[str]] = None
assigned_team_ids: Optional[List[str]] = None
assigned_key_ids: Optional[List[str]] = None
description: str | None = None
access_model_names: list[str] | None = None
access_mcp_server_ids: list[str] | None = None
access_agent_ids: list[str] | None = None
assigned_team_ids: list[str] | None = None
assigned_key_ids: list[str] | None = None
class AccessGroupUpdateRequest(BaseModel):
access_group_name: Optional[str] = None
description: Optional[str] = None
access_model_names: Optional[List[str]] = None
access_mcp_server_ids: Optional[List[str]] = None
access_agent_ids: Optional[List[str]] = None
assigned_team_ids: Optional[List[str]] = None
assigned_key_ids: Optional[List[str]] = None
access_group_name: str | None = None
description: str | None = None
access_model_names: list[str] | None = None
access_mcp_server_ids: list[str] | None = None
access_agent_ids: list[str] | None = None
assigned_team_ids: list[str] | None = None
assigned_key_ids: list[str] | None = None
class AccessGroupResponse(BaseModel):
access_group_id: str
access_group_name: str
description: Optional[str] = None
access_model_names: List[str]
access_mcp_server_ids: List[str]
access_agent_ids: List[str]
assigned_team_ids: List[str]
assigned_key_ids: List[str]
description: str | None = None
access_model_names: list[str]
access_mcp_server_ids: list[str]
access_agent_ids: list[str]
assigned_team_ids: list[str]
assigned_key_ids: list[str]
created_at: datetime
created_by: Optional[str] = None
created_by: str | None = None
updated_at: datetime
updated_by: Optional[str] = None
updated_by: str | None = None

View file

@ -1,5 +1,5 @@
from datetime import datetime
from typing import Any, Dict, Final, List, Literal, Optional, TYPE_CHECKING, Union
from typing import TYPE_CHECKING, Any, Final, Literal
from pydantic import BaseModel, PrivateAttr
from typing_extensions import Required, TypedDict
@ -23,26 +23,26 @@ class AgentExtension(TypedDict, total=False):
"""A declaration of a protocol extension supported by an Agent."""
uri: str # required
description: Optional[str]
required: Optional[bool]
params: Optional[Dict[str, Any]]
description: str | None
required: bool | None
params: dict[str, Any] | None
# AgentCapabilities
class AgentCapabilities(TypedDict, total=False):
"""Defines optional capabilities supported by an agent."""
streaming: Optional[bool]
pushNotifications: Optional[bool]
stateTransitionHistory: Optional[bool]
extensions: Optional[List[AgentExtension]]
streaming: bool | None
pushNotifications: bool | None
stateTransitionHistory: bool | None
extensions: list[AgentExtension] | None
# SecurityScheme types
class SecuritySchemeBase(TypedDict, total=False):
"""Base properties shared by all security scheme objects."""
description: Optional[str]
description: str | None
class APIKeySecurityScheme(SecuritySchemeBase, total=False):
@ -58,7 +58,7 @@ class HTTPAuthSecurityScheme(SecuritySchemeBase, total=False):
type: Required[Literal["http"]]
scheme: Required[str]
bearerFormat: Optional[str]
bearerFormat: str | None
class MutualTLSSecurityScheme(SecuritySchemeBase, total=False):
@ -70,10 +70,10 @@ class MutualTLSSecurityScheme(SecuritySchemeBase, total=False):
class OAuthFlows(TypedDict, total=False):
"""Defines the configuration for the supported OAuth 2.0 flows."""
authorizationCode: Optional[Dict[str, Any]]
clientCredentials: Optional[Dict[str, Any]]
implicit: Optional[Dict[str, Any]]
password: Optional[Dict[str, Any]]
authorizationCode: dict[str, Any] | None
clientCredentials: dict[str, Any] | None
implicit: dict[str, Any] | None
password: dict[str, Any] | None
class OAuth2SecurityScheme(SecuritySchemeBase, total=False):
@ -81,7 +81,7 @@ class OAuth2SecurityScheme(SecuritySchemeBase, total=False):
type: Required[Literal["oauth2"]]
flows: Required[OAuthFlows]
oauth2MetadataUrl: Optional[str]
oauth2MetadataUrl: str | None
class OpenIdConnectSecurityScheme(SecuritySchemeBase, total=False):
@ -92,13 +92,13 @@ class OpenIdConnectSecurityScheme(SecuritySchemeBase, total=False):
# Union of all security schemes
SecurityScheme = Union[
APIKeySecurityScheme,
HTTPAuthSecurityScheme,
OAuth2SecurityScheme,
OpenIdConnectSecurityScheme,
MutualTLSSecurityScheme,
]
SecurityScheme = (
APIKeySecurityScheme
| HTTPAuthSecurityScheme
| OAuth2SecurityScheme
| OpenIdConnectSecurityScheme
| MutualTLSSecurityScheme
)
# AgentSkill
@ -108,11 +108,11 @@ class AgentSkill(TypedDict, total=False):
id: str # required
name: str # required
description: str # required
tags: List[str] # required
examples: Optional[List[str]]
inputModes: Optional[List[str]]
outputModes: Optional[List[str]]
security: Optional[List[Dict[str, List[str]]]]
tags: list[str] # required
examples: list[str] | None
inputModes: list[str] | None
outputModes: list[str] | None
security: list[dict[str, list[str]]] | None
# AgentInterface
@ -129,7 +129,7 @@ class AgentCardSignature(TypedDict, total=False):
protected: str # required
signature: str # required
header: Optional[Dict[str, Any]]
header: dict[str, Any] | None
# AgentCard
@ -147,20 +147,20 @@ class AgentCard(TypedDict, total=False):
url: str
version: str
capabilities: AgentCapabilities
defaultInputModes: List[str]
defaultOutputModes: List[str]
skills: List[AgentSkill]
defaultInputModes: list[str]
defaultOutputModes: list[str]
skills: list[AgentSkill]
# Optional fields
preferredTransport: Optional[str]
additionalInterfaces: Optional[List[AgentInterface]]
iconUrl: Optional[str]
provider: Optional[AgentProvider]
documentationUrl: Optional[str]
securitySchemes: Optional[Dict[str, SecurityScheme]]
security: Optional[List[Dict[str, List[str]]]]
supportsAuthenticatedExtendedCard: Optional[bool]
signatures: Optional[List[AgentCardSignature]]
preferredTransport: str | None
additionalInterfaces: list[AgentInterface] | None
iconUrl: str | None
provider: AgentProvider | None
documentationUrl: str | None
securitySchemes: dict[str, SecurityScheme] | None
security: list[dict[str, list[str]]] | None
supportsAuthenticatedExtendedCard: bool | None
signatures: list[AgentCardSignature] | None
class AugmentedAgentCard(AgentCard):
@ -169,37 +169,37 @@ class AugmentedAgentCard(AgentCard):
# Object permission shape for agent MCP tool access (mirrors LiteLLM_ObjectPermissionBase)
class AgentObjectPermission(TypedDict, total=False):
mcp_servers: Optional[List[str]]
mcp_access_groups: Optional[List[str]]
mcp_tool_permissions: Optional[Dict[str, List[str]]]
models: Optional[List[str]]
agents: Optional[List[str]]
mcp_servers: list[str] | None
mcp_access_groups: list[str] | None
mcp_tool_permissions: dict[str, list[str]] | None
models: list[str] | None
agents: list[str] | None
class AgentConfig(TypedDict, total=False):
agent_name: Required[str]
agent_card_params: Required[AgentCard]
litellm_params: Dict[str, Any] # allow for any future litellm params
litellm_params: dict[str, Any] # allow for any future litellm params
object_permission: AgentObjectPermission
tpm_limit: Optional[int]
rpm_limit: Optional[int]
session_tpm_limit: Optional[int]
session_rpm_limit: Optional[int]
static_headers: Optional[Dict[str, str]]
extra_headers: Optional[List[str]]
tpm_limit: int | None
rpm_limit: int | None
session_tpm_limit: int | None
session_rpm_limit: int | None
static_headers: dict[str, str] | None
extra_headers: list[str] | None
class PatchAgentRequest(TypedDict, total=False):
agent_name: str
agent_card_params: AgentCard
litellm_params: Dict[str, Any]
litellm_params: dict[str, Any]
object_permission: AgentObjectPermission
tpm_limit: Optional[int]
rpm_limit: Optional[int]
session_tpm_limit: Optional[int]
session_rpm_limit: Optional[int]
static_headers: Optional[Dict[str, str]]
extra_headers: Optional[List[str]]
tpm_limit: int | None
rpm_limit: int | None
session_tpm_limit: int | None
session_rpm_limit: int | None
static_headers: dict[str, str] | None
extra_headers: list[str] | None
# Request/Response models for CRUD endpoints
@ -207,32 +207,32 @@ class PatchAgentRequest(TypedDict, total=False):
class AgentKeySummary(BaseModel):
token: str
key_alias: Optional[str] = None
key_name: Optional[str] = None
key_alias: str | None = None
key_name: str | None = None
class AgentResponse(BaseModel):
agent_id: str
agent_name: str
litellm_params: Optional[Dict[str, Any]] = None
agent_card_params: Dict[str, Any]
object_permission: Optional[Dict[str, Any]] = None
spend: Optional[float] = None
tpm_limit: Optional[int] = None
rpm_limit: Optional[int] = None
session_tpm_limit: Optional[int] = None
session_rpm_limit: Optional[int] = None
static_headers: Optional[Dict[str, str]] = None
extra_headers: Optional[List[str]] = None
keys: Optional[List[AgentKeySummary]] = None
created_at: Optional[datetime] = None
updated_at: Optional[datetime] = None
created_by: Optional[str] = None
updated_by: Optional[str] = None
litellm_params: dict[str, Any] | None = None
agent_card_params: dict[str, Any]
object_permission: dict[str, Any] | None = None
spend: float | None = None
tpm_limit: int | None = None
rpm_limit: int | None = None
session_tpm_limit: int | None = None
session_rpm_limit: int | None = None
static_headers: dict[str, str] | None = None
extra_headers: list[str] | None = None
keys: list[AgentKeySummary] | None = None
created_at: datetime | None = None
updated_at: datetime | None = None
created_by: str | None = None
updated_by: str | None = None
class ListAgentsResponse(BaseModel):
agents: List[AgentResponse]
agents: list[AgentResponse]
class AgentCreateResponse(LiteLLMPydanticObjectBase):
@ -246,8 +246,8 @@ class AgentCreateResponse(LiteLLMPydanticObjectBase):
are preserved via extra="allow".
"""
id: Optional[str] = None
name: Optional[str] = None
id: str | None = None
name: str | None = None
model_config = {"extra": "allow"}
_hidden_params: dict = PrivateAttr(default_factory=dict)
@ -274,8 +274,8 @@ class AgentListResponse(LiteLLMPydanticObjectBase):
a plain dict so no fields are silently dropped.
"""
agents: List[Dict[str, Any]] = []
next_page_token: Optional[str] = None
agents: list[dict[str, Any]] = []
next_page_token: str | None = None
model_config = {"extra": "allow"}
_hidden_params: dict = PrivateAttr(default_factory=dict)
@ -288,8 +288,8 @@ class AgentVersionsResponse(LiteLLMPydanticObjectBase):
field of the form ``agents/{agent_id}/versions/{uuid}``.
"""
agent_versions: List[Dict[str, Any]] = []
next_page_token: Optional[str] = None
agent_versions: list[dict[str, Any]] = []
next_page_token: str | None = None
model_config = {"extra": "allow"}
_hidden_params: dict = PrivateAttr(default_factory=dict)
@ -297,18 +297,18 @@ class AgentVersionsResponse(LiteLLMPydanticObjectBase):
class AgentMakePublicResponse(BaseModel):
message: str
public_agent_groups: List[str]
public_agent_groups: list[str]
updated_by: str
class MakeAgentsPublicRequest(BaseModel):
agent_ids: List[str]
agent_ids: list[str]
def _normalize_a2a_jsonrpc_response(
response_dict: Dict[str, Any],
request_id: Optional[Any] = None,
) -> Dict[str, Any]:
response_dict: dict[str, Any],
request_id: Any | None = None,
) -> dict[str, Any]:
"""
Ensure JSON-RPC responses include ``id`` when the caller supplied one.
@ -333,11 +333,11 @@ class LiteLLMSendMessageResponse(LiteLLMPydanticObjectBase):
# A2A response fields
id: str
jsonrpc: str = "2.0"
result: Optional[Dict[str, Any]] = None
error: Optional[Dict[str, Any]] = None
result: dict[str, Any] | None = None
error: dict[str, Any] | None = None
# LiteLLM usage tracking
usage: Optional[Dict[str, Any]] = None
usage: dict[str, Any] | None = None
model_config = {"extra": "allow"}
@ -348,7 +348,7 @@ class LiteLLMSendMessageResponse(LiteLLMPydanticObjectBase):
def from_a2a_response(
cls,
response: "SendMessageResponse",
request_id: Optional[Any] = None,
request_id: Any | None = None,
) -> "LiteLLMSendMessageResponse":
"""
Create a LiteLLMSendMessageResponse from an a2a SDK SendMessageResponse.
@ -367,8 +367,8 @@ class LiteLLMSendMessageResponse(LiteLLMPydanticObjectBase):
@classmethod
def from_dict(
cls,
response_dict: Dict[str, Any],
request_id: Optional[Any] = None,
response_dict: dict[str, Any],
request_id: Any | None = None,
) -> "LiteLLMSendMessageResponse":
"""
Create a LiteLLMSendMessageResponse from a dict.

View file

@ -1,5 +1,5 @@
from enum import Enum
from typing import Any, Dict, Final, List, Literal, Optional, Union
from typing import Any, Final, Literal, Optional, Union
from pydantic import BaseModel
from typing_extensions import TypedDict
@ -40,7 +40,7 @@ class RedisPipelineIncrementOperation(TypedDict):
key: str
increment_value: float
ttl: Optional[int]
ttl: int | None
class RedisPipelineSetOperation(TypedDict):
@ -50,7 +50,7 @@ class RedisPipelineSetOperation(TypedDict):
key: str
value: Any
ttl: Optional[int]
ttl: int | None
class RedisPipelineRpushOperation(TypedDict):
@ -59,7 +59,7 @@ class RedisPipelineRpushOperation(TypedDict):
"""
key: str
values: List[Any]
values: list[Any]
class RedisPipelineLpopOperation(TypedDict):
@ -68,23 +68,23 @@ class RedisPipelineLpopOperation(TypedDict):
"""
key: str
count: Optional[int]
count: int | None
DynamicCacheControl = TypedDict(
"DynamicCacheControl",
{
# Will cache the response for the user-defined amount of time (in seconds).
"ttl": Optional[int],
"ttl": int | None,
# Namespace to use for caching
"namespace": Optional[str],
"namespace": str | None,
# Max Age to use for caching
"s-maxage": Optional[int],
"s-max-age": Optional[int],
"s-maxage": int | None,
"s-max-age": int | None,
# Will not return a cached response, but instead call the actual endpoint.
"no-cache": Optional[bool],
"no-cache": bool | None,
# Will not store the response in the cache.
"no-store": Optional[bool],
"no-store": bool | None,
},
)
@ -92,12 +92,12 @@ DynamicCacheControl = TypedDict(
class CachePingResponse(BaseModel):
status: str
cache_type: str
ping_response: Optional[bool] = None
set_cache_response: Optional[str] = None
litellm_cache_params: Optional[str] = None
ping_response: bool | None = None
set_cache_response: str | None = None
litellm_cache_params: str | None = None
# intentionally a dict, since we run masker.mask_dict() on HealthCheckCacheParams
health_check_cache_params: Optional[dict] = None
health_check_cache_params: dict | None = None
class HealthCheckCacheParams(BaseModel):
@ -105,19 +105,19 @@ class HealthCheckCacheParams(BaseModel):
Cache Params returned on /cache/ping call
"""
host: Optional[str] = None
port: Optional[Union[str, int]] = None
redis_kwargs: Optional[Dict[str, Any]] = None
namespace: Optional[str] = None
redis_version: Optional[Union[str, int, float]] = None
host: str | None = None
port: str | int | None = None
redis_kwargs: dict[str, Any] | None = None
namespace: str | None = None
redis_version: str | int | float | None = None
class CachedEmbedding(TypedDict):
"""Type definition for cached embedding objects"""
embedding: Optional[List[float]]
index: Optional[int]
object: Optional[str]
model: Optional[str]
prompt_tokens: Optional[int]
prompt_tokens_details: Optional[dict]
embedding: list[float] | None
index: int | None
object: str | None
model: str | None
prompt_tokens: int | None
prompt_tokens_details: dict | None

View file

@ -1,10 +1,11 @@
from __future__ import annotations
from collections.abc import Callable, Coroutine, Iterable
from dataclasses import dataclass
from typing import Any, Callable, Coroutine, Final, Iterable, List, Optional, TYPE_CHECKING, Union
from typing import TYPE_CHECKING, Any, Literal, Union
from pydantic import BaseModel, ConfigDict
from typing_extensions import Literal, Required, TypedDict
from typing_extensions import Required, TypedDict
if TYPE_CHECKING:
import httpx
@ -57,11 +58,11 @@ class ChatCompletionContentPartImageParam(TypedDict, total=False):
"""The type of the content part."""
ChatCompletionContentPartParam = Union[ChatCompletionContentPartTextParam, ChatCompletionContentPartImageParam]
ChatCompletionContentPartParam = ChatCompletionContentPartTextParam | ChatCompletionContentPartImageParam
class ChatCompletionUserMessageParam(TypedDict, total=False):
content: Required[Union[str, Iterable[ChatCompletionContentPartParam]]]
content: Required[str | Iterable[ChatCompletionContentPartParam]]
"""The contents of the user message."""
role: Required[Literal["user"]]
@ -102,7 +103,7 @@ class Function(TypedDict, total=False):
class ChatCompletionToolMessageParam(TypedDict, total=False):
content: Required[Union[str, Iterable[ChatCompletionContentPartParam]]]
content: Required[str | Iterable[ChatCompletionContentPartParam]]
"""The contents of the tool message."""
role: Required[Literal["tool"]]
@ -113,7 +114,7 @@ class ChatCompletionToolMessageParam(TypedDict, total=False):
class ChatCompletionFunctionMessageParam(TypedDict, total=False):
content: Required[Union[str, Iterable[ChatCompletionContentPartParam]]]
content: Required[str | Iterable[ChatCompletionContentPartParam]]
"""The contents of the function message."""
name: Required[str]
@ -138,7 +139,7 @@ class ChatCompletionAssistantMessageParam(TypedDict, total=False):
role: Required[Literal["assistant"]]
"""The role of the messages author, in this case `assistant`."""
content: Optional[str]
content: str | None
"""The contents of the assistant message.
Required unless `tool_calls` or `function_call` is specified.
@ -162,42 +163,42 @@ class ChatCompletionAssistantMessageParam(TypedDict, total=False):
"""The tool calls generated by the model, such as function calls."""
ChatCompletionMessageParam = Union[
ChatCompletionSystemMessageParam,
ChatCompletionUserMessageParam,
ChatCompletionAssistantMessageParam,
ChatCompletionFunctionMessageParam,
ChatCompletionToolMessageParam,
]
ChatCompletionMessageParam = (
ChatCompletionSystemMessageParam
| ChatCompletionUserMessageParam
| ChatCompletionAssistantMessageParam
| ChatCompletionFunctionMessageParam
| ChatCompletionToolMessageParam
)
class CompletionRequest(BaseModel):
model: str
messages: List[ChatCompletionMessageParam] = []
timeout: Optional[Union[float, int]] = None
temperature: Optional[float] = None
top_p: Optional[float] = None
n: Optional[int] = None
stream: Optional[bool] = None
stop: Optional[dict] = None
max_tokens: Optional[int] = None
presence_penalty: Optional[float] = None
frequency_penalty: Optional[float] = None
logit_bias: Optional[dict] = None
user: Optional[str] = None
response_format: Optional[dict] = None
seed: Optional[int] = None
tools: Optional[List[str]] = None
tool_choice: Optional[str] = None
logprobs: Optional[bool] = None
top_logprobs: Optional[int] = None
deployment_id: Optional[str] = None
functions: Optional[List[str]] = None
function_call: Optional[str] = None
base_url: Optional[str] = None
api_version: Optional[str] = None
api_key: Optional[str] = None
model_list: Optional[List[str]] = None
messages: list[ChatCompletionMessageParam] = []
timeout: float | int | None = None
temperature: float | None = None
top_p: float | None = None
n: int | None = None
stream: bool | None = None
stop: dict | None = None
max_tokens: int | None = None
presence_penalty: float | None = None
frequency_penalty: float | None = None
logit_bias: dict | None = None
user: str | None = None
response_format: dict | None = None
seed: int | None = None
tools: list[str] | None = None
tool_choice: str | None = None
logprobs: bool | None = None
top_logprobs: int | None = None
deployment_id: str | None = None
functions: list[str] | None = None
function_call: str | None = None
base_url: str | None = None
api_version: str | None = None
api_key: str | None = None
model_list: list[str] | None = None
model_config = ConfigDict(protected_namespaces=(), extra="allow")
@ -206,34 +207,34 @@ class CompletionRequest(BaseModel):
class _CompletionDispatchContext:
_azure_detection_model: str
acompletion: bool
api_base: Optional[str]
api_key: Optional[str]
api_version: Optional[str]
api_base: str | None
api_key: str | None
api_version: str | None
client: Any
custom_llm_provider: str
custom_prompt_dict: dict
extra_headers: Optional[dict]
extra_headers: dict | None
headers: dict
hf_model_name: Optional[str]
hf_model_name: str | None
kwargs: dict
litellm_params: dict
logger_fn: Optional[Callable]
logger_fn: Callable | None
logging: LiteLLMLoggingObj
max_retries: Optional[int]
max_tokens: Optional[int]
max_retries: int | None
max_tokens: int | None
messages: list
metadata: Optional[dict]
metadata: dict | None
model: str
model_response: ModelResponse
optional_params: dict
organization: Optional[str]
provider_config: Optional[BaseConfig]
shared_session: Optional[ClientSession]
stream: Optional[bool]
temperature: Optional[float]
organization: str | None
provider_config: BaseConfig | None
shared_session: ClientSession | None
stream: bool | None
temperature: float | None
text_completion: bool
timeout: Optional[Union[float, str, httpx.Timeout]]
top_p: Optional[float]
timeout: float | str | httpx.Timeout | None
top_p: float | None
_CompletionDispatchResult = Union[

View file

@ -5,18 +5,18 @@ Type definitions for litellm.compress().
import sys
if sys.version_info >= (3, 11):
from typing import Dict, List, NotRequired, TypedDict
from typing import NotRequired, TypedDict
else:
from typing import Dict, List, TypedDict
from typing import TypedDict
from typing_extensions import NotRequired
class CompressedResult(TypedDict):
messages: List[dict] # compressed messages (stubs replace low-relevance messages)
messages: list[dict] # compressed messages (stubs replace low-relevance messages)
original_tokens: int # token count before compression
compressed_tokens: int # token count after compression
compression_ratio: float # fraction reduced, e.g. 0.6 means 60% reduction
cache: Dict[str, str] # key -> original content (for retrieval tool responses)
tools: List[dict] # [litellm_content_retrieve tool definition]
cache: dict[str, str] # key -> original content (for retrieval tool responses)
tools: list[dict] # [litellm_content_retrieve tool definition]
compression_skipped_reason: NotRequired[str]

View file

@ -1,4 +1,4 @@
from typing import Any, Dict, List, Literal, Optional
from typing import Any, Literal
from pydantic import BaseModel
from typing_extensions import TypedDict
@ -18,12 +18,12 @@ class ContainerObject(BaseModel):
object: Literal["container"]
created_at: int
status: str
expires_after: Optional[ExpiresAfter] = None
last_active_at: Optional[int] = None
name: Optional[str] = None
_hidden_params: Dict[str, Any] = {}
expires_after: ExpiresAfter | None = None
last_active_at: int | None = None
name: str | None = None
_hidden_params: dict[str, Any] = {}
def __contains__(self, key):
def __contains__(self, key) -> bool:
# Define custom behavior for the 'in' operator
return hasattr(self, key)
@ -50,7 +50,7 @@ class DeleteContainerResult(BaseModel):
object: Literal["container.deleted"]
deleted: bool
def __contains__(self, key):
def __contains__(self, key) -> bool:
return hasattr(self, key)
def get(self, key, default=None):
@ -70,12 +70,12 @@ class ContainerListResponse(BaseModel):
"""Response object for list containers request."""
object: Literal["list"]
data: List[ContainerObject]
first_id: Optional[str] = None
last_id: Optional[str] = None
data: list[ContainerObject]
first_id: str | None = None
last_id: str | None = None
has_more: bool
def __contains__(self, key):
def __contains__(self, key) -> bool:
return hasattr(self, key)
def get(self, key, default=None):
@ -98,10 +98,10 @@ class ContainerCreateOptionalRequestParams(TypedDict, total=False):
Params here: https://platform.openai.com/docs/api-reference/containers/create
"""
expires_after: Optional[Dict[str, Any]] # ExpiresAfter object
file_ids: Optional[List[str]]
extra_headers: Optional[Dict[str, str]]
extra_body: Optional[Dict[str, str]]
expires_after: dict[str, Any] | None # ExpiresAfter object
file_ids: list[str] | None
extra_headers: dict[str, str] | None
extra_body: dict[str, str] | None
class ContainerCreateRequestParams(ContainerCreateOptionalRequestParams, total=False):
@ -121,11 +121,11 @@ class ContainerListOptionalRequestParams(TypedDict, total=False):
Params here: https://platform.openai.com/docs/api-reference/containers/list
"""
after: Optional[str]
limit: Optional[int]
order: Optional[str]
extra_headers: Optional[Dict[str, str]]
extra_query: Optional[Dict[str, str]]
after: str | None
limit: int | None
order: str | None
extra_headers: dict[str, str] | None
extra_query: dict[str, str] | None
class ContainerFileObject(BaseModel):
@ -134,13 +134,13 @@ class ContainerFileObject(BaseModel):
id: str
object: Literal["container.file", "container_file"] # OpenAI returns "container.file"
container_id: str
bytes: Optional[int] = None # Can be null for some files
bytes: int | None = None # Can be null for some files
created_at: int
path: str
source: str
_hidden_params: Dict[str, Any] = {}
_hidden_params: dict[str, Any] = {}
def __contains__(self, key):
def __contains__(self, key) -> bool:
return hasattr(self, key)
def get(self, key, default=None):
@ -160,12 +160,12 @@ class ContainerFileListResponse(BaseModel):
"""Response object for list container files request."""
object: Literal["list"]
data: List[ContainerFileObject]
first_id: Optional[str] = None
last_id: Optional[str] = None
data: list[ContainerFileObject]
first_id: str | None = None
last_id: str | None = None
has_more: bool
def __contains__(self, key):
def __contains__(self, key) -> bool:
return hasattr(self, key)
def get(self, key, default=None):
@ -189,7 +189,7 @@ class DeleteContainerFileResponse(BaseModel):
object: Literal["container.file.deleted", "container_file.deleted"]
deleted: bool
def __contains__(self, key):
def __contains__(self, key) -> bool:
return hasattr(self, key)
def get(self, key, default=None):

View file

@ -1,21 +1,19 @@
from typing import List, Optional, Union
from pydantic import BaseModel, ConfigDict
class EmbeddingRequest(BaseModel):
model: str
input: List[str] = []
input: list[str] = []
timeout: int = 600
api_base: Optional[str] = None
api_version: Optional[str] = None
api_key: Optional[str] = None
api_type: Optional[str] = None
api_base: str | None = None
api_version: str | None = None
api_key: str | None = None
api_type: str | None = None
caching: bool = False
user: Optional[str] = None
custom_llm_provider: Optional[Union[str, dict]] = None
litellm_call_id: Optional[str] = None
litellm_logging_obj: Optional[dict] = None
logger_fn: Optional[str] = None
user: str | None = None
custom_llm_provider: str | dict | None = None
litellm_call_id: str | None = None
litellm_logging_obj: dict | None = None
logger_fn: str | None = None
model_config = ConfigDict(extra="allow")

Some files were not shown because too many files have changed in this diff Show more