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/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/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/caching/evicted_client_closer.py b/litellm/caching/evicted_client_closer.py new file mode 100644 index 00000000000..eee7e2ea289 --- /dev/null +++ b/litellm/caching/evicted_client_closer.py @@ -0,0 +1,276 @@ +""" +Deferred close of HTTP/SDK clients that the LLM client cache has evicted. + +Eviction only drops the cache's reference to a client. Every OpenAI/Azure SDK +client is a reference cycle (each resource namespace holds the client back), so +an evicted client and its pooled TCP connections survive until a generational +collection runs, which under load is thousands of requests later. + +Closing at eviction time is not an option: a request that was handed the client +just before it was evicted is still using it, and closing it underneath that +request raises ``RuntimeError: Cannot send a request, as the client has been +closed.`` + +So an evicted client is closed once two conditions hold. A grace window must +have passed since its eviction, which covers a request that holds the client +but is momentarily not on the wire, and the client must report no connection in +flight. The second condition is what keeps the first honest: a request may run +for ``litellm.request_timeout`` seconds, 6000 by default, and a streaming +response is bounded only by how long the upstream keeps sending, so no deadline +on its own can promise that a request has finished. + +Only clients litellm itself created are closed; a client the caller supplied is +left alone because litellm does not own its lifecycle. + +A client that closes synchronously is closed from wherever the cache is next +used. One whose close is a coroutine needs the event loop it was evicted on, so +it waits for a call from that loop rather than having work scheduled onto a loop +it does not belong to. Queued clients are therefore bucketed by what it takes to +close them, and each bucket is ordered by deadline, so a reap walks the entries +that are due rather than the whole queue. + +The queue holds its clients weakly, so waiting out a grace window never keeps +alive anything the collector would have reclaimed first. +""" + +import asyncio +import contextlib +import inspect +import threading +import time +import weakref +from collections import deque +from collections.abc import Awaitable, Callable, Iterator +from dataclasses import dataclass, replace +from typing import Final + +from litellm.constants import ( + EVICTED_LLM_CLIENT_CLOSE_GRACE_SECONDS, + EVICTED_LLM_CLIENT_CLOSE_MAX_PENDING, +) + +_CLOSABLE_ANYWHERE: Final = "closable-anywhere" +_CLOSABLE_ON_ANY_LOOP: Final = "closable-on-any-loop" + +_BucketKey = str | int + + +@dataclass(frozen=True, slots=True) +class _PendingClose: + """A queued close. + + The client is held weakly, so queueing one never keeps alive anything the + collector would otherwise have reclaimed first. + + ``needs_loop`` is set for a client whose close is a coroutine; those can only + be closed from the event loop they were evicted on, recorded in ``loop_id``. + A client that closes synchronously carries neither constraint. + """ + + client_ref: "weakref.ref[object]" + loop_id: int | None + needs_loop: bool + close_after: float + + +def _bucket_key(pending: _PendingClose) -> _BucketKey: + """Which reaps can close this entry: any at all, any running a loop, or one loop's.""" + if not pending.needs_loop: + return _CLOSABLE_ANYWHERE + if pending.loop_id is None: + return _CLOSABLE_ON_ANY_LOOP + return pending.loop_id + + +def _running_loop_id() -> int | None: + try: + return id(asyncio.get_running_loop()) + except RuntimeError: + return None + + +def _close_function(client: object) -> Callable[[], object] | None: + close_fn: Final[Callable[[], object] | None] = getattr(client, "aclose", None) or getattr(client, "close", None) + return close_fn + + +def _transport_of(client: object) -> object: + """The httpx transport behind an SDK wrapper, a litellm handler, or a bare client.""" + for holder in (getattr(client, "_client", None), getattr(client, "client", None), client): + transport: object = getattr(holder, "_transport", None) + if transport is not None: + return transport + return None + + +def _connection_is_idle(connection: object) -> bool: + """A pooled connection is idle unless it is servicing a request.""" + is_idle: Final[object] = getattr(connection, "is_idle", None) + return bool(is_idle()) if callable(is_idle) else True + + +def _pool_has_busy_connection(transport: object) -> bool | None: + """Whether the httpcore pool behind the transport is servicing a request. + + ``None`` when there is no such pool, so the caller can ask the other backend. + """ + pooled: Final[object] = getattr(getattr(transport, "_pool", None), "connections", None) + if not isinstance(pooled, (list, tuple)): + return None + return any( + not _connection_is_idle(connection) # pyright: ignore[reportUnknownArgumentType] # untyped pool list + for connection in pooled # pyright: ignore[reportUnknownVariableType] # untyped pool list + ) + + +def _has_connection_in_flight(client: object) -> bool: + """Whether the client is servicing a request right now. + + Both connection backends litellm uses already account for the connections + they have handed out, so this reads the client's own lease accounting rather + than inferring it from elapsed time: httpcore reports a non-idle connection + for the whole of a response including a stream, and aiohttp holds the + connection in ``_acquired`` over the same span. + + A client that cannot answer is reported as idle, which leaves the grace + window as the only guard, exactly as it was before this check existed. + """ + try: + transport: Final = _transport_of(client) + pooled_busy: Final = _pool_has_busy_connection(transport) + if pooled_busy is not None: + return pooled_busy + session: Final[object] = getattr(transport, "client", None) + return bool(getattr(getattr(session, "connector", None), "_acquired", None)) + except Exception: # noqa: BLE001 - a client that cannot report its state is treated as idle + return False + + +async def _close_quietly(closing: Awaitable[object]) -> None: + with contextlib.suppress(Exception): + await closing + + +class EvictedClientCloser: + """Closes evicted, litellm-owned clients once they are idle and out of grace.""" + + def __init__( + self, + grace_seconds: float = EVICTED_LLM_CLIENT_CLOSE_GRACE_SECONDS, + max_pending: int = EVICTED_LLM_CLIENT_CLOSE_MAX_PENDING, + clock: Callable[[], float] = time.monotonic, + ) -> None: + self._grace_seconds = grace_seconds + self._max_pending = max_pending + self._clock = clock + self._owned: weakref.WeakSet[object] = weakref.WeakSet() + self._buckets: dict[_BucketKey, deque[_PendingClose]] = {} # mutable-ok: deadline-ordered queues + self._pending_count = 0 + self._queue_lock = threading.Lock() # the cache is reachable from every worker thread's loop + self._close_tasks: set[asyncio.Task[None]] = set() # mutable-ok: strong refs to running closes + + def mark_owned(self, client: object) -> None: + """Record that litellm created this client, so it may be closed on eviction.""" + try: + self._owned.add(client) + except TypeError: + pass # values that cannot be weak-referenced are never litellm clients + + def _is_owned(self, client: object) -> bool: + try: + return client in self._owned + except TypeError: + return False # unhashable values are never litellm clients + + def schedule(self, client: object) -> None: + """Queue an evicted client for closing once it is idle and out of grace. + + Past ``max_pending`` the client is left to the collector instead, so a + workload that churns the cache cannot grow this queue without bound. + Every queued entry comes due within one grace window, so the capacity it + occupies is returned within that window rather than held. + """ + if client is None or not self._is_owned(client): + return + close_fn: Final = _close_function(client) + if close_fn is None: + return + if self._pending_count >= self._max_pending: + return + self._enqueue( + _PendingClose( + client_ref=weakref.ref(client), + loop_id=_running_loop_id(), + needs_loop=inspect.iscoroutinefunction(close_fn), + close_after=self._clock() + self._grace_seconds, + ) + ) + + def reap(self) -> None: + """Close every queued client that is due, idle, and closable from here. + + Called from the cache's read path, so the empty-queue exit comes first and + the work done past it is proportional to what is due, not to the queue. + """ + if not self._pending_count: + return + now: Final = self._clock() + for pending in self._take_due(_running_loop_id(), now): + client = pending.client_ref() + if client is None: + continue + if _has_connection_in_flight(client): + self._enqueue(replace(pending, close_after=now + self._grace_seconds)) + continue + self._close(client) + + @property + def pending_count(self) -> int: + return self._pending_count + + def _enqueue(self, pending: _PendingClose) -> None: + """Append to the entry's bucket, dropping any dead entries it queues behind. + + Deadlines only ever move forward, so appending keeps each bucket ordered + by deadline, and entries whose client the collector already took sit at + the front rather than having to be searched for. + """ + with self._queue_lock: + bucket: Final = self._buckets.setdefault(_bucket_key(pending), deque()) # mutable-ok: FIFO by design + while bucket and bucket[0].client_ref() is None: + bucket.popleft() + self._pending_count -= 1 + bucket.append(pending) + self._pending_count += 1 + + def _take_due(self, loop_id: int | None, now: float) -> tuple[_PendingClose, ...]: + buckets = (_CLOSABLE_ANYWHERE,) if loop_id is None else (_CLOSABLE_ANYWHERE, _CLOSABLE_ON_ANY_LOOP, loop_id) + with self._queue_lock: + return tuple(pending for key in buckets for pending in self._drain_locked(key, now)) + + def _drain_locked(self, key: _BucketKey, now: float) -> Iterator[_PendingClose]: + bucket: Final = self._buckets.get(key) + if bucket is None: + return + while bucket and bucket[0].close_after <= now: + self._pending_count -= 1 + yield bucket.popleft() + if not bucket: + del self._buckets[key] + + def _close(self, client: object) -> None: + close_fn: Final = _close_function(client) + if close_fn is None: + return + try: + closing: Final = close_fn() + except Exception: # noqa: BLE001 - a discarded client's close must never surface to callers + return + if not inspect.isawaitable(closing): + return + task: Final = asyncio.get_running_loop().create_task(_close_quietly(closing)) + self._close_tasks.add(task) + task.add_done_callback(self._close_tasks.discard) + + +default_evicted_client_closer: Final = EvictedClientCloser() diff --git a/litellm/caching/llm_caching_handler.py b/litellm/caching/llm_caching_handler.py index 7d072a40195..a89e43b78b4 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, + ) -> 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/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/azure/common_utils.py b/litellm/llms/azure/common_utils.py index 28ed6ef9681..1ce83e226e7 100644 --- a/litellm/llms/azure/common_utils.py +++ b/litellm/llms/azure/common_utils.py @@ -427,88 +427,95 @@ class BaseAzureLLM(BaseOpenAILLM): f"|azure_password={hashlib.sha256(_azure_password.encode()).hexdigest() if isinstance(_azure_password, str) else None}" f"|azure_scope={_lp.get('azure_scope')}" ) - if client is None: - cached_client: Final = self.get_cached_openai_client( - client_initialization_params=client_initialization_params, - client_type="azure", - ) - if cached_client: - if isinstance(cached_client, (AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI)): - return cached_client - - azure_client_params: Final = self.initialize_azure_sdk_client( - litellm_params=litellm_params or {}, - api_key=api_key, - api_base=api_base, - model_name=model, - api_version=api_version, - is_async=_is_async, - ) - - # For Azure v1 API, use standard OpenAI client instead of AzureOpenAI - # See: https://learn.microsoft.com/en-us/azure/ai-services/openai/reference#api-specs - if self._is_azure_v1_api_version(api_version): - # Extract only params that OpenAI client accepts - # Always use /openai/v1/ regardless of whether user passed "v1", "latest", or "preview" - # The OpenAI client accepts a callable for `api_key` and re-invokes it - # on every request (via `_refresh_api_key`), so passing - # `azure_ad_token_provider` directly preserves Azure AD token refresh - # behavior that the regular AzureOpenAI client provides. - v1_api_key: str | Callable[[], Any] | None = ( - azure_client_params.get("api_key") - or azure_client_params.get("azure_ad_token_provider") - or azure_client_params.get("azure_ad_token") - ) - if _is_async is True and callable(v1_api_key): - # AsyncOpenAI expects an async provider; wrap the sync provider - # returned by azure-identity. Offload to a thread so a token - # refresh (blocking HTTP call to AAD on cache miss) does not - # stall the event loop. - _sync_provider: Final = v1_api_key - - async def _async_v1_api_key() -> str: - return await asyncio.to_thread(_sync_provider) - - v1_api_key = _async_v1_api_key - - v1_params: Final[dict[str, Any]] = { - "api_key": v1_api_key, - "base_url": f"{api_base}/openai/v1/", - } - if "timeout" in azure_client_params: - v1_params["timeout"] = azure_client_params["timeout"] - if "max_retries" in azure_client_params: - v1_params["max_retries"] = azure_client_params["max_retries"] - if "http_client" in azure_client_params: - v1_params["http_client"] = azure_client_params["http_client"] - - verbose_logger.debug("Using Azure v1 API with base_url: %s", v1_params["base_url"]) - - if _is_async is True: - openai_client = AsyncOpenAI(**v1_params) - else: - openai_client = OpenAI(**v1_params) - else: - # Traditional Azure API uses AzureOpenAI client - if _is_async is True: - openai_client = AsyncAzureOpenAI(**azure_client_params) - else: - openai_client = AzureOpenAI(**azure_client_params) - else: - openai_client = client + if client is not None: if ( api_version is not None - and isinstance(openai_client, (AzureOpenAI, AsyncAzureOpenAI)) - and isinstance(openai_client._custom_query, dict) + and isinstance(client, (AzureOpenAI, AsyncAzureOpenAI)) + and isinstance(client._custom_query, dict) ): # set api_version to version passed by user - openai_client._custom_query.setdefault("api-version", api_version) + client._custom_query.setdefault("api-version", api_version) + self.set_cached_openai_client( + openai_client=client, + client_initialization_params=client_initialization_params, + client_type="azure", + litellm_owned_client=False, + ) + return client + + cached_client: Final = self.get_cached_openai_client( + client_initialization_params=client_initialization_params, + client_type="azure", + ) + if cached_client: + if isinstance(cached_client, (AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI)): + return cached_client + + azure_client_params: Final = self.initialize_azure_sdk_client( + litellm_params=litellm_params or {}, + api_key=api_key, + api_base=api_base, + model_name=model, + api_version=api_version, + is_async=_is_async, + ) + + # For Azure v1 API, use standard OpenAI client instead of AzureOpenAI + # See: https://learn.microsoft.com/en-us/azure/ai-services/openai/reference#api-specs + if self._is_azure_v1_api_version(api_version): + # Extract only params that OpenAI client accepts + # Always use /openai/v1/ regardless of whether user passed "v1", "latest", or "preview" + # The OpenAI client accepts a callable for `api_key` and re-invokes it + # on every request (via `_refresh_api_key`), so passing + # `azure_ad_token_provider` directly preserves Azure AD token refresh + # behavior that the regular AzureOpenAI client provides. + v1_api_key: str | Callable[[], Any] | None = ( + azure_client_params.get("api_key") + or azure_client_params.get("azure_ad_token_provider") + or azure_client_params.get("azure_ad_token") + ) + if _is_async is True and callable(v1_api_key): + # AsyncOpenAI expects an async provider; wrap the sync provider + # returned by azure-identity. Offload to a thread so a token + # refresh (blocking HTTP call to AAD on cache miss) does not + # stall the event loop. + _sync_provider: Final = v1_api_key + + async def _async_v1_api_key() -> str: + return await asyncio.to_thread(_sync_provider) + + v1_api_key = _async_v1_api_key + + v1_params: Final[dict[str, Any]] = { + "api_key": v1_api_key, + "base_url": f"{api_base}/openai/v1/", + } + if "timeout" in azure_client_params: + v1_params["timeout"] = azure_client_params["timeout"] + if "max_retries" in azure_client_params: + v1_params["max_retries"] = azure_client_params["max_retries"] + if "http_client" in azure_client_params: + v1_params["http_client"] = azure_client_params["http_client"] + + verbose_logger.debug("Using Azure v1 API with base_url: %s", v1_params["base_url"]) + + if _is_async is True: + openai_client = AsyncOpenAI(**v1_params) + else: + openai_client = OpenAI(**v1_params) + else: + # Traditional Azure API uses AzureOpenAI client + if _is_async is True: + openai_client = AsyncAzureOpenAI(**azure_client_params) + else: + openai_client = AzureOpenAI(**azure_client_params) # save client in-memory cache self.set_cached_openai_client( openai_client=openai_client, client_initialization_params=client_initialization_params, client_type="azure", + litellm_owned_client=self.owns_wrapped_http_client(azure_client_params.get("http_client")), ) return openai_client 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/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 3270332a6b0..c1421d0f969 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 bbb4c203460..82ebee3962e 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 6c3aec2452c..b46e49ac99e 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/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/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/_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/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/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/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..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): @@ -8876,7 +8884,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 +11487,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/litellm/proxy/utils.py b/litellm/proxy/utils.py index 5dc2da84de7..8f9bafbe22c 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 @@ -5004,7 +5004,7 @@ class PrismaClient: import traceback error_msg: Final = f"LiteLLM Prisma Client Exception health_check(): {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 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, 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/caching/test_evicted_client_closer.py b/tests/test_litellm/caching/test_evicted_client_closer.py new file mode 100644 index 00000000000..939be5f3d6b --- /dev/null +++ b/tests/test_litellm/caching/test_evicted_client_closer.py @@ -0,0 +1,411 @@ +""" +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() + held_client = handler.client + + 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 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 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() + + +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/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 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 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/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) 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"} diff --git a/ui/litellm-dashboard/eslint-budgets.json b/ui/litellm-dashboard/eslint-budgets.json index 3526d71ce90..c4f078f2ff2 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 } + "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 } } 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/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", () => { 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() ?? "");