initial mcp auth with special header

This commit is contained in:
wagnerjt 2025-06-17 19:45:20 -07:00
parent a8196159b9
commit b22fb6d12e
6 changed files with 206 additions and 3 deletions

View file

@ -0,0 +1,148 @@
from typing import List, Optional
import base64
from datetime import timedelta
from mcp import ClientSession
from mcp.types import Tool as MCPTool
from mcp.types import CallToolResult as MCPCallToolResult
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
from mcp.client.sse import sse_client
from mcp.client.streamable_http import streamablehttp_client
from litellm.types.mcp import MCPAuth, MCPTransport, MCPTransportType, MCPAuthType
def to_basic_auth(auth_value: str) -> str:
"""Convert auth value to Basic Auth format."""
return base64.b64encode(auth_value.encode("utf-8")).decode()
class MCPClient:
"""
MCP Client supporting:
SSE and HTTP transports
Authentication via Bearer token, Basic Auth, or API Key
Tool calling with error handling and result parsing
"""
def __init__(
self,
server_url: str,
transport_type: MCPTransportType = MCPTransport.http,
auth_type: MCPAuthType = None,
auth_value: Optional[str] = None,
timeout: float = 60.0,
):
self.server_url: str = server_url
self.transport_type: MCPTransport = transport_type
self.auth_type: MCPAuthType = auth_type
self.timeout: float = timeout
self._mcp_auth_value: Optional[str] = None
self._session: Optional[ClientSession] = None
self._context = None
self._transport_ctx = None
self._transport = None
self._session_ctx = None
# handle the basic auth value if provided
if auth_value:
self.update_auth_value(auth_value)
async def __aenter__(self):
"""
Enable async context manager support.
Initializes the transport and session.
"""
headers = self._get_auth_headers()
if self.transport_type == MCPTransport.sse:
self._transport_ctx = sse_client(
url=self.server_url,
timeout=self.timeout,
headers=headers,
)
self._transport = await self._transport_ctx.__aenter__()
self._session_ctx = ClientSession(self._transport[0], self._transport[1])
self._session = await self._session_ctx.__aenter__()
await self._session.initialize()
else:
self._transport_ctx = streamablehttp_client(
url=self.server_url,
timeout=timedelta(seconds=self.timeout),
headers=headers,
)
self._transport = await self._transport_ctx.__aenter__()
self._session_ctx = ClientSession(self._transport[0], self._transport[1])
self._session = await self._session_ctx.__aenter__()
await self._session.initialize()
return self
async def __aexit__(self, exc_type, exc_val, exc_tb):
"""Cleanup when exiting context manager."""
if self._session:
await self._session_ctx.__aexit__(exc_type, exc_val, exc_tb) # type: ignore
if self._transport_ctx:
await self._transport_ctx.__aexit__(exc_type, exc_val, exc_tb)
async def disconnect(self):
"""Clean up session and connections."""
if self._session:
try:
# Ensure session is properly closed
await self._session.close() # type: ignore
except Exception:
pass
self._session = None
if self._context:
try:
await self._context.__aexit__(None, None, None) # type: ignore
except Exception:
pass
self._context = None
def update_auth_value(self, mcp_auth_value: str):
"""
Set the authentication header for the MCP client.
"""
if self.auth_type == MCPAuth.basic:
# Assuming mcp_auth_value is in format "username:password", convert it when updating
mcp_auth_value = to_basic_auth(mcp_auth_value)
self._mcp_auth_value = mcp_auth_value
def _get_auth_headers(self) -> dict:
"""Generate authentication headers based on auth type."""
if not self._mcp_auth_value:
return {}
if self.auth_type == MCPAuth.bearer_token:
return {"Authorization": f"Bearer {self._mcp_auth_value}"}
elif self.auth_type == MCPAuth.basic:
return {"Authorization": f"Basic {self._mcp_auth_value}"}
elif self.auth_type == MCPAuth.api_key:
return {"X-API-Key": self._mcp_auth_value}
return {}
async def list_tools(self) -> List[MCPTool]:
"""List available tools from the server."""
if not self._session:
await self.connect() # type: ignore
result = await self._session.list_tools()
if hasattr(result, "tools") and result.tools:
return result.tools
return result
async def call_tool(
self, call_tool_request_params: MCPCallToolRequestParams
) -> MCPCallToolResult:
"""
Call an MCP Tool.
"""
if not self._session:
await self.connect() # type: ignore
tool_result = await self._session.call_tool(
name=call_tool_request_params.name,
arguments=call_tool_request_params.arguments,
)
return tool_result

