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 #[serde(default, skip_serializing_if = "Option::is_none")]
47 pub message_position: Option<MessagePosition>,
48}
49
50impl CallOrigin {
51 pub fn new(carrier: Principal, call_key: impl Into<String>) -> Self {
54 Self {
55 carrier,
56 call_key: call_key.into(),
57 message_position: None,
58 }
59 }
60
61 pub fn with_message_position(mut self, position: MessagePosition) -> Self {
63 self.message_position = Some(position);
64 self
65 }
66}
67
68#[derive(Clone, Debug, Deserialize, PartialEq, Eq, Serialize)]
79#[non_exhaustive]
80pub struct MessagePosition {
81 pub message_id: String,
84 pub index: u32,
87}
88
89impl MessagePosition {
90 pub fn new(message_id: impl Into<String>, index: u32) -> Self {
92 Self {
93 message_id: message_id.into(),
94 index,
95 }
96 }
97}
98
99#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
101pub struct ToolCallRequest {
102 pub name: String,
104 pub arguments: Value,
108 #[serde(default, skip_serializing_if = "Option::is_none")]
123 pub tool_call_id: Option<String>,
124 #[serde(default, skip_serializing_if = "Option::is_none")]
127 pub progress_token: Option<Value>,
128 #[serde(default, skip_serializing_if = "Option::is_none")]
138 pub call_key: Option<String>,
139 #[serde(default, skip_serializing_if = "Option::is_none")]
146 pub schema_pin: Option<String>,
147 #[serde(default, skip_serializing_if = "Option::is_none")]
158 pub preset: Option<String>,
159 #[serde(default, skip_serializing_if = "Option::is_none")]
164 pub origin: Option<CallOrigin>,
165}
166
167impl ToolCallRequest {
168 pub fn new(name: impl Into<String>, arguments: Value) -> Self {
171 Self {
172 name: name.into(),
173 arguments,
174 tool_call_id: None,
175 progress_token: None,
176 call_key: None,
177 schema_pin: None,
178 preset: None,
179 origin: None,
180 }
181 }
182}
183
184pub const CALL_KEY_FIELD: &str = "call_key";
187
188pub const SCHEMA_PIN_FIELD: &str = "schema_pin";
191
192pub const PRESET_FIELD: &str = "preset";
195
196pub const PRESET_MAX_LEN: usize = 64;
198
199#[derive(Clone, Debug, PartialEq, Eq)]
202#[non_exhaustive]
203pub enum PresetError {
204 Empty,
205 TooLong { length: usize },
206 InvalidCharacter { index: usize },
207}
208
209impl PresetError {
210 pub fn field(&self) -> &'static str {
211 PRESET_FIELD
212 }
213}
214
215impl std::fmt::Display for PresetError {
216 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
217 match self {
218 Self::Empty => write!(f, "preset must not be empty"),
219 Self::TooLong { length } => write!(
220 f,
221 "preset is {length} bytes; at most {PRESET_MAX_LEN} are allowed"
222 ),
223 Self::InvalidCharacter { index } => {
224 write!(
225 f,
226 "preset has a character at byte {index} outside [a-z0-9_-]"
227 )
228 }
229 }
230 }
231}
232
233impl std::error::Error for PresetError {}
234
235pub const ORIGIN_CALL_KEY_FIELD: &str = "origin.call_key";
238
239pub const ORIGIN_MESSAGE_ID_FIELD: &str = "origin.message_position.message_id";
243
244pub const OPAQUE_FIELD_MAX_LEN: usize = 256;
247
248pub const CALL_KEY_MAX_LEN: usize = OPAQUE_FIELD_MAX_LEN;
250
251pub const SCHEMA_PIN_MAX_LEN: usize = OPAQUE_FIELD_MAX_LEN;
253
254#[derive(Clone, Debug, PartialEq, Eq)]
258pub enum OpaqueFieldError {
259 Empty { field: &'static str },
261 TooLong { field: &'static str, length: usize },
263 InvalidCharacter { field: &'static str, index: usize },
265}
266
267pub type CallKeyError = OpaqueFieldError;
270
271impl OpaqueFieldError {
272 pub fn field(&self) -> &'static str {
274 match self {
275 Self::Empty { field }
276 | Self::TooLong { field, .. }
277 | Self::InvalidCharacter { field, .. } => field,
278 }
279 }
280}
281
282impl std::fmt::Display for OpaqueFieldError {
283 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
284 match self {
285 Self::Empty { field } => write!(f, "{field} must not be empty"),
286 Self::TooLong { field, length } => write!(
287 f,
288 "{field} is {length} bytes; at most {OPAQUE_FIELD_MAX_LEN} are allowed"
289 ),
290 Self::InvalidCharacter { field, index } => write!(
291 f,
292 "{field} has a character at byte {index} outside printable ASCII \
293 (0x21 to 0x7E; space is not allowed)"
294 ),
295 }
296 }
297}
298
299impl std::error::Error for OpaqueFieldError {}
300
301pub(crate) fn validate_opaque_field(
312 field: &'static str,
313 value: &str,
314) -> Result<(), OpaqueFieldError> {
315 if value.is_empty() {
316 return Err(OpaqueFieldError::Empty { field });
317 }
318 if value.len() > OPAQUE_FIELD_MAX_LEN {
319 return Err(OpaqueFieldError::TooLong {
320 field,
321 length: value.len(),
322 });
323 }
324 if let Some(index) = value
325 .bytes()
326 .position(|byte| !(0x21..=0x7e).contains(&byte))
327 {
328 return Err(OpaqueFieldError::InvalidCharacter { field, index });
329 }
330 Ok(())
331}
332
333pub fn validate_call_key(key: &str) -> Result<(), CallKeyError> {
336 validate_opaque_field(CALL_KEY_FIELD, key)
337}
338
339pub fn validate_schema_pin(pin: &str) -> Result<(), OpaqueFieldError> {
342 validate_opaque_field(SCHEMA_PIN_FIELD, pin)
343}
344
345pub fn validate_preset(preset: &str) -> Result<(), PresetError> {
348 if preset.is_empty() {
349 return Err(PresetError::Empty);
350 }
351 if preset.len() > PRESET_MAX_LEN {
352 return Err(PresetError::TooLong {
353 length: preset.len(),
354 });
355 }
356 if let Some(index) = preset.bytes().position(|byte| {
357 !(byte.is_ascii_lowercase() || byte.is_ascii_digit() || b"_-".contains(&byte))
358 }) {
359 return Err(PresetError::InvalidCharacter { index });
360 }
361 Ok(())
362}
363
364pub fn validate_call_origin(origin: &CallOrigin) -> Result<(), OpaqueFieldError> {
369 validate_opaque_field(ORIGIN_CALL_KEY_FIELD, &origin.call_key)?;
370 if let Some(position) = &origin.message_position {
371 validate_opaque_field(ORIGIN_MESSAGE_ID_FIELD, &position.message_id)?;
372 }
373 Ok(())
374}
375
376#[cfg(test)]
377mod tests {
378 use super::*;
379 use serde_json::json;
380
381 #[test]
382 fn omitted_optionals_decode_as_none() {
383 let request: ToolCallRequest =
384 serde_json::from_value(json!({ "name": "grep", "arguments": { "q": "x" } }))
385 .expect("two-field body decodes");
386 assert_eq!(request.tool_call_id, None);
387 assert_eq!(request.progress_token, None);
388 assert_eq!(request.call_key, None);
389 assert_eq!(request.schema_pin, None);
390 assert_eq!(request.preset, None);
391 assert_eq!(request.origin, None);
392 }
393
394 #[test]
395 fn call_key_round_trips_as_a_top_level_member() {
396 let request = ToolCallRequest {
397 name: "grep".to_string(),
398 arguments: json!({ "q": "x" }),
399 tool_call_id: None,
400 progress_token: None,
401 call_key: Some("run-7:call-3".to_string()),
402 schema_pin: None,
403 preset: None,
404 origin: None,
405 };
406 let encoded = serde_json::to_value(&request).expect("encode");
407 assert_eq!(
408 encoded,
409 json!({ "name": "grep", "arguments": { "q": "x" }, "call_key": "run-7:call-3" })
410 );
411 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
412 assert_eq!(decoded, request);
413 }
414
415 #[test]
416 fn a_request_without_a_call_key_omits_the_member_and_round_trips() {
417 let request = ToolCallRequest::new("grep", json!({}));
418 let encoded = serde_json::to_value(&request).expect("encode");
419 assert!(encoded.get("call_key").is_none(), "{encoded}");
420 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
421 assert_eq!(decoded.call_key, None);
422 assert_eq!(decoded, request);
423 }
424
425 #[test]
428 fn opaque_field_bounds_are_one_to_256_printable_non_space_ascii() {
429 type Validate = fn(&str) -> Result<(), OpaqueFieldError>;
430 let validators: [(&str, Validate); 2] = [
431 (CALL_KEY_FIELD, validate_call_key),
432 (SCHEMA_PIN_FIELD, validate_schema_pin),
433 ];
434 for (field, validate) in validators {
435 assert_eq!(validate(""), Err(OpaqueFieldError::Empty { field }));
436 assert_eq!(validate("k"), Ok(()));
437 assert_eq!(validate(&"k".repeat(256)), Ok(()));
438 assert_eq!(
439 validate(&"k".repeat(257)),
440 Err(OpaqueFieldError::TooLong { field, length: 257 })
441 );
442 assert_eq!(validate("!~"), Ok(()), "both ends of 0x21..=0x7E");
443 for bad in ["ké", "a\tb", "a\u{7f}", "a b"] {
444 assert_eq!(
445 validate(bad),
446 Err(OpaqueFieldError::InvalidCharacter { field, index: 1 }),
447 "{field}: {bad:?}"
448 );
449 }
450 let error = validate("").unwrap_err();
451 assert_eq!(error.field(), field);
452 assert!(error.to_string().starts_with(field), "{error}");
453 }
454 }
455
456 #[test]
457 fn schema_pin_round_trips_as_a_top_level_member() {
458 let request = ToolCallRequest {
459 schema_pin: Some("sha256:0f1e2d".to_string()),
460 ..ToolCallRequest::new("grep", json!({ "q": "x" }))
461 };
462 let encoded = serde_json::to_value(&request).expect("encode");
463 assert_eq!(
464 encoded,
465 json!({ "name": "grep", "arguments": { "q": "x" }, "schema_pin": "sha256:0f1e2d" })
466 );
467 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
468 assert_eq!(decoded, request);
469 }
470
471 #[test]
472 fn preset_round_trips_outside_arguments_and_absence_keeps_the_bytes() {
473 let mut request = ToolCallRequest::new("grep", json!({ "q": "x" }));
474 let absent = serde_json::to_string(&request).unwrap();
475 assert_eq!(absent, r#"{"name":"grep","arguments":{"q":"x"}}"#);
476 let decoded: ToolCallRequest = serde_json::from_str(&absent).unwrap();
477 assert_eq!(decoded, request);
478 assert_eq!(decoded.preset, None);
479
480 request.preset = Some("read_only-2".to_string());
481 let encoded = serde_json::to_value(&request).unwrap();
482 assert_eq!(
483 encoded,
484 json!({
485 "name": "grep", "arguments": { "q": "x" }, "preset": "read_only-2"
486 })
487 );
488 assert_eq!(
489 serde_json::from_value::<ToolCallRequest>(encoded).unwrap(),
490 request
491 );
492 }
493
494 #[test]
495 fn preset_bounds_are_one_to_64_lowercase_digits_underscore_or_hyphen() {
496 assert_eq!(validate_preset(""), Err(PresetError::Empty));
497 assert_eq!(validate_preset("a"), Ok(()));
498 assert_eq!(validate_preset(&"a".repeat(64)), Ok(()));
499 assert_eq!(
500 validate_preset(&"a".repeat(65)),
501 Err(PresetError::TooLong { length: 65 })
502 );
503 assert_eq!(validate_preset("a0_-z9"), Ok(()));
504 for bad in ["A", ".", "é", " "] {
505 let error = validate_preset(bad).unwrap_err();
506 assert_eq!(error, PresetError::InvalidCharacter { index: 0 });
507 assert_eq!(error.field(), "preset");
508 assert!(error.to_string().starts_with("preset"));
509 let refusal =
510 crate::ErrorBody::new(crate::error_codes::INVALID_REQUEST, error.to_string())
511 .with_detail(json!({ "field": error.field() }));
512 assert_eq!(refusal.code, "invalid_request");
513 assert_eq!(refusal.detail.unwrap()["field"], "preset");
514 }
515 }
516
517 #[test]
518 fn a_request_without_a_schema_pin_omits_the_member_and_round_trips() {
519 let request = ToolCallRequest::new("grep", json!({}));
520 let encoded = serde_json::to_value(&request).expect("encode");
521 assert!(encoded.get("schema_pin").is_none(), "{encoded}");
522 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
523 assert_eq!(decoded.schema_pin, None);
524 assert_eq!(decoded, request);
525 }
526
527 #[test]
528 fn none_optionals_are_omitted_so_the_wire_matches_the_two_field_shape() {
529 let request = ToolCallRequest::new("grep", json!({ "q": "x" }));
530 let encoded = serde_json::to_value(&request).expect("encode");
531 assert_eq!(
532 encoded,
533 json!({ "name": "grep", "arguments": { "q": "x" } })
534 );
535 }
536
537 #[test]
538 fn tool_call_id_round_trips() {
539 let request = ToolCallRequest {
540 name: "grep".to_string(),
541 arguments: json!({ "q": "x" }),
542 tool_call_id: Some("wal-intent-42".to_string()),
543 progress_token: None,
544 call_key: None,
545 schema_pin: None,
546 preset: None,
547 origin: None,
548 };
549 let encoded = serde_json::to_value(&request).expect("encode");
550 assert_eq!(encoded["tool_call_id"], json!("wal-intent-42"));
551 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
552 assert_eq!(decoded, request);
553 }
554
555 #[test]
556 fn unknown_members_do_not_fail_a_provider_decode() {
557 let request: ToolCallRequest = serde_json::from_value(json!({
560 "name": "grep",
561 "arguments": {},
562 "some_future_key": { "nested": true }
563 }))
564 .expect("unknown members are tolerated");
565 assert_eq!(request.name, "grep");
566 }
567
568 fn relayed_origin() -> CallOrigin {
569 CallOrigin::new(
570 Principal::Reserved {
571 module_id: "broca".to_string(),
572 },
573 "broca:run-7/call-3",
574 )
575 }
576
577 #[test]
580 fn origin_round_trips_as_a_top_level_member_with_a_tagged_carrier() {
581 let request = ToolCallRequest {
582 call_key: Some("pf:relay/991".to_string()),
583 origin: Some(relayed_origin()),
584 ..ToolCallRequest::new("grep", json!({ "q": "x" }))
585 };
586 let encoded = serde_json::to_value(&request).expect("encode");
587 assert_eq!(
588 encoded,
589 json!({
590 "name": "grep",
591 "arguments": { "q": "x" },
592 "call_key": "pf:relay/991",
593 "origin": {
594 "carrier": { "kind": "reserved", "module_id": "broca" },
595 "call_key": "broca:run-7/call-3"
596 }
597 })
598 );
599 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
600 assert_eq!(decoded, request);
601 }
602
603 #[test]
606 fn message_position_round_trips_and_is_omitted_when_unknown() {
607 let origin = relayed_origin().with_message_position(MessagePosition::new("msg_01AbC", 2));
608 let encoded = serde_json::to_value(&origin).expect("encode");
609 assert_eq!(
610 encoded,
611 json!({
612 "carrier": { "kind": "reserved", "module_id": "broca" },
613 "call_key": "broca:run-7/call-3",
614 "message_position": { "message_id": "msg_01AbC", "index": 2 }
615 })
616 );
617 let decoded: CallOrigin = serde_json::from_value(encoded).expect("decode");
618 assert_eq!(decoded, origin);
619
620 let without = serde_json::to_value(relayed_origin()).expect("encode");
621 assert!(without.get("message_position").is_none(), "{without}");
622 let decoded: CallOrigin = serde_json::from_value(without).expect("decode");
623 assert_eq!(decoded.message_position, None);
624 }
625
626 #[test]
629 fn a_message_position_missing_either_member_does_not_decode() {
630 for position in [json!({ "index": 0 }), json!({ "message_id": "m" })] {
631 let body = json!({
632 "carrier": { "kind": "direct" },
633 "call_key": "k",
634 "message_position": position
635 });
636 assert!(
637 serde_json::from_value::<CallOrigin>(body.clone()).is_err(),
638 "{body}"
639 );
640 }
641 }
642
643 #[test]
646 fn unknown_members_inside_a_message_position_are_tolerated() {
647 let decoded: CallOrigin = serde_json::from_value(json!({
648 "carrier": { "kind": "direct" },
649 "call_key": "k",
650 "message_position": { "message_id": "m", "index": 4, "later": true }
651 }))
652 .expect("unknown members inside a position are tolerated");
653 assert_eq!(decoded.message_position, Some(MessagePosition::new("m", 4)));
654 }
655
656 #[test]
657 fn a_malformed_message_id_is_refused_by_its_field() {
658 for bad in [
659 String::new(),
660 "has space".to_string(),
661 "x".repeat(OPAQUE_FIELD_MAX_LEN + 1),
662 ] {
663 let origin = CallOrigin::new(Principal::Direct, "k")
664 .with_message_position(MessagePosition::new(bad.clone(), 0));
665 let error = validate_call_origin(&origin).expect_err(&bad);
666 assert_eq!(error.field(), ORIGIN_MESSAGE_ID_FIELD, "{bad:?}");
667 }
668 let good = CallOrigin::new(Principal::Direct, "k")
669 .with_message_position(MessagePosition::new("msg_01AbC", u32::MAX));
670 assert_eq!(validate_call_origin(&good), Ok(()));
671 }
672
673 #[test]
674 fn a_request_without_an_origin_omits_the_member_and_decodes_as_none() {
675 let request = ToolCallRequest::new("grep", json!({}));
676 let encoded = serde_json::to_value(&request).expect("encode");
677 assert!(encoded.get("origin").is_none(), "{encoded}");
678 let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
679 assert_eq!(decoded.origin, None);
680 assert_eq!(decoded, request);
681 }
682
683 #[test]
684 fn a_request_with_an_origin_and_an_unknown_member_still_decodes() {
685 let request: ToolCallRequest = serde_json::from_value(json!({
686 "name": "grep",
687 "arguments": {},
688 "origin": {
689 "carrier": { "kind": "direct" },
690 "call_key": "k",
691 "some_future_origin_key": 1
692 },
693 "some_future_key": { "nested": true }
694 }))
695 .expect("unknown members are tolerated");
696 assert_eq!(
697 request.origin,
698 Some(CallOrigin::new(Principal::Direct, "k"))
699 );
700 }
701
702 #[test]
703 fn call_origin_key_refusals_name_the_origin_call_key_field() {
704 let carrier = Principal::Direct;
705 let field = ORIGIN_CALL_KEY_FIELD;
706 assert_eq!(field, "origin.call_key");
707 let cases = [
708 (String::new(), OpaqueFieldError::Empty { field }),
709 (
710 "k".repeat(257),
711 OpaqueFieldError::TooLong { field, length: 257 },
712 ),
713 (
714 "pf:relay 991".to_string(),
715 OpaqueFieldError::InvalidCharacter { field, index: 8 },
716 ),
717 ];
718 for (key, expected) in cases {
719 let error = validate_call_origin(&CallOrigin::new(carrier.clone(), key.clone()))
720 .expect_err("malformed origin key is refused");
721 assert_eq!(error, expected, "{key:?}");
722 assert_eq!(error.field(), "origin.call_key", "{key:?}");
723 }
724 }
725
726 #[test]
727 fn call_origin_accepts_every_carrier_kind() {
728 let carriers = [
729 Principal::Reserved {
730 module_id: "prefrontal-core".to_string(),
731 },
732 Principal::Direct,
733 Principal::Unverified,
734 ];
735 for carrier in carriers {
736 assert_eq!(
737 validate_call_origin(&CallOrigin::new(carrier.clone(), "pf:relay/991")),
738 Ok(()),
739 "{carrier:?}"
740 );
741 }
742 }
743}