use std::path::{Path, PathBuf};
use std::time::Duration;
use anyhow::Context;
use serde::de::DeserializeOwned;
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::net::UnixStream;
use super::bootstrap::ensure_daemon_running;
use super::ipc::{DaemonRequest, DaemonRequestPayload, DaemonResponse, ErrorKind};
use super::socket::daemon_socket_paths;
#[derive(Debug, Clone)]
pub struct DaemonClient {
socket: PathBuf,
}
pub async fn daemon_client() -> anyhow::Result<DaemonClient> {
let paths = daemon_socket_paths()?;
connect_primary_or_sibling(paths.canonical(), paths.legacy())
.await
.ok_or_else(|| daemon_client_unavailable(paths.canonical().display()))
}
pub async fn client_for(socket: Option<&Path>) -> anyhow::Result<DaemonClient> {
client_for_with(socket, ensure_daemon_running).await
}
pub async fn client_for_with<F, Fut>(
socket: Option<&Path>,
respawn: F,
) -> anyhow::Result<DaemonClient>
where
F: FnOnce() -> Fut,
Fut: std::future::Future<Output = anyhow::Result<()>>,
{
match socket {
Some(path) => connect_at(path).await,
None => {
let paths = daemon_socket_paths()?;
connect_or_spawn(paths.canonical(), paths.legacy(), respawn).await
}
}
}
pub(crate) async fn connect_primary_or_sibling(
canonical: &Path,
legacy: Option<&Path>,
) -> Option<DaemonClient> {
if let Ok(client) = connect_at(canonical).await {
return Some(client);
}
let legacy = legacy?;
let client = connect_at(legacy).await.ok()?;
tracing::debug!(legacy = %legacy.display(), "adopted live legacy daemon (#218)");
Some(client)
}
async fn connect_or_spawn<F, Fut>(
primary: &Path,
sibling: Option<&Path>,
respawn: F,
) -> anyhow::Result<DaemonClient>
where
F: FnOnce() -> Fut,
Fut: std::future::Future<Output = anyhow::Result<()>>,
{
match connect_primary_or_sibling(primary, sibling).await {
Some(client) => Ok(client),
None => {
respawn().await?;
connect_at(primary).await
}
}
}
pub fn daemon_client_unavailable(reason: impl std::fmt::Display) -> anyhow::Error {
anyhow::anyhow!("daemon unavailable: {reason} (start it with `hallouminate daemon`)")
}
pub async fn connect_at(socket: &Path) -> anyhow::Result<DaemonClient> {
UnixStream::connect(socket).await.with_context(|| {
format!(
"daemon unavailable: cannot connect to {} \
(start it with `hallouminate daemon`)",
socket.display()
)
})?;
Ok(DaemonClient {
socket: socket.to_path_buf(),
})
}
impl DaemonClient {
pub fn socket_path(&self) -> &Path {
&self.socket
}
pub async fn call_raw(&self, req: DaemonRequest) -> anyhow::Result<DaemonResponse> {
let mut stream = UnixStream::connect(&self.socket).await.map_err(|e| {
daemon_client_unavailable(format!("connect to {} failed: {e}", self.socket.display()))
})?;
let mut text = serde_json::to_string(&req)?;
text.push('\n');
stream.write_all(text.as_bytes()).await.map_err(|e| {
daemon_client_unavailable(format!("write to {} failed: {e}", self.socket.display()))
})?;
stream.flush().await.map_err(|e| {
daemon_client_unavailable(format!("flush {} failed: {e}", self.socket.display()))
})?;
let (read_half, _) = stream.into_split();
let mut reader = BufReader::new(read_half);
let mut line = String::new();
let n = reader.read_line(&mut line).await.map_err(|e| {
daemon_client_unavailable(format!("read from {} failed: {e}", self.socket.display()))
})?;
if n == 0 {
return Err(daemon_client_unavailable(format!(
"daemon at {} closed the connection before responding",
self.socket.display(),
)));
}
let response: DaemonResponse = serde_json::from_str(line.trim_end()).map_err(|e| {
daemon_client_unavailable(format!(
"invalid daemon response from {}: {e} (response: {line:?})",
self.socket.display(),
))
})?;
Ok(response)
}
pub async fn call_raw_with_timeout(
&self,
req: DaemonRequest,
timeout: Duration,
) -> anyhow::Result<DaemonResponse> {
match tokio::time::timeout(timeout, self.call_raw(req)).await {
Ok(result) => result,
Err(_elapsed) => Err(DaemonRpcError::retryable(format!(
"no response from {} within {}s; the daemon may be busy — retry",
self.socket.display(),
timeout.as_secs(),
))
.into()),
}
}
pub async fn call<T: DeserializeOwned>(&self, req: DaemonRequest) -> anyhow::Result<T> {
let timeout = timeout_for(&req.payload);
match self.call_raw_with_timeout(req, timeout).await? {
DaemonResponse::Ok { result } => serde_json::from_value(result)
.map_err(|e| anyhow::anyhow!("daemon returned unexpected payload: {e}")),
DaemonResponse::Err { kind, message } => match kind {
ErrorKind::InvalidParams => Err(DaemonRpcError::invalid_params(message).into()),
ErrorKind::Internal => Err(DaemonRpcError::internal(message).into()),
ErrorKind::Retryable => Err(DaemonRpcError::retryable(message).into()),
},
}
}
}
fn timeout_for(payload: &DaemonRequestPayload) -> Duration {
match payload {
DaemonRequestPayload::Ground(_) => Duration::from_secs(120),
DaemonRequestPayload::Index(_)
| DaemonRequestPayload::AddMarkdown(_)
| DaemonRequestPayload::DeleteMarkdown(_) => Duration::from_secs(15 * 60),
DaemonRequestPayload::Ping
| DaemonRequestPayload::ListCorpora
| DaemonRequestPayload::ListFiles(_)
| DaemonRequestPayload::ListTree(_)
| DaemonRequestPayload::ReadMarkdown(_)
| DaemonRequestPayload::Backlinks(_)
| DaemonRequestPayload::CorpusStats { .. }
| DaemonRequestPayload::Status
| DaemonRequestPayload::Shutdown => Duration::from_secs(60),
}
}
#[derive(Debug)]
pub struct DaemonRpcError {
pub kind: ErrorKind,
pub message: String,
}
impl DaemonRpcError {
pub fn invalid_params(msg: impl Into<String>) -> Self {
Self {
kind: ErrorKind::InvalidParams,
message: msg.into(),
}
}
pub fn retryable(msg: impl Into<String>) -> Self {
Self {
kind: ErrorKind::Retryable,
message: msg.into(),
}
}
pub fn internal(msg: impl Into<String>) -> Self {
Self {
kind: ErrorKind::Internal,
message: msg.into(),
}
}
}
impl std::fmt::Display for DaemonRpcError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.message)
}
}
impl std::error::Error for DaemonRpcError {}
#[cfg(test)]
mod tests {
use super::super::ipc::{
AddMarkdownRequest, BacklinksRequest, DaemonRequestPayload, DeleteMarkdownRequest,
GroundRequest, IndexRequest, ListFilesRequest, ListTreeRequest, ReadMarkdownRequest,
};
use super::*;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
#[tokio::test]
async fn client_for_with_explicit_socket_never_respawns() {
let tmp = tempfile::tempdir().expect("tempdir");
let missing = tmp.path().join("never.sock");
let calls = Arc::new(AtomicUsize::new(0));
let calls_ref = Arc::clone(&calls);
let result = client_for_with(Some(&missing), || {
calls_ref.fetch_add(1, Ordering::SeqCst);
async { anyhow::Ok(()) }
})
.await;
result.expect_err("connect to a missing socket must fail");
assert_eq!(
calls.load(Ordering::SeqCst),
0,
"explicit-socket path must never invoke the respawn step",
);
}
#[tokio::test]
async fn connect_or_spawn_prefers_live_sibling_over_spawning() {
let tmp = tempfile::tempdir().expect("tempdir");
let primary = tmp.path().join("primary.sock");
let sibling = tmp.path().join("sibling.sock");
let listener = tokio::net::UnixListener::bind(&sibling).expect("bind sibling");
tokio::spawn(async move {
loop {
let Ok((stream, _)) = listener.accept().await else {
break;
};
drop(stream);
}
});
let respawn_calls = Arc::new(AtomicUsize::new(0));
let calls = Arc::clone(&respawn_calls);
let client = connect_or_spawn(&primary, Some(&sibling), || {
calls.fetch_add(1, Ordering::SeqCst);
async { anyhow::Ok(()) }
})
.await
.expect("must connect to the live sibling instead of spawning");
assert_eq!(client.socket_path(), sibling.as_path());
assert_eq!(
respawn_calls.load(Ordering::SeqCst),
0,
"must not spawn a daemon when a sibling already answers",
);
}
#[tokio::test]
async fn connect_or_spawn_spawns_when_no_sibling_candidate() {
let tmp = tempfile::tempdir().expect("tempdir");
let primary = tmp.path().join("primary.sock");
let respawn_calls = Arc::new(AtomicUsize::new(0));
let calls = Arc::clone(&respawn_calls);
let result = connect_or_spawn(&primary, None, || {
calls.fetch_add(1, Ordering::SeqCst);
async { anyhow::Ok(()) }
})
.await;
result.expect_err("a still-missing socket after respawn must fail");
assert_eq!(
respawn_calls.load(Ordering::SeqCst),
1,
"must respawn when there is no sibling candidate to try",
);
}
#[tokio::test]
async fn connect_or_spawn_spawns_when_sibling_also_dead() {
let tmp = tempfile::tempdir().expect("tempdir");
let primary = tmp.path().join("primary.sock");
let sibling = tmp.path().join("sibling.sock");
let respawn_calls = Arc::new(AtomicUsize::new(0));
let calls = Arc::clone(&respawn_calls);
let result = connect_or_spawn(&primary, Some(&sibling), || {
calls.fetch_add(1, Ordering::SeqCst);
async { anyhow::Ok(()) }
})
.await;
result.expect_err("both candidates dead must still fail");
assert_eq!(respawn_calls.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn call_raw_with_timeout_returns_err_when_server_never_replies() {
let tmp = tempfile::tempdir().expect("tempdir");
let sock_path = tmp.path().join("silent.sock");
let listener = tokio::net::UnixListener::bind(&sock_path).expect("bind");
tokio::spawn(async move {
let (_stream, _addr) = listener.accept().await.expect("accept");
std::future::pending::<()>().await;
});
let client = connect_at(&sock_path).await.expect("connect");
let started = std::time::Instant::now();
let result = client
.call_raw_with_timeout(
DaemonRequest {
cwd: PathBuf::from("."),
payload: DaemonRequestPayload::Ping,
},
Duration::from_millis(100),
)
.await;
result.expect_err("a silent server must time out, not hang");
assert!(
started.elapsed() < Duration::from_secs(5),
"call_raw_with_timeout must not block past its deadline",
);
}
#[test]
fn timeout_for_classifies_by_request_class() {
let read_class = Duration::from_secs(60);
let ground_class = Duration::from_secs(120);
let mutation_class = Duration::from_secs(15 * 60);
assert_eq!(timeout_for(&DaemonRequestPayload::Ping), read_class);
assert_eq!(timeout_for(&DaemonRequestPayload::ListCorpora), read_class);
assert_eq!(
timeout_for(&DaemonRequestPayload::ListFiles(ListFilesRequest {
corpus: None
})),
read_class,
);
assert_eq!(
timeout_for(&DaemonRequestPayload::ListTree(ListTreeRequest {
corpus: None
})),
read_class,
);
assert_eq!(
timeout_for(&DaemonRequestPayload::ReadMarkdown(ReadMarkdownRequest {
corpus: None,
path: "x.md".to_string(),
})),
read_class,
);
assert_eq!(
timeout_for(&DaemonRequestPayload::Backlinks(BacklinksRequest {
corpus: None,
path: "x.md".to_string(),
})),
read_class,
);
assert_eq!(
timeout_for(&DaemonRequestPayload::CorpusStats { corpus: None }),
read_class,
);
assert_eq!(timeout_for(&DaemonRequestPayload::Shutdown), read_class);
assert_eq!(
timeout_for(&DaemonRequestPayload::Ground(GroundRequest {
query: "q".to_string(),
corpus: None,
top_files: None,
chunks_per_file: None,
limit: None,
snippet_chars: None,
})),
ground_class,
);
assert_eq!(
timeout_for(&DaemonRequestPayload::Index(IndexRequest {
corpus: None,
paths_from: None,
strict: false,
})),
mutation_class,
);
assert_eq!(
timeout_for(&DaemonRequestPayload::AddMarkdown(
AddMarkdownRequest::default()
)),
mutation_class,
);
assert_eq!(
timeout_for(&DaemonRequestPayload::DeleteMarkdown(
DeleteMarkdownRequest {
corpus: "c".to_string(),
path: "x.md".to_string(),
}
)),
mutation_class,
);
}
#[tokio::test]
async fn rpc_timeout_is_typed_retryable() {
let tmp = tempfile::tempdir().expect("tempdir");
let sock_path = tmp.path().join("silent.sock");
let listener = tokio::net::UnixListener::bind(&sock_path).expect("bind");
tokio::spawn(async move {
let (_stream, _addr) = listener.accept().await.expect("accept");
std::future::pending::<()>().await;
});
let client = connect_at(&sock_path).await.expect("connect");
let err = client
.call_raw_with_timeout(
DaemonRequest {
cwd: PathBuf::from("."),
payload: DaemonRequestPayload::Ping,
},
Duration::from_millis(100),
)
.await
.expect_err("a silent server must time out");
let rpc = err
.downcast_ref::<DaemonRpcError>()
.expect("timeout must surface as a typed DaemonRpcError");
assert_eq!(rpc.kind, ErrorKind::Retryable);
assert!(
rpc.message.contains("retry"),
"CLI callers see only the message, so it must say the error is \
retryable: {}",
rpc.message,
);
}
#[tokio::test]
async fn transport_eof_is_not_typed_retryable() {
let tmp = tempfile::tempdir().expect("tempdir");
let sock_path = tmp.path().join("eof.sock");
let listener = tokio::net::UnixListener::bind(&sock_path).expect("bind");
tokio::spawn(async move {
loop {
let Ok((stream, _)) = listener.accept().await else {
break;
};
drop(stream);
}
});
let client = connect_at(&sock_path).await.expect("connect");
let err = client
.call_raw_with_timeout(
DaemonRequest {
cwd: PathBuf::from("."),
payload: DaemonRequestPayload::Ping,
},
Duration::from_secs(5),
)
.await
.expect_err("an immediate EOF must fail");
assert!(
err.downcast_ref::<DaemonRpcError>().is_none(),
"transport EOF must not be classified retryable: {err:#}",
);
}
}