Merge origin/litellm_internal_staging into litellm_health_check_planned_recreate
Some checks failed
LiteLLM Rust / rustfmt, clippy, test (push) Has been cancelled
Terraform Provider / gofmt, vet, build, test (push) Has been cancelled
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Has been cancelled

This commit is contained in:
Devin AI 2026-08-05 19:30:47 +00:00
commit 2f64cc6c03
92 changed files with 2707 additions and 442 deletions

View file

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

View file

View file

@ -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

View file

@ -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

View file

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

View file

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

View file

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

View file

@ -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"):

View file

@ -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:

View file

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

View file

@ -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:

View file

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

View file

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

View file

@ -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

View file

@ -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)

View file

@ -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 (

View file

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

View file

@ -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

View file

@ -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

View file

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

View file

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

View file

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

View file

@ -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)

View file

@ -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

View file

@ -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})

View file

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

View file

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

View file

@ -78,6 +78,7 @@ from litellm.proxy.management_endpoints.common_utils import (
_is_user_team_admin,
_set_object_metadata_field,
_team_member_has_permission,
_user_has_admin_view,
validate_finite_spend,
)
from litellm.proxy.management_endpoints.model_management_endpoints import (
@ -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:

View file

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

View file

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

View file

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

View file

@ -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 <migration_name>'")
raise e
prisma_client = PrismaClient(database_url=database_url, proxy_logging_obj=proxy_logging_obj)
## Start RDS IAM token refresh background task if enabled ##
# This proactively refreshes IAM tokens before they expire,
# preventing the 15-minute connection failure bug (#16220)
if hasattr(prisma_client, "db") and hasattr(prisma_client.db, "start_token_refresh_task"):
await prisma_client.db.start_token_refresh_task()
try:
await prisma_client.connect()
except Exception as e:
if "P3018" in str(e) or "P3009" in str(e):
verbose_proxy_logger.debug("CRITICAL: DATABASE MIGRATION FAILED")
verbose_proxy_logger.debug("Your database is in a 'dirty' state.")
verbose_proxy_logger.debug("FIX: Run 'prisma migrate resolve --applied <migration_name>'")
raise e
## Add necessary views to proxy ##
asyncio.create_task(
prisma_client.check_view_exists()
) # check if all necessary views exist. Don't block execution
connected_client = prisma_client
asyncio.create_task(
prisma_client._set_spend_logs_row_count_in_proxy_state()
) # set the spend logs row count in proxy state. Don't block execution
## Start RDS IAM token refresh background task if enabled ##
# This proactively refreshes IAM tokens before they expire,
# preventing the 15-minute connection failure bug (#16220)
if hasattr(prisma_client, "db") and hasattr(prisma_client.db, "start_token_refresh_task"):
await prisma_client.db.start_token_refresh_task()
# run a health check to ensure the DB is ready
if get_secret_bool("DISABLE_PRISMA_HEALTH_CHECK_ON_STARTUP", False) is not True:
await prisma_client.health_check()
## Add necessary views to proxy ##
asyncio.create_task(
prisma_client.check_view_exists()
) # check if all necessary views exist. Don't block execution
asyncio.create_task(
prisma_client._set_spend_logs_row_count_in_proxy_state()
) # set the spend logs row count in proxy state. Don't block execution
if hasattr(prisma_client, "start_db_health_watchdog_task"):
await prisma_client.start_db_health_watchdog_task()
# run a health check to ensure the DB is ready
if get_secret_bool("DISABLE_PRISMA_HEALTH_CHECK_ON_STARTUP", False) is not True:
await prisma_client.health_check()
if hasattr(prisma_client, "start_db_health_watchdog_task"):
await prisma_client.start_db_health_watchdog_task()
return prisma_client
except Exception as e:
PrismaDBExceptionHandler.handle_db_exception(e)
return None
if connected_client is not None:
verbose_proxy_logger.warning(
"Retaining the connected Prisma client after a post-connect startup step failed: %s. "
"The DB health watchdog keeps probing and reconnects once the database recovers.",
e,
)
return connected_client
@classmethod
def _init_dd_tracer(cls):
@ -8876,7 +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:

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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):

View file

@ -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,
)

View file

@ -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

View file

@ -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,
)

View file

@ -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,

View file

