diff --git a/tests/integration/mcp/test_mcp_agent_365_guardrail.py b/tests/integration/mcp/test_mcp_agent_365_guardrail.py index 6f8aea462f8..f3c923cbf10 100644 --- a/tests/integration/mcp/test_mcp_agent_365_guardrail.py +++ b/tests/integration/mcp/test_mcp_agent_365_guardrail.py @@ -38,6 +38,10 @@ GUARDRAIL_ROWS: Final = ( FALLBACKS: Final = (None, "fail_open", "fail_closed") +def _origin(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + + def _generic_guardrail_outage(request: Request) -> Reply: assert request.target == "/beta/litellm_basic_guardrail_api", request.target return Reply(status=503, body=json.dumps({"error": "synthetic sibling guardrail outage"}).encode()) @@ -136,17 +140,21 @@ def test_a_missing_or_malformed_caller_bearer_blocks_on_every_entry_point_whatev ) -> None: with _rig(gateway, tmp_path, fallback) as rig: assert f"{rig.alias}-add" in rig.caller().list_tools().tools, "the catalog needs only the virtual key" + metadata_url: Final = f"{_origin(rig.candidate)}/.well-known/oauth-protected-resource/{rig.alias}/mcp" for entry in ENTRY_POINTS: - missing: Final = rig.caller(entry).call(f"{rig.alias}-add", {"entry": entry}, server_id=rig.server_id) - malformed: Final = rig.caller(entry, "not-a-jws").call( - f"{rig.alias}-add", {"entry": entry}, server_id=rig.server_id - ) - for label, outcome in (("without a bearer", missing), ("opaque bearer", malformed)): - assert outcome.error is not None, f"{entry} {label}: {outcome.raw}" + for label, bearer in (("without a bearer", None), ("opaque bearer", "not-a-jws")): + caller: Final = rig.caller(entry, bearer) if entry == "server_mcp": - assert outcome.status == 401, f"{entry} {label} skips the connect sign-in challenge: {outcome.raw}" - else: - assert REJECTED in outcome.raw, f"{entry} {label}: {outcome.raw}" + challenged: Final = caller.rpc( + "tools/call", {"name": f"{rig.alias}-add", "arguments": {"entry": entry}} + ) + assert challenged.status_code == 401, f"{entry} {label}: {challenged.text}" + authenticate: Final = challenged.headers.get("www-authenticate", "") + assert f'resource_metadata="{metadata_url}"' in authenticate, f"{entry} {label}: {authenticate!r}" + assert 'error="invalid_token"' in authenticate, f"{entry} {label}: {authenticate!r}" + continue + outcome: Final = caller.call(f"{rig.alias}-add", {"entry": entry}, server_id=rig.server_id) + assert outcome.error is not None and REJECTED in outcome.raw, f"{entry} {label}: {outcome.raw}" assert rig.upstream_tool_names() == () expected: Final = 2 * (len(ENTRY_POINTS) - 1) assert rig.guardrail_statuses("call_mcp_tool", expected) == ["guardrail_intervened"] * expected