diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index c2f7bf7d531..3a715b80a2d 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -83,6 +83,8 @@ _MCP_GUARDRAIL_REJECTIONS: Final = ( HTTPException, ) +_CLIENT_FORWARDED_TOKEN_AUTH_TYPES: Final = frozenset((MCPAuth.true_passthrough, MCPAuth.oauth_delegate)) + def _connection_error_message(exc: BaseException, url: str | None, timeout_seconds: float) -> str: reference: Final = uuid4().hex @@ -1186,6 +1188,11 @@ if MCP_AVAILABLE: ) if target_server is not None: user_oauth_extra_headers = await _get_user_oauth_extra_headers(target_server, user_api_key_dict) + caller_oauth2_headers: Final = ( + MCPRequestHandler._get_oauth2_headers_from_headers(request.headers) + if target_server is not None and target_server.auth_type in _CLIENT_FORWARDED_TOKEN_AUTH_TYPES + else None + ) # Call execute_mcp_tool directly (permission checks already done) _tool_start_time: Final = datetime.now() @@ -1197,7 +1204,7 @@ if MCP_AVAILABLE: user_api_key_auth=data.get("user_api_key_auth"), mcp_auth_header=data.get("mcp_auth_header"), mcp_server_auth_headers=data.get("mcp_server_auth_headers"), - oauth2_headers=user_oauth_extra_headers or data.get("oauth2_headers"), + oauth2_headers=user_oauth_extra_headers or caller_oauth2_headers, raw_headers=data.get("raw_headers"), client_ip=IPAddressUtils.get_mcp_client_ip(request), litellm_logging_obj=data.get("litellm_logging_obj"), diff --git a/tests/integration/mcp/test_mcp_oauth_flows.py b/tests/integration/mcp/test_mcp_oauth_flows.py index 4efb78dacc2..7a60c8ede30 100644 --- a/tests/integration/mcp/test_mcp_oauth_flows.py +++ b/tests/integration/mcp/test_mcp_oauth_flows.py @@ -215,10 +215,7 @@ def test_delegated_auth_forwards_the_callers_bearer_untouched(gateway: Gateway, peer.drain() outcome: Final = caller.call(f"{alias}-add", ADD, identity if entry in ("mcp", "root", "sse", "rest") else None) assert outcome.ok, outcome.raw - seen: Final = _authorizations(peer) - if seen == (None,) and entry == "rest": - pytest.skip("BUG: /mcp-rest/tools/call drops the caller's Authorization on an oauth_delegate server") - assert seen == (f"Bearer {token}".encode(),), seen + assert _authorizations(peer) == (f"Bearer {token}".encode(),) @dataclass(frozen=True, slots=True) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index 233a8cc96ba..b25538d1814 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -2534,8 +2534,83 @@ class TestCallToolRestAPI: assert captured["name"] == "demo-tool" assert captured["arguments"] == {"foo": "bar"} assert captured["allowed_mcp_servers"] == [stub_server] + assert captured["oauth2_headers"] is None fire_logging.assert_awaited_once() + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("auth_type", "per_user_oauth", "expected"), + [ + ("oauth_delegate", None, {"Authorization": "Bearer user-subject-token"}), + ( + "oauth_delegate", + {"Authorization": "Bearer per-user-oauth-token"}, + {"Authorization": "Bearer per-user-oauth-token"}, + ), + ("oauth2", None, None), + ], + ) + async def test_forwards_callers_bearer_as_oauth2_headers(self, monkeypatch, auth_type, per_user_oauth, expected): + """A distinct caller Authorization rides oauth2_headers to execute_mcp_tool only for + client-forwarded-token servers, with a per-user OAuth token still taking precedence. + A gateway-managed oauth2 server never sees the caller's bearer.""" + + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return ["server-1"] + + class StubServer: + server_id = "server-1" + alias = "server-1" + server_name = "server-1" + name = "stub" + allowed_tools = None + mcp_info = {"server_name": "stub"} + available_on_public_internet = True + + stub_server = StubServer() + stub_server.auth_type = auth_type + + async def fake_add_litellm_data_to_request(**kwargs): + return kwargs.get("data", {}) + + async def fake_get_user_oauth_extra_headers(server, user_api_key_dict, prefetched_creds=None): + return per_user_oauth + + captured = {} + + async def fake_execute_mcp_tool(**kwargs): + captured.update(kwargs) + return {"result": "ok"} + + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, "get_allowed_mcp_servers", fake_get_allowed_mcp_servers + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_mcp_server_by_id", + lambda server_id: stub_server if server_id == "server-1" else None, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.add_litellm_data_to_request", fake_add_litellm_data_to_request) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_config", {}, raising=False) + monkeypatch.setattr(rest_endpoints, "_get_user_oauth_extra_headers", fake_get_user_oauth_extra_headers) + monkeypatch.setattr(rest_endpoints, "execute_mcp_tool", fake_execute_mcp_tool) + monkeypatch.setattr( + rest_endpoints, "_fire_mcp_tool_call_logging", AsyncMock(side_effect=RuntimeError("logging failed")) + ) + + request = _build_request( + {"x-litellm-api-key": "sk-admission-key", "authorization": "Bearer user-subject-token"}, + path="/mcp-rest/tools/call", + method="POST", + json_body={"server_id": "server-1", "name": "demo-tool", "arguments": {}}, + ) + + result = await rest_endpoints.call_tool_rest_api(request, user_api_key_dict=UserAPIKeyAuth()) + + assert result == {"result": "ok"} + assert captured["oauth2_headers"] == expected + assert captured["raw_headers"]["authorization"] == "Bearer user-subject-token" + async def test_returns_guardrail_rewritten_tool_result(self, monkeypatch): """A post_mcp_call guardrail rewrite of the tool result must reach the REST caller, not the raw result the upstream server returned.""" @@ -2847,7 +2922,9 @@ class TestCallToolRestAPI: @pytest.mark.parametrize("raise_site", ["pre_call_hook", "execute_mcp_tool"]) @pytest.mark.parametrize("custom_code", [False, True]) - async def test_guardrail_block_runs_failure_logging_before_http_translation(self, monkeypatch, raise_site, custom_code): + async def test_guardrail_block_runs_failure_logging_before_http_translation( + self, monkeypatch, raise_site, custom_code + ): """A pre_mcp_call guardrail block, whether raised by the pre-call hook or from inside execute_mcp_tool, must reach proxy_logging_obj.post_call_failure_hook (the only path that writes the failure spend-log row) with the logging object's failure payload already built, @@ -2940,7 +3017,9 @@ class TestCallToolRestAPI: assert exc_info.value.status_code == 400 if custom_code: assert exc_info.value.detail == { - "error": "guardrail_violation", "message": "Content blocked", "guardrail_name": "block-all" + "error": "guardrail_violation", + "message": "Content blocked", + "guardrail_name": "block-all", } else: assert exc_info.value is guardrail_error @@ -3118,7 +3197,10 @@ class TestCallToolRestAPI: @pytest.mark.parametrize("selected", [False, True]) @pytest.mark.parametrize("action", ["block", "modify"]) async def test_request_selected_tool_specific_guardrail_applies_to_virtual_execution( - monkeypatch: pytest.MonkeyPatch, virtual: bool, selected: bool, action: str, + monkeypatch: pytest.MonkeyPatch, + virtual: bool, + selected: bool, + action: str, ) -> None: import litellm from litellm.caching.caching import DualCache @@ -3129,16 +3211,23 @@ async def test_request_selected_tool_specific_guardrail_applies_to_virtual_execu from litellm.proxy.utils import ProxyLogging guardrail: Final = CustomCodeGuardrail( - guardrail_name="block-resolved-tool", event_hook="pre_mcp_call", default_on=False, - custom_code='def apply_guardrail(inputs, request_data, input_type):\n' + guardrail_name="block-resolved-tool", + event_hook="pre_mcp_call", + default_on=False, + custom_code="def apply_guardrail(inputs, request_data, input_type):\n" ' if inputs.get("tools", [{}])[0].get("function", {}).get("name") == "execute":\n' f' return {{"action": "{action}", "reason": "resolved tool blocked", "texts": ["redacted"]}}\n' - ' return allow()\n', + " return allow()\n", ) manager: Final = mcp_server_manager.MCPServerManager() managed_server: Final = MCPServer( - server_id="observer", name="observer", server_name="observer", transport="http", - url="https://observer.example/mcp", spec_path="observer.json", auth_type="none", + server_id="observer", + name="observer", + server_name="observer", + transport="http", + url="https://observer.example/mcp", + spec_path="observer.json", + auth_type="none", ) manager.registry = {"observer": managed_server} manager.tool_name_to_mcp_server_name_mapping = {"observer-execute": "observer"} @@ -3161,18 +3250,23 @@ async def test_request_selected_tool_specific_guardrail_applies_to_virtual_execu monkeypatch.setattr(proxy_server, "proxy_config", {}) monkeypatch.setattr(proxy_server, "general_settings", {}) caller: Final = UserAPIKeyAuth( - api_key="hashed-key", request_route="/mcp-rest/tools/call", + api_key="hashed-key", + request_route="/mcp-rest/tools/call", object_permission=LiteLLM_ObjectPermissionTable( - object_permission_id="virtual-test", mcp_servers=["observer"], mcp_tool_search_enabled=True, + object_permission_id="virtual-test", + mcp_servers=["observer"], + mcp_tool_search_enabled=True, ), ) request: Final = _build_request( - path="/mcp-rest/tools/call", method="POST", + path="/mcp-rest/tools/call", + method="POST", json_body={ "name": "mcp_tool_call" if virtual else "observer-execute", "server_id": "observer", "arguments": {"tool_name": "observer-execute", "arguments": {"q": "confidential"}} - if virtual else {"q": "confidential"}, + if virtual + else {"q": "confidential"}, "guardrails": ["block-resolved-tool"] if selected else [], }, )