fix(mcp): bound the Agent 365 Entra exchange by request_timeout and cover moved OBO and passthrough connect shapes
Some checks failed
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run
Terraform Provider / gofmt, vet, build, test (push) Has been cancelled
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Has been cancelled

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-10-03 00:29:25 +00:00
parent 6b8d3cbf15
commit 3e7efe95d6
5 changed files with 116 additions and 9 deletions

View file

@ -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,

View file

@ -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

View file

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

View file

@ -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

View file

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