use serde::{Deserialize, Serialize};
use crate::tool::ToolProvenance;
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
#[serde(rename_all = "snake_case")]
pub enum SystemCacheType {
Ephemeral,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct SystemCacheMarker {
pub offset: usize,
pub length: usize,
pub cache_type: SystemCacheType,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct ToolCall {
pub id: String,
pub name: String,
pub input: serde_json::Value,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct MessageToolResult {
pub tool_use_id: String,
pub content: String,
pub is_error: bool,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
#[serde(tag = "role", rename_all = "snake_case")]
pub enum Message {
System {
content: String,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
cache_markers: Vec<SystemCacheMarker>,
},
Context {
content: String,
},
User {
content: String,
},
Assistant {
content: String,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
tool_calls: Vec<ToolCall>,
},
Tool {
result: MessageToolResult,
},
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
#[serde(rename_all = "snake_case")]
pub enum StopReason {
EndTurn,
ToolUse,
MaxTokens,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq)]
pub struct Usage {
pub input_tokens: u32,
pub output_tokens: u32,
#[serde(default)]
pub cache_read_tokens: u32,
#[serde(default)]
pub cache_write_tokens: u32,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
#[serde(tag = "type", rename_all = "snake_case")]
#[non_exhaustive]
pub enum AgentEvent {
TurnStart,
TextDelta {
delta: String,
},
ToolCallStarted {
name: String,
},
ToolCall {
call: ToolCall,
provenance: ToolProvenance,
summary_fields: Vec<String>,
},
ToolResult {
result: MessageToolResult,
},
TurnEnd {
stop_reason: StopReason,
usage: Usage,
},
Error {
message: String,
},
}
#[cfg(test)]
#[allow(warnings)]
#[allow(warnings)]
#[allow(warnings)]
#[allow(warnings)]
mod tests {
use super::*;
#[test]
fn message_roundtrip_user() {
let msg = Message::User {
content: "Hello, world!".into(),
};
let json = serde_json::to_string(&msg).unwrap();
let decoded: Message = serde_json::from_str(&json).unwrap();
assert_eq!(msg, decoded);
}
#[test]
fn message_roundtrip_assistant() {
let msg = Message::Assistant {
content: "I can help with that.".into(),
tool_calls: vec![ToolCall {
id: "call_1".into(),
name: "read_file".into(),
input: serde_json::json!({"path": "src/main.rs"}),
}],
};
let json = serde_json::to_string(&msg).unwrap();
let decoded: Message = serde_json::from_str(&json).unwrap();
assert_eq!(msg, decoded);
}
#[test]
fn message_roundtrip_tool() {
let msg = Message::Tool {
result: MessageToolResult {
tool_use_id: "call_1".into(),
content: "fn main() {}".into(),
is_error: false,
},
};
let json = serde_json::to_string(&msg).unwrap();
let decoded: Message = serde_json::from_str(&json).unwrap();
assert_eq!(msg, decoded);
}
#[test]
fn event_roundtrip() {
let events = vec![
AgentEvent::TurnStart,
AgentEvent::TextDelta {
delta: "Hello".into(),
},
AgentEvent::ToolCall {
call: ToolCall {
id: "c1".into(),
name: "bash".into(),
input: serde_json::json!({"command": "ls"}),
},
provenance: ToolProvenance::Native,
summary_fields: vec![],
},
AgentEvent::ToolResult {
result: MessageToolResult {
tool_use_id: "c1".into(),
content: "file.rs".into(),
is_error: false,
},
},
AgentEvent::TurnEnd {
stop_reason: StopReason::EndTurn,
usage: Usage {
input_tokens: 100,
output_tokens: 50,
cache_read_tokens: 80,
cache_write_tokens: 20,
},
},
AgentEvent::Error {
message: "something failed".into(),
},
];
for event in events {
let json = serde_json::to_string(&event).unwrap();
let decoded: AgentEvent = serde_json::from_str(&json).unwrap();
assert_eq!(event, decoded);
}
}
}
pub fn extract_tool_calls_from_text(text: &str) -> Vec<ToolCall> {
let mut calls = Vec::new();
let mut remaining = text;
while let Some(start) = remaining.find("```json-tool") {
let inner_start = start + "```json-tool".len();
let inner = remaining[inner_start..].trim_start();
let end = inner.find("```").unwrap_or(inner.len());
let content = inner[..end].trim();
if let Ok(obj) = serde_json::from_str::<serde_json::Value>(content)
&& let (Some(name), Some(args)) = (obj["name"].as_str(), Some(obj["args"].clone()))
{
calls.push(ToolCall {
id: format!("tc_{}", calls.len()),
name: name.to_string(),
input: args,
});
}
remaining = &inner[end..];
if end + 3 < remaining.len() {
remaining = &remaining[3..];
} else {
break;
}
}
calls
}
pub fn strip_tool_syntax(text: &str) -> String {
let mut result = text.to_string();
while let Some(start) = result.find("```json-tool") {
let inner_start = start + "```json-tool".len();
let inner = &result[inner_start..];
let end = inner_start + inner.find("```").unwrap_or(inner.len()) + 3;
result.replace_range(start..end, "");
}
result.trim().to_string()
}