use std::collections::HashMap;
use std::io::Write as _;
use std::sync::Arc;
use std::sync::mpsc::{Receiver, Sender};
use std::time::Instant;
use cudarc::driver::CudaSlice;
use memra_engine::Engine;
use memra_engine::cache::Cache;
use memra_engine::decode::{GenParams, StopReason};
use memra_engine::hybrid::HybridModel;
use memra_engine::sampler::{Sampler, SamplerConfig};
use memra_gguf::GgufFile;
use memra_tokenizer::Tokenizer;
pub const MAX_ACTIVE: usize = 4;
pub const MAX_NEW_CTX_BOUNDED: usize = usize::MAX;
const PREFILL_TICK_T: usize = 1024;
struct LoadedModel {
model: HybridModel,
tok: Tokenizer,
eos_id: u32,
from_dir: bool,
constraints: std::cell::OnceCell<Result<crate::constrained::ConstraintFactory, String>>,
}
#[derive(Debug, Clone)]
pub enum Event {
Token { id: u32, text: String },
Done { stop_reason: String, n_tokens: usize, n_prompt: usize, n_cached: usize,
elapsed_s: f64, spec: Option<SpecUsage> },
Error(String),
}
#[derive(Debug, Clone, Copy, Default)]
pub struct SpecUsage {
pub rounds: u64,
pub drafted: u64,
pub accepted: u64,
}
pub struct Request {
pub model: String,
pub prompt_ids: Vec<u32>, pub prompt_text: String,
pub chat: bool,
pub chat_turns: Vec<memra_tokenizer::chat::Turn>,
pub tools_json: Vec<String>,
pub think: memra_tokenizer::chat::ThinkMode,
pub params: GenParams,
pub sampler_cfg: SamplerConfig,
pub stop_strings: Vec<String>,
pub trace_id: Option<String>,
pub cache_ns: String,
pub affinity: Option<String>,
pub lane: crate::lanes::Lane,
pub oom_retries: u32,
pub grammar: Option<crate::constrained::GrammarSpec>,
pub tx: tokio::sync::mpsc::UnboundedSender<Event>,
}
#[derive(Debug, Clone, Default)]
pub struct ModelCaps {
pub tools_branch: bool,
pub qwen_think: bool,
pub think_switch: bool,
pub chat_ok: bool,
pub context_length: usize,
pub tokenizer: String,
pub instruct_type: Option<String>,
}
pub enum Cmd {
Generate(Box<Request>),
}
pub static PENDING_ADMITS: std::sync::atomic::AtomicUsize =
std::sync::atomic::AtomicUsize::new(0);
#[derive(Clone, Default)]
pub struct Metrics {
pub admitted: u64,
pub completed: u64,
pub tokens_out: u64,
pub step_p50_ms: f32,
pub step_p99_ms: f32,
pub prompt_tokens_in: u64,
pub cached_tokens_in: u64,
pub prefix_hits: u64,
pub prefix_entries: u64,
pub prefix_bytes: u64,
pub lane_admitted: [u64; 3],
pub lane_shed: [u64; 3],
pub lane_completed: [u64; 3],
pub lane_tokens: [u64; 3],
pub batch_size_last: usize,
pub spec: HashMap<String, memra_engine::spec::SpecTelemetry>,
}
pub type SharedMetrics = std::sync::Arc<std::sync::Mutex<Metrics>>;
use crate::lanes::StepStats;
struct ReuseEntry {
fed: Vec<u32>,
cache: Cache,
last_logits: Vec<f32>,
cap: usize,
}
struct SpecReuseEntry {
sess: memra_engine::spec::SpecSession,
committed_text: String,
affinity: Option<String>,
fingerprint: Vec<u64>,
}
fn reuse_pool_per_model() -> usize {
static N: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*N.get_or_init(|| std::env::var("MEMRA_REUSE_POOL").ok()
.and_then(|v| v.parse().ok()).unwrap_or(2))
}
const REUSE_MIN_PREFIX: usize = 16;
#[derive(Default)]
struct SpecSizing {
evict_first: std::collections::HashSet<String>,
learned_ctx: HashMap<String, usize>,
}
const SPEC_SHRINK_SLACK: usize = 64;
const SPEC_SHRINK_RESERVE: usize = 1536 << 20;
type PoolKey = (String, String);
fn ns_suffix(ns: &str) -> String {
if ns.is_empty() { String::new() } else { format!(", ns {ns:?}") }
}
const FP_WINDOW: usize = 8;
const FP_MIN_SEGMENTS: usize = 3;
fn affinity_enabled() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_AFFINITY").map(|v| v != "0").unwrap_or(true))
}
fn fnv1a(seed: u64, toks: &[u32]) -> u64 {
let mut h = seed;
for &t in toks {
for b in t.to_le_bytes() {
h ^= b as u64;
h = h.wrapping_mul(0x100000001b3);
}
}
h
}
fn conversation_fingerprint(
toks: &[u32],
is_boundary: &dyn Fn(u32) -> bool,
drop_live: bool,
) -> Vec<u64> {
let mut segs: Vec<(usize, usize)> = Vec::new();
let mut start = 0usize;
for (i, &t) in toks.iter().enumerate() {
if is_boundary(t) && i > start {
segs.push((start, i));
start = i;
}
}
if start < toks.len() {
segs.push((start, toks.len()));
}
if drop_live && !segs.is_empty() {
segs.pop();
}
segs.iter()
.map(|&(lo, hi)| {
let seg = &toks[lo..hi];
let head = &seg[..FP_WINDOW.min(seg.len())];
let tail = &seg[seg.len().saturating_sub(FP_WINDOW)..];
fnv1a(fnv1a(0xcbf29ce484222325, head), tail)
})
.collect()
}
fn fingerprint_affinity(a: &[u64], b: &[u64]) -> usize {
a.iter().zip(b).take_while(|(x, y)| x == y).count()
}
#[derive(PartialEq, Eq, Debug)]
enum AffinityMatch {
Exact { suffix_from: usize },
Diverged { at: usize },
}
fn affinity_match(prompt: &[u32], committed: &[u32]) -> AffinityMatch {
let n = committed.len().min(prompt.len());
for i in 0..n {
if prompt[i] != committed[i] {
return AffinityMatch::Diverged { at: i };
}
}
if prompt.len() < committed.len() {
return AffinityMatch::Diverged { at: prompt.len() };
}
AffinityMatch::Exact { suffix_from: committed.len() }
}
const PREFIX_CACHE_MIN_TOKENS: usize = 64;
fn prefix_cache_budget_bytes() -> usize {
static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*B.get_or_init(|| {
std::env::var("MEMRA_PREFIX_CACHE_MB").ok()
.and_then(|v| v.parse::<usize>().ok()).unwrap_or(256)
.saturating_mul(1024 * 1024)
})
}
fn serve_batching() -> bool {
static B: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*B.get_or_init(|| std::env::var("MEMRA_SERVE_BATCH").map(|v| v != "0").unwrap_or(true))
}
fn serve_spec_enabled() -> bool {
static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*S.get_or_init(|| std::env::var("MEMRA_SERVE_SPEC").map(|v| v != "0").unwrap_or(true))
}
fn admit_reserve_override() -> Option<usize> {
static O: std::sync::OnceLock<Option<usize>> = std::sync::OnceLock::new();
*O.get_or_init(|| {
std::env::var("MEMRA_ADMIT_RESERVE_MB").ok()
.and_then(|v| v.parse::<usize>().ok())
.map(|mb| {
eprintln!("[admit-oom] WARN: MEMRA_ADMIT_RESERVE_MB={mb} overrides the \
{}MB transient reserve (teeth/diagnostics door — NOT a tuning knob)",
SPEC_SHRINK_RESERVE / (1 << 20));
mb * (1 << 20)
})
})
}
const STEP_OOM_MAX_RETRIES: u32 = 3;
fn step_oom_retries() -> u32 {
static R: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
*R.get_or_init(|| {
std::env::var("MEMRA_STEP_OOM_RETRIES").ok().and_then(|v| v.parse().ok())
.unwrap_or(STEP_OOM_MAX_RETRIES)
})
}
fn is_cuda_oom(err: &str) -> bool {
err.contains("CUDA_ERROR_OUT_OF_MEMORY") || err.contains("out of memory")
}
struct PrefixPlane {
k: CudaSlice<u8>,
v: CudaSlice<u8>,
len: usize,
}
struct PrefixEntry {
toks: Vec<u32>,
kv: Vec<Option<PrefixPlane>>,
conv: Vec<Option<CudaSlice<f32>>>,
ssm: Vec<Option<CudaSlice<f32>>>,
pos: usize,
last_logits: Vec<f32>,
bytes: usize,
last_use: Instant,
id: u64,
}
#[derive(Default)]
struct PrefixCache {
entries: HashMap<PoolKey, Vec<PrefixEntry>>,
lru: std::collections::BTreeMap<(Instant, u64), (PoolKey, usize)>,
next_id: u64,
total_bytes: usize,
hits: u64,
misses: u64,
inserts: u64,
evictions: u64,
hit_tokens: u64,
}
impl PrefixCache {
fn lcp(a: &[u32], b: &[u32]) -> usize {
a.iter().zip(b.iter()).take_while(|(x, y)| x == y).count()
}
fn n_entries(&self) -> usize {
self.entries.values().map(|p| p.len()).sum()
}
fn lookup(&self, key: &PoolKey, prompt: &[u32]) -> Option<usize> {
let pool = self.entries.get(key)?;
let mut best: Option<(usize, usize)> = None;
for (i, e) in pool.iter().enumerate() {
let n = e.toks.len();
if n >= PREFIX_CACHE_MIN_TOKENS && n <= prompt.len() && prompt[..n] == e.toks[..]
&& best.is_none_or(|(_, bn)| n > bn)
{
best = Some((i, n));
}
}
best.map(|(i, _)| i)
}
fn best_lcp(&self, key: &PoolKey, prompt: &[u32]) -> usize {
self.entries.get(key)
.map(|pool| pool.iter().map(|e| Self::lcp(&e.toks, prompt)).max().unwrap_or(0))
.unwrap_or(0)
}
fn has_covering(&self, key: &PoolKey, prompt: &[u32]) -> bool {
self.entries.get(key).is_some_and(|pool| pool.iter().any(|e| {
let n = e.toks.len();
n >= PREFIX_CACHE_MIN_TOKENS && n <= prompt.len() && prompt[..n] == e.toks[..]
}))
}
fn has_key(&self, key: &PoolKey, toks: &[u32]) -> bool {
self.entries.get(key).is_some_and(|pool| pool.iter().any(|e| e.toks[..] == *toks))
}
fn lru_key(e: &PrefixEntry) -> (Instant, u64) {
(e.last_use, e.id)
}
fn touch(&mut self, key: &PoolKey, i: usize) {
if let Some(e) = self.entries.get_mut(key).and_then(|p| p.get_mut(i)) {
self.lru.remove(&(e.last_use, e.id));
e.last_use = Instant::now();
self.lru.insert((e.last_use, e.id), (key.clone(), i));
}
}
fn remove_at(&mut self, key: &PoolKey, i: usize) -> Option<PrefixEntry> {
let pool = self.entries.get_mut(key)?;
if i >= pool.len() {
return None;
}
let dead = pool.swap_remove(i);
self.lru.remove(&Self::lru_key(&dead));
if let Some(moved) = pool.get(i) {
self.lru.insert(Self::lru_key(moved), (key.clone(), i));
}
if pool.is_empty() {
self.entries.remove(key);
}
Some(dead)
}
fn insert(&mut self, key: &PoolKey, e: PrefixEntry, why: &str) {
self.insert_with_budget(key, e, why, prefix_cache_budget_bytes());
}
fn insert_with_budget(&mut self, key: &PoolKey, mut e: PrefixEntry, why: &str, budget: usize) {
if e.bytes > budget {
eprintln!("[prefix-cache] skip {why} insert: entry {:.1}MB > budget {:.0}MB",
e.bytes as f64 / 1e6, budget as f64 / 1e6);
return;
}
if self.has_key(key, &e.toks) {
return; }
self.total_bytes += e.bytes;
self.inserts += 1;
eprintln!("[prefix-cache] insert ({why}): {} tokens, {:.1}MB (resident {:.1}MB / {:.0}MB, model {}{})",
e.toks.len(), e.bytes as f64 / 1e6,
self.total_bytes as f64 / 1e6, budget as f64 / 1e6,
key.0, ns_suffix(&key.1));
e.id = self.next_id;
self.next_id += 1;
let lk = Self::lru_key(&e);
let idx = {
let pool = self.entries.entry(key.clone()).or_default();
pool.push(e);
pool.len() - 1
};
self.lru.insert(lk, (key.clone(), idx));
while self.total_bytes > budget {
let Some((k, i)) = self.lru.values().next().cloned() else { break };
let Some(dead) = self.remove_at(&k, i) else { break };
self.total_bytes = self.total_bytes.saturating_sub(dead.bytes);
self.evictions += 1;
eprintln!("[prefix-cache] evict (LRU): {} tokens, {:.1}MB (model {}{})",
dead.toks.len(), dead.bytes as f64 / 1e6, k.0, ns_suffix(&k.1));
}
}
fn evict_all(&mut self) -> usize {
let n = self.n_entries();
self.entries.clear();
self.lru.clear();
self.total_bytes = 0;
self.evictions += n as u64;
n
}
}
fn prefix_snapshot(
engine: &Engine,
cache: &Cache,
toks: &[u32],
last_logits: &[f32],
) -> Result<PrefixEntry, Box<dyn std::error::Error>> {
let n = cache.kv.len();
let mut kv = Vec::with_capacity(n);
let mut conv = Vec::with_capacity(n);
let mut ssm = Vec::with_capacity(n);
let mut bytes = 0usize;
for il in 0..n {
match &cache.kv[il] {
Some(l) => {
let kb = l.len * l.k_tok_bytes;
let vb = l.len * l.v_tok_bytes;
let mut k = engine.alloc_u8(kb.max(1))?;
let mut v = engine.alloc_u8(vb.max(1))?;
if kb > 0 { engine.copy_u8_into(&mut k, 0, &l.k, kb)?; }
if vb > 0 { engine.copy_u8_into(&mut v, 0, &l.v, vb)?; }
bytes += kb + vb;
kv.push(Some(PrefixPlane { k, v, len: l.len }));
}
None => kv.push(None),
}
match &cache.recur[il] {
Some(r) => {
conv.push(Some(engine.clone_dtod(&r.conv_state)?));
ssm.push(Some(engine.clone_dtod(&r.ssm_state)?));
bytes += (r.conv_state.len() + r.ssm_state.len()) * 4;
}
None => {
conv.push(None);
ssm.push(None);
}
}
}
Ok(PrefixEntry {
toks: toks.to_vec(),
kv,
conv,
ssm,
pos: cache.pos,
last_logits: last_logits.to_vec(),
bytes,
last_use: Instant::now(),
id: 0, })
}
fn prefix_restore(
engine: &Engine,
cache: &mut Cache,
e: &PrefixEntry,
) -> Result<(), Box<dyn std::error::Error>> {
if cache.kv.len() != e.kv.len() {
return Err(format!("prefix entry layer count {} != cache {}", e.kv.len(), cache.kv.len()).into());
}
for il in 0..cache.kv.len() {
match (cache.kv[il].as_mut(), &e.kv[il]) {
(Some(dst), Some(src)) => {
let kb = src.len * dst.k_tok_bytes;
let vb = src.len * dst.v_tok_bytes;
if kb > 0 { engine.copy_u8_into(&mut dst.k, 0, &src.k, kb)?; }
if vb > 0 { engine.copy_u8_into(&mut dst.v, 0, &src.v, vb)?; }
dst.len = src.len;
engine.set_i32_one(&mut dst.len_d, src.len as i32)?;
}
(None, None) => {}
_ => return Err(format!("prefix entry layer {il} kind mismatch").into()),
}
match (cache.recur[il].as_mut(), &e.conv[il], &e.ssm[il]) {
(Some(dst), Some(c), Some(s)) => {
engine.copy_into(&mut dst.conv_state, 0, c, c.len())?;
engine.copy_into(&mut dst.ssm_state, 0, s, s.len())?;
}
(None, None, None) => {}
_ => return Err(format!("prefix entry recur {il} mismatch").into()),
}
}
cache.pos = e.pos;
Ok(())
}
fn prefix_insert_from_session(engine: &Engine, px: &mut PrefixCache, s: &Session, why: &str) {
let Some(cache) = s.cache.as_ref() else { return };
if s.last_logits.is_empty() {
return;
}
match prefix_snapshot(engine, cache, &s.fed, &s.last_logits) {
Ok(e) => px.insert(&s.pool_key(), e, why),
Err(err) => eprintln!("[prefix-cache] snapshot failed ({err}); prefix not cached"),
}
}
fn maybe_prefix_seed(engine: &Engine, px: &mut PrefixCache, s: &mut Session) {
if !s.seed_prefix {
return;
}
s.seed_prefix = false;
if s.n_cached > 0 || s.cache.is_none() || s.fed.len() < PREFIX_CACHE_MIN_TOKENS {
return;
}
if px.has_covering(&s.pool_key(), &s.fed) {
return; }
prefix_insert_from_session(engine, px, s, "seed");
}
struct ReplayPlan {
prompt_ids: Vec<u32>,
prompt_text: String,
chat: bool,
chat_turns: Vec<memra_tokenizer::chat::Turn>,
tools_json: Vec<String>,
think: memra_tokenizer::chat::ThinkMode,
params: GenParams,
sampler_cfg: SamplerConfig,
grammar: Option<crate::constrained::GrammarSpec>,
}
struct Session {
model: String,
cache_ns: String,
affinity: Option<String>,
lane: crate::lanes::Lane,
cache: Option<Cache>,
spec: Option<memra_engine::spec::SpecSession>,
graph: Option<memra_engine::decode::GraphSession>,
graph_pending: Option<u32>,
oom_retries: u32,
replay: Box<ReplayPlan>,
spec_drafted: usize,
spec_accepted: usize,
spec_rounds: u64,
sampler: Sampler,
last_logits: Vec<f32>,
device_next: Option<u32>,
constraint: Option<crate::constrained::SessionConstraint>,
mask_dev: Option<CudaSlice<u32>>,
mask_words: usize,
fed: Vec<u32>,
prefill_queue: std::collections::VecDeque<u32>,
prefill_done: bool,
generated: Vec<u32>,
params: GenParams,
stop_strings: Vec<String>,
trace_id: Option<String>,
emitted_bytes: usize,
budget: usize, n_prompt: usize,
n_cached: usize,
snapshot_at: Option<usize>,
seed_prefix: bool,
tx: tokio::sync::mpsc::UnboundedSender<Event>,
t0: Instant,
}
impl Session {
fn pool_key(&self) -> PoolKey {
(self.model.clone(), self.cache_ns.clone())
}
}
pub fn run(
models: Vec<(String, String, Option<String>)>,
rx: Receiver<Cmd>,
ready_tx: Sender<Result<(Vec<String>, HashMap<String, ModelCaps>), String>>,
metrics: SharedMetrics,
) {
let engine = match Engine::new(0) {
Ok(e) => e,
Err(err) => { let _ = ready_tx.send(Err(format!("Engine::new failed: {err}"))); return; }
};
let fast = std::env::var("MEMRA_FAST").as_deref() != Ok("0");
eprintln!("[worker] Engine ready (MEMRA_FAST={})", fast);
let mut loaded: HashMap<String, LoadedModel> = HashMap::new();
let mut order: Vec<String> = Vec::new();
for (name, path, draft) in &models {
eprintln!("[worker] loading model {name:?} <- {path}");
let from_dir = std::path::Path::new(path).is_dir();
let (model, tok) = if from_dir {
let dir = std::path::Path::new(path);
let (src, tok_dir): (Box<dyn memra_gguf::source::TensorSource>, std::path::PathBuf) =
if dir.join("manifest.json").exists() {
let repack = match memra_gguf::source::Hy3RepackSource::open(dir) {
Ok(source) => source,
Err(err) => { let _ = ready_tx.send(Err(format!("open {path}: {err}"))); return; }
};
let tok_dir = repack.source_dir()
.filter(|source| source.join("tokenizer.json").exists())
.unwrap_or(dir).to_path_buf();
(Box::new(repack), tok_dir)
} else {
let st = match memra_gguf::source::SafetensorsSource::open(dir) {
Ok(source) => source,
Err(err) => { let _ = ready_tx.send(Err(format!("open {path}: {err}"))); return; }
};
(Box::new(st), dir.to_path_buf())
};
let model = match HybridModel::load_from_source(&engine, src.as_ref()) {
Ok(m) => m,
Err(err) => { let _ = ready_tx.send(Err(format!("load {name}: {err}"))); return; }
};
let tok = match Tokenizer::from_hf_dir(&tok_dir) {
Ok(t) => t,
Err(err) => { let _ = ready_tx.send(Err(format!("tokenizer {name}: {err}"))); return; }
};
(model, tok)
} else {
let g = match GgufFile::open(path) {
Ok(g) => g,
Err(err) => { let _ = ready_tx.send(Err(format!("open {path}: {err}"))); return; }
};
let model = match HybridModel::load(&engine, &g) {
Ok(m) => m,
Err(err) => { let _ = ready_tx.send(Err(format!("load {name}: {err}"))); return; }
};
let tok = match Tokenizer::from_gguf(&g) {
Ok(t) => t,
Err(err) => { let _ = ready_tx.send(Err(format!("tokenizer {name}: {err}"))); return; }
};
(model, tok)
};
let model = {
let mut model = model;
if let Some(dpath) = draft {
let dg = match GgufFile::open(dpath) {
Ok(g) => g,
Err(err) => { let _ = ready_tx.send(Err(format!("draft {name}: {err}"))); return; }
};
match memra_engine::hybrid::MtpHead::load_draft(&engine, &dg, &model.cfg) {
Ok(head) => {
eprintln!("[worker] {name}: regime draft attached ({dpath})");
model.mtp = Some(head);
}
Err(err) => { let _ = ready_tx.send(Err(format!("draft {name}: {err}"))); return; }
}
}
model
};
let eos_id = tok.eos_id();
eprintln!("[worker] loaded {name:?}: {} layers, eos={eos_id}", model.cfg.n_layer);
loaded.insert(name.clone(), LoadedModel {
model, tok, eos_id, from_dir, constraints: std::cell::OnceCell::new(),
});
order.push(name.clone());
}
let caps: HashMap<String, ModelCaps> = loaded.iter().map(|(n, lm)| {
let t = lm.tok.chat_template();
let caps = ModelCaps {
tools_branch: t.is_some_and(|t| t.contains("<tools>")
&& !t.contains("hy_User") && !t.contains("<|turn>")),
qwen_think: t.is_some_and(|t| t.contains("<think>") && t.contains("add_generation_prompt")),
think_switch: t.is_some_and(|t| t.contains("enable_thinking")),
chat_ok: t.is_some() || !lm.from_dir,
context_length: lm.model.cfg.context_length as usize,
tokenizer: lm.tok.pre().to_string(),
instruct_type: t.and_then(|t| {
if t.contains("<|im_start|>") { Some("chatml".to_string()) }
else if t.contains("<start_of_turn>") { Some("gemma".to_string()) }
else { None }
}),
};
eprintln!("[worker] {n}: template caps tools={} think={} think_switch={} chat_ok={} \
ctx={} tok={:?} instruct={:?}",
caps.tools_branch, caps.qwen_think, caps.think_switch, caps.chat_ok,
caps.context_length, caps.tokenizer, caps.instruct_type);
(n.clone(), caps)
}).collect();
let _ = ready_tx.send(Ok((order.clone(), caps)));
let chunk_caps: HashMap<String, usize> =
loaded.iter().map(|(n, lm)| (n.clone(), chunk_cap_for(lm))).collect();
for (n, c) in &chunk_caps {
eprintln!("[worker] {n}: decode chunk cap {c}{}",
if *c > 8 { " (exact-16 tier)" } else { "" });
}
let mut active: Vec<Session> = Vec::new();
let mut queue: std::collections::VecDeque<Box<Request>> = std::collections::VecDeque::new();
let mut reuse: HashMap<PoolKey, Vec<ReuseEntry>> = HashMap::new();
let mut spec_reuse: HashMap<PoolKey, Vec<SpecReuseEntry>> = HashMap::new();
let mut spec_sizing = SpecSizing::default();
let mut px = PrefixCache::default();
if prefix_cache_budget_bytes() > 0 && serve_batching() {
eprintln!("[prefix-cache] on: budget {:.0}MB (MEMRA_PREFIX_CACHE_MB), min prefix {} tokens",
prefix_cache_budget_bytes() as f64 / 1e6, PREFIX_CACHE_MIN_TOKENS);
}
let mut session_vram_cost: HashMap<String, usize> = HashMap::new();
let policy = crate::lanes::LanePolicy::from_env();
let mut step_stats = StepStats::new(
std::env::var("MEMRA_LANE_WINDOW_S").ok().and_then(|v| v.parse().ok()).unwrap_or(30.0));
let mut n_admitted = 0u64;
let mut n_completed = 0u64;
let mut n_tokens_out = 0u64;
let mut n_prompt_in = 0u64;
let mut n_cached_in = 0u64;
let mut lane_admitted = [0u64; 3];
let mut lane_shed = [0u64; 3];
let mut lane_completed = [0u64; 3];
let mut lane_tokens = [0u64; 3];
let mut last_batch = 0usize;
let mut spec_telem: HashMap<String, memra_engine::spec::SpecTelemetry> = HashMap::new();
let mut spec_telem_dirty = false;
let mut last_interactive_decode = Instant::now();
let mut tick_n: u64 = 0;
loop {
if active.is_empty() && queue.is_empty() {
match rx.recv() {
Ok(cmd) => handle_cmd(cmd, &loaded, &order, &mut queue),
Err(_) => break, }
}
loop {
match rx.try_recv() {
Ok(cmd) => handle_cmd(cmd, &loaded, &order, &mut queue),
Err(std::sync::mpsc::TryRecvError::Empty) => break,
Err(std::sync::mpsc::TryRecvError::Disconnected) => {
if active.is_empty() { return; } else { break; }
}
}
}
let max_active = if confidence_trace_enabled() { 1 } else { MAX_ACTIVE };
let mut requeue: std::collections::VecDeque<Box<Request>> = Default::default();
let mut vram_defers = 0usize;
while let Some(req) = queue.pop_front() {
if req.tx.is_closed() {
eprintln!("[abort] client disconnected while queued (model {:?}); dropped",
req.model);
continue;
}
let lane = req.lane;
let batching_on = std::env::var("MEMRA_SERVE_BATCH").map(|v| v != "0").unwrap_or(true);
let cap = if lane == crate::lanes::Lane::Interactive {
if batching_on {
std::env::var("MEMRA_MAX_SESSIONS").ok()
.and_then(|v| v.parse().ok()).unwrap_or(64)
} else {
max_active
}
} else {
policy.max_sessions[lane.idx()]
};
let lane_count = active.iter().filter(|s| s.lane == lane).count();
if lane_count >= cap {
if lane == crate::lanes::Lane::Interactive {
requeue.push_back(req); } else {
lane_shed[lane.idx()] += 1;
let _ = req.tx.send(Event::Error(format!(
"shed:{}:lane at capacity, retry", lane.as_str())));
}
continue;
}
let interactive_active_or_waiting = active.iter()
.any(|s| s.lane == crate::lanes::Lane::Interactive);
let starved = interactive_active_or_waiting
&& last_interactive_decode.elapsed().as_secs_f32() * 1000.0 > policy.slo_p99_ms;
if !policy.admit(lane, &mut step_stats, starved) {
lane_shed[lane.idx()] += 1;
let _ = req.tx.send(Event::Error(format!(
"shed:{}:interactive p99 over budget, retry", lane.as_str())));
continue;
}
if !active.is_empty() {
if let (Some(&cost), Ok((free, _))) =
(session_vram_cost.get(&req.model), engine.ctx().mem_get_info()) {
let reserve = if serve_spec_enabled()
&& loaded.get(&req.model).is_some_and(|lm| lm.model.mtp.is_some())
{
admit_reserve_override().unwrap_or(SPEC_SHRINK_RESERVE)
} else {
cost
};
let free = free.saturating_add(engine.pool_cached_bytes());
if free < cost.saturating_add(reserve) {
if vram_defers == 0 {
let (res, used) = engine.pool_reserved_used();
let parked: usize = spec_reuse.values().map(|v| v.len()).sum();
eprintln!("[admit-oom] VRAM defer: {} active, effective free \
{:.0}MB (driver + {:.0}MB pool-cached) < cost {:.0}MB \
+ reserve {:.0}MB — queueing (FIFO) \
[pool res {:.0}MB used {:.0}MB; parked spec sessions {}; \
plain reuse {}; queue {}]",
active.len(), free as f64 / 1e6,
engine.pool_cached_bytes() as f64 / 1e6,
cost as f64 / 1e6, reserve as f64 / 1e6,
res as f64 / 1e6, used as f64 / 1e6, parked,
reuse.values().map(|v| v.len()).sum::<usize>(),
queue.len() + requeue.len());
}
vram_defers += 1;
requeue.push_back(req); continue;
}
}
}
let model_key = req.model.clone();
let free_before = engine.ctx().mem_get_info().map(|(f, _)| f).ok();
match admit(&engine, &loaded, &mut reuse, &mut spec_reuse, &mut spec_sizing,
&mut px, *req) {
Ok(s) => {
n_admitted += 1;
lane_admitted[lane.idx()] += 1;
n_prompt_in += s.n_prompt as u64;
n_cached_in += s.n_cached as u64;
active.push(s);
if !session_vram_cost.contains_key(&model_key) {
if let (Some(fb), Ok((fa, _))) = (free_before, engine.ctx().mem_get_info()) {
let cost = fb.saturating_sub(fa);
if cost > 0 {
eprintln!("[worker] observed session VRAM cost for {model_key:?}: \
{:.0}MB (admission gate = 2x)", cost as f64 / 1e6);
session_vram_cost.insert(model_key, cost);
}
}
}
}
Err((tx, msg)) => { let _ = tx.send(Event::Error(msg)); }
}
}
queue = requeue;
let batching = serve_batching();
let mut finished: Vec<usize> = Vec::new();
let mut requeue_oom: std::collections::VecDeque<Box<Request>> = Default::default();
for (i, s) in active.iter().enumerate() {
if s.tx.is_closed() {
abort_log(s);
finished.push(i);
}
}
if !batching {
for i in 0..active.len() {
if finished.contains(&i) { continue; }
match step_session(&engine, &loaded, &mut active[i], &mut spec_telem) {
Ok(true) => {}
Ok(false) => finished.push(i),
Err(err) => {
let _ = active[i].tx.send(Event::Error(format!("step error: {err}")));
finished.push(i);
}
}
}
} else {
let gs_on = {
static G: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*G.get_or_init(|| std::env::var("MEMRA_SERVE_GS").map(|v| v != "0").unwrap_or(true))
};
if gs_on && active.len() > 1 {
for i in 0..active.len() {
if finished.contains(&i) || active[i].graph.is_none() { continue; }
let s = &mut active[i];
let g = s.graph.take().unwrap();
s.cache = Some(g.cache);
if let Some(pend) = s.graph_pending.take() {
let (cont, _) = advance_token_emit(&loaded, s, pend);
if !cont {
finished.push(i);
} else {
let lm = &loaded[&s.model];
match lm.model.decode_step(&engine, pend, s.cache.as_mut().unwrap()) {
Ok(l) => { s.last_logits = l; s.fed.push(pend); }
Err(err) => {
let _ = s.tx.send(Event::Error(format!("degrade: {err}")));
finished.push(i);
}
}
}
}
}
}
if gs_on && active.len() == 1 && !finished.contains(&0) {
let s = &mut active[0];
let gs_min: usize = std::env::var("MEMRA_GS_MIN").ok()
.and_then(|v| v.parse().ok()).unwrap_or(384);
let constr_graph_ok = s.constraint.is_none()
|| (!constrain_host() && devsample_meta(s).is_some());
if s.graph.is_none() && s.spec.is_none() && s.sampler.is_greedy()
&& constr_graph_ok
&& s.lane == crate::lanes::Lane::Interactive
&& s.budget >= gs_min
&& s.prefill_done && s.generated.is_empty() && s.cache.is_some()
&& !s.last_logits.is_empty()
{
let lm = &loaded[&s.model];
let (first, mask0) = match s.constraint.as_mut() {
Some(c) => match c.compute_mask() {
Ok(m) => {
let mut row = s.last_logits.clone();
crate::constrained::apply_mask(&m, &mut row);
(memra_engine::forward::argmax(&row) as u32, Some(m))
}
Err(err) => {
let _ = s.tx.send(Event::Error(format!("constraint mask: {err}")));
finished.push(0);
(0, None)
}
},
None => (memra_engine::forward::argmax(&s.last_logits) as u32, None),
};
if !finished.contains(&0) {
let cache = s.cache.take().unwrap();
match lm.model.graph_session_from_cache_masked(
&engine, cache, first, s.budget + 2,
mask0.as_ref().map(|m| m.as_slice())) {
Ok((g, first)) => {
s.graph = Some(g);
s.graph_pending = Some(first);
}
Err(err) => {
let _ = s.tx.send(Event::Error(format!("graph promote failed: {err}")));
finished.push(0);
}
}
}
}
let s = &mut active[0];
if let Some(pend) = s.graph_pending.take() {
let t_g = Instant::now();
let (cont, _) = advance_token_emit(&loaded, s, pend);
if !cont {
finished.push(0);
} else {
s.fed.push(pend);
let mut mask_err = None;
if let Some(c) = s.constraint.as_mut() {
match c.compute_mask() {
Ok(m) => {
if let Err(err) = s.graph.as_mut().unwrap()
.upload_mask(&engine, m.as_slice()) {
mask_err = Some(err.to_string());
}
}
Err(err) => mask_err = Some(err),
}
}
if let Some(err) = mask_err {
let _ = s.tx.send(Event::Error(format!("constraint mask: {err}")));
finished.push(0);
} else {
let lm = &loaded[&s.model];
let at_budget = s.graph.as_ref()
.is_some_and(|g| g.cache.pos + 1 >= g.bucket_max);
let g = s.graph.as_mut().unwrap();
match g.step(&engine, &lm.model) {
Ok(next) => { s.graph_pending = Some(next); }
Err(err) if at_budget => {
eprintln!("[worker] graph session capture budget reached \
(model {}): {err}", s.model);
finish(s, StopReason::MaxNew);
finished.push(0);
}
Err(err) => {
eprintln!("[worker] graph session step FAILED \
(model {}): {err}", s.model);
let _ = s.tx.send(Event::Error(
format!("graph step failed: {err}")));
finished.push(0);
}
}
n_tokens_out += 1;
lane_tokens[0] += 1;
step_stats.record(t_g.elapsed().as_secs_f32() * 1000.0);
last_interactive_decode = Instant::now();
}
}
}
}
let mut spec_order: Vec<usize> = (0..active.len())
.filter(|&i| active[i].spec.is_some())
.collect();
let admit_yield_on = {
static Y: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*Y.get_or_init(|| std::env::var("MEMRA_ADMIT_YIELD").as_deref() != Ok("0"))
};
if admit_yield_on {
spec_order.sort_by_key(|&i| !active[i].generated.is_empty());
}
for i in spec_order {
if finished.contains(&i) { continue; }
match step_session(&engine, &loaded, &mut active[i], &mut spec_telem) {
Ok(true) => {}
Ok(false) => finished.push(i),
Err(err) if is_cuda_oom(&err.to_string())
&& step_oom_retries() > 0
&& active[i].generated.is_empty()
&& active[i].oom_retries < step_oom_retries() =>
{
let n_active = active.len();
let s = &mut active[i];
s.oom_retries += 1;
eprintln!("[admit-oom] step OOM parked session back to queue \
(model {}, retry {}/{}, {n_active} active): {err}",
s.model, s.oom_retries, step_oom_retries());
match park_requeue(&loaded, s) {
Some(req) => { requeue_oom.push_back(req); finished.push(i); }
None => {
let _ = s.tx.send(Event::Error(format!("step error: {err}")));
finished.push(i);
}
}
}
Err(err) => {
if is_cuda_oom(&err.to_string()) {
eprintln!("[admit-oom] step OOM NOT parked (model {}, retries \
{}/{}, generated {}): reporting honestly",
active[i].model, active[i].oom_retries,
step_oom_retries(), active[i].generated.len());
}
let _ = active[i].tx.send(Event::Error(format!("step error: {err}")));
finished.push(i);
}
}
}
let budgets = policy.prefill_budget;
let (cand, held) = 'pb: loop {
let pb_max: usize = std::env::var("MEMRA_PRIME_BATCH").ok()
.and_then(|v| v.parse().ok()).unwrap_or(6);
let pb_maxt: usize = std::env::var("MEMRA_PRIME_BATCH_MAX_T").ok()
.and_then(|v| v.parse().ok()).unwrap_or(2048);
let min_t = memra_engine::hybrid_forward::PRIME_MIN_T.max(2);
let mut cand: Vec<usize> = Vec::new();
let mut cand_model: Option<String> = None;
if pb_max >= 2 && !confidence_trace_enabled() {
for i in 0..active.len() {
if finished.contains(&i) { continue; }
let s = &active[i];
let ql = s.prefill_queue.len();
if s.spec.is_none() && !s.prefill_done && s.graph.is_none()
&& s.lane == crate::lanes::Lane::Interactive
&& s.fed.is_empty()
&& s.cache.as_ref().is_some_and(|c| c.pos == 0)
&& s.snapshot_at.is_none()
&& ql >= min_t && ql <= pb_maxt && ql <= budgets[0]
&& cand_model.as_ref().is_none_or(|m| *m == s.model)
{
cand_model.get_or_insert_with(|| s.model.clone());
cand.push(i);
if cand.len() == pb_max { break; }
}
}
}
let hold_ms: u64 = std::env::var("MEMRA_PRIME_BATCH_HOLD_MS").ok()
.and_then(|v| v.parse().ok()).unwrap_or(4);
let mut held = false;
if cand.len() == 1 && hold_ms > 0 {
let s = &active[cand[0]];
if s.t0.elapsed().as_millis() < hold_ms as u128 {
held = true;
}
}
let mut fired = false;
if cand.len() >= 2 {
let prompts: Vec<Vec<u32>> = cand.iter()
.map(|&i| active[i].prefill_queue.drain(..).collect())
.collect();
let prompt_refs: Vec<&[u32]> = prompts.iter().map(|p| p.as_slice()).collect();
let mut cache_refs: Vec<&mut memra_engine::cache::Cache> = active.iter_mut()
.enumerate()
.filter(|(i, _)| cand.contains(i))
.map(|(_, s)| s.cache.as_mut().unwrap())
.collect();
let lm = &loaded[cand_model.as_ref().unwrap()];
let t_pb = Instant::now();
match lm.model.prime_cache_batch(&engine, &prompt_refs, &mut cache_refs) {
Ok(outs) => {
let toks: usize = prompts.iter().map(|p| p.len()).sum();
eprintln!("[prime-batch] B={} tokens={} in {:.1}ms",
cand.len(), toks, t_pb.elapsed().as_secs_f64() * 1e3);
for ((&i, prompt), (l, _h, _x)) in
cand.iter().zip(&prompts).zip(outs)
{
let s = &mut active[i];
s.last_logits = l;
for &tok in prompt { s.fed.push(tok); s.sampler.accept(tok); }
s.prefill_done = true;
maybe_prefix_seed(&engine, &mut px, s);
}
fired = true;
}
Err(err) => {
eprintln!("[prime-batch] failed ({err}); single primes serve");
for (&i, prompt) in cand.iter().zip(&prompts) {
active[i].prefill_queue = prompt.iter().copied().collect();
}
}
}
}
if fired { continue 'pb; }
break 'pb (cand, held);
};
for i in 0..active.len() {
if finished.contains(&i) { continue; }
if held && cand.first() == Some(&i) { continue; } let s = &mut active[i];
if s.spec.is_some() || s.prefill_done { continue; }
if s.lane != crate::lanes::Lane::Interactive { continue; }
match prefill_tick(&engine, &loaded, &mut px, s, budgets[0]) {
Ok(_) => {}
Err(err) => {
let _ = s.tx.send(Event::Error(format!("prefill error: {err}")));
finished.push(i);
}
}
}
let t_decode = Instant::now();
let mut decoding: Vec<usize> = (0..active.len())
.filter(|&i| !finished.contains(&i)
&& active[i].spec.is_none() && active[i].prefill_done
&& active[i].cache.is_some())
.collect();
decoding.sort_by_key(|&i| active[i].lane.idx());
let mut had_interactive = false;
let mut ready: Vec<(usize, u32)> = Vec::new();
for &i in &decoding {
let (cont, next) = advance_sample_emit(&loaded, &mut active[i]);
match (cont, next) {
(false, _) => finished.push(i),
(true, Some(t)) => {
if let Err(err) = stage_grammar_mask(&engine, &mut active[i]) {
let _ = active[i].tx.send(Event::Error(
format!("constraint mask: {err}")));
finished.push(i);
continue;
}
had_interactive |= active[i].lane == crate::lanes::Lane::Interactive;
ready.push((i, t));
}
(true, None) => {} }
}
for chunk in group_chunks(&active, &ready, &chunk_caps) {
let toks: Vec<u32> = chunk.iter().map(|&(_, t)| t).collect();
let idxs: Vec<usize> = chunk.iter().map(|&(i, _)| i).collect();
let model_name = active[idxs[0]].model.clone();
let lm = &loaded[&model_name];
let samp: Vec<Option<(f32, u64, u32)>> = idxs
.iter()
.map(|&i| {
let s = &active[i];
if s.constraint.is_some() && s.mask_words == 0 {
return None;
}
devsample_meta(s)
})
.collect();
let mask_ptrs: Vec<Option<(*const CudaSlice<u32>, usize)>> = idxs
.iter()
.map(|&i| {
let s = &active[i];
if s.mask_words > 0 {
s.mask_dev.as_ref().map(|d| (d as *const _, s.mask_words))
} else {
None
}
})
.collect();
let logits = {
let mut caches: Vec<&mut Cache> = Vec::with_capacity(idxs.len());
let base = active.as_mut_ptr();
for &i in &idxs {
let s = unsafe { &mut *base.add(i) };
caches.push(s.cache.as_mut().unwrap());
}
let masks: Vec<Option<(&CudaSlice<u32>, usize)>> = mask_ptrs
.iter()
.map(|m| m.map(|(p, w)| (unsafe { &*p }, w)))
.collect();
lm.model.decode_step_batch_sampled_lean_masked(
&engine, &toks, &mut caches, &samp, &masks, serve_leanlogits())
};
match logits {
Ok((rows, next_toks)) => {
for (k, &i) in idxs.iter().enumerate() {
active[i].last_logits = rows[k].clone();
active[i].device_next = next_toks[k];
active[i].fed.push(toks[k]);
n_tokens_out += 1;
lane_tokens[active[i].lane.idx()] += 1;
}
}
Err(err) => {
for &i in &idxs {
let _ = active[i].tx.send(Event::Error(format!("batch step: {err}")));
finished.push(i);
}
}
}
}
if had_interactive {
last_interactive_decode = Instant::now();
}
last_batch = ready.len();
if std::env::var("MEMRA_TICK_TRACE").as_deref() == Ok("1") {
let n_int = active.iter()
.filter(|s| s.lane == crate::lanes::Lane::Interactive).count();
let n_pref = active.iter().filter(|s| !s.prefill_done).count();
eprintln!("[tick] act={} int={} priming={} ready={} decode_ms={:.1}",
active.len(), n_int, n_pref, ready.len(),
t_decode.elapsed().as_secs_f32() * 1000.0);
}
let decode_ms = t_decode.elapsed().as_secs_f32() * 1000.0;
let headroom_ms = (policy.slo_p99_ms - decode_ms).max(0.0);
let prime_tok_per_ms: f32 = std::env::var("MEMRA_PRIME_TOK_PER_MS").ok()
.and_then(|v| v.parse().ok()).unwrap_or(8.0);
let adaptive_cap = (headroom_ms * prime_tok_per_ms) as usize;
let mut dark_batched = false;
{
let min_t = memra_engine::hybrid_forward::PRIME_MIN_T.max(2);
let mut dcand: Vec<usize> = Vec::new();
let mut dmodel: Option<String> = None;
let mut dlane: Option<usize> = None;
let mut dsum = 0usize;
for i in 0..active.len() {
if finished.contains(&i) { continue; }
let s = &active[i];
let li = s.lane.idx();
let ql = s.prefill_queue.len();
if li == 0 || budgets[li] == 0 { continue; }
if s.spec.is_some() || s.prefill_done || s.graph.is_some()
|| s.snapshot_at.is_some()
|| !s.cache.as_ref().is_some_and(|c| c.pos == s.fed.len()) { continue; }
if !s.fed.is_empty() && loaded[&s.model].model.cfg.gemma4.is_some() { continue; }
let cap = budgets[li].min(adaptive_cap);
if ql < min_t || dsum + ql > cap { continue; }
if dlane.is_some_and(|l| l != li) { continue; }
if dmodel.as_ref().is_some_and(|m| *m != s.model) { continue; }
dlane.get_or_insert(li);
dmodel.get_or_insert_with(|| s.model.clone());
dsum += ql;
dcand.push(i);
}
if dcand.len() >= 2 {
let prompts: Vec<Vec<u32>> = dcand.iter()
.map(|&i| active[i].prefill_queue.drain(..).collect())
.collect();
let prompt_refs: Vec<&[u32]> = prompts.iter().map(|p| p.as_slice()).collect();
let mut cache_refs: Vec<&mut memra_engine::cache::Cache> = active.iter_mut()
.enumerate()
.filter(|(i, _)| dcand.contains(i))
.map(|(_, s)| s.cache.as_mut().unwrap())
.collect();
let lm = &loaded[dmodel.as_ref().unwrap()];
match lm.model.prime_cache_batch(&engine, &prompt_refs, &mut cache_refs) {
Ok(outs) => {
let ncar = dcand.iter()
.filter(|&&i| !active[i].fed.is_empty()).count();
eprintln!("[prime-batch dark] lane={} B={} tokens={dsum} carried={ncar}",
dlane.unwrap(), dcand.len());
for ((&i, prompt), (l, _h, _x)) in
dcand.iter().zip(&prompts).zip(outs)
{
let s = &mut active[i];
s.last_logits = l;
for &tok in prompt { s.fed.push(tok); s.sampler.accept(tok); }
s.prefill_done = true;
}
}
Err(err) => {
eprintln!("[prime-batch dark] failed ({err}); chunks serve");
for (&i, prompt) in dcand.iter().zip(&prompts) {
active[i].prefill_queue = prompt.iter().copied().collect();
}
dcand.clear();
}
}
dark_batched = !dcand.is_empty(); }
}
for i in 0..active.len() {
if dark_batched { break; }
if finished.contains(&i) { continue; }
let s = &mut active[i];
if s.spec.is_some() || s.prefill_done { continue; }
let li = s.lane.idx();
if li == 0 || budgets[li] == 0 { continue; }
let chunk = budgets[li].min(adaptive_cap);
if chunk < memra_engine::hybrid_forward::PRIME_MIN_T { break; }
if let Err(err) = prefill_tick(&engine, &loaded, &mut px, s, chunk) {
let _ = s.tx.send(Event::Error(format!("prefill error: {err}")));
finished.push(i);
}
break; }
if had_interactive {
step_stats.record(t_decode.elapsed().as_secs_f32() * 1000.0);
}
}
finished.sort_unstable();
finished.dedup();
for &i in finished.iter().rev() {
let s = active.remove(i);
let pool_key = s.pool_key(); n_completed += 1;
lane_completed[s.lane.idx()] += 1;
if s.spec_rounds > 0 { spec_telem_dirty = true; } if let Some(mut sess) = s.spec {
if sess.pending_tok.is_some() {
if let Err(err) = loaded[&s.model].model.spec_flush_pending(&engine, &mut sess) {
eprintln!("[worker] spec pending flush failed ({err}); dropping session");
continue;
}
}
if sess.committed.len() >= REUSE_MIN_PREFIX && sess.next_pred.is_some() {
let toks = &sess.committed;
let skip = loaded[&s.model].tok.bos_id()
.map(|b| toks.first() == Some(&b)).unwrap_or(false) as usize;
let committed_text = loaded[&s.model].tok.decode_special(&toks[skip..], true);
let tok = &loaded[&s.model].tok;
let fingerprint = conversation_fingerprint(
toks, &|t| tok.token_is_control(t), false);
let pool = spec_reuse.entry(pool_key).or_default();
while pool.len() >= reuse_pool_per_model().max(1) { pool.remove(0); }
if reuse_pool_per_model() > 0 {
pool.push(SpecReuseEntry {
sess, committed_text, affinity: s.affinity, fingerprint,
});
}
}
} else if s.fed.len() >= REUSE_MIN_PREFIX && s.prefill_done {
if let Some(cache) = s.cache {
let last_logits = if s.last_logits.is_empty() {
cache.last_logits_dev.as_ref()
.and_then(|d| engine.dtoh(d).ok())
.unwrap_or_default()
} else {
s.last_logits
};
if !last_logits.is_empty() {
let pool = reuse.entry(pool_key).or_default();
while pool.len() >= reuse_pool_per_model().max(1) { pool.remove(0); }
let cap = cache.max_ctx;
if reuse_pool_per_model() > 0 {
pool.push(ReuseEntry {
fed: s.fed, cache, last_logits, cap,
});
}
}
}
}
}
while let Some(req) = requeue_oom.pop_back() {
queue.push_front(req);
}
tick_n = tick_n.wrapping_add(1);
if tick_n % 32 == 0 || spec_telem_dirty { if let Ok(mut m) = metrics.lock() {
spec_telem_dirty = false;
m.admitted = n_admitted;
m.completed = n_completed;
m.tokens_out = n_tokens_out;
m.step_p50_ms = step_stats.p(50.0).unwrap_or(0.0);
m.step_p99_ms = step_stats.p(99.0).unwrap_or(0.0);
m.prompt_tokens_in = n_prompt_in;
m.cached_tokens_in = n_cached_in;
m.prefix_hits = px.hits;
m.prefix_entries = px.n_entries() as u64;
m.prefix_bytes = px.total_bytes as u64;
m.lane_admitted = lane_admitted;
m.lane_shed = lane_shed;
m.lane_completed = lane_completed;
m.lane_tokens = lane_tokens;
m.batch_size_last = last_batch;
m.spec = spec_telem.clone();
} }
if !finished.is_empty() && std::env::var("MEMRA_SPILL_STATS").as_deref() == Ok("1") {
if let Some((reads, bytes, errors, short, fallbacks, waits, ring_full)) =
engine.moe_pread_stats() {
eprintln!("[spill-pread] snapshot reads={reads} bytes={bytes} errors={errors} \
short_reads={short} fallbacks={fallbacks} buffer_waits={waits} \
ring_full={ring_full}");
}
if let Some((hits, misses, staged_bytes, slots)) = engine.moe_cache_stats() {
let accesses = hits.saturating_add(misses);
let hit_rate = if accesses == 0 {
0.0
} else {
100.0 * hits as f64 / accesses as f64
};
eprintln!("[moe-cache] snapshot hits={hits} misses={misses} \
hit_rate={hit_rate:.3} staged_bytes={staged_bytes} slots={slots}");
}
}
}
}
fn handle_cmd(
cmd: Cmd,
loaded: &HashMap<String, LoadedModel>,
order: &[String],
queue: &mut std::collections::VecDeque<Box<Request>>,
) {
let _ = PENDING_ADMITS.fetch_update(
std::sync::atomic::Ordering::AcqRel,
std::sync::atomic::Ordering::Acquire,
|v| v.checked_sub(1),
);
match cmd {
Cmd::Generate(req) => {
if !loaded.contains_key(&req.model) {
let _ = req.tx.send(Event::Error(format!(
"unknown model {:?}; loaded: {:?}", req.model, order)));
return;
}
queue.push_back(req);
}
}
}
fn park_requeue(loaded: &HashMap<String, LoadedModel>, s: &Session) -> Option<Box<Request>> {
let p = &s.replay;
if p.prompt_ids.is_empty() && p.prompt_text.is_empty() && p.chat_turns.is_empty() {
return None;
}
debug_assert!(loaded.contains_key(&s.model), "parked session's model must still be loaded");
Some(Box::new(Request {
model: s.model.clone(),
prompt_ids: p.prompt_ids.clone(),
prompt_text: p.prompt_text.clone(),
chat: p.chat,
chat_turns: p.chat_turns.clone(),
tools_json: p.tools_json.clone(),
think: p.think,
params: p.params.clone(),
sampler_cfg: p.sampler_cfg.clone(),
stop_strings: s.stop_strings.clone(),
trace_id: s.trace_id.clone(),
cache_ns: s.cache_ns.clone(),
affinity: s.affinity.clone(),
lane: s.lane,
grammar: p.grammar.clone(),
oom_retries: s.oom_retries,
tx: s.tx.clone(),
}))
}
fn admit(
engine: &Engine,
loaded: &HashMap<String, LoadedModel>,
reuse: &mut HashMap<PoolKey, Vec<ReuseEntry>>,
spec_reuse: &mut HashMap<PoolKey, Vec<SpecReuseEntry>>,
spec_sizing: &mut SpecSizing,
px: &mut PrefixCache,
req: Request,
) -> Result<Session, (tokio::sync::mpsc::UnboundedSender<Event>, String)> {
let lm = &loaded[&req.model];
let pool_key: PoolKey = (req.model.clone(), req.cache_ns.clone());
let replay = Box::new(ReplayPlan {
prompt_ids: req.prompt_ids.clone(),
prompt_text: req.prompt_text.clone(),
chat: req.chat,
chat_turns: req.chat_turns.clone(),
tools_json: req.tools_json.clone(),
think: req.think,
params: req.params.clone(),
sampler_cfg: req.sampler_cfg.clone(),
grammar: req.grammar.clone(),
});
let req_oom_retries = req.oom_retries;
let prompt: Vec<u32> = if !req.prompt_ids.is_empty() {
req.prompt_ids.clone()
} else if !req.chat_turns.is_empty() {
let plain = req.tools_json.is_empty()
&& req.think == memra_tokenizer::chat::ThinkMode::Default
&& req.chat_turns.iter().all(|t| t.role != "tool" && t.tool_calls.is_empty());
let rendered = if plain {
let messages: Vec<_> = req.chat_turns.iter()
.map(|t| (t.role.as_str(), t.content.as_str()))
.collect();
lm.tok.apply_chat_template(&messages, true)
} else {
match lm.tok.apply_chat_template_tools(&req.chat_turns, true,
&req.tools_json, req.think) {
Ok(rendered) => rendered,
Err(err) => return Err((req.tx, format!("chat template: {err}"))),
}
};
lm.tok.encode(&rendered, true)
} else if req.chat {
let rendered = lm.tok.apply_chat_template(&[("user", req.prompt_text.as_str())], true);
lm.tok.encode(&rendered, true)
} else {
lm.tok.encode(&req.prompt_text, true)
};
if prompt.is_empty() {
return Err((req.tx, "empty prompt after tokenization".into()));
}
let ctx_floor: usize = std::env::var("MEMRA_CTX").ok().and_then(|v| v.parse().ok()).unwrap_or(8192);
let ctx_cap = match (req.params.max_ctx, req.params.max_new) {
(Some(c), _) => c.max(ctx_floor),
(None, MAX_NEW_CTX_BOUNDED) => {
let model_ctx = lm.model.cfg.context_length as usize;
let mut c = ctx_floor;
if prompt.len() + 16 > c { c = prompt.len().saturating_add(ctx_floor); }
if model_ctx > 0 { c = c.min(model_ctx); }
c
}
(None, max_new) => (prompt.len() + max_new + 8).max(ctx_floor),
};
if prompt.len() >= ctx_cap {
return Err((req.tx, format!(
"prompt ({} tok) >= context cap ({})", prompt.len(), ctx_cap)));
}
let room = ctx_cap - prompt.len();
let budget = req.params.max_new.min(room);
let need = prompt.len().saturating_add(budget).saturating_add(SPEC_SHRINK_SLACK);
let mut reused: Option<ReuseEntry> = None;
let reuse_on = !confidence_trace_enabled()
&& std::env::var("MEMRA_KV_REUSE").map(|v| v != "0").unwrap_or(true);
if let (true, Some(pool)) = (reuse_on, reuse.get_mut(&pool_key)) {
if let Some(idx) = pool.iter().rposition(|e|
e.fed.len() >= REUSE_MIN_PREFIX && e.cap >= ctx_cap
&& prompt.len() >= e.fed.len() && prompt.starts_with(&e.fed)) {
reused = Some(pool.remove(idx));
}
}
let serve_spec = !confidence_trace_enabled()
&& std::env::var("MEMRA_SERVE_SPEC").map(|v| v != "0").unwrap_or(true);
let mut sampler = Sampler::new(req.sampler_cfg);
let greedy_penalized = sampler.is_greedy()
&& (sampler.penalty_repeat() != 1.0 || sampler.penalty_freq() != 0.0
|| sampler.penalty_present() != 0.0);
let constraint = match &req.grammar {
None => None,
Some(spec) => {
let factory = lm.constraints.get_or_init(||
crate::constrained::ConstraintFactory::new(&lm.tok));
match factory {
Err(err) => return Err((req.tx, format!("constrained decoding: {err}"))),
Ok(f) => {
let sc = f.matcher(spec);
if let Some(err) = sc.error() {
return Err((req.tx, format!("response_format: {err}")));
}
Some(sc)
}
}
}
};
let spec_eligible = serve_spec
&& (constraint.is_none() || (sampler.is_greedy() && !constrain_host()))
&& (sampler.is_greedy() || sampler.temperature() > 0.0)
&& !greedy_penalized
&& lm.model.mtp.is_some();
let prefix_on = reuse_on && serve_batching() && prefix_cache_budget_bytes() > 0;
let mut prefix_hit = false;
let mut snapshot_at: Option<usize> = None;
let mut seed_prefix = false;
if prefix_on && reused.is_none() && !spec_eligible {
if let Some(i) = px.lookup(&pool_key, &prompt) {
let restored = {
let e = &px.entries[&pool_key][i];
match Cache::new(engine, &lm.model.cfg, ctx_cap) {
Ok(mut c) => match prefix_restore(engine, &mut c, e) {
Ok(()) => Ok(ReuseEntry {
fed: e.toks.clone(),
cache: c,
last_logits: e.last_logits.clone(),
cap: ctx_cap,
}),
Err(err) => Err(format!("restore failed: {err}")),
},
Err(err) => Err(format!("session cache alloc failed: {err}")),
}
};
match restored {
Ok(entry) => {
px.touch(&pool_key, i); px.hits += 1;
px.hit_tokens += entry.fed.len() as u64;
prefix_hit = true;
eprintln!("[prefix-cache] hit: {} of {} prompt tokens from cache (model {})",
entry.fed.len(), prompt.len(), req.model);
reused = Some(entry);
}
Err(msg) => {
if msg.starts_with("session cache alloc failed") {
let n = px.evict_all();
eprintln!("[prefix-cache] {msg}; evicted {n} entries, cold path serves");
} else {
eprintln!("[prefix-cache] {msg}; cold path serves");
}
}
}
}
if reused.is_none() {
px.misses += 1;
let l = px.best_lcp(&pool_key, &prompt);
if l >= PREFIX_CACHE_MIN_TOKENS && l < prompt.len()
&& !px.has_key(&pool_key, &prompt[..l])
{
snapshot_at = Some(l);
}
if prompt.len() >= PREFIX_CACHE_MIN_TOKENS {
seed_prefix = true; }
}
}
let (cache, seed_fed, seed_logits) = match reused {
Some(e) => {
if !prefix_hit {
eprintln!("[worker] kv-reuse: {} of {} prompt tokens resumed (model {})",
e.fed.len(), prompt.len(), req.model);
}
(Some(e.cache), e.fed, e.last_logits)
}
None => (None, Vec::new(), Vec::new()),
};
let mut params = req.params;
if !params.eos.contains(&lm.eos_id) { params.eos.push(lm.eos_id); }
for &t in &seed_fed { sampler.accept(t); }
let suffix: Vec<u32> = prompt[seed_fed.len()..].to_vec();
let prefill_done_at_admit = suffix.is_empty();
let mut spec_resumed = 0usize;
let mut text_suffix: Option<Vec<u32>> = None;
let spec = if spec_eligible && seed_fed.is_empty() {
let mut affinity_rewound: Option<(usize, &'static str)> = None;
let resumed = if constraint.is_some() { None } else {
spec_reuse.get_mut(&pool_key).and_then(|pool| {
if let Some(idx) = pool.iter().rposition(|e|
e.sess.cache_max_ctx() >= ctx_cap
&& prompt.len() >= e.sess.committed.len()
&& prompt.starts_with(&e.sess.committed)) {
return Some(pool.remove(idx).sess);
}
if !req.prompt_text.is_empty() {
if let Some(idx) = pool.iter().rposition(|e|
e.sess.cache_max_ctx() >= ctx_cap
&& req.prompt_text.len() >= e.committed_text.len()
&& req.prompt_text.starts_with(e.committed_text.as_str())) {
let e = pool.remove(idx);
let rem = &req.prompt_text[e.committed_text.len()..];
text_suffix = Some(lm.tok.encode(rem, false));
return Some(e.sess);
}
}
if !affinity_enabled() { return None; }
let req_fp = conversation_fingerprint(
&prompt, &|t| lm.tok.token_is_control(t), true);
let mut why: String = "empty pool".into();
let cand = pool.iter().enumerate().rev().find(|(_, e)| {
if e.sess.cache_max_ctx() < need {
why = format!("no room (session ctx {} < need {need})",
e.sess.cache_max_ctx());
return false;
}
let Some(pos) = e.sess.rewind_pos() else {
why = "no turn checkpoint retained".into(); return false;
};
if pos == 0 { why = "checkpoint at 0".into(); return false; }
match affinity_match(&prompt, &e.sess.committed[..pos]) {
AffinityMatch::Exact { suffix_from } if suffix_from == pos => {}
AffinityMatch::Diverged { at } => {
why = format!("history diverged at {at} of checkpoint {pos}");
return false;
}
_ => { why = "diff did not land on the checkpoint".into(); return false; }
}
if prompt.len() == pos { why = "empty suffix".into(); return false; }
let ok = match (&req.affinity, &e.affinity) {
(Some(a), Some(b)) if a == b => true,
(Some(_), _) | (_, Some(_)) => false,
_ => fingerprint_affinity(&req_fp, &e.fingerprint) >= FP_MIN_SEGMENTS,
};
if !ok { why = "identity did not nominate".into(); }
ok
}).map(|(i, e)| (i, e.affinity.is_some()));
if cand.is_none() && !pool.is_empty() {
eprintln!("[worker] spec-affinity: declined ({why}; {} parked, {} prompt \
tokens; model {})", pool.len(), prompt.len(), req.model);
}
if let Some((idx, explicit)) = cand {
let mut e = pool.remove(idx);
match lm.model.spec_rewind_to_checkpoint(engine, &mut e.sess) {
Ok(Some(pos)) => {
affinity_rewound =
Some((pos, if explicit { "explicit" } else { "fingerprint" }));
return Some(e.sess);
}
Ok(None) => {}
Err(err) => eprintln!("[worker] affinity rewind failed ({err}); \
dropping session, full prime"),
}
return None;
}
None
})};
match resumed {
Some(mut sess) => {
sess.reset_graph_fallback_on_resume();
spec_resumed = sess.committed.len();
match affinity_rewound {
Some((pos, tier)) => eprintln!(
"[worker] spec-affinity: rewound to {pos} of {} prompt tokens \
({tier}; priming {} suffix; model {})",
prompt.len(), prompt.len() - pos, req.model),
None => eprintln!(
"[worker] spec-reuse: {} committed tokens resumed{} (model {})",
spec_resumed,
if text_suffix.is_some() { " [text-prefix]" } else { "" }, req.model),
}
Some(sess)
}
None => {
if spec_sizing.evict_first.contains(&req.model) {
if let Some(n) = spec_reuse.get_mut(&pool_key)
.map(|p| { let n = p.len(); p.clear(); n }).filter(|&n| n > 0)
{
eprintln!("[worker] spec pool evicted ({n}) pre-alloc \
(learned VRAM-tight; model {})", req.model);
}
}
match lm.model.new_session(engine, ctx_cap) {
Ok(sess) => Some(sess),
Err(first_err) => {
let evicted = spec_reuse.get_mut(&pool_key)
.map(|p| { let n = p.len(); p.clear(); n }).unwrap_or(0);
if evicted > 0 {
spec_sizing.evict_first.insert(req.model.clone());
eprintln!("[worker] spec pool evicted ({evicted}) after alloc \
failure; retrying (evict-first learned)");
}
let retried = if evicted > 0 {
lm.model.new_session(engine, ctx_cap).ok()
} else { None };
match retried {
Some(sess) => Some(sess),
None => {
let mut sess = None;
if need <= ctx_cap {
let mut ask = spec_sizing.learned_ctx.get(&req.model)
.copied().unwrap_or(ctx_cap / 2)
.clamp(need, ctx_cap);
loop {
let landed = match lm.model.new_session(engine, ask) {
Ok(s) => {
let proven = spec_sizing.learned_ctx.get(&req.model)
.is_some_and(|&l| ask <= l);
let ok = lm.model.ensure_embed_resident(engine).is_ok()
&& (proven
|| engine.alloc_u8_uninit(SPEC_SHRINK_RESERVE).is_ok());
if ok { Some(s) } else { drop(s); None }
}
Err(_) => None,
};
match landed {
Some(s) => {
eprintln!("[worker] spec session right-sized: \
ctx {ask} of {ctx_cap} (prompt {} + \
budget {budget}; model {})",
prompt.len(), req.model);
spec_sizing.learned_ctx.insert(req.model.clone(), ask);
sess = Some(s);
break;
}
None if ask > need => { ask = (ask / 2).max(need); }
None => break,
}
}
}
if sess.is_none() {
eprintln!("[worker] spec session alloc failed ({first_err}); \
tokenwise path");
}
sess
}
}
}
}
}
}
} else { None };
if spec_resumed > 0 {
match (&spec, &text_suffix) {
(Some(sess), Some(_)) => { for &t in &sess.committed { sampler.accept(t); } }
_ => { for &t in &prompt[..spec_resumed] { sampler.accept(t); } }
}
}
let cache = match (&spec, cache) {
(Some(_), c) => c, (None, Some(c)) => Some(c),
(None, None) => match Cache::new(engine, &lm.model.cfg, ctx_cap) {
Ok(c) => Some(c),
Err(err) => {
let evicted = px.evict_all();
if evicted > 0 {
eprintln!("[prefix-cache] evicted {evicted} entries after cache alloc failure; retrying");
match Cache::new(engine, &lm.model.cfg, ctx_cap) {
Ok(c) => Some(c),
Err(err) => return Err((req.tx, format!("cache alloc failed: {err}"))),
}
} else {
return Err((req.tx, format!("cache alloc failed: {err}")));
}
}
},
};
let (n_prompt, n_cached) = if spec_resumed > 0 {
let suffix_len = text_suffix.as_ref().map(|t| t.len())
.unwrap_or_else(|| prompt.len() - spec_resumed);
(spec_resumed + suffix_len, spec_resumed)
} else {
(prompt.len(), seed_fed.len())
};
Ok(Session {
model: req.model,
cache_ns: req.cache_ns,
affinity: req.affinity,
lane: req.lane,
cache,
sampler,
spec,
graph: None,
graph_pending: None,
oom_retries: req_oom_retries,
replay,
spec_drafted: 0,
spec_accepted: 0,
spec_rounds: 0,
last_logits: seed_logits,
device_next: None,
constraint,
mask_dev: None,
mask_words: 0,
fed: seed_fed,
prefill_queue: if let Some(ts) = text_suffix { ts.into_iter().collect() }
else if spec_resumed > 0 { prompt[spec_resumed..].to_vec().into_iter().collect() }
else { suffix.into_iter().collect() },
prefill_done: prefill_done_at_admit,
generated: Vec::new(),
params,
stop_strings: req.stop_strings,
trace_id: req.trace_id,
emitted_bytes: 0,
budget,
n_prompt,
n_cached,
snapshot_at,
seed_prefix,
tx: req.tx,
t0: Instant::now(),
})
}
fn utf8_delta(decoded: &[u8], emitted_bytes: &mut usize) -> String {
if *emitted_bytes > decoded.len() {
return String::new();
}
let mut cursor = *emitted_bytes;
let mut delta = String::new();
while cursor < decoded.len() {
match std::str::from_utf8(&decoded[cursor..]) {
Ok(text) => {
delta.push_str(text);
cursor = decoded.len();
}
Err(err) => {
let valid = err.valid_up_to();
if valid != 0 {
delta.push_str(unsafe {
std::str::from_utf8_unchecked(&decoded[cursor..cursor + valid])
});
cursor += valid;
}
match err.error_len() {
None => break,
Some(invalid) => {
delta.push('\u{fffd}');
cursor += invalid;
}
}
}
}
}
*emitted_bytes = cursor;
delta
}
fn prefill_tick(
engine: &Engine,
loaded: &HashMap<String, LoadedModel>,
px: &mut PrefixCache,
s: &mut Session,
budget: usize,
) -> Result<usize, Box<dyn std::error::Error>> {
let lm = &loaded[&s.model];
let q = s.prefill_queue.len();
if q == 0 {
s.prefill_done = true;
maybe_prefix_seed(engine, px, s);
return Ok(0);
}
let mut consumed = 0usize;
let bound_rem = s.snapshot_at.map(|b| b - s.fed.len());
if !confidence_trace_enabled()
&& q >= memra_engine::hybrid_forward::PRIME_MIN_T.max(2)
&& budget >= memra_engine::hybrid_forward::PRIME_MIN_T
&& bound_rem.is_none_or(|r| r >= memra_engine::hybrid_forward::PRIME_MIN_T)
{
let mut take = q.min(budget);
if q - take > 0 && q - take < memra_engine::hybrid_forward::PRIME_MIN_T {
take = if q <= budget { q } else { take };
}
if let Some(r) = bound_rem {
if take >= r {
take = r; } else if r - take < memra_engine::hybrid_forward::PRIME_MIN_T {
take = (r - memra_engine::hybrid_forward::PRIME_MIN_T)
.max(memra_engine::hybrid_forward::PRIME_MIN_T);
}
}
let chunk: Vec<u32> = s.prefill_queue.drain(..take).collect();
let (l, _h, _x) = lm.model.prime_cache(engine, &chunk, s.cache.as_mut().unwrap())?;
s.last_logits = l;
for &tok in &chunk { s.fed.push(tok); s.sampler.accept(tok); }
consumed = take;
} else if let Some(tok) = s.prefill_queue.pop_front() {
s.last_logits = lm.model.decode_step(engine, tok, s.cache.as_mut().unwrap())?;
if let Some(&target) = s.prefill_queue.front() {
write_confidence_trace(s, tok, target, &s.last_logits)?;
}
s.fed.push(tok);
s.sampler.accept(tok);
consumed = 1;
}
if s.snapshot_at == Some(s.fed.len()) {
s.snapshot_at = None;
prefix_insert_from_session(engine, px, s, "lcp-split");
}
if s.prefill_queue.is_empty() {
s.prefill_done = true;
maybe_prefix_seed(engine, px, s);
}
Ok(consumed)
}
fn stage_grammar_mask(engine: &Engine, s: &mut Session) -> Result<(), String> {
s.mask_words = 0;
if s.constraint.is_none() || constrain_host() || devsample_meta(s).is_none() {
return Ok(());
}
let mask = s.constraint.as_mut().unwrap().compute_mask()?;
let words = mask.as_slice();
match s.mask_dev.as_mut() {
Some(d) if d.len() >= words.len() => {
engine.htod_u32_into(d, words).map_err(|e| e.to_string())?;
}
_ => {
let mut d = engine.alloc_u32_zeroed(words.len()).map_err(|e| e.to_string())?;
engine.htod_u32_into(&mut d, words).map_err(|e| e.to_string())?;
s.mask_dev = Some(d);
}
}
s.mask_words = words.len();
Ok(())
}
fn advance_sample_emit(
loaded: &HashMap<String, LoadedModel>,
s: &mut Session,
) -> (bool, Option<u32>) {
let lm = &loaded[&s.model];
if s.generated.len() >= s.budget {
finish(s, StopReason::MaxNew);
return (false, None);
}
let next = match (s.device_next.take(), s.constraint.as_mut()) {
(Some(t), _) => t,
(None, Some(c)) => {
let mut row = s.last_logits.clone();
if let Err(err) = c.mask_logits(&mut row) {
let _ = s.tx.send(Event::Error(format!("constraint mask: {err}")));
return (false, None);
}
s.sampler.sample(&row)
}
(None, None) => s.sampler.sample(&s.last_logits),
};
s.sampler.accept(next);
s.generated.push(next);
if s.params.eos.contains(&next) {
finish(s, StopReason::Eos);
return (false, None);
}
if let Some(c) = s.constraint.as_mut() {
if let Err(err) = c.consume(next) {
let _ = s.tx.send(Event::Error(format!("constraint advance: {err}")));
return (false, None);
}
}
let decoded = lm.tok.decode_bytes_special(&s.generated, true);
let delta = utf8_delta(&decoded, &mut s.emitted_bytes);
let full = String::from_utf8_lossy(&decoded);
if s.tx.send(Event::Token { id: next, text: delta }).is_err() {
abort_log(s);
return (false, None);
}
if !s.stop_strings.is_empty() && s.stop_strings.iter().any(|ss| full.contains(ss.as_str())) {
finish(s, StopReason::Callback);
return (false, None);
}
if s.cache.as_ref().map(|c| c.pos >= c.max_ctx).unwrap_or(false) {
finish(s, StopReason::ContextFull);
return (false, None);
}
(true, Some(next))
}
fn advance_token_emit(
loaded: &HashMap<String, LoadedModel>,
s: &mut Session,
tok: u32,
) -> (bool, ()) {
let lm = &loaded[&s.model];
if s.generated.len() >= s.budget {
finish(s, StopReason::MaxNew);
return (false, ());
}
s.sampler.accept(tok);
s.generated.push(tok);
if s.params.eos.contains(&tok) {
finish(s, StopReason::Eos);
return (false, ());
}
if let Some(c) = s.constraint.as_mut() {
if let Err(err) = c.consume(tok) {
let _ = s.tx.send(Event::Error(format!("constraint advance: {err}")));
return (false, ());
}
}
let decoded = lm.tok.decode_bytes_special(&s.generated, true);
let delta = utf8_delta(&decoded, &mut s.emitted_bytes);
let full = String::from_utf8_lossy(&decoded);
if s.tx.send(Event::Token { id: tok, text: delta }).is_err() {
abort_log(s);
return (false, ());
}
if !s.stop_strings.is_empty() && s.stop_strings.iter().any(|ss| full.contains(ss.as_str())) {
finish(s, StopReason::Callback);
return (false, ());
}
(true, ())
}
fn serve_devsample() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_SERVE_DEVSAMPLE").as_deref() != Ok("0"))
}
fn serve_leanlogits() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_SERVE_LEANLOGITS").as_deref() != Ok("0"))
}
fn constrain_host() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_CONSTRAIN_HOST").as_deref() == Ok("1"))
}
fn devsample_meta(s: &Session) -> Option<(f32, u64, u32)> {
if !serve_devsample() {
return None;
}
let sm = &s.sampler;
let no_pen = sm.penalty_repeat() == 1.0
&& sm.penalty_freq() == 0.0
&& sm.penalty_present() == 0.0;
if !no_pen || sm.top_k() != 0 || sm.top_p() < 1.0 || sm.min_p() > 0.0 {
return None;
}
if sm.is_greedy() {
Some((0.0, 0, 0))
} else {
Some((sm.temperature(), sm.seed(), s.generated.len() as u32))
}
}
fn chunk_cap_for(lm: &LoadedModel) -> usize {
if let Some(c) = std::env::var("MEMRA_DECODE_BATCH_CAP").ok().and_then(|v| v.parse().ok()) {
return usize::clamp(c, 1, 32);
}
if lm.model.decode_batch_exact16_ok() { 16 } else { 8 }
}
fn group_chunks(
active: &[Session],
ready: &[(usize, u32)],
caps: &HashMap<String, usize>,
) -> Vec<Vec<(usize, u32)>> {
let mut chunks: Vec<Vec<(usize, u32)>> = Vec::new();
for &(i, t) in ready {
let model = &active[i].model;
let cap = caps.get(model).copied().unwrap_or(8);
match chunks.last_mut() {
Some(c) if c.len() < cap && active[c[0].0].model == *model => c.push((i, t)),
_ => chunks.push(vec![(i, t)]),
}
}
chunks
}
fn step_session(
engine: &Engine,
loaded: &HashMap<String, LoadedModel>,
s: &mut Session,
spec_telem: &mut HashMap<String, memra_engine::spec::SpecTelemetry>,
) -> Result<bool, Box<dyn std::error::Error>> {
let lm = &loaded[&s.model];
if let Some(spec) = s.spec.as_mut() {
let burst_t: usize = std::env::var("MEMRA_SPEC_BURST").ok()
.and_then(|v| v.parse().ok()).unwrap_or(32);
let k: usize = std::env::var("MEMRA_SPEC_K").ok().and_then(|v| v.parse().ok()).unwrap_or(3);
let room = s.budget.saturating_sub(s.generated.len()).min(burst_t);
if room == 0 { finish(s, StopReason::MaxNew); return Ok(false); }
let suffix: Vec<u32> = s.prefill_queue.drain(..).collect();
s.prefill_done = true;
if suffix.is_empty() && spec.next_pred.is_none() && spec.pending_tok.is_none() {
finish(s, StopReason::MaxNew); return Ok(false);
}
let sampling = if s.sampler.temperature() > 0.0 {
Some(memra_engine::spec::SpecSampling {
temp: s.sampler.temperature(),
seed: s.sampler.seed(),
top_k: s.sampler.top_k() as i32,
top_p: s.sampler.top_p(),
min_p: s.sampler.min_p(),
penalty_last_n: s.sampler.penalty_last_n(),
penalty_repeat: s.sampler.penalty_repeat(),
penalty_freq: s.sampler.penalty_freq(),
penalty_present: s.sampler.penalty_present(),
})
} else { None };
let telem_before = spec.telem;
let per_burst_emit = std::env::var("MEMRA_SSE_PER_BURST").as_deref() == Ok("1");
let admit_yield = std::env::var("MEMRA_ADMIT_YIELD").as_deref() != Ok("0");
let mut vis: Vec<u32> = s.generated.clone();
let mut cursor = s.emitted_bytes;
let mut eos_seen = false;
let mut send_ok = true;
let flush_tx = s.tx.clone();
let eos_ids = s.params.eos.clone();
let tok_ref = &lm.tok;
let mut flush_cb = |slice: &[u32]| -> bool {
let keep = !admit_yield
|| PENDING_ADMITS.load(std::sync::atomic::Ordering::Acquire) == 0;
if per_burst_emit || eos_seen || slice.is_empty() {
return keep;
}
let mut last_id = 0u32;
for &t in slice {
if eos_ids.contains(&t) {
eos_seen = true;
break;
}
vis.push(t);
last_id = t;
}
if !send_ok {
return keep; }
let decoded = tok_ref.decode_bytes_special(&vis, true);
let delta = utf8_delta(&decoded, &mut cursor);
if !delta.is_empty()
&& flush_tx.send(Event::Token { id: last_id, text: delta }).is_err()
{
send_ok = false;
}
keep
};
let on_commit: Option<&mut dyn FnMut(&[u32]) -> bool> =
if per_burst_emit && !admit_yield { None } else { Some(&mut flush_cb) };
let (burst, d, a) = match s.constraint.as_mut() {
Some(c) => {
let mut g = crate::constrained::SpecGrammar::new(c, lm.eos_id);
lm.model.generate_spec_session_constrained(
engine, spec, &suffix, room, k, sampling, Some(&mut g), on_commit)?
}
None => lm.model.generate_spec_session_sampled(
engine, spec, &suffix, room, k, sampling, on_commit)?,
};
let telem_delta = spec.telem.delta_since(&telem_before);
spec_telem.entry(s.model.clone()).or_default().merge(&telem_delta);
s.spec_rounds += telem_delta.rounds;
s.spec_drafted += d;
s.spec_accepted += a;
if d > 0 {
eprintln!("[spec-acc] ctx={} burst={}/{} cum={}/{}={:.3}",
s.fed.len() + suffix.len(), a, d, s.spec_accepted, s.spec_drafted,
s.spec_accepted as f64 / s.spec_drafted.max(1) as f64);
}
for &tok in &suffix { s.fed.push(tok); s.sampler.accept(tok); }
let mut stop: Option<StopReason> = None;
for &tok in &burst {
s.sampler.accept(tok);
s.generated.push(tok);
s.fed.push(tok);
if s.params.eos.contains(&tok) { stop = Some(StopReason::Eos); break; }
}
let visible = match stop {
Some(StopReason::Eos) => &s.generated[..s.generated.len() - 1],
_ => &s.generated[..],
};
s.emitted_bytes = cursor;
if !send_ok {
abort_log(s);
return Ok(false);
}
let decoded = lm.tok.decode_bytes_special(visible, true);
let delta = utf8_delta(&decoded, &mut s.emitted_bytes);
let full = String::from_utf8_lossy(&decoded);
if !delta.is_empty()
&& s.tx.send(Event::Token { id: *burst.last().unwrap_or(&0), text: delta }).is_err()
{
abort_log(s);
return Ok(false);
}
if stop.is_none() && !s.stop_strings.is_empty()
&& s.stop_strings.iter().any(|ss| full.contains(ss.as_str())) {
stop = Some(StopReason::Callback);
}
if stop.is_none() && s.generated.len() >= s.budget { stop = Some(StopReason::MaxNew); }
if stop.is_none() && spec.committed.len() + k + 3 >= spec.cache_max_ctx() {
stop = Some(StopReason::ContextFull);
}
if let Some(r) = stop { finish(s, r); return Ok(false); }
return Ok(true);
}
if !s.prefill_done {
let q = s.prefill_queue.len();
if !confidence_trace_enabled() && q >= memra_engine::hybrid_forward::PRIME_MIN_T.max(2) {
let mut take = q.min(PREFILL_TICK_T);
if q - take > 0 && q - take < memra_engine::hybrid_forward::PRIME_MIN_T { take = q; }
let chunk: Vec<u32> = s.prefill_queue.drain(..take).collect();
let (l, _h, _x) = lm.model.prime_cache(engine, &chunk, s.cache.as_mut().unwrap())?;
s.last_logits = l;
for &tok in &chunk { s.fed.push(tok); s.sampler.accept(tok); }
} else if let Some(tok) = s.prefill_queue.pop_front() {
s.last_logits = lm.model.decode_step(engine, tok, s.cache.as_mut().unwrap())?;
if let Some(&target) = s.prefill_queue.front() {
write_confidence_trace(s, tok, target, &s.last_logits)?;
}
s.fed.push(tok);
s.sampler.accept(tok);
}
if s.prefill_queue.is_empty() { s.prefill_done = true; }
return Ok(true);
}
if s.generated.len() >= s.budget {
finish(s, StopReason::MaxNew);
return Ok(false);
}
let next = match (s.device_next.take(), s.constraint.as_mut()) {
(Some(t), _) => t,
(None, Some(c)) => {
let mut row = s.last_logits.clone();
c.mask_logits(&mut row).map_err(|e| format!("constraint mask: {e}"))?;
s.sampler.sample(&row)
}
(None, None) => s.sampler.sample(&s.last_logits),
};
s.sampler.accept(next);
s.generated.push(next);
if s.params.eos.contains(&next) {
finish(s, StopReason::Eos);
return Ok(false);
}
if let Some(c) = s.constraint.as_mut() {
c.consume(next).map_err(|e| format!("constraint advance: {e}"))?;
}
let decoded = lm.tok.decode_bytes_special(&s.generated, true);
let delta = utf8_delta(&decoded, &mut s.emitted_bytes);
let full = String::from_utf8_lossy(&decoded);
if s.tx.send(Event::Token { id: next, text: delta }).is_err() {
abort_log(s);
return Ok(false);
}
if !s.stop_strings.is_empty() && s.stop_strings.iter().any(|ss| full.contains(ss.as_str())) {
finish(s, StopReason::Callback);
return Ok(false);
}
if s.cache.as_ref().map(|c| c.pos >= c.max_ctx).unwrap_or(false) {
finish(s, StopReason::ContextFull);
return Ok(false);
}
s.last_logits = lm.model.decode_step(engine, next, s.cache.as_mut().unwrap())?;
s.fed.push(next);
Ok(true)
}
fn confidence_trace_enabled() -> bool {
std::env::var("MEMRA_CONFIDENCE_TRACE").is_ok()
}
#[derive(Debug)]
struct ConfidenceSummary {
reference_logprob: f64,
top1_token: u32,
top1_correct: bool,
top1_top2_margin: f32,
entropy: f64,
}
fn summarize_confidence(logits: &[f32], target: u32) -> Result<ConfidenceSummary, String> {
let target = target as usize;
if logits.is_empty() || target >= logits.len() {
return Err(format!("target token {target} outside {} logits", logits.len()));
}
let mut top1 = (0usize, f32::NEG_INFINITY);
let mut top2 = f32::NEG_INFINITY;
for (index, &logit) in logits.iter().enumerate() {
if logit > top1.1 {
top2 = top1.1;
top1 = (index, logit);
} else if logit > top2 {
top2 = logit;
}
}
let max_logit = top1.1 as f64;
let mut sum_exp = 0.0f64;
let mut weighted_logit = 0.0f64;
for &logit in logits {
let exp = ((logit as f64) - max_logit).exp();
sum_exp += exp;
weighted_logit += exp * logit as f64;
}
let logsumexp = max_logit + sum_exp.ln();
Ok(ConfidenceSummary {
reference_logprob: logits[target] as f64 - logsumexp,
top1_token: top1.0 as u32,
top1_correct: top1.0 == target,
top1_top2_margin: top1.1 - top2,
entropy: logsumexp - weighted_logit / sum_exp,
})
}
fn write_confidence_trace(
session: &Session,
input_token: u32,
target_token: u32,
logits: &[f32],
) -> Result<(), Box<dyn std::error::Error>> {
let Ok(path) = std::env::var("MEMRA_CONFIDENCE_TRACE") else { return Ok(()) };
let summary = summarize_confidence(logits, target_token).map_err(std::io::Error::other)?;
let record = serde_json::json!({
"format": "memra-token-confidence-v1",
"trace_id": session.trace_id,
"input_position": session.fed.len(),
"input_token": input_token,
"target_token": target_token,
"reference_logprob": summary.reference_logprob,
"top1_token": summary.top1_token,
"top1_correct": summary.top1_correct,
"top1_top2_margin": summary.top1_top2_margin,
"entropy": summary.entropy,
});
let mut file = std::fs::OpenOptions::new().create(true).append(true).open(path)?;
writeln!(file, "{record}")?;
Ok(())
}
fn abort_log(s: &Session) {
eprintln!("[abort] client disconnected: model {:?}, prompt {} ({} cached), \
{} generated — billed to abort point, {:.2}s",
s.model, s.n_prompt, s.n_cached, s.generated.len(),
s.t0.elapsed().as_secs_f64());
}
fn finish(s: &Session, reason: StopReason) {
let elapsed = s.t0.elapsed().as_secs_f64();
if let Some(c) = s.constraint.as_ref() {
if c.steps > 0 {
eprintln!("[constrained] {}: {} masked steps, mask total {:.2} ms ({:.3} ms/step)",
s.model, c.steps, c.mask_ns as f64 / 1e6,
c.mask_ns as f64 / 1e6 / c.steps as f64);
}
if c.spec_clones > 0 {
eprintln!("[draft-mask] {}: {} clones {:.2} ms ({:.3} ms/clone), \
{} draft masks {:.2} ms ({:.3} ms/mask)",
s.model, c.spec_clones, c.spec_ns as f64 / 1e6,
c.spec_ns as f64 / 1e6 / c.spec_clones as f64,
c.draft_masks, c.draft_mask_ns as f64 / 1e6,
c.draft_mask_ns as f64 / 1e6 / c.draft_masks.max(1) as f64);
}
}
let reason = format!("{reason:?}");
let spec = (s.spec_rounds > 0).then(|| SpecUsage {
rounds: s.spec_rounds,
drafted: s.spec_drafted as u64,
accepted: s.spec_accepted as u64,
});
let _ = s.tx.send(Event::Done {
stop_reason: reason,
n_tokens: s.generated.len(),
n_prompt: s.n_prompt,
n_cached: s.n_cached,
elapsed_s: elapsed,
spec,
});
}
#[allow(clippy::type_complexity)]
pub fn spawn(models: Vec<(String, String, Option<String>)>)
-> Result<(Sender<Cmd>, Arc<Vec<String>>, Arc<HashMap<String, ModelCaps>>, SharedMetrics), String> {
let (cmd_tx, cmd_rx) = std::sync::mpsc::channel::<Cmd>();
let (ready_tx, ready_rx) =
std::sync::mpsc::channel::<Result<(Vec<String>, HashMap<String, ModelCaps>), String>>();
let metrics: SharedMetrics = Default::default();
let m2 = metrics.clone();
std::thread::Builder::new()
.name("memra-gpu-worker".into())
.spawn(move || run(models, cmd_rx, ready_tx, m2))
.map_err(|e| format!("spawn worker thread: {e}"))?;
match ready_rx.recv() {
Ok(Ok((names, caps))) => Ok((cmd_tx, Arc::new(names), Arc::new(caps), metrics)),
Ok(Err(err)) => Err(err),
Err(_) => Err("worker died during init".into()),
}
}
#[cfg(test)]
mod tests {
use super::{summarize_confidence, utf8_delta};
use super::{PoolKey, PrefixCache, PrefixEntry, PREFIX_CACHE_MIN_TOKENS};
fn entry(toks: Vec<u32>) -> PrefixEntry {
PrefixEntry {
toks,
kv: Vec::new(),
conv: Vec::new(),
ssm: Vec::new(),
pos: 0,
last_logits: vec![0.0],
bytes: 1,
last_use: std::time::Instant::now(),
id: 0,
}
}
fn key(ns: &str) -> PoolKey {
("m".to_string(), ns.to_string())
}
fn toks(n: usize) -> Vec<u32> {
(0..n as u32).collect()
}
#[test]
fn prefix_cache_same_namespace_same_prefix_hits() {
let mut px = PrefixCache::default();
let prefix = toks(PREFIX_CACHE_MIN_TOKENS);
px.insert(&key("tenant-a"), entry(prefix.clone()), "test");
assert!(px.lookup(&key("tenant-a"), &toks(PREFIX_CACHE_MIN_TOKENS + 32)).is_some());
assert!(px.has_covering(&key("tenant-a"), &prefix));
assert_eq!(px.best_lcp(&key("tenant-a"), &prefix), prefix.len());
}
#[test]
fn prefix_cache_namespaces_isolate_both_directions() {
let mut px = PrefixCache::default();
let prompt = toks(PREFIX_CACHE_MIN_TOKENS + 32);
px.insert(&key("tenant-a"), entry(toks(PREFIX_CACHE_MIN_TOKENS)), "test");
assert!(px.lookup(&key("tenant-b"), &prompt).is_none());
assert!(px.lookup(&key(""), &prompt).is_none());
assert_eq!(px.best_lcp(&key("tenant-b"), &prompt), 0);
assert!(!px.has_covering(&key("tenant-b"), &prompt));
px.insert(&key("tenant-b"), entry(toks(PREFIX_CACHE_MIN_TOKENS)), "test");
assert_eq!(px.n_entries(), 2);
assert!(px.lookup(&key("tenant-a"), &prompt).is_some());
assert!(px.lookup(&key("tenant-b"), &prompt).is_some());
assert!(px.lookup(&key("tenant-c"), &prompt).is_none());
}
#[test]
fn prefix_cache_default_namespace_preserves_single_tenant_behavior() {
let mut px = PrefixCache::default();
let short = toks(PREFIX_CACHE_MIN_TOKENS);
let long = toks(PREFIX_CACHE_MIN_TOKENS + 16);
px.insert(&key(""), entry(short.clone()), "test");
px.insert(&key(""), entry(long.clone()), "test");
px.insert(&key(""), entry(long.clone()), "test"); assert_eq!(px.n_entries(), 2);
let hit = px.lookup(&key(""), &toks(PREFIX_CACHE_MIN_TOKENS + 64)).unwrap();
assert_eq!(px.entries[&key("")][hit].toks.len(), long.len());
assert!(px.lookup(&key(""), &toks(PREFIX_CACHE_MIN_TOKENS - 1)).is_none());
}
fn entry_b(ident: u32, bytes: usize) -> PrefixEntry {
PrefixEntry {
toks: vec![ident],
kv: Vec::new(),
conv: Vec::new(),
ssm: Vec::new(),
pos: 0,
last_logits: vec![0.0],
bytes,
last_use: next_instant(),
id: 0,
}
}
fn next_instant() -> std::time::Instant {
let t = std::time::Instant::now();
loop {
let u = std::time::Instant::now();
if u > t {
return u;
}
}
}
struct OldModel {
entries: Vec<(u32, usize, u64)>,
total: usize,
clock: u64,
victims: Vec<u32>,
}
impl OldModel {
fn insert(&mut self, ident: u32, bytes: usize, budget: usize) {
if bytes > budget {
return;
}
self.clock += 1;
self.entries.push((ident, bytes, self.clock));
self.total += bytes;
while self.total > budget {
let Some(&(v, b, _)) = self.entries.iter().min_by_key(|&&(_, _, o)| o) else {
break;
};
self.entries.retain(|&(i, _, _)| i != v);
self.total -= b;
self.victims.push(v);
}
}
fn touch(&mut self, ident: u32) {
self.clock += 1;
if let Some(e) = self.entries.iter_mut().find(|e| e.0 == ident) {
e.2 = self.clock;
}
}
fn survivors(&self) -> Vec<u32> {
let mut v: Vec<u32> = self.entries.iter().map(|e| e.0).collect();
v.sort_unstable();
v
}
}
fn px_survivors(px: &PrefixCache) -> Vec<u32> {
let mut v: Vec<u32> = px.entries.values().flatten().map(|e| e.toks[0]).collect();
v.sort_unstable();
v
}
#[test]
fn prefix_cache_eviction_matches_old_policy_on_recorded_pattern() {
const BUDGET: usize = 24;
let mut px = PrefixCache::default();
let mut old = OldModel { entries: Vec::new(), total: 0, clock: 0, victims: Vec::new() };
let mut rng: u64 = 0x9E3779B97F4A7C15;
let mut step = || {
rng = rng.wrapping_mul(6364136223846793005).wrapping_add(1442695040888963407);
(rng >> 33) as usize
};
let namespaces = ["", "tenant-a", "tenant-b"];
let mut ident: u32 = 0;
let mut placed: Vec<(u32, PoolKey)> = Vec::new(); for _ in 0..400 {
let survivors = px_survivors(&px);
if step() % 3 == 2 && !survivors.is_empty() {
let tgt = survivors[step() % survivors.len()];
let k = placed.iter().find(|(i, _)| *i == tgt).unwrap().1.clone();
let idx = px.entries[&k].iter().position(|e| e.toks[0] == tgt).unwrap();
next_instant(); px.touch(&k, idx);
old.touch(tgt);
} else {
let k = key(namespaces[step() % namespaces.len()]);
let bytes = 1 + step() % 8;
px.insert_with_budget(&k, entry_b(ident, bytes), "test", BUDGET);
old.insert(ident, bytes, BUDGET);
placed.push((ident, k));
ident += 1;
}
assert_eq!(px_survivors(&px), old.survivors());
assert_eq!(px.total_bytes, old.total);
}
assert_eq!(px.evictions as usize, old.victims.len());
assert!(old.victims.len() > 50, "pattern too tame to prove anything: {} evictions",
old.victims.len());
}
#[test]
fn prefix_cache_touch_rescues_the_would_be_victim() {
let mut px = PrefixCache::default();
px.insert_with_budget(&key(""), entry_b(0, 4), "test", 8);
px.insert_with_budget(&key(""), entry_b(1, 4), "test", 8);
let idx = px.entries[&key("")].iter().position(|e| e.toks[0] == 0).unwrap();
next_instant();
px.touch(&key(""), idx);
px.insert_with_budget(&key(""), entry_b(2, 4), "test", 8);
assert_eq!(px_survivors(&px), vec![0, 2], "touched 0 must survive, untouched 1 evicts");
}
#[test]
fn prefix_cache_eviction_large_pool_flush_smoke() {
const E: usize = 10_000;
let mut px = PrefixCache::default();
for i in 0..E {
px.insert_with_budget(&key(""), entry_b(i as u32, 1), "test", E);
}
assert_eq!(px.n_entries(), E);
let t0 = std::time::Instant::now();
px.insert_with_budget(&key(""), entry_b(u32::MAX, E / 2), "test", E);
let dt = t0.elapsed();
assert_eq!(px.n_entries(), E / 2 + 1);
assert_eq!(px.evictions as usize, E / 2);
assert_eq!(px.total_bytes, E);
let survivors = px_survivors(&px);
assert!(survivors.contains(&u32::MAX));
assert!(!survivors.contains(&0) && !survivors.contains(&((E / 2 - 1) as u32)),
"victims must be the oldest half");
assert!(survivors.contains(&((E / 2) as u32)), "newest half survives");
assert!(dt < std::time::Duration::from_secs(2),
"large-E flush took {dt:?} — eviction is scaling with pool size again");
}
const IM: u32 = 1000;
fn is_marker(t: u32) -> bool {
t == IM
}
fn convo(segs: &[&[u32]]) -> Vec<u32> {
let mut v = Vec::new();
for s in segs {
v.push(IM);
v.extend_from_slice(s);
}
v
}
fn fp(toks: &[u32]) -> Vec<u64> {
super::conversation_fingerprint(toks, &is_marker, true)
}
fn fp_parked(toks: &[u32]) -> Vec<u64> {
super::conversation_fingerprint(toks, &is_marker, false)
}
fn shared(a: &[u64], b: &[u64]) -> usize {
super::fingerprint_affinity(a, b)
}
fn body(tag: u32, n: usize) -> Vec<u32> {
(0..n as u32).map(|i| tag * 100 + i).collect()
}
#[test]
fn fingerprint_survives_an_assistant_interior_rewrite() {
let sys = body(1, 24);
let user1 = body(2, 24);
let mut asst1 = body(3, 40);
let user2 = body(4, 24);
let live = body(9, 8);
let before = convo(&[&sys, &user1, &asst1, &user2, &live]);
asst1.drain(super::FP_WINDOW..asst1.len() - super::FP_WINDOW);
let after = convo(&[&sys, &user1, &asst1, &user2, &live]);
assert_ne!(before, after, "the rewrite must actually change the token stream");
assert!(!after.starts_with(&before[..before.len() - 1]),
"the rewrite must break plain prefix-extension (else the old probe would hit)");
assert_eq!(fp(&before), fp(&after));
assert!(fp(&before).len() >= super::FP_MIN_SEGMENTS);
}
#[test]
fn fingerprint_nominates_the_parked_session_across_a_rewritten_turn() {
let (sys, user1, user2, live) = (body(1, 24), body(2, 24), body(4, 24), body(9, 8));
let mut asst1 = body(3, 40);
let parked = convo(&[&sys, &user1, &asst1]);
asst1.drain(super::FP_WINDOW..asst1.len() - super::FP_WINDOW);
let request = convo(&[&sys, &user1, &asst1, &user2, &live]);
let n = shared(&fp(&request), &fp_parked(&parked));
assert_eq!(n, 3, "system + user1 + rewritten assistant1 all match");
assert!(n >= super::FP_MIN_SEGMENTS, "clears the nomination bar");
}
#[test]
fn fingerprint_degrades_gracefully_when_a_rewrite_reaches_a_head_window() {
let (sys, user1, user2, live) = (body(1, 24), body(2, 24), body(4, 24), body(9, 8));
let asst1 = body(3, 40);
let parked = convo(&[&sys, &user1, &asst1, &user2]);
let mut wrecked = asst1.clone();
wrecked.drain(..super::FP_WINDOW); let request = convo(&[&sys, &user1, &wrecked, &user2, &live]);
let n = shared(&fp(&request), &fp_parked(&parked));
assert_eq!(n, 2, "shared run ends at the damaged segment, not at zero");
}
#[test]
fn fingerprint_ignores_the_live_turn() {
let (sys, user1, asst1) = (body(1, 24), body(2, 24), body(3, 24));
let turn_a = convo(&[&sys, &user1, &asst1, &body(7, 12)]);
let turn_b = convo(&[&sys, &user1, &asst1, &body(8, 30)]);
assert_eq!(fp(&turn_a), fp(&turn_b));
}
#[test]
fn fingerprint_separates_different_conversations() {
let (sys, user1, asst1, live) = (body(1, 24), body(2, 24), body(3, 24), body(9, 8));
let base = fp(&convo(&[&sys, &user1, &asst1, &live]));
let other_sys = fp(&convo(&[&body(5, 24), &user1, &asst1, &live]));
let other_user = fp(&convo(&[&sys, &body(6, 24), &asst1, &live]));
assert_eq!(shared(&base, &other_sys), 0, "different system prompt: nothing shared");
assert_eq!(shared(&base, &other_user), 1, "only the system prompt is shared");
assert!(shared(&base, &other_user) < super::FP_MIN_SEGMENTS, "below the bar");
}
#[test]
fn fingerprint_declines_short_generic_openers() {
let sys = body(1, 24);
let a = fp(&convo(&[&sys, &body(2, 24)]));
let b = fp(&convo(&[&sys, &body(7, 24)]));
assert!(shared(&a, &b) < super::FP_MIN_SEGMENTS);
let long = convo(&[&sys, &body(2, 24), &body(3, 24), &body(9, 8)]);
assert!(shared(&fp(&long), &fp(&long)) >= super::FP_MIN_SEGMENTS);
}
#[test]
fn fingerprint_handles_a_prompt_with_no_markers() {
assert!(fp(&toks(512)).is_empty());
assert!(shared(&fp(&toks(512)), &fp_parked(&toks(512))) < super::FP_MIN_SEGMENTS);
}
#[test]
fn affinity_resume_requires_the_whole_committed_prefix() {
use super::{affinity_match, AffinityMatch};
assert_eq!(
affinity_match(&toks(100), &toks(60)),
AffinityMatch::Exact { suffix_from: 60 }
);
assert_eq!(
affinity_match(&toks(60), &toks(60)),
AffinityMatch::Exact { suffix_from: 60 }
);
}
#[test]
fn affinity_refuses_to_resume_across_a_committed_range_divergence() {
use super::{affinity_match, AffinityMatch};
let mut prompt = toks(100);
prompt[42] = 999;
assert_eq!(
affinity_match(&prompt, &toks(60)),
AffinityMatch::Diverged { at: 42 }
);
assert_eq!(
affinity_match(&toks(40), &toks(60)),
AffinityMatch::Diverged { at: 40 }
);
}
#[test]
fn affinity_room_test_accepts_f5_right_sized_sessions() {
let (prompt_len, budget, ctx_cap) = (12_000usize, 512usize, 131_072usize);
let need = prompt_len + budget + super::SPEC_SHRINK_SLACK;
let laddered = 16_384usize; assert!(laddered < ctx_cap, "the ladder lands below the cap (else no interaction)");
assert!(laddered >= need, "and still covers what this request needs -> eligible");
assert!(8_192 < need);
}
#[test]
fn streaming_utf8_waits_for_a_complete_multibyte_sequence() {
let mut emitted = 0;
assert_eq!(utf8_delta(b"caf\xc3", &mut emitted), "caf");
assert_eq!(emitted, 3);
assert_eq!(utf8_delta(b"caf\xc3\xa9\n", &mut emitted), "é\n");
assert_eq!(emitted, 6);
}
#[test]
fn streaming_utf8_consumes_truly_invalid_bytes_once() {
let mut emitted = 0;
assert_eq!(utf8_delta(b"a\xffb", &mut emitted), "a\u{fffd}b");
assert_eq!(emitted, 3);
assert_eq!(utf8_delta(b"a\xffbc", &mut emitted), "c");
}
#[test]
fn confidence_summary_tracks_reference_and_margin() {
let summary = summarize_confidence(&[0.0, 2.0, 1.0], 1).unwrap();
assert_eq!(summary.top1_token, 1);
assert!(summary.top1_correct);
assert!((summary.top1_top2_margin - 1.0).abs() < 1e-6);
let expected = 2.0f64 - (0.0f64.exp() + 2.0f64.exp() + 1.0f64.exp()).ln();
assert!((summary.reference_logprob - expected).abs() < 1e-12);
assert!(summary.entropy > 0.0);
}
}