//! MCP 协议客户端 —— 与 MCP Server 通过 JSON-RPC over stdio 通信。 //! //! 当前 Phase 2 实现 stdio transport。`StreamableHttp` 枚举变体已预留, //! 但实际实现推迟到后续版本。 //! //! ## 协议版本 //! //! 实现遵循 MCP 协议版本 2025-03-26。 use std::collections::HashMap; use std::process::Stdio; use std::sync::Arc; use std::sync::atomic::{AtomicBool, Ordering}; use std::time::Duration; use async_trait::async_trait; use serde::{Deserialize, Serialize}; use serde_json::{Value, json}; use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader}; use tokio::process::{Child, ChildStdin, ChildStdout, Command}; use tokio::sync::{Mutex, oneshot}; use crate::llm::types::tool::ToolDef; use crate::tools::base::{BaseTool, ToolContext, ToolRef}; use crate::tools::error::ToolError; /// MCP 协议版本。 const MCP_VERSION: &str = "2025-03-26"; /// MCP 传输方式。 #[derive(Debug, Clone)] pub enum McpTransport { /// 通过子进程 stdin/stdout 通信。 Stdio { /// 启动命令(如 `"npx"`)。 command: String, /// 命令参数(如 `["-y", "@modelcontextprotocol/server-filesystem", "/tmp"]`)。 args: Vec, }, /// Streamable HTTP 传输(MCP 2025-03-26 引入,替代已废弃的 HTTP+SSE)。 /// /// 当前 Phase 2 预留枚举变体,调用方法会返回 `ToolError::McpError`。 StreamableHttp { /// MCP 端点 URL。 url: String, /// 可选的 HTTP 头(如 Authorization)。 headers: Option>, }, } /// JSON-RPC 请求。 #[derive(Debug, Serialize, Deserialize)] struct JsonRpcRequest { jsonrpc: &'static str, id: u64, method: String, #[serde(skip_serializing_if = "Option::is_none")] params: Option, } impl JsonRpcRequest { fn new(id: u64, method: impl Into, params: Option) -> Self { Self { jsonrpc: "2.0", id, method: method.into(), params, } } } /// JSON-RPC 响应。 #[derive(Debug, Serialize, Deserialize)] struct JsonRpcResponse { jsonrpc: String, id: u64, #[serde(default)] result: Option, #[serde(default)] error: Option, } #[derive(Debug, Serialize, Deserialize)] struct JsonRpcError { code: i32, message: String, #[serde(default)] data: Option, } /// MCP 子进程运行时状态。 struct ChildProcessState { child: Child, stdin: ChildStdin, pending: HashMap>>, next_id: u64, } impl ChildProcessState { fn next_id(&mut self) -> u64 { self.next_id += 1; self.next_id } } /// MCP Server 暴露的工具(缓存结构)。 #[derive(Debug, Clone)] struct McpTool { name: String, description: Option, input_schema: Value, } /// MCP 客户端 —— 与 MCP 服务器通信。 pub struct McpClient { transport: McpTransport, server_name: String, /// 已初始化的工具列表(缓存)。 tools: Vec, /// 是否已初始化。 initialized: AtomicBool, /// 超时时间(秒)。 timeout_secs: u64, /// 子进程运行时状态(`connect()` 后创建,`close()` 后取回)。 process: Option>>, } impl std::fmt::Debug for McpClient { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("McpClient") .field("server_name", &self.server_name) .field("initialized", &self.initialized.load(Ordering::SeqCst)) .field("tool_count", &self.tools.len()) .finish() } } impl McpClient { /// 创建一个 MCP 客户端。 pub fn new(server_name: impl Into, transport: McpTransport) -> Self { Self { transport, server_name: server_name.into(), tools: Vec::new(), initialized: AtomicBool::new(false), timeout_secs: 30, process: None, } } /// 设置超时时间(秒)。 pub fn with_timeout(mut self, secs: u64) -> Self { self.timeout_secs = secs; self } /// 检查是否已连接。 pub fn is_initialized(&self) -> bool { self.initialized.load(Ordering::SeqCst) } /// 连接并初始化(发送 initialize 请求)。 pub async fn connect(&mut self) -> Result<(), ToolError> { if self.is_initialized() { return Ok(()); } match &self.transport { McpTransport::Stdio { command, args } => { let mut cmd = Command::new(command); cmd.args(args) .stdin(Stdio::piped()) .stdout(Stdio::piped()) .stderr(Stdio::piped()); #[cfg(unix)] cmd.kill_on_drop(true); #[cfg(windows)] cmd.creation_flags(0x08000000); // CREATE_NO_WINDOW let mut child = cmd .spawn() .map_err(|e| ToolError::McpError(format!("启动 MCP 子进程失败: {e}")))?; let stdin = child .stdin .take() .ok_or_else(|| ToolError::McpError("无法获取子进程 stdin".into()))?; let stdout = child .stdout .take() .ok_or_else(|| ToolError::McpError("无法获取子进程 stdout".into()))?; // 启动 reader task 持续读取 stdout let pending: HashMap>> = HashMap::new(); let state = Arc::new(Mutex::new(ChildProcessState { child, stdin, pending, next_id: 0, })); // 启动后台 reader let pending_arc = Arc::clone(&state); tokio::spawn(async move { Self::read_loop(BufReader::new(stdout), pending_arc).await; }); self.process = Some(state); } McpTransport::StreamableHttp { .. } => { return Err(ToolError::McpError( "StreamableHttp transport 尚未实现".into(), )); } } // 发送 initialize 请求 let init_params = json!({ "protocolVersion": MCP_VERSION, "capabilities": {}, "clientInfo": { "name": "agcore", "version": env!("CARGO_PKG_VERSION") } }); let _response = self.send_request("initialize", Some(init_params)).await?; // 发送 initialized 通知(无 id) self.send_notification("notifications/initialized", Some(json!({}))) .await?; self.initialized.store(true, Ordering::SeqCst); Ok(()) } /// 列出服务器支持的工具(调用 `tools/list`)。 pub async fn list_tools(&mut self) -> Result, ToolError> { if !self.is_initialized() { return Err(ToolError::McpNotInitialized(self.server_name.clone())); } let response = self.send_request("tools/list", None).await?; let tools_value = response .get("tools") .ok_or_else(|| ToolError::McpError("tools/list 响应缺少 tools 字段".into()))?; let tools_arr = tools_value .as_array() .ok_or_else(|| ToolError::McpError("tools/list 响应 tools 字段不是数组".into()))?; self.tools.clear(); let mut defs = Vec::with_capacity(tools_arr.len()); for tool in tools_arr { let name = tool .get("name") .and_then(|v| v.as_str()) .ok_or_else(|| ToolError::McpError("工具缺少 name 字段".into()))? .to_string(); let description = tool .get("description") .and_then(|v| v.as_str()) .map(|s| s.to_string()); let input_schema = tool .get("inputSchema") .cloned() .unwrap_or_else(|| json!({"type": "object", "properties": {}})); self.tools.push(McpTool { name: name.clone(), description: description.clone(), input_schema: input_schema.clone(), }); defs.push(ToolDef { name, description, parameters: input_schema, }); } Ok(defs) } /// 调用一个工具(调用 `tools/call`)。 pub async fn call_tool(&self, name: &str, args: Value) -> Result { if !self.is_initialized() { return Err(ToolError::McpNotInitialized(self.server_name.clone())); } let params = json!({ "name": name, "arguments": args, }); let response = self.send_request("tools/call", Some(params)).await?; // 解析 content 字段 if let Some(content) = response.get("content").and_then(|c| c.as_array()) { // 收集所有 text 内容 let mut combined = String::new(); for item in content { if let Some(text) = item.get("text").and_then(|t| t.as_str()) { if !combined.is_empty() { combined.push('\n'); } combined.push_str(text); } } if !combined.is_empty() { return Ok(Value::String(combined)); } } // 如果没有 content 字段,尝试直接返回 is_error 标记 if let Some(true) = response.get("isError").and_then(|v| v.as_bool()) { return Err(ToolError::ExecutionFailed( name.to_string(), "MCP 工具返回 isError=true".into(), )); } // 回退:返回完整响应 Ok(response) } /// 关闭连接(终止子进程)。 pub async fn close(&mut self) -> Result<(), ToolError> { if !self.is_initialized() { return Ok(()); } // 尝试发送 shutdown(不强制要求响应) let _ = self.send_notification("shutdown", None).await; if let Some(state) = self.process.take() { let mut state = state.lock().await; // 优雅等待 5 秒 let graceful = tokio::time::timeout(Duration::from_secs(5), state.child.wait()).await; if graceful.is_err() { // 超时则强杀 let _ = state.child.kill().await; } } self.initialized.store(false, Ordering::SeqCst); self.tools.clear(); Ok(()) } /// 将 MCP 客户端转换为 `BaseTool` 适配器列表(用于注册到 `ToolRegistry`)。 /// /// **注意**:返回的适配器持有 `Arc`,但 `McpClient` 的可变性 /// (如 `list_tools` 刷新缓存)会通过 `Mutex` 处理。当前适配器仅缓存 /// 转换时的工具列表,不感知后续刷新。 pub fn into_tools(self) -> Vec { let mut tools = Vec::with_capacity(self.tools.len()); for mcp_tool in self.tools { let tool = McpToolAdapter { client: McpClientHandle::Empty, name: mcp_tool.name, description: mcp_tool.description.unwrap_or_default(), parameters: mcp_tool.input_schema, }; tools.push(Arc::new(tool) as ToolRef); } tools } async fn send_request(&self, method: &str, params: Option) -> Result { let state_arc = self .process .as_ref() .ok_or_else(|| ToolError::McpNotInitialized(self.server_name.clone()))? .clone(); let (id, request_json) = { let mut state = state_arc.lock().await; let id = state.next_id(); let req = JsonRpcRequest::new(id, method, params); let json = serde_json::to_string(&req) .map_err(|e| ToolError::McpError(format!("序列化请求失败: {e}")))?; (id, json) }; // 注册 oneshot 等待响应 let (tx, rx) = oneshot::channel(); { let mut state = state_arc.lock().await; state.pending.insert(id, tx); } // 写入请求 { let mut state = state_arc.lock().await; state .stdin .write_all(request_json.as_bytes()) .await .map_err(|e| ToolError::McpError(format!("写入请求失败: {e}")))?; state .stdin .write_all(b"\n") .await .map_err(|e| ToolError::McpError(format!("写入换行失败: {e}")))?; state .stdin .flush() .await .map_err(|e| ToolError::McpError(format!("flush stdin 失败: {e}")))?; } // 等待响应(带超时) tokio::time::timeout(Duration::from_secs(self.timeout_secs), rx) .await .map_err(|_| { // 超时:清理 pending let state_arc = state_arc.clone(); tokio::spawn(async move { let mut state = state_arc.lock().await; state.pending.remove(&id); }); ToolError::McpTimeout(method.to_string()) })? .map_err(|_| ToolError::McpError("response channel 关闭".into()))? } async fn send_notification( &self, method: &str, params: Option, ) -> Result<(), ToolError> { let state_arc = self .process .as_ref() .ok_or_else(|| ToolError::McpNotInitialized(self.server_name.clone()))? .clone(); let notification = json!({ "jsonrpc": "2.0", "method": method, "params": params, }); let json = serde_json::to_string(¬ification) .map_err(|e| ToolError::McpError(format!("序列化通知失败: {e}")))?; let mut state = state_arc.lock().await; state .stdin .write_all(json.as_bytes()) .await .map_err(|e| ToolError::McpError(format!("写入通知失败: {e}")))?; state .stdin .write_all(b"\n") .await .map_err(|e| ToolError::McpError(format!("写入换行失败: {e}")))?; state .stdin .flush() .await .map_err(|e| ToolError::McpError(format!("flush stdin 失败: {e}")))?; Ok(()) } /// 持续读取 stdout,将响应分发到对应的 oneshot sender。 async fn read_loop(mut reader: BufReader, state: Arc>) { let mut line = String::new(); loop { line.clear(); match reader.read_line(&mut line).await { Ok(0) => { // EOF:通知所有 pending 失败 let mut state = state.lock().await; for (_, tx) in state.pending.drain() { let _ = tx.send(Err(ToolError::McpError("子进程退出".into()))); } break; } Ok(_) => { let trimmed = line.trim(); if trimmed.is_empty() { continue; } // 尝试解析为 JSON-RPC 响应 let parsed: Result = serde_json::from_str(trimmed); if let Ok(response) = parsed { let value = if let Some(err) = response.error { Err(ToolError::McpError(format!( "[{}] {}", err.code, err.message ))) } else { Ok(response.result.unwrap_or(Value::Null)) }; let mut state = state.lock().await; if let Some(tx) = state.pending.remove(&response.id) { let _ = tx.send(value); } } // 非响应消息(通知、request from server)忽略 } Err(e) => { tracing::warn!("MCP read_loop error: {e}"); let mut state = state.lock().await; for (_, tx) in state.pending.drain() { let _ = tx.send(Err(ToolError::McpError(format!("读取失败: {e}")))); } break; } } } } } /// MCP 工具适配器 —— 将 MCP 工具包装为 `BaseTool`。 struct McpToolAdapter { /// 持有 client 的弱引用。实际生产中应使用 `Arc`, /// 但当前 Phase 2 实现不直接持有可变的 `McpClient`。 /// 标记为 unused 但保留字段以展示扩展路径。 #[allow(dead_code)] client: McpClientHandle, name: String, description: String, parameters: Value, } #[allow(dead_code)] enum McpClientHandle { Empty, // Future: Shared(Arc), } #[async_trait] impl BaseTool for McpToolAdapter { fn name(&self) -> &str { &self.name } fn description(&self) -> &str { &self.description } fn parameters(&self) -> Value { self.parameters.clone() } async fn execute(&self, _args: Value, _ctx: &ToolContext<'_>) -> Result { // 当前 Phase 2 实现的简化:McpToolAdapter 不持有活跃 MCP 连接。 // 实际生产中应持有 Arc 并通过 mcp.call_tool() 执行。 // 这里返回错误,提示需要通过其他方式调用 MCP 工具。 Err(ToolError::McpError(format!( "MCP 工具 '{}' 需要活跃的 McpClient 引用(当前 Phase 2 简化实现)", self.name ))) } } #[cfg(test)] mod tests { use super::*; #[test] fn test_transport_debug() { let transport = McpTransport::Stdio { command: "echo".to_string(), args: vec!["hello".to_string()], }; let formatted = format!("{transport:?}"); assert!(formatted.contains("echo")); } #[test] fn test_client_creation() { let transport = McpTransport::Stdio { command: "test".to_string(), args: vec![], }; let client = McpClient::new("test-server", transport).with_timeout(60); assert_eq!(client.server_name, "test-server"); assert_eq!(client.timeout_secs, 60); assert!(!client.is_initialized()); } #[test] fn test_jsonrpc_request_serialize() { let req = JsonRpcRequest::new(42, "test", Some(json!({"a": 1}))); let s = serde_json::to_string(&req).unwrap(); assert!(s.contains("\"jsonrpc\":\"2.0\"")); assert!(s.contains("\"id\":42")); assert!(s.contains("\"method\":\"test\"")); } #[test] fn test_jsonrpc_response_parse_ok() { let s = r#"{"jsonrpc":"2.0","id":1,"result":{"foo":"bar"}}"#; let resp: JsonRpcResponse = serde_json::from_str(s).unwrap(); assert_eq!(resp.id, 1); assert!(resp.result.is_some()); assert!(resp.error.is_none()); } #[test] fn test_jsonrpc_response_parse_error() { let s = r#"{"jsonrpc":"2.0","id":1,"error":{"code":-32601,"message":"Method not found"}}"#; let resp: JsonRpcResponse = serde_json::from_str(s).unwrap(); assert_eq!(resp.id, 1); assert!(resp.result.is_none()); let err = resp.error.unwrap(); assert_eq!(err.code, -32601); } #[tokio::test] async fn test_streamable_http_not_implemented() { let mut client = McpClient::new( "http-server", McpTransport::StreamableHttp { url: "https://example.com/mcp".to_string(), headers: None, }, ); let result = client.connect().await; // 当前 Phase 2 返回未实现错误 assert!(result.is_err()); assert!(matches!(result, Err(ToolError::McpError(_)))); } }