use serde::{Deserialize, Serialize};
use serde_json::Value;
#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
pub struct ToolCallRequest {
pub name: String,
pub arguments: Value,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tool_call_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub progress_token: Option<Value>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub call_key: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub schema_pin: Option<String>,
}
impl ToolCallRequest {
pub fn new(name: impl Into<String>, arguments: Value) -> Self {
Self {
name: name.into(),
arguments,
tool_call_id: None,
progress_token: None,
call_key: None,
schema_pin: None,
}
}
}
pub const CALL_KEY_FIELD: &str = "call_key";
pub const SCHEMA_PIN_FIELD: &str = "schema_pin";
pub const OPAQUE_FIELD_MAX_LEN: usize = 256;
pub const CALL_KEY_MAX_LEN: usize = OPAQUE_FIELD_MAX_LEN;
pub const SCHEMA_PIN_MAX_LEN: usize = OPAQUE_FIELD_MAX_LEN;
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum OpaqueFieldError {
Empty { field: &'static str },
TooLong { field: &'static str, length: usize },
InvalidCharacter { field: &'static str, index: usize },
}
pub type CallKeyError = OpaqueFieldError;
impl OpaqueFieldError {
pub fn field(&self) -> &'static str {
match self {
Self::Empty { field }
| Self::TooLong { field, .. }
| Self::InvalidCharacter { field, .. } => field,
}
}
}
impl std::fmt::Display for OpaqueFieldError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Empty { field } => write!(f, "{field} must not be empty"),
Self::TooLong { field, length } => write!(
f,
"{field} is {length} bytes; at most {OPAQUE_FIELD_MAX_LEN} are allowed"
),
Self::InvalidCharacter { field, index } => write!(
f,
"{field} has a character at byte {index} outside printable ASCII \
(0x21 to 0x7E; space is not allowed)"
),
}
}
}
impl std::error::Error for OpaqueFieldError {}
fn validate_opaque_field(field: &'static str, value: &str) -> Result<(), OpaqueFieldError> {
if value.is_empty() {
return Err(OpaqueFieldError::Empty { field });
}
if value.len() > OPAQUE_FIELD_MAX_LEN {
return Err(OpaqueFieldError::TooLong {
field,
length: value.len(),
});
}
if let Some(index) = value
.bytes()
.position(|byte| !(0x21..=0x7e).contains(&byte))
{
return Err(OpaqueFieldError::InvalidCharacter { field, index });
}
Ok(())
}
pub fn validate_call_key(key: &str) -> Result<(), CallKeyError> {
validate_opaque_field(CALL_KEY_FIELD, key)
}
pub fn validate_schema_pin(pin: &str) -> Result<(), OpaqueFieldError> {
validate_opaque_field(SCHEMA_PIN_FIELD, pin)
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn omitted_optionals_decode_as_none() {
let request: ToolCallRequest =
serde_json::from_value(json!({ "name": "grep", "arguments": { "q": "x" } }))
.expect("two-field body decodes");
assert_eq!(request.tool_call_id, None);
assert_eq!(request.progress_token, None);
assert_eq!(request.call_key, None);
assert_eq!(request.schema_pin, None);
}
#[test]
fn call_key_round_trips_as_a_top_level_member() {
let request = ToolCallRequest {
name: "grep".to_string(),
arguments: json!({ "q": "x" }),
tool_call_id: None,
progress_token: None,
call_key: Some("run-7:call-3".to_string()),
schema_pin: None,
};
let encoded = serde_json::to_value(&request).expect("encode");
assert_eq!(
encoded,
json!({ "name": "grep", "arguments": { "q": "x" }, "call_key": "run-7:call-3" })
);
let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
assert_eq!(decoded, request);
}
#[test]
fn a_request_without_a_call_key_omits_the_member_and_round_trips() {
let request = ToolCallRequest::new("grep", json!({}));
let encoded = serde_json::to_value(&request).expect("encode");
assert!(encoded.get("call_key").is_none(), "{encoded}");
let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
assert_eq!(decoded.call_key, None);
assert_eq!(decoded, request);
}
#[test]
fn opaque_field_bounds_are_one_to_256_printable_non_space_ascii() {
type Validate = fn(&str) -> Result<(), OpaqueFieldError>;
let validators: [(&str, Validate); 2] = [
(CALL_KEY_FIELD, validate_call_key),
(SCHEMA_PIN_FIELD, validate_schema_pin),
];
for (field, validate) in validators {
assert_eq!(validate(""), Err(OpaqueFieldError::Empty { field }));
assert_eq!(validate("k"), Ok(()));
assert_eq!(validate(&"k".repeat(256)), Ok(()));
assert_eq!(
validate(&"k".repeat(257)),
Err(OpaqueFieldError::TooLong { field, length: 257 })
);
assert_eq!(validate("!~"), Ok(()), "both ends of 0x21..=0x7E");
for bad in ["ké", "a\tb", "a\u{7f}", "a b"] {
assert_eq!(
validate(bad),
Err(OpaqueFieldError::InvalidCharacter { field, index: 1 }),
"{field}: {bad:?}"
);
}
let error = validate("").unwrap_err();
assert_eq!(error.field(), field);
assert!(error.to_string().starts_with(field), "{error}");
}
}
#[test]
fn schema_pin_round_trips_as_a_top_level_member() {
let request = ToolCallRequest {
schema_pin: Some("sha256:0f1e2d".to_string()),
..ToolCallRequest::new("grep", json!({ "q": "x" }))
};
let encoded = serde_json::to_value(&request).expect("encode");
assert_eq!(
encoded,
json!({ "name": "grep", "arguments": { "q": "x" }, "schema_pin": "sha256:0f1e2d" })
);
let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
assert_eq!(decoded, request);
}
#[test]
fn a_request_without_a_schema_pin_omits_the_member_and_round_trips() {
let request = ToolCallRequest::new("grep", json!({}));
let encoded = serde_json::to_value(&request).expect("encode");
assert!(encoded.get("schema_pin").is_none(), "{encoded}");
let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
assert_eq!(decoded.schema_pin, None);
assert_eq!(decoded, request);
}
#[test]
fn none_optionals_are_omitted_so_the_wire_matches_the_two_field_shape() {
let request = ToolCallRequest::new("grep", json!({ "q": "x" }));
let encoded = serde_json::to_value(&request).expect("encode");
assert_eq!(
encoded,
json!({ "name": "grep", "arguments": { "q": "x" } })
);
}
#[test]
fn tool_call_id_round_trips() {
let request = ToolCallRequest {
name: "grep".to_string(),
arguments: json!({ "q": "x" }),
tool_call_id: Some("wal-intent-42".to_string()),
progress_token: None,
call_key: None,
schema_pin: None,
};
let encoded = serde_json::to_value(&request).expect("encode");
assert_eq!(encoded["tool_call_id"], json!("wal-intent-42"));
let decoded: ToolCallRequest = serde_json::from_value(encoded).expect("decode");
assert_eq!(decoded, request);
}
#[test]
fn unknown_members_do_not_fail_a_provider_decode() {
let request: ToolCallRequest = serde_json::from_value(json!({
"name": "grep",
"arguments": {},
"some_future_key": { "nested": true }
}))
.expect("unknown members are tolerated");
assert_eq!(request.name, "grep");
}
}