mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
initial mcp auth with special header
This commit is contained in:
parent
a8196159b9
commit
b22fb6d12e
6 changed files with 206 additions and 3 deletions
|
|
@ -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
|
||||
|
||||
|
|
@ -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),
|
||||
):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
47
litellm/types/mcp.py
Normal 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)
|
||||
Loading…
Add table
Reference in a new issue