fix(mcp): forward caller bearer on REST oauth_delegate tool calls (#42787)

* fix(mcp): forward caller bearer on REST oauth_delegate tool calls

Co-Authored-By: bot_apk <apk@cognition.ai>

* fix(mcp): only forward caller bearer on REST for client-forwarded-token servers

Co-Authored-By: bot_apk <apk@cognition.ai>

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: bot_apk <apk@cognition.ai>
This commit is contained in:
devin-ai-integration[bot] 2026-09-23 14:27:54 -07:00 • committed by GitHub
parent dc14e76147
commit 49d0ece934
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 115 additions and 17 deletions

View file

@ -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"),

View file

@ -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)

View file

@ -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 [],
},
)