mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
4be107e779
commit
692086a0ae
6 changed files with 147 additions and 16 deletions
|
|
@ -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]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue