omni_dev/daemon/
client.rs1use std::path::{Path, PathBuf};
6
7use anyhow::{bail, Context, Result};
8use futures::{SinkExt, StreamExt};
9use tokio::net::UnixStream;
10use tokio_util::codec::{Framed, LinesCodec};
11
12use super::protocol::{DaemonEnvelope, DaemonReply, StatusReport, MAX_LINE_BYTES};
13
14#[derive(Debug, Clone)]
18pub struct DaemonClient {
19 socket_path: PathBuf,
20}
21
22impl DaemonClient {
23 pub fn new(socket_path: impl Into<PathBuf>) -> Self {
25 Self {
26 socket_path: socket_path.into(),
27 }
28 }
29
30 pub fn socket_path(&self) -> &Path {
32 &self.socket_path
33 }
34
35 pub async fn request(&self, envelope: DaemonEnvelope) -> Result<DaemonReply> {
37 let stream = UnixStream::connect(&self.socket_path)
38 .await
39 .with_context(|| {
40 format!(
41 "failed to connect to daemon socket {} (is the daemon running?)",
42 self.socket_path.display()
43 )
44 })?;
45 let mut framed = Framed::new(stream, LinesCodec::new_with_max_length(MAX_LINE_BYTES));
46 let line = serde_json::to_string(&envelope).context("failed to encode daemon request")?;
47 framed
48 .send(line)
49 .await
50 .context("failed to send daemon request")?;
51 let response = framed
52 .next()
53 .await
54 .context("daemon closed the connection without replying")?
55 .context("failed to read daemon reply")?;
56 serde_json::from_str(&response).context("failed to decode daemon reply")
57 }
58
59 async fn request_ok(&self, envelope: DaemonEnvelope) -> Result<serde_json::Value> {
62 let reply = self.request(envelope).await?;
63 if reply.ok {
64 Ok(reply.payload)
65 } else {
66 bail!(
67 "daemon returned an error: {}",
68 reply.error.as_deref().unwrap_or("unknown error")
69 )
70 }
71 }
72
73 pub async fn ping(&self) -> Result<()> {
75 self.request_ok(DaemonEnvelope::builtin("ping"))
76 .await
77 .map(|_| ())
78 }
79
80 pub async fn version(&self) -> Result<Option<String>> {
86 let payload = self.request_ok(DaemonEnvelope::builtin("ping")).await?;
87 Ok(payload
88 .get("version")
89 .and_then(serde_json::Value::as_str)
90 .map(str::to_string))
91 }
92
93 pub async fn status(&self) -> Result<StatusReport> {
95 let payload = self.request_ok(DaemonEnvelope::builtin("status")).await?;
96 serde_json::from_value(payload).context("failed to decode daemon status report")
97 }
98
99 pub async fn shutdown(&self) -> Result<()> {
101 self.request_ok(DaemonEnvelope::builtin("shutdown"))
102 .await
103 .map(|_| ())
104 }
105}
106
107#[cfg(test)]
108#[allow(clippy::unwrap_used, clippy::expect_used)]
109mod tests {
110 use super::*;
111 use crate::daemon::testutil::fake_daemon_reply;
112
113 #[tokio::test]
114 async fn version_reads_the_version_from_a_ping_reply() {
115 let (_dir, sock, server) = fake_daemon_reply(
116 serde_json::json!({ "ok": true, "payload": { "pong": true, "version": "1.2.3" } }),
117 );
118 let version = DaemonClient::new(&sock).version().await.unwrap();
119 assert_eq!(version.as_deref(), Some("1.2.3"));
120 server.await.unwrap();
121 }
122
123 #[tokio::test]
124 async fn version_is_none_for_a_pre_1113_daemon_without_a_version_field() {
125 let (_dir, sock, server) =
128 fake_daemon_reply(serde_json::json!({ "ok": true, "payload": { "pong": true } }));
129 let version = DaemonClient::new(&sock).version().await.unwrap();
130 assert!(version.is_none());
131 server.await.unwrap();
132 }
133
134 #[tokio::test]
135 async fn request_ok_maps_an_error_reply_to_an_err() {
136 let (_dir, sock, server) =
137 fake_daemon_reply(serde_json::json!({ "ok": false, "error": "boom" }));
138 let err = DaemonClient::new(&sock).version().await.unwrap_err();
139 assert!(err.to_string().contains("boom"), "{err}");
140 server.await.unwrap();
141 }
142}