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:
devin-ai-integration[bot] 2026-08-27 12:44:36 -07:00 committed by GitHub
parent 452254963e
commit e16aa9f512
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 268 additions and 29 deletions

View file

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

View file

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