mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
feat(mcp/v2): graft bearer_token and thread the inbound token (fixes passthrough/exchange 401)
bearer_token: map it in to_server_spec to ApiKeyConfig(Authorization, Bearer, SharedKey), reusing the api_key arm; it stops deferring and produces v1-identical headers. Inbound token: every v2 op now extracts the caller token (the inherited _extract_bearer_token) and threads it as subject_token -> to_subject -> Subject.inbound_token. This fixes a latent bug: passthrough and token_exchange were already mapped in to_server_spec (so they ran v2, not v1) but no op ever passed subject_token, so inbound_token was always None and their arms rejected EVERY request with 401, even ones carrying a valid caller token. The token is always extracted and passed; only the passthrough and token_exchange arms read it, so m2m and the other modes ignore it and credential isolation holds. The mcp_auth_header override stays deferred to v1 (its own deprecation track); basic/authorization/token and OpenAPI also remain deferred. Validated: bearer_token parity (v2 headers byte-identical to v1); passthrough WITH a caller token now succeeds end-to-end against an in-process server (the fix), WITHOUT a token fails closed with 401 (correct); 11 manager + 28 bridge tests pass.
This commit is contained in:
parent
b28a636fdc
commit
6ccb5aab76
4 changed files with 105 additions and 8 deletions
|
|
@ -13,7 +13,7 @@ strangler scaffolding that shrinks to zero as modes migrate, then is deleted wit
|
|||
``_create_mcp_client``) and ``_v2_connection`` (the permanent resolve()+build seam that returns a
|
||||
configured ``UpstreamConnection``). ``super()`` is reserved for what v2 does not own yet:
|
||||
unmapped/misconfigured modes, OpenAPI tools (registry, S1.8), the per-request ``mcp_auth_header``
|
||||
override/inbound-token path, and the JWT-signer guardrail.
|
||||
override path, and the JWT-signer guardrail.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
|
@ -66,8 +66,8 @@ class MCPServerManagerV2(MCPServerManager):
|
|||
"""Whether this request falls back to v1 (super()) instead of the v2 egress path.
|
||||
|
||||
True for what v2 does not own yet: unmapped/misconfigured modes (to_server_spec is None,
|
||||
e.g. bearer_token), OpenAPI servers (registry, not a connection), the per-request
|
||||
mcp_auth_header override/inbound-token path, and the JWT-signer guardrail. Temporary
|
||||
e.g. basic/token), OpenAPI servers (registry, not a connection), the per-request
|
||||
mcp_auth_header override path, and the JWT-signer guardrail. Temporary
|
||||
strangler scaffolding: returns True for fewer cases as modes migrate, then is deleted
|
||||
(with _create_mcp_client) once nothing is left on v1.
|
||||
"""
|
||||
|
|
@ -209,6 +209,7 @@ class MCPServerManagerV2(MCPServerManager):
|
|||
user_api_key_auth,
|
||||
raw_headers=raw_headers,
|
||||
extra_headers=extra_headers,
|
||||
subject_token=self._extract_bearer_token(None, raw_headers),
|
||||
)
|
||||
if isinstance(conn, Error):
|
||||
self._egress_list_failure(server, conn.error)
|
||||
|
|
@ -234,7 +235,11 @@ class MCPServerManagerV2(MCPServerManager):
|
|||
server, mcp_auth_header, extra_headers, add_prefix, raw_headers
|
||||
)
|
||||
conn = await self._v2_connection(
|
||||
server, None, raw_headers=raw_headers, extra_headers=extra_headers
|
||||
server,
|
||||
None,
|
||||
raw_headers=raw_headers,
|
||||
extra_headers=extra_headers,
|
||||
subject_token=self._extract_bearer_token(None, raw_headers),
|
||||
)
|
||||
if isinstance(conn, Error):
|
||||
self._egress_list_failure(server, conn.error)
|
||||
|
|
@ -260,7 +265,11 @@ class MCPServerManagerV2(MCPServerManager):
|
|||
server, mcp_auth_header, extra_headers, add_prefix, raw_headers
|
||||
)
|
||||
conn = await self._v2_connection(
|
||||
server, None, raw_headers=raw_headers, extra_headers=extra_headers
|
||||
server,
|
||||
None,
|
||||
raw_headers=raw_headers,
|
||||
extra_headers=extra_headers,
|
||||
subject_token=self._extract_bearer_token(None, raw_headers),
|
||||
)
|
||||
if isinstance(conn, Error):
|
||||
self._egress_list_failure(server, conn.error)
|
||||
|
|
@ -286,7 +295,11 @@ class MCPServerManagerV2(MCPServerManager):
|
|||
server, mcp_auth_header, extra_headers, add_prefix, raw_headers
|
||||
)
|
||||
conn = await self._v2_connection(
|
||||
server, None, raw_headers=raw_headers, extra_headers=extra_headers
|
||||
server,
|
||||
None,
|
||||
raw_headers=raw_headers,
|
||||
extra_headers=extra_headers,
|
||||
subject_token=self._extract_bearer_token(None, raw_headers),
|
||||
)
|
||||
if isinstance(conn, Error):
|
||||
self._egress_list_failure(server, conn.error)
|
||||
|
|
@ -319,6 +332,7 @@ class MCPServerManagerV2(MCPServerManager):
|
|||
raw_headers=raw_headers,
|
||||
extra_headers=extra_headers,
|
||||
raise_on_missing_env=True,
|
||||
subject_token=self._extract_bearer_token(None, raw_headers),
|
||||
)
|
||||
if isinstance(conn, Error):
|
||||
self._egress_item_failure(server, conn.error)
|
||||
|
|
@ -353,6 +367,7 @@ class MCPServerManagerV2(MCPServerManager):
|
|||
raw_headers=raw_headers,
|
||||
extra_headers=extra_headers,
|
||||
raise_on_missing_env=True,
|
||||
subject_token=self._extract_bearer_token(None, raw_headers),
|
||||
)
|
||||
if isinstance(conn, Error):
|
||||
self._egress_item_failure(server, conn.error)
|
||||
|
|
@ -416,6 +431,7 @@ class MCPServerManagerV2(MCPServerManager):
|
|||
raw_headers=raw_headers,
|
||||
forward_caller_headers=True,
|
||||
raise_on_missing_env=True,
|
||||
subject_token=self._extract_bearer_token(oauth2_headers, raw_headers),
|
||||
)
|
||||
if isinstance(conn, Error):
|
||||
self._egress_item_failure(mcp_server, conn.error)
|
||||
|
|
|
|||
|
|
@ -146,6 +146,19 @@ def to_server_spec(server: MCPServer) -> Optional[ServerSpec]:
|
|||
key_source=SharedKey(value=SecretStr(token)),
|
||||
),
|
||||
)
|
||||
if server.auth_type == MCPAuth.bearer_token:
|
||||
token = server.authentication_token
|
||||
if not token:
|
||||
return None # bearer_token with no token: let v1 handle it (parity-safe)
|
||||
return ServerSpec(
|
||||
server_id=server.server_id,
|
||||
resource=resource,
|
||||
config=ApiKeyConfig(
|
||||
header_name="Authorization",
|
||||
value_prefix="Bearer",
|
||||
key_source=SharedKey(value=SecretStr(token)),
|
||||
),
|
||||
)
|
||||
if server.auth_type == MCPAuth.aws_sigv4:
|
||||
# Reuse the SigV4 config builder; the resolver's aws_sigv4 arm turns it into the botocore
|
||||
# signer (an httpx.Auth) that signs each upstream request.
|
||||
|
|
|
|||
|
|
@ -210,3 +210,60 @@ async def test_v2_override_calls_tool_via_upstream_connection(echo_server_url):
|
|||
user_api_key_auth=None,
|
||||
)
|
||||
assert any("echo: hi" in getattr(c, "text", "") for c in result.content)
|
||||
|
||||
|
||||
def _passthrough_server(url):
|
||||
return MCPServer(
|
||||
server_id="pt",
|
||||
name="pt",
|
||||
transport=MCPTransport.http,
|
||||
url=url,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
delegate_auth_to_upstream=True,
|
||||
client_id="cid",
|
||||
authorization_url="https://idp/auth",
|
||||
token_url="https://idp/token",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_v2_passthrough_forwards_inbound_token(echo_server_url):
|
||||
# Passthrough: the caller token is extracted and threaded as inbound_token, so the v2 call
|
||||
# reaches the upstream (the no-auth echo server ignores the forwarded bearer and serves it).
|
||||
manager = MCPServerManagerV2()
|
||||
result = await manager._open_and_call_tool(
|
||||
_passthrough_server(echo_server_url),
|
||||
"echo",
|
||||
{"text": "hi"},
|
||||
mcp_auth_header=None,
|
||||
mcp_server_auth_headers=None,
|
||||
oauth2_headers={"Authorization": "Bearer caller-token"},
|
||||
raw_headers=None,
|
||||
hook_extra_headers=None,
|
||||
host_progress_callback=None,
|
||||
user_api_key_auth=None,
|
||||
)
|
||||
assert any("echo: hi" in getattr(c, "text", "") for c in result.content)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_v2_passthrough_without_token_fails_closed(echo_server_url):
|
||||
# No caller token -> inbound_token is None -> the passthrough arm fails closed (401), surfaced
|
||||
# as MCPUpstreamAuthError. Before the inbound-token plumbing, the token was never threaded
|
||||
# through, so every passthrough call hit this path.
|
||||
from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError
|
||||
|
||||
manager = MCPServerManagerV2()
|
||||
with pytest.raises(MCPUpstreamAuthError):
|
||||
await manager._open_and_call_tool(
|
||||
_passthrough_server(echo_server_url),
|
||||
"echo",
|
||||
{"text": "hi"},
|
||||
mcp_auth_header=None,
|
||||
mcp_server_auth_headers=None,
|
||||
oauth2_headers=None,
|
||||
raw_headers=None,
|
||||
hook_extra_headers=None,
|
||||
host_progress_callback=None,
|
||||
user_api_key_auth=None,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -72,8 +72,19 @@ async def test_none_attaches_no_auth(v2_on):
|
|||
|
||||
|
||||
async def test_non_grafted_mode_defers_to_v1(v2_on):
|
||||
# bearer_token is not grafted yet -> v2 returns None so v1 handles it
|
||||
assert await resolve_v2_auth_value(_server(MCPAuth.bearer_token, "k")) is None
|
||||
# basic is not grafted yet -> v2 returns None so v1 handles it
|
||||
assert await resolve_v2_auth_value(_server(MCPAuth.basic, "k")) is None
|
||||
|
||||
|
||||
async def test_bearer_token_parity(v2_on):
|
||||
token = "up-secret"
|
||||
server = _server(MCPAuth.bearer_token, token)
|
||||
v2_value = await resolve_v2_auth_value(server)
|
||||
assert v2_value == {"Authorization": f"Bearer {token}"}
|
||||
# byte-identical to v1's final upstream headers
|
||||
assert _v1_headers(MCPAuth.bearer_token, token) == _v1_headers(
|
||||
MCPAuth.bearer_token, v2_value
|
||||
)
|
||||
|
||||
|
||||
async def test_api_key_without_token_defers_to_v1(v2_on):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue