From f95367db5f3bb772e80e2d472b91d5f7eae6302f Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 4 Aug 2026 18:01:35 -0700 Subject: [PATCH 01/11] Revert "revert: "fix(caching): close evicted LLM clients so their connections are reclaimed (#35492)"" This reverts commit adb9a53ba1b5d4281936f686b56d6007230b785c. --- litellm/caching/evicted_client_closer.py | 277 ++++++++++++ litellm/caching/llm_caching_handler.py | 45 +- litellm/constants.py | 10 + litellm/llms/azure/common_utils.py | 2 + litellm/llms/custom_httpx/http_handler.py | 2 + litellm/llms/openai/common_utils.py | 23 +- litellm/llms/openai/openai.py | 10 +- .../caching/test_evicted_client_closer.py | 409 ++++++++++++++++++ .../caching/test_llm_caching_handler.py | 66 +++ .../llms/azure/test_azure_common_utils.py | 71 +++ .../llms/openai/test_openai_common_utils.py | 72 +++ 11 files changed, 976 insertions(+), 11 deletions(-) create mode 100644 litellm/caching/evicted_client_closer.py create mode 100644 tests/test_litellm/caching/test_evicted_client_closer.py diff --git a/litellm/caching/evicted_client_closer.py b/litellm/caching/evicted_client_closer.py new file mode 100644 index 00000000000..c895669be2b --- /dev/null +++ b/litellm/caching/evicted_client_closer.py @@ -0,0 +1,277 @@ +""" +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 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: + try: + await closing + except Exception: # noqa: BLE001 - a discarded client's close must never surface to callers + pass + + +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() diff --git a/litellm/caching/llm_caching_handler.py b/litellm/caching/llm_caching_handler.py index 7d072a40195..6fa5963c99b 100644 --- a/litellm/caching/llm_caching_handler.py +++ b/litellm/caching/llm_caching_handler.py @@ -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, + ): + 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) diff --git a/litellm/constants.py b/litellm/constants.py index 264f595027f..0c7316455d6 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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)) diff --git a/litellm/llms/azure/common_utils.py b/litellm/llms/azure/common_utils.py index 9e613ae4eb4..25dd9698624 100644 --- a/litellm/llms/azure/common_utils.py +++ b/litellm/llms/azure/common_utils.py @@ -509,6 +509,8 @@ class BaseAzureLLM(BaseOpenAILLM): openai_client=openai_client, client_initialization_params=client_initialization_params, client_type="azure", + litellm_owned_client=client is None + and self.owns_wrapped_http_client(azure_client_params.get("http_client")), ) return openai_client diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 726392577e5..ed52d17c81b 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -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 diff --git a/litellm/llms/openai/common_utils.py b/litellm/llms/openai/common_utils.py index 5c5e78c062d..527f44b930f 100644 --- a/litellm/llms/openai/common_utils.py +++ b/litellm/llms/openai/common_utils.py @@ -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 diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index 998319f3e85..3c6846823e2 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -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 diff --git a/tests/test_litellm/caching/test_evicted_client_closer.py b/tests/test_litellm/caching/test_evicted_client_closer.py new file mode 100644 index 00000000000..a08fd58079d --- /dev/null +++ b/tests/test_litellm/caching/test_evicted_client_closer.py @@ -0,0 +1,409 @@ +""" +Tests for EvictedClientCloser. + +An evicted client must stay open long enough for a request that already holds it +to finish, and must then actually be closed, otherwise its connection pool is +retained until a generational collection runs. A client the caller supplied is +never closed, because litellm does not own its lifecycle. +""" + +import asyncio +import gc +import weakref + +import httpx +import pytest + +from litellm.caching.evicted_client_closer import EvictedClientCloser +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + + +class FakeClock: + """Hand-advanced monotonic clock, so grace windows need no real waiting.""" + + def __init__(self) -> None: + self.now = 1000.0 + + def __call__(self) -> float: + return self.now + + def advance(self, seconds: float) -> None: + self.now += seconds + + +class AsyncClient: + def __init__(self) -> None: + self.closed = False + + async def close(self) -> None: + self.closed = True + + +class SyncClient: + def __init__(self) -> None: + self.closed = False + + def close(self) -> None: + self.closed = True + + +class CountingDeadline(float): + """A clock reading that tallies every deadline comparison made against it. + + Deadline comparisons are the work a reap does, so counting them says whether + that work tracks the entries that are due or the size of the whole queue. + """ + + comparisons = 0 + + def __add__(self, other: float) -> "CountingDeadline": + return CountingDeadline(float(self) + other) + + def __le__(self, other: float) -> bool: + CountingDeadline.comparisons += 1 + return float(self) <= float(other) + + def __gt__(self, other: float) -> bool: + CountingDeadline.comparisons += 1 + return float(self) > float(other) + + +def make_closer(clock: FakeClock, grace_seconds: float = 60.0) -> EvictedClientCloser: + return EvictedClientCloser(grace_seconds=grace_seconds, clock=clock) + + +async def _trickling_upstream(reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None: + """Serves a chunked body slowly, so a request stays on the wire long enough to observe.""" + await reader.read(4096) + writer.write(b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n") + await writer.drain() + for _ in range(6): + writer.write(b"5\r\nhello\r\n") + await writer.drain() + await asyncio.sleep(0.1) + writer.write(b"0\r\n\r\n") + await writer.drain() + + +@pytest.mark.asyncio +async def test_owned_client_is_closed_once_the_grace_window_elapses(): + clock = FakeClock() + closer = make_closer(clock) + client = AsyncClient() + + closer.mark_owned(client) + closer.schedule(client) + clock.advance(61.0) + closer.reap() + await asyncio.sleep(0.05) + + assert client.closed is True + assert closer.pending_count == 0 + + +@pytest.mark.asyncio +async def test_owned_client_stays_open_inside_the_grace_window(): + """A request handed the client just before eviction is still using it.""" + clock = FakeClock() + closer = make_closer(clock) + client = AsyncClient() + + closer.mark_owned(client) + closer.schedule(client) + clock.advance(59.0) + closer.reap() + await asyncio.sleep(0.05) + + assert client.closed is False + assert closer.pending_count == 1 + + +@pytest.mark.asyncio +async def test_caller_supplied_client_is_never_closed(): + clock = FakeClock() + closer = make_closer(clock) + client = AsyncClient() + + closer.schedule(client) + clock.advance(3600.0) + closer.reap() + await asyncio.sleep(0.05) + + assert client.closed is False + assert closer.pending_count == 0 + + +@pytest.mark.asyncio +async def test_sync_client_is_closed_once_the_grace_window_elapses(): + clock = FakeClock() + closer = make_closer(clock) + client = SyncClient() + + closer.mark_owned(client) + closer.schedule(client) + clock.advance(61.0) + closer.reap() + + assert client.closed is True + + +@pytest.mark.asyncio +async def test_a_failing_close_does_not_propagate_or_block_the_others(): + class ExplodingClient: + async def close(self) -> None: + raise RuntimeError("connection already gone") + + clock = FakeClock() + closer = make_closer(clock) + exploding, healthy = ExplodingClient(), AsyncClient() + + for client in (exploding, healthy): + closer.mark_owned(client) + closer.schedule(client) + clock.advance(61.0) + closer.reap() + await asyncio.sleep(0.05) + + assert healthy.closed is True + + +@pytest.mark.asyncio +async def test_an_unhashable_cached_value_does_not_break_eviction(): + """The cache holds arbitrary values; an ownership test must never raise on one.""" + + class Unhashable: + __hash__ = None # pyright: ignore[reportAssignmentType] # unhashable by construction + + clock = FakeClock() + closer = make_closer(clock) + + closer.mark_owned(Unhashable()) + closer.schedule(Unhashable()) + + assert closer.pending_count == 0 + + +@pytest.mark.asyncio +async def test_values_with_nothing_to_close_are_never_queued(): + """The cache holds plain values too; those have nothing to reclaim.""" + + class NotAClient: + pass + + clock = FakeClock() + closer = make_closer(clock) + value = NotAClient() + + closer.mark_owned(value) + closer.schedule(value) + + assert closer.pending_count == 0 + + +@pytest.mark.asyncio +async def test_a_queued_client_is_not_kept_alive_by_the_queue(): + """Waiting out a grace window must not retain what the collector would free first.""" + clock = FakeClock() + closer = make_closer(clock) + client = AsyncClient() + gone = weakref.ref(client) + + closer.mark_owned(client) + closer.schedule(client) + del client + gc.collect() + + assert gone() is None, "the pending queue is holding the client alive" + + clock.advance(61.0) + closer.reap() + assert closer.pending_count == 0 + + +def test_sync_client_evicted_outside_an_event_loop_is_still_closed(): + """The sync httpx handler is cached and evicted from call sites with no loop.""" + clock = FakeClock() + closer = make_closer(clock) + client = SyncClient() + + closer.mark_owned(client) + closer.schedule(client) + assert closer.pending_count == 1 + + clock.advance(61.0) + closer.reap() + + assert client.closed is True + assert closer.pending_count == 0 + + +@pytest.mark.asyncio +async def test_an_async_client_waits_for_a_loop_rather_than_being_dropped(): + clock = FakeClock() + closer = make_closer(clock) + client = AsyncClient() + closer.mark_owned(client) + + def schedule_outside_a_loop() -> None: + closer.schedule(client) + clock.advance(61.0) + closer.reap() + + await asyncio.to_thread(schedule_outside_a_loop) + assert client.closed is False, "no loop was running, so it could not have been closed" + assert closer.pending_count == 1 + + closer.reap() + await asyncio.sleep(0.05) + + assert client.closed is True + + +@pytest.mark.asyncio +async def test_a_client_evicted_on_another_event_loop_is_left_alone(): + """Closing a client bound to a different loop would schedule work on that loop.""" + clock = FakeClock() + closer = make_closer(clock) + client = AsyncClient() + closer.mark_owned(client) + + def schedule_on_its_own_loop() -> None: + asyncio.run(_schedule()) + + async def _schedule() -> None: + closer.schedule(client) + + await asyncio.to_thread(schedule_on_its_own_loop) + assert closer.pending_count == 1 + + clock.advance(61.0) + closer.reap() + await asyncio.sleep(0.05) + + assert client.closed is False + assert closer.pending_count == 1 + + +@pytest.mark.asyncio +async def test_a_client_serving_a_request_is_not_closed_when_its_grace_window_ends(): + """The grace window on its own cannot promise that a request has finished. + + ``litellm.request_timeout`` defaults to 6000 seconds and a streaming response + is bounded only by how long the upstream keeps sending, so a client past its + deadline is closed only once its own pool reports nothing in flight. + """ + server = await asyncio.start_server(_trickling_upstream, "127.0.0.1", 0) + port = server.sockets[0].getsockname()[1] + clock = FakeClock() + closer = make_closer(clock) + client = httpx.AsyncClient() + + closer.mark_owned(client) + closer.schedule(client) + + async def read_the_stream() -> int: + received = 0 + async with client.stream("GET", f"http://127.0.0.1:{port}/") as response: + async for chunk in response.aiter_bytes(): + received += len(chunk) + return received + + streaming = asyncio.create_task(read_the_stream()) + await asyncio.sleep(0.25) # the request is on the wire + clock.advance(3600.0) # and its grace window is long gone + closer.reap() + await asyncio.sleep(0.05) + + assert client.is_closed is False, "closed a client that was serving a request" + assert await streaming > 0, "the in-flight request did not survive the reap" + + clock.advance(3600.0) + closer.reap() + await asyncio.sleep(0.05) + + assert client.is_closed is True, "an idle client past its grace window must be closed" + assert closer.pending_count == 0 + server.close() + + +@pytest.mark.asyncio +async def test_the_aiohttp_backed_handler_is_not_closed_mid_request(): + """The default async path is aiohttp-backed, whose pool accounts for its own leases.""" + server = await asyncio.start_server(_trickling_upstream, "127.0.0.1", 0) + port = server.sockets[0].getsockname()[1] + clock = FakeClock() + closer = make_closer(clock) + handler = AsyncHTTPHandler() + + closer.mark_owned(handler) + closer.schedule(handler) + + request = asyncio.create_task(handler.get(f"http://127.0.0.1:{port}/")) + await asyncio.sleep(0.25) + clock.advance(3600.0) + closer.reap() + await asyncio.sleep(0.05) + + assert handler.client.is_closed is False, "closed a handler that was serving a request" + assert (await request).status_code == 200 + + clock.advance(3600.0) + closer.reap() + await asyncio.sleep(0.05) + + assert handler.client.is_closed is True + server.close() + + +def test_the_pending_queue_cannot_grow_past_its_bound(): + """A caller that churns the client cache must not be able to grow this queue.""" + clock = FakeClock() + closer = EvictedClientCloser(grace_seconds=60.0, max_pending=8, clock=clock) + clients = tuple(SyncClient() for _ in range(50)) + + for client in clients: + closer.mark_owned(client) + closer.schedule(client) + + assert closer.pending_count == 8, "the queue grew past max_pending" + + clock.advance(61.0) + closer.reap() + + assert closer.pending_count == 0 + assert sum(client.closed for client in clients) == 8, "everything queued should have been closed" + + +def test_a_reap_looks_at_what_is_due_rather_than_at_the_whole_queue(): + """Sustained churn evicts a client per request, and every read of the cache reaps. + + So the cost of a reap has to track the entries that are due, not the length of + the queue; a reap that filters the whole queue makes the pair quadratic. Each + bucket is ordered by deadline, so an up-to-date reap compares one entry per + bucket and stops. Counting the comparisons measures that directly, where a + wall-clock budget would only measure the machine. + """ + evictions = 1_000 + clock = FakeClock() + closer = EvictedClientCloser( + grace_seconds=60.0, + max_pending=evictions, + clock=lambda: CountingDeadline(clock.now), + ) + clients = tuple(SyncClient() for _ in range(evictions)) + for client in clients: + closer.mark_owned(client) + + CountingDeadline.comparisons = 0 + for client in clients: + closer.schedule(client) + closer.reap() # nothing is due yet, which is the hot path + clock.advance(61.0) + closer.reap() + + assert closer.pending_count == 0 + assert all(client.closed for client in clients) + assert CountingDeadline.comparisons < 10 * evictions, ( + f"{CountingDeadline.comparisons} deadline comparisons for {evictions} evictions; " + "a reap is walking the whole queue" + ) diff --git a/tests/test_litellm/caching/test_llm_caching_handler.py b/tests/test_litellm/caching/test_llm_caching_handler.py index 8e6a94945b0..5f0e82dbb80 100644 --- a/tests/test_litellm/caching/test_llm_caching_handler.py +++ b/tests/test_litellm/caching/test_llm_caching_handler.py @@ -19,6 +19,7 @@ sys.path.insert( 0, os.path.abspath("../../..") ) # Adds the parent directory to the system path +from litellm.caching.evicted_client_closer import EvictedClientCloser from litellm.caching.llm_caching_handler import LLMClientCache @@ -156,6 +157,71 @@ def test_remove_key_no_event_loop(): assert "test-key" not in cache.cache_dict +class _FakeClock: + """Hand-advanced monotonic clock, so grace windows need no real waiting.""" + + def __init__(self) -> None: + self.now = 1000.0 + + def __call__(self) -> float: + return self.now + + def advance(self, seconds: float) -> None: + self.now += seconds + + +@pytest.mark.asyncio +async def test_evicted_litellm_owned_client_is_closed_once_the_grace_window_elapses(): + """ + Eviction only drops the cache's reference. The SDK clients are reference + cycles, so without an explicit close the client keeps its connection pool + open until a generational collection runs. + """ + clock = _FakeClock() + cache = LLMClientCache( + max_size_in_memory=2, + evicted_client_closer=EvictedClientCloser(grace_seconds=60.0, clock=clock), + ) + + client = MockAsyncClient() + cache.set_cache("client-key", client, litellm_owned_client=True, ttl=600) + + cache.ttl_dict = {key: 0 for key in cache.ttl_dict} + cache.expiration_heap = [(0, key) for _, key in cache.expiration_heap] + cache.evict_cache() + await asyncio.sleep(0.1) + assert client.closed is False, "an in-flight request may still hold the client" + + clock.advance(61.0) + cache.get_cache("any-key") + await asyncio.sleep(0.1) + + assert client.closed is True + + +@pytest.mark.asyncio +async def test_evicted_caller_supplied_client_is_never_closed(): + """litellm does not own a client the caller passed in, so it must stay open.""" + clock = _FakeClock() + cache = LLMClientCache( + max_size_in_memory=2, + evicted_client_closer=EvictedClientCloser(grace_seconds=60.0, clock=clock), + ) + + client = MockAsyncClient() + cache.set_cache("client-key", client, ttl=600) + + cache.ttl_dict = {key: 0 for key in cache.ttl_dict} + cache.expiration_heap = [(0, key) for _, key in cache.expiration_heap] + cache.evict_cache() + + clock.advance(3600.0) + cache.get_cache("any-key") + await asyncio.sleep(0.1) + + assert client.closed is False + + def test_remove_key_removes_plain_values(): """ _remove_key correctly removes non-client values (strings, dicts, etc.). diff --git a/tests/test_litellm/llms/azure/test_azure_common_utils.py b/tests/test_litellm/llms/azure/test_azure_common_utils.py index c0446a6cfba..85db11fdb24 100644 --- a/tests/test_litellm/llms/azure/test_azure_common_utils.py +++ b/tests/test_litellm/llms/azure/test_azure_common_utils.py @@ -2034,3 +2034,74 @@ def test_azure_traditional_api_uses_azure_openai_client(): assert isinstance( async_client, AsyncAzureOpenAI ), f"Expected AsyncAzureOpenAI client for api_version={api_version}" + + +def test_evicting_an_azure_client_built_on_the_callers_session_leaves_it_open(monkeypatch): + """`initialize_azure_sdk_client` puts `litellm.aclient_session` on the SDK client. + + That session belongs to the caller. `AsyncAzureOpenAI.close()` closes whatever + http client it was handed, so treating the wrapper as litellm's to close would + close the caller's shared session out from under them. + """ + import httpx + + from litellm.caching.evicted_client_closer import EvictedClientCloser + from litellm.caching.llm_caching_handler import LLMClientCache + + shared_session = httpx.AsyncClient() + closer = EvictedClientCloser(grace_seconds=0.0) + monkeypatch.setattr(litellm, "aclient_session", shared_session) + monkeypatch.setattr( + litellm, + "in_memory_llm_clients_cache", + LLMClientCache(evicted_client_closer=closer), + ) + + wrapper = BaseAzureLLM().get_azure_openai_client( + api_key="not-a-real-key", + api_base="https://litellm.openai.azure.com", + api_version="2024-02-01", + litellm_params={}, + _is_async=True, + ) + + assert wrapper is not None + assert wrapper._client is shared_session, "the wrapper should be built on the caller's session" + + closer.schedule(wrapper) + closer.reap() + + assert closer.pending_count == 0, "a wrapper around the caller's session must never be queued" + assert shared_session.is_closed is False, "closed the session the caller configured" + + +def test_an_azure_client_litellm_built_its_own_http_client_for_is_still_closed(monkeypatch): + """The ownership check must not turn the reclaim off for the ordinary case.""" + from litellm.caching.evicted_client_closer import EvictedClientCloser + from litellm.caching.llm_caching_handler import LLMClientCache + + closer = EvictedClientCloser(grace_seconds=0.0) + monkeypatch.setattr(litellm, "aclient_session", None) + monkeypatch.setattr(litellm, "client_session", None) + monkeypatch.setattr( + litellm, + "in_memory_llm_clients_cache", + LLMClientCache(evicted_client_closer=closer), + ) + + wrapper = BaseAzureLLM().get_azure_openai_client( + api_key="not-a-real-key", + api_base="https://litellm.openai.azure.com", + api_version="2024-02-01", + litellm_params={}, + _is_async=False, + ) + + assert wrapper is not None + closer.schedule(wrapper) + + assert closer.pending_count == 1, "litellm built this client's http client, so it owns it" + + closer.reap() + + assert wrapper.is_closed() is True diff --git a/tests/test_litellm/llms/openai/test_openai_common_utils.py b/tests/test_litellm/llms/openai/test_openai_common_utils.py index ce25f7e9af6..a099b5c659f 100644 --- a/tests/test_litellm/llms/openai/test_openai_common_utils.py +++ b/tests/test_litellm/llms/openai/test_openai_common_utils.py @@ -175,3 +175,75 @@ def test_get_openai_client_cache_key(client_type): ) assert isinstance(key, str) assert "api_key=sk-test" in key + + +def test_evicting_a_client_built_on_the_callers_session_leaves_that_session_open(monkeypatch): + """`litellm.aclient_session` belongs to the caller, who goes on using it. + + `_get_async_http_client` hands that session straight back, so the SDK client + litellm builds around it is only a wrapper. The SDK's `close()` closes + whatever http client it was given, so treating the wrapper as litellm's to + close would close the caller's shared session out from under them. + """ + import httpx + + from litellm.caching.evicted_client_closer import EvictedClientCloser + from litellm.caching.llm_caching_handler import LLMClientCache + from litellm.llms.openai.openai import OpenAIChatCompletion + + shared_session = httpx.AsyncClient() + closer = EvictedClientCloser(grace_seconds=0.0) + monkeypatch.setattr(litellm, "aclient_session", shared_session) + monkeypatch.setattr( + litellm, + "in_memory_llm_clients_cache", + LLMClientCache(evicted_client_closer=closer), + ) + + wrapper = OpenAIChatCompletion()._get_openai_client( + is_async=True, + api_key="sk-not-a-real-key", + api_base="https://api.openai.com/v1", + max_retries=2, + ) + + assert wrapper is not None + assert wrapper._client is shared_session, "the wrapper should be built on the caller's session" + + closer.schedule(wrapper) + closer.reap() + + assert closer.pending_count == 0, "a wrapper around the caller's session must never be queued" + assert shared_session.is_closed is False, "closed the session the caller configured" + + +def test_a_client_litellm_built_its_own_http_client_for_is_still_closed(monkeypatch): + """The ownership check must not turn the reclaim off for the ordinary case.""" + from litellm.caching.evicted_client_closer import EvictedClientCloser + from litellm.caching.llm_caching_handler import LLMClientCache + from litellm.llms.openai.openai import OpenAIChatCompletion + + closer = EvictedClientCloser(grace_seconds=0.0) + monkeypatch.setattr(litellm, "aclient_session", None) + monkeypatch.setattr(litellm, "client_session", None) + monkeypatch.setattr( + litellm, + "in_memory_llm_clients_cache", + LLMClientCache(evicted_client_closer=closer), + ) + + wrapper = OpenAIChatCompletion()._get_openai_client( + is_async=False, + api_key="sk-not-a-real-key", + api_base="https://api.openai.com/v1", + max_retries=2, + ) + + assert wrapper is not None + closer.schedule(wrapper) + + assert closer.pending_count == 1, "litellm built this client's http client, so it owns it" + + closer.reap() + + assert wrapper.is_closed() is True From d45e2bc34ec45418b29026d89a770582d5cfc4b5 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 4 Aug 2026 18:03:16 -0700 Subject: [PATCH 02/11] test(caching): align closer tests with self-healing handlers The re-landed closer test asserted a reaped handler's client stays closed; with #35862 the handler heals on next access, so the test now pins the inner client up front and asserts the heal as the contract. Also adds an end-to-end regression test that evicts an init-held handler through LLMClientCache, waits out the grace close, and proves the next request succeeds. --- .../caching/test_evicted_client_closer.py | 6 ++-- .../llms/custom_httpx/test_http_handler.py | 34 +++++++++++++++++++ 2 files changed, 38 insertions(+), 2 deletions(-) diff --git a/tests/test_litellm/caching/test_evicted_client_closer.py b/tests/test_litellm/caching/test_evicted_client_closer.py index a08fd58079d..939be5f3d6b 100644 --- a/tests/test_litellm/caching/test_evicted_client_closer.py +++ b/tests/test_litellm/caching/test_evicted_client_closer.py @@ -334,6 +334,7 @@ async def test_the_aiohttp_backed_handler_is_not_closed_mid_request(): clock = FakeClock() closer = make_closer(clock) handler = AsyncHTTPHandler() + held_client = handler.client closer.mark_owned(handler) closer.schedule(handler) @@ -344,14 +345,15 @@ async def test_the_aiohttp_backed_handler_is_not_closed_mid_request(): closer.reap() await asyncio.sleep(0.05) - assert handler.client.is_closed is False, "closed a handler that was serving a request" + assert held_client.is_closed is False, "closed a handler that was serving a request" assert (await request).status_code == 200 clock.advance(3600.0) closer.reap() await asyncio.sleep(0.05) - assert handler.client.is_closed is True + assert held_client.is_closed is True + assert handler.client.is_closed is False, "a held handler must self-heal after its evicted client is closed" server.close() diff --git a/tests/test_litellm/llms/custom_httpx/test_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_http_handler.py index 35db698ad76..b4921558ded 100644 --- a/tests/test_litellm/llms/custom_httpx/test_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_http_handler.py @@ -838,6 +838,40 @@ async def test_init_held_async_handler_survives_external_client_close(): await handler.close() +@pytest.mark.asyncio +async def test_init_held_async_handler_survives_evicted_client_close(): + from litellm.caching.evicted_client_closer import EvictedClientCloser + from litellm.caching.llm_caching_handler import LLMClientCache + + cache = LLMClientCache(evicted_client_closer=EvictedClientCloser(grace_seconds=0)) + handler = AsyncHTTPHandler(timeout=42.5) + held_client = handler.client + cache.set_cache("init-held-handler", handler, litellm_owned_client=True, ttl=0) + await asyncio.sleep(0.02) + assert cache.get_cache("init-held-handler") is None + await asyncio.sleep(0.05) + assert held_client.is_closed + + async def respond(reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None: + await _read_http_request(reader) + writer.write(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nok") + await writer.drain() + writer.close() + + server = await asyncio.start_server(respond, "127.0.0.1", 0) + port = server.sockets[0].getsockname()[1] + try: + response = await handler.post(f"http://127.0.0.1:{port}/v1/compress", json={"messages": []}) + finally: + server.close() + await server.wait_closed() + + assert response.status_code == 200 + assert handler.client is not held_client + assert handler.client.timeout == httpx.Timeout(42.5) + await handler.close() + + def test_init_held_sync_handler_recreates_closed_client(): from http.server import BaseHTTPRequestHandler, HTTPServer From 9ee23a20996238dedcdc970aea633715113dd967 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 4 Aug 2026 18:12:59 -0700 Subject: [PATCH 03/11] fix(azure): return a provided client early in create_azure_client The restored #35492 code referenced azure_client_params outside the branch that binds it, guarded only by a client-is-None short circuit. That is safe at runtime but basedpyright cannot correlate the two checks, so the ratchet gate flagged it as a net-new possibly-unbound reference. Early-returning the provided-client path leaves azure_client_params bound on every path that reaches the ownership check and removes the need for the short circuit. --- litellm/llms/azure/common_utils.py | 153 +++++++++++++++-------------- 1 file changed, 79 insertions(+), 74 deletions(-) diff --git a/litellm/llms/azure/common_utils.py b/litellm/llms/azure/common_utils.py index 25dd9698624..b48820d75c4 100644 --- a/litellm/llms/azure/common_utils.py +++ b/litellm/llms/azure/common_utils.py @@ -427,90 +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) # type: ignore - else: - openai_client = OpenAI(**v1_params) # type: ignore - 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) # type: ignore - 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) # type: ignore + else: + openai_client = OpenAI(**v1_params) # type: ignore + 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) # type: ignore # 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=client is None - and self.owns_wrapped_http_client(azure_client_params.get("http_client")), + litellm_owned_client=self.owns_wrapped_http_client(azure_client_params.get("http_client")), ) return openai_client From 469d5126f69aac6e8cd9eb7d8c3346dc5f357e04 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 5 Aug 2026 10:23:02 -0700 Subject: [PATCH 04/11] fix(lint): bring basedpyright rule counts back under their budget limits --- .../enterprise_callbacks/__init__.py | 0 .../pagerduty/__init__.py | 0 .../send_emails/__init__.py | 0 .../integrations/__init__.py | 0 .../litellm_core_utils/__init__.py | 0 .../proxy/hooks/__init__.py | 0 .../proxy/vector_stores/__init__.py | 0 enterprise/litellm_enterprise/py.typed | 0 .../litellm_enterprise/types/__init__.py | 0 .../types/enterprise_callbacks/__init__.py | 0 .../types/proxy/__init__.py | 0 .../litellm_proxy_extras/py.typed | 0 litellm/__init__.py | 36 +++++++++---------- litellm/_lazy_imports.py | 10 +++--- .../litellm_core_utils/audio_utils/utils.py | 3 +- .../prompt_templates/factory.py | 24 ++++++------- .../bedrock/chat/converse_transformation.py | 10 +++--- litellm/llms/xai/oauth.py | 2 +- .../mcp_server/elicitation_handler.py | 11 +++++- .../mcp_server/sampling_handler.py | 10 +++++- .../proxy/_experimental/mcp_server/server.py | 10 +++--- .../example_config_yaml/custom_guardrail.py | 3 +- .../example_config_yaml/custom_handler.py | 6 ++-- litellm/types/adapter.py | 4 +-- litellm/types/google_genai/main.py | 4 +-- litellm/types/integrations/argilla.py | 5 +-- litellm/types/llms/anthropic_skills.py | 6 ++-- litellm/types/llms/azure_ai.py | 2 +- litellm/types/llms/custom_llm.py | 4 +-- litellm/types/llms/databricks.py | 11 ++---- litellm/types/llms/ollama.py | 10 +----- litellm/types/llms/openrouter.py | 4 +-- litellm/types/llms/rerank.py | 11 +----- .../internal_user_endpoints.py | 5 ++- 34 files changed, 86 insertions(+), 105 deletions(-) create mode 100644 enterprise/litellm_enterprise/enterprise_callbacks/__init__.py create mode 100644 enterprise/litellm_enterprise/enterprise_callbacks/pagerduty/__init__.py create mode 100644 enterprise/litellm_enterprise/enterprise_callbacks/send_emails/__init__.py create mode 100644 enterprise/litellm_enterprise/integrations/__init__.py create mode 100644 enterprise/litellm_enterprise/litellm_core_utils/__init__.py create mode 100644 enterprise/litellm_enterprise/proxy/hooks/__init__.py create mode 100644 enterprise/litellm_enterprise/proxy/vector_stores/__init__.py create mode 100644 enterprise/litellm_enterprise/py.typed create mode 100644 enterprise/litellm_enterprise/types/__init__.py create mode 100644 enterprise/litellm_enterprise/types/enterprise_callbacks/__init__.py create mode 100644 enterprise/litellm_enterprise/types/proxy/__init__.py create mode 100644 litellm-proxy-extras/litellm_proxy_extras/py.typed diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/__init__.py b/enterprise/litellm_enterprise/enterprise_callbacks/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/pagerduty/__init__.py b/enterprise/litellm_enterprise/enterprise_callbacks/pagerduty/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/__init__.py b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/enterprise/litellm_enterprise/integrations/__init__.py b/enterprise/litellm_enterprise/integrations/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/enterprise/litellm_enterprise/litellm_core_utils/__init__.py b/enterprise/litellm_enterprise/litellm_core_utils/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/enterprise/litellm_enterprise/proxy/hooks/__init__.py b/enterprise/litellm_enterprise/proxy/hooks/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/enterprise/litellm_enterprise/proxy/vector_stores/__init__.py b/enterprise/litellm_enterprise/proxy/vector_stores/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/enterprise/litellm_enterprise/py.typed b/enterprise/litellm_enterprise/py.typed new file mode 100644 index 00000000000..e69de29bb2d diff --git a/enterprise/litellm_enterprise/types/__init__.py b/enterprise/litellm_enterprise/types/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/enterprise/litellm_enterprise/types/enterprise_callbacks/__init__.py b/enterprise/litellm_enterprise/types/enterprise_callbacks/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/enterprise/litellm_enterprise/types/proxy/__init__.py b/enterprise/litellm_enterprise/types/proxy/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm-proxy-extras/litellm_proxy_extras/py.typed b/litellm-proxy-extras/litellm_proxy_extras/py.typed new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/__init__.py b/litellm/__init__.py index 319da4e25eb..89310120768 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -2150,9 +2150,9 @@ def __getattr__(name: str) -> Any: # Lazy load encoding from main.py to avoid heavy tiktoken import if name == "encoding": - from ._lazy_imports import _get_litellm_globals + from ._lazy_imports import get_litellm_globals - _globals = _get_litellm_globals() + _globals = get_litellm_globals() # Check if already cached if "encoding" not in _globals: from .main import encoding as _encoding @@ -2162,9 +2162,9 @@ def __getattr__(name: str) -> Any: # Lazy load bedrock_tool_name_mappings instance if name == "bedrock_tool_name_mappings": - from ._lazy_imports import _get_litellm_globals + from ._lazy_imports import get_litellm_globals - _globals = _get_litellm_globals() + _globals = get_litellm_globals() # Check if already cached if "bedrock_tool_name_mappings" not in _globals: from .llms.bedrock.chat.invoke_handler import ( @@ -2176,9 +2176,9 @@ def __getattr__(name: str) -> Any: # Lazy load AzureOpenAIError exception class if name == "AzureOpenAIError": - from ._lazy_imports import _get_litellm_globals + from ._lazy_imports import get_litellm_globals - _globals = _get_litellm_globals() + _globals = get_litellm_globals() # Check if already cached if "AzureOpenAIError" not in _globals: from .llms.azure.common_utils import AzureOpenAIError as _AzureOpenAIError @@ -2188,9 +2188,9 @@ def __getattr__(name: str) -> Any: # Lazy load openaiOSeriesConfig instance if name == "openaiOSeriesConfig": - from ._lazy_imports import _get_litellm_globals + from ._lazy_imports import get_litellm_globals - _globals = _get_litellm_globals() + _globals = get_litellm_globals() if "openaiOSeriesConfig" not in _globals: # Import the config class and instantiate it config_class = __getattr__("OpenAIOSeriesConfig") @@ -2206,9 +2206,9 @@ def __getattr__(name: str) -> Any: "nvidiaNimEmbeddingConfig": "NvidiaNimEmbeddingConfig", } if name in _config_instances: - from ._lazy_imports import _get_litellm_globals + from ._lazy_imports import get_litellm_globals - _globals = _get_litellm_globals() + _globals = get_litellm_globals() if name not in _globals: # Import the config class and instantiate it config_class = __getattr__(_config_instances[name]) @@ -2221,9 +2221,9 @@ def __getattr__(name: str) -> Any: # Lazy load provider_list if name == "provider_list": - from ._lazy_imports import _get_litellm_globals + from ._lazy_imports import get_litellm_globals - _globals = _get_litellm_globals() + _globals = get_litellm_globals() # Check if already cached if "provider_list" not in _globals: # LlmProviders is eagerly imported above, so we can import it directly @@ -2234,9 +2234,9 @@ def __getattr__(name: str) -> Any: # Lazy load priority_reservation_settings instance if name == "priority_reservation_settings": - from ._lazy_imports import _get_litellm_globals + from ._lazy_imports import get_litellm_globals - _globals = _get_litellm_globals() + _globals = get_litellm_globals() # Check if already cached if "priority_reservation_settings" not in _globals: # Import the class and instantiate it @@ -2246,9 +2246,9 @@ def __getattr__(name: str) -> Any: # Lazy load logging_callback_manager instance if name == "logging_callback_manager": - from ._lazy_imports import _get_litellm_globals + from ._lazy_imports import get_litellm_globals - _globals = _get_litellm_globals() + _globals = get_litellm_globals() # Check if already cached if "logging_callback_manager" not in _globals: # Import the class and instantiate it @@ -2258,9 +2258,9 @@ def __getattr__(name: str) -> Any: # Lazy load _service_logger module if name == "_service_logger": - from ._lazy_imports import _get_litellm_globals + from ._lazy_imports import get_litellm_globals - _globals = _get_litellm_globals() + _globals = get_litellm_globals() # Check if already cached if "_service_logger" not in _globals: # Import the module lazily diff --git a/litellm/_lazy_imports.py b/litellm/_lazy_imports.py index 63142ee4f2f..933464d3f23 100644 --- a/litellm/_lazy_imports.py +++ b/litellm/_lazy_imports.py @@ -54,7 +54,7 @@ from ._lazy_imports_registry import ( ) -def _get_litellm_globals() -> dict: +def get_litellm_globals() -> dict: """ Get the globals dictionary of the litellm module. @@ -233,7 +233,7 @@ def _generic_lazy_import(name: str, import_map: dict[str, tuple[str, str]], cate raise AttributeError(f"{category} lazy import: unknown attribute {name!r}") # Step 2: Get the cache (where we store imported things) - _globals: Final = _get_litellm_globals() + _globals: Final = get_litellm_globals() # Step 3: If we've already imported it, just return the cached version if name in _globals: @@ -332,7 +332,7 @@ def _lazy_import_utils_module(name: str) -> Any: Handler for utils module lazy imports. This uses a custom implementation because utils module needs to use - _get_utils_globals() instead of _get_litellm_globals() for caching. + _get_utils_globals() instead of get_litellm_globals() for caching. """ # Check if this attribute exists in our map if name not in _UTILS_MODULE_IMPORT_MAP: @@ -379,7 +379,7 @@ def _lazy_import_llm_client_cache(name: str) -> Any: - "in_memory_llm_clients_cache" is a singleton instance of that class So we need custom logic to handle both cases. """ - _globals: Final = _get_litellm_globals() + _globals: Final = get_litellm_globals() # If already cached, return it if name in _globals: @@ -412,7 +412,7 @@ def _lazy_import_http_handlers(name: str) -> Any: - They need configuration (timeout, etc.) from the module globals - They use factory functions instead of direct instantiation """ - _globals: Final = _get_litellm_globals() + _globals: Final = get_litellm_globals() if name == "module_level_aclient": # Create an async HTTP client using the factory function diff --git a/litellm/litellm_core_utils/audio_utils/utils.py b/litellm/litellm_core_utils/audio_utils/utils.py index 0f9addb16f7..3b3775a8fe6 100644 --- a/litellm/litellm_core_utils/audio_utils/utils.py +++ b/litellm/litellm_core_utils/audio_utils/utils.py @@ -180,6 +180,7 @@ def get_audio_file_content_hash(file_obj: FileTypes) -> str: if isinstance(file_obj, tuple): if len(file_obj) < 2: fallback_filename = str(file_obj[0]) if len(file_obj) > 0 else None + file_content_obj = None else: fallback_filename = str(file_obj[0]) if file_obj[0] is not None else None file_content_obj = file_obj[1] @@ -206,7 +207,7 @@ def get_audio_file_content_hash(file_obj: FileTypes) -> str: except OSError: fallback_filename = str(file_content_obj) file_content = None - elif hasattr(file_content_obj, "read"): + elif file_content_obj is not None and hasattr(file_content_obj, "read"): try: current_position: Final = file_content_obj.tell() if hasattr(file_content_obj, "tell") else None if hasattr(file_content_obj, "seek"): diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index d51d31eaa3b..3a1a426eaa9 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -3684,7 +3684,7 @@ def _convert_to_bedrock_tool_call_invoke( # cache_control applies to the whole original # tool call; attach after the last split block. if tool.get("cache_control", None) is not None: - _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( + _cache_point_block = litellm.AmazonConverseConfig().get_cache_point_block( {"cache_control": tool["cache_control"]}, block_type="content_block", model=model, @@ -3701,7 +3701,7 @@ def _convert_to_bedrock_tool_call_invoke( # Check for cache_control and add a separate cachePoint block if tool.get("cache_control", None) is not None: - cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( + cache_point_block = litellm.AmazonConverseConfig().get_cache_point_block( {"cache_control": tool["cache_control"]}, block_type="content_block", model=model, @@ -4360,7 +4360,7 @@ class BedrockConverseMessagesProcessor: elif element["type"] == "document": _part = BedrockConverseMessagesProcessor._process_document_message(element) _parts.append(_part) - _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( + _cache_point_block = litellm.AmazonConverseConfig().get_cache_point_block( message_block=cast(OpenAIMessageContentListBlock, element), block_type="content_block", model=model, @@ -4370,7 +4370,7 @@ class BedrockConverseMessagesProcessor: user_content.extend(_parts) elif message_block["content"] and isinstance(message_block["content"], str): _part = BedrockContentBlock(text=messages[msg_i]["content"]) - _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( + _cache_point_block = litellm.AmazonConverseConfig().get_cache_point_block( message_block, block_type="content_block", model=model ) user_content.append(_part) @@ -4417,7 +4417,7 @@ class BedrockConverseMessagesProcessor: # Add a separate cachePoint block if cache_control is present if tool_msg_cache_control is not None: - cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( + cache_point_block = litellm.AmazonConverseConfig().get_cache_point_block( {"cache_control": tool_msg_cache_control}, block_type="content_block", model=model, @@ -4496,7 +4496,7 @@ class BedrockConverseMessagesProcessor: assistants_part = await BedrockImageProcessor.process_image_async(image_url=image_url) assistants_parts.append(assistants_part) # Add cache point block for assistant content elements - _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( + _cache_point_block = litellm.AmazonConverseConfig().get_cache_point_block( message_block=cast(OpenAIMessageContentListBlock, element), block_type="content_block", model=model, @@ -4510,7 +4510,7 @@ class BedrockConverseMessagesProcessor: assistant_content.append(BedrockContentBlock(text=_assistant_content)) # If content is empty/whitespace, skip it (don't add a placeholder) # Add cache point block for assistant string content - _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( + _cache_point_block = litellm.AmazonConverseConfig().get_cache_point_block( assistant_message_block, block_type="content_block", model=model ) if _cache_point_block is not None: @@ -4733,7 +4733,7 @@ def _bedrock_converse_messages_pt( elif element["type"] == "document": _part = BedrockConverseMessagesProcessor._process_document_message(element) _parts.append(_part) - _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( + _cache_point_block = litellm.AmazonConverseConfig().get_cache_point_block( message_block=cast(OpenAIMessageContentListBlock, element), block_type="content_block", model=model, @@ -4743,7 +4743,7 @@ def _bedrock_converse_messages_pt( user_content.extend(_parts) elif message_block["content"] and isinstance(message_block["content"], str): _part = BedrockContentBlock(text=messages[msg_i]["content"]) - _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( + _cache_point_block = litellm.AmazonConverseConfig().get_cache_point_block( message_block, block_type="content_block", model=model ) user_content.append(_part) @@ -4792,7 +4792,7 @@ def _bedrock_converse_messages_pt( # Add a separate cachePoint block if cache_control is present if tool_msg_cache_control is not None: - cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( + cache_point_block = litellm.AmazonConverseConfig().get_cache_point_block( {"cache_control": tool_msg_cache_control}, block_type="content_block", model=model, @@ -4874,7 +4874,7 @@ def _bedrock_converse_messages_pt( assistants_part = BedrockImageProcessor.process_image_sync(image_url=image_url) assistants_parts.append(assistants_part) # Add cache point block for assistant content elements - _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( + _cache_point_block = litellm.AmazonConverseConfig().get_cache_point_block( message_block=cast(OpenAIMessageContentListBlock, element), block_type="content_block", model=model, @@ -4887,7 +4887,7 @@ def _bedrock_converse_messages_pt( if _assistant_content.strip(): assistant_content.append(BedrockContentBlock(text=_assistant_content)) # Add cache point block for assistant string content - _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( + _cache_point_block = litellm.AmazonConverseConfig().get_cache_point_block( assistant_message_block, block_type="content_block", model=model ) if _cache_point_block is not None: diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 0b1689b8ee4..193987a3543 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -1081,7 +1081,7 @@ class AmazonConverseConfig(BaseConfig): optional_params["maxTokens"] = thinking_token_budget + DEFAULT_MAX_TOKENS @overload - def _get_cache_point_block( + def get_cache_point_block( self, message_block: OpenAIMessageContentListBlock | ChatCompletionUserMessage @@ -1093,7 +1093,7 @@ class AmazonConverseConfig(BaseConfig): pass @overload - def _get_cache_point_block( + def get_cache_point_block( self, message_block: OpenAIMessageContentListBlock | ChatCompletionUserMessage @@ -1104,7 +1104,7 @@ class AmazonConverseConfig(BaseConfig): ) -> ContentBlock | None: pass - def _get_cache_point_block( + def get_cache_point_block( self, message_block: OpenAIMessageContentListBlock | ChatCompletionUserMessage @@ -1149,14 +1149,14 @@ class AmazonConverseConfig(BaseConfig): system_prompt_indices.append(idx) if isinstance(message["content"], str) and message["content"]: system_content_blocks.append(SystemContentBlock(text=message["content"])) - cache_block = self._get_cache_point_block(message, block_type="system", model=model) + cache_block = self.get_cache_point_block(message, block_type="system", model=model) if cache_block: system_content_blocks.append(cache_block) elif isinstance(message["content"], list): for m in message["content"]: if m.get("type") == "text" and m.get("text"): system_content_blocks.append(SystemContentBlock(text=m["text"])) - cache_block = self._get_cache_point_block(m, block_type="system", model=model) + cache_block = self.get_cache_point_block(m, block_type="system", model=model) if cache_block: system_content_blocks.append(cache_block) if len(system_prompt_indices) > 0: diff --git a/litellm/llms/xai/oauth.py b/litellm/llms/xai/oauth.py index 8f303e9585f..37dae93a725 100644 --- a/litellm/llms/xai/oauth.py +++ b/litellm/llms/xai/oauth.py @@ -40,7 +40,7 @@ class XAIOAuthLoginRequiredError(XAIOAuthError): class _CallbackHandler(BaseHTTPRequestHandler): - server: "_CallbackServer" + server: "_CallbackServer" # pyright: ignore[reportIncompatibleVariableOverride] # stdlib stubs type server as BaseServer; _CallbackServer is the only server this handler is registered on def do_GET(self) -> None: parsed: Final = urlparse(self.path) diff --git a/litellm/proxy/_experimental/mcp_server/elicitation_handler.py b/litellm/proxy/_experimental/mcp_server/elicitation_handler.py index 66c262a6eb9..ce7e963f55f 100644 --- a/litellm/proxy/_experimental/mcp_server/elicitation_handler.py +++ b/litellm/proxy/_experimental/mcp_server/elicitation_handler.py @@ -9,10 +9,19 @@ MCP Spec Reference: https://modelcontextprotocol.io/specification/2025-11-25/client/elicitation """ -from typing import Any, Final, Union +from typing import TYPE_CHECKING, Any, Final, Union from litellm._logging import verbose_logger +if TYPE_CHECKING: + from mcp.types import ( + ElicitRequestFormParams, + ElicitRequestParams, + ElicitRequestURLParams, + ElicitResult, + ErrorData, + ) + # Guard imports that require the mcp package try: from mcp.types import ( diff --git a/litellm/proxy/_experimental/mcp_server/sampling_handler.py b/litellm/proxy/_experimental/mcp_server/sampling_handler.py index 45490385df8..0f5c02bd781 100644 --- a/litellm/proxy/_experimental/mcp_server/sampling_handler.py +++ b/litellm/proxy/_experimental/mcp_server/sampling_handler.py @@ -18,7 +18,15 @@ if typing.TYPE_CHECKING: from fastapi import Request from mcp.client.session import ClientSession from mcp.shared.context import RequestContext - from mcp.types import ContentBlock, SamplingMessageContentBlock + from mcp.types import ( + ContentBlock, + CreateMessageResult, + CreateMessageResultWithTools, + ErrorData, + SamplingMessageContentBlock, + TextContent, + ToolUseContent, + ) from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.utils import ProxyLogging diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index f2267bbcf7f..ef7bfd4b4f6 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -79,6 +79,8 @@ from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall from litellm.utils import Rules, client, function_setup if TYPE_CHECKING: + from mcp.server.session import ServerSession as _McpServerSession + from litellm.proxy._experimental.mcp_server.db import OAuthCredentialPayload # Short-lived in-memory cache for BYOK credentials. @@ -144,10 +146,6 @@ try: # Robust auth lookup keyed by session_object. _session_obj_auth_storage: "weakref.WeakKeyDictionary[Any, MCPAuthenticatedUser]" = weakref.WeakKeyDictionary() - - active_mcp_session_var: Final[contextvars.ContextVar[_McpServerSession | None]] = contextvars.ContextVar( - "active_mcp_session", default=None - ) except ImportError as e: verbose_logger.debug("MCP module not found: %s", e) MCP_AVAILABLE = False @@ -163,6 +161,10 @@ except ImportError as e: Server = None TextResourceContents = None +active_mcp_session_var: Final[contextvars.ContextVar["_McpServerSession | None"]] = contextvars.ContextVar( + "active_mcp_session", default=None +) + # Global variables to track initialization _SESSION_MANAGERS_INITIALIZED = False diff --git a/litellm/proxy/example_config_yaml/custom_guardrail.py b/litellm/proxy/example_config_yaml/custom_guardrail.py index 2f53bb4675a..979976ddfc4 100644 --- a/litellm/proxy/example_config_yaml/custom_guardrail.py +++ b/litellm/proxy/example_config_yaml/custom_guardrail.py @@ -1,11 +1,10 @@ -from typing import Any, Dict, Final, List, Literal, Optional, Union +from typing import Dict, Final, Optional, Union import litellm from litellm._logging import verbose_proxy_logger from litellm.caching.caching import DualCache from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.guardrails.guardrail_helpers import should_proceed_based_on_metadata from litellm.types.utils import CallTypesLiteral # Global counter for tracking which guardrail was called (for load balancing tests) diff --git a/litellm/proxy/example_config_yaml/custom_handler.py b/litellm/proxy/example_config_yaml/custom_handler.py index 3bf998c726a..c0483dd3304 100644 --- a/litellm/proxy/example_config_yaml/custom_handler.py +++ b/litellm/proxy/example_config_yaml/custom_handler.py @@ -1,9 +1,7 @@ -import time -from typing import Any, Final, Optional +from typing import Final import litellm -from litellm import CustomLLM, ImageObject, ImageResponse, completion, get_llm_provider -from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm import CustomLLM from litellm.types.utils import ModelResponse diff --git a/litellm/types/adapter.py b/litellm/types/adapter.py index 2995cfbc1c2..924fabcb86d 100644 --- a/litellm/types/adapter.py +++ b/litellm/types/adapter.py @@ -1,6 +1,4 @@ -from typing import List - -from typing_extensions import Dict, Required, TypedDict, override +from typing_extensions import TypedDict from litellm.integrations.custom_logger import CustomLogger diff --git a/litellm/types/google_genai/main.py b/litellm/types/google_genai/main.py index 467db318057..876a4d4533e 100644 --- a/litellm/types/google_genai/main.py +++ b/litellm/types/google_genai/main.py @@ -1,8 +1,6 @@ # Import types from the Google GenAI SDK -from typing import TYPE_CHECKING, Any, Dict, List, Optional, TypeAlias +from typing import TYPE_CHECKING, Any, Dict, Optional -from pydantic import BaseModel -from typing_extensions import TypedDict from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject diff --git a/litellm/types/integrations/argilla.py b/litellm/types/integrations/argilla.py index 52dad347304..2def010a722 100644 --- a/litellm/types/integrations/argilla.py +++ b/litellm/types/integrations/argilla.py @@ -1,7 +1,4 @@ -import os -from datetime import datetime as dt -from enum import Enum -from typing import Any, Dict, Final, List, Literal, Optional, Set +from typing import Any, Dict, Final, List from typing_extensions import TypedDict diff --git a/litellm/types/llms/anthropic_skills.py b/litellm/types/llms/anthropic_skills.py index 22257888493..0659b499bcc 100644 --- a/litellm/types/llms/anthropic_skills.py +++ b/litellm/types/llms/anthropic_skills.py @@ -2,10 +2,10 @@ Type definitions for Anthropic Skills API """ -from typing import Any, Dict, List, Literal, Optional, Union +from typing import Any, Dict, List, Optional -from pydantic import BaseModel, Field -from typing_extensions import Required, TypedDict +from pydantic import BaseModel +from typing_extensions import TypedDict # Skills API Request Types diff --git a/litellm/types/llms/azure_ai.py b/litellm/types/llms/azure_ai.py index ddc9dbe3c55..49b7349c67e 100644 --- a/litellm/types/llms/azure_ai.py +++ b/litellm/types/llms/azure_ai.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, Final, Iterable, List, Literal, Optional, Union +from typing import List, Literal from typing_extensions import Required, TypedDict diff --git a/litellm/types/llms/custom_llm.py b/litellm/types/llms/custom_llm.py index d5499a41944..e57a7a28007 100644 --- a/litellm/types/llms/custom_llm.py +++ b/litellm/types/llms/custom_llm.py @@ -1,6 +1,4 @@ -from typing import List - -from typing_extensions import Dict, Required, TypedDict, override +from typing_extensions import TypedDict from litellm.llms.custom_llm import CustomLLM diff --git a/litellm/types/llms/databricks.py b/litellm/types/llms/databricks.py index 46f988ae4a0..c2bd0aa92bd 100644 --- a/litellm/types/llms/databricks.py +++ b/litellm/types/llms/databricks.py @@ -1,19 +1,12 @@ -import json -from typing import Any, Dict, Final, List, Literal, Optional, Union +from typing import Any, Dict, List, Literal, Optional, Union from pydantic import BaseModel from typing_extensions import ( - Protocol, Required, - Self, TypedDict, - TypeGuard, - get_origin, - override, - runtime_checkable, ) -from .openai import ChatCompletionToolCallChunk, ChatCompletionUsageBlock +from .openai import ChatCompletionUsageBlock class GenericStreamingChunk(TypedDict, total=False): diff --git a/litellm/types/llms/ollama.py b/litellm/types/llms/ollama.py index ca28120dd9d..9fcb6b755bd 100644 --- a/litellm/types/llms/ollama.py +++ b/litellm/types/llms/ollama.py @@ -1,16 +1,8 @@ -import json -from typing import Any, List, Optional, Union +from typing import List -from pydantic import BaseModel from typing_extensions import ( - Protocol, Required, - Self, TypedDict, - TypeGuard, - get_origin, - override, - runtime_checkable, ) diff --git a/litellm/types/llms/openrouter.py b/litellm/types/llms/openrouter.py index 39ed7e104fb..73bf647d4ea 100644 --- a/litellm/types/llms/openrouter.py +++ b/litellm/types/llms/openrouter.py @@ -1,6 +1,4 @@ -import json -from enum import Enum -from typing import Any, Dict, List, Literal, Optional, Tuple, Union +from typing import Dict from typing_extensions import TypedDict diff --git a/litellm/types/llms/rerank.py b/litellm/types/llms/rerank.py index fac093161c1..83cdb1caa0b 100644 --- a/litellm/types/llms/rerank.py +++ b/litellm/types/llms/rerank.py @@ -1,16 +1,7 @@ -import json -from enum import Enum -from typing import Any, Dict, List, Literal, Optional, Tuple, Union +from typing import Optional from typing_extensions import ( - Protocol, - Required, - Self, TypedDict, - TypeGuard, - get_origin, - override, - runtime_checkable, ) diff --git a/litellm/types/proxy/management_endpoints/internal_user_endpoints.py b/litellm/types/proxy/management_endpoints/internal_user_endpoints.py index faf2660a6f8..16f4c45e2f4 100644 --- a/litellm/types/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/types/proxy/management_endpoints/internal_user_endpoints.py @@ -1,7 +1,6 @@ -from typing import Any, Dict, Final, List, Literal, Optional, Union +from typing import Any, Dict, Final, List, Optional -from fastapi import HTTPException -from pydantic import BaseModel, EmailStr, field_validator +from pydantic import BaseModel, field_validator from litellm.proxy._types import ( LiteLLM_UserTableWithKeyCount, From 64f83a23e1874ef98d06302279edba59865f6e8f Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 5 Aug 2026 10:44:02 -0700 Subject: [PATCH 05/11] Revert "chore(ui): zero stale headroom on local dashboard eslint budgets" --- ui/litellm-dashboard/eslint-budgets.json | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/ui/litellm-dashboard/eslint-budgets.json b/ui/litellm-dashboard/eslint-budgets.json index 3526d71ce90..f08e1bb6160 100644 --- a/ui/litellm-dashboard/eslint-budgets.json +++ b/ui/litellm-dashboard/eslint-budgets.json @@ -1,8 +1,8 @@ { "@typescript-eslint/no-explicit-any": { "max": 2040, "target": 1500 }, - "no-console": { "max": 12, "target": 0 }, - "complexity": { "max": 121, "target": 80 }, - "max-depth": { "max": 55, "target": 30 }, - "local/no-large-inline-object-arg": { "max": 469, "target": 300 }, - "local/no-long-condition-chain": { "max": 217, "target": 120 } + "no-console": { "max": 484, "target": 0 }, + "complexity": { "max": 140, "target": 80 }, + "max-depth": { "max": 70, "target": 30 }, + "local/no-large-inline-object-arg": { "max": 560, "target": 300 }, + "local/no-long-condition-chain": { "max": 265, "target": 120 } } From 85aad29885fcf556b0e2dbf202c03167e87fec90 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 5 Aug 2026 10:45:56 -0700 Subject: [PATCH 06/11] chore: make no-console max 12 --- ui/litellm-dashboard/eslint-budgets.json | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ui/litellm-dashboard/eslint-budgets.json b/ui/litellm-dashboard/eslint-budgets.json index f08e1bb6160..c4f078f2ff2 100644 --- a/ui/litellm-dashboard/eslint-budgets.json +++ b/ui/litellm-dashboard/eslint-budgets.json @@ -1,6 +1,6 @@ { "@typescript-eslint/no-explicit-any": { "max": 2040, "target": 1500 }, - "no-console": { "max": 484, "target": 0 }, + "no-console": { "max": 12, "target": 0 }, "complexity": { "max": 140, "target": 80 }, "max-depth": { "max": 70, "target": 30 }, "local/no-large-inline-object-arg": { "max": 560, "target": 300 }, From ebc31f7e066e7966b4c37edbfdb2d2635272fb58 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 5 Aug 2026 11:08:19 -0700 Subject: [PATCH 07/11] fix(caching): clear strict-lint budget breaches in re-landed closer code --- litellm/caching/evicted_client_closer.py | 5 ++--- litellm/caching/llm_caching_handler.py | 2 +- 2 files changed, 3 insertions(+), 4 deletions(-) diff --git a/litellm/caching/evicted_client_closer.py b/litellm/caching/evicted_client_closer.py index c895669be2b..eee7e2ea289 100644 --- a/litellm/caching/evicted_client_closer.py +++ b/litellm/caching/evicted_client_closer.py @@ -34,6 +34,7 @@ alive anything the collector would have reclaimed first. """ import asyncio +import contextlib import inspect import threading import time @@ -146,10 +147,8 @@ def _has_connection_in_flight(client: object) -> bool: async def _close_quietly(closing: Awaitable[object]) -> None: - try: + with contextlib.suppress(Exception): await closing - except Exception: # noqa: BLE001 - a discarded client's close must never surface to callers - pass class EvictedClientCloser: diff --git a/litellm/caching/llm_caching_handler.py b/litellm/caching/llm_caching_handler.py index 6fa5963c99b..a89e43b78b4 100644 --- a/litellm/caching/llm_caching_handler.py +++ b/litellm/caching/llm_caching_handler.py @@ -29,7 +29,7 @@ class LLMClientCache(InMemoryCache): 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, From 2792887e47d698edaeb8a2d0ad1abfc9576609a6 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 5 Aug 2026 11:33:55 -0700 Subject: [PATCH 08/11] fix(proxy): give proxy_admin_viewer read parity with proxy_admin (#35851) * fix(proxy): give proxy_admin_viewer read parity with proxy_admin Route-level checks already default-allow management GETs for the viewer role, but ~15 handlers compared user_role to PROXY_ADMIN only, dropping viewers into regular-user scoping (/key/list, /user/info, /model/info, guardrails, prompts, agents, memory, workflows, MCP catalog, coordination redis settings, credential migration check, enterprise projects). Swap those read paths to user_api_key_has_admin_view; write gates unchanged. The dashboard now presents the viewer session as Admin for all gating (effectiveSessionRole) so every page fetches with admin visibility, with userRoleLabel/isViewOnly preserving the account-menu label and the playground cost guard. The server remains the write authority. * refactor(agents): remove side-effectful health_check param from GET /v1/agents Addresses a security review finding on the admin viewer read parity change: listing agents with health_check=true made the proxy issue a server-side GET to every agent URL, so a read-scoped caller could trigger request fan-out beyond their object permissions. The list endpoint is now a pure read for every role. Removes the query param, the URL probing helper and its timeouts, the AgentHealthCheck httpx provider tag, and the dashboard's Health Check toggle. Requests still passing health_check=true get the full list back with the param ignored. * fix(proxy): keep credential encryption check proxy_admin only The residual scan behind GET /credentials/migrate-encryption/check loads every model, credential, MCP, team, and verification-token row and runs a decryption attempt on each stored value. Extending it to proxy_admin_viewer let a read-only account repeatedly trigger deployment-wide scans, so the route keeps its original full-admin gate. * fix(agents): restore health_check, keep list fast path proxy_admin only Restores the agent health_check feature exactly as before this PR: the query param, the URL probing helper, the httpx provider tag, and the dashboard toggle all return, so existing callers keep the filtering contract. The viewer expansion is instead reverted at its source: the GET /v1/agents admin fast path stays PROXY_ADMIN only, so a proxy_admin_viewer goes through the object-permission scoped branch as before and cannot fan out health checks beyond their allowlist. The viewer read of a single agent stays viewer-inclusive since it has no side effects. --- .../management_endpoints/project_endpoints.py | 4 +- .../mcp_server/rest_endpoints.py | 10 +- litellm/proxy/agent_endpoints/endpoints.py | 13 +- litellm/proxy/auth/auth_checks.py | 3 - litellm/proxy/auth/route_checks.py | 6 +- .../proxy/guardrails/guardrail_endpoints.py | 4 +- .../coordination_redis_endpoints.py | 4 +- .../internal_user_endpoints.py | 11 +- .../key_management_endpoints.py | 3 +- .../workflow_management_endpoints.py | 19 ++- litellm/proxy/memory/memory_endpoints.py | 3 +- litellm/proxy/prompts/prompt_endpoints.py | 25 ++-- litellm/proxy/proxy_server.py | 4 +- .../proxy/agent_endpoints/test_endpoints.py | 87 +++++++++++ .../proxy/auth/test_auth_checks.py | 22 +++ .../proxy/auth/test_route_checks.py | 54 +++++++ .../guardrails/test_guardrail_endpoints.py | 136 +++++++++++++++++ .../test_coordination_redis_endpoints.py | 44 ++++++ .../test_internal_user_endpoints.py | 47 ++++-- .../test_key_management_endpoints.py | 119 +++++++++++++++ .../test_workflow_management_endpoints.py | 132 +++++++++++++++- .../proxy/memory/test_memory_endpoints.py | 74 ++++++++- .../proxy/prompts/test_prompt_endpoints.py | 141 ++++++++++++++++++ .../test_team_model_name_translation.py | 42 ++++++ .../(dashboard)/hooks/useAuthorized.test.ts | 40 +++++ .../app/(dashboard)/hooks/useAuthorized.ts | 6 +- .../app/(dashboard)/playground/page.test.tsx | 1 + .../src/app/(dashboard)/playground/page.tsx | 5 +- .../Navbar/UserDropdown/UserDropdown.test.tsx | 12 +- .../Navbar/UserDropdown/UserDropdown.tsx | 2 +- .../SidebarAccountMenu.test.tsx | 12 +- .../SidebarAccountMenu/SidebarAccountMenu.tsx | 2 +- .../src/components/leftnav.test.tsx | 10 +- .../src/components/leftnav.tsx | 3 +- .../src/components/user_dashboard.tsx | 28 +--- .../src/contexts/AuthContext.tsx | 4 +- ui/litellm-dashboard/src/utils/roles.test.ts | 64 ++++++++ ui/litellm-dashboard/src/utils/roles.ts | 12 ++ 38 files changed, 1094 insertions(+), 114 deletions(-) diff --git a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py index 9d668985eb8..1f693526d1f 100644 --- a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py +++ b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py @@ -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( diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 3fb8e6fe9bb..76618e0f742 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -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 diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py index f729d422d1d..1f9c6e1cc05 100644 --- a/litellm/proxy/agent_endpoints/endpoints.py +++ b/litellm/proxy/agent_endpoints/endpoints.py @@ -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, ) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 5ef4eb471ad..f17c9fff31a 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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 diff --git a/litellm/proxy/auth/route_checks.py b/litellm/proxy/auth/route_checks.py index 8a34438b141..04eb7ab326b 100644 --- a/litellm/proxy/auth/route_checks.py +++ b/litellm/proxy/auth/route_checks.py @@ -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}", diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index 79c4362055e..aef5f2deac4 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -212,7 +212,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 +944,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}) diff --git a/litellm/proxy/management_endpoints/coordination_redis_endpoints.py b/litellm/proxy/management_endpoints/coordination_redis_endpoints.py index 2cfb5cd8793..fe9a613656d 100644 --- a/litellm/proxy/management_endpoints/coordination_redis_endpoints.py +++ b/litellm/proxy/management_endpoints/coordination_redis_endpoints.py @@ -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) diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 97b4ec76c50..cefc7371ce6 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -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 diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index e4def45892b..068429890c8 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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 ( @@ -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: diff --git a/litellm/proxy/management_endpoints/workflow_management_endpoints.py b/litellm/proxy/management_endpoints/workflow_management_endpoints.py index 7e2c7404199..70a6cc507f5 100644 --- a/litellm/proxy/management_endpoints/workflow_management_endpoints.py +++ b/litellm/proxy/management_endpoints/workflow_management_endpoints.py @@ -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( diff --git a/litellm/proxy/memory/memory_endpoints.py b/litellm/proxy/memory/memory_endpoints.py index 33d131bf3b2..987823d987f 100644 --- a/litellm/proxy/memory/memory_endpoints.py +++ b/litellm/proxy/memory/memory_endpoints.py @@ -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: diff --git a/litellm/proxy/prompts/prompt_endpoints.py b/litellm/proxy/prompts/prompt_endpoints.py index 4ac88f87596..d8e9f8dfaee 100644 --- a/litellm/proxy/prompts/prompt_endpoints.py +++ b/litellm/proxy/prompts/prompt_endpoints.py @@ -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}", diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 3cb2f795c61..61c6ce22a91 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -8876,7 +8876,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 +11479,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: diff --git a/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py index bcd3333baf9..3e097711ad7 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py @@ -368,6 +368,17 @@ class TestAgentByIdKeyRedaction: assert resp.status_code == 200 assert resp.json()["keys"] is None + def test_view_only_admin_reads_a_denied_agent_but_still_without_keys(self): + """proxy_admin_viewer skips the per-agent object_permission gate (denied + here) yet stays on the redacted response path.""" + with patch( + "litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.is_agent_allowed", + AsyncMock(return_value=False), + ): + resp = self._get_as(LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY) + assert resp.status_code == 200 + assert resp.json()["keys"] is None + # ---------- RBAC enforcement tests ---------- @@ -469,6 +480,82 @@ class TestAgentRBACInternalUserViewOnly: assert resp.status_code == 403 +class TestAgentRBACProxyAdminViewOnly: + """Read-only proxy admins go through the object-permission scoped branch on + GET /v1/agents (the admin fast path stays full PROXY_ADMIN only, so viewers + cannot fan out health checks beyond their allowlist), and secret unredaction + also stays gated on full PROXY_ADMIN.""" + + @pytest.fixture(autouse=True) + def _setup(self, monkeypatch): + from litellm.proxy.agent_endpoints import agent_registry as ar_mod + + self.viewer_client = _make_app_with_role(LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY) + self.admin_client = _make_app_with_role(LitellmUserRoles.PROXY_ADMIN) + self.agents = [ + AgentResponse( + agent_id=f"agent-{index}", + agent_name=f"Agent {index}", + agent_card_params=_sample_agent_card_params(), + litellm_params={"api_key": "sk-super-secret-agent-key"}, + ) + for index in (1, 2) + ] + self.mock_registry = MagicMock() + self.mock_registry.get_agent_list = MagicMock(return_value=self.agents) + monkeypatch.setattr(ar_mod, "global_agent_registry", self.mock_registry) + + self.allowed_agents_spy = AsyncMock(return_value=["someone-elses-agent"]) + monkeypatch.setattr( + "litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.get_allowed_agents", + self.allowed_agents_spy, + ) + + def _list_agents(self, test_client: TestClient): + key_row = MagicMock() + key_row.token = "hash-aaa" + key_row.agent_id = "agent-1" + key_row.key_alias = "primary" + key_row.key_name = "sk-...aaa" + + with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.db.litellm_agentstable.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[key_row] + ) + return test_client.get("/v1/agents", headers={"Authorization": "Bearer k"}) + + def test_should_scope_view_only_admin_to_allowed_agents(self): + """The key/team allowlist here excludes every registered agent; a viewer + on the admin fast path would see everything, so an empty response pins + that viewers stay in the scoped branch.""" + resp = self._list_agents(self.viewer_client) + + assert resp.status_code == 200 + assert resp.json() == [] + self.allowed_agents_spy.assert_awaited_once() + + def test_should_still_redact_secrets_for_view_only_admin(self): + """An unrestricted viewer (empty allowlist means no restrictions) sees the + same agents as an admin but with keys stripped and litellm_params masked.""" + self.allowed_agents_spy.return_value = [] + viewer_resp = self._list_agents(self.viewer_client) + admin_resp = self._list_agents(self.admin_client) + + assert viewer_resp.status_code == 200 + viewer_by_id = {agent["agent_id"]: agent for agent in viewer_resp.json()} + assert set(viewer_by_id) == {"agent-1", "agent-2"} + assert viewer_by_id["agent-1"]["keys"] is None + assert "sk-super-secret-agent-key" not in viewer_resp.text + + admin_by_id = {agent["agent_id"]: agent for agent in admin_resp.json()} + assert admin_by_id["agent-1"]["keys"][0]["token"] == "hash-aaa" + assert ( + admin_by_id["agent-1"]["litellm_params"]["api_key"] + == "sk-super-secret-agent-key" + ) + + class TestAgentRBACProxyAdmin: """Proxy admins should have full CRUD access to agents.""" diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index d1b5395c73d..a5211ba83e7 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -5462,3 +5462,25 @@ async def test_get_project_object_db_fetch_returns_cached_obj(): assert isinstance(result, LiteLLM_ProjectTableCachedObj) assert result.project_id == "p-1" assert result.project_alias == "proj" + + +def test_is_user_proxy_admin_rejects_view_only_admin(): + """This predicate skips `non_proxy_admin_allowed_routes_check` entirely, so an + Admin Viewer answering True here would gain every write route. Read parity for + that role belongs in the route checks, never here.""" + from litellm.proxy.auth.auth_checks import _is_user_proxy_admin + + viewer = LiteLLM_UserTable( + user_id="viewer_user", + user_email="viewer@example.com", + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + ) + admin = LiteLLM_UserTable( + user_id="admin_user", + user_email="admin@example.com", + user_role=LitellmUserRoles.PROXY_ADMIN.value, + ) + + assert _is_user_proxy_admin(user_obj=viewer) is False + assert _is_user_proxy_admin(user_obj=admin) is True + assert _is_user_proxy_admin(user_obj=None) is False diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index 87f5187b5a1..9285b997efc 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -3192,3 +3192,57 @@ def test_internal_user_blocked_from_search_tool_writes(route): assert "Only proxy admin" in str(exc_info.value) assert f"Route={route}" in str(exc_info.value) assert "Your role=internal_user" in str(exc_info.value) + + +def test_proxy_admin_viewer_can_read_another_users_info(): + """Admin Viewer has read parity with Proxy Admin, so the /user/info + key-ownership gate must not apply to it โ€” the Users page reads every row.""" + user_obj = LiteLLM_UserTable( + user_id="viewer_user", + user_email="viewer@example.com", + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + ) + valid_token = UserAPIKeyAuth( + user_id="viewer_user", + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + ) + request = MagicMock(spec=Request) + request.query_params = {"user_id": "some_other_user"} + + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=user_obj, + _user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + route="/user/info", + request=request, + valid_token=valid_token, + request_data={}, + ) + + +def test_internal_user_still_blocked_from_another_users_info(): + """The Admin Viewer carve-out above must stay scoped to that role; internal + users keep hitting the ownership 403.""" + user_obj = LiteLLM_UserTable( + user_id="internal_user", + user_email="user@example.com", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + valid_token = UserAPIKeyAuth( + user_id="internal_user", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + request = MagicMock(spec=Request) + request.query_params = {"user_id": "some_other_user"} + + with pytest.raises(HTTPException) as exc_info: + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=user_obj, + _user_role=LitellmUserRoles.INTERNAL_USER.value, + route="/user/info", + request=request, + valid_token=valid_token, + request_data={}, + ) + + assert exc_info.value.status_code == 403 + assert "key not allowed to access this user's info" in str(exc_info.value.detail) diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index 1c452e2fb6c..e1dd6b7d48b 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -339,6 +339,109 @@ async def test_list_guardrails_v2_masks_sensitive_data_in_config_guardrails(mock assert params["mode"] == "during_call" +@pytest.mark.asyncio +async def test_list_guardrails_v2_admin_viewer_sees_guardrails_of_teams_they_are_not_in( + mocker, +): + """ + proxy_admin_viewer reads the same unscoped list as proxy_admin: a team-owned + guardrail must surface even though the viewer belongs to no teams. + """ + other_team_guardrail = { + "guardrail_id": "other-team-guardrail", + "guardrail_name": "Other Team Guardrail", + "litellm_params": {"guardrail": "bedrock", "mode": "pre_call"}, + "guardrail_info": {"description": "owned by a team the viewer is not in"}, + "team_id": "team-viewer-is-not-in", + "created_at": datetime.now(), + "updated_at": datetime.now(), + } + + mock_prisma_client = mocker.Mock() + mock_prisma_client.db = mocker.Mock() + mock_prisma_client.db.litellm_guardrailstable = mocker.Mock() + mock_prisma_client.db.litellm_guardrailstable.find_many = AsyncMock( + return_value=[other_team_guardrail] + ) + + mock_in_memory_handler = mocker.Mock() + mock_in_memory_handler.list_in_memory_guardrails.return_value = [] + + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mocker.patch( + "litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER", + mock_in_memory_handler, + ) + mock_get_user_team_ids = mocker.patch( + "litellm.proxy.guardrails.guardrail_endpoints._get_user_team_ids", + AsyncMock(return_value=[]), + ) + + viewer_auth = UserAPIKeyAuth( + user_id="viewer-1", user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY + ) + response = await list_guardrails_v2(user_api_key_dict=viewer_auth) + + assert [g.guardrail_id for g in response.guardrails] == ["other-team-guardrail"] + mock_get_user_team_ids.assert_not_called() + + +@pytest.mark.asyncio +async def test_list_guardrails_v2_masks_sensitive_data_for_admin_viewer(mocker): + """ + Read parity for proxy_admin_viewer must not also hand out unmasked secrets. + The guardrail is team-owned so it only reaches the viewer via the admin path. + """ + other_team_guardrail_with_secrets = { + "guardrail_id": "other-team-secret-guardrail", + "guardrail_name": "Other Team Guardrail with Secrets", + "litellm_params": { + "guardrail": "azure/text_moderations", + "mode": "pre_call", + "api_key": "sk-viewer-must-not-see-this", + }, + "guardrail_info": {}, + "team_id": "team-viewer-is-not-in", + "created_at": datetime.now(), + "updated_at": datetime.now(), + } + + mock_prisma_client = mocker.Mock() + mock_prisma_client.db = mocker.Mock() + mock_prisma_client.db.litellm_guardrailstable = mocker.Mock() + mock_prisma_client.db.litellm_guardrailstable.find_many = AsyncMock( + return_value=[other_team_guardrail_with_secrets] + ) + + mock_in_memory_handler = mocker.Mock() + mock_in_memory_handler.list_in_memory_guardrails.return_value = [] + + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mocker.patch( + "litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER", + mock_in_memory_handler, + ) + mocker.patch( + "litellm.proxy.guardrails.guardrail_endpoints._get_user_team_ids", + AsyncMock(return_value=[]), + ) + + viewer_auth = UserAPIKeyAuth( + user_id="viewer-1", user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY + ) + response = await list_guardrails_v2(user_api_key_dict=viewer_auth) + + guardrail = next( + g + for g in response.guardrails + if g.guardrail_id == "other-team-secret-guardrail" + ) + params = guardrail.litellm_params.model_dump() + assert params["api_key"] != "sk-viewer-must-not-see-this" + assert "****" in str(params["api_key"]) + assert params["guardrail"] == "azure/text_moderations" + + @pytest.mark.asyncio async def test_get_guardrail_info_from_db(mocker, mock_prisma_client): """Test getting guardrail info from DB""" @@ -2037,6 +2140,39 @@ async def test_get_guardrail_submission_non_admin_other_team_forbidden(mocker): assert exc_info.value.status_code == 403 +@pytest.mark.asyncio +async def test_get_guardrail_submission_admin_viewer_other_team_allowed(mocker): + """proxy_admin_viewer reads any team's submission without the membership check.""" + mock_prisma = mocker.Mock() + row = mocker.Mock( + guardrail_id="sub-1", + guardrail_name="team-guard", + status="pending_review", + team_id="team-other", + litellm_params={}, + guardrail_info={}, + submitted_at=None, + reviewed_at=None, + created_at=datetime.now(), + updated_at=datetime.now(), + ) + mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=row) + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) + mock_get_user_team_ids = mocker.patch( + "litellm.proxy.guardrails.guardrail_endpoints._get_user_team_ids", + AsyncMock(return_value=[]), + ) + user = UserAPIKeyAuth( + user_id="viewer-1", user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY + ) + + result = await get_guardrail_submission("sub-1", user) + + assert result.guardrail_id == "sub-1" + assert result.team_id == "team-other" + mock_get_user_team_ids.assert_not_called() + + @pytest.mark.asyncio async def test_approve_guardrail_submission_success(mocker): """Approve sets status to active and initializes guardrail in memory.""" diff --git a/tests/test_litellm/proxy/management_endpoints/test_coordination_redis_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_coordination_redis_endpoints.py index 4e6bfc4c063..2e78a4ca0e3 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_coordination_redis_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_coordination_redis_endpoints.py @@ -210,6 +210,27 @@ async def test_get_rejects_non_admin(): assert exc_info.value.status_code == 403 +@pytest.mark.asyncio +async def test_get_allows_proxy_admin_viewer(): + """proxy_admin_viewer has READ parity with proxy_admin; credentials stay redacted.""" + with ( + patch( + "litellm.proxy.proxy_server.prisma_client", + _prisma_with_general_settings({"coordination_redis": _SAVED_SETTINGS}), + ), + patch("litellm.proxy.proxy_server.proxy_config", _proxy_config()), + ): + response = await get_coordination_redis_settings( + user_api_key_dict=UserAPIKeyAuth( + api_key="hashed", user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY + ) + ) + + assert response.source == "coordination_redis" + assert response.values["host"] == "coord-redis.example.com" + assert response.values["password"] == _REDACTED_VALUE + + def test_fields_cover_every_coordination_redis_param(): """The declarative field list drives the Admin UI form; it must stay in sync with the model the backend validates against.""" @@ -437,6 +458,18 @@ async def test_update_rejects_non_admin(): assert exc_info.value.status_code == 403 +@pytest.mark.asyncio +async def test_update_rejects_proxy_admin_viewer(): + """READ parity for proxy_admin_viewer must not leak into the save endpoint.""" + with pytest.raises(HTTPException) as exc_info: + await update_coordination_redis_settings( + request=CoordinationRedisSettingsRequest(settings={"host": "coord-redis.example.com"}), + user_api_key_dict=UserAPIKeyAuth(api_key="hashed", user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY), + litellm_changed_by=None, + ) + assert exc_info.value.status_code == 403 + + # โ”€โ”€ POST /coordination_redis/settings/test โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€ @@ -575,3 +608,14 @@ async def test_connection_test_rejects_non_admin(): user_api_key_dict=UserAPIKeyAuth(api_key="hashed", user_role=LitellmUserRoles.INTERNAL_USER), ) assert exc_info.value.status_code == 403 + + +@pytest.mark.asyncio +async def test_connection_test_rejects_proxy_admin_viewer(): + """Dialing a caller-supplied Redis is a write-shaped action; viewers stay out.""" + with pytest.raises(HTTPException) as exc_info: + await check_coordination_redis_connection( + request=CoordinationRedisSettingsRequest(settings={"host": "coord-redis.example.com"}), + user_api_key_dict=UserAPIKeyAuth(api_key="hashed", user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY), + ) + assert exc_info.value.status_code == 403 diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index a37f7ca764d..aab9a0b4fd0 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -1383,6 +1383,39 @@ async def test_user_info_nonexistent_user(mocker): assert f"User {nonexistent_user_id} not found" in str(exc_info.value.message) +@pytest.mark.asyncio +async def test_user_info_no_user_id_view_only_admin_gets_proxy_admin_payload(mocker): + """PROXY_ADMIN_VIEW_ONLY must take the proxy-admin branch; otherwise /user/info + silently narrows to the viewer's own row instead of the whole tenant.""" + from fastapi import Request + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth, UserInfoResponse + from litellm.proxy.management_endpoints.internal_user_endpoints import user_info + + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.get_data = mocker.AsyncMock(return_value=None) + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + admin_payload = UserInfoResponse(user_id=None, user_info=None, keys=[], teams=[]) + mock_get_user_info_for_proxy_admin = mocker.AsyncMock(return_value=admin_payload) + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints._get_user_info_for_proxy_admin", + mock_get_user_info_for_proxy_admin, + ) + + viewer = UserAPIKeyAuth( + user_id="viewer", user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value + ) + mock_request = mocker.MagicMock(spec=Request) + + response = await user_info( + user_id=None, user_api_key_dict=viewer, request=mock_request + ) + + mock_get_user_info_for_proxy_admin.assert_awaited_once_with(user_api_key_dict=viewer) + assert response is admin_payload + + @pytest.mark.asyncio async def test_new_user_default_teams_flow(mocker): """ @@ -3213,13 +3246,9 @@ def test_enforce_user_info_access_admin_bypass(): _enforce_user_info_access(user_id="someone_else", user_api_key_dict=admin) -def test_enforce_user_info_access_view_only_admin_blocked_from_other_users(): - """PROXY_ADMIN_VIEW_ONLY is not a true admin for /user/info โ€” the upstream - route check applies the same `user_id == valid_token.user_id` rule, so the - re-check here must mirror that and deny cross-user lookups.""" - import pytest - from fastapi import HTTPException - +def test_enforce_user_info_access_view_only_admin_can_read_other_users(): + """PROXY_ADMIN_VIEW_ONLY has read parity with PROXY_ADMIN, so the ownership + re-check must wave it through for another user's id.""" from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.management_endpoints.internal_user_endpoints import ( _enforce_user_info_access, @@ -3229,9 +3258,7 @@ def test_enforce_user_info_access_view_only_admin_blocked_from_other_users(): user_id="viewer", user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, ) - with pytest.raises(HTTPException) as exc_info: - _enforce_user_info_access(user_id="someone_else", user_api_key_dict=viewer) - assert exc_info.value.status_code == 403 + _enforce_user_info_access(user_id="someone_else", user_api_key_dict=viewer) def test_enforce_user_info_access_view_only_admin_can_read_own(): diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index cf9aa477112..e8709f3af34 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -8006,6 +8006,74 @@ async def test_validate_key_list_check_key_hash_not_found(): assert "Key Hash not found" in exc_info.value.message +@pytest.mark.asyncio +async def test_validate_key_list_check_proxy_admin_viewer_skips_db_lookup(): + """proxy_admin_viewer takes the same unscoped read fast-path as proxy_admin, so no + user row is fetched and none of the user/team scoping filters apply.""" + mock_prisma_client = AsyncMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=LiteLLM_UserTable( + user_id="viewer-user", + user_email="viewer@example.com", + teams=[], + organization_memberships=[], + ) + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, + user_id="viewer-user", + ) + + result = await validate_key_list_check( + user_api_key_dict=user_api_key_dict, + user_id="someone-else", + team_id="team-viewer-is-not-in", + organization_id=None, + key_alias=None, + key_hash=None, + prisma_client=mock_prisma_client, + ) + + assert result is None + mock_prisma_client.db.litellm_usertable.find_unique.assert_not_awaited() + assert mock_prisma_client.mock_calls == [] + + +@pytest.mark.asyncio +async def test_validate_key_list_check_internal_user_cannot_query_other_user(): + """Admin-view parity must not leak past the admin roles: an internal user still + cannot list another user's keys.""" + mock_prisma_client = AsyncMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=LiteLLM_UserTable( + user_id="test-user", + user_email="test@example.com", + teams=[], + organization_memberships=[], + ) + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="test-user", + ) + + with pytest.raises(ProxyException) as exc_info: + await validate_key_list_check( + user_api_key_dict=user_api_key_dict, + user_id="other-user", + team_id=None, + organization_id=None, + key_alias=None, + key_hash=None, + prisma_client=mock_prisma_client, + ) + + assert exc_info.value.code == "403" + assert "not authorized to check another user's keys" in exc_info.value.message + + @pytest.mark.asyncio async def test_key_with_budget_id_does_not_store_budget_duration(): """ @@ -15323,3 +15391,54 @@ async def test_rotate_master_key_rotates_sso_identity_assertions( prisma_client=mock_prisma_client, new_master_key="sk-new-master-key", ) + + +@pytest.mark.asyncio +async def test_check_encryption_endpoint_rejects_proxy_admin_viewer(): + """The residual scan walks and decrypt-classifies every credential-bearing table, + so it stays proxy_admin-only despite being read-only.""" + from litellm.proxy.management_endpoints import credential_migration as cm + from litellm.proxy.management_endpoints.key_management_endpoints import ( + check_encryption_endpoint, + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, + user_id="viewer-user", + ) + mock_check = AsyncMock(return_value=cm.MigrationReport()) + + with patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch.object( + cm, "check_encryption", mock_check + ): + with pytest.raises(HTTPException) as exc_info: + await check_encryption_endpoint(user_api_key_dict=user_api_key_dict) + + assert exc_info.value.status_code == 403 + mock_check.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_migrate_encryption_endpoint_rejects_proxy_admin_viewer(): + """The re-encryption write sibling is also proxy_admin-only.""" + from litellm.proxy.management_endpoints import credential_migration as cm + from litellm.proxy.management_endpoints.key_management_endpoints import ( + migrate_encryption_endpoint, + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, + user_id="viewer-user", + ) + mock_migrate = AsyncMock(return_value=cm.MigrationReport()) + + with patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch.object( + cm, "migrate_encryption", mock_migrate + ): + with pytest.raises(HTTPException) as exc_info: + await migrate_encryption_endpoint( + user_api_key_dict=user_api_key_dict, dry_run=False + ) + + assert exc_info.value.status_code == 403 + mock_migrate.assert_not_awaited() diff --git a/tests/test_litellm/proxy/management_endpoints/test_workflow_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_workflow_management_endpoints.py index a337ff6d888..27adb3e0892 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_workflow_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_workflow_management_endpoints.py @@ -3,19 +3,25 @@ Unit tests for workflow management endpoints (/v1/workflows/runs/*). Uses FastAPI TestClient with a mocked prisma_client. """ +import asyncio import os import sys from datetime import datetime, timezone from typing import Any from unittest.mock import AsyncMock, MagicMock, patch -from fastapi import FastAPI +import pytest +from fastapi import FastAPI, HTTPException from fastapi.testclient import TestClient from prisma.errors import UniqueViolationError sys.path.insert(0, os.path.abspath("../../..")) -from litellm.proxy.management_endpoints.workflow_management_endpoints import router +from litellm.proxy.management_endpoints.workflow_management_endpoints import ( + _read_scope_caller, + _require_run, + router, +) # --------------------------------------------------------------------------- @@ -140,6 +146,31 @@ def _override_auth_user_with_token(token: str = "tok-abc") -> Any: return auth +def _override_auth_admin_viewer(token: str = "tok-viewer") -> Any: + """Viewer carries a real token, so a re-scoped read path would be observable.""" + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + + auth = UserAPIKeyAuth( + api_key="sk-viewer", + user_id="viewer-1", + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, + ) + auth.token = token + return auth + + +def _override_auth_internal_user(token: str = "tok-internal") -> Any: + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + + auth = UserAPIKeyAuth( + api_key="sk-internal", + user_id="user-2", + user_role=LitellmUserRoles.INTERNAL_USER, + ) + auth.token = token + return auth + + # --------------------------------------------------------------------------- # Tests # --------------------------------------------------------------------------- @@ -609,3 +640,100 @@ class TestTenantIsolation: resp = client.get("/v1/workflows/runs/run-1") assert resp.status_code == 200 + + +class TestAdminViewerReadParity: + """proxy_admin_viewer reads every run; write paths stay on the strict admin gate.""" + + def _make_app_with_auth(self, auth_fn): + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + self._prisma = _make_prisma_client() + app = _make_app() + app.dependency_overrides[user_api_key_auth] = auth_fn + return TestClient(app, raise_server_exceptions=True) + + def test_read_scope_caller_drops_scope_for_admin_viewer_only(self): + """None means 'no ownership filter'; every other non-admin role keeps its caller.""" + internal = _override_auth_internal_user() + assert _read_scope_caller(_override_auth_admin_viewer()) is None + assert _read_scope_caller(internal) is internal + + @patch("litellm.proxy.proxy_server.prisma_client") + def test_admin_viewer_list_not_scoped(self, mock_pc): + client = self._make_app_with_auth(_override_auth_admin_viewer) + mock_pc.db = self._prisma.db + self._prisma.db.litellm_workflowrun.find_many = AsyncMock(return_value=[]) + + resp = client.get("/v1/workflows/runs") + assert resp.status_code == 200 + call_kwargs = self._prisma.db.litellm_workflowrun.find_many.call_args[1] + assert "created_by" not in call_kwargs["where"] + + @patch("litellm.proxy.proxy_server.prisma_client") + def test_admin_viewer_get_other_owners_run_succeeds(self, mock_pc): + client = self._make_app_with_auth(_override_auth_admin_viewer) + mock_pc.db = self._prisma.db + self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( + return_value=_make_run(created_by="tok-other-owner") + ) + + resp = client.get("/v1/workflows/runs/run-1") + assert resp.status_code == 200 + + @patch("litellm.proxy.proxy_server.prisma_client") + def test_admin_viewer_lists_other_owners_events(self, mock_pc): + client = self._make_app_with_auth(_override_auth_admin_viewer) + mock_pc.db = self._prisma.db + self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( + return_value=_make_run(created_by="tok-other-owner") + ) + self._prisma.db.litellm_workflowevent.find_many = AsyncMock( + return_value=[_make_event(sequence_number=0)] + ) + + resp = client.get("/v1/workflows/runs/run-1/events") + assert resp.status_code == 200 + assert resp.json()["count"] == 1 + + @patch("litellm.proxy.proxy_server.prisma_client") + def test_admin_viewer_lists_other_owners_messages(self, mock_pc): + client = self._make_app_with_auth(_override_auth_admin_viewer) + mock_pc.db = self._prisma.db + self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( + return_value=_make_run(created_by="tok-other-owner") + ) + self._prisma.db.litellm_workflowmessage.find_many = AsyncMock( + return_value=[_make_message(sequence_number=0)] + ) + + resp = client.get("/v1/workflows/runs/run-1/messages") + assert resp.status_code == 200 + assert resp.json()["count"] == 1 + + @patch("litellm.proxy.proxy_server.prisma_client") + def test_admin_viewer_cannot_update_other_owners_run(self, mock_pc): + """Read parity must not become write parity: PATCH still passes the caller through.""" + client = self._make_app_with_auth(_override_auth_admin_viewer) + mock_pc.db = self._prisma.db + self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( + return_value=_make_run(created_by="tok-other-owner") + ) + self._prisma.db.litellm_workflowrun.update = AsyncMock( + return_value=_make_run(status="completed") + ) + + resp = client.patch("/v1/workflows/runs/run-1", json={"status": "completed"}) + assert resp.status_code == 404 + self._prisma.db.litellm_workflowrun.update.assert_not_awaited() + + def test_require_run_still_scopes_when_handed_a_viewer(self): + """Only read callers pass None; the helper itself never loosened.""" + prisma = _make_prisma_client() + prisma.db.litellm_workflowrun.find_unique = AsyncMock( + return_value=_make_run(created_by="tok-other-owner") + ) + + with pytest.raises(HTTPException) as exc_info: + asyncio.run(_require_run(prisma, "run-1", _override_auth_admin_viewer())) + assert exc_info.value.status_code == 404 diff --git a/tests/test_litellm/proxy/memory/test_memory_endpoints.py b/tests/test_litellm/proxy/memory/test_memory_endpoints.py index ca011c77af8..ec81ef2ff7a 100644 --- a/tests/test_litellm/proxy/memory/test_memory_endpoints.py +++ b/tests/test_litellm/proxy/memory/test_memory_endpoints.py @@ -19,7 +19,7 @@ from fastapi.testclient import TestClient sys.path.insert(0, os.path.abspath("../../..")) from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth -from litellm.proxy.memory.memory_endpoints import router +from litellm.proxy.memory.memory_endpoints import _visibility_filter, router def _make_row( @@ -218,6 +218,14 @@ def _admin_auth() -> UserAPIKeyAuth: ) +def _admin_viewer_auth() -> UserAPIKeyAuth: + return UserAPIKeyAuth( + api_key="sk-viewer", + user_id="viewer", + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, + ) + + def _patch_prisma(prisma: Any): """Patch the endpoint module's _require_prisma to return our fake.""" return patch( @@ -913,3 +921,67 @@ class TestMemoryEndpoints: with _patch_prisma(self.prisma): resp = client.delete("/v1/memory/notes") assert resp.status_code == 404 + + def test_visibility_filter_unscoped_for_admin_viewer(self): + """ + proxy_admin_viewer reads with the same unscoped filter as proxy_admin; + every other role stays row-restricted. + """ + assert _visibility_filter(_admin_viewer_auth()) is None + assert _visibility_filter(_user_auth("user-a", "team-a")) is not None + + def test_list_memory_admin_viewer_sees_all(self): + """Read parity end-to-end: the viewer's own user_id/team_id must not filter the list.""" + table = self.prisma.db.litellm_memorytable + table.rows.extend( + [ + _make_row(memory_id="m1", key="a", user_id="user-a", team_id=None), + _make_row(memory_id="m2", key="b", user_id="user-b", team_id="team-b"), + ] + ) + client = _make_client(_admin_viewer_auth()) + with _patch_prisma(self.prisma): + resp = client.get("/v1/memory") + assert resp.status_code == 200, resp.text + body = resp.json() + assert {m["key"] for m in body["memories"]} == {"a", "b"} + assert body["total"] == 2 + + def test_put_memory_admin_viewer_cannot_overwrite_foreign_row(self): + """ + Read parity must not become write parity: the viewer now SEES this row + (403, not 404) but `_assert_write_access` still refuses the write. + """ + table = self.prisma.db.litellm_memorytable + table.rows.append( + _make_row( + memory_id="m1", + key="user_role", + value="A's notes", + user_id="user-a", + team_id="team-a", + ) + ) + client = _make_client(_admin_viewer_auth()) + with _patch_prisma(self.prisma): + resp = client.put("/v1/memory/user_role", json={"value": "viewer overwrite"}) + assert resp.status_code == 403, resp.text + assert table.rows[0].value == "A's notes" + + def test_delete_memory_admin_viewer_cannot_delete_foreign_row(self): + """Same write gate as the PUT case, for DELETE.""" + table = self.prisma.db.litellm_memorytable + table.rows.append( + _make_row( + memory_id="m1", + key="user_role", + value="A's notes", + user_id="user-a", + team_id="team-a", + ) + ) + client = _make_client(_admin_viewer_auth()) + with _patch_prisma(self.prisma): + resp = client.delete("/v1/memory/user_role") + assert resp.status_code == 403, resp.text + assert len(table.rows) == 1 diff --git a/tests/test_litellm/proxy/prompts/test_prompt_endpoints.py b/tests/test_litellm/proxy/prompts/test_prompt_endpoints.py index 39b6bce46fa..57ad6acae3b 100644 --- a/tests/test_litellm/proxy/prompts/test_prompt_endpoints.py +++ b/tests/test_litellm/proxy/prompts/test_prompt_endpoints.py @@ -319,3 +319,144 @@ class TestPromptVersionsEndpoint: assert exc_info.value.status_code == 404 assert "No versions found" in exc_info.value.detail + + +class TestAdminViewerReadAccess: + """ + proxy_admin_viewer has READ parity with proxy_admin on the prompt read endpoints + """ + + @pytest.mark.asyncio + async def test_list_prompts_returns_all_prompts_for_admin_viewer(self): + """A role without admin view falls through to the empty-list branch here.""" + from unittest.mock import patch + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.prompts.prompt_endpoints import list_prompts + + viewer = UserAPIKeyAuth( + api_key="test_key", user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY + ) + + mock_prompts = { + "jack.v1": PromptSpec( + prompt_id="jack.v1", + litellm_params=PromptLiteLLMParams( + prompt_id="jack", + prompt_integration="dotprompt", + dotprompt_content="v1", + ), + prompt_info=PromptInfo(prompt_type="db"), + ), + "jack.v2": PromptSpec( + prompt_id="jack.v2", + litellm_params=PromptLiteLLMParams( + prompt_id="jack", + prompt_integration="dotprompt", + dotprompt_content="v2", + ), + prompt_info=PromptInfo(prompt_type="db"), + ), + "jane.v1": PromptSpec( + prompt_id="jane.v1", + litellm_params=PromptLiteLLMParams( + prompt_id="jane", + prompt_integration="dotprompt", + dotprompt_content="jane", + ), + prompt_info=PromptInfo(prompt_type="db"), + ), + } + + with patch( + "litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY" + ) as mock_registry: + mock_registry.IN_MEMORY_PROMPTS = mock_prompts + + response = await list_prompts(user_api_key_dict=viewer) + + assert sorted(p.prompt_id for p in response.prompts) == ["jack", "jane"] + jack = next(p for p in response.prompts if p.prompt_id == "jack") + assert jack.litellm_params.dotprompt_content == "v2" + + @pytest.mark.asyncio + async def test_get_prompt_versions_allows_admin_viewer(self): + """Version history used to 403 anyone who was not exactly proxy_admin.""" + from unittest.mock import patch + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.prompts.prompt_endpoints import get_prompt_versions + + viewer = UserAPIKeyAuth( + api_key="test_key", user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY + ) + + mock_prompts = { + "jack.v1": PromptSpec( + prompt_id="jack.v1", + litellm_params=PromptLiteLLMParams( + prompt_id="jack", + prompt_integration="dotprompt", + dotprompt_content="v1", + ), + prompt_info=PromptInfo(prompt_type="db"), + ), + "jack.v2": PromptSpec( + prompt_id="jack.v2", + litellm_params=PromptLiteLLMParams( + prompt_id="jack", + prompt_integration="dotprompt", + dotprompt_content="v2", + ), + prompt_info=PromptInfo(prompt_type="db"), + ), + } + + with ( + patch("litellm.proxy.proxy_server.prisma_client", None), + patch( + "litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY" + ) as mock_registry, + ): + mock_registry.IN_MEMORY_PROMPTS = mock_prompts + + response = await get_prompt_versions( + prompt_id="jack", user_api_key_dict=viewer + ) + + assert [p.version for p in response.prompts] == [2, 1] + + @pytest.mark.asyncio + async def test_get_prompt_info_allows_admin_viewer(self): + """Prompt info used to 403 anyone who was not exactly proxy_admin.""" + from unittest.mock import patch + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.prompts.prompt_endpoints import get_prompt_info + + viewer = UserAPIKeyAuth( + api_key="test_key", user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY + ) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", None), + patch( + "litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY" + ) as mock_registry, + ): + mock_registry.get_prompt_by_id.return_value = PromptSpec( + prompt_id="jack.v2", + litellm_params=PromptLiteLLMParams( + prompt_id="jack", + prompt_integration="dotprompt", + dotprompt_content="v2", + ), + prompt_info=PromptInfo(prompt_type="db"), + ) + mock_registry.IN_MEMORY_PROMPTS = {"jack.v1": {}, "jack.v2": {}} + mock_registry.get_prompt_callback_by_id.return_value = None + + response = await get_prompt_info(prompt_id="jack", user_api_key_dict=viewer) + + assert response.prompt_spec.prompt_id == "jack" + assert response.prompt_spec.version == 2 diff --git a/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py b/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py index 577af3dcffc..e73f1d08cb5 100644 --- a/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py +++ b/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py @@ -538,6 +538,48 @@ async def test_populate_team_access_sets_direct_access_false_by_default(monkeypa assert by_id["global-id-1"]["model_info"]["direct_access"] is True +@pytest.mark.asyncio +async def test_populate_team_access_gives_view_only_admin_full_admin_scope(monkeypatch): + """proxy_admin_viewer reads with admin scope - every team ("*") plus direct access + to all non-team models - instead of being narrowed to its own user row.""" + team_row = _team_row() + global_row = { + "model_name": "gpt-4o", + "litellm_params": {"model": "gpt-4o"}, + "model_info": {"id": "global-id-1", "db_model": False}, + } + router = MagicMock() + router.get_model_ids.return_value = ["global-id-1"] + + get_all_team_models = AsyncMock(return_value={"byok-id-1": ["team-abc-123"]}) + monkeypatch.setattr(ps, "get_all_team_models", get_all_team_models) + + prisma_client = MagicMock() + prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=LiteLLM_UserTable(user_id="viewer", teams=[], models=[]) + ) + + viewer = UserAPIKeyAuth( + user_id="viewer", + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, + team_models=[], + ) + result = await ps._populate_team_access_on_models( + user_api_key_dict=viewer, + prisma_client=prisma_client, + llm_router=router, + all_models=[team_row, global_row], + ) + + assert get_all_team_models.await_args.kwargs["user_teams"] == "*" + router.get_model_ids.assert_called_once_with(exclude_team_models=True) + prisma_client.db.litellm_usertable.find_unique.assert_not_awaited() + + by_id = {m["model_info"]["id"]: m for m in result} + assert by_id["byok-id-1"]["model_info"]["access_via_team_ids"] == ["team-abc-123"] + assert by_id["global-id-1"]["model_info"]["direct_access"] is True + + @pytest.mark.asyncio async def test_model_info_v1_team_id_without_db_fails_fast(monkeypatch): """`teamId` without a connected DB raises 500 before any enrichment work runs.""" diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useAuthorized.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useAuthorized.test.ts index 567ca911458..bb14f6c3d21 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useAuthorized.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useAuthorized.test.ts @@ -152,6 +152,8 @@ describe("useAuthorized", () => { expect(result.current.userId).toBe("user-1"); expect(result.current.userEmail).toBe("user@example.com"); expect(result.current.userRole).toBe("Admin"); + expect(result.current.userRoleLabel).toBe("Admin"); + expect(result.current.isViewOnly).toBe(false); expect(result.current.premiumUser).toBe(true); expect(result.current.disabledPersonalKeyCreation).toBe(false); expect(result.current.showSSOBanner).toBe(true); @@ -159,6 +161,44 @@ describe("useAuthorized", () => { expect(clearTokenCookiesMock).not.toHaveBeenCalled(); }); + it("should present proxy_admin_viewer as Admin while flagging it view-only", async () => { + getUiConfigMock.mockResolvedValue({ + server_root_path: "/", + proxy_base_url: null, + auto_redirect_to_sso: false, + admin_ui_disabled: false, + sso_configured: false, + }); + + const decodedPayload = { + key: "api-key-456", + user_id: "user-2", + user_email: "viewer@example.com", + user_role: "proxy_admin_viewer", + premium_user: true, + disabled_non_admin_personal_key_creation: false, + login_method: "username_password", + }; + + decodeTokenMock.mockReturnValue(decodedPayload); + checkTokenValidityMock.mockReturnValue(true); + + const token = createJwt(decodedPayload); + document.cookie = `token=${token}; path=/;`; + + const { result } = renderHook(() => useAuthorized(), { wrapper }); + + await waitFor(() => { + expect(result.current.token).toBe(token); + }); + + expect(result.current.userRole).toBe("Admin"); + expect(result.current.userRoleLabel).toBe("Admin Viewer"); + expect(result.current.isViewOnly).toBe(true); + expect(replaceMock).not.toHaveBeenCalled(); + expect(clearTokenCookiesMock).not.toHaveBeenCalled(); + }); + it("should clear cookies and redirect on an invalid token", async () => { getUiConfigMock.mockResolvedValue({ server_root_path: "/", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useAuthorized.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useAuthorized.ts index bb22ebf5edc..40d1ec09d1f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useAuthorized.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useAuthorized.ts @@ -5,7 +5,7 @@ import { clearTokenCookies, getCookie } from "@/utils/cookieUtils"; import { checkTokenValidity, decodeToken } from "@/utils/jwtUtils"; import { buildLoginUrlWithReturn, getLoginUrl, storeReturnUrl } from "@/utils/returnUrlUtils"; import { useCallback, useEffect, useMemo } from "react"; -import { formatUserRole } from "@/utils/roles"; +import { effectiveSessionRole, formatUserRole, isViewOnlySessionRole } from "@/utils/roles"; import { useUIConfig } from "./uiConfig/useUIConfig"; const useAuthorized = () => { @@ -45,7 +45,9 @@ const useAuthorized = () => { accessToken: decoded?.key ?? null, userId: decoded?.user_id ?? null, userEmail: decoded?.user_email ?? null, - userRole: formatUserRole(decoded?.user_role), + userRole: effectiveSessionRole(decoded?.user_role), + userRoleLabel: formatUserRole(decoded?.user_role), + isViewOnly: isViewOnlySessionRole(decoded?.user_role), premiumUser: decoded?.premium_user ?? null, disabledPersonalKeyCreation: decoded?.disabled_non_admin_personal_key_creation ?? null, showSSOBanner: decoded?.login_method === "username_password", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/playground/page.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/playground/page.test.tsx index 54e99d9db29..85e19d7d251 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/playground/page.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/playground/page.test.tsx @@ -10,6 +10,7 @@ vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ accessToken: "sk-test", userId: "user-1", userRole: authState.userRole, + isViewOnly: ["Admin Viewer", "Internal Viewer"].includes(authState.userRole), disabledPersonalKeyCreation: false, }), })); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/playground/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/playground/page.tsx index 8986084b1a7..a4ea85311c4 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/playground/page.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/playground/page.tsx @@ -9,7 +9,6 @@ import { TabGroup, TabList, Tab, TabPanels, TabPanel } from "@tremor/react"; import { DeprecationBanner } from "@/components/DeprecationBanner"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import { fetchProxySettings } from "@/utils/proxyUtils"; -import { isViewOnlyRole } from "@/utils/roles"; interface ProxySettings { PROXY_BASE_URL?: string; @@ -17,7 +16,7 @@ interface ProxySettings { } export default function PlaygroundPage() { - const { accessToken, userRole, userId, disabledPersonalKeyCreation, token } = useAuthorized(); + const { accessToken, userRole, userId, disabledPersonalKeyCreation, token, isViewOnly } = useAuthorized(); const [proxySettings, setProxySettings] = useState(undefined); useEffect(() => { @@ -36,7 +35,7 @@ export default function PlaygroundPage() { initializeProxySettings(); }, [accessToken]); - if (isViewOnlyRole(userRole)) { + if (isViewOnly) { return (

Access Denied

diff --git a/ui/litellm-dashboard/src/components/Navbar/UserDropdown/UserDropdown.test.tsx b/ui/litellm-dashboard/src/components/Navbar/UserDropdown/UserDropdown.test.tsx index 31ddae31798..cad5ced340e 100644 --- a/ui/litellm-dashboard/src/components/Navbar/UserDropdown/UserDropdown.test.tsx +++ b/ui/litellm-dashboard/src/components/Navbar/UserDropdown/UserDropdown.test.tsx @@ -6,7 +6,7 @@ import UserDropdown from "./UserDropdown"; let mockUseAuthorizedImpl = () => ({ userId: "test-user-id", userEmail: "test@example.com", - userRole: "Admin", + userRoleLabel: "Admin", premiumUser: false, }); @@ -44,7 +44,7 @@ describe("UserDropdown", () => { mockUseAuthorizedImpl = () => ({ userId: "test-user-id", userEmail: "test@example.com", - userRole: "Admin", + userRoleLabel: "Admin", premiumUser: false, }); mockUseDisableShowPromptsImpl = () => false; @@ -115,7 +115,7 @@ describe("UserDropdown", () => { mockUseAuthorizedImpl = () => ({ userId: "test-user-id", userEmail: "test@example.com", - userRole: "Admin", + userRoleLabel: "Admin", premiumUser: true, }); @@ -238,7 +238,7 @@ describe("UserDropdown", () => { mockUseAuthorizedImpl = () => ({ userId: "default_user_id", userEmail: null as any, - userRole: "Admin", + userRoleLabel: "Admin", premiumUser: false, }); renderWithProviders(); @@ -250,7 +250,7 @@ describe("UserDropdown", () => { mockUseAuthorizedImpl = () => ({ userId: "test-user-id", userEmail: null as any, - userRole: "Admin", + userRoleLabel: "Admin", premiumUser: false, }); @@ -268,7 +268,7 @@ describe("UserDropdown", () => { mockUseAuthorizedImpl = () => ({ userId: null as any, userEmail: "test@example.com", - userRole: "Admin", + userRoleLabel: "Admin", premiumUser: false, }); diff --git a/ui/litellm-dashboard/src/components/Navbar/UserDropdown/UserDropdown.tsx b/ui/litellm-dashboard/src/components/Navbar/UserDropdown/UserDropdown.tsx index a71fc1b97a8..28e981c57a1 100644 --- a/ui/litellm-dashboard/src/components/Navbar/UserDropdown/UserDropdown.tsx +++ b/ui/litellm-dashboard/src/components/Navbar/UserDropdown/UserDropdown.tsx @@ -69,7 +69,7 @@ interface UserDropdownProps { } const UserDropdown: React.FC = ({ onLogout, variant = "navbar", collapsed = false }) => { - const { userId, userEmail, userRole, premiumUser } = useAuthorized(); + const { userId, userEmail, userRoleLabel: userRole, premiumUser } = useAuthorized(); const disableShowPrompts = useDisableShowPrompts(); const disableBlogPosts = useDisableBlogPosts(); const disableBouncingIcon = useDisableBouncingIcon(); diff --git a/ui/litellm-dashboard/src/components/SidebarAccountMenu/SidebarAccountMenu.test.tsx b/ui/litellm-dashboard/src/components/SidebarAccountMenu/SidebarAccountMenu.test.tsx index 1e4eb5b5af4..9d56a889ed4 100644 --- a/ui/litellm-dashboard/src/components/SidebarAccountMenu/SidebarAccountMenu.test.tsx +++ b/ui/litellm-dashboard/src/components/SidebarAccountMenu/SidebarAccountMenu.test.tsx @@ -6,7 +6,7 @@ import SidebarAccountMenu from "./SidebarAccountMenu"; interface AuthMock { userId: string | null; userEmail: string | null; - userRole: string; + userRoleLabel: string; premiumUser: boolean; accessToken: string; } @@ -14,7 +14,7 @@ interface AuthMock { let mockUseAuthorizedImpl: () => AuthMock = () => ({ userId: "test-user-id", userEmail: "test@example.com", - userRole: "Admin", + userRoleLabel: "Admin", premiumUser: false, accessToken: "test-token", }); @@ -74,7 +74,7 @@ describe("SidebarAccountMenu", () => { mockUseAuthorizedImpl = () => ({ userId: "test-user-id", userEmail: "test@example.com", - userRole: "Admin", + userRoleLabel: "Admin", premiumUser: false, accessToken: "test-token", }); @@ -127,7 +127,7 @@ describe("SidebarAccountMenu", () => { mockUseAuthorizedImpl = () => ({ userId: "test-user-id", userEmail: "test@example.com", - userRole: "Admin", + userRoleLabel: "Admin", premiumUser: true, accessToken: "test-token", }); @@ -273,7 +273,7 @@ describe("SidebarAccountMenu", () => { mockUseAuthorizedImpl = () => ({ userId: "default_user_id", userEmail: null, - userRole: "Admin", + userRoleLabel: "Admin", premiumUser: false, accessToken: "test-token", }); @@ -286,7 +286,7 @@ describe("SidebarAccountMenu", () => { mockUseAuthorizedImpl = () => ({ userId: "test-user-id", userEmail: null, - userRole: "Admin", + userRoleLabel: "Admin", premiumUser: false, accessToken: "test-token", }); diff --git a/ui/litellm-dashboard/src/components/SidebarAccountMenu/SidebarAccountMenu.tsx b/ui/litellm-dashboard/src/components/SidebarAccountMenu/SidebarAccountMenu.tsx index b7a16bcf09a..d1bed9370b4 100644 --- a/ui/litellm-dashboard/src/components/SidebarAccountMenu/SidebarAccountMenu.tsx +++ b/ui/litellm-dashboard/src/components/SidebarAccountMenu/SidebarAccountMenu.tsx @@ -81,7 +81,7 @@ interface SidebarAccountMenuProps { } const SidebarAccountMenu: React.FC = ({ onLogout, collapsed = false }) => { - const { userId, userEmail, userRole, premiumUser, accessToken } = useAuthorized(); + const { userId, userEmail, userRoleLabel: userRole, premiumUser, accessToken } = useAuthorized(); const { data: healthData } = useHealthReadinessDetails(accessToken); const version = healthData?.litellm_version; const disableShowPrompts = useDisableShowPrompts(); diff --git a/ui/litellm-dashboard/src/components/leftnav.test.tsx b/ui/litellm-dashboard/src/components/leftnav.test.tsx index e07d0bb26eb..a5b273a0f56 100644 --- a/ui/litellm-dashboard/src/components/leftnav.test.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.test.tsx @@ -19,6 +19,7 @@ const { mockUseAuthorized, mockUseOrganizations } = vi.hoisted(() => { userId: "test-user-id", accessToken: "test-access-token", userRole: "admin", + isViewOnly: false, token: "test-token", userEmail: "test@example.com", premiumUser: false, @@ -156,12 +157,15 @@ describe("Sidebar (leftnav)", () => { describe("Admin Viewer parity", () => { // Admin Viewer follows a "read parity with Proxy Admin, no writes, no - // cost-incurring actions" rule. Playground stays hidden (incurs LLM - // cost); Models + Endpoints and Agents must be visible read-only. + // cost-incurring actions" rule. The session hook presents the viewer as + // an admin (`userRole: "admin"`) with `isViewOnly: true`; Playground + // stays hidden (incurs LLM cost) via the isViewOnly flag, while every + // admin page (Models + Endpoints, Agents, Logs, ...) is visible read-only. const adminViewerAuth = { userId: "admin-viewer-user-id", accessToken: "test-access-token", - userRole: "admin_viewer", + userRole: "admin", + isViewOnly: true, token: "test-token", userEmail: "viewer@example.com", premiumUser: false, diff --git a/ui/litellm-dashboard/src/components/leftnav.tsx b/ui/litellm-dashboard/src/components/leftnav.tsx index cd92fc5bedb..f08092d0e38 100644 --- a/ui/litellm-dashboard/src/components/leftnav.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.tsx @@ -407,7 +407,7 @@ const Sidebar_: React.FC = ({ disableVectorStoresForInternalUsers, allowVectorStoresForTeamAdmins, }) => { - const { userId, accessToken, userRole } = useAuthorized(); + const { userId, accessToken, userRole, isViewOnly } = useAuthorized(); const { data: organizations } = useOrganizations(); const { data: teams } = useTeams(); const { logoUrl } = useTheme(); @@ -449,6 +449,7 @@ const Sidebar_: React.FC = ({ return items .map((item) => ({ ...item, children: item.children ? filterItemsByRole(item.children) : undefined })) .filter((item) => { + if (item.key === "llm-playground" && isViewOnly) return false; if (item.key === "organizations" || item.key === "users") { const hasRoleAccess = !item.roles || item.roles.includes(userRole) || isOrgAdmin; if (!hasRoleAccess) return false; diff --git a/ui/litellm-dashboard/src/components/user_dashboard.tsx b/ui/litellm-dashboard/src/components/user_dashboard.tsx index 1b8afecb619..1ed4e1d0bba 100644 --- a/ui/litellm-dashboard/src/components/user_dashboard.tsx +++ b/ui/litellm-dashboard/src/components/user_dashboard.tsx @@ -5,6 +5,7 @@ import { jwtDecode } from "jwt-decode"; import React, { useEffect, useState } from "react"; import { fetchTeams } from "./common_components/fetch_teams"; import { KeyResponse, Team } from "./key_team_helpers/key_list"; +import { effectiveSessionRole } from "@/utils/roles"; import { getProxyBaseUrl, getProxyUISettings, @@ -97,30 +98,6 @@ const UserDashboard: React.FC = ({ return () => window.removeEventListener("beforeunload", handleBeforeUnload); }, []); - function formatUserRole(userRole: string) { - if (!userRole) { - return "Undefined Role"; - } - switch (userRole.toLowerCase()) { - case "app_owner": - return "App Owner"; - case "demo_app_owner": - return "App Owner"; - case "proxy_admin": - return "Admin"; - case "proxy_admin_viewer": - return "Admin Viewer"; - case "app_user": - return "App User"; - case "internal_user": - return "Internal User"; - case "internal_user_viewer": - return "Internal Viewer"; - default: - return "Unknown Role"; - } - } - // console.log(`selectedTeam: ${Object.entries(selectedTeam)}`); // Moved useEffect inside the component and used a condition to run fetch only if the params are available useEffect(() => { @@ -134,8 +111,7 @@ const UserDashboard: React.FC = ({ // check if userRole is defined if (decoded.user_role) { - const formattedUserRole = formatUserRole(decoded.user_role); - setUserRole(formattedUserRole); + setUserRole(effectiveSessionRole(decoded.user_role)); } else { } diff --git a/ui/litellm-dashboard/src/contexts/AuthContext.tsx b/ui/litellm-dashboard/src/contexts/AuthContext.tsx index 3693d858952..123feb18a6c 100644 --- a/ui/litellm-dashboard/src/contexts/AuthContext.tsx +++ b/ui/litellm-dashboard/src/contexts/AuthContext.tsx @@ -4,7 +4,7 @@ import React, { createContext, useContext, useEffect, useState } from "react"; import { jwtDecode } from "jwt-decode"; import { clearTokenCookies, getCookie } from "@/utils/cookieUtils"; import { isJwtExpired } from "@/utils/jwtUtils"; -import { formatUserRole } from "@/utils/roles"; +import { effectiveSessionRole } from "@/utils/roles"; import { getUiConfig, setGlobalLitellmHeaderName } from "@/components/networking"; function deleteCookie(name: string, path = "/") { @@ -107,7 +107,7 @@ export function AuthProvider({ children }: { children: React.ReactNode }) { setDisabledPersonalKeyCreation(decoded.disabled_non_admin_personal_key_creation); if (decoded.user_role) { - setUserRole(formatUserRole(decoded.user_role)); + setUserRole(effectiveSessionRole(decoded.user_role)); } if (decoded.user_email) { setUserEmail(decoded.user_email); diff --git a/ui/litellm-dashboard/src/utils/roles.test.ts b/ui/litellm-dashboard/src/utils/roles.test.ts index 9a8a5a9c0c4..83f633bc299 100644 --- a/ui/litellm-dashboard/src/utils/roles.test.ts +++ b/ui/litellm-dashboard/src/utils/roles.test.ts @@ -1,9 +1,11 @@ import { describe, it, expect } from "vitest"; import { + effectiveSessionRole, isAdminRole, isProxyAdminRole, isUserTeamAdminForAnyTeam, isUserTeamAdminForSingleTeam, + isViewOnlySessionRole, rolesAllowedToViewWriteScopedPages, rolesWithWriteAccess, } from "./roles"; @@ -172,4 +174,66 @@ describe("roles", () => { expect(rolesAllowedToViewWriteScopedPages.length).toBeGreaterThan(rolesWithWriteAccess.length); }); }); + + describe("effectiveSessionRole", () => { + it("normalizes proxy_admin_viewer to Admin", () => { + expect(effectiveSessionRole("proxy_admin_viewer")).toBe("Admin"); + }); + + it("keeps proxy_admin as Admin", () => { + expect(effectiveSessionRole("proxy_admin")).toBe("Admin"); + }); + + it("gives proxy_admin_viewer the same session role as proxy_admin", () => { + expect(effectiveSessionRole("proxy_admin_viewer")).toBe(effectiveSessionRole("proxy_admin")); + }); + + it("lets a normalized proxy_admin_viewer pass admin-tier role gates", () => { + expect(rolesWithWriteAccess).toContain(effectiveSessionRole("proxy_admin_viewer")); + }); + + it("does not collapse internal_user_viewer into an admin role", () => { + expect(effectiveSessionRole("internal_user_viewer")).toBe("Internal Viewer"); + expect(rolesWithWriteAccess).not.toContain(effectiveSessionRole("internal_user_viewer")); + }); + + it("leaves other roles untouched", () => { + expect(effectiveSessionRole("internal_user")).toBe("Internal User"); + expect(effectiveSessionRole("org_admin")).toBe("Org Admin"); + }); + + it("returns Undefined Role for a missing role", () => { + expect(effectiveSessionRole(undefined)).toBe("Undefined Role"); + expect(effectiveSessionRole("")).toBe("Undefined Role"); + }); + }); + + describe("isViewOnlySessionRole", () => { + it("returns true for proxy_admin_viewer", () => { + expect(isViewOnlySessionRole("proxy_admin_viewer")).toBe(true); + }); + + it("returns false for proxy_admin", () => { + expect(isViewOnlySessionRole("proxy_admin")).toBe(false); + }); + + it("returns true for internal_user_viewer", () => { + expect(isViewOnlySessionRole("internal_user_viewer")).toBe(true); + }); + + it("returns false for internal_user and org_admin", () => { + expect(isViewOnlySessionRole("internal_user")).toBe(false); + expect(isViewOnlySessionRole("org_admin")).toBe(false); + }); + + it("returns false for a missing role", () => { + expect(isViewOnlySessionRole(undefined)).toBe(false); + expect(isViewOnlySessionRole("")).toBe(false); + }); + + it("stays true for proxy_admin_viewer even though its session role reads as Admin", () => { + expect(effectiveSessionRole("proxy_admin_viewer")).toBe("Admin"); + expect(isViewOnlySessionRole("proxy_admin_viewer")).toBe(true); + }); + }); }); diff --git a/ui/litellm-dashboard/src/utils/roles.ts b/ui/litellm-dashboard/src/utils/roles.ts index 2137d6cfaf2..8d226313f78 100644 --- a/ui/litellm-dashboard/src/utils/roles.ts +++ b/ui/litellm-dashboard/src/utils/roles.ts @@ -65,3 +65,15 @@ export const formatUserRole = (userRole: string): string => { return "Unknown Role"; } }; + +const viewOnlyRawRoles = ["proxy_admin_viewer", "internal_user_viewer", "internal_viewer"]; + +export const effectiveSessionRole = (rawUserRole?: string): string => { + if (rawUserRole?.toLowerCase() === "proxy_admin_viewer") { + return "Admin"; + } + return formatUserRole(rawUserRole ?? ""); +}; + +export const isViewOnlySessionRole = (rawUserRole?: string): boolean => + viewOnlyRawRoles.includes(rawUserRole?.toLowerCase() ?? ""); From d3d30353aa4957518b1b55141a4c0a402a1604c1 Mon Sep 17 00:00:00 2001 From: tin-berri Date: Wed, 5 Aug 2026 11:51:24 -0700 Subject: [PATCH 09/11] refactor(ui): remove the three dashboard lint-budget violations added by #35893 (#35960) PR #35929 zeroed the eslint budget headroom while #35893 added UI code in parallel, so staging went over budget by one complexity violation and two no-large-inline-object-arg violations, failing frontend-lint on every UI-touching PR until #35964 reverted the ratchet. This removes the three violations at the source so the budgets can ratchet back down: the submit-blocked-reason chain in add_auto_router_tab moves to a module-level helper, taking the component arrow from complexity 21 to 18, and the two four-property object literals in build_complexity_router_config.test.ts move into named variables. No behavior change; the touched suites pass (101 tests) --- .../add_model/add_auto_router_tab.tsx | 30 +++++++++++++------ .../build_complexity_router_config.test.ts | 10 +++---- 2 files changed, 25 insertions(+), 15 deletions(-) diff --git a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx index a7516b2a3a1..eeabfc681c8 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx @@ -102,6 +102,21 @@ const tierConfigSummary = (tiers: ComplexityTiers): string => { return parts.length > 0 ? parts.join(" ยท ") : "No tiers configured yet"; }; +// Why the submit is unavailable, or null when it is available. The button reads this to disable +// itself and to say what is missing, so the two can never give different answers. Checks the +// config actually being built, not which preset (if any) it came from: a preset only ever +// prefills once (handlePresetChange), and everything after that is edited exactly like Custom. +const getSubmitBlockedReason = ( + config: ComplexityRouterConfigValue, + keywordTierRules: KeywordTierRule[], + referencedModelsParams: Parameters[0], + availableModelSet: Set, +): string | null => + getMissingTiersError(config.tiers) ?? + getTierLabelsError(config.tier_labels) ?? + getKeywordTierRulesError(keywordTierRules) ?? + getReferencedModelsError(referencedModelsParams, availableModelSet); + const AddAutoRouterTab: React.FC = ({ handleOk, accessToken, @@ -222,15 +237,12 @@ const AddAutoRouterTab: React.FC = ({ embeddingModel, }; - // Why the submit is unavailable, or null when it is available. The button reads this to disable - // itself and to say what is missing, so the two can never give different answers. Checks the - // config actually being built, not which preset (if any) it came from: a preset only ever - // prefills once (handlePresetChange), and everything after that is edited exactly like Custom. - const submitBlockedReason = - getMissingTiersError(complexityRouterConfig.tiers) ?? - getTierLabelsError(complexityRouterConfig.tier_labels) ?? - getKeywordTierRulesError(keywordTierRules) ?? - getReferencedModelsError(referencedModelsParams, availableModelSet); + const submitBlockedReason = getSubmitBlockedReason( + complexityRouterConfig, + keywordTierRules, + referencedModelsParams, + availableModelSet, + ); const complexityRouterConfigParams: BuildComplexityRouterConfigParams = { tiers: complexityRouterConfig.tiers, diff --git a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts index bcbf50bbdea..187ec7070f2 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts @@ -437,9 +437,8 @@ describe("getTierLabelsError", () => { }); it("accepts a full distinct rename", () => { - expect( - getTierLabelsError({ SIMPLE: "Cheap", MEDIUM: "Standard", COMPLEX: "Premium", REASONING: "Deep" }), - ).toBeNull(); + const fullRename = { SIMPLE: "Cheap", MEDIUM: "Standard", COMPLEX: "Premium", REASONING: "Deep" }; + expect(getTierLabelsError(fullRename)).toBeNull(); }); it("rejects two tiers sharing a name, which would be ambiguous in the logs", () => { @@ -473,9 +472,8 @@ describe("hydrateTierLabels", () => { }); it("drops non-string and blank values a hand-edited config could hold", () => { - expect(hydrateTierLabels({ SIMPLE: 7, MEDIUM: " ", COMPLEX: null, REASONING: "Deep" })).toEqual({ - REASONING: "Deep", - }); + const handEdited = { SIMPLE: 7, MEDIUM: " ", COMPLEX: null, REASONING: "Deep" }; + expect(hydrateTierLabels(handEdited)).toEqual({ REASONING: "Deep" }); }); it("ignores keys that are not tiers", () => { From 0b8c58735da79ba1b4f257bad843b6d77f5f152f Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Wed, 5 Aug 2026 12:03:15 -0700 Subject: [PATCH 10/11] fix(ci): make the env-key doc gate see get_secret_bool reads (#35833) The gate only matched os.getenv(, litellm.get_secret( and litellm.get_secret_str(, so a bare get_secret_bool("X") matched nothing and the key bypassed the documentation requirement entirely. Add a fourth pattern for get_secret_bool, with or without the litellm. prefix, and a negative lookbehind so an unrelated receiver's .get_secret*( call is not mistaken for an env var read. Extraction and table parsing move into functions behind a __main__ guard so the patterns can be unit tested; the script is still invoked exactly the same way by CI. This surfaces 13 keys the gate never checked, 8 of which have no reference row yet. --- tests/documentation_tests/test_env_keys.py | 156 +++++++++----------- tests/test_litellm/test_env_key_doc_gate.py | 103 +++++++++++++ 2 files changed, 172 insertions(+), 87 deletions(-) create mode 100644 tests/test_litellm/test_env_key_doc_gate.py diff --git a/tests/documentation_tests/test_env_keys.py b/tests/documentation_tests/test_env_keys.py index 3bf2c88a848..31ba7ca9379 100644 --- a/tests/documentation_tests/test_env_keys.py +++ b/tests/documentation_tests/test_env_keys.py @@ -1,20 +1,19 @@ import os import re +from collections.abc import Iterator # Define the base directory for the litellm repository and documentation path repo_base = "./litellm" # Change this to your actual path -# Regular expressions to capture the keys used in os.getenv() and litellm.get_secret() -getenv_pattern = re.compile(r'os\.getenv\(\s*[\'"]([^\'"]+)[\'"]\s*(?:,\s*[^)]*)?\)') -get_secret_pattern = re.compile( - r'litellm\.get_secret\(\s*[\'"]([^\'"]+)[\'"]\s*(?:,\s*[^)]*|,\s*default_value=[^)]*)?\)' -) -get_secret_str_pattern = re.compile( - r'litellm\.get_secret_str\(\s*[\'"]([^\'"]+)[\'"]\s*(?:,\s*[^)]*|,\s*default_value=[^)]*)?\)' -) +_GETENV_ARGS = r"""\(\s*['"]([^'"]+)['"]\s*(?:,\s*[^)]*)?\)""" +_GET_SECRET_ARGS = r"""\(\s*['"]([^'"]+)['"]\s*(?:,\s*[^)]*|,\s*default_value=[^)]*)?\)""" -# Set to store unique keys from the code -env_keys = set() +ENV_KEY_PATTERNS: tuple[re.Pattern[str], ...] = ( + re.compile(r"os\.getenv" + _GETENV_ARGS), + re.compile(r"litellm\.get_secret" + _GET_SECRET_ARGS), + re.compile(r"litellm\.get_secret_str" + _GET_SECRET_ARGS), + re.compile(r"(? frozenset[str]: + """Return every documentable env var name read by the given Python source.""" + return frozenset( + match for pattern in ENV_KEY_PATTERNS for match in pattern.findall(source) if match not in EXCLUDED_KEYS ) -print(f"documented_keys: {documented_keys}") -# Compare and find undocumented keys -undocumented_keys = env_keys - documented_keys +def collect_env_keys(base_dir: str) -> frozenset[str]: + """Return every documentable env var name read anywhere under ``base_dir``.""" + return frozenset(key for file_path in _python_files(base_dir) for key in extract_env_keys(_read_text(file_path))) -# Print results -print("Keys expected in 'environment settings' (found in code):") -for key in sorted(env_keys): - print(key) -if undocumented_keys: - raise Exception( - f"\nKeys not documented in 'environment settings - Reference': {undocumented_keys}" +def _python_files(base_dir: str) -> Iterator[str]: + for root, dirs, files in os.walk(base_dir): + # Skip dependency/venv directories - prevents picking up env vars from installed packages + dirs[:] = [d for d in dirs if d not in SKIP_DIRS] + yield from (os.path.join(root, name) for name in files if name.endswith(".py")) + + +def _read_text(file_path: str) -> str: + with open(file_path, "r", encoding="utf-8") as f: + return f.read() + + +def extract_documented_keys(docs_content: str) -> frozenset[str]: + """Return the key names listed in the 'environment variables - Reference' table.""" + section = re.search( + r"### environment variables - Reference(.*?)(?=\n###|\Z)", + docs_content, + re.DOTALL | re.MULTILINE, ) -else: - print( - "\nAll keys are documented in 'environment settings - Reference'. - {}".format( - env_keys - ) + if section is None: + return frozenset() + # Match | KEY_NAME | description | - capture first column only + return frozenset( + match.group(1).strip() + for match in (re.match(r"^\|\s*([A-Z_][A-Z0-9_]*)\s*\|", line) for line in section.group(1).split("\n")) + if match is not None ) + + +def main() -> None: + env_keys = collect_env_keys(repo_base) + print(env_keys) + + docs_path = "./docs/my-website/docs/proxy/config_settings.md" # Path to the documentation + try: + documented_keys = extract_documented_keys(_read_text(docs_path)) + except Exception as e: + raise Exception(f"Error reading documentation: {e}, \n repo base - {os.listdir('./')}") + + print(f"documented_keys: {documented_keys}") + undocumented_keys = env_keys - documented_keys + + print("Keys expected in 'environment settings' (found in code):") + for key in sorted(env_keys): + print(key) + + if undocumented_keys: + raise Exception(f"\nKeys not documented in 'environment settings - Reference': {sorted(undocumented_keys)}") + print(f"\nAll keys are documented in 'environment settings - Reference'. - {env_keys}") + + +if __name__ == "__main__": + main() diff --git a/tests/test_litellm/test_env_key_doc_gate.py b/tests/test_litellm/test_env_key_doc_gate.py new file mode 100644 index 00000000000..aabda09a441 --- /dev/null +++ b/tests/test_litellm/test_env_key_doc_gate.py @@ -0,0 +1,103 @@ +"""Tests for the env-var extraction used by tests/documentation_tests/test_env_keys.py. + +That script is the CI gate that fails when a user-facing environment variable read +under litellm/ has no row in the docs reference table. It only sees a key if one of its +patterns matches the call, so a call shape the patterns miss silently bypasses the gate. +Each supported shape is asserted here, along with the shapes that must not be treated as +env var reads, so narrowing a pattern makes a test fail instead of quietly reopening the +hole. +""" + +import importlib.util +import sys +from pathlib import Path + +_REPO_ROOT = Path(__file__).resolve().parents[2] +_MODULE_PATH = _REPO_ROOT / "tests" / "documentation_tests" / "test_env_keys.py" +_spec = importlib.util.spec_from_file_location("documentation_test_env_keys", _MODULE_PATH) +assert _spec is not None and _spec.loader is not None +gate = importlib.util.module_from_spec(_spec) +sys.modules[_spec.name] = gate +_spec.loader.exec_module(gate) + + +def test_bare_get_secret_bool_is_captured() -> None: + assert gate.extract_env_keys('flag = get_secret_bool("QSTASH_FLUSH_ON_BOOT")') == {"QSTASH_FLUSH_ON_BOOT"} + + +def test_get_secret_bool_with_default_is_captured() -> None: + assert gate.extract_env_keys('if get_secret_bool("QSTASH_FLUSH_ON_BOOT", False) is not True:') == { + "QSTASH_FLUSH_ON_BOOT" + } + + +def test_get_secret_bool_with_keyword_default_is_captured() -> None: + assert gate.extract_env_keys('get_secret_bool("QSTASH_FLUSH_ON_BOOT", default_value=False)') == { + "QSTASH_FLUSH_ON_BOOT" + } + + +def test_litellm_prefixed_get_secret_bool_is_captured() -> None: + assert gate.extract_env_keys('litellm.get_secret_bool("QSTASH_FLUSH_ON_BOOT")') == {"QSTASH_FLUSH_ON_BOOT"} + + +def test_previously_supported_call_shapes_are_still_captured() -> None: + source = "\n".join( + ( + 'os.getenv("QSTASH_ALPHA")', + 'os.getenv("QSTASH_BRAVO", "fallback")', + 'litellm.get_secret("QSTASH_CHARLIE")', + 'litellm.get_secret_str("QSTASH_DELTA", default_value=None)', + ) + ) + assert gate.extract_env_keys(source) == {"QSTASH_ALPHA", "QSTASH_BRAVO", "QSTASH_CHARLIE", "QSTASH_DELTA"} + + +def test_get_secret_calls_on_unrelated_objects_are_not_env_reads() -> None: + source = "\n".join( + ( + 'vault_client.get_secret("QSTASH_ALPHA")', + 'self.get_secret_str("QSTASH_BRAVO")', + 'provider.get_secret_bool("QSTASH_CHARLIE")', + ) + ) + assert gate.extract_env_keys(source) == frozenset() + + +def test_similarly_named_helpers_are_not_env_reads() -> None: + assert gate.extract_env_keys('get_secret_bundle("QSTASH_ALPHA")') == frozenset() + + +def test_non_literal_arguments_are_not_env_reads() -> None: + assert gate.extract_env_keys("get_secret_bool(flag_name)") == frozenset() + + +def test_excluded_keys_are_filtered_for_every_call_shape() -> None: + source = "\n".join( + ( + 'os.getenv("TERM_PROGRAM")', + 'get_secret_bool("LITELLM_RUST")', + 'litellm.get_secret_str("MAVVRIK_FOCUS_FREQUENCY")', + ) + ) + assert gate.extract_env_keys(source) == frozenset() + + +def test_documented_keys_are_read_from_the_reference_table_only() -> None: + docs = "\n".join( + ( + "### general_settings - Reference", + "| BEFORE_THE_TABLE | not the env var table", + "", + "### environment variables - Reference", + "", + "| Name | Description |", + "|------|-------------|", + "| QSTASH_ALPHA | first key", + "| QSTASH_BRAVO | second key", + "", + "### another section - Reference", + "| AFTER_THE_TABLE | also not the env var table", + ) + ) + assert gate.extract_documented_keys(docs) == {"QSTASH_ALPHA", "QSTASH_BRAVO"} From 2a9843e649a4336927646c237632b97acc451e59 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Wed, 5 Aug 2026 12:27:49 -0700 Subject: [PATCH 11/11] fix(proxy): keep the connected DB client when a startup health check fails (#35837) `_setup_prisma_client` ran `connect()`, then a `SELECT 1` health check, then armed the DB health watchdog. Any failure fell into one handler that, with `allow_requests_on_db_unavailable` set, swallowed the error and returned None, which the caller assigns to the module-level `prisma_client`. A single transient timeout on that health check therefore discarded a client that had already connected, for the life of the process, and skipped the watchdog that exists to reconnect it. The watchdog now starts before the health check, and a swallowed post-connect failure returns the connected client instead of None. A client whose `connect()` failed is still discarded, and startup still hard-fails when `allow_requests_on_db_unavailable` is not set. The same check also misreported its own failure. `health_check()` labelled its error `disconnect()`, a copy-paste from the real `disconnect()` below it, so grepping the logs for the health check turned up nothing and read as "the check never ran". Both it and the sibling `connect()` failure reported through `print_verbose`, which reaches `verbose_proxy_logger.debug` and otherwise prints only under the deprecated `litellm.set_verbose`, leaving a startup-blocking database fault invisible at the verbosity operators actually run. Both now log at warning under their own names. The proxy logger's handler carries the secret redaction filter, so a connection string in the exception text is redacted exactly as it was on the old print path. --- litellm/proxy/proxy_server.py | 72 ++++++----- litellm/proxy/utils.py | 6 +- tests/test_litellm/proxy/test_proxy_server.py | 120 ++++++++++++++++++ tests/test_litellm/proxy/test_proxy_utils.py | 76 +++++++++++ 4 files changed, 239 insertions(+), 35 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 61c6ce22a91..8d6b930591b 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -8664,48 +8664,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 '") - 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 '") + 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): diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 8d638dedff8..7717b4da1af 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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 diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index ede93dc0c58..4a491ec0cff 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -11018,3 +11018,123 @@ def test_startup_is_silent_when_mock_testing_params_disabled(caplog): ProxyStartupEvent._warn_if_mock_testing_params_enabled(general_settings={}) assert MOCK_TESTING_CONFIG_KEY not in caplog.text + + +def _mock_startup_prisma_client(health_check_error=None, connect_error=None): + client = MagicMock() + client.connect = AsyncMock(side_effect=connect_error) + client.db.start_token_refresh_task = AsyncMock() + client.check_view_exists = AsyncMock() + client._set_spend_logs_row_count_in_proxy_state = AsyncMock() + client.start_db_health_watchdog_task = AsyncMock() + client.health_check = AsyncMock(side_effect=health_check_error) + return client + + +async def _run_setup_prisma_client(mock_client): + from litellm.proxy.proxy_server import ProxyStartupEvent + + with patch.object(proxy_server_module, "PrismaClient", return_value=mock_client): + result = await ProxyStartupEvent._setup_prisma_client( + database_url="postgresql://litellm:litellm@localhost:5432/litellm", + proxy_logging_obj=MagicMock(), + user_api_key_cache=DualCache(), + ) + await asyncio.sleep(0.05) + return result + + +@pytest.mark.asyncio +async def test_setup_prisma_client_retains_connected_client_when_startup_health_check_fails( + monkeypatch, +): + """A transient failure of the startup ``SELECT 1`` must not discard a client + whose ``connect()`` already succeeded. + + Discarding it assigns ``None`` to the module-level ``prisma_client`` for the + life of the process, so a database that came back a second later is never + used again until the proxy is restarted.""" + monkeypatch.setenv("DISABLE_PRISMA_HEALTH_CHECK_ON_STARTUP", "False") + monkeypatch.setattr( + proxy_server_module, + "general_settings", + {"allow_requests_on_db_unavailable": True}, + ) + + mock_client = _mock_startup_prisma_client( + health_check_error=httpx.ReadTimeout("startup health check timed out") + ) + result = await _run_setup_prisma_client(mock_client) + + assert mock_client.connect.await_count == 1 + assert mock_client.health_check.await_count == 1 + assert result is mock_client + + +@pytest.mark.asyncio +async def test_setup_prisma_client_arms_health_watchdog_before_startup_health_check( + monkeypatch, +): + """The health watchdog is the only thing that reconnects a dropped DB, so it + has to be armed before the startup health check can fail. + + Armed after, the single failure it exists to recover from is exactly the one + that skips it, and recovery never happens.""" + monkeypatch.setenv("DISABLE_PRISMA_HEALTH_CHECK_ON_STARTUP", "False") + monkeypatch.setattr( + proxy_server_module, + "general_settings", + {"allow_requests_on_db_unavailable": True}, + ) + + mock_client = _mock_startup_prisma_client( + health_check_error=httpx.ReadTimeout("startup health check timed out") + ) + call_order = MagicMock() + call_order.attach_mock(mock_client.start_db_health_watchdog_task, "watchdog") + call_order.attach_mock(mock_client.health_check, "health_check") + + await _run_setup_prisma_client(mock_client) + + assert mock_client.start_db_health_watchdog_task.await_count == 1 + assert [call[0] for call in call_order.mock_calls] == ["watchdog", "health_check"] + + +@pytest.mark.asyncio +async def test_setup_prisma_client_raises_when_db_unavailable_is_not_allowed(monkeypatch): + """Without ``allow_requests_on_db_unavailable`` a failed startup health check + must still hard-fail startup. Retaining the client is a fallback for + operators who opted into serving traffic without a database, never a way to + boot a proxy whose DB never answered.""" + monkeypatch.setenv("DISABLE_PRISMA_HEALTH_CHECK_ON_STARTUP", "False") + monkeypatch.setattr( + proxy_server_module, + "general_settings", + {"allow_requests_on_db_unavailable": False}, + ) + + mock_client = _mock_startup_prisma_client( + health_check_error=httpx.ReadTimeout("startup health check timed out") + ) + with pytest.raises(httpx.ReadTimeout): + await _run_setup_prisma_client(mock_client) + + +@pytest.mark.asyncio +async def test_setup_prisma_client_returns_none_when_connect_itself_fails(monkeypatch): + """Retaining only ever applies to a client that connected. If ``connect()`` + failed there is no usable client and no watchdog to recover it, so the caller + must still get ``None``.""" + monkeypatch.setenv("DISABLE_PRISMA_HEALTH_CHECK_ON_STARTUP", "False") + monkeypatch.setattr( + proxy_server_module, + "general_settings", + {"allow_requests_on_db_unavailable": True}, + ) + + mock_client = _mock_startup_prisma_client(connect_error=httpx.ConnectError("connection refused")) + result = await _run_setup_prisma_client(mock_client) + + assert result is None + assert mock_client.start_db_health_watchdog_task.await_count == 0 + assert mock_client.health_check.await_count == 0 diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index 3421751d962..abd6220144b 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -1085,3 +1085,79 @@ async def test_post_mcp_call_hook_propagates_guardrail_block(restore_callbacks): request_data={"mcp_tool_name": "echo"}, user_api_key_dict=None, ) + + +@pytest.mark.asyncio +async def test_prisma_health_check_failure_names_itself_at_operator_visible_level(caplog): + """A failing DB health check has to name the check that failed, at a level + operators actually run at. + + Reporting it as ``disconnect()`` sends anyone grepping the logs to the wrong + function and reads as "the check never ran", and reporting it only at debug + level hides a database fault behind a flag nobody enables in production.""" + import logging + from unittest.mock import AsyncMock + + from litellm.proxy.utils import PrismaClient + + client = MagicMock() + client.db.query_raw = AsyncMock(side_effect=Exception("connection refused")) + client.proxy_logging_obj.failure_handler = AsyncMock() + + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + with pytest.raises(Exception, match="connection refused"): + await PrismaClient.health_check(client) + + assert "health_check()" in caplog.text + assert "disconnect()" not in caplog.text + assert "connection refused" in caplog.text + + +@pytest.mark.asyncio +async def test_prisma_connect_failure_is_reported_at_operator_visible_level(caplog): + """The sibling connect failure is labelled correctly but was equally + invisible. A database the proxy could not connect to at startup must not be + a debug-only record.""" + import logging + from unittest.mock import AsyncMock + + from litellm.proxy.utils import PrismaClient + + client = MagicMock() + client.db.is_connected = MagicMock(return_value=False) + client.db.connect = AsyncMock(side_effect=Exception("could not reach database")) + client.proxy_logging_obj.failure_handler = AsyncMock() + + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + with pytest.raises(Exception, match="could not reach database"): + await PrismaClient.connect(client) + + assert "connect()" in caplog.text + assert "could not reach database" in caplog.text + + +@pytest.mark.asyncio +async def test_prisma_health_check_failure_redacts_database_credentials(caplog): + """Raising the level must not widen what reaches the logs. The exception + text can carry a full connection string, so the credential has to be gone + from the emitted record.""" + import logging + from unittest.mock import AsyncMock + + from litellm.proxy.utils import PrismaClient + + client = MagicMock() + client.db.query_raw = AsyncMock( + side_effect=Exception("could not connect to postgresql://admin:hunter2@db.internal:5432/litellm") + ) + client.proxy_logging_obj.failure_handler = AsyncMock() + + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + with pytest.raises(Exception): + await PrismaClient.health_check(client) + + emitted = [record.getMessage() for record in caplog.records if record.name == "LiteLLM Proxy"] + + assert emitted + assert all("hunter2" not in message for message in emitted) + assert any("postgresql://REDACTED@db.internal" in message for message in emitted)