diff --git a/.github/workflows/codspeed.yml b/.github/workflows/codspeed.yml index 54a8e53d7a3..a69e50b5753 100644 --- a/.github/workflows/codspeed.yml +++ b/.github/workflows/codspeed.yml @@ -5,10 +5,24 @@ on: branches: - main - litellm_internal_staging + paths: + - "litellm/**" + - "tests/benchmarks/**" + - "pyproject.toml" + - "uv.lock" + - ".github/workflows/codspeed.yml" + - ".github/actions/setup-uv-with-retries/**" pull_request: branches: - main - litellm_internal_staging + paths: + - "litellm/**" + - "tests/benchmarks/**" + - "pyproject.toml" + - "uv.lock" + - ".github/workflows/codspeed.yml" + - ".github/actions/setup-uv-with-retries/**" # Allow CodSpeed to trigger backtest performance analysis # in order to generate initial data workflow_dispatch: diff --git a/litellm/__init__.py b/litellm/__init__.py index 55821012df9..3f8c742c5a2 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -211,6 +211,9 @@ filter_invalid_headers: Optional[bool] = False add_user_information_to_llm_headers: Optional[bool] = ( None # adds user_id, team_id, token hash (params from StandardLoggingMetadata) to request headers ) +overwrite_user_with_key_hash: bool = ( + False # force the outgoing `user` param to the hashed api key, so providers see a stable, tamper-proof id +) store_audit_logs = False # Enterprise feature, allow users to see audit logs skip_system_message_in_guardrail: bool = False skip_tool_message_in_guardrail: bool = False diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index a40a8e1389c..96aed20529f 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -2191,6 +2191,13 @@ def batch_cost_calculator( return total_prompt_cost, total_completion_cost +def _summable_prompt_token_fields(prompt_tokens_details: BaseModel) -> List[str]: + field_names = list(type(prompt_tokens_details).model_fields) + if getattr(prompt_tokens_details, "cache_write_tokens", None) is None: + return field_names + return [attr for attr in field_names if attr != "cache_creation_tokens"] + + class BaseTokenUsageProcessor: @staticmethod def combine_usage_objects(usage_objects: List[Usage]) -> Usage: @@ -2225,7 +2232,7 @@ class BaseTokenUsageProcessor: # Check what keys exist in the model's prompt_tokens_details # Access model_fields on the class, not the instance, to avoid Pydantic 2.11+ deprecation warnings - for attr in type(usage.prompt_tokens_details).model_fields: + for attr in _summable_prompt_token_fields(usage.prompt_tokens_details): if ( hasattr(usage.prompt_tokens_details, attr) and not attr.startswith("_") diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index ce0c8285221..f9d887a19b0 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -470,8 +470,8 @@ def _parse_prompt_tokens_details(usage: Usage) -> PromptTokensDetailsResult: cache_creation_tokens = ( cast( Optional[int], - getattr(usage.prompt_tokens_details, "cache_creation_tokens", 0) - or getattr(usage.prompt_tokens_details, "cache_write_tokens", 0), + getattr(usage.prompt_tokens_details, "cache_write_tokens", 0) + or getattr(usage.prompt_tokens_details, "cache_creation_tokens", 0), ) or 0 ) @@ -920,10 +920,6 @@ def get_token_type_cost_breakdown( cache_read_tokens = prompt_tokens_details["cache_hit_tokens"] cache_creation_tokens = prompt_tokens_details["cache_creation_tokens"] cache_creation_token_details = prompt_tokens_details["cache_creation_token_details"] - # Some OpenAI-compatible providers (e.g. kimi-k2) report cache-write tokens - # under `cache_write_tokens`; mirror the total-cost normalization path. - if not cache_creation_tokens: - cache_creation_tokens = _coerce_token_count(getattr(usage.prompt_tokens_details, "cache_write_tokens", 0)) # Fall back to the private top-level counters the Usage constructor mirrors cache # tokens onto, so providers/callers that bypass prompt_tokens_details are covered. if not cache_read_tokens: diff --git a/litellm/proxy/_experimental/mcp_server/auth/token_exchange.py b/litellm/proxy/_experimental/mcp_server/auth/token_exchange.py deleted file mode 100644 index cd41dd648ee..00000000000 --- a/litellm/proxy/_experimental/mcp_server/auth/token_exchange.py +++ /dev/null @@ -1,192 +0,0 @@ -""" -OAuth 2.0 Token Exchange (RFC 8693) handler for MCP servers. - -Exchanges a user's incoming JWT (subject_token) for a scoped access token -at an IDP's token exchange endpoint. The exchanged token is then used to -authenticate requests to the upstream MCP server. - -See: https://datatracker.ietf.org/doc/html/rfc8693 -""" - -import asyncio -import hashlib -import weakref -from typing import TYPE_CHECKING, Dict, Tuple - -import httpx - -from litellm._logging import verbose_logger -from litellm.caching.in_memory_cache import InMemoryCache -from litellm.constants import ( - MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL, - MCP_OAUTH2_TOKEN_CACHE_MIN_TTL, - MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS, - MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE, -) -from litellm.llms.custom_httpx.http_handler import get_async_httpx_client -from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import ( - build_token_endpoint_client_auth, -) -from litellm.types.llms.custom_http import httpxSpecialProvider -from litellm.types.mcp import DEFAULT_SUBJECT_TOKEN_TYPE - -if TYPE_CHECKING: - from litellm.types.mcp_server.mcp_server_manager import MCPServer - -# RFC 8693 grant type constant -TOKEN_EXCHANGE_GRANT_TYPE = "urn:ietf:params:oauth:grant-type:token-exchange" - - -class TokenExchangeHandler: - """Handles OAuth 2.0 Token Exchange (RFC 8693) for MCP servers. - - Caches exchanged tokens keyed by ``hash(subject_token + server_id)`` so - repeated calls with the same user token skip the IDP round-trip. - """ - - def __init__(self) -> None: - self._cache = InMemoryCache( - max_size_in_memory=MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE, - default_ttl=MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL, - ) - # WeakValueDictionary so locks are GC'd once no coroutine holds a reference, - # preventing unbounded growth with many rotating user tokens. - self._locks: weakref.WeakValueDictionary[str, asyncio.Lock] = weakref.WeakValueDictionary() - - def _get_lock(self, cache_key: str) -> asyncio.Lock: - lock = self._locks.get(cache_key) - if lock is None: - lock = asyncio.Lock() - self._locks[cache_key] = lock - return lock - - @staticmethod - def _cache_key(subject_token: str, server_id: str) -> str: - raw = f"{subject_token}:{server_id}" - return hashlib.sha256(raw.encode()).hexdigest() - - async def exchange_token( - self, - subject_token: str, - server: "MCPServer", - ) -> str: - """Exchange *subject_token* for a scoped access token. - - Returns the exchanged ``access_token`` string (suitable for a - ``Bearer`` header). - - Raises ``ValueError`` on configuration or IDP errors. - """ - cache_key = self._cache_key(subject_token, server.server_id) - - # Fast path - cached = self._cache.get_cache(cache_key) - if cached is not None: - return cached - - # Slow path — one exchange at a time per (user, server) pair - async with self._get_lock(cache_key): - cached = self._cache.get_cache(cache_key) - if cached is not None: - return cached - - token, ttl = await self._do_exchange(subject_token, server) - self._cache.set_cache(cache_key, token, ttl=ttl) - return token - - async def _do_exchange( - self, - subject_token: str, - server: "MCPServer", - ) -> Tuple[str, int]: - """POST to the token exchange endpoint with RFC 8693 parameters. - - Returns ``(access_token, ttl_seconds)``. - """ - endpoint = server.token_exchange_endpoint or server.token_url - if not endpoint: - raise ValueError( - f"MCP server '{server.server_id}' has auth_type=oauth2_token_exchange " - f"but no token_exchange_endpoint or token_url configured" - ) - if not server.client_id or not server.client_secret: - raise ValueError( - f"MCP server '{server.server_id}' has auth_type=oauth2_token_exchange " - f"but missing client_id or client_secret" - ) - - client_auth = build_token_endpoint_client_auth( - auth_method=server.token_endpoint_auth_method, - client_id=server.client_id, - client_secret=server.client_secret, - ) - data: Dict[str, str] = { - "grant_type": TOKEN_EXCHANGE_GRANT_TYPE, - "subject_token": subject_token, - "subject_token_type": server.subject_token_type or DEFAULT_SUBJECT_TOKEN_TYPE, - **client_auth.body, - } - if server.audience: - data["audience"] = server.audience - if server.scopes: - data["scope"] = " ".join(server.scopes) - - verbose_logger.debug( - "Exchanging token for MCP server %s at %s (audience=%s)", - server.server_id, - endpoint, - server.audience, - ) - - client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) - post_kwargs = {"data": data, **({"headers": client_auth.headers} if client_auth.headers else {})} - try: - response = await client.post(endpoint, **post_kwargs) - response.raise_for_status() - except httpx.HTTPStatusError as exc: - verbose_logger.debug( - "Token exchange IDP error for MCP server %s (status %d)", - server.server_id, - exc.response.status_code, - ) - raise ValueError( - f"Token exchange for MCP server '{server.server_id}' failed with status {exc.response.status_code}" - ) from exc - - body = response.json() - if not isinstance(body, dict): - raise ValueError( - f"Token exchange response for MCP server '{server.server_id}' " - f"returned non-object JSON (got {type(body).__name__})" - ) - - access_token = body.get("access_token") - if not access_token: - raise ValueError(f"Token exchange response for MCP server '{server.server_id}' missing 'access_token'") - - raw_expires_in = body.get("expires_in") - try: - expires_in = int(raw_expires_in) if raw_expires_in is not None else MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL - except (TypeError, ValueError): - expires_in = MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL - - ttl = max( - expires_in - MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS, - MCP_OAUTH2_TOKEN_CACHE_MIN_TTL, - ) - - verbose_logger.info( - "Token exchange succeeded for MCP server %s (expires in %ds)", - server.server_id, - expires_in, - ) - return access_token, ttl - - def invalidate(self, subject_token: str, server_id: str) -> None: - """Remove a cached exchanged token (e.g. after a 401).""" - cache_key = self._cache_key(subject_token, server_id) - self._cache.delete_cache(cache_key) - - -# Module-level singleton -mcp_token_exchange_handler = TokenExchangeHandler() 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 f1fcc95c532..a27d6b92843 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 @@ -25,6 +25,9 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credenti from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( EnvelopeIdentity, ) +from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import ( + is_session_bearer_shaped, +) from litellm.proxy._types import ( UI_TEAM_ID, LiteLLM_TeamTable, @@ -124,6 +127,27 @@ def _has_client_supplied_mcp_auth( return bool(mcp_auth_header) or bool(mcp_server_auth_headers) +def _is_mcp_admitted_user_subject(user_api_key_auth: UserAPIKeyAuth | None) -> bool: + """True when this auth is a keyless subject admitted by the gateway session / bridge user + path, as opposed to a JWT or other keyless auth that merely lacks a ``team_id``. + + Reads the server-only ``mcp_admitted_user_subject`` field, set only by ``_reload_admitted_user``. It + is deliberately NOT a ``metadata`` key, which is caller-controlled at key creation and so forgeable + on a personal key to gain the team grant union or dodge the egress scrub; this field cannot be.""" + return user_api_key_auth is not None and user_api_key_auth.mcp_admitted_user_subject is True + + +def _is_aggregate_mcp_scope(route: str, mcp_servers: list[str] | None) -> bool: + """True when a request targets the aggregate ``/mcp`` endpoint rather than any named + server. Named targets arrive either through ``x-mcp-servers`` (``mcp_servers``) or a + path segment (``/mcp/{server}`` / ``/{server}/mcp``); the aggregate scope has neither. + The gateway-DCR session arm and challenge fire only here, so a per-server flow is never + affected.""" + if mcp_servers: + return False + return len(MCPRequestHandler._extract_target_server_names_from_path(route)) == 0 + + def _is_aggregate_gateway_dcr_challenge_scope( route: str, mcp_servers: list[str] | None, @@ -141,11 +165,9 @@ def _is_aggregate_gateway_dcr_challenge_scope( client. Fails closed to the original admission error otherwise.""" if not _is_litellm_auth_admission_error(exc): return False - if mcp_servers: - return False if _has_client_supplied_mcp_auth(mcp_auth_header, mcp_server_auth_headers): return False - return len(MCPRequestHandler._extract_target_server_names_from_path(route)) == 0 + return _is_aggregate_mcp_scope(route, mcp_servers) def _aggregate_gateway_dcr_challenge(request: Request, invalid_token: bool) -> HTTPException: @@ -362,6 +384,19 @@ class MCPRequestHandler: request=request, route=request_route, ) + elif ( + _is_aggregate_mcp_scope(request_route, mcp_servers) + and oauth2_headers + and is_session_bearer_shaped(oauth2_headers["Authorization"]) + ): + # A gateway DCR session bearer at the aggregate /mcp scope: open the identity-only session + # token and admit under the live litellm user. One that does not open fails closed with the + # aggregate invalid_token challenge; a non-session bearer falls through to the oauth2 arm. + validated_user_api_key_auth = await MCPRequestHandler._admit_gateway_session( + authorization_value=oauth2_headers["Authorization"], + request=request, + route=request_route, + ) elif oauth2_headers: # Authorization on a non-delegated server: the bearer must be a real # LiteLLM credential, so a failed validation is a genuine 401/403 and @@ -392,15 +427,80 @@ class MCPRequestHandler: bearer_presented=False, ) + # Leak-defense (single chokepoint): a gateway admission credential (session bearer or bridge + # envelope) is NEVER a valid upstream token. Scrub it from EVERY egress context so no + # client-forwarded, OBO, or passthrough path can send it upstream for replay. Anchored to the + # credential SHAPE, so a legitimate upstream/passthrough token is forwarded unchanged. + raw_headers = dict(headers) + ( + oauth2_headers, + raw_headers, + mcp_auth_header, + mcp_server_auth_headers, + ) = MCPRequestHandler._scrub_gateway_admission_credentials( + admitted=_is_mcp_admitted_user_subject(validated_user_api_key_auth), + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + ) + return ( validated_user_api_key_auth, mcp_auth_header, mcp_servers, mcp_server_auth_headers, oauth2_headers, - dict(headers), + raw_headers, ) + @staticmethod + def _is_gateway_admission_credential(value: str | None) -> bool: + """True when a header value is a gateway admission credential — a session bearer or bridge + envelope. It proves who signed in to the GATEWAY, never a valid UPSTREAM token, so it must never + be forwarded (a hostile upstream could capture and replay it against the aggregate ``/mcp`` scope).""" + return value is not None and (is_session_bearer_shaped(value) or is_bridge_envelope_shaped(value)) + + @staticmethod + def _scrub_gateway_admission_credentials( + admitted: bool, + oauth2_headers: dict[str, str] | None, + raw_headers: dict[str, str], + mcp_auth_header: str | None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None, + ) -> tuple[dict[str, str] | None, dict[str, str], str | None, dict[str, dict[str, str]] | None]: + """Remove any gateway admission credential from EVERY egress header context, keyed on the credential + SHAPE: top-level ``Authorization`` (oauth2 + raw), the deprecated ``x-mcp-auth``, and per-server + ``x-mcp-{alias}-authorization``. A legitimate upstream/passthrough token is never gateway-shaped so + it survives (including the real upstream token the bridge arm injects per-server); an admitted + subject's top-level Authorization is dropped unconditionally as defense-in-depth.""" + cred = MCPRequestHandler._is_gateway_admission_credential + + # 1. Top-level Authorization → oauth2_headers. + authz = oauth2_headers.get("Authorization") if oauth2_headers else None + if admitted or cred(authz): + oauth2_headers = None + + # 2. raw_headers: drop the admitted subject's Authorization, and ANY header whose value is a + # gateway credential (covers x-mcp-auth and x-mcp-{alias}-authorization in their raw form). + raw_headers = { + k: v for k, v in raw_headers.items() if not ((admitted and k.lower() == "authorization") or cred(v)) + } + + # 3. Deprecated x-mcp-auth value. + if cred(mcp_auth_header): + mcp_auth_header = None + + # 4. Per-server x-mcp-{alias}-authorization values (drop the value, then any now-empty server dict). + if mcp_server_auth_headers: + stripped = { + alias: {h: val for h, val in hdrs.items() if not cred(val)} + for alias, hdrs in mcp_server_auth_headers.items() + } + mcp_server_auth_headers = {alias: hdrs for alias, hdrs in stripped.items() if hdrs} + + return oauth2_headers, raw_headers, mcp_auth_header, mcp_server_auth_headers + @staticmethod def _extract_target_server_names_from_path(path: str) -> List[str]: """ @@ -626,6 +726,62 @@ class MCPRequestHandler: case _: assert_never(result) + @staticmethod + async def _admit_gateway_session( + authorization_value: str, + request: Request, + route: str, + ) -> UserAPIKeyAuth: + """Open a gateway DCR session bearer and admit the live litellm user it references. + + Identity-only sibling of :meth:`_admit_dcr_bridge_delegate`: the session token seals no + upstream credential (those are vaulted per user, resolved at egress), so authorization is + resolved fresh via :meth:`_reload_admitted_user` + the centralized policy gate rather than a + mint-time snapshot. Pre-DB gates (size, IP, route allowlist) run first, mirroring the standard + pipeline. Fails closed with the aggregate ``invalid_token`` challenge on an expired, tampered, + foreign, or refresh token, or a missing/deactivated/policy-rejected user.""" + from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import ( + NotSessionBearer, + SessionBearerAdmitted, + SessionBearerInvalid, + resolve_session_bearer, + session_keys_from_master_key, + ) + from litellm.proxy.proxy_server import master_key + + if not master_key: + raise HTTPException(status_code=500, detail="Server misconfigured: master_key is not set") + + await MCPRequestHandler._run_pre_db_read_auth_checks(request=request, route=route) + + keys = session_keys_from_master_key(master_key) + result = resolve_session_bearer(authorization_value, keys, datetime.now(timezone.utc)) + match result: + case SessionBearerAdmitted(): + try: + admitted = await MCPRequestHandler._reload_admitted_user(result.principal.user_id) + await MCPRequestHandler._enforce_admitted_live_policy( + admitted=admitted, request=request, route=route + ) + except HTTPException as exc: + # A cryptographically valid bearer whose referenced user is now missing or + # SCIM-deactivated is an invalid_token at the aggregate scope: relay the RFC 9728 + # challenge so the DCR client re-authorizes, matching the SessionBearerInvalid + # arm, instead of a bare 401 with no WWW-Authenticate. A 503 (DB outage) is a + # transient availability failure, not an auth failure, so it passes through. + if exc.status_code == 401: + raise _aggregate_gateway_dcr_challenge(request, invalid_token=True) from exc + raise + return admitted + case SessionBearerInvalid(): + raise _aggregate_gateway_dcr_challenge(request, invalid_token=True) + case NotSessionBearer(): + # Unreachable: the arm is entered only for an is_session_bearer_shaped + # value. Kept for match exhaustiveness and fails closed regardless. + raise _aggregate_gateway_dcr_challenge(request, invalid_token=True) + case _: + assert_never(result) + @staticmethod async def _run_pre_db_read_auth_checks(request: Request, route: str) -> None: """Run the proxy-wide gates ``user_api_key_auth`` applies before any key lookup: the @@ -666,30 +822,18 @@ class MCPRequestHandler: @staticmethod async def _reload_admitted_user(user_id: str) -> UserAPIKeyAuth: - """Reload the live user an interactively-minted envelope references and admit them as - themselves. + """Reload the live user an interactively-minted envelope references and admit them as themselves. - The DCR client authenticates via SSO at the bridged authorize, which yields a user - subject rather than a virtual key, so the envelope admits under the user's own - identity: the reloaded ``user_id`` and the user's own MCP object permission ride on the - returned ``UserAPIKeyAuth``, and the SAME ``get_allowed_mcp_servers`` the key path uses then - computes which servers the user may reach, so the user's litellm MCP grants and access groups - gate the request exactly as a key's do. Only the user's OWN object permission is bound: a - ``UserAPIKeyAuth`` carries a single ``team_id`` while a user may belong to many teams, so - team-inherited MCP grants for a user are a follow-up (they need a many-teams union - ``get_allowed_mcp_servers`` does not do off one auth object). The caller's centralized policy - gate enforces the user's live budget and org state, and a SCIM-deactivated owner fails closed. + The user's own object permission and ``org_id`` ride on the returned ``UserAPIKeyAuth``, and the + SAME ``get_allowed_mcp_servers`` the key path uses gates the request. The ``mcp_admitted_user_subject`` + marker (set below) makes that resolver union the servers the user reaches through ANY of their teams + on top of these direct grants, each source bounded by ITS OWN org, so a user spanning organizations + cannot leak one org's servers past another's ceiling. - Error handling mirrors the key path's retryable-503 contract, but ``get_user_object`` defeats a - type-based check: where ``get_key_object`` raises a typed ``ProxyException`` for a missing key - and lets a DB outage propagate raw, ``get_user_object`` catches every DB failure and re-raises a - bare ``ValueError``, so a missing user and a real outage look identical and the original error - survives only as ``__context__``. ``_raise_503_if_db_unavailable`` therefore walks the cause - chain: a transient DB outage still surfaces as a retryable 503, while a missing user, or any - other non-outage resolution failure, fails closed as a 401 rather than an opaque 500. The - object-permission load shares this one boundary, so an outage there is classified the same - way (``get_object_permission`` itself swallows a failed load to ``None``, matching how - ``get_key_object`` best-effort-loads a key's object permission).""" + Error handling: ``get_user_object`` catches every DB failure and re-raises a bare ``ValueError``, so a + missing user and a real outage look identical (the cause survives only as ``__context__``). + ``_raise_503_if_db_unavailable`` walks the cause chain so an outage stays a retryable 503 while any + other failure fails closed as 401, not an opaque 500; the object-permission load shares that boundary.""" from litellm.proxy.auth.auth_checks import get_object_permission, get_user_object from litellm.proxy.proxy_server import prisma_client, user_api_key_cache @@ -721,12 +865,75 @@ class MCPRequestHandler: raise HTTPException(status_code=401, detail="Invalid or expired credential") if isinstance(user_object.metadata, dict) and user_object.metadata.get("scim_active") is False: raise HTTPException(status_code=401, detail="Invalid or expired credential") - return UserAPIKeyAuth( + admitted = UserAPIKeyAuth( user_id=user_object.user_id, user_role=user_object.user_role, + org_id=user_object.organization_id, object_permission=object_permission, object_permission_id=user_object.object_permission_id, + # Copy the live user's rate limits, as the standard user-subject path does: the parallel + # limiter reads these off the auth object and treats None as unlimited, so a keyless subject + # with them unset would outrun its user RPM/TPM. (Per-team mcp_rpm_limit is stamped below; + # per-KEY limits do not apply, there being no key.) + user_tpm_limit=user_object.tpm_limit, + user_rpm_limit=user_object.rpm_limit, ) + # Server-only marker, set AFTER construction: the before-validator strips it from any validated + # input, so caller-supplied data (key metadata, JWT claims) can never forge it. + admitted.mcp_admitted_user_subject = True + # Carry each granting team's per-server mcp_rpm_limit: this subject reaches servers through + # several teams under its own identity, so without this a cross-team user outruns every team's + # limit. Resolved from the same roster-checked sources as the grant union, so a team throttles + # only what it granted. + admitted.mcp_source_team_rpm_limits = await MCPRequestHandler._admitted_subject_team_rpm_limits(admitted) + return admitted + + @staticmethod + async def _admitted_subject_team_rpm_limits(auth: UserAPIKeyAuth) -> dict[str, dict[str, int]] | None: + """``team_id -> mcp_rpm_limit`` for every team this subject reaches servers through, each map + filtered to the servers THAT team's grant actually reaches. + + A limit rides the same scope as the access it bounds, so a roster team is charged only for a + server its OWN grant reaches (never one the user reaches through a different team, which would + drain a bucket shared by that team's keys for access it never provided). Grant scope comes from + the SAME ``get_allowed_mcp_servers(source)`` authorization uses; limit-map keys are names/aliases + so each is resolved to an id via ``expand_permission_list`` before the membership check. Returns + None (no descriptors) when nothing applies; a lookup failure narrows to None rather than raising, + since rate limiting must not deny a request authorization already allowed.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + try: + limits: dict[str, dict[str, int]] = {} + source_grants = await MCPRequestHandler.admitted_source_grants(auth) + for source, granted_ids in source_grants: + if not source.team_id: + continue + team_obj = await MCPRequestHandler._roster_team_object(source.team_id, auth) + team_limit = (team_obj.metadata or {}).get("mcp_rpm_limit") if team_obj is not None else None + if not isinstance(team_limit, dict) or not team_limit: + continue + applicable: dict[str, int] = {} + for server_name, rpm in team_limit.items(): + for server_id in global_mcp_server_manager.expand_permission_list([server_name]): + if server_id not in granted_ids: + continue + # Charge ONLY the source billing attributes the call to (same owner), so one + # cross-team user cannot drain several teams' shared buckets on a single call, + # and a server the user's OWN grant reaches charges no team bucket. + attributed = await MCPRequestHandler.attributing_source_for_server( + auth, server_id, source_grants=source_grants + ) + if attributed is not None and attributed.team_id == source.team_id: + applicable[server_name] = rpm + break + if applicable: + limits[source.team_id] = applicable + return limits or None + except Exception as e: # noqa: BLE001 # throttling metadata must never fail an allowed request + verbose_logger.warning(f"Failed to resolve per-team MCP rpm limits for admitted subject: {str(e)}") + return None @staticmethod async def _reload_admitted_key(key_hash: str) -> UserAPIKeyAuth: @@ -1075,6 +1282,8 @@ class MCPRequestHandler: @staticmethod async def get_allowed_mcp_servers( user_api_key_auth: Optional[UserAPIKeyAuth] = None, + *, + keyless_source: bool = False, ) -> List[str]: """ Get list of allowed MCP servers for the given user/key based on permissions. @@ -1096,6 +1305,13 @@ class MCPRequestHandler: from litellm.proxy.proxy_server import general_settings try: + # A keyless admitted subject resolves per source BEFORE any single-source rule here. Ordering + # matters: the no_mcp_servers opt-out below reads the caller's own object_permission, so above + # this branch a user's own opt-out would wrongly zero their TEAMS' grants too (each source is + # independent; an opt-out silences only its own source, inside the recursive call). + if _is_mcp_admitted_user_subject(user_api_key_auth) and user_api_key_auth is not None: + return await MCPRequestHandler._resolve_admitted_subject_servers(user_api_key_auth) + # Get allowed servers from key and team allowed_mcp_servers_for_key = await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth) @@ -1125,8 +1341,18 @@ class MCPRequestHandler: # team's by default. With require_key_mcp_access_defined the # team is a ceiling rather than a default, so the key must # grant servers explicitly (or via an access group) to reach - # any — it inherits none. - base = set() if general_settings.get("require_key_mcp_access_defined", False) else team_set + # any — it inherits none. That ceiling is for VIRTUAL KEYS that + # can declare their own access; a keyless gateway/bridge-admitted + # user has no key to declare access on — team membership IS their + # only access path — so the flag must not zero their team grants. + # A keyless admitted subject returned above and never reaches this virtual-key ceiling, + # so require_key_mcp_access_defined can only ever zero a real key's inherited team grants. + # ``keyless_source`` marks one grant source of an admitted subject, which has no key + # to declare access on, so the flag must not zero its team grants. + require_key_access = ( + general_settings.get("require_key_mcp_access_defined", False) and not keyless_source + ) + base = team_set if not require_key_access else set() else: base = key_set & team_set # both restrict → intersect @@ -1185,24 +1411,290 @@ class MCPRequestHandler: ######################################################### # Apply org-level ceiling if org_id is set ######################################################### - if user_api_key_auth and user_api_key_auth.org_id: - allowed_mcp_servers_for_org = await MCPRequestHandler._get_allowed_mcp_servers_for_org( - user_api_key_auth - ) - if len(allowed_mcp_servers_for_org) > 0: - if has_lower_level_mcp_restrictions: - # Lower-level restrictions exist, so org can only cap them. - allowed_mcp_servers = [s for s in allowed_mcp_servers if s in allowed_mcp_servers_for_org] - else: - # No lower-level restrictions → org list becomes the ceiling - allowed_mcp_servers = allowed_mcp_servers_for_org - verbose_logger.debug(f"Applied org ceiling filter. Final allowed servers: {allowed_mcp_servers}") + allowed_mcp_servers = await MCPRequestHandler._apply_primary_org_ceiling( + allowed_mcp_servers, + user_api_key_auth, + has_lower_level_mcp_restrictions, + keyless_source=keyless_source, + ) return list(set(allowed_mcp_servers)) except Exception as e: verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}") return [] + @staticmethod + async def _apply_primary_org_ceiling( + allowed_mcp_servers: list[str], + user_api_key_auth: UserAPIKeyAuth | None, + has_lower_level_mcp_restrictions: bool, + keyless_source: bool = False, + ) -> list[str]: + """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.""" + 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) + if allowed_mcp_servers_for_org is None: + verbose_logger.warning( + f"MCP org ceiling unresolved for org_id={user_api_key_auth.org_id!r}; " + f"{'denying (keyless admitted subject)' if keyless_source else 'leaving uncapped (key auth)'}" + ) + return [] if keyless_source else allowed_mcp_servers + if len(allowed_mcp_servers_for_org) == 0: + return allowed_mcp_servers + if has_lower_level_mcp_restrictions or keyless_source: + # Org can only cap lower-level restrictions. A keyless admitted source ALWAYS takes this + # arm: its model unions GRANTS, so an org list may only narrow a source, never become one. + capped = [s for s in allowed_mcp_servers if s in allowed_mcp_servers_for_org] + else: + # No lower-level restrictions → org list becomes the ceiling. + capped = allowed_mcp_servers_for_org + verbose_logger.debug(f"Applied org ceiling filter. Final allowed servers: {capped}") + return capped + + @staticmethod + def _scoped_source_auth( + auth: UserAPIKeyAuth, + *, + team_id: str | None, + org_id: str | None, + carry_user_grants: bool, + ) -> UserAPIKeyAuth: + """A plain, UNMARKED auth describing ONE grant source of an admitted subject. + + Only the fields the resolver consults are carried; everything else is left at its default on + purpose: no ``api_key``/``token`` (not a key), no budget/spend/rate-limit (the subject's own + user-level limits meter the request, and per-source copies would double descriptors), no + ``user_role`` (an admin role would grant every server at the server-manager wrapper). The + admission marker cannot be set via the constructor (a before-validator pops it), so each source + resolves as an ordinary caller and cannot re-enter the admitted path.""" + scoped = UserAPIKeyAuth( + user_id=auth.user_id, + team_id=team_id, + org_id=org_id, + parent_otel_span=auth.parent_otel_span, + ) + if carry_user_grants: + # The user's OWN grants. A team source carries none of these (the resolver loads the team's + # own object_permission from team_id); mixing them in would widen the team with grants it never made. + scoped.object_permission = auth.object_permission + scoped.object_permission_id = auth.object_permission_id + scoped.access_group_ids = auth.access_group_ids + return scoped + + @staticmethod + async def _admitted_subject_sources(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]: + """The independent sources a keyless admitted subject reaches MCP servers through: their own + direct grants, plus every team they are a live roster member of. + + Each team source carries that TEAM's org (falling back to the user's), so the canonical resolver + applies the team's OWN owning-org ceiling — a cross-org user's teams are each bounded by their + own org, not the caller's home org. Roster membership is checked HERE (not per resolution) + because a user's cached ``teams`` array can name a team whose ``members_with_roles`` no longer + lists them; the roster is the source of truth for revocation.""" + from litellm.proxy.proxy_server import prisma_client + + sources = [ + MCPRequestHandler._scoped_source_auth(auth, team_id=None, org_id=auth.org_id, carry_user_grants=True) + ] + if not auth.user_id or prisma_client is None: + return sources + for team_id in await MCPRequestHandler._resolve_user_team_ids(auth.user_id, auth): + team_obj = await MCPRequestHandler._roster_team_object(team_id, auth) + if team_obj is None: + continue + sources.append( + MCPRequestHandler._scoped_source_auth( + auth, + team_id=team_id, + org_id=team_obj.organization_id or auth.org_id, + carry_user_grants=False, + ) + ) + return sources + + @staticmethod + async def _roster_team_object(team_id: str, auth: UserAPIKeyAuth) -> LiteLLM_TeamTable | None: + """The team row for ``team_id``, but ONLY when ``auth``'s user is a live roster member of it. + + The single owner of "is this team really one of this subject's sources", so the grant union + and the per-team rate limits cannot disagree about which teams count. A team lingering in the + user's cached ``teams`` array whose ``members_with_roles`` no longer lists them returns None + here, which is what revokes both its grants and its throttle in one place.""" + from litellm.proxy.auth.auth_checks import get_team_object + from litellm.proxy.proxy_server import ( + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) + + if prisma_client is None or not auth.user_id: + return None + try: + team_obj: LiteLLM_TeamTable | None = await get_team_object( + team_id=team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=auth.parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) + except Exception as e: # noqa: BLE001 # per-source isolation: one team's blip must not deny the others + # Fault isolation is per SOURCE: an unresolvable team contributes nothing (fail closed for + # it alone, access only narrows) while every other source stands. Raising would collapse the + # whole union to deny-all over one momentarily-unreadable row. + verbose_logger.warning(f"MCP admitted-subject source team {team_id!r} unresolvable, skipping: {str(e)}") + return None + if team_obj is None: + return None + member_user_ids = {getattr(m, "user_id", None) for m in (team_obj.members_with_roles or [])} - {None} + if auth.user_id not in member_user_ids: + return None + # A team (or its owning org) over budget is not a live grantor, exactly as it is not for a key + # pinned to it. Enforced via the SAME owners the key path uses (_team_max_budget_check / + # _organization_max_budget_check), targeted at the TEAM's org through the scoped source view, so + # no consumer of the source list ever sees an over-budget team. This is ENFORCEMENT of an + # already-exceeded state; ATTRIBUTION of new spend stays with the user (documented deferral). + from litellm.exceptions import BudgetExceededError + from litellm.proxy.auth.auth_checks import ( + _organization_max_budget_check, + _team_max_budget_check, + ) + + source_view = MCPRequestHandler._scoped_source_auth( + auth, team_id=team_id, org_id=team_obj.organization_id or auth.org_id, carry_user_grants=False + ) + try: + await _team_max_budget_check( + team_object=team_obj, valid_token=source_view, proxy_logging_obj=proxy_logging_obj + ) + await _organization_max_budget_check( + valid_token=source_view, + team_object=team_obj, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + except BudgetExceededError as e: + verbose_logger.info(f"MCP admitted-subject source team {team_id!r} over budget, not a grantor: {str(e)}") + return None + except Exception as e: # noqa: BLE001 # per-source isolation: a budget-check fault narrows, never raises + verbose_logger.warning(f"MCP budget check failed for source team {team_id!r}, skipping source: {str(e)}") + return None + return team_obj + + @staticmethod + async def admitted_source_grants(auth: UserAPIKeyAuth) -> list[tuple[UserAPIKeyAuth, set[str]]]: + """``(source, the servers that source grants)`` for every source of an admitted subject. + + THE owner of "which source reaches which server". The reachable union, the per-team throttle + scope, the tool union and billing attribution are all just different reads of this one + answer — computing it separately per consumer is how they drift (a throttle map scoped by + roster instead of by grant charged unrelated teams' buckets).""" + return [ + (source, set(await MCPRequestHandler.get_allowed_mcp_servers(source, keyless_source=True))) + for source in await MCPRequestHandler._admitted_subject_sources(auth) + ] + + @staticmethod + async def _resolve_admitted_subject_servers(auth: UserAPIKeyAuth) -> list[str]: + """Union of what each of the admitted subject's sources reaches, each answered by the + canonical resolver so no rule is reimplemented for this caller shape.""" + reachable: set[str] = set() + for _source, granted in await MCPRequestHandler.admitted_source_grants(auth): + reachable.update(granted) + return list(reachable) + + @staticmethod + async def billing_auth_for_tool_call(auth: UserAPIKeyAuth, tool_name: str) -> UserAPIKeyAuth: + """The auth object a tool call's SPEND should be recorded against. + + ``auth`` unchanged for any non-admitted caller (key/JWT billing byte-identical). For an admitted + subject whose call is reached through a team's grant, a copy carrying that team's ``team_id`` and + owning ``org_id`` so the team's budget accumulates and the right org is charged. Falls back to + user-level attribution (rather than guessing a team) when the tool name does not resolve to a + server, reusing the manager's own tool-name lookup.""" + if not _is_mcp_admitted_user_subject(auth): + return auth + try: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + server = global_mcp_server_manager._get_mcp_server_from_tool_name(tool_name) + if server is None: + return auth + source = await MCPRequestHandler.attributing_source_for_server(auth, server.server_id) + if source is None or not source.team_id: + return auth + billed = auth.model_copy() + billed.team_id = source.team_id + billed.org_id = source.org_id + return billed + except Exception as e: # noqa: BLE001 # attribution must never fail an authorized call + verbose_logger.warning(f"MCP billing attribution failed for {tool_name!r}, billing the user: {str(e)}") + return auth + + @staticmethod + async def attributing_source_for_server( + auth: UserAPIKeyAuth, + server_id: str, + source_grants: list[tuple[UserAPIKeyAuth, set[str]]] | None = None, + ) -> UserAPIKeyAuth | None: + """The source a billable call to ``server_id`` is attributed to, or None to bill the caller as + themselves (their own grant reaches it, or nothing does). + + The rule: a user's OWN grant is not "through a team", so it bills the user; otherwise the call + bills a granting team, deterministically the lowest ``team_id`` when several grant the server so + the pick is stable rather than dict-ordering-dependent. Reads the one grant owner, so the billed + team is always one that actually granted the server (restoring the team budget accrual and + owning-org charge that a keyless, team_id-less subject otherwise skipped).""" + source_grants = source_grants or await MCPRequestHandler.admitted_source_grants(auth) + granting = [(source, granted) for source, granted in source_grants if server_id in granted] + if not granting: + return None + for source, _granted in granting: + if source.team_id is None: + return None # the user's own grant reaches it: their spend, their org + return min((source for source, _ in granting), key=lambda s: s.team_id or "") + + @staticmethod + async def _resolve_admitted_subject_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None: + """Effective tool allowlist on ``server_id`` for an admitted subject, as the union over the + sources that actually grant that server. + + A source that does not grant the server contributes nothing, so its tool rules cannot leak + onto a server reached through a different source. A source that grants the server with no + tool restriction means the user can use every tool on it, so allow-all wins the union. When + no source grants the server the result is ``[]`` — deny all, fail closed.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + # An OPEN channel (allow_all_keys, the user's own BYOM) makes the server REACHABLE through the + # user, though no grant source names it — without this the union returns [], listable but + # uninvokable. Reachability is ALL it confers, NOT a ceiling waiver: the user's own + # mcp_tool_permissions and org tool ceiling still bind, exactly as a key's do on an allow_all server. + reachable_via_open_channel = server_id in await global_mcp_server_manager.operator_open_server_ids(auth) + + allowed: set[str] = set() + for source, granted in await MCPRequestHandler.admitted_source_grants(auth): + # The open channel is evaluated against the user's OWN source (team_id is None), so that + # source's restrictions apply to it; a team's rules never ride an open-channel server. + if server_id not in granted and not (reachable_via_open_channel and source.team_id is None): + continue + tools = await MCPRequestHandler.get_allowed_tools_for_server(server_id, source, keyless_source=True) + if tools is None: + return None + allowed.update(tools) + return sorted(allowed) + @staticmethod def _get_key_object_permission( user_api_key_auth: Optional[UserAPIKeyAuth] = None, @@ -1262,6 +1754,8 @@ class MCPRequestHandler: async def get_allowed_tools_for_server( server_id: str, user_api_key_auth: Optional[UserAPIKeyAuth] = None, + *, + keyless_source: bool = False, ) -> Optional[List[str]]: """ Get list of allowed tool names for a specific server based on key/team permissions. @@ -1278,6 +1772,12 @@ class MCPRequestHandler: return None try: + # FIRST statement, mirroring get_allowed_mcp_servers: a keyless admitted subject resolves per + # source and shares nothing with the single-credential prelude below. Ordering is the invariant: + # sat after the prelude, a fault in a lookup the subject never uses denied tools its teams grant. + if _is_mcp_admitted_user_subject(user_api_key_auth): + return await MCPRequestHandler._resolve_admitted_subject_tools(server_id, user_api_key_auth) + # Get key and team object permissions (already loaded in main auth flow) key_obj_perm = MCPRequestHandler._get_key_object_permission(user_api_key_auth) team_obj_perm = await MCPRequestHandler._get_team_object_permission(user_api_key_auth) @@ -1331,42 +1831,73 @@ class MCPRequestHandler: # No team restrictions → use key restrictions allowed_tools = cast(List[str], key_tools) - # Intersect with agent's tool permissions if agent_id is set - if user_api_key_auth.agent_id: - # Pre-fetch agent object_permission once to avoid duplicate DB query - agent_obj_perm = await MCPRequestHandler._get_agent_object_permission(user_api_key_auth) - agent_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server( - server_id=server_id, - user_api_key_auth=user_api_key_auth, - agent_object_permission=agent_obj_perm, - ) - if agent_tools is not None: - if allowed_tools is not None: - allowed_tools = list(set(allowed_tools) & set(agent_tools)) - else: - allowed_tools = agent_tools - - # Apply org-level tool ceiling if org_id is set - if user_api_key_auth.org_id: - # _get_org_object_permission uses user_api_key_cache, so this is not a - # fresh DB round-trip when get_allowed_mcp_servers was already called. - org_obj_perm = await MCPRequestHandler._get_org_object_permission(user_api_key_auth) - org_tools = ( - global_mcp_server_manager.expand_tool_permissions(org_obj_perm.mcp_tool_permissions).get(server_id) - if org_obj_perm and org_obj_perm.mcp_tool_permissions - else None - ) - if org_tools is not None: - if allowed_tools is not None: - allowed_tools = list(set(allowed_tools) & set(org_tools)) - else: - allowed_tools = list(org_tools) - - return allowed_tools + return await MCPRequestHandler._apply_agent_and_org_tool_ceilings( + allowed_tools, server_id, user_api_key_auth, keyless_source=keyless_source + ) except Exception as e: verbose_logger.warning(f"Failed to get allowed tools for server: {str(e)}") - return None + # 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 + + @staticmethod + async def _apply_agent_and_org_tool_ceilings( + allowed_tools: list[str] | None, + server_id: str, + user_api_key_auth: UserAPIKeyAuth, + keyless_source: bool = False, + ) -> list[str] | None: + """Narrow a key/team tool allowlist by the agent's tool permissions and the caller's org tool + ceiling. Each level only intersects; None at a level means no restriction from it. + + An UNRESOLVABLE org ceiling is decided per caller shape, mirroring the servers axis: a key stays + fail-open (skip the org step, keep the key/team/agent restrictions; letting the raise escape + would collapse them to allow-all, WIDER than before the fault), while a keyless source re-raises + so the outer handler denies that one source (its only org bound is this ceiling).""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + if user_api_key_auth.agent_id: + # Pre-fetch agent object_permission once to avoid a duplicate DB query. + agent_obj_perm = await MCPRequestHandler._get_agent_object_permission(user_api_key_auth) + agent_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server( + server_id=server_id, + user_api_key_auth=user_api_key_auth, + agent_object_permission=agent_obj_perm, + ) + if agent_tools is not None: + allowed_tools = ( + list(set(allowed_tools) & set(agent_tools)) if allowed_tools is not None else agent_tools + ) + + if user_api_key_auth.org_id: + # _get_org_object_permission uses user_api_key_cache, so this is not a fresh DB round-trip + # when get_allowed_mcp_servers was already called. + 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: + raise + verbose_logger.warning( + f"MCP org tool ceiling unresolvable for org_id={user_api_key_auth.org_id!r}; " + f"skipping org intersect, key/team/agent restrictions stand: {str(e)}" + ) + return allowed_tools + org_tools = ( + global_mcp_server_manager.expand_tool_permissions(org_obj_perm.mcp_tool_permissions).get(server_id) + if org_obj_perm and org_obj_perm.mcp_tool_permissions + else None + ) + if org_tools is not None: + allowed_tools = ( + list(set(allowed_tools) & set(org_tools)) if allowed_tools is not None else list(org_tools) + ) + + return allowed_tools @staticmethod async def is_tool_allowed_for_server( @@ -1537,10 +2068,96 @@ class MCPRequestHandler: @staticmethod async def _get_allowed_mcp_servers_for_team( - user_api_key_auth: Optional[UserAPIKeyAuth] = None, - ) -> List[str]: + user_api_key_auth: UserAPIKeyAuth | None = None, + ) -> list[str]: + """Get allowed MCP servers a caller inherits from the team it is pinned to. + + Exactly one team, or none. A subject that reaches servers through SEVERAL teams does not + fan out here: it is resolved one source per team in ``_resolve_admitted_subject_servers``, + and each of those sources pins a single ``team_id`` before reaching this point. Keeping the + fan-out here as well would be a second multi-team path to drift from that one. """ - Get allowed MCP servers for a team. + team_ids = await MCPRequestHandler._team_ids_for_mcp_grant(user_api_key_auth) + if not team_ids: + return [] + return await MCPRequestHandler._allowed_mcp_servers_for_single_team(team_ids[0], user_api_key_auth) + + @staticmethod + async def _team_ids_for_mcp_grant(user_api_key_auth: UserAPIKeyAuth | None) -> list[str]: + """The team ids whose MCP grants a caller inherits. + + A caller with an explicit ``team_id`` uses that single team; every other caller inherits no + team grants. That covers key auth and JWT auth (a keyless ``user_id`` auth with no team_id, + which must NOT silently gain the union across every team the user belongs to), and it covers + each single-source auth an admitted subject fans out into — those pin a team_id, so they land + on the first branch. The admitted subject itself never reaches here: it resolves per source + in ``_resolve_admitted_subject_servers`` before this point. The ``UI_TEAM_ID`` sentinel + resolves to no teams exactly as before.""" + if user_api_key_auth is None or not user_api_key_auth.team_id: + return [] + return [] if user_api_key_auth.team_id == UI_TEAM_ID else [user_api_key_auth.team_id] + + @staticmethod + async def _resolve_user_team_ids(user_id: str, user_api_key_auth: UserAPIKeyAuth) -> list[str]: + """The distinct team ids a user belongs to, from the live user record. Returns [] on + no DB, a missing user, or any resolution failure so a lookup blip narrows access + rather than raising; the caller's direct grants still apply.""" + from litellm.proxy.auth.auth_checks import get_user_object + from litellm.proxy.proxy_server import ( + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) + + if prisma_client is None: + return [] + try: + user_object = await get_user_object( + user_id=user_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + user_id_upsert=False, + parent_otel_span=user_api_key_auth.parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) + except Exception as e: # noqa: BLE001 # a team-resolution blip narrows access, never raises + verbose_logger.warning(f"Failed to resolve user teams for MCP grant: {str(e)}") + return [] + if user_object is None or not user_object.teams: + return [] + return list(dict.fromkeys(t for t in user_object.teams if t and t != UI_TEAM_ID)) + + @staticmethod + async def _team_granted_servers(team_obj: LiteLLM_TeamTable, team_access_group_servers: list[str]) -> set[str]: + """The raw MCP-server set a team grants (before any org ceiling): its object_permission (direct + ``mcp_servers``, the ``all_proxy_servers`` sentinel → the full registry, legacy access groups, + tool-perm-referenced servers) unioned with its unified ``access_group_ids`` servers.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + object_permissions = team_obj.object_permission + if object_permissions is None: + return set(team_access_group_servers) + if SpecialMCPServerName.all_proxy_servers.value in (object_permissions.mcp_servers or []): + return set(global_mcp_server_manager.get_registry().keys()) + legacy_access_group_servers = await MCPRequestHandler._get_mcp_servers_from_access_groups( + object_permissions.mcp_access_groups or [] + ) + return ( + set(global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or [])) + | set(legacy_access_group_servers) + | set(global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys()) + | set(team_access_group_servers) + ) + + @staticmethod + async def _allowed_mcp_servers_for_single_team( + team_id: str, + user_api_key_auth: UserAPIKeyAuth | None, + ) -> list[str]: + """Allowed MCP servers granted by ONE team (its raw grant, then capped by the team's own org + for a keyless admitted subject). Unions two sources: - Legacy team.object_permission (mcp_servers, mcp_access_groups, @@ -1551,9 +2168,6 @@ class MCPRequestHandler: the gate (no assigned_team_ids check needed here). """ try: - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - global_mcp_server_manager, - ) from litellm.proxy.auth.auth_checks import ( _get_mcp_server_ids_from_access_groups, get_team_object, @@ -1564,22 +2178,24 @@ class MCPRequestHandler: user_api_key_cache, ) - if user_api_key_auth is None or not user_api_key_auth.team_id or prisma_client is None: + if not team_id or team_id == UI_TEAM_ID or prisma_client is None: return [] - if user_api_key_auth.team_id == UI_TEAM_ID: - return [] - - team_obj: Optional[LiteLLM_TeamTable] = await get_team_object( - team_id=user_api_key_auth.team_id, + parent_otel_span = user_api_key_auth.parent_otel_span if user_api_key_auth is not None else None + team_obj: LiteLLM_TeamTable | None = await get_team_object( + team_id=team_id, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, - parent_otel_span=user_api_key_auth.parent_otel_span, + parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, ) if team_obj is None: return [] - + if team_obj.blocked: + # A blocked team grants nothing. The central policy gate enforces this for a key + # pinned to a single team_id, but a keyless admitted identity (no team_id) unions + # across all of its teams and would otherwise inherit a blocked team's MCP grants. + return [] team_access_group_servers = await _get_mcp_server_ids_from_access_groups( access_group_ids=team_obj.access_group_ids or [], prisma_client=prisma_client, @@ -1587,27 +2203,8 @@ class MCPRequestHandler: proxy_logging_obj=proxy_logging_obj, ) - object_permissions = team_obj.object_permission - if object_permissions is None: - return list(set(team_access_group_servers)) - - if SpecialMCPServerName.all_proxy_servers.value in (object_permissions.mcp_servers or []): - return list(global_mcp_server_manager.get_registry().keys()) - - direct_mcp_servers = global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or []) - - legacy_access_group_servers = await MCPRequestHandler._get_mcp_servers_from_access_groups( - object_permissions.mcp_access_groups or [] - ) - - tool_perm_servers = list( - global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys() - ) - - all_servers = ( - direct_mcp_servers + legacy_access_group_servers + tool_perm_servers + team_access_group_servers - ) - return list(set(all_servers)) + servers = await MCPRequestHandler._team_granted_servers(team_obj, team_access_group_servers) + return list(servers) except Exception as e: verbose_logger.warning(f"Failed to get allowed MCP servers for team: {str(e)}") return [] @@ -1621,7 +2218,11 @@ class MCPRequestHandler: ``get_object_permission`` helpers so MCP requests share the same ``user_api_key_cache`` entries as the rest of the proxy. """ - from litellm.proxy.auth.auth_checks import get_object_permission, get_org_object + from litellm.proxy.auth.auth_checks import ( + OrganizationNotFoundError, + get_object_permission, + get_org_object, + ) from litellm.proxy.proxy_server import ( prisma_client, proxy_logging_obj, @@ -1635,6 +2236,8 @@ class MCPRequestHandler: verbose_logger.debug("prisma_client is None") return None + # A team's organization_id can point at a deleted or not-yet-synced row; get_org_object raises + # OrganizationNotFoundError for that. That is a determinate ABSENCE (no ceiling), handled below. try: org_obj = await get_org_object( org_id=user_api_key_auth.org_id, @@ -1643,21 +2246,32 @@ class MCPRequestHandler: parent_otel_span=user_api_key_auth.parent_otel_span, proxy_logging_obj=proxy_logging_obj, ) - - if org_obj is None or not org_obj.object_permission_id: - return None - - return await get_object_permission( - 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, - ) - except Exception as e: - verbose_logger.warning(f"Failed to get org object permission: {str(e)}") + except OrganizationNotFoundError as e: + # CONFIRMED absent: places no ceiling. Every OTHER exception propagates as an unresolvable + # ceiling (denies for a keyless source, fail-open for a key); catching bare Exception here + # would treat a DB outage as "no org" and silently drop a real ceiling for its duration. + verbose_logger.debug(f"MCP org ceiling: org {user_api_key_auth.org_id!r} does not exist: {e}") return None + 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( + 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, + ) + 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, @@ -1692,8 +2306,10 @@ class MCPRequestHandler: all_servers = direct_mcp_servers + access_group_servers + tool_perm_servers return list(set(all_servers)) 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. verbose_logger.warning(f"Failed to get allowed MCP servers for org: {str(e)}") - return [] + return None @staticmethod async def _get_allowed_mcp_servers_for_end_user( diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index ee3196b539c..26241119dd8 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -33,6 +33,7 @@ from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( _finish_bridge_mint, _prepare_bridge_mint, _prepare_bridge_refresh, + _reload_active_user_by_id, ) from litellm.proxy._experimental.mcp_server.faults import ( CallerRejected, @@ -43,6 +44,14 @@ from litellm.proxy._experimental.mcp_server.faults import ( dcr_fault_detail, render_token_fault, ) +from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import ( + aggregate_authorize, + aggregate_token, + complete_connect_flow, + is_gateway_dcr_client_id, + register_aggregate_client, + relative_request_url, +) from litellm.proxy._experimental.mcp_server.oauth_utils import ( TOKEN_NO_CACHE_HEADERS, get_request_base_url, @@ -324,14 +333,25 @@ def redeem_passthrough_authorization_code( return sealed +def _session_cookie_user_id(request: Request) -> str | None: + """The signed-in litellm user for a browser request, or ``None``. Thin wrapper so the + aggregate DCR flow's verbs receive the identity as a plain value instead of parsing + cookies themselves.""" + from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( # noqa: PLC0415 # circular import at module load + _user_id_from_session_cookie, + ) + + return _user_id_from_session_cookie(request) + + def _redirect_to_litellm_login(request: Request) -> RedirectResponse: """Send an unauthenticated browser through litellm login before the interactive bridge authorize can capture its identity. The bridge oauth_delegate flow seals the SSO user into the gateway code, - so a session is required; without one there is nothing to bind. After login the user re-initiates - the connection, which then finds the session cookie (the seamless return-to round-trip, which is - origin-validated against the control-plane URL, is a follow-up).""" + so a session is required; without one there is nothing to bind. A same-origin relative + ``return_to`` (honored by the SSO callback) brings the browser straight back to this authorize + request after login instead of stranding it on the dashboard.""" base_url = get_request_base_url(request) - return RedirectResponse(f"{base_url}/sso/key/generate") + return RedirectResponse(f"{base_url}/sso/key/generate?{urlencode({'return_to': relative_request_url(request)})}") # LIT-4197: some upstream authorization servers reject an over-long ``state`` @@ -606,6 +626,35 @@ def _raise_if_not_oauth2(mcp_server: MCPServer) -> None: ) +def _endpoint_not_configured_detail( + mcp_server: MCPServer, + endpoint_label: str, + manual_remedy: str, + issuer_remedy: str, +) -> str: + """The 400 detail for an unresolved OAuth endpoint, naming the likely cause for this server's + shape (LIT-4658): an anchored issuer whose metadata fell short, a configured (possibly + misconfigured) server url whose discovery failed, or no discovery source at all. Kept free of + URLs and issuer values because these endpoints are reachable pre-auth.""" + if mcp_server.issuer_is_anchored: + return ( + f"MCP server {endpoint_label} is not configured. Endpoint discovery anchored on the configured " + f"Issuer (RFC 8414) failed or its metadata did not include this endpoint; check the proxy logs " + f"for 'MCP OAuth' warnings from server load, verify the Issuer, or {manual_remedy}." + ) + if mcp_server.url: + return ( + f"MCP server {endpoint_label} is not configured. OAuth endpoint discovery against the configured " + f"server url did not resolve it; the url may be misconfigured. Check the proxy logs for " + f"'MCP OAuth' warnings from server load, verify the server url, or {manual_remedy}, or " + f"{issuer_remedy}." + ) + return ( + f"MCP server {endpoint_label} is not configured. Servers with no url (OpenAPI spec or stdio) run no " + f"resource discovery, so {manual_remedy}, or {issuer_remedy}." + ) + + def _raise_unless_oauth2_discovery_server( mcp_server: Optional[MCPServer], mcp_server_name: Optional[str], @@ -707,10 +756,11 @@ async def authorize_with_server( if mcp_server.authorization_url is None: raise HTTPException( status_code=400, - detail=( - "MCP server authorization url is not configured. Servers with no url (OpenAPI " - "spec or stdio) run no resource discovery, so set Authorization URL and Token URL " - "manually, or set Issuer to discover them from the identity provider (RFC 8414)." + detail=_endpoint_not_configured_detail( + mcp_server, + "authorization url", + "set Authorization URL and Token URL manually", + "set Issuer to discover them from the identity provider (RFC 8414)", ), ) @@ -828,10 +878,11 @@ async def exchange_token_with_server( if mcp_server.token_url is None: raise HTTPException( status_code=400, - detail=( - "MCP server token url is not configured. Servers with no url (OpenAPI spec or " - "stdio) run no resource discovery, so set Token URL manually, or set Issuer to " - "discover it from the identity provider (RFC 8414)." + detail=_endpoint_not_configured_detail( + mcp_server, + "token url", + "set Token URL manually", + "set Issuer to discover it from the identity provider (RFC 8414)", ), ) @@ -1540,10 +1591,11 @@ async def register_client_with_server( if mcp_server.authorization_url is None: raise HTTPException( status_code=400, - detail=( - "MCP server authorization url is not configured. Servers with no url (OpenAPI " - "spec or stdio) run no resource discovery, so set Authorization URL and Token URL " - "manually, or set Issuer to discover them from the identity provider (RFC 8414)." + detail=_endpoint_not_configured_detail( + mcp_server, + "authorization url", + "set Authorization URL and Token URL manually", + "set Issuer to discover them from the identity provider (RFC 8414)", ), ) @@ -1601,6 +1653,18 @@ async def authorize( global_mcp_server_manager, ) + if mcp_server_name is None and client_id and is_gateway_dcr_client_id(client_id): + return aggregate_authorize( + request=request, + client_id=client_id, + redirect_uri=redirect_uri, + state=state, + code_challenge=code_challenge, + code_challenge_method=code_challenge_method, + response_type=response_type, + session_user_id=_session_cookie_user_id(request), + ) + lookup_name: Optional[str] = mcp_server_name or client_id client_ip = IPAddressUtils.get_mcp_client_ip(request) mcp_server = ( @@ -1664,6 +1728,25 @@ async def token_endpoint( global_mcp_server_manager, ) + if mcp_server_name is None and is_gateway_dcr_client_id(client_id): + from litellm.proxy.proxy_server import ( # noqa: PLC0415 # circular import at module load + master_key, + user_api_key_cache, + ) + + return await aggregate_token( + request=request, + grant_type=grant_type, + code=code, + redirect_uri=redirect_uri, + client_id=client_id, + code_verifier=code_verifier, + refresh_token=refresh_token, + master_key=master_key, + reload_user=_reload_active_user_by_id, + cache=user_api_key_cache, + ) + lookup_name = mcp_server_name or client_id client_ip = IPAddressUtils.get_mcp_client_ip(request) mcp_server = global_mcp_server_manager.get_mcp_server_by_name(lookup_name, client_ip=client_ip) @@ -1685,6 +1768,21 @@ async def token_endpoint( ) +@router.post("/authorize/complete") +async def authorize_complete(request: Request, flow: str = Form(...)): + """Finish an aggregate connect flow: mint the gateway authorization code for the + signed-in user and redirect back to the DCR client. POST plus the per-flow HttpOnly + cookie set at /authorize; an anonymous or bad-flow request just 400s.""" + from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415 # circular import at module load + + return await complete_connect_flow( + request=request, + flow_handle=flow, + session_user_id=_session_cookie_user_id(request), + cache=user_api_key_cache, + ) + + # Per RFC 6749 §4.1.2.1, an IdP that rejects an OAuth authorization request # redirects back to the configured redirect URI with ``error`` / # ``error_description`` / ``error_uri`` query params and no ``code``. The MCP @@ -2422,6 +2520,13 @@ async def register_client(request: Request, mcp_server_name: Optional[str] = Non } client_ip = IPAddressUtils.get_mcp_client_ip(request) if not mcp_server_name: + # A real DCR request carries redirect_uris (RFC 7591): route it to the aggregate DCR + # endpoint the aggregate authorization-server metadata advertises. A single-server + # deployment registers at /{server}/register instead (its bare-origin discovery + # advertises that), so this does not affect it. A request without redirect_uris is not + # a DCR request, so the legacy single-server-or-dummy fallback is kept for it. + if data.get("redirect_uris"): + return await register_aggregate_client(request=request, request_body=data) resolved = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip) if resolved: return await register_client_with_server( diff --git a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py new file mode 100644 index 00000000000..58233c4c9e5 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py @@ -0,0 +1,637 @@ +"""The gateway-level DCR flow for the aggregate ``/mcp`` endpoint (``mcp_gateway_dcr``). + +An OAuth-only DCR client (Claude Desktop, Claude Code, MCP Inspector) pointed at the +aggregate ``/mcp`` endpoint discovers the gateway as its authorization server (PR 1 of +this track) and then walks the flow implemented here: + +1. ``POST /register``: stateless dynamic client registration. The ``client_id`` IS the + registration: the client's redirect URIs are sealed into it with the repo's + authenticated symmetric helper, so nothing is persisted and a forged or tampered + client_id simply fails to open. Clients are always public (``token_endpoint_auth_method + "none"``); PKCE S256 is what protects the code. +2. ``GET /authorize``: validates the client and redirect URI, requires S256 PKCE, and + interposes LiteLLM sign-in. Without a session cookie the browser is sent through + ``/sso/key/generate`` with a same-origin ``return_to`` so it lands back here after + login. With a session, the flow parameters and the SSO user are sealed into a per-flow + HttpOnly cookie (the same pattern as the upstream OAuth state relay) and the browser is + sent to the connect page, where the user authorizes individual servers (vaulting those + tokens server-side) before finishing. +3. ``POST /authorize/complete``: the deliberate finish step. A POST (not GET) bound to the + SameSite=Lax flow cookie, so a cross-site link cannot silently mint a code with the + victim's session, and the signed-in user must match the user sealed into the flow. + Mints a short-lived, single-use, gateway-sealed authorization code and redirects to the + client's registered redirect URI. +4. ``POST /token``: exchanges the code (PKCE-verified, client- and redirect-bound, + single-use) for the identity-only session tokens of + :mod:`.outbound_credentials.session_token`, re-validating that the litellm user is + still active first; the ``refresh_token`` grant rotates the pair the same way. + +Nothing here stores state server-side except the single-use code guard (a TTL cache +entry). Every sealed value is authenticated encryption over the proxy salt/master key +family, opened totally (bad input maps to an OAuth error, never a raise), and every +identity is a stable reference re-validated live at mint, refresh, and (in the admission +PR) tool-call time. Upstream server credentials never appear anywhere in this flow; they +are vaulted per user by the existing ``/v1/mcp`` authorize endpoints and resolved at +egress by user id. +""" + +from __future__ import annotations + +import hashlib +import hmac +import secrets +from base64 import urlsafe_b64encode +from collections.abc import Mapping +from datetime import datetime, timezone +from typing import Awaitable, Callable, Literal, TypeVar +from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse + +from fastapi import HTTPException, Request +from fastapi.responses import JSONResponse, RedirectResponse, Response +from pydantic import BaseModel, ConfigDict, Field, ValidationError +from typing_extensions import assert_never + +from litellm._logging import verbose_logger +from litellm.caching.caching import DualCache +from litellm.proxy._experimental.mcp_server.oauth_utils import ( + TOKEN_NO_CACHE_HEADERS, + get_request_base_url, + is_loopback_redirect_host, + validate_redirect_uri_shape, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import ( + SessionRefreshOpened, + open_session_refresh_bearer, + session_keys_from_master_key, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import ( + SESSION_REFRESH_TTL_SECONDS, + MintedSessionToken, + SessionKeys, + SessionPrincipal, + mint_session_refresh_token, + mint_session_token, +) +from litellm.proxy.common_utils.encrypt_decrypt_utils import ( + decrypt_value_helper, + encrypt_value_helper, +) + +GATEWAY_DCR_CLIENT_ID_PREFIX = "llm_dcrc_" +"""Marker prefix on every gateway-issued DCR client_id so the root authorize/token +endpoints can route an aggregate-flow request without decrypting, and existing per-server +flows (whose client_ids are upstream-issued) are never captured by the aggregate arm.""" + +GATEWAY_AUTH_CODE_PREFIX = "llm_gcode_" +"""Marker prefix on the gateway-sealed authorization code, distinct from the bridge +``llm_bcode_`` so neither flow can consume the other's codes.""" + +CONNECT_FLOW_COOKIE_PREFIX = "mcp_connect_flow_" +"""Per-flow HttpOnly cookie holding the sealed connect flow, keyed by a short random +handle carried in the connect-page URL (the same handle-plus-cookie pattern as the +``mcp_oauth_state_`` upstream relay, for the same reasons: replica-safe with no +server-side session store, and the sealed value never appears in a URL).""" + +CONNECT_FLOW_TTL_SECONDS = 600 +GATEWAY_AUTH_CODE_TTL_SECONDS = 120 +_CLAIM_TTL_BUFFER_SECONDS = 60 +_USED_CODE_CACHE_PREFIX = "mcp_gateway_dcr_code_used:" +_USED_FLOW_CACHE_PREFIX = "mcp_gateway_dcr_flow_used:" +_USED_REFRESH_CACHE_PREFIX = "mcp_gateway_dcr_refresh_used:" + +MAX_REDIRECT_URIS = 3 +MAX_REDIRECT_URI_LENGTH = 256 +MAX_CLIENT_ID_LENGTH = 2048 +"""Registration bounds. They exist to bound the sealed client_id, which rides inside +every session-token claim set: 3 URIs of 256 bytes seal to roughly 1.2KB, comfortably +under this cap and under the session token's own 4KB ceiling. Claude Desktop and MCP +Inspector register one or two redirect URIs.""" + +MAX_STATE_LENGTH = 1024 +"""Bound on the client ``state`` sealed into the flow cookie and echoed on the auth-code +redirect. An unbounded ``state`` can push the sealed cookie past the browser's ~4KB cap +(silently dropped, breaking the flow); spec clients send a short opaque value.""" + +MIN_CODE_VERIFIER_LENGTH = 43 +MAX_CODE_VERIFIER_LENGTH = 128 +"""RFC 7636 section 4.1 bounds for the PKCE ``code_verifier``. Enforced so an out-of-range +verifier gets a clean ``invalid_request`` instead of an opaque PKCE-mismatch.""" + +_UNPREFIXED = "" +"""Prefix for a sealed value that carries no wire marker because it is never routed by +prefix (the connect flow lives only in its own per-handle cookie, opened by that one +handle). Named so the empty-string argument to ``_seal`` / ``_open_sealed`` reads as +deliberate rather than a typo.""" + +_CLIENT_RECORD_DEBUG_KEY = "gateway_dcr_client" +_CONNECT_FLOW_DEBUG_KEY = "gateway_connect_flow" +_AUTH_CODE_DEBUG_KEY = "gateway_authorization_code" + +ReloadUserFailure = Literal["unresolvable", "unavailable", "no_active_key"] +ReloadUser = Callable[[str], Awaitable[ReloadUserFailure | None]] +"""Injected live-user revalidation (the token endpoint's mirror of admission): +``None`` means the user is active; ``unavailable`` is a retryable DB outage; anything +else fails the grant closed.""" + + +class GatewayDcrClient(BaseModel): + """The registration record sealed into a gateway DCR ``client_id``. + + ``extra="forbid"`` so a sealed value of another type (an auth code, a connect flow) + that happened to decrypt under the shared key can never validate as a client record: + cross-type confusion is rejected at the model boundary, not left to differing required + fields.""" + + model_config = ConfigDict(frozen=True, extra="forbid") + redirect_uris: tuple[str, ...] = Field(min_length=1, max_length=MAX_REDIRECT_URIS) + iat: int + + +class _ConnectFlow(BaseModel): + """One in-flight authorize: the SSO user it belongs to and the client parameters + needed to mint the code at the finish step. Sealed into the per-flow cookie. ``jti`` + makes the flow single-use at complete; ``extra="forbid"`` rejects cross-type + confusion.""" + + model_config = ConfigDict(frozen=True, extra="forbid") + user_id: str = Field(min_length=1) + client_id: str = Field(min_length=1) + redirect_uri: str = Field(min_length=1) + state: str + code_challenge: str = Field(min_length=1) + jti: str = Field(min_length=1) + exp: int + + +class _GatewayAuthCode(BaseModel): + """The gateway-sealed authorization code: the user consent it represents and the + bindings the token endpoint must verify (client, redirect URI, PKCE challenge), + plus a ``jti`` for the single-use guard. ``extra="forbid"`` rejects cross-type + confusion.""" + + model_config = ConfigDict(frozen=True, extra="forbid") + user_id: str = Field(min_length=1) + client_id: str = Field(min_length=1) + redirect_uri: str = Field(min_length=1) + code_challenge: str = Field(min_length=1) + jti: str = Field(min_length=1) + iat: int + exp: int + + +def is_gateway_dcr_client_id(client_id: str | None) -> bool: + """Cheap prefix routing test so the root endpoints only enter the aggregate arm for + clients this flow registered; every other client_id keeps today's behavior.""" + return client_id is not None and client_id.startswith(GATEWAY_DCR_CLIENT_ID_PREFIX) + + +def _oauth_error(status_code: int, error: str, description: str) -> JSONResponse: + """RFC 6749 section 5.2 / RFC 7591 section 3.2.2 error body. Descriptions carry no + token, code, or URL material so they are safe to relay to any client.""" + return JSONResponse( + status_code=status_code, + content={"error": error, "error_description": description}, + headers=TOKEN_NO_CACHE_HEADERS, + ) + + +def _seal(prefix: str, payload: BaseModel) -> str: + return prefix + encrypt_value_helper(payload.model_dump_json()) + + +_SealedModelT = TypeVar("_SealedModelT", bound=BaseModel) + + +def _open_sealed(value: str, prefix: str, model: type[_SealedModelT], debug_key: str) -> _SealedModelT | None: + """Open a sealed value totally: anything that is not prefix-shaped, does not decrypt, + or does not validate returns ``None`` for the caller to map onto an OAuth error.""" + if not value.startswith(prefix): + return None + decrypted = decrypt_value_helper(value[len(prefix) :], debug_key, return_original_value=False) + if not isinstance(decrypted, str): + return None + try: + return model.model_validate_json(decrypted) + except ValidationError: + return None + + +def open_gateway_dcr_client(client_id: str) -> GatewayDcrClient | None: + return _open_sealed(client_id, GATEWAY_DCR_CLIENT_ID_PREFIX, GatewayDcrClient, _CLIENT_RECORD_DEBUG_KEY) + + +async def register_aggregate_client(request: Request, request_body: Mapping[str, object]) -> Response: + """RFC 7591 dynamic registration against the gateway itself, statelessly. + + Only ``redirect_uris`` is authoritative; every client is registered as a public + ``token_endpoint_auth_method "none"`` client regardless of what it asked for (RFC + 7591 lets the server override metadata), because the gateway never issues client + secrets: possession of a secret would add nothing over the mandatory S256 PKCE, and a + stateless registration has nowhere to keep one. Nothing is persisted, so open + registration cannot be used to fill storage. + + Redirect-URI *hygiene* is not decided here: :func:`validate_redirect_uri_shape` is + the single owner of that rule across the MCP OAuth surface, so allowlisted native + callbacks (``cursor://``) are accepted and fragments, missing hosts, userinfo + (``https://claude.ai@attacker.example/cb``) and backslash hosts are rejected exactly + as they are on /authorize and /callback. + + What this endpoint does decide is its own trust policy, which is deliberately wider + than :func:`validate_trusted_redirect_uri`'s: registration is *public*, so any https + client may register (that is what lets a hosted MCP client register at all), and the + controls are mandatory S256 PKCE plus the consent screen showing the client origin. + http is confined to loopback per RFC 8252 section 7.3. + """ + raw_uris = request_body.get("redirect_uris") + if not isinstance(raw_uris, list) or not raw_uris or len(raw_uris) > MAX_REDIRECT_URIS: + return _oauth_error( + 400, + "invalid_redirect_uri", + f"redirect_uris must be a list of 1 to {MAX_REDIRECT_URIS} URIs", + ) + if not all(isinstance(uri, str) and len(uri) <= MAX_REDIRECT_URI_LENGTH for uri in raw_uris): + return _oauth_error( + 400, + "invalid_redirect_uri", + f"each redirect URI must be a string of at most {MAX_REDIRECT_URI_LENGTH} characters", + ) + for uri in raw_uris: + parsed = urlparse(uri) + try: + if validate_redirect_uri_shape(parsed): + continue # allowlisted native callback, e.g. cursor:// + except HTTPException as exc: + # The shared validator speaks HTTP; RFC 7591 registration answers with an OAuth + # error object, so translate the shape without re-deciding the rule. + return _oauth_error(400, "invalid_redirect_uri", str(exc.detail)) + if parsed.scheme == "https" or (parsed.scheme == "http" and is_loopback_redirect_host(parsed)): + continue + return _oauth_error( + 400, + "invalid_redirect_uri", + "each redirect URI must be https, http on a loopback host, or a registered native callback", + ) + now = datetime.now(timezone.utc) + client_id = _seal( + GATEWAY_DCR_CLIENT_ID_PREFIX, GatewayDcrClient(redirect_uris=tuple(raw_uris), iat=int(now.timestamp())) + ) + if len(client_id) > MAX_CLIENT_ID_LENGTH: + return _oauth_error(400, "invalid_client_metadata", "registered metadata is too large") + return JSONResponse( + status_code=201, + content={ + "client_id": client_id, + "client_id_issued_at": int(now.timestamp()), + "redirect_uris": list(raw_uris), + "token_endpoint_auth_method": "none", + "grant_types": ["authorization_code", "refresh_token"], + "response_types": ["code"], + }, + ) + + +def _flow_cookie_name(handle: str) -> str: + return f"{CONNECT_FLOW_COOKIE_PREFIX}{handle}" + + +def _cookie_path_and_secure(request: Request) -> tuple[str, bool]: + parsed = urlparse(get_request_base_url(request)) + return parsed.path or "/", parsed.scheme == "https" + + +def _append_query_params(url: str, params: dict[str, str]) -> str: + parsed = urlparse(url) + query = parse_qsl(parsed.query, keep_blank_values=True) + list(params.items()) + return urlunparse(parsed._replace(query=urlencode(query))) + + +def relative_request_url(request: Request) -> str: + """The request's own path and query as a same-origin ``return_to`` target for the + login round-trip; relative by construction, so it can never leave the gateway.""" + path = request.url.path + return f"{path}?{request.url.query}" if request.url.query else path + + +def aggregate_authorize( + request: Request, + client_id: str, + redirect_uri: str, + state: str, + code_challenge: str | None, + code_challenge_method: str | None, + response_type: str | None, + session_user_id: str | None, +) -> Response: + """The aggregate authorize verb: validate the client, require S256 PKCE, interpose + LiteLLM sign-in, and hand the browser to the connect page with the flow sealed into a + per-flow cookie. + + Validation failures respond directly with 400 and never redirect: per RFC 6749 + section 4.1.2.1 an unvalidated redirect URI must not receive an error redirect, and + once the client is at fault there is no trusted place to send the browser. + """ + client = open_gateway_dcr_client(client_id) + if client is None: + return _oauth_error(400, "invalid_client", "unknown or malformed client_id") + if redirect_uri not in client.redirect_uris: + return _oauth_error(400, "invalid_request", "redirect_uri is not registered for this client") + if response_type != "code": + return _oauth_error(400, "unsupported_response_type", "response_type must be 'code'") + if not code_challenge or code_challenge_method != "S256": + return _oauth_error( + 400, + "invalid_request", + "PKCE is required: send code_challenge with code_challenge_method=S256", + ) + if len(state) > MAX_STATE_LENGTH: + return _oauth_error(400, "invalid_request", f"state must be at most {MAX_STATE_LENGTH} characters") + base_url = get_request_base_url(request) + if session_user_id is None: + login_url = f"{base_url}/sso/key/generate?{urlencode({'return_to': relative_request_url(request)})}" + return RedirectResponse(login_url, status_code=303) + now = datetime.now(timezone.utc) + handle = secrets.token_urlsafe(24) + flow = _ConnectFlow( + user_id=session_user_id, + client_id=client_id, + redirect_uri=redirect_uri, + state=state, + code_challenge=code_challenge, + jti=secrets.token_urlsafe(24), + exp=int(now.timestamp()) + CONNECT_FLOW_TTL_SECONDS, + ) + connect_url = _append_query_params( + f"{base_url}/ui/chat/integrations", + {"connect_flow": handle, "connect_client": _origin_only(redirect_uri)}, + ) + response = RedirectResponse(connect_url, status_code=303) + path, secure = _cookie_path_and_secure(request) + response.set_cookie( + key=_flow_cookie_name(handle), + value=_seal(_UNPREFIXED, flow), + max_age=CONNECT_FLOW_TTL_SECONDS, + path=path, + secure=secure, + httponly=True, + samesite="lax", + ) + return response + + +def _origin_only(url: str) -> str: + """Scheme+host for display on the connect page; never the full redirect URI, whose + path or query could carry values that do not belong in a page URL or logs.""" + parsed = urlparse(url) + return f"{parsed.scheme}://{parsed.netloc}" if parsed.netloc else "" + + +async def complete_connect_flow( + request: Request, + flow_handle: str, + session_user_id: str | None, + cache: DualCache, +) -> Response: + """The deliberate finish step of the connect flow: mint the gateway authorization + code and send the browser back to the client. + + Reached by POST so a cross-site GET cannot trigger it, and bound to the HttpOnly + per-flow cookie plus an exact match between the signed-in user and the user sealed + into the flow: a link crafted by another party dies here with ``access_denied`` + instead of minting a code for the victim's identity. The flow is single-use (an atomic + claim on its ``jti``), so a double-submit cannot mint two codes from one sign-in. + """ + sealed_flow = request.cookies.get(_flow_cookie_name(flow_handle)) + if sealed_flow is None: + return _oauth_error(400, "invalid_request", "unknown or expired connect flow") + flow = _open_sealed(sealed_flow, _UNPREFIXED, _ConnectFlow, _CONNECT_FLOW_DEBUG_KEY) + if flow is None: + return _oauth_error(400, "invalid_request", "unknown or expired connect flow") + now = datetime.now(timezone.utc) + if now.timestamp() >= flow.exp: + return _oauth_error(400, "invalid_request", "the connect flow has expired; restart the connection") + if session_user_id is None: + return _oauth_error(401, "login_required", "sign in to LiteLLM to finish connecting") + if session_user_id != flow.user_id: + return _oauth_error(403, "access_denied", "the signed-in user does not match this connect flow") + if not await _SingleUseGuard(cache).claim( + f"{_USED_FLOW_CACHE_PREFIX}{flow.jti}", CONNECT_FLOW_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS + ): + return _oauth_error(400, "invalid_request", "this connect flow was already completed; restart the connection") + code = _seal( + GATEWAY_AUTH_CODE_PREFIX, + _GatewayAuthCode( + user_id=flow.user_id, + client_id=flow.client_id, + redirect_uri=flow.redirect_uri, + code_challenge=flow.code_challenge, + jti=secrets.token_urlsafe(24), + iat=int(now.timestamp()), + exp=int(now.timestamp()) + GATEWAY_AUTH_CODE_TTL_SECONDS, + ), + ) + params = {"code": code, **({"state": flow.state} if flow.state else {})} + response = RedirectResponse(_append_query_params(flow.redirect_uri, params), status_code=303) + path, secure = _cookie_path_and_secure(request) + response.delete_cookie(key=_flow_cookie_name(flow_handle), path=path, secure=secure, httponly=True, samesite="lax") + return response + + +def _pkce_verifier_matches(code_verifier: str, code_challenge: str) -> bool: + """RFC 7636 S256 verification, total over hostile input. The comparison is over bytes + so a non-ASCII ``code_challenge`` (which reaches here unvalidated from the client's + authorize request) simply fails to match instead of raising ``TypeError`` the way + ``hmac.compare_digest`` does on two ``str`` with non-ASCII content. The verifier is + ASCII per spec; a compliant client's challenge is base64url and matches.""" + digest = hashlib.sha256(code_verifier.encode("ascii", "replace")).digest() + computed = urlsafe_b64encode(digest).rstrip(b"=") + return hmac.compare_digest(computed, code_challenge.encode("utf-8")) + + +class _SingleUseGuard: + """Atomic single-use claim for a one-time id (an auth-code, connect-flow ``jti``, or refresh-token + ``jti``) over the injected proxy cache. + + Uses an atomic increment rather than a get-then-set: two concurrent redemptions of the same id + cannot both observe "unused", because exactly one increment returns 1. The claim IS the gate, so it + fails closed. Crucially, the increment must be recorded in a backend SHARED across replicas, or the + single-use property is per-worker only (each replica's in-memory counter returns 1, so a captured + id replays through a different worker): + + - When a Redis backend is configured it is the SOLE authority: the claim goes straight to Redis + (``INCR`` is atomic across replicas), and any Redis fault fails the claim CLOSED — it never falls + back to the per-worker in-memory count (``DualCache.async_increment_cache`` does fall back, which + is exactly the replay window this avoids). + - With no Redis configured (single-replica) the in-memory increment is authoritative within the one + process. A multi-worker deployment must run Redis for the guarantee to hold across workers. + + The id's own TTL is the outer bound. For the auth code, PKCE binding is the primary defense against + interception; this makes the RFC 6749 4.1.2 single-use property reliable on top of it.""" + + def __init__(self, cache: DualCache) -> None: + self._cache = cache + + async def claim(self, key: str, ttl_seconds: int) -> bool: + """Atomically claim ``key``. ``True`` iff this caller is the first (increment to 1); ``False`` + on a replay (>1) or when the claim could not be recorded in the shared backend (fail closed).""" + from litellm.proxy.proxy_server import redis_usage_cache # noqa: PLC0415 # circular import at module load + + # Resolve the shared authority HERE rather than trusting the injected cache: callers pass + # user_api_key_cache, which only carries a redis_cache when enable_redis_auth_cache is set + # (off by default), so a guard that read its injected cache silently degraded every claim to + # a per-worker count on a stock multi-worker deployment. redis_usage_cache is the store the + # proxy already treats as cross-worker, so no call site can wire the guarantee away. + redis_cache = redis_usage_cache or getattr(self._cache, "redis_cache", None) + if redis_cache is not None: + # Shared, atomic authority for multi-replica deployments. Claim ONLY against Redis and fail + # CLOSED on any Redis fault (async_increment re-raises) rather than fall back to the + # per-worker in-memory count, which would let each replica observe count==1 and replay the id. + try: + count = await redis_cache.async_increment(key, 1, ttl=ttl_seconds) + except Exception as e: # noqa: BLE001 # ANY Redis fault fails the single-use claim closed + verbose_logger.warning( + "mcp gateway single-use claim: shared cache backend unavailable, failing closed: %s", e + ) + return False + return count == 1 + # No shared backend configured (single-replica): the in-memory increment is authoritative. + count = await self._cache.async_increment_cache(key, 1, ttl=ttl_seconds, local_only=True) + return count == 1 + + +def _session_token_pair(principal: SessionPrincipal, keys: SessionKeys, now: datetime) -> Response: + access = mint_session_token(principal, keys, now) + refresh = mint_session_refresh_token(principal, keys, now) + if not isinstance(access, MintedSessionToken) or not isinstance(refresh, MintedSessionToken): + return _oauth_error(500, "server_error", "failed to mint the session credential") + return JSONResponse( + status_code=200, + content={ + "access_token": access.token.get_secret_value(), + "token_type": "Bearer", + "expires_in": int((access.expires_at - now).total_seconds()), + "refresh_token": refresh.token.get_secret_value(), + }, + headers=TOKEN_NO_CACHE_HEADERS, + ) + + +def _reload_failure_response(failure: ReloadUserFailure) -> Response: + """Map the live-user revalidation failure onto its OAuth error, exhaustively, so a new + ``ReloadUserFailure`` member is a type error here rather than silently 400ing.""" + match failure: + case "unavailable": + return _oauth_error(503, "temporarily_unavailable", "the gateway database is unavailable; retry") + case "unresolvable": + return _oauth_error(500, "server_error", "the gateway is not configured to resolve users") + case "no_active_key": + return _oauth_error(400, "invalid_grant", "the user for this grant is no longer active") + case _: + assert_never(failure) + + +async def aggregate_token( + request: Request, + grant_type: str, + code: str | None, + redirect_uri: str | None, + client_id: str, + code_verifier: str | None, + refresh_token: str | None, + master_key: str | None, + reload_user: ReloadUser, + cache: DualCache, +) -> Response: + """The aggregate token verb: authorization_code and refresh_token grants for the + identity-only session pair. Every path re-validates the litellm user live before + minting, so a deactivated user cannot obtain or renew a session.""" + if master_key is None: + verbose_logger.error("mcp_gateway_dcr token grant rejected: no master_key configured") + return _oauth_error(500, "server_error", "the gateway has no master key configured") + keys = session_keys_from_master_key(master_key) + now = datetime.now(timezone.utc) + if grant_type == "authorization_code": + return await _authorization_code_grant( + code=code, + redirect_uri=redirect_uri, + client_id=client_id, + code_verifier=code_verifier, + keys=keys, + now=now, + reload_user=reload_user, + guard=_SingleUseGuard(cache), + ) + if grant_type == "refresh_token": + return await _refresh_token_grant( + refresh_token=refresh_token, + client_id=client_id, + keys=keys, + now=now, + reload_user=reload_user, + guard=_SingleUseGuard(cache), + ) + return _oauth_error(400, "unsupported_grant_type", "grant_type must be authorization_code or refresh_token") + + +async def _authorization_code_grant( + code: str | None, + redirect_uri: str | None, + client_id: str, + code_verifier: str | None, + keys: SessionKeys, + now: datetime, + reload_user: ReloadUser, + guard: _SingleUseGuard, +) -> Response: + if not code or not redirect_uri or not code_verifier: + return _oauth_error(400, "invalid_request", "code, redirect_uri, and code_verifier are required") + if not MIN_CODE_VERIFIER_LENGTH <= len(code_verifier) <= MAX_CODE_VERIFIER_LENGTH: + return _oauth_error(400, "invalid_request", "code_verifier must be 43 to 128 characters (RFC 7636)") + parsed = _open_sealed(code, GATEWAY_AUTH_CODE_PREFIX, _GatewayAuthCode, _AUTH_CODE_DEBUG_KEY) + if parsed is None: + return _oauth_error(400, "invalid_grant", "the authorization code is invalid") + if now.timestamp() >= parsed.exp: + return _oauth_error(400, "invalid_grant", "the authorization code has expired") + if client_id != parsed.client_id or redirect_uri != parsed.redirect_uri: + return _oauth_error(400, "invalid_grant", "the authorization code was issued to a different client") + if not _pkce_verifier_matches(code_verifier, parsed.code_challenge): + return _oauth_error(400, "invalid_grant", "PKCE verification failed") + # Revalidate the user BEFORE claiming the code, so a transient DB outage (a retryable + # 503) does not consume a still-valid code and force the client to restart sign-in. + failure = await reload_user(parsed.user_id) + if failure is not None: + return _reload_failure_response(failure) + # Atomic single-use claim is the gate: on a concurrent double-redeem exactly one caller + # wins, and a claim that cannot be recorded fails closed. + if not await guard.claim( + f"{_USED_CODE_CACHE_PREFIX}{parsed.jti}", GATEWAY_AUTH_CODE_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS + ): + return _oauth_error(400, "invalid_grant", "the authorization code was already used") + return _session_token_pair(SessionPrincipal(user_id=parsed.user_id, client_id=client_id), keys, now) + + +async def _refresh_token_grant( + refresh_token: str | None, + client_id: str, + keys: SessionKeys, + now: datetime, + reload_user: ReloadUser, + guard: _SingleUseGuard, +) -> Response: + if not refresh_token: + return _oauth_error(400, "invalid_request", "refresh_token is required") + opened = open_session_refresh_bearer(refresh_token, keys, now, expected_client_id=client_id) + if not isinstance(opened, SessionRefreshOpened): + return _oauth_error(400, "invalid_grant", "the refresh token is invalid for this client") + failure = await reload_user(opened.principal.user_id) + if failure is not None: + return _reload_failure_response(failure) + # Refresh-token rotation (OAuth 2.0 Security BCP section 4.13): the presented refresh token is + # single-use. Claim its jti before issuing the replacement pair, so a captured or replayed + # refresh token cannot mint a second pair after the legitimate holder rotated. Claimed AFTER + # user revalidation so a transient DB 503 does not burn a still-valid token; a claim that + # cannot be recorded fails closed, exactly like the authorization-code path. + if not await guard.claim( + f"{_USED_REFRESH_CACHE_PREFIX}{opened.jti}", SESSION_REFRESH_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS + ): + return _oauth_error(400, "invalid_grant", "the refresh token was already used") + return _session_token_pair(opened.principal, keys, now) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 90b70dd01f2..0ee74960293 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -13,6 +13,7 @@ import json import os import re import time +from collections.abc import Sequence from contextlib import asynccontextmanager from typing import Any, AsyncIterator, Callable, Literal, Optional, Union, cast from urllib.parse import urlparse @@ -49,6 +50,10 @@ from litellm.litellm_core_utils.url_utils import SSRFError, async_safe_get from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, + _is_mcp_admitted_user_subject, +) +from litellm.proxy._experimental.mcp_server.elicitation_handler import ( + MCP_ELICITATION_AVAILABLE, ) from litellm.proxy._experimental.mcp_server.exceptions import ( MCPServerListError, @@ -59,17 +64,14 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ( raise_classified_list_failure, upstream_auth_challenge, ) -from litellm.proxy._experimental.mcp_server.elicitation_handler import ( - MCP_ELICITATION_AVAILABLE, -) -from litellm.proxy._experimental.mcp_server.sampling_handler import ( - MCP_SAMPLING_AVAILABLE, -) from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( MCPPerUserTokenCache, mcp_per_user_token_cache, resolve_mcp_auth, ) +from litellm.proxy._experimental.mcp_server.oauth_utils import ( + _redact_mcp_resource_url, +) from litellm.proxy._experimental.mcp_server.outbound_credentials import ( Error, Ok, @@ -100,6 +102,9 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( ServerSpec, TokenExchangeConfig, ) +from litellm.proxy._experimental.mcp_server.sampling_handler import ( + MCP_SAMPLING_AVAILABLE, +) from litellm.proxy._experimental.mcp_server.utils import ( MCP_TOOL_PREFIX_SEPARATOR, MCPMissingUserEnvVarsError, @@ -143,11 +148,9 @@ from litellm.types.mcp_server.mcp_server_manager import ( from litellm.types.utils import CallTypes try: - from mcp.shared.tool_name_validation import ( - validate_tool_name, # pyright: ignore[reportAssignmentType] - ) from mcp.shared.tool_name_validation import ( SEP_986_URL, + validate_tool_name, # pyright: ignore[reportAssignmentType] ) except ImportError: from pydantic import BaseModel @@ -408,6 +411,88 @@ def _restrict_discovery_to_corroborated_authorization_server( return metadata.model_copy(update={"token_url": None, "registration_url": None}) +def _redacted_origin_list(urls: Sequence[str]) -> str: + return ", ".join(_redact_mcp_resource_url(url) or "" for url in urls) + + +def _sanitized_error_text(exc: Exception) -> str: + return re.sub(r"https?://\S+", "", str(exc))[:200] + + +def _discovery_failure_leaves_needs_unresolved( + *, + needs_authorization_url: bool, + needs_token_url: bool, + manual_authorization_url: str | None, + manual_token_url: str | None, +) -> bool: + return (needs_authorization_url and not manual_authorization_url) or (needs_token_url and not manual_token_url) + + +def _warn_oauth_endpoints_unresolved( + *, + server_ref: str, + server_url: str | None, + discovery_attempted: bool, + issuer_anchored: bool, + metadata: MCPOAuthMetadata | None, + needs_authorization_url: bool, + needs_token_url: bool, + manual_authorization_url: str | None, + manual_token_url: str | None, +) -> None: + """Log one actionable warning when a server that depends on OAuth endpoint discovery finishes a + build without the endpoints that its flows need (LIT-4658). + + This is the operator-facing signal for a misconfigured server url: discovery failures themselves + are logged where they happen (``_descovery_metadata``), and this names WHICH server is affected, + which endpoints stayed unresolved after manual configuration was considered, and the remedies. + Scopes never trigger the warning on their own: scope-less metadata is normal for many servers and + warning on it every rebuild would be noise. Callers own the per-flow policy of which endpoints + are needed (client_credentials never needs authorization_url; OBO needs only token_url); the + issuer-anchored arm is excluded here because it has its own RFC 8414 §3.3 warning. + """ + if issuer_anchored: + return + unresolved = tuple( + field + for field, needed, value in ( + ( + "authorization_url", + needs_authorization_url, + manual_authorization_url or (metadata.authorization_url if metadata else None), + ), + ( + "token_url", + needs_token_url, + manual_token_url or (metadata.token_url if metadata else None), + ), + ) + if needed and not value + ) + if not unresolved: + return + if discovery_attempted: + verbose_logger.warning( + "MCP server %s: OAuth endpoint discovery left %s unresolved (server url origin: %s). OAuth flows " + "that need them will fail with 'not configured' errors until they resolve. Check the preceding " + "'MCP OAuth' log lines for why discovery failed, verify the configured server url, or set the " + "unresolved endpoint urls manually, or set issuer to discover them from the identity provider " + "(RFC 8414)", + server_ref, + ", ".join(unresolved), + _redact_mcp_resource_url(server_url) or "", + ) + return + verbose_logger.warning( + "MCP server %s uses OAuth but has no discovery source (no server url or pinned issuer), and %s not " + "set manually. Set the missing endpoint urls on the server, or set issuer to discover them from the " + "identity provider (RFC 8414)", + server_ref, + " and ".join(unresolved) + (" is" if len(unresolved) == 1 else " are"), + ) + + def invalidate_user_env_vars_cache(user_id: str, server_id: str) -> None: """Drop a cached entry after the user stores or clears their env var values so the next request reads the fresh value instead of a stale one.""" @@ -884,10 +969,10 @@ def _create_sampling_callback(user_api_key_auth: Optional[Any] = None): return None async def _sampling_callback(context, params): + import litellm from litellm.proxy._experimental.mcp_server.sampling_handler import ( handle_sampling_create_message, ) - import litellm from litellm.proxy._experimental.mcp_server.server import ( get_active_auth_context, ) @@ -1284,6 +1369,15 @@ class MCPServerManager: should_discover = _has_oauth_discovery_source(server_url, use_issuer_anchor) and ( is_discovery_auth_type or obo_needs_discovery ) + config_oauth2_flow = server_config.get("oauth2_flow", None) + needs_authorization_url = is_discovery_auth_type and config_oauth2_flow != "client_credentials" + needs_token_url = is_discovery_auth_type or obo_needs_discovery + warn_on_empty_discovery = _discovery_failure_leaves_needs_unresolved( + needs_authorization_url=needs_authorization_url, + needs_token_url=needs_token_url, + manual_authorization_url=manual_authorization_url, + manual_token_url=manual_token_url, + ) if not should_discover: mcp_oauth_metadata = None elif use_issuer_anchor and manual_issuer is not None: @@ -1292,6 +1386,7 @@ class MCPServerManager: mcp_oauth_metadata = await self._descovery_metadata( server_url=server_url, allow_origin_fallback=is_discovery_auth_type, + warn_when_no_metadata=warn_on_empty_discovery, ) if use_issuer_anchor: @@ -1326,7 +1421,6 @@ class MCPServerManager: ) effective_issuer = manual_issuer or discovered_issuer - config_oauth2_flow = server_config.get("oauth2_flow", None) if auth_type == MCPAuth.oauth2 and config_oauth2_flow not in ( "client_credentials", "authorization_code", @@ -1358,6 +1452,18 @@ class MCPServerManager: "authorization-code flow." ) + _warn_oauth_endpoints_unresolved( + server_ref=server_name or server_id, + server_url=server_url, + discovery_attempted=should_discover, + issuer_anchored=use_issuer_anchor, + metadata=gated_oauth_metadata, + needs_authorization_url=needs_authorization_url, + needs_token_url=needs_token_url, + manual_authorization_url=manual_authorization_url, + manual_token_url=manual_token_url, + ) + new_server = MCPServer( server_id=server_id, name=name_for_prefix, @@ -1485,14 +1591,12 @@ class MCPServerManager: from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( build_input_schema, create_tool_function, + load_openapi_spec_async, + resolve_operation_params, ) from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( get_base_url as get_openapi_base_url, ) - from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - load_openapi_spec_async, - resolve_operation_params, - ) from litellm.proxy._experimental.mcp_server.tool_registry import ( global_mcp_tool_registry, ) @@ -1681,10 +1785,20 @@ class MCPServerManager: scopes: Optional[list[str]], token_exchange_endpoint: Optional[str], ) -> Optional[MCPOAuthMetadata]: + obo_needs_discovery = self._obo_needs_endpoint_discovery(auth_type, token_exchange_endpoint, manual_token_url) + needs_authorization_url = ( + is_discovery_auth_type and getattr(mcp_server, "oauth2_flow", None) != "client_credentials" + ) + needs_token_url = is_discovery_auth_type or obo_needs_discovery + warn_on_empty_discovery = _discovery_failure_leaves_needs_unresolved( + needs_authorization_url=needs_authorization_url, + needs_token_url=needs_token_url, + manual_authorization_url=manual_authorization_url, + manual_token_url=manual_token_url, + ) has_all_upstream_oauth_fields = bool(manual_authorization_url and manual_token_url and scopes) needs_discovery = _has_oauth_discovery_source(server_url, use_issuer_anchor) and ( - (is_discovery_auth_type and not has_all_upstream_oauth_fields) - or self._obo_needs_endpoint_discovery(auth_type, token_exchange_endpoint, manual_token_url) + (is_discovery_auth_type and not has_all_upstream_oauth_fields) or obo_needs_discovery ) if not needs_discovery: mcp_oauth_metadata: Optional[MCPOAuthMetadata] = None @@ -1694,24 +1808,32 @@ class MCPServerManager: mcp_oauth_metadata = await self._descovery_metadata( server_url=server_url, # type: ignore[arg-type] allow_origin_fallback=is_discovery_auth_type, - ) - if needs_discovery and not use_issuer_anchor and mcp_oauth_metadata is None: - verbose_logger.warning( - "MCP OAuth discovery yielded no metadata for server %s (%s); " - "OAuth endpoints/scopes stay unresolved until a rebuild succeeds", - mcp_server.server_id, - server_url, + warn_when_no_metadata=warn_on_empty_discovery, ) if use_issuer_anchor: return mcp_oauth_metadata - if is_discovery_auth_type: - return _restrict_discovery_to_corroborated_authorization_server( + gated_metadata = ( + _restrict_discovery_to_corroborated_authorization_server( mcp_oauth_metadata, manual_authorization_url, mcp_server.server_id, bool(getattr(mcp_server, "dcr_bridge", None)), ) - return mcp_oauth_metadata + if is_discovery_auth_type + else mcp_oauth_metadata + ) + _warn_oauth_endpoints_unresolved( + server_ref=mcp_server.alias or mcp_server.server_name or mcp_server.server_id, + server_url=server_url, + discovery_attempted=needs_discovery, + issuer_anchored=False, + metadata=gated_metadata, + needs_authorization_url=needs_authorization_url, + needs_token_url=needs_token_url, + manual_authorization_url=manual_authorization_url, + manual_token_url=manual_token_url, + ) + return gated_metadata async def build_mcp_server_from_table( self, @@ -2197,6 +2319,56 @@ class MCPServerManager: return [server_id for server_id in submitted_server_ids if self.get_mcp_server_by_id(server_id) is not None] + async def operator_open_server_ids( + self, + user_api_key_auth: UserAPIKeyAuth | None = None, + *, + allow_all_server_ids: list[str] | None = None, + submitted_server_ids: list[str] | None = None, + ) -> set: + """Servers reachable through OPEN channels rather than a grant: operator-opened + ``allow_all_keys`` servers, plus the caller's own active BYOM submissions when the caller + carries no explicit ``mcp_servers`` scope. + + The single owner of that question for BOTH axes. The server union in + ``get_allowed_mcp_servers`` adds these ids, and the admitted subject's tool resolution asks + the same question to treat an open-channel server as default-open for tools — exactly how a + virtual key experiences it. Encoding the channel membership twice is how a server ends up + listable but uninvokable. + + Empty inside a toolset scope: toolset_mcp_route / dynamic_mcp_route set + ``_mcp_active_toolset_id`` before calling the handler, pinning the request to the toolset's + own servers (checking op.mcp_toolsets==[] instead would false-positive on DB-default rows + where Postgres initialises the column to ARRAY[]::TEXT[]). + + ``allow_all_server_ids`` / ``submitted_server_ids`` are injectable so the server union, + which precomputes both for its fallback path, does not compute them twice.""" + from litellm.proxy._experimental.mcp_server.mcp_context import ( # noqa: PLC0415 + _mcp_active_toolset_id, + ) + + if _mcp_active_toolset_id.get() is not None: + return set() + if allow_all_server_ids is None: + allow_all_server_ids = self.get_allow_all_keys_server_ids() + open_ids = set(allow_all_server_ids) + key_object_permission = user_api_key_auth.object_permission if user_api_key_auth else None + # "Explicitly scoped, so do not widen with BYOM" is a rule about a CREDENTIAL that carries + # its own mcp_servers list. It does not describe a keyless admitted subject: its + # object_permission is the user's own row, whose mcp_servers column is [] by DB default, so + # applying this rule would hide almost every admitted user's OWN submitted servers. Their + # submissions are theirs by authorship, and their scope comes from the per-source union. + has_explicit_object_permission = ( + not _is_mcp_admitted_user_subject(user_api_key_auth) + and key_object_permission is not None + and (key_object_permission.mcp_servers is not None) + ) + if not has_explicit_object_permission: + if submitted_server_ids is None: + submitted_server_ids = await self._get_active_submitted_mcp_server_ids_for_user(user_api_key_auth) + open_ids.update(submitted_server_ids) + return open_ids + async def get_allowed_mcp_servers(self, user_api_key_auth: Optional[UserAPIKeyAuth] = None) -> list[str]: """ Get the allowed MCP Servers for the user. @@ -2210,11 +2382,22 @@ class MCPServerManager: allow_all_server_ids = self.get_allow_all_keys_server_ids() + # A keyless admitted subject is resolved per grant source, and channel decisions that are + # absolute for a scoped KEY credential are not absolute for it: its own opt-out silences its + # own source (handled per source in the resolver), never its teams' grants, and its admin + # role does not swallow the grant model — a session bearer is a third-party client + # credential, not the dashboard, so an admin signing in through the connect flow gets their + # grants like anyone else rather than handing the client the full registry ahead of every + # per-team org ceiling. + is_admitted_subject = _is_mcp_admitted_user_subject(user_api_key_auth) + # The key explicitly opted out of every MCP server. Return zero before # layering on allow_all_keys or submitted servers so the opt-out is absolute. key_object_permission = user_api_key_auth.object_permission if user_api_key_auth else None - if key_object_permission is not None and ( - SpecialMCPServerNames.no_mcp_servers.value in (key_object_permission.mcp_servers or []) + if ( + not is_admitted_subject + and key_object_permission is not None + and (SpecialMCPServerNames.no_mcp_servers.value in (key_object_permission.mcp_servers or [])) ): return [] @@ -2234,8 +2417,14 @@ class MCPServerManager: ) try: - # If admin but NO explicit object permission, get all servers - if user_api_key_auth and _user_has_admin_view(user_api_key_auth) and not has_explicit_object_permission: + # If admin but NO explicit object permission, get all servers (never for an admitted + # subject — see is_admitted_subject above) + if ( + user_api_key_auth + and not is_admitted_subject + and _user_has_admin_view(user_api_key_auth) + and not has_explicit_object_permission + ): verbose_logger.debug("Admin user without explicit object_permission - returning all servers") return list(self.get_registry().keys()) @@ -2243,20 +2432,14 @@ class MCPServerManager: allowed_mcp_servers = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth) verbose_logger.debug(f"Allowed MCP Servers for user api key auth: {allowed_mcp_servers}") combined_servers = set(allowed_mcp_servers) - # Only skip allow_all_keys servers when the request is inside a toolset - # scope. toolset_mcp_route / dynamic_mcp_route set _mcp_active_toolset_id - # before calling the handler — that ContextVar is the reliable signal. - # Using op.mcp_toolsets==[] would false-positive on DB-default rows where - # Postgres initialises the column to ARRAY[]::TEXT[]. - from litellm.proxy._experimental.mcp_server.mcp_context import ( # noqa: PLC0415 - _mcp_active_toolset_id, + combined_servers.update( + await self.operator_open_server_ids( + user_api_key_auth, + allow_all_server_ids=allow_all_server_ids, + submitted_server_ids=submitted_server_ids, + ) ) - in_toolset_scope = _mcp_active_toolset_id.get() is not None - if not in_toolset_scope: - combined_servers.update(allow_all_server_ids) - combined_servers.update(submitted_server_ids) - # For anonymous callers (no user_id, no role), also surface any # servers the operator has opted into upstream-delegated auth. # These servers handle their own auth at the upstream level, so @@ -2903,9 +3086,7 @@ class MCPServerManager: ) ): spec = None - auth_value = ( - await resolve_mcp_auth(server, mcp_auth_header, subject_token=subject_token) if spec is None else None - ) + auth_value = await resolve_mcp_auth(server, mcp_auth_header) if spec is None else None # Create sampling and elicitation callbacks for this client sampling_cb = _create_sampling_callback(user_api_key_auth=user_api_key_auth) if server.allow_sampling else None @@ -3430,6 +3611,7 @@ class MCPServerManager: server_url: str, *, allow_origin_fallback: bool = True, + warn_when_no_metadata: bool = False, ) -> Optional[MCPOAuthMetadata]: """Discover OAuth metadata by following RFC 9728 (protected resource metadata discovery). @@ -3438,8 +3620,32 @@ class MCPServerManager: it (a human sees the redirect), but token_exchange (OBO) sets it False so the gateway never exchanges a subject token against an endpoint it inferred rather than one explicitly configured or authoritatively advertised via RFC 9728 / RFC 8414. - """ + ``warn_when_no_metadata`` makes an all-empty result log one WARNING with the per-step attempt + outcomes (LIT-4658), so a misconfigured server url is diagnosable from default-level logs. The + server loaders set it; the issuer-anchored resource-scopes lookup keeps it off because empty + scopes are not a fault there. + """ + metadata, attempts = await self._discover_metadata_recording_attempts( + server_url, allow_origin_fallback=allow_origin_fallback + ) + if metadata is None and warn_when_no_metadata: + verbose_logger.warning( + "MCP OAuth endpoint discovery against %s found no authorization server metadata. Attempts: %s. " + "The MCP server url may be misconfigured, or the upstream may not support OAuth discovery " + "(RFC 9728 / RFC 8414)", + _redact_mcp_resource_url(server_url) or "", + "; ".join(attempts) if attempts else "none recorded", + ) + return metadata + + async def _discover_metadata_recording_attempts( + self, + server_url: str, + *, + allow_origin_fallback: bool, + ) -> tuple[MCPOAuthMetadata | None, tuple[str, ...]]: + origin = _redact_mcp_resource_url(server_url) or "" try: client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) response = await client.get(server_url) @@ -3452,67 +3658,112 @@ class MCPServerManager: if metadata is None and not resource_scopes and authorization_servers and response.status_code == 200: verbose_logger.warning( "MCP OAuth discovery for %s received 200 OK without RFC 9728 challenge and no discoverable authorization metadata.", - server_url, + origin, ) + attempts = ( + f"GET {origin}: HTTP {response.status_code} (no RFC 9728 challenge)", + *( + ("well-known protected-resource lookup found no authorization servers",) + if not authorization_servers + else () + ), + *( + (f"authorization server metadata fetch failed for: {_redacted_origin_list(authorization_servers)}",) + if authorization_servers and metadata is None + else () + ), + ) if metadata is None and resource_scopes: - return MCPOAuthMetadata(scopes=resource_scopes) + return MCPOAuthMetadata(scopes=resource_scopes), attempts if metadata is not None and resource_scopes: metadata.scopes = resource_scopes - return metadata + return metadata, attempts except HTTPStatusError as exc: - verbose_logger.debug( - "MCP OAuth discovery for %s received status error: %s", - server_url, - exc, - ) - - header_value: Optional[str] = None - if exc.response is not None: - header_value = exc.response.headers.get("WWW-Authenticate") or exc.response.headers.get( - "www-authenticate" - ) - - resource_metadata_url, scopes = self._parse_www_authenticate_header(header_value) - - authorization_servers = [] - resource_scopes = None - if resource_metadata_url: - ( - authorization_servers, - resource_scopes, - ) = await self._fetch_oauth_metadata_from_resource(resource_metadata_url, server_url) - else: - ( - authorization_servers, - resource_scopes, - ) = await self._attempt_well_known_discovery(server_url) - - metadata = None - used_origin_fallback = False - if allow_origin_fallback and not authorization_servers: - try: - parsed_url = urlparse(server_url) - if parsed_url.scheme and parsed_url.netloc: - authorization_servers = [f"{parsed_url.scheme}://{parsed_url.netloc}"] - used_origin_fallback = True - except Exception: - authorization_servers = [] - - if authorization_servers: - metadata = await self._fetch_authorization_server_metadata(authorization_servers, server_url) - if metadata is not None and used_origin_fallback: - metadata.from_origin_fallback = True - - preferred_scopes = scopes or resource_scopes - if metadata is None and preferred_scopes: - metadata = MCPOAuthMetadata(scopes=preferred_scopes) - elif metadata is not None and preferred_scopes: - metadata.scopes = preferred_scopes - - return metadata + return await self._discover_after_status_error(server_url, exc, allow_origin_fallback=allow_origin_fallback) except Exception as exc: # pragma: no cover - network/transient issues verbose_logger.debug("MCP OAuth discovery failed for %s: %s", server_url, exc) - return None + return None, (f"GET {origin}: {type(exc).__name__}: {_sanitized_error_text(exc)}",) + + async def _discover_after_status_error( + self, + server_url: str, + exc: HTTPStatusError, + *, + allow_origin_fallback: bool, + ) -> tuple[MCPOAuthMetadata | None, tuple[str, ...]]: + origin = _redact_mcp_resource_url(server_url) or "" + verbose_logger.debug( + "MCP OAuth discovery for %s received status error: %s", + server_url, + exc, + ) + + header_value: Optional[str] = None + if exc.response is not None: + header_value = exc.response.headers.get("WWW-Authenticate") or exc.response.headers.get("www-authenticate") + status_attempt = ( + f"GET {origin}: HTTP {exc.response.status_code}" + if exc.response is not None + else f"GET {origin}: status error" + ) + + resource_metadata_url, scopes = self._parse_www_authenticate_header(header_value) + + authorization_servers = [] + resource_scopes = None + if resource_metadata_url: + ( + authorization_servers, + resource_scopes, + ) = await self._fetch_oauth_metadata_from_resource(resource_metadata_url, server_url) + lookup_attempt = ( + None + if authorization_servers + else "challenge-advertised resource metadata yielded no authorization servers" + ) + else: + ( + authorization_servers, + resource_scopes, + ) = await self._attempt_well_known_discovery(server_url) + lookup_attempt = ( + None + if authorization_servers + else "no challenge-advertised resource metadata; well-known protected-resource lookup found no authorization servers" + ) + + metadata = None + used_origin_fallback = False + if allow_origin_fallback and not authorization_servers: + try: + parsed_url = urlparse(server_url) + if parsed_url.scheme and parsed_url.netloc: + authorization_servers = [f"{parsed_url.scheme}://{parsed_url.netloc}"] + used_origin_fallback = True + except Exception: + authorization_servers = [] + + fallback_attempt = None + if authorization_servers: + metadata = await self._fetch_authorization_server_metadata(authorization_servers, server_url) + if metadata is not None and used_origin_fallback: + metadata.from_origin_fallback = True + if metadata is None: + fallback_attempt = ( + f"origin fallback: no authorization server metadata at {origin}" + if used_origin_fallback + else f"authorization server metadata fetch failed for: {_redacted_origin_list(authorization_servers)}" + ) + + attempts = tuple(entry for entry in (status_attempt, lookup_attempt, fallback_attempt) if entry) + + preferred_scopes = scopes or resource_scopes + if metadata is None and preferred_scopes: + return MCPOAuthMetadata(scopes=preferred_scopes), attempts + if metadata is not None and preferred_scopes: + metadata.scopes = preferred_scopes + + return metadata, attempts def _parse_www_authenticate_header(self, header_value: Optional[str]) -> tuple[Optional[str], Optional[list[str]]]: if not header_value: diff --git a/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py b/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py index 43fe3999291..a6acaf8e1d6 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py +++ b/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py @@ -26,7 +26,6 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, ) -from litellm.proxy._experimental.mcp_server.auth import token_exchange from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import ( build_token_endpoint_client_auth, ) @@ -58,17 +57,12 @@ class MCPOAuth2TokenCache(InMemoryCache): def _has_client_credentials_config(server: "MCPServer") -> bool: return bool(server.client_id and server.client_secret and server.token_url) - async def async_get_token( - self, - server: "MCPServer", - *, - require_client_credentials_flow: bool = True, - ) -> Optional[str]: + async def async_get_token(self, server: "MCPServer") -> Optional[str]: """Return a valid access token, fetching or refreshing as needed. Returns ``None`` when the server lacks client credentials config. """ - if require_client_credentials_flow and not server.has_client_credentials: + if not server.has_client_credentials: return None if not self._has_client_credentials_config(server): return None @@ -278,36 +272,16 @@ mcp_per_user_token_cache = MCPPerUserTokenCache() async def resolve_mcp_auth( server: "MCPServer", mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None, - subject_token: Optional[str] = None, ) -> Optional[Union[str, Dict[str, str]]]: """Resolve the auth value for an MCP server. Priority: 1. ``mcp_auth_header`` — per-request/per-user override - 2. OAuth2 Token Exchange (OBO / RFC 8693) — exchange user token for scoped token - 3. OAuth2 client_credentials token — auto-fetched and cached - 4. ``server.authentication_token`` — static token from config/DB + 2. OAuth2 client_credentials token — auto-fetched and cached + 3. ``server.authentication_token`` — static token from config/DB """ if mcp_auth_header: return mcp_auth_header - if server.has_token_exchange_config: - if subject_token: - return await token_exchange.mcp_token_exchange_handler.exchange_token(subject_token, server) - # No subject_token — fall back to client_credentials using the same client - # credentials and token_url so M2M scenarios still work. - if server.client_id and server.client_secret and server.token_url: - return await mcp_oauth2_token_cache.async_get_token( - server, - require_client_credentials_flow=False, - ) - # OBO configured but no subject_token and missing client credentials — warn - # rather than silently proceeding unauthenticated. - verbose_logger.warning( - "MCP server '%s' is configured for token exchange (OBO) but no subject_token " - "was provided and client credentials (client_id/client_secret/token_url) are " - "incomplete. The request will proceed without authentication.", - server.server_id, - ) if server.has_client_credentials: return await mcp_oauth2_token_cache.async_get_token(server) return server.authentication_token diff --git a/litellm/proxy/_experimental/mcp_server/oauth_utils.py b/litellm/proxy/_experimental/mcp_server/oauth_utils.py index 53686e329bb..9b7760a30d7 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_utils.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_utils.py @@ -4,7 +4,7 @@ import os from ipaddress import ip_address from typing import Any, Dict, List, NoReturn, Optional -from urllib.parse import ParseResult, urlparse, urlunparse +from urllib.parse import ParseResult, urlparse, urlsplit, urlunparse, urlunsplit from fastapi import HTTPException, Request @@ -70,6 +70,29 @@ def _origin_label(scheme: str, netloc: str) -> str: return f"{scheme}://{netloc}" if netloc else f"{scheme}://" +def _redact_mcp_resource_url(url: Optional[str]) -> Optional[str]: + """Reduce an MCP server URL to its origin (scheme + host + port) for logging. + + Everything else is dropped: userinfo (``user:pass@``), the query string, the + fragment, and the path, because hosted MCP servers routinely embed the + credential in the path (e.g. ``/mcp/s/``) and this value is persisted + in spend-log metadata that a caller who can invoke the tool can read back. + Returns None when the URL has no host to identify (nothing safe to log). + """ + if not isinstance(url, str) or not url: + return None + try: + parts = urlsplit(url) + hostname = parts.hostname + port = parts.port + except ValueError: + return None + if not hostname: + return None + netloc = f"{hostname}:{port}" if port else hostname + return urlunsplit((parts.scheme, netloc, "", "", "")) or None + + def _resolve_proxy_base_url_env() -> Optional[str]: global _warned_invalid_proxy_base_url configured = os.environ.get("PROXY_BASE_URL", "").strip() @@ -343,8 +366,36 @@ def _parse_redirect_uri_for_validation(redirect_uri: str) -> ParseResult: ) -def _validate_trusted_http_redirect_shape(parsed: ParseResult) -> bool: - """Return True when ``parsed`` is an allowlisted native callback (caller may return).""" +def is_loopback_redirect_host(parsed: ParseResult) -> bool: + """True when the redirect host is loopback (RFC 8252 section 7.3). + + Shared by every redirect-URI policy in the MCP OAuth surface so that none of them + hand-rolls its own host list: a literal ``("localhost", "127.0.0.1", "::1")`` tuple + silently misses the rest of 127.0.0.0/8 and IPv6-mapped forms. + """ + host = (parsed.hostname or "").lower() + if host == "localhost": + return True + try: + return ip_address(host).is_loopback + except ValueError: + return False + + +def validate_redirect_uri_shape(parsed: ParseResult) -> bool: + """Validate redirect-URI *hygiene* and resolve allowlisted native callbacks. + + Returns True when ``parsed`` is an allowlisted native callback (the caller may accept + it outright); returns False for http/https, leaving the trust decision to the caller; + raises for a URI that no policy should ever accept (bad scheme, fragment, missing + host, userinfo, backslash in the host). + + This is deliberately separate from :func:`validate_trusted_redirect_uri`, which adds + the *first-party* trust policy (same-origin, loopback, ops allowlist) appropriate to + the proxy's own OAuth endpoints. Public dynamic-client registration accepts any https + client and relies on PKCE plus the consent screen instead, so it shares this hygiene + rule but not that trust policy. + """ if parsed.scheme not in ("http", "https"): if _matches_trusted_native_redirect_uri(parsed): return True @@ -396,14 +447,8 @@ def _trusted_redirect_uri_is_allowed( ): return True - host = (parsed.hostname or "").lower() - if host == "localhost": + if is_loopback_redirect_host(parsed): return True - try: - if ip_address(host).is_loopback: - return True - except ValueError: - pass if parsed.scheme == "https": for entry in _parse_trusted_redirect_origins(): @@ -522,7 +567,7 @@ def validate_trusted_redirect_uri(request: Request, redirect_uri: str) -> None: :func:`validate_loopback_redirect_uri`. """ parsed = _parse_redirect_uri_for_validation(redirect_uri) - if _validate_trusted_http_redirect_shape(parsed): + if validate_redirect_uri_shape(parsed): return redirect_netloc = _strip_default_port(parsed.scheme, parsed.netloc) proxy_base = _resolve_proxy_base_for_redirect(request) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_credentials.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_credentials.py index 08d5cc8b1f1..8844d8c8ad0 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_credentials.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_credentials.py @@ -149,6 +149,7 @@ class SessionRefreshOpened(BaseModel): model_config = ConfigDict(frozen=True) tag: Literal["opened"] = "opened" principal: SessionPrincipal + jti: str class SessionRefreshInvalid(BaseModel): @@ -187,4 +188,4 @@ def open_session_refresh_bearer( return SessionRefreshInvalid() if opened.principal.client_id != expected_client_id: return SessionRefreshInvalid() - return SessionRefreshOpened(principal=opened.principal) + return SessionRefreshOpened(principal=opened.principal, jti=opened.jti) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py index 9325428f049..4ccbcd1a511 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py @@ -113,10 +113,12 @@ class MintedSessionToken(BaseModel): class OpenedSessionToken(BaseModel): - """A validated session token of either kind: the principal it was minted for.""" + """A validated session token of either kind: the principal it was minted for, plus the + ``jti`` so the token endpoint can enforce single-use rotation on a refresh token.""" model_config = ConfigDict(frozen=True) principal: SessionPrincipal + jti: str class SessionTokenTooLarge(BaseModel): @@ -320,7 +322,9 @@ def _open( return SessionMalformed() if now.timestamp() >= claims.exp: return SessionExpired() - return OpenedSessionToken(principal=SessionPrincipal(user_id=claims.user_id, client_id=claims.client_id)) + return OpenedSessionToken( + principal=SessionPrincipal(user_id=claims.user_id, client_id=claims.client_id), jti=claims.jti + ) def _decode_claims( diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 94271c54f4b..26e4176e09b 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -230,16 +230,33 @@ if MCP_AVAILABLE: return server_auth return mcp_auth_header - def _get_oauth2_server_ids(allowed_server_ids: List[str]) -> Set[str]: - """Return the subset of *allowed_server_ids* whose servers use OAuth2 auth. + def _is_v1_resolved_oauth2_server(server: Optional[MCPServer]) -> bool: + """Whether this server's per-user OAuth2 token is still resolved by v1. - Used as a cheap pre-flight check to skip bulk credential fetching when no - OAuth2 servers are involved in the current request. + A server the v2 resolver owns reads its stored token from the resolver at connect + time and drops any Authorization built for it here, so the v1 lookup would be a DB + round-trip whose result is discarded. Mirrors the same guard on the protocol listing + path and in ``_resolve_oauth2_headers_for_tool_call``. + """ + from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( + to_server_spec, + ) + + if getattr(server, "auth_type", None) != MCPAuth.oauth2: + return False + return to_server_spec(server) is None + + def _v1_resolved_oauth2_server_ids(allowed_server_ids: List[str]) -> Set[str]: + """Return the subset of *allowed_server_ids* whose per-user OAuth2 token is still + resolved by v1. + + Used as a cheap pre-flight check to skip bulk credential fetching when no such + server is involved in the current request. """ return { sid for sid in allowed_server_ids - if getattr(global_mcp_server_manager.get_mcp_server_by_id(sid), "auth_type", None) == MCPAuth.oauth2 + if _is_v1_resolved_oauth2_server(global_mcp_server_manager.get_mcp_server_by_id(sid)) } async def _get_user_oauth_extra_headers( @@ -253,11 +270,13 @@ if MCP_AVAILABLE: the MCP server the same way the admin "Add MCP / Authorize and Fetch" flow does. Returns None for non-OAuth2 servers or when no credential is stored. + A server the v2 resolver owns is skipped; see ``_is_v1_resolved_oauth2_server``. + Args: prefetched_creds: Optional dict keyed by server_id with credential payloads. When provided, avoids a per-server DB round-trip. """ - if getattr(server, "auth_type", None) != MCPAuth.oauth2: + if not _is_v1_resolved_oauth2_server(server): return None user_id = getattr(user_api_key_dict, "user_id", None) server_id = getattr(server, "server_id", None) @@ -320,38 +339,6 @@ if MCP_AVAILABLE: verbose_logger.warning(f"_prefetch_user_oauth_creds: failed to prefetch for user={user_id}: {e}") return {} - async def _get_bulk_user_oauth_headers( - user_api_key_dict: UserAPIKeyAuth, - ) -> Dict[str, Dict[str, str]]: - """ - Fetch ALL OAuth2 credentials for the current user in a single DB query and - return a mapping of server_id → {"Authorization": "Bearer "}. - - This is the batch alternative to calling _get_user_oauth_extra_headers - per-server inside a loop (N+1 DB queries). - """ - user_id = getattr(user_api_key_dict, "user_id", None) - if not user_id: - return {} - try: - from litellm.proxy._experimental.mcp_server.db import ( - list_user_oauth_credentials, - ) - from litellm.proxy.utils import get_prisma_client_or_throw - - prisma_client = get_prisma_client_or_throw( - "Database not connected. Connect a database to use OAuth2 MCP tools." - ) - creds = await list_user_oauth_credentials(prisma_client, user_id) - return { - c["server_id"]: {"Authorization": f"Bearer {c['access_token']}"} - for c in creds - if c.get("access_token") and c.get("server_id") - } - except Exception: - verbose_logger.debug("Failed to bulk-fetch OAuth credentials", exc_info=True) - return {} - def _create_tool_response_objects(tools, server: MCPServer): """Helper function to create tool response objects. @@ -825,7 +812,7 @@ if MCP_AVAILABLE: # to avoid an unnecessary DB round-trip on requests with no OAuth2 MCP servers. prefetched_oauth_creds = ( await _prefetch_user_oauth_creds(user_api_key_dict) - if _get_oauth2_server_ids(allowed_server_ids) + if _v1_resolved_oauth2_server_ids(allowed_server_ids) else {} ) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 396dd6c7dc7..4fca4406a6f 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -27,7 +27,6 @@ from typing import ( Union, cast, ) -from urllib.parse import urlsplit, urlunsplit import httpx from fastapi import FastAPI, HTTPException @@ -59,6 +58,9 @@ from litellm.proxy._experimental.mcp_server.mcp_context import ( _mcp_gateway_server_name, ) from litellm.proxy._experimental.mcp_server.mcp_debug import MCPDebug +from litellm.proxy._experimental.mcp_server.oauth_utils import ( + _redact_mcp_resource_url, +) from litellm.proxy._experimental.mcp_server.utils import ( LITELLM_MCP_SERVER_DESCRIPTION, LITELLM_MCP_SERVER_NAME, @@ -106,27 +108,6 @@ _MAX_STATEFUL_SESSIONS_PER_OWNER = 100 _MCP_ROUTING_PEEK_MAX_BYTES = 4096 -def _redact_mcp_resource_url(url: Optional[str]) -> Optional[str]: - """Reduce an MCP server URL to its origin (scheme + host + port) for logging. - - Everything else is dropped: userinfo (``user:pass@``), the query string, the - fragment, and the path, because hosted MCP servers routinely embed the - credential in the path (e.g. ``/mcp/s/``) and this value is persisted - in spend-log metadata that a caller who can invoke the tool can read back. - Returns None when the URL has no host to identify (nothing safe to log). - """ - if not isinstance(url, str) or not url: - return None - try: - parts = urlsplit(url) - except ValueError: - return None - if not parts.hostname: - return None - netloc = f"{parts.hostname}:{parts.port}" if parts.port else parts.hostname - return urlunsplit((parts.scheme, netloc, "", "", "")) or None - - def _invalidate_byok_cred_cache(user_id: str, server_id: str) -> None: """Remove a (user_id, server_id) entry from the BYOK credential cache. @@ -977,7 +958,17 @@ if MCP_AVAILABLE: data = await add_litellm_data_to_request( data=body_data, request=request, - user_api_key_dict=user_api_key_auth, + # Bill a team-derived call to the team that granted it. A keyless admitted + # subject carries no team_id, so spend skipped team updates entirely and + # charged the user's PRIMARY org — the granting team's budget never + # accumulated (so it could never begin to block) and, cross-org, the wrong + # organization was charged. This is the ACCOUNTING half; the enforcement + # half (an already-over-budget team stops granting) lives in the source gate. + # Authorization is unaffected: it ran before this, and the union is resolved + # from the untouched auth object passed to call_mcp_tool below. + user_api_key_dict=await MCPRequestHandler.billing_auth_for_tool_call( + user_api_key_auth, tool_name=name + ), proxy_config=proxy_config, ) else: diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 67d935e34e3..7e8d08e7cad 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -26762,6 +26762,113 @@ "title": "ToolPolicyUpdateResponse", "type": "object" }, + "ToolSpendDailyEntry": { + "description": "Spend attributed to one tool on one UTC day.", + "properties": { + "call_count": { + "default": 0, + "title": "Call Count", + "type": "integer" + }, + "date": { + "title": "Date", + "type": "string" + }, + "spend": { + "default": 0.0, + "title": "Spend", + "type": "number" + }, + "tool_name": { + "title": "Tool Name", + "type": "string" + } + }, + "required": [ + "date", + "tool_name" + ], + "title": "ToolSpendDailyEntry", + "type": "object" + }, + "ToolSpendEntry": { + "description": "Total spend attributed to one tool over the requested window.", + "properties": { + "call_count": { + "default": 0, + "title": "Call Count", + "type": "integer" + }, + "spend": { + "default": 0.0, + "description": "Attributed spend: a request that used several tools counts its full spend toward each of them", + "title": "Spend", + "type": "number" + }, + "tool_name": { + "title": "Tool Name", + "type": "string" + }, + "total_tokens": { + "default": 0, + "title": "Total Tokens", + "type": "integer" + } + }, + "required": [ + "tool_name" + ], + "title": "ToolSpendEntry", + "type": "object" + }, + "ToolSpendResponse": { + "properties": { + "by_tool": { + "items": { + "$ref": "#/components/schemas/ToolSpendEntry" + }, + "title": "By Tool", + "type": "array" + }, + "daily": { + "items": { + "$ref": "#/components/schemas/ToolSpendDailyEntry" + }, + "title": "Daily", + "type": "array" + }, + "end_date": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "End Date" + }, + "start_date": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Start Date" + }, + "total_spend": { + "default": 0.0, + "description": "Deduplicated spend of every request that called at least one tool in the window; less than the sum of per-tool attributed spend whenever multi-tool requests exist", + "title": "Total Spend", + "type": "number" + } + }, + "title": "ToolSpendResponse", + "type": "object" + }, "ToolUsageLogEntry": { "description": "One spend log row for a tool call (for UI \"recent logs\" table).", "properties": { @@ -26858,6 +26965,13 @@ }, "ValidationError": { "properties": { + "ctx": { + "title": "Context", + "type": "object" + }, + "input": { + "title": "Input" + }, "loc": { "items": { "anyOf": [ @@ -27301,6 +27415,81 @@ ] } }, + "/v1/tool/spend": { + "get": { + "description": "Spend attributed to each tool over a date range, for the Cost Optimization dashboard.\n\nJoins ``LiteLLM_SpendLogToolIndex`` (which tool names ran on which request) to\n``LiteLLM_SpendLogs`` (what the request cost). A request that used multiple tools\ncounts its full spend toward each of those tools, so per-tool numbers are\nattributions. ``total_spend`` is the deduplicated spend of every request that\ncalled at least one tool in the window, so it never double counts.", + "operationId": "get_tool_spend_v1_tool_spend_get", + "parameters": [ + { + "description": "YYYY-MM-DD (defaults to 30 days ago)", + "in": "query", + "name": "start_date", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "description": "YYYY-MM-DD (defaults to 30 days ago)", + "title": "Start Date" + } + }, + { + "description": "YYYY-MM-DD (defaults to today)", + "in": "query", + "name": "end_date", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "description": "YYYY-MM-DD (defaults to today)", + "title": "End Date" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ToolSpendResponse" + } + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Get Tool Spend", + "tags": [ + "tools" + ] + } + }, "/v1/tool/{tool_name}": { "get": { "description": "Get details for a single tool.", diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index b73841c4793..98efadc10a8 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2605,6 +2605,28 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob user_max_budget: Optional[float] = None request_route: Optional[str] = None is_session_token: bool = False + # Server-only marker set exclusively by the MCP gateway admission path + # (_reload_admitted_user) for a keyless user-subject admitted via a gateway DCR session + # bearer or bridge envelope. Not a DB column and never populated from caller-controlled key + # metadata or JWT claims, so it cannot be forged to gain the team-inherited MCP grant union + # or to escape the caller-Authorization egress scrub. exclude=True keeps it out of serialization. + mcp_admitted_user_subject: bool = Field(default=False, exclude=True) + # team_id -> that team's mcp_rpm_limit map, for a keyless admitted subject that reaches MCP + # servers through several teams at once and therefore has no single team_id for the limiter to + # key off. Server-only and stripped from validated input for the same reason as the marker + # above: a forged entry would let a caller pick which team's rpm bucket it is charged against. + mcp_source_team_rpm_limits: dict[str, dict[str, int]] | None = Field(default=None, exclude=True) + via_virtual_key: bool = Field( + default=False, + exclude=True, + description=( + "Server-only marker set exclusively by the DB virtual-key and master-key auth paths via " + "post-construction assignment. Stripped from validated input so custom auth handlers, JWT " + "claims, or key metadata cannot forge it. Gates overwrite_user_with_key_hash stamping: only " + "a credential the proxy itself validated as a key may be forwarded as the provider-facing " + "user id." + ), + ) budget_reservation: Optional[Dict[str, Any]] = Field(default=None, exclude=True) budget_throttle_pct: Optional[float] = Field(default=None, exclude=True) user: Optional[Any] = None # Expanded user object when expand=user is used @@ -2625,6 +2647,12 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob # If values is already an instance (not a dict), return it as-is if not isinstance(values, dict): return values + # mcp_admitted_user_subject is a server-only marker, set ONLY by the MCP gateway admission + # path via post-construction assignment. Strip it from any validated input (constructor + # kwargs, model_validate, a JWT/key claim splat) so it can never be forged from caller data. + values.pop("mcp_admitted_user_subject", None) + values.pop("mcp_source_team_rpm_limits", None) + values.pop("via_virtual_key", None) if values.get("api_key") is not None: values.update({"token": cls._safe_hash_litellm_api_key(values.get("api_key"))}) if isinstance(values.get("api_key"), str): diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 99a867a5d07..ce82ca74267 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -2661,6 +2661,15 @@ async def get_managed_vector_store_rows_by_uuids( return result +class OrganizationNotFoundError(Exception): + """The organization row is CONFIRMED absent, as opposed to a lookup that failed. + + Subclasses Exception so every existing except Exception caller keeps its current + behavior; it exists so a caller that wants to treat "no such org" as "no restriction" can do + that WITHOUT also swallowing an outage and silently dropping a real org ceiling. + """ + + @log_db_metrics async def get_org_object( org_id: str, @@ -2707,25 +2716,30 @@ async def get_org_object( query_kwargs["include"] = {"litellm_budget_table": True} response = await OrganizationRepository(prisma_client).table.find_unique(**query_kwargs) - - if response is None: - raise Exception - - _org_obj = LiteLLM_OrganizationTable(**response.model_dump()) - # Cache the result - await user_api_key_cache.async_set_cache( - key=cache_key, - value=_org_obj, - model_type=LiteLLM_OrganizationTable, - ttl=DEFAULT_IN_MEMORY_TTL, - ) - - return _org_obj except Exception: - raise Exception( + # An operational failure (DB down, timeout, cache fault) is NOT the same fact as a confirmed + # missing row, and relabelling it as "doesn't exist" made every caller unable to tell them + # apart — a caller that treats absence as "this org places no restriction" then drops a real + # org ceiling during an outage. Propagate the real error; callers that already catch + # Exception are unaffected. + raise + + if response is None: + raise OrganizationNotFoundError( f"Organization doesn't exist in db. Organization={org_id}. Create organization via `/organization/new` call." ) + _org_obj = LiteLLM_OrganizationTable(**response.model_dump()) + # Cache the result + await user_api_key_cache.async_set_cache( + key=cache_key, + value=_org_obj, + model_type=LiteLLM_OrganizationTable, + ttl=DEFAULT_IN_MEMORY_TTL, + ) + + return _org_obj + async def _get_resources_from_access_groups( access_group_ids: List[str], diff --git a/litellm/proxy/auth/login_utils.py b/litellm/proxy/auth/login_utils.py index 11f12e597b9..f35d94c986e 100644 --- a/litellm/proxy/auth/login_utils.py +++ b/litellm/proxy/auth/login_utils.py @@ -7,12 +7,15 @@ login endpoints (e.g., /login and /v2/login). import os import secrets +from datetime import datetime, timedelta, timezone from typing import Literal, Optional, cast +import jwt from fastapi import HTTPException import litellm from litellm.constants import LITELLM_PROXY_ADMIN_NAME, LITELLM_UI_SESSION_DURATION +from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.proxy._types import ( LiteLLM_UserTable, LitellmUserRoles, @@ -313,6 +316,29 @@ async def authenticate_user( ) +def _ui_session_exp_timestamp() -> int: + """The ``exp`` claim (unix seconds) for a UI session cookie, ``LITELLM_UI_SESSION_DURATION`` + from now. The virtual key sealed inside the cookie already expires after this same + duration; stamping the JWT itself gives the cookie the bounded lifetime the dashboard's + client-side expiry check and the server-side session-cookie readers both assume, instead + of a token that stays signature-valid until the master key rotates.""" + ttl_seconds = duration_in_seconds(LITELLM_UI_SESSION_DURATION) + return int((datetime.now(timezone.utc) + timedelta(seconds=ttl_seconds)).timestamp()) + + +def encode_ui_session_jwt(returned_ui_token_object: ReturnedUITokenObject, master_key: str) -> str: + """Encode a UI session cookie JWT with a bounded ``exp``. + + The single choke point every UI login path (SSO and username/password /login, /v2, + /v3) uses to mint the ``token`` cookie, so the cookie's lifetime is set in exactly one + place and cannot drift between paths. Without the ``exp`` the cookie is valid until the + master key rotates, and the session-cookie readers that require a bounded lifetime + (the MCP interactive sign-in) reject it. + """ + claims = {**cast(dict, returned_ui_token_object), "exp": _ui_session_exp_timestamp()} + return jwt.encode(claims, master_key, algorithm="HS256") + + def create_ui_token_object( login_result: LoginResult, general_settings: dict, diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 83a8a69511b..709df5e64df 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1497,6 +1497,13 @@ async def _user_api_key_auth_builder( check_cache_only=True, ).resolve(hashed_token=hash_token(api_key)) ) + # Key-cache entries are written only after the proxy validated a + # virtual key or the master key, but via_virtual_key is exclude=True + # so serialization drops it; restore it at this trusted boundary. + # The UI-login JWT fallback below constructs its token from a + # decrypted blob, not this cache, and stays unmarked. + if isinstance(valid_token, UserAPIKeyAuth): + valid_token.via_virtual_key = True except Exception: verbose_logger.debug("api key not found in cache.") valid_token = None @@ -1614,6 +1621,7 @@ async def _user_api_key_auth_builder( _user_api_key_obj = update_valid_token_with_end_user_params( valid_token=_user_api_key_obj, end_user_params=end_user_params ) + _user_api_key_obj.via_virtual_key = True return _user_api_key_obj @@ -2021,7 +2029,7 @@ async def _user_api_key_auth_builder( # No token was found when looking up in the DB raise Exception("Invalid proxy server token passed") if valid_token_dict is not None: - return await _return_user_api_key_auth_obj( + virtual_key_auth_obj = await _return_user_api_key_auth_obj( user_obj=user_obj, api_key=api_key, parent_otel_span=parent_otel_span, @@ -2029,6 +2037,8 @@ async def _user_api_key_auth_builder( route=route, start_time=start_time, ) + virtual_key_auth_obj.via_virtual_key = True + return virtual_key_auth_obj except Exception as e: return await UserAPIKeyAuthExceptionHandler._handle_authentication_error( e=e, @@ -2442,6 +2452,7 @@ async def _reserve_budget_after_common_checks( end_user_id=end_user_id, end_user_object=end_user_object, skip_user_budget_on_team_key=general_settings.get("skip_user_budget_on_team_key") is True, + fail_closed_budget_enforcement=general_settings.get("fail_closed_budget_enforcement") is True, ) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 22ea9fe176a..b2216488db2 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -1781,28 +1781,38 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): """ from litellm.proxy.auth.auth_utils import get_team_mcp_rpm_limit - if not mcp_server_name or not user_api_key_dict.team_id: + if not mcp_server_name: return - mcp_rpm_limit = get_team_mcp_rpm_limit(user_api_key_dict) - if not mcp_rpm_limit: - return + # Which teams' buckets does this call charge? A key is pinned to exactly one team. A keyless + # MCP-admitted subject reaches servers through SEVERAL teams at once and has no team_id, so + # without the second source below its calls charged no team bucket at all and it outran every + # team's mcp_rpm_limit. Every applicable team is charged rather than one being picked: the + # limiter enforces all descriptors, so each team's own ceiling binds on a call made through + # its grant, and there is no arbitrary attribution when several teams grant the same server. + team_limits: list[tuple[str | None, dict[str, int] | None]] = [] + if user_api_key_dict.team_id: + team_limits.append((user_api_key_dict.team_id, get_team_mcp_rpm_limit(user_api_key_dict))) + for source_team_id, source_limit in (user_api_key_dict.mcp_source_team_rpm_limits or {}).items(): + team_limits.append((source_team_id, source_limit)) - server_rpm_limit = mcp_rpm_limit.get(mcp_server_name) - if server_rpm_limit is None: - return - - descriptors.append( - RateLimitDescriptor( - key="mcp_per_team", - value=f"{user_api_key_dict.team_id}:{mcp_server_name}", - rate_limit={ - "requests_per_unit": server_rpm_limit, - "tokens_per_unit": None, - "window_size": self.window_size, - }, + for team_id, mcp_rpm_limit in team_limits: + if not team_id or not mcp_rpm_limit: + continue + server_rpm_limit = mcp_rpm_limit.get(mcp_server_name) + if server_rpm_limit is None: + continue + descriptors.append( + RateLimitDescriptor( + key="mcp_per_team", + value=f"{team_id}:{mcp_server_name}", + rate_limit={ + "requests_per_unit": server_rpm_limit, + "tokens_per_unit": None, + "window_size": self.window_size, + }, + ) ) - ) def _should_enforce_rate_limit( self, diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 9d9ef28ec9b..a4cc4a62009 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -13,7 +13,7 @@ from starlette.datastructures import Headers import litellm from litellm._logging import verbose_logger, verbose_proxy_logger from litellm._service_logger import ServiceLogging -from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY +from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS, PRE_CALL_EXECUTED_GUARDRAILS_KEY from litellm.litellm_core_utils.credential_accessor import CredentialAccessor from litellm.litellm_core_utils.initialize_dynamic_callback_params import ( iter_client_callback_metadata_dicts, @@ -48,6 +48,24 @@ _EXPLICIT_SESSION_HEADERS = frozenset({"x-litellm-trace-id", "x-litellm-session- # Session-id values must be non-empty strings of alphanumerics, hyphens, or underscores # (covers UUIDs and most common session-id formats). _SESSION_ID_VALUE_RE = re.compile(r"^[a-zA-Z0-9_\-]{8,}$") + +_SHA256_HEX_RE = re.compile(r"^[0-9a-f]{64}$") + + +def _stampable_key_hash(user_api_key_dict: UserAPIKeyAuth) -> str | None: + """Only proxy-validated keys are stamped, proven by the unforgeable + via_virtual_key marker AND a known non-secret shape: the sha256 hex digest + UserAPIKeyAuth stores virtual keys in, or the master key's stable alias. + Custom-auth credentials arrive raw (never forward auth material) and hashed + JWTs rotate on re-issue (useless as a stable ban id), so both are skipped.""" + api_key = user_api_key_dict.api_key + if not user_api_key_dict.via_virtual_key or api_key is None: + return None + if api_key == LITELLM_PROXY_MASTER_KEY_ALIAS or _SHA256_HEX_RE.fullmatch(api_key): + return api_key + return None + + _ANTHROPIC_SESSION_ID_VALUE_RE = re.compile(r"^[a-zA-Z0-9_\-]+$") @@ -1447,6 +1465,11 @@ async def add_litellm_data_to_request( if "user" not in data: data["user"] = user + if litellm.overwrite_user_with_key_hash is True: + stampable_hash = _stampable_key_hash(user_api_key_dict) + if stampable_hash is not None: + data["user"] = stampable_hash + data["secret_fields"] = SecretFields(raw_headers=_raw_headers) ## Dynamic api version (Azure OpenAI endpoints) ## diff --git a/litellm/proxy/management_endpoints/tool_management_endpoints.py b/litellm/proxy/management_endpoints/tool_management_endpoints.py index 9d71761f115..ca606e07cee 100644 --- a/litellm/proxy/management_endpoints/tool_management_endpoints.py +++ b/litellm/proxy/management_endpoints/tool_management_endpoints.py @@ -10,16 +10,18 @@ POST /v1/tool/policy - Update the input_policy / output_policy for a """ import uuid -from datetime import datetime, timezone -from typing import TYPE_CHECKING, Any, List, Optional +from datetime import datetime, timedelta, timezone +from itertools import groupby +from typing import TYPE_CHECKING, Annotated, Any, List, Optional from fastapi import APIRouter, Depends, HTTPException, Query +from pydantic import BaseModel, TypeAdapter if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient from litellm._logging import verbose_proxy_logger -from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth +from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.repositories.object_permission_repository import ObjectPermissionRepository from litellm.repositories.table_repositories import ( @@ -39,6 +41,9 @@ from litellm.types.tool_management import ( ToolPolicyOptionsResponse, ToolPolicyUpdateRequest, ToolPolicyUpdateResponse, + ToolSpendDailyEntry, + ToolSpendEntry, + ToolSpendResponse, ToolUsageLogEntry, ToolUsageLogsResponse, ) @@ -124,6 +129,147 @@ async def list_tools( raise HTTPException(status_code=500, detail=str(e)) +def _parse_day_start(value: str | None) -> datetime | None: + if not value: + return None + try: + return datetime.strptime(value.strip(), "%Y-%m-%d").replace(tzinfo=timezone.utc) + except ValueError: + raise HTTPException( + status_code=400, + detail=f"Invalid date format: {value}. Expected: 'YYYY-MM-DD'", + ) + + +class _ToolSpendRow(BaseModel): + date: str + tool_name: str + call_count: int + spend: float + total_tokens: int + + +class _RequestTotalRow(BaseModel): + total_spend: float + + +_TOOL_SPEND_ROWS = TypeAdapter(list[_ToolSpendRow]) +_REQUEST_TOTAL_ROWS = TypeAdapter(list[_RequestTotalRow]) + + +def _summarize_tool(name: str, grp: tuple[_ToolSpendRow, ...]) -> ToolSpendEntry: + return ToolSpendEntry( + tool_name=name, + spend=sum(r.spend for r in grp), + call_count=sum(r.call_count for r in grp), + total_tokens=sum(r.total_tokens for r in grp), + ) + + +def _build_tool_spend_response( + rows: list[_ToolSpendRow], + total_spend: float, + start_date: str, + end_date: str, +) -> ToolSpendResponse: + daily = [ + ToolSpendDailyEntry(date=r.date, tool_name=r.tool_name, spend=r.spend, call_count=r.call_count) for r in rows + ] + grouped = groupby(sorted(rows, key=lambda r: r.tool_name), key=lambda r: r.tool_name) + by_tool = sorted( + (_summarize_tool(name, tuple(grp)) for name, grp in grouped), + key=lambda e: e.spend, + reverse=True, + ) + return ToolSpendResponse( + by_tool=by_tool, + daily=daily, + total_spend=total_spend, + start_date=start_date, + end_date=end_date, + ) + + +@router.get( + "/v1/tool/spend", + tags=["tool management"], + dependencies=[Depends(user_api_key_auth)], + response_model=ToolSpendResponse, +) +async def get_tool_spend( + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + start_date: Annotated[str | None, Query(description="YYYY-MM-DD (defaults to 30 days ago)")] = None, + end_date: Annotated[str | None, Query(description="YYYY-MM-DD (defaults to today)")] = None, +): + """ + Spend attributed to each tool over a date range, for the Cost Optimization dashboard. + + Joins ``LiteLLM_SpendLogToolIndex`` (which tool names ran on which request) to + ``LiteLLM_SpendLogs`` (what the request cost). A request that used multiple tools + counts its full spend toward each of those tools, so per-tool numbers are + attributions. ``total_spend`` is the deduplicated spend of every request that + called at least one tool in the window, so it never double counts. + """ + from litellm.proxy.proxy_server import prisma_client + + if user_api_key_dict.user_role not in ( + LitellmUserRoles.PROXY_ADMIN, + LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, + ): + raise HTTPException( + status_code=403, + detail="Only proxy admin roles can view tool spend across the deployment", + ) + + if prisma_client is None: + raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) + + now = datetime.now(timezone.utc) + end_day = _parse_day_start(end_date) + start_dt = _parse_day_start(start_date) or ((end_day or now) - timedelta(days=30)) + end_exclusive = (end_day + timedelta(days=1)) if end_day else now + + rows = await prisma_client.db.query_raw( + """ + SELECT to_char(ti.start_time, 'YYYY-MM-DD') AS date, + ti.tool_name AS tool_name, + COUNT(*)::int AS call_count, + COALESCE(SUM(sl.spend), 0)::double precision AS spend, + COALESCE(SUM(sl.total_tokens), 0)::bigint AS total_tokens + FROM "LiteLLM_SpendLogToolIndex" ti + JOIN "LiteLLM_SpendLogs" sl ON sl.request_id = ti.request_id + WHERE ti.start_time >= ($1::timestamptz AT TIME ZONE 'UTC') + AND ti.start_time < ($2::timestamptz AT TIME ZONE 'UTC') + GROUP BY date, ti.tool_name + ORDER BY date ASC, spend DESC + """, + start_dt.isoformat(), + end_exclusive.isoformat(), + ) + totals = await prisma_client.db.query_raw( + """ + SELECT COALESCE(SUM(sl.spend), 0)::double precision AS total_spend + FROM "LiteLLM_SpendLogs" sl + WHERE EXISTS ( + SELECT 1 + FROM "LiteLLM_SpendLogToolIndex" ti + WHERE ti.request_id = sl.request_id + AND ti.start_time >= ($1::timestamptz AT TIME ZONE 'UTC') + AND ti.start_time < ($2::timestamptz AT TIME ZONE 'UTC') + ) + """, + start_dt.isoformat(), + end_exclusive.isoformat(), + ) + total_rows = _REQUEST_TOTAL_ROWS.validate_python(totals or []) + return _build_tool_spend_response( + rows=_TOOL_SPEND_ROWS.validate_python(rows or []), + total_spend=total_rows[0].total_spend if total_rows else 0.0, + start_date=start_dt.strftime("%Y-%m-%d"), + end_date=(end_day or now).strftime("%Y-%m-%d"), + ) + + @router.get( "/v1/tool/{tool_name:path}/detail", tags=["tool management"], diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index de988a0140f..31b98bf20e4 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -36,7 +36,7 @@ if TYPE_CHECKING: import httpx import jwt -from fastapi import APIRouter, Depends, Header, HTTPException, Request, status +from fastapi import APIRouter, Depends, Header, HTTPException, Request, Response, status from fastapi.responses import RedirectResponse import litellm @@ -965,15 +965,8 @@ async def google_login( state=cli_state, request=request, ) - if return_to is not None and sso_redirect is not None: - if SSOAuthenticationHandler._validate_return_to(return_to): - sso_redirect.set_cookie( - key="litellm_cp_return_to", - value=return_to, - max_age=600, - httponly=True, - samesite="lax", - ) + if sso_redirect is not None: + _persist_return_to_cookie(sso_redirect, return_to) return sso_redirect from fastapi.responses import HTMLResponse @@ -982,13 +975,19 @@ async def google_login( os.getenv("LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT", "false").lower() == "true" or general_settings.get("hide_default_credentials_hint", False) is True ) - return HTMLResponse( + form_response = HTMLResponse( content=build_ui_login_form( show_deprecation_banner=True, hide_default_credentials_hint=hide_default_credentials_hint, ), status_code=200, ) + # Preserve return_to across the username/password sign-in too, via the SAME shared, never-raising + # helper the SSO branch uses, so /login can resume the connect flow instead of dead-ending at the + # dashboard. One implementation → the two sign-in branches cannot diverge (and the login form always + # renders, since the helper never raises on a bad return_to). + _persist_return_to_cookie(form_response, return_to) + return form_response def generic_response_convertor( @@ -2418,6 +2417,92 @@ async def sso_readiness(): ) +def _is_same_origin_return_path(return_to: str) -> bool: + """True for a strictly relative return path that stays on the gateway's own origin by + construction, and is therefore safe to honor without a configured ``control_plane_url``. + Used by the MCP gateway DCR authorize round-trip so a browser sent through login lands + back on the authorize request. + + Requires a single leading ``/`` (not protocol-relative ``//``), no backslash (browsers + fold ``\\`` to ``/``, so ``/\\evil.com`` would escape the origin), and no control or + whitespace characters. Rejecting control chars keeps a ``\\r\\n``/tab-bearing value out + of the redirect ``Location`` and the ``litellm_cp_return_to`` cookie entirely, rather + than relying on downstream header encoding to neutralize it.""" + if not return_to.startswith("/") or return_to.startswith("//") or "\\" in return_to: + return False + return not any(ord(ch) < 0x20 or ch in (" ", "\x7f") for ch in return_to) + + +async def _sso_return_to_redirect( + return_to: str | None, + jwt_token: str, + redis_usage_cache, + user_api_key_cache, +) -> RedirectResponse | None: + """Resolve the post-SSO redirect for a ``return_to``, or None to fall through to the dashboard. + + Two arms, both clearing the one-shot ``litellm_cp_return_to`` cookie: + - **Same-origin relative path** (the MCP gateway DCR authorize round-trip): set the session cookie + exactly like the dashboard path, then send the browser back where it came from. + - **Control-plane cross-origin** (``control_plane_url``): stash the JWT behind a single-use opaque + code (60s TTL) so the token never lands in browser history/logs; the control plane redeems it via + ``POST /v3/login/exchange``. + + Extracted from ``get_redirect_response_from_openid`` to keep that method inside the complexity + budget; behavior is identical to the inline arms it replaces (including letting + ``_validate_return_to`` raise for a mismatched absolute return_to, as before).""" + if return_to is None: + return None + + if _is_same_origin_return_path(return_to): + redirect_response = RedirectResponse(url=return_to, status_code=303) + redirect_response.set_cookie(key="token", value=jwt_token) + redirect_response.delete_cookie("litellm_cp_return_to") + return redirect_response + + if SSOAuthenticationHandler._validate_return_to(return_to): + code = secrets.token_urlsafe(32) + cache_key = f"login_code:{code}" + cache_value = {"token": jwt_token, "redirect_url": return_to} + if redis_usage_cache is not None: + await redis_usage_cache.async_set_cache(key=cache_key, value=cache_value, ttl=60) + else: + await user_api_key_cache.async_set_cache(key=cache_key, value=cache_value, ttl=60) + + separator = "&" if "?" in return_to else "?" + redirect_url = return_to + separator + urlencode({"login": "success", "code": code}) + verbose_proxy_logger.info("Cross-origin SSO: redirecting to control plane with login code") + redirect_response = RedirectResponse(url=redirect_url, status_code=303) + redirect_response.delete_cookie("litellm_cp_return_to") + return redirect_response + + return None + + +def _persist_return_to_cookie(response: Response, return_to: str | None) -> None: + """Best-effort: persist a SAFE ``return_to`` on ``response`` as the one-shot ``litellm_cp_return_to`` + cookie so ANY sign-in path — SSO / Okta / generic OR the username/password form — can resume there + afterwards. THIS is the single source of truth, called by every sign-in branch so they cannot + diverge (a per-branch reimplementation is exactly how the two drifted before). Honors a strictly + relative same-origin path, and (when ``control_plane_url`` is configured) a return_to matching that + origin. It NEVER raises: a mismatched or invalid ``return_to`` is simply not stored, so it can never + block sign-in — the login entrypoint must always render.""" + if return_to is None: + return + try: + safe = _is_same_origin_return_path(return_to) or SSOAuthenticationHandler._validate_return_to(return_to) + except HTTPException: + return # a non-matching absolute return_to is ignored, never blocks sign-in + if safe: + response.set_cookie( + key="litellm_cp_return_to", + value=return_to, + max_age=600, + httponly=True, + samesite="lax", + ) + + class SSOAuthenticationHandler: """ Handler for SSO Authentication across all SSO providers @@ -3055,7 +3140,6 @@ class SSOAuthenticationHandler: return_to: Optional[str] = None, sso_assertion: SSOIdentityAssertion | None = None, ) -> RedirectResponse: - import jwt from litellm.proxy.proxy_server import ( general_settings, @@ -3219,30 +3303,21 @@ class SSOAuthenticationHandler: server_root_path=get_server_root_path(), ) - jwt_token = jwt.encode( - cast(dict, returned_ui_token_object), - master_key or "", - algorithm="HS256", + from litellm.proxy.auth.login_utils import encode_ui_session_jwt + + jwt_token = encode_ui_session_jwt(returned_ui_token_object, master_key or "") + + # Post-SSO return_to handling (the same-origin DCR round-trip and the control-plane + # cross-origin code exchange) lives in one shared helper so this method stays inside the + # complexity budget. None falls through to the dashboard redirect below. + return_to_redirect = await _sso_return_to_redirect( + return_to=return_to, + jwt_token=jwt_token, + redis_usage_cache=redis_usage_cache, + user_api_key_cache=user_api_key_cache, ) - - # Control-plane cross-origin: store JWT behind a single-use opaque - # code (60s TTL) so the token never appears in browser history / logs. - # The control plane redeems it via POST /v3/login/exchange. - if return_to is not None and SSOAuthenticationHandler._validate_return_to(return_to): - code = secrets.token_urlsafe(32) - cache_key = f"login_code:{code}" - cache_value = {"token": jwt_token, "redirect_url": return_to} - if redis_usage_cache is not None: - await redis_usage_cache.async_set_cache(key=cache_key, value=cache_value, ttl=60) - else: - await user_api_key_cache.async_set_cache(key=cache_key, value=cache_value, ttl=60) - - separator = "&" if "?" in return_to else "?" - redirect_url = return_to + separator + urlencode({"login": "success", "code": code}) - verbose_proxy_logger.info("Cross-origin SSO: redirecting to control plane with login code") - redirect_response = RedirectResponse(url=redirect_url, status_code=303) - redirect_response.delete_cookie("litellm_cp_return_to") - return redirect_response + if return_to_redirect is not None: + return return_to_redirect if user_id is not None and isinstance(user_id, str): litellm_dashboard_ui += "?login=success" diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 32845763f22..a20b557e38b 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -13472,7 +13472,7 @@ async def fallback_login(request: Request): @router.post("/login", include_in_schema=False) # hidden since this is a helper for UI sso login async def login(request: Request): global premium_user, general_settings, master_key - from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object + from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object, encode_ui_session_jwt from litellm.proxy.utils import get_custom_url form = await request.form() @@ -13495,13 +13495,7 @@ async def login(request: Request): ) # Generate JWT token - import jwt - - jwt_token = jwt.encode( - cast(dict, returned_ui_token_object), - cast(str, master_key), - algorithm="HS256", - ) + jwt_token = encode_ui_session_jwt(returned_ui_token_object, cast(str, master_key)) # Build redirect URL litellm_dashboard_ui = get_custom_url(str(request.base_url)) @@ -13511,16 +13505,51 @@ async def login(request: Request): litellm_dashboard_ui += "/ui/" litellm_dashboard_ui += "?login=success" + # Honor a same-origin return_to preserved by the sign-in page (e.g. the aggregate DCR connect flow's + # authorize round-trip), mirroring the SSO callback; otherwise land on the dashboard. Gated by + # _is_same_origin_return_path (strictly relative path) so it can never be an open redirect, and the + # one-shot cookie is cleared after use. + from litellm.proxy.management_endpoints.ui_sso import _sso_return_to_redirect + + # Resume through the SAME resumer the SSO callback uses, rather than a second, narrower arm. + # _persist_return_to_cookie stores both shapes it accepts (a relative same-origin path AND a + # control_plane_url-matching absolute URL); honoring only the relative one here silently dropped + # the control-plane case, landing the user on the dashboard. One function decides how a stored + # return_to is honored for EVERY sign-in branch, so the write and read sets cannot diverge: it + # sets the token cookie on the same-origin arm and hands off via a one-time login code on the + # cross-origin arm, and clears the one-shot cookie in both. + cp_return_to = request.cookies.get("litellm_cp_return_to") + if cp_return_to: + try: + resumed = await _sso_return_to_redirect( + return_to=cp_return_to, + jwt_token=jwt_token, + redis_usage_cache=redis_usage_cache, + user_api_key_cache=user_api_key_cache, + ) + except Exception: # noqa: BLE001 # resuming must NEVER block a completed sign-in + # The symmetric half of _persist_return_to_cookie's "never raises" contract. The resumer + # rejects a return_to that no longer matches control_plane_url (a config change between + # the cookie's write and this read), and the user has ALREADY authenticated here — + # failing their login over a stale one-shot cookie is the worst possible outcome. Land + # on the dashboard instead; the cookie is cleared below either way. + verbose_proxy_logger.info("Ignoring stale litellm_cp_return_to cookie; landing on dashboard") + resumed = None + if resumed is not None: + return resumed + # Create redirect response with cookie redirect_response = RedirectResponse(url=litellm_dashboard_ui, status_code=303) redirect_response.set_cookie(key="token", value=jwt_token) + if cp_return_to: + redirect_response.delete_cookie(key="litellm_cp_return_to") return redirect_response @router.post("/v2/login", include_in_schema=False) # hidden helper for UI logins via API async def login_v2(request: Request): global premium_user, general_settings, master_key - from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object + from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object, encode_ui_session_jwt from litellm.proxy.utils import get_custom_url try: @@ -13541,13 +13570,7 @@ async def login_v2(request: Request): premium_user=premium_user, ) - import jwt - - jwt_token = jwt.encode( - cast(dict, returned_ui_token_object), - cast(str, master_key), - algorithm="HS256", - ) + jwt_token = encode_ui_session_jwt(returned_ui_token_object, cast(str, master_key)) litellm_dashboard_ui = get_custom_url(str(request.base_url)) if litellm_dashboard_ui.endswith("/"): @@ -13591,7 +13614,7 @@ async def login_v2(request: Request): ) # control-plane login — always returns token in body for cross-origin use async def login_v3(request: Request): global premium_user, general_settings, master_key - from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object + from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object, encode_ui_session_jwt from litellm.proxy.utils import get_custom_url try: @@ -13620,13 +13643,7 @@ async def login_v3(request: Request): premium_user=premium_user, ) - import jwt - - jwt_token = jwt.encode( - cast(dict, returned_ui_token_object), - cast(str, master_key), - algorithm="HS256", - ) + jwt_token = encode_ui_session_jwt(returned_ui_token_object, cast(str, master_key)) litellm_dashboard_ui = get_custom_url(str(request.base_url)) if litellm_dashboard_ui.endswith("/"): diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 80fd8a1594e..013873179c6 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -4,7 +4,9 @@ import asyncio import json from dataclasses import dataclass from datetime import datetime, timedelta, timezone -from typing import Any, Dict, List, Mapping, Optional, Sequence, cast +from typing import Any, Dict, List, Mapping, NoReturn, Optional, Sequence, cast + +from fastapi import HTTPException, status import litellm from litellm._logging import verbose_proxy_logger @@ -59,6 +61,22 @@ class _CounterReservationUnavailable(Exception): super().__init__("Counter reservation unavailable") +def _raise_reservation_unavailable(counter_key: str) -> NoReturn: + verbose_proxy_logger.warning( + "fail_closed_budget_enforcement: rejecting request — budget reservation for %s could not be written", + counter_key, + ) + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail=( + "Budget enforcement unavailable: the budget reservation could not " + "be written to the spend counter backend, and " + "fail_closed_budget_enforcement is enabled, so the request was " + "rejected to avoid exceeding the configured budget. Retry shortly." + ), + ) + + def get_reserved_counter_keys(budget_reservation: Optional[dict]) -> set: if not budget_reservation: return set() @@ -138,6 +156,7 @@ async def reserve_budget_for_request( end_user_id: Optional[str] = None, end_user_object: Optional[Any] = None, skip_user_budget_on_team_key: bool = False, + fail_closed_budget_enforcement: bool = False, ) -> Optional[dict]: if valid_token is None or not RouteChecks.is_llm_api_route(route=route): return None @@ -193,6 +212,8 @@ async def reserve_budget_for_request( default_reserved_cost=reservation_cost, ) applied_entries.remove(entry) + if fail_closed_budget_enforcement: + _raise_reservation_unavailable(counter_key=counter.counter_key) continue if reserved_value is not None: diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index 55b50e7d9ff..0c525ee9466 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -1695,13 +1695,10 @@ async def ui_view_spend_logs( code=status.HTTP_401_UNAUTHORIZED, ) - if start_date is None or end_date is None: - raise ProxyException( - message="Start date and end date are required", - type="bad_request", - param="None", - code=status.HTTP_400_BAD_REQUEST, - ) + # Inline import — auth_utils participates in a proxy import cycle. + from litellm.proxy.auth.auth_utils import get_request_route # noqa: PLC0415 + + is_v2 = "/spend/logs/v2" in get_request_route(request) # Validate sort_by and sort_order valid_sort_fields = { @@ -1729,36 +1726,50 @@ async def ui_view_spend_logs( ) try: - # Inline import — auth_utils participates in a proxy import cycle. - from litellm.proxy.auth.auth_utils import get_request_route # noqa: PLC0415 + is_admin_view = _is_admin_view_safe(user_api_key_dict=user_api_key_dict) + is_request_id_lookup = request_id is not None and not is_v2 - is_v2 = "/spend/logs/v2" in get_request_route(request) - formats = ["%Y-%m-%d %H:%M:%S", "%Y-%m-%d"] if is_v2 else ["%Y-%m-%d %H:%M:%S"] + if is_request_id_lookup: + # request_id is the @id primary key: it identifies a single row, so a + # time window is meaningless. The dashboard always sends a default 24h + # window, which hid ids copied from an older page (LIT-3981). Drop the + # window for the id lookup so it resolves across all time; every other + # query, including the public v2 route, still requires one (below). + start_date_obj: datetime | None = None + end_date_obj: datetime | None = None + else: + if start_date is None or end_date is None: + raise ProxyException( + message="Start date and end date are required", + type="bad_request", + param="None", + code=status.HTTP_400_BAD_REQUEST, + ) + formats = ["%Y-%m-%d %H:%M:%S", "%Y-%m-%d"] if is_v2 else ["%Y-%m-%d %H:%M:%S"] - def parse_date(date_str: str) -> datetime: - date_str = date_str.strip() - for fmt in formats: - try: - return datetime.strptime(date_str, fmt).replace(tzinfo=timezone.utc) - except ValueError: - continue - expected = "'YYYY-MM-DD' or 'YYYY-MM-DD HH:MM:SS'" if is_v2 else "'YYYY-MM-DD HH:MM:SS'" - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail=f"Invalid date format: {date_str}. Expected: {expected}", - ) + def parse_date(date_str: str) -> datetime: + date_str = date_str.strip() + for fmt in formats: + try: + return datetime.strptime(date_str, fmt).replace(tzinfo=timezone.utc) + except ValueError: + continue + expected = "'YYYY-MM-DD' or 'YYYY-MM-DD HH:MM:SS'" if is_v2 else "'YYYY-MM-DD HH:MM:SS'" + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=f"Invalid date format: {date_str}. Expected: {expected}", + ) - start_date_obj = parse_date(start_date) - end_date_obj = parse_date(end_date) - - # Convert to ISO format strings for Prisma - start_date_iso = start_date_obj.isoformat() # Already in UTC, no need to add Z - end_date_iso = end_date_obj.isoformat() # Already in UTC, no need to add Z + start_date_obj = parse_date(start_date) + end_date_obj = parse_date(end_date) # Build where conditions - where_conditions: dict[str, Any] = { - "startTime": {"gte": start_date_iso, "lte": end_date_iso}, - } + where_conditions: dict[str, Any] = {} + if start_date_obj is not None and end_date_obj is not None: + where_conditions["startTime"] = { + "gte": start_date_obj.isoformat(), # Already in UTC, no need to add Z + "lte": end_date_obj.isoformat(), + } if team_id is not None: where_conditions["team_id"] = team_id @@ -1827,9 +1838,19 @@ async def ui_view_spend_logs( where_conditions["spend"]["gte"] = min_spend if max_spend is not None: where_conditions["spend"]["lte"] = max_spend - is_admin_view = _is_admin_view_safe(user_api_key_dict=user_api_key_dict) + # A request_id lookup drops the date window, so a non-admin could otherwise + # reach any single row by id; require they own it, mirroring the detail + # endpoint. That ownership check fully authorizes the one row, so the + # general scoping below is skipped for id lookups. Scoped to the UI route + # so the public v2 contract is unchanged. + if request_id is not None and not is_v2 and not is_admin_view: + await _assert_user_can_view_request_id( + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + request_id=request_id, + ) permitted_team_ids: List[str] | None = None - if not is_admin_view: + if not is_request_id_lookup and not is_admin_view: if team_id is not None: can_view_team = await _can_team_member_view_log( prisma_client=prisma_client, @@ -1875,15 +1896,16 @@ async def ui_view_spend_logs( sql_params: List[Any] = [] p = 1 # parameter index counter - # Date range (always present). Wrap the param side with - # `AT TIME ZONE 'UTC'` so comparison against the plain `timestamp` - # column does not depend on the DB session timezone (see #22529). - sql_conditions.append(f"\"startTime\" >= (${p}::timestamptz AT TIME ZONE 'UTC')") - sql_params.append(start_date_obj) - p += 1 - sql_conditions.append(f"\"startTime\" <= (${p}::timestamptz AT TIME ZONE 'UTC')") - sql_params.append(end_date_obj) - p += 1 + # Date range. Wrap the param side with `AT TIME ZONE 'UTC'` so comparison + # against the plain `timestamp` column does not depend on the DB session + # timezone (see #22529). Absent for a request_id-only lookup (see above). + if start_date_obj is not None and end_date_obj is not None: + sql_conditions.append(f"\"startTime\" >= (${p}::timestamptz AT TIME ZONE 'UTC')") + sql_params.append(start_date_obj) + p += 1 + sql_conditions.append(f"\"startTime\" <= (${p}::timestamptz AT TIME ZONE 'UTC')") + sql_params.append(end_date_obj) + p += 1 # Equality filters - read effective values from where_conditions (post-authorization) for sql_col, wc_key in [ diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index 50f6f791bc2..a6105b6dff9 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -374,12 +374,22 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs if isinstance(v, BaseModel): v = v.model_dump() additional_usage_values.update({k: v}) - if "cache_read_input_tokens" not in additional_usage_values: - prompt_tokens_details = additional_usage_values.get("prompt_tokens_details") - if isinstance(prompt_tokens_details, dict): + prompt_tokens_details = additional_usage_values.get("prompt_tokens_details") + if not isinstance(prompt_tokens_details, dict): + usage_object = clean_metadata.get("usage_object") + if isinstance(usage_object, dict): + prompt_tokens_details = usage_object.get("prompt_tokens_details") + if isinstance(prompt_tokens_details, dict): + if "cache_read_input_tokens" not in additional_usage_values: cached_tokens = prompt_tokens_details.get("cached_tokens") if isinstance(cached_tokens, int) and cached_tokens > 0: additional_usage_values["cache_read_input_tokens"] = cached_tokens + if "cache_creation_input_tokens" not in additional_usage_values: + cache_write_tokens = prompt_tokens_details.get("cache_write_tokens") or prompt_tokens_details.get( + "cache_creation_tokens" + ) + if isinstance(cache_write_tokens, int) and cache_write_tokens > 0: + additional_usage_values["cache_creation_input_tokens"] = cache_write_tokens clean_metadata["additional_usage_values"] = additional_usage_values if litellm.cache is not None: diff --git a/litellm/responses/utils.py b/litellm/responses/utils.py index 67a0ed75fdc..12c890ec91d 100644 --- a/litellm/responses/utils.py +++ b/litellm/responses/utils.py @@ -1042,21 +1042,14 @@ class ResponseAPILoggingUtils: prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None if response_api_usage.input_tokens_details: if isinstance(response_api_usage.input_tokens_details, dict): - input_tokens_details = dict(response_api_usage.input_tokens_details) - cache_write_tokens = input_tokens_details.pop("cache_write_tokens", None) - if input_tokens_details.get("cache_creation_tokens") is None and cache_write_tokens is not None: - input_tokens_details["cache_creation_tokens"] = cache_write_tokens - prompt_tokens_details = PromptTokensDetailsWrapper(**input_tokens_details) + prompt_tokens_details = PromptTokensDetailsWrapper(**response_api_usage.input_tokens_details) else: prompt_tokens_details = PromptTokensDetailsWrapper( cached_tokens=getattr(response_api_usage.input_tokens_details, "cached_tokens", None), audio_tokens=getattr(response_api_usage.input_tokens_details, "audio_tokens", None), text_tokens=getattr(response_api_usage.input_tokens_details, "text_tokens", None), image_tokens=getattr(response_api_usage.input_tokens_details, "image_tokens", None), - cache_creation_tokens=getattr( - response_api_usage.input_tokens_details, "cache_creation_tokens", None - ) - or getattr(response_api_usage.input_tokens_details, "cache_write_tokens", None), + cache_write_tokens=getattr(response_api_usage.input_tokens_details, "cache_write_tokens", None), ) completion_tokens_details: Optional[CompletionTokensDetailsWrapper] = None output_tokens_details = getattr(response_api_usage, "output_tokens_details", None) diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 8ae974b19a6..b0af22e7c3f 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -261,12 +261,3 @@ class MCPServer(BaseModel): if self.oauth_passthrough is not True: return False return any(h.lower() == "authorization" for h in self.extra_headers) - - @property - def has_token_exchange_config(self) -> bool: - """True if this server is configured for OAuth2 token exchange (OBO / RFC 8693).""" - return ( - self.auth_type == MCPAuth.oauth2_token_exchange - and bool(self.client_id and self.client_secret) - and bool(self.token_exchange_endpoint or self.token_url) - ) diff --git a/litellm/types/tool_management.py b/litellm/types/tool_management.py index 1c5e1df9e9a..71ec412e8ef 100644 --- a/litellm/types/tool_management.py +++ b/litellm/types/tool_management.py @@ -98,3 +98,38 @@ class ToolUsageLogsResponse(BaseModel): total: int page: int page_size: int + + +class ToolSpendEntry(BaseModel): + """Total spend attributed to one tool over the requested window.""" + + tool_name: str + spend: float = Field( + 0.0, + description="Attributed spend: a request that used several tools counts its full spend toward each of them", + ) + call_count: int = 0 + total_tokens: int = 0 + + +class ToolSpendDailyEntry(BaseModel): + """Spend attributed to one tool on one UTC day.""" + + date: str + tool_name: str + spend: float = 0.0 + call_count: int = 0 + + +class ToolSpendResponse(BaseModel): + by_tool: List[ToolSpendEntry] = Field(default_factory=list) + daily: List[ToolSpendDailyEntry] = Field(default_factory=list) + total_spend: float = Field( + 0.0, + description=( + "Deduplicated spend of every request that called at least one tool in the window; " + "less than the sum of per-tool attributed spend whenever multi-tool requests exist" + ), + ) + start_date: str | None = None + end_date: str | None = None diff --git a/litellm/types/utils.py b/litellm/types/utils.py index c6c835ae0c2..c84f22a4b76 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -1537,14 +1537,27 @@ class PromptTokensDetailsWrapper( audio_length_seconds: Optional[float] = None """Length of audio sent to the model. Used for multimodal embeddings priced per audio-second.""" + cache_write_tokens: Optional[int] = None + """Number of cache write (creation) tokens sent to the model. OpenAI naming (prompt_tokens_details.cache_write_tokens); this is the canonical field.""" + cache_creation_tokens: Optional[int] = None - """Number of cache creation tokens sent to the model. Used for Anthropic prompt caching.""" + """Number of cache creation tokens sent to the model. Anthropic/Bedrock naming; kept in sync with cache_write_tokens (assigning either mirrors to the other).""" cache_creation_token_details: Optional[CacheCreationTokenDetails] = None """Details of cache creation tokens sent to the model. Used for tracking 5m/1h cache creation tokens for Anthropic prompt caching.""" + def __setattr__(self, name: str, value: object) -> None: + super().__setattr__(name, value) + if name == "cache_write_tokens": + super().__setattr__("cache_creation_tokens", value) + elif name == "cache_creation_tokens": + super().__setattr__("cache_write_tokens", value) + def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) + self.cache_write_tokens = ( + self.cache_write_tokens if self.cache_write_tokens is not None else self.cache_creation_tokens + ) if self.character_count is None: del self.character_count if self.image_count is None: @@ -1557,6 +1570,8 @@ class PromptTokensDetailsWrapper( del self.web_search_requests if self.tool_use_tokens is None: del self.tool_use_tokens + if self.cache_write_tokens is None: + del self.cache_write_tokens if self.cache_creation_tokens is None: del self.cache_creation_tokens if self.cache_creation_token_details is None: @@ -1665,10 +1680,10 @@ class Usage(SafeAttributeModel, CompletionUsage): if "cache_creation_input_tokens" in params and isinstance(params["cache_creation_input_tokens"], int): if _prompt_tokens_details is None: _prompt_tokens_details = PromptTokensDetailsWrapper( - cache_creation_tokens=params["cache_creation_input_tokens"] + cache_write_tokens=params["cache_creation_input_tokens"] ) else: - _prompt_tokens_details.cache_creation_tokens = params["cache_creation_input_tokens"] + _prompt_tokens_details.cache_write_tokens = params["cache_creation_input_tokens"] super().__init__( prompt_tokens=prompt_tokens or 0, diff --git a/pyproject.toml b/pyproject.toml index 62bd37c3db6..a448ab042b8 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -198,6 +198,7 @@ e2e-dev = [ "playwright==1.61.0", "websockets>=15.0.1,<16.0", "locust==2.45.0", + "mcp>=1.28.1,<2.0", ] proxy-dev = [ "prisma==0.11.0", diff --git a/tests/code_coverage_tests/check_e2e_no_raw_requests.py b/tests/code_coverage_tests/check_e2e_no_raw_requests.py index e70e83652d1..fe6a77fc26c 100644 --- a/tests/code_coverage_tests/check_e2e_no_raw_requests.py +++ b/tests/code_coverage_tests/check_e2e_no_raw_requests.py @@ -2,8 +2,10 @@ raw HTTP client imports (requests, urllib.request, httpx, aiohttp, http.client) are banned in suite code. Importing requests' exception types for catching is fine anywhere; a small allowlist grandfathers the files that legitimately make raw calls -(the transport itself, the root conftest liveness probe, and the claude_code version -resolver's constant registry URL fetch). Referenced by tests/e2e/CLAUDE.md.""" +(the transport itself, the root conftest liveness probe, the claude_code version +resolver's constant registry URL fetch, and the mcp OAuth client, whose httpx +client is the object the official mcp SDK's streamable_http_client requires and so +cannot go through the sync requests transport). Referenced by tests/e2e/CLAUDE.md.""" from __future__ import annotations @@ -19,6 +21,7 @@ ALLOWED_RAW_CLIENT_FILES = { "e2e_http.py": ("requests",), "conftest.py": ("requests",), "claude_code/pr_gate_version_resolver.py": ("urllib.request",), + "mcp/oauth_chat_client.py": ("httpx",), } EXCEPTION_ONLY_NAMES = frozenset({"RequestException", "ConnectionError", "Timeout", "HTTPError"}) diff --git a/tests/e2e/CLAUDE.md b/tests/e2e/CLAUDE.md index 0e39664e358..17aee22560c 100644 --- a/tests/e2e/CLAUDE.md +++ b/tests/e2e/CLAUDE.md @@ -14,7 +14,7 @@ Each subdirectory under `tests/e2e/` is one suite, scoped to an endpoint family - `quota_management/` - quota enforcement and accounting, one subfolder per behavior: `ratelimit/` (rpm/tpm blocks, window reset, pacing headers on live traffic), `budgets/` (budget definition, enforcement, and reset windows: key, team, tag, soft, multi-window), and `spend_tracking/` (spend logging and cost attribution on `/spend/*`) - `management/` - key/team/user/organization management routes: create/update/delete persistence via the info routes, team membership, and llm-only-key route denials (API surface; not Playwright) - `a2a/` - the A2A (agent-to-agent) surface: admin registration via `/v1/agents`, proxy-fronted card discovery at `/.well-known/agent-card.json`, and JSON-RPC `message/send` invocation, driving agents backed by the litellm completion bridge (a real provider) and asserting protocol-version normalization (0.3 vs 1.0) -- `mcp/` - the MCP server surface over api_key auth against the real Datadog remote MCP server only (see "MCP suite: real Datadog only" below) +- `mcp/` - the MCP server surface over api_key auth against the real Datadog remote MCP server (see "MCP suite: real Datadog only" below); plus the gateway-managed OAuth (authorization_code) path exercised through `/chat/completions`, the one behavior Datadog's static-header auth cannot reach, seeding the per-user upstream token via the interactive authorize dance driven with the mcp SDK's own OAuth client (headless-browser consent from a saved session) and asserting the completion lists and executes the server's tools with the stored per-user token - `logging/` - logging-integration delivery (datadog and friends) - `security/` - secret handling and log-leak protection - `router/` - routing and reliability behavior (fallbacks, cooldowns) @@ -33,6 +33,7 @@ Every test under `tests/e2e/mcp/` must exercise the proxy against the real Datad - Prefer calling real Datadog tools that prove the product path (e.g. `search_datadog_logs` for list/call and permission denials). Seed a unique marker (`e2e-datadog-mcp-*`) in a chat completion when you need a log the tool can find; dual-read with `dd_logs` from conftest when delivery matters - Delete the MCP server (and any keys) through `resources.defer` the same way every other suite tears down - If a new MCP behavior cannot be covered with Datadog's tool surface, say so in the PR and get agreement before inventing another upstream; the default is always Datadog +- The one standing exception is `test_mcp_chat_completion_oauth_e2e.py`. Datadog authenticates with the static `DD-API-KEY` / `DD-APPLICATION-KEY` headers and exposes no authorize/token dance at all, so it cannot exercise gateway-managed OAuth or per-user token seeding in any form. That test drives a real Linear MCP server instead; it is still a real remote upstream, so the no-mock, no-fixture rule above holds unchanged ## Lay the pattern down in a class diff --git a/tests/e2e/batches/test_batches_e2e.py b/tests/e2e/batches/test_batches_e2e.py index b0c53becb6b..f9cd2a3f15f 100644 --- a/tests/e2e/batches/test_batches_e2e.py +++ b/tests/e2e/batches/test_batches_e2e.py @@ -23,7 +23,7 @@ from typing import Callable import pytest -from e2e_config import require_env, unique_marker +from e2e_config import unique_marker from batch_client import ( UPLOAD_FILENAME, @@ -702,14 +702,7 @@ class TestBedrockBatchAssumeRole: def test_unified_batch_create_with_assume_role( self, client: BatchClient, resources: ResourceManager ) -> None: - (role_arn,) = require_env("AWS_ROLE_NAME") - require_env( - "AWS_ACCESS_KEY_ID", - "AWS_SECRET_ACCESS_KEY", - "AWS_REGION", - "AWS_BATCH_S3_BUCKET", - "AWS_BATCH_ROLE_ARN", - ) + role_arn = os.environ["AWS_ROLE_NAME"] session_name = f"e2e-batch-sts-{unique_marker()}"[:64] model_name = batch_model_name("bedrock-sts-batch") @@ -819,7 +812,7 @@ class TestHostedVllmBatch: def test_unified_file_and_batch_create( self, client: BatchClient, resources: ResourceManager ) -> None: - (api_base,) = require_env("HOSTED_VLLM_API_BASE") + api_base = os.environ["HOSTED_VLLM_API_BASE"] api_key = (os.environ.get("HOSTED_VLLM_API_KEY") or "").strip() or None model_id = ( os.environ.get("HOSTED_VLLM_MODEL") or "meta-llama/Llama-3.2-3B-Instruct" diff --git a/tests/e2e/e2e_config.py b/tests/e2e/e2e_config.py index e7c48690c0a..feed680bd4a 100644 --- a/tests/e2e/e2e_config.py +++ b/tests/e2e/e2e_config.py @@ -38,6 +38,9 @@ UI_BASE_URL = os.environ.get("E2E_UI_BASE_URL", PROXY_BASE_URL).rstrip("/") CHEAP_ANTHROPIC_MODEL = os.environ.get("E2E_CHEAP_ANTHROPIC_MODEL", "claude-haiku-4-5") CHEAP_OPENAI_MODEL = os.environ.get("E2E_CHEAP_OPENAI_MODEL", "gpt-5.5") +LINEAR_MCP_URL = os.environ.get("E2E_LINEAR_MCP_URL", "https://mcp.linear.app/mcp") +LINEAR_STORAGE_STATE = os.environ.get("E2E_LINEAR_STORAGE_STATE", "") + # Jaeger query API of the compose stack's OTEL trace destination (the `jaeger` # service in docker-compose.yml maps it to host 16686). Trace-completeness tests # read exported spans back through it. @@ -99,22 +102,6 @@ ANOMALY_SPEND_SETTLE_SECONDS = float( ) -def require_env(*names: str) -> tuple[str, ...]: - """Return the non-empty values for each env name, or hard-fail naming which are missing. - - Live e2e never skips for missing credentials: a missing key is a red run so - ops knows the suite cannot prove the product path. - """ - missing = tuple(name for name in names if not (os.environ.get(name) or "").strip()) - if missing: - joined = ", ".join(missing) - raise AssertionError( - f"missing required env for e2e: {joined}. " - "Add them to tests/e2e/.env locally and to litellm ops for stage/CI." - ) - return tuple((os.environ.get(name) or "").strip() for name in names) - - def datadog_mcp_url(*, toolsets: str = "core") -> str: """Regional Datadog remote MCP endpoint for this process's DD_SITE. diff --git a/tests/e2e/guardrails/test_bedrock_guardrail_e2e.py b/tests/e2e/guardrails/test_bedrock_guardrail_e2e.py index 9e41b8808e8..a2408f0021e 100644 --- a/tests/e2e/guardrails/test_bedrock_guardrail_e2e.py +++ b/tests/e2e/guardrails/test_bedrock_guardrail_e2e.py @@ -8,9 +8,11 @@ a 200 means the guardrail never ran. from __future__ import annotations +import os + import pytest -from e2e_config import require_env, unique_marker +from e2e_config import unique_marker from e2e_http import UnknownApiError from guardrails_client import GuardrailsClient from lifecycle import ResourceManager @@ -33,11 +35,8 @@ class TestBedrockGuardrail: def test_bedrock_pre_call_blocks_harmful_prompt( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: - (identifier, version) = require_env( - "BEDROCK_GUARDRAIL_IDENTIFIER", - "BEDROCK_GUARDRAIL_VERSION", - ) - require_env("AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY", "AWS_REGION") + identifier = os.environ["BEDROCK_GUARDRAIL_IDENTIFIER"] + version = os.environ["BEDROCK_GUARDRAIL_VERSION"] name = f"e2e-bedrock-guard-{unique_marker()}" guardrail_id = client.create_bedrock_guardrail( diff --git a/tests/e2e/guardrails/test_block_code_execution_guardrail_e2e.py b/tests/e2e/guardrails/test_block_code_execution_guardrail_e2e.py index e36fc7c3f9d..de087b190d0 100644 --- a/tests/e2e/guardrails/test_block_code_execution_guardrail_e2e.py +++ b/tests/e2e/guardrails/test_block_code_execution_guardrail_e2e.py @@ -16,7 +16,7 @@ from __future__ import annotations import pytest -from e2e_config import require_env, unique_marker +from e2e_config import unique_marker from e2e_http import unwrap from guardrails_client import BlockCodeExecutionParamsBody, GuardrailsClient from lifecycle import ResourceManager @@ -46,7 +46,6 @@ class TestBlockCodeExecutionGuardrail: def test_blocks_execution_request_but_allows_explanation( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: - require_env("GEMINI_API_KEY") model = client.create_backend_model(resources, prefix="e2e-blockcode-backend") name = f"e2e-block-code-{unique_marker()}" diff --git a/tests/e2e/guardrails/test_openai_moderation_guardrail_e2e.py b/tests/e2e/guardrails/test_openai_moderation_guardrail_e2e.py index 4e2fcbf8fba..39950259fb5 100644 --- a/tests/e2e/guardrails/test_openai_moderation_guardrail_e2e.py +++ b/tests/e2e/guardrails/test_openai_moderation_guardrail_e2e.py @@ -14,7 +14,7 @@ from __future__ import annotations import pytest -from e2e_config import require_env, unique_marker +from e2e_config import unique_marker from e2e_http import UnknownApiError, unwrap from guardrails_client import GuardrailsClient, OpenAIModerationParamsBody from lifecycle import ResourceManager @@ -34,7 +34,6 @@ class TestOpenAIModerationGuardrail: def test_moderation_blocks_flagged_input( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: - require_env("OPENAI_API_KEY", "GEMINI_API_KEY") model = client.create_backend_model(resources, prefix="e2e-moderation-backend") name = f"e2e-openai-moderation-{unique_marker()}" diff --git a/tests/e2e/guardrails/test_presidio_guardrail_e2e.py b/tests/e2e/guardrails/test_presidio_guardrail_e2e.py index a911f387382..d103714b1dd 100644 --- a/tests/e2e/guardrails/test_presidio_guardrail_e2e.py +++ b/tests/e2e/guardrails/test_presidio_guardrail_e2e.py @@ -25,11 +25,12 @@ The chat backend is a gemini deployment created for the test. from __future__ import annotations +import os import time import pytest -from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, require_env, unique_marker +from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, unique_marker from e2e_http import NoBody, require_successful_call, unwrap from guardrails_client import GuardrailMode, GuardrailsClient, PresidioParamsBody from lifecycle import ResourceManager @@ -88,9 +89,8 @@ def _poll_logged_prompt(reader: OtelReader, *, call_id: str, genai_span: str) -> def _presidio_params( mode: GuardrailMode, *, apply_to_output: bool = False, logging_only: bool = False ) -> PresidioParamsBody: - analyzer, anonymizer = require_env( - "PRESIDIO_ANALYZER_API_BASE", "PRESIDIO_ANONYMIZER_API_BASE" - ) + analyzer = os.environ["PRESIDIO_ANALYZER_API_BASE"] + anonymizer = os.environ["PRESIDIO_ANONYMIZER_API_BASE"] return PresidioParamsBody( mode=mode, default_on=False, @@ -124,7 +124,6 @@ class TestPresidioGuardrail: def test_pre_call_masks_pii_before_the_model_sees_it( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: - require_env("GEMINI_API_KEY") model = client.create_backend_model(resources, prefix="e2e-presidio-pre") name = f"e2e-presidio-pre-{unique_marker()}" guardrail_id = client.register(name, _presidio_params("pre_call")) @@ -149,7 +148,6 @@ class TestPresidioGuardrail: def test_post_call_masks_pii_in_model_output( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: - require_env("GEMINI_API_KEY") model = client.create_backend_model(resources, prefix="e2e-presidio-post") name = f"e2e-presidio-post-{unique_marker()}" guardrail_id = client.register(name, _presidio_params("post_call", apply_to_output=True)) @@ -173,7 +171,6 @@ class TestPresidioGuardrail: def test_logging_only_masks_the_logged_prompt( self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: - require_env("GEMINI_API_KEY") _require_otel_v2_active(client) reader = build_otel_reader() diff --git a/tests/e2e/llm_translation/test_chat_completions_regression_e2e.py b/tests/e2e/llm_translation/test_chat_completions_regression_e2e.py index 8d3622e441a..af0e782e224 100644 --- a/tests/e2e/llm_translation/test_chat_completions_regression_e2e.py +++ b/tests/e2e/llm_translation/test_chat_completions_regression_e2e.py @@ -21,7 +21,7 @@ import os import pytest from pydantic import BaseModel -from e2e_config import require_env, unique_marker +from e2e_config import unique_marker from e2e_http import StreamingResponse, unwrap from lifecycle import ResourceManager from models import ( @@ -250,7 +250,7 @@ class TestCohereChat: def test_cohere_chat_returns_content( self, client: PassthroughClient, resources: ResourceManager ) -> None: - (cohere_key,) = require_env("COHERE_API_KEY") + cohere_key = os.environ["COHERE_API_KEY"] model = f"e2e-cohere-chat-{unique_marker()}" model_id = client.proxy.create_model( model, @@ -343,7 +343,7 @@ class TestHostedVllmChat: def test_hosted_vllm_chat_returns_content( self, client: PassthroughClient, resources: ResourceManager ) -> None: - (api_base,) = require_env("HOSTED_VLLM_API_BASE") + api_base = os.environ["HOSTED_VLLM_API_BASE"] api_key = (os.environ.get("HOSTED_VLLM_API_KEY") or "").strip() or None backend = ( os.environ.get("HOSTED_VLLM_MODEL") or "meta-llama/Llama-3.2-3B-Instruct" @@ -395,7 +395,6 @@ class TestOpenAIChatCompletions: def test_openai_chat_streams_real_content( self, client: PassthroughClient, resources: ResourceManager ) -> None: - require_env("OPENAI_API_KEY") model = f"e2e-openai-chat-{unique_marker()}" model_id = client.proxy.create_model( model, LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY") @@ -423,7 +422,6 @@ class TestOpenAIChatCompletions: def test_openai_chat_logs_cost( self, client: PassthroughClient, resources: ResourceManager ) -> None: - require_env("OPENAI_API_KEY") model = f"e2e-openai-cost-{unique_marker()}" model_id = client.proxy.create_model( model, LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY") @@ -457,7 +455,6 @@ class TestOpenAIChatCompletions: def test_openai_chat_returns_tool_call( self, client: PassthroughClient, resources: ResourceManager ) -> None: - require_env("OPENAI_API_KEY") model = f"e2e-openai-tool-{unique_marker()}" model_id = client.proxy.create_model( model, LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY") @@ -488,7 +485,6 @@ class TestOpenAIChatCompletions: def test_openai_chat_structured_output_conforms_to_schema( self, client: PassthroughClient, resources: ResourceManager ) -> None: - require_env("OPENAI_API_KEY") model = f"e2e-openai-schema-{unique_marker()}" model_id = client.proxy.create_model( model, LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY") @@ -522,7 +518,6 @@ class TestOpenAIChatCompletions: def test_openai_chat_reasoning_reports_reasoning_tokens( self, client: PassthroughClient, resources: ResourceManager ) -> None: - require_env("OPENAI_API_KEY") model = f"e2e-openai-reasoning-{unique_marker()}" model_id = client.proxy.create_model( model, LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY") @@ -561,7 +556,6 @@ class TestOpenAIChatCompletions: def test_openai_chat_vision_describes_image( self, client: PassthroughClient, resources: ResourceManager ) -> None: - require_env("OPENAI_API_KEY") model = f"e2e-openai-vision-{unique_marker()}" model_id = client.proxy.create_model( model, LiteLLMParamsBody(model=OPENAI_VISION_BACKEND, api_key="os.environ/OPENAI_API_KEY") @@ -579,7 +573,6 @@ class TestOpenAIChatCompletions: def test_openai_chat_prompt_cache_hits_on_repeat( self, client: PassthroughClient, resources: ResourceManager ) -> None: - require_env("OPENAI_API_KEY") model = f"e2e-openai-cache-{unique_marker()}" model_id = client.proxy.create_model( model, LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY") @@ -610,7 +603,6 @@ class TestOpenAIChatCompletions: def test_openai_chat_streams_tool_call( self, client: PassthroughClient, resources: ResourceManager ) -> None: - require_env("OPENAI_API_KEY") model = f"e2e-openai-tool-stream-{unique_marker()}" model_id = client.proxy.create_model( model, LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY") @@ -645,7 +637,6 @@ class TestBedrockConverseChatCompletions: """ def _register(self, client: PassthroughClient, resources: ResourceManager, prefix: str) -> str: - require_env("AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY", "AWS_REGION") model = f"{prefix}-{unique_marker()}" model_id = client.proxy.create_model(model, _bedrock_params()) resources.defer(lambda: client.proxy.delete_model(model_id)) diff --git a/tests/e2e/llm_translation/test_image_generation_e2e.py b/tests/e2e/llm_translation/test_image_generation_e2e.py index 1ba78a7e083..45861d1e93a 100644 --- a/tests/e2e/llm_translation/test_image_generation_e2e.py +++ b/tests/e2e/llm_translation/test_image_generation_e2e.py @@ -9,7 +9,7 @@ from __future__ import annotations import pytest -from e2e_config import require_env, unique_marker +from e2e_config import unique_marker from e2e_http import require_successful_call from endpoints_client import EndpointsClient, ImagesResult from lifecycle import ResourceManager @@ -50,7 +50,6 @@ class TestImageGeneration: def test_bedrock_image_generation_returns_image( self, endpoints_client: EndpointsClient, resources: ResourceManager ) -> None: - require_env("AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY", "AWS_REGION") model = f"e2e-bedrock-image-{unique_marker()}" model_id = endpoints_client.create_model( model, diff --git a/tests/e2e/llm_translation/test_messages_e2e.py b/tests/e2e/llm_translation/test_messages_e2e.py index 44376218c6b..ef6ba5b95d3 100644 --- a/tests/e2e/llm_translation/test_messages_e2e.py +++ b/tests/e2e/llm_translation/test_messages_e2e.py @@ -10,7 +10,7 @@ from __future__ import annotations import pytest -from e2e_config import require_env, unique_marker +from e2e_config import unique_marker from e2e_http import require_successful_call, unwrap from endpoints_client import EndpointsClient, MessagesResult from lifecycle import ResourceManager @@ -73,7 +73,6 @@ class TestAnthropicMessages: def test_messages_logs_cost_matching_the_response_header( self, endpoints_client: EndpointsClient, resources: ResourceManager ) -> None: - require_env("ANTHROPIC_API_KEY") model = f"e2e-messages-cost-{unique_marker()}" model_id = endpoints_client.create_model( model, diff --git a/tests/e2e/llm_translation/test_rerank_e2e.py b/tests/e2e/llm_translation/test_rerank_e2e.py index 0857ff65a52..c3614251e77 100644 --- a/tests/e2e/llm_translation/test_rerank_e2e.py +++ b/tests/e2e/llm_translation/test_rerank_e2e.py @@ -9,7 +9,7 @@ from __future__ import annotations import pytest -from e2e_config import require_env, unique_marker +from e2e_config import unique_marker from e2e_http import require_successful_call from endpoints_client import EndpointsClient, RerankResult from lifecycle import ResourceManager @@ -56,7 +56,6 @@ class TestRerank: def test_bedrock_rerank_scores_top_n( self, endpoints_client: EndpointsClient, resources: ResourceManager ) -> None: - require_env("AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY", "AWS_REGION") model = f"e2e-bedrock-rerank-{unique_marker()}" model_id = endpoints_client.create_model( model, diff --git a/tests/e2e/llm_translation/test_responses_e2e.py b/tests/e2e/llm_translation/test_responses_e2e.py index d24d2b53b71..0b2ffce5b2a 100644 --- a/tests/e2e/llm_translation/test_responses_e2e.py +++ b/tests/e2e/llm_translation/test_responses_e2e.py @@ -13,7 +13,7 @@ from typing import cast import pytest from pydantic import BaseModel, ValidationError -from e2e_config import require_env, unique_marker +from e2e_config import unique_marker from e2e_http import require_successful_call from endpoints_client import ( EndpointsClient, @@ -255,7 +255,6 @@ class TestResponses: def test_responses_bedrock_returns_completion( self, endpoints_client: EndpointsClient, resources: ResourceManager ) -> None: - require_env("AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY", "AWS_REGION") model = f"e2e-responses-{unique_marker()}" model_id = endpoints_client.create_model(model, _bedrock_params()) resources.defer(lambda: endpoints_client.delete_model(model_id)) @@ -270,7 +269,6 @@ class TestResponses: def test_responses_bedrock_returns_function_call( self, endpoints_client: EndpointsClient, resources: ResourceManager ) -> None: - require_env("AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY", "AWS_REGION") model = f"e2e-responses-{unique_marker()}" model_id = endpoints_client.create_model(model, _bedrock_params()) resources.defer(lambda: endpoints_client.delete_model(model_id)) diff --git a/tests/e2e/llm_translation/test_responses_metadata_e2e.py b/tests/e2e/llm_translation/test_responses_metadata_e2e.py index 6cf24348095..df854dcfa19 100644 --- a/tests/e2e/llm_translation/test_responses_metadata_e2e.py +++ b/tests/e2e/llm_translation/test_responses_metadata_e2e.py @@ -14,7 +14,7 @@ import time import pytest from pydantic import BaseModel, ConfigDict -from e2e_config import require_env, unique_marker +from e2e_config import unique_marker from e2e_http import require_successful_call from endpoints_client import EndpointsClient, ResponsesResult from lifecycle import ResourceManager @@ -42,7 +42,7 @@ class RedisKeyInfo(BaseModel): def _redis_scan(marker: str) -> tuple[RedisKeyInfo, ...]: import redis - (host,) = require_env("REDIS_HOST") + host = os.environ["REDIS_HOST"] port = int((os.environ.get("REDIS_PORT") or "6379").strip() or "6379") try: with socket.create_connection((host, port), timeout=3): diff --git a/tests/e2e/mcp/linear_session_capture.py b/tests/e2e/mcp/linear_session_capture.py new file mode 100644 index 00000000000..1c867e17e22 --- /dev/null +++ b/tests/e2e/mcp/linear_session_capture.py @@ -0,0 +1,58 @@ +"""One-time helper to capture a logged-in Linear browser session for the +real-Linear MCP e2e test. + +The real-Linear test drives the genuine gateway-managed authorization_code +dance against ``mcp.linear.app``. The only step that cannot be scripted is +Linear's login (magic link / SSO), so a human authenticates once here and the +resulting session (cookies + local storage) is persisted to disk. The e2e test +then loads that session in a headless Playwright context and clicks Approve on +Linear's consent screen every run, with no human and no login automation. + +Run it with the e2e venv, log into Linear in the window that opens, then return +to the terminal and press Enter: + + LITELLM=~/litellm-mcpe2e + "$LITELLM"/.venv/bin/python "$LITELLM"/tests/e2e/mcp/linear_session_capture.py + +The session is written to ``E2E_LINEAR_STORAGE_STATE`` (default +``~/.litellm-e2e/linear_storage_state.json``), outside the repo. It is a +secret: never commit it. Re-run this whenever Linear expires the session. +""" + +from __future__ import annotations + +import os +from pathlib import Path + +from playwright.sync_api import sync_playwright + +DEFAULT_STATE_PATH = Path.home() / ".litellm-e2e" / "linear_storage_state.json" + + +def capture(state_path: Path) -> None: + """Open a headed browser at Linear, wait for the human to log in, then save + the authenticated session to ``state_path``.""" + state_path.parent.mkdir(parents=True, exist_ok=True) + with sync_playwright() as playwright: + browser = playwright.chromium.launch(headless=False) + context = browser.new_context() + page = context.new_page() + page.goto("https://linear.app/login", wait_until="domcontentloaded") + print("\n" + "=" * 72) + print("Log into Linear in the browser window that just opened.") + print("If Linear emails you a magic link, paste the link into THIS window's") + print("address bar (opening it in your default browser won't capture the") + print("session). Google SSO works too as long as you complete it here.") + print("When your Linear workspace has loaded, come back and press Enter.") + print("=" * 72) + input("Press Enter once you are logged in... ") + page.goto("https://mcp.linear.app/", wait_until="domcontentloaded") + context.storage_state(path=str(state_path)) + browser.close() + print(f"\nSaved Linear session to {state_path}") + print("Point the e2e test at it with:") + print(f' export E2E_LINEAR_STORAGE_STATE="{state_path}"') + + +if __name__ == "__main__": + capture(Path(os.environ.get("E2E_LINEAR_STORAGE_STATE", str(DEFAULT_STATE_PATH)))) diff --git a/tests/e2e/mcp/oauth_chat_client.py b/tests/e2e/mcp/oauth_chat_client.py new file mode 100644 index 00000000000..2eaf512cfa5 --- /dev/null +++ b/tests/e2e/mcp/oauth_chat_client.py @@ -0,0 +1,271 @@ +"""Client for the mcp chat-completion OAuth e2e suite. + +Registers a gateway-managed OAuth (authorization_code) MCP server, seeds the +per-user upstream token by driving the interactive authorize dance with the +official mcp SDK's OAuthClientProvider (the browser leg is a headless Chromium +primed with a human's saved Linear session), then exercises the server through +/chat/completions, where the gateway lists and executes its tools with the +stored per-user token. + +Management routes (/v1/mcp/server CRUD, /chat/completions) go through the +shared ProxyClient transport. The MCP protocol used to seed the token goes through +the mcp SDK, the same library production MCP hosts run. +""" + +from __future__ import annotations + +import asyncio +import re +import time +from dataclasses import dataclass +from typing import TYPE_CHECKING +from urllib.parse import parse_qsl + +import httpx +import pytest +from mcp import ClientSession +from mcp.client.auth import OAuthClientProvider +from mcp.client.streamable_http import streamable_http_client +from mcp.shared.auth import OAuthClientInformationFull, OAuthClientMetadata, OAuthToken + +from e2e_config import PROXY_BASE_URL, REQUEST_TIMEOUT +from proxy_client import ProxyClient +from e2e_http import AuthHeaders, NoBody, unwrap +from models import ChatBody, ChatResponse, McpServerCreateBody, McpServerInfo + +if TYPE_CHECKING: + from playwright.async_api import Route + +# Where the "browser" lands at the end of the authorize dance. Nothing listens +# here: the route interceptor short-circuits the final redirect and reads the +# code/state off its query string, exactly like a desktop MCP host intercepting +# its loopback redirect. +OAUTH_CLIENT_REDIRECT_URI = "http://127.0.0.1:53682/e2e/callback" +BROWSER_CONSENT_TIMEOUT = 60.0 + + +def _mcp_url(alias: str) -> str: + return f"{PROXY_BASE_URL}/{alias}/mcp" + + +class InMemoryTokenStorage: + """The mcp SDK's TokenStorage protocol, in memory for one dance: the + DCR-registered client and the gateway tokens minted for it.""" + + def __init__(self) -> None: + self._tokens: OAuthToken | None = None + self._client_info: OAuthClientInformationFull | None = None + + async def get_tokens(self) -> OAuthToken | None: + return self._tokens + + async def set_tokens(self, tokens: OAuthToken) -> None: + self._tokens = tokens + + async def get_client_info(self) -> OAuthClientInformationFull | None: + return self._client_info + + async def set_client_info(self, client_info: OAuthClientInformationFull) -> None: + self._client_info = client_info + + +async def _browser_follow_authorize(start_url: str, storage_state_path: str) -> tuple[str, str | None]: + """Play the browser's role for a real upstream whose authorize endpoint + serves an interactive consent page (Linear). A headless Chromium primed + with a human's saved Linear session opens the gateway authorize URL and + clicks through Linear's consent screens (the mcp.linear.app Approve form, + then the linear.app workspace-selection page), riding the rest of the chain + (Linear -> gateway callback -> host redirect_uri). The final hop is + intercepted and short-circuited, since nothing listens there, and its + code/state are read off the query string.""" + from playwright.async_api import async_playwright + + captured: dict[str, str] = {} # mutable-ok: hand-off from the request listener + trail: list[str] = [] # mutable-ok: navigation diagnostics for a failed dance + + def _note_request(request: object) -> None: + url = getattr(request, "url", "") + if url.startswith(OAUTH_CLIENT_REDIRECT_URI) and "url" not in captured: + captured["url"] = url + + async def _swallow_redirect(route: "Route") -> None: + await route.fulfill(status=200, content_type="text/plain", body="ok") + + async with async_playwright() as playwright: + browser = await playwright.chromium.launch(headless=True) + context = await browser.new_context(storage_state=storage_state_path) + await context.route(re.compile(re.escape(OAUTH_CLIENT_REDIRECT_URI) + r".*"), _swallow_redirect) + page = await context.new_page() + page.on("request", _note_request) + page.on("framenavigated", lambda frame: trail.append(frame.url.split("?", 1)[0])) + await page.goto(start_url, wait_until="domcontentloaded") + deadline = time.monotonic() + BROWSER_CONSENT_TIMEOUT + while "url" not in captured and time.monotonic() < deadline: + try: + await page.wait_for_load_state("networkidle", timeout=8000) + except Exception: # noqa: BLE001 - a busy consent page never idles; fall through and try to advance it + pass + if "url" in captured: + break + control = page.locator( + 'button[name="action"][value="approve"], button:has-text("Authorize"), ' + 'button:has-text("Allow"), button:has-text("@"), a:has-text("@")' + ).first + try: + await control.click(timeout=5000) + except Exception: # noqa: BLE001 - nothing to advance yet; loop and re-check + await asyncio.sleep(0.5) + final_url = page.url + await browser.close() + + landing = captured.get("url") + assert landing is not None, ( + f"consent flow never reached {OAUTH_CLIENT_REDIRECT_URI}; " + f"final={final_url.split('?', 1)[0]!r}; trail={trail[-6:]}" + ) + params = dict(parse_qsl(httpx.URL(landing).query.decode())) + assert "code" in params, f"client redirect_uri carried no code: {landing}" + return params["code"], params.get("state") + + +def _oauth_provider(url: str, storage: InMemoryTokenStorage, storage_state_path: str) -> OAuthClientProvider: + """The SDK's real OAuth machinery (RFC 9728/8414 discovery, RFC 7591 DCR, + PKCE, token exchange) with the browser leg driven by Playwright against the + upstream's consent screen.""" + code_holder: dict[str, str | None] = {} # mutable-ok: hand-off between the two SDK callbacks + + async def redirect_handler(authorize_url: str) -> None: + code, state = await _browser_follow_authorize(authorize_url, storage_state_path) + code_holder["code"] = code + code_holder["state"] = state + + async def callback_handler() -> tuple[str, str | None]: + code = code_holder.get("code") + assert code is not None, "callback_handler ran before the authorize redirect completed" + return code, code_holder.get("state") + + return OAuthClientProvider( + server_url=url, + client_metadata=OAuthClientMetadata.model_validate( + { + "redirect_uris": [OAUTH_CLIENT_REDIRECT_URI], + "token_endpoint_auth_method": "none", + "grant_types": ["authorization_code", "refresh_token"], + "response_types": ["code"], + "client_name": "e2e-mcp-host", + } + ), + storage=storage, + redirect_handler=redirect_handler, + callback_handler=callback_handler, + ) + + +class _HeaderInjectingTransport(httpx.AsyncBaseTransport): + """Adds the caller's LiteLLM key header to every outgoing SDK request + (discovery, DCR, token exchange), so the gateway resolves which user to + store the upstream token for from the key on the token exchange, exactly + like a production MCP host configured with a LiteLLM key header.""" + + def __init__(self, inner: httpx.AsyncBaseTransport, headers: dict[str, str]) -> None: + self._inner = inner + self._headers = headers + + async def handle_async_request(self, request: httpx.Request) -> httpx.Response: + for name, value in self._headers.items(): + if name not in request.headers: + request.headers[name] = value + return await self._inner.handle_async_request(request) + + +def _oauth_http_client(headers: dict[str, str], auth: OAuthClientProvider) -> httpx.AsyncClient: + return httpx.AsyncClient( + headers=headers, + auth=auth, + timeout=httpx.Timeout(REQUEST_TIMEOUT), + follow_redirects=True, + transport=_HeaderInjectingTransport(httpx.AsyncHTTPTransport(), headers), + ) + + +async def _seed_via_dance( + url: str, headers: dict[str, str], storage: InMemoryTokenStorage, storage_state_path: str +) -> tuple[str, ...]: + async with _oauth_http_client(headers, _oauth_provider(url, storage, storage_state_path)) as http_client: + async with streamable_http_client(url, http_client=http_client) as (read, write, _): + async with ClientSession(read, write) as session: + await session.initialize() + listed = await session.list_tools() + return tuple(sorted(tool.name for tool in listed.tools)) + + +@dataclass(frozen=True, slots=True) +class ChatMcpClient: + proxy: ProxyClient + + def create_server(self, body: McpServerCreateBody) -> McpServerInfo: + return unwrap( + self.proxy.transport.post( + "/v1/mcp/server", + headers=self.proxy.transport.master, + json=body, + response_type=McpServerInfo, + ) + ) + + def server_info(self, server_id: str) -> McpServerInfo: + return unwrap( + self.proxy.transport.get( + f"/v1/mcp/server/{server_id}", + headers=self.proxy.transport.master, + params=NoBody(), + response_type=McpServerInfo, + ) + ) + + def delete_server(self, server_id: str) -> None: + _ = self.proxy.transport.delete( + f"/v1/mcp/server/{server_id}", + headers=self.proxy.transport.master, + json=NoBody(), + response_type=NoBody, + ) + + def seed_user_token(self, alias: str, key: str, storage_state_path: str) -> tuple[str, ...]: + """Drive the interactive authorize dance for `key`'s user so the gateway + stores their upstream token, retried to the shared deadline since the + just-created server and key propagate asynchronously. The LiteLLM key + rides x-litellm-api-key so the gateway binds the token to that user. + Returns the upstream tool names the dance listed, proof the token works.""" + headers = {"x-litellm-api-key": f"Bearer {key}"} + storage = InMemoryTokenStorage() + deadline = time.monotonic() + self.proxy.poll_timeout + last_error: Exception | None = None + while time.monotonic() < deadline: + try: + return asyncio.run(_seed_via_dance(_mcp_url(alias), headers, storage, storage_state_path)) + except Exception as exc: # noqa: BLE001 - retried to the deadline; the last error surfaces below + last_error = exc + time.sleep(self.proxy.poll_interval) + pytest.fail( + f"authorize dance for {alias!r} never completed within {self.proxy.poll_timeout}s; " + f"last error: {last_error!r}" + ) + + def chat_with_mcp(self, headers: AuthHeaders, body: ChatBody) -> ChatResponse: + """POST /chat/completions carrying the LiteLLM key in `headers` (either + ingress form) with an MCP server attached in `body.tools`. The gateway + resolves the user from the key and lists/executes the server's tools + with that user's stored upstream token.""" + return unwrap( + self.proxy.transport.post( + "/chat/completions", + headers=headers, + json=body, + response_type=ChatResponse, + ) + ) + + +def build_chat_client(proxy: ProxyClient) -> ChatMcpClient: + return ChatMcpClient(proxy=proxy) diff --git a/tests/e2e/mcp/test_mcp_chat_completion_oauth_e2e.py b/tests/e2e/mcp/test_mcp_chat_completion_oauth_e2e.py new file mode 100644 index 00000000000..01e94f7b86f --- /dev/null +++ b/tests/e2e/mcp/test_mcp_chat_completion_oauth_e2e.py @@ -0,0 +1,197 @@ +"""On-demand e2e: a chat completion drives a gateway-managed OAuth MCP server. + +The real end-user flow for MCP over an OAuth server: a user registers a Linear +authorization_code server, authorizes it once so the gateway stores their +upstream token, then sends a normal /chat/completions request with the Linear +MCP attached. The gateway resolves the user from the LiteLLM key, lists Linear's +tools with the stored per-user token, lets the model call one, executes it +upstream with that token, and returns the answer. This is proven against the +real Linear MCP server (mcp.linear.app) and a real Anthropic model, once per +documented ingress header (x-litellm-api-key and Authorization). + +The authorize dance is seeded through the mcp SDK's OAuthClientProvider; the one +step Linear cannot auto-approve is the human consent, so it is captured once out +of band (mcp/linear_session_capture.py) into a saved browser session and a +headless Chromium clicks Approve every run. The test therefore skips unless +E2E_LINEAR_STORAGE_STATE points at that session, so it never runs on the per-PR +CI path; it is a nightly/on-demand real-server smoke test. + +Fail-before-fix: without the stored per-user token the gateway lists no Linear +tools, so mcp_list_tools comes back empty, nothing is called, and the +assertions fail; a served, called, non-empty Linear tool proves the gateway +pulled and used the user's token. +""" + +from __future__ import annotations + +import os + +import pytest + +from e2e_config import CHEAP_ANTHROPIC_MODEL, LINEAR_MCP_URL, LINEAR_STORAGE_STATE, unique_marker +from e2e_http import AuthHeaders +from lifecycle import ResourceManager +from models import ChatBody, ChatMessage, KeyGenerateBody, McpChatTool, McpServerCreateBody, ObjectPermission +from proxy_client import ProxyClient + +pytest.importorskip("mcp", reason="mcp SDK not installed; run `uv sync --inexact --group e2e-dev`") +pytest.importorskip( + "playwright.async_api", + reason="playwright not installed; run `uv pip install playwright` and `playwright install chromium`", +) + +from oauth_chat_client import ChatMcpClient, build_chat_client # noqa: E402 # imports follow the importorskip guards + +pytestmark = [ + pytest.mark.e2e, + pytest.mark.skipif( + not LINEAR_STORAGE_STATE or not os.path.exists(LINEAR_STORAGE_STATE), + reason="set E2E_LINEAR_STORAGE_STATE to a Linear session captured via mcp/linear_session_capture.py", + ), +] + +# Pinned from a live dance during verification (never guessed); the gateway +# prefixes every upstream tool name with the server alias. list_teams is a +# read-only Linear tool that takes no arguments and returns the caller's teams. +LINEAR_READONLY_TOOL = "list_teams" +LINEAR_PROMPT = "Use the list_teams tool to list my Linear teams, then reply with the name of one of them." + + +@pytest.fixture(scope="session") +def chat_client(proxy: ProxyClient) -> ChatMcpClient: + return build_chat_client(proxy) + + +class TestMcpChatCompletionOauth: + """A scoped internal-user key on a real Linear authorization_code server, + used through /chat/completions once per ingress header: the gateway pulls + the user's stored upstream token, lists and executes Linear's tools during + the completion, and returns the answer.""" + + @pytest.mark.covers("mcp.list_tools.oauth.succeeds") + @pytest.mark.covers("mcp.call_tool.oauth.succeeds") + def test_chat_completion_uses_linear_with_x_litellm_api_key_header( + self, chat_client: ChatMcpClient, resources: ResourceManager + ) -> None: + marker = unique_marker() + alias = f"e2elinear{marker}" + created = chat_client.create_server( + McpServerCreateBody( + alias=alias, + url=LINEAR_MCP_URL, + allow_all_keys=False, + auth_type="oauth2", + oauth2_flow="authorization_code", + ) + ) + resources.defer(lambda: chat_client.delete_server(created.server_id)) + + stored = chat_client.server_info(created.server_id) + assert stored.auth_type == "oauth2" + assert stored.oauth2_flow == "authorization_code" + assert stored.allow_all_keys is False + + key = chat_client.proxy.generate_key( + KeyGenerateBody( + user_id="e2e-test-user", + object_permission=ObjectPermission(mcp_servers=[created.server_id]), + ) + ) + resources.defer(lambda: chat_client.proxy.delete_key(key)) + + seeded = chat_client.seed_user_token(alias, key, LINEAR_STORAGE_STATE) + assert f"{alias}-{LINEAR_READONLY_TOOL}" in seeded, ( + f"the authorize dance listed {seeded}, expected it to include {alias}-{LINEAR_READONLY_TOOL}" + ) + + response = chat_client.chat_with_mcp( + AuthHeaders.model_validate({"x-litellm-api-key": f"Bearer {key}"}), + ChatBody( + model=CHEAP_ANTHROPIC_MODEL, + messages=[ChatMessage(role="user", content=LINEAR_PROMPT)], + tools=[ + McpChatTool( + server_url=f"litellm_proxy/mcp/{alias}", + server_label=alias, + require_approval="never", + ) + ], + ), + ) + + message = response.choices[0].message + assert message is not None and message.content, f"completion returned no answer: {response}" + meta = message.provider_specific_fields + assert meta is not None, f"no MCP metadata on the completion: {response}" + listed = {t.function.name for t in (meta.mcp_list_tools or []) if t.function} + assert f"{alias}-{LINEAR_READONLY_TOOL}" in listed, ( + f"the gateway listed {sorted(listed)}, expected the stored token to surface {alias}-{LINEAR_READONLY_TOOL}" + ) + results = [r for r in (meta.mcp_call_results or []) if r.name == f"{alias}-{LINEAR_READONLY_TOOL}"] + assert results and results[0].result, ( + f"Linear tool {alias}-{LINEAR_READONLY_TOOL} was not executed with a result: {meta.mcp_call_results}" + ) + + @pytest.mark.covers("mcp.list_tools.oauth.succeeds") + @pytest.mark.covers("mcp.call_tool.oauth.succeeds") + def test_chat_completion_uses_linear_with_authorization_bearer_header( + self, chat_client: ChatMcpClient, resources: ResourceManager + ) -> None: + marker = unique_marker() + alias = f"e2elinear{marker}" + created = chat_client.create_server( + McpServerCreateBody( + alias=alias, + url=LINEAR_MCP_URL, + allow_all_keys=False, + auth_type="oauth2", + oauth2_flow="authorization_code", + ) + ) + resources.defer(lambda: chat_client.delete_server(created.server_id)) + + stored = chat_client.server_info(created.server_id) + assert stored.auth_type == "oauth2" + assert stored.oauth2_flow == "authorization_code" + assert stored.allow_all_keys is False + + key = chat_client.proxy.generate_key( + KeyGenerateBody( + user_id="e2e-test-user", + object_permission=ObjectPermission(mcp_servers=[created.server_id]), + ) + ) + resources.defer(lambda: chat_client.proxy.delete_key(key)) + + seeded = chat_client.seed_user_token(alias, key, LINEAR_STORAGE_STATE) + assert f"{alias}-{LINEAR_READONLY_TOOL}" in seeded, ( + f"the authorize dance listed {seeded}, expected it to include {alias}-{LINEAR_READONLY_TOOL}" + ) + + response = chat_client.chat_with_mcp( + AuthHeaders.model_validate({"authorization": f"Bearer {key}"}), + ChatBody( + model=CHEAP_ANTHROPIC_MODEL, + messages=[ChatMessage(role="user", content=LINEAR_PROMPT)], + tools=[ + McpChatTool( + server_url=f"litellm_proxy/mcp/{alias}", + server_label=alias, + require_approval="never", + ) + ], + ), + ) + + message = response.choices[0].message + assert message is not None and message.content, f"completion returned no answer: {response}" + meta = message.provider_specific_fields + assert meta is not None, f"no MCP metadata on the completion: {response}" + listed = {t.function.name for t in (meta.mcp_list_tools or []) if t.function} + assert f"{alias}-{LINEAR_READONLY_TOOL}" in listed, ( + f"the gateway listed {sorted(listed)}, expected the stored token to surface {alias}-{LINEAR_READONLY_TOOL}" + ) + results = [r for r in (meta.mcp_call_results or []) if r.name == f"{alias}-{LINEAR_READONLY_TOOL}"] + assert results and results[0].result, ( + f"Linear tool {alias}-{LINEAR_READONLY_TOOL} was not executed with a result: {meta.mcp_call_results}" + ) diff --git a/tests/e2e/models.py b/tests/e2e/models.py index b3ea9346180..d21920cf848 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -6,6 +6,7 @@ response validates without mirroring every proxy field. No untyped dicts. from __future__ import annotations +from collections.abc import Sequence from datetime import datetime from typing import Literal @@ -196,6 +197,19 @@ class ChatTool(BaseModel): function: ChatToolFunction +class McpChatTool(BaseModel): + """An MCP server attached to a chat completion (OpenAI `type: "mcp"` tool). + `server_url` selects the gateway-registered server by its alias suffix; with + `require_approval="never"` the gateway lists, calls, and feeds the server's + tools back to the model in one agentic turn.""" + + type: Literal["mcp"] = "mcp" + server_url: str + require_approval: str + server_label: str | None = None + allowed_tools: list[str] | None = None + + class ChatBody(BaseModel): model: str messages: list[ChatMessage] @@ -206,7 +220,7 @@ class ChatBody(BaseModel): reasoning_effort: str | None = None thinking: ThinkingParam | None = None service_tier: str | None = None - tools: list[ChatTool] | None = None + tools: Sequence[ChatTool | McpChatTool] | None = None tool_choice: str | None = None guardrails: list[str] | None = None response_format: dict[str, object] | None = None @@ -242,10 +256,46 @@ class ToolCall(BaseModel): function: ToolCallFunction = ToolCallFunction() +class McpToolFunctionRef(BaseModel): + name: str + + +class McpListedTool(BaseModel): + """One entry of `mcp_list_tools`: a tool the gateway listed from the + attached MCP server and exposed to the model, in OpenAI function shape.""" + + function: McpToolFunctionRef | None = None + + +class McpToolCall(BaseModel): + """One entry of `mcp_tool_calls`: a tool the model asked the gateway to run.""" + + function: McpToolFunctionRef | None = None + + +class McpCallResult(BaseModel): + """One entry of `mcp_call_results`: what the gateway got back from executing + a tool upstream on the caller's behalf.""" + + name: str | None = None + result: str | None = None + + +class McpResponseMetadata(BaseModel): + """`choices[].message.provider_specific_fields` MCP section: which tools the + gateway listed from the attached server, which the model called, and their + results. Populated only when the completion drove an MCP server.""" + + mcp_list_tools: list[McpListedTool] | None = None + mcp_tool_calls: list[McpToolCall] | None = None + mcp_call_results: list[McpCallResult] | None = None + + class OutMessage(BaseModel): content: str | None = None reasoning_content: str | None = None tool_calls: list[ToolCall] | None = None + provider_specific_fields: McpResponseMetadata | None = None class ChatChoice(BaseModel): @@ -355,6 +405,36 @@ class CountTokensResponse(BaseModel): input_tokens: int +# ---------- mcp servers ---------- + + +class McpServerCreateBody(BaseModel): + """POST /v1/mcp/server. For a gateway-managed OAuth server, `auth_type` is + `oauth2` and `oauth2_flow` is `authorization_code`; the upstream endpoints + are discovered and registered via DCR when left unset. `allow_all_keys` + false scopes the server to keys granted it through object_permission.""" + + alias: str + url: str + transport: str = "http" + allow_all_keys: bool = True + auth_type: str | None = None + oauth2_flow: Literal["client_credentials", "authorization_code"] | None = None + authorization_url: str | None = None + token_url: str | None = None + + +class McpServerInfo(BaseModel): + """Response of POST /v1/mcp/server and GET /v1/mcp/server/{server_id}.""" + + server_id: str + alias: str | None = None + url: str | None = None + auth_type: str | None = None + oauth2_flow: str | None = None + allow_all_keys: bool | None = None + + class EmbedBody(BaseModel): model: str input: str diff --git a/tests/e2e/quota_management/ratelimit/test_redis_backed_ratelimit_e2e.py b/tests/e2e/quota_management/ratelimit/test_redis_backed_ratelimit_e2e.py index ed6f0ce3b2c..a88f0ca546a 100644 --- a/tests/e2e/quota_management/ratelimit/test_redis_backed_ratelimit_e2e.py +++ b/tests/e2e/quota_management/ratelimit/test_redis_backed_ratelimit_e2e.py @@ -11,7 +11,7 @@ import socket import pytest -from e2e_config import require_env, unique_marker +from e2e_config import unique_marker from e2e_http import require_successful_call from lifecycle import ResourceManager from models import KeyGenerateBody, LiteLLMParamsBody @@ -23,7 +23,7 @@ BACKEND = "anthropic/claude-haiku-4-5-20251001" def _require_redis_reachable() -> None: - (host,) = require_env("REDIS_HOST") + host = os.environ["REDIS_HOST"] port = int((os.environ.get("REDIS_PORT") or "6379").strip() or "6379") try: with socket.create_connection((host, port), timeout=3): diff --git a/tests/e2e/quota_management/ratelimit/test_redis_circuit_breaker_e2e.py b/tests/e2e/quota_management/ratelimit/test_redis_circuit_breaker_e2e.py index b509ae000f5..3e1bc662470 100644 --- a/tests/e2e/quota_management/ratelimit/test_redis_circuit_breaker_e2e.py +++ b/tests/e2e/quota_management/ratelimit/test_redis_circuit_breaker_e2e.py @@ -13,7 +13,7 @@ from concurrent.futures import ThreadPoolExecutor, as_completed import pytest -from e2e_config import require_env, unique_marker +from e2e_config import unique_marker from e2e_http import require_successful_call from lifecycle import ResourceManager from models import KeyGenerateBody, LiteLLMParamsBody @@ -28,7 +28,7 @@ RECOVERY_TIMEOUT = float( def _require_redis() -> None: - (host,) = require_env("REDIS_HOST") + host = os.environ["REDIS_HOST"] port = int((os.environ.get("REDIS_PORT") or "6379").strip() or "6379") try: with socket.create_connection((host, port), timeout=3): diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index 890f625f90f..ab114f93cb0 100644 --- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -2271,6 +2271,37 @@ def test_token_type_cost_breakdown_reads_cache_write_tokens(): ) +def test_generic_cost_per_token_openai_cache_write_tokens_gpt_5_6(): + """ + Regression: OpenAI gpt-5.6 reports cache-write tokens under + prompt_tokens_details.cache_write_tokens (not the Anthropic cache_creation_tokens + name). Those tokens must be billed at the cache-write rate rather than the plain + input rate. Customer report: cache creation tokens were never counted for the + GPT-5.6 series, so cost was undercounted on cache-write requests. + """ + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + + model = "gpt-5.6" + usage = Usage( + prompt_tokens=1000, + completion_tokens=10, + total_tokens=1010, + prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=0, cache_write_tokens=800), + ) + + assert usage.prompt_tokens_details.cache_write_tokens == 800 + assert usage.prompt_tokens_details.cache_creation_tokens == 800 + + prompt_cost, _ = generic_cost_per_token(model=model, usage=usage, custom_llm_provider="openai") + + info = litellm.get_model_info(model=model, custom_llm_provider="openai") + expected_prompt = (1000 - 800) * info["input_cost_per_token"] + 800 * info["cache_creation_input_token_cost"] + assert prompt_cost == pytest.approx(expected_prompt) + assert info["cache_creation_input_token_cost"] > info["input_cost_per_token"] + assert prompt_cost > 1000 * info["input_cost_per_token"] + + def test_token_type_cost_breakdown_reconciles_with_generic_total(): """ Both-ways check: the reasoning subset must sum with the remaining (text) output @@ -2326,6 +2357,65 @@ def test_token_type_cost_breakdown_zero_without_special_tokens(): ) +@pytest.mark.parametrize( + "raw_usage, expect_read, expect_write", + [ + ( + { + "input_tokens": 5000, + "output_tokens": 10, + "total_tokens": 5010, + "input_tokens_details": {"cached_tokens": 0, "cache_write_tokens": 4012}, + }, + False, + True, + ), + ( + { + "input_tokens": 5000, + "output_tokens": 10, + "total_tokens": 5010, + "input_tokens_details": {"cached_tokens": 4012, "cache_write_tokens": 0}, + }, + True, + False, + ), + ], +) +def test_token_type_cost_breakdown_openai_responses_api_cache_write_read( + raw_usage, expect_read, expect_write +): + """Regression for #34309: OpenAI Responses API reports cache tokens under + input_tokens_details.{cached_tokens, cache_write_tokens}, not the Anthropic-style + top-level cache_creation_input_tokens. The itemized breakdown must still populate + cache_read_cost / cache_creation_cost from the transformed usage.""" + from litellm.responses.utils import ResponseAPILoggingUtils + + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + + model = "gpt-5.6" + usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(raw_usage) + + breakdown = get_token_type_cost_breakdown( + model=model, custom_llm_provider="openai", usage=usage + ) + + info = litellm.get_model_info(model=model, custom_llm_provider="openai") + if expect_write: + assert breakdown.cache_creation_cost == pytest.approx( + 4012 * info["cache_creation_input_token_cost"] + ) + assert breakdown.cache_creation_cost > 0 + assert breakdown.cache_read_cost == 0.0 + if expect_read: + assert breakdown.cache_read_cost == pytest.approx( + 4012 * info["cache_read_input_token_cost"] + ) + assert breakdown.cache_read_cost > 0 + assert breakdown.cache_creation_cost == 0.0 + + def test_token_type_cost_breakdown_handles_unknown_model_gracefully(): """A model with no pricing must yield zeros, never raise.""" breakdown = get_token_type_cost_breakdown( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_token_exchange.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_token_exchange.py deleted file mode 100644 index d2aa58e29ea..00000000000 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_token_exchange.py +++ /dev/null @@ -1,539 +0,0 @@ -""" -Tests for OAuth 2.0 Token Exchange (RFC 8693) handler for MCP servers. - -Covers: exchange flow, caching, error handling, resolve_mcp_auth integration, -bearer token extraction, and config loading. -""" - -from unittest.mock import AsyncMock, MagicMock, patch - -import httpx -import pytest - -from litellm.proxy._experimental.mcp_server.auth.token_exchange import ( - TOKEN_EXCHANGE_GRANT_TYPE, - TokenExchangeHandler, -) -from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - MCPServerManager, -) -from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( - resolve_mcp_auth, -) -from litellm.proxy._types import LiteLLM_MCPServerTable, MCPTransport -from litellm.types.mcp import MCPAuth -from litellm.types.mcp_server.mcp_server_manager import MCPServer - - -def _obo_server(**overrides) -> MCPServer: - defaults = dict( - server_id="srv-obo-1", - name="test-obo", - url="https://mcp.example.com/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2_token_exchange, - client_id="litellm-client-id", - client_secret="litellm-client-secret", - token_exchange_endpoint="https://idp.example.com/oauth2/token", - audience="api://mcp-server", - scopes=["mcp.tools.read", "mcp.tools.execute"], - ) - defaults.update(overrides) - return MCPServer(**defaults) - - -def _exchange_response(token="exchanged-tok-abc", expires_in=3600): - resp = MagicMock() - resp.json.return_value = { - "access_token": token, - "token_type": "Bearer", - "expires_in": expires_in, - } - resp.raise_for_status = MagicMock() - resp.text = "" - return resp - - -# ── Exchange Flow ── - - -@pytest.mark.asyncio -async def test_exchange_token_success(): - """Token exchange sends correct RFC 8693 parameters and returns access_token.""" - handler = TokenExchangeHandler() - server = _obo_server() - mock_client = AsyncMock() - mock_client.post.return_value = _exchange_response("scoped-token-1") - - with patch( - "litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client", - return_value=mock_client, - ): - result = await handler.exchange_token("user-jwt-xyz", server) - - assert result == "scoped-token-1" - mock_client.post.assert_called_once() - - _, kwargs = mock_client.post.call_args - data = kwargs["data"] - assert data["grant_type"] == TOKEN_EXCHANGE_GRANT_TYPE - assert data["subject_token"] == "user-jwt-xyz" - assert data["subject_token_type"] == "urn:ietf:params:oauth:token-type:access_token" - assert data["audience"] == "api://mcp-server" - assert data["scope"] == "mcp.tools.read mcp.tools.execute" - assert data["client_id"] == "litellm-client-id" - assert data["client_secret"] == "litellm-client-secret" - - -@pytest.mark.asyncio -async def test_exchange_token_no_audience(): - """When audience is None, it is omitted from the request.""" - handler = TokenExchangeHandler() - server = _obo_server(audience=None) - mock_client = AsyncMock() - mock_client.post.return_value = _exchange_response() - - with patch( - "litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client", - return_value=mock_client, - ): - await handler.exchange_token("user-jwt", server) - - _, kwargs = mock_client.post.call_args - assert "audience" not in kwargs["data"] - - -@pytest.mark.asyncio -async def test_exchange_token_no_scopes(): - """When scopes is None, scope param is omitted from the request.""" - handler = TokenExchangeHandler() - server = _obo_server(scopes=None) - mock_client = AsyncMock() - mock_client.post.return_value = _exchange_response() - - with patch( - "litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client", - return_value=mock_client, - ): - await handler.exchange_token("user-jwt", server) - - _, kwargs = mock_client.post.call_args - assert "scope" not in kwargs["data"] - - -# ── Caching ── - - -@pytest.mark.asyncio -async def test_exchange_token_cached(): - """Second call with same user token uses cache — only 1 HTTP POST.""" - handler = TokenExchangeHandler() - server = _obo_server() - mock_client = AsyncMock() - mock_client.post.return_value = _exchange_response("cached-exchange-tok") - - with patch( - "litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client", - return_value=mock_client, - ): - t1 = await handler.exchange_token("same-jwt", server) - t2 = await handler.exchange_token("same-jwt", server) - - assert t1 == t2 == "cached-exchange-tok" - assert mock_client.post.call_count == 1 - - -@pytest.mark.asyncio -async def test_different_user_tokens_not_shared(): - """Different user JWTs get different exchanged tokens.""" - handler = TokenExchangeHandler() - server = _obo_server() - call_count = 0 - - async def mock_post(url, data=None): - nonlocal call_count - call_count += 1 - resp = MagicMock() - resp.json.return_value = { - "access_token": f"exchanged-{call_count}", - "expires_in": 3600, - } - resp.raise_for_status = MagicMock() - return resp - - mock_client = AsyncMock() - mock_client.post = mock_post - - with patch( - "litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client", - return_value=mock_client, - ): - t1 = await handler.exchange_token("user-a-jwt", server) - t2 = await handler.exchange_token("user-b-jwt", server) - - assert t1 == "exchanged-1" - assert t2 == "exchanged-2" - assert call_count == 2 - - -# ── Error Handling ── - - -@pytest.mark.asyncio -async def test_exchange_token_http_error(): - """HTTP errors from the IDP are wrapped in a ValueError.""" - handler = TokenExchangeHandler() - server = _obo_server() - mock_response = MagicMock() - mock_response.status_code = 400 - mock_response.text = "invalid_grant" - mock_response.raise_for_status.side_effect = httpx.HTTPStatusError( - "Bad Request", - request=MagicMock(), - response=mock_response, - ) - mock_client = AsyncMock() - mock_client.post.return_value = mock_response - - with ( - patch( - "litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client", - return_value=mock_client, - ), - pytest.raises(ValueError, match="failed with status 400"), - ): - await handler.exchange_token("bad-jwt", server) - - -@pytest.mark.asyncio -async def test_exchange_token_http_error_does_not_log_response_body(): - """Raw IDP error bodies are not logged because they can contain credentials.""" - handler = TokenExchangeHandler() - server = _obo_server() - raw_response_body = "client_secret=do-not-log" - mock_response = MagicMock() - mock_response.status_code = 401 - mock_response.text = raw_response_body - mock_response.raise_for_status.side_effect = httpx.HTTPStatusError( - "Unauthorized", - request=MagicMock(), - response=mock_response, - ) - mock_client = AsyncMock() - mock_client.post.return_value = mock_response - - with ( - patch( - "litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client", - return_value=mock_client, - ), - patch( - "litellm.proxy._experimental.mcp_server.auth.token_exchange.verbose_logger.debug" - ) as mock_debug, - pytest.raises(ValueError, match="failed with status 401"), - ): - await handler.exchange_token("bad-jwt", server) - - logged_values = " ".join( - str(value) - for call in mock_debug.call_args_list - for value in [*call.args, *call.kwargs.values()] - ) - assert raw_response_body not in logged_values - - -@pytest.mark.asyncio -async def test_exchange_token_missing_access_token(): - """Response without access_token raises ValueError.""" - handler = TokenExchangeHandler() - server = _obo_server() - resp = MagicMock() - resp.json.return_value = {"token_type": "Bearer"} - resp.raise_for_status = MagicMock() - mock_client = AsyncMock() - mock_client.post.return_value = resp - - with ( - patch( - "litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client", - return_value=mock_client, - ), - pytest.raises(ValueError, match="missing 'access_token'"), - ): - await handler.exchange_token("jwt", server) - - -@pytest.mark.asyncio -async def test_exchange_token_missing_endpoint(): - """Missing token_exchange_endpoint and token_url raises ValueError.""" - handler = TokenExchangeHandler() - server = _obo_server(token_exchange_endpoint=None, token_url=None) - - with pytest.raises(ValueError, match="no token_exchange_endpoint or token_url"): - await handler.exchange_token("jwt", server) - - -@pytest.mark.asyncio -async def test_exchange_token_missing_credentials(): - """Missing client_id or client_secret raises ValueError.""" - handler = TokenExchangeHandler() - server = _obo_server(client_id=None, client_secret=None) - # has_token_exchange_config will be False, so we call _do_exchange directly - with pytest.raises(ValueError, match="missing client_id or client_secret"): - await handler._do_exchange("jwt", server) - - -# ── resolve_mcp_auth Integration ── - - -@pytest.mark.asyncio -async def test_resolve_mcp_auth_with_token_exchange(): - """resolve_mcp_auth delegates to token exchange when server has OBO config and subject_token provided.""" - server = _obo_server() - mock_handler = AsyncMock() - mock_handler.exchange_token.return_value = "obo-scoped-token" - - with patch( - "litellm.proxy._experimental.mcp_server.auth.token_exchange.mcp_token_exchange_handler", - mock_handler, - ): - result = await resolve_mcp_auth(server, subject_token="user-jwt") - - assert result == "obo-scoped-token" - mock_handler.exchange_token.assert_called_once_with("user-jwt", server) - - -@pytest.mark.asyncio -async def test_resolve_mcp_auth_obo_without_subject_token_falls_through(): - """Without a subject_token, resolve_mcp_auth falls through to client_credentials.""" - server = _obo_server( - token_url="https://auth.example.com/token", - ) - mock_client = AsyncMock() - mock_client.post.return_value = _exchange_response("cc-token") - - with patch( - "litellm.proxy._experimental.mcp_server.oauth2_token_cache.get_async_httpx_client", - return_value=mock_client, - ): - result = await resolve_mcp_auth(server, subject_token=None) - - # Falls through to client_credentials since subject_token is None - # The server has client_id/client_secret/token_url so has_client_credentials is True - assert result == "cc-token" - - -@pytest.mark.asyncio -async def test_resolve_mcp_auth_obo_without_subject_token_uses_cached_client_credentials(): - """The M2M fallback for OBO servers reuses the client_credentials cache.""" - server = _obo_server( - server_id="srv-obo-m2m-cache", - token_url="https://auth.example.com/token", - ) - mock_client = AsyncMock() - mock_client.post.return_value = _exchange_response("cached-cc-token") - - with patch( - "litellm.proxy._experimental.mcp_server.oauth2_token_cache.get_async_httpx_client", - return_value=mock_client, - ): - first = await resolve_mcp_auth(server, subject_token=None) - second = await resolve_mcp_auth(server, subject_token=None) - - assert first == second == "cached-cc-token" - mock_client.post.assert_called_once() - - -@pytest.mark.asyncio -async def test_resolve_mcp_auth_header_beats_obo(): - """An explicit mcp_auth_header takes priority over OBO token exchange.""" - server = _obo_server() - result = await resolve_mcp_auth( - server, mcp_auth_header="Bearer override", subject_token="user-jwt" - ) - assert result == "Bearer override" - - -# ── Bearer Token Extraction ── - - -def test_extract_bearer_token_from_oauth2_headers(): - """Extracts token from oauth2_headers Authorization header.""" - result = MCPServerManager._extract_bearer_token( - oauth2_headers={"Authorization": "Bearer my-jwt-token"}, - raw_headers=None, - ) - assert result == "my-jwt-token" - - -def test_extract_bearer_token_from_raw_headers(): - """Falls back to raw_headers when oauth2_headers missing.""" - result = MCPServerManager._extract_bearer_token( - oauth2_headers=None, - raw_headers={"authorization": "Bearer raw-jwt"}, - ) - assert result == "raw-jwt" - - -def test_extract_bearer_token_no_bearer_prefix(): - """Returns token as-is when no Bearer prefix.""" - result = MCPServerManager._extract_bearer_token( - oauth2_headers={"Authorization": "some-opaque-token"}, - raw_headers=None, - ) - assert result == "some-opaque-token" - - -def test_extract_bearer_token_none(): - """Returns None when no auth headers present.""" - result = MCPServerManager._extract_bearer_token( - oauth2_headers=None, - raw_headers=None, - ) - assert result is None - - -# ── MCPServer Properties ── - - -def test_has_token_exchange_config_true(): - """has_token_exchange_config is True for a fully configured OBO server.""" - server = _obo_server() - assert server.has_token_exchange_config is True - - -def test_has_token_exchange_config_false_wrong_auth_type(): - """has_token_exchange_config is False when auth_type is not oauth2_token_exchange.""" - server = _obo_server(auth_type=MCPAuth.oauth2) - assert server.has_token_exchange_config is False - - -def test_has_token_exchange_config_false_missing_creds(): - """has_token_exchange_config is False when client_id/client_secret missing.""" - server = _obo_server(client_id=None) - assert server.has_token_exchange_config is False - - -def test_has_token_exchange_config_uses_token_url_fallback(): - """has_token_exchange_config is True when token_url is set instead of token_exchange_endpoint.""" - server = _obo_server( - token_exchange_endpoint=None, - token_url="https://idp.example.com/token", - ) - assert server.has_token_exchange_config is True - - -# ── Config Loading ── - - -@pytest.mark.asyncio -async def test_config_loading_token_exchange_fields(): - """load_servers_from_config correctly maps OBO config fields to MCPServer.""" - manager = MCPServerManager() - config = { - "my_obo_server": { - "url": "https://mcp.example.com/mcp", - "transport": "http", - "auth_type": "oauth2_token_exchange", - "client_id": "my-client", - "client_secret": "my-secret", - "token_exchange_endpoint": "https://idp.example.com/oauth2/token", - "audience": "api://my-mcp", - "scopes": ["read", "write"], - "subject_token_type": "urn:ietf:params:oauth:token-type:jwt", - } - } - await manager.load_servers_from_config(config) - - servers = list(manager.config_mcp_servers.values()) - assert len(servers) == 1 - - server = servers[0] - assert server.auth_type == MCPAuth.oauth2_token_exchange - assert server.token_exchange_endpoint == "https://idp.example.com/oauth2/token" - assert server.audience == "api://my-mcp" - assert server.subject_token_type == "urn:ietf:params:oauth:token-type:jwt" - assert server.client_id == "my-client" - assert server.client_secret == "my-secret" - assert server.scopes == ["read", "write"] - assert server.has_token_exchange_config is True - - -@pytest.mark.asyncio -async def test_config_loading_default_subject_token_type(): - """subject_token_type defaults to access_token when not specified in config.""" - manager = MCPServerManager() - config = { - "obo_defaults": { - "url": "https://mcp.example.com/mcp", - "transport": "http", - "auth_type": "oauth2_token_exchange", - "client_id": "cid", - "client_secret": "csec", - "token_exchange_endpoint": "https://idp.example.com/token", - } - } - await manager.load_servers_from_config(config) - - server = list(manager.config_mcp_servers.values())[0] - assert server.subject_token_type == "urn:ietf:params:oauth:token-type:access_token" - - -@pytest.mark.asyncio -async def test_database_loading_token_exchange_scopes_from_credentials(): - """DB-loaded OBO server credentials retain configured scopes.""" - manager = MCPServerManager() - db_server = LiteLLM_MCPServerTable( - server_id="srv-obo-db", - server_name="obo_db_server", - url="https://mcp.example.com/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2_token_exchange, - credentials={ - "client_id": "db-client", - "client_secret": "db-secret", - "token_exchange_endpoint": "https://idp.example.com/oauth2/token", - "audience": "api://db-mcp", - "scopes": ["db.read", "db.write"], - }, - ) - - server = await manager.build_mcp_server_from_table( - db_server, - credentials_are_encrypted=False, - ) - - assert server.auth_type == MCPAuth.oauth2_token_exchange - assert server.client_id == "db-client" - assert server.client_secret == "db-secret" - assert server.token_exchange_endpoint == "https://idp.example.com/oauth2/token" - assert server.audience == "api://db-mcp" - assert server.scopes == ["db.read", "db.write"] - - -@pytest.mark.asyncio -async def test_exchange_token_uses_client_secret_basic_when_configured(): - """LIT-4091: token exchange with token_endpoint_auth_method=client_secret_basic sends the - client credentials as HTTP Basic and omits client_secret from the body.""" - import base64 - - handler = TokenExchangeHandler() - server = _obo_server( - server_id="srv-obo-basic", token_endpoint_auth_method="client_secret_basic" - ) - mock_client = AsyncMock() - mock_client.post.return_value = _exchange_response("scoped-basic") - - with patch( - "litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client", - return_value=mock_client, - ): - result = await handler.exchange_token("user-jwt-basic", server) - - assert result == "scoped-basic" - _, kwargs = mock_client.post.call_args - expected = "Basic " + base64.b64encode(b"litellm-client-id:litellm-client-secret").decode() - assert kwargs["headers"]["Authorization"] == expected - assert "client_secret" not in kwargs["data"] - assert "client_id" not in kwargs["data"] - assert kwargs["data"]["grant_type"] == TOKEN_EXCHANGE_GRANT_TYPE 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 7b05b8c9dd0..b3c0dcd1681 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 @@ -15,6 +15,7 @@ from starlette.datastructures import Headers from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, + _is_mcp_admitted_user_subject, ) from litellm.proxy._types import ( LiteLLM_ObjectPermissionTable, @@ -206,6 +207,31 @@ class TestMCPRequestHandler: assert sorted(result) == sorted(expected) + async def test_admitted_subject_not_zeroed_by_require_key_mcp_access_defined(self): + """10x-flow regression: with require_key_mcp_access_defined ON (team = ceiling for keys), a + keyless gateway/bridge-admitted subject whose ONLY access path is team membership must still + inherit the team's servers. The flag zeros empty *virtual keys* that must declare their own + access; a keyless admitted user has no key to declare it on, so it must not be zeroed.""" + auth = UserAPIKeyAuth(api_key=None, user_id="sso-user") + auth.mcp_admitted_user_subject = True + with ( + patch.object( + MCPRequestHandler, "_get_allowed_mcp_servers_for_key", new_callable=AsyncMock, return_value=[] + ), + patch.object( + MCPRequestHandler, + "_get_allowed_mcp_servers_for_team", + new_callable=AsyncMock, + return_value=["team_server1", "team_server2"], + ), + patch.object( + MCPRequestHandler, "_get_key_access_group_mcp_server_extras", new_callable=AsyncMock, return_value=[] + ), + patch("litellm.proxy.proxy_server.general_settings", {"require_key_mcp_access_defined": True}), + ): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert sorted(result) == ["team_server1", "team_server2"] + @pytest.mark.parametrize( "key_servers,grants,expected,scenario", [ @@ -5302,6 +5328,7 @@ class TestMCPDcrBridgeDelegateAdmission: self._patch_user_reload( return_value=MagicMock( user_id="sso-user-7", + organization_id=None, metadata={"scim_active": True}, user_role=None, object_permission=None, @@ -5344,6 +5371,7 @@ class TestMCPDcrBridgeDelegateAdmission: self._patch_user_reload( return_value=MagicMock( user_id="sso-user-7", + organization_id=None, metadata={"scim_active": True}, user_role=None, object_permission=object_permission, @@ -5421,7 +5449,9 @@ class TestMCPDcrBridgeDelegateAdmission: with ( patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), - self._patch_user_reload(return_value=MagicMock(user_id="offboarded-user", metadata={"scim_active": False})), + self._patch_user_reload( + return_value=MagicMock(user_id="offboarded-user", organization_id=None, metadata={"scim_active": False}) + ), ): mock_mgr.get_mcp_server_by_name.return_value = self._bridge_delegate_server() with pytest.raises(HTTPException) as exc_info: @@ -6204,7 +6234,9 @@ class TestAggregateGatewayDcrChallenge: with pytest.raises(HTTPException) as exc_info: await MCPRequestHandler.process_mcp_request(self._scope()) www_authenticate = (exc_info.value.headers or {})["WWW-Authenticate"] - assert 'resource_metadata="http://testserver/.well-known/oauth-protected-resource/litellm/mcp"' in www_authenticate + assert ( + 'resource_metadata="http://testserver/.well-known/oauth-protected-resource/litellm/mcp"' in www_authenticate + ) async def test_no_challenge_for_explicit_litellm_key(self): """An explicit x-litellm-api-key declares a litellm-key client; a typo @@ -6225,9 +6257,7 @@ class TestAggregateGatewayDcrChallenge: patch(self._AUTH_PATCH_TARGET, side_effect=self._auth_401()), ): with pytest.raises(ProxyException): - await MCPRequestHandler.process_mcp_request( - self._scope(extra_headers=((b"x-mcp-servers", b"github"),)) - ) + await MCPRequestHandler.process_mcp_request(self._scope(extra_headers=((b"x-mcp-servers", b"github"),))) async def test_no_challenge_for_path_named_server(self): """/mcp/{server} targets one server; the aggregate challenge must not @@ -6261,3 +6291,1357 @@ class TestAggregateGatewayDcrChallenge: with pytest.raises(ProxyException) as exc_info: await MCPRequestHandler.process_mcp_request(self._scope()) assert str(exc_info.value.code) == "500" + + +@pytest.mark.asyncio +class TestGatewaySessionAdmission: + """The aggregate /mcp session-bearer admission arm (mcp_gateway_dcr). A valid session + token admits under the LIVE litellm user it references; an invalid/expired/refresh/foreign + token fails closed with the aggregate invalid_token challenge; the arm fires ONLY at the + aggregate scope, never for named servers or per-server flows.""" + + _MASTER_KEY = "sk-gateway-session-admission-master-key" + + def _session_bearer(self, user_id="sso-user-42", client_id="llm_dcrc_abc"): + from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import ( + session_keys_from_master_key, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import ( + SessionPrincipal, + mint_session_token, + mint_session_refresh_token, + ) + + keys = session_keys_from_master_key(self._MASTER_KEY) + principal = SessionPrincipal(user_id=user_id, client_id=client_id) + return mint_session_token, mint_session_refresh_token, principal, keys + + def _access_token(self, **kw): + from datetime import datetime, timezone + + mint, _refresh, principal, keys = self._session_bearer(**kw) + return mint(principal, keys, datetime(2030, 1, 1, tzinfo=timezone.utc)).token.get_secret_value() + + def _scope(self, bearer, path="/mcp", extra_headers=()): + return { + "type": "http", + "method": "POST", + "path": path, + "headers": [(b"host", b"testserver"), (b"authorization", f"Bearer {bearer}".encode()), *extra_headers], + } + + @staticmethod + @contextlib.contextmanager + def _patch_user_reload(*, user_id, active=True, organization_id=None, tpm_limit=None, rpm_limit=None): + get_user_object = AsyncMock( + return_value=MagicMock( + user_id=user_id, + organization_id=organization_id, + metadata={"scim_active": active} if not active else {"scim_active": True}, + user_role=None, + object_permission=None, + object_permission_id=None, + tpm_limit=tpm_limit, + rpm_limit=rpm_limit, + ) + ) + with ( + patch("litellm.proxy.auth.auth_checks.get_user_object", get_user_object), + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + ): + yield get_user_object + + async def test_session_admission_binds_org_id_so_the_org_ceiling_applies(self): + """The admitted auth carries the user's org_id, so get_allowed_mcp_servers keeps the + org-level MCP ceiling in force for a gateway session instead of skipping it.""" + token = self._access_token(user_id="org-user") + with ( + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + new_callable=AsyncMock, + ), + self._patch_user_reload(user_id="org-user", organization_id="org-123"), + ): + auth_result, *_rest = await MCPRequestHandler.process_mcp_request(self._scope(token)) + assert auth_result.org_id == "org-123" + + async def test_session_admission_copies_user_rate_limits(self): + """Security regression: the reconstructed auth must carry the live user's RPM/TPM, exactly as + the standard user-subject path does. The parallel limiter reads these off the auth object and + treats None as unlimited, so a keyless subject with them unset would invoke tools past their + configured user rate limits.""" + token = self._access_token(user_id="rl-user") + with ( + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + new_callable=AsyncMock, + ), + self._patch_user_reload(user_id="rl-user", tpm_limit=1000, rpm_limit=50), + ): + auth_result, *_rest = await MCPRequestHandler.process_mcp_request(self._scope(token)) + assert auth_result.user_tpm_limit == 1000 + assert auth_result.user_rpm_limit == 50 + + async def test_valid_session_admits_under_live_user_at_aggregate_scope(self): + token = self._access_token(user_id="sso-user-42") + with ( + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + new_callable=AsyncMock, + ) as mock_auth, + self._patch_user_reload(user_id="sso-user-42") as get_user_object, + ): + auth_result, _h, _servers, mcp_server_auth_headers, _o, _r = await MCPRequestHandler.process_mcp_request( + self._scope(token) + ) + assert get_user_object.await_args.kwargs["user_id"] == "sso-user-42" + assert auth_result.user_id == "sso-user-42" + mock_auth.assert_not_called() + # Identity-only admission injects no per-server upstream credential (unlike the + # bridge envelope arm); the headers dict is whatever the request carried, here empty. + assert not mcp_server_auth_headers + + @pytest.mark.parametrize( + "scenario, expect_challenge", + [("expired", True), ("tampered", False), ("refresh_at_tool_edge", False), ("foreign_key", False)], + ) + async def test_bad_session_bearer_fails_closed(self, scenario, expect_challenge): + # Every non-admissible session-shaped bearer fails closed with 401; a valid-but-unusable one + # (expired) additionally carries the invalid_token challenge so the DCR client re-authorizes. + from datetime import datetime, timezone + + if scenario == "expired": + mint, _refresh, principal, keys = self._session_bearer() + bearer = mint(principal, keys, datetime(2020, 1, 1, tzinfo=timezone.utc)).token.get_secret_value() + elif scenario == "tampered": + token = self._access_token() + bearer = token[:-3] + ("aaa" if not token.endswith("aaa") else "bbb") + elif scenario == "refresh_at_tool_edge": + _mint, refresh, principal, keys = self._session_bearer() + bearer = refresh(principal, keys, datetime(2030, 1, 1, tzinfo=timezone.utc)).token.get_secret_value() + else: # foreign_key: minted under the real master key, presented while the proxy uses another + bearer = self._access_token() + master_key = "sk-a-totally-different-master-key" if scenario == "foreign_key" else self._MASTER_KEY + with patch("litellm.proxy.proxy_server.master_key", master_key): + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(self._scope(bearer)) + assert exc_info.value.status_code == 401 + if expect_challenge: + assert 'error="invalid_token"' in (exc_info.value.headers or {})["WWW-Authenticate"] + + async def test_deactivated_user_fails_with_invalid_token_challenge(self): + """A cryptographically valid bearer whose referenced user is SCIM-deactivated must fail with + the aggregate invalid_token challenge (WWW-Authenticate), matching the expired/tampered arms, + so the DCR client re-authorizes instead of getting a bare 401 with no challenge.""" + token = self._access_token(user_id="offboarded-user") + with ( + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + new_callable=AsyncMock, + ), + self._patch_user_reload(user_id="offboarded-user", active=False), + ): + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(self._scope(token)) + assert exc_info.value.status_code == 401 + assert 'error="invalid_token"' in (exc_info.value.headers or {})["WWW-Authenticate"] + + async def test_session_bearer_scrubbed_from_egress_header_contexts(self): + """Security regression (credential leak): after a keyless session admission, the session + bearer must be removed from BOTH returned egress header contexts (oauth2_headers and the raw + headers) so no passthrough/OBO egress can forward it upstream for replay as this user.""" + token = self._access_token(user_id="sso-user-42") + with ( + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + new_callable=AsyncMock, + ), + self._patch_user_reload(user_id="sso-user-42"), + ): + _auth, _h, _servers, _msah, oauth2_headers, raw_headers = await MCPRequestHandler.process_mcp_request( + self._scope(token) + ) + # the request carried "Authorization: Bearer "; both egress contexts must be scrubbed + assert oauth2_headers is None + assert not any(k.lower() == "authorization" for k in (raw_headers or {})) + + async def test_arm_does_not_fire_for_named_server(self): + """A session-shaped bearer aimed at a named server (path scope) does not enter the + aggregate arm; it is treated as an ordinary bearer on that server.""" + token = self._access_token() + with ( + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + new_callable=AsyncMock, + side_effect=ProxyException(message="bad key", type="auth_error", param="api_key", code=401), + ) as mock_auth, + ): + with pytest.raises((HTTPException, ProxyException)): + await MCPRequestHandler.process_mcp_request(self._scope(token, path="/mcp/github")) + mock_auth.assert_called_once() + + +def _make_team(team_id, mcp_servers, *, org_id=None, tool_perms=None, members=("sso-user",)): + from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable, Member + + return LiteLLM_TeamTable( + team_id=team_id, + organization_id=org_id, + members_with_roles=[Member(user_id=u, role="user") for u in members], + access_group_ids=[], + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id=f"op-{team_id}", mcp_servers=mcp_servers, mcp_tool_permissions=tool_perms + ), + ) + + +def _make_admitted_subject(user_id, *, org_id=None, own_servers=None, own_tool_perms=None): + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + op = None + if own_servers is not None or own_tool_perms is not None: + op = LiteLLM_ObjectPermissionTable( + object_permission_id=f"userop-{user_id}", + mcp_servers=own_servers or [], + mcp_tool_permissions=own_tool_perms, + ) + auth = UserAPIKeyAuth(user_id=user_id, api_key=None, org_id=org_id, object_permission=op) + auth.mcp_admitted_user_subject = True + return auth + + +@pytest.mark.asyncio +class TestUserSubjectTeamUnion: + """_get_allowed_mcp_servers_for_team unions across ALL a user's teams for a keyless + user-subject caller (the gateway DCR session bearer and bridge user-envelope), while a + key-based caller keeps its single-team behavior byte-identically.""" + + @contextlib.contextmanager + def _patch(self, *, teams_by_id, user_teams=None, orgs_by_id=None): + async def _get_team_object(team_id, **kw): + return teams_by_id.get(team_id) + + async def _get_user_object(user_id, **kw): + return MagicMock(user_id=user_id, teams=user_teams or []) + + async def _get_org_object(org_id, **kw): + return (orgs_by_id or {}).get(org_id) + + async def _spend_from_fallback(counter_key, fallback_spend, max_budget=None, **kw): + # The budget owners read cross-pod spend Redis-first with the row's spend as fallback; + # unit tests have no Redis, so the fallback IS the spend. + return fallback_spend + + with ( + patch("litellm.proxy.auth.auth_checks.get_team_object", _get_team_object), + patch("litellm.proxy.auth.auth_checks.get_user_object", _get_user_object), + patch("litellm.proxy.auth.auth_checks.get_org_object", _get_org_object), + patch("litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", AsyncMock(return_value=[])), + patch("litellm.proxy.proxy_server.get_current_spend", _spend_from_fallback), + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + ): + yield + + async def test_keyless_user_unions_servers_across_all_their_teams(self): + teams = {"team-a": _make_team("team-a", ["srv1", "srv2"]), "team-b": _make_team("team-b", ["srv2", "srv3"])} + auth = _make_admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"]): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert set(result) == {"srv1", "srv2", "srv3"} + + async def test_key_based_caller_uses_single_team_only(self): + """A key-based caller (api_key set) with a team_id sees ONLY that team, even though the + same user belongs to other teams: key auth must be byte-identical to before.""" + teams = {"team-a": _make_team("team-a", ["srv1"]), "team-b": _make_team("team-b", ["srv2", "srv3"])} + auth = UserAPIKeyAuth(user_id="sso-user", api_key="sk-hash", team_id="team-a") + with self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"]): + result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(auth) + assert set(result) == {"srv1"} + + async def test_keyless_user_with_explicit_team_id_uses_that_team_only(self): + """A keyless caller that already pins a team_id (not the user-subject fan-out shape) + resolves only that team; the union is strictly for the no-team-id user-subject case.""" + teams = {"team-a": _make_team("team-a", ["srv1"]), "team-b": _make_team("team-b", ["srv2"])} + auth = UserAPIKeyAuth(user_id="sso-user", api_key=None, team_id="team-a") + with self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"]): + result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(auth) + assert set(result) == {"srv1"} + + async def test_keyless_user_with_no_teams_gets_nothing_from_teams(self): + auth = _make_admitted_subject("lonely-user") + with self._patch(teams_by_id={}, user_teams=[]): + result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(auth) + assert result == [] + + async def test_ui_session_team_id_still_resolves_to_nothing(self): + from litellm.proxy._types import UI_TEAM_ID + + auth = UserAPIKeyAuth(user_id="dash-user", api_key="sk-hash", team_id=UI_TEAM_ID) + with self._patch(teams_by_id={}, user_teams=["team-a"]): + result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(auth) + assert result == [] + + async def test_team_ids_helper_gates_on_shape(self): + from litellm.proxy._types import UI_TEAM_ID + + # key-based with team -> that team + assert await MCPRequestHandler._team_ids_for_mcp_grant( + UserAPIKeyAuth(api_key="sk", team_id="t1", user_id="u") + ) == ["t1"] + # An admitted subject never fans out HERE: it resolves one source per team first, and each of + # those pins a team_id, so this helper only ever answers the single-team question. The fan-out + # itself is _admitted_subject_sources' job, asserted below. + with self._patch(teams_by_id={}, user_teams=["t2", "t3"]): + assert await MCPRequestHandler._team_ids_for_mcp_grant(_make_admitted_subject("u")) == [] + # keyless, no user_id -> nothing + assert await MCPRequestHandler._team_ids_for_mcp_grant(UserAPIKeyAuth(api_key=None)) == [] + # keyless with a user_id but NOT admission-marked (JWT auth) -> nothing (unchanged behavior) + with self._patch(teams_by_id={}, user_teams=["t2", "t3"]): + assert ( + await MCPRequestHandler._team_ids_for_mcp_grant(UserAPIKeyAuth(api_key=None, user_id="jwt-user")) == [] + ) + # UI sentinel -> nothing + assert ( + await MCPRequestHandler._team_ids_for_mcp_grant( + UserAPIKeyAuth(api_key="sk", team_id=UI_TEAM_ID, user_id="u") + ) + == [] + ) + + async def test_org_outage_is_not_treated_as_a_missing_org(self): + """A CONFIRMED-absent org places no ceiling; a FAILED lookup must not be read as the same + fact. get_org_object used to relabel every error as "doesn't exist", so a DB outage silently + dropped a real org's ceiling for as long as it lasted. Absent -> the team's grant stands; + outage -> the keyless source denies.""" + from litellm.proxy.auth.auth_checks import OrganizationNotFoundError + + teams = {"t1": _make_team("t1", ["srv1"])} + teams["t1"].organization_id = "org-a" + auth = _make_admitted_subject("sso-user") + + absent = AsyncMock(side_effect=OrganizationNotFoundError("Organization doesn't exist in db.")) + with self._patch(teams_by_id=teams, user_teams=["t1"]): + with patch("litellm.proxy.auth.auth_checks.get_org_object", absent): + reachable = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert set(reachable) == {"srv1"}, "a deleted org places no ceiling" + + outage = AsyncMock(side_effect=RuntimeError("connection reset by peer")) + with self._patch(teams_by_id=teams, user_teams=["t1"]): + with patch("litellm.proxy.auth.auth_checks.get_org_object", outage): + reachable = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert reachable == [], "an unresolvable ceiling must deny a keyless source, not be skipped" + + async def test_org_ceiling_fault_fails_closed_for_admitted_but_open_for_keys(self): + """An unresolvable org ceiling is NOT the same fact as "this org places no restriction". + + For a virtual key the ceiling is one of several bounds and a DB blip must not lock working + keys out, so it stays fail-open. For a keyless admitted subject the per-source org ceiling is + the ONLY org bound, so dropping it on a fault would widen a cross-org user to servers their + team's org forbids. That is escalation, not an availability blip, so it fails closed.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + boom = AsyncMock(side_effect=RuntimeError("org lookup exploded")) + # The subject must actually REACH something, or the assertion passes either way and pins + # nothing (a fail-open mutant survived an earlier version of this test for exactly that). + auth = _make_admitted_subject("sso-user") + auth.org_id = "org-a" + auth.object_permission = LiteLLM_ObjectPermissionTable(object_permission_id="op-u", mcp_servers=["srv1"]) + with self._patch(teams_by_id={}, user_teams=[]): + assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"srv1"} # control + with patch.object(MCPRequestHandler, "_get_org_object_permission", boom): + admitted = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert admitted == [], "admitted subject must fail CLOSED when its org ceiling cannot resolve" + + key_auth = UserAPIKeyAuth(user_id="u", api_key="sk-hash", team_id="t1", org_id="org-a") + with self._patch(teams_by_id={"t1": _make_team("t1", ["srv1"])}, user_teams=[]): + with patch.object(MCPRequestHandler, "_get_org_object_permission", boom): + keyed = await MCPRequestHandler.get_allowed_mcp_servers(key_auth) + assert set(keyed) == {"srv1"}, "key auth must keep its long-standing fail-open behavior" + + async def test_only_the_attributing_team_bucket_is_charged(self): + """A team's mcp_rpm_limit bounds that team's SHARED bucket. Charging every granting team let + one cross-team user drain several teams' buckets on a single call, blocking their other + members for access those teams did not provide. Exactly one source is charged, and it is the + SAME source billing picks — one owner for both, so they cannot disagree.""" + from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + + t1 = _make_team("t1", ["srv1"]) + t1.metadata = {"mcp_rpm_limit": {"srv1": 5}} + t2 = _make_team("t2", ["srv1"]) + t2.metadata = {"mcp_rpm_limit": {"srv1": 9}} + auth = _make_admitted_subject("sso-user") + with self._patch(teams_by_id={"t1": t1, "t2": t2}, user_teams=["t1", "t2"]): + auth.mcp_source_team_rpm_limits = await MCPRequestHandler._admitted_subject_team_rpm_limits(auth) + billed = await MCPRequestHandler.attributing_source_for_server(auth, "srv1") + assert auth.mcp_source_team_rpm_limits == {"t1": {"srv1": 5}}, "t2's shared bucket is untouched" + assert billed is not None and billed.team_id == "t1", "throttling and billing pick the same source" + + descriptors: list = [] + limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=MagicMock()) + limiter._add_mcp_per_team_rate_limit_descriptor(auth, "srv1", descriptors) + charged = {d["value"]: d["rate_limit"]["requests_per_unit"] for d in descriptors} + assert charged == {"t1:srv1": 5}, "only the attributing team's bucket is charged" + + async def test_direct_user_grant_charges_no_team_bucket(self): + """When the user's OWN grant reaches the server, no team provided the access, so no team + bucket may be charged — the user's own rpm/tpm is what bounds them. Mirrors billing, which + bills the user and their own org for exactly this case.""" + t1 = _make_team("t1", ["srv1"]) + t1.metadata = {"mcp_rpm_limit": {"srv1": 5}} + auth = _make_admitted_subject("sso-user") + auth.object_permission = LiteLLM_ObjectPermissionTable(object_permission_id="op-u", mcp_servers=["srv1"]) + with self._patch(teams_by_id={"t1": t1}, user_teams=["t1"]): + limits = await MCPRequestHandler._admitted_subject_team_rpm_limits(auth) + billed = await MCPRequestHandler.attributing_source_for_server(auth, "srv1") + assert limits is None, "a direct user grant must not charge any team's shared bucket" + assert billed is None, "and billing agrees: the user is billed, not a team" + + def _manager_with(self, server_ids, allow_all=()): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.types.mcp import MCPTransport + + manager = MCPServerManager() + for sid in server_ids: + manager.registry[sid] = MCPServer( + server_id=sid, + name=sid, + server_name=sid, + url="https://example.com/mcp", + transport=MCPTransport.http, + allow_all_keys=sid in allow_all, + ) + manager._get_active_submitted_mcp_server_ids_for_user = AsyncMock(return_value=[]) + return manager + + async def test_team_derived_call_bills_the_granting_team_and_its_org(self): + """ACCOUNTING half of team budgets. Without attribution the admitted auth kept team_id=None, + so spend skipped team updates (the team's budget never accumulated, so it could never begin + to block) and charged the user's PRIMARY org rather than the org owning the granting team.""" + t_grant = _make_team("t-grant", ["srv1"]) + t_grant.organization_id = "org-team" + auth = _make_admitted_subject("sso-user") + auth.org_id = "org-user-primary" + with self._patch(teams_by_id={"t-grant": t_grant}, user_teams=["t-grant"]): + source = await MCPRequestHandler.attributing_source_for_server(auth, "srv1") + assert source is not None and source.team_id == "t-grant" + assert source.org_id == "org-team", "the granting team's org is charged, not the user's primary" + assert auth.team_id is None and auth.org_id == "org-user-primary", "authz object untouched" + + async def test_billing_auth_carries_team_and_org_onto_the_spend_object(self): + """Asserted on billing_auth_for_tool_call itself, not on the source it picks: the source + already carries the team's org by construction, so asserting there leaves the copy step + unpinned (a mutant dropping org_id survived exactly that). This is the object spend reads.""" + t_grant = _make_team("t-grant", ["srv1"]) + t_grant.organization_id = "org-team" + auth = _make_admitted_subject("sso-user") + auth.org_id = "org-user-primary" + server = MagicMock(server_id="srv1") + with self._patch(teams_by_id={"t-grant": t_grant}, user_teams=["t-grant"]): + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager._get_mcp_server_from_tool_name", + MagicMock(return_value=server), + ): + billed = await MCPRequestHandler.billing_auth_for_tool_call(auth, tool_name="t-grant/tool_a") + assert (billed.team_id, billed.org_id) == ("t-grant", "org-team") + assert (auth.team_id, auth.org_id) == (None, "org-user-primary"), "authz object must be untouched" + + async def test_own_grant_bills_the_user_not_a_team(self): + """A server the user's OWN grant reaches is not reached "through a team", so it bills the + user and their own org — attributing it to an unrelated team the user happens to belong to + would charge that team for access it never provided.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + t_other = _make_team("t-other", ["srv1"]) + auth = _make_admitted_subject("sso-user") + auth.object_permission = LiteLLM_ObjectPermissionTable(object_permission_id="op-u", mcp_servers=["srv1"]) + with self._patch(teams_by_id={"t-other": t_other}, user_teams=["t-other"]): + assert await MCPRequestHandler.attributing_source_for_server(auth, "srv1") is None + + async def test_billing_attribution_is_deterministic_across_several_granting_teams(self): + """When several teams grant the same server the pick must be stable and reproducible rather + than dependent on dict/roster ordering, or the same call bills different teams run to run.""" + teams = {"t-b": _make_team("t-b", ["srv1"]), "t-a": _make_team("t-a", ["srv1"])} + auth = _make_admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["t-b", "t-a"]): + first = await MCPRequestHandler.attributing_source_for_server(auth, "srv1") + with self._patch(teams_by_id=teams, user_teams=["t-a", "t-b"]): + second = await MCPRequestHandler.attributing_source_for_server(auth, "srv1") + assert first is not None and first.team_id == "t-a" + assert second is not None and second.team_id == "t-a", "roster order must not change who is billed" + + async def test_billing_auth_leaves_non_admitted_callers_untouched(self): + """Key and JWT billing must be byte-identical: the attribution wrapper returns the very same + object for anything that is not a keyless admitted subject.""" + key_auth = UserAPIKeyAuth(user_id="u", api_key="sk-hash", team_id="t1", org_id="org-a") + assert await MCPRequestHandler.billing_auth_for_tool_call(key_auth, tool_name="srv1-tool") is key_auth + + async def test_admitted_tools_never_run_the_single_credential_prelude(self): + """ORDERING is the invariant: the admitted branch is the FIRST statement of the tools + resolver, exactly as in the servers resolver. A fault in a lookup the subject never uses + (its own mcp_toolsets) must not reach it at all — when this branch sat after the prelude, + such a fault hit the fail-closed handler and denied tools its teams did grant.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + teams = {"t1": _make_team("t1", ["srv1"], tool_perms={"srv1": ["read"]})} + auth = _make_admitted_subject("sso-user") + # The subject must carry a toolset, or the prelude never resolves one and the fault below is + # unreachable — the branch could sit anywhere and the test would still pass (it did). + auth.object_permission = LiteLLM_ObjectPermissionTable(object_permission_id="op-u", mcp_toolsets=["ts-1"]) + boom = AsyncMock(side_effect=RuntimeError("toolset resolution exploded")) + with self._patch(teams_by_id=teams, user_teams=["t1"]): + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.resolve_toolset_tool_permissions", + boom, + ): + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth) + # The fault DOES fire, correctly, inside the subject's own source (which carries its + # toolsets) — that source contributes nothing. What must not happen is the top-level + # prelude running it first and denying the team's grant through the fail-closed handler. + assert tools == ["read"], "a fault in the subject's own toolsets must not deny its team's tools" + + async def test_admitted_own_byom_servers_stay_open(self): + """BYOM suppression-by-explicit-scope is a rule about a CREDENTIAL carrying its own + mcp_servers list. An admitted subject's object_permission is the user's own row, whose + mcp_servers column is [] by DB default — applying the rule would hide almost every admitted + user's OWN submitted servers. A key with an explicit scope still gets no BYOM widening.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + manager = self._manager_with(["srv-byom"]) + manager._get_active_submitted_mcp_server_ids_for_user = AsyncMock(return_value=["srv-byom"]) + db_default_perm = LiteLLM_ObjectPermissionTable(object_permission_id="op-u", mcp_servers=[]) + + admitted = _make_admitted_subject("sso-user") + admitted.object_permission = db_default_perm + scoped_key = UserAPIKeyAuth(user_id="u", api_key="sk-hash", object_permission=db_default_perm) + + assert await manager.operator_open_server_ids(admitted) == {"srv-byom"} + assert await manager.operator_open_server_ids(scoped_key) == set(), "explicit key scope still suppresses BYOM" + + async def test_admitted_admin_is_scoped_to_grants_not_full_registry(self): + """The wrapper's admin short-circuit hands the FULL registry to any admin-role auth before + the grant union or the per-team org ceilings run. A session bearer is a third-party client + credential, not the dashboard: an admin signing in through the connect flow gets their + grants like anyone else. A real admin key keeps the dashboard behavior unchanged.""" + from litellm.proxy._types import LitellmUserRoles + + manager = self._manager_with(["srv-granted", "srv-secret"]) + admitted = _make_admitted_subject("admin-user") + admitted.user_role = LitellmUserRoles.PROXY_ADMIN + with patch.object(MCPRequestHandler, "get_allowed_mcp_servers", AsyncMock(return_value=["srv-granted"])): + admitted_view = set(await manager.get_allowed_mcp_servers(admitted)) + key_admin_view = set( + await manager.get_allowed_mcp_servers( + UserAPIKeyAuth(user_id="admin-user", api_key="sk-hash", user_role=LitellmUserRoles.PROXY_ADMIN) + ) + ) + assert admitted_view == {"srv-granted"}, "an admitted admin gets their grants, not the registry" + assert key_admin_view == {"srv-granted", "srv-secret"}, "admin KEY behavior must be unchanged" + + async def test_admitted_opt_out_via_wrapper_keeps_team_servers(self): + """The wrapper's no_mcp_servers early-return is a KEY rule (a scoped credential's opt-out is + absolute). The admitted subject's opt-out silences only its own source, which the resolver + enforces per source — the wrapper must defer to it, or the resolver-level rule is dead code + on the production path.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable, SpecialMCPServerNames + + manager = self._manager_with(["srv-team"]) + opt_out = LiteLLM_ObjectPermissionTable( + object_permission_id="op-u", mcp_servers=[SpecialMCPServerNames.no_mcp_servers.value] + ) + admitted = _make_admitted_subject("sso-user") + admitted.object_permission = opt_out + with patch.object(MCPRequestHandler, "get_allowed_mcp_servers", AsyncMock(return_value=["srv-team"])): + admitted_view = set(await manager.get_allowed_mcp_servers(admitted)) + key_view = await manager.get_allowed_mcp_servers( + UserAPIKeyAuth(user_id="u", api_key="sk-hash", object_permission=opt_out) + ) + assert "srv-team" in admitted_view, "user opt-out must not zero team grants on the wrapper path" + assert key_view == [], "a key's opt-out stays absolute" + + async def test_open_channel_confers_reachability_not_a_ceiling_waiver(self): + """An open channel (allow_all_keys / own BYOM) makes a server REACHABLE. It is not a waiver + of the ceilings that bound it: the user's own mcp_tool_permissions still apply, exactly as a + virtual key's key_tools do on the same allow_all server. Returning None outright let a + session holder invoke tools their own policy excludes.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + auth = _make_admitted_subject("sso-user") + # The user is restricted to `read` on srv-open, and NO grant source names that server — + # it is reachable only through the open channel, which is exactly the bypass path. + auth.object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="op-u", mcp_servers=[], mcp_tool_permissions={"srv-open": ["read"]} + ) + open_ids = AsyncMock(return_value={"srv-open"}) + with self._patch(teams_by_id={}, user_teams=[]): + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.operator_open_server_ids", + open_ids, + ): + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv-open", auth) + assert tools == ["read"], "the user's own tool policy must still bind on an open-channel server" + + async def test_open_channel_server_gets_default_open_tools_for_admitted(self): + """A server reachable through an open channel (allow_all_keys / own BYOM) is granted by NO + source, so the source union alone returns [] — listable but uninvokable. The tools axis asks + the same open-channel owner the server union uses, so the server is default-open for tools + exactly as a virtual key experiences it.""" + auth = _make_admitted_subject("sso-user") + open_ids = AsyncMock(return_value={"srv-open"}) + with self._patch(teams_by_id={}, user_teams=[]): + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.operator_open_server_ids", + open_ids, + ): + open_tools = await MCPRequestHandler.get_allowed_tools_for_server("srv-open", auth) + closed_tools = await MCPRequestHandler.get_allowed_tools_for_server("srv-ungranted", auth) + assert open_tools is None, "open-channel server must be default-open for tools" + assert closed_tools == [], "a server no source or channel grants stays deny-all" + + async def test_over_budget_team_grants_nothing_and_healthy_team_stands(self): + """Budget ENFORCEMENT is the sibling of blocked: a team that has already exceeded its + max_budget is rejected outright for a virtual key pinned to it (common_checks), so it must + not keep granting servers, tools or throttle scope to a keyless union subject either. + Enforced through the SAME owner the key path uses (_team_max_budget_check). Distinct from + budget ATTRIBUTION of new spend, which stays with the user (documented deferral).""" + t_over = _make_team("t-over", ["srv1"]) + t_over.max_budget = 10.0 + t_over.spend = 11.0 + t_ok = _make_team("t-ok", ["srv2"]) + t_ok.max_budget = 10.0 + t_ok.spend = 1.0 + auth = _make_admitted_subject("sso-user") + with self._patch(teams_by_id={"t-over": t_over, "t-ok": t_ok}, user_teams=["t-over", "t-ok"]): + servers = set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) + limits = await MCPRequestHandler._admitted_subject_team_rpm_limits(auth) + assert servers == {"srv2"}, "an over-budget team must stop granting; the healthy team stands" + assert limits is None, "an over-budget team is not a source, so it stamps no throttle either" + + async def test_team_in_over_budget_org_grants_nothing(self): + """The org axis of the same rule, judged against the TEAM's own org (not the caller's + primary): a team owned by an org over its budget grants nothing, exactly as a key in that + org is rejected by _organization_max_budget_check.""" + t_in_broke_org = _make_team("t-b", ["srv1"]) + t_in_broke_org.organization_id = "org-broke" + # object_permission_id=None: the org has NO MCP ceiling, so the source is denied by the + # budget gate alone. A truthy auto-Mock id here made an earlier version of this test pass + # through the org-CEILING fault path with the budget gate deleted — vacuous. + org = MagicMock(object_permission_id=None, litellm_budget_table=MagicMock(max_budget=5.0), spend=9.0) + auth = _make_admitted_subject("sso-user") + with self._patch(teams_by_id={"t-b": t_in_broke_org}, user_teams=["t-b"], orgs_by_id={"org-broke": org}): + servers = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert servers == [], "a team in an over-budget org must not grant through the union" + + async def test_one_faulting_team_does_not_deny_the_other_sources(self): + """The unit of fault isolation is the SOURCE. One team's row being momentarily unreadable + contributes nothing for THAT team (access only narrows) while the user's own grants and every + other resolvable team stand — it must not collapse the whole union to deny-all on either + axis.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + t_ok = _make_team("t-ok", ["srv1"], tool_perms={"srv1": ["read"]}) + auth = _make_admitted_subject("sso-user") + auth.object_permission = LiteLLM_ObjectPermissionTable(object_permission_id="op-u", mcp_servers=["srv-own"]) + teams = {"t-ok": t_ok} # t-boom absent from the map -> our patched get_team_object RAISES for it + + async def _team_or_boom(team_id, **kw): + if team_id not in teams: + raise RuntimeError(f"transient DB blip loading {team_id}") + return teams[team_id] + + with self._patch(teams_by_id=teams, user_teams=["t-boom", "t-ok"]): + with patch("litellm.proxy.auth.auth_checks.get_team_object", _team_or_boom): + servers = set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth) + assert servers == {"srv-own", "srv1"}, "healthy sources must stand when one team faults" + assert tools == ["read"], "the healthy team's tool grant must survive the other team's fault" + + async def test_key_org_tool_ceiling_fault_keeps_key_restrictions(self): + """Virtual-key tools axis mirrors its servers axis on an unresolvable org ceiling: the org + intersect is SKIPPED and the key's own tool restrictions stand. Letting the fault escape + collapsed the whole resolution to None (allow-all), which is fail-open WIDER than before the + fault — key restrictions must never be dropped by an org lookup blip.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + key_auth = UserAPIKeyAuth(user_id="u", api_key="sk-hash", org_id="org-a") + key_auth.object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="op-k", mcp_servers=["srv1"], mcp_tool_permissions={"srv1": ["read"]} + ) + boom = AsyncMock(side_effect=RuntimeError("org permission load exploded")) + with self._patch(teams_by_id={}, user_teams=[]): + with patch.object(MCPRequestHandler, "_get_org_object_permission", boom): + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", key_auth) + assert tools == ["read"], "key tool restrictions must survive an unresolvable org ceiling" + + async def test_team_rpm_limit_binds_only_within_that_teams_grant_scope(self): + """A limit rides the same scope as the access it bounds. A roster team is charged ONLY for + servers its own grant reaches: not for a server the user reaches through a DIFFERENT team + (else this user's calls drain a bucket shared by that team's keys for access the team never + provided), not for map entries beyond its grant, and never when the team is blocked.""" + # t-granting grants srv1 and limits it; also names srv9 in its map, which it does NOT grant. + t_granting = _make_team("t-granting", ["srv1"]) + t_granting.metadata = {"mcp_rpm_limit": {"srv1": 5, "srv9": 7}} + # t-other grants only srv2 but retains limit metadata for srv1 -> must not be charged for it. + t_other = _make_team("t-other", ["srv2"]) + t_other.metadata = {"mcp_rpm_limit": {"srv1": 3}} + # t-blocked grants srv1 and limits it, but is blocked -> grants nothing, charges nothing. + t_blocked = _make_team("t-blocked", ["srv1"]) + t_blocked.metadata = {"mcp_rpm_limit": {"srv1": 2}} + t_blocked.blocked = True + + auth = _make_admitted_subject("sso-user") + teams = {"t-granting": t_granting, "t-other": t_other, "t-blocked": t_blocked} + with self._patch(teams_by_id=teams, user_teams=["t-granting", "t-other", "t-blocked"]): + limits = await MCPRequestHandler._admitted_subject_team_rpm_limits(auth) + + assert limits == {"t-granting": {"srv1": 5}}, ( + "only the granting team's bucket, and only for the server it grants" + ) + + async def test_non_roster_team_rpm_limit_does_not_apply(self): + """The roster gates grants and throttles through one owner, so a team the user was removed + from neither grants servers nor gets charged for their calls.""" + stale = _make_team("t-stale", ["srv1"], members=("someone-else",)) + stale.metadata = {"mcp_rpm_limit": {"srv1": 1}} + auth = _make_admitted_subject("sso-user") + with self._patch(teams_by_id={"t-stale": stale}, user_teams=["t-stale"]): + limits = await MCPRequestHandler._admitted_subject_team_rpm_limits(auth) + assert limits is None + + async def test_org_list_caps_a_source_but_never_becomes_a_grant(self): + """The admitted model is a union of GRANTS, so an org allowlist may only narrow what a source + already grants. For a virtual key with no lower-level restriction the org list legitimately + BECOMES the allowed set, and inheriting that arm would hand every admitted user with an + org_id their whole org's server list with no direct or team grant behind it.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + auth = _make_admitted_subject("sso-user") + auth.org_id = "org-a" # org allows srv1+srv2; the user and their teams grant NOTHING + org_perm = AsyncMock( + return_value=LiteLLM_ObjectPermissionTable(object_permission_id="op-org-a", mcp_servers=["srv1", "srv2"]) + ) + with self._patch(teams_by_id={}, user_teams=[]): + with patch.object(MCPRequestHandler, "_get_org_object_permission", org_perm): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert result == [], "an org ceiling must not grant servers the user was never granted" + + async def test_tool_ceiling_fails_closed_when_a_SOURCE_faults(self): + """Each source is resolved through an UNMARKED auth, so a fault under a source must still + deny. Returning None there would win the union as allow-all and drop every team/org tool + ceiling on a DB blip -- the marker alone only covers faults raised before the fan-out.""" + auth = _make_admitted_subject("sso-user") + teams = {"t1": _make_team("t1", ["srv1"])} + # Fault INSIDE the tool resolution only. Faulting something the server path also uses would + # make the source grant nothing, so the union would return [] without the tool path ever + # running -- the test would pass while pinning nothing (an earlier version did exactly that). + boom = AsyncMock(side_effect=RuntimeError("org tool ceiling exploded")) + with self._patch(teams_by_id=teams, user_teams=["t1"]): + assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == ["srv1"] # control: granted + with patch.object(MCPRequestHandler, "_apply_agent_and_org_tool_ceilings", boom): + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth) + assert tools == [], "a source-level fault must deny tools, never collapse to allow-all" + + async def test_own_opt_out_silences_only_that_source_not_the_teams(self): + """no_mcp_servers on the USER's own grants opts that source out. It must not zero the teams: + the sources are independent, so an opt-out on one silences one. (The same sentinel on a + virtual KEY still overrides team inheritance -- that is the key ceiling model, unchanged.)""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable, SpecialMCPServerNames + + auth = _make_admitted_subject("sso-user") + auth.object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="op-user", mcp_servers=[SpecialMCPServerNames.no_mcp_servers.value] + ) + with self._patch(teams_by_id={"t1": _make_team("t1", ["srv1"])}, user_teams=["t1"]): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert set(result) == {"srv1"}, "the user's own opt-out must not zero their team's grants" + + async def test_sources_fan_out_per_team_and_drop_non_roster_teams(self): + """The fan-out lives here now. One source per grant source: the user's own grants (no team_id, + carrying their object_permission) plus each team they are a LIVE roster member of. A team that + lingers in the user's cached `teams` array but no longer lists them in members_with_roles is + dropped, which is what revokes access after a team_member_delete the user row hasn't caught up + on. Each team source carries that team's own org, which is what makes the shared resolver apply + the team's owning-org ceiling rather than the caller's home org.""" + teams = { + "t-member": _make_team("t-member", ["srv1"], members=("sso-user",)), + "t-stale": _make_team("t-stale", ["srv2"], members=("someone-else",)), + } + teams["t-member"].organization_id = "org-a" + auth = _make_admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["t-member", "t-stale"]): + sources = await MCPRequestHandler._admitted_subject_sources(auth) + + assert [(s.team_id, s.org_id) for s in sources] == [(None, None), ("t-member", "org-a")] + # The user's own source carries their grants; a team source must NOT, or the team would be + # widened by grants the team never made. + assert sources[0].object_permission is auth.object_permission + assert sources[1].object_permission is None + # Every source is an ordinary caller, so it cannot re-enter the admitted fan-out. + assert all(not s.mcp_admitted_user_subject for s in sources) + # Nothing that meters or elevates the request may ride along onto a per-source clone. + assert all(s.api_key is None and s.user_role is None for s in sources) + + async def test_jwt_keyless_user_without_team_claim_does_not_union(self): + """Regression for the review finding: a JWT-authenticated caller is also keyless with a + user_id and (with no team claim) no team_id, but it is NOT admission-marked, so it must + keep its prior behavior of inheriting no team grants rather than silently gaining the + union across every team the user belongs to.""" + teams = {"team-a": _make_team("team-a", ["srv1"]), "team-b": _make_team("team-b", ["srv2"])} + jwt_auth = UserAPIKeyAuth(user_id="jwt-user", api_key=None) # no admission marker + with self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"]): + result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(jwt_auth) + assert result == [] + + async def test_forged_metadata_marker_on_a_real_key_grants_no_union(self): + """Security regression (forged admission marker): the admitted-subject marker is a + server-only ``UserAPIKeyAuth`` field, NOT a metadata key, precisely because virtual-key + metadata is caller-controlled at key creation. A user who sets + ``mcp_admitted_user_subject: true`` in their own key's metadata (api_key present, no + team_id) must NOT be treated as an admitted subject and must gain no cross-team union.""" + teams = {"team-a": _make_team("team-a", ["srv1"]), "team-b": _make_team("team-b", ["srv2"])} + forged = UserAPIKeyAuth( + user_id="attacker", + api_key="sk-real-key", + metadata={"mcp_admitted_user_subject": True}, # caller-forged marker in key metadata + ) + assert _is_mcp_admitted_user_subject(forged) is False + with self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"]): + assert await MCPRequestHandler._team_ids_for_mcp_grant(forged) == [] + assert await MCPRequestHandler._get_allowed_mcp_servers_for_team(forged) == [] + + async def test_admitted_subject_team_tool_restriction_binds(self): + """Security regression (team tool restrictions bypassed): a keyless admitted subject whose + granting team restricts ``srv1`` to ``{tool_a}`` must NOT receive allow-all on srv1. The + single-team-id tool lookup returns None (allow-all) for a keyless multi-team user, dropping + the exclusion; the union across granting teams restores it.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable, Member + + team = LiteLLM_TeamTable( + team_id="team-a", + members_with_roles=[Member(user_id="sso-user", role="user")], + access_group_ids=[], + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="op-team-a", + mcp_servers=["srv1"], + mcp_tool_permissions={"srv1": ["tool_a"]}, + ), + ) + auth = _make_admitted_subject("sso-user") + with self._patch(teams_by_id={"team-a": team}, user_teams=["team-a"]): + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth) + assert tools == ["tool_a"] + + async def test_blocked_team_grants_no_servers_to_admitted_subject(self): + """Security regression: a blocked team grants nothing. The central policy gate enforces this + for a key pinned to a single team_id, but a keyless admitted subject unions across ALL its + teams (no team_id), so a blocked team's MCP grants must be dropped at the per-team resolver.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable, Member + + blocked = LiteLLM_TeamTable( + team_id="team-blocked", + blocked=True, + members_with_roles=[Member(user_id="sso-user", role="user")], + access_group_ids=[], + object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-blk", mcp_servers=["srv-secret"]), + ) + teams = {"team-ok": _make_team("team-ok", ["srv-ok"]), "team-blocked": blocked} + auth = _make_admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-ok", "team-blocked"]): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert set(result) == {"srv-ok"} + + async def test_admitted_subject_not_on_team_roster_gets_no_grant(self): + """Security regression (membership containment): a keyless subject whose user_id is NOT on a + team's roster inherits nothing from it, even when the team id lingers in the user's (stale or + cached) teams array. The team roster is the source of truth, so a removed or foreign + membership revokes access at the union rather than granting it.""" + teams = {"team-x": _make_team("team-x", ["srv-x"], members=("someone-else",))} + auth = _make_admitted_subject("sso-user") # in user.teams for team-x, but NOT on its roster + with self._patch(teams_by_id=teams, user_teams=["team-x"]): + result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(auth) + assert result == [] + + async def test_tool_resolution_fails_closed_on_db_error(self): + """Security regression: ANY error resolving the tool allowlist for a keyless admitted subject + must DENY the server's tools ([]) rather than collapse to allow-all (None). Patches an await + OUTSIDE the multi-team fan-out (the team-object lookup) to prove the whole function fails + closed, not just the one helper — mirroring the fail-closed server path.""" + auth = _make_admitted_subject("sso-user") + with patch.object( + MCPRequestHandler, + "_get_team_object_permission", + new=AsyncMock(side_effect=RuntimeError("db blip")), + ): + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth) + assert tools == [] + + async def test_admission_marker_cannot_be_set_from_validated_input(self): + """Defense-in-depth: the mcp_admitted_user_subject marker is server-only. Supplying it in any + validated input (constructor kwargs OR model_validate, e.g. a future JWT/key claim splat) is + stripped by the before-validator, so ONLY the admission path's post-construction assignment + can set it.""" + via_kwarg = UserAPIKeyAuth(user_id="u", api_key=None, mcp_admitted_user_subject=True) + via_validate = UserAPIKeyAuth.model_validate({"user_id": "u", "mcp_admitted_user_subject": True}) + assert via_kwarg.mcp_admitted_user_subject is False + assert via_validate.mcp_admitted_user_subject is False + assert _is_mcp_admitted_user_subject(via_kwarg) is False + assert _is_mcp_admitted_user_subject(via_validate) is False + + +@pytest.mark.asyncio +class TestAdmittedSubjectPerTeamOrgCap: + """A keyless admitted subject unions grants across teams that may span organizations. Each team's + grant (servers AND tools) is capped by that team's OWN org, and the user's direct grants by the + user's own org — never the caller's primary org applied over the whole cross-org union. Guards the + Veria 'team grants bypass their owning policies' finding.""" + + #: sentinel for org_perms: org has an object_permission_id but its load returns None (a swallowed + #: DB error / dangling id), which _object_permission_for_org must treat as fail-closed. + LOAD_FAILS = "__load_fails__" + + @contextlib.contextmanager + def _patch(self, *, teams_by_id, user_teams, org_perms=None, registry=None): + """org_perms: {org_id: LiteLLM_ObjectPermissionTable | None | LOAD_FAILS}. + - table → org exists, ceiling = that permission. + - None → org exists but carries no object_permission (no ceiling). + - LOAD_FAILS → org exists with an object_permission_id, but the permission load returns None. + - org_id ABSENT from the map → org row missing: get_org_object RAISES a bare Exception, exactly + as production does (it does NOT return None or raise HTTPException).""" + org_perms = org_perms or {} + + async def _get_team_object(team_id, **kw): + return teams_by_id.get(team_id) + + async def _get_user_object(user_id, **kw): + return MagicMock(user_id=user_id, teams=user_teams) + + async def _get_org_object(org_id, **kw): + if org_id not in org_perms: + from litellm.proxy.auth.auth_checks import OrganizationNotFoundError + + # matches production: a CONFIRMED-absent org raises this specific type, so callers + # can tell it apart from an outage (a bare Exception now means "lookup failed"). + raise OrganizationNotFoundError(f"Organization doesn't exist. Org={org_id}.") + op = org_perms[org_id] + has_permission_id = op is not None # a table OR LOAD_FAILS carries an id; None does not + return MagicMock( + organization_id=org_id, + object_permission_id=(f"orgop-{org_id}" if has_permission_id else None), + # Real typed values: the budget owners compare these, and a bare MagicMock attribute + # would explode the comparison and silently drop the source (bare-Mock rule). + litellm_budget_table=None, + spend=0.0, + ) + + async def _get_object_permission(object_permission_id, **kw): + for oid, op in org_perms.items(): + if op is not None and op != self.LOAD_FAILS and object_permission_id == f"orgop-{oid}": + return op + return None # LOAD_FAILS (or an unknown id) → None, simulating get_object_permission's swallow + + cms = [ + patch("litellm.proxy.auth.auth_checks.get_team_object", _get_team_object), + patch("litellm.proxy.auth.auth_checks.get_user_object", _get_user_object), + patch("litellm.proxy.auth.auth_checks.get_org_object", _get_org_object), + patch("litellm.proxy.auth.auth_checks.get_object_permission", _get_object_permission), + patch( + "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + AsyncMock(return_value=[]), + ), + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + ] + if registry is not None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + # registry may be a list of bare server_ids (MagicMock servers) OR a dict of + # {server_id: server_obj} for tests that need real alias/name resolution (config servers). + reg = registry if isinstance(registry, dict) else {s: MagicMock() for s in registry} + cms.append(patch.object(global_mcp_server_manager, "get_registry", return_value=reg)) + with contextlib.ExitStack() as es: + for cm in cms: + es.enter_context(cm) + yield + + # ---- server axis ---- + + async def test_team_grant_capped_by_its_own_org(self): + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + teams = {"team-a": _make_team("team-a", ["srv1", "srv2"], org_id="org-a")} + org_perms = {"org-a": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-a", mcp_servers=["srv1"])} + auth = _make_admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms=org_perms): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert set(result) == {"srv1"} # srv2 capped out by org-a's ceiling + + async def test_cross_org_teams_each_capped_by_own_org(self): + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + teams = { + "team-a": _make_team("team-a", ["srv1", "srv2"], org_id="org-a"), + "team-b": _make_team("team-b", ["srv3", "srv4"], org_id="org-b"), + } + org_perms = { + "org-a": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-a", mcp_servers=["srv1"]), + "org-b": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-b", mcp_servers=["srv3"]), + } + auth = _make_admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"], org_perms=org_perms): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert set(result) == {"srv1", "srv3"} # each team clipped by its OWN org, then unioned + + async def test_all_proxy_grant_capped_by_org(self): + from litellm.proxy._types import LiteLLM_ObjectPermissionTable, SpecialMCPServerName + + teams = {"team-a": _make_team("team-a", [SpecialMCPServerName.all_proxy_servers.value], org_id="org-a")} + org_perms = {"org-a": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-a", mcp_servers=["srv1"])} + auth = _make_admitted_subject("sso-user") + with self._patch( + teams_by_id=teams, user_teams=["team-a"], org_perms=org_perms, registry=["srv1", "srv2", "srv3"] + ): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + # all_proxy expands to the whole registry, then org-a caps to {srv1} — the cell the old partial + # patch missed (it returned the full registry before capping). + assert set(result) == {"srv1"} + + async def test_org_row_without_object_permission_does_not_cap(self): + teams = {"team-a": _make_team("team-a", ["srv1", "srv2"], org_id="org-a")} + auth = _make_admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms={"org-a": None}): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert set(result) == {"srv1", "srv2"} # empty ceiling = no restriction + + async def test_direct_grants_unioned_with_team_and_capped_by_user_org(self): + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + teams = {"team-a": _make_team("team-a", ["srv1"], org_id="org-a")} + org_perms = { + "org-a": None, # the team's org imposes no ceiling + "org-u": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-u", mcp_servers=["srvD", "srv1"]), + } + auth = _make_admitted_subject("sso-user", org_id="org-u", own_servers=["srvD", "srvX"]) + with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms=org_perms): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + # direct {srvD,srvX} ∩ user-org {srvD,srv1} = {srvD}; UNIONed with team {srv1} (not intersected). + # srvX capped out by the user's org; team's srv1 NOT clipped by the user's primary org. + assert set(result) == {"srvD", "srv1"} + + async def test_single_team_key_uses_primary_org_cap_not_per_team(self): + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + # A KEY (not admitted): the per-team org cap must NOT fire; the top-level primary-org cap applies, + # byte-identical to before. team-a (org-a) grants {srv1,srv2}; the key's primary org is org-k. + teams = {"team-a": _make_team("team-a", ["srv1", "srv2"], org_id="org-a")} + org_perms = { + "org-a": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-a", mcp_servers=["srv2"]), + "org-k": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-k", mcp_servers=["srv1"]), + } + key_auth = UserAPIKeyAuth(user_id="u", api_key="sk-hash", team_id="team-a", org_id="org-k") + with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms=org_perms): + result = await MCPRequestHandler.get_allowed_mcp_servers(key_auth) + # If the per-team (org-a) cap wrongly fired, team-a would clip to {srv2} then org-k → {} (empty). + # Correct key behavior: no per-team cap; primary-org (org-k) cap → {srv1}. + assert set(result) == {"srv1"} + + # ---- tool axis ---- + + async def test_org_tool_ceiling_binds_when_team_places_no_tool_restriction(self): + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + # team grants srv1 with NO tool restriction; org-a restricts srv1's tools to {tool_a}. + teams = {"team-a": _make_team("team-a", ["srv1"], org_id="org-a")} + org_perms = { + "org-a": LiteLLM_ObjectPermissionTable( + object_permission_id="orgop-org-a", + mcp_servers=["srv1"], + mcp_tool_permissions={"srv1": ["tool_a"]}, + ) + } + auth = _make_admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms=org_perms): + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth) + # Without the per-team org tool ceiling this would be None (all tools) — org-a's tool ceiling + # would be bypassed exactly like the server case. + assert tools == ["tool_a"] + + async def test_tool_union_across_cross_org_teams(self): + teams = { + "team-a": _make_team("team-a", ["srv1"], org_id="org-a", tool_perms={"srv1": ["t1"]}), + "team-b": _make_team("team-b", ["srv1"], org_id="org-b", tool_perms={"srv1": ["t2"]}), + } + auth = _make_admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"], org_perms={"org-a": None, "org-b": None}): + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth) + assert set(tools) == {"t1", "t2"} + + async def test_tool_deny_all_when_team_grant_and_org_tool_ceiling_disjoint(self): + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + teams = {"team-a": _make_team("team-a", ["srv1"], org_id="org-a", tool_perms={"srv1": ["t1"]})} + org_perms = { + "org-a": LiteLLM_ObjectPermissionTable( + object_permission_id="orgop-org-a", + mcp_servers=["srv1"], + mcp_tool_permissions={"srv1": ["t2"]}, + ) + } + auth = _make_admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms=org_perms): + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth) + # team {t1} ∩ org {t2} = {} → deny every tool ([]), NOT allow-all (None). + assert tools == [] + + # ---- error contract (adversarial-review findings) ---- + + async def test_missing_org_row_is_treated_as_no_ceiling_not_lockout(self): + """A team's organization_id may point to an org row that no longer exists (deleted / not yet + synced). get_org_object RAISES a bare Exception for that; it must be treated as 'no ceiling' and + must NOT lock the admitted subject out of the team's grants (parity with the key path, which + tolerates a deleted org).""" + teams = {"team-a": _make_team("team-a", ["srv1", "srv2"], org_id="org-gone")} + auth = _make_admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms={}): # org-gone absent → raises + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert set(result) == {"srv1", "srv2"} + + async def test_org_permission_load_failure_fails_closed(self): + """The org carries an object_permission_id but the permission load returns None (a swallowed DB + error / dangling id). The ceiling cannot be verified, so the admitted subject must fail CLOSED + for that team — NOT skip the ceiling, which would leak org-forbidden servers.""" + teams = {"team-a": _make_team("team-a", ["srv1", "srv2"], org_id="org-a")} + auth = _make_admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms={"org-a": self.LOAD_FAILS}): + # Asserted through the PUBLIC resolver: the per-source org ceiling is applied there now, + # so calling the single-team helper would return [] for an admitted subject either way + # and pin nothing. + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert result == [] # fail closed, not {srv1, srv2} + + # ---- open bot-thread findings (2026-07-21 re-review) ---- + + async def test_org_less_team_grant_capped_by_user_primary_org(self): + """HIGH (cursor): a team with NO organization_id must still be bounded by the user's PRIMARY + org — otherwise, since admitted subjects skip the top-level primary-org cap, an org-less team's + grant would bypass every org ceiling and reach servers the user's home org forbids.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + teams = {"team-noorg": _make_team("team-noorg", ["srv1", "srv2"], org_id=None)} + org_perms = {"org-U": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-U", mcp_servers=["srv1"])} + auth = _make_admitted_subject("sso-user", org_id="org-U") + with self._patch(teams_by_id=teams, user_teams=["team-noorg"], org_perms=org_perms): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + # org-less team falls back to the user's primary org (org-U → {srv1}); srv2 capped out. + assert set(result) == {"srv1"} + + async def test_tool_empty_contributions_fails_closed(self): + """MEDIUM (greptile/cursor): when no source in the tool-resolution view grants the server (a + TOCTOU/cache-lag inconsistency on a server that passed the server gate), the admitted path must + fail CLOSED (deny all tools = []), NOT allow-all (None).""" + teams = {"team-a": _make_team("team-a", ["srv1"], org_id="org-a")} + auth = _make_admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-a"], org_perms={"org-a": None}): + # 'srv-nobody' is granted by neither the team nor the user directly → empty contributions. + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv-nobody", auth) + assert tools == [] + + async def test_tool_no_db_honors_in_memory_direct_restriction(self): + """MEDIUM (cursor): with no DB, the tool path must still honor the user's OWN in-memory + object_permission tool restriction (resolvable without a DB) rather than blanket-allow (None).""" + auth = _make_admitted_subject( + "sso-user", own_servers=["srv1"], own_tool_perms={"srv1": ["t1"]} + ) # no org_id, direct grant of srv1 restricted to {t1} + with patch("litellm.proxy.proxy_server.prisma_client", None): + tools = await MCPRequestHandler.get_allowed_tools_for_server("srv1", auth) + assert tools == ["t1"] # in-memory restriction honored, not widened to all tools + + # ---- config.yaml-defined servers (incl. OAuth) ---- + + async def test_config_defined_oauth_server_by_alias_reached_and_org_capped(self): + """A config.yaml-defined MCP OAuth server flows through the SAME resolution as a DB server: + the team grant (and the org ceiling) reference it by ALIAS, expand_permission_list resolves it + via the config+DB registry union to its server_id, and the per-team org cap applies identically. + (The config server's OAuth *client* persistence is #33768 — an orthogonal egress concern; this + pins the grant/reachability side of the 10x flow for config-defined servers.)""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + cfg_server = MagicMock() + cfg_server.server_id = "cfg-oauth-1" + cfg_server.alias = "linear_cfg" + cfg_server.server_name = "linear_cfg" + cfg_server.name = "linear_cfg" + + # team grants the config server BY ALIAS alongside a DB-style bare id; org-a's ceiling lists + # ONLY the config server (also by alias). + teams = {"team-a": _make_team("team-a", ["linear_cfg", "srv-db"], org_id="org-a")} + org_perms = { + "org-a": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-a", mcp_servers=["linear_cfg"]) + } + auth = _make_admitted_subject("sso-user") + with self._patch( + teams_by_id=teams, + user_teams=["team-a"], + org_perms=org_perms, + registry={"cfg-oauth-1": cfg_server}, + ): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + # 'linear_cfg' alias resolves to the config server_id and survives org-a's ceiling; 'srv-db' + # (not in org-a's allowlist) is capped out — same per-team org cap, config server included. + assert set(result) == {"cfg-oauth-1"} + + async def test_config_oauth_server_alias_resolution_feeds_the_org_cap(self): + """A config-defined OAuth server granted BY ALIAS whose OWN org forbids it is capped out — AND the + cap is proven to run on RESOLVED server_ids, not raw strings. A control config server, granted by + alias and allowed by the org via its RESOLVED id, must survive: that inclusion is impossible unless + expand_permission_list resolved the grant alias to the id the ceiling lists, so a broken alias path + yields {} and FAILS this test — whereas a bare `assert empty` would pass even if resolution never + ran (the weakness Cursor flagged).""" + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + forbidden = MagicMock() # granted by alias, but its org forbids it → must be capped out + forbidden.server_id = "cfg-oauth-1" + forbidden.alias = forbidden.server_name = forbidden.name = "linear_cfg" + control = MagicMock() # granted by alias, allowed by the org via its RESOLVED id → must survive + control.server_id = "control-id" + control.alias = control.server_name = control.name = "control_alias" + + teams = {"team-b": _make_team("team-b", ["linear_cfg", "control_alias"], org_id="org-b")} + # org-b's ceiling allows ONLY the control server, referenced by its RESOLVED server_id. + org_perms = { + "org-b": LiteLLM_ObjectPermissionTable(object_permission_id="orgop-org-b", mcp_servers=["control-id"]) + } + auth = _make_admitted_subject("sso-user") + with self._patch( + teams_by_id=teams, + user_teams=["team-b"], + org_perms=org_perms, + registry={"cfg-oauth-1": forbidden, "control-id": control}, + ): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + # control survives ('control_alias' resolved to 'control-id', matching the id-based ceiling); the + # forbidden config server ('cfg-oauth-1') is capped out. A broken alias path → {} → fails here. + assert set(result) == {"control-id"} + + +@pytest.mark.asyncio +class TestSessionBearerEgressScrub: + """The gateway session bearer / bridge envelope is an admission credential, never an upstream token. + The leak-defense scrub is anchored to the credential SHAPE, so a session-shaped Authorization is + stripped from every egress context even when it reaches a non-aggregate scope that never set the + admission marker (design-review finding: a session bearer misdirected to a per-server true_passthrough + path would otherwise be forwarded upstream verbatim and replayed against the aggregate endpoint).""" + + async def test_session_bearer_misdirected_to_passthrough_is_scrubbed(self): + from litellm.types.mcp import MCPAuth + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/passthrough_server", + "headers": [(b"authorization", b"Bearer llm_session_synthetic-shaped-token")], + } + ttp_server = MagicMock() + ttp_server.auth_type = MCPAuth.true_passthrough + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + new_callable=AsyncMock, + ) as mock_auth, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ttp_server + (_auth, _mah, _srv, _sah, oauth2_headers, raw_headers) = await MCPRequestHandler.process_mcp_request(scope) + + mock_auth.assert_not_called() # true_passthrough → LiteLLM auth skipped (anonymous arm, no marker) + assert oauth2_headers is None # session-shaped bearer scrubbed from oauth2 egress + assert all(k.lower() != "authorization" for k in raw_headers) # ...and from raw egress headers + + async def test_legitimate_upstream_token_is_not_scrubbed(self): + """A genuine upstream/passthrough token is never session- or envelope-shaped, so the shape-anchored + scrub must leave it intact for forwarding (guards against over-stripping).""" + from litellm.types.mcp import MCPAuth + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/passthrough_server", + "headers": [(b"authorization", b"Bearer real-upstream-opaque-token-xyz")], + } + ttp_server = MagicMock() + ttp_server.auth_type = MCPAuth.true_passthrough + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + new_callable=AsyncMock, + ), + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ttp_server + (_auth, _mah, _srv, _sah, oauth2_headers, _raw) = await MCPRequestHandler.process_mcp_request(scope) + + assert oauth2_headers.get("Authorization") == "Bearer real-upstream-opaque-token-xyz" + + async def test_scrub_removes_gateway_credential_from_every_egress_context(self): + """The scrub is anchored to the credential SHAPE and covers ALL egress contexts, not just + Authorization: a session bearer placed in x-mcp-auth OR a per-server x-mcp-{alias}-authorization + header is stripped too (the High-severity gap: those were forwarded upstream before).""" + sess = "Bearer llm_session_abc" + oauth2, raw, mcp_auth, per_server = MCPRequestHandler._scrub_gateway_admission_credentials( + admitted=False, + oauth2_headers={"Authorization": sess}, + raw_headers={ + "authorization": sess, + "x-mcp-auth": "llm_session_xyz", + "x-mcp-github-authorization": "llm_session_ghi", + }, + mcp_auth_header="llm_session_xyz", + mcp_server_auth_headers={"github": {"Authorization": "llm_session_ghi"}}, + ) + assert oauth2 is None + assert "authorization" not in {k.lower() for k in raw} + assert all("llm_session_" not in v for v in raw.values()) # x-mcp-auth + per-server raw values gone + assert mcp_auth is None # deprecated x-mcp-auth value scrubbed + assert per_server == {} # per-server session bearer removed → now-empty server dict dropped + + async def test_scrub_keeps_real_upstream_tokens(self): + """A legitimate upstream token is never session-/envelope-shaped, so every context is forwarded + unchanged — guards against over-stripping a real credential the caller meant for the upstream.""" + oauth2, raw, mcp_auth, per_server = MCPRequestHandler._scrub_gateway_admission_credentials( + admitted=False, + oauth2_headers={"Authorization": "Bearer real-upstream-xyz"}, + raw_headers={"authorization": "Bearer real-upstream-xyz", "x-mcp-github-authorization": "Bearer gh_real"}, + mcp_auth_header="some-api-key-123", + mcp_server_auth_headers={"github": {"Authorization": "Bearer gh_real"}}, + ) + assert oauth2 == {"Authorization": "Bearer real-upstream-xyz"} + assert raw["authorization"] == "Bearer real-upstream-xyz" + assert mcp_auth == "some-api-key-123" + assert per_server == {"github": {"Authorization": "Bearer gh_real"}} + + async def test_scrub_admitted_drops_authorization_but_keeps_injected_upstream_token(self): + """An admitted subject's top-level Authorization is dropped unconditionally, while the real + upstream token the bridge arm INJECTS into a per-server header (not gateway-shaped) survives.""" + oauth2, raw, mcp_auth, per_server = MCPRequestHandler._scrub_gateway_admission_credentials( + admitted=True, + oauth2_headers={"Authorization": "Bearer llm_session_abc"}, + raw_headers={"authorization": "Bearer llm_session_abc"}, + mcp_auth_header=None, + mcp_server_auth_headers={"github": {"Authorization": "Bearer gh_injected_upstream"}}, + ) + assert oauth2 is None + assert "authorization" not in {k.lower() for k in raw} + assert per_server == {"github": {"Authorization": "Bearer gh_injected_upstream"}} diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 692e5340f48..a61e9de3281 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -3302,19 +3302,20 @@ async def test_token_root_does_not_resolve_private_server_for_external_client(): @pytest.mark.asyncio -async def test_register_root_resolves_single_oauth2_server(): - """When /register is hit without server name and exactly 1 OAuth2 server exists, resolve it.""" - try: - from fastapi import Request +async def test_register_root_does_aggregate_dcr_not_single_server_resolution(): + """Root /register is the aggregate DCR endpoint: it mints a stateless llm_dcrc_ client + from the request's redirect_uris and does NOT resolve a single configured oauth2 server + (a single-server deployment registers at /{server}/register instead).""" + import json - from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( - register_client, - ) - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - global_mcp_server_manager, - ) - except ImportError: - pytest.skip("MCP discoverable endpoints not available") + from fastapi import Request + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + register_client, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) global_mcp_server_manager.registry.clear() oauth2_server = _create_oauth2_server() @@ -3325,33 +3326,37 @@ async def test_register_root_resolves_single_oauth2_server(): mock_request.headers = {} try: - with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", - new=AsyncMock(return_value={}), + with ( + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + new=AsyncMock(return_value={"redirect_uris": ["https://claude.ai/cb"]}), + ), + patch("litellm.proxy.proxy_server.master_key", "sk-test-salt-for-lit3637"), ): - result = await register_client(request=mock_request, mcp_server_name=None) + response = await register_client(request=mock_request, mcp_server_name=None) - # Should resolve to the single server and return its name as client_id - assert result["client_id"] == "test_oauth" - assert "redirect_uris" in result + body = json.loads(response.body) + assert body["client_id"].startswith("llm_dcrc_") + assert body["client_id"] != "test_oauth" + assert body["token_endpoint_auth_method"] == "none" finally: global_mcp_server_manager.registry.clear() @pytest.mark.asyncio -async def test_register_root_does_not_resolve_private_server_for_external_client(): - """Root /register must not reveal or use a hidden MCP server.""" - try: - from fastapi import Request +async def test_register_root_does_not_leak_a_private_server(): + """Root /register never resolves or reveals a configured server, so a private one cannot + leak to an external caller: it always mints the aggregate DCR client instead.""" + import json - from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( - register_client, - ) - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - global_mcp_server_manager, - ) - except ImportError: - pytest.skip("MCP discoverable endpoints not available") + from fastapi import Request + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + register_client, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) global_mcp_server_manager.registry.clear() oauth2_server = _create_oauth2_server(available_on_public_internet=False) @@ -3365,17 +3370,19 @@ async def test_register_root_does_not_resolve_private_server_for_external_client with ( patch( "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", - new=AsyncMock(return_value={}), + new=AsyncMock(return_value={"redirect_uris": ["https://claude.ai/cb"]}), ), patch( "litellm.proxy._experimental.mcp_server.discoverable_endpoints.IPAddressUtils.get_mcp_client_ip", return_value="198.51.100.10", ), + patch("litellm.proxy.proxy_server.master_key", "sk-test-salt-for-lit3637"), ): - result = await register_client(request=mock_request, mcp_server_name=None) + response = await register_client(request=mock_request, mcp_server_name=None) - assert result["client_id"] == "dummy_client" - assert result["redirect_uris"] == ["https://llm.example.com/callback"] + body = json.loads(response.body) + assert body["client_id"].startswith("llm_dcrc_") + assert "test_oauth" not in body["client_id"] finally: global_mcp_server_manager.registry.clear() @@ -5155,7 +5162,10 @@ async def test_bridge_refresh_grant_with_non_envelope_is_invalid_grant_before_up def _mint_test_refresh_envelope( - server_id="bridge_srv", key_hash="hashed-litellm-key-77", upstream_refresh="UPSTREAM-REFRESH", identity=None, + server_id="bridge_srv", + key_hash="hashed-litellm-key-77", + upstream_refresh="UPSTREAM-REFRESH", + identity=None, scope=None, ): """Mint a refresh envelope the way the producer does, for driving the refresh_token grant in tests. @@ -5178,7 +5188,9 @@ def _mint_test_refresh_envelope( keys = envelope_keys_from_master_key(_BRIDGE_MASTER_KEY) identity = identity if identity is not None else key_hash_identity(server_id=server_id, key_hash=key_hash) sealed = build_bridge_refresh_token_response( - identity, RefreshCredential(refresh_token=SecretStr(upstream_refresh), scope=scope), keys, + identity, + RefreshCredential(refresh_token=SecretStr(upstream_refresh), scope=scope), + keys, datetime.now(timezone.utc), ) assert isinstance(sealed, SealedEnvelope) @@ -5413,7 +5425,10 @@ async def test_bridge_refresh_re_requests_the_sealed_scope_when_client_omits_it( ) captured: dict = {} response = await _refresh_for_bridge_server( - server, refresh_env, {"access_token": "NEW-ACCESS", "token_type": "Bearer", "expires_in": 3600}, None, + server, + refresh_env, + {"access_token": "NEW-ACCESS", "token_type": "Bearer", "expires_in": 3600}, + None, fake_client_out=captured, ) @@ -5553,7 +5568,9 @@ async def test_bridge_refresh_upstream_invalid_grant_maps_to_invalid_grant(): error_response = MagicMock() error_response.status_code = 400 error_response.text = '{"error": "invalid_grant", "error_description": "refresh token expired"}' - error_response.json = MagicMock(return_value={"error": "invalid_grant", "error_description": "refresh token expired"}) + error_response.json = MagicMock( + return_value={"error": "invalid_grant", "error_description": "refresh token expired"} + ) error_response.raise_for_status = MagicMock( side_effect=httpx.HTTPStatusError("bad", request=MagicMock(), response=error_response) ) @@ -7183,7 +7200,9 @@ def _upstream_token_response(status_code: int, *, json_body: object = None, text return httpx.Response(status_code, text=text_body, request=request) -async def _exchange_with_upstream_response(upstream_response, *, server_client_id="web-client.apps.googleusercontent.com"): +async def _exchange_with_upstream_response( + upstream_response, *, server_client_id="web-client.apps.googleusercontent.com" +): """Run the raw (non-bridge) authorization_code exchange against a canned upstream token-endpoint response and return what the gateway would hand the client. ``server_client_id=None`` models the caller-supplied-credentials flow (no stored client on the server).""" @@ -7334,9 +7353,7 @@ async def test_token_exchange_bounds_relayed_error_fields(): async def test_token_exchange_200_without_access_token_is_502_not_keyerror(): """A 200 whose body has no usable access_token used to KeyError into a 500; the raw arm now answers 502 with the same wording as the bridge arm's no_upstream_token rejection.""" - response = await _exchange_with_upstream_response( - _upstream_token_response(200, json_body={"token_type": "Bearer"}) - ) + response = await _exchange_with_upstream_response(_upstream_token_response(200, json_body={"token_type": "Bearer"})) assert response.status_code == 502 body = json.loads(response.body) @@ -7357,7 +7374,9 @@ async def test_token_exchange_relays_rejection_when_http_client_raises(): ) raising_client = MagicMock() raising_client.post = AsyncMock( - side_effect=httpx.HTTPStatusError("Client error '401 Unauthorized'", request=rejection.request, response=rejection) + side_effect=httpx.HTTPStatusError( + "Client error '401 Unauthorized'", request=rejection.request, response=rejection + ) ) from fastapi import Request @@ -7422,7 +7441,9 @@ async def test_register_relays_rejection_when_http_client_raises(): ) raising_client = MagicMock() raising_client.post = AsyncMock( - side_effect=httpx.HTTPStatusError("Client error '400 Bad Request'", request=rejection.request, response=rejection) + side_effect=httpx.HTTPStatusError( + "Client error '400 Bad Request'", request=rejection.request, response=rejection + ) ) oauth2_server = _bridge_server(auth_type=MCPAuth.oauth2, dcr_bridge=None) @@ -7800,9 +7821,7 @@ async def test_hydrate_does_not_overwrite_explicit_config_client_id(): auth_type=MCPAuth.oauth2, client_id="explicit-from-config", ) - store_read = AsyncMock( - return_value={"client_id": "stale-store-client", "client_secret": "x", "redirect_uris": []} - ) + store_read = AsyncMock(return_value={"client_id": "stale-store-client", "client_secret": "x", "redirect_uris": []}) with ( patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()), patch( @@ -8019,9 +8038,7 @@ async def test_bare_origin_discovery_resolves_single_server_not_aggregate(): mock_request.headers = {} try: - authorization_response = _build_oauth_authorization_server_response( - request=mock_request, mcp_server_name=None - ) + authorization_response = _build_oauth_authorization_server_response(request=mock_request, mcp_server_name=None) resource_response = await _build_oauth_protected_resource_response( request=mock_request, mcp_server_name=None, use_standard_pattern=True ) @@ -8033,6 +8050,65 @@ async def test_bare_origin_discovery_resolves_single_server_not_aggregate(): global_mcp_server_manager.registry.clear() +def test_gateway_dcr_flow_routing_engages_only_for_llm_dcrc_clients(monkeypatch): + """The aggregate DCR arms engage for llm_dcrc_ client_ids (register always mints one, + authorize/token route into the aggregate flow); a non-gateway client_id keeps the + per-server behavior, and /authorize/complete exists but 400s without a valid flow.""" + from fastapi import FastAPI + from fastapi.testclient import TestClient + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import router + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-for-lit3637") + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-test-salt-for-lit3637", raising=False) + global_mcp_server_manager.registry.clear() + app = FastAPI() + app.include_router(router) + client = TestClient(app) + + registered = client.post("/register", json={"redirect_uris": ["https://claude.ai/cb"]}) + assert registered.status_code == 201 + assert registered.json()["client_id"].startswith("llm_dcrc_") + assert registered.json()["token_endpoint_auth_method"] == "none" + + authorize_params = { + "client_id": "llm_dcrc_bogus", + "redirect_uri": "https://claude.ai/cb", + "response_type": "code", + "code_challenge": "c" * 43, + "code_challenge_method": "S256", + } + bogus_client = client.get("/authorize", params=authorize_params) + assert bogus_client.status_code == 400 + assert bogus_client.json()["error"] == "invalid_client" + + no_cookie = client.post("/authorize/complete", data={"flow": "h"}) + assert no_cookie.status_code == 400 + assert no_cookie.json()["error"] == "invalid_request" + + token_response = client.post( + "/token", + data={ + "grant_type": "authorization_code", + "client_id": "llm_dcrc_bogus", + "code": "x", + "redirect_uri": "https://claude.ai/cb", + "code_verifier": "v" * 43, + }, + ) + assert token_response.status_code == 400 + assert token_response.json()["error"] == "invalid_grant" + + upstream_shaped = client.post( + "/token", + data={"grant_type": "authorization_code", "client_id": "regular-upstream-client", "code": "x"}, + ) + assert upstream_shaped.status_code == 404 + + @pytest.mark.asyncio async def test_authorize_wall_names_the_fix_for_urlless_servers(): """LIT-4629: the authorize wall previously said only "authorization url is not set" with no @@ -8150,6 +8226,118 @@ async def test_register_wall_names_the_fix_for_urlless_servers(): assert "Issuer" in detail_text +@pytest.mark.asyncio +async def test_authorize_wall_points_at_discovery_failure_for_url_servers(): + """LIT-4658: a server WITH a url that still has no authorization_url got here because OAuth + discovery against that url failed (typically a misconfigured url); the old detail blamed + "servers with no url", sending the operator down the wrong path. The detail must now name the + discovery failure and point at the proxy logs where LIT-4658's warnings carry the reason.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + authorize_with_server, + ) + from litellm.types.mcp import MCPAuth, MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="typo-url-wall", + name="typo_wall", + server_name="typo_wall", + url="https://typo-host.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + ) + mock_request = MagicMock() + mock_request.base_url = "https://litellm.example.com/" + mock_request.headers = {} + + with pytest.raises(HTTPException) as exc_info: + await authorize_with_server( + request=mock_request, + mcp_server=server, + client_id="client", + redirect_uri="http://localhost/callback", + ) + assert exc_info.value.status_code == 400 + detail_text = str(exc_info.value.detail) + assert "may be misconfigured" in detail_text + assert "proxy logs" in detail_text + assert "Servers with no url" not in detail_text + assert "typo-host.example.com" not in detail_text + + +@pytest.mark.asyncio +async def test_token_wall_points_at_discovery_failure_for_url_servers(): + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + exchange_token_with_server, + ) + from litellm.types.mcp import MCPAuth, MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="typo-url-token-wall", + name="typo_token_wall", + server_name="typo_token_wall", + url="https://typo-host.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + authorization_url="https://idp.example.com/authorize", + ) + mock_request = MagicMock() + mock_request.base_url = "https://litellm.example.com/" + mock_request.headers = {} + + with pytest.raises(HTTPException) as exc_info: + await exchange_token_with_server( + request=mock_request, + mcp_server=server, + grant_type="authorization_code", + code="auth-code", + redirect_uri="http://localhost/callback", + client_id="client", + client_secret=None, + code_verifier="verifier", + ) + assert exc_info.value.status_code == 400 + detail_text = str(exc_info.value.detail) + assert "token url is not configured" in detail_text + assert "may be misconfigured" in detail_text + assert "Servers with no url" not in detail_text + + +@pytest.mark.asyncio +async def test_authorize_wall_names_the_issuer_for_anchored_servers(): + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + authorize_with_server, + ) + from litellm.types.mcp import MCPAuth, MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="anchored-wall", + name="anchored_wall", + server_name="anchored_wall", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + issuer="https://idp.example.com", + issuer_is_anchored=True, + ) + mock_request = MagicMock() + mock_request.base_url = "https://litellm.example.com/" + mock_request.headers = {} + + with pytest.raises(HTTPException) as exc_info: + await authorize_with_server( + request=mock_request, + mcp_server=server, + client_id="client", + redirect_uri="http://localhost/callback", + ) + assert exc_info.value.status_code == 400 + detail_text = str(exc_info.value.detail) + assert "verify the Issuer" in detail_text + assert "Servers with no url" not in detail_text + assert "idp.example.com" not in detail_text def test_passthrough_authorization_code_round_trips_and_rejects_hostile_input(): """The passthrough gateway code seals and recovers the ephemeral DCR client and upstream code, and is total over hostile input: a raw upstream code opens to None, and a tampered or diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py new file mode 100644 index 00000000000..375ec022115 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py @@ -0,0 +1,590 @@ +"""Tests for the aggregate gateway DCR flow (register, authorize, complete, token).""" + +import hashlib +import json +from base64 import urlsafe_b64encode +from datetime import datetime, timedelta, timezone +from http.cookies import SimpleCookie +from urllib.parse import parse_qs, urlparse + +import pytest +from starlette.requests import Request + +from litellm.caching.caching import DualCache +from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import ( + CONNECT_FLOW_COOKIE_PREFIX, + GATEWAY_AUTH_CODE_PREFIX, + GATEWAY_AUTH_CODE_TTL_SECONDS, + GATEWAY_DCR_CLIENT_ID_PREFIX, + _GatewayAuthCode, + _seal, + aggregate_authorize, + aggregate_token, + complete_connect_flow, + is_gateway_dcr_client_id, + open_gateway_dcr_client, + register_aggregate_client, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import ( + resolve_session_bearer, + session_keys_from_master_key, + SessionBearerAdmitted, +) + +MASTER_KEY = "sk-gateway-dcr-flow-tests" +REDIRECT_URI = "https://claude.ai/api/mcp/auth_callback" +CODE_VERIFIER = "verifier-" + "v" * 43 +CODE_CHALLENGE = urlsafe_b64encode(hashlib.sha256(CODE_VERIFIER.encode("ascii")).digest()).rstrip(b"=").decode("ascii") + + +@pytest.fixture(autouse=True) +def _salt_key(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", MASTER_KEY) + + +def _request(path="/authorize", query="", cookies=None, method="GET"): + cookie_header = [] + if cookies: + cookie = SimpleCookie() + for name, value in cookies.items(): + cookie[name] = value + cookie_header = [(b"cookie", cookie.output(header="", sep="; ").strip().encode())] + return Request( + { + "type": "http", + "method": method, + "scheme": "https", + "path": path, + "query_string": query.encode(), + "headers": [(b"host", b"llm.example.com"), *cookie_header], + } + ) + + +async def _register(redirect_uris) -> dict: + response = await register_aggregate_client( + request=_request(path="/register", method="POST"), request_body={"redirect_uris": redirect_uris} + ) + return json.loads(response.body) + + +async def _reload_user_active(user_id: str): + return None + + +@pytest.mark.asyncio +async def test_register_mints_stateless_public_client(): + body = await _register([REDIRECT_URI]) + assert body["token_endpoint_auth_method"] == "none" + assert "client_secret" not in body + assert body["redirect_uris"] == [REDIRECT_URI] + assert is_gateway_dcr_client_id(body["client_id"]) + record = open_gateway_dcr_client(body["client_id"]) + assert record is not None + assert record.redirect_uris == (REDIRECT_URI,) + + +@pytest.mark.asyncio +async def test_register_allows_loopback_http_for_dev_clients(): + body = await _register(["http://localhost:6274/oauth/callback"]) + assert is_gateway_dcr_client_id(body["client_id"]) + + +@pytest.mark.parametrize( + "code_challenge", + ["short", "", "p" * 300, "ünïcode-challenge", "AAAA" * 20], +) +def test_pkce_mismatched_challenge_returns_false_never_raises(code_challenge): + """A wrong-length or non-ASCII code_challenge must VERIFY FALSE, not raise. + + Pins the reason this compares bytes rather than str: hmac.compare_digest raises TypeError on + two str with non-ASCII content, but on bytes of unequal length it simply returns False. A + review flagged this as an unhandled 500 on length mismatch; encoding both sides to bytes is + exactly what makes that impossible, so the claim is pinned here rather than in a comment.""" + from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import _pkce_verifier_matches + + assert _pkce_verifier_matches("a" * 43, code_challenge) is False + + +@pytest.mark.asyncio +async def test_register_allows_allowlisted_native_callback(): + """Native MCP clients register a private-use scheme, not https. Registration shares + the one redirect-URI shape owner with /authorize, so the callback the allowlist + already trusts there is registrable here rather than rejected as non-https.""" + body = await _register(["cursor://anysphere.cursor-mcp/oauth/callback"]) + assert is_gateway_dcr_client_id(body["client_id"]) + record = open_gateway_dcr_client(body["client_id"]) + assert record is not None + assert record.redirect_uris == ("cursor://anysphere.cursor-mcp/oauth/callback",) + + +@pytest.mark.asyncio +async def test_register_rejects_userinfo_spoofed_origin(): + """``https://claude.ai@attacker.example/cb`` parses with netloc + ``claude.ai@attacker.example``, so a naive origin display on the consent screen reads + as claude.ai while the code would be delivered to attacker.example. Rejected at + registration, which is the only way such a URI could enter a sealed client.""" + response = await register_aggregate_client( + request=_request(path="/register", method="POST"), + request_body={"redirect_uris": ["https://claude.ai@attacker.example/callback"]}, + ) + assert response.status_code == 400 + assert json.loads(response.body)["error"] == "invalid_redirect_uri" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "redirect_uris", + [ + [], + "not-a-list", + ["http://evil.example.com/callback"], + ["https://claude.ai/cb#fragment"], + ["ftp://claude.ai/cb"], + ["https://a.example.com/" + "p" * 300], + ["https://a.example.com/1", "https://a.example.com/2", "https://a.example.com/3", "https://a.example.com/4"], + [12345], + ], +) +async def test_register_rejects_bad_redirect_uris(redirect_uris): + response = await register_aggregate_client( + request=_request(path="/register", method="POST"), request_body={"redirect_uris": redirect_uris} + ) + assert response.status_code == 400 + assert json.loads(response.body)["error"] in ("invalid_redirect_uri", "invalid_client_metadata") + + +@pytest.mark.asyncio +async def test_tampered_client_id_does_not_open(): + body = await _register([REDIRECT_URI]) + tampered = body["client_id"][:-4] + "AAAA" + assert open_gateway_dcr_client(tampered) is None + assert open_gateway_dcr_client("llm_dcrc_garbage") is None + assert open_gateway_dcr_client("other_prefix") is None + + +def _authorize( + client_id, session_user_id, redirect_uri=REDIRECT_URI, challenge=CODE_CHALLENGE, method="S256", response_type="code" +): + return aggregate_authorize( + request=_request(query=f"client_id={client_id}"), + client_id=client_id, + redirect_uri=redirect_uri, + state="client-state-123", + code_challenge=challenge, + code_challenge_method=method, + response_type=response_type, + session_user_id=session_user_id, + ) + + +@pytest.mark.asyncio +async def test_authorize_validation_failures_never_redirect_to_client(): + client_id = (await _register([REDIRECT_URI]))["client_id"] + for response, expected_error in ( + (_authorize("llm_dcrc_bogus", "u1"), "invalid_client"), + (_authorize(client_id, "u1", redirect_uri="https://attacker.example.com/cb"), "invalid_request"), + (_authorize(client_id, "u1", response_type="token"), "unsupported_response_type"), + (_authorize(client_id, "u1", challenge=None), "invalid_request"), + (_authorize(client_id, "u1", method="plain"), "invalid_request"), + ): + assert response.status_code == 400 + assert json.loads(response.body)["error"] == expected_error + + +@pytest.mark.asyncio +async def test_authorize_without_session_redirects_to_login_with_return_to(): + client_id = (await _register([REDIRECT_URI]))["client_id"] + response = _authorize(client_id, session_user_id=None) + assert response.status_code == 303 + location = response.headers["location"] + assert location.startswith("https://llm.example.com/sso/key/generate?return_to=") + assert "return_to=%2Fauthorize" in location + + +@pytest.mark.asyncio +async def test_authorize_with_session_hands_browser_to_connect_page_with_flow_cookie(): + client_id = (await _register([REDIRECT_URI]))["client_id"] + response = _authorize(client_id, session_user_id="u1") + assert response.status_code == 303 + location = urlparse(response.headers["location"]) + assert location.path == "/ui/chat/integrations" + params = parse_qs(location.query) + handle = params["connect_flow"][0] + assert params["connect_client"] == ["https://claude.ai"] + set_cookie = response.headers["set-cookie"] + assert f"{CONNECT_FLOW_COOKIE_PREFIX}{handle}" in set_cookie + assert "HttpOnly" in set_cookie + return handle, set_cookie + + +def _flow_cookie_from(response) -> tuple: + location = urlparse(response.headers["location"]) + handle = parse_qs(location.query)["connect_flow"][0] + cookie = SimpleCookie() + cookie.load(response.headers["set-cookie"]) + name = f"{CONNECT_FLOW_COOKIE_PREFIX}{handle}" + return handle, {name: cookie[name].value} + + +@pytest.mark.asyncio +async def test_full_walk_register_authorize_complete_token_and_replay(): + """The whole front door on one deterministic walk: register -> authorize -> + complete -> token, then the security edges on the same artifacts (user mismatch, + PKCE mismatch, single-use replay, refresh rotation, cross-client refresh).""" + client_id = (await _register([REDIRECT_URI]))["client_id"] + authorize_response = _authorize(client_id, session_user_id="u1") + handle, cookies = _flow_cookie_from(authorize_response) + + denied = await complete_connect_flow( + request=_request("/authorize/complete", cookies=cookies, method="POST"), + flow_handle=handle, + session_user_id="attacker", + cache=DualCache(), + ) + assert denied.status_code == 403 + + anonymous = await complete_connect_flow( + request=_request("/authorize/complete", cookies=cookies, method="POST"), + flow_handle=handle, + session_user_id=None, + cache=DualCache(), + ) + assert anonymous.status_code == 401 + + completed = await complete_connect_flow( + request=_request("/authorize/complete", cookies=cookies, method="POST"), + flow_handle=handle, + session_user_id="u1", + cache=DualCache(), + ) + assert completed.status_code == 303 + redirect = urlparse(completed.headers["location"]) + assert f"{redirect.scheme}://{redirect.netloc}{redirect.path}" == REDIRECT_URI + params = parse_qs(redirect.query) + assert params["state"] == ["client-state-123"] + code = params["code"][0] + assert code.startswith(GATEWAY_AUTH_CODE_PREFIX) + + cache = DualCache() + + async def _token(**overrides): + arguments = { + "request": _request("/token", method="POST"), + "grant_type": "authorization_code", + "code": code, + "redirect_uri": REDIRECT_URI, + "client_id": client_id, + "code_verifier": CODE_VERIFIER, + "refresh_token": None, + "master_key": MASTER_KEY, + "reload_user": _reload_user_active, + "cache": cache, + } + return await aggregate_token(**{**arguments, **overrides}) + + wrong_verifier = await _token(code_verifier="wrong-" + "w" * 43) + assert json.loads(wrong_verifier.body)["error"] == "invalid_grant" + + wrong_client = await _token(client_id=(await _register([REDIRECT_URI]))["client_id"]) + assert json.loads(wrong_client.body)["error"] == "invalid_grant" + + token_response = await _token() + assert token_response.status_code == 200 + payload = json.loads(token_response.body) + assert payload["token_type"] == "Bearer" + assert 0 < payload["expires_in"] <= 3600 + + keys = session_keys_from_master_key(MASTER_KEY) + admitted = resolve_session_bearer(f"Bearer {payload['access_token']}", keys, datetime.now(timezone.utc)) + assert isinstance(admitted, SessionBearerAdmitted) + assert admitted.principal.user_id == "u1" + assert admitted.principal.client_id == client_id + + replay = await _token() + assert json.loads(replay.body)["error"] == "invalid_grant" + + refreshed = await _token(grant_type="refresh_token", code=None, refresh_token=payload["refresh_token"]) + assert refreshed.status_code == 200 + rotated = json.loads(refreshed.body) + assert rotated["refresh_token"] != payload["refresh_token"] + + # Rotation is single-use: replaying the now-consumed refresh token cannot mint a second pair + # (a captured token is dead once the legitimate holder has rotated). + replayed = await _token(grant_type="refresh_token", code=None, refresh_token=payload["refresh_token"]) + assert json.loads(replayed.body)["error"] == "invalid_grant" + assert "already used" in json.loads(replayed.body).get("error_description", "") + + cross_client = await _token( + grant_type="refresh_token", + code=None, + refresh_token=payload["refresh_token"], + client_id=(await _register([REDIRECT_URI]))["client_id"], + ) + assert json.loads(cross_client.body)["error"] == "invalid_grant" + + +@pytest.mark.asyncio +async def test_complete_rejects_missing_tampered_and_expired_flows(): + missing = await complete_connect_flow( + request=_request("/authorize/complete", method="POST"), + flow_handle="nope", + session_user_id="u1", + cache=DualCache(), + ) + assert missing.status_code == 400 + + tampered = await complete_connect_flow( + request=_request("/authorize/complete", cookies={f"{CONNECT_FLOW_COOKIE_PREFIX}h1": "garbage"}, method="POST"), + flow_handle="h1", + session_user_id="u1", + cache=DualCache(), + ) + assert tampered.status_code == 400 + + +@pytest.mark.asyncio +async def test_token_rejects_expired_code_and_missing_configuration(): + expired_code = _seal( + GATEWAY_AUTH_CODE_PREFIX, + _GatewayAuthCode( + user_id="u1", + client_id="llm_dcrc_x", + redirect_uri=REDIRECT_URI, + code_challenge=CODE_CHALLENGE, + jti="jti-1", + iat=int((datetime.now(timezone.utc) - timedelta(seconds=500)).timestamp()), + exp=int((datetime.now(timezone.utc) - timedelta(seconds=500 - GATEWAY_AUTH_CODE_TTL_SECONDS)).timestamp()), + ), + ) + response = await aggregate_token( + request=_request("/token", method="POST"), + grant_type="authorization_code", + code=expired_code, + redirect_uri=REDIRECT_URI, + client_id="llm_dcrc_x", + code_verifier=CODE_VERIFIER, + refresh_token=None, + master_key=MASTER_KEY, + reload_user=_reload_user_active, + cache=DualCache(), + ) + assert json.loads(response.body)["error"] == "invalid_grant" + + no_master_key = await aggregate_token( + request=_request("/token", method="POST"), + grant_type="authorization_code", + code="llm_gcode_x", + redirect_uri=REDIRECT_URI, + client_id="llm_dcrc_x", + code_verifier=CODE_VERIFIER, + refresh_token=None, + master_key=None, + reload_user=_reload_user_active, + cache=DualCache(), + ) + assert no_master_key.status_code == 500 + assert json.loads(no_master_key.body)["error"] == "server_error" + + unsupported = await aggregate_token( + request=_request("/token", method="POST"), + grant_type="password", + code=None, + redirect_uri=None, + client_id="llm_dcrc_x", + code_verifier=None, + refresh_token=None, + master_key=MASTER_KEY, + reload_user=_reload_user_active, + cache=DualCache(), + ) + assert json.loads(unsupported.body)["error"] == "unsupported_grant_type" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "failure,expected_status,expected_error", + [ + ("no_active_key", 400, "invalid_grant"), + ("unavailable", 503, "temporarily_unavailable"), + ("unresolvable", 500, "server_error"), + ], +) +async def test_token_gates_on_live_user_revalidation(failure, expected_status, expected_error): + client_id = (await _register([REDIRECT_URI]))["client_id"] + authorize_response = _authorize(client_id, session_user_id="deactivated-user") + handle, cookies = _flow_cookie_from(authorize_response) + completed = await complete_connect_flow( + request=_request("/authorize/complete", cookies=cookies, method="POST"), + flow_handle=handle, + session_user_id="deactivated-user", + cache=DualCache(), + ) + code = parse_qs(urlparse(completed.headers["location"]).query)["code"][0] + + async def _reload_user_failing(user_id: str): + return failure + + response = await aggregate_token( + request=_request("/token", method="POST"), + grant_type="authorization_code", + code=code, + redirect_uri=REDIRECT_URI, + client_id=client_id, + code_verifier=CODE_VERIFIER, + refresh_token=None, + master_key=MASTER_KEY, + reload_user=_reload_user_failing, + cache=DualCache(), + ) + assert response.status_code == expected_status + assert json.loads(response.body)["error"] == expected_error + + +@pytest.mark.asyncio +async def test_flow_is_single_use_shared_cache_rejects_second_complete(): + """A double-submit of the finish step mints only ONE code: the second complete over the + same cache fails invalid_request (atomic flow claim), so one sign-in cannot yield two codes.""" + cache = DualCache() + client_id = (await _register([REDIRECT_URI]))["client_id"] + handle, cookies = _flow_cookie_from(_authorize(client_id, session_user_id="u1")) + + first = await complete_connect_flow( + request=_request("/authorize/complete", cookies=cookies, method="POST"), + flow_handle=handle, + session_user_id="u1", + cache=cache, + ) + assert first.status_code == 303 + second = await complete_connect_flow( + request=_request("/authorize/complete", cookies=cookies, method="POST"), + flow_handle=handle, + session_user_id="u1", + cache=cache, + ) + assert second.status_code == 400 + assert json.loads(second.body)["error"] == "invalid_request" + + +@pytest.mark.asyncio +async def test_token_rejects_out_of_range_code_verifier(): + """RFC 7636: a code_verifier outside 43-128 chars is invalid_request, not a confusing + invalid_grant PKCE-mismatch.""" + for bad in ["short", "x" * 200]: + response = await aggregate_token( + request=_request("/token", method="POST"), + grant_type="authorization_code", + code="llm_gcode_whatever", + redirect_uri=REDIRECT_URI, + client_id="llm_dcrc_x", + code_verifier=bad, + refresh_token=None, + master_key=MASTER_KEY, + reload_user=_reload_user_active, + cache=DualCache(), + ) + assert response.status_code == 400 + assert json.loads(response.body)["error"] == "invalid_request" + + +@pytest.mark.asyncio +async def test_authorize_rejects_over_long_state(): + client_id = (await _register([REDIRECT_URI]))["client_id"] + response = aggregate_authorize( + request=_request(query=f"client_id={client_id}"), + client_id=client_id, + redirect_uri=REDIRECT_URI, + state="s" * 2000, + code_challenge=CODE_CHALLENGE, + code_challenge_method="S256", + response_type="code", + session_user_id="u1", + ) + assert response.status_code == 400 + assert json.loads(response.body)["error"] == "invalid_request" + + +@pytest.mark.asyncio +async def test_non_ascii_code_challenge_fails_grant_not_500(): + """A non-ASCII code_challenge (unvalidated from the client) must yield a clean + invalid_grant, never a TypeError-driven 500 (bytes comparison, not str).""" + client_id = (await _register([REDIRECT_URI]))["client_id"] + # Seal a code carrying a non-ASCII challenge directly (authorize requires S256 shape, + # but the challenge charset is not validated there, so this state is reachable). + from datetime import datetime, timezone + + code = _seal( + GATEWAY_AUTH_CODE_PREFIX, + _GatewayAuthCode( + user_id="u1", + client_id=client_id, + redirect_uri=REDIRECT_URI, + code_challenge="challenge-with-€-non-ascii", + jti="jti-x", + iat=int(datetime.now(timezone.utc).timestamp()), + exp=int(datetime.now(timezone.utc).timestamp()) + 120, + ), + ) + response = await aggregate_token( + request=_request("/token", method="POST"), + grant_type="authorization_code", + code=code, + redirect_uri=REDIRECT_URI, + client_id=client_id, + code_verifier=CODE_VERIFIER, + refresh_token=None, + master_key=MASTER_KEY, + reload_user=_reload_user_active, + cache=DualCache(), + ) + assert response.status_code == 400 + assert json.loads(response.body)["error"] == "invalid_grant" + + +@pytest.mark.asyncio +async def test_single_use_guard_in_memory_is_single_use_within_process(): + """No Redis configured (single-replica): the in-memory increment is authoritative — the first claim + wins, a replay of the same id loses.""" + from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import _SingleUseGuard + + guard = _SingleUseGuard(DualCache()) # redis_cache is None + assert await guard.claim("jti-inmem", 60) is True + assert await guard.claim("jti-inmem", 60) is False # replay of the same id + + +@pytest.mark.asyncio +async def test_single_use_guard_uses_redis_as_sole_authority_when_configured(): + """With Redis configured it is the SOLE authority: the shared INCR result decides the claim (1 → + first caller, >1 → replay), and the per-worker in-memory count is never consulted.""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import _SingleUseGuard + + cache = DualCache() + cache.redis_cache = MagicMock() + cache.redis_cache.async_increment = AsyncMock(return_value=1) + # in-memory must NOT be consulted when Redis is configured — poison it so any fallback is visible. + cache.async_increment_cache = AsyncMock(side_effect=AssertionError("must not fall back to in-memory")) + + guard = _SingleUseGuard(cache) + assert await guard.claim("jti-redis", 60) is True + cache.redis_cache.async_increment = AsyncMock(return_value=2) + assert await guard.claim("jti-redis", 60) is False # Redis says 2 → replay + + +@pytest.mark.asyncio +async def test_single_use_guard_fails_closed_when_redis_errors(): + """A Redis fault must fail the claim CLOSED (refuse the id) rather than fall back to the per-worker + in-memory count — which would let each replica observe count==1 and replay the one-time id (the + Cursor/Veria replay-across-workers finding).""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import _SingleUseGuard + + cache = DualCache() + cache.redis_cache = MagicMock() + cache.redis_cache.async_increment = AsyncMock(side_effect=ConnectionError("redis down")) + cache.async_increment_cache = AsyncMock(return_value=1) # would fail OPEN if the guard fell back + + guard = _SingleUseGuard(cache) + assert await guard.claim("jti-fault", 60) is False # fail closed, not a fallback count of 1 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index ae4f12fc1e1..dff1f1d87c7 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -7375,6 +7375,10 @@ async def test_call_tool_with_legacy_db_m2m_server_resolves_oauth2_flow(): ("", None), ("not a url", None), ("http://[::1", None), + # urlsplit validates the port lazily on attribute access, so a malformed port must not + # raise out of the helper: the server loaders call it while warning about exactly this + # kind of typo'd url (LIT-4658) + ("https://example.com:bad/mcp", None), ], ) def test_redact_mcp_resource_url_strips_credentials(url, expected): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index a5cb16822cf..42b6cbee1c4 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -2333,6 +2333,59 @@ class TestMCPServerManager: assert emitted.headers["Authorization"] == "Bearer upstream-token" assert not kwargs["extra_headers"] or "authorization" not in {k.lower() for k in kwargs["extra_headers"]} + @pytest.mark.asyncio + async def test_create_mcp_client_token_exchange_never_falls_back_to_v1(self): + """A configured OBO server is owned end to end by the v2 token_exchange arm, even when the + caller supplies an x-mcp-* override. This is what makes the v1 OBO handler unreachable, so if + it ever defers to v1 again the deleted handler is silently needed back.""" + from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( + OAuthToken, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.resolver import ( + UpstreamCredentialProvider, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok + + class _StubExchanger: + def __init__(self): + self.subject_tokens = [] + + async def exchange(self, subject_token, server, config, *, tenant_id=""): + self.subject_tokens.append(subject_token) + return Ok(OAuthToken(access_token="exchanged-token")) + + async def invalidate(self, subject_token, server, config, *, tenant_id=""): + return None + + exchanger = _StubExchanger() + manager = MCPServerManager() + server = MCPServer( + server_id="obo-egress", + name="obo", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2_token_exchange, + client_id="gateway-client", + client_secret="gateway-secret", + token_exchange_endpoint="https://idp.example.com/oauth2/token", + ) + with ( + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.resolve_mcp_auth", + new_callable=AsyncMock, + ) as mock_resolve, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient") as mock_client_cls, + ): + await manager._create_mcp_client( + server=server, + mcp_auth_header="Bearer caller-override", + subject_token="eyJ-subject-token", + cred_provider=UpstreamCredentialProvider(token_exchanger=exchanger), + ) + mock_resolve.assert_not_awaited() + assert exchanger.subject_tokens == ["eyJ-subject-token"] + assert self._emitted_authorization(mock_client_cls) == "Bearer exchanged-token" + @staticmethod def _emitted_authorization(mock_client_cls) -> str: kwargs = mock_client_cls.call_args.kwargs @@ -3031,7 +3084,7 @@ class TestMCPServerManager: registration_url="https://discovered.example.com/register", ) - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True): + async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): assert server_url == "https://example.com/mcp" # oauth2 (browser flow) keeps the origin fallback; only OBO disables it. assert allow_origin_fallback is True @@ -5426,7 +5479,7 @@ class TestMCPServerTimestamps: manager = MCPServerManager() calls: list[bool] = [] - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True): + async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): calls.append(allow_origin_fallback) return MCPOAuthMetadata( scopes=None, @@ -5461,7 +5514,7 @@ class TestMCPServerTimestamps: manager = MCPServerManager() calls: list[str] = [] - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True): + async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): calls.append(server_url) raise AssertionError("discovery must not run when token_exchange_endpoint is configured") @@ -5491,7 +5544,7 @@ class TestMCPServerTimestamps: back to the row, so the next rebuild skips discovery instead of re-running it every time.""" manager = MCPServerManager() - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True): + async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): assert server_url == "https://example.com/mcp" assert allow_origin_fallback is False # OBO never guesses the origin return MCPOAuthMetadata( @@ -5602,7 +5655,7 @@ class TestMCPServerTimestamps: _dcr_bridge_relays_client_registration keys off that column.""" manager = MCPServerManager() - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True): + async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): assert allow_origin_fallback is True return MCPOAuthMetadata( scopes=["mcp.read", "mcp.write"], @@ -5817,7 +5870,7 @@ class TestMCPServerTimestamps: persist_discovered_endpoints=False neither the oauth2 nor the OBO write-back may fire.""" manager = MCPServerManager() - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True): + async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): return MCPOAuthMetadata( scopes=["s1"], authorization_url="https://idp.example.com/authorize", @@ -8539,7 +8592,7 @@ class TestOBOEndpointDiscovery: ) seen = [] - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True): + async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): seen.append((server_url, allow_origin_fallback)) return discovered @@ -8567,7 +8620,7 @@ class TestOBOEndpointDiscovery: async def test_config_obo_with_configured_endpoint_skips_discovery(self): manager = MCPServerManager() - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True): + async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): raise AssertionError("discovery must not run when the endpoint is configured") manager._descovery_metadata = fake_discovery # type: ignore[attr-defined] @@ -9028,3 +9081,162 @@ class TestUrllessIssuerDiscovery: anchored.assert_awaited_once_with("https://idp.example.com", None) resource_rooted.assert_not_awaited() assert built.token_url == "https://idp.example.com/token" + + +class TestDiscoveryFailureLogging: + """LIT-4658: a misconfigured MCP server url must be diagnosable from default-level server logs. + + Discovery failures used to die at debug level and the config-load path emitted no warning at + all, so the only operator-facing signal was the bare 400 at /authorize.""" + + def _connect_error_client(self, url: str) -> MagicMock: + client = MagicMock() + client.get = AsyncMock( + side_effect=httpx.ConnectError(f"[Errno 8] nodename nor servname provided for {url}") + ) + return client + + @pytest.mark.asyncio + async def test_descovery_metadata_warns_with_redacted_attempts_on_connect_error(self, caplog): + manager = MCPServerManager() + secret_url = "https://typo-host.example.com/mcp/s/PATHSECRET/mcp" + with ( + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client", + return_value=self._connect_error_client(secret_url), + ), + caplog.at_level(logging.WARNING, logger="LiteLLM"), + ): + result = await manager._descovery_metadata(secret_url, warn_when_no_metadata=True) + assert result is None + assert "found no authorization server metadata" in caplog.text + assert "ConnectError" in caplog.text + assert "https://typo-host.example.com" in caplog.text + # hosted MCP urls embed credentials in the path; neither the url nor the exception + # text may leak it into warning-level logs + assert "PATHSECRET" not in caplog.text + + @pytest.mark.asyncio + async def test_descovery_metadata_stays_silent_without_warn_flag(self, caplog): + manager = MCPServerManager() + url = "https://typo-host.example.com/mcp" + with ( + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client", + return_value=self._connect_error_client(url), + ), + caplog.at_level(logging.WARNING, logger="LiteLLM"), + ): + result = await manager._descovery_metadata(url) + assert result is None + assert "found no authorization server metadata" not in caplog.text + + @pytest.mark.asyncio + async def test_descovery_metadata_attempt_trail_names_each_failed_step(self, caplog): + manager = MCPServerManager() + url = "https://real-host.example.com/mcp-typo" + client = MagicMock() + client.get = AsyncMock( + return_value=httpx.Response(404, request=httpx.Request("GET", url)) + ) + with ( + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client", + return_value=client, + ), + caplog.at_level(logging.WARNING, logger="LiteLLM"), + ): + result = await manager._descovery_metadata(url, warn_when_no_metadata=True) + assert result is None + assert "HTTP 404" in caplog.text + assert "well-known protected-resource lookup found no authorization servers" in caplog.text + assert "origin fallback" in caplog.text + + @pytest.mark.asyncio + async def test_load_servers_from_config_warns_when_endpoints_unresolved(self, caplog): + manager = MCPServerManager() + manager._descovery_metadata = AsyncMock(return_value=None) # type: ignore[attr-defined] + config = { + "typo_server": { + "url": "https://typo.example.com/mcp", + "transport": MCPTransport.http, + "auth_type": MCPAuth.oauth2, + "oauth2_flow": "authorization_code", + } + } + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + await manager.load_servers_from_config(config) + assert "typo_server" in caplog.text + assert "authorization_url, token_url" in caplog.text + assert "unresolved" in caplog.text + assert "verify the configured server url" in caplog.text + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "extra_config", + [ + { + "authorization_url": "https://idp.example.com/auth", + "token_url": "https://idp.example.com/token", + }, + { + "oauth2_flow": "client_credentials", + "token_url": "https://idp.example.com/token", + "client_id": "cid", + "client_secret": "csec", + }, + ], + ) + async def test_load_servers_from_config_silent_when_flow_needs_covered(self, caplog, extra_config): + """Manually covered endpoints and M2M servers (which never need authorization_url) must not + warn on every reload; the warning is a misconfiguration signal, not discovery telemetry.""" + manager = MCPServerManager() + manager._descovery_metadata = AsyncMock(return_value=None) # type: ignore[attr-defined] + config = { + "covered_server": { + "url": "https://up.example.com/mcp", + "transport": MCPTransport.http, + "auth_type": MCPAuth.oauth2, + "oauth2_flow": "authorization_code", + **extra_config, + } + } + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + await manager.load_servers_from_config(config) + assert "unresolved" not in caplog.text + assert "no discovery source" not in caplog.text + + @pytest.mark.asyncio + async def test_config_server_without_discovery_source_warns_about_missing_endpoints(self, caplog): + manager = MCPServerManager() + manager._register_openapi_tools = AsyncMock() # type: ignore[attr-defined] + config = { + "spec_only": { + "spec_path": "https://example.com/openapi.yaml", + "transport": MCPTransport.http, + "auth_type": MCPAuth.oauth2, + "oauth2_flow": "authorization_code", + } + } + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + await manager.load_servers_from_config(config) + assert "no discovery source" in caplog.text + assert "authorization_url and token_url are not set manually" in caplog.text + + @pytest.mark.asyncio + async def test_db_build_warns_when_discovery_fails_for_oauth2_row(self, caplog): + manager = MCPServerManager() + manager._descovery_metadata = AsyncMock(return_value=None) # type: ignore[attr-defined] + record = LiteLLM_MCPServerTable( + server_id="typo-row-1", + server_name="typo_row", + url="https://typo.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + ) + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + await manager.build_mcp_server_from_table(record, credentials_are_encrypted=False) + assert "typo_row" in caplog.text + assert "authorization_url, token_url" in caplog.text + assert "unresolved" in caplog.text diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py index 1da44029b5c..0e442102e53 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py @@ -237,7 +237,7 @@ class TestListToolRestApiWithToolSearch: return_value={}, ), patch( - "litellm.proxy._experimental.mcp_server.rest_endpoints._get_oauth2_server_ids", + "litellm.proxy._experimental.mcp_server.rest_endpoints._v1_resolved_oauth2_server_ids", return_value=[], ), patch( @@ -316,7 +316,7 @@ class TestListToolRestApiWithToolSearch: return_value={}, ), patch( - "litellm.proxy._experimental.mcp_server.rest_endpoints._get_oauth2_server_ids", + "litellm.proxy._experimental.mcp_server.rest_endpoints._v1_resolved_oauth2_server_ids", return_value=[], ), patch( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index d4ba66c4381..5c9612a055e 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -2783,3 +2783,68 @@ class TestRestListToolsetFiltering: ) assert [tool.name for tool in result] == ["lookup_status"] + + +class TestV1ResolvedOauth2Gate: + """The REST surface must stop resolving per-user OAuth2 tokens for servers the v2 resolver owns. + + ``_resolve_v2_auth`` drops any Authorization built here for an ``authorization_code`` server and + injects the resolver's own token, so the v1 lookup was a DB round-trip whose result was discarded. + A server that still defers to v1 (upstream-delegated oauth2) must keep resolving, which is what + makes these assertions non-vacuous. + """ + + @staticmethod + def _oauth2_server(*, delegate_auth_to_upstream: bool) -> Any: + from litellm.proxy._experimental.mcp_server.server import MCPServer + from litellm.types.mcp import MCPTransport + + return MCPServer( + server_id="oauth2-srv", + name="oauth2-srv", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=delegate_auth_to_upstream, + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "delegate_auth_to_upstream, expected_headers, expected_lookups", + [ + (False, None, 0), + (True, {"Authorization": "Bearer stored-token"}, 1), + ], + ) + async def test_user_oauth_headers_skip_v2_owned_servers( + self, delegate_auth_to_upstream, expected_headers, expected_lookups, monkeypatch + ): + from litellm.proxy._experimental.mcp_server import db as mcp_db + + server = self._oauth2_server(delegate_auth_to_upstream=delegate_auth_to_upstream) + resolve_token = AsyncMock(return_value={"access_token": "stored-token"}) + monkeypatch.setattr(mcp_db, "resolve_valid_user_oauth_token", resolve_token) + + headers = await rest_endpoints._get_user_oauth_extra_headers( + server, + UserAPIKeyAuth(user_id="alice", api_key="sk-1234"), + prefetched_creds={"oauth2-srv": {"access_token": "stored-token"}}, + ) + + assert headers == expected_headers + assert resolve_token.await_count == expected_lookups + + def test_prefetch_preflight_only_counts_v1_resolved_servers(self, monkeypatch): + v2_owned = self._oauth2_server(delegate_auth_to_upstream=False) + v1_resolved = self._oauth2_server(delegate_auth_to_upstream=True) + v1_resolved.server_id = "delegate-srv" + registry = {"oauth2-srv": v2_owned, "delegate-srv": v1_resolved} + + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_mcp_server_by_id", + lambda server_id: registry.get(server_id), + ) + + assert rest_endpoints._v1_resolved_oauth2_server_ids(["oauth2-srv"]) == set() + assert rest_endpoints._v1_resolved_oauth2_server_ids(["oauth2-srv", "delegate-srv"]) == {"delegate-srv"} diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index 288e2533b72..c589014f276 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -176,9 +176,7 @@ async def test_authenticate_user_invalid_credentials(): mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) - with patch.dict( - os.environ, {"UI_USERNAME": ui_username, "UI_PASSWORD": "correct-password"} - ): + with patch.dict(os.environ, {"UI_USERNAME": ui_username, "UI_PASSWORD": "correct-password"}): with pytest.raises(ProxyException) as exc_info: await authenticate_user( username=ui_username, @@ -227,9 +225,7 @@ async def test_authenticate_user_wrong_password(): ) mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_usertable.find_first = AsyncMock( - return_value=mock_user - ) + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=mock_user) with patch.dict( os.environ, @@ -279,9 +275,7 @@ async def test_authenticate_user_email_case_insensitive_login(): return None mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_usertable.find_first = AsyncMock( - side_effect=mock_find_first - ) + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(side_effect=mock_find_first) with patch.dict( os.environ, @@ -334,9 +328,7 @@ async def test_authenticate_user_database_required_for_admin(): mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) - with patch.dict( - os.environ, {"UI_USERNAME": ui_username, "UI_PASSWORD": ui_password} - ): + with patch.dict(os.environ, {"UI_USERNAME": ui_username, "UI_PASSWORD": ui_password}): with patch( "litellm.proxy.auth.login_utils.user_update", new_callable=AsyncMock, @@ -429,9 +421,7 @@ def test_authenticate_user_non_ascii_direct_comparison(): assert result is True # And correctly returns False for different passwords - result = secrets.compare_digest( - password.encode("utf-8"), "different£pass".encode("utf-8") - ) + result = secrets.compare_digest(password.encode("utf-8"), "different£pass".encode("utf-8")) assert result is False @@ -531,9 +521,7 @@ async def test_authenticate_user_database_login_with_non_ascii_password(): return None mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_usertable.find_first = AsyncMock( - side_effect=mock_find_first - ) + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(side_effect=mock_find_first) with patch.dict( os.environ, @@ -559,3 +547,58 @@ async def test_authenticate_user_database_login_with_non_ascii_password(): assert isinstance(result, LoginResult) assert result.user_id == "test-user-123" assert result.user_email == user_email + + +class TestEncodeUiSessionJwt: + """The UI session cookie must carry a bounded exp so it does not stay + signature-valid until the master key rotates, and so the session-cookie readers + that require a bounded lifetime (the MCP interactive sign-in) accept it.""" + + def _decode(self, token: str) -> dict: + import jwt + + return jwt.decode(token, "sk-master-for-tests", algorithms=["HS256"]) + + def test_encoded_cookie_carries_bounded_exp(self): + import time + + from litellm.proxy.auth.login_utils import encode_ui_session_jwt + + token_object = {"user_id": "u1", "key": "sk-abc", "login_method": "username_password"} + with patch("litellm.proxy.auth.login_utils.LITELLM_UI_SESSION_DURATION", "24h"): + token = encode_ui_session_jwt(token_object, "sk-master-for-tests") + claims = self._decode(token) + assert claims["user_id"] == "u1" + assert claims["login_method"] == "username_password" + remaining = claims["exp"] - int(time.time()) + assert 23 * 3600 < remaining <= 24 * 3600 + + def test_duration_is_honored_from_env(self): + import time + + from litellm.proxy.auth.login_utils import encode_ui_session_jwt + + with patch("litellm.proxy.auth.login_utils.LITELLM_UI_SESSION_DURATION", "1h"): + token = encode_ui_session_jwt({"user_id": "u1"}, "sk-master-for-tests") + remaining = self._decode(token)["exp"] - int(time.time()) + assert 0 < remaining <= 3600 + + def test_cookie_is_accepted_by_the_exp_requiring_session_reader(self): + """The regression this change exists for: before it, the UI cookie carried no + exp and _user_id_from_session_cookie (require=["exp"]) rejected every real login, + so the MCP interactive sign-in could never capture identity. A cookie minted by + this helper must now be accepted.""" + from unittest.mock import MagicMock + + from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( + _user_id_from_session_cookie, + ) + from litellm.proxy.auth.login_utils import encode_ui_session_jwt + + token_object = {"user_id": "cornell-user", "key": "sk-abc", "login_method": "sso"} + with patch("litellm.proxy.auth.login_utils.LITELLM_UI_SESSION_DURATION", "24h"): + token = encode_ui_session_jwt(token_object, "sk-master-for-tests") + request = MagicMock() + request.cookies = {"token": token} + with patch("litellm.proxy.proxy_server.master_key", "sk-master-for-tests"): + assert _user_id_from_session_cookie(request) == "cornell-user" diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 2c1948adca1..0359d974d19 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -179,6 +179,45 @@ async def test_budget_reservation_runs_when_not_disabled(): assert user_api_key_auth_obj.budget_reservation == reservation +@pytest.mark.asyncio +@pytest.mark.parametrize( + "general_settings,expected_flag", + [ + ({"fail_closed_budget_enforcement": True}, True), + ({}, False), + ], +) +async def test_fail_closed_budget_enforcement_reaches_reservation( + general_settings, expected_flag +): + """#33923: the strict flag must be threaded into reserve_budget_for_request so a + failed reservation write can reject instead of failing open.""" + user_api_key_auth_obj = UserAPIKeyAuth(token="test_token") + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.reserve_budget_for_request", + new=AsyncMock(return_value=None), + ) as mock_reserve: + await _reserve_budget_after_common_checks( + user_api_key_auth_obj=user_api_key_auth_obj, + request_data={"model": "gpt-4o"}, + route="/v1/chat/completions", + llm_router=None, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + skip_budget_checks=False, + general_settings=general_settings, + ) + + assert ( + mock_reserve.await_args.kwargs["fail_closed_budget_enforcement"] + is expected_flag + ) + + @pytest.mark.asyncio async def test_should_not_reuse_cached_key_object_for_request_state(): key_cache = DualCache() @@ -1290,6 +1329,250 @@ async def test_scim_deactivated_user_key_is_rejected(): setattr(_proxy_server_mod, attr, val) +@pytest.mark.asyncio +async def test_cached_proxy_admin_key_sets_via_virtual_key_marker(): + """Cached PROXY_ADMIN auth objects early-return before the marked DB and + master-key returns, and cache serialization drops the exclude=True marker; + the cache-hit boundary must restore it or cached admin traffic silently + bypasses overwrite_user_with_key_hash stamping.""" + from fastapi import Request + from starlette.datastructures import URL + + from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.proxy_server import hash_token + + api_key = "sk-cached-admin-marker-test" + hashed_key = hash_token(api_key) + + cached_token = UserAPIKeyAuth( + api_key=api_key, + token=hashed_key, + user_id="cached-admin-user", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + assert cached_token.via_virtual_key is False + + mock_cache = AsyncMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.delete_cache = MagicMock() + + mock_proxy_logging_obj = MagicMock() + mock_proxy_logging_obj.internal_usage_cache = MagicMock() + mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock() + mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = ( + AsyncMock() + ) + mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) + + import litellm.proxy.proxy_server as _proxy_server_mod + + _attrs_to_set = { + "prisma_client": MagicMock(), + "user_api_key_cache": mock_cache, + "proxy_logging_obj": mock_proxy_logging_obj, + "master_key": "sk-master-key", + "general_settings": {}, + "llm_model_list": [], + "llm_router": None, + "open_telemetry_logger": None, + "model_max_budget_limiter": MagicMock(), + "user_custom_auth": None, + "jwt_handler": None, + "litellm_proxy_admin_name": "admin", + } + _original_values = { + attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set + } + try: + for attr, val in _attrs_to_set.items(): + setattr(_proxy_server_mod, attr, val) + + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + + with patch( + "litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key", + new_callable=AsyncMock, + return_value=cached_token, + ): + result = await _user_api_key_auth_builder( + request=request, + api_key=f"Bearer {api_key}", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={}, + ) + + assert isinstance(result, UserAPIKeyAuth) + assert result.user_role == LitellmUserRoles.PROXY_ADMIN + assert result.via_virtual_key is True + assert result.api_key == hashed_key + finally: + for attr, val in _original_values.items(): + setattr(_proxy_server_mod, attr, val) + + +@pytest.mark.asyncio +async def test_master_key_auth_sets_via_virtual_key_marker(): + """Master-key requests must also be stamped by overwrite_user_with_key_hash; + the auth path substitutes the stable alias for api_key and must mark the + result as proxy-validated.""" + from fastapi import Request + from starlette.datastructures import URL + + from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS + from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + + master_key = "sk-master-key" + + mock_cache = AsyncMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.delete_cache = MagicMock() + + mock_proxy_logging_obj = MagicMock() + mock_proxy_logging_obj.internal_usage_cache = MagicMock() + mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock() + mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = ( + AsyncMock() + ) + mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) + + import litellm.proxy.proxy_server as _proxy_server_mod + + _attrs_to_set = { + "prisma_client": MagicMock(), + "user_api_key_cache": mock_cache, + "proxy_logging_obj": mock_proxy_logging_obj, + "master_key": master_key, + "general_settings": {}, + "llm_model_list": [], + "llm_router": None, + "open_telemetry_logger": None, + "model_max_budget_limiter": MagicMock(), + "user_custom_auth": None, + "jwt_handler": None, + "litellm_proxy_admin_name": "admin", + } + _original_values = { + attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set + } + try: + for attr, val in _attrs_to_set.items(): + setattr(_proxy_server_mod, attr, val) + + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + + result = await _user_api_key_auth_builder( + request=request, + api_key=f"Bearer {master_key}", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={}, + ) + + assert isinstance(result, UserAPIKeyAuth) + assert result.via_virtual_key is True + assert result.api_key == LITELLM_PROXY_MASTER_KEY_ALIAS + finally: + for attr, val in _original_values.items(): + setattr(_proxy_server_mod, attr, val) + + +@pytest.mark.asyncio +async def test_db_virtual_key_auth_sets_via_virtual_key_marker(): + """via_virtual_key gates overwrite_user_with_key_hash stamping and is + forge-stripped from validated input, so the DB auth path setting it by + post-construction assignment is the only thing that turns stamping on.""" + from fastapi import Request + from starlette.datastructures import URL + + from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.proxy_server import hash_token + + api_key = "sk-via-virtual-key-marker-test" + hashed_key = hash_token(api_key) + + valid_token = UserAPIKeyAuth( + api_key=api_key, + token=hashed_key, + user_id="marker-test-user", + ) + + mock_cache = AsyncMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.delete_cache = MagicMock() + + mock_proxy_logging_obj = MagicMock() + mock_proxy_logging_obj.internal_usage_cache = MagicMock() + mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock() + mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = ( + AsyncMock() + ) + mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) + + mock_prisma_client = MagicMock() + + import litellm.proxy.proxy_server as _proxy_server_mod + + _attrs_to_set = { + "prisma_client": mock_prisma_client, + "user_api_key_cache": mock_cache, + "proxy_logging_obj": mock_proxy_logging_obj, + "master_key": "sk-master-key", + "general_settings": {}, + "llm_model_list": [], + "llm_router": None, + "open_telemetry_logger": None, + "model_max_budget_limiter": MagicMock(), + "user_custom_auth": None, + "jwt_handler": None, + "litellm_proxy_admin_name": "admin", + } + _original_values = { + attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set + } + try: + for attr, val in _attrs_to_set.items(): + setattr(_proxy_server_mod, attr, val) + + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + + with ( + patch( + "litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key", + new_callable=AsyncMock, + return_value=valid_token, + ), + patch( + "litellm.proxy.auth.user_api_key_auth.get_user_object", + new_callable=AsyncMock, + return_value=None, + ), + ): + result = await _user_api_key_auth_builder( + request=request, + api_key=f"Bearer {api_key}", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={}, + ) + + assert isinstance(result, UserAPIKeyAuth) + assert result.via_virtual_key is True + assert result.api_key == hashed_key + finally: + for attr, val in _original_values.items(): + setattr(_proxy_server_mod, attr, val) + + @pytest.mark.asyncio async def test_return_user_api_key_auth_obj_user_spend_and_budget(): """ diff --git a/tests/test_litellm/proxy/management_endpoints/test_tool_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_tool_management_endpoints.py index cf80ee5dee5..351f125052d 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_tool_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_tool_management_endpoints.py @@ -13,12 +13,17 @@ from datetime import datetime, timezone from typing import Optional from unittest.mock import AsyncMock, MagicMock, patch +import pytest from fastapi import FastAPI from fastapi.testclient import TestClient sys.path.insert(0, os.path.abspath("../../..")) -from litellm.proxy.management_endpoints.tool_management_endpoints import router +from litellm.proxy.management_endpoints.tool_management_endpoints import ( + _build_tool_spend_response, + _ToolSpendRow, + router, +) from litellm.types.tool_management import LiteLLM_ToolTableRow # --- helpers --- @@ -50,9 +55,9 @@ def _make_app() -> FastAPI: # Stub the auth dependency so we don't need a real proxy running. def _override_auth(): - from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth - return UserAPIKeyAuth(api_key="sk-test", user_id="admin") + return UserAPIKeyAuth(api_key="sk-test", user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) # A real (non-None) prisma stub for truthiness checks. @@ -147,3 +152,117 @@ class TestToolManagementEndpoints: json={"tool_name": "my_tool", "input_policy": "invalid_value"}, ) assert resp.status_code == 422 + + def test_tool_spend_route_not_shadowed_by_get_tool(self): + prisma = MagicMock() + prisma.db.query_raw = AsyncMock(return_value=[]) + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + resp = self.client.get("/v1/tool/spend") + assert resp.status_code == 200 + assert resp.json()["by_tool"] == [] + + def test_tool_spend_aggregates_and_sorts(self): + rows = [ + {"date": "2026-07-01", "tool_name": "search", "call_count": 2, "spend": 1.0, "total_tokens": 100}, + {"date": "2026-07-02", "tool_name": "search", "call_count": 1, "spend": 4.0, "total_tokens": 50}, + {"date": "2026-07-01", "tool_name": "read_file", "call_count": 3, "spend": 2.0, "total_tokens": 300}, + ] + prisma = MagicMock() + prisma.db.query_raw = AsyncMock(side_effect=[rows, [{"total_spend": 5.5}]]) + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + resp = self.client.get("/v1/tool/spend?start_date=2026-07-01&end_date=2026-07-02") + assert resp.status_code == 200 + body = resp.json() + assert [t["tool_name"] for t in body["by_tool"]] == ["search", "read_file"] + search = body["by_tool"][0] + assert search["spend"] == 5.0 + assert search["call_count"] == 3 + assert search["total_tokens"] == 150 + assert len(body["daily"]) == 3 + assert body["start_date"] == "2026-07-01" + assert body["end_date"] == "2026-07-02" + assert body["total_spend"] == 5.5 + + @patch("litellm.proxy.proxy_server.prisma_client", None) + def test_tool_spend_no_db_returns_500(self): + resp = self.client.get("/v1/tool/spend") + assert resp.status_code == 500 + + def test_tool_spend_end_date_is_inclusive_via_exclusive_next_day_bound(self): + prisma = MagicMock() + prisma.db.query_raw = AsyncMock(return_value=[]) + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + resp = self.client.get("/v1/tool/spend?start_date=2026-07-01&end_date=2026-07-02") + assert resp.status_code == 200 + expected_binds = ( + datetime(2026, 7, 1, tzinfo=timezone.utc).isoformat(), + datetime(2026, 7, 3, tzinfo=timezone.utc).isoformat(), + ) + assert prisma.db.query_raw.await_count == 2 + for call in prisma.db.query_raw.await_args_list: + assert tuple(call.args[1:]) == expected_binds + assert resp.json()["end_date"] == "2026-07-02" + + @pytest.mark.parametrize( + "query", + [ + "start_date=not-a-date", + "start_date=2026-02-30", + "start_date=07/01/2026", + "end_date=2026-13-01", + "end_date=20260701", + ], + ) + def test_tool_spend_malformed_date_returns_400(self, query: str): + prisma = MagicMock() + prisma.db.query_raw = AsyncMock(return_value=[]) + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + resp = self.client.get(f"/v1/tool/spend?{query}") + assert resp.status_code == 400 + assert "Invalid date format" in resp.json()["detail"] + prisma.db.query_raw.assert_not_awaited() + + def test_tool_spend_non_admin_returns_403(self): + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + app = _make_app() + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + api_key="sk-user", user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER + ) + client = TestClient(app, raise_server_exceptions=True) + prisma = MagicMock() + prisma.db.query_raw = AsyncMock(return_value=[]) + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + resp = client.get("/v1/tool/spend") + assert resp.status_code == 403 + prisma.db.query_raw.assert_not_awaited() + + +def _spend_row(date: str, tool_name: str, spend: float, call_count: int = 1, total_tokens: int = 10) -> _ToolSpendRow: + return _ToolSpendRow(date=date, tool_name=tool_name, call_count=call_count, spend=spend, total_tokens=total_tokens) + + +class TestBuildToolSpendResponse: + def test_multi_tool_attribution_double_counts_per_tool_but_not_total(self): + rows = [ + _spend_row("2026-07-01", "a", spend=3.0), + _spend_row("2026-07-01", "b", spend=3.0), + ] + resp = _build_tool_spend_response(rows, total_spend=3.0, start_date="2026-07-01", end_date="2026-07-01") + by_tool = {t.tool_name: t.spend for t in resp.by_tool} + assert by_tool == {"a": 3.0, "b": 3.0} + assert resp.total_spend == 3.0 + + def test_groups_across_days_and_sorts_by_spend(self): + rows = [ + _spend_row("2026-07-01", "b", spend=1.0, call_count=2, total_tokens=100), + _spend_row("2026-07-02", "b", spend=4.0, call_count=1, total_tokens=50), + _spend_row("2026-07-01", "a", spend=2.0, call_count=3, total_tokens=300), + ] + resp = _build_tool_spend_response(rows, total_spend=7.0, start_date="2026-07-01", end_date="2026-07-02") + assert [(t.tool_name, t.spend, t.call_count, t.total_tokens) for t in resp.by_tool] == [ + ("b", 5.0, 3, 150), + ("a", 2.0, 3, 300), + ] + assert len(resp.daily) == 3 diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index c693017e134..63a47428780 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -7763,3 +7763,80 @@ async def test_cli_completion_persists_assertion_under_db_user_id(): retain_mock.assert_awaited_once_with(user_id="cli-user-id", assertion=assertion) assert response.status_code == 200 + + +class TestSameOriginReturnPath: + """The same-origin relative return_to arm added for the MCP gateway DCR authorize + round-trip: only strictly relative paths qualify, so login can never redirect the + browser off the gateway origin.""" + + def test_accepts_relative_paths(self): + from litellm.proxy.management_endpoints.ui_sso import _is_same_origin_return_path + + assert _is_same_origin_return_path("/authorize?client_id=llm_dcrc_x&state=s") is True + assert _is_same_origin_return_path("/some_server/authorize") is True + + def test_rejects_absolute_protocol_relative_and_backslash_paths(self): + from litellm.proxy.management_endpoints.ui_sso import _is_same_origin_return_path + + assert _is_same_origin_return_path("https://evil.example.com/authorize") is False + assert _is_same_origin_return_path("//evil.example.com/authorize") is False + assert _is_same_origin_return_path("/\\evil.example.com") is False + assert _is_same_origin_return_path("javascript:alert(1)") is False + assert _is_same_origin_return_path("") is False + + +class TestPersistReturnToCookieSharedHelper: + """The single shared return_to helper used by EVERY sign-in branch (SSO / Okta / generic AND the + username/password form). It must be best-effort and NEVER raise — a bad return_to can never block + sign-in. Regression: the password form previously 400'd because it called _validate_return_to + directly (which raises for a non-matching absolute return_to when control_plane_url is set).""" + + @staticmethod + def _cookie(resp) -> str: + return resp.headers.get("set-cookie", "") + + def test_sets_cookie_for_same_origin_relative_path(self, monkeypatch): + from fastapi import Response + + from litellm.proxy.management_endpoints.ui_sso import _persist_return_to_cookie + + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {}) + resp = Response() + _persist_return_to_cookie(resp, "/mcp/authorize?client_id=llm_dcrc_abc") + assert "litellm_cp_return_to=" in self._cookie(resp) + + def test_bad_absolute_with_control_plane_configured_does_not_raise_and_is_not_stored(self, monkeypatch): + """THE regression: a non-matching absolute return_to with control_plane_url set must NOT raise + (it did, blocking the login form) and must NOT be stored — sign-in proceeds.""" + from fastapi import Response + + from litellm.proxy.management_endpoints.ui_sso import _persist_return_to_cookie + + monkeypatch.setattr( + "litellm.proxy.proxy_server.general_settings", {"control_plane_url": "https://cp.example.com"} + ) + resp = Response() + _persist_return_to_cookie(resp, "https://evil.example.com/steal") # must not raise + assert "litellm_cp_return_to=" not in self._cookie(resp) + + def test_none_return_to_is_a_noop(self): + from fastapi import Response + + from litellm.proxy.management_endpoints.ui_sso import _persist_return_to_cookie + + resp = Response() + _persist_return_to_cookie(resp, None) + assert "litellm_cp_return_to=" not in self._cookie(resp) + + def test_control_plane_matching_absolute_is_stored(self, monkeypatch): + from fastapi import Response + + from litellm.proxy.management_endpoints.ui_sso import _persist_return_to_cookie + + monkeypatch.setattr( + "litellm.proxy.proxy_server.general_settings", {"control_plane_url": "https://cp.example.com"} + ) + resp = Response() + _persist_return_to_cookie(resp, "https://cp.example.com/ui?page=models") + assert "litellm_cp_return_to=" in self._cookie(resp) diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py b/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py index f0250bbe1a6..a75d5bd5730 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py @@ -49,9 +49,7 @@ def _install_login_mocks(monkeypatch, raise_on_auth: bool = False) -> None: } monkeypatch.setattr("litellm.proxy.auth.login_utils.authenticate_user", _fake_auth) - monkeypatch.setattr( - "litellm.proxy.auth.login_utils.create_ui_token_object", _fake_token_object - ) + monkeypatch.setattr("litellm.proxy.auth.login_utils.create_ui_token_object", _fake_token_object) monkeypatch.setattr(ps, "master_key", "sk-test-master") monkeypatch.setattr(ps, "general_settings", {}) monkeypatch.setattr(ps, "premium_user", False) @@ -69,9 +67,7 @@ def test_fallback_login_returns_html_form(client, monkeypatch): body_lower = response.text.lower() shape = { "status": response.status_code, - "content_type_html": response.headers.get("content-type", "").startswith( - "text/html" - ), + "content_type_html": response.headers.get("content-type", "").startswith("text/html"), "has_form": "", "token": ""} + assert normalize(response.json(), volatile=frozenset({"token", "redirect_url"})) == { + "redirect_url": "", + "token": "", + } body = response.json() set_cookie = response.headers.get("set-cookie", "") shape = { "redirect_url_has_ui": "/ui/" in body.get("redirect_url", ""), - "redirect_url_has_login_success": "login=success" - in body.get("redirect_url", ""), + "redirect_url_has_login_success": "login=success" in body.get("redirect_url", ""), "token_in_body": bool(body.get("token")), "token_cookie_set": "token=" in set_cookie, } @@ -264,9 +254,7 @@ def test_v3_login_success_returns_code(client, monkeypatch): from litellm.proxy import proxy_server as ps _install_login_mocks(monkeypatch) - monkeypatch.setattr( - ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"} - ) + monkeypatch.setattr(ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"}) # Force the local (non-redis) cache path monkeypatch.setattr(ps, "redis_usage_cache", None) fake_cache = MagicMock() @@ -301,9 +289,7 @@ def test_v3_login_authenticate_failure_500(client, monkeypatch): from litellm.proxy import proxy_server as ps _install_login_mocks(monkeypatch, raise_on_auth=True) - monkeypatch.setattr( - ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"} - ) + monkeypatch.setattr(ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"}) response = client.post( "/v3/login", @@ -337,9 +323,7 @@ def test_v3_login_exchange_missing_code_400(client, monkeypatch): """Error path: missing 'code' in body -> 400 with 'Missing' message.""" from litellm.proxy import proxy_server as ps - monkeypatch.setattr( - ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"} - ) + monkeypatch.setattr(ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"}) response = client.post("/v3/login/exchange", json={}) assert response.status_code == 400 @@ -352,9 +336,7 @@ def test_v3_login_exchange_invalid_code_401(client, monkeypatch): """Error path: code that isn't in cache -> 401 'Invalid or expired'.""" from litellm.proxy import proxy_server as ps - monkeypatch.setattr( - ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"} - ) + monkeypatch.setattr(ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"}) monkeypatch.setattr(ps, "redis_usage_cache", None) fake_cache = MagicMock() fake_cache.async_get_cache = AsyncMock(return_value=None) @@ -372,9 +354,7 @@ def test_v3_login_exchange_success_returns_token_and_redirect(client, monkeypatc """Pin: valid code -> JSON {token, redirect_url} + token cookie + cache deleted (single-use).""" from litellm.proxy import proxy_server as ps - monkeypatch.setattr( - ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"} - ) + monkeypatch.setattr(ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"}) monkeypatch.setattr(ps, "redis_usage_cache", None) cached_payload = { @@ -388,9 +368,10 @@ def test_v3_login_exchange_success_returns_token_and_redirect(client, monkeypatc response = client.post("/v3/login/exchange", json={"code": "valid-code"}) assert response.status_code == 200 - assert normalize( - response.json(), volatile=frozenset({"token", "redirect_url"}) - ) == {"token": "", "redirect_url": ""} + assert normalize(response.json(), volatile=frozenset({"token", "redirect_url"})) == { + "token": "", + "redirect_url": "", + } body = response.json() set_cookie = response.headers.get("set-cookie", "") shape = { @@ -405,3 +386,77 @@ def test_v3_login_exchange_success_returns_token_and_redirect(client, monkeypatc "token_cookie_set": True, "cache_deleted_once": True, } + + +def test_login_form_honors_same_origin_return_to_cookie(client, monkeypatch): + """The aggregate DCR connect flow preserves a same-origin return_to in the litellm_cp_return_to + cookie; /login must RESUME there after password sign-in instead of dead-ending at the dashboard.""" + _install_login_mocks(monkeypatch) + return_to = "/mcp/authorize?client_id=llm_dcrc_abc&response_type=code" + response = client.post( + "/login", + data={"username": "admin", "password": "password"}, + cookies={"litellm_cp_return_to": return_to}, + follow_redirects=False, + ) + assert response.status_code == 303 + assert response.headers.get("location", "") == return_to # resumed the connect flow, not the dashboard + assert "token=" in response.headers.get("set-cookie", "") + + +def test_login_form_honors_control_plane_return_to_cookie(client, monkeypatch): + """/login resumes through the SAME resumer the SSO callback uses, so it honors BOTH shapes + _persist_return_to_cookie is willing to store. Honoring only the relative one silently dropped + a control-plane return_to and landed the user on the dashboard.""" + import litellm.proxy.proxy_server as ps + + _install_login_mocks(monkeypatch) + monkeypatch.setitem(ps.general_settings, "control_plane_url", "https://cp.example.com") + response = client.post( + "/login", + data={"username": "admin", "password": "password"}, + cookies={"litellm_cp_return_to": "https://cp.example.com/console"}, + follow_redirects=False, + ) + location = response.headers.get("location", "") + assert response.status_code == 303 + assert location.startswith("https://cp.example.com/console") + # Cross-origin arm hands the JWT off via a one-time code rather than a cookie. + assert "code=" in location and "login=success" in location + assert "token=" not in response.headers.get("set-cookie", "") + + +def test_login_form_survives_stale_control_plane_return_to(client, monkeypatch): + """A stale one-shot cookie must NEVER fail a completed sign-in. The resumer rejects a return_to + that no longer matches control_plane_url (a config change between the cookie's write and this + read); the user has already authenticated, so land on the dashboard instead of erroring.""" + import litellm.proxy.proxy_server as ps + + _install_login_mocks(monkeypatch) + monkeypatch.setitem(ps.general_settings, "control_plane_url", "https://new-cp.example.com") + response = client.post( + "/login", + data={"username": "admin", "password": "password"}, + cookies={"litellm_cp_return_to": "https://old-cp.example.com/console"}, + follow_redirects=False, + ) + assert response.status_code == 303, "login must not break on a stale return_to cookie" + location = response.headers.get("location", "") + assert "old-cp.example.com" not in location + assert "/ui/" in location + + +def test_login_form_ignores_open_redirect_return_to(client, monkeypatch): + """A non-same-origin return_to (open-redirect attempt) is rejected — /login falls back to the + dashboard rather than honoring an absolute/foreign URL.""" + _install_login_mocks(monkeypatch) + response = client.post( + "/login", + data={"username": "admin", "password": "password"}, + cookies={"litellm_cp_return_to": "https://evil.example.com/steal"}, + follow_redirects=False, + ) + assert response.status_code == 303 + location = response.headers.get("location", "") + assert "evil.example.com" not in location + assert "/ui/" in location # dashboard fallback diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index db72a7fb38c..67945436987 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -1628,6 +1628,226 @@ async def test_ui_view_spend_logs_date_range_filter(client, monkeypatch): assert data["data"][0]["id"] == "log2" +@pytest.mark.asyncio +async def test_ui_view_spend_logs_request_id_lookup_ignores_date_window( + client, monkeypatch +): + """ + LIT-3981: a request_id lookup on the UI route resolves across all time even + when the caller sends a date window that excludes the log (the dashboard + always sends a window). The window is dropped and request_id alone scopes + the query. Pre-fix the window was always applied, so an id from an older + page returned nothing. + """ + today = datetime.datetime.now(timezone.utc) + mock_spend_logs = [ + { + "id": "log_old", + "request_id": "req-old", + "api_key": "sk-test-key", + "user": "test_user_1", + "team_id": "team1", + "spend": 0.05, + "startTime": (today - datetime.timedelta(days=90)).isoformat(), + "model": "gpt-4", + }, + ] + + captured: dict = {} + + def filter_fn(where): + captured["where"] = where + rows = _filter_logs_by_date_range(mock_spend_logs, where) + if where.get("request_id"): + rows = [r for r in rows if r["request_id"] == where["request_id"]] + return rows + + monkeypatch.setattr( + "litellm.proxy.proxy_server.prisma_client", + make_ui_spend_logs_mock_prisma(mock_spend_logs, filter_fn), + ) + + # A 5-day window that EXCLUDES the 90-day-old log, as the dashboard sends. + start_date = (today - datetime.timedelta(days=5)).strftime("%Y-%m-%d %H:%M:%S") + end_date = today.strftime("%Y-%m-%d %H:%M:%S") + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN + ) + try: + response = client.get( + "/spend/logs/ui", + params={ + "request_id": "req-old", + "start_date": start_date, + "end_date": end_date, + }, + headers={"Authorization": "Bearer sk-test"}, + ) + assert response.status_code == 200 + data = response.json() + assert data["total"] == 1 + assert data["data"][0]["request_id"] == "req-old" + # Query dropped the time window and scoped solely by the primary key. + assert "startTime" not in captured["where"] + assert captured["where"]["request_id"] == "req-old" + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_ui_view_spend_logs_requires_dates_without_request_id( + client, monkeypatch +): + """The date window stays mandatory on the UI route when no request_id is set.""" + monkeypatch.setattr( + "litellm.proxy.proxy_server.prisma_client", + make_ui_spend_logs_mock_prisma([], lambda where: []), + ) + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN + ) + try: + response = client.get( + "/spend/logs/ui", headers={"Authorization": "Bearer sk-test"} + ) + assert response.status_code == 400 + assert "date" in response.text.lower() + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_spend_logs_v2_still_requires_dates_with_request_id(client, monkeypatch): + """The public /spend/logs/v2 contract is unchanged: dates remain required even + when request_id is supplied. Only the internal UI route relaxes the window.""" + monkeypatch.setattr( + "litellm.proxy.proxy_server.prisma_client", + make_ui_spend_logs_mock_prisma([], lambda where: []), + ) + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN + ) + try: + response = client.get( + "/spend/logs/v2", + params={"request_id": "req-old"}, + headers={"Authorization": "Bearer sk-test"}, + ) + assert response.status_code == 400 + assert "date" in response.text.lower() + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_ui_view_spend_logs_request_id_blocks_non_owner(client, monkeypatch): + """A non-admin looking up a request_id they do not own is rejected (403), so + the relaxed date window cannot read another tenant's log by id.""" + + class _ForeignRow: + user = "other_user" + team_id = None + + class _SpendLogs: + async def find_unique(self, where, include=None): + return _ForeignRow() + + class _DB: + def __init__(self): + self.litellm_spendlogs = _SpendLogs() + + class _Prisma: + def __init__(self): + self.db = _DB() + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", _Prisma()) + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, user_id="user_1" + ) + try: + response = client.get( + "/spend/logs/ui", + params={"request_id": "foreign-req"}, + headers={"Authorization": "Bearer sk-test"}, + ) + assert response.status_code == 403 + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_ui_view_spend_logs_request_id_owner_scoped_by_id_only( + client, monkeypatch +): + """A non-admin owner looking up their own request_id resolves across all time. + The ownership check authorizes the single row, so the query drops both the date + window and the general user/team scoping and filters by the primary key alone; + without that skip an internal user would have a `user`/`OR` clause added.""" + today = datetime.datetime.now(timezone.utc) + mock_spend_logs = [ + { + "id": "log_old", + "request_id": "req-old", + "api_key": "sk-test-key", + "user": "user_1", + "team_id": "team1", + "spend": 0.05, + "startTime": (today - datetime.timedelta(days=90)).isoformat(), + "model": "gpt-4", + }, + ] + + captured: dict = {} + + def filter_fn(where): + captured["where"] = where + rows = _filter_logs_by_date_range(mock_spend_logs, where) + if where.get("request_id"): + rows = [r for r in rows if r["request_id"] == where["request_id"]] + return rows + + mock_prisma = make_ui_spend_logs_mock_prisma(mock_spend_logs, filter_fn) + + class _OwnedRow: + user = "user_1" + team_id = "team1" + + async def _find_unique(where, include=None): + return _OwnedRow() + + mock_prisma.db.find_unique = _find_unique + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) + + # A 5-day window that EXCLUDES the 90-day-old log, as the dashboard sends. + start_date = (today - datetime.timedelta(days=5)).strftime("%Y-%m-%d %H:%M:%S") + end_date = today.strftime("%Y-%m-%d %H:%M:%S") + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, user_id="user_1" + ) + try: + response = client.get( + "/spend/logs/ui", + params={ + "request_id": "req-old", + "start_date": start_date, + "end_date": end_date, + }, + headers={"Authorization": "Bearer sk-test"}, + ) + assert response.status_code == 200 + data = response.json() + assert data["total"] == 1 + assert data["data"][0]["request_id"] == "req-old" + assert "startTime" not in captured["where"] + assert captured["where"]["request_id"] == "req-old" + assert "user" not in captured["where"] + assert "OR" not in captured["where"] + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + @pytest.mark.asyncio async def test_ui_view_spend_logs_unauthorized(client): # Test without authorization header diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py index 5d10bb33751..cc1e2943c8f 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py @@ -109,6 +109,143 @@ def test_get_logging_payload_does_not_map_missing_or_zero_cached_tokens(prompt_t assert "cache_read_input_tokens" not in additional_usage_values +def test_get_logging_payload_maps_openai_cache_write_tokens_to_cache_creation_input_tokens(): + additional_usage_values = _get_additional_usage_values_for_usage( + litellm.Usage( + prompt_tokens=1000, + completion_tokens=2, + total_tokens=1002, + prompt_tokens_details={"cached_tokens": 0, "cache_write_tokens": 800}, + ) + ) + + assert additional_usage_values["cache_creation_input_tokens"] == 800 + assert additional_usage_values["prompt_tokens_details"]["cache_write_tokens"] == 800 + + +def test_get_logging_payload_preserves_anthropic_cache_creation_input_tokens(): + additional_usage_values = _get_additional_usage_values_for_usage( + litellm.Usage( + prompt_tokens=1000, + completion_tokens=2, + total_tokens=1002, + cache_creation_input_tokens=300, + ) + ) + + assert additional_usage_values["cache_creation_input_tokens"] == 300 + + +@pytest.mark.parametrize( + "prompt_tokens_details", + [None, {"cached_tokens": 100}, {"cached_tokens": 100, "cache_write_tokens": 0}], +) +def test_get_logging_payload_does_not_map_missing_or_zero_cache_write_tokens(prompt_tokens_details): + additional_usage_values = _get_additional_usage_values_for_usage( + litellm.Usage( + prompt_tokens=10, + completion_tokens=2, + total_tokens=12, + prompt_tokens_details=prompt_tokens_details, + ) + ) + + assert "cache_creation_input_tokens" not in additional_usage_values + + +def _make_standard_logging_payload_with_usage_object(usage_object: dict) -> StandardLoggingPayload: + return StandardLoggingPayload( + id="test-id-responses", + call_type="responses", + stream=False, + response_cost=0.02, + status="success", + total_tokens=1010, + prompt_tokens=1000, + completion_tokens=10, + startTime=1234567890.0, + endTime=1234567891.0, + completionStartTime=None, + model_map_information=StandardLoggingModelInformation(model_map_key="gpt-5.6", model_map_value=None), + model="gpt-5.6", + model_id="model-123", + model_group="openai", + custom_llm_provider="openai", + api_base="https://api.openai.com", + metadata=StandardLoggingMetadata( + user_api_key_hash="test_hash", + user_api_key_alias=None, + user_api_key_team_id=None, + user_api_key_org_id=None, + user_api_key_user_id=None, + user_api_key_team_alias=None, + spend_logs_metadata=None, + requester_ip_address=None, + requester_metadata=None, + user_api_key_end_user_id=None, + usage_object=usage_object, + ), + cache_hit=False, + cache_key=None, + saved_cache_cost=0.0, + request_tags=[], + end_user=None, + requester_ip_address=None, + messages=[], + response={}, + error_str=None, + model_parameters={}, + hidden_params=StandardLoggingHiddenParams( + model_id="model-123", + cache_key=None, + api_base="https://api.openai.com", + response_cost="0.02", + litellm_overhead_time_ms=None, + additional_headers=None, + batch_models=None, + litellm_model_name=None, + usage_object=None, + ), + ) + + +def test_get_logging_payload_maps_responses_api_cache_write_tokens_from_usage_object(): + """Responses API (/v1/responses) usage is not chat-Usage-shaped, so + additional_usage_values can't derive cache tokens from response_obj.usage. + The Admin UI Logs "Cache Creation Tokens" row reads + additional_usage_values.cache_creation_input_tokens, so it must be filled + from the normalized standard_logging usage_object (LIT-4633).""" + standard_logging_payload = _make_standard_logging_payload_with_usage_object( + usage_object={ + "prompt_tokens": 1000, + "completion_tokens": 10, + "total_tokens": 1010, + "prompt_tokens_details": {"cached_tokens": 0, "cache_write_tokens": 800, "cache_creation_tokens": 800}, + } + ) + payload = get_logging_payload( + kwargs={ + "model": "gpt-5.6", + "call_type": "responses", + "litellm_params": {"metadata": {"user_api_key": "test-key"}}, + "standard_logging_object": standard_logging_payload, + }, + response_obj={ + "id": "resp-test", + "usage": { + "input_tokens": 1000, + "output_tokens": 10, + "total_tokens": 1010, + "input_tokens_details": {"cached_tokens": 0, "cache_write_tokens": 800}, + }, + }, + start_time=datetime.datetime.now(timezone.utc), + end_time=datetime.datetime.now(timezone.utc), + ) + additional_usage_values = json.loads(payload["metadata"])["additional_usage_values"] + assert additional_usage_values["cache_creation_input_tokens"] == 800 + + def test_sanitize_request_body_for_spend_logs_payload_basic(): request_body = { "messages": [{"role": "user", "content": "Hello, how are you?"}], diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 1db76aed61d..e6ccbc579b5 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -3,6 +3,7 @@ from datetime import datetime, timedelta, timezone from unittest.mock import AsyncMock, MagicMock, patch import pytest +from fastapi import HTTPException import litellm from litellm.caching.dual_cache import DualCache @@ -1562,6 +1563,100 @@ async def test_should_skip_reservation_when_counter_increment_fails( ) +@pytest.mark.asyncio +async def test_should_raise_503_when_counter_increment_fails_and_fail_closed( + spend_counter_state, + monkeypatch, +): + """#33923: with fail_closed_budget_enforcement on, a failed reservation write + must reject instead of silently degrading to read-time-only enforcement.""" + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-reserve-fail-closed", + spend=0.0, + max_budget=1.0, + ) + + async def fail_increment_cache(*args, **kwargs): + raise RuntimeError("counter unavailable") + + monkeypatch.setattr(counter_cache, "async_increment_cache", fail_increment_cache) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.5, + ): + with pytest.raises(HTTPException) as exc_info: + await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + fail_closed_budget_enforcement=True, + ) + + assert exc_info.value.status_code == 503 + assert ( + counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-reserve-fail-closed" + ) + is None + ) + + +@pytest.mark.asyncio +async def test_fail_closed_releases_earlier_counters_before_503( + spend_counter_state, +): + """#33923: when a later counter's reservation write fails in strict mode, the + counters that already reserved must be released before the 503 propagates.""" + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-fail-closed-release", + spend=0.0, + max_budget=1.0, + budget_limits=[ + { + "budget_duration": "1h", + "max_budget": 1.0, + } + ], + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.5, + ): + with pytest.raises(HTTPException) as exc_info: + await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + fail_closed_budget_enforcement=True, + ) + + assert exc_info.value.status_code == 503 + assert ( + counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-fail-closed-release" + ) + == 0.0 + ) + + @pytest.mark.asyncio async def test_should_skip_reservation_when_counter_initialization_fails( spend_counter_state, diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 8bee7e9f33b..1437899f561 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -5226,3 +5226,198 @@ async def test_add_litellm_data_to_request_unions_metadata_tags_with_header_tags tags = updated["litellm_metadata"]["tags"] assert "header-tag" in tags assert "body-tag" in tags + + +def _make_chat_request_mock() -> MagicMock: + return _make_request_mock("/v1/chat/completions", {"Content-Type": "application/json"}) + + +@pytest.mark.asyncio +async def test_overwrite_user_with_key_hash_clobbers_caller_supplied_user(monkeypatch): + """The flag exists so providers can ban by a tamper-proof id; a caller-chosen + `user` must never survive, and the raw sk- key must never be forwarded.""" + from litellm.proxy._types import hash_token + + monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True) + + raw_key = "sk-overwrite-user-test-1234" + user_api_key_dict = UserAPIKeyAuth(api_key=raw_key) + user_api_key_dict.via_virtual_key = True + data = {"model": "gpt-4o", "user": "attacker-chosen-id"} + + updated_data = await add_litellm_data_to_request( + data=data, + request=_make_chat_request_mock(), + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated_data["user"] == hash_token(raw_key) + assert updated_data["user"] != "attacker-chosen-id" + assert raw_key not in updated_data["user"] + + +@pytest.mark.asyncio +async def test_overwrite_user_with_key_hash_sets_user_when_absent(monkeypatch): + from litellm.proxy._types import hash_token + + monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True) + + raw_key = "sk-overwrite-user-test-5678" + user_api_key_dict = UserAPIKeyAuth(api_key=raw_key) + user_api_key_dict.via_virtual_key = True + data = {"model": "gpt-4o"} + + updated_data = await add_litellm_data_to_request( + data=data, + request=_make_chat_request_mock(), + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated_data["user"] == hash_token(raw_key) + + +@pytest.mark.asyncio +async def test_overwrite_user_with_key_hash_disabled_preserves_caller_user(): + assert litellm.overwrite_user_with_key_hash is False + + user_api_key_dict = UserAPIKeyAuth(api_key="sk-overwrite-user-test-9999") + user_api_key_dict.via_virtual_key = True + data = {"model": "gpt-4o", "user": "caller-chosen-id"} + + updated_data = await add_litellm_data_to_request( + data=data, + request=_make_chat_request_mock(), + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated_data["user"] == "caller-chosen-id" + + +@pytest.mark.asyncio +async def test_overwrite_user_with_key_hash_skips_custom_auth_credential(monkeypatch): + """Custom-auth credentials are not sk-prefixed or JWTs, so UserAPIKeyAuth stores + them raw; the stamp must skip them entirely so auth material never leaks.""" + monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True) + + raw_credential = "my-custom-auth-credential-abc123" + user_api_key_dict = UserAPIKeyAuth(api_key=raw_credential) + assert user_api_key_dict.api_key == raw_credential + + updated_data = await add_litellm_data_to_request( + data={"model": "gpt-4o", "user": "caller-chosen-id"}, + request=_make_chat_request_mock(), + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated_data["user"] == "caller-chosen-id" + + +@pytest.mark.asyncio +async def test_overwrite_user_with_key_hash_skips_jwt_auth(monkeypatch): + """A hashed JWT rotates on every token re-issue, so it is useless as a stable + ban id; JWT-authenticated requests are not stamped.""" + from litellm.proxy._types import hash_token + + monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True) + + hashed_jwt = f"hashed-jwt-{hash_token('some-jwt-token')}" + user_api_key_dict = UserAPIKeyAuth(api_key=hashed_jwt) + + updated_data = await add_litellm_data_to_request( + data={"model": "gpt-4o", "user": "caller-chosen-id"}, + request=_make_chat_request_mock(), + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated_data["user"] == "caller-chosen-id" + + +@pytest.mark.asyncio +async def test_overwrite_user_with_key_hash_skips_hex_shaped_custom_credential(monkeypatch): + """A custom-auth credential that happens to be 64 hex chars is indistinguishable + from a key hash by shape alone; only the server-set via_virtual_key marker may + authorize stamping, so this raw credential must never be forwarded.""" + monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True) + + hex_shaped_credential = "a" * 64 + user_api_key_dict = UserAPIKeyAuth(api_key=hex_shaped_credential) + assert user_api_key_dict.api_key == hex_shaped_credential + assert user_api_key_dict.via_virtual_key is False + + updated_data = await add_litellm_data_to_request( + data={"model": "gpt-4o", "user": "caller-chosen-id"}, + request=_make_chat_request_mock(), + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated_data["user"] == "caller-chosen-id" + + +def test_via_virtual_key_cannot_be_forged_from_validated_input(): + from_kwargs = UserAPIKeyAuth(api_key="b" * 64, via_virtual_key=True) + assert from_kwargs.via_virtual_key is False + + from_dict = UserAPIKeyAuth.model_validate({"api_key": "b" * 64, "via_virtual_key": True}) + assert from_dict.via_virtual_key is False + + +@pytest.mark.asyncio +async def test_overwrite_user_with_key_hash_stamps_master_key_alias(monkeypatch): + """Master-key requests carry the stable alias instead of a hash (so the master + key never propagates anywhere); the alias is the stampable id for them.""" + from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS + + monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True) + + user_api_key_dict = UserAPIKeyAuth(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS) + user_api_key_dict.via_virtual_key = True + + updated_data = await add_litellm_data_to_request( + data={"model": "gpt-4o", "user": "attacker-chosen-id"}, + request=_make_chat_request_mock(), + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated_data["user"] == LITELLM_PROXY_MASTER_KEY_ALIAS + + +@pytest.mark.asyncio +async def test_overwrite_user_with_key_hash_rejects_alias_without_marker(monkeypatch): + from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS + + monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True) + + user_api_key_dict = UserAPIKeyAuth(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS) + assert user_api_key_dict.via_virtual_key is False + + updated_data = await add_litellm_data_to_request( + data={"model": "gpt-4o", "user": "caller-chosen-id"}, + request=_make_chat_request_mock(), + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated_data["user"] == "caller-chosen-id" diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 5b67780dc58..bad76864ca7 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -127,11 +127,15 @@ def test_login_v2_returns_redirect_url_and_sets_cookie(monkeypatch): general_settings={}, premium_user=False, ) - mock_jwt_encode.assert_called_once_with( - {"user_id": "test-user"}, - "test-master-key", - algorithm="HS256", - ) + mock_jwt_encode.assert_called_once() + payload, secret = mock_jwt_encode.call_args.args + # The UI session token carries a bounded-lifetime `exp` claim (dynamic timestamp), alongside + # the user_id; assert its presence rather than an exact expiry value. + assert payload["user_id"] == "test-user" + assert isinstance(payload.get("exp"), int) and payload["exp"] > 0 + assert set(payload.keys()) == {"user_id", "exp"} + assert secret == "test-master-key" + assert mock_jwt_encode.call_args.kwargs == {"algorithm": "HS256"} def test_login_v2_returns_json_on_proxy_exception(monkeypatch): diff --git a/tests/test_litellm/responses/test_responses_utils.py b/tests/test_litellm/responses/test_responses_utils.py index e4db8022f01..0141cf5d96a 100644 --- a/tests/test_litellm/responses/test_responses_utils.py +++ b/tests/test_litellm/responses/test_responses_utils.py @@ -369,6 +369,32 @@ class TestResponseAPILoggingUtils: assert result.completion_tokens_details.image_tokens == 272 assert result.completion_tokens_details.text_tokens == 100 + def test_transform_response_api_usage_maps_cache_write_tokens(self): + """Responses API (/v1/responses) cache-write tokens must survive the usage transform. + + gpt-5.6 returns usage.input_tokens_details.cache_write_tokens (an extra field + not typed on InputTokensDetails). Before the fix the transform rebuilt the token + details and dropped it, leaving the cache-creation metric empty (LIT-4633). + """ + usage = { + "input_tokens": 10062, + "output_tokens": 16, + "total_tokens": 10078, + "input_tokens_details": { + "cached_tokens": 0, + "cache_write_tokens": 10059, + }, + } + + result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( + usage + ) + + assert result.prompt_tokens_details is not None + assert result.prompt_tokens_details.cache_write_tokens == 10059 + assert result.prompt_tokens_details.cache_creation_tokens == 10059 + assert result.prompt_tokens_details.cached_tokens == 0 + def test_transform_response_api_usage_mixed_details(self): """Test transformation handles mixed token details (cached + image + audio).""" # Setup - hypothetical usage with mixed token types @@ -461,64 +487,6 @@ class TestResponseAPILoggingUtils: assert result.completion_tokens_details.text_tokens == 20 assert result.completion_tokens_details.audio_tokens is None - def test_transform_response_api_usage_maps_cache_write_tokens_dict(self): - """Regression for LIT-4725 / #33772: the Responses API reports cache writes under - input_tokens_details.cache_write_tokens. The chat-shaped usage must carry them on - cache_creation_tokens so cost is computed identically to /chat/completions. - model_construct keeps input_tokens_details a raw dict so the dict branch runs.""" - from litellm.types.llms.openai import ResponseAPIUsage - - usage = ResponseAPIUsage.model_construct( - input_tokens=10_000, - output_tokens=20, - total_tokens=10_020, - input_tokens_details={"cached_tokens": 2_000, "cache_write_tokens": 8_000}, - ) - - result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) - - assert result.prompt_tokens_details is not None - assert result.prompt_tokens_details.cached_tokens == 2_000 - assert result.prompt_tokens_details.cache_creation_tokens == 8_000 - assert getattr(result.prompt_tokens_details, "cache_write_tokens", None) is None - - def test_transform_response_api_usage_cache_creation_tokens_precedence_dict(self): - """When both cache_creation_tokens and cache_write_tokens are present, the explicit - cache_creation_tokens wins (they describe the same tokens under different names).""" - from litellm.types.llms.openai import ResponseAPIUsage - - usage = ResponseAPIUsage.model_construct( - input_tokens=10_000, - output_tokens=20, - total_tokens=10_020, - input_tokens_details={"cache_creation_tokens": 5_000, "cache_write_tokens": 8_000}, - ) - - result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) - - assert result.prompt_tokens_details is not None - assert result.prompt_tokens_details.cache_creation_tokens == 5_000 - - def test_transform_response_api_usage_maps_cache_write_tokens_object(self): - """Object-path counterpart: a ResponseAPIUsage whose input_tokens_details object - carries cache_write_tokens must still land on cache_creation_tokens.""" - from litellm.types.llms.openai import InputTokensDetails, ResponseAPIUsage - - input_tokens_details = InputTokensDetails(cached_tokens=2_000) - input_tokens_details.cache_write_tokens = 8_000 - usage = ResponseAPIUsage( - input_tokens=10_000, - output_tokens=20, - total_tokens=10_020, - input_tokens_details=input_tokens_details, - ) - - result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) - - assert result.prompt_tokens_details is not None - assert result.prompt_tokens_details.cached_tokens == 2_000 - assert result.prompt_tokens_details.cache_creation_tokens == 8_000 - class TestResponsesAPIProviderSpecificParams: """ diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 9636db4f4cd..276ee96ed65 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -12,6 +12,7 @@ from pydantic import BaseModel import litellm from litellm.cost_calculator import ( + BaseTokenUsageProcessor, RealtimeAPITokenUsageProcessor, completion_cost, cost_per_token, @@ -3479,3 +3480,32 @@ def test_batch_cost_calculator_cache_creation_falls_back_to_input_rate(): ) assert prompt_cost == pytest.approx((1000 * 3e-6 + 8000 * 3e-7 + 2000 * 3e-6) / 2) + + +def test_combine_usage_objects_sums_mirrored_cache_write_fields_once(): + """ + cache_write_tokens and cache_creation_tokens mirror each other on + PromptTokensDetailsWrapper, so field-iterating aggregation must sum the pair + once: a single 50-token usage stays 50 and two combine to 100, not double. + """ + single = Usage( + prompt_tokens=100, + completion_tokens=10, + total_tokens=110, + prompt_tokens_details=PromptTokensDetailsWrapper(cache_write_tokens=50), + ) + combined = BaseTokenUsageProcessor.combine_usage_objects([single]) + assert combined.prompt_tokens_details is not None + assert combined.prompt_tokens_details.cache_write_tokens == 50 + assert combined.prompt_tokens_details.cache_creation_tokens == 50 + + anthropic_style = Usage( + prompt_tokens=100, + completion_tokens=10, + total_tokens=110, + cache_creation_input_tokens=50, + ) + combined_pair = BaseTokenUsageProcessor.combine_usage_objects([anthropic_style, anthropic_style]) + assert combined_pair.prompt_tokens_details is not None + assert combined_pair.prompt_tokens_details.cache_write_tokens == 100 + assert combined_pair.prompt_tokens_details.cache_creation_tokens == 100 diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index a1a9448cc58..edc0cfed63e 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -17,7 +17,9 @@ from litellm.types.utils import ( Delta, LlmProviders, ModelResponseStream, + PromptTokensDetailsWrapper, StreamingChoices, + Usage, ) from litellm.utils import ( ProviderConfigManager, @@ -34,6 +36,57 @@ from litellm.utils import ( # Adds the parent directory to the system path +def test_usage_openai_cache_write_tokens_populates_both_names(): + """OpenAI reports cache-write tokens as prompt_tokens_details.cache_write_tokens. + The Usage constructor must expose it under both cache_write_tokens (canonical, + OpenAI naming) and cache_creation_tokens (legacy, Anthropic naming).""" + usage = Usage( + prompt_tokens=1000, + completion_tokens=10, + total_tokens=1010, + prompt_tokens_details={"cached_tokens": 0, "cache_write_tokens": 800}, + ) + assert usage.prompt_tokens_details.cache_write_tokens == 800 + assert usage.prompt_tokens_details.cache_creation_tokens == 800 + + +def test_usage_anthropic_cache_creation_maps_to_cache_write_tokens(): + """Anthropic/Bedrock report the top-level cache_creation_input_tokens field. + It must be normalized onto the OpenAI cache_write_tokens name as well as the + legacy cache_creation_tokens name.""" + usage = Usage( + prompt_tokens=500, + completion_tokens=50, + total_tokens=550, + cache_creation_input_tokens=300, + cache_read_input_tokens=120, + ) + assert usage.prompt_tokens_details.cache_write_tokens == 300 + assert usage.prompt_tokens_details.cache_creation_tokens == 300 + assert usage.prompt_tokens_details.cached_tokens == 120 + + +def test_prompt_tokens_details_no_cache_write_tokens_when_absent(): + """A read-only cache hit (no cache write) must not surface cache-write fields.""" + details = PromptTokensDetailsWrapper(cached_tokens=800) + assert details.cached_tokens == 800 + assert not hasattr(details, "cache_write_tokens") + assert not hasattr(details, "cache_creation_tokens") + + +def test_prompt_tokens_details_cache_write_creation_stay_in_sync_on_assignment(): + """Assigning either name after construction must mirror to the other, so a + caller that sets only one field can't leave the pair silently out of sync.""" + details = PromptTokensDetailsWrapper(cache_write_tokens=100) + assert details.cache_write_tokens == details.cache_creation_tokens == 100 + + details.cache_write_tokens = 250 + assert details.cache_write_tokens == details.cache_creation_tokens == 250 + + details.cache_creation_tokens = 375 + assert details.cache_write_tokens == details.cache_creation_tokens == 375 + + @pytest.fixture def local_model_cost_map(monkeypatch): original_model_cost = litellm.model_cost diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index 12f19eaa1fe..ec1e3ac05ba 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -37,16 +37,6 @@ "count": 1 } }, - "src/app/(dashboard)/agents/_components/AgentsPanel.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, - "src/app/(dashboard)/agents/_components/AgentsTable.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/app/(dashboard)/agents/_components/add_agent_form.tsx": { "local/filename-pascal-case": { "count": 1 @@ -71,9 +61,6 @@ "local/filename-pascal-case": { "count": 1 }, - "no-restricted-imports": { - "count": 1 - }, "react-hooks/refs": { "count": 3 }, @@ -84,9 +71,6 @@ "src/app/(dashboard)/agents/_components/agent_cost_view.tsx": { "local/filename-pascal-case": { "count": 1 - }, - "no-restricted-imports": { - "count": 2 } }, "src/app/(dashboard)/agents/_components/agent_form_fields.tsx": { @@ -123,9 +107,6 @@ }, "no-nested-ternary": { "count": 1 - }, - "no-restricted-imports": { - "count": 1 } }, "src/app/(dashboard)/agents/_components/cost_config_fields.tsx": { @@ -1143,25 +1124,6 @@ "count": 1 } }, - "src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.test.tsx": { - "max-params": { - "count": 1 - }, - "unused-imports/no-unused-imports": { - "count": 1 - } - }, - "src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx": { - "local/no-complex-jsx-arrow": { - "count": 2 - }, - "no-restricted-imports": { - "count": 2 - }, - "react-hooks/set-state-in-effect": { - "count": 3 - } - }, "src/app/(dashboard)/models-and-endpoints/components/ModelRetrySettingsTab.test.tsx": { "react/display-name": { "count": 1 @@ -1180,6 +1142,11 @@ "count": 1 } }, + "src/app/(dashboard)/models-and-endpoints/layout.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/app/(dashboard)/models-and-endpoints/utils/modelDataTransformer.ts": { "prefer-const": { "count": 6 @@ -2134,9 +2101,6 @@ "no-nested-ternary": { "count": 1 }, - "no-restricted-imports": { - "count": 1 - }, "no-restricted-syntax": { "count": 3 }, @@ -2371,11 +2335,6 @@ "count": 1 } }, - "src/components/KeyAliasSelect/PaginatedKeyAliasSelect/PaginatedKeyAliasSelect.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/components/LicenseExpiryBanner.tsx": { "no-restricted-imports": { "count": 1 @@ -2386,14 +2345,6 @@ "count": 1 } }, - "src/components/ModelSelect/PaginatedModelSelect/PaginatedModelSelect.tsx": { - "no-nested-ternary": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - } - }, "src/components/Navbar/BlogDropdown/BlogDropdown.test.tsx": { "max-nested-callbacks": { "count": 12 @@ -2972,7 +2923,7 @@ "count": 1 }, "no-nested-ternary": { - "count": 7 + "count": 6 } }, "src/components/chat/MCPConnectPicker.tsx": { @@ -3486,17 +3437,6 @@ "count": 1 } }, - "src/components/model_dashboard/all_models_table.tsx": { - "local/filename-pascal-case": { - "count": 1 - }, - "no-nested-ternary": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - } - }, "src/components/model_filters.tsx": { "local/filename-pascal-case": { "count": 1 @@ -3544,17 +3484,6 @@ "count": 1 } }, - "src/components/molecules/filter.tsx": { - "local/filename-pascal-case": { - "count": 1 - }, - "no-nested-ternary": { - "count": 2 - }, - "no-restricted-imports": { - "count": 1 - } - }, "src/components/molecules/message_manager.tsx": { "local/filename-pascal-case": { "count": 1 @@ -3563,28 +3492,6 @@ "count": 2 } }, - "src/components/molecules/models/columns.test.tsx": { - "no-restricted-imports": { - "count": 1 - }, - "react/display-name": { - "count": 1 - } - }, - "src/components/molecules/models/columns.tsx": { - "local/filename-pascal-case": { - "count": 1 - }, - "max-params": { - "count": 1 - }, - "no-nested-ternary": { - "count": 2 - }, - "no-restricted-imports": { - "count": 2 - } - }, "src/components/molecules/notifications_manager.test.tsx": { "no-restricted-imports": { "count": 1 @@ -3689,9 +3596,6 @@ }, "no-restricted-imports": { "count": 3 - }, - "unused-imports/no-unused-imports": { - "count": 1 } }, "src/components/page_utils.test.ts": { @@ -4182,6 +4086,11 @@ "count": 1 } }, + "src/components/ui/hover-card.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, "src/components/ui/input-group.tsx": { "local/filename-pascal-case": { "count": 1 @@ -4454,17 +4363,6 @@ "count": 2 } }, - "src/components/view_logs/LogsTableToolbar.tsx": { - "local/no-complex-jsx-arrow": { - "count": 1 - }, - "no-nested-ternary": { - "count": 4 - }, - "no-restricted-imports": { - "count": 1 - } - }, "src/components/view_logs/ToolsSection/FormattedToolView.tsx": { "no-restricted-imports": { "count": 1 @@ -4493,9 +4391,6 @@ "src/components/view_logs/columns.tsx": { "local/filename-pascal-case": { "count": 1 - }, - "no-restricted-imports": { - "count": 1 } }, "src/components/view_logs/index.tsx": { @@ -4504,9 +4399,6 @@ }, "no-restricted-imports": { "count": 1 - }, - "react-hooks/set-state-in-effect": { - "count": 1 } }, "src/components/view_logs/log_filter_logic.tsx": { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.test.tsx index 441d300436a..d873687378b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.test.tsx @@ -140,8 +140,9 @@ describe("AgentsPanel", () => { await user.click(await screen.findByTestId("agent-actions-agent-9")); await user.click(await screen.findByTestId("agent-action-delete")); - const modal = await screen.findByRole("dialog"); - await user.click(within(modal).getByRole("button", { name: /^delete$/i })); + const confirmPrompt = await screen.findByText(/are you sure you want to delete agent: Doomed Agent\?/i); + const confirmDialog = confirmPrompt.closest('[role="dialog"],[role="alertdialog"]') as HTMLElement; + await user.click(within(confirmDialog).getByRole("button", { name: /^delete$/i })); await waitFor(() => { expect(networking.deleteAgentCall).toHaveBeenCalledWith("test-token", "agent-9"); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.tsx index a4a71530c84..4459ee0c377 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.tsx @@ -1,6 +1,5 @@ import React, { useState, useEffect } from "react"; -import { Modal, Alert } from "antd"; -import { Plus } from "lucide-react"; +import { Info, Plus } from "lucide-react"; import { getAgentsList, deleteAgentCall } from "@/components/networking"; import AddAgentForm from "./add_agent_form"; import { isAdminRole } from "@/utils/roles"; @@ -9,6 +8,16 @@ import AgentsTable from "./AgentsTable"; import NotificationsManager from "@/components/molecules/notifications_manager"; import { Agent } from "@/components/agents/types"; import { Team } from "@/components/key_team_helpers/key_list"; +import { Alert, AlertDescription, AlertTitle } from "@/components/shared/Alert"; +import { + AlertDialog, + AlertDialogCancel, + AlertDialogContent, + AlertDialogDescription, + AlertDialogFooter, + AlertDialogHeader, + AlertDialogTitle, +} from "@/components/ui/alert-dialog"; import { Button } from "@/components/ui/button"; interface AgentsPanelProps { @@ -130,17 +139,18 @@ const AgentsPanel: React.FC = ({ accessToken, userRole, teams

Agents

-

+

List of A2A-spec agents that are available to be used in your organization. Go to AI Hub, to make agents public.

- + + + Why do agents need keys? + + Keys scope access to an agent and allow it to call MCP tools. Assign a key when creating an agent or from + the Virtual Keys page. + + {isAdmin && (
+ + + )}
); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTable.tsx index 824ae47f3e6..359cb49b910 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTable.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTable.tsx @@ -1,13 +1,13 @@ "use client"; import { SortingState } from "@tanstack/react-table"; -import { Tooltip, Switch } from "antd"; -import { CheckCircleOutlined } from "@ant-design/icons"; -import { Bot } from "lucide-react"; +import { Bot, CircleCheck } from "lucide-react"; import React, { useMemo, useState } from "react"; import { Agent } from "@/components/agents/types"; import { DataTable } from "@/components/shared/DataTable"; +import { Switch } from "@/components/ui/switch"; +import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip"; import { getAgentsTableColumns } from "./AgentsTableColumns"; @@ -67,18 +67,27 @@ const AgentsTable: React.FC = ({ size="compact" toolbar={() => (
- -
- - Health Check - + + + + Health Check + +
+ } /> -
- + When enabled, only agents with reachable URLs are shown + +
)} /> diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_card_discovery.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_card_discovery.test.tsx index 4ee6332c54e..7858bdb1cd4 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_card_discovery.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_card_discovery.test.tsx @@ -126,10 +126,7 @@ describe("AgentCardDiscovery", () => { expect(initialSelection.upstream_url).toBe("https://upstream.example.com"); expect(initialSelection.selected_card.skills).toHaveLength(2); - const summarizeLabel = screen.getByText("Summarize").closest("label"); - expect(summarizeLabel).toBeTruthy(); - const summarizeCheckbox = summarizeLabel!.querySelector("input[type='checkbox']") as HTMLInputElement; - await user.click(summarizeCheckbox); + await user.click(screen.getByRole("checkbox", { name: /Summarize/i })); await waitFor(() => { const latest = onApply.mock.calls.at(-1)?.[0]; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_card_discovery.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_card_discovery.tsx index e979b2dbe3f..017a9928f8b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_card_discovery.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_card_discovery.tsx @@ -1,18 +1,20 @@ "use client"; import React, { useCallback, useEffect, useMemo, useRef, useState } from "react"; -import { Alert, Button, Checkbox, Collapse, Empty, Input, Space, Spin, Switch, Tag, Tooltip, Typography } from "antd"; -// Empty is used in the skills panel below. -import { - CheckCircleTwoTone, - InfoCircleOutlined, - LinkOutlined, - ReloadOutlined, - SearchOutlined, -} from "@ant-design/icons"; +import { ChevronDown, CircleAlert, CircleCheck, Info, Link as LinkIcon, RotateCw, Search, X } from "lucide-react"; import { useDebouncedCallback } from "@tanstack/react-pacer/debouncer"; import { DiscoveredAgentCard, discoverAgentCardCall } from "@/components/networking"; +import { Alert, AlertAction, AlertDescription, AlertTitle } from "@/components/shared/Alert"; +import { Badge } from "@/components/ui/badge"; +import { Button } from "@/components/ui/button"; +import { Checkbox } from "@/components/ui/checkbox"; +import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible"; +import { Input } from "@/components/ui/input"; +import { Switch } from "@/components/ui/switch"; +import { Textarea } from "@/components/ui/textarea"; +import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip"; +import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; import { ALLOWED_CAPABILITY_KEYS, selectionsFromSavedAgentCard, @@ -20,9 +22,6 @@ import { skillId, } from "./agent_discovery_utils"; -const { Text, Paragraph } = Typography; -const { Panel } = Collapse; - const DISCOVERY_DEBOUNCE_WAIT_MS = 400; export interface DiscoveredAgentCardSelection { @@ -243,102 +242,115 @@ const AgentCardDiscovery: React.FC = ({ const skillCount = card?.skills?.length ?? 0; const selectedSkillCount = selectedSkillIds.size; + const renderDiscoverIcon = () => { + if (loading) return ; + if (card) return ; + return ; + }; + const discoverLabel = card ? "Re-discover" : "Discover"; + return ( -
-
- - Discover from agent URL - - - +
+
+ + Discover from agent URL + + + + + + } + /> + + LiteLLM will fetch /.well-known/agent-card.json from this URL and let you pick which skills and + capabilities to expose through the proxy. + + +
{isParentDriven ? ( <> - +

Using the connection details you entered above. We'll fetch: - -

+

+
{discoveryRequest!.display_url || effectiveUrl || ( - Fill in the fields above first + Fill in the fields above first )}
-
) : ( <> - +

Paste the upstream agent's base URL. We'll try /.well-known/agent-card.json,{" "} /.well-known/agent.json, and /agent.json in order. - +

- +
setManualUrl(e.target.value)} - onPressEnter={handleDiscover} - allowClear + onKeyDown={(e) => { + if (e.key === "Enter") handleDiscover(); + }} disabled={loading} /> - - +
)} {error && ( - setError(null)} - /> + + + Discovery failed + {error} + + + + )} {loading && !card && (
- +
)} {card && ( -
-
- - - Upstream card loaded - {card.version && v{card.version}} - {card.provider?.organization && {card.provider.organization}} - +
+
+ + Upstream card loaded + {card.version && v{card.version}} + {card.provider?.organization && {card.provider.organization}}
-
+
- + setEditedName(e.target.value)} placeholder="Agent name" />
- - Description +