use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use serde_json::Value as JsonValue;
use super::metadata::EventMetadata;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ProgressEvent {
#[serde(flatten)]
pub metadata: EventMetadata,
#[serde(rename = "toolUseID")]
pub tool_use_id: Option<String>,
#[serde(rename = "parentToolUseID")]
pub parent_tool_use_id: Option<String>,
pub data: ProgressData,
}
impl ProgressEvent {
pub fn uuid(&self) -> &str {
&self.metadata.uuid
}
pub fn parent_uuid(&self) -> Option<&str> {
self.metadata.parent_uuid.as_deref()
}
pub fn timestamp(&self) -> DateTime<Utc> {
self.metadata.timestamp
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum ProgressData {
BashProgress(BashProgressData),
HookProgress(HookProgressData),
AgentProgress(AgentProgressData),
QueryUpdate(QueryUpdateData),
SearchResultsReceived(SearchResultsData),
WaitingForTask(WaitingForTaskData),
#[serde(other)]
Unknown,
}
impl ProgressData {
pub fn agent_id(&self) -> Option<&str> {
match self {
Self::AgentProgress(data) => Some(&data.agent_id),
_ => None,
}
}
pub fn agent_prompt(&self) -> Option<&str> {
match self {
Self::AgentProgress(data) => Some(&data.prompt),
_ => None,
}
}
pub fn normalized_messages(&self) -> Option<&[NormalizedMessage]> {
match self {
Self::BashProgress(data) => Some(&data.normalized_messages),
Self::HookProgress(data) => Some(&data.normalized_messages),
Self::AgentProgress(data) => Some(&data.normalized_messages),
_ => None,
}
}
pub fn as_bash_progress(&self) -> Option<&BashProgressData> {
match self {
Self::BashProgress(data) => Some(data),
_ => None,
}
}
pub fn as_agent_progress(&self) -> Option<&AgentProgressData> {
match self {
Self::AgentProgress(data) => Some(data),
_ => None,
}
}
pub fn as_hook_progress(&self) -> Option<&HookProgressData> {
match self {
Self::HookProgress(data) => Some(data),
_ => None,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BashProgressData {
pub output: String,
#[serde(rename = "fullOutput")]
pub full_output: String,
#[serde(rename = "elapsedTimeSeconds")]
pub elapsed_time_seconds: u64,
#[serde(rename = "totalLines")]
pub total_lines: u64,
pub message: JsonValue,
#[serde(rename = "normalizedMessages")]
pub normalized_messages: Vec<NormalizedMessage>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct HookProgressData {
#[serde(rename = "hookEvent")]
pub hook_event: String,
#[serde(rename = "hookName")]
pub hook_name: String,
pub command: String,
pub message: JsonValue,
#[serde(rename = "normalizedMessages")]
pub normalized_messages: Vec<NormalizedMessage>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AgentProgressData {
#[serde(rename = "agentId")]
pub agent_id: String,
pub prompt: String,
pub message: JsonValue,
#[serde(rename = "normalizedMessages")]
pub normalized_messages: Vec<NormalizedMessage>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct QueryUpdateData {
pub query: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SearchResultsData {
#[serde(flatten)]
pub results: JsonValue,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WaitingForTaskData {
#[serde(rename = "taskDescription")]
pub task_description: String,
#[serde(rename = "taskType")]
pub task_type: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "lowercase")]
pub enum NormalizedMessage {
#[serde(rename = "user")]
User(NormalizedUserMessage),
#[serde(rename = "assistant")]
Assistant(NormalizedAssistantMessage),
#[serde(rename = "progress")]
Progress(JsonValue),
#[serde(rename = "attachment")]
Attachment(NormalizedAttachment),
#[serde(rename = "system")]
System(JsonValue),
#[serde(other)]
Unknown,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NormalizedUserMessage {
pub uuid: String,
#[serde(rename = "parentUuid")]
pub parent_uuid: Option<String>,
pub timestamp: DateTime<Utc>,
pub message: super::message::MessageContent,
#[serde(rename = "userType")]
pub user_type: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NormalizedAssistantMessage {
pub uuid: String,
#[serde(rename = "parentUuid")]
pub parent_uuid: Option<String>,
pub timestamp: DateTime<Utc>,
pub message: JsonValue,
pub model: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NormalizedAttachment {
#[serde(flatten)]
pub attachment_type: super::attachment::AttachmentType,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse_bash_progress() {
let json = r#"{
"type": "bash_progress",
"output": "Compiling...",
"fullOutput": "Building...\nCompiling...",
"elapsedTimeSeconds": 5,
"totalLines": 2,
"message": {},
"normalizedMessages": []
}"#;
let data: ProgressData = serde_json::from_str(json).unwrap();
assert!(matches!(data, ProgressData::BashProgress(_)));
if let ProgressData::BashProgress(bash) = data {
assert_eq!(bash.output, "Compiling...");
assert_eq!(bash.full_output, "Building...\nCompiling...");
assert_eq!(bash.elapsed_time_seconds, 5);
assert_eq!(bash.total_lines, 2);
}
}
#[test]
fn test_parse_agent_progress() {
let json = r#"{
"type": "agent_progress",
"agentId": "abc123",
"prompt": "Implement feature X",
"message": {},
"normalizedMessages": []
}"#;
let data: ProgressData = serde_json::from_str(json).unwrap();
assert!(matches!(data, ProgressData::AgentProgress(_)));
assert_eq!(data.agent_id(), Some("abc123"));
assert_eq!(data.agent_prompt(), Some("Implement feature X"));
}
#[test]
fn test_parse_hook_progress() {
let json = r#"{
"type": "hook_progress",
"hookEvent": "pre-tool-use",
"hookName": "pre-commit",
"command": "./hooks/pre-commit.sh",
"message": {},
"normalizedMessages": []
}"#;
let data: ProgressData = serde_json::from_str(json).unwrap();
assert!(matches!(data, ProgressData::HookProgress(_)));
if let ProgressData::HookProgress(hook) = data {
assert_eq!(hook.hook_event, "pre-tool-use");
assert_eq!(hook.hook_name, "pre-commit");
assert_eq!(hook.command, "./hooks/pre-commit.sh");
}
}
#[test]
fn test_parse_query_update() {
let json = r#"{
"type": "query_update",
"query": "rust async best practices"
}"#;
let data: ProgressData = serde_json::from_str(json).unwrap();
assert!(matches!(data, ProgressData::QueryUpdate(_)));
if let ProgressData::QueryUpdate(query) = data {
assert_eq!(query.query, "rust async best practices");
}
}
#[test]
fn test_parse_waiting_for_task() {
let json = r#"{
"type": "waiting_for_task",
"taskDescription": "Write Bitstamp tests Phase 3",
"taskType": "local_agent"
}"#;
let data: ProgressData = serde_json::from_str(json).unwrap();
assert!(matches!(data, ProgressData::WaitingForTask(_)));
if let ProgressData::WaitingForTask(task) = data {
assert_eq!(task.task_description, "Write Bitstamp tests Phase 3");
assert_eq!(task.task_type, "local_agent");
}
}
#[test]
fn test_normalized_messages_access() {
let json = r#"{
"type": "bash_progress",
"output": "test",
"fullOutput": "test",
"elapsedTimeSeconds": 1,
"totalLines": 1,
"message": {},
"normalizedMessages": []
}"#;
let data: ProgressData = serde_json::from_str(json).unwrap();
assert!(data.normalized_messages().is_some());
assert_eq!(data.normalized_messages().unwrap().len(), 0);
}
}