From 212f5e482242a88c7e697ad3c9b7e8084e3eb7e0 Mon Sep 17 00:00:00 2001 From: Simo Lin Date: Sat, 25 Oct 2025 21:26:28 -0700 Subject: [PATCH] [router] MCP Manager Refactoring - Flat Architecture with Connection Pooling (#12097) --- sgl-router/README.md | 115 +++ .../py_src/sglang_router/router_args.py | 9 + sgl-router/src/app_context.rs | 46 + sgl-router/src/config/builder.rs | 40 +- sgl-router/src/config/types.rs | 5 + sgl-router/src/core/job_queue.rs | 87 +- sgl-router/src/core/workflow/mod.rs | 5 +- .../core/workflow/steps/mcp_registration.rs | 304 ++++++ sgl-router/src/core/workflow/steps/mod.rs | 6 + sgl-router/src/lib.rs | 5 + sgl-router/src/main.rs | 4 + sgl-router/src/mcp/client_manager.rs | 556 ----------- sgl-router/src/mcp/config.rs | 479 ++++++++++ sgl-router/src/mcp/connection_pool.rs | 448 +++++++++ sgl-router/src/mcp/inventory.rs | 620 ++++++++++++ sgl-router/src/mcp/manager.rs | 893 ++++++++++++++++++ sgl-router/src/mcp/mod.rs | 15 +- sgl-router/src/mcp/proxy.rs | 253 +++++ sgl-router/src/routers/factory.rs | 9 +- .../src/routers/grpc/responses/handlers.rs | 28 +- .../src/routers/grpc/responses/tool_loop.rs | 13 +- sgl-router/src/routers/grpc/router.rs | 30 +- sgl-router/src/routers/openai/mcp.rs | 58 +- sgl-router/src/routers/openai/router.rs | 82 +- sgl-router/src/routers/openai/streaming.rs | 29 +- sgl-router/src/server.rs | 28 +- sgl-router/src/service_discovery.rs | 1 + sgl-router/tests/api_endpoints_test.rs | 4 +- sgl-router/tests/common/mod.rs | 130 ++- sgl-router/tests/common/test_app.rs | 56 ++ sgl-router/tests/mcp_test.rs | 145 ++- sgl-router/tests/request_formats_test.rs | 2 +- sgl-router/tests/responses_api_test.rs | 24 +- sgl-router/tests/streaming_tests.rs | 2 +- sgl-router/tests/test_openai_routing.rs | 155 +-- sgl-router/tests/test_pd_routing.rs | 4 +- 36 files changed, 3838 insertions(+), 852 deletions(-) create mode 100644 sgl-router/src/core/workflow/steps/mcp_registration.rs delete mode 100644 sgl-router/src/mcp/client_manager.rs create mode 100644 sgl-router/src/mcp/connection_pool.rs create mode 100644 sgl-router/src/mcp/inventory.rs create mode 100644 sgl-router/src/mcp/manager.rs create mode 100644 sgl-router/src/mcp/proxy.rs diff --git a/sgl-router/README.md b/sgl-router/README.md index da0890edf..5ba7887cd 100644 --- a/sgl-router/README.md +++ b/sgl-router/README.md @@ -206,6 +206,121 @@ python3 -m sglang_router.launch_router \ - Provide exactly one `--worker-urls` entry per router instance. - The Rust binary supports the same flags (`./target/release/sglang-router --backend openai ...`). +### MCP Integration +The SGL Model Gateway provides native Model Context Protocol (MCP) client integration, enabling tool calling across STDIO, SSE, and Streamable transports. MCP servers are configured via a YAML configuration file and registered at startup through the workflow engine. + +#### Basic Usage +```bash +# Rust binary +./target/release/sglang-router \ + --mcp-config-path /path/to/mcp-config.yaml \ + --worker-urls http://worker1:8000 + +# Python launcher +python3 -m sglang_router.launch_router \ + --mcp-config-path /path/to/mcp-config.yaml \ + --worker-urls http://worker1:8000 +``` + +#### MCP Configuration File +Create an MCP configuration file to define servers, transports, and connection settings: + +```yaml +servers: + - name: "filesystem" + command: "npx" + args: ["-y", "@modelcontextprotocol/server-filesystem", "/tmp"] + required: false + + - name: "github" + url: "https://api.github.com/mcp" + token: "ghp_xxxxx" + transport: "sse" + required: false + + - name: "custom-tools" + url: "https://tools.example.com/mcp" + transport: "streamable" + required: true + +pool: + max_connections: 100 + idle_timeout: 300 # seconds + +proxy: + http: "http://proxy.internal:8080" + https: "https://proxy.internal:8443" + no_proxy: "localhost,127.0.0.1,*.internal" + +inventory: + enable_refresh: true + tool_ttl: 300 # seconds - how long tools are considered fresh + refresh_interval: 300 # seconds - background refresh interval +``` + +#### Configuration Options + +**Server Configuration** (`servers` array): +- `name`: Unique identifier for the MCP server +- `command` + `args`: For STDIO transport (local process execution) +- `url`: For SSE or Streamable transports (HTTP/HTTPS endpoints) +- `token`: Optional authentication token for HTTP-based transports +- `transport`: Protocol type (`"sse"` or `"streamable"`; STDIO is inferred from `command`) +- `required`: If `true`, router fails to start if server is unreachable (default: `false`) +- `envs`: Environment variables for STDIO processes (optional) +- `proxy`: Per-server proxy override (set to `null` to bypass global proxy) + +**Connection Pool** (`pool`): +- `max_connections`: Maximum pooled connections for dynamic servers (default: 100) +- `idle_timeout`: Idle connection timeout in seconds before cleanup (default: 300) + +**Proxy Configuration** (`proxy`): +- `http`/`https`: Proxy URLs for MCP server connections (not LLM traffic) +- `no_proxy`: Comma-separated hosts to exclude from proxying (supports wildcards) +- **Note**: Proxy settings are currently ignored for `streamable` transport. Use STDIO or SSE transports if proxy support is required. + +**Inventory Settings** (`inventory`): +- `enable_refresh`: Enable automatic background refresh of tool inventory (default: true) +- `tool_ttl`: Tool cache TTL in seconds - how long tools are considered fresh (default: 300) +- `refresh_interval`: Background refresh interval in seconds - proactive inventory refresh (default: 300) + +#### Transport Types + +**STDIO** (Local Process): +```yaml +name: "local-tools" +command: "python" +args: ["-m", "my_mcp_server"] +envs: + API_KEY: "secret" + DEBUG: "true" +``` + +**SSE** (Server-Sent Events): +```yaml +name: "remote-sse" +url: "https://mcp.example.com/events" +token: "bearer-token" +transport: "sse" +``` + +**Streamable** (Bidirectional Streaming): +```yaml +name: "streaming-tools" +url: "https://mcp.example.com/stream" +transport: "streamable" +required: true +``` + +#### Server Lifecycle +- MCP servers are registered via the workflow engine with retry logic (100 attempts, 2-hour timeout for STDIO servers) +- Discovery phase identifies tools, prompts, and resources +- Tool inventory is cached with configurable TTL and periodic refresh +- Failed optional servers log warnings; required servers halt startup +- Static servers (from config) are permanent; dynamic servers (per-request) use connection pooling + +Check Prometheus metrics for MCP activity (`mcp_*` metrics) and workflow job status via the admin API. + ### Python Launcher (Router + Workers) Launch router and SGLang worker processes together; `launch_server` spins up workers (HTTP or gRPC) and the router in one shot. ```bash diff --git a/sgl-router/py_src/sglang_router/router_args.py b/sgl-router/py_src/sglang_router/router_args.py index 1de752d8f..53f804e04 100644 --- a/sgl-router/py_src/sglang_router/router_args.py +++ b/sgl-router/py_src/sglang_router/router_args.py @@ -94,6 +94,8 @@ class RouterArgs: tokenizer_cache_l1_max_memory: int = 50 * 1024 * 1024 # 50MB reasoning_parser: Optional[str] = None tool_call_parser: Optional[str] = None + # MCP server configuration + mcp_config_path: Optional[str] = None # Backend selection backend: str = "sglang" # History backend configuration @@ -512,6 +514,13 @@ class RouterArgs: default=None, help="Specify the parser for handling tool-call interactions", ) + # MCP server configuration + parser.add_argument( + f"--{prefix}mcp-config-path", + type=str, + default=None, + help="Path to MCP (Model Context Protocol) server configuration file", + ) # Backend selection parser.add_argument( f"--{prefix}backend", diff --git a/sgl-router/src/app_context.rs b/sgl-router/src/app_context.rs index b4ad9b1d4..a0f96ea4b 100644 --- a/sgl-router/src/app_context.rs +++ b/sgl-router/src/app_context.rs @@ -13,6 +13,7 @@ use crate::{ create_storage, SharedConversationItemStorage, SharedConversationStorage, SharedResponseStorage, }, + mcp::McpManager, middleware::TokenBucket, policies::PolicyRegistry, reasoning_parser::ParserFactory as ReasoningParserFactory, @@ -56,6 +57,7 @@ pub struct AppContext { pub configured_tool_parser: Option, pub worker_job_queue: Arc>>, pub workflow_engine: Arc>>, + pub mcp_manager: Arc>>, } pub struct AppContextBuilder { @@ -74,6 +76,7 @@ pub struct AppContextBuilder { load_monitor: Option>, worker_job_queue: Option>>>, workflow_engine: Option>>>, + mcp_manager: Option>>>, } impl AppContext { @@ -112,6 +115,7 @@ impl AppContextBuilder { load_monitor: None, worker_job_queue: None, workflow_engine: None, + mcp_manager: None, } } @@ -196,6 +200,11 @@ impl AppContextBuilder { self } + pub fn mcp_manager(mut self, mcp_manager: Arc>>) -> Self { + self.mcp_manager = Some(mcp_manager); + self + } + pub fn build(self) -> Result { let router_config = self .router_config @@ -235,6 +244,9 @@ impl AppContextBuilder { workflow_engine: self .workflow_engine .ok_or(AppContextBuildError("workflow_engine"))?, + mcp_manager: self + .mcp_manager + .ok_or(AppContextBuildError("mcp_manager"))?, }) } @@ -256,6 +268,8 @@ impl AppContextBuilder { .with_load_monitor(&router_config) .with_worker_job_queue() .with_workflow_engine() + .with_mcp_manager(&router_config) + .await? .router_config(router_config)) } @@ -457,6 +471,38 @@ impl AppContextBuilder { self.workflow_engine = Some(Arc::new(OnceLock::new())); self } + + /// Create and initialize MCP manager with empty config + /// + /// This initializes the MCP manager with an empty config and default settings. + /// MCP servers will be registered later via the InitializeMcpServers job. + async fn with_mcp_manager(mut self, _router_config: &RouterConfig) -> Result { + // Create OnceLock container + let mcp_manager_lock = Arc::new(OnceLock::new()); + + // Always create with empty config and defaults + info!("Initializing MCP manager with empty config and default settings (5 min TTL, 100 max connections)"); + + let empty_config = crate::mcp::McpConfig { + servers: Vec::new(), + pool: Default::default(), + proxy: None, + warmup: Vec::new(), + inventory: Default::default(), + }; + + let manager = McpManager::with_defaults(empty_config) + .await + .map_err(|e| format!("Failed to initialize MCP manager with defaults: {}", e))?; + + // Store the initialized manager in the OnceLock + mcp_manager_lock + .set(Arc::new(manager)) + .map_err(|_| "Failed to set MCP manager in OnceLock".to_string())?; + + self.mcp_manager = Some(mcp_manager_lock); + Ok(self) + } } impl Default for AppContextBuilder { diff --git a/sgl-router/src/config/builder.rs b/sgl-router/src/config/builder.rs index e6aa66d87..95dd5f197 100644 --- a/sgl-router/src/config/builder.rs +++ b/sgl-router/src/config/builder.rs @@ -3,7 +3,7 @@ use super::{ HistoryBackend, MetricsConfig, OracleConfig, PolicyConfig, RetryConfig, RouterConfig, RoutingMode, TokenizerCacheConfig, }; -use crate::core::ConnectionMode; +use crate::{core::ConnectionMode, mcp::McpConfig}; /// Builder for RouterConfig that wraps the config itself /// This eliminates field duplication and stays in sync automatically @@ -14,6 +14,7 @@ pub struct RouterConfigBuilder { client_cert_path: Option, client_key_path: Option, ca_cert_paths: Vec, + mcp_config_path: Option, } impl RouterConfigBuilder { @@ -29,6 +30,7 @@ impl RouterConfigBuilder { client_cert_path: None, client_key_path: None, ca_cert_paths: Vec::new(), + mcp_config_path: None, } } @@ -620,6 +622,21 @@ impl RouterConfigBuilder { self } + // ==================== MCP Configuration ==================== + + /// Set MCP server configuration file path + /// The config file will be loaded during build() + pub fn mcp_config_path>(mut self, path: S) -> Self { + self.mcp_config_path = Some(path.into()); + self + } + + /// Set MCP server configuration file path if Some + pub fn maybe_mcp_config_path(mut self, path: Option>) -> Self { + self.mcp_config_path = path.map(|p| p.into()); + self + } + // ==================== Builder Methods ==================== /// Build the RouterConfig, validating if requested @@ -637,6 +654,9 @@ impl RouterConfigBuilder { // Read mTLS certificates from paths if provided self = self.read_mtls_certificates()?; + // Read MCP config from path if provided + self = self.read_mcp_config()?; + let config: RouterConfig = self.into(); if validate { config.validate()?; @@ -695,6 +715,24 @@ impl RouterConfigBuilder { Ok(self) } + + /// Internal method to read MCP config from path + fn read_mcp_config(mut self) -> ConfigResult { + if let Some(mcp_config_path) = &self.mcp_config_path { + let contents = std::fs::read_to_string(mcp_config_path).map_err(|e| { + ConfigError::ValidationFailed { + reason: format!("Failed to read MCP config from {}: {}", mcp_config_path, e), + } + })?; + let mcp_config: McpConfig = + serde_yaml::from_str(&contents).map_err(|e| ConfigError::ValidationFailed { + reason: format!("Failed to parse MCP config from {}: {}", mcp_config_path, e), + })?; + self.config.mcp_config = Some(mcp_config); + } + + Ok(self) + } } impl From for RouterConfig { diff --git a/sgl-router/src/config/types.rs b/sgl-router/src/config/types.rs index 1be1b515f..d621cfa37 100644 --- a/sgl-router/src/config/types.rs +++ b/sgl-router/src/config/types.rs @@ -93,6 +93,10 @@ pub struct RouterConfig { /// Loaded from ca_cert_paths during config creation #[serde(default)] pub ca_certificates: Vec>, + /// MCP server configuration (loaded from mcp_config_path during config creation) + /// This is loaded from the config file path and stored here for runtime use + #[serde(skip)] + pub mcp_config: Option, } /// Tokenizer cache configuration @@ -508,6 +512,7 @@ impl Default for RouterConfig { tokenizer_cache: TokenizerCacheConfig::default(), client_identity: None, ca_certificates: vec![], + mcp_config: None, } } } diff --git a/sgl-router/src/core/job_queue.rs b/sgl-router/src/core/job_queue.rs index ac396c7ec..c8e9edf0d 100644 --- a/sgl-router/src/core/job_queue.rs +++ b/sgl-router/src/core/job_queue.rs @@ -17,9 +17,10 @@ use crate::{ app_context::AppContext, config::{RouterConfig, RoutingMode}, core::workflow::{ - steps::WorkerRemovalRequest, WorkflowContext, WorkflowEngine, WorkflowId, - WorkflowInstanceId, WorkflowStatus, + steps::{McpServerConfigRequest, WorkerRemovalRequest}, + WorkflowContext, WorkflowEngine, WorkflowId, WorkflowInstanceId, WorkflowStatus, }, + mcp::McpConfig, metrics::RouterMetrics, protocols::worker_spec::{JobStatus, WorkerConfigRequest}, }; @@ -30,6 +31,8 @@ pub enum Job { AddWorker { config: Box }, RemoveWorker { url: String }, InitializeWorkersFromConfig { router_config: Box }, + InitializeMcpServers { mcp_config: Box }, + RegisterMcpServer { config: Box }, } impl Job { @@ -39,15 +42,19 @@ impl Job { Job::AddWorker { .. } => "AddWorker", Job::RemoveWorker { .. } => "RemoveWorker", Job::InitializeWorkersFromConfig { .. } => "InitializeWorkersFromConfig", + Job::InitializeMcpServers { .. } => "InitializeMcpServers", + Job::RegisterMcpServer { .. } => "RegisterMcpServer", } } - /// Get worker URL for logging + /// Get worker URL or MCP server name for logging pub fn worker_url(&self) -> &str { match self { Job::AddWorker { config } => &config.url, Job::RemoveWorker { url } => url, Job::InitializeWorkersFromConfig { .. } => "startup", + Job::InitializeMcpServers { .. } => "startup", + Job::RegisterMcpServer { config } => &config.name, } } } @@ -421,6 +428,64 @@ impl JobQueue { Ok(format!("Submitted {} AddWorker jobs", worker_count)) } + Job::InitializeMcpServers { mcp_config } => { + let mut server_count = 0; + + debug!( + "Creating RegisterMcpServer jobs for {} MCP servers from config", + mcp_config.servers.len() + ); + + // Submit RegisterMcpServer jobs for each server in the config + for server_config in &mcp_config.servers { + let mcp_server_request = McpServerConfigRequest { + name: server_config.name.clone(), + config: server_config.clone(), + }; + + let job = Job::RegisterMcpServer { + config: Box::new(mcp_server_request), + }; + + if let Some(queue) = context.worker_job_queue.get() { + queue.submit(job).await.map_err(|e| { + format!( + "Failed to submit RegisterMcpServer job for '{}': {}", + server_config.name, e + ) + })?; + server_count += 1; + } else { + return Err("JobQueue not available".to_string()); + } + } + + Ok(format!("Submitted {} RegisterMcpServer jobs", server_count)) + } + Job::RegisterMcpServer { config } => { + let engine = context + .workflow_engine + .get() + .ok_or_else(|| "Workflow engine not initialized".to_string())?; + + let instance_id = + Self::start_mcp_registration_workflow(engine, config, context).await?; + + debug!( + "Started MCP registration workflow for {} (instance: {})", + config.name, instance_id + ); + + let timeout_duration = Duration::from_secs(7200 + 30); // 2hr + margin + + Self::wait_for_workflow_completion( + engine, + instance_id, + &config.name, + timeout_duration, + ) + .await + } } } @@ -461,6 +526,22 @@ impl JobQueue { .map_err(|e| format!("Failed to start worker removal workflow: {:?}", e)) } + /// Start MCP server registration workflow + async fn start_mcp_registration_workflow( + engine: &Arc, + config: &McpServerConfigRequest, + context: &Arc, + ) -> Result { + let mut workflow_context = WorkflowContext::new(WorkflowInstanceId::new()); + workflow_context.set("mcp_server_config", config.clone()); + workflow_context.set_arc("app_context", Arc::clone(context)); + + engine + .start_workflow(WorkflowId::new("mcp_registration"), workflow_context) + .await + .map_err(|e| format!("Failed to start MCP registration workflow: {:?}", e)) + } + /// Wait for workflow completion with adaptive polling async fn wait_for_workflow_completion( engine: &Arc, diff --git a/sgl-router/src/core/workflow/mod.rs b/sgl-router/src/core/workflow/mod.rs index f1c3293a2..320ffe053 100644 --- a/sgl-router/src/core/workflow/mod.rs +++ b/sgl-router/src/core/workflow/mod.rs @@ -14,5 +14,8 @@ pub use engine::WorkflowEngine; pub use event::{EventBus, EventSubscriber, LoggingSubscriber, WorkflowEvent}; pub use executor::{FunctionStep, StepExecutor}; pub use state::WorkflowStateStore; -pub use steps::{create_worker_registration_workflow, create_worker_removal_workflow}; +pub use steps::{ + create_mcp_registration_workflow, create_worker_registration_workflow, + create_worker_removal_workflow, +}; pub use types::*; diff --git a/sgl-router/src/core/workflow/steps/mcp_registration.rs b/sgl-router/src/core/workflow/steps/mcp_registration.rs new file mode 100644 index 000000000..516749159 --- /dev/null +++ b/sgl-router/src/core/workflow/steps/mcp_registration.rs @@ -0,0 +1,304 @@ +//! MCP server registration workflow steps +//! +//! Each step is atomic and performs a single operation in the MCP server registration process. +//! Updated for flat manager architecture - single McpManager manages all clients directly. +//! +//! Workflow order: +//! 1. ConnectMcpServer - Establish connection to MCP server using McpManager::connect_server() +//! 2. DiscoverMcpInventory - Discover and cache inventory using McpManager::load_server_inventory() +//! 3. RegisterMcpServer - Register McpClient in McpManager's client map + +use std::{sync::Arc, time::Duration}; + +use async_trait::async_trait; +use rmcp::{service::RunningService, RoleClient}; +use tracing::{debug, error, info, warn}; + +use crate::{ + app_context::AppContext, + core::workflow::*, + mcp::{config::McpServerConfig, manager::McpManager}, +}; + +/// MCP server connection configuration +#[derive(Debug, Clone)] +pub struct McpServerConfigRequest { + /// Server name (unique identifier) + pub name: String, + /// Server configuration (transport, proxy, etc.) + pub config: McpServerConfig, +} + +impl McpServerConfigRequest { + /// Check if this server is required for router startup + pub fn is_required(&self) -> bool { + self.config.required + } +} + +/// Step 1: Connect to MCP server +/// +/// This step establishes a connection to the MCP server using the flat manager architecture. +/// The connection is retried aggressively (100 attempts) with a long timeout (2 hours) +/// to handle slow-starting servers or network issues. +pub struct ConnectMcpServerStep; + +#[async_trait] +impl StepExecutor for ConnectMcpServerStep { + async fn execute(&self, context: &mut WorkflowContext) -> WorkflowResult { + let config_request: Arc = context + .get("mcp_server_config") + .ok_or_else(|| WorkflowError::ContextValueNotFound("mcp_server_config".to_string()))?; + let app_context: Arc = context + .get("app_context") + .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; + + debug!("Connecting to MCP server: {}", config_request.name); + + // Get proxy config from router_config if available, otherwise fall back to env + let proxy_config = app_context + .router_config + .mcp_config + .as_ref() + .and_then(|cfg| cfg.proxy.as_ref()); + + // Connect to MCP server + let client = McpManager::connect_server(&config_request.config, proxy_config) + .await + .map_err(|e| WorkflowError::StepFailed { + step_id: StepId::new("connect_mcp_server"), + message: format!( + "Failed to connect to MCP server {}: {}", + config_request.name, e + ), + })?; + + info!( + "Successfully connected to MCP server: {}", + config_request.name + ); + + // Store client in context (context.set() will wrap in Arc) + context.set("mcp_client", client); + + Ok(StepResult::Success) + } + + fn is_retryable(&self, _error: &WorkflowError) -> bool { + true // Connection failures are retryable + } +} + +/// Step 2: Discover MCP inventory (tools, prompts, resources) +/// +/// This step queries the MCP server for its capabilities using McpManager::load_server_inventory(). +/// - Tools: Available function calls +/// - Prompts: Reusable prompt templates +/// - Resources: Accessible files/data +pub struct DiscoverMcpInventoryStep; + +#[async_trait] +impl StepExecutor for DiscoverMcpInventoryStep { + async fn execute(&self, context: &mut WorkflowContext) -> WorkflowResult { + use rmcp::{service::RunningService, RoleClient}; + + let config_request: Arc = context + .get("mcp_server_config") + .ok_or_else(|| WorkflowError::ContextValueNotFound("mcp_server_config".to_string()))?; + let app_context: Arc = context + .get("app_context") + .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; + let mcp_client: Arc> = context + .get("mcp_client") + .ok_or_else(|| WorkflowError::ContextValueNotFound("mcp_client".to_string()))?; + + debug!( + "Discovering inventory for MCP server: {}", + config_request.name + ); + + // Get shared ToolInventory from McpManager + let mcp_manager = + app_context + .mcp_manager + .get() + .ok_or_else(|| WorkflowError::StepFailed { + step_id: StepId::new("discover_mcp_inventory"), + message: "MCP manager not initialized".to_string(), + })?; + + let inventory = mcp_manager.inventory(); + + // Use the public load_server_inventory method + McpManager::load_server_inventory(&inventory, &config_request.name, &mcp_client).await; + + info!("Completed inventory discovery for {}", config_request.name); + + Ok(StepResult::Success) + } + + fn is_retryable(&self, _error: &WorkflowError) -> bool { + true // Discovery failures are retryable + } +} + +/// Step 3: Register MCP server in manager +/// +/// This step adds the MCP client to the McpManager's client map so it can be +/// used for tool calls and inventory management. +pub struct RegisterMcpServerStep; + +#[async_trait] +impl StepExecutor for RegisterMcpServerStep { + async fn execute(&self, context: &mut WorkflowContext) -> WorkflowResult { + use rmcp::{service::RunningService, RoleClient}; + + let config_request: Arc = context + .get("mcp_server_config") + .ok_or_else(|| WorkflowError::ContextValueNotFound("mcp_server_config".to_string()))?; + let app_context: Arc = context + .get("app_context") + .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; + let mcp_client: Arc> = context + .get("mcp_client") + .ok_or_else(|| WorkflowError::ContextValueNotFound("mcp_client".to_string()))?; + + debug!("Registering MCP server: {}", config_request.name); + + // Get MCP manager from app context + let mcp_manager = + app_context + .mcp_manager + .get() + .ok_or_else(|| WorkflowError::StepFailed { + step_id: StepId::new("register_mcp_server"), + message: "MCP manager not initialized".to_string(), + })?; + + // Register the client in the manager's client map + mcp_manager.register_static_server(config_request.name.clone(), mcp_client); + + info!("Registered MCP server: {}", config_request.name); + + Ok(StepResult::Success) + } + + fn is_retryable(&self, _error: &WorkflowError) -> bool { + false // Registration is a simple operation, not retryable + } +} + +/// Step 4: Validate registration based on required flag +/// +/// This step checks if the server is marked as required. If the server is required +/// but wasn't successfully registered (client not in context), this step fails the workflow. +/// For optional servers, this step always succeeds, allowing the workflow to complete +/// even if earlier steps failed. +pub struct ValidateRegistrationStep; + +#[async_trait] +impl StepExecutor for ValidateRegistrationStep { + async fn execute(&self, context: &mut WorkflowContext) -> WorkflowResult { + let config_request: Arc = context + .get("mcp_server_config") + .ok_or_else(|| WorkflowError::ContextValueNotFound("mcp_server_config".to_string()))?; + + let client_registered = context + .get::>>("mcp_client") + .is_some(); + + if client_registered { + info!( + "MCP server '{}' registered successfully", + config_request.name + ); + return Ok(StepResult::Success); + } + + if config_request.is_required() { + error!( + "Required MCP server '{}' failed to register", + config_request.name + ); + Err(WorkflowError::StepFailed { + step_id: StepId::new("validate_registration"), + message: format!( + "Required MCP server '{}' failed to register", + config_request.name + ), + }) + } else { + warn!( + "Optional MCP server '{}' failed to register, continuing workflow", + config_request.name + ); + Ok(StepResult::Success) + } + } + + fn is_retryable(&self, _error: &WorkflowError) -> bool { + false + } +} + +/// Create MCP server registration workflow +/// +/// This workflow adapts its failure behavior based on the `required` field in the server config: +/// - If `required == true`: Uses FailWorkflow - router startup fails if server cannot be reached +/// - If `required == false` (default): Uses ContinueNextStep - logs warning but continues +/// +/// Workflow configuration: +/// - ConnectMcpServer: 100 retries, 2hr timeout (aggressive retry for slow servers) +/// - DiscoverMcpInventory: 3 retries, 10s timeout (discovery + caching) +/// - RegisterMcpServer: No retry, 5s timeout (fast registration) +/// - ValidateRegistration: Final validation step +pub fn create_mcp_registration_workflow() -> WorkflowDefinition { + WorkflowDefinition::new("mcp_registration", "MCP Server Registration") + .add_step( + StepDefinition::new( + "connect_mcp_server", + "Connect to MCP Server", + Arc::new(ConnectMcpServerStep), + ) + .with_retry(RetryPolicy { + max_attempts: 100, + backoff: BackoffStrategy::Linear { + increment: Duration::from_secs(1), + max: Duration::from_secs(5), + }, + }) + .with_timeout(Duration::from_secs(7200)) // 2 hours + .with_failure_action(FailureAction::ContinueNextStep), + ) + .add_step( + StepDefinition::new( + "discover_mcp_inventory", + "Discover and Cache MCP Inventory", + Arc::new(DiscoverMcpInventoryStep), + ) + .with_retry(RetryPolicy { + max_attempts: 3, + backoff: BackoffStrategy::Fixed(Duration::from_secs(1)), + }) + .with_timeout(Duration::from_secs(10)) + .with_failure_action(FailureAction::ContinueNextStep), + ) + .add_step( + StepDefinition::new( + "register_mcp_server", + "Register MCP Server", + Arc::new(RegisterMcpServerStep), + ) + .with_timeout(Duration::from_secs(5)) + .with_failure_action(FailureAction::ContinueNextStep), + ) + .add_step( + StepDefinition::new( + "validate_registration", + "Validate MCP Registration", + Arc::new(ValidateRegistrationStep), + ) + .with_timeout(Duration::from_secs(1)) + .with_failure_action(FailureAction::FailWorkflow), + ) +} diff --git a/sgl-router/src/core/workflow/steps/mod.rs b/sgl-router/src/core/workflow/steps/mod.rs index 9de153023..10b5c5f40 100644 --- a/sgl-router/src/core/workflow/steps/mod.rs +++ b/sgl-router/src/core/workflow/steps/mod.rs @@ -3,11 +3,17 @@ //! This module contains concrete step implementations for various workflows: //! - Worker registration and activation //! - Worker removal +//! - MCP server registration //! - Future: Tokenizer fetching, LoRA updates, etc. +pub mod mcp_registration; pub mod worker_registration; pub mod worker_removal; +pub use mcp_registration::{ + create_mcp_registration_workflow, ConnectMcpServerStep, DiscoverMcpInventoryStep, + McpServerConfigRequest, RegisterMcpServerStep, ValidateRegistrationStep, +}; pub use worker_registration::{ create_worker_registration_workflow, ActivateWorkerStep, CreateWorkerStep, DetectConnectionModeStep, DiscoverMetadataStep, RegisterWorkerStep, UpdatePoliciesStep, diff --git a/sgl-router/src/lib.rs b/sgl-router/src/lib.rs index 054386a3f..be6aac9db 100644 --- a/sgl-router/src/lib.rs +++ b/sgl-router/src/lib.rs @@ -205,6 +205,7 @@ struct Router { tokenizer_cache_l1_max_memory: usize, reasoning_parser: Option, tool_call_parser: Option, + mcp_config_path: Option, backend: BackendType, history_backend: HistoryBackendType, oracle_config: Option, @@ -360,6 +361,7 @@ impl Router { .maybe_oracle(oracle) .maybe_reasoning_parser(self.reasoning_parser.as_ref()) .maybe_tool_call_parser(self.tool_call_parser.as_ref()) + .maybe_mcp_config_path(self.mcp_config_path.as_ref()) .dp_aware(self.dp_aware) .retries(!self.disable_retries) .circuit_breaker(!self.disable_circuit_breaker) @@ -440,6 +442,7 @@ impl Router { tokenizer_cache_l1_max_memory = 52428800, reasoning_parser = None, tool_call_parser = None, + mcp_config_path = None, backend = BackendType::Sglang, history_backend = HistoryBackendType::Memory, oracle_config = None, @@ -512,6 +515,7 @@ impl Router { tokenizer_cache_l1_max_memory: usize, reasoning_parser: Option, tool_call_parser: Option, + mcp_config_path: Option, backend: BackendType, history_backend: HistoryBackendType, oracle_config: Option, @@ -598,6 +602,7 @@ impl Router { tokenizer_cache_l1_max_memory, reasoning_parser, tool_call_parser, + mcp_config_path, backend, history_backend, oracle_config, diff --git a/sgl-router/src/main.rs b/sgl-router/src/main.rs index 12d9a278e..12987cb41 100644 --- a/sgl-router/src/main.rs +++ b/sgl-router/src/main.rs @@ -315,6 +315,9 @@ struct CliArgs { #[arg(long)] tool_call_parser: Option, + + #[arg(long)] + mcp_config_path: Option, } enum OracleConnectSource { @@ -594,6 +597,7 @@ impl CliArgs { .maybe_oracle(oracle) .maybe_reasoning_parser(self.reasoning_parser.as_ref()) .maybe_tool_call_parser(self.tool_call_parser.as_ref()) + .maybe_mcp_config_path(self.mcp_config_path.as_ref()) .dp_aware(self.dp_aware) .retries(!self.disable_retries) .circuit_breaker(!self.disable_circuit_breaker) diff --git a/sgl-router/src/mcp/client_manager.rs b/sgl-router/src/mcp/client_manager.rs deleted file mode 100644 index 91c958e95..000000000 --- a/sgl-router/src/mcp/client_manager.rs +++ /dev/null @@ -1,556 +0,0 @@ -use std::{borrow::Cow, collections::HashMap, time::Duration}; - -use backoff::ExponentialBackoffBuilder; -use dashmap::DashMap; -use rmcp::{ - model::{ - CallToolRequestParam, GetPromptRequestParam, GetPromptResult, Prompt, - ReadResourceRequestParam, ReadResourceResult, Resource, Tool as McpTool, - }, - service::RunningService, - transport::{ - sse_client::SseClientConfig, streamable_http_client::StreamableHttpClientTransportConfig, - ConfigureCommandExt, SseClientTransport, StreamableHttpClientTransport, TokioChildProcess, - }, - RoleClient, ServiceExt, -}; -use serde::{Deserialize, Serialize}; - -use crate::mcp::{ - config::{McpConfig, McpServerConfig, McpTransport}, - error::{McpError, McpResult}, -}; - -/// Information about an available tool -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ToolInfo { - pub name: String, - pub description: String, - pub server: String, - pub parameters: Option, -} - -/// Information about an available prompt -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct PromptInfo { - pub name: String, - pub description: Option, - pub server: String, - pub arguments: Option>, -} - -/// Information about an available resource -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ResourceInfo { - pub uri: String, - pub name: String, - pub description: Option, - pub mime_type: Option, - pub server: String, -} - -/// Manages MCP client connections and tool execution -pub struct McpClientManager { - /// Map of server_name -> MCP client - clients: HashMap>, - /// Map of tool_name -> (server_name, tool_definition) - tools: DashMap, - /// Map of prompt_name -> (server_name, prompt_definition) - prompts: DashMap, - /// Map of resource_uri -> (server_name, resource_definition) - resources: DashMap, -} - -impl McpClientManager { - /// Create a new manager and connect to all configured servers - pub async fn new(config: McpConfig) -> McpResult { - let mut mgr = Self { - clients: HashMap::new(), - tools: DashMap::new(), - prompts: DashMap::new(), - resources: DashMap::new(), - }; - - for server_config in config.servers { - match Self::connect_server(&server_config).await { - Ok(client) => { - mgr.load_server_inventory(&server_config.name, &client) - .await; - mgr.clients.insert(server_config.name.clone(), client); - } - Err(e) => { - tracing::error!( - "Failed to connect to server '{}': {}", - server_config.name, - e - ); - } - } - } - - if mgr.clients.is_empty() { - return Err(McpError::ConnectionFailed( - "Failed to connect to any MCP servers".to_string(), - )); - } - Ok(mgr) - } - - /// Discover and cache tools/prompts/resources for a connected server - async fn load_server_inventory( - &self, - server_name: &str, - client: &RunningService, - ) { - // Tools - match client.peer().list_all_tools().await { - Ok(ts) => { - tracing::info!("Discovered {} tools from '{}'", ts.len(), server_name); - for t in ts { - if self.tools.contains_key(t.name.as_ref()) { - tracing::warn!( - "Tool '{}' from server '{}' is overwriting an existing tool.", - &t.name, - server_name - ); - } - self.tools - .insert(t.name.to_string(), (server_name.to_string(), t)); - } - } - Err(e) => tracing::warn!("Failed to list tools from '{}': {}", server_name, e), - } - - // Prompts - match client.peer().list_all_prompts().await { - Ok(ps) => { - tracing::info!("Discovered {} prompts from '{}'", ps.len(), server_name); - for p in ps { - if self.prompts.contains_key(&p.name) { - tracing::warn!( - "Prompt '{}' from server '{}' is overwriting an existing prompt.", - &p.name, - server_name - ); - } - self.prompts - .insert(p.name.clone(), (server_name.to_string(), p)); - } - } - Err(e) => tracing::debug!("No prompts or failed to list on '{}': {}", server_name, e), - } - - // Resources - match client.peer().list_all_resources().await { - Ok(rs) => { - tracing::info!("Discovered {} resources from '{}'", rs.len(), server_name); - for r in rs { - if self.resources.contains_key(&r.uri) { - tracing::warn!( - "Resource '{}' from server '{}' is overwriting an existing resource.", - &r.uri, - server_name - ); - } - self.resources - .insert(r.uri.clone(), (server_name.to_string(), r)); - } - } - Err(e) => tracing::debug!("No resources or failed to list on '{}': {}", server_name, e), - } - } - - /// Connect to a single MCP server with retry logic for remote transports - async fn connect_server(config: &McpServerConfig) -> McpResult> { - let needs_retry = matches!( - &config.transport, - McpTransport::Sse { .. } | McpTransport::Streamable { .. } - ); - if needs_retry { - Self::connect_server_with_retry(config).await - } else { - Self::connect_server_impl(config).await - } - } - - /// Connect with exponential backoff retry for remote servers - async fn connect_server_with_retry( - config: &McpServerConfig, - ) -> McpResult> { - let backoff = ExponentialBackoffBuilder::new() - .with_initial_interval(Duration::from_secs(1)) - .with_max_interval(Duration::from_secs(30)) - .with_max_elapsed_time(Some(Duration::from_secs(30))) - .build(); - - backoff::future::retry(backoff, || async { - match Self::connect_server_impl(config).await { - Ok(client) => Ok(client), - Err(e) => { - if Self::is_permanent_error(&e) { - tracing::error!( - "Permanent error connecting to '{}': {} - not retrying", - config.name, - e - ); - Err(backoff::Error::permanent(e)) - } else { - tracing::warn!("Failed to connect to '{}', retrying: {}", config.name, e); - Err(backoff::Error::transient(e)) - } - } - } - }) - .await - } - - /// Determine if an error is permanent (should not retry) or transient (should retry) - fn is_permanent_error(error: &McpError) -> bool { - match error { - McpError::Config(_) => true, - McpError::Auth(_) => true, - McpError::ServerNotFound(_) => true, - McpError::Transport(_) => true, - McpError::ConnectionFailed(msg) => { - msg.contains("initialize") - || msg.contains("connection closed") - || msg.contains("connection refused") - || msg.contains("invalid URL") - || msg.contains("not found") - } - // Tool-related errors shouldn't occur during connection - _ => false, - } - } - - /// Internal implementation of server connection - async fn connect_server_impl( - config: &McpServerConfig, - ) -> McpResult> { - tracing::info!( - "Connecting to MCP server '{}' via {:?}", - config.name, - config.transport - ); - - match &config.transport { - McpTransport::Stdio { - command, - args, - envs, - } => { - let transport = TokioChildProcess::new( - tokio::process::Command::new(command).configure(|cmd| { - cmd.args(args) - .envs(envs.iter()) - .stderr(std::process::Stdio::inherit()); - }), - ) - .map_err(|e| McpError::Transport(format!("create stdio transport: {}", e)))?; - - let client = ().serve(transport).await.map_err(|e| { - McpError::ConnectionFailed(format!("initialize stdio client: {}", e)) - })?; - - tracing::info!("Connected to stdio server '{}'", config.name); - Ok(client) - } - - McpTransport::Sse { url, token } => { - let transport = if let Some(tok) = token { - let client = reqwest::Client::builder() - .default_headers({ - let mut headers = reqwest::header::HeaderMap::new(); - headers.insert( - reqwest::header::AUTHORIZATION, - format!("Bearer {}", tok).parse().map_err(|e| { - McpError::Transport(format!("auth token: {}", e)) - })?, - ); - headers - }) - .build() - .map_err(|e| McpError::Transport(format!("build HTTP client: {}", e)))?; - - let cfg = SseClientConfig { - sse_endpoint: url.clone().into(), - ..Default::default() - }; - - SseClientTransport::start_with_client(client, cfg) - .await - .map_err(|e| McpError::Transport(format!("create SSE transport: {}", e)))? - } else { - SseClientTransport::start(url.as_str()) - .await - .map_err(|e| McpError::Transport(format!("create SSE transport: {}", e)))? - }; - - let client = ().serve(transport).await.map_err(|e| { - McpError::ConnectionFailed(format!("initialize SSE client: {}", e)) - })?; - - tracing::info!("Connected to SSE server '{}' at {}", config.name, url); - Ok(client) - } - - McpTransport::Streamable { url, token } => { - let transport = if let Some(tok) = token { - let mut cfg = StreamableHttpClientTransportConfig::with_uri(url.as_str()); - cfg.auth_header = Some(format!("Bearer {}", tok)); - StreamableHttpClientTransport::from_config(cfg) - } else { - StreamableHttpClientTransport::from_uri(url.as_str()) - }; - - let client = ().serve(transport).await.map_err(|e| { - McpError::ConnectionFailed(format!("initialize streamable client: {}", e)) - })?; - - tracing::info!( - "Connected to streamable HTTP server '{}' at {}", - config.name, - url - ); - Ok(client) - } - } - } - - fn client_for(&self, server_name: &str) -> McpResult<&RunningService> { - self.clients - .get(server_name) - .ok_or_else(|| McpError::ServerNotFound(server_name.to_string())) - } - - fn tool_entry(&self, name: &str) -> McpResult<(String, McpTool)> { - self.tools - .get(name) - .map(|e| e.value().clone()) - .ok_or_else(|| McpError::ToolNotFound(name.to_string())) - } - - fn prompt_entry(&self, name: &str) -> McpResult<(String, Prompt)> { - self.prompts - .get(name) - .map(|e| e.value().clone()) - .ok_or_else(|| McpError::PromptNotFound(name.to_string())) - } - - fn resource_entry(&self, uri: &str) -> McpResult<(String, Resource)> { - self.resources - .get(uri) - .map(|e| e.value().clone()) - .ok_or_else(|| McpError::ResourceNotFound(uri.to_string())) - } - - /// Call a tool by name - pub async fn call_tool( - &self, - tool_name: &str, - arguments: Option>, - ) -> McpResult { - let (server_name, _tool) = self.tool_entry(tool_name)?; - let client = self.client_for(&server_name)?; - - tracing::debug!("Calling tool '{}' on '{}'", tool_name, server_name); - - client - .peer() - .call_tool(CallToolRequestParam { - name: Cow::Owned(tool_name.to_string()), - arguments, - }) - .await - .map_err(|e| McpError::ToolExecution(format!("Tool call failed: {}", e))) - } - - /// Get all available tools - pub fn list_tools(&self) -> Vec { - self.tools - .iter() - .map(|entry| { - let tool_name = entry.key().clone(); - let (server_name, tool) = entry.value(); - ToolInfo { - name: tool_name, - description: tool.description.as_deref().unwrap_or_default().to_string(), - server: server_name.clone(), - parameters: Some(serde_json::Value::Object((*tool.input_schema).clone())), - } - }) - .collect() - } - - /// Get a specific tool by name - pub fn get_tool(&self, name: &str) -> Option { - self.tools.get(name).map(|entry| { - let (server_name, tool) = entry.value(); - ToolInfo { - name: name.to_string(), - description: tool.description.as_deref().unwrap_or_default().to_string(), - server: server_name.clone(), - parameters: Some(serde_json::Value::Object((*tool.input_schema).clone())), - } - }) - } - - /// Check if a tool exists - pub fn has_tool(&self, name: &str) -> bool { - self.tools.contains_key(name) - } - - /// Get list of connected servers - pub fn list_servers(&self) -> Vec { - self.clients.keys().cloned().collect() - } - - /// Get a prompt by name with arguments - pub async fn get_prompt( - &self, - prompt_name: &str, - arguments: Option>, - ) -> McpResult { - let (server_name, _prompt) = self.prompt_entry(prompt_name)?; - let client = self.client_for(&server_name)?; - - tracing::debug!("Getting prompt '{}' from '{}'", prompt_name, server_name); - - client - .peer() - .get_prompt(GetPromptRequestParam { - name: prompt_name.to_string(), - arguments, - }) - .await - .map_err(|e| McpError::ToolExecution(format!("Failed to get prompt: {}", e))) - } - - /// List all available prompts - pub fn list_prompts(&self) -> Vec { - self.prompts - .iter() - .map(|entry| { - let name = entry.key().clone(); - let (server_name, prompt) = entry.value(); - PromptInfo { - name, - description: prompt.description.clone(), - server: server_name.clone(), - arguments: prompt - .arguments - .clone() - .map(|args| args.into_iter().map(|arg| serde_json::json!(arg)).collect()), - } - }) - .collect() - } - - /// Get a specific prompt info by name - pub fn get_prompt_info(&self, name: &str) -> Option { - self.prompts.get(name).map(|entry| { - let (server_name, prompt) = entry.value(); - PromptInfo { - name: name.to_string(), - description: prompt.description.clone(), - server: server_name.clone(), - arguments: prompt - .arguments - .clone() - .map(|args| args.into_iter().map(|arg| serde_json::json!(arg)).collect()), - } - }) - } - - /// Read a resource by URI - pub async fn read_resource(&self, uri: &str) -> McpResult { - let (server_name, _resource) = self.resource_entry(uri)?; - let client = self.client_for(&server_name)?; - - tracing::debug!("Reading resource '{}' from '{}'", uri, server_name); - - client - .peer() - .read_resource(ReadResourceRequestParam { - uri: uri.to_string(), - }) - .await - .map_err(|e| McpError::ToolExecution(format!("Failed to read resource: {}", e))) - } - - /// List all available resources - pub fn list_resources(&self) -> Vec { - self.resources - .iter() - .map(|entry| { - let uri = entry.key().clone(); - let (server_name, resource) = entry.value(); - ResourceInfo { - uri, - name: resource.name.clone(), - description: resource.description.clone(), - mime_type: resource.mime_type.clone(), - server: server_name.clone(), - } - }) - .collect() - } - - /// Get a specific resource info by URI - pub fn get_resource_info(&self, uri: &str) -> Option { - self.resources.get(uri).map(|entry| { - let (server_name, resource) = entry.value(); - ResourceInfo { - uri: uri.to_string(), - name: resource.name.clone(), - description: resource.description.clone(), - mime_type: resource.mime_type.clone(), - server: server_name.clone(), - } - }) - } - - /// Subscribe to resource changes - pub async fn subscribe_resource(&self, uri: &str) -> McpResult<()> { - let (server_name, _resource) = self.resource_entry(uri)?; - let client = self.client_for(&server_name)?; - - tracing::debug!("Subscribing to '{}' on '{}'", uri, server_name); - - client - .peer() - .subscribe(rmcp::model::SubscribeRequestParam { - uri: uri.to_string(), - }) - .await - .map_err(|e| McpError::ToolExecution(format!("Failed to subscribe: {}", e))) - } - - /// Unsubscribe from resource changes - pub async fn unsubscribe_resource(&self, uri: &str) -> McpResult<()> { - let (server_name, _resource) = self.resource_entry(uri)?; - let client = self.client_for(&server_name)?; - - tracing::debug!("Unsubscribing from '{}' on '{}'", uri, server_name); - - client - .peer() - .unsubscribe(rmcp::model::UnsubscribeRequestParam { - uri: uri.to_string(), - }) - .await - .map_err(|e| McpError::ToolExecution(format!("Failed to unsubscribe: {}", e))) - } - - /// Disconnect from all servers (for cleanup) - pub async fn shutdown(&mut self) { - for (name, client) in self.clients.drain() { - if let Err(e) = client.cancel().await { - tracing::warn!("Error disconnecting from '{}': {}", name, e); - } - } - self.tools.clear(); - self.prompts.clear(); - self.resources.clear(); - } -} diff --git a/sgl-router/src/mcp/config.rs b/sgl-router/src/mcp/config.rs index e94208e5b..92590521d 100644 --- a/sgl-router/src/mcp/config.rs +++ b/sgl-router/src/mcp/config.rs @@ -2,9 +2,63 @@ use std::collections::HashMap; use serde::{Deserialize, Serialize}; +// ============================================================================ +// MCP Data Structures +// ============================================================================ + +/// Information about an available tool +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ToolInfo { + pub name: String, + pub description: String, + pub server: String, + pub parameters: Option, +} + +/// Information about an available prompt +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct PromptInfo { + pub name: String, + pub description: Option, + pub server: String, + pub arguments: Option>, +} + +/// Information about an available resource +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ResourceInfo { + pub uri: String, + pub name: String, + pub description: Option, + pub mime_type: Option, + pub server: String, +} + +// ============================================================================ +// Configuration Structures +// ============================================================================ + #[derive(Debug, Clone, Deserialize, Serialize)] pub struct McpConfig { + /// Static MCP servers (loaded at startup) pub servers: Vec, + + /// Connection pool settings + #[serde(default)] + pub pool: McpPoolConfig, + + /// Global MCP proxy configuration (default for all servers) + /// Can be overridden per-server + #[serde(default)] + pub proxy: Option, + + /// Pre-warm these connections at startup + #[serde(default)] + pub warmup: Vec, + + /// Tool inventory refresh settings + #[serde(default)] + pub inventory: InventoryConfig, } #[derive(Debug, Clone, Deserialize, Serialize)] @@ -12,6 +66,17 @@ pub struct McpServerConfig { pub name: String, #[serde(flatten)] pub transport: McpTransport, + + /// Per-server proxy override (overrides global proxy) + /// Set to `null` in YAML to force direct connection (no proxy) + #[serde(default)] + pub proxy: Option, + + /// Whether this server is required for router startup + /// - true: Router startup fails if this server cannot be reached + /// - false: Log warning but continue (default) + #[serde(default)] + pub required: bool, } #[derive(Debug, Clone, Deserialize, Serialize)] @@ -36,6 +101,144 @@ pub enum McpTransport { }, } +/// MCP-specific proxy configuration (does NOT affect LLM API traffic) +#[derive(Debug, Clone, Deserialize, Serialize)] +pub struct McpProxyConfig { + /// HTTP proxy URL (e.g., "http://proxy.internal:8080") + pub http: Option, + + /// HTTPS proxy URL + pub https: Option, + + /// Comma-separated hosts to exclude from proxying + /// Example: "localhost,127.0.0.1,*.internal,10.*" + pub no_proxy: Option, + + /// Custom proxy authentication (if needed) + #[serde(skip_serializing_if = "Option::is_none")] + pub username: Option, + + #[serde(skip_serializing_if = "Option::is_none")] + pub password: Option, +} + +/// Connection pool configuration +#[derive(Debug, Clone, Deserialize, Serialize)] +pub struct McpPoolConfig { + /// Maximum cached connections per server URL + #[serde(default = "default_max_connections")] + pub max_connections: usize, + + /// Idle timeout before closing connection (seconds) + #[serde(default = "default_idle_timeout")] + pub idle_timeout: u64, +} + +/// Tool inventory refresh configuration +#[derive(Debug, Clone, Deserialize, Serialize)] +pub struct InventoryConfig { + /// Enable automatic tool inventory refresh + #[serde(default = "default_true")] + pub enable_refresh: bool, + + /// Tool cache TTL (seconds) - how long tools are considered fresh + #[serde(default = "default_tool_ttl")] + pub tool_ttl: u64, + + /// Background refresh interval (seconds) - proactive refresh + #[serde(default = "default_refresh_interval")] + pub refresh_interval: u64, + + /// Refresh on tool call failure (try refreshing if tool not found) + #[serde(default = "default_true")] + pub refresh_on_error: bool, +} + +/// Pre-warm server connections at startup +#[derive(Debug, Clone, Deserialize, Serialize)] +pub struct WarmupServer { + /// Server URL + pub url: String, + + /// Server label/name + pub label: String, + + /// Optional authentication token + #[serde(skip_serializing_if = "Option::is_none")] + pub token: Option, +} + +// Default value functions +fn default_max_connections() -> usize { + 100 +} + +fn default_idle_timeout() -> u64 { + 300 // 5 minutes +} + +fn default_true() -> bool { + true +} + +fn default_tool_ttl() -> u64 { + 300 // 5 minutes +} + +fn default_refresh_interval() -> u64 { + 60 // 1 minute +} + +// Default implementations +impl Default for McpPoolConfig { + fn default() -> Self { + Self { + max_connections: default_max_connections(), + idle_timeout: default_idle_timeout(), + } + } +} + +impl Default for InventoryConfig { + fn default() -> Self { + Self { + enable_refresh: true, + tool_ttl: default_tool_ttl(), + refresh_interval: default_refresh_interval(), + refresh_on_error: true, + } + } +} + +impl McpProxyConfig { + /// Load proxy config from standard environment variables + pub fn from_env() -> Option { + let http = std::env::var("MCP_HTTP_PROXY") + .ok() + .or_else(|| std::env::var("HTTP_PROXY").ok()); + + let https = std::env::var("MCP_HTTPS_PROXY") + .ok() + .or_else(|| std::env::var("HTTPS_PROXY").ok()); + + let no_proxy = std::env::var("MCP_NO_PROXY") + .ok() + .or_else(|| std::env::var("NO_PROXY").ok()); + + if http.is_some() || https.is_some() { + Some(Self { + http, + https, + no_proxy, + username: None, + password: None, + }) + } else { + None + } + } +} + impl McpConfig { /// Load configuration from a YAML file pub async fn from_file(path: &str) -> Result> { @@ -50,4 +253,280 @@ impl McpConfig { // For now, return None to indicate env config not implemented None } + + /// Merge with environment-based proxy config + pub fn with_env_proxy(mut self) -> Self { + if self.proxy.is_none() { + self.proxy = McpProxyConfig::from_env(); + } + self + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_default_pool_config() { + let config = McpPoolConfig::default(); + assert_eq!(config.max_connections, 100); + assert_eq!(config.idle_timeout, 300); + } + + #[test] + fn test_default_inventory_config() { + let config = InventoryConfig::default(); + assert!(config.enable_refresh); + assert_eq!(config.tool_ttl, 300); + assert_eq!(config.refresh_interval, 60); + assert!(config.refresh_on_error); + } + + #[test] + fn test_proxy_from_env_empty() { + // Ensure no proxy env vars are set for this test + std::env::remove_var("MCP_HTTP_PROXY"); + std::env::remove_var("MCP_HTTPS_PROXY"); + std::env::remove_var("HTTP_PROXY"); + std::env::remove_var("HTTPS_PROXY"); + + let proxy = McpProxyConfig::from_env(); + assert!(proxy.is_none(), "Should return None when no env vars set"); + } + + #[test] + fn test_proxy_from_env_with_vars() { + std::env::set_var("MCP_HTTP_PROXY", "http://test-proxy:8080"); + std::env::set_var("MCP_NO_PROXY", "localhost,127.0.0.1"); + + let proxy = McpProxyConfig::from_env(); + assert!(proxy.is_some(), "Should return Some when env vars set"); + + let proxy = proxy.unwrap(); + assert_eq!(proxy.http.as_ref().unwrap(), "http://test-proxy:8080"); + assert_eq!(proxy.no_proxy.as_ref().unwrap(), "localhost,127.0.0.1"); + + // Cleanup + std::env::remove_var("MCP_HTTP_PROXY"); + std::env::remove_var("MCP_NO_PROXY"); + } + + #[tokio::test] + async fn test_yaml_minimal_config() { + let yaml = r#" +servers: + - name: "test-server" + protocol: sse + url: "http://localhost:3000/sse" +"#; + + let config: McpConfig = serde_yaml::from_str(yaml).expect("Failed to parse YAML"); + assert_eq!(config.servers.len(), 1); + assert_eq!(config.servers[0].name, "test-server"); + assert!(!config.servers[0].required); // Should default to false + assert!(config.servers[0].proxy.is_none()); // Should default to None + assert_eq!(config.pool.max_connections, 100); // Should use default + assert_eq!(config.inventory.tool_ttl, 300); // Should use default + } + + #[tokio::test] + async fn test_yaml_full_config() { + let yaml = r#" +# Global proxy configuration +proxy: + http: "http://global-proxy:8080" + https: "http://global-proxy:8080" + no_proxy: "localhost,127.0.0.1,*.internal" + +# Connection pool settings +pool: + max_connections: 50 + idle_timeout: 600 + +# Tool inventory settings +inventory: + enable_refresh: true + tool_ttl: 600 + refresh_interval: 120 + refresh_on_error: true + +# Static servers +servers: + - name: "required-server" + protocol: sse + url: "https://api.example.com/sse" + token: "secret-token" + required: true + + - name: "optional-server" + protocol: stdio + command: "mcp-server" + args: ["--port", "3000"] + required: false + proxy: + http: "http://server-specific-proxy:9090" + +# Pre-warm connections +warmup: + - url: "http://localhost:3000/sse" + label: "local-dev" +"#; + + let config: McpConfig = serde_yaml::from_str(yaml).expect("Failed to parse YAML"); + + // Check global proxy + assert!(config.proxy.is_some()); + let global_proxy = config.proxy.as_ref().unwrap(); + assert_eq!( + global_proxy.http.as_ref().unwrap(), + "http://global-proxy:8080" + ); + + // Check pool config + assert_eq!(config.pool.max_connections, 50); + assert_eq!(config.pool.idle_timeout, 600); + + // Check inventory config + assert_eq!(config.inventory.tool_ttl, 600); + assert_eq!(config.inventory.refresh_interval, 120); + + // Check servers + assert_eq!(config.servers.len(), 2); + + // Required server + assert_eq!(config.servers[0].name, "required-server"); + assert!(config.servers[0].required); + assert!(config.servers[0].proxy.is_none()); // Inherits global proxy + + // Optional server with custom proxy + assert_eq!(config.servers[1].name, "optional-server"); + assert!(!config.servers[1].required); + assert!(config.servers[1].proxy.is_some()); + assert_eq!( + config.servers[1] + .proxy + .as_ref() + .unwrap() + .http + .as_ref() + .unwrap(), + "http://server-specific-proxy:9090" + ); + + // Check warmup + assert_eq!(config.warmup.len(), 1); + assert_eq!(config.warmup[0].label, "local-dev"); + } + + #[tokio::test] + async fn test_yaml_backward_compatibility() { + // Old config format should still work + let yaml = r#" +servers: + - name: "legacy-server" + protocol: sse + url: "http://localhost:3000/sse" +"#; + + let config: McpConfig = serde_yaml::from_str(yaml).expect("Failed to parse old format"); + assert_eq!(config.servers.len(), 1); + assert_eq!(config.servers[0].name, "legacy-server"); + assert!(!config.servers[0].required); // New field defaults to false + assert!(config.servers[0].proxy.is_none()); // New field defaults to None + assert!(config.proxy.is_none()); // New field defaults to None + assert!(config.warmup.is_empty()); // New field defaults to empty + } + + #[tokio::test] + async fn test_yaml_null_proxy_override() { + // Test that explicit null in YAML sets proxy to None + let yaml = r#" +proxy: + http: "http://global-proxy:8080" + +servers: + - name: "direct-connection" + protocol: sse + url: "http://localhost:3000/sse" + proxy: null +"#; + + let config: McpConfig = serde_yaml::from_str(yaml).expect("Failed to parse YAML"); + assert!(config.proxy.is_some()); // Global proxy set + assert_eq!(config.servers.len(), 1); + assert!(config.servers[0].proxy.is_none()); // Explicitly set to null + } + + #[test] + fn test_transport_stdio() { + let yaml = r#" +name: "test" +protocol: stdio +command: "mcp-server" +args: ["--port", "3000"] +envs: + VAR1: "value1" + VAR2: "value2" +"#; + + let config: McpServerConfig = serde_yaml::from_str(yaml).expect("Failed to parse stdio"); + assert_eq!(config.name, "test"); + + match config.transport { + McpTransport::Stdio { + command, + args, + envs, + } => { + assert_eq!(command, "mcp-server"); + assert_eq!(args.len(), 2); + assert_eq!(args[0], "--port"); + assert_eq!(envs.get("VAR1").unwrap(), "value1"); + } + _ => panic!("Expected Stdio transport"), + } + } + + #[test] + fn test_transport_sse() { + let yaml = r#" +name: "test" +protocol: sse +url: "http://localhost:3000/sse" +token: "secret" +"#; + + let config: McpServerConfig = serde_yaml::from_str(yaml).expect("Failed to parse sse"); + assert_eq!(config.name, "test"); + + match config.transport { + McpTransport::Sse { url, token } => { + assert_eq!(url, "http://localhost:3000/sse"); + assert_eq!(token.unwrap(), "secret"); + } + _ => panic!("Expected Sse transport"), + } + } + + #[test] + fn test_transport_streamable() { + let yaml = r#" +name: "test" +protocol: streamable +url: "http://localhost:3000" +"#; + + let config: McpServerConfig = + serde_yaml::from_str(yaml).expect("Failed to parse streamable"); + assert_eq!(config.name, "test"); + + match config.transport { + McpTransport::Streamable { url, token } => { + assert_eq!(url, "http://localhost:3000"); + assert!(token.is_none()); + } + _ => panic!("Expected Streamable transport"), + } + } } diff --git a/sgl-router/src/mcp/connection_pool.rs b/sgl-router/src/mcp/connection_pool.rs new file mode 100644 index 000000000..5a4f774e2 --- /dev/null +++ b/sgl-router/src/mcp/connection_pool.rs @@ -0,0 +1,448 @@ +// MCP Connection Pool +// +// This module provides connection pooling for dynamic MCP servers (per-request). +// Connections are cached and reused to avoid 70-650ms connection overhead on every request. +// +// Performance target: +// - First request: 70-650ms (connection establishment) +// - Subsequent requests: <1ms (cache hit) +// - 90%+ reduction in per-request overhead + +use std::{ + sync::Arc, + time::{Duration, Instant}, +}; + +use dashmap::DashMap; +use rmcp::{service::RunningService, RoleClient}; + +use crate::mcp::{ + config::{McpProxyConfig, McpServerConfig}, + error::McpResult, +}; + +/// Type alias for MCP client +type McpClient = RunningService; + +/// Cached MCP connection with metadata +#[derive(Clone)] +pub struct CachedConnection { + /// The MCP client instance + pub client: Arc, + /// Last time this connection was accessed + pub last_used: Instant, + /// Server configuration used to create this connection + pub config: McpServerConfig, +} + +impl CachedConnection { + /// Create a new cached connection + pub fn new(client: Arc, config: McpServerConfig) -> Self { + Self { + client, + last_used: Instant::now(), + config, + } + } + + /// Update last_used timestamp + pub fn touch(&mut self) { + self.last_used = Instant::now(); + } + + /// Check if connection has been idle for longer than TTL + pub fn is_idle(&self, idle_ttl: Duration) -> bool { + self.last_used.elapsed() > idle_ttl + } +} + +/// Connection pool for dynamic MCP servers +/// +/// Provides thread-safe connection pooling with automatic cleanup of idle connections. +/// Connections are keyed by server URL and reused across requests. +pub struct McpConnectionPool { + /// Map of server_url -> cached connection + connections: DashMap, + + /// Idle connection TTL (connections unused for this duration are cleaned up) + idle_ttl: Duration, + + /// Maximum number of cached connections (prevents unbounded growth) + max_connections: usize, + + /// Global proxy configuration (applied to all dynamic servers) + /// Can be overridden per-server via McpServerConfig.proxy + global_proxy: Option, +} + +impl McpConnectionPool { + /// Create a new connection pool with default settings + /// + /// Default settings: + /// - idle_ttl: 300 seconds (5 minutes) + /// - max_connections: 100 + /// - global_proxy: Loaded from environment variables (MCP_HTTP_PROXY, etc.) + pub fn new() -> Self { + Self { + connections: DashMap::new(), + idle_ttl: Duration::from_secs(300), + max_connections: 100, + global_proxy: McpProxyConfig::from_env(), + } + } + + /// Create a new connection pool with custom settings + pub fn with_config(idle_ttl: Duration, max_connections: usize) -> Self { + Self { + connections: DashMap::new(), + idle_ttl, + max_connections, + global_proxy: McpProxyConfig::from_env(), + } + } + + /// Create a new connection pool with full custom configuration + pub fn with_full_config( + idle_ttl: Duration, + max_connections: usize, + global_proxy: Option, + ) -> Self { + Self { + connections: DashMap::new(), + idle_ttl, + max_connections, + global_proxy, + } + } + + /// Get an existing connection or create a new one + /// + /// This method: + /// 1. Checks if a connection exists for the given URL + /// 2. If exists and fresh, updates last_used and returns it (fast path <1ms) + /// 3. If not exists or stale, creates new connection (slow path 70-650ms) + /// + /// # Arguments + /// * `server_url` - The MCP server URL (used as cache key) + /// * `server_config` - Server configuration (used to create new connection if needed) + /// * `connect_fn` - Async function to create a new client connection + /// + /// # Returns + /// Arc to the MCP client, either from cache or newly created + pub async fn get_or_create( + &self, + server_url: &str, + server_config: McpServerConfig, + connect_fn: F, + ) -> McpResult> + where + F: FnOnce(McpServerConfig, Option) -> Fut, + Fut: std::future::Future>, + { + // Fast path: Check if connection exists and is still fresh + if let Some(mut entry) = self.connections.get_mut(server_url) { + let cached = entry.value_mut(); + + // Check if connection is still within TTL + if !cached.is_idle(self.idle_ttl) { + // Update last_used and return cached connection + cached.touch(); + return Ok(Arc::clone(&cached.client)); + } + + // Connection is stale, drop it and create new one + drop(entry); + self.connections.remove(server_url); + } + + // Slow path: Create new connection + // Enforce max_connections limit + if self.connections.len() >= self.max_connections { + self.cleanup_idle_connections(); + + // If still at limit after cleanup, remove oldest connection + if self.connections.len() >= self.max_connections { + if let Some(oldest_key) = self.find_oldest_connection() { + self.connections.remove(&oldest_key); + } + } + } + + // Create new MCP client using the provided connect function + let client = connect_fn(server_config.clone(), self.global_proxy.clone()).await?; + let client_arc = Arc::new(client); + + // Cache the new connection + let cached = CachedConnection::new(Arc::clone(&client_arc), server_config); + self.connections.insert(server_url.to_string(), cached); + + Ok(client_arc) + } + + /// Remove all idle connections that have exceeded the TTL + /// + /// This method is called: + /// - Automatically when max_connections limit is reached + /// - Can be called manually by background cleanup task + pub fn cleanup_idle_connections(&self) { + let now = Instant::now(); + self.connections + .retain(|_, cached| now.duration_since(cached.last_used) < self.idle_ttl); + } + + /// Find the oldest connection (by last_used timestamp) + /// + /// Used for eviction when max_connections is reached and cleanup didn't free space + fn find_oldest_connection(&self) -> Option { + self.connections + .iter() + .min_by_key(|entry| entry.value().last_used) + .map(|entry| entry.key().clone()) + } + + /// Get current number of cached connections + pub fn len(&self) -> usize { + self.connections.len() + } + + /// Check if pool is empty + pub fn is_empty(&self) -> bool { + self.connections.is_empty() + } + + /// Clear all connections (useful for tests) + pub fn clear(&self) { + self.connections.clear(); + } + + /// Get connection statistics + pub fn stats(&self) -> PoolStats { + let total = self.connections.len(); + let idle_count = self + .connections + .iter() + .filter(|entry| entry.value().is_idle(self.idle_ttl)) + .count(); + + PoolStats { + total_connections: total, + active_connections: total - idle_count, + idle_connections: idle_count, + } + } +} + +impl Default for McpConnectionPool { + fn default() -> Self { + Self::new() + } +} + +/// Connection pool statistics +#[derive(Debug, Clone)] +pub struct PoolStats { + pub total_connections: usize, + pub active_connections: usize, + pub idle_connections: usize, +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::mcp::McpTransport; + + // Helper to create test server config + fn create_test_config(url: &str) -> McpServerConfig { + McpServerConfig { + name: "test_server".to_string(), + transport: McpTransport::Streamable { + url: url.to_string(), + token: None, + }, + proxy: None, + required: false, + } + } + + #[tokio::test] + async fn test_pool_creation() { + let pool = McpConnectionPool::new(); + assert_eq!(pool.len(), 0); + assert!(pool.is_empty()); + } + + #[test] + #[allow(invalid_value)] + fn test_cached_connection_touch() { + let config = create_test_config("http://localhost:3000"); + let client: Arc = Arc::new(unsafe { + // SAFETY: This is only for testing the CachedConnection struct + std::mem::MaybeUninit::zeroed().assume_init() + }); + let mut cached = CachedConnection::new(client.clone(), config); + + let first_time = cached.last_used; + std::thread::sleep(Duration::from_millis(10)); + cached.touch(); + assert!(cached.last_used > first_time); + + // Prevent drop of invalid Arc (would segfault) + std::mem::forget(client); + } + + #[test] + #[allow(invalid_value)] + fn test_cached_connection_is_idle() { + let config = create_test_config("http://localhost:3000"); + let client: Arc = Arc::new(unsafe { + // SAFETY: This is only for testing the CachedConnection struct + std::mem::MaybeUninit::zeroed().assume_init() + }); + let cached = CachedConnection::new(client.clone(), config); + + // Fresh connection should not be idle + assert!(!cached.is_idle(Duration::from_secs(1))); + + // Wait and check + std::thread::sleep(Duration::from_millis(100)); + assert!(cached.is_idle(Duration::from_millis(50))); + + // Prevent drop of invalid Arc (would segfault) + std::mem::forget(client); + } + + #[test] + fn test_pool_stats() { + let pool = McpConnectionPool::with_config(Duration::from_millis(100), 10); + + let stats = pool.stats(); + assert_eq!(stats.total_connections, 0); + assert_eq!(stats.active_connections, 0); + assert_eq!(stats.idle_connections, 0); + } + + #[test] + #[allow(invalid_value)] + fn test_cleanup_idle_connections() { + let pool = McpConnectionPool::with_config(Duration::from_millis(50), 10); + + // Initially empty + assert_eq!(pool.len(), 0); + + // Add a connection manually for testing + let config = create_test_config("http://localhost:3000"); + let client: Arc = + Arc::new(unsafe { std::mem::MaybeUninit::zeroed().assume_init() }); + let cached = CachedConnection::new(client.clone(), config); + pool.connections + .insert("http://localhost:3000".to_string(), cached); + + assert_eq!(pool.len(), 1); + + // Wait for TTL to expire + std::thread::sleep(Duration::from_millis(100)); + + // Cleanup should remove idle connection + pool.cleanup_idle_connections(); + assert_eq!(pool.len(), 0); + + // Prevent drop of invalid Arc (would segfault) + std::mem::forget(client); + } + + #[test] + #[allow(invalid_value)] + fn test_find_oldest_connection() { + let pool = McpConnectionPool::new(); + + // Collect clients to forget at end + let mut clients = Vec::new(); + + // Add connections with different timestamps + for i in 0..3 { + let url = format!("http://localhost:{}", 3000 + i); + let config = create_test_config(&url); + let client: Arc = + Arc::new(unsafe { std::mem::MaybeUninit::zeroed().assume_init() }); + let cached = CachedConnection::new(client.clone(), config); + pool.connections.insert(url, cached); + clients.push(client); + std::thread::sleep(Duration::from_millis(10)); + } + + // Oldest should be the first one + let oldest = pool.find_oldest_connection(); + assert!(oldest.is_some()); + assert_eq!(oldest.unwrap(), "http://localhost:3000"); + + // Prevent drop of invalid Arcs (would segfault) + for client in clients { + std::mem::forget(client); + } + } + + #[test] + #[allow(invalid_value)] + fn test_pool_clear() { + let pool = McpConnectionPool::new(); + + // Add a connection + let config = create_test_config("http://localhost:3000"); + let client: Arc = + Arc::new(unsafe { std::mem::MaybeUninit::zeroed().assume_init() }); + let cached = CachedConnection::new(client.clone(), config); + pool.connections + .insert("http://localhost:3000".to_string(), cached); + + assert_eq!(pool.len(), 1); + + pool.clear(); + assert_eq!(pool.len(), 0); + assert!(pool.is_empty()); + + // Prevent drop of invalid Arc (would segfault) + std::mem::forget(client); + } + + #[test] + fn test_pool_with_global_proxy() { + use crate::mcp::McpProxyConfig; + + // Create proxy config + let proxy = McpProxyConfig { + http: Some("http://proxy.example.com:8080".to_string()), + https: None, + no_proxy: Some("localhost,127.0.0.1".to_string()), + username: None, + password: None, + }; + + // Create pool with proxy + let pool = + McpConnectionPool::with_full_config(Duration::from_secs(300), 100, Some(proxy.clone())); + + // Verify proxy is stored + assert!(pool.global_proxy.is_some()); + let stored_proxy = pool.global_proxy.as_ref().unwrap(); + assert_eq!( + stored_proxy.http.as_ref().unwrap(), + "http://proxy.example.com:8080" + ); + assert_eq!( + stored_proxy.no_proxy.as_ref().unwrap(), + "localhost,127.0.0.1" + ); + } + + #[test] + fn test_pool_proxy_from_env() { + // Note: This test depends on environment variables + // In production, proxy is loaded from MCP_HTTP_PROXY or HTTP_PROXY env vars + let pool = McpConnectionPool::new(); + + // Pool should either have proxy from env or None + // We can't assert specific value since it depends on test environment + // Just verify it doesn't panic + assert!(pool.global_proxy.is_some() || pool.global_proxy.is_none()); + } +} diff --git a/sgl-router/src/mcp/inventory.rs b/sgl-router/src/mcp/inventory.rs new file mode 100644 index 000000000..cead2399d --- /dev/null +++ b/sgl-router/src/mcp/inventory.rs @@ -0,0 +1,620 @@ +// MCP Tool Inventory with TTL-based Caching +// +// This module provides TTL-based caching for MCP tools, prompts, and resources. +// Tools are cached with timestamps and automatically expire after the configured TTL. +// Background refresh tasks can proactively update the inventory. + +use std::time::{Duration, Instant}; + +use dashmap::DashMap; + +use crate::mcp::config::{PromptInfo, ResourceInfo, ToolInfo}; + +/// Cached tool with metadata +#[derive(Clone)] +pub struct CachedTool { + pub server_name: String, + pub tool: ToolInfo, + pub cached_at: Instant, +} + +/// Cached prompt with metadata +#[derive(Clone)] +pub struct CachedPrompt { + pub server_name: String, + pub prompt: PromptInfo, + pub cached_at: Instant, +} + +/// Cached resource with metadata +#[derive(Clone)] +pub struct CachedResource { + pub server_name: String, + pub resource: ResourceInfo, + pub cached_at: Instant, +} + +/// Tool inventory with TTL-based caching +/// +/// Provides thread-safe caching of MCP tools, prompts, and resources with automatic expiration. +/// Entries are timestamped and can be queried with TTL validation. +pub struct ToolInventory { + /// Map of tool_name -> cached tool + tools: DashMap, + + /// Map of prompt_name -> cached prompt + prompts: DashMap, + + /// Map of resource_uri -> cached resource + resources: DashMap, + + /// Tool cache TTL + tool_ttl: Duration, + + /// Last refresh time per server + server_refresh_times: DashMap, +} + +impl ToolInventory { + /// Create a new tool inventory with the specified TTL + pub fn new(tool_ttl: Duration) -> Self { + Self { + tools: DashMap::new(), + prompts: DashMap::new(), + resources: DashMap::new(), + tool_ttl, + server_refresh_times: DashMap::new(), + } + } + + // ============================================================================ + // Tool Methods + // ============================================================================ + + /// Get a tool if it exists and is fresh (within TTL) + /// + /// Returns None if the tool doesn't exist or has expired. + pub fn get_tool(&self, tool_name: &str) -> Option<(String, ToolInfo)> { + self.tools.get(tool_name).and_then(|entry| { + let cached = entry.value(); + + // Check if still fresh + if cached.cached_at.elapsed() < self.tool_ttl { + Some((cached.server_name.clone(), cached.tool.clone())) + } else { + // Expired - will be removed by cleanup + None + } + }) + } + + /// Check if tool exists (regardless of TTL) + pub fn has_tool(&self, tool_name: &str) -> bool { + self.tools.contains_key(tool_name) + } + + /// Insert or update a tool + pub fn insert_tool(&self, tool_name: String, server_name: String, tool: ToolInfo) { + self.tools.insert( + tool_name, + CachedTool { + server_name, + tool, + cached_at: Instant::now(), + }, + ); + } + + /// Get all tools (fresh only) + pub fn list_tools(&self) -> Vec<(String, String, ToolInfo)> { + let now = Instant::now(); + self.tools + .iter() + .filter_map(|entry| { + let (name, cached) = entry.pair(); + if now.duration_since(cached.cached_at) < self.tool_ttl { + Some(( + name.clone(), + cached.server_name.clone(), + cached.tool.clone(), + )) + } else { + None + } + }) + .collect() + } + + // ============================================================================ + // Prompt Methods + // ============================================================================ + + /// Get a prompt if it exists and is fresh (within TTL) + pub fn get_prompt(&self, prompt_name: &str) -> Option<(String, PromptInfo)> { + self.prompts.get(prompt_name).and_then(|entry| { + let cached = entry.value(); + + // Check if still fresh + if cached.cached_at.elapsed() < self.tool_ttl { + Some((cached.server_name.clone(), cached.prompt.clone())) + } else { + None + } + }) + } + + /// Check if prompt exists (regardless of TTL) + pub fn has_prompt(&self, prompt_name: &str) -> bool { + self.prompts.contains_key(prompt_name) + } + + /// Insert or update a prompt + pub fn insert_prompt(&self, prompt_name: String, server_name: String, prompt: PromptInfo) { + self.prompts.insert( + prompt_name, + CachedPrompt { + server_name, + prompt, + cached_at: Instant::now(), + }, + ); + } + + /// Get all prompts (fresh only) + pub fn list_prompts(&self) -> Vec<(String, String, PromptInfo)> { + let now = Instant::now(); + self.prompts + .iter() + .filter_map(|entry| { + let (name, cached) = entry.pair(); + if now.duration_since(cached.cached_at) < self.tool_ttl { + Some(( + name.clone(), + cached.server_name.clone(), + cached.prompt.clone(), + )) + } else { + None + } + }) + .collect() + } + + // ============================================================================ + // Resource Methods + // ============================================================================ + + /// Get a resource if it exists and is fresh (within TTL) + pub fn get_resource(&self, resource_uri: &str) -> Option<(String, ResourceInfo)> { + self.resources.get(resource_uri).and_then(|entry| { + let cached = entry.value(); + + // Check if still fresh + if cached.cached_at.elapsed() < self.tool_ttl { + Some((cached.server_name.clone(), cached.resource.clone())) + } else { + None + } + }) + } + + /// Check if resource exists (regardless of TTL) + pub fn has_resource(&self, resource_uri: &str) -> bool { + self.resources.contains_key(resource_uri) + } + + /// Insert or update a resource + pub fn insert_resource( + &self, + resource_uri: String, + server_name: String, + resource: ResourceInfo, + ) { + self.resources.insert( + resource_uri, + CachedResource { + server_name, + resource, + cached_at: Instant::now(), + }, + ); + } + + /// Get all resources (fresh only) + pub fn list_resources(&self) -> Vec<(String, String, ResourceInfo)> { + let now = Instant::now(); + self.resources + .iter() + .filter_map(|entry| { + let (uri, cached) = entry.pair(); + if now.duration_since(cached.cached_at) < self.tool_ttl { + Some(( + uri.clone(), + cached.server_name.clone(), + cached.resource.clone(), + )) + } else { + None + } + }) + .collect() + } + + // ============================================================================ + // Server Management Methods + // ============================================================================ + + /// Clear all cached items for a specific server (before refresh) + pub fn clear_server_tools(&self, server_name: &str) { + self.tools + .retain(|_, cached| cached.server_name != server_name); + self.prompts + .retain(|_, cached| cached.server_name != server_name); + self.resources + .retain(|_, cached| cached.server_name != server_name); + } + + /// Mark server as refreshed + pub fn mark_refreshed(&self, server_name: &str) { + self.server_refresh_times + .insert(server_name.to_string(), Instant::now()); + } + + /// Check if server needs refresh based on refresh interval + pub fn needs_refresh(&self, server_name: &str, refresh_interval: Duration) -> bool { + self.server_refresh_times + .get(server_name) + .map(|t| t.elapsed() > refresh_interval) + .unwrap_or(true) // Never refreshed = needs refresh + } + + /// Get last refresh time for a server + pub fn last_refresh(&self, server_name: &str) -> Option { + self.server_refresh_times + .get(server_name) + .map(|t| *t.value()) + } + + // ============================================================================ + // Cleanup Methods + // ============================================================================ + + /// Cleanup expired entries + /// + /// Removes all tools, prompts, and resources that have exceeded their TTL. + /// Should be called periodically by a background task. + pub fn cleanup_expired(&self) { + let now = Instant::now(); + + // Remove expired tools + self.tools + .retain(|_, cached| now.duration_since(cached.cached_at) < self.tool_ttl); + + // Remove expired prompts + self.prompts + .retain(|_, cached| now.duration_since(cached.cached_at) < self.tool_ttl); + + // Remove expired resources + self.resources + .retain(|_, cached| now.duration_since(cached.cached_at) < self.tool_ttl); + } + + /// Get count of cached items + pub fn counts(&self) -> (usize, usize, usize) { + (self.tools.len(), self.prompts.len(), self.resources.len()) + } + + /// Clear all cached items + pub fn clear_all(&self) { + self.tools.clear(); + self.prompts.clear(); + self.resources.clear(); + self.server_refresh_times.clear(); + } +} + +#[cfg(test)] +mod tests { + use super::*; + + // Helper to create a test tool + fn create_test_tool(name: &str) -> ToolInfo { + ToolInfo { + name: name.to_string(), + description: format!("Test tool: {}", name), + server: "test_server".to_string(), + parameters: Some(serde_json::json!({ + "type": "object", + "properties": {} + })), + } + } + + // Helper to create a test prompt + fn create_test_prompt(name: &str) -> PromptInfo { + PromptInfo { + name: name.to_string(), + description: Some(format!("Test prompt: {}", name)), + server: "test_server".to_string(), + arguments: None, + } + } + + // Helper to create a test resource + fn create_test_resource(uri: &str) -> ResourceInfo { + ResourceInfo { + uri: uri.to_string(), + name: uri.to_string(), + description: Some(format!("Test resource: {}", uri)), + mime_type: Some("text/plain".to_string()), + server: "test_server".to_string(), + } + } + + #[test] + fn test_tool_insert_and_get() { + let inventory = ToolInventory::new(Duration::from_secs(60)); + let tool = create_test_tool("test_tool"); + + inventory.insert_tool("test_tool".to_string(), "server1".to_string(), tool.clone()); + + let result = inventory.get_tool("test_tool"); + assert!(result.is_some()); + + let (server_name, retrieved_tool) = result.unwrap(); + assert_eq!(server_name, "server1"); + assert_eq!(retrieved_tool.name, "test_tool"); + } + + #[test] + fn test_tool_expiration() { + let inventory = ToolInventory::new(Duration::from_millis(100)); + let tool = create_test_tool("expiring_tool"); + + inventory.insert_tool( + "expiring_tool".to_string(), + "server1".to_string(), + tool.clone(), + ); + + // Should be available immediately + assert!(inventory.get_tool("expiring_tool").is_some()); + + // Wait for expiration + std::thread::sleep(Duration::from_millis(150)); + + // Should be expired now + assert!(inventory.get_tool("expiring_tool").is_none()); + } + + #[test] + fn test_has_tool() { + let inventory = ToolInventory::new(Duration::from_secs(60)); + let tool = create_test_tool("check_tool"); + + assert!(!inventory.has_tool("check_tool")); + + inventory.insert_tool("check_tool".to_string(), "server1".to_string(), tool); + + assert!(inventory.has_tool("check_tool")); + } + + #[test] + fn test_list_tools() { + let inventory = ToolInventory::new(Duration::from_secs(60)); + + inventory.insert_tool( + "tool1".to_string(), + "server1".to_string(), + create_test_tool("tool1"), + ); + inventory.insert_tool( + "tool2".to_string(), + "server1".to_string(), + create_test_tool("tool2"), + ); + inventory.insert_tool( + "tool3".to_string(), + "server2".to_string(), + create_test_tool("tool3"), + ); + + let tools = inventory.list_tools(); + assert_eq!(tools.len(), 3); + } + + #[test] + fn test_list_tools_filters_expired() { + let inventory = ToolInventory::new(Duration::from_millis(100)); + + inventory.insert_tool( + "tool1".to_string(), + "server1".to_string(), + create_test_tool("tool1"), + ); + + // Should have 1 tool + assert_eq!(inventory.list_tools().len(), 1); + + // Wait for expiration + std::thread::sleep(Duration::from_millis(150)); + + // Should have 0 tools (filtered out) + assert_eq!(inventory.list_tools().len(), 0); + } + + #[test] + fn test_clear_server_tools() { + let inventory = ToolInventory::new(Duration::from_secs(60)); + + inventory.insert_tool( + "tool1".to_string(), + "server1".to_string(), + create_test_tool("tool1"), + ); + inventory.insert_tool( + "tool2".to_string(), + "server2".to_string(), + create_test_tool("tool2"), + ); + + assert_eq!(inventory.list_tools().len(), 2); + + inventory.clear_server_tools("server1"); + + let tools = inventory.list_tools(); + assert_eq!(tools.len(), 1); + assert_eq!(tools[0].0, "tool2"); + } + + #[test] + fn test_server_refresh_tracking() { + let inventory = ToolInventory::new(Duration::from_secs(60)); + + // Never refreshed + assert!(inventory.needs_refresh("server1", Duration::from_secs(10))); + + // Mark as refreshed + inventory.mark_refreshed("server1"); + + // Should not need refresh immediately + assert!(!inventory.needs_refresh("server1", Duration::from_secs(10))); + + // Wait and check again + std::thread::sleep(Duration::from_millis(100)); + assert!(inventory.needs_refresh("server1", Duration::from_millis(50))); + } + + #[test] + fn test_cleanup_expired() { + let inventory = ToolInventory::new(Duration::from_millis(100)); + + inventory.insert_tool( + "tool1".to_string(), + "server1".to_string(), + create_test_tool("tool1"), + ); + inventory.insert_tool( + "tool2".to_string(), + "server1".to_string(), + create_test_tool("tool2"), + ); + + let (tools, _, _) = inventory.counts(); + assert_eq!(tools, 2); + + // Wait for expiration + std::thread::sleep(Duration::from_millis(150)); + + // Cleanup expired entries + inventory.cleanup_expired(); + + let (tools, _, _) = inventory.counts(); + assert_eq!(tools, 0); + } + + #[test] + fn test_prompt_operations() { + let inventory = ToolInventory::new(Duration::from_secs(60)); + let prompt = create_test_prompt("test_prompt"); + + inventory.insert_prompt( + "test_prompt".to_string(), + "server1".to_string(), + prompt.clone(), + ); + + assert!(inventory.has_prompt("test_prompt")); + + let result = inventory.get_prompt("test_prompt"); + assert!(result.is_some()); + + let (server_name, retrieved_prompt) = result.unwrap(); + assert_eq!(server_name, "server1"); + assert_eq!(retrieved_prompt.name, "test_prompt"); + } + + #[test] + fn test_resource_operations() { + let inventory = ToolInventory::new(Duration::from_secs(60)); + let resource = create_test_resource("file:///test.txt"); + + inventory.insert_resource( + "file:///test.txt".to_string(), + "server1".to_string(), + resource.clone(), + ); + + assert!(inventory.has_resource("file:///test.txt")); + + let result = inventory.get_resource("file:///test.txt"); + assert!(result.is_some()); + + let (server_name, retrieved_resource) = result.unwrap(); + assert_eq!(server_name, "server1"); + assert_eq!(retrieved_resource.uri, "file:///test.txt"); + } + + #[tokio::test] + async fn test_concurrent_access() { + use std::sync::Arc; + + let inventory = Arc::new(ToolInventory::new(Duration::from_secs(60))); + + // Spawn multiple tasks that insert tools concurrently + let mut handles = vec![]; + for i in 0..10 { + let inv = Arc::clone(&inventory); + let handle = tokio::spawn(async move { + let tool = create_test_tool(&format!("tool_{}", i)); + inv.insert_tool(format!("tool_{}", i), format!("server_{}", i % 3), tool); + }); + handles.push(handle); + } + + // Wait for all tasks to complete + for handle in handles { + handle.await.unwrap(); + } + + // Should have 10 tools + let (tools, _, _) = inventory.counts(); + assert_eq!(tools, 10); + } + + #[test] + fn test_clear_all() { + let inventory = ToolInventory::new(Duration::from_secs(60)); + + inventory.insert_tool( + "tool1".to_string(), + "server1".to_string(), + create_test_tool("tool1"), + ); + inventory.insert_prompt( + "prompt1".to_string(), + "server1".to_string(), + create_test_prompt("prompt1"), + ); + inventory.insert_resource( + "res1".to_string(), + "server1".to_string(), + create_test_resource("res1"), + ); + + inventory.mark_refreshed("server1"); + + let (tools, prompts, resources) = inventory.counts(); + assert_eq!(tools, 1); + assert_eq!(prompts, 1); + assert_eq!(resources, 1); + + inventory.clear_all(); + + let (tools, prompts, resources) = inventory.counts(); + assert_eq!(tools, 0); + assert_eq!(prompts, 0); + assert_eq!(resources, 0); + assert!(inventory.last_refresh("server1").is_none()); + } +} diff --git a/sgl-router/src/mcp/manager.rs b/sgl-router/src/mcp/manager.rs new file mode 100644 index 000000000..4bbd18533 --- /dev/null +++ b/sgl-router/src/mcp/manager.rs @@ -0,0 +1,893 @@ +//! Refactored MCP Manager - Single flat structure for all MCP operations +//! +//! This replaces the previous hierarchy: +//! - McpManager (wrapper for static/dynamic distinction) +//! - McpClientManager (manages multiple clients) +//! - McpClient (actual client) +//! +//! New flat structure: +//! - McpManager (single component handling all MCP concerns) +//! - McpClient (actual client to one server) + +use std::{borrow::Cow, sync::Arc, time::Duration}; + +use backoff::ExponentialBackoffBuilder; +use dashmap::DashMap; +use rmcp::{ + model::{ + CallToolRequestParam, CallToolResult, GetPromptRequestParam, GetPromptResult, + ReadResourceRequestParam, ReadResourceResult, SubscribeRequestParam, + UnsubscribeRequestParam, + }, + service::RunningService, + transport::{ + sse_client::SseClientConfig, streamable_http_client::StreamableHttpClientTransportConfig, + ConfigureCommandExt, SseClientTransport, StreamableHttpClientTransport, TokioChildProcess, + }, + RoleClient, ServiceExt, +}; +use serde_json::Map; +use tracing::{debug, error, info, warn}; + +use crate::mcp::{ + config::{ + McpConfig, McpProxyConfig, McpServerConfig, McpTransport, PromptInfo, ResourceInfo, + ToolInfo, + }, + connection_pool::McpConnectionPool, + error::{McpError, McpResult}, + inventory::ToolInventory, +}; +/// Type alias for MCP client +type McpClient = RunningService; + +/// Unified MCP Manager - handles all MCP operations +/// +/// This single component manages: +/// - Client connections (both static and dynamic) +/// - Tool inventory and caching +/// - Connection pooling +/// - Background refresh +/// - Tool/prompt/resource operations +pub struct McpManager { + /// All MCP clients (static + dynamic) + /// Key: server_name for static, server_url for dynamic + /// Using DashMap for concurrent access + clients: Arc>>, + + /// Track which servers are static (from config) + /// Using DashMap for thread-safe mutation during workflow registration + static_servers: Arc>, + + /// Shared tool inventory with TTL and caching + inventory: Arc, + + /// Connection pool for dynamic servers (TTL-based cleanup) + connection_pool: Arc, + + /// Original config for static servers (kept for potential future use) + _config: McpConfig, +} + +impl McpManager { + /// Create a new MCP manager with custom TTLs + pub async fn new( + config: McpConfig, + tool_ttl: Duration, + pool_idle_ttl: Duration, + pool_max_connections: usize, + ) -> McpResult { + // Create shared inventory + let inventory = Arc::new(ToolInventory::new(tool_ttl)); + + // Create connection pool + let connection_pool = Arc::new(McpConnectionPool::with_config( + pool_idle_ttl, + pool_max_connections, + )); + + // Create manager structure + let clients = Arc::new(DashMap::new()); + let static_servers = Arc::new(DashMap::new()); + + // Get global proxy config for all servers + let global_proxy = config.proxy.as_ref(); + + // Connect to all static servers from config + for server_config in &config.servers { + static_servers.insert(server_config.name.clone(), ()); + + match Self::connect_server(server_config, global_proxy).await { + Ok(client) => { + let client_arc = Arc::new(client); + // Load inventory for this server + Self::load_server_inventory(&inventory, &server_config.name, &client_arc).await; + clients.insert(server_config.name.clone(), client_arc); + info!("Connected to static server '{}'", server_config.name); + } + Err(e) => { + error!( + "Failed to connect to static server '{}': {}", + server_config.name, e + ); + } + } + } + + if static_servers.is_empty() || clients.is_empty() { + warn!("No static MCP servers connected"); + } + + Ok(Self { + clients, + static_servers, + inventory, + connection_pool, + _config: config, + }) + } + + /// Create with default settings (300s TTL, 300s idle, 100 max connections) + pub async fn with_defaults(config: McpConfig) -> McpResult { + Self::new( + config, + Duration::from_secs(300), + Duration::from_secs(300), + 100, + ) + .await + } + + // ======================================================================== + // Client Management + // ======================================================================== + + /// Get a client by server name (static or dynamic) + pub async fn get_client(&self, server_name: &str) -> Option> { + self.clients.get(server_name).map(|e| Arc::clone(e.value())) + } + + /// Get or create a dynamic client from server config + pub async fn get_or_create_client( + &self, + server_config: McpServerConfig, + ) -> McpResult> { + // Check if client already exists + let server_key = Self::server_key(&server_config); + + if let Some(client) = self.clients.get(&server_key) { + return Ok(Arc::clone(client.value())); + } + + // Client doesn't exist, create new one via connection pool + let client = self + .connection_pool + .get_or_create( + &server_key, + server_config, + |config, global_proxy| async move { + Self::connect_server(&config, global_proxy.as_ref()).await + }, + ) + .await?; + + // Store in clients map + self.clients.insert(server_key, Arc::clone(&client)); + + Ok(client) + } + + /// List all static server names + pub fn list_static_servers(&self) -> Vec { + self.static_servers + .iter() + .map(|e| e.key().clone()) + .collect() + } + + /// Check if a server is static + pub fn is_static_server(&self, server_name: &str) -> bool { + self.static_servers.contains_key(server_name) + } + + /// Register a static server (called by workflow system) + /// + /// This method registers a static MCP server that was configured and connected + /// via the workflow system. Static servers are never removed during runtime. + /// + /// # Arguments + /// * `name` - Unique server name (from config) + /// * `client` - Connected MCP client + pub fn register_static_server(&self, name: String, client: Arc) { + // Insert into clients map + self.clients.insert(name.clone(), client); + + // Mark as static server (for background refresh and stats) + self.static_servers.insert(name.clone(), ()); + + info!("Registered static MCP server: {}", name); + } + + // ======================================================================== + // Tool Operations (delegate to clients via inventory) + // ======================================================================== + + /// List all available tools from all servers + pub fn list_tools(&self) -> Vec { + self.inventory + .list_tools() + .into_iter() + .map(|(_tool_name, _server_name, tool_info)| tool_info) + .collect() + } + + /// Call a tool by name + pub async fn call_tool( + &self, + tool_name: &str, + args: Option>, + ) -> McpResult { + // Get server that owns this tool + let (server_name, _tool_info) = self + .inventory + .get_tool(tool_name) + .ok_or_else(|| McpError::ToolNotFound(tool_name.to_string()))?; + + // Get client for that server + let client = self + .get_client(&server_name) + .await + .ok_or_else(|| McpError::ServerNotFound(server_name.clone()))?; + + // Call the tool + let request = CallToolRequestParam { + name: Cow::Owned(tool_name.to_string()), + arguments: args, + }; + + client + .call_tool(request) + .await + .map_err(|e| McpError::ToolExecution(format!("Failed to call tool: {}", e))) + } + + /// Get a tool by name + pub fn get_tool(&self, tool_name: &str) -> Option { + self.inventory + .get_tool(tool_name) + .map(|(_server_name, tool_info)| tool_info) + } + + // ======================================================================== + // Prompt Operations + // ======================================================================== + + /// Get a prompt by name + pub async fn get_prompt( + &self, + prompt_name: &str, + args: Option>, + ) -> McpResult { + // Get server that owns this prompt + let (server_name, _prompt_info) = self + .inventory + .get_prompt(prompt_name) + .ok_or_else(|| McpError::PromptNotFound(prompt_name.to_string()))?; + + // Get client for that server + let client = self + .get_client(&server_name) + .await + .ok_or_else(|| McpError::ServerNotFound(server_name.clone()))?; + + // Get the prompt + let request = GetPromptRequestParam { + name: prompt_name.to_string(), + arguments: args, + }; + + client + .get_prompt(request) + .await + .map_err(|e| McpError::Transport(format!("Failed to get prompt: {}", e))) + } + + /// List all available prompts + pub fn list_prompts(&self) -> Vec { + self.inventory + .list_prompts() + .into_iter() + .map(|(_prompt_name, _server_name, prompt_info)| prompt_info) + .collect() + } + + // ======================================================================== + // Resource Operations + // ======================================================================== + + /// Read a resource by URI + pub async fn read_resource(&self, uri: &str) -> McpResult { + // Get server that owns this resource + let (server_name, _resource_info) = self + .inventory + .get_resource(uri) + .ok_or_else(|| McpError::ResourceNotFound(uri.to_string()))?; + + // Get client for that server + let client = self + .get_client(&server_name) + .await + .ok_or_else(|| McpError::ServerNotFound(server_name.clone()))?; + + // Read the resource + let request = ReadResourceRequestParam { + uri: uri.to_string(), + }; + + client + .read_resource(request) + .await + .map_err(|e| McpError::Transport(format!("Failed to read resource: {}", e))) + } + + /// List all available resources + pub fn list_resources(&self) -> Vec { + self.inventory + .list_resources() + .into_iter() + .map(|(_resource_uri, _server_name, resource_info)| resource_info) + .collect() + } + + // ======================================================================== + // Inventory Management + // ======================================================================== + + /// Refresh inventory for a specific server + pub async fn refresh_server_inventory(&self, server_name: &str) -> McpResult<()> { + let client = self + .get_client(server_name) + .await + .ok_or_else(|| McpError::ServerNotFound(server_name.to_string()))?; + + info!("Refreshing inventory for server: {}", server_name); + self.load_server_inventory_internal(server_name, &client) + .await; + Ok(()) + } + + /// Start background refresh for all static servers + pub fn spawn_background_refresh_all( + self: Arc, + refresh_interval: Duration, + ) -> tokio::task::JoinHandle<()> { + tokio::spawn(async move { + let mut interval = tokio::time::interval(refresh_interval); + interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); + + loop { + interval.tick().await; + + let server_names = self.list_static_servers(); + + if !server_names.is_empty() { + debug!( + "Background refresh: Refreshing {} static server(s)", + server_names.len() + ); + + for server_name in server_names { + if let Err(e) = self.refresh_server_inventory(&server_name).await { + warn!("Background refresh failed for '{}': {}", server_name, e); + } + } + + debug!("Background refresh: Completed refresh cycle"); + } + } + }) + } + + // ======================================================================== + // Additional Tool/Prompt/Resource Methods + // ======================================================================== + + /// Check if a tool exists + pub fn has_tool(&self, name: &str) -> bool { + self.inventory.has_tool(name) + } + + /// Get prompt info by name + pub fn get_prompt_info(&self, name: &str) -> Option { + self.inventory.get_prompt(name).map(|(_server, info)| info) + } + + /// Get resource info by URI + pub fn get_resource_info(&self, uri: &str) -> Option { + self.inventory.get_resource(uri).map(|(_server, info)| info) + } + + /// Subscribe to resource changes + pub async fn subscribe_resource(&self, uri: &str) -> McpResult<()> { + let (server_name, _resource_info) = self + .inventory + .get_resource(uri) + .ok_or_else(|| McpError::ResourceNotFound(uri.to_string()))?; + + let client = self + .get_client(&server_name) + .await + .ok_or_else(|| McpError::ServerNotFound(server_name.clone()))?; + + debug!("Subscribing to '{}' on '{}'", uri, server_name); + + client + .peer() + .subscribe(SubscribeRequestParam { + uri: uri.to_string(), + }) + .await + .map_err(|e| McpError::ToolExecution(format!("Failed to subscribe: {}", e))) + } + + /// Unsubscribe from resource changes + pub async fn unsubscribe_resource(&self, uri: &str) -> McpResult<()> { + let (server_name, _resource_info) = self + .inventory + .get_resource(uri) + .ok_or_else(|| McpError::ResourceNotFound(uri.to_string()))?; + + let client = self + .get_client(&server_name) + .await + .ok_or_else(|| McpError::ServerNotFound(server_name.clone()))?; + + debug!("Unsubscribing from '{}' on '{}'", uri, server_name); + + client + .peer() + .unsubscribe(UnsubscribeRequestParam { + uri: uri.to_string(), + }) + .await + .map_err(|e| McpError::ToolExecution(format!("Failed to unsubscribe: {}", e))) + } + + /// List all connected servers + pub fn list_servers(&self) -> Vec { + self.clients.iter().map(|e| e.key().clone()).collect() + } + + /// Disconnect from all servers (for cleanup) + pub async fn shutdown(&self) { + let keys: Vec = self.clients.iter().map(|e| e.key().clone()).collect(); + + for name in keys { + if let Some((_, client)) = self.clients.remove(&name) { + // Try to unwrap Arc to call cancel + match Arc::try_unwrap(client) { + Ok(client) => { + if let Err(e) = client.cancel().await { + warn!("Error disconnecting from '{}': {}", name, e); + } + } + Err(_) => { + warn!("Could not shutdown '{}': client still in use", name); + } + } + } + } + } + + // ======================================================================== + // Statistics and Accessors + // ======================================================================== + + /// Get statistics about the manager + pub fn stats(&self) -> McpManagerStats { + let (tools, prompts, resources) = self.inventory.counts(); + McpManagerStats { + static_server_count: self.static_servers.len(), + pool_stats: self.connection_pool.stats(), + tool_count: tools, + prompt_count: prompts, + resource_count: resources, + } + } + + /// Get the shared tool inventory + pub fn inventory(&self) -> Arc { + Arc::clone(&self.inventory) + } + + /// Get the connection pool + pub fn connection_pool(&self) -> Arc { + Arc::clone(&self.connection_pool) + } + + // ======================================================================== + // Internal Helper Methods + // ======================================================================== + + /// Static helper for loading inventory (for new()) + /// Discover and cache tools/prompts/resources for a connected server + /// + /// This method is public to allow workflow-based inventory loading. + /// It discovers all tools, prompts, and resources from the client and caches them in the inventory. + pub async fn load_server_inventory( + inventory: &Arc, + server_name: &str, + client: &Arc, + ) { + // Tools + match client.peer().list_all_tools().await { + Ok(ts) => { + info!("Discovered {} tools from '{}'", ts.len(), server_name); + for t in ts { + let tool_info = ToolInfo { + name: t.name.to_string(), + description: t.description.as_deref().unwrap_or_default().to_string(), + server: server_name.to_string(), + parameters: Some(serde_json::Value::Object((*t.input_schema).clone())), + }; + inventory.insert_tool(t.name.to_string(), server_name.to_string(), tool_info); + } + } + Err(e) => warn!("Failed to list tools from '{}': {}", server_name, e), + } + + // Prompts + match client.peer().list_all_prompts().await { + Ok(ps) => { + info!("Discovered {} prompts from '{}'", ps.len(), server_name); + for p in ps { + let prompt_info = PromptInfo { + name: p.name.clone(), + description: p.description.clone(), + server: server_name.to_string(), + arguments: p.arguments.clone().map(|args| { + args.into_iter().map(|arg| serde_json::json!(arg)).collect() + }), + }; + inventory.insert_prompt(p.name.clone(), server_name.to_string(), prompt_info); + } + } + Err(e) => debug!("No prompts or failed to list on '{}': {}", server_name, e), + } + + // Resources + match client.peer().list_all_resources().await { + Ok(rs) => { + info!("Discovered {} resources from '{}'", rs.len(), server_name); + for r in rs { + let resource_info = ResourceInfo { + uri: r.uri.clone(), + name: r.name.clone(), + description: r.description.clone(), + mime_type: r.mime_type.clone(), + server: server_name.to_string(), + }; + inventory.insert_resource( + r.uri.clone(), + server_name.to_string(), + resource_info, + ); + } + } + Err(e) => debug!("No resources or failed to list on '{}': {}", server_name, e), + } + + // Mark server as refreshed + inventory.mark_refreshed(server_name); + } + + /// Discover and cache tools/prompts/resources for a connected server (internal wrapper) + async fn load_server_inventory_internal(&self, server_name: &str, client: &McpClient) { + // Tools + match client.peer().list_all_tools().await { + Ok(ts) => { + info!("Discovered {} tools from '{}'", ts.len(), server_name); + for t in ts { + let tool_info = ToolInfo { + name: t.name.to_string(), + description: t.description.as_deref().unwrap_or_default().to_string(), + server: server_name.to_string(), + parameters: Some(serde_json::Value::Object((*t.input_schema).clone())), + }; + self.inventory.insert_tool( + t.name.to_string(), + server_name.to_string(), + tool_info, + ); + } + } + Err(e) => warn!("Failed to list tools from '{}': {}", server_name, e), + } + + // Prompts + match client.peer().list_all_prompts().await { + Ok(ps) => { + info!("Discovered {} prompts from '{}'", ps.len(), server_name); + for p in ps { + let prompt_info = PromptInfo { + name: p.name.clone(), + description: p.description.clone(), + server: server_name.to_string(), + arguments: p.arguments.clone().map(|args| { + args.into_iter().map(|arg| serde_json::json!(arg)).collect() + }), + }; + self.inventory.insert_prompt( + p.name.clone(), + server_name.to_string(), + prompt_info, + ); + } + } + Err(e) => debug!("No prompts or failed to list on '{}': {}", server_name, e), + } + + // Resources + match client.peer().list_all_resources().await { + Ok(rs) => { + info!("Discovered {} resources from '{}'", rs.len(), server_name); + for r in rs { + let resource_info = ResourceInfo { + uri: r.uri.clone(), + name: r.name.clone(), + description: r.description.clone(), + mime_type: r.mime_type.clone(), + server: server_name.to_string(), + }; + self.inventory.insert_resource( + r.uri.clone(), + server_name.to_string(), + resource_info, + ); + } + } + Err(e) => debug!("No resources or failed to list on '{}': {}", server_name, e), + } + + // Mark server as refreshed + self.inventory.mark_refreshed(server_name); + } + + // ======================================================================== + // Connection Logic (from client_manager.rs) + // ======================================================================== + + /// Connect to an MCP server + /// + /// This method is public to allow workflow-based server registration at runtime. + /// It handles connection with automatic retry for network-based transports (SSE/Streamable). + pub async fn connect_server( + config: &McpServerConfig, + global_proxy: Option<&McpProxyConfig>, + ) -> McpResult { + let needs_retry = matches!( + &config.transport, + McpTransport::Sse { .. } | McpTransport::Streamable { .. } + ); + if needs_retry { + Self::connect_server_with_retry(config, global_proxy).await + } else { + Self::connect_server_impl(config, global_proxy).await + } + } + + /// Connect with exponential backoff retry for remote servers + async fn connect_server_with_retry( + config: &McpServerConfig, + global_proxy: Option<&McpProxyConfig>, + ) -> McpResult { + let backoff = ExponentialBackoffBuilder::new() + .with_initial_interval(Duration::from_secs(1)) + .with_max_interval(Duration::from_secs(30)) + .with_max_elapsed_time(Some(Duration::from_secs(30))) + .build(); + + backoff::future::retry(backoff, || async { + match Self::connect_server_impl(config, global_proxy).await { + Ok(client) => Ok(client), + Err(e) => { + if Self::is_permanent_error(&e) { + error!( + "Permanent error connecting to '{}': {} - not retrying", + config.name, e + ); + Err(backoff::Error::permanent(e)) + } else { + warn!("Failed to connect to '{}', retrying: {}", config.name, e); + Err(backoff::Error::transient(e)) + } + } + } + }) + .await + } + + /// Determine if an error is permanent (should not retry) or transient + fn is_permanent_error(error: &McpError) -> bool { + match error { + McpError::Config(_) => true, + McpError::Auth(_) => true, + McpError::ServerNotFound(_) => true, + McpError::Transport(_) => true, + McpError::ConnectionFailed(msg) => { + msg.contains("initialize") + || msg.contains("connection closed") + || msg.contains("connection refused") + || msg.contains("invalid URL") + || msg.contains("not found") + } + _ => false, + } + } + + /// Internal implementation of server connection (stdio/sse/streamable) + async fn connect_server_impl( + config: &McpServerConfig, + global_proxy: Option<&McpProxyConfig>, + ) -> McpResult { + info!( + "Connecting to MCP server '{}' via {:?}", + config.name, config.transport + ); + + match &config.transport { + McpTransport::Stdio { + command, + args, + envs, + } => { + let transport = TokioChildProcess::new( + tokio::process::Command::new(command).configure(|cmd| { + cmd.args(args) + .envs(envs.iter()) + .stderr(std::process::Stdio::inherit()); + }), + ) + .map_err(|e| McpError::Transport(format!("create stdio transport: {}", e)))?; + + let client = ().serve(transport).await.map_err(|e| { + McpError::ConnectionFailed(format!("initialize stdio client: {}", e)) + })?; + + info!("Connected to stdio server '{}'", config.name); + Ok(client) + } + + McpTransport::Sse { url, token } => { + // Resolve proxy configuration + let proxy_config = crate::mcp::proxy::resolve_proxy_config(config, global_proxy); + + // Create HTTP client with proxy support + let client = if token.is_some() { + let mut builder = reqwest::Client::builder() + .timeout(Duration::from_secs(30)) + .connect_timeout(Duration::from_secs(10)); + + // Apply proxy configuration using proxy.rs helper + if let Some(proxy_cfg) = proxy_config { + builder = crate::mcp::proxy::apply_proxy_to_builder(builder, proxy_cfg)?; + } + + // Add Authorization header + builder = builder.default_headers({ + let mut headers = reqwest::header::HeaderMap::new(); + headers.insert( + reqwest::header::AUTHORIZATION, + format!("Bearer {}", token.as_ref().unwrap()) + .parse() + .map_err(|e| McpError::Transport(format!("auth token: {}", e)))?, + ); + headers + }); + + builder + .build() + .map_err(|e| McpError::Transport(format!("build HTTP client: {}", e)))? + } else { + crate::mcp::proxy::create_http_client(proxy_config)? + }; + + let cfg = SseClientConfig { + sse_endpoint: url.clone().into(), + ..Default::default() + }; + + let transport = SseClientTransport::start_with_client(client, cfg) + .await + .map_err(|e| McpError::Transport(format!("create SSE transport: {}", e)))?; + + let client = ().serve(transport).await.map_err(|e| { + McpError::ConnectionFailed(format!("initialize SSE client: {}", e)) + })?; + + info!("Connected to SSE server '{}' at {}", config.name, url); + Ok(client) + } + + McpTransport::Streamable { url, token } => { + // Note: Streamable transport doesn't support proxy yet + let _proxy_config = crate::mcp::proxy::resolve_proxy_config(config, global_proxy); + if _proxy_config.is_some() { + warn!( + "Proxy configuration detected but not supported for Streamable transport on server '{}'", + config.name + ); + } + + let transport = if let Some(tok) = token { + let mut cfg = StreamableHttpClientTransportConfig::with_uri(url.as_str()); + cfg.auth_header = Some(format!("Bearer {}", tok)); + StreamableHttpClientTransport::from_config(cfg) + } else { + StreamableHttpClientTransport::from_uri(url.as_str()) + }; + + let client = ().serve(transport).await.map_err(|e| { + McpError::ConnectionFailed(format!("initialize streamable client: {}", e)) + })?; + + info!( + "Connected to streamable HTTP server '{}' at {}", + config.name, url + ); + Ok(client) + } + } + } + + /// Generate a unique key for a server config + fn server_key(config: &McpServerConfig) -> String { + // Extract URL from transport or use name + match &config.transport { + McpTransport::Streamable { url, .. } => url.clone(), + McpTransport::Sse { url, .. } => url.clone(), + McpTransport::Stdio { command, .. } => command.clone(), + } + } +} + +/// Statistics about the MCP manager +#[derive(Debug, Clone)] +pub struct McpManagerStats { + /// Number of static servers registered + pub static_server_count: usize, + /// Connection pool statistics + pub pool_stats: crate::mcp::connection_pool::PoolStats, + /// Number of cached tools + pub tool_count: usize, + /// Number of cached prompts + pub prompt_count: usize, + /// Number of cached resources + pub resource_count: usize, +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn test_manager_creation() { + let config = McpConfig { + servers: vec![], + pool: Default::default(), + proxy: None, + warmup: vec![], + inventory: Default::default(), + }; + + let manager = McpManager::new( + config, + Duration::from_secs(300), + Duration::from_secs(300), + 100, + ) + .await + .unwrap(); + assert_eq!(manager.list_static_servers().len(), 0); + } +} diff --git a/sgl-router/src/mcp/mod.rs b/sgl-router/src/mcp/mod.rs index 6cebc4c7d..1257c44cb 100644 --- a/sgl-router/src/mcp/mod.rs +++ b/sgl-router/src/mcp/mod.rs @@ -7,12 +7,21 @@ // - Resources: File/data access with subscription support // - OAuth: Secure authentication for remote servers -pub mod client_manager; pub mod config; +pub mod connection_pool; pub mod error; +pub mod inventory; +pub mod manager; pub mod oauth; +pub mod proxy; // Re-export the main types for convenience -pub use client_manager::{McpClientManager, PromptInfo, ResourceInfo, ToolInfo}; -pub use config::{McpConfig, McpServerConfig, McpTransport}; +pub use config::{ + InventoryConfig, McpConfig, McpPoolConfig, McpProxyConfig, McpServerConfig, McpTransport, + PromptInfo, ResourceInfo, ToolInfo, WarmupServer, +}; +pub use connection_pool::{CachedConnection, McpConnectionPool, PoolStats}; pub use error::{McpError, McpResult}; +pub use inventory::ToolInventory; +pub use manager::{McpManager, McpManagerStats}; +pub use proxy::{create_http_client, resolve_proxy_config}; diff --git a/sgl-router/src/mcp/proxy.rs b/sgl-router/src/mcp/proxy.rs new file mode 100644 index 000000000..676ad911d --- /dev/null +++ b/sgl-router/src/mcp/proxy.rs @@ -0,0 +1,253 @@ +// MCP Proxy Configuration and Resolution +// +// This module provides proxy configuration resolution and HTTP client creation +// for MCP server connections. Proxy settings are MCP-specific and do NOT affect +// LLM API traffic. + +use std::time::Duration; + +use crate::mcp::{McpError, McpProxyConfig, McpResult, McpServerConfig}; + +/// Resolve proxy configuration for a server +/// Priority: server.proxy > global.proxy > None +/// +/// # Arguments +/// * `server_config` - Server-specific configuration +/// * `global_proxy` - Global proxy configuration from McpConfig +/// +/// # Returns +/// The resolved proxy configuration, or None for direct connection +pub fn resolve_proxy_config<'a>( + server_config: &'a McpServerConfig, + global_proxy: Option<&'a McpProxyConfig>, +) -> Option<&'a McpProxyConfig> { + // Priority 1: Check if server has explicit proxy config + // Note: server.proxy = Some(config) uses that config + // server.proxy = None (set explicitly in YAML as null) forces direct connection + // server.proxy not set (field missing) falls back to global + if server_config.proxy.is_some() { + server_config.proxy.as_ref() + } else { + // Priority 2: Fall back to global proxy + global_proxy + } +} + +/// Apply proxy configuration to a ClientBuilder +/// +/// This is a reusable helper that applies proxy settings without building the client, +/// allowing additional configuration (like auth headers) to be added afterward. +/// +/// # Arguments +/// * `builder` - The reqwest::ClientBuilder to configure +/// * `proxy_config` - The proxy configuration to apply +/// +/// # Returns +/// The configured builder or error +pub fn apply_proxy_to_builder( + mut builder: reqwest::ClientBuilder, + proxy_cfg: &McpProxyConfig, +) -> McpResult { + // Configure HTTP proxy + if let Some(ref http_proxy) = proxy_cfg.http { + let mut proxy = reqwest::Proxy::http(http_proxy) + .map_err(|e| McpError::Config(format!("Invalid HTTP proxy: {}", e)))?; + + // Apply no_proxy exclusions + if let Some(ref no_proxy) = proxy_cfg.no_proxy { + proxy = proxy.no_proxy(reqwest::NoProxy::from_string(no_proxy)); + } + + // Apply authentication if configured + if let (Some(ref username), Some(ref password)) = (&proxy_cfg.username, &proxy_cfg.password) + { + proxy = proxy.basic_auth(username, password); + } + + builder = builder.proxy(proxy); + } + + // Configure HTTPS proxy + if let Some(ref https_proxy) = proxy_cfg.https { + let mut proxy = reqwest::Proxy::https(https_proxy) + .map_err(|e| McpError::Config(format!("Invalid HTTPS proxy: {}", e)))?; + + // Apply no_proxy exclusions + if let Some(ref no_proxy) = proxy_cfg.no_proxy { + proxy = proxy.no_proxy(reqwest::NoProxy::from_string(no_proxy)); + } + + // Apply authentication if configured + if let (Some(ref username), Some(ref password)) = (&proxy_cfg.username, &proxy_cfg.password) + { + proxy = proxy.basic_auth(username, password); + } + + builder = builder.proxy(proxy); + } + + Ok(builder) +} + +/// Create HTTP client with MCP-specific proxy configuration +/// +/// # Arguments +/// * `proxy_config` - Optional proxy configuration to apply +/// +/// # Returns +/// A configured reqwest::Client or error +pub fn create_http_client(proxy_config: Option<&McpProxyConfig>) -> McpResult { + let mut builder = reqwest::Client::builder() + .timeout(Duration::from_secs(30)) + .connect_timeout(Duration::from_secs(10)); + + // Apply MCP-specific proxy if configured + if let Some(proxy_cfg) = proxy_config { + builder = apply_proxy_to_builder(builder, proxy_cfg)?; + } + + builder + .build() + .map_err(|e| McpError::Transport(format!("Failed to build HTTP client: {}", e))) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::mcp::McpTransport; + + #[test] + fn test_resolve_proxy_no_config() { + let server = McpServerConfig { + name: "test".to_string(), + transport: McpTransport::Sse { + url: "http://localhost:3000/sse".to_string(), + token: None, + }, + proxy: None, + required: false, + }; + + let result = resolve_proxy_config(&server, None); + assert!( + result.is_none(), + "Should return None when no proxy configured" + ); + } + + #[test] + fn test_resolve_proxy_global_only() { + let server = McpServerConfig { + name: "test".to_string(), + transport: McpTransport::Sse { + url: "http://localhost:3000/sse".to_string(), + token: None, + }, + proxy: None, + required: false, + }; + + let global = McpProxyConfig { + http: Some("http://global-proxy:8080".to_string()), + https: None, + no_proxy: None, + username: None, + password: None, + }; + + let result = resolve_proxy_config(&server, Some(&global)); + assert!(result.is_some(), "Should use global proxy"); + assert_eq!( + result.unwrap().http.as_ref().unwrap(), + "http://global-proxy:8080" + ); + } + + #[test] + fn test_resolve_proxy_server_override() { + let server_proxy = McpProxyConfig { + http: Some("http://server-proxy:9090".to_string()), + https: None, + no_proxy: None, + username: None, + password: None, + }; + + let server = McpServerConfig { + name: "test".to_string(), + transport: McpTransport::Sse { + url: "http://localhost:3000/sse".to_string(), + token: None, + }, + proxy: Some(server_proxy), + required: false, + }; + + let global = McpProxyConfig { + http: Some("http://global-proxy:8080".to_string()), + https: None, + no_proxy: None, + username: None, + password: None, + }; + + let result = resolve_proxy_config(&server, Some(&global)); + assert!(result.is_some(), "Should use server-specific proxy"); + assert_eq!( + result.unwrap().http.as_ref().unwrap(), + "http://server-proxy:9090", + "Server proxy should override global" + ); + } + + #[test] + fn test_create_http_client_no_proxy() { + let client = create_http_client(None); + assert!(client.is_ok(), "Should create client without proxy"); + } + + #[test] + fn test_create_http_client_with_proxy() { + let proxy = McpProxyConfig { + http: Some("http://proxy.example.com:8080".to_string()), + https: None, + no_proxy: Some("localhost,127.0.0.1".to_string()), + username: None, + password: None, + }; + + let client = create_http_client(Some(&proxy)); + assert!(client.is_ok(), "Should create client with proxy"); + } + + #[test] + fn test_create_http_client_with_auth() { + let proxy = McpProxyConfig { + http: Some("http://proxy.example.com:8080".to_string()), + https: None, + no_proxy: None, + username: Some("user".to_string()), + password: Some("pass".to_string()), + }; + + let client = create_http_client(Some(&proxy)); + assert!( + client.is_ok(), + "Should create client with proxy authentication" + ); + } + + #[test] + fn test_create_http_client_invalid_proxy() { + let proxy = McpProxyConfig { + http: Some("://invalid".to_string()), // Invalid URL format + https: None, + no_proxy: None, + username: None, + password: None, + }; + + let client = create_http_client(Some(&proxy)); + assert!(client.is_err(), "Should fail with invalid proxy URL"); + } +} diff --git a/sgl-router/src/routers/factory.rs b/sgl-router/src/routers/factory.rs index 54c29eea9..5f6c962f5 100644 --- a/sgl-router/src/routers/factory.rs +++ b/sgl-router/src/routers/factory.rs @@ -127,14 +127,7 @@ impl RouterFactory { return Err("OpenAI mode requires at least one worker URL".to_string()); } - let router = OpenAIRouter::new( - worker_urls, - Some(ctx.router_config.circuit_breaker.clone()), - ctx.response_storage.clone(), - ctx.conversation_storage.clone(), - ctx.conversation_item_storage.clone(), - ) - .await?; + let router = OpenAIRouter::new(worker_urls, ctx).await?; Ok(Box::new(router)) } diff --git a/sgl-router/src/routers/grpc/responses/handlers.rs b/sgl-router/src/routers/grpc/responses/handlers.rs index 060ce1bac..c4d79dc7f 100644 --- a/sgl-router/src/routers/grpc/responses/handlers.rs +++ b/sgl-router/src/routers/grpc/responses/handlers.rs @@ -65,6 +65,7 @@ pub async fn route_responses( response_storage: SharedResponseStorage, conversation_storage: SharedConversationStorage, conversation_item_storage: SharedConversationItemStorage, + mcp_manager: Arc, background_tasks: Arc>>, ) -> Response { // 1. Validate mutually exclusive parameters @@ -113,6 +114,7 @@ pub async fn route_responses( response_storage, conversation_storage, conversation_item_storage, + mcp_manager, ) .await } else if is_background { @@ -125,6 +127,7 @@ pub async fn route_responses( response_storage, conversation_storage, conversation_item_storage, + mcp_manager, background_tasks, ) .await @@ -138,6 +141,7 @@ pub async fn route_responses( response_storage, conversation_storage, conversation_item_storage, + mcp_manager, None, // No response_id for sync None, // No background_tasks for sync ) @@ -167,6 +171,7 @@ async fn route_responses_sync( response_storage: SharedResponseStorage, conversation_storage: SharedConversationStorage, conversation_item_storage: SharedConversationItemStorage, + mcp_manager: Arc, response_id: Option, background_tasks: Option>>>, ) -> Response { @@ -179,6 +184,7 @@ async fn route_responses_sync( response_storage, conversation_storage, conversation_item_storage, + mcp_manager, response_id, background_tasks, ) @@ -209,6 +215,7 @@ async fn route_responses_internal( response_storage: SharedResponseStorage, conversation_storage: SharedConversationStorage, conversation_item_storage: SharedConversationItemStorage, + mcp_manager: Arc, response_id: Option, background_tasks: Option>>>, ) -> Result { @@ -223,7 +230,10 @@ async fn route_responses_internal( // 2. Check if request has MCP tools - if so, use tool loop let responses_response = if let Some(tools) = &request.tools { - if let Some(mcp_manager) = create_mcp_manager_from_request(tools).await { + // Try to create dynamic MCP client from request tools using the manager + if let Some(request_mcp_manager) = + create_mcp_manager_from_request(&mcp_manager, tools).await + { debug!("MCP tools detected, using tool loop"); // Execute with MCP tool loop @@ -234,13 +244,14 @@ async fn route_responses_internal( headers, model_id, components, - mcp_manager, + request_mcp_manager, response_id.clone(), background_tasks, ) .await? } else { - // No MCP manager, execute normally + debug!("Failed to create MCP client from request tools"); + // Fall through to non-MCP execution execute_without_mcp( pipeline, &modified_request, @@ -303,6 +314,7 @@ async fn route_responses_background( response_storage: SharedResponseStorage, conversation_storage: SharedConversationStorage, conversation_item_storage: SharedConversationItemStorage, + mcp_manager: Arc, background_tasks: Arc>>, ) -> Response { // Generate response_id for background tracking @@ -365,6 +377,7 @@ async fn route_responses_background( let response_storage_clone = response_storage.clone(); let conversation_storage_clone = conversation_storage.clone(); let conversation_item_storage_clone = conversation_item_storage.clone(); + let mcp_manager_clone = mcp_manager.clone(); let response_id_clone = response_id.clone(); let background_tasks_clone = background_tasks.clone(); @@ -382,6 +395,7 @@ async fn route_responses_background( response_storage_clone, conversation_storage_clone, conversation_item_storage_clone, + mcp_manager_clone, Some(response_id_clone.clone()), Some(background_tasks_clone.clone()), ) @@ -434,6 +448,7 @@ async fn route_responses_streaming( response_storage: SharedResponseStorage, conversation_storage: SharedConversationStorage, conversation_item_storage: SharedConversationItemStorage, + mcp_manager: Arc, ) -> Response { // 1. Load conversation history let modified_request = match load_conversation_history( @@ -461,7 +476,10 @@ async fn route_responses_streaming( // 2. Check if request has MCP tools - if so, use streaming tool loop if let Some(tools) = &request.tools { - if let Some(mcp_manager) = create_mcp_manager_from_request(tools).await { + // Try to create dynamic MCP client from request tools using the manager + if let Some(request_mcp_manager) = + create_mcp_manager_from_request(&mcp_manager, tools).await + { debug!("MCP tools detected in streaming mode, using streaming tool loop"); return execute_tool_loop_streaming( @@ -471,7 +489,7 @@ async fn route_responses_streaming( headers, model_id, components, - mcp_manager, + request_mcp_manager, response_storage, conversation_storage, conversation_item_storage, diff --git a/sgl-router/src/routers/grpc/responses/tool_loop.rs b/sgl-router/src/routers/grpc/responses/tool_loop.rs index 01601c7c8..8555aa166 100644 --- a/sgl-router/src/routers/grpc/responses/tool_loop.rs +++ b/sgl-router/src/routers/grpc/responses/tool_loop.rs @@ -24,12 +24,11 @@ use super::{ types::BackgroundTaskInfo, }; /// This is a re-export of the shared implementation from openai::mcp -pub(super) use crate::routers::openai::mcp::mcp_manager_from_request_tools as create_mcp_manager_from_request; +pub(super) use crate::routers::openai::mcp::ensure_request_mcp_client as create_mcp_manager_from_request; use crate::{ data_connector::{ SharedConversationItemStorage, SharedConversationStorage, SharedResponseStorage, }, - mcp::McpClientManager, protocols::{ chat::ChatCompletionResponse, common::{Tool, ToolChoice, ToolChoiceValue}, @@ -102,7 +101,7 @@ fn extract_all_tool_calls_from_chat( /// Execute an MCP tool call async fn execute_mcp_call( - mcp_mgr: &Arc, + mcp_mgr: &Arc, tool_name: &str, args_json_str: &str, ) -> Result { @@ -222,7 +221,7 @@ fn generate_mcp_id(prefix: &str) -> String { /// Build mcp_list_tools output item fn build_mcp_list_tools_item( - mcp: &Arc, + mcp: &Arc, server_label: &str, ) -> ResponseOutputItem { let tools = mcp.list_tools(); @@ -287,7 +286,7 @@ pub(super) async fn execute_tool_loop( headers: Option, model_id: Option, components: Arc, - mcp_manager: Arc, + mcp_manager: Arc, response_id: Option, background_tasks: Option>>>, ) -> Result { @@ -507,7 +506,7 @@ pub(super) async fn execute_tool_loop_streaming( headers: Option, model_id: Option, components: Arc, - mcp_manager: Arc, + mcp_manager: Arc, response_storage: SharedResponseStorage, conversation_storage: SharedConversationStorage, conversation_item_storage: SharedConversationItemStorage, @@ -598,7 +597,7 @@ async fn execute_tool_loop_streaming_internal( headers: Option, model_id: Option, components: Arc, - mcp_manager: Arc, + mcp_manager: Arc, server_label: String, _response_storage: SharedResponseStorage, _conversation_storage: SharedConversationStorage, diff --git a/sgl-router/src/routers/grpc/router.rs b/sgl-router/src/routers/grpc/router.rs index 24e183d72..0b8b75fc2 100644 --- a/sgl-router/src/routers/grpc/router.rs +++ b/sgl-router/src/routers/grpc/router.rs @@ -24,6 +24,7 @@ use crate::{ data_connector::{ SharedConversationItemStorage, SharedConversationStorage, SharedResponseStorage, }, + mcp::McpManager, policies::PolicyRegistry, protocols::{ chat::ChatCompletionRequest, @@ -60,8 +61,7 @@ pub struct GrpcRouter { response_storage: SharedResponseStorage, conversation_storage: SharedConversationStorage, conversation_item_storage: SharedConversationItemStorage, - // Optional MCP manager for tool execution (enabled via SGLANG_MCP_CONFIG env var) - mcp_manager: Option>, + mcp_manager: Arc, // Background task handles for cancellation support (includes gRPC client for Python abort) background_tasks: Arc>>, } @@ -94,25 +94,12 @@ impl GrpcRouter { let conversation_storage = ctx.conversation_storage.clone(); let conversation_item_storage = ctx.conversation_item_storage.clone(); - // Optional MCP manager activation via env var path (config-driven gate) - let mcp_manager = match std::env::var("SGLANG_MCP_CONFIG").ok() { - Some(path) if !path.trim().is_empty() => { - match crate::mcp::McpConfig::from_file(&path).await { - Ok(cfg) => match crate::mcp::McpClientManager::new(cfg).await { - Ok(mgr) => Some(Arc::new(mgr)), - Err(err) => { - tracing::warn!("Failed to initialize MCP manager: {}", err); - None - } - }, - Err(err) => { - tracing::warn!("Failed to load MCP config from '{}': {}", path, err); - None - } - } - } - _ => None, - }; + // Get MCP manager from app context + let mcp_manager = ctx + .mcp_manager + .get() + .ok_or_else(|| "gRPC router requires MCP manager".to_string())? + .clone(); // Create shared components for pipeline let shared_components = Arc::new(SharedComponents { @@ -285,6 +272,7 @@ impl RouterTrait for GrpcRouter { self.response_storage.clone(), self.conversation_storage.clone(), self.conversation_item_storage.clone(), + self.mcp_manager.clone(), self.background_tasks.clone(), ) .await diff --git a/sgl-router/src/routers/openai/mcp.rs b/sgl-router/src/routers/openai/mcp.rs index fb0b2d327..ef860655d 100644 --- a/sgl-router/src/routers/openai/mcp.rs +++ b/sgl-router/src/routers/openai/mcp.rs @@ -18,7 +18,7 @@ use tracing::{info, warn}; use super::utils::{event_types, generate_id}; use crate::{ - mcp::McpClientManager, + mcp, protocols::responses::{ResponseInput, ResponseTool, ResponseToolType, ResponsesRequest}, routers::header_utils::apply_request_headers, }; @@ -128,10 +128,19 @@ impl FunctionCallInProgress { // MCP Manager Integration // ============================================================================ -/// Build a request-scoped MCP manager from request tools, if present. -pub async fn mcp_manager_from_request_tools( +/// Ensure a dynamic MCP client exists for request-scoped tools. +/// +/// This function parses request tools to extract MCP server configuration, +/// then ensures a dynamic client exists in the McpManager via `get_or_create_client()`. +/// The McpManager itself is returned (cloned Arc) for convenience, though the main +/// purpose is the side effect of registering the dynamic client. +/// +/// Returns Some(manager) if a dynamic MCP tool was found and client was created/retrieved, +/// None if no MCP tools were found or connection failed. +pub async fn ensure_request_mcp_client( + mcp_manager: &Arc, tools: &[ResponseTool], -) -> Option> { +) -> Option> { let tool = tools .iter() .find(|t| matches!(t.r#type, ResponseToolType::Mcp) && t.server_url.is_some())?; @@ -149,23 +158,30 @@ pub async fn mcp_manager_from_request_tools( .unwrap_or_else(|| "request-mcp".to_string()); let token = tool.authorization.clone(); let transport = if server_url.contains("/sse") { - crate::mcp::McpTransport::Sse { - url: server_url, + mcp::McpTransport::Sse { + url: server_url.clone(), token, } } else { - crate::mcp::McpTransport::Streamable { - url: server_url, + mcp::McpTransport::Streamable { + url: server_url.clone(), token, } }; - let cfg = crate::mcp::McpConfig { - servers: vec![crate::mcp::McpServerConfig { name, transport }], + + // Create server config + let server_config = mcp::McpServerConfig { + name, + transport, + proxy: None, + required: false, }; - match McpClientManager::new(cfg).await { - Ok(mgr) => Some(Arc::new(mgr)), + + // Use McpManager to get or create dynamic client + match mcp_manager.get_or_create_client(server_config).await { + Ok(_client) => Some(mcp_manager.clone()), Err(err) => { - warn!("Failed to initialize request-scoped MCP manager: {}", err); + warn!("Failed to get/create MCP connection: {}", err); None } } @@ -177,7 +193,7 @@ pub async fn mcp_manager_from_request_tools( /// Execute an MCP tool call pub(super) async fn execute_mcp_call( - mcp_mgr: &Arc, + mcp_mgr: &Arc, tool_name: &str, args_json_str: &str, ) -> Result<(String, String), String> { @@ -204,7 +220,7 @@ pub(super) async fn execute_mcp_call( /// Returns false if client disconnected during execution pub(super) async fn execute_streaming_tool_calls( pending_calls: Vec, - active_mcp: &Arc, + active_mcp: &Arc, tx: &mpsc::UnboundedSender>, state: &mut ToolLoopState, server_label: &str, @@ -269,7 +285,7 @@ pub(super) async fn execute_streaming_tool_calls( /// Transform payload to replace MCP tools with function tools for streaming pub(super) fn prepare_mcp_payload_for_streaming( payload: &mut Value, - active_mcp: &Arc, + active_mcp: &Arc, ) { if let Some(obj) = payload.as_object_mut() { // Remove any non-function tools from outgoing payload @@ -377,7 +393,7 @@ pub(super) fn build_resume_payload( /// Returns false if client disconnected pub(super) fn send_mcp_list_tools_events( tx: &mpsc::UnboundedSender>, - mcp: &Arc, + mcp: &Arc, server_label: &str, output_index: usize, sequence_number: &mut u64, @@ -533,7 +549,7 @@ pub(super) fn send_mcp_call_completion_events_with_error( pub(super) fn inject_mcp_metadata_streaming( response: &mut Value, state: &ToolLoopState, - mcp: &Arc, + mcp: &Arc, server_label: &str, ) { if let Some(output_array) = response.get_mut("output").and_then(|v| v.as_array_mut()) { @@ -573,7 +589,7 @@ pub(super) async fn execute_tool_loop( headers: Option<&HeaderMap>, initial_payload: Value, original_body: &ResponsesRequest, - active_mcp: &Arc, + active_mcp: &Arc, config: &McpLoopConfig, ) -> Result { let mut state = ToolLoopState::new(original_body.input.clone()); @@ -734,7 +750,7 @@ pub(super) fn build_incomplete_response( mut response: Value, state: ToolLoopState, reason: &str, - active_mcp: &Arc, + active_mcp: &Arc, original_body: &ResponsesRequest, ) -> Result { let obj = response @@ -837,7 +853,7 @@ pub(super) fn build_incomplete_response( // ============================================================================ /// Build an mcp_list_tools output item -pub(super) fn build_mcp_list_tools_item(mcp: &Arc, server_label: &str) -> Value { +pub(super) fn build_mcp_list_tools_item(mcp: &Arc, server_label: &str) -> Value { let tools = mcp.list_tools(); let tools_json: Vec = tools .iter() diff --git a/sgl-router/src/routers/openai/router.rs b/sgl-router/src/routers/openai/router.rs index 66ee551b7..8ea008e56 100644 --- a/sgl-router/src/routers/openai/router.rs +++ b/sgl-router/src/routers/openai/router.rs @@ -28,7 +28,7 @@ use super::conversations::{ }; use super::{ mcp::{ - execute_tool_loop, mcp_manager_from_request_tools, prepare_mcp_payload_for_streaming, + ensure_request_mcp_client, execute_tool_loop, prepare_mcp_payload_for_streaming, McpLoopConfig, }, responses::{mask_tools_as_mcp, patch_streaming_response_json}, @@ -36,12 +36,12 @@ use super::{ utils::{apply_provider_headers, extract_auth_header, probe_endpoint_for_model}, }; use crate::{ - config::CircuitBreakerConfig, core::{CircuitBreaker, CircuitBreakerConfig as CoreCircuitBreakerConfig}, data_connector::{ ConversationId, ListParams, ResponseId, SharedConversationItemStorage, SharedConversationStorage, SharedResponseStorage, SortOrder, }, + mcp::McpManager, protocols::{ chat::ChatCompletionRequest, classify::ClassifyRequest, @@ -86,8 +86,8 @@ pub struct OpenAIRouter { conversation_storage: SharedConversationStorage, /// Conversation item storage backend conversation_item_storage: SharedConversationItemStorage, - /// Optional MCP manager (enabled via config presence) - mcp_manager: Option>, + /// MCP manager (handles both static and dynamic servers) + mcp_manager: Arc, } impl std::fmt::Debug for OpenAIRouter { @@ -109,15 +109,10 @@ impl OpenAIRouter { /// Create a new OpenAI router pub async fn new( worker_urls: Vec, - circuit_breaker_config: Option, - response_storage: SharedResponseStorage, - conversation_storage: SharedConversationStorage, - conversation_item_storage: SharedConversationItemStorage, + ctx: &Arc, ) -> Result { - let client = reqwest::Client::builder() - .timeout(Duration::from_secs(300)) - .build() - .map_err(|e| format!("Failed to create HTTP client: {}", e))?; + // Use HTTP client from AppContext + let client = ctx.client.clone(); // Normalize URLs (remove trailing slashes) let worker_urls: Vec = worker_urls @@ -125,37 +120,23 @@ impl OpenAIRouter { .map(|url| url.trim_end_matches('/').to_string()) .collect(); - // Convert circuit breaker config - let core_cb_config = circuit_breaker_config - .map(|cb| CoreCircuitBreakerConfig { - failure_threshold: cb.failure_threshold, - success_threshold: cb.success_threshold, - timeout_duration: Duration::from_secs(cb.timeout_duration_secs), - window_duration: Duration::from_secs(cb.window_duration_secs), - }) - .unwrap_or_default(); + // Convert circuit breaker config from AppContext + let cb = &ctx.router_config.circuit_breaker; + let core_cb_config = CoreCircuitBreakerConfig { + failure_threshold: cb.failure_threshold, + success_threshold: cb.success_threshold, + timeout_duration: Duration::from_secs(cb.timeout_duration_secs), + window_duration: Duration::from_secs(cb.window_duration_secs), + }; let circuit_breaker = CircuitBreaker::with_config(core_cb_config); - // Optional MCP manager activation via env var path (config-driven gate) - let mcp_manager = match std::env::var("SGLANG_MCP_CONFIG").ok() { - Some(path) if !path.trim().is_empty() => { - match crate::mcp::McpConfig::from_file(&path).await { - Ok(cfg) => match crate::mcp::McpClientManager::new(cfg).await { - Ok(mgr) => Some(Arc::new(mgr)), - Err(err) => { - warn!("Failed to initialize MCP manager: {}", err); - None - } - }, - Err(err) => { - warn!("Failed to load MCP config from '{}': {}", path, err); - None - } - } - } - _ => None, - }; + // Get MCP manager from AppContext (must be initialized) + let mcp_manager = ctx + .mcp_manager + .get() + .ok_or_else(|| "MCP manager not initialized in AppContext".to_string())? + .clone(); Ok(Self { client, @@ -163,9 +144,9 @@ impl OpenAIRouter { model_cache: Arc::new(DashMap::new()), circuit_breaker, healthy: AtomicBool::new(true), - response_storage, - conversation_storage, - conversation_item_storage, + response_storage: ctx.response_storage.clone(), + conversation_storage: ctx.conversation_storage.clone(), + conversation_item_storage: ctx.conversation_item_storage.clone(), mcp_manager, }) } @@ -241,12 +222,17 @@ impl OpenAIRouter { original_previous_response_id: Option, ) -> Response { // Check if MCP is active for this request - let req_mcp_manager = if let Some(ref tools) = original_body.tools { - mcp_manager_from_request_tools(tools.as_slice()).await - } else { + // Ensure dynamic client is created if needed + if let Some(ref tools) = original_body.tools { + ensure_request_mcp_client(&self.mcp_manager, tools.as_slice()).await; + } + + // Use the tool loop if the manager has any tools available (static or dynamic). + let active_mcp = if self.mcp_manager.list_tools().is_empty() { None + } else { + Some(&self.mcp_manager) }; - let active_mcp = req_mcp_manager.as_ref().or(self.mcp_manager.as_ref()); let mut response_json: Value; @@ -984,7 +970,7 @@ impl crate::routers::RouterTrait for OpenAIRouter { handle_streaming_response( &self.client, &self.circuit_breaker, - self.mcp_manager.as_ref(), + Some(&self.mcp_manager), self.response_storage.clone(), self.conversation_storage.clone(), self.conversation_item_storage.clone(), diff --git a/sgl-router/src/routers/openai/streaming.rs b/sgl-router/src/routers/openai/streaming.rs index 804144446..e1e976d8f 100644 --- a/sgl-router/src/routers/openai/streaming.rs +++ b/sgl-router/src/routers/openai/streaming.rs @@ -25,8 +25,8 @@ use tracing::warn; use super::conversations::persist_conversation_items; use super::{ mcp::{ - build_resume_payload, execute_streaming_tool_calls, inject_mcp_metadata_streaming, - mcp_manager_from_request_tools, prepare_mcp_payload_for_streaming, + build_resume_payload, ensure_request_mcp_client, execute_streaming_tool_calls, + inject_mcp_metadata_streaming, prepare_mcp_payload_for_streaming, send_mcp_list_tools_events, McpLoopConfig, ToolLoopState, }, responses::{mask_tools_as_mcp, patch_streaming_response_json, rewrite_streaming_block}, @@ -907,7 +907,7 @@ pub(super) fn send_final_response_event( tx: &mpsc::UnboundedSender>, sequence_number: &mut u64, state: &ToolLoopState, - active_mcp: Option<&Arc>, + active_mcp: Option<&Arc>, original_request: &ResponsesRequest, previous_response_id: Option<&str>, server_label: &str, @@ -1138,7 +1138,7 @@ pub(super) async fn handle_streaming_with_tool_interception( mut payload: Value, original_body: &ResponsesRequest, original_previous_response_id: Option, - active_mcp: &Arc, + active_mcp: &Arc, ) -> Response { // Transform MCP tools to function tools in payload prepare_mcp_payload_for_streaming(&mut payload, active_mcp); @@ -1491,7 +1491,7 @@ pub(super) async fn handle_streaming_with_tool_interception( pub(super) async fn handle_streaming_response( client: &reqwest::Client, circuit_breaker: &crate::core::CircuitBreaker, - mcp_manager: Option<&Arc>, + mcp_manager: Option<&Arc>, response_storage: SharedResponseStorage, conversation_storage: SharedConversationStorage, conversation_item_storage: SharedConversationItemStorage, @@ -1502,12 +1502,19 @@ pub(super) async fn handle_streaming_response( original_previous_response_id: Option, ) -> Response { // Check if MCP is active for this request - let req_mcp_manager = if let Some(ref tools) = original_body.tools { - mcp_manager_from_request_tools(tools.as_slice()).await - } else { - None - }; - let active_mcp = req_mcp_manager.as_ref().or(mcp_manager); + // Ensure dynamic client is created if needed + if let (Some(manager), Some(ref tools)) = (mcp_manager, &original_body.tools) { + ensure_request_mcp_client(manager, tools.as_slice()).await; + } + + // Use the tool loop if the manager has any tools available (static or dynamic). + let active_mcp = mcp_manager.and_then(|mgr| { + if mgr.list_tools().is_empty() { + None + } else { + Some(mgr) + } + }); // If no MCP is active, use simple pass-through streaming if active_mcp.is_none() { diff --git a/sgl-router/src/server.rs b/sgl-router/src/server.rs index de9ff25ed..951957b9f 100644 --- a/sgl-router/src/server.rs +++ b/sgl-router/src/server.rs @@ -24,8 +24,8 @@ use crate::{ core::{ worker_to_info, workflow::{ - create_worker_registration_workflow, create_worker_removal_workflow, LoggingSubscriber, - WorkflowEngine, + create_mcp_registration_workflow, create_worker_registration_workflow, + create_worker_removal_workflow, LoggingSubscriber, WorkflowEngine, }, Job, JobQueue, JobQueueConfig, WorkerManager, WorkerType, }, @@ -739,11 +739,12 @@ pub async fn startup(config: ServerConfig) -> Result<(), Box Result<(), Box Arc { +pub async fn create_test_context(config: RouterConfig) -> Arc { let client = reqwest::Client::new(); // Initialize rate limiter @@ -62,9 +62,10 @@ pub fn create_test_context(config: RouterConfig) -> Arc { config.worker_startup_check_interval_secs, ))); - // Create empty OnceLock for worker job queue and workflow engine + // Create empty OnceLock for worker job queue, workflow engine, and mcp manager let worker_job_queue = Arc::new(OnceLock::new()); let workflow_engine = Arc::new(OnceLock::new()); + let mcp_manager_lock = Arc::new(OnceLock::new()); let app_context = Arc::new( AppContext::builder() @@ -82,6 +83,7 @@ pub fn create_test_context(config: RouterConfig) -> Arc { .load_monitor(load_monitor) .worker_job_queue(worker_job_queue) .workflow_engine(workflow_engine) + .mcp_manager(mcp_manager_lock) .build() .unwrap(), ); @@ -109,6 +111,130 @@ pub fn create_test_context(config: RouterConfig) -> Arc { .set(engine) .expect("WorkflowEngine should only be initialized once"); + // Initialize MCP manager with empty config + use sglang_router_rs::mcp::{McpConfig, McpManager}; + let empty_config = McpConfig { + servers: vec![], + pool: Default::default(), + proxy: None, + warmup: vec![], + inventory: Default::default(), + }; + let mcp_manager = McpManager::with_defaults(empty_config) + .await + .expect("Failed to create MCP manager"); + app_context + .mcp_manager + .set(Arc::new(mcp_manager)) + .ok() + .expect("McpManager should only be initialized once"); + + app_context +} + +/// Helper function to create AppContext for tests with MCP config from file +pub async fn create_test_context_with_mcp_config( + config: RouterConfig, + mcp_config_path: &str, +) -> Arc { + use sglang_router_rs::mcp::{McpConfig, McpManager}; + + let client = reqwest::Client::new(); + + // Initialize rate limiter + let rate_limiter = match config.max_concurrent_requests { + n if n <= 0 => None, + n => { + let rate_limit_tokens = config + .rate_limit_tokens_per_second + .filter(|&t| t > 0) + .unwrap_or(n); + Some(Arc::new(TokenBucket::new( + n as usize, + rate_limit_tokens as usize, + ))) + } + }; + + // Initialize registries + let worker_registry = Arc::new(WorkerRegistry::new()); + let policy_registry = Arc::new(PolicyRegistry::new(config.policy.clone())); + + // Initialize storage backends (Memory for tests) + let response_storage = Arc::new(MemoryResponseStorage::new()); + let conversation_storage = Arc::new(MemoryConversationStorage::new()); + let conversation_item_storage = Arc::new(MemoryConversationItemStorage::new()); + + // Initialize load monitor + let load_monitor = Some(Arc::new(LoadMonitor::new( + worker_registry.clone(), + policy_registry.clone(), + client.clone(), + config.worker_startup_check_interval_secs, + ))); + + // Create empty OnceLock for worker job queue, workflow engine, and mcp manager + let worker_job_queue = Arc::new(OnceLock::new()); + let workflow_engine = Arc::new(OnceLock::new()); + let mcp_manager_lock = Arc::new(OnceLock::new()); + + let app_context = Arc::new( + AppContext::builder() + .router_config(config) + .client(client) + .rate_limiter(rate_limiter) + .tokenizer(None) // tokenizer + .reasoning_parser_factory(None) // reasoning_parser_factory + .tool_parser_factory(None) // tool_parser_factory + .worker_registry(worker_registry) + .policy_registry(policy_registry) + .response_storage(response_storage) + .conversation_storage(conversation_storage) + .conversation_item_storage(conversation_item_storage) + .load_monitor(load_monitor) + .worker_job_queue(worker_job_queue) + .workflow_engine(workflow_engine) + .mcp_manager(mcp_manager_lock) + .build() + .unwrap(), + ); + + // Initialize JobQueue after AppContext is created + let weak_context = Arc::downgrade(&app_context); + let job_queue = sglang_router_rs::core::JobQueue::new( + sglang_router_rs::core::JobQueueConfig::default(), + weak_context, + ); + app_context + .worker_job_queue + .set(job_queue) + .expect("JobQueue should only be initialized once"); + + // Initialize WorkflowEngine and register workflows + use sglang_router_rs::core::workflow::{ + create_worker_registration_workflow, create_worker_removal_workflow, WorkflowEngine, + }; + let engine = Arc::new(WorkflowEngine::new()); + engine.register_workflow(create_worker_registration_workflow()); + engine.register_workflow(create_worker_removal_workflow()); + app_context + .workflow_engine + .set(engine) + .expect("WorkflowEngine should only be initialized once"); + + // Initialize MCP manager from config file + let mcp_config = McpConfig::from_file(mcp_config_path) + .await + .expect("Failed to load MCP config from file"); + let mcp_manager = McpManager::with_defaults(mcp_config) + .await + .expect("Failed to create MCP manager"); + app_context + .mcp_manager + .set(Arc::new(mcp_manager)) + .ok() + .expect("McpManager should only be initialized once"); + app_context } diff --git a/sgl-router/tests/common/test_app.rs b/sgl-router/tests/common/test_app.rs index 41577b950..4cc2f6671 100644 --- a/sgl-router/tests/common/test_app.rs +++ b/sgl-router/tests/common/test_app.rs @@ -9,6 +9,7 @@ use sglang_router_rs::{ data_connector::{ MemoryConversationItemStorage, MemoryConversationStorage, MemoryResponseStorage, }, + mcp::{McpConfig, McpManager}, middleware::{AuthConfig, TokenBucket}, policies::PolicyRegistry, routers::RouterTrait, @@ -153,3 +154,58 @@ pub fn create_test_app_with_context( router_config.cors_allowed_origins.clone(), ) } + +/// Create a minimal test AppContext for unit tests +#[allow(dead_code)] +pub async fn create_test_app_context() -> Arc { + let router_config = RouterConfig::default(); + let client = Client::new(); + + // Initialize empty OnceLocks + let worker_job_queue = Arc::new(OnceLock::new()); + let workflow_engine = Arc::new(OnceLock::new()); + + // Initialize MCP manager with empty config + let mcp_manager_lock = Arc::new(OnceLock::new()); + let empty_config = McpConfig { + servers: vec![], + pool: Default::default(), + proxy: None, + warmup: vec![], + inventory: Default::default(), + }; + let mcp_manager = McpManager::with_defaults(empty_config) + .await + .expect("Failed to create MCP manager"); + mcp_manager_lock.set(Arc::new(mcp_manager)).ok(); + + // Initialize registries + let worker_registry = Arc::new(WorkerRegistry::new()); + let policy_registry = Arc::new(PolicyRegistry::new(router_config.policy.clone())); + + // Initialize storage backends + let response_storage = Arc::new(MemoryResponseStorage::new()); + let conversation_storage = Arc::new(MemoryConversationStorage::new()); + let conversation_item_storage = Arc::new(MemoryConversationItemStorage::new()); + + Arc::new( + AppContext::builder() + .router_config(router_config) + .client(client) + .rate_limiter(None) + .tokenizer(None) + .reasoning_parser_factory(None) + .tool_parser_factory(None) + .worker_registry(worker_registry) + .policy_registry(policy_registry) + .response_storage(response_storage) + .conversation_storage(conversation_storage) + .conversation_item_storage(conversation_item_storage) + .load_monitor(None) + .worker_job_queue(worker_job_queue) + .workflow_engine(workflow_engine) + .mcp_manager(mcp_manager_lock) + .build() + .unwrap(), + ) +} diff --git a/sgl-router/tests/mcp_test.rs b/sgl-router/tests/mcp_test.rs index fb1c4404c..5327fff8c 100644 --- a/sgl-router/tests/mcp_test.rs +++ b/sgl-router/tests/mcp_test.rs @@ -13,7 +13,7 @@ use std::collections::HashMap; use common::mock_mcp_server::MockMCPServer; use serde_json::json; -use sglang_router_rs::mcp::{McpClientManager, McpConfig, McpError, McpServerConfig, McpTransport}; +use sglang_router_rs::mcp::{McpConfig, McpError, McpManager, McpServerConfig, McpTransport}; /// Create a new mock server for testing (each test gets its own) async fn create_mock_server() -> MockMCPServer { @@ -26,11 +26,23 @@ async fn create_mock_server() -> MockMCPServer { #[tokio::test] async fn test_mcp_server_initialization() { - let config = McpConfig { servers: vec![] }; + let config = McpConfig { + servers: vec![], + pool: Default::default(), + proxy: None, + warmup: Vec::new(), + inventory: Default::default(), + }; - // Should fail with no servers - let result = McpClientManager::new(config).await; - assert!(result.is_err(), "Should fail with no servers configured"); + // Should succeed but with no connected servers (empty config is allowed) + let result = McpManager::with_defaults(config).await; + assert!(result.is_ok(), "Should succeed with empty config"); + + let manager = result.unwrap(); + let servers = manager.list_servers(); + assert_eq!(servers.len(), 0, "Should have no servers"); + let tools = manager.list_tools(); + assert_eq!(tools.len(), 0, "Should have no tools"); } #[tokio::test] @@ -44,13 +56,19 @@ async fn test_server_connection_with_mock() { url: mock_server.url(), token: None, }, + proxy: None, + required: false, }], + pool: Default::default(), + proxy: None, + warmup: Vec::new(), + inventory: Default::default(), }; - let result = McpClientManager::new(config).await; + let result = McpManager::with_defaults(config).await; assert!(result.is_ok(), "Should connect to mock server"); - let mut manager = result.unwrap(); + let manager = result.unwrap(); let servers = manager.list_servers(); assert_eq!(servers.len(), 1); @@ -76,10 +94,16 @@ async fn test_tool_availability_checking() { url: mock_server.url(), token: None, }, + proxy: None, + required: false, }], + pool: Default::default(), + proxy: None, + warmup: Vec::new(), + inventory: Default::default(), }; - let mut manager = McpClientManager::new(config).await.unwrap(); + let manager = McpManager::with_defaults(config).await.unwrap(); let test_tools = vec!["brave_web_search", "brave_local_search", "calculator"]; for tool in test_tools { @@ -119,6 +143,8 @@ async fn test_multi_server_connection() { url: mock_server1.url(), token: None, }, + proxy: None, + required: false, }, McpServerConfig { name: "mock_server_2".to_string(), @@ -126,15 +152,21 @@ async fn test_multi_server_connection() { url: mock_server2.url(), token: None, }, + proxy: None, + required: false, }, ], + pool: Default::default(), + proxy: None, + warmup: Vec::new(), + inventory: Default::default(), }; // Note: This will fail to connect to both servers in the current implementation // since they return the same tools. The manager will connect to the first one. - let result = McpClientManager::new(config).await; + let result = McpManager::with_defaults(config).await; - if let Ok(mut manager) = result { + if let Ok(manager) = result { let servers = manager.list_servers(); assert!(!servers.is_empty(), "Should have at least one server"); @@ -156,10 +188,16 @@ async fn test_tool_execution_with_mock() { url: mock_server.url(), token: None, }, + proxy: None, + required: false, }], + pool: Default::default(), + proxy: None, + warmup: Vec::new(), + inventory: Default::default(), }; - let mut manager = McpClientManager::new(config).await.unwrap(); + let manager = McpManager::with_defaults(config).await.unwrap(); let result = manager .call_tool( @@ -207,10 +245,16 @@ async fn test_concurrent_tool_execution() { url: mock_server.url(), token: None, }, + proxy: None, + required: false, }], + pool: Default::default(), + proxy: None, + warmup: Vec::new(), + inventory: Default::default(), }; - let mut manager = McpClientManager::new(config).await.unwrap(); + let manager = McpManager::with_defaults(config).await.unwrap(); // Execute tools sequentially (true concurrent execution would require Arc) let tool_calls = vec![ @@ -244,10 +288,16 @@ async fn test_tool_execution_errors() { url: mock_server.url(), token: None, }, + proxy: None, + required: false, }], + pool: Default::default(), + proxy: None, + warmup: Vec::new(), + inventory: Default::default(), }; - let mut manager = McpClientManager::new(config).await.unwrap(); + let manager = McpManager::with_defaults(config).await.unwrap(); // Try to call unknown tool let result = manager @@ -275,23 +325,25 @@ async fn test_connection_without_server() { args: vec![], envs: HashMap::new(), }, + proxy: None, + required: false, }], + pool: Default::default(), + proxy: None, + warmup: Vec::new(), + inventory: Default::default(), }; - let result = McpClientManager::new(config).await; - assert!(result.is_err(), "Should fail when no server is running"); + let result = McpManager::with_defaults(config).await; + // Manager succeeds but no servers are connected (errors are logged) + assert!( + result.is_ok(), + "Manager should succeed even if servers fail to connect" + ); - if let Err(e) = result { - let error_msg = e.to_string(); - assert!( - error_msg.contains("Failed to connect") - || error_msg.contains("Connection") - || error_msg.contains("failed") - || error_msg.contains("error"), - "Error should indicate failure: {}", - error_msg - ); - } + let manager = result.unwrap(); + let servers = manager.list_servers(); + assert_eq!(servers.len(), 0, "Should have no connected servers"); } // Schema Validation Tests @@ -307,10 +359,16 @@ async fn test_tool_info_structure() { url: mock_server.url(), token: None, }, + proxy: None, + required: false, }], + pool: Default::default(), + proxy: None, + warmup: Vec::new(), + inventory: Default::default(), }; - let manager = McpClientManager::new(config).await.unwrap(); + let manager = McpManager::with_defaults(config).await.unwrap(); let tools = manager.list_tools(); let brave_search = tools @@ -337,12 +395,25 @@ async fn test_sse_connection() { args: vec!["--sse".to_string()], envs: HashMap::new(), }, + proxy: None, + required: false, }], + pool: Default::default(), + proxy: None, + warmup: Vec::new(), + inventory: Default::default(), }; - // This will fail immediately without retry - let result = McpClientManager::new(config).await; - assert!(result.is_err(), "Should fail for non-existent SSE server"); + // Manager succeeds but no servers are connected (errors are logged) + let result = McpManager::with_defaults(config).await; + assert!( + result.is_ok(), + "Manager should succeed even if SSE server fails to connect" + ); + + let manager = result.unwrap(); + let servers = manager.list_servers(); + assert_eq!(servers.len(), 0, "Should have no connected servers"); } // Connection Type Tests @@ -356,6 +427,8 @@ async fn test_transport_types() { url: "http://localhost:8080/mcp".to_string(), token: Some("auth_token".to_string()), }, + proxy: None, + required: false, }; assert_eq!(http_config.name, "http_server"); @@ -366,6 +439,8 @@ async fn test_transport_types() { url: "http://localhost:8081/sse".to_string(), token: None, }, + proxy: None, + required: false, }; assert_eq!(sse_config.name, "sse_server"); @@ -377,6 +452,8 @@ async fn test_transport_types() { args: vec!["--port".to_string(), "8082".to_string()], envs: HashMap::new(), }, + proxy: None, + required: false, }; assert_eq!(stdio_config.name, "stdio_server"); } @@ -395,11 +472,17 @@ async fn test_complete_workflow() { url: mock_server.url(), token: None, }, + proxy: None, + required: false, }], + pool: Default::default(), + proxy: None, + warmup: Vec::new(), + inventory: Default::default(), }; // 2. Connect to server - let mut manager = McpClientManager::new(config) + let manager = McpManager::with_defaults(config) .await .expect("Should connect to mock server"); diff --git a/sgl-router/tests/request_formats_test.rs b/sgl-router/tests/request_formats_test.rs index 76feafa70..e7e40b78d 100644 --- a/sgl-router/tests/request_formats_test.rs +++ b/sgl-router/tests/request_formats_test.rs @@ -44,7 +44,7 @@ impl TestContext { worker_urls: worker_urls.clone(), }; - let app_context = common::create_test_context(config.clone()); + let app_context = common::create_test_context(config.clone()).await; let router = RouterFactory::create_router(&app_context).await.unwrap(); let router = Arc::from(router); diff --git a/sgl-router/tests/responses_api_test.rs b/sgl-router/tests/responses_api_test.rs index 78a8b596e..92ced5ff7 100644 --- a/sgl-router/tests/responses_api_test.rs +++ b/sgl-router/tests/responses_api_test.rs @@ -55,8 +55,9 @@ async fn test_non_streaming_mcp_minimal_e2e_with_persistence() { .queue_timeout_secs(5) .build_unchecked(); - // Create router and context - let ctx = common::create_test_context(router_cfg); + // Create router and context with MCP config from file + let ctx = + common::create_test_context_with_mcp_config(router_cfg, cfg_path.to_str().unwrap()).await; let router = RouterFactory::create_router(&ctx).await.expect("router"); // Build a simple ResponsesRequest that will trigger the tool call @@ -230,7 +231,7 @@ async fn test_conversations_crud_basic() { .queue_timeout_secs(5) .build_unchecked(); - let ctx = common::create_test_context(router_cfg); + let ctx = common::create_test_context(router_cfg).await; let router = RouterFactory::create_router(&ctx).await.expect("router"); // Create @@ -540,7 +541,7 @@ async fn test_multi_turn_loop_with_mcp() { .queue_timeout_secs(5) .build_unchecked(); - let ctx = common::create_test_context(router_cfg); + let ctx = common::create_test_context(router_cfg).await; let router = RouterFactory::create_router(&ctx).await.expect("router"); // Build request with MCP tools @@ -691,7 +692,7 @@ async fn test_max_tool_calls_limit() { .queue_timeout_secs(5) .build_unchecked(); - let ctx = common::create_test_context(router_cfg); + let ctx = common::create_test_context(router_cfg).await; let router = RouterFactory::create_router(&ctx).await.expect("router"); let req = ResponsesRequest { @@ -808,7 +809,8 @@ async fn setup_streaming_mcp_test() -> ( .queue_timeout_secs(5) .build_unchecked(); - let ctx = common::create_test_context(router_cfg); + let ctx = + common::create_test_context_with_mcp_config(router_cfg, cfg_path.to_str().unwrap()).await; let router = RouterFactory::create_router(&ctx).await.expect("router"); (mcp, worker, router, dir) @@ -1224,7 +1226,7 @@ async fn test_conversation_items_create_and_get() { .queue_timeout_secs(5) .build_unchecked(); - let ctx = common::create_test_context(router_cfg); + let ctx = common::create_test_context(router_cfg).await; let router = RouterFactory::create_router(&ctx).await.expect("router"); // Create conversation @@ -1300,7 +1302,7 @@ async fn test_conversation_items_delete() { .queue_timeout_secs(5) .build_unchecked(); - let ctx = common::create_test_context(router_cfg); + let ctx = common::create_test_context(router_cfg).await; let router = RouterFactory::create_router(&ctx).await.expect("router"); // Create conversation @@ -1382,7 +1384,7 @@ async fn test_conversation_items_max_limit() { .queue_timeout_secs(5) .build_unchecked(); - let ctx = common::create_test_context(router_cfg); + let ctx = common::create_test_context(router_cfg).await; let router = RouterFactory::create_router(&ctx).await.expect("router"); // Create conversation @@ -1434,7 +1436,7 @@ async fn test_conversation_items_unsupported_type() { .queue_timeout_secs(5) .build_unchecked(); - let ctx = common::create_test_context(router_cfg); + let ctx = common::create_test_context(router_cfg).await; let router = RouterFactory::create_router(&ctx).await.expect("router"); // Create conversation @@ -1485,7 +1487,7 @@ async fn test_conversation_items_multi_conversation_sharing() { .queue_timeout_secs(5) .build_unchecked(); - let ctx = common::create_test_context(router_cfg); + let ctx = common::create_test_context(router_cfg).await; let router = RouterFactory::create_router(&ctx).await.expect("router"); // Create two conversations diff --git a/sgl-router/tests/streaming_tests.rs b/sgl-router/tests/streaming_tests.rs index 81aed7e74..701f218b1 100644 --- a/sgl-router/tests/streaming_tests.rs +++ b/sgl-router/tests/streaming_tests.rs @@ -45,7 +45,7 @@ impl TestContext { worker_urls: worker_urls.clone(), }; - let app_context = common::create_test_context(config.clone()); + let app_context = common::create_test_context(config.clone()).await; let router = RouterFactory::create_router(&app_context).await.unwrap(); let router = Arc::from(router); diff --git a/sgl-router/tests/test_openai_routing.rs b/sgl-router/tests/test_openai_routing.rs index c759e23dc..38ec3d163 100644 --- a/sgl-router/tests/test_openai_routing.rs +++ b/sgl-router/tests/test_openai_routing.rs @@ -21,10 +21,7 @@ use sglang_router_rs::{ config::{ ConfigError, ConfigValidator, HistoryBackend, OracleConfig, RouterConfig, RoutingMode, }, - data_connector::{ - MemoryConversationItemStorage, MemoryConversationStorage, MemoryResponseStorage, - ResponseId, ResponseStorage, StoredResponse, - }, + data_connector::{ResponseId, StoredResponse}, protocols::{ chat::{ChatCompletionRequest, ChatMessage, UserMessageContent}, common::StringOrArray, @@ -98,14 +95,8 @@ fn create_minimal_completion_request() -> CompletionRequest { /// Test basic OpenAI router creation and configuration #[tokio::test] async fn test_openai_router_creation() { - let router = OpenAIRouter::new( - vec!["https://api.openai.com".to_string()], - None, - Arc::new(MemoryResponseStorage::new()), - Arc::new(MemoryConversationStorage::new()), - Arc::new(MemoryConversationItemStorage::new()), - ) - .await; + let ctx = common::test_app::create_test_app_context().await; + let router = OpenAIRouter::new(vec!["https://api.openai.com".to_string()], &ctx).await; assert!(router.is_ok(), "Router creation should succeed"); @@ -117,15 +108,10 @@ async fn test_openai_router_creation() { /// Test server info endpoint #[tokio::test] async fn test_openai_router_server_info() { - let router = OpenAIRouter::new( - vec!["https://api.openai.com".to_string()], - None, - Arc::new(MemoryResponseStorage::new()), - Arc::new(MemoryConversationStorage::new()), - Arc::new(MemoryConversationItemStorage::new()), - ) - .await - .unwrap(); + let ctx = common::test_app::create_test_app_context().await; + let router = OpenAIRouter::new(vec!["https://api.openai.com".to_string()], &ctx) + .await + .unwrap(); let req = Request::builder() .method(Method::GET) @@ -148,15 +134,10 @@ async fn test_openai_router_server_info() { async fn test_openai_router_models() { // Use mock server for deterministic models response let mock_server = MockOpenAIServer::new().await; - let router = OpenAIRouter::new( - vec![mock_server.base_url()], - None, - Arc::new(MemoryResponseStorage::new()), - Arc::new(MemoryConversationStorage::new()), - Arc::new(MemoryConversationItemStorage::new()), - ) - .await - .unwrap(); + let ctx = common::test_app::create_test_app_context().await; + let router = OpenAIRouter::new(vec![mock_server.base_url()], &ctx) + .await + .unwrap(); let req = Request::builder() .method(Method::GET) @@ -226,17 +207,12 @@ async fn test_openai_router_responses_with_mock() { }); let base_url = format!("http://{}", addr); - let storage = Arc::new(MemoryResponseStorage::new()); - let router = OpenAIRouter::new( - vec![base_url], - None, - storage.clone(), - Arc::new(MemoryConversationStorage::new()), - Arc::new(MemoryConversationItemStorage::new()), - ) - .await - .unwrap(); + let ctx = common::test_app::create_test_app_context().await; + let router = OpenAIRouter::new(vec![base_url], &ctx).await.unwrap(); + + // Get storage from context (router uses this, not a separate storage) + let storage = ctx.response_storage.clone(); let request1 = ResponsesRequest { model: "gpt-4o-mini".to_string(), @@ -495,25 +471,18 @@ async fn test_openai_router_responses_streaming_with_mock() { }); let base_url = format!("http://{}", addr); - let storage = Arc::new(MemoryResponseStorage::new()); - // Seed a previous response so previous_response_id logic has data to pull from. + let ctx = common::test_app::create_test_app_context().await; + let router = OpenAIRouter::new(vec![base_url], &ctx).await.unwrap(); + + // Get storage from context and seed a previous response + let storage = ctx.response_storage.clone(); let mut previous = StoredResponse::new(None); previous.id = ResponseId::from("resp_prev_chain"); previous.input = serde_json::json!("Earlier bedtime question"); previous.output = serde_json::json!("Earlier answer"); storage.store_response(previous).await.unwrap(); - let router = OpenAIRouter::new( - vec![base_url], - None, - storage.clone(), - Arc::new(MemoryConversationStorage::new()), - Arc::new(MemoryConversationItemStorage::new()), - ) - .await - .unwrap(); - let mut metadata = HashMap::new(); metadata.insert("topic".to_string(), json!("unicorns")); @@ -611,7 +580,7 @@ async fn test_router_factory_openai_mode() { let router_config = RouterConfig::new(routing_mode, sglang_router_rs::config::PolicyConfig::Random); - let app_context = common::create_test_context(router_config); + let app_context = common::create_test_context(router_config).await; let router = sglang_router_rs::routers::RouterFactory::create_router(&app_context).await; assert!( @@ -626,15 +595,10 @@ async fn test_router_factory_openai_mode() { /// Test that unsupported endpoints return proper error codes #[tokio::test] async fn test_unsupported_endpoints() { - let router = OpenAIRouter::new( - vec!["https://api.openai.com".to_string()], - None, - Arc::new(MemoryResponseStorage::new()), - Arc::new(MemoryConversationStorage::new()), - Arc::new(MemoryConversationItemStorage::new()), - ) - .await - .unwrap(); + let ctx = common::test_app::create_test_app_context().await; + let router = OpenAIRouter::new(vec!["https://api.openai.com".to_string()], &ctx) + .await + .unwrap(); let generate_request = GenerateRequest { text: Some("Hello world".to_string()), @@ -690,16 +654,9 @@ async fn test_openai_router_chat_completion_with_mock() { let mock_server = MockOpenAIServer::new().await; let base_url = mock_server.base_url(); + let ctx = common::test_app::create_test_app_context().await; // Create router pointing to mock server - let router = OpenAIRouter::new( - vec![base_url], - None, - Arc::new(MemoryResponseStorage::new()), - Arc::new(MemoryConversationStorage::new()), - Arc::new(MemoryConversationItemStorage::new()), - ) - .await - .unwrap(); + let router = OpenAIRouter::new(vec![base_url], &ctx).await.unwrap(); // Create a minimal chat completion request let mut chat_request = create_minimal_chat_request(); @@ -732,16 +689,9 @@ async fn test_openai_e2e_with_server() { let mock_server = MockOpenAIServer::new().await; let base_url = mock_server.base_url(); + let ctx = common::test_app::create_test_app_context().await; // Create router - let router = OpenAIRouter::new( - vec![base_url], - None, - Arc::new(MemoryResponseStorage::new()), - Arc::new(MemoryConversationStorage::new()), - Arc::new(MemoryConversationItemStorage::new()), - ) - .await - .unwrap(); + let router = OpenAIRouter::new(vec![base_url], &ctx).await.unwrap(); // Create Axum app with chat completions endpoint let app = Router::new().route( @@ -804,15 +754,8 @@ async fn test_openai_e2e_with_server() { async fn test_openai_router_chat_streaming_with_mock() { let mock_server = MockOpenAIServer::new().await; let base_url = mock_server.base_url(); - let router = OpenAIRouter::new( - vec![base_url], - None, - Arc::new(MemoryResponseStorage::new()), - Arc::new(MemoryConversationStorage::new()), - Arc::new(MemoryConversationItemStorage::new()), - ) - .await - .unwrap(); + let ctx = common::test_app::create_test_app_context().await; + let router = OpenAIRouter::new(vec![base_url], &ctx).await.unwrap(); // Build a streaming chat request let val = json!({ @@ -850,23 +793,10 @@ async fn test_openai_router_chat_streaming_with_mock() { /// Test circuit breaker functionality #[tokio::test] async fn test_openai_router_circuit_breaker() { - // Create router with circuit breaker config - let cb_config = sglang_router_rs::config::CircuitBreakerConfig { - failure_threshold: 2, - success_threshold: 1, - timeout_duration_secs: 1, - window_duration_secs: 10, - }; - - let router = OpenAIRouter::new( - vec!["http://invalid-url-that-will-fail".to_string()], - Some(cb_config), - Arc::new(MemoryResponseStorage::new()), - Arc::new(MemoryConversationStorage::new()), - Arc::new(MemoryConversationItemStorage::new()), - ) - .await - .unwrap(); + let ctx = common::test_app::create_test_app_context().await; + let router = OpenAIRouter::new(vec!["http://invalid-url-that-will-fail".to_string()], &ctx) + .await + .unwrap(); let chat_request = create_minimal_chat_request(); @@ -887,15 +817,10 @@ async fn test_openai_router_models_auth_forwarding() { // Start a mock server that requires Authorization let expected_auth = "Bearer test-token".to_string(); let mock_server = MockOpenAIServer::new_with_auth(Some(expected_auth.clone())).await; - let router = OpenAIRouter::new( - vec![mock_server.base_url()], - None, - Arc::new(MemoryResponseStorage::new()), - Arc::new(MemoryConversationStorage::new()), - Arc::new(MemoryConversationItemStorage::new()), - ) - .await - .unwrap(); + let ctx = common::test_app::create_test_app_context().await; + let router = OpenAIRouter::new(vec![mock_server.base_url()], &ctx) + .await + .unwrap(); // 1) Without auth header -> expect 200 with empty model list // (multi-endpoint aggregation silently skips failed endpoints) diff --git a/sgl-router/tests/test_pd_routing.rs b/sgl-router/tests/test_pd_routing.rs index a0c1dec0b..4a6ba7504 100644 --- a/sgl-router/tests/test_pd_routing.rs +++ b/sgl-router/tests/test_pd_routing.rs @@ -218,9 +218,10 @@ mod test_pd_routing { config.worker_startup_check_interval_secs, ))); - // Create empty OnceLock for worker job queue and workflow engine + // Create empty OnceLock for worker job queue, workflow engine, and mcp manager let worker_job_queue = Arc::new(OnceLock::new()); let workflow_engine = Arc::new(OnceLock::new()); + let mcp_manager = Arc::new(OnceLock::new()); Arc::new( AppContext::builder() @@ -238,6 +239,7 @@ mod test_pd_routing { .load_monitor(load_monitor) .worker_job_queue(worker_job_queue) .workflow_engine(workflow_engine) + .mcp_manager(mcp_manager) .build() .unwrap(), )