feat(tools): 添加工具系统框架与 MCP 协议客户端

This commit is contained in:
徐涛
2026-06-07 10:57:15 +08:00
parent e598f6d3ee
commit b6e7acfb0f
9 changed files with 2034 additions and 1 deletions
+640
View File
@@ -0,0 +1,640 @@
//! 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::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::Duration;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use serde_json::{json, Value};
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::process::{Child, ChildStdin, ChildStdout, Command};
use tokio::sync::{oneshot, Mutex};
use crate::llm::types::ToolDefinition;
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<ToolDefinition>, 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(ToolDefinition {
name,
description,
parameters: input_schema,
strict: None,
});
}
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(_))));
}
}