diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index e69de29bb2d..ced904d46ca 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -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 + diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 2b983bd7e8c..10cce518c83 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -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), ): """ diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 9dcb516ef5b..34bf19abcb2 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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 diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 14414e1c5ea..0d6f9d763c7 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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): diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 959ff64f69f..ce8372ac386 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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: diff --git a/litellm/types/mcp.py b/litellm/types/mcp.py new file mode 100644 index 00000000000..7baac95b22a --- /dev/null +++ b/litellm/types/mcp.py @@ -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) \ No newline at end of file