use std::collections::HashMap;
use std::path::PathBuf;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Instant;
use super::engine::Engine;
use super::schema::OverflowPolicy;
use crate::core::hardware::HardwareProfile;
use crate::inference::models::bert::config::PoolingType;
use crate::inference::models::bert::weights::LoadedBertWeights;
use crate::inference::models::bert::BertConfig;
use crate::inference::models::nomic_bert::{LoadedNomicBertWeights, NomicBertConfig};
use crate::inference::vision::mmproj::{ArchProfile, MmprojConfig};
use crate::inference::vision::mmproj_weights::LoadedMmprojWeights;
use crate::serve::cache::ModelCache;
use crate::serve::multi_model::{
DefaultModelLoader, HotSwapManager, LoadedPool, RestoreErrorKind, RestoreOutcome,
SpillErrorKind, SpillOutcome,
};
use crate::serve::quant_select::QuantType;
#[derive(Debug, Clone)]
pub struct ServerConfig {
pub host: String,
pub port: u16,
pub auth_token: Option<String>,
pub cors_allowed_origins: Vec<String>,
pub queue_capacity: usize,
pub max_concurrent_requests: usize,
pub request_timeout_seconds: u64,
pub default_overflow_policy: OverflowPolicy,
pub cache_dir: Option<PathBuf>,
pub system_fingerprint: Option<String>,
}
impl Default for ServerConfig {
fn default() -> Self {
Self {
host: "127.0.0.1".to_string(),
port: 8080,
auth_token: None,
cors_allowed_origins: Vec::new(),
queue_capacity: 32,
max_concurrent_requests: 0,
request_timeout_seconds: 0,
default_overflow_policy: OverflowPolicy::Summarize,
cache_dir: default_cache_dir(),
system_fingerprint: None,
}
}
}
pub fn default_cache_dir() -> Option<PathBuf> {
std::env::var_os("HOME")
.map(PathBuf::from)
.map(|h| h.join(".cache").join("hf2q"))
}
fn synthetic_cache_root() -> PathBuf {
use std::sync::atomic::{AtomicU64, Ordering};
static COUNTER: AtomicU64 = AtomicU64::new(0);
let id = COUNTER.fetch_add(1, Ordering::Relaxed);
let pid = std::process::id();
let mut p = std::env::temp_dir();
p.push(format!("hf2q-test-cache-{pid}-{id}"));
p
}
#[derive(Debug, Default)]
pub struct ServerMetrics {
pub requests_total: AtomicU64,
pub chat_completions_started: AtomicU64,
pub chat_completions_completed: AtomicU64,
pub chat_completions_queue_full: AtomicU64,
pub sse_cancellations: Arc<AtomicU64>,
pub decode_tokens_total: AtomicU64,
pub prompt_tokens_total: AtomicU64,
pub requests_rejected_total: AtomicU64,
}
impl ServerMetrics {
pub fn sse_cancellations_counter_arc(&self) -> Arc<AtomicU64> {
Arc::clone(&self.sse_cancellations)
}
}
pub const KV_SPILL_OUTCOMES: &[&str] = &["success", "codec_err", "io_err", "parity_fail"];
const KV_OUTCOME_SUCCESS: usize = 0;
const KV_OUTCOME_CODEC_ERR: usize = 1;
const KV_OUTCOME_IO_ERR: usize = 2;
const KV_OUTCOME_PARITY_FAIL: usize = 3;
const KV_OUTCOME_COUNT: usize = 4;
#[derive(Debug, Default)]
pub struct KvSpillCounters {
spills: std::sync::Mutex<HashMap<(String, String), [AtomicU64; KV_OUTCOME_COUNT]>>,
restores: std::sync::Mutex<HashMap<(String, String), [AtomicU64; KV_OUTCOME_COUNT]>>,
server_timing_enabled: AtomicBool,
quarantines: [AtomicU64; KV_QUARANTINE_REASON_COUNT],
evictions: [AtomicU64; KV_EVICTION_TRIGGER_COUNT],
lcp_lookups_total: AtomicU64,
lcp_detected_total: AtomicU64,
}
pub use crate::serve::kv_persist::metrics::{
KvCacheMetricsSink, KvQuarantineReason, KV_EVICTION_TRIGGERS, KV_EVICTION_TRIGGER_COUNT,
KV_QUARANTINE_REASONS, KV_QUARANTINE_REASON_COUNT,
};
const KV_EVICTION_TRIGGER_BUDGET_OVERFLOW: usize = 0;
impl KvSpillCounters {
pub fn new() -> Self {
Self::default()
}
fn spill_outcome_index(outcome: SpillOutcome) -> Option<usize> {
match outcome {
SpillOutcome::Skipped => None,
SpillOutcome::EnqueuedBlocks(_) => Some(KV_OUTCOME_SUCCESS),
SpillOutcome::Error(SpillErrorKind::CodecErr) => Some(KV_OUTCOME_CODEC_ERR),
SpillOutcome::Error(SpillErrorKind::IoErr) => Some(KV_OUTCOME_IO_ERR),
SpillOutcome::Error(SpillErrorKind::ParityFail) => Some(KV_OUTCOME_PARITY_FAIL),
}
}
fn restore_outcome_index(outcome: RestoreOutcome) -> Option<usize> {
match outcome {
RestoreOutcome::Skipped => None,
RestoreOutcome::RestoredBlocks(_) => Some(KV_OUTCOME_SUCCESS),
RestoreOutcome::Error(RestoreErrorKind::CodecErr) => Some(KV_OUTCOME_CODEC_ERR),
RestoreOutcome::Error(RestoreErrorKind::IoErr) => Some(KV_OUTCOME_IO_ERR),
RestoreOutcome::Error(RestoreErrorKind::ParityFail) => Some(KV_OUTCOME_PARITY_FAIL),
}
}
fn new_row() -> [AtomicU64; KV_OUTCOME_COUNT] {
[
AtomicU64::new(0),
AtomicU64::new(0),
AtomicU64::new(0),
AtomicU64::new(0),
]
}
pub fn record_spill(&self, repo: &str, quant: QuantType, outcome: SpillOutcome) {
let Some(idx) = Self::spill_outcome_index(outcome) else {
return; };
let key = (repo.to_string(), quant.as_str().to_string());
let mut guard = self.spills.lock().expect("kv_spill_counters poisoned");
let row = guard.entry(key).or_insert_with(Self::new_row);
row[idx].fetch_add(1, Ordering::Relaxed);
}
pub fn record_restore(&self, repo: &str, quant: QuantType, outcome: RestoreOutcome) {
let Some(idx) = Self::restore_outcome_index(outcome) else {
return; };
let key = (repo.to_string(), quant.as_str().to_string());
let mut guard = self.restores.lock().expect("kv_spill_counters poisoned");
let row = guard.entry(key).or_insert_with(Self::new_row);
row[idx].fetch_add(1, Ordering::Relaxed);
}
pub fn snapshot_spills(&self) -> Vec<((String, String), [u64; KV_OUTCOME_COUNT])> {
let guard = self.spills.lock().expect("kv_spill_counters poisoned");
let mut out: Vec<_> = guard
.iter()
.map(|(k, row)| {
(
k.clone(),
[
row[0].load(Ordering::Relaxed),
row[1].load(Ordering::Relaxed),
row[2].load(Ordering::Relaxed),
row[3].load(Ordering::Relaxed),
],
)
})
.collect();
out.sort_by(|a, b| a.0.cmp(&b.0));
out
}
pub fn snapshot_restores(&self) -> Vec<((String, String), [u64; KV_OUTCOME_COUNT])> {
let guard = self.restores.lock().expect("kv_spill_counters poisoned");
let mut out: Vec<_> = guard
.iter()
.map(|(k, row)| {
(
k.clone(),
[
row[0].load(Ordering::Relaxed),
row[1].load(Ordering::Relaxed),
row[2].load(Ordering::Relaxed),
row[3].load(Ordering::Relaxed),
],
)
})
.collect();
out.sort_by(|a, b| a.0.cmp(&b.0));
out
}
pub fn server_timing_enabled(&self) -> bool {
self.server_timing_enabled.load(Ordering::Acquire)
}
pub fn set_server_timing_enabled(&self, enabled: bool) {
self.server_timing_enabled.store(enabled, Ordering::Release);
}
pub fn snapshot_quarantines(&self) -> [u64; KV_QUARANTINE_REASON_COUNT] {
let mut out = [0u64; KV_QUARANTINE_REASON_COUNT];
for (i, slot) in out.iter_mut().enumerate() {
*slot = self.quarantines[i].load(Ordering::Relaxed);
}
out
}
pub fn snapshot_evictions(&self) -> [u64; KV_EVICTION_TRIGGER_COUNT] {
[self.evictions[KV_EVICTION_TRIGGER_BUDGET_OVERFLOW].load(Ordering::Relaxed)]
}
pub fn snapshot_lcp(&self) -> (u64, u64) {
(
self.lcp_lookups_total.load(Ordering::Relaxed),
self.lcp_detected_total.load(Ordering::Relaxed),
)
}
}
impl KvCacheMetricsSink for KvSpillCounters {
fn record_quarantine(&self, reason: KvQuarantineReason) {
self.quarantines[reason.index()].fetch_add(1, Ordering::Relaxed);
}
fn record_eviction_budget_overflow(&self) {
self.evictions[KV_EVICTION_TRIGGER_BUDGET_OVERFLOW].fetch_add(1, Ordering::Relaxed);
}
fn record_lcp_probe(&self, detected_k: Option<usize>) {
self.lcp_lookups_total.fetch_add(1, Ordering::Relaxed);
if detected_k.is_some() {
self.lcp_detected_total.fetch_add(1, Ordering::Relaxed);
}
}
}
#[derive(Clone)]
pub struct AppState {
pub config: Arc<ServerConfig>,
pub started_at: Arc<Instant>,
pub ready_for_gen: Arc<AtomicBool>,
pub request_counter: Arc<AtomicU64>,
pub pool: Arc<std::sync::RwLock<HotSwapManager<Engine>>>,
pub cache: Arc<std::sync::Mutex<ModelCache>>,
pub hardware: Arc<HardwareProfile>,
pub no_integrity: bool,
pub engine_queue_capacity: usize,
pub default_model: Option<String>,
pub embedding_config: Option<EmbeddingModel>,
pub embedding_registry: Option<Arc<std::sync::Mutex<mlx_native::KernelRegistry>>>,
pub mmproj: Option<LoadedMmproj>,
pub metrics: Arc<ServerMetrics>,
pub kv_spill_counters: Arc<KvSpillCounters>,
pub kv_disk_store: Option<Arc<crate::serve::kv_persist::DiskBlockStore>>,
pub kv_spiller: Option<
Arc<crate::serve::kv_persist::BlockPrefixCacheSpiller<crate::serve::api::engine::Engine>>,
>,
}
#[derive(Clone)]
pub struct EmbeddingModel {
pub gguf_path: PathBuf,
pub vocab: Arc<crate::inference::models::bert::BertVocab>,
pub tokenizer: Arc<crate::inference::models::bert::BertWpmTokenizer>,
pub model_id: String,
pub arch: Option<EmbeddingArch>,
}
#[derive(Debug, Clone)]
pub enum EmbeddingArch {
Bert {
config: BertConfig,
weights: Arc<LoadedBertWeights>,
},
NomicBert {
config: NomicBertConfig,
weights: Arc<LoadedNomicBertWeights>,
},
}
impl EmbeddingArch {
pub fn hidden_size(&self) -> usize {
match self {
Self::Bert { config, .. } => config.hidden_size,
Self::NomicBert { config, .. } => config.hidden_size,
}
}
pub fn max_position_embeddings(&self) -> usize {
match self {
Self::Bert { config, .. } => config.max_position_embeddings,
Self::NomicBert { config, .. } => config.max_position_embeddings,
}
}
pub fn pooling_type(&self) -> PoolingType {
match self {
Self::Bert { config, .. } => config.pooling_type,
Self::NomicBert { config, .. } => config.pooling_type,
}
}
pub fn arch_name(&self) -> &'static str {
match self {
Self::Bert { .. } => "bert",
Self::NomicBert { .. } => "nomic-bert",
}
}
}
impl std::fmt::Debug for EmbeddingModel {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("EmbeddingModel")
.field("gguf_path", &self.gguf_path)
.field("arch", &self.arch.as_ref().map(|a| a.arch_name()))
.field("hidden", &self.arch.as_ref().map(|a| a.hidden_size()))
.field("vocab_len", &self.vocab.len())
.field("model_id", &self.model_id)
.finish()
}
}
impl EmbeddingModel {
pub fn encode(&self, input: &str, add_special_tokens: bool) -> Vec<u32> {
self.tokenizer.encode(input, add_special_tokens)
}
}
#[derive(Debug, Clone)]
pub struct LoadedMmproj {
pub gguf_path: PathBuf,
pub config: MmprojConfig,
pub arch: ArchProfile,
pub weights: Arc<LoadedMmprojWeights>,
pub model_id: String,
}
impl AppState {
pub fn new_for_serve(
config: ServerConfig,
no_integrity: bool,
engine_queue_capacity: usize,
default_model: Option<String>,
) -> anyhow::Result<Self> {
let hardware = crate::core::hardware::HardwareProfiler::detect()
.map_err(|e| anyhow::anyhow!("hardware detection: {e}"))?;
let cache = ModelCache::open()?;
let pool = LoadedPool::from_hardware(&hardware);
let kv_spill_counters = Arc::new(KvSpillCounters::new());
let mut manager = HotSwapManager::new(pool, Arc::new(DefaultModelLoader));
manager.set_kv_counters(Arc::clone(&kv_spill_counters));
Ok(Self {
config: Arc::new(config),
started_at: Arc::new(Instant::now()),
ready_for_gen: Arc::new(AtomicBool::new(true)),
request_counter: Arc::new(AtomicU64::new(0)),
pool: Arc::new(std::sync::RwLock::new(manager)),
cache: Arc::new(std::sync::Mutex::new(cache)),
hardware: Arc::new(hardware),
no_integrity,
engine_queue_capacity,
default_model,
embedding_config: None,
embedding_registry: None,
mmproj: None,
metrics: Arc::new(ServerMetrics::default()),
kv_spill_counters,
kv_disk_store: None,
kv_spiller: None,
})
}
pub fn new(config: ServerConfig) -> Self {
let pool = LoadedPool::with_capacity_and_budget(3, 1u64 << 30);
let kv_spill_counters_test = Arc::new(KvSpillCounters::new());
let mut manager = HotSwapManager::new(pool, Arc::new(DefaultModelLoader));
manager.set_kv_counters(Arc::clone(&kv_spill_counters_test));
let cache = ModelCache::open_at(synthetic_cache_root())
.expect("open synthetic cache for AppState::new (test path)");
let hardware = HardwareProfile {
chip_model: "Synthetic-Test".into(),
total_memory_bytes: 16u64 << 30,
available_memory_bytes: 16u64 << 30,
performance_cores: 8,
efficiency_cores: 4,
total_cores: 12,
memory_bandwidth_gbs: 400.0,
};
Self {
config: Arc::new(config),
started_at: Arc::new(Instant::now()),
ready_for_gen: Arc::new(AtomicBool::new(true)),
request_counter: Arc::new(AtomicU64::new(0)),
pool: Arc::new(std::sync::RwLock::new(manager)),
cache: Arc::new(std::sync::Mutex::new(cache)),
hardware: Arc::new(hardware),
no_integrity: false,
engine_queue_capacity: 32,
default_model: None,
embedding_config: None,
embedding_registry: None,
mmproj: None,
metrics: Arc::new(ServerMetrics::default()),
kv_spill_counters: kv_spill_counters_test,
kv_disk_store: None,
kv_spiller: None,
}
}
pub fn with_default_model(mut self, default_model: Option<String>) -> Self {
self.default_model = default_model;
self
}
pub fn with_kv_disk_store(
mut self,
store: Arc<crate::serve::kv_persist::DiskBlockStore>,
) -> Self {
self.kv_disk_store = Some(store);
self
}
pub fn with_kv_spiller(
mut self,
spiller: Arc<
crate::serve::kv_persist::BlockPrefixCacheSpiller<crate::serve::api::engine::Engine>,
>,
) -> Self {
self.kv_spiller = Some(spiller);
self
}
pub fn with_embedding_model(mut self, em: EmbeddingModel) -> Self {
self.embedding_config = Some(em);
self
}
pub fn with_embedding_registry(
mut self,
registry: Arc<std::sync::Mutex<mlx_native::KernelRegistry>>,
) -> Self {
self.embedding_registry = Some(registry);
self
}
pub fn with_mmproj(mut self, m: LoadedMmproj) -> Self {
self.mmproj = Some(m);
self
}
pub fn uptime_seconds(&self) -> u64 {
self.started_at.elapsed().as_secs()
}
pub fn mark_ready_for_gen(&self) {
self.ready_for_gen.store(true, Ordering::Release);
}
pub fn mark_not_ready(&self) {
self.ready_for_gen.store(false, Ordering::Release);
}
pub fn is_ready_for_gen(&self) -> bool {
self.ready_for_gen.load(Ordering::Acquire)
}
pub fn next_request_seq(&self) -> u64 {
self.request_counter.fetch_add(1, Ordering::Relaxed)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn default_config_uses_localhost() {
let cfg = ServerConfig::default();
assert_eq!(cfg.host, "127.0.0.1");
assert_eq!(cfg.port, 8080);
assert!(cfg.auth_token.is_none());
assert!(cfg.cors_allowed_origins.is_empty());
assert_eq!(cfg.queue_capacity, 32);
assert_eq!(cfg.default_overflow_policy, OverflowPolicy::Summarize);
}
#[test]
fn app_state_starts_ready_in_iter_2() {
let state = AppState::new(ServerConfig::default());
assert!(state.is_ready_for_gen());
assert_eq!(state.uptime_seconds(), 0);
}
#[test]
fn mark_not_ready_flips_to_false_then_back() {
let state = AppState::new(ServerConfig::default());
assert!(state.is_ready_for_gen());
state.mark_not_ready();
assert!(!state.is_ready_for_gen());
state.mark_ready_for_gen();
assert!(state.is_ready_for_gen());
}
#[test]
fn request_seq_is_monotonic() {
let state = AppState::new(ServerConfig::default());
let a = state.next_request_seq();
let b = state.next_request_seq();
let c = state.next_request_seq();
assert_eq!(a + 1, b);
assert_eq!(b + 1, c);
}
#[test]
fn embedding_model_encode_round_trips_hello() {
use crate::inference::models::bert::{
build_wordpiece_tokenizer, BertSpecialTokens, BertVocab,
};
let vocab = BertVocab {
tokens: vec![
"[UNK]".into(),
"[CLS]".into(),
"[SEP]".into(),
"[PAD]".into(),
"\u{2581}hello".into(),
"\u{2581}world".into(),
],
specials: BertSpecialTokens {
cls: 1,
sep: 2,
pad: 3,
unk: 0,
mask: 0,
},
};
let tokenizer = build_wordpiece_tokenizer(&vocab).expect("build");
let em = EmbeddingModel {
gguf_path: "/tmp/synthetic.gguf".into(),
vocab: Arc::new(vocab.clone()),
tokenizer: Arc::new(crate::inference::models::bert::BertWpmTokenizer::new(
&vocab,
)),
model_id: "synthetic-embed".into(),
arch: None,
};
let _ = tokenizer; let ids = em.encode("hello world", false);
assert!(ids.contains(&4), "expected 'hello'=4 in {:?}", ids);
assert!(ids.contains(&5), "expected 'world'=5 in {:?}", ids);
}
#[test]
fn with_mmproj_attaches_descriptor_to_state() {
use crate::inference::vision::mmproj::{MmprojConfig, ProjectorType};
let cfg = MmprojConfig {
image_size: 896,
patch_size: 14,
num_patches_side: 64,
hidden_size: 1152,
intermediate_size: 4304,
num_attention_heads: 16,
num_hidden_layers: 27,
layer_norm_eps: 1e-6,
projector: ProjectorType::Mlp,
image_mean: [0.5, 0.5, 0.5],
image_std: [0.5, 0.5, 0.5],
spatial_merge_size: None,
projection_dim: None,
deepstack_indexes: None,
};
let device = mlx_native::MlxDevice::new().expect("create device");
let m = LoadedMmproj {
gguf_path: "/tmp/synthetic-mmproj.gguf".into(),
config: cfg.clone(),
arch: ArchProfile::Gemma4Siglip,
weights: Arc::new(LoadedMmprojWeights::empty(device)),
model_id: "synthetic-mmproj".into(),
};
let state = AppState::new(ServerConfig::default()).with_mmproj(m);
let attached = state.mmproj.as_ref().expect("mmproj should be Some");
assert_eq!(attached.model_id, "synthetic-mmproj");
assert_eq!(
attached.gguf_path.file_name().unwrap(),
"synthetic-mmproj.gguf"
);
assert_eq!(attached.config, cfg);
assert_eq!(attached.arch, ArchProfile::Gemma4Siglip);
assert!(attached.config.projector.is_supported());
}
#[test]
fn kv_spill_counters_lcp_probe_records_lookups_unconditionally() {
let counters = KvSpillCounters::new();
counters.record_lcp_probe(None);
counters.record_lcp_probe(None);
counters.record_lcp_probe(None);
let (lookups, detected) = counters.snapshot_lcp();
assert_eq!(lookups, 3, "every probe must bump lookups_total");
assert_eq!(detected, 0, "None outcome must NOT bump detected_total");
}
#[test]
fn kv_spill_counters_lcp_probe_increments_detected_on_some() {
let counters = KvSpillCounters::new();
counters.record_lcp_probe(Some(5));
counters.record_lcp_probe(Some(127));
counters.record_lcp_probe(None); counters.record_lcp_probe(Some(42));
let (lookups, detected) = counters.snapshot_lcp();
assert_eq!(lookups, 4, "all 4 probes must bump lookups_total");
assert_eq!(
detected, 3,
"3 Some outcomes must bump detected_total; 1 None must not"
);
}
#[test]
fn kv_spill_counters_lcp_probe_starts_at_zero() {
let counters = KvSpillCounters::new();
assert_eq!(counters.snapshot_lcp(), (0, 0));
}
}