@ -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"(?<![\w.])(?:litellm\.)?get_secret_bool" + _GET_SECRET_ARGS),
)
# Terminal/environment detection variables that should not be documented
# These are internal variables used for terminal detection, not user-configurable settings
@ -48,6 +47,8 @@ EXCLUDED_TERMINAL_VARS = {
"ALACRITTY_SOCKET",
}
EXCLUDED_KEYS = frozenset(EXCLUDED_TERMINAL_VARS | EXCLUDED_GUARD_ONLY_VARS | EXCLUDED_ROLLOUT_FLAGS)
# Directories to skip (dependencies, venvs, caches) - only scan litellm source
SKIP_DIRS = {
".venv",
@ -61,88 +62,69 @@ SKIP_DIRS = {
"build",
}
# Walk through all files in the litellm repo to find references of os.getenv() and litellm.get_secret()
for root, dirs, files in os.walk(repo_base):
# Skip dependency/venv directories - prevents picking up env vars from installed packages
dirs[:] = [d for d in dirs if d not in SKIP_DIRS]
for file in files:
if file.endswith(".py"): # Only process Python files
file_path = os.path.join(root, file)
with open(file_path, "r", encoding="utf-8") as f:
content = f.read()
# Find all keys using os.getenv()
getenv_matches = getenv_pattern.findall(content)
env_keys.update(
match
for match in getenv_matches
if match not in EXCLUDED_TERMINAL_VARS
and match not in EXCLUDED_GUARD_ONLY_VARS
and match not in EXCLUDED_ROLLOUT_FLAGS
) # Extract only the key part, excluding terminal vars
# Find all keys using litellm.get_secret()
get_secret_matches = get_secret_pattern.findall(content)
env_keys.update(match for match in get_secret_matches)
# Find all keys using litellm.get_secret_str()
get_secret_str_matches = get_secret_str_pattern.findall(content)
env_keys.update(match for match in get_secret_str_matches)
# Print the unique keys found
print(env_keys)
# Parse the documentation to extract documented keys
repo_base = "./"
print(os.listdir(repo_base))
docs_path = (
"./docs/my-website/docs/proxy/config_settings.md" # Path to the documentation
)
documented_keys = set()
try:
with open(docs_path, "r", encoding="utf-8") as docs_file:
content = docs_file.read()
print(f"content: {content}")
# Find the section titled "general_settings - Reference"
general_settings_section = re.search(
r"### environment variables - Reference(.*?)(?=\n###|\Z)",
content,
re.DOTALL | re.MULTILINE,
)
print(f"general_settings_section: {general_settings_section}")
if general_settings_section:
# Extract the table rows - only first column (key name) from each row
table_content = general_settings_section.group(1)
for line in table_content.split("\n"):
# Match | KEY_NAME | description | - capture first column only
match = re.match(r"^\|\s*([A-Z_][A-Z0-9_]*)\s*\|", line)
if match:
documented_keys.add(match.group(1).strip())
except Exception as e:
raise Exception(
f"Error reading documentation: {e}, \n repo base - {os.listdir(repo_base)}"
def extract_env_keys(source: str) -> 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()

View file

@ -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"
)

View file

@ -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.).

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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."""

View file

@ -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

View file

@ -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)

View file

@ -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."""

View file

@ -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

View file

@ -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():

View file

@ -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()

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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."""

View file

@ -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

View file

@ -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)

View file

@ -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"}

View file

@ -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 }
}

View file

@ -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: "/",

View file

@ -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",

View file

@ -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,
}),
}));

View file

@ -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<ProxySettings | undefined>(undefined);
useEffect(() => {
@ -36,7 +35,7 @@ export default function PlaygroundPage() {
initializeProxySettings();
}, [accessToken]);
if (isViewOnlyRole(userRole)) {
if (isViewOnly) {
return (
<div className="flex h-full w-full flex-col items-center justify-center gap-2 p-8 text-center">
<h1 className="text-2xl font-semibold">Access Denied</h1>

View file

@ -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(<UserDropdown onLogout={mockOnLogout} />);
@ -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,
});

View file

@ -69,7 +69,7 @@ interface UserDropdownProps {
}
const UserDropdown: React.FC<UserDropdownProps> = ({ 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();

View file

@ -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",
});

View file

@ -81,7 +81,7 @@ interface SidebarAccountMenuProps {
}
const SidebarAccountMenu: React.FC<SidebarAccountMenuProps> = ({ 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();

View file

@ -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<typeof getReferencedModelsError>[0],
availableModelSet: Set<string>,
): string | null =>
getMissingTiersError(config.tiers) ??
getTierLabelsError(config.tier_labels) ??
getKeywordTierRulesError(keywordTierRules) ??
getReferencedModelsError(referencedModelsParams, availableModelSet);
const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({
handleOk,
accessToken,
@ -222,15 +237,12 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({
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,

View file

@ -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", () => {

View file

@ -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,

View file

@ -407,7 +407,7 @@ const Sidebar_: React.FC<SidebarProps> = ({
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<SidebarProps> = ({
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;

View file

@ -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<UserDashboardProps> = ({
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<UserDashboardProps> = ({
// check if userRole is defined
if (decoded.user_role) {
const formattedUserRole = formatUserRole(decoded.user_role);
setUserRole(formattedUserRole);
setUserRole(effectiveSessionRole(decoded.user_role));
} else {
}

View file

@ -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);

View file

@ -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);
});
});
});

View file

@ -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() ?? "");