use std::net::SocketAddr;
use std::path::PathBuf;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::Duration;
use parking_lot::Mutex;
use serde::Deserialize;
use tiny_http::{Header, Method, Request, Response, Server};
use crate::catalog::{Catalog, CatalogModel};
use crate::engine::Engine;
use crate::host::{HostError, ModelHost, ModelStatus};
use crate::job_gate::JobGate;
use crate::lifecycle::ModelState;
use crate::local::{chat_on_lane, run_image, run_kind, LocalError, LocalImageRequest};
use crate::runtime::{JobOutcome, WorkerObservers};
use crate::stt_stream::tokens::StreamTokens;
use crate::types::{
AudioSttParams, AudioTtsParams, ChatMessage, LlmParams, Task, TaskKind, TaskResult, VideoParams,
};
const TRACE_TARGET: &str = "studio_worker::local_api";
const POLL: Duration = Duration::from_millis(200);
pub const MAX_BODY_BYTES: usize = 1024 * 1024;
#[derive(Debug, PartialEq, Eq)]
enum Denial {
Host(String),
Origin(String),
Token,
}
fn host_is_loopback(host: &str) -> bool {
let bare = if let Some(rest) = host.strip_prefix('[') {
match rest.split_once(']') {
Some((addr, _port)) => addr,
None => return false,
}
} else {
host.rsplit_once(':').map(|(h, _)| h).unwrap_or(host)
};
bare.eq_ignore_ascii_case("localhost") || bare == "127.0.0.1" || bare == "::1"
}
fn origin_is_loopback(origin: &str) -> bool {
let rest = origin
.strip_prefix("https://")
.or_else(|| origin.strip_prefix("http://"));
match rest {
Some(host) => host_is_loopback(host),
None => false,
}
}
fn token_matches(presented: &str, expected: &str) -> bool {
let (a, b) = (presented.as_bytes(), expected.as_bytes());
if a.len() != b.len() {
return false;
}
a.iter().zip(b).fold(0u8, |acc, (x, y)| acc | (x ^ y)) == 0
}
fn deny_reason(
host: Option<&str>,
origin: Option<&str>,
authorization: Option<&str>,
token: &str,
) -> Option<Denial> {
if let Some(host) = host {
if !host_is_loopback(host) {
return Some(Denial::Host(host.to_string()));
}
}
if let Some(origin) = origin {
if !origin_is_loopback(origin) {
return Some(Denial::Origin(origin.to_string()));
}
}
let presented = authorization.and_then(|a| {
a.strip_prefix("Bearer ")
.or_else(|| a.strip_prefix("bearer "))
});
match presented {
Some(presented) if token_matches(presented, token) => None,
_ => Some(Denial::Token),
}
}
pub struct LocalApi {
engine: Arc<dyn Engine>,
catalog: Arc<Mutex<Catalog>>,
catalog_path: Option<PathBuf>,
observers: WorkerObservers,
server: Server,
addr: SocketAddr,
token: String,
gate: JobGate,
models_root: Option<PathBuf>,
services: ModelServices,
control: Option<crate::control::DaemonControl>,
}
#[derive(Clone)]
pub struct ModelServices {
pub host: ModelHost,
pub tokens: Arc<StreamTokens>,
pub stream_port: Arc<std::sync::atomic::AtomicU16>,
}
impl ModelServices {
pub fn new(host: ModelHost) -> Self {
Self {
host,
tokens: Arc::new(StreamTokens::default()),
stream_port: Arc::new(std::sync::atomic::AtomicU16::new(0)),
}
}
}
#[derive(Deserialize)]
#[serde(rename_all = "camelCase")]
struct StreamTokenBody {
model: String,
#[serde(default)]
ttl_secs: Option<i64>,
}
const DEFAULT_STREAM_TOKEN_TTL_SECS: i64 = 600;
#[derive(Deserialize)]
#[serde(rename_all = "camelCase")]
struct ImageBody {
prompt: String,
#[serde(default)]
model: Option<String>,
#[serde(default)]
negative_prompt: Option<String>,
#[serde(default)]
width: Option<u32>,
#[serde(default)]
height: Option<u32>,
#[serde(default)]
steps: Option<u32>,
#[serde(default)]
seed: Option<u64>,
#[serde(default)]
ext: Option<String>,
}
#[derive(Deserialize)]
struct ChatBody {
#[serde(default)]
model: Option<String>,
messages: Vec<ChatMessageBody>,
#[serde(default)]
max_tokens: Option<u32>,
#[serde(default)]
temperature: Option<f32>,
#[serde(default)]
top_p: Option<f32>,
#[serde(default)]
stop: Option<Vec<String>>,
#[serde(default)]
chat_template_kwargs: Option<serde_json::Map<String, serde_json::Value>>,
}
#[derive(Deserialize)]
struct ChatMessageBody {
role: String,
content: String,
}
#[derive(Deserialize)]
#[serde(rename_all = "camelCase")]
struct TtsBody {
text: String,
#[serde(default)]
model: Option<String>,
#[serde(default)]
voice: Option<String>,
#[serde(default)]
speed: Option<f32>,
#[serde(default)]
language: Option<String>,
#[serde(default)]
ext: Option<String>,
}
#[derive(Deserialize)]
#[serde(rename_all = "camelCase")]
struct SttBody {
input_url: String,
#[serde(default)]
model: Option<String>,
#[serde(default)]
language: Option<String>,
}
#[derive(Deserialize)]
#[serde(rename_all = "camelCase")]
struct VideoBody {
prompt: String,
#[serde(default)]
model: Option<String>,
#[serde(default)]
negative_prompt: Option<String>,
#[serde(default)]
seconds: Option<f32>,
#[serde(default)]
width: Option<u32>,
#[serde(default)]
height: Option<u32>,
#[serde(default)]
ext: Option<String>,
}
impl LocalApi {
#[allow(clippy::too_many_arguments)]
pub fn bind(
addr: &str,
engine: Arc<dyn Engine>,
catalog: Arc<Mutex<Catalog>>,
catalog_path: Option<PathBuf>,
observers: WorkerObservers,
token: String,
gate: JobGate,
models_root: Option<PathBuf>,
services: ModelServices,
) -> anyhow::Result<Self> {
anyhow::ensure!(
!token.is_empty(),
"local api: refusing to serve with an empty token"
);
let server =
Server::http(addr).map_err(|e| anyhow::anyhow!("local api bind {addr}: {e}"))?;
let addr = server
.server_addr()
.to_ip()
.ok_or_else(|| anyhow::anyhow!("local api: non-ip listen address"))?;
Ok(Self {
engine,
catalog,
catalog_path,
observers,
server,
addr,
token,
gate,
models_root,
services,
control: None,
})
}
pub fn with_control(mut self, control: crate::control::DaemonControl) -> Self {
self.control = Some(control);
self
}
pub fn local_addr(&self) -> SocketAddr {
self.addr
}
pub fn url(&self) -> String {
format!("http://{}", self.addr)
}
pub fn serve(&self, stop: &AtomicBool) {
const WORKERS: usize = 8;
std::thread::scope(|scope| {
for _ in 0..WORKERS {
scope.spawn(|| {
while !stop.load(Ordering::Relaxed) {
match self.server.recv_timeout(POLL) {
Ok(Some(request)) => self.route(request),
Ok(None) => {}
Err(err) => {
tracing::warn!(target: TRACE_TARGET, error = %err, "local api recv error");
break;
}
}
}
});
}
});
}
fn route(&self, request: Request) {
let method = request.method().clone();
let url = request.url().to_string();
let path = url.split('?').next().unwrap_or("/");
if !(method == Method::Get && path == "/healthz") {
let header = |name: &'static str| {
request
.headers()
.iter()
.find(|h| h.field.equiv(name))
.map(|h| h.value.as_str().to_string())
};
let denial = deny_reason(
header("host").as_deref(),
header("origin").as_deref(),
header("authorization").as_deref(),
&self.token,
);
if let Some(denial) = denial {
let (status, body) = match &denial {
Denial::Host(host) => {
(403, format!("forbidden: non-loopback Host header {host:?}"))
}
Denial::Origin(origin) => (
403,
format!("forbidden: cross-site request from Origin {origin:?}"),
),
Denial::Token => (
401,
"missing or invalid Authorization bearer token; local clients \
can read the current token from the local-api.json discovery \
file in the worker's config directory"
.to_string(),
),
};
tracing::warn!(
target: TRACE_TARGET,
op = "deny",
method = %method,
path,
status,
reason = ?denial,
"local api request denied"
);
if let Err(err) = respond(request, status, "text/plain", body.as_bytes()) {
tracing::warn!(target: TRACE_TARGET, error = %err, "local api respond error");
}
return;
}
}
let outcome = match (&method, path) {
(Method::Get, "/healthz") => self.handle_healthz(request),
(Method::Post, "/image") => self.handle_image(request),
(Method::Post, "/v1/chat/completions") => self.handle_chat(request),
(Method::Post, "/tts") => self.handle_tts(request),
(Method::Post, "/stt") => self.handle_stt(request),
(Method::Post, "/video") => self.handle_video(request),
(Method::Get, "/models") => self.handle_list_models(request),
(Method::Post, "/models") => self.handle_add_model(request),
(Method::Get, "/jobs") => self.handle_jobs(request),
(Method::Post, "/stream-tokens") => self.handle_stream_token(request),
(_, p) if p.starts_with("/daemon/") => self.handle_daemon(request, &method, &url),
(Method::Get, p) if job_route(p, "/log").is_some() => {
let id = job_route(p, "/log").unwrap_or_default().to_string();
self.handle_job_log(request, &id)
}
(Method::Get, p) if job_route(p, "/thumbnail").is_some() => {
let id = job_route(p, "/thumbnail").unwrap_or_default().to_string();
self.handle_job_thumbnail(request, &id)
}
(Method::Get, p) if lifecycle_route(p, "/state").is_some() => {
let id = lifecycle_route(p, "/state").unwrap_or_default().to_string();
self.respond_lifecycle(request, self.services.host.status(&id), 200)
}
(Method::Post, p) if lifecycle_route(p, "/load").is_some() => {
let id = lifecycle_route(p, "/load").unwrap_or_default().to_string();
self.respond_lifecycle(request, self.services.host.load(&id), 202)
}
(Method::Post, p) if lifecycle_route(p, "/unload").is_some() => {
let id = lifecycle_route(p, "/unload")
.unwrap_or_default()
.to_string();
self.respond_lifecycle(request, self.services.host.unload(&id), 202)
}
(Method::Delete, p) if p.starts_with("/models/") => {
let id = p.trim_start_matches("/models/").to_string();
self.handle_delete_model(request, &id)
}
_ => respond(request, 404, "text/plain", b"not found"),
};
if let Err(err) = outcome {
tracing::warn!(target: TRACE_TARGET, error = %err, "local api respond error");
}
}
fn handle_healthz(&self, request: Request) -> std::io::Result<()> {
let free_bytes = self
.models_root
.as_deref()
.and_then(|root| fs4::available_space(root).ok());
let gpu = self
.observers
.gpu_runtime
.lock()
.clone()
.map(|g| serde_json::json!({ "ok": g.ok, "detail": g.detail }));
let body = serde_json::json!({
"ok": true,
"version": crate::AGENT_VERSION,
"busy": self.gate.is_busy(),
"engine": self.engine.name(),
"modelsRoot": self.models_root.as_ref().map(|p| p.display().to_string()),
"modelsRootFreeBytes": free_bytes,
"gpuRuntime": gpu,
});
match serde_json::to_vec(&body) {
Ok(bytes) => respond(request, 200, "application/json", &bytes),
Err(_) => respond(request, 200, "application/json", b"{\"ok\":true}"),
}
}
fn handle_image(&self, mut request: Request) -> std::io::Result<()> {
let body = match read_body(&mut request)? {
BodyOutcome::Ok(body) => body,
BodyOutcome::TooLarge => return respond_too_large(request),
};
let parsed: ImageBody = match serde_json::from_str(&body) {
Ok(parsed) => parsed,
Err(err) => {
return respond(
request,
400,
"text/plain",
format!("bad json: {err}").as_bytes(),
)
}
};
let req = LocalImageRequest {
prompt: parsed.prompt,
model: parsed.model,
negative_prompt: parsed.negative_prompt,
width: parsed.width,
height: parsed.height,
steps: parsed.steps,
seed: parsed.seed,
ext: parsed.ext,
};
let Some(_reservation) = self.gate.try_reserve() else {
return respond_busy(request);
};
let catalog = self.catalog.lock().clone();
match run_image(self.engine.as_ref(), &catalog, &self.observers, &req) {
Ok(TaskResult::Image { bytes, ext }) => {
respond(request, 200, content_type_for(&ext), &bytes)
}
Ok(_) => respond(request, 500, "text/plain", b"unexpected non-image result"),
Err(err) => respond_local_err(request, err),
}
}
fn handle_chat(&self, mut request: Request) -> std::io::Result<()> {
let body = match read_body(&mut request)? {
BodyOutcome::Ok(body) => body,
BodyOutcome::TooLarge => return respond_too_large(request),
};
let parsed: ChatBody = match serde_json::from_str(&body) {
Ok(p) => p,
Err(err) => {
return respond(
request,
400,
"text/plain",
format!("bad json: {err}").as_bytes(),
)
}
};
let prompt_preview = parsed
.messages
.last()
.map(|m| m.content.clone())
.unwrap_or_default();
let params = LlmParams {
messages: parsed
.messages
.into_iter()
.map(|m| ChatMessage {
role: m.role,
content: m.content,
})
.collect(),
max_tokens: parsed.max_tokens.unwrap_or(512),
temperature: parsed.temperature.unwrap_or(0.7),
top_p: parsed.top_p,
stop: parsed.stop,
chat_template_kwargs: parsed.chat_template_kwargs,
..Default::default()
};
let catalog = self.catalog.lock().clone();
if let Some(result) = chat_on_lane(
&self.services.host,
&catalog,
&self.observers,
parsed.model.as_deref(),
&prompt_preview,
params.clone(),
) {
return respond_llm(request, result);
}
let Some(_reservation) = self.gate.try_reserve() else {
return respond_busy(request);
};
let outcome = run_kind(
self.engine.as_ref(),
&catalog,
&self.observers,
TaskKind::Llm,
parsed.model.as_deref(),
&prompt_preview,
Task::Llm(params),
);
respond_llm(request, outcome)
}
fn handle_tts(&self, mut request: Request) -> std::io::Result<()> {
let body = match read_body(&mut request)? {
BodyOutcome::Ok(body) => body,
BodyOutcome::TooLarge => return respond_too_large(request),
};
let parsed: TtsBody = match serde_json::from_str(&body) {
Ok(p) => p,
Err(err) => {
return respond(
request,
400,
"text/plain",
format!("bad json: {err}").as_bytes(),
)
}
};
let preview = parsed.text.clone();
let params = AudioTtsParams {
text: parsed.text,
voice: parsed.voice.unwrap_or_else(|| "default".into()),
speed: parsed.speed,
language: parsed.language,
ext: parsed.ext.unwrap_or_else(|| "wav".into()),
};
let Some(_reservation) = self.gate.try_reserve() else {
return respond_busy(request);
};
let catalog = self.catalog.lock().clone();
match run_kind(
self.engine.as_ref(),
&catalog,
&self.observers,
TaskKind::AudioTts,
parsed.model.as_deref(),
&preview,
Task::AudioTts(params),
) {
Ok(TaskResult::AudioTts { bytes, ext }) => {
respond(request, 200, content_type_for(&ext), &bytes)
}
Ok(_) => respond(request, 500, "text/plain", b"unexpected non-audio result"),
Err(err) => respond_local_err(request, err),
}
}
fn handle_stt(&self, mut request: Request) -> std::io::Result<()> {
let body = match read_body(&mut request)? {
BodyOutcome::Ok(body) => body,
BodyOutcome::TooLarge => return respond_too_large(request),
};
let parsed: SttBody = match serde_json::from_str(&body) {
Ok(p) => p,
Err(err) => {
return respond(
request,
400,
"text/plain",
format!("bad json: {err}").as_bytes(),
)
}
};
let preview = parsed.input_url.clone();
let params = AudioSttParams {
input_url: parsed.input_url,
language: parsed.language,
..Default::default()
};
let Some(_reservation) = self.gate.try_reserve() else {
return respond_busy(request);
};
let catalog = self.catalog.lock().clone();
match run_kind(
self.engine.as_ref(),
&catalog,
&self.observers,
TaskKind::AudioStt,
parsed.model.as_deref(),
&preview,
Task::AudioStt(params),
) {
Ok(TaskResult::AudioStt { json }) => match serde_json::to_vec(&json) {
Ok(bytes) => respond(request, 200, "application/json", &bytes),
Err(e) => respond(request, 500, "text/plain", e.to_string().as_bytes()),
},
Ok(_) => respond(
request,
500,
"text/plain",
b"unexpected non-transcript result",
),
Err(err) => respond_local_err(request, err),
}
}
fn handle_video(&self, mut request: Request) -> std::io::Result<()> {
let body = match read_body(&mut request)? {
BodyOutcome::Ok(body) => body,
BodyOutcome::TooLarge => return respond_too_large(request),
};
let parsed: VideoBody = match serde_json::from_str(&body) {
Ok(p) => p,
Err(err) => {
return respond(
request,
400,
"text/plain",
format!("bad json: {err}").as_bytes(),
)
}
};
let preview = parsed.prompt.clone();
let params = VideoParams {
prompt: parsed.prompt,
negative_prompt: parsed.negative_prompt,
seconds: parsed.seconds.unwrap_or(2.0),
width: parsed.width.unwrap_or(256),
height: parsed.height.unwrap_or(256),
ext: parsed.ext.unwrap_or_else(|| "mp4".into()),
..Default::default()
};
let Some(_reservation) = self.gate.try_reserve() else {
return respond_busy(request);
};
let catalog = self.catalog.lock().clone();
match run_kind(
self.engine.as_ref(),
&catalog,
&self.observers,
TaskKind::Video,
parsed.model.as_deref(),
&preview,
Task::Video(params),
) {
Ok(TaskResult::Video { bytes, ext }) => {
respond(request, 200, content_type_for(&ext), &bytes)
}
Ok(_) => respond(request, 500, "text/plain", b"unexpected non-video result"),
Err(err) => respond_local_err(request, err),
}
}
fn handle_list_models(&self, request: Request) -> std::io::Result<()> {
let models = self.catalog.lock().models.clone();
let statuses = self.services.host.statuses();
let listed: Vec<serde_json::Value> = models
.iter()
.map(|model| {
let mut value = serde_json::to_value(model).unwrap_or_default();
if let (Some(obj), Some(status)) = (
value.as_object_mut(),
statuses.iter().find(|s| s.id == model.id),
) {
obj.insert("state".into(), status.state.name().into());
obj.insert("resident".into(), status.resident.into());
obj.insert("since".into(), status.since.to_rfc3339().into());
obj.insert("loadable".into(), self.services.host.can_load(model).into());
if let ModelState::Failed { reason } = &status.state {
obj.insert("error".into(), reason.clone().into());
}
}
value
})
.collect();
match serde_json::to_vec(&listed) {
Ok(body) => respond(request, 200, "application/json", &body),
Err(err) => respond(request, 500, "text/plain", err.to_string().as_bytes()),
}
}
fn handle_add_model(&self, mut request: Request) -> std::io::Result<()> {
let body = match read_body(&mut request)? {
BodyOutcome::Ok(body) => body,
BodyOutcome::TooLarge => return respond_too_large(request),
};
let model: CatalogModel = match serde_json::from_str(&body) {
Ok(model) => model,
Err(err) => {
return respond(
request,
400,
"text/plain",
format!("bad model: {err}").as_bytes(),
)
}
};
let saved = {
let mut catalog = self.catalog.lock();
catalog.upsert(model);
self.persist(&catalog)
};
match saved {
Ok(()) => respond(request, 200, "application/json", b"{\"ok\":true}"),
Err(err) => respond(request, 500, "text/plain", err.to_string().as_bytes()),
}
}
fn handle_delete_model(&self, request: Request, id: &str) -> std::io::Result<()> {
if let Err(err) = self.services.host.unload(id) {
if !matches!(err, HostError::UnknownModel(_)) {
return respond_json(
request,
500,
&serde_json::json!({ "error": "unload_failed", "message": err.to_string() }),
);
}
}
let (existed, saved) = {
let mut catalog = self.catalog.lock();
let existed = catalog.remove(id);
(existed, self.persist(&catalog))
};
if !existed {
return respond(request, 404, "text/plain", b"no such model");
}
match saved {
Ok(()) => respond(request, 200, "application/json", b"{\"ok\":true}"),
Err(err) => respond(request, 500, "text/plain", err.to_string().as_bytes()),
}
}
fn respond_lifecycle(
&self,
request: Request,
outcome: Result<ModelStatus, HostError>,
pending_status: u16,
) -> std::io::Result<()> {
match outcome {
Ok(status) => {
let code = match status.state {
ModelState::Loading | ModelState::Unloading => pending_status,
_ => 200,
};
respond_json(request, code, &status_json(&status))
}
Err(err) => {
let (code, body) = match &err {
HostError::UnknownModel(_) => {
(404, serde_json::json!({ "error": "unknown_model" }))
}
HostError::Disabled(_) => {
(400, serde_json::json!({ "error": "model_disabled" }))
}
HostError::Refused(r) => (
409,
serde_json::json!({
"error": "insufficient_memory",
"neededGib": r.needed_gib,
"freeGib": r.free_gib,
"marginGib": r.margin_gib,
}),
),
HostError::NotLoaded { state, .. } => (
409,
serde_json::json!({ "error": "model_not_loaded", "state": state }),
),
HostError::Persist(_) => {
(500, serde_json::json!({ "error": "residency_not_saved" }))
}
HostError::LaneBusy(_) => (409, serde_json::json!({ "error": "model_busy" })),
};
let mut body = body;
body["message"] = err.to_string().into();
tracing::warn!(
target: TRACE_TARGET,
op = "lifecycle",
status = code,
error = %err,
"lifecycle request refused"
);
respond_json(request, code, &body)
}
}
}
fn handle_stream_token(&self, mut request: Request) -> std::io::Result<()> {
let body = match read_body(&mut request)? {
BodyOutcome::Ok(body) => body,
BodyOutcome::TooLarge => return respond_too_large(request),
};
let parsed: StreamTokenBody = match serde_json::from_str(&body) {
Ok(p) => p,
Err(err) => {
return respond_json(
request,
400,
&serde_json::json!({ "error": "bad_request", "message": err.to_string() }),
)
}
};
let model = self.catalog.lock().get(&parsed.model).cloned();
let Some(model) = model else {
return respond_json(
request,
404,
&serde_json::json!({ "error": "unknown_model" }),
);
};
if model.source.engine != crate::types::ModelEngine::Parakeet {
return respond_json(
request,
400,
&serde_json::json!({ "error": "not_a_stream_model" }),
);
}
let port = self.services.stream_port.load(Ordering::SeqCst);
if port == 0 {
return respond_json(
request,
503,
&serde_json::json!({ "error": "stream_listener_down" }),
);
}
let ttl =
chrono::Duration::seconds(parsed.ttl_secs.unwrap_or(DEFAULT_STREAM_TOKEN_TTL_SECS));
let grant = self
.services
.tokens
.mint(&model.id, ttl, chrono::Utc::now());
tracing::info!(
target: TRACE_TARGET,
op = "stream_token",
model = %model.id,
expires_at = %grant.expires_at,
"stream token minted"
);
respond_json(
request,
200,
&serde_json::json!({
"token": grant.token,
"model": grant.model,
"expiresAt": grant.expires_at.to_rfc3339(),
"port": port,
"path": crate::stt_stream::server::STREAM_PATH,
}),
)
}
fn handle_jobs(&self, request: Request) -> std::io::Result<()> {
let jobs: Vec<serde_json::Value> = self
.observers
.local_jobs
.lock()
.iter()
.map(|job| {
let (status, reason) = match &job.outcome {
JobOutcome::Completed => ("completed", None),
JobOutcome::Failed { reason } => ("failed", Some(reason.clone())),
};
serde_json::json!({
"jobId": job.job_id,
"kind": job.kind.as_str(),
"model": job.model,
"prompt": job.prompt,
"status": status,
"reason": reason,
"startedAt": job.started_at.to_rfc3339(),
"finishedAt": job.finished_at.to_rfc3339(),
})
})
.collect();
match serde_json::to_vec(&jobs) {
Ok(body) => respond(request, 200, "application/json", &body),
Err(err) => respond(request, 500, "text/plain", err.to_string().as_bytes()),
}
}
fn handle_daemon(
&self,
mut request: Request,
method: &Method,
url: &str,
) -> std::io::Result<()> {
let Some(control) = &self.control else {
return respond_json(
request,
503,
&serde_json::json!({ "error": "daemon_control_unavailable" }),
);
};
let path = url.split('?').next().unwrap_or("/");
match (method, path) {
(Method::Get, "/daemon/status") => {
let status = control.status(&self.observers, self.gate.is_busy());
respond_serialised(request, 200, &status)
}
(Method::Get, "/daemon/logs") => {
let after = query_param(url, "after")
.and_then(|v| v.parse::<u64>().ok())
.unwrap_or(0);
let (entries, seq) = crate::runtime::recent_logs_after(&self.observers, after);
respond_serialised(request, 200, &crate::daemon_api::LogsPage { entries, seq })
}
(Method::Post, "/daemon/pause") => {
let paused = control.set_paused(true);
respond_json(request, 200, &serde_json::json!({ "paused": paused }))
}
(Method::Post, "/daemon/resume") => {
let paused = control.set_paused(false);
respond_json(request, 200, &serde_json::json!({ "paused": paused }))
}
(Method::Get, "/daemon/config") => {
respond_serialised(request, 200, &control.editable_config())
}
(Method::Put, "/daemon/config") => {
let body = match read_body(&mut request)? {
BodyOutcome::Ok(body) => body,
BodyOutcome::TooLarge => return respond_too_large(request),
};
let edit: crate::daemon_api::EditableConfig = match serde_json::from_str(&body) {
Ok(edit) => edit,
Err(err) => {
return respond_error(request, 400, "bad_request", &err.to_string())
}
};
match control.update_config(edit) {
Ok(saved) => respond_serialised(request, 200, &saved),
Err(err @ crate::control::ControlError::Invalid(_)) => {
respond_error(request, 400, "invalid_config", &err.to_string())
}
Err(err) => respond_error(request, 500, "config_not_saved", &err.to_string()),
}
}
(Method::Post, "/daemon/registration/reset") => {
match control.request_registration_reset() {
Ok(()) => respond_json(request, 202, &serde_json::json!({ "ok": true })),
Err(err) => respond_error(request, 409, "not_rejected", &err.to_string()),
}
}
(Method::Post, "/daemon/shutdown") => {
control.shutdown();
respond_json(request, 202, &serde_json::json!({ "ok": true }))
}
_ => respond_error(request, 404, "not_found", "no such daemon route"),
}
}
fn handle_job_log(&self, request: Request, id: &str) -> std::io::Result<()> {
match crate::job_log::global().get(id) {
Some(log) => respond_serialised(request, 200, &log),
None => respond_error(request, 404, "unknown_job", "no log captured for that job"),
}
}
fn handle_job_thumbnail(&self, request: Request, id: &str) -> std::io::Result<()> {
match self.observers.thumbnails.get(id) {
Some(png) => respond(request, 200, "image/png", &png),
None => respond_error(request, 404, "no_thumbnail", "no thumbnail for that job"),
}
}
fn persist(&self, catalog: &Catalog) -> std::io::Result<()> {
match &self.catalog_path {
Some(path) => catalog.save(path),
None => Ok(()),
}
}
}
enum BodyOutcome {
Ok(String),
TooLarge,
}
fn read_body(request: &mut Request) -> std::io::Result<BodyOutcome> {
if matches!(request.body_length(), Some(len) if len > MAX_BODY_BYTES) {
return Ok(BodyOutcome::TooLarge);
}
let mut body = String::new();
use std::io::Read as _;
request
.as_reader()
.take(MAX_BODY_BYTES as u64 + 1)
.read_to_string(&mut body)?;
if body.len() > MAX_BODY_BYTES {
return Ok(BodyOutcome::TooLarge);
}
Ok(BodyOutcome::Ok(body))
}
fn respond_too_large(request: Request) -> std::io::Result<()> {
respond(
request,
413,
"text/plain",
format!("request body exceeds {MAX_BODY_BYTES} bytes").as_bytes(),
)
}
fn respond_busy(request: Request) -> std::io::Result<()> {
let retry = Header::from_bytes(b"Retry-After".as_slice(), b"2".as_slice())
.expect("static Retry-After header is valid");
let response = Response::from_data(
b"worker is busy with another job (studio or local); retry shortly".to_vec(),
)
.with_status_code(503)
.with_header(retry)
.with_header(
Header::from_bytes(b"Content-Type".as_slice(), b"text/plain".as_slice())
.expect("static content-type header is valid"),
);
request.respond(response)
}
pub fn write_discovery_file(path: &std::path::Path, url: &str, token: &str) -> anyhow::Result<()> {
let body = serde_json::json!({ "url": url, "token": token });
let text = serde_json::to_string_pretty(&body)?;
crate::config::write_atomic(path, text.as_bytes())?;
tracing::info!(
target: TRACE_TARGET,
op = "discovery",
path = %path.display(),
url,
"local api discovery file written"
);
Ok(())
}
pub fn remove_discovery_file(path: &std::path::Path) {
if let Err(e) = std::fs::remove_file(path) {
if e.kind() != std::io::ErrorKind::NotFound {
tracing::warn!(
target: TRACE_TARGET,
op = "discovery",
path = %path.display(),
error = %e,
"failed to remove local api discovery file"
);
}
}
}
fn content_type_for(ext: &str) -> &'static str {
match ext.to_ascii_lowercase().as_str() {
"webp" => "image/webp",
"png" => "image/png",
"jpg" | "jpeg" => "image/jpeg",
"gif" => "image/gif",
"wav" => "audio/wav",
"mp3" => "audio/mpeg",
"ogg" | "opus" => "audio/ogg",
"flac" => "audio/flac",
"mp4" => "video/mp4",
"webm" => "video/webm",
_ => "application/octet-stream",
}
}
fn respond_local_err(request: Request, err: LocalError) -> std::io::Result<()> {
let status = match err {
LocalError::Engine(_) => 500,
_ => 400,
};
respond(request, status, "text/plain", err.to_string().as_bytes())
}
fn respond_llm(request: Request, outcome: Result<TaskResult, LocalError>) -> std::io::Result<()> {
match outcome {
Ok(TaskResult::Llm { json }) => match serde_json::to_vec(&json) {
Ok(bytes) => respond(request, 200, "application/json", &bytes),
Err(e) => respond(request, 500, "text/plain", e.to_string().as_bytes()),
},
Ok(_) => respond(request, 500, "text/plain", b"unexpected non-llm result"),
Err(err) => respond_local_err(request, err),
}
}
fn lifecycle_route<'a>(path: &'a str, suffix: &str) -> Option<&'a str> {
let id = path.strip_prefix("/models/")?.strip_suffix(suffix)?;
(!id.is_empty() && !id.contains('/')).then_some(id)
}
fn job_route<'a>(path: &'a str, suffix: &str) -> Option<&'a str> {
let id = path.strip_prefix("/jobs/")?.strip_suffix(suffix)?;
(!id.is_empty() && !id.contains('/')).then_some(id)
}
fn query_param<'a>(url: &'a str, name: &str) -> Option<&'a str> {
url.split_once('?')?
.1
.split('&')
.find_map(|pair| pair.strip_prefix(name)?.strip_prefix('='))
}
fn status_json(status: &ModelStatus) -> serde_json::Value {
let mut body = serde_json::json!({
"id": status.id,
"state": status.state.name(),
"resident": status.resident,
"since": status.since.to_rfc3339(),
});
if let ModelState::Failed { reason } = &status.state {
body["error"] = reason.clone().into();
}
body
}
fn respond_serialised<T: serde::Serialize>(
request: Request,
status: u16,
body: &T,
) -> std::io::Result<()> {
match serde_json::to_vec(body) {
Ok(bytes) => respond(request, status, "application/json", &bytes),
Err(err) => respond(request, 500, "text/plain", err.to_string().as_bytes()),
}
}
fn respond_error(request: Request, status: u16, code: &str, message: &str) -> std::io::Result<()> {
respond_serialised(
request,
status,
&crate::daemon_api::ErrorBody {
error: code.to_string(),
message: Some(message.to_string()),
},
)
}
fn respond_json(request: Request, status: u16, body: &serde_json::Value) -> std::io::Result<()> {
let bytes = serde_json::to_vec(body).unwrap_or_else(|_| b"{}".to_vec());
respond(request, status, "application/json", &bytes)
}
fn respond(request: Request, status: u16, content_type: &str, body: &[u8]) -> std::io::Result<()> {
let header = Header::from_bytes(b"Content-Type".as_slice(), content_type.as_bytes())
.expect("static content-type header is valid");
let response = Response::from_data(body)
.with_status_code(status)
.with_header(header);
request.respond(response)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::catalog::CatalogModel;
use crate::engine::{EngineCapabilities, SyntheticEngine};
use crate::types::{ModelCliDefaults, ModelEngine, ModelSource, Task, TaskKind};
struct SlowEngine {
inner: SyntheticEngine,
delay: std::time::Duration,
}
impl Engine for SlowEngine {
fn name(&self) -> &'static str {
"slow"
}
fn capabilities(&self) -> EngineCapabilities {
self.inner.capabilities()
}
fn dispatch(&self, model: &str, task: Task) -> anyhow::Result<TaskResult> {
std::thread::sleep(self.delay);
self.inner.dispatch(model, task)
}
}
fn synthetic_model_of(id: &str, kind: TaskKind) -> CatalogModel {
CatalogModel {
kind,
..synthetic_model(id)
}
}
fn multi_kind_catalog() -> Catalog {
Catalog {
models: vec![
synthetic_model_of("img", TaskKind::Image),
synthetic_model_of("chat", TaskKind::Llm),
synthetic_model_of("tts", TaskKind::AudioTts),
synthetic_model_of("stt", TaskKind::AudioStt),
synthetic_model_of("vid", TaskKind::Video),
],
..Default::default()
}
}
fn synthetic_model(id: &str) -> CatalogModel {
CatalogModel {
id: id.into(),
display_name: id.into(),
kind: TaskKind::Image,
vram_gb_estimate: 0.0,
description: None,
source: ModelSource {
engine: ModelEngine::Synthetic,
files: vec![],
cli_defaults: ModelCliDefaults {
cfg_scale: 1.0,
steps: 4,
width: 64,
height: 64,
..Default::default()
},
},
enabled: true,
origin: "local".into(),
exclusive_group: None,
}
}
const TEST_TOKEN: &str = "test-token-0123456789abcdef";
struct Harness {
url: String,
observers: WorkerObservers,
host: crate::host::ModelHost,
services: ModelServices,
stop: Arc<AtomicBool>,
handle: Option<std::thread::JoinHandle<()>>,
}
impl Harness {
fn start(catalog: Catalog) -> Self {
Self::start_with_gate(catalog, JobGate::new())
}
fn start_with_gate(catalog: Catalog, gate: JobGate) -> Self {
Self::start_full(catalog, gate, 20.0)
}
fn start_with_free(catalog: Catalog, free_gib: f32) -> Self {
Self::start_full(catalog, JobGate::new(), free_gib)
}
fn start_full(catalog: Catalog, gate: JobGate, free_gib: f32) -> Self {
let engine: Arc<dyn Engine> = Arc::new(SyntheticEngine::new());
let observers = WorkerObservers::default();
let catalog = Arc::new(Mutex::new(catalog));
let host = crate::host::ModelHost::new(
catalog.clone(),
Arc::new(crate::test_support::InstantRuntime),
Arc::new(crate::test_support::FixedProbe(free_gib)),
crate::residency::Residency::load_for_serving(None),
);
let services = ModelServices::new(host.clone());
services
.stream_port
.store(4798, std::sync::atomic::Ordering::SeqCst);
let api = LocalApi::bind(
"127.0.0.1:0",
engine,
catalog,
None,
observers.clone(),
TEST_TOKEN.to_string(),
gate.clone(),
None,
services.clone(),
)
.unwrap();
let url = api.url();
let stop = Arc::new(AtomicBool::new(false));
let stop_thread = stop.clone();
let handle = std::thread::spawn(move || api.serve(&stop_thread));
Harness {
url,
observers,
host,
services,
stop,
handle: Some(handle),
}
}
fn post(&self, path: &str) -> reqwest::blocking::RequestBuilder {
reqwest::blocking::Client::new()
.post(format!("{}{}", self.url, path))
.bearer_auth(TEST_TOKEN)
}
fn get(&self, path: &str) -> reqwest::blocking::RequestBuilder {
reqwest::blocking::Client::new()
.get(format!("{}{}", self.url, path))
.bearer_auth(TEST_TOKEN)
}
}
impl Drop for Harness {
fn drop(&mut self) {
self.stop.store(true, Ordering::Relaxed);
if let Some(handle) = self.handle.take() {
let _ = handle.join();
}
}
}
fn test_host(catalog: &Arc<Mutex<Catalog>>) -> crate::host::ModelHost {
crate::host::ModelHost::new(
catalog.clone(),
Arc::new(crate::test_support::InstantRuntime),
Arc::new(crate::test_support::FixedProbe(20.0)),
crate::residency::Residency::load_for_serving(None),
)
}
fn seeded_catalog() -> Catalog {
Catalog {
models: vec![synthetic_model("synthetic-img")],
..Default::default()
}
}
#[test]
fn post_image_returns_image_bytes_and_records_job() {
let h = Harness::start(seeded_catalog());
let res = h
.post("/image")
.json(&serde_json::json!({ "prompt": "a blue bird" }))
.send()
.unwrap();
assert_eq!(res.status(), 200);
assert_eq!(res.headers()["content-type"], "image/webp");
let bytes = res.bytes().unwrap();
assert!(!bytes.is_empty());
assert_eq!(h.observers.local_jobs.lock().len(), 1);
}
#[test]
fn post_image_honours_requested_ext() {
let h = Harness::start(seeded_catalog());
let res = h
.post("/image")
.json(&serde_json::json!({ "prompt": "x", "ext": "png" }))
.send()
.unwrap();
assert_eq!(res.status(), 200);
assert_eq!(res.headers()["content-type"], "image/png");
}
#[test]
fn get_models_lists_catalog() {
let h = Harness::start(seeded_catalog());
let body = h.get("/models").send().unwrap().text().unwrap();
assert!(body.contains("synthetic-img"));
}
fn json(res: reqwest::blocking::Response) -> (u16, serde_json::Value) {
let status = res.status().as_u16();
(status, res.json().unwrap())
}
fn wait_state(h: &Harness, id: &str, want: &str) -> serde_json::Value {
let deadline = std::time::Instant::now() + std::time::Duration::from_secs(5);
loop {
let (status, body) = json(h.get(&format!("/models/{id}/state")).send().unwrap());
assert_eq!(status, 200, "{body}");
if body["state"] == want {
return body;
}
assert!(
std::time::Instant::now() < deadline,
"never reached {want}: {body}"
);
std::thread::sleep(std::time::Duration::from_millis(10));
}
}
#[test]
fn get_models_carries_state_and_residency() {
let h = Harness::start(seeded_catalog());
let (status, body) = json(h.get("/models").send().unwrap());
assert_eq!(status, 200);
let first = &body.as_array().unwrap()[0];
assert_eq!(first["state"], "unloaded");
assert_eq!(first["resident"], false);
assert!(first["id"].is_string(), "catalogue fields stay: {first}");
}
#[test]
fn load_reaches_loaded_and_marks_resident() {
let h = Harness::start(seeded_catalog());
let (status, body) = json(h.post("/models/synthetic-img/load").send().unwrap());
assert!(status == 202 || status == 200, "{status} {body}");
assert_eq!(body["id"], "synthetic-img");
assert_eq!(body["resident"], true);
let body = wait_state(&h, "synthetic-img", "loaded");
assert!(body["since"].is_string());
let (status, _) = json(h.post("/models/synthetic-img/load").send().unwrap());
assert_eq!(status, 200, "loading a loaded model is a no-op");
}
#[test]
fn unload_frees_and_clears_residency() {
let h = Harness::start(seeded_catalog());
h.post("/models/synthetic-img/load").send().unwrap();
wait_state(&h, "synthetic-img", "loaded");
let (status, body) = json(h.post("/models/synthetic-img/unload").send().unwrap());
assert!(status == 202 || status == 200, "{status} {body}");
assert_eq!(body["resident"], false);
wait_state(&h, "synthetic-img", "unloaded");
let (status, _) = json(h.post("/models/synthetic-img/unload").send().unwrap());
assert_eq!(status, 200, "unloading an unloaded model is a no-op");
}
#[test]
fn a_load_that_does_not_fit_is_a_409_with_the_numbers() {
let mut catalog = seeded_catalog();
catalog.models[0].vram_gb_estimate = 8.0;
let h = Harness::start_with_free(catalog, 4.0);
let (status, body) = json(h.post("/models/synthetic-img/load").send().unwrap());
assert_eq!(status, 409, "{body}");
assert_eq!(body["error"], "insufficient_memory");
assert_eq!(body["neededGib"], 8.0);
assert_eq!(body["freeGib"], 4.0);
assert!(body["marginGib"].is_number());
wait_state(&h, "synthetic-img", "unloaded");
}
#[test]
fn unknown_models_are_404_on_every_lifecycle_route() {
let h = Harness::start(seeded_catalog());
for res in [
h.get("/models/nope/state").send().unwrap(),
h.post("/models/nope/load").send().unwrap(),
h.post("/models/nope/unload").send().unwrap(),
] {
let (status, body) = json(res);
assert_eq!(status, 404);
assert_eq!(body["error"], "unknown_model");
}
}
#[test]
fn a_disabled_model_cannot_be_loaded() {
let mut catalog = seeded_catalog();
catalog.models[0].enabled = false;
let h = Harness::start(catalog);
let (status, body) = json(h.post("/models/synthetic-img/load").send().unwrap());
assert_eq!(status, 400);
assert_eq!(body["error"], "model_disabled");
}
#[test]
fn lifecycle_routes_need_the_token() {
let h = Harness::start(seeded_catalog());
let res = reqwest::blocking::Client::new()
.post(format!("{}/models/synthetic-img/load", h.url))
.send()
.unwrap();
assert_eq!(res.status(), 401);
wait_state(&h, "synthetic-img", "unloaded");
}
#[test]
fn deleting_a_loaded_model_unloads_it_first() {
let h = Harness::start(seeded_catalog());
h.post("/models/synthetic-img/load").send().unwrap();
wait_state(&h, "synthetic-img", "loaded");
let res = reqwest::blocking::Client::new()
.delete(format!("{}/models/synthetic-img", h.url))
.bearer_auth(TEST_TOKEN)
.send()
.unwrap();
assert_eq!(res.status(), 200);
let unloaded = h.host.wait_for(
"synthetic-img",
|s| *s == crate::lifecycle::ModelState::Unloaded,
std::time::Duration::from_secs(5),
);
assert!(unloaded.is_some(), "weights freed after delete");
assert_eq!(h.host.loaded_gib(), 0.0);
}
fn llm_catalog() -> Catalog {
Catalog {
models: vec![synthetic_model_of("chat-llm", TaskKind::Llm)],
..Default::default()
}
}
fn chat(h: &Harness, body: serde_json::Value) -> (u16, serde_json::Value) {
json(h.post("/v1/chat/completions").json(&body).send().unwrap())
}
#[test]
fn chat_is_served_on_the_lane_of_a_loaded_model() {
let h = Harness::start(llm_catalog());
h.post("/models/chat-llm/load").send().unwrap();
wait_state(&h, "chat-llm", "loaded");
let (status, body) = chat(
&h,
serde_json::json!({ "model": "chat-llm", "messages": [{ "role": "user", "content": "hi" }] }),
);
assert_eq!(status, 200, "{body}");
assert_eq!(body["choices"][0]["message"]["content"], "resident:hi");
}
#[test]
fn chat_uses_the_default_llm_when_it_is_loaded() {
let h = Harness::start(llm_catalog());
h.post("/models/chat-llm/load").send().unwrap();
wait_state(&h, "chat-llm", "loaded");
let (_, body) = chat(
&h,
serde_json::json!({ "messages": [{ "role": "user", "content": "yo" }] }),
);
assert_eq!(body["choices"][0]["message"]["content"], "resident:yo");
}
#[test]
fn chat_template_kwargs_reach_the_model() {
let h = Harness::start(llm_catalog());
h.post("/models/chat-llm/load").send().unwrap();
wait_state(&h, "chat-llm", "loaded");
let (_, body) = chat(
&h,
serde_json::json!({
"messages": [{ "role": "user", "content": "hi" }],
"chat_template_kwargs": { "enable_thinking": false },
}),
);
assert_eq!(
body["kwargs"],
serde_json::json!({ "enable_thinking": false })
);
}
#[test]
fn chat_on_an_unloaded_model_runs_as_a_transient_job() {
let h = Harness::start(llm_catalog());
let (status, body) = chat(
&h,
serde_json::json!({ "model": "chat-llm", "messages": [{ "role": "user", "content": "hi" }] }),
);
assert_eq!(status, 200, "{body}");
let content = body["choices"][0]["message"]["content"].as_str().unwrap();
assert!(!content.starts_with("resident:"), "{content}");
}
#[test]
fn a_resident_chat_is_recorded_as_a_local_job_and_skips_the_job_gate() {
let gate = JobGate::new();
let h = Harness::start_with_gate(llm_catalog(), gate.clone());
h.post("/models/chat-llm/load").send().unwrap();
wait_state(&h, "chat-llm", "loaded");
let _held = gate.try_reserve().expect("a transient job holds the gate");
let (status, _) = chat(
&h,
serde_json::json!({ "model": "chat-llm", "messages": [{ "role": "user", "content": "lane" }] }),
);
assert_eq!(status, 200, "a loaded model serves on its own lane");
let jobs = h.observers.local_jobs.lock().clone();
let last = jobs.front().expect("recorded");
assert_eq!(last.model, "chat-llm");
assert_eq!(last.prompt, "lane");
}
fn stream_catalog() -> Catalog {
let mut stt = synthetic_model_of("stt-a", TaskKind::AudioStt);
stt.source.engine = crate::types::ModelEngine::Parakeet;
Catalog {
models: vec![stt, synthetic_model_of("chat-llm", TaskKind::Llm)],
..Default::default()
}
}
#[test]
fn stream_tokens_are_minted_for_streaming_models() {
let h = Harness::start(stream_catalog());
let (status, body) = json(
h.post("/stream-tokens")
.json(&serde_json::json!({ "model": "stt-a", "ttlSecs": 600 }))
.send()
.unwrap(),
);
assert_eq!(status, 200, "{body}");
let token = body["token"].as_str().unwrap();
assert_eq!(token.len(), 64);
assert_eq!(body["model"], "stt-a");
assert_eq!(body["port"], 4798);
assert_eq!(body["path"], "/transcribe");
assert!(body["expiresAt"].is_string());
assert_eq!(
h.services.tokens.check(token, chrono::Utc::now()),
Ok("stt-a".to_string()),
"the listener accepts it"
);
}
#[test]
fn stream_tokens_are_refused_for_other_models() {
let h = Harness::start(stream_catalog());
let (status, body) = json(
h.post("/stream-tokens")
.json(&serde_json::json!({ "model": "chat-llm" }))
.send()
.unwrap(),
);
assert_eq!(status, 400);
assert_eq!(body["error"], "not_a_stream_model");
let (status, body) = json(
h.post("/stream-tokens")
.json(&serde_json::json!({ "model": "nope" }))
.send()
.unwrap(),
);
assert_eq!(status, 404);
assert_eq!(body["error"], "unknown_model");
}
#[test]
fn stream_tokens_need_a_running_listener() {
let mut h = Harness::start(stream_catalog());
h.services
.stream_port
.store(0, std::sync::atomic::Ordering::SeqCst);
let (status, body) = json(
h.post("/stream-tokens")
.json(&serde_json::json!({ "model": "stt-a" }))
.send()
.unwrap(),
);
assert_eq!(status, 503);
assert_eq!(body["error"], "stream_listener_down");
let _ = &mut h;
}
#[test]
fn stream_tokens_need_the_install_token() {
let h = Harness::start(stream_catalog());
let res = reqwest::blocking::Client::new()
.post(format!("{}/stream-tokens", h.url))
.json(&serde_json::json!({ "model": "stt-a" }))
.send()
.unwrap();
assert_eq!(res.status(), 401);
}
#[test]
fn post_models_adds_a_model_then_lists_it() {
let h = Harness::start(seeded_catalog());
let res = h
.post("/models")
.json(&synthetic_model("added-model"))
.send()
.unwrap();
assert_eq!(res.status(), 200);
let body = h.get("/models").send().unwrap().text().unwrap();
assert!(body.contains("added-model"));
}
#[test]
fn unknown_model_is_a_400() {
let h = Harness::start(seeded_catalog());
let res = h
.post("/image")
.json(&serde_json::json!({ "prompt": "x", "model": "nope" }))
.send()
.unwrap();
assert_eq!(res.status(), 400);
}
#[test]
fn invalid_json_is_a_400() {
let h = Harness::start(seeded_catalog());
let res = h
.post("/image")
.body("not json")
.header("content-type", "application/json")
.send()
.unwrap();
assert_eq!(res.status(), 400);
}
#[test]
fn healthz_reports_a_runtime_snapshot() {
let h = Harness::start(seeded_catalog());
let body: serde_json::Value = reqwest::blocking::get(format!("{}/healthz", h.url))
.unwrap()
.json()
.unwrap();
assert_eq!(body["ok"], true);
assert_eq!(body["version"], crate::AGENT_VERSION);
assert_eq!(body["busy"], false);
assert_eq!(body["engine"], "synthetic");
let raw = serde_json::to_string(&body).unwrap();
assert!(
!raw.contains(TEST_TOKEN),
"healthz must not carry the token"
);
}
#[test]
fn healthz_surfaces_gpu_runtime_when_probed() {
let h = Harness::start(seeded_catalog());
crate::runtime::set_gpu_runtime_status(
&h.observers,
Err(anyhow::anyhow!(
"Vulkan runtime not available: install libvulkan1"
)),
);
let body: serde_json::Value = reqwest::blocking::get(format!("{}/healthz", h.url))
.unwrap()
.json()
.unwrap();
assert_eq!(body["gpuRuntime"]["ok"], false);
assert!(body["gpuRuntime"]["detail"]
.as_str()
.unwrap()
.contains("libvulkan1"));
}
#[test]
fn chat_completions_returns_an_openai_shaped_body() {
let h = Harness::start(multi_kind_catalog());
let res = h
.post("/v1/chat/completions")
.json(&serde_json::json!({
"messages": [{"role": "user", "content": "hello there"}],
"max_tokens": 16
}))
.send()
.unwrap();
assert_eq!(res.status(), 200);
assert_eq!(res.headers()["content-type"], "application/json");
let body: serde_json::Value = res.json().unwrap();
assert!(
body.get("choices").is_some() || body.get("object").is_some(),
"expected an OpenAI-ish body, got: {body}"
);
assert!(h
.observers
.local_jobs
.lock()
.iter()
.any(|j| j.kind == TaskKind::Llm));
}
#[test]
fn tts_returns_audio_bytes() {
let h = Harness::start(multi_kind_catalog());
let res = h
.post("/tts")
.json(&serde_json::json!({ "text": "read this aloud" }))
.send()
.unwrap();
assert_eq!(res.status(), 200);
assert_eq!(res.headers()["content-type"], "audio/wav");
assert!(!res.bytes().unwrap().is_empty());
}
#[test]
fn stt_returns_a_transcript_json() {
let h = Harness::start(multi_kind_catalog());
let res = h
.post("/stt")
.json(&serde_json::json!({ "inputUrl": "https://example.com/a.wav" }))
.send()
.unwrap();
assert_eq!(res.status(), 200);
assert_eq!(res.headers()["content-type"], "application/json");
}
#[test]
fn video_returns_bytes() {
let h = Harness::start(multi_kind_catalog());
let res = h
.post("/video")
.json(&serde_json::json!({ "prompt": "a tiny dragon" }))
.send()
.unwrap();
assert_eq!(res.status(), 200);
assert!(!res.bytes().unwrap().is_empty());
}
#[test]
fn chat_without_an_llm_model_is_a_400() {
let h = Harness::start(seeded_catalog());
let res = h
.post("/v1/chat/completions")
.json(&serde_json::json!({
"messages": [{"role": "user", "content": "hi"}]
}))
.send()
.unwrap();
assert_eq!(res.status(), 400);
assert!(res.text().unwrap().contains("llm"));
}
#[test]
fn chat_endpoint_respects_the_busy_gate() {
let gate = JobGate::new();
let h = Harness::start_with_gate(multi_kind_catalog(), gate.clone());
let _held = gate.try_reserve().unwrap();
let res = h
.post("/v1/chat/completions")
.json(&serde_json::json!({ "messages": [{"role":"user","content":"x"}] }))
.send()
.unwrap();
assert_eq!(res.status(), 503);
}
#[test]
fn jobs_endpoint_reports_after_generation() {
let h = Harness::start(seeded_catalog());
h.post("/image")
.json(&serde_json::json!({ "prompt": "x" }))
.send()
.unwrap();
let body = h.get("/jobs").send().unwrap().text().unwrap();
assert!(body.contains("\"completed\""));
assert!(body.contains("synthetic-img"));
}
#[test]
fn routes_reject_requests_without_a_token() {
let h = Harness::start(seeded_catalog());
let client = reqwest::blocking::Client::new();
let cases: Vec<(reqwest::blocking::RequestBuilder, &str)> = vec![
(
client
.post(format!("{}/image", h.url))
.json(&serde_json::json!({ "prompt": "x" })),
"POST /image",
),
(client.get(format!("{}/models", h.url)), "GET /models"),
(
client
.post(format!("{}/models", h.url))
.json(&synthetic_model("evil")),
"POST /models",
),
(
client.delete(format!("{}/models/synthetic-img", h.url)),
"DELETE /models",
),
(client.get(format!("{}/jobs", h.url)), "GET /jobs"),
];
for (req, name) in cases {
let res = req.send().unwrap();
assert_eq!(res.status(), 401, "{name} must require the token");
let body = res.text().unwrap();
assert!(
body.contains("local-api.json"),
"{name}: the 401 must point at the discovery file, got: {body}"
);
}
let body = h.get("/models").send().unwrap().text().unwrap();
assert!(!body.contains("evil"));
assert!(body.contains("synthetic-img"));
}
#[test]
fn routes_reject_a_wrong_token() {
let h = Harness::start(seeded_catalog());
let res = reqwest::blocking::Client::new()
.get(format!("{}/models", h.url))
.bearer_auth("wrong-token")
.send()
.unwrap();
assert_eq!(res.status(), 401);
}
#[test]
fn daemon_routes_need_daemon_control() {
let h = Harness::start(multi_kind_catalog());
let resp = h.get("/daemon/status").send().unwrap();
assert_eq!(resp.status(), 503);
let body: serde_json::Value = resp.json().unwrap();
assert_eq!(body["error"], "daemon_control_unavailable");
}
#[test]
fn daemon_routes_need_the_token() {
let h = Harness::start(multi_kind_catalog());
let resp = reqwest::blocking::Client::new()
.get(format!("{}/daemon/status", h.url))
.send()
.unwrap();
assert_eq!(resp.status(), 401);
}
#[test]
fn an_unknown_daemon_route_is_not_found() {
let daemon = crate::test_support::DaemonHarness::start();
let resp = reqwest::blocking::Client::new()
.get(format!("{}/daemon/nope", daemon.url))
.bearer_auth(crate::test_support::HARNESS_TOKEN)
.send()
.unwrap();
assert_eq!(resp.status(), 404);
}
#[test]
fn a_malformed_config_body_is_a_bad_request() {
let daemon = crate::test_support::DaemonHarness::start();
let resp = reqwest::blocking::Client::new()
.put(format!("{}/daemon/config", daemon.url))
.bearer_auth(crate::test_support::HARNESS_TOKEN)
.body("{")
.send()
.unwrap();
assert_eq!(resp.status(), 400);
}
#[test]
fn the_models_listing_carries_since_and_the_failure() {
let daemon = crate::test_support::DaemonHarness::start();
let models = daemon.client().models().unwrap();
assert!(models.iter().all(|m| m.since.is_some() && m.loadable));
assert!(models.iter().all(|m| m.error.is_none()));
}
#[test]
fn job_routes_and_query_params_parse() {
assert_eq!(job_route("/jobs/local-1/log", "/log"), Some("local-1"));
assert_eq!(job_route("/jobs//log", "/log"), None);
assert_eq!(job_route("/jobs/a/b/log", "/log"), None);
assert_eq!(query_param("/daemon/logs?after=12", "after"), Some("12"));
assert_eq!(query_param("/daemon/logs?x=1&after=3", "after"), Some("3"));
assert_eq!(query_param("/daemon/logs?afterx=3", "after"), None);
assert_eq!(query_param("/daemon/logs", "after"), None);
}
#[test]
fn healthz_needs_no_token() {
let h = Harness::start(seeded_catalog());
let res = reqwest::blocking::get(format!("{}/healthz", h.url)).unwrap();
assert_eq!(res.status(), 200);
}
#[test]
fn healthz_answers_while_a_generation_is_in_flight() {
let engine: Arc<dyn Engine> = Arc::new(SlowEngine {
inner: SyntheticEngine::new(),
delay: std::time::Duration::from_millis(400),
});
let observers = WorkerObservers::default();
let catalog = Arc::new(Mutex::new(seeded_catalog()));
let api = LocalApi::bind(
"127.0.0.1:0",
engine,
catalog.clone(),
None,
observers,
TEST_TOKEN.to_string(),
JobGate::new(),
None,
ModelServices::new(test_host(&catalog)),
)
.unwrap();
let url = api.url();
let stop = Arc::new(AtomicBool::new(false));
let stop_thread = stop.clone();
let handle = std::thread::spawn(move || api.serve(&stop_thread));
let gen_url = url.clone();
let gen = std::thread::spawn(move || {
reqwest::blocking::Client::new()
.post(format!("{gen_url}/image"))
.bearer_auth(TEST_TOKEN)
.json(&serde_json::json!({ "prompt": "slow" }))
.timeout(std::time::Duration::from_secs(5))
.send()
.unwrap()
.status()
.as_u16()
});
std::thread::sleep(std::time::Duration::from_millis(100));
let start = std::time::Instant::now();
let health = reqwest::blocking::get(format!("{url}/healthz")).unwrap();
let elapsed = start.elapsed();
assert_eq!(health.status(), 200);
assert!(
elapsed < std::time::Duration::from_millis(250),
"healthz blocked behind the generation ({elapsed:?}); the pool isn't concurrent"
);
let body: serde_json::Value = health.json().unwrap();
assert_eq!(body["busy"], true, "a running job must show busy=true");
assert_eq!(gen.join().unwrap(), 200, "the generation still succeeds");
stop.store(true, Ordering::Relaxed);
let _ = handle.join();
}
#[test]
fn non_loopback_host_header_is_forbidden_even_with_a_token() {
let h = Harness::start(seeded_catalog());
let res = h
.get("/models")
.header("host", "evil.example:4787")
.send()
.unwrap();
assert_eq!(res.status(), 403);
assert!(res.text().unwrap().contains("Host"));
}
#[test]
fn cross_site_origin_is_forbidden_even_with_a_token() {
let h = Harness::start(seeded_catalog());
let res = h
.post("/image")
.header("origin", "https://evil.example")
.json(&serde_json::json!({ "prompt": "x" }))
.send()
.unwrap();
assert_eq!(res.status(), 403);
assert!(res.text().unwrap().contains("Origin"));
}
#[test]
fn loopback_origin_is_allowed() {
let h = Harness::start(seeded_catalog());
let res = h
.get("/models")
.header("origin", "http://localhost:5173")
.send()
.unwrap();
assert_eq!(res.status(), 200);
}
#[test]
fn oversized_body_is_a_413() {
let h = Harness::start(seeded_catalog());
let big = "x".repeat(MAX_BODY_BYTES + 1);
let res = h.post("/image").body(big).send().unwrap();
assert_eq!(res.status(), 413);
}
#[test]
fn body_at_the_cap_is_still_read() {
let h = Harness::start(seeded_catalog());
let exact = "x".repeat(MAX_BODY_BYTES);
let res = h.post("/image").body(exact).send().unwrap();
assert_eq!(res.status(), 400);
}
#[test]
fn bind_refuses_an_empty_token() {
let engine: Arc<dyn Engine> = Arc::new(SyntheticEngine::new());
let catalog = Arc::new(Mutex::new(seeded_catalog()));
let err = LocalApi::bind(
"127.0.0.1:0",
engine,
catalog.clone(),
None,
WorkerObservers::default(),
String::new(),
JobGate::new(),
None,
ModelServices::new(test_host(&catalog)),
)
.err()
.expect("empty token must be refused")
.to_string();
assert!(err.contains("empty token"), "got: {err}");
}
#[test]
fn post_image_returns_503_when_the_shared_gate_is_held() {
let gate = JobGate::new();
let h = Harness::start_with_gate(seeded_catalog(), gate.clone());
let reservation = gate.try_reserve().expect("pre-hold the slot");
let res = h
.post("/image")
.json(&serde_json::json!({ "prompt": "x" }))
.send()
.unwrap();
assert_eq!(res.status(), 503);
assert_eq!(res.headers()["retry-after"], "2");
drop(reservation);
let res = h
.post("/image")
.json(&serde_json::json!({ "prompt": "x" }))
.send()
.unwrap();
assert_eq!(res.status(), 200);
}
#[test]
fn host_is_loopback_accepts_only_loopback_shapes() {
for ok in [
"127.0.0.1",
"127.0.0.1:4787",
"localhost",
"LOCALHOST:80",
"[::1]",
"[::1]:4787",
] {
assert!(host_is_loopback(ok), "{ok} should be loopback");
}
for bad in [
"evil.example",
"evil.example:4787",
"127.0.0.1.evil.example",
"192.168.1.10:4787",
"[::2]:4787",
"[::1",
"",
] {
assert!(!host_is_loopback(bad), "{bad} should be rejected");
}
}
#[test]
fn origin_is_loopback_accepts_only_loopback_origins() {
for ok in [
"http://127.0.0.1:4787",
"http://localhost:5173",
"https://localhost",
"http://[::1]:3000",
] {
assert!(origin_is_loopback(ok), "{ok} should be allowed");
}
for bad in [
"https://evil.example",
"http://192.168.1.10",
"null",
"file://",
"chrome-extension://abc",
"",
] {
assert!(!origin_is_loopback(bad), "{bad} should be rejected");
}
}
#[test]
fn deny_reason_orders_host_origin_then_token() {
let t = "tok";
assert!(matches!(
deny_reason(Some("evil.example"), Some("https://evil.example"), None, t),
Some(Denial::Host(_))
));
assert!(matches!(
deny_reason(Some("127.0.0.1"), Some("https://evil.example"), None, t),
Some(Denial::Origin(_))
));
assert_eq!(
deny_reason(Some("127.0.0.1"), None, None, t),
Some(Denial::Token)
);
assert_eq!(
deny_reason(None, None, Some("Basic dXNlcjpwdw=="), t),
Some(Denial::Token)
);
assert_eq!(deny_reason(None, None, Some("Bearer tok"), t), None);
assert_eq!(deny_reason(None, None, Some("bearer tok"), t), None);
}
#[test]
fn token_matches_is_exact() {
assert!(token_matches("abc", "abc"));
assert!(!token_matches("abd", "abc"));
assert!(!token_matches("ab", "abc"));
assert!(!token_matches("", "abc"));
}
#[test]
fn discovery_file_round_trips_and_is_owner_only() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("local-api.json");
write_discovery_file(&path, "http://127.0.0.1:4787", "tok-123").unwrap();
let parsed: serde_json::Value =
serde_json::from_str(&std::fs::read_to_string(&path).unwrap()).unwrap();
assert_eq!(parsed["url"], "http://127.0.0.1:4787");
assert_eq!(parsed["token"], "tok-123");
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
let mode = std::fs::metadata(&path).unwrap().permissions().mode();
assert_eq!(
mode & 0o077,
0,
"discovery file carries the token and must be owner-only, got {mode:o}"
);
}
remove_discovery_file(&path);
assert!(!path.exists());
remove_discovery_file(&path);
}
}