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}