PicoBot/src/session/turn.rs

648 lines
20 KiB
Rust

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<String>,
},
}
#[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<TurnBlock>,
pub usage: Option<Usage>,
pub error: Option<String>,
}
pub type TurnSnapshot = TurnState;
struct TurnControllerInner {
state: TurnState,
snapshots: watch::Sender<Arc<TurnSnapshot>>,
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<Usage>,
error: Option<String>,
) -> 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<Mutex<TurnControllerInner>>,
}
impl TurnController {
pub fn start(
session_id: impl Into<String>,
message_id: impl Into<String>,
) -> (Self, TurnEmitter, watch::Receiver<Arc<TurnSnapshot>>) {
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<Mutex<TurnControllerInner>> = 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<TurnSnapshot> {
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<Usage>) -> bool {
lock_unpoisoned(&self.inner).transition_terminal(TurnStatus::Completed, usage, None)
}
pub fn cancel(&self, reason: Option<String>) -> bool {
lock_unpoisoned(&self.inner).transition_terminal(TurnStatus::Cancelled, None, reason)
}
pub fn fail(&self, error: impl Into<String>) -> bool {
lock_unpoisoned(&self.inner).transition_terminal(
TurnStatus::Failed,
None,
Some(error.into()),
)
}
}
fn lock_unpoisoned<T>(mutex: &Mutex<T>) -> 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<Arc<TurnSnapshot>>,
) {
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());
}
}