fix(mcp): satisfy type discipline lint gate

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yassin 2026-09-02 16:04:53 +00:00
parent 50874f9fc9
commit 35db36ee93
4 changed files with 31 additions and 20 deletions

View file

@ -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,

View file

@ -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."

View file

@ -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 "

View file

@ -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