feat: ToolError migration to ErrorData (#4051)
This commit is contained in:
@@ -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())
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user