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>,
}
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,
}
}
}
pub const CALL_KEY_FIELD: &str = "call_key";
pub const CALL_KEY_MAX_LEN: usize = 256;
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum CallKeyError {
Empty,
TooLong { length: usize },
InvalidCharacter { index: usize },
}
impl CallKeyError {
pub fn field(&self) -> &'static str {
CALL_KEY_FIELD
}
}
impl std::fmt::Display for CallKeyError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Empty => write!(f, "{CALL_KEY_FIELD} must not be empty"),
Self::TooLong { length } => write!(
f,
"{CALL_KEY_FIELD} is {length} bytes; at most {CALL_KEY_MAX_LEN} are allowed"
),
Self::InvalidCharacter { index } => write!(
f,
"{CALL_KEY_FIELD} has a character at byte {index} outside printable ASCII \
(0x21 to 0x7E; space is not allowed)"
),
}
}
}
impl std::error::Error for CallKeyError {}
pub fn validate_call_key(key: &str) -> Result<(), CallKeyError> {
if key.is_empty() {
return Err(CallKeyError::Empty);
}
if key.len() > CALL_KEY_MAX_LEN {
return Err(CallKeyError::TooLong { length: key.len() });
}
if let Some(index) = key.bytes().position(|byte| !(0x21..=0x7e).contains(&byte)) {
return Err(CallKeyError::InvalidCharacter { index });
}
Ok(())
}
#[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);
}
#[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()),
};
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 call_key_bounds_are_one_to_256_printable_non_space_ascii() {
assert_eq!(validate_call_key(""), Err(CallKeyError::Empty));
assert_eq!(validate_call_key("k"), Ok(()));
assert_eq!(validate_call_key(&"k".repeat(256)), Ok(()));
assert_eq!(
validate_call_key(&"k".repeat(257)),
Err(CallKeyError::TooLong { length: 257 })
);
assert_eq!(validate_call_key("!~"), Ok(()), "both ends of 0x21..=0x7E");
assert_eq!(
validate_call_key("ké"),
Err(CallKeyError::InvalidCharacter { index: 1 })
);
assert_eq!(
validate_call_key("a\tb"),
Err(CallKeyError::InvalidCharacter { index: 1 })
);
assert_eq!(
validate_call_key("a\u{7f}"),
Err(CallKeyError::InvalidCharacter { index: 1 })
);
assert_eq!(
validate_call_key("a b"),
Err(CallKeyError::InvalidCharacter { index: 1 })
);
assert_eq!(CallKeyError::Empty.field(), "call_key");
}
#[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,
};
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");
}
}