mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat(mcp): graft v2 resolver onto _create_mcp_client for migrated modes
Wire the none + api_key static-family resolver arms from PR4a onto v1's live request path. In _create_mcp_client's HTTP/SSE branch, to_server_spec decides per mode: a migrated mode resolves through the injected UpstreamCredentialProvider and feeds the resulting httpx.Auth into the new resolved_auth slot; every other mode returns None and falls through to the unchanged v1 construction. resolve_mcp_auth now runs only when the mode defers, so a migrated server skips the v1 token-exchange / M2M I/O. stdio is untouched: auth_type/auth_value never reach the upstream on the stdio path (_get_auth_headers is HTTP/SSE only), so there is nothing to graft there. No v1 code is deleted yet; resolve_mcp_auth's static return still backs stdio and the not-yet-migrated modes until later PRs retire it.
This commit is contained in:
parent
161d438791
commit
9f2b336abe
1 changed files with 41 additions and 5 deletions
|
|
@ -56,6 +56,16 @@ 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
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials import (
|
||||
Error,
|
||||
Ok,
|
||||
UpstreamCredentialProvider,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
|
||||
raise_public,
|
||||
to_server_spec,
|
||||
to_subject,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
MCP_TOOL_PREFIX_SEPARATOR,
|
||||
MCPMissingUserEnvVarsError,
|
||||
|
|
@ -511,7 +521,8 @@ class MCPServerManager:
|
|||
return "client_credentials"
|
||||
return None
|
||||
|
||||
def __init__(self):
|
||||
def __init__(self, cred_provider: Optional[UpstreamCredentialProvider] = None):
|
||||
self._cred_provider = cred_provider or UpstreamCredentialProvider()
|
||||
self.registry: Dict[str, MCPServer] = {}
|
||||
self.config_mcp_servers: Dict[str, MCPServer] = {}
|
||||
"""
|
||||
|
|
@ -1942,11 +1953,13 @@ class MCPServerManager:
|
|||
Returns:
|
||||
Configured MCP client instance.
|
||||
"""
|
||||
auth_value = await resolve_mcp_auth(
|
||||
server, mcp_auth_header, subject_token=subject_token
|
||||
)
|
||||
|
||||
transport = server.transport or MCPTransport.sse
|
||||
spec = None if transport == MCPTransport.stdio else to_server_spec(server)
|
||||
auth_value = (
|
||||
await resolve_mcp_auth(server, mcp_auth_header, subject_token=subject_token)
|
||||
if spec is None
|
||||
else None
|
||||
)
|
||||
|
||||
# Create sampling and elicitation callbacks for this client
|
||||
sampling_cb = (
|
||||
|
|
@ -2017,6 +2030,29 @@ class MCPServerManager:
|
|||
# For HTTP/SSE transports
|
||||
server_url = server.url or ""
|
||||
|
||||
if spec is not None:
|
||||
match await self._cred_provider.resolve_credentials(
|
||||
to_subject(user_api_key_auth, subject_token), spec
|
||||
):
|
||||
case Ok(auth):
|
||||
resolved_auth = auth
|
||||
case Error(err):
|
||||
raise_public(err)
|
||||
return MCPClient(
|
||||
server_url=server_url,
|
||||
transport_type=transport,
|
||||
auth_type=server.auth_type,
|
||||
timeout=(
|
||||
server.timeout
|
||||
if server.timeout is not None
|
||||
else MCP_CLIENT_TIMEOUT
|
||||
),
|
||||
extra_headers=extra_headers,
|
||||
resolved_auth=resolved_auth,
|
||||
sampling_callback=sampling_cb,
|
||||
elicitation_callback=elicitation_cb,
|
||||
)
|
||||
|
||||
# Create SigV4 auth if configured
|
||||
aws_auth = None
|
||||
if server.auth_type == MCPAuth.aws_sigv4:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue