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")]
115 pub preset: Option<String>,
116 #[serde(default, skip_serializing_if = "Option::is_none")]
121 pub origin: Option<CallOrigin>,
122}
123
124impl ToolCallRequest {
125 pub fn new(name: impl Into<String>, arguments: Value) -> Self {
128 Self {
129 name: name.into(),
130 arguments,
131 tool_call_id: None,
132 progress_token: None,
133 call_key: None,
134 schema_pin: None,
135 preset: None,
136 origin: None,
137 }
138 }
139}
140
141pub const CALL_KEY_FIELD: &str = "call_key";
144
145pub const SCHEMA_PIN_FIELD: &str = "schema_pin";
148
149pub const PRESET_FIELD: &str = "preset";
152
153pub const PRESET_MAX_LEN: usize = 64;
155
156#[derive(Clone, Debug, PartialEq, Eq)]
159#[non_exhaustive]
160pub enum PresetError {
161 Empty,
162 TooLong { length: usize },
163 InvalidCharacter { index: usize },
164}
165
166impl PresetError {
167 pub fn field(&self) -> &'static str {
168 PRESET_FIELD
169 }
170}
171
172impl std::fmt::Display for PresetError {
173 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
174 match self {
175 Self::Empty => write!(f, "preset must not be empty"),
176 Self::TooLong { length } => write!(
177 f,
178 "preset is {length} bytes; at most {PRESET_MAX_LEN} are allowed"
179 ),
180 Self::InvalidCharacter { index } => {
181 write!(
182 f,
183 "preset has a character at byte {index} outside [a-z0-9_-]"
184 )
185 }
186 }
187 }
188}
189
190impl std::error::Error for PresetError {}
191
192pub const ORIGIN_CALL_KEY_FIELD: &str = "origin.call_key";
195
196pub const OPAQUE_FIELD_MAX_LEN: usize = 256;
199
200pub const CALL_KEY_MAX_LEN: usize = OPAQUE_FIELD_MAX_LEN;
202
203pub const SCHEMA_PIN_MAX_LEN: usize = OPAQUE_FIELD_MAX_LEN;
205
206#[derive(Clone, Debug, PartialEq, Eq)]
210pub enum OpaqueFieldError {
211 Empty { field: &'static str },
213 TooLong { field: &'static str, length: usize },
215 InvalidCharacter { field: &'static str, index: usize },
217}
218
219pub type CallKeyError = OpaqueFieldError;
222
223impl OpaqueFieldError {
224 pub fn field(&self) -> &'static str {
226 match self {
227 Self::Empty { field }
228 | Self::TooLong { field, .. }
229 | Self::InvalidCharacter { field, .. } => field,
230 }
231 }
232}
233
234impl std::fmt::Display for OpaqueFieldError {
235 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
236 match self {
237 Self::Empty { field } => write!(f, "{field} must not be empty"),
238 Self::TooLong { field, length } => write!(
239 f,
240 "{field} is {length} bytes; at most {OPAQUE_FIELD_MAX_LEN} are allowed"
241 ),
242 Self::InvalidCharacter { field, index } => write!(
243 f,
244 "{field} has a character at byte {index} outside printable ASCII \
245 (0x21 to 0x7E; space is not allowed)"
246 ),
247 }
248 }
249}
250
251impl std::error::Error for OpaqueFieldError {}
252
253pub(crate) fn validate_opaque_field(
264 field: &'static str,
265 value: &str,
266) -> Result<(), OpaqueFieldError> {
267 if value.is_empty() {
268 return Err(OpaqueFieldError::Empty { field });
269 }
270 if value.len() > OPAQUE_FIELD_MAX_LEN {
271 return Err(OpaqueFieldError::TooLong {
272 field,
273 length: value.len(),
274 });
275 }
276 if let Some(index) = value
277 .bytes()
278 .position(|byte| !(0x21..=0x7e).contains(&byte))
279 {
280 return Err(OpaqueFieldError::InvalidCharacter { field, index });
281 }
282 Ok(())
283}
284
285pub fn validate_call_key(key: &str) -> Result<(), CallKeyError> {
288 validate_opaque_field(CALL_KEY_FIELD, key)
289}
290
291pub fn validate_schema_pin(pin: &str) -> Result<(), OpaqueFieldError> {
294 validate_opaque_field(SCHEMA_PIN_FIELD, pin)
295}
296
297pub fn validate_preset(preset: &str) -> Result<(), PresetError> {
300 if preset.is_empty() {
301 return Err(PresetError::Empty);
302 }
303 if preset.len() > PRESET_MAX_LEN {
304 return Err(PresetError::TooLong {
305 length: preset.len(),
306 });
307 }
308 if let Some(index) = preset.bytes().position(|byte| {
309 !(byte.is_ascii_lowercase() || byte.is_ascii_digit() || b"_-".contains(&byte))
310 }) {
311 return Err(PresetError::InvalidCharacter { index });
312 }
313 Ok(())
314}
315
316pub fn validate_call_origin(origin: &CallOrigin) -> Result<(), OpaqueFieldError> {
320 validate_opaque_field(ORIGIN_CALL_KEY_FIELD, &origin.call_key)
321}
322
323#[cfg(test)]
324mod tests {
325 use super::*;
326 use serde_json::json;
327
328 #[test]
329 fn omitted_optionals_decode_as_none() {
330 let request: ToolCallRequest =
331 serde_json::from_value(json!({ "name": "grep", "arguments": { "q": "x" } }))
332 .expect("two-field body decodes");
333 assert_eq!(request.tool_call_id, None);
334 assert_eq!(request.progress_token, None);
335 assert_eq!(request.call_key, None);
336 assert_eq!(request.schema_pin, None);
337 assert_eq!(request.preset, None);
338 assert_eq!(request.origin, None);
339 }
340
341 #[test]
342 fn call_key_round_trips_as_a_top_level_member() {
343 let request = ToolCallRequest {
344 name: "grep".to_string(),
345 arguments: json!({ "q": "x" }),
346 tool_call_id: None,
347 progress_token: None,
348 call_key: Some("run-7:call-3".to_string()),
349 schema_pin: None,
350 preset: None,
351 origin: None,
352 };
353 let encoded = serde_json::to_value(&request).expect("encode");
354 assert_eq!(
355 encoded,
356 json!({ "name": "grep", "arguments": { "q": "x" }, "call_key": "run-7:call-3" })
357 );
358 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
359 assert_eq!(decoded, request);
360 }
361
362 #[test]
363 fn a_request_without_a_call_key_omits_the_member_and_round_trips() {
364 let request = ToolCallRequest::new("grep", json!({}));
365 let encoded = serde_json::to_value(&request).expect("encode");
366 assert!(encoded.get("call_key").is_none(), "{encoded}");
367 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
368 assert_eq!(decoded.call_key, None);
369 assert_eq!(decoded, request);
370 }
371
372 #[test]
375 fn opaque_field_bounds_are_one_to_256_printable_non_space_ascii() {
376 type Validate = fn(&str) -> Result<(), OpaqueFieldError>;
377 let validators: [(&str, Validate); 2] = [
378 (CALL_KEY_FIELD, validate_call_key),
379 (SCHEMA_PIN_FIELD, validate_schema_pin),
380 ];
381 for (field, validate) in validators {
382 assert_eq!(validate(""), Err(OpaqueFieldError::Empty { field }));
383 assert_eq!(validate("k"), Ok(()));
384 assert_eq!(validate(&"k".repeat(256)), Ok(()));
385 assert_eq!(
386 validate(&"k".repeat(257)),
387 Err(OpaqueFieldError::TooLong { field, length: 257 })
388 );
389 assert_eq!(validate("!~"), Ok(()), "both ends of 0x21..=0x7E");
390 for bad in ["ké", "a\tb", "a\u{7f}", "a b"] {
391 assert_eq!(
392 validate(bad),
393 Err(OpaqueFieldError::InvalidCharacter { field, index: 1 }),
394 "{field}: {bad:?}"
395 );
396 }
397 let error = validate("").unwrap_err();
398 assert_eq!(error.field(), field);
399 assert!(error.to_string().starts_with(field), "{error}");
400 }
401 }
402
403 #[test]
404 fn schema_pin_round_trips_as_a_top_level_member() {
405 let request = ToolCallRequest {
406 schema_pin: Some("sha256:0f1e2d".to_string()),
407 ..ToolCallRequest::new("grep", json!({ "q": "x" }))
408 };
409 let encoded = serde_json::to_value(&request).expect("encode");
410 assert_eq!(
411 encoded,
412 json!({ "name": "grep", "arguments": { "q": "x" }, "schema_pin": "sha256:0f1e2d" })
413 );
414 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
415 assert_eq!(decoded, request);
416 }
417
418 #[test]
419 fn preset_round_trips_outside_arguments_and_absence_keeps_the_bytes() {
420 let mut request = ToolCallRequest::new("grep", json!({ "q": "x" }));
421 let absent = serde_json::to_string(&request).unwrap();
422 assert_eq!(absent, r#"{"name":"grep","arguments":{"q":"x"}}"#);
423 let decoded: ToolCallRequest = serde_json::from_str(&absent).unwrap();
424 assert_eq!(decoded, request);
425 assert_eq!(decoded.preset, None);
426
427 request.preset = Some("read_only-2".to_string());
428 let encoded = serde_json::to_value(&request).unwrap();
429 assert_eq!(
430 encoded,
431 json!({
432 "name": "grep", "arguments": { "q": "x" }, "preset": "read_only-2"
433 })
434 );
435 assert_eq!(
436 serde_json::from_value::<ToolCallRequest>(encoded).unwrap(),
437 request
438 );
439 }
440
441 #[test]
442 fn preset_bounds_are_one_to_64_lowercase_digits_underscore_or_hyphen() {
443 assert_eq!(validate_preset(""), Err(PresetError::Empty));
444 assert_eq!(validate_preset("a"), Ok(()));
445 assert_eq!(validate_preset(&"a".repeat(64)), Ok(()));
446 assert_eq!(
447 validate_preset(&"a".repeat(65)),
448 Err(PresetError::TooLong { length: 65 })
449 );
450 assert_eq!(validate_preset("a0_-z9"), Ok(()));
451 for bad in ["A", ".", "é", " "] {
452 let error = validate_preset(bad).unwrap_err();
453 assert_eq!(error, PresetError::InvalidCharacter { index: 0 });
454 assert_eq!(error.field(), "preset");
455 assert!(error.to_string().starts_with("preset"));
456 let refusal =
457 crate::ErrorBody::new(crate::error_codes::INVALID_REQUEST, error.to_string())
458 .with_detail(json!({ "field": error.field() }));
459 assert_eq!(refusal.code, "invalid_request");
460 assert_eq!(refusal.detail.unwrap()["field"], "preset");
461 }
462 }
463
464 #[test]
465 fn a_request_without_a_schema_pin_omits_the_member_and_round_trips() {
466 let request = ToolCallRequest::new("grep", json!({}));
467 let encoded = serde_json::to_value(&request).expect("encode");
468 assert!(encoded.get("schema_pin").is_none(), "{encoded}");
469 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
470 assert_eq!(decoded.schema_pin, None);
471 assert_eq!(decoded, request);
472 }
473
474 #[test]
475 fn none_optionals_are_omitted_so_the_wire_matches_the_two_field_shape() {
476 let request = ToolCallRequest::new("grep", json!({ "q": "x" }));
477 let encoded = serde_json::to_value(&request).expect("encode");
478 assert_eq!(
479 encoded,
480 json!({ "name": "grep", "arguments": { "q": "x" } })
481 );
482 }
483
484 #[test]
485 fn tool_call_id_round_trips() {
486 let request = ToolCallRequest {
487 name: "grep".to_string(),
488 arguments: json!({ "q": "x" }),
489 tool_call_id: Some("wal-intent-42".to_string()),
490 progress_token: None,
491 call_key: None,
492 schema_pin: None,
493 preset: None,
494 origin: None,
495 };
496 let encoded = serde_json::to_value(&request).expect("encode");
497 assert_eq!(encoded["tool_call_id"], json!("wal-intent-42"));
498 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
499 assert_eq!(decoded, request);
500 }
501
502 #[test]
503 fn unknown_members_do_not_fail_a_provider_decode() {
504 let request: ToolCallRequest = serde_json::from_value(json!({
507 "name": "grep",
508 "arguments": {},
509 "some_future_key": { "nested": true }
510 }))
511 .expect("unknown members are tolerated");
512 assert_eq!(request.name, "grep");
513 }
514
515 fn relayed_origin() -> CallOrigin {
516 CallOrigin::new(
517 Principal::Reserved {
518 module_id: "broca".to_string(),
519 },
520 "broca:run-7/call-3",
521 )
522 }
523
524 #[test]
527 fn origin_round_trips_as_a_top_level_member_with_a_tagged_carrier() {
528 let request = ToolCallRequest {
529 call_key: Some("pf:relay/991".to_string()),
530 origin: Some(relayed_origin()),
531 ..ToolCallRequest::new("grep", json!({ "q": "x" }))
532 };
533 let encoded = serde_json::to_value(&request).expect("encode");
534 assert_eq!(
535 encoded,
536 json!({
537 "name": "grep",
538 "arguments": { "q": "x" },
539 "call_key": "pf:relay/991",
540 "origin": {
541 "carrier": { "kind": "reserved", "module_id": "broca" },
542 "call_key": "broca:run-7/call-3"
543 }
544 })
545 );
546 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
547 assert_eq!(decoded, request);
548 }
549
550 #[test]
551 fn a_request_without_an_origin_omits_the_member_and_decodes_as_none() {
552 let request = ToolCallRequest::new("grep", json!({}));
553 let encoded = serde_json::to_value(&request).expect("encode");
554 assert!(encoded.get("origin").is_none(), "{encoded}");
555 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
556 assert_eq!(decoded.origin, None);
557 assert_eq!(decoded, request);
558 }
559
560 #[test]
561 fn a_request_with_an_origin_and_an_unknown_member_still_decodes() {
562 let request: ToolCallRequest = serde_json::from_value(json!({
563 "name": "grep",
564 "arguments": {},
565 "origin": {
566 "carrier": { "kind": "direct" },
567 "call_key": "k",
568 "some_future_origin_key": 1
569 },
570 "some_future_key": { "nested": true }
571 }))
572 .expect("unknown members are tolerated");
573 assert_eq!(
574 request.origin,
575 Some(CallOrigin::new(Principal::Direct, "k"))
576 );
577 }
578
579 #[test]
580 fn call_origin_key_refusals_name_the_origin_call_key_field() {
581 let carrier = Principal::Direct;
582 let field = ORIGIN_CALL_KEY_FIELD;
583 assert_eq!(field, "origin.call_key");
584 let cases = [
585 (String::new(), OpaqueFieldError::Empty { field }),
586 (
587 "k".repeat(257),
588 OpaqueFieldError::TooLong { field, length: 257 },
589 ),
590 (
591 "pf:relay 991".to_string(),
592 OpaqueFieldError::InvalidCharacter { field, index: 8 },
593 ),
594 ];
595 for (key, expected) in cases {
596 let error = validate_call_origin(&CallOrigin::new(carrier.clone(), key.clone()))
597 .expect_err("malformed origin key is refused");
598 assert_eq!(error, expected, "{key:?}");
599 assert_eq!(error.field(), "origin.call_key", "{key:?}");
600 }
601 }
602
603 #[test]
604 fn call_origin_accepts_every_carrier_kind() {
605 let carriers = [
606 Principal::Reserved {
607 module_id: "prefrontal-core".to_string(),
608 },
609 Principal::Direct,
610 Principal::Unverified,
611 ];
612 for carrier in carriers {
613 assert_eq!(
614 validate_call_origin(&CallOrigin::new(carrier.clone(), "pf:relay/991")),
615 Ok(()),
616 "{carrier:?}"
617 );
618 }
619 }
620}