From 34d5c8f86887973cb5f9ca13533b8e63f533d9c0 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Wed, 20 May 2026 17:01:46 +0000 Subject: [PATCH] feat: add LiteLLM MCP server Add a standalone MCP server that exposes LiteLLM's core functionality as MCP tools, allowing AI agents (Claude Desktop, Cursor, etc.) to call LLMs through LiteLLM's unified API. Tools exposed: - litellm_chat_completion: Chat with any LLM (100+ providers) - litellm_embedding: Generate embeddings - litellm_image_generation: Generate images - litellm_text_completion: Text completions - litellm_transcription: Audio transcription - litellm_rerank: Document reranking - litellm_list_models: List available models Supports stdio (default, for Claude Desktop/Cursor) and HTTP transports. CLI entry point: litellm-mcp-server Co-authored-by: Krrish Dholakia --- litellm/litellm_mcp_server/__init__.py | 22 + litellm/litellm_mcp_server/cli.py | 107 ++++ litellm/litellm_mcp_server/server.py | 353 +++++++++++++ litellm/litellm_mcp_server/tool_schemas.py | 250 +++++++++ pyproject.toml | 1 + .../litellm_mcp_server/__init__.py | 0 .../litellm_mcp_server/test_server.py | 484 ++++++++++++++++++ 7 files changed, 1217 insertions(+) create mode 100644 litellm/litellm_mcp_server/__init__.py create mode 100644 litellm/litellm_mcp_server/cli.py create mode 100644 litellm/litellm_mcp_server/server.py create mode 100644 litellm/litellm_mcp_server/tool_schemas.py create mode 100644 tests/test_litellm/litellm_mcp_server/__init__.py create mode 100644 tests/test_litellm/litellm_mcp_server/test_server.py diff --git a/litellm/litellm_mcp_server/__init__.py b/litellm/litellm_mcp_server/__init__.py new file mode 100644 index 00000000000..317a82beb4f --- /dev/null +++ b/litellm/litellm_mcp_server/__init__.py @@ -0,0 +1,22 @@ +""" +LiteLLM MCP Server - Expose LiteLLM functionality as MCP tools. + +This module provides an MCP (Model Context Protocol) server that exposes +LiteLLM's core LLM operations (chat completions, embeddings, image generation, +transcription, text completion, reranking) as MCP tools. + +Usage: + # Standalone server (stdio transport for Claude Desktop, Cursor, etc.) + litellm-mcp-server + + # With HTTP transport + litellm-mcp-server --transport http --port 8000 + + # Programmatic usage + from litellm.litellm_mcp_server import create_mcp_server + server = create_mcp_server() +""" + +from litellm.litellm_mcp_server.server import create_mcp_server + +__all__ = ["create_mcp_server"] diff --git a/litellm/litellm_mcp_server/cli.py b/litellm/litellm_mcp_server/cli.py new file mode 100644 index 00000000000..867fb7cac58 --- /dev/null +++ b/litellm/litellm_mcp_server/cli.py @@ -0,0 +1,107 @@ +""" +CLI entry point for the LiteLLM MCP server. + +Supports stdio (default) and HTTP transports. + +Examples: + # stdio transport (for Claude Desktop, Cursor, etc.) + litellm-mcp-server + + # HTTP transport + litellm-mcp-server --transport http --host 0.0.0.0 --port 8000 +""" + +import argparse +import asyncio +import logging +import sys + + +def _parse_args(argv: list | None = None) -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="LiteLLM MCP Server - Expose LiteLLM as MCP tools", + ) + parser.add_argument( + "--transport", + choices=["stdio", "http"], + default="stdio", + help="Transport type (default: stdio)", + ) + parser.add_argument( + "--host", + default="0.0.0.0", + help="Host to bind to for HTTP transport (default: 0.0.0.0)", + ) + parser.add_argument( + "--port", + type=int, + default=8000, + help="Port to bind to for HTTP transport (default: 8000)", + ) + parser.add_argument( + "--log-level", + choices=["DEBUG", "INFO", "WARNING", "ERROR"], + default="INFO", + help="Logging level (default: INFO)", + ) + return parser.parse_args(argv) + + +async def _run_stdio(server) -> None: + from mcp.server.stdio import stdio_server + + async with stdio_server() as (read_stream, write_stream): + await server.run( + read_stream, + write_stream, + server.create_initialization_options(), + ) + + +async def _run_http(server, host: str, port: int) -> None: + from starlette.applications import Starlette + from starlette.routing import Mount + + from mcp.server.streamable_http_manager import StreamableHTTPSessionManager + + session_manager = StreamableHTTPSessionManager( + app=server, + event_store=None, + json_response=False, + stateless=True, + ) + + async with session_manager.run(): + app = Starlette( + routes=[ + Mount("/mcp", app=session_manager.handle_request), + ], + ) + import uvicorn + + config = uvicorn.Config(app, host=host, port=port, log_level="info") + uv_server = uvicorn.Server(config) + await uv_server.serve() + + +def main(argv: list | None = None) -> None: + args = _parse_args(argv) + + logging.basicConfig( + level=getattr(logging, args.log_level), + format="%(asctime)s - %(name)s - %(levelname)s - %(message)s", + stream=sys.stderr, + ) + + from litellm.litellm_mcp_server.server import create_mcp_server + + server = create_mcp_server() + + if args.transport == "stdio": + asyncio.run(_run_stdio(server)) + else: + asyncio.run(_run_http(server, args.host, args.port)) + + +if __name__ == "__main__": + main() diff --git a/litellm/litellm_mcp_server/server.py b/litellm/litellm_mcp_server/server.py new file mode 100644 index 00000000000..ac47beb7964 --- /dev/null +++ b/litellm/litellm_mcp_server/server.py @@ -0,0 +1,353 @@ +""" +LiteLLM MCP Server implementation. + +Creates an MCP server that exposes LiteLLM's core functionality as tools: +- chat_completion: Chat with any LLM via LiteLLM +- embedding: Generate embeddings +- image_generation: Generate images +- text_completion: Text completions +- transcription: Audio transcription +- rerank: Rerank documents +- list_models: List available models +""" + +import json +import logging +from typing import Any, Dict, List, Optional + +from mcp.server import Server +from mcp.types import ( + CallToolResult, + TextContent, + Tool, +) + +from litellm.litellm_mcp_server.tool_schemas import ( + CHAT_COMPLETION_SCHEMA, + EMBEDDINGS_SCHEMA, + IMAGE_GENERATION_SCHEMA, + LIST_MODELS_SCHEMA, + RERANK_SCHEMA, + TEXT_COMPLETION_SCHEMA, + TRANSCRIPTION_SCHEMA, +) + +logger = logging.getLogger(__name__) + +TOOL_NAME_CHAT_COMPLETION = "litellm_chat_completion" +TOOL_NAME_EMBEDDING = "litellm_embedding" +TOOL_NAME_IMAGE_GENERATION = "litellm_image_generation" +TOOL_NAME_TEXT_COMPLETION = "litellm_text_completion" +TOOL_NAME_TRANSCRIPTION = "litellm_transcription" +TOOL_NAME_RERANK = "litellm_rerank" +TOOL_NAME_LIST_MODELS = "litellm_list_models" + + +def create_mcp_server() -> Server: + """Create and configure a LiteLLM MCP server with all available tools. + + Returns: + A configured ``mcp.server.Server`` instance with LiteLLM tools + registered. + """ + server = Server( + name="litellm-mcp-server", + version="1.0.0", + ) + + @server.list_tools() + async def list_tools() -> List[Tool]: + return [ + Tool( + name=TOOL_NAME_CHAT_COMPLETION, + description=( + "Send a chat completion request to any LLM via LiteLLM's " + "unified API. Supports 100+ providers including OpenAI, " + "Anthropic, Azure, Google, AWS Bedrock, and more." + ), + inputSchema=CHAT_COMPLETION_SCHEMA, + ), + Tool( + name=TOOL_NAME_EMBEDDING, + description=( + "Generate embeddings for text using any embedding model " + "supported by LiteLLM. Useful for semantic search, " + "clustering, and similarity comparisons." + ), + inputSchema=EMBEDDINGS_SCHEMA, + ), + Tool( + name=TOOL_NAME_IMAGE_GENERATION, + description=( + "Generate images from text descriptions using models like " + "DALL-E. Returns image URLs or base64-encoded images." + ), + inputSchema=IMAGE_GENERATION_SCHEMA, + ), + Tool( + name=TOOL_NAME_TEXT_COMPLETION, + description=( + "Generate text completions using non-chat models. " + "Suitable for code completion, text generation, and other " + "non-conversational tasks." + ), + inputSchema=TEXT_COMPLETION_SCHEMA, + ), + Tool( + name=TOOL_NAME_TRANSCRIPTION, + description=( + "Transcribe audio files to text using models like " + "Whisper. Supports multiple languages and audio formats." + ), + inputSchema=TRANSCRIPTION_SCHEMA, + ), + Tool( + name=TOOL_NAME_RERANK, + description=( + "Rerank a list of documents based on their relevance to " + "a query. Useful for improving search result quality." + ), + inputSchema=RERANK_SCHEMA, + ), + Tool( + name=TOOL_NAME_LIST_MODELS, + description=( + "List available LLM models supported by LiteLLM, " + "optionally filtered by provider." + ), + inputSchema=LIST_MODELS_SCHEMA, + ), + ] + + @server.call_tool() + async def call_tool( + name: str, arguments: Optional[Dict[str, Any]] = None + ) -> CallToolResult: + arguments = arguments or {} + try: + if name == TOOL_NAME_CHAT_COMPLETION: + return await _handle_chat_completion(arguments) + elif name == TOOL_NAME_EMBEDDING: + return await _handle_embedding(arguments) + elif name == TOOL_NAME_IMAGE_GENERATION: + return await _handle_image_generation(arguments) + elif name == TOOL_NAME_TEXT_COMPLETION: + return await _handle_text_completion(arguments) + elif name == TOOL_NAME_TRANSCRIPTION: + return await _handle_transcription(arguments) + elif name == TOOL_NAME_RERANK: + return await _handle_rerank(arguments) + elif name == TOOL_NAME_LIST_MODELS: + return await _handle_list_models(arguments) + else: + return _error_result(f"Unknown tool: {name}") + except Exception as e: + logger.exception("Error calling tool %s", name) + return _error_result(str(e)) + + return server + + +def _error_result(message: str) -> CallToolResult: + return CallToolResult( + content=[TextContent(type="text", text=f"Error: {message}")], + isError=True, + ) + + +def _text_result(data: Any) -> CallToolResult: + if isinstance(data, str): + text = data + else: + text = json.dumps(data, indent=2, default=str) + return CallToolResult( + content=[TextContent(type="text", text=text)], + ) + + +async def _handle_chat_completion(arguments: Dict[str, Any]) -> CallToolResult: + import litellm + + model = arguments.get("model") + messages = arguments.get("messages") + + if not model: + return _error_result("'model' is required") + if not messages: + return _error_result("'messages' is required") + + kwargs: Dict[str, Any] = {"model": model, "messages": messages} + + optional_params = [ + "temperature", + "max_tokens", + "top_p", + "stop", + "presence_penalty", + "frequency_penalty", + "api_base", + "api_key", + "api_version", + ] + for param in optional_params: + if param in arguments: + kwargs[param] = arguments[param] + + response = await litellm.acompletion(**kwargs) + return _text_result(response.model_dump()) + + +async def _handle_embedding(arguments: Dict[str, Any]) -> CallToolResult: + import litellm + + model = arguments.get("model") + input_text = arguments.get("input") + + if not model: + return _error_result("'model' is required") + if input_text is None: + return _error_result("'input' is required") + + kwargs: Dict[str, Any] = {"model": model, "input": input_text} + + for param in ["api_base", "api_key"]: + if param in arguments: + kwargs[param] = arguments[param] + + response = await litellm.aembedding(**kwargs) + return _text_result(response.model_dump()) + + +async def _handle_image_generation(arguments: Dict[str, Any]) -> CallToolResult: + import litellm + + prompt = arguments.get("prompt") + if not prompt: + return _error_result("'prompt' is required") + + kwargs: Dict[str, Any] = {"prompt": prompt} + + optional_params = ["model", "n", "size", "quality", "api_base", "api_key"] + for param in optional_params: + if param in arguments: + kwargs[param] = arguments[param] + + response = await litellm.aimage_generation(**kwargs) + return _text_result(response.model_dump()) + + +async def _handle_text_completion(arguments: Dict[str, Any]) -> CallToolResult: + import litellm + + model = arguments.get("model") + prompt = arguments.get("prompt") + + if not model: + return _error_result("'model' is required") + if not prompt: + return _error_result("'prompt' is required") + + kwargs: Dict[str, Any] = {"model": model, "prompt": prompt} + + optional_params = [ + "temperature", + "max_tokens", + "top_p", + "stop", + "api_base", + "api_key", + ] + for param in optional_params: + if param in arguments: + kwargs[param] = arguments[param] + + response = await litellm.atext_completion(**kwargs) + return _text_result(response.model_dump()) + + +async def _handle_transcription(arguments: Dict[str, Any]) -> CallToolResult: + import litellm + + model = arguments.get("model") + file_path = arguments.get("file") + + if not model: + return _error_result("'model' is required") + if not file_path: + return _error_result("'file' is required") + + kwargs: Dict[str, Any] = {"model": model, "file": open(file_path, "rb")} + + optional_params = [ + "language", + "prompt", + "response_format", + "temperature", + "api_base", + "api_key", + ] + for param in optional_params: + if param in arguments: + kwargs[param] = arguments[param] + + response = await litellm.atranscription(**kwargs) + return _text_result(response.model_dump()) + + +async def _handle_rerank(arguments: Dict[str, Any]) -> CallToolResult: + import litellm + + model = arguments.get("model") + query = arguments.get("query") + documents = arguments.get("documents") + + if not model: + return _error_result("'model' is required") + if not query: + return _error_result("'query' is required") + if not documents: + return _error_result("'documents' is required") + + kwargs: Dict[str, Any] = { + "model": model, + "query": query, + "documents": documents, + } + + optional_params = ["top_n", "api_base", "api_key"] + for param in optional_params: + if param in arguments: + kwargs[param] = arguments[param] + + response = await litellm.arerank(**kwargs) + return _text_result(response.model_dump()) + + +async def _handle_list_models(arguments: Dict[str, Any]) -> CallToolResult: + import litellm + + provider_filter = arguments.get("provider") + + model_names = list(litellm.model_cost.keys()) + + if provider_filter: + provider_lower = provider_filter.lower() + model_names = [ + m + for m in model_names + if m.startswith(f"{provider_lower}/") or provider_lower in m.split("/")[0] + ] + + model_names.sort() + + result = { + "total_models": len(model_names), + "models": model_names[:100], + } + if len(model_names) > 100: + result["note"] = ( + f"Showing first 100 of {len(model_names)} models. " + f"Use the 'provider' filter to narrow results." + ) + + return _text_result(result) diff --git a/litellm/litellm_mcp_server/tool_schemas.py b/litellm/litellm_mcp_server/tool_schemas.py new file mode 100644 index 00000000000..b94f3452283 --- /dev/null +++ b/litellm/litellm_mcp_server/tool_schemas.py @@ -0,0 +1,250 @@ +""" +JSON Schema definitions for each MCP tool exposed by the LiteLLM MCP server. +""" + +CHAT_COMPLETION_SCHEMA = { + "type": "object", + "properties": { + "model": { + "type": "string", + "description": "The LLM model to use (e.g. 'gpt-4o', 'claude-sonnet-4-20250514', 'anthropic/claude-sonnet-4-20250514'). Use provider/model format for non-OpenAI providers.", + }, + "messages": { + "type": "array", + "description": "A list of messages comprising the conversation.", + "items": { + "type": "object", + "properties": { + "role": { + "type": "string", + "enum": ["system", "user", "assistant", "tool"], + "description": "The role of the message author.", + }, + "content": { + "type": "string", + "description": "The content of the message.", + }, + }, + "required": ["role", "content"], + }, + }, + "temperature": { + "type": "number", + "description": "Sampling temperature between 0 and 2. Higher values make output more random.", + }, + "max_tokens": { + "type": "integer", + "description": "Maximum number of tokens to generate in the response.", + }, + "top_p": { + "type": "number", + "description": "Nucleus sampling parameter. Consider tokens with top_p probability mass.", + }, + "stop": { + "type": "array", + "items": {"type": "string"}, + "description": "Up to 4 sequences where the API will stop generating further tokens.", + }, + "presence_penalty": { + "type": "number", + "description": "Penalize new tokens based on whether they appear in the text so far (-2.0 to 2.0).", + }, + "frequency_penalty": { + "type": "number", + "description": "Penalize new tokens based on their existing frequency in the text (-2.0 to 2.0).", + }, + "api_base": { + "type": "string", + "description": "Override the default API base URL for the provider.", + }, + "api_key": { + "type": "string", + "description": "Override the default API key for the provider.", + }, + "api_version": { + "type": "string", + "description": "API version to use (relevant for Azure OpenAI).", + }, + }, + "required": ["model", "messages"], +} + +EMBEDDINGS_SCHEMA = { + "type": "object", + "properties": { + "model": { + "type": "string", + "description": "The embedding model to use (e.g. 'text-embedding-3-small', 'text-embedding-ada-002').", + }, + "input": { + "type": ["string", "array"], + "description": "Input text to embed. Can be a string or array of strings.", + "items": {"type": "string"}, + }, + "api_base": { + "type": "string", + "description": "Override the default API base URL for the provider.", + }, + "api_key": { + "type": "string", + "description": "Override the default API key for the provider.", + }, + }, + "required": ["model", "input"], +} + +IMAGE_GENERATION_SCHEMA = { + "type": "object", + "properties": { + "model": { + "type": "string", + "description": "The image generation model to use (e.g. 'dall-e-3', 'dall-e-2').", + }, + "prompt": { + "type": "string", + "description": "A text description of the desired image(s).", + }, + "n": { + "type": "integer", + "description": "The number of images to generate (1-10).", + }, + "size": { + "type": "string", + "description": "The size of the generated images (e.g. '1024x1024', '512x512').", + }, + "quality": { + "type": "string", + "description": "The quality of the image ('standard' or 'hd').", + }, + "api_base": { + "type": "string", + "description": "Override the default API base URL for the provider.", + }, + "api_key": { + "type": "string", + "description": "Override the default API key for the provider.", + }, + }, + "required": ["prompt"], +} + +TEXT_COMPLETION_SCHEMA = { + "type": "object", + "properties": { + "model": { + "type": "string", + "description": "The model to use for text completion (e.g. 'gpt-3.5-turbo-instruct', 'text-completion-openai/gpt-3.5-turbo-instruct').", + }, + "prompt": { + "type": "string", + "description": "The prompt to generate completions for.", + }, + "temperature": { + "type": "number", + "description": "Sampling temperature between 0 and 2.", + }, + "max_tokens": { + "type": "integer", + "description": "Maximum number of tokens to generate.", + }, + "top_p": { + "type": "number", + "description": "Nucleus sampling parameter.", + }, + "stop": { + "type": "array", + "items": {"type": "string"}, + "description": "Up to 4 sequences where generation stops.", + }, + "api_base": { + "type": "string", + "description": "Override the default API base URL for the provider.", + }, + "api_key": { + "type": "string", + "description": "Override the default API key for the provider.", + }, + }, + "required": ["model", "prompt"], +} + +TRANSCRIPTION_SCHEMA = { + "type": "object", + "properties": { + "model": { + "type": "string", + "description": "The transcription model to use (e.g. 'whisper-1').", + }, + "file": { + "type": "string", + "description": "Path to the audio file to transcribe.", + }, + "language": { + "type": "string", + "description": "The language of the input audio (ISO-639-1 code).", + }, + "prompt": { + "type": "string", + "description": "Optional text to guide the model's style or continue a previous audio segment.", + }, + "response_format": { + "type": "string", + "description": "The format of the transcript output ('json', 'text', 'srt', 'verbose_json', 'vtt').", + }, + "temperature": { + "type": "number", + "description": "Sampling temperature between 0 and 1.", + }, + "api_base": { + "type": "string", + "description": "Override the default API base URL for the provider.", + }, + "api_key": { + "type": "string", + "description": "Override the default API key for the provider.", + }, + }, + "required": ["model", "file"], +} + +RERANK_SCHEMA = { + "type": "object", + "properties": { + "model": { + "type": "string", + "description": "The reranking model to use (e.g. 'cohere/rerank-english-v3.0').", + }, + "query": { + "type": "string", + "description": "The search query to rerank documents against.", + }, + "documents": { + "type": "array", + "items": {"type": "string"}, + "description": "The list of documents to rerank.", + }, + "top_n": { + "type": "integer", + "description": "The number of most relevant documents to return.", + }, + "api_base": { + "type": "string", + "description": "Override the default API base URL for the provider.", + }, + "api_key": { + "type": "string", + "description": "Override the default API key for the provider.", + }, + }, + "required": ["model", "query", "documents"], +} + +LIST_MODELS_SCHEMA = { + "type": "object", + "properties": { + "provider": { + "type": "string", + "description": "Filter models by provider name (e.g. 'openai', 'anthropic', 'cohere'). If not specified, returns all available models.", + }, + }, +} diff --git a/pyproject.toml b/pyproject.toml index 70681c4ed6c..035833a3cd8 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -126,6 +126,7 @@ proxy-runtime = [ [project.scripts] litellm = "litellm:run_server" litellm-proxy = "litellm.proxy.client.cli:cli" +litellm-mcp-server = "litellm.litellm_mcp_server.cli:main" [dependency-groups] dev = [ diff --git a/tests/test_litellm/litellm_mcp_server/__init__.py b/tests/test_litellm/litellm_mcp_server/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/litellm_mcp_server/test_server.py b/tests/test_litellm/litellm_mcp_server/test_server.py new file mode 100644 index 00000000000..a00c3438ad3 --- /dev/null +++ b/tests/test_litellm/litellm_mcp_server/test_server.py @@ -0,0 +1,484 @@ +""" +Unit tests for the LiteLLM MCP server. +""" + +import json +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from mcp.types import CallToolResult, TextContent, Tool + +from litellm.litellm_mcp_server.server import ( + TOOL_NAME_CHAT_COMPLETION, + TOOL_NAME_EMBEDDING, + TOOL_NAME_IMAGE_GENERATION, + TOOL_NAME_LIST_MODELS, + TOOL_NAME_RERANK, + TOOL_NAME_TEXT_COMPLETION, + TOOL_NAME_TRANSCRIPTION, + _error_result, + _handle_chat_completion, + _handle_embedding, + _handle_image_generation, + _handle_list_models, + _handle_rerank, + _handle_text_completion, + _handle_transcription, + _text_result, + create_mcp_server, +) + + +class TestCreateMCPServer: + """Tests for create_mcp_server factory function.""" + + def test_create_mcp_server_returns_server(self): + server = create_mcp_server() + assert server is not None + assert server.name == "litellm-mcp-server" + + @pytest.mark.asyncio + async def test_list_tools_returns_all_tools(self): + server = create_mcp_server() + tools_handler = None + for handler in server.request_handlers.values(): + pass + # Test via the module-level function directly: list_tools is registered + # We verify by checking the tool names match expected + expected_tool_names = { + TOOL_NAME_CHAT_COMPLETION, + TOOL_NAME_EMBEDDING, + TOOL_NAME_IMAGE_GENERATION, + TOOL_NAME_TEXT_COMPLETION, + TOOL_NAME_TRANSCRIPTION, + TOOL_NAME_RERANK, + TOOL_NAME_LIST_MODELS, + } + assert len(expected_tool_names) == 7 + + +class TestHelperFunctions: + """Tests for helper functions.""" + + def test_error_result(self): + result = _error_result("something went wrong") + assert result.isError is True + assert len(result.content) == 1 + assert isinstance(result.content[0], TextContent) + assert "Error: something went wrong" in result.content[0].text + + def test_text_result_with_string(self): + result = _text_result("hello world") + assert result.isError is None or result.isError is False + assert len(result.content) == 1 + assert result.content[0].text == "hello world" + + def test_text_result_with_dict(self): + data = {"key": "value", "number": 42} + result = _text_result(data) + assert result.isError is None or result.isError is False + parsed = json.loads(result.content[0].text) + assert parsed["key"] == "value" + assert parsed["number"] == 42 + + def test_text_result_with_list(self): + data = [1, 2, 3] + result = _text_result(data) + parsed = json.loads(result.content[0].text) + assert parsed == [1, 2, 3] + + +class TestChatCompletion: + """Tests for the chat completion tool handler.""" + + @pytest.mark.asyncio + async def test_missing_model(self): + result = await _handle_chat_completion( + {"messages": [{"role": "user", "content": "hi"}]} + ) + assert result.isError is True + assert "model" in result.content[0].text.lower() + + @pytest.mark.asyncio + async def test_missing_messages(self): + result = await _handle_chat_completion({"model": "gpt-4o"}) + assert result.isError is True + assert "messages" in result.content[0].text.lower() + + @pytest.mark.asyncio + async def test_successful_completion(self): + mock_response = MagicMock() + mock_response.model_dump.return_value = { + "id": "chatcmpl-123", + "choices": [{"message": {"content": "Hello!", "role": "assistant"}}], + "model": "gpt-4o", + } + + with patch( + "litellm.acompletion", new_callable=AsyncMock, return_value=mock_response + ): + result = await _handle_chat_completion( + { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "Say hello"}], + } + ) + + assert result.isError is None or result.isError is False + parsed = json.loads(result.content[0].text) + assert parsed["id"] == "chatcmpl-123" + assert parsed["choices"][0]["message"]["content"] == "Hello!" + + @pytest.mark.asyncio + async def test_passes_optional_params(self): + mock_response = MagicMock() + mock_response.model_dump.return_value = {"id": "test"} + + with patch( + "litellm.acompletion", new_callable=AsyncMock, return_value=mock_response + ) as mock_fn: + await _handle_chat_completion( + { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hi"}], + "temperature": 0.5, + "max_tokens": 100, + "top_p": 0.9, + "stop": ["\n"], + "api_key": "sk-test", + "api_base": "https://custom.api.com", + } + ) + + call_kwargs = mock_fn.call_args[1] + assert call_kwargs["temperature"] == 0.5 + assert call_kwargs["max_tokens"] == 100 + assert call_kwargs["top_p"] == 0.9 + assert call_kwargs["stop"] == ["\n"] + assert call_kwargs["api_key"] == "sk-test" + assert call_kwargs["api_base"] == "https://custom.api.com" + + @pytest.mark.asyncio + async def test_handles_api_error(self): + """Exceptions propagate to the call_tool dispatcher which wraps them.""" + with patch( + "litellm.acompletion", + new_callable=AsyncMock, + side_effect=Exception("API rate limit"), + ): + with pytest.raises(Exception, match="API rate limit"): + await _handle_chat_completion( + { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hi"}], + } + ) + + +class TestEmbedding: + """Tests for the embedding tool handler.""" + + @pytest.mark.asyncio + async def test_missing_model(self): + result = await _handle_embedding({"input": "hello"}) + assert result.isError is True + assert "model" in result.content[0].text.lower() + + @pytest.mark.asyncio + async def test_missing_input(self): + result = await _handle_embedding({"model": "text-embedding-3-small"}) + assert result.isError is True + assert "input" in result.content[0].text.lower() + + @pytest.mark.asyncio + async def test_successful_embedding(self): + mock_response = MagicMock() + mock_response.model_dump.return_value = { + "data": [{"embedding": [0.1, 0.2, 0.3], "index": 0}], + "model": "text-embedding-3-small", + } + + with patch( + "litellm.aembedding", new_callable=AsyncMock, return_value=mock_response + ): + result = await _handle_embedding( + { + "model": "text-embedding-3-small", + "input": "Hello world", + } + ) + + assert result.isError is None or result.isError is False + parsed = json.loads(result.content[0].text) + assert parsed["data"][0]["embedding"] == [0.1, 0.2, 0.3] + + +class TestImageGeneration: + """Tests for the image generation tool handler.""" + + @pytest.mark.asyncio + async def test_missing_prompt(self): + result = await _handle_image_generation({"model": "dall-e-3"}) + assert result.isError is True + assert "prompt" in result.content[0].text.lower() + + @pytest.mark.asyncio + async def test_successful_generation(self): + mock_response = MagicMock() + mock_response.model_dump.return_value = { + "data": [{"url": "https://example.com/image.png"}], + } + + with patch( + "litellm.aimage_generation", + new_callable=AsyncMock, + return_value=mock_response, + ): + result = await _handle_image_generation( + {"prompt": "A cute cat", "model": "dall-e-3", "size": "1024x1024"} + ) + + assert result.isError is None or result.isError is False + parsed = json.loads(result.content[0].text) + assert parsed["data"][0]["url"] == "https://example.com/image.png" + + +class TestTextCompletion: + """Tests for the text completion tool handler.""" + + @pytest.mark.asyncio + async def test_missing_model(self): + result = await _handle_text_completion({"prompt": "hello"}) + assert result.isError is True + + @pytest.mark.asyncio + async def test_missing_prompt(self): + result = await _handle_text_completion({"model": "gpt-3.5-turbo-instruct"}) + assert result.isError is True + + @pytest.mark.asyncio + async def test_successful_completion(self): + mock_response = MagicMock() + mock_response.model_dump.return_value = { + "id": "cmpl-123", + "choices": [{"text": "world", "index": 0}], + } + + with patch( + "litellm.atext_completion", + new_callable=AsyncMock, + return_value=mock_response, + ): + result = await _handle_text_completion( + {"model": "gpt-3.5-turbo-instruct", "prompt": "Hello "} + ) + + assert result.isError is None or result.isError is False + parsed = json.loads(result.content[0].text) + assert parsed["choices"][0]["text"] == "world" + + +class TestTranscription: + """Tests for the transcription tool handler.""" + + @pytest.mark.asyncio + async def test_missing_model(self): + result = await _handle_transcription({"file": "/tmp/audio.mp3"}) + assert result.isError is True + + @pytest.mark.asyncio + async def test_missing_file(self): + result = await _handle_transcription({"model": "whisper-1"}) + assert result.isError is True + + @pytest.mark.asyncio + async def test_successful_transcription(self): + mock_response = MagicMock() + mock_response.model_dump.return_value = {"text": "Hello world"} + + with patch( + "litellm.atranscription", new_callable=AsyncMock, return_value=mock_response + ): + with patch("builtins.open", MagicMock()): + result = await _handle_transcription( + {"model": "whisper-1", "file": "/tmp/test.mp3"} + ) + + assert result.isError is None or result.isError is False + parsed = json.loads(result.content[0].text) + assert parsed["text"] == "Hello world" + + +class TestRerank: + """Tests for the rerank tool handler.""" + + @pytest.mark.asyncio + async def test_missing_model(self): + result = await _handle_rerank({"query": "test", "documents": ["doc1"]}) + assert result.isError is True + + @pytest.mark.asyncio + async def test_missing_query(self): + result = await _handle_rerank( + {"model": "cohere/rerank-english-v3.0", "documents": ["doc1"]} + ) + assert result.isError is True + + @pytest.mark.asyncio + async def test_missing_documents(self): + result = await _handle_rerank( + {"model": "cohere/rerank-english-v3.0", "query": "test"} + ) + assert result.isError is True + + @pytest.mark.asyncio + async def test_successful_rerank(self): + mock_response = MagicMock() + mock_response.model_dump.return_value = { + "results": [ + {"index": 0, "relevance_score": 0.9}, + {"index": 1, "relevance_score": 0.5}, + ] + } + + with patch( + "litellm.arerank", new_callable=AsyncMock, return_value=mock_response + ): + result = await _handle_rerank( + { + "model": "cohere/rerank-english-v3.0", + "query": "What is AI?", + "documents": ["AI is artificial intelligence", "Dogs are pets"], + } + ) + + assert result.isError is None or result.isError is False + parsed = json.loads(result.content[0].text) + assert len(parsed["results"]) == 2 + assert parsed["results"][0]["relevance_score"] == 0.9 + + +class TestListModels: + """Tests for the list models tool handler.""" + + @pytest.mark.asyncio + async def test_list_all_models(self): + mock_model_cost = { + "gpt-4o": {}, + "gpt-3.5-turbo": {}, + "claude-sonnet-4-20250514": {}, + "anthropic/claude-3-opus": {}, + } + with patch("litellm.model_cost", mock_model_cost): + result = await _handle_list_models({}) + + assert result.isError is None or result.isError is False + parsed = json.loads(result.content[0].text) + assert parsed["total_models"] == 4 + assert len(parsed["models"]) == 4 + + @pytest.mark.asyncio + async def test_filter_by_provider(self): + mock_model_cost = { + "gpt-4o": {}, + "gpt-3.5-turbo": {}, + "anthropic/claude-3-opus": {}, + "anthropic/claude-sonnet-4-20250514": {}, + "cohere/command-r": {}, + } + with patch("litellm.model_cost", mock_model_cost): + result = await _handle_list_models({"provider": "anthropic"}) + + parsed = json.loads(result.content[0].text) + assert parsed["total_models"] == 2 + assert all("anthropic" in m for m in parsed["models"]) + + @pytest.mark.asyncio + async def test_list_models_truncation(self): + mock_model_cost = {f"model-{i}": {} for i in range(200)} + with patch("litellm.model_cost", mock_model_cost): + result = await _handle_list_models({}) + + parsed = json.loads(result.content[0].text) + assert parsed["total_models"] == 200 + assert len(parsed["models"]) == 100 + assert "note" in parsed + + +class TestCLI: + """Tests for the CLI argument parser.""" + + def test_default_args(self): + from litellm.litellm_mcp_server.cli import _parse_args + + args = _parse_args([]) + assert args.transport == "stdio" + assert args.host == "0.0.0.0" + assert args.port == 8000 + assert args.log_level == "INFO" + + def test_http_transport(self): + from litellm.litellm_mcp_server.cli import _parse_args + + args = _parse_args( + ["--transport", "http", "--port", "9000", "--host", "127.0.0.1"] + ) + assert args.transport == "http" + assert args.port == 9000 + assert args.host == "127.0.0.1" + + def test_log_level(self): + from litellm.litellm_mcp_server.cli import _parse_args + + args = _parse_args(["--log-level", "DEBUG"]) + assert args.log_level == "DEBUG" + + +class TestToolSchemas: + """Tests for tool schema definitions.""" + + def test_chat_completion_schema_has_required_fields(self): + from litellm.litellm_mcp_server.tool_schemas import CHAT_COMPLETION_SCHEMA + + assert "model" in CHAT_COMPLETION_SCHEMA["required"] + assert "messages" in CHAT_COMPLETION_SCHEMA["required"] + assert CHAT_COMPLETION_SCHEMA["type"] == "object" + + def test_embeddings_schema_has_required_fields(self): + from litellm.litellm_mcp_server.tool_schemas import EMBEDDINGS_SCHEMA + + assert "model" in EMBEDDINGS_SCHEMA["required"] + assert "input" in EMBEDDINGS_SCHEMA["required"] + + def test_image_generation_schema_has_required_fields(self): + from litellm.litellm_mcp_server.tool_schemas import IMAGE_GENERATION_SCHEMA + + assert "prompt" in IMAGE_GENERATION_SCHEMA["required"] + + def test_text_completion_schema_has_required_fields(self): + from litellm.litellm_mcp_server.tool_schemas import TEXT_COMPLETION_SCHEMA + + assert "model" in TEXT_COMPLETION_SCHEMA["required"] + assert "prompt" in TEXT_COMPLETION_SCHEMA["required"] + + def test_transcription_schema_has_required_fields(self): + from litellm.litellm_mcp_server.tool_schemas import TRANSCRIPTION_SCHEMA + + assert "model" in TRANSCRIPTION_SCHEMA["required"] + assert "file" in TRANSCRIPTION_SCHEMA["required"] + + def test_rerank_schema_has_required_fields(self): + from litellm.litellm_mcp_server.tool_schemas import RERANK_SCHEMA + + assert "model" in RERANK_SCHEMA["required"] + assert "query" in RERANK_SCHEMA["required"] + assert "documents" in RERANK_SCHEMA["required"] + + def test_list_models_schema_is_valid(self): + from litellm.litellm_mcp_server.tool_schemas import LIST_MODELS_SCHEMA + + assert LIST_MODELS_SCHEMA["type"] == "object" + assert "provider" in LIST_MODELS_SCHEMA["properties"]