feat: ToolError migration to ErrorData (#4051)

This commit is contained in:
Alex Hancock
2025-08-12 16:18:41 -04:00
committed by GitHub
parent 88b013194c
commit bd1eff52a4
35 changed files with 2459 additions and 1336 deletions
+81 -43
View File
@@ -1,11 +1,11 @@
use mcp_core::ToolError;
use rmcp::model::Content;
use rmcp::model::Tool;
use rmcp::model::{Content, ErrorCode, ErrorData};
use anyhow::{Context, Result};
use async_trait::async_trait;
use serde::Serialize;
use serde_json::Value;
use std::borrow::Cow;
use std::collections::HashMap;
use std::collections::VecDeque;
use std::env;
@@ -32,11 +32,11 @@ pub enum RouterToolSelectionStrategy {
#[async_trait]
pub trait RouterToolSelector: Send + Sync {
async fn select_tools(&self, params: Value) -> Result<Vec<Content>, ToolError>;
async fn index_tools(&self, tools: &[Tool], extension_name: &str) -> Result<(), ToolError>;
async fn remove_tool(&self, tool_name: &str) -> Result<(), ToolError>;
async fn record_tool_call(&self, tool_name: &str) -> Result<(), ToolError>;
async fn get_recent_tool_calls(&self, limit: usize) -> Result<Vec<String>, ToolError>;
async fn select_tools(&self, params: Value) -> Result<Vec<Content>, ErrorData>;
async fn index_tools(&self, tools: &[Tool], extension_name: &str) -> Result<(), ErrorData>;
async fn remove_tool(&self, tool_name: &str) -> Result<(), ErrorData>;
async fn record_tool_call(&self, tool_name: &str) -> Result<(), ErrorData>;
async fn get_recent_tool_calls(&self, limit: usize) -> Result<Vec<String>, ErrorData>;
fn selector_type(&self) -> RouterToolSelectionStrategy;
}
@@ -80,11 +80,15 @@ impl VectorToolSelector {
#[async_trait]
impl RouterToolSelector for VectorToolSelector {
async fn select_tools(&self, params: Value) -> Result<Vec<Content>, ToolError> {
async fn select_tools(&self, params: Value) -> Result<Vec<Content>, ErrorData> {
let query = params
.get("query")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("Missing 'query' parameter".to_string()))?;
.ok_or_else(|| ErrorData {
code: ErrorCode::INVALID_PARAMS,
message: Cow::from("Missing 'query' parameter"),
data: None,
})?;
let k = params.get("k").and_then(|v| v.as_u64()).unwrap_or(5) as usize;
@@ -93,29 +97,38 @@ impl RouterToolSelector for VectorToolSelector {
// Check if provider supports embeddings
if !self.embedding_provider.supports_embeddings() {
return Err(ToolError::ExecutionError(
"Embedding provider does not support embeddings".to_string(),
));
return Err(ErrorData {
code: ErrorCode::INTERNAL_ERROR,
message: Cow::from("Embedding provider does not support embeddings"),
data: None,
});
}
let embeddings = self
.embedding_provider
.create_embeddings(vec![query.to_string()])
.await
.map_err(|e| {
ToolError::ExecutionError(format!("Failed to generate query embedding: {}", e))
.map_err(|e| ErrorData {
code: ErrorCode::INTERNAL_ERROR,
message: Cow::from(format!("Failed to generate query embedding: {}", e)),
data: None,
})?;
let query_embedding = embeddings
.into_iter()
.next()
.ok_or_else(|| ToolError::ExecutionError("No embedding returned".to_string()))?;
let query_embedding = embeddings.into_iter().next().ok_or_else(|| ErrorData {
code: ErrorCode::INTERNAL_ERROR,
message: Cow::from("No embedding returned"),
data: None,
})?;
let vector_db = self.vector_db.read().await;
let tools = vector_db
.search_tools(query_embedding, k, extension_name)
.await
.map_err(|e| ToolError::ExecutionError(format!("Failed to search tools: {}", e)))?;
.map_err(|e| ErrorData {
code: ErrorCode::INTERNAL_ERROR,
message: Cow::from(format!("Failed to search tools: {}", e)),
data: None,
})?;
let selected_tools: Vec<Content> = tools
.into_iter()
@@ -131,7 +144,7 @@ impl RouterToolSelector for VectorToolSelector {
Ok(selected_tools)
}
async fn index_tools(&self, tools: &[Tool], extension_name: &str) -> Result<(), ToolError> {
async fn index_tools(&self, tools: &[Tool], extension_name: &str) -> Result<(), ErrorData> {
let texts_to_embed: Vec<String> = tools
.iter()
.map(|tool| {
@@ -150,17 +163,21 @@ impl RouterToolSelector for VectorToolSelector {
.collect();
if !self.embedding_provider.supports_embeddings() {
return Err(ToolError::ExecutionError(
"Embedding provider does not support embeddings".to_string(),
));
return Err(ErrorData {
code: ErrorCode::INTERNAL_ERROR,
message: Cow::from("Embedding provider does not support embeddings"),
data: None,
});
}
let embeddings = self
.embedding_provider
.create_embeddings(texts_to_embed)
.await
.map_err(|e| {
ToolError::ExecutionError(format!("Failed to generate tool embeddings: {}", e))
.map_err(|e| ErrorData {
code: ErrorCode::INTERNAL_ERROR,
message: Cow::from(format!("Failed to generate tool embeddings: {}", e)),
data: None,
})?;
// Create tool records
@@ -194,8 +211,10 @@ impl RouterToolSelector for VectorToolSelector {
let existing_tools = vector_db
.search_tools(record.vector.clone(), 1, Some(&record.extension_name))
.await
.map_err(|e| {
ToolError::ExecutionError(format!("Failed to search for existing tools: {}", e))
.map_err(|e| ErrorData {
code: ErrorCode::INTERNAL_ERROR,
message: Cow::from(format!("Failed to search for existing tools: {}", e)),
data: None,
})?;
// Only add if no exact match found
@@ -212,21 +231,30 @@ impl RouterToolSelector for VectorToolSelector {
vector_db
.index_tools(new_tool_records)
.await
.map_err(|e| ToolError::ExecutionError(format!("Failed to index tools: {}", e)))?;
.map_err(|e| ErrorData {
code: ErrorCode::INTERNAL_ERROR,
message: Cow::from(format!("Failed to index tools: {}", e)),
data: None,
})?;
}
Ok(())
}
async fn remove_tool(&self, tool_name: &str) -> Result<(), ToolError> {
async fn remove_tool(&self, tool_name: &str) -> Result<(), ErrorData> {
let vector_db = self.vector_db.read().await;
vector_db.remove_tool(tool_name).await.map_err(|e| {
ToolError::ExecutionError(format!("Failed to remove tool {}: {}", tool_name, e))
})?;
vector_db
.remove_tool(tool_name)
.await
.map_err(|e| ErrorData {
code: ErrorCode::INTERNAL_ERROR,
message: Cow::from(format!("Failed to remove tool {}: {}", tool_name, e)),
data: None,
})?;
Ok(())
}
async fn record_tool_call(&self, tool_name: &str) -> Result<(), ToolError> {
async fn record_tool_call(&self, tool_name: &str) -> Result<(), ErrorData> {
let mut recent_calls = self.recent_tool_calls.write().await;
if recent_calls.len() >= 100 {
recent_calls.pop_front();
@@ -235,7 +263,7 @@ impl RouterToolSelector for VectorToolSelector {
Ok(())
}
async fn get_recent_tool_calls(&self, limit: usize) -> Result<Vec<String>, ToolError> {
async fn get_recent_tool_calls(&self, limit: usize) -> Result<Vec<String>, ErrorData> {
let recent_calls = self.recent_tool_calls.read().await;
Ok(recent_calls.iter().rev().take(limit).cloned().collect())
}
@@ -263,11 +291,15 @@ impl LLMToolSelector {
#[async_trait]
impl RouterToolSelector for LLMToolSelector {
async fn select_tools(&self, params: Value) -> Result<Vec<Content>, ToolError> {
async fn select_tools(&self, params: Value) -> Result<Vec<Content>, ErrorData> {
let query = params
.get("query")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("Missing 'query' parameter".to_string()))?;
.ok_or_else(|| ErrorData {
code: ErrorCode::INVALID_PARAMS,
message: Cow::from("Missing 'query' parameter"),
data: None,
})?;
let extension_name = params
.get("extension_name")
@@ -297,8 +329,10 @@ impl RouterToolSelector for LLMToolSelector {
};
let user_prompt =
render_global_file("router_tool_selector.md", &context).map_err(|e| {
ToolError::ExecutionError(format!("Failed to render prompt template: {}", e))
render_global_file("router_tool_selector.md", &context).map_err(|e| ErrorData {
code: ErrorCode::INTERNAL_ERROR,
message: Cow::from(format!("Failed to render prompt template: {}", e)),
data: None,
})?;
let user_message = Message::user().with_text(&user_prompt);
@@ -306,7 +340,11 @@ impl RouterToolSelector for LLMToolSelector {
.llm_provider
.complete("", &[user_message], &[])
.await
.map_err(|e| ToolError::ExecutionError(format!("Failed to search tools: {}", e)))?;
.map_err(|e| ErrorData {
code: ErrorCode::INTERNAL_ERROR,
message: Cow::from(format!("Failed to search tools: {}", e)),
data: None,
})?;
// Extract just the message content from the response
let (message, _usage) = response;
@@ -325,7 +363,7 @@ impl RouterToolSelector for LLMToolSelector {
}
}
async fn index_tools(&self, tools: &[Tool], extension_name: &str) -> Result<(), ToolError> {
async fn index_tools(&self, tools: &[Tool], extension_name: &str) -> Result<(), ErrorData> {
let mut tool_strings = self.tool_strings.write().await;
for tool in tools {
@@ -354,7 +392,7 @@ impl RouterToolSelector for LLMToolSelector {
Ok(())
}
async fn remove_tool(&self, tool_name: &str) -> Result<(), ToolError> {
async fn remove_tool(&self, tool_name: &str) -> Result<(), ErrorData> {
let mut tool_strings = self.tool_strings.write().await;
if let Some(extension_name) = tool_name.split("__").next() {
tool_strings.remove(extension_name);
@@ -362,7 +400,7 @@ impl RouterToolSelector for LLMToolSelector {
Ok(())
}
async fn record_tool_call(&self, tool_name: &str) -> Result<(), ToolError> {
async fn record_tool_call(&self, tool_name: &str) -> Result<(), ErrorData> {
let mut recent_calls = self.recent_tool_calls.write().await;
if recent_calls.len() >= 100 {
recent_calls.pop_front();
@@ -371,7 +409,7 @@ impl RouterToolSelector for LLMToolSelector {
Ok(())
}
async fn get_recent_tool_calls(&self, limit: usize) -> Result<Vec<String>, ToolError> {
async fn get_recent_tool_calls(&self, limit: usize) -> Result<Vec<String>, ErrorData> {
let recent_calls = self.recent_tool_calls.read().await;
Ok(recent_calls.iter().rev().take(limit).cloned().collect())
}