mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(mcp): caller Authorization must not override the stored per-user OAuth token
A caller with a valid x-litellm-api-key could include their own "Authorization: Bearer <chosen>" header and have the proxy execute tools against that bearer instead of the user's stored OAuth credential. For a v2-migrated authorization_code server the caller's Authorization was seeded into extra_headers, and the graft's apply-if-absent then dropped the resolved per-user token in its favor. v1 prevented this by overwriting a stale client Authorization with the stored token; this restores that precedence on both egress paths (connect + call_tool). - _should_strip_caller_authorization: also strip for migrated per-user OAuth (authorization_code) servers - the v2 resolver injects the stored token, so a caller-forwarded Authorization must not be forwarded upstream. Delegate / pass-through (to_server_spec is None) keep forwarding the caller's bearer. - both seed sites (_prepare_mcp_server_headers, _call_regular_mcp_tool) drop only the Authorization from the caller's oauth2_headers (via _without_authorization), keeping any other forwarded header and any hook/static Authorization (which still wins, as in v1). - regression test for the call_tool path; updated the two tests that asserted the old (vulnerable) forwarding to assert the secure behavior. Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
c813b594b7
commit
596abecb09
4 changed files with 146 additions and 21 deletions
|
|
@ -190,6 +190,10 @@ def _should_strip_caller_authorization(
|
|||
Strip rules:
|
||||
- **M2M (client_credentials) servers**: never forward the caller's
|
||||
``Authorization`` — the proxy fetches its own upstream token.
|
||||
- **Migrated per-user OAuth (authorization_code) servers**: never forward
|
||||
the caller's ``Authorization`` — the v2 resolver injects the stored
|
||||
per-user token, so a caller-supplied bearer cannot override another
|
||||
user's stored credential. Delegate / pass-through keep forwarding it.
|
||||
- **OAuth pass-through servers**: strip when the ``Authorization``
|
||||
header is actually the LiteLLM API key — either because admission
|
||||
validated it (``user_api_key_auth.api_key`` is set) and the caller
|
||||
|
|
@ -202,6 +206,15 @@ def _should_strip_caller_authorization(
|
|||
"""
|
||||
if mcp_server.has_client_credentials:
|
||||
return True
|
||||
if (
|
||||
mcp_server.auth_type == MCPAuth.oauth2
|
||||
and to_server_spec(mcp_server) is not None
|
||||
):
|
||||
# Migrated per-user OAuth (authorization_code): the v2 resolver injects the
|
||||
# stored token, so a caller-forwarded Authorization must not be forwarded
|
||||
# upstream — it would override another user's stored credential. Delegate and
|
||||
# pass-through return None from to_server_spec and keep forwarding the bearer.
|
||||
return True
|
||||
if not mcp_server.is_oauth_passthrough:
|
||||
return False
|
||||
|
||||
|
|
@ -221,6 +234,18 @@ def _should_strip_caller_authorization(
|
|||
)
|
||||
|
||||
|
||||
def _without_authorization(
|
||||
headers: Optional[dict[str, str]],
|
||||
) -> Optional[dict[str, str]]:
|
||||
"""A copy of ``headers`` with any ``Authorization`` key removed (case-insensitive), or
|
||||
None if nothing remains. Drops only the credential, keeping other forwarded headers.
|
||||
"""
|
||||
if not headers:
|
||||
return None
|
||||
filtered = {k: v for k, v in headers.items() if k.lower() != "authorization"}
|
||||
return filtered or None
|
||||
|
||||
|
||||
def _extract_upstream_auth_failure(
|
||||
exc: BaseException,
|
||||
) -> Optional[Tuple[int, Optional[str]]]:
|
||||
|
|
@ -3509,6 +3534,16 @@ class MCPServerManager:
|
|||
extra_headers = None
|
||||
else:
|
||||
extra_headers = oauth2_headers
|
||||
# Migrated authorization_code: the v2 resolver injects the stored per-user
|
||||
# token, so drop the caller-forwarded Authorization (apply-if-absent would
|
||||
# otherwise let it shadow the resolved token). Delegate keeps it. Centralized
|
||||
# via _should_strip_caller_authorization to match _prepare_mcp_server_headers.
|
||||
if extra_headers and _should_strip_caller_authorization(
|
||||
mcp_server=mcp_server,
|
||||
raw_headers=raw_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
):
|
||||
extra_headers = _without_authorization(extra_headers)
|
||||
|
||||
if mcp_server.extra_headers and raw_headers:
|
||||
if extra_headers is None:
|
||||
|
|
|
|||
|
|
@ -283,6 +283,7 @@ if MCP_AVAILABLE:
|
|||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
MCPServerManager,
|
||||
_should_strip_caller_authorization,
|
||||
_without_authorization,
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
|
|
@ -1435,6 +1436,16 @@ if MCP_AVAILABLE:
|
|||
else:
|
||||
# Copy to avoid mutating the original dict (important for parallel fetching)
|
||||
extra_headers = oauth2_headers.copy() if oauth2_headers else None
|
||||
# Migrated authorization_code: the v2 resolver injects the stored per-user
|
||||
# token, so drop the caller-forwarded Authorization (apply-if-absent would
|
||||
# otherwise let it shadow the resolved token). Delegate keeps it. Centralized
|
||||
# via _should_strip_caller_authorization to match _call_regular_mcp_tool.
|
||||
if extra_headers and _should_strip_caller_authorization(
|
||||
mcp_server=server,
|
||||
raw_headers=raw_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
):
|
||||
extra_headers = _without_authorization(extra_headers)
|
||||
|
||||
if server.extra_headers and raw_headers:
|
||||
if extra_headers is None:
|
||||
|
|
|
|||
|
|
@ -284,8 +284,11 @@ def test_prepare_mcp_server_headers_oauth2_m2m_omits_litellm_caller_authorizatio
|
|||
assert extra_headers is None
|
||||
|
||||
|
||||
def test_prepare_mcp_server_headers_oauth2_interactive_copies_oauth2_headers():
|
||||
"""Interactive OAuth still forwards the user's OAuth token in extra_headers."""
|
||||
def test_prepare_mcp_server_headers_oauth2_interactive_drops_caller_authorization():
|
||||
"""A v2-migrated interactive OAuth (authorization_code) server must NOT forward the
|
||||
caller's Authorization: the resolver injects the stored per-user token, so a
|
||||
caller-supplied bearer must not override another user's stored credential. Non-auth
|
||||
headers are still carried; only the credential is dropped."""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
_prepare_mcp_server_headers,
|
||||
|
|
@ -293,7 +296,7 @@ def test_prepare_mcp_server_headers_oauth2_interactive_copies_oauth2_headers():
|
|||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
user_oauth = {"Authorization": "Bearer upstream-user-token"}
|
||||
caller_oauth = {"Authorization": "Bearer caller-supplied-token"}
|
||||
|
||||
server = MCPServer(
|
||||
server_id="3lo-server",
|
||||
|
|
@ -307,12 +310,13 @@ def test_prepare_mcp_server_headers_oauth2_interactive_copies_oauth2_headers():
|
|||
server=server,
|
||||
mcp_server_auth_headers=None,
|
||||
mcp_auth_header=None,
|
||||
oauth2_headers=user_oauth,
|
||||
oauth2_headers=caller_oauth,
|
||||
raw_headers=None,
|
||||
)
|
||||
|
||||
assert server_auth_header is None
|
||||
assert extra_headers == user_oauth
|
||||
# Caller's Authorization is dropped (only key present) -> extra_headers is None.
|
||||
assert extra_headers is None
|
||||
|
||||
|
||||
def test_prepare_mcp_server_headers_m2m_skips_authorization_from_raw_extra_headers():
|
||||
|
|
@ -2813,8 +2817,10 @@ async def test_mcp_routing_with_conflicting_alias_and_group_name():
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.no_parallel
|
||||
async def test_oauth2_headers_passed_to_mcp_client():
|
||||
"""Test that OAuth2 headers are properly passed through to the MCP client for OAuth2 servers like github_mcp"""
|
||||
async def test_oauth2_caller_headers_not_forwarded_for_migrated_server():
|
||||
"""A v2-migrated authorization_code server (like github_mcp) must NOT forward the
|
||||
caller's oauth2 Authorization to the MCP client — the resolver injects the stored
|
||||
per-user token, so a caller-supplied bearer cannot override another user's credential."""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
|
|
@ -2928,20 +2934,13 @@ async def test_oauth2_headers_passed_to_mcp_client():
|
|||
assert captured_client_args["server"].server_id == oauth2_server.server_id
|
||||
assert captured_client_args["server"].auth_type == MCPAuth.oauth2
|
||||
|
||||
# Most importantly: verify that OAuth2 headers were passed as extra_headers
|
||||
assert (
|
||||
captured_client_args["extra_headers"] is not None
|
||||
), "Expected extra_headers to be passed for OAuth2 server"
|
||||
assert (
|
||||
captured_client_args["extra_headers"] == oauth2_headers
|
||||
), f"Expected OAuth2 headers to be passed as extra_headers, got {captured_client_args['extra_headers']}"
|
||||
|
||||
# Verify the Authorization header specifically
|
||||
assert "Authorization" in captured_client_args["extra_headers"]
|
||||
assert (
|
||||
captured_client_args["extra_headers"]["Authorization"]
|
||||
== "Bearer github_oauth_token_12345"
|
||||
)
|
||||
# Security: a v2-migrated authorization_code server must NOT forward the caller's
|
||||
# oauth2 Authorization upstream. The v2 resolver injects the stored per-user token,
|
||||
# so a caller-supplied bearer cannot override another user's stored credential.
|
||||
extra_headers = captured_client_args["extra_headers"]
|
||||
assert extra_headers is None or "Authorization" not in {
|
||||
k.lower() for k in extra_headers
|
||||
}, f"Caller Authorization must not be forwarded, got {extra_headers}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -31,6 +31,8 @@ from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
|||
_deserialize_json_dict,
|
||||
_deserialize_json_list,
|
||||
_normalize_mcp_server_cost_info,
|
||||
_should_strip_caller_authorization,
|
||||
_without_authorization,
|
||||
)
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_MCPServerTable,
|
||||
|
|
@ -600,6 +602,84 @@ class TestMCPServerManager:
|
|||
|
||||
assert captured_extra_headers == {"x-request-id": "req-123"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_regular_mcp_tool_v2_authz_code_drops_caller_authorization(
|
||||
self,
|
||||
):
|
||||
"""A v2-migrated per-user OAuth (authorization_code) server must NOT seed a
|
||||
caller-forwarded Authorization into extra_headers — the resolver injects the
|
||||
stored per-user token, and apply-if-absent would otherwise let the caller's
|
||||
header override another user's stored credential (matches v1's overwrite)."""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
manager = MCPServerManager()
|
||||
# oauth2, not M2M, not delegate => to_server_spec maps it to AuthorizationCodeConfig
|
||||
server = MCPServer(
|
||||
server_id="server-authz-code-call",
|
||||
name="authz-code-server",
|
||||
url="https://example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
)
|
||||
# Migrated authorization_code => the centralized strip decision says drop the
|
||||
# caller's Authorization (the v2 resolver injects the stored token).
|
||||
assert (
|
||||
_should_strip_caller_authorization(
|
||||
mcp_server=server, raw_headers=None, user_api_key_auth=None
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.call_tool = AsyncMock(
|
||||
return_value=CallToolResult(content=[], isError=False)
|
||||
)
|
||||
captured_extra_headers = "unset"
|
||||
|
||||
async def capture_create_mcp_client(
|
||||
server,
|
||||
mcp_auth_header,
|
||||
extra_headers,
|
||||
stdio_env,
|
||||
subject_token=None,
|
||||
**kwargs,
|
||||
): # pragma: no cover - helper
|
||||
nonlocal captured_extra_headers
|
||||
captured_extra_headers = extra_headers
|
||||
return mock_client
|
||||
|
||||
manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client)
|
||||
|
||||
await manager._call_regular_mcp_tool(
|
||||
mcp_server=server,
|
||||
original_tool_name="tool",
|
||||
arguments={},
|
||||
tasks=[],
|
||||
mcp_auth_header=None,
|
||||
mcp_server_auth_headers=None,
|
||||
oauth2_headers={"Authorization": "Bearer caller-supplied-token"},
|
||||
raw_headers={"authorization": "Bearer caller-supplied-token"},
|
||||
proxy_logging_obj=None,
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-key"),
|
||||
)
|
||||
|
||||
# The caller's Authorization must not reach extra_headers; the v2 resolver is
|
||||
# the sole Authorization source for this server.
|
||||
assert captured_extra_headers != "unset"
|
||||
if captured_extra_headers:
|
||||
assert "authorization" not in {k.lower() for k in captured_extra_headers}
|
||||
|
||||
def test_without_authorization_drops_only_the_credential(self):
|
||||
# None / empty -> None
|
||||
assert _without_authorization(None) is None
|
||||
assert _without_authorization({}) is None
|
||||
# Only Authorization present -> nothing left -> None (case-insensitive)
|
||||
assert _without_authorization({"authorization": "Bearer x"}) is None
|
||||
# Authorization dropped, other headers kept
|
||||
assert _without_authorization(
|
||||
{"Authorization": "Bearer x", "X-Trace-Id": "t"}
|
||||
) == {"X-Trace-Id": "t"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_regular_mcp_tool_passthrough_forwards_authorization_with_admission_header(
|
||||
self,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue