1//! A command enum served over MCP. 2 3use std::{future::Future, marker::PhantomData, sync::Arc}; 4 5use rmcp::{ 6 ErrorData, RoleServer, ServerHandler, ServiceExt, 7 model::{ 8 CallToolRequestParams, CallToolResponse, CallToolResult, ContentBlock, Implementation, 9 ListToolsResult, PaginatedRequestParams, ServerCapabilities, ServerConfig, Tool, 10 }, 11 service::RequestContext, 12}; 13use schemars::JsonSchema; 14use serde::de::DeserializeOwned; 15use serde_json::Value; 16 17use crate::reflect::{CommandTool, ReflectError, command_from_call, command_tools}; 18 19/// How a server introduces itself to a client. 20#[derive(Debug, Clone)] 21pub struct ServerIdentity { 22 pub name: String, 23 pub version: String, 24 /// What the server is for, shown to the model alongside the tool list. 25 pub instructions: String, 26} 27 28/// Why serving stopped. 29#[derive(Debug)] 30pub struct ServeError(String); 31 32impl std::fmt::Display for ServeError { 33 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { 34 f.write_str(&self.0) 35 } 36} 37 38impl std::error::Error for ServeError {} 39 40/// Serves every variant of the command enum `C` as a tool, running each call 41/// through `run`. 42/// 43/// `run` returns the command's result as JSON, or a message the model can 44/// act on. Neither is a protocol error: a refused command is still an answer. 45pub struct CommandServer<C, H> { 46 identity: ServerIdentity, 47 tag: String, 48 tools: Vec<Tool>, 49 run: H, 50 command: PhantomData<fn() -> C>, 51} 52 53impl<C, H, F> CommandServer<C, H> 54where 55 C: JsonSchema + DeserializeOwned + Send + 'static, 56 H: Fn(C) -> F + Send + Sync + 'static, 57 F: Future<Output = Result<Value, String>> + Send + 'static, 58{ 59 /// `tag` is the enum's serde tag field (`#[serde(tag = "…")]`). 60 pub fn new(identity: ServerIdentity, tag: &str, run: H) -> Result<Self, ReflectError> { 61 Ok(Self { 62 identity, 63 tag: tag.to_owned(), 64 tools: command_tools::<C>(tag)?.into_iter().map(tool).collect(), 65 run, 66 command: PhantomData, 67 }) 68 } 69 70 /// Drop the named commands from the tool list: the ones that make no 71 /// sense as a single call and answer, such as a subscription. 72 pub fn without(mut self, names: &[&str]) -> Self { 73 self.tools.retain(|tool| !names.contains(&tool.name.as_ref())); 74 self 75 } 76 77 /// The tools this server lists, by name. 78 pub fn tool_names(&self) -> Vec<&str> { 79 self.tools.iter().map(|tool| tool.name.as_ref()).collect() 80 } 81 82 /// Serve on stdin/stdout until the client closes the connection. 83 pub async fn serve_stdio(self) -> Result<(), ServeError> { 84 let running = self 85 .serve(rmcp::transport::stdio()) 86 .await 87 .map_err(|why| ServeError(why.to_string()))?; 88 running 89 .waiting() 90 .await 91 .map_err(|why| ServeError(why.to_string()))?; 92 Ok(()) 93 } 94 95 async fn call(&self, name: &str, arguments: Option<rmcp::model::JsonObject>) -> CallToolResult { 96 let command = match command_from_call::<C>(&self.tag, name, arguments) { 97 Ok(command) => command, 98 Err(why) => return failure(format!("invalid arguments for `{name}`: {why}")), 99 }; 100 match (self.run)(command).await { 101 Ok(result) => CallToolResult::success(vec![ContentBlock::text(result.to_string())]), 102 Err(why) => failure(why), 103 } 104 } 105} 106 107fn tool(command: CommandTool) -> Tool { 108 Tool::new( 109 command.name, 110 command.description.unwrap_or_default(), 111 Arc::new(command.input_schema), 112 ) 113} 114 115fn failure(message: String) -> CallToolResult { 116 CallToolResult::error(vec![ContentBlock::text(message)]) 117} 118 119impl<C, H, F> ServerHandler for CommandServer<C, H> 120where 121 C: JsonSchema + DeserializeOwned + Send + 'static, 122 H: Fn(C) -> F + Send + Sync + 'static, 123 F: Future<Output = Result<Value, String>> + Send + 'static, 124{ 125 fn get_info(&self) -> ServerConfig { 126 ServerConfig::new(ServerCapabilities::builder().enable_tools().build()) 127 .with_server_info(Implementation::new( 128 self.identity.name.clone(), 129 self.identity.version.clone(), 130 )) 131 .with_instructions(self.identity.instructions.clone()) 132 } 133 134 async fn list_tools( 135 &self, 136 _request: Option<PaginatedRequestParams>, 137 _context: RequestContext<RoleServer>, 138 ) -> Result<ListToolsResult, ErrorData> { 139 Ok(ListToolsResult::with_all_items(self.tools.clone())) 140 } 141 142 fn get_tool(&self, name: &str) -> Option<Tool> { 143 self.tools.iter().find(|tool| tool.name == name).cloned() 144 } 145 146 async fn call_tool( 147 &self, 148 request: CallToolRequestParams, 149 _context: RequestContext<RoleServer>, 150 ) -> Result<CallToolResponse, ErrorData> { 151 if self.get_tool(&request.name).is_none() { 152 return Err(ErrorData::invalid_params( 153 format!("no tool named `{}`", request.name), 154 None, 155 )); 156 } 157 Ok(self.call(&request.name, request.arguments).await.into()) 158 } 159} 160 161#[cfg(test)] 162mod tests { 163 use super::*; 164 use serde::Deserialize; 165 166 #[derive(Debug, Deserialize, JsonSchema)] 167 #[serde(tag = "type", rename_all = "snake_case")] 168 enum Command { 169 /// Add one to a number. 170 Increment { value: i64 }, 171 /// Always refused. 172 Refuse, 173 /// A stream, not a call. 174 Subscribe, 175 } 176 177 fn identity() -> ServerIdentity { 178 ServerIdentity { 179 name: "test".into(), 180 version: "0".into(), 181 instructions: "test".into(), 182 } 183 } 184 185 async fn run(command: Command) -> Result<Value, String> { 186 match command { 187 Command::Increment { value } => Ok(serde_json::json!({ "value": value + 1 })), 188 Command::Refuse => Err("refused".into()), 189 Command::Subscribe => Ok(Value::Null), 190 } 191 } 192 193 fn text(result: &CallToolResult) -> String { 194 serde_json::to_value(&result.content).unwrap()[0]["text"] 195 .as_str() 196 .unwrap() 197 .to_owned() 198 } 199 200 #[test] 201 fn lists_every_command_except_the_excluded() { 202 let server = CommandServer::new(identity(), "type", run) 203 .unwrap() 204 .without(&["subscribe"]); 205 assert_eq!(server.tool_names(), ["increment", "refuse"]); 206 } 207 208 #[tokio::test] 209 async fn a_call_runs_its_command() { 210 let server = CommandServer::new(identity(), "type", run).unwrap(); 211 let arguments = serde_json::json!({ "value": 41 }); 212 let result = server.call("increment", arguments.as_object().cloned()).await; 213 assert_eq!(result.is_error, Some(false)); 214 assert_eq!(text(&result), r#"{"value":42}"#); 215 } 216 217 #[tokio::test] 218 async fn a_refusal_and_bad_arguments_are_tool_errors() { 219 let server = CommandServer::new(identity(), "type", run).unwrap(); 220 let refused = server.call("refuse", None).await; 221 assert_eq!(refused.is_error, Some(true)); 222 assert_eq!(text(&refused), "refused"); 223 224 let invalid = server.call("increment", None).await; 225 assert_eq!(invalid.is_error, Some(true)); 226 assert!(text(&invalid).starts_with("invalid arguments for `increment`")); 227 } 228}