fix(mcp): restore the raw bearer hook kwarg, exact-name-first resolution, and the base OBO challenge gate

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-26 23:21:02 +00:00
parent 4be107e779
commit 692086a0ae
6 changed files with 147 additions and 16 deletions

View file

@ -5971,8 +5971,14 @@ class MCPServerManager:
if proxy_logging_obj is None:
return hook_result
# Admission credentials are never handed to guardrails as the caller's assertion.
incoming_bearer_token: Final = self._extract_subject_token(None, raw_headers, user_api_key_auth)
inbound_authorization: Final = next(
(v for k, v in (raw_headers or {}).items() if isinstance(k, str) and k.lower() == "authorization"),
"",
)
incoming_bearer_token: Final = (
inbound_authorization[len("bearer ") :] if inbound_authorization.lower().startswith("bearer ") else None
)
incoming_subject_token: Final = self._extract_subject_token(None, raw_headers, user_api_key_auth)
pre_hook_kwargs: Final = {
"guardrail_context": guardrail_context,
@ -5988,6 +5994,7 @@ class MCPServerManager:
),
"user_api_key_hash": (getattr(user_api_key_auth, "api_key_hash", None) if user_api_key_auth else None),
"incoming_bearer_token": incoming_bearer_token,
"incoming_subject_token": incoming_subject_token,
"headers": logging_safe_mcp_headers(raw_headers),
"tool_description": tool.description if tool is not None else None,
"tool_input_schema": tool.input_schema if tool is not None else None,
@ -7192,17 +7199,16 @@ class MCPServerManager:
return None
def get_mcp_server_answering_to(self, name: str, client_ip: str | None = None) -> MCPServer | None:
"""The server a scoped ``/mcp/{name}`` connect resolves to, matched the way the router matches
it: case-insensitive over server_id, name and every published prefix form, then the exact
name lookup as the fallback."""
return next(
"""The server a scoped ``/mcp/{name}`` connect resolves to: the alias-first exact lookup, then
the router's case-insensitive prefix match."""
return self.get_mcp_server_by_name(name, client_ip=client_ip) or next(
(
server
for server in self.get_filtered_registry(client_ip).values()
if server_answers_to_name(server, name)
),
None,
) or self.get_mcp_server_by_name(name, client_ip=client_ip)
)
def get_filtered_registry(self, client_ip: str | None = None) -> dict[str, MCPServer]:
"""

View file

@ -1715,8 +1715,8 @@ if MCP_AVAILABLE:
continue
# Caller sign-in: challenge at connect because a tool-call-time 401 is wrapped into a
# JSON-RPC error and the WWW-Authenticate header is lost. Non-OBO gates fire only on a
# single-server connect the key's grant admits.
# JSON-RPC error and the WWW-Authenticate header is lost. OBO keeps its connect gate;
# guardrail-only gates fire only on a single-server connect the key's grant admits.
sign_in = caller_sign_in_for(server, user_api_key_auth) if server is not None else None
if (
server
@ -1726,7 +1726,7 @@ if MCP_AVAILABLE:
)
is None
and (
server.auth_type == MCPAuth.oauth2_token_exchange
(server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers)
or await _key_granted_single_server(server, mcp_servers, user_api_key_auth, client_ip)
)
):

View file

@ -195,7 +195,7 @@ class Agent365Guardrail(CustomGuardrail):
return data
tool_name: Final = str(data.get("mcp_tool_name") or "")
assertion: Final = entra_assertion(data.get("incoming_bearer_token"))
assertion: Final = entra_assertion(data.get("incoming_subject_token"))
if assertion is None:
self._handle_caller_fault(
data=data,

View file

@ -10270,6 +10270,50 @@ class TestOboPreflightScopedToAllowedServers:
)
class TestOboChallengeGateKeepsBaseConnectRules:
"""An OBO connect carrying a bearer in oauth2_headers is challenged by the exchange path, not the
preemptive gate, so a multi-server connect or a single-server connect with any bearer at all must
not be refused before the session opens."""
LITELLM_KEY_BEARER = {"Authorization": "Bearer sk-1234"}
async def _run(self, servers: list[MCPServer], mcp_servers: list[str]) -> None:
from litellm.proxy._experimental.mcp_server import server as server_module
with (
patch.object( # test-quality-ok: route wiring must use the manager's configured server
mcp_operations.global_mcp_server_manager,
"get_mcp_server_answering_to",
return_value=servers[0],
),
patch.object( # test-quality-ok: allowed-set resolution needs the DB; the test controls its answer
mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=servers)
),
):
await server_module._raise_preemptive_401_for_unauthenticated_servers(
scope={"type": "http", "method": "POST", "path": "/mcp/obo", "headers": []},
mcp_servers=mcp_servers,
oauth2_headers=self.LITELLM_KEY_BEARER,
mcp_server_auth_headers=None,
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"},
)
@pytest.mark.asyncio
async def test_multi_server_connect_with_any_bearer_is_not_preemptively_challenged(self):
obo = _make_obo_server("obo")
catalog = MCPServer(server_id="id-catalog", name="catalog", alias="catalog", transport=MCPTransport.http)
await self._run([obo, catalog], ["obo", "catalog"])
@pytest.mark.asyncio
async def test_single_obo_connect_with_litellm_key_bearer_still_challenges(self):
with pytest.raises(HTTPException) as exc:
await self._run([_make_obo_server("obo")], ["obo"])
assert exc.value.status_code == 401
@pytest.mark.asyncio
async def test_post_mcp_call_guardrails_return_the_rewritten_result():
"""The result a post_mcp_call guardrail rewrote must be what the caller sends back."""

View file

@ -6186,6 +6186,24 @@ class TestMCPServerManager:
assert manager._resolve_mcp_server_for_tool_call("zapier", "create_zap") is zapier
assert manager._resolve_mcp_server_for_tool_call("other", "create_zap") is other
@pytest.mark.parametrize("gh_first", [True, False], ids=["gh-listed-first", "gh-public-listed-first"])
def test_answering_to_prefers_alias_over_earlier_prefix_match(self, gh_first):
manager = MCPServerManager()
gh = MCPServer(
server_id="gh-id", name="gh", server_name="gh", transport=MCPTransport.http, auth_type=MCPAuth.oauth2
)
gh_public = MCPServer(
server_id="gh-public-id", name="gh_public", server_name="gh_public", alias="gh", transport=MCPTransport.http
)
manager.registry = (
{"gh-id": gh, "gh-public-id": gh_public} if gh_first else {"gh-public-id": gh_public, "gh-id": gh}
)
assert manager.get_mcp_server_answering_to("gh") is gh_public
assert manager.get_mcp_server_answering_to("gh-public-id") is gh_public
assert manager.get_mcp_server_answering_to("GH_PUBLIC") is gh_public
assert manager.get_mcp_server_answering_to("gh-id") is gh
def test_remove_server_drops_only_its_own_tool_mapping_rows(self):
manager = self._manager_with_deepwiki_and_huggingface()
@ -6813,6 +6831,59 @@ class TestMCPServerManager:
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
@pytest.mark.parametrize(
("raw_headers", "api_key", "expected_bearer", "expected_subject"),
[
pytest.param(
{"authorization": "Bearer eyJ.x.y"},
"eyJ.x.y",
"eyJ.x.y",
None,
id="idp-token-as-admission-stays-raw-bearer",
),
pytest.param(
{"authorization": "Bearer sk-1234"},
"sk-1234",
"sk-1234",
None,
id="litellm-key-as-bearer-is-not-a-subject",
),
pytest.param(
{"x-litellm-api-key": "sk-1234", "authorization": "Bearer eyJ.x.y"},
"sk-1234",
"eyJ.x.y",
"eyJ.x.y",
id="key-admission-plus-idp-bearer-subject",
),
],
)
async def test_pre_call_tool_check_separates_raw_bearer_from_subject(
self, raw_headers, api_key, expected_bearer, expected_subject
):
manager = MCPServerManager()
server = MCPServer(
server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv", allowed_tools=None
)
proxy_logging = MagicMock()
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={})
proxy_logging.pre_call_hook = AsyncMock(return_value=None)
await manager.pre_call_tool_check(
server_name="srv",
name="turn",
arguments={},
user_api_key_auth=UserAPIKeyAuth(api_key=api_key, user_id="u"),
proxy_logging_obj=proxy_logging,
server=server,
raw_headers=raw_headers,
)
kwargs: Final = proxy_logging._create_mcp_request_object_from_kwargs.call_args.args[0]
assert kwargs["incoming_bearer_token"] == expected_bearer
assert kwargs["incoming_subject_token"] == expected_subject
@pytest.mark.asyncio
async def test_check_tool_permission_for_key_team_allows_permitted_tool(self):
"""

