mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-17 23:51:30 +00:00
feat(mcp): graft aws_sigv4 to the v2 resolver via the aws_auth seam
aws_sigv4 signs every request, so it cannot ride the header-extraction seam the other grafted modes use; it grafts at MCPClient.aws_auth instead. _create_mcp_client now asks the bridge's resolve_v2_aws_auth(server) for the signer when the flag is on and falls back to v1's MCPSigV4Auth otherwise. The bridge wires HttpxSigV4Signer into the provider and maps a v1 aws_sigv4 server's aws_* fields onto AwsSigV4Config (assume-role / static keys / ambient), deferring to v1 for the one shape v2 can't yet represent (assume-role with explicit base keys). Tests cover the credential-source mapping, the signer signing a real request, the role+base-keys defer-to-v1 case, the flag-off and non-aws defer paths. 83 tests pass; the bridge typechecks clean and the manager imports the guarded hook.
This commit is contained in:
parent
5dd99ca52f
commit
230ca7d195
3 changed files with 196 additions and 11 deletions
|
|
@ -56,6 +56,15 @@ from litellm.proxy._experimental.mcp_server.sampling_handler import (
|
|||
MCP_SAMPLING_AVAILABLE,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import resolve_mcp_auth
|
||||
|
||||
# v2 aws_sigv4 graft: SigV4 signs per request, so it attaches via MCPClient.aws_auth rather than
|
||||
# the header seam resolve_mcp_auth uses. Guarded so a missing optional dep degrades to v1.
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.v2_resolver_bridge import (
|
||||
resolve_v2_aws_auth,
|
||||
)
|
||||
except ImportError:
|
||||
resolve_v2_aws_auth = None # type: ignore[assignment]
|
||||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
MCP_TOOL_PREFIX_SEPARATOR,
|
||||
MCPMissingUserEnvVarsError,
|
||||
|
|
@ -2005,18 +2014,22 @@ class MCPServerManager:
|
|||
# For HTTP/SSE transports
|
||||
server_url = server.url or ""
|
||||
|
||||
# Create SigV4 auth if configured
|
||||
# Create SigV4 auth if configured. The v2 resolver owns this when the flag is on
|
||||
# (returns the signer to attach as aws_auth); otherwise v1 builds it.
|
||||
aws_auth = None
|
||||
if server.auth_type == MCPAuth.aws_sigv4:
|
||||
aws_auth = MCPSigV4Auth(
|
||||
aws_access_key_id=server.aws_access_key_id,
|
||||
aws_secret_access_key=server.aws_secret_access_key,
|
||||
aws_session_token=server.aws_session_token,
|
||||
aws_region_name=server.aws_region_name,
|
||||
aws_service_name=server.aws_service_name,
|
||||
aws_role_name=server.aws_role_name,
|
||||
aws_session_name=server.aws_session_name,
|
||||
)
|
||||
if resolve_v2_aws_auth is not None:
|
||||
aws_auth = await resolve_v2_aws_auth(server)
|
||||
if aws_auth is None:
|
||||
aws_auth = MCPSigV4Auth(
|
||||
aws_access_key_id=server.aws_access_key_id,
|
||||
aws_secret_access_key=server.aws_secret_access_key,
|
||||
aws_session_token=server.aws_session_token,
|
||||
aws_region_name=server.aws_region_name,
|
||||
aws_service_name=server.aws_service_name,
|
||||
aws_role_name=server.aws_role_name,
|
||||
aws_session_name=server.aws_session_name,
|
||||
)
|
||||
|
||||
return MCPClient(
|
||||
server_url=server_url,
|
||||
|
|
|
|||
|
|
@ -22,6 +22,7 @@ from pydantic import SecretStr
|
|||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy._experimental.mcp_server.v2_port_bodies import (
|
||||
HttpxClientCredentialsFetcher,
|
||||
HttpxSigV4Signer,
|
||||
)
|
||||
from litellm.proxy.gateway.mcp.outbound_credentials.clock import SystemClock
|
||||
from litellm.proxy.gateway.mcp.outbound_credentials.credential_store import (
|
||||
|
|
@ -38,7 +39,9 @@ from litellm.proxy.gateway.mcp.outbound_credentials.token_store import (
|
|||
StoredToken,
|
||||
)
|
||||
from litellm.proxy.gateway.mcp.outbound_credentials.types import (
|
||||
Ambient,
|
||||
ApiKeyConfig,
|
||||
AssumeRole,
|
||||
AuthorizationCodeConfig,
|
||||
AwsSigV4Config,
|
||||
ClientCredentialsConfig,
|
||||
|
|
@ -46,6 +49,7 @@ from litellm.proxy.gateway.mcp.outbound_credentials.types import (
|
|||
NoneConfig,
|
||||
ServerSpec,
|
||||
SharedKey,
|
||||
StaticKeys,
|
||||
Subject,
|
||||
TokenExchangeConfig,
|
||||
)
|
||||
|
|
@ -99,7 +103,7 @@ def _provider() -> UpstreamCredentialProvider:
|
|||
service_token_store=InMemoryServiceTokenStore(),
|
||||
client_credentials_fetcher=HttpxClientCredentialsFetcher(),
|
||||
token_exchanger=unwired,
|
||||
signer_factory=unwired,
|
||||
signer_factory=HttpxSigV4Signer(),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -179,3 +183,72 @@ async def resolve_v2_auth_value(server: MCPServer) -> Optional[Dict[str, str]]:
|
|||
f"attached {sorted(headers)}" if headers else "no auth header",
|
||||
)
|
||||
return headers
|
||||
|
||||
|
||||
def _to_aws_sigv4_config(server: MCPServer) -> Optional[AwsSigV4Config]:
|
||||
region = server.aws_region_name or "us-east-1"
|
||||
service = server.aws_service_name or "bedrock-agentcore"
|
||||
if server.aws_role_name:
|
||||
# v2's AssumeRole assumes via the ambient/default chain. v1 also supports assuming with
|
||||
# explicit base keys, which v2 can't represent yet, so defer that case to v1.
|
||||
if server.aws_access_key_id or server.aws_secret_access_key:
|
||||
return None
|
||||
return AwsSigV4Config(
|
||||
region=region,
|
||||
service=service,
|
||||
credentials=AssumeRole(
|
||||
role_arn=server.aws_role_name,
|
||||
session_name=server.aws_session_name,
|
||||
),
|
||||
)
|
||||
if server.aws_access_key_id and server.aws_secret_access_key:
|
||||
return AwsSigV4Config(
|
||||
region=region,
|
||||
service=service,
|
||||
credentials=StaticKeys(
|
||||
access_key_id=server.aws_access_key_id,
|
||||
secret_access_key=SecretStr(server.aws_secret_access_key),
|
||||
session_token=(
|
||||
SecretStr(server.aws_session_token)
|
||||
if server.aws_session_token
|
||||
else None
|
||||
),
|
||||
),
|
||||
)
|
||||
return AwsSigV4Config(region=region, service=service, credentials=Ambient())
|
||||
|
||||
|
||||
async def resolve_v2_aws_auth(server: MCPServer) -> Optional[httpx.Auth]:
|
||||
"""Resolve `aws_sigv4` via the v2 resolver into the SigV4 signer, or None to defer to v1.
|
||||
|
||||
SigV4 signs every request, so (unlike the header modes) the result is attached to the
|
||||
connection as `MCPClient.aws_auth` rather than extracted into a header.
|
||||
"""
|
||||
if not v2_resolver_enabled():
|
||||
return None
|
||||
if server.auth_type != MCPAuth.aws_sigv4:
|
||||
return None
|
||||
config = _to_aws_sigv4_config(server)
|
||||
if config is None:
|
||||
return None
|
||||
result = await _provider().resolve(
|
||||
Subject(tenant_id="", subject_id="", inbound_token=None),
|
||||
ServerSpec(
|
||||
server_id=server.server_id,
|
||||
resource=server.url or server.server_id,
|
||||
config=config,
|
||||
),
|
||||
)
|
||||
if isinstance(result, Error):
|
||||
verbose_logger.warning(
|
||||
"v2 MCP resolver failed to build aws_sigv4 signer for server %s: %s; "
|
||||
"falling back to v1",
|
||||
server.server_id,
|
||||
result.error.summary,
|
||||
)
|
||||
return None
|
||||
verbose_logger.info(
|
||||
"v2 MCP resolver handled server %s (auth_type=aws_sigv4): attached SigV4 signer",
|
||||
server.server_id,
|
||||
)
|
||||
return result.ok
|
||||
|
|
|
|||
|
|
@ -5,12 +5,14 @@ headers must be byte-identical to what v1 produces, and every other mode (or any
|
|||
fall back to v1 unchanged.
|
||||
"""
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.experimental_mcp_client.client import MCPClient
|
||||
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import resolve_mcp_auth
|
||||
from litellm.proxy._experimental.mcp_server.v2_resolver_bridge import (
|
||||
resolve_v2_auth_value,
|
||||
resolve_v2_aws_auth,
|
||||
)
|
||||
from litellm.types.mcp import MCPAuth, MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
|
@ -139,3 +141,100 @@ async def test_client_credentials_graft_end_to_end(v2_on, monkeypatch):
|
|||
assert await resolve_v2_auth_value(_m2m_server()) == {
|
||||
"Authorization": "Bearer m2m-tok"
|
||||
}
|
||||
|
||||
|
||||
def _aws_server(**aws):
|
||||
return MCPServer(
|
||||
server_id="aws",
|
||||
name="aws",
|
||||
transport=MCPTransport.http,
|
||||
url="https://svc.us-east-1.amazonaws.com/mcp",
|
||||
auth_type=MCPAuth.aws_sigv4,
|
||||
**aws,
|
||||
)
|
||||
|
||||
|
||||
async def test_aws_sigv4_flag_off_defers_to_v1(v2_off):
|
||||
server = _aws_server(aws_access_key_id="AKIA", aws_secret_access_key="s")
|
||||
assert await resolve_v2_aws_auth(server) is None
|
||||
|
||||
|
||||
async def test_aws_sigv4_non_aws_server_returns_none(v2_on):
|
||||
assert await resolve_v2_aws_auth(_server(MCPAuth.api_key, "k")) is None
|
||||
|
||||
|
||||
async def test_aws_sigv4_static_keys_returns_signing_auth(v2_on):
|
||||
server = _aws_server(
|
||||
aws_access_key_id="AKIATEST",
|
||||
aws_secret_access_key="secret",
|
||||
aws_region_name="us-east-1",
|
||||
)
|
||||
auth = await resolve_v2_aws_auth(server)
|
||||
assert auth is not None
|
||||
req = httpx.Request(
|
||||
"POST", "https://svc.us-east-1.amazonaws.com/mcp", content=b"{}"
|
||||
)
|
||||
signed = next(auth.auth_flow(req))
|
||||
assert signed.headers["Authorization"].startswith(
|
||||
"AWS4-HMAC-SHA256 Credential=AKIATEST/"
|
||||
)
|
||||
assert "X-Amz-Date" in signed.headers
|
||||
|
||||
|
||||
async def test_aws_sigv4_config_maps_static_keys(v2_on):
|
||||
from litellm.proxy._experimental.mcp_server.v2_resolver_bridge import (
|
||||
_to_aws_sigv4_config,
|
||||
)
|
||||
from litellm.proxy.gateway.mcp.outbound_credentials.types import StaticKeys
|
||||
|
||||
cfg = _to_aws_sigv4_config(
|
||||
_aws_server(
|
||||
aws_access_key_id="AKIA",
|
||||
aws_secret_access_key="s",
|
||||
aws_region_name="eu-west-1",
|
||||
)
|
||||
)
|
||||
assert cfg is not None
|
||||
assert cfg.region == "eu-west-1"
|
||||
assert isinstance(cfg.credentials, StaticKeys)
|
||||
assert cfg.credentials.access_key_id == "AKIA"
|
||||
|
||||
|
||||
async def test_aws_sigv4_config_maps_assume_role(v2_on):
|
||||
from litellm.proxy._experimental.mcp_server.v2_resolver_bridge import (
|
||||
_to_aws_sigv4_config,
|
||||
)
|
||||
from litellm.proxy.gateway.mcp.outbound_credentials.types import AssumeRole
|
||||
|
||||
cfg = _to_aws_sigv4_config(
|
||||
_aws_server(aws_role_name="arn:aws:iam::1:role/r", aws_session_name="sess")
|
||||
)
|
||||
assert cfg is not None
|
||||
assert isinstance(cfg.credentials, AssumeRole)
|
||||
assert cfg.credentials.role_arn == "arn:aws:iam::1:role/r"
|
||||
|
||||
|
||||
async def test_aws_sigv4_config_role_with_base_keys_defers_to_v1(v2_on):
|
||||
from litellm.proxy._experimental.mcp_server.v2_resolver_bridge import (
|
||||
_to_aws_sigv4_config,
|
||||
)
|
||||
|
||||
cfg = _to_aws_sigv4_config(
|
||||
_aws_server(
|
||||
aws_role_name="arn:aws:iam::1:role/r",
|
||||
aws_access_key_id="AKIA",
|
||||
aws_secret_access_key="s",
|
||||
)
|
||||
)
|
||||
assert cfg is None
|
||||
|
||||
|
||||
async def test_aws_sigv4_config_defaults_to_ambient(v2_on):
|
||||
from litellm.proxy._experimental.mcp_server.v2_resolver_bridge import (
|
||||
_to_aws_sigv4_config,
|
||||
)
|
||||
from litellm.proxy.gateway.mcp.outbound_credentials.types import Ambient
|
||||
|
||||
cfg = _to_aws_sigv4_config(_aws_server())
|
||||
assert cfg is not None
|
||||
assert isinstance(cfg.credentials, Ambient)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue