A command enum served over MCP.
3use std::{future::Future, marker::PhantomData, sync::Arc};
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};
How a server introduces itself to a client.
What the server is for, shown to the model alongside the tool list.
Serves every variant of the command enum C as a tool, running each call
through run.
run returns the command's result as JSON, or a message the model can
act on. Neither is a protocol error: a refused command is still an answer.
tag is the enum's serde tag field (#[serde(tag = "…")]).
Drop the named commands from the tool list: the ones that make no sense as a single call and answer, such as a subscription.
The tools this server lists, by name.
Serve on stdin/stdout until the client closes the connection.
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 {
Add one to a number.
170 Increment { value: i64 },
Always refused.
172 Refuse,
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}