use std::sync::Arc;
use rmcp::handler::server::ServerHandler;
use rmcp::model::{
CallToolRequestParams, CallToolResult, ContentBlock, Implementation, ListToolsResult,
PaginatedRequestParams, ServerCapabilities, ServerInfo, Tool,
};
use rmcp::service::{NotificationContext, RequestContext, RoleServer};
use rmcp::transport::io::stdio;
use rmcp::ErrorData as RmcpError;
use rmcp::{serve_server, ServiceExt};
use rtb_app::app::App;
use rtb_app::command::{Command, BUILTIN_COMMANDS};
use crate::error::{McpError, Result};
use crate::transport::Transport;
pub struct McpServer {
app: App,
tools: Arc<Vec<RegisteredTool>>,
transport: Transport,
}
#[derive(Clone)]
struct RegisteredTool {
name: &'static str,
about: &'static str,
aliases: &'static [&'static str],
schema: serde_json::Map<String, serde_json::Value>,
factory: fn() -> Box<dyn Command>,
}
impl McpServer {
#[must_use]
pub fn new(app: App, transport: Transport) -> Self {
let mut tools = Vec::new();
for factory in BUILTIN_COMMANDS {
let cmd = factory();
if !cmd.mcp_exposed() {
continue;
}
let spec = cmd.spec();
let schema = match cmd.mcp_input_schema() {
Some(serde_json::Value::Object(map)) => map,
Some(other) => {
let mut map = serde_json::Map::new();
map.insert("type".into(), serde_json::Value::String("object".into()));
let _ = other; map
}
None => {
let mut map = serde_json::Map::new();
map.insert("type".into(), serde_json::Value::String("object".into()));
map
}
};
tools.push(RegisteredTool {
name: spec.name,
about: spec.about,
aliases: spec.aliases,
schema,
factory: *factory,
});
}
Self { app, tools: Arc::new(tools), transport }
}
#[must_use]
pub fn tool_count(&self) -> usize {
self.tools.len()
}
pub fn tool_manifest(&self) -> impl Iterator<Item = (&str, &str, serde_json::Value)> + '_ {
self.tools.iter().map(|t| (t.name, t.about, serde_json::Value::Object(t.schema.clone())))
}
pub async fn dispatch(&self, name: &str) -> Result<()> {
let tool = self
.tools
.iter()
.find(|t| t.name == name || t.aliases.contains(&name))
.ok_or_else(|| McpError::Protocol(format!("unknown MCP tool: {name}")))?;
let cmd = (tool.factory)();
match cmd.run(self.app.clone()).await {
Ok(()) => Ok(()),
Err(e) => {
Err(McpError::Command { command: tool.name.to_string(), message: e.to_string() })
}
}
}
pub async fn serve(self) -> Result<()> {
match self.transport.clone() {
Transport::Stdio => self.serve_stdio().await,
Transport::Sse { .. } | Transport::Http { .. } => Err(McpError::Transport(
"SSE / streamable HTTP transports are not yet implemented in rtb-mcp v0.1; \
use --transport stdio"
.to_string(),
)),
}
}
async fn serve_stdio(self) -> Result<()> {
let (stdin, stdout) = stdio();
self.serve_with_pipe(stdin, stdout).await
}
pub async fn serve_with_pipe<R, W>(self, read: R, write: W) -> Result<()>
where
R: tokio::io::AsyncRead + Send + Sync + Unpin + 'static,
W: tokio::io::AsyncWrite + Send + Sync + Unpin + 'static,
{
let shutdown = self.app.shutdown.clone();
let handler = McpHandler::from_server(&self);
let running =
handler.serve((read, write)).await.map_err(|e| McpError::Transport(e.to_string()))?;
let cancel = running.cancellation_token();
tokio::select! {
res = running.waiting() => {
res.map(|_| ()).map_err(|e| McpError::Transport(e.to_string()))
}
() = shutdown.cancelled() => {
cancel.cancel();
Ok(())
}
}
}
}
#[derive(Clone)]
struct McpHandler {
app: App,
tools: Arc<Vec<RegisteredTool>>,
server_name: String,
server_version: String,
}
impl McpHandler {
fn from_server(server: &McpServer) -> Self {
Self {
app: server.app.clone(),
tools: server.tools.clone(),
server_name: server.app.metadata.name.clone(),
server_version: server.app.version.version.to_string(),
}
}
fn render_tools(&self) -> Vec<Tool> {
self.tools
.iter()
.map(|t| Tool::new(t.name.to_string(), t.about.to_string(), t.schema.clone()))
.collect()
}
fn find_tool(&self, name: &str) -> Option<RegisteredTool> {
self.tools.iter().find(|t| t.name == name || t.aliases.contains(&name)).cloned()
}
}
impl ServerHandler for McpHandler {
fn get_info(&self) -> ServerInfo {
let mut info = ServerInfo::new(ServerCapabilities::builder().enable_tools().build());
info.server_info =
Implementation::new(self.server_name.clone(), self.server_version.clone());
info
}
async fn list_tools(
&self,
_request: Option<PaginatedRequestParams>,
_context: RequestContext<RoleServer>,
) -> std::result::Result<ListToolsResult, RmcpError> {
Ok(ListToolsResult { tools: self.render_tools(), next_cursor: None, meta: None })
}
async fn call_tool(
&self,
request: CallToolRequestParams,
_context: RequestContext<RoleServer>,
) -> std::result::Result<CallToolResult, RmcpError> {
let Some(tool) = self.find_tool(request.name.as_ref()) else {
return Err(RmcpError::invalid_params(
format!("unknown MCP tool: {}", request.name),
None,
));
};
let cmd = (tool.factory)();
match cmd.run(self.app.clone()).await {
Ok(()) => {
Ok(CallToolResult::success(vec![ContentBlock::text(format!("{} ok", tool.name))]))
}
Err(e) => {
Ok(CallToolResult::error(vec![ContentBlock::text(format!("{}: {}", tool.name, e))]))
}
}
}
async fn on_initialized(&self, _context: NotificationContext<RoleServer>) {}
}
#[doc(hidden)]
#[allow(dead_code)]
const fn _link_check() {
let _ = serve_server::<McpHandler, (tokio::io::Stdin, tokio::io::Stdout), _, _>;
}