use std::collections::HashMap;
use itertools::Itertools;
use rmcp::ErrorData as McpError;
use rmcp::RoleServer;
use rmcp::ServerHandler;
use rmcp::model::CallToolRequestParams;
use rmcp::model::CallToolResponse;
use rmcp::model::ListToolsResult;
use rmcp::model::PaginatedRequestParams;
use rmcp::model::ServerCapabilities;
use rmcp::model::ServerInfo;
use rmcp::model::Tool;
use rmcp::service::RequestContext;
use super::tool;
use super::tool::ToolDef;
pub(crate) struct McpService {
tool_defs: HashMap<String, ToolDef>,
tools: Vec<Tool>,
}
impl McpService {
pub(crate) fn new() -> Self {
let all_defs = tool::get_all_tool_definitions();
let tool_defs = all_defs
.iter()
.map(|tool_def| (tool_def.name().to_string(), tool_def.clone()))
.collect();
let tools: Vec<_> = all_defs
.iter()
.map(ToolDef::to_tool)
.sorted_by_key(|tool| {
tool.annotations
.as_ref()
.and_then(|ann| ann.title.as_ref())
.map_or_else(|| tool.name.as_ref(), String::as_str)
.to_string()
})
.collect();
Self { tool_defs, tools }
}
fn get_tool_def(&self, name: &str) -> Option<&ToolDef> { self.tool_defs.get(name) }
fn list_mcp_tools(&self) -> ListToolsResult {
ListToolsResult::with_all_items(self.tools.clone())
}
}
impl ServerHandler for McpService {
fn get_info(&self) -> ServerInfo {
let mut info = rmcp::model::ServerInfo::default();
info.capabilities = ServerCapabilities::builder().enable_tools().build();
info
}
async fn list_tools(
&self,
_: Option<PaginatedRequestParams>,
_: RequestContext<RoleServer>,
) -> Result<ListToolsResult, McpError> {
Ok(self.list_mcp_tools())
}
async fn call_tool(
&self,
request: CallToolRequestParams,
_: RequestContext<RoleServer>,
) -> Result<CallToolResponse, McpError> {
let tool_def = self.get_tool_def(&request.name).ok_or_else(|| {
McpError::invalid_params(format!("unknown tool: {}", request.name), None)
})?;
tool_def.call_tool(request).await.map(Into::into)
}
}