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 7c68224ab1a..a5f41e540e2 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 @@ -81,7 +81,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] + url: str, form: dict[str, str], client_auth_headers: dict[str, str], *, timeout: float | None = None ) -> dict[str, object] | None: from litellm.llms.custom_httpx.http_handler import ( # noqa: PLC0415 get_async_httpx_client, # pyright: ignore @@ -95,7 +95,9 @@ async def _post_exchange_endpoint( headers: Final = {"Accept": "application/json", **client_auth_headers} try: client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) # pyright: ignore - response: Final = await client.post(url, headers=headers, data=form) # pyright: ignore + response: Final = await client.post( # pyright: ignore[reportUnknownMemberType] # untyped handler + url, headers=headers, data=form, timeout=timeout + ) response.raise_for_status() # pyright: ignore parsed: Final[object] = response.json() # pyright: ignore except httpx.HTTPStatusError as status_err: @@ -133,9 +135,12 @@ async def _post_exchange_endpoint( return parsed # pyright: ignore -def build_token_exchanger() -> OboTokenExchanger: +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) + return OboTokenExchanger( - _post_exchange_endpoint, + post, cache=InMemoryTokenCacheBackend(max_size=MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE), default_ttl_seconds=MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL, min_ttl_seconds=MCP_OAUTH2_TOKEN_CACHE_MIN_TTL, 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 ac09821feb3..7c374bd8306 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -185,7 +185,11 @@ class Agent365Guardrail(CustomGuardrail): resource=AGENT_365_PROD_API_BASE, config=self._exchange_config, ) - self._token_exchanger: Final = token_exchanger if token_exchanger is not None else build_token_exchanger() + self._token_exchanger: Final = ( + token_exchanger + if token_exchanger is not None + else build_token_exchanger(request_timeout=self.request_timeout) + ) verbose_proxy_logger.info("Initialized Microsoft Agent 365 guardrail: %s", guardrail_name) @staticmethod 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 589c9ce16f5..b18a0ae48de 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 @@ -9,7 +9,7 @@ from unittest.mock import patch import pytest from pydantic import SecretStr -from litellm.proxy._experimental.mcp_server.outbound_credentials import Error, ServerSpec +from litellm.proxy._experimental.mcp_server.outbound_credentials import Error, Ok, ServerSpec from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_provider import ( _post_exchange_endpoint, build_token_exchanger, @@ -38,7 +38,7 @@ def _client_raising_status(status: int, body: object): raise httpx.HTTPStatusError("bad request", request=request, response=response) class _Client: - async def post(self, url, headers, data): + async def post(self, url, headers, data, timeout=None): return _Resp() return _Client() @@ -53,6 +53,36 @@ 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] = [] + 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) + assert isinstance(result, Ok) + assert seen == [request_timeout] + + @pytest.mark.asyncio async def test_post_returns_none_on_transport_error(): with patch(_HTTP_CLIENT, side_effect=RuntimeError("boom")): @@ -70,7 +100,7 @@ async def test_post_parses_json_body_on_success(): return {"access_token": "x", "expires_in": 60} class _Client: - async def post(self, url, headers, data): + async def post(self, url, headers, data, timeout=None): return _Resp() with patch(_HTTP_CLIENT, return_value=_Client()): @@ -142,7 +172,7 @@ async def test_post_returns_none_on_non_object_json(payload): return payload class _Client: - async def post(self, url, headers, data): + async def post(self, url, headers, data, timeout=None): return _Resp() with patch(_HTTP_CLIENT, return_value=_Client()): 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 bfca3a74aad..ed94de2b90e 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 @@ -10499,6 +10499,43 @@ class TestPreemptive401ModeAware: assert "/gwx" in exact_header assert moved_header == exact_header.replace("/gwx", f"/{requested}") + @pytest.mark.asyncio + @pytest.mark.parametrize("kind", ["plain_obo", "oauth_passthrough"]) + @pytest.mark.parametrize("shape", ["alias_case", "server_id", "x_mcp_servers"]) + async def test_moved_obo_and_passthrough_shapes_get_the_exact_name_routes_challenge(self, kind, shape, monkeypatch): + monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) + server = ( + _make_obo_server("obx") + if kind == "plain_obo" + else MCPServer( + server_id="id-obx", + name="obx", + alias="obx", + server_name="obx", + url="https://obx.test/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + mcp_info={"server_name": "obx"}, + ) + ) + requested, path, exact_path = { + "alias_case": ("OBX", "/mcp/OBX", "/mcp/obx"), + "server_id": (server.server_id, f"/mcp/{server.server_id}", "/mcp/obx"), + "x_mcp_servers": ("OBX", "/mcp", "/mcp"), + }[shape] + + exact = await self._connect_with_a_grant(server, "obx", exact_path) + moved = await self._connect_with_a_grant(server, requested, path) + + assert exact.status_code == 401 + assert (moved.status_code, moved.detail) == (exact.status_code, exact.detail) + exact_header = {k.lower(): v for k, v in (exact.headers or {}).items()}["www-authenticate"] + moved_header = {k.lower(): v for k, v in (moved.headers or {}).items()}["www-authenticate"] + assert "/obx" in exact_header + assert moved_header == exact_header.replace("/obx", f"/{requested}") + @pytest.mark.asyncio async def test_aggregate_connect_without_a_server_selection_is_not_challenged(self): from litellm.proxy._experimental.mcp_server import server as server_module 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 ff94d311f67..d5d0fa6c812 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -1352,6 +1352,37 @@ class TestPreflightCallerSignIn: assert verdict == Rejected(detail="the provided assertion has expired", claims="step-up") + @pytest.mark.asyncio + async def test_configured_timeout_bounds_the_entra_exchange_leg(self): + seen: Final[list[object]] = [] + + 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) + + assert verdict == SignedIn() + assert seen == [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)