diff --git a/.circleci/config.yml b/.circleci/config.yml index 26268ba3573..c8d9419c93c 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -66,7 +66,7 @@ jobs: pip install python-multipart pip install google-cloud-aiplatform pip install prometheus-client==0.20.0 - pip install "pydantic==2.7.1" + pip install "pydantic==2.10.2" pip install "diskcache==5.6.1" pip install "Pillow==10.3.0" pip install "jsonschema==4.22.0" @@ -185,7 +185,7 @@ jobs: pip install python-multipart pip install google-cloud-aiplatform pip install prometheus-client==0.20.0 - pip install "pydantic==2.7.1" + pip install "pydantic==2.10.2" pip install "diskcache==5.6.1" pip install "Pillow==10.3.0" pip install "jsonschema==4.22.0" @@ -285,7 +285,7 @@ jobs: pip install python-multipart pip install google-cloud-aiplatform pip install prometheus-client==0.20.0 - pip install "pydantic==2.7.1" + pip install "pydantic==2.10.2" pip install "diskcache==5.6.1" pip install "Pillow==10.3.0" pip install "jsonschema==4.22.0" @@ -530,7 +530,7 @@ jobs: pip install python-multipart pip install google-cloud-aiplatform pip install prometheus-client==0.20.0 - pip install "pydantic==2.7.1" + pip install "pydantic==2.10.2" pip install "diskcache==5.6.1" pip install "Pillow==10.3.0" pip install "jsonschema==4.22.0" @@ -625,7 +625,13 @@ jobs: python -m pip install -r requirements.txt pip install "pytest==7.3.1" pip install "pytest-retry==1.6.3" + pip install "pytest-cov==5.0.0" pip install "pytest-asyncio==0.21.1" + pip install "respx==0.21.1" + - run: + name: Show current pydantic version + command: | + python -m pip show pydantic # Run pytest and generate JUnit XML report - run: name: Run tests @@ -700,8 +706,8 @@ jobs: pip install "pytest-cov==5.0.0" pip install "pytest-asyncio==0.21.1" pip install "respx==0.21.1" - pip install "pydantic==2.7.2" - pip install "mcp==1.4.1" + pip install "pydantic==2.10.2" + pip install "mcp==1.5.0" # Run pytest and generate JUnit XML report - run: name: Run tests @@ -788,8 +794,8 @@ jobs: pip install "pytest-asyncio==0.21.1" pip install "respx==0.21.1" pip install "hypercorn==0.17.3" - pip install "pydantic==2.7.2" - pip install "mcp==1.4.1" + pip install "pydantic==2.10.2" + pip install "mcp==1.5.0" # Run pytest and generate JUnit XML report - run: name: Run tests @@ -870,6 +876,7 @@ jobs: name: Install Dependencies command: | python -m pip install --upgrade pip + pip install numpydoc python -m pip install -r requirements.txt pip install "respx==0.21.1" pip install "pytest==7.3.1" @@ -878,7 +885,6 @@ jobs: pip install "pytest-cov==5.0.0" pip install "google-generativeai==0.3.2" pip install "google-cloud-aiplatform==1.43.0" - pip install numpydoc # Run pytest and generate JUnit XML report - run: name: Run tests @@ -1054,8 +1060,9 @@ jobs: pip install click pip install "boto3==1.34.34" pip install jinja2 - pip install tokenizers=="0.20.0" - pip install uvloop==0.21.0 + pip install "tokenizers==0.20.0" + pip install "uvloop==0.21.0" + pip install "mcp==1.5.0" pip install jsonschema - run: name: Run tests @@ -1448,6 +1455,7 @@ jobs: pip install "boto3==1.34.34" pip install "aioboto3==12.3.0" pip install langchain + pip install "langchain_mcp_adapters==0.0.5" pip install "langfuse>=2.0.0" pip install "logfire==0.29.0" pip install numpydoc @@ -2014,7 +2022,7 @@ jobs: pip install "openai==1.68.2" pip install "assemblyai==0.37.0" python -m pip install --upgrade pip - pip install "pydantic==2.7.1" + pip install "pydantic==2.10.2" pip install "pytest==7.3.1" pip install "pytest-mock==3.12.0" pip install "pytest-asyncio==0.21.1" @@ -2289,7 +2297,7 @@ jobs: pip install aiohttp pip install "openai==1.68.2" python -m pip install --upgrade pip - pip install "pydantic==2.7.1" + pip install "pydantic==2.10.2" pip install "pytest==7.3.1" pip install "pytest-mock==3.12.0" pip install "pytest-asyncio==0.21.1" diff --git a/.circleci/requirements.txt b/.circleci/requirements.txt index cada0f605e6..5dece1fc8b6 100644 --- a/.circleci/requirements.txt +++ b/.circleci/requirements.txt @@ -8,7 +8,8 @@ redis==5.2.1 redisvl==0.4.1 anthropic orjson==3.9.15 -pydantic==2.7.1 +pydantic==2.10.2 google-cloud-aiplatform==1.43.0 fastapi-sso==0.10.0 uvloop==0.21.0 +mcp==1.5.0 # for MCP server diff --git a/docker-compose.yml b/docker-compose.yml index d16ec6ed20b..1cab80ee4cc 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -1,35 +1,5 @@ version: "3.11" services: - litellm: - build: - context: . - args: - target: runtime - image: ghcr.io/berriai/litellm:main-stable - ######################################### - ## Uncomment these lines to start proxy with a config.yaml file ## - # volumes: - # - ./config.yaml:/app/config.yaml <<- this is missing in the docker-compose file currently - # command: - # - "--config=/app/config.yaml" - ############################################## - ports: - - "4000:4000" # Map the container port to the host, change the host port if necessary - environment: - DATABASE_URL: "postgresql://llmproxy:dbpassword9090@db:5432/litellm" - STORE_MODEL_IN_DB: "True" # allows adding models to proxy via UI - env_file: - - .env # Load local .env file - depends_on: - - db # Indicates that this service depends on the 'db' service, ensuring 'db' starts first - healthcheck: # Defines the health check configuration for the container - test: [ "CMD", "curl", "-f", "http://localhost:4000/health/liveliness || exit 1" ] # Command to execute for health check - interval: 30s # Perform health check every 30 seconds - timeout: 10s # Health check command times out after 10 seconds - retries: 3 # Retry up to 3 times if health check fails - start_period: 40s # Wait 40 seconds after container start before beginning health checks - - db: image: postgres:16 restart: always @@ -46,25 +16,3 @@ services: interval: 1s timeout: 5s retries: 10 - - prometheus: - image: prom/prometheus - volumes: - - prometheus_data:/prometheus - - ./prometheus.yml:/etc/prometheus/prometheus.yml - ports: - - "9090:9090" - command: - - '--config.file=/etc/prometheus/prometheus.yml' - - '--storage.tsdb.path=/prometheus' - - '--storage.tsdb.retention.time=15d' - restart: always - -volumes: - prometheus_data: - driver: local - postgres_data: - name: litellm_postgres_data # Named volume for Postgres data persistence - - -# ...rest of your docker-compose config if any diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index 5aeee715d14..2e5d0f2a0a6 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -304,6 +304,7 @@ const sidebars = { "image_variations", ] }, + "mcp", { type: "category", label: "/audio", diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py new file mode 100644 index 00000000000..02a9303be5d --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -0,0 +1,110 @@ +""" +LiteLLM MCP Server Routes +""" + +import asyncio +from typing import Any, Dict, List, Union + +from anyio import BrokenResourceError +from fastapi import APIRouter, HTTPException, Request +from fastapi.responses import StreamingResponse +from mcp.server import NotificationOptions, Server +from mcp.server.models import InitializationOptions +from mcp.types import EmbeddedResource as MCPEmbeddedResource +from mcp.types import ImageContent as MCPImageContent +from mcp.types import TextContent as MCPTextContent +from mcp.types import Tool as MCPTool +from pydantic import ValidationError + +from litellm._logging import verbose_logger +from litellm.proxy._experimental.mcp_server.tool_registry import ( + global_mcp_tool_registry, +) + +from .sse_transport import SseServerTransport + +######################################################## +############ Initialize the MCP Server ################# +######################################################## +router = APIRouter( + prefix="/mcp", + tags=["mcp"], +) +server: Server = Server("litellm-mcp-server") +sse: SseServerTransport = SseServerTransport("/mcp/sse/messages") + +######################################################## +############### MCP Server Routes ####################### +######################################################## + + +@server.list_tools() +async def list_tools() -> list[MCPTool]: + """ + List all available tools + """ + tools = [] + for tool in global_mcp_tool_registry.list_tools(): + tools.append( + MCPTool( + name=tool.name, + description=tool.description, + inputSchema=tool.input_schema, + ) + ) + + return tools + + +@server.call_tool() +async def handle_call_tool( + name: str, arguments: Dict[str, Any] | None +) -> List[Union[MCPTextContent, MCPImageContent, MCPEmbeddedResource]]: + """ + Call a specific tool with the provided arguments + """ + tool = global_mcp_tool_registry.get_tool(name) + if not tool: + raise HTTPException(status_code=404, detail=f"Tool '{name}' not found") + if arguments is None: + raise HTTPException(status_code=400, detail="Request arguments are required") + + try: + result = tool.handler(**arguments) + return [MCPTextContent(text=str(result), type="text")] + except Exception as e: + return [MCPTextContent(text=f"Error: {str(e)}", type="text")] + + +@router.get("/", response_class=StreamingResponse) +async def handle_sse(request: Request): + verbose_logger.info("new incoming SSE connection established") + async with sse.connect_sse(request) as streams: + try: + await server.run(streams[0], streams[1], options) + except BrokenResourceError: + pass + except asyncio.CancelledError: + pass + except ValidationError: + pass + except Exception: + raise + await request.close() + + +@router.post("/sse/messages") +async def handle_messages(request: Request): + verbose_logger.info("incoming SSE message received") + await sse.handle_post_message(request.scope, request.receive, request._send) + await request.close() + + +options = InitializationOptions( + server_name="litellm-mcp-server", + server_version="0.1.0", + capabilities=server.get_capabilities( + notification_options=NotificationOptions(), + experimental_capabilities={}, + ), +) diff --git a/litellm/proxy/_experimental/mcp_server/sse_transport.py b/litellm/proxy/_experimental/mcp_server/sse_transport.py new file mode 100644 index 00000000000..63ffd403c66 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/sse_transport.py @@ -0,0 +1,150 @@ +""" +This is a modification of code from: https://github.com/SecretiveShell/MCP-Bridge/blob/master/mcp_bridge/mcp_server/sse_transport.py + +Credit to the maintainers of SecretiveShell for their SSE Transport implementation + +""" + +from contextlib import asynccontextmanager +from typing import Any +from urllib.parse import quote +from uuid import UUID, uuid4 + +import anyio +import mcp.types as types +from anyio.streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream +from fastapi.requests import Request +from fastapi.responses import Response +from pydantic import ValidationError +from sse_starlette import EventSourceResponse +from starlette.types import Receive, Scope, Send + +from litellm._logging import verbose_logger + + +class SseServerTransport: + """ + SSE server transport for MCP. This class provides _two_ ASGI applications, + suitable to be used with a framework like Starlette and a server like Hypercorn: + + 1. connect_sse() is an ASGI application which receives incoming GET requests, + and sets up a new SSE stream to send server messages to the client. + 2. handle_post_message() is an ASGI application which receives incoming POST + requests, which should contain client messages that link to a + previously-established SSE session. + """ + + _endpoint: str + _read_stream_writers: dict[ + UUID, MemoryObjectSendStream[types.JSONRPCMessage | Exception] + ] + + def __init__(self, endpoint: str) -> None: + """ + Creates a new SSE server transport, which will direct the client to POST + messages to the relative or absolute URL given. + """ + + super().__init__() + self._endpoint = endpoint + self._read_stream_writers = {} + verbose_logger.debug( + f"SseServerTransport initialized with endpoint: {endpoint}" + ) + + @asynccontextmanager + async def connect_sse(self, request: Request): + if request.scope["type"] != "http": + verbose_logger.error("connect_sse received non-HTTP request") + raise ValueError("connect_sse can only handle HTTP requests") + + verbose_logger.debug("Setting up SSE connection") + read_stream: MemoryObjectReceiveStream[types.JSONRPCMessage | Exception] + read_stream_writer: MemoryObjectSendStream[types.JSONRPCMessage | Exception] + + write_stream: MemoryObjectSendStream[types.JSONRPCMessage] + write_stream_reader: MemoryObjectReceiveStream[types.JSONRPCMessage] + + read_stream_writer, read_stream = anyio.create_memory_object_stream(0) + write_stream, write_stream_reader = anyio.create_memory_object_stream(0) + + session_id = uuid4() + session_uri = f"{quote(self._endpoint)}?session_id={session_id.hex}" + self._read_stream_writers[session_id] = read_stream_writer + verbose_logger.debug(f"Created new session with ID: {session_id}") + + sse_stream_writer: MemoryObjectSendStream[dict[str, Any]] + sse_stream_reader: MemoryObjectReceiveStream[dict[str, Any]] + sse_stream_writer, sse_stream_reader = anyio.create_memory_object_stream( + 0, dict[str, Any] + ) + + async def sse_writer(): + verbose_logger.debug("Starting SSE writer") + async with sse_stream_writer, write_stream_reader: + await sse_stream_writer.send({"event": "endpoint", "data": session_uri}) + verbose_logger.debug(f"Sent endpoint event: {session_uri}") + + async for message in write_stream_reader: + verbose_logger.debug(f"Sending message via SSE: {message}") + await sse_stream_writer.send( + { + "event": "message", + "data": message.model_dump_json( + by_alias=True, exclude_none=True + ), + } + ) + + async with anyio.create_task_group() as tg: + response = EventSourceResponse( + content=sse_stream_reader, data_sender_callable=sse_writer + ) + verbose_logger.debug("Starting SSE response task") + tg.start_soon(response, request.scope, request.receive, request._send) + + verbose_logger.debug("Yielding read and write streams") + yield (read_stream, write_stream) + + async def handle_post_message( + self, scope: Scope, receive: Receive, send: Send + ) -> Response: + verbose_logger.debug("Handling POST message") + request = Request(scope, receive) + + session_id_param = request.query_params.get("session_id") + if session_id_param is None: + verbose_logger.warning("Received request without session_id") + response = Response("session_id is required", status_code=400) + return response + + try: + session_id = UUID(hex=session_id_param) + verbose_logger.debug(f"Parsed session ID: {session_id}") + except ValueError: + verbose_logger.warning(f"Received invalid session ID: {session_id_param}") + response = Response("Invalid session ID", status_code=400) + return response + + writer = self._read_stream_writers.get(session_id) + if not writer: + verbose_logger.warning(f"Could not find session for ID: {session_id}") + response = Response("Could not find session", status_code=404) + return response + + json = await request.json() + verbose_logger.debug(f"Received JSON: {json}") + + try: + message = types.JSONRPCMessage.model_validate(json) + verbose_logger.debug(f"Validated client message: {message}") + except ValidationError as err: + verbose_logger.error(f"Failed to parse message: {err}") + response = Response("Could not parse message", status_code=400) + await writer.send(err) + return response + + verbose_logger.debug(f"Sending message to writer: {message}") + response = Response("Accepted", status_code=202) + await writer.send(message) + return response diff --git a/litellm/proxy/_experimental/mcp_server/tool_registry.py b/litellm/proxy/_experimental/mcp_server/tool_registry.py new file mode 100644 index 00000000000..c08b7979683 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/tool_registry.py @@ -0,0 +1,103 @@ +import json +from typing import Any, Callable, Dict, List, Optional + +from litellm._logging import verbose_logger +from litellm.proxy.types_utils.utils import get_instance_fn +from litellm.types.mcp_server.tool_registry import MCPTool + + +class MCPToolRegistry: + """ + A registry for managing MCP tools + """ + + def __init__(self): + # Registry to store all registered tools + self.tools: Dict[str, MCPTool] = {} + + def register_tool( + self, + name: str, + description: str, + input_schema: Dict[str, Any], + handler: Callable, + ) -> None: + """ + Register a new tool in the registry + """ + self.tools[name] = MCPTool( + name=name, + description=description, + input_schema=input_schema, + handler=handler, + ) + verbose_logger.debug(f"Registered tool: {name}") + + def get_tool(self, name: str) -> Optional[MCPTool]: + """ + Get a tool from the registry by name + """ + return self.tools.get(name) + + def list_tools(self) -> List[MCPTool]: + """ + List all registered tools + """ + return list(self.tools.values()) + + def load_tools_from_config( + self, mcp_tools_config: Optional[Dict[str, Any]] = None + ) -> None: + """ + Load and register tools from the proxy config + + Args: + mcp_tools_config: The mcp_tools config from the proxy config + """ + if mcp_tools_config is None: + raise ValueError( + "mcp_tools_config is required, please set `mcp_tools` in your proxy config" + ) + + for tool_config in mcp_tools_config: + if not isinstance(tool_config, dict): + raise ValueError("mcp_tools_config must be a list of dictionaries") + + name = tool_config.get("name") + description = tool_config.get("description") + input_schema = tool_config.get("input_schema", {}) + handler_name = tool_config.get("handler") + + if not all([name, description, handler_name]): + continue + + # Try to resolve the handler + # First check if it's a module path (e.g., "module.submodule.function") + if handler_name is None: + raise ValueError(f"handler is required for tool {name}") + handler = get_instance_fn(handler_name) + + if handler is None: + verbose_logger.warning( + f"Warning: Could not find handler {handler_name} for tool {name}" + ) + continue + + # Register the tool + if name is None: + raise ValueError(f"name is required for tool {name}") + if description is None: + raise ValueError(f"description is required for tool {name}") + + self.register_tool( + name=name, + description=description, + input_schema=input_schema, + handler=handler, + ) + verbose_logger.debug( + "all registered tools: %s", json.dumps(self.tools, indent=4, default=str) + ) + + +global_mcp_tool_registry = MCPToolRegistry() diff --git a/litellm/proxy/mcp_tools.py b/litellm/proxy/mcp_tools.py new file mode 100644 index 00000000000..eb63e6b43bc --- /dev/null +++ b/litellm/proxy/mcp_tools.py @@ -0,0 +1,35 @@ +from typing import Any, Dict, Optional + + +def get_current_time(params: Optional[Dict[str, Any]] = None) -> str: + """ + Get the current time (hardcoded sample implementation) + + Args: + params: Optional dictionary with parameters + - format: The format of the time to return (e.g., "short") + + Returns: + A string representing the current time + """ + # Hardcoded time value for sample implementation + if params and params.get("format") == "short": + return "10:30 AM" + return "10:30:45 AM" + + +def get_current_date(params: Optional[Dict[str, Any]] = None) -> str: + """ + Get the current date (hardcoded sample implementation) + + Args: + params: Optional dictionary with parameters + - format: The format of the date to return (e.g., "short") + + Returns: + A string representing the current date + """ + # Hardcoded date value for sample implementation + if params and params.get("format") == "short": + return "Oct 15" + return "October 15, 2023" diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index 0877a02a74d..26ce6cb8f84 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -7,5 +7,31 @@ model_list: api_key: os.environ/AZURE_API_KEY -litellm_settings: - callbacks: ["custom_prompt_management.x42_prompt_management"] + +mcp_tools: + - name: "get_current_time" + description: "Get the current time" + input_schema: { + "type": "object", + "properties": { + "format": { + "type": "string", + "description": "The format of the time to return", + "enum": ["short"] + } + } + } + handler: "mcp_tools.get_current_time" + - name: "get_current_date" + description: "Get the current date" + input_schema: { + "type": "object", + "properties": { + "format": { + "type": "string", + "description": "The format of the date to return", + "enum": ["short"] + } + } + } + handler: "mcp_tools.get_current_date" diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index cac416e75f1..568253735c1 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -126,6 +126,10 @@ from litellm.litellm_core_utils.core_helpers import ( from litellm.litellm_core_utils.credential_accessor import CredentialAccessor from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.proxy._experimental.mcp_server.server import router as mcp_router +from litellm.proxy._experimental.mcp_server.tool_registry import ( + global_mcp_tool_registry, +) from litellm.proxy._types import * from litellm.proxy.analytics_endpoints.analytics_endpoints import ( router as analytics_router, @@ -2153,6 +2157,11 @@ class ProxyConfig: all_guardrails=guardrails_v2, config_file_path=config_file_path ) + ## MCP TOOLS + mcp_tools_config = config.get("mcp_tools", None) + if mcp_tools_config: + global_mcp_tool_registry.load_tools_from_config(mcp_tools_config) + ## CREDENTIALS credential_list_dict = self.load_credential_list(config=config) litellm.credential_list = credential_list_dict @@ -8162,6 +8171,7 @@ app.include_router(rerank_router) app.include_router(fine_tuning_router) app.include_router(credential_router) app.include_router(llm_passthrough_router) +app.include_router(mcp_router) app.include_router(anthropic_router) app.include_router(langfuse_router) app.include_router(pass_through_router) diff --git a/litellm/types/mcp_server/tool_registry.py b/litellm/types/mcp_server/tool_registry.py new file mode 100644 index 00000000000..f2c1cf1a30b --- /dev/null +++ b/litellm/types/mcp_server/tool_registry.py @@ -0,0 +1,35 @@ +from typing import Any, Callable, Dict, List, Optional + +from pydantic import BaseModel + + +class MCPTool(BaseModel): + name: str + description: str + input_schema: Dict[str, Any] + handler: Callable + + class Config: + arbitrary_types_allowed = True + + +class ToolSchema(BaseModel): + name: str + description: str + inputSchema: Dict[str, Any] + + +class ListToolsResponse(BaseModel): + tools: List[ToolSchema] + nextCursor: Optional[str] = None + _meta: Optional[Dict[str, Any]] = None + + +class CallToolRequest(BaseModel): + method: str = "tools/call" + params: Dict[str, Any] + + +class ContentItem(BaseModel): + type: str + text: Optional[str] = None diff --git a/poetry.lock b/poetry.lock index fc2b4743bf0..9273f91accc 100644 --- a/poetry.lock +++ b/poetry.lock @@ -1,4 +1,4 @@ -# This file is automatically @generated by Poetry 2.0.1 and should not be changed by hand. +# This file is automatically @generated by Poetry 2.0.0 and should not be changed by hand. [[package]] name = "aiohappyeyeballs" @@ -1154,7 +1154,7 @@ description = "HTTP/2-based RPC framework" optional = true python-versions = ">=3.8" groups = ["main"] -markers = "extra == \"extra-proxy\" and python_version < \"3.11\"" +markers = "extra == \"extra-proxy\"" files = [ {file = "grpcio-1.70.0-cp310-cp310-linux_armv7l.whl", hash = "sha256:95469d1977429f45fe7df441f586521361e235982a0b39e33841549143ae2851"}, {file = "grpcio-1.70.0-cp310-cp310-macosx_12_0_universal2.whl", hash = "sha256:ed9718f17fbdb472e33b869c77a16d0b55e166b100ec57b016dc7de9c8d236bf"}, @@ -1216,71 +1216,6 @@ files = [ [package.extras] protobuf = ["grpcio-tools (>=1.70.0)"] -[[package]] -name = "grpcio" -version = "1.71.0" -description = "HTTP/2-based RPC framework" -optional = true -python-versions = ">=3.9" -groups = ["main"] -markers = "python_version >= \"3.11\" and extra == \"extra-proxy\"" -files = [ - {file = "grpcio-1.71.0-cp310-cp310-linux_armv7l.whl", hash = "sha256:c200cb6f2393468142eb50ab19613229dcc7829b5ccee8b658a36005f6669fdd"}, - {file = "grpcio-1.71.0-cp310-cp310-macosx_12_0_universal2.whl", hash = "sha256:b2266862c5ad664a380fbbcdbdb8289d71464c42a8c29053820ee78ba0119e5d"}, - {file = "grpcio-1.71.0-cp310-cp310-manylinux_2_17_aarch64.whl", hash = "sha256:0ab8b2864396663a5b0b0d6d79495657ae85fa37dcb6498a2669d067c65c11ea"}, - {file = "grpcio-1.71.0-cp310-cp310-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:c30f393f9d5ff00a71bb56de4aa75b8fe91b161aeb61d39528db6b768d7eac69"}, - {file = "grpcio-1.71.0-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:f250ff44843d9a0615e350c77f890082102a0318d66a99540f54769c8766ab73"}, - {file = "grpcio-1.71.0-cp310-cp310-musllinux_1_1_aarch64.whl", hash = "sha256:e6d8de076528f7c43a2f576bc311799f89d795aa6c9b637377cc2b1616473804"}, - {file = "grpcio-1.71.0-cp310-cp310-musllinux_1_1_i686.whl", hash = "sha256:9b91879d6da1605811ebc60d21ab6a7e4bae6c35f6b63a061d61eb818c8168f6"}, - {file = "grpcio-1.71.0-cp310-cp310-musllinux_1_1_x86_64.whl", hash = "sha256:f71574afdf944e6652203cd1badcda195b2a27d9c83e6d88dc1ce3cfb73b31a5"}, - {file = "grpcio-1.71.0-cp310-cp310-win32.whl", hash = "sha256:8997d6785e93308f277884ee6899ba63baafa0dfb4729748200fcc537858a509"}, - {file = "grpcio-1.71.0-cp310-cp310-win_amd64.whl", hash = "sha256:7d6ac9481d9d0d129224f6d5934d5832c4b1cddb96b59e7eba8416868909786a"}, - {file = "grpcio-1.71.0-cp311-cp311-linux_armv7l.whl", hash = "sha256:d6aa986318c36508dc1d5001a3ff169a15b99b9f96ef5e98e13522c506b37eef"}, - {file = "grpcio-1.71.0-cp311-cp311-macosx_10_14_universal2.whl", hash = "sha256:d2c170247315f2d7e5798a22358e982ad6eeb68fa20cf7a820bb74c11f0736e7"}, - {file = "grpcio-1.71.0-cp311-cp311-manylinux_2_17_aarch64.whl", hash = "sha256:e6f83a583ed0a5b08c5bc7a3fe860bb3c2eac1f03f1f63e0bc2091325605d2b7"}, - {file = "grpcio-1.71.0-cp311-cp311-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:4be74ddeeb92cc87190e0e376dbc8fc7736dbb6d3d454f2fa1f5be1dee26b9d7"}, - {file = "grpcio-1.71.0-cp311-cp311-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:4dd0dfbe4d5eb1fcfec9490ca13f82b089a309dc3678e2edabc144051270a66e"}, - {file = "grpcio-1.71.0-cp311-cp311-musllinux_1_1_aarch64.whl", hash = "sha256:a2242d6950dc892afdf9e951ed7ff89473aaf744b7d5727ad56bdaace363722b"}, - {file = "grpcio-1.71.0-cp311-cp311-musllinux_1_1_i686.whl", hash = "sha256:0fa05ee31a20456b13ae49ad2e5d585265f71dd19fbd9ef983c28f926d45d0a7"}, - {file = "grpcio-1.71.0-cp311-cp311-musllinux_1_1_x86_64.whl", hash = "sha256:3d081e859fb1ebe176de33fc3adb26c7d46b8812f906042705346b314bde32c3"}, - {file = "grpcio-1.71.0-cp311-cp311-win32.whl", hash = "sha256:d6de81c9c00c8a23047136b11794b3584cdc1460ed7cbc10eada50614baa1444"}, - {file = "grpcio-1.71.0-cp311-cp311-win_amd64.whl", hash = "sha256:24e867651fc67717b6f896d5f0cac0ec863a8b5fb7d6441c2ab428f52c651c6b"}, - {file = "grpcio-1.71.0-cp312-cp312-linux_armv7l.whl", hash = "sha256:0ff35c8d807c1c7531d3002be03221ff9ae15712b53ab46e2a0b4bb271f38537"}, - {file = "grpcio-1.71.0-cp312-cp312-macosx_10_14_universal2.whl", hash = "sha256:b78a99cd1ece4be92ab7c07765a0b038194ded2e0a26fd654591ee136088d8d7"}, - {file = "grpcio-1.71.0-cp312-cp312-manylinux_2_17_aarch64.whl", hash = "sha256:dc1a1231ed23caac1de9f943d031f1bc38d0f69d2a3b243ea0d664fc1fbd7fec"}, - {file = "grpcio-1.71.0-cp312-cp312-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:e6beeea5566092c5e3c4896c6d1d307fb46b1d4bdf3e70c8340b190a69198594"}, - {file = "grpcio-1.71.0-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:d5170929109450a2c031cfe87d6716f2fae39695ad5335d9106ae88cc32dc84c"}, - {file = "grpcio-1.71.0-cp312-cp312-musllinux_1_1_aarch64.whl", hash = "sha256:5b08d03ace7aca7b2fadd4baf291139b4a5f058805a8327bfe9aece7253b6d67"}, - {file = "grpcio-1.71.0-cp312-cp312-musllinux_1_1_i686.whl", hash = "sha256:f903017db76bf9cc2b2d8bdd37bf04b505bbccad6be8a81e1542206875d0e9db"}, - {file = "grpcio-1.71.0-cp312-cp312-musllinux_1_1_x86_64.whl", hash = "sha256:469f42a0b410883185eab4689060a20488a1a0a00f8bbb3cbc1061197b4c5a79"}, - {file = "grpcio-1.71.0-cp312-cp312-win32.whl", hash = "sha256:ad9f30838550695b5eb302add33f21f7301b882937460dd24f24b3cc5a95067a"}, - {file = "grpcio-1.71.0-cp312-cp312-win_amd64.whl", hash = "sha256:652350609332de6dac4ece254e5d7e1ff834e203d6afb769601f286886f6f3a8"}, - {file = "grpcio-1.71.0-cp313-cp313-linux_armv7l.whl", hash = "sha256:cebc1b34ba40a312ab480ccdb396ff3c529377a2fce72c45a741f7215bfe8379"}, - {file = "grpcio-1.71.0-cp313-cp313-macosx_10_14_universal2.whl", hash = "sha256:85da336e3649a3d2171e82f696b5cad2c6231fdd5bad52616476235681bee5b3"}, - {file = "grpcio-1.71.0-cp313-cp313-manylinux_2_17_aarch64.whl", hash = "sha256:f9a412f55bb6e8f3bb000e020dbc1e709627dcb3a56f6431fa7076b4c1aab0db"}, - {file = "grpcio-1.71.0-cp313-cp313-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:47be9584729534660416f6d2a3108aaeac1122f6b5bdbf9fd823e11fe6fbaa29"}, - {file = "grpcio-1.71.0-cp313-cp313-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:7c9c80ac6091c916db81131d50926a93ab162a7e97e4428ffc186b6e80d6dda4"}, - {file = "grpcio-1.71.0-cp313-cp313-musllinux_1_1_aarch64.whl", hash = "sha256:789d5e2a3a15419374b7b45cd680b1e83bbc1e52b9086e49308e2c0b5bbae6e3"}, - {file = "grpcio-1.71.0-cp313-cp313-musllinux_1_1_i686.whl", hash = "sha256:1be857615e26a86d7363e8a163fade914595c81fec962b3d514a4b1e8760467b"}, - {file = "grpcio-1.71.0-cp313-cp313-musllinux_1_1_x86_64.whl", hash = "sha256:a76d39b5fafd79ed604c4be0a869ec3581a172a707e2a8d7a4858cb05a5a7637"}, - {file = "grpcio-1.71.0-cp313-cp313-win32.whl", hash = "sha256:74258dce215cb1995083daa17b379a1a5a87d275387b7ffe137f1d5131e2cfbb"}, - {file = "grpcio-1.71.0-cp313-cp313-win_amd64.whl", hash = "sha256:22c3bc8d488c039a199f7a003a38cb7635db6656fa96437a8accde8322ce2366"}, - {file = "grpcio-1.71.0-cp39-cp39-linux_armv7l.whl", hash = "sha256:c6a0a28450c16809f94e0b5bfe52cabff63e7e4b97b44123ebf77f448534d07d"}, - {file = "grpcio-1.71.0-cp39-cp39-macosx_10_14_universal2.whl", hash = "sha256:a371e6b6a5379d3692cc4ea1cb92754d2a47bdddeee755d3203d1f84ae08e03e"}, - {file = "grpcio-1.71.0-cp39-cp39-manylinux_2_17_aarch64.whl", hash = "sha256:39983a9245d37394fd59de71e88c4b295eb510a3555e0a847d9965088cdbd033"}, - {file = "grpcio-1.71.0-cp39-cp39-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:9182e0063112e55e74ee7584769ec5a0b4f18252c35787f48738627e23a62b97"}, - {file = "grpcio-1.71.0-cp39-cp39-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:693bc706c031aeb848849b9d1c6b63ae6bcc64057984bb91a542332b75aa4c3d"}, - {file = "grpcio-1.71.0-cp39-cp39-musllinux_1_1_aarch64.whl", hash = "sha256:20e8f653abd5ec606be69540f57289274c9ca503ed38388481e98fa396ed0b41"}, - {file = "grpcio-1.71.0-cp39-cp39-musllinux_1_1_i686.whl", hash = "sha256:8700a2a57771cc43ea295296330daaddc0d93c088f0a35cc969292b6db959bf3"}, - {file = "grpcio-1.71.0-cp39-cp39-musllinux_1_1_x86_64.whl", hash = "sha256:d35a95f05a8a2cbe8e02be137740138b3b2ea5f80bd004444e4f9a1ffc511e32"}, - {file = "grpcio-1.71.0-cp39-cp39-win32.whl", hash = "sha256:f9c30c464cb2ddfbc2ddf9400287701270fdc0f14be5f08a1e3939f1e749b455"}, - {file = "grpcio-1.71.0-cp39-cp39-win_amd64.whl", hash = "sha256:63e41b91032f298b3e973b3fa4093cbbc620c875e2da7b93e249d4728b54559a"}, - {file = "grpcio-1.71.0.tar.gz", hash = "sha256:2b85f7820475ad3edec209d3d89a7909ada16caab05d3f2e08a7e8ae3200a55c"}, -] - -[package.extras] -protobuf = ["grpcio-tools (>=1.71.0)"] - [[package]] name = "grpcio-status" version = "1.70.0" @@ -1288,7 +1223,7 @@ description = "Status proto mapping for gRPC" optional = true python-versions = ">=3.8" groups = ["main"] -markers = "extra == \"extra-proxy\" and python_version < \"3.11\"" +markers = "extra == \"extra-proxy\"" files = [ {file = "grpcio_status-1.70.0-py3-none-any.whl", hash = "sha256:fc5a2ae2b9b1c1969cc49f3262676e6854aa2398ec69cb5bd6c47cd501904a85"}, {file = "grpcio_status-1.70.0.tar.gz", hash = "sha256:0e7b42816512433b18b9d764285ff029bde059e9d41f8fe10a60631bd8348101"}, @@ -1299,24 +1234,6 @@ googleapis-common-protos = ">=1.5.5" grpcio = ">=1.70.0" protobuf = ">=5.26.1,<6.0dev" -[[package]] -name = "grpcio-status" -version = "1.71.0" -description = "Status proto mapping for gRPC" -optional = true -python-versions = ">=3.9" -groups = ["main"] -markers = "python_version >= \"3.11\" and extra == \"extra-proxy\"" -files = [ - {file = "grpcio_status-1.71.0-py3-none-any.whl", hash = "sha256:843934ef8c09e3e858952887467f8256aac3910c55f077a359a65b2b3cde3e68"}, - {file = "grpcio_status-1.71.0.tar.gz", hash = "sha256:11405fed67b68f406b3f3c7c5ae5104a79d2d309666d10d61b152e91d28fb968"}, -] - -[package.dependencies] -googleapis-common-protos = ">=1.5.5" -grpcio = ">=1.71.0" -protobuf = ">=5.26.1,<6.0dev" - [[package]] name = "gunicorn" version = "23.0.0" @@ -1399,6 +1316,19 @@ http2 = ["h2 (>=3,<5)"] socks = ["socksio (==1.*)"] zstd = ["zstandard (>=0.18.0)"] +[[package]] +name = "httpx-sse" +version = "0.4.0" +description = "Consume Server-Sent Event (SSE) messages with HTTPX." +optional = true +python-versions = ">=3.8" +groups = ["main"] +markers = "python_version >= \"3.10\" and extra == \"proxy\"" +files = [ + {file = "httpx-sse-0.4.0.tar.gz", hash = "sha256:1e81a3a3070ce322add1d3529ed42eb5f70817f45ed6ec915ab753f961139721"}, + {file = "httpx_sse-0.4.0-py3-none-any.whl", hash = "sha256:f329af6eae57eaa2bdfd962b42524764af68075ea87370a2de920af5341e318f"}, +] + [[package]] name = "huggingface-hub" version = "0.29.3" @@ -1777,6 +1707,34 @@ files = [ {file = "mccabe-0.7.0.tar.gz", hash = "sha256:348e0240c33b60bbdf4e523192ef919f28cb2c3d7d5c7794f74009290f236325"}, ] +[[package]] +name = "mcp" +version = "1.5.0" +description = "Model Context Protocol SDK" +optional = true +python-versions = ">=3.10" +groups = ["main"] +markers = "python_version >= \"3.10\" and extra == \"proxy\"" +files = [ + {file = "mcp-1.5.0-py3-none-any.whl", hash = "sha256:51c3f35ce93cb702f7513c12406bbea9665ef75a08db909200b07da9db641527"}, + {file = "mcp-1.5.0.tar.gz", hash = "sha256:5b2766c05e68e01a2034875e250139839498c61792163a7b221fc170c12f5aa9"}, +] + +[package.dependencies] +anyio = ">=4.5" +httpx = ">=0.27" +httpx-sse = ">=0.4" +pydantic = ">=2.7.2,<3.0.0" +pydantic-settings = ">=2.5.2" +sse-starlette = ">=1.6.1" +starlette = ">=0.27" +uvicorn = ">=0.23.1" + +[package.extras] +cli = ["python-dotenv (>=1.0.0)", "typer (>=0.12.4)"] +rich = ["rich (>=13.9.4)"] +ws = ["websockets (>=15.0.1)"] + [[package]] name = "ml-dtypes" version = "0.4.1" @@ -2729,6 +2687,28 @@ files = [ [package.dependencies] typing-extensions = ">=4.6.0,<4.7.0 || >4.7.0" +[[package]] +name = "pydantic-settings" +version = "2.8.1" +description = "Settings management using Pydantic" +optional = true +python-versions = ">=3.8" +groups = ["main"] +markers = "python_version >= \"3.10\" and extra == \"proxy\"" +files = [ + {file = "pydantic_settings-2.8.1-py3-none-any.whl", hash = "sha256:81942d5ac3d905f7f3ee1a70df5dfb62d5569c12f51a5a647defc1c3d9ee2e9c"}, + {file = "pydantic_settings-2.8.1.tar.gz", hash = "sha256:d5c663dfbe9db9d5e1c646b2e161da12f0d734d422ee56f567d0ea2cee4e8585"}, +] + +[package.dependencies] +pydantic = ">=2.7.0" +python-dotenv = ">=0.21.0" + +[package.extras] +azure-key-vault = ["azure-identity (>=1.16.0)", "azure-keyvault-secrets (>=4.8.0)"] +toml = ["tomli (>=2.0.1)"] +yaml = ["pyyaml (>=6.0.1)"] + [[package]] name = "pyflakes" version = "3.1.0" @@ -3394,6 +3374,27 @@ files = [ {file = "sniffio-1.3.1.tar.gz", hash = "sha256:f4324edc670a0f49750a81b895f35c3adb843cca46f0530f79fc1babb23789dc"}, ] +[[package]] +name = "sse-starlette" +version = "2.1.3" +description = "SSE plugin for Starlette" +optional = true +python-versions = ">=3.8" +groups = ["main"] +markers = "python_version >= \"3.10\" and extra == \"proxy\"" +files = [ + {file = "sse_starlette-2.1.3-py3-none-any.whl", hash = "sha256:8ec846438b4665b9e8c560fcdea6bc8081a3abf7942faa95e5a744999d219772"}, + {file = "sse_starlette-2.1.3.tar.gz", hash = "sha256:9cd27eb35319e1414e3d2558ee7414487f9529ce3b3cf9b21434fd110e017169"}, +] + +[package.dependencies] +anyio = "*" +starlette = "*" +uvicorn = "*" + +[package.extras] +examples = ["fastapi"] + [[package]] name = "starlette" version = "0.44.0" @@ -3999,9 +4000,9 @@ type = ["pytest-mypy"] [extras] extra-proxy = ["azure-identity", "azure-keyvault-secrets", "google-cloud-kms", "prisma", "redisvl", "resend"] -proxy = ["PyJWT", "apscheduler", "backoff", "boto3", "cryptography", "fastapi", "fastapi-sso", "gunicorn", "orjson", "pynacl", "python-multipart", "pyyaml", "rq", "uvicorn", "uvloop", "websockets"] +proxy = ["PyJWT", "apscheduler", "backoff", "boto3", "cryptography", "fastapi", "fastapi-sso", "gunicorn", "mcp", "orjson", "pynacl", "python-multipart", "pyyaml", "rq", "uvicorn", "uvloop", "websockets"] [metadata] lock-version = "2.1" python-versions = ">=3.8.1,<4.0, !=3.9.7" -content-hash = "6850286db1cedd6507c4688767fde27c2f8cc8e657a0a0d792656664eec63d5d" +content-hash = "9c863b11189227a035a9130c8872de44fe7c5e1e32b47569a56af86e3f6570c5" diff --git a/pyproject.toml b/pyproject.toml index f163180dcef..f984b43d675 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -54,6 +54,7 @@ pynacl = {version = "^1.5.0", optional = true} websockets = {version = "^13.1.0", optional = true} boto3 = {version = "1.34.34", optional = true} redisvl = {version = "^0.4.1", optional = true, markers = "python_version >= '3.9' and python_version < '3.14'"} +mcp = {version = "1.5.0", optional = true, python = ">=3.10"} [tool.poetry.extras] proxy = [ @@ -72,7 +73,8 @@ proxy = [ "cryptography", "pynacl", "websockets", - "boto3" + "boto3", + "mcp" ] extra_proxy = [ diff --git a/requirements.txt b/requirements.txt index 621a4d1dd21..c28a2e5cf24 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ # LITELLM PROXY DEPENDENCIES # -anyio==4.4.0 # openai + http req. +anyio==4.5.0 # openai + http req. httpx==0.27.0 # Pin Httpx dependency openai==1.68.2 # openai req. fastapi==0.115.5 # server dep @@ -16,6 +16,7 @@ mangum==0.17.0 # for aws lambda functions pynacl==1.5.0 # for encrypting keys google-cloud-aiplatform==1.47.0 # for vertex ai calls anthropic[vertex]==0.21.3 +mcp==1.5.0 # for MCP server google-generativeai==0.5.0 # for vertex ai calls async_generator==1.10.0 # for async ollama calls langfuse==2.45.0 # for langfuse self-hosted logging @@ -48,7 +49,7 @@ jinja2==3.1.6 # for prompt templates aiohttp==3.10.2 # for network calls aioboto3==12.3.0 # for async sagemaker calls tenacity==8.2.3 # for retrying requests, when litellm.num_retries set -pydantic==2.10.0 # proxy + openai req. +pydantic==2.10.2 # proxy + openai req. jsonschema==4.22.0 # validating json schema websockets==13.1.0 # for realtime API #### \ No newline at end of file diff --git a/tests/litellm/proxy/experimental/mcp_server/test_tool_registry.py b/tests/litellm/proxy/experimental/mcp_server/test_tool_registry.py new file mode 100644 index 00000000000..d5ba9744c7d --- /dev/null +++ b/tests/litellm/proxy/experimental/mcp_server/test_tool_registry.py @@ -0,0 +1,95 @@ +import json +import os +import sys + +import pytest + +sys.path.insert( + 0, os.path.abspath("../../..") +) # Adds the parent directory to the system path + +from litellm.proxy._experimental.mcp_server.tool_registry import MCPToolRegistry + + +# Test handler function +def example_handler(input_data): + return {"result": input_data} + + +def test_register_and_get_tool(): + registry = MCPToolRegistry() + + # Test registering a tool + registry.register_tool( + name="test_tool", + description="A test tool", + input_schema={"type": "object", "properties": {"test": {"type": "string"}}}, + handler=example_handler, + ) + + # Test getting the registered tool + tool = registry.get_tool("test_tool") + assert tool is not None + assert tool.name == "test_tool" + assert tool.description == "A test tool" + assert callable(tool.handler) + + # Test getting non-existent tool + assert registry.get_tool("non_existent") is None + + +def test_list_tools(): + registry = MCPToolRegistry() + + # Register multiple tools + registry.register_tool( + name="tool1", description="Tool 1", input_schema={}, handler=example_handler + ) + registry.register_tool( + name="tool2", description="Tool 2", input_schema={}, handler=example_handler + ) + + # Test listing tools + tools = registry.list_tools() + assert len(tools) == 2 + assert {tool.name for tool in tools} == {"tool1", "tool2"} + + +def test_load_tools_from_config(): + registry = MCPToolRegistry() + + # Test valid config + valid_config = [ + { + "name": "config_tool", + "description": "A tool from config", + "input_schema": {"type": "object"}, + "handler": "test_tool_registry.example_handler", + } + ] + + registry.load_tools_from_config(valid_config) + assert "config_tool" in registry.tools + assert registry.tools["config_tool"].name == "config_tool" + assert registry.tools["config_tool"].description == "A tool from config" + assert callable(registry.tools["config_tool"].handler) + + +def test_tool_execution(): + registry = MCPToolRegistry() + + # Register a tool + registry.register_tool( + name="echo", + description="Echo the input", + input_schema={"type": "object"}, + handler=example_handler, + ) + + # Get and execute the tool + tool = registry.get_tool("echo") + assert tool is not None + + test_input = {"message": "hello"} + result = tool.handler(test_input) + assert result == {"result": test_input} diff --git a/tests/pass_through_tests/test_mcp_routes.py b/tests/pass_through_tests/test_mcp_routes.py new file mode 100644 index 00000000000..687efe6195d --- /dev/null +++ b/tests/pass_through_tests/test_mcp_routes.py @@ -0,0 +1,35 @@ +# Create server parameters for stdio connection +import asyncio +import os + +from langchain_mcp_adapters.tools import load_mcp_tools +from langchain_openai import ChatOpenAI +from langgraph.prebuilt import create_react_agent +from mcp import ClientSession +from mcp.client.sse import sse_client + + +async def main(): + model = ChatOpenAI(model="gpt-4o", api_key="sk-12") + + async with sse_client(url="http://localhost:4000/mcp/") as (read, write): + async with ClientSession(read, write) as session: + # Initialize the connection + print("Initializing session") + await session.initialize() + print("Session initialized") + + # Get tools + print("Loading tools") + tools = await load_mcp_tools(session) + print("Tools loaded") + print(tools) + + # # Create and run the agent + # agent = create_react_agent(model, tools) + # agent_response = await agent.ainvoke({"messages": "what's (3 + 5) x 12?"}) + + +# Run the async function +if __name__ == "__main__": + asyncio.run(main())