Files
agcore/src/tools/mcp.rs
T
徐涛 4cf5918b9c refactor(llm): 移除 ToolDefinition 别名与遗留 deprecation 抑制点
完成 ToolDef IR 切换的最后清理:移除 ToolDefinition 别名与
OpenaiToolDefinition 的公共 re-export;各调用点(cycle/registry/mcp/agent)
直接使用 ToolDef 并清理对应的 #[allow(deprecated)] 抑制点;新增
roundtrip 测试验证 ToolDef 序列化兼容性。
2026-07-05 10:32:35 +08:00

624 lines
20 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! 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<String>,
},
/// Streamable HTTP 传输(MCP 2025-03-26 引入,替代已废弃的 HTTP+SSE)。
///
/// 当前 Phase 2 预留枚举变体,调用方法会返回 `ToolError::McpError`。
StreamableHttp {
/// MCP 端点 URL。
url: String,
/// 可选的 HTTP 头(如 Authorization)。
headers: Option<Vec<(String, String)>>,
},
}
/// JSON-RPC 请求。
#[derive(Debug, Serialize, Deserialize)]
struct JsonRpcRequest {
jsonrpc: &'static str,
id: u64,
method: String,
#[serde(skip_serializing_if = "Option::is_none")]
params: Option<Value>,
}
impl JsonRpcRequest {
fn new(id: u64, method: impl Into<String>, params: Option<Value>) -> 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<Value>,
#[serde(default)]
error: Option<JsonRpcError>,
}
#[derive(Debug, Serialize, Deserialize)]
struct JsonRpcError {
code: i32,
message: String,
#[serde(default)]
data: Option<Value>,
}
/// MCP 子进程运行时状态。
struct ChildProcessState {
child: Child,
stdin: ChildStdin,
pending: HashMap<u64, oneshot::Sender<Result<Value, ToolError>>>,
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<String>,
input_schema: Value,
}
/// MCP 客户端 —— 与 MCP 服务器通信。
pub struct McpClient {
transport: McpTransport,
server_name: String,
/// 已初始化的工具列表(缓存)。
tools: Vec<McpTool>,
/// 是否已初始化。
initialized: AtomicBool,
/// 超时时间(秒)。
timeout_secs: u64,
/// 子进程运行时状态(`connect()` 后创建,`close()` 后取回)。
process: Option<Arc<Mutex<ChildProcessState>>>,
}
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<String>, 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<u64, oneshot::Sender<Result<Value, ToolError>>> =
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<Vec<ToolDef>, 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<Value, ToolError> {
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>`,但 `McpClient` 的可变性
/// (如 `list_tools` 刷新缓存)会通过 `Mutex` 处理。当前适配器仅缓存
/// 转换时的工具列表,不感知后续刷新。
pub fn into_tools(self) -> Vec<ToolRef> {
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<Value>) -> Result<Value, ToolError> {
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<Value>,
) -> 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(&notification)
.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<ChildStdout>, state: Arc<Mutex<ChildProcessState>>) {
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<JsonRpcResponse, _> = 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<McpClient>`
/// 但当前 Phase 2 实现不直接持有可变的 `McpClient`。
/// 标记为 unused 但保留字段以展示扩展路径。
#[allow(dead_code)]
client: McpClientHandle,
name: String,
description: String,
parameters: Value,
}
#[allow(dead_code)]
enum McpClientHandle {
Empty,
// Future: Shared(Arc<McpClient>),
}
#[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<Value, ToolError> {
// 当前 Phase 2 实现的简化:McpToolAdapter 不持有活跃 MCP 连接。
// 实际生产中应持有 Arc<McpClient> 并通过 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(_))));
}
}