fix(mcp): answer Agent 365 exchange faults inside the tools/call envelope with the merge base's reasons

The connect-time sign-in preflight raised its fail-closed 503 on every method of a single-server route, so a post-connect tools/call whose Entra exchange hit a gateway fault got a bare 503 with no JSON-RPC envelope, no hook run and no guardrail Logs row. The gate now reads the body before the preflight and only raises the 503 while the request is an initialize (or the SSE GET); any other request reaches the tool-call hook, which answers 200 isError with the verdict and writes the failure row as the merge base did

The guardrail posts to the Entra token endpoint through its own handler again, bounded by request_timeout, and classifies the answer itself through the shared OAuth error reader, so the gateway-fault reason keeps the OAuth code (invalid_client, invalid_scope, ...), a malformed caller assertion (AADSTS 50027xx) stays a 401 caller fault reading rejected (invalid_client), and transport, HTTP 500, non-JSON and missing access_token answers read as the merge base's distinct reasons instead of one collapsed sentence. build_token_exchanger takes the HTTP post as a dependency in place of request_timeout

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-10-03 05:34:55 +00:00
parent 3e7efe95d6
commit 7603412660
7 changed files with 406 additions and 164 deletions

View file

@ -162,9 +162,12 @@ async def preflight_caller_sign_in(
*,
root_path: str,
resource_metadata: str | None,
connecting: bool,
) -> None:
"""Run every provider's connect-time check against the subject token, so a bearer the IdP will
reject surfaces as a challenge here rather than a JSON-RPC error at the first tool call."""
reject surfaces as a challenge here rather than a JSON-RPC error at the first tool call. A fail-closed
provider outage is the connect's 503 only while ``connecting``; on an open session the tool-call hook
answers it inside the JSON-RPC envelope, with its guardrail Logs row."""
from fastapi import HTTPException # noqa: PLC0415 # lazy: fastapi import stays off the cold path
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph
@ -181,9 +184,11 @@ async def preflight_caller_sign_in(
raise_token_exchange_challenge(
server, root_path=root_path, claims=claims, resource_metadata=resource_metadata
)
case Unavailable(detail=detail, fail_open=True):
case Unavailable(fail_open=True):
continue
case Unavailable(detail=detail, fail_open=False):
case Unavailable(detail=detail, fail_open=False) if connecting:
raise HTTPException(status_code=503, detail=detail)
case Unavailable():
continue
case _ as verdict:
assert_never(verdict)

View file

@ -25,6 +25,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_sto
InMemoryTokenCacheBackend,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import (
ExchangeHttpPost,
OboTokenExchanger,
SubjectTokenRejected,
TokenExchangeClientError,
@ -39,7 +40,7 @@ _INVALID_ASSERTION_AADSTS_PREFIX: Final = "50027"
@dataclass(frozen=True, slots=True)
class _OAuthErrorBody:
class OAuthErrorBody:
error: str | None
claims: str | None
error_codes: tuple[str, ...]
@ -53,7 +54,7 @@ class _OAuthErrorBody:
return self.error
def _oauth_error_fields(response: httpx.Response) -> _OAuthErrorBody:
def oauth_error_fields(response: httpx.Response) -> OAuthErrorBody:
"""Read the RFC 6749 5.2 ``error`` code, the IdP's step-up ``claims`` blob and Entra's
``error_codes`` sub-codes from a token-endpoint error body, None or empty for whatever is absent.
@ -65,13 +66,13 @@ def _oauth_error_fields(response: httpx.Response) -> _OAuthErrorBody:
try:
body: Final[object] = response.json()
except Exception: # noqa: BLE001
return _OAuthErrorBody(error=None, claims=None, error_codes=())
return OAuthErrorBody(error=None, claims=None, error_codes=())
if not isinstance(body, dict):
return _OAuthErrorBody(error=None, claims=None, error_codes=())
return OAuthErrorBody(error=None, claims=None, error_codes=())
code: Final = body.get("error")
claims: Final = body.get("claims")
raw_codes: Final = body.get("error_codes")
return _OAuthErrorBody(
return OAuthErrorBody(
error=code if isinstance(code, str) else None,
claims=claims if isinstance(claims, str) and claims else None,
error_codes=tuple(str(c) for c in raw_codes if isinstance(c, (int, str)))
@ -81,7 +82,7 @@ def _oauth_error_fields(response: httpx.Response) -> _OAuthErrorBody:
async def _post_exchange_endpoint(
url: str, form: dict[str, str], client_auth_headers: dict[str, str], *, timeout: float | None = None
url: str, form: dict[str, str], client_auth_headers: dict[str, str]
) -> dict[str, object] | None:
from litellm.llms.custom_httpx.http_handler import ( # noqa: PLC0415
get_async_httpx_client, # pyright: ignore
@ -96,7 +97,7 @@ async def _post_exchange_endpoint(
try:
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) # pyright: ignore
response: Final = await client.post( # pyright: ignore[reportUnknownMemberType] # untyped handler
url, headers=headers, data=form, timeout=timeout
url, headers=headers, data=form
)
response.raise_for_status() # pyright: ignore
parsed: Final[object] = response.json() # pyright: ignore
@ -108,7 +109,7 @@ async def _post_exchange_endpoint(
verbose_logger.warning("MCP token exchange throttled or timed out (HTTP %d)", status_code)
return None
if 400 <= status_code < 500:
oauth_error: Final = _oauth_error_fields(status_err.response)
oauth_error: Final = oauth_error_fields(status_err.response)
gateway_fault: Final = oauth_error.gateway_fault
if gateway_fault is not None:
verbose_logger.warning(
@ -135,10 +136,7 @@ async def _post_exchange_endpoint(
return parsed # pyright: ignore
def build_token_exchanger(*, request_timeout: float | None = None) -> OboTokenExchanger:
async def post(url: str, form: dict[str, str], client_auth_headers: dict[str, str]) -> dict[str, object] | None:
return await _post_exchange_endpoint(url, form, client_auth_headers, timeout=request_timeout)
def build_token_exchanger(*, post: ExchangeHttpPost = _post_exchange_endpoint) -> OboTokenExchanger:
return OboTokenExchanger(
post,
cache=InMemoryTokenCacheBackend(max_size=MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE),

View file

@ -1608,6 +1608,7 @@ if MCP_AVAILABLE:
mcp_server_auth_headers: dict[str, dict[str, str]] | None,
user_api_key_auth: UserAPIKeyAuth | None,
client_ip: str | None,
connecting: bool,
allowed_server_ids: set[str] | None = None,
raw_headers: Mapping[str, str] | None = None,
) -> None:
@ -1779,6 +1780,7 @@ if MCP_AVAILABLE:
subject_token,
root_path=get_request_root_path(),
resource_metadata=resource_metadata,
connecting=connecting,
)
# Exchange-backed modes (token_exchange's OBO mint, id_jag's stored-assertion mint): run
@ -2070,6 +2072,22 @@ if MCP_AVAILABLE:
user_api_key_auth = await _apply_toolset_scope(user_api_key_auth, active_toolset_id)
toolset_allowed_server_ids = await _toolset_server_ids(active_toolset_id)
consumed_messages, body = (
await _read_request_body_for_routing(receive) if scope.get("method") == "POST" else ([], b"")
)
is_initialize: Final = _is_initialize_request(body)
# Replay body messages if we consumed them for peeking
original_receive: Final = receive
if consumed_messages:
async def wrapped_receive():
if consumed_messages:
return consumed_messages.pop(0)
return await original_receive()
receive = wrapped_receive
# https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response
# Must run after toolset scoping so the challenge set is derived
# from the fully-authorized server set: a passthrough server that
@ -2082,6 +2100,7 @@ if MCP_AVAILABLE:
mcp_server_auth_headers=mcp_server_auth_headers,
user_api_key_auth=user_api_key_auth,
client_ip=_client_ip,
connecting=is_initialize,
allowed_server_ids=toolset_allowed_server_ids,
raw_headers=raw_headers,
)
@ -2117,8 +2136,6 @@ if MCP_AVAILABLE:
# - No session ID + initialize → stateful (so client gets mcp-session-id)
# - No session ID + other → stateless (curl, Inspector, Notion)
session_id = _get_session_id_from_scope(scope)
is_initialize = False
consumed_messages: list[Message] = []
# Owner-binding: a live stateful session may only be driven by the
# caller that created it. Reject mismatches with 403 so a leaked
@ -2126,8 +2143,7 @@ if MCP_AVAILABLE:
#
# Run before ``_handle_stale_mcp_session`` so a non-owner cannot
# force-clean another caller's residual tracking entries via a
# stale DELETE, and before peeking the request body so the 403
# response sees a pristine ``receive`` channel.
# stale DELETE.
if session_id:
expected_owner: Final = _stateful_session_owners.get(session_id)
request_owner = _owner_fingerprint_for(user_api_key_auth, oauth2_headers, _client_ip)
@ -2156,11 +2172,6 @@ if MCP_AVAILABLE:
return
session_id = _get_session_id_from_scope(scope)
body = b""
if scope.get("method") == "POST":
consumed_messages, body = await _read_request_body_for_routing(receive)
is_initialize = _is_initialize_request(body)
use_stateful: Final = bool(session_id or is_initialize)
target_manager: Final = session_manager_stateful if use_stateful else session_manager_stateless
@ -2189,17 +2200,6 @@ if MCP_AVAILABLE:
await too_many_response(scope, receive, send)
return
# Replay body messages if we consumed them for peeking
original_receive: Final = receive
if consumed_messages:
async def wrapped_receive():
if consumed_messages:
return consumed_messages.pop(0)
return await original_receive()
receive = wrapped_receive
# Serialize requests on the same stateful session so concurrent
# callers don't clobber each other's auth context mid-flight.
#
@ -2429,6 +2429,7 @@ if MCP_AVAILABLE:
mcp_server_auth_headers=mcp_server_auth_headers,
user_api_key_auth=user_api_key_auth,
client_ip=_sse_client_ip,
connecting=scope["method"] == "GET",
allowed_server_ids=toolset_allowed_server_ids,
raw_headers=raw_headers,
)

View file

@ -42,8 +42,12 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_sto
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok, Result
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_provider import (
build_token_exchanger,
oauth_error_fields,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import (
SubjectTokenRejected,
TokenExchanger,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import TokenExchanger
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
CredError,
ServerSpec,
@ -72,12 +76,12 @@ MCP_SESSION_ID_HEADER: Final = "mcp-session-id"
DEFENDER_STATUS_EVALUATED: Final = "Evaluated"
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])
_JSON_OBJECT_ADAPTER: Final = TypeAdapter(dict[str, object])
def _parse_tool_input_schema(raw: object) -> Mapping[str, object] | None:
try:
return _TOOL_INPUT_SCHEMA_ADAPTER.validate_python(raw)
return _JSON_OBJECT_ADAPTER.validate_python(raw)
except ValidationError:
return None
@ -130,12 +134,32 @@ class _BlockedDetail(TypedDict):
correlation_id: ReadOnly[str | None]
class Agent365TokenExchangeError(Exception):
"""Entra refused the gateway's own client credentials, scope or resource; the caller cannot fix that by
signing in again, so it is the gateway's outage, never a 401."""
def __init__(self, error_code: str) -> None:
super().__init__(error_code)
self.error_code = error_code
class Agent365MalformedResponseError(Exception):
pass
class Agent365ThrottledError(Exception):
def __init__(self, status_code: int) -> None:
super().__init__(f"HTTP {status_code}")
self.status_code = status_code
def _gateway_fault_reason(error_code: str) -> str:
return (
f"Entra rejected the gateway's own Agent 365 credentials ({error_code}); "
"check the guardrail's client_id and client_secret"
)
class Agent365Guardrail(CustomGuardrail):
"""Pre-MCP-call guardrail enforcing Microsoft Agent 365 tool-evaluation verdicts.
@ -188,7 +212,7 @@ class Agent365Guardrail(CustomGuardrail):
self._token_exchanger: Final = (
token_exchanger
if token_exchanger is not None
else build_token_exchanger(request_timeout=self.request_timeout)
else build_token_exchanger(post=self._post_entra_token_endpoint)
)
verbose_proxy_logger.info("Initialized Microsoft Agent 365 guardrail: %s", guardrail_name)
@ -230,12 +254,25 @@ class Agent365Guardrail(CustomGuardrail):
try:
exchange_result: Final = await self._exchange_caller_assertion(assertion)
except Agent365TokenExchangeError as exc:
return self._handle_unavailable(
data=data, tool_name=tool_name, reason=_gateway_fault_reason(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,
)
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
@ -248,15 +285,6 @@ class Agent365Guardrail(CustomGuardrail):
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,
@ -487,6 +515,48 @@ class Agent365Guardrail(CustomGuardrail):
assertion, self._exchange_server, self._exchange_config, tenant_id=self.tenant_id
)
async def _post_entra_token_endpoint(
self,
url: str,
form: dict[str, str], # mutable-ok: ExchangeHttpPost contract
client_auth_headers: dict[str, str], # mutable-ok: ExchangeHttpPost contract
) -> dict[str, object] | None:
"""The exchanger's HTTP edge for this guardrail: every way Entra can fail keeps its own exception, so
the verdict reason the Logs row carries names the OAuth error code or the transport fault."""
response: Final = await self._post_allowing_error_status(
url=url,
data=form,
headers={"Content-Type": "application/x-www-form-urlencoded", **client_auth_headers},
)
if response.status_code in (408, 429):
raise Agent365ThrottledError(status_code=response.status_code)
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[object] = response.json()
except ValueError as exc:
raise Agent365MalformedResponseError("the Entra token endpoint returned a non-JSON body") from exc
try:
body: Final = _JSON_OBJECT_ADAPTER.validate_python(parsed_body)
except ValidationError as exc:
raise Agent365MalformedResponseError("the Entra token endpoint returned a non-object JSON body") from exc
if response.status_code >= 400:
oauth_error: Final = oauth_error_fields(response)
gateway_fault: Final = oauth_error.gateway_fault
if gateway_fault is not None:
raise Agent365TokenExchangeError(error_code=gateway_fault)
raise SubjectTokenRejected(oauth_error.error or "invalid_grant", claims=oauth_error.claims)
if "access_token" not in body:
raise Agent365MalformedResponseError("the Entra token endpoint returned no access_token")
access_token: Final = body["access_token"]
if not isinstance(access_token, str) or not access_token:
raise Agent365MalformedResponseError("the Entra token endpoint returned a non-string access_token")
return body
async def preflight_caller_sign_in(
self, server: MCPServer, user_api_key_auth: "UserAPIKeyAuth | None", subject_token: str
) -> CallerSignInPreflight:
@ -496,25 +566,25 @@ class Agent365Guardrail(CustomGuardrail):
assertion: Final = entra_assertion(subject_token)
if assertion is None:
return Rejected(detail="the caller's bearer is not an Entra token; sign in with Entra and retry")
fail_open: Final = self.unreachable_fallback == "fail_open"
try:
exchange_result: Final = await self._exchange_caller_assertion(assertion)
except Agent365TokenExchangeError as exc:
return Unavailable(detail=_gateway_fault_reason(exc.error_code), fail_open=fail_open)
except Agent365ThrottledError as exc:
return Unavailable(detail=f"the Entra token endpoint returned HTTP {exc.status_code}", fail_open=False)
except (httpx.HTTPError, LitellmTimeout, TimeoutError) as exc:
return Unavailable(
detail=f"the Entra token endpoint could not be reached ({type(exc).__name__})",
fail_open=self.unreachable_fallback == "fail_open",
detail=f"the Entra token endpoint could not be reached ({type(exc).__name__})", fail_open=fail_open
)
except Agent365MalformedResponseError as exc:
return Unavailable(detail=str(exc), fail_open=fail_open)
if isinstance(exchange_result, Ok):
return SignedIn()
error: Final = exchange_result.error
if error.tag == "unauthorized":
return Rejected(detail=error.unauthorized.detail, claims=error.unauthorized.claims)
detail: Final = (
f"Entra rejected the gateway's own Agent 365 credentials ({error.misconfigured}); "
"check the guardrail's client_id and client_secret"
if error.tag == "misconfigured"
else f"the Entra token exchange failed ({error.summary})"
)
return Unavailable(detail=detail, fail_open=self.unreachable_fallback == "fail_open")
return Unavailable(detail=f"the Entra token exchange failed ({error.summary})", fail_open=fail_open)
async def _post_allowing_error_status(
self,

View file

@ -53,34 +53,23 @@ def test_build_gives_each_caller_an_independent_cache():
assert build_token_exchanger() is not build_token_exchanger()
def _recording_client(seen: list[float | None]):
class _Resp:
def raise_for_status(self) -> None:
return None
def json(self) -> dict[str, object]:
return {"access_token": "x", "expires_in": 60}
class _Client:
async def post(self, url, headers, data, timeout=None):
seen.append(timeout)
return _Resp()
return _Client()
@pytest.mark.asyncio
@pytest.mark.parametrize("request_timeout", [0.5, None], ids=["bounded", "handler_default"])
async def test_built_exchanger_posts_with_the_configured_request_timeout(request_timeout):
seen: list[float | None] = []
async def test_build_token_exchanger_drives_the_injected_http_edge():
seen: list[tuple[str, dict[str, str]]] = []
async def post(url: str, form: dict[str, str], client_auth_headers: dict[str, str]) -> dict[str, object] | None:
seen.append((url, form))
return {"access_token": "x", "expires_in": 60}
config = TokenExchangeConfig(
token_exchange_endpoint="https://idp/token", client_id="cid", client_secret=SecretStr("csec")
)
server = ServerSpec(server_id="srv", resource="https://up.example.com", config=config)
with patch(_HTTP_CLIENT, return_value=_recording_client(seen)):
result = await build_token_exchanger(request_timeout=request_timeout).exchange("jwt", server, config)
result = await build_token_exchanger(post=post).exchange("jwt", server, config)
assert isinstance(result, Ok)
assert seen == [request_timeout]
assert result.ok.access_token == "x"
assert [url for url, _ in seen] == ["https://idp/token"]
assert seen[0][1]["subject_token"] == "jwt"
@pytest.mark.asyncio

View file

@ -8933,6 +8933,7 @@ class TestConnectPreflightRoutesLikeTheScopedRouter:
user_api_key_auth=UserAPIKeyAuth(api_key="sk-granted-docs", user_id="u-1"),
client_ip=None,
raw_headers={"x-litellm-api-key": "sk-granted-docs"},
connecting=True,
)
@pytest.mark.asyncio
@ -9021,6 +9022,7 @@ class TestConnectPreflightRoutesLikeTheScopedRouter:
user_api_key_auth=UserAPIKeyAuth(api_key="sk-collision", user_id="u-1"),
client_ip=None,
raw_headers={"x-litellm-api-key": "sk-collision"},
connecting=True,
)
@pytest.mark.asyncio
@ -9128,6 +9130,7 @@ class TestConnectPreflightRoutesLikeTheScopedRouter:
user_api_key_auth=UserAPIKeyAuth(api_key="sk-group", user_id="u-1"),
client_ip=None,
raw_headers={"x-litellm-api-key": "sk-group"},
connecting=True,
)
assert outcome is None, (
@ -10450,6 +10453,7 @@ class TestPreemptive401ModeAware:
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"),
client_ip=None,
connecting=True,
)
async def _connect_with_a_grant(self, server, requested: str, path: str) -> HTTPException:
@ -10472,6 +10476,7 @@ class TestPreemptive401ModeAware:
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"),
client_ip=None,
connecting=True,
)
return exc.value
@ -10553,6 +10558,7 @@ class TestPreemptive401ModeAware:
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"),
client_ip=None,
connecting=True,
)
assert outcome is None, "an unselected aggregate connect must pass without a sign-in challenge"
@ -10655,6 +10661,7 @@ class TestPreemptive401ModeAware:
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"),
client_ip=None,
connecting=True,
)
assert exc.value.status_code == 401
@ -10763,6 +10770,7 @@ class TestSingleServerPreflightReachesIdJag:
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
client_ip=None,
connecting=True,
)
@pytest.mark.asyncio
@ -10821,6 +10829,7 @@ class TestSingleServerPreflightReachesIdJag:
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
client_ip=None,
connecting=True,
)
assert exc.value.status_code == 401
@ -10889,6 +10898,7 @@ class TestOboPreflightScopedToAllowedServers:
"x-litellm-api-key": user_api_key_auth.api_key if user_api_key_auth else "",
"authorization": self.SUBJECT_HEADERS["Authorization"],
},
connecting=True,
)
return allowed_lookup, preflight
@ -10951,6 +10961,7 @@ class TestOboChallengeGateKeepsBaseConnectRules:
user_api_key_auth=UserAPIKeyAuth(api_key="sk-1234", user_id="u-1"),
client_ip=None,
raw_headers={"x-litellm-api-key": "sk-1234", "authorization": "Bearer sk-1234"},
connecting=True,
)
@pytest.mark.asyncio
@ -11795,6 +11806,7 @@ class TestConnectChallengeResolver:
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
client_ip=None,
connecting=True,
)
finally:
litellm.logging_callback_manager.remove_callback_from_list_by_object(
@ -11836,6 +11848,7 @@ class TestConnectChallengeResolver:
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
client_ip=None,
connecting=True,
)
finally:
litellm.logging_callback_manager.remove_callback_from_list_by_object(
@ -11888,6 +11901,7 @@ class TestConnectChallengeResolver:
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
client_ip=None,
connecting=True,
)
assert exc.value.status_code == 401
@ -11955,6 +11969,7 @@ class TestConnectChallengeResolver:
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
client_ip=client_ip,
connecting=True,
)
finally:
litellm.logging_callback_manager.remove_callback_from_list_by_object(
@ -12003,6 +12018,7 @@ class TestConnectChallengeResolver:
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
client_ip=None,
connecting=True,
)
finally:
litellm.logging_callback_manager.remove_callback_from_list_by_object(
@ -12019,7 +12035,7 @@ class TestConnectSignInPreflight:
"""A subject token the provider rejects must fail at connect with the RFC 9728 challenge, not as a
JSON-RPC error on every tools/call where ``WWW-Authenticate`` is lost."""
async def _connect(self, route_names, guardrail, allowed, raw_headers=None):
async def _connect(self, route_names, guardrail, allowed, raw_headers=None, connecting=True):
from litellm.proxy._experimental.mcp_server import server as server_module
server = _catalog_server()
@ -12044,6 +12060,7 @@ class TestConnectSignInPreflight:
client_ip=None,
raw_headers=raw_headers
or {"x-litellm-api-key": "sk-litellm-virtual-key", "authorization": "Bearer entra.jwt.token"},
connecting=connecting,
)
@pytest.mark.asyncio
@ -12091,6 +12108,97 @@ class TestConnectSignInPreflight:
assert exc.value.status_code == 503
assert exc.value.detail == "the Entra token endpoint could not be reached"
@pytest.mark.asyncio
@pytest.mark.parametrize(
("rpc_method", "session_id", "reaches"),
[
("initialize", None, None),
("tools/call", "open-session-1", "stateful"),
("tools/call", None, "stateless"),
],
ids=["connect_503", "open_session_tools_call", "stateless_tools_call"],
)
async def test_fail_closed_outage_is_the_connects_503_only_on_initialize(self, rpc_method, session_id, reaches):
"""Only the ``initialize`` POST turns a fail-closed provider outage into the connect's 503. Every other
JSON-RPC POST on the gated route must reach the session manager with its body intact, so the tools/call
hook answers the outage inside the result envelope and writes the guardrail Logs row, as base did."""
from litellm.proxy._experimental.mcp_server import server as server_module
from litellm.proxy._experimental.mcp_server.caller_sign_in import Unavailable
server = _catalog_server()
guardrail = _CallerSignInGuardrail(
guardrail_name="sign-in-stub",
preflight_result=Unavailable("Entra rejected the gateway's own Agent 365 credentials", fail_open=False),
)
body = json.dumps({"jsonrpc": "2.0", "id": 7, "method": rpc_method, "params": {}}).encode()
session_headers = [(b"mcp-session-id", session_id.encode())] if session_id else []
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/catalog",
"headers": [(b"content-type", b"application/json"), (b"authorization", b"Bearer entra.jwt.token")]
+ session_headers,
}
receive = AsyncMock(return_value={"type": "http.request", "body": body, "more_body": False})
send = AsyncMock()
delivered: dict[str, bytes] = {} # mutable-ok: records which manager saw the replayed body
def _recorder(manager: str):
async def handle(_scope, replayed_receive, _send):
delivered[manager] = (await replayed_receive())["body"]
return handle
litellm.logging_callback_manager.add_litellm_callback(guardrail)
try:
with (
patch.object(
server_module,
"extract_mcp_auth_context",
AsyncMock(
return_value=(
UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
None,
["catalog"],
None,
None,
{"x-litellm-api-key": "sk-litellm-virtual-key", "authorization": "Bearer entra.jwt.token"},
)
),
),
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])),
patch.object(server_module, "_SESSION_MANAGERS_INITIALIZED", True),
patch.object(server_module.session_manager_stateful, "handle_request", _recorder("stateful")),
patch.object(server_module.session_manager_stateless, "handle_request", _recorder("stateless")),
patch.object(
server_module.session_manager_stateful,
"_server_instances",
{session_id: MagicMock()} if session_id else {},
),
):
if reaches is None:
with pytest.raises(HTTPException) as exc:
await server_module.handle_streamable_http_mcp(scope, receive, send)
assert exc.value.status_code == 503
assert exc.value.detail == "Entra rejected the gateway's own Agent 365 credentials"
else:
await server_module.handle_streamable_http_mcp(scope, receive, send)
finally:
litellm.logging_callback_manager.remove_callback_from_list_by_object(
litellm.callbacks, guardrail, require_self=False
)
if session_id:
server_module._remove_stateful_session_tracking(session_id)
assert guardrail.preflight_calls == ["entra.jwt.token"]
assert delivered == ({} if reaches is None else {reaches: body})
send.assert_not_awaited()
@pytest.mark.asyncio
async def test_multi_server_connect_never_awaits_the_preflight(self):
server = _catalog_server()

View file

@ -3,7 +3,6 @@ import time
import uuid
from types import SimpleNamespace
from typing import Any, Final
from unittest.mock import patch
import httpx
import pytest
@ -24,10 +23,6 @@ from litellm.proxy._experimental.mcp_server.caller_sign_in import (
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import OAuthToken
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok, Result
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_provider import (
_post_exchange_endpoint,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import OboTokenExchanger
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
CredError,
ServerSpec,
@ -67,25 +62,6 @@ def _response(status_code: int, payload: Any = None, text: str | None = None) ->
return httpx.Response(status_code=status_code, text=text or "", request=request)
_HTTP_CLIENT: Final = "litellm.llms.custom_httpx.http_handler.get_async_httpx_client"
def _entra_rejecting_with(body: dict[str, object]) -> object:
"""An httpx client whose token POST raises the HTTPStatusError the real exchanger classifies."""
request: Final = httpx.Request("POST", TOKEN_URL)
response: Final = httpx.Response(401, json=body, request=request)
class _Resp:
def raise_for_status(self) -> None:
raise httpx.HTTPStatusError("unauthorized", request=request, response=response)
class _Client:
async def post(self, *args: object, **kwargs: object) -> _Resp:
return _Resp()
return _Client()
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
@ -230,6 +206,24 @@ def _default_fallback_guardrail(handler: FakeHandler, exchanger: StubTokenExchan
)
def _entra_driven_guardrail(
handler: FakeHandler, *, unreachable_fallback: str = "fail_closed", request_timeout: float = 10.0
) -> Agent365Guardrail:
"""A guardrail whose Entra exchange runs through the real exchanger and the guardrail's own HTTP edge, so
``handler`` answers the token POST first and the evaluate POST after it."""
return Agent365Guardrail(
guardrail_name="agent-365-guard",
tenant_id="tenant-abc",
client_id="client-xyz",
client_secret="secret-123",
unreachable_fallback=unreachable_fallback,
request_timeout=request_timeout,
async_handler=handler,
event_hook="pre_mcp_call",
default_on=True,
)
def _server(**overrides: Any) -> MCPServer:
kwargs: Final[dict] = {
"server_id": "outlook-id",
@ -793,31 +787,32 @@ class TestUnreachableFallback:
"""Entra answers a garbled or unverifiable caller assertion with invalid_client AADSTS5002723, the
same top-level code as a wrong gateway secret. The sub-code makes it the caller's 401 challenge,
never the fail-open Unscanned pass and never a 503 that blames the gateway credentials."""
exchanger: Final = OboTokenExchanger(_post_exchange_endpoint)
handler: Final = FakeHandler([])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open")
data: Final = _mcp_data()
with (
patch(
_HTTP_CLIENT,
return_value=_entra_rejecting_with(
handler: Final = FakeHandler(
[
_response(
400,
{
"error": "invalid_client",
"error_description": "AADSTS5002723: Invalid JWT token. Token is not well formed.",
"error_codes": [5002723],
}
),
),
pytest.raises(HTTPException) as exc_info,
):
},
)
]
)
guardrail: Final = _entra_driven_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 == 401
assert "On-Behalf-Of token exchange was rejected" in exc_info.value.detail["message"]
info: Final = _guardrail_info(data)
assert info["guardrail_status"] == "guardrail_intervened"
assert info["guardrail_response"]["verdict"] == "Rejected"
assert "client_secret" not in info["guardrail_response"]["reason"]
assert handler.calls == []
assert (
info["guardrail_response"]["reason"]
== "the Entra On-Behalf-Of token exchange was rejected (invalid_client)"
)
assert [call.url for call in handler.calls] == [TOKEN_URL]
@pytest.mark.asyncio
async def test_evaluate_4xx_blocks_even_fail_open(self):
@ -847,22 +842,36 @@ class TestUnreachableFallback:
@pytest.mark.asyncio
@pytest.mark.parametrize(
"error_code", ["invalid_client", "unauthorized_client", "invalid_scope", "invalid_resource"]
("entra", "error_code"),
[
(_response(400, {"error": "invalid_scope"}), "invalid_scope"),
(_response(401, {"error": "invalid_client"}), "invalid_client"),
(_response(400, {"error": "invalid_client", "error_codes": [7000215]}), "invalid_client"),
(_response(400, {"error": "unauthorized_client"}), "unauthorized_client"),
],
ids=["invalid_scope", "invalid_client_401", "invalid_client_wrong_secret", "unauthorized_client"],
)
async def test_gateway_credential_rejection_is_unavailable_not_a_caller_401(self, error_code: str):
exchanger: Final = StubTokenExchanger([Error(CredError.of_misconfigured(error_code))])
handler: Final = FakeHandler([])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
async def test_gateway_credential_rejection_is_unavailable_not_a_caller_401(
self, entra: httpx.Response, error_code: str
):
"""The verdict reason carries the OAuth error code Entra answered with, the text the Logs row and the
503 detail show an admin, not a generic exchanger summary."""
handler: Final = FakeHandler([entra])
guardrail: Final = _entra_driven_guardrail(handler)
data: Final = _mcp_data()
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, data)
assert exc_info.value.status_code == 503
assert exc_info.value.headers is None or "WWW-Authenticate" not in exc_info.value.headers
assert f"({error_code})" in exc_info.value.detail["message"]
info: Final = _guardrail_info(data)
assert info["guardrail_status"] == "guardrail_failed_to_respond"
assert info["guardrail_response"]["verdict"] == "Unavailable"
assert error_code in info["guardrail_response"]["reason"]
assert "client_secret" in info["guardrail_response"]["reason"]
assert info["guardrail_response"]["reason"] == (
f"Entra rejected the gateway's own Agent 365 credentials ({error_code}); "
"check the guardrail's client_id and client_secret"
)
assert [call.url for call in handler.calls] == [TOKEN_URL]
@pytest.mark.asyncio
async def test_caller_rejection_reason_does_not_blame_the_gateway_credentials(self):
@ -879,16 +888,16 @@ class TestUnreachableFallback:
@pytest.mark.asyncio
async def test_gateway_credential_rejection_follows_fail_open(self):
exchanger: Final = StubTokenExchanger([Error(CredError.of_misconfigured("invalid_client"))])
handler: Final = FakeHandler([])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open")
handler: Final = FakeHandler([_response(401, {"error": "invalid_client"})])
guardrail: Final = _entra_driven_guardrail(handler, 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 "invalid_client" in info["guardrail_response"]["reason"]
assert "(invalid_client)" in info["guardrail_response"]["reason"]
assert [call.url for call in handler.calls] == [TOKEN_URL]
@pytest.mark.asyncio
async def test_exchange_upstream_unavailable_is_unavailable_with_the_summary_not_a_caller_401(self):
@ -937,6 +946,92 @@ class TestUnreachableFallback:
assert info["guardrail_response"]["verdict"] == "Unscanned"
class TestEntraTokenEndpointReasons:
"""Every way the Entra token endpoint can fail keeps its own verdict reason, since that text is what the
guardrail Logs row and the 503 detail carry."""
@pytest.mark.asyncio
@pytest.mark.parametrize(
("entra", "reason"),
[
(
_response(500, text="<html>gateway</html>"),
"the Entra token endpoint could not be reached (HTTPStatusError)",
),
(httpx.ConnectError("refused"), "the Entra token endpoint could not be reached (ConnectError)"),
(_response(200, text="<html>waf page</html>"), "the Entra token endpoint returned a non-JSON body"),
(_response(200, payload=["x"]), "the Entra token endpoint returned a non-object JSON body"),
(_response(200, payload={"token_type": "Bearer"}), "the Entra token endpoint returned no access_token"),
(
_response(200, payload={"access_token": 7}),
"the Entra token endpoint returned a non-string access_token",
),
],
ids=["http_500", "connect_error", "non_json", "non_object", "no_access_token", "non_string_access_token"],
)
async def test_unavailable_reason_names_the_fault(self, entra: object, reason: str):
handler: Final = FakeHandler([entra])
guardrail: Final = _entra_driven_guardrail(handler)
data: Final = _mcp_data()
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, data)
assert exc_info.value.status_code == 503
assert reason in exc_info.value.detail["message"]
info: Final = _guardrail_info(data)
assert info["guardrail_status"] == "guardrail_failed_to_respond"
assert info["guardrail_response"]["verdict"] == "Unavailable"
assert info["guardrail_response"]["reason"] == reason
assert [call.url for call in handler.calls] == [TOKEN_URL]
@pytest.mark.asyncio
@pytest.mark.parametrize("error_code", ["invalid_grant", "interaction_required", "invalid_resource"])
async def test_caller_rejection_reason_names_the_oauth_code(self, error_code: str):
handler: Final = FakeHandler([_response(400, {"error": error_code, "error_codes": [700082]})])
guardrail: Final = _entra_driven_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 == 401
assert _guardrail_info(data)["guardrail_response"]["reason"] == (
f"the Entra On-Behalf-Of token exchange was rejected ({error_code})"
)
@pytest.mark.asyncio
@pytest.mark.parametrize("unreachable_fallback", ["fail_closed", "fail_open"])
async def test_throttled_token_endpoint_blocks_regardless_of_fallback(self, unreachable_fallback: str):
handler: Final = FakeHandler([_response(429, text="slow down"), _response(429, text="slow down")])
guardrail: Final = _entra_driven_guardrail(handler, unreachable_fallback=unreachable_fallback)
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_response"]["verdict"] == "Throttled"
assert info["guardrail_response"]["reason"] == "the Entra token endpoint returned HTTP 429"
connect: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION)
assert connect == Unavailable(detail="the Entra token endpoint returned HTTP 429", fail_open=False)
@pytest.mark.asyncio
async def test_preflight_gateway_fault_detail_names_the_oauth_code(self):
handler: Final = FakeHandler([_response(400, {"error": "invalid_scope"})])
guardrail: Final = _entra_driven_guardrail(handler)
verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION)
assert verdict == Unavailable(
detail=(
"Entra rejected the gateway's own Agent 365 credentials (invalid_scope); "
"check the guardrail's client_id and client_secret"
),
fail_open=False,
)
@pytest.mark.asyncio
async def test_preflight_unavailable_detail_names_the_fault(self):
handler: Final = FakeHandler([_response(200, text="<html>waf page</html>")])
guardrail: Final = _entra_driven_guardrail(handler, unreachable_fallback="fail_open")
verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION)
assert verdict == Unavailable(detail="the Entra token endpoint returned a non-JSON body", fail_open=True)
class TestOboTokenCache:
@pytest.mark.asyncio
async def test_same_assertion_reuses_token(self):
@ -1354,48 +1449,24 @@ class TestPreflightCallerSignIn:
@pytest.mark.asyncio
async def test_configured_timeout_bounds_the_entra_exchange_leg(self):
seen: Final[list[object]] = []
handler: Final = FakeHandler([_response(200, {"access_token": "exchanged", "expires_in": 3600})])
guardrail: Final = _entra_driven_guardrail(handler, request_timeout=0.5)
class _Resp:
def raise_for_status(self) -> None:
return None
def json(self) -> dict[str, object]:
return {"access_token": "exchanged", "expires_in": 3600}
class _Client:
async def post(self, *args: object, **kwargs: object) -> _Resp:
seen.append(kwargs.get("timeout"))
return _Resp()
guardrail: Final = Agent365Guardrail(
guardrail_name="a365",
tenant_id="tenant-abc",
client_id="cid",
client_secret="csecret",
request_timeout=0.5,
async_handler=FakeHandler([]),
)
with patch(_HTTP_CLIENT, return_value=_Client()):
verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION)
verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION)
assert verdict == SignedIn()
assert seen == [0.5], "the Entra token POST must carry the guardrail's own request_timeout"
assert [(call.url, call.timeout) for call in handler.calls] == [(TOKEN_URL, 0.5)], (
"the Entra token POST must carry the guardrail's own request_timeout"
)
@pytest.mark.asyncio
async def test_malformed_assertion_is_rejected_at_connect_even_fail_open(self):
exchanger: Final = OboTokenExchanger(_post_exchange_endpoint)
guardrail: Final = _make_guardrail(FakeHandler([]), exchanger=exchanger, unreachable_fallback="fail_open")
handler: Final = FakeHandler([_response(401, {"error": "invalid_client", "error_codes": [5002723]})])
guardrail: Final = _entra_driven_guardrail(handler, unreachable_fallback="fail_open")
with patch(
_HTTP_CLIENT,
return_value=_entra_rejecting_with({"error": "invalid_client", "error_codes": [5002723]}),
):
verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION)
verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION)
assert isinstance(verdict, Rejected)
assert "client_secret" not in verdict.detail
assert verdict == Rejected(detail="invalid_client", claims=None)
@pytest.mark.asyncio
async def test_misconfigured_fail_closed_is_unavailable(self):