1use 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 #[serde(default, skip_serializing_if = "Option::is_none")]
68 pub schema_pin: Option<String>,
69}
70
71impl ToolCallRequest {
72 pub fn new(name: impl Into<String>, arguments: Value) -> Self {
75 Self {
76 name: name.into(),
77 arguments,
78 tool_call_id: None,
79 progress_token: None,
80 call_key: None,
81 schema_pin: None,
82 }
83 }
84}
85
86pub const CALL_KEY_FIELD: &str = "call_key";
89
90pub const SCHEMA_PIN_FIELD: &str = "schema_pin";
93
94pub const OPAQUE_FIELD_MAX_LEN: usize = 256;
97
98pub const CALL_KEY_MAX_LEN: usize = OPAQUE_FIELD_MAX_LEN;
100
101pub const SCHEMA_PIN_MAX_LEN: usize = OPAQUE_FIELD_MAX_LEN;
103
104#[derive(Clone, Debug, PartialEq, Eq)]
108pub enum OpaqueFieldError {
109 Empty { field: &'static str },
111 TooLong { field: &'static str, length: usize },
113 InvalidCharacter { field: &'static str, index: usize },
115}
116
117pub type CallKeyError = OpaqueFieldError;
120
121impl OpaqueFieldError {
122 pub fn field(&self) -> &'static str {
124 match self {
125 Self::Empty { field }
126 | Self::TooLong { field, .. }
127 | Self::InvalidCharacter { field, .. } => field,
128 }
129 }
130}
131
132impl std::fmt::Display for OpaqueFieldError {
133 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
134 match self {
135 Self::Empty { field } => write!(f, "{field} must not be empty"),
136 Self::TooLong { field, length } => write!(
137 f,
138 "{field} is {length} bytes; at most {OPAQUE_FIELD_MAX_LEN} are allowed"
139 ),
140 Self::InvalidCharacter { field, index } => write!(
141 f,
142 "{field} has a character at byte {index} outside printable ASCII \
143 (0x21 to 0x7E; space is not allowed)"
144 ),
145 }
146 }
147}
148
149impl std::error::Error for OpaqueFieldError {}
150
151fn validate_opaque_field(field: &'static str, value: &str) -> Result<(), OpaqueFieldError> {
162 if value.is_empty() {
163 return Err(OpaqueFieldError::Empty { field });
164 }
165 if value.len() > OPAQUE_FIELD_MAX_LEN {
166 return Err(OpaqueFieldError::TooLong {
167 field,
168 length: value.len(),
169 });
170 }
171 if let Some(index) = value
172 .bytes()
173 .position(|byte| !(0x21..=0x7e).contains(&byte))
174 {
175 return Err(OpaqueFieldError::InvalidCharacter { field, index });
176 }
177 Ok(())
178}
179
180pub fn validate_call_key(key: &str) -> Result<(), CallKeyError> {
183 validate_opaque_field(CALL_KEY_FIELD, key)
184}
185
186pub fn validate_schema_pin(pin: &str) -> Result<(), OpaqueFieldError> {
189 validate_opaque_field(SCHEMA_PIN_FIELD, pin)
190}
191
192#[cfg(test)]
193mod tests {
194 use super::*;
195 use serde_json::json;
196
197 #[test]
198 fn omitted_optionals_decode_as_none() {
199 let request: ToolCallRequest =
200 serde_json::from_value(json!({ "name": "grep", "arguments": { "q": "x" } }))
201 .expect("two-field body decodes");
202 assert_eq!(request.tool_call_id, None);
203 assert_eq!(request.progress_token, None);
204 assert_eq!(request.call_key, None);
205 assert_eq!(request.schema_pin, None);
206 }
207
208 #[test]
209 fn call_key_round_trips_as_a_top_level_member() {
210 let request = ToolCallRequest {
211 name: "grep".to_string(),
212 arguments: json!({ "q": "x" }),
213 tool_call_id: None,
214 progress_token: None,
215 call_key: Some("run-7:call-3".to_string()),
216 schema_pin: None,
217 };
218 let encoded = serde_json::to_value(&request).expect("encode");
219 assert_eq!(
220 encoded,
221 json!({ "name": "grep", "arguments": { "q": "x" }, "call_key": "run-7:call-3" })
222 );
223 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
224 assert_eq!(decoded, request);
225 }
226
227 #[test]
228 fn a_request_without_a_call_key_omits_the_member_and_round_trips() {
229 let request = ToolCallRequest::new("grep", json!({}));
230 let encoded = serde_json::to_value(&request).expect("encode");
231 assert!(encoded.get("call_key").is_none(), "{encoded}");
232 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
233 assert_eq!(decoded.call_key, None);
234 assert_eq!(decoded, request);
235 }
236
237 #[test]
240 fn opaque_field_bounds_are_one_to_256_printable_non_space_ascii() {
241 type Validate = fn(&str) -> Result<(), OpaqueFieldError>;
242 let validators: [(&str, Validate); 2] = [
243 (CALL_KEY_FIELD, validate_call_key),
244 (SCHEMA_PIN_FIELD, validate_schema_pin),
245 ];
246 for (field, validate) in validators {
247 assert_eq!(validate(""), Err(OpaqueFieldError::Empty { field }));
248 assert_eq!(validate("k"), Ok(()));
249 assert_eq!(validate(&"k".repeat(256)), Ok(()));
250 assert_eq!(
251 validate(&"k".repeat(257)),
252 Err(OpaqueFieldError::TooLong { field, length: 257 })
253 );
254 assert_eq!(validate("!~"), Ok(()), "both ends of 0x21..=0x7E");
255 for bad in ["ké", "a\tb", "a\u{7f}", "a b"] {
256 assert_eq!(
257 validate(bad),
258 Err(OpaqueFieldError::InvalidCharacter { field, index: 1 }),
259 "{field}: {bad:?}"
260 );
261 }
262 let error = validate("").unwrap_err();
263 assert_eq!(error.field(), field);
264 assert!(error.to_string().starts_with(field), "{error}");
265 }
266 }
267
268 #[test]
269 fn schema_pin_round_trips_as_a_top_level_member() {
270 let request = ToolCallRequest {
271 schema_pin: Some("sha256:0f1e2d".to_string()),
272 ..ToolCallRequest::new("grep", json!({ "q": "x" }))
273 };
274 let encoded = serde_json::to_value(&request).expect("encode");
275 assert_eq!(
276 encoded,
277 json!({ "name": "grep", "arguments": { "q": "x" }, "schema_pin": "sha256:0f1e2d" })
278 );
279 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
280 assert_eq!(decoded, request);
281 }
282
283 #[test]
284 fn a_request_without_a_schema_pin_omits_the_member_and_round_trips() {
285 let request = ToolCallRequest::new("grep", json!({}));
286 let encoded = serde_json::to_value(&request).expect("encode");
287 assert!(encoded.get("schema_pin").is_none(), "{encoded}");
288 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
289 assert_eq!(decoded.schema_pin, None);
290 assert_eq!(decoded, request);
291 }
292
293 #[test]
294 fn none_optionals_are_omitted_so_the_wire_matches_the_two_field_shape() {
295 let request = ToolCallRequest::new("grep", json!({ "q": "x" }));
296 let encoded = serde_json::to_value(&request).expect("encode");
297 assert_eq!(
298 encoded,
299 json!({ "name": "grep", "arguments": { "q": "x" } })
300 );
301 }
302
303 #[test]
304 fn tool_call_id_round_trips() {
305 let request = ToolCallRequest {
306 name: "grep".to_string(),
307 arguments: json!({ "q": "x" }),
308 tool_call_id: Some("wal-intent-42".to_string()),
309 progress_token: None,
310 call_key: None,
311 schema_pin: None,
312 };
313 let encoded = serde_json::to_value(&request).expect("encode");
314 assert_eq!(encoded["tool_call_id"], json!("wal-intent-42"));
315 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
316 assert_eq!(decoded, request);
317 }
318
319 #[test]
320 fn unknown_members_do_not_fail_a_provider_decode() {
321 let request: ToolCallRequest = serde_json::from_value(json!({
324 "name": "grep",
325 "arguments": {},
326 "some_future_key": { "nested": true }
327 }))
328 .expect("unknown members are tolerated");
329 assert_eq!(request.name, "grep");
330 }
331}