1use std::net::SocketAddr;
31use std::path::PathBuf;
32use std::sync::atomic::{AtomicBool, Ordering};
33use std::sync::Arc;
34use std::time::Duration;
35
36use parking_lot::Mutex;
37use serde::Deserialize;
38use tiny_http::{Header, Method, Request, Response, Server};
39
40use crate::catalog::{Catalog, CatalogModel};
41use crate::engine::Engine;
42use crate::host::{HostError, ModelHost, ModelStatus};
43use crate::job_gate::JobGate;
44use crate::lifecycle::ModelState;
45use crate::local::{chat_on_lane, run_image, run_kind, LocalError, LocalImageRequest};
46use crate::runtime::{JobOutcome, WorkerObservers};
47use crate::stt_stream::tokens::StreamTokens;
48use crate::types::{
49 AudioSttParams, AudioTtsParams, ChatMessage, LlmParams, Task, TaskKind, TaskResult, VideoParams,
50};
51
52const TRACE_TARGET: &str = "studio_worker::local_api";
53const POLL: Duration = Duration::from_millis(200);
54
55pub const MAX_BODY_BYTES: usize = 1024 * 1024;
58
59#[derive(Debug, PartialEq, Eq)]
61enum Denial {
62 Host(String),
64 Origin(String),
66 Token,
68}
69
70fn host_is_loopback(host: &str) -> bool {
73 let bare = if let Some(rest) = host.strip_prefix('[') {
76 match rest.split_once(']') {
77 Some((addr, _port)) => addr,
78 None => return false,
79 }
80 } else {
81 host.rsplit_once(':').map(|(h, _)| h).unwrap_or(host)
82 };
83 bare.eq_ignore_ascii_case("localhost") || bare == "127.0.0.1" || bare == "::1"
84}
85
86fn origin_is_loopback(origin: &str) -> bool {
90 let rest = origin
91 .strip_prefix("https://")
92 .or_else(|| origin.strip_prefix("http://"));
93 match rest {
94 Some(host) => host_is_loopback(host),
95 None => false,
96 }
97}
98
99fn token_matches(presented: &str, expected: &str) -> bool {
102 let (a, b) = (presented.as_bytes(), expected.as_bytes());
103 if a.len() != b.len() {
104 return false;
105 }
106 a.iter().zip(b).fold(0u8, |acc, (x, y)| acc | (x ^ y)) == 0
107}
108
109fn deny_reason(
114 host: Option<&str>,
115 origin: Option<&str>,
116 authorization: Option<&str>,
117 token: &str,
118) -> Option<Denial> {
119 if let Some(host) = host {
120 if !host_is_loopback(host) {
121 return Some(Denial::Host(host.to_string()));
122 }
123 }
124 if let Some(origin) = origin {
125 if !origin_is_loopback(origin) {
126 return Some(Denial::Origin(origin.to_string()));
127 }
128 }
129 let presented = authorization.and_then(|a| {
130 a.strip_prefix("Bearer ")
131 .or_else(|| a.strip_prefix("bearer "))
132 });
133 match presented {
134 Some(presented) if token_matches(presented, token) => None,
135 _ => Some(Denial::Token),
136 }
137}
138
139pub struct LocalApi {
141 engine: Arc<dyn Engine>,
142 catalog: Arc<Mutex<Catalog>>,
143 catalog_path: Option<PathBuf>,
144 observers: WorkerObservers,
145 server: Server,
146 addr: SocketAddr,
147 token: String,
149 gate: JobGate,
152 models_root: Option<PathBuf>,
155 services: ModelServices,
157 control: Option<crate::control::DaemonControl>,
160}
161
162#[derive(Clone)]
165pub struct ModelServices {
166 pub host: ModelHost,
167 pub tokens: Arc<StreamTokens>,
168 pub stream_port: Arc<std::sync::atomic::AtomicU16>,
170}
171
172impl ModelServices {
173 pub fn new(host: ModelHost) -> Self {
174 Self {
175 host,
176 tokens: Arc::new(StreamTokens::default()),
177 stream_port: Arc::new(std::sync::atomic::AtomicU16::new(0)),
178 }
179 }
180}
181
182#[derive(Deserialize)]
183#[serde(rename_all = "camelCase")]
184struct StreamTokenBody {
185 model: String,
186 #[serde(default)]
187 ttl_secs: Option<i64>,
188}
189
190const DEFAULT_STREAM_TOKEN_TTL_SECS: i64 = 600;
194
195#[derive(Deserialize)]
196#[serde(rename_all = "camelCase")]
197struct ImageBody {
198 prompt: String,
199 #[serde(default)]
200 model: Option<String>,
201 #[serde(default)]
202 negative_prompt: Option<String>,
203 #[serde(default)]
204 width: Option<u32>,
205 #[serde(default)]
206 height: Option<u32>,
207 #[serde(default)]
208 steps: Option<u32>,
209 #[serde(default)]
210 seed: Option<u64>,
211 #[serde(default)]
212 ext: Option<String>,
213}
214
215#[derive(Deserialize)]
217struct ChatBody {
218 #[serde(default)]
219 model: Option<String>,
220 messages: Vec<ChatMessageBody>,
221 #[serde(default)]
222 max_tokens: Option<u32>,
223 #[serde(default)]
224 temperature: Option<f32>,
225 #[serde(default)]
226 top_p: Option<f32>,
227 #[serde(default)]
228 stop: Option<Vec<String>>,
229 #[serde(default)]
231 chat_template_kwargs: Option<serde_json::Map<String, serde_json::Value>>,
232 #[serde(default)]
234 stream: bool,
235}
236
237#[derive(Deserialize)]
239struct TokenizeBody {
240 content: String,
241 #[serde(default)]
242 model: Option<String>,
243 #[serde(default)]
244 add_special: bool,
245}
246
247const STREAM_BUFFER_CHUNKS: usize = 256;
250
251struct ChannelBody {
254 rx: std::sync::mpsc::Receiver<Vec<u8>>,
255 buf: Vec<u8>,
256 pos: usize,
257}
258
259impl std::io::Read for ChannelBody {
260 fn read(&mut self, out: &mut [u8]) -> std::io::Result<usize> {
261 if self.pos == self.buf.len() {
262 match self.rx.recv() {
263 Ok(next) => {
264 self.buf = next;
265 self.pos = 0;
266 }
267 Err(_) => return Ok(0),
268 }
269 }
270 let n = out.len().min(self.buf.len() - self.pos);
271 out[..n].copy_from_slice(&self.buf[self.pos..self.pos + n]);
272 self.pos += n;
273 Ok(n)
274 }
275}
276
277#[derive(Deserialize)]
278struct ChatMessageBody {
279 role: String,
280 content: String,
281}
282
283#[derive(Deserialize)]
284#[serde(rename_all = "camelCase")]
285struct TtsBody {
286 text: String,
287 #[serde(default)]
288 model: Option<String>,
289 #[serde(default)]
290 voice: Option<String>,
291 #[serde(default)]
292 speed: Option<f32>,
293 #[serde(default)]
294 language: Option<String>,
295 #[serde(default)]
296 ext: Option<String>,
297}
298
299#[derive(Deserialize)]
300#[serde(rename_all = "camelCase")]
301struct SttBody {
302 input_url: String,
303 #[serde(default)]
304 model: Option<String>,
305 #[serde(default)]
306 language: Option<String>,
307}
308
309#[derive(Deserialize)]
310#[serde(rename_all = "camelCase")]
311struct VideoBody {
312 prompt: String,
313 #[serde(default)]
314 model: Option<String>,
315 #[serde(default)]
316 negative_prompt: Option<String>,
317 #[serde(default)]
318 seconds: Option<f32>,
319 #[serde(default)]
320 width: Option<u32>,
321 #[serde(default)]
322 height: Option<u32>,
323 #[serde(default)]
324 ext: Option<String>,
325}
326
327impl LocalApi {
328 #[allow(clippy::too_many_arguments)]
333 pub fn bind(
334 addr: &str,
335 engine: Arc<dyn Engine>,
336 catalog: Arc<Mutex<Catalog>>,
337 catalog_path: Option<PathBuf>,
338 observers: WorkerObservers,
339 token: String,
340 gate: JobGate,
341 models_root: Option<PathBuf>,
342 services: ModelServices,
343 ) -> anyhow::Result<Self> {
344 anyhow::ensure!(
345 !token.is_empty(),
346 "local api: refusing to serve with an empty token"
347 );
348 let server =
349 Server::http(addr).map_err(|e| anyhow::anyhow!("local api bind {addr}: {e}"))?;
350 let addr = server
351 .server_addr()
352 .to_ip()
353 .ok_or_else(|| anyhow::anyhow!("local api: non-ip listen address"))?;
354 Ok(Self {
355 engine,
356 catalog,
357 catalog_path,
358 observers,
359 server,
360 addr,
361 token,
362 gate,
363 models_root,
364 services,
365 control: None,
366 })
367 }
368
369 pub fn with_control(mut self, control: crate::control::DaemonControl) -> Self {
371 self.control = Some(control);
372 self
373 }
374
375 pub fn local_addr(&self) -> SocketAddr {
377 self.addr
378 }
379
380 pub fn url(&self) -> String {
382 format!("http://{}", self.addr)
383 }
384
385 pub fn serve(&self, stop: &AtomicBool) {
396 const WORKERS: usize = 8;
399 std::thread::scope(|scope| {
400 for _ in 0..WORKERS {
401 scope.spawn(|| {
402 while !stop.load(Ordering::Relaxed) {
403 match self.server.recv_timeout(POLL) {
404 Ok(Some(request)) => self.route(request),
405 Ok(None) => {}
406 Err(err) => {
407 tracing::warn!(target: TRACE_TARGET, error = %err, "local api recv error");
408 break;
409 }
410 }
411 }
412 });
413 }
414 });
415 }
416
417 fn route(&self, request: Request) {
418 let method = request.method().clone();
419 let url = request.url().to_string();
420 let path = url.split('?').next().unwrap_or("/");
421
422 if !(method == Method::Get && path == "/healthz") {
425 let header = |name: &'static str| {
426 request
427 .headers()
428 .iter()
429 .find(|h| h.field.equiv(name))
430 .map(|h| h.value.as_str().to_string())
431 };
432 let denial = deny_reason(
433 header("host").as_deref(),
434 header("origin").as_deref(),
435 header("authorization").as_deref(),
436 &self.token,
437 );
438 if let Some(denial) = denial {
439 let (status, body) = match &denial {
440 Denial::Host(host) => {
441 (403, format!("forbidden: non-loopback Host header {host:?}"))
442 }
443 Denial::Origin(origin) => (
444 403,
445 format!("forbidden: cross-site request from Origin {origin:?}"),
446 ),
447 Denial::Token => (
448 401,
449 "missing or invalid Authorization bearer token; local clients \
450 can read the current token from the local-api.json discovery \
451 file in the worker's config directory"
452 .to_string(),
453 ),
454 };
455 tracing::warn!(
456 target: TRACE_TARGET,
457 op = "deny",
458 method = %method,
459 path,
460 status,
461 reason = ?denial,
462 "local api request denied"
463 );
464 if let Err(err) = respond(request, status, "text/plain", body.as_bytes()) {
465 tracing::warn!(target: TRACE_TARGET, error = %err, "local api respond error");
466 }
467 return;
468 }
469 }
470
471 let outcome = match (&method, path) {
472 (Method::Get, "/healthz") => self.handle_healthz(request),
473 (Method::Post, "/image") => self.handle_image(request),
474 (Method::Post, "/v1/chat/completions") => self.handle_chat(request),
475 (Method::Post, "/tokenize") => self.handle_tokenize(request),
476 (Method::Post, "/tts") => self.handle_tts(request),
477 (Method::Post, "/stt") => self.handle_stt(request),
478 (Method::Post, "/video") => self.handle_video(request),
479 (Method::Get, "/models") => self.handle_list_models(request),
480 (Method::Post, "/models") => self.handle_add_model(request),
481 (Method::Get, "/jobs") => self.handle_jobs(request),
482 (Method::Post, "/stream-tokens") => self.handle_stream_token(request),
483 (_, p) if p.starts_with("/daemon/") => self.handle_daemon(request, &method, &url),
484 (Method::Get, p) if job_route(p, "/log").is_some() => {
485 let id = job_route(p, "/log").unwrap_or_default().to_string();
486 self.handle_job_log(request, &id)
487 }
488 (Method::Get, p) if job_route(p, "/thumbnail").is_some() => {
489 let id = job_route(p, "/thumbnail").unwrap_or_default().to_string();
490 self.handle_job_thumbnail(request, &id)
491 }
492 (Method::Get, p) if lifecycle_route(p, "/state").is_some() => {
493 let id = lifecycle_route(p, "/state").unwrap_or_default().to_string();
494 self.respond_lifecycle(request, self.services.host.status(&id), 200)
495 }
496 (Method::Post, p) if lifecycle_route(p, "/load").is_some() => {
497 let id = lifecycle_route(p, "/load").unwrap_or_default().to_string();
498 self.respond_lifecycle(request, self.services.host.load(&id), 202)
499 }
500 (Method::Post, p) if lifecycle_route(p, "/unload").is_some() => {
501 let id = lifecycle_route(p, "/unload")
502 .unwrap_or_default()
503 .to_string();
504 self.respond_lifecycle(request, self.services.host.unload(&id), 202)
505 }
506 (Method::Delete, p) if p.starts_with("/models/") => {
507 let id = p.trim_start_matches("/models/").to_string();
508 self.handle_delete_model(request, &id)
509 }
510 _ => respond(request, 404, "text/plain", b"not found"),
511 };
512 if let Err(err) = outcome {
513 tracing::warn!(target: TRACE_TARGET, error = %err, "local api respond error");
514 }
515 }
516
517 fn handle_healthz(&self, request: Request) -> std::io::Result<()> {
523 let free_bytes = self
524 .models_root
525 .as_deref()
526 .and_then(|root| fs4::available_space(root).ok());
527 let gpu = self
528 .observers
529 .gpu_runtime
530 .lock()
531 .clone()
532 .map(|g| serde_json::json!({ "ok": g.ok, "detail": g.detail }));
533 let body = serde_json::json!({
534 "ok": true,
535 "version": crate::AGENT_VERSION,
536 "busy": self.gate.is_busy(),
537 "engine": self.engine.name(),
538 "modelsRoot": self.models_root.as_ref().map(|p| p.display().to_string()),
539 "modelsRootFreeBytes": free_bytes,
540 "gpuRuntime": gpu,
541 });
542 match serde_json::to_vec(&body) {
543 Ok(bytes) => respond(request, 200, "application/json", &bytes),
544 Err(_) => respond(request, 200, "application/json", b"{\"ok\":true}"),
546 }
547 }
548
549 fn handle_image(&self, mut request: Request) -> std::io::Result<()> {
550 let body = match read_body(&mut request)? {
551 BodyOutcome::Ok(body) => body,
552 BodyOutcome::TooLarge => return respond_too_large(request),
553 };
554 let parsed: ImageBody = match serde_json::from_str(&body) {
555 Ok(parsed) => parsed,
556 Err(err) => {
557 return respond(
558 request,
559 400,
560 "text/plain",
561 format!("bad json: {err}").as_bytes(),
562 )
563 }
564 };
565 let req = LocalImageRequest {
566 prompt: parsed.prompt,
567 model: parsed.model,
568 negative_prompt: parsed.negative_prompt,
569 width: parsed.width,
570 height: parsed.height,
571 steps: parsed.steps,
572 seed: parsed.seed,
573 ext: parsed.ext,
574 };
575
576 let Some(_reservation) = self.gate.try_reserve() else {
580 return respond_busy(request);
581 };
582
583 let catalog = self.catalog.lock().clone();
584 match run_image(self.engine.as_ref(), &catalog, &self.observers, &req) {
585 Ok(TaskResult::Image { bytes, ext }) => {
586 respond(request, 200, content_type_for(&ext), &bytes)
587 }
588 Ok(_) => respond(request, 500, "text/plain", b"unexpected non-image result"),
589 Err(err) => respond_local_err(request, err),
590 }
591 }
592
593 fn handle_chat(&self, mut request: Request) -> std::io::Result<()> {
598 let body = match read_body(&mut request)? {
599 BodyOutcome::Ok(body) => body,
600 BodyOutcome::TooLarge => return respond_too_large(request),
601 };
602 let parsed: ChatBody = match serde_json::from_str(&body) {
603 Ok(p) => p,
604 Err(err) => {
605 return respond(
606 request,
607 400,
608 "text/plain",
609 format!("bad json: {err}").as_bytes(),
610 )
611 }
612 };
613 let prompt_preview = parsed
614 .messages
615 .last()
616 .map(|m| m.content.clone())
617 .unwrap_or_default();
618 let params = LlmParams {
619 messages: parsed
620 .messages
621 .into_iter()
622 .map(|m| ChatMessage {
623 role: m.role,
624 content: m.content,
625 })
626 .collect(),
627 max_tokens: parsed.max_tokens.unwrap_or(512),
628 temperature: parsed.temperature.unwrap_or(0.7),
629 top_p: parsed.top_p,
630 stop: parsed.stop,
631 chat_template_kwargs: parsed.chat_template_kwargs,
632 ..Default::default()
633 };
634 let catalog = self.catalog.lock().clone();
635 if parsed.stream {
636 return self.stream_chat(
637 request,
638 &catalog,
639 parsed.model.as_deref(),
640 &prompt_preview,
641 params,
642 );
643 }
644 if let Some(result) = chat_on_lane(
646 &self.services.host,
647 &catalog,
648 &self.observers,
649 parsed.model.as_deref(),
650 &prompt_preview,
651 params.clone(),
652 ) {
653 return respond_llm(request, result);
654 }
655 let Some(_reservation) = self.gate.try_reserve() else {
656 return respond_busy(request);
657 };
658 let outcome = run_kind(
659 self.engine.as_ref(),
660 &catalog,
661 &self.observers,
662 TaskKind::Llm,
663 parsed.model.as_deref(),
664 &prompt_preview,
665 Task::Llm(params),
666 );
667 respond_llm(request, outcome)
668 }
669
670 fn stream_chat(
673 &self,
674 request: Request,
675 catalog: &Catalog,
676 model_id: Option<&str>,
677 prompt_preview: &str,
678 params: LlmParams,
679 ) -> std::io::Result<()> {
680 let model = match crate::local::resolve_llm(catalog, model_id) {
681 Ok(model) => model.id.clone(),
682 Err(err) => return respond_local_err(request, err),
683 };
684 if let Err(refusal) = self.require_loaded(&model) {
685 return respond_json(request, 409, &refusal);
686 }
687 let (tx, rx) = std::sync::mpsc::sync_channel::<Vec<u8>>(STREAM_BUFFER_CHUNKS);
688 let host = self.services.host.clone();
689 let observers = self.observers.clone();
690 let preview = prompt_preview.to_string();
691 std::thread::spawn(move || {
692 crate::local::stream_on_lane(
693 &host,
694 &observers,
695 &model,
696 &preview,
697 params,
698 &mut |bytes| tx.send(bytes).is_ok(),
699 );
700 });
701 let body = ChannelBody {
702 rx,
703 buf: Vec::new(),
704 pos: 0,
705 };
706 let headers = vec![
707 Header::from_bytes("content-type", "text/event-stream").expect("static header"),
708 Header::from_bytes("cache-control", "no-cache").expect("static header"),
709 ];
710 request.respond(Response::new(200.into(), headers, body, None, None))
711 }
712
713 fn require_loaded(&self, model: &str) -> Result<(), serde_json::Value> {
715 match self.services.host.status(model) {
716 Ok(status) if status.state.serves() => Ok(()),
717 Ok(status) => Err(serde_json::json!({
718 "error": "model_not_loaded",
719 "state": status.state.name(),
720 "message": format!("model {model} is not loaded; POST /models/{model}/load first"),
721 })),
722 Err(err) => {
723 Err(serde_json::json!({ "error": "model_not_loaded", "message": err.to_string() }))
724 }
725 }
726 }
727
728 fn handle_tokenize(&self, mut request: Request) -> std::io::Result<()> {
731 let body = match read_body(&mut request)? {
732 BodyOutcome::Ok(body) => body,
733 BodyOutcome::TooLarge => return respond_too_large(request),
734 };
735 let parsed: TokenizeBody = match serde_json::from_str(&body) {
736 Ok(p) => p,
737 Err(err) => {
738 return respond_json(
739 request,
740 400,
741 &serde_json::json!({ "error": "bad_request", "message": err.to_string() }),
742 )
743 }
744 };
745 let model = {
746 let catalog = self.catalog.lock();
747 match crate::local::resolve_llm(&catalog, parsed.model.as_deref()) {
748 Ok(model) => model.id.clone(),
749 Err(err @ LocalError::UnknownModel(_)) => {
750 return respond_json(
751 request,
752 404,
753 &serde_json::json!({ "error": "unknown_model", "message": err.to_string() }),
754 )
755 }
756 Err(err) => {
757 return respond_json(
758 request,
759 400,
760 &serde_json::json!({ "error": "bad_request", "message": err.to_string() }),
761 )
762 }
763 }
764 };
765 if let Err(refusal) = self.require_loaded(&model) {
766 return respond_json(request, 409, &refusal);
767 }
768 match crate::local::tokenize_on_lane(
769 &self.services.host,
770 &model,
771 &parsed.content,
772 parsed.add_special,
773 ) {
774 Ok(Ok(tokens)) => respond_json(request, 200, &serde_json::json!({ "tokens": tokens })),
775 Ok(Err(err)) => respond_json(
776 request,
777 500,
778 &serde_json::json!({ "error": "tokenize_failed", "message": format!("{err:#}") }),
779 ),
780 Err(err) => respond_json(
781 request,
782 409,
783 &serde_json::json!({ "error": "model_not_loaded", "message": err.to_string() }),
784 ),
785 }
786 }
787
788 fn handle_tts(&self, mut request: Request) -> std::io::Result<()> {
789 let body = match read_body(&mut request)? {
790 BodyOutcome::Ok(body) => body,
791 BodyOutcome::TooLarge => return respond_too_large(request),
792 };
793 let parsed: TtsBody = match serde_json::from_str(&body) {
794 Ok(p) => p,
795 Err(err) => {
796 return respond(
797 request,
798 400,
799 "text/plain",
800 format!("bad json: {err}").as_bytes(),
801 )
802 }
803 };
804 let preview = parsed.text.clone();
805 let params = AudioTtsParams {
806 text: parsed.text,
807 voice: parsed.voice.unwrap_or_else(|| "default".into()),
808 speed: parsed.speed,
809 language: parsed.language,
810 ext: parsed.ext.unwrap_or_else(|| "wav".into()),
811 };
812 let Some(_reservation) = self.gate.try_reserve() else {
813 return respond_busy(request);
814 };
815 let catalog = self.catalog.lock().clone();
816 match run_kind(
817 self.engine.as_ref(),
818 &catalog,
819 &self.observers,
820 TaskKind::AudioTts,
821 parsed.model.as_deref(),
822 &preview,
823 Task::AudioTts(params),
824 ) {
825 Ok(TaskResult::AudioTts { bytes, ext }) => {
826 respond(request, 200, content_type_for(&ext), &bytes)
827 }
828 Ok(_) => respond(request, 500, "text/plain", b"unexpected non-audio result"),
829 Err(err) => respond_local_err(request, err),
830 }
831 }
832
833 fn handle_stt(&self, mut request: Request) -> std::io::Result<()> {
834 let body = match read_body(&mut request)? {
835 BodyOutcome::Ok(body) => body,
836 BodyOutcome::TooLarge => return respond_too_large(request),
837 };
838 let parsed: SttBody = match serde_json::from_str(&body) {
839 Ok(p) => p,
840 Err(err) => {
841 return respond(
842 request,
843 400,
844 "text/plain",
845 format!("bad json: {err}").as_bytes(),
846 )
847 }
848 };
849 let preview = parsed.input_url.clone();
850 let params = AudioSttParams {
851 input_url: parsed.input_url,
852 language: parsed.language,
853 ..Default::default()
854 };
855 let Some(_reservation) = self.gate.try_reserve() else {
856 return respond_busy(request);
857 };
858 let catalog = self.catalog.lock().clone();
859 match run_kind(
860 self.engine.as_ref(),
861 &catalog,
862 &self.observers,
863 TaskKind::AudioStt,
864 parsed.model.as_deref(),
865 &preview,
866 Task::AudioStt(params),
867 ) {
868 Ok(TaskResult::AudioStt { json }) => match serde_json::to_vec(&json) {
869 Ok(bytes) => respond(request, 200, "application/json", &bytes),
870 Err(e) => respond(request, 500, "text/plain", e.to_string().as_bytes()),
871 },
872 Ok(_) => respond(
873 request,
874 500,
875 "text/plain",
876 b"unexpected non-transcript result",
877 ),
878 Err(err) => respond_local_err(request, err),
879 }
880 }
881
882 fn handle_video(&self, mut request: Request) -> std::io::Result<()> {
883 let body = match read_body(&mut request)? {
884 BodyOutcome::Ok(body) => body,
885 BodyOutcome::TooLarge => return respond_too_large(request),
886 };
887 let parsed: VideoBody = match serde_json::from_str(&body) {
888 Ok(p) => p,
889 Err(err) => {
890 return respond(
891 request,
892 400,
893 "text/plain",
894 format!("bad json: {err}").as_bytes(),
895 )
896 }
897 };
898 let preview = parsed.prompt.clone();
899 let params = VideoParams {
900 prompt: parsed.prompt,
901 negative_prompt: parsed.negative_prompt,
902 seconds: parsed.seconds.unwrap_or(2.0),
903 width: parsed.width.unwrap_or(256),
904 height: parsed.height.unwrap_or(256),
905 ext: parsed.ext.unwrap_or_else(|| "mp4".into()),
906 ..Default::default()
907 };
908 let Some(_reservation) = self.gate.try_reserve() else {
909 return respond_busy(request);
910 };
911 let catalog = self.catalog.lock().clone();
912 match run_kind(
913 self.engine.as_ref(),
914 &catalog,
915 &self.observers,
916 TaskKind::Video,
917 parsed.model.as_deref(),
918 &preview,
919 Task::Video(params),
920 ) {
921 Ok(TaskResult::Video { bytes, ext }) => {
922 respond(request, 200, content_type_for(&ext), &bytes)
923 }
924 Ok(_) => respond(request, 500, "text/plain", b"unexpected non-video result"),
925 Err(err) => respond_local_err(request, err),
926 }
927 }
928
929 fn handle_list_models(&self, request: Request) -> std::io::Result<()> {
930 let models = self.catalog.lock().models.clone();
931 let statuses = self.services.host.statuses();
932 let listed: Vec<serde_json::Value> = models
933 .iter()
934 .map(|model| {
935 let mut value = serde_json::to_value(model).unwrap_or_default();
936 if let (Some(obj), Some(status)) = (
937 value.as_object_mut(),
938 statuses.iter().find(|s| s.id == model.id),
939 ) {
940 obj.insert("state".into(), status.state.name().into());
941 obj.insert("resident".into(), status.resident.into());
942 obj.insert("since".into(), status.since.to_rfc3339().into());
943 obj.insert("loadable".into(), self.services.host.can_load(model).into());
944 if let ModelState::Failed { reason } = &status.state {
945 obj.insert("error".into(), reason.clone().into());
946 }
947 }
948 value
949 })
950 .collect();
951 match serde_json::to_vec(&listed) {
952 Ok(body) => respond(request, 200, "application/json", &body),
953 Err(err) => respond(request, 500, "text/plain", err.to_string().as_bytes()),
954 }
955 }
956
957 fn handle_add_model(&self, mut request: Request) -> std::io::Result<()> {
958 let body = match read_body(&mut request)? {
959 BodyOutcome::Ok(body) => body,
960 BodyOutcome::TooLarge => return respond_too_large(request),
961 };
962 let model: CatalogModel = match serde_json::from_str(&body) {
963 Ok(model) => model,
964 Err(err) => {
965 return respond(
966 request,
967 400,
968 "text/plain",
969 format!("bad model: {err}").as_bytes(),
970 )
971 }
972 };
973 let saved = {
974 let mut catalog = self.catalog.lock();
975 catalog.upsert(model);
976 self.persist(&catalog)
977 };
978 match saved {
979 Ok(()) => respond(request, 200, "application/json", b"{\"ok\":true}"),
980 Err(err) => respond(request, 500, "text/plain", err.to_string().as_bytes()),
981 }
982 }
983
984 fn handle_delete_model(&self, request: Request, id: &str) -> std::io::Result<()> {
985 if let Err(err) = self.services.host.unload(id) {
987 if !matches!(err, HostError::UnknownModel(_)) {
988 return respond_json(
989 request,
990 500,
991 &serde_json::json!({ "error": "unload_failed", "message": err.to_string() }),
992 );
993 }
994 }
995 let (existed, saved) = {
996 let mut catalog = self.catalog.lock();
997 let existed = catalog.remove(id);
998 (existed, self.persist(&catalog))
999 };
1000 if !existed {
1001 return respond(request, 404, "text/plain", b"no such model");
1002 }
1003 match saved {
1004 Ok(()) => respond(request, 200, "application/json", b"{\"ok\":true}"),
1005 Err(err) => respond(request, 500, "text/plain", err.to_string().as_bytes()),
1006 }
1007 }
1008
1009 fn respond_lifecycle(
1012 &self,
1013 request: Request,
1014 outcome: Result<ModelStatus, HostError>,
1015 pending_status: u16,
1016 ) -> std::io::Result<()> {
1017 match outcome {
1018 Ok(status) => {
1019 let code = match status.state {
1020 ModelState::Loading | ModelState::Unloading => pending_status,
1021 _ => 200,
1022 };
1023 respond_json(request, code, &status_json(&status))
1024 }
1025 Err(err) => {
1026 let (code, body) = match &err {
1027 HostError::UnknownModel(_) => {
1028 (404, serde_json::json!({ "error": "unknown_model" }))
1029 }
1030 HostError::Disabled(_) => {
1031 (400, serde_json::json!({ "error": "model_disabled" }))
1032 }
1033 HostError::Refused(r) => (
1034 409,
1035 serde_json::json!({
1036 "error": "insufficient_memory",
1037 "neededGib": r.needed_gib,
1038 "freeGib": r.free_gib,
1039 "marginGib": r.margin_gib,
1040 }),
1041 ),
1042 HostError::NotLoaded { state, .. } => (
1043 409,
1044 serde_json::json!({ "error": "model_not_loaded", "state": state }),
1045 ),
1046 HostError::Persist(_) => {
1047 (500, serde_json::json!({ "error": "residency_not_saved" }))
1048 }
1049 HostError::LaneBusy(_) => (409, serde_json::json!({ "error": "model_busy" })),
1050 };
1051 let mut body = body;
1052 body["message"] = err.to_string().into();
1053 tracing::warn!(
1054 target: TRACE_TARGET,
1055 op = "lifecycle",
1056 status = code,
1057 error = %err,
1058 "lifecycle request refused"
1059 );
1060 respond_json(request, code, &body)
1061 }
1062 }
1063 }
1064
1065 fn handle_stream_token(&self, mut request: Request) -> std::io::Result<()> {
1068 let body = match read_body(&mut request)? {
1069 BodyOutcome::Ok(body) => body,
1070 BodyOutcome::TooLarge => return respond_too_large(request),
1071 };
1072 let parsed: StreamTokenBody = match serde_json::from_str(&body) {
1073 Ok(p) => p,
1074 Err(err) => {
1075 return respond_json(
1076 request,
1077 400,
1078 &serde_json::json!({ "error": "bad_request", "message": err.to_string() }),
1079 )
1080 }
1081 };
1082 let model = self.catalog.lock().get(&parsed.model).cloned();
1083 let Some(model) = model else {
1084 return respond_json(
1085 request,
1086 404,
1087 &serde_json::json!({ "error": "unknown_model" }),
1088 );
1089 };
1090 if model.source.engine != crate::types::ModelEngine::Parakeet {
1091 return respond_json(
1092 request,
1093 400,
1094 &serde_json::json!({ "error": "not_a_stream_model" }),
1095 );
1096 }
1097 let port = self.services.stream_port.load(Ordering::SeqCst);
1098 if port == 0 {
1099 return respond_json(
1100 request,
1101 503,
1102 &serde_json::json!({ "error": "stream_listener_down" }),
1103 );
1104 }
1105 let ttl =
1106 chrono::Duration::seconds(parsed.ttl_secs.unwrap_or(DEFAULT_STREAM_TOKEN_TTL_SECS));
1107 let grant = self
1108 .services
1109 .tokens
1110 .mint(&model.id, ttl, chrono::Utc::now());
1111 tracing::info!(
1112 target: TRACE_TARGET,
1113 op = "stream_token",
1114 model = %model.id,
1115 expires_at = %grant.expires_at,
1116 "stream token minted"
1117 );
1118 respond_json(
1119 request,
1120 200,
1121 &serde_json::json!({
1122 "token": grant.token,
1123 "model": grant.model,
1124 "expiresAt": grant.expires_at.to_rfc3339(),
1125 "port": port,
1126 "path": crate::stt_stream::server::STREAM_PATH,
1127 }),
1128 )
1129 }
1130
1131 fn handle_jobs(&self, request: Request) -> std::io::Result<()> {
1132 let jobs: Vec<serde_json::Value> = self
1133 .observers
1134 .local_jobs
1135 .lock()
1136 .iter()
1137 .map(|job| {
1138 let (status, reason) = match &job.outcome {
1139 JobOutcome::Completed => ("completed", None),
1140 JobOutcome::Failed { reason } => ("failed", Some(reason.clone())),
1141 };
1142 serde_json::json!({
1143 "jobId": job.job_id,
1144 "kind": job.kind.as_str(),
1145 "model": job.model,
1146 "prompt": job.prompt,
1147 "status": status,
1148 "reason": reason,
1149 "startedAt": job.started_at.to_rfc3339(),
1150 "finishedAt": job.finished_at.to_rfc3339(),
1151 })
1152 })
1153 .collect();
1154 match serde_json::to_vec(&jobs) {
1155 Ok(body) => respond(request, 200, "application/json", &body),
1156 Err(err) => respond(request, 500, "text/plain", err.to_string().as_bytes()),
1157 }
1158 }
1159
1160 fn handle_daemon(
1162 &self,
1163 mut request: Request,
1164 method: &Method,
1165 url: &str,
1166 ) -> std::io::Result<()> {
1167 let Some(control) = &self.control else {
1168 return respond_json(
1169 request,
1170 503,
1171 &serde_json::json!({ "error": "daemon_control_unavailable" }),
1172 );
1173 };
1174 let path = url.split('?').next().unwrap_or("/");
1175 match (method, path) {
1176 (Method::Get, "/daemon/status") => {
1177 let status = control.status(&self.observers, self.gate.is_busy());
1178 respond_serialised(request, 200, &status)
1179 }
1180 (Method::Get, "/daemon/logs") => {
1181 let after = query_param(url, "after")
1182 .and_then(|v| v.parse::<u64>().ok())
1183 .unwrap_or(0);
1184 let (entries, seq) = crate::runtime::recent_logs_after(&self.observers, after);
1185 respond_serialised(request, 200, &crate::daemon_api::LogsPage { entries, seq })
1186 }
1187 (Method::Post, "/daemon/pause") => {
1188 let paused = control.set_paused(true);
1189 respond_json(request, 200, &serde_json::json!({ "paused": paused }))
1190 }
1191 (Method::Post, "/daemon/resume") => {
1192 let paused = control.set_paused(false);
1193 respond_json(request, 200, &serde_json::json!({ "paused": paused }))
1194 }
1195 (Method::Get, "/daemon/config") => {
1196 respond_serialised(request, 200, &control.editable_config())
1197 }
1198 (Method::Put, "/daemon/config") => {
1199 let body = match read_body(&mut request)? {
1200 BodyOutcome::Ok(body) => body,
1201 BodyOutcome::TooLarge => return respond_too_large(request),
1202 };
1203 let edit: crate::daemon_api::EditableConfig = match serde_json::from_str(&body) {
1204 Ok(edit) => edit,
1205 Err(err) => {
1206 return respond_error(request, 400, "bad_request", &err.to_string())
1207 }
1208 };
1209 match control.update_config(edit) {
1210 Ok(saved) => respond_serialised(request, 200, &saved),
1211 Err(err @ crate::control::ControlError::Invalid(_)) => {
1212 respond_error(request, 400, "invalid_config", &err.to_string())
1213 }
1214 Err(err) => respond_error(request, 500, "config_not_saved", &err.to_string()),
1215 }
1216 }
1217 (Method::Post, "/daemon/registration/reset") => {
1218 match control.request_registration_reset() {
1219 Ok(()) => respond_json(request, 202, &serde_json::json!({ "ok": true })),
1220 Err(err) => respond_error(request, 409, "not_rejected", &err.to_string()),
1221 }
1222 }
1223 (Method::Post, "/daemon/shutdown") => {
1224 control.shutdown();
1225 respond_json(request, 202, &serde_json::json!({ "ok": true }))
1226 }
1227 _ => respond_error(request, 404, "not_found", "no such daemon route"),
1228 }
1229 }
1230
1231 fn handle_job_log(&self, request: Request, id: &str) -> std::io::Result<()> {
1232 match crate::job_log::global().get(id) {
1233 Some(log) => respond_serialised(request, 200, &log),
1234 None => respond_error(request, 404, "unknown_job", "no log captured for that job"),
1235 }
1236 }
1237
1238 fn handle_job_thumbnail(&self, request: Request, id: &str) -> std::io::Result<()> {
1239 match self.observers.thumbnails.get(id) {
1240 Some(png) => respond(request, 200, "image/png", &png),
1241 None => respond_error(request, 404, "no_thumbnail", "no thumbnail for that job"),
1242 }
1243 }
1244
1245 fn persist(&self, catalog: &Catalog) -> std::io::Result<()> {
1246 match &self.catalog_path {
1247 Some(path) => catalog.save(path),
1248 None => Ok(()),
1249 }
1250 }
1251}
1252
1253enum BodyOutcome {
1255 Ok(String),
1256 TooLarge,
1257}
1258
1259fn read_body(request: &mut Request) -> std::io::Result<BodyOutcome> {
1260 if matches!(request.body_length(), Some(len) if len > MAX_BODY_BYTES) {
1262 return Ok(BodyOutcome::TooLarge);
1263 }
1264 let mut body = String::new();
1267 use std::io::Read as _;
1268 request
1269 .as_reader()
1270 .take(MAX_BODY_BYTES as u64 + 1)
1271 .read_to_string(&mut body)?;
1272 if body.len() > MAX_BODY_BYTES {
1273 return Ok(BodyOutcome::TooLarge);
1274 }
1275 Ok(BodyOutcome::Ok(body))
1276}
1277
1278fn respond_too_large(request: Request) -> std::io::Result<()> {
1279 respond(
1280 request,
1281 413,
1282 "text/plain",
1283 format!("request body exceeds {MAX_BODY_BYTES} bytes").as_bytes(),
1284 )
1285}
1286
1287fn respond_busy(request: Request) -> std::io::Result<()> {
1291 let retry = Header::from_bytes(b"Retry-After".as_slice(), b"2".as_slice())
1292 .expect("static Retry-After header is valid");
1293 let response = Response::from_data(
1294 b"worker is busy with another job (studio or local); retry shortly".to_vec(),
1295 )
1296 .with_status_code(503)
1297 .with_header(retry)
1298 .with_header(
1299 Header::from_bytes(b"Content-Type".as_slice(), b"text/plain".as_slice())
1300 .expect("static content-type header is valid"),
1301 );
1302 request.respond(response)
1303}
1304
1305pub fn write_discovery_file(path: &std::path::Path, url: &str, token: &str) -> anyhow::Result<()> {
1310 let body = serde_json::json!({ "url": url, "token": token });
1311 let text = serde_json::to_string_pretty(&body)?;
1312 crate::config::write_atomic(path, text.as_bytes())?;
1313 tracing::info!(
1314 target: TRACE_TARGET,
1315 op = "discovery",
1316 path = %path.display(),
1317 url,
1318 "local api discovery file written"
1319 );
1320 Ok(())
1321}
1322
1323pub fn remove_discovery_file(path: &std::path::Path) {
1328 if let Err(e) = std::fs::remove_file(path) {
1329 if e.kind() != std::io::ErrorKind::NotFound {
1330 tracing::warn!(
1331 target: TRACE_TARGET,
1332 op = "discovery",
1333 path = %path.display(),
1334 error = %e,
1335 "failed to remove local api discovery file"
1336 );
1337 }
1338 }
1339}
1340
1341fn content_type_for(ext: &str) -> &'static str {
1342 match ext.to_ascii_lowercase().as_str() {
1343 "webp" => "image/webp",
1344 "png" => "image/png",
1345 "jpg" | "jpeg" => "image/jpeg",
1346 "gif" => "image/gif",
1347 "wav" => "audio/wav",
1348 "mp3" => "audio/mpeg",
1349 "ogg" | "opus" => "audio/ogg",
1350 "flac" => "audio/flac",
1351 "mp4" => "video/mp4",
1352 "webm" => "video/webm",
1353 _ => "application/octet-stream",
1354 }
1355}
1356
1357fn respond_local_err(request: Request, err: LocalError) -> std::io::Result<()> {
1361 let status = match err {
1362 LocalError::Engine(_) => 500,
1363 _ => 400,
1364 };
1365 respond(request, status, "text/plain", err.to_string().as_bytes())
1366}
1367
1368fn respond_llm(request: Request, outcome: Result<TaskResult, LocalError>) -> std::io::Result<()> {
1369 match outcome {
1370 Ok(TaskResult::Llm { json }) => match serde_json::to_vec(&json) {
1371 Ok(bytes) => respond(request, 200, "application/json", &bytes),
1372 Err(e) => respond(request, 500, "text/plain", e.to_string().as_bytes()),
1373 },
1374 Ok(_) => respond(request, 500, "text/plain", b"unexpected non-llm result"),
1375 Err(err) => respond_local_err(request, err),
1376 }
1377}
1378
1379fn lifecycle_route<'a>(path: &'a str, suffix: &str) -> Option<&'a str> {
1381 let id = path.strip_prefix("/models/")?.strip_suffix(suffix)?;
1382 (!id.is_empty() && !id.contains('/')).then_some(id)
1383}
1384
1385fn job_route<'a>(path: &'a str, suffix: &str) -> Option<&'a str> {
1387 let id = path.strip_prefix("/jobs/")?.strip_suffix(suffix)?;
1388 (!id.is_empty() && !id.contains('/')).then_some(id)
1389}
1390
1391fn query_param<'a>(url: &'a str, name: &str) -> Option<&'a str> {
1393 url.split_once('?')?
1394 .1
1395 .split('&')
1396 .find_map(|pair| pair.strip_prefix(name)?.strip_prefix('='))
1397}
1398
1399fn status_json(status: &ModelStatus) -> serde_json::Value {
1401 let mut body = serde_json::json!({
1402 "id": status.id,
1403 "state": status.state.name(),
1404 "resident": status.resident,
1405 "since": status.since.to_rfc3339(),
1406 });
1407 if let ModelState::Failed { reason } = &status.state {
1408 body["error"] = reason.clone().into();
1409 }
1410 body
1411}
1412
1413fn respond_serialised<T: serde::Serialize>(
1414 request: Request,
1415 status: u16,
1416 body: &T,
1417) -> std::io::Result<()> {
1418 match serde_json::to_vec(body) {
1419 Ok(bytes) => respond(request, status, "application/json", &bytes),
1420 Err(err) => respond(request, 500, "text/plain", err.to_string().as_bytes()),
1421 }
1422}
1423
1424fn respond_error(request: Request, status: u16, code: &str, message: &str) -> std::io::Result<()> {
1425 respond_serialised(
1426 request,
1427 status,
1428 &crate::daemon_api::ErrorBody {
1429 error: code.to_string(),
1430 message: Some(message.to_string()),
1431 },
1432 )
1433}
1434
1435fn respond_json(request: Request, status: u16, body: &serde_json::Value) -> std::io::Result<()> {
1436 let bytes = serde_json::to_vec(body).unwrap_or_else(|_| b"{}".to_vec());
1437 respond(request, status, "application/json", &bytes)
1438}
1439
1440fn respond(request: Request, status: u16, content_type: &str, body: &[u8]) -> std::io::Result<()> {
1441 let header = Header::from_bytes(b"Content-Type".as_slice(), content_type.as_bytes())
1442 .expect("static content-type header is valid");
1443 let response = Response::from_data(body)
1444 .with_status_code(status)
1445 .with_header(header);
1446 request.respond(response)
1447}
1448
1449#[cfg(test)]
1450mod tests {
1451 use super::*;
1452 use crate::catalog::CatalogModel;
1453 use crate::engine::{EngineCapabilities, SyntheticEngine};
1454 use crate::types::{ModelCliDefaults, ModelEngine, ModelSource, Task, TaskKind};
1455
1456 struct SlowEngine {
1460 inner: SyntheticEngine,
1461 delay: std::time::Duration,
1462 }
1463
1464 impl Engine for SlowEngine {
1465 fn name(&self) -> &'static str {
1466 "slow"
1467 }
1468 fn capabilities(&self) -> EngineCapabilities {
1469 self.inner.capabilities()
1470 }
1471 fn dispatch(&self, model: &str, task: Task) -> anyhow::Result<TaskResult> {
1472 std::thread::sleep(self.delay);
1473 self.inner.dispatch(model, task)
1474 }
1475 }
1476
1477 fn synthetic_model_of(id: &str, kind: TaskKind) -> CatalogModel {
1478 CatalogModel {
1479 kind,
1480 ..synthetic_model(id)
1481 }
1482 }
1483
1484 fn multi_kind_catalog() -> Catalog {
1487 Catalog {
1488 models: vec![
1489 synthetic_model_of("img", TaskKind::Image),
1490 synthetic_model_of("chat", TaskKind::Llm),
1491 synthetic_model_of("tts", TaskKind::AudioTts),
1492 synthetic_model_of("stt", TaskKind::AudioStt),
1493 synthetic_model_of("vid", TaskKind::Video),
1494 ],
1495 ..Default::default()
1496 }
1497 }
1498
1499 fn synthetic_model(id: &str) -> CatalogModel {
1500 CatalogModel {
1501 id: id.into(),
1502 display_name: id.into(),
1503 kind: TaskKind::Image,
1504 vram_gb_estimate: 0.0,
1505 description: None,
1506 source: ModelSource {
1507 engine: ModelEngine::Synthetic,
1508 files: vec![],
1509 cli_defaults: ModelCliDefaults {
1510 cfg_scale: 1.0,
1511 steps: 4,
1512 width: 64,
1513 height: 64,
1514 ..Default::default()
1515 },
1516 },
1517 enabled: true,
1518 origin: "local".into(),
1519 exclusive_group: None,
1520 }
1521 }
1522
1523 const TEST_TOKEN: &str = "test-token-0123456789abcdef";
1524
1525 struct Harness {
1526 url: String,
1527 observers: WorkerObservers,
1528 host: crate::host::ModelHost,
1529 services: ModelServices,
1530 stop: Arc<AtomicBool>,
1531 handle: Option<std::thread::JoinHandle<()>>,
1532 }
1533
1534 impl Harness {
1535 fn start(catalog: Catalog) -> Self {
1536 Self::start_with_gate(catalog, JobGate::new())
1537 }
1538
1539 fn start_with_gate(catalog: Catalog, gate: JobGate) -> Self {
1540 Self::start_full(catalog, gate, 20.0)
1541 }
1542
1543 fn start_with_free(catalog: Catalog, free_gib: f32) -> Self {
1545 Self::start_full(catalog, JobGate::new(), free_gib)
1546 }
1547
1548 fn start_full(catalog: Catalog, gate: JobGate, free_gib: f32) -> Self {
1549 let engine: Arc<dyn Engine> = Arc::new(SyntheticEngine::new());
1550 let observers = WorkerObservers::default();
1551 let catalog = Arc::new(Mutex::new(catalog));
1552 let host = crate::host::ModelHost::new(
1553 catalog.clone(),
1554 Arc::new(crate::test_support::InstantRuntime),
1555 Arc::new(crate::test_support::FixedProbe(free_gib)),
1556 crate::residency::Residency::load_for_serving(None),
1557 );
1558 let services = ModelServices::new(host.clone());
1559 services
1560 .stream_port
1561 .store(4798, std::sync::atomic::Ordering::SeqCst);
1562 let api = LocalApi::bind(
1563 "127.0.0.1:0",
1564 engine,
1565 catalog,
1566 None,
1567 observers.clone(),
1568 TEST_TOKEN.to_string(),
1569 gate.clone(),
1570 None,
1571 services.clone(),
1572 )
1573 .unwrap();
1574 let url = api.url();
1575 let stop = Arc::new(AtomicBool::new(false));
1576 let stop_thread = stop.clone();
1577 let handle = std::thread::spawn(move || api.serve(&stop_thread));
1578 Harness {
1579 url,
1580 observers,
1581 host,
1582 services,
1583 stop,
1584 handle: Some(handle),
1585 }
1586 }
1587
1588 fn post(&self, path: &str) -> reqwest::blocking::RequestBuilder {
1590 reqwest::blocking::Client::new()
1591 .post(format!("{}{}", self.url, path))
1592 .bearer_auth(TEST_TOKEN)
1593 }
1594
1595 fn get(&self, path: &str) -> reqwest::blocking::RequestBuilder {
1597 reqwest::blocking::Client::new()
1598 .get(format!("{}{}", self.url, path))
1599 .bearer_auth(TEST_TOKEN)
1600 }
1601 }
1602
1603 impl Drop for Harness {
1604 fn drop(&mut self) {
1605 self.stop.store(true, Ordering::Relaxed);
1606 if let Some(handle) = self.handle.take() {
1607 let _ = handle.join();
1608 }
1609 }
1610 }
1611
1612 fn test_host(catalog: &Arc<Mutex<Catalog>>) -> crate::host::ModelHost {
1613 crate::host::ModelHost::new(
1614 catalog.clone(),
1615 Arc::new(crate::test_support::InstantRuntime),
1616 Arc::new(crate::test_support::FixedProbe(20.0)),
1617 crate::residency::Residency::load_for_serving(None),
1618 )
1619 }
1620
1621 fn seeded_catalog() -> Catalog {
1622 Catalog {
1623 models: vec![synthetic_model("synthetic-img")],
1624 ..Default::default()
1625 }
1626 }
1627
1628 #[test]
1629 fn post_image_returns_image_bytes_and_records_job() {
1630 let h = Harness::start(seeded_catalog());
1631
1632 let res = h
1633 .post("/image")
1634 .json(&serde_json::json!({ "prompt": "a blue bird" }))
1635 .send()
1636 .unwrap();
1637 assert_eq!(res.status(), 200);
1638 assert_eq!(res.headers()["content-type"], "image/webp");
1639 let bytes = res.bytes().unwrap();
1640 assert!(!bytes.is_empty());
1641
1642 assert_eq!(h.observers.local_jobs.lock().len(), 1);
1643 }
1644
1645 #[test]
1646 fn post_image_honours_requested_ext() {
1647 let h = Harness::start(seeded_catalog());
1648 let res = h
1649 .post("/image")
1650 .json(&serde_json::json!({ "prompt": "x", "ext": "png" }))
1651 .send()
1652 .unwrap();
1653 assert_eq!(res.status(), 200);
1654 assert_eq!(res.headers()["content-type"], "image/png");
1655 }
1656
1657 #[test]
1658 fn get_models_lists_catalog() {
1659 let h = Harness::start(seeded_catalog());
1660 let body = h.get("/models").send().unwrap().text().unwrap();
1661 assert!(body.contains("synthetic-img"));
1662 }
1663
1664 fn json(res: reqwest::blocking::Response) -> (u16, serde_json::Value) {
1665 let status = res.status().as_u16();
1666 (status, res.json().unwrap())
1667 }
1668
1669 fn wait_state(h: &Harness, id: &str, want: &str) -> serde_json::Value {
1670 let deadline = std::time::Instant::now() + std::time::Duration::from_secs(5);
1671 loop {
1672 let (status, body) = json(h.get(&format!("/models/{id}/state")).send().unwrap());
1673 assert_eq!(status, 200, "{body}");
1674 if body["state"] == want {
1675 return body;
1676 }
1677 assert!(
1678 std::time::Instant::now() < deadline,
1679 "never reached {want}: {body}"
1680 );
1681 std::thread::sleep(std::time::Duration::from_millis(10));
1682 }
1683 }
1684
1685 #[test]
1686 fn get_models_carries_state_and_residency() {
1687 let h = Harness::start(seeded_catalog());
1688 let (status, body) = json(h.get("/models").send().unwrap());
1689 assert_eq!(status, 200);
1690 let first = &body.as_array().unwrap()[0];
1691 assert_eq!(first["state"], "unloaded");
1692 assert_eq!(first["resident"], false);
1693 assert!(first["id"].is_string(), "catalogue fields stay: {first}");
1694 }
1695
1696 #[test]
1697 fn load_reaches_loaded_and_marks_resident() {
1698 let h = Harness::start(seeded_catalog());
1699 let (status, body) = json(h.post("/models/synthetic-img/load").send().unwrap());
1700 assert!(status == 202 || status == 200, "{status} {body}");
1701 assert_eq!(body["id"], "synthetic-img");
1702 assert_eq!(body["resident"], true);
1703 let body = wait_state(&h, "synthetic-img", "loaded");
1704 assert!(body["since"].is_string());
1705 let (status, _) = json(h.post("/models/synthetic-img/load").send().unwrap());
1706 assert_eq!(status, 200, "loading a loaded model is a no-op");
1707 }
1708
1709 #[test]
1710 fn unload_frees_and_clears_residency() {
1711 let h = Harness::start(seeded_catalog());
1712 h.post("/models/synthetic-img/load").send().unwrap();
1713 wait_state(&h, "synthetic-img", "loaded");
1714 let (status, body) = json(h.post("/models/synthetic-img/unload").send().unwrap());
1715 assert!(status == 202 || status == 200, "{status} {body}");
1716 assert_eq!(body["resident"], false);
1717 wait_state(&h, "synthetic-img", "unloaded");
1718 let (status, _) = json(h.post("/models/synthetic-img/unload").send().unwrap());
1719 assert_eq!(status, 200, "unloading an unloaded model is a no-op");
1720 }
1721
1722 #[test]
1723 fn a_load_that_does_not_fit_is_a_409_with_the_numbers() {
1724 let mut catalog = seeded_catalog();
1725 catalog.models[0].vram_gb_estimate = 8.0;
1726 let h = Harness::start_with_free(catalog, 4.0);
1727 let (status, body) = json(h.post("/models/synthetic-img/load").send().unwrap());
1728 assert_eq!(status, 409, "{body}");
1729 assert_eq!(body["error"], "insufficient_memory");
1730 assert_eq!(body["neededGib"], 8.0);
1731 assert_eq!(body["freeGib"], 4.0);
1732 assert!(body["marginGib"].is_number());
1733 wait_state(&h, "synthetic-img", "unloaded");
1734 }
1735
1736 #[test]
1737 fn unknown_models_are_404_on_every_lifecycle_route() {
1738 let h = Harness::start(seeded_catalog());
1739 for res in [
1740 h.get("/models/nope/state").send().unwrap(),
1741 h.post("/models/nope/load").send().unwrap(),
1742 h.post("/models/nope/unload").send().unwrap(),
1743 ] {
1744 let (status, body) = json(res);
1745 assert_eq!(status, 404);
1746 assert_eq!(body["error"], "unknown_model");
1747 }
1748 }
1749
1750 #[test]
1751 fn a_disabled_model_cannot_be_loaded() {
1752 let mut catalog = seeded_catalog();
1753 catalog.models[0].enabled = false;
1754 let h = Harness::start(catalog);
1755 let (status, body) = json(h.post("/models/synthetic-img/load").send().unwrap());
1756 assert_eq!(status, 400);
1757 assert_eq!(body["error"], "model_disabled");
1758 }
1759
1760 #[test]
1761 fn lifecycle_routes_need_the_token() {
1762 let h = Harness::start(seeded_catalog());
1763 let res = reqwest::blocking::Client::new()
1764 .post(format!("{}/models/synthetic-img/load", h.url))
1765 .send()
1766 .unwrap();
1767 assert_eq!(res.status(), 401);
1768 wait_state(&h, "synthetic-img", "unloaded");
1769 }
1770
1771 #[test]
1772 fn deleting_a_loaded_model_unloads_it_first() {
1773 let h = Harness::start(seeded_catalog());
1774 h.post("/models/synthetic-img/load").send().unwrap();
1775 wait_state(&h, "synthetic-img", "loaded");
1776 let res = reqwest::blocking::Client::new()
1777 .delete(format!("{}/models/synthetic-img", h.url))
1778 .bearer_auth(TEST_TOKEN)
1779 .send()
1780 .unwrap();
1781 assert_eq!(res.status(), 200);
1782 let unloaded = h.host.wait_for(
1783 "synthetic-img",
1784 |s| *s == crate::lifecycle::ModelState::Unloaded,
1785 std::time::Duration::from_secs(5),
1786 );
1787 assert!(unloaded.is_some(), "weights freed after delete");
1788 assert_eq!(h.host.loaded_gib(), 0.0);
1789 }
1790
1791 fn llm_catalog() -> Catalog {
1792 Catalog {
1793 models: vec![synthetic_model_of("chat-llm", TaskKind::Llm)],
1794 ..Default::default()
1795 }
1796 }
1797
1798 fn chat(h: &Harness, body: serde_json::Value) -> (u16, serde_json::Value) {
1799 json(h.post("/v1/chat/completions").json(&body).send().unwrap())
1800 }
1801
1802 #[test]
1803 fn chat_is_served_on_the_lane_of_a_loaded_model() {
1804 let h = Harness::start(llm_catalog());
1805 h.post("/models/chat-llm/load").send().unwrap();
1806 wait_state(&h, "chat-llm", "loaded");
1807 let (status, body) = chat(
1808 &h,
1809 serde_json::json!({ "model": "chat-llm", "messages": [{ "role": "user", "content": "hi" }] }),
1810 );
1811 assert_eq!(status, 200, "{body}");
1812 assert_eq!(body["choices"][0]["message"]["content"], "resident:hi");
1813 }
1814
1815 #[test]
1816 fn chat_uses_the_default_llm_when_it_is_loaded() {
1817 let h = Harness::start(llm_catalog());
1818 h.post("/models/chat-llm/load").send().unwrap();
1819 wait_state(&h, "chat-llm", "loaded");
1820 let (_, body) = chat(
1821 &h,
1822 serde_json::json!({ "messages": [{ "role": "user", "content": "yo" }] }),
1823 );
1824 assert_eq!(body["choices"][0]["message"]["content"], "resident:yo");
1825 }
1826
1827 fn sse(body: &str) -> (Vec<serde_json::Value>, bool) {
1829 let data: Vec<&str> = body
1830 .split("\n\n")
1831 .filter_map(|e| e.strip_prefix("data: "))
1832 .collect();
1833 let done = data.last() == Some(&"[DONE]");
1834 let frames = data
1835 .iter()
1836 .filter(|d| **d != "[DONE]")
1837 .map(|d| serde_json::from_str(d).unwrap())
1838 .collect();
1839 (frames, done)
1840 }
1841
1842 #[test]
1843 fn a_loaded_model_streams_its_answer_as_server_sent_events() {
1844 let h = Harness::start(llm_catalog());
1845 h.post("/models/chat-llm/load").send().unwrap();
1846 wait_state(&h, "chat-llm", "loaded");
1847 let res = h
1848 .post("/v1/chat/completions")
1849 .json(&serde_json::json!({
1850 "model": "chat-llm",
1851 "stream": true,
1852 "messages": [{ "role": "user", "content": "hi there" }],
1853 }))
1854 .send()
1855 .unwrap();
1856 assert_eq!(res.status(), 200);
1857 assert_eq!(res.headers()["content-type"], "text/event-stream");
1858 let (frames, done) = sse(&res.text().unwrap());
1859 assert!(done, "ends with [DONE]");
1860 let text: String = frames
1861 .iter()
1862 .filter_map(|f| f["choices"][0]["delta"]["content"].as_str())
1863 .collect();
1864 assert_eq!(text, "resident:hi there");
1865 let last = frames.last().unwrap();
1866 assert_eq!(last["choices"][0]["finish_reason"], "stop");
1867 assert_eq!(last["usage"]["total_tokens"], 5);
1868 assert!(frames.len() > 3, "more than one content chunk: {frames:?}");
1869 let jobs = h.observers.local_jobs.lock().clone();
1870 assert_eq!(jobs.front().unwrap().model, "chat-llm");
1871 }
1872
1873 #[test]
1874 fn streaming_needs_the_model_loaded() {
1875 let h = Harness::start(llm_catalog());
1876 let (status, body) = chat(
1877 &h,
1878 serde_json::json!({ "stream": true, "messages": [{ "role": "user", "content": "x" }] }),
1879 );
1880 assert_eq!(status, 409, "{body}");
1881 assert_eq!(body["error"], "model_not_loaded");
1882 }
1883
1884 #[test]
1885 fn streaming_an_unknown_model_is_a_400() {
1886 let h = Harness::start(llm_catalog());
1887 let res = h
1888 .post("/v1/chat/completions")
1889 .json(&serde_json::json!({ "model": "nope", "stream": true, "messages": [] }))
1890 .send()
1891 .unwrap();
1892 assert_eq!(res.status(), 400);
1893 }
1894
1895 #[test]
1896 fn tokenize_counts_with_the_loaded_model() {
1897 let h = Harness::start(llm_catalog());
1898 h.post("/models/chat-llm/load").send().unwrap();
1899 wait_state(&h, "chat-llm", "loaded");
1900 let (status, body) = json(
1901 h.post("/tokenize")
1902 .json(&serde_json::json!({ "content": "hello" }))
1903 .send()
1904 .unwrap(),
1905 );
1906 assert_eq!(status, 200, "{body}");
1907 assert_eq!(body["tokens"].as_array().unwrap().len(), 5);
1908 }
1909
1910 #[test]
1911 fn tokenize_refuses_what_it_cannot_do() {
1912 let h = Harness::start(llm_catalog());
1913 let (status, body) = json(
1914 h.post("/tokenize")
1915 .json(&serde_json::json!({ "content": "x" }))
1916 .send()
1917 .unwrap(),
1918 );
1919 assert_eq!(status, 409, "{body}");
1920 assert_eq!(body["error"], "model_not_loaded");
1921 let (status, body) = json(
1922 h.post("/tokenize")
1923 .json(&serde_json::json!({ "content": "x", "model": "nope" }))
1924 .send()
1925 .unwrap(),
1926 );
1927 assert_eq!(status, 404, "{body}");
1928 let res = h.post("/tokenize").body("no").send().unwrap();
1929 assert_eq!(res.status(), 400);
1930 let res = reqwest::blocking::Client::new()
1931 .post(format!("{}/tokenize", h.url))
1932 .json(&serde_json::json!({ "content": "x" }))
1933 .send()
1934 .unwrap();
1935 assert_eq!(res.status(), 401);
1936 }
1937
1938 #[test]
1939 fn chat_template_kwargs_reach_the_model() {
1940 let h = Harness::start(llm_catalog());
1941 h.post("/models/chat-llm/load").send().unwrap();
1942 wait_state(&h, "chat-llm", "loaded");
1943 let (_, body) = chat(
1944 &h,
1945 serde_json::json!({
1946 "messages": [{ "role": "user", "content": "hi" }],
1947 "chat_template_kwargs": { "enable_thinking": false },
1948 }),
1949 );
1950 assert_eq!(
1951 body["kwargs"],
1952 serde_json::json!({ "enable_thinking": false })
1953 );
1954 }
1955
1956 #[test]
1957 fn chat_on_an_unloaded_model_runs_as_a_transient_job() {
1958 let h = Harness::start(llm_catalog());
1959 let (status, body) = chat(
1960 &h,
1961 serde_json::json!({ "model": "chat-llm", "messages": [{ "role": "user", "content": "hi" }] }),
1962 );
1963 assert_eq!(status, 200, "{body}");
1964 let content = body["choices"][0]["message"]["content"].as_str().unwrap();
1965 assert!(!content.starts_with("resident:"), "{content}");
1966 }
1967
1968 #[test]
1969 fn a_resident_chat_is_recorded_as_a_local_job_and_skips_the_job_gate() {
1970 let gate = JobGate::new();
1971 let h = Harness::start_with_gate(llm_catalog(), gate.clone());
1972 h.post("/models/chat-llm/load").send().unwrap();
1973 wait_state(&h, "chat-llm", "loaded");
1974 let _held = gate.try_reserve().expect("a transient job holds the gate");
1975 let (status, _) = chat(
1976 &h,
1977 serde_json::json!({ "model": "chat-llm", "messages": [{ "role": "user", "content": "lane" }] }),
1978 );
1979 assert_eq!(status, 200, "a loaded model serves on its own lane");
1980 let jobs = h.observers.local_jobs.lock().clone();
1981 let last = jobs.front().expect("recorded");
1982 assert_eq!(last.model, "chat-llm");
1983 assert_eq!(last.prompt, "lane");
1984 }
1985
1986 fn stream_catalog() -> Catalog {
1987 let mut stt = synthetic_model_of("stt-a", TaskKind::AudioStt);
1988 stt.source.engine = crate::types::ModelEngine::Parakeet;
1989 Catalog {
1990 models: vec![stt, synthetic_model_of("chat-llm", TaskKind::Llm)],
1991 ..Default::default()
1992 }
1993 }
1994
1995 #[test]
1996 fn stream_tokens_are_minted_for_streaming_models() {
1997 let h = Harness::start(stream_catalog());
1998 let (status, body) = json(
1999 h.post("/stream-tokens")
2000 .json(&serde_json::json!({ "model": "stt-a", "ttlSecs": 600 }))
2001 .send()
2002 .unwrap(),
2003 );
2004 assert_eq!(status, 200, "{body}");
2005 let token = body["token"].as_str().unwrap();
2006 assert_eq!(token.len(), 64);
2007 assert_eq!(body["model"], "stt-a");
2008 assert_eq!(body["port"], 4798);
2009 assert_eq!(body["path"], "/transcribe");
2010 assert!(body["expiresAt"].is_string());
2011 assert_eq!(
2012 h.services.tokens.check(token, chrono::Utc::now()),
2013 Ok("stt-a".to_string()),
2014 "the listener accepts it"
2015 );
2016 }
2017
2018 #[test]
2019 fn stream_tokens_are_refused_for_other_models() {
2020 let h = Harness::start(stream_catalog());
2021 let (status, body) = json(
2022 h.post("/stream-tokens")
2023 .json(&serde_json::json!({ "model": "chat-llm" }))
2024 .send()
2025 .unwrap(),
2026 );
2027 assert_eq!(status, 400);
2028 assert_eq!(body["error"], "not_a_stream_model");
2029 let (status, body) = json(
2030 h.post("/stream-tokens")
2031 .json(&serde_json::json!({ "model": "nope" }))
2032 .send()
2033 .unwrap(),
2034 );
2035 assert_eq!(status, 404);
2036 assert_eq!(body["error"], "unknown_model");
2037 }
2038
2039 #[test]
2040 fn stream_tokens_need_a_running_listener() {
2041 let mut h = Harness::start(stream_catalog());
2042 h.services
2043 .stream_port
2044 .store(0, std::sync::atomic::Ordering::SeqCst);
2045 let (status, body) = json(
2046 h.post("/stream-tokens")
2047 .json(&serde_json::json!({ "model": "stt-a" }))
2048 .send()
2049 .unwrap(),
2050 );
2051 assert_eq!(status, 503);
2052 assert_eq!(body["error"], "stream_listener_down");
2053 let _ = &mut h;
2054 }
2055
2056 #[test]
2057 fn stream_tokens_need_the_install_token() {
2058 let h = Harness::start(stream_catalog());
2059 let res = reqwest::blocking::Client::new()
2060 .post(format!("{}/stream-tokens", h.url))
2061 .json(&serde_json::json!({ "model": "stt-a" }))
2062 .send()
2063 .unwrap();
2064 assert_eq!(res.status(), 401);
2065 }
2066
2067 #[test]
2068 fn post_models_adds_a_model_then_lists_it() {
2069 let h = Harness::start(seeded_catalog());
2070 let res = h
2071 .post("/models")
2072 .json(&synthetic_model("added-model"))
2073 .send()
2074 .unwrap();
2075 assert_eq!(res.status(), 200);
2076
2077 let body = h.get("/models").send().unwrap().text().unwrap();
2078 assert!(body.contains("added-model"));
2079 }
2080
2081 #[test]
2082 fn unknown_model_is_a_400() {
2083 let h = Harness::start(seeded_catalog());
2084 let res = h
2085 .post("/image")
2086 .json(&serde_json::json!({ "prompt": "x", "model": "nope" }))
2087 .send()
2088 .unwrap();
2089 assert_eq!(res.status(), 400);
2090 }
2091
2092 #[test]
2093 fn invalid_json_is_a_400() {
2094 let h = Harness::start(seeded_catalog());
2095 let res = h
2096 .post("/image")
2097 .body("not json")
2098 .header("content-type", "application/json")
2099 .send()
2100 .unwrap();
2101 assert_eq!(res.status(), 400);
2102 }
2103
2104 #[test]
2105 fn healthz_reports_a_runtime_snapshot() {
2106 let h = Harness::start(seeded_catalog());
2107 let body: serde_json::Value = reqwest::blocking::get(format!("{}/healthz", h.url))
2108 .unwrap()
2109 .json()
2110 .unwrap();
2111 assert_eq!(body["ok"], true);
2112 assert_eq!(body["version"], crate::AGENT_VERSION);
2113 assert_eq!(body["busy"], false);
2114 assert_eq!(body["engine"], "synthetic");
2115 let raw = serde_json::to_string(&body).unwrap();
2117 assert!(
2118 !raw.contains(TEST_TOKEN),
2119 "healthz must not carry the token"
2120 );
2121 }
2122
2123 #[test]
2124 fn healthz_surfaces_gpu_runtime_when_probed() {
2125 let h = Harness::start(seeded_catalog());
2126 crate::runtime::set_gpu_runtime_status(
2128 &h.observers,
2129 Err(anyhow::anyhow!(
2130 "Vulkan runtime not available: install libvulkan1"
2131 )),
2132 );
2133 let body: serde_json::Value = reqwest::blocking::get(format!("{}/healthz", h.url))
2134 .unwrap()
2135 .json()
2136 .unwrap();
2137 assert_eq!(body["gpuRuntime"]["ok"], false);
2138 assert!(body["gpuRuntime"]["detail"]
2139 .as_str()
2140 .unwrap()
2141 .contains("libvulkan1"));
2142 }
2143
2144 #[test]
2151 fn chat_completions_returns_an_openai_shaped_body() {
2152 let h = Harness::start(multi_kind_catalog());
2153 let res = h
2154 .post("/v1/chat/completions")
2155 .json(&serde_json::json!({
2156 "messages": [{"role": "user", "content": "hello there"}],
2157 "max_tokens": 16
2158 }))
2159 .send()
2160 .unwrap();
2161 assert_eq!(res.status(), 200);
2162 assert_eq!(res.headers()["content-type"], "application/json");
2163 let body: serde_json::Value = res.json().unwrap();
2164 assert!(
2166 body.get("choices").is_some() || body.get("object").is_some(),
2167 "expected an OpenAI-ish body, got: {body}"
2168 );
2169 assert!(h
2171 .observers
2172 .local_jobs
2173 .lock()
2174 .iter()
2175 .any(|j| j.kind == TaskKind::Llm));
2176 }
2177
2178 #[test]
2179 fn tts_returns_audio_bytes() {
2180 let h = Harness::start(multi_kind_catalog());
2181 let res = h
2182 .post("/tts")
2183 .json(&serde_json::json!({ "text": "read this aloud" }))
2184 .send()
2185 .unwrap();
2186 assert_eq!(res.status(), 200);
2187 assert_eq!(res.headers()["content-type"], "audio/wav");
2188 assert!(!res.bytes().unwrap().is_empty());
2189 }
2190
2191 #[test]
2192 fn stt_returns_a_transcript_json() {
2193 let h = Harness::start(multi_kind_catalog());
2194 let res = h
2195 .post("/stt")
2196 .json(&serde_json::json!({ "inputUrl": "https://example.com/a.wav" }))
2197 .send()
2198 .unwrap();
2199 assert_eq!(res.status(), 200);
2200 assert_eq!(res.headers()["content-type"], "application/json");
2201 }
2202
2203 #[test]
2204 fn video_returns_bytes() {
2205 let h = Harness::start(multi_kind_catalog());
2206 let res = h
2207 .post("/video")
2208 .json(&serde_json::json!({ "prompt": "a tiny dragon" }))
2209 .send()
2210 .unwrap();
2211 assert_eq!(res.status(), 200);
2212 assert!(!res.bytes().unwrap().is_empty());
2213 }
2214
2215 #[test]
2216 fn chat_without_an_llm_model_is_a_400() {
2217 let h = Harness::start(seeded_catalog());
2220 let res = h
2221 .post("/v1/chat/completions")
2222 .json(&serde_json::json!({
2223 "messages": [{"role": "user", "content": "hi"}]
2224 }))
2225 .send()
2226 .unwrap();
2227 assert_eq!(res.status(), 400);
2228 assert!(res.text().unwrap().contains("llm"));
2229 }
2230
2231 #[test]
2232 fn chat_endpoint_respects_the_busy_gate() {
2233 let gate = JobGate::new();
2234 let h = Harness::start_with_gate(multi_kind_catalog(), gate.clone());
2235 let _held = gate.try_reserve().unwrap();
2236 let res = h
2237 .post("/v1/chat/completions")
2238 .json(&serde_json::json!({ "messages": [{"role":"user","content":"x"}] }))
2239 .send()
2240 .unwrap();
2241 assert_eq!(res.status(), 503);
2242 }
2243
2244 #[test]
2245 fn jobs_endpoint_reports_after_generation() {
2246 let h = Harness::start(seeded_catalog());
2247 h.post("/image")
2248 .json(&serde_json::json!({ "prompt": "x" }))
2249 .send()
2250 .unwrap();
2251 let body = h.get("/jobs").send().unwrap().text().unwrap();
2252 assert!(body.contains("\"completed\""));
2253 assert!(body.contains("synthetic-img"));
2254 }
2255
2256 #[test]
2264 fn routes_reject_requests_without_a_token() {
2265 let h = Harness::start(seeded_catalog());
2266 let client = reqwest::blocking::Client::new();
2267 let cases: Vec<(reqwest::blocking::RequestBuilder, &str)> = vec![
2268 (
2269 client
2270 .post(format!("{}/image", h.url))
2271 .json(&serde_json::json!({ "prompt": "x" })),
2272 "POST /image",
2273 ),
2274 (client.get(format!("{}/models", h.url)), "GET /models"),
2275 (
2276 client
2277 .post(format!("{}/models", h.url))
2278 .json(&synthetic_model("evil")),
2279 "POST /models",
2280 ),
2281 (
2282 client.delete(format!("{}/models/synthetic-img", h.url)),
2283 "DELETE /models",
2284 ),
2285 (client.get(format!("{}/jobs", h.url)), "GET /jobs"),
2286 ];
2287 for (req, name) in cases {
2288 let res = req.send().unwrap();
2289 assert_eq!(res.status(), 401, "{name} must require the token");
2290 let body = res.text().unwrap();
2291 assert!(
2292 body.contains("local-api.json"),
2293 "{name}: the 401 must point at the discovery file, got: {body}"
2294 );
2295 }
2296 let body = h.get("/models").send().unwrap().text().unwrap();
2298 assert!(!body.contains("evil"));
2299 assert!(body.contains("synthetic-img"));
2300 }
2301
2302 #[test]
2303 fn routes_reject_a_wrong_token() {
2304 let h = Harness::start(seeded_catalog());
2305 let res = reqwest::blocking::Client::new()
2306 .get(format!("{}/models", h.url))
2307 .bearer_auth("wrong-token")
2308 .send()
2309 .unwrap();
2310 assert_eq!(res.status(), 401);
2311 }
2312
2313 #[test]
2314 fn daemon_routes_need_daemon_control() {
2315 let h = Harness::start(multi_kind_catalog());
2316 let resp = h.get("/daemon/status").send().unwrap();
2317 assert_eq!(resp.status(), 503);
2318 let body: serde_json::Value = resp.json().unwrap();
2319 assert_eq!(body["error"], "daemon_control_unavailable");
2320 }
2321
2322 #[test]
2323 fn daemon_routes_need_the_token() {
2324 let h = Harness::start(multi_kind_catalog());
2325 let resp = reqwest::blocking::Client::new()
2326 .get(format!("{}/daemon/status", h.url))
2327 .send()
2328 .unwrap();
2329 assert_eq!(resp.status(), 401);
2330 }
2331
2332 #[test]
2333 fn an_unknown_daemon_route_is_not_found() {
2334 let daemon = crate::test_support::DaemonHarness::start();
2335 let resp = reqwest::blocking::Client::new()
2336 .get(format!("{}/daemon/nope", daemon.url))
2337 .bearer_auth(crate::test_support::HARNESS_TOKEN)
2338 .send()
2339 .unwrap();
2340 assert_eq!(resp.status(), 404);
2341 }
2342
2343 #[test]
2344 fn a_malformed_config_body_is_a_bad_request() {
2345 let daemon = crate::test_support::DaemonHarness::start();
2346 let resp = reqwest::blocking::Client::new()
2347 .put(format!("{}/daemon/config", daemon.url))
2348 .bearer_auth(crate::test_support::HARNESS_TOKEN)
2349 .body("{")
2350 .send()
2351 .unwrap();
2352 assert_eq!(resp.status(), 400);
2353 }
2354
2355 #[test]
2356 fn the_models_listing_carries_since_and_the_failure() {
2357 let daemon = crate::test_support::DaemonHarness::start();
2358 let models = daemon.client().models().unwrap();
2359 assert!(models.iter().all(|m| m.since.is_some() && m.loadable));
2360 assert!(models.iter().all(|m| m.error.is_none()));
2361 }
2362
2363 #[test]
2364 fn job_routes_and_query_params_parse() {
2365 assert_eq!(job_route("/jobs/local-1/log", "/log"), Some("local-1"));
2366 assert_eq!(job_route("/jobs//log", "/log"), None);
2367 assert_eq!(job_route("/jobs/a/b/log", "/log"), None);
2368 assert_eq!(query_param("/daemon/logs?after=12", "after"), Some("12"));
2369 assert_eq!(query_param("/daemon/logs?x=1&after=3", "after"), Some("3"));
2370 assert_eq!(query_param("/daemon/logs?afterx=3", "after"), None);
2371 assert_eq!(query_param("/daemon/logs", "after"), None);
2372 }
2373
2374 #[test]
2375 fn healthz_needs_no_token() {
2376 let h = Harness::start(seeded_catalog());
2377 let res = reqwest::blocking::get(format!("{}/healthz", h.url)).unwrap();
2378 assert_eq!(res.status(), 200);
2379 }
2380
2381 #[test]
2382 fn healthz_answers_while_a_generation_is_in_flight() {
2383 let engine: Arc<dyn Engine> = Arc::new(SlowEngine {
2388 inner: SyntheticEngine::new(),
2389 delay: std::time::Duration::from_millis(400),
2390 });
2391 let observers = WorkerObservers::default();
2392 let catalog = Arc::new(Mutex::new(seeded_catalog()));
2393 let api = LocalApi::bind(
2394 "127.0.0.1:0",
2395 engine,
2396 catalog.clone(),
2397 None,
2398 observers,
2399 TEST_TOKEN.to_string(),
2400 JobGate::new(),
2401 None,
2402 ModelServices::new(test_host(&catalog)),
2403 )
2404 .unwrap();
2405 let url = api.url();
2406 let stop = Arc::new(AtomicBool::new(false));
2407 let stop_thread = stop.clone();
2408 let handle = std::thread::spawn(move || api.serve(&stop_thread));
2409
2410 let gen_url = url.clone();
2412 let gen = std::thread::spawn(move || {
2413 reqwest::blocking::Client::new()
2414 .post(format!("{gen_url}/image"))
2415 .bearer_auth(TEST_TOKEN)
2416 .json(&serde_json::json!({ "prompt": "slow" }))
2417 .timeout(std::time::Duration::from_secs(5))
2418 .send()
2419 .unwrap()
2420 .status()
2421 .as_u16()
2422 });
2423
2424 std::thread::sleep(std::time::Duration::from_millis(100));
2427 let start = std::time::Instant::now();
2428 let health = reqwest::blocking::get(format!("{url}/healthz")).unwrap();
2429 let elapsed = start.elapsed();
2430 assert_eq!(health.status(), 200);
2431 assert!(
2432 elapsed < std::time::Duration::from_millis(250),
2433 "healthz blocked behind the generation ({elapsed:?}); the pool isn't concurrent"
2434 );
2435 let body: serde_json::Value = health.json().unwrap();
2437 assert_eq!(body["busy"], true, "a running job must show busy=true");
2438
2439 assert_eq!(gen.join().unwrap(), 200, "the generation still succeeds");
2440 stop.store(true, Ordering::Relaxed);
2441 let _ = handle.join();
2442 }
2443
2444 #[test]
2445 fn non_loopback_host_header_is_forbidden_even_with_a_token() {
2446 let h = Harness::start(seeded_catalog());
2449 let res = h
2450 .get("/models")
2451 .header("host", "evil.example:4787")
2452 .send()
2453 .unwrap();
2454 assert_eq!(res.status(), 403);
2455 assert!(res.text().unwrap().contains("Host"));
2456 }
2457
2458 #[test]
2459 fn cross_site_origin_is_forbidden_even_with_a_token() {
2460 let h = Harness::start(seeded_catalog());
2463 let res = h
2464 .post("/image")
2465 .header("origin", "https://evil.example")
2466 .json(&serde_json::json!({ "prompt": "x" }))
2467 .send()
2468 .unwrap();
2469 assert_eq!(res.status(), 403);
2470 assert!(res.text().unwrap().contains("Origin"));
2471 }
2472
2473 #[test]
2474 fn loopback_origin_is_allowed() {
2475 let h = Harness::start(seeded_catalog());
2478 let res = h
2479 .get("/models")
2480 .header("origin", "http://localhost:5173")
2481 .send()
2482 .unwrap();
2483 assert_eq!(res.status(), 200);
2484 }
2485
2486 #[test]
2487 fn oversized_body_is_a_413() {
2488 let h = Harness::start(seeded_catalog());
2489 let big = "x".repeat(MAX_BODY_BYTES + 1);
2490 let res = h.post("/image").body(big).send().unwrap();
2491 assert_eq!(res.status(), 413);
2492 }
2493
2494 #[test]
2495 fn body_at_the_cap_is_still_read() {
2496 let h = Harness::start(seeded_catalog());
2500 let exact = "x".repeat(MAX_BODY_BYTES);
2501 let res = h.post("/image").body(exact).send().unwrap();
2502 assert_eq!(res.status(), 400);
2503 }
2504
2505 #[test]
2506 fn bind_refuses_an_empty_token() {
2507 let engine: Arc<dyn Engine> = Arc::new(SyntheticEngine::new());
2508 let catalog = Arc::new(Mutex::new(seeded_catalog()));
2509 let err = LocalApi::bind(
2510 "127.0.0.1:0",
2511 engine,
2512 catalog.clone(),
2513 None,
2514 WorkerObservers::default(),
2515 String::new(),
2516 JobGate::new(),
2517 None,
2518 ModelServices::new(test_host(&catalog)),
2519 )
2520 .err()
2521 .expect("empty token must be refused")
2522 .to_string();
2523 assert!(err.contains("empty token"), "got: {err}");
2524 }
2525
2526 #[test]
2527 fn post_image_returns_503_when_the_shared_gate_is_held() {
2528 let gate = JobGate::new();
2532 let h = Harness::start_with_gate(seeded_catalog(), gate.clone());
2533 let reservation = gate.try_reserve().expect("pre-hold the slot");
2534 let res = h
2535 .post("/image")
2536 .json(&serde_json::json!({ "prompt": "x" }))
2537 .send()
2538 .unwrap();
2539 assert_eq!(res.status(), 503);
2540 assert_eq!(res.headers()["retry-after"], "2");
2541
2542 drop(reservation);
2545 let res = h
2546 .post("/image")
2547 .json(&serde_json::json!({ "prompt": "x" }))
2548 .send()
2549 .unwrap();
2550 assert_eq!(res.status(), 200);
2551 }
2552
2553 #[test]
2558 fn host_is_loopback_accepts_only_loopback_shapes() {
2559 for ok in [
2560 "127.0.0.1",
2561 "127.0.0.1:4787",
2562 "localhost",
2563 "LOCALHOST:80",
2564 "[::1]",
2565 "[::1]:4787",
2566 ] {
2567 assert!(host_is_loopback(ok), "{ok} should be loopback");
2568 }
2569 for bad in [
2570 "evil.example",
2571 "evil.example:4787",
2572 "127.0.0.1.evil.example",
2573 "192.168.1.10:4787",
2574 "[::2]:4787",
2575 "[::1",
2576 "",
2577 ] {
2578 assert!(!host_is_loopback(bad), "{bad} should be rejected");
2579 }
2580 }
2581
2582 #[test]
2583 fn origin_is_loopback_accepts_only_loopback_origins() {
2584 for ok in [
2585 "http://127.0.0.1:4787",
2586 "http://localhost:5173",
2587 "https://localhost",
2588 "http://[::1]:3000",
2589 ] {
2590 assert!(origin_is_loopback(ok), "{ok} should be allowed");
2591 }
2592 for bad in [
2593 "https://evil.example",
2594 "http://192.168.1.10",
2595 "null",
2596 "file://",
2597 "chrome-extension://abc",
2598 "",
2599 ] {
2600 assert!(!origin_is_loopback(bad), "{bad} should be rejected");
2601 }
2602 }
2603
2604 #[test]
2605 fn deny_reason_orders_host_origin_then_token() {
2606 let t = "tok";
2607 assert!(matches!(
2609 deny_reason(Some("evil.example"), Some("https://evil.example"), None, t),
2610 Some(Denial::Host(_))
2611 ));
2612 assert!(matches!(
2614 deny_reason(Some("127.0.0.1"), Some("https://evil.example"), None, t),
2615 Some(Denial::Origin(_))
2616 ));
2617 assert_eq!(
2619 deny_reason(Some("127.0.0.1"), None, None, t),
2620 Some(Denial::Token)
2621 );
2622 assert_eq!(
2624 deny_reason(None, None, Some("Basic dXNlcjpwdw=="), t),
2625 Some(Denial::Token)
2626 );
2627 assert_eq!(deny_reason(None, None, Some("Bearer tok"), t), None);
2629 assert_eq!(deny_reason(None, None, Some("bearer tok"), t), None);
2631 }
2632
2633 #[test]
2634 fn token_matches_is_exact() {
2635 assert!(token_matches("abc", "abc"));
2636 assert!(!token_matches("abd", "abc"));
2637 assert!(!token_matches("ab", "abc"));
2638 assert!(!token_matches("", "abc"));
2639 }
2640
2641 #[test]
2646 fn discovery_file_round_trips_and_is_owner_only() {
2647 let dir = tempfile::tempdir().unwrap();
2648 let path = dir.path().join("local-api.json");
2649 write_discovery_file(&path, "http://127.0.0.1:4787", "tok-123").unwrap();
2650
2651 let parsed: serde_json::Value =
2652 serde_json::from_str(&std::fs::read_to_string(&path).unwrap()).unwrap();
2653 assert_eq!(parsed["url"], "http://127.0.0.1:4787");
2654 assert_eq!(parsed["token"], "tok-123");
2655
2656 #[cfg(unix)]
2657 {
2658 use std::os::unix::fs::PermissionsExt;
2659 let mode = std::fs::metadata(&path).unwrap().permissions().mode();
2660 assert_eq!(
2661 mode & 0o077,
2662 0,
2663 "discovery file carries the token and must be owner-only, got {mode:o}"
2664 );
2665 }
2666
2667 remove_discovery_file(&path);
2668 assert!(!path.exists());
2669 remove_discovery_file(&path);
2671 }
2672}