use std::collections::HashMap;
use std::process::Child;
use std::sync::Mutex;
use std::time::{Duration, Instant};
const DEFAULT_MAX_LOADED: usize = 2;
const READY_TIMEOUT: Duration = Duration::from_secs(300);
struct Loaded {
port: u16,
child: Child,
last_used: Instant,
}
pub struct LoaderHooks {
pub fetch: FetchFn,
pub spawn: SpawnFn,
pub wait_ready: Box<dyn Fn(u16) -> Result<(), String> + Send>,
}
pub type FetchFn = Box<dyn Fn(&str) -> Result<std::path::PathBuf, String> + Send>;
pub type SpawnFn = Box<dyn Fn(&std::path::Path, u16) -> Result<Child, String> + Send>;
pub struct GeneralLoader {
hooks: LoaderHooks,
state: Mutex<LoaderState>,
}
struct LoaderState {
loaded: HashMap<String, Loaded>,
next_port: u16,
max_loaded: usize,
}
impl GeneralLoader {
pub fn new(hooks: LoaderHooks, first_backend_port: u16, max_loaded: usize) -> Self {
GeneralLoader {
hooks,
state: Mutex::new(LoaderState {
loaded: HashMap::new(),
next_port: first_backend_port,
max_loaded: max_loaded.max(1),
}),
}
}
pub fn max_loaded_from_env() -> usize {
std::env::var("ZAKURO_GENERAL_MAX_MODELS")
.ok()
.and_then(|v| v.parse::<usize>().ok())
.filter(|n| *n >= 1)
.unwrap_or(DEFAULT_MAX_LOADED)
}
pub fn ensure_loaded(&self, model_uuid: &str) -> Result<u16, String> {
let mut st = self
.state
.lock()
.map_err(|_| "loader state poisoned".to_string())?;
if let Some(entry) = st.loaded.get_mut(model_uuid) {
let alive = entry.child.try_wait().map(|x| x.is_none()).unwrap_or(false);
if alive {
entry.last_used = Instant::now();
return Ok(entry.port);
}
st.loaded.remove(model_uuid);
}
while st.loaded.len() >= st.max_loaded {
let lru = st
.loaded
.iter()
.min_by_key(|(_, l)| l.last_used)
.map(|(k, _)| k.clone());
match lru {
Some(uuid) => {
if let Some(mut old) = st.loaded.remove(&uuid) {
let _ = old.child.kill();
let _ = old.child.wait();
}
}
None => break,
}
}
let gguf = (self.hooks.fetch)(model_uuid)?;
let port = st.next_port;
st.next_port = st.next_port.wrapping_add(1);
let mut child = (self.hooks.spawn)(&gguf, port)?;
if let Err(e) = (self.hooks.wait_ready)(port) {
let _ = child.kill();
let _ = child.wait();
return Err(format!(
"llama-server for {model_uuid} never became ready: {e}"
));
}
st.loaded.insert(
model_uuid.to_string(),
Loaded {
port,
child,
last_used: Instant::now(),
},
);
Ok(port)
}
pub fn resident(&self) -> Vec<String> {
self.state
.lock()
.map(|st| st.loaded.keys().cloned().collect())
.unwrap_or_default()
}
}
pub fn production_hooks(
agent: ureq::Agent,
api_url: String,
auth_bearer: Option<String>,
) -> LoaderHooks {
LoaderHooks {
fetch: Box::new(move |uuid| {
crate::serve::fetch_model_gguf(&agent, &api_url, uuid, auth_bearer.as_deref())
}),
spawn: Box::new(|gguf, port| {
std::process::Command::new("llama-server")
.arg("--model")
.arg(gguf)
.arg("--host")
.arg("0.0.0.0")
.arg("--port")
.arg(port.to_string())
.spawn()
.map_err(|e| format!("spawning llama-server: {e} (is it installed and on PATH?)"))
}),
wait_ready: Box::new(|port| {
let deadline = Instant::now() + READY_TIMEOUT;
let probe = ureq::Agent::new_with_config(
ureq::Agent::config_builder()
.timeout_global(Some(Duration::from_secs(2)))
.build(),
);
loop {
if probe
.get(&format!("http://127.0.0.1:{port}/health"))
.call()
.is_ok()
{
return Ok(());
}
if Instant::now() >= deadline {
return Err(format!("no /health within {}s", READY_TIMEOUT.as_secs()));
}
std::thread::sleep(Duration::from_millis(500));
}
}),
}
}
pub fn model_uuid_from_body(body: &serde_json::Value) -> Option<String> {
body.get("model")
.and_then(|m| m.as_str())
.and_then(crate::model_uri::parse_model_uri)
}
pub fn run_general_server(loader: GeneralLoader, bind_port: u16) -> Result<(), String> {
let server = tiny_http::Server::http(("0.0.0.0", bind_port))
.map_err(|e| format!("binding general provider on :{bind_port}: {e}"))?;
println!(" general provider listening on :{bind_port} (on-demand model loading)");
let proxy = ureq::Agent::new_with_config(
ureq::Agent::config_builder()
.timeout_global(Some(Duration::from_secs(600)))
.build(),
);
for mut request in server.incoming_requests() {
let respond = |request: tiny_http::Request, status: u16, body: serde_json::Value| {
let data = body.to_string();
let response = tiny_http::Response::from_string(data)
.with_status_code(status)
.with_header(
tiny_http::Header::from_bytes(&b"Content-Type"[..], &b"application/json"[..])
.expect("static header"),
);
let _ = request.respond(response);
};
let url = request.url().to_string();
if url == "/health" {
respond(request, 200, serde_json::json!({"status": "ok"}));
continue;
}
if !url.starts_with("/v1/chat/completions") {
respond(request, 404, serde_json::json!({"error": "not found"}));
continue;
}
let mut body_str = String::new();
if std::io::Read::read_to_string(request.as_reader(), &mut body_str).is_err() {
respond(
request,
400,
serde_json::json!({"error": "unreadable body"}),
);
continue;
}
let body: serde_json::Value = match serde_json::from_str(&body_str) {
Ok(v) => v,
Err(_) => {
respond(request, 400, serde_json::json!({"error": "invalid json"}));
continue;
}
};
let uuid = match model_uuid_from_body(&body) {
Some(u) => u,
None => {
respond(
request,
400,
serde_json::json!({"error": "model field must be a model uuid (zc://<uuid> or bare)"}),
);
continue;
}
};
let port = match loader.ensure_loaded(&uuid) {
Ok(p) => p,
Err(e) => {
eprintln!(" [GENERAL] load failed for {uuid}: {e}");
respond(
request,
503,
serde_json::json!({"error": format!("model load failed: {e}")}),
);
continue;
}
};
match proxy
.post(&format!("http://127.0.0.1:{port}/v1/chat/completions"))
.header("Content-Type", "application/json")
.send(&body_str[..])
{
Ok(mut upstream) => {
let status = upstream.status().as_u16();
let text = upstream.body_mut().read_to_string().unwrap_or_default();
let json: serde_json::Value = serde_json::from_str(&text)
.unwrap_or_else(|_| serde_json::json!({"error": "bad upstream body"}));
respond(request, status, json);
}
Err(e) => {
eprintln!(" [GENERAL] proxy to :{port} failed: {e}");
respond(
request,
502,
serde_json::json!({"error": "backend request failed"}),
);
}
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
fn sleeper() -> Child {
std::process::Command::new("sleep")
.arg("300")
.spawn()
.expect("spawn sleep")
}
fn test_loader(
max: usize,
fetches: Arc<AtomicUsize>,
spawns: Arc<AtomicUsize>,
) -> GeneralLoader {
let hooks = LoaderHooks {
fetch: Box::new(move |_uuid| {
fetches.fetch_add(1, Ordering::SeqCst);
Ok(std::path::PathBuf::from("/dev/null"))
}),
spawn: Box::new(move |_path, _port| {
spawns.fetch_add(1, Ordering::SeqCst);
Ok(sleeper())
}),
wait_ready: Box::new(|_port| Ok(())),
};
GeneralLoader::new(hooks, 9500, max)
}
const A: &str = "aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa";
const B: &str = "bbbbbbbb-bbbb-bbbb-bbbb-bbbbbbbbbbbb";
const C: &str = "cccccccc-cccc-cccc-cccc-cccccccccccc";
#[test]
fn second_request_reuses_the_resident_server() {
let fetches = Arc::new(AtomicUsize::new(0));
let spawns = Arc::new(AtomicUsize::new(0));
let loader = test_loader(2, fetches.clone(), spawns.clone());
let p1 = loader.ensure_loaded(A).unwrap();
let p2 = loader.ensure_loaded(A).unwrap();
assert_eq!(p1, p2);
assert_eq!(fetches.load(Ordering::SeqCst), 1, "one fetch, then cache");
assert_eq!(spawns.load(Ordering::SeqCst), 1, "one spawn, then reuse");
}
#[test]
fn lru_eviction_at_the_cap() {
let loader = test_loader(
2,
Arc::new(AtomicUsize::new(0)),
Arc::new(AtomicUsize::new(0)),
);
loader.ensure_loaded(A).unwrap();
loader.ensure_loaded(B).unwrap();
loader.ensure_loaded(A).unwrap();
loader.ensure_loaded(C).unwrap();
let mut resident = loader.resident();
resident.sort();
assert_eq!(
resident,
vec![A.to_string(), C.to_string()],
"B evicted as LRU"
);
}
#[test]
fn dead_backend_is_respawned_not_proxied_into() {
let spawns = Arc::new(AtomicUsize::new(0));
let loader = test_loader(2, Arc::new(AtomicUsize::new(0)), spawns.clone());
loader.ensure_loaded(A).unwrap();
{
let mut st = loader.state.lock().unwrap();
let entry = st.loaded.get_mut(A).unwrap();
entry.child.kill().unwrap();
entry.child.wait().unwrap();
}
let p = loader.ensure_loaded(A).unwrap();
assert_eq!(spawns.load(Ordering::SeqCst), 2, "dead server respawned");
assert!(p >= 9500);
}
#[test]
fn failed_readiness_kills_the_spawn_and_errors() {
let hooks = LoaderHooks {
fetch: Box::new(|_| Ok(std::path::PathBuf::from("/dev/null"))),
spawn: Box::new(|_, _| Ok(sleeper())),
wait_ready: Box::new(|_| Err("never healthy".into())),
};
let loader = GeneralLoader::new(hooks, 9500, 2);
assert!(loader.ensure_loaded(A).is_err());
assert!(
loader.resident().is_empty(),
"failed spawn must not stay resident"
);
}
#[test]
fn model_uuid_accepted_in_all_spellings() {
for m in [
format!("\"{A}\""),
format!("\"zc://{A}\""),
format!("\"zc://model-{A}\""),
] {
let body: serde_json::Value =
serde_json::from_str(&format!("{{\"model\":{m}}}")).unwrap();
assert_eq!(model_uuid_from_body(&body).as_deref(), Some(A));
}
let bad: serde_json::Value = serde_json::json!({"model": "llama-3-8b"});
assert_eq!(model_uuid_from_body(&bad), None);
}
}