View file

@ -5,7 +5,7 @@ from fastapi import APIRouter, Depends, Query, Request
from litellm._logging import verbose_logger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth, mcp_auth_header
MCP_AVAILABLE: bool = True
try:
@ -37,6 +37,7 @@ if MCP_AVAILABLE:
server_id: Optional[str] = Query(
None, description="The server id to list tools for"
),
mcp_auth: Optional[str] = mcp_auth_header,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> List[ListMCPToolsRestAPIResponseObject]:
"""
@ -88,6 +89,7 @@ if MCP_AVAILABLE:
@router.post("/tools/call", dependencies=[Depends(user_api_key_auth)])
async def call_tool_rest_api(
request: Request,
mcp_auth: Optional[str] = mcp_auth_header,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""

View file

@ -238,7 +238,7 @@ if MCP_AVAILABLE:
@client
async def call_mcp_tool(
name: str, arguments: Optional[Dict[str, Any]] = None, **kwargs: Any
name: str, arguments: Optional[Dict[str, Any]] = None, mcp_auth_header: Optional[str] = None, **kwargs: Any
) -> List[Union[MCPTextContent, MCPImageContent, MCPEmbeddedResource]]:
"""
Call a specific tool with the provided arguments

View file

@ -2688,6 +2688,7 @@ class SpecialHeaders(enum.Enum):
google_ai_studio_authorization = "x-goog-api-key"
azure_apim_authorization = "Ocp-Apim-Subscription-Key"
custom_litellm_api_key = "x-litellm-api-key"
mcp_auth = "x-mcp-auth"
class LitellmDataForBackendLLMCall(TypedDict, total=False):

View file

@ -13,7 +13,7 @@ from datetime import datetime, timezone
from typing import List, Optional, Tuple, cast
import fastapi
from fastapi import HTTPException, Request, WebSocket, status
from fastapi import HTTPException, Request, WebSocket, status, Header
from fastapi.security.api_key import APIKeyHeader
import litellm
@ -101,6 +101,11 @@ azure_apim_header = APIKeyHeader(
auto_error=False,
description="The default name of the subscription key header of Azure",
)
mcp_auth_header: Optional[str] = Header(
name=SpecialHeaders.mcp_auth.value,
default=None,
description="MCP Auth header to be passed to the mcp servers",
)
def _get_bearer_token_or_received_api_key(api_key: str) -> str:

47
litellm/types/mcp.py Normal file
View file

@ -0,0 +1,47 @@
import enum
from typing import Literal, Optional, TypedDict
from pydantic import BaseModel, ConfigDict
class MCPTransport(str, enum.Enum):
sse = "sse"
http = "http"
class MCPSpecVersion(str, enum.Enum):
nov_2024 = "2024-11-05"
mar_2025 = "2025-03-26"
class MCPAuth(str, enum.Enum):
none = "none"
api_key = "api_key"
bearer_token = "bearer_token"
basic = "basic"
# MCP Literals
MCPTransportType = Literal[MCPTransport.sse, MCPTransport.http]
MCPSpecVersionType = Literal[MCPSpecVersion.nov_2024, MCPSpecVersion.mar_2025]
MCPAuthType = Optional[
Literal[MCPAuth.none, MCPAuth.api_key, MCPAuth.bearer_token, MCPAuth.basic]
]
class MCPInfo(TypedDict, total=False):
server_name: str
description: Optional[str]
logo_url: Optional[str]
class MCPServer(BaseModel):
server_id: str
name: str
url: str
# TODO: alter the types to be the Literal explicit
transport: MCPTransportType
spec_version: MCPSpecVersionType
auth_type: Optional[MCPAuthType] = None
mcp_info: Optional[MCPInfo] = None
model_config = ConfigDict(arbitrary_types_allowed=True)