use std::path::{Path, PathBuf};
use anyhow::{bail, Context, Result};
use futures::{SinkExt, StreamExt};
use tokio::net::UnixStream;
use tokio_util::codec::{Framed, LinesCodec};
use super::protocol::{DaemonEnvelope, DaemonReply, StatusReport, MAX_LINE_BYTES};
#[derive(Debug, Clone)]
pub struct DaemonClient {
socket_path: PathBuf,
}
impl DaemonClient {
pub fn new(socket_path: impl Into<PathBuf>) -> Self {
Self {
socket_path: socket_path.into(),
}
}
pub fn socket_path(&self) -> &Path {
&self.socket_path
}
pub async fn request(&self, envelope: DaemonEnvelope) -> Result<DaemonReply> {
let stream = UnixStream::connect(&self.socket_path)
.await
.with_context(|| {
format!(
"failed to connect to daemon socket {} (is the daemon running?)",
self.socket_path.display()
)
})?;
let mut framed = Framed::new(stream, LinesCodec::new_with_max_length(MAX_LINE_BYTES));
let line = serde_json::to_string(&envelope).context("failed to encode daemon request")?;
framed
.send(line)
.await
.context("failed to send daemon request")?;
let response = framed
.next()
.await
.context("daemon closed the connection without replying")?
.context("failed to read daemon reply")?;
serde_json::from_str(&response).context("failed to decode daemon reply")
}
pub async fn subscribe(&self, envelope: DaemonEnvelope) -> Result<DaemonSubscription> {
let stream = UnixStream::connect(&self.socket_path)
.await
.with_context(|| {
format!(
"failed to connect to daemon socket {} (is the daemon running?)",
self.socket_path.display()
)
})?;
let mut framed = Framed::new(stream, LinesCodec::new_with_max_length(MAX_LINE_BYTES));
let line = serde_json::to_string(&envelope).context("failed to encode daemon request")?;
framed
.send(line)
.await
.context("failed to send daemon subscription request")?;
Ok(DaemonSubscription { framed })
}
async fn request_ok(&self, envelope: DaemonEnvelope) -> Result<serde_json::Value> {
let reply = self.request(envelope).await?;
if reply.ok {
Ok(reply.payload)
} else {
bail!(
"daemon returned an error: {}",
reply.error.as_deref().unwrap_or("unknown error")
)
}
}
pub async fn ping(&self) -> Result<()> {
self.request_ok(DaemonEnvelope::builtin("ping"))
.await
.map(|_| ())
}
pub async fn version(&self) -> Result<Option<String>> {
let payload = self.request_ok(DaemonEnvelope::builtin("ping")).await?;
Ok(payload
.get("version")
.and_then(serde_json::Value::as_str)
.map(str::to_string))
}
pub async fn status(&self) -> Result<StatusReport> {
let payload = self.request_ok(DaemonEnvelope::builtin("status")).await?;
serde_json::from_value(payload).context("failed to decode daemon status report")
}
pub async fn shutdown(&self) -> Result<()> {
self.request_ok(DaemonEnvelope::builtin("shutdown"))
.await
.map(|_| ())
}
}
#[derive(Debug)]
pub struct DaemonSubscription {
framed: Framed<UnixStream, LinesCodec>,
}
impl DaemonSubscription {
pub async fn next(&mut self) -> Option<Result<DaemonReply>> {
match self.framed.next().await {
Some(Ok(line)) => {
Some(serde_json::from_str(&line).context("failed to decode daemon stream frame"))
}
Some(Err(e)) => Some(Err(
anyhow::Error::new(e).context("failed to read daemon stream")
)),
None => None,
}
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)]
mod tests {
use super::*;
use crate::daemon::testutil::{fake_daemon_reply, fake_daemon_stream};
#[tokio::test]
async fn version_reads_the_version_from_a_ping_reply() {
let (_dir, sock, server) = fake_daemon_reply(
serde_json::json!({ "ok": true, "payload": { "pong": true, "version": "1.2.3" } }),
);
let version = DaemonClient::new(&sock).version().await.unwrap();
assert_eq!(version.as_deref(), Some("1.2.3"));
server.await.unwrap();
}
#[tokio::test]
async fn version_is_none_for_a_pre_1113_daemon_without_a_version_field() {
let (_dir, sock, server) =
fake_daemon_reply(serde_json::json!({ "ok": true, "payload": { "pong": true } }));
let version = DaemonClient::new(&sock).version().await.unwrap();
assert!(version.is_none());
server.await.unwrap();
}
#[tokio::test]
async fn request_ok_maps_an_error_reply_to_an_err() {
let (_dir, sock, server) =
fake_daemon_reply(serde_json::json!({ "ok": false, "error": "boom" }));
let err = DaemonClient::new(&sock).version().await.unwrap_err();
assert!(err.to_string().contains("boom"), "{err}");
server.await.unwrap();
}
#[tokio::test]
async fn subscribe_yields_each_pushed_frame_then_none_on_close() {
let (_dir, sock, server) = fake_daemon_stream(vec![
serde_json::json!({ "ok": true, "payload": { "repos": [], "show_closed": true } }),
serde_json::json!({ "ok": true, "payload": { "repos": [], "show_closed": false } }),
]);
let mut sub = DaemonClient::new(&sock)
.subscribe(DaemonEnvelope::service(
"worktrees",
"subscribe",
serde_json::Value::Null,
))
.await
.unwrap();
let first = sub.next().await.unwrap().unwrap();
assert!(first.ok);
assert_eq!(first.payload["show_closed"], serde_json::json!(true));
let second = sub.next().await.unwrap().unwrap();
assert_eq!(second.payload["show_closed"], serde_json::json!(false));
assert!(sub.next().await.is_none());
server.await.unwrap();
}
#[tokio::test]
async fn subscribe_next_surfaces_a_read_error_on_an_oversized_frame() {
let big = "x".repeat(crate::daemon::protocol::MAX_LINE_BYTES + 1);
let (_dir, sock, server) = fake_daemon_stream(vec![serde_json::Value::String(big)]);
let mut sub = DaemonClient::new(&sock)
.subscribe(DaemonEnvelope::service(
"worktrees",
"subscribe",
serde_json::Value::Null,
))
.await
.unwrap();
let frame = sub.next().await;
assert!(
matches!(frame, Some(Err(_))),
"an oversized frame should surface as a read error"
);
server.await.unwrap();
}
#[tokio::test]
async fn subscribe_errors_when_the_daemon_is_unreachable() {
let err = DaemonClient::new("/nonexistent/omni-dev-sub.sock")
.subscribe(DaemonEnvelope::service(
"worktrees",
"subscribe",
serde_json::Value::Null,
))
.await
.unwrap_err();
assert!(err.to_string().contains("failed to connect"), "{err}");
}
}