From 7603412660b3455239ede2300feedc3b7f92a8fc Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 05:34:55 +0000 Subject: [PATCH] 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> --- .../mcp_server/caller_sign_in.py | 11 +- .../token_exchange_provider.py | 22 +- .../proxy/_experimental/mcp_server/server.py | 41 +-- .../guardrail_hooks/agent_365/agent_365.py | 114 +++++++-- .../test_token_exchange_provider.py | 33 +-- .../test_mcp_server_tool_calls_and_headers.py | 110 +++++++- .../guardrail_hooks/test_agent_365.py | 239 ++++++++++++------ 7 files changed, 406 insertions(+), 164 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py index 1d417c01dbd..74637972f75 100644 --- a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py +++ b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py @@ -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) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_exchange_provider.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_exchange_provider.py index a5f41e540e2..7ce84624847 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_exchange_provider.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_exchange_provider.py @@ -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), diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 861fe36eacc..42ef852c932 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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, ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py index 7c374bd8306..dbc041ec806 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -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, diff --git a/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_token_exchange_provider.py b/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_token_exchange_provider.py index b18a0ae48de..3e714a73a98 100644 --- a/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_token_exchange_provider.py +++ b/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_token_exchange_provider.py @@ -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 diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index ed94de2b90e..2d510e70bc6 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -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() diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py index d5d0fa6c812..0d5f5b3ad8b 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -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="gateway"), + "the Entra token endpoint could not be reached (HTTPStatusError)", + ), + (httpx.ConnectError("refused"), "the Entra token endpoint could not be reached (ConnectError)"), + (_response(200, text="waf page"), "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="waf page")]) + 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):