mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
feat(mcp): fold gateway sign-in into a caller sign-in contract on the challenge path
Replace the parallel gateway sign-in provider registry with CallerSignInProvider, merged into the existing oauth2_token_exchange challenge: gates key off caller_sign_in_for() returning non-None, the resolver answers by server_id, case-insensitive name, and short prefix like the router, keeps_caller_authorization covers OBO servers, and the Agent 365 guardrail exchanges through an injected TokenExchanger. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
9c68ca4412
commit
4cadb402e5
15 changed files with 1005 additions and 477 deletions
124
litellm/proxy/_experimental/mcp_server/caller_sign_in.py
Normal file
124
litellm/proxy/_experimental/mcp_server/caller_sign_in.py
Normal file
|
|
@ -0,0 +1,124 @@
|
|||
"""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
|
||||
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Final, Protocol, cast, runtime_checkable
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
|
||||
|
||||
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, ...]
|
||||
|
||||
|
||||
@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)."""
|
||||
...
|
||||
|
||||
|
||||
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]] = cast(Mapping[str, object], general_settings)
|
||||
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 = [
|
||||
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(issuer for contribution in contributions for issuer in contribution.issuers))
|
||||
scopes: Final = tuple(dict.fromkeys(scope for contribution in contributions for scope in contribution.scopes))
|
||||
return CallerSignIn(issuers=issuers, scopes=scopes)
|
||||
|
|
@ -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,
|
||||
|
|
@ -2580,9 +2581,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 +2604,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": list(sign_in.issuers),
|
||||
"resource": resource_url,
|
||||
"scopes_supported": (mcp_server.scopes if mcp_server.scopes else []),
|
||||
"scopes_supported": list(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)
|
||||
|
|
|
|||
|
|
@ -168,6 +168,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,
|
||||
)
|
||||
|
|
@ -4241,6 +4242,12 @@ 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,
|
||||
)
|
||||
|
||||
if subject_token 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())
|
||||
return
|
||||
resolved_server: Final = await self.ensure_oauth_metadata_discovered(server)
|
||||
spec: Final = _to_server_spec_fail_closed(resolved_server)
|
||||
|
|
@ -5964,13 +5971,8 @@ 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 ") :]
|
||||
# Admission credentials are never handed to guardrails as the caller's assertion.
|
||||
incoming_bearer_token: Final = self._extract_subject_token(None, raw_headers, user_api_key_auth)
|
||||
|
||||
pre_hook_kwargs: Final = {
|
||||
"guardrail_context": guardrail_context,
|
||||
|
|
@ -7189,6 +7191,18 @@ class MCPServerManager:
|
|||
return server
|
||||
return None
|
||||
|
||||
def get_mcp_server_answering_to(self, name: str, client_ip: str | None = None) -> MCPServer | None:
|
||||
"""The server a scoped ``/mcp/{name}`` connect resolves to, matched the way the router matches
|
||||
it: case-insensitive over server_id, name and every published prefix form."""
|
||||
return next(
|
||||
(
|
||||
server
|
||||
for server in self.get_filtered_registry(client_ip).values()
|
||||
if 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,
|
||||
)
|
||||
|
|
@ -496,8 +497,7 @@ 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)
|
||||
|
||||
|
||||
async def raise_denied_scoped_mcp_access(
|
||||
|
|
@ -1768,7 +1768,11 @@ def _challenge_missing_token_exchange_subject(
|
|||
warm path already raises. Gated to servers the key may reach so an unauthorized caller
|
||||
learns nothing about the catalog.
|
||||
"""
|
||||
if server is None or server.auth_type != MCPAuth.oauth2_token_exchange:
|
||||
from litellm.proxy._experimental.mcp_server.caller_sign_in import (
|
||||
caller_sign_in_for, # noqa: PLC0415 # lazy: caller_sign_in pulls the proxy graph
|
||||
)
|
||||
|
||||
if server is None or caller_sign_in_for(server, user_api_key_auth) is None:
|
||||
return
|
||||
if requested_server is not None and requested_server.server_id != server.server_id:
|
||||
return
|
||||
|
|
|
|||
|
|
@ -79,6 +79,11 @@ async def _post_exchange_endpoint(
|
|||
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:
|
||||
|
|
|
|||
|
|
@ -36,6 +36,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,
|
||||
|
|
@ -1580,6 +1581,21 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
return user_api_key_auth.model_copy(update={"object_permission": updated_op, "mcp_toolset_id": toolset_id})
|
||||
|
||||
async def _key_granted_single_server(
|
||||
server: MCPServer,
|
||||
mcp_servers: Sequence[str] | None,
|
||||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
client_ip: str | None,
|
||||
) -> bool:
|
||||
"""Sign-in challenges are issued 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."""
|
||||
if len(mcp_servers or []) != 1:
|
||||
return False
|
||||
allowed: Final = await operations._get_allowed_mcp_servers(
|
||||
user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip
|
||||
)
|
||||
return any(granted.server_id == server.server_id for granted in allowed)
|
||||
|
||||
async def _raise_preemptive_401_for_unauthenticated_servers(
|
||||
scope: Scope,
|
||||
mcp_servers: list[str] | None,
|
||||
|
|
@ -1602,7 +1618,7 @@ 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)
|
||||
server = operations.global_mcp_server_manager.get_mcp_server_answering_to(server_name, client_ip=client_ip)
|
||||
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
|
||||
|
|
@ -1698,12 +1714,22 @@ 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. Non-OBO gates fire only on a
|
||||
# single-server connect the key's grant admits.
|
||||
sign_in = caller_sign_in_for(server, user_api_key_auth) if server is not None else None
|
||||
if (
|
||||
server
|
||||
and sign_in is not None
|
||||
and operations.global_mcp_server_manager._extract_subject_token( # pyright: ignore[reportPrivateUsage] # the manager owns the subject/admission filter shared with the preflight
|
||||
oauth2_headers, raw_headers, user_api_key_auth
|
||||
)
|
||||
is None
|
||||
and (
|
||||
server.auth_type == MCPAuth.oauth2_token_exchange
|
||||
or await _key_granted_single_server(server, mcp_servers, user_api_key_auth, client_ip)
|
||||
)
|
||||
):
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph
|
||||
raise_token_exchange_challenge,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -370,6 +370,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.
|
||||
|
|
|
|||
|
|
@ -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,18 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_provider import (
|
||||
build_token_exchanger,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import TokenExchanger
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
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,38 +57,14 @@ 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 ()
|
||||
|
||||
|
||||
def _parse_tool_input_schema(raw: object) -> Mapping[str, object] | None:
|
||||
|
|
@ -129,27 +113,6 @@ class _BlockedDetail(TypedDict):
|
|||
correlation_id: ReadOnly[str | None]
|
||||
|
||||
|
||||
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
|
||||
|
||||
@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)
|
||||
|
||||
|
||||
class Agent365MalformedResponseError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class Agent365ThrottledError(Exception):
|
||||
def __init__(self, status_code: int) -> None:
|
||||
super().__init__(f"HTTP {status_code}")
|
||||
|
|
@ -173,6 +136,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 +156,19 @@ 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()
|
||||
verbose_proxy_logger.info("Initialized Microsoft Agent 365 guardrail: %s", guardrail_name)
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -233,42 +208,40 @@ class Agent365Guardrail(CustomGuardrail):
|
|||
)
|
||||
|
||||
try:
|
||||
obo_token: Final = await self._get_obo_token(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})",
|
||||
)
|
||||
except Agent365ThrottledError as exc:
|
||||
self._handle_throttled(
|
||||
data=data,
|
||||
tool_name=tool_name,
|
||||
reason=f"the Entra token endpoint returned HTTP {exc.status_code}",
|
||||
latency_ms=None,
|
||||
)
|
||||
exchange_result: Final = await self._exchange_caller_assertion(assertion)
|
||||
except (httpx.HTTPError, LitellmTimeout, TimeoutError) as exc:
|
||||
return self._handle_unavailable(
|
||||
data=data,
|
||||
tool_name=tool_name,
|
||||
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),
|
||||
)
|
||||
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 "misconfigured":
|
||||
return self._handle_unavailable(
|
||||
data=data,
|
||||
tool_name=tool_name,
|
||||
reason=(
|
||||
f"Entra rejected the gateway's own Agent 365 credentials ({error.misconfigured}); "
|
||||
"check the guardrail's client_id and client_secret"
|
||||
),
|
||||
)
|
||||
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 +257,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 +281,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,62 +430,29 @@ 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]
|
||||
|
||||
response: Final = await self._post_allowing_error_status(
|
||||
url=TOKEN_ENDPOINT_TEMPLATE.format(tenant_id=self.tenant_id),
|
||||
data={ # mutable-ok: OAuth form body; AsyncHTTPHandler.post requires dict
|
||||
"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"}, # mutable-ok: httpx header dict
|
||||
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 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
|
||||
if user_api_key_auth is not None:
|
||||
probe: Final[dict[str, Mapping[str, object]]] = { # pyright: ignore[reportUnknownVariableType] # UserAPIKeyAuth metadata dicts are untyped
|
||||
"metadata": {
|
||||
"user_api_key_metadata": user_api_key_auth.metadata, # pyright: ignore[reportUnknownMemberType] # UserAPIKeyAuth.metadata is a raw dict
|
||||
"user_api_key_team_metadata": user_api_key_auth.team_metadata, # 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=(GATEWAY_SCOPE_TEMPLATE.format(client_id=self.client_id),),
|
||||
)
|
||||
if response.status_code in (408, 429):
|
||||
raise Agent365ThrottledError(status_code=response.status_code)
|
||||
if response.status_code >= 500:
|
||||
raise httpx.HTTPStatusError(
|
||||
f"Entra token endpoint returned {response.status_code}",
|
||||
request=response.request,
|
||||
response=response,
|
||||
)
|
||||
try:
|
||||
parsed_body: Final = 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
|
||||
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")),
|
||||
)
|
||||
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:
|
||||
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
|
||||
|
||||
async def _post_allowing_error_status(
|
||||
self,
|
||||
|
|
@ -577,11 +519,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
|
||||
|
|
|
|||
|
|
@ -313,10 +313,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,
|
||||
|
|
@ -326,11 +326,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
|
||||
|
|
|
|||
|
|
@ -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, ServerSpec
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_provider import (
|
||||
_post_exchange_endpoint,
|
||||
build_token_exchanger,
|
||||
|
|
@ -17,17 +19,19 @@ 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:
|
||||
|
|
@ -80,7 +84,7 @@ 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"}, {})
|
||||
|
||||
|
|
@ -93,7 +97,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"}, {})
|
||||
|
||||
|
|
@ -128,7 +132,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 +141,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 +152,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"
|
||||
|
|
|
|||
|
|
@ -0,0 +1,134 @@
|
|||
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,
|
||||
CallerSignInProvider,
|
||||
caller_sign_in_for,
|
||||
)
|
||||
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,))
|
||||
|
||||
|
||||
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]
|
||||
|
|
@ -7288,7 +7288,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,15 +7307,15 @@ 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"],
|
||||
"resource": _OBO_RESOURCE,
|
||||
|
|
@ -7323,32 +7323,32 @@ def test_obo_protected_resource_response_names_jwt_issuers():
|
|||
}
|
||||
|
||||
|
||||
def test_obo_protected_resource_response_scopes_default_empty():
|
||||
def test_caller_sign_in_protected_resource_response_scopes_default_empty():
|
||||
"""A scopeless OBO server reports scopes_supported as [] 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)
|
||||
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 +7360,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
|
||||
|
|
|
|||
|
|
@ -1,4 +1,3 @@
|
|||
from litellm.proxy._experimental.mcp_server import operations as mcp_operations
|
||||
import asyncio
|
||||
import contextlib
|
||||
import contextvars
|
||||
|
|
@ -29,6 +28,9 @@ from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERS
|
|||
from pydantic import TypeAdapter
|
||||
from starlette.types import Message, Receive, Scope, Send
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.proxy._experimental.mcp_server import operations as mcp_operations
|
||||
from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
|
||||
from litellm.proxy._types import (
|
||||
|
|
@ -1159,6 +1161,7 @@ async def test_mcp_read_resource_success():
|
|||
)
|
||||
async def test_read_resource_preserves_content_metadata(_mcp_request_ctx, kind, metadata):
|
||||
from mcp.types import ReadResourceRequestParams
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import operations, server
|
||||
|
||||
uri: Final = "https://example.com/resource"
|
||||
|
|
@ -9881,7 +9884,7 @@ class TestPreemptive401ModeAware:
|
|||
with (
|
||||
patch.object(
|
||||
mcp_operations.global_mcp_server_manager,
|
||||
"get_mcp_server_by_name",
|
||||
"get_mcp_server_answering_to",
|
||||
return_value=server,
|
||||
),
|
||||
patch.object(
|
||||
|
|
@ -9980,7 +9983,7 @@ class TestPreemptive401ModeAware:
|
|||
patch.dict(os.environ, {"SERVER_ROOT_PATH": "/litellm"}),
|
||||
patch.object(
|
||||
mcp_operations.global_mcp_server_manager,
|
||||
"get_mcp_server_by_name",
|
||||
"get_mcp_server_answering_to",
|
||||
return_value=server,
|
||||
),
|
||||
patch.object(
|
||||
|
|
@ -10087,7 +10090,7 @@ class TestSingleServerPreflightReachesIdJag:
|
|||
with (
|
||||
patch.object( # test-quality-ok: route wiring must use the manager's configured server
|
||||
mcp_operations.global_mcp_server_manager,
|
||||
"get_mcp_server_by_name",
|
||||
"get_mcp_server_answering_to",
|
||||
return_value=server,
|
||||
),
|
||||
patch.object( # test-quality-ok: route wiring must invoke the manager preflight
|
||||
|
|
@ -10147,7 +10150,7 @@ class TestSingleServerPreflightReachesIdJag:
|
|||
with (
|
||||
patch.object( # test-quality-ok: route wiring must use the manager's configured server
|
||||
mcp_operations.global_mcp_server_manager,
|
||||
"get_mcp_server_by_name",
|
||||
"get_mcp_server_answering_to",
|
||||
return_value=token_exchange,
|
||||
),
|
||||
patch.object( # test-quality-ok: route wiring must invoke the manager preflight
|
||||
|
|
@ -10212,7 +10215,7 @@ class TestOboPreflightScopedToAllowedServers:
|
|||
preflight = AsyncMock()
|
||||
with (
|
||||
patch.object( # test-quality-ok: route handler reads the module-level manager, no injection seam
|
||||
mcp_operations.global_mcp_server_manager, "get_mcp_server_by_name", return_value=requested
|
||||
mcp_operations.global_mcp_server_manager, "get_mcp_server_answering_to", return_value=requested
|
||||
),
|
||||
patch.object( # test-quality-ok: the exchanger is the observable; a real one would call an IdP
|
||||
mcp_operations.global_mcp_server_manager, "preflight_token_exchange", preflight
|
||||
|
|
@ -10228,6 +10231,10 @@ class TestOboPreflightScopedToAllowedServers:
|
|||
mcp_server_auth_headers=None,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
client_ip="10.0.0.7",
|
||||
raw_headers={
|
||||
"x-litellm-api-key": user_api_key_auth.api_key if user_api_key_auth else "",
|
||||
"authorization": self.SUBJECT_HEADERS["Authorization"],
|
||||
},
|
||||
)
|
||||
return allowed_lookup, preflight
|
||||
|
||||
|
|
@ -10253,7 +10260,13 @@ class TestOboPreflightScopedToAllowedServers:
|
|||
_, preflight = await self._run(requested, allowed=[requested], user_api_key_auth=key)
|
||||
|
||||
preflight.assert_awaited_once_with(
|
||||
server=requested, oauth2_headers=self.SUBJECT_HEADERS, user_api_key_auth=key, raw_headers=None
|
||||
server=requested,
|
||||
oauth2_headers=self.SUBJECT_HEADERS,
|
||||
user_api_key_auth=key,
|
||||
raw_headers={
|
||||
"x-litellm-api-key": key.api_key,
|
||||
"authorization": self.SUBJECT_HEADERS["Authorization"],
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -10774,8 +10787,8 @@ async def test_native_listing_preserves_empty_result_on_auth_failure(_mcp_reques
|
|||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("failure_hook_raises", [False, True])
|
||||
async def test_tool_listing_preserves_permission_denial_when_failure_logging_fails(failure_hook_raises):
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
|
||||
auth = UserAPIKeyAuth(user_id="denied-caller")
|
||||
denial = HTTPException(status_code=403, detail="scope denied")
|
||||
|
|
@ -10811,8 +10824,9 @@ async def test_legacy_sse_mount_emits_message_endpoint(
|
|||
) -> None:
|
||||
from starlette.applications import Starlette
|
||||
from starlette.routing import Mount
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import server as mcp_server
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing
|
||||
|
||||
app: Final = Starlette(routes=[Mount("/mcp", app=mcp_server.app)])
|
||||
incoming: Final[asyncio.Queue[Message]] = asyncio.Queue()
|
||||
|
|
@ -10970,6 +10984,7 @@ def test_protocol_header_respects_configured_advertisement(revision, rejected):
|
|||
@pytest.mark.asyncio
|
||||
async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ctx):
|
||||
from mcp.types import DiscoverResult, RequestParams, ServerCapabilities
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import server
|
||||
|
||||
expected = DiscoverResult(supported_versions=["2025-11-25"], capabilities=ServerCapabilities())
|
||||
|
|
@ -10988,3 +11003,151 @@ async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ct
|
|||
context = dispatched.await_args.args[1]
|
||||
assert context.user_api_key_auth.user_id == "discover-caller"
|
||||
assert context.mcp_servers == ("allowed",)
|
||||
|
||||
|
||||
def _catalog_server() -> MCPServer:
|
||||
return 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"},
|
||||
)
|
||||
|
||||
|
||||
class _CallerSignInGuardrail(CustomGuardrail):
|
||||
def caller_sign_in(self, server, user_api_key_auth):
|
||||
from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn
|
||||
|
||||
return CallerSignIn(issuers=("https://idp.test",), scopes=("scope-a",))
|
||||
|
||||
|
||||
class TestConnectChallengeResolver:
|
||||
"""The connect-time sign-in challenge must resolve the server the same way the router resolves
|
||||
``/mcp/{name}``: server_id, case-insensitive name, and the short prefix all reach it."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"route_name",
|
||||
[
|
||||
"catalog-server-id-001",
|
||||
"CATALOG",
|
||||
],
|
||||
ids=["server_id", "uppercase_name"],
|
||||
)
|
||||
async def test_provider_gated_server_challenged_on_every_route_spelling(self, route_name):
|
||||
from litellm.proxy._experimental.mcp_server import server as server_module
|
||||
|
||||
server = _catalog_server()
|
||||
guardrail = _CallerSignInGuardrail(guardrail_name="sign-in-stub")
|
||||
litellm.logging_callback_manager.add_litellm_callback(guardrail)
|
||||
preflight = AsyncMock()
|
||||
try:
|
||||
with (
|
||||
patch.object(
|
||||
mcp_operations.global_mcp_server_manager,
|
||||
"get_filtered_registry",
|
||||
return_value={server.server_id: server},
|
||||
),
|
||||
patch.object(
|
||||
mcp_operations.global_mcp_server_manager,
|
||||
"preflight_token_exchange",
|
||||
preflight,
|
||||
),
|
||||
patch.object(
|
||||
mcp_operations,
|
||||
"_get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=[server]),
|
||||
),
|
||||
pytest.raises(HTTPException) as exc,
|
||||
):
|
||||
await server_module._raise_preemptive_401_for_unauthenticated_servers(
|
||||
scope={"type": "http", "method": "POST", "path": f"/mcp/{route_name}", "headers": []},
|
||||
mcp_servers=[route_name],
|
||||
oauth2_headers=None,
|
||||
mcp_server_auth_headers=None,
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
|
||||
client_ip=None,
|
||||
)
|
||||
finally:
|
||||
litellm.logging_callback_manager.remove_callback_from_list_by_object(
|
||||
litellm.callbacks, guardrail, require_self=False
|
||||
)
|
||||
|
||||
assert exc.value.status_code == 401
|
||||
headers = exc.value.headers or {}
|
||||
assert "resource_metadata" in (headers.get("WWW-Authenticate") or "")
|
||||
preflight.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_provider_gated_server_challenged_on_short_prefix_route(self):
|
||||
from litellm.proxy._experimental.mcp_server import server as server_module
|
||||
from litellm.proxy._experimental.mcp_server.utils import compute_short_server_prefix
|
||||
|
||||
server = _catalog_server()
|
||||
route_name = compute_short_server_prefix(server.server_id)
|
||||
guardrail = _CallerSignInGuardrail(guardrail_name="sign-in-stub")
|
||||
litellm.logging_callback_manager.add_litellm_callback(guardrail)
|
||||
try:
|
||||
with (
|
||||
patch.object(
|
||||
mcp_operations.global_mcp_server_manager,
|
||||
"get_filtered_registry",
|
||||
return_value={server.server_id: server},
|
||||
),
|
||||
patch.object(
|
||||
mcp_operations,
|
||||
"_get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=[server]),
|
||||
),
|
||||
pytest.raises(HTTPException) as exc,
|
||||
):
|
||||
await server_module._raise_preemptive_401_for_unauthenticated_servers(
|
||||
scope={"type": "http", "method": "POST", "path": f"/mcp/{route_name}", "headers": []},
|
||||
mcp_servers=[route_name],
|
||||
oauth2_headers=None,
|
||||
mcp_server_auth_headers=None,
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
|
||||
client_ip=None,
|
||||
)
|
||||
finally:
|
||||
litellm.logging_callback_manager.remove_callback_from_list_by_object(
|
||||
litellm.callbacks, guardrail, require_self=False
|
||||
)
|
||||
|
||||
assert exc.value.status_code == 401
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_obo_challenge_www_authenticate_matches_main_byte_for_byte(self, monkeypatch):
|
||||
"""The provider redesign must not change what an OBO server challenges with: the relative
|
||||
RFC 9728 resource_metadata path plus the RFC 6750 invalid_token triple, exactly as main."""
|
||||
from litellm.proxy._experimental.mcp_server import server as server_module
|
||||
|
||||
monkeypatch.delenv("SERVER_ROOT_PATH", raising=False)
|
||||
obo = _make_obo_server("obo")
|
||||
with (
|
||||
patch.object(
|
||||
mcp_operations.global_mcp_server_manager,
|
||||
"get_mcp_server_answering_to",
|
||||
return_value=obo,
|
||||
),
|
||||
pytest.raises(HTTPException) as exc,
|
||||
):
|
||||
await server_module._raise_preemptive_401_for_unauthenticated_servers(
|
||||
scope={"type": "http", "method": "POST", "path": "/mcp/obo", "headers": []},
|
||||
mcp_servers=["obo"],
|
||||
oauth2_headers=None,
|
||||
mcp_server_auth_headers=None,
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
|
||||
client_ip=None,
|
||||
)
|
||||
|
||||
assert exc.value.status_code == 401
|
||||
assert (exc.value.headers or {}).get("WWW-Authenticate") == (
|
||||
'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/obo", '
|
||||
'error="invalid_token", '
|
||||
'error_description="Missing or invalid subject token; authenticate with the IdP and retry"'
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1464,6 +1464,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",
|
||||
|
|
|
|||
|
|
@ -14,6 +14,14 @@ from litellm.exceptions import Timeout as LitellmTimeout
|
|||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.secret_redaction import redact_string
|
||||
from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn, caller_sign_in_for
|
||||
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.types import (
|
||||
CredError,
|
||||
ServerSpec,
|
||||
TokenExchangeConfig,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.guardrails.guardrail_hooks.agent_365 import (
|
||||
Agent365Guardrail,
|
||||
|
|
@ -27,9 +35,12 @@ from litellm.types.guardrails import (
|
|||
LitellmParams,
|
||||
SupportedGuardrailIntegrations,
|
||||
)
|
||||
from litellm.types.mcp import MCPAuth, MCPTransport
|
||||
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,
|
||||
AGENT_365_SCOPE_NAME,
|
||||
Agent365GuardrailConfigModel,
|
||||
)
|
||||
|
||||
|
|
@ -45,8 +56,46 @@ def _response(status_code: int, payload: Any = None, text: str | None = None) ->
|
|||
return httpx.Response(status_code=status_code, text=text or "", request=request)
|
||||
|
||||
|
||||
def _token_response(access_token: str = "obo-access-token", expires_in: int = 3599) -> httpx.Response:
|
||||
return _response(200, {"access_token": access_token, "expires_in": expires_in})
|
||||
class StubTokenExchanger:
|
||||
"""The TokenExchanger the guardrail is injected with in tests: programmed Result queue plus a
|
||||
per-subject cache honoring ``expires_at``, so cache and evaluate-401-invalidate behavior is
|
||||
exercised the way the real OboTokenExchanger drives it."""
|
||||
|
||||
def __init__(self, results: list[Result[OAuthToken, CredError] | BaseException] | None = None):
|
||||
self._results = list(results or [])
|
||||
self._cache: dict[str, OAuthToken] = {}
|
||||
self.calls: list[tuple[str, ServerSpec, TokenExchangeConfig]] = []
|
||||
self.invalidations: list[str] = []
|
||||
|
||||
async def exchange(
|
||||
self, subject_token: str, server: ServerSpec, config: TokenExchangeConfig, *, tenant_id: str = ""
|
||||
) -> Result[OAuthToken, CredError]:
|
||||
cached: Final = self._cache.get(subject_token)
|
||||
if cached is not None and (cached.expires_at is None or cached.expires_at > time.time()):
|
||||
return Ok(cached)
|
||||
self.calls.append((subject_token, server, config))
|
||||
if not self._results:
|
||||
raise AssertionError("StubTokenExchanger ran out of programmed results")
|
||||
result = self._results.pop(0)
|
||||
if isinstance(result, BaseException):
|
||||
raise result
|
||||
if isinstance(result, Ok):
|
||||
self._cache[subject_token] = result.ok
|
||||
return result
|
||||
|
||||
async def invalidate(
|
||||
self, subject_token: str, server: ServerSpec, config: TokenExchangeConfig, *, tenant_id: str = ""
|
||||
) -> None:
|
||||
self.invalidations.append(subject_token)
|
||||
self._cache.pop(subject_token, None)
|
||||
|
||||
|
||||
def _ok_exchange(access_token: str = "obo-access-token", expires_in: int = 3599) -> Ok[OAuthToken, CredError]:
|
||||
return Ok(OAuthToken(access_token=access_token, expires_at=time.time() + expires_in))
|
||||
|
||||
|
||||
def _obo_ok(access_token: str = "obo-access-token") -> list[Result[OAuthToken, CredError]]:
|
||||
return [_ok_exchange(access_token)]
|
||||
|
||||
|
||||
def _allow_response(correlation_id: str = "corr-1") -> httpx.Response:
|
||||
|
|
@ -121,7 +170,9 @@ class FakeHandler:
|
|||
def _make_guardrail(
|
||||
handler: FakeHandler,
|
||||
*,
|
||||
exchanger: StubTokenExchanger | None = None,
|
||||
unreachable_fallback: str = "fail_closed",
|
||||
default_on: bool = True,
|
||||
) -> Agent365Guardrail:
|
||||
return Agent365Guardrail(
|
||||
guardrail_name="agent-365-guard",
|
||||
|
|
@ -130,21 +181,28 @@ def _make_guardrail(
|
|||
client_secret="secret-123",
|
||||
unreachable_fallback=unreachable_fallback,
|
||||
async_handler=handler,
|
||||
token_exchanger=exchanger if exchanger is not None else StubTokenExchanger(_obo_ok()),
|
||||
event_hook="pre_mcp_call",
|
||||
default_on=True,
|
||||
default_on=default_on,
|
||||
)
|
||||
|
||||
|
||||
def _default_fallback_guardrail(handler: FakeHandler) -> Agent365Guardrail:
|
||||
return Agent365Guardrail(
|
||||
guardrail_name="agent-365-guard",
|
||||
tenant_id="tenant-abc",
|
||||
client_id="client-xyz",
|
||||
client_secret="secret-123",
|
||||
async_handler=handler,
|
||||
event_hook="pre_mcp_call",
|
||||
default_on=True,
|
||||
)
|
||||
def _default_fallback_guardrail(
|
||||
handler: FakeHandler, exchanger: StubTokenExchanger | None = None
|
||||
) -> Agent365Guardrail:
|
||||
return _make_guardrail(handler, exchanger=exchanger)
|
||||
|
||||
|
||||
def _server(**overrides: Any) -> MCPServer:
|
||||
kwargs: Final[dict] = {
|
||||
"server_id": "outlook-id",
|
||||
"name": "outlook_mcp",
|
||||
"server_name": "outlook_mcp",
|
||||
"transport": MCPTransport.http,
|
||||
"url": "https://outlook.test/mcp",
|
||||
}
|
||||
kwargs.update(overrides)
|
||||
return MCPServer(**kwargs)
|
||||
|
||||
|
||||
def _mcp_data(**overrides: Any) -> dict:
|
||||
|
|
@ -257,14 +315,13 @@ class TestInitializeGuardrail:
|
|||
resource_app_id="00000000-0000-0000-0000-000000000000",
|
||||
agent_id="yaml-agent",
|
||||
)
|
||||
handler: Final = FakeHandler([_token_response(), _allow_response()])
|
||||
handler: Final = FakeHandler([_allow_response()])
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
guardrail: Final = initialize_guardrail(params, {"guardrail_name": "a365-stale"}, async_handler=handler)
|
||||
initialize_guardrail(params, {"guardrail_name": "a365-stale"}, async_handler=handler)
|
||||
assert "ignoring api_base, resource_app_id, agent_id" in caplog.text
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
await _run(guardrail, _mcp_data())
|
||||
token_call, evaluate_call = handler.calls
|
||||
assert token_call.url == TOKEN_URL
|
||||
assert token_call.data["scope"] == f"{AGENT_365_PROD_RESOURCE_APP_ID}/ThreatProtection.Evaluate.All"
|
||||
evaluate_call: Final = handler.calls[0]
|
||||
assert evaluate_call.url == EVALUATE_URL
|
||||
assert evaluate_call.json["agentId"] == "my-agent-key"
|
||||
|
||||
|
|
@ -305,7 +362,7 @@ def _guardrail_info(data: dict) -> dict:
|
|||
class TestAllowFlow:
|
||||
@pytest.mark.asyncio
|
||||
async def test_allowed_call_passes_through(self):
|
||||
handler: Final = FakeHandler([_token_response(), _allow_response()])
|
||||
handler: Final = FakeHandler([_allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
data: Final = _mcp_data()
|
||||
result: Final = await _run(guardrail, data)
|
||||
|
|
@ -319,25 +376,27 @@ class TestAllowFlow:
|
|||
assert info["guardrail_response"]["latency_ms"] >= 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_obo_exchange_form(self):
|
||||
handler: Final = FakeHandler([_token_response(), _allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
async def test_obo_exchange_uses_the_built_entra_obo_config(self):
|
||||
exchanger: Final = StubTokenExchanger(_obo_ok())
|
||||
handler: Final = FakeHandler([_allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
|
||||
await _run(guardrail, _mcp_data())
|
||||
token_call: Final = handler.calls[0]
|
||||
assert token_call.url == TOKEN_URL
|
||||
assert token_call.data["grant_type"] == "urn:ietf:params:oauth:grant-type:jwt-bearer"
|
||||
assert token_call.data["requested_token_use"] == "on_behalf_of"
|
||||
assert token_call.data["assertion"] == FAKE_ASSERTION
|
||||
assert token_call.data["client_id"] == "client-xyz"
|
||||
assert token_call.data["client_secret"] == "secret-123"
|
||||
assert token_call.data["scope"] == f"{AGENT_365_PROD_RESOURCE_APP_ID}/ThreatProtection.Evaluate.All"
|
||||
subject_token, server, config = exchanger.calls[0]
|
||||
assert subject_token == FAKE_ASSERTION
|
||||
assert server.server_id == "agent-365:tenant-abc"
|
||||
assert server.resource == AGENT_365_PROD_API_BASE
|
||||
assert config.profile == "entra_obo"
|
||||
assert config.token_exchange_endpoint == TOKEN_URL
|
||||
assert config.client_id == "client-xyz"
|
||||
assert config.client_secret is not None and config.client_secret.get_secret_value() == "secret-123"
|
||||
assert config.scopes == (f"{AGENT_365_PROD_RESOURCE_APP_ID}/{AGENT_365_SCOPE_NAME}",)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_evaluate_payload(self):
|
||||
handler: Final = FakeHandler([_token_response(), _allow_response()])
|
||||
handler: Final = FakeHandler([_allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
await _run(guardrail, _mcp_data())
|
||||
evaluate_call: Final = handler.calls[1]
|
||||
evaluate_call: Final = handler.calls[0]
|
||||
assert evaluate_call.url == EVALUATE_URL
|
||||
assert evaluate_call.headers["Authorization"] == "Bearer obo-access-token"
|
||||
assert evaluate_call.json["tool"] == {"name": "send_email"}
|
||||
|
|
@ -348,11 +407,11 @@ class TestAllowFlow:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_evaluate_payload_includes_listed_tool_metadata(self):
|
||||
handler: Final = FakeHandler([_token_response(), _allow_response()])
|
||||
handler: Final = FakeHandler([_allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
schema: Final = {"type": "object", "properties": {"to": {"type": "string"}}, "required": ["to"]}
|
||||
await _run(guardrail, _mcp_data(mcp_tool_description="Send an email", mcp_input_schema=schema))
|
||||
assert handler.calls[1].json["tool"] == {
|
||||
await _run(guardrail, _mcp_data(mcp_tool_description="Send an email", mcp_tool_input_schema=schema))
|
||||
assert handler.calls[0].json["tool"] == {
|
||||
"name": "send_email",
|
||||
"description": "Send an email",
|
||||
"inputSchema": schema,
|
||||
|
|
@ -364,10 +423,17 @@ class TestAllowFlow:
|
|||
[(None, None), ("", None), (None, ["not", "a", "schema"]), (42, "type: object")],
|
||||
)
|
||||
async def test_evaluate_payload_omits_missing_or_malformed_tool_metadata(self, description, schema):
|
||||
handler: Final = FakeHandler([_token_response(), _allow_response()])
|
||||
handler: Final = FakeHandler([_allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
await _run(guardrail, _mcp_data(mcp_tool_description=description, mcp_input_schema=schema))
|
||||
assert handler.calls[1].json["tool"] == {"name": "send_email"}
|
||||
await _run(guardrail, _mcp_data(mcp_tool_description=description, mcp_tool_input_schema=schema))
|
||||
assert handler.calls[0].json["tool"] == {"name": "send_email"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_id_falls_back_to_key_alias(self):
|
||||
handler: Final = FakeHandler([_allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
await _run(guardrail, _mcp_data())
|
||||
assert handler.calls[0].json["agentId"] == "my-agent-key"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_mcp_call_type_skipped(self):
|
||||
|
|
@ -385,26 +451,26 @@ class TestConversationId:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_two_calls_in_one_session_share_the_conversation_id(self):
|
||||
handler: Final = FakeHandler([_token_response(), _allow_response(), _allow_response()])
|
||||
handler: Final = FakeHandler([_allow_response(), _allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
for call_id in ("call-1", "call-2"):
|
||||
await _run(
|
||||
guardrail,
|
||||
_mcp_data(litellm_call_id=call_id, litellm_logging_obj=_logging_obj(call_id, mcp_session_id="sess-A")),
|
||||
)
|
||||
assert [call.json["conversationId"] for call in handler.calls[1:]] == ["sess-A", "sess-A"]
|
||||
assert [call.json["conversationId"] for call in handler.calls[:]] == ["sess-A", "sess-A"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_server_recorded_session_beats_the_client_header(self):
|
||||
handler: Final = FakeHandler([_token_response(), _allow_response()])
|
||||
handler: Final = FakeHandler([_allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
data: Final = _mcp_data(litellm_logging_obj=_logging_obj("call-id-1", mcp_session_id="sess-from-logging"))
|
||||
await _run(guardrail, data)
|
||||
assert handler.calls[1].json["conversationId"] == "sess-from-logging"
|
||||
assert handler.calls[0].json["conversationId"] == "sess-from-logging"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sessionless_call_falls_back_to_the_request_call_id(self):
|
||||
handler: Final = FakeHandler([_token_response(), _allow_response()])
|
||||
handler: Final = FakeHandler([_allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
data: Final = _mcp_data(
|
||||
metadata={"headers": {}},
|
||||
|
|
@ -412,37 +478,37 @@ class TestConversationId:
|
|||
litellm_logging_obj=_logging_obj("call-id-from-logging"),
|
||||
)
|
||||
await _run(guardrail, data)
|
||||
assert handler.calls[1].json["conversationId"] == "call-id-from-data"
|
||||
assert handler.calls[0].json["conversationId"] == "call-id-from-data"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sessionless_call_without_request_call_id_uses_the_logging_call_id(self):
|
||||
handler: Final = FakeHandler([_token_response(), _allow_response()])
|
||||
handler: Final = FakeHandler([_allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
data: Final = _mcp_data(metadata={"headers": {}}, litellm_logging_obj=_logging_obj("call-id-from-logging"))
|
||||
await _run(guardrail, data)
|
||||
assert handler.calls[1].json["conversationId"] == "call-id-from-logging"
|
||||
assert handler.calls[0].json["conversationId"] == "call-id-from-logging"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_id_header_case_insensitive(self):
|
||||
handler: Final = FakeHandler([_token_response(), _allow_response()])
|
||||
handler: Final = FakeHandler([_allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
data: Final = _mcp_data(metadata={"headers": {"Mcp-Session-Id": "sess-CASED"}})
|
||||
await _run(guardrail, data)
|
||||
assert handler.calls[1].json["conversationId"] == "sess-CASED"
|
||||
assert handler.calls[0].json["conversationId"] == "sess-CASED"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generates_uuid_when_no_identifier_available(self):
|
||||
handler: Final = FakeHandler([_token_response(), _allow_response()])
|
||||
handler: Final = FakeHandler([_allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
await _run(guardrail, _mcp_data(metadata={"headers": {}}, litellm_logging_obj=_logging_obj("")))
|
||||
conversation_id: Final = handler.calls[1].json["conversationId"]
|
||||
conversation_id: Final = handler.calls[0].json["conversationId"]
|
||||
assert uuid.UUID(conversation_id).version == 4
|
||||
|
||||
|
||||
class TestBlockFlow:
|
||||
@pytest.mark.asyncio
|
||||
async def test_blocked_call_raises_400(self):
|
||||
handler: Final = FakeHandler([_token_response(), _block_response(message="Injection detected")])
|
||||
handler: Final = FakeHandler([_block_response(message="Injection detected")])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
data: Final = _mcp_data()
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
|
|
@ -458,7 +524,7 @@ class TestBlockFlow:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_blocked_even_with_fail_open(self):
|
||||
handler: Final = FakeHandler([_token_response(), _block_response()])
|
||||
handler: Final = FakeHandler([_block_response()])
|
||||
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _run(guardrail, _mcp_data())
|
||||
|
|
@ -467,7 +533,7 @@ class TestBlockFlow:
|
|||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("status", ["Skipped", "FailedOpen"])
|
||||
async def test_explicit_block_wins_over_non_evaluated_status(self, status):
|
||||
handler: Final = FakeHandler([_token_response(), _block_response(status=status)])
|
||||
handler: Final = FakeHandler([_block_response(status=status)])
|
||||
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
|
||||
data: Final = _mcp_data()
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
|
|
@ -483,7 +549,7 @@ class TestDefenderNotEvaluated:
|
|||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("status", ["Skipped", "FailedOpen"])
|
||||
async def test_fail_closed_blocks_allowed_but_unevaluated_call(self, status):
|
||||
handler: Final = FakeHandler([_token_response(), _not_evaluated_response(status)])
|
||||
handler: Final = FakeHandler([_not_evaluated_response(status)])
|
||||
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_closed")
|
||||
data: Final = _mcp_data()
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
|
|
@ -500,7 +566,7 @@ class TestDefenderNotEvaluated:
|
|||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("status", ["Skipped", "FailedOpen"])
|
||||
async def test_fail_open_allows_unevaluated_call_as_unscanned(self, status):
|
||||
handler: Final = FakeHandler([_token_response(), _not_evaluated_response(status)])
|
||||
handler: Final = FakeHandler([_not_evaluated_response(status)])
|
||||
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
|
||||
data: Final = _mcp_data()
|
||||
result: Final = await _run(guardrail, data)
|
||||
|
|
@ -514,7 +580,7 @@ class TestDefenderNotEvaluated:
|
|||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("payload", [{"allowed": True}, {"allowed": True, "defender": {"verdict": "Allow"}}])
|
||||
async def test_allowed_without_defender_status_is_not_an_evaluated_allow(self, payload):
|
||||
handler: Final = FakeHandler([_token_response(), _response(200, payload)])
|
||||
handler: Final = FakeHandler([_response(200, payload)])
|
||||
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_closed")
|
||||
data: Final = _mcp_data()
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
|
|
@ -525,7 +591,7 @@ class TestDefenderNotEvaluated:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_http_400_always_blocks_even_fail_open(self):
|
||||
handler: Final = FakeHandler([_token_response(), _response(400, text="Bad request: serverName missing")])
|
||||
handler: Final = FakeHandler([_response(400, text="Bad request: serverName missing")])
|
||||
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _run(guardrail, _mcp_data())
|
||||
|
|
@ -534,18 +600,20 @@ class TestDefenderNotEvaluated:
|
|||
|
||||
|
||||
AVAILABILITY_FAILURES: Final = (
|
||||
pytest.param([_token_response(), httpx.ReadTimeout("timed out")], id="evaluate-timeout"),
|
||||
pytest.param([_token_response(), _response(502, text="bad gateway")], id="evaluate-5xx"),
|
||||
pytest.param([_token_response(), _not_evaluated_response("Skipped")], id="evaluate-skipped"),
|
||||
pytest.param([_response(503, text="entra down")], id="entra-5xx"),
|
||||
pytest.param([httpx.ReadTimeout("timed out")], _obo_ok(), id="evaluate-timeout"),
|
||||
pytest.param([_response(502, text="bad gateway")], _obo_ok(), id="evaluate-5xx"),
|
||||
pytest.param([_not_evaluated_response("Skipped")], _obo_ok(), id="evaluate-skipped"),
|
||||
pytest.param([], [Error(CredError.of_upstream_unavailable("entra down"))], id="entra-5xx"),
|
||||
)
|
||||
|
||||
|
||||
class TestFailOpenOptIn:
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("responses", AVAILABILITY_FAILURES)
|
||||
async def test_constructor_default_blocks_each_availability_failure_with_503(self, responses):
|
||||
guardrail: Final = _default_fallback_guardrail(FakeHandler(responses))
|
||||
@pytest.mark.parametrize(("responses", "exchange_results"), AVAILABILITY_FAILURES)
|
||||
async def test_constructor_default_blocks_each_availability_failure_with_503(self, responses, exchange_results):
|
||||
guardrail: Final = _default_fallback_guardrail(
|
||||
FakeHandler(responses), StubTokenExchanger(exchange_results)
|
||||
)
|
||||
assert guardrail.unreachable_fallback == "fail_closed"
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _run(guardrail, _mcp_data())
|
||||
|
|
@ -553,9 +621,15 @@ class TestFailOpenOptIn:
|
|||
assert "fail_closed" in exc_info.value.detail["message"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("responses", AVAILABILITY_FAILURES)
|
||||
async def test_opted_in_fail_open_lets_each_availability_failure_through_as_failed_to_respond(self, responses):
|
||||
guardrail: Final = _make_guardrail(FakeHandler(responses), unreachable_fallback="fail_open")
|
||||
@pytest.mark.parametrize(("responses", "exchange_results"), AVAILABILITY_FAILURES)
|
||||
async def test_opted_in_fail_open_lets_each_availability_failure_through_as_failed_to_respond(
|
||||
self, responses, exchange_results
|
||||
):
|
||||
guardrail: Final = _make_guardrail(
|
||||
FakeHandler(responses),
|
||||
exchanger=StubTokenExchanger(exchange_results),
|
||||
unreachable_fallback="fail_open",
|
||||
)
|
||||
data: Final = _mcp_data()
|
||||
assert await _run(guardrail, data) is data
|
||||
info: Final = _guardrail_info(data)
|
||||
|
|
@ -564,7 +638,7 @@ class TestFailOpenOptIn:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_opted_in_fail_open_logs_the_unscanned_call_at_error_level(self, caplog):
|
||||
handler: Final = FakeHandler([_token_response(), httpx.ReadTimeout("timed out")])
|
||||
handler: Final = FakeHandler([httpx.ReadTimeout("timed out")])
|
||||
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
|
||||
with caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"):
|
||||
await _run(guardrail, _mcp_data())
|
||||
|
|
@ -573,7 +647,7 @@ class TestFailOpenOptIn:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_opted_in_fail_open_still_blocks_a_policy_block(self):
|
||||
handler: Final = FakeHandler([_token_response(), _block_response()])
|
||||
handler: Final = FakeHandler([_block_response()])
|
||||
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _run(guardrail, _mcp_data())
|
||||
|
|
@ -584,10 +658,7 @@ class TestUnreachableFallback:
|
|||
@pytest.mark.asyncio
|
||||
async def test_evaluate_litellm_timeout_fail_closed(self):
|
||||
handler: Final = FakeHandler(
|
||||
[
|
||||
_token_response(),
|
||||
LitellmTimeout(message="Connection timed out", model="default-model-name", llm_provider="httpx"),
|
||||
]
|
||||
[LitellmTimeout(message="Connection timed out", model="default-model-name", llm_provider="httpx")]
|
||||
)
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
|
|
@ -596,7 +667,7 @@ class TestUnreachableFallback:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_evaluate_timeout_fail_closed(self):
|
||||
handler: Final = FakeHandler([_token_response(), httpx.ReadTimeout("timed out")])
|
||||
handler: Final = FakeHandler([httpx.ReadTimeout("timed out")])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _run(guardrail, _mcp_data())
|
||||
|
|
@ -605,7 +676,7 @@ class TestUnreachableFallback:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_evaluate_timeout_fail_open(self):
|
||||
handler: Final = FakeHandler([_token_response(), httpx.ReadTimeout("timed out")])
|
||||
handler: Final = FakeHandler([httpx.ReadTimeout("timed out")])
|
||||
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
|
||||
data: Final = _mcp_data()
|
||||
result: Final = await _run(guardrail, data)
|
||||
|
|
@ -616,7 +687,7 @@ class TestUnreachableFallback:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_evaluate_5xx_fail_closed(self):
|
||||
handler: Final = FakeHandler([_token_response(), _response(502, text="bad gateway")])
|
||||
handler: Final = FakeHandler([_response(502, text="bad gateway")])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _run(guardrail, _mcp_data())
|
||||
|
|
@ -634,11 +705,13 @@ class TestUnreachableFallback:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_jwt_bearer_token_fail_closed(self):
|
||||
exchanger: Final = StubTokenExchanger()
|
||||
handler: Final = FakeHandler([])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _run(guardrail, _mcp_data(incoming_bearer_token="sk-litellm-virtual-key"))
|
||||
assert exc_info.value.status_code == 401
|
||||
assert exchanger.calls == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_bearer_token_blocks_even_fail_open(self):
|
||||
|
|
@ -655,18 +728,19 @@ class TestUnreachableFallback:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_obo_rejected_blocks_even_fail_open(self):
|
||||
handler: Final = FakeHandler(
|
||||
[_response(400, {"error": "invalid_grant", "error_description": "AADSTS50013: bad assertion"})]
|
||||
exchanger: Final = StubTokenExchanger(
|
||||
[Error(CredError.of_unauthorized("IdP rejected the subject token (HTTP 400)", claims="invalid_grant"))]
|
||||
)
|
||||
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
|
||||
handler: Final = FakeHandler([])
|
||||
guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open")
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _run(guardrail, _mcp_data())
|
||||
assert exc_info.value.status_code == 401
|
||||
assert "invalid_grant" in exc_info.value.detail["message"]
|
||||
assert "On-Behalf-Of token exchange was rejected" in exc_info.value.detail["message"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_evaluate_4xx_blocks_even_fail_open(self):
|
||||
handler: Final = FakeHandler([_token_response(), _response(403, text="obo token lacks the scope")])
|
||||
handler: Final = FakeHandler([_response(403, text="obo token lacks the scope")])
|
||||
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
|
||||
data: Final = _mcp_data()
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
|
|
@ -680,24 +754,24 @@ class TestUnreachableFallback:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_obo_rejected_fail_closed(self):
|
||||
handler: Final = FakeHandler(
|
||||
[_response(400, {"error": "invalid_grant", "error_description": "AADSTS50013: bad assertion"})]
|
||||
exchanger: Final = StubTokenExchanger(
|
||||
[Error(CredError.of_unauthorized("IdP rejected the subject token (HTTP 400)"))]
|
||||
)
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
handler: Final = FakeHandler([])
|
||||
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _run(guardrail, _mcp_data())
|
||||
assert exc_info.value.status_code == 401
|
||||
assert "invalid_grant" in exc_info.value.detail["message"]
|
||||
assert "On-Behalf-Of token exchange was rejected" in exc_info.value.detail["message"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"error_code", ["invalid_client", "unauthorized_client", "invalid_scope", "invalid_resource"]
|
||||
)
|
||||
async def test_gateway_credential_rejection_is_unavailable_not_a_caller_401(self, error_code: str):
|
||||
handler: Final = FakeHandler(
|
||||
[_response(401, {"error": error_code, "error_description": "AADSTS7000215: invalid client secret"})]
|
||||
)
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
exchanger: Final = StubTokenExchanger([Error(CredError.of_misconfigured(error_code))])
|
||||
handler: Final = FakeHandler([])
|
||||
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
|
||||
data: Final = _mcp_data()
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _run(guardrail, data)
|
||||
|
|
@ -710,23 +784,12 @@ class TestUnreachableFallback:
|
|||
assert "client_secret" in info["guardrail_response"]["reason"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("aadsts_code", [5002710, 5002723], ids=["malformed-header", "no-kid"])
|
||||
async def test_malformed_assertion_reported_as_invalid_client_is_a_caller_401(self, aadsts_code: int):
|
||||
"""Entra answers ``invalid_client`` for a forged or garbled assertion (AADSTS50027xx) exactly as for a
|
||||
bad gateway secret; the sub-code is what says the caller, not the gateway, has to fix it."""
|
||||
handler: Final = FakeHandler(
|
||||
[
|
||||
_response(
|
||||
401,
|
||||
{
|
||||
"error": "invalid_client",
|
||||
"error_description": f"AADSTS{aadsts_code}: Invalid JWT token.",
|
||||
"error_codes": [aadsts_code],
|
||||
},
|
||||
)
|
||||
]
|
||||
async def test_caller_rejection_reason_does_not_blame_the_gateway_credentials(self):
|
||||
exchanger: Final = StubTokenExchanger(
|
||||
[Error(CredError.of_unauthorized("IdP rejected the subject token (HTTP 401)"))]
|
||||
)
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
handler: Final = FakeHandler([])
|
||||
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
|
||||
data: Final = _mcp_data()
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _run(guardrail, data)
|
||||
|
|
@ -735,10 +798,9 @@ class TestUnreachableFallback:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gateway_credential_rejection_follows_fail_open(self):
|
||||
handler: Final = FakeHandler(
|
||||
[_response(401, {"error": "invalid_client", "error_description": "AADSTS7000215: invalid client secret"})]
|
||||
)
|
||||
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
|
||||
exchanger: Final = StubTokenExchanger([Error(CredError.of_misconfigured("invalid_client"))])
|
||||
handler: Final = FakeHandler([])
|
||||
guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open")
|
||||
data: Final = _mcp_data()
|
||||
result: Final = await _run(guardrail, data)
|
||||
assert result is data
|
||||
|
|
@ -747,10 +809,43 @@ class TestUnreachableFallback:
|
|||
assert info["guardrail_response"]["verdict"] == "Unscanned"
|
||||
assert "invalid_client" in info["guardrail_response"]["reason"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchange_upstream_unavailable_is_unavailable_with_the_summary_not_a_caller_401(self):
|
||||
exchanger: Final = StubTokenExchanger(
|
||||
[Error(CredError.of_upstream_unavailable("token exchange did not return a usable access token"))]
|
||||
)
|
||||
handler: Final = FakeHandler([])
|
||||
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
|
||||
data: Final = _mcp_data()
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _run(guardrail, data)
|
||||
assert exc_info.value.status_code == 503
|
||||
info: Final = _guardrail_info(data)
|
||||
assert info["guardrail_status"] == "guardrail_failed_to_respond"
|
||||
assert info["guardrail_response"]["verdict"] == "Unavailable"
|
||||
assert (
|
||||
info["guardrail_response"]["reason"]
|
||||
== "the Entra token exchange failed (upstream unavailable: token exchange did not return a usable access token)"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchange_upstream_unavailable_follows_fail_open(self):
|
||||
exchanger: Final = StubTokenExchanger([Error(CredError.of_upstream_unavailable("Entra throttled the exchange"))])
|
||||
handler: Final = FakeHandler([])
|
||||
guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open")
|
||||
data: Final = _mcp_data()
|
||||
result: Final = await _run(guardrail, data)
|
||||
assert result is data
|
||||
info: Final = _guardrail_info(data)
|
||||
assert info["guardrail_status"] == "guardrail_failed_to_respond"
|
||||
assert info["guardrail_response"]["verdict"] == "Unscanned"
|
||||
assert "Entra throttled the exchange" in info["guardrail_response"]["reason"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_obo_endpoint_5xx_fail_open(self):
|
||||
handler: Final = FakeHandler([_response(503, text="entra down")])
|
||||
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
|
||||
exchanger: Final = StubTokenExchanger([Error(CredError.of_upstream_unavailable("entra down"))])
|
||||
handler: Final = FakeHandler([])
|
||||
guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open")
|
||||
data: Final = _mcp_data()
|
||||
result: Final = await _run(guardrail, data)
|
||||
assert result is data
|
||||
|
|
@ -762,47 +857,33 @@ class TestUnreachableFallback:
|
|||
class TestOboTokenCache:
|
||||
@pytest.mark.asyncio
|
||||
async def test_same_assertion_reuses_token(self):
|
||||
handler: Final = FakeHandler([_token_response(), _allow_response(), _allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
exchanger: Final = StubTokenExchanger(_obo_ok())
|
||||
handler: Final = FakeHandler([_allow_response(), _allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
|
||||
await _run(guardrail, _mcp_data())
|
||||
await _run(guardrail, _mcp_data())
|
||||
token_calls: Final = [c for c in handler.calls if c.url == TOKEN_URL]
|
||||
assert len(token_calls) == 1
|
||||
assert len(exchanger.calls) == 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_different_assertions_get_distinct_tokens(self):
|
||||
other_assertion: Final = "eyJhbGciOi.eyJvdGhlciI.b3RoZXJzaWc"
|
||||
handler: Final = FakeHandler(
|
||||
[
|
||||
_token_response(access_token="token-a"),
|
||||
_allow_response(),
|
||||
_token_response(access_token="token-b"),
|
||||
_allow_response(),
|
||||
]
|
||||
)
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
exchanger: Final = StubTokenExchanger([_ok_exchange("token-a"), _ok_exchange("token-b")])
|
||||
handler: Final = FakeHandler([_allow_response(), _allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
|
||||
await _run(guardrail, _mcp_data())
|
||||
await _run(guardrail, _mcp_data(incoming_bearer_token=other_assertion))
|
||||
token_calls: Final = [c for c in handler.calls if c.url == TOKEN_URL]
|
||||
assert len(token_calls) == 2
|
||||
assert handler.calls[3].headers["Authorization"] == "Bearer token-b"
|
||||
assert len(exchanger.calls) == 2
|
||||
assert handler.calls[1].headers["Authorization"] == "Bearer token-b"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_expired_token_refreshed(self):
|
||||
handler: Final = FakeHandler(
|
||||
[
|
||||
_token_response(access_token="short-lived", expires_in=1),
|
||||
_allow_response(),
|
||||
_token_response(access_token="fresh"),
|
||||
_allow_response(),
|
||||
]
|
||||
)
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
exchanger: Final = StubTokenExchanger([_ok_exchange("short-lived", -1), _ok_exchange("fresh")])
|
||||
handler: Final = FakeHandler([_allow_response(), _allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
|
||||
await _run(guardrail, _mcp_data())
|
||||
await _run(guardrail, _mcp_data())
|
||||
token_calls: Final = [c for c in handler.calls if c.url == TOKEN_URL]
|
||||
assert len(token_calls) == 2
|
||||
assert handler.calls[3].headers["Authorization"] == "Bearer fresh"
|
||||
assert len(exchanger.calls) == 2
|
||||
assert handler.calls[1].headers["Authorization"] == "Bearer fresh"
|
||||
|
||||
|
||||
class TestEarlyPhasePassthrough:
|
||||
|
|
@ -835,34 +916,27 @@ class TestRegistryDiscovery:
|
|||
|
||||
class TestMalformedResponses:
|
||||
@pytest.mark.asyncio
|
||||
async def test_obo_html_body_fail_open(self):
|
||||
handler: Final = FakeHandler([_response(200, text="<html>blocked by egress proxy</html>")])
|
||||
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
|
||||
async def test_obo_upstream_unavailable_fail_open(self):
|
||||
exchanger: Final = StubTokenExchanger([Error(CredError.of_upstream_unavailable("entra returned no token"))])
|
||||
handler: Final = FakeHandler([])
|
||||
guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open")
|
||||
data: Final = _mcp_data()
|
||||
result: Final = await _run(guardrail, data)
|
||||
assert result is data
|
||||
assert _guardrail_info(data)["guardrail_response"]["verdict"] == "Unscanned"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_obo_html_body_fail_closed(self):
|
||||
handler: Final = FakeHandler([_response(200, text="<html>outage</html>")])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _run(guardrail, _mcp_data())
|
||||
assert exc_info.value.status_code == 503
|
||||
assert "non-JSON" in exc_info.value.detail["message"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_obo_non_object_json_fail_closed(self):
|
||||
handler: Final = FakeHandler([_response(200, ["not", "a", "dict"])])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
async def test_obo_upstream_unavailable_fail_closed(self):
|
||||
exchanger: Final = StubTokenExchanger([Error(CredError.of_upstream_unavailable("entra returned no token"))])
|
||||
handler: Final = FakeHandler([])
|
||||
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _run(guardrail, _mcp_data())
|
||||
assert exc_info.value.status_code == 503
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_evaluate_html_body_fail_open(self):
|
||||
handler: Final = FakeHandler([_token_response(), _response(200, text="<html>waf page</html>")])
|
||||
handler: Final = FakeHandler([_response(200, text="<html>waf page</html>")])
|
||||
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
|
||||
data: Final = _mcp_data()
|
||||
result: Final = await _run(guardrail, data)
|
||||
|
|
@ -871,7 +945,7 @@ class TestMalformedResponses:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_evaluate_html_body_fail_closed(self):
|
||||
handler: Final = FakeHandler([_token_response(), _response(200, text="<html>waf page</html>")])
|
||||
handler: Final = FakeHandler([_response(200, text="<html>waf page</html>")])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _run(guardrail, _mcp_data())
|
||||
|
|
@ -879,7 +953,7 @@ class TestMalformedResponses:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_evaluate_non_object_json_fail_closed(self):
|
||||
handler: Final = FakeHandler([_token_response(), _response(200, "allowed")])
|
||||
handler: Final = FakeHandler([_response(200, "allowed")])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _run(guardrail, _mcp_data())
|
||||
|
|
@ -892,7 +966,7 @@ class TestMalformedResponses:
|
|||
ids=["missing", "null", "string-true", "int-one", "string-false"],
|
||||
)
|
||||
async def test_evaluate_non_boolean_allowed_fail_closed(self, verdict: dict):
|
||||
handler: Final = FakeHandler([_token_response(), _response(200, verdict)])
|
||||
handler: Final = FakeHandler([_response(200, verdict)])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
data: Final = _mcp_data()
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
|
|
@ -908,7 +982,7 @@ class TestMalformedResponses:
|
|||
ids=["missing", "null", "string-true", "int-one", "string-false"],
|
||||
)
|
||||
async def test_evaluate_non_boolean_allowed_fail_open(self, verdict: dict):
|
||||
handler: Final = FakeHandler([_token_response(), _response(200, verdict)])
|
||||
handler: Final = FakeHandler([_response(200, verdict)])
|
||||
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
|
||||
data: Final = _mcp_data()
|
||||
result: Final = await _run(guardrail, data)
|
||||
|
|
@ -917,21 +991,21 @@ class TestMalformedResponses:
|
|||
assert _guardrail_info(data)["guardrail_status"] == "guardrail_failed_to_respond"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bad_expires_in_still_allows(self):
|
||||
handler: Final = FakeHandler(
|
||||
[_response(200, {"access_token": "tok-1", "expires_in": "soon"}), _allow_response()]
|
||||
)
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
async def test_obo_token_without_expiry_still_allows(self):
|
||||
exchanger: Final = StubTokenExchanger([Ok(OAuthToken(access_token="tok-1", expires_at=None))])
|
||||
handler: Final = FakeHandler([_allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
|
||||
data: Final = _mcp_data()
|
||||
result: Final = await _run(guardrail, data)
|
||||
assert result is data
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_obo_litellm_timeout_fail_open(self):
|
||||
handler: Final = FakeHandler(
|
||||
exchanger: Final = StubTokenExchanger(
|
||||
[LitellmTimeout(message="Connection timed out", model="default-model-name", llm_provider="httpx")]
|
||||
)
|
||||
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
|
||||
handler: Final = FakeHandler([])
|
||||
guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open")
|
||||
data: Final = _mcp_data()
|
||||
result: Final = await _run(guardrail, data)
|
||||
assert result is data
|
||||
|
|
@ -940,28 +1014,18 @@ class TestMalformedResponses:
|
|||
|
||||
class TestDeltaHardening:
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_string_access_token_fail_closed(self):
|
||||
handler: Final = FakeHandler([_response(200, {"access_token": None, "expires_in": 3599})])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _run(guardrail, _mcp_data())
|
||||
assert exc_info.value.status_code == 503
|
||||
assert "access_token" in exc_info.value.detail["message"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_numeric_string_expires_in_honored(self):
|
||||
handler: Final = FakeHandler(
|
||||
[_response(200, {"access_token": "tok-9", "expires_in": "120"}), _allow_response()]
|
||||
)
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
async def test_unexpired_exchange_result_is_reused(self):
|
||||
exchanger: Final = StubTokenExchanger([_ok_exchange("tok-9", 120)])
|
||||
handler: Final = FakeHandler([_allow_response(), _allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
|
||||
await _run(guardrail, _mcp_data())
|
||||
entries: Final = list(guardrail._obo_token_cache.values())
|
||||
assert len(entries) == 1
|
||||
assert entries[0][1] - time.time() < 200
|
||||
await _run(guardrail, _mcp_data())
|
||||
assert len(exchanger.calls) == 1
|
||||
assert handler.calls[1].headers["Authorization"] == "Bearer tok-9"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_evaluate_400_records_intervention(self):
|
||||
handler: Final = FakeHandler([_token_response(), _response(400, text="bad request shape")])
|
||||
handler: Final = FakeHandler([_response(400, text="bad request shape")])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
data: Final = _mcp_data()
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
|
|
@ -975,25 +1039,19 @@ class TestDeltaHardening:
|
|||
class TestVeriaHardening:
|
||||
@pytest.mark.asyncio
|
||||
async def test_evaluate_401_evicts_cached_obo_token(self):
|
||||
handler: Final = FakeHandler(
|
||||
[
|
||||
_token_response(),
|
||||
_response(401, text="token expired"),
|
||||
_token_response(access_token="tok-2"),
|
||||
_allow_response(),
|
||||
]
|
||||
)
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
exchanger: Final = StubTokenExchanger(_obo_ok() + _obo_ok("tok-2"))
|
||||
handler: Final = FakeHandler([_response(401, text="token expired"), _allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
|
||||
with pytest.raises(HTTPException):
|
||||
await _run(guardrail, _mcp_data())
|
||||
result: Final = await _run(guardrail, _mcp_data())
|
||||
assert result is not None
|
||||
token_calls: Final = [c for c in handler.calls if c.url == TOKEN_URL]
|
||||
assert len(token_calls) == 2
|
||||
assert exchanger.invalidations == [FAKE_ASSERTION]
|
||||
assert len(exchanger.calls) == 2
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_evaluate_429_blocks_even_fail_open_as_throttled(self):
|
||||
handler: Final = FakeHandler([_token_response(), _response(429, text="slow down")])
|
||||
handler: Final = FakeHandler([_response(429, text="slow down")])
|
||||
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
|
||||
data: Final = _mcp_data()
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
|
|
@ -1006,7 +1064,7 @@ class TestVeriaHardening:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_evaluate_500_is_unavailable(self):
|
||||
handler: Final = FakeHandler([_token_response(), _response(500, text="oops")])
|
||||
handler: Final = FakeHandler([_response(500, text="oops")])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _run(guardrail, _mcp_data())
|
||||
|
|
@ -1014,51 +1072,23 @@ class TestVeriaHardening:
|
|||
assert "500" in exc_info.value.detail["message"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_endpoint_429_blocks_even_fail_open_as_throttled(self):
|
||||
handler: Final = FakeHandler(
|
||||
[_response(429, {"error": "temporarily_throttled", "error_description": "AADSTS90056"})]
|
||||
async def test_token_endpoint_unauthorized_is_a_caller_401_even_fail_open(self):
|
||||
exchanger: Final = StubTokenExchanger(
|
||||
[Error(CredError.of_unauthorized("IdP rejected the subject token (HTTP 429)"))]
|
||||
)
|
||||
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
|
||||
handler: Final = FakeHandler([])
|
||||
guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open")
|
||||
data: Final = _mcp_data()
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _run(guardrail, data)
|
||||
assert exc_info.value.status_code == 503
|
||||
assert "429" in exc_info.value.detail["message"]
|
||||
assert exc_info.value.status_code == 401
|
||||
info: Final = _guardrail_info(data)
|
||||
assert info["guardrail_status"] == "guardrail_failed_to_respond"
|
||||
assert info["guardrail_response"]["verdict"] == "Throttled"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_endpoint_408_non_json_blocks_as_throttled(self):
|
||||
handler: Final = FakeHandler([_response(408, text="Request Timeout")])
|
||||
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
|
||||
data: Final = _mcp_data()
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _run(guardrail, data)
|
||||
assert exc_info.value.status_code == 503
|
||||
assert _guardrail_info(data)["guardrail_response"]["verdict"] == "Throttled"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_endpoint_4xx_html_stays_infra_fail_open(self):
|
||||
handler: Final = FakeHandler([_response(403, text="<html>waf block page</html>")])
|
||||
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
|
||||
data: Final = _mcp_data()
|
||||
result: Final = await _run(guardrail, data)
|
||||
assert result is data
|
||||
assert _guardrail_info(data)["guardrail_response"]["verdict"] == "Unscanned"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_entra_200_missing_access_token_is_malformed(self):
|
||||
handler: Final = FakeHandler([_response(200, {"token_type": "Bearer"})])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _run(guardrail, _mcp_data())
|
||||
assert exc_info.value.status_code == 503
|
||||
assert "access_token" in exc_info.value.detail["message"]
|
||||
assert info["guardrail_status"] == "guardrail_intervened"
|
||||
assert info["guardrail_response"]["verdict"] == "Rejected"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_evaluate_5xx_fail_open_allows_unscanned_once(self):
|
||||
handler: Final = FakeHandler([_token_response(), _response(502, text='{"error": "bad gateway"}')])
|
||||
handler: Final = FakeHandler([_response(502, text='{"error": "bad gateway"}')])
|
||||
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
|
||||
data: Final = _mcp_data()
|
||||
result: Final = await _run(guardrail, data)
|
||||
|
|
@ -1099,7 +1129,7 @@ class TestFinalArgumentsEvaluated:
|
|||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("agent_365_first", [True, False], ids=["agent_365_then_masker", "masker_then_agent_365"])
|
||||
async def test_agent_365_receives_the_arguments_sent_upstream(self, agent_365_first: bool):
|
||||
handler: Final = FakeHandler([_token_response(), _allow_response()])
|
||||
handler: Final = FakeHandler([_allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
masker: Final = _ArgumentMasker("arg-rewrite")
|
||||
registered: Final = (guardrail, masker) if agent_365_first else (masker, guardrail)
|
||||
|
|
@ -1116,4 +1146,61 @@ class TestFinalArgumentsEvaluated:
|
|||
litellm.callbacks, callback, require_self=False
|
||||
)
|
||||
assert result["modified_arguments"] == {"turn": "please [REWRITE_ME_REDACTED] now"}
|
||||
assert handler.calls[1].json["arguments"] == {"turn": "please [REWRITE_ME_REDACTED] now"}
|
||||
assert handler.calls[0].json["arguments"] == {"turn": "please [REWRITE_ME_REDACTED] now"}
|
||||
|
||||
|
||||
class TestCallerSignIn:
|
||||
"""The guardrail is a CallerSignInProvider: it tells the MCP connect path which issuer and scope
|
||||
the caller must sign in for before the first tool call, and only for servers that keep the
|
||||
caller's Authorization with the gateway."""
|
||||
|
||||
def test_gated_server_advertises_entra_issuer_and_gateway_scope(self):
|
||||
guardrail: Final = _make_guardrail(FakeHandler([]))
|
||||
sign_in: Final = guardrail.caller_sign_in(_server(), None)
|
||||
assert sign_in == CallerSignIn(
|
||||
issuers=("https://login.microsoftonline.com/tenant-abc/v2.0",),
|
||||
scopes=("api://client-xyz/access_as_user",),
|
||||
)
|
||||
|
||||
def test_default_off_guardrail_does_not_gate(self):
|
||||
guardrail: Final = _make_guardrail(FakeHandler([]), default_on=False)
|
||||
assert guardrail.caller_sign_in(_server(), None) is None
|
||||
|
||||
def test_server_that_fills_authorization_itself_is_not_gated(self):
|
||||
guardrail: Final = _make_guardrail(FakeHandler([]))
|
||||
assert guardrail.caller_sign_in(_server(auth_type=MCPAuth.oauth2), None) is None
|
||||
assert guardrail.caller_sign_in(_server(extra_headers=["authorization"]), None) is None
|
||||
|
||||
def test_opted_out_key_does_not_gate(self):
|
||||
class _OptedOut(Agent365Guardrail):
|
||||
def should_run_guardrail(self, data, event_type) -> bool:
|
||||
return False
|
||||
|
||||
guardrail: Final = _OptedOut(
|
||||
guardrail_name="a365-off",
|
||||
tenant_id="tenant-abc",
|
||||
client_id="client-xyz",
|
||||
client_secret="secret-123",
|
||||
token_exchanger=StubTokenExchanger(),
|
||||
default_on=True,
|
||||
)
|
||||
assert guardrail.caller_sign_in(_server(), UserAPIKeyAuth(api_key="k", user_id="u-1")) is None
|
||||
assert guardrail.caller_sign_in(_server(), None) is not None
|
||||
|
||||
def test_obo_server_with_provider_advertises_both_issuers_and_scopes(self, monkeypatch):
|
||||
monkeypatch.setenv("JWT_ISSUER", "https://jwt-idp.test")
|
||||
guardrail: Final = _make_guardrail(FakeHandler([]))
|
||||
litellm.logging_callback_manager.add_litellm_callback(guardrail)
|
||||
try:
|
||||
server: Final = _server(auth_type=MCPAuth.oauth2_token_exchange, scopes=["read"])
|
||||
sign_in: Final = caller_sign_in_for(server, None)
|
||||
finally:
|
||||
litellm.logging_callback_manager.remove_callback_from_list_by_object(
|
||||
litellm.callbacks, guardrail, require_self=False
|
||||
)
|
||||
assert sign_in is not None
|
||||
assert sign_in.issuers == (
|
||||
"https://jwt-idp.test",
|
||||
"https://login.microsoftonline.com/tenant-abc/v2.0",
|
||||
)
|
||||
assert sign_in.scopes == ("read", "api://client-xyz/access_as_user")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue