use crate::broker::worker::{ProviderType, WorkerHeartbeat, WorkerRegistration};
use std::process::{Child, Command};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::Duration;
pub const HEARTBEAT_INTERVAL: Duration = Duration::from_secs(15);
#[derive(Debug, Clone, PartialEq)]
pub struct ServeArgs {
pub model_addresses: Vec<crate::model_uri::ModelAddress>,
pub general: bool,
pub price_per_mtok: f64,
pub base_port: u16,
pub api_url: Option<String>,
pub advertise: Option<String>,
}
impl Default for ServeArgs {
fn default() -> Self {
ServeArgs {
model_addresses: Vec::new(),
general: false,
price_per_mtok: 0.0,
base_port: 8600,
api_url: None,
advertise: None,
}
}
}
pub fn parse_serve_args(args: &[String]) -> Result<ServeArgs, String> {
let mut out = ServeArgs::default();
let mut it = args.iter().peekable();
while let Some(arg) = it.next() {
match arg.as_str() {
"--general" => out.general = true,
"--price" => {
let v = it
.next()
.ok_or_else(|| "--price requires a value".to_string())?;
out.price_per_mtok = v
.parse::<f64>()
.map_err(|_| format!("invalid --price value '{v}'"))?;
}
"--port" => {
let v = it
.next()
.ok_or_else(|| "--port requires a value".to_string())?;
out.base_port = v
.parse::<u16>()
.map_err(|_| format!("invalid --port value '{v}'"))?;
}
"--api-url" => {
let v = it
.next()
.ok_or_else(|| "--api-url requires a value".to_string())?;
out.api_url = Some(v.clone());
}
"--advertise" => {
let v = it
.next()
.ok_or_else(|| "--advertise requires a value".to_string())?;
out.advertise = Some(v.clone());
}
other => {
let addr = crate::model_uri::parse_model_address(other)
.ok_or_else(|| format!("not a model address: '{other}'"))?;
out.model_addresses.push(addr);
}
}
}
if !out.general && out.model_addresses.is_empty() {
return Err(
"zc serve requires at least one zc://<owner>/<name>, zc://<uuid>, or --general"
.to_string(),
);
}
Ok(out)
}
pub fn worker_uri(advertise: Option<&str>, port: u16) -> String {
match advertise {
None => format!("http://127.0.0.1:{port}"),
Some(host) => {
let host = host.trim_end_matches('/');
if host.contains("://") {
format!("{host}:{port}")
} else {
format!("http://{host}:{port}")
}
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct ResolvedModelFile {
pub digest: String,
pub file_path: String,
}
pub fn select_gguf_file<'a>(file_paths: &[&'a str]) -> Option<&'a str> {
file_paths.iter().find(|p| p.ends_with(".gguf")).copied()
}
pub fn resolve_model_file(model_json: &serde_json::Value) -> Result<ResolvedModelFile, String> {
let version = model_json
.get("latest_version")
.ok_or_else(|| "model has no latest_version (not published/ready?)".to_string())?;
let digest = version
.get("digest")
.and_then(|d| d.as_str())
.ok_or_else(|| "latest_version missing digest".to_string())?
.to_string();
let files = version
.get("files")
.and_then(|f| f.as_array())
.ok_or_else(|| "latest_version missing files".to_string())?;
let paths: Vec<&str> = files
.iter()
.filter_map(|f| f.get("path").and_then(|p| p.as_str()))
.collect();
let file_path = select_gguf_file(&paths)
.ok_or_else(|| "no .gguf file in latest version".to_string())?
.to_string();
Ok(ResolvedModelFile { digest, file_path })
}
pub fn download_url(api_url: &str, model_uuid: &str, resolved: &ResolvedModelFile) -> String {
format!(
"{}/api/models/{}/versions/{}/files/{}",
api_url.trim_end_matches('/'),
model_uuid,
resolved.digest,
resolved.file_path
)
}
pub fn worker_registration_for(
worker_name: &str,
local_uri: &str,
model_uuid: &str,
price_per_mtok: f64,
) -> WorkerRegistration {
WorkerRegistration {
name: worker_name.to_string(),
uri: local_uri.to_string(),
worker_type: "llm".to_string(),
resources: Default::default(),
pricing: Default::default(),
tags: vec!["model-inference".to_string()],
max_timeout_secs: 0.0,
hardware: Default::default(),
wireguard_ip: None,
is_docker: None,
source_node: None,
explicit_local: false,
provider_type: ProviderType::Specialized,
served_models: vec![model_uuid.to_string()],
price_per_mtok,
}
}
pub fn worker_registration_general(
worker_name: &str,
local_uri: &str,
price_per_mtok: f64,
) -> WorkerRegistration {
WorkerRegistration {
name: worker_name.to_string(),
uri: local_uri.to_string(),
worker_type: "llm".to_string(),
resources: Default::default(),
pricing: Default::default(),
tags: vec!["model-inference".to_string(), "general".to_string()],
max_timeout_secs: 0.0,
hardware: Default::default(),
wireguard_ip: None,
is_docker: None,
source_node: None,
explicit_local: false,
provider_type: ProviderType::General,
served_models: vec![crate::broker::worker::MODEL_WILDCARD.to_string()],
price_per_mtok,
}
}
pub fn general_provider_note() -> &'static str {
"Note: --general registers this worker as a general provider with the \
broker (served_models: [\"*\"]), but zc does not yet implement on-demand \
multi-model serving (dynamically pulling+launching llama-server for \
whatever model a request names). That dynamic loader is a Phase-2 \
follow-up. Pass one or more zc://{model-uuid} addresses to actually \
serve models today."
}
fn download_to_file(
agent: &ureq::Agent,
url: &str,
dest_path: &std::path::Path,
) -> Result<(), String> {
if let Some(parent) = dest_path.parent() {
std::fs::create_dir_all(parent)
.map_err(|e| format!("creating {}: {e}", parent.display()))?;
}
let mut response = agent
.get(url)
.call()
.map_err(|e| format!("downloading {url}: {e}"))?;
let mut reader = response.body_mut().as_reader();
let mut file = std::fs::File::create(dest_path)
.map_err(|e| format!("creating {}: {e}", dest_path.display()))?;
std::io::copy(&mut reader, &mut file)
.map_err(|e| format!("writing {}: {e}", dest_path.display()))?;
Ok(())
}
fn cache_path(model_uuid: &str, resolved: &ResolvedModelFile) -> std::path::PathBuf {
let home = std::env::var("ZAKURO_HOME").unwrap_or_else(|_| {
std::env::var("HOME")
.map(|h| format!("{h}/.zakuro"))
.unwrap_or_else(|_| "/tmp/.zakuro".to_string())
});
std::path::PathBuf::from(home)
.join("models")
.join(model_uuid)
.join(&resolved.digest)
.join(&resolved.file_path)
}
pub fn fetch_model_gguf(
agent: &ureq::Agent,
api_url: &str,
model_uuid: &str,
auth_bearer: Option<&str>,
) -> Result<std::path::PathBuf, String> {
let model_endpoint = format!(
"{}/api/models/{}",
api_url.trim_end_matches('/'),
model_uuid
);
let mut req = agent.get(&model_endpoint);
if let Some(token) = auth_bearer {
req = req.header("Authorization", &format!("Bearer {token}"));
}
let body = req
.call()
.map_err(|e| format!("fetching model {model_uuid}: {e}"))?
.body_mut()
.read_to_string()
.map_err(|e| format!("reading model {model_uuid} response: {e}"))?;
let model_json: serde_json::Value = serde_json::from_str(&body)
.map_err(|e| format!("parsing model {model_uuid} response: {e}"))?;
let resolved = resolve_model_file(&model_json)?;
let dest = cache_path(model_uuid, &resolved);
if !dest.exists() {
let url = download_url(api_url, model_uuid, &resolved);
download_to_file(agent, &url, &dest)?;
}
Ok(dest)
}
#[allow(clippy::too_many_arguments)]
pub fn serve_one_model(
agent: &ureq::Agent,
api_url: &str,
broker_url: &str,
model_uuid: &str,
port: u16,
price_per_mtok: f64,
auth_bearer: Option<&str>,
advertise: Option<&str>,
) -> Result<(Child, String), String> {
let dest = fetch_model_gguf(agent, api_url, model_uuid, auth_bearer)?;
let child = Command::new("llama-server")
.arg("--model")
.arg(&dest)
.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?)"))?;
let advertised_uri = worker_uri(advertise, port);
let worker_name = format!("zc-serve-{model_uuid}");
let registration =
worker_registration_for(&worker_name, &advertised_uri, model_uuid, price_per_mtok);
let worker_id = register_worker(agent, broker_url, ®istration)?;
spawn_heartbeat_loop(agent.clone(), broker_url.to_string(), worker_id, &child);
Ok((child, worker_name))
}
pub fn register_worker(
agent: &ureq::Agent,
broker_url: &str,
registration: &WorkerRegistration,
) -> Result<String, String> {
let url = format!("{}/workers", broker_url.trim_end_matches('/'));
let mut response = agent
.post(&url)
.header("Content-Type", "application/json")
.send_json(registration)
.map_err(|e| format!("registering worker with broker {broker_url}: {e}"))?;
let body = response
.body_mut()
.read_to_string()
.map_err(|e| format!("reading registration response: {e}"))?;
let worker: serde_json::Value =
serde_json::from_str(&body).map_err(|e| format!("parsing registration response: {e}"))?;
worker
.get("id")
.and_then(|v| v.as_str())
.map(|s| s.to_string())
.ok_or_else(|| "registration response missing worker id".to_string())
}
pub fn send_heartbeat(
agent: &ureq::Agent,
broker_url: &str,
worker_id: &str,
) -> Result<(), String> {
let url = format!("{}/workers/heartbeat", broker_url.trim_end_matches('/'));
let heartbeat = WorkerHeartbeat {
worker_id: worker_id.to_string(),
resources: None,
active_requests: None,
status: None,
max_timeout_secs: None,
};
agent
.post(&url)
.header("Content-Type", "application/json")
.send_json(&heartbeat)
.map_err(|e| format!("sending heartbeat to broker {broker_url}: {e}"))?;
Ok(())
}
pub(crate) fn run_heartbeat_loop(
stop: &AtomicBool,
interval: Duration,
mut sleep_fn: impl FnMut(Duration),
mut send_fn: impl FnMut(),
) {
while !stop.load(Ordering::Relaxed) {
sleep_fn(interval);
if stop.load(Ordering::Relaxed) {
break;
}
send_fn();
}
}
fn spawn_heartbeat_loop(agent: ureq::Agent, broker_url: String, worker_id: String, child: &Child) {
let stop = Arc::new(AtomicBool::new(false));
let pid = child.id();
let stop_hb = stop.clone();
let agent_hb = agent.clone();
let broker_hb = broker_url.clone();
let worker_hb = worker_id.clone();
let _ = std::thread::Builder::new()
.name("zc-serve-heartbeat".into())
.spawn(move || {
run_heartbeat_loop(&stop_hb, HEARTBEAT_INTERVAL, std::thread::sleep, || {
if let Err(e) = send_heartbeat(&agent_hb, &broker_hb, &worker_hb) {
eprintln!(" [SERVE] heartbeat failed: {e}");
}
});
});
std::thread::Builder::new()
.name("zc-serve-heartbeat-watch".into())
.spawn(move || {
while process_alive(pid) {
std::thread::sleep(Duration::from_secs(1));
}
stop.store(true, Ordering::Relaxed);
})
.ok();
}
#[cfg(unix)]
fn process_alive(pid: u32) -> bool {
unsafe { libc::kill(pid as libc::pid_t, 0) == 0 }
}
#[cfg(not(unix))]
fn process_alive(_pid: u32) -> bool {
true
}
#[cfg(test)]
mod tests {
use super::*;
fn uuids(v: &[&str]) -> Vec<crate::model_uri::ModelAddress> {
v.iter()
.map(|s| crate::model_uri::ModelAddress::Uuid(s.to_string()))
.collect()
}
#[test]
fn parses_single_specialized_model() {
let args = vec!["zc://ae5f3db4-437a-40d1-93ec-c8258315d69a".to_string()];
let parsed = parse_serve_args(&args).unwrap();
assert!(!parsed.general);
assert_eq!(
parsed.model_addresses,
uuids(&["ae5f3db4-437a-40d1-93ec-c8258315d69a"])
);
assert_eq!(parsed.base_port, 8600);
assert_eq!(parsed.price_per_mtok, 0.0);
}
#[test]
fn parses_multiple_models_with_price_and_port() {
let args = vec![
"zc://ae5f3db4-437a-40d1-93ec-c8258315d69a".to_string(),
"zc://bb5f3db4-437a-40d1-93ec-c8258315d69b".to_string(),
"--price".to_string(),
"2.5".to_string(),
"--port".to_string(),
"9100".to_string(),
];
let parsed = parse_serve_args(&args).unwrap();
assert_eq!(parsed.model_addresses.len(), 2);
assert_eq!(parsed.price_per_mtok, 2.5);
assert_eq!(parsed.base_port, 9100);
}
#[test]
fn parses_general_flag_alone() {
let args = vec!["--general".to_string()];
let parsed = parse_serve_args(&args).unwrap();
assert!(parsed.general);
assert!(parsed.model_addresses.is_empty());
}
#[test]
fn parses_api_url_override() {
let args = vec![
"--general".to_string(),
"--api-url".to_string(),
"https://stg.api.zakuro-ai.com".to_string(),
];
let parsed = parse_serve_args(&args).unwrap();
assert_eq!(
parsed.api_url.as_deref(),
Some("https://stg.api.zakuro-ai.com")
);
}
#[test]
fn parses_advertise_flag() {
let args = vec![
"--general".to_string(),
"--advertise".to_string(),
"192.168.0.167".to_string(),
];
let parsed = parse_serve_args(&args).unwrap();
assert_eq!(parsed.advertise.as_deref(), Some("192.168.0.167"));
}
#[test]
fn worker_uri_defaults_to_loopback() {
assert_eq!(worker_uri(None, 8600), "http://127.0.0.1:8600");
}
#[test]
fn worker_uri_bare_host_gets_scheme_and_port() {
assert_eq!(
worker_uri(Some("192.168.0.167"), 8601),
"http://192.168.0.167:8601"
);
}
#[test]
fn worker_uri_keeps_explicit_scheme() {
assert_eq!(
worker_uri(Some("https://node7.mesh/"), 8600),
"https://node7.mesh:8600"
);
}
#[test]
fn rejects_no_models_and_no_general() {
let args: Vec<String> = vec![];
assert!(parse_serve_args(&args).is_err());
}
#[test]
fn rejects_garbage_model_address() {
let args = vec!["not-a-model".to_string()];
assert!(parse_serve_args(&args).is_err());
}
#[test]
fn rejects_missing_price_value() {
let args = vec!["--price".to_string()];
assert!(parse_serve_args(&args).is_err());
}
#[test]
fn select_gguf_file_picks_the_gguf_among_others() {
let files = vec!["README.md", "tokenizer.json", "model.gguf"];
assert_eq!(select_gguf_file(&files), Some("model.gguf"));
}
#[test]
fn select_gguf_file_none_when_absent() {
let files = vec!["README.md", "tokenizer.json"];
assert_eq!(select_gguf_file(&files), None);
}
#[test]
fn resolve_model_file_from_stubbed_json() {
let json = serde_json::json!({
"latest_version": {
"digest": "sha256:abc123",
"files": [
{"path": "README.md"},
{"path": "model.gguf"}
]
}
});
let resolved = resolve_model_file(&json).unwrap();
assert_eq!(resolved.digest, "sha256:abc123");
assert_eq!(resolved.file_path, "model.gguf");
}
#[test]
fn resolve_model_file_errors_without_gguf() {
let json = serde_json::json!({
"latest_version": {
"digest": "sha256:abc123",
"files": [{"path": "README.md"}]
}
});
assert!(resolve_model_file(&json).is_err());
}
#[test]
fn resolve_model_file_errors_without_latest_version() {
let json = serde_json::json!({});
assert!(resolve_model_file(&json).is_err());
}
#[test]
fn download_url_construction() {
let resolved = ResolvedModelFile {
digest: "sha256:abc123".to_string(),
file_path: "model.gguf".to_string(),
};
let url = download_url(
"https://my.zakuro-ai.com/",
"ae5f3db4-437a-40d1-93ec-c8258315d69a",
&resolved,
);
assert_eq!(
url,
"https://my.zakuro-ai.com/api/models/ae5f3db4-437a-40d1-93ec-c8258315d69a/versions/sha256:abc123/files/model.gguf"
);
}
#[test]
fn worker_registration_for_specialized_model() {
let reg = worker_registration_for(
"zc-serve-ae5f3db4",
"http://127.0.0.1:8600",
"ae5f3db4-437a-40d1-93ec-c8258315d69a",
1.5,
);
assert_eq!(reg.provider_type, ProviderType::Specialized);
assert_eq!(
reg.served_models,
vec!["ae5f3db4-437a-40d1-93ec-c8258315d69a".to_string()]
);
assert_eq!(reg.uri, "http://127.0.0.1:8600");
assert_eq!(reg.price_per_mtok, 1.5);
assert_eq!(reg.name, "zc-serve-ae5f3db4");
}
#[test]
fn worker_registration_general_uses_wildcard() {
let reg = worker_registration_general("zc-serve-general", "http://127.0.0.1:8600", 0.0);
assert_eq!(reg.provider_type, ProviderType::General);
assert_eq!(reg.served_models, vec!["*".to_string()]);
}
#[test]
fn heartbeat_loop_sends_once_per_interval_until_stopped() {
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
let stop = AtomicBool::new(false);
let sends = AtomicUsize::new(0);
let sleeps = AtomicUsize::new(0);
const N: usize = 5;
run_heartbeat_loop(
&stop,
Duration::from_secs(15),
|d| {
assert_eq!(d, Duration::from_secs(15));
if sleeps.fetch_add(1, Ordering::Relaxed) + 1 >= N {
stop.store(true, Ordering::Relaxed);
}
},
|| {
sends.fetch_add(1, Ordering::Relaxed);
},
);
assert_eq!(sleeps.load(Ordering::Relaxed), N);
assert_eq!(sends.load(Ordering::Relaxed), N - 1);
}
#[test]
fn heartbeat_loop_never_sends_if_already_stopped() {
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
let stop = AtomicBool::new(true);
let sends = AtomicUsize::new(0);
run_heartbeat_loop(
&stop,
Duration::from_millis(1),
|_| panic!("must not sleep when already stopped"),
|| {
sends.fetch_add(1, Ordering::Relaxed);
},
);
assert_eq!(sends.load(Ordering::Relaxed), 0);
}
#[test]
fn heartbeat_interval_is_well_under_broker_worker_timeout() {
assert!(HEARTBEAT_INTERVAL < Duration::from_secs(30));
assert_eq!(HEARTBEAT_INTERVAL, Duration::from_secs(15));
}
#[test]
fn send_heartbeat_builds_dedicated_heartbeat_payload() {
let agent = ureq::Agent::new_with_config(
ureq::Agent::config_builder()
.timeout_connect(Some(Duration::from_millis(50)))
.build(),
);
let err = send_heartbeat(&agent, "http://127.0.0.1:1", "worker-123").unwrap_err();
assert!(
err.contains("sending heartbeat to broker"),
"unexpected error: {err}"
);
}
}