use crate::types::*;
use crate::mcp_tool::{DiscoveredMcpTool, McpToolFactory};
use crate::types::McpServerConfig;
use crate::tool_registry::ToolRegistry;
use crate::workspace_context::WorkspaceContext;
use crate::prompt_registry::PromptRegistry;
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
use anyhow::Result;
use serde_json;
use tracing::{info, warn, error};
use std::time::Duration;
use tokio::time::timeout;
use rmcp::{
transport::TokioChildProcess,
ServiceExt,
};
use tokio::process::Command;
pub const MCP_DEFAULT_TIMEOUT_MSEC: u64 = 10 * 1000;
#[derive(Debug, Clone, PartialEq)]
pub enum McpServerStatus {
Disconnected,
Connecting,
Connected,
}
#[derive(Debug, Clone, PartialEq)]
pub enum McpDiscoveryState {
NotStarted,
InProgress,
Completed,
}
#[derive(Debug, Clone)]
pub struct DiscoveredMcpPrompt {
pub name: String,
pub description: Option<String>,
pub arguments: Option<Vec<serde_json::Value>>,
pub server_name: String,
}
#[derive(Clone)]
pub struct McpClient {
name: String,
version: String,
server_name: String,
status: McpServerStatus,
transport: Option<Arc<dyn std::any::Any + Send + Sync>>,
}
impl McpClient {
pub fn new(name: String, version: String, server_name: String) -> Self {
Self {
name,
version,
server_name,
status: McpServerStatus::Disconnected,
transport: None,
}
}
pub async fn connect(&mut self, server_config: &McpServerConfig) -> Result<()> {
info!("连接到MCP服务器: {}", self.server_name);
self.status = McpServerStatus::Connecting;
let transport = self.create_transport(server_config).await?;
self.transport = Some(Arc::new(transport));
self.status = McpServerStatus::Connected;
info!("MCP客户端连接成功: {}", self.server_name);
Ok(())
}
async fn create_transport(&self, server_config: &McpServerConfig) -> Result<Box<dyn std::any::Any + Send + Sync>> {
if server_config.command.is_some() {
let command = server_config.command.as_ref()
.ok_or_else(|| anyhow::anyhow!("child_process传输需要command配置"))?;
let args = server_config.args.clone().unwrap_or_default();
let env = server_config.env.clone().unwrap_or_default();
let cwd = server_config.cwd.clone();
info!("创建child_process传输: {} {:?}", command, args);
info!("环境变量: {:?}", env);
info!("工作目录: {:?}", cwd);
let mut cmd = if cfg!(target_os = "windows") && command == "npx" {
let mut cmd = Command::new("cmd");
cmd.args(&["/c", "npx"]);
cmd.args(&args);
cmd
} else {
let mut cmd = Command::new(command);
cmd.args(&args);
cmd
};
for (key, value) in env {
cmd.env(key, value);
}
if let Some(work_dir) = cwd {
cmd.current_dir(work_dir);
}
info!("正在创建child_process传输...");
let transport = tokio::task::spawn_blocking(move || {
TokioChildProcess::new(cmd)
}).await
.map_err(|e| anyhow::anyhow!("创建child_process传输任务失败: {}", e))?
.map_err(|e| {
error!("创建child_process传输失败 - 命令: {} 参数: {:?} 错误: {}", command, args, e);
anyhow::anyhow!("创建child_process传输失败: {}", e)
})?;
info!("child_process传输创建成功");
info!("成功创建child_process传输");
Ok(Box::new(transport))
} else if server_config.url.is_some() {
let url = server_config.url.as_ref().unwrap();
Err(anyhow::anyhow!("HTTP/SSE传输暂时不支持,URL: {}", url))
} else {
Err(anyhow::anyhow!("无效的MCP服务器配置:需要command或url"))
}
}
pub async fn list_tools(&self) -> Result<Vec<MockTool>> {
if self.transport.is_none() {
warn!("MCP客户端 {} 未连接,无法发现工具", self.server_name);
return Ok(vec![]);
}
info!("开始从MCP服务器 {} 发现工具", self.server_name);
match self.discover_real_tools().await {
Ok(real_tools) => {
info!("从MCP服务器 {} 发现 {} 个真实工具", self.server_name, real_tools.len());
Ok(real_tools)
}
Err(e) => {
warn!("从MCP服务器 {} 发现真实工具失败: {},返回空列表", self.server_name, e);
Ok(vec![])
}
}
}
async fn discover_real_tools(&self) -> Result<Vec<MockTool>> {
let transport = self.transport.as_ref()
.ok_or_else(|| anyhow::anyhow!("传输层未初始化"))?;
info!("尝试从MCP服务器 {} 发现真实工具", self.server_name);
match self.try_real_mcp_discovery(transport).await {
Ok(real_tools) => {
if !real_tools.is_empty() {
info!("从MCP服务器 {} 发现 {} 个真实工具", self.server_name, real_tools.len());
return Ok(real_tools);
}
}
Err(e) => {
warn!("真实MCP工具发现失败: {},回退到模拟工具", e);
}
}
warn!("从MCP服务器 {} 未发现任何工具", self.server_name);
Ok(vec![])
}
async fn try_real_mcp_discovery(&self, transport: &Arc<dyn std::any::Any + Send + Sync>) -> Result<Vec<MockTool>> {
info!("开始真正的MCP工具发现流程");
let tools = self.discover_tools_via_process().await?;
if !tools.is_empty() {
info!("通过进程通信发现 {} 个工具", tools.len());
return Ok(tools);
}
warn!("真正的MCP工具发现失败,返回空列表");
Ok(vec![])
}
async fn discover_tools_via_process(&self) -> Result<Vec<MockTool>> {
use std::process::{Command, Stdio};
use std::io::{Write, BufRead, BufReader};
use serde_json::{json, Value};
use tokio::time::{sleep, Duration};
let (command, args) = match self.server_name.as_str() {
"filesystem" => ("npx", vec!["@modelcontextprotocol/server-filesystem", "."]),
"memory" => ("npx", vec!["-y", "@modelcontextprotocol/server-memory"]),
"payment" => ("python", vec!["-m", "blockchain_payment_mcp.server"]),
_ => {
warn!("未知的服务器类型: {}", self.server_name);
return Ok(vec![]);
}
};
info!("启动MCP服务器进程: {} {:?}", command, args);
let mut child = if cfg!(target_os = "windows") && command == "npx" {
Command::new("cmd")
.args(&["/c", "npx"])
.args(&args)
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.spawn()?
} else {
Command::new(command)
.args(&args)
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.spawn()?
};
let stdin = child.stdin.as_mut().unwrap();
let stdout = child.stdout.take().unwrap();
sleep(Duration::from_millis(2000)).await;
let init_request = json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2024-11-05",
"capabilities": {
"tools": {}
},
"clientInfo": {
"name": "alou-mcp-client",
"version": "0.1.0"
}
}
});
writeln!(stdin, "{}", init_request)?;
stdin.flush()?;
sleep(Duration::from_millis(1000)).await;
let tools_request = json!({
"jsonrpc": "2.0",
"id": 2,
"method": "tools/list"
});
writeln!(stdin, "{}", tools_request)?;
stdin.flush()?;
let reader = BufReader::new(stdout);
let mut tools = Vec::new();
for line in reader.lines() {
let line = line?;
if let Ok(json) = serde_json::from_str::<Value>(&line) {
if let Some(result) = json.get("result") {
if let Some(tools_array) = result.get("tools") {
if let Some(tools_list) = tools_array.as_array() {
for tool_json in tools_list {
if let Some(tool) = self.parse_mcp_tool(tool_json) {
tools.push(tool);
}
}
break; }
}
}
}
}
let _ = child.kill();
Ok(tools)
}
fn parse_mcp_tool(&self, tool_json: &serde_json::Value) -> Option<MockTool> {
let name = tool_json.get("name")?.as_str()?.to_string();
let description = tool_json.get("description").and_then(|d| d.as_str()).map(|s| s.to_string());
let input_schema = tool_json.get("inputSchema").cloned();
Some(MockTool {
name,
description,
input_schema,
})
}
pub async fn close(&mut self) -> Result<()> {
info!("关闭MCP客户端连接: {}", self.server_name);
self.transport = None;
self.status = McpServerStatus::Disconnected;
Ok(())
}
pub fn get_status(&self) -> &McpServerStatus {
&self.status
}
pub fn get_server_name(&self) -> &str {
&self.server_name
}
pub async fn list_prompts(&self) -> Result<Vec<MockPrompt>> {
Ok(vec![])
}
pub async fn get_prompt(&self, name: &str, arguments: HashMap<String, serde_json::Value>) -> Result<MockPromptResult> {
Ok(MockPromptResult {
description: Some(format!("模拟提示: {}", name)),
messages: vec![],
})
}
pub async fn call_tool(&self, name: &str, arguments: HashMap<String, serde_json::Value>) -> Result<MockToolResult> {
Ok(MockToolResult {
content: format!("模拟调用工具: {} 参数: {:?}", name, arguments),
is_error: false,
})
}
}
#[derive(Debug, Clone)]
pub struct MockTool {
pub name: String,
pub description: Option<String>,
pub input_schema: Option<serde_json::Value>,
}
#[derive(Debug, Clone)]
pub struct MockToolResult {
pub content: String,
pub is_error: bool,
}
#[derive(Debug, Clone)]
pub struct MockPrompt {
pub name: String,
pub description: Option<String>,
pub arguments: Option<Vec<serde_json::Value>>,
}
#[derive(Debug, Clone)]
pub struct MockPromptResult {
pub description: Option<String>,
pub messages: Vec<serde_json::Value>,
}
pub struct McpClientManager {
clients: Arc<RwLock<HashMap<String, Arc<RwLock<McpClient>>>>>,
server_statuses: Arc<RwLock<HashMap<String, McpServerStatus>>>,
discovery_state: Arc<RwLock<McpDiscoveryState>>,
}
impl McpClientManager {
pub fn new() -> Self {
Self {
clients: Arc::new(RwLock::new(HashMap::new())),
server_statuses: Arc::new(RwLock::new(HashMap::new())),
discovery_state: Arc::new(RwLock::new(McpDiscoveryState::NotStarted)),
}
}
pub async fn discover_mcp_tools(
&self,
mcp_servers: HashMap<String, McpServerConfig>,
tool_registry: Arc<ToolRegistry>,
prompt_registry: Arc<PromptRegistry>,
debug_mode: bool,
workspace_context: Arc<dyn WorkspaceContext + Send + Sync>,
) -> Result<()> {
let mut discovery_state = self.discovery_state.write().await;
*discovery_state = McpDiscoveryState::InProgress;
drop(discovery_state);
let discovery_promises = mcp_servers.into_iter().map(|(server_name, server_config)| {
self.connect_and_discover(
server_name,
server_config,
tool_registry.clone(),
prompt_registry.clone(),
debug_mode,
workspace_context.clone(),
)
});
let results = futures::future::join_all(discovery_promises).await;
let mut success_count = 0;
for result in results {
if result.is_ok() {
success_count += 1;
}
}
let mut discovery_state = self.discovery_state.write().await;
*discovery_state = McpDiscoveryState::Completed;
info!("MCP工具发现完成,成功连接 {} 个服务器", success_count);
Ok(())
}
async fn connect_and_discover(
&self,
server_name: String,
server_config: McpServerConfig,
tool_registry: Arc<ToolRegistry>,
prompt_registry: Arc<PromptRegistry>,
debug_mode: bool,
_workspace_context: Arc<dyn WorkspaceContext + Send + Sync>,
) -> Result<()> {
info!("连接MCP服务器: {}", server_name);
{
let mut statuses = self.server_statuses.write().await;
statuses.insert(server_name.clone(), McpServerStatus::Connecting);
}
let mut mcp_client = McpClient::new(
"alou-mcp-client".to_string(),
"0.1.0".to_string(),
server_name.clone(),
);
let timeout_duration = Duration::from_millis(server_config.timeout.unwrap_or(MCP_DEFAULT_TIMEOUT_MSEC));
let connect_result = timeout(timeout_duration, mcp_client.connect(&server_config)).await;
match connect_result {
Ok(Ok(())) => {
info!("已连接到MCP服务器: {}", server_name);
}
Ok(Err(e)) => {
error!("连接MCP服务器 '{}' 失败: {}", server_name, e);
error!("服务器配置: command={:?}, args={:?}", server_config.command, server_config.args);
let mut statuses = self.server_statuses.write().await;
statuses.insert(server_name, McpServerStatus::Disconnected);
return Err(e);
}
Err(_) => {
error!("连接MCP服务器 '{}' 超时 ({}秒)", server_name, timeout_duration.as_secs());
error!("服务器配置: command={:?}, args={:?}", server_config.command, server_config.args);
let mut statuses = self.server_statuses.write().await;
statuses.insert(server_name, McpServerStatus::Disconnected);
return Err(anyhow::anyhow!("连接超时"));
}
}
let prompts: Vec<DiscoveredMcpPrompt> = Vec::new(); let tools = self.discover_tools(&server_name, &server_config, &mcp_client).await?;
info!("发现 {} 个工具和 {} 个提示", tools.len(), prompts.len());
if tools.is_empty() && prompts.is_empty() {
return Err(anyhow::anyhow!("服务器上未发现任何工具或提示"));
}
{
let mut statuses = self.server_statuses.write().await;
statuses.insert(server_name.clone(), McpServerStatus::Connected);
}
for tool in tools {
info!("注册工具: {}", tool.name());
tool_registry.register_mcp_tool(Box::new(tool)).await?;
}
{
let mut clients = self.clients.write().await;
clients.insert(server_name.clone(), Arc::new(RwLock::new(mcp_client)));
}
info!("完成连接到MCP服务器: {}", server_name);
Ok(())
}
async fn discover_tools(
&self,
server_name: &str,
server_config: &McpServerConfig,
mcp_client: &McpClient,
) -> Result<Vec<DiscoveredMcpTool>> {
let tools = mcp_client.list_tools().await?;
let mut discovered_tools = Vec::new();
for tool in tools {
if !self.is_enabled(&tool, server_name, server_config) {
continue;
}
if !self.has_valid_types(&tool) {
warn!(
"跳过工具 '{}' 从MCP服务器 '{}',因为其参数模式中缺少类型。请向MCP服务器的所有者提交问题。",
tool.name, server_name
);
continue;
}
let discovered_tool = McpToolFactory::create_discovered_tool(
Arc::new(RwLock::new(mcp_client.clone())),
server_name.to_string(),
tool.name.clone(),
tool.description.unwrap_or_default(),
tool.input_schema.unwrap_or_else(|| serde_json::json!({"type": "object", "properties": {}})),
server_config.timeout,
server_config.trust,
);
discovered_tools.push(discovered_tool);
}
Ok(discovered_tools)
}
async fn discover_prompts(
&self,
server_name: &str,
mcp_client: &McpClient,
prompt_registry: &mut PromptRegistry,
) -> Result<Vec<DiscoveredMcpPrompt>> {
let prompts = mcp_client.list_prompts().await?;
let mut discovered_prompts = Vec::new();
for prompt in prompts {
let discovered_prompt = DiscoveredMcpPrompt {
name: prompt.name.clone(),
description: prompt.description,
arguments: prompt.arguments,
server_name: server_name.to_string(),
};
let registry_prompt = crate::prompt_registry::DiscoveredMcpPrompt {
name: discovered_prompt.name.clone(),
description: discovered_prompt.description.clone(),
arguments: discovered_prompt.arguments.clone(),
server_name: discovered_prompt.server_name.clone(),
invoke: Box::new(|_args| {
Ok(crate::prompt_registry::GetPromptResult {
description: Some("Mock prompt result".to_string()),
messages: vec![],
})
}),
};
prompt_registry.register_prompt(registry_prompt);
discovered_prompts.push(discovered_prompt);
}
Ok(discovered_prompts)
}
fn is_enabled(&self, tool: &MockTool, server_name: &str, server_config: &McpServerConfig) -> bool {
let include_tools = &server_config.include_tools;
let exclude_tools = &server_config.exclude_tools;
if let Some(exclude_list) = exclude_tools {
if exclude_list.contains(&tool.name) {
return false;
}
}
match include_tools {
Some(include_list) => {
include_list.iter().any(|tool_name| {
tool_name == &tool.name || tool_name.starts_with(&format!("{}(", tool.name))
})
}
None => true, }
}
fn has_valid_types(&self, tool: &MockTool) -> bool {
if let Some(schema) = &tool.input_schema {
self.validate_schema_types(schema)
} else {
false
}
}
fn validate_schema_types(&self, schema: &serde_json::Value) -> bool {
if let Some(obj) = schema.as_object() {
if obj.contains_key("type") {
return true;
}
for keyword in &["anyOf", "allOf", "oneOf"] {
if let Some(array) = obj.get(&keyword.to_string()).and_then(|v| v.as_array()) {
return array.iter().all(|sub_schema| self.validate_schema_types(sub_schema));
}
}
}
false
}
pub async fn get_server_status(&self, server_name: &str) -> McpServerStatus {
let statuses = self.server_statuses.read().await;
statuses.get(server_name).cloned().unwrap_or(McpServerStatus::Disconnected)
}
pub async fn get_all_server_statuses(&self) -> HashMap<String, McpServerStatus> {
let statuses = self.server_statuses.read().await;
statuses.clone()
}
pub async fn get_discovery_state(&self) -> McpDiscoveryState {
let state = self.discovery_state.read().await;
state.clone()
}
pub async fn get_client(&self, server_name: &str) -> Option<Arc<RwLock<McpClient>>> {
let clients = self.clients.read().await;
clients.get(server_name).cloned()
}
pub async fn close_all(&self) -> Result<()> {
let mut clients = self.clients.write().await;
for (server_name, client) in clients.iter_mut() {
if let Ok(mut client_guard) = client.try_write() {
if let Err(e) = client_guard.close().await {
error!("关闭MCP客户端 '{}' 失败: {}", server_name, e);
}
}
}
clients.clear();
let mut statuses = self.server_statuses.write().await;
statuses.clear();
let mut discovery_state = self.discovery_state.write().await;
*discovery_state = McpDiscoveryState::NotStarted;
Ok(())
}
}
impl Default for McpClientManager {
fn default() -> Self {
Self::new()
}
}
fn get_error_message(error: &dyn std::error::Error) -> String {
error.to_string()
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_mcp_client_creation() {
let client = McpClient::new("test-client".to_string(), "1.0.0".to_string(), "test-server".to_string());
assert_eq!(client.name, "test-client");
assert_eq!(client.version, "1.0.0");
assert_eq!(client.server_name, "test-server");
}
#[tokio::test]
async fn test_mcp_client_manager_creation() {
let manager = McpClientManager::new();
let state = manager.get_discovery_state().await;
assert_eq!(state, McpDiscoveryState::NotStarted);
}
#[test]
fn test_validate_schema_types() {
let manager = McpClientManager::new();
let valid_schema = serde_json::json!({
"type": "object",
"properties": {
"name": {"type": "string"}
}
});
assert!(manager.validate_schema_types(&valid_schema));
let invalid_schema = serde_json::json!({
"properties": {
"name": {"type": "string"}
}
});
assert!(!manager.validate_schema_types(&invalid_schema));
}
}