mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
fix(mcp): keep upstream OAuth Authorization when jwt signer hook injects one on tools/call (#38555)
* fix(mcp): keep upstream OAuth Authorization when jwt signer hook injects one on tools/call Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): only treat server credential as occupying Authorization when it maps to that header Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin <yassin@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
452254963e
commit
e16aa9f512
2 changed files with 268 additions and 29 deletions
|
|
@ -5256,7 +5256,9 @@ class MCPServerManager:
|
|||
proxy_logging_obj: Optional ProxyLogging object for hook integration
|
||||
host_progress_callback: Optional callback for progress updates
|
||||
hook_extra_headers: Optional headers injected by pre_mcp_call guardrail
|
||||
hooks. Merged last (highest priority) into outbound request headers.
|
||||
hooks. Merged last into outbound request headers, except a hook
|
||||
Authorization header is dropped when an upstream credential already
|
||||
occupies the Authorization slot.
|
||||
|
||||
Returns:
|
||||
CallToolResult from the MCP server
|
||||
|
|
@ -5347,27 +5349,26 @@ class MCPServerManager:
|
|||
if hook_extra_headers:
|
||||
if extra_headers is None:
|
||||
extra_headers = {}
|
||||
if "Authorization" in hook_extra_headers:
|
||||
if "Authorization" in extra_headers:
|
||||
verbose_logger.warning(
|
||||
"MCPServerManager: hook_extra_headers 'Authorization' will overwrite "
|
||||
"the existing Authorization header from static_headers. "
|
||||
"The hook JWT will take precedence."
|
||||
)
|
||||
elif server_auth_header is not None:
|
||||
# server_auth_header is passed separately to _create_mcp_client as
|
||||
# auth_value. Both will reach the upstream server — warn so admins
|
||||
# know two Authorization credentials are being sent.
|
||||
verbose_logger.warning(
|
||||
"MCPServerManager: hook_extra_headers injects 'Authorization' while "
|
||||
"server '%s' already has a configured authentication_token. "
|
||||
"Both credentials will be sent; the hook header is in extra_headers "
|
||||
"and the server token is in auth_value — the upstream server decides "
|
||||
"which one wins. Consider unsetting authentication_token if you want "
|
||||
"the hook JWT to be the sole credential.",
|
||||
mcp_server.server_name or mcp_server.name,
|
||||
)
|
||||
extra_headers.update(hook_extra_headers)
|
||||
hook_has_authorization: Final = any(k.lower() == "authorization" for k in hook_extra_headers)
|
||||
existing_has_authorization: Final = any(k.lower() == "authorization" for k in extra_headers)
|
||||
server_auth_occupies_authorization: Final = (
|
||||
any(k.lower() == "authorization" for k in server_auth_header)
|
||||
if isinstance(server_auth_header, dict)
|
||||
else server_auth_header is not None and mcp_server.auth_type != MCPAuth.api_key
|
||||
)
|
||||
if hook_has_authorization and (existing_has_authorization or server_auth_occupies_authorization):
|
||||
# Mirror the tools/list signer guard: an upstream credential (user OAuth,
|
||||
# static header, or configured authentication_token) already occupies the
|
||||
# Authorization slot, so the hook must not replace it.
|
||||
verbose_logger.warning(
|
||||
"MCPServerManager: dropping hook-injected 'Authorization' header for "
|
||||
"server '%s' because an upstream credential already occupies the "
|
||||
"Authorization slot; the existing credential is kept.",
|
||||
mcp_server.server_name or mcp_server.name,
|
||||
)
|
||||
extra_headers.update({k: v for k, v in hook_extra_headers.items() if k.lower() != "authorization"})
|
||||
else:
|
||||
extra_headers.update(hook_extra_headers)
|
||||
|
||||
# Reset to None if no headers were actually added
|
||||
if extra_headers is not None and len(extra_headers) == 0:
|
||||
|
|
|
|||
|
|
@ -5,7 +5,8 @@ Validates that:
|
|||
1. _convert_mcp_hook_response_to_kwargs extracts extra_headers from hook response
|
||||
2. pre_call_tool_check returns hook-provided extra_headers AND modified arguments
|
||||
3. call_tool flows hook headers and modified arguments downstream
|
||||
4. Hook-provided headers take highest priority (merge after static_headers)
|
||||
4. Hook-provided headers merge after static_headers, but a hook Authorization
|
||||
header never displaces an existing upstream Authorization credential
|
||||
5. OpenAPI-backed servers log a warning and continue (skip injection) when hook headers are present
|
||||
6. JWT claims are propagated in both standard and virtual-key fast paths
|
||||
7. Backward compatibility: hooks without extra_headers continue to work
|
||||
|
|
@ -487,8 +488,8 @@ class TestHookHeaderMergePriority:
|
|||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hook_headers_override_static_headers(self):
|
||||
"""Hook headers should take precedence over static_headers."""
|
||||
async def test_hook_authorization_does_not_override_static_authorization(self):
|
||||
"""A hook Authorization must not displace a static_headers Authorization (LIT-6321)."""
|
||||
manager = MCPServerManager()
|
||||
server = self._make_server(static_headers={"Authorization": "Bearer static-token", "X-Static": "yes"})
|
||||
|
||||
|
|
@ -521,7 +522,7 @@ class TestHookHeaderMergePriority:
|
|||
pass
|
||||
|
||||
headers = captured_extra_headers.get("value", {})
|
||||
assert headers["Authorization"] == "Bearer hook-signed-jwt"
|
||||
assert headers["Authorization"] == "Bearer static-token"
|
||||
assert headers["X-Static"] == "yes"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -560,8 +561,8 @@ class TestHookHeaderMergePriority:
|
|||
assert headers == {"X-Static": "static-value"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hook_headers_merge_with_oauth2(self):
|
||||
"""Hook headers merge on top of OAuth2 headers."""
|
||||
async def test_hook_authorization_does_not_override_oauth2_authorization(self):
|
||||
"""tools/call keeps the user's OAuth Authorization; only non-auth hook headers merge (LIT-6321)."""
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="test-id",
|
||||
|
|
@ -570,6 +571,8 @@ class TestHookHeaderMergePriority:
|
|||
url="https://example.com",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
oauth2_flow="authorization_code",
|
||||
delegate_auth_to_upstream=True,
|
||||
)
|
||||
|
||||
captured_extra_headers: Dict[str, Any] = {}
|
||||
|
|
@ -605,10 +608,245 @@ class TestHookHeaderMergePriority:
|
|||
pass
|
||||
|
||||
headers = captured_extra_headers.get("value", {})
|
||||
assert headers["Authorization"] == "Bearer hook-jwt"
|
||||
assert headers["Authorization"] == "Bearer oauth2-token"
|
||||
assert headers["X-OAuth"] == "yes"
|
||||
assert headers["X-Trace-Id"] == "trace-123"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hook_authorization_used_when_no_upstream_credential(self):
|
||||
"""With no upstream credential, the signer JWT is still injected."""
|
||||
manager = MCPServerManager()
|
||||
server = self._make_server()
|
||||
|
||||
captured_extra_headers: Dict[str, Optional[Dict[str, str]]] = {}
|
||||
|
||||
async def fake_create_mcp_client(server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs):
|
||||
captured_extra_headers["value"] = extra_headers
|
||||
mock_client = MagicMock()
|
||||
mock_client.call_tool = AsyncMock(return_value=MagicMock())
|
||||
return mock_client
|
||||
|
||||
with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client):
|
||||
with patch.object(manager, "_build_stdio_env", return_value=None):
|
||||
try:
|
||||
await manager._call_regular_mcp_tool(
|
||||
mcp_server=server,
|
||||
original_tool_name="test_tool",
|
||||
arguments={"key": "val"},
|
||||
tasks=[],
|
||||
mcp_auth_header=None,
|
||||
mcp_server_auth_headers=None,
|
||||
oauth2_headers=None,
|
||||
raw_headers=None,
|
||||
proxy_logging_obj=None,
|
||||
hook_extra_headers={"Authorization": "Bearer hook-jwt"},
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
headers = captured_extra_headers.get("value") or {}
|
||||
assert headers["Authorization"] == "Bearer hook-jwt"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hook_authorization_dropped_when_server_auth_header_present(self):
|
||||
"""With a configured authentication_token (auth_value), the hook Authorization is dropped."""
|
||||
manager = MCPServerManager()
|
||||
server = self._make_server()
|
||||
|
||||
captured: Dict[str, object] = {}
|
||||
|
||||
async def fake_create_mcp_client(server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs):
|
||||
captured["extra_headers"] = extra_headers
|
||||
captured["mcp_auth_header"] = mcp_auth_header
|
||||
mock_client = MagicMock()
|
||||
mock_client.call_tool = AsyncMock(return_value=MagicMock())
|
||||
return mock_client
|
||||
|
||||
with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client):
|
||||
with patch.object(manager, "_build_stdio_env", return_value=None):
|
||||
try:
|
||||
await manager._call_regular_mcp_tool(
|
||||
mcp_server=server,
|
||||
original_tool_name="test_tool",
|
||||
arguments={"key": "val"},
|
||||
tasks=[],
|
||||
mcp_auth_header="server-static-token",
|
||||
mcp_server_auth_headers=None,
|
||||
oauth2_headers=None,
|
||||
raw_headers=None,
|
||||
proxy_logging_obj=None,
|
||||
hook_extra_headers={
|
||||
"Authorization": "Bearer hook-jwt",
|
||||
"X-Trace-Id": "trace-123",
|
||||
},
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
headers = captured.get("extra_headers") or {}
|
||||
assert isinstance(headers, dict)
|
||||
assert "Authorization" not in headers
|
||||
assert headers.get("X-Trace-Id") == "trace-123"
|
||||
assert captured.get("mcp_auth_header") == "server-static-token"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hook_authorization_case_insensitive_conflict(self):
|
||||
"""Authorization conflicts are matched case-insensitively."""
|
||||
manager = MCPServerManager()
|
||||
server = self._make_server(static_headers={"authorization": "Bearer static-token"})
|
||||
|
||||
captured_extra_headers: Dict[str, Optional[Dict[str, str]]] = {}
|
||||
|
||||
async def fake_create_mcp_client(server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs):
|
||||
captured_extra_headers["value"] = extra_headers
|
||||
mock_client = MagicMock()
|
||||
mock_client.call_tool = AsyncMock(return_value=MagicMock())
|
||||
return mock_client
|
||||
|
||||
with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client):
|
||||
with patch.object(manager, "_build_stdio_env", return_value=None):
|
||||
try:
|
||||
await manager._call_regular_mcp_tool(
|
||||
mcp_server=server,
|
||||
original_tool_name="test_tool",
|
||||
arguments={"key": "val"},
|
||||
tasks=[],
|
||||
mcp_auth_header=None,
|
||||
mcp_server_auth_headers=None,
|
||||
oauth2_headers=None,
|
||||
raw_headers=None,
|
||||
proxy_logging_obj=None,
|
||||
hook_extra_headers={"Authorization": "Bearer hook-jwt"},
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
headers = captured_extra_headers.get("value") or {}
|
||||
assert headers.get("authorization") == "Bearer static-token"
|
||||
assert "Authorization" not in headers
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hook_authorization_kept_with_api_key_server_credential(self):
|
||||
"""An api_key credential maps to X-API-Key, so the hook Authorization is kept."""
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="test-id",
|
||||
name="Test Server",
|
||||
server_name="test_server",
|
||||
url="https://example.com",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.api_key,
|
||||
)
|
||||
|
||||
captured: Dict[str, object] = {}
|
||||
|
||||
async def fake_create_mcp_client(server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs):
|
||||
captured["extra_headers"] = extra_headers
|
||||
captured["mcp_auth_header"] = mcp_auth_header
|
||||
mock_client = MagicMock()
|
||||
mock_client.call_tool = AsyncMock(return_value=MagicMock())
|
||||
return mock_client
|
||||
|
||||
with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client):
|
||||
with patch.object(manager, "_build_stdio_env", return_value=None):
|
||||
try:
|
||||
await manager._call_regular_mcp_tool(
|
||||
mcp_server=server,
|
||||
original_tool_name="test_tool",
|
||||
arguments={"key": "val"},
|
||||
tasks=[],
|
||||
mcp_auth_header="server-api-key",
|
||||
mcp_server_auth_headers=None,
|
||||
oauth2_headers=None,
|
||||
raw_headers=None,
|
||||
proxy_logging_obj=None,
|
||||
hook_extra_headers={"Authorization": "Bearer hook-jwt"},
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
headers = captured.get("extra_headers") or {}
|
||||
assert isinstance(headers, dict)
|
||||
assert headers.get("Authorization") == "Bearer hook-jwt"
|
||||
assert captured.get("mcp_auth_header") == "server-api-key"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hook_authorization_kept_with_non_authorization_server_header_dict(self):
|
||||
"""A per-server header dict without Authorization does not block the hook JWT."""
|
||||
manager = MCPServerManager()
|
||||
server = self._make_server()
|
||||
|
||||
captured: Dict[str, object] = {}
|
||||
|
||||
async def fake_create_mcp_client(server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs):
|
||||
captured["extra_headers"] = extra_headers
|
||||
captured["mcp_auth_header"] = mcp_auth_header
|
||||
mock_client = MagicMock()
|
||||
mock_client.call_tool = AsyncMock(return_value=MagicMock())
|
||||
return mock_client
|
||||
|
||||
with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client):
|
||||
with patch.object(manager, "_build_stdio_env", return_value=None):
|
||||
try:
|
||||
await manager._call_regular_mcp_tool(
|
||||
mcp_server=server,
|
||||
original_tool_name="test_tool",
|
||||
arguments={"key": "val"},
|
||||
tasks=[],
|
||||
mcp_auth_header=None,
|
||||
mcp_server_auth_headers={"test_server": {"X-API-Key": "per-server-key"}},
|
||||
oauth2_headers=None,
|
||||
raw_headers=None,
|
||||
proxy_logging_obj=None,
|
||||
hook_extra_headers={"Authorization": "Bearer hook-jwt"},
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
headers = captured.get("extra_headers") or {}
|
||||
assert isinstance(headers, dict)
|
||||
assert headers.get("Authorization") == "Bearer hook-jwt"
|
||||
assert captured.get("mcp_auth_header") == {"X-API-Key": "per-server-key"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hook_authorization_dropped_with_authorization_server_header_dict(self):
|
||||
"""A per-server header dict carrying Authorization blocks the hook JWT."""
|
||||
manager = MCPServerManager()
|
||||
server = self._make_server()
|
||||
|
||||
captured: Dict[str, object] = {}
|
||||
|
||||
async def fake_create_mcp_client(server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs):
|
||||
captured["extra_headers"] = extra_headers
|
||||
captured["mcp_auth_header"] = mcp_auth_header
|
||||
mock_client = MagicMock()
|
||||
mock_client.call_tool = AsyncMock(return_value=MagicMock())
|
||||
return mock_client
|
||||
|
||||
with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client):
|
||||
with patch.object(manager, "_build_stdio_env", return_value=None):
|
||||
try:
|
||||
await manager._call_regular_mcp_tool(
|
||||
mcp_server=server,
|
||||
original_tool_name="test_tool",
|
||||
arguments={"key": "val"},
|
||||
tasks=[],
|
||||
mcp_auth_header=None,
|
||||
mcp_server_auth_headers={"test_server": {"authorization": "Bearer per-server-token"}},
|
||||
oauth2_headers=None,
|
||||
raw_headers=None,
|
||||
proxy_logging_obj=None,
|
||||
hook_extra_headers={"Authorization": "Bearer hook-jwt", "X-Trace-Id": "trace-123"},
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
headers = captured.get("extra_headers") or {}
|
||||
assert isinstance(headers, dict)
|
||||
assert "Authorization" not in headers
|
||||
assert headers.get("X-Trace-Id") == "trace-123"
|
||||
assert captured.get("mcp_auth_header") == {"authorization": "Bearer per-server-token"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_m2m_oauth2_does_not_forward_litellm_caller_authorization(self):
|
||||
"""M2M must not put caller Bearer (LiteLLM API key) into extra_headers (#23652)."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue