diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index e4c0258a..1173b0ef 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -52,6 +52,7 @@ jobs: # Run tests that don't need server - fast feedback! python tests/test_ci_quick.py python -m pytest tests/test_plugins.py -v --tb=short || python tests/test_plugins.py + python -m pytest tests/test_proxy_plugin.py -v --tb=short python tests/test_approaches.py python tests/test_reasoning_simple.py python tests/test_batching.py diff --git a/README.md b/README.md index a5061401..9f44b5b7 100644 --- a/README.md +++ b/README.md @@ -377,7 +377,7 @@ The Model Context Protocol (MCP) plugin enables OptiLLM to connect with MCP serv OptiLLM supports both **local** and **remote** MCP servers through multiple transport methods: - **stdio**: Local servers (traditional) - **SSE**: Remote servers via Server-Sent Events -- **WebSocket**: Remote servers via WebSocket connections +- **Streamable HTTP**: Remote servers via the MCP Streamable HTTP transport #### What is MCP? @@ -449,14 +449,17 @@ The [Model Context Protocol](https://modelcontextprotocol.io/) (MCP) is an open } ``` -**Remote Server (WebSocket) - New Feature:** +**Remote Server (Streamable HTTP):** ```json { "mcpServers": { - "remote-ws": { - "transport": "websocket", - "url": "wss://api.example.com/mcp", - "description": "Remote WebSocket MCP server" + "remote-http": { + "transport": "streamable_http", + "url": "https://api.example.com/mcp", + "headers": { + "Authorization": "Bearer ${API_TOKEN}" + }, + "description": "Remote Streamable HTTP MCP server" } }, "log_level": "INFO" @@ -482,8 +485,8 @@ The [Model Context Protocol](https://modelcontextprotocol.io/) (MCP) is an open "description": "GitHub MCP server" }, "remote-api": { - "transport": "websocket", - "url": "wss://api.company.com/mcp", + "transport": "streamable_http", + "url": "https://api.company.com/mcp", "description": "Company internal MCP server" } }, @@ -495,7 +498,7 @@ The [Model Context Protocol](https://modelcontextprotocol.io/) (MCP) is an open **Common Parameters:** - **Server name**: A unique identifier for the server (e.g., "filesystem", "github") -- **transport**: Transport method - "stdio" (default), "sse", or "websocket" +- **transport**: Transport method - "stdio" (default), "sse", or "streamable_http" - **description** (optional): Description of the server's functionality - **timeout** (optional): Connection timeout in seconds (default: 5.0) @@ -509,8 +512,12 @@ The [Model Context Protocol](https://modelcontextprotocol.io/) (MCP) is an open - **headers** (optional): HTTP headers for authentication - **sse_read_timeout** (optional): SSE read timeout in seconds (default: 300.0) -**websocket Transport (WebSocket):** -- **url**: The WebSocket endpoint URL +**streamable_http Transport (Streamable HTTP):** +- **url**: The MCP endpoint URL +- **headers** (optional): HTTP headers for authentication +- **sse_read_timeout** (optional): Read timeout in seconds (default: 300.0) + +The `websocket` transport was removed in MCP Python SDK 2.x and is no longer supported; use `streamable_http` or `sse` instead. **Environment Variable Expansion:** Headers and other string values support environment variable expansion using `${VARIABLE_NAME}` syntax. This is especially useful for API keys: @@ -536,12 +543,12 @@ You can use any of the [official MCP servers](https://modelcontextprotocol.io/ex - **SQLite**: `@modelcontextprotocol/server-sqlite` - SQLite database access - **Brave Search**: `@modelcontextprotocol/server-brave-search` - Web search capabilities -##### Remote MCP Servers (SSE/WebSocket transport) +##### Remote MCP Servers (SSE/Streamable HTTP transport) Remote servers provide centralized access without requiring local installation: - **GitHub MCP Server**: `https://api.githubcopilot.com/mcp` - Repository management, issue tracking, and code analysis -- **Third-party servers**: Any MCP server that supports SSE or WebSocket protocols +- **Third-party servers**: Any MCP server that supports the SSE or Streamable HTTP transports ##### Example: Comprehensive Configuration @@ -627,7 +634,7 @@ Check this log file for connection issues, tool execution errors, and other diag 2. **Access denied**: For filesystem operations, ensure the paths specified in the configuration are accessible to the process. -**Remote Server Issues (SSE/WebSocket transport):** +**Remote Server Issues (SSE/Streamable HTTP transport):** 3. **Connection timeout**: Remote servers may take longer to connect. Increase the `timeout` value in your configuration. @@ -641,7 +648,7 @@ Check this log file for connection issues, tool execution errors, and other diag 7. **Method not found**: Some servers don't implement all MCP capabilities (tools, resources, prompts). Verify which capabilities the server supports. -8. **Transport not supported**: Ensure you're using a supported transport: "stdio", "sse", or "websocket". +8. **Transport not supported**: Ensure you're using a supported transport: "stdio", "sse", or "streamable_http". **Example: Testing GitHub MCP Connection** diff --git a/optillm/__init__.py b/optillm/__init__.py index 0aad5ab0..bdb79053 100644 --- a/optillm/__init__.py +++ b/optillm/__init__.py @@ -1,5 +1,5 @@ # Version information -__version__ = "0.3.22" +__version__ = "0.4.0" import os as _os diff --git a/optillm/plugins/mcp_plugin.py b/optillm/plugins/mcp_plugin.py index 49983eeb..a3638d1c 100644 --- a/optillm/plugins/mcp_plugin.py +++ b/optillm/plugins/mcp_plugin.py @@ -9,6 +9,7 @@ import json import logging import asyncio +import contextlib import sys import time import re @@ -22,9 +23,16 @@ from mcp import ClientSession, StdioServerParameters from mcp.client.stdio import stdio_client from mcp.client.sse import sse_client -from mcp.client.websocket import websocket_client +from mcp.client.streamable_http import streamable_http_client, create_mcp_http_client import mcp.types as types -from mcp.shared.exceptions import McpError +from mcp.shared.exceptions import MCPError +import httpx2 + +# Transport names accepted in mcp_config.json for Streamable HTTP servers +STREAMABLE_HTTP_TRANSPORTS = ("streamable_http", "streamable-http", "http") +# The WebSocket client transport was removed in mcp 2.x +WEBSOCKET_REMOVED_ERROR = ("WebSocket transport is no longer supported (removed in mcp 2.x); " + "use \"streamable_http\" or \"sse\" instead") # Configure logging LOG_DIR = Path.home() / ".optillm" / "logs" @@ -118,14 +126,14 @@ def find_executable(cmd: str) -> Optional[str]: @dataclass class ServerConfig: """Configuration for a single MCP server""" - # Transport type: "stdio" (default), "sse", or "websocket" + # Transport type: "stdio" (default), "sse", or "streamable_http" transport: str = "stdio" # For stdio transport command: Optional[str] = None args: List[str] = None - # For remote transports (SSE/WebSocket) + # For remote transports (SSE/Streamable HTTP) url: Optional[str] = None headers: Dict[str, str] = None @@ -228,14 +236,22 @@ def create_default_config(self) -> bool: logger.error(f"Error creating default configuration: {e}") return False +def _dump_params(message: Any) -> Any: + """Return a message's params as plain data for logging""" + params = getattr(message, "params", None) + if hasattr(params, "model_dump"): + return params.model_dump(mode="json", by_alias=True, exclude_none=True) + return params + # Create a custom ClientSession that logs all communication class LoggingClientSession(ClientSession): """A ClientSession that logs all communication""" async def send_request(self, *args, **kwargs): """Log and forward requests""" - method = args[0] - params = args[1] if len(args) > 1 else None + request = args[0] + method = getattr(request, "method", request) + params = _dump_params(request) log_mcp_message("REQUEST", method, params) try: @@ -248,8 +264,9 @@ async def send_request(self, *args, **kwargs): async def send_notification(self, *args, **kwargs): """Log and forward notifications""" - method = args[0] - params = args[1] if len(args) > 1 else None + notification = args[0] + method = getattr(notification, "method", notification) + params = _dump_params(notification) log_mcp_message("NOTIFICATION", method, params) try: @@ -258,6 +275,33 @@ async def send_notification(self, *args, **kwargs): log_mcp_message("ERROR", method, error=str(e)) raise +def expand_headers(headers: Dict[str, str]) -> Dict[str, str]: + """Expand ${ENV_VAR} header values from the environment""" + expanded_headers = {} + for key, value in headers.items(): + if isinstance(value, str) and value.startswith("${") and value.endswith("}"): + env_var = value[2:-1] + expanded_value = os.environ.get(env_var) + if expanded_value: + expanded_headers[key] = expanded_value + else: + logger.warning(f"Environment variable {env_var} not found for header {key}") + else: + expanded_headers[key] = value + return expanded_headers + +@contextlib.asynccontextmanager +async def _streamable_http_session(config: ServerConfig): + """Open a logging client session over Streamable HTTP with the configured headers and timeouts""" + http_client = create_mcp_http_client( + headers=expand_headers(config.headers), + timeout=httpx2.Timeout(config.timeout, read=config.sse_read_timeout), + ) + async with http_client: + async with streamable_http_client(config.url, http_client=http_client) as (read_stream, write_stream): + async with LoggingClientSession(read_stream, write_stream) as session: + yield session + class MCPServer: """Represents a connection to an MCP server""" @@ -287,7 +331,7 @@ async def connect_stdio(self, session: LoggingClientSession) -> bool: server_capabilities = result.capabilities # Discover tools if supported - if hasattr(server_capabilities, "tools"): + if getattr(server_capabilities, "tools", None) is not None: self.has_tools_capability = True logger.info(f"Discovering tools for {self.server_name}") try: @@ -295,11 +339,11 @@ async def connect_stdio(self, session: LoggingClientSession) -> bool: self.tools = tools_result.tools logger.info(f"Found {len(self.tools)} tools") logger.debug(f"Tools details: {[t.name for t in self.tools]}") - except McpError as e: + except MCPError as e: logger.warning(f"Failed to list tools: {e}") # Discover resources if supported - if hasattr(server_capabilities, "resources"): + if getattr(server_capabilities, "resources", None) is not None: self.has_resources_capability = True logger.info(f"Discovering resources for {self.server_name}") try: @@ -307,11 +351,11 @@ async def connect_stdio(self, session: LoggingClientSession) -> bool: self.resources = resources_result.resources logger.info(f"Found {len(self.resources)} resources") logger.debug(f"Resources details: {[r.uri for r in self.resources]}") - except McpError as e: + except MCPError as e: logger.warning(f"Failed to list resources: {e}") # Discover prompts if supported - if hasattr(server_capabilities, "prompts"): + if getattr(server_capabilities, "prompts", None) is not None: self.has_prompts_capability = True logger.info(f"Discovering prompts for {self.server_name}") try: @@ -319,7 +363,7 @@ async def connect_stdio(self, session: LoggingClientSession) -> bool: self.prompts = prompts_result.prompts logger.info(f"Found {len(self.prompts)} prompts") logger.debug(f"Prompts details: {[p.name for p in self.prompts]}") - except McpError as e: + except MCPError as e: logger.warning(f"Failed to list prompts: {e}") logger.info(f"Server {self.server_name} capabilities: " @@ -332,6 +376,24 @@ async def connect_stdio(self, session: LoggingClientSession) -> bool: logger.error(traceback.format_exc()) return False + async def connect_streamable_http(self) -> bool: + """Connect to server using Streamable HTTP transport and discover capabilities""" + logger.info(f"Connecting to Streamable HTTP server: {self.server_name}") + logger.debug(f"Streamable HTTP URL: {self.config.url}") + + if not self.config.url: + logger.error(f"Streamable HTTP transport requires URL for server {self.server_name}") + return False + + try: + async with _streamable_http_session(self.config) as session: + return await self.connect_stdio(session) + + except Exception as e: + logger.error(f"Error connecting to Streamable HTTP server {self.server_name}: {e}") + logger.error(traceback.format_exc()) + return False + async def connect_sse(self) -> bool: """Connect to server using SSE transport and discover capabilities""" logger.info(f"Connecting to SSE server: {self.server_name}") @@ -343,18 +405,7 @@ async def connect_sse(self) -> bool: return False try: - # Expand environment variables in headers - expanded_headers = {} - for key, value in self.config.headers.items(): - if isinstance(value, str) and value.startswith("${") and value.endswith("}"): - env_var = value[2:-1] - expanded_value = os.environ.get(env_var) - if expanded_value: - expanded_headers[key] = expanded_value - else: - logger.warning(f"Environment variable {env_var} not found for header {key}") - else: - expanded_headers[key] = value + expanded_headers = expand_headers(self.config.headers) async with sse_client( url=self.config.url, @@ -379,15 +430,8 @@ async def connect_websocket(self) -> bool: logger.error(f"WebSocket transport requires URL for server {self.server_name}") return False - try: - async with websocket_client(self.config.url) as (read_stream, write_stream): - async with LoggingClientSession(read_stream, write_stream) as session: - return await self.connect_stdio(session) - - except Exception as e: - logger.error(f"Error connecting to WebSocket server {self.server_name}: {e}") - logger.error(traceback.format_exc()) - return False + logger.error(f"Server {self.server_name}: {WEBSOCKET_REMOVED_ERROR}") + return False async def connect_stdio_native(self) -> bool: """Connect using stdio transport with local executable""" @@ -479,6 +523,8 @@ async def connect_and_discover(self) -> bool: success = await self.connect_stdio_native() elif self.config.transport == "sse": success = await self.connect_sse() + elif self.config.transport in STREAMABLE_HTTP_TRANSPORTS: + success = await self.connect_streamable_http() elif self.config.transport == "websocket": success = await self.connect_websocket() else: @@ -532,7 +578,7 @@ async def initialize(self) -> bool: "server": server_name, "name": tool.name, "description": tool.description, - "input_schema": tool.inputSchema + "input_schema": tool.input_schema } self.all_tools.append(tool_info) logger.debug(f"Cached tool: {tool_info}") @@ -634,6 +680,10 @@ async def execute_tool_with_session(session: LoggingClientSession, tool_name: st logger.info(f"Calling tool {tool_name} with arguments: {arguments}") result = await session.call_tool(tool_name, arguments) + if not isinstance(result, types.CallToolResult): + # e.g. InputRequiredResult: the tool needs interactive input we can't provide + return {"error": f"Tool {tool_name} returned an unsupported result type: {type(result).__name__}"} + # Process the result content_results = [] for content in result.content: @@ -647,13 +697,13 @@ async def execute_tool_with_session(session: LoggingClientSession, tool_name: st content_results.append({ "type": "image", "data": content.data, - "mimeType": content.mimeType + "mimeType": content.mime_type }) - logger.debug(f"Tool result (image): {content.mimeType}") + logger.debug(f"Tool result (image): {content.mime_type}") return { "result": content_results, - "is_error": result.isError + "is_error": result.is_error } except Exception as e: @@ -692,6 +742,8 @@ async def execute_tool(server_name: str, tool_name: str, arguments: Dict[str, An return await execute_tool_stdio(server_config, tool_name, arguments) elif server_config.transport == "sse": return await execute_tool_sse(server_config, tool_name, arguments) + elif server_config.transport in STREAMABLE_HTTP_TRANSPORTS: + return await execute_tool_streamable_http(server_config, tool_name, arguments) elif server_config.transport == "websocket": return await execute_tool_websocket(server_config, tool_name, arguments) else: @@ -745,18 +797,7 @@ async def execute_tool_sse(server_config: ServerConfig, tool_name: str, argument return {"error": "SSE transport requires URL"} try: - # Expand environment variables in headers - expanded_headers = {} - for key, value in server_config.headers.items(): - if isinstance(value, str) and value.startswith("${") and value.endswith("}"): - env_var = value[2:-1] - expanded_value = os.environ.get(env_var) - if expanded_value: - expanded_headers[key] = expanded_value - else: - logger.warning(f"Environment variable {env_var} not found for header {key}") - else: - expanded_headers[key] = value + expanded_headers = expand_headers(server_config.headers) logger.debug(f" URL: {server_config.url}") logger.debug(f" Headers: {list(expanded_headers.keys())}") @@ -780,17 +821,23 @@ async def execute_tool_websocket(server_config: ServerConfig, tool_name: str, ar if not server_config.url: return {"error": "WebSocket transport requires URL"} + return {"error": WEBSOCKET_REMOVED_ERROR} + +async def execute_tool_streamable_http(server_config: ServerConfig, tool_name: str, arguments: Dict[str, Any]) -> Dict[str, Any]: + """Execute tool using Streamable HTTP transport""" + if not server_config.url: + return {"error": "Streamable HTTP transport requires URL"} + try: logger.debug(f" URL: {server_config.url}") - async with websocket_client(server_config.url) as (read_stream, write_stream): - async with LoggingClientSession(read_stream, write_stream) as session: - return await execute_tool_with_session(session, tool_name, arguments) + async with _streamable_http_session(server_config) as session: + return await execute_tool_with_session(session, tool_name, arguments) except Exception as e: - logger.error(f"Error with WebSocket tool execution: {e}") + logger.error(f"Error with Streamable HTTP tool execution: {e}") logger.error(traceback.format_exc()) - return {"error": f"Error executing tool via WebSocket: {str(e)}"} + return {"error": f"Error executing tool via Streamable HTTP: {str(e)}"} async def run(system_prompt: str, initial_query: str, client, model: str) -> Tuple[str, int]: """ diff --git a/optillm/plugins/proxy/README.md b/optillm/plugins/proxy/README.md index 87c6b262..6c6b72f5 100644 --- a/optillm/plugins/proxy/README.md +++ b/optillm/plugins/proxy/README.md @@ -25,7 +25,7 @@ optillm --version ### 1. Create Configuration -Create `~/.optillm/proxy_config.yaml`: +Create `~/.optillm/proxy_config.yaml`. If it doesn't exist, an empty template is created at `~/.optillm/proxy_config.yaml` and requests go to the server's default client (`--base-url`) until you add providers. ```yaml providers: @@ -68,6 +68,22 @@ optillm optillm --approach proxy --port 8000 ``` +With `--approach proxy`, `/v1/models` lists the models reported by your configured providers (plus any `model_map` aliases), so you don't need to set `--base-url` as well. + +#### Local servers and agent clients + +The proxy passes the original messages through unchanged, including `tools`, assistant `tool_calls` and `tool` results, so coding agents (e.g. Crush) work through it. For a local llama.cpp server: + +```yaml +providers: + - name: llamacpp + base_url: http://localhost:8080/v1 + api_key: none + +timeouts: + request: 300 # local models on long agent prompts can take well over the 30s default +``` + ### 3. Usage Examples #### Method 1: Using --approach proxy (Recommended) @@ -188,7 +204,7 @@ queue: **How it works:** - **Request Timeout**: Each request to a provider has a maximum time limit. If exceeded, the request is cancelled and the next provider is tried. - **Queue Management**: Limits concurrent requests to prevent memory exhaustion. New requests wait up to `queue.timeout` seconds before being rejected. -- **Automatic Failover**: When a provider times out, it's marked unhealthy and the request automatically fails over to the next available provider. +- **Automatic Failover**: When a provider times out or returns a server error, it's marked unhealthy and the request automatically fails over to the next available provider. Rejected requests (4xx) don't mark a provider unhealthy. If every provider is unhealthy, the proxy still retries them before falling back to the server's default client. - **Protection**: Prevents slow backends from causing queue buildup that can crash the proxy server. ### Per-Provider Concurrency Limits @@ -311,11 +327,11 @@ providers: ### Logging -Enable detailed logging for debugging: +Proxy logs follow the server log level (`--log debug` or `OPTILLM_LOG=debug`). To use a different level for the proxy only, set it in the config: ```yaml monitoring: - log_level: DEBUG # Options: DEBUG, INFO, WARNING, ERROR + log_level: DEBUG # Optional. Options: DEBUG, INFO, WARNING, ERROR track_latency: true track_errors: true ``` @@ -361,8 +377,7 @@ When `track_latency` is enabled, the proxy logs: Enable debug logging to see detailed routing decisions: ```bash -export OPTILLM_LOG_LEVEL=DEBUG -python optillm.py +optillm --approach proxy --log debug ``` ## Best Practices diff --git a/optillm/plugins/proxy/client.py b/optillm/plugins/proxy/client.py index f26bf151..5fbc0f5e 100644 --- a/optillm/plugins/proxy/client.py +++ b/optillm/plugins/proxy/client.py @@ -101,6 +101,24 @@ def available_slots(self) -> Optional[int]: # Note: _value is internal but there's no public method to check availability return self._semaphore._value +def list_provider_models(config: Dict) -> List[Dict[str, Any]]: + """ + Collect the models exposed by every configured provider, plus any + model_map aliases, de-duplicated by id. Unreachable providers are skipped. + """ + models = {} + for provider_config in config.get('providers', []): + provider = Provider(provider_config) + try: + for model in provider.client.models.list(timeout=config.get('timeouts', {}).get('connect', 5)).data: + model_dict = model.model_dump() if hasattr(model, 'model_dump') else dict(model) + models.setdefault(model_dict['id'], model_dict) + except Exception as e: + logger.warning(f"Could not list models from provider {provider.name}: {e}") + for alias in provider.model_map: + models.setdefault(alias, {"id": alias, "object": "model", "created": 0, "owned_by": provider.name}) + return list(models.values()) + class ProxyClient: """OpenAI-compatible client that proxies to multiple providers""" @@ -275,6 +293,12 @@ def create(self, **kwargs): if not healthy_providers: logger.warning("No healthy providers, trying fallback providers") healthy_providers = self.proxy_client.fallback_providers + + if not healthy_providers: + # Still prefer the configured providers over the server's + # default client, which may point somewhere else entirely + logger.warning("No fallback providers configured, retrying all providers") + healthy_providers = self.proxy_client.active_providers # Try routing through healthy providers while healthy_providers: @@ -339,7 +363,13 @@ def create(self, **kwargs): except Exception as e: logger.error(f"Provider {provider.name} failed: {e}") errors.append((provider.name, str(e))) - + + # A rejected request (4xx other than timeout/rate limit) + # says nothing about the provider's health + status_code = getattr(e, 'status_code', None) + if status_code is not None and 400 <= status_code < 500 and status_code not in (408, 429): + continue + # Mark provider as unhealthy if self.proxy_client.track_errors: provider.is_healthy = False diff --git a/optillm/plugins/proxy/config.py b/optillm/plugins/proxy/config.py index a827a164..ef31cf3c 100644 --- a/optillm/plugins/proxy/config.py +++ b/optillm/plugins/proxy/config.py @@ -33,11 +33,12 @@ def load(cls, path: str = None, force_reload: bool = False) -> Dict[str, Any]: return cls._cached_config if not path: - # Priority order for config files + # Priority order for config files. The bundled example_config.yaml + # is deliberately not used: its sample providers would silently + # receive traffic meant for the user's own endpoints. config_locations = [ Path.home() / ".optillm" / "proxy_config.yaml", Path.home() / ".optillm" / "proxy_config.yml", - Path(__file__).parent / "example_config.yaml", ] for config_path in config_locations: @@ -152,7 +153,6 @@ def _apply_defaults(config: Dict) -> Dict: # Monitoring defaults monitoring = config['monitoring'] - monitoring.setdefault('log_level', 'INFO') monitoring.setdefault('track_latency', True) monitoring.setdefault('track_errors', True) @@ -253,7 +253,7 @@ def _create_default(path: Path): timeout: 60 # Maximum time in queue (seconds) monitoring: - log_level: INFO + # log_level: DEBUG # Overrides the server --log level for proxy logs track_latency: true track_errors: true @@ -281,7 +281,6 @@ def _get_minimal_config() -> Dict: 'timeout': 60 }, 'monitoring': { - 'log_level': 'INFO', 'track_latency': False, 'track_errors': True } diff --git a/optillm/plugins/proxy/example_config.yaml b/optillm/plugins/proxy/example_config.yaml index a1b89346..af5cea88 100644 --- a/optillm/plugins/proxy/example_config.yaml +++ b/optillm/plugins/proxy/example_config.yaml @@ -61,7 +61,7 @@ routing: # Monitoring and logging monitoring: - log_level: INFO # DEBUG, INFO, WARNING, ERROR + # log_level: DEBUG # Overrides the server --log level for proxy logs track_latency: true # Track request latencies track_errors: true # Track and log errors diff --git a/optillm/plugins/proxy/routing.py b/optillm/plugins/proxy/routing.py index e330ab06..9daeef42 100644 --- a/optillm/plugins/proxy/routing.py +++ b/optillm/plugins/proxy/routing.py @@ -6,11 +6,7 @@ from typing import List, Optional from abc import ABC, abstractmethod -# Configure logging for this module logger = logging.getLogger(__name__) -# Ensure we show debug messages -logging.basicConfig() -logger.setLevel(logging.DEBUG) class Router(ABC): """Abstract base class for routing strategies""" diff --git a/optillm/plugins/proxy_plugin.py b/optillm/plugins/proxy_plugin.py index edcdc474..151b0421 100644 --- a/optillm/plugins/proxy_plugin.py +++ b/optillm/plugins/proxy_plugin.py @@ -5,8 +5,7 @@ with health monitoring, failover, and support for wrapping other approaches. """ import logging -import threading -from typing import Tuple, Optional, Dict +from typing import Tuple, Optional from optillm.plugins.proxy.config import ProxyConfig from optillm.plugins.proxy.client import ProxyClient from optillm.plugins.proxy.approach_handler import ApproachHandler @@ -14,88 +13,32 @@ SLUG = "proxy" logger = logging.getLogger(__name__) -# Configure logging based on environment -import os -log_level = os.environ.get('OPTILLM_LOG_LEVEL', 'INFO') -logging.basicConfig(level=getattr(logging, log_level)) - # Global proxy client cache to maintain state between requests _proxy_client_cache = {} -# Global cache for system message support per provider-model combination -_system_message_support_cache: Dict[str, bool] = {} -_cache_lock = threading.RLock() +def _apply_log_level(config: dict): + """Apply monitoring.log_level from the proxy config, if set, to the proxy loggers.""" + log_level = config.get('monitoring', {}).get('log_level') + if log_level: + level = getattr(logging, str(log_level).upper(), None) + if isinstance(level, int): + logger.setLevel(level) + logging.getLogger('optillm.plugins.proxy').setLevel(level) -def _test_system_message_support(proxy_client, model: str) -> bool: +def _request_messages(system_prompt: str, initial_query: str, messages: Optional[list]) -> list: """ - Test if a model supports system messages by making a minimal test request. - Returns True if supported, False otherwise. + Use the original request messages when available so multi-turn structure, + tool_calls and tool results are preserved. Otherwise rebuild them. """ - try: - # Try a minimal system message request - test_response = proxy_client.chat.completions.create( - model=model, - messages=[ - {"role": "system", "content": "test"}, - {"role": "user", "content": "hi"} - ], - max_tokens=1, # Minimal token generation - temperature=0 - ) - return True - except Exception as e: - error_msg = str(e).lower() - # Check for known system message rejection patterns - if any(pattern in error_msg for pattern in [ - "developer instruction", - "system message", - "not enabled", - "not supported" - ]): - logger.info(f"Model {model} does not support system messages: {str(e)[:100]}") - return False - else: - # If it's a different error, assume system messages are supported - # but something else went wrong (rate limit, timeout, etc.) - logger.debug(f"System message test failed for {model}, assuming supported: {str(e)[:100]}") - return True - -def _get_system_message_support(proxy_client, model: str) -> bool: - """ - Get cached system message support status, testing if not cached. - Thread-safe with locking. - """ - # Create a unique cache key based on model and base_url - cache_key = f"{getattr(proxy_client, '_base_identifier', 'default')}:{model}" - - with _cache_lock: - if cache_key not in _system_message_support_cache: - logger.debug(f"Testing system message support for {model}") - _system_message_support_cache[cache_key] = _test_system_message_support(proxy_client, model) - - return _system_message_support_cache[cache_key] + if messages: + return messages + return [ + {"role": "system", "content": system_prompt}, + {"role": "user", "content": initial_query} + ] -def _format_messages_for_model(system_prompt: str, initial_query: str, - supports_system_messages: bool) -> list: - """ - Format messages based on whether the model supports system messages. - """ - if supports_system_messages: - return [ - {"role": "system", "content": system_prompt}, - {"role": "user", "content": initial_query} - ] - else: - # Merge system prompt into user message - if system_prompt.strip(): - combined_message = f"{system_prompt}\n\nUser: {initial_query}" - else: - combined_message = initial_query - - return [{"role": "user", "content": combined_message}] - -def run(system_prompt: str, initial_query: str, client, model: str, - request_config: dict = None) -> Tuple[str, int]: +def run(system_prompt: str, initial_query: str, client, model: str, + request_config: dict = None, messages: list = None) -> Tuple[str, int]: """ Main proxy plugin entry point. @@ -110,6 +53,7 @@ def run(system_prompt: str, initial_query: str, client, model: str, client: Original OpenAI client (used as fallback) model: Model identifier request_config: Additional request configuration + messages: Original request messages, forwarded as-is for direct routing Returns: Tuple of (response_text, token_count) @@ -117,19 +61,17 @@ def run(system_prompt: str, initial_query: str, client, model: str, try: # Load configuration config = ProxyConfig.load() + _apply_log_level(config) if not config.get('providers'): - logger.warning("No providers configured, falling back to original client") + logger.warning(f"No providers configured in {ProxyConfig._config_path}, falling back to original client") # Strip stream parameter to force complete response api_config = dict(request_config or {}) api_config.pop('stream', None) response = client.chat.completions.create( model=model, - messages=[ - {"role": "system", "content": system_prompt}, - {"role": "user", "content": initial_query} - ], + messages=_request_messages(system_prompt, initial_query, messages), **api_config ) # Return full response dict to preserve all usage information @@ -197,18 +139,10 @@ def run(system_prompt: str, initial_query: str, client, model: str, logger.info(f"Proxy routing approach/plugin: {potential_approach}") return result - # Direct proxy execution with dynamic system message support detection + # Direct proxy execution. ProxyClient detects per-provider system + # message support and merges the system prompt when needed. logger.info(f"Direct proxy routing for model: {model}") - - # Test and cache system message support for this model - supports_system_messages = _get_system_message_support(proxy_client, model) - - # Format messages based on system message support - messages = _format_messages_for_model(system_prompt, initial_query, supports_system_messages) - - if not supports_system_messages: - logger.info(f"Using fallback message formatting for {model} (no system message support)") - + # Strip stream parameter to force complete response # server.py will handle converting to SSE streaming format if needed api_config = dict(request_config or {}) @@ -216,7 +150,7 @@ def run(system_prompt: str, initial_query: str, client, model: str, response = proxy_client.chat.completions.create( model=model, - messages=messages, + messages=_request_messages(system_prompt, initial_query, messages), **api_config ) @@ -234,10 +168,7 @@ def run(system_prompt: str, initial_query: str, client, model: str, response = client.chat.completions.create( model=model, - messages=[ - {"role": "system", "content": system_prompt}, - {"role": "user", "content": initial_query} - ], + messages=_request_messages(system_prompt, initial_query, messages), **api_config ) # Return full response dict to preserve all usage information diff --git a/optillm/server.py b/optillm/server.py index c9697369..8ed597fe 100644 --- a/optillm/server.py +++ b/optillm/server.py @@ -393,7 +393,7 @@ def parse_combined_approach(model: str, known_approaches: list, plugin_approache return operation, approaches, actual_model -def execute_single_approach(approach, system_prompt, initial_query, client, model, request_config: dict = None, request_id: str = None): +def execute_single_approach(approach, system_prompt, initial_query, client, model, request_config: dict = None, request_id: str = None, messages: list = None): if approach in known_approaches: if approach == 'none': # Use the request_config that was already prepared and passed to this function @@ -403,12 +403,15 @@ def execute_single_approach(approach, system_prompt, initial_query, client, mode # Note: 'n' is NOT removed - the none_approach passes it to the client which handles multiple completions kwargs.pop('stream', None) # stream is handled by proxy() - # Reconstruct original messages from system_prompt and initial_query - messages = [] - if system_prompt: - messages.append({"role": "system", "content": system_prompt}) - if initial_query: - messages.append({"role": "user", "content": initial_query}) + # Pass the original messages through when available so multi-turn + # structure, tool_calls and tool results reach the provider intact. + # Otherwise reconstruct them from system_prompt and initial_query. + if not messages: + messages = [] + if system_prompt: + messages.append({"role": "system", "content": system_prompt}) + if initial_query: + messages.append({"role": "user", "content": initial_query}) logger.debug(f"none_approach kwargs: {kwargs}") response = none_approach(original_messages=messages, client=client, model=model, request_id=request_id, **kwargs) @@ -481,9 +484,13 @@ def execute_single_approach(approach, system_prompt, initial_query, client, mode loop.close() else: # For synchronous functions, call directly + plugin_kwargs = {} + if 'messages' in sig.parameters: + # Plugin wants the original request messages (e.g. proxy passthrough) + plugin_kwargs['messages'] = messages if 'request_config' in sig.parameters: # Plugin supports request_config - return plugin_func(system_prompt, initial_query, client, model, request_config=request_config) + return plugin_func(system_prompt, initial_query, client, model, request_config=request_config, **plugin_kwargs) else: # Legacy plugin without request_config support return plugin_func(system_prompt, initial_query, client, model) @@ -509,7 +516,7 @@ async def run_approach(approach): return list(responses), sum(tokens) def execute_n_times(n: int, approaches, operation: str, system_prompt: str, initial_query: str, client: Any, model: str, - request_config: dict = None, request_id: str = None) -> Tuple[Union[str, List[str]], int]: + request_config: dict = None, request_id: str = None, messages: list = None) -> Tuple[Union[str, List[str]], int]: """ Execute the pipeline n times and return n responses. @@ -530,7 +537,7 @@ def execute_n_times(n: int, approaches, operation: str, system_prompt: str, init for _ in range(n): if operation == 'SINGLE': - response, tokens = execute_single_approach(approaches[0], system_prompt, initial_query, client, model, request_config, request_id) + response, tokens = execute_single_approach(approaches[0], system_prompt, initial_query, client, model, request_config, request_id, messages) elif operation == 'AND': response, tokens = execute_combined_approaches(approaches, system_prompt, initial_query, client, model, request_config) elif operation == 'OR': @@ -584,20 +591,44 @@ def generate_streaming_response(final_response, model): # Yield the final message to indicate the stream has ended yield "data: [DONE]\n\n" -def extract_contents(response_obj): - contents = [] - # Handle both single response and list of responses - responses = response_obj if isinstance(response_obj, list) else [response_obj] +def generate_streaming_completion(completion, model, include_usage=False): + """ + Convert a full (non-streamed) chat completion dict into SSE chunks. + + Unlike generate_streaming_response, this keeps tool_calls, reasoning + content, finish_reason and usage so agent clients (Crush, Cline, etc.) + get the same information they would from a streamed provider response. + """ + response_id = completion.get('id') or f"chatcmpl-{int(time.time()*1000)}" + created = completion.get('created') or int(time.time()) + model = completion.get('model') or model + + def chunk(choices, **extra): + return "data: " + json.dumps({ + "id": response_id, + "object": "chat.completion.chunk", + "created": created, + "model": model, + "choices": choices, + **extra, + }) + "\n\n" - for response in responses: - # Extract content from first choice if it exists - if (response.get('choices') and - len(response['choices']) > 0 and - response['choices'][0].get('message') and - response['choices'][0]['message'].get('content')): - contents.append(response['choices'][0]['message']['content']) + for choice in completion.get('choices') or []: + message = choice.get('message') or {} + delta = {"role": message.get('role') or "assistant"} + for key in ('content', 'reasoning_content', 'refusal'): + if message.get(key) is not None: + delta[key] = message[key] + if message.get('tool_calls'): + # Streamed tool call deltas must carry their position in the list + delta['tool_calls'] = [{"index": i, **tool_call} for i, tool_call in enumerate(message['tool_calls'])] + yield chunk([{"index": choice.get('index', 0), "delta": delta, "finish_reason": None}]) + yield chunk([{"index": choice.get('index', 0), "delta": {}, "finish_reason": choice.get('finish_reason') or "stop"}]) + + if include_usage and completion.get('usage'): + yield chunk([], usage=completion['usage']) - return contents + yield "data: [DONE]\n\n" def parse_conversation(messages): system_prompt = "" @@ -605,19 +636,23 @@ def parse_conversation(messages): optillm_approach = None for message in messages: - role = message['role'] - content = message['content'] + role = message.get('role') + # Assistant messages carrying tool_calls have content None + content = message.get('content') or '' # Handle content that could be a list or string if isinstance(content, list): # Extract text content from the list text_content = ' '.join( - item['text'] for item in content + item.get('text', '') for item in content if isinstance(item, dict) and item.get('type') == 'text' ) else: text_content = content + if role == 'assistant' and not text_content: + continue + if role == 'system': system_prompt, optillm_approach = extract_optillm_approach(text_content) elif role == 'user': @@ -686,6 +721,26 @@ def extract_optillm_approach(content): return content, approach return content, None +def strip_optillm_approach_tags(messages): + """ + Return a copy of the request messages with tags removed, + keeping roles, tool_calls and tool results intact for passthrough. + """ + stripped = [] + for message in messages: + message = dict(message) + content = message.get('content') + if isinstance(content, str): + message['content'], _ = extract_optillm_approach(content) + elif isinstance(content, list): + message['content'] = [ + {**item, 'text': extract_optillm_approach(item['text'])[0]} + if isinstance(item, dict) and isinstance(item.get('text'), str) else item + for item in content + ] + stripped.append(message) + return stripped + # Optional API key configuration to secure the proxy @app.before_request def check_api_key(): @@ -726,8 +781,12 @@ def proxy(): max_completion_tokens = data.get('max_completion_tokens') max_tokens = data.get('max_tokens') + # stream_options only applies to our own SSE output; upstream calls are + # never streamed and providers reject stream_options without stream=true + include_usage = bool((data.get('stream_options') or {}).get('include_usage')) + # Explicit keys that we are already handling - explicit_keys = {'stream', 'messages', 'model', 'n', 'response_format', 'max_completion_tokens', 'max_tokens'} + explicit_keys = {'stream', 'stream_options', 'messages', 'model', 'n', 'response_format', 'max_completion_tokens', 'max_tokens'} # Copy the rest into request_config request_config = {k: v for k, v in data.items() if k not in explicit_keys} @@ -761,6 +820,7 @@ def proxy(): # params into every later request). See issue #304. system_prompt, initial_query, message_optillm_approach = parse_conversation(messages) + passthrough_messages = strip_optillm_approach_tags(messages) if message_optillm_approach: optillm_approach = message_optillm_approach @@ -834,7 +894,7 @@ def proxy(): if operation == 'SINGLE' and approaches[0] == 'none': # Pass through the request including the n parameter - result, completion_tokens = execute_single_approach(approaches[0], system_prompt, initial_query, client, model, request_config, request_id) + result, completion_tokens = execute_single_approach(approaches[0], system_prompt, initial_query, client, model, request_config, request_id, passthrough_messages) logger.debug(f'Direct proxy response: {result}') @@ -846,7 +906,7 @@ def proxy(): if stream: if request_id: logger.info(f'Request {request_id}: Completed (streaming response)') - return Response(generate_streaming_response(extract_contents(result), model), content_type='text/event-stream') + return Response(generate_streaming_completion(result, model, include_usage), content_type='text/event-stream') else : if request_id: logger.info(f'Request {request_id}: Completed') @@ -857,7 +917,7 @@ def proxy(): raise ValueError("'none' approach cannot be combined with other approaches") # Handle non-none approaches with n attempts - response, completion_tokens = execute_n_times(n, approaches, operation, system_prompt, initial_query, client, model, request_config, request_id) + response, completion_tokens = execute_n_times(n, approaches, operation, system_prompt, initial_query, client, model, request_config, request_id, passthrough_messages) # Check if the response is a full dict (like from proxy plugin or none approach) if operation == 'SINGLE' and isinstance(response, dict) and 'choices' in response and 'usage' in response: @@ -869,7 +929,7 @@ def proxy(): if stream: if request_id: logger.info(f'Request {request_id}: Completed (streaming response)') - return Response(generate_streaming_response(extract_contents(response), model), content_type='text/event-stream') + return Response(generate_streaming_completion(response, model, include_usage), content_type='text/event-stream') else: if request_id: logger.info(f'Request {request_id}: Completed') @@ -957,7 +1017,18 @@ def proxy_models(): logger.info('Received request to /v1/models') default_client, API_KEY = get_config() try: - if server_config['base_url']: + proxy_models_data = None + if server_config['approach'] == 'proxy': + # List models from the proxy plugin's providers, not --base-url + from optillm.plugins.proxy.config import ProxyConfig + from optillm.plugins.proxy.client import list_provider_models + proxy_config = ProxyConfig.load() + if proxy_config.get('providers'): + proxy_models_data = {"object": "list", "data": list_provider_models(proxy_config)} + + if proxy_models_data is not None: + models_data = proxy_models_data + elif server_config['base_url']: client = OpenAI(api_key=API_KEY, base_url=server_config['base_url']) # For external API, fetch models using the OpenAI client models_response = client.models.list() @@ -1018,7 +1089,7 @@ def parse_args(): ("--return-full-response", "OPTILLM_RETURN_FULL_RESPONSE", bool, False, "Return the full response including the CoT with tags"), ("--host", "OPTILLM_HOST", str, "127.0.0.1", "Host address to bind the server to (use 0.0.0.0 to allow external connections)"), ("--port", "OPTILLM_PORT", int, 8000, "Specify the port to run the proxy"), - ("--log", "OPTILLM_LOG", str, "info", "Specify the logging level", list(logging_levels.keys())), + ("--log", "OPTILLM_LOG", str.lower, "info", "Specify the logging level", list(logging_levels.keys())), ("--launch-gui", "OPTILLM_LAUNCH_GUI", bool, False, "Launch a Gradio chat interface"), ("--plugins-dir", "OPTILLM_PLUGINS_DIR", str, "", "Path to the plugins directory"), ("--log-conversations", "OPTILLM_LOG_CONVERSATIONS", bool, False, "Enable conversation logging with full metadata"), @@ -1263,6 +1334,8 @@ def process_batch_requests(batch_requests): # Set logging level from user request logging_level = server_config['log'] if logging_level in logging_levels.keys(): + # Set the root logger so approach and plugin loggers inherit the level too + logging.getLogger().setLevel(logging_levels[logging_level]) logger.setLevel(logging_levels[logging_level]) # Initialize conversation logger if enabled diff --git a/pyproject.toml b/pyproject.toml index 7e2159ca..28f35b2b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "optillm" -version = "0.3.22" +version = "0.4.0" description = "An optimizing inference proxy for LLMs." readme = "README.md" license = "Apache-2.0" @@ -48,7 +48,7 @@ dependencies = [ "cerebras_cloud_sdk", "outlines[transformers]>=1.2.3", "sentencepiece", - "mcp", + "mcp>=2,<3", "adaptive-classifier", "datasets", "math-verify", diff --git a/requirements.txt b/requirements.txt index 04b2181d..b120cc12 100644 --- a/requirements.txt +++ b/requirements.txt @@ -35,7 +35,7 @@ outlines[transformers]>=1.2.3 sentencepiece adaptive-classifier datasets -mcp +mcp>=2,<3 # MLX support for Apple Silicon optimization mlx-lm>=0.24.0; platform_machine=="arm64" and sys_platform=="darwin" math-verify diff --git a/tests/test_mcp_plugin.py b/tests/test_mcp_plugin.py index 09f56461..621d6f1d 100644 --- a/tests/test_mcp_plugin.py +++ b/tests/test_mcp_plugin.py @@ -22,7 +22,7 @@ from optillm.plugins.mcp_plugin import ( ServerConfig, MCPServer, MCPConfigManager, MCPServerManager, execute_tool, execute_tool_stdio, execute_tool_sse, execute_tool_websocket, - LoggingClientSession, SLUG + execute_tool_streamable_http, execute_tool_with_session, LoggingClientSession, SLUG ) @@ -208,6 +208,14 @@ async def test_connect_websocket_validation(self): result = await server.connect_websocket() assert not result + async def test_connect_streamable_http_validation(self): + """Test Streamable HTTP connection validation""" + config = ServerConfig(transport="streamable_http") # No URL + server = MCPServer("test_server", config) + + result = await server.connect_streamable_http() + assert not result + async def test_connect_and_discover_unsupported_transport(self): """Test unsupported transport type""" config = ServerConfig(transport="invalid") @@ -301,13 +309,36 @@ async def test_execute_tool_sse_no_url(self): assert "error" in result assert "requires URL" in result["error"] - async def test_execute_tool_websocket_no_url(self): - """Test WebSocket tool execution without URL""" - config = ServerConfig(transport="websocket") # No URL + async def test_execute_tool_websocket_unsupported(self): + """WebSocket transport was removed in mcp 2.x and reports a clear error""" + config = ServerConfig(transport="websocket", url="ws://localhost:1234") result = await execute_tool_websocket(config, "test_tool", {}) assert "error" in result + assert "no longer supported" in result["error"] + + async def test_execute_tool_streamable_http_no_url(self): + """Test Streamable HTTP tool execution without URL""" + config = ServerConfig(transport="streamable_http") # No URL + result = await execute_tool_streamable_http(config, "test_tool", {}) + assert "error" in result assert "requires URL" in result["error"] + async def test_execute_tool_with_session_result(self): + """Tool results are read from mcp 2.x snake_case fields""" + import mcp.types as types + session = AsyncMock() + session.call_tool.return_value = types.CallToolResult( + content=[types.TextContent(type="text", text="hello"), + types.ImageContent(type="image", data="aGk=", mime_type="image/png")], + is_error=False, + ) + result = await execute_tool_with_session(session, "test_tool", {}) + assert result == { + "result": [{"type": "text", "text": "hello"}, + {"type": "image", "data": "aGk=", "mimeType": "image/png"}], + "is_error": False, + } + class TestMCPServerManager: """Test MCP server manager functionality""" @@ -410,9 +441,11 @@ def test_required_imports(self): """Test that required modules can be imported""" try: from mcp.client.sse import sse_client - from mcp.client.websocket import websocket_client + from mcp.client.streamable_http import streamable_http_client + from mcp.shared.exceptions import MCPError assert sse_client is not None - assert websocket_client is not None + assert streamable_http_client is not None + assert MCPError is not None except ImportError as e: pytest.fail(f"Required MCP imports failed: {e}") @@ -504,6 +537,7 @@ async def run_async_tests(): 'test_connect_stdio_validation', 'test_connect_sse_validation', 'test_connect_websocket_validation', + 'test_connect_streamable_http_validation', 'test_connect_and_discover_unsupported_transport' ] @@ -523,7 +557,9 @@ async def run_async_tests(): 'test_execute_tool_unsupported_transport', 'test_execute_tool_stdio_no_command', 'test_execute_tool_sse_no_url', - 'test_execute_tool_websocket_no_url' + 'test_execute_tool_websocket_unsupported', + 'test_execute_tool_streamable_http_no_url', + 'test_execute_tool_with_session_result' ] for method_name in tool_methods: diff --git a/tests/test_plugins.py b/tests/test_plugins.py index e7f6b744..ace9d2ef 100644 --- a/tests/test_plugins.py +++ b/tests/test_plugins.py @@ -238,13 +238,16 @@ def test_proxy_plugin_token_counts(): } mock_client.chat.completions.create.return_value = mock_response - # Run the proxy plugin - result, _ = plugin.run( - system_prompt="Test system", - initial_query="Test query", - client=mock_client, - model="test-model" - ) + # Run the proxy plugin with no providers configured so it uses mock_client, + # regardless of any ~/.optillm/proxy_config.yaml on the machine + from unittest.mock import patch + with patch.object(plugin.ProxyConfig, 'load', return_value={'providers': []}): + result, _ = plugin.run( + system_prompt="Test system", + initial_query="Test query", + client=mock_client, + model="test-model" + ) # Verify the result contains all token counts assert isinstance(result, dict), "Result should be a dictionary" diff --git a/tests/test_proxy_plugin.py b/tests/test_proxy_plugin.py new file mode 100644 index 00000000..e075221b --- /dev/null +++ b/tests/test_proxy_plugin.py @@ -0,0 +1,180 @@ +#!/usr/bin/env python3 +""" +Tests for the proxy plugin and agent-style (tool calling) passthrough. +Regression tests for issue #330. +""" + +import json +import os +import sys +import tempfile +import unittest +from pathlib import Path +from unittest.mock import MagicMock, PropertyMock, patch + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +import optillm.server as server +from optillm.plugins.proxy.client import ProxyClient, list_provider_models +from optillm.plugins.proxy.config import ProxyConfig + +TOOLS = [{"type": "function", "function": {"name": "ls", "parameters": {"type": "object", "properties": {}}}}] + +AGENT_MESSAGES = [ + {"role": "system", "content": "You are an agent."}, + {"role": "user", "content": [{"type": "text", "text": "list files"}]}, + {"role": "assistant", "content": None, "tool_calls": [ + {"id": "call_1", "type": "function", "function": {"name": "ls", "arguments": "{}"}}]}, + {"role": "tool", "tool_call_id": "call_1", "content": "a.py"}, +] + +TOOL_CALL_COMPLETION = { + "id": "chatcmpl-1", "object": "chat.completion", "created": 1, "model": "local-model", + "choices": [{"index": 0, "finish_reason": "tool_calls", "message": { + "role": "assistant", "content": None, + "tool_calls": [{"id": "call_2", "type": "function", "function": {"name": "ls", "arguments": "{}"}}]}}], + "usage": {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5}, +} + + +def parse_sse(body): + chunks = [] + for line in body.splitlines(): + if line.startswith("data: ") and line != "data: [DONE]": + chunks.append(json.loads(line[len("data: "):])) + return chunks + + +class TestMessageHandling(unittest.TestCase): + def test_parse_conversation_handles_tool_messages(self): + system_prompt, initial_query, approach = server.parse_conversation(AGENT_MESSAGES) + self.assertEqual(system_prompt, "You are an agent.") + self.assertEqual(initial_query, "User: list files") + self.assertIsNone(approach) + + def test_strip_optillm_approach_tags_keeps_structure(self): + messages = [{"role": "user", "content": "proxy hi"}] + AGENT_MESSAGES[2:] + stripped = server.strip_optillm_approach_tags(messages) + self.assertEqual(stripped[0]["content"], "hi") + self.assertEqual(stripped[1:], AGENT_MESSAGES[2:]) + # Original request is not mutated + self.assertIn("", messages[0]["content"]) + + def test_streaming_completion_keeps_tool_calls_and_usage(self): + chunks = parse_sse("".join(server.generate_streaming_completion(TOOL_CALL_COMPLETION, "m", include_usage=True))) + delta = chunks[0]["choices"][0]["delta"] + self.assertEqual(delta["tool_calls"][0]["index"], 0) + self.assertEqual(delta["tool_calls"][0]["function"]["name"], "ls") + self.assertNotIn("content", delta) + self.assertEqual(chunks[1]["choices"][0]["finish_reason"], "tool_calls") + self.assertEqual(chunks[2]["usage"]["total_tokens"], 5) + + +class TestChatCompletionsPassthrough(unittest.TestCase): + def setUp(self): + self.original_config = server.server_config.copy() + self.app = server.app.test_client() + self.upstream = MagicMock() + self.upstream.chat.completions.create.return_value = TOOL_CALL_COMPLETION + + def tearDown(self): + server.server_config.clear() + server.server_config.update(self.original_config) + + def test_none_approach_forwards_original_messages(self): + with patch.object(server, "get_config", return_value=(self.upstream, "key")): + resp = self.app.post("/v1/chat/completions", json={ + "model": "local-model", "messages": AGENT_MESSAGES, "tools": TOOLS, + "stream": True, "stream_options": {"include_usage": True}}) + + kwargs = self.upstream.chat.completions.create.call_args.kwargs + self.assertEqual(kwargs["messages"][2]["tool_calls"], AGENT_MESSAGES[2]["tool_calls"]) + self.assertEqual(kwargs["messages"][3]["role"], "tool") + self.assertEqual(kwargs["tools"], TOOLS) + self.assertNotIn("stream", kwargs) + self.assertNotIn("stream_options", kwargs) + + chunks = parse_sse(resp.get_data(as_text=True)) + self.assertEqual(chunks[0]["choices"][0]["delta"]["tool_calls"][0]["function"]["name"], "ls") + self.assertEqual(chunks[1]["choices"][0]["finish_reason"], "tool_calls") + + +class TestProxyConfig(unittest.TestCase): + def setUp(self): + ProxyConfig._cached_config = None + self.tmp = tempfile.TemporaryDirectory() + + def tearDown(self): + ProxyConfig._cached_config = None + self.tmp.cleanup() + + def test_does_not_fall_back_to_bundled_example(self): + with patch.object(Path, "home", return_value=Path(self.tmp.name)): + config = ProxyConfig.load() + self.assertEqual(config["providers"], []) + self.assertTrue((Path(self.tmp.name) / ".optillm" / "proxy_config.yaml").exists()) + + +class TestProxyClientFailover(unittest.TestCase): + def make_client(self, fallback): + config = ProxyConfig._validate_config(ProxyConfig._apply_defaults({ + "providers": [{"name": "local", "base_url": "http://localhost:8080/v1", "api_key": "none"}], + "routing": {"health_check": {"enabled": False}}, + })) + client = ProxyClient(config, fallback_client=fallback) + provider = client.providers[0] + provider._client = MagicMock() + return client, provider + + def test_client_error_does_not_mark_provider_unhealthy(self): + client, provider = self.make_client(fallback=None) + error = Exception("bad request") + error.status_code = 400 + provider._client.chat.completions.create.side_effect = error + with self.assertRaises(Exception): + client.chat.completions.create(model="m", messages=[{"role": "user", "content": "hi"}]) + self.assertTrue(provider.is_healthy) + + def test_unhealthy_provider_is_retried_before_default_client(self): + fallback = MagicMock() + client, provider = self.make_client(fallback=fallback) + provider.is_healthy = False + provider._client.chat.completions.create.return_value = "from provider" + result = client.chat.completions.create(model="m", messages=[{"role": "user", "content": "hi"}]) + self.assertEqual(result, "from provider") + fallback.chat.completions.create.assert_not_called() + + +class TestProxyModels(unittest.TestCase): + CONFIG = {"providers": [{"name": "local", "base_url": "http://localhost:8080/v1", "api_key": "none", + "model_map": {"alias": "local-model"}}]} + + def fake_provider_client(self): + model = MagicMock() + model.model_dump.return_value = {"id": "local-model", "object": "model", "created": 0, "owned_by": "llamacpp"} + fake = MagicMock() + fake.models.list.return_value.data = [model] + return fake + + def test_list_provider_models_includes_aliases(self): + with patch("optillm.plugins.proxy.client.Provider.client", new_callable=PropertyMock, + return_value=self.fake_provider_client()): + models = list_provider_models(self.CONFIG) + self.assertEqual([m["id"] for m in models], ["local-model", "alias"]) + + def test_models_endpoint_uses_proxy_providers(self): + original_config = server.server_config.copy() + server.server_config.update({"approach": "proxy", "base_url": ""}) + try: + with patch.object(server, "get_config", return_value=(MagicMock(), "key")), \ + patch.object(ProxyConfig, "load", return_value=self.CONFIG), \ + patch("optillm.plugins.proxy.client.list_provider_models", return_value=[{"id": "local-model"}]): + resp = server.app.test_client().get("/v1/models") + finally: + server.server_config.clear() + server.server_config.update(original_config) + self.assertEqual(resp.get_json()["data"], [{"id": "local-model"}]) + + +if __name__ == "__main__": + unittest.main()