mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
feat(mcp): add true_passthrough and oauth_delegate auth modes
Introduce two first-class MCP server auth_type values that make LiteLLM's role in upstream authentication explicit, added alongside the existing delegate_auth_to_upstream / oauth_passthrough flags without changing their behavior. true_passthrough is a transparent proxy: LiteLLM performs no admission auth, requires no x-litellm-api-key, mints/stores/refreshes nothing, and forwards the client's Authorization to the upstream exactly as received. oauth_delegate keeps normal LiteLLM admission (x-litellm-api-key / SSO / JWT) and then forwards the client's separate upstream Authorization unchanged; the admission credential is never forwarded upstream. Both modes forward the caller's token via the existing extra_headers path and defer egress credential resolution to v1 (the v2 to_server_spec returns None for them). Upstream 401/403 responses are surfaced rather than swallowed so upstream OAuth challenges are preserved. Servers in either mode require per-user auth, so userless health checks are skipped.
This commit is contained in:
parent
db2402754a
commit
50c90281f4
12 changed files with 676 additions and 446 deletions
|
|
@ -220,6 +220,12 @@ class MCPRequestHandler:
|
|||
# when EVERY target is auth_type=oauth2 with delegate_auth_to_upstream
|
||||
# set; fails closed otherwise.
|
||||
validated_user_api_key_auth = UserAPIKeyAuth()
|
||||
elif MCPRequestHandler._target_servers_are_true_passthrough(
|
||||
path=request_route,
|
||||
mcp_servers=mcp_servers,
|
||||
client_ip=IPAddressUtils.get_mcp_client_ip(request),
|
||||
):
|
||||
validated_user_api_key_auth = UserAPIKeyAuth()
|
||||
elif oauth2_headers:
|
||||
# Authorization on a non-delegated server: the bearer must be a real
|
||||
# LiteLLM credential, so a failed validation is a genuine 401/403 and
|
||||
|
|
@ -399,6 +405,33 @@ class MCPRequestHandler:
|
|||
return False
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def _target_servers_are_true_passthrough(
|
||||
path: str, mcp_servers: Optional[list[str]], client_ip: Optional[str]
|
||||
) -> bool:
|
||||
"""
|
||||
True only when EVERY MCP server the request targets is ``auth_type == true_passthrough``.
|
||||
Fails closed when any target does not opt in or cannot be resolved.
|
||||
|
||||
Used by :meth:`process_mcp_request` to skip LiteLLM admission auth entirely: the gateway is a
|
||||
transparent proxy and the caller's ``Authorization`` is an upstream token, never a LiteLLM key.
|
||||
Mirrors :meth:`_target_servers_delegate_auth_to_upstream`; a mixed-target request keeps normal auth.
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.types.mcp import MCPAuth
|
||||
|
||||
target_names = MCPRequestHandler._resolve_target_server_names(path=path, mcp_servers_header=mcp_servers)
|
||||
if not target_names:
|
||||
return False
|
||||
|
||||
for name in target_names:
|
||||
server = global_mcp_server_manager.get_mcp_server_by_name(name, client_ip=client_ip)
|
||||
if server is None or server.auth_type != MCPAuth.true_passthrough:
|
||||
return False
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def _resolve_target_server_names(path: str, mcp_servers_header: Optional[List[str]]) -> List[str]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -78,6 +78,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
AuthorizationCodeConfig,
|
||||
PassthroughConfig,
|
||||
ServerSpec,
|
||||
TokenExchangeConfig,
|
||||
)
|
||||
|
|
@ -213,6 +214,13 @@ def _should_strip_caller_authorization(
|
|||
pass-through cold-start case (RFC 9728) the bearer in
|
||||
``Authorization`` is the upstream OAuth token and must be
|
||||
forwarded, so we keep it.
|
||||
- **oauth_delegate servers**: admission always runs and there is no
|
||||
anonymous path, so the caller's separate ``Authorization`` is
|
||||
forwarded only when a distinct ``x-litellm-api-key`` carried
|
||||
admission. Without that header the ``Authorization`` *was* the
|
||||
admission credential — a virtual key, an IdP JWT, or an SSO / OIDC /
|
||||
session token whose ``api_key`` is ``None`` — and must never reach
|
||||
the upstream, so it is stripped regardless of the ``api_key`` value.
|
||||
"""
|
||||
if mcp_server.auth_type == MCPAuth.oauth2_token_exchange:
|
||||
# OBO: the inbound Authorization is the subject token. It is exchanged at the IdP and only the
|
||||
|
|
@ -226,11 +234,13 @@ def _should_strip_caller_authorization(
|
|||
# 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:
|
||||
if not (mcp_server.is_oauth_passthrough or mcp_server.is_oauth_delegate):
|
||||
return False
|
||||
|
||||
normalized_raw_headers = {str(k).lower(): v for k, v in (raw_headers or {}).items() if isinstance(k, str)}
|
||||
has_explicit_litellm_admission_header = normalized_raw_headers.get("x-litellm-api-key") is not None
|
||||
if mcp_server.is_oauth_delegate:
|
||||
return not has_explicit_litellm_admission_header
|
||||
admission_consumed_authorization_as_litellm_key = (
|
||||
user_api_key_auth is not None
|
||||
and bool(getattr(user_api_key_auth, "api_key", None))
|
||||
|
|
@ -323,6 +333,41 @@ async def _resolve_byok_mcp_auth_header(
|
|||
return mcp_auth_header
|
||||
|
||||
|
||||
def _client_forwarded_authorization_headers(
|
||||
mcp_server: MCPServer,
|
||||
oauth2_headers: Optional[dict[str, str]],
|
||||
raw_headers: Optional[dict[str, str]],
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth],
|
||||
) -> Optional[dict[str, str]]:
|
||||
"""Egress headers for the client-forwarded-token modes (``true_passthrough`` / ``oauth_delegate``).
|
||||
|
||||
Forwards the caller's ``Authorization`` to the upstream, stripped when
|
||||
``_should_strip_caller_authorization`` says it was consumed as the LiteLLM admission key. Shared by
|
||||
``_call_regular_mcp_tool`` and ``server.py``'s ``_prepare_mcp_server_headers`` so the two egress
|
||||
paths cannot drift, mirroring the ``_should_strip_caller_authorization`` split.
|
||||
"""
|
||||
extra_headers = oauth2_headers.copy() if oauth2_headers else None
|
||||
if extra_headers and _should_strip_caller_authorization(
|
||||
mcp_server=mcp_server,
|
||||
raw_headers=raw_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
):
|
||||
return _without_authorization(extra_headers)
|
||||
return extra_headers
|
||||
|
||||
|
||||
def _take_forwarded_authorization(
|
||||
headers: Optional[dict[str, str]],
|
||||
) -> tuple[Optional[str], Optional[dict[str, str]]]:
|
||||
"""Pop the ``Authorization`` value out of ``headers`` (case-insensitive), returning it with the
|
||||
remaining headers, so the passthrough resolver arm is the single Authorization source rather than
|
||||
the header also riding in ``extra_headers`` (which the resolved auth would then defer to)."""
|
||||
if not headers:
|
||||
return None, headers
|
||||
value = next((v for k, v in headers.items() if k.lower() == "authorization"), None)
|
||||
return value, _without_authorization(headers)
|
||||
|
||||
|
||||
def _extract_upstream_auth_failure(
|
||||
exc: BaseException,
|
||||
) -> Optional[tuple[int, Optional[str]]]:
|
||||
|
|
@ -1623,15 +1668,18 @@ class MCPServerManager:
|
|||
delegate_server_ids = [
|
||||
server.server_id
|
||||
for server in self.get_registry().values()
|
||||
if getattr(server, "auth_type", None) == MCPAuth.oauth2
|
||||
and getattr(server, "delegate_auth_to_upstream", False) is True
|
||||
# M2M servers must not be exposed anonymously: an
|
||||
# unauthenticated caller would get LiteLLM to proxy tool
|
||||
# calls using its stored client_credentials. Resolve the flow
|
||||
# rather than reading has_client_credentials so an unstamped
|
||||
# M2M-shape row (null column, verbatim-read as non-M2M) still
|
||||
# fails closed here, matching the anonymous-delegate auth gate.
|
||||
and MCPServerManager.effective_oauth2_flow(server) != "client_credentials"
|
||||
if (
|
||||
getattr(server, "auth_type", None) == MCPAuth.oauth2
|
||||
and getattr(server, "delegate_auth_to_upstream", False) is True
|
||||
# M2M servers must not be exposed anonymously: an
|
||||
# unauthenticated caller would get LiteLLM to proxy tool
|
||||
# calls using its stored client_credentials. Resolve the flow
|
||||
# rather than reading has_client_credentials so an unstamped
|
||||
# M2M-shape row (null column, verbatim-read as non-M2M) still
|
||||
# fails closed here, matching the anonymous-delegate auth gate.
|
||||
and MCPServerManager.effective_oauth2_flow(server) != "client_credentials"
|
||||
)
|
||||
or getattr(server, "auth_type", None) == MCPAuth.true_passthrough
|
||||
]
|
||||
combined_servers.update(delegate_server_ids)
|
||||
|
||||
|
|
@ -2229,16 +2277,17 @@ class MCPServerManager:
|
|||
spec = None if transport == MCPTransport.stdio else to_server_spec(server)
|
||||
provider = cred_provider or self._cred_provider
|
||||
# A caller-supplied per-request override (mcp_auth_header / x-mcp-*) defers to the v1 path
|
||||
# so it wins - except for the per-user modes the v2 resolver owns (authorization_code's
|
||||
# stored token and token_exchange's RFC 8693 minted token). A caller must not be able to
|
||||
# substitute another user's stored credential, nor silently disable the OBO exchange and
|
||||
# forward an arbitrary bearer upstream, so we keep the v2 spec and ignore the override for
|
||||
# both; the REST tools preview supplies its not-yet-persisted token through the resolver
|
||||
# (cred_provider), never this path.
|
||||
# so it wins - except for the modes the v2 resolver owns per-caller (authorization_code's
|
||||
# stored token, token_exchange's RFC 8693 minted token, and the passthrough modes'
|
||||
# forwarded caller token). A caller must not be able to substitute another user's stored
|
||||
# credential, nor silently disable the OBO exchange and forward an arbitrary bearer
|
||||
# upstream, so we keep the v2 spec and ignore the override for these; the REST tools
|
||||
# preview supplies its not-yet-persisted token through the resolver (cred_provider),
|
||||
# never this path.
|
||||
if (
|
||||
spec is not None
|
||||
and mcp_auth_header
|
||||
and not isinstance(spec.config, (AuthorizationCodeConfig, TokenExchangeConfig))
|
||||
and not isinstance(spec.config, (AuthorizationCodeConfig, PassthroughConfig, TokenExchangeConfig))
|
||||
):
|
||||
spec = None
|
||||
auth_value = (
|
||||
|
|
@ -2305,11 +2354,14 @@ class MCPServerManager:
|
|||
server_url = server.url or ""
|
||||
|
||||
if spec is not None:
|
||||
inbound_token = subject_token
|
||||
if isinstance(spec.config, PassthroughConfig):
|
||||
inbound_token, extra_headers = _take_forwarded_authorization(extra_headers)
|
||||
resolved_auth, extra_headers = await self._resolve_v2_auth(
|
||||
server=server,
|
||||
spec=spec,
|
||||
provider=provider,
|
||||
subject_token=subject_token,
|
||||
subject_token=inbound_token,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
extra_headers=extra_headers,
|
||||
)
|
||||
|
|
@ -3730,6 +3782,13 @@ class MCPServerManager:
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
):
|
||||
extra_headers = _without_authorization(extra_headers)
|
||||
elif mcp_server.is_true_passthrough or mcp_server.is_oauth_delegate:
|
||||
extra_headers = _client_forwarded_authorization_headers(
|
||||
mcp_server=mcp_server,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
if mcp_server.extra_headers and raw_headers:
|
||||
if extra_headers is None:
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
|||
AuthorizationCodeConfig,
|
||||
CredError,
|
||||
NoneConfig,
|
||||
PassthroughConfig,
|
||||
ServerSpec,
|
||||
SharedKey,
|
||||
Subject,
|
||||
|
|
@ -62,9 +63,10 @@ def to_server_spec(server: MCPServer) -> Optional[ServerSpec]:
|
|||
an ``assert_never`` tail, so a newly added auth mode fails the type gate here until it is
|
||||
explicitly mapped or explicitly deferred, rather than silently falling through to v1. Live
|
||||
modes: ``none``, the static-header family (``api_key`` plus the Authorization schemes,
|
||||
all shared-key), ``oauth2`` per-user tokens (``authorization_code``), and
|
||||
``oauth2_token_exchange`` (OBO); client_credentials (M2M), delegated/passthrough
|
||||
oauth2, and SigV4 return None and stay on v1.
|
||||
all shared-key), ``oauth2`` per-user tokens (``authorization_code``), ``oauth2_token_exchange``
|
||||
(OBO), and the client-forwarded token modes ``true_passthrough`` / ``oauth_delegate``
|
||||
(``PassthroughConfig``); client_credentials (M2M), delegated/passthrough oauth2, and SigV4
|
||||
return None and stay on v1.
|
||||
"""
|
||||
if server.is_byok:
|
||||
return None # per-user BYOK source not migrated yet -> defer to v1 (any auth_type)
|
||||
|
|
@ -94,6 +96,8 @@ def to_server_spec(server: MCPServer) -> Optional[ServerSpec]:
|
|||
)
|
||||
# client_credentials (M2M) and delegate/passthrough oauth2 stay on v1
|
||||
return None
|
||||
case MCPAuth.true_passthrough | MCPAuth.oauth_delegate:
|
||||
return ServerSpec(server_id=server.server_id, resource=resource, config=PassthroughConfig())
|
||||
case MCPAuth.oauth2_token_exchange:
|
||||
return _token_exchange_spec(server, resource)
|
||||
case MCPAuth.aws_sigv4:
|
||||
|
|
|
|||
|
|
@ -7,10 +7,11 @@ no precedence cascade. It is wildcard-free with an `assert_never` tail, so addin
|
|||
an arm fails the type gate (basedpyright `reportMatchNotExhaustive`); a bypassed gate fails loudly
|
||||
at runtime instead of returning `None`.
|
||||
|
||||
`none` and `api_key` (shared-key source) are live, as is `authorization_code`, which reads the
|
||||
user's token from the injected `OAuthTokenStore`, and `token_exchange`, which swaps the caller's
|
||||
inbound token through the injected `TokenExchanger`. The remaining arms are `not_implemented` stubs
|
||||
that each land in a follow-up PR with their seam. Pure v2: no imports from v1.
|
||||
`none`, `api_key` (shared-key source), and `passthrough` (forwards the caller's own inbound token)
|
||||
are live, as is `authorization_code`, which reads the user's token from the injected
|
||||
`OAuthTokenStore`, and `token_exchange`, which swaps the caller's inbound token through the
|
||||
injected `TokenExchanger`. The remaining arms are `not_implemented` stubs that each land in a
|
||||
follow-up PR with their seam. Pure v2: no imports from v1.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
|
@ -97,7 +98,7 @@ class UpstreamCredentialProvider:
|
|||
case ApiKeyConfig() as config:
|
||||
return self._api_key(config)
|
||||
case PassthroughConfig():
|
||||
return _not_implemented(AuthSpecKind.passthrough)
|
||||
return self._passthrough(subject)
|
||||
case ClientCredentialsConfig():
|
||||
return _not_implemented(AuthSpecKind.client_credentials)
|
||||
case TokenExchangeConfig() as config:
|
||||
|
|
@ -118,6 +119,18 @@ class UpstreamCredentialProvider:
|
|||
"""
|
||||
return await self._authz_token(subject, server) is not None
|
||||
|
||||
def _passthrough(self, subject: Subject) -> Result[httpx.Auth, CredError]:
|
||||
"""Forward the caller's own upstream credential verbatim; the gateway mints nothing.
|
||||
|
||||
The inbound token is the caller's already-disambiguated ``Authorization`` (never the LiteLLM
|
||||
admission credential; the edge adapter drops that before building the ``Subject``). When it is
|
||||
absent the request is sent unauthenticated so the upstream's own 401 surfaces, rather than the
|
||||
gateway challenging on the upstream's behalf.
|
||||
"""
|
||||
if subject.inbound_token is None:
|
||||
return Ok(NoOpAuth())
|
||||
return Ok(StaticHeaderAuth(subject.inbound_token.get_secret_value(), header_name="Authorization"))
|
||||
|
||||
def _api_key(self, config: ApiKeyConfig) -> Result[httpx.Auth, CredError]:
|
||||
match config.key_source:
|
||||
case SharedKey() as source:
|
||||
|
|
|
|||
|
|
@ -331,6 +331,7 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
MCPServerManager,
|
||||
_client_forwarded_authorization_headers,
|
||||
_should_strip_caller_authorization,
|
||||
_without_authorization,
|
||||
global_mcp_server_manager,
|
||||
|
|
@ -1561,6 +1562,13 @@ if MCP_AVAILABLE:
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
):
|
||||
extra_headers = _without_authorization(extra_headers)
|
||||
elif server.is_true_passthrough or server.is_oauth_delegate:
|
||||
extra_headers = _client_forwarded_authorization_headers(
|
||||
mcp_server=server,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
if server.extra_headers and raw_headers:
|
||||
if extra_headers is None:
|
||||
|
|
|
|||
|
|
@ -38,6 +38,8 @@ class MCPAuth(str, enum.Enum):
|
|||
aws_sigv4 = "aws_sigv4"
|
||||
token = "token"
|
||||
oauth2_token_exchange = "oauth2_token_exchange"
|
||||
true_passthrough = "true_passthrough"
|
||||
oauth_delegate = "oauth_delegate"
|
||||
|
||||
|
||||
# RFC 8693 default subject_token_type. A NULL column / omitted config key means
|
||||
|
|
@ -60,6 +62,8 @@ MCPAuthType = Optional[
|
|||
MCPAuth.aws_sigv4,
|
||||
MCPAuth.token,
|
||||
MCPAuth.oauth2_token_exchange,
|
||||
MCPAuth.true_passthrough,
|
||||
MCPAuth.oauth_delegate,
|
||||
]
|
||||
]
|
||||
|
||||
|
|
|
|||
|
|
@ -152,6 +152,18 @@ class MCPServer(BaseModel):
|
|||
"""True if this is an OAuth2 server that relies on per-user tokens (no client_credentials)."""
|
||||
return self.auth_type == MCPAuth.oauth2 and not self.has_client_credentials
|
||||
|
||||
@property
|
||||
def is_true_passthrough(self) -> bool:
|
||||
"""True for the transparent-proxy mode: LiteLLM performs no admission auth and forwards the
|
||||
client's ``Authorization`` to the upstream unchanged."""
|
||||
return self.auth_type == MCPAuth.true_passthrough
|
||||
|
||||
@property
|
||||
def is_oauth_delegate(self) -> bool:
|
||||
"""True for the delegated-upstream-OAuth mode: LiteLLM still admits the caller (API key / SSO /
|
||||
JWT) but forwards the caller's separate upstream ``Authorization`` unchanged, minting nothing."""
|
||||
return self.auth_type == MCPAuth.oauth_delegate
|
||||
|
||||
@property
|
||||
def requires_per_user_auth(self) -> bool:
|
||||
"""
|
||||
|
|
@ -167,6 +179,9 @@ class MCPServer(BaseModel):
|
|||
if self.needs_user_oauth_token:
|
||||
return True
|
||||
|
||||
if self.is_true_passthrough or self.is_oauth_delegate:
|
||||
return True
|
||||
|
||||
# PAT passthrough: auth_type is none but extra_headers includes auth headers
|
||||
if self.auth_type == MCPAuth.none and self.extra_headers:
|
||||
auth_header_names = {"authorization", "x-api-key", "api-key", "apikey"}
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -23,6 +23,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
|||
AuthorizationCodeConfig,
|
||||
CredError,
|
||||
NoneConfig,
|
||||
PassthroughConfig,
|
||||
SharedKey,
|
||||
TokenExchangeConfig,
|
||||
)
|
||||
|
|
@ -234,6 +235,12 @@ def test_token_exchange_empty_subject_token_type_normalizes_to_default():
|
|||
assert spec.config.subject_token_type == "urn:ietf:params:oauth:token-type:access_token"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("auth_type", [MCPAuth.true_passthrough, MCPAuth.oauth_delegate])
|
||||
def test_client_forwarded_modes_map_to_passthrough_config(auth_type):
|
||||
spec = to_server_spec(_server(auth_type=auth_type))
|
||||
assert spec is not None and isinstance(spec.config, PassthroughConfig)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"server",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
"""Tests for the resolver dispatch: live arms produce auth, stubbed arms fail closed.
|
||||
|
||||
`none`, `api_key` (shared-key source), `authorization_code`, and `token_exchange` are implemented;
|
||||
every other arm, plus the `api_key` BYOK source, returns a typed `not_implemented` error until its
|
||||
mode lands. Parametrizing the stubs over one config each also guards reachability: a dropped `case`
|
||||
would hit `assert_never` and raise instead of returning the stub.
|
||||
`none`, `api_key` (shared-key source), `passthrough`, `authorization_code`, and `token_exchange` are
|
||||
implemented; every other arm, plus the `api_key` BYOK source, returns a typed `not_implemented` error
|
||||
until its mode lands. Parametrizing the stubs over one config each also guards reachability: a dropped
|
||||
`case` would hit `assert_never` and raise instead of returning the stub.
|
||||
"""
|
||||
|
||||
import httpx
|
||||
|
|
@ -253,9 +253,24 @@ async def test_token_exchange_without_an_exchanger_fails_closed():
|
|||
assert result.error.tag == "misconfigured"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_passthrough_forwards_the_inbound_token_verbatim():
|
||||
subject = Subject(tenant_id="", subject_id="", inbound_token=SecretStr("Bearer upstream-xyz"))
|
||||
result = await UpstreamCredentialProvider().resolve_credentials(subject, _spec(PassthroughConfig()))
|
||||
assert isinstance(result, Ok)
|
||||
assert isinstance(result.ok, StaticHeaderAuth)
|
||||
assert _emitted(result.ok)["Authorization"] == "Bearer upstream-xyz"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_passthrough_without_inbound_token_is_a_no_op():
|
||||
result = await UpstreamCredentialProvider().resolve_credentials(_SUBJECT, _spec(PassthroughConfig()))
|
||||
assert isinstance(result, Ok)
|
||||
assert isinstance(result.ok, NoOpAuth)
|
||||
|
||||
|
||||
_STUBBED = [
|
||||
("api_key_byok", ApiKeyConfig(key_source=Byok())),
|
||||
("passthrough", PassthroughConfig()),
|
||||
("client_credentials", ClientCredentialsConfig()),
|
||||
("aws_sigv4", AwsSigV4Config(region="us-east-1")),
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1267,6 +1267,229 @@ class TestMCPServerManager:
|
|||
|
||||
assert captured_extra_headers == {"Authorization": "Bearer upstream-oauth-bearer"}
|
||||
|
||||
async def _capture_call_extra_headers(self, server, oauth2_headers, raw_headers, user_api_key_auth):
|
||||
manager = MCPServerManager()
|
||||
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
|
||||
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=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
proxy_logging_obj=None,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
return captured["extra_headers"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_regular_mcp_tool_true_passthrough_forwards_authorization(self):
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
server = MCPServer(
|
||||
server_id="server-true-passthrough",
|
||||
name="tp-server",
|
||||
url="https://example.com",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.true_passthrough,
|
||||
)
|
||||
extra_headers = await self._capture_call_extra_headers(
|
||||
server,
|
||||
oauth2_headers={"Authorization": "Bearer upstream-token"},
|
||||
raw_headers={"authorization": "Bearer upstream-token"},
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key=None),
|
||||
)
|
||||
assert extra_headers == {"Authorization": "Bearer upstream-token"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_regular_mcp_tool_oauth_delegate_forwards_separate_authorization(self):
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
server = MCPServer(
|
||||
server_id="server-oauth-delegate",
|
||||
name="od-server",
|
||||
url="https://example.com",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth_delegate,
|
||||
)
|
||||
extra_headers = await self._capture_call_extra_headers(
|
||||
server,
|
||||
oauth2_headers={"Authorization": "Bearer upstream-token"},
|
||||
raw_headers={
|
||||
"x-litellm-api-key": "Bearer sk-litellm-key",
|
||||
"authorization": "Bearer upstream-token",
|
||||
},
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-key"),
|
||||
)
|
||||
assert extra_headers == {"Authorization": "Bearer upstream-token"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_regular_mcp_tool_oauth_delegate_never_forwards_admission_key(self):
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
server = MCPServer(
|
||||
server_id="server-oauth-delegate-leak",
|
||||
name="od-server",
|
||||
url="https://example.com",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth_delegate,
|
||||
)
|
||||
extra_headers = await self._capture_call_extra_headers(
|
||||
server,
|
||||
oauth2_headers={"Authorization": "Bearer sk-litellm-key"},
|
||||
raw_headers={"authorization": "Bearer sk-litellm-key"},
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-key"),
|
||||
)
|
||||
assert not extra_headers or "authorization" not in {k.lower() for k in extra_headers}
|
||||
|
||||
def test_should_strip_caller_authorization_new_modes(self):
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
true_passthrough = MCPServer(
|
||||
server_id="tp",
|
||||
name="tp",
|
||||
url="https://example.com",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.true_passthrough,
|
||||
)
|
||||
assert (
|
||||
_should_strip_caller_authorization(
|
||||
mcp_server=true_passthrough,
|
||||
raw_headers={"authorization": "Bearer upstream"},
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key=None),
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
oauth_delegate = MCPServer(
|
||||
server_id="od",
|
||||
name="od",
|
||||
url="https://example.com",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth_delegate,
|
||||
)
|
||||
assert (
|
||||
_should_strip_caller_authorization(
|
||||
mcp_server=oauth_delegate,
|
||||
raw_headers={
|
||||
"x-litellm-api-key": "Bearer sk-litellm-key",
|
||||
"authorization": "Bearer upstream",
|
||||
},
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-key"),
|
||||
)
|
||||
is False
|
||||
)
|
||||
assert (
|
||||
_should_strip_caller_authorization(
|
||||
mcp_server=oauth_delegate,
|
||||
raw_headers={"authorization": "Bearer sk-litellm-key"},
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-key"),
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
def test_should_strip_authorization_for_oauth_delegate_admitted_via_jwt_without_api_key(self):
|
||||
"""JWT / SSO / OIDC / session admission yields a UserAPIKeyAuth with a user_id but
|
||||
api_key=None; the caller's Authorization was that credential and must be stripped for
|
||||
oauth_delegate when no separate x-litellm-api-key carried admission (LIT-3794-class leak)."""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
oauth_delegate = MCPServer(
|
||||
server_id="od-jwt",
|
||||
name="od",
|
||||
url="https://example.com",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth_delegate,
|
||||
)
|
||||
assert (
|
||||
_should_strip_caller_authorization(
|
||||
mcp_server=oauth_delegate,
|
||||
raw_headers={"authorization": "Bearer eyJ-idp-jwt"},
|
||||
user_api_key_auth=UserAPIKeyAuth(user_id="alice", api_key=None),
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert (
|
||||
_should_strip_caller_authorization(
|
||||
mcp_server=oauth_delegate,
|
||||
raw_headers={
|
||||
"x-litellm-api-key": "Bearer sk-1234",
|
||||
"authorization": "Bearer upstream",
|
||||
},
|
||||
user_api_key_auth=UserAPIKeyAuth(user_id="alice", api_key=None),
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_regular_mcp_tool_oauth_delegate_never_forwards_jwt_admission(self):
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
server = MCPServer(
|
||||
server_id="od-jwt-e2e",
|
||||
name="od",
|
||||
url="https://example.com",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth_delegate,
|
||||
)
|
||||
extra_headers = await self._capture_call_extra_headers(
|
||||
server,
|
||||
oauth2_headers={"Authorization": "Bearer eyJ-idp-jwt"},
|
||||
raw_headers={"authorization": "Bearer eyJ-idp-jwt"},
|
||||
user_api_key_auth=UserAPIKeyAuth(user_id="alice", api_key=None),
|
||||
)
|
||||
assert not extra_headers or "authorization" not in {k.lower() for k in extra_headers}
|
||||
|
||||
def test_new_passthrough_modes_require_per_user_auth(self):
|
||||
for auth_type in (MCPAuth.true_passthrough, MCPAuth.oauth_delegate):
|
||||
server = MCPServer(
|
||||
server_id="s",
|
||||
name="s",
|
||||
url="https://example.com",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=auth_type,
|
||||
)
|
||||
assert server.requires_per_user_auth is True
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_mcp_client_forwarded_modes_use_the_passthrough_arm(self):
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="tp-egress",
|
||||
name="tp",
|
||||
url="https://example.com",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.true_passthrough,
|
||||
)
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.resolve_mcp_auth",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_resolve,
|
||||
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient") as mock_client_cls,
|
||||
):
|
||||
await manager._create_mcp_client(server=server, extra_headers={"Authorization": "Bearer upstream-token"})
|
||||
mock_resolve.assert_not_awaited()
|
||||
kwargs = mock_client_cls.call_args.kwargs
|
||||
emitted = httpx.Request("GET", "https://example.com/mcp")
|
||||
flow = kwargs["resolved_auth"].auth_flow(emitted)
|
||||
next(flow)
|
||||
flow.close()
|
||||
assert emitted.headers["Authorization"] == "Bearer upstream-token"
|
||||
assert not kwargs["extra_headers"] or "authorization" not in {k.lower() for k in kwargs["extra_headers"]}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_prompts_from_server_success(self):
|
||||
"""Ensure prompts are fetched and prefixed when requested."""
|
||||
|
|
|
|||
2
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
2
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -27008,7 +27008,7 @@ export interface components {
|
|||
/** Alias */
|
||||
alias?: string | null;
|
||||
/** Auth Type */
|
||||
auth_type?: ("none" | "api_key" | "bearer_token" | "basic" | "authorization" | "oauth2" | "aws_sigv4" | "token" | "oauth2_token_exchange") | null;
|
||||
auth_type?: ("none" | "api_key" | "bearer_token" | "basic" | "authorization" | "oauth2" | "aws_sigv4" | "token" | "oauth2_token_exchange" | "true_passthrough" | "oauth_delegate") | null;
|
||||
/** Mcp Info */
|
||||
mcp_info?: {
|
||||
[key: string]: unknown;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue