mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge 55210dd1af into 9068c2441d
This commit is contained in:
commit
ca3db98365
26 changed files with 4615 additions and 545 deletions
196
litellm/proxy/_experimental/mcp_server/caller_sign_in.py
Normal file
196
litellm/proxy/_experimental/mcp_server/caller_sign_in.py
Normal file
|
|
@ -0,0 +1,196 @@
|
|||
"""Caller-side sign-in requirements for MCP connects.
|
||||
|
||||
A guardrail that evaluates tool calls in the caller's own identity (an On-Behalf-Of exchange of the caller's
|
||||
bearer) needs the caller signed in with its issuer before the first tool call, and a tool call's JSON-RPC
|
||||
error cannot carry ``WWW-Authenticate``. ``token_exchange`` (OBO) servers have the same need: the caller must
|
||||
present a subject token the gateway can exchange. Both cases share one contract: a connect that carries no
|
||||
usable subject answers 401 with the RFC 9728 challenge, and the protected-resource metadata advertises the
|
||||
issuers and scopes the caller signs in for.
|
||||
|
||||
Guardrails implement :class:`CallerSignInProvider`; :func:`caller_sign_in_for` merges the OBO server's own
|
||||
requirement with every registered provider's so the challenge and the metadata always agree. The MCP package
|
||||
never imports a concrete provider.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import itertools
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Final, Protocol, runtime_checkable
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
|
||||
from typing_extensions import assert_never
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.types.mcp import MCPAuth
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CallerSignIn:
|
||||
"""The issuers a caller signs in with and the scopes it requests before calling a server."""
|
||||
|
||||
issuers: tuple[str, ...]
|
||||
scopes: tuple[str, ...]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SignedIn:
|
||||
"""The subject token the caller presented satisfies this provider's sign-in."""
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Rejected:
|
||||
"""The caller's identity provider rejected the presented subject token."""
|
||||
|
||||
detail: str
|
||||
claims: str | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Unavailable:
|
||||
"""The provider could not reach a verdict; ``fail_open`` is the provider's own fallback policy."""
|
||||
|
||||
detail: str
|
||||
fail_open: bool
|
||||
|
||||
|
||||
CallerSignInPreflight = SignedIn | Rejected | Unavailable
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class CallerSignInProvider(Protocol):
|
||||
def caller_sign_in(self, server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None) -> CallerSignIn | None:
|
||||
"""The sign-in this provider requires of callers hitting ``server``; ``None`` when it does not gate
|
||||
the server for this caller (``user_api_key_auth=None`` is the anonymous metadata fetch that follows
|
||||
a challenge)."""
|
||||
...
|
||||
|
||||
async def preflight_caller_sign_in(
|
||||
self, server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None, subject_token: str
|
||||
) -> CallerSignInPreflight:
|
||||
"""Validate ``subject_token`` against this provider at connect time, where a challenge's
|
||||
``WWW-Authenticate`` still reaches the client."""
|
||||
...
|
||||
|
||||
|
||||
class _JwtIssuerEntry(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore")
|
||||
|
||||
issuer: str | None = None
|
||||
|
||||
|
||||
class _JwtAuthConfig(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore")
|
||||
|
||||
issuers: list[_JwtIssuerEntry] = [] # mutable-ok: pydantic copies the default per instance
|
||||
|
||||
|
||||
_JWT_AUTH_ADAPTER: Final = TypeAdapter(_JwtAuthConfig)
|
||||
|
||||
|
||||
def _jwt_auth_issuer_entries(jwtauth: object) -> tuple[_JwtIssuerEntry, ...]:
|
||||
try:
|
||||
return tuple(_JWT_AUTH_ADAPTER.validate_python(jwtauth, from_attributes=True).issuers)
|
||||
except ValidationError:
|
||||
return ()
|
||||
|
||||
|
||||
def _providers() -> tuple[CallerSignInProvider, ...]:
|
||||
return tuple(
|
||||
callback
|
||||
for callback in litellm.logging_callback_manager.get_custom_loggers_for_type(CustomGuardrail)
|
||||
if isinstance(callback, CallerSignInProvider)
|
||||
)
|
||||
|
||||
|
||||
def jwt_auth_issuers() -> tuple[str, ...]:
|
||||
"""The OAuth issuer identifier(s) LiteLLM's JWT auth trusts, for the OBO PRM authorization_servers.
|
||||
|
||||
In token_exchange the IdP that issues the subject JWT is the same one LiteLLM validates it
|
||||
against, so OBO discovery points clients at the JWT-auth issuer to obtain a subject token.
|
||||
Sourced from ``JWT_ISSUER`` and any configured ``litellm_jwtauth.issuers``.
|
||||
"""
|
||||
import os # noqa: PLC0415
|
||||
|
||||
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # lazy: proxy_server pulls the whole proxy graph
|
||||
general_settings, # pyright: ignore[reportUnknownVariableType] # proxy_server.general_settings is a raw untyped dict
|
||||
)
|
||||
|
||||
env_issuer: Final = os.getenv("JWT_ISSUER")
|
||||
env: Final[tuple[str, ...]] = (env_issuer,) if env_issuer else ()
|
||||
|
||||
settings: Final[Mapping[str, object]] = general_settings if isinstance(general_settings, Mapping) else {}
|
||||
configured: Final = tuple(
|
||||
entry.issuer for entry in _jwt_auth_issuer_entries(settings.get("litellm_jwtauth")) if entry.issuer
|
||||
)
|
||||
return tuple(dict.fromkeys((*env, *configured)))
|
||||
|
||||
|
||||
def caller_sign_in_for(server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None) -> CallerSignIn | None:
|
||||
"""The merged sign-in requirement for ``server``: the OBO server's own issuer/scopes plus every
|
||||
registered provider's contribution. ``None`` when nothing requires sign-in, which is also the gate the
|
||||
connect-time challenge branches on."""
|
||||
contributions: Final = tuple(
|
||||
contribution
|
||||
for contribution in (
|
||||
*(
|
||||
(CallerSignIn(issuers=jwt_auth_issuers(), scopes=tuple(server.scopes or ())),)
|
||||
if server.auth_type == MCPAuth.oauth2_token_exchange
|
||||
else ()
|
||||
),
|
||||
*(provider.caller_sign_in(server, user_api_key_auth) for provider in _providers()),
|
||||
)
|
||||
if contribution is not None
|
||||
)
|
||||
if not contributions:
|
||||
return None
|
||||
issuers: Final = tuple(dict.fromkeys(itertools.chain.from_iterable(c.issuers for c in contributions)))
|
||||
scopes: Final = tuple(dict.fromkeys(itertools.chain.from_iterable(c.scopes for c in contributions)))
|
||||
return CallerSignIn(issuers=issuers, scopes=scopes)
|
||||
|
||||
|
||||
async def preflight_caller_sign_in(
|
||||
server: MCPServer,
|
||||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
subject_token: str,
|
||||
*,
|
||||
root_path: str,
|
||||
resource_metadata: str | None,
|
||||
connecting: Callable[[], Awaitable[bool]],
|
||||
) -> None:
|
||||
"""Run every provider's connect-time check against the subject token while ``connecting``, so a bearer
|
||||
the IdP will reject surfaces as a challenge here rather than a JSON-RPC error at the first tool call, and
|
||||
a fail-closed provider outage is the connect's 503. On an open session nothing is exchanged here: the
|
||||
tool-call hook runs the one exchange and answers inside the JSON-RPC envelope, with its guardrail Logs
|
||||
row."""
|
||||
from fastapi import HTTPException # noqa: PLC0415 # lazy: fastapi import stays off the cold path
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph
|
||||
raise_token_exchange_challenge,
|
||||
)
|
||||
|
||||
gating: Final = tuple(
|
||||
provider for provider in _providers() if provider.caller_sign_in(server, user_api_key_auth) is not None
|
||||
)
|
||||
if not gating or not await connecting():
|
||||
return
|
||||
for provider in gating:
|
||||
match await provider.preflight_caller_sign_in(server, user_api_key_auth, subject_token):
|
||||
case SignedIn():
|
||||
continue
|
||||
case Rejected(detail=_, claims=claims):
|
||||
raise_token_exchange_challenge(
|
||||
server, root_path=root_path, claims=claims, resource_metadata=resource_metadata
|
||||
)
|
||||
case Unavailable(fail_open=True):
|
||||
continue
|
||||
case Unavailable(detail=detail, fail_open=False):
|
||||
raise HTTPException(status_code=503, detail=detail)
|
||||
case _ as verdict:
|
||||
assert_never(verdict)
|
||||
|
|
@ -37,6 +37,7 @@ from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
|
|||
can_store_oauth_credential,
|
||||
oauth_authorization_uses_gateway_credential,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.caller_sign_in import caller_sign_in_for
|
||||
from litellm.proxy._experimental.mcp_server.faults import (
|
||||
CallerRejected,
|
||||
CredentialSource,
|
||||
|
|
@ -527,10 +528,7 @@ def _resolve_mcp_server_by_name_or_id(lookup: str, client_ip: str | None) -> MCP
|
|||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
by_name: Final = global_mcp_server_manager.get_mcp_server_by_name(lookup, client_ip=client_ip)
|
||||
if by_name is not None:
|
||||
return by_name
|
||||
return global_mcp_server_manager.get_mcp_server_by_id(lookup, client_ip=client_ip)
|
||||
return global_mcp_server_manager.get_mcp_server_answering_to(lookup, client_ip=client_ip)
|
||||
|
||||
|
||||
def _resolve_oauth2_server_for_root_endpoints(
|
||||
|
|
@ -2580,9 +2578,9 @@ async def _build_oauth_protected_resource_response(
|
|||
detail=(f"Upstream oauth-protected-resource metadata unavailable for MCP server {mcp_server.name!r}"),
|
||||
)
|
||||
|
||||
obo_response: Final = _obo_protected_resource_response(mcp_server, resource_url)
|
||||
if obo_response is not None:
|
||||
return obo_response
|
||||
sign_in_response: Final = _caller_sign_in_protected_resource_response(mcp_server, resource_url)
|
||||
if sign_in_response is not None:
|
||||
return sign_in_response
|
||||
|
||||
if mcp_server is not None and mcp_server.advertises_gateway_authorization_server:
|
||||
return {
|
||||
|
|
@ -2603,51 +2601,30 @@ async def _build_oauth_protected_resource_response(
|
|||
}
|
||||
|
||||
|
||||
def _obo_protected_resource_response(mcp_server: MCPServer | None, resource_url: str) -> dict | None:
|
||||
"""The OBO (token_exchange) PRM, or None when this server is not OBO / no issuer is configured.
|
||||
def _caller_sign_in_protected_resource_response(
|
||||
mcp_server: MCPServer | None, resource_url: str
|
||||
) -> dict[str, object] | None:
|
||||
"""The caller sign-in PRM: the OBO issuer(s) LiteLLM trusts merged with every registered
|
||||
``CallerSignInProvider``'s contribution, or None when no sign-in gates this server.
|
||||
|
||||
The client SSOs with the IdP to obtain a subject token, which LiteLLM then exchanges, so discovery
|
||||
points at the JWT-auth issuer(s) LiteLLM trusts (the same IdP that issues and validates the
|
||||
subject), not the gateway. None falls the caller back to the gateway default so discovery still
|
||||
returns metadata; it just can't name the IdP.
|
||||
The client SSOs with the IdP to obtain a subject token, which LiteLLM then exchanges (or a
|
||||
guardrail consumes directly), so discovery points at the issuer(s), not the gateway. None falls
|
||||
the caller back to the gateway default so discovery still returns metadata; it just can't name
|
||||
the IdP. The anonymous metadata fetch passes ``user_api_key_auth=None`` because it cannot see
|
||||
which key selected a provider.
|
||||
"""
|
||||
if mcp_server is None or mcp_server.auth_type != MCPAuth.oauth2_token_exchange:
|
||||
if mcp_server is None:
|
||||
return None
|
||||
issuers: Final = _jwt_auth_issuers()
|
||||
if not issuers:
|
||||
sign_in: Final = caller_sign_in_for(mcp_server, None)
|
||||
if sign_in is None or not sign_in.issuers:
|
||||
return None
|
||||
return {
|
||||
"authorization_servers": issuers,
|
||||
"authorization_servers": sign_in.issuers,
|
||||
"resource": resource_url,
|
||||
"scopes_supported": (mcp_server.scopes if mcp_server.scopes else []),
|
||||
"scopes_supported": sign_in.scopes,
|
||||
}
|
||||
|
||||
|
||||
def _jwt_auth_issuers() -> list:
|
||||
"""The OAuth issuer identifier(s) LiteLLM's JWT auth trusts, for the OBO PRM authorization_servers.
|
||||
|
||||
In token_exchange the IdP that issues the subject JWT is the same one LiteLLM validates it
|
||||
against, so OBO discovery points clients at the JWT-auth issuer to obtain a subject token.
|
||||
Sourced from ``JWT_ISSUER`` and any configured ``litellm_jwtauth.issuers``.
|
||||
"""
|
||||
import os # noqa: PLC0415
|
||||
|
||||
from litellm.proxy.proxy_server import general_settings # noqa: PLC0415
|
||||
|
||||
issuers: Final[list] = []
|
||||
env_issuer: Final = os.getenv("JWT_ISSUER")
|
||||
if env_issuer:
|
||||
issuers.append(env_issuer)
|
||||
|
||||
jwtauth: Final = general_settings.get("litellm_jwtauth") if isinstance(general_settings, Mapping) else None
|
||||
raw_issuers: Final = jwtauth.get("issuers") if isinstance(jwtauth, dict) else getattr(jwtauth, "issuers", None)
|
||||
for cfg in raw_issuers or []:
|
||||
issuer = cfg.get("issuer") if isinstance(cfg, dict) else getattr(cfg, "issuer", None)
|
||||
if issuer and issuer not in issuers:
|
||||
issuers.append(issuer)
|
||||
return issuers
|
||||
|
||||
|
||||
@router.get("/.well-known/oauth-protected-resource")
|
||||
def oauth_protected_resource_root(request: Request) -> dict[str, str | tuple[str, ...]]:
|
||||
request_base_url: Final = get_request_base_url(request)
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ import json
|
|||
import math
|
||||
import os
|
||||
import re
|
||||
import secrets
|
||||
import time
|
||||
from collections.abc import (
|
||||
AsyncIterator,
|
||||
|
|
@ -169,6 +170,7 @@ from litellm.proxy._experimental.mcp_server.utils import (
|
|||
normalize_server_name,
|
||||
openapi_tool_name,
|
||||
parse_admin_env_vars,
|
||||
server_answers_to_name,
|
||||
strip_known_server_prefix,
|
||||
validate_mcp_server_name,
|
||||
)
|
||||
|
|
@ -1163,6 +1165,10 @@ def _raw_header_value(raw_headers: Mapping[str, str] | None, name: str) -> str |
|
|||
return next((v for k, v in (raw_headers or {}).items() if isinstance(k, str) and k.lower() == name), None)
|
||||
|
||||
|
||||
def _is_master_key(bearer: str, master_key: str | None) -> bool:
|
||||
return bool(master_key) and secrets.compare_digest(bearer.encode(), (master_key or "").encode())
|
||||
|
||||
|
||||
def _has_explicit_litellm_admission_header(raw_headers: Mapping[str, str] | None) -> bool:
|
||||
"""Admission only consumes a non-empty ``x-litellm-api-key``; an empty one falls back to ``Authorization``."""
|
||||
return bool(_raw_header_value(raw_headers, "x-litellm-api-key"))
|
||||
|
|
@ -3897,6 +3903,31 @@ class MCPServerManager:
|
|||
return None
|
||||
return bearer
|
||||
|
||||
@staticmethod
|
||||
def _caller_sign_in_subject_token(
|
||||
oauth2_headers: Mapping[str, str] | None,
|
||||
raw_headers: Mapping[str, str] | None,
|
||||
) -> str | None:
|
||||
"""The ``Bearer`` credential a caller sign-in provider validates. An admission that consumed
|
||||
``Authorization`` (custom auth, built-in OAuth2, JWT) did so on the caller's own IdP token, so that token
|
||||
is the subject; any other scheme, a LiteLLM key (virtual or master) and a bearer repeating
|
||||
``x-litellm-api-key`` are withheld."""
|
||||
from litellm.proxy.proxy_server import master_key # noqa: PLC0415 # circular import
|
||||
|
||||
authorization: Final = (oauth2_headers or {}).get("Authorization") or _raw_header_value(
|
||||
raw_headers, "authorization"
|
||||
)
|
||||
scheme_and_credential: Final = (authorization or "").split(None, 1)
|
||||
if len(scheme_and_credential) != 2 or scheme_and_credential[0].lower() != "bearer":
|
||||
return None
|
||||
bearer: Final = scheme_and_credential[1]
|
||||
if bearer.startswith(LITELLM_VIRTUAL_KEY_PREFIX) or _is_master_key(bearer, master_key):
|
||||
return None
|
||||
admission_header: Final = _raw_header_value(raw_headers, "x-litellm-api-key")
|
||||
if admission_header and strip_auth_scheme(admission_header, "Bearer") == bearer:
|
||||
return None
|
||||
return bearer
|
||||
|
||||
def _obo_subject_token(
|
||||
self,
|
||||
server: MCPServer,
|
||||
|
|
@ -4109,6 +4140,7 @@ class MCPServerManager:
|
|||
oauth2_headers: dict[str, str] | None,
|
||||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
raw_headers: Mapping[str, str] | None = None,
|
||||
resource_metadata: str | None = None,
|
||||
) -> None:
|
||||
"""Mint an exchange-backed server's upstream credential at the transport edge.
|
||||
|
||||
|
|
@ -4140,13 +4172,24 @@ class MCPServerManager:
|
|||
if subject_token is not None:
|
||||
return
|
||||
case _:
|
||||
from litellm.proxy._experimental.mcp_server.caller_sign_in import ( # noqa: PLC0415 # lazy: provider discovery pulls the guardrail registry
|
||||
caller_sign_in_for,
|
||||
)
|
||||
|
||||
sign_in_subject: Final = self._caller_sign_in_subject_token(oauth2_headers, raw_headers)
|
||||
if sign_in_subject is None and caller_sign_in_for(server, user_api_key_auth) is not None:
|
||||
raise_token_exchange_challenge(
|
||||
server, root_path=get_request_root_path(), resource_metadata=resource_metadata
|
||||
)
|
||||
return
|
||||
resolved_server: Final = await self.ensure_oauth_metadata_discovered(server)
|
||||
spec: Final = to_server_spec_fail_closed(resolved_server)
|
||||
if spec is None or not isinstance(spec.config, (TokenExchangeConfig, IdJagConfig)):
|
||||
return
|
||||
if subject_token is None and isinstance(spec.config, TokenExchangeConfig):
|
||||
raise_token_exchange_challenge(resolved_server, root_path=get_request_root_path())
|
||||
raise_token_exchange_challenge(
|
||||
resolved_server, root_path=get_request_root_path(), resource_metadata=resource_metadata
|
||||
)
|
||||
match await self._cred_provider.resolve_credentials(to_subject(user_api_key_auth, subject_token), spec):
|
||||
case Ok(_):
|
||||
return
|
||||
|
|
@ -4156,6 +4199,7 @@ class MCPServerManager:
|
|||
resolved_server,
|
||||
root_path=get_request_root_path(),
|
||||
claims=err.unauthorized.claims,
|
||||
resource_metadata=resource_metadata,
|
||||
)
|
||||
raise_public(err)
|
||||
|
||||
|
|
@ -5747,13 +5791,11 @@ class MCPServerManager:
|
|||
if proxy_logging_obj is None:
|
||||
return hook_result
|
||||
|
||||
# Extract incoming Bearer token from raw request headers so
|
||||
# guardrails like MCPJWTSigner can verify + re-sign it (FR-5).
|
||||
normalized_raw: Final = {k.lower(): v for k, v in (raw_headers or {}).items()}
|
||||
incoming_bearer_token: str | None = None
|
||||
auth_hdr: Final = normalized_raw.get("authorization", "")
|
||||
if auth_hdr.lower().startswith("bearer "):
|
||||
incoming_bearer_token = auth_hdr[len("bearer ") :]
|
||||
inbound_authorization: Final = _raw_header_value(raw_headers, "authorization") or ""
|
||||
incoming_bearer_token: Final = (
|
||||
inbound_authorization[len("bearer ") :] if inbound_authorization.lower().startswith("bearer ") else None
|
||||
)
|
||||
incoming_subject_token: Final = self._caller_sign_in_subject_token(None, raw_headers)
|
||||
|
||||
pre_hook_kwargs: Final = {
|
||||
"guardrail_context": guardrail_context,
|
||||
|
|
@ -5769,6 +5811,7 @@ class MCPServerManager:
|
|||
),
|
||||
"user_api_key_hash": (getattr(user_api_key_auth, "api_key_hash", None) if user_api_key_auth else None),
|
||||
"incoming_bearer_token": incoming_bearer_token,
|
||||
"incoming_subject_token": incoming_subject_token,
|
||||
"headers": logging_safe_mcp_headers(raw_headers),
|
||||
"tool_description": tool.description if tool is not None else None,
|
||||
"tool_input_schema": tool.input_schema if tool is not None else None,
|
||||
|
|
@ -6981,6 +7024,70 @@ class MCPServerManager:
|
|||
return server
|
||||
return None
|
||||
|
||||
def get_mcp_server_answering_to(
|
||||
self, name: str, client_ip: str | None = None, *, among: Sequence[MCPServer] | None = None
|
||||
) -> MCPServer | None:
|
||||
"""The one server a ``/mcp/{name}`` segment denotes, shared by the connect preflight, the scoped
|
||||
router, and RFC 9728 discovery so all three name the same server: the exact ``get_mcp_server_by_name``
|
||||
priority first, then the exact ``server_id``, then the name priority case-insensitively, then any prefix
|
||||
form routing accepts. A name that denotes a server hidden from ``client_ip`` resolves to ``None`` at the
|
||||
pass that found it: it never falls through to a looser pass that could name another server. ``among``
|
||||
runs the same passes over those servers alone instead of the registry, which is how the scoped router
|
||||
picks the caller's granted server answering to ``name``."""
|
||||
if among is not None:
|
||||
return self._server_among_answering_to(name, tuple(among), client_ip)
|
||||
exact: Final = self.get_mcp_server_by_name(name)
|
||||
if exact is not None:
|
||||
return exact if self._is_server_accessible_from_ip(exact, client_ip) else None
|
||||
by_id: Final = self.get_mcp_server_by_id(name)
|
||||
if by_id is not None:
|
||||
return by_id if self._is_server_accessible_from_ip(by_id, client_ip) else None
|
||||
requested: Final = name.lower()
|
||||
servers: Final = tuple(self.get_registry().values())
|
||||
identifiers: Final[tuple[Callable[[MCPServer], str | None], ...]] = (
|
||||
lambda server: server.alias,
|
||||
lambda server: server.server_name,
|
||||
lambda server: server.name,
|
||||
)
|
||||
for identifier in identifiers:
|
||||
if (found := next((s for s in servers if (identifier(s) or "").lower() == requested), None)) is not None:
|
||||
return found if self._is_server_accessible_from_ip(found, client_ip) else None
|
||||
return next(
|
||||
(
|
||||
server
|
||||
for server in self.get_filtered_registry(client_ip).values()
|
||||
if server_answers_to_name(server, name)
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
def _server_among_answering_to(
|
||||
self, name: str, servers: Sequence[MCPServer], client_ip: str | None
|
||||
) -> MCPServer | None:
|
||||
"""``get_mcp_server_answering_to`` over ``servers`` instead of the registry: the same passes in the
|
||||
same order, with a server hidden from ``client_ip`` resolving to ``None`` at the pass that found it."""
|
||||
requested: Final = name.lower()
|
||||
passes: Final[tuple[Callable[[MCPServer], bool], ...]] = (
|
||||
lambda server: server.alias == name,
|
||||
lambda server: server.server_name == name,
|
||||
lambda server: server.name == name,
|
||||
lambda server: server.server_id == name,
|
||||
lambda server: (server.alias or "").lower() == requested,
|
||||
lambda server: (server.server_name or "").lower() == requested,
|
||||
lambda server: (server.name or "").lower() == requested,
|
||||
)
|
||||
for matches in passes:
|
||||
if (found := next((server for server in servers if matches(server)), None)) is not None:
|
||||
return found if self._is_server_accessible_from_ip(found, client_ip) else None
|
||||
return next(
|
||||
(
|
||||
server
|
||||
for server in servers
|
||||
if self._is_server_accessible_from_ip(server, client_ip) and server_answers_to_name(server, name)
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
def get_filtered_registry(self, client_ip: str | None = None) -> dict[str, MCPServer]:
|
||||
"""
|
||||
Get registry filtered by client IP access control.
|
||||
|
|
|
|||
|
|
@ -120,6 +120,7 @@ from litellm.proxy._experimental.mcp_server.utils import (
|
|||
logging_safe_mcp_headers,
|
||||
match_known_tool_name,
|
||||
normalize_server_name,
|
||||
server_answers_to_name,
|
||||
split_server_prefix_from_name,
|
||||
strip_known_server_prefix,
|
||||
)
|
||||
|
|
@ -435,6 +436,7 @@ async def _dispatch_virtual_mcp_tool(
|
|||
async def _get_allowed_mcp_servers_from_mcp_server_names(
|
||||
mcp_servers: Sequence[str] | None,
|
||||
allowed_mcp_servers: list[MCPServer],
|
||||
client_ip: str | None = None,
|
||||
) -> list[MCPServer]:
|
||||
"""
|
||||
Get the filtered MCP servers from the MCP server names.
|
||||
|
|
@ -451,15 +453,10 @@ async def _get_allowed_mcp_servers_from_mcp_server_names(
|
|||
# Filter servers based on mcp_servers parameter if provided
|
||||
if mcp_servers is not None:
|
||||
for server_or_group in mcp_servers:
|
||||
server_name_matched = False
|
||||
|
||||
for server in allowed_mcp_servers:
|
||||
if server and _server_answers_to(server, server_or_group):
|
||||
filtered_server[server.server_id] = server
|
||||
server_name_matched = True
|
||||
break
|
||||
|
||||
if not server_name_matched:
|
||||
scoped = _scoped_server(server_or_group, allowed_mcp_servers, client_ip)
|
||||
if scoped is not None:
|
||||
filtered_server[scoped.server_id] = scoped
|
||||
else:
|
||||
try:
|
||||
access_group_server_ids = await MCPRequestHandler._get_mcp_servers_from_access_groups(
|
||||
[server_or_group]
|
||||
|
|
@ -493,8 +490,20 @@ def _http_detail_message(detail: object) -> str:
|
|||
|
||||
|
||||
def _server_answers_to(server: MCPServer, name: str) -> bool:
|
||||
requested: Final = name.lower()
|
||||
return any(requested == known.lower() for known in iter_known_server_prefixes(server) if known)
|
||||
return server_answers_to_name(server, name)
|
||||
|
||||
|
||||
def _scoped_server(name: str, allowed_mcp_servers: Sequence[MCPServer], client_ip: str | None) -> MCPServer | None:
|
||||
"""The granted server a scoped ``name`` selects for the caller: the registry's own pass order run over
|
||||
``allowed_mcp_servers`` alone, so a granted server wins over an ungranted alias or case variant the registry
|
||||
would pick. ``None`` when no granted server answers, or when the registry's own pick for ``name`` is a server
|
||||
hidden from ``client_ip``; the caller then retries the name as an access group it holds."""
|
||||
if (
|
||||
global_mcp_server_manager.get_mcp_server_answering_to(name, client_ip=client_ip) is None
|
||||
and global_mcp_server_manager.get_mcp_server_answering_to(name) is not None
|
||||
):
|
||||
return None
|
||||
return global_mcp_server_manager.get_mcp_server_answering_to(name, client_ip=client_ip, among=allowed_mcp_servers)
|
||||
|
||||
|
||||
async def raise_denied_scoped_mcp_access(
|
||||
|
|
@ -675,6 +684,7 @@ async def _get_allowed_mcp_servers(
|
|||
allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names(
|
||||
mcp_servers=mcp_servers,
|
||||
allowed_mcp_servers=allowed_mcp_servers,
|
||||
client_ip=client_ip,
|
||||
)
|
||||
|
||||
return allowed_mcp_servers
|
||||
|
|
@ -1776,7 +1786,9 @@ def _challenge_missing_token_exchange_subject(
|
|||
The listing that fills a cold catalog absorbs the upstream 401 by design, so without this
|
||||
check a missing subject surfaces as an unknown-tool error instead of the challenge the
|
||||
warm path already raises. Gated to servers the key may reach so an unauthorized caller
|
||||
learns nothing about the catalog.
|
||||
learns nothing about the catalog. Guardrail-only sign-in is challenged at connect instead: a
|
||||
tool call's JSON-RPC error drops ``WWW-Authenticate``, so the guardrail's own rejection is
|
||||
the more useful answer there.
|
||||
"""
|
||||
if server is None or server.auth_type != MCPAuth.oauth2_token_exchange:
|
||||
return
|
||||
|
|
@ -2439,6 +2451,7 @@ async def call_mcp_tool(
|
|||
allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names(
|
||||
mcp_servers=mcp_servers,
|
||||
allowed_mcp_servers=allowed_mcp_servers,
|
||||
client_ip=client_ip,
|
||||
)
|
||||
if mcp_servers and not allowed_mcp_servers:
|
||||
await raise_denied_scoped_mcp_access(
|
||||
|
|
|
|||
|
|
@ -348,6 +348,7 @@ def raise_token_exchange_challenge(
|
|||
*,
|
||||
root_path: str,
|
||||
claims: str | None = None,
|
||||
resource_metadata: str | None = None,
|
||||
) -> NoReturn:
|
||||
"""Raise the RFC 9728 / RFC 6750 challenge an OBO (``token_exchange``) server returns when the
|
||||
caller's subject token is missing or the IdP rejected it.
|
||||
|
|
@ -365,8 +366,12 @@ def raise_token_exchange_challenge(
|
|||
``error="invalid_token"`` and is byte-identical to the static one. Both the error value (one of
|
||||
two literals) and the base64 claims draw from a fixed alphabet, so nothing from the IdP body
|
||||
reaches the header unescaped.
|
||||
|
||||
``resource_metadata`` is the absolute metadata URL of the route the client connected on (RFC 9728
|
||||
5.1 names the parameter a URL, and the MCP SDK fetches it verbatim), supplied by the connect gate
|
||||
that still holds the request; without it the challenge falls back to the alias's relative path.
|
||||
"""
|
||||
resource_metadata: Final = oauth_protected_resource_path(root_path, server)
|
||||
metadata_url: Final = resource_metadata or oauth_protected_resource_path(root_path, server)
|
||||
encoded_claims: Final = base64.b64encode(claims.encode()).decode() if claims else None
|
||||
error: Final = "insufficient_claims" if encoded_claims else "invalid_token"
|
||||
error_description: Final = (
|
||||
|
|
@ -376,7 +381,7 @@ def raise_token_exchange_challenge(
|
|||
)
|
||||
www_authenticate: Final = ", ".join(
|
||||
(
|
||||
f'Bearer resource_metadata="{resource_metadata}"',
|
||||
f'Bearer resource_metadata="{metadata_url}"',
|
||||
f'error="{error}"',
|
||||
f'error_description="{error_description}"',
|
||||
*((f'claims="{encoded_claims}"',) if encoded_claims else ()),
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ call), so it needs no lazy wrapper.
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
|
|
@ -24,6 +25,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_sto
|
|||
InMemoryTokenCacheBackend,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import (
|
||||
ExchangeHttpPost,
|
||||
OboTokenExchanger,
|
||||
SubjectTokenRejected,
|
||||
TokenExchangeClientError,
|
||||
|
|
@ -34,11 +36,27 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger
|
|||
_GATEWAY_FAULT_OAUTH_ERRORS: Final = frozenset(
|
||||
{"invalid_client", "unauthorized_client", "unsupported_grant_type", "invalid_target", "invalid_scope"}
|
||||
)
|
||||
_INVALID_ASSERTION_AADSTS_PREFIX: Final = "50027"
|
||||
|
||||
|
||||
def _oauth_error_fields(response: httpx.Response) -> tuple[str | None, str | None]:
|
||||
"""Read the RFC 6749 5.2 ``error`` code and the IdP's step-up ``claims`` blob from a
|
||||
token-endpoint error body, as ``(error, claims)`` with None for whatever is absent.
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class OAuthErrorBody:
|
||||
error: str | None
|
||||
claims: str | None
|
||||
error_codes: tuple[str, ...]
|
||||
|
||||
@property
|
||||
def gateway_fault(self) -> str | None:
|
||||
if self.error is None or self.error not in _GATEWAY_FAULT_OAUTH_ERRORS:
|
||||
return None
|
||||
if any(code.startswith(_INVALID_ASSERTION_AADSTS_PREFIX) for code in self.error_codes):
|
||||
return None
|
||||
return self.error
|
||||
|
||||
|
||||
def oauth_error_fields(response: httpx.Response) -> OAuthErrorBody:
|
||||
"""Read the RFC 6749 5.2 ``error`` code, the IdP's step-up ``claims`` blob and Entra's
|
||||
``error_codes`` sub-codes from a token-endpoint error body, None or empty for whatever is absent.
|
||||
|
||||
``claims`` is the Entra Conditional Access / CAE challenge (a JSON string the client must
|
||||
replay to the IdP to satisfy the step-up); it is the caller's own requirement, not an IdP
|
||||
|
|
@ -48,14 +66,18 @@ def _oauth_error_fields(response: httpx.Response) -> tuple[str | None, str | Non
|
|||
try:
|
||||
body: Final[object] = response.json()
|
||||
except Exception: # noqa: BLE001
|
||||
return None, None
|
||||
return OAuthErrorBody(error=None, claims=None, error_codes=())
|
||||
if not isinstance(body, dict):
|
||||
return None, None
|
||||
return OAuthErrorBody(error=None, claims=None, error_codes=())
|
||||
code: Final = body.get("error")
|
||||
claims: Final = body.get("claims")
|
||||
return (
|
||||
code if isinstance(code, str) else None,
|
||||
claims if isinstance(claims, str) and claims else None,
|
||||
raw_codes: Final = body.get("error_codes")
|
||||
return OAuthErrorBody(
|
||||
error=code if isinstance(code, str) else None,
|
||||
claims=claims if isinstance(claims, str) and claims else None,
|
||||
error_codes=tuple(str(c) for c in raw_codes if isinstance(c, (int, str)))
|
||||
if isinstance(raw_codes, list)
|
||||
else (),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -74,24 +96,32 @@ async def _post_exchange_endpoint(
|
|||
headers: Final = {"Accept": "application/json", **client_auth_headers}
|
||||
try:
|
||||
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) # pyright: ignore
|
||||
response: Final = await client.post(url, headers=headers, data=form) # pyright: ignore
|
||||
response: Final = await client.post( # pyright: ignore[reportUnknownMemberType] # untyped handler
|
||||
url, headers=headers, data=form
|
||||
)
|
||||
response.raise_for_status() # pyright: ignore
|
||||
parsed: Final[object] = response.json() # pyright: ignore
|
||||
except httpx.HTTPStatusError as status_err:
|
||||
status_code: Final = status_err.response.status_code
|
||||
if status_code in (408, 429):
|
||||
# Retry hints, not subject rejections: the IdP is shedding load, so a 401 would tell the
|
||||
# caller to sign in again for nothing; surface it like a transport failure.
|
||||
verbose_logger.warning("MCP token exchange throttled or timed out (HTTP %d)", status_code)
|
||||
return None
|
||||
if 400 <= status_code < 500:
|
||||
oauth_error, claims = _oauth_error_fields(status_err.response)
|
||||
if oauth_error in _GATEWAY_FAULT_OAUTH_ERRORS:
|
||||
oauth_error: Final = oauth_error_fields(status_err.response)
|
||||
gateway_fault: Final = oauth_error.gateway_fault
|
||||
if gateway_fault is not None:
|
||||
verbose_logger.warning(
|
||||
"MCP token exchange rejected as %s (HTTP %d); check the gateway client credentials, "
|
||||
"audience, and scope for this server",
|
||||
oauth_error,
|
||||
gateway_fault,
|
||||
status_code,
|
||||
)
|
||||
raise TokenExchangeClientError(oauth_error) from status_err
|
||||
raise TokenExchangeClientError(gateway_fault) from status_err
|
||||
raise SubjectTokenRejected(
|
||||
f"IdP rejected the subject token (HTTP {status_code})",
|
||||
claims=claims,
|
||||
claims=oauth_error.claims,
|
||||
) from status_err
|
||||
verbose_logger.warning("MCP token exchange request failed: %s", status_err)
|
||||
return None
|
||||
|
|
@ -106,9 +136,9 @@ async def _post_exchange_endpoint(
|
|||
return parsed # pyright: ignore
|
||||
|
||||
|
||||
def build_token_exchanger() -> OboTokenExchanger:
|
||||
def build_token_exchanger(*, post: ExchangeHttpPost = _post_exchange_endpoint) -> OboTokenExchanger:
|
||||
return OboTokenExchanger(
|
||||
_post_exchange_endpoint,
|
||||
post,
|
||||
cache=InMemoryTokenCacheBackend(max_size=MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE),
|
||||
default_ttl_seconds=MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL,
|
||||
min_ttl_seconds=MCP_OAUTH2_TOKEN_CACHE_MIN_TTL,
|
||||
|
|
|
|||
|
|
@ -8,12 +8,13 @@ import asyncio
|
|||
import contextlib
|
||||
import contextvars
|
||||
import hashlib
|
||||
import itertools
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
import types
|
||||
from collections import Counter
|
||||
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Iterable, Mapping, Sequence
|
||||
from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Iterable, Iterator, Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Final, NoReturn, Protocol
|
||||
|
||||
import httpx
|
||||
|
|
@ -36,6 +37,7 @@ 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.caller_sign_in import caller_sign_in_for
|
||||
from litellm.proxy._experimental.mcp_server.client_allowlist import (
|
||||
MCPClientAllowlist,
|
||||
check_mcp_client_allowed,
|
||||
|
|
@ -62,6 +64,7 @@ from litellm.proxy._experimental.mcp_server.mcp_debug import (
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
_redact_mcp_resource_url,
|
||||
get_passthrough_resource_metadata_url,
|
||||
get_passthrough_www_authenticate,
|
||||
get_route_relative_request_path,
|
||||
well_known_root_suffix,
|
||||
|
|
@ -1424,6 +1427,41 @@ if MCP_AVAILABLE:
|
|||
|
||||
return consumed_messages, b"".join(body_chunks)
|
||||
|
||||
class _ConnectBodyPeek:
|
||||
"""Reads a session-less ``POST`` body only once a gate asks whether it is ``initialize``, so a challenge
|
||||
that needs no body still answers before the body arrives; consumed messages replay through ``receive``."""
|
||||
|
||||
def __init__(self, receive: Receive, peekable: bool) -> None:
|
||||
self._receive: Final = receive
|
||||
self._peekable: Final = peekable
|
||||
self._body: bytes | None = None
|
||||
self._replay: Iterator[Message] = iter(())
|
||||
|
||||
async def read(self) -> bytes:
|
||||
messages, body = await _read_request_body_for_routing(self._receive)
|
||||
self._replay = itertools.chain(self._replay, messages)
|
||||
return body
|
||||
|
||||
async def body(self) -> bytes:
|
||||
if not self._peekable:
|
||||
return b""
|
||||
if self._body is None:
|
||||
self._body = await self.read()
|
||||
return self._body
|
||||
|
||||
async def connecting(self) -> bool:
|
||||
return _is_initialize_request(await self.body())
|
||||
|
||||
async def receive(self) -> Message:
|
||||
replayed: Final = next(self._replay, None)
|
||||
return replayed if replayed is not None else await self._receive()
|
||||
|
||||
def _known_connecting(value: bool) -> Callable[[], Awaitable[bool]]:
|
||||
async def answer() -> bool:
|
||||
return value
|
||||
|
||||
return answer
|
||||
|
||||
async def _handle_stale_mcp_session(
|
||||
scope: Scope,
|
||||
receive: Receive,
|
||||
|
|
@ -1606,6 +1644,7 @@ if MCP_AVAILABLE:
|
|||
mcp_server_auth_headers: dict[str, dict[str, str]] | None,
|
||||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
client_ip: str | None,
|
||||
connecting: Callable[[], Awaitable[bool]],
|
||||
allowed_server_ids: set[str] | None = None,
|
||||
raw_headers: Mapping[str, str] | None = None,
|
||||
) -> None:
|
||||
|
|
@ -1621,7 +1660,28 @@ if MCP_AVAILABLE:
|
|||
a server it will be 403'd on immediately after authentication.
|
||||
"""
|
||||
for server_name in mcp_servers or []:
|
||||
server = operations.global_mcp_server_manager.get_mcp_server_by_name(server_name, client_ip=client_ip)
|
||||
registry_pick = operations.global_mcp_server_manager.get_mcp_server_answering_to(
|
||||
server_name, client_ip=client_ip
|
||||
)
|
||||
allowed_single = (
|
||||
await operations._get_allowed_mcp_servers(
|
||||
user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip
|
||||
)
|
||||
if registry_pick and mcp_servers is not None and len(mcp_servers) == 1
|
||||
else ()
|
||||
)
|
||||
granted = (
|
||||
operations.global_mcp_server_manager.get_mcp_server_answering_to(
|
||||
server_name, client_ip=client_ip, among=allowed_single
|
||||
)
|
||||
if allowed_single
|
||||
else None
|
||||
)
|
||||
server = granted if granted is not None else registry_pick
|
||||
granted_single = granted is not None
|
||||
obo_without_subject = (
|
||||
server is not None and server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers
|
||||
)
|
||||
if server is not None and allowed_server_ids is not None and server.server_id not in allowed_server_ids:
|
||||
# Caller's narrowed scope excludes this server — skip the
|
||||
# preemptive challenge and let downstream authorization
|
||||
|
|
@ -1717,12 +1777,21 @@ if MCP_AVAILABLE:
|
|||
# reaches the token_exchange / pass-through blocks below.
|
||||
continue
|
||||
|
||||
# token_exchange (OBO): the caller supplied no subject token. Challenge at connect
|
||||
# (transport level, where WWW-Authenticate survives) with the RFC 9728 resource_metadata
|
||||
# so the client discovers the IdP, SSOs, and retries with a subject token, which LiteLLM
|
||||
# then exchanges. A tool-call-time 401 would be wrapped into a JSON-RPC error and the
|
||||
# header lost, so the discovery flow needs this pre-emptive challenge.
|
||||
if server and server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers:
|
||||
# Caller sign-in: challenge at connect because a tool-call-time 401 is wrapped into a
|
||||
# JSON-RPC error and the WWW-Authenticate header is lost. OBO keeps its connect gate;
|
||||
# guardrail-only gates fire only on a single-server connect the key's grant admits, so a
|
||||
# key without access gets the grant's 403 instead of a sign-in it could not use. The one
|
||||
# admission lookup above serves the challenge, the sign-in preflight and the exchange.
|
||||
sign_in = caller_sign_in_for(server, user_api_key_auth) if server is not None else None
|
||||
resource_metadata = get_passthrough_resource_metadata_url(scope, server_name)
|
||||
subject_token = (
|
||||
operations.global_mcp_server_manager._caller_sign_in_subject_token( # pyright: ignore[reportPrivateUsage] # the manager owns the subject/admission filter shared with the preflight
|
||||
oauth2_headers, raw_headers
|
||||
)
|
||||
if server is not None
|
||||
else None
|
||||
)
|
||||
if server and sign_in is not None and subject_token is None and (obo_without_subject or granted_single):
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph
|
||||
raise_token_exchange_challenge,
|
||||
)
|
||||
|
|
@ -1730,7 +1799,25 @@ if MCP_AVAILABLE:
|
|||
get_request_root_path,
|
||||
)
|
||||
|
||||
raise_token_exchange_challenge(server, root_path=get_request_root_path())
|
||||
raise_token_exchange_challenge(
|
||||
server, root_path=get_request_root_path(), resource_metadata=resource_metadata
|
||||
)
|
||||
if server and sign_in is not None and subject_token is not None and granted_single:
|
||||
from litellm.proxy._experimental.mcp_server.caller_sign_in import ( # noqa: PLC0415 # lazy: provider discovery pulls the guardrail registry
|
||||
preflight_caller_sign_in,
|
||||
)
|
||||
from litellm.proxy.middleware.per_request_root_path_middleware import ( # noqa: PLC0415 # lazy: middleware imports proxy utils
|
||||
get_request_root_path,
|
||||
)
|
||||
|
||||
await preflight_caller_sign_in(
|
||||
server,
|
||||
user_api_key_auth,
|
||||
subject_token,
|
||||
root_path=get_request_root_path(),
|
||||
resource_metadata=resource_metadata,
|
||||
connecting=connecting,
|
||||
)
|
||||
|
||||
# Exchange-backed modes (token_exchange's OBO mint, id_jag's stored-assertion mint): run
|
||||
# the exchange here at the transport edge, so a rejected subject raises the RFC 9728
|
||||
|
|
@ -1739,22 +1826,13 @@ if MCP_AVAILABLE:
|
|||
# and what each mints from. Gated to single-server routes the key may reach; the
|
||||
# multi-server aggregate keeps absorbing per-server auth failures so one bad server
|
||||
# cannot 401 the whole connect.
|
||||
if (
|
||||
server
|
||||
and len(mcp_servers or []) == 1
|
||||
and server.server_id
|
||||
in frozenset(
|
||||
allowed.server_id
|
||||
for allowed in await operations._get_allowed_mcp_servers(
|
||||
user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip
|
||||
)
|
||||
)
|
||||
):
|
||||
if server and granted_single:
|
||||
await operations.global_mcp_server_manager.preflight_token_exchange(
|
||||
server=server,
|
||||
oauth2_headers=oauth2_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
raw_headers=raw_headers,
|
||||
resource_metadata=resource_metadata,
|
||||
)
|
||||
|
||||
# Pass-through OAuth: when the admin has opted a server into
|
||||
|
|
@ -2030,6 +2108,20 @@ if MCP_AVAILABLE:
|
|||
user_api_key_auth = await _apply_toolset_scope(user_api_key_auth, active_toolset_id)
|
||||
toolset_allowed_server_ids = await _toolset_server_ids(active_toolset_id)
|
||||
|
||||
named_session_id: Final = _get_session_id_from_scope(scope)
|
||||
names_live_session: Final = (
|
||||
named_session_id is not None and named_session_id in _stateful_server_instances()
|
||||
)
|
||||
request_owner: Final = _owner_fingerprint_for(user_api_key_auth, oauth2_headers, _client_ip)
|
||||
expected_owner: Final = (
|
||||
_stateful_session_owners.get(named_session_id) if named_session_id is not None else None
|
||||
)
|
||||
owner_mismatch: Final = expected_owner is not None and expected_owner != request_owner
|
||||
connect_peek: Final = _ConnectBodyPeek(
|
||||
receive, peekable=scope.get("method") == "POST" and not names_live_session and not owner_mismatch
|
||||
)
|
||||
receive = connect_peek.receive
|
||||
|
||||
# https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response
|
||||
# Must run after toolset scoping so the challenge set is derived
|
||||
# from the fully-authorized server set: a passthrough server that
|
||||
|
|
@ -2042,6 +2134,7 @@ if MCP_AVAILABLE:
|
|||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
client_ip=_client_ip,
|
||||
connecting=connect_peek.connecting,
|
||||
allowed_server_ids=toolset_allowed_server_ids,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
|
|
@ -2077,8 +2170,6 @@ if MCP_AVAILABLE:
|
|||
# - No session ID + initialize → stateful (so client gets mcp-session-id)
|
||||
# - No session ID + other → stateless (curl, Inspector, Notion)
|
||||
session_id = _get_session_id_from_scope(scope)
|
||||
is_initialize = False
|
||||
consumed_messages: list[Message] = []
|
||||
|
||||
# Owner-binding: a live stateful session may only be driven by the
|
||||
# caller that created it. Reject mismatches with 403 so a leaked
|
||||
|
|
@ -2086,12 +2177,9 @@ if MCP_AVAILABLE:
|
|||
#
|
||||
# Run before ``_handle_stale_mcp_session`` so a non-owner cannot
|
||||
# force-clean another caller's residual tracking entries via a
|
||||
# stale DELETE, and before peeking the request body so the 403
|
||||
# response sees a pristine ``receive`` channel.
|
||||
# stale DELETE.
|
||||
if session_id:
|
||||
expected_owner: Final = _stateful_session_owners.get(session_id)
|
||||
request_owner = _owner_fingerprint_for(user_api_key_auth, oauth2_headers, _client_ip)
|
||||
if expected_owner is not None and expected_owner != request_owner:
|
||||
if owner_mismatch:
|
||||
verbose_logger.warning(
|
||||
"Rejecting MCP request: session '%s' owner mismatch.",
|
||||
session_id,
|
||||
|
|
@ -2116,10 +2204,11 @@ if MCP_AVAILABLE:
|
|||
return
|
||||
session_id = _get_session_id_from_scope(scope)
|
||||
|
||||
body = b""
|
||||
if scope.get("method") == "POST":
|
||||
consumed_messages, body = await _read_request_body_for_routing(receive)
|
||||
is_initialize = _is_initialize_request(body)
|
||||
session_body: Final = (
|
||||
await connect_peek.read() if scope.get("method") == "POST" and names_live_session else b""
|
||||
)
|
||||
body: Final = await connect_peek.body() or session_body
|
||||
is_initialize: Final = _is_initialize_request(body)
|
||||
|
||||
use_stateful: Final = bool(session_id or is_initialize)
|
||||
target_manager: Final = session_manager_stateful if use_stateful else session_manager_stateless
|
||||
|
|
@ -2134,7 +2223,6 @@ if MCP_AVAILABLE:
|
|||
# session. Cap how many a single caller can hold so an authenticated
|
||||
# client cannot spam `initialize` and exhaust memory.
|
||||
if is_initialize and not session_id:
|
||||
request_owner = _owner_fingerprint_for(user_api_key_auth, oauth2_headers, _client_ip)
|
||||
if not await _enforce_stateful_session_cap_for_owner(request_owner):
|
||||
verbose_logger.warning(
|
||||
"Rejecting MCP initialize: caller already holds the maximum number of active stateful sessions."
|
||||
|
|
@ -2149,17 +2237,6 @@ if MCP_AVAILABLE:
|
|||
await too_many_response(scope, receive, send)
|
||||
return
|
||||
|
||||
# Replay body messages if we consumed them for peeking
|
||||
original_receive: Final = receive
|
||||
if consumed_messages:
|
||||
|
||||
async def wrapped_receive():
|
||||
if consumed_messages:
|
||||
return consumed_messages.pop(0)
|
||||
return await original_receive()
|
||||
|
||||
receive = wrapped_receive
|
||||
|
||||
# Serialize requests on the same stateful session so concurrent
|
||||
# callers don't clobber each other's auth context mid-flight.
|
||||
#
|
||||
|
|
@ -2389,6 +2466,7 @@ if MCP_AVAILABLE:
|
|||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
client_ip=_sse_client_ip,
|
||||
connecting=_known_connecting(scope["method"] == "GET"),
|
||||
allowed_server_ids=toolset_allowed_server_ids,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -362,6 +362,14 @@ def iter_known_server_prefixes(server: _McpServerLike) -> Iterator[str]:
|
|||
yield from _emit(server_id)
|
||||
|
||||
|
||||
def server_answers_to_name(server: _McpServerLike, name: str) -> bool:
|
||||
"""Whether a scoped ``/mcp/{name}`` connect resolves to ``server``: case-insensitive over every prefix
|
||||
form routing accepts (alias, server_name, server_id, short prefix), the same match
|
||||
``_server_answers_to`` applies when the router scopes a request."""
|
||||
requested: Final = name.lower()
|
||||
return any(requested == known.lower() for known in iter_known_server_prefixes(server) if known)
|
||||
|
||||
|
||||
def iter_known_tool_name_spellings(tool_name: str, server: MCPServer) -> Iterator[str]:
|
||||
"""Yield every name that denotes the bare ``tool_name`` on ``server``: the bare name,
|
||||
then its wire spelling under each prefix ``iter_known_server_prefixes`` accepts.
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ from .agent_365 import Agent365Guardrail
|
|||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import TokenExchanger
|
||||
from litellm.types.guardrails import Guardrail, LitellmParams
|
||||
|
||||
|
||||
|
|
@ -15,6 +16,7 @@ def initialize_guardrail(
|
|||
guardrail: "Guardrail",
|
||||
*,
|
||||
async_handler: "AsyncHTTPHandler | None" = None,
|
||||
token_exchanger: "TokenExchanger | None" = None,
|
||||
) -> Agent365Guardrail:
|
||||
import litellm
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
|
@ -64,6 +66,7 @@ def initialize_guardrail(
|
|||
request_timeout=litellm_params.timeout if litellm_params.timeout is not None else 10.0,
|
||||
unreachable_fallback=litellm_params.unreachable_fallback,
|
||||
async_handler=async_handler,
|
||||
token_exchanger=token_exchanger,
|
||||
event_hook=litellm_params.mode,
|
||||
default_on=litellm_params.default_on,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -9,17 +9,14 @@ exchanged for a delegated Agent 365 token, so Defender evaluates and audits
|
|||
as the signed-in user.
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from collections import OrderedDict
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, ClassVar, Final, Literal, NoReturn
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException
|
||||
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError
|
||||
from pydantic import BaseModel, ConfigDict, Field, SecretStr, TypeAdapter, ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -34,7 +31,30 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.caller_sign_in import (
|
||||
CallerSignIn,
|
||||
CallerSignInPreflight,
|
||||
Rejected,
|
||||
SignedIn,
|
||||
Unavailable,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import OAuthToken
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok, Result
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_provider import (
|
||||
build_token_exchanger,
|
||||
oauth_error_fields,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import (
|
||||
SubjectTokenRejected,
|
||||
TokenExchanger,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
CredError,
|
||||
ServerSpec,
|
||||
TokenExchangeConfig,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.agent_365 import (
|
||||
AGENT_365_PROD_API_BASE,
|
||||
AGENT_365_PROD_RESOURCE_APP_ID,
|
||||
|
|
@ -49,43 +69,19 @@ if TYPE_CHECKING:
|
|||
from litellm.types.utils import GuardrailStatus
|
||||
|
||||
TOKEN_ENDPOINT_TEMPLATE: Final = "https://login.microsoftonline.com/{tenant_id}/oauth2/v2.0/token"
|
||||
ENTRA_ISSUER_TEMPLATE: Final = "https://login.microsoftonline.com/{tenant_id}/v2.0"
|
||||
EVALUATE_URL: Final = f"{AGENT_365_PROD_API_BASE}/agents/tool-evaluation/evaluate"
|
||||
OBO_SCOPE: Final = f"{AGENT_365_PROD_RESOURCE_APP_ID}/{AGENT_365_SCOPE_NAME}"
|
||||
MCP_SESSION_ID_HEADER: Final = "mcp-session-id"
|
||||
DEFENDER_STATUS_EVALUATED: Final = "Evaluated"
|
||||
_GATEWAY_OWNED_TOKEN_ERRORS: Final = frozenset(
|
||||
{"invalid_client", "unauthorized_client", "invalid_scope", "invalid_resource"}
|
||||
)
|
||||
# Entra reports a malformed or unverifiable assertion as ``invalid_client`` too; only its AADSTS50027xx
|
||||
# (InvalidJwtToken) sub-codes tell that apart from a bad gateway secret.
|
||||
_INVALID_ASSERTION_AADSTS_PREFIX: Final = "50027"
|
||||
_AADSTS_CODES_ADAPTER: Final = TypeAdapter(tuple[int, ...])
|
||||
GATEWAY_SCOPE_TEMPLATE: Final = "api://{client_id}/access_as_user"
|
||||
_MCP_CALL_TYPES: Final[tuple[str, ...]] = ("mcp_call", "call_mcp_tool")
|
||||
_TOOL_INPUT_SCHEMA_ADAPTER: Final = TypeAdapter(dict[str, object])
|
||||
_OBO_CACHE_MAX_ENTRIES: Final = 1000
|
||||
_DEFAULT_TOKEN_TTL_SECONDS: Final = 3599.0
|
||||
_TOKEN_EXPIRY_SLACK_SECONDS: Final = 60.0
|
||||
|
||||
|
||||
def _parse_expires_in(raw: object) -> float:
|
||||
if not isinstance(raw, (int, float, str)):
|
||||
return _DEFAULT_TOKEN_TTL_SECONDS
|
||||
try:
|
||||
return float(raw)
|
||||
except ValueError:
|
||||
return _DEFAULT_TOKEN_TTL_SECONDS
|
||||
|
||||
|
||||
def _parse_aadsts_codes(raw: object) -> tuple[int, ...]:
|
||||
try:
|
||||
return _AADSTS_CODES_ADAPTER.validate_python(raw)
|
||||
except ValidationError:
|
||||
return ()
|
||||
_JSON_OBJECT_ADAPTER: Final = TypeAdapter(dict[str, object])
|
||||
|
||||
|
||||
def _parse_tool_input_schema(raw: object) -> Mapping[str, object] | None:
|
||||
try:
|
||||
return _TOOL_INPUT_SCHEMA_ADAPTER.validate_python(raw)
|
||||
return _JSON_OBJECT_ADAPTER.validate_python(raw)
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
|
@ -108,6 +104,15 @@ class _EvaluateResponse(TypedDict, total=False):
|
|||
correlationId: ReadOnly[str]
|
||||
|
||||
|
||||
class _AdmissionMetadata(TypedDict):
|
||||
user_api_key_metadata: ReadOnly[dict | None]
|
||||
user_api_key_team_metadata: ReadOnly[dict | None]
|
||||
|
||||
|
||||
class _AdmissionProbe(TypedDict):
|
||||
metadata: ReadOnly[_AdmissionMetadata]
|
||||
|
||||
|
||||
class _ToolReference(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
|
|
@ -130,20 +135,12 @@ class _BlockedDetail(TypedDict):
|
|||
|
||||
|
||||
class Agent365TokenExchangeError(Exception):
|
||||
def __init__(self, status_code: int, error_code: str, description: str, aadsts_codes: tuple[int, ...] = ()) -> None:
|
||||
super().__init__(f"{error_code}: {description}")
|
||||
self.status_code = status_code
|
||||
self.error_code = error_code
|
||||
self.description = description
|
||||
self.aadsts_codes = aadsts_codes
|
||||
"""Entra refused the gateway's own client credentials, scope or resource; the caller cannot fix that by
|
||||
signing in again, so it is the gateway's outage, never a 401."""
|
||||
|
||||
@property
|
||||
def gateway_owned(self) -> bool:
|
||||
"""Whether the gateway's own client credentials, scope or resource were refused, as opposed to the
|
||||
caller's assertion. The caller cannot fix a gateway-owned rejection by signing in again."""
|
||||
if self.error_code not in _GATEWAY_OWNED_TOKEN_ERRORS:
|
||||
return False
|
||||
return not any(str(code).startswith(_INVALID_ASSERTION_AADSTS_PREFIX) for code in self.aadsts_codes)
|
||||
def __init__(self, error_code: str) -> None:
|
||||
super().__init__(error_code)
|
||||
self.error_code = error_code
|
||||
|
||||
|
||||
class Agent365MalformedResponseError(Exception):
|
||||
|
|
@ -156,6 +153,13 @@ class Agent365ThrottledError(Exception):
|
|||
self.status_code = status_code
|
||||
|
||||
|
||||
def _gateway_fault_reason(error_code: str) -> str:
|
||||
return (
|
||||
f"Entra rejected the gateway's own Agent 365 credentials ({error_code}); "
|
||||
"check the guardrail's client_id and client_secret"
|
||||
)
|
||||
|
||||
|
||||
class Agent365Guardrail(CustomGuardrail):
|
||||
"""Pre-MCP-call guardrail enforcing Microsoft Agent 365 tool-evaluation verdicts.
|
||||
|
||||
|
|
@ -173,6 +177,7 @@ class Agent365Guardrail(CustomGuardrail):
|
|||
request_timeout: float = 10.0,
|
||||
unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed",
|
||||
async_handler: AsyncHTTPHandler | None = None,
|
||||
token_exchanger: TokenExchanger | None = None,
|
||||
**kwargs, # noqa: ANN003 # kwargs-ok: forwarded verbatim to CustomGuardrail (event_hook, default_on)
|
||||
) -> None:
|
||||
super().__init__(
|
||||
|
|
@ -192,8 +197,23 @@ class Agent365Guardrail(CustomGuardrail):
|
|||
self.async_handler = async_handler or get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.GuardrailCallback
|
||||
)
|
||||
self._obo_token_cache: OrderedDict[str, tuple[str, float]] = OrderedDict() # mutable-ok: lock-guarded LRU
|
||||
self._obo_cache_lock = threading.Lock()
|
||||
self._exchange_config: Final = TokenExchangeConfig(
|
||||
profile="entra_obo",
|
||||
token_exchange_endpoint=TOKEN_ENDPOINT_TEMPLATE.format(tenant_id=tenant_id),
|
||||
client_id=client_id,
|
||||
client_secret=SecretStr(client_secret),
|
||||
scopes=(OBO_SCOPE,),
|
||||
)
|
||||
self._exchange_server: Final = ServerSpec(
|
||||
server_id=f"agent-365:{tenant_id}",
|
||||
resource=AGENT_365_PROD_API_BASE,
|
||||
config=self._exchange_config,
|
||||
)
|
||||
self._token_exchanger: Final = (
|
||||
token_exchanger
|
||||
if token_exchanger is not None
|
||||
else build_token_exchanger(post=self._post_entra_token_endpoint)
|
||||
)
|
||||
verbose_proxy_logger.info("Initialized Microsoft Agent 365 guardrail: %s", guardrail_name)
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -220,7 +240,7 @@ class Agent365Guardrail(CustomGuardrail):
|
|||
return data
|
||||
|
||||
tool_name: Final = str(data.get("mcp_tool_name") or "")
|
||||
assertion: Final = entra_assertion(data.get("incoming_bearer_token"))
|
||||
assertion: Final = entra_assertion(data.get("incoming_subject_token"))
|
||||
if assertion is None:
|
||||
self._handle_caller_fault(
|
||||
data=data,
|
||||
|
|
@ -233,22 +253,10 @@ class Agent365Guardrail(CustomGuardrail):
|
|||
)
|
||||
|
||||
try:
|
||||
obo_token: Final = await self._get_obo_token(assertion)
|
||||
exchange_result: Final = await self._exchange_caller_assertion(assertion)
|
||||
except Agent365TokenExchangeError as exc:
|
||||
if exc.gateway_owned:
|
||||
return self._handle_unavailable(
|
||||
data=data,
|
||||
tool_name=tool_name,
|
||||
reason=(
|
||||
f"Entra rejected the gateway's own Agent 365 credentials ({exc.error_code}); "
|
||||
"check the guardrail's client_id and client_secret"
|
||||
),
|
||||
)
|
||||
self._handle_caller_fault(
|
||||
data=data,
|
||||
tool_name=tool_name,
|
||||
status_code=401,
|
||||
reason=f"the Entra On-Behalf-Of token exchange was rejected ({exc.error_code})",
|
||||
return self._handle_unavailable(
|
||||
data=data, tool_name=tool_name, reason=_gateway_fault_reason(exc.error_code)
|
||||
)
|
||||
except Agent365ThrottledError as exc:
|
||||
self._handle_throttled(
|
||||
|
|
@ -264,11 +272,25 @@ class Agent365Guardrail(CustomGuardrail):
|
|||
reason=f"the Entra token endpoint could not be reached ({type(exc).__name__})",
|
||||
)
|
||||
except Agent365MalformedResponseError as exc:
|
||||
return self._handle_unavailable(
|
||||
data=data,
|
||||
tool_name=tool_name,
|
||||
reason=str(exc),
|
||||
)
|
||||
return self._handle_unavailable(data=data, tool_name=tool_name, reason=str(exc))
|
||||
match exchange_result:
|
||||
case Ok(token):
|
||||
obo_token: Final = token.access_token
|
||||
case Error(error):
|
||||
match error.tag:
|
||||
case "unauthorized":
|
||||
self._handle_caller_fault(
|
||||
data=data,
|
||||
tool_name=tool_name,
|
||||
status_code=401,
|
||||
reason=f"the Entra On-Behalf-Of token exchange was rejected ({error.unauthorized.detail})",
|
||||
)
|
||||
case _:
|
||||
return self._handle_unavailable(
|
||||
data=data,
|
||||
tool_name=tool_name,
|
||||
reason=f"the Entra token exchange failed ({error.summary})",
|
||||
)
|
||||
|
||||
start: Final = time.perf_counter()
|
||||
try:
|
||||
|
|
@ -284,14 +306,14 @@ class Agent365Guardrail(CustomGuardrail):
|
|||
reason=f"the Agent 365 endpoint could not be reached ({type(exc).__name__})",
|
||||
)
|
||||
latency_ms: Final = (time.perf_counter() - start) * 1000.0
|
||||
fallback: Final = self._handle_evaluate_error(
|
||||
fallback: Final = await self._handle_evaluate_error(
|
||||
data=data, tool_name=tool_name, assertion=assertion, response=response, latency_ms=latency_ms
|
||||
)
|
||||
if fallback is not None:
|
||||
return fallback
|
||||
return self._enforce_verdict(data=data, tool_name=tool_name, response=response, latency_ms=latency_ms)
|
||||
|
||||
def _handle_evaluate_error(
|
||||
async def _handle_evaluate_error(
|
||||
self,
|
||||
data: dict, # mutable-ok: guardrail logging appends into the request metadata in place
|
||||
tool_name: str,
|
||||
|
|
@ -308,7 +330,9 @@ class Agent365Guardrail(CustomGuardrail):
|
|||
)
|
||||
if 400 <= response.status_code < 500:
|
||||
if response.status_code == 401:
|
||||
self._evict_obo_token(assertion)
|
||||
await self._token_exchanger.invalidate(
|
||||
assertion, self._exchange_server, self._exchange_config, tenant_id=self.tenant_id
|
||||
)
|
||||
self._record_verdict(
|
||||
data=data,
|
||||
verdict="Rejected",
|
||||
|
|
@ -455,26 +479,54 @@ class Agent365Guardrail(CustomGuardrail):
|
|||
return call_id
|
||||
return str(uuid.uuid4())
|
||||
|
||||
async def _get_obo_token(self, assertion: str) -> str:
|
||||
cache_key: Final = hashlib.sha256(assertion.encode("utf-8")).hexdigest()
|
||||
now: Final = time.time()
|
||||
with self._obo_cache_lock:
|
||||
cached: Final = self._obo_token_cache.get(cache_key)
|
||||
if cached and cached[1] > now + _TOKEN_EXPIRY_SLACK_SECONDS:
|
||||
self._obo_token_cache.move_to_end(cache_key)
|
||||
return cached[0]
|
||||
def caller_sign_in(self, server: MCPServer, user_api_key_auth: "UserAPIKeyAuth | None") -> CallerSignIn | None:
|
||||
"""The Entra sign-in this guardrail requires of callers: only a ``default_on`` guardrail whose mode gates a
|
||||
tagless MCP connect and that the caller's key or team has not opted out of, because the anonymous
|
||||
metadata fetch that follows a challenge cannot see which key selected a guardrail and would advertise
|
||||
the wrong issuer. Only servers that leave the caller's top-level ``Authorization`` with the gateway
|
||||
qualify: a forwarded API-key header travels upstream in its own slot and does not displace the Entra
|
||||
assertion."""
|
||||
if not (self.default_on and server.keeps_caller_authorization):
|
||||
return None
|
||||
probe: Final[_AdmissionProbe] = {
|
||||
"metadata": {
|
||||
"user_api_key_metadata": user_api_key_auth.metadata if user_api_key_auth else None, # pyright: ignore[reportUnknownMemberType] # UserAPIKeyAuth.metadata is a raw dict
|
||||
"user_api_key_team_metadata": user_api_key_auth.team_metadata if user_api_key_auth else None, # pyright: ignore[reportUnknownMemberType] # UserAPIKeyAuth.team_metadata is a raw dict
|
||||
}
|
||||
}
|
||||
if (
|
||||
self.should_run_guardrail( # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # should_run_guardrail takes an untyped data dict
|
||||
data=probe, event_type=GuardrailEventHooks.pre_mcp_call
|
||||
)
|
||||
is not True
|
||||
):
|
||||
return None
|
||||
return CallerSignIn(
|
||||
issuers=(ENTRA_ISSUER_TEMPLATE.format(tenant_id=self.tenant_id),),
|
||||
scopes=tuple(server.scopes)
|
||||
if server.scopes
|
||||
else (GATEWAY_SCOPE_TEMPLATE.format(client_id=self.client_id),),
|
||||
)
|
||||
|
||||
async def _exchange_caller_assertion(self, assertion: str) -> Result[OAuthToken, CredError]:
|
||||
"""The one Entra OBO exchange call both the tool-call path and the connect preflight run; a
|
||||
successful result is cached by the exchanger, so the session reuses what the preflight minted."""
|
||||
return await self._token_exchanger.exchange(
|
||||
assertion, self._exchange_server, self._exchange_config, tenant_id=self.tenant_id
|
||||
)
|
||||
|
||||
async def _post_entra_token_endpoint(
|
||||
self,
|
||||
url: str,
|
||||
form: dict[str, str], # mutable-ok: ExchangeHttpPost contract
|
||||
client_auth_headers: dict[str, str], # mutable-ok: ExchangeHttpPost contract
|
||||
) -> dict[str, object] | None:
|
||||
"""The exchanger's HTTP edge for this guardrail: every way Entra can fail keeps its own exception, so
|
||||
the verdict reason the Logs row carries names the OAuth error code or the transport fault."""
|
||||
response: Final = await self._post_allowing_error_status(
|
||||
url=TOKEN_ENDPOINT_TEMPLATE.format(tenant_id=self.tenant_id),
|
||||
data={
|
||||
"grant_type": "urn:ietf:params:oauth:grant-type:jwt-bearer",
|
||||
"client_id": self.client_id,
|
||||
"client_secret": self.client_secret,
|
||||
"assertion": assertion,
|
||||
"scope": OBO_SCOPE,
|
||||
"requested_token_use": "on_behalf_of",
|
||||
},
|
||||
headers={"Content-Type": "application/x-www-form-urlencoded"},
|
||||
url=url,
|
||||
data=form,
|
||||
headers={"Content-Type": "application/x-www-form-urlencoded", **client_auth_headers},
|
||||
)
|
||||
if response.status_code in (408, 429):
|
||||
raise Agent365ThrottledError(status_code=response.status_code)
|
||||
|
|
@ -485,32 +537,54 @@ class Agent365Guardrail(CustomGuardrail):
|
|||
response=response,
|
||||
)
|
||||
try:
|
||||
parsed_body: Final = response.json()
|
||||
parsed_body: Final[object] = response.json()
|
||||
except ValueError as exc:
|
||||
raise Agent365MalformedResponseError("the Entra token endpoint returned a non-JSON body") from exc
|
||||
if not isinstance(parsed_body, dict):
|
||||
raise Agent365MalformedResponseError("the Entra token endpoint returned a non-object JSON body")
|
||||
body: Final = parsed_body
|
||||
try:
|
||||
body: Final = _JSON_OBJECT_ADAPTER.validate_python(parsed_body)
|
||||
except ValidationError as exc:
|
||||
raise Agent365MalformedResponseError("the Entra token endpoint returned a non-object JSON body") from exc
|
||||
if response.status_code >= 400:
|
||||
raise Agent365TokenExchangeError(
|
||||
status_code=response.status_code,
|
||||
error_code=str(body.get("error", "invalid_grant")),
|
||||
description=str(body.get("error_description", ""))[:512],
|
||||
aadsts_codes=_parse_aadsts_codes(body.get("error_codes")),
|
||||
)
|
||||
oauth_error: Final = oauth_error_fields(response)
|
||||
gateway_fault: Final = oauth_error.gateway_fault
|
||||
if gateway_fault is not None:
|
||||
raise Agent365TokenExchangeError(error_code=gateway_fault)
|
||||
raise SubjectTokenRejected(oauth_error.error or "invalid_grant", claims=oauth_error.claims)
|
||||
if "access_token" not in body:
|
||||
raise Agent365MalformedResponseError("the Entra token endpoint returned no access_token")
|
||||
raw_access_token: Final = body.get("access_token")
|
||||
if not isinstance(raw_access_token, str) or not raw_access_token:
|
||||
access_token: Final = body["access_token"]
|
||||
if not isinstance(access_token, str) or not access_token:
|
||||
raise Agent365MalformedResponseError("the Entra token endpoint returned a non-string access_token")
|
||||
access_token: Final = raw_access_token
|
||||
expires_at: Final = time.time() + _parse_expires_in(body.get("expires_in", 3599))
|
||||
with self._obo_cache_lock:
|
||||
self._obo_token_cache[cache_key] = (access_token, expires_at)
|
||||
self._obo_token_cache.move_to_end(cache_key)
|
||||
while len(self._obo_token_cache) > _OBO_CACHE_MAX_ENTRIES:
|
||||
self._obo_token_cache.popitem(last=False)
|
||||
return access_token
|
||||
return body
|
||||
|
||||
async def preflight_caller_sign_in(
|
||||
self, server: MCPServer, user_api_key_auth: "UserAPIKeyAuth | None", subject_token: str
|
||||
) -> CallerSignInPreflight:
|
||||
"""The connect-time check the preemptive gate runs: a bearer Entra rejects, or one it could never
|
||||
accept, gets the sign-in challenge here, where ``WWW-Authenticate`` still reaches the client, instead of
|
||||
surfacing as a JSON-RPC error on every tools/call. ``subject_token=None`` stays the challenge gate's job."""
|
||||
assertion: Final = entra_assertion(subject_token)
|
||||
if assertion is None:
|
||||
return Rejected(detail="the caller's bearer is not an Entra token; sign in with Entra and retry")
|
||||
fail_open: Final = self.unreachable_fallback == "fail_open"
|
||||
try:
|
||||
exchange_result: Final = await self._exchange_caller_assertion(assertion)
|
||||
except Agent365TokenExchangeError as exc:
|
||||
return Unavailable(detail=_gateway_fault_reason(exc.error_code), fail_open=fail_open)
|
||||
except Agent365ThrottledError as exc:
|
||||
return Unavailable(detail=f"the Entra token endpoint returned HTTP {exc.status_code}", fail_open=False)
|
||||
except (httpx.HTTPError, LitellmTimeout, TimeoutError) as exc:
|
||||
return Unavailable(
|
||||
detail=f"the Entra token endpoint could not be reached ({type(exc).__name__})", fail_open=fail_open
|
||||
)
|
||||
except Agent365MalformedResponseError as exc:
|
||||
return Unavailable(detail=str(exc), fail_open=fail_open)
|
||||
if isinstance(exchange_result, Ok):
|
||||
return SignedIn()
|
||||
error: Final = exchange_result.error
|
||||
if error.tag == "unauthorized":
|
||||
return Rejected(detail=error.unauthorized.detail, claims=error.unauthorized.claims)
|
||||
return Unavailable(detail=f"the Entra token exchange failed ({error.summary})", fail_open=fail_open)
|
||||
|
||||
async def _post_allowing_error_status(
|
||||
self,
|
||||
|
|
@ -577,11 +651,6 @@ class Agent365Guardrail(CustomGuardrail):
|
|||
}
|
||||
raise HTTPException(status_code=503, detail=throttled_detail)
|
||||
|
||||
def _evict_obo_token(self, assertion: str) -> None:
|
||||
cache_key: Final = hashlib.sha256(assertion.encode("utf-8")).hexdigest()
|
||||
with self._obo_cache_lock:
|
||||
self._obo_token_cache.pop(cache_key, None)
|
||||
|
||||
def _handle_unavailable(
|
||||
self,
|
||||
data: dict, # mutable-ok: guardrail logging appends into the request metadata in place
|
||||
|
|
|
|||
|
|
@ -1518,6 +1518,7 @@ class ProxyLogging:
|
|||
# (e.g. MCPJWTSigner) to independently verify the caller's identity
|
||||
# before re-signing an outbound token (FR-5 verify+re-sign).
|
||||
"incoming_bearer_token": kwargs.get("incoming_bearer_token"),
|
||||
"incoming_subject_token": kwargs.get("incoming_subject_token"),
|
||||
"metadata": synthetic_metadata,
|
||||
}
|
||||
user_api_key_auth: Final = kwargs.get("user_api_key_auth")
|
||||
|
|
|
|||
|
|
@ -315,10 +315,10 @@ class MCPServer(BaseModel):
|
|||
return self.per_server_oauth_discovery and self.auth_type == MCPAuth.oauth2 and not self.has_client_credentials
|
||||
|
||||
@property
|
||||
def advertises_gateway_authorization_server(self) -> bool:
|
||||
"""Whether named discovery should advertise the aggregate gateway authorization server."""
|
||||
if self.auth_type == MCPAuth.oauth2:
|
||||
return self.is_gateway_managed_oauth2 and not self.uses_per_server_oauth_relay
|
||||
def keeps_caller_authorization(self) -> bool:
|
||||
"""Whether the caller's top-level ``Authorization`` stays with the gateway: the server neither relays
|
||||
it upstream nor runs an OAuth mode that fills that slot itself, so a gateway guardrail may consume it
|
||||
as the caller's own assertion. Forwarding a separate API-key header leaves the slot untouched."""
|
||||
if self.auth_type not in (
|
||||
None,
|
||||
MCPAuth.none,
|
||||
|
|
@ -328,11 +328,20 @@ class MCPServer(BaseModel):
|
|||
MCPAuth.authorization,
|
||||
MCPAuth.token,
|
||||
MCPAuth.aws_sigv4,
|
||||
MCPAuth.oauth2_token_exchange,
|
||||
):
|
||||
return False
|
||||
return not any(
|
||||
header.lower() in ("authorization", "x-api-key", "api-key", "apikey")
|
||||
for header in (self.extra_headers or ())
|
||||
return not any(header.lower() == "authorization" for header in (self.extra_headers or ()))
|
||||
|
||||
@property
|
||||
def advertises_gateway_authorization_server(self) -> bool:
|
||||
"""Whether named discovery should advertise the aggregate gateway authorization server."""
|
||||
if self.auth_type == MCPAuth.oauth2:
|
||||
return self.is_gateway_managed_oauth2 and not self.uses_per_server_oauth_relay
|
||||
if self.auth_type == MCPAuth.oauth2_token_exchange:
|
||||
return False
|
||||
return self.keeps_caller_authorization and not any(
|
||||
header.lower() in ("x-api-key", "api-key", "apikey") for header in (self.extra_headers or ())
|
||||
)
|
||||
|
||||
@property
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import json
|
||||
import uuid
|
||||
from typing import Final
|
||||
|
||||
|
|
@ -92,6 +93,36 @@ def test_subject_grant_lists_only_reachable_tools_and_denies_the_rest(
|
|||
assert not any(name.startswith(denied_alias) for name in denied_listed.tools), denied_listed.tools
|
||||
|
||||
|
||||
@pytest.mark.parametrize("entry", ("mcp", "server_mcp"))
|
||||
def test_access_group_named_like_an_ungranted_server_still_routes_the_groups_servers(
|
||||
gateway: Gateway, entry: EntryPoint
|
||||
) -> None:
|
||||
with peer_of("http") as shadow, peer_of("http") as member, gateway.scenario() as scenario:
|
||||
group: Final = "docs" + uuid.uuid4().hex[:8]
|
||||
member_alias: Final = "mem" + uuid.uuid4().hex[:8]
|
||||
register_mcp(scenario, shadow, group)
|
||||
register_mcp(scenario, member, member_alias, mcp_access_groups=[group])
|
||||
key: Final = scenario.key(object_permission={"mcp_access_groups": [group]})
|
||||
unmatched: Final = "none" + uuid.uuid4().hex[:8]
|
||||
denied: Final = McpCaller(
|
||||
gateway, key, entry, unmatched, {"x-mcp-servers": unmatched} if entry == "mcp" else {}
|
||||
).list_tools()
|
||||
assert (denied.status, json.loads(denied.error or "null"), denied.tools) == (
|
||||
(
|
||||
200,
|
||||
{"code": -32600, "message": f"The key is not allowed to access the requested MCP servers: {unmatched}"},
|
||||
(),
|
||||
)
|
||||
if entry == "mcp"
|
||||
else (404, {"detail": f"MCP server, toolset, or access group '{unmatched}' not found"}, ())
|
||||
), denied.raw
|
||||
selection: Final = {"x-mcp-servers": group} if entry == "mcp" else {}
|
||||
listed: Final = McpCaller(gateway, key, entry, group, selection).list_tools()
|
||||
assert listed.ok, listed.raw
|
||||
assert set(listed.tools) == {f"{member_alias}-{tool}" for tool in ("add", "multiply", "fail")}, listed.tools
|
||||
assert tool_calls(shadow.drain()) == ()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("entry", ("mcp", "server_mcp", "rest"))
|
||||
def test_key_without_any_grant_sees_no_scoped_server(gateway: Gateway, entry: EntryPoint) -> None:
|
||||
with peer_of("http") as peer, gateway.scenario() as scenario:
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ from integration._support.client import Gateway, eventually
|
|||
from integration._support.database import read_rows
|
||||
from integration._support.mcp import (
|
||||
ENTRY_POINTS,
|
||||
INITIALIZE,
|
||||
EntryPoint,
|
||||
McpCaller,
|
||||
McpPeer,
|
||||
|
|
@ -38,6 +39,10 @@ GUARDRAIL_ROWS: Final = (
|
|||
FALLBACKS: Final = (None, "fail_open", "fail_closed")
|
||||
|
||||
|
||||
def _origin(gateway: Gateway) -> str:
|
||||
return str(gateway.client.base_url).rstrip("/")
|
||||
|
||||
|
||||
def _generic_guardrail_outage(request: Request) -> Reply:
|
||||
assert request.target == "/beta/litellm_basic_guardrail_api", request.target
|
||||
return Reply(status=503, body=json.dumps({"error": "synthetic sibling guardrail outage"}).encode())
|
||||
|
|
@ -136,15 +141,26 @@ def test_a_missing_or_malformed_caller_bearer_blocks_on_every_entry_point_whatev
|
|||
) -> None:
|
||||
with _rig(gateway, tmp_path, fallback) as rig:
|
||||
assert f"{rig.alias}-add" in rig.caller().list_tools().tools, "the catalog needs only the virtual key"
|
||||
metadata_url: Final = f"{_origin(rig.candidate)}/.well-known/oauth-protected-resource/{rig.alias}/mcp"
|
||||
for entry in ENTRY_POINTS:
|
||||
missing: Final = rig.caller(entry).call(f"{rig.alias}-add", {"entry": entry}, server_id=rig.server_id)
|
||||
assert missing.error is not None and REJECTED in missing.raw, f"{entry} without a bearer: {missing.raw}"
|
||||
malformed: Final = rig.caller(entry, "not-a-jws").call(
|
||||
f"{rig.alias}-add", {"entry": entry}, server_id=rig.server_id
|
||||
)
|
||||
assert malformed.error is not None and REJECTED in malformed.raw, f"{entry} opaque bearer: {malformed.raw}"
|
||||
for label, bearer in (("without a bearer", None), ("opaque bearer", "not-a-jws")):
|
||||
caller: Final = rig.caller(entry, bearer)
|
||||
if entry == "server_mcp":
|
||||
challenged: Final = caller.rpc("initialize", INITIALIZE)
|
||||
assert challenged.status_code == 401, f"{entry} {label}: {challenged.text}"
|
||||
authenticate: Final = challenged.headers.get("www-authenticate", "")
|
||||
assert f'resource_metadata="{metadata_url}"' in authenticate, f"{entry} {label}: {authenticate!r}"
|
||||
assert 'error="invalid_token"' in authenticate, f"{entry} {label}: {authenticate!r}"
|
||||
if bearer is None:
|
||||
unsigned: Final = caller.rpc(
|
||||
"tools/call", {"name": f"{rig.alias}-add", "arguments": {"entry": entry}}
|
||||
)
|
||||
assert unsigned.status_code == 401, f"{entry} {label}: {unsigned.text}"
|
||||
continue
|
||||
outcome: Final = caller.call(f"{rig.alias}-add", {"entry": entry}, server_id=rig.server_id)
|
||||
assert outcome.error is not None and REJECTED in outcome.raw, f"{entry} {label}: {outcome.raw}"
|
||||
assert rig.upstream_tool_names() == ()
|
||||
expected: Final = 2 * len(ENTRY_POINTS)
|
||||
expected: Final = 2 * len(ENTRY_POINTS) - 1
|
||||
assert rig.guardrail_statuses("call_mcp_tool", expected) == ["guardrail_intervened"] * expected
|
||||
|
||||
|
||||
|
|
|
|||
463
tests/integration/mcp/test_mcp_caller_sign_in.py
Normal file
463
tests/integration/mcp/test_mcp_caller_sign_in.py
Normal file
|
|
@ -0,0 +1,463 @@
|
|||
import json
|
||||
import uuid
|
||||
from collections.abc import Mapping
|
||||
from pathlib import Path
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import yaml
|
||||
from cryptography.hazmat.primitives.asymmetric import rsa
|
||||
from integration._support.client import Gateway
|
||||
from integration._support.mcp import (
|
||||
INITIALIZE,
|
||||
McpCaller,
|
||||
forget_mcp,
|
||||
mcp_peer,
|
||||
register_mcp,
|
||||
tool_calls,
|
||||
)
|
||||
from integration._support.process import owned_proxy
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from jwt import algorithms as jwt_algorithms
|
||||
|
||||
ADD: Final = {"a": 2, "b": 3}
|
||||
ACCEPT: Final = {"Accept": "application/json, text/event-stream"}
|
||||
|
||||
|
||||
def _rpc(
|
||||
gateway: Gateway,
|
||||
path: str,
|
||||
key: str,
|
||||
headers: dict[str, str],
|
||||
method: str = "initialize",
|
||||
params: Mapping[str, object] = INITIALIZE,
|
||||
) -> httpx.Response:
|
||||
return gateway.client.post(
|
||||
path,
|
||||
json={"jsonrpc": "2.0", "id": 1, "method": method, "params": dict(params)},
|
||||
headers={"x-litellm-api-key": key, **ACCEPT, **headers},
|
||||
)
|
||||
|
||||
|
||||
def _sse_data(response: httpx.Response) -> str:
|
||||
return next(line[5:].strip() for line in response.text.splitlines() if line.startswith("data:"))
|
||||
|
||||
|
||||
def _advertised(gateway: Gateway, segment: str) -> tuple[int, tuple[str, ...], object]:
|
||||
response: Final = gateway.client.get(f"/.well-known/oauth-protected-resource/mcp/{segment}")
|
||||
document: Final = response.json()
|
||||
issuers: Final = tuple(
|
||||
str(issuer).removesuffix(f"/{segment}") for issuer in document.get("authorization_servers", ())
|
||||
)
|
||||
return response.status_code, issuers, document.get("scopes_supported")
|
||||
|
||||
|
||||
def _origin(gateway: Gateway) -> str:
|
||||
return str(gateway.client.base_url).rstrip("/")
|
||||
|
||||
|
||||
def _sign_in_config(
|
||||
guardrail_params: dict[str, object], path: Path, general_settings: Mapping[str, object] = MappingProxyType({})
|
||||
) -> Path:
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["guardrails"] = [{"guardrail_name": "signin" + uuid.uuid4().hex, "litellm_params": guardrail_params}]
|
||||
config["general_settings"] = {**config.get("general_settings", {}), **general_settings}
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
return path
|
||||
|
||||
|
||||
def test_multi_server_connect_with_a_litellm_key_in_authorization_admits_an_obo_server(
|
||||
gateway: Gateway,
|
||||
) -> None:
|
||||
with mcp_peer() as obo_peer, mcp_peer() as math_peer, gateway.scenario() as scenario:
|
||||
obo_alias: Final = "obo" + uuid.uuid4().hex[:8]
|
||||
math_alias: Final = "math" + uuid.uuid4().hex[:8]
|
||||
obo_id: Final = register_mcp(
|
||||
scenario,
|
||||
obo_peer,
|
||||
obo_alias,
|
||||
auth_type="oauth2_token_exchange",
|
||||
token_exchange_endpoint="http://127.0.0.1:9/token",
|
||||
credentials={"client_id": "obo-client", "client_secret": "obo-secret"},
|
||||
)
|
||||
math_id: Final = register_mcp(scenario, math_peer, math_alias)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [obo_id, math_id]})
|
||||
caller: Final = McpCaller(
|
||||
gateway,
|
||||
key,
|
||||
"root",
|
||||
headers={
|
||||
"x-mcp-servers": f"{obo_alias},{math_alias}",
|
||||
"Authorization": f"Bearer {key}",
|
||||
},
|
||||
)
|
||||
init: Final = caller.initialize()
|
||||
assert init.ok, init.raw
|
||||
listed: Final = caller.list_tools()
|
||||
assert listed.ok, listed.raw
|
||||
assert any(name.endswith("add") for name in listed.tools), listed.tools
|
||||
|
||||
|
||||
def test_alias_first_lookup_wins_over_a_server_whose_name_matches_the_alias(gateway: Gateway, tmp_path: Path) -> None:
|
||||
with mcp_peer() as math_peer, mcp_peer() as obo_peer:
|
||||
name: Final = "gh" + uuid.uuid4().hex[:6]
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["mcp_servers"] = {
|
||||
name: {
|
||||
"transport": "http",
|
||||
"url": obo_peer.url,
|
||||
"auth_type": "oauth2_token_exchange",
|
||||
"token_exchange_endpoint": "http://127.0.0.1:9/token",
|
||||
"credentials": {"client_id": "obo-client", "client_secret": "obo-secret"},
|
||||
}
|
||||
}
|
||||
path: Final = tmp_path / "collision.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
with (
|
||||
owned_proxy(gateway, tmp_path, {}, config=path) as candidate,
|
||||
candidate.scenario() as scenario,
|
||||
):
|
||||
second: Final = candidate.request(
|
||||
"POST",
|
||||
"/v1/mcp/server",
|
||||
{"server_name": name + "_public", "alias": name, **math_peer.registration()},
|
||||
)
|
||||
assert second.status_code == 201, second.text
|
||||
second_id: Final = second.json()["server_id"]
|
||||
scenario.cleanups.callback(forget_mcp, candidate, second_id)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [second_id]})
|
||||
|
||||
init: Final = _rpc(candidate, f"/mcp/{name}", key, {})
|
||||
assert init.status_code == 200, init.text
|
||||
listing: Final = candidate.client.post(
|
||||
f"/mcp/{name}",
|
||||
json={"jsonrpc": "2.0", "id": 2, "method": "tools/list", "params": {}},
|
||||
headers={"x-litellm-api-key": key, **ACCEPT},
|
||||
)
|
||||
assert listing.status_code == 200, listing.text
|
||||
body: Final = json.loads(
|
||||
next(line[5:].strip() for line in listing.text.splitlines() if line.startswith("data:"))
|
||||
)
|
||||
names: Final = {tool["name"] for tool in body["result"]["tools"]}
|
||||
assert any(name.endswith("add") for name in names), names
|
||||
|
||||
|
||||
def test_exact_name_wins_over_a_case_folded_config_alias_for_connect_discovery_and_calls(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
with mcp_peer() as obo_peer, mcp_peer() as math_peer:
|
||||
stem: Final = "gh" + uuid.uuid4().hex[:6]
|
||||
cased: Final = stem.capitalize()
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["mcp_servers"] = {
|
||||
stem + "_obo": {
|
||||
"alias": stem,
|
||||
"transport": "http",
|
||||
"url": obo_peer.url,
|
||||
"auth_type": "oauth2_token_exchange",
|
||||
"token_exchange_endpoint": "http://127.0.0.1:9/token",
|
||||
"credentials": {"client_id": "obo-client", "client_secret": "obo-secret"},
|
||||
}
|
||||
}
|
||||
path: Final = tmp_path / "case-collision.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario:
|
||||
by_name: Final = register_mcp(scenario, math_peer, cased)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [by_name]})
|
||||
|
||||
init: Final = _rpc(candidate, f"/mcp/{cased}", key, {})
|
||||
assert init.status_code == 200, init.text
|
||||
listed: Final = _rpc(candidate, f"/mcp/{cased}", key, {}, method="tools/list")
|
||||
assert listed.status_code == 200, listed.text
|
||||
tools: Final = json.loads(_sse_data(listed))["result"]["tools"]
|
||||
add: Final = next(tool["name"] for tool in tools if tool["name"].endswith("add"))
|
||||
called: Final = _rpc(
|
||||
candidate, f"/mcp/{cased}", key, {}, method="tools/call", params={"name": add, "arguments": ADD}
|
||||
)
|
||||
assert called.status_code == 200, called.text
|
||||
assert json.loads(_sse_data(called))["result"]["content"][0]["text"] == "5", called.text
|
||||
assert len(tool_calls(math_peer.drain())) == 1
|
||||
assert tool_calls(obo_peer.drain()) == ()
|
||||
assert _advertised(candidate, cased) == _advertised(candidate, by_name)
|
||||
|
||||
same_name: Final = _rpc(candidate, f"/mcp/{stem}", key, {})
|
||||
assert same_name.status_code == 200, same_name.text
|
||||
assert "www-authenticate" not in same_name.headers, same_name.headers
|
||||
via_alias: Final = _rpc(
|
||||
candidate, f"/mcp/{stem}", key, {}, method="tools/call", params={"name": add, "arguments": ADD}
|
||||
)
|
||||
assert json.loads(_sse_data(via_alias))["result"]["content"][0]["text"] == "5", via_alias.text
|
||||
assert len(tool_calls(math_peer.drain())) == 1
|
||||
assert tool_calls(obo_peer.drain()) == ()
|
||||
|
||||
challenged: Final = _rpc(candidate, f"/mcp/{stem}", scenario.key(), {})
|
||||
assert challenged.status_code == 401, challenged.text
|
||||
assert (
|
||||
f'resource_metadata="{_origin(candidate)}/.well-known/oauth-protected-resource/mcp/{stem}"'
|
||||
in challenged.headers.get("www-authenticate", "")
|
||||
)
|
||||
assert _advertised(candidate, stem) == _advertised(candidate, stem + "_obo")
|
||||
assert _advertised(candidate, stem) != _advertised(candidate, cased)
|
||||
|
||||
|
||||
def test_name_of_a_server_hidden_from_an_external_ip_does_not_reroute_to_a_case_variant(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
stem: Final = "gh" + uuid.uuid4().hex[:6]
|
||||
cased: Final = stem.capitalize()
|
||||
external: Final = {"X-Forwarded-For": "203.0.113.7"}
|
||||
with mcp_peer() as hidden_peer, mcp_peer() as public_peer:
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["general_settings"] = {
|
||||
**config.get("general_settings", {}),
|
||||
"use_x_forwarded_for": True,
|
||||
"mcp_trusted_proxy_ranges": ["127.0.0.0/8"],
|
||||
}
|
||||
config["mcp_servers"] = {
|
||||
stem: {"transport": "http", "url": hidden_peer.url, "available_on_public_internet": False}
|
||||
}
|
||||
path: Final = tmp_path / "external-ip.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario:
|
||||
public_id: Final = register_mcp(scenario, public_peer, cased)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [public_id]})
|
||||
|
||||
hidden_name: Final = _rpc(candidate, f"/mcp/{stem}", key, external)
|
||||
assert hidden_name.status_code == 403, hidden_name.text
|
||||
assert "www-authenticate" not in hidden_name.headers, hidden_name.headers
|
||||
assert (
|
||||
candidate.client.get(f"/.well-known/oauth-protected-resource/mcp/{stem}", headers=external).status_code
|
||||
== 404
|
||||
)
|
||||
assert _rpc(candidate, f"/mcp/{stem}", key, {}).status_code == 200
|
||||
|
||||
own_name: Final = _rpc(candidate, f"/mcp/{cased}", key, external)
|
||||
assert own_name.status_code == 200, own_name.text
|
||||
listed: Final = _rpc(candidate, f"/mcp/{cased}", key, external, method="tools/list")
|
||||
add: Final = next(
|
||||
tool["name"]
|
||||
for tool in json.loads(_sse_data(listed))["result"]["tools"]
|
||||
if tool["name"].endswith("add")
|
||||
)
|
||||
called: Final = _rpc(
|
||||
candidate, f"/mcp/{cased}", key, external, method="tools/call", params={"name": add, "arguments": ADD}
|
||||
)
|
||||
assert called.status_code == 200, called.text
|
||||
assert json.loads(_sse_data(called))["result"]["content"][0]["text"] == "5", called.text
|
||||
assert len(tool_calls(public_peer.drain())) == 1
|
||||
assert tool_calls(hidden_peer.drain()) == ()
|
||||
|
||||
|
||||
def test_jwt_signer_verifies_the_bearer_that_admitted_the_call(gateway: Gateway, tmp_path: Path) -> None:
|
||||
signer_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048)
|
||||
jwk: Final = json.loads(jwt_algorithms.RSAAlgorithm.to_jwk(signer_key.public_key()))
|
||||
jwk["kid"] = "idp"
|
||||
holder: Final = []
|
||||
|
||||
def idp(request: Request) -> Reply:
|
||||
issuer: Final = holder[0].url
|
||||
if request.target.endswith("/.well-known/openid-configuration"):
|
||||
return Reply(body=json.dumps({"issuer": issuer, "jwks_uri": issuer + "/jwks"}).encode())
|
||||
return Reply(body=json.dumps({"keys": [jwk]}).encode())
|
||||
|
||||
def introspect(request: Request) -> Reply:
|
||||
return Reply(body=b'{"active": true, "sub": "subject-1"}')
|
||||
|
||||
with (
|
||||
mcp_peer() as peer,
|
||||
wire_server(idp) as verify_idp,
|
||||
wire_server(introspect) as introspect_stub,
|
||||
gateway.scenario() as scenario,
|
||||
):
|
||||
holder.append(verify_idp)
|
||||
alias: Final = "sig" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias)
|
||||
config: Final = _sign_in_config(
|
||||
{
|
||||
"guardrail": "mcp_jwt_signer",
|
||||
"mode": "pre_mcp_call",
|
||||
"default_on": True,
|
||||
"access_token_discovery_uri": verify_idp.url + "/.well-known/openid-configuration",
|
||||
"token_introspection_endpoint": introspect_stub.url,
|
||||
"required_claims": ["sub"],
|
||||
},
|
||||
tmp_path / "signer.yaml",
|
||||
)
|
||||
with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as owned:
|
||||
key: Final = owned.key(object_permission={"mcp_servers": [identity]})
|
||||
|
||||
def call(headers: Mapping[str, str]) -> httpx.Response:
|
||||
return candidate.client.post(
|
||||
"/mcp-rest/tools/call",
|
||||
headers=dict(headers),
|
||||
json={"name": "add", "arguments": ADD, "server_id": identity},
|
||||
)
|
||||
|
||||
admitted: Final = call({"Authorization": f"Bearer {key}"})
|
||||
assert admitted.status_code == 200, admitted.text
|
||||
assert len(tool_calls(peer.drain())) == 1
|
||||
probes: Final = tuple(item for item in introspect_stub.drain() if item.body)
|
||||
assert any(key.encode() in probe.body for probe in probes), probes
|
||||
|
||||
split: Final = call({"x-litellm-api-key": key})
|
||||
assert split.status_code == 403, split.text
|
||||
assert introspect_stub.drain() == ()
|
||||
|
||||
|
||||
AGENT_365_PARAMS: Final = {
|
||||
"guardrail": "agent_365",
|
||||
"mode": "pre_mcp_call",
|
||||
"default_on": True,
|
||||
"tenant_id": "00000000-0000-0000-0000-000000000000",
|
||||
"client_id": "22222222-2222-2222-2222-222222222222",
|
||||
"client_secret": "secret",
|
||||
}
|
||||
ENTRA_ISSUER: Final = "https://login.microsoftonline.com/00000000-0000-0000-0000-000000000000/v2.0"
|
||||
GATEWAY_SCOPE: Final = "api://22222222-2222-2222-2222-222222222222/access_as_user"
|
||||
|
||||
|
||||
def test_agent_365_gated_server_challenges_at_connect_and_advertises_entra(gateway: Gateway, tmp_path: Path) -> None:
|
||||
config: Final = _sign_in_config(dict(AGENT_365_PARAMS), tmp_path / "agent365.yaml")
|
||||
with (
|
||||
owned_proxy(gateway, tmp_path, {}, config=config) as candidate,
|
||||
mcp_peer() as peer,
|
||||
candidate.scenario() as scenario,
|
||||
):
|
||||
alias: Final = "a365" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias)
|
||||
granted: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
denied: Final = scenario.key(object_permission={"mcp_servers": ["no-mcp-servers"]})
|
||||
|
||||
challenged: Final = _rpc(candidate, f"/mcp/{alias}", granted, {})
|
||||
assert challenged.status_code == 401, challenged.text
|
||||
authenticate: Final = challenged.headers.get("www-authenticate", "")
|
||||
assert (
|
||||
f'resource_metadata="{_origin(candidate)}/.well-known/oauth-protected-resource/mcp/{alias}"' in authenticate
|
||||
)
|
||||
assert 'error="invalid_token"' in authenticate
|
||||
|
||||
opaque: Final = _rpc(candidate, f"/mcp/{alias}", granted, {"Authorization": "Bearer not-a-jws"})
|
||||
assert opaque.status_code == 401, opaque.text
|
||||
assert (
|
||||
f'resource_metadata="{_origin(candidate)}/.well-known/oauth-protected-resource/mcp/{alias}"'
|
||||
in opaque.headers.get("www-authenticate", "")
|
||||
)
|
||||
|
||||
discovery: Final = candidate.client.get(f"/.well-known/oauth-protected-resource/mcp/{alias}")
|
||||
assert discovery.status_code == 200, discovery.text
|
||||
document: Final = discovery.json()
|
||||
assert document["authorization_servers"] == [ENTRA_ISSUER]
|
||||
assert document["scopes_supported"] == [GATEWAY_SCOPE]
|
||||
|
||||
refused: Final = _rpc(candidate, f"/mcp/{alias}", denied, {})
|
||||
assert refused.status_code == 403, refused.text
|
||||
assert "www-authenticate" not in refused.headers
|
||||
assert tool_calls(peer.drain()) == ()
|
||||
|
||||
|
||||
def test_agent_365_prm_advertises_the_servers_configured_scopes(gateway: Gateway, tmp_path: Path) -> None:
|
||||
config: Final = _sign_in_config(dict(AGENT_365_PARAMS), tmp_path / "agent365-scopes.yaml")
|
||||
with (
|
||||
owned_proxy(gateway, tmp_path, {}, config=config) as candidate,
|
||||
mcp_peer() as peer,
|
||||
candidate.scenario() as scenario,
|
||||
):
|
||||
scoped: Final = "a365" + uuid.uuid4().hex[:8]
|
||||
register_mcp(
|
||||
scenario,
|
||||
peer,
|
||||
scoped,
|
||||
credentials={"scopes": ["https://example/mcp/scoped/access_as_user", "offline_access"]},
|
||||
)
|
||||
unscoped: Final = "a365" + uuid.uuid4().hex[:8]
|
||||
register_mcp(scenario, peer, unscoped)
|
||||
|
||||
assert _advertised(candidate, scoped) == (
|
||||
200,
|
||||
(ENTRA_ISSUER,),
|
||||
["https://example/mcp/scoped/access_as_user", "offline_access"],
|
||||
)
|
||||
assert _advertised(candidate, unscoped) == (200, (ENTRA_ISSUER,), [GATEWAY_SCOPE])
|
||||
|
||||
|
||||
def test_challenge_names_the_server_first_route_the_client_connected_on(gateway: Gateway, tmp_path: Path) -> None:
|
||||
config: Final = _sign_in_config(dict(AGENT_365_PARAMS), tmp_path / "agent365-route.yaml")
|
||||
with (
|
||||
owned_proxy(gateway, tmp_path, {}, config=config) as candidate,
|
||||
mcp_peer() as peer,
|
||||
candidate.scenario() as scenario,
|
||||
):
|
||||
alias: Final = "a365" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias)
|
||||
granted: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
metadata_url: Final = f"{_origin(candidate)}/.well-known/oauth-protected-resource/{alias}/mcp"
|
||||
|
||||
challenged: Final = _rpc(candidate, f"/{alias}/mcp", granted, {})
|
||||
assert challenged.status_code == 401, challenged.text
|
||||
assert f'resource_metadata="{metadata_url}"' in challenged.headers.get("www-authenticate", "")
|
||||
|
||||
document: Final = httpx.get(metadata_url, timeout=15).json()
|
||||
assert document["resource"] == f"{_origin(candidate)}/{alias}/mcp", document
|
||||
assert document["authorization_servers"] == [ENTRA_ISSUER]
|
||||
assert tool_calls(peer.drain()) == ()
|
||||
|
||||
|
||||
def test_challenge_names_the_forwarded_origin_only_from_a_trusted_proxy(gateway: Gateway, tmp_path: Path) -> None:
|
||||
config: Final = _sign_in_config(
|
||||
dict(AGENT_365_PARAMS),
|
||||
tmp_path / "agent365-forwarded.yaml",
|
||||
{"use_x_forwarded_for": True, "mcp_trusted_proxy_ranges": ["127.0.0.0/8"]},
|
||||
)
|
||||
with (
|
||||
owned_proxy(gateway, tmp_path, {}, config=config) as candidate,
|
||||
mcp_peer() as peer,
|
||||
candidate.scenario() as scenario,
|
||||
):
|
||||
alias: Final = "a365" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias)
|
||||
granted: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
forwarded: Final = {"X-Forwarded-Proto": "https", "X-Forwarded-Host": "public.example"}
|
||||
|
||||
challenged: Final = _rpc(candidate, f"/mcp/{alias}", granted, forwarded)
|
||||
assert challenged.status_code == 401, challenged.text
|
||||
assert (
|
||||
f'resource_metadata="https://public.example/.well-known/oauth-protected-resource/mcp/{alias}"'
|
||||
in challenged.headers.get("www-authenticate", "")
|
||||
)
|
||||
|
||||
plain: Final = _rpc(candidate, f"/mcp/{alias}", granted, {})
|
||||
assert plain.status_code == 401, plain.text
|
||||
assert (
|
||||
f'resource_metadata="{_origin(candidate)}/.well-known/oauth-protected-resource/mcp/{alias}"'
|
||||
in plain.headers.get("www-authenticate", "")
|
||||
)
|
||||
|
||||
|
||||
def test_challenge_and_prm_resolve_the_connected_case_variant(gateway: Gateway, tmp_path: Path) -> None:
|
||||
config: Final = _sign_in_config(dict(AGENT_365_PARAMS), tmp_path / "agent365-case.yaml")
|
||||
with (
|
||||
owned_proxy(gateway, tmp_path, {}, config=config) as candidate,
|
||||
mcp_peer() as peer,
|
||||
candidate.scenario() as scenario,
|
||||
):
|
||||
alias: Final = "a365" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias)
|
||||
granted: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
connected_as: Final = alias.upper()
|
||||
|
||||
challenged: Final = _rpc(candidate, f"/mcp/{connected_as}", granted, {})
|
||||
assert challenged.status_code == 401, challenged.text
|
||||
authenticate: Final = challenged.headers.get("www-authenticate", "")
|
||||
assert (
|
||||
f'resource_metadata="{_origin(candidate)}/.well-known/oauth-protected-resource/mcp/{connected_as}"'
|
||||
in authenticate
|
||||
)
|
||||
assert 'error="invalid_token"' in authenticate
|
||||
|
||||
discovery: Final = candidate.client.get(f"/.well-known/oauth-protected-resource/mcp/{connected_as}")
|
||||
assert discovery.status_code == 200, discovery.text
|
||||
document: Final = discovery.json()
|
||||
assert document["authorization_servers"] == [ENTRA_ISSUER]
|
||||
assert document["scopes_supported"] == [GATEWAY_SCOPE]
|
||||
assert tool_calls(peer.drain()) == ()
|
||||
|
|
@ -1,6 +1,8 @@
|
|||
import base64
|
||||
import hashlib
|
||||
import json
|
||||
import secrets
|
||||
import socket
|
||||
import time
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
|
|
@ -27,6 +29,7 @@ from integration._support.mcp import (
|
|||
)
|
||||
from integration._support.mcp_grants import create_toolset
|
||||
from integration._support.oauth_server import AuthorizationServer, oauth_server
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
|
||||
ADD: Final = {"a": 2, "b": 3}
|
||||
CLIENT_REDIRECT: Final = "http://127.0.0.1:9/cb"
|
||||
|
|
@ -203,6 +206,105 @@ def test_token_exchange_without_a_subject_token_is_rejected_before_any_upstream_
|
|||
_assert_subject_token_challenge(as_subject, alias)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("status", (429, 408), ids=("throttled", "timed-out"))
|
||||
def test_a_throttled_token_exchange_is_an_outage_not_a_sign_in_challenge(gateway: Gateway, status: int) -> None:
|
||||
def shedding_idp(request: Request) -> Reply:
|
||||
assert request.method == "POST" and request.target == "/token", request
|
||||
return Reply(status=status, body=json.dumps({"error": "temporarily_unavailable"}).encode())
|
||||
|
||||
with mcp_peer() as peer, wire_server(shedding_idp) as idp, gateway.scenario() as scenario:
|
||||
alias: Final = "te" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(
|
||||
scenario,
|
||||
peer,
|
||||
alias,
|
||||
auth_type="oauth2_token_exchange",
|
||||
token_exchange_endpoint=idp.url + "/token",
|
||||
credentials={"client_id": "te-client", "client_secret": "te-secret"},
|
||||
)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
caller: Final = McpCaller(
|
||||
gateway, key, "server_mcp", alias, headers={"Authorization": "Bearer subject-" + uuid.uuid4().hex}
|
||||
)
|
||||
peer.drain()
|
||||
response: Final = caller.rpc("tools/call", {"name": f"{alias}-add", "arguments": ADD})
|
||||
assert response.status_code == 503, (response.status_code, response.text, dict(response.headers))
|
||||
assert "www-authenticate" not in response.headers, dict(response.headers)
|
||||
assert len(idp.drain()) == 1
|
||||
assert tool_calls(peer.drain()) == ()
|
||||
|
||||
|
||||
def test_a_malformed_caller_assertion_is_a_sign_in_challenge_not_an_outage(gateway: Gateway) -> None:
|
||||
def entra_like_idp(request: Request) -> Reply:
|
||||
assert request.method == "POST" and request.target == "/token", request
|
||||
body: Final = {"error": "invalid_client", "error_codes": [5002723], "error_description": "Invalid JWT token"}
|
||||
return Reply(status=401, body=json.dumps(body).encode())
|
||||
|
||||
with mcp_peer() as peer, wire_server(entra_like_idp) as idp, gateway.scenario() as scenario:
|
||||
alias: Final = "te" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(
|
||||
scenario,
|
||||
peer,
|
||||
alias,
|
||||
auth_type="oauth2_token_exchange",
|
||||
token_exchange_endpoint=idp.url + "/token",
|
||||
credentials={"client_id": "te-client", "client_secret": "te-secret"},
|
||||
)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
caller: Final = McpCaller(
|
||||
gateway,
|
||||
key,
|
||||
"server_mcp",
|
||||
alias,
|
||||
headers={"Authorization": "Bearer eyJhbGciOiJSUzI1NiJ9.eyJhdWQiOiJ3cm9uZyJ9.c2ln"},
|
||||
)
|
||||
peer.drain()
|
||||
response: Final = caller.rpc("tools/call", {"name": f"{alias}-add", "arguments": ADD})
|
||||
assert response.status_code == 401, (response.status_code, response.text, dict(response.headers))
|
||||
challenge: Final = response.headers["www-authenticate"]
|
||||
origin: Final = str(gateway.client.base_url).rstrip("/")
|
||||
assert f'resource_metadata="{origin}/.well-known/oauth-protected-resource/{alias}/mcp"' in challenge, challenge
|
||||
assert 'error="invalid_token"' in challenge, challenge
|
||||
assert len(idp.drain()) == 1
|
||||
assert tool_calls(peer.drain()) == ()
|
||||
|
||||
|
||||
def test_a_refused_token_exchange_answers_before_the_request_body_arrives(gateway: Gateway) -> None:
|
||||
"""The exchange preflight runs on the headers alone, so a client whose body is still in flight gets the
|
||||
challenge at once instead of the gateway waiting for bytes it will never use."""
|
||||
|
||||
def entra_like_idp(request: Request) -> Reply:
|
||||
body: Final = {"error": "invalid_client", "error_codes": [5002723], "error_description": "Invalid JWT token"}
|
||||
return Reply(status=401, body=json.dumps(body).encode())
|
||||
|
||||
with mcp_peer() as peer, wire_server(entra_like_idp) as idp, gateway.scenario() as scenario:
|
||||
alias: Final = "te" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(
|
||||
scenario,
|
||||
peer,
|
||||
alias,
|
||||
auth_type="oauth2_token_exchange",
|
||||
token_exchange_endpoint=idp.url + "/token",
|
||||
credentials={"client_id": "te-client", "client_secret": "te-secret"},
|
||||
)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
origin: Final = urlsplit(str(gateway.client.base_url))
|
||||
head: Final = (
|
||||
f"POST /mcp/{alias} HTTP/1.1\r\nHost: {origin.netloc}\r\nx-litellm-api-key: {key}\r\n"
|
||||
"Authorization: Bearer eyJhbGciOiJSUzI1NiJ9.eyJhdWQiOiJ3cm9uZyJ9.c2ln\r\n"
|
||||
"Accept: application/json, text/event-stream\r\nContent-Type: application/json\r\n"
|
||||
"Content-Length: 4096\r\n\r\n"
|
||||
)
|
||||
started: Final = time.monotonic()
|
||||
with socket.create_connection((origin.hostname or "127.0.0.1", origin.port or 80), timeout=10) as raw:
|
||||
raw.sendall(head.encode())
|
||||
status_line: Final = raw.recv(4096).split(b"\r\n", 1)[0]
|
||||
assert status_line == b"HTTP/1.1 401 Unauthorized", status_line
|
||||
assert time.monotonic() - started < 5
|
||||
assert len(idp.drain()) == 1
|
||||
assert tool_calls(peer.drain()) == ()
|
||||
|
||||
|
||||
def _assert_subject_token_challenge(response: httpx.Response, alias: str) -> None:
|
||||
assert response.status_code == 401, response.text
|
||||
challenge: Final = response.headers["www-authenticate"]
|
||||
|
|
|
|||
|
|
@ -615,6 +615,21 @@ def test_raise_token_exchange_challenge_is_rfc9728_invalid_token():
|
|||
assert "error_description=" in www
|
||||
|
||||
|
||||
def test_raise_token_exchange_challenge_advertises_the_connected_route_metadata_url():
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
|
||||
raise_token_exchange_challenge,
|
||||
)
|
||||
|
||||
connected: Final = "https://gw.example/.well-known/oauth-protected-resource/obo-srv/mcp"
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
raise_token_exchange_challenge(_server(alias="obo-srv"), root_path="/", resource_metadata=connected)
|
||||
assert exc_info.value.headers["WWW-Authenticate"] == (
|
||||
f'Bearer resource_metadata="{connected}", '
|
||||
'error="invalid_token", '
|
||||
'error_description="Missing or invalid subject token; authenticate with the IdP and retry"'
|
||||
)
|
||||
|
||||
|
||||
def test_raise_token_exchange_challenge_includes_server_root_path(monkeypatch):
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
|
||||
raise_token_exchange_challenge,
|
||||
|
|
|
|||
|
|
@ -7,7 +7,9 @@ the I/O edge that maps any transport/HTTP failure to None and parses a JSON body
|
|||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
from pydantic import SecretStr
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials import Error, Ok, ServerSpec
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_provider import (
|
||||
_post_exchange_endpoint,
|
||||
build_token_exchanger,
|
||||
|
|
@ -17,24 +19,26 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger
|
|||
SubjectTokenRejected,
|
||||
TokenExchangeClientError,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import TokenExchangeConfig
|
||||
|
||||
_HTTP_CLIENT = "litellm.llms.custom_httpx.http_handler.get_async_httpx_client"
|
||||
|
||||
|
||||
def _client_raising_4xx(body: object):
|
||||
"""An httpx client whose POST returns a 4xx whose ``raise_for_status`` raises an HTTPStatusError
|
||||
carrying ``body`` as its JSON, so the RFC 6749 error-code classification can be driven."""
|
||||
def _client_raising_status(status: int, body: object):
|
||||
"""An httpx client whose POST returns ``status`` whose ``raise_for_status`` raises an
|
||||
HTTPStatusError carrying ``body`` as its JSON, so the RFC 6749 error-code classification can be
|
||||
driven."""
|
||||
import httpx
|
||||
|
||||
request = httpx.Request("POST", "https://idp/token")
|
||||
response = httpx.Response(400, json=body, request=request)
|
||||
response = httpx.Response(status, json=body, request=request)
|
||||
|
||||
class _Resp:
|
||||
def raise_for_status(self) -> None:
|
||||
raise httpx.HTTPStatusError("bad request", request=request, response=response)
|
||||
|
||||
class _Client:
|
||||
async def post(self, url, headers, data):
|
||||
async def post(self, url, headers, data, timeout=None):
|
||||
return _Resp()
|
||||
|
||||
return _Client()
|
||||
|
|
@ -49,6 +53,25 @@ def test_build_gives_each_caller_an_independent_cache():
|
|||
assert build_token_exchanger() is not build_token_exchanger()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_token_exchanger_drives_the_injected_http_edge():
|
||||
seen: list[tuple[str, dict[str, str]]] = []
|
||||
|
||||
async def post(url: str, form: dict[str, str], client_auth_headers: dict[str, str]) -> dict[str, object] | None:
|
||||
seen.append((url, form))
|
||||
return {"access_token": "x", "expires_in": 60}
|
||||
|
||||
config = TokenExchangeConfig(
|
||||
token_exchange_endpoint="https://idp/token", client_id="cid", client_secret=SecretStr("csec")
|
||||
)
|
||||
server = ServerSpec(server_id="srv", resource="https://up.example.com", config=config)
|
||||
result = await build_token_exchanger(post=post).exchange("jwt", server, config)
|
||||
assert isinstance(result, Ok)
|
||||
assert result.ok.access_token == "x"
|
||||
assert [url for url, _ in seen] == ["https://idp/token"]
|
||||
assert seen[0][1]["subject_token"] == "jwt"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_returns_none_on_transport_error():
|
||||
with patch(_HTTP_CLIENT, side_effect=RuntimeError("boom")):
|
||||
|
|
@ -66,7 +89,7 @@ async def test_post_parses_json_body_on_success():
|
|||
return {"access_token": "x", "expires_in": 60}
|
||||
|
||||
class _Client:
|
||||
async def post(self, url, headers, data):
|
||||
async def post(self, url, headers, data, timeout=None):
|
||||
return _Resp()
|
||||
|
||||
with patch(_HTTP_CLIENT, return_value=_Client()):
|
||||
|
|
@ -80,7 +103,35 @@ async def test_post_parses_json_body_on_success():
|
|||
)
|
||||
async def test_post_maps_gateway_fault_4xx_to_client_error(code):
|
||||
# RFC 6749 5.2 gateway-fault codes must raise TokenExchangeClientError (-> 500), not the caller 401.
|
||||
with patch(_HTTP_CLIENT, return_value=_client_raising_4xx({"error": code})):
|
||||
with patch(_HTTP_CLIENT, return_value=_client_raising_status(400, {"error": code})):
|
||||
with pytest.raises(TokenExchangeClientError):
|
||||
await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("aadsts_code", [5002723, "5002710"], ids=["invalid_jwt", "no_kid_as_string"])
|
||||
async def test_post_maps_invalid_client_with_an_aadsts_50027xx_code_to_subject_rejected(aadsts_code):
|
||||
# Entra reports a malformed or unverifiable caller assertion as invalid_client with an AADSTS50027xx
|
||||
# sub-code (the same top-level error it uses for a bad gateway secret); that one is the caller's 401.
|
||||
body = {
|
||||
"error": "invalid_client",
|
||||
"error_description": f"AADSTS{aadsts_code}: Invalid JWT token.",
|
||||
"error_codes": [aadsts_code],
|
||||
}
|
||||
with patch(_HTTP_CLIENT, return_value=_client_raising_status(401, body)):
|
||||
with pytest.raises(SubjectTokenRejected):
|
||||
await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"error_codes",
|
||||
[[7000215], [5002723.0], "5002723", None],
|
||||
ids=["bad_secret_code", "float_code", "codes_not_a_list", "no_codes"],
|
||||
)
|
||||
async def test_post_keeps_invalid_client_without_an_assertion_code_as_client_error(error_codes):
|
||||
body = {"error": "invalid_client", **({} if error_codes is None else {"error_codes": error_codes})}
|
||||
with patch(_HTTP_CLIENT, return_value=_client_raising_status(401, body)):
|
||||
with pytest.raises(TokenExchangeClientError):
|
||||
await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {})
|
||||
|
||||
|
|
@ -93,7 +144,7 @@ async def test_post_maps_gateway_fault_4xx_to_client_error(code):
|
|||
)
|
||||
async def test_post_maps_subject_fault_4xx_to_subject_rejected(body):
|
||||
# A subject-fault code (or an unparseable/absent error) is the caller's problem -> SubjectTokenRejected (401).
|
||||
with patch(_HTTP_CLIENT, return_value=_client_raising_4xx(body)):
|
||||
with patch(_HTTP_CLIENT, return_value=_client_raising_status(400, body)):
|
||||
with pytest.raises(SubjectTokenRejected):
|
||||
await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {})
|
||||
|
||||
|
|
@ -110,7 +161,7 @@ async def test_post_returns_none_on_non_object_json(payload):
|
|||
return payload
|
||||
|
||||
class _Client:
|
||||
async def post(self, url, headers, data):
|
||||
async def post(self, url, headers, data, timeout=None):
|
||||
return _Resp()
|
||||
|
||||
with patch(_HTTP_CLIENT, return_value=_Client()):
|
||||
|
|
@ -128,7 +179,7 @@ async def test_post_threads_step_up_error_and_claims_into_subject_rejected():
|
|||
"error_description": "AADSTS50079: the user must enroll MFA",
|
||||
"claims": claims,
|
||||
}
|
||||
with patch(_HTTP_CLIENT, return_value=_client_raising_4xx(body)):
|
||||
with patch(_HTTP_CLIENT, return_value=_client_raising_status(400, body)):
|
||||
with pytest.raises(SubjectTokenRejected) as exc_info:
|
||||
await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {})
|
||||
assert exc_info.value.claims == claims
|
||||
|
|
@ -137,7 +188,7 @@ async def test_post_threads_step_up_error_and_claims_into_subject_rejected():
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_subject_rejection_without_claims_carries_none_claims():
|
||||
with patch(_HTTP_CLIENT, return_value=_client_raising_4xx({"error": "invalid_grant"})):
|
||||
with patch(_HTTP_CLIENT, return_value=_client_raising_status(400, {"error": "invalid_grant"})):
|
||||
with pytest.raises(SubjectTokenRejected) as exc_info:
|
||||
await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {})
|
||||
assert exc_info.value.claims is None
|
||||
|
|
@ -148,6 +199,26 @@ async def test_post_gateway_fault_still_wins_when_claims_are_present():
|
|||
# A gateway-fault code stays a 500-class TokenExchangeClientError even if the body carries
|
||||
# claims; the caller cannot fix invalid_client by stepping up.
|
||||
body = {"error": "invalid_client", "claims": '{"access_token":{}}'}
|
||||
with patch(_HTTP_CLIENT, return_value=_client_raising_4xx(body)):
|
||||
with patch(_HTTP_CLIENT, return_value=_client_raising_status(400, body)):
|
||||
with pytest.raises(TokenExchangeClientError):
|
||||
await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {})
|
||||
|
||||
|
||||
_CONFIG = TokenExchangeConfig(
|
||||
token_exchange_endpoint="https://idp.example.com/token",
|
||||
client_id="cid",
|
||||
client_secret=SecretStr("csec"),
|
||||
scopes=("s1",),
|
||||
)
|
||||
_SERVER = ServerSpec(server_id="srv", resource="https://up.example.com", config=_CONFIG)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("status", [408, 429])
|
||||
async def test_exchange_maps_throttled_or_timed_out_4xx_to_upstream_unavailable(status):
|
||||
# 408/429 are the IdP shedding load, not the caller presenting a bad subject: the exchange must
|
||||
# surface upstream_unavailable (503-class, retryable) and never tell the caller to sign in again.
|
||||
with patch(_HTTP_CLIENT, return_value=_client_raising_status(status, {"error": "temporarily_unavailable"})):
|
||||
result = await OboTokenExchanger(_post_exchange_endpoint).exchange("caller-jwt", _SERVER, _CONFIG)
|
||||
assert isinstance(result, Error)
|
||||
assert result.error.tag == "upstream_unavailable"
|
||||
|
|
|
|||
179
tests/unit/proxy/_experimental/mcp_server/test_caller_sign_in.py
Normal file
179
tests/unit/proxy/_experimental/mcp_server/test_caller_sign_in.py
Normal file
|
|
@ -0,0 +1,179 @@
|
|||
import asyncio
|
||||
from collections.abc import Iterator, Mapping
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.proxy._experimental.mcp_server.caller_sign_in import (
|
||||
CallerSignIn,
|
||||
CallerSignInPreflight,
|
||||
CallerSignInProvider,
|
||||
SignedIn,
|
||||
caller_sign_in_for,
|
||||
preflight_caller_sign_in,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.mcp import MCPAuth, MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
|
||||
class _SignInGuardrail(CustomGuardrail):
|
||||
def __init__(self, issuer: str, scope: str, gated: bool = True) -> None:
|
||||
super().__init__(guardrail_name=f"sign-in-{issuer}")
|
||||
self.issuer: Final = issuer
|
||||
self.scope: Final = scope
|
||||
self.gated: Final = gated
|
||||
self.seen: list[tuple[str, str | None]] = [] # mutable-ok: call recorder
|
||||
|
||||
def caller_sign_in(self, server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None) -> CallerSignIn | None:
|
||||
self.seen.append((server.name, user_api_key_auth.user_id if user_api_key_auth else None))
|
||||
if not self.gated:
|
||||
return None
|
||||
return CallerSignIn(issuers=(self.issuer,), scopes=(self.scope,))
|
||||
|
||||
async def preflight_caller_sign_in(
|
||||
self, server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None, subject_token: str
|
||||
) -> CallerSignInPreflight:
|
||||
return SignedIn()
|
||||
|
||||
|
||||
class _PlainGuardrail(CustomGuardrail):
|
||||
def __init__(self) -> None:
|
||||
super().__init__(guardrail_name="plain")
|
||||
|
||||
|
||||
def _server(auth_type: MCPAuth | None = None, scopes: list[str] | None = None) -> MCPServer:
|
||||
return MCPServer(
|
||||
server_id="tools-id",
|
||||
name="tools",
|
||||
server_name="tools",
|
||||
transport=MCPTransport.http,
|
||||
url="https://tools.test/mcp",
|
||||
auth_type=auth_type,
|
||||
scopes=scopes,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def registered() -> Iterator[tuple[_SignInGuardrail, _SignInGuardrail]]:
|
||||
first: Final = _SignInGuardrail("https://idp-a.test", "scope-a")
|
||||
second: Final = _SignInGuardrail("https://idp-b.test", "scope-b")
|
||||
plain: Final = _PlainGuardrail()
|
||||
for callback in (first, plain, second):
|
||||
litellm.logging_callback_manager.add_litellm_callback(callback)
|
||||
try:
|
||||
yield first, second
|
||||
finally:
|
||||
for callback in (first, plain, second):
|
||||
litellm.logging_callback_manager.remove_callback_from_list_by_object(
|
||||
litellm.callbacks, callback, require_self=False
|
||||
)
|
||||
|
||||
|
||||
def test_protocol_matches_only_guardrails_implementing_the_hook():
|
||||
assert isinstance(_SignInGuardrail("i", "s"), CallerSignInProvider)
|
||||
assert not isinstance(_PlainGuardrail(), CallerSignInProvider)
|
||||
|
||||
|
||||
def test_no_registered_provider_and_non_obo_advertises_nothing():
|
||||
assert caller_sign_in_for(_server(), None) is None
|
||||
|
||||
|
||||
def test_registered_providers_merge_in_order_and_dedupe(registered):
|
||||
first, second = registered
|
||||
sign_in: Final = caller_sign_in_for(_server(), None)
|
||||
|
||||
assert sign_in is not None
|
||||
assert sign_in.issuers == ("https://idp-a.test", "https://idp-b.test")
|
||||
assert sign_in.scopes == ("scope-a", "scope-b")
|
||||
assert first.seen == [("tools", None)]
|
||||
assert second.seen == [("tools", None)]
|
||||
|
||||
|
||||
def test_provider_returning_none_contributes_nothing(registered):
|
||||
ungated = _SignInGuardrail("https://idp-c.test", "scope-c", gated=False)
|
||||
litellm.logging_callback_manager.add_litellm_callback(ungated)
|
||||
try:
|
||||
sign_in: Final = caller_sign_in_for(_server(), None)
|
||||
assert sign_in is not None
|
||||
assert "https://idp-c.test" not in sign_in.issuers
|
||||
finally:
|
||||
litellm.logging_callback_manager.remove_callback_from_list_by_object(
|
||||
litellm.callbacks, ungated, require_self=False
|
||||
)
|
||||
|
||||
|
||||
def test_obo_server_contributes_jwt_issuers_and_own_scopes(monkeypatch):
|
||||
monkeypatch.setenv("JWT_ISSUER", "https://jwt-idp.test")
|
||||
sign_in: Final = caller_sign_in_for(_server(auth_type=MCPAuth.oauth2_token_exchange, scopes=["read"]), None)
|
||||
assert sign_in == CallerSignIn(issuers=("https://jwt-idp.test",), scopes=("read",))
|
||||
|
||||
|
||||
def test_obo_server_and_provider_merge_and_dedupe(monkeypatch, registered):
|
||||
monkeypatch.setenv("JWT_ISSUER", "https://idp-a.test")
|
||||
sign_in: Final = caller_sign_in_for(
|
||||
_server(auth_type=MCPAuth.oauth2_token_exchange, scopes=["read", "scope-a"]), None
|
||||
)
|
||||
assert sign_in is not None
|
||||
assert sign_in.issuers == ("https://idp-a.test", "https://idp-b.test")
|
||||
assert sign_in.scopes == ("read", "scope-a", "scope-b")
|
||||
|
||||
|
||||
def test_obo_server_without_jwt_issuer_still_signs_in_when_a_provider_gates(registered):
|
||||
sign_in: Final = caller_sign_in_for(_server(auth_type=MCPAuth.oauth2_token_exchange), None)
|
||||
assert sign_in is not None
|
||||
assert sign_in.issuers == ("https://idp-a.test", "https://idp-b.test")
|
||||
|
||||
|
||||
def test_oauth_utils_strips_the_route_relative_root_path():
|
||||
"""Regression: Starlette sets ``app_root_path`` to ``""`` on an unmounted app, so the strip must
|
||||
fall back to ``root_path`` (which is where ``/mcp`` lands when the MCP app is mounted)."""
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import get_route_relative_request_path
|
||||
|
||||
scope: Final[Mapping[str, object]] = {
|
||||
"type": "http",
|
||||
"path": "/mcp/catalog",
|
||||
"root_path": "/mcp",
|
||||
"app_root_path": "",
|
||||
}
|
||||
assert get_route_relative_request_path(scope) == "/catalog" # pyright: ignore[reportArgumentType]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_preflight_with_no_gating_provider_never_reads_the_request_body(monkeypatch):
|
||||
"""A plain OBO connect has nothing to pre-flight, so the exchange answers before the body arrives, as it
|
||||
did before the sign-in seam; reading the body first would stall a client that sends its headers early."""
|
||||
monkeypatch.setenv("JWT_ISSUER", "https://jwt-idp.test")
|
||||
body_read: Final = asyncio.Event()
|
||||
|
||||
async def connecting() -> bool:
|
||||
body_read.set()
|
||||
return True
|
||||
|
||||
await preflight_caller_sign_in(
|
||||
_server(auth_type=MCPAuth.oauth2_token_exchange),
|
||||
None,
|
||||
"sub.ject.jws",
|
||||
root_path="",
|
||||
resource_metadata=None,
|
||||
connecting=connecting,
|
||||
)
|
||||
|
||||
assert not body_read.is_set()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_preflight_with_a_gating_provider_reads_the_body_to_tell_a_connect_apart(registered):
|
||||
body_read: Final = asyncio.Event()
|
||||
|
||||
async def connecting() -> bool:
|
||||
body_read.set()
|
||||
return True
|
||||
|
||||
await preflight_caller_sign_in(
|
||||
_server(), None, "sub.ject.jws", root_path="", resource_metadata=None, connecting=connecting
|
||||
)
|
||||
|
||||
assert body_read.is_set()
|
||||
|
|
@ -7,6 +7,7 @@ from base64 import urlsafe_b64encode
|
|||
from datetime import datetime, timedelta, timezone
|
||||
from typing import TYPE_CHECKING, Final
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from urllib.parse import parse_qs, urlparse
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -392,6 +393,84 @@ def trust_xff():
|
|||
yield
|
||||
|
||||
|
||||
def _registered_gateway_oauth2_server():
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
server: Final = MCPServer(
|
||||
server_id="oid-7f3a",
|
||||
name="gwx",
|
||||
server_name="gwx",
|
||||
alias="gwx",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
oauth2_flow="authorization_code",
|
||||
client_id="gw-client",
|
||||
client_secret="gw-secret",
|
||||
authorization_url="https://provider.com/oauth/authorize",
|
||||
token_url="https://provider.com/oauth/token",
|
||||
scopes=["read"],
|
||||
)
|
||||
global_mcp_server_manager.registry.clear()
|
||||
global_mcp_server_manager.registry[server.server_id] = server
|
||||
return server
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("lookup", ["GWX", "oid-7f3a"], ids=["alias_case", "server_id"])
|
||||
async def test_authorization_server_doc_for_a_moved_lookup_matches_the_exact_name_doc(lookup):
|
||||
"""A case variant or server id now resolves like the exact name, so its AS metadata is the
|
||||
exact-name doc with the requested spelling in the issuer and endpoint paths."""
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
oauth_authorization_server_mcp_standard,
|
||||
)
|
||||
|
||||
_registered_gateway_oauth2_server()
|
||||
request: Final = _mock_callback_request("http://litellm.example.com/")
|
||||
|
||||
exact: Final = await oauth_authorization_server_mcp_standard(request=request, mcp_server_name="gwx")
|
||||
moved: Final = await oauth_authorization_server_mcp_standard(request=request, mcp_server_name=lookup)
|
||||
|
||||
assert exact["issuer"] == "http://litellm.example.com/mcp/gwx"
|
||||
assert moved == {
|
||||
key: value.replace("/gwx", f"/{lookup}") if isinstance(value, str) else value for key, value in exact.items()
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("lookup", ["GWX", "oid-7f3a"], ids=["alias_case", "server_id"])
|
||||
async def test_authorize_relay_for_a_moved_lookup_redirects_like_the_exact_name(lookup):
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import authorize
|
||||
|
||||
_registered_gateway_oauth2_server()
|
||||
request: Final = _mock_callback_request("http://litellm.example.com/")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper", return_value="sealed"
|
||||
):
|
||||
exact: Final = await authorize(
|
||||
request=request, mcp_server_name="gwx", redirect_uri="http://127.0.0.1:60108/callback", state="s1"
|
||||
)
|
||||
moved: Final = await authorize(
|
||||
request=request, mcp_server_name=lookup, redirect_uri="http://127.0.0.1:60108/callback", state="s1"
|
||||
)
|
||||
|
||||
exact_target: Final = urlparse(exact.headers["location"])
|
||||
moved_target: Final = urlparse(moved.headers["location"])
|
||||
exact_query: Final = parse_qs(exact_target.query)
|
||||
moved_query: Final = parse_qs(moved_target.query)
|
||||
assert exact.status_code == 307
|
||||
assert exact_target._replace(query="") == urlparse("https://provider.com/oauth/authorize")
|
||||
assert exact_query["client_id"] == ["gw-client"]
|
||||
assert len(exact_query.pop("state")) == 1 and len(moved_query.pop("state")) == 1
|
||||
assert (moved.status_code, moved_target._replace(query=""), moved_query) == (
|
||||
exact.status_code,
|
||||
exact_target._replace(query=""),
|
||||
exact_query,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authorize_endpoint_includes_response_type():
|
||||
"""Test that authorize endpoint includes response_type=code parameter (fixes #15684)"""
|
||||
|
|
@ -3642,8 +3721,8 @@ async def test_authorize_resolves_server_by_id_when_name_lookup_fails():
|
|||
|
||||
assert response.status_code == 307
|
||||
assert "https://provider.com/oauth/authorize" in response.headers["location"]
|
||||
by_name.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_id.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_name.assert_called_once_with(server.server_id)
|
||||
by_id.assert_called_once_with(server.server_id)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -3679,8 +3758,8 @@ async def test_token_endpoint_resolves_server_by_id_when_name_lookup_fails():
|
|||
)
|
||||
|
||||
assert json.loads(result.body)["access_token"] == "token"
|
||||
by_name.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_id.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_name.assert_called_once_with(server.server_id)
|
||||
by_id.assert_called_once_with(server.server_id)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -3711,8 +3790,8 @@ async def test_register_client_resolves_server_by_id_when_name_lookup_fails():
|
|||
result = await discoverable_endpoints.register_client(request=request, mcp_server_name=server.server_id)
|
||||
|
||||
assert json.loads(result.body)["client_id"] == "registered-client"
|
||||
by_name.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_id.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_name.assert_called_once_with(server.server_id)
|
||||
by_id.assert_called_once_with(server.server_id)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -3739,8 +3818,58 @@ async def test_protected_resource_metadata_resolves_server_by_id_when_name_looku
|
|||
|
||||
assert result["authorization_servers"] == ["https://llm.example.com/mcp"]
|
||||
assert result["resource"] == f"https://llm.example.com/mcp/{server.server_id}"
|
||||
by_name.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_id.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_name.assert_called_once_with(server.server_id)
|
||||
by_id.assert_called_once_with(server.server_id)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_protected_resource_metadata_resolves_the_connected_case_variant():
|
||||
"""The challenge points clients at the segment they connected with (``/mcp/CATALOG``), so the PRM
|
||||
route must resolve that same segment through the answering-to fallback."""
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
from litellm.types.mcp import MCPAuth, MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
server = MCPServer(
|
||||
server_id="catalog-server-id-001",
|
||||
name="catalog",
|
||||
alias="catalog",
|
||||
server_name="catalog",
|
||||
url="https://catalog.test/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
mcp_info={"server_name": "catalog"},
|
||||
)
|
||||
sign_in: Final = CallerSignIn(
|
||||
issuers=("https://login.microsoftonline.com/00000000-0000-0000-0000-000000000000/v2.0",),
|
||||
scopes=("api://22222222-2222-2222-2222-222222222222/access_as_user",),
|
||||
)
|
||||
request = MagicMock(spec=Request)
|
||||
request.base_url = "https://llm.example.com/"
|
||||
request.headers = {}
|
||||
|
||||
with (
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_name", return_value=None), # test-quality-ok: resolver seam
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=None), # test-quality-ok: resolver seam
|
||||
patch.object(
|
||||
global_mcp_server_manager, "get_filtered_registry", return_value={server.server_id: server}
|
||||
), # test-quality-ok: resolver seam
|
||||
patch.object(discoverable_endpoints, "caller_sign_in_for", return_value=sign_in), # test-quality-ok: provider seam
|
||||
):
|
||||
result = await discoverable_endpoints._build_oauth_protected_resource_response(
|
||||
request=request,
|
||||
mcp_server_name="CATALOG",
|
||||
use_standard_pattern=True,
|
||||
)
|
||||
|
||||
assert result["authorization_servers"] == (
|
||||
"https://login.microsoftonline.com/00000000-0000-0000-0000-000000000000/v2.0",
|
||||
)
|
||||
assert result["resource"] == "https://llm.example.com/mcp/CATALOG"
|
||||
|
||||
|
||||
def test_authorization_server_metadata_resolves_server_by_id_when_name_lookup_fails():
|
||||
|
|
@ -3765,8 +3894,8 @@ def test_authorization_server_metadata_resolves_server_by_id_when_name_lookup_fa
|
|||
|
||||
assert result["scopes_supported"] == server.scopes
|
||||
assert result["issuer"] == f"https://llm.example.com/{server.server_id}"
|
||||
by_name.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_id.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_name.assert_called_once_with(server.server_id)
|
||||
by_id.assert_called_once_with(server.server_id)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -7288,7 +7417,7 @@ async def test_token_exchange_persists_for_oauth2():
|
|||
# -------------------------------------------------------------------
|
||||
|
||||
_OBO_RESOURCE = "https://litellm.example.com/mcp/obo_mcp"
|
||||
_PATCH_ISSUERS = "litellm.proxy._experimental.mcp_server.discoverable_endpoints._jwt_auth_issuers"
|
||||
_PATCH_ISSUERS = "litellm.proxy._experimental.mcp_server.caller_sign_in.jwt_auth_issuers"
|
||||
|
||||
|
||||
def _obo_server(scopes=None):
|
||||
|
|
@ -7307,48 +7436,48 @@ def _obo_server(scopes=None):
|
|||
)
|
||||
|
||||
|
||||
def test_obo_protected_resource_response_names_jwt_issuers():
|
||||
def test_caller_sign_in_protected_resource_response_names_jwt_issuers():
|
||||
"""An OBO server's PRM points authorization_servers at the configured JWT issuers (the IdP that
|
||||
mints and validates the subject token), with the gateway resource echoed back."""
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
_obo_protected_resource_response,
|
||||
_caller_sign_in_protected_resource_response,
|
||||
)
|
||||
|
||||
with patch(_PATCH_ISSUERS, return_value=["https://idp.example.com"]):
|
||||
response = _obo_protected_resource_response(_obo_server(scopes=["read"]), _OBO_RESOURCE)
|
||||
response = _caller_sign_in_protected_resource_response(_obo_server(scopes=["read"]), _OBO_RESOURCE)
|
||||
assert response == {
|
||||
"authorization_servers": ["https://idp.example.com"],
|
||||
"authorization_servers": ("https://idp.example.com",),
|
||||
"resource": _OBO_RESOURCE,
|
||||
"scopes_supported": ["read"],
|
||||
"scopes_supported": ("read",),
|
||||
}
|
||||
|
||||
|
||||
def test_obo_protected_resource_response_scopes_default_empty():
|
||||
"""A scopeless OBO server reports scopes_supported as [] rather than None."""
|
||||
def test_caller_sign_in_protected_resource_response_scopes_default_empty():
|
||||
"""A scopeless OBO server reports scopes_supported as an empty array rather than None."""
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
_obo_protected_resource_response,
|
||||
_caller_sign_in_protected_resource_response,
|
||||
)
|
||||
|
||||
with patch(_PATCH_ISSUERS, return_value=["https://idp.example.com"]):
|
||||
response = _obo_protected_resource_response(_obo_server(scopes=None), _OBO_RESOURCE)
|
||||
assert response["scopes_supported"] == []
|
||||
response = _caller_sign_in_protected_resource_response(_obo_server(scopes=None), _OBO_RESOURCE)
|
||||
assert response["scopes_supported"] == ()
|
||||
|
||||
|
||||
def test_obo_protected_resource_response_falls_back_when_no_issuer():
|
||||
def test_caller_sign_in_protected_resource_response_falls_back_when_no_issuer():
|
||||
"""With no JWT issuer configured, the OBO branch returns None so the caller falls back to the
|
||||
gateway-default PRM (discovery still works, it just can't name the IdP)."""
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
_obo_protected_resource_response,
|
||||
_caller_sign_in_protected_resource_response,
|
||||
)
|
||||
|
||||
with patch(_PATCH_ISSUERS, return_value=[]):
|
||||
assert _obo_protected_resource_response(_obo_server(), _OBO_RESOURCE) is None
|
||||
assert _caller_sign_in_protected_resource_response(_obo_server(), _OBO_RESOURCE) is None
|
||||
|
||||
|
||||
def test_obo_protected_resource_response_ignores_non_obo_server():
|
||||
def test_caller_sign_in_protected_resource_response_ignores_non_obo_server():
|
||||
"""Non-OBO servers are not handled by this branch (returns None -> gateway default)."""
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
_obo_protected_resource_response,
|
||||
_caller_sign_in_protected_resource_response,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp import MCPAuth
|
||||
|
|
@ -7360,7 +7489,7 @@ def test_obo_protected_resource_response_ignores_non_obo_server():
|
|||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
)
|
||||
assert _obo_protected_resource_response(oauth2_server, _OBO_RESOURCE) is None
|
||||
assert _caller_sign_in_protected_resource_response(oauth2_server, _OBO_RESOURCE) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -7390,7 +7519,7 @@ async def test_build_oauth_protected_resource_response_obo_end_to_end():
|
|||
mcp_server_name="obo_mcp",
|
||||
use_standard_pattern=True,
|
||||
)
|
||||
assert response["authorization_servers"] == ["https://idp.example.com"]
|
||||
assert response["authorization_servers"] == ("https://idp.example.com",)
|
||||
assert response["resource"] == "https://litellm.example.com/mcp/obo_mcp"
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
|
|
|||
|
|
@ -924,6 +924,7 @@ async def test_get_tools_from_mcp_servers():
|
|||
mock_manager.get_mcp_server_by_id = lambda server_id: (
|
||||
mock_server_1 if server_id == "server1_id" else mock_server_2
|
||||
)
|
||||
mock_manager.get_mcp_server_answering_to = MCPServerManager().get_mcp_server_answering_to
|
||||
mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1])
|
||||
# Mock filter_server_ids_by_ip_with_info to return input unchanged (no IP filtering in test)
|
||||
mock_manager.filter_server_ids_by_ip_with_info = MagicMock(
|
||||
|
|
@ -1004,6 +1005,7 @@ async def test_get_tools_from_mcp_servers():
|
|||
if server_id == "server1_id"
|
||||
else (mock_server_2 if server_id == "server2_id" else mock_server_3)
|
||||
)
|
||||
mock_manager.get_mcp_server_answering_to = MagicMock(return_value=None)
|
||||
mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1])
|
||||
# Mock filter_server_ids_by_ip_with_info to return input unchanged (no IP filtering in test)
|
||||
mock_manager.filter_server_ids_by_ip_with_info = MagicMock(
|
||||
|
|
|
|||
|
|
@ -1637,6 +1637,19 @@ class TestMCPServerManager:
|
|||
assert server.uses_per_server_oauth_relay is True
|
||||
assert server.advertises_gateway_authorization_server is False
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_servers_from_config_does_not_advertise_gateway_as_for_token_exchange(self):
|
||||
# keeps_caller_authorization includes oauth2_token_exchange so a sign-in provider may gate it,
|
||||
# but named discovery must still fall to the server's own PRM rather than the aggregate AS.
|
||||
manager = MCPServerManager()
|
||||
|
||||
with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)):
|
||||
await manager.load_servers_from_config(self._client_forwarded_config(MCPAuth.oauth2_token_exchange))
|
||||
|
||||
server = next(iter(manager.config_mcp_servers.values()))
|
||||
assert server.keeps_caller_authorization is True
|
||||
assert server.advertises_gateway_authorization_server is False
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"config",
|
||||
|
|
@ -3090,6 +3103,34 @@ class TestMCPServerManager:
|
|||
www_authenticate = headers.get("WWW-Authenticate") or headers.get("www-authenticate") or ""
|
||||
assert "resource_metadata" in www_authenticate
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_preflight_rejected_subject_challenge_names_the_connected_segment(self):
|
||||
"""A subject rejected on ``/mcp/<server_id>`` must point resource_metadata at that same
|
||||
segment, the way the sign-in preflight does, so the client's discovery fetch resolves."""
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import CredError
|
||||
|
||||
class _FakeProvider:
|
||||
async def resolve_credentials(self, subject, server):
|
||||
return Error(CredError.of_unauthorized("subject token rejected by the IdP"))
|
||||
|
||||
manager = MCPServerManager(cred_provider=_FakeProvider())
|
||||
server = self._token_exchange_server("te-preflight-segment")
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await manager.preflight_token_exchange(
|
||||
server=server,
|
||||
oauth2_headers={"Authorization": "Bearer rejected-subject"},
|
||||
user_api_key_auth=None,
|
||||
resource_metadata=f"http://gw.test/.well-known/oauth-protected-resource/mcp/{server.server_id}",
|
||||
)
|
||||
headers = exc_info.value.headers or {}
|
||||
www_authenticate = headers.get("WWW-Authenticate") or headers.get("www-authenticate") or ""
|
||||
assert (
|
||||
f'resource_metadata="http://gw.test/.well-known/oauth-protected-resource/mcp/{server.server_id}"'
|
||||
in www_authenticate
|
||||
), www_authenticate
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_preflight_token_exchange_maps_gateway_fault_to_public_status(self):
|
||||
"""A gateway-fault CredError (e.g. invalid_client) must surface its public status (500)
|
||||
|
|
@ -6346,6 +6387,94 @@ class TestMCPServerManager:
|
|||
assert manager._resolve_mcp_server_for_tool_call("zapier", "create_zap") is zapier
|
||||
assert manager._resolve_mcp_server_for_tool_call("other", "create_zap") is other
|
||||
|
||||
@pytest.mark.parametrize("gh_first", [True, False], ids=["gh-listed-first", "gh-public-listed-first"])
|
||||
def test_answering_to_prefers_alias_over_earlier_prefix_match(self, gh_first):
|
||||
manager = MCPServerManager()
|
||||
gh = MCPServer(
|
||||
server_id="gh-id", name="gh", server_name="gh", transport=MCPTransport.http, auth_type=MCPAuth.oauth2
|
||||
)
|
||||
gh_public = MCPServer(
|
||||
server_id="gh-public-id", name="gh_public", server_name="gh_public", alias="gh", transport=MCPTransport.http
|
||||
)
|
||||
manager.registry = (
|
||||
{"gh-id": gh, "gh-public-id": gh_public} if gh_first else {"gh-public-id": gh_public, "gh-id": gh}
|
||||
)
|
||||
|
||||
assert manager.get_mcp_server_answering_to("gh") is gh_public
|
||||
assert manager.get_mcp_server_answering_to("GH") is gh_public
|
||||
assert manager.get_mcp_server_answering_to("Gh_Public") is gh_public
|
||||
assert manager.get_mcp_server_answering_to("gh-public-id") is gh_public
|
||||
assert manager.get_mcp_server_answering_to("GH_PUBLIC") is gh_public
|
||||
assert manager.get_mcp_server_answering_to("gh-id") is gh
|
||||
|
||||
@pytest.mark.parametrize("gh_first", [True, False], ids=["alias-listed-first", "server-name-listed-first"])
|
||||
def test_answering_to_agrees_with_exact_name_before_case_folding(self, gh_first):
|
||||
manager = MCPServerManager()
|
||||
by_alias = MCPServer(
|
||||
server_id="a-id",
|
||||
name="a",
|
||||
server_name="a",
|
||||
alias="gh",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
)
|
||||
by_server_name = MCPServer(server_id="b-id", name="b", server_name="Gh", transport=MCPTransport.http)
|
||||
manager.registry = (
|
||||
{"a-id": by_alias, "b-id": by_server_name} if gh_first else {"b-id": by_server_name, "a-id": by_alias}
|
||||
)
|
||||
|
||||
for name in ("Gh", "gh"):
|
||||
assert manager.get_mcp_server_answering_to(name) is manager.get_mcp_server_by_name(name), name
|
||||
assert manager.get_mcp_server_answering_to("Gh") is by_server_name
|
||||
assert manager.get_mcp_server_answering_to("gh") is by_alias
|
||||
assert manager.get_mcp_server_answering_to("GH") is by_alias
|
||||
|
||||
@pytest.mark.parametrize("hidden_first", [True, False], ids=["hidden-listed-first", "public-listed-first"])
|
||||
def test_answering_to_never_reroutes_a_name_hidden_from_an_ip_to_a_case_variant(self, hidden_first):
|
||||
manager = MCPServerManager()
|
||||
hidden = MCPServer(
|
||||
server_id="p-id",
|
||||
name="gh",
|
||||
server_name="gh",
|
||||
transport=MCPTransport.http,
|
||||
available_on_public_internet=False,
|
||||
)
|
||||
public = MCPServer(server_id="u-id", name="u", server_name="u", alias="Gh", transport=MCPTransport.http)
|
||||
manager.registry = {"p-id": hidden, "u-id": public} if hidden_first else {"u-id": public, "p-id": hidden}
|
||||
|
||||
assert manager.get_mcp_server_answering_to("gh", client_ip="203.0.113.7") is None
|
||||
assert manager.get_mcp_server_answering_to("p-id", client_ip="203.0.113.7") is None
|
||||
assert manager.get_mcp_server_answering_to("gh", client_ip="10.0.0.7") is hidden
|
||||
assert manager.get_mcp_server_answering_to("Gh", client_ip="203.0.113.7") is public
|
||||
|
||||
@pytest.mark.parametrize("pinned_first", [True, False], ids=["pinned-id-listed-first", "alias-listed-first"])
|
||||
def test_answering_to_and_discovery_agree_on_a_pinned_id_that_another_alias_case_folds_to(self, pinned_first):
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
pinned = MCPServer(server_id="foo", name="pinned", server_name="pinned", transport=MCPTransport.http)
|
||||
by_alias = MCPServer(
|
||||
server_id="b-id",
|
||||
name="b",
|
||||
server_name="b",
|
||||
alias="Foo",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
)
|
||||
registry = {"foo": pinned, "b-id": by_alias} if pinned_first else {"b-id": by_alias, "foo": pinned}
|
||||
global_mcp_server_manager.registry.clear()
|
||||
global_mcp_server_manager.registry.update(registry)
|
||||
try:
|
||||
for name in ("foo", "Foo", "FOO"):
|
||||
connected = global_mcp_server_manager.get_mcp_server_answering_to(name)
|
||||
discovered = discoverable_endpoints._resolve_mcp_server_by_name_or_id(name, client_ip=None)
|
||||
assert discovered is connected, name
|
||||
assert global_mcp_server_manager.get_mcp_server_answering_to("foo") is pinned
|
||||
assert global_mcp_server_manager.get_mcp_server_answering_to("Foo") is by_alias
|
||||
assert global_mcp_server_manager.get_mcp_server_answering_to("FOO") is by_alias
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
def test_remove_server_drops_only_its_own_tool_mapping_rows(self):
|
||||
manager = self._manager_with_deepwiki_and_huggingface()
|
||||
|
||||
|
|
@ -6973,6 +7102,185 @@ class TestMCPServerManager:
|
|||
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("raw_headers", "api_key", "expected_bearer", "expected_subject"),
|
||||
[
|
||||
pytest.param(
|
||||
{"authorization": "Bearer sk-1234"},
|
||||
"sk-1234",
|
||||
"sk-1234",
|
||||
None,
|
||||
id="litellm-key-as-bearer-is-not-a-subject",
|
||||
),
|
||||
pytest.param(
|
||||
{"x-litellm-api-key": "sk-1234", "authorization": "Bearer eyJ.x.y"},
|
||||
"sk-1234",
|
||||
"eyJ.x.y",
|
||||
"eyJ.x.y",
|
||||
id="key-admission-plus-idp-bearer-subject",
|
||||
),
|
||||
pytest.param(
|
||||
{"x-litellm-api-key": "sk-1234", "authorization": "bearer eyJ.x.y"},
|
||||
"sk-1234",
|
||||
"eyJ.x.y",
|
||||
"eyJ.x.y",
|
||||
id="lowercase-bearer-scheme-is-the-subject",
|
||||
),
|
||||
pytest.param(
|
||||
{"x-litellm-api-key": "sk-1234", "authorization": "Basic a.b.c"},
|
||||
"sk-1234",
|
||||
None,
|
||||
None,
|
||||
id="non-bearer-scheme-is-not-a-subject",
|
||||
),
|
||||
pytest.param(
|
||||
{"x-litellm-api-key": "sk-1234", "authorization": "Digest x.y.z"},
|
||||
"sk-1234",
|
||||
None,
|
||||
None,
|
||||
id="digest-scheme-is-not-a-subject",
|
||||
),
|
||||
pytest.param(
|
||||
{"x-litellm-api-key": "sk-1234", "authorization": "eyJ.x.y"},
|
||||
"sk-1234",
|
||||
None,
|
||||
None,
|
||||
id="scheme-less-value-is-not-a-subject",
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_pre_call_tool_check_separates_raw_bearer_from_subject(
|
||||
self, raw_headers, api_key, expected_bearer, expected_subject
|
||||
):
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv", allowed_tools=None
|
||||
)
|
||||
proxy_logging = MagicMock()
|
||||
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
|
||||
proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={})
|
||||
proxy_logging.pre_call_hook = AsyncMock(return_value=None)
|
||||
|
||||
await manager.pre_call_tool_check(
|
||||
server_name="srv",
|
||||
name="turn",
|
||||
arguments={},
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key=api_key, user_id="u"),
|
||||
proxy_logging_obj=proxy_logging,
|
||||
server=server,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
|
||||
kwargs: Final = proxy_logging._create_mcp_request_object_from_kwargs.call_args.args[0]
|
||||
assert kwargs["incoming_bearer_token"] == expected_bearer
|
||||
assert kwargs["incoming_subject_token"] == expected_subject
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("raw_headers", "expected_subject"),
|
||||
[
|
||||
pytest.param({"authorization": "Bearer gw.master.key"}, None, id="master-key-alone-is-not-a-subject"),
|
||||
pytest.param(
|
||||
{"x-litellm-api-key": "sk-1234", "authorization": "Bearer gw.master.key"},
|
||||
None,
|
||||
id="master-key-next-to-a-virtual-key-is-not-a-subject",
|
||||
),
|
||||
pytest.param(
|
||||
{"x-litellm-api-key": "sk-1234", "authorization": "Bearer gw.other.jws"},
|
||||
"gw.other.jws",
|
||||
id="a-dotted-bearer-that-is-not-the-master-key-is-the-subject",
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_pre_call_tool_check_withholds_a_dotted_master_key_from_the_subject(
|
||||
self, raw_headers, expected_subject
|
||||
):
|
||||
"""A master key is a LiteLLM credential whatever its shape, so even one with the two dots of a
|
||||
compact JWS never becomes the sign-in subject a provider would send to its IdP."""
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv", allowed_tools=None
|
||||
)
|
||||
proxy_logging = MagicMock()
|
||||
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
|
||||
proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={})
|
||||
proxy_logging.pre_call_hook = AsyncMock(return_value=None)
|
||||
|
||||
with patch("litellm.proxy.proxy_server.master_key", "gw.master.key"):
|
||||
await manager.pre_call_tool_check(
|
||||
server_name="srv",
|
||||
name="turn",
|
||||
arguments={},
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-1234", user_id="u"),
|
||||
proxy_logging_obj=proxy_logging,
|
||||
server=server,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
|
||||
kwargs: Final = proxy_logging._create_mcp_request_object_from_kwargs.call_args.args[0]
|
||||
assert kwargs["incoming_bearer_token"] == raw_headers["authorization"].removeprefix("Bearer ")
|
||||
assert kwargs["incoming_subject_token"] == expected_subject
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("raw_headers", "api_key", "custom_auth", "expected_subject"),
|
||||
[
|
||||
pytest.param(
|
||||
{"authorization": "Bearer eyJ.x.y"},
|
||||
"eyJ.x.y",
|
||||
True,
|
||||
"eyJ.x.y",
|
||||
id="custom-auth-idp-bearer-is-the-subject",
|
||||
),
|
||||
pytest.param(
|
||||
{"authorization": "Bearer eyJ.x.y"},
|
||||
"eyJ.x.y",
|
||||
False,
|
||||
"eyJ.x.y",
|
||||
id="built-in-oauth2-admission-bearer-is-the-subject",
|
||||
),
|
||||
pytest.param({"authorization": "Bearer sk-1234"}, "sk-1234", True, None, id="virtual-key-is-not-a-subject"),
|
||||
pytest.param(
|
||||
{"x-litellm-api-key": "ca-key", "authorization": "Bearer ca-key"},
|
||||
"ca-key",
|
||||
True,
|
||||
None,
|
||||
id="explicit-key-admission-repeated-in-authorization-is-not-a-subject",
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_pre_call_tool_check_hands_sign_in_the_bearer_that_admitted_the_caller(
|
||||
self, raw_headers, api_key, custom_auth, expected_subject
|
||||
):
|
||||
"""Custom auth and the built-in OAuth2 admission both admit the caller on its own IdP token in
|
||||
``Authorization`` with no ``x-litellm-api-key`` and record it as ``api_key``; that token is the
|
||||
sign-in subject as the raw bearer was before the subject split."""
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv", allowed_tools=None
|
||||
)
|
||||
admitted = UserAPIKeyAuth(api_key=api_key, user_id="u")
|
||||
admitted.authenticated_by_custom_auth = custom_auth
|
||||
proxy_logging = MagicMock()
|
||||
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
|
||||
proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={})
|
||||
proxy_logging.pre_call_hook = AsyncMock(return_value=None)
|
||||
|
||||
await manager.pre_call_tool_check(
|
||||
server_name="srv",
|
||||
name="turn",
|
||||
arguments={},
|
||||
user_api_key_auth=admitted,
|
||||
proxy_logging_obj=proxy_logging,
|
||||
server=server,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
|
||||
kwargs: Final = proxy_logging._create_mcp_request_object_from_kwargs.call_args.args[0]
|
||||
assert kwargs["incoming_bearer_token"] == api_key
|
||||
assert kwargs["incoming_subject_token"] == expected_subject
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_tool_permission_for_key_team_allows_permitted_tool(self):
|
||||
"""
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -650,6 +650,10 @@ async def test_per_user_oauth_missing_stored_token_returns_preemptive_401():
|
|||
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name",
|
||||
return_value=oauth_server,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to",
|
||||
return_value=oauth_server,
|
||||
),
|
||||
patch.object(
|
||||
session_manager_stateless,
|
||||
"handle_request",
|
||||
|
|
@ -738,6 +742,10 @@ async def test_admitted_subject_missing_stored_token_challenged_with_resource_me
|
|||
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name",
|
||||
return_value=oauth_server,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to",
|
||||
return_value=oauth_server,
|
||||
),
|
||||
patch.object(
|
||||
session_manager_stateless,
|
||||
"handle_request",
|
||||
|
|
@ -1033,6 +1041,10 @@ async def test_per_user_oauth_with_stored_token_skips_preemptive_401():
|
|||
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name",
|
||||
return_value=oauth_server,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to",
|
||||
return_value=oauth_server,
|
||||
),
|
||||
patch.object(
|
||||
session_manager_stateless,
|
||||
"handle_request",
|
||||
|
|
@ -1136,6 +1148,10 @@ async def test_handle_streamable_http_mcp_delegated_server_without_token_returns
|
|||
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name",
|
||||
return_value=delegated_server,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to",
|
||||
return_value=delegated_server,
|
||||
),
|
||||
patch.object(
|
||||
session_manager_stateful,
|
||||
"handle_request",
|
||||
|
|
@ -1224,6 +1240,10 @@ async def test_handle_streamable_http_mcp_token_exchange_without_subject_returns
|
|||
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name",
|
||||
return_value=obo_server,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to",
|
||||
return_value=obo_server,
|
||||
),
|
||||
patch.object(
|
||||
session_manager_stateful,
|
||||
"handle_request",
|
||||
|
|
@ -1323,6 +1343,10 @@ async def test_handle_streamable_http_mcp_oauth_delegate_without_token_returns_g
|
|||
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name",
|
||||
return_value=od_server,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to",
|
||||
return_value=od_server,
|
||||
),
|
||||
patch.object(
|
||||
session_manager_stateful,
|
||||
"handle_request",
|
||||
|
|
@ -1580,6 +1604,10 @@ async def test_handle_streamable_http_mcp_true_passthrough_without_token_surface
|
|||
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name",
|
||||
return_value=tp_server,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to",
|
||||
return_value=tp_server,
|
||||
),
|
||||
patch.object(
|
||||
session_manager_stateful,
|
||||
"handle_request",
|
||||
|
|
@ -1648,6 +1676,10 @@ async def test_handle_streamable_http_mcp_true_passthrough_dcr_bridge_challenges
|
|||
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name",
|
||||
return_value=bridge_server,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to",
|
||||
return_value=bridge_server,
|
||||
),
|
||||
patch.object(
|
||||
session_manager_stateful,
|
||||
"handle_request",
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -39,6 +39,7 @@ def test_convert_mcp_to_llm_format_returns_synthetic_data(proxy_logging, make_mc
|
|||
"user_api_key_hash": "hash",
|
||||
"user_api_key_request_route": "/mcp",
|
||||
"incoming_bearer_token": "tok",
|
||||
"incoming_subject_token": "a.b.c",
|
||||
},
|
||||
)
|
||||
snapshot = {
|
||||
|
|
@ -47,6 +48,7 @@ def test_convert_mcp_to_llm_format_returns_synthetic_data(proxy_logging, make_mc
|
|||
"mcp_tool_name": out["mcp_tool_name"],
|
||||
"mcp_arguments": out["mcp_arguments"],
|
||||
"incoming_bearer_token": out["incoming_bearer_token"],
|
||||
"incoming_subject_token": out["incoming_subject_token"],
|
||||
"message_role": out["messages"][0]["role"],
|
||||
}
|
||||
assert snapshot == {
|
||||
|
|
@ -55,6 +57,7 @@ def test_convert_mcp_to_llm_format_returns_synthetic_data(proxy_logging, make_mc
|
|||
"mcp_tool_name": "search",
|
||||
"mcp_arguments": {"q": "hello"},
|
||||
"incoming_bearer_token": "tok",
|
||||
"incoming_subject_token": "a.b.c",
|
||||
"message_role": "user",
|
||||
}
|
||||
|
||||
|
|
@ -66,12 +69,14 @@ def test_convert_mcp_to_llm_format_defaults_model(proxy_logging, make_mcp_reques
|
|||
"model": out["model"],
|
||||
"mcp_tool_name": out["mcp_tool_name"],
|
||||
"incoming_bearer_token": out["incoming_bearer_token"],
|
||||
"incoming_subject_token": out["incoming_subject_token"],
|
||||
"user_id": out["user_api_key_user_id"],
|
||||
}
|
||||
assert snapshot == {
|
||||
"model": "mcp-tool-call",
|
||||
"mcp_tool_name": "calculator",
|
||||
"incoming_bearer_token": None,
|
||||
"incoming_subject_token": None,
|
||||
"user_id": None,
|
||||
}
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue