Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_daily_any_cleanup_08_04_2026

This commit is contained in:
mateo-berri 2026-08-05 12:43:11 -07:00
commit 2bd6af1b64
319 changed files with 7655 additions and 5649 deletions

View file

@ -1,6 +1,6 @@
{
"reportAny": {
"limit": 29207
"limit": 29204
},
"reportArgumentType": {
"limit": 2635
@ -21,10 +21,10 @@
"limit": 215
},
"reportDuplicateImport": {
"limit": 24
"limit": 19
},
"reportExplicitAny": {
"limit": 9231
"limit": 9227
},
"reportFunctionMemberAccess": {
"limit": 7
@ -105,13 +105,13 @@
"limit": 113
},
"reportUnknownMemberType": {
"limit": 40345
"limit": 40340
},
"reportUnknownParameterType": {
"limit": 20293
},
"reportUnknownVariableType": {
"limit": 31895
"limit": 31797
},
"reportUnnecessaryCast": {
"limit": 122
@ -126,7 +126,7 @@
"limit": 865
},
"reportUntypedBaseClass": {
"limit": 165
"limit": 72
},
"reportUntypedFunctionDecorator": {
"limit": 33
@ -138,9 +138,9 @@
"limit": 139
},
"reportUnusedImport": {
"limit": 587
"limit": 555
},
"reportUnusedVariable": {
"limit": 147
"limit": 146
}
}

View file

@ -831,7 +831,7 @@ async def project_info(
)
# Check if user has access to this project (admin or team member)
is_admin = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
is_admin = user_api_key_has_admin_view(user_api_key_dict)
is_team_member = False
if project.team_id and user_api_key_dict.user_id:
@ -886,7 +886,7 @@ async def list_projects(
)
# If proxy admin, get all projects
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
if user_api_key_has_admin_view(user_api_key_dict):
projects: Sequence[
prisma_models.LiteLLM_ProjectTable
] = await prisma_client.db.litellm_projecttable.find_many(

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

@ -0,0 +1,276 @@
"""
Deferred close of HTTP/SDK clients that the LLM client cache has evicted.
Eviction only drops the cache's reference to a client. Every OpenAI/Azure SDK
client is a reference cycle (each resource namespace holds the client back), so
an evicted client and its pooled TCP connections survive until a generational
collection runs, which under load is thousands of requests later.
Closing at eviction time is not an option: a request that was handed the client
just before it was evicted is still using it, and closing it underneath that
request raises ``RuntimeError: Cannot send a request, as the client has been
closed.``
So an evicted client is closed once two conditions hold. A grace window must
have passed since its eviction, which covers a request that holds the client
but is momentarily not on the wire, and the client must report no connection in
flight. The second condition is what keeps the first honest: a request may run
for ``litellm.request_timeout`` seconds, 6000 by default, and a streaming
response is bounded only by how long the upstream keeps sending, so no deadline
on its own can promise that a request has finished.
Only clients litellm itself created are closed; a client the caller supplied is
left alone because litellm does not own its lifecycle.
A client that closes synchronously is closed from wherever the cache is next
used. One whose close is a coroutine needs the event loop it was evicted on, so
it waits for a call from that loop rather than having work scheduled onto a loop
it does not belong to. Queued clients are therefore bucketed by what it takes to
close them, and each bucket is ordered by deadline, so a reap walks the entries
that are due rather than the whole queue.
The queue holds its clients weakly, so waiting out a grace window never keeps
alive anything the collector would have reclaimed first.
"""
import asyncio
import contextlib
import inspect
import threading
import time
import weakref
from collections import deque
from collections.abc import Awaitable, Callable, Iterator
from dataclasses import dataclass, replace
from typing import Final
from litellm.constants import (
EVICTED_LLM_CLIENT_CLOSE_GRACE_SECONDS,
EVICTED_LLM_CLIENT_CLOSE_MAX_PENDING,
)
_CLOSABLE_ANYWHERE: Final = "closable-anywhere"
_CLOSABLE_ON_ANY_LOOP: Final = "closable-on-any-loop"
_BucketKey = str | int
@dataclass(frozen=True, slots=True)
class _PendingClose:
"""A queued close.
The client is held weakly, so queueing one never keeps alive anything the
collector would otherwise have reclaimed first.
``needs_loop`` is set for a client whose close is a coroutine; those can only
be closed from the event loop they were evicted on, recorded in ``loop_id``.
A client that closes synchronously carries neither constraint.
"""
client_ref: "weakref.ref[object]"
loop_id: int | None
needs_loop: bool
close_after: float
def _bucket_key(pending: _PendingClose) -> _BucketKey:
"""Which reaps can close this entry: any at all, any running a loop, or one loop's."""
if not pending.needs_loop:
return _CLOSABLE_ANYWHERE
if pending.loop_id is None:
return _CLOSABLE_ON_ANY_LOOP
return pending.loop_id
def _running_loop_id() -> int | None:
try:
return id(asyncio.get_running_loop())
except RuntimeError:
return None
def _close_function(client: object) -> Callable[[], object] | None:
close_fn: Final[Callable[[], object] | None] = getattr(client, "aclose", None) or getattr(client, "close", None)
return close_fn
def _transport_of(client: object) -> object:
"""The httpx transport behind an SDK wrapper, a litellm handler, or a bare client."""
for holder in (getattr(client, "_client", None), getattr(client, "client", None), client):
transport: object = getattr(holder, "_transport", None)
if transport is not None:
return transport
return None
def _connection_is_idle(connection: object) -> bool:
"""A pooled connection is idle unless it is servicing a request."""
is_idle: Final[object] = getattr(connection, "is_idle", None)
return bool(is_idle()) if callable(is_idle) else True
def _pool_has_busy_connection(transport: object) -> bool | None:
"""Whether the httpcore pool behind the transport is servicing a request.
``None`` when there is no such pool, so the caller can ask the other backend.
"""
pooled: Final[object] = getattr(getattr(transport, "_pool", None), "connections", None)
if not isinstance(pooled, (list, tuple)):
return None
return any(
not _connection_is_idle(connection) # pyright: ignore[reportUnknownArgumentType] # untyped pool list
for connection in pooled # pyright: ignore[reportUnknownVariableType] # untyped pool list
)
def _has_connection_in_flight(client: object) -> bool:
"""Whether the client is servicing a request right now.
Both connection backends litellm uses already account for the connections
they have handed out, so this reads the client's own lease accounting rather
than inferring it from elapsed time: httpcore reports a non-idle connection
for the whole of a response including a stream, and aiohttp holds the
connection in ``_acquired`` over the same span.
A client that cannot answer is reported as idle, which leaves the grace
window as the only guard, exactly as it was before this check existed.
"""
try:
transport: Final = _transport_of(client)
pooled_busy: Final = _pool_has_busy_connection(transport)
if pooled_busy is not None:
return pooled_busy
session: Final[object] = getattr(transport, "client", None)
return bool(getattr(getattr(session, "connector", None), "_acquired", None))
except Exception: # noqa: BLE001 - a client that cannot report its state is treated as idle
return False
async def _close_quietly(closing: Awaitable[object]) -> None:
with contextlib.suppress(Exception):
await closing
class EvictedClientCloser:
"""Closes evicted, litellm-owned clients once they are idle and out of grace."""
def __init__(
self,
grace_seconds: float = EVICTED_LLM_CLIENT_CLOSE_GRACE_SECONDS,
max_pending: int = EVICTED_LLM_CLIENT_CLOSE_MAX_PENDING,
clock: Callable[[], float] = time.monotonic,
) -> None:
self._grace_seconds = grace_seconds
self._max_pending = max_pending
self._clock = clock
self._owned: weakref.WeakSet[object] = weakref.WeakSet()
self._buckets: dict[_BucketKey, deque[_PendingClose]] = {} # mutable-ok: deadline-ordered queues
self._pending_count = 0
self._queue_lock = threading.Lock() # the cache is reachable from every worker thread's loop
self._close_tasks: set[asyncio.Task[None]] = set() # mutable-ok: strong refs to running closes
def mark_owned(self, client: object) -> None:
"""Record that litellm created this client, so it may be closed on eviction."""
try:
self._owned.add(client)
except TypeError:
pass # values that cannot be weak-referenced are never litellm clients
def _is_owned(self, client: object) -> bool:
try:
return client in self._owned
except TypeError:
return False # unhashable values are never litellm clients
def schedule(self, client: object) -> None:
"""Queue an evicted client for closing once it is idle and out of grace.
Past ``max_pending`` the client is left to the collector instead, so a
workload that churns the cache cannot grow this queue without bound.
Every queued entry comes due within one grace window, so the capacity it
occupies is returned within that window rather than held.
"""
if client is None or not self._is_owned(client):
return
close_fn: Final = _close_function(client)
if close_fn is None:
return
if self._pending_count >= self._max_pending:
return
self._enqueue(
_PendingClose(
client_ref=weakref.ref(client),
loop_id=_running_loop_id(),
needs_loop=inspect.iscoroutinefunction(close_fn),
close_after=self._clock() + self._grace_seconds,
)
)
def reap(self) -> None:
"""Close every queued client that is due, idle, and closable from here.
Called from the cache's read path, so the empty-queue exit comes first and
the work done past it is proportional to what is due, not to the queue.
"""
if not self._pending_count:
return
now: Final = self._clock()
for pending in self._take_due(_running_loop_id(), now):
client = pending.client_ref()
if client is None:
continue
if _has_connection_in_flight(client):
self._enqueue(replace(pending, close_after=now + self._grace_seconds))
continue
self._close(client)
@property
def pending_count(self) -> int:
return self._pending_count
def _enqueue(self, pending: _PendingClose) -> None:
"""Append to the entry's bucket, dropping any dead entries it queues behind.
Deadlines only ever move forward, so appending keeps each bucket ordered
by deadline, and entries whose client the collector already took sit at
the front rather than having to be searched for.
"""
with self._queue_lock:
bucket: Final = self._buckets.setdefault(_bucket_key(pending), deque()) # mutable-ok: FIFO by design
while bucket and bucket[0].client_ref() is None:
bucket.popleft()
self._pending_count -= 1
bucket.append(pending)
self._pending_count += 1
def _take_due(self, loop_id: int | None, now: float) -> tuple[_PendingClose, ...]:
buckets = (_CLOSABLE_ANYWHERE,) if loop_id is None else (_CLOSABLE_ANYWHERE, _CLOSABLE_ON_ANY_LOOP, loop_id)
with self._queue_lock:
return tuple(pending for key in buckets for pending in self._drain_locked(key, now))
def _drain_locked(self, key: _BucketKey, now: float) -> Iterator[_PendingClose]:
bucket: Final = self._buckets.get(key)
if bucket is None:
return
while bucket and bucket[0].close_after <= now:
self._pending_count -= 1
yield bucket.popleft()
if not bucket:
del self._buckets[key]
def _close(self, client: object) -> None:
close_fn: Final = _close_function(client)
if close_fn is None:
return
try:
closing: Final = close_fn()
except Exception: # noqa: BLE001 - a discarded client's close must never surface to callers
return
if not inspect.isawaitable(closing):
return
task: Final = asyncio.get_running_loop().create_task(_close_quietly(closing))
self._close_tasks.add(task)
task.add_done_callback(self._close_tasks.discard)
default_evicted_client_closer: Final = EvictedClientCloser()

View file

@ -5,21 +5,44 @@ Add the event loop to the cache key, to prevent event loop closed errors.
import asyncio
from typing import Final
from .evicted_client_closer import EvictedClientCloser, default_evicted_client_closer
from .in_memory_cache import InMemoryCache
class LLMClientCache(InMemoryCache):
"""Cache for LLM HTTP clients (OpenAI, Azure, httpx, etc.).
IMPORTANT: This cache intentionally does NOT close clients on eviction.
Evicted clients may still be in use by in-flight requests. Closing them
eagerly causes ``RuntimeError: Cannot send a request, as the client has
been closed.`` errors in production after the TTL (1 hour) expires.
An evicted client is never closed on the spot: a request handed the client
just before eviction is still using it, and closing it there raises
``RuntimeError: Cannot send a request, as the client has been closed.``
Clients that are no longer referenced will be garbage-collected normally.
For explicit shutdown cleanup, use ``close_litellm_async_clients()``.
Nor can eviction be left to rely on garbage collection. The SDK clients are
reference cycles, so an evicted client and its open TCP connections survive
until a generational collection runs. Instead a client litellm created is
handed to ``EvictedClientCloser``, which closes it once a grace window has
passed. Clients the caller supplied are left untouched.
"""
def __init__(
self,
max_size_in_memory: int | None = 200,
default_ttl: int | None = 600,
max_size_per_item: int | None = 1024,
evicted_client_closer: EvictedClientCloser | None = None,
) -> None:
super().__init__(
max_size_in_memory=max_size_in_memory,
default_ttl=default_ttl,
max_size_per_item=max_size_per_item,
)
self.evicted_client_closer = evicted_client_closer or default_evicted_client_closer
def _remove_key(self, key: str) -> None:
evicted: Final[object] = self.cache_dict.get(key)
super()._remove_key(key)
self.evicted_client_closer.schedule(evicted)
self.evicted_client_closer.reap()
def update_cache_key_with_event_loop(self, key):
"""
Add the event loop to the cache key, to prevent event loop closed errors.
@ -32,16 +55,22 @@ class LLMClientCache(InMemoryCache):
except RuntimeError: # handle no current running event loop
return key
def set_cache(self, key, value, **kwargs):
def set_cache(self, key: str, value: object, litellm_owned_client: bool = False, **kwargs):
"""``litellm_owned_client`` marks a client litellm built, so it may be closed once evicted."""
if litellm_owned_client:
self.evicted_client_closer.mark_owned(value)
key = self.update_cache_key_with_event_loop(key)
return super().set_cache(key, value, **kwargs)
async def async_set_cache(self, key, value, **kwargs):
async def async_set_cache(self, key: str, value: object, litellm_owned_client: bool = False, **kwargs):
if litellm_owned_client:
self.evicted_client_closer.mark_owned(value)
key = self.update_cache_key_with_event_loop(key)
return await super().async_set_cache(key, value, **kwargs)
def get_cache(self, key, **kwargs):
key = self.update_cache_key_with_event_loop(key)
self.evicted_client_closer.reap()
return super().get_cache(key, **kwargs)

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

@ -197,6 +197,16 @@ RUNWAYML_POLLING_TIMEOUT = int(os.getenv("RUNWAYML_POLLING_TIMEOUT", 600)) # 10
########## Networking constants ##############################################################
_DEFAULT_TTL_FOR_HTTPX_CLIENTS: Final = 3600 # 1 hour, re-use the same httpx client for 1 hour
# The earliest an evicted, litellm-created client may be closed. A request handed the
# client just before eviction is still using it, so nothing is closed inside this window;
# past it, the client is closed once it reports no connection in flight.
EVICTED_LLM_CLIENT_CLOSE_GRACE_SECONDS: Final = 900
# How many evicted clients may be queued for closing at once. Past this, an evicted client
# is left to the collector rather than letting a cache-churning workload grow the queue
# without bound. Each queued entry is ~100 bytes and comes due within one grace window.
EVICTED_LLM_CLIENT_CLOSE_MAX_PENDING: Final = 10_000
# Aiohttp connection pooling - prevents memory leaks from unbounded connection growth
# Set to 0 for unlimited (not recommended for production)
AIOHTTP_CONNECTOR_LIMIT: Final = int(os.getenv("AIOHTTP_CONNECTOR_LIMIT", 1000))

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

@ -427,88 +427,95 @@ class BaseAzureLLM(BaseOpenAILLM):
f"|azure_password={hashlib.sha256(_azure_password.encode()).hexdigest() if isinstance(_azure_password, str) else None}"
f"|azure_scope={_lp.get('azure_scope')}"
)
if client is None:
cached_client: Final = self.get_cached_openai_client(
client_initialization_params=client_initialization_params,
client_type="azure",
)
if cached_client:
if isinstance(cached_client, (AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI)):
return cached_client
azure_client_params: Final = self.initialize_azure_sdk_client(
litellm_params=litellm_params or {},
api_key=api_key,
api_base=api_base,
model_name=model,
api_version=api_version,
is_async=_is_async,
)
# For Azure v1 API, use standard OpenAI client instead of AzureOpenAI
# See: https://learn.microsoft.com/en-us/azure/ai-services/openai/reference#api-specs
if self._is_azure_v1_api_version(api_version):
# Extract only params that OpenAI client accepts
# Always use /openai/v1/ regardless of whether user passed "v1", "latest", or "preview"
# The OpenAI client accepts a callable for `api_key` and re-invokes it
# on every request (via `_refresh_api_key`), so passing
# `azure_ad_token_provider` directly preserves Azure AD token refresh
# behavior that the regular AzureOpenAI client provides.
v1_api_key: str | Callable[[], Any] | None = (
azure_client_params.get("api_key")
or azure_client_params.get("azure_ad_token_provider")
or azure_client_params.get("azure_ad_token")
)
if _is_async is True and callable(v1_api_key):
# AsyncOpenAI expects an async provider; wrap the sync provider
# returned by azure-identity. Offload to a thread so a token
# refresh (blocking HTTP call to AAD on cache miss) does not
# stall the event loop.
_sync_provider: Final = v1_api_key
async def _async_v1_api_key() -> str:
return await asyncio.to_thread(_sync_provider)
v1_api_key = _async_v1_api_key
v1_params: Final[dict[str, Any]] = {
"api_key": v1_api_key,
"base_url": f"{api_base}/openai/v1/",
}
if "timeout" in azure_client_params:
v1_params["timeout"] = azure_client_params["timeout"]
if "max_retries" in azure_client_params:
v1_params["max_retries"] = azure_client_params["max_retries"]
if "http_client" in azure_client_params:
v1_params["http_client"] = azure_client_params["http_client"]
verbose_logger.debug("Using Azure v1 API with base_url: %s", v1_params["base_url"])
if _is_async is True:
openai_client = AsyncOpenAI(**v1_params)
else:
openai_client = OpenAI(**v1_params)
else:
# Traditional Azure API uses AzureOpenAI client
if _is_async is True:
openai_client = AsyncAzureOpenAI(**azure_client_params)
else:
openai_client = AzureOpenAI(**azure_client_params)
else:
openai_client = client
if client is not None:
if (
api_version is not None
and isinstance(openai_client, (AzureOpenAI, AsyncAzureOpenAI))
and isinstance(openai_client._custom_query, dict)
and isinstance(client, (AzureOpenAI, AsyncAzureOpenAI))
and isinstance(client._custom_query, dict)
):
# set api_version to version passed by user
openai_client._custom_query.setdefault("api-version", api_version)
client._custom_query.setdefault("api-version", api_version)
self.set_cached_openai_client(
openai_client=client,
client_initialization_params=client_initialization_params,
client_type="azure",
litellm_owned_client=False,
)
return client
cached_client: Final = self.get_cached_openai_client(
client_initialization_params=client_initialization_params,
client_type="azure",
)
if cached_client:
if isinstance(cached_client, (AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI)):
return cached_client
azure_client_params: Final = self.initialize_azure_sdk_client(
litellm_params=litellm_params or {},
api_key=api_key,
api_base=api_base,
model_name=model,
api_version=api_version,
is_async=_is_async,
)
# For Azure v1 API, use standard OpenAI client instead of AzureOpenAI
# See: https://learn.microsoft.com/en-us/azure/ai-services/openai/reference#api-specs
if self._is_azure_v1_api_version(api_version):
# Extract only params that OpenAI client accepts
# Always use /openai/v1/ regardless of whether user passed "v1", "latest", or "preview"
# The OpenAI client accepts a callable for `api_key` and re-invokes it
# on every request (via `_refresh_api_key`), so passing
# `azure_ad_token_provider` directly preserves Azure AD token refresh
# behavior that the regular AzureOpenAI client provides.
v1_api_key: str | Callable[[], Any] | None = (
azure_client_params.get("api_key")
or azure_client_params.get("azure_ad_token_provider")
or azure_client_params.get("azure_ad_token")
)
if _is_async is True and callable(v1_api_key):
# AsyncOpenAI expects an async provider; wrap the sync provider
# returned by azure-identity. Offload to a thread so a token
# refresh (blocking HTTP call to AAD on cache miss) does not
# stall the event loop.
_sync_provider: Final = v1_api_key
async def _async_v1_api_key() -> str:
return await asyncio.to_thread(_sync_provider)
v1_api_key = _async_v1_api_key
v1_params: Final[dict[str, Any]] = {
"api_key": v1_api_key,
"base_url": f"{api_base}/openai/v1/",
}
if "timeout" in azure_client_params:
v1_params["timeout"] = azure_client_params["timeout"]
if "max_retries" in azure_client_params:
v1_params["max_retries"] = azure_client_params["max_retries"]
if "http_client" in azure_client_params:
v1_params["http_client"] = azure_client_params["http_client"]
verbose_logger.debug("Using Azure v1 API with base_url: %s", v1_params["base_url"])
if _is_async is True:
openai_client = AsyncOpenAI(**v1_params)
else:
openai_client = OpenAI(**v1_params)
else:
# Traditional Azure API uses AzureOpenAI client
if _is_async is True:
openai_client = AsyncAzureOpenAI(**azure_client_params)
else:
openai_client = AzureOpenAI(**azure_client_params)
# save client in-memory cache
self.set_cached_openai_client(
openai_client=openai_client,
client_initialization_params=client_initialization_params,
client_type="azure",
litellm_owned_client=self.owns_wrapped_http_client(azure_client_params.get("http_client")),
)
return openai_client

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

@ -1441,6 +1441,7 @@ def get_async_httpx_client(
key=_cache_key_name,
value=_new_client,
ttl=_DEFAULT_TTL_FOR_HTTPX_CLIENTS,
litellm_owned_client=True,
)
return _new_client
@ -1486,5 +1487,6 @@ def _get_httpx_client(params: dict | None = None) -> HTTPHandler:
key=_cache_key_name,
value=_new_client,
ttl=_DEFAULT_TTL_FOR_HTTPX_CLIENTS,
litellm_owned_client=True,
)
return _new_client

View file

@ -128,13 +128,33 @@ class BaseOpenAILLM:
_cached_client: Final = litellm.in_memory_llm_clients_cache.get_cache(_cache_key)
return _cached_client
@staticmethod
def owns_wrapped_http_client(http_client: httpx.Client | httpx.AsyncClient | None) -> bool:
"""Whether litellm may close an SDK client built around ``http_client``.
``_get_async_http_client`` / ``_get_sync_http_client`` hand back
``litellm.aclient_session`` / ``litellm.client_session`` when the caller
configured one. The SDK's ``close()`` closes whatever http client it was
given, so an SDK client wrapping one of those shared sessions must never be
closed on eviction; the caller goes on using the session. ``None`` means the
SDK built its own http client, which litellm does own.
"""
if http_client is None:
return True
return http_client is not litellm.aclient_session and http_client is not litellm.client_session
@staticmethod
def set_cached_openai_client(
openai_client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI,
client_type: Literal["openai", "azure"],
client_initialization_params: dict,
litellm_owned_client: bool = False,
):
"""Stores the OpenAI client in the in-memory cache for _DEFAULT_TTL_FOR_HTTPX_CLIENTS SECONDS"""
"""Stores the OpenAI client in the in-memory cache for _DEFAULT_TTL_FOR_HTTPX_CLIENTS SECONDS
``litellm_owned_client`` says litellm built this client, so the cache may close it once it
is evicted. A client the caller supplied stays open, since litellm does not own it.
"""
_cache_key: Final = BaseOpenAILLM.get_openai_client_cache_key(
client_initialization_params=client_initialization_params,
client_type=client_type,
@ -143,6 +163,7 @@ class BaseOpenAILLM:
key=_cache_key,
value=openai_client,
ttl=_DEFAULT_TTL_FOR_HTTPX_CLIENTS,
litellm_owned_client=litellm_owned_client,
)
@staticmethod

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(
@ -360,11 +360,16 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
if cached_client:
if isinstance(cached_client, OpenAI) or isinstance(cached_client, AsyncOpenAI):
return cached_client
http_client: Final[httpx.Client | httpx.AsyncClient | None] = (
OpenAIChatCompletion._get_async_http_client(shared_session=shared_session)
if is_async
else OpenAIChatCompletion._get_sync_http_client()
)
if is_async:
_new_client: OpenAI | AsyncOpenAI = AsyncOpenAI(
api_key=api_key,
base_url=api_base,
http_client=OpenAIChatCompletion._get_async_http_client(shared_session=shared_session),
http_client=http_client,
timeout=timeout,
max_retries=max_retries,
organization=organization,
@ -373,7 +378,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
_new_client = OpenAI(
api_key=api_key,
base_url=api_base,
http_client=OpenAIChatCompletion._get_sync_http_client(),
http_client=http_client,
timeout=timeout,
max_retries=max_retries,
organization=organization,
@ -384,6 +389,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
openai_client=_new_client,
client_initialization_params=client_initialization_params,
client_type="openai",
litellm_owned_client=self.owns_wrapped_http_client(http_client),
)
return _new_client
@ -402,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
@ -439,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
@ -474,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).
@ -1288,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,
@ -2842,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,
@ -2881,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,
@ -2906,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,
@ -2933,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,
@ -2955,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,
@ -2978,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

@ -30,7 +30,11 @@ from litellm.proxy._experimental.mcp_server.utils import (
get_server_prefix,
merge_mcp_headers,
)
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy._types import (
LitellmUserRoles,
UserAPIKeyAuth,
user_api_key_has_admin_view,
)
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
@ -738,9 +742,7 @@ if MCP_AVAILABLE:
# The full catalog (allowlist filter skipped) is admin-only so the
# REST endpoint can't be used to enumerate deliberately-disabled tools.
apply_tool_filters: Final = not (
include_disabled_tools and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
)
apply_tool_filters: Final = not (include_disabled_tools and user_api_key_has_admin_view(user_api_key_dict))
if server_id is None:
server_id = mcp_server_name

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

@ -21,7 +21,12 @@ import litellm
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.litellm_logging import _get_masked_values
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy._types import (
CommonProxyErrors,
LitellmUserRoles,
UserAPIKeyAuth,
user_api_key_has_admin_view,
)
from litellm.proxy.a2a.agent_card import (
SUPPORTED_A2A_PROTOCOL_VERSIONS,
merge_agent_card,
@ -468,11 +473,7 @@ async def get_agent_by_id(
"""
await check_feature_access_for_user(user_api_key_dict, "agents")
is_admin = (
user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
)
if not is_admin:
if not user_api_key_has_admin_view(user_api_key_dict):
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import (
AgentRequestHandler,
)

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
@ -832,9 +832,6 @@ def _is_user_proxy_admin(user_obj: LiteLLM_UserTable | None):
if user_obj.user_role is not None and user_obj.user_role == LitellmUserRoles.PROXY_ADMIN.value:
return True
if user_obj.user_role is not None and user_obj.user_role == LitellmUserRoles.PROXY_ADMIN.value:
return True
return False

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

@ -260,7 +260,11 @@ class RouteChecks:
query_params: Final = request.query_params
user_id: Final = query_params.get("user_id")
verbose_proxy_logger.debug("user_id: %s & valid_token.user_id: %s", user_id, valid_token.user_id)
if user_id and user_id != valid_token.user_id:
if (
user_id
and user_id != valid_token.user_id
and _user_role != LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value
):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=f"key not allowed to access this user's info. user_id={user_id}, key's user_id={valid_token.user_id}",

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
@ -212,7 +213,7 @@ async def list_guardrails_v2(
from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER
from litellm.proxy.proxy_server import prisma_client
is_admin: Final = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
is_admin: Final = _user_has_admin_view(user_api_key_dict)
try:
guardrails = (
@ -944,7 +945,7 @@ async def get_guardrail_submission(
if prisma_client is None:
raise HTTPException(status_code=500, detail="Prisma client not initialized")
is_admin: Final = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
is_admin: Final = _user_has_admin_view(user_api_key_dict)
try:
row: Final = await _guardrails_table(prisma_client).find_unique(where={"guardrail_id": guardrail_id})
@ -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

@ -20,7 +20,6 @@ from typing import (
Protocol,
TypeAlias,
TypedDict,
Union,
)
from litellm import DualCache
@ -59,7 +58,7 @@ if TYPE_CHECKING:
from litellm.types.agents import AgentResponse
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

@ -33,6 +33,7 @@ from litellm.proxy._types import (
LitellmTableNames,
LitellmUserRoles,
UserAPIKeyAuth,
user_api_key_has_admin_view,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.utils import invalidate_config_param
@ -302,7 +303,8 @@ async def get_coordination_redis_settings(
- fields: all configurable settings with their metadata (type, description, default, section)
- source: "coordination_redis" | "cache_backend" | "environment" | null
"""
_enforce_proxy_admin(user_api_key_dict)
if not user_api_key_has_admin_view(user_api_key_dict):
_enforce_proxy_admin(user_api_key_dict)
settings: Final = await _current_coordination_redis_settings()
source: Final = _coordination_redis_source(settings)

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", "")
@ -714,11 +714,10 @@ def _enforce_user_info_access(user_id: str | None, user_api_key_dict: UserAPIKey
"""
if user_id is None:
return
# Only true proxy admin bypasses ownership. PROXY_ADMIN_VIEW_ONLY is
# subject to the same `user_id == valid_token.user_id` rule that
# `RouteChecks.non_proxy_admin_allowed_routes_check` applies upstream
# for the `/user/info` route.
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
# Admin-view roles (PROXY_ADMIN and PROXY_ADMIN_VIEW_ONLY) bypass
# ownership, mirroring the `/user/info` carve-out that
# `RouteChecks.non_proxy_admin_allowed_routes_check` applies upstream.
if _user_has_admin_view(user_api_key_dict):
return
if user_id == user_api_key_dict.user_id:
return
@ -862,7 +861,7 @@ async def user_info(
raise Exception(
"Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys"
)
if user_id is None and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
if user_id is None and _user_has_admin_view(user_api_key_dict):
return await _get_user_info_for_proxy_admin(user_api_key_dict=user_api_key_dict)
elif user_id is None:
user_id = user_api_key_dict.user_id

View file

@ -78,6 +78,7 @@ from litellm.proxy.management_endpoints.common_utils import (
_is_user_team_admin,
_set_object_metadata_field,
_team_member_has_permission,
_user_has_admin_view,
validate_finite_spend,
)
from litellm.proxy.management_endpoints.model_management_endpoints import (
@ -214,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
@ -497,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,
@ -751,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:
"""
@ -1160,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],
@ -1231,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"
@ -1270,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.
@ -1295,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.
@ -1311,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:
"""
@ -1347,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:
@ -1397,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.
@ -1430,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.
@ -1486,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:
"""
@ -1943,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)
@ -5102,7 +5103,7 @@ async def validate_key_list_check(
key_hash: str | None,
prisma_client: PrismaClient,
) -> LiteLLM_UserTable | None:
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value:
if _user_has_admin_view(user_api_key_dict):
return None
if user_api_key_dict.user_id is None:
@ -5671,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:
@ -5683,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
@ -5917,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

@ -25,7 +25,12 @@ except ImportError:
from pydantic import BaseModel
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy._types import (
CommonProxyErrors,
LitellmUserRoles,
UserAPIKeyAuth,
user_api_key_has_admin_view,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.repositories.table_repositories import (
WorkflowEventRepository,
@ -47,6 +52,10 @@ def _is_admin(user_api_key_dict: UserAPIKeyAuth) -> bool:
return user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
def _read_scope_caller(user_api_key_dict: UserAPIKeyAuth) -> UserAPIKeyAuth | None:
return None if user_api_key_has_admin_view(user_api_key_dict) else user_api_key_dict
def _caller_key(user_api_key_dict: UserAPIKeyAuth) -> str | None:
"""Return the hashed key token that identifies this caller, or None for master key."""
return user_api_key_dict.token
@ -199,7 +208,7 @@ async def list_workflow_runs(
where["status"] = {"in": statuses} if len(statuses) > 1 else statuses[0]
# Non-admin callers are scoped to their own key.
if not _is_admin(user_api_key_dict):
if not user_api_key_has_admin_view(user_api_key_dict):
caller: Final = _caller_key(user_api_key_dict)
if caller:
where["created_by"] = caller
@ -238,7 +247,7 @@ async def get_workflow_run(
)
if run is None:
raise HTTPException(status_code=404, detail=f"Run '{run_id}' not found")
if not _is_admin(user_api_key_dict):
if not user_api_key_has_admin_view(user_api_key_dict):
caller: Final = _caller_key(user_api_key_dict)
if not caller or run.created_by != caller:
raise HTTPException(status_code=404, detail=f"Run '{run_id}' not found")
@ -377,7 +386,7 @@ async def list_workflow_events(
if prisma_client is None:
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
await _require_run(prisma_client, run_id, user_api_key_dict)
await _require_run(prisma_client, run_id, _read_scope_caller(user_api_key_dict))
try:
events: Final = await WorkflowEventRepository(prisma_client).table.find_many(
@ -461,7 +470,7 @@ async def list_workflow_messages(
if prisma_client is None:
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
await _require_run(prisma_client, run_id, user_api_key_dict)
await _require_run(prisma_client, run_id, _read_scope_caller(user_api_key_dict))
try:
messages: Final = await WorkflowMessageRepository(prisma_client).table.find_many(

View file

@ -27,6 +27,7 @@ from litellm.proxy._types import (
CommonProxyErrors,
LitellmUserRoles,
UserAPIKeyAuth,
user_api_key_has_admin_view,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.repositories.table_repositories import MemoryRepository
@ -66,7 +67,7 @@ def _visibility_filter(user_api_key_dict: UserAPIKeyAuth) -> dict | None:
Prisma `where` fragment restricting rows to those the caller can see.
Returns None for admins (no restriction).
"""
if _is_admin(user_api_key_dict):
if user_api_key_has_admin_view(user_api_key_dict):
return None
ors: Final[list[dict]] = []
if user_api_key_dict.user_id:

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

@ -18,7 +18,12 @@ from fastapi import (
from pydantic import BaseModel
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy._types import (
CommonProxyErrors,
LitellmUserRoles,
UserAPIKeyAuth,
user_api_key_has_admin_view,
)
from litellm.proxy.auth.auth_utils import is_request_body_safe
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.path_utils import safe_filename
@ -317,7 +322,6 @@ async def list_prompts(
}
```
"""
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
# check key metadata for prompts
@ -347,10 +351,7 @@ async def list_prompts(
prompt_list.append(prompt_copy)
return ListPromptsResponse(prompts=prompt_list)
# check if user is proxy admin - show all prompts
if user_api_key_dict.user_role is not None and (
user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
):
if user_api_key_has_admin_view(user_api_key_dict):
# Get all prompts and filter to show only the latest version of each
all_prompts = list(IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS.values())
if environment:
@ -422,10 +423,7 @@ async def get_prompt_versions(
from litellm.proxy.proxy_server import prisma_client
# Only allow proxy admins to view version history
if user_api_key_dict.user_role is None or (
user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN
and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value
):
if not user_api_key_has_admin_view(user_api_key_dict):
raise HTTPException(status_code=403, detail="Only proxy admins can view prompt versions")
base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id)
@ -581,12 +579,7 @@ async def get_prompt_info(
prompts = cast(list[str] | None, user_api_key_dict.metadata.get("prompts", None))
if prompts is not None and prompt_id not in prompts:
raise HTTPException(status_code=400, detail=f"Prompt {prompt_id} not found")
if user_api_key_dict.user_role is not None and (
user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
):
pass
else:
if not user_api_key_has_admin_view(user_api_key_dict):
raise HTTPException(
status_code=403,
detail=f"You are not authorized to access this prompt. Your role - {user_api_key_dict.user_role}, Your key's prompts - {prompts}",

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
)
@ -8664,48 +8663,56 @@ class ProxyStartupEvent:
- Sets up prisma client
- Adds necessary views to proxy
"""
connected_client: PrismaClient | None = None
try:
prisma_client: PrismaClient | None = None
if database_url is not None:
try:
prisma_client = PrismaClient(database_url=database_url, proxy_logging_obj=proxy_logging_obj)
except Exception as e:
raise e
if database_url is None:
return None
try:
await prisma_client.connect()
except Exception as e:
if "P3018" in str(e) or "P3009" in str(e):
verbose_proxy_logger.debug("CRITICAL: DATABASE MIGRATION FAILED")
verbose_proxy_logger.debug("Your database is in a 'dirty' state.")
verbose_proxy_logger.debug("FIX: Run 'prisma migrate resolve --applied <migration_name>'")
raise e
prisma_client = PrismaClient(database_url=database_url, proxy_logging_obj=proxy_logging_obj)
## Start RDS IAM token refresh background task if enabled ##
# This proactively refreshes IAM tokens before they expire,
# preventing the 15-minute connection failure bug (#16220)
if hasattr(prisma_client, "db") and hasattr(prisma_client.db, "start_token_refresh_task"):
await prisma_client.db.start_token_refresh_task()
try:
await prisma_client.connect()
except Exception as e:
if "P3018" in str(e) or "P3009" in str(e):
verbose_proxy_logger.debug("CRITICAL: DATABASE MIGRATION FAILED")
verbose_proxy_logger.debug("Your database is in a 'dirty' state.")
verbose_proxy_logger.debug("FIX: Run 'prisma migrate resolve --applied <migration_name>'")
raise e
## Add necessary views to proxy ##
asyncio.create_task(
prisma_client.check_view_exists()
) # check if all necessary views exist. Don't block execution
connected_client = prisma_client
asyncio.create_task(
prisma_client._set_spend_logs_row_count_in_proxy_state()
) # set the spend logs row count in proxy state. Don't block execution
## Start RDS IAM token refresh background task if enabled ##
# This proactively refreshes IAM tokens before they expire,
# preventing the 15-minute connection failure bug (#16220)
if hasattr(prisma_client, "db") and hasattr(prisma_client.db, "start_token_refresh_task"):
await prisma_client.db.start_token_refresh_task()
# run a health check to ensure the DB is ready
if get_secret_bool("DISABLE_PRISMA_HEALTH_CHECK_ON_STARTUP", False) is not True:
await prisma_client.health_check()
## Add necessary views to proxy ##
asyncio.create_task(
prisma_client.check_view_exists()
) # check if all necessary views exist. Don't block execution
asyncio.create_task(
prisma_client._set_spend_logs_row_count_in_proxy_state()
) # set the spend logs row count in proxy state. Don't block execution
if hasattr(prisma_client, "start_db_health_watchdog_task"):
await prisma_client.start_db_health_watchdog_task()
# run a health check to ensure the DB is ready
if get_secret_bool("DISABLE_PRISMA_HEALTH_CHECK_ON_STARTUP", False) is not True:
await prisma_client.health_check()
if hasattr(prisma_client, "start_db_health_watchdog_task"):
await prisma_client.start_db_health_watchdog_task()
return prisma_client
except Exception as e:
PrismaDBExceptionHandler.handle_db_exception(e)
return None
if connected_client is not None:
verbose_proxy_logger.warning(
"Retaining the connected Prisma client after a post-connect startup step failed: %s. "
"The DB health watchdog keeps probing and reconnects once the database recovers.",
e,
)
return connected_client
@classmethod
def _init_dd_tracer(cls):
@ -8876,7 +8883,7 @@ async def model_list(
# Check if scope=expand is requested and user has admin privileges
should_expand_scope = False
if scope == "expand":
should_expand_scope = await _user_has_admin_privileges(
should_expand_scope = _user_has_admin_view(user_api_key_dict) or await _user_has_admin_privileges(
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
@ -11479,7 +11486,7 @@ async def _populate_team_access_on_models(
"""
user_teams: list[str] | Literal["*"] | None = None
direct_access_models: list[str] = []
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
if _user_has_admin_view(user_api_key_dict):
user_teams = "*"
direct_access_models = llm_router.get_model_ids(exclude_team_models=True) # has access to all models
elif user_api_key_dict.user_id is not None:
@ -15188,7 +15195,7 @@ async def get_config_general_settings(
)
GeneralSettingsUILiteLLMValue = Union[float, bool, str, None]
GeneralSettingsUILiteLLMValue = float | bool | str | None
class GeneralSettingsUILiteLLMFieldSpec(TypedDict):
@ -16122,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
@ -4269,7 +4269,7 @@ class PrismaClient:
import traceback
error_msg: Final = f"LiteLLM Prisma Client Exception connect(): {e}"
print_verbose(error_msg)
verbose_proxy_logger.warning(error_msg)
error_traceback: Final = error_msg + "\n" + traceback.format_exc()
end_time: Final = time.time()
_duration: Final = end_time - start_time
@ -4987,8 +4987,8 @@ class PrismaClient:
except Exception as e:
import traceback
error_msg: Final = f"LiteLLM Prisma Client Exception disconnect(): {e}"
print_verbose(error_msg)
error_msg: Final = f"LiteLLM Prisma Client Exception health_check(): {e}"
verbose_proxy_logger.warning(error_msg)
error_traceback: Final = error_msg + "\n" + traceback.format_exc()
end_time: Final = time.time()
_duration: Final = end_time - start_time

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
@ -6353,10 +6352,7 @@ 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}"
@ -7489,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]
@ -8238,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
@ -9483,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"]:

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

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