use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
use crate::error::Result;
#[derive(Debug, Clone)]
pub struct FieldDescription {
pub name: String,
pub field_type: String,
pub description: String,
pub required: bool,
}
pub trait Deserializable: Serialize + for<'de> Deserialize<'de> + Send + Sync {
fn from_partial(partial: serde_json::Value) -> Result<Self>
where
Self: Sized;
fn validate_complete(&self) -> Result<()>;
fn field_descriptions() -> Vec<FieldDescription>
where
Self: Sized;
}
#[derive(Debug)]
pub struct ParsedTool {
pub tool_name: String,
pub tool_use_id: String,
pub value: serde_json::Value,
}
pub type ToolConstructor = Box<dyn Fn(serde_json::Value) -> Result<ParsedTool> + Send + Sync>;
#[derive(Clone)]
pub struct NativeToolParser {
tool_constructors: Arc<HashMap<String, Arc<ToolConstructor>>>,
}
impl NativeToolParser {
pub fn new(tool_constructors: HashMap<String, Arc<ToolConstructor>>) -> Self {
Self {
tool_constructors: Arc::new(tool_constructors),
}
}
pub fn with_tool<T: Deserializable + Serialize + 'static>(tool_name: &str) -> Self {
let mut constructors = HashMap::new();
let name = tool_name.to_string();
let constructor: ToolConstructor = Box::new(move |json: serde_json::Value| {
T::from_partial(json.clone()).map(|tool| ParsedTool {
tool_name: name.clone(),
tool_use_id: String::new(),
value: serde_json::to_value(&tool).unwrap_or(json),
})
});
constructors.insert(tool_name.to_string(), Arc::new(constructor));
Self::new(constructors)
}
pub fn parse_tool(&self, tool_name: &str, delta_json: &str, tool_id: &str) -> ParsedToolResult {
let constructor = match self.tool_constructors.get(tool_name) {
Some(ctor) => ctor,
None => {
return ParsedToolResult {
value: None,
error: Some(crate::error::Error::NonRetryable(format!(
"Tool '{}' not found in registry",
tool_name
))),
is_partial: true,
raw_output: delta_json.to_string(),
};
}
};
let partial_data = match serde_json::from_str::<serde_json::Value>(delta_json) {
Ok(data) => data,
Err(_) => {
match constructor(serde_json::json!({})) {
Ok(mut empty_tool) => {
empty_tool.tool_use_id = tool_id.to_string();
return ParsedToolResult {
value: Some(empty_tool),
error: None,
is_partial: true,
raw_output: delta_json.to_string(),
};
}
Err(e) => {
return ParsedToolResult {
value: None,
error: Some(e),
is_partial: true,
raw_output: delta_json.to_string(),
};
}
}
}
};
match constructor(partial_data) {
Ok(mut parsed_tool) => {
parsed_tool.tool_use_id = tool_id.to_string();
ParsedToolResult {
value: Some(parsed_tool),
error: None,
is_partial: true,
raw_output: delta_json.to_string(),
}
}
Err(e) => {
match constructor(serde_json::json!({})) {
Ok(mut empty_tool) => {
empty_tool.tool_use_id = tool_id.to_string();
ParsedToolResult {
value: Some(empty_tool),
error: Some(e),
is_partial: true,
raw_output: delta_json.to_string(),
}
}
Err(e2) => ParsedToolResult {
value: None,
error: Some(e2),
is_partial: true,
raw_output: delta_json.to_string(),
},
}
}
}
}
pub fn extract_tool_name(&self, tool_call_data: &serde_json::Value) -> String {
if let Some(function) = tool_call_data.get("function") {
if let Some(name) = function.get("name") {
return name.as_str().unwrap_or("").to_string();
}
}
String::new()
}
pub fn extract_tool_id(&self, tool_call_data: &serde_json::Value) -> Option<String> {
tool_call_data
.get("id")
.and_then(|v| v.as_str())
.map(|s| s.to_string())
}
pub fn extract_arguments(&self, tool_call_data: &serde_json::Value) -> String {
if let Some(function) = tool_call_data.get("function") {
if let Some(args) = function.get("arguments") {
return args.as_str().unwrap_or("{}").to_string();
}
}
"{}".to_string()
}
}
impl Default for NativeToolParser {
fn default() -> Self {
Self::new(HashMap::new())
}
}
#[derive(Debug)]
pub struct ParsedToolResult {
pub value: Option<ParsedTool>,
pub error: Option<crate::error::Error>,
pub is_partial: bool,
pub raw_output: String,
}
impl ParsedToolResult {
pub fn success(&self) -> bool {
self.value.is_some()
}
}
#[derive(Debug)]
pub struct ParserResult<T> {
pub value: Option<T>,
pub error: Option<crate::error::Error>,
pub is_partial: bool,
pub raw_output: String,
}
impl<T> ParserResult<T> {
pub fn success(&self) -> bool {
self.value.is_some()
}
}
pub struct TypeParser<T> {
_phantom: std::marker::PhantomData<T>,
}
impl<T> TypeParser<T>
where
T: for<'de> Deserialize<'de>,
{
pub fn new() -> Self {
Self {
_phantom: std::marker::PhantomData,
}
}
pub fn parse(&self, json: &str) -> Result<T> {
serde_json::from_str(json)
.map_err(|e| crate::error::Error::NonRetryable(format!("Failed to parse JSON: {}", e)))
}
pub fn parse_value(&self, value: serde_json::Value) -> Result<T> {
serde_json::from_value(value).map_err(|e| {
crate::error::Error::NonRetryable(format!("Failed to parse JSON value: {}", e))
})
}
}
impl<T> Default for TypeParser<T>
where
T: for<'de> Deserialize<'de>,
{
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Debug, Serialize, Deserialize, PartialEq)]
struct TestStruct {
#[serde(default)]
name: String,
#[serde(default)]
age: u32,
#[serde(default)]
optional: Option<String>,
}
impl Deserializable for TestStruct {
fn from_partial(partial: serde_json::Value) -> Result<Self> {
serde_json::from_value(partial)
.map_err(|e| crate::error::Error::NonRetryable(format!("Parse error: {}", e)))
}
fn validate_complete(&self) -> Result<()> {
if self.name.is_empty() {
return Err(crate::error::Error::NonRetryable(
"name is required".to_string(),
));
}
Ok(())
}
fn field_descriptions() -> Vec<FieldDescription> {
vec![
FieldDescription {
name: "name".to_string(),
field_type: "string".to_string(),
description: "Person's name".to_string(),
required: true,
},
FieldDescription {
name: "age".to_string(),
field_type: "number".to_string(),
description: "Person's age".to_string(),
required: true,
},
FieldDescription {
name: "optional".to_string(),
field_type: "string".to_string(),
description: "Optional field".to_string(),
required: false,
},
]
}
}
#[test]
fn test_native_tool_parser_new() {
let parser = NativeToolParser::new(HashMap::new());
assert!(parser.tool_constructors.is_empty());
}
#[test]
fn test_native_tool_parser_extract_tool_name() {
let parser = NativeToolParser::new(HashMap::new());
let tool_data = serde_json::json!({
"function": {"name": "GetWeather"}
});
let name = parser.extract_tool_name(&tool_data);
assert_eq!(name, "GetWeather");
}
#[test]
fn test_native_tool_parser_extract_arguments() {
let parser = NativeToolParser::new(HashMap::new());
let tool_data = serde_json::json!({
"function": {"arguments": "{\"city\":\"SF\"}"}
});
let args = parser.extract_arguments(&tool_data);
assert_eq!(args, "{\"city\":\"SF\"}");
}
#[test]
fn test_native_tool_parser_extract_tool_id() {
let parser = NativeToolParser::new(HashMap::new());
let tool_data = serde_json::json!({
"id": "tool_123"
});
let id = parser.extract_tool_id(&tool_data);
assert_eq!(id, Some("tool_123".to_string()));
}
#[test]
fn test_native_tool_parser_parse_tool_complete() {
let parser = NativeToolParser::with_tool::<TestStruct>("TestTool");
let json = r#"{"name":"Alice","age":30}"#;
let result = parser.parse_tool("TestTool", json, "tool_1");
assert!(result.success());
let parsed = result.value.unwrap();
assert_eq!(parsed.tool_name, "TestTool");
assert_eq!(parsed.tool_use_id, "tool_1");
let obj: TestStruct = serde_json::from_value(parsed.value).unwrap();
assert_eq!(obj.name, "Alice");
assert_eq!(obj.age, 30);
}
#[test]
fn test_native_tool_parser_parse_tool_incomplete() {
let parser = NativeToolParser::with_tool::<TestStruct>("TestTool");
let json = r#"{"name":"Alice","age":"#;
let result = parser.parse_tool("TestTool", json, "tool_1");
assert!(result.is_partial);
assert!(result.value.is_some());
}
#[test]
fn test_native_tool_parser_parse_tool_not_found() {
let parser = NativeToolParser::new(HashMap::new());
let json = r#"{"name":"Alice"}"#;
let result = parser.parse_tool("UnknownTool", json, "tool_1");
assert!(!result.success());
assert!(result.error.is_some());
}
#[test]
fn test_type_parser_parse() {
let parser = TypeParser::<TestStruct>::new();
let json = r#"{"name":"Bob","age":25}"#;
let result = parser.parse(json);
assert!(result.is_ok());
let obj = result.unwrap();
assert_eq!(obj.name, "Bob");
assert_eq!(obj.age, 25);
}
#[test]
fn test_type_parser_parse_value() {
let parser = TypeParser::<TestStruct>::new();
let value = serde_json::json!({"name":"Charlie","age":35});
let result = parser.parse_value(value);
assert!(result.is_ok());
let obj = result.unwrap();
assert_eq!(obj.name, "Charlie");
assert_eq!(obj.age, 35);
}
#[test]
fn test_type_parser_parse_invalid() {
let parser = TypeParser::<TestStruct>::new();
let json = r#"{"invalid":"json"#;
let result = parser.parse(json);
assert!(result.is_err());
}
#[test]
fn test_deserializable_from_partial() {
let partial = serde_json::json!({"name":"Dave","age":40});
let result = TestStruct::from_partial(partial);
assert!(result.is_ok());
let obj = result.unwrap();
assert_eq!(obj.name, "Dave");
assert_eq!(obj.age, 40);
}
#[test]
fn test_deserializable_validate_complete() {
let valid = TestStruct {
name: "Eve".to_string(),
age: 28,
optional: None,
};
assert!(valid.validate_complete().is_ok());
let invalid = TestStruct {
name: "".to_string(),
age: 28,
optional: None,
};
assert!(invalid.validate_complete().is_err());
}
#[test]
fn test_deserializable_field_descriptions() {
let descs = TestStruct::field_descriptions();
assert_eq!(descs.len(), 3);
assert_eq!(descs[0].name, "name");
assert_eq!(descs[1].name, "age");
assert_eq!(descs[2].name, "optional");
assert!(descs[0].required);
assert!(descs[1].required);
assert!(!descs[2].required);
}
#[test]
fn test_parsed_tool_metadata() {
let parser = NativeToolParser::with_tool::<TestStruct>("GetPerson");
let json = r#"{"name":"Frank","age":42}"#;
let result = parser.parse_tool("GetPerson", json, "call_abc123");
assert!(result.success());
let parsed = result.value.unwrap();
assert_eq!(parsed.tool_name, "GetPerson");
assert_eq!(parsed.tool_use_id, "call_abc123");
let obj: TestStruct = serde_json::from_value(parsed.value).unwrap();
assert_eq!(obj.name, "Frank");
assert_eq!(obj.age, 42);
}
#[test]
fn test_parser_with_multiple_tools() {
let mut constructors = HashMap::new();
let name1 = "Tool1".to_string();
let constructor1: ToolConstructor = Box::new(move |json: serde_json::Value| {
TestStruct::from_partial(json.clone()).map(|tool| ParsedTool {
tool_name: name1.clone(),
tool_use_id: String::new(),
value: serde_json::to_value(&tool).unwrap_or(json),
})
});
constructors.insert("Tool1".to_string(), Arc::new(constructor1));
let parser = NativeToolParser::new(constructors);
let result = parser.parse_tool("Tool1", r#"{"name":"George","age":50}"#, "id1");
assert!(result.success());
assert_eq!(result.value.unwrap().tool_name, "Tool1");
let result = parser.parse_tool("Tool2", r#"{}"#, "id2");
assert!(!result.success());
assert!(result.error.is_some());
}
}