mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge origin/litellm_internal_staging into litellm_health_check_planned_recreate
This commit is contained in:
commit
2f64cc6c03
92 changed files with 2707 additions and 442 deletions
0
enterprise/litellm_enterprise/integrations/__init__.py
Normal file
0
enterprise/litellm_enterprise/integrations/__init__.py
Normal file
0
enterprise/litellm_enterprise/proxy/hooks/__init__.py
Normal file
0
enterprise/litellm_enterprise/proxy/hooks/__init__.py
Normal 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(
|
||||
|
|
|
|||
0
enterprise/litellm_enterprise/py.typed
Normal file
0
enterprise/litellm_enterprise/py.typed
Normal file
0
enterprise/litellm_enterprise/types/__init__.py
Normal file
0
enterprise/litellm_enterprise/types/__init__.py
Normal file
0
enterprise/litellm_enterprise/types/proxy/__init__.py
Normal file
0
enterprise/litellm_enterprise/types/proxy/__init__.py
Normal file
0
litellm-proxy-extras/litellm_proxy_extras/py.typed
Normal file
0
litellm-proxy-extras/litellm_proxy_extras/py.typed
Normal 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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
276
litellm/caching/evicted_client_closer.py
Normal file
276
litellm/caching/evicted_client_closer.py
Normal file
|
|
@ -0,0 +1,276 @@
|
|||
"""
|
||||
Deferred close of HTTP/SDK clients that the LLM client cache has evicted.
|
||||
|
||||
Eviction only drops the cache's reference to a client. Every OpenAI/Azure SDK
|
||||
client is a reference cycle (each resource namespace holds the client back), so
|
||||
an evicted client and its pooled TCP connections survive until a generational
|
||||
collection runs, which under load is thousands of requests later.
|
||||
|
||||
Closing at eviction time is not an option: a request that was handed the client
|
||||
just before it was evicted is still using it, and closing it underneath that
|
||||
request raises ``RuntimeError: Cannot send a request, as the client has been
|
||||
closed.``
|
||||
|
||||
So an evicted client is closed once two conditions hold. A grace window must
|
||||
have passed since its eviction, which covers a request that holds the client
|
||||
but is momentarily not on the wire, and the client must report no connection in
|
||||
flight. The second condition is what keeps the first honest: a request may run
|
||||
for ``litellm.request_timeout`` seconds, 6000 by default, and a streaming
|
||||
response is bounded only by how long the upstream keeps sending, so no deadline
|
||||
on its own can promise that a request has finished.
|
||||
|
||||
Only clients litellm itself created are closed; a client the caller supplied is
|
||||
left alone because litellm does not own its lifecycle.
|
||||
|
||||
A client that closes synchronously is closed from wherever the cache is next
|
||||
used. One whose close is a coroutine needs the event loop it was evicted on, so
|
||||
it waits for a call from that loop rather than having work scheduled onto a loop
|
||||
it does not belong to. Queued clients are therefore bucketed by what it takes to
|
||||
close them, and each bucket is ordered by deadline, so a reap walks the entries
|
||||
that are due rather than the whole queue.
|
||||
|
||||
The queue holds its clients weakly, so waiting out a grace window never keeps
|
||||
alive anything the collector would have reclaimed first.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import inspect
|
||||
import threading
|
||||
import time
|
||||
import weakref
|
||||
from collections import deque
|
||||
from collections.abc import Awaitable, Callable, Iterator
|
||||
from dataclasses import dataclass, replace
|
||||
from typing import Final
|
||||
|
||||
from litellm.constants import (
|
||||
EVICTED_LLM_CLIENT_CLOSE_GRACE_SECONDS,
|
||||
EVICTED_LLM_CLIENT_CLOSE_MAX_PENDING,
|
||||
)
|
||||
|
||||
_CLOSABLE_ANYWHERE: Final = "closable-anywhere"
|
||||
_CLOSABLE_ON_ANY_LOOP: Final = "closable-on-any-loop"
|
||||
|
||||
_BucketKey = str | int
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _PendingClose:
|
||||
"""A queued close.
|
||||
|
||||
The client is held weakly, so queueing one never keeps alive anything the
|
||||
collector would otherwise have reclaimed first.
|
||||
|
||||
``needs_loop`` is set for a client whose close is a coroutine; those can only
|
||||
be closed from the event loop they were evicted on, recorded in ``loop_id``.
|
||||
A client that closes synchronously carries neither constraint.
|
||||
"""
|
||||
|
||||
client_ref: "weakref.ref[object]"
|
||||
loop_id: int | None
|
||||
needs_loop: bool
|
||||
close_after: float
|
||||
|
||||
|
||||
def _bucket_key(pending: _PendingClose) -> _BucketKey:
|
||||
"""Which reaps can close this entry: any at all, any running a loop, or one loop's."""
|
||||
if not pending.needs_loop:
|
||||
return _CLOSABLE_ANYWHERE
|
||||
if pending.loop_id is None:
|
||||
return _CLOSABLE_ON_ANY_LOOP
|
||||
return pending.loop_id
|
||||
|
||||
|
||||
def _running_loop_id() -> int | None:
|
||||
try:
|
||||
return id(asyncio.get_running_loop())
|
||||
except RuntimeError:
|
||||
return None
|
||||
|
||||
|
||||
def _close_function(client: object) -> Callable[[], object] | None:
|
||||
close_fn: Final[Callable[[], object] | None] = getattr(client, "aclose", None) or getattr(client, "close", None)
|
||||
return close_fn
|
||||
|
||||
|
||||
def _transport_of(client: object) -> object:
|
||||
"""The httpx transport behind an SDK wrapper, a litellm handler, or a bare client."""
|
||||
for holder in (getattr(client, "_client", None), getattr(client, "client", None), client):
|
||||
transport: object = getattr(holder, "_transport", None)
|
||||
if transport is not None:
|
||||
return transport
|
||||
return None
|
||||
|
||||
|
||||
def _connection_is_idle(connection: object) -> bool:
|
||||
"""A pooled connection is idle unless it is servicing a request."""
|
||||
is_idle: Final[object] = getattr(connection, "is_idle", None)
|
||||
return bool(is_idle()) if callable(is_idle) else True
|
||||
|
||||
|
||||
def _pool_has_busy_connection(transport: object) -> bool | None:
|
||||
"""Whether the httpcore pool behind the transport is servicing a request.
|
||||
|
||||
``None`` when there is no such pool, so the caller can ask the other backend.
|
||||
"""
|
||||
pooled: Final[object] = getattr(getattr(transport, "_pool", None), "connections", None)
|
||||
if not isinstance(pooled, (list, tuple)):
|
||||
return None
|
||||
return any(
|
||||
not _connection_is_idle(connection) # pyright: ignore[reportUnknownArgumentType] # untyped pool list
|
||||
for connection in pooled # pyright: ignore[reportUnknownVariableType] # untyped pool list
|
||||
)
|
||||
|
||||
|
||||
def _has_connection_in_flight(client: object) -> bool:
|
||||
"""Whether the client is servicing a request right now.
|
||||
|
||||
Both connection backends litellm uses already account for the connections
|
||||
they have handed out, so this reads the client's own lease accounting rather
|
||||
than inferring it from elapsed time: httpcore reports a non-idle connection
|
||||
for the whole of a response including a stream, and aiohttp holds the
|
||||
connection in ``_acquired`` over the same span.
|
||||
|
||||
A client that cannot answer is reported as idle, which leaves the grace
|
||||
window as the only guard, exactly as it was before this check existed.
|
||||
"""
|
||||
try:
|
||||
transport: Final = _transport_of(client)
|
||||
pooled_busy: Final = _pool_has_busy_connection(transport)
|
||||
if pooled_busy is not None:
|
||||
return pooled_busy
|
||||
session: Final[object] = getattr(transport, "client", None)
|
||||
return bool(getattr(getattr(session, "connector", None), "_acquired", None))
|
||||
except Exception: # noqa: BLE001 - a client that cannot report its state is treated as idle
|
||||
return False
|
||||
|
||||
|
||||
async def _close_quietly(closing: Awaitable[object]) -> None:
|
||||
with contextlib.suppress(Exception):
|
||||
await closing
|
||||
|
||||
|
||||
class EvictedClientCloser:
|
||||
"""Closes evicted, litellm-owned clients once they are idle and out of grace."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
grace_seconds: float = EVICTED_LLM_CLIENT_CLOSE_GRACE_SECONDS,
|
||||
max_pending: int = EVICTED_LLM_CLIENT_CLOSE_MAX_PENDING,
|
||||
clock: Callable[[], float] = time.monotonic,
|
||||
) -> None:
|
||||
self._grace_seconds = grace_seconds
|
||||
self._max_pending = max_pending
|
||||
self._clock = clock
|
||||
self._owned: weakref.WeakSet[object] = weakref.WeakSet()
|
||||
self._buckets: dict[_BucketKey, deque[_PendingClose]] = {} # mutable-ok: deadline-ordered queues
|
||||
self._pending_count = 0
|
||||
self._queue_lock = threading.Lock() # the cache is reachable from every worker thread's loop
|
||||
self._close_tasks: set[asyncio.Task[None]] = set() # mutable-ok: strong refs to running closes
|
||||
|
||||
def mark_owned(self, client: object) -> None:
|
||||
"""Record that litellm created this client, so it may be closed on eviction."""
|
||||
try:
|
||||
self._owned.add(client)
|
||||
except TypeError:
|
||||
pass # values that cannot be weak-referenced are never litellm clients
|
||||
|
||||
def _is_owned(self, client: object) -> bool:
|
||||
try:
|
||||
return client in self._owned
|
||||
except TypeError:
|
||||
return False # unhashable values are never litellm clients
|
||||
|
||||
def schedule(self, client: object) -> None:
|
||||
"""Queue an evicted client for closing once it is idle and out of grace.
|
||||
|
||||
Past ``max_pending`` the client is left to the collector instead, so a
|
||||
workload that churns the cache cannot grow this queue without bound.
|
||||
Every queued entry comes due within one grace window, so the capacity it
|
||||
occupies is returned within that window rather than held.
|
||||
"""
|
||||
if client is None or not self._is_owned(client):
|
||||
return
|
||||
close_fn: Final = _close_function(client)
|
||||
if close_fn is None:
|
||||
return
|
||||
if self._pending_count >= self._max_pending:
|
||||
return
|
||||
self._enqueue(
|
||||
_PendingClose(
|
||||
client_ref=weakref.ref(client),
|
||||
loop_id=_running_loop_id(),
|
||||
needs_loop=inspect.iscoroutinefunction(close_fn),
|
||||
close_after=self._clock() + self._grace_seconds,
|
||||
)
|
||||
)
|
||||
|
||||
def reap(self) -> None:
|
||||
"""Close every queued client that is due, idle, and closable from here.
|
||||
|
||||
Called from the cache's read path, so the empty-queue exit comes first and
|
||||
the work done past it is proportional to what is due, not to the queue.
|
||||
"""
|
||||
if not self._pending_count:
|
||||
return
|
||||
now: Final = self._clock()
|
||||
for pending in self._take_due(_running_loop_id(), now):
|
||||
client = pending.client_ref()
|
||||
if client is None:
|
||||
continue
|
||||
if _has_connection_in_flight(client):
|
||||
self._enqueue(replace(pending, close_after=now + self._grace_seconds))
|
||||
continue
|
||||
self._close(client)
|
||||
|
||||
@property
|
||||
def pending_count(self) -> int:
|
||||
return self._pending_count
|
||||
|
||||
def _enqueue(self, pending: _PendingClose) -> None:
|
||||
"""Append to the entry's bucket, dropping any dead entries it queues behind.
|
||||
|
||||
Deadlines only ever move forward, so appending keeps each bucket ordered
|
||||
by deadline, and entries whose client the collector already took sit at
|
||||
the front rather than having to be searched for.
|
||||
"""
|
||||
with self._queue_lock:
|
||||
bucket: Final = self._buckets.setdefault(_bucket_key(pending), deque()) # mutable-ok: FIFO by design
|
||||
while bucket and bucket[0].client_ref() is None:
|
||||
bucket.popleft()
|
||||
self._pending_count -= 1
|
||||
bucket.append(pending)
|
||||
self._pending_count += 1
|
||||
|
||||
def _take_due(self, loop_id: int | None, now: float) -> tuple[_PendingClose, ...]:
|
||||
buckets = (_CLOSABLE_ANYWHERE,) if loop_id is None else (_CLOSABLE_ANYWHERE, _CLOSABLE_ON_ANY_LOOP, loop_id)
|
||||
with self._queue_lock:
|
||||
return tuple(pending for key in buckets for pending in self._drain_locked(key, now))
|
||||
|
||||
def _drain_locked(self, key: _BucketKey, now: float) -> Iterator[_PendingClose]:
|
||||
bucket: Final = self._buckets.get(key)
|
||||
if bucket is None:
|
||||
return
|
||||
while bucket and bucket[0].close_after <= now:
|
||||
self._pending_count -= 1
|
||||
yield bucket.popleft()
|
||||
if not bucket:
|
||||
del self._buckets[key]
|
||||
|
||||
def _close(self, client: object) -> None:
|
||||
close_fn: Final = _close_function(client)
|
||||
if close_fn is None:
|
||||
return
|
||||
try:
|
||||
closing: Final = close_fn()
|
||||
except Exception: # noqa: BLE001 - a discarded client's close must never surface to callers
|
||||
return
|
||||
if not inspect.isawaitable(closing):
|
||||
return
|
||||
task: Final = asyncio.get_running_loop().create_task(_close_quietly(closing))
|
||||
self._close_tasks.add(task)
|
||||
task.add_done_callback(self._close_tasks.discard)
|
||||
|
||||
|
||||
default_evicted_client_closer: Final = EvictedClientCloser()
|
||||
|
|
@ -5,21 +5,44 @@ Add the event loop to the cache key, to prevent event loop closed errors.
|
|||
import asyncio
|
||||
from typing import Final
|
||||
|
||||
from .evicted_client_closer import EvictedClientCloser, default_evicted_client_closer
|
||||
from .in_memory_cache import InMemoryCache
|
||||
|
||||
|
||||
class LLMClientCache(InMemoryCache):
|
||||
"""Cache for LLM HTTP clients (OpenAI, Azure, httpx, etc.).
|
||||
|
||||
IMPORTANT: This cache intentionally does NOT close clients on eviction.
|
||||
Evicted clients may still be in use by in-flight requests. Closing them
|
||||
eagerly causes ``RuntimeError: Cannot send a request, as the client has
|
||||
been closed.`` errors in production after the TTL (1 hour) expires.
|
||||
An evicted client is never closed on the spot: a request handed the client
|
||||
just before eviction is still using it, and closing it there raises
|
||||
``RuntimeError: Cannot send a request, as the client has been closed.``
|
||||
|
||||
Clients that are no longer referenced will be garbage-collected normally.
|
||||
For explicit shutdown cleanup, use ``close_litellm_async_clients()``.
|
||||
Nor can eviction be left to rely on garbage collection. The SDK clients are
|
||||
reference cycles, so an evicted client and its open TCP connections survive
|
||||
until a generational collection runs. Instead a client litellm created is
|
||||
handed to ``EvictedClientCloser``, which closes it once a grace window has
|
||||
passed. Clients the caller supplied are left untouched.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
max_size_in_memory: int | None = 200,
|
||||
default_ttl: int | None = 600,
|
||||
max_size_per_item: int | None = 1024,
|
||||
evicted_client_closer: EvictedClientCloser | None = None,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
max_size_in_memory=max_size_in_memory,
|
||||
default_ttl=default_ttl,
|
||||
max_size_per_item=max_size_per_item,
|
||||
)
|
||||
self.evicted_client_closer = evicted_client_closer or default_evicted_client_closer
|
||||
|
||||
def _remove_key(self, key: str) -> None:
|
||||
evicted: Final[object] = self.cache_dict.get(key)
|
||||
super()._remove_key(key)
|
||||
self.evicted_client_closer.schedule(evicted)
|
||||
self.evicted_client_closer.reap()
|
||||
|
||||
def update_cache_key_with_event_loop(self, key):
|
||||
"""
|
||||
Add the event loop to the cache key, to prevent event loop closed errors.
|
||||
|
|
@ -32,16 +55,22 @@ class LLMClientCache(InMemoryCache):
|
|||
except RuntimeError: # handle no current running event loop
|
||||
return key
|
||||
|
||||
def set_cache(self, key, value, **kwargs):
|
||||
def set_cache(self, key: str, value: object, litellm_owned_client: bool = False, **kwargs):
|
||||
"""``litellm_owned_client`` marks a client litellm built, so it may be closed once evicted."""
|
||||
if litellm_owned_client:
|
||||
self.evicted_client_closer.mark_owned(value)
|
||||
key = self.update_cache_key_with_event_loop(key)
|
||||
return super().set_cache(key, value, **kwargs)
|
||||
|
||||
async def async_set_cache(self, key, value, **kwargs):
|
||||
async def async_set_cache(self, key: str, value: object, litellm_owned_client: bool = False, **kwargs):
|
||||
if litellm_owned_client:
|
||||
self.evicted_client_closer.mark_owned(value)
|
||||
key = self.update_cache_key_with_event_loop(key)
|
||||
return await super().async_set_cache(key, value, **kwargs)
|
||||
|
||||
def get_cache(self, key, **kwargs):
|
||||
key = self.update_cache_key_with_event_loop(key)
|
||||
self.evicted_client_closer.reap()
|
||||
|
||||
return super().get_cache(key, **kwargs)
|
||||
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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"):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -1441,6 +1441,7 @@ def get_async_httpx_client(
|
|||
key=_cache_key_name,
|
||||
value=_new_client,
|
||||
ttl=_DEFAULT_TTL_FOR_HTTPX_CLIENTS,
|
||||
litellm_owned_client=True,
|
||||
)
|
||||
return _new_client
|
||||
|
||||
|
|
@ -1486,5 +1487,6 @@ def _get_httpx_client(params: dict | None = None) -> HTTPHandler:
|
|||
key=_cache_key_name,
|
||||
value=_new_client,
|
||||
ttl=_DEFAULT_TTL_FOR_HTTPX_CLIENTS,
|
||||
litellm_owned_client=True,
|
||||
)
|
||||
return _new_client
|
||||
|
|
|
|||
|
|
@ -128,13 +128,33 @@ class BaseOpenAILLM:
|
|||
_cached_client: Final = litellm.in_memory_llm_clients_cache.get_cache(_cache_key)
|
||||
return _cached_client
|
||||
|
||||
@staticmethod
|
||||
def owns_wrapped_http_client(http_client: httpx.Client | httpx.AsyncClient | None) -> bool:
|
||||
"""Whether litellm may close an SDK client built around ``http_client``.
|
||||
|
||||
``_get_async_http_client`` / ``_get_sync_http_client`` hand back
|
||||
``litellm.aclient_session`` / ``litellm.client_session`` when the caller
|
||||
configured one. The SDK's ``close()`` closes whatever http client it was
|
||||
given, so an SDK client wrapping one of those shared sessions must never be
|
||||
closed on eviction; the caller goes on using the session. ``None`` means the
|
||||
SDK built its own http client, which litellm does own.
|
||||
"""
|
||||
if http_client is None:
|
||||
return True
|
||||
return http_client is not litellm.aclient_session and http_client is not litellm.client_session
|
||||
|
||||
@staticmethod
|
||||
def set_cached_openai_client(
|
||||
openai_client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI,
|
||||
client_type: Literal["openai", "azure"],
|
||||
client_initialization_params: dict,
|
||||
litellm_owned_client: bool = False,
|
||||
):
|
||||
"""Stores the OpenAI client in the in-memory cache for _DEFAULT_TTL_FOR_HTTPX_CLIENTS SECONDS"""
|
||||
"""Stores the OpenAI client in the in-memory cache for _DEFAULT_TTL_FOR_HTTPX_CLIENTS SECONDS
|
||||
|
||||
``litellm_owned_client`` says litellm built this client, so the cache may close it once it
|
||||
is evicted. A client the caller supplied stays open, since litellm does not own it.
|
||||
"""
|
||||
_cache_key: Final = BaseOpenAILLM.get_openai_client_cache_key(
|
||||
client_initialization_params=client_initialization_params,
|
||||
client_type=client_type,
|
||||
|
|
@ -143,6 +163,7 @@ class BaseOpenAILLM:
|
|||
key=_cache_key,
|
||||
value=openai_client,
|
||||
ttl=_DEFAULT_TTL_FOR_HTTPX_CLIENTS,
|
||||
litellm_owned_client=litellm_owned_client,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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}",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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})
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -714,11 +714,10 @@ def _enforce_user_info_access(user_id: str | None, user_api_key_dict: UserAPIKey
|
|||
"""
|
||||
if user_id is None:
|
||||
return
|
||||
# Only true proxy admin bypasses ownership. PROXY_ADMIN_VIEW_ONLY is
|
||||
# subject to the same `user_id == valid_token.user_id` rule that
|
||||
# `RouteChecks.non_proxy_admin_allowed_routes_check` applies upstream
|
||||
# for the `/user/info` route.
|
||||
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
|
||||
# Admin-view roles (PROXY_ADMIN and PROXY_ADMIN_VIEW_ONLY) bypass
|
||||
# ownership, mirroring the `/user/info` carve-out that
|
||||
# `RouteChecks.non_proxy_admin_allowed_routes_check` applies upstream.
|
||||
if _user_has_admin_view(user_api_key_dict):
|
||||
return
|
||||
if user_id == user_api_key_dict.user_id:
|
||||
return
|
||||
|
|
@ -862,7 +861,7 @@ async def user_info(
|
|||
raise Exception(
|
||||
"Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys"
|
||||
)
|
||||
if user_id is None and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
|
||||
if user_id is None and _user_has_admin_view(user_api_key_dict):
|
||||
return await _get_user_info_for_proxy_admin(user_api_key_dict=user_api_key_dict)
|
||||
elif user_id is None:
|
||||
user_id = user_api_key_dict.user_id
|
||||
|
|
|
|||
|
|
@ -78,6 +78,7 @@ from litellm.proxy.management_endpoints.common_utils import (
|
|||
_is_user_team_admin,
|
||||
_set_object_metadata_field,
|
||||
_team_member_has_permission,
|
||||
_user_has_admin_view,
|
||||
validate_finite_spend,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -25,7 +25,12 @@ except ImportError:
|
|||
from pydantic import BaseModel
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy._types import (
|
||||
CommonProxyErrors,
|
||||
LitellmUserRoles,
|
||||
UserAPIKeyAuth,
|
||||
user_api_key_has_admin_view,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.repositories.table_repositories import (
|
||||
WorkflowEventRepository,
|
||||
|
|
@ -47,6 +52,10 @@ def _is_admin(user_api_key_dict: UserAPIKeyAuth) -> bool:
|
|||
return user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
|
||||
|
||||
|
||||
def _read_scope_caller(user_api_key_dict: UserAPIKeyAuth) -> UserAPIKeyAuth | None:
|
||||
return None if user_api_key_has_admin_view(user_api_key_dict) else user_api_key_dict
|
||||
|
||||
|
||||
def _caller_key(user_api_key_dict: UserAPIKeyAuth) -> str | None:
|
||||
"""Return the hashed key token that identifies this caller, or None for master key."""
|
||||
return user_api_key_dict.token
|
||||
|
|
@ -199,7 +208,7 @@ async def list_workflow_runs(
|
|||
where["status"] = {"in": statuses} if len(statuses) > 1 else statuses[0]
|
||||
|
||||
# Non-admin callers are scoped to their own key.
|
||||
if not _is_admin(user_api_key_dict):
|
||||
if not user_api_key_has_admin_view(user_api_key_dict):
|
||||
caller: Final = _caller_key(user_api_key_dict)
|
||||
if caller:
|
||||
where["created_by"] = caller
|
||||
|
|
@ -238,7 +247,7 @@ async def get_workflow_run(
|
|||
)
|
||||
if run is None:
|
||||
raise HTTPException(status_code=404, detail=f"Run '{run_id}' not found")
|
||||
if not _is_admin(user_api_key_dict):
|
||||
if not user_api_key_has_admin_view(user_api_key_dict):
|
||||
caller: Final = _caller_key(user_api_key_dict)
|
||||
if not caller or run.created_by != caller:
|
||||
raise HTTPException(status_code=404, detail=f"Run '{run_id}' not found")
|
||||
|
|
@ -377,7 +386,7 @@ async def list_workflow_events(
|
|||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
|
||||
|
||||
await _require_run(prisma_client, run_id, user_api_key_dict)
|
||||
await _require_run(prisma_client, run_id, _read_scope_caller(user_api_key_dict))
|
||||
|
||||
try:
|
||||
events: Final = await WorkflowEventRepository(prisma_client).table.find_many(
|
||||
|
|
@ -461,7 +470,7 @@ async def list_workflow_messages(
|
|||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
|
||||
|
||||
await _require_run(prisma_client, run_id, user_api_key_dict)
|
||||
await _require_run(prisma_client, run_id, _read_scope_caller(user_api_key_dict))
|
||||
|
||||
try:
|
||||
messages: Final = await WorkflowMessageRepository(prisma_client).table.find_many(
|
||||
|
|
|
|||
|
|
@ -27,6 +27,7 @@ from litellm.proxy._types import (
|
|||
CommonProxyErrors,
|
||||
LitellmUserRoles,
|
||||
UserAPIKeyAuth,
|
||||
user_api_key_has_admin_view,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.repositories.table_repositories import MemoryRepository
|
||||
|
|
@ -66,7 +67,7 @@ def _visibility_filter(user_api_key_dict: UserAPIKeyAuth) -> dict | None:
|
|||
Prisma `where` fragment restricting rows to those the caller can see.
|
||||
Returns None for admins (no restriction).
|
||||
"""
|
||||
if _is_admin(user_api_key_dict):
|
||||
if user_api_key_has_admin_view(user_api_key_dict):
|
||||
return None
|
||||
ors: Final[list[dict]] = []
|
||||
if user_api_key_dict.user_id:
|
||||
|
|
|
|||
|
|
@ -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}",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
411
tests/test_litellm/caching/test_evicted_client_closer.py
Normal file
411
tests/test_litellm/caching/test_evicted_client_closer.py
Normal 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"
|
||||
)
|
||||
|
|
@ -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.).
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
103
tests/test_litellm/test_env_key_doc_gate.py
Normal file
103
tests/test_litellm/test_env_key_doc_gate.py
Normal 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"}
|
||||
|
|
@ -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 }
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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: "/",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}),
|
||||
}));
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
});
|
||||
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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", () => {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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() ?? "");
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue