use crate::daemon::DaemonCommand;
use crate::sessions::SessionCommand;
use choreo_proto::{
ClientMessage, ContextConfig, DaemonMessage, ProtoError, SessionEvent, read_message,
write_message,
};
use std::io::{self, BufReader, BufWriter, Write};
use std::net::{Shutdown, TcpStream};
#[cfg(unix)]
use std::os::unix::net::UnixStream;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::mpsc;
use std::time::{Duration, Instant};
use tracing::{debug, error, info, warn};
#[cfg(windows)]
use uds_windows::UnixStream;
const WRITER_JOIN_GRACE: Duration = Duration::from_secs(5);
const WRITER_WRITE_TIMEOUT: Duration = Duration::from_secs(5);
trait ConnectionWriter {
fn send_message(&mut self, msg: &DaemonMessage) -> Result<(), String>;
fn shutdown(&mut self);
}
impl ConnectionWriter for BufWriter<UnixStream> {
fn send_message(&mut self, msg: &DaemonMessage) -> Result<(), String> {
write_message(self, msg).map_err(|e| e.to_string())?;
self.flush().map_err(|e| e.to_string())
}
fn shutdown(&mut self) {
let _ = self.get_ref().shutdown(Shutdown::Both);
}
}
impl ConnectionWriter for choreo_transport::noise::NoiseStream {
fn send_message(&mut self, msg: &DaemonMessage) -> Result<(), String> {
self.send_daemon_message(msg).map_err(|e| e.to_string())
}
fn shutdown(&mut self) {
let _ = self.get_ref().shutdown(Shutdown::Both);
}
}
struct ChannelConnectionWriter {
tx: Option<crossbeam_channel::Sender<DaemonMessage>>,
}
impl ChannelConnectionWriter {
fn new(tx: crossbeam_channel::Sender<DaemonMessage>) -> Self {
Self { tx: Some(tx) }
}
}
impl ConnectionWriter for ChannelConnectionWriter {
fn send_message(&mut self, msg: &DaemonMessage) -> Result<(), String> {
match &self.tx {
Some(tx) => tx.send(msg.clone()).map_err(|_| {
"embedded client receiver dropped".to_string()
}),
None => Err("embedded writer already shut down".to_string()),
}
}
fn shutdown(&mut self) {
self.tx = None;
}
}
fn writer_thread<W: ConnectionWriter>(
mut writer: W,
rx: crossbeam_channel::Receiver<DaemonMessage>,
bytes: Arc<AtomicUsize>,
global: Arc<AtomicUsize>,
) {
for msg in &rx {
let size = msg.approx_wire_size();
if let Err(e) = writer.send_message(&msg) {
warn!("writer thread error: {e}");
bytes.fetch_sub(size, Ordering::Relaxed);
global.fetch_sub(size, Ordering::Relaxed);
writer.shutdown();
break;
}
bytes.fetch_sub(size, Ordering::Relaxed);
global.fetch_sub(size, Ordering::Relaxed);
if matches!(msg, DaemonMessage::ShuttingDown | DaemonMessage::Evicted) {
writer.shutdown();
break;
}
}
for msg in rx.try_iter() {
let size = msg.approx_wire_size();
bytes.fetch_sub(size, Ordering::Relaxed);
global.fetch_sub(size, Ordering::Relaxed);
}
}
pub(crate) fn register_client_writer(
daemon_tx: &mpsc::Sender<DaemonCommand>,
) -> (
u64,
crate::broadcast::SubscriberSink,
crossbeam_channel::Receiver<DaemonMessage>,
) {
let (writer_tx, writer_rx) = crossbeam_channel::unbounded::<DaemonMessage>();
let sink = crate::broadcast::SubscriberSink::new(writer_tx);
let client_id = rand::random::<u64>();
let _ = daemon_tx.send(DaemonCommand::RegisterClientWriter {
client_id,
writer: sink.clone(),
});
(client_id, sink, writer_rx)
}
fn send_to_writer(ctx: &ClientCtx, msg: DaemonMessage) {
ctx.writer.send_accounted(&msg, ctx.global_lag);
}
struct ClientCtx<'a> {
writer: &'a crate::broadcast::SubscriberSink,
global_lag: &'a AtomicUsize,
daemon_tx: &'a mpsc::Sender<DaemonCommand>,
attached_session_id: &'a mut Option<u64>,
attached_session_tx: &'a mut Option<mpsc::Sender<SessionCommand>>,
client_id: u64,
is_unix: bool,
}
fn cleanup_client(
attached_session_tx: Option<mpsc::Sender<SessionCommand>>,
client_id: u64,
daemon_tx: &mpsc::Sender<DaemonCommand>,
writer: crate::broadcast::SubscriberSink,
writer_handle: std::thread::JoinHandle<()>,
) {
if let Some(ref tx) = attached_session_tx {
let _ = tx.send(SessionCommand::Detach { client_id });
}
let _ = daemon_tx.send(DaemonCommand::ClientDisconnected { client_id });
drop(writer);
crate::server::lifecycle::join_thread_bounded(
writer_handle,
Instant::now() + WRITER_JOIN_GRACE,
);
crate::metrics::record_client_disconnected();
}
fn dispatch_client_message(msg: ClientMessage, ctx: &mut ClientCtx) -> io::Result<()> {
match msg {
ClientMessage::CreateSession {
title,
parent_session_id,
working_dir,
context_config,
account_name,
selected_model,
reasoning_effort,
} => {
if !handle_client_create_session(
title,
parent_session_id,
working_dir,
context_config,
account_name,
selected_model,
reasoning_effort,
ctx,
) {
return Err(io::Error::new(
io::ErrorKind::ConnectionAborted,
"daemon disconnected",
));
}
}
ClientMessage::AttachSession { session_id } => {
if !handle_client_attach_session(session_id, ctx) {
return Err(io::Error::new(
io::ErrorKind::ConnectionAborted,
"daemon disconnected",
));
}
}
ClientMessage::ListSessions => {
debug!("client {}: ListSessions", ctx.client_id);
let (reply, rx) = mpsc::channel();
let _ = ctx.daemon_tx.send(DaemonCommand::ListSessions { reply });
if let Ok(sessions) = rx.recv() {
send_to_writer(ctx, DaemonMessage::Sessions { sessions });
}
}
ClientMessage::SubscribeSessionsSummary => {
let _ = ctx
.daemon_tx
.send(DaemonCommand::RegisterSummarySubscriber {
client_id: ctx.client_id,
writer: ctx.writer.clone(),
});
}
ClientMessage::UnsubscribeSessionsSummary => {
let _ = ctx
.daemon_tx
.send(DaemonCommand::UnregisterSummarySubscriber {
client_id: ctx.client_id,
});
}
ClientMessage::RunInput { request_id, input } => {
debug!("client {}: RunInput id={}", ctx.client_id, request_id);
if let Some(tx) = ctx.attached_session_tx {
let _ = tx.send(SessionCommand::RunInput { request_id, input });
} else {
send_to_writer(
ctx,
DaemonMessage::Session {
session_id: None,
event: SessionEvent::Failed {
request_id,
error: "no session attached".to_string(),
},
},
);
}
}
ClientMessage::Cancel { request_id } => {
debug!("client {}: Cancel id={}", ctx.client_id, request_id);
if let Some(session_id) = *ctx.attached_session_id {
let _ = ctx.daemon_tx.send(DaemonCommand::CancelRequest {
session_id,
request_id,
});
}
}
ClientMessage::Undo => {
debug!("client {}: Undo", ctx.client_id);
if let Some(tx) = ctx.attached_session_tx {
let _ = tx.send(SessionCommand::Undo);
}
}
ClientMessage::Redo => {
debug!("client {}: Redo", ctx.client_id);
if let Some(tx) = ctx.attached_session_tx {
let _ = tx.send(SessionCommand::Redo);
}
}
ClientMessage::ContinueGeneration { request_id } => {
debug!(
"client {}: ContinueGeneration id={}",
ctx.client_id, request_id
);
if let Some(tx) = ctx.attached_session_tx {
let _ = tx.send(SessionCommand::RunInput {
request_id,
input: b"Continue.".to_vec(),
});
} else {
send_to_writer(
ctx,
DaemonMessage::Session {
session_id: None,
event: SessionEvent::Failed {
request_id,
error: "no session attached".to_string(),
},
},
);
}
}
ClientMessage::Ping => {
debug!("client {}: Ping", ctx.client_id);
send_to_writer(ctx, DaemonMessage::Pong);
}
ClientMessage::SetModel { model } => {
info!(
"client {}: SetModel model={} attached={}",
ctx.client_id,
model,
ctx.attached_session_tx.is_some()
);
if let Some(tx) = ctx.attached_session_tx {
let _ = tx.send(SessionCommand::SetModel { model });
} else {
send_to_writer(
ctx,
DaemonMessage::Session {
session_id: None,
event: SessionEvent::ModelSelectionFailed {
model,
error: "no session attached".to_string(),
},
},
);
}
}
ClientMessage::SetReasoningEffort { effort } => {
info!(
"client {}: SetReasoningEffort effort={} attached={}",
ctx.client_id,
effort,
ctx.attached_session_tx.is_some()
);
if let Some(tx) = ctx.attached_session_tx {
let _ = tx.send(SessionCommand::SetReasoningEffort { effort });
} else {
send_to_writer(
ctx,
DaemonMessage::Session {
session_id: None,
event: SessionEvent::ReasoningEffortSetFailed {
effort,
error: "no session attached".to_string(),
},
},
);
}
}
ClientMessage::GetReasoningEffort => {
if let Some(tx) = ctx.attached_session_tx {
let (reply, rx) = mpsc::channel();
let _ = tx.send(SessionCommand::GetReasoningEffort { reply });
if let Ok(effort) = rx.recv() {
send_to_writer(
ctx,
DaemonMessage::Session {
session_id: *ctx.attached_session_id,
event: SessionEvent::ReasoningEffortSet { effort },
},
);
}
} else {
send_to_writer(
ctx,
DaemonMessage::Session {
session_id: None,
event: SessionEvent::ReasoningEffortSet {
effort: "off".to_string(),
},
},
);
}
}
ClientMessage::Unlock { private_key } => {
info!("client {}: Unlock", ctx.client_id);
handle_unlock_sync(ctx, private_key);
}
ClientMessage::BindKeystore { key } => {
info!("client {}: BindKeystore", ctx.client_id);
handle_bind_keystore_sync(ctx, key);
}
ClientMessage::Lock => {
info!("client {}: Lock", ctx.client_id);
handle_lock_sync(ctx);
}
ClientMessage::AddCredential {
service,
encrypted_payload,
unlock_key,
} => {
info!(
"client {}: AddCredential service={}",
ctx.client_id, service
);
handle_add_credential_sync(ctx, service, encrypted_payload, unlock_key);
}
ClientMessage::RemoveCredential { service } => {
info!(
"client {}: RemoveCredential service={}",
ctx.client_id, service
);
handle_remove_credential_sync(ctx, service);
}
ClientMessage::AclAdd { pubkey } => {
info!("client {}: AclAdd (local={})", ctx.client_id, ctx.is_unix);
handle_acl_add_sync(ctx, pubkey);
}
ClientMessage::ListModels => {
debug!("client {}: ListModels", ctx.client_id);
handle_list_models_sync(ctx, *ctx.attached_session_id);
}
ClientMessage::RefreshModels { force } => {
debug!("client {}: RefreshModels force={}", ctx.client_id, force);
handle_refresh_models_sync(ctx, force);
}
ClientMessage::DeleteSession { session_id } => {
info!("client {}: DeleteSession id={}", ctx.client_id, session_id);
handle_delete_session_sync(ctx, session_id);
}
ClientMessage::GetCredential { service } => {
handle_get_credential_sync(ctx, service);
}
ClientMessage::AddAccount {
name,
provider,
base_url,
streaming,
retry_max_attempts,
connect_timeout_secs,
request_timeout_secs,
total_timeout_secs,
} => {
let result = request_daemon(ctx.daemon_tx, |reply| DaemonCommand::AddAccountCmd {
name: name.clone(),
provider,
base_url,
streaming,
retry_max_attempts,
connect_timeout_secs,
request_timeout_secs,
total_timeout_secs,
reply,
});
match result {
Ok(Ok(())) => {
send_to_writer(ctx, DaemonMessage::AccountAdded { name });
}
Ok(Err(e)) => {
send_to_writer(ctx, DaemonMessage::AccountAddFailed { name, error: e });
}
Err(_) => warn!("daemon disconnected while handling add account"),
}
}
ClientMessage::RemoveAccount { name } => {
let result = request_daemon(ctx.daemon_tx, |reply| DaemonCommand::RemoveAccountCmd {
name: name.clone(),
reply,
});
match result {
Ok(Ok(())) => {
send_to_writer(ctx, DaemonMessage::AccountRemoved { name });
}
Ok(Err(e)) => {
send_to_writer(ctx, DaemonMessage::AccountRemoveFailed { name, error: e });
}
Err(_) => warn!("daemon disconnected while handling remove account"),
}
}
ClientMessage::ListAccounts => {
let result = request_daemon(ctx.daemon_tx, |reply| DaemonCommand::ListAccountsCmd {
reply,
});
match result {
Ok(Ok(accounts)) => {
send_to_writer(ctx, DaemonMessage::Accounts { accounts });
}
Ok(Err(e)) => {
send_to_writer(ctx, DaemonMessage::AccountListFailed { error: e });
}
Err(_) => warn!("daemon disconnected while handling list accounts"),
}
}
ClientMessage::SetSessionAccount { name } => {
handle_client_set_session_account(name, ctx);
}
ClientMessage::SubscribeAllActivity => {
let _ = ctx
.daemon_tx
.send(DaemonCommand::RegisterActivitySubscriber {
client_id: ctx.client_id,
writer: ctx.writer.clone(),
});
}
ClientMessage::UnsubscribeAllActivity => {
let _ = ctx
.daemon_tx
.send(DaemonCommand::UnregisterActivitySubscriber {
client_id: ctx.client_id,
});
}
_ => {
warn!(
"unhandled client message: {:?}",
std::mem::discriminant(&msg)
);
}
}
Ok(())
}
pub(crate) struct ClientConn {
daemon_tx: mpsc::Sender<DaemonCommand>,
writer: crate::broadcast::SubscriberSink,
global_lag: Arc<AtomicUsize>,
client_id: u64,
is_unix: bool,
attached_session_id: Option<u64>,
attached_session_tx: Option<mpsc::Sender<SessionCommand>>,
writer_handle: std::thread::JoinHandle<()>,
}
impl ClientConn {
fn new<W: ConnectionWriter + Send + 'static>(
daemon_tx: mpsc::Sender<DaemonCommand>,
writer: crate::broadcast::SubscriberSink,
writer_buf: W,
writer_rx: crossbeam_channel::Receiver<DaemonMessage>,
global_lag: Arc<AtomicUsize>,
client_id: u64,
is_unix: bool,
) -> Self {
let bytes = Arc::clone(&writer.bytes_in_flight);
let global = Arc::clone(&global_lag);
let writer_handle =
std::thread::spawn(move || writer_thread(writer_buf, writer_rx, bytes, global));
Self {
daemon_tx,
writer,
global_lag,
client_id,
is_unix,
attached_session_id: None,
attached_session_tx: None,
writer_handle,
}
}
pub(crate) fn dispatch(&mut self, msg: ClientMessage) -> io::Result<()> {
let mut ctx = ClientCtx {
writer: &self.writer,
global_lag: &self.global_lag,
daemon_tx: &self.daemon_tx,
attached_session_id: &mut self.attached_session_id,
attached_session_tx: &mut self.attached_session_tx,
client_id: self.client_id,
is_unix: self.is_unix,
};
dispatch_client_message(msg, &mut ctx)
}
pub(crate) fn finish(self) {
cleanup_client(
self.attached_session_tx,
self.client_id,
&self.daemon_tx,
self.writer,
self.writer_handle,
);
}
}
pub(crate) fn client_thread(
stream: UnixStream,
daemon_tx: mpsc::Sender<DaemonCommand>,
client_id: u64,
writer: crate::broadcast::SubscriberSink,
writer_rx: crossbeam_channel::Receiver<DaemonMessage>,
global_lag: Arc<AtomicUsize>,
) -> io::Result<()> {
stream.set_write_timeout(Some(WRITER_WRITE_TIMEOUT))?;
let reader = BufReader::new(stream.try_clone()?);
let writer_buf = BufWriter::new(stream);
let mut conn = ClientConn::new(
daemon_tx, writer, writer_buf, writer_rx, global_lag, client_id, true,
);
info!("client connected: id={}", client_id);
crate::metrics::record_client_connected();
let mut reader = reader;
loop {
match read_message::<_, ClientMessage>(&mut reader) {
Ok(msg) => {
if let Err(e) = conn.dispatch(msg) {
debug!("daemon disconnected: {e}");
break;
}
}
Err(ProtoError::Io(e))
if matches!(
e.kind(),
io::ErrorKind::UnexpectedEof | io::ErrorKind::ConnectionReset
) =>
{
debug!("client disconnected");
break;
}
Err(e) => {
error!(error = %e, "failed to read client message");
break;
}
}
}
conn.finish();
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn tcp_handshake_and_client_thread(
mut tcp: TcpStream,
transport_sk: [u8; 32],
acl: Arc<crate::server::acl::SharedAcl>,
daemon_tx: mpsc::Sender<DaemonCommand>,
client_id: u64,
writer: crate::broadcast::SubscriberSink,
writer_rx: crossbeam_channel::Receiver<DaemonMessage>,
global_lag: Arc<AtomicUsize>,
) -> io::Result<()> {
let preamble = match choreo_transport::handshake::read_handshake_preamble(&mut tcp) {
Ok(p) => p,
Err(e) => {
warn!(
error = %e,
"TCP client never sent a valid handshake-mode preamble; closing"
);
let _ = daemon_tx.send(DaemonCommand::ClientDisconnected { client_id });
return Ok(());
}
};
let handshake_result = match preamble {
choreo_transport::handshake::PREAMBLE_IK => {
debug!("TCP client selected Noise IK handshake");
choreo_transport::handshake::handshake_responder(tcp, &transport_sk, |pk| {
acl.contains(pk)
})
}
choreo_transport::handshake::PREAMBLE_XX => {
debug!("TCP client selected Noise XX (first-contact) handshake");
choreo_transport::handshake::handshake_responder_xx(tcp, &transport_sk, |pk| {
acl.contains(pk)
})
}
other => {
warn!(
preamble = other,
"unknown handshake-mode preamble byte; closing connection"
);
let _ = daemon_tx.send(DaemonCommand::ClientDisconnected { client_id });
return Ok(()); }
};
let noise = match handshake_result {
Ok(noise) => noise,
Err(e) => {
error!(error = %e, "Noise handshake rejected");
let _ = daemon_tx.send(DaemonCommand::ClientDisconnected { client_id });
return Ok(());
}
};
tcp_client_thread(noise, daemon_tx, client_id, writer, writer_rx, global_lag)
}
pub(crate) fn tcp_client_thread(
noise: choreo_transport::noise::NoiseStream,
daemon_tx: mpsc::Sender<DaemonCommand>,
client_id: u64,
writer: crate::broadcast::SubscriberSink,
writer_rx: crossbeam_channel::Receiver<DaemonMessage>,
global_lag: Arc<AtomicUsize>,
) -> io::Result<()> {
noise
.get_ref()
.set_write_timeout(Some(WRITER_WRITE_TIMEOUT))?;
let writer_buf = noise.try_clone()?;
let mut conn = ClientConn::new(
daemon_tx, writer, writer_buf, writer_rx, global_lag, client_id, false,
);
info!("TCP client connected: id={}", client_id);
crate::metrics::record_client_connected();
let mut reader = noise;
loop {
match reader.recv_client_message() {
Ok(msg) => {
if let Err(e) = conn.dispatch(msg) {
debug!("daemon disconnected: {e}");
break;
}
}
Err(choreo_transport::error::TransportError::ConnectionClosed) => {
info!("TCP client closed connection");
break;
}
Err(e) => {
error!(error = %e, "failed to read client message");
break;
}
}
}
conn.finish();
Ok(())
}
pub(crate) struct EmbeddedConnArgs {
pub client_rx: crossbeam_channel::Receiver<ClientMessage>,
pub out_tx: crossbeam_channel::Sender<DaemonMessage>,
pub daemon_tx: mpsc::Sender<DaemonCommand>,
pub client_id: u64,
pub writer: crate::broadcast::SubscriberSink,
pub writer_rx: crossbeam_channel::Receiver<DaemonMessage>,
pub global_lag: Arc<AtomicUsize>,
}
pub(crate) fn embedded_client_thread(args: EmbeddedConnArgs) -> io::Result<()> {
let EmbeddedConnArgs {
client_rx,
out_tx,
daemon_tx,
client_id,
writer,
writer_rx,
global_lag,
} = args;
let mut conn = ClientConn::new(
daemon_tx,
writer,
ChannelConnectionWriter::new(out_tx),
writer_rx,
global_lag,
client_id,
true,
);
info!("embedded client connected: id={}", client_id);
crate::metrics::record_client_connected();
for msg in client_rx {
if let Err(e) = conn.dispatch(msg) {
debug!("daemon disconnected: {e}");
break;
}
}
info!("embedded client disconnected: id={}", client_id);
conn.finish();
Ok(())
}
fn switch_attached_session(
new_session_id: u64,
session_tx: mpsc::Sender<SessionCommand>,
ctx: &mut ClientCtx,
) {
if Some(new_session_id) != *ctx.attached_session_id
&& let Some(old_tx) = ctx.attached_session_tx.as_ref()
{
let _ = old_tx.send(SessionCommand::Detach {
client_id: ctx.client_id,
});
}
let _ = session_tx.send(SessionCommand::Attach {
client_id: ctx.client_id,
tx: ctx.writer.clone(),
});
*ctx.attached_session_tx = Some(session_tx);
*ctx.attached_session_id = Some(new_session_id);
}
#[expect(clippy::too_many_arguments)]
fn handle_client_create_session(
title: Option<String>,
parent_session_id: Option<u64>,
working_dir: Option<String>,
context_config: Option<ContextConfig>,
account_name: Option<String>,
selected_model: Option<String>,
reasoning_effort: Option<String>,
ctx: &mut ClientCtx,
) -> bool {
info!("client {}: CreateSession", ctx.client_id);
let cwd_str = working_dir.clone();
let (reply, rx) = mpsc::channel();
let _ = ctx.daemon_tx.send(DaemonCommand::CreateSession {
title: title.clone(),
parent_session_id,
working_dir: working_dir.map(std::path::PathBuf::from),
reasoning_effort: reasoning_effort.clone(),
selected_model: selected_model.clone(),
context_config,
account_name: account_name.clone(),
active_tool_groups: Vec::new(),
reply,
});
match rx.recv() {
Ok(Ok((sid, _session_tx))) => {
send_to_writer(
ctx,
DaemonMessage::Session {
session_id: Some(sid),
event: SessionEvent::SessionCreated {
title,
parent_session_id,
working_dir: cwd_str,
account_name,
selected_model,
reasoning_effort,
},
},
);
}
Ok(Err(e)) => {
send_to_writer(
ctx,
DaemonMessage::Session {
session_id: None,
event: SessionEvent::SessionFailed {
operation: "create_session".into(),
error: e.to_string(),
},
},
);
}
Err(_) => return false,
}
true
}
fn handle_client_attach_session(session_id: u64, ctx: &mut ClientCtx) -> bool {
info!("client {}: AttachSession id={}", ctx.client_id, session_id);
let (reply, rx) = mpsc::channel();
let _ = ctx
.daemon_tx
.send(DaemonCommand::AttachSession { session_id, reply });
match rx.recv() {
Ok(Ok(session_tx)) => {
send_to_writer(
ctx,
DaemonMessage::Session {
session_id: Some(session_id),
event: SessionEvent::SessionAttached,
},
);
switch_attached_session(session_id, session_tx, ctx);
}
Ok(Err(e)) => {
send_to_writer(
ctx,
DaemonMessage::Session {
session_id: None,
event: SessionEvent::SessionFailed {
operation: "attach_session".into(),
error: e.to_string(),
},
},
);
}
Err(_) => return false,
}
true
}
fn handle_client_set_session_account(name: String, ctx: &mut ClientCtx) {
if let Some(tx) = ctx.attached_session_tx.as_ref() {
let (reply, rx) = mpsc::channel();
let _ = ctx.daemon_tx.send(DaemonCommand::AccountExists {
name: name.clone(),
reply,
});
match rx.recv() {
Ok(true) => {
let _ = tx.send(SessionCommand::SetAccount { name });
}
_ => {
send_to_writer(
ctx,
DaemonMessage::Session {
session_id: *ctx.attached_session_id,
event: SessionEvent::SessionFailed {
operation: "set_account".into(),
error: format!("account '{name}' not found"),
},
},
);
}
}
} else {
send_to_writer(
ctx,
DaemonMessage::Session {
session_id: None,
event: SessionEvent::SessionFailed {
operation: "set_account".into(),
error: "no session attached".to_string(),
},
},
);
}
}
fn request_daemon<R>(
daemon_tx: &mpsc::Sender<DaemonCommand>,
make_cmd: impl FnOnce(mpsc::Sender<R>) -> DaemonCommand,
) -> Result<R, mpsc::RecvError> {
let (reply, rx) = mpsc::channel();
if daemon_tx.send(make_cmd(reply)).is_err() {
return Err(mpsc::RecvError);
}
rx.recv()
}
fn handle_unlock_sync(ctx: &mut ClientCtx, private_key: Vec<u8>) {
let result = request_daemon(ctx.daemon_tx, |reply| DaemonCommand::Unlock {
private_key,
client_writer: Some(ctx.writer.clone()),
reply,
});
if result.is_err() {
warn!("daemon disconnected while handling unlock");
}
}
fn handle_bind_keystore_sync(ctx: &mut ClientCtx, key: Vec<u8>) {
let result = request_daemon(ctx.daemon_tx, |reply| DaemonCommand::BindKeystore {
key,
client_writer: Some(ctx.writer.clone()),
reply,
});
if result.is_err() {
warn!("daemon disconnected while handling bind keystore");
}
}
fn handle_lock_sync(ctx: &mut ClientCtx) {
let result = request_daemon(ctx.daemon_tx, |reply| DaemonCommand::Lock { reply });
match result {
Ok(Ok(())) => {
send_to_writer(ctx, DaemonMessage::Locked);
}
Ok(Err(e)) => {
send_to_writer(ctx, DaemonMessage::LockedError { error: e });
}
Err(_) => warn!("daemon disconnected while handling lock"),
}
}
fn handle_list_models_sync(ctx: &mut ClientCtx, attached_session_id: Option<u64>) {
let result = request_daemon(ctx.daemon_tx, |reply| DaemonCommand::ListModels {
session_id: attached_session_id,
reply,
});
match result {
Ok(Ok((models, selected_model))) => {
send_to_writer(
ctx,
DaemonMessage::Models {
models,
selected_model,
},
);
}
Ok(Err(e)) => {
send_to_writer(ctx, DaemonMessage::ModelsFailed { error: e });
}
Err(_) => warn!("daemon disconnected while handling list models"),
}
}
fn handle_refresh_models_sync(ctx: &mut ClientCtx, force: bool) {
let result = request_daemon(ctx.daemon_tx, |reply| DaemonCommand::RefreshModels {
force,
reply,
});
match result {
Ok(Ok(report)) => {
send_to_writer(
ctx,
DaemonMessage::ModelsRefreshed {
providers: report.providers,
models: report.models,
status: report.status,
},
);
}
Ok(Err(e)) => {
send_to_writer(ctx, DaemonMessage::ModelsRefreshFailed { error: e });
}
Err(_) => warn!("daemon disconnected while handling refresh models"),
}
}
fn handle_get_credential_sync(ctx: &mut ClientCtx, service: String) {
let result = request_daemon(ctx.daemon_tx, |reply| DaemonCommand::GetCredential {
service: service.clone(),
reply,
});
match result {
Ok(Some(key)) => {
send_to_writer(
ctx,
DaemonMessage::Credential {
service,
key: Some(key),
},
);
}
Ok(None) => {
send_to_writer(ctx, DaemonMessage::Credential { service, key: None });
}
Err(_) => warn!("daemon disconnected while handling get credential"),
}
}
fn handle_delete_session_sync(ctx: &mut ClientCtx, session_id: u64) {
let result = request_daemon(ctx.daemon_tx, |reply| DaemonCommand::DeleteSession {
session_id,
reply,
});
match result {
Ok(Ok(())) => {
}
Ok(Err(e)) => {
send_to_writer(
ctx,
DaemonMessage::Session {
session_id: Some(session_id),
event: SessionEvent::SessionDeleteFailed {
error: e.to_string(),
},
},
);
}
Err(_) => warn!("daemon disconnected while handling delete session"),
}
}
fn handle_add_credential_sync(
ctx: &mut ClientCtx,
service: String,
encrypted_payload: Vec<u8>,
unlock_key: Vec<u8>,
) {
let result = request_daemon(ctx.daemon_tx, |reply| DaemonCommand::SaveCredential {
service: service.clone(),
encrypted_blob: encrypted_payload,
unlock_key,
client_writer: Some(ctx.writer.clone()),
reply,
});
if result.is_err() {
warn!("daemon disconnected while handling add credential");
}
}
fn handle_acl_add_sync(ctx: &mut ClientCtx, pubkey: String) {
if !ctx.is_unix {
warn!(
"client {}: AclAdd refused: remote connections cannot change the ACL",
ctx.client_id
);
send_to_writer(
ctx,
DaemonMessage::AclAddResult {
ok: false,
message: "ACL changes are only permitted from local connections".to_string(),
},
);
return;
}
let result = request_daemon(ctx.daemon_tx, |reply| DaemonCommand::AclAddCmd {
pubkey: pubkey.clone(),
reply,
});
match result {
Ok(Ok(count)) => {
send_to_writer(
ctx,
DaemonMessage::AclAddResult {
ok: true,
message: format!("client key authorized ({count} client(s) now trusted)"),
},
);
}
Ok(Err(e)) => {
send_to_writer(
ctx,
DaemonMessage::AclAddResult {
ok: false,
message: e,
},
);
}
Err(_) => warn!("daemon disconnected while handling acl add"),
}
}
fn handle_remove_credential_sync(ctx: &mut ClientCtx, service: String) {
let result = request_daemon(ctx.daemon_tx, |reply| DaemonCommand::RemoveCredentialCmd {
service: service.clone(),
reply,
});
match result {
Ok(Ok(())) => {
send_to_writer(ctx, DaemonMessage::CredentialRemoved { service });
}
Ok(Err(e)) => {
send_to_writer(
ctx,
DaemonMessage::CredentialRemoveFailed { service, error: e },
);
}
Err(_) => warn!("daemon disconnected while handling remove credential"),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::broadcast::test_sink;
struct MockConnectionWriter {
sent: mpsc::Sender<DaemonMessage>,
shutdown_tx: mpsc::Sender<()>,
fail_on: Option<usize>,
calls: usize,
}
impl ConnectionWriter for MockConnectionWriter {
fn send_message(&mut self, msg: &DaemonMessage) -> Result<(), String> {
self.calls += 1;
if self.fail_on == Some(self.calls) {
return Err("mock write failure".to_string());
}
let _ = self.sent.send(msg.clone());
Ok(())
}
fn shutdown(&mut self) {
let _ = self.shutdown_tx.send(());
}
}
fn mock_writer(
fail_on: Option<usize>,
) -> (
MockConnectionWriter,
mpsc::Receiver<DaemonMessage>,
mpsc::Receiver<()>,
) {
let (sent_tx, sent_rx) = mpsc::channel();
let (shutdown_tx, shutdown_rx) = mpsc::channel();
(
MockConnectionWriter {
sent: sent_tx,
shutdown_tx,
fail_on,
calls: 0,
},
sent_rx,
shutdown_rx,
)
}
#[test]
fn writer_thread_flushes_shutting_down_then_shuts_down_and_stops() {
let (tx, rx) = crossbeam_channel::unbounded::<DaemonMessage>();
let (writer, sent_rx, shutdown_rx) = mock_writer(None);
let bytes = Arc::new(AtomicUsize::new(0));
let global = Arc::new(AtomicUsize::new(0));
let handle = std::thread::spawn({
let bytes = Arc::clone(&bytes);
let global = Arc::clone(&global);
move || writer_thread(writer, rx, bytes, global)
});
tx.send(DaemonMessage::Pong).unwrap();
tx.send(DaemonMessage::ShuttingDown).unwrap();
tx.send(DaemonMessage::Pong).unwrap();
handle.join().expect("writer thread panicked");
let written: Vec<_> = sent_rx.try_iter().collect();
assert_eq!(
written,
vec![DaemonMessage::Pong, DaemonMessage::ShuttingDown],
"ShuttingDown must be flushed in order, then draining must stop"
);
assert!(
shutdown_rx.try_recv().is_ok(),
"the writer thread must close the socket itself after ShuttingDown"
);
}
#[test]
fn writer_thread_flushes_evicted_then_shuts_down_and_stops() {
let (tx, rx) = crossbeam_channel::unbounded::<DaemonMessage>();
let (writer, sent_rx, shutdown_rx) = mock_writer(None);
let bytes = Arc::new(AtomicUsize::new(0));
let global = Arc::new(AtomicUsize::new(0));
let handle = std::thread::spawn({
let bytes = Arc::clone(&bytes);
let global = Arc::clone(&global);
move || writer_thread(writer, rx, bytes, global)
});
tx.send(DaemonMessage::Pong).unwrap();
tx.send(DaemonMessage::Evicted).unwrap();
tx.send(DaemonMessage::Pong).unwrap();
handle.join().expect("writer thread panicked");
let written: Vec<_> = sent_rx.try_iter().collect();
assert_eq!(
written,
vec![DaemonMessage::Pong, DaemonMessage::Evicted],
"Evicted must be flushed in order, then draining must stop"
);
assert!(
shutdown_rx.try_recv().is_ok(),
"the writer thread must close the socket itself after Evicted"
);
}
#[test]
fn writer_thread_stops_and_shuts_down_on_send_error() {
let (tx, rx) = crossbeam_channel::unbounded::<DaemonMessage>();
let (writer, sent_rx, shutdown_rx) = mock_writer(Some(2));
let bytes = Arc::new(AtomicUsize::new(0));
let global = Arc::new(AtomicUsize::new(0));
let handle = std::thread::spawn({
let bytes = Arc::clone(&bytes);
let global = Arc::clone(&global);
move || writer_thread(writer, rx, bytes, global)
});
tx.send(DaemonMessage::Pong).unwrap();
tx.send(DaemonMessage::Pong).unwrap();
tx.send(DaemonMessage::ShuttingDown).unwrap();
drop(tx);
handle.join().expect("writer thread panicked");
let written: Vec<_> = sent_rx.try_iter().collect();
assert_eq!(written.len(), 1, "writer must stop at the failing message");
assert!(
shutdown_rx.try_recv().is_ok(),
"writer must shut the socket down on a send error so the reader is unblocked"
);
}
#[test]
fn writer_thread_exits_cleanly_on_disconnect() {
let (tx, rx) = crossbeam_channel::unbounded::<DaemonMessage>();
let (writer, sent_rx, shutdown_rx) = mock_writer(None);
let bytes = Arc::new(AtomicUsize::new(0));
let global = Arc::new(AtomicUsize::new(0));
let handle = std::thread::spawn({
let bytes = Arc::clone(&bytes);
let global = Arc::clone(&global);
move || writer_thread(writer, rx, bytes, global)
});
tx.send(DaemonMessage::Pong).unwrap();
drop(tx);
handle.join().expect("writer thread panicked");
let written: Vec<_> = sent_rx.try_iter().collect();
assert_eq!(written, vec![DaemonMessage::Pong]);
assert!(shutdown_rx.try_recv().is_err());
}
#[test]
fn writer_thread_decrements_byte_counters_per_message() {
let (tx, rx) = crossbeam_channel::unbounded::<DaemonMessage>();
let (writer, _sent_rx, _shutdown_rx) = mock_writer(None);
let bytes = Arc::new(AtomicUsize::new(0));
let global = Arc::new(AtomicUsize::new(0));
let handle = std::thread::spawn({
let bytes = Arc::clone(&bytes);
let global = Arc::clone(&global);
move || writer_thread(writer, rx, bytes, global)
});
let m1 = DaemonMessage::Session {
session_id: Some(1),
event: SessionEvent::Failed {
request_id: 1,
error: "a".repeat(100),
},
};
let m2 = DaemonMessage::Session {
session_id: Some(2),
event: SessionEvent::Failed {
request_id: 2,
error: "b".repeat(50),
},
};
let s1 = m1.approx_wire_size();
let s2 = m2.approx_wire_size();
bytes.fetch_add(s1 + s2, Ordering::Relaxed);
global.fetch_add(s1 + s2, Ordering::Relaxed);
tx.send(m1).unwrap();
tx.send(m2).unwrap();
drop(tx);
handle.join().expect("writer thread panicked");
assert_eq!(
bytes.load(Ordering::Relaxed),
0,
"every dequeued message must decrement the per-client counter"
);
assert_eq!(
global.load(Ordering::Relaxed),
0,
"every dequeued message must decrement the daemon-wide counter"
);
}
#[test]
fn writer_thread_drains_and_decrements_abandoned_backlog() {
let (tx, rx) = crossbeam_channel::unbounded::<DaemonMessage>();
let (writer, sent_rx, _shutdown_rx) = mock_writer(None);
let bytes = Arc::new(AtomicUsize::new(0));
let global = Arc::new(AtomicUsize::new(0));
let handle = std::thread::spawn({
let bytes = Arc::clone(&bytes);
let global = Arc::clone(&global);
move || writer_thread(writer, rx, bytes, global)
});
let m1 = DaemonMessage::Session {
session_id: Some(1),
event: SessionEvent::Failed {
request_id: 1,
error: "a".repeat(100),
},
};
let m2 = DaemonMessage::Session {
session_id: Some(2),
event: SessionEvent::Failed {
request_id: 2,
error: "b".repeat(50),
},
};
let m3 = DaemonMessage::Session {
session_id: Some(3),
event: SessionEvent::Failed {
request_id: 3,
error: "c".repeat(25),
},
};
let s1 = m1.approx_wire_size();
let s2 = m2.approx_wire_size();
let s3 = m3.approx_wire_size();
let evicted_size = DaemonMessage::Evicted.approx_wire_size();
let total = s1 + s2 + s3 + evicted_size;
bytes.fetch_add(total, Ordering::Relaxed);
global.fetch_add(total, Ordering::Relaxed);
tx.send(m1.clone()).unwrap();
tx.send(DaemonMessage::Evicted).unwrap();
tx.send(m2).unwrap();
tx.send(m3).unwrap();
handle.join().expect("writer thread panicked");
let written: Vec<_> = sent_rx.try_iter().collect();
assert_eq!(
written,
vec![m1.clone(), DaemonMessage::Evicted],
"messages behind the advisory must never be written"
);
assert_eq!(
bytes.load(Ordering::Relaxed),
0,
"abandoned backlog must be decremented from the per-client counter"
);
assert_eq!(
global.load(Ordering::Relaxed),
0,
"abandoned backlog must be decremented from the daemon-wide counter"
);
}
#[test]
fn handle_acl_add_sync_refuses_remote_clients_without_dialing_daemon() {
let (daemon_tx, daemon_rx) = mpsc::channel();
let (sink, writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let mut none_id = None;
let mut none_tx = None;
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut none_id,
attached_session_tx: &mut none_tx,
client_id: 7,
is_unix: false, };
handle_acl_add_sync(
&mut ctx,
"AQIDBAUGBwgJCgsMDQ4PEBESExQVFhcYGRobHB0eHyA=".to_string(),
);
let msg = writer_rx.recv().unwrap();
match msg {
DaemonMessage::AclAddResult { ok: false, message } => {
assert!(
message.contains("local connections"),
"the refusal must explain the trust boundary, got: {message}"
);
}
other => panic!("expected AclAddResult refusal, got {other:?}"),
}
assert!(
daemon_rx.try_recv().is_err(),
"a remote AclAdd must never reach the daemon command loop"
);
}
#[test]
fn handle_unlock_sync_ok() {
let (daemon_tx, daemon_rx) = mpsc::channel();
let (sink, writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let mut none_id = None;
let mut none_tx = None;
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut none_id,
attached_session_tx: &mut none_tx,
client_id: 0,
is_unix: true,
};
let lag = Arc::clone(&global_lag);
std::thread::spawn(move || {
if let Ok(DaemonCommand::Unlock {
client_writer,
reply,
..
}) = daemon_rx.recv()
{
if let Some(w) = &client_writer {
w.send_accounted(&DaemonMessage::Unlocked, &lag);
}
let _ = reply.send(());
}
});
handle_unlock_sync(&mut ctx, vec![0u8; 32]);
let msg = writer_rx.recv().unwrap();
assert!(matches!(msg, DaemonMessage::Unlocked));
}
#[test]
fn handle_unlock_sync_err() {
let (daemon_tx, daemon_rx) = mpsc::channel();
let (sink, writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let mut none_id = None;
let mut none_tx = None;
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut none_id,
attached_session_tx: &mut none_tx,
client_id: 0,
is_unix: true,
};
let lag = Arc::clone(&global_lag);
std::thread::spawn(move || {
if let Ok(DaemonCommand::Unlock {
client_writer,
reply,
..
}) = daemon_rx.recv()
{
if let Some(w) = &client_writer {
w.send_accounted(
&DaemonMessage::LockedError {
error: "wrong password".to_string(),
},
&lag,
);
}
let _ = reply.send(());
}
});
handle_unlock_sync(&mut ctx, vec![0u8; 32]);
let msg = writer_rx.recv().unwrap();
assert!(matches!(msg, DaemonMessage::LockedError { .. }));
if let DaemonMessage::LockedError { error } = &msg {
assert_eq!(error, "wrong password");
}
}
#[test]
fn handle_unlock_sync_disconnected() {
let (daemon_tx, daemon_rx) = mpsc::channel::<DaemonCommand>();
let (sink, writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let mut none_id = None;
let mut none_tx = None;
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut none_id,
attached_session_tx: &mut none_tx,
client_id: 0,
is_unix: true,
};
drop(daemon_rx);
handle_unlock_sync(&mut ctx, vec![0u8; 32]);
assert!(writer_rx.try_recv().is_err());
}
#[test]
fn handle_lock_sync_ok_replies_locked() {
let (daemon_tx, daemon_rx) = mpsc::channel();
let (sink, writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let mut none_id = None;
let mut none_tx = None;
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut none_id,
attached_session_tx: &mut none_tx,
client_id: 0,
is_unix: true,
};
std::thread::spawn(move || {
if let Ok(DaemonCommand::Lock { reply }) = daemon_rx.recv() {
let _ = reply.send(Ok(()));
}
});
handle_lock_sync(&mut ctx);
let msg = writer_rx.recv().unwrap();
assert!(matches!(msg, DaemonMessage::Locked));
}
#[test]
fn handle_lock_sync_err_replies_locked_error() {
let (daemon_tx, daemon_rx) = mpsc::channel();
let (sink, writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let mut none_id = None;
let mut none_tx = None;
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut none_id,
attached_session_tx: &mut none_tx,
client_id: 0,
is_unix: true,
};
std::thread::spawn(move || {
if let Ok(DaemonCommand::Lock { reply }) = daemon_rx.recv() {
let _ = reply.send(Err("cannot lock".into()));
}
});
handle_lock_sync(&mut ctx);
let msg = writer_rx.recv().unwrap();
assert!(matches!(msg, DaemonMessage::LockedError { .. }));
}
#[test]
fn handle_list_models_sync_ok() {
let (daemon_tx, daemon_rx) = mpsc::channel();
let (sink, writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let mut none_id = None;
let mut none_tx = None;
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut none_id,
attached_session_tx: &mut none_tx,
client_id: 0,
is_unix: true,
};
std::thread::spawn(move || {
if let Ok(DaemonCommand::ListModels { reply, .. }) = daemon_rx.recv() {
let _ = reply.send(Ok((
vec!["gpt-4".into(), "gpt-3.5".into()],
Some("gpt-4".into()),
)));
}
});
handle_list_models_sync(&mut ctx, None);
let msg = writer_rx.recv().unwrap();
assert!(matches!(msg, DaemonMessage::Models { .. }));
}
#[test]
fn handle_refresh_models_sync_ok() {
let (daemon_tx, daemon_rx) = mpsc::channel();
let (sink, writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let mut none_id = None;
let mut none_tx = None;
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut none_id,
attached_session_tx: &mut none_tx,
client_id: 0,
is_unix: true,
};
std::thread::spawn(move || {
if let Ok(DaemonCommand::RefreshModels { force, reply }) = daemon_rx.recv() {
assert!(force);
let _ = reply.send(Ok(crate::catalog::RefreshReport {
providers: 208,
models: 1234,
status: choreo_proto::RefreshStatus::Updated,
}));
}
});
handle_refresh_models_sync(&mut ctx, true);
let msg = writer_rx.recv().unwrap();
assert!(matches!(
&msg,
DaemonMessage::ModelsRefreshed {
providers: 208,
models: 1234,
status: choreo_proto::RefreshStatus::Updated,
}
));
}
#[test]
fn handle_refresh_models_sync_err() {
let (daemon_tx, daemon_rx) = mpsc::channel();
let (sink, writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let mut none_id = None;
let mut none_tx = None;
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut none_id,
attached_session_tx: &mut none_tx,
client_id: 0,
is_unix: true,
};
std::thread::spawn(move || {
if let Ok(DaemonCommand::RefreshModels { reply, .. }) = daemon_rx.recv() {
let _ = reply.send(Err("daemon is locked".into()));
}
});
handle_refresh_models_sync(&mut ctx, false);
let msg = writer_rx.recv().unwrap();
assert!(
matches!(&msg, DaemonMessage::ModelsRefreshFailed { error } if error == "daemon is locked")
);
}
#[test]
fn handle_list_models_sync_err() {
let (daemon_tx, daemon_rx) = mpsc::channel();
let (sink, writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let mut none_id = None;
let mut none_tx = None;
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut none_id,
attached_session_tx: &mut none_tx,
client_id: 0,
is_unix: true,
};
std::thread::spawn(move || {
if let Ok(DaemonCommand::ListModels { reply, .. }) = daemon_rx.recv() {
let _ = reply.send(Err("daemon is locked".into()));
}
});
handle_list_models_sync(&mut ctx, None);
let msg = writer_rx.recv().unwrap();
assert!(matches!(msg, DaemonMessage::ModelsFailed { .. }));
if let DaemonMessage::ModelsFailed { error } = &msg {
assert_eq!(error, "daemon is locked");
}
}
#[test]
fn handle_get_credential_sync_some() {
let (daemon_tx, daemon_rx) = mpsc::channel();
let (sink, writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let mut none_id = None;
let mut none_tx = None;
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut none_id,
attached_session_tx: &mut none_tx,
client_id: 0,
is_unix: true,
};
std::thread::spawn(move || {
if let Ok(DaemonCommand::GetCredential { service, reply }) = daemon_rx.recv() {
assert_eq!(service, "openai");
let _ = reply.send(Some("sk-123".into()));
}
});
handle_get_credential_sync(&mut ctx, "openai".into());
let msg = writer_rx.recv().unwrap();
assert!(matches!(msg, DaemonMessage::Credential { .. }));
if let DaemonMessage::Credential { service, key } = &msg {
assert_eq!(service, "openai");
assert_eq!(key.as_deref(), Some("sk-123"));
}
}
#[test]
fn handle_get_credential_sync_none() {
let (daemon_tx, daemon_rx) = mpsc::channel();
let (sink, writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let mut none_id = None;
let mut none_tx = None;
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut none_id,
attached_session_tx: &mut none_tx,
client_id: 0,
is_unix: true,
};
std::thread::spawn(move || {
if let Ok(DaemonCommand::GetCredential { service, reply }) = daemon_rx.recv() {
assert_eq!(service, "openai");
let _ = reply.send(None);
}
});
handle_get_credential_sync(&mut ctx, "openai".into());
let msg = writer_rx.recv().unwrap();
assert!(matches!(msg, DaemonMessage::Credential { .. }));
if let DaemonMessage::Credential { service, key } = &msg {
assert_eq!(service, "openai");
assert!(key.is_none());
}
}
#[test]
fn switch_session_to_different_sends_detach_to_old() {
let (old_tx, old_rx) = mpsc::channel();
let (new_tx, new_rx) = mpsc::channel::<SessionCommand>();
let (sink, _writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let (daemon_tx, _daemon_rx) = mpsc::channel();
let mut attached_id = Some(1u64);
let mut attached_tx = Some(old_tx);
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut attached_id,
attached_session_tx: &mut attached_tx,
client_id: 42,
is_unix: true,
};
switch_attached_session(2, new_tx, &mut ctx);
assert!(matches!(
old_rx.try_recv().ok(),
Some(SessionCommand::Detach { client_id: 42 })
));
assert!(matches!(
new_rx.try_recv().ok(),
Some(SessionCommand::Attach { client_id: 42, .. })
));
assert_eq!(attached_id, Some(2));
}
#[test]
fn switch_session_same_skips_detach() {
let (old_tx, old_rx) = mpsc::channel();
let (new_tx, new_rx) = mpsc::channel::<SessionCommand>();
let (sink, _writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let (daemon_tx, _daemon_rx) = mpsc::channel();
let mut attached_id = Some(1u64);
let mut attached_tx = Some(old_tx);
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut attached_id,
attached_session_tx: &mut attached_tx,
client_id: 42,
is_unix: true,
};
switch_attached_session(1, new_tx, &mut ctx);
assert!(old_rx.try_recv().is_err());
assert!(matches!(
new_rx.try_recv().ok(),
Some(SessionCommand::Attach { client_id: 42, .. })
));
assert_eq!(attached_id, Some(1));
}
#[test]
fn handle_delete_session_sync_success_no_message_sent() {
let (daemon_tx, daemon_rx) = mpsc::channel();
let (sink, writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let mut none_id = None;
let mut none_tx = None;
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut none_id,
attached_session_tx: &mut none_tx,
client_id: 0,
is_unix: true,
};
std::thread::spawn(move || {
if let Ok(DaemonCommand::DeleteSession { reply, .. }) = daemon_rx.recv() {
let _ = reply.send(Ok(()));
}
});
handle_delete_session_sync(&mut ctx, 42);
assert!(writer_rx.try_recv().is_err());
}
#[test]
fn handle_delete_session_sync_error() {
let (daemon_tx, daemon_rx) = mpsc::channel();
let (sink, writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let mut none_id = None;
let mut none_tx = None;
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut none_id,
attached_session_tx: &mut none_tx,
client_id: 0,
is_unix: true,
};
std::thread::spawn(move || {
if let Ok(DaemonCommand::DeleteSession { reply, .. }) = daemon_rx.recv() {
let _ = reply.send(Err(io::Error::other("db error")));
}
});
handle_delete_session_sync(&mut ctx, 42);
let msg = writer_rx.recv().unwrap();
assert!(matches!(
msg,
DaemonMessage::Session {
event: SessionEvent::SessionDeleteFailed { .. },
..
}
));
if let DaemonMessage::Session {
session_id: Some(session_id),
event: SessionEvent::SessionDeleteFailed { error },
} = &msg
{
assert_eq!(*session_id, 42);
assert_eq!(error, "db error");
}
}
#[test]
fn handle_delete_session_sync_disconnected() {
let (daemon_tx, daemon_rx) = mpsc::channel::<DaemonCommand>();
let (sink, writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let mut none_id = None;
let mut none_tx = None;
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut none_id,
attached_session_tx: &mut none_tx,
client_id: 0,
is_unix: true,
};
drop(daemon_rx);
handle_delete_session_sync(&mut ctx, 42);
assert!(writer_rx.try_recv().is_err());
}
#[test]
fn switch_session_from_none_no_detach() {
let (new_tx, new_rx) = mpsc::channel::<SessionCommand>();
let (sink, _writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let (daemon_tx, _daemon_rx) = mpsc::channel();
let mut attached_id: Option<u64> = None;
let mut attached_tx: Option<mpsc::Sender<SessionCommand>> = None;
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut attached_id,
attached_session_tx: &mut attached_tx,
client_id: 42,
is_unix: true,
};
switch_attached_session(1, new_tx, &mut ctx);
assert_eq!(attached_id, Some(1));
assert!(matches!(
new_rx.try_recv().ok(),
Some(SessionCommand::Attach { client_id: 42, .. })
));
}
#[test]
fn dispatch_undo_when_attached_sends_undo_command() {
let (daemon_tx, _daemon_rx) = mpsc::channel();
let (sink, _writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let (session_tx, session_rx) = mpsc::channel();
let mut attached_id = Some(1u64);
let mut attached_tx = Some(session_tx);
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut attached_id,
attached_session_tx: &mut attached_tx,
client_id: 42,
is_unix: true,
};
dispatch_client_message(ClientMessage::Undo, &mut ctx).unwrap();
assert!(matches!(
session_rx.try_recv().ok(),
Some(SessionCommand::Undo)
));
}
#[test]
fn dispatch_undo_when_not_attached_is_noop() {
let (daemon_tx, _daemon_rx) = mpsc::channel::<DaemonCommand>();
let (sink, writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let mut none_id = None;
let mut none_tx = None;
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut none_id,
attached_session_tx: &mut none_tx,
client_id: 0,
is_unix: true,
};
dispatch_client_message(ClientMessage::Undo, &mut ctx).unwrap();
assert!(writer_rx.try_recv().is_err());
}
#[test]
fn dispatch_redo_when_attached_sends_redo_command() {
let (daemon_tx, _daemon_rx) = mpsc::channel();
let (sink, _writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let (session_tx, session_rx) = mpsc::channel();
let mut attached_id = Some(1u64);
let mut attached_tx = Some(session_tx);
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut attached_id,
attached_session_tx: &mut attached_tx,
client_id: 42,
is_unix: true,
};
dispatch_client_message(ClientMessage::Redo, &mut ctx).unwrap();
assert!(matches!(
session_rx.try_recv().ok(),
Some(SessionCommand::Redo)
));
}
#[test]
fn dispatch_redo_when_not_attached_is_noop() {
let (daemon_tx, _daemon_rx) = mpsc::channel::<DaemonCommand>();
let (sink, writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let mut none_id = None;
let mut none_tx = None;
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut none_id,
attached_session_tx: &mut none_tx,
client_id: 0,
is_unix: true,
};
dispatch_client_message(ClientMessage::Redo, &mut ctx).unwrap();
assert!(writer_rx.try_recv().is_err());
}
#[test]
fn dispatch_continue_generation_when_attached_sends_run_input() {
let (daemon_tx, _daemon_rx) = mpsc::channel();
let (sink, _writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let (session_tx, session_rx) = mpsc::channel();
let mut attached_id = Some(1u64);
let mut attached_tx = Some(session_tx);
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut attached_id,
attached_session_tx: &mut attached_tx,
client_id: 42,
is_unix: true,
};
dispatch_client_message(
ClientMessage::ContinueGeneration { request_id: 7 },
&mut ctx,
)
.unwrap();
let cmd = session_rx.try_recv().expect("should receive RunInput");
assert!(matches!(
&cmd,
SessionCommand::RunInput {
request_id: 7,
input,
} if input == b"Continue."
));
}
#[test]
fn dispatch_continue_generation_when_not_attached_sends_failed() {
let (daemon_tx, _daemon_rx) = mpsc::channel::<DaemonCommand>();
let (sink, writer_rx) = test_sink();
let global_lag = Arc::new(AtomicUsize::new(0));
let mut none_id = None;
let mut none_tx = None;
let mut ctx = ClientCtx {
writer: &sink,
global_lag: &global_lag,
daemon_tx: &daemon_tx,
attached_session_id: &mut none_id,
attached_session_tx: &mut none_tx,
client_id: 0,
is_unix: true,
};
dispatch_client_message(
ClientMessage::ContinueGeneration { request_id: 7 },
&mut ctx,
)
.unwrap();
let msg = writer_rx.recv().expect("should receive Failed");
assert!(matches!(
&msg,
DaemonMessage::Session {
session_id: None,
event: SessionEvent::Failed {
request_id: 7,
error,
},
} if error == "no session attached"
));
}
#[test]
fn channel_writer_forwards_values_and_receiver_sees_close() {
let (tx, rx) = crossbeam_channel::unbounded::<DaemonMessage>();
let mut writer = ChannelConnectionWriter::new(tx);
writer.send_message(&DaemonMessage::Pong).unwrap();
drop(writer);
assert!(matches!(rx.recv(), Ok(DaemonMessage::Pong)));
assert!(
rx.recv().is_err(),
"dropping the sender must close the receiver"
);
}
#[test]
fn channel_writer_send_after_shutdown_errors() {
let (tx, rx) = crossbeam_channel::unbounded::<DaemonMessage>();
let mut writer = ChannelConnectionWriter::new(tx);
writer.shutdown();
assert!(
writer.send_message(&DaemonMessage::Pong).is_err(),
"sending after shutdown must error (the writer thread never does this on the \
graceful path, but the contract must hold)"
);
assert!(rx.recv().is_err(), "shutdown must close the receiver");
}
#[test]
fn writer_thread_delivers_shutting_down_before_channel_close() {
let (tx, rx) = crossbeam_channel::unbounded::<DaemonMessage>();
let (out_tx, out_rx) = crossbeam_channel::unbounded();
let bytes = Arc::new(AtomicUsize::new(0));
let global = Arc::new(AtomicUsize::new(0));
let handle = std::thread::spawn({
let bytes = Arc::clone(&bytes);
let global = Arc::clone(&global);
move || writer_thread(ChannelConnectionWriter::new(out_tx), rx, bytes, global)
});
tx.send(DaemonMessage::Pong).unwrap();
tx.send(DaemonMessage::ShuttingDown).unwrap();
tx.send(DaemonMessage::Pong).unwrap();
handle.join().expect("writer thread panicked");
assert!(matches!(out_rx.recv(), Ok(DaemonMessage::Pong)));
assert!(
matches!(out_rx.recv(), Ok(DaemonMessage::ShuttingDown)),
"ShuttingDown must be delivered BEFORE the channel closes"
);
assert!(
out_rx.recv().is_err(),
"after the notification the channel must be closed (recv errors)"
);
}
}