use std::collections::HashMap;
use crate::protocol::{StepUpdate, StepUpdateSource, StepUpdateState, StepUpdateTarget};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[non_exhaustive]
pub enum StepKind {
Message,
ListDirectory,
FindFile,
SearchDirectory,
ViewFile,
CreateFile,
EditFile,
RunCommand,
Compaction,
InvokeSubagent,
GenerateImage,
SearchWeb,
ReadUrlContent,
McpTool,
CustomTool,
Finish,
Error,
ToolConfirmationRequest,
QuestionsRequest,
}
#[derive(Debug, Clone)]
pub struct Step {
pub trajectory_id: String,
pub step_index: u32,
pub state: StepUpdateState,
pub source: StepUpdateSource,
pub target: StepUpdateTarget,
pub kind: StepKind,
pub text: String,
pub thinking: String,
pub text_delta: String,
pub error_message: Option<String>,
pub update: StepUpdate,
}
impl Step {
pub fn id(&self) -> String {
format!("{}:{}", self.trajectory_id, self.step_index)
}
pub fn is_final(&self) -> bool {
matches!(self.state, StepUpdateState::Done | StepUpdateState::Error)
}
pub fn text(&self) -> Option<&str> {
Some(self.text.as_str()).filter(|t| !t.is_empty())
}
pub fn user_facing_text(&self) -> Option<&str> {
match self.target {
StepUpdateTarget::User => self.text(),
_ => None,
}
}
}
#[derive(Debug, Default)]
pub struct StepAssembler {
buffers: HashMap<(String, u32), Buffer>,
main_trajectory: Option<String>,
}
#[derive(Debug, Default)]
struct Buffer {
text: String,
thinking: String,
}
impl StepAssembler {
pub fn new(main_trajectory: Option<String>) -> Self {
Self {
buffers: HashMap::new(),
main_trajectory,
}
}
pub fn main_trajectory(&self) -> Option<&str> {
self.main_trajectory.as_deref()
}
pub fn is_main(&self, trajectory_id: &str) -> bool {
self.main_trajectory.as_deref() == Some(trajectory_id)
}
pub fn ingest(&mut self, update: StepUpdate) -> Step {
let trajectory_id = update.trajectory_id.clone().unwrap_or_default();
let step_index = update.step_index.unwrap_or_default();
if self.main_trajectory.is_none() && !trajectory_id.is_empty() {
self.main_trajectory = Some(trajectory_id.clone());
}
let buffer = self
.buffers
.entry((trajectory_id.clone(), step_index))
.or_default();
let text_delta = update.text_delta.clone().unwrap_or_default();
if !text_delta.is_empty() {
buffer.text.push_str(&text_delta);
}
if let Some(text) = update.text.as_deref().filter(|t| !t.is_empty()) {
buffer.text = text.to_string();
}
if let Some(delta) = update.thinking_delta.as_deref().filter(|t| !t.is_empty()) {
buffer.thinking.push_str(delta);
}
if let Some(thinking) = update.thinking.as_deref().filter(|t| !t.is_empty()) {
buffer.thinking = thinking.to_string();
}
let step = Step {
trajectory_id,
step_index,
state: update.state.clone().unwrap_or_default(),
source: update.source.clone().unwrap_or_default(),
target: update.target.clone().unwrap_or_default(),
kind: classify(&update),
text: buffer.text.clone(),
thinking: buffer.thinking.clone(),
text_delta,
error_message: update.error_message.clone().filter(|m| !m.is_empty()),
update,
};
if step.is_final() {
self.buffers
.remove(&(step.trajectory_id.clone(), step.step_index));
}
step
}
}
fn classify(update: &StepUpdate) -> StepKind {
if update.tool_confirmation_request.is_some() {
StepKind::ToolConfirmationRequest
} else if update.questions_request.is_some() {
StepKind::QuestionsRequest
} else if update.error.is_some() {
StepKind::Error
} else if update.finish.is_some() {
StepKind::Finish
} else if update.list_directory.is_some() {
StepKind::ListDirectory
} else if update.find_file.is_some() {
StepKind::FindFile
} else if update.search_directory.is_some() {
StepKind::SearchDirectory
} else if update.view_file.is_some() {
StepKind::ViewFile
} else if update.create_file.is_some() {
StepKind::CreateFile
} else if update.edit_file.is_some() {
StepKind::EditFile
} else if update.run_command.is_some() {
StepKind::RunCommand
} else if update.compaction.is_some() {
StepKind::Compaction
} else if update.invoke_subagent.is_some() {
StepKind::InvokeSubagent
} else if update.generate_image.is_some() {
StepKind::GenerateImage
} else if update.search_web.is_some() {
StepKind::SearchWeb
} else if update.read_url_content.is_some() {
StepKind::ReadUrlContent
} else if update.mcp_tool.is_some() {
StepKind::McpTool
} else if update.custom_tool.is_some() {
StepKind::CustomTool
} else {
StepKind::Message
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::protocol::{ActionRunCommand, UserQuestionsRequest};
fn delta(trajectory: &str, index: u32, text: &str, state: StepUpdateState) -> StepUpdate {
StepUpdate {
trajectory_id: Some(trajectory.into()),
step_index: Some(index),
state: Some(state),
text_delta: Some(text.into()),
..Default::default()
}
}
#[test]
fn deltas_accumulate_within_a_step() {
let mut a = StepAssembler::new(None);
a.ingest(delta("t1", 0, "Hel", StepUpdateState::Active));
let step = a.ingest(delta("t1", 0, "lo", StepUpdateState::Active));
assert_eq!(step.text, "Hello");
assert_eq!(step.text_delta, "lo");
assert!(!step.is_final());
}
#[test]
fn the_final_text_replaces_the_accumulated_deltas() {
let mut a = StepAssembler::new(None);
a.ingest(delta("t1", 0, "Hel", StepUpdateState::Active));
let step = a.ingest(StepUpdate {
trajectory_id: Some("t1".into()),
step_index: Some(0),
state: Some(StepUpdateState::Done),
text: Some("Hello, world".into()),
..Default::default()
});
assert_eq!(step.text, "Hello, world");
assert!(step.is_final());
}
#[test]
fn concurrent_trajectories_do_not_bleed_into_each_other() {
let mut a = StepAssembler::new(Some("main".into()));
a.ingest(delta("main", 0, "main-", StepUpdateState::Active));
a.ingest(delta("sub", 0, "sub-", StepUpdateState::Active));
let main = a.ingest(delta("main", 0, "text", StepUpdateState::Active));
let sub = a.ingest(delta("sub", 0, "text", StepUpdateState::Active));
assert_eq!(main.text, "main-text");
assert_eq!(sub.text, "sub-text");
assert!(a.is_main("main"));
assert!(!a.is_main("sub"));
}
#[test]
fn the_first_trajectory_seen_becomes_the_main_one() {
let mut a = StepAssembler::new(None);
a.ingest(delta("first", 0, "x", StepUpdateState::Active));
assert_eq!(a.main_trajectory(), Some("first"));
a.ingest(delta("second", 0, "y", StepUpdateState::Active));
assert_eq!(a.main_trajectory(), Some("first"));
}
#[test]
fn a_settled_step_releases_its_buffer() {
let mut a = StepAssembler::new(None);
a.ingest(delta("t1", 0, "hi", StepUpdateState::Done));
assert!(a.buffers.is_empty());
}
#[test]
fn actions_classify_by_their_populated_member() {
let update = StepUpdate {
run_command: Some(ActionRunCommand {
command_line: Some("ls".into()),
..Default::default()
}),
..Default::default()
};
assert_eq!(classify(&update), StepKind::RunCommand);
assert_eq!(classify(&StepUpdate::default()), StepKind::Message);
}
#[test]
fn a_pending_question_outranks_its_action() {
let update = StepUpdate {
run_command: Some(ActionRunCommand::default()),
questions_request: Some(UserQuestionsRequest::default()),
..Default::default()
};
assert_eq!(classify(&update), StepKind::QuestionsRequest);
}
#[test]
fn only_user_targeted_text_is_user_facing() {
let mut a = StepAssembler::new(None);
let to_model = a.ingest(StepUpdate {
trajectory_id: Some("t".into()),
target: Some(StepUpdateTarget::Model),
text: Some("echo of the prompt".into()),
..Default::default()
});
assert_eq!(to_model.user_facing_text(), None);
assert_eq!(to_model.text(), Some("echo of the prompt"));
}
}