Skip to main content

subc_protocol/
tool_call.rs

1//! The body of a tool-call `REQUEST` frame on a bound route.
2//!
3//! The daemon splices route frames without reading their bodies, so this
4//! shape is a contract between consumers (the MCP gateway, model runners)
5//! and provider modules, not something the daemon enforces. Before this
6//! type existed every consumer carried its own struct and every provider its
7//! own reader, and the fields drifted: the gateway sent `progress_token`,
8//! a model runner sent only `name` and `arguments`, and a provider that
9//! needed the caller's tool-call id had no field to read it from.
10//!
11//! Decoding is deliberately tolerant of unknown members: a provider must
12//! never refuse a call because a newer consumer added a key it does not
13//! know. Omitted optionals decode as `None`; `None` optionals are omitted on
14//! the wire, so a body carrying neither optional serializes exactly as the
15//! two-field shape older consumers already send — this type is drop-in for
16//! them without a wire change.
17
18use serde::{Deserialize, Serialize};
19use serde_json::Value;
20
21use crate::Principal;
22
23/// Who really asked for a call that another module relays.
24///
25/// When a module forwards a tool call on behalf of a different caller, the
26/// provider sees the forwarding module as the route's principal and the
27/// forwarder's own `call_key`. This records the caller behind it, so the
28/// provider can attribute the call in its logs and ledgers.
29///
30/// It is a claim made by the forwarding module, not something the daemon
31/// checks: a provider may record it but must never grant or refuse anything
32/// because of it. Authority stays with the route's own principal.
33///
34/// `#[non_exhaustive]` so a later member is an additive change; build it with
35/// [`CallOrigin::new`].
36#[derive(Clone, Debug, Deserialize, PartialEq, Eq, Serialize)]
37#[non_exhaustive]
38pub struct CallOrigin {
39    /// The original caller, in the same form the daemon stamps on a route.
40    pub carrier: Principal,
41    /// The original caller's key for the call, with the same bounds as
42    /// [`ToolCallRequest::call_key`]; check it with [`validate_call_origin`].
43    pub call_key: String,
44}
45
46impl CallOrigin {
47    /// An origin naming `carrier` as the caller and `call_key` as its key.
48    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/// A tool invocation as carried on a route `REQUEST` frame.
57#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
58pub struct ToolCallRequest {
59    /// The provider's bare manifest tool name (no gateway prefix).
60    pub name: String,
61    /// The arguments exactly as the caller supplied them; consumers never
62    /// translate them, and the provider's manifest schema is what accepts
63    /// or rejects their shape.
64    pub arguments: Value,
65    /// The consumer's own identifier for this call, minted by whatever
66    /// dispatched it (a model runner's WAL intent id, a gateway request id).
67    /// Opaque to the daemon and to subc; unique per call on the consumer's
68    /// side, so a provider's at-most-once fence can key on it directly
69    /// instead of synthesizing an id from the call's contents.
70    ///
71    /// `None` is a statement about the PRODUCER, not the call: it means this
72    /// consumer did not supply an id, never that the call has no identity.
73    /// A reader must not collapse the two — the moment a legacy producer is
74    /// on the other end, treating `None` as "no id exists" and synthesizing
75    /// one silently reproduces exactly the failure this field exists to end.
76    /// A reader that synthesizes a fallback id when this is `None` must
77    /// record that the fallback fired (a fallback that never reports firing
78    /// is indistinguishable from a working component that is quietly wrong).
79    #[serde(default, skip_serializing_if = "Option::is_none")]
80    pub tool_call_id: Option<String>,
81    /// An MCP progress token the consumer wants progress notifications
82    /// correlated to, when the caller requested progress. Opaque here.
83    #[serde(default, skip_serializing_if = "Option::is_none")]
84    pub progress_token: Option<Value>,
85    /// A key the consumer chose for this call, which a provider may use to
86    /// recognise the same call arriving twice. Opaque to the daemon; a
87    /// provider checks only its shape, with [`validate_call_key`], and answers
88    /// a malformed one with `invalid_request` naming the field
89    /// [`CALL_KEY_FIELD`].
90    ///
91    /// This struct is deliberately not `#[non_exhaustive]`: a consumer that
92    /// builds it field by field must decide what key, if any, to send, so a
93    /// new field here is meant to stop its struct literal compiling.
94    #[serde(default, skip_serializing_if = "Option::is_none")]
95    pub call_key: Option<String>,
96    /// Which version of the tool's schema the consumer built its arguments
97    /// against, so a provider can tell a call made against a schema it no
98    /// longer serves. Opaque here: the tool-provider role defines what it
99    /// holds. A provider checks only its shape, with [`validate_schema_pin`],
100    /// and answers a malformed one with `invalid_request` naming the field
101    /// [`SCHEMA_PIN_FIELD`].
102    #[serde(default, skip_serializing_if = "Option::is_none")]
103    pub schema_pin: Option<String>,
104    /// The caller behind this call when the consumer is relaying it for
105    /// someone else; `None` when the consumer is the caller. For attribution
106    /// only, never authority: see [`CallOrigin`]. A provider checks it with
107    /// [`validate_call_origin`].
108    #[serde(default, skip_serializing_if = "Option::is_none")]
109    pub origin: Option<CallOrigin>,
110}
111
112impl ToolCallRequest {
113    /// A call with no consumer id, no progress token, no call key, no schema
114    /// pin and no origin — the shape older two-field consumers send.
115    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
128/// The wire name of [`ToolCallRequest::call_key`], for the `field` of the
129/// `invalid_request` error a provider returns when the key is malformed.
130pub const CALL_KEY_FIELD: &str = "call_key";
131
132/// The wire name of [`ToolCallRequest::schema_pin`], for the `field` of the
133/// `invalid_request` error a provider returns when the pin is malformed.
134pub const SCHEMA_PIN_FIELD: &str = "schema_pin";
135
136/// The wire path of [`CallOrigin::call_key`] inside a request, for the `field`
137/// of the `invalid_request` error a provider returns when it is malformed.
138pub const ORIGIN_CALL_KEY_FIELD: &str = "origin.call_key";
139
140/// The longest opaque token field accepted, in bytes (every accepted byte is
141/// one ASCII character). Shared by `call_key` and `schema_pin`.
142pub const OPAQUE_FIELD_MAX_LEN: usize = 256;
143
144/// The longest `call_key` accepted.
145pub const CALL_KEY_MAX_LEN: usize = OPAQUE_FIELD_MAX_LEN;
146
147/// The longest `schema_pin` accepted.
148pub const SCHEMA_PIN_MAX_LEN: usize = OPAQUE_FIELD_MAX_LEN;
149
150/// Why an opaque token field (`call_key`, `schema_pin`) was refused. Every
151/// variant names the request field it is about, so a provider can put it in
152/// its `invalid_request` error without tracking which check ran.
153#[derive(Clone, Debug, PartialEq, Eq)]
154pub enum OpaqueFieldError {
155    /// The value was the empty string. An absent value is `None`, never `""`.
156    Empty { field: &'static str },
157    /// The value was longer than [`OPAQUE_FIELD_MAX_LEN`] bytes.
158    TooLong { field: &'static str, length: usize },
159    /// The byte at `index` is not printable, non-space ASCII.
160    InvalidCharacter { field: &'static str, index: usize },
161}
162
163/// The error [`validate_call_key`] returns. Kept as a name so code that only
164/// formats the error or reads [`OpaqueFieldError::field`] keeps compiling.
165pub type CallKeyError = OpaqueFieldError;
166
167impl OpaqueFieldError {
168    /// The request field the error is about.
169    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
197/// Check an opaque token field: 1 to [`OPAQUE_FIELD_MAX_LEN`] characters, each
198/// printable ASCII from 0x21 to 0x7E. `field` is the wire name the error
199/// reports.
200///
201/// Space (0x20) is refused. Providers compare these values byte for byte and
202/// write them into logs and ledgers, where a leading or trailing space is
203/// invisible: two values that differ only by one would read as the same value
204/// and act as different ones. Every other printable character is allowed, so
205/// a consumer can use its existing ids (UUIDs, `prefix:id` forms, base64,
206/// digests) unchanged.
207fn 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
226/// Check a `call_key` with the shared opaque-field rule; errors name
227/// [`CALL_KEY_FIELD`].
228pub fn validate_call_key(key: &str) -> Result<(), CallKeyError> {
229    validate_opaque_field(CALL_KEY_FIELD, key)
230}
231
232/// Check a `schema_pin` with the shared opaque-field rule; errors name
233/// [`SCHEMA_PIN_FIELD`].
234pub fn validate_schema_pin(pin: &str) -> Result<(), OpaqueFieldError> {
235    validate_opaque_field(SCHEMA_PIN_FIELD, pin)
236}
237
238/// Check a [`CallOrigin`]: its `call_key` must pass the shared opaque-field
239/// rule, and errors name [`ORIGIN_CALL_KEY_FIELD`]. Every [`Principal`] is
240/// accepted as the carrier.
241pub 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    /// The same bounds for both fields, with each refusal naming its own
293    /// field: one validator, two names.
294    #[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        // A newer consumer added a key this provider has never heard of; the
378        // call must still decode rather than refuse.
379        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    /// The carrier travels as the tagged `Principal` object the daemon stamps
398    /// on a route, never as a bare string.
399    #[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}