From 35db36ee936643d36391eab1cb3c9f5031284b0d Mon Sep 17 00:00:00 2001 From: yassin Date: Wed, 2 Sep 2026 16:04:53 +0000 Subject: [PATCH] fix(mcp): satisfy type discipline lint gate Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/discoverable_endpoints.py | 4 +- .../mcp_server/oauth_identity_binding.py | 39 ++++++++++++------- .../mcp_management_endpoints.py | 6 +-- .../types/mcp_server/mcp_server_manager.py | 2 +- 4 files changed, 31 insertions(+), 20 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 88f2c47f5fc..5f54bbb4958 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -1048,13 +1048,13 @@ async def exchange_token_with_server( refresh_request_scope = scope or bridge_upstream_scope if refresh_request_scope: token_data["scope"] = refresh_request_scope - refresh_ownership = ( + refresh_ownership = ( # rebind-ok: grant-specific branches assign one ownership value RefreshOwnershipProven() if bridge_upstream_refresh is not None else RefreshTokenPresented(upstream_refresh_token) ) else: - refresh_ownership = None + refresh_ownership = None # rebind-ok: grant-specific branches assign one ownership value if not code: raise HTTPException( status_code=400, diff --git a/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py b/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py index d4339c15e1a..b73227f3413 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py @@ -8,12 +8,13 @@ claim is compared to the caller's trusted LiteLLM identity. Mismatches fail clos and are logged in audit mode. """ -from collections.abc import Awaitable, Callable, Mapping +from collections.abc import Awaitable, Callable, Mapping, Sequence from dataclasses import dataclass -from typing import Final, Literal +from typing import Final, Literal, TypeAlias import jwt from fastapi import HTTPException +from jwt.types import Options from typing_extensions import assert_never from litellm._logging import verbose_logger @@ -36,11 +37,20 @@ _ALLOWED_ID_TOKEN_ALGORITHMS: Final = ( _JWKS_CACHE_TTL_SECONDS: Final = 3600 _jwks_cache: Final = InMemoryCache(default_ttl=_JWKS_CACHE_TTL_SECONDS) -JwksFetcher = Callable[[MCPOAuthIdentityBinding], Awaitable[list[Mapping[str, object]]]] -CallerPrincipalLoader = Callable[[str, MCPOAuthIdentityBinding], Awaitable[str | None]] -StoredRefreshTokenLoader = Callable[[str, str], Awaitable[str | None]] +JwksFetcher: TypeAlias = Callable[ + [MCPOAuthIdentityBinding], # mutable-ok: Callable parameter syntax requires a list + Awaitable[Sequence[Mapping[str, object]]], +] +CallerPrincipalLoader: TypeAlias = Callable[ + [str, MCPOAuthIdentityBinding], # mutable-ok: Callable parameter syntax requires a list + Awaitable[str | None], +] +StoredRefreshTokenLoader: TypeAlias = Callable[ + [str, str], # mutable-ok: Callable parameter syntax requires a list + Awaitable[str | None], +] -_RejectionCode = Literal["oauth_principal_mismatch", "oauth_identity_binding_failed"] +_RejectionCode: TypeAlias = Literal["oauth_principal_mismatch", "oauth_identity_binding_failed"] @dataclass(frozen=True, slots=True) @@ -59,10 +69,10 @@ class RefreshTokenPresented: refresh_token: str -RefreshOwnership = RefreshOwnershipProven | RefreshTokenPresented | None +RefreshOwnership: TypeAlias = RefreshOwnershipProven | RefreshTokenPresented | None -async def _fetch_issuer_jwks(binding: MCPOAuthIdentityBinding) -> list[Mapping[str, object]]: +async def _fetch_issuer_jwks(binding: MCPOAuthIdentityBinding) -> Sequence[Mapping[str, object]]: jwks_url: Final[str] = binding.jwks_url or await _discover_jwks_url(binding.issuer) cached: Final = await _jwks_cache.async_get_cache(jwks_url) if isinstance(cached, list): @@ -90,12 +100,12 @@ async def _discover_jwks_url(issuer: str) -> str: return jwks_uri -def _select_signing_key(id_token: str, keys: list[Mapping[str, object]]) -> "jwt.PyJWK | _BindingRejection": +def _select_signing_key(id_token: str, keys: Sequence[Mapping[str, object]]) -> "jwt.PyJWK | _BindingRejection": header: Final = jwt.get_unverified_header(id_token) kid: Final = header.get("kid") for key in keys: if kid is None or key.get("kid") == kid: - return jwt.PyJWK(dict(key)) + return jwt.PyJWK(dict(key)) # mutable-ok: PyJWT requires a concrete JWK dictionary return _BindingRejection( code="oauth_identity_binding_failed", description=f"id_token signing key (kid={kid!r}) not found in the issuer's JWKS", @@ -108,13 +118,14 @@ def _decode_id_token( signing_key: "jwt.PyJWK", ) -> "Mapping[str, object] | _BindingRejection": try: + decode_options: Final[Options] = {"require": ("iss", "exp")} return jwt.decode( id_token, signing_key.key, - algorithms=list(_ALLOWED_ID_TOKEN_ALGORITHMS), + algorithms=_ALLOWED_ID_TOKEN_ALGORITHMS, issuer=binding.issuer, audience=binding.audiences, - options={"require": ["iss", "exp"]}, + options=decode_options, ) except jwt.InvalidTokenError as exc: return _BindingRejection( @@ -160,10 +171,10 @@ async def _load_caller_principal(litellm_user_id: str, binding: MCPOAuthIdentity async def _load_stored_refresh_token(litellm_user_id: str, server_id: str) -> str | None: try: - from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 + from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # keep database imports lazy get_user_oauth_credential, ) - from litellm.proxy.utils import get_prisma_client_or_throw # noqa: PLC0415 + from litellm.proxy.utils import get_prisma_client_or_throw # noqa: PLC0415 # keep database imports lazy prisma_client: Final = get_prisma_client_or_throw( "Database not connected. Cannot verify OAuth refresh token ownership." diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 92a6997abff..109030a2df5 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -259,7 +259,7 @@ if MCP_AVAILABLE: if normalize_upstream_header_name(raw) is None: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, - detail={ + detail={ # mutable-ok: FastAPI exception detail requires a JSON-serializable dictionary "error": ( f"Invalid upstream_token_header {raw!r}: must be a valid HTTP header name " "(RFC 7230 token, e.g. 'esb-oauth')" @@ -2256,7 +2256,7 @@ if MCP_AVAILABLE: """Persist the OAuth2 access token obtained by the calling user.""" prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") await _authorize_and_fetch_mcp_server(prisma_client, user_api_key_dict, server_id) - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # keep manager import lazy global_mcp_server_manager as _manager, ) @@ -2267,7 +2267,7 @@ if MCP_AVAILABLE: if binding is not None and binding.mode == "enforce": raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, - detail={ + detail={ # mutable-ok: FastAPI exception detail requires a JSON-serializable dictionary "error": "oauth_identity_binding_enforced", "error_description": ( "Direct credential storage is disabled for this server: its OAuth identity " diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index cebc7decac4..c38f1c366b1 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -52,7 +52,7 @@ class MCPOAuthIdentityBinding(BaseModel): mode: Literal["disabled", "audit", "enforce"] = "disabled" issuer: str jwks_url: str | None = None - audiences: list[str] = Field(min_length=1) + audiences: list[str] = Field(min_length=1) # mutable-ok: public Pydantic schema requires list values principal_claim: str = "email" caller_field: Literal["user_email", "user_id"] = "user_email" require_email_verified: bool = True