use crate::types::*;
use crate::tools::BaseDeclarativeTool;
use crate::types::Tool;
use crate::mcp_client::McpClient;
use async_trait::async_trait;
use std::collections::HashMap;
use std::sync::Arc;
use serde_json;
use anyhow::Result;
use tokio_util::sync::CancellationToken;
use tokio::sync::RwLock;
pub struct DiscoveredMcpToolInvocation {
mcp_client: Arc<RwLock<McpClient>>,
server_name: String,
server_tool_name: String,
display_name: String,
timeout: Option<u64>,
trust: Option<bool>,
params: HashMap<String, serde_json::Value>,
allowlist: Arc<tokio::sync::RwLock<HashMap<String, bool>>>,
}
impl DiscoveredMcpToolInvocation {
pub fn new(
mcp_client: Arc<RwLock<McpClient>>,
server_name: String,
server_tool_name: String,
display_name: String,
timeout: Option<u64>,
trust: Option<bool>,
params: HashMap<String, serde_json::Value>,
) -> Self {
Self {
mcp_client,
server_name,
server_tool_name,
display_name,
timeout,
trust,
params,
allowlist: Arc::new(tokio::sync::RwLock::new(HashMap::new())),
}
}
}
#[async_trait]
impl ToolInvocation for DiscoveredMcpToolInvocation {
fn name(&self) -> &str {
&self.server_tool_name
}
fn params(&self) -> &HashMap<String, serde_json::Value> {
&self.params
}
async fn should_confirm_execute(&self, _abort_signal: &CancellationToken) -> Result<Option<ToolCallConfirmationDetails>, Box<dyn std::error::Error + Send + Sync>> {
let server_allowlist_key = &self.server_name;
let tool_allowlist_key = format!("{}.{}", self.server_name, self.server_tool_name);
if self.trust.unwrap_or(false) {
return Ok(None); }
let allowlist = self.allowlist.read().await;
if allowlist.contains_key(server_allowlist_key) || allowlist.contains_key(&tool_allowlist_key) {
return Ok(None); }
let confirmation_details = ToolCallConfirmationDetails {
tool_name: self.server_tool_name.clone(),
params: self.params.clone(),
};
Ok(Some(confirmation_details))
}
async fn execute(&self) -> Result<ToolResultContent, Box<dyn std::error::Error + Send + Sync>> {
let client = self.mcp_client.read().await;
let result = client.call_tool(&self.server_tool_name, self.params.clone()).await?;
let content = if result.is_error {
format!("工具调用错误: {}", result.content)
} else {
result.content
};
Ok(ToolResultContent {
content,
mime_type: None,
llm_content: None,
return_display: None,
})
}
fn get_description(&self) -> &str {
&self.display_name
}
}
pub struct DiscoveredMcpTool {
mcp_client: Arc<RwLock<McpClient>>,
server_name: String,
server_tool_name: String,
description: String,
parameter_schema: serde_json::Value,
timeout: Option<u64>,
trust: Option<bool>,
base_tool: BaseDeclarativeTool,
}
impl DiscoveredMcpTool {
pub fn new(
mcp_client: Arc<RwLock<McpClient>>,
server_name: String,
server_tool_name: String,
description: String,
parameter_schema: serde_json::Value,
timeout: Option<u64>,
trust: Option<bool>,
name_override: Option<String>,
) -> Self {
let name = name_override.unwrap_or_else(|| generate_valid_name(&server_tool_name));
let display_name = format!("{} ({} MCP Server)", server_tool_name, server_name);
let base_tool = BaseDeclarativeTool::new(
name,
display_name,
description.clone(),
Kind::Other,
parameter_schema.clone(),
true, false, );
Self {
mcp_client,
server_name,
server_tool_name,
description,
parameter_schema,
timeout,
trust,
base_tool,
}
}
pub fn as_fully_qualified_tool(&self) -> Self {
let qualified_name = format!("{}__{}", self.server_name, self.server_tool_name);
Self::new(
self.mcp_client.clone(),
self.server_name.clone(),
self.server_tool_name.clone(),
self.description.clone(),
self.parameter_schema.clone(),
self.timeout,
self.trust,
Some(qualified_name),
)
}
pub fn server_name(&self) -> &str {
&self.server_name
}
pub fn server_tool_name(&self) -> &str {
&self.server_tool_name
}
}
#[async_trait]
impl Tool for DiscoveredMcpTool {
fn name(&self) -> &str {
&self.base_tool.name
}
fn description(&self) -> &str {
&self.base_tool.description
}
fn display_name(&self) -> &str {
&self.base_tool.display_name
}
fn kind(&self) -> Kind {
self.base_tool.kind.clone()
}
fn parameter_schema(&self) -> &serde_json::Value {
&self.base_tool.parameter_schema
}
fn is_output_markdown(&self) -> bool {
self.base_tool.is_output_markdown
}
fn can_update_output(&self) -> bool {
self.base_tool.can_update_output
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
async fn should_confirm_execute(&self, abort_signal: &CancellationToken) -> Result<Option<ToolCallConfirmationDetails>, Box<dyn std::error::Error + Send + Sync>> {
let invocation = self.create_invocation(HashMap::new());
invocation.should_confirm_execute(abort_signal).await
}
async fn execute(&self, params: HashMap<String, serde_json::Value>) -> Result<ToolResultContent, Box<dyn std::error::Error + Send + Sync>> {
let invocation = self.create_invocation(params);
invocation.execute().await
}
async fn build_and_execute(&self, params: HashMap<String, serde_json::Value>, abort_signal: Option<&CancellationToken>) -> Result<ToolResultContent, Box<dyn std::error::Error + Send + Sync>> {
let invocation = self.create_invocation(params);
if let Some(signal) = abort_signal {
invocation.should_confirm_execute(signal).await?;
}
invocation.execute().await
}
}
impl DiscoveredMcpTool {
fn create_invocation(&self, params: HashMap<String, serde_json::Value>) -> DiscoveredMcpToolInvocation {
DiscoveredMcpToolInvocation::new(
self.mcp_client.clone(),
self.server_name.clone(),
self.server_tool_name.clone(),
self.display_name().to_string(),
self.timeout,
self.trust,
params,
)
}
}
pub fn generate_valid_name(name: &str) -> String {
let mut valid_toolname = name
.chars()
.map(|c| {
if c.is_alphanumeric() || c == '_' || c == '-' || c == '.' {
c
} else {
'_'
}
})
.collect::<String>();
if valid_toolname.len() > 63 {
valid_toolname = format!(
"{}___{}",
&valid_toolname[..28],
&valid_toolname[valid_toolname.len() - 32..]
);
}
valid_toolname
}
pub struct MockCallableTool;
impl MockCallableTool {
pub fn create_mock_client(server_name: String) -> Arc<RwLock<McpClient>> {
let client = McpClient::new(
"mock-client".to_string(),
"1.0.0".to_string(),
server_name,
);
Arc::new(RwLock::new(client))
}
}
pub struct McpToolFactory;
impl McpToolFactory {
pub fn create_discovered_tool(
mcp_client: Arc<RwLock<McpClient>>,
server_name: String,
server_tool_name: String,
description: String,
parameter_schema: serde_json::Value,
timeout: Option<u64>,
trust: Option<bool>,
) -> DiscoveredMcpTool {
DiscoveredMcpTool::new(
mcp_client,
server_name,
server_tool_name,
description,
parameter_schema,
timeout,
trust,
None,
)
}
pub fn create_fully_qualified_tool(
mcp_client: Arc<RwLock<McpClient>>,
server_name: String,
server_tool_name: String,
description: String,
parameter_schema: serde_json::Value,
timeout: Option<u64>,
trust: Option<bool>,
) -> DiscoveredMcpTool {
let tool = DiscoveredMcpTool::new(
mcp_client,
server_name,
server_tool_name,
description,
parameter_schema,
timeout,
trust,
None,
);
tool.as_fully_qualified_tool()
}
pub fn create_mock_tool(
server_name: String,
server_tool_name: String,
description: String,
parameter_schema: serde_json::Value,
) -> DiscoveredMcpTool {
let mock_client = MockCallableTool::create_mock_client(server_name.clone());
Self::create_discovered_tool(
mock_client,
server_name,
server_tool_name,
description,
parameter_schema,
Some(30000),
Some(true),
)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
#[test]
fn test_generate_valid_name() {
assert_eq!(generate_valid_name("test-tool"), "test-tool");
assert_eq!(generate_valid_name("test tool"), "test_tool");
assert_eq!(generate_valid_name("test@tool#"), "test_tool_");
}
#[test]
fn test_discovered_mcp_tool_creation() {
let mock_client = MockCallableTool::create_mock_client("test_server".to_string());
let tool = McpToolFactory::create_discovered_tool(
mock_client,
"test_server".to_string(),
"test_tool".to_string(),
"Test tool description".to_string(),
serde_json::json!({"type": "object"}),
Some(30000),
Some(false),
);
assert_eq!(tool.server_name(), "test_server");
assert_eq!(tool.server_tool_name(), "test_tool");
}
#[test]
fn test_mock_tool_creation() {
let tool = McpToolFactory::create_mock_tool(
"test_server".to_string(),
"test_tool".to_string(),
"Test tool description".to_string(),
serde_json::json!({"type": "object"}),
);
assert_eq!(tool.server_name(), "test_server");
assert_eq!(tool.server_tool_name(), "test_tool");
}
#[tokio::test]
async fn test_tool_execution() {
let tool = McpToolFactory::create_mock_tool(
"test_server".to_string(),
"test_tool".to_string(),
"Test tool description".to_string(),
serde_json::json!({
"type": "object",
"properties": {
"message": {"type": "string"}
}
}),
);
let mut params = HashMap::new();
params.insert("message".to_string(), serde_json::json!("Hello, World!"));
let result = tool.execute(params).await;
assert!(result.is_ok() || result.is_err());
}
}