use super::*;
use crate::broadcast::test_sink;
use crate::sessions::SessionMetadata;
use choreo_proto::{DaemonMessage, SessionEvent, SessionStatus};
use std::collections::HashMap;
use std::sync::mpsc;
use std::time::{Duration, Instant};
pub(super) fn make_daemon_state() -> (DaemonState, mpsc::Receiver<DaemonCommand>) {
let (daemon_tx, daemon_rx) = mpsc::channel();
let dir = tempfile::tempdir().unwrap();
let db = Arc::new(redb::Database::create(dir.path().join("test.redb")).unwrap());
let tool_registry = crate::tools::ToolRegistry::new().build();
let config_dir: &'static tempfile::TempDir = Box::leak(Box::new(tempfile::tempdir().unwrap()));
let accounts_path = config_dir.path().join("accounts.toml");
let state = DaemonState {
next_session_id: 1,
max_turns: 10,
active_sessions: HashMap::new(),
session_metadata: HashMap::new(),
deleted_sessions: HashSet::new(),
children: HashMap::new(),
accounts: AccountManager::load(&accounts_path).unwrap(),
daemon_registry: choreo_ai_protocols::SocketRegistry::default(),
session_registries: HashMap::new(),
credentials: HashMap::new(),
x_credentials: None,
locked: true,
db,
tool_registry,
daemon_tx,
summary_subscribers: HashMap::new(),
client_writers: HashMap::new(),
activity_subscribers: HashMap::new(),
client_subscribed_sessions: HashMap::new(),
global_lag: Arc::new(AtomicUsize::new(0)),
lag_limits: LagLimits::default(),
model_cache: HashMap::new(),
model_prefetch_in_flight: HashSet::new(),
mcp_manager: crate::mcp::McpManager::empty(),
maintenance_tx: None,
acl: None,
catalog_paths: CatalogPaths::default(),
};
(state, daemon_rx)
}
fn seed_credentialed_account(
state: &mut DaemonState,
name: &str,
provider_slug: &str,
) -> crate::accounts::AccountConfig {
seed_credentialed_account_with_url(state, name, provider_slug, None)
}
fn seed_credentialed_account_with_url(
state: &mut DaemonState,
name: &str,
provider_slug: &str,
base_url: Option<String>,
) -> crate::accounts::AccountConfig {
let mut config = crate::accounts::AccountConfig::simple(name, provider_slug);
config.base_url = base_url;
config.retry_max_attempts = Some(1);
config.connect_timeout_secs = Some(2);
config.request_timeout_secs = Some(5);
config.total_timeout_secs = Some(10);
state.accounts.add(config.clone()).unwrap();
state.credentials.insert(
name.to_string(),
ServiceCredential::ApiKey {
key: "test-key".to_string(),
},
);
config
}
fn dead_base_url() -> String {
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let _ = listener.set_nonblocking(true);
format!("http://{}", listener.local_addr().unwrap())
}
mod suspend_tests {
use super::*;
use choreo_ai_protocols::SocketRegistry;
use std::os::unix::net::UnixStream;
fn registered_pair(registry: &SocketRegistry) -> UnixStream {
let (a, b) = UnixStream::pair().unwrap();
let dup = a.try_clone().unwrap();
registry.register(dup);
drop(a);
b
}
#[test]
fn sleep_force_closes_registered_sockets() {
let registry = SocketRegistry::new();
let mut peer = registered_pair(®istry);
assert_eq!(registry.registered_count(), 1);
let empty: std::collections::HashMap<u64, SocketRegistry> = HashMap::new();
handle_suspend_event(&SuspendEvent::Sleep, ®istry, &empty);
assert_eq!(registry.registered_count(), 0);
let mut buf = [0u8; 1];
let n = std::io::Read::read(&mut peer, &mut buf).unwrap();
assert_eq!(n, 0, "peer must see EOF after sleep force-close");
}
#[test]
fn sleep_force_closes_session_registries_too() {
let daemon_registry = SocketRegistry::new();
let mut sessions: HashMap<u64, SocketRegistry> = HashMap::new();
let mut peers = Vec::new();
for id in [1u64, 2, 3] {
let r = SocketRegistry::new();
peers.push(registered_pair(&r));
assert_eq!(r.registered_count(), 1);
sessions.insert(id, r);
}
assert_eq!(daemon_registry.registered_count(), 0);
handle_suspend_event(&SuspendEvent::Sleep, &daemon_registry, &sessions);
assert_eq!(daemon_registry.registered_count(), 0);
for (id, r) in &sessions {
assert_eq!(r.registered_count(), 0, "session {id} registry cleared");
}
for peer in &mut peers {
let mut buf = [0u8; 1];
let n = std::io::Read::read(peer, &mut buf).unwrap();
assert_eq!(n, 0, "session peer must see EOF after sleep force-close");
}
}
#[test]
fn wake_prunes_dead_but_keeps_live_sockets() {
let registry = SocketRegistry::new();
let _live_keep = live_pair(®istry);
dead_pair(®istry); let mut sessions: HashMap<u64, SocketRegistry> = HashMap::new();
let session_registry = SocketRegistry::new();
let _session_live = live_pair(&session_registry);
dead_pair(&session_registry);
sessions.insert(1, session_registry);
handle_suspend_event(&SuspendEvent::Wake, ®istry, &sessions);
assert_eq!(registry.registered_count(), 1);
assert_eq!(sessions[&1].registered_count(), 1);
}
fn live_pair(registry: &SocketRegistry) -> (UnixStream, UnixStream) {
let (a, b) = UnixStream::pair().expect("unix pair");
registry.register(a.try_clone().expect("dup"));
(a, b)
}
fn dead_pair(registry: &SocketRegistry) {
let (a, b) = UnixStream::pair().expect("unix pair");
registry.register(a.try_clone().expect("dup"));
drop((a, b)); }
}
mod cancel_isolation_tests {
use super::*;
use choreo_ai_protocols::SocketRegistry;
use std::os::unix::net::UnixStream;
fn register_pair(registry: &SocketRegistry) -> (UnixStream, UnixStream) {
let (a, b) = UnixStream::pair().unwrap();
registry.register(a.try_clone().unwrap());
(a, b)
}
fn seed_session(state: &mut DaemonState, id: u64) -> SocketRegistry {
let registry = SocketRegistry::default();
state.session_registries.insert(id, registry.clone());
let (cmd_tx, _cmd_rx) = mpsc::channel();
state.active_sessions.insert(
id,
ActiveSessionEntry {
cmd_tx,
handle: std::thread::Builder::new()
.spawn(|| ())
.expect("spawn placeholder thread"),
},
);
registry
}
fn peer_saw_close(b: &mut UnixStream) -> bool {
use std::time::Duration;
let _ = b.set_read_timeout(Some(Duration::from_secs(1)));
let mut buf = [0u8; 1];
matches!(std::io::Read::read(b, &mut buf), Ok(0))
}
#[test]
fn cancel_of_session_a_leaves_session_b_registry_untouched() {
let (mut state, _daemon_rx) = make_daemon_state();
let (_a_handle, mut a_peer) = {
let r = seed_session(&mut state, 1);
register_pair(&r)
};
let (mut b_handle, _b_peer) = {
let r = seed_session(&mut state, 2);
register_pair(&r)
};
state.handle_cancel_request(1, 7);
assert!(peer_saw_close(&mut a_peer), "cancelled session sees EOF");
assert_eq!(state.session_registries[&2].registered_count(), 1);
use std::io::Write;
b_handle
.write_all(b"x")
.expect("uncancelled session's socket must survive the cancel");
}
#[test]
fn cancel_of_parent_closes_children_registries() {
let (mut state, _daemon_rx) = make_daemon_state();
let (parent_handle, mut parent_peer) = {
let r = seed_session(&mut state, 1);
register_pair(&r)
};
let (child_handle, mut child_peer) = {
let r = seed_session(&mut state, 11);
register_pair(&r)
};
state.children.insert(1, vec![11]);
state.handle_cancel_request(1, 7);
assert!(peer_saw_close(&mut parent_peer));
assert!(peer_saw_close(&mut child_peer));
let (outside, _outside_peer) = {
let r = seed_session(&mut state, 2);
register_pair(&r)
};
assert_eq!(state.session_registries[&2].registered_count(), 1);
drop((parent_handle, child_handle, outside));
}
#[test]
fn cancel_of_unknown_session_is_a_noop() {
let (mut state, _daemon_rx) = make_daemon_state();
let (handle, mut peer) = {
let r = seed_session(&mut state, 1);
register_pair(&r)
};
state.handle_cancel_request(999, 7);
assert_eq!(state.session_registries[&1].registered_count(), 1);
drop(handle);
assert!(!peer_saw_close(&mut peer));
}
}
#[test]
fn handle_evict_client_removes_from_maps_and_sends_advisory() {
let (mut state, _rx) = make_daemon_state();
let (sink, rx) = test_sink();
state.client_writers.insert(7, sink.clone());
state.summary_subscribers.insert(7, sink.clone());
state.activity_subscribers.insert(7, sink.clone());
state
.client_subscribed_sessions
.insert(7, HashSet::from([1]));
state.handle_command(DaemonCommand::EvictClient { client_id: 7 });
assert!(!state.client_writers.contains_key(&7));
assert!(!state.summary_subscribers.contains_key(&7));
assert!(!state.activity_subscribers.contains_key(&7));
assert!(!state.client_subscribed_sessions.contains_key(&7));
assert_eq!(rx.recv().unwrap(), DaemonMessage::Evicted);
}
#[test]
fn handle_evict_client_is_idempotent_for_unknown_client() {
let (mut state, _rx) = make_daemon_state();
state.handle_command(DaemonCommand::EvictClient { client_id: 999 });
assert!(state.client_writers.is_empty());
}
#[test]
fn handle_evict_largest_lagging_evicts_biggest_backlog() {
let (mut state, _rx) = make_daemon_state();
let (sink_small, _) = test_sink();
sink_small.bytes_in_flight.store(10, Ordering::Relaxed);
state.client_writers.insert(1, sink_small);
let (sink_big, _) = test_sink();
sink_big.bytes_in_flight.store(1_000, Ordering::Relaxed);
state.client_writers.insert(2, sink_big);
state.handle_command(DaemonCommand::EvictLargestLagging);
assert!(
!state.client_writers.contains_key(&2),
"the largest backlog must be evicted"
);
assert!(state.client_writers.contains_key(&1));
}
#[test]
fn handle_evict_largest_lagging_noop_when_all_healthy() {
let (mut state, _rx) = make_daemon_state();
let (sink, _) = test_sink();
state.client_writers.insert(1, sink);
state.handle_command(DaemonCommand::EvictLargestLagging);
assert!(state.client_writers.contains_key(&1));
}
#[test]
fn handle_list_sessions_empty() {
let (mut state, _rx) = make_daemon_state();
let (reply, rx) = mpsc::channel();
state.handle_command(DaemonCommand::ListSessions { reply });
let sessions = rx.recv().unwrap();
assert!(sessions.is_empty());
}
#[test]
fn handle_list_sessions_with_metadata() {
let (mut state, _rx) = make_daemon_state();
state.session_metadata.insert(
1,
SessionMetadata {
title: Some("test".into()),
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
created_at: 1000,
last_modified: 1000,
turn_count: 3,
status: SessionStatus::Inactive,
active_tool_groups: vec!["core".into()],
account_name: None,
accumulated_usage: TokenUsage::default(),
context_window: None,
last_prompt_tokens: None,
},
);
let (reply, rx) = mpsc::channel();
state.handle_command(DaemonCommand::ListSessions { reply });
let sessions: Vec<SessionSummary> = rx.recv().unwrap();
assert_eq!(sessions.len(), 1);
assert_eq!(sessions[0].session_id, 1);
assert_eq!(sessions[0].title.as_deref(), Some("test"));
}
#[test]
fn handle_list_sessions_orders_by_last_modified_desc() {
let (mut state, _rx) = make_daemon_state();
for (id, created, modified) in [(1, 1000, 1000), (2, 2000, 9000), (3, 3000, 5000)] {
state.session_metadata.insert(
id,
SessionMetadata {
title: Some(format!("s{id}")),
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
created_at: created,
last_modified: modified,
turn_count: 0,
status: SessionStatus::Inactive,
active_tool_groups: vec![],
account_name: None,
accumulated_usage: TokenUsage::default(),
context_window: None,
last_prompt_tokens: None,
},
);
}
let (reply, rx) = mpsc::channel();
state.handle_command(DaemonCommand::ListSessions { reply });
let sessions: Vec<SessionSummary> = rx.recv().unwrap();
let ids: Vec<u64> = sessions.iter().map(|s| s.session_id).collect();
assert_eq!(ids, vec![2, 3, 1]);
}
#[test]
fn handle_list_sessions_tiebreaks_by_session_id_desc() {
let (mut state, _rx) = make_daemon_state();
for id in [1u64, 2, 3] {
state.session_metadata.insert(
id,
SessionMetadata {
title: None,
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
created_at: id as i64 * 1000,
last_modified: 5000,
turn_count: 0,
status: SessionStatus::Inactive,
active_tool_groups: vec![],
account_name: None,
accumulated_usage: TokenUsage::default(),
context_window: None,
last_prompt_tokens: None,
},
);
}
let (reply, rx) = mpsc::channel();
state.handle_command(DaemonCommand::ListSessions { reply });
let sessions: Vec<SessionSummary> = rx.recv().unwrap();
let ids: Vec<u64> = sessions.iter().map(|s| s.session_id).collect();
assert_eq!(ids, vec![3, 2, 1]);
}
#[test]
fn handle_get_session_missing() {
let (mut state, _rx) = make_daemon_state();
let (reply, rx) = mpsc::channel();
state.handle_command(DaemonCommand::GetSession {
session_id: 1,
reply,
});
let result = rx.recv().unwrap();
assert!(result.is_none());
}
#[test]
fn handle_update_metadata() {
let (mut state, _rx) = make_daemon_state();
state.session_metadata.insert(
1,
SessionMetadata {
title: Some("original".into()),
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
created_at: 1000,
last_modified: 1000,
turn_count: 0,
status: SessionStatus::Inactive,
active_tool_groups: vec!["core".into()],
account_name: None,
accumulated_usage: TokenUsage::default(),
context_window: None,
last_prompt_tokens: None,
},
);
let new_meta = SessionMetadata {
title: Some("updated".into()),
selected_model: Some("gpt-4".into()),
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
created_at: 2000,
last_modified: 2000,
turn_count: 5,
status: SessionStatus::Inference,
active_tool_groups: vec!["core".into(), "git".into()],
account_name: None,
accumulated_usage: TokenUsage::default(),
context_window: None,
last_prompt_tokens: None,
};
state.handle_command(DaemonCommand::UpdateMetadata {
session_id: 1,
metadata: new_meta.clone(),
});
let stored = state.session_metadata.get(&1).unwrap();
assert_eq!(stored.title.as_deref(), Some("updated"));
assert_eq!(stored.selected_model.as_deref(), Some("gpt-4"));
assert_eq!(stored.turn_count, 5);
assert_eq!(stored.status, SessionStatus::Inference);
}
#[test]
fn handle_update_metadata_preserves_sleeping_status_after_exit() {
let (mut state, _rx) = make_daemon_state();
state.session_metadata.insert(
1,
SessionMetadata {
title: Some("exited".into()),
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
created_at: 1000,
last_modified: 5000,
turn_count: 3,
status: SessionStatus::Sleeping,
active_tool_groups: vec![],
account_name: None,
accumulated_usage: TokenUsage::default(),
context_window: None,
last_prompt_tokens: None,
},
);
let stale = SessionMetadata {
title: Some("exited".into()),
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
created_at: 1000,
last_modified: 5000,
turn_count: 4,
status: SessionStatus::Inactive,
active_tool_groups: vec![],
account_name: None,
accumulated_usage: TokenUsage::default(),
context_window: None,
last_prompt_tokens: None,
};
state.handle_command(DaemonCommand::UpdateMetadata {
session_id: 1,
metadata: stale,
});
let stored = state.session_metadata.get(&1).unwrap();
assert_eq!(
stored.status,
SessionStatus::Sleeping,
"exited session must not regress to a stale status"
);
assert_eq!(stored.turn_count, 4);
}
#[test]
fn handle_session_exited_nonexistent() {
let (mut state, _rx) = make_daemon_state();
state.handle_command(DaemonCommand::SessionExited { session_id: 999 });
assert!(!state.session_metadata.contains_key(&999));
}
#[test]
fn handle_get_credential_locked() {
let (mut state, _rx) = make_daemon_state();
let (reply, rx) = mpsc::channel();
state.handle_command(DaemonCommand::GetCredential {
service: "openai".into(),
reply,
});
let key = rx.recv().unwrap();
assert!(key.is_none());
}
#[test]
fn handle_register_unregister_subscriber() {
let (mut state, _rx) = make_daemon_state();
let (tx, _rx_sub) = test_sink();
assert!(!state.summary_subscribers.contains_key(&42));
state.handle_command(DaemonCommand::RegisterSummarySubscriber {
client_id: 42,
writer: tx,
});
assert!(state.summary_subscribers.contains_key(&42));
state.handle_command(DaemonCommand::UnregisterSummarySubscriber { client_id: 42 });
assert!(!state.summary_subscribers.contains_key(&42));
}
#[test]
fn handle_broadcast_session_status() {
let (mut state, _rx) = make_daemon_state();
state.session_metadata.insert(
42,
SessionMetadata {
title: None,
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
created_at: 1000,
last_modified: 1000,
turn_count: 0,
status: SessionStatus::Inactive,
active_tool_groups: vec![],
account_name: None,
accumulated_usage: TokenUsage::default(),
context_window: None,
last_prompt_tokens: None,
},
);
let (tx, rx) = test_sink();
state.handle_command(DaemonCommand::RegisterSummarySubscriber {
client_id: 1,
writer: tx,
});
state.handle_command(DaemonCommand::BroadcastSessionStatus {
session_id: 42,
status: SessionStatus::Inference,
});
let msg = rx.recv().unwrap();
assert!(matches!(
msg,
DaemonMessage::Session {
session_id: Some(42),
event: SessionEvent::SessionStatusChanged {
status: SessionStatus::Inference,
..
},
}
));
let meta = state.session_metadata.get(&42).expect("index updated");
assert_eq!(meta.status, SessionStatus::Inference);
assert_eq!(
meta.last_modified, 1000,
"status transitions must not bump last_modified \
(only completed requests and explicit edits do)"
);
}
#[test]
fn handle_broadcast_session_status_dedups_against_session_and_activity_subscribers() {
let (mut state, _rx) = make_daemon_state();
state.session_metadata.insert(
42,
SessionMetadata {
title: None,
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
created_at: 1000,
last_modified: 1000,
turn_count: 0,
status: SessionStatus::Inactive,
active_tool_groups: vec![],
account_name: None,
accumulated_usage: TokenUsage::default(),
context_window: None,
last_prompt_tokens: None,
},
);
let (tx1, rx1) = test_sink();
state.handle_command(DaemonCommand::RegisterSummarySubscriber {
client_id: 1,
writer: tx1,
});
state.handle_command(DaemonCommand::TrackSessionSubscription {
client_id: 1,
session_id: 42,
});
let (tx2, rx2) = test_sink();
state.handle_command(DaemonCommand::RegisterSummarySubscriber {
client_id: 2,
writer: tx2.clone(),
});
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 2,
writer: tx2,
});
drain_send_on_subscribe(&rx2);
let (tx3, rx3) = test_sink();
state.handle_command(DaemonCommand::RegisterSummarySubscriber {
client_id: 3,
writer: tx3,
});
state.handle_command(DaemonCommand::BroadcastSessionStatus {
session_id: 42,
status: SessionStatus::Inference,
});
assert!(
rx1.try_recv().is_err(),
"session subscriber must not get a duplicate via the summary fan-out"
);
assert!(
rx2.try_recv().is_err(),
"activity subscriber must not get a duplicate via the summary fan-out"
);
let msg = rx3.recv().unwrap();
assert!(matches!(
msg,
DaemonMessage::Session {
session_id: Some(42),
event: SessionEvent::SessionStatusChanged {
status: SessionStatus::Inference,
..
},
}
));
assert!(
rx3.try_recv().is_err(),
"summary-only client gets exactly one copy"
);
}
#[test]
fn handle_create_session_succeeds_when_locked() {
let (mut state, _rx) = make_daemon_state();
let (reply, rx) = mpsc::channel();
state.handle_command(DaemonCommand::CreateSession {
title: None,
parent_session_id: None,
working_dir: None,
reasoning_effort: None,
selected_model: None,
context_config: None,
account_name: None,
active_tool_groups: Vec::new(),
reply,
});
let result = rx.recv().unwrap();
assert!(
result.is_ok(),
"CreateSession should succeed even when locked: {:?}",
result.err()
);
}
#[test]
fn handle_delete_session_succeeds_when_locked() {
let (mut state, _rx) = make_daemon_state();
let (reply, rx) = mpsc::channel();
state.handle_command(DaemonCommand::DeleteSession {
session_id: 1,
reply,
});
let result = rx.recv().unwrap();
assert!(
result.is_ok(),
"DeleteSession should succeed even when locked: {:?}",
result.err()
);
}
#[test]
fn handle_attach_session_rejects_deleted_session() {
let (mut state, _daemon_rx) = make_daemon_state();
let record = SessionRecord {
title: Some("ghost".into()),
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
turn_count: 0,
created_at: 1000,
last_modified: 1000,
active_tool_groups: vec![],
context_config: ContextConfig::default(),
account_name: None,
last_response_id: None,
last_response_id_producer: None,
};
db::write_session(&state.db, 1, &record).unwrap();
state.deleted_sessions.insert(1);
let (reply, rx) = mpsc::channel();
state.handle_command(DaemonCommand::AttachSession {
session_id: 1,
reply,
});
let result = rx.recv().unwrap();
assert!(
result.is_err(),
"deleted session must not be resurrected via attach"
);
assert!(!state.session_metadata.contains_key(&1));
assert!(!state.active_sessions.contains_key(&1));
}
#[test]
fn session_exited_finalizes_pending_delete() {
let (mut state, daemon_rx) = make_daemon_state();
let record = SessionRecord {
title: Some("doomed".into()),
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
turn_count: 0,
created_at: 1000,
last_modified: 1000,
active_tool_groups: vec![],
context_config: ContextConfig::default(),
account_name: None,
last_response_id: None,
last_response_id_producer: None,
};
db::write_session(&state.db, 7, &record).unwrap();
db::mark_session_deleted(&state.db, 7).unwrap();
state.deleted_sessions.insert(7);
state.handle_command(DaemonCommand::SessionExited { session_id: 7 });
match daemon_rx.recv() {
Ok(DaemonCommand::SessionDeleteFinalized { session_id: 7 }) => {
state.handle_command(DaemonCommand::SessionDeleteFinalized { session_id: 7 });
}
other => panic!(
"expected SessionDeleteFinalized for session 7, got {:?}",
std::mem::discriminant(&other)
),
}
assert!(!state.deleted_sessions.contains(&7));
assert!(db::read_session(&state.db, 7).unwrap().is_none());
assert_eq!(
db::purge_tombstoned_sessions(&state.db).unwrap(),
0,
"tombstone must be cleared once the record is deleted"
);
}
#[test]
fn delete_finished_session_guards_against_straggler_resurrection() {
let (mut state, daemon_rx) = make_daemon_state();
let record = SessionRecord {
title: Some("doomed".into()),
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
turn_count: 0,
created_at: 1000,
last_modified: 1000,
active_tool_groups: vec![],
context_config: ContextConfig::default(),
account_name: None,
last_response_id: None,
last_response_id_producer: None,
};
db::write_session(&state.db, 12, &record).unwrap();
db::mark_session_deleted(&state.db, 12).unwrap();
state.delete_finished_session(12).unwrap();
assert!(db::read_session(&state.db, 12).unwrap().is_none());
assert!(
state.deleted_sessions.contains(&12),
"fast path must set the deleted marker so stragglers cannot resurrect the session"
);
assert!(!state.session_metadata.contains_key(&12));
assert_eq!(
db::purge_tombstoned_sessions(&state.db).unwrap(),
0,
"stale tombstone must be cleared on the fast-path delete"
);
state.handle_command(DaemonCommand::UpdateMetadata {
session_id: 12,
metadata: SessionMetadata {
title: Some("doomed".into()),
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
created_at: 1000,
last_modified: 2000,
turn_count: 0,
status: SessionStatus::Inactive,
active_tool_groups: vec![],
account_name: None,
accumulated_usage: TokenUsage::default(),
context_window: None,
last_prompt_tokens: None,
},
});
assert!(
!state.session_metadata.contains_key(&12),
"straggler UpdateMetadata must not resurrect a deleted session"
);
state.handle_command(DaemonCommand::SessionExited { session_id: 12 });
match daemon_rx.recv() {
Ok(DaemonCommand::SessionDeleteFinalized { session_id: 12 }) => {
state.handle_command(DaemonCommand::SessionDeleteFinalized { session_id: 12 });
}
other => panic!(
"expected SessionDeleteFinalized for session 12, got {:?}",
std::mem::discriminant(&other)
),
}
assert!(
!state.deleted_sessions.contains(&12),
"marker must be dropped once the finalize confirms the record is gone"
);
}
#[test]
fn delete_session_clears_stale_tombstone_when_no_live_thread() {
let (mut state, _daemon_rx) = make_daemon_state();
let record = SessionRecord {
title: Some("stale".into()),
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
turn_count: 0,
created_at: 1000,
last_modified: 1000,
active_tool_groups: vec![],
context_config: ContextConfig::default(),
account_name: None,
last_response_id: None,
last_response_id_producer: None,
};
db::write_session(&state.db, 3, &record).unwrap();
db::mark_session_deleted(&state.db, 3).unwrap();
let (reply, rx) = mpsc::channel();
state.handle_command(DaemonCommand::DeleteSession {
session_id: 3,
reply,
});
assert!(rx.recv().unwrap().is_ok());
assert!(db::read_session(&state.db, 3).unwrap().is_none());
assert_eq!(
db::purge_tombstoned_sessions(&state.db).unwrap(),
0,
"stale tombstone must be cleared on immediate delete"
);
assert!(!state.deleted_sessions.contains(&3));
}
#[test]
fn delete_session_defers_when_thread_alive() {
let (mut state, _daemon_rx) = make_daemon_state();
let record = SessionRecord {
title: Some("deferred".into()),
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
turn_count: 0,
created_at: 1000,
last_modified: 1000,
active_tool_groups: vec![],
context_config: ContextConfig::default(),
account_name: None,
last_response_id: None,
last_response_id_producer: None,
};
db::write_session(&state.db, 4, &record).unwrap();
let (_cmd_rx, release_tx) = insert_active_session(&mut state, 4);
let (reply, rx) = mpsc::channel();
state.handle_command(DaemonCommand::DeleteSession {
session_id: 4,
reply,
});
assert!(rx.recv().unwrap().is_ok());
assert!(db::read_session(&state.db, 4).unwrap().is_some());
assert!(state.deleted_sessions.contains(&4));
assert_eq!(
db::purge_tombstoned_sessions(&state.db).unwrap(),
1,
"deferred delete must write a tombstone"
);
drop(release_tx);
}
#[test]
fn delete_session_keeps_tombstone_while_a_delete_is_pending() {
let (mut state, _daemon_rx) = make_daemon_state();
let record = SessionRecord {
title: Some("double-deleted".into()),
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
turn_count: 0,
created_at: 1000,
last_modified: 1000,
active_tool_groups: vec![],
context_config: ContextConfig::default(),
account_name: None,
last_response_id: None,
last_response_id_producer: None,
};
db::write_session(&state.db, 5, &record).unwrap();
let (_cmd_rx, release_tx) = insert_active_session(&mut state, 5);
let (reply, rx) = mpsc::channel();
state.handle_command(DaemonCommand::DeleteSession {
session_id: 5,
reply,
});
assert!(rx.recv().unwrap().is_ok());
assert!(state.deleted_sessions.contains(&5));
let (reply, rx) = mpsc::channel();
state.handle_command(DaemonCommand::DeleteSession {
session_id: 5,
reply,
});
assert!(rx.recv().unwrap().is_ok());
assert!(
state.deleted_sessions.contains(&5),
"marker must stay while the deferred delete is pending"
);
assert_eq!(
db::purge_tombstoned_sessions(&state.db).unwrap(),
1,
"the pending delete's tombstone must not be swept by a second delete"
);
drop(release_tx);
}
#[test]
fn session_delete_finalized_drops_marker() {
let (mut state, _daemon_rx) = make_daemon_state();
state.deleted_sessions.insert(9);
state.handle_command(DaemonCommand::SessionDeleteFinalized { session_id: 9 });
assert!(!state.deleted_sessions.contains(&9));
}
#[test]
fn session_exited_does_not_delete_non_deleted_session() {
let (mut state, _daemon_rx) = make_daemon_state();
let record = SessionRecord {
title: Some("alive".into()),
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
turn_count: 0,
created_at: 1000,
last_modified: 1000,
active_tool_groups: vec![],
context_config: ContextConfig::default(),
account_name: None,
last_response_id: None,
last_response_id_producer: None,
};
db::write_session(&state.db, 8, &record).unwrap();
state.handle_command(DaemonCommand::SessionExited { session_id: 8 });
assert!(db::read_session(&state.db, 8).unwrap().is_some());
}
#[test]
fn broadcast_sends_to_subscriber() {
let (mut state, _rx) = make_daemon_state();
let (tx, rx) = test_sink();
state.summary_subscribers.insert(1, tx);
let msg = DaemonMessage::Session {
session_id: Some(42),
event: SessionEvent::SessionDeleted,
};
state.broadcast(msg.clone());
let received = rx.recv().unwrap();
assert_eq!(received, msg);
assert!(state.summary_subscribers.contains_key(&1));
}
#[test]
fn broadcast_removes_disconnected_subscriber() {
let (mut state, _rx) = make_daemon_state();
let (tx, rx) = test_sink();
state.summary_subscribers.insert(1, tx);
drop(rx); state.broadcast(DaemonMessage::Session {
session_id: Some(42),
event: SessionEvent::SessionDeleted,
});
assert!(!state.summary_subscribers.contains_key(&1));
}
#[test]
fn broadcast_enqueues_losslessly_and_evicts_over_lag_client() {
let (mut state, _rx) = make_daemon_state();
state.lag_limits = LagLimits {
per_client_cap: 16,
global_budget: usize::MAX,
};
let (sink, rx) = test_sink();
state.summary_subscribers.insert(7, sink.clone());
state.client_writers.insert(7, sink);
let msg = DaemonMessage::Session {
session_id: Some(42),
event: SessionEvent::SessionDeleted,
};
state.broadcast(msg.clone());
assert_eq!(rx.recv().unwrap(), msg);
assert!(
!state.summary_subscribers.contains_key(&7),
"over-lag subscriber must be evicted from the summary map"
);
assert!(
!state.client_writers.contains_key(&7),
"over-lag subscriber must be evicted from the writer registry"
);
}
#[test]
#[serial_test::serial(catalog)]
fn broadcast_lifecycle_delivers_to_summary_and_activity_exactly_once_per_client() {
let (mut state, _rx) = make_daemon_state();
let (tx1, rx1) = test_sink();
let (tx2, rx2) = test_sink();
let (tx3, rx3) = test_sink();
state.handle_command(DaemonCommand::RegisterSummarySubscriber {
client_id: 1,
writer: tx1,
});
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 2,
writer: tx2,
});
state.handle_command(DaemonCommand::RegisterSummarySubscriber {
client_id: 3,
writer: tx3.clone(),
});
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 3,
writer: tx3,
});
drain_send_on_subscribe(&rx2);
drain_send_on_subscribe(&rx3);
let msg = DaemonMessage::Session {
session_id: Some(42),
event: SessionEvent::SessionDeleted,
};
state.broadcast(msg.clone());
assert_eq!(rx1.recv().unwrap(), msg);
assert_eq!(rx2.recv().unwrap(), msg);
assert_eq!(rx3.recv().unwrap(), msg);
assert!(
rx3.try_recv().is_err(),
"a summary+activity client must receive the lifecycle event exactly once"
);
}
#[test]
fn handle_validate_model_allows_through_when_no_session() {
let (mut state, _rx) = make_daemon_state();
let (reply, rx) = mpsc::channel();
state.handle_command(DaemonCommand::ValidateModel {
session_id: 999,
model: "gpt-4".into(),
reply,
});
let result = rx.recv().unwrap();
assert_eq!(result, Ok(()));
}
#[test]
fn handle_validate_model_rejects_when_no_provider() {
let (mut state, _rx) = make_daemon_state();
state.session_metadata.insert(
1,
SessionMetadata {
title: None,
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
created_at: 1000,
last_modified: 1000,
turn_count: 0,
status: SessionStatus::Sleeping,
active_tool_groups: vec![],
account_name: Some("locked-account".into()),
accumulated_usage: TokenUsage::default(),
context_window: None,
last_prompt_tokens: None,
},
);
let (reply, rx) = mpsc::channel();
state.handle_command(DaemonCommand::ValidateModel {
session_id: 1,
model: "gpt-4".into(),
reply,
});
let result = rx.recv().unwrap();
assert!(result.is_err());
let err = result.unwrap_err();
assert!(err.contains("locked"), "error should mention locked daemon");
assert!(
err.contains("locked-account"),
"error should mention the account"
);
}
#[test]
fn handle_validate_model_rejects_unknown_model() {
let (mut state, _rx) = make_daemon_state();
state.session_metadata.insert(
1,
SessionMetadata {
title: None,
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
created_at: 1000,
last_modified: 1000,
turn_count: 0,
status: SessionStatus::Inactive,
active_tool_groups: vec![],
account_name: Some("test-account".into()),
accumulated_usage: TokenUsage::default(),
context_window: None,
last_prompt_tokens: None,
},
);
seed_credentialed_account(&mut state, "test-account", "openai");
state.model_cache.insert(
"test-account".into(),
(vec!["gpt-4".into(), "gpt-3.5".into()], Instant::now()),
);
let (reply, rx) = mpsc::channel();
state.handle_command(DaemonCommand::ValidateModel {
session_id: 1,
model: "nonexistent-model".into(),
reply,
});
let result = rx.recv().unwrap();
assert!(result.is_err());
let err = result.unwrap_err();
assert!(
err.contains("nonexistent-model"),
"error should mention the model name"
);
assert!(err.contains("gpt-4"), "error should list available models");
}
#[test]
fn handle_validate_model_allows_known_model() {
let (mut state, _rx) = make_daemon_state();
state.session_metadata.insert(
1,
SessionMetadata {
title: None,
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
created_at: 1000,
last_modified: 1000,
turn_count: 0,
status: SessionStatus::Inactive,
active_tool_groups: vec![],
account_name: Some("test-account".into()),
accumulated_usage: TokenUsage::default(),
context_window: None,
last_prompt_tokens: None,
},
);
seed_credentialed_account(&mut state, "test-account", "openai");
state.model_cache.insert(
"test-account".into(),
(vec!["gpt-4".into(), "gpt-3.5".into()], Instant::now()),
);
let (reply, rx) = mpsc::channel();
state.handle_command(DaemonCommand::ValidateModel {
session_id: 1,
model: "gpt-4".into(),
reply,
});
let result = rx.recv().unwrap();
assert_eq!(result, Ok(()));
}
#[test]
fn handle_set_session_title_forwards_to_session() {
let (mut state, _daemon_rx) = make_daemon_state();
let (cmd_tx, cmd_rx) = mpsc::channel();
let (handle_tx, handle_rx) = std::sync::mpsc::channel::<()>();
let handle = std::thread::spawn(move || {
let _ = handle_rx.recv();
});
state
.active_sessions
.insert(1, ActiveSessionEntry { cmd_tx, handle });
state.session_metadata.insert(
1,
SessionMetadata {
title: Some("old title".into()),
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
created_at: 1000,
last_modified: 1000,
turn_count: 0,
status: SessionStatus::Inactive,
active_tool_groups: vec!["core".into()],
account_name: None,
accumulated_usage: TokenUsage::default(),
context_window: None,
last_prompt_tokens: None,
},
);
state.handle_command(DaemonCommand::SetSessionTitle {
session_id: 1,
title: "new title".into(),
});
match cmd_rx.try_recv() {
Ok(SessionCommand::SetTitle { title }) => {
assert_eq!(title, "new title");
}
Ok(_) => {
panic!("expected SetTitle, got a different SessionCommand variant");
}
Err(e) => {
panic!("expected SetTitle, got error: {e}");
}
}
let _ = handle_tx.send(());
}
#[test]
fn handle_set_session_title_nonexistent_session_logs_warning() {
let (mut state, _rx) = make_daemon_state();
state.handle_command(DaemonCommand::SetSessionTitle {
session_id: 999,
title: "ghost title".into(),
});
}
fn insert_active_session(
state: &mut DaemonState,
session_id: u64,
) -> (mpsc::Receiver<SessionCommand>, mpsc::Sender<()>) {
let (cmd_tx, cmd_rx) = mpsc::channel();
let (release_tx, release_rx) = mpsc::channel::<()>();
let handle = std::thread::spawn(move || {
let _ = release_rx.recv();
});
state
.active_sessions
.insert(session_id, ActiveSessionEntry { cmd_tx, handle });
(cmd_rx, release_tx)
}
#[test]
fn handle_set_working_dir_forwards_to_session() {
let (mut state, _daemon_rx) = make_daemon_state();
let (cmd_rx, release_tx) = insert_active_session(&mut state, 1);
state.handle_command(DaemonCommand::SetWorkingDir {
session_id: 1,
path: PathBuf::from("/tmp"),
reply: mpsc::channel().0,
});
match cmd_rx.try_recv() {
Ok(SessionCommand::SetWorkingDir { path, .. }) => {
assert_eq!(path, PathBuf::from("/tmp"));
}
Ok(_) => panic!("expected SetWorkingDir, got a different SessionCommand variant"),
Err(e) => panic!("expected SetWorkingDir, got error: {e}"),
}
let _ = release_tx.send(());
}
#[test]
fn handle_set_working_dir_nonexistent_session_replies_error() {
let (mut state, _daemon_rx) = make_daemon_state();
let (reply_tx, reply_rx) = mpsc::channel();
state.handle_command(DaemonCommand::SetWorkingDir {
session_id: 999,
path: PathBuf::from("/tmp"),
reply: reply_tx,
});
match reply_rx.recv() {
Ok(Err(msg)) => assert!(msg.contains("not active"), "unexpected msg: {msg}"),
Ok(Ok(_)) => panic!("expected an error reply for an inactive session"),
Err(e) => panic!("expected error reply, got {e:?}"),
}
}
#[test]
fn handle_load_tools_forwards_to_session() {
let (mut state, _daemon_rx) = make_daemon_state();
let (cmd_rx, release_tx) = insert_active_session(&mut state, 1);
state.handle_command(DaemonCommand::LoadTools {
session_id: 1,
groups: vec!["x".into()],
reply: mpsc::channel().0,
});
match cmd_rx.try_recv() {
Ok(SessionCommand::LoadTools { groups, .. }) => {
assert_eq!(groups, vec!["x"]);
}
Ok(_) => panic!("expected LoadTools, got a different SessionCommand variant"),
Err(e) => panic!("expected LoadTools, got error: {e}"),
}
let _ = release_tx.send(());
}
#[test]
fn handle_load_tools_nonexistent_session_replies_error() {
let (mut state, _daemon_rx) = make_daemon_state();
let (reply_tx, reply_rx) = mpsc::channel();
state.handle_command(DaemonCommand::LoadTools {
session_id: 999,
groups: vec!["x".into()],
reply: reply_tx,
});
match reply_rx.recv() {
Ok(Err(msg)) => assert!(msg.contains("not active"), "unexpected msg: {msg}"),
Ok(Ok(_)) => panic!("expected an error reply for an inactive session"),
Err(e) => panic!("expected error reply, got {e:?}"),
}
}
#[test]
fn handle_unload_tools_forwards_to_session() {
let (mut state, _daemon_rx) = make_daemon_state();
let (cmd_rx, release_tx) = insert_active_session(&mut state, 1);
state.handle_command(DaemonCommand::UnloadTools {
session_id: 1,
groups: vec!["x".into()],
reply: mpsc::channel().0,
});
match cmd_rx.try_recv() {
Ok(SessionCommand::UnloadTools { groups, .. }) => {
assert_eq!(groups, vec!["x"]);
}
Ok(_) => panic!("expected UnloadTools, got a different SessionCommand variant"),
Err(e) => panic!("expected UnloadTools, got error: {e}"),
}
let _ = release_tx.send(());
}
#[test]
fn handle_unload_tools_nonexistent_session_replies_error() {
let (mut state, _daemon_rx) = make_daemon_state();
let (reply_tx, reply_rx) = mpsc::channel();
state.handle_command(DaemonCommand::UnloadTools {
session_id: 999,
groups: vec!["x".into()],
reply: reply_tx,
});
match reply_rx.recv() {
Ok(Err(msg)) => assert!(msg.contains("not active"), "unexpected msg: {msg}"),
Ok(Ok(_)) => panic!("expected an error reply for an inactive session"),
Err(e) => panic!("expected error reply, got {e:?}"),
}
}
fn drain_send_on_subscribe(rx: &crossbeam_channel::Receiver<DaemonMessage>) {
let msg = rx.recv().unwrap();
assert!(
matches!(&msg, DaemonMessage::CatalogUpdated { providers } if !providers.is_empty()),
"expected the send-on-subscribe CatalogUpdated, got {msg:?}",
);
match rx.recv().unwrap() {
DaemonMessage::Locked | DaemonMessage::Unlocked => {}
other => panic!("expected the send-on-subscribe lock state, got {other:?}"),
}
}
const ACL_KEY_A: [u8; 32] = [1u8; 32];
const ACL_KEY_B: [u8; 32] = [2u8; 32];
fn acl_b64(key: &[u8; 32]) -> String {
use base64::Engine as _;
base64::engine::general_purpose::STANDARD.encode(key)
}
fn make_acl_state() -> (DaemonState, tempfile::TempDir) {
let (mut state, _rx) = make_daemon_state();
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("authorized_clients.toml");
std::fs::write(
&path,
format!("[[client]]\npubkey = \"{}\"\n", acl_b64(&ACL_KEY_A)),
)
.unwrap();
state.acl = Some(crate::server::acl::SharedAcl::load(&path));
(state, dir)
}
#[test]
fn handle_acl_add_enrolls_key_updates_file_and_broadcasts() {
let (mut state, dir) = make_acl_state();
let acl_path = state.acl.as_ref().unwrap().path().to_path_buf();
let _ = &dir;
let (writer, writer_rx) = test_sink();
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 1,
writer,
});
drain_send_on_subscribe(&writer_rx);
let (reply_tx, reply_rx) = mpsc::channel();
state.handle_command(DaemonCommand::AclAddCmd {
pubkey: acl_b64(&ACL_KEY_B),
reply: reply_tx,
});
assert_eq!(reply_rx.recv().unwrap().unwrap(), 2);
let file = std::fs::read_to_string(&acl_path).unwrap();
assert!(file.contains(&acl_b64(&ACL_KEY_A)), "existing key survives");
assert!(file.contains(&acl_b64(&ACL_KEY_B)), "new key written");
assert!(state.acl.as_ref().unwrap().contains(&ACL_KEY_B));
match writer_rx.recv().unwrap() {
DaemonMessage::AclUpdated { clients } => assert_eq!(clients, 2),
other => panic!("expected AclUpdated broadcast, got {other:?}"),
}
}
#[test]
fn handle_acl_add_is_idempotent_for_an_already_trusted_key() {
let (mut state, dir) = make_acl_state();
let acl_path = state.acl.as_ref().unwrap().path().to_path_buf();
let before = std::fs::read_to_string(&acl_path).unwrap();
let _ = &dir;
let (reply_tx, reply_rx) = mpsc::channel();
state.handle_command(DaemonCommand::AclAddCmd {
pubkey: acl_b64(&ACL_KEY_A),
reply: reply_tx,
});
assert_eq!(reply_rx.recv().unwrap().unwrap(), 1);
assert_eq!(
std::fs::read_to_string(&acl_path).unwrap(),
before,
"re-adding a trusted key must not rewrite the file"
);
}
#[test]
fn handle_acl_add_rejects_a_bad_key() {
let (mut state, _dir) = make_acl_state();
let (reply_tx, reply_rx) = mpsc::channel();
state.handle_command(DaemonCommand::AclAddCmd {
pubkey: "not-base64!!!".to_string(),
reply: reply_tx,
});
assert!(reply_rx.recv().unwrap().is_err(), "bad base64 must fail");
use base64::Engine as _;
let short = base64::engine::general_purpose::STANDARD.encode([9u8; 16]);
let (reply_tx, reply_rx) = mpsc::channel();
state.handle_command(DaemonCommand::AclAddCmd {
pubkey: short,
reply: reply_tx,
});
assert!(reply_rx.recv().unwrap().is_err(), "wrong length must fail");
}
#[test]
fn handle_accounts_reload_applies_external_change_and_broadcasts() {
let (mut state, _rx) = make_daemon_state();
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("accounts.toml");
state.accounts = AccountManager::load(&path).unwrap();
assert!(state.accounts.is_empty(), "fresh manager starts empty");
let (writer, writer_rx) = test_sink();
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 1,
writer,
});
drain_send_on_subscribe(&writer_rx);
std::fs::write(
&path,
"[[account]]\nname = \"alpha\"\nprovider = \"openai\"\n\n[[account]]\nname = \"beta\"\nprovider = \"anthropic\"\n",
)
.unwrap();
state.handle_command(DaemonCommand::AccountsReload);
assert!(state.accounts.contains("alpha"));
assert!(state.accounts.contains("beta"));
match writer_rx.recv().unwrap() {
DaemonMessage::Accounts { accounts } => {
let names: Vec<&str> = accounts.iter().map(|a| a.name.as_str()).collect();
assert!(names.contains(&"alpha"), "broadcast carries alpha");
assert!(names.contains(&"beta"), "broadcast carries beta");
}
other => panic!("expected Accounts broadcast, got {other:?}"),
}
}
#[test]
fn handle_accounts_reload_noops_when_logically_unchanged() {
let (mut state, _rx) = make_daemon_state();
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("accounts.toml");
std::fs::write(
&path,
"[[account]]\nname = \"alpha\"\nprovider = \"openai\"\n",
)
.unwrap();
state.accounts = AccountManager::load(&path).unwrap();
let (writer, writer_rx) = test_sink();
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 1,
writer,
});
drain_send_on_subscribe(&writer_rx);
state.accounts.save().unwrap();
state.handle_command(DaemonCommand::AccountsReload);
assert!(
writer_rx.try_recv().is_err(),
"no broadcast for a logically-unchanged reload"
);
assert_eq!(state.accounts.names(), vec!["alpha".to_string()]);
}
fn insert_active_session_with_account(
state: &mut DaemonState,
session_id: u64,
account: &str,
) -> (mpsc::Receiver<SessionCommand>, mpsc::Sender<()>) {
state.session_metadata.insert(
session_id,
SessionMetadata {
title: None,
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
created_at: 1000,
last_modified: 1000,
turn_count: 0,
status: SessionStatus::Inactive,
active_tool_groups: vec![],
account_name: Some(account.to_string()),
accumulated_usage: TokenUsage::default(),
context_window: None,
last_prompt_tokens: None,
},
);
insert_active_session(state, session_id)
}
#[test]
fn handle_accounts_reload_invalidates_session_clients_of_removed_account() {
let (mut state, _rx) = make_daemon_state();
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("accounts.toml");
state.accounts = AccountManager::load(&path).unwrap();
std::fs::write(
&path,
"[[account]]\nname = \"keep\"\nprovider = \"openai\"\n\n[[account]]\nname = \"gone\"\nprovider = \"anthropic\"\n",
)
.unwrap();
seed_credentialed_account(&mut state, "keep", "openai");
seed_credentialed_account(&mut state, "gone", "anthropic");
let (keep_cmd_rx, keep_release) = insert_active_session_with_account(&mut state, 1, "keep");
let (gone_cmd_rx, gone_release) = insert_active_session_with_account(&mut state, 2, "gone");
let saved = std::fs::read_to_string(&path).unwrap();
let edited: String = format!(
"[[account]]{}",
saved
.split("[[account]]")
.skip(1)
.filter(|section| !section.contains("name = \"gone\""))
.collect::<Vec<_>>()
.join("[[account]]")
);
std::fs::write(&path, edited).unwrap();
state.handle_command(DaemonCommand::AccountsReload);
assert!(state.accounts.contains("keep"));
assert!(!state.accounts.contains("gone"));
assert!(
matches!(gone_cmd_rx.try_recv(), Ok(SessionCommand::DropProvider)),
"removed account's session client must be invalidated"
);
assert!(
keep_cmd_rx.try_recv().is_err(),
"session bound to an untouched account must NOT be invalidated — \
its cached client and connection pool stay warm"
);
drop(keep_release);
drop(gone_release);
}
#[test]
fn handle_accounts_reload_drops_stale_client_for_modified_account() {
let (mut state, _rx) = make_daemon_state();
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("accounts.toml");
state.accounts = AccountManager::load(&path).unwrap();
seed_credentialed_account(&mut state, "keep", "openai");
state.accounts = AccountManager::load(&path).unwrap();
let (cmd_rx, release) = insert_active_session_with_account(&mut state, 1, "keep");
std::fs::write(
&path,
"[[account]]\nname = \"keep\"\nprovider = \"bogus\"\n",
)
.unwrap();
state.handle_command(DaemonCommand::AccountsReload);
assert_eq!(state.accounts.get("keep").unwrap().provider, "bogus");
assert!(
matches!(cmd_rx.try_recv(), Ok(SessionCommand::DropProvider)),
"modified account's session client must be invalidated"
);
drop(release);
}
#[test]
fn handle_accounts_reload_noops_without_a_real_path() {
let (mut state, _rx) = make_daemon_state();
state.accounts = AccountManager::empty();
state.handle_command(DaemonCommand::AccountsReload);
assert!(state.accounts.is_empty());
}
#[test]
#[serial_test::serial(catalog)]
fn handle_register_activity_subscriber_adds_to_map() {
let (mut state, _rx) = make_daemon_state();
let (tx, _) = test_sink();
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 10,
writer: tx,
});
assert!(state.activity_subscribers.contains_key(&10));
}
#[test]
#[serial_test::serial(catalog)]
fn handle_register_activity_subscriber_replaces_existing() {
let (mut state, _rx) = make_daemon_state();
let (tx1, _) = test_sink();
let (tx2, _) = test_sink();
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 10,
writer: tx1,
});
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 10,
writer: tx2,
});
assert!(state.activity_subscribers.contains_key(&10));
}
#[test]
#[serial_test::serial(catalog)]
fn handle_unregister_activity_subscriber_preserves_session_tracking() {
let (mut state, _rx) = make_daemon_state();
let (tx, _) = test_sink();
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 10,
writer: tx,
});
state.handle_command(DaemonCommand::TrackSessionSubscription {
client_id: 10,
session_id: 42,
});
state.handle_command(DaemonCommand::UnregisterActivitySubscriber { client_id: 10 });
assert!(!state.activity_subscribers.contains_key(&10));
assert!(state.client_subscribed_sessions.contains_key(&10));
let sessions = state.client_subscribed_sessions.get(&10).unwrap();
assert!(sessions.contains(&42));
assert_eq!(sessions.len(), 1);
}
#[test]
#[serial_test::serial(catalog)]
fn handle_client_disconnected_clears_all_tracking() {
let (mut state, _rx) = make_daemon_state();
let (tx, _) = test_sink();
state.handle_command(DaemonCommand::RegisterSummarySubscriber {
client_id: 10,
writer: tx.clone(),
});
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 10,
writer: tx,
});
state.handle_command(DaemonCommand::TrackSessionSubscription {
client_id: 10,
session_id: 1,
});
state.handle_command(DaemonCommand::TrackSessionSubscription {
client_id: 10,
session_id: 2,
});
assert!(state.summary_subscribers.contains_key(&10));
assert!(state.activity_subscribers.contains_key(&10));
assert!(state.client_subscribed_sessions.contains_key(&10));
state.handle_command(DaemonCommand::ClientDisconnected { client_id: 10 });
assert!(!state.summary_subscribers.contains_key(&10));
assert!(!state.activity_subscribers.contains_key(&10));
assert!(!state.client_subscribed_sessions.contains_key(&10));
}
#[test]
fn handle_client_disconnected_noop_for_unknown_client() {
let (mut state, _rx) = make_daemon_state();
state.handle_command(DaemonCommand::ClientDisconnected { client_id: 999 });
}
#[test]
fn handle_track_session_subscription_adds_entry() {
let (mut state, _rx) = make_daemon_state();
state.handle_command(DaemonCommand::TrackSessionSubscription {
client_id: 10,
session_id: 42,
});
let sessions = state
.client_subscribed_sessions
.get(&10)
.expect("client should have entry");
assert!(sessions.contains(&42));
assert_eq!(sessions.len(), 1);
}
#[test]
fn handle_track_session_subscription_idempotent_re_attach() {
let (mut state, _rx) = make_daemon_state();
state.handle_command(DaemonCommand::TrackSessionSubscription {
client_id: 10,
session_id: 42,
});
state.handle_command(DaemonCommand::TrackSessionSubscription {
client_id: 10,
session_id: 42,
});
let sessions = state
.client_subscribed_sessions
.get(&10)
.expect("client should have entry");
assert!(sessions.contains(&42));
assert_eq!(sessions.len(), 1, "should not duplicate session_id");
}
#[test]
fn handle_track_session_subscription_tracks_multiple_sessions() {
let (mut state, _rx) = make_daemon_state();
state.handle_command(DaemonCommand::TrackSessionSubscription {
client_id: 10,
session_id: 42,
});
state.handle_command(DaemonCommand::TrackSessionSubscription {
client_id: 10,
session_id: 99,
});
let sessions = state
.client_subscribed_sessions
.get(&10)
.expect("client should have entry");
assert!(sessions.contains(&42));
assert!(sessions.contains(&99));
assert_eq!(sessions.len(), 2);
}
#[test]
fn handle_untrack_session_subscription_removes_session() {
let (mut state, _rx) = make_daemon_state();
state.handle_command(DaemonCommand::TrackSessionSubscription {
client_id: 10,
session_id: 42,
});
state.handle_command(DaemonCommand::TrackSessionSubscription {
client_id: 10,
session_id: 99,
});
state.handle_command(DaemonCommand::UntrackSessionSubscription {
client_id: 10,
session_id: 42,
});
let sessions = state
.client_subscribed_sessions
.get(&10)
.expect("client should still have entry");
assert!(!sessions.contains(&42));
assert!(sessions.contains(&99));
assert_eq!(sessions.len(), 1);
}
#[test]
fn handle_untrack_session_subscription_removes_client_when_empty() {
let (mut state, _rx) = make_daemon_state();
state.handle_command(DaemonCommand::TrackSessionSubscription {
client_id: 10,
session_id: 42,
});
state.handle_command(DaemonCommand::UntrackSessionSubscription {
client_id: 10,
session_id: 42,
});
assert!(!state.client_subscribed_sessions.contains_key(&10));
}
#[test]
fn handle_untrack_session_subscription_noop_for_unknown_session() {
let (mut state, _rx) = make_daemon_state();
state.handle_command(DaemonCommand::UntrackSessionSubscription {
client_id: 10,
session_id: 42,
});
assert!(!state.client_subscribed_sessions.contains_key(&10));
}
#[test]
fn handle_untrack_session_subscription_noop_for_unknown_client() {
let (mut state, _rx) = make_daemon_state();
state.handle_command(DaemonCommand::UntrackSessionSubscription {
client_id: 999,
session_id: 42,
});
}
#[test]
#[serial_test::serial(catalog)]
fn handle_broadcast_activity_sends_to_subscriber() {
let (mut state, _rx) = make_daemon_state();
let (tx, rx) = test_sink();
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 10,
writer: tx,
});
drain_send_on_subscribe(&rx);
let msg = DaemonMessage::Session {
session_id: Some(1),
event: SessionEvent::OutputChunk {
request_id: 5,
stream: choreo_proto::OutputStream::Answer,
data: b"hello".to_vec(),
},
};
state.handle_command(DaemonCommand::BroadcastActivity {
session_id: Some(1),
msg: msg.clone(),
});
let received = rx.recv().unwrap();
assert_eq!(received, msg);
assert!(state.activity_subscribers.contains_key(&10));
}
#[test]
#[serial_test::serial(catalog)]
fn handle_broadcast_activity_skips_dedup_for_session_subscriber() {
let (mut state, _rx) = make_daemon_state();
let (tx, rx) = test_sink();
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 10,
writer: tx,
});
drain_send_on_subscribe(&rx);
state.handle_command(DaemonCommand::TrackSessionSubscription {
client_id: 10,
session_id: 1,
});
let msg = DaemonMessage::Session {
session_id: Some(1),
event: SessionEvent::OutputChunk {
request_id: 5,
stream: choreo_proto::OutputStream::Answer,
data: b"hello".to_vec(),
},
};
state.handle_command(DaemonCommand::BroadcastActivity {
session_id: Some(1),
msg,
});
assert!(
rx.try_recv().is_err(),
"message should have been suppressed for session subscriber"
);
assert!(state.activity_subscribers.contains_key(&10));
}
#[test]
#[serial_test::serial(catalog)]
fn handle_broadcast_activity_no_dedup_for_different_session() {
let (mut state, _rx) = make_daemon_state();
let (tx, rx) = test_sink();
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 10,
writer: tx,
});
drain_send_on_subscribe(&rx);
state.handle_command(DaemonCommand::TrackSessionSubscription {
client_id: 10,
session_id: 1,
});
let msg = DaemonMessage::Session {
session_id: Some(2),
event: SessionEvent::OutputChunk {
request_id: 5,
stream: choreo_proto::OutputStream::Answer,
data: b"hello".to_vec(),
},
};
state.handle_command(DaemonCommand::BroadcastActivity {
session_id: Some(2),
msg: msg.clone(),
});
let received = rx.recv().unwrap();
assert_eq!(received, msg);
}
#[test]
#[serial_test::serial(catalog)]
fn handle_broadcast_activity_sends_when_no_session_id() {
let (mut state, _rx) = make_daemon_state();
let (tx, rx) = test_sink();
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 10,
writer: tx,
});
drain_send_on_subscribe(&rx);
let msg = DaemonMessage::Models {
models: vec!["gpt-4".into()],
selected_model: Some("gpt-4".into()),
};
state.handle_command(DaemonCommand::BroadcastActivity {
session_id: None,
msg: msg.clone(),
});
let received = rx.recv().unwrap();
assert_eq!(received, msg);
}
#[test]
#[serial_test::serial(catalog)]
fn handle_broadcast_activity_removes_disconnected_subscriber() {
let (mut state, _rx) = make_daemon_state();
let (tx, rx) = test_sink();
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 10,
writer: tx,
});
drop(rx);
let msg = DaemonMessage::Session {
session_id: Some(1),
event: SessionEvent::SessionStatusChanged {
status: SessionStatus::Inactive,
last_modified: 0,
},
};
state.handle_command(DaemonCommand::BroadcastActivity {
session_id: Some(1),
msg,
});
assert!(!state.activity_subscribers.contains_key(&10));
}
#[test]
#[serial_test::serial(catalog)]
fn handle_broadcast_activity_evicts_over_lag_subscriber() {
let (mut state, _rx) = make_daemon_state();
state.lag_limits = LagLimits {
per_client_cap: 16,
global_budget: usize::MAX,
};
let (tx, rx) = test_sink();
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 10,
writer: tx.clone(),
});
state.client_writers.insert(10, tx);
drain_send_on_subscribe(&rx);
let broadcast = DaemonMessage::Session {
session_id: Some(7),
event: SessionEvent::OutputChunk {
request_id: 99,
stream: choreo_proto::OutputStream::Answer,
data: b"hello".to_vec(),
},
};
state.handle_command(DaemonCommand::BroadcastActivity {
session_id: Some(7),
msg: broadcast.clone(),
});
assert_eq!(rx.recv().unwrap(), broadcast);
assert!(
!state.activity_subscribers.contains_key(&10),
"over-lag subscriber must be evicted from the activity map"
);
assert!(
!state.client_writers.contains_key(&10),
"over-lag subscriber must be evicted from the writer registry"
);
}
#[test]
#[serial_test::serial(catalog)]
fn handle_broadcast_activity_handles_multiple_clients() {
let (mut state, _rx) = make_daemon_state();
let (tx1, rx1) = test_sink();
let (tx2, rx2) = test_sink();
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 10,
writer: tx1,
});
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 20,
writer: tx2,
});
drain_send_on_subscribe(&rx1);
drain_send_on_subscribe(&rx2);
state.handle_command(DaemonCommand::TrackSessionSubscription {
client_id: 10,
session_id: 1,
});
let msg = DaemonMessage::Session {
session_id: Some(1),
event: SessionEvent::OutputChunk {
request_id: 5,
stream: choreo_proto::OutputStream::Answer,
data: b"data".to_vec(),
},
};
state.handle_command(DaemonCommand::BroadcastActivity {
session_id: Some(1),
msg: msg.clone(),
});
assert!(
rx1.try_recv().is_err(),
"client 10 is a session subscriber, should be suppressed"
);
let received = rx2.recv().unwrap();
assert_eq!(received, msg);
}
#[test]
#[serial_test::serial(catalog)]
fn handle_broadcast_activity_dedup_keyed_on_command_origin_not_message_shape() {
let (mut state, _rx) = make_daemon_state();
let (tx, rx) = test_sink();
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 10,
writer: tx,
});
drain_send_on_subscribe(&rx);
state.handle_command(DaemonCommand::TrackSessionSubscription {
client_id: 10,
session_id: 42,
});
let msg = DaemonMessage::Sessions { sessions: vec![] };
state.handle_command(DaemonCommand::BroadcastActivity {
session_id: Some(42),
msg,
});
assert!(
rx.try_recv().is_err(),
"message should have been suppressed: origin came from the command, not the payload"
);
assert!(state.activity_subscribers.contains_key(&10));
}
#[test]
fn broadcast_origin_contract_requires_agreeing_provenance() {
assert!(
super::subscriber_handlers::violates_broadcast_origin_contract(
Some(42),
&DaemonMessage::Sessions { sessions: vec![] },
)
);
assert!(
super::subscriber_handlers::violates_broadcast_origin_contract(
Some(42),
&DaemonMessage::CatalogUpdated { providers: vec![] },
)
);
let session_msg = DaemonMessage::Session {
session_id: Some(42),
event: SessionEvent::OutputChunk {
request_id: 1,
stream: choreo_proto::OutputStream::Answer,
data: vec![],
},
};
let other_session_msg = DaemonMessage::Session {
session_id: Some(7),
event: SessionEvent::OutputChunk {
request_id: 1,
stream: choreo_proto::OutputStream::Answer,
data: vec![],
},
};
assert!(
super::subscriber_handlers::violates_broadcast_origin_contract(
Some(42),
&other_session_msg,
)
);
assert!(
super::subscriber_handlers::violates_broadcast_origin_contract(
Some(42),
&DaemonMessage::Session {
session_id: None,
event: SessionEvent::Failed {
request_id: 1,
error: "no session attached".into(),
},
},
)
);
assert!(super::subscriber_handlers::violates_broadcast_origin_contract(None, &session_msg,));
assert!(
!super::subscriber_handlers::violates_broadcast_origin_contract(Some(42), &session_msg,)
);
assert!(
!super::subscriber_handlers::violates_broadcast_origin_contract(
None,
&DaemonMessage::CatalogUpdated { providers: vec![] },
)
);
assert!(
!super::subscriber_handlers::violates_broadcast_origin_contract(
None,
&DaemonMessage::Session {
session_id: None,
event: SessionEvent::Failed {
request_id: 1,
error: "no session attached".into(),
},
},
)
);
}
fn bundled_catalog() -> Vec<choreo_ai_protocols::ProviderEntry> {
merge_overlay(
&choreo_ai_protocols::load_bundled_base(),
bundled_overlay_src(),
)
}
struct RestoreBundledCatalogOnDrop;
impl Drop for RestoreBundledCatalogOnDrop {
fn drop(&mut self) {
replace_catalog(bundled_catalog());
}
}
fn tiny_base() -> Vec<choreo_ai_protocols::ProviderEntry> {
vec![choreo_ai_protocols::ProviderEntry {
slug: "tiny-test".into(),
display_name: "Tiny Test".into(),
protocol: choreo_ai_protocols::ProviderProtocol::OpenAi {
max_tokens_field: choreo_ai_protocols::MaxTokensField::MaxCompletionTokens,
},
base_url: "https://tiny.example/v1".into(),
default_model: "tiny-1".into(),
models: vec![choreo_ai_protocols::ModelEntry {
model: "tiny-1".into(),
context_window: 4096,
reasoning_supported: true,
max_output_tokens: 2048,
..Default::default()
}],
}]
}
#[test]
#[serial_test::serial(catalog)]
fn refresh_models_without_maintenance_thread_replies_error() {
let (mut state, _rx) = make_daemon_state();
let (reply, rx) = mpsc::channel();
state.handle_command(DaemonCommand::RefreshModels {
force: false,
reply,
});
let result = rx.recv().unwrap();
assert!(result.is_err(), "no maintenance thread → error reply");
let err = result.unwrap_err();
assert!(
err.contains("maintenance thread"),
"unexpected error: {err}"
);
}
#[test]
#[serial_test::serial(catalog)]
fn refresh_models_with_dead_maintenance_thread_replies_error() {
let (mut state, _rx) = make_daemon_state();
let (maintenance_tx, maintenance_rx) = crossbeam_channel::unbounded::<MaintenanceEvent>();
drop(maintenance_rx); state.maintenance_tx = Some(maintenance_tx);
let (reply, rx) = mpsc::channel();
state.handle_command(DaemonCommand::RefreshModels {
force: false,
reply,
});
let result = rx.recv().unwrap();
assert!(result.is_err(), "dead maintenance thread → error reply");
let err = result.unwrap_err();
assert!(
err.contains("maintenance thread"),
"unexpected error: {err}"
);
}
#[test]
#[serial_test::serial(catalog)]
fn refresh_models_forwards_to_maintenance_thread() {
let (mut state, _rx) = make_daemon_state();
let (maintenance_tx, maintenance_rx) = crossbeam_channel::unbounded();
state.maintenance_tx = Some(maintenance_tx);
let (reply, _reply_rx) = mpsc::channel();
state.handle_command(DaemonCommand::RefreshModels { force: true, reply });
let msg = maintenance_rx.recv().unwrap();
match msg {
MaintenanceEvent::RefreshNow { force, .. } => assert!(force),
}
}
#[test]
#[serial_test::serial(catalog)]
fn catalog_base_changed_swaps_broadcasts_and_replies() {
let _restore = RestoreBundledCatalogOnDrop;
let (mut state, _rx) = make_daemon_state();
let (writer_tx, writer_rx) = test_sink();
state.activity_subscribers.insert(1, writer_tx);
let (reply, reply_rx) = mpsc::channel();
state.handle_command(DaemonCommand::CatalogBaseChanged {
base: tiny_base(),
etag: Some("\"v42\"".into()),
user_overlay: None,
persist: false,
reply: vec![RefreshRequester {
force: false,
tx: reply,
}],
});
assert_eq!(
choreo_ai_protocols::lookup_provider("tiny-test")
.expect("swapped catalog")
.slug,
"tiny-test"
);
let broadcast = writer_rx.recv().unwrap();
assert!(matches!(
&broadcast,
DaemonMessage::CatalogUpdated { providers } if providers.iter().any(|p| p.slug == "tiny-test")
));
let report = reply_rx.recv().unwrap().expect("refresh succeeds");
assert!(report.providers > 1, "overlay-only providers must survive");
assert!(report.models >= 1);
assert_eq!(report.status, RefreshStatus::Updated);
}
#[test]
#[serial_test::serial(catalog)]
fn catalog_base_changed_user_overlay_merges_on_top() {
let _restore = RestoreBundledCatalogOnDrop;
let (mut state, _rx) = make_daemon_state();
let overlay = r#"
[provider.tiny-test]
display_name = "Renamed By User"
[provider.user-only]
display_name = "User Only"
protocol = "openai"
base_url = "https://user.example/v1"
default_model = "u-1"
[provider.user-only.models."u-1"]
context_window = 1024
"#;
state.handle_command(DaemonCommand::CatalogBaseChanged {
base: tiny_base(),
etag: None,
user_overlay: Some(overlay.to_string()),
persist: false,
reply: Vec::new(),
});
let renamed = choreo_ai_protocols::lookup_provider("tiny-test").expect("tiny-test present");
assert_eq!(renamed.display_name, "Renamed By User");
let user_only =
choreo_ai_protocols::lookup_provider("user-only").expect("user overlay provider");
assert_eq!(user_only.display_name, "User Only");
assert_eq!(user_only.models.len(), 1);
}
#[test]
#[serial_test::serial(catalog)]
fn catalog_base_changed_empty_base_still_yields_overlay_only_providers() {
let _restore = RestoreBundledCatalogOnDrop;
let (mut state, _rx) = make_daemon_state();
let (reply, reply_rx) = mpsc::channel();
state.handle_command(DaemonCommand::CatalogBaseChanged {
base: Vec::new(),
etag: None,
user_overlay: None,
persist: false,
reply: vec![RefreshRequester {
force: false,
tx: reply,
}],
});
let report = reply_rx.recv().unwrap().expect("refresh succeeds");
assert!(report.providers > 1, "overlay-only providers must survive");
assert_eq!(report.status, RefreshStatus::Updated);
assert!(choreo_ai_protocols::lookup_provider("ollama").is_some());
}
#[test]
#[serial_test::serial(catalog)]
fn catalog_not_modified_replies_up_to_date_with_current_counts() {
let (mut state, _rx) = make_daemon_state();
let (reply, reply_rx) = mpsc::channel();
state.handle_command(DaemonCommand::CatalogNotModified {
reply: vec![RefreshRequester {
force: true,
tx: reply,
}],
});
let report = reply_rx.recv().unwrap().expect("304 reply is Ok");
assert_eq!(report.status, RefreshStatus::UpToDate);
assert!(
report.providers > 0,
"counts come from the currently swapped catalog"
);
}
#[test]
fn send_catalog_reply_individualizes_status_per_requester() {
let (forced_tx, forced_rx) = mpsc::channel();
let (plain_tx, plain_rx) = mpsc::channel();
send_catalog_reply(
vec![
RefreshRequester {
force: true,
tx: forced_tx,
},
RefreshRequester {
force: false,
tx: plain_tx,
},
],
208,
1234,
);
let forced = forced_rx.recv().unwrap().expect("forced reply is Ok");
assert_eq!(forced.status, RefreshStatus::Forced);
assert_eq!((forced.providers, forced.models), (208, 1234));
let plain = plain_rx.recv().unwrap().expect("plain reply is Ok");
assert_eq!(plain.status, RefreshStatus::Updated);
}
#[test]
#[serial_test::serial(catalog)]
fn activity_subscriber_gets_current_provider_list_on_register() {
let (mut state, _rx) = make_daemon_state();
let (writer_tx, writer_rx) = test_sink();
state.handle_register_activity_subscriber(1, writer_tx);
let msg = writer_rx.recv().unwrap();
match &msg {
DaemonMessage::CatalogUpdated { providers } => {
assert!(!providers.is_empty());
assert!(providers.iter().any(|p| p.slug == "openai"));
}
other => panic!("expected CatalogUpdated, got {other:?}"),
}
}
#[test]
fn activity_subscriber_gets_current_lock_state_on_register() {
let (mut state, _rx) = make_daemon_state();
let (writer_tx, writer_rx) = test_sink();
state.handle_register_activity_subscriber(1, writer_tx);
let msg = writer_rx.recv().unwrap(); assert!(matches!(&msg, DaemonMessage::CatalogUpdated { .. }));
match writer_rx.recv().unwrap() {
DaemonMessage::Locked => {}
other => panic!("expected subscribe-time Locked, got {other:?}"),
}
state.locked = false;
let (writer_tx2, writer_rx2) = test_sink();
state.handle_register_activity_subscriber(2, writer_tx2);
let _ = writer_rx2.recv().unwrap(); match writer_rx2.recv().unwrap() {
DaemonMessage::Unlocked => {}
other => panic!("expected subscribe-time Unlocked, got {other:?}"),
}
}
#[test]
fn broadcast_lock_state_sends_current_state_to_all_activity_subscribers() {
let (mut state, _rx) = make_daemon_state();
let (writer_a, rx_a) = test_sink();
let (writer_b, rx_b) = test_sink();
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 1,
writer: writer_a,
});
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 2,
writer: writer_b,
});
drain_send_on_subscribe(&rx_a);
drain_send_on_subscribe(&rx_b);
state.locked = false;
state.broadcast_lock_state();
assert!(matches!(rx_a.recv().unwrap(), DaemonMessage::Unlocked));
assert!(matches!(rx_b.recv().unwrap(), DaemonMessage::Unlocked));
state.locked = true;
state.broadcast_lock_state();
assert!(matches!(rx_a.recv().unwrap(), DaemonMessage::Locked));
assert!(matches!(rx_b.recv().unwrap(), DaemonMessage::Locked));
}
#[test]
fn handle_lock_clears_credentials_latches_locked_and_broadcasts() {
let (mut state, _rx) = make_daemon_state();
let (writer, writer_rx) = test_sink();
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 1,
writer,
});
drain_send_on_subscribe(&writer_rx);
state.locked = false;
state.credentials.insert(
"openai".to_string(),
ServiceCredential::ApiKey {
key: "sk-secret".to_string(),
},
);
let (cmd_rx, release) = insert_active_session_with_account(&mut state, 5, "openai");
state.x_credentials = Some(ServiceCredential::ApiKey {
key: "x-secret".to_string(),
});
let (reply, reply_rx) = mpsc::channel();
state.handle_command(DaemonCommand::Lock { reply });
assert!(reply_rx.recv().unwrap().is_ok());
assert!(state.locked, "/lock must latch the locked state");
assert!(state.credentials.is_empty(), "credentials cleared");
assert!(state.x_credentials.is_none(), "x credential cleared");
assert!(
matches!(cmd_rx.try_recv(), Ok(SessionCommand::DropProvider)),
"/lock must invalidate live session clients"
);
drop(release);
match writer_rx.recv().unwrap() {
DaemonMessage::Locked => {}
other => panic!("expected Locked transition broadcast, got {other:?}"),
}
}
#[test]
fn handle_lock_when_already_locked_does_not_rebroadcast() {
let (mut state, _rx) = make_daemon_state();
let (writer, writer_rx) = test_sink();
state.handle_command(DaemonCommand::RegisterActivitySubscriber {
client_id: 1,
writer,
});
drain_send_on_subscribe(&writer_rx);
let (reply, reply_rx) = mpsc::channel();
state.handle_command(DaemonCommand::Lock { reply });
assert!(reply_rx.recv().unwrap().is_ok());
assert!(state.locked);
assert!(
writer_rx.try_recv().is_err(),
"locking an already-locked daemon must not re-broadcast the Locked state"
);
}
#[test]
fn should_prefetch_models_gates_on_account_flight_and_freshness() {
let (mut state, _rx) = make_daemon_state();
assert!(!state.should_prefetch_models("acct"));
seed_credentialed_account(&mut state, "acct", "openai");
assert!(state.should_prefetch_models("acct"));
state.model_prefetch_in_flight.insert("acct".into());
assert!(!state.should_prefetch_models("acct"));
state.model_prefetch_in_flight.clear();
state
.model_cache
.insert("acct".into(), (vec!["m".into()], Instant::now()));
assert!(!state.should_prefetch_models("acct"));
state.model_cache.insert(
"acct".into(),
(
vec!["m".into()],
Instant::now()
.checked_sub(MODEL_CACHE_TTL)
.unwrap()
.checked_sub(Duration::from_secs(1))
.unwrap(),
),
);
assert!(state.should_prefetch_models("acct"));
}
#[test]
fn handle_model_prefetch_result_success_populates_cache_and_releases_guard() {
let (mut state, _rx) = make_daemon_state();
seed_credentialed_account(&mut state, "acct", "openai");
state.model_prefetch_in_flight.insert("acct".into());
state.handle_command(DaemonCommand::ModelPrefetchResult {
account: "acct".into(),
result: Ok(vec!["m1".into(), "m2".into()]),
});
assert!(!state.model_prefetch_in_flight.contains("acct"));
let (models, cached_at) = state.model_cache.get("acct").expect("cache populated");
assert_eq!(models, &["m1".to_string(), "m2".to_string()]);
assert!(cached_at.elapsed() < MODEL_CACHE_TTL);
}
#[test]
fn handle_model_prefetch_result_failure_releases_guard_without_caching() {
let (mut state, _rx) = make_daemon_state();
state.model_prefetch_in_flight.insert("acct".into());
state.handle_command(DaemonCommand::ModelPrefetchResult {
account: "acct".into(),
result: Err("provider unreachable".into()),
});
assert!(!state.model_prefetch_in_flight.contains("acct"));
assert!(!state.model_cache.contains_key("acct"));
}
#[test]
fn update_metadata_account_change_spawns_background_prefetch() {
let (mut state, rx) = make_daemon_state();
seed_credentialed_account_with_url(&mut state, "acct", "openai", Some(dead_base_url()));
state.session_metadata.insert(
1,
SessionMetadata {
title: Some("s".into()),
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
created_at: 1000,
last_modified: 1000,
turn_count: 0,
status: SessionStatus::Inactive,
active_tool_groups: vec![],
account_name: None,
accumulated_usage: TokenUsage::default(),
context_window: None,
last_prompt_tokens: None,
},
);
let mut meta = state.session_metadata.get(&1).unwrap().clone();
meta.account_name = Some("acct".into());
state.handle_command(DaemonCommand::UpdateMetadata {
session_id: 1,
metadata: meta.clone(),
});
assert!(
state.model_prefetch_in_flight.contains("acct"),
"account change must spawn a prefetch"
);
let msg = rx.recv().unwrap();
assert!(
matches!(
&msg,
DaemonCommand::ModelPrefetchResult { account, result }
if account == "acct" && result.is_err()
),
"expected ModelPrefetchResult for 'acct' with Err (failing provider)"
);
state.handle_command(msg);
assert!(!state.model_prefetch_in_flight.contains("acct"));
state.handle_command(DaemonCommand::UpdateMetadata {
session_id: 1,
metadata: meta.clone(),
});
assert!(
rx.try_recv().is_err(),
"no prefetch thread may be spawned for an unchanged account"
);
}
#[test]
fn create_session_with_account_spawns_background_prefetch() {
let (mut state, _rx) = make_daemon_state();
seed_credentialed_account_with_url(&mut state, "acct", "openai", Some(dead_base_url()));
let (reply, rx) = mpsc::channel();
state.handle_command(DaemonCommand::CreateSession {
title: None,
parent_session_id: None,
working_dir: None,
reasoning_effort: None,
selected_model: None,
context_config: None,
account_name: Some("acct".into()),
active_tool_groups: Vec::new(),
reply,
});
rx.recv().unwrap().expect("session created");
assert!(
state.model_prefetch_in_flight.contains("acct"),
"session create with an account must spawn a prefetch"
);
}
fn state_with_session_account(account: &str) -> (DaemonState, mpsc::Receiver<DaemonCommand>) {
let (mut state, rx) = make_daemon_state();
seed_credentialed_account_with_url(&mut state, account, "openai", Some(dead_base_url()));
state.session_metadata.insert(
1,
SessionMetadata {
title: Some("s".into()),
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
created_at: 1000,
last_modified: 1000,
turn_count: 0,
status: SessionStatus::Inactive,
active_tool_groups: vec![],
account_name: Some(account.to_string()),
accumulated_usage: TokenUsage::default(),
context_window: None,
last_prompt_tokens: None,
},
);
(state, rx)
}
#[test]
fn list_models_serves_stale_cache_without_duplicate_fetch_while_prefetch_in_flight() {
let (mut state, rx) = state_with_session_account("acct");
state.model_cache.insert(
"acct".into(),
(
vec!["old-model".into()],
Instant::now()
.checked_sub(MODEL_CACHE_TTL)
.unwrap()
.checked_sub(Duration::from_secs(1))
.unwrap(),
),
);
state.model_prefetch_in_flight.insert("acct".into());
let (models, _) = handle_list_models_inner(&mut state, Some(1)).expect("stale list served");
assert_eq!(models, vec!["old-model".to_string()]);
assert!(rx.try_recv().is_err(), "no duplicate prefetch spawned");
}
#[test]
fn list_models_with_cold_cache_triggers_background_prefetch_and_reports_warming() {
let (mut state, rx) = state_with_session_account("acct");
let err = handle_list_models_inner(&mut state, Some(1)).expect_err("cold cache → warming");
assert!(err.contains("warming"), "unexpected error: {err}");
assert!(
state.model_prefetch_in_flight.contains("acct"),
"a background prefetch must have been spawned"
);
let msg = rx.recv().unwrap();
assert!(matches!(
&msg,
DaemonCommand::ModelPrefetchResult { account, result }
if account == "acct" && result.is_err()
));
state.handle_command(msg);
assert!(!state.model_prefetch_in_flight.contains("acct"));
assert!(!state.model_cache.contains_key("acct"));
}
#[test]
fn list_models_with_stale_cache_and_no_prefetch_serves_stale_and_warms_background() {
let (mut state, rx) = state_with_session_account("acct");
state.model_cache.insert(
"acct".into(),
(
vec!["old-model".into()],
Instant::now()
.checked_sub(MODEL_CACHE_TTL)
.unwrap()
.checked_sub(Duration::from_secs(1))
.unwrap(),
),
);
let (models, _) = handle_list_models_inner(&mut state, Some(1)).expect("stale list served");
assert_eq!(models, vec!["old-model".to_string()]);
assert!(
state.model_prefetch_in_flight.contains("acct"),
"stale cache must trigger a background refresh"
);
let _ = rx.recv().unwrap();
}
#[test]
fn prefetch_result_for_removed_account_is_discarded_not_cached() {
let (mut state, _rx) = make_daemon_state();
state.model_prefetch_in_flight.insert("acct".into());
state.handle_command(DaemonCommand::ModelPrefetchResult {
account: "acct".into(),
result: Ok(vec!["stale".into()]),
});
assert!(!state.model_prefetch_in_flight.contains("acct"));
assert!(!state.model_cache.contains_key("acct"), "result discarded");
}
#[test]
fn failed_fetch_releases_in_flight_guard_with_error() {
let (mut state, rx) = make_daemon_state();
seed_credentialed_account_with_url(&mut state, "acct", "openai", Some(dead_base_url()));
state.maybe_spawn_model_prefetch("acct");
assert!(
state.model_prefetch_in_flight.contains("acct"),
"prefetch spawned"
);
let msg = rx.recv().unwrap();
assert!(
matches!(
&msg,
DaemonCommand::ModelPrefetchResult { account, result }
if account == "acct" && result.is_err()
),
"expected a fetch failure reported as an Err"
);
if let DaemonCommand::ModelPrefetchResult { result, .. } = &msg {
assert!(
result.as_ref().is_err(),
"expected an Err result, got {result:?}"
);
}
state.handle_command(msg);
assert!(
!state.model_prefetch_in_flight.contains("acct"),
"guard released"
);
assert!(!state.model_cache.contains_key("acct"));
}
use choreo_keystore::ServiceCredential as TestCred;
use x25519_dalek::StaticSecret as TestSecret;
fn test_pub(key: [u8; 32]) -> [u8; 32] {
*x25519_dalek::PublicKey::from(&TestSecret::from(key)).as_bytes()
}
#[test]
fn bind_keystore_adopts_on_unbound_and_runs_unlock_tail() {
let (mut state, _rx) = make_daemon_state();
let key: [u8; 32] = [3u8; 32];
let (writer, writer_rx) = test_sink();
let (reply, reply_rx) = mpsc::channel();
state.handle_command(DaemonCommand::BindKeystore {
key: key.to_vec(),
client_writer: Some(writer),
reply,
});
reply_rx.recv().unwrap();
assert!(matches!(writer_rx.recv().unwrap(), DaemonMessage::Bound));
assert_eq!(
db::get_keystore_binding(&state.db).unwrap(),
Some(test_pub(key))
);
assert!(!state.locked, "BindKeystore must run the unlock tail");
}
#[test]
fn bind_keystore_on_bound_keystore_rejects_wrong_key_without_overwrite() {
let (mut state, _rx) = make_daemon_state();
let key_a: [u8; 32] = [3u8; 32];
let key_b: [u8; 32] = [4u8; 32];
let (writer, writer_rx) = test_sink();
let (reply, reply_rx) = mpsc::channel();
state.handle_command(DaemonCommand::BindKeystore {
key: key_a.to_vec(),
client_writer: Some(writer),
reply,
});
reply_rx.recv().unwrap();
assert!(matches!(writer_rx.recv().unwrap(), DaemonMessage::Bound));
let (reply, reply_rx) = mpsc::channel();
state.handle_command(DaemonCommand::BindKeystore {
key: key_b.to_vec(),
client_writer: None, reply,
});
reply_rx.recv().unwrap();
assert_eq!(
db::get_keystore_binding(&state.db).unwrap(),
Some(test_pub(key_a))
);
}
#[test]
fn unlock_on_unbound_keystore_is_refused_and_does_not_adopt() {
let (mut state, _rx) = make_daemon_state();
let key: [u8; 32] = [5u8; 32];
let (writer, writer_rx) = test_sink();
let (reply, reply_rx) = mpsc::channel();
state.handle_command(DaemonCommand::Unlock {
private_key: key.to_vec(),
client_writer: Some(writer),
reply,
});
reply_rx.recv().unwrap();
assert!(matches!(
writer_rx.recv().unwrap(),
DaemonMessage::KeystoreUnbound { .. }
));
assert_eq!(db::get_keystore_binding(&state.db).unwrap(), None);
assert!(state.locked, "a refused unlock must not change lock state");
}
#[test]
fn add_credential_on_unbound_keystore_is_refused_without_binding_or_persist() {
let (mut state, _rx) = make_daemon_state();
let key: [u8; 32] = [6u8; 32];
let (writer, writer_rx) = test_sink();
let (reply, reply_rx) = mpsc::channel();
state.handle_command(DaemonCommand::SaveCredential {
service: "svc".to_string(),
encrypted_blob: vec![1, 2, 3],
unlock_key: key.to_vec(),
client_writer: Some(writer),
reply,
});
reply_rx.recv().unwrap();
assert!(matches!(
writer_rx.recv().unwrap(),
DaemonMessage::KeystoreUnbound { .. }
));
assert_eq!(db::get_keystore_binding(&state.db).unwrap(), None);
assert!(db::get_all_credential_blobs(&state.db).unwrap().is_empty());
assert!(state.locked);
}
#[test]
fn add_credential_verify_only_implicitly_unlocks_bound_keystore() {
let (mut state, _rx) = make_daemon_state();
let (sub_writer, sub_rx) = test_sink();
state.handle_register_activity_subscriber(1, sub_writer);
drain_send_on_subscribe(&sub_rx);
let key: [u8; 32] = [7u8; 32];
let (reply, reply_rx) = mpsc::channel();
state.handle_command(DaemonCommand::BindKeystore {
key: key.to_vec(),
client_writer: None,
reply,
});
reply_rx.recv().unwrap();
assert!(matches!(sub_rx.recv().unwrap(), DaemonMessage::Unlocked));
state.locked = true;
let derived = test_pub(key);
let blob = choreo_keystore::crypto::encrypt_with_public_key(
&derived,
&postcard::to_allocvec(&TestCred::ApiKey { key: "k".into() }).unwrap(),
)
.unwrap();
let (writer, writer_rx) = test_sink();
let (reply, reply_rx) = mpsc::channel();
state.handle_command(DaemonCommand::SaveCredential {
service: "svc".to_string(),
encrypted_blob: blob,
unlock_key: key.to_vec(),
client_writer: Some(writer),
reply,
});
reply_rx.recv().unwrap();
assert!(matches!(writer_rx.recv().unwrap(), DaemonMessage::Unlocked));
assert!(matches!(
writer_rx.recv().unwrap(),
DaemonMessage::CredentialAdded { .. }
));
assert!(!state.locked, "valid AddCredential implicitly unlocks");
assert!(matches!(
state.credentials.get("svc"),
Some(TestCred::ApiKey { key }) if key == "k"
));
assert!(matches!(sub_rx.recv().unwrap(), DaemonMessage::Unlocked));
}
#[test]
fn add_credential_on_bound_keystore_rejects_wrong_key_blob() {
let (mut state, _rx) = make_daemon_state();
let key: [u8; 32] = [8u8; 32];
let (reply, reply_rx) = mpsc::channel();
state.handle_command(DaemonCommand::BindKeystore {
key: key.to_vec(),
client_writer: None,
reply,
});
reply_rx.recv().unwrap();
state.locked = true;
let other: [u8; 32] = [9u8; 32];
let (writer, writer_rx) = test_sink();
let (reply, reply_rx) = mpsc::channel();
state.handle_command(DaemonCommand::SaveCredential {
service: "svc".to_string(),
encrypted_blob: vec![1, 2, 3],
unlock_key: other.to_vec(),
client_writer: Some(writer),
reply,
});
reply_rx.recv().unwrap();
assert!(
matches!(&writer_rx.recv().unwrap(),
DaemonMessage::CredentialAddFailed { error, .. } if error.contains("does not match")),
"mismatched key must be rejected with CredentialAddFailed"
);
assert!(db::get_all_credential_blobs(&state.db).unwrap().is_empty());
assert!(state.locked);
}