rustcrates.git / mcp / src / reflect.rs
1//! A tagged command enum, read as a list of tools.
2
3use schemars::JsonSchema;
4use serde::de::DeserializeOwned;
5use serde_json::{Map, Value};
6
7/// One variant of a command enum, as a tool.
8#[derive(Debug, Clone, PartialEq)]
9pub struct CommandTool {
10    /// The variant's tag value, e.g. `focus_left`.
11    pub name: String,
12    /// The variant's doc comment.
13    pub description: Option<String>,
14    /// JSON Schema of the variant's fields, without the tag.
15    pub input_schema: Map<String, Value>,
16}
17
18/// Why an enum cannot be read as tools.
19#[derive(Debug, Clone, PartialEq)]
20pub enum ReflectError {
21    /// The type's schema is not a list of alternatives, so it is not an enum
22    /// with one variant per command.
23    NotAnEnum,
24    /// A variant has no constant string under the tag field, so no tool name
25    /// identifies it. The index is its position in the enum.
26    UntaggedVariant { index: usize },
27    /// Two variants share a tag value.
28    DuplicateName(String),
29}
30
31impl std::fmt::Display for ReflectError {
32    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
33        match self {
34            Self::NotAnEnum => write!(f, "the command type's schema is not an enum of variants"),
35            Self::UntaggedVariant { index } => {
36                write!(f, "variant {index} of the command enum has no constant tag")
37            }
38            Self::DuplicateName(name) => write!(f, "two commands are both named `{name}`"),
39        }
40    }
41}
42
43impl std::error::Error for ReflectError {}
44
45/// Every variant of `C` as a tool. `C` must be a serde enum tagged by the
46/// field `tag` (`#[serde(tag = "…")]`).
47pub fn command_tools<C: JsonSchema>(tag: &str) -> Result<Vec<CommandTool>, ReflectError> {
48    let root = schemars::schema_for!(C).to_value();
49    let variants = root
50        .get("oneOf")
51        .or_else(|| root.get("anyOf"))
52        .and_then(Value::as_array)
53        .ok_or(ReflectError::NotAnEnum)?;
54    let definitions = root.get("$defs");
55
56    let mut tools: Vec<CommandTool> = Vec::with_capacity(variants.len());
57    for (index, variant) in variants.iter().enumerate() {
58        let tool = variant_tool(variant, tag, definitions)
59            .ok_or(ReflectError::UntaggedVariant { index })?;
60        if tools.iter().any(|seen| seen.name == tool.name) {
61            return Err(ReflectError::DuplicateName(tool.name));
62        }
63        tools.push(tool);
64    }
65    Ok(tools)
66}
67
68fn variant_tool(variant: &Value, tag: &str, definitions: Option<&Value>) -> Option<CommandTool> {
69    let properties = variant.get("properties")?.as_object()?;
70    let name = tag_value(properties.get(tag)?)?.to_owned();
71
72    let mut fields = properties.clone();
73    fields.remove(tag);
74    let required: Vec<Value> = variant
75        .get("required")
76        .and_then(Value::as_array)
77        .map(|names| {
78            names
79                .iter()
80                .filter(|name| name.as_str() != Some(tag))
81                .cloned()
82                .collect()
83        })
84        .unwrap_or_default();
85
86    let mut input_schema = Map::new();
87    input_schema.insert("type".into(), "object".into());
88    input_schema.insert("properties".into(), Value::Object(fields));
89    if !required.is_empty() {
90        input_schema.insert("required".into(), Value::Array(required));
91    }
92    // Field types defined elsewhere in the enum's schema are referenced by
93    // `$ref`, which resolves against the schema it appears in.
94    if let Some(definitions) = definitions {
95        input_schema.insert("$defs".into(), definitions.clone());
96    }
97
98    Some(CommandTool {
99        name,
100        description: variant
101            .get("description")
102            .and_then(Value::as_str)
103            .map(str::to_owned),
104        input_schema,
105    })
106}
107
108/// The tag's constant: `const`, or an `enum` of exactly one string.
109fn tag_value(tag_schema: &Value) -> Option<&str> {
110    if let Some(constant) = tag_schema.get("const") {
111        return constant.as_str();
112    }
113    match tag_schema.get("enum")?.as_array()?.as_slice() {
114        [only] => only.as_str(),
115        _ => None,
116    }
117}
118
119/// The command a tool call names: the call's arguments with the tag put back.
120pub fn command_from_call<C: DeserializeOwned>(
121    tag: &str,
122    name: &str,
123    arguments: Option<Map<String, Value>>,
124) -> Result<C, serde_json::Error> {
125    let mut command = arguments.unwrap_or_default();
126    command.insert(tag.to_owned(), Value::String(name.to_owned()));
127    serde_json::from_value(Value::Object(command))
128}
129
130#[cfg(test)]
131mod tests {
132    use super::*;
133    use serde::Deserialize;
134
135    #[derive(Debug, PartialEq, Deserialize, JsonSchema)]
136    #[serde(rename_all = "snake_case")]
137    enum Edge {
138        Left,
139        Right,
140    }
141
142    /// Commands a window manager takes.
143    #[derive(Debug, PartialEq, Deserialize, JsonSchema)]
144    #[serde(tag = "type", rename_all = "snake_case")]
145    enum Command {
146        /// Focus the column to the left.
147        FocusLeft,
148        /// Resize the focused column.
149        Resize {
150            /// Width delta in pixels.
151            delta: i32,
152        },
153        /// Move to an edge, optionally on a named monitor.
154        MoveTo { edge: Edge, monitor: Option<String> },
155    }
156
157    fn tool(name: &str) -> CommandTool {
158        command_tools::<Command>("type")
159            .unwrap()
160            .into_iter()
161            .find(|tool| tool.name == name)
162            .unwrap()
163    }
164
165    #[test]
166    fn every_variant_is_a_tool_named_by_its_tag() {
167        let names: Vec<String> = command_tools::<Command>("type")
168            .unwrap()
169            .into_iter()
170            .map(|tool| tool.name)
171            .collect();
172        assert_eq!(names, ["focus_left", "resize", "move_to"]);
173    }
174
175    #[test]
176    fn the_doc_comment_is_the_description() {
177        assert_eq!(
178            tool("focus_left").description.as_deref(),
179            Some("Focus the column to the left.")
180        );
181    }
182
183    #[test]
184    fn a_unit_variant_takes_no_arguments() {
185        let schema = tool("focus_left").input_schema;
186        assert_eq!(schema["properties"], serde_json::json!({}));
187        assert!(!schema.contains_key("required"));
188    }
189
190    #[test]
191    fn fields_are_arguments_and_the_tag_is_not() {
192        let schema = tool("resize").input_schema;
193        assert_eq!(schema["properties"]["delta"]["type"], "integer");
194        assert_eq!(schema["properties"]["delta"]["description"], "Width delta in pixels.");
195        assert!(schema["properties"].get("type").is_none());
196        assert_eq!(schema["required"], serde_json::json!(["delta"]));
197    }
198
199    #[test]
200    fn optional_fields_are_not_required_and_named_types_resolve() {
201        let schema = tool("move_to").input_schema;
202        assert_eq!(schema["required"], serde_json::json!(["edge"]));
203        let reference = schema["properties"]["edge"]["$ref"].as_str().unwrap();
204        let definition = reference.strip_prefix("#/$defs/").unwrap();
205        assert!(schema["$defs"].get(definition).is_some());
206    }
207
208    #[test]
209    fn a_call_becomes_its_command() {
210        let arguments = serde_json::json!({ "delta": -40 });
211        let command: Command =
212            command_from_call("type", "resize", arguments.as_object().cloned()).unwrap();
213        assert_eq!(command, Command::Resize { delta: -40 });
214
215        let command: Command = command_from_call("type", "focus_left", None).unwrap();
216        assert_eq!(command, Command::FocusLeft);
217    }
218
219    #[test]
220    fn a_call_with_wrong_arguments_is_refused() {
221        let arguments = serde_json::json!({ "delta": "wide" });
222        assert!(command_from_call::<Command>("type", "resize", arguments.as_object().cloned()).is_err());
223    }
224
225    #[test]
226    fn a_struct_is_not_an_enum() {
227        #[derive(JsonSchema)]
228        struct NotACommand {
229            #[allow(dead_code)]
230            delta: i32,
231        }
232        assert_eq!(command_tools::<NotACommand>("type"), Err(ReflectError::NotAnEnum));
233    }
234
235    #[test]
236    fn an_untagged_enum_is_refused() {
237        #[derive(JsonSchema)]
238        #[allow(dead_code)]
239        enum External {
240            One { delta: i32 },
241        }
242        assert_eq!(
243            command_tools::<External>("type"),
244            Err(ReflectError::UntaggedVariant { index: 0 })
245        );
246    }
247}