mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_daily_any_cleanup_07_31_2026b
This commit is contained in:
commit
375d7199bf
9 changed files with 2278 additions and 131 deletions
|
|
@ -1,10 +1,11 @@
|
|||
import asyncio
|
||||
import concurrent.futures
|
||||
import contextlib
|
||||
import os
|
||||
import ssl
|
||||
import typing
|
||||
import urllib.request
|
||||
from typing import Any, Callable, Dict, Optional, Union
|
||||
from typing import Any, Callable, ClassVar, Dict, Optional, Union
|
||||
|
||||
import aiohttp
|
||||
import aiohttp.client_exceptions
|
||||
|
|
@ -138,6 +139,11 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
Credit to: https://github.com/karpetrosyan/httpx-aiohttp for this implementation
|
||||
"""
|
||||
|
||||
# Strong references to scheduled session-close tasks. A bare
|
||||
# asyncio.create_task() result may be garbage-collected before it runs,
|
||||
# leaving the recycled session unclosed ("Unclosed client session").
|
||||
_background_close_tasks: ClassVar[set["asyncio.Task[None]"]] = set() # mutable-ok: strong refs for pending closes
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
client: Union[ClientSession, Callable[[], ClientSession]],
|
||||
|
|
@ -164,6 +170,92 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
self._owns_session = True
|
||||
return session
|
||||
|
||||
@classmethod
|
||||
def _on_close_task_done(cls, task: "asyncio.Task[None]") -> None:
|
||||
cls._background_close_tasks.discard(task)
|
||||
if task.cancelled():
|
||||
return
|
||||
exc = task.exception()
|
||||
if exc is not None:
|
||||
verbose_logger.debug("Error closing recycled aiohttp session: %s", exc)
|
||||
|
||||
@staticmethod
|
||||
def _on_threadsafe_close_done(future: "concurrent.futures.Future[None]") -> None:
|
||||
if future.cancelled():
|
||||
return
|
||||
exc = future.exception()
|
||||
if exc is not None:
|
||||
verbose_logger.debug("Error closing recycled aiohttp session on its own loop: %s", exc)
|
||||
|
||||
@staticmethod
|
||||
def _mark_connector_closed(session: ClientSession) -> None:
|
||||
"""Synchronously dispose a session whose event loop is gone.
|
||||
|
||||
An async close can no longer run on a closed loop. BaseConnector._close
|
||||
is the same synchronous teardown aiohttp's own finalizer (__del__)
|
||||
uses: it is guarded for closed loops, releases pooled connections, and
|
||||
flips the flags that ClientSession.closed / BaseConnector.closed read -
|
||||
so no "Unclosed client session" / "Unclosed connector" warnings reach
|
||||
the event-loop exception handler at garbage collection.
|
||||
"""
|
||||
connector = getattr(session, "_connector", None)
|
||||
close_sync = getattr(connector, "_close", None)
|
||||
if not callable(close_sync):
|
||||
return
|
||||
try:
|
||||
close_sync()
|
||||
except (RuntimeError, AttributeError, OSError) as e:
|
||||
verbose_logger.debug("Best-effort connector close failed: %s", e)
|
||||
|
||||
def _close_recycled_session(self, session: ClientSession) -> None:
|
||||
"""Deterministically dispose a ClientSession this transport is replacing.
|
||||
|
||||
Covers the three lifecycles a recycled session can be in:
|
||||
- its loop is the current running loop: schedule an async close and keep
|
||||
a strong reference to the task until it completes;
|
||||
- its loop is still running elsewhere (e.g. another thread): hand the
|
||||
close to that loop thread-safely;
|
||||
- its loop is stopped or closed, or there is no running loop: fall
|
||||
back to the synchronous finalizer-safe teardown.
|
||||
"""
|
||||
if session.closed:
|
||||
return
|
||||
|
||||
session_loop = getattr(session, "_loop", None)
|
||||
try:
|
||||
current_loop: Optional[asyncio.AbstractEventLoop] = asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
current_loop = None
|
||||
|
||||
if session_loop is not None and session_loop is not current_loop:
|
||||
if not session_loop.is_closed() and session_loop.is_running():
|
||||
# The session's loop is running somewhere else (e.g. another
|
||||
# thread): closing from here would touch that loop's internals
|
||||
# unsafely; hand the close to its own loop.
|
||||
try:
|
||||
future = asyncio.run_coroutine_threadsafe(session.close(), session_loop)
|
||||
except RuntimeError as e: # loop shut down between the checks
|
||||
verbose_logger.debug("Threadsafe session close failed: %s", e)
|
||||
self._mark_connector_closed(session)
|
||||
else:
|
||||
future.add_done_callback(self._on_threadsafe_close_done)
|
||||
return
|
||||
|
||||
# Foreign loop that is stopped or closed: an async close can no
|
||||
# longer run there, and running it on the current loop would touch
|
||||
# another loop's internals. Dispose synchronously instead.
|
||||
self._mark_connector_closed(session)
|
||||
return
|
||||
|
||||
if current_loop is None:
|
||||
self._mark_connector_closed(session)
|
||||
return
|
||||
|
||||
task = current_loop.create_task(session.close())
|
||||
cls = type(self)
|
||||
cls._background_close_tasks.add(task)
|
||||
task.add_done_callback(cls._on_close_task_done)
|
||||
|
||||
def _get_valid_client_session(self) -> ClientSession:
|
||||
"""
|
||||
Helper to get a valid ClientSession for the current event loop.
|
||||
|
|
@ -193,21 +285,25 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
# Close old session to prevent leaks
|
||||
old_session = self.client
|
||||
try:
|
||||
if self._owns_session and not old_session.closed:
|
||||
try:
|
||||
asyncio.create_task(old_session.close())
|
||||
except RuntimeError:
|
||||
# Different event loop - can't schedule task, rely on GC
|
||||
verbose_logger.debug("Old session from different loop, relying on GC")
|
||||
if self._owns_session:
|
||||
self._close_recycled_session(old_session)
|
||||
except Exception as e:
|
||||
verbose_logger.debug(f"Error closing old session: {e}")
|
||||
|
||||
# Create a new session in the current event loop
|
||||
self.client = self._rebuild_session()
|
||||
|
||||
except (RuntimeError, AttributeError):
|
||||
# If we can't check the loop or session is invalid, recreate it
|
||||
except (RuntimeError, AttributeError) as e:
|
||||
# If we can't check the loop or session is invalid, recreate it,
|
||||
# but still dispose of the session being replaced.
|
||||
old_session = self.client
|
||||
if self._owns_session:
|
||||
try:
|
||||
self._close_recycled_session(old_session)
|
||||
except (RuntimeError, AttributeError, OSError) as close_error:
|
||||
verbose_logger.debug(f"Error closing old session: {close_error}")
|
||||
self.client = self._rebuild_session()
|
||||
verbose_logger.debug(f"Error checking session loop, created new session: {e}")
|
||||
|
||||
return self.client
|
||||
|
||||
|
|
@ -301,7 +397,14 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
# Handle the case where session was closed between our check and actual use
|
||||
if "Session is closed" in str(e):
|
||||
verbose_logger.debug(f"Session closed during request, retrying with new session: {e}")
|
||||
# Force creation of a new session
|
||||
# Dispose of the session that actually faulted. Do NOT read
|
||||
# self.client here: a concurrent task may already have
|
||||
# replaced it with a healthy session that must stay open.
|
||||
# Guarded by isinstance: factory-injected sessions may be
|
||||
# duck-typed test doubles without a close() coroutine.
|
||||
# Read _owns_session before _rebuild_session() claims ownership.
|
||||
if self._owns_session and isinstance(client_session, ClientSession):
|
||||
self._close_recycled_session(client_session)
|
||||
self.client = self._rebuild_session()
|
||||
client_session = self.client
|
||||
|
||||
|
|
|
|||
|
|
@ -67,6 +67,18 @@ def _as_list(values: Sequence[str] | None) -> list[str] | None: # mutable-ok: r
|
|||
return None if values is None else list(values)
|
||||
|
||||
|
||||
class UnloadableEntitlementError(Exception):
|
||||
"""A principal's row NAMES an ``object_permission_id`` whose contents could not be read.
|
||||
|
||||
Raised only where there is POSITIVE evidence an entitlement exists, so every caller must DENY
|
||||
rather than fall back to "this level places no restriction": a ceiling we know exists but cannot
|
||||
read would otherwise silently widen the caller for as long as the fault lasts.
|
||||
|
||||
Deliberately distinct from a lookup that fails before the principal's entitlement is known at
|
||||
all. Not knowing whether someone is entitled is the state that existed before the level did, so
|
||||
it places no ceiling; denying there would refuse MCP to every caller during a cold-cache fault."""
|
||||
|
||||
|
||||
def _parse_mcp_server_names_from_path(path: str, mcp_servers_header: Optional[List[str]] = None) -> Optional[List[str]]:
|
||||
"""Resolve the single MCP server name a cold-start passthrough bypass may
|
||||
target. Delegates parsing to
|
||||
|
|
@ -292,6 +304,22 @@ class MCPRequestHandler:
|
|||
3. Header extraction and validation
|
||||
|
||||
Utilizes the main `user_api_key_auth` function to validate authentication
|
||||
|
||||
Entitlement-fault contract (``get_allowed_mcp_servers`` / ``get_allowed_tools_for_server``)
|
||||
------------------------------------------------------------------------------------------
|
||||
Every level (key, team, end user, agent, org) answers "which servers/tools does this level
|
||||
permit", and a level that answers nothing places no restriction. A lookup FAULT is not that
|
||||
answer, and the two callers resolve it differently on purpose:
|
||||
|
||||
- A keyless gateway-admitted subject fails CLOSED on any fault at any level. Each of its grant
|
||||
sources is resolved independently and unioned, so a fault that returned "no restriction" would
|
||||
win the union as allow-all, and its per-source org ceiling is the ONLY org bound it has.
|
||||
- Key auth fails closed only where there is POSITIVE evidence an entitlement exists: a principal
|
||||
row that NAMES an ``object_permission_id`` we cannot load is a known entitlement with unknown
|
||||
contents (``UnloadableEntitlementError`` -> deny). A fault so early we cannot tell whether the
|
||||
principal is entitled at all leaves no ceiling, because that is the state that existed before
|
||||
the level did; denying there would refuse MCP to every caller, most of whom have no entitlement
|
||||
configured, for the duration of a cold-cache or DB fault.
|
||||
"""
|
||||
|
||||
LITELLM_API_KEY_HEADER_NAME_PRIMARY = SpecialHeaders.custom_litellm_api_key.value
|
||||
|
|
@ -1348,6 +1376,9 @@ class MCPRequestHandler:
|
|||
has an explicit MCP server list, the combined key/team/end_user/agent result is
|
||||
capped to that list. If the org has no list, no extra restriction is applied.
|
||||
|
||||
A level that cannot answer is NOT a level that permits everything; see the class docstring
|
||||
for how each caller shape resolves an entitlement fault.
|
||||
|
||||
Returns:
|
||||
List[str]: List of allowed MCP servers by server id
|
||||
"""
|
||||
|
|
@ -1478,7 +1509,12 @@ class MCPRequestHandler:
|
|||
|
||||
return list(set(allowed_mcp_servers))
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}")
|
||||
if isinstance(e, UnloadableEntitlementError):
|
||||
# A ceiling we KNOW exists and cannot read. Denying is the only answer that does not
|
||||
# widen this caller past what an operator configured, for both caller shapes.
|
||||
verbose_logger.warning(f"Denying MCP access, entitlement unreadable: {str(e)}")
|
||||
else:
|
||||
verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}")
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -1491,11 +1527,15 @@ class MCPRequestHandler:
|
|||
"""Cap the resolved server list by this caller's org ceiling: an explicit org list intersects
|
||||
lower-level restrictions (else becomes the ceiling); no org or an empty list leaves it unchanged.
|
||||
|
||||
``keyless_source`` governs both divergences for a keyless admitted source. An UNRESOLVABLE ceiling
|
||||
fails CLOSED for it (its only org bound is this ceiling, so dropping it on a fault would escalate a
|
||||
cross-org user) while a key stays fail-open. And an org list may only ever INTERSECT a source (the
|
||||
admitted model unions grants, so a ceiling must not become one), whereas for a key it may
|
||||
substitute, that being the key ceiling model."""
|
||||
``keyless_source`` governs both divergences for a keyless admitted source. An INDETERMINATE ceiling
|
||||
(we cannot tell whether the org restricts at all) fails CLOSED for it (its only org bound is this
|
||||
ceiling, so dropping it on a fault would escalate a cross-org user) while a key stays fail-open. And
|
||||
an org list may only ever INTERSECT a source (the admitted model unions grants, so a ceiling must not
|
||||
become one), whereas for a key it may substitute, that being the key ceiling model.
|
||||
|
||||
The fail-open arm is reached only for an INDETERMINATE fault: a ceiling the org NAMES but that
|
||||
cannot be read raises out of ``_get_allowed_mcp_servers_for_org`` and never arrives here as
|
||||
``None``, so key auth cannot silently shed a ceiling an operator did configure."""
|
||||
if not (user_api_key_auth and user_api_key_auth.org_id):
|
||||
return allowed_mcp_servers
|
||||
allowed_mcp_servers_for_org = await MCPRequestHandler._get_allowed_mcp_servers_for_org(user_api_key_auth)
|
||||
|
|
@ -1900,12 +1940,19 @@ class MCPRequestHandler:
|
|||
)
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"Failed to get allowed tools for server: {str(e)}")
|
||||
# An entitlement known to exist but unreadable denies for BOTH caller shapes, so [] rather
|
||||
# than the None (allow-all) key auth gets for an indeterminate fault.
|
||||
unreadable_entitlement = isinstance(e, UnloadableEntitlementError)
|
||||
if unreadable_entitlement:
|
||||
verbose_logger.warning(f"Denying MCP tools, entitlement unreadable: {str(e)}")
|
||||
else:
|
||||
verbose_logger.warning(f"Failed to get allowed tools for server: {str(e)}")
|
||||
# Fail CLOSED for a keyless admitted subject: ANY error must deny the server's tools ([]),
|
||||
# not collapse to allow-all (None); key/JWT auth keeps its prior allow-all-on-error. Both
|
||||
# keyless_source AND the marker are needed: each source resolves through an UNMARKED auth, so
|
||||
# without keyless_source a fault under a source returns None and wins the union as allow-all.
|
||||
return [] if (keyless_source or _is_mcp_admitted_user_subject(user_api_key_auth)) else None
|
||||
deny_all = unreadable_entitlement or keyless_source or _is_mcp_admitted_user_subject(user_api_key_auth)
|
||||
return [] if deny_all else None
|
||||
|
||||
@staticmethod
|
||||
async def _apply_agent_and_org_tool_ceilings(
|
||||
|
|
@ -1944,7 +1991,9 @@ class MCPRequestHandler:
|
|||
try:
|
||||
org_obj_perm = await MCPRequestHandler._get_org_object_permission(user_api_key_auth)
|
||||
except Exception as e: # noqa: BLE001 # unresolvable org ceiling, decided per caller shape
|
||||
if keyless_source:
|
||||
# A ceiling the org NAMES but that cannot be read denies at every caller shape; only an
|
||||
# INDETERMINATE fault (we cannot tell whether a ceiling exists) keeps key auth open.
|
||||
if keyless_source or isinstance(e, UnloadableEntitlementError):
|
||||
raise
|
||||
verbose_logger.warning(
|
||||
f"MCP org tool ceiling unresolvable for org_id={user_api_key_auth.org_id!r}; "
|
||||
|
|
@ -2275,18 +2324,54 @@ class MCPRequestHandler:
|
|||
verbose_logger.warning(f"Failed to get allowed MCP servers for team: {str(e)}")
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
async def _load_named_object_permission(
|
||||
principal: str,
|
||||
object_permission_id: str,
|
||||
prisma_client: "PrismaClient",
|
||||
user_api_key_auth: UserAPIKeyAuth,
|
||||
) -> LiteLLM_ObjectPermissionTable:
|
||||
"""Load the object permission a principal's row NAMES, or raise ``UnloadableEntitlementError``.
|
||||
|
||||
The single place that fault is minted, so end user, agent and org cannot drift on what counts
|
||||
as "known entitlement, unknown contents". ``get_object_permission`` answers None for both an
|
||||
absent row and a failed read, and neither is evidence the principal is unrestricted: the link
|
||||
proves an entitlement was configured, so both must deny."""
|
||||
from litellm.proxy.auth.auth_checks import get_object_permission
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache
|
||||
|
||||
unloadable = UnloadableEntitlementError(
|
||||
f"{principal} names object_permission_id {object_permission_id!r} which could not be loaded"
|
||||
)
|
||||
try:
|
||||
object_permission = await get_object_permission(
|
||||
object_permission_id=object_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # a named entitlement we cannot read denies, whatever the read failed with
|
||||
raise unloadable from e
|
||||
if object_permission is None:
|
||||
raise unloadable
|
||||
return object_permission
|
||||
|
||||
@staticmethod
|
||||
async def _get_org_object_permission(
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
):
|
||||
) -> LiteLLM_ObjectPermissionTable | None:
|
||||
"""
|
||||
Get org object_permission via the established ``get_org_object`` /
|
||||
``get_object_permission`` helpers so MCP requests share the same
|
||||
``user_api_key_cache`` entries as the rest of the proxy.
|
||||
|
||||
``None`` means the org places NO ceiling: no ``org_id``, no DB, or an org row naming no
|
||||
permission. A row that NAMES one it cannot load raises ``UnloadableEntitlementError``;
|
||||
every other lookup failure propagates as itself, leaving the ceiling merely unresolved.
|
||||
"""
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
OrganizationNotFoundError,
|
||||
get_object_permission,
|
||||
get_org_object,
|
||||
)
|
||||
from litellm.proxy.proxy_server import (
|
||||
|
|
@ -2322,31 +2407,29 @@ class MCPRequestHandler:
|
|||
if org_obj is None or not org_obj.object_permission_id:
|
||||
return None
|
||||
|
||||
# The org NAMES a permission; failing to read it is INDETERMINATE and must not collapse into the
|
||||
# None that means "no ceiling". Raise and let each caller pick fail-open or fail-closed.
|
||||
object_permission = await get_object_permission(
|
||||
# The org NAMES a permission; failing to read it is a KNOWN ceiling with unknown contents and
|
||||
# must not collapse into the None that means "no ceiling". Raising denies at every caller shape.
|
||||
return await MCPRequestHandler._load_named_object_permission(
|
||||
principal=f"org {user_api_key_auth.org_id!r}",
|
||||
object_permission_id=org_obj.object_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
if object_permission is None:
|
||||
raise ValueError(
|
||||
f"org {user_api_key_auth.org_id!r} names object_permission_id "
|
||||
f"{org_obj.object_permission_id!r} which could not be loaded"
|
||||
)
|
||||
return object_permission
|
||||
|
||||
@staticmethod
|
||||
async def _get_allowed_mcp_servers_for_org(
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
) -> List[str]:
|
||||
) -> list[str] | None:
|
||||
"""
|
||||
Get allowed MCP servers for an organization.
|
||||
|
||||
Returns the MCP servers from the org's object_permission.
|
||||
An empty result means the org places no restriction (allow-all from this level).
|
||||
An empty result means the org places no restriction (allow-all from this level), ``None``
|
||||
that the ceiling could not be resolved, which the caller decides per shape.
|
||||
|
||||
A ceiling the org NAMES but we cannot read is neither: it raises out of here so both caller
|
||||
shapes deny, because dropping a ceiling known to exist is exactly the silent widening the
|
||||
level is there to prevent.
|
||||
"""
|
||||
try:
|
||||
object_permissions = await MCPRequestHandler._get_org_object_permission(user_api_key_auth)
|
||||
|
|
@ -2374,34 +2457,28 @@ class MCPRequestHandler:
|
|||
except Exception as e:
|
||||
# None = ceiling UNRESOLVED, distinct from [] = org places no restriction. Collapsing them
|
||||
# let a DB fault silently drop a ceiling; the caller picks fail-open/closed from this signal.
|
||||
# A NAMED-but-unreadable ceiling is a stronger fact than "unresolved" and denies everywhere.
|
||||
if isinstance(e, UnloadableEntitlementError):
|
||||
raise
|
||||
verbose_logger.warning(f"Failed to get allowed MCP servers for org: {str(e)}")
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
async def _get_allowed_mcp_servers_for_end_user(
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
) -> List[str]:
|
||||
"""
|
||||
Get allowed MCP servers for an end user.
|
||||
async def _get_end_user_object_permission(
|
||||
user_api_key_auth: UserAPIKeyAuth,
|
||||
prisma_client: "PrismaClient",
|
||||
) -> LiteLLM_ObjectPermissionTable | None:
|
||||
"""The end user's own object_permission, or ``None`` when this level places no restriction.
|
||||
|
||||
Returns the MCP servers from the end_user's object_permission.
|
||||
"""
|
||||
``None`` covers an end user row that is absent or names no permission, and an end user we
|
||||
could not resolve at all (``get_end_user_object`` answers None for an absent row AND for a
|
||||
failed read, so this level genuinely cannot tell those apart). A row that DOES name a
|
||||
permission we cannot load raises ``UnloadableEntitlementError``: the link is positive
|
||||
evidence of an entitlement, so its contents may not be assumed empty."""
|
||||
from litellm.proxy.auth.auth_checks import get_end_user_object
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if not user_api_key_auth or not user_api_key_auth.end_user_id:
|
||||
return []
|
||||
|
||||
if prisma_client is None:
|
||||
verbose_logger.debug("prisma_client is None")
|
||||
return []
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache
|
||||
|
||||
try:
|
||||
# Use optimized get_end_user_object function with caching
|
||||
end_user_obj = await get_end_user_object(
|
||||
end_user_id=user_api_key_auth.end_user_id,
|
||||
prisma_client=prisma_client,
|
||||
|
|
@ -2410,29 +2487,65 @@ class MCPRequestHandler:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
route="/mcp",
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # entitlement unknown, not known-absent: no ceiling, as before this level
|
||||
verbose_logger.warning(f"Failed to resolve end_user for MCP permissions: {str(e)}")
|
||||
return None
|
||||
|
||||
if end_user_obj is None or end_user_obj.object_permission is None:
|
||||
return []
|
||||
if end_user_obj is None:
|
||||
return None
|
||||
if end_user_obj.object_permission is not None:
|
||||
return end_user_obj.object_permission
|
||||
if not end_user_obj.object_permission_id:
|
||||
return None
|
||||
# The row NAMES a permission the relation did not carry. One shared (cached) lookup decides
|
||||
# whether it is readable; an unreadable one denies rather than reading as "no restriction".
|
||||
return await MCPRequestHandler._load_named_object_permission(
|
||||
principal=f"end user {user_api_key_auth.end_user_id!r}",
|
||||
object_permission_id=end_user_obj.object_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _get_allowed_mcp_servers_for_end_user(
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
) -> List[str]:
|
||||
"""
|
||||
Get allowed MCP servers for an end user.
|
||||
|
||||
Returns the MCP servers from the end_user's object_permission; an empty result means this
|
||||
level places no restriction. An entitlement the end user row NAMES but that cannot be read
|
||||
raises ``UnloadableEntitlementError`` out of here so the resolver denies.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if not user_api_key_auth or not user_api_key_auth.end_user_id:
|
||||
return []
|
||||
|
||||
if prisma_client is None:
|
||||
verbose_logger.debug("prisma_client is None")
|
||||
return []
|
||||
|
||||
object_permission = await MCPRequestHandler._get_end_user_object_permission(user_api_key_auth, prisma_client)
|
||||
if object_permission is None:
|
||||
return []
|
||||
|
||||
try:
|
||||
# Permission entries may be server_ids OR names/aliases — expand to ids.
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(
|
||||
end_user_obj.object_permission.mcp_servers or []
|
||||
)
|
||||
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(object_permission.mcp_servers or [])
|
||||
|
||||
# Get MCP servers from access groups
|
||||
access_group_servers = await MCPRequestHandler._get_mcp_servers_from_access_groups(
|
||||
end_user_obj.object_permission.mcp_access_groups or []
|
||||
object_permission.mcp_access_groups or []
|
||||
)
|
||||
|
||||
# servers referenced in tool permissions should also be accessible
|
||||
tool_perm_servers = list(
|
||||
global_mcp_server_manager.expand_tool_permissions(
|
||||
end_user_obj.object_permission.mcp_tool_permissions
|
||||
).keys()
|
||||
global_mcp_server_manager.expand_tool_permissions(object_permission.mcp_tool_permissions).keys()
|
||||
)
|
||||
|
||||
# Combine all lists
|
||||
|
|
@ -2643,22 +2756,51 @@ class MCPRequestHandler:
|
|||
# don't re-query the DB on every MCP request for that agent.
|
||||
_AGENT_NO_PERMISSION_SENTINEL = "__agent_no_mcp_permission__"
|
||||
|
||||
@staticmethod
|
||||
async def _agent_object_permission_id(agent_id: str, prisma_client: "PrismaClient") -> str | None:
|
||||
"""The permission row this agent's row links to, or ``None`` when it links none.
|
||||
|
||||
Caches the link (with a sentinel for "links none") so an agent without an entitlement costs
|
||||
no DB read per MCP request. A read that fails also answers ``None``: not knowing whether the
|
||||
agent is entitled is the state that existed before this level, so it places no ceiling. Only
|
||||
a link we DID resolve can make the caller deny."""
|
||||
from litellm.proxy.proxy_server import user_api_key_cache
|
||||
|
||||
cache_key = f"agent_object_permission_id:{agent_id}"
|
||||
try:
|
||||
cached: object = await user_api_key_cache.async_get_cache(key=cache_key)
|
||||
if cached == MCPRequestHandler._AGENT_NO_PERMISSION_SENTINEL:
|
||||
return None
|
||||
if isinstance(cached, str) and cached:
|
||||
return cached
|
||||
agent_row = await AgentsRepository(prisma_client).table.find_unique(where={"agent_id": agent_id})
|
||||
linked: object = getattr(agent_row, "object_permission_id", None) if agent_row is not None else None
|
||||
object_permission_id = linked if isinstance(linked, str) and linked else None
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=cache_key,
|
||||
value=object_permission_id or MCPRequestHandler._AGENT_NO_PERMISSION_SENTINEL,
|
||||
ttl=get_management_object_ttl(user_api_key_cache),
|
||||
)
|
||||
return object_permission_id
|
||||
except Exception as e: # noqa: BLE001 # entitlement unknown, not known-absent: no ceiling, as before this level
|
||||
verbose_logger.warning(f"Failed to resolve object_permission_id for agent {agent_id!r}: {str(e)}")
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
async def _get_agent_object_permission(
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
):
|
||||
) -> LiteLLM_ObjectPermissionTable | None:
|
||||
"""
|
||||
Get agent object_permission via the established ``get_object_permission``
|
||||
helper. Caches the ``agent_id -> object_permission_id`` mapping so we
|
||||
avoid re-reading the agent row on every request, and reuses the shared
|
||||
``object_permission_id`` cache populated by the org / team / key paths.
|
||||
|
||||
``None`` means the agent places NO restriction: no ``agent_id``, no DB, or an agent linking
|
||||
no permission. An agent that LINKS one we cannot load raises ``UnloadableEntitlementError``,
|
||||
since a known entitlement with unknown contents must deny rather than read as unrestricted.
|
||||
"""
|
||||
from litellm.proxy.auth.auth_checks import get_object_permission
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if not user_api_key_auth or not user_api_key_auth.agent_id:
|
||||
return None
|
||||
|
|
@ -2668,40 +2810,17 @@ class MCPRequestHandler:
|
|||
return None
|
||||
|
||||
agent_id = user_api_key_auth.agent_id
|
||||
cache_key = f"agent_object_permission_id:{agent_id}"
|
||||
|
||||
try:
|
||||
object_permission_id: Optional[str] = await user_api_key_cache.async_get_cache(key=cache_key)
|
||||
|
||||
if object_permission_id == MCPRequestHandler._AGENT_NO_PERMISSION_SENTINEL:
|
||||
return None
|
||||
|
||||
if object_permission_id is None:
|
||||
agent_row = await AgentsRepository(prisma_client).table.find_unique(
|
||||
where={"agent_id": agent_id},
|
||||
)
|
||||
object_permission_id = (
|
||||
getattr(agent_row, "object_permission_id", None) if agent_row is not None else None
|
||||
)
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=cache_key,
|
||||
value=object_permission_id or MCPRequestHandler._AGENT_NO_PERMISSION_SENTINEL,
|
||||
ttl=get_management_object_ttl(user_api_key_cache),
|
||||
)
|
||||
if not object_permission_id:
|
||||
return None
|
||||
|
||||
return await get_object_permission(
|
||||
object_permission_id=object_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"Failed to get agent object permission: {str(e)}")
|
||||
object_permission_id = await MCPRequestHandler._agent_object_permission_id(agent_id, prisma_client)
|
||||
if object_permission_id is None:
|
||||
return None
|
||||
|
||||
return await MCPRequestHandler._load_named_object_permission(
|
||||
principal=f"agent {agent_id!r}",
|
||||
object_permission_id=object_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _get_allowed_mcp_servers_for_agent(
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
|
|
@ -2711,7 +2830,9 @@ class MCPRequestHandler:
|
|||
Get allowed MCP servers for an agent (from the agent's object_permission).
|
||||
|
||||
Returns the MCP servers from the agent's object_permission.
|
||||
If agent has no object_permission, returns [] (no extra restriction).
|
||||
If agent has no object_permission, returns [] (no extra restriction). An entitlement the
|
||||
agent LINKS but that cannot be read raises ``UnloadableEntitlementError`` out of here so the
|
||||
resolver denies.
|
||||
|
||||
Args:
|
||||
user_api_key_auth: User auth with agent_id
|
||||
|
|
@ -2721,13 +2842,13 @@ class MCPRequestHandler:
|
|||
if not user_api_key_auth or not user_api_key_auth.agent_id:
|
||||
return []
|
||||
|
||||
try:
|
||||
obj_perm = agent_object_permission
|
||||
if obj_perm is None:
|
||||
obj_perm = await MCPRequestHandler._get_agent_object_permission(user_api_key_auth)
|
||||
if obj_perm is None:
|
||||
return []
|
||||
obj_perm = agent_object_permission
|
||||
if obj_perm is None:
|
||||
obj_perm = await MCPRequestHandler._get_agent_object_permission(user_api_key_auth)
|
||||
if obj_perm is None:
|
||||
return []
|
||||
|
||||
try:
|
||||
direct_mcp_servers = getattr(obj_perm, "mcp_servers", None) or []
|
||||
if isinstance(direct_mcp_servers, str):
|
||||
direct_mcp_servers = []
|
||||
|
|
@ -2757,7 +2878,9 @@ class MCPRequestHandler:
|
|||
) -> Optional[List[str]]:
|
||||
"""
|
||||
Get allowed tool names for a server from the agent's object_permission.
|
||||
Returns None if agent has no tool restrictions for this server.
|
||||
Returns None if agent has no tool restrictions for this server. An entitlement the agent
|
||||
LINKS but that cannot be read raises ``UnloadableEntitlementError`` out of here, which the
|
||||
tool resolver turns into deny-all for the server rather than an unrestricted tool list.
|
||||
|
||||
Args:
|
||||
server_id: Server ID to check permissions for
|
||||
|
|
@ -2768,13 +2891,13 @@ class MCPRequestHandler:
|
|||
if not user_api_key_auth or not user_api_key_auth.agent_id:
|
||||
return None
|
||||
|
||||
try:
|
||||
obj_perm = agent_object_permission
|
||||
if obj_perm is None:
|
||||
obj_perm = await MCPRequestHandler._get_agent_object_permission(user_api_key_auth)
|
||||
if obj_perm is None:
|
||||
return None
|
||||
obj_perm = agent_object_permission
|
||||
if obj_perm is None:
|
||||
obj_perm = await MCPRequestHandler._get_agent_object_permission(user_api_key_auth)
|
||||
if obj_perm is None:
|
||||
return None
|
||||
|
||||
try:
|
||||
mcp_tool_permissions = getattr(obj_perm, "mcp_tool_permissions", None)
|
||||
if not mcp_tool_permissions or not isinstance(mcp_tool_permissions, dict):
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ from fastapi.dependencies.utils import get_flat_dependant
|
|||
from fastapi.responses import JSONResponse
|
||||
|
||||
from litellm.types.proxy.management_endpoints.management_v1 import (
|
||||
ListLinks,
|
||||
PageLinks,
|
||||
ProblemDetail,
|
||||
)
|
||||
|
|
@ -43,6 +44,21 @@ def _declared_query_params(request: Request) -> frozenset[str]:
|
|||
return frozenset(field.alias for field in get_flat_dependant(dependant, skip_repeats=True).query_params)
|
||||
|
||||
|
||||
def escape_like(value: str) -> str:
|
||||
"""Escape LIKE/ILIKE metacharacters. Ids routinely contain `_`, which is a wildcard unescaped."""
|
||||
return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
|
||||
|
||||
|
||||
def unknown_query_param_problem(unknown: tuple[str, ...], allowed: tuple[str, ...]) -> ProblemDetail:
|
||||
return ProblemDetail(
|
||||
type=f"{PROBLEM_TYPE_BASE}unknown-query-parameter",
|
||||
title="Unknown query parameter",
|
||||
status=400,
|
||||
detail=f"Unrecognized query parameter(s): {', '.join(unknown)}.",
|
||||
allowed=sorted(allowed),
|
||||
)
|
||||
|
||||
|
||||
async def reject_unknown_query_params(request: Request) -> None:
|
||||
"""Reject any query param the route did not declare.
|
||||
|
||||
|
|
@ -53,15 +69,7 @@ async def reject_unknown_query_params(request: Request) -> None:
|
|||
unknown: tuple[str, ...] = tuple(sorted(name for name in request.query_params if name not in declared))
|
||||
if not unknown:
|
||||
return
|
||||
raise ManagementProblem(
|
||||
ProblemDetail(
|
||||
type=f"{PROBLEM_TYPE_BASE}unknown-query-parameter",
|
||||
title="Unknown query parameter",
|
||||
status=400,
|
||||
detail=f"Unrecognized query parameter(s): {', '.join(unknown)}.",
|
||||
allowed=sorted(declared),
|
||||
)
|
||||
)
|
||||
raise ManagementProblem(unknown_query_param_problem(unknown=unknown, allowed=tuple(sorted(declared))))
|
||||
|
||||
|
||||
def _page_url(request: Request, page: int) -> str:
|
||||
|
|
@ -75,3 +83,15 @@ def build_page_links(request: Request, page: int, has_more: bool) -> PageLinks:
|
|||
prev=_page_url(request, page - 1) if page > 1 else None,
|
||||
next=_page_url(request, page + 1) if has_more else None,
|
||||
)
|
||||
|
||||
|
||||
def build_list_links(request: Request, page: int, total_pages: int) -> ListLinks:
|
||||
"""Page-mode links. `last` clamps to page 1 on an empty result set so every link still resolves."""
|
||||
last = max(total_pages, 1)
|
||||
return ListLinks(
|
||||
self_link=_page_url(request, page),
|
||||
first=_page_url(request, 1),
|
||||
prev=_page_url(request, page - 1) if page > 1 else None,
|
||||
next=_page_url(request, page + 1) if page < last else None,
|
||||
last=_page_url(request, last),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,522 @@
|
|||
"""Generic list handling for `/management/v1` collection routes.
|
||||
|
||||
A resource declares a `ListSpec`; `build_query_plan` turns query parameters into a
|
||||
`QueryPlan` or an RFC 9457 problem without touching a database, and `handle_list`
|
||||
runs that plan through an injected `ListExecutor`. Keeping the planning pure is what
|
||||
lets a caller assert the plan as a value instead of asserting against a live Prisma
|
||||
client, and it keeps this module free of any database dependency.
|
||||
|
||||
A plan's `where` is a tuple of frozen `Predicate`s rather than a backend-shaped
|
||||
mapping, so the framework never has to know which query builder executes it and a
|
||||
planned predicate cannot be rewritten afterwards. `where_sql` renders one for a
|
||||
raw-SQL executor with every caller-supplied value bound to a placeholder.
|
||||
"""
|
||||
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from math import ceil
|
||||
from typing import Generic, Literal, Protocol, TypeVar
|
||||
|
||||
from fastapi import Request
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.management_endpoints.management_v1.common import (
|
||||
PROBLEM_TYPE_BASE,
|
||||
ManagementProblem,
|
||||
build_list_links,
|
||||
escape_like,
|
||||
unknown_query_param_problem,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.management_v1 import (
|
||||
ListMeta,
|
||||
ListResponse,
|
||||
ProblemDetail,
|
||||
)
|
||||
|
||||
ComparisonOp = Literal["eq", "gte", "lte", "gt", "lt", "contains", "not"]
|
||||
# `is_null` is not in the design doc's operator set. It is here because there is no
|
||||
# other way to ask for "max_budget IS NULL", and a table that renders nulls as
|
||||
# "Unlimited" has to be able to filter on them.
|
||||
FilterOp = ComparisonOp | Literal["in", "is_null"]
|
||||
|
||||
FilterType = type[str] | type[int] | type[float] | type[datetime]
|
||||
FilterValue = str | int | float | datetime
|
||||
|
||||
PAGE_PARAM = "page"
|
||||
PAGE_SIZE_PARAM = "page_size"
|
||||
SORT_PARAM = "sort"
|
||||
SEARCH_PARAM = "q"
|
||||
|
||||
TRow = TypeVar("TRow")
|
||||
TRow_co = TypeVar("TRow_co", covariant=True)
|
||||
TOut = TypeVar("TOut")
|
||||
|
||||
_FILTER_OP_ADAPTER: TypeAdapter[FilterOp] = TypeAdapter(FilterOp)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Compare:
|
||||
"""`field <op> value`."""
|
||||
|
||||
field: str
|
||||
op: ComparisonOp
|
||||
value: FilterValue
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Within:
|
||||
"""`field IN (values)`."""
|
||||
|
||||
field: str
|
||||
values: tuple[FilterValue, ...]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class IsNull:
|
||||
"""`field IS NULL`, or `IS NOT NULL` when negated."""
|
||||
|
||||
field: str
|
||||
negated: bool
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AnyOf:
|
||||
"""Disjunction of its clauses. `?q=` is the only producer today."""
|
||||
|
||||
clauses: tuple["Predicate", ...]
|
||||
|
||||
|
||||
Predicate = Compare | Within | IsNull | AnyOf
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class FilterSpec:
|
||||
type: FilterType
|
||||
ops: frozenset[FilterOp]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SortKey:
|
||||
field: str
|
||||
descending: bool
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ScopeAll:
|
||||
"""The caller may read every row of the resource."""
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ScopeWhere:
|
||||
"""The caller may read the rows matching every predicate in `where`."""
|
||||
|
||||
where: tuple[Predicate, ...]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ScopeDenied:
|
||||
"""The caller may read no rows at all, and should be told so rather than shown an empty page."""
|
||||
|
||||
reason: str
|
||||
|
||||
|
||||
Scope = ScopeAll | ScopeWhere | ScopeDenied
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ListSpec(Generic[TRow, TOut]):
|
||||
resource: str
|
||||
sortable: frozenset[str]
|
||||
searchable: frozenset[str]
|
||||
filters: Mapping[str, FilterSpec]
|
||||
default_sort: tuple[SortKey, ...]
|
||||
default_page_size: int
|
||||
max_page_size: int
|
||||
scope: Callable[[UserAPIKeyAuth], Scope]
|
||||
serialize: Callable[[TRow], TOut]
|
||||
tiebreaker: str
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""A malformed spec is a programming error at import time, so this raises rather than
|
||||
returning a problem: there is no request in flight and no caller to answer."""
|
||||
if not 1 <= self.default_page_size <= self.max_page_size:
|
||||
raise ValueError(
|
||||
f"{self.resource}: default_page_size must be between 1 and max_page_size "
|
||||
f"({self.max_page_size}), got {self.default_page_size}. A default above the cap "
|
||||
f"would serve more rows than the resource allows whenever page_size is omitted."
|
||||
)
|
||||
if not self.tiebreaker:
|
||||
raise ValueError(f"{self.resource}: tiebreaker is required; it is the final sort key on every query.")
|
||||
undeclared = tuple(sorted(frozenset(key.field for key in self.default_sort) - self.sortable))
|
||||
if undeclared:
|
||||
raise ValueError(f"{self.resource}: default_sort orders by non-sortable field(s): {', '.join(undeclared)}.")
|
||||
non_text = tuple(
|
||||
sorted(field for field, spec in self.filters.items() if "contains" in spec.ops and spec.type is not str)
|
||||
)
|
||||
if non_text:
|
||||
raise ValueError(
|
||||
f"{self.resource}: contains renders as ILIKE and is only meaningful on text columns, "
|
||||
f"but is declared on: {', '.join(non_text)}."
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class QueryPlan:
|
||||
"""`where` is an implicit AND, ordered scope-first; `order` always ends with the spec's tiebreaker."""
|
||||
|
||||
where: tuple[Predicate, ...]
|
||||
order: tuple[SortKey, ...]
|
||||
skip: int
|
||||
take: int
|
||||
|
||||
|
||||
class ListExecutor(Protocol[TRow_co]):
|
||||
"""The database half of a list, injected so this module never imports Prisma."""
|
||||
|
||||
async def count(self, where: tuple[Predicate, ...]) -> int: ...
|
||||
|
||||
async def find_many(self, plan: QueryPlan) -> Sequence[TRow_co]: ...
|
||||
|
||||
|
||||
def order_by_sql(order: tuple[SortKey, ...]) -> str:
|
||||
"""`ORDER BY` body for a plan, NULLS LAST in both directions.
|
||||
|
||||
Postgres sorts nulls last ascending but first descending, so an unqualified flip of
|
||||
the sort direction drags every "Unlimited" row to the top of the table. Every field
|
||||
reaching here is either a member of `ListSpec.sortable` (caller-supplied sort is
|
||||
checked against it, `default_sort` at construction) or the spec's `tiebreaker`, so
|
||||
these are developer-declared column names, never caller-controlled text.
|
||||
"""
|
||||
return ", ".join(f'"{key.field}" {"DESC" if key.descending else "ASC"} NULLS LAST' for key in order)
|
||||
|
||||
|
||||
def _sql_operator(op: ComparisonOp) -> str:
|
||||
match op:
|
||||
case "eq":
|
||||
return "="
|
||||
case "not":
|
||||
return "<>"
|
||||
case "gte":
|
||||
return ">="
|
||||
case "lte":
|
||||
return "<="
|
||||
case "gt":
|
||||
return ">"
|
||||
case "lt":
|
||||
return "<"
|
||||
case "contains":
|
||||
return "ILIKE"
|
||||
case _:
|
||||
assert_never(op)
|
||||
|
||||
|
||||
def _render(predicate: Predicate, index: int) -> tuple[str, tuple[object, ...]]:
|
||||
match predicate:
|
||||
case IsNull(field=field, negated=negated):
|
||||
return f'"{field}" IS {"NOT NULL" if negated else "NULL"}', ()
|
||||
case Within(field=field, values=values):
|
||||
placeholders = ", ".join(f"${index + offset}" for offset in range(len(values)))
|
||||
return f'"{field}" IN ({placeholders})', values
|
||||
case AnyOf(clauses=clauses):
|
||||
rendered, params = _render_all(clauses, index)
|
||||
return f"({' OR '.join(rendered)})", params
|
||||
case Compare(field=field, op="contains", value=value):
|
||||
return f"\"{field}\" ILIKE ${index} ESCAPE '\\'", (f"%{escape_like(str(value))}%",)
|
||||
case Compare(field=field, op=op, value=value):
|
||||
return f'"{field}" {_sql_operator(op)} ${index}', (value,)
|
||||
case _:
|
||||
assert_never(predicate)
|
||||
|
||||
|
||||
def _render_all(predicates: tuple[Predicate, ...], index: int) -> tuple[tuple[str, ...], tuple[object, ...]]:
|
||||
if not predicates:
|
||||
return (), ()
|
||||
head, head_params = _render(predicates[0], index)
|
||||
tail, tail_params = _render_all(predicates[1:], index + len(head_params))
|
||||
return (head, *tail), head_params + tail_params
|
||||
|
||||
|
||||
def where_sql(where: tuple[Predicate, ...], first_index: int = 1) -> tuple[str, tuple[object, ...]]:
|
||||
"""`WHERE` body and its bind parameters, numbered from `first_index`.
|
||||
|
||||
Returns `("", ())` when there is nothing to filter on. Every caller-supplied value
|
||||
becomes a `$n` placeholder rather than being written into the SQL text; only column
|
||||
names reach the text, and those come from the spec's own declarations.
|
||||
"""
|
||||
clauses, params = _render_all(where, first_index)
|
||||
return " AND ".join(clauses), params
|
||||
|
||||
|
||||
def _problem(slug: str, title: str, status: int, detail: str, allowed: tuple[str, ...] | None = None) -> ProblemDetail:
|
||||
return ProblemDetail(
|
||||
type=f"{PROBLEM_TYPE_BASE}{slug}",
|
||||
title=title,
|
||||
status=status,
|
||||
detail=detail,
|
||||
allowed=sorted(allowed) if allowed is not None else None,
|
||||
)
|
||||
|
||||
|
||||
def _invalid(detail: str) -> ProblemDetail:
|
||||
return _problem("invalid-query-parameter", "Invalid query parameter", 400, detail)
|
||||
|
||||
|
||||
def _parse_filter_key(name: str) -> tuple[str, FilterOp] | None:
|
||||
"""`filter[max_budget][gte]` -> `("max_budget", "gte")`; bare `filter[status]` -> `("status", "eq")`."""
|
||||
if not name.startswith("filter[") or not name.endswith("]"):
|
||||
return None
|
||||
field, separator, raw_op = name[len("filter[") : -1].partition("][")
|
||||
if not separator:
|
||||
return field, "eq"
|
||||
try:
|
||||
return field, _FILTER_OP_ADAPTER.validate_python(raw_op)
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
def _is_known_param(spec: ListSpec[TRow, TOut], name: str) -> bool:
|
||||
if name in (PAGE_PARAM, PAGE_SIZE_PARAM):
|
||||
return True
|
||||
if name == SORT_PARAM:
|
||||
return bool(spec.sortable)
|
||||
if name == SEARCH_PARAM:
|
||||
return bool(spec.searchable)
|
||||
parsed = _parse_filter_key(name)
|
||||
return parsed is not None and parsed[0] in spec.filters
|
||||
|
||||
|
||||
def _allowed_params(spec: ListSpec[TRow, TOut]) -> tuple[str, ...]:
|
||||
return tuple(
|
||||
sorted(
|
||||
(PAGE_PARAM, PAGE_SIZE_PARAM)
|
||||
+ ((SORT_PARAM,) if spec.sortable else ())
|
||||
+ ((SEARCH_PARAM,) if spec.searchable else ())
|
||||
+ tuple(
|
||||
f"filter[{field}]" if op == "eq" else f"filter[{field}][{op}]"
|
||||
for field, filter_spec in spec.filters.items()
|
||||
for op in filter_spec.ops
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _parse_positive_int(name: str, raw: str) -> int | ProblemDetail:
|
||||
try:
|
||||
value = int(raw)
|
||||
except ValueError:
|
||||
return _invalid(f"'{name}' must be an integer.")
|
||||
if value < 1:
|
||||
return _invalid(f"'{name}' must be 1 or greater.")
|
||||
return value
|
||||
|
||||
|
||||
def _parse_page(params: Mapping[str, str]) -> int | ProblemDetail:
|
||||
raw = params.get(PAGE_PARAM)
|
||||
return 1 if raw is None else _parse_positive_int(PAGE_PARAM, raw)
|
||||
|
||||
|
||||
def _parse_page_size(spec: ListSpec[TRow, TOut], params: Mapping[str, str]) -> int | ProblemDetail:
|
||||
raw = params.get(PAGE_SIZE_PARAM)
|
||||
if raw is None:
|
||||
return spec.default_page_size
|
||||
value = _parse_positive_int(PAGE_SIZE_PARAM, raw)
|
||||
if isinstance(value, ProblemDetail):
|
||||
return value
|
||||
return min(value, spec.max_page_size)
|
||||
|
||||
|
||||
def _parse_sort(spec: ListSpec[TRow, TOut], params: Mapping[str, str]) -> tuple[SortKey, ...] | ProblemDetail:
|
||||
raw = params.get(SORT_PARAM)
|
||||
if raw is None:
|
||||
return spec.default_sort
|
||||
segments = tuple(segment.strip() for segment in raw.split(","))
|
||||
keys = tuple(
|
||||
SortKey(field=segment[1:] if segment.startswith("-") else segment, descending=segment.startswith("-"))
|
||||
for segment in segments
|
||||
)
|
||||
rejected = tuple(sorted(frozenset(key.field for key in keys) - spec.sortable))
|
||||
if rejected:
|
||||
return _problem(
|
||||
"invalid-sort-field",
|
||||
"Invalid sort field",
|
||||
400,
|
||||
f"Cannot sort {spec.resource} by: {', '.join(repr(field) for field in rejected)}.",
|
||||
tuple(spec.sortable),
|
||||
)
|
||||
return keys
|
||||
|
||||
|
||||
def _to_utc(value: datetime) -> datetime:
|
||||
return value.replace(tzinfo=timezone.utc) if value.tzinfo is None else value.astimezone(timezone.utc)
|
||||
|
||||
|
||||
def _coerce(field: str, op: FilterOp, raw: str, target: FilterType) -> FilterValue | ProblemDetail:
|
||||
try:
|
||||
if target is str:
|
||||
return raw
|
||||
if target is int:
|
||||
return int(raw)
|
||||
if target is float:
|
||||
return float(raw)
|
||||
return _to_utc(datetime.fromisoformat(raw[:-1] + "+00:00" if raw.endswith("Z") else raw))
|
||||
except ValueError:
|
||||
return _invalid(f"'filter[{field}][{op}]' is not a valid {target.__name__}: {raw!r}.")
|
||||
|
||||
|
||||
def _null_predicate(field: str, raw: str) -> Predicate | ProblemDetail:
|
||||
if raw.lower() == "true":
|
||||
return IsNull(field=field, negated=False)
|
||||
if raw.lower() == "false":
|
||||
return IsNull(field=field, negated=True)
|
||||
return _invalid(f"'filter[{field}][is_null]' must be 'true' or 'false'.")
|
||||
|
||||
|
||||
def _within_predicate(field: str, raw: str, target: FilterType) -> Predicate | ProblemDetail:
|
||||
coerced = tuple(_coerce(field, "in", item.strip(), target) for item in raw.split(","))
|
||||
problems = tuple(item for item in coerced if isinstance(item, ProblemDetail))
|
||||
if problems:
|
||||
return problems[0]
|
||||
return Within(field=field, values=tuple(item for item in coerced if not isinstance(item, ProblemDetail)))
|
||||
|
||||
|
||||
def _parse_filter(field: str, op: FilterOp, raw: str, filter_spec: FilterSpec) -> Predicate | ProblemDetail:
|
||||
if op not in filter_spec.ops:
|
||||
return _problem(
|
||||
"unsupported-filter-operator",
|
||||
"Unsupported filter operator",
|
||||
400,
|
||||
f"Operator '{op}' is not supported on '{field}'.",
|
||||
tuple(filter_spec.ops),
|
||||
)
|
||||
if op == "is_null":
|
||||
return _null_predicate(field, raw)
|
||||
if op == "in":
|
||||
return _within_predicate(field, raw, filter_spec.type)
|
||||
value = _coerce(field, op, raw, filter_spec.type)
|
||||
if isinstance(value, ProblemDetail):
|
||||
return value
|
||||
return Compare(field=field, op=op, value=value)
|
||||
|
||||
|
||||
def _parse_filters(spec: ListSpec[TRow, TOut], params: Mapping[str, str]) -> tuple[Predicate, ...] | ProblemDetail:
|
||||
keys = tuple(
|
||||
(name, parsed)
|
||||
for name in sorted(params)
|
||||
if (parsed := _parse_filter_key(name)) is not None and parsed[0] in spec.filters
|
||||
)
|
||||
parsed = tuple(_parse_filter(field, op, params[name], spec.filters[field]) for name, (field, op) in keys)
|
||||
problems = tuple(item for item in parsed if isinstance(item, ProblemDetail))
|
||||
if problems:
|
||||
return problems[0]
|
||||
return tuple(item for item in parsed if not isinstance(item, ProblemDetail))
|
||||
|
||||
|
||||
def _search_predicate(spec: ListSpec[TRow, TOut], params: Mapping[str, str]) -> Predicate | None:
|
||||
raw = params.get(SEARCH_PARAM)
|
||||
if not raw:
|
||||
return None
|
||||
return AnyOf(clauses=tuple(Compare(field=field, op="contains", value=raw) for field in sorted(spec.searchable)))
|
||||
|
||||
|
||||
def _scope_predicates(scope: Scope) -> tuple[Predicate, ...] | ProblemDetail:
|
||||
match scope:
|
||||
case ScopeAll():
|
||||
return ()
|
||||
case ScopeWhere(where=where):
|
||||
return where
|
||||
case ScopeDenied(reason=reason):
|
||||
return _problem("forbidden", "Forbidden", 403, reason)
|
||||
case _:
|
||||
assert_never(scope)
|
||||
|
||||
|
||||
def build_query_plan(
|
||||
spec: ListSpec[TRow, TOut],
|
||||
params: Mapping[str, str],
|
||||
caller: UserAPIKeyAuth,
|
||||
) -> QueryPlan | ProblemDetail:
|
||||
"""Turn query parameters into a plan, or into the problem that explains why they are not one."""
|
||||
scope_predicates = _scope_predicates(spec.scope(caller))
|
||||
if isinstance(scope_predicates, ProblemDetail):
|
||||
return scope_predicates
|
||||
|
||||
unknown = tuple(sorted(name for name in params if not _is_known_param(spec, name)))
|
||||
if unknown:
|
||||
return unknown_query_param_problem(unknown=unknown, allowed=_allowed_params(spec))
|
||||
|
||||
page = _parse_page(params)
|
||||
if isinstance(page, ProblemDetail):
|
||||
return page
|
||||
|
||||
page_size = _parse_page_size(spec, params)
|
||||
if isinstance(page_size, ProblemDetail):
|
||||
return page_size
|
||||
|
||||
sort = _parse_sort(spec, params)
|
||||
if isinstance(sort, ProblemDetail):
|
||||
return sort
|
||||
|
||||
filters = _parse_filters(spec, params)
|
||||
if isinstance(filters, ProblemDetail):
|
||||
return filters
|
||||
|
||||
search = _search_predicate(spec, params)
|
||||
return QueryPlan(
|
||||
# Scope first: conjuncts a caller filter sits behind and cannot replace.
|
||||
where=scope_predicates + filters + ((search,) if search is not None else ()),
|
||||
# Ordering by an all-null column without a unique final key lets Postgres return
|
||||
# the same row on two different pages.
|
||||
order=sort + (SortKey(field=spec.tiebreaker, descending=False),),
|
||||
skip=(page - 1) * page_size,
|
||||
take=page_size,
|
||||
)
|
||||
|
||||
|
||||
def _duplicate_params(request: Request) -> tuple[str, ...]:
|
||||
names = tuple(name for name, _ in request.query_params.multi_items())
|
||||
return tuple(sorted(frozenset(name for name in names if names.count(name) > 1)))
|
||||
|
||||
|
||||
async def handle_list(
|
||||
spec: ListSpec[TRow, TOut],
|
||||
executor: ListExecutor[TRow],
|
||||
request: Request,
|
||||
caller: UserAPIKeyAuth,
|
||||
) -> ListResponse[TOut]:
|
||||
"""Plan, execute, count, serialize, envelope. Failures reach the client as RFC 9457 problems."""
|
||||
plan = build_query_plan(spec=spec, params=request.query_params, caller=caller)
|
||||
if isinstance(plan, ProblemDetail):
|
||||
raise ManagementProblem(plan)
|
||||
|
||||
# Checked here rather than in build_query_plan because a Mapping[str, str] cannot
|
||||
# represent a repeat: query_params.get() silently keeps the last one, so ?page=1&page=999
|
||||
# would page from 999 without the caller ever being told which value won.
|
||||
duplicates = _duplicate_params(request)
|
||||
if duplicates:
|
||||
raise ManagementProblem(
|
||||
_problem(
|
||||
"duplicate-query-parameter",
|
||||
"Duplicate query parameter",
|
||||
400,
|
||||
f"Repeated query parameter(s): {', '.join(duplicates)}. Each may appear once; "
|
||||
f"use a comma-separated list for multiple sort keys or filter values.",
|
||||
)
|
||||
)
|
||||
|
||||
total_count = await executor.count(plan.where)
|
||||
rows = await executor.find_many(plan)
|
||||
total_pages = ceil(total_count / plan.take)
|
||||
page = plan.skip // plan.take + 1
|
||||
return ListResponse[TOut](
|
||||
data=tuple(spec.serialize(row) for row in rows),
|
||||
meta=ListMeta(
|
||||
total_count=total_count,
|
||||
page=page,
|
||||
page_size=plan.take,
|
||||
total_pages=total_pages,
|
||||
),
|
||||
links=build_list_links(request=request, page=page, total_pages=total_pages),
|
||||
)
|
||||
|
|
@ -13,6 +13,7 @@ from litellm.proxy.management_endpoints.management_v1.common import (
|
|||
PROBLEM_TYPE_BASE,
|
||||
ManagementProblem,
|
||||
build_page_links,
|
||||
escape_like,
|
||||
reject_unknown_query_params,
|
||||
)
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
|
@ -34,10 +35,6 @@ def _as_utc(value: datetime) -> datetime:
|
|||
return value.replace(tzinfo=timezone.utc) if value.tzinfo is None else value.astimezone(timezone.utc)
|
||||
|
||||
|
||||
def _escape_like(value: str) -> str:
|
||||
return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
|
||||
|
||||
|
||||
async def _end_user_scope_clause(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
prisma_client: PrismaClient,
|
||||
|
|
@ -133,7 +130,7 @@ async def list_spend_log_end_users(
|
|||
)
|
||||
|
||||
window_params: tuple[Any, ...] = (_as_utc(start_time), _as_utc(end_time))
|
||||
search_params: tuple[Any, ...] = (f"%{_escape_like(q)}%",) if q else ()
|
||||
search_params: tuple[Any, ...] = (f"%{escape_like(q)}%",) if q else ()
|
||||
search_clause = (f"end_user ILIKE ${len(window_params) + 1} ESCAPE '\\'",) if q else ()
|
||||
|
||||
scope_clause, scope_params = await _end_user_scope_clause(
|
||||
|
|
|
|||
|
|
@ -1,7 +1,11 @@
|
|||
"""Shared response shapes for the `/management/v1` control-plane surface."""
|
||||
|
||||
from typing import Generic, TypeVar
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
TOut = TypeVar("TOut")
|
||||
|
||||
|
||||
class ProblemDetail(BaseModel):
|
||||
"""RFC 9457 problem details, served as `application/problem+json`."""
|
||||
|
|
@ -37,3 +41,33 @@ class FacetListResponse(BaseModel):
|
|||
data: list[str]
|
||||
meta: PageMeta
|
||||
links: PageLinks
|
||||
|
||||
|
||||
class ListMeta(BaseModel):
|
||||
"""Page-mode counterpart to `PageMeta`: an entity list pays for the COUNT(*) so the table can show a page count."""
|
||||
|
||||
total_count: int
|
||||
page: int
|
||||
page_size: int
|
||||
total_pages: int
|
||||
|
||||
|
||||
class ListLinks(BaseModel):
|
||||
"""Page-mode counterpart to `PageLinks`. `first`/`last` are knowable here because the total count is."""
|
||||
|
||||
model_config = ConfigDict(populate_by_name=True)
|
||||
|
||||
self_link: str = Field(alias="self")
|
||||
first: str
|
||||
prev: str | None = None
|
||||
next: str | None = None
|
||||
last: str
|
||||
|
||||
|
||||
class ListResponse(BaseModel, Generic[TOut]):
|
||||
"""Rows stay flat: JSON:API's `{type, id, attributes}` wrapper is a deliberate deviation, so every
|
||||
dashboard column accessor would otherwise have to go through `.attributes`."""
|
||||
|
||||
data: list[TOut]
|
||||
meta: ListMeta
|
||||
links: ListLinks
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import asyncio
|
||||
import concurrent.futures
|
||||
import os
|
||||
import sys
|
||||
|
||||
|
|
@ -827,3 +828,305 @@ async def test_stale_loop_rebuild_does_not_close_unowned_session():
|
|||
shared_session._loop = running_loop
|
||||
other_loop.close()
|
||||
await shared_session.close()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Recycled-session leak tests (#24230)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def _new_session() -> aiohttp.ClientSession:
|
||||
return aiohttp.ClientSession()
|
||||
|
||||
|
||||
def _make_session_on_dead_loop() -> aiohttp.ClientSession:
|
||||
"""Create a ClientSession bound to an event loop that is then closed.
|
||||
|
||||
Runs in a worker thread: the caller may already be inside a running
|
||||
event loop, where a nested run_until_complete is forbidden.
|
||||
"""
|
||||
import threading
|
||||
|
||||
result: dict = {}
|
||||
|
||||
def build() -> None:
|
||||
loop = asyncio.new_event_loop()
|
||||
try:
|
||||
result["session"] = loop.run_until_complete(_new_session())
|
||||
finally:
|
||||
loop.close()
|
||||
|
||||
thread = threading.Thread(target=build)
|
||||
thread.start()
|
||||
thread.join(5)
|
||||
return result["session"]
|
||||
|
||||
|
||||
def _flaky_get_running_loop_factory():
|
||||
"""get_running_loop stand-in that fails once, then delegates.
|
||||
|
||||
Reproduces #24230: a transient loop-inspection failure sends
|
||||
_get_valid_client_session into its (RuntimeError, AttributeError)
|
||||
fallback branch.
|
||||
"""
|
||||
real_get_running_loop = asyncio.get_running_loop
|
||||
calls = {"count": 0}
|
||||
|
||||
def flaky():
|
||||
calls["count"] += 1
|
||||
if calls["count"] == 1:
|
||||
raise RuntimeError("simulated loop inspection failure")
|
||||
return real_get_running_loop()
|
||||
|
||||
return flaky
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fallback_recreate_closes_previous_session():
|
||||
"""
|
||||
Regression test for #24230: when loop inspection fails and the fallback
|
||||
branch recreates the session, the replaced session must still be closed -
|
||||
not silently abandoned to the garbage collector.
|
||||
"""
|
||||
from unittest.mock import patch
|
||||
|
||||
old_session = aiohttp.ClientSession()
|
||||
transport = LiteLLMAiohttpTransport(client=lambda: aiohttp.ClientSession())
|
||||
transport.client = old_session
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.aiohttp_transport.asyncio.get_running_loop",
|
||||
side_effect=_flaky_get_running_loop_factory(),
|
||||
):
|
||||
new_session = transport._get_valid_client_session()
|
||||
|
||||
try:
|
||||
assert new_session is not old_session
|
||||
for _ in range(3):
|
||||
await asyncio.sleep(0)
|
||||
assert old_session.closed, "replaced session must be closed, not leaked"
|
||||
finally:
|
||||
await new_session.close()
|
||||
if not old_session.closed:
|
||||
await old_session.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_replaced_session_emits_no_unclosed_warnings():
|
||||
"""
|
||||
Regression test for #24230: a session replaced by the fallback branch must
|
||||
not surface "Unclosed client session" / "Unclosed connector" warnings when
|
||||
the garbage collector finalizes it.
|
||||
"""
|
||||
import gc
|
||||
import warnings as warnings_mod
|
||||
from unittest.mock import patch
|
||||
|
||||
old_session = aiohttp.ClientSession()
|
||||
transport = LiteLLMAiohttpTransport(client=lambda: aiohttp.ClientSession())
|
||||
transport.client = old_session
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.aiohttp_transport.asyncio.get_running_loop",
|
||||
side_effect=_flaky_get_running_loop_factory(),
|
||||
):
|
||||
new_session = transport._get_valid_client_session()
|
||||
|
||||
try:
|
||||
for _ in range(3):
|
||||
await asyncio.sleep(0)
|
||||
|
||||
del old_session
|
||||
with warnings_mod.catch_warnings(record=True) as caught:
|
||||
warnings_mod.simplefilter("always")
|
||||
gc.collect()
|
||||
|
||||
unclosed = [
|
||||
str(w.message)
|
||||
for w in caught
|
||||
if "Unclosed client session" in str(w.message) or "Unclosed connector" in str(w.message)
|
||||
]
|
||||
assert not unclosed, f"leaked session warnings: {unclosed}"
|
||||
finally:
|
||||
await new_session.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dead_loop_session_closed_synchronously_on_recycle():
|
||||
"""
|
||||
Regression test for #24230: a session whose event loop is already closed
|
||||
cannot run an async close anywhere. Recycling it must dispose of it
|
||||
deterministically, the session reads closed as soon as the recycle
|
||||
returns, so no finalizer warning window remains.
|
||||
"""
|
||||
old_session = _make_session_on_dead_loop()
|
||||
transport = LiteLLMAiohttpTransport(client=lambda: aiohttp.ClientSession())
|
||||
transport.client = old_session
|
||||
|
||||
new_session = transport._get_valid_client_session()
|
||||
|
||||
try:
|
||||
assert new_session is not old_session
|
||||
assert old_session.closed, "session from a closed loop must be disposed synchronously at recycle"
|
||||
finally:
|
||||
await new_session.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_close_task_strongly_referenced_until_done():
|
||||
"""
|
||||
Regression test for #24230: scheduled session-close tasks must be strongly
|
||||
referenced (and pruned on completion) so they cannot be garbage-collected
|
||||
before they run.
|
||||
"""
|
||||
old_session = aiohttp.ClientSession()
|
||||
transport = LiteLLMAiohttpTransport(client=lambda: aiohttp.ClientSession())
|
||||
|
||||
transport._close_recycled_session(old_session)
|
||||
|
||||
assert LiteLLMAiohttpTransport._background_close_tasks, "close task must be strongly referenced while pending"
|
||||
for _ in range(5):
|
||||
await asyncio.sleep(0)
|
||||
assert old_session.closed
|
||||
assert not LiteLLMAiohttpTransport._background_close_tasks, "completed close tasks must be pruned from the registry"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_from_other_running_loop_closed_threadsafe():
|
||||
"""
|
||||
Regression test for #24230: a session that belongs to a loop still running
|
||||
in another thread must be closed on its own loop (thread-safe), not driven
|
||||
from the current loop.
|
||||
"""
|
||||
import threading
|
||||
import time
|
||||
|
||||
ready = threading.Event()
|
||||
holder: dict = {}
|
||||
|
||||
def worker() -> None:
|
||||
loop = asyncio.new_event_loop()
|
||||
holder["loop"] = loop
|
||||
|
||||
async def make() -> None:
|
||||
holder["session"] = aiohttp.ClientSession()
|
||||
|
||||
loop.run_until_complete(make())
|
||||
ready.set()
|
||||
loop.run_forever()
|
||||
loop.close()
|
||||
|
||||
thread = threading.Thread(target=worker, daemon=True)
|
||||
thread.start()
|
||||
assert ready.wait(5), "worker loop failed to start"
|
||||
|
||||
transport = LiteLLMAiohttpTransport(client=lambda: aiohttp.ClientSession())
|
||||
transport.client = holder["session"]
|
||||
|
||||
new_session = transport._get_valid_client_session()
|
||||
|
||||
try:
|
||||
deadline = time.monotonic() + 5
|
||||
while not holder["session"].closed and time.monotonic() < deadline:
|
||||
await asyncio.sleep(0.01)
|
||||
assert holder["session"].closed, "foreign-loop session was never closed"
|
||||
finally:
|
||||
holder["loop"].call_soon_threadsafe(holder["loop"].stop)
|
||||
thread.join(5)
|
||||
await new_session.close()
|
||||
|
||||
|
||||
def test_threadsafe_close_done_callback_tolerates_cancelled_future():
|
||||
"""
|
||||
Regression test for #24230 (review finding): when the foreign loop stops
|
||||
before the handed-off close coroutine runs, asyncio cancels the
|
||||
concurrent.futures.Future. The done-callback must return quietly instead
|
||||
of letting future.exception() raise CancelledError (a BaseException that
|
||||
escapes _invoke_callbacks and crashes the foreign loop's thread).
|
||||
"""
|
||||
future: "concurrent.futures.Future[None]" = concurrent.futures.Future()
|
||||
future.cancel()
|
||||
|
||||
LiteLLMAiohttpTransport._on_threadsafe_close_done(future)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_closed_retry_does_not_close_concurrent_replacement():
|
||||
"""
|
||||
Regression test for #24230 (review finding): when the "Session is closed"
|
||||
retry fires, the handler must dispose the session that actually faulted,
|
||||
not self.client - a concurrent task may already have replaced self.client
|
||||
with a healthy session, which must stay open.
|
||||
"""
|
||||
from unittest.mock import patch
|
||||
|
||||
faulted_session = aiohttp.ClientSession()
|
||||
healthy_replacement = aiohttp.ClientSession()
|
||||
transport = LiteLLMAiohttpTransport(client=lambda: aiohttp.ClientSession())
|
||||
transport.client = faulted_session
|
||||
|
||||
calls = {"n": 0}
|
||||
|
||||
async def fake_make_request(*args, **kwargs):
|
||||
calls["n"] += 1
|
||||
if calls["n"] == 1:
|
||||
# simulate a concurrent task replacing the shared session between
|
||||
# the failed await and the exception handler
|
||||
transport.client = healthy_replacement
|
||||
raise RuntimeError("Session is closed")
|
||||
raise StopAsyncIteration("stop after retry dispatch")
|
||||
|
||||
with patch.object(transport, "_make_aiohttp_request", side_effect=fake_make_request):
|
||||
with pytest.raises(Exception):
|
||||
await transport.handle_async_request(httpx.Request("GET", "http://example.com"))
|
||||
|
||||
try:
|
||||
assert not healthy_replacement.closed, "concurrent replacement session must not be closed by the retry handler"
|
||||
for _ in range(3):
|
||||
await asyncio.sleep(0)
|
||||
assert faulted_session.closed, "the faulted session must be disposed"
|
||||
finally:
|
||||
await faulted_session.close()
|
||||
await healthy_replacement.close()
|
||||
new_session = transport.client
|
||||
if isinstance(new_session, aiohttp.ClientSession):
|
||||
await new_session.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stopped_loop_session_disposed_synchronously_on_recycle():
|
||||
"""
|
||||
Regression test for #24230 (review finding): a session whose loop is
|
||||
stopped but not yet closed cannot safely run an async close on another
|
||||
loop, and nothing will ever process a close handed to the stopped loop.
|
||||
Recycling must dispose it synchronously, like the closed-loop case.
|
||||
"""
|
||||
import threading
|
||||
|
||||
result: dict = {}
|
||||
|
||||
def build() -> None:
|
||||
loop = asyncio.new_event_loop()
|
||||
|
||||
async def make() -> None:
|
||||
result["session"] = aiohttp.ClientSession()
|
||||
|
||||
loop.run_until_complete(make())
|
||||
result["loop"] = loop # stopped, deliberately NOT closed
|
||||
|
||||
thread = threading.Thread(target=build)
|
||||
thread.start()
|
||||
thread.join(5)
|
||||
|
||||
old_session = result["session"]
|
||||
transport = LiteLLMAiohttpTransport(client=lambda: aiohttp.ClientSession())
|
||||
transport.client = old_session
|
||||
|
||||
new_session = transport._get_valid_client_session()
|
||||
|
||||
try:
|
||||
assert new_session is not old_session
|
||||
assert old_session.closed, "session from a stopped (not yet closed) loop must be disposed synchronously"
|
||||
finally:
|
||||
await new_session.close()
|
||||
result["loop"].close()
|
||||
|
|
|
|||
|
|
@ -8191,3 +8191,177 @@ class TestGetUserObjectPermission:
|
|||
async def test_no_user_id_places_no_ceiling(self):
|
||||
assert await MCPRequestHandler._get_user_object_permission(UserAPIKeyAuth(api_key="sk-test")) is None
|
||||
assert await MCPRequestHandler._get_user_object_permission(None) is None
|
||||
|
||||
|
||||
def _key_auth_reaching(server, *, tools=None, **fields):
|
||||
"""A key-authenticated caller whose OWN key grant reaches ``server`` (and optionally its ``tools``).
|
||||
|
||||
The key grant is the thing an upper-level entitlement fault must not silently hand back: every
|
||||
test below asserts against what this key reaches when the level under test cannot be resolved.
|
||||
"""
|
||||
return UserAPIKeyAuth(
|
||||
api_key="sk-hash",
|
||||
user_id="u1",
|
||||
object_permission=LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="op-key",
|
||||
mcp_servers=[server],
|
||||
mcp_tool_permissions={server: tools} if tools else None,
|
||||
),
|
||||
**fields,
|
||||
)
|
||||
|
||||
|
||||
def _agent_prisma(object_permission_id=None, side_effect=None):
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_agentstable.find_unique = AsyncMock(
|
||||
return_value=MagicMock(object_permission_id=object_permission_id),
|
||||
side_effect=side_effect,
|
||||
)
|
||||
return prisma_client
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _entitlement_fault_globals(prisma_client=None):
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", prisma_client or MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
|
||||
):
|
||||
yield
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestEntitlementFaultSemantics:
|
||||
"""Each entitlement level distinguishes two fault classes for a KEY-authenticated caller.
|
||||
|
||||
A principal row that NAMES an object_permission we cannot load is a known entitlement with
|
||||
unknown contents, so the level denies rather than handing back the wider key scope. A lookup
|
||||
that fails before we can tell whether the principal is entitled at all leaves no ceiling, which
|
||||
is the state that existed before the level did; denying there would refuse MCP to the majority
|
||||
of callers, who have no such entitlement configured, for the duration of a cold-cache fault.
|
||||
"""
|
||||
|
||||
async def test_end_user_named_but_unloadable_permission_denies(self):
|
||||
end_user = MagicMock(object_permission=None, object_permission_id="op-eu")
|
||||
auth = _key_auth_reaching("srv1", end_user_id="eu-1")
|
||||
with _entitlement_fault_globals():
|
||||
with (
|
||||
patch("litellm.proxy.auth.auth_checks.get_end_user_object", AsyncMock(return_value=end_user)),
|
||||
patch("litellm.proxy.auth.auth_checks.get_object_permission", AsyncMock(return_value=None)),
|
||||
):
|
||||
allowed = await MCPRequestHandler.get_allowed_mcp_servers(auth)
|
||||
assert allowed == [], "an end-user entitlement we know exists but cannot read must deny"
|
||||
|
||||
async def test_end_user_without_an_entitlement_places_no_ceiling(self):
|
||||
"""The three shapes that are NOT evidence of an entitlement: an end user row linking no
|
||||
permission, no end user row at all, and a lookup that blew up before answering either."""
|
||||
auth = _key_auth_reaching("srv1", end_user_id="eu-1")
|
||||
linked_none = MagicMock(object_permission=None, object_permission_id=None)
|
||||
for lookup, shape in (
|
||||
(AsyncMock(return_value=linked_none), "row links no permission"),
|
||||
(AsyncMock(return_value=None), "no end user row"),
|
||||
(AsyncMock(side_effect=RuntimeError("connection reset by peer")), "lookup failed"),
|
||||
):
|
||||
with _entitlement_fault_globals():
|
||||
with patch("litellm.proxy.auth.auth_checks.get_end_user_object", lookup):
|
||||
allowed = await MCPRequestHandler.get_allowed_mcp_servers(auth)
|
||||
assert set(allowed) == {"srv1"}, f"{shape}: no evidence of an entitlement, so no ceiling"
|
||||
|
||||
async def test_agent_named_but_unloadable_permission_denies(self):
|
||||
auth = _key_auth_reaching("srv1", agent_id="agent-unloadable")
|
||||
with _entitlement_fault_globals(_agent_prisma(object_permission_id="op-agent")):
|
||||
with patch("litellm.proxy.auth.auth_checks.get_object_permission", AsyncMock(return_value=None)):
|
||||
allowed = await MCPRequestHandler.get_allowed_mcp_servers(auth)
|
||||
assert allowed == [], "an agent entitlement we know exists but cannot read must deny"
|
||||
|
||||
async def test_agent_without_an_entitlement_places_no_ceiling(self):
|
||||
"""An agent row linking no permission, and an agent row we could not read at all."""
|
||||
for prisma_client, agent_id, shape in (
|
||||
(_agent_prisma(object_permission_id=None), "agent-unlinked", "agent links no permission"),
|
||||
(_agent_prisma(side_effect=RuntimeError("connection reset by peer")), "agent-unread", "row read failed"),
|
||||
):
|
||||
auth = _key_auth_reaching("srv1", agent_id=agent_id)
|
||||
with _entitlement_fault_globals(prisma_client):
|
||||
allowed = await MCPRequestHandler.get_allowed_mcp_servers(auth)
|
||||
assert set(allowed) == {"srv1"}, f"{shape}: no evidence of an entitlement, so no ceiling"
|
||||
|
||||
async def test_agent_named_but_unloadable_permission_denies_tools(self):
|
||||
"""The tools axis denies with [] rather than the None (allow-all) key auth gets for an
|
||||
indeterminate fault, so an unreadable agent entitlement cannot widen the key's tool scope."""
|
||||
auth = _key_auth_reaching("srv1", tools=["tool_a"], agent_id="agent-tools-unloadable")
|
||||
with _entitlement_fault_globals(_agent_prisma(object_permission_id="op-agent")):
|
||||
with patch("litellm.proxy.auth.auth_checks.get_object_permission", AsyncMock(return_value=None)):
|
||||
tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth)
|
||||
assert tools == [], "an agent entitlement we know exists but cannot read must deny its tools"
|
||||
|
||||
async def test_org_named_but_unloadable_ceiling_denies(self):
|
||||
auth = _key_auth_reaching("srv1", org_id="org-a")
|
||||
org = MagicMock(object_permission_id="op-org")
|
||||
with _entitlement_fault_globals():
|
||||
with (
|
||||
patch("litellm.proxy.auth.auth_checks.get_org_object", AsyncMock(return_value=org)),
|
||||
patch("litellm.proxy.auth.auth_checks.get_object_permission", AsyncMock(return_value=None)),
|
||||
):
|
||||
allowed = await MCPRequestHandler.get_allowed_mcp_servers(auth)
|
||||
assert allowed == [], "an org ceiling we know exists but cannot read must deny, key auth included"
|
||||
|
||||
async def test_org_named_but_unloadable_ceiling_denies_tools(self):
|
||||
auth = _key_auth_reaching("srv1", tools=["tool_a"], org_id="org-a")
|
||||
org = MagicMock(object_permission_id="op-org")
|
||||
with _entitlement_fault_globals():
|
||||
with (
|
||||
patch("litellm.proxy.auth.auth_checks.get_org_object", AsyncMock(return_value=org)),
|
||||
patch("litellm.proxy.auth.auth_checks.get_object_permission", AsyncMock(return_value=None)),
|
||||
):
|
||||
tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth)
|
||||
assert tools == [], "an org tool ceiling we know exists but cannot read must deny its tools"
|
||||
|
||||
async def test_org_without_a_resolvable_entitlement_places_no_ceiling(self):
|
||||
"""A deleted org and an org lookup that failed are both cases where we cannot point at a
|
||||
ceiling; key auth keeps its long-standing fail-open behavior for them."""
|
||||
from litellm.proxy.auth.auth_checks import OrganizationNotFoundError
|
||||
|
||||
auth = _key_auth_reaching("srv1", org_id="org-a")
|
||||
for lookup, shape in (
|
||||
(AsyncMock(return_value=MagicMock(object_permission_id=None)), "org names no permission"),
|
||||
(AsyncMock(side_effect=OrganizationNotFoundError("Organization doesn't exist in db.")), "org deleted"),
|
||||
(AsyncMock(side_effect=RuntimeError("connection reset by peer")), "org lookup failed"),
|
||||
):
|
||||
with _entitlement_fault_globals():
|
||||
with patch("litellm.proxy.auth.auth_checks.get_org_object", lookup):
|
||||
allowed = await MCPRequestHandler.get_allowed_mcp_servers(auth)
|
||||
assert set(allowed) == {"srv1"}, f"{shape}: no ceiling we can point at, so key auth stays open"
|
||||
|
||||
async def test_keyless_org_ceiling_denies_on_either_fault_class(self):
|
||||
"""The keyless gateway-admitted path is untouched: it already denied on ANY org-ceiling
|
||||
fault, and still denies on both classes, because a per-source org ceiling is the only org
|
||||
bound a keyless subject has and an unbounded source would win the union."""
|
||||
auth = _make_admitted_subject("sso-user", org_id="org-a", own_servers=["srv1"])
|
||||
org = MagicMock(object_permission_id="op-org")
|
||||
with _entitlement_fault_globals():
|
||||
with patch("litellm.proxy.auth.auth_checks.get_org_object", AsyncMock(return_value=org)):
|
||||
with patch("litellm.proxy.auth.auth_checks.get_object_permission", AsyncMock(return_value=None)):
|
||||
named_unloadable = await MCPRequestHandler.get_allowed_mcp_servers(auth)
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks.get_org_object",
|
||||
AsyncMock(side_effect=RuntimeError("connection reset by peer")),
|
||||
):
|
||||
indeterminate = await MCPRequestHandler.get_allowed_mcp_servers(auth)
|
||||
assert named_unloadable == [] and indeterminate == []
|
||||
|
||||
async def test_keyless_source_never_consults_the_end_user_or_agent_levels(self):
|
||||
"""A keyless subject's grant sources carry neither end_user_id nor agent_id, so neither
|
||||
level runs for it and neither new deny can reach its union. Pinned because a source that
|
||||
DID consult them would fail closed on a fault and silently drop a team's grants."""
|
||||
auth = _make_admitted_subject("sso-user", own_servers=["srv1"])
|
||||
auth.end_user_id = "eu-1"
|
||||
auth.agent_id = "agent-unloadable"
|
||||
with _entitlement_fault_globals(_agent_prisma(object_permission_id="op-agent")):
|
||||
with (
|
||||
patch("litellm.proxy.auth.auth_checks.get_end_user_object", AsyncMock(side_effect=AssertionError)),
|
||||
patch("litellm.proxy.auth.auth_checks.get_object_permission", AsyncMock(return_value=None)),
|
||||
):
|
||||
allowed = await MCPRequestHandler.get_allowed_mcp_servers(auth)
|
||||
assert set(allowed) == {"srv1"}
|
||||
|
|
|
|||
|
|
@ -0,0 +1,871 @@
|
|||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass, replace
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import pytest
|
||||
from fastapi import Request
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.management_endpoints.management_v1.common import (
|
||||
MANAGEMENT_V1_PREFIX,
|
||||
PROBLEM_TYPE_BASE,
|
||||
ManagementProblem,
|
||||
build_page_links,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.management_v1.list_framework import (
|
||||
AnyOf,
|
||||
Compare,
|
||||
FilterSpec,
|
||||
IsNull,
|
||||
ListSpec,
|
||||
QueryPlan,
|
||||
ScopeAll,
|
||||
ScopeDenied,
|
||||
ScopeWhere,
|
||||
SortKey,
|
||||
Within,
|
||||
build_query_plan,
|
||||
handle_list,
|
||||
order_by_sql,
|
||||
where_sql,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.management_v1 import (
|
||||
PageLinks,
|
||||
PageMeta,
|
||||
ProblemDetail,
|
||||
)
|
||||
|
||||
BUDGETS_PATH = f"{MANAGEMENT_V1_PREFIX}/budgets"
|
||||
CALLER = UserAPIKeyAuth(user_id="caller-1")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BudgetRow:
|
||||
budget_id: str
|
||||
max_budget: float | None
|
||||
created_by: str
|
||||
|
||||
|
||||
class BudgetOut(BaseModel):
|
||||
budget_id: str
|
||||
max_budget: float | None
|
||||
|
||||
|
||||
def _serialize(row: BudgetRow) -> BudgetOut:
|
||||
return BudgetOut(budget_id=row.budget_id, max_budget=row.max_budget)
|
||||
|
||||
|
||||
def _spec(
|
||||
scope=lambda caller: ScopeAll(),
|
||||
searchable=frozenset({"budget_id", "created_by"}),
|
||||
sortable=frozenset({"max_budget", "created_at", "budget_id"}),
|
||||
) -> ListSpec[BudgetRow, BudgetOut]:
|
||||
return ListSpec(
|
||||
resource="budgets",
|
||||
sortable=sortable,
|
||||
searchable=searchable,
|
||||
filters={
|
||||
"max_budget": FilterSpec(type=float, ops=frozenset({"eq", "gte", "lte", "is_null"})),
|
||||
"created_at": FilterSpec(type=datetime, ops=frozenset({"gte", "lte"})),
|
||||
"created_by": FilterSpec(type=str, ops=frozenset({"eq", "in", "contains"})),
|
||||
"tpm_limit": FilterSpec(type=int, ops=frozenset({"eq"})),
|
||||
},
|
||||
default_sort=(SortKey(field="created_at", descending=True),),
|
||||
default_page_size=25,
|
||||
max_page_size=100,
|
||||
scope=scope,
|
||||
serialize=_serialize,
|
||||
tiebreaker="budget_id",
|
||||
)
|
||||
|
||||
|
||||
def _spec_with(**overrides) -> ListSpec[BudgetRow, BudgetOut]:
|
||||
"""`replace` re-runs `__init__`, so the spec's own validation applies to the override."""
|
||||
return replace(_spec(), **overrides)
|
||||
|
||||
|
||||
class RecordingExecutor:
|
||||
"""In-memory stand-in for the Prisma-backed executor PR 2 supplies."""
|
||||
|
||||
def __init__(self, rows: tuple[BudgetRow, ...], total_count: int | None = None) -> None:
|
||||
self.rows = rows
|
||||
self.total_count = len(rows) if total_count is None else total_count
|
||||
self.plan: QueryPlan | None = None
|
||||
self.count_where: tuple[object, ...] | None = None
|
||||
|
||||
async def count(self, where: tuple[object, ...]) -> int:
|
||||
self.count_where = where
|
||||
return self.total_count
|
||||
|
||||
async def find_many(self, plan: QueryPlan) -> Sequence[BudgetRow]:
|
||||
self.plan = plan
|
||||
return self.rows[plan.skip : plan.skip + plan.take]
|
||||
|
||||
|
||||
def _request(query: str = "") -> Request:
|
||||
return Request(
|
||||
{
|
||||
"type": "http",
|
||||
"method": "GET",
|
||||
"scheme": "http",
|
||||
"root_path": "",
|
||||
"path": BUDGETS_PATH,
|
||||
"query_string": query.encode(),
|
||||
"headers": [(b"host", b"testserver")],
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _plan(query_params: Mapping[str, str], spec: ListSpec[BudgetRow, BudgetOut] | None = None) -> QueryPlan:
|
||||
result = build_query_plan(spec=spec or _spec(), params=query_params, caller=CALLER)
|
||||
assert isinstance(result, QueryPlan), result
|
||||
return result
|
||||
|
||||
|
||||
def _problem(query_params: Mapping[str, str], spec: ListSpec[BudgetRow, BudgetOut] | None = None) -> ProblemDetail:
|
||||
result = build_query_plan(spec=spec or _spec(), params=query_params, caller=CALLER)
|
||||
assert isinstance(result, ProblemDetail), result
|
||||
return result
|
||||
|
||||
|
||||
def _conjuncts(plan: QueryPlan) -> tuple[object, ...]:
|
||||
return plan.where
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- invariant 1
|
||||
|
||||
|
||||
def test_appends_the_tiebreaker_to_the_default_sort():
|
||||
"""Without a unique final key, ordering by an all-null column lets Postgres hand the
|
||||
same row back on two different pages."""
|
||||
assert _plan({}).order == (SortKey(field="created_at", descending=True), SortKey(field="budget_id", descending=False))
|
||||
|
||||
|
||||
def test_appends_the_tiebreaker_to_an_explicit_multi_key_sort():
|
||||
order = _plan({"sort": "-max_budget,created_at"}).order
|
||||
|
||||
assert len(order) == 3
|
||||
assert order[-1] == SortKey(field="budget_id", descending=False)
|
||||
|
||||
|
||||
def test_appends_the_tiebreaker_even_when_the_caller_already_sorts_by_it():
|
||||
"""Deduplicating it away is the tempting simplification, and it is the one that
|
||||
reintroduces a non-total order the moment the leading key stops being unique."""
|
||||
assert _plan({"sort": "-budget_id"}).order == (
|
||||
SortKey(field="budget_id", descending=True),
|
||||
SortKey(field="budget_id", descending=False),
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- invariant 2
|
||||
|
||||
|
||||
def test_orders_nulls_last_in_both_directions():
|
||||
"""Postgres sorts nulls last ascending but first descending, so flipping the sort
|
||||
direction on max_budget would otherwise float every "Unlimited" row to the top."""
|
||||
sql = order_by_sql((SortKey(field="max_budget", descending=True), SortKey(field="budget_id", descending=False)))
|
||||
|
||||
assert sql == '"max_budget" DESC NULLS LAST, "budget_id" ASC NULLS LAST'
|
||||
|
||||
|
||||
def test_order_sql_covers_every_key_in_the_plan():
|
||||
sql = order_by_sql(_plan({"sort": "-max_budget,created_at"}).order)
|
||||
|
||||
assert sql.count("NULLS LAST") == 3
|
||||
assert sql == '"max_budget" DESC NULLS LAST, "created_at" ASC NULLS LAST, "budget_id" ASC NULLS LAST'
|
||||
|
||||
|
||||
# ------------------------------------------------------------ where rendering
|
||||
|
||||
|
||||
def test_every_caller_value_is_bound_not_interpolated():
|
||||
"""The one property that keeps a filter value from reaching the SQL text. A value that
|
||||
looks like SQL has to come back as a parameter, never as part of the statement."""
|
||||
sql, params = where_sql((Compare(field="budget_id", op="eq", value="'; DROP TABLE x --"),))
|
||||
|
||||
assert sql == '"budget_id" = $1'
|
||||
assert params == ("'; DROP TABLE x --",)
|
||||
assert "DROP" not in sql
|
||||
|
||||
|
||||
def test_placeholders_are_numbered_across_the_whole_plan():
|
||||
"""A predicate that binds several values has to advance the counter by that many, or
|
||||
every later predicate reads the wrong parameter."""
|
||||
sql, params = where_sql(
|
||||
(
|
||||
Compare(field="created_by", op="eq", value="alice"),
|
||||
Within(field="budget_id", values=("a", "b", "c")),
|
||||
Compare(field="max_budget", op="gte", value=5.0),
|
||||
)
|
||||
)
|
||||
|
||||
assert sql == '"created_by" = $1 AND "budget_id" IN ($2, $3, $4) AND "max_budget" >= $5'
|
||||
assert params == ("alice", "a", "b", "c", 5.0)
|
||||
|
||||
|
||||
def test_placeholder_numbering_can_start_past_earlier_parameters():
|
||||
sql, params = where_sql((Compare(field="created_by", op="eq", value="alice"),), first_index=4)
|
||||
|
||||
assert sql == '"created_by" = $4'
|
||||
assert params == ("alice",)
|
||||
|
||||
|
||||
def test_is_null_binds_no_parameter_and_does_not_consume_a_placeholder():
|
||||
sql, params = where_sql(
|
||||
(IsNull(field="max_budget", negated=False), Compare(field="created_by", op="eq", value="alice"))
|
||||
)
|
||||
|
||||
assert sql == '"max_budget" IS NULL AND "created_by" = $1'
|
||||
assert params == ("alice",)
|
||||
|
||||
|
||||
def test_is_null_negated_renders_is_not_null():
|
||||
assert where_sql((IsNull(field="max_budget", negated=True),))[0] == '"max_budget" IS NOT NULL'
|
||||
|
||||
|
||||
def test_a_search_renders_as_a_parenthesised_or():
|
||||
"""Without the parentheses the OR would bind looser than the surrounding ANDs and the
|
||||
scope predicate would stop constraining the search branch."""
|
||||
sql, params = where_sql(
|
||||
(
|
||||
Compare(field="created_by", op="eq", value="alice"),
|
||||
AnyOf(
|
||||
clauses=(
|
||||
Compare(field="budget_id", op="contains", value="prod"),
|
||||
Compare(field="created_by", op="contains", value="prod"),
|
||||
)
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
assert sql == (
|
||||
'"created_by" = $1 AND ('
|
||||
"\"budget_id\" ILIKE $2 ESCAPE '\\'"
|
||||
" OR "
|
||||
"\"created_by\" ILIKE $3 ESCAPE '\\'"
|
||||
")"
|
||||
)
|
||||
assert params == ("alice", "%prod%", "%prod%")
|
||||
|
||||
|
||||
def test_contains_escapes_like_metacharacters():
|
||||
"""Budget ids routinely contain '_', which is a single-character wildcard unescaped."""
|
||||
_, params = where_sql((Compare(field="budget_id", op="contains", value="device_id%"),))
|
||||
|
||||
assert params == (r"%device\_id\%%",)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("op", "operator"),
|
||||
[("eq", "="), ("not", "<>"), ("gte", ">="), ("lte", "<="), ("gt", ">"), ("lt", "<")],
|
||||
)
|
||||
def test_each_comparison_operator_renders_its_sql_spelling(op, operator):
|
||||
assert where_sql((Compare(field="max_budget", op=op, value=1),))[0] == f'"max_budget" {operator} $1'
|
||||
|
||||
|
||||
def test_an_empty_plan_renders_no_where_body():
|
||||
assert where_sql(()) == ("", ())
|
||||
|
||||
|
||||
def test_a_planned_filter_renders_end_to_end():
|
||||
"""Ties the parser to the renderer: what build_query_plan produces is what executes."""
|
||||
sql, params = where_sql(_plan({"filter[max_budget][is_null]": "true", "q": "prod"}).where)
|
||||
|
||||
assert sql == (
|
||||
'"max_budget" IS NULL AND ('
|
||||
"\"budget_id\" ILIKE $1 ESCAPE '\\'"
|
||||
" OR "
|
||||
"\"created_by\" ILIKE $2 ESCAPE '\\'"
|
||||
")"
|
||||
)
|
||||
assert params == ("%prod%", "%prod%")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- invariant 3
|
||||
|
||||
|
||||
def test_the_scope_predicate_is_the_first_conjunct():
|
||||
spec = _spec(scope=lambda caller: ScopeWhere(where=(Compare(field="created_by", op="eq", value=caller.user_id),)))
|
||||
|
||||
conjuncts = _conjuncts(_plan({"filter[max_budget][gte]": "5"}, spec=spec))
|
||||
|
||||
assert conjuncts[0] == Compare(field="created_by", op="eq", value="caller-1")
|
||||
|
||||
|
||||
def test_a_caller_filter_cannot_replace_the_scope_predicate():
|
||||
"""The failure this guards is a `{**scope, **filters}` merge: a caller filtering on
|
||||
the scoped column would silently overwrite the scope and read another user's rows."""
|
||||
spec = _spec(scope=lambda caller: ScopeWhere(where=(Compare(field="created_by", op="eq", value=caller.user_id),)))
|
||||
|
||||
conjuncts = _conjuncts(_plan({"filter[created_by][eq]": "someone-else"}, spec=spec))
|
||||
|
||||
assert conjuncts[0] == Compare(field="created_by", op="eq", value="caller-1")
|
||||
assert Compare(field="created_by", op="eq", value="someone-else") in conjuncts
|
||||
assert len(conjuncts) == 2
|
||||
|
||||
|
||||
def test_the_scope_predicate_survives_a_search():
|
||||
spec = _spec(scope=lambda caller: ScopeWhere(where=(Compare(field="created_by", op="eq", value=caller.user_id),)))
|
||||
|
||||
conjuncts = _conjuncts(_plan({"q": "prod"}, spec=spec))
|
||||
|
||||
assert conjuncts[0] == Compare(field="created_by", op="eq", value="caller-1")
|
||||
assert any(isinstance(conjunct, AnyOf) for conjunct in conjuncts)
|
||||
|
||||
|
||||
def test_an_unscoped_caller_gets_no_scope_conjunct():
|
||||
assert _plan({"filter[max_budget][gte]": "5"}).where == (Compare(field="max_budget", op="gte", value=5.0),)
|
||||
|
||||
|
||||
def test_an_unfiltered_unscoped_list_has_an_empty_where():
|
||||
assert _plan({}).where == ()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- invariant 4
|
||||
|
||||
|
||||
def test_a_denied_scope_is_a_403_problem():
|
||||
spec = _spec(scope=lambda caller: ScopeDenied(reason="Only a proxy admin can list budgets."))
|
||||
|
||||
problem = _problem({}, spec=spec)
|
||||
|
||||
assert problem.status == 403
|
||||
assert problem.type == f"{PROBLEM_TYPE_BASE}forbidden"
|
||||
assert problem.detail == "Only a proxy admin can list budgets."
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_denied_scope_never_reaches_the_database():
|
||||
"""A 200 with an empty list would tell the caller the resource is empty rather than
|
||||
that they cannot read it, and would still pay for the query."""
|
||||
spec = _spec(scope=lambda caller: ScopeDenied(reason="nope"))
|
||||
executor = RecordingExecutor(rows=(BudgetRow(budget_id="b1", max_budget=None, created_by="x"),))
|
||||
|
||||
with pytest.raises(ManagementProblem) as raised:
|
||||
await handle_list(spec=spec, executor=executor, request=_request(), caller=CALLER)
|
||||
|
||||
assert raised.value.problem.status == 403
|
||||
assert executor.plan is None
|
||||
assert executor.count_where is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- invariant 5
|
||||
|
||||
|
||||
def test_page_size_falls_back_to_the_spec_default():
|
||||
assert _plan({}).take == 25
|
||||
|
||||
|
||||
def test_page_size_is_clamped_to_the_spec_maximum():
|
||||
"""Clamped rather than rejected: an over-large page is a UI bug, not a caller error,
|
||||
but serving it would let one request read the whole table."""
|
||||
assert _plan({"page_size": "100000"}).take == 100
|
||||
|
||||
|
||||
def test_page_offsets_by_page_size():
|
||||
plan = _plan({"page": "3", "page_size": "10"})
|
||||
|
||||
assert (plan.skip, plan.take) == (20, 10)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("page", ["0", "-1"], ids=["zero", "negative"])
|
||||
def test_page_below_one_is_rejected(page):
|
||||
problem = _problem({"page": page})
|
||||
|
||||
assert problem.status == 400
|
||||
assert problem.type == f"{PROBLEM_TYPE_BASE}invalid-query-parameter"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"params",
|
||||
[{"page": "one"}, {"page_size": "many"}, {"page_size": "0"}],
|
||||
ids=["page-not-an-int", "page-size-not-an-int", "page-size-zero"],
|
||||
)
|
||||
def test_unusable_paging_values_are_rejected(params):
|
||||
assert _problem(params).status == 400
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- invariant 6
|
||||
|
||||
|
||||
def test_an_unknown_query_parameter_is_rejected_with_the_allowed_set():
|
||||
problem = _problem({"page_sizee": "10"})
|
||||
|
||||
assert problem.status == 400
|
||||
assert problem.type == f"{PROBLEM_TYPE_BASE}unknown-query-parameter"
|
||||
assert "page_sizee" in problem.detail
|
||||
assert problem.allowed is not None
|
||||
assert "page_size" in problem.allowed
|
||||
assert "filter[max_budget][gte]" in problem.allowed
|
||||
|
||||
|
||||
def test_the_allowed_set_enumerates_only_operators_the_field_declares():
|
||||
problem = _problem({"nope": "1"})
|
||||
|
||||
assert problem.allowed is not None
|
||||
assert "filter[max_budget][is_null]" in problem.allowed
|
||||
assert "filter[tpm_limit][gte]" not in problem.allowed
|
||||
assert "filter[created_by][in]" in problem.allowed
|
||||
|
||||
|
||||
def test_a_filter_on_an_undeclared_field_is_an_unknown_parameter():
|
||||
problem = _problem({"filter[secret_column][eq]": "x"})
|
||||
|
||||
assert problem.type == f"{PROBLEM_TYPE_BASE}unknown-query-parameter"
|
||||
assert "filter[secret_column][eq]" in problem.detail
|
||||
|
||||
|
||||
def test_every_declared_parameter_is_accepted():
|
||||
"""Guards the unknown-param check against rejecting the spec's own contract."""
|
||||
plan = _plan(
|
||||
{
|
||||
"page": "2",
|
||||
"page_size": "10",
|
||||
"sort": "-max_budget",
|
||||
"q": "prod",
|
||||
"filter[max_budget][gte]": "5",
|
||||
"filter[created_by][in]": "a,b",
|
||||
}
|
||||
)
|
||||
|
||||
assert plan.take == 10
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- invariant 7
|
||||
|
||||
|
||||
def test_sorting_by_an_undeclared_field_is_rejected():
|
||||
problem = _problem({"sort": "api_key"})
|
||||
|
||||
assert problem.status == 400
|
||||
assert problem.type == f"{PROBLEM_TYPE_BASE}invalid-sort-field"
|
||||
assert problem.allowed == ["budget_id", "created_at", "max_budget"]
|
||||
assert "api_key" in problem.detail
|
||||
|
||||
|
||||
def test_one_bad_key_rejects_the_whole_multi_key_sort():
|
||||
"""Dropping the unknown key and sorting by the rest would silently return a
|
||||
differently-ordered page than the one asked for."""
|
||||
assert _problem({"sort": "-created_at,api_key"}).type == f"{PROBLEM_TYPE_BASE}invalid-sort-field"
|
||||
|
||||
|
||||
def test_a_double_dash_prefix_is_not_a_descending_sort():
|
||||
assert _problem({"sort": "--created_at"}).type == f"{PROBLEM_TYPE_BASE}invalid-sort-field"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- invariant 8
|
||||
|
||||
|
||||
def test_an_operator_the_field_does_not_declare_is_rejected():
|
||||
problem = _problem({"filter[max_budget][contains]": "5"})
|
||||
|
||||
assert problem.status == 400
|
||||
assert problem.type == f"{PROBLEM_TYPE_BASE}unsupported-filter-operator"
|
||||
assert problem.allowed == ["eq", "gte", "is_null", "lte"]
|
||||
assert "contains" in problem.detail
|
||||
|
||||
|
||||
def test_the_same_operator_is_accepted_on_a_field_that_declares_it():
|
||||
"""Pins the rejection to the field's own operator set rather than a global denylist."""
|
||||
conjuncts = _conjuncts(_plan({"filter[created_by][contains]": "ops"}))
|
||||
|
||||
assert conjuncts == (Compare(field="created_by", op="contains", value="ops"),)
|
||||
|
||||
|
||||
def test_a_string_that_is_not_an_operator_at_all_is_an_unknown_parameter():
|
||||
"""`gt3` is a typo, not an operator the field withheld, so the useful reply is the
|
||||
parameter list rather than this field's operator set."""
|
||||
problem = _problem({"filter[max_budget][gt3]": "5"})
|
||||
|
||||
assert problem.type == f"{PROBLEM_TYPE_BASE}unknown-query-parameter"
|
||||
assert problem.allowed is not None
|
||||
assert "filter[max_budget][gte]" in problem.allowed
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- invariant 9
|
||||
|
||||
|
||||
def test_search_against_a_spec_with_nothing_searchable_is_rejected():
|
||||
"""A silently-empty search filter returns the unfiltered table, which reads as
|
||||
"no results were filtered out" rather than "this resource cannot be searched"."""
|
||||
problem = _problem({"q": "prod"}, spec=_spec(searchable=frozenset()))
|
||||
|
||||
assert problem.status == 400
|
||||
assert problem.type == f"{PROBLEM_TYPE_BASE}unknown-query-parameter"
|
||||
assert problem.allowed is not None
|
||||
assert "q" not in problem.allowed
|
||||
|
||||
|
||||
def test_search_is_a_case_insensitive_or_across_every_searchable_field():
|
||||
conjuncts = _conjuncts(_plan({"q": "Prod"}))
|
||||
|
||||
assert conjuncts == (
|
||||
AnyOf(
|
||||
clauses=(
|
||||
Compare(field="budget_id", op="contains", value="Prod"),
|
||||
Compare(field="created_by", op="contains", value="Prod"),
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_an_empty_search_string_adds_no_filter():
|
||||
assert _plan({"q": ""}).where == ()
|
||||
|
||||
|
||||
# --------------------------------------------------------------- invariant 10
|
||||
|
||||
|
||||
def test_multi_key_sort_parses_the_json_api_grammar():
|
||||
order = _plan({"sort": "-created_at,budget_id,-max_budget"}).order
|
||||
|
||||
assert order[:3] == (
|
||||
SortKey(field="created_at", descending=True),
|
||||
SortKey(field="budget_id", descending=False),
|
||||
SortKey(field="max_budget", descending=True),
|
||||
)
|
||||
|
||||
|
||||
def test_sort_segments_tolerate_surrounding_whitespace():
|
||||
assert _plan({"sort": "-created_at, budget_id"}).order[:2] == (
|
||||
SortKey(field="created_at", descending=True),
|
||||
SortKey(field="budget_id", descending=False),
|
||||
)
|
||||
|
||||
|
||||
# ------------------------------------------------------------- filter parsing
|
||||
|
||||
|
||||
def test_comparison_operators_become_prisma_range_fragments():
|
||||
conjuncts = _conjuncts(_plan({"filter[max_budget][gte]": "5", "filter[max_budget][lte]": "50"}))
|
||||
|
||||
assert conjuncts == (
|
||||
Compare(field="max_budget", op="gte", value=5.0),
|
||||
Compare(field="max_budget", op="lte", value=50.0),
|
||||
)
|
||||
|
||||
|
||||
def test_eq_is_a_bare_value_not_a_wrapped_one():
|
||||
assert _conjuncts(_plan({"filter[tpm_limit][eq]": "100"})) == (Compare(field="tpm_limit", op="eq", value=100),)
|
||||
|
||||
|
||||
def test_a_filter_with_no_operator_bracket_means_eq():
|
||||
"""`filter[status]=active` is the design doc's canonical spelling for equality;
|
||||
only the non-eq operators carry a second bracket."""
|
||||
assert _conjuncts(_plan({"filter[tpm_limit]": "100"})) == (Compare(field="tpm_limit", op="eq", value=100),)
|
||||
|
||||
|
||||
def test_the_bare_form_and_the_explicit_eq_form_agree():
|
||||
assert _plan({"filter[created_by]": "alice"}) == _plan({"filter[created_by][eq]": "alice"})
|
||||
|
||||
|
||||
def test_the_bare_form_still_coerces_to_the_declared_type():
|
||||
assert _problem({"filter[tpm_limit]": "1.5"}).status == 400
|
||||
|
||||
|
||||
def test_the_bare_form_is_rejected_on_a_field_that_does_not_declare_eq():
|
||||
"""The shorthand is sugar for the eq operator, not a bypass around the operator set."""
|
||||
problem = _problem({"filter[created_at]": "2026-07-23T00:00:00Z"})
|
||||
|
||||
assert problem.type == f"{PROBLEM_TYPE_BASE}unsupported-filter-operator"
|
||||
assert problem.allowed == ["gte", "lte"]
|
||||
|
||||
|
||||
def test_the_allowed_set_advertises_the_bare_spelling_for_eq():
|
||||
allowed = _problem({"nope": "1"}).allowed
|
||||
|
||||
assert allowed is not None
|
||||
assert "filter[max_budget]" in allowed
|
||||
assert "filter[max_budget][eq]" not in allowed
|
||||
assert "filter[created_at][gte]" in allowed
|
||||
assert "filter[created_at]" not in allowed
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"name",
|
||||
["filter[]", "filter[a][b][c]", "filter[a][", "filter", "filter[a][gte", "filter[max_budget]]["],
|
||||
ids=["empty", "triple", "unbalanced", "bare-word", "unterminated", "bracketed-field"],
|
||||
)
|
||||
def test_malformed_filter_keys_are_unknown_parameters_not_eq_filters(name):
|
||||
"""A malformed key must not fall through to the bare-eq branch and silently filter
|
||||
on a field nobody declared. `field in spec.filters` is the gate that makes this hold,
|
||||
which is also why the parser needs no separate well-formedness guard."""
|
||||
assert _problem({name: "x"}).type == f"{PROBLEM_TYPE_BASE}unknown-query-parameter"
|
||||
|
||||
|
||||
def test_in_splits_on_commas_and_coerces_every_member():
|
||||
assert _conjuncts(_plan({"filter[created_by][in]": "alice, bob"})) == (
|
||||
Within(field="created_by", values=("alice", "bob")),
|
||||
)
|
||||
|
||||
|
||||
def test_is_null_true_matches_rows_with_no_budget():
|
||||
"""Budgets renders a null max_budget as "Unlimited"; without is_null there is no way
|
||||
to ask for those rows."""
|
||||
assert _conjuncts(_plan({"filter[max_budget][is_null]": "true"})) == (
|
||||
IsNull(field="max_budget", negated=False),
|
||||
)
|
||||
|
||||
|
||||
def test_is_null_false_matches_rows_that_have_one():
|
||||
assert _conjuncts(_plan({"filter[max_budget][is_null]": "false"})) == (
|
||||
IsNull(field="max_budget", negated=True),
|
||||
)
|
||||
|
||||
|
||||
def test_is_null_rejects_a_non_boolean():
|
||||
assert _problem({"filter[max_budget][is_null]": "maybe"}).status == 400
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"params",
|
||||
[
|
||||
{"filter[max_budget][gte]": "lots"},
|
||||
{"filter[tpm_limit][eq]": "1.5"},
|
||||
{"filter[created_at][gte]": "yesterday"},
|
||||
{"filter[created_by][in]": "alice,"},
|
||||
],
|
||||
ids=["float", "int", "datetime", "in-member"],
|
||||
)
|
||||
def test_a_value_that_does_not_match_the_declared_type_is_rejected(params):
|
||||
numeric_in = _spec_with(filters={**_spec().filters, "created_by": FilterSpec(type=int, ops=frozenset({"in"}))})
|
||||
target = numeric_in if "filter[created_by][in]" in params else _spec()
|
||||
|
||||
assert _problem(params, spec=target).status == 400
|
||||
|
||||
|
||||
def test_a_datetime_filter_is_normalised_to_utc():
|
||||
"""The dashboard sends both offset-bearing and naive timestamps; reading a naive one
|
||||
as server-local time would shift the window off what the table is showing."""
|
||||
with_offset = _conjuncts(_plan({"filter[created_at][gte]": "2026-07-23T02:00:00+02:00"}))
|
||||
naive = _conjuncts(_plan({"filter[created_at][gte]": "2026-07-23 00:00:00"}))
|
||||
|
||||
assert with_offset == (Compare(field="created_at", op="gte", value=datetime(2026, 7, 23, tzinfo=timezone.utc)),)
|
||||
assert naive == with_offset
|
||||
|
||||
|
||||
def test_filters_are_ordered_deterministically():
|
||||
"""Two requests differing only in query-string order must plan identically, or the
|
||||
plan stops being a comparable value."""
|
||||
forwards = _plan({"filter[created_by][eq]": "a", "filter[max_budget][gte]": "5"})
|
||||
backwards = _plan({"filter[max_budget][gte]": "5", "filter[created_by][eq]": "a"})
|
||||
|
||||
assert forwards == backwards
|
||||
|
||||
|
||||
# ------------------------------------------------------------------- envelope
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_the_page_mode_envelope():
|
||||
executor = RecordingExecutor(
|
||||
rows=tuple(BudgetRow(budget_id=f"b{i}", max_budget=float(i), created_by="u") for i in range(10)),
|
||||
total_count=42,
|
||||
)
|
||||
|
||||
response = await handle_list(
|
||||
spec=_spec(), executor=executor, request=_request("page=2&page_size=5"), caller=CALLER
|
||||
)
|
||||
body = response.model_dump(by_alias=True)
|
||||
|
||||
assert body["meta"] == {"total_count": 42, "page": 2, "page_size": 5, "total_pages": 9}
|
||||
assert set(body) == {"data", "meta", "links"}
|
||||
assert "has_more" not in body["meta"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_serializes_rows_flat_without_a_json_api_resource_wrapper():
|
||||
executor = RecordingExecutor(rows=(BudgetRow(budget_id="b1", max_budget=None, created_by="u"),))
|
||||
|
||||
response = await handle_list(spec=_spec(), executor=executor, request=_request(), caller=CALLER)
|
||||
body = response.model_dump(by_alias=True)
|
||||
|
||||
assert body["data"] == [{"budget_id": "b1", "max_budget": None}]
|
||||
assert "attributes" not in body["data"][0]
|
||||
assert "created_by" not in body["data"][0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_links_let_a_client_page_without_building_urls():
|
||||
executor = RecordingExecutor(rows=(), total_count=42)
|
||||
|
||||
response = await handle_list(
|
||||
spec=_spec(), executor=executor, request=_request("page=2&page_size=5"), caller=CALLER
|
||||
)
|
||||
links = response.model_dump(by_alias=True)["links"]
|
||||
|
||||
assert links["self"] == f"{BUDGETS_PATH}?page_size=5&page=2"
|
||||
assert links["first"] == f"{BUDGETS_PATH}?page_size=5&page=1"
|
||||
assert links["prev"] == f"{BUDGETS_PATH}?page_size=5&page=1"
|
||||
assert links["next"] == f"{BUDGETS_PATH}?page_size=5&page=3"
|
||||
assert links["last"] == f"{BUDGETS_PATH}?page_size=5&page=9"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_last_page_has_no_next_link():
|
||||
executor = RecordingExecutor(rows=(), total_count=10)
|
||||
|
||||
response = await handle_list(
|
||||
spec=_spec(), executor=executor, request=_request("page=2&page_size=5"), caller=CALLER
|
||||
)
|
||||
links = response.model_dump(by_alias=True)["links"]
|
||||
|
||||
assert links["next"] is None
|
||||
assert links["prev"] == f"{BUDGETS_PATH}?page_size=5&page=1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_an_empty_result_set_still_resolves_every_link():
|
||||
executor = RecordingExecutor(rows=(), total_count=0)
|
||||
|
||||
response = await handle_list(spec=_spec(), executor=executor, request=_request(), caller=CALLER)
|
||||
body = response.model_dump(by_alias=True)
|
||||
|
||||
assert body["data"] == []
|
||||
assert body["meta"]["total_pages"] == 0
|
||||
assert body["links"]["first"] == body["links"]["last"] == f"{BUDGETS_PATH}?page=1"
|
||||
assert body["links"]["next"] is None
|
||||
assert body["links"]["prev"] is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_executor_counts_the_same_predicate_it_reads():
|
||||
"""Counting a wider predicate than the read inflates total_pages and hands the UI
|
||||
pages that are always empty."""
|
||||
executor = RecordingExecutor(rows=(), total_count=3)
|
||||
spec = _spec(scope=lambda caller: ScopeWhere(where=(Compare(field="created_by", op="eq", value=caller.user_id),)))
|
||||
|
||||
await handle_list(spec=spec, executor=executor, request=_request("filter[max_budget][gte]=5"), caller=CALLER)
|
||||
|
||||
assert executor.plan is not None
|
||||
assert executor.count_where == executor.plan.where
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_rejected_request_is_raised_as_a_problem_before_any_query():
|
||||
executor = RecordingExecutor(rows=())
|
||||
|
||||
with pytest.raises(ManagementProblem) as raised:
|
||||
await handle_list(spec=_spec(), executor=executor, request=_request("sort=api_key"), caller=CALLER)
|
||||
|
||||
assert raised.value.problem.status == 400
|
||||
assert executor.count_where is None
|
||||
|
||||
|
||||
# ------------------------------------------------------- spec construction
|
||||
|
||||
|
||||
def test_a_default_page_size_above_the_cap_is_rejected_at_construction():
|
||||
"""The cap is only enforced on a supplied page_size, so a default above it would serve
|
||||
more rows than the resource allows on exactly the request that omits page_size."""
|
||||
with pytest.raises(ValueError, match="default_page_size"):
|
||||
_spec_with(default_page_size=200, max_page_size=100)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"overrides",
|
||||
[{"default_page_size": 0}, {"default_page_size": -5}, {"max_page_size": 0}],
|
||||
ids=["zero-default", "negative-default", "zero-cap"],
|
||||
)
|
||||
def test_a_non_positive_page_size_is_rejected_at_construction(overrides):
|
||||
"""take=0 divides by zero when handle_list computes total_pages, so the resource would
|
||||
500 on every request instead of failing when it is registered."""
|
||||
with pytest.raises(ValueError, match="default_page_size"):
|
||||
_spec_with(**overrides)
|
||||
|
||||
|
||||
def test_a_default_sort_on_a_non_sortable_field_is_rejected_at_construction():
|
||||
"""Caller-supplied sort is validated against `sortable`; default_sort is not read from
|
||||
the request, so without this it reaches order_by_sql and yields invalid SQL."""
|
||||
with pytest.raises(ValueError, match="default_sort"):
|
||||
_spec_with(default_sort=(SortKey(field="not_a_column", descending=True),))
|
||||
|
||||
|
||||
def test_an_empty_tiebreaker_is_rejected_at_construction():
|
||||
with pytest.raises(ValueError, match="tiebreaker"):
|
||||
_spec_with(tiebreaker="")
|
||||
|
||||
|
||||
def test_a_page_size_equal_to_the_cap_is_a_valid_spec():
|
||||
"""Guards the bound against being tightened into an off-by-one that bans max==default."""
|
||||
assert _spec_with(default_page_size=100, max_page_size=100).default_page_size == 100
|
||||
|
||||
|
||||
# ------------------------------------------------------ repeated parameters
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_repeated_query_parameter_is_rejected():
|
||||
"""Starlette keeps the last value, so ?page=1&page=999 would page from 999 with nothing
|
||||
telling the caller which one won. The doc rejects silently-altered params for this reason."""
|
||||
executor = RecordingExecutor(rows=())
|
||||
|
||||
with pytest.raises(ManagementProblem) as raised:
|
||||
await handle_list(spec=_spec(), executor=executor, request=_request("page=1&page=999"), caller=CALLER)
|
||||
|
||||
assert raised.value.problem.status == 400
|
||||
assert raised.value.problem.type == f"{PROBLEM_TYPE_BASE}duplicate-query-parameter"
|
||||
assert "page" in raised.value.problem.detail
|
||||
assert executor.count_where is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_repeated_filter_parameter_is_rejected():
|
||||
executor = RecordingExecutor(rows=())
|
||||
|
||||
with pytest.raises(ManagementProblem) as raised:
|
||||
await handle_list(
|
||||
spec=_spec(),
|
||||
executor=executor,
|
||||
request=_request("filter[created_by][eq]=alice&filter[created_by][eq]=bob"),
|
||||
caller=CALLER,
|
||||
)
|
||||
|
||||
assert raised.value.problem.type == f"{PROBLEM_TYPE_BASE}duplicate-query-parameter"
|
||||
assert "filter[created_by][eq]" in raised.value.problem.detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_distinct_parameters_are_not_treated_as_duplicates():
|
||||
"""Guards the check against rejecting two different operators on one field, which is
|
||||
how a range filter is expressed."""
|
||||
executor = RecordingExecutor(rows=(), total_count=0)
|
||||
|
||||
response = await handle_list(
|
||||
spec=_spec(),
|
||||
executor=executor,
|
||||
request=_request("filter[max_budget][gte]=5&filter[max_budget][lte]=50&page=2"),
|
||||
caller=CALLER,
|
||||
)
|
||||
|
||||
assert response.meta.page == 2
|
||||
assert executor.count_where is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_denied_scope_outranks_a_duplicate_parameter():
|
||||
"""Permission is the stronger statement about the caller, so it is answered first."""
|
||||
spec = _spec(scope=lambda caller: ScopeDenied(reason="nope"))
|
||||
executor = RecordingExecutor(rows=())
|
||||
|
||||
with pytest.raises(ManagementProblem) as raised:
|
||||
await handle_list(spec=spec, executor=executor, request=_request("page=1&page=2"), caller=CALLER)
|
||||
|
||||
assert raised.value.problem.status == 403
|
||||
|
||||
|
||||
# --------------------------------------------------- facet-mode regression
|
||||
|
||||
|
||||
def test_the_facet_page_shapes_are_untouched_by_page_mode():
|
||||
"""The live facet endpoint reports `has_more` and has no first/last, because it
|
||||
deliberately skips the COUNT(*). Folding it into the page-mode shapes would either
|
||||
break its response or make every keystroke pay for a full-table count."""
|
||||
assert set(PageMeta.model_fields) == {"page", "page_size", "has_more"}
|
||||
assert set(PageLinks.model_fields) == {"self_link", "prev", "next"}
|
||||
|
||||
links = build_page_links(request=_request("q=ac&page=2"), page=2, has_more=True).model_dump(by_alias=True)
|
||||
|
||||
assert set(links) == {"self", "prev", "next"}
|
||||
assert links["next"] == "/management/v1/budgets?q=ac&page=3"
|
||||
Loading…
Add table
Reference in a new issue