diff --git a/litellm/llms/custom_httpx/aiohttp_transport.py b/litellm/llms/custom_httpx/aiohttp_transport.py index df5b10b3bdc..2c5f455692c 100644 --- a/litellm/llms/custom_httpx/aiohttp_transport.py +++ b/litellm/llms/custom_httpx/aiohttp_transport.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index dc796837277..81983cc62fd 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -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 diff --git a/litellm/proxy/management_endpoints/management_v1/common.py b/litellm/proxy/management_endpoints/management_v1/common.py index c0e7f49f2e9..daa2c60ac5e 100644 --- a/litellm/proxy/management_endpoints/management_v1/common.py +++ b/litellm/proxy/management_endpoints/management_v1/common.py @@ -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), + ) diff --git a/litellm/proxy/management_endpoints/management_v1/list_framework.py b/litellm/proxy/management_endpoints/management_v1/list_framework.py new file mode 100644 index 00000000000..3e4b9131d1e --- /dev/null +++ b/litellm/proxy/management_endpoints/management_v1/list_framework.py @@ -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 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), + ) diff --git a/litellm/proxy/management_endpoints/management_v1/spend_logs.py b/litellm/proxy/management_endpoints/management_v1/spend_logs.py index c11a14bbfea..ccde3c4112c 100644 --- a/litellm/proxy/management_endpoints/management_v1/spend_logs.py +++ b/litellm/proxy/management_endpoints/management_v1/spend_logs.py @@ -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( diff --git a/litellm/types/proxy/management_endpoints/management_v1.py b/litellm/types/proxy/management_endpoints/management_v1.py index 2aecc54f114..b2244f6eb9b 100644 --- a/litellm/types/proxy/management_endpoints/management_v1.py +++ b/litellm/types/proxy/management_endpoints/management_v1.py @@ -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 diff --git a/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py b/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py index 0550898c73d..b0b092a541f 100644 --- a/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py +++ b/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py @@ -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() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 89b8f018e5c..0b95a497882 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -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"} diff --git a/tests/test_litellm/proxy/management_endpoints/management_v1/test_list_framework.py b/tests/test_litellm/proxy/management_endpoints/management_v1/test_list_framework.py new file mode 100644 index 00000000000..35bd5517361 --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/management_v1/test_list_framework.py @@ -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"