mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
50874f9fc9
commit
35db36ee93
4 changed files with 31 additions and 20 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
|
|
|
|||
|
|
@ -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 "
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue