Skip to main content

studio_worker/
local_api.rs

1//! Always-on local HTTP API for image generation (127.0.0.1 only, no auth).
2//!
3//! Synchronous: `POST /image` blocks until the engine finishes and returns the
4//! image bytes. Models come from the local [`Catalog`], which the operator can
5//! extend at runtime via `POST /models` — the same `ModelSource` shape the
6//! studio uses. Every job is recorded into the in-app local queue.
7
8use std::net::SocketAddr;
9use std::path::PathBuf;
10use std::sync::atomic::{AtomicBool, Ordering};
11use std::sync::Arc;
12use std::time::Duration;
13
14use parking_lot::Mutex;
15use serde::Deserialize;
16use tiny_http::{Header, Method, Request, Response, Server};
17
18use crate::catalog::{Catalog, CatalogModel};
19use crate::engine::Engine;
20use crate::local::{run_image, LocalError, LocalImageRequest};
21use crate::runtime::{JobOutcome, WorkerObservers};
22use crate::types::TaskResult;
23
24const TRACE_TARGET: &str = "studio_worker::local_api";
25const POLL: Duration = Duration::from_millis(200);
26
27/// The local image API server, bound but not yet serving.
28pub struct LocalApi {
29    engine: Arc<dyn Engine>,
30    catalog: Arc<Mutex<Catalog>>,
31    catalog_path: Option<PathBuf>,
32    observers: WorkerObservers,
33    server: Server,
34    addr: SocketAddr,
35}
36
37#[derive(Deserialize)]
38#[serde(rename_all = "camelCase")]
39struct ImageBody {
40    prompt: String,
41    #[serde(default)]
42    model: Option<String>,
43    #[serde(default)]
44    negative_prompt: Option<String>,
45    #[serde(default)]
46    width: Option<u32>,
47    #[serde(default)]
48    height: Option<u32>,
49    #[serde(default)]
50    steps: Option<u32>,
51    #[serde(default)]
52    seed: Option<u64>,
53    #[serde(default)]
54    ext: Option<String>,
55}
56
57impl LocalApi {
58    /// Bind to `addr` (e.g. `127.0.0.1:0` for an ephemeral port).
59    pub fn bind(
60        addr: &str,
61        engine: Arc<dyn Engine>,
62        catalog: Arc<Mutex<Catalog>>,
63        catalog_path: Option<PathBuf>,
64        observers: WorkerObservers,
65    ) -> anyhow::Result<Self> {
66        let server =
67            Server::http(addr).map_err(|e| anyhow::anyhow!("local api bind {addr}: {e}"))?;
68        let addr = server
69            .server_addr()
70            .to_ip()
71            .ok_or_else(|| anyhow::anyhow!("local api: non-ip listen address"))?;
72        Ok(Self {
73            engine,
74            catalog,
75            catalog_path,
76            observers,
77            server,
78            addr,
79        })
80    }
81
82    /// The bound socket address.
83    pub fn local_addr(&self) -> SocketAddr {
84        self.addr
85    }
86
87    /// The base URL the API is reachable at.
88    pub fn url(&self) -> String {
89        format!("http://{}", self.addr)
90    }
91
92    /// Serve requests until `stop` is set.
93    pub fn serve(&self, stop: &AtomicBool) {
94        while !stop.load(Ordering::Relaxed) {
95            match self.server.recv_timeout(POLL) {
96                Ok(Some(request)) => self.route(request),
97                Ok(None) => {}
98                Err(err) => {
99                    tracing::warn!(target: TRACE_TARGET, error = %err, "local api recv error");
100                    break;
101                }
102            }
103        }
104    }
105
106    fn route(&self, request: Request) {
107        let method = request.method().clone();
108        let url = request.url().to_string();
109        let path = url.split('?').next().unwrap_or("/");
110
111        let outcome = match (&method, path) {
112            (Method::Get, "/healthz") => {
113                respond(request, 200, "application/json", b"{\"ok\":true}")
114            }
115            (Method::Post, "/image") => self.handle_image(request),
116            (Method::Get, "/models") => self.handle_list_models(request),
117            (Method::Post, "/models") => self.handle_add_model(request),
118            (Method::Get, "/jobs") => self.handle_jobs(request),
119            (Method::Delete, p) if p.starts_with("/models/") => {
120                let id = p.trim_start_matches("/models/").to_string();
121                self.handle_delete_model(request, &id)
122            }
123            _ => respond(request, 404, "text/plain", b"not found"),
124        };
125        if let Err(err) = outcome {
126            tracing::warn!(target: TRACE_TARGET, error = %err, "local api respond error");
127        }
128    }
129
130    fn handle_image(&self, mut request: Request) -> std::io::Result<()> {
131        let body = read_body(&mut request)?;
132        let parsed: ImageBody = match serde_json::from_str(&body) {
133            Ok(parsed) => parsed,
134            Err(err) => {
135                return respond(
136                    request,
137                    400,
138                    "text/plain",
139                    format!("bad json: {err}").as_bytes(),
140                )
141            }
142        };
143        let req = LocalImageRequest {
144            prompt: parsed.prompt,
145            model: parsed.model,
146            negative_prompt: parsed.negative_prompt,
147            width: parsed.width,
148            height: parsed.height,
149            steps: parsed.steps,
150            seed: parsed.seed,
151            ext: parsed.ext,
152        };
153
154        let catalog = self.catalog.lock().clone();
155        match run_image(self.engine.as_ref(), &catalog, &self.observers, &req) {
156            Ok(TaskResult::Image { bytes, ext }) => {
157                respond(request, 200, content_type_for(&ext), &bytes)
158            }
159            Ok(_) => respond(request, 500, "text/plain", b"unexpected non-image result"),
160            Err(err) => {
161                let status = match err {
162                    LocalError::Engine(_) => 500,
163                    _ => 400,
164                };
165                respond(request, status, "text/plain", err.to_string().as_bytes())
166            }
167        }
168    }
169
170    fn handle_list_models(&self, request: Request) -> std::io::Result<()> {
171        let catalog = self.catalog.lock();
172        match serde_json::to_vec(&catalog.models) {
173            Ok(body) => respond(request, 200, "application/json", &body),
174            Err(err) => respond(request, 500, "text/plain", err.to_string().as_bytes()),
175        }
176    }
177
178    fn handle_add_model(&self, mut request: Request) -> std::io::Result<()> {
179        let body = read_body(&mut request)?;
180        let model: CatalogModel = match serde_json::from_str(&body) {
181            Ok(model) => model,
182            Err(err) => {
183                return respond(
184                    request,
185                    400,
186                    "text/plain",
187                    format!("bad model: {err}").as_bytes(),
188                )
189            }
190        };
191        let saved = {
192            let mut catalog = self.catalog.lock();
193            catalog.upsert(model);
194            self.persist(&catalog)
195        };
196        match saved {
197            Ok(()) => respond(request, 200, "application/json", b"{\"ok\":true}"),
198            Err(err) => respond(request, 500, "text/plain", err.to_string().as_bytes()),
199        }
200    }
201
202    fn handle_delete_model(&self, request: Request, id: &str) -> std::io::Result<()> {
203        let (existed, saved) = {
204            let mut catalog = self.catalog.lock();
205            let existed = catalog.remove(id);
206            (existed, self.persist(&catalog))
207        };
208        if !existed {
209            return respond(request, 404, "text/plain", b"no such model");
210        }
211        match saved {
212            Ok(()) => respond(request, 200, "application/json", b"{\"ok\":true}"),
213            Err(err) => respond(request, 500, "text/plain", err.to_string().as_bytes()),
214        }
215    }
216
217    fn handle_jobs(&self, request: Request) -> std::io::Result<()> {
218        let jobs: Vec<serde_json::Value> = self
219            .observers
220            .local_jobs
221            .lock()
222            .iter()
223            .map(|job| {
224                let (status, reason) = match &job.outcome {
225                    JobOutcome::Completed => ("completed", None),
226                    JobOutcome::Failed { reason } => ("failed", Some(reason.clone())),
227                };
228                serde_json::json!({
229                    "jobId": job.job_id,
230                    "kind": job.kind.as_str(),
231                    "model": job.model,
232                    "prompt": job.prompt,
233                    "status": status,
234                    "reason": reason,
235                    "startedAt": job.started_at.to_rfc3339(),
236                    "finishedAt": job.finished_at.to_rfc3339(),
237                })
238            })
239            .collect();
240        match serde_json::to_vec(&jobs) {
241            Ok(body) => respond(request, 200, "application/json", &body),
242            Err(err) => respond(request, 500, "text/plain", err.to_string().as_bytes()),
243        }
244    }
245
246    fn persist(&self, catalog: &Catalog) -> std::io::Result<()> {
247        match &self.catalog_path {
248            Some(path) => catalog.save(path),
249            None => Ok(()),
250        }
251    }
252}
253
254fn read_body(request: &mut Request) -> std::io::Result<String> {
255    let mut body = String::new();
256    request.as_reader().read_to_string(&mut body)?;
257    Ok(body)
258}
259
260fn content_type_for(ext: &str) -> &'static str {
261    match ext.to_ascii_lowercase().as_str() {
262        "webp" => "image/webp",
263        "png" => "image/png",
264        "jpg" | "jpeg" => "image/jpeg",
265        "gif" => "image/gif",
266        _ => "application/octet-stream",
267    }
268}
269
270fn respond(request: Request, status: u16, content_type: &str, body: &[u8]) -> std::io::Result<()> {
271    let header = Header::from_bytes(b"Content-Type".as_slice(), content_type.as_bytes())
272        .expect("static content-type header is valid");
273    let response = Response::from_data(body)
274        .with_status_code(status)
275        .with_header(header);
276    request.respond(response)
277}
278
279#[cfg(test)]
280mod tests {
281    use super::*;
282    use crate::catalog::CatalogModel;
283    use crate::engine::SyntheticEngine;
284    use crate::types::{ModelCliDefaults, ModelEngine, ModelSource, TaskKind};
285
286    fn synthetic_model(id: &str) -> CatalogModel {
287        CatalogModel {
288            id: id.into(),
289            display_name: id.into(),
290            kind: TaskKind::Image,
291            vram_gb_estimate: 0.0,
292            description: None,
293            source: ModelSource {
294                engine: ModelEngine::Synthetic,
295                files: vec![],
296                cli_defaults: ModelCliDefaults {
297                    cfg_scale: 1.0,
298                    steps: 4,
299                    width: 64,
300                    height: 64,
301                    ..Default::default()
302                },
303            },
304            enabled: true,
305        }
306    }
307
308    struct Harness {
309        url: String,
310        observers: WorkerObservers,
311        stop: Arc<AtomicBool>,
312        handle: Option<std::thread::JoinHandle<()>>,
313    }
314
315    impl Harness {
316        fn start(catalog: Catalog) -> Self {
317            let engine: Arc<dyn Engine> = Arc::new(SyntheticEngine::new());
318            let observers = WorkerObservers::default();
319            let api = LocalApi::bind(
320                "127.0.0.1:0",
321                engine,
322                Arc::new(Mutex::new(catalog)),
323                None,
324                observers.clone(),
325            )
326            .unwrap();
327            let url = api.url();
328            let stop = Arc::new(AtomicBool::new(false));
329            let stop_thread = stop.clone();
330            let handle = std::thread::spawn(move || api.serve(&stop_thread));
331            Harness {
332                url,
333                observers,
334                stop,
335                handle: Some(handle),
336            }
337        }
338    }
339
340    impl Drop for Harness {
341        fn drop(&mut self) {
342            self.stop.store(true, Ordering::Relaxed);
343            if let Some(handle) = self.handle.take() {
344                let _ = handle.join();
345            }
346        }
347    }
348
349    fn seeded_catalog() -> Catalog {
350        Catalog {
351            models: vec![synthetic_model("synthetic-img")],
352        }
353    }
354
355    #[test]
356    fn post_image_returns_image_bytes_and_records_job() {
357        let h = Harness::start(seeded_catalog());
358        let client = reqwest::blocking::Client::new();
359
360        let res = client
361            .post(format!("{}/image", h.url))
362            .json(&serde_json::json!({ "prompt": "a blue bird" }))
363            .send()
364            .unwrap();
365        assert_eq!(res.status(), 200);
366        assert_eq!(res.headers()["content-type"], "image/webp");
367        let bytes = res.bytes().unwrap();
368        assert!(!bytes.is_empty());
369
370        assert_eq!(h.observers.local_jobs.lock().len(), 1);
371    }
372
373    #[test]
374    fn post_image_honours_requested_ext() {
375        let h = Harness::start(seeded_catalog());
376        let client = reqwest::blocking::Client::new();
377        let res = client
378            .post(format!("{}/image", h.url))
379            .json(&serde_json::json!({ "prompt": "x", "ext": "png" }))
380            .send()
381            .unwrap();
382        assert_eq!(res.status(), 200);
383        assert_eq!(res.headers()["content-type"], "image/png");
384    }
385
386    #[test]
387    fn get_models_lists_catalog() {
388        let h = Harness::start(seeded_catalog());
389        let body = reqwest::blocking::get(format!("{}/models", h.url))
390            .unwrap()
391            .text()
392            .unwrap();
393        assert!(body.contains("synthetic-img"));
394    }
395
396    #[test]
397    fn post_models_adds_a_model_then_lists_it() {
398        let h = Harness::start(seeded_catalog());
399        let client = reqwest::blocking::Client::new();
400        let res = client
401            .post(format!("{}/models", h.url))
402            .json(&synthetic_model("added-model"))
403            .send()
404            .unwrap();
405        assert_eq!(res.status(), 200);
406
407        let body = reqwest::blocking::get(format!("{}/models", h.url))
408            .unwrap()
409            .text()
410            .unwrap();
411        assert!(body.contains("added-model"));
412    }
413
414    #[test]
415    fn unknown_model_is_a_400() {
416        let h = Harness::start(seeded_catalog());
417        let client = reqwest::blocking::Client::new();
418        let res = client
419            .post(format!("{}/image", h.url))
420            .json(&serde_json::json!({ "prompt": "x", "model": "nope" }))
421            .send()
422            .unwrap();
423        assert_eq!(res.status(), 400);
424    }
425
426    #[test]
427    fn invalid_json_is_a_400() {
428        let h = Harness::start(seeded_catalog());
429        let client = reqwest::blocking::Client::new();
430        let res = client
431            .post(format!("{}/image", h.url))
432            .body("not json")
433            .header("content-type", "application/json")
434            .send()
435            .unwrap();
436        assert_eq!(res.status(), 400);
437    }
438
439    #[test]
440    fn healthz_ok() {
441        let h = Harness::start(seeded_catalog());
442        let res = reqwest::blocking::get(format!("{}/healthz", h.url)).unwrap();
443        assert_eq!(res.status(), 200);
444    }
445
446    #[test]
447    fn jobs_endpoint_reports_after_generation() {
448        let h = Harness::start(seeded_catalog());
449        let client = reqwest::blocking::Client::new();
450        client
451            .post(format!("{}/image", h.url))
452            .json(&serde_json::json!({ "prompt": "x" }))
453            .send()
454            .unwrap();
455        let body = reqwest::blocking::get(format!("{}/jobs", h.url))
456            .unwrap()
457            .text()
458            .unwrap();
459        assert!(body.contains("\"completed\""));
460        assert!(body.contains("synthetic-img"));
461    }
462}