use chrono::{DateTime, Duration, Utc};
use dashmap::DashMap;
use serde::{Deserialize, Serialize};
use std::collections::{HashSet, VecDeque};
use std::sync::Mutex;
pub const MODEL_WILDCARD: &str = "*";
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum ProviderType {
General,
#[default]
Specialized,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum WorkerStatus {
Healthy,
Busy,
Unhealthy,
Draining,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WorkerResources {
pub cpus_total: f64,
pub cpus_available: f64,
pub memory_total: u64,
pub memory_available: u64,
pub gpus_total: u32,
pub gpus_available: u32,
}
impl Default for WorkerResources {
fn default() -> Self {
Self {
cpus_total: 1.0,
cpus_available: 1.0,
memory_total: 1024 * 1024 * 1024, memory_available: 1024 * 1024 * 1024,
gpus_total: 0,
gpus_available: 0,
}
}
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub struct BrokerResources {
pub cpus_available: u64,
pub memory_available: u64,
pub gpus_available: u32,
pub healthy_workers: u32,
}
pub(crate) fn default_price_per_hour() -> f64 {
3.6
}
fn default_min_charge() -> f64 {
0.001
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WorkerPricing {
#[serde(default = "default_price_per_hour")]
pub price_per_hour: f64,
#[serde(default = "default_min_charge")]
pub min_charge: f64,
}
impl Default for WorkerPricing {
fn default() -> Self {
Self {
price_per_hour: 3.6, min_charge: 0.001, }
}
}
impl WorkerPricing {
pub fn estimate_cost(&self, duration_secs: f64) -> f64 {
(self.price_per_hour / 3600.0 * duration_secs).max(self.min_charge)
}
pub fn price_score(&self) -> f64 {
self.estimate_cost(10.0)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WorkerQuotas {
pub per_5h: u64,
pub per_week: u64,
pub per_month: u64,
}
impl Default for WorkerQuotas {
fn default() -> Self {
Self {
per_5h: std::env::var("ZAKURO_QUOTA_5H")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0),
per_week: std::env::var("ZAKURO_QUOTA_WEEK")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0),
per_month: std::env::var("ZAKURO_QUOTA_MONTH")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0),
}
}
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct HardwareInfo {
#[serde(default)]
pub gpu_model: Option<String>,
#[serde(default)]
pub gpu_vram_gb: Option<u32>,
#[serde(default)]
pub cpu_model: Option<String>,
#[serde(default)]
pub storage_gb: Option<u32>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Worker {
pub id: String,
pub name: String,
pub uri: String,
pub worker_type: String,
pub status: WorkerStatus,
pub resources: WorkerResources,
pub pricing: WorkerPricing,
pub last_heartbeat: DateTime<Utc>,
pub active_requests: u32,
pub total_requests: u64,
pub avg_latency_ms: f64,
pub tags: Vec<String>,
#[serde(default)]
pub max_timeout_secs: f64,
#[serde(default)]
pub hardware: HardwareInfo,
#[serde(default)]
pub wireguard_ip: Option<String>,
#[serde(default)]
pub is_docker: Option<bool>,
#[serde(default)]
pub source_node: Option<String>,
#[serde(default)]
pub explicit_local: bool,
#[serde(default)]
pub node_fp: String,
#[serde(default)]
pub slot: String,
#[serde(default)]
pub provider_type: ProviderType,
#[serde(default)]
pub served_models: Vec<String>,
#[serde(default)]
pub price_per_mtok: f64,
}
impl Worker {
pub fn zc_uri(&self) -> String {
format!("zc://worker-{}-{}", self.node_fp, self.slot)
}
pub fn new(id: String, name: String, uri: String, worker_type: String) -> Self {
Self {
id,
name,
uri,
worker_type,
status: WorkerStatus::Healthy,
resources: WorkerResources::default(),
pricing: WorkerPricing::default(),
last_heartbeat: Utc::now(),
active_requests: 0,
total_requests: 0,
avg_latency_ms: 0.0,
tags: Vec::new(),
max_timeout_secs: 0.0, hardware: HardwareInfo::default(),
wireguard_ip: None,
is_docker: None,
source_node: None,
explicit_local: false,
node_fp: String::new(),
slot: String::new(),
provider_type: ProviderType::default(),
served_models: Vec::new(),
price_per_mtok: 0.0,
}
}
pub fn can_handle(&self, cpus: f64, memory_bytes: u64, gpus: u32) -> bool {
self.status == WorkerStatus::Healthy
&& self.resources.cpus_available >= cpus
&& self.resources.memory_available >= memory_bytes
&& self.resources.gpus_available >= gpus
}
pub fn heartbeat(&mut self) {
self.last_heartbeat = Utc::now();
}
pub fn is_stale(&self, timeout_secs: i64) -> bool {
let elapsed = Utc::now().signed_duration_since(self.last_heartbeat);
elapsed.num_seconds() > timeout_secs
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WorkerRegistration {
pub name: String,
pub uri: String,
pub worker_type: String,
#[serde(default)]
pub resources: WorkerResources,
#[serde(default)]
pub pricing: WorkerPricing,
#[serde(default)]
pub tags: Vec<String>,
#[serde(default)]
pub max_timeout_secs: f64,
#[serde(default)]
pub hardware: HardwareInfo,
#[serde(default)]
pub wireguard_ip: Option<String>,
#[serde(default)]
pub is_docker: Option<bool>,
#[serde(default)]
pub source_node: Option<String>,
#[serde(default)]
pub explicit_local: bool,
#[serde(default)]
pub provider_type: ProviderType,
#[serde(default)]
pub served_models: Vec<String>,
#[serde(default)]
pub price_per_mtok: f64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WorkerHeartbeat {
pub worker_id: String,
#[serde(default)]
pub resources: Option<WorkerResources>,
#[serde(default)]
pub active_requests: Option<u32>,
#[serde(default)]
pub status: Option<WorkerStatus>,
#[serde(default)]
pub max_timeout_secs: Option<f64>,
}
#[derive(Debug)]
pub struct WorkerRegistry {
workers: DashMap<String, Worker>,
request_history: DashMap<String, Mutex<VecDeque<DateTime<Utc>>>>,
quotas: DashMap<String, WorkerQuotas>,
model_index: DashMap<String, HashSet<String>>,
general_workers: DashMap<String, ()>,
}
impl WorkerRegistry {
pub fn new() -> Self {
Self {
workers: DashMap::new(),
request_history: DashMap::new(),
quotas: DashMap::new(),
model_index: DashMap::new(),
general_workers: DashMap::new(),
}
}
fn index_worker(&self, worker: &Worker) {
match worker.provider_type {
ProviderType::General => {
self.general_workers.insert(worker.id.clone(), ());
}
ProviderType::Specialized => {
for model_uuid in &worker.served_models {
if model_uuid == MODEL_WILDCARD {
self.general_workers.insert(worker.id.clone(), ());
continue;
}
self.model_index
.entry(model_uuid.clone())
.or_default()
.insert(worker.id.clone());
}
}
}
}
fn deindex_worker(&self, worker_id: &str) {
self.general_workers.remove(worker_id);
self.model_index.retain(|_, ids| {
ids.remove(worker_id);
!ids.is_empty()
});
}
pub fn workers_serving(&self, model_uuid: &str) -> Vec<Worker> {
let mut ids: HashSet<String> = self
.model_index
.get(model_uuid)
.map(|entry| entry.clone())
.unwrap_or_default();
for entry in self.general_workers.iter() {
ids.insert(entry.key().clone());
}
ids.into_iter().filter_map(|id| self.get(&id)).collect()
}
pub fn register(&self, registration: WorkerRegistration) -> Worker {
let id = uuid::Uuid::new_v4().to_string();
let wireguard_ip = registration.wireguard_ip.clone().or_else(|| {
let uri_ip = registration
.uri
.strip_prefix("http://")
.or_else(|| registration.uri.strip_prefix("https://"))
.and_then(|rest| {
let host = rest.split('/').next().unwrap_or(rest);
let ip = host.split(':').next().unwrap_or(host);
if ip.is_empty() {
None
} else {
Some(ip.to_string())
}
});
match uri_ip.as_deref() {
Some("127.0.0.1") | Some("::1") | Some("localhost") => {
super::discovery::get_mesh_ip().or(uri_ip)
}
_ => uri_ip,
}
});
let mut worker = Worker::new(
id.clone(),
registration.name,
registration.uri,
registration.worker_type,
);
worker.resources = registration.resources;
worker.pricing = registration.pricing;
worker.tags = registration.tags;
worker.max_timeout_secs = registration.max_timeout_secs;
worker.hardware = registration.hardware;
worker.wireguard_ip = wireguard_ip;
worker.is_docker = registration
.is_docker
.or_else(|| Some(std::path::Path::new("/.dockerenv").exists()));
worker.source_node = registration.source_node;
worker.explicit_local = registration.explicit_local;
worker.provider_type = registration.provider_type;
worker.served_models = registration.served_models;
worker.price_per_mtok = registration.price_per_mtok;
worker.slot = worker
.uri
.strip_prefix("http://")
.or_else(|| worker.uri.strip_prefix("https://"))
.and_then(|rest| rest.split('/').next())
.and_then(|host| host.rsplit(':').next())
.unwrap_or_default()
.to_string();
self.workers.insert(id.clone(), worker.clone());
self.request_history
.insert(id.clone(), Mutex::new(VecDeque::new()));
self.quotas.insert(id.clone(), WorkerQuotas::default());
self.index_worker(&worker);
worker
}
pub fn set_node_fp(&self, worker_id: &str, node_fp: &str) -> Option<Worker> {
self.workers.get_mut(worker_id).map(|mut w| {
w.node_fp = node_fp.to_string();
w.clone()
})
}
pub fn requests_in_window(&self, worker_id: &str, window: Duration) -> u64 {
let cutoff = Utc::now() - window;
self.request_history
.get(worker_id)
.map(|entry| {
let ring = entry.lock().unwrap();
ring.iter().filter(|&&ts| ts >= cutoff).count() as u64
})
.unwrap_or(0)
}
pub fn get_quotas(&self, worker_id: &str) -> WorkerQuotas {
self.quotas
.get(worker_id)
.map(|q| q.clone())
.unwrap_or_default()
}
pub fn heartbeat(&self, heartbeat: WorkerHeartbeat) -> Option<Worker> {
self.workers.get_mut(&heartbeat.worker_id).map(|mut w| {
w.heartbeat();
if let Some(resources) = heartbeat.resources {
w.resources = resources;
}
if let Some(active) = heartbeat.active_requests {
w.active_requests = active;
}
if let Some(status) = heartbeat.status {
w.status = status;
}
if let Some(max_timeout) = heartbeat.max_timeout_secs {
w.max_timeout_secs = max_timeout;
}
w.clone()
})
}
pub fn refresh_heartbeat(&self, id: &str) {
if let Some(mut w) = self.workers.get_mut(id) {
w.heartbeat();
if w.status == WorkerStatus::Unhealthy {
w.status = WorkerStatus::Healthy;
}
}
}
pub fn update_resources(&self, id: &str, resources: WorkerResources, hardware: HardwareInfo) {
if let Some(mut w) = self.workers.get_mut(id) {
w.heartbeat();
w.resources = resources;
if hardware.storage_gb.is_some() {
w.hardware.storage_gb = hardware.storage_gb;
}
if w.status == WorkerStatus::Unhealthy {
w.status = WorkerStatus::Healthy;
}
}
}
pub fn get(&self, id: &str) -> Option<Worker> {
self.workers.get(id).map(|w| w.clone())
}
pub fn remove(&self, id: &str) -> Option<Worker> {
self.deindex_worker(id);
self.workers.remove(id).map(|(_, w)| w)
}
pub fn list(&self) -> Vec<Worker> {
self.workers.iter().map(|w| w.clone()).collect()
}
pub fn aggregate_available(&self) -> BrokerResources {
self.healthy()
.into_iter()
.fold(BrokerResources::default(), |mut acc, w| {
let cpus = w.resources.cpus_available.max(0.0) as u64;
acc.cpus_available = acc.cpus_available.saturating_add(cpus);
acc.memory_available = acc
.memory_available
.saturating_add(w.resources.memory_available);
acc.gpus_available = acc
.gpus_available
.saturating_add(w.resources.gpus_available);
acc.healthy_workers = acc.healthy_workers.saturating_add(1);
acc
})
}
pub fn healthy(&self) -> Vec<Worker> {
self.workers
.iter()
.filter(|w| w.status == WorkerStatus::Healthy)
.map(|w| w.clone())
.collect()
}
pub fn healthy_ids(&self) -> Vec<String> {
self.workers
.iter()
.filter(|w| w.status == WorkerStatus::Healthy)
.map(|w| w.key().clone())
.collect()
}
pub fn find_capable(&self, cpus: f64, memory_bytes: u64, gpus: u32) -> Vec<Worker> {
self.workers
.iter()
.filter(|w| w.can_handle(cpus, memory_bytes, gpus))
.map(|w| w.clone())
.collect()
}
pub fn mark_stale(&self, timeout_secs: i64) {
for mut entry in self.workers.iter_mut() {
if entry.is_stale(timeout_secs) && entry.status == WorkerStatus::Healthy {
entry.status = WorkerStatus::Unhealthy;
}
}
}
pub fn remove_stale(&self, timeout_secs: i64) -> Vec<String> {
let to_remove: Vec<String> = self
.workers
.iter()
.filter(|w| w.status == WorkerStatus::Unhealthy && w.is_stale(timeout_secs))
.map(|w| w.id.clone())
.collect();
for id in &to_remove {
self.deindex_worker(id);
self.workers.remove(id);
self.request_history.remove(id);
self.quotas.remove(id);
}
to_remove
}
pub fn try_reserve_quota(&self, worker_id: &str) -> bool {
let quotas = self.get_quotas(worker_id);
let all_unlimited = quotas.per_5h == 0 && quotas.per_week == 0 && quotas.per_month == 0;
match self.request_history.get(worker_id) {
None => true, Some(entry) => {
let mut ring = entry.lock().unwrap();
let now = Utc::now();
if !all_unlimited {
if quotas.per_5h > 0 {
let cutoff = now - Duration::hours(5);
let count = ring.iter().filter(|&&ts| ts >= cutoff).count() as u64;
if count >= quotas.per_5h {
return false;
}
}
if quotas.per_week > 0 {
let cutoff = now - Duration::weeks(1);
let count = ring.iter().filter(|&&ts| ts >= cutoff).count() as u64;
if count >= quotas.per_week {
return false;
}
}
if quotas.per_month > 0 {
let cutoff = now - Duration::days(30);
let count = ring.iter().filter(|&&ts| ts >= cutoff).count() as u64;
if count >= quotas.per_month {
return false;
}
}
}
ring.push_back(now);
let prune_cutoff = now - Duration::days(30);
while ring.front().map(|&ts| ts < prune_cutoff).unwrap_or(false) {
ring.pop_front();
}
true
}
}
}
pub fn cancel_quota_reservation(&self, worker_id: &str) {
if let Some(entry) = self.request_history.get(worker_id) {
let mut ring = entry.lock().unwrap();
ring.pop_back();
}
}
pub fn record_request(&self, worker_id: &str, duration_ms: f64, _success: bool) {
if let Some(mut worker) = self.workers.get_mut(worker_id) {
worker.total_requests += 1;
let alpha = 0.1;
worker.avg_latency_ms = alpha * duration_ms + (1.0 - alpha) * worker.avg_latency_ms;
if worker.active_requests > 0 {
worker.active_requests -= 1;
}
}
}
pub fn increment_active(&self, worker_id: &str) {
if let Some(mut worker) = self.workers.get_mut(worker_id) {
worker.active_requests += 1;
}
}
pub fn decrement_active(&self, worker_id: &str) {
if let Some(mut worker) = self.workers.get_mut(worker_id) {
if worker.active_requests > 0 {
worker.active_requests -= 1;
}
}
}
pub fn active_guard(&self, worker_id: &str) -> ActiveGuard<'_> {
self.increment_active(worker_id);
ActiveGuard {
registry: self,
worker_id: worker_id.to_string(),
armed: true,
}
}
pub fn count(&self) -> usize {
self.workers.len()
}
pub fn mark_unhealthy(&self, worker_id: &str) {
if let Some(mut worker) = self.workers.get_mut(worker_id) {
worker.status = WorkerStatus::Unhealthy;
}
}
pub fn resolve_model(&self, model_uuid: &str) -> Result<Worker, ModelResolveError> {
let candidates: Vec<Worker> = self
.workers_serving(model_uuid)
.into_iter()
.filter(|w| w.status == WorkerStatus::Healthy)
.collect();
let (specialized, general): (Vec<Worker>, Vec<Worker>) = candidates
.into_iter()
.partition(|w| w.provider_type == ProviderType::Specialized);
let pick_cheapest = |mut tier: Vec<Worker>| -> Option<Worker> {
tier.sort_by(|a, b| {
a.price_per_mtok
.partial_cmp(&b.price_per_mtok)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.name.cmp(&b.name))
});
tier.into_iter().next()
};
if !specialized.is_empty() {
return Ok(pick_cheapest(specialized).expect("non-empty specialized tier"));
}
if !general.is_empty() {
return Ok(pick_cheapest(general).expect("non-empty general tier"));
}
Err(ModelResolveError::NoProvider {
model_uuid: model_uuid.to_string(),
})
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ModelResolveError {
NoProvider { model_uuid: String },
}
impl std::fmt::Display for ModelResolveError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::NoProvider { model_uuid } => {
write!(f, "no provider serving zc://{}", model_uuid)
}
}
}
}
impl std::error::Error for ModelResolveError {}
impl Default for WorkerRegistry {
fn default() -> Self {
Self::new()
}
}
pub struct ActiveGuard<'a> {
registry: &'a WorkerRegistry,
worker_id: String,
armed: bool,
}
impl<'a> ActiveGuard<'a> {
pub fn disarm(&mut self) {
self.armed = false;
}
}
impl<'a> Drop for ActiveGuard<'a> {
fn drop(&mut self) {
if self.armed {
self.registry.decrement_active(&self.worker_id);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn worker_zc_uri_and_source_node_default() {
let mut w = Worker::new(
"id1".into(),
"worker-abc123".into(),
"http://10.13.13.21:3960".into(),
"zakuro".into(),
);
w.node_fp = "abc123".into();
w.slot = "3960".into();
assert_eq!(w.zc_uri(), "zc://worker-abc123-3960");
assert!(w.source_node.is_none());
w.source_node = Some("node-i9".into());
assert_eq!(w.source_node.as_deref(), Some("node-i9"));
}
#[test]
fn worker_zc_uri_is_key_derived_and_collision_free() {
let mut w1 = Worker::new(
"id1".into(),
"worker-lxd".into(),
"http://127.0.0.1:3960".into(),
"zakuro".into(),
);
let mut w2 = Worker::new(
"id2".into(),
"worker-lxd".into(),
"http://127.0.0.1:3960".into(),
"zakuro".into(),
);
w1.node_fp = "aaaaaaaaaaaaaaaa".into();
w1.slot = "3960".into();
w2.node_fp = "bbbbbbbbbbbbbbbb".into();
w2.slot = "3960".into();
assert_eq!(w1.zc_uri(), "zc://worker-aaaaaaaaaaaaaaaa-3960");
assert_ne!(w1.zc_uri(), w2.zc_uri()); }
fn registration(name: &str, uri: &str) -> WorkerRegistration {
WorkerRegistration {
name: name.to_string(),
uri: uri.to_string(),
worker_type: "zakuro".to_string(),
resources: WorkerResources::default(),
pricing: WorkerPricing::default(),
tags: vec![],
max_timeout_secs: 0.0,
hardware: HardwareInfo::default(),
wireguard_ip: None,
is_docker: None,
source_node: None,
explicit_local: false,
provider_type: ProviderType::default(),
served_models: vec![],
price_per_mtok: 0.0,
}
}
fn model_registration(
name: &str,
uri: &str,
provider_type: ProviderType,
served_models: Vec<&str>,
) -> WorkerRegistration {
let mut r = registration(name, uri);
r.provider_type = provider_type;
r.served_models = served_models.into_iter().map(String::from).collect();
r
}
#[test]
fn test_workers_serving_specialized_matches_only_its_model() {
let reg = WorkerRegistry::new();
let w = reg.register(model_registration(
"specialist",
"http://127.0.0.1:4001",
ProviderType::Specialized,
vec!["uuid-a"],
));
let matches = reg.workers_serving("uuid-a");
assert_eq!(matches.len(), 1);
assert_eq!(matches[0].id, w.id);
let no_matches = reg.workers_serving("uuid-b");
assert!(no_matches.is_empty());
}
#[test]
fn test_workers_serving_general_matches_any_model() {
let reg = WorkerRegistry::new();
let w = reg.register(model_registration(
"generalist",
"http://127.0.0.1:4002",
ProviderType::General,
vec!["*"],
));
for uuid in ["uuid-a", "uuid-b", "any-random-uuid"] {
let matches = reg.workers_serving(uuid);
assert_eq!(matches.len(), 1, "general worker should serve {uuid}");
assert_eq!(matches[0].id, w.id);
}
}
#[test]
fn test_workers_serving_unions_specialized_and_general() {
let reg = WorkerRegistry::new();
let specialist = reg.register(model_registration(
"specialist",
"http://127.0.0.1:4003",
ProviderType::Specialized,
vec!["uuid-a"],
));
let generalist = reg.register(model_registration(
"generalist",
"http://127.0.0.1:4004",
ProviderType::General,
vec![],
));
let mut ids: Vec<String> = reg
.workers_serving("uuid-a")
.into_iter()
.map(|w| w.id)
.collect();
ids.sort();
let mut expected = vec![specialist.id.clone(), generalist.id.clone()];
expected.sort();
assert_eq!(ids, expected);
}
#[test]
fn test_reregister_updates_model_index() {
let reg = WorkerRegistry::new();
let w1 = reg.register(model_registration(
"flip",
"http://127.0.0.1:4005",
ProviderType::Specialized,
vec!["uuid-old"],
));
assert_eq!(reg.workers_serving("uuid-old").len(), 1);
reg.remove(&w1.id);
assert!(
reg.workers_serving("uuid-old").is_empty(),
"removed worker must drop out of the old model's index entry"
);
let w2 = reg.register(model_registration(
"flip",
"http://127.0.0.1:4005",
ProviderType::Specialized,
vec!["uuid-new"],
));
assert!(reg.workers_serving("uuid-old").is_empty());
let matches = reg.workers_serving("uuid-new");
assert_eq!(matches.len(), 1);
assert_eq!(matches[0].id, w2.id);
}
#[test]
fn test_remove_drops_worker_from_general_index() {
let reg = WorkerRegistry::new();
let w = reg.register(model_registration(
"generalist",
"http://127.0.0.1:4006",
ProviderType::General,
vec![],
));
assert_eq!(reg.workers_serving("any-uuid").len(), 1);
reg.remove(&w.id);
assert!(reg.workers_serving("any-uuid").is_empty());
}
#[test]
fn test_legacy_worker_json_without_new_fields_defaults_back_compat() {
let legacy_json = r#"{
"id": "legacy-id",
"name": "legacy-worker",
"uri": "http://127.0.0.1:3960",
"worker_type": "zakuro",
"status": "healthy",
"resources": {
"cpus_total": 1.0,
"cpus_available": 1.0,
"memory_total": 1073741824,
"memory_available": 1073741824,
"gpus_total": 0,
"gpus_available": 0
},
"pricing": {
"price_per_hour": 3.6,
"min_charge": 0.001
},
"last_heartbeat": "2026-01-01T00:00:00Z",
"active_requests": 0,
"total_requests": 0,
"avg_latency_ms": 0.0,
"tags": []
}"#;
let w: Worker = serde_json::from_str(legacy_json).expect("legacy JSON must deserialize");
assert_eq!(w.provider_type, ProviderType::Specialized);
assert!(w.served_models.is_empty());
assert_eq!(w.price_per_mtok, 0.0);
}
#[test]
fn test_legacy_worker_registration_json_without_new_fields_defaults_back_compat() {
let legacy_json = r#"{
"name": "legacy-worker",
"uri": "http://127.0.0.1:3960",
"worker_type": "zakuro"
}"#;
let r: WorkerRegistration =
serde_json::from_str(legacy_json).expect("legacy JSON must deserialize");
assert_eq!(r.provider_type, ProviderType::Specialized);
assert!(r.served_models.is_empty());
assert_eq!(r.price_per_mtok, 0.0);
}
#[test]
fn test_register_assigns_unique_id() {
let reg = WorkerRegistry::new();
let w1 = reg.register(registration("w1", "http://127.0.0.1:3960"));
let w2 = reg.register(registration("w2", "http://127.0.0.1:3961"));
assert!(!w1.id.is_empty());
assert_ne!(w1.id, w2.id);
assert_eq!(reg.count(), 2);
}
#[test]
fn test_register_stores_fields() {
let reg = WorkerRegistry::new();
let worker = reg.register(registration("my-worker", "http://10.0.0.1:3960"));
let got = reg.get(&worker.id).unwrap();
assert_eq!(got.name, "my-worker");
assert_eq!(got.uri, "http://10.0.0.1:3960");
assert_eq!(got.status, WorkerStatus::Healthy);
assert_eq!(got.active_requests, 0);
}
#[test]
fn test_remove_worker() {
let reg = WorkerRegistry::new();
let w = reg.register(registration("w", "http://127.0.0.1:3960"));
assert!(reg.get(&w.id).is_some());
let removed = reg.remove(&w.id);
assert!(removed.is_some());
assert!(reg.get(&w.id).is_none());
assert_eq!(reg.count(), 0);
}
#[test]
fn test_heartbeat_updates_resources_and_status() {
let reg = WorkerRegistry::new();
let w = reg.register(registration("w", "http://127.0.0.1:3960"));
let hb = WorkerHeartbeat {
worker_id: w.id.clone(),
resources: Some(WorkerResources {
cpus_total: 8.0,
cpus_available: 4.0,
memory_total: 16 * 1024 * 1024 * 1024,
memory_available: 8 * 1024 * 1024 * 1024,
gpus_total: 2,
gpus_available: 1,
}),
active_requests: Some(5),
status: Some(WorkerStatus::Busy),
max_timeout_secs: Some(120.0),
};
let updated = reg.heartbeat(hb).unwrap();
assert_eq!(updated.resources.cpus_available, 4.0);
assert_eq!(updated.resources.gpus_available, 1);
assert_eq!(updated.active_requests, 5);
assert_eq!(updated.status, WorkerStatus::Busy);
assert_eq!(updated.max_timeout_secs, 120.0);
}
#[test]
fn test_heartbeat_unknown_worker_returns_none() {
let reg = WorkerRegistry::new();
let hb = WorkerHeartbeat {
worker_id: "nonexistent".to_string(),
resources: None,
active_requests: None,
status: None,
max_timeout_secs: None,
};
assert!(reg.heartbeat(hb).is_none());
}
#[test]
fn test_mark_stale_sets_unhealthy() {
let reg = WorkerRegistry::new();
let w = reg.register(registration("w", "http://127.0.0.1:3960"));
reg.mark_stale(-1);
let got = reg.get(&w.id).unwrap();
assert_eq!(got.status, WorkerStatus::Unhealthy);
}
#[test]
fn test_refresh_heartbeat_restores_healthy() {
let reg = WorkerRegistry::new();
let w = reg.register(registration("w", "http://127.0.0.1:3960"));
reg.mark_stale(-1);
assert_eq!(reg.get(&w.id).unwrap().status, WorkerStatus::Unhealthy);
reg.refresh_heartbeat(&w.id);
assert_eq!(reg.get(&w.id).unwrap().status, WorkerStatus::Healthy);
}
#[test]
fn aggregate_available_sums_healthy_workers() {
let reg = WorkerRegistry::new();
let mut r1 = registration("w1", "http://127.0.0.1:3960");
r1.resources = WorkerResources {
cpus_total: 4.0,
cpus_available: 4.0,
memory_total: 8_000_000_000,
memory_available: 8_000_000_000,
gpus_total: 1,
gpus_available: 1,
};
let mut r2 = registration("w2", "http://127.0.0.1:3961");
r2.resources = WorkerResources {
cpus_total: 2.0,
cpus_available: 2.0,
memory_total: 4_000_000_000,
memory_available: 4_000_000_000,
gpus_total: 0,
gpus_available: 0,
};
reg.register(r1);
reg.register(r2);
let agg = reg.aggregate_available();
assert_eq!(agg.cpus_available, 6);
assert_eq!(agg.memory_available, 12_000_000_000);
assert_eq!(agg.gpus_available, 1);
assert_eq!(agg.healthy_workers, 2);
}
#[test]
fn aggregate_available_excludes_unhealthy_and_ignores_empty() {
let reg = WorkerRegistry::new();
let empty = reg.aggregate_available();
assert_eq!(empty, BrokerResources::default());
let w = reg.register(registration("sick", "http://127.0.0.1:3962"));
reg.mark_unhealthy(&w.id);
let agg = reg.aggregate_available();
assert_eq!(agg.healthy_workers, 0);
assert_eq!(agg.cpus_available, 0);
}
#[test]
fn test_healthy_list_excludes_unhealthy() {
let reg = WorkerRegistry::new();
let w1 = reg.register(registration("healthy", "http://127.0.0.1:3960"));
let w2 = reg.register(registration("sick", "http://127.0.0.1:3961"));
reg.mark_unhealthy(&w2.id);
let healthy = reg.healthy();
assert_eq!(healthy.len(), 1);
assert_eq!(healthy[0].id, w1.id);
}
#[test]
fn test_increment_active_and_decrement_on_record() {
let reg = WorkerRegistry::new();
let w = reg.register(registration("w", "http://127.0.0.1:3960"));
reg.increment_active(&w.id);
reg.increment_active(&w.id);
assert_eq!(reg.get(&w.id).unwrap().active_requests, 2);
reg.record_request(&w.id, 100.0, true);
assert_eq!(reg.get(&w.id).unwrap().active_requests, 1);
}
#[test]
fn active_guard_decrements_on_drop_even_without_record() {
let reg = WorkerRegistry::new();
let w = reg.register(registration("w1", "http://127.0.0.1:3960"));
{
let _g = reg.active_guard(&w.id);
assert_eq!(reg.get(&w.id).unwrap().active_requests, 1);
} assert_eq!(
reg.get(&w.id).unwrap().active_requests,
0,
"active_requests must return to 0 when the guard drops on an error path"
);
}
#[test]
fn active_guard_disarm_skips_decrement() {
let reg = WorkerRegistry::new();
let w = reg.register(registration("w2", "http://127.0.0.1:3960"));
{
let mut g = reg.active_guard(&w.id);
assert_eq!(reg.get(&w.id).unwrap().active_requests, 1);
reg.record_request(&w.id, 10.0, true);
g.disarm();
}
assert_eq!(
reg.get(&w.id).unwrap().active_requests,
0,
"disarmed guard must not double-decrement after record_request"
);
}
#[test]
fn test_record_request_latency_ema() {
let reg = WorkerRegistry::new();
let w = reg.register(registration("w", "http://127.0.0.1:3960"));
reg.record_request(&w.id, 200.0, true);
let got = reg.get(&w.id).unwrap();
assert!((got.avg_latency_ms - 20.0).abs() < 0.001);
assert_eq!(got.total_requests, 1);
}
#[test]
fn test_find_capable_filters_by_resources() {
let reg = WorkerRegistry::new();
let mut r_small = registration("small", "http://127.0.0.1:3960");
r_small.resources = WorkerResources {
cpus_total: 2.0,
cpus_available: 2.0,
memory_total: 2 * 1024 * 1024 * 1024,
memory_available: 2 * 1024 * 1024 * 1024,
gpus_total: 0,
gpus_available: 0,
};
let mut r_large = registration("large", "http://127.0.0.1:3961");
r_large.resources = WorkerResources {
cpus_total: 32.0,
cpus_available: 32.0,
memory_total: 128 * 1024 * 1024 * 1024,
memory_available: 128 * 1024 * 1024 * 1024,
gpus_total: 4,
gpus_available: 4,
};
reg.register(r_small);
reg.register(r_large);
let capable = reg.find_capable(16.0, 1024 * 1024 * 1024, 0);
assert_eq!(capable.len(), 1);
assert_eq!(capable[0].name, "large");
assert_eq!(reg.find_capable(1.0, 512 * 1024 * 1024, 0).len(), 2);
assert_eq!(reg.find_capable(1.0, 512 * 1024 * 1024, 8).len(), 0);
}
#[test]
fn test_can_handle_exact_match() {
let w = Worker::new(
"id".to_string(),
"w".to_string(),
"http://x".to_string(),
"zakuro".to_string(),
);
assert!(w.can_handle(1.0, 1024 * 1024 * 1024, 0));
}
#[test]
fn test_can_handle_over_cpu_fails() {
let w = Worker::new(
"id".to_string(),
"w".to_string(),
"http://x".to_string(),
"zakuro".to_string(),
);
assert!(!w.can_handle(2.0, 512 * 1024 * 1024, 0));
}
#[test]
fn test_can_handle_over_memory_fails() {
let w = Worker::new(
"id".to_string(),
"w".to_string(),
"http://x".to_string(),
"zakuro".to_string(),
);
assert!(!w.can_handle(0.5, 2 * 1024 * 1024 * 1024, 0));
}
#[test]
fn test_can_handle_gpu_required_but_none_fails() {
let w = Worker::new(
"id".to_string(),
"w".to_string(),
"http://x".to_string(),
"zakuro".to_string(),
);
assert!(!w.can_handle(0.5, 512 * 1024 * 1024, 1));
}
#[test]
fn test_can_handle_unhealthy_worker_fails() {
let mut w = Worker::new(
"id".to_string(),
"w".to_string(),
"http://x".to_string(),
"zakuro".to_string(),
);
w.status = WorkerStatus::Unhealthy;
assert!(!w.can_handle(0.1, 1024, 0));
}
#[test]
fn test_pricing_cost_formula() {
let p = WorkerPricing {
price_per_hour: 3.6,
min_charge: 0.001,
};
let cost = p.estimate_cost(10.0);
assert!((cost - 0.010).abs() < 0.0001);
}
#[test]
fn test_pricing_min_charge_enforced() {
let p = WorkerPricing {
price_per_hour: 0.0,
min_charge: 0.005,
};
let cost = p.estimate_cost(0.001);
assert_eq!(cost, 0.005);
}
fn priced_model_registration(
name: &str,
uri: &str,
provider_type: ProviderType,
served_models: Vec<&str>,
price_per_mtok: f64,
) -> WorkerRegistration {
let mut r = model_registration(name, uri, provider_type, served_models);
r.price_per_mtok = price_per_mtok;
r
}
#[test]
fn test_resolve_model_prefers_specialized_over_general() {
let reg = WorkerRegistry::new();
let specialized = reg.register(priced_model_registration(
"specialist",
"http://127.0.0.1:5001",
ProviderType::Specialized,
vec!["uuid-a"],
10.0,
));
reg.register(priced_model_registration(
"generalist",
"http://127.0.0.1:5002",
ProviderType::General,
vec![],
1.0, ));
let resolved = reg.resolve_model("uuid-a").unwrap();
assert_eq!(resolved.id, specialized.id);
}
#[test]
fn test_resolve_model_cheapest_among_specialized() {
let reg = WorkerRegistry::new();
reg.register(priced_model_registration(
"expensive",
"http://127.0.0.1:5003",
ProviderType::Specialized,
vec!["uuid-a"],
20.0,
));
let cheap = reg.register(priced_model_registration(
"cheap",
"http://127.0.0.1:5004",
ProviderType::Specialized,
vec!["uuid-a"],
5.0,
));
let resolved = reg.resolve_model("uuid-a").unwrap();
assert_eq!(resolved.id, cheap.id);
}
#[test]
fn test_resolve_model_general_fallback_cheapest() {
let reg = WorkerRegistry::new();
reg.register(priced_model_registration(
"general-expensive",
"http://127.0.0.1:5005",
ProviderType::General,
vec![],
8.0,
));
let cheap_general = reg.register(priced_model_registration(
"general-cheap",
"http://127.0.0.1:5006",
ProviderType::General,
vec![],
2.0,
));
let resolved = reg.resolve_model("uuid-z").unwrap();
assert_eq!(resolved.id, cheap_general.id);
}
#[test]
fn test_resolve_model_excludes_unhealthy_specialized_falls_back_to_general() {
let reg = WorkerRegistry::new();
let stale_specialist = reg.register(priced_model_registration(
"stale-specialist",
"http://127.0.0.1:5007",
ProviderType::Specialized,
vec!["uuid-a"],
1.0, ));
reg.mark_unhealthy(&stale_specialist.id);
let healthy_general = reg.register(priced_model_registration(
"healthy-general",
"http://127.0.0.1:5008",
ProviderType::General,
vec![],
5.0,
));
let resolved = reg.resolve_model("uuid-a").unwrap();
assert_eq!(resolved.id, healthy_general.id);
}
#[test]
fn test_resolve_model_no_provider_when_no_healthy_candidate() {
let reg = WorkerRegistry::new();
let w = reg.register(priced_model_registration(
"only-specialist",
"http://127.0.0.1:5009",
ProviderType::Specialized,
vec!["uuid-a"],
1.0,
));
reg.mark_unhealthy(&w.id);
let err = reg.resolve_model("uuid-a").unwrap_err();
assert_eq!(
err,
ModelResolveError::NoProvider {
model_uuid: "uuid-a".to_string()
}
);
assert_eq!(err.to_string(), "no provider serving zc://uuid-a");
}
#[test]
fn test_pricing_price_score_weighted() {
let p = WorkerPricing::default();
let score = p.price_score();
assert!((score - 0.010).abs() < 0.0001);
}
}