mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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:
parent
dc14e76147
commit
49d0ece934
3 changed files with 115 additions and 17 deletions
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 [],
|
||||
},
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue