1use serde::{Deserialize, Serialize};
19use serde_json::Value;
20
21use crate::Principal;
22
23#[derive(Clone, Debug, Deserialize, PartialEq, Eq, Serialize)]
37#[non_exhaustive]
38pub struct CallOrigin {
39 pub carrier: Principal,
41 pub call_key: String,
44}
45
46impl CallOrigin {
47 pub fn new(carrier: Principal, call_key: impl Into<String>) -> Self {
49 Self {
50 carrier,
51 call_key: call_key.into(),
52 }
53 }
54}
55
56#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
58pub struct ToolCallRequest {
59 pub name: String,
61 pub arguments: Value,
65 #[serde(default, skip_serializing_if = "Option::is_none")]
80 pub tool_call_id: Option<String>,
81 #[serde(default, skip_serializing_if = "Option::is_none")]
84 pub progress_token: Option<Value>,
85 #[serde(default, skip_serializing_if = "Option::is_none")]
95 pub call_key: Option<String>,
96 #[serde(default, skip_serializing_if = "Option::is_none")]
103 pub schema_pin: Option<String>,
104 #[serde(default, skip_serializing_if = "Option::is_none")]
109 pub origin: Option<CallOrigin>,
110}
111
112impl ToolCallRequest {
113 pub fn new(name: impl Into<String>, arguments: Value) -> Self {
116 Self {
117 name: name.into(),
118 arguments,
119 tool_call_id: None,
120 progress_token: None,
121 call_key: None,
122 schema_pin: None,
123 origin: None,
124 }
125 }
126}
127
128pub const CALL_KEY_FIELD: &str = "call_key";
131
132pub const SCHEMA_PIN_FIELD: &str = "schema_pin";
135
136pub const ORIGIN_CALL_KEY_FIELD: &str = "origin.call_key";
139
140pub const OPAQUE_FIELD_MAX_LEN: usize = 256;
143
144pub const CALL_KEY_MAX_LEN: usize = OPAQUE_FIELD_MAX_LEN;
146
147pub const SCHEMA_PIN_MAX_LEN: usize = OPAQUE_FIELD_MAX_LEN;
149
150#[derive(Clone, Debug, PartialEq, Eq)]
154pub enum OpaqueFieldError {
155 Empty { field: &'static str },
157 TooLong { field: &'static str, length: usize },
159 InvalidCharacter { field: &'static str, index: usize },
161}
162
163pub type CallKeyError = OpaqueFieldError;
166
167impl OpaqueFieldError {
168 pub fn field(&self) -> &'static str {
170 match self {
171 Self::Empty { field }
172 | Self::TooLong { field, .. }
173 | Self::InvalidCharacter { field, .. } => field,
174 }
175 }
176}
177
178impl std::fmt::Display for OpaqueFieldError {
179 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
180 match self {
181 Self::Empty { field } => write!(f, "{field} must not be empty"),
182 Self::TooLong { field, length } => write!(
183 f,
184 "{field} is {length} bytes; at most {OPAQUE_FIELD_MAX_LEN} are allowed"
185 ),
186 Self::InvalidCharacter { field, index } => write!(
187 f,
188 "{field} has a character at byte {index} outside printable ASCII \
189 (0x21 to 0x7E; space is not allowed)"
190 ),
191 }
192 }
193}
194
195impl std::error::Error for OpaqueFieldError {}
196
197fn validate_opaque_field(field: &'static str, value: &str) -> Result<(), OpaqueFieldError> {
208 if value.is_empty() {
209 return Err(OpaqueFieldError::Empty { field });
210 }
211 if value.len() > OPAQUE_FIELD_MAX_LEN {
212 return Err(OpaqueFieldError::TooLong {
213 field,
214 length: value.len(),
215 });
216 }
217 if let Some(index) = value
218 .bytes()
219 .position(|byte| !(0x21..=0x7e).contains(&byte))
220 {
221 return Err(OpaqueFieldError::InvalidCharacter { field, index });
222 }
223 Ok(())
224}
225
226pub fn validate_call_key(key: &str) -> Result<(), CallKeyError> {
229 validate_opaque_field(CALL_KEY_FIELD, key)
230}
231
232pub fn validate_schema_pin(pin: &str) -> Result<(), OpaqueFieldError> {
235 validate_opaque_field(SCHEMA_PIN_FIELD, pin)
236}
237
238pub fn validate_call_origin(origin: &CallOrigin) -> Result<(), OpaqueFieldError> {
242 validate_opaque_field(ORIGIN_CALL_KEY_FIELD, &origin.call_key)
243}
244
245#[cfg(test)]
246mod tests {
247 use super::*;
248 use serde_json::json;
249
250 #[test]
251 fn omitted_optionals_decode_as_none() {
252 let request: ToolCallRequest =
253 serde_json::from_value(json!({ "name": "grep", "arguments": { "q": "x" } }))
254 .expect("two-field body decodes");
255 assert_eq!(request.tool_call_id, None);
256 assert_eq!(request.progress_token, None);
257 assert_eq!(request.call_key, None);
258 assert_eq!(request.schema_pin, None);
259 assert_eq!(request.origin, None);
260 }
261
262 #[test]
263 fn call_key_round_trips_as_a_top_level_member() {
264 let request = ToolCallRequest {
265 name: "grep".to_string(),
266 arguments: json!({ "q": "x" }),
267 tool_call_id: None,
268 progress_token: None,
269 call_key: Some("run-7:call-3".to_string()),
270 schema_pin: None,
271 origin: None,
272 };
273 let encoded = serde_json::to_value(&request).expect("encode");
274 assert_eq!(
275 encoded,
276 json!({ "name": "grep", "arguments": { "q": "x" }, "call_key": "run-7:call-3" })
277 );
278 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
279 assert_eq!(decoded, request);
280 }
281
282 #[test]
283 fn a_request_without_a_call_key_omits_the_member_and_round_trips() {
284 let request = ToolCallRequest::new("grep", json!({}));
285 let encoded = serde_json::to_value(&request).expect("encode");
286 assert!(encoded.get("call_key").is_none(), "{encoded}");
287 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
288 assert_eq!(decoded.call_key, None);
289 assert_eq!(decoded, request);
290 }
291
292 #[test]
295 fn opaque_field_bounds_are_one_to_256_printable_non_space_ascii() {
296 type Validate = fn(&str) -> Result<(), OpaqueFieldError>;
297 let validators: [(&str, Validate); 2] = [
298 (CALL_KEY_FIELD, validate_call_key),
299 (SCHEMA_PIN_FIELD, validate_schema_pin),
300 ];
301 for (field, validate) in validators {
302 assert_eq!(validate(""), Err(OpaqueFieldError::Empty { field }));
303 assert_eq!(validate("k"), Ok(()));
304 assert_eq!(validate(&"k".repeat(256)), Ok(()));
305 assert_eq!(
306 validate(&"k".repeat(257)),
307 Err(OpaqueFieldError::TooLong { field, length: 257 })
308 );
309 assert_eq!(validate("!~"), Ok(()), "both ends of 0x21..=0x7E");
310 for bad in ["ké", "a\tb", "a\u{7f}", "a b"] {
311 assert_eq!(
312 validate(bad),
313 Err(OpaqueFieldError::InvalidCharacter { field, index: 1 }),
314 "{field}: {bad:?}"
315 );
316 }
317 let error = validate("").unwrap_err();
318 assert_eq!(error.field(), field);
319 assert!(error.to_string().starts_with(field), "{error}");
320 }
321 }
322
323 #[test]
324 fn schema_pin_round_trips_as_a_top_level_member() {
325 let request = ToolCallRequest {
326 schema_pin: Some("sha256:0f1e2d".to_string()),
327 ..ToolCallRequest::new("grep", json!({ "q": "x" }))
328 };
329 let encoded = serde_json::to_value(&request).expect("encode");
330 assert_eq!(
331 encoded,
332 json!({ "name": "grep", "arguments": { "q": "x" }, "schema_pin": "sha256:0f1e2d" })
333 );
334 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
335 assert_eq!(decoded, request);
336 }
337
338 #[test]
339 fn a_request_without_a_schema_pin_omits_the_member_and_round_trips() {
340 let request = ToolCallRequest::new("grep", json!({}));
341 let encoded = serde_json::to_value(&request).expect("encode");
342 assert!(encoded.get("schema_pin").is_none(), "{encoded}");
343 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
344 assert_eq!(decoded.schema_pin, None);
345 assert_eq!(decoded, request);
346 }
347
348 #[test]
349 fn none_optionals_are_omitted_so_the_wire_matches_the_two_field_shape() {
350 let request = ToolCallRequest::new("grep", json!({ "q": "x" }));
351 let encoded = serde_json::to_value(&request).expect("encode");
352 assert_eq!(
353 encoded,
354 json!({ "name": "grep", "arguments": { "q": "x" } })
355 );
356 }
357
358 #[test]
359 fn tool_call_id_round_trips() {
360 let request = ToolCallRequest {
361 name: "grep".to_string(),
362 arguments: json!({ "q": "x" }),
363 tool_call_id: Some("wal-intent-42".to_string()),
364 progress_token: None,
365 call_key: None,
366 schema_pin: None,
367 origin: None,
368 };
369 let encoded = serde_json::to_value(&request).expect("encode");
370 assert_eq!(encoded["tool_call_id"], json!("wal-intent-42"));
371 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
372 assert_eq!(decoded, request);
373 }
374
375 #[test]
376 fn unknown_members_do_not_fail_a_provider_decode() {
377 let request: ToolCallRequest = serde_json::from_value(json!({
380 "name": "grep",
381 "arguments": {},
382 "some_future_key": { "nested": true }
383 }))
384 .expect("unknown members are tolerated");
385 assert_eq!(request.name, "grep");
386 }
387
388 fn relayed_origin() -> CallOrigin {
389 CallOrigin::new(
390 Principal::Reserved {
391 module_id: "broca".to_string(),
392 },
393 "broca:run-7/call-3",
394 )
395 }
396
397 #[test]
400 fn origin_round_trips_as_a_top_level_member_with_a_tagged_carrier() {
401 let request = ToolCallRequest {
402 call_key: Some("pf:relay/991".to_string()),
403 origin: Some(relayed_origin()),
404 ..ToolCallRequest::new("grep", json!({ "q": "x" }))
405 };
406 let encoded = serde_json::to_value(&request).expect("encode");
407 assert_eq!(
408 encoded,
409 json!({
410 "name": "grep",
411 "arguments": { "q": "x" },
412 "call_key": "pf:relay/991",
413 "origin": {
414 "carrier": { "kind": "reserved", "module_id": "broca" },
415 "call_key": "broca:run-7/call-3"
416 }
417 })
418 );
419 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
420 assert_eq!(decoded, request);
421 }
422
423 #[test]
424 fn a_request_without_an_origin_omits_the_member_and_decodes_as_none() {
425 let request = ToolCallRequest::new("grep", json!({}));
426 let encoded = serde_json::to_value(&request).expect("encode");
427 assert!(encoded.get("origin").is_none(), "{encoded}");
428 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
429 assert_eq!(decoded.origin, None);
430 assert_eq!(decoded, request);
431 }
432
433 #[test]
434 fn a_request_with_an_origin_and_an_unknown_member_still_decodes() {
435 let request: ToolCallRequest = serde_json::from_value(json!({
436 "name": "grep",
437 "arguments": {},
438 "origin": {
439 "carrier": { "kind": "direct" },
440 "call_key": "k",
441 "some_future_origin_key": 1
442 },
443 "some_future_key": { "nested": true }
444 }))
445 .expect("unknown members are tolerated");
446 assert_eq!(
447 request.origin,
448 Some(CallOrigin::new(Principal::Direct, "k"))
449 );
450 }
451
452 #[test]
453 fn call_origin_key_refusals_name_the_origin_call_key_field() {
454 let carrier = Principal::Direct;
455 let field = ORIGIN_CALL_KEY_FIELD;
456 assert_eq!(field, "origin.call_key");
457 let cases = [
458 (String::new(), OpaqueFieldError::Empty { field }),
459 (
460 "k".repeat(257),
461 OpaqueFieldError::TooLong { field, length: 257 },
462 ),
463 (
464 "pf:relay 991".to_string(),
465 OpaqueFieldError::InvalidCharacter { field, index: 8 },
466 ),
467 ];
468 for (key, expected) in cases {
469 let error = validate_call_origin(&CallOrigin::new(carrier.clone(), key.clone()))
470 .expect_err("malformed origin key is refused");
471 assert_eq!(error, expected, "{key:?}");
472 assert_eq!(error.field(), "origin.call_key", "{key:?}");
473 }
474 }
475
476 #[test]
477 fn call_origin_accepts_every_carrier_kind() {
478 let carriers = [
479 Principal::Reserved {
480 module_id: "prefrontal-core".to_string(),
481 },
482 Principal::Direct,
483 Principal::Unverified,
484 ];
485 for carrier in carriers {
486 assert_eq!(
487 validate_call_origin(&CallOrigin::new(carrier.clone(), "pf:relay/991")),
488 Ok(()),
489 "{carrier:?}"
490 );
491 }
492 }
493}