View file

@ -210,7 +210,7 @@ def _mcp_data(**overrides: Any) -> dict:
"mcp_tool_name": "send_email",
"mcp_arguments": {"to": "user@example.com", "body": "hello"},
"mcp_server_name": "outlook_mcp",
"incoming_bearer_token": FAKE_ASSERTION,
"incoming_subject_token": FAKE_ASSERTION,
"metadata": {"headers": {"mcp-session-id": "sess-123"}},
}
data.update(overrides)
@ -699,17 +699,27 @@ class TestUnreachableFallback:
handler: Final = FakeHandler([])
guardrail: Final = _make_guardrail(handler)
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, _mcp_data(incoming_bearer_token=None))
await _run(guardrail, _mcp_data(incoming_subject_token=None))
assert exc_info.value.status_code == 401
assert handler.calls == []
@pytest.mark.asyncio
async def test_raw_bearer_without_subject_token_is_no_bearer(self):
exchanger: Final = StubTokenExchanger()
handler: Final = FakeHandler([])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, _mcp_data(incoming_subject_token=None, incoming_bearer_token=FAKE_ASSERTION))
assert exc_info.value.status_code == 401
assert exchanger.calls == []
@pytest.mark.asyncio
async def test_non_jwt_bearer_token_fail_closed(self):
exchanger: Final = StubTokenExchanger()
handler: Final = FakeHandler([])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, _mcp_data(incoming_bearer_token="sk-litellm-virtual-key"))
await _run(guardrail, _mcp_data(incoming_subject_token="sk-litellm-virtual-key"))
assert exc_info.value.status_code == 401
assert exchanger.calls == []
@ -717,7 +727,7 @@ class TestUnreachableFallback:
async def test_missing_bearer_token_blocks_even_fail_open(self):
handler: Final = FakeHandler([])
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
data: Final = _mcp_data(incoming_bearer_token=None)
data: Final = _mcp_data(incoming_subject_token=None)
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, data)
assert exc_info.value.status_code == 401
@ -871,7 +881,7 @@ class TestOboTokenCache:
handler: Final = FakeHandler([_allow_response(), _allow_response()])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
await _run(guardrail, _mcp_data())
await _run(guardrail, _mcp_data(incoming_bearer_token=other_assertion))
await _run(guardrail, _mcp_data(incoming_subject_token=other_assertion))
assert len(exchanger.calls) == 2
assert handler.calls[1].headers["Authorization"] == "Bearer token-b"