1use 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
27pub 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 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 pub fn local_addr(&self) -> SocketAddr {
84 self.addr
85 }
86
87 pub fn url(&self) -> String {
89 format!("http://{}", self.addr)
90 }
91
92 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}