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}