diff --git a/crates/agent/src/cli.rs b/crates/agent/src/cli.rs index db13d0ecb..7980e5572 100644 --- a/crates/agent/src/cli.rs +++ b/crates/agent/src/cli.rs @@ -1,5 +1,5 @@ use crate::{ - AnthropicProfile, EventData, EventKind, GeminiProfile, LocalExecutionEnvironment, OpenAiProfile, + AgentEvent, AnthropicProfile, GeminiProfile, LocalExecutionEnvironment, OpenAiProfile, ProviderProfile, Session, SessionConfig, ToolApprovalFn, Turn, }; use clap::{Parser, ValueEnum}; @@ -348,8 +348,8 @@ pub async fn run() -> anyhow::Result<()> { tokio::spawn(async move { let s = styles; while let Ok(event) = rx.recv().await { - match (&event.kind, &event.data) { - (EventKind::ToolCallStart, EventData::ToolCall { tool_name, arguments, .. }) => { + match &event.event { + AgentEvent::ToolCallStarted { tool_name, arguments, .. } => { eprintln!( " {dim}\u{25cf}{reset} {bold}{cyan}{tool_name}{reset}{dim}({args}){reset}", dim = s.dim, @@ -359,12 +359,9 @@ pub async fn run() -> anyhow::Result<()> { args = format_tool_args(arguments, &cwd_str), ); } - ( - EventKind::ToolCallEnd, - EventData::ToolCallEnd { - tool_name, output, is_error, .. - }, - ) if verbose => { + AgentEvent::ToolCallCompleted { + tool_name, output, is_error, .. + } if verbose => { let label = if *is_error { "tool error" } else { "tool result" }; eprintln!( " {}[{label}] {tool_name}:{}\n{}", @@ -374,7 +371,7 @@ pub async fn run() -> anyhow::Result<()> { .unwrap_or_else(|_| output.to_string()), ); } - (EventKind::Error, EventData::Error { error }) => { + AgentEvent::Error { error } => { eprintln!( " {red}\u{2717} {error}{reset}", red = s.red, diff --git a/crates/agent/src/event.rs b/crates/agent/src/event.rs index 625498339..fd4a73722 100644 --- a/crates/agent/src/event.rs +++ b/crates/agent/src/event.rs @@ -1,4 +1,4 @@ -use crate::types::{EventData, EventKind, SessionEvent}; +use crate::types::{AgentEvent, SessionEvent}; use std::time::SystemTime; use tokio::sync::broadcast; @@ -14,15 +14,14 @@ impl EventEmitter { Self { sender } } - pub fn emit(&self, kind: EventKind, session_id: String, data: EventData) { - let event = SessionEvent { - kind, + pub fn emit(&self, session_id: String, event: AgentEvent) { + let wrapped = SessionEvent { + event, timestamp: SystemTime::now(), session_id, - data, }; // Ignore send error (no receivers) - let _ = self.sender.send(event); + let _ = self.sender.send(wrapped); } #[must_use] @@ -46,12 +45,11 @@ mod tests { let emitter = EventEmitter::new(); let mut receiver = emitter.subscribe(); - emitter.emit(EventKind::SessionStart, "sess-1".into(), EventData::Empty); + emitter.emit("sess-1".into(), AgentEvent::SessionStarted); let event = receiver.recv().await.unwrap(); - assert_eq!(event.kind, EventKind::SessionStart); + assert!(matches!(event.event, AgentEvent::SessionStarted)); assert_eq!(event.session_id, "sess-1"); - assert!(matches!(event.data, EventData::Empty)); } #[tokio::test] @@ -60,17 +58,15 @@ mod tests { let mut receiver = emitter.subscribe(); emitter.emit( - EventKind::Error, "sess-2".into(), - EventData::Error { + AgentEvent::Error { error: "something went wrong".into(), }, ); let event = receiver.recv().await.unwrap(); - assert_eq!(event.kind, EventKind::Error); assert!( - matches!(&event.data, EventData::Error { error } if error == "something went wrong") + matches!(&event.event, AgentEvent::Error { error } if error == "something went wrong") ); } @@ -80,12 +76,12 @@ mod tests { let mut rx1 = emitter.subscribe(); let mut rx2 = emitter.subscribe(); - emitter.emit(EventKind::SessionEnd, "sess-3".into(), EventData::Empty); + emitter.emit("sess-3".into(), AgentEvent::SessionEnded); let e1 = rx1.recv().await.unwrap(); let e2 = rx2.recv().await.unwrap(); - assert_eq!(e1.kind, EventKind::SessionEnd); - assert_eq!(e2.kind, EventKind::SessionEnd); + assert!(matches!(e1.event, AgentEvent::SessionEnded)); + assert!(matches!(e2.event, AgentEvent::SessionEnded)); assert_eq!(e1.session_id, "sess-3"); assert_eq!(e2.session_id, "sess-3"); } @@ -94,9 +90,8 @@ mod tests { fn emit_without_subscribers_does_not_panic() { let emitter = EventEmitter::new(); emitter.emit( - EventKind::Error, "sess-4".into(), - EventData::Error { + AgentEvent::Error { error: "test".into(), }, ); diff --git a/crates/agent/src/lib.rs b/crates/agent/src/lib.rs index 395ea2d06..dae6b8586 100644 --- a/crates/agent/src/lib.rs +++ b/crates/agent/src/lib.rs @@ -39,7 +39,7 @@ pub use tools::{ make_shell_tool_with_config, make_write_file_tool, }; pub use truncation::{truncate_lines, truncate_output, truncate_tool_output, TruncationMode}; -pub use types::{EventData, EventKind, SessionEvent, SessionState, Turn}; +pub use types::{AgentEvent, SessionEvent, SessionState, Turn}; #[cfg(test)] pub(crate) mod test_support; diff --git a/crates/agent/src/session.rs b/crates/agent/src/session.rs index 353600cc6..d2a5ef346 100644 --- a/crates/agent/src/session.rs +++ b/crates/agent/src/session.rs @@ -9,7 +9,7 @@ use crate::project_docs::discover_project_docs; use crate::provider_profile::ProviderProfile; use crate::tool_registry::ToolRegistry; use crate::truncation::truncate_tool_output; -use crate::types::{EventData, EventKind, SessionState, Turn}; +use crate::types::{AgentEvent, SessionState, Turn}; use std::collections::VecDeque; use std::sync::{Arc, Mutex}; use std::time::SystemTime; @@ -65,7 +65,7 @@ impl Session { /// Call before `process_input`. pub async fn initialize(&mut self) { self.event_emitter - .emit(EventKind::SessionStart, self.id.clone(), EventData::Empty); + .emit(self.id.clone(), AgentEvent::SessionStarted); let doc_root = self .config @@ -176,7 +176,7 @@ impl Session { if self.state != SessionState::Closed { self.state = SessionState::Closed; self.event_emitter - .emit(EventKind::SessionEnd, self.id.clone(), EventData::Empty); + .emit(self.id.clone(), AgentEvent::SessionEnded); } } @@ -228,7 +228,7 @@ impl Session { timestamp: SystemTime::now(), }); self.event_emitter - .emit(EventKind::UserInput, self.id.clone(), EventData::Empty); + .emit(self.id.clone(), AgentEvent::UserInput); // Drain steering queue before first LLM call self.drain_steering(); @@ -247,14 +247,14 @@ impl Session { // Check max_tool_rounds_per_input if round_count >= self.config.max_tool_rounds_per_input { self.event_emitter - .emit(EventKind::TurnLimit, self.id.clone(), EventData::Empty); + .emit(self.id.clone(), AgentEvent::TurnLimitReached); break; } // Check max_turns if self.config.max_turns > 0 && self.history.turns().len() >= self.config.max_turns { self.event_emitter - .emit(EventKind::TurnLimit, self.id.clone(), EventData::Empty); + .emit(self.id.clone(), AgentEvent::TurnLimitReached); break; } @@ -268,20 +268,16 @@ impl Session { let request = self.build_request(&system_prompt); // Emit AssistantTextStart before LLM call - self.event_emitter.emit( - EventKind::AssistantTextStart, - self.id.clone(), - EventData::Empty, - ); + self.event_emitter + .emit(self.id.clone(), AgentEvent::AssistantTextStart); // Call LLM (streaming) let mut event_stream = match self.llm_client.stream(&request).await { Ok(stream) => stream, Err(err) => { self.event_emitter.emit( - EventKind::Error, self.id.clone(), - EventData::Error { + AgentEvent::Error { error: err.to_string(), }, ); @@ -299,9 +295,8 @@ impl Session { Ok(event) => { if let StreamEvent::TextDelta { ref delta, .. } = event { self.event_emitter.emit( - EventKind::AssistantTextDelta, self.id.clone(), - EventData::TextDelta { + AgentEvent::TextDelta { delta: delta.clone(), }, ); @@ -310,9 +305,8 @@ impl Session { } Err(err) => { self.event_emitter.emit( - EventKind::Error, self.id.clone(), - EventData::Error { + AgentEvent::Error { error: err.to_string(), }, ); @@ -370,11 +364,16 @@ impl Session { timestamp: SystemTime::now(), }); - // Emit AssistantTextEnd + // Emit AssistantMessage with enriched data from the response self.event_emitter.emit( - EventKind::AssistantTextEnd, self.id.clone(), - EventData::Empty, + AgentEvent::AssistantMessage { + text: text.clone(), + model: response.model.clone(), + input_tokens: response.usage.input_tokens, + output_tokens: response.usage.output_tokens, + tool_call_count: tool_calls.len(), + }, ); // Check context window usage @@ -417,11 +416,8 @@ impl Session { content: "WARNING: Loop detected. You appear to be repeating the same tool calls. Please try a different approach or ask for clarification.".to_string(), timestamp: SystemTime::now(), }); - self.event_emitter.emit( - EventKind::LoopDetection, - self.id.clone(), - EventData::Empty, - ); + self.event_emitter + .emit(self.id.clone(), AgentEvent::LoopDetected); } } @@ -440,11 +436,8 @@ impl Session { content: msg, timestamp: SystemTime::now(), }); - self.event_emitter.emit( - EventKind::SteeringInjected, - self.id.clone(), - EventData::Empty, - ); + self.event_emitter + .emit(self.id.clone(), AgentEvent::SteeringInjected); } } @@ -506,9 +499,8 @@ impl Session { } self.event_emitter.emit( - EventKind::ToolCallStart, self.id.clone(), - EventData::ToolCall { + AgentEvent::ToolCallStarted { tool_name: tc.name.clone(), tool_call_id: tc.id.clone(), arguments: tc.arguments.clone(), @@ -527,17 +519,15 @@ impl Session { .await; self.event_emitter.emit( - EventKind::ToolCallOutputDelta, self.id.clone(), - EventData::TextDelta { + AgentEvent::ToolCallOutputDelta { delta: result.content.to_string(), }, ); self.event_emitter.emit( - EventKind::ToolCallEnd, self.id.clone(), - EventData::ToolCallEnd { + AgentEvent::ToolCallCompleted { tool_name: tc.name.clone(), tool_call_id: tc.id.clone(), output: result.content.clone(), @@ -574,9 +564,8 @@ impl Session { let tc = tc.clone(); async move { emitter.emit( - EventKind::ToolCallStart, session_id.clone(), - EventData::ToolCall { + AgentEvent::ToolCallStarted { tool_name: tc.name.clone(), tool_call_id: tc.id.clone(), arguments: tc.arguments.clone(), @@ -595,17 +584,15 @@ impl Session { .await; emitter.emit( - EventKind::ToolCallOutputDelta, session_id.clone(), - EventData::TextDelta { + AgentEvent::ToolCallOutputDelta { delta: result.content.to_string(), }, ); emitter.emit( - EventKind::ToolCallEnd, session_id, - EventData::ToolCallEnd { + AgentEvent::ToolCallCompleted { tool_name: tc.name.clone(), tool_call_id: tc.id.clone(), output: result.content.clone(), @@ -663,9 +650,8 @@ impl Session { if estimated_tokens > threshold { self.event_emitter.emit( - EventKind::ContextWindowWarning, self.id.clone(), - EventData::ContextWarning { + AgentEvent::ContextWindowWarning { estimated_tokens, context_window_size: context_window, usage_percent: estimated_tokens * 100 / context_window, @@ -959,13 +945,13 @@ mod tests { // Collect events let mut events = Vec::new(); while let Ok(event) = rx.try_recv() { - events.push(event.kind.clone()); + events.push(event); } - assert!(events.contains(&EventKind::SessionStart)); - assert!(events.contains(&EventKind::UserInput)); - assert!(events.contains(&EventKind::AssistantTextEnd)); - assert!(events.contains(&EventKind::SessionEnd)); + assert!(events.iter().any(|e| matches!(e.event, AgentEvent::SessionStarted))); + assert!(events.iter().any(|e| matches!(e.event, AgentEvent::UserInput))); + assert!(events.iter().any(|e| matches!(e.event, AgentEvent::AssistantMessage { .. }))); + assert!(events.iter().any(|e| matches!(e.event, AgentEvent::SessionEnded))); } #[tokio::test] @@ -985,17 +971,17 @@ mod tests { let mut tool_end_events = Vec::new(); while let Ok(event) = rx.try_recv() { - if event.kind == EventKind::ToolCallEnd { + if matches!(event.event, AgentEvent::ToolCallCompleted { .. }) { tool_end_events.push(event); } } assert_eq!(tool_end_events.len(), 1); - match &tool_end_events[0].data { - EventData::ToolCallEnd { output, .. } => { + match &tool_end_events[0].event { + AgentEvent::ToolCallCompleted { output, .. } => { assert_eq!(output, &serde_json::json!("echo: hello world")); } - _ => panic!("Expected ToolCallEnd event data"), + _ => panic!("Expected ToolCallCompleted event"), } } @@ -1073,10 +1059,10 @@ mod tests { session.process_input("Keep echoing").await.unwrap(); - // Check for LoopDetection event + // Check for LoopDetected event let mut found_loop_detection = false; while let Ok(event) = rx.try_recv() { - if event.kind == EventKind::LoopDetection { + if matches!(event.event, AgentEvent::LoopDetected) { found_loop_detection = true; } } @@ -1238,14 +1224,14 @@ mod tests { let result = session.process_input("Hello").await; assert!(matches!(result, Err(AgentError::SessionClosed))); - // No SessionStart event should have been emitted + // No SessionStarted event should have been emitted let mut events = Vec::new(); while let Ok(event) = rx.try_recv() { - events.push(event.kind.clone()); + events.push(event); } assert!( - !events.contains(&EventKind::SessionStart), - "SessionStart should not be emitted for a closed session" + !events.iter().any(|e| matches!(e.event, AgentEvent::SessionStarted)), + "SessionStarted should not be emitted for a closed session" ); } @@ -1289,13 +1275,13 @@ mod tests { panic!("Expected ToolResults turn at index 2"); } - // Verify ToolCallStart and ToolCallEnd events for all 3 calls + // Verify ToolCallStarted and ToolCallCompleted events for all 3 calls let mut start_count = 0; let mut end_count = 0; while let Ok(event) = rx.try_recv() { - match event.kind { - EventKind::ToolCallStart => start_count += 1, - EventKind::ToolCallEnd => end_count += 1, + match &event.event { + AgentEvent::ToolCallStarted { .. } => start_count += 1, + AgentEvent::ToolCallCompleted { .. } => end_count += 1, _ => {} } } @@ -1327,17 +1313,13 @@ mod tests { let mut found_warning = false; while let Ok(event) = rx.try_recv() { - if event.kind == EventKind::ContextWindowWarning { + if let AgentEvent::ContextWindowWarning { + context_window_size, + .. + } = &event.event + { found_warning = true; - match &event.data { - EventData::ContextWarning { - context_window_size, - .. - } => { - assert_eq!(*context_window_size, 100); - } - _ => panic!("Expected ContextWarning event data"), - } + assert_eq!(*context_window_size, 100); } } assert!(found_warning); @@ -1380,7 +1362,7 @@ mod tests { let mut found_warning = false; while let Ok(event) = rx.try_recv() { - if event.kind == EventKind::ContextWindowWarning { + if matches!(event.event, AgentEvent::ContextWindowWarning { .. }) { found_warning = true; } } @@ -1486,14 +1468,14 @@ mod tests { let mut session_start_count = 0; let mut session_end_count = 0; while let Ok(event) = rx.try_recv() { - if event.kind == EventKind::SessionStart { + if matches!(event.event, AgentEvent::SessionStarted) { session_start_count += 1; } - if event.kind == EventKind::SessionEnd { + if matches!(event.event, AgentEvent::SessionEnded) { session_end_count += 1; } } - // SESSION_START is emitted once during initialize(), SESSION_END once during close() + // SessionStarted is emitted once during initialize(), SessionEnded once during close() assert_eq!(session_start_count, 1); assert_eq!(session_end_count, 1); } @@ -1681,17 +1663,17 @@ mod tests { let mut tool_end_events = Vec::new(); while let Ok(event) = rx.try_recv() { - if event.kind == EventKind::ToolCallEnd { + if matches!(event.event, AgentEvent::ToolCallCompleted { .. }) { tool_end_events.push(event); } } assert_eq!(tool_end_events.len(), 1); - match &tool_end_events[0].data { - EventData::ToolCallEnd { is_error, .. } => { - assert!(is_error, "ToolCallEnd event should have is_error: true"); + match &tool_end_events[0].event { + AgentEvent::ToolCallCompleted { is_error, .. } => { + assert!(is_error, "ToolCallCompleted event should have is_error: true"); } - _ => panic!("Expected ToolCallEnd event data"), + _ => panic!("Expected ToolCallCompleted event"), } } @@ -1704,10 +1686,8 @@ mod tests { let mut deltas = Vec::new(); while let Ok(event) = rx.try_recv() { - if event.kind == EventKind::AssistantTextDelta { - if let EventData::TextDelta { delta } = &event.data { - deltas.push(delta.clone()); - } + if let AgentEvent::TextDelta { delta } = &event.event { + deltas.push(delta.clone()); } } diff --git a/crates/agent/src/types.rs b/crates/agent/src/types.rs index 26bf094a4..0b9f0402a 100644 --- a/crates/agent/src/types.rs +++ b/crates/agent/src/types.rs @@ -43,33 +43,31 @@ pub enum SessionState { Closed, } -#[derive(Debug, Clone, PartialEq, Eq)] -pub enum EventKind { - SessionStart, - SessionEnd, +#[derive(Debug, Clone)] +pub enum AgentEvent { + SessionStarted, + SessionEnded, UserInput, AssistantTextStart, - AssistantTextDelta, - AssistantTextEnd, - ToolCallStart, - ToolCallOutputDelta, - ToolCallEnd, - SteeringInjected, - TurnLimit, - LoopDetection, - ContextWindowWarning, - Error, -} - -#[derive(Debug, Clone)] -pub enum EventData { - Empty, - ToolCall { + AssistantMessage { + text: String, + model: String, + input_tokens: i64, + output_tokens: i64, + tool_call_count: usize, + }, + TextDelta { + delta: String, + }, + ToolCallStarted { tool_name: String, tool_call_id: String, arguments: serde_json::Value, }, - ToolCallEnd { + ToolCallOutputDelta { + delta: String, + }, + ToolCallCompleted { tool_name: String, tool_call_id: String, output: serde_json::Value, @@ -78,22 +76,21 @@ pub enum EventData { Error { error: String, }, - TextDelta { - delta: String, - }, - ContextWarning { + ContextWindowWarning { estimated_tokens: usize, context_window_size: usize, usage_percent: usize, }, + LoopDetected, + TurnLimitReached, + SteeringInjected, } #[derive(Debug, Clone)] pub struct SessionEvent { - pub kind: EventKind, + pub event: AgentEvent, pub timestamp: SystemTime, pub session_id: String, - pub data: EventData, } #[cfg(test)] @@ -103,13 +100,23 @@ mod tests { #[test] fn session_event_construction() { let event = SessionEvent { - kind: EventKind::SessionStart, + event: AgentEvent::SessionStarted, timestamp: SystemTime::now(), session_id: "sess_1".into(), - data: EventData::Empty, }; - assert_eq!(event.kind, EventKind::SessionStart); + assert!(matches!(event.event, AgentEvent::SessionStarted)); assert_eq!(event.session_id, "sess_1"); - assert!(matches!(event.data, EventData::Empty)); + } + + #[test] + fn agent_event_assistant_message() { + let event = AgentEvent::AssistantMessage { + text: "Hello".into(), + model: "test-model".into(), + input_tokens: 100, + output_tokens: 50, + tool_call_count: 2, + }; + assert!(matches!(event, AgentEvent::AssistantMessage { tool_call_count: 2, .. })); } } diff --git a/crates/attractor/src/cli/backend.rs b/crates/attractor/src/cli/backend.rs index 73b5356f1..042d6acb5 100644 --- a/crates/attractor/src/cli/backend.rs +++ b/crates/attractor/src/cli/backend.rs @@ -4,7 +4,7 @@ use std::sync::Arc; use async_trait::async_trait; use agent::{ - AnthropicProfile, DockerConfig, DockerExecutionEnvironment, EventData, EventKind, + AgentEvent, AnthropicProfile, DockerConfig, DockerExecutionEnvironment, ExecutionEnvironment, GeminiProfile, LocalExecutionEnvironment, OpenAiProfile, ProviderProfile, Session, SessionConfig, Turn, }; @@ -62,6 +62,7 @@ impl CodergenBackend for AgentBackend { prompt: &str, _context: &Context, _thread_id: Option<&str>, + emitter: &Arc, ) -> Result { let client = Client::from_env() .await @@ -90,23 +91,99 @@ impl CodergenBackend for AgentBackend { let mut session = Session::new(client, profile, exec_env, config); - // Subscribe to session events for real-time tool status on stderr. + // Subscribe to session events: forward to pipeline emitter and optionally print to stderr. let verbose = self.verbose; - if verbose >= 1 { - let node_id = node.id.clone(); - let styles = self.styles; - let mut rx = session.subscribe(); - tokio::spawn(async move { - while let Ok(event) = rx.recv().await { - match (&event.kind, &event.data) { - ( - EventKind::ToolCallStart, - EventData::ToolCall { - tool_name, - arguments, - .. - }, - ) => { + let node_id = node.id.clone(); + let styles = self.styles; + let pipeline_emitter = Arc::clone(emitter); + let mut rx = session.subscribe(); + tokio::spawn(async move { + use crate::event::PipelineEvent; + while let Ok(event) = rx.recv().await { + // Forward agent events to pipeline events + match &event.event { + AgentEvent::AssistantMessage { + text, + model, + input_tokens, + output_tokens, + tool_call_count, + } => { + pipeline_emitter.emit(&PipelineEvent::AssistantMessage { + stage: node_id.clone(), + text: text.clone(), + model: model.clone(), + input_tokens: *input_tokens, + output_tokens: *output_tokens, + tool_call_count: *tool_call_count, + }); + } + AgentEvent::ToolCallStarted { + tool_name, + tool_call_id, + arguments, + } => { + pipeline_emitter.emit(&PipelineEvent::ToolCallStarted { + stage: node_id.clone(), + tool_name: tool_name.clone(), + tool_call_id: tool_call_id.clone(), + arguments: arguments.clone(), + }); + } + AgentEvent::ToolCallCompleted { + tool_name, + tool_call_id, + output, + is_error, + } => { + pipeline_emitter.emit(&PipelineEvent::ToolCallCompleted { + stage: node_id.clone(), + tool_name: tool_name.clone(), + tool_call_id: tool_call_id.clone(), + output: output.clone(), + is_error: *is_error, + }); + } + AgentEvent::Error { error } => { + pipeline_emitter.emit(&PipelineEvent::SessionError { + stage: node_id.clone(), + error: error.clone(), + }); + } + AgentEvent::ContextWindowWarning { + estimated_tokens, + context_window_size, + usage_percent, + } => { + pipeline_emitter.emit(&PipelineEvent::ContextWindowWarning { + stage: node_id.clone(), + estimated_tokens: *estimated_tokens, + context_window_size: *context_window_size, + usage_percent: *usage_percent, + }); + } + AgentEvent::LoopDetected => { + pipeline_emitter.emit(&PipelineEvent::LoopDetected { + stage: node_id.clone(), + }); + } + AgentEvent::TurnLimitReached => { + pipeline_emitter.emit(&PipelineEvent::TurnLimitReached { + stage: node_id.clone(), + }); + } + // Streaming events and session lifecycle not forwarded + _ => {} + } + + // Verbose stderr printing (gated on verbosity) + if verbose >= 1 { + match &event.event { + AgentEvent::ToolCallStarted { + tool_name, + arguments, + .. + } => { eprintln!( "{dim}[{node_id}]{reset} {dim}\u{25cf}{reset} {bold}{cyan}{tool_name}{reset}{dim}({args}){reset}", dim = styles.dim, @@ -116,15 +193,12 @@ impl CodergenBackend for AgentBackend { args = format_tool_args(arguments), ); } - ( - EventKind::ToolCallEnd, - EventData::ToolCallEnd { - tool_name, - output, - is_error, - .. - }, - ) if verbose >= 2 => { + AgentEvent::ToolCallCompleted { + tool_name, + output, + is_error, + .. + } if verbose >= 2 => { let label = if *is_error { "error" } else { "result" }; eprintln!( "{dim}[{node_id}] [{label}] {tool_name}:{reset}\n{}", @@ -134,7 +208,7 @@ impl CodergenBackend for AgentBackend { reset = styles.reset, ); } - (EventKind::Error, EventData::Error { error }) => { + AgentEvent::Error { error } => { eprintln!( "{dim}[{node_id}]{reset} {red}\u{2717} {error}{reset}", dim = styles.dim, @@ -145,8 +219,14 @@ impl CodergenBackend for AgentBackend { _ => {} } } - }); - } + } + }); + + // Emit Prompt event before processing + emitter.emit(&crate::event::PipelineEvent::Prompt { + stage: node.id.clone(), + text: prompt.to_string(), + }); session.initialize().await; session.process_input(prompt).await.map_err(|e| { diff --git a/crates/attractor/src/cli/mod.rs b/crates/attractor/src/cli/mod.rs index ce0ff2671..8a3d21614 100644 --- a/crates/attractor/src/cli/mod.rs +++ b/crates/attractor/src/cli/mod.rs @@ -269,6 +269,53 @@ pub fn format_event_summary(event: &PipelineEvent, styles: &Styles) -> String { PipelineEvent::CheckpointSaved { node_id } => { format!("[CHECKPOINT_SAVED] node={node_id}") } + PipelineEvent::Prompt { stage, text } => { + let truncated = if text.len() > 80 { &text[..80] } else { text }; + format!("[PROMPT] stage={stage} text=\"{truncated}\"") + } + PipelineEvent::AssistantMessage { + stage, + model, + input_tokens, + output_tokens, + tool_call_count, + .. + } => { + let total = input_tokens + output_tokens; + let tokens_str = format_tokens_human(total); + format!("[ASSISTANT_MESSAGE] stage={stage} model={model} tokens={tokens_str} tool_calls={tool_call_count}") + } + PipelineEvent::ToolCallStarted { + stage, + tool_name, + .. + } => { + format!("[TOOL_CALL_STARTED] stage={stage} tool={tool_name}") + } + PipelineEvent::ToolCallCompleted { + stage, + tool_name, + is_error, + .. + } => { + format!("[TOOL_CALL_COMPLETED] stage={stage} tool={tool_name} is_error={is_error}") + } + PipelineEvent::SessionError { stage, error } => { + format!("[SESSION_ERROR] stage={stage} error=\"{error}\"") + } + PipelineEvent::ContextWindowWarning { + stage, + usage_percent, + .. + } => { + format!("[CONTEXT_WINDOW_WARNING] stage={stage} usage={usage_percent}%") + } + PipelineEvent::LoopDetected { stage } => { + format!("[LOOP_DETECTED] stage={stage}") + } + PipelineEvent::TurnLimitReached { stage } => { + format!("[TURN_LIMIT_REACHED] stage={stage}") + } }; format!("{dim}{body}{reset}", dim = styles.dim, reset = styles.reset) } @@ -389,6 +436,63 @@ pub fn format_event_detail(event: &PipelineEvent, styles: &Styles) -> String { "{d}── CHECKPOINT_SAVED ─────────────────────────{r}\n {d}node_id:{r} {node_id}\n" ) } + PipelineEvent::Prompt { stage, text } => { + format!("{d}── PROMPT ───────────────────────────────────{r}\n {d}stage:{r} {stage}\n {d}text:{r}\n{text}\n") + } + PipelineEvent::AssistantMessage { + stage, + text, + model, + input_tokens, + output_tokens, + tool_call_count, + } => { + let total = input_tokens + output_tokens; + let truncated = if text.len() > 200 { &text[..200] } else { text.as_str() }; + format!("{d}── ASSISTANT_MESSAGE ────────────────────────{r}\n {d}stage:{r} {stage}\n {d}model:{r} {model}\n {d}tokens:{r} {} ({} in / {} out)\n {d}tool_calls:{r} {tool_call_count}\n {d}text:{r} {truncated}\n", + format_tokens_human(total), + format_tokens_human(*input_tokens), + format_tokens_human(*output_tokens), + ) + } + PipelineEvent::ToolCallStarted { + stage, + tool_name, + tool_call_id, + arguments, + } => { + let args_str = serde_json::to_string(arguments).unwrap_or_else(|_| arguments.to_string()); + let truncated = if args_str.len() > 200 { &args_str[..200] } else { &args_str }; + format!("{d}── TOOL_CALL_STARTED ────────────────────────{r}\n {d}stage:{r} {stage}\n {d}tool_name:{r} {tool_name}\n {d}tool_call_id:{r} {tool_call_id}\n {d}arguments:{r} {truncated}\n") + } + PipelineEvent::ToolCallCompleted { + stage, + tool_name, + tool_call_id, + output, + is_error, + } => { + let output_str = serde_json::to_string(output).unwrap_or_else(|_| output.to_string()); + let truncated = if output_str.len() > 200 { &output_str[..200] } else { &output_str }; + format!("{d}── TOOL_CALL_COMPLETED ──────────────────────{r}\n {d}stage:{r} {stage}\n {d}tool_name:{r} {tool_name}\n {d}tool_call_id:{r} {tool_call_id}\n {d}is_error:{r} {is_error}\n {d}output:{r} {truncated}\n") + } + PipelineEvent::SessionError { stage, error } => { + format!("{d}── SESSION_ERROR ────────────────────────────{r}\n {d}stage:{r} {stage}\n {d}error:{r} {error}\n") + } + PipelineEvent::ContextWindowWarning { + stage, + estimated_tokens, + context_window_size, + usage_percent, + } => { + format!("{d}── CONTEXT_WINDOW_WARNING ───────────────────{r}\n {d}stage:{r} {stage}\n {d}estimated_tokens:{r} {estimated_tokens}\n {d}context_window_size:{r} {context_window_size}\n {d}usage_percent:{r} {usage_percent}%\n") + } + PipelineEvent::LoopDetected { stage } => { + format!("{d}── LOOP_DETECTED ────────────────────────────{r}\n {d}stage:{r} {stage}\n") + } + PipelineEvent::TurnLimitReached { stage } => { + format!("{d}── TURN_LIMIT_REACHED ───────────────────────{r}\n {d}stage:{r} {stage}\n") + } } } diff --git a/crates/attractor/src/event.rs b/crates/attractor/src/event.rs index 1d0a55871..d00513633 100644 --- a/crates/attractor/src/event.rs +++ b/crates/attractor/src/event.rs @@ -77,6 +77,47 @@ pub enum PipelineEvent { CheckpointSaved { node_id: String, }, + Prompt { + stage: String, + text: String, + }, + AssistantMessage { + stage: String, + text: String, + model: String, + input_tokens: i64, + output_tokens: i64, + tool_call_count: usize, + }, + ToolCallStarted { + stage: String, + tool_name: String, + tool_call_id: String, + arguments: serde_json::Value, + }, + ToolCallCompleted { + stage: String, + tool_name: String, + tool_call_id: String, + output: serde_json::Value, + is_error: bool, + }, + SessionError { + stage: String, + error: String, + }, + ContextWindowWarning { + stage: String, + estimated_tokens: usize, + context_window_size: usize, + usage_percent: usize, + }, + LoopDetected { + stage: String, + }, + TurnLimitReached { + stage: String, + }, } /// Listener callback type for pipeline events. @@ -168,4 +209,37 @@ mod tests { let emitter = EventEmitter::default(); assert_eq!(emitter.listeners.len(), 0); } + + #[test] + fn llm_conversation_event_serialization() { + let event = PipelineEvent::ToolCallStarted { + stage: "plan".to_string(), + tool_name: "read_file".to_string(), + tool_call_id: "call_1".to_string(), + arguments: serde_json::json!({"path": "/tmp/test.txt"}), + }; + let json = serde_json::to_string(&event).unwrap(); + assert!(json.contains("ToolCallStarted")); + assert!(json.contains("read_file")); + assert!(json.contains("plan")); + + // Verify round-trip + let deserialized: PipelineEvent = serde_json::from_str(&json).unwrap(); + assert!(matches!(deserialized, PipelineEvent::ToolCallStarted { stage, .. } if stage == "plan")); + } + + #[test] + fn assistant_message_event_serialization() { + let event = PipelineEvent::AssistantMessage { + stage: "code".to_string(), + text: "Here is the implementation".to_string(), + model: "claude-opus-4-6".to_string(), + input_tokens: 1000, + output_tokens: 500, + tool_call_count: 3, + }; + let json = serde_json::to_string(&event).unwrap(); + assert!(json.contains("AssistantMessage")); + assert!(json.contains("claude-opus-4-6")); + } } diff --git a/crates/attractor/src/handler/codergen.rs b/crates/attractor/src/handler/codergen.rs index 9a85b8955..f9f405682 100644 --- a/crates/attractor/src/handler/codergen.rs +++ b/crates/attractor/src/handler/codergen.rs @@ -1,9 +1,11 @@ use std::path::Path; +use std::sync::Arc; use async_trait::async_trait; use crate::context::Context; use crate::error::AttractorError; +use crate::event::EventEmitter; use crate::graph::{Graph, Node}; use crate::outcome::{Outcome, StageUsage}; @@ -27,6 +29,7 @@ pub trait CodergenBackend: Send + Sync { prompt: &str, context: &Context, thread_id: Option<&str>, + emitter: &Arc, ) -> Result; } @@ -175,7 +178,7 @@ impl Handler for CodergenHandler { context: &Context, graph: &Graph, logs_root: &Path, - _services: &EngineServices, + services: &EngineServices, ) -> Result { // 1. Build prompt let raw_prompt = node @@ -203,7 +206,7 @@ impl Handler for CodergenHandler { .get("internal.thread_id") .and_then(|v| v.as_str().map(String::from)); let (response_text, stage_usage) = if let Some(backend) = &self.backend { - match backend.run(node, &prompt, context, thread_id.as_deref()).await { + match backend.run(node, &prompt, context, thread_id.as_deref(), &services.emitter).await { Ok(CodergenResult::Full(outcome)) => { let status_json = serde_json::to_string_pretty(&outcome) .unwrap_or_else(|_| "{}".to_string()); @@ -521,6 +524,7 @@ mod tests { _prompt: &str, _context: &Context, thread_id: Option<&str>, + _emitter: &Arc, ) -> Result { *self.captured_thread_id.lock().unwrap() = Some(thread_id.map(String::from)); @@ -566,6 +570,7 @@ mod tests { _prompt: &str, _context: &Context, thread_id: Option<&str>, + _emitter: &Arc, ) -> Result { *self.captured_thread_id.lock().unwrap() = Some(thread_id.map(String::from)); @@ -606,6 +611,7 @@ mod tests { _prompt: &str, _context: &Context, _thread_id: Option<&str>, + _emitter: &Arc, ) -> Result { Err(AttractorError::Handler("Request timed out".to_string())) } @@ -717,6 +723,7 @@ Some text in between. _prompt: &str, _context: &Context, _thread_id: Option<&str>, + _emitter: &Arc, ) -> Result { Err(AttractorError::Validation("bad config".to_string())) } diff --git a/crates/attractor/src/handler/fan_in.rs b/crates/attractor/src/handler/fan_in.rs index 43af6230e..e40db4846 100644 --- a/crates/attractor/src/handler/fan_in.rs +++ b/crates/attractor/src/handler/fan_in.rs @@ -1,9 +1,11 @@ use std::path::Path; +use std::sync::Arc; use async_trait::async_trait; use crate::context::Context; use crate::error::AttractorError; +use crate::event::EventEmitter; use crate::graph::{Graph, Node}; use crate::outcome::Outcome; @@ -30,7 +32,7 @@ impl Handler for FanInHandler { context: &Context, _graph: &Graph, logs_root: &Path, - _services: &EngineServices, + services: &EngineServices, ) -> Result { let results = context.get("parallel.results"); let Some(results) = results else { @@ -40,7 +42,7 @@ impl Handler for FanInHandler { let prompt = node.prompt().filter(|p| !p.is_empty()); let best = if let (Some(prompt_text), Some(backend)) = (prompt, &self.backend) { - llm_evaluate(backend.as_ref(), prompt_text, &results, context, logs_root, &node.id).await? + llm_evaluate(backend.as_ref(), prompt_text, &results, context, logs_root, &node.id, &services.emitter).await? } else { heuristic_select(&results) }; @@ -153,6 +155,7 @@ async fn llm_evaluate( context: &Context, logs_root: &Path, node_id: &str, + emitter: &Arc, ) -> Result { let results_text = serde_json::to_string_pretty(results) .unwrap_or_else(|_| results.to_string()); @@ -171,7 +174,7 @@ async fn llm_evaluate( let eval_node = Node::new("fan_in_eval"); // Fan-in evaluation runs outside a thread context, so pass None - match backend.run(&eval_node, &full_prompt, context, None).await { + match backend.run(&eval_node, &full_prompt, context, None, emitter).await { Ok(CodergenResult::Full(outcome)) => { // If the backend returned a full Outcome, extract best_id from context_updates let best_id = outcome @@ -367,6 +370,7 @@ mod tests { _prompt: &str, _context: &Context, _thread_id: Option<&str>, + _emitter: &Arc, ) -> Result { // Return text that contains the ID "branch_b" Ok(CodergenResult::Text { text: "The best candidate is branch_b".to_string(), usage: None }) diff --git a/crates/attractor/tests/integration.rs b/crates/attractor/tests/integration.rs index 570e697e3..1864a8688 100644 --- a/crates/attractor/tests/integration.rs +++ b/crates/attractor/tests/integration.rs @@ -1056,6 +1056,7 @@ impl CodergenBackend for MockCodergenBackend { prompt: &str, _context: &Context, _thread_id: Option<&str>, + _emitter: &Arc, ) -> Result { Ok(CodergenResult::Text { text: format!( @@ -5079,6 +5080,7 @@ mod real_llm { prompt: &str, _context: &Context, _thread_id: Option<&str>, + _emitter: &Arc, ) -> Result { let request = Request { model: self.model.clone(),