mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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
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:
parent
6b8d3cbf15
commit
3e7efe95d6
5 changed files with 116 additions and 9 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue