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):