use crate::context::{LoadedSkill, SkillMeta};
use crate::daemon::DaemonCommand;
use crate::db::{self, SessionRecord, write_session_retry, write_turn_retry};
use crate::providers::InferenceProvider;
use crate::requests::run_agent_loop;
use crate::tools::ToolRegistry;
use choreo_ai_protocols::model_reasoning_capability;
use choreo_proto::{
AssistantToolCallRecord, ContextConfig, DaemonMessage, DisplayedImageRecord, ReasoningArtifact,
ReasoningProducer, SessionStatus, SessionSummary, TimestampMs, TokenUsage, ToolResultRecord,
Turn,
};
use std::collections::{BTreeMap, HashMap, HashSet};
use std::io;
use std::path::PathBuf;
use std::sync::{Arc, mpsc};
use tracing::{debug, error, info, trace, warn};
use unicode_segmentation::UnicodeSegmentation;
pub(crate) const CANCEL_ALL: u32 = 0;
pub(crate) const MAX_TITLE_CHARS: usize = 200;
pub(crate) const SESSION_SHUTDOWN_GRACE: std::time::Duration = std::time::Duration::from_secs(5);
pub(crate) fn join_session_shutdown(handle: std::thread::JoinHandle<()>, session_id: u64) -> bool {
poll_join_with_grace(
handle,
session_id,
SESSION_SHUTDOWN_GRACE,
std::time::Instant::now,
std::thread::sleep,
)
}
fn poll_join_with_grace<F, S>(
handle: std::thread::JoinHandle<()>,
session_id: u64,
grace: std::time::Duration,
now: F,
sleep: S,
) -> bool
where
F: FnMut() -> std::time::Instant,
S: FnMut(std::time::Duration),
{
let handle = std::cell::RefCell::new(Some(handle));
shutdown_join_poll(
session_id,
grace,
|| handle.borrow().as_ref().is_some_and(|h| h.is_finished()),
|| {
if let Some(h) = handle.borrow_mut().take() {
let _ = h.join();
}
},
now,
sleep,
)
}
fn shutdown_join_poll(
session_id: u64,
grace: std::time::Duration,
mut finished: impl FnMut() -> bool,
mut reap: impl FnMut(),
mut now: impl FnMut() -> std::time::Instant,
mut sleep: impl FnMut(std::time::Duration),
) -> bool {
let deadline = now() + grace;
loop {
if finished() {
reap();
return true;
}
let remaining = deadline.saturating_duration_since(now());
if remaining.is_zero() {
tracing::warn!(
session_id,
grace_ms = grace.as_millis(),
"session thread did not exit within shutdown grace period; abandoning join \
(process exit will reap the thread; completed turns are already persisted)",
);
return false;
}
sleep(remaining.min(std::time::Duration::from_millis(50)));
}
}
#[cfg(feature = "test-utils")]
#[doc(hidden)]
pub fn join_session_shutdown_with_grace_for_test(
handle: std::thread::JoinHandle<()>,
session_id: u64,
grace: std::time::Duration,
) -> bool {
poll_join_with_grace(
handle,
session_id,
grace,
std::time::Instant::now,
std::thread::sleep,
)
}
pub enum SessionCommand {
RunInput {
request_id: u32,
input: Vec<u8>,
},
RunChildInput {
request_id: u32,
user_text: Option<String>,
reply: std::sync::mpsc::Sender<io::Result<ChildResult>>,
},
Cancel {
request_id: u32,
},
SetModel {
model: String,
},
StatusChanged(SessionStatus),
Attach {
client_id: u64,
tx: std::sync::mpsc::SyncSender<DaemonMessage>,
},
Detach {
client_id: u64,
},
GetSummary {
reply: std::sync::mpsc::Sender<SessionSummary>,
},
RequestFinished {
request_id: u32,
snapshot: SessionSnapshot,
},
Broadcast(DaemonMessage),
SetTitle {
title: String,
},
SetWorkingDir {
path: PathBuf,
reply: mpsc::Sender<Result<String, String>>,
},
LoadTools {
groups: Vec<String>,
reply: mpsc::Sender<Result<String, String>>,
},
UnloadTools {
groups: Vec<String>,
reply: mpsc::Sender<Result<String, String>>,
},
SetAccount {
name: String,
},
SetReasoningEffort {
effort: String,
},
GetReasoningEffort {
reply: mpsc::Sender<String>,
},
Undo,
Redo,
Shutdown,
}
#[derive(Clone)]
pub struct RequestContext {
pub cmd_tx: mpsc::Sender<SessionCommand>,
pub session_id: u64,
pub db: Arc<redb::Database>,
pub tool_registry: Arc<ToolRegistry>,
pub daemon_tx: mpsc::Sender<DaemonCommand>,
pub max_turns: u32,
}
pub struct ChildResult {
pub output: String,
pub is_error: bool,
}
#[derive(Debug, Clone)]
pub struct SessionMetadata {
pub title: Option<String>,
pub selected_model: Option<String>,
pub reasoning_effort: Option<String>,
pub parent_session_id: Option<u64>,
pub working_dir: Option<String>,
pub created_at: i64,
pub last_modified: i64,
pub turn_count: u32,
pub status: SessionStatus,
pub active_tool_groups: Vec<String>,
pub account_name: Option<String>,
pub accumulated_usage: TokenUsage,
pub context_window: Option<u32>,
pub last_prompt_tokens: Option<u32>,
}
impl From<SessionRecord> for SessionMetadata {
fn from(record: SessionRecord) -> Self {
let config = SessionConfig {
title: record.title,
selected_model: record.selected_model,
reasoning_effort: record.reasoning_effort,
parent_session_id: record.parent_session_id,
working_dir: record.working_dir.map(PathBuf::from),
created_at: record.created_at,
last_modified: record.last_modified,
status: SessionStatus::Sleeping,
active_tool_groups: record.active_tool_groups.into_iter().collect(),
context_config: record.context_config,
account_name: record.account_name,
accumulated_usage: TokenUsage::default(),
context_window: None,
last_prompt_tokens: None,
last_response_id: record.last_response_id,
last_response_id_producer: record.last_response_id_producer,
};
let mut meta = SessionMetadata::from(&config);
meta.turn_count = record.turn_count;
meta
}
}
impl From<SessionMetadata> for SessionRecord {
fn from(meta: SessionMetadata) -> Self {
SessionRecord {
title: meta.title,
selected_model: meta.selected_model,
reasoning_effort: meta.reasoning_effort,
parent_session_id: meta.parent_session_id,
working_dir: meta.working_dir,
turn_count: meta.turn_count,
created_at: meta.created_at,
last_modified: meta.last_modified,
active_tool_groups: meta.active_tool_groups,
context_config: ContextConfig::default(),
account_name: meta.account_name,
last_response_id: None,
last_response_id_producer: None,
}
}
}
impl From<&SessionState> for SessionMetadata {
fn from(state: &SessionState) -> Self {
let mut meta = SessionMetadata::from(&state.config);
meta.turn_count = state.turns.len() as u32;
meta
}
}
impl SessionMetadata {
pub fn to_summary(&self, session_id: u64) -> SessionSummary {
SessionSummary {
session_id,
title: self.title.clone(),
selected_model: self.selected_model.clone(),
reasoning_effort: self.reasoning_effort.clone(),
parent_session_id: self.parent_session_id,
working_dir: self.working_dir.clone(),
created_at: self.created_at,
last_modified: self.last_modified,
turn_count: self.turn_count,
status: self.status.clone(),
active_tool_groups: self.active_tool_groups.clone(),
account_name: self.account_name.clone(),
token_usage: Some(self.accumulated_usage),
context_window: self.context_window,
last_prompt_tokens: self.last_prompt_tokens,
}
}
}
#[derive(Clone, Debug)]
pub struct SessionConfig {
pub title: Option<String>,
pub selected_model: Option<String>,
pub reasoning_effort: Option<String>,
pub parent_session_id: Option<u64>,
pub working_dir: Option<PathBuf>,
pub created_at: i64,
pub last_modified: i64,
pub status: SessionStatus,
pub active_tool_groups: HashSet<String>,
pub context_config: ContextConfig,
pub account_name: Option<String>,
pub accumulated_usage: TokenUsage,
pub context_window: Option<u32>,
pub last_prompt_tokens: Option<u32>,
pub last_response_id: Option<String>,
pub last_response_id_producer: Option<ReasoningProducer>,
}
impl Default for SessionConfig {
fn default() -> Self {
Self {
title: None,
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
created_at: 0,
last_modified: 0,
status: SessionStatus::Inactive,
active_tool_groups: HashSet::new(),
context_config: ContextConfig::default(),
account_name: None,
accumulated_usage: TokenUsage::default(),
context_window: None,
last_prompt_tokens: None,
last_response_id: None,
last_response_id_producer: None,
}
}
}
impl SessionConfig {
fn apply_worker_snapshot(&mut self, snapshot: &SessionConfig) {
self.accumulated_usage = snapshot.accumulated_usage;
self.context_window = snapshot.context_window;
self.last_prompt_tokens = snapshot.last_prompt_tokens;
self.last_response_id = snapshot.last_response_id.clone();
self.last_response_id_producer = snapshot.last_response_id_producer.clone();
}
}
impl From<&SessionConfig> for SessionMetadata {
fn from(config: &SessionConfig) -> Self {
SessionMetadata {
title: config.title.clone(),
selected_model: config.selected_model.clone(),
reasoning_effort: config.reasoning_effort.clone(),
parent_session_id: config.parent_session_id,
working_dir: config.working_dir.as_ref().map(|p| p.display().to_string()),
created_at: config.created_at,
last_modified: config.last_modified,
turn_count: 0,
status: config.status.clone(),
active_tool_groups: config.active_tool_groups.iter().cloned().collect(),
account_name: config.account_name.clone(),
accumulated_usage: config.accumulated_usage,
context_window: config.context_window,
last_prompt_tokens: config.last_prompt_tokens,
}
}
}
impl From<&SessionState> for SessionRecord {
fn from(state: &SessionState) -> Self {
let meta: SessionMetadata = state.into();
let mut record: SessionRecord = meta.into();
record.context_config = state.config.context_config.clone();
record.last_response_id = state.config.last_response_id.clone();
record.last_response_id_producer = state.config.last_response_id_producer.clone();
record
}
}
#[derive(Clone)]
pub struct SessionSnapshot {
pub config: SessionConfig,
pub turns: BTreeMap<u32, Turn>,
pub loaded_skill_bodies: Vec<LoadedSkill>,
pub context_cache: Option<(u64, Arc<String>)>,
pub discovered_skills: Option<Vec<SkillMeta>>,
}
pub(crate) struct ActiveRequest {
pub(crate) cancel_tx: crossbeam_channel::Sender<()>,
pub(crate) turn_id: u32,
}
pub struct ActiveSessionEntry {
pub cmd_tx: mpsc::Sender<SessionCommand>,
pub handle: std::thread::JoinHandle<()>,
}
pub struct SessionState {
pub config: SessionConfig,
pub next_turn_id: u32,
last_undo_turn_ids: Option<Vec<u32>>,
pub turns: BTreeMap<u32, Turn>,
subscribers: HashMap<u64, std::sync::mpsc::SyncSender<DaemonMessage>>,
pub(crate) active_requests: BTreeMap<u32, ActiveRequest>,
pub provider: Option<InferenceProvider>,
pub loaded_skill_bodies: Vec<LoadedSkill>,
pub context_cache: Option<(u64, Arc<String>)>,
pub discovered_skills: Option<Vec<SkillMeta>>,
}
#[derive(Debug, Clone, Default)]
pub struct AssistantResponse {
pub text: Option<String>,
pub reasoning: Option<String>,
pub tool_calls: Vec<AssistantToolCallRecord>,
pub token_usage: Option<TokenUsage>,
pub reasoning_artifact: Option<ReasoningArtifact>,
pub reasoning_producer: Option<ReasoningProducer>,
}
impl SessionState {
fn resolve_context_window_if_missing(
&mut self,
session_id: u64,
daemon_tx: &mpsc::Sender<DaemonCommand>,
) {
if self.config.context_window.is_some() {
return;
}
let (Some(model), Some(provider)) = (&self.config.selected_model, &self.provider) else {
return;
};
if let Some(cw) = provider.resolve_context_window(model) {
debug!(
"session {}: re-resolved context_window={} for model={}",
session_id, cw, model
);
self.config.context_window = Some(cw);
broadcast(
&mut self.subscribers,
daemon_tx,
DaemonMessage::ContextWindowResolved {
session_id,
context_window: cw,
},
);
}
}
fn snapshot(&self) -> SessionSnapshot {
SessionSnapshot {
config: self.config.clone(),
turns: self.turns.clone(),
loaded_skill_bodies: self.loaded_skill_bodies.clone(),
context_cache: self.context_cache.clone(),
discovered_skills: self.discovered_skills.clone(),
}
}
fn from_snapshot(
snapshot: SessionSnapshot,
subscribers: HashMap<u64, std::sync::mpsc::SyncSender<DaemonMessage>>,
) -> Self {
let turn_count = snapshot.turns.len() as u32;
Self {
config: snapshot.config,
next_turn_id: turn_count,
last_undo_turn_ids: None,
turns: snapshot.turns,
subscribers,
active_requests: BTreeMap::new(),
provider: None,
loaded_skill_bodies: snapshot.loaded_skill_bodies,
context_cache: snapshot.context_cache,
discovered_skills: snapshot.discovered_skills,
}
}
pub(crate) fn session_state_message(&self, session_id: u64) -> DaemonMessage {
let reasoning_capability = self.config.selected_model.as_ref().and_then(|model| {
let slug = self.provider.as_ref()?.provider_slug();
Some(model_reasoning_capability(slug, model))
});
DaemonMessage::SessionState {
session_id,
title: self.config.title.clone(),
selected_model: self.config.selected_model.clone(),
parent_session_id: self.config.parent_session_id,
working_dir: self
.config
.working_dir
.as_ref()
.map(|p| p.display().to_string()),
turns: self
.turns
.iter()
.map(|(&turn_id, turn)| (turn_id, turn_for_client(turn)))
.collect(),
active_tool_groups: self.config.active_tool_groups.iter().cloned().collect(),
token_usage: Some(self.config.accumulated_usage),
context_window: self.config.context_window,
last_prompt_tokens: self.config.last_prompt_tokens,
status: self.config.status.clone(),
reasoning_effort: self.config.reasoning_effort.clone(),
reasoning_capability,
}
}
pub fn start_turn(&mut self, user_text: Option<String>) -> (u32, Turn) {
if user_text.is_some() {
self.last_undo_turn_ids = None;
}
let turn_id = self.next_turn_id;
self.next_turn_id += 1;
let turn = Turn {
created_at: TimestampMs::now(),
undone: false,
error: None,
user_text,
assistant_text: None,
assistant_reasoning: None,
tool_calls: Vec::new(),
token_usage: None,
tool_results: Vec::new(),
displayed_images: Vec::new(),
reasoning_artifact: None,
reasoning_producer: None,
};
self.turns.insert(turn_id, turn.clone());
(turn_id, turn)
}
pub fn set_assistant_response(&mut self, turn_id: u32, response: AssistantResponse) {
if let Some(turn) = self.turns.get_mut(&turn_id) {
turn.assistant_text = response.text;
turn.assistant_reasoning = response.reasoning;
turn.tool_calls = response.tool_calls;
turn.token_usage = response.token_usage;
turn.reasoning_artifact = response.reasoning_artifact;
turn.reasoning_producer = response.reasoning_producer;
}
}
pub fn seed_tool_results(&mut self, turn_id: u32, tool_calls: &[AssistantToolCallRecord]) {
if let Some(turn) = self.turns.get_mut(&turn_id) {
turn.tool_results = tool_calls
.iter()
.map(|tc| ToolResultRecord {
call_id: tc.call_id.clone(),
name: tc.name.clone(),
content: String::new(),
is_error: false,
invocation_description: String::new(),
})
.collect();
}
}
pub fn update_tool_result(
&mut self,
turn_id: u32,
call_id: &str,
name: String,
content: String,
is_error: bool,
invocation_description: String,
) {
if let Some(turn) = self.turns.get_mut(&turn_id)
&& let Some(record) = turn.tool_results.iter_mut().find(|r| r.call_id == call_id)
{
record.name = name;
record.content = content;
record.is_error = is_error;
record.invocation_description = invocation_description;
}
}
pub fn mark_unexecuted_tool_results(&mut self, turn_id: u32, executed: &HashSet<String>) {
if let Some(turn) = self.turns.get_mut(&turn_id) {
for record in &mut turn.tool_results {
if !executed.contains(&record.call_id) {
record.content = "[cancelled — result not recorded]".to_string();
record.is_error = true;
}
}
}
}
pub fn add_displayed_image(&mut self, turn_id: u32, record: DisplayedImageRecord) {
if let Some(turn) = self.turns.get_mut(&turn_id) {
turn.displayed_images.push(record);
}
}
pub fn set_turn_error(&mut self, turn_id: u32, error: String) {
if let Some(turn) = self.turns.get_mut(&turn_id) {
turn.error = Some(error);
}
}
pub fn finalize_turn(
&mut self,
db: &redb::Database,
session_id: u64,
turn_id: u32,
) -> io::Result<()> {
if let Some(turn) = self.turns.get(&turn_id) {
write_turn_retry(db, session_id, turn_id, turn)
.map_err(|e| io::Error::other(format!("failed to persist turn {turn_id}: {e}")))?;
}
Ok(())
}
pub fn undo_turns(&mut self) -> Option<Vec<u32>> {
let target = self
.turns
.iter()
.rev()
.find(|(_, t)| !t.undone && t.user_text.is_some())
.map(|(&id, _)| id)?;
let to_undo: Vec<u32> = self.turns.range(target..).map(|(&id, _)| id).collect();
for &id in &to_undo {
if let Some(turn) = self.turns.get_mut(&id) {
turn.undone = true;
}
}
self.last_undo_turn_ids = Some(to_undo.clone());
Some(to_undo)
}
pub fn redo_turns(&mut self) -> Option<BTreeMap<u32, Turn>> {
let ids = self.last_undo_turn_ids.take()?;
let mut restored = BTreeMap::new();
for &id in &ids {
if let Some(turn) = self.turns.get_mut(&id) {
turn.undone = false;
restored.insert(id, turn.clone());
}
}
Some(restored)
}
pub fn empty() -> Self {
Self {
config: SessionConfig::default(),
next_turn_id: 0,
last_undo_turn_ids: None,
turns: BTreeMap::new(),
subscribers: HashMap::new(),
active_requests: BTreeMap::new(),
provider: None,
loaded_skill_bodies: Vec::new(),
context_cache: None,
discovered_skills: None,
}
}
}
pub(crate) fn turn_for_client(turn: &Turn) -> Turn {
let mut clone = turn.clone();
clone.reasoning_artifact = None;
clone.reasoning_producer = None;
clone
}
fn broadcast(
subscribers: &mut HashMap<u64, std::sync::mpsc::SyncSender<DaemonMessage>>,
daemon_tx: &mpsc::Sender<DaemonCommand>,
message: DaemonMessage,
) {
let _ = daemon_tx.send(DaemonCommand::BroadcastActivity(message.clone()));
subscribers.retain(|client_id, tx| {
crate::broadcast::try_send_keep_on_full(tx, *client_id, "session", &message)
});
}
fn fail_request(
subscribers: &mut HashMap<u64, std::sync::mpsc::SyncSender<DaemonMessage>>,
daemon_tx: &mpsc::Sender<DaemonCommand>,
session_id: u64,
request_id: u32,
error: impl Into<String>,
) -> bool {
broadcast(
subscribers,
daemon_tx,
DaemonMessage::Started {
session_id,
request_id,
turn_id: 0,
estimated_prompt_tokens: 0,
},
);
broadcast(
subscribers,
daemon_tx,
DaemonMessage::Failed {
session_id,
request_id,
error: error.into(),
},
);
false
}
fn persist_session_metadata(state: &mut SessionState, ctx: &RequestContext, label: &str) {
let now = TimestampMs::now().as_millis();
state.config.last_modified = state.config.last_modified.max(now);
let _ = ctx.daemon_tx.send(DaemonCommand::UpdateMetadata {
session_id: ctx.session_id,
metadata: SessionMetadata::from(&*state),
});
let record = SessionRecord::from(&*state);
if let Err(e) = write_session_retry(&ctx.db, ctx.session_id, &record) {
warn!(error = %e, "failed to persist session record after {label}");
}
}
pub fn session_main(
rx: std::sync::mpsc::Receiver<SessionCommand>,
provider: Option<InferenceProvider>,
account_name: Option<String>,
init_record: Option<SessionRecord>,
ctx: RequestContext,
) {
let config = SessionConfig {
title: init_record.as_ref().and_then(|r| r.title.clone()),
selected_model: init_record.as_ref().and_then(|r| r.selected_model.clone()),
reasoning_effort: init_record
.as_ref()
.and_then(|r| r.reasoning_effort.clone()),
parent_session_id: init_record.as_ref().and_then(|r| r.parent_session_id),
working_dir: init_record
.as_ref()
.and_then(|r| r.working_dir.as_ref().map(PathBuf::from)),
created_at: init_record
.as_ref()
.map(|r| r.created_at)
.unwrap_or_else(|| TimestampMs::now().as_millis()),
last_modified: init_record
.as_ref()
.map(|r| r.last_modified)
.unwrap_or_else(|| TimestampMs::now().as_millis()),
status: SessionStatus::Inactive,
active_tool_groups: init_record
.as_ref()
.map(|r| r.active_tool_groups.iter().cloned().collect())
.filter(|cats: &HashSet<String>| !cats.is_empty())
.unwrap_or_else(|| {
HashSet::from(["core".to_string(), "git".to_string(), "shell".to_string()])
}),
context_config: init_record
.as_ref()
.map(|r| r.context_config.clone())
.unwrap_or_default(),
account_name,
accumulated_usage: TokenUsage::default(),
context_window: None,
last_prompt_tokens: None,
last_response_id: init_record
.as_ref()
.and_then(|r| r.last_response_id.clone()),
last_response_id_producer: init_record
.as_ref()
.and_then(|r| r.last_response_id_producer.clone()),
};
let mut state = SessionState {
config,
provider,
..SessionState::empty()
};
state.resolve_context_window_if_missing(ctx.session_id, &ctx.daemon_tx);
match db::read_turns(&ctx.db, ctx.session_id) {
Ok(turns) => {
for (turn_id, turn) in turns {
state.turns.insert(turn_id, turn);
state.next_turn_id = state.next_turn_id.max(turn_id + 1);
}
let mut accumulated_usage = TokenUsage::default();
let mut last_prompt_tokens = None;
for turn in state.turns.values() {
if let Some(u) = turn.token_usage {
accumulated_usage.input_tokens += u.input_tokens;
accumulated_usage.output_tokens += u.output_tokens;
accumulated_usage.total_tokens += u.total_tokens;
last_prompt_tokens = Some(u.input_tokens);
}
}
state.config.accumulated_usage = accumulated_usage;
state.config.last_prompt_tokens = last_prompt_tokens;
trace!(
last_prompt_tokens,
?accumulated_usage,
"reconstructed token state from turns after daemon restart"
);
}
Err(e) => warn!(ctx.session_id, error = %e, "failed to load turns from DB"),
}
let _ = ctx.daemon_tx.send(DaemonCommand::UpdateMetadata {
session_id: ctx.session_id,
metadata: SessionMetadata::from(&state),
});
info!("session {} started", ctx.session_id);
let mut shutdown_requested = false;
while let Ok(cmd) = rx.recv() {
if process_command(cmd, &mut state, &mut shutdown_requested, &ctx) {
break;
}
}
info!("session {} exiting", ctx.session_id);
persist_and_exit(&state, &ctx.db, ctx.session_id, &ctx.daemon_tx);
}
fn process_command(
cmd: SessionCommand,
state: &mut SessionState,
shutdown_requested: &mut bool,
ctx: &RequestContext,
) -> bool {
match cmd {
SessionCommand::RunInput { request_id, input } => {
handle_run_input(request_id, input, state, shutdown_requested, ctx)
}
SessionCommand::RunChildInput {
request_id,
user_text,
reply,
} => handle_run_child_input(request_id, user_text, reply, state, shutdown_requested, ctx),
SessionCommand::Cancel { request_id } => handle_cancel(request_id, state, ctx),
SessionCommand::SetModel { model } => handle_set_model(model, state, ctx),
SessionCommand::StatusChanged(new_status) => handle_status_changed(new_status, state, ctx),
SessionCommand::Attach { client_id, tx } => handle_attach(client_id, tx, state, ctx),
SessionCommand::Detach { client_id } => {
handle_detach(client_id, state, shutdown_requested, ctx)
}
SessionCommand::GetSummary { reply } => handle_get_summary(reply, state, ctx),
SessionCommand::RequestFinished {
request_id,
snapshot,
} => handle_request_finished(request_id, snapshot, state, shutdown_requested, ctx),
SessionCommand::Broadcast(message) => handle_broadcast(message, state, ctx),
SessionCommand::SetTitle { title } => handle_set_title(title, state, ctx),
SessionCommand::SetWorkingDir { path, reply } => {
handle_set_working_dir(path, reply, state, ctx)
}
SessionCommand::LoadTools { groups, reply } => handle_load_tools(groups, reply, state, ctx),
SessionCommand::UnloadTools { groups, reply } => {
handle_unload_tools(groups, reply, state, ctx)
}
SessionCommand::SetAccount { name } => handle_set_account(name, state, ctx),
SessionCommand::SetReasoningEffort { effort } => {
handle_set_reasoning_effort(effort, state, ctx)
}
SessionCommand::GetReasoningEffort { reply } => {
handle_get_reasoning_effort(reply, state, ctx)
}
SessionCommand::Undo => handle_undo(state, ctx),
SessionCommand::Redo => handle_redo(state, ctx),
SessionCommand::Shutdown => handle_shutdown(state, shutdown_requested, ctx),
}
}
fn handle_run_input(
request_id: u32,
input: Vec<u8>,
state: &mut SessionState,
shutdown_requested: &mut bool,
ctx: &RequestContext,
) -> bool {
debug!("session {}: RunInput id={}", ctx.session_id, request_id);
let text = String::from_utf8_lossy(&input).trim().to_string();
info!(
session_id = ctx.session_id,
input_len = text.len(),
input_preview = %text.chars().take(120).collect::<String>(),
"session received input",
);
if text.is_empty() {
return fail_request(
&mut state.subscribers,
&ctx.daemon_tx,
ctx.session_id,
request_id,
"empty input",
);
}
let provider = if let Some(p) = state.provider.as_ref() {
p.clone()
} else if let Some(ref name) = state.config.account_name {
let (reply, rx) = mpsc::channel();
let _ = ctx.daemon_tx.send(DaemonCommand::ResolveProviderCmd {
account: name.clone(),
reply,
});
match rx.recv() {
Ok(Some(provider)) => {
state.provider = Some(provider);
state.resolve_context_window_if_missing(ctx.session_id, &ctx.daemon_tx);
let Some(p) = state.provider.as_ref() else {
return fail_request(
&mut state.subscribers,
&ctx.daemon_tx,
ctx.session_id,
request_id,
"internal error: provider not set after resolution".to_string(),
);
};
p.clone()
}
_ => {
return fail_request(
&mut state.subscribers,
&ctx.daemon_tx,
ctx.session_id,
request_id,
format!(
"no credential stored for account '{name}' — add one via the AI Providers page or /add-key"
),
);
}
}
} else {
return fail_request(
&mut state.subscribers,
&ctx.daemon_tx,
ctx.session_id,
request_id,
"no account configured on this session — use /account <name> to set one",
);
};
let model = match &state.config.selected_model {
Some(m) => m.clone(),
None => {
return fail_request(
&mut state.subscribers,
&ctx.daemon_tx,
ctx.session_id,
request_id,
"no model selected",
);
}
};
if *shutdown_requested {
return fail_request(
&mut state.subscribers,
&ctx.daemon_tx,
ctx.session_id,
request_id,
"session is shutting down",
);
}
if !state.active_requests.is_empty() {
return fail_request(
&mut state.subscribers,
&ctx.daemon_tx,
ctx.session_id,
request_id,
"session already has an active request",
);
}
broadcast(
&mut state.subscribers,
&ctx.daemon_tx,
DaemonMessage::Started {
session_id: ctx.session_id,
request_id,
turn_id: state.next_turn_id,
estimated_prompt_tokens: 0,
},
);
let (cancel_tx, cancel_rx) = crossbeam_channel::unbounded::<()>();
state.active_requests.insert(
request_id,
ActiveRequest {
cancel_tx,
turn_id: state.next_turn_id,
},
);
let mut worker_session = SessionState::from_snapshot(state.snapshot(), HashMap::new());
let ctx = ctx.clone();
let user_text = Some(text);
std::thread::spawn(move || {
let _ = run_request_worker(
request_id,
provider,
&mut worker_session,
model,
cancel_rx,
ctx,
None,
user_text,
);
});
false
}
fn handle_run_child_input(
request_id: u32,
user_text: Option<String>,
reply: std::sync::mpsc::Sender<io::Result<ChildResult>>,
state: &mut SessionState,
shutdown_requested: &mut bool,
ctx: &RequestContext,
) -> bool {
let Some(provider) = state.provider.as_ref() else {
let _ = reply.send(Err(io::Error::other("daemon locked")));
return false;
};
let model = state.config.selected_model.clone().unwrap_or_default();
if *shutdown_requested {
let _ = reply.send(Err(io::Error::other("session is shutting down")));
return false;
}
if !state.active_requests.is_empty() {
let _ = reply.send(Err(io::Error::other(
"session already has an active request",
)));
return false;
}
broadcast(
&mut state.subscribers,
&ctx.daemon_tx,
DaemonMessage::Started {
session_id: ctx.session_id,
request_id,
turn_id: state.next_turn_id,
estimated_prompt_tokens: 0,
},
);
let (cancel_tx, cancel_rx) = crossbeam_channel::unbounded::<()>();
state.active_requests.insert(
request_id,
ActiveRequest {
cancel_tx,
turn_id: state.next_turn_id,
},
);
let mut worker_session = SessionState::from_snapshot(state.snapshot(), HashMap::new());
let ctx = ctx.clone();
let provider = provider.clone();
std::thread::spawn(move || {
let result = run_request_worker(
request_id,
provider,
&mut worker_session,
model,
cancel_rx,
ctx,
Some(reply),
user_text,
);
let _ = result;
});
false
}
fn handle_cancel(request_id: u32, state: &mut SessionState, ctx: &RequestContext) -> bool {
let targets: Vec<u32> = if request_id == 0 {
state.active_requests.keys().copied().collect()
} else {
vec![request_id]
};
for rid in targets {
if let Some(active) = state.active_requests.get(&rid) {
let _ = active.cancel_tx.send(());
broadcast(
&mut state.subscribers,
&ctx.daemon_tx,
DaemonMessage::Cancelled {
session_id: ctx.session_id,
request_id: rid,
},
);
}
}
false
}
fn handle_set_model(model: String, state: &mut SessionState, ctx: &RequestContext) -> bool {
info!("session {}: SetModel model={}", ctx.session_id, model);
if let Err(msg) = validate_model_via_daemon(&model, ctx) {
warn!(
"session {}: model '{model}' rejected: {msg}",
ctx.session_id
);
broadcast(
&mut state.subscribers,
&ctx.daemon_tx,
DaemonMessage::ModelSelectionFailed {
session_id: ctx.session_id,
model,
error: msg,
},
);
return false;
}
state.config.selected_model = Some(model.clone());
let cw = state
.provider
.as_ref()
.and_then(|p| p.resolve_context_window(&model));
debug!(
"session {}: resolved context_window={:?} for model={}",
ctx.session_id, cw, model
);
state.config.context_window = cw;
if let Some(cw) = cw {
broadcast(
&mut state.subscribers,
&ctx.daemon_tx,
DaemonMessage::ContextWindowResolved {
session_id: ctx.session_id,
context_window: cw,
},
);
}
let capability = state
.provider
.as_ref()
.map(|p| model_reasoning_capability(p.provider_slug(), &model));
if let Some(ref cap) = capability
&& let Some(ref effort) = state.config.reasoning_effort
&& effort != "off"
&& !cap.available_effort_levels.iter().any(|l| l == effort)
{
warn!(
session_id = ctx.session_id,
old_effort = %effort,
"reasoning effort not supported by new model, resetting to 'off'",
);
state.config.reasoning_effort = Some("off".to_string());
broadcast(
&mut state.subscribers,
&ctx.daemon_tx,
DaemonMessage::ReasoningEffortSet {
session_id: ctx.session_id,
effort: "off".to_string(),
},
);
}
debug!(
"session {}: broadcasting ModelSelected model={}",
ctx.session_id, model
);
broadcast(
&mut state.subscribers,
&ctx.daemon_tx,
DaemonMessage::ModelSelected {
session_id: ctx.session_id,
model: model.clone(),
reasoning_capability: capability,
},
);
persist_session_metadata(state, ctx, "SetModel");
false
}
fn validate_model_via_daemon(model: &str, ctx: &RequestContext) -> Result<(), String> {
let (reply, rx) = mpsc::channel();
if ctx
.daemon_tx
.send(DaemonCommand::ValidateModel {
session_id: ctx.session_id,
model: model.to_string(),
reply,
})
.is_err()
{
warn!(
"session {}: daemon disconnected during model validation for '{model}'",
ctx.session_id
);
return Ok(());
}
match rx.recv() {
Ok(Ok(())) => Ok(()),
Ok(Err(msg)) => Err(msg),
Err(_) => {
warn!(
"session {}: daemon disconnected while waiting for model validation \
of '{model}', allowing through",
ctx.session_id
);
Ok(())
}
}
}
fn handle_status_changed(
new_status: SessionStatus,
state: &mut SessionState,
ctx: &RequestContext,
) -> bool {
state.config.status = new_status.clone();
let last_modified = state.config.last_modified;
broadcast(
&mut state.subscribers,
&ctx.daemon_tx,
DaemonMessage::SessionStatusChanged {
session_id: ctx.session_id,
status: new_status.clone(),
last_modified,
},
);
let _ = ctx.daemon_tx.send(DaemonCommand::BroadcastSessionStatus {
session_id: ctx.session_id,
status: new_status,
});
false
}
fn handle_attach(
client_id: u64,
tx: std::sync::mpsc::SyncSender<DaemonMessage>,
state: &mut SessionState,
ctx: &RequestContext,
) -> bool {
info!("session {}: client {} attached", ctx.session_id, client_id);
state.subscribers.insert(client_id, tx);
let _ = ctx.daemon_tx.send(DaemonCommand::TrackSessionSubscription {
client_id,
session_id: ctx.session_id,
});
if !state.active_requests.is_empty()
&& let Some(tx) = state.subscribers.get(&client_id)
{
for (&request_id, active) in &state.active_requests {
let _ = tx.try_send(DaemonMessage::Started {
session_id: ctx.session_id,
request_id,
turn_id: active.turn_id,
estimated_prompt_tokens: 0,
});
}
}
let snapshot = state.session_state_message(ctx.session_id);
if let Some(tx) = state.subscribers.get(&client_id) {
match tx.try_send(snapshot) {
Ok(()) => {}
Err(mpsc::TrySendError::Full(_)) => {
warn!(
"session {}: dropped attach snapshot for client {}: writer buffer full",
ctx.session_id, client_id
);
crate::metrics::record_broadcast_dropped("attach");
}
Err(mpsc::TrySendError::Disconnected(_)) => {
debug!(
"session {}: attach snapshot for client {} dropped: receiver gone",
ctx.session_id, client_id
);
}
}
}
false
}
fn handle_detach(
client_id: u64,
state: &mut SessionState,
shutdown_requested: &bool,
ctx: &RequestContext,
) -> bool {
info!("session {}: client {} detached", ctx.session_id, client_id);
state.subscribers.remove(&client_id);
let _ = ctx
.daemon_tx
.send(DaemonCommand::UntrackSessionSubscription {
client_id,
session_id: ctx.session_id,
});
state.active_requests.is_empty() && (state.subscribers.is_empty() || *shutdown_requested)
}
fn handle_get_summary(
reply: std::sync::mpsc::Sender<SessionSummary>,
state: &SessionState,
ctx: &RequestContext,
) -> bool {
let _ = reply.send(SessionSummary {
session_id: ctx.session_id,
title: state.config.title.clone(),
selected_model: state.config.selected_model.clone(),
reasoning_effort: state.config.reasoning_effort.clone(),
parent_session_id: state.config.parent_session_id,
working_dir: state
.config
.working_dir
.as_ref()
.map(|p| p.display().to_string()),
created_at: state.config.created_at,
last_modified: state.config.last_modified,
turn_count: state.turns.len() as u32,
status: state.config.status.clone(),
active_tool_groups: state.config.active_tool_groups.iter().cloned().collect(),
account_name: state.config.account_name.clone(),
token_usage: Some(state.config.accumulated_usage),
context_window: state.config.context_window,
last_prompt_tokens: state.config.last_prompt_tokens,
});
false
}
fn handle_request_finished(
request_id: u32,
mut snapshot: SessionSnapshot,
state: &mut SessionState,
shutdown_requested: &bool,
ctx: &RequestContext,
) -> bool {
let undo_during_request = snapshot.turns.iter().any(|(&turn_id, snap_turn)| {
state
.turns
.get(&turn_id)
.is_some_and(|state_turn| state_turn.undone && !snap_turn.undone)
});
if undo_during_request {
debug!(
session_id = ctx.session_id,
request_id,
"undo landed while request was in flight; dropping stale response-id chain from worker snapshot",
);
snapshot.config.last_response_id = None;
snapshot.config.last_response_id_producer = None;
}
state.config.apply_worker_snapshot(&snapshot.config);
let last_modified = TimestampMs::now().as_millis();
state.config.last_modified = state.config.last_modified.max(last_modified);
let record = SessionRecord::from(&*state);
if let Err(e) = write_session_retry(&ctx.db, ctx.session_id, &record) {
warn!(error = %e, "failed to persist session config after request");
}
state.loaded_skill_bodies = snapshot.loaded_skill_bodies;
state.context_cache = snapshot.context_cache;
state.discovered_skills = snapshot.discovered_skills;
for (&turn_id, turn) in &snapshot.turns {
let is_new = !state.turns.contains_key(&turn_id);
if !is_new
&& let Some(state_turn) = state.turns.get(&turn_id)
&& state_turn.undone
&& !turn.undone
{
continue;
}
state.turns.insert(turn_id, turn.clone());
if is_new {
if let Err(e) = write_turn_retry(&ctx.db, ctx.session_id, turn_id, turn) {
tracing::warn!(turn_id, error = %e, "failed to persist turn");
}
} else if state.turns.get(&turn_id).is_some_and(|t| t != turn) {
if let Err(e) = write_turn_retry(&ctx.db, ctx.session_id, turn_id, turn) {
tracing::warn!(turn_id, error = %e, "failed to persist updated turn");
}
}
}
if let Some(max_id) = snapshot.turns.keys().max() {
state.next_turn_id = state.next_turn_id.max(max_id + 1);
}
state.active_requests.remove(&request_id);
state.config.status = SessionStatus::Inactive;
let _ = ctx.daemon_tx.send(DaemonCommand::UpdateMetadata {
session_id: ctx.session_id,
metadata: SessionMetadata::from(&*state),
});
broadcast(
&mut state.subscribers,
&ctx.daemon_tx,
DaemonMessage::SessionStatusChanged {
session_id: ctx.session_id,
status: SessionStatus::Inactive,
last_modified,
},
);
let _ = ctx.daemon_tx.send(DaemonCommand::BroadcastSessionStatus {
session_id: ctx.session_id,
status: SessionStatus::Inactive,
});
state.active_requests.is_empty() && (state.subscribers.is_empty() || *shutdown_requested)
}
fn handle_broadcast(
message: DaemonMessage,
state: &mut SessionState,
ctx: &RequestContext,
) -> bool {
broadcast(&mut state.subscribers, &ctx.daemon_tx, message);
false
}
fn handle_set_title(title: String, state: &mut SessionState, ctx: &RequestContext) -> bool {
if title.graphemes(true).count() > MAX_TITLE_CHARS {
warn!(
session_id = ctx.session_id,
length = title.graphemes(true).count(),
max = MAX_TITLE_CHARS,
"rejecting SetTitle: title too long (defense-in-depth)",
);
return false;
}
info!(
session_id = ctx.session_id,
old_title = ?state.config.title,
new_title = %title,
"session title changed",
);
state.config.title = Some(title.clone());
broadcast(
&mut state.subscribers,
&ctx.daemon_tx,
DaemonMessage::SessionTitleSet {
session_id: ctx.session_id,
title: title.clone(),
},
);
persist_session_metadata(state, ctx, "SetTitle");
false
}
fn handle_set_working_dir(
path: PathBuf,
reply: mpsc::Sender<Result<String, String>>,
state: &mut SessionState,
ctx: &RequestContext,
) -> bool {
info!(
session_id = ctx.session_id,
old_path = ?state.config.working_dir,
new_path = %path.display(),
"session working directory changed",
);
state.config.working_dir = Some(path.clone());
state.discovered_skills = None;
broadcast(
&mut state.subscribers,
&ctx.daemon_tx,
DaemonMessage::SessionWorkingDirSet {
session_id: ctx.session_id,
path: Some(path.to_string_lossy().into_owned()),
},
);
persist_session_metadata(state, ctx, "SetWorkingDir");
let _ = reply.send(Ok(path.to_string_lossy().into_owned()));
false
}
fn handle_load_tools(
groups: Vec<String>,
reply: mpsc::Sender<Result<String, String>>,
state: &mut SessionState,
ctx: &RequestContext,
) -> bool {
info!(session_id = ctx.session_id, groups = ?groups, "session load_tools");
let known = ctx.tool_registry.known_group_names();
if let Some(unknown) = crate::tools::unknown_group_names(&groups, &known) {
let _ = reply.send(Err(format!(
"Unknown tool group(s): {}",
unknown.join(", ")
)));
return false;
}
let result =
crate::tools::load_tools::apply_load_tools(&mut state.config.active_tool_groups, &groups);
let session_state = state.session_state_message(ctx.session_id);
broadcast(&mut state.subscribers, &ctx.daemon_tx, session_state);
persist_session_metadata(state, ctx, "LoadTools");
let _ = reply.send(Ok(result));
false
}
fn handle_unload_tools(
groups: Vec<String>,
reply: mpsc::Sender<Result<String, String>>,
state: &mut SessionState,
ctx: &RequestContext,
) -> bool {
info!(session_id = ctx.session_id, groups = ?groups, "session unload_tools");
let known = ctx.tool_registry.known_group_names();
if let Some(unknown) = crate::tools::unknown_group_names(&groups, &known) {
let _ = reply.send(Err(format!(
"Unknown tool group(s): {}",
unknown.join(", ")
)));
return false;
}
let result = crate::tools::unload_tools::apply_unload_tools(
&mut state.config.active_tool_groups,
&groups,
);
let session_state = state.session_state_message(ctx.session_id);
broadcast(&mut state.subscribers, &ctx.daemon_tx, session_state);
persist_session_metadata(state, ctx, "UnloadTools");
let _ = reply.send(Ok(result));
false
}
fn handle_set_account(name: String, state: &mut SessionState, ctx: &RequestContext) -> bool {
info!("session {}: SetAccount account={}", ctx.session_id, name);
let (reply, rx) = mpsc::channel();
let _ = ctx.daemon_tx.send(DaemonCommand::ResolveProviderCmd {
account: name.clone(),
reply,
});
if let Ok(Some(provider)) = rx.recv() {
if let Some(ref model) = state.config.selected_model {
let cw = provider.resolve_context_window(model);
debug!(
"session {}: re-resolved context_window={:?} after account change for model={}",
ctx.session_id, cw, model
);
state.config.context_window = cw;
if let Some(cw) = cw {
broadcast(
&mut state.subscribers,
&ctx.daemon_tx,
DaemonMessage::ContextWindowResolved {
session_id: ctx.session_id,
context_window: cw,
},
);
}
}
state.provider = Some(provider);
}
state.config.account_name = Some(name.clone());
broadcast(
&mut state.subscribers,
&ctx.daemon_tx,
DaemonMessage::SessionAccountSet {
session_id: ctx.session_id,
account: name,
},
);
persist_session_metadata(state, ctx, "SetAccount");
false
}
fn handle_set_reasoning_effort(
effort: String,
state: &mut SessionState,
ctx: &RequestContext,
) -> bool {
if effort.len() > 64 {
let msg = format!("reasoning effort slug too long ({} bytes)", effort.len());
warn!(session_id = ctx.session_id, error = %msg, "reasoning effort rejected");
broadcast(
&mut state.subscribers,
&ctx.daemon_tx,
DaemonMessage::ReasoningEffortSetFailed {
session_id: ctx.session_id,
effort,
error: msg,
},
);
return false;
}
let capability = state.config.selected_model.as_ref().and_then(|model| {
let slug = state.provider.as_ref()?.provider_slug();
Some(model_reasoning_capability(slug, model))
});
let valid = effort == "off"
|| capability
.as_ref()
.map(|c| c.available_effort_levels.contains(&effort))
.unwrap_or(false)
|| capability.is_none();
if valid {
state.config.reasoning_effort = Some(effort.clone());
info!(
session_id = ctx.session_id,
effort = %effort,
model = ?state.config.selected_model,
"reasoning effort set",
);
broadcast(
&mut state.subscribers,
&ctx.daemon_tx,
DaemonMessage::ReasoningEffortSet {
session_id: ctx.session_id,
effort,
},
);
return false;
} else {
let model = state.config.selected_model.as_deref().unwrap_or("(none)");
let msg = format!("model '{model}' does not support reasoning effort '{effort}'");
warn!(session_id = ctx.session_id, error = %msg, "reasoning effort rejected");
broadcast(
&mut state.subscribers,
&ctx.daemon_tx,
DaemonMessage::ReasoningEffortSetFailed {
session_id: ctx.session_id,
effort,
error: msg,
},
);
}
false
}
fn handle_get_reasoning_effort(
reply: mpsc::Sender<String>,
state: &SessionState,
ctx: &RequestContext,
) -> bool {
let _ = ctx;
let current = state
.config
.reasoning_effort
.clone()
.unwrap_or_else(|| "off".to_string());
let _ = reply.send(current);
false
}
fn handle_undo(state: &mut SessionState, ctx: &RequestContext) -> bool {
let Some(turn_ids) = state.undo_turns() else {
debug!(
session_id = ctx.session_id,
"undo requested but no user turn to undo",
);
return false;
};
info!(
session_id = ctx.session_id,
turn_count = turn_ids.len(),
"undo: marked turns as undone",
);
if state.config.last_response_id.is_some() {
state.config.last_response_id = None;
state.config.last_response_id_producer = None;
let record = SessionRecord::from(&*state);
if let Err(e) = write_session_retry(&ctx.db, ctx.session_id, &record) {
tracing::warn!(error = %e, "failed to persist session record after Undo");
}
}
for &id in &turn_ids {
if let Some(turn) = state.turns.get(&id)
&& let Err(e) = write_turn_retry(&ctx.db, ctx.session_id, id, turn)
{
tracing::warn!(turn_id = id, error = %e, "failed to persist undone turn");
}
}
broadcast(
&mut state.subscribers,
&ctx.daemon_tx,
DaemonMessage::TurnsUndone {
session_id: ctx.session_id,
turn_ids,
},
);
false
}
fn handle_redo(state: &mut SessionState, ctx: &RequestContext) -> bool {
let Some(turns) = state.redo_turns() else {
debug!(
session_id = ctx.session_id,
"redo requested but nothing to redo (no prior undo, or new input after undo)",
);
return false;
};
info!(
session_id = ctx.session_id,
turn_count = turns.len(),
"redo: restored previously-undone turns",
);
for (&id, turn) in &turns {
if let Err(e) = write_turn_retry(&ctx.db, ctx.session_id, id, turn) {
tracing::warn!(turn_id = id, error = %e, "failed to persist redone turn");
}
}
broadcast(
&mut state.subscribers,
&ctx.daemon_tx,
DaemonMessage::TurnsRedone {
session_id: ctx.session_id,
turns: turns
.iter()
.map(|(&turn_id, turn)| (turn_id, turn_for_client(turn)))
.collect(),
},
);
false
}
fn handle_shutdown(
state: &mut SessionState,
shutdown_requested: &mut bool,
ctx: &RequestContext,
) -> bool {
*shutdown_requested = true;
for (&request_id, active) in &state.active_requests {
let _ = active.cancel_tx.send(());
broadcast(
&mut state.subscribers,
&ctx.daemon_tx,
DaemonMessage::Cancelled {
session_id: ctx.session_id,
request_id,
},
);
}
state.active_requests.is_empty()
}
#[expect(clippy::too_many_arguments)]
fn run_request_worker(
request_id: u32,
client: InferenceProvider,
session: &mut SessionState,
model: String,
cancel_rx: crossbeam_channel::Receiver<()>,
ctx: RequestContext,
child_reply: Option<mpsc::Sender<io::Result<ChildResult>>>,
user_text: Option<String>,
) -> io::Result<()> {
let request_start = std::time::Instant::now();
let initial_snapshot = session.snapshot();
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
run_agent_loop(
&client, session, &model, request_id, &cancel_rx, &ctx, user_text,
)
}));
let (outcome, snapshot) = match result {
Ok(Ok(true)) => (RequestOutcome::Cancelled, session.snapshot()),
Ok(Ok(false)) => (RequestOutcome::Done, session.snapshot()),
Ok(Err(e)) => (RequestOutcome::Failed(e), session.snapshot()),
Err(_) => (
RequestOutcome::Failed(io::Error::other("request worker panicked")),
initial_snapshot,
),
};
let req_status = match &outcome {
RequestOutcome::Done => "done",
RequestOutcome::Failed(_) => "failed",
RequestOutcome::Cancelled => "cancelled",
};
crate::metrics::record_request_total(req_status);
crate::metrics::record_request_duration(req_status, request_start.elapsed().as_secs_f64());
match &outcome {
RequestOutcome::Done => {
info!(session_id = ctx.session_id, request_id, "request completed");
let usage = &session.config.accumulated_usage;
debug!(
session_id = ctx.session_id,
request_id,
input_tokens = usage.input_tokens,
output_tokens = usage.output_tokens,
total_tokens = usage.total_tokens,
"broadcasting Done with accumulated token usage"
);
let _ = ctx
.cmd_tx
.send(SessionCommand::Broadcast(DaemonMessage::Done {
session_id: ctx.session_id,
request_id,
token_usage: Some(*usage),
last_prompt_tokens: session.config.last_prompt_tokens,
}));
}
RequestOutcome::Failed(error) => {
info!(session_id = ctx.session_id, request_id, error = %error, "request failed");
let _ = ctx
.cmd_tx
.send(SessionCommand::Broadcast(DaemonMessage::Failed {
session_id: ctx.session_id,
request_id,
error: error.to_string(),
}));
}
RequestOutcome::Cancelled => {
info!(session_id = ctx.session_id, request_id, "request cancelled");
}
}
if let Some(reply) = child_reply {
let child_result = match &outcome {
RequestOutcome::Done => {
let output = session
.turns
.values()
.filter_map(|t| t.assistant_text.clone())
.collect::<Vec<_>>()
.join("\n");
Ok(ChildResult {
output,
is_error: false,
})
}
RequestOutcome::Failed(error) => Ok(ChildResult {
output: error.to_string(),
is_error: true,
}),
RequestOutcome::Cancelled => Ok(ChildResult {
output: "request cancelled".to_string(),
is_error: true,
}),
};
let _ = reply.send(child_result);
}
let _ = ctx.cmd_tx.send(SessionCommand::RequestFinished {
request_id,
snapshot,
});
Ok(())
}
enum RequestOutcome {
Done,
Failed(io::Error),
Cancelled,
}
fn persist_and_exit(
state: &SessionState,
db: &redb::Database,
session_id: u64,
daemon_tx: &mpsc::Sender<DaemonCommand>,
) {
let record: SessionRecord = SessionRecord::from(state);
if let Err(e) = write_session_retry(db, session_id, &record) {
error!(
"persist_and_exit: failed to persist session {}: {e}",
session_id
);
}
let _ = daemon_tx.send(DaemonCommand::SessionExited { session_id });
}
#[cfg(test)]
mod tests {
use super::*;
use crate::server::connection::SUBSCRIBER_CHANNEL_CAPACITY;
use crate::tools::ToolRegistry;
use choreo_proto::SessionStatus;
use std::collections::HashMap;
use tempfile::tempdir;
fn test_state() -> SessionState {
let mut turns = BTreeMap::new();
turns.insert(
0,
Turn {
created_at: TimestampMs::now(),
undone: false,
error: None,
user_text: Some("hello".into()),
assistant_text: Some("hi".into()),
assistant_reasoning: None,
tool_calls: Vec::new(),
token_usage: None,
tool_results: Vec::new(),
displayed_images: Vec::new(),
reasoning_artifact: None,
reasoning_producer: None,
},
);
SessionState {
config: SessionConfig {
title: Some("test session".into()),
selected_model: Some("gpt-4".into()),
reasoning_effort: None,
parent_session_id: None,
working_dir: Some(std::path::PathBuf::from("/tmp")),
created_at: 1000,
last_modified: 1000,
status: SessionStatus::Inactive,
active_tool_groups: ["core".into(), "shell".into()].into(),
context_config: ContextConfig::default(),
account_name: None,
accumulated_usage: TokenUsage::default(),
context_window: None,
last_prompt_tokens: None,
last_response_id: None,
last_response_id_producer: None,
},
next_turn_id: 1,
last_undo_turn_ids: None,
turns,
loaded_skill_bodies: Vec::new(),
context_cache: None,
discovered_skills: None,
subscribers: HashMap::new(),
active_requests: BTreeMap::new(),
provider: None,
}
}
#[test]
fn set_assistant_response_stores_artifact_and_producer() {
let mut state = SessionState::empty();
let (tid, _) = state.start_turn(Some("hello".into()));
let artifact = ReasoningArtifact::ChatReasoning {
field: choreo_proto::ChatReasoningField::ReasoningContent,
bytes: b"thinking".to_vec(),
};
let producer = ReasoningProducer {
provider_slug: "deepseek".into(),
model: "deepseek-v4-pro".into(),
};
state.set_assistant_response(
tid,
AssistantResponse {
text: Some("hi".into()),
reasoning_artifact: Some(artifact.clone()),
reasoning_producer: Some(producer.clone()),
..Default::default()
},
);
let turn = state.turns.get(&tid).expect("turn exists");
assert_eq!(turn.reasoning_artifact, Some(artifact));
assert_eq!(turn.reasoning_producer, Some(producer));
}
#[test]
fn set_assistant_response_no_artifact_keeps_turn_clean() {
let mut state = SessionState::empty();
let (tid, _) = state.start_turn(Some("hello".into()));
state.set_assistant_response(
tid,
AssistantResponse {
text: Some("hi".into()),
..Default::default()
},
);
let turn = state.turns.get(&tid).expect("turn exists");
assert_eq!(turn.reasoning_artifact, None);
assert_eq!(turn.reasoning_producer, None);
}
#[test]
fn turn_for_client_strips_artifact_and_producer() {
let artifact = ReasoningArtifact::ChatReasoning {
field: choreo_proto::ChatReasoningField::ReasoningContent,
bytes: b"thinking".to_vec(),
};
let producer = ReasoningProducer {
provider_slug: "deepseek".into(),
model: "deepseek-v4-pro".into(),
};
let mut state = SessionState::empty();
let (tid, _) = state.start_turn(Some("hello".into()));
state.set_assistant_response(
tid,
AssistantResponse {
text: Some("hi".into()),
reasoning: Some("thinking out loud".into()),
reasoning_artifact: Some(artifact),
reasoning_producer: Some(producer),
..Default::default()
},
);
let authoritative = state.turns.get(&tid).expect("turn exists");
let client = turn_for_client(authoritative);
assert_eq!(client.reasoning_artifact, None);
assert_eq!(client.reasoning_producer, None);
assert_eq!(client.assistant_text.as_deref(), Some("hi"));
assert_eq!(
client.assistant_reasoning.as_deref(),
Some("thinking out loud")
);
assert_eq!(client.user_text.as_deref(), Some("hello"));
assert!(authoritative.reasoning_artifact.is_some());
assert!(authoritative.reasoning_producer.is_some());
}
#[test]
fn session_state_message_strips_artifacts_from_turns() {
let artifact = ReasoningArtifact::ChatReasoning {
field: choreo_proto::ChatReasoningField::ReasoningContent,
bytes: b"thinking".to_vec(),
};
let producer = ReasoningProducer {
provider_slug: "deepseek".into(),
model: "deepseek-v4-pro".into(),
};
let mut state = SessionState::empty();
let (tid, _) = state.start_turn(Some("hello".into()));
state.set_assistant_response(
tid,
AssistantResponse {
text: Some("hi".into()),
reasoning: Some("thinking out loud".into()),
reasoning_artifact: Some(artifact),
reasoning_producer: Some(producer),
..Default::default()
},
);
let DaemonMessage::SessionState { turns, .. } = state.session_state_message(7) else {
panic!("expected SessionState message");
};
let client_turn = turns.get(&tid).expect("turn present in message");
assert_eq!(client_turn.reasoning_artifact, None);
assert_eq!(client_turn.reasoning_producer, None);
assert_eq!(client_turn.assistant_text.as_deref(), Some("hi"));
assert_eq!(
client_turn.assistant_reasoning.as_deref(),
Some("thinking out loud")
);
let authoritative = state.turns.get(&tid).expect("turn exists");
assert!(authoritative.reasoning_artifact.is_some());
assert!(authoritative.reasoning_producer.is_some());
}
#[test]
fn session_record_carries_last_response_id_from_config() {
let mut state = SessionState::empty();
state.config.last_response_id = Some("resp_9".into());
state.config.last_response_id_producer = Some(ReasoningProducer {
provider_slug: "openai".into(),
model: "gpt-5.4".into(),
});
let record = SessionRecord::from(&state);
assert_eq!(record.last_response_id.as_deref(), Some("resp_9"));
assert_eq!(
record
.last_response_id_producer
.as_ref()
.map(|p| p.model.as_str()),
Some("gpt-5.4"),
"producer must survive the state → record conversion",
);
let restored = SessionConfig {
last_response_id: record.last_response_id.clone(),
last_response_id_producer: record.last_response_id_producer.clone(),
..SessionConfig::default()
};
assert_eq!(restored.last_response_id.as_deref(), Some("resp_9"));
assert_eq!(restored.last_response_id_producer.unwrap().model, "gpt-5.4",);
}
fn broadcast_setup() -> (SessionState, RequestContext) {
let dir = tempdir().unwrap();
let db = Arc::new(redb::Database::create(dir.path().join("test.redb")).unwrap());
let tool_registry = ToolRegistry::new().build();
let (daemon_tx, _) = mpsc::channel();
let (cmd_tx, _) = mpsc::channel();
let ctx = RequestContext {
cmd_tx,
session_id: 1,
db,
tool_registry,
daemon_tx,
max_turns: 0,
};
(test_state(), ctx)
}
#[test]
fn session_state_round_trip_metadata() {
let state = test_state();
let meta: SessionMetadata = (&state).into();
assert_eq!(meta.title, state.config.title);
assert_eq!(meta.selected_model, state.config.selected_model);
assert_eq!(meta.turn_count, 1);
assert_eq!(meta.status, state.config.status);
}
#[test]
fn session_state_to_record() {
let state = test_state();
let record: SessionRecord = (&state).into();
assert_eq!(record.title, state.config.title);
assert_eq!(record.selected_model, state.config.selected_model);
assert_eq!(record.turn_count, 1);
}
#[test]
fn apply_worker_snapshot_preserves_main_loop_config_mutations() {
let mut state = test_state();
state.config.working_dir = Some(PathBuf::from("/main-loop-wd"));
state.config.active_tool_groups.insert("x".into());
state.config.title = Some("main-loop title".into());
let mut snapshot = state.config.clone();
snapshot.working_dir = Some(PathBuf::from("/stale-wd"));
snapshot.active_tool_groups = ["core".into()].into_iter().collect();
snapshot.title = Some("stale title".into());
snapshot.accumulated_usage = TokenUsage {
input_tokens: 10,
output_tokens: 5,
total_tokens: 15,
};
snapshot.context_window = Some(8192);
snapshot.last_prompt_tokens = Some(10);
state.config.apply_worker_snapshot(&snapshot);
assert_eq!(state.config.accumulated_usage.total_tokens, 15);
assert_eq!(state.config.context_window, Some(8192));
assert_eq!(state.config.last_prompt_tokens, Some(10));
assert_eq!(
state.config.working_dir,
Some(PathBuf::from("/main-loop-wd"))
);
assert!(
state.config.active_tool_groups.contains("x"),
"active_tool_groups must not be clobbered by the worker snapshot"
);
assert_eq!(state.config.title.as_deref(), Some("main-loop title"));
}
#[test]
fn broadcast_delivers_message_to_all_subscribers() {
let (tx1, rx1) = mpsc::sync_channel(SUBSCRIBER_CHANNEL_CAPACITY);
let (tx2, rx2) = mpsc::sync_channel(SUBSCRIBER_CHANNEL_CAPACITY);
let (mut state, ctx) = broadcast_setup();
state.subscribers.insert(10, tx1);
state.subscribers.insert(20, tx2);
let mut shutdown = false;
process_command(
SessionCommand::Broadcast(DaemonMessage::Done {
session_id: ctx.session_id,
request_id: 5,
token_usage: None,
last_prompt_tokens: None,
}),
&mut state,
&mut shutdown,
&ctx,
);
assert_eq!(
rx1.recv().unwrap(),
DaemonMessage::Done {
session_id: ctx.session_id,
request_id: 5,
token_usage: None,
last_prompt_tokens: None,
}
);
assert_eq!(
rx2.recv().unwrap(),
DaemonMessage::Done {
session_id: ctx.session_id,
request_id: 5,
token_usage: None,
last_prompt_tokens: None,
}
);
assert!(!shutdown);
}
#[test]
fn broadcast_with_no_subscribers_does_not_panic() {
let (mut state, ctx) = broadcast_setup();
let mut shutdown = false;
process_command(
SessionCommand::Broadcast(DaemonMessage::Done {
session_id: ctx.session_id,
request_id: 0,
token_usage: None,
last_prompt_tokens: None,
}),
&mut state,
&mut shutdown,
&ctx,
);
assert!(!shutdown);
}
#[test]
fn broadcast_handles_disconnected_subscriber_gracefully() {
let (tx, _rx) = mpsc::sync_channel(SUBSCRIBER_CHANNEL_CAPACITY);
drop(_rx);
let (mut state, ctx) = broadcast_setup();
state.subscribers.insert(99, tx);
let mut shutdown = false;
process_command(
SessionCommand::Broadcast(DaemonMessage::Pong),
&mut state,
&mut shutdown,
&ctx,
);
assert!(!shutdown);
}
#[test]
fn broadcast_keeps_slow_subscriber_on_full_buffer() {
let (mut state, ctx) = broadcast_setup();
let (tx, rx) = mpsc::sync_channel::<DaemonMessage>(1);
state.subscribers.insert(10, tx);
let filler = DaemonMessage::Pong;
state
.subscribers
.get(&10)
.expect("subscriber registered")
.send(filler.clone())
.expect("filler fits in capacity-1 channel");
let broadcast = DaemonMessage::Done {
session_id: ctx.session_id,
request_id: 5,
token_usage: None,
last_prompt_tokens: None,
};
let mut shutdown = false;
process_command(
SessionCommand::Broadcast(broadcast.clone()),
&mut state,
&mut shutdown,
&ctx,
);
assert!(
state.subscribers.contains_key(&10),
"a full buffer must not evict the session subscriber"
);
assert_eq!(rx.recv().unwrap(), filler);
assert!(rx.try_recv().is_err(), "message dropped while buffer full");
process_command(
SessionCommand::Broadcast(broadcast.clone()),
&mut state,
&mut shutdown,
&ctx,
);
assert_eq!(rx.recv().unwrap(), broadcast);
assert!(state.subscribers.contains_key(&10));
assert!(!shutdown);
}
#[test]
fn set_working_dir_updates_config_and_broadcasts() {
let (mut state, ctx) = broadcast_setup();
let (tx, rx) = mpsc::sync_channel(SUBSCRIBER_CHANNEL_CAPACITY);
state.subscribers.insert(10, tx);
let (reply_tx, reply_rx) = mpsc::channel();
state.discovered_skills = Some(Vec::new());
let new_path = PathBuf::from("/tmp/new-wd");
let mut shutdown = false;
process_command(
SessionCommand::SetWorkingDir {
path: new_path.clone(),
reply: reply_tx,
},
&mut state,
&mut shutdown,
&ctx,
);
assert_eq!(state.config.working_dir, Some(new_path));
assert!(
state.discovered_skills.is_none(),
"skill cache must be invalidated on working-dir change"
);
assert!(!shutdown);
match rx.recv().unwrap() {
DaemonMessage::SessionWorkingDirSet { session_id, path } => {
assert_eq!(session_id, ctx.session_id);
assert_eq!(path.as_deref(), Some("/tmp/new-wd"));
}
other => panic!("expected SessionWorkingDirSet, got {:?}", other),
}
match reply_rx.recv() {
Ok(Ok(msg)) => assert_eq!(msg, "/tmp/new-wd"),
Ok(Err(e)) => panic!("expected success reply, got error: {e}"),
Err(e) => panic!("expected reply, got {e:?}"),
}
}
#[test]
fn load_tools_updates_active_groups_and_replies() {
let (mut state, ctx) = broadcast_setup();
let (reply_tx, reply_rx) = mpsc::channel();
let mut shutdown = false;
process_command(
SessionCommand::LoadTools {
groups: vec!["x".into()],
reply: reply_tx,
},
&mut state,
&mut shutdown,
&ctx,
);
assert!(state.config.active_tool_groups.contains("x"));
assert!(!shutdown);
match reply_rx.recv() {
Ok(Ok(msg)) => assert_eq!(msg, "Activated tool groups: x"),
Ok(Err(e)) => panic!("expected success reply, got error: {e}"),
Err(e) => panic!("expected reply, got {e:?}"),
}
}
#[test]
fn load_tools_skips_already_active_in_reply() {
let (mut state, ctx) = broadcast_setup();
state.config.active_tool_groups.insert("shell".into());
let (reply_tx, reply_rx) = mpsc::channel();
let mut shutdown = false;
process_command(
SessionCommand::LoadTools {
groups: vec!["shell".into()],
reply: reply_tx,
},
&mut state,
&mut shutdown,
&ctx,
);
match reply_rx.recv() {
Ok(Ok(msg)) => {
assert_eq!(msg, "All specified groups were already active.")
}
Ok(Err(e)) => panic!("expected success reply, got error: {e}"),
Err(e) => panic!("expected reply, got {e:?}"),
}
}
#[test]
fn unload_tools_updates_active_groups_and_replies() {
let (mut state, ctx) = broadcast_setup();
state.config.active_tool_groups.insert("x".into());
let (reply_tx, reply_rx) = mpsc::channel();
let mut shutdown = false;
process_command(
SessionCommand::UnloadTools {
groups: vec!["x".into()],
reply: reply_tx,
},
&mut state,
&mut shutdown,
&ctx,
);
assert!(!state.config.active_tool_groups.contains("x"));
assert!(!shutdown);
match reply_rx.recv() {
Ok(Ok(msg)) => assert_eq!(msg, "Deactivated tool groups: x"),
Ok(Err(e)) => panic!("expected success reply, got error: {e}"),
Err(e) => panic!("expected reply, got {e:?}"),
}
}
#[test]
fn unload_tools_protects_core() {
let (mut state, ctx) = broadcast_setup();
let (reply_tx, reply_rx) = mpsc::channel();
let mut shutdown = false;
process_command(
SessionCommand::UnloadTools {
groups: vec!["core".into()],
reply: reply_tx,
},
&mut state,
&mut shutdown,
&ctx,
);
assert!(state.config.active_tool_groups.contains("core"));
match reply_rx.recv() {
Ok(Ok(msg)) => assert_eq!(msg, "The 'core' group cannot be unloaded."),
Ok(Err(e)) => panic!("expected success reply, got error: {e}"),
Err(e) => panic!("expected reply, got {e:?}"),
}
}
#[test]
fn load_tools_rejects_unknown_group() {
let (mut state, ctx) = broadcast_setup();
let (reply_tx, reply_rx) = mpsc::channel();
let mut shutdown = false;
process_command(
SessionCommand::LoadTools {
groups: vec!["not-a-real-group".into()],
reply: reply_tx,
},
&mut state,
&mut shutdown,
&ctx,
);
assert!(!state.config.active_tool_groups.contains("not-a-real-group"));
match reply_rx.recv() {
Ok(Err(msg)) => {
assert!(msg.contains("Unknown tool group(s): not-a-real-group"))
}
Ok(Ok(msg)) => panic!("expected error reply, got success: {msg}"),
Err(e) => panic!("expected reply, got {e:?}"),
}
}
#[test]
fn unload_tools_rejects_unknown_group() {
let (mut state, ctx) = broadcast_setup();
state.config.active_tool_groups.insert("git".into());
let (reply_tx, reply_rx) = mpsc::channel();
let mut shutdown = false;
process_command(
SessionCommand::UnloadTools {
groups: vec!["not-a-real-group".into()],
reply: reply_tx,
},
&mut state,
&mut shutdown,
&ctx,
);
assert!(state.config.active_tool_groups.contains("git"));
match reply_rx.recv() {
Ok(Err(msg)) => {
assert!(msg.contains("Unknown tool group(s): not-a-real-group"))
}
Ok(Ok(msg)) => panic!("expected error reply, got success: {msg}"),
Err(e) => panic!("expected reply, got {e:?}"),
}
}
#[test]
fn cancel_sends_through_channel() {
let (cancel_tx, cancel_rx) = crossbeam_channel::unbounded::<()>();
let (mut state, ctx) = broadcast_setup();
state.active_requests.insert(
1,
ActiveRequest {
cancel_tx,
turn_id: 1,
},
);
let mut shutdown = false;
process_command(
SessionCommand::Cancel { request_id: 1 },
&mut state,
&mut shutdown,
&ctx,
);
assert!(cancel_rx.try_recv().is_ok());
assert!(!shutdown);
}
#[test]
fn shutdown_cancels_all_active_requests() {
let (cancel_tx1, cancel_rx1) = crossbeam_channel::unbounded::<()>();
let (cancel_tx2, cancel_rx2) = crossbeam_channel::unbounded::<()>();
let (mut state, ctx) = broadcast_setup();
state.active_requests.insert(
1,
ActiveRequest {
cancel_tx: cancel_tx1,
turn_id: 1,
},
);
state.active_requests.insert(
2,
ActiveRequest {
cancel_tx: cancel_tx2,
turn_id: 2,
},
);
let mut shutdown = false;
process_command(SessionCommand::Shutdown, &mut state, &mut shutdown, &ctx);
assert!(shutdown);
assert!(cancel_rx1.try_recv().is_ok());
assert!(cancel_rx2.try_recv().is_ok());
}
#[test]
fn shutdown_with_empty_active_requests_returns_true() {
let (mut state, ctx) = broadcast_setup();
let mut shutdown = false;
let should_exit =
process_command(SessionCommand::Shutdown, &mut state, &mut shutdown, &ctx);
assert!(shutdown);
assert!(should_exit);
}
#[test]
fn accumulated_usage_starts_at_zero() {
let state = SessionState::empty();
assert_eq!(state.config.accumulated_usage.input_tokens, 0);
assert_eq!(state.config.accumulated_usage.output_tokens, 0);
assert_eq!(state.config.accumulated_usage.total_tokens, 0);
}
#[test]
fn accumulated_usage_reconstructed_from_turns() {
let mut state = SessionState::empty();
state.turns.insert(
0,
Turn {
created_at: TimestampMs::now(),
undone: false,
error: None,
user_text: Some("hello".into()),
assistant_text: Some("hi".into()),
assistant_reasoning: None,
tool_calls: Vec::new(),
token_usage: None,
tool_results: Vec::new(),
displayed_images: Vec::new(),
reasoning_artifact: None,
reasoning_producer: None,
},
);
state.turns.insert(
1,
Turn {
created_at: TimestampMs::now(),
undone: false,
error: None,
user_text: Some("turn 2".into()),
assistant_text: Some("response 2".into()),
assistant_reasoning: None,
tool_calls: Vec::new(),
token_usage: Some(TokenUsage {
input_tokens: 10,
output_tokens: 20,
total_tokens: 30,
}),
tool_results: Vec::new(),
displayed_images: Vec::new(),
reasoning_artifact: None,
reasoning_producer: None,
},
);
state.turns.insert(
2,
Turn {
created_at: TimestampMs::now(),
undone: false,
error: None,
user_text: Some("turn 3".into()),
assistant_text: Some("response 3".into()),
assistant_reasoning: None,
tool_calls: Vec::new(),
token_usage: Some(TokenUsage {
input_tokens: 100,
output_tokens: 50,
total_tokens: 150,
}),
tool_results: Vec::new(),
displayed_images: Vec::new(),
reasoning_artifact: None,
reasoning_producer: None,
},
);
state.turns.insert(
3,
Turn {
created_at: TimestampMs::now(),
undone: false,
error: None,
user_text: Some("no usage".into()),
assistant_text: None,
assistant_reasoning: None,
tool_calls: Vec::new(),
token_usage: None,
tool_results: Vec::new(),
displayed_images: Vec::new(),
reasoning_artifact: None,
reasoning_producer: None,
},
);
let mut accumulated_usage = TokenUsage::default();
let mut last_prompt_tokens = None;
for turn in state.turns.values() {
if let Some(u) = turn.token_usage {
accumulated_usage.input_tokens += u.input_tokens;
accumulated_usage.output_tokens += u.output_tokens;
accumulated_usage.total_tokens += u.total_tokens;
last_prompt_tokens = Some(u.input_tokens);
}
}
state.config.accumulated_usage = accumulated_usage;
state.config.last_prompt_tokens = last_prompt_tokens;
assert_eq!(state.config.accumulated_usage.input_tokens, 110);
assert_eq!(state.config.accumulated_usage.output_tokens, 70);
assert_eq!(state.config.accumulated_usage.total_tokens, 180);
assert_eq!(state.config.last_prompt_tokens, Some(100));
}
#[test]
fn last_prompt_tokens_from_latest_usage_turn() {
let mut state = SessionState::empty();
state.turns.insert(
0,
Turn {
created_at: TimestampMs::now(),
undone: false,
error: None,
user_text: Some("no usage".into()),
assistant_text: None,
assistant_reasoning: None,
tool_calls: Vec::new(),
token_usage: None,
tool_results: Vec::new(),
displayed_images: Vec::new(),
reasoning_artifact: None,
reasoning_producer: None,
},
);
state.turns.insert(
1,
Turn {
created_at: TimestampMs::now(),
undone: false,
error: None,
user_text: Some("first".into()),
assistant_text: Some("response".into()),
assistant_reasoning: None,
tool_calls: Vec::new(),
token_usage: Some(TokenUsage {
input_tokens: 5,
output_tokens: 10,
total_tokens: 15,
}),
tool_results: Vec::new(),
displayed_images: Vec::new(),
reasoning_artifact: None,
reasoning_producer: None,
},
);
state.turns.insert(
2,
Turn {
created_at: TimestampMs::now(),
undone: false,
error: None,
user_text: Some("second".into()),
assistant_text: Some("response 2".into()),
assistant_reasoning: None,
tool_calls: Vec::new(),
token_usage: Some(TokenUsage {
input_tokens: 42,
output_tokens: 7,
total_tokens: 49,
}),
tool_results: Vec::new(),
displayed_images: Vec::new(),
reasoning_artifact: None,
reasoning_producer: None,
},
);
let mut accumulated_usage = TokenUsage::default();
let mut last_prompt_tokens = None;
for turn in state.turns.values() {
if let Some(u) = turn.token_usage {
accumulated_usage.input_tokens += u.input_tokens;
accumulated_usage.output_tokens += u.output_tokens;
accumulated_usage.total_tokens += u.total_tokens;
last_prompt_tokens = Some(u.input_tokens);
}
}
state.config.accumulated_usage = accumulated_usage;
state.config.last_prompt_tokens = last_prompt_tokens;
assert_eq!(state.config.accumulated_usage.input_tokens, 47);
assert_eq!(state.config.accumulated_usage.output_tokens, 17);
assert_eq!(state.config.accumulated_usage.total_tokens, 64);
assert_eq!(state.config.last_prompt_tokens, Some(42));
}
#[test]
fn last_prompt_tokens_none_when_no_turns_have_usage() {
let mut state = SessionState::empty();
state.turns.insert(
0,
Turn {
created_at: TimestampMs::now(),
undone: false,
error: None,
user_text: Some("no usage".into()),
assistant_text: None,
assistant_reasoning: None,
tool_calls: Vec::new(),
token_usage: None,
tool_results: Vec::new(),
displayed_images: Vec::new(),
reasoning_artifact: None,
reasoning_producer: None,
},
);
let mut accumulated_usage = TokenUsage::default();
let mut last_prompt_tokens = None;
for turn in state.turns.values() {
if let Some(u) = turn.token_usage {
accumulated_usage.input_tokens += u.input_tokens;
accumulated_usage.output_tokens += u.output_tokens;
accumulated_usage.total_tokens += u.total_tokens;
last_prompt_tokens = Some(u.input_tokens);
}
}
state.config.accumulated_usage = accumulated_usage;
state.config.last_prompt_tokens = last_prompt_tokens;
assert_eq!(state.config.accumulated_usage.input_tokens, 0);
assert_eq!(state.config.last_prompt_tokens, None);
}
#[test]
fn accumulated_usage_in_snapshot() {
let mut state = SessionState::empty();
state.config.accumulated_usage = TokenUsage {
input_tokens: 50,
output_tokens: 25,
total_tokens: 75,
};
let snap = state.snapshot();
assert_eq!(snap.config.accumulated_usage.input_tokens, 50);
assert_eq!(snap.config.accumulated_usage.output_tokens, 25);
assert_eq!(snap.config.accumulated_usage.total_tokens, 75);
}
#[test]
fn accumulated_usage_in_session_summary() {
let (mut state, ctx) = broadcast_setup();
state.config.accumulated_usage = TokenUsage {
input_tokens: 80,
output_tokens: 40,
total_tokens: 120,
};
let (reply, rx) = mpsc::channel();
let mut shutdown = false;
process_command(
SessionCommand::GetSummary { reply },
&mut state,
&mut shutdown,
&ctx,
);
let summary: SessionSummary = rx.recv().unwrap();
let summary_usage = summary
.token_usage
.expect("token_usage should be present in SessionSummary");
assert_eq!(summary_usage.input_tokens, 80);
assert_eq!(summary_usage.output_tokens, 40);
assert_eq!(summary_usage.total_tokens, 120);
}
#[test]
fn accumulated_usage_in_attach_snapshot() {
let (mut state, ctx) = broadcast_setup();
state.config.accumulated_usage = TokenUsage {
input_tokens: 30,
output_tokens: 15,
total_tokens: 45,
};
let (sub_tx, sub_rx) = mpsc::sync_channel(SUBSCRIBER_CHANNEL_CAPACITY);
let mut shutdown = false;
process_command(
SessionCommand::Attach {
client_id: 42,
tx: sub_tx,
},
&mut state,
&mut shutdown,
&ctx,
);
let msg = sub_rx.recv().unwrap();
match msg {
DaemonMessage::SessionState { token_usage, .. } => {
let usage = token_usage.expect("token_usage in SessionState");
assert_eq!(usage.input_tokens, 30);
assert_eq!(usage.output_tokens, 15);
assert_eq!(usage.total_tokens, 45);
}
other => panic!("expected SessionState, got {other:?}"),
}
}
#[test]
fn attach_with_active_requests_sends_started_to_new_subscriber() {
let (mut state, ctx) = broadcast_setup();
let (cancel_tx1, _cancel_rx1) = crossbeam_channel::unbounded::<()>();
let (cancel_tx2, _cancel_rx2) = crossbeam_channel::unbounded::<()>();
state.active_requests.insert(
10,
ActiveRequest {
cancel_tx: cancel_tx1,
turn_id: 3,
},
);
state.active_requests.insert(
20,
ActiveRequest {
cancel_tx: cancel_tx2,
turn_id: 7,
},
);
let (sub_tx, sub_rx) = mpsc::sync_channel(SUBSCRIBER_CHANNEL_CAPACITY);
let mut shutdown = false;
process_command(
SessionCommand::Attach {
client_id: 42,
tx: sub_tx,
},
&mut state,
&mut shutdown,
&ctx,
);
match sub_rx.recv().unwrap() {
DaemonMessage::Started {
session_id: 1,
request_id: 10,
turn_id: 3,
estimated_prompt_tokens: 0,
} => {}
other => panic!("expected Started(10, turn=3), got {other:?}"),
}
match sub_rx.recv().unwrap() {
DaemonMessage::Started {
session_id: 1,
request_id: 20,
turn_id: 7,
estimated_prompt_tokens: 0,
} => {}
other => panic!("expected Started(20, turn=7), got {other:?}"),
}
match sub_rx.recv().unwrap() {
DaemonMessage::SessionState { .. } => {}
other => panic!("expected SessionState, got {other:?}"),
}
assert!(!shutdown);
}
#[test]
fn attach_without_active_requests_does_not_send_started() {
let (mut state, ctx) = broadcast_setup();
let (sub_tx, sub_rx) = mpsc::sync_channel(SUBSCRIBER_CHANNEL_CAPACITY);
let mut shutdown = false;
process_command(
SessionCommand::Attach {
client_id: 42,
tx: sub_tx,
},
&mut state,
&mut shutdown,
&ctx,
);
match sub_rx.recv().unwrap() {
DaemonMessage::SessionState { .. } => {}
other => panic!("expected SessionState, got {other:?}"),
}
assert!(sub_rx.try_recv().is_err());
assert!(!shutdown);
}
#[test]
fn start_turn_assigns_increasing_ids() {
let mut state = SessionState::empty();
let (id0, _) = state.start_turn(Some("first".into()));
let (id1, _) = state.start_turn(Some("second".into()));
assert_eq!(id0, 0);
assert_eq!(id1, 1);
assert_eq!(state.turns.len(), 2);
}
#[test]
fn seed_then_update_tool_results_preserves_call_order() {
let mut state = SessionState::empty();
let (tid, _) = state.start_turn(Some("run tools".into()));
let calls = vec![
AssistantToolCallRecord {
call_id: "a".into(),
name: "read_file".into(),
arguments_json: "{}".into(),
},
AssistantToolCallRecord {
call_id: "b".into(),
name: "grep".into(),
arguments_json: "{}".into(),
},
AssistantToolCallRecord {
call_id: "c".into(),
name: "sh".into(),
arguments_json: "{}".into(),
},
];
state.seed_tool_results(tid, &calls);
let order_of = |state: &SessionState| {
state
.turns
.get(&tid)
.map(|t| {
t.tool_results
.iter()
.map(|r| r.call_id.clone())
.collect::<Vec<_>>()
})
.unwrap_or_default()
};
assert_eq!(order_of(&state), vec!["a", "b", "c"]);
state.update_tool_result(tid, "c", "sh".into(), "c-out".into(), false, String::new());
assert_eq!(order_of(&state), vec!["a", "b", "c"]);
assert_eq!(state.turns[&tid].tool_results[2].content, "c-out");
state.update_tool_result(
tid,
"a",
"read_file".into(),
"a-out".into(),
false,
String::new(),
);
state.update_tool_result(
tid,
"b",
"grep".into(),
"b-out".into(),
false,
String::new(),
);
assert_eq!(order_of(&state), vec!["a", "b", "c"]);
assert_eq!(state.turns[&tid].tool_results[0].content, "a-out");
assert_eq!(state.turns[&tid].tool_results[1].content, "b-out");
state.update_tool_result(tid, "b", "grep".into(), "boom".into(), true, String::new());
assert_eq!(order_of(&state), vec!["a", "b", "c"]);
assert!(state.turns[&tid].tool_results[1].is_error);
assert_eq!(state.turns[&tid].tool_results[1].content, "boom");
}
#[test]
fn update_tool_result_unknown_call_id_is_noop() {
let mut state = SessionState::empty();
let (tid, _) = state.start_turn(None);
state.seed_tool_results(tid, &[]);
state.update_tool_result(
tid,
"ghost",
"read_file".into(),
"x".into(),
false,
String::new(),
);
assert!(state.turns[&tid].tool_results.is_empty());
}
#[test]
fn mark_unexecuted_tool_results_marks_only_unexecuted() {
let mut state = SessionState::empty();
let (tid, _) = state.start_turn(Some("run tools".into()));
let calls = vec![
AssistantToolCallRecord {
call_id: "a".into(),
name: "read_file".into(),
arguments_json: "{}".into(),
},
AssistantToolCallRecord {
call_id: "b".into(),
name: "grep".into(),
arguments_json: "{}".into(),
},
AssistantToolCallRecord {
call_id: "c".into(),
name: "sh".into(),
arguments_json: "{}".into(),
},
];
state.seed_tool_results(tid, &calls);
state.update_tool_result(
tid,
"a",
"read_file".into(),
"a-out".into(),
false,
String::new(),
);
let executed = HashSet::from(["a".to_string()]);
state.mark_unexecuted_tool_results(tid, &executed);
let results = &state.turns[&tid].tool_results;
assert_eq!(results[0].content, "a-out");
assert!(!results[0].is_error);
assert_eq!(results[1].content, "[cancelled — result not recorded]");
assert!(results[1].is_error);
assert_eq!(results[2].content, "[cancelled — result not recorded]");
assert!(results[2].is_error);
}
#[test]
fn mark_unexecuted_tool_results_preserves_recorded_error_results() {
let mut state = SessionState::empty();
let (tid, _) = state.start_turn(Some("run tools".into()));
let calls = vec![AssistantToolCallRecord {
call_id: "a".into(),
name: "sh".into(),
arguments_json: "{}".into(),
}];
state.seed_tool_results(tid, &calls);
state.update_tool_result(
tid,
"a",
"sh".into(),
"timed out".into(),
true,
String::new(),
);
state.mark_unexecuted_tool_results(tid, &HashSet::from(["a".to_string()]));
assert_eq!(state.turns[&tid].tool_results[0].content, "timed out");
assert!(state.turns[&tid].tool_results[0].is_error);
}
#[test]
fn undo_turns_marks_range_and_returns_ids() {
let mut state = SessionState::empty();
let _ = state.start_turn(Some("user 1".into()));
let _ = state.start_turn(Some("user 2".into()));
assert!(state.turns.values().all(|t| !t.undone));
let ids = state
.undo_turns()
.expect("undo_turns should find a user turn");
assert_eq!(ids.len(), 1, "only the most recent user turn");
assert!(state.turns.get(&1).unwrap().undone);
}
#[test]
fn undo_turns_returns_none_when_no_user_turn() {
let mut state = SessionState::empty();
let _ = state.start_turn(None); assert!(state.undo_turns().is_none());
}
#[test]
fn redo_turns_restores_undone_turns() {
let mut state = SessionState::empty();
let _ = state.start_turn(Some("user".into()));
let ids = state.undo_turns().expect("undo succeeds");
assert!(!ids.is_empty());
let restored = state.redo_turns().expect("redo succeeds");
assert_eq!(restored.len(), ids.len());
assert!(state.turns.values().all(|t| !t.undone));
}
#[test]
fn redo_turns_returns_none_when_nothing_to_redo() {
let mut state = SessionState::empty();
let _ = state.start_turn(Some("user".into()));
assert!(state.redo_turns().is_none());
}
#[test]
fn redo_turns_cleared_by_new_turn_start() {
let mut state = SessionState::empty();
let _ = state.start_turn(Some("first".into()));
state.undo_turns();
let _ = state.start_turn(Some("second".into()));
assert!(state.redo_turns().is_none());
}
#[test]
fn undo_clears_last_response_id_for_chain_invalidation() {
let (mut state, ctx) = broadcast_setup();
state.config.last_response_id = Some("resp_9".into());
state.config.last_response_id_producer = Some(ReasoningProducer {
provider_slug: "openai".into(),
model: "gpt-5.4".into(),
});
let _ = state.start_turn(Some("user 2".into()));
let mut shutdown = false;
process_command(SessionCommand::Undo, &mut state, &mut shutdown, &ctx);
assert_eq!(state.config.last_response_id, None);
assert_eq!(state.config.last_response_id_producer, None);
let record = SessionRecord::from(&state);
assert_eq!(record.last_response_id, None);
assert!(!shutdown);
}
#[test]
fn undo_without_response_id_leaves_session_untouched() {
let (mut state, ctx) = broadcast_setup();
let _ = state.start_turn(Some("user 2".into()));
let mut shutdown = false;
process_command(SessionCommand::Undo, &mut state, &mut shutdown, &ctx);
assert_eq!(state.config.last_response_id, None);
assert_eq!(state.config.last_response_id_producer, None);
assert!(state.turns.get(&1).expect("turn exists").undone);
assert!(!state.turns.get(&0).expect("turn exists").undone);
assert!(!shutdown);
}
#[test]
fn request_finished_after_in_flight_undo_preserves_chain_break_and_undone_turns() {
let (mut state, ctx) = broadcast_setup();
state.config.last_response_id = Some("resp_9".into());
state.config.last_response_id_producer = Some(ReasoningProducer {
provider_slug: "openai".into(),
model: "gpt-5.4".into(),
});
let _ = state.start_turn(Some("user 2".into()));
let mut shutdown = false;
process_command(SessionCommand::Undo, &mut state, &mut shutdown, &ctx);
assert_eq!(state.config.last_response_id, None);
let mut snapshot = state.snapshot();
snapshot.config.last_response_id = Some("resp_9".into());
snapshot.config.last_response_id_producer = Some(ReasoningProducer {
provider_slug: "openai".into(),
model: "gpt-5.4".into(),
});
for turn in snapshot.turns.values_mut() {
turn.undone = false;
}
let mut in_flight = snapshot.turns.get(&0).cloned().expect("seeded turn");
in_flight.user_text = Some("user 3 (in-flight)".into());
snapshot.turns.insert(2, in_flight);
process_command(
SessionCommand::RequestFinished {
request_id: 1,
snapshot,
},
&mut state,
&mut shutdown,
&ctx,
);
assert_eq!(state.config.last_response_id, None);
assert_eq!(state.config.last_response_id_producer, None);
let record = SessionRecord::from(&state);
assert_eq!(record.last_response_id, None);
assert!(state.turns.get(&1).expect("turn exists").undone);
assert_eq!(
state
.turns
.get(&2)
.expect("turn exists")
.user_text
.as_deref(),
Some("user 3 (in-flight)"),
);
assert!(!shutdown);
}
#[test]
fn request_finished_without_undo_restores_chain_id() {
let (mut state, ctx) = broadcast_setup();
let mut snapshot = state.snapshot();
snapshot.config.last_response_id = Some("resp_10".into());
snapshot.config.last_response_id_producer = Some(ReasoningProducer {
provider_slug: "openai".into(),
model: "gpt-5.4".into(),
});
let mut shutdown = false;
process_command(
SessionCommand::RequestFinished {
request_id: 1,
snapshot,
},
&mut state,
&mut shutdown,
&ctx,
);
assert_eq!(state.config.last_response_id.as_deref(), Some("resp_10"));
assert_eq!(
state
.config
.last_response_id_producer
.as_ref()
.unwrap()
.model,
"gpt-5.4",
);
assert!(!shutdown);
}
#[test]
fn loaded_skill_bodies_default_is_empty() {
let state = SessionState::empty();
assert!(state.loaded_skill_bodies.is_empty());
}
#[test]
fn context_cache_default_is_none() {
let state = SessionState::empty();
assert!(state.context_cache.is_none());
}
#[test]
fn loaded_skill_bodies_survives_snapshot_round_trip() {
let mut state = SessionState::empty();
state.loaded_skill_bodies.push(LoadedSkill {
name: "test".to_string(),
body: "body content".to_string(),
});
let snap = state.snapshot();
assert_eq!(snap.loaded_skill_bodies.len(), 1);
assert_eq!(snap.loaded_skill_bodies[0].name, "test");
let restored = SessionState::from_snapshot(snap, HashMap::new());
assert_eq!(restored.loaded_skill_bodies.len(), 1);
assert_eq!(restored.loaded_skill_bodies[0].name, "test");
assert_eq!(restored.loaded_skill_bodies[0].body, "body content");
}
#[test]
fn context_cache_survives_snapshot_round_trip() {
let mut state = SessionState::empty();
state.context_cache = Some((42, Arc::new("cached content".to_string())));
let snap = state.snapshot();
assert_eq!(
snap.context_cache,
Some((42, Arc::new("cached content".to_string())))
);
let restored = SessionState::from_snapshot(snap, HashMap::new());
assert_eq!(
restored.context_cache,
Some((42, Arc::new("cached content".to_string())))
);
}
#[test]
fn shutdown_join_poll_joins_when_finished() {
let mut reaped = false;
let exited = shutdown_join_poll(
1,
std::time::Duration::from_millis(30),
|| true,
|| reaped = true,
std::time::Instant::now,
|_| panic!("must not sleep when already finished"),
);
assert!(exited, "finished thread must be joined successfully");
assert!(reaped, "the reap callback must run when the check passes");
}
#[test]
fn shutdown_join_poll_abandons_after_deadline() {
let base = std::time::Instant::now();
let mut elapsed = std::time::Duration::ZERO;
let mut clock = move || {
elapsed += std::time::Duration::from_millis(10);
base + elapsed
};
let mut slept: Vec<std::time::Duration> = Vec::new();
let exited = shutdown_join_poll(
1,
std::time::Duration::from_millis(30),
|| false,
|| panic!("must not reap a thread that has not finished"),
&mut clock,
|d| slept.push(d),
);
assert!(!exited, "stuck thread must be abandoned, not joined");
assert!(!slept.is_empty(), "the poll loop must sleep while waiting");
assert!(
slept
.iter()
.all(|d| *d <= std::time::Duration::from_millis(50)),
"each sleep must respect the 50 ms poll cap"
);
}
#[test]
fn poll_join_with_grace_abandons_stuck_thread() {
let (tx, rx) = mpsc::channel::<()>();
let handle = std::thread::spawn(move || {
let _ = rx.recv();
});
let base = std::time::Instant::now();
let mut elapsed = std::time::Duration::ZERO;
let mut clock = move || {
elapsed += std::time::Duration::from_millis(10);
base + elapsed
};
let mut no_sleep = |_: std::time::Duration| {};
let exited = poll_join_with_grace(
handle,
1,
std::time::Duration::from_millis(30),
&mut clock,
&mut no_sleep,
);
assert!(!exited, "stuck thread must be abandoned, not joined");
drop(tx);
}
}