subc_protocol/
tool_call.rs1use serde::{Deserialize, Serialize};
19use serde_json::Value;
20
21#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
23pub struct ToolCallRequest {
24 pub name: String,
26 pub arguments: Value,
30 #[serde(default, skip_serializing_if = "Option::is_none")]
45 pub tool_call_id: Option<String>,
46 #[serde(default, skip_serializing_if = "Option::is_none")]
49 pub progress_token: Option<Value>,
50 #[serde(default, skip_serializing_if = "Option::is_none")]
60 pub call_key: Option<String>,
61}
62
63impl ToolCallRequest {
64 pub fn new(name: impl Into<String>, arguments: Value) -> Self {
67 Self {
68 name: name.into(),
69 arguments,
70 tool_call_id: None,
71 progress_token: None,
72 call_key: None,
73 }
74 }
75}
76
77pub const CALL_KEY_FIELD: &str = "call_key";
80
81pub const CALL_KEY_MAX_LEN: usize = 256;
84
85#[derive(Clone, Debug, PartialEq, Eq)]
87pub enum CallKeyError {
88 Empty,
90 TooLong { length: usize },
92 InvalidCharacter { index: usize },
94}
95
96impl CallKeyError {
97 pub fn field(&self) -> &'static str {
99 CALL_KEY_FIELD
100 }
101}
102
103impl std::fmt::Display for CallKeyError {
104 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
105 match self {
106 Self::Empty => write!(f, "{CALL_KEY_FIELD} must not be empty"),
107 Self::TooLong { length } => write!(
108 f,
109 "{CALL_KEY_FIELD} is {length} bytes; at most {CALL_KEY_MAX_LEN} are allowed"
110 ),
111 Self::InvalidCharacter { index } => write!(
112 f,
113 "{CALL_KEY_FIELD} has a character at byte {index} outside printable ASCII \
114 (0x21 to 0x7E; space is not allowed)"
115 ),
116 }
117 }
118}
119
120impl std::error::Error for CallKeyError {}
121
122pub fn validate_call_key(key: &str) -> Result<(), CallKeyError> {
131 if key.is_empty() {
132 return Err(CallKeyError::Empty);
133 }
134 if key.len() > CALL_KEY_MAX_LEN {
135 return Err(CallKeyError::TooLong { length: key.len() });
136 }
137 if let Some(index) = key.bytes().position(|byte| !(0x21..=0x7e).contains(&byte)) {
138 return Err(CallKeyError::InvalidCharacter { index });
139 }
140 Ok(())
141}
142
143#[cfg(test)]
144mod tests {
145 use super::*;
146 use serde_json::json;
147
148 #[test]
149 fn omitted_optionals_decode_as_none() {
150 let request: ToolCallRequest =
151 serde_json::from_value(json!({ "name": "grep", "arguments": { "q": "x" } }))
152 .expect("two-field body decodes");
153 assert_eq!(request.tool_call_id, None);
154 assert_eq!(request.progress_token, None);
155 assert_eq!(request.call_key, None);
156 }
157
158 #[test]
159 fn call_key_round_trips_as_a_top_level_member() {
160 let request = ToolCallRequest {
161 name: "grep".to_string(),
162 arguments: json!({ "q": "x" }),
163 tool_call_id: None,
164 progress_token: None,
165 call_key: Some("run-7:call-3".to_string()),
166 };
167 let encoded = serde_json::to_value(&request).expect("encode");
168 assert_eq!(
169 encoded,
170 json!({ "name": "grep", "arguments": { "q": "x" }, "call_key": "run-7:call-3" })
171 );
172 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
173 assert_eq!(decoded, request);
174 }
175
176 #[test]
177 fn a_request_without_a_call_key_omits_the_member_and_round_trips() {
178 let request = ToolCallRequest::new("grep", json!({}));
179 let encoded = serde_json::to_value(&request).expect("encode");
180 assert!(encoded.get("call_key").is_none(), "{encoded}");
181 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
182 assert_eq!(decoded.call_key, None);
183 assert_eq!(decoded, request);
184 }
185
186 #[test]
187 fn call_key_bounds_are_one_to_256_printable_non_space_ascii() {
188 assert_eq!(validate_call_key(""), Err(CallKeyError::Empty));
189 assert_eq!(validate_call_key("k"), Ok(()));
190 assert_eq!(validate_call_key(&"k".repeat(256)), Ok(()));
191 assert_eq!(
192 validate_call_key(&"k".repeat(257)),
193 Err(CallKeyError::TooLong { length: 257 })
194 );
195 assert_eq!(validate_call_key("!~"), Ok(()), "both ends of 0x21..=0x7E");
196 assert_eq!(
197 validate_call_key("ké"),
198 Err(CallKeyError::InvalidCharacter { index: 1 })
199 );
200 assert_eq!(
201 validate_call_key("a\tb"),
202 Err(CallKeyError::InvalidCharacter { index: 1 })
203 );
204 assert_eq!(
205 validate_call_key("a\u{7f}"),
206 Err(CallKeyError::InvalidCharacter { index: 1 })
207 );
208 assert_eq!(
209 validate_call_key("a b"),
210 Err(CallKeyError::InvalidCharacter { index: 1 })
211 );
212 assert_eq!(CallKeyError::Empty.field(), "call_key");
213 }
214
215 #[test]
216 fn none_optionals_are_omitted_so_the_wire_matches_the_two_field_shape() {
217 let request = ToolCallRequest::new("grep", json!({ "q": "x" }));
218 let encoded = serde_json::to_value(&request).expect("encode");
219 assert_eq!(
220 encoded,
221 json!({ "name": "grep", "arguments": { "q": "x" } })
222 );
223 }
224
225 #[test]
226 fn tool_call_id_round_trips() {
227 let request = ToolCallRequest {
228 name: "grep".to_string(),
229 arguments: json!({ "q": "x" }),
230 tool_call_id: Some("wal-intent-42".to_string()),
231 progress_token: None,
232 call_key: None,
233 };
234 let encoded = serde_json::to_value(&request).expect("encode");
235 assert_eq!(encoded["tool_call_id"], json!("wal-intent-42"));
236 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
237 assert_eq!(decoded, request);
238 }
239
240 #[test]
241 fn unknown_members_do_not_fail_a_provider_decode() {
242 let request: ToolCallRequest = serde_json::from_value(json!({
245 "name": "grep",
246 "arguments": {},
247 "some_future_key": { "nested": true }
248 }))
249 .expect("unknown members are tolerated");
250 assert_eq!(request.name, "grep");
251 }
252}