Skip to main content

studio_worker/
daemon_client.rs

1//! Blocking client for the daemon's local API, used by the tray UI.
2//!
3//! The daemon publishes its URL and bearer token in the discovery file
4//! (`<config dir>/local-api.json`); [`DaemonClient::discover`] reads it
5//! afresh, so a restarted daemon on a new port is found again.
6
7use std::path::Path;
8use std::time::Duration;
9
10use serde::de::DeserializeOwned;
11use serde::Deserialize;
12
13use crate::daemon_api::{DaemonStatus, EditableConfig, ErrorBody, LogsPage, ModelEntry};
14use crate::job_log::JobLog;
15
16/// Per-request timeout.  Every route the UI calls answers at once; a
17/// daemon that takes longer is treated as unreachable.
18pub const REQUEST_TIMEOUT: Duration = Duration::from_secs(5);
19
20/// Why a call to the daemon failed.
21#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
22pub enum ClientError {
23    /// No discovery file, or the daemon did not answer.
24    #[error("daemon not reachable: {0}")]
25    Unreachable(String),
26    /// The daemon answered with an error.
27    #[error("{message} ({code}, HTTP {status})")]
28    Refused {
29        status: u16,
30        code: String,
31        message: String,
32    },
33    /// The daemon answered something this UI cannot read.
34    #[error("unexpected answer from the daemon: {0}")]
35    Unreadable(String),
36}
37
38#[derive(Deserialize)]
39struct Discovery {
40    url: String,
41    token: String,
42}
43
44/// A client bound to one daemon's URL and token.
45#[derive(Clone)]
46pub struct DaemonClient {
47    http: reqwest::blocking::Client,
48    url: String,
49    token: String,
50}
51
52impl DaemonClient {
53    pub fn new(url: &str, token: &str) -> Result<Self, ClientError> {
54        let http = reqwest::blocking::Client::builder()
55            .timeout(REQUEST_TIMEOUT)
56            .build()
57            .map_err(|e| ClientError::Unreachable(e.to_string()))?;
58        Ok(Self {
59            http,
60            url: url.trim_end_matches('/').to_string(),
61            token: token.to_string(),
62        })
63    }
64
65    /// A client for the daemon serving the config at `config_path`, from its
66    /// discovery file.
67    pub fn discover(config_path: &Path) -> Result<Self, ClientError> {
68        let path = crate::config::local_api_discovery_path_for(config_path)
69            .ok_or_else(|| ClientError::Unreachable("no discovery path".into()))?;
70        let text = std::fs::read_to_string(&path).map_err(|e| {
71            ClientError::Unreachable(format!("no discovery file at {}: {e}", path.display()))
72        })?;
73        let discovery: Discovery = serde_json::from_str(&text).map_err(|e| {
74            ClientError::Unreachable(format!("unreadable discovery file {}: {e}", path.display()))
75        })?;
76        Self::new(&discovery.url, &discovery.token)
77    }
78
79    /// The daemon's base URL.
80    pub fn url(&self) -> &str {
81        &self.url
82    }
83
84    fn send(&self, request: reqwest::blocking::RequestBuilder) -> Result<Vec<u8>, ClientError> {
85        let response = request
86            .bearer_auth(&self.token)
87            .send()
88            .map_err(|e| ClientError::Unreachable(e.to_string()))?;
89        let status = response.status().as_u16();
90        let body = response
91            .bytes()
92            .map_err(|e| ClientError::Unreachable(e.to_string()))?
93            .to_vec();
94        if (200..300).contains(&status) {
95            return Ok(body);
96        }
97        let (code, message) = match serde_json::from_slice::<ErrorBody>(&body) {
98            Ok(err) => (err.error, err.message.unwrap_or_default()),
99            Err(_) => (
100                "http_error".to_string(),
101                String::from_utf8_lossy(&body).trim().to_string(),
102            ),
103        };
104        Err(ClientError::Refused {
105            status,
106            code,
107            message,
108        })
109    }
110
111    fn json<T: DeserializeOwned>(
112        &self,
113        request: reqwest::blocking::RequestBuilder,
114    ) -> Result<T, ClientError> {
115        let body = self.send(request)?;
116        serde_json::from_slice(&body).map_err(|e| ClientError::Unreadable(e.to_string()))
117    }
118
119    fn get(&self, path: &str) -> reqwest::blocking::RequestBuilder {
120        self.http.get(format!("{}{path}", self.url))
121    }
122
123    fn post(&self, path: &str) -> reqwest::blocking::RequestBuilder {
124        self.http.post(format!("{}{path}", self.url))
125    }
126
127    pub fn status(&self) -> Result<DaemonStatus, ClientError> {
128        self.json(self.get("/daemon/status"))
129    }
130
131    pub fn logs(&self, after: u64) -> Result<LogsPage, ClientError> {
132        self.json(self.get(&format!("/daemon/logs?after={after}")))
133    }
134
135    pub fn models(&self) -> Result<Vec<ModelEntry>, ClientError> {
136        self.json(self.get("/models"))
137    }
138
139    /// The job's log; `None` when the daemon captured none.
140    pub fn job_log(&self, job_id: &str) -> Result<Option<JobLog>, ClientError> {
141        not_found_is_none(self.json(self.get(&format!("/jobs/{job_id}/log"))))
142    }
143
144    /// The job's PNG thumbnail; `None` when it has none.
145    pub fn thumbnail(&self, job_id: &str) -> Result<Option<Vec<u8>>, ClientError> {
146        not_found_is_none(self.send(self.get(&format!("/jobs/{job_id}/thumbnail"))))
147    }
148
149    /// Pause (`true`) or resume claiming studio jobs.
150    pub fn set_paused(&self, paused: bool) -> Result<(), ClientError> {
151        let path = if paused {
152            "/daemon/pause"
153        } else {
154            "/daemon/resume"
155        };
156        self.send(self.post(path)).map(drop)
157    }
158
159    pub fn put_config(&self, edit: &EditableConfig) -> Result<EditableConfig, ClientError> {
160        self.json(
161            self.http
162                .put(format!("{}/daemon/config", self.url))
163                .json(edit),
164        )
165    }
166
167    /// Load a model; answers its lifecycle state.
168    pub fn load_model(&self, id: &str) -> Result<String, ClientError> {
169        self.lifecycle(id, "load")
170    }
171
172    /// Unload a model; answers its lifecycle state.
173    pub fn unload_model(&self, id: &str) -> Result<String, ClientError> {
174        self.lifecycle(id, "unload")
175    }
176
177    fn lifecycle(&self, id: &str, verb: &str) -> Result<String, ClientError> {
178        #[derive(Deserialize)]
179        struct State {
180            state: String,
181        }
182        let state: State = self.json(self.post(&format!("/models/{id}/{verb}")))?;
183        Ok(state.state)
184    }
185
186    pub fn reset_registration(&self) -> Result<(), ClientError> {
187        self.send(self.post("/daemon/registration/reset")).map(drop)
188    }
189
190    pub fn shutdown(&self) -> Result<(), ClientError> {
191        self.send(self.post("/daemon/shutdown")).map(drop)
192    }
193}
194
195fn not_found_is_none<T>(result: Result<T, ClientError>) -> Result<Option<T>, ClientError> {
196    match result {
197        Ok(value) => Ok(Some(value)),
198        Err(ClientError::Refused { status: 404, .. }) => Ok(None),
199        Err(err) => Err(err),
200    }
201}
202
203#[cfg(test)]
204mod tests {
205    use super::*;
206    use crate::test_support::DaemonHarness;
207
208    #[test]
209    fn discovery_finds_the_daemon_and_reads_its_status() {
210        let daemon = DaemonHarness::start();
211        let client = DaemonClient::discover(&daemon.config_path).unwrap();
212        assert_eq!(client.url(), daemon.url);
213        let status = client.status().unwrap();
214        assert_eq!(status.version, crate::AGENT_VERSION);
215    }
216
217    #[test]
218    fn a_missing_discovery_file_is_unreachable() {
219        let dir = tempfile::tempdir().unwrap();
220        let err = DaemonClient::discover(&dir.path().join("config.toml"))
221            .err()
222            .unwrap();
223        assert!(matches!(err, ClientError::Unreachable(_)), "{err}");
224    }
225
226    #[test]
227    fn an_unreadable_discovery_file_is_unreachable() {
228        let dir = tempfile::tempdir().unwrap();
229        std::fs::write(dir.path().join("local-api.json"), "{").unwrap();
230        let err = DaemonClient::discover(&dir.path().join("config.toml"))
231            .err()
232            .unwrap();
233        assert!(
234            err.to_string().contains("unreadable discovery file"),
235            "{err}"
236        );
237    }
238
239    #[test]
240    fn a_closed_port_is_unreachable() {
241        let port = std::net::TcpListener::bind("127.0.0.1:0")
242            .unwrap()
243            .local_addr()
244            .unwrap()
245            .port();
246        let client = DaemonClient::new(&format!("http://127.0.0.1:{port}"), "t").unwrap();
247        assert!(matches!(client.status(), Err(ClientError::Unreachable(_))));
248    }
249
250    #[test]
251    fn a_wrong_token_is_refused_with_the_status() {
252        let daemon = DaemonHarness::start();
253        let client = DaemonClient::new(&daemon.url, "wrong").unwrap();
254        let err = client.status().unwrap_err();
255        assert!(
256            matches!(err, ClientError::Refused { status: 401, ref code, .. } if code == "http_error"),
257            "{err}"
258        );
259    }
260
261    #[test]
262    fn pause_and_resume_reach_the_daemon() {
263        let daemon = DaemonHarness::start();
264        let client = daemon.client();
265        client.set_paused(true).unwrap();
266        assert!(client.status().unwrap().paused);
267        client.set_paused(false).unwrap();
268        assert!(!client.status().unwrap().paused);
269    }
270
271    #[test]
272    fn a_config_update_round_trips_and_a_bad_one_is_refused() {
273        let daemon = DaemonHarness::start();
274        let client = daemon.client();
275        let mut edit = client.status().unwrap().config;
276        edit.vram_threshold_gb = 9.0;
277        assert_eq!(client.put_config(&edit).unwrap().vram_threshold_gb, 9.0);
278
279        edit.auto_update_interval_secs = 1;
280        let err = client.put_config(&edit).unwrap_err();
281        assert!(
282            matches!(err, ClientError::Refused { status: 400, ref code, .. } if code == "invalid_config"),
283            "{err}"
284        );
285    }
286
287    #[test]
288    fn models_load_and_unload() {
289        let daemon = DaemonHarness::start();
290        let client = daemon.client();
291        let models = client.models().unwrap();
292        assert_eq!(models[0].id, "chat");
293        assert_eq!(models[0].state, "unloaded");
294        client.load_model("chat").unwrap();
295        daemon.wait_state("chat", "loaded");
296        assert!(client.models().unwrap()[0].resident);
297        client.unload_model("chat").unwrap();
298        daemon.wait_state("chat", "unloaded");
299        let err = client.load_model("nope").unwrap_err();
300        assert!(
301            matches!(err, ClientError::Refused { status: 404, ref code, .. } if code == "unknown_model"),
302            "{err}"
303        );
304    }
305
306    #[test]
307    fn a_job_log_and_thumbnail_are_fetched_or_absent() {
308        crate::test_support::install_job_log_capture();
309        let daemon = DaemonHarness::start();
310        let client = daemon.client();
311        let job_id = daemon.run_image_job();
312        let log = client.job_log(&job_id).unwrap().expect("a log");
313        assert!(log
314            .lines
315            .iter()
316            .any(|l| l.message.starts_with("job finished")));
317        let png = client.thumbnail(&job_id).unwrap().expect("a thumbnail");
318        assert!(image::load_from_memory(&png).is_ok());
319        assert_eq!(client.job_log("no-such-job").unwrap(), None);
320        assert_eq!(client.thumbnail("no-such-job").unwrap(), None);
321    }
322
323    #[test]
324    fn logs_page_after_a_sequence_number() {
325        let daemon = DaemonHarness::start();
326        let client = daemon.client();
327        daemon.push_log("first");
328        daemon.push_log("second");
329        let page = client.logs(0).unwrap();
330        assert_eq!(page.seq, 2);
331        assert_eq!(page.entries.len(), 2);
332        let page = client.logs(1).unwrap();
333        assert_eq!(page.entries[0].message, "second");
334    }
335
336    #[test]
337    fn a_registration_reset_is_refused_unless_rejected_and_shutdown_stops() {
338        let daemon = DaemonHarness::start();
339        let client = daemon.client();
340        let err = client.reset_registration().unwrap_err();
341        assert!(
342            matches!(err, ClientError::Refused { status: 409, .. }),
343            "{err}"
344        );
345        *daemon.control.registration.lock() =
346            crate::auto_register::RegistrationState::Rejected { reason: "x".into() };
347        client.reset_registration().unwrap();
348        client.shutdown().unwrap();
349        assert!(daemon
350            .control
351            .stop
352            .load(std::sync::atomic::Ordering::SeqCst));
353    }
354}