mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_daily_any_cleanup_08_04_2026
This commit is contained in:
commit
2bd6af1b64
319 changed files with 7655 additions and 5649 deletions
|
|
@ -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
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
276
litellm/caching/evicted_client_closer.py
Normal file
276
litellm/caching/evicted_client_closer.py
Normal 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()
|
||||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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, ...]:
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 ################
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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[
|
||||
|
|
|
|||
|
|
@ -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"})
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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}",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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}",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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"]:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue