use std::sync::{Arc, Mutex, MutexGuard, Weak}; use serde::{Deserialize, Serialize}; use tokio::sync::watch; use crate::agent::{TurnEmitError, TurnEmitter, TurnEvent}; use crate::providers::Usage; #[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)] #[serde(transparent)] pub struct TurnId(pub String); impl TurnId { pub fn new() -> Self { Self(uuid::Uuid::new_v4().to_string()) } } impl Default for TurnId { fn default() -> Self { Self::new() } } #[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)] #[serde(transparent)] pub struct BlockId(pub String); impl BlockId { fn new() -> Self { Self(uuid::Uuid::new_v4().to_string()) } } #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] #[serde(rename_all = "snake_case")] pub enum TurnStatus { Running, Completed, Cancelled, Failed, } #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] #[serde(rename_all = "snake_case")] pub enum TurnPhase { Queued, Reasoning, Responding, Acting, Finalizing, } #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] #[serde(rename_all = "snake_case")] pub enum ToolStatus { Running, Completed, Cancelled, Failed, } #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] #[serde(tag = "type", rename_all = "snake_case")] pub enum TurnBlock { Reasoning { id: BlockId, iteration: u32, text: String, }, Assistant { id: BlockId, iteration: u32, text: String, }, Tool { id: String, iteration: u32, name: String, arguments: serde_json::Value, status: ToolStatus, preview: Option, }, } #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub struct TurnState { pub id: TurnId, pub session_id: String, pub message_id: String, pub revision: u64, pub status: TurnStatus, pub phase: TurnPhase, pub blocks: Vec, pub usage: Option, pub error: Option, } pub type TurnSnapshot = TurnState; struct TurnControllerInner { state: TurnState, snapshots: watch::Sender>, text_segment_open: bool, } impl TurnControllerInner { fn emit(&mut self, event: TurnEvent) -> Result<(), TurnEmitError> { if self.state.status != TurnStatus::Running { return Err(TurnEmitError::Inactive); } let changed = match event { TurnEvent::ReasoningDelta { iteration, delta } => { if delta.is_empty() { false } else { self.state.phase = TurnPhase::Reasoning; match self.state.blocks.last_mut() { Some(TurnBlock::Reasoning { iteration: current, text, .. }) if *current == iteration => text.push_str(&delta), _ => self.state.blocks.push(TurnBlock::Reasoning { id: BlockId::new(), iteration, text: delta, }), } true } } TurnEvent::TextDelta { iteration, delta } => { if delta.is_empty() { false } else { self.state.phase = TurnPhase::Responding; if self.text_segment_open { match self.state.blocks.last_mut() { Some(TurnBlock::Assistant { iteration: current, text, .. }) if *current == iteration => text.push_str(&delta), _ => { self.push_text_block(iteration, delta); } } } else { self.push_text_block(iteration, delta); } self.text_segment_open = true; true } } TurnEvent::TextSegmentFinished { .. } => { self.text_segment_open = false; false } TurnEvent::ToolStarted { iteration, call } => { if self .state .blocks .iter() .any(|block| matches!(block, TurnBlock::Tool { id, .. } if id == &call.id)) { return Err(TurnEmitError::DuplicateTool(call.id)); } self.text_segment_open = false; self.state.phase = TurnPhase::Acting; self.state.blocks.push(TurnBlock::Tool { id: call.id, iteration, name: call.name, arguments: call.arguments, status: ToolStatus::Running, preview: None, }); true } TurnEvent::ToolFinished { iteration, call_id, success, preview, } => { let Some(TurnBlock::Tool { status, preview: current_preview, .. }) = self.state.blocks.iter_mut().find(|block| { matches!(block, TurnBlock::Tool { id, iteration: current, .. } if id == &call_id && *current == iteration) }) else { return Err(TurnEmitError::UnknownTool(call_id)); }; *status = if success { ToolStatus::Completed } else { ToolStatus::Failed }; *current_preview = preview; true } }; if changed { self.publish(); } Ok(()) } fn push_text_block(&mut self, iteration: u32, text: String) { self.state.blocks.push(TurnBlock::Assistant { id: BlockId::new(), iteration, text, }); } fn publish(&mut self) { self.state.revision = self.state.revision.wrapping_add(1); self.snapshots.send_replace(Arc::new(self.state.clone())); } fn transition_terminal( &mut self, status: TurnStatus, usage: Option, error: Option, ) -> bool { if self.state.status != TurnStatus::Running { return false; } self.text_segment_open = false; if status == TurnStatus::Cancelled { for block in &mut self.state.blocks { if let TurnBlock::Tool { status: tool_status, .. } = block && *tool_status == ToolStatus::Running { *tool_status = ToolStatus::Cancelled; } } } self.state.status = status; self.state.phase = TurnPhase::Finalizing; self.state.usage = usage; self.state.error = error; self.publish(); true } } /// The sole writer for one running turn's presentation state. /// /// Mutation is synchronous and bounded to a small in-memory reduction. This /// lets AgentLoop emit facts without creating an unbounded token queue or a /// reducer background task. pub struct TurnController { inner: Arc>, } impl TurnController { pub fn start( session_id: impl Into, message_id: impl Into, ) -> (Self, TurnEmitter, watch::Receiver>) { let initial = TurnState { id: TurnId::new(), session_id: session_id.into(), message_id: message_id.into(), revision: 0, status: TurnStatus::Running, phase: TurnPhase::Queued, blocks: Vec::new(), usage: None, error: None, }; let (snapshots, receiver) = watch::channel(Arc::new(initial.clone())); let inner = Arc::new(Mutex::new(TurnControllerInner { state: initial, snapshots, text_segment_open: false, })); let weak: Weak> = Arc::downgrade(&inner); let emitter = TurnEmitter::new(move |event| { let Some(inner) = weak.upgrade() else { return Err(TurnEmitError::Inactive); }; lock_unpoisoned(&inner).emit(event) }); (Self { inner }, emitter, receiver) } pub fn snapshot(&self) -> Arc { Arc::new(lock_unpoisoned(&self.inner).state.clone()) } pub fn begin_finalizing(&self) -> bool { let mut inner = lock_unpoisoned(&self.inner); if inner.state.status != TurnStatus::Running || inner.state.phase == TurnPhase::Finalizing { return false; } inner.text_segment_open = false; inner.state.phase = TurnPhase::Finalizing; inner.publish(); true } pub fn complete(&self, usage: Option) -> bool { lock_unpoisoned(&self.inner).transition_terminal(TurnStatus::Completed, usage, None) } pub fn cancel(&self, reason: Option) -> bool { lock_unpoisoned(&self.inner).transition_terminal(TurnStatus::Cancelled, None, reason) } pub fn fail(&self, error: impl Into) -> bool { lock_unpoisoned(&self.inner).transition_terminal( TurnStatus::Failed, None, Some(error.into()), ) } } fn lock_unpoisoned(mutex: &Mutex) -> MutexGuard<'_, T> { mutex .lock() .unwrap_or_else(|poisoned| poisoned.into_inner()) } #[cfg(test)] mod tests { use super::*; use crate::providers::ToolCall; fn start() -> ( TurnController, TurnEmitter, watch::Receiver>, ) { TurnController::start("cli:test:dialog", "assistant-message") } #[test] fn ordered_blocks_preserve_reasoning_text_tool_and_iterations() { let (controller, emitter, _) = start(); emitter .emit(TurnEvent::ReasoningDelta { iteration: 0, delta: "plan ".into(), }) .unwrap(); emitter .emit(TurnEvent::ReasoningDelta { iteration: 0, delta: "step".into(), }) .unwrap(); emitter .emit(TurnEvent::TextDelta { iteration: 0, delta: "checking".into(), }) .unwrap(); emitter .emit(TurnEvent::ToolStarted { iteration: 0, call: ToolCall { id: "call-1".into(), name: "bash".into(), arguments: serde_json::json!({"cmd":"pwd"}), }, }) .unwrap(); emitter .emit(TurnEvent::ToolFinished { iteration: 0, call_id: "call-1".into(), success: true, preview: Some("/tmp".into()), }) .unwrap(); emitter .emit(TurnEvent::ReasoningDelta { iteration: 1, delta: "done thinking".into(), }) .unwrap(); emitter .emit(TurnEvent::TextDelta { iteration: 1, delta: "final".into(), }) .unwrap(); let snapshot = controller.snapshot(); assert_eq!(snapshot.revision, 7); assert_eq!(snapshot.phase, TurnPhase::Responding); assert_eq!(snapshot.blocks.len(), 5); assert!(matches!( &snapshot.blocks[0], TurnBlock::Reasoning { iteration: 0, text, .. } if text == "plan step" )); assert!(matches!( &snapshot.blocks[1], TurnBlock::Assistant { iteration: 0, text, .. } if text == "checking" )); assert!(matches!( &snapshot.blocks[2], TurnBlock::Tool { id, status: ToolStatus::Completed, preview: Some(preview), .. } if id == "call-1" && preview == "/tmp" )); assert!(matches!( &snapshot.blocks[3], TurnBlock::Reasoning { iteration: 1, .. } )); assert!(matches!( &snapshot.blocks[4], TurnBlock::Assistant { iteration: 1, .. } )); } #[test] fn explicit_segment_boundary_prevents_text_coalescing() { let (controller, emitter, _) = start(); emitter .emit(TurnEvent::TextDelta { iteration: 0, delta: "before".into(), }) .unwrap(); emitter .emit(TurnEvent::TextSegmentFinished { iteration: 0 }) .unwrap(); emitter .emit(TurnEvent::TextDelta { iteration: 0, delta: "after".into(), }) .unwrap(); let snapshot = controller.snapshot(); assert_eq!(snapshot.revision, 2); assert_eq!(snapshot.blocks.len(), 2); assert!( matches!(&snapshot.blocks[0], TurnBlock::Assistant { text, .. } if text == "before") ); assert!( matches!(&snapshot.blocks[1], TurnBlock::Assistant { text, .. } if text == "after") ); } #[test] fn revisions_are_monotonic_and_watch_is_latest_wins() { let (_controller, emitter, receiver) = start(); for delta in ["a", "b", "c"] { emitter .emit(TurnEvent::TextDelta { iteration: 0, delta: delta.into(), }) .unwrap(); } let latest = receiver.borrow().clone(); assert_eq!(latest.revision, 3); assert!(matches!(&latest.blocks[0], TurnBlock::Assistant { text, .. } if text == "abc")); } #[test] fn terminal_state_rejects_late_events_and_is_idempotent() { let (controller, emitter, receiver) = start(); emitter .emit(TurnEvent::TextDelta { iteration: 0, delta: "saved".into(), }) .unwrap(); assert!(controller.begin_finalizing()); assert!(controller.complete(None)); assert!(!controller.complete(None)); assert_eq!( emitter.emit(TurnEvent::TextDelta { iteration: 0, delta: "late".into(), }), Err(TurnEmitError::Inactive) ); let latest = receiver.borrow().clone(); assert_eq!(latest.status, TurnStatus::Completed); assert_eq!(latest.phase, TurnPhase::Finalizing); assert_eq!(latest.revision, 3); assert!(matches!(&latest.blocks[0], TurnBlock::Assistant { text, .. } if text == "saved")); } #[test] fn invalid_tool_transitions_do_not_publish() { let (_controller, emitter, receiver) = start(); let unknown = emitter.emit(TurnEvent::ToolFinished { iteration: 0, call_id: "missing".into(), success: false, preview: None, }); assert_eq!(unknown, Err(TurnEmitError::UnknownTool("missing".into()))); assert_eq!(receiver.borrow().revision, 0); let call = ToolCall { id: "same".into(), name: "bash".into(), arguments: serde_json::json!({}), }; emitter .emit(TurnEvent::ToolStarted { iteration: 0, call: call.clone(), }) .unwrap(); assert_eq!( emitter.emit(TurnEvent::ToolStarted { iteration: 0, call }), Err(TurnEmitError::DuplicateTool("same".into())) ); assert_eq!(receiver.borrow().revision, 1); } #[test] fn parallel_tools_update_independently() { let (controller, emitter, _) = start(); for id in ["first", "second"] { emitter .emit(TurnEvent::ToolStarted { iteration: 0, call: ToolCall { id: id.into(), name: "bash".into(), arguments: serde_json::json!({"cmd": id}), }, }) .unwrap(); } emitter .emit(TurnEvent::ToolFinished { iteration: 0, call_id: "second".into(), success: false, preview: Some("failed".into()), }) .unwrap(); let snapshot = controller.snapshot(); assert!(matches!( &snapshot.blocks[0], TurnBlock::Tool { id, status: ToolStatus::Running, .. } if id == "first" )); assert!(matches!( &snapshot.blocks[1], TurnBlock::Tool { id, status: ToolStatus::Failed, .. } if id == "second" )); } #[test] fn empty_deltas_are_noops_and_cancel_reason_is_terminal() { let (controller, emitter, receiver) = start(); emitter .emit(TurnEvent::ReasoningDelta { iteration: 0, delta: String::new(), }) .unwrap(); emitter .emit(TurnEvent::TextDelta { iteration: 0, delta: String::new(), }) .unwrap(); assert_eq!(receiver.borrow().revision, 0); assert!(controller.cancel(Some("stopped by user".into()))); let snapshot = controller.snapshot(); assert_eq!(snapshot.status, TurnStatus::Cancelled); assert_eq!(snapshot.error.as_deref(), Some("stopped by user")); assert_eq!(snapshot.revision, 1); } #[test] fn cancelling_turn_marks_running_tools_cancelled() { let (controller, emitter, _receiver) = start(); emitter .emit(TurnEvent::ToolStarted { iteration: 0, call: ToolCall { id: "sleep-call".into(), name: "sleep".into(), arguments: serde_json::json!({"seconds": 60}), }, }) .unwrap(); assert!(controller.cancel(Some("stopped by user".into()))); let snapshot = controller.snapshot(); assert_eq!(snapshot.status, TurnStatus::Cancelled); assert!(matches!( &snapshot.blocks[0], TurnBlock::Tool { id, status: ToolStatus::Cancelled, .. } if id == "sleep-call" )); } #[test] fn failure_is_published_as_structured_terminal_state() { let (controller, _emitter, _) = start(); assert!(controller.fail("provider disconnected")); let snapshot = controller.snapshot(); assert_eq!(snapshot.status, TurnStatus::Failed); assert_eq!(snapshot.error.as_deref(), Some("provider disconnected")); assert_eq!(snapshot.phase, TurnPhase::Finalizing); } #[test] fn deactivated_emitter_drops_events_before_stale_worker_reduction() { let (controller, emitter, _) = start(); emitter.deactivate(); emitter .emit(TurnEvent::TextDelta { iteration: 0, delta: "late".into(), }) .unwrap(); assert_eq!(controller.snapshot().revision, 0); assert!(controller.snapshot().blocks.is_empty()); } }