From 473f43dfbf08797c2239c645147bd722030373be Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Fri, 31 Jul 2026 08:49:22 -0700 Subject: [PATCH 1/3] fix(mcp): deny MCP access when a named entitlement cannot be read (#35160) An MCP permission level answers which servers and tools it permits, and a level that answers nothing places no restriction. Key auth was reading a lookup FAULT as that same answer, so the end user, agent and org ceilings quietly disappeared for as long as one lasted, while the keyless gateway-admitted path failed closed on the very same fault. Those levels now separate the two fault classes the user level already did. A principal row that NAMES an object_permission_id whose contents cannot be read is a known entitlement with unknown contents, so it denies. A lookup that fails before we can tell whether the principal is entitled at all still places no ceiling, that being the state which existed before the level did; denying there would refuse MCP to the majority of callers, who have no such entitlement configured. The keyless path is unchanged. Resolves LIT-4960 --- .../mcp_server/auth/user_api_key_auth_mcp.py | 337 ++++++++++++------ .../auth/test_user_api_key_auth_mcp.py | 174 +++++++++ 2 files changed, 404 insertions(+), 107 deletions(-) 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/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"} From 416e39815474b35acdfa02d18f619d84f8101582 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Fri, 31 Jul 2026 09:26:30 -0700 Subject: [PATCH 2/3] feat(proxy): add generic list handler for /management/v1 (#35308) * feat(proxy): add generic list handler for /management/v1 Adds the ListSpec/QueryPlan machinery the control-plane list endpoints are meant to share, so a resource declares what it exposes instead of hand-rolling its own paging, sorting and filter parsing. build_query_plan is pure: it turns query parameters into a QueryPlan or an RFC 9457 problem without any I/O, which is what lets the plan be asserted as a value. The database half is a ListExecutor protocol injected by the caller, so this module has no Prisma dependency at all. Four things the framework guarantees rather than leaving to each resource: the spec's unique tiebreaker is always the final sort key, so pages cannot repeat rows when the leading column is all nulls; ordering is NULLS LAST in both directions, since Postgres otherwise floats empty values to the top the moment the sort direction flips; the scope predicate is a separate conjunct ahead of every caller filter, so a filter on a scoped column cannot widen it; and a denied scope is a 403 problem rather than a 200 with an empty list. No route and no consumer yet; budgets registers against it next. The facet endpoint's has_more shapes are untouched, and a test pins them so page mode cannot quietly absorb them. * fix(proxy): accept the bare filter[field] form in the list framework Section 5 of the design doc spells equality without an operator bracket (`?filter[status]=active`, and `/management/v1/keys?filter[team_id]=` in the sub-resource paragraph); only the other operators carry a second bracket. The parser only understood `filter[field][op]`, so the canonical spelling came back as an unknown query parameter. `filter[field]` now resolves to the field's `eq` operator, which means it still goes through the declared operator set rather than around it: a field that does not offer `eq` rejects the shorthand. The allowed-parameter list advertises the bare spelling for `eq` and the bracketed one for everything else. Drops two guards from the key parser that could not fire. Operator validation already rejects every malformed operator, and `field in spec.filters` already rejects every field nobody declared, so a well-formedness check on top of them was unreachable; the tests cover the malformed keys directly instead. * fix(proxy): validate list specs at construction and reject repeated params Two gaps a review flagged on the list framework. The page-size cap was only enforced against a supplied page_size, so a spec whose default_page_size exceeded its max_page_size served more rows than the resource allows on exactly the request that omits the parameter. A default of zero was worse: it reached the total_pages division and made the resource 500 on every request. ListSpec now validates 1 <= default_page_size <= max_page_size when it is built, so a misconfigured resource fails as it is registered rather than per request. default_sort is checked against sortable for the same reason; caller-supplied sort was already validated, but the default never passed through that path and a typo there reached the ORDER BY clause untouched. Raising is right here despite the usual model-failures-as-values rule: there is no request in flight and no caller to answer. Repeated query parameters silently collapsed to their last value, so ?page=1&page=999 paged from 999 and a repeated sort key quietly won, which is the same silently-altered-semantics failure the surface already rejects unknown parameters to avoid. They are now a 400. The check lives in handle_list rather than build_query_plan because a Mapping[str, str] cannot represent a repeat at all; the boundary that can see one is the boundary that rejects it. A denied scope still outranks it, matching every other rejection here. Also corrects the order_by_sql docstring, which claimed every field reaching it had been validated against sortable. That held for caller-supplied sort only. * refactor(proxy): model list predicates as frozen values instead of dicts The LIT002 budget rejected the framework: building a where-fragment meant a dict literal per operator, and a dict keyed by a column name chosen at runtime cannot be frozen into a TypedDict or a dataclass field, so there was no spelling of the old shape the rule would accept. Replacing the fragments with a tagged union removes the construction entirely. A plan's where is now a tuple of frozen Compare / Within / IsNull / AnyOf, matched exhaustively, and the field name is a value rather than a key. That also retires the Mapping[str, object] the plan used to carry, which said nothing about what was inside it and left the fragment shape as a convention two sides had to keep agreeing on. Scope predicates take the same type, so a resource declares its row filter in the same vocabulary rather than hand-rolling a backend dict. where_sql renders a plan for a raw-SQL executor, binding every caller-supplied value to a numbered placeholder and writing only spec-declared column names into the statement. It is the counterpart to order_by_sql, which already existed for the same reason: nulls ordering forces the executor onto raw SQL, so the escaping and placeholder arithmetic belong in one reviewed place rather than in each consumer. Also folds the two remaining mutable builds out of the module (set comprehensions and Counter to frozenset/tuple, the serialized page to a tuple pydantic coerces), and lifts the LIKE escaper into common.py so the facet endpoint and the framework share one copy instead of two that can drift. No behavioural change to the facet endpoint; its tests, including the one pinning the escaping, pass untouched. --- .../management_v1/common.py | 38 +- .../management_v1/list_framework.py | 522 +++++++++++ .../management_v1/spend_logs.py | 7 +- .../management_endpoints/management_v1.py | 34 + .../management_v1/test_list_framework.py | 871 ++++++++++++++++++ 5 files changed, 1458 insertions(+), 14 deletions(-) create mode 100644 litellm/proxy/management_endpoints/management_v1/list_framework.py create mode 100644 tests/test_litellm/proxy/management_endpoints/management_v1/test_list_framework.py 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/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" From 16507f11742144aa6c66c8239a05279b992cd597 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Fri, 31 Jul 2026 09:48:08 -0700 Subject: [PATCH 3/3] fix(aiohttp): dispose recycled client sessions deterministically (#33428) * fix(aiohttp): dispose recycled client sessions deterministically LiteLLMAiohttpTransport replaced its cached aiohttp.ClientSession on loop-mismatch, loop-inspection failure, and "Session is closed" retry without reliably closing the previous session: - the close task from asyncio.create_task() was never referenced, so it could be garbage-collected before running; - the (RuntimeError, AttributeError) fallback branch replaced the session without closing it at all; - sessions bound to a closed event loop were abandoned to the GC ("rely on GC"), and sessions bound to a loop running in another thread were closed from the wrong loop. Replaced sessions surfaced as intermittent "Unclosed client session" / "Unclosed connector" errors from the event-loop exception handler at GC time. _close_recycled_session() now covers the three lifecycles a recycled session can be in: same-loop closes keep a strong task reference until completion; sessions owned by a loop running elsewhere are closed on their own loop via run_coroutine_threadsafe; sessions whose loop is gone are disposed synchronously through the connector teardown that aiohttp's own finalizer uses, which releases pooled connections and silences the finalizer warnings. Fixes #24230 * fix(aiohttp): guard threadsafe close callback against cancelled futures --------- Co-authored-by: Anmol Jaiswal <68013660+anmolg1997@users.noreply.github.com> --- .../llms/custom_httpx/aiohttp_transport.py | 123 ++++++- .../custom_httpx/test_aiohttp_transport.py | 303 ++++++++++++++++++ 2 files changed, 416 insertions(+), 10 deletions(-) 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/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()