use std::collections::{HashMap, VecDeque};
use std::io::Write as _;
use std::sync::Arc;
use std::sync::mpsc::{Receiver, Sender};
use std::time::{Duration, 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;
const SOLO_PREFILL_TICK_T: usize = 8192;
fn interactive_prefill_budget(
configured: usize,
configured_explicitly: bool,
sole_unfinished: bool,
fresh: bool,
queued: usize,
) -> usize {
if configured_explicitly || !sole_unfinished || !fresh {
return configured;
}
let mut widened = queued.min(SOLO_PREFILL_TICK_T);
let tail = queued - widened;
if tail > 0 && tail < memra_engine::hybrid_forward::PRIME_MIN_T {
widened = queued;
}
configured.max(widened)
}
fn tick_trace_enabled() -> bool {
static ENABLED: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ENABLED.get_or_init(|| std::env::var("MEMRA_TICK_TRACE").as_deref() == Ok("1"))
}
struct LoadedModel {
model: HybridModel,
tok: Arc<Tokenizer>,
eos_id: u32,
from_dir: bool,
constraints: crate::constrained::ConstraintCompiler,
}
#[derive(Debug, Clone)]
pub enum Event {
PromptUsage { n_prompt: usize, n_cached: usize },
Token { id: u32, text: String },
TokenSnapshot(Vec<u32>),
Done { stop_reason: String, n_tokens: usize, n_prompt: usize, n_cached: usize,
elapsed_s: f64, spec: Option<SpecUsage> },
Error(EngineError),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ErrClass {
InvalidRequest,
ContextLength,
ModelNotFound,
RateLimit,
Overloaded,
Engine,
}
#[derive(Debug, Clone)]
pub struct EngineError {
pub class: ErrClass,
pub message: String,
pub param: Option<&'static str>,
}
impl EngineError {
pub fn invalid_param(message: impl Into<String>, param: &'static str) -> Self {
Self { class: ErrClass::InvalidRequest, message: message.into(), param: Some(param) }
}
pub fn context_length(message: impl Into<String>) -> Self {
Self { class: ErrClass::ContextLength, message: message.into(), param: Some("messages") }
}
pub fn model_not_found(message: impl Into<String>) -> Self {
Self { class: ErrClass::ModelNotFound, message: message.into(), param: Some("model") }
}
pub fn rate_limit(message: impl Into<String>) -> Self {
Self { class: ErrClass::RateLimit, message: message.into(), param: None }
}
pub fn overloaded(message: impl Into<String>) -> Self {
Self { class: ErrClass::Overloaded, message: message.into(), param: None }
}
pub fn engine(message: impl Into<String>) -> Self {
let message = message.into();
let class = if is_cuda_oom(&message) { ErrClass::Overloaded } else { ErrClass::Engine };
Self { class, message, param: None }
}
}
#[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 reasoning_effort: Option<String>,
pub params: GenParams,
pub sampler_cfg: SamplerConfig,
pub stop_strings: Vec<String>,
pub trace_id: Option<String>,
pub max_prompt_tokens: Option<usize>,
pub cache_ns: String,
pub affinity: Option<String>,
pub lane: crate::lanes::Lane,
pub oom_retries: u32,
pub spec_k_replay: Option<usize>,
pub grammar: Option<crate::constrained::GrammarSpec>,
pub(crate) prepared_constraint: Option<crate::constrained::SessionConstraint>,
pub(crate) constraint_ready:
Option<tokio::sync::oneshot::Sender<Result<(), EngineError>>>,
pub(crate) prepared_prompt: Option<Vec<u32>>,
pub ttft: Option<Arc<crate::ttft::Trace>>,
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 effort_levels: bool,
pub gemma_think: bool,
pub chat_temperature_default: Option<f32>,
pub chat_top_p_default: Option<f32>,
}
pub enum Cmd {
Generate(Box<Request>),
}
const CONSTRAINT_RESULT_POLL: Duration = Duration::from_millis(5);
struct PendingConstraintCompile {
request: Box<Request>,
deadline: Instant,
}
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 prefix_misses: u64,
pub prefix_inserts: u64,
pub prefix_evictions: u64,
pub prefix_hit_tokens: u64,
pub admission_session_defers: u64,
pub admission_vram_defers: u64,
pub step_oom_parks: u64,
pub continuation_pool_hits: u64,
pub continuation_pool_evictions: u64,
pub plain_affinity_rewinds: u64,
pub spec_pool_hits: u64,
pub spec_pool_misses: u64,
pub spec_pool_affinity_rewinds: u64,
pub spec_pool_evictions: u64,
pub active_sessions: u64,
pub queued_requests: u64,
pub continuation_pool_entries: u64,
pub spec_pool_entries: u64,
pub cuda_driver_free_bytes: u64,
pub cuda_pool_reserved_bytes: u64,
pub cuda_pool_used_bytes: u64,
pub cuda_pool_cached_bytes: u64,
pub lcp_hist: [u64; 11],
pub ns_tokens: HashMap<String, [u64; 2]>,
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 spec_window: HashMap<String, memra_engine::spec::SpecTelemetry>,
pub adsd_suspect_total: HashMap<String, u64>,
pub constraint_compiler_fail_closed:
HashMap<String, Arc<std::sync::atomic::AtomicBool>>,
}
pub type SharedMetrics = std::sync::Arc<std::sync::Mutex<Metrics>>;
use crate::lanes::{Lane, StepStats};
struct ReuseEntry {
fed: Vec<u32>,
cache: Cache,
last_logits: Vec<f32>,
cap: usize,
ckpt: Option<PlainCheckpoint>,
affinity: Option<String>,
fingerprint: Vec<u64>,
parked_at: Instant,
}
struct PlainCheckpoint {
snap: memra_engine::cache::CacheSnapshot,
pos: usize,
last_logits: Vec<f32>,
}
struct SpecReuseEntry {
sess: memra_engine::spec::SpecSession,
committed_text: String,
affinity: Option<String>,
fingerprint: Vec<u64>,
parked_at: Instant,
}
trait ParkedEntryAge {
fn parked_at(&self) -> Instant;
}
impl ParkedEntryAge for ReuseEntry {
fn parked_at(&self) -> Instant { self.parked_at }
}
impl ParkedEntryAge for SpecReuseEntry {
fn parked_at(&self) -> Instant { self.parked_at }
}
fn reuse_pool_per_namespace() -> 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 DEFAULT_REUSE_POOL_GLOBAL_CAP: usize = 16;
fn reuse_pool_global_cap() -> usize {
static N: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*N.get_or_init(|| std::env::var("MEMRA_REUSE_POOL_GLOBAL_CAP").ok()
.and_then(|v| v.parse().ok()).unwrap_or(DEFAULT_REUSE_POOL_GLOBAL_CAP))
}
const REUSE_MIN_PREFIX: usize = 16;
#[derive(Default)]
struct SpecSizing {
evict_first: std::collections::HashSet<String>,
learned_ctx: HashMap<String, usize>,
}
#[derive(Default)]
struct ReuseMetrics {
continuation_hits: u64,
continuation_evictions: u64,
spec_hits: u64,
spec_misses: u64,
spec_affinity_rewinds: u64,
spec_evictions: u64,
plain_affinity_rewinds: u64,
}
fn context_cache_bytes(
bytes_per_token: usize,
ring_bytes_per_token: usize,
ring_rows: usize,
ctx_cap: usize,
) -> usize {
let flat = bytes_per_token.saturating_sub(ring_bytes_per_token);
flat.saturating_mul(ctx_cap).saturating_add(
ring_bytes_per_token.saturating_mul(ctx_cap.min(ring_rows)),
)
}
#[derive(Debug)]
struct AdmissionCostModel {
plain_bytes_per_token: usize,
spec_bytes_per_token: usize,
plain_ring_bytes_per_token: usize,
spec_ring_bytes_per_token: usize,
ring_rows: usize,
activation_bytes: usize,
last_logged: Option<(usize, bool, usize)>,
}
impl AdmissionCostModel {
fn new(model: &HybridModel) -> Self {
let (plain_bytes_per_token, plain_ring_bytes_per_token, plain_ring_rows) =
model.plain_session_kv_shape();
let (spec_bytes_per_token, spec_ring_bytes_per_token, spec_ring_rows) =
model.spec_session_kv_shape();
debug_assert!(plain_ring_rows == 0 || spec_ring_rows == 0 || plain_ring_rows == spec_ring_rows);
Self {
plain_bytes_per_token,
spec_bytes_per_token,
plain_ring_bytes_per_token,
spec_ring_bytes_per_token,
ring_rows: plain_ring_rows.max(spec_ring_rows),
activation_bytes: 0,
last_logged: None,
}
}
fn bytes_per_token(&self, spec: bool) -> usize {
if spec {
self.spec_bytes_per_token
} else {
self.plain_bytes_per_token
}
}
fn ring_bytes_per_token(&self, spec: bool) -> usize {
if spec {
self.spec_ring_bytes_per_token
} else {
self.plain_ring_bytes_per_token
}
}
fn context_bytes(&self, ctx_cap: usize, spec: bool) -> usize {
context_cache_bytes(
self.bytes_per_token(spec),
self.ring_bytes_per_token(spec),
self.ring_rows,
ctx_cap,
)
}
fn estimate(&self, ctx_cap: usize, spec: bool) -> usize {
self.context_bytes(ctx_cap, spec)
.saturating_add(self.activation_bytes)
}
fn observe(&mut self, observed_bytes: usize, ctx_cap: usize, spec: bool) -> Option<usize> {
let context = self.context_bytes(ctx_cap, spec);
let residual = observed_bytes.saturating_sub(context);
if residual > self.activation_bytes {
self.activation_bytes = residual;
Some(residual)
} else {
None
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct RequestShape {
ctx_cap: usize,
budget: usize,
need: usize,
}
impl RequestShape {
fn admission_cap(self) -> usize {
self.ctx_cap.max(self.need)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum ParkedPool {
Plain,
Spec,
}
#[derive(Debug, Clone, PartialEq, Eq)]
struct ParkedCandidate {
pool: ParkedPool,
key: PoolKey,
index: usize,
parked_at: Instant,
}
fn older_parked(
current: Option<ParkedCandidate>,
candidate: ParkedCandidate,
) -> Option<ParkedCandidate> {
match current {
Some(oldest) if oldest.parked_at <= candidate.parked_at => Some(oldest),
_ => Some(candidate),
}
}
fn oldest_parked_candidate(
candidates: impl IntoIterator<Item = ParkedCandidate>,
) -> Option<ParkedCandidate> {
candidates.into_iter().fold(None, older_parked)
}
fn oldest_parked<P: ParkedEntryAge, S: ParkedEntryAge>(
reuse: &HashMap<PoolKey, Vec<P>>,
spec_reuse: &HashMap<PoolKey, Vec<S>>,
) -> Option<ParkedCandidate> {
let plain = reuse.iter().flat_map(|(key, pool)| {
pool.iter().enumerate().map(move |(index, entry)| ParkedCandidate {
pool: ParkedPool::Plain,
key: key.clone(),
index,
parked_at: entry.parked_at(),
})
});
let spec = spec_reuse.iter().flat_map(|(key, pool)| {
pool.iter().enumerate().map(move |(index, entry)| ParkedCandidate {
pool: ParkedPool::Spec,
key: key.clone(),
index,
parked_at: entry.parked_at(),
})
});
oldest_parked_candidate(plain.chain(spec))
}
fn evict_oldest_parked<P: ParkedEntryAge, S: ParkedEntryAge>(
reuse: &mut HashMap<PoolKey, Vec<P>>,
spec_reuse: &mut HashMap<PoolKey, Vec<S>>,
metrics: &mut ReuseMetrics,
) -> Option<ParkedPool> {
let candidate = oldest_parked(reuse, spec_reuse)?;
match candidate.pool {
ParkedPool::Plain => {
let empty = {
let pool = reuse.get_mut(&candidate.key)
.expect("oldest plain parked entry vanished");
drop(pool.remove(candidate.index));
pool.is_empty()
};
if empty {
reuse.remove(&candidate.key);
}
metrics.continuation_evictions += 1;
}
ParkedPool::Spec => {
let empty = {
let pool = spec_reuse.get_mut(&candidate.key)
.expect("oldest spec parked entry vanished");
drop(pool.remove(candidate.index));
pool.is_empty()
};
if empty {
spec_reuse.remove(&candidate.key);
}
metrics.spec_evictions += 1;
}
}
Some(candidate.pool)
}
fn parked_entry_count<P, S>(
reuse: &HashMap<PoolKey, Vec<P>>,
spec_reuse: &HashMap<PoolKey, Vec<S>>,
) -> usize {
reuse.values().map(Vec::len).sum::<usize>()
+ spec_reuse.values().map(Vec::len).sum::<usize>()
}
fn trim_parked_namespace<T>(pool: Option<&mut Vec<T>>, cap: usize, evictions: &mut u64) {
if let Some(pool) = pool {
while pool.len() >= cap {
drop(pool.remove(0));
*evictions += 1;
}
}
}
fn prepare_park<P: ParkedEntryAge, S: ParkedEntryAge>(
target: ParkedPool,
key: &PoolKey,
reuse: &mut HashMap<PoolKey, Vec<P>>,
spec_reuse: &mut HashMap<PoolKey, Vec<S>>,
metrics: &mut ReuseMetrics,
per_namespace_cap: usize,
global_cap: usize,
) -> bool {
if per_namespace_cap == 0 || global_cap == 0 {
return false;
}
match target {
ParkedPool::Plain => trim_parked_namespace(
reuse.get_mut(key), per_namespace_cap, &mut metrics.continuation_evictions),
ParkedPool::Spec => trim_parked_namespace(
spec_reuse.get_mut(key), per_namespace_cap, &mut metrics.spec_evictions),
}
while parked_entry_count(reuse, spec_reuse) >= global_cap {
if evict_oldest_parked(reuse, spec_reuse, metrics).is_none() {
break;
}
}
true
}
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 {
if memra_engine::pp::pp_host_bounce_active() {
return false;
}
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() }
}
fn affinity_resume_target(
prompt: &[u32],
committed: &[u32],
checkpoint_pos: usize,
parked_cap: usize,
request_cap: usize,
identity_matches: bool,
) -> Result<usize, String> {
if checkpoint_pos == 0 || checkpoint_pos > committed.len() {
return Err(format!(
"bad checkpoint pos {checkpoint_pos} of {}",
committed.len(),
));
}
match affinity_match(prompt, &committed[..checkpoint_pos]) {
AffinityMatch::Exact { suffix_from } if suffix_from == checkpoint_pos => {}
AffinityMatch::Diverged { at } => {
return Err(format!(
"history diverged at {at} of checkpoint {checkpoint_pos}",
));
}
_ => return Err("diff did not land on the checkpoint".into()),
}
if prompt.len() == checkpoint_pos {
return Err("empty suffix".into());
}
if !identity_matches {
return Err("identity did not nominate".into());
}
Ok(parked_cap.max(request_cap))
}
fn plain_checkpoint_boundary(prompt: &[u32], is_control: &dyn Fn(u32) -> bool) -> Option<usize> {
let n = prompt.len();
if n <= REUSE_MIN_PREFIX + PLAIN_CKPT_RAW_GUARD {
return None;
}
if let Some(last_marker) = prompt.iter().rposition(|&t| is_control(t)) {
if last_marker > REUSE_MIN_PREFIX && last_marker < n {
return Some(last_marker);
}
}
let b = n - PLAIN_CKPT_RAW_GUARD;
if b > REUSE_MIN_PREFIX { Some(b) } else { None }
}
const PLAIN_CKPT_RAW_GUARD: usize = 16;
fn plain_ckpt_nominatable(prompt: &[u32], is_control: &dyn Fn(u32) -> bool) -> bool {
conversation_fingerprint(prompt, is_control, false).len() >= FP_MIN_SEGMENTS
}
const PREFIX_CACHE_MIN_TOKENS: usize = 64;
const METER_TENANT_CAP: usize = 256;
pub const SPEC_METRICS_WINDOW_S: f32 = 30.0;
const SPEC_METRICS_WINDOW_MAX_SAMPLES: usize = 16_384;
struct SpecTelemetryWindow {
samples: VecDeque<(Instant, memra_engine::spec::SpecTelemetry)>,
total: memra_engine::spec::SpecTelemetry,
window: Duration,
}
impl SpecTelemetryWindow {
fn new(window_s: f32) -> Self {
Self {
samples: VecDeque::new(),
total: Default::default(),
window: Duration::from_secs_f32(window_s),
}
}
fn push(&mut self, delta: memra_engine::spec::SpecTelemetry) {
self.push_at(Instant::now(), delta);
}
fn push_at(&mut self, now: Instant, delta: memra_engine::spec::SpecTelemetry) {
if delta.rounds == 0 {
return;
}
self.samples.push_back((now, delta));
self.total.merge(&delta);
while self.samples.len() > SPEC_METRICS_WINDOW_MAX_SAMPLES {
let (_, old) = self.samples.pop_front().unwrap();
self.total = self.total.delta_since(&old);
}
self.evict_at(now);
}
fn evict_at(&mut self, now: Instant) {
while let Some((at, _)) = self.samples.front() {
if now.saturating_duration_since(*at) <= self.window {
break;
}
let (_, old) = self.samples.pop_front().unwrap();
self.total = self.total.delta_since(&old);
}
}
fn snapshot_at(&mut self, now: Instant) -> memra_engine::spec::SpecTelemetry {
self.evict_at(now);
self.total
}
}
struct SpecMetricState {
lifetime: HashMap<String, memra_engine::spec::SpecTelemetry>,
windows: HashMap<String, SpecTelemetryWindow>,
window_s: f32,
}
impl SpecMetricState {
fn new(window_s: f32) -> Self {
Self { lifetime: HashMap::new(), windows: HashMap::new(), window_s }
}
fn record(&mut self, model: &str, delta: memra_engine::spec::SpecTelemetry) {
if delta.rounds == 0 {
return;
}
self.lifetime.entry(model.to_string()).or_default().merge(&delta);
self.windows.entry(model.to_string())
.or_insert_with(|| SpecTelemetryWindow::new(self.window_s))
.push(delta);
}
fn window_snapshots(&mut self) -> HashMap<String, memra_engine::spec::SpecTelemetry> {
let now = Instant::now();
self.windows.iter_mut().filter_map(|(model, window)| {
let snapshot = window.snapshot_at(now);
(snapshot.rounds > 0).then(|| (model.clone(), snapshot))
}).collect()
}
}
const ADSD_TENANT_WINDOW: usize = 8;
const ADSD_MODEL_WINDOW: usize = 64;
const ADSD_BASELINE_MIN_SAMPLES: usize = 16;
const ADSD_BASELINE_MIN_DRAFTED: u64 = 512;
const ADSD_TENANT_MIN_DRAFTED: u64 = 128;
const ADSD_Z_THRESHOLD: f64 = -3.0;
const ADSD_MIN_RATE_DROP: f64 = 0.20;
const ADSD_REARM_RATE_DROP: f64 = 0.10;
const ADSD_SUSTAINED_OBSERVATIONS: u8 = 3;
#[derive(Default)]
struct AcceptanceWindow {
samples: VecDeque<(u64, u64)>,
accepted: u64,
drafted: u64,
}
impl AcceptanceWindow {
fn push(&mut self, accepted: u64, drafted: u64, cap: usize) {
self.samples.push_back((accepted, drafted));
self.accepted += accepted;
self.drafted += drafted;
while self.samples.len() > cap {
let (old_accepted, old_drafted) = self.samples.pop_front().unwrap();
self.accepted = self.accepted.saturating_sub(old_accepted);
self.drafted = self.drafted.saturating_sub(old_drafted);
}
}
fn rate(&self) -> f64 {
self.accepted as f64 / self.drafted.max(1) as f64
}
}
#[derive(Default)]
struct ModelAcceptanceWindow {
samples: VecDeque<(String, u64, u64)>,
}
impl ModelAcceptanceWindow {
fn push(&mut self, tenant: &str, accepted: u64, drafted: u64) {
self.samples.push_back((tenant.to_string(), accepted, drafted));
while self.samples.len() > ADSD_MODEL_WINDOW {
self.samples.pop_front();
}
}
fn baseline_excluding(&self, tenant: &str) -> Option<(u64, u64)> {
let mut samples = 0;
let mut accepted = 0;
let mut drafted = 0;
for (sample_tenant, sample_accepted, sample_drafted) in &self.samples {
if sample_tenant == tenant {
continue;
}
samples += 1;
accepted += sample_accepted;
drafted += sample_drafted;
}
(samples >= ADSD_BASELINE_MIN_SAMPLES && drafted >= ADSD_BASELINE_MIN_DRAFTED)
.then_some((accepted, drafted))
}
fn historical_baseline(&self, tenant: &str) -> Option<(u64, u64)> {
let mut recent = ADSD_TENANT_WINDOW.saturating_sub(1);
let mut samples = 0;
let mut accepted = 0;
let mut drafted = 0;
for (sample_tenant, sample_accepted, sample_drafted) in self.samples.iter().rev() {
if sample_tenant != tenant {
continue;
}
if recent > 0 {
recent -= 1;
continue;
}
samples += 1;
accepted += sample_accepted;
drafted += sample_drafted;
}
(samples >= ADSD_BASELINE_MIN_SAMPLES && drafted >= ADSD_BASELINE_MIN_DRAFTED)
.then_some((accepted, drafted))
}
}
#[derive(Default)]
struct TenantAcceptance {
window: AcceptanceWindow,
anomalous_observations: u8,
incident_latched: bool,
latched_baseline: Option<(u64, u64)>,
}
#[derive(Debug)]
struct AdsdSuspect {
model: String,
tenant: String,
baseline_rate: f64,
tenant_rate: f64,
z_score: f64,
drafted: u64,
}
#[derive(Default)]
struct AdsdDetector {
model_windows: HashMap<String, ModelAcceptanceWindow>,
tenant_windows: HashMap<(String, String), TenantAcceptance>,
suspect_total: HashMap<String, u64>,
}
impl AdsdDetector {
fn observe(
&mut self,
model: &str,
tenant: &str,
accepted: u64,
drafted: u64,
) -> Option<AdsdSuspect> {
if drafted == 0 || accepted > drafted {
return None;
}
let mut key = (model.to_string(), tenant.to_string());
if !self.tenant_windows.contains_key(&key)
&& self.tenant_windows.len() >= METER_TENANT_CAP
{
key.1 = "(other)".into();
}
let cross_baseline = self.model_windows.get(model)
.and_then(|window| window.baseline_excluding(&key.1));
let historical_baseline = if cross_baseline.is_none() {
self.model_windows.get(model)
.and_then(|window| window.historical_baseline(&key.1))
} else {
None
};
let state = self.tenant_windows.entry(key.clone()).or_default();
state.window.push(accepted, drafted, ADSD_TENANT_WINDOW);
let baseline = cross_baseline.or_else(|| {
if state.incident_latched {
state.latched_baseline.or(historical_baseline)
} else {
historical_baseline
}
});
let mut suspect = None;
if let Some((baseline_accepted, baseline_drafted)) = baseline {
if state.window.samples.len() == ADSD_TENANT_WINDOW
&& state.window.drafted >= ADSD_TENANT_MIN_DRAFTED
{
let baseline_rate = baseline_accepted as f64 / baseline_drafted as f64;
let tenant_rate = state.window.rate();
let pooled_rate = (baseline_accepted as f64 + state.window.accepted as f64)
/ (baseline_drafted as f64 + state.window.drafted as f64);
let variance = (pooled_rate * (1.0 - pooled_rate)
* (1.0 / baseline_drafted as f64 + 1.0 / state.window.drafted as f64))
.max(f64::EPSILON);
let z_score = (tenant_rate - baseline_rate) / variance.sqrt();
let rate_drop = baseline_rate - tenant_rate;
let anomalous = rate_drop >= ADSD_MIN_RATE_DROP
&& z_score <= ADSD_Z_THRESHOLD;
if anomalous {
state.anomalous_observations =
state.anomalous_observations.saturating_add(1);
if state.anomalous_observations >= ADSD_SUSTAINED_OBSERVATIONS
&& !state.incident_latched
{
state.incident_latched = true;
state.latched_baseline = Some((baseline_accepted, baseline_drafted));
suspect = Some(AdsdSuspect {
model: model.to_string(),
tenant: key.1.clone(),
baseline_rate,
tenant_rate,
z_score,
drafted: state.window.drafted,
});
}
} else {
state.anomalous_observations = 0;
if rate_drop <= ADSD_REARM_RATE_DROP {
state.incident_latched = false;
state.latched_baseline = None;
}
}
}
}
self.model_windows.entry(model.to_string()).or_default()
.push(&key.1, accepted, drafted);
if let Some(event) = suspect.as_ref() {
*self.suspect_total.entry(event.tenant.clone()).or_default() += 1;
}
suspect
}
}
fn meter_account(ns_tokens: &mut HashMap<String, [u64; 2]>, cache_ns: &str,
n_prompt: u64, n_cached: u64) {
let mk = crate::auth::meter_key(cache_ns);
let row = if ns_tokens.contains_key(mk) || ns_tokens.len() < METER_TENANT_CAP {
ns_tokens.entry(mk.to_string()).or_default()
} else {
ns_tokens.entry("(other)".to_string()).or_default()
};
row[0] += n_prompt;
row[1] += n_cached;
}
fn meter_cached_credit(ns_tokens: &mut HashMap<String, [u64; 2]>, cache_ns: &str,
n_cached: u64) {
meter_account(ns_tokens, cache_ns, 0, n_cached);
}
fn prefix_cache_budget_bytes() -> usize {
static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*B.get_or_init(|| {
let configured = std::env::var("MEMRA_PREFIX_CACHE_MB").ok()
.and_then(|v| v.parse::<usize>().ok()).unwrap_or(256)
.saturating_mul(1024 * 1024);
if configured > 0 && memra_engine::pp::pp_host_bounce_active() {
0
} else {
configured
}
})
}
fn prefix_dedup_enabled() -> bool {
static D: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*D.get_or_init(|| std::env::var("MEMRA_PREFIX_DEDUP").as_deref() != Ok("0"))
}
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 {
if memra_engine::pp::pp_host_bounce_active() {
return false;
}
static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
let armed = *S.get_or_init(|| {
std::env::var("MEMRA_SERVE_SPEC")
.map(|v| v != "0")
.unwrap_or(true)
});
if !armed {
return false;
}
match spec_k_pin() {
Some(k) => k > 0,
None => !spec_gate_on() || spec_gate_low() > 0,
}
}
fn spec_gate_on() -> bool {
static G: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*G.get_or_init(|| std::env::var("MEMRA_SPEC_GATE").as_deref() != Ok("0"))
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct SpecGateThresholds {
low: usize,
high: usize,
raw_high: usize,
pp2_default: bool,
low_overridden: bool,
high_overridden: bool,
high_clamped: bool,
}
fn spec_gate_defaults(pp2: bool) -> (usize, usize) {
if pp2 { (0, 1) } else { (2, 4) }
}
fn resolve_spec_gate_thresholds(
pp2: bool,
low_override: Option<usize>,
high_override: Option<usize>,
) -> SpecGateThresholds {
let (default_low, default_high) = spec_gate_defaults(pp2);
let low = low_override.unwrap_or(default_low);
let raw_high = high_override.unwrap_or(default_high);
let high_clamped = raw_high <= low;
let high = if high_clamped {
low.saturating_add(1)
} else {
raw_high
};
SpecGateThresholds {
low,
high,
raw_high,
pp2_default: pp2,
low_overridden: low_override.is_some(),
high_overridden: high_override.is_some(),
high_clamped,
}
}
fn spec_gate_pp2_placement() -> bool {
let exactly_two_stages = std::env::var("MEMRA_PP_STAGES")
.ok()
.and_then(|v| v.parse::<usize>().ok())
.is_some_and(|n| n == 2);
exactly_two_stages && memra_engine::pp::pp_sharded_cross_device()
}
fn spec_gate_thresholds() -> &'static SpecGateThresholds {
static T: std::sync::OnceLock<SpecGateThresholds> = std::sync::OnceLock::new();
T.get_or_init(|| {
let low_override = std::env::var("MEMRA_SPEC_GATE_LOW")
.ok()
.and_then(|v| v.parse().ok());
let high_override = std::env::var("MEMRA_SPEC_GATE_HIGH")
.ok()
.and_then(|v| v.parse().ok());
let thresholds =
resolve_spec_gate_thresholds(spec_gate_pp2_placement(), low_override, high_override);
if thresholds.high_clamped {
eprintln!(
"[spec-gate] WARN: MEMRA_SPEC_GATE_HIGH={} <= LOW={} leaves no hysteresis \
band (mode thrash); clamped to {}",
thresholds.raw_high, thresholds.low, thresholds.high
);
}
thresholds
})
}
const SPEC_K_LONG_PROMPT_MIN: usize = 1024;
const SPEC_K_LONG_CACHE_MIN: usize = 1024;
const SPEC_K_COLD_SHORT: usize = 3;
const SPEC_K_COLD_LONG: usize = 3;
const SPEC_K_CACHED_LONG: usize = 2;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum SpecKReason {
OperatorPin,
Replay,
Placement,
Concurrency,
CachedLong,
ColdShort,
ColdLong,
}
impl SpecKReason {
fn as_str(self) -> &'static str {
match self {
Self::OperatorPin => "operator-pin",
Self::Replay => "oom-replay",
Self::Placement => "pp2-placement",
Self::Concurrency => "concurrency",
Self::CachedLong => "cached-long",
Self::ColdShort => "cold-short",
Self::ColdLong => "cold-long",
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct SpecKDecision {
k: usize,
reason: SpecKReason,
}
fn parse_spec_k_pin(raw: Option<&str>) -> Result<Option<usize>, String> {
raw.map(|value| {
value.parse::<usize>().map_err(|_| {
format!("MEMRA_SPEC_K={value:?} is not a non-negative integer")
})
}).transpose()
}
fn spec_k_pin() -> Option<usize> {
static K: std::sync::OnceLock<Option<usize>> = std::sync::OnceLock::new();
*K.get_or_init(|| {
match parse_spec_k_pin(std::env::var("MEMRA_SPEC_K").ok().as_deref()) {
Ok(k) => k,
Err(err) => {
eprintln!("[spec-k] WARN: {err}; using automatic policy");
None
}
}
})
}
fn choose_spec_k(
pin: Option<usize>,
gate_on: bool,
thresholds: SpecGateThresholds,
projected_active: usize,
prompt_tokens: usize,
cached_tokens: usize,
) -> SpecKDecision {
if let Some(k) = pin {
return SpecKDecision { k, reason: SpecKReason::OperatorPin };
}
if gate_on && projected_active > thresholds.low {
let reason = if thresholds.pp2_default
&& !thresholds.low_overridden
&& thresholds.low == 0
{
SpecKReason::Placement
} else {
SpecKReason::Concurrency
};
return SpecKDecision { k: 0, reason };
}
if prompt_tokens >= SPEC_K_LONG_PROMPT_MIN
&& cached_tokens >= SPEC_K_LONG_CACHE_MIN
{
return SpecKDecision { k: SPEC_K_CACHED_LONG, reason: SpecKReason::CachedLong };
}
if prompt_tokens < SPEC_K_LONG_PROMPT_MIN {
SpecKDecision { k: SPEC_K_COLD_SHORT, reason: SpecKReason::ColdShort }
} else {
SpecKDecision { k: SPEC_K_COLD_LONG, reason: SpecKReason::ColdLong }
}
}
fn log_spec_gate_policy() {
if let Some(k) = spec_k_pin() {
eprintln!(
"[spec-k] operator pin K={k}: automatic placement/concurrency/prompt policy \
and automatic demotion disabled"
);
return;
}
if !spec_gate_on() {
eprintln!("[spec-gate] policy disabled by MEMRA_SPEC_GATE=0: always-spec");
} else {
let thresholds = spec_gate_thresholds();
let placement = if thresholds.pp2_default {
"pp2-cross-device"
} else {
"single-or-non-pp2"
};
let source = if thresholds.low_overridden || thresholds.high_overridden {
"env-resolved"
} else {
"placement-default"
};
let admission = if thresholds.low == 0 { "off" } else { "on" };
eprintln!(
"[spec-gate] policy placement={placement} LOW={} HIGH={} source={source} \
spec-admission={admission}",
thresholds.low, thresholds.high
);
}
eprintln!(
"[spec-k] automatic table: prompt<{} -> K={}; cold-long -> K={}; \
prompt>= {} and cached>= {} -> K={}",
SPEC_K_LONG_PROMPT_MIN, SPEC_K_COLD_SHORT, SPEC_K_COLD_LONG,
SPEC_K_LONG_PROMPT_MIN, SPEC_K_LONG_CACHE_MIN, SPEC_K_CACHED_LONG
);
}
fn spec_gate_low() -> usize {
spec_gate_thresholds().low
}
fn spec_gate_high() -> usize {
spec_gate_thresholds().high
}
#[derive(Debug, PartialEq, Eq)]
pub enum DraftVerdict {
Attached,
NoDrafterExternalMtpArch,
NoDrafterQuiet,
}
pub fn draft_verdict(
has_drafter: bool,
external_mtp_arch: bool,
) -> DraftVerdict {
if has_drafter {
return DraftVerdict::Attached;
}
if external_mtp_arch {
DraftVerdict::NoDrafterExternalMtpArch
} else {
DraftVerdict::NoDrafterQuiet
}
}
pub fn draft_verdict_message(v: &DraftVerdict, name: &str, path: &str) -> Option<String> {
match v {
DraftVerdict::Attached | DraftVerdict::NoDrafterQuiet => None,
DraftVerdict::NoDrafterExternalMtpArch => Some(format!(
"[worker] WARN: {name}: step35: no MTP drafter attached — serving plain decode, \
no speculative decoding. This arch ships its MTP/NextN head in a SEPARATE GGUF, \
so the trunk's nextn_predict_layers=0 is expected and does NOT mean the model \
has no drafter. Attach with MEMRA_MODELS=\"{name}={path}+/path/to/\
Step3.7-flash-mtp-Q8_0.gguf\" (the same '+draft' convention every regime drafter \
uses; docs/DRAFT-REGIME.md)."
)),
}
}
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)
})
})
}
fn admission_reserve(spec_capable: bool, cost: usize, override_bytes: Option<usize>) -> usize {
let floor = override_bytes.unwrap_or(SPEC_SHRINK_RESERVE);
if spec_capable { floor } else { cost.min(floor) }
}
fn dual_pp_boundary_slot_bytes(wave_cap: usize, n_embd: usize) -> usize {
wave_cap
.saturating_mul(n_embd)
.saturating_mul(std::mem::size_of::<f32>())
}
fn admission_required(cost: usize, reserve: usize) -> usize {
cost.saturating_add(reserve)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct DualPpStageAdmission {
session_bytes: usize,
reserve_bytes: usize,
boundary_bytes: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct DualPpDeviceRequirement {
device: usize,
session_bytes: usize,
reserve_bytes: usize,
boundary_bytes: usize,
}
impl DualPpDeviceRequirement {
fn required(self) -> usize {
self.session_bytes
.saturating_add(self.reserve_bytes)
.saturating_add(self.boundary_bytes)
}
fn add_stage(&mut self, stage: DualPpStageAdmission) {
self.session_bytes = self.session_bytes.saturating_add(stage.session_bytes);
self.reserve_bytes = self.reserve_bytes.saturating_add(stage.reserve_bytes);
self.boundary_bytes = self.boundary_bytes.saturating_add(stage.boundary_bytes);
}
}
fn dual_pp_device_requirements(
devices: [usize; 2],
stages: [DualPpStageAdmission; 2],
) -> Vec<DualPpDeviceRequirement> {
let mut requirements: Vec<DualPpDeviceRequirement> = Vec::with_capacity(2);
for (device, stage) in devices.into_iter().zip(stages) {
if let Some(existing) = requirements
.iter_mut()
.find(|requirement| requirement.device == device)
{
existing.add_stage(stage);
} else {
requirements.push(DualPpDeviceRequirement {
device,
session_bytes: stage.session_bytes,
reserve_bytes: stage.reserve_bytes,
boundary_bytes: stage.boundary_bytes,
});
}
}
requirements
}
fn dual_pp_stage_context_bytes(
model: &HybridModel,
fence: &[usize],
ctx_cap: usize,
spec: bool,
) -> Option<[usize; 2]> {
if fence.len() != 3 || fence[0] != 0 {
return None;
}
let n_layers = model.cfg.n_layer as usize;
if fence[1] > fence[2] || fence[2] > n_layers {
return None;
}
let ring_rows = memra_engine::cache::cache_ring_row_cap(&model.cfg);
let mut stage_bytes = [0usize; 2];
for stage in 0..2 {
let lo = fence[stage];
let hi = if stage == 1 { n_layers } else { fence[stage + 1] };
let bytes_per_token =
memra_engine::cache::cache_bytes_per_token_for_layers(&model.cfg, lo, hi);
let ring_bytes_per_token =
memra_engine::cache::cache_ring_bytes_per_token_for_layers(&model.cfg, lo, hi);
stage_bytes[stage] = context_cache_bytes(
bytes_per_token,
ring_bytes_per_token,
ring_rows,
ctx_cap,
);
}
if spec {
let (plain, plain_ring, _) = model.plain_session_kv_shape();
let (spec, spec_ring, _) = model.spec_session_kv_shape();
stage_bytes[1] = stage_bytes[1].saturating_add(context_cache_bytes(
spec.saturating_sub(plain),
spec_ring.saturating_sub(plain_ring),
ring_rows,
ctx_cap,
));
}
let (total, total_ring, _) = if spec {
model.spec_session_kv_shape()
} else {
model.plain_session_kv_shape()
};
debug_assert_eq!(
stage_bytes[0].saturating_add(stage_bytes[1]),
context_cache_bytes(total, total_ring, ring_rows, ctx_cap),
"PP stage-local KV admission must partition the aggregate cache geometry",
);
Some(stage_bytes)
}
fn dual_pp_stage_admission(
context_bytes: [usize; 2],
activation_bytes: usize,
reserve_bytes: usize,
boundary_slot_bytes: usize,
) -> [DualPpStageAdmission; 2] {
[
DualPpStageAdmission {
session_bytes: context_bytes[0].saturating_add(activation_bytes),
reserve_bytes,
boundary_bytes: 0,
},
DualPpStageAdmission {
session_bytes: context_bytes[1].saturating_add(activation_bytes),
reserve_bytes,
boundary_bytes: boundary_slot_bytes.saturating_mul(2),
},
]
}
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")
}
fn alloc_with_single_reclaim_retry<T, E>(
mut alloc: impl FnMut() -> Result<T, E>,
mut reclaim: impl FnMut(&E) -> bool,
) -> Result<T, E> {
match alloc() {
Ok(value) => Ok(value),
Err(first_err) if reclaim(&first_err) => alloc(),
Err(err) => Err(err),
}
}
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,
pins: usize,
}
#[derive(Clone, Debug, PartialEq, Eq)]
struct PrefixPin {
key: PoolKey,
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,
lcp_hist: [u64; 11],
}
pub const LCP_HIST_EDGES: [usize; 11] = [0, 1, 16, 32, 64, 128, 256, 512, 1024, 2048, 4096];
#[derive(Debug, PartialEq, Eq)]
struct PrefixFanoutCandidate {
active_idx: usize,
key: PoolKey,
prompt: Vec<u32>,
}
#[derive(Debug, PartialEq, Eq)]
struct PrefixFanoutGroup {
members: Vec<usize>,
prefix_len: usize,
}
fn prefix_fanout_groups(
candidates: &[PrefixFanoutCandidate],
prefix_cap: usize,
) -> Vec<PrefixFanoutGroup> {
if prefix_cap < PREFIX_CACHE_MIN_TOKENS {
return Vec::new();
}
let mut used = vec![false; candidates.len()];
let mut out = Vec::new();
for i in 0..candidates.len() {
if used[i] || candidates[i].prompt.len() < PREFIX_CACHE_MIN_TOKENS {
continue;
}
let mut group = vec![i];
for j in i + 1..candidates.len() {
if used[j]
|| candidates[j].key != candidates[i].key
|| candidates[j].prompt.len() < PREFIX_CACHE_MIN_TOKENS
|| candidates[j].prompt[..PREFIX_CACHE_MIN_TOKENS]
!= candidates[i].prompt[..PREFIX_CACHE_MIN_TOKENS]
{
continue;
}
group.push(j);
}
if group.len() < 2 {
continue;
}
let mut lcp = group.iter()
.map(|&j| candidates[j].prompt.len())
.min()
.unwrap_or(0);
for &j in group.iter().skip(1) {
lcp = PrefixCache::lcp(
&candidates[i].prompt[..lcp],
&candidates[j].prompt[..lcp],
);
}
let prefix_len = lcp.min(prefix_cap);
if prefix_len < PREFIX_CACHE_MIN_TOKENS {
continue;
}
for &j in &group {
used[j] = true;
}
out.push(PrefixFanoutGroup {
members: group.iter().map(|&j| candidates[j].active_idx).collect(),
prefix_len,
});
}
out
}
impl PrefixCache {
fn lcp(a: &[u32], b: &[u32]) -> usize {
a.iter().zip(b.iter()).take_while(|(x, y)| x == y).count()
}
fn lcp_bucket(n: usize) -> usize {
LCP_HIST_EDGES.iter().rposition(|&e| n >= e).unwrap_or(0)
}
fn record_lcp(&mut self, n: usize) {
self.lcp_hist[Self::lcp_bucket(n)] += 1;
}
fn promote_miss_to_hit(&mut self, miss_lcp: usize, hit_len: usize) {
self.misses = self.misses.checked_sub(1).unwrap_or_else(|| {
eprintln!("[prefix-cache] WARN: fanout promoted a miss that was never recorded \
(misses underflow) — metric slip, not fatal");
0
});
let bucket = Self::lcp_bucket(miss_lcp);
self.lcp_hist[bucket] = self.lcp_hist[bucket].saturating_sub(1);
self.hits += 1;
self.hit_tokens += hit_len as u64;
self.record_lcp(hit_len);
}
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 key_index(&self, key: &PoolKey, toks: &[u32]) -> Option<usize> {
self.entries.get(key)?.iter().position(|e| e.toks[..] == *toks)
}
fn id_index(&self, pin: &PrefixPin) -> Option<usize> {
self.entries.get(&pin.key)?.iter().position(|e| e.id == pin.id)
}
fn lru_key(e: &PrefixEntry) -> (Instant, u64) {
(e.last_use, e.id)
}
#[cfg_attr(not(test), allow(dead_code))]
fn touch(&mut self, key: &PoolKey, i: usize) {
let Some((old_lru, pinned)) = self.entries.get(key).and_then(|p| p.get(i))
.map(|e| (Self::lru_key(e), e.pins > 0)) else { return };
if !pinned {
self.lru.remove(&old_lru);
}
let e = &mut self.entries.get_mut(key).unwrap()[i];
e.last_use = Instant::now();
if e.pins == 0 {
self.lru.insert(Self::lru_key(e), (key.clone(), i));
}
}
fn pin_n(&mut self, key: &PoolKey, i: usize, n: usize) -> Option<PrefixPin> {
if n == 0 {
return None;
}
let (old_lru, id, was_unpinned) = {
let e = self.entries.get(key)?.get(i)?;
(Self::lru_key(e), e.id, e.pins == 0)
};
if was_unpinned {
self.lru.remove(&old_lru);
}
let e = &mut self.entries.get_mut(key)?[i];
e.pins = e.pins.checked_add(n).expect("prefix pin refcount overflow");
e.last_use = Instant::now();
Some(PrefixPin { key: key.clone(), id })
}
fn pin(&mut self, key: &PoolKey, i: usize) -> Option<PrefixPin> {
self.pin_n(key, i, 1)
}
fn unpin(&mut self, pin: &PrefixPin) -> bool {
let Some(i) = self.id_index(pin) else { return false };
let e = &mut self.entries.get_mut(&pin.key).unwrap()[i];
if e.pins == 0 {
return false;
}
e.pins -= 1;
if e.pins == 0 {
e.last_use = Instant::now();
self.lru.insert(Self::lru_key(e), (pin.key.clone(), i));
}
true
}
fn pinned_bytes(&self) -> usize {
self.entries.values().flatten().filter(|e| e.pins > 0).map(|e| e.bytes).sum()
}
fn remove_at(&mut self, key: &PoolKey, i: usize) -> Option<PrefixEntry> {
let pool = self.entries.get_mut(key)?;
if i >= pool.len() || pool[i].pins > 0 {
return None;
}
let dead = pool.swap_remove(i);
self.lru.remove(&Self::lru_key(&dead));
if let Some(moved) = pool.get(i) {
if moved.pins == 0 {
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_pinned(
&mut self,
key: &PoolKey,
e: PrefixEntry,
why: &str,
pins: usize,
) -> Option<PrefixPin> {
let id = self.insert_with_budget_pins(
key, e, why, prefix_cache_budget_bytes(), pins)?;
Some(PrefixPin { key: key.clone(), id })
}
fn insert_with_budget(&mut self, key: &PoolKey, e: PrefixEntry, why: &str, budget: usize) {
let _ = self.insert_with_budget_pins(key, e, why, budget, 0);
}
fn insert_with_budget_pins(
&mut self,
key: &PoolKey,
mut e: PrefixEntry,
why: &str,
budget: usize,
initial_pins: usize,
) -> Option<u64> {
if let Some(i) = self.key_index(key, &e.toks) {
return if initial_pins > 0 {
self.pin_n(key, i, initial_pins).map(|pin| pin.id)
} else {
None
};
}
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 None;
}
if initial_pins > 0 && e.bytes > budget.saturating_sub(self.pinned_bytes()) {
eprintln!("[prefix-cache] skip pinned {why} insert: entry {:.1}MB cannot fit \
beside {:.1}MB already pinned (budget {:.0}MB)",
e.bytes as f64 / 1e6, self.pinned_bytes() as f64 / 1e6,
budget as f64 / 1e6);
return None;
}
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;
e.pins = initial_pins;
let inserted_id = e.id;
let lk = Self::lru_key(&e);
let idx = {
let pool = self.entries.entry(key.clone()).or_default();
pool.push(e);
pool.len() - 1
};
if initial_pins == 0 {
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));
}
self.entries.get(key)
.and_then(|pool| pool.iter().find(|entry| entry.id == inserted_id))
.map(|_| inserted_id)
}
fn evict_all(&mut self) -> usize {
let mut n = 0usize;
while let Some((key, i)) = self.lru.values().next().cloned() {
let Some(dead) = self.remove_at(&key, i) else { break };
self.total_bytes = self.total_bytes.saturating_sub(dead.bytes);
n += 1;
}
self.evictions += n as u64;
n
}
}
fn retire_prefix_pin(px: &mut PrefixCache, prefix_pin: &mut Option<PrefixPin>) {
if let Some(pin) = prefix_pin.take() {
if !px.unpin(&pin) {
eprintln!("[prefix-cache] warning: retired session held a missing prefix pin");
}
}
}
fn prefix_snapshot(
engine: &Engine,
cache: &Cache,
toks: &[u32],
last_logits: &[f32],
) -> Result<PrefixEntry, Box<dyn std::error::Error>> {
if cache.has_swa_ring() {
return Err("SWA ring sessions do not support flat-history prefix snapshots".into());
}
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, pins: 0,
})
}
fn prefix_restore(
engine: &Engine,
cache: &mut Cache,
e: &PrefixEntry,
) -> Result<(), Box<dyn std::error::Error>> {
if cache.has_swa_ring() {
return Err("SWA ring sessions do not support flat-history prefix restores".into());
}
if cache.kv.len() != e.kv.len() {
return Err(format!("prefix entry layer count {} != cache {}", e.kv.len(), cache.kv.len()).into());
}
let cache_cap = cache.max_ctx;
for il in 0..cache.kv.len() {
match (cache.kv[il].as_mut(), &e.kv[il]) {
(Some(dst), Some(src)) => {
if src.len > cache_cap {
return Err(format!(
"prefix entry layer {il} len {} exceeds cache capacity {cache_cap}",
src.len).into());
}
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_plain_checkpoint(engine: &Engine, s: &mut Session) {
let Some(at) = s.ckpt_at else { return };
if s.fed.len() != at {
return;
}
s.ckpt_at = None; let Some(cache) = s.cache.as_ref() else { return };
debug_assert_eq!(cache.pos, at, "plain checkpoint must sit at the fed boundary");
match cache.snapshot(engine) {
Ok(snap) => {
s.ckpt_snap = Some(PlainCheckpoint {
snap,
pos: at,
last_logits: s.last_logits.clone(),
});
}
Err(err) => {
if std::env::var("MEMRA_DEBUG_AFFINITY").is_ok() {
eprintln!("[plain-affinity] checkpoint capture skipped ({err}); \
next turn re-primes in full");
}
}
}
}
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,
reasoning_effort: Option<String>,
params: GenParams,
sampler_cfg: SamplerConfig,
grammar: Option<crate::constrained::GrammarSpec>,
max_prompt_tokens: Option<usize>,
}
struct Session {
model: String,
spec_k: usize,
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>,
tokens_emitted: usize,
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>,
ckpt_at: Option<usize>,
ckpt_snap: Option<PlainCheckpoint>,
prefix_miss_lcp: Option<usize>,
seed_prefix: bool,
prefix_pin: Option<PrefixPin>,
tx: tokio::sync::mpsc::UnboundedSender<Event>,
ttft: Option<Arc<crate::ttft::Trace>>,
t0: Instant,
}
impl Session {
fn pool_key(&self) -> PoolKey {
(self.model.clone(), self.cache_ns.clone())
}
}
fn worker_device(pp_devices: Option<&str>) -> Result<usize, String> {
let Some(devices) = pp_devices.filter(|v| !v.trim().is_empty()) else {
return Ok(0);
};
let mut last = 0usize;
for part in devices.split(',') {
let part = part.trim();
last = part.parse::<usize>().map_err(|_| {
format!(
"MEMRA_PP_DEVICES={devices} has invalid device {part:?} \
(want <d0>,..,<dN-1> e.g. 0,1)"
)
})?;
}
Ok(last)
}
fn optipipe_controller_threshold(raw: Option<&str>) -> Result<Option<f32>, String> {
let Some(raw) = raw else { return Ok(None); };
let threshold = raw.parse::<f32>().map_err(|err| {
format!("MEMRA_OPTI_CONTROLLER_Q={raw:?} is not a float: {err}")
})?;
if !threshold.is_finite() || !(0.0..=1.0).contains(&threshold) {
return Err(format!(
"MEMRA_OPTI_CONTROLLER_Q={raw:?} is outside the inclusive [0,1] range"
));
}
Ok(Some(threshold))
}
pub fn run(
models: Vec<(String, String, Option<String>)>,
rx: &Receiver<Cmd>,
ready_tx: Sender<Result<(Vec<String>, HashMap<String, ModelCaps>), String>>,
metrics: SharedMetrics,
health: crate::health::SharedHealth,
) {
let pp_devices = std::env::var("MEMRA_PP_DEVICES").ok();
let device = match worker_device(pp_devices.as_deref()) {
Ok(device) => device,
Err(err) => { let _ = ready_tx.send(Err(err)); return; }
};
let engine = match Engine::new(device) {
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 (device={device}, MEMRA_FAST={})", fast);
match optipipe_controller_threshold(
std::env::var("MEMRA_OPTI_CONTROLLER_Q").ok().as_deref(),
) {
Ok(Some(threshold)) => {
memra_engine::spec::set_optipipe_controller_threshold(threshold);
eprintln!(
"[worker] OPTIPIPE increment-2 diagnostic controller armed q_threshold={threshold:.3}"
);
}
Ok(None) => {}
Err(err) => { let _ = ready_tx.send(Err(err)); return; }
}
log_spec_gate_policy();
let (constraint_result_tx, constraint_result_rx) =
std::sync::mpsc::channel::<crate::constrained::ConstraintCompileResult>();
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} (drafter path {dpath:?} was requested via the \
MEMRA_MODELS '+draft' attach — refusing to start rather than \
silently serving plain decode)")));
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} (drafter path {dpath:?} was requested via the \
MEMRA_MODELS '+draft' attach — refusing to start rather than \
silently serving plain decode)")));
return;
}
}
}
model
};
let verdict = draft_verdict(model.mtp.is_some(), model.cfg.arch.is_step35());
if let Some(msg) = draft_verdict_message(&verdict, name, path) {
eprintln!("{msg}");
}
let eos_id = tok.eos_id();
eprintln!("[worker] loaded {name:?}: {} layers, eos={eos_id}", model.cfg.n_layer);
let tok = Arc::new(tok);
let constraints = match crate::constrained::ConstraintCompiler::spawn(
name,
tok.clone(),
constraint_result_tx.clone(),
&metrics,
) {
Ok(compiler) => compiler,
Err(err) => { let _ = ready_tx.send(Err(err)); return; }
};
loaded.insert(name.clone(), LoadedModel {
model, tok, eos_id, from_dir, constraints,
});
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 }
}),
effort_levels: t.is_some_and(|t| t.contains("reasoning_effort is defined")),
gemma_think: t.is_some_and(|t| t.contains("<|channel>")),
chat_temperature_default: lm.model.cfg.arch.is_step35().then_some(0.5),
chat_top_p_default: lm.model.cfg.arch.is_step35().then_some(0.9),
};
eprintln!("[worker] {n}: template caps tools={} think={} think_switch={} chat_ok={} \
effort_levels={} gemma_think={} ctx={} tok={:?} instruct={:?} \
chat_defaults={:?}/{:?}",
caps.tools_branch, caps.qwen_think, caps.think_switch, caps.chat_ok,
caps.effort_levels, caps.gemma_think, caps.context_length, caps.tokenizer,
caps.instruct_type, caps.chat_temperature_default, caps.chat_top_p_default);
(n.clone(), caps)
}).collect();
let _ = ready_tx.send(Ok((order.clone(), caps)));
health.mark_ready();
let chunk_policies: HashMap<String, DecodeChunkPolicy> = loaded
.iter()
.map(|(n, lm)| (n.clone(), decode_chunk_policy(lm)))
.collect();
for (n, policy) in &chunk_policies {
eprintln!(
"[worker] {n}: decode wave cap {}; scheduler tick cap {}{}",
policy.wave_cap,
policy.tick_cap(),
if policy.dual { " (dual PP, default-off arm)" }
else if policy.wave_cap > 8 { " (exact-16 tier)" }
else { "" },
);
}
let eager_only: std::collections::HashSet<String> = loaded.iter()
.filter(|(_, lm)| eager_only_model(lm))
.map(|(n, _)| n.clone()).collect();
for n in &eager_only {
eprintln!("[worker] {n}: EAGER-ONLY serving (gemma4 class — no batched decode arm): \
per-session eager decode, monolithic prefill, no graph promotion, \
no prime batching");
}
let mut active: Vec<Session> = Vec::new();
let mut queue: std::collections::VecDeque<Box<Request>> = std::collections::VecDeque::new();
let mut pending_constraints: HashMap<u64, PendingConstraintCompile> = HashMap::new();
let mut next_constraint_id = 0u64;
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 reuse_metrics = ReuseMetrics::default();
let mut px = PrefixCache::default();
if memra_engine::pp::pp_host_bounce_active() {
eprintln!(
"[pp] MEMRA_PP_HOST_BOUNCE=1 safety doors: speculative PP, cross-device prefix \
snapshots, and plain-affinity checkpoints disabled (they retain peer reads/copies)"
);
}
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 admission_costs: HashMap<String, AdmissionCostModel> = loaded
.iter()
.map(|(name, lm)| (name.clone(), AdmissionCostModel::new(&lm.model)))
.collect();
for (name, cost) in &admission_costs {
if cost.ring_rows > 0 {
eprintln!(
"[admission] {name:?}: plain {} B/token ({} capped at {} rows), spec {} \
B/token ({} capped); fixed residual learns from effective-free deltas",
cost.plain_bytes_per_token,
cost.plain_ring_bytes_per_token,
cost.ring_rows,
cost.spec_bytes_per_token,
cost.spec_ring_bytes_per_token,
);
} else {
eprintln!(
"[admission] {name:?}: plain {} B/token, spec {} B/token; fixed residual learns \
from effective-free deltas",
cost.plain_bytes_per_token,
cost.spec_bytes_per_token,
);
}
}
let policy = crate::lanes::LanePolicy::from_env();
let prefill_tick_explicit = std::env::var_os("MEMRA_PREFILL_TICK").is_some();
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 n_session_defers = 0u64;
let mut n_vram_defers = 0u64;
let mut n_step_oom_parks = 0u64;
let mut ns_tokens: HashMap<String, [u64; 2]> = HashMap::new();
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_metrics = SpecMetricState::new(SPEC_METRICS_WINDOW_S);
let mut spec_telem_dirty = false;
let mut adsd_detector = AdsdDetector::default();
let mut last_interactive_decode = Instant::now();
let mut tick_n: u64 = 0;
let mut n_demoted = 0u64;
loop {
if let Err(err) = memra_engine::pp::service_runtime_peer_probe(&engine) {
panic!("runtime peer-probe safety failure: {err}");
}
resolve_constraint_compiles(
&constraint_result_rx,
&mut pending_constraints,
&mut queue,
);
expire_constraint_compiles(&mut pending_constraints, Instant::now());
if active.is_empty() && queue.is_empty() {
if pending_constraints.is_empty() {
health.set_phase(crate::health::PHASE_IDLE);
match rx.recv() {
Ok(cmd) => {
health.set_phase(crate::health::PHASE_BUSY);
handle_cmd(cmd, &loaded, &order, &mut queue);
}
Err(_) => break, }
} else {
health.beat_busy();
let wait = constraint_poll_wait(&pending_constraints, Instant::now());
match rx.recv_timeout(wait) {
Ok(cmd) => handle_cmd(cmd, &loaded, &order, &mut queue),
Err(std::sync::mpsc::RecvTimeoutError::Timeout) => {}
Err(std::sync::mpsc::RecvTimeoutError::Disconnected) => return,
}
}
}
health.beat_busy();
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; }
}
}
}
resolve_constraint_compiles(
&constraint_result_rx,
&mut pending_constraints,
&mut queue,
);
expire_constraint_compiles(&mut pending_constraints, Instant::now());
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(mut req) = queue.pop_front() {
if req.tx.is_closed() {
eprintln!("[abort] client disconnected while queued (model {:?}); dropped",
req.model);
continue;
}
if req.grammar.is_some() && req.prepared_constraint.is_none() {
let spec = req.grammar.take().expect("grammar checked above");
let deadline = Instant::now() +
crate::constrained::CONSTRAINT_COMPILE_TIMEOUT;
let id = next_constraint_id;
next_constraint_id = next_constraint_id.wrapping_add(1);
match loaded[&req.model].constraints.try_submit(id, spec, deadline) {
Ok(()) => {
pending_constraints.insert(id, PendingConstraintCompile {
request: req,
deadline,
});
}
Err(crate::constrained::ConstraintSubmitError::Busy) => {
fail_request(req, EngineError::overloaded(
"response_format compiler is busy; retry",
));
}
Err(crate::constrained::ConstraintSubmitError::Closed) => {
fail_request(req, EngineError::engine(
"response_format compiler is unavailable",
));
}
Err(crate::constrained::ConstraintSubmitError::AbandonedWorkerLimit) => {
fail_request(req, constraint_worker_limit_error());
}
}
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 {
n_session_defers += 1;
requeue.push_back(req); } else {
lane_shed[lane.idx()] += 1;
let _ = req.tx.send(Event::Error(EngineError::rate_limit(format!(
"lane {} is 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(EngineError::rate_limit(format!(
"lane {} shed: interactive p99 over budget, retry", lane.as_str()))));
continue;
}
let shape = match prepare_request(&loaded, &mut req) {
Ok(shape) => shape,
Err(err) => {
let _ = req.tx.send(Event::Error(err));
continue;
}
};
let model_key = req.model.clone();
let prompt_len = req.prepared_prompt.as_ref().unwrap().len();
let estimate_spec = admission_request_may_spec(
&loaded[&model_key],
&req,
active.len() + 1,
prompt_len,
);
let admission_cap = shape.admission_cap();
let decode_policy = chunk_policies.get(&model_key)
.expect("loaded model missing decode chunk policy");
let (cost, bytes_per_token, activation_bytes, log_estimate) = {
let model = admission_costs.get_mut(&model_key)
.expect("loaded model missing admission cost model");
let cost = model.estimate(admission_cap, estimate_spec);
let key = (admission_cap, estimate_spec, cost);
let log = model.last_logged != Some(key);
if log {
model.last_logged = Some(key);
}
(cost, model.bytes_per_token(estimate_spec), model.activation_bytes, log)
};
if log_estimate {
eprintln!(
"[admission] request cost: model={model_key:?} ctx={} path={} = {} \
B/token x ctx + {:.0}MB fixed = {:.0}MB",
admission_cap,
if estimate_spec { "spec" } else { "plain" },
bytes_per_token,
activation_bytes as f64 / 1e6,
cost as f64 / 1e6,
);
}
if !active.is_empty() {
let reserve = admission_reserve(
serve_spec_enabled()
&& loaded.get(&req.model).is_some_and(|lm| lm.model.mtp.is_some()),
cost,
admit_reserve_override(),
);
let dual_stages = if decode_policy.dual {
let model = &loaded[&model_key].model;
let n_trunk = (model.cfg.n_layer - model.cfg.nextn_predict_layers) as usize;
let fence = memra_engine::pp::pp_cuts(n_trunk)
.expect("dual decode policy missing PP stage fence");
let context_bytes = dual_pp_stage_context_bytes(
model,
&fence,
admission_cap,
estimate_spec,
)
.expect("dual decode policy requires exactly two PP stages");
let boundary_slot_bytes = dual_pp_boundary_slot_bytes(
decode_policy.wave_cap,
model.cfg.n_embd as usize,
);
let stages = dual_pp_stage_admission(
context_bytes,
activation_bytes,
reserve,
boundary_slot_bytes,
);
if log_estimate {
eprintln!(
"[admission] dual PP per-stage plan: stage0 session {:.0}MB + \
reserve {:.0}MB; stage1 session {:.0}MB + reserve {:.0}MB + \
two boundary slots {:.3}MB",
stages[0].session_bytes as f64 / 1e6,
stages[0].reserve_bytes as f64 / 1e6,
stages[1].session_bytes as f64 / 1e6,
stages[1].reserve_bytes as f64 / 1e6,
stages[1].boundary_bytes as f64 / 1e6,
);
}
Some(stages)
} else {
None
};
let required = admission_required(cost, reserve);
if let Some(mut headroom) = admission_headroom(&engine, dual_stages) {
if !headroom.sufficient(required) {
let free_before_reclaim = headroom.limiting_free_bytes();
let mut evicted_plain = 0usize;
let mut evicted_spec = 0usize;
while !headroom.sufficient(required) {
match evict_oldest_parked(
&mut reuse,
&mut spec_reuse,
&mut reuse_metrics,
) {
Some(ParkedPool::Plain) => evicted_plain += 1,
Some(ParkedPool::Spec) => evicted_spec += 1,
None => break,
}
let Some(next_headroom) = admission_headroom(&engine, dual_stages)
else { break };
headroom = next_headroom;
}
if evicted_plain + evicted_spec > 0 {
eprintln!(
"[admit-oom] reclaim-on-defer: evicted {} plain + {} spec \
parked sessions (global LRU); effective free {:.0}MB -> {:.0}MB",
evicted_plain,
evicted_spec,
free_before_reclaim as f64 / 1e6,
headroom.limiting_free_bytes() as f64 / 1e6,
);
}
}
if !headroom.sufficient(required) {
if vram_defers == 0 {
let parked: usize = spec_reuse.values().map(|v| v.len()).sum();
match &headroom {
AdmissionHeadroom::Dual(devices) => {
let budgets = devices
.iter()
.map(|device| format!(
"dev{} free {:.0}MB (pool-cached {:.0}MB) vs \
required {:.0}MB = session {:.0}MB + reserve \
{:.0}MB + boundary {:.3}MB [pool res {:.0}MB \
used {:.0}MB]",
device.requirement.device,
device.free_bytes as f64 / 1e6,
device.pool_cached_bytes as f64 / 1e6,
device.requirement.required() as f64 / 1e6,
device.requirement.session_bytes as f64 / 1e6,
device.requirement.reserve_bytes as f64 / 1e6,
device.requirement.boundary_bytes as f64 / 1e6,
device.pool_reserved_bytes as f64 / 1e6,
device.pool_used_bytes as f64 / 1e6,
))
.collect::<Vec<_>>()
.join("; ");
eprintln!(
"[admit-oom] VRAM defer: {} active, dual PP device \
budgets [{budgets}] — queueing (FIFO) [parked spec \
sessions {}; plain reuse {}; queue {}]",
active.len(),
parked,
reuse.values().map(|v| v.len()).sum::<usize>(),
queue.len() + requeue.len(),
);
}
AdmissionHeadroom::Primary { .. } => {
let (res, used) = engine.pool_reserved_used();
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(),
headroom.limiting_free_bytes() as f64 / 1e6,
headroom.primary_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;
n_vram_defers += 1;
requeue.push_back(req); continue;
}
}
}
let free_before = effective_free_bytes(&engine).map(|(free, _)| free);
match admit(&engine, &loaded, &mut reuse, &mut spec_reuse, &mut spec_sizing,
&mut reuse_metrics, &mut px, active.len(), *req, shape) {
Ok(s) => {
let actual_spec = s.spec.is_some();
let actual_ctx = s.spec.as_ref().map(|spec| spec.cache_max_ctx())
.or_else(|| s.cache.as_ref().map(|cache| cache.max_ctx))
.unwrap_or(shape.ctx_cap);
if let (Some(before), Some((after, _))) =
(free_before, effective_free_bytes(&engine))
{
let observed = before.saturating_sub(after);
let model = admission_costs.get_mut(&model_key)
.expect("loaded model missing admission cost model");
if let Some(residual) = model.observe(observed, actual_ctx, actual_spec) {
model.last_logged = None;
eprintln!(
"[admission] fixed residual high-water: model={model_key:?} \
{:.0}MB (observed {:.0}MB at ctx {} path={})",
residual as f64 / 1e6,
observed as f64 / 1e6,
actual_ctx,
if actual_spec { "spec" } else { "plain" },
);
}
}
n_admitted += 1;
lane_admitted[lane.idx()] += 1;
n_prompt_in += s.n_prompt as u64;
n_cached_in += s.n_cached as u64;
let _ = s.tx.send(Event::PromptUsage {
n_prompt: s.n_prompt,
n_cached: s.n_cached,
});
meter_account(&mut ns_tokens, &s.cache_ns,
s.n_prompt as u64, s.n_cached as u64);
active.push(s);
}
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; }
let generated_before = active[i].generated.len();
let lane = active[i].lane;
let step_started = Instant::now();
let step_result = step_session(&engine, &loaded, &mut active[i], &mut spec_metrics);
record_output_progress(
generated_before,
active[i].generated.len(),
lane,
step_started.elapsed().as_secs_f32() * 1000.0,
&mut n_tokens_out,
&mut lane_tokens,
&mut step_stats,
&mut last_interactive_decode,
);
match step_result {
Ok(true) => {}
Ok(false) => finished.push(i),
Err(err) => {
let _ = active[i].tx.send(Event::Error(EngineError::engine(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 generated_before = s.generated.len();
let lane = s.lane;
let (cont, _) = advance_token_emit(&loaded, s, pend);
let emitted = record_output_tokens(
generated_before,
s.generated.len(),
lane,
&mut n_tokens_out,
&mut lane_tokens,
);
if emitted > 0 && lane == crate::lanes::Lane::Interactive {
last_interactive_decode = Instant::now();
}
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(EngineError::engine(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()
&& !eager_only.contains(&s.model)
&& loaded[&s.model].model.cfg.step35.is_none()
&& 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(EngineError::engine(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(EngineError::engine(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 generated_before = s.generated.len();
let lane = s.lane;
let (cont, _) = advance_token_emit(&loaded, s, pend);
let emitted = record_output_tokens(
generated_before,
s.generated.len(),
lane,
&mut n_tokens_out,
&mut lane_tokens,
);
if emitted > 0 && lane == crate::lanes::Lane::Interactive {
last_interactive_decode = Instant::now();
}
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(EngineError::engine(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(
EngineError::engine(format!("graph step failed: {err}"))));
finished.push(0);
}
}
step_stats.record(t_g.elapsed().as_secs_f32() * 1000.0);
}
}
}
}
let demote_at: Option<usize> = {
static D: std::sync::OnceLock<Option<usize>> = std::sync::OnceLock::new();
*D.get_or_init(|| std::env::var("MEMRA_SPEC_DEMOTE_AT").ok()
.and_then(|v| v.parse().ok()))
};
let automatic_demote = spec_k_pin().is_none() && spec_gate_on();
if automatic_demote || demote_at.is_some() {
let n_live = active.len() - finished.len();
let forced = demote_at.is_some_and(|n| {
active.iter().enumerate().any(|(i, s)| {
!finished.contains(&i) && s.spec.is_some() && s.generated.len() >= n
})
});
if n_live >= spec_gate_high() || forced {
for i in 0..active.len() {
if finished.contains(&i) { continue; }
let s = &mut active[i];
if s.spec.is_none() { continue; }
if !s.sampler.is_greedy() || s.constraint.is_some() { continue; }
if let Some(n) = demote_at {
if s.generated.len() < n { continue; }
} else if n_live < spec_gate_high() {
continue;
}
let sess = s.spec.as_ref().unwrap();
if sess.committed_len() == 0
|| (!sess.demote_ready() && !sess.has_pending()) { continue; }
let mut sess = s.spec.take().unwrap();
let lm = &loaded[&s.model];
if sess.has_pending() {
if let Err(err) = lm.model.spec_flush_pending(&engine, &mut sess) {
eprintln!("[spec-gate] demote flush FAILED (model {}): {err}",
s.model);
let _ = s.tx.send(Event::Error(EngineError::engine(format!(
"spec demote flush failed: {err}"))));
finished.push(i);
continue;
}
}
if !sess.demote_ready() {
eprintln!("[spec-gate] demote SKIPPED: session not in handoff shape \
after flush (model {}); staying on spec", s.model);
s.spec = Some(sess);
continue;
}
let committed = sess.committed_len();
let Some((cache, next)) = sess.into_demoted() else { continue };
debug_assert_eq!(cache.pos, s.fed.len(),
"demote handoff: cache rows != fed tokens");
s.cache = Some(cache);
s.device_next = Some(next);
s.spec_k = 0;
s.prefill_done = true;
s.last_logits.clear();
n_demoted += 1;
let why = match demote_at {
Some(n) => format!("FORCED at DEMOTE_AT={n} (test door)"),
None => format!("{n_live} active >= HIGH={}", spec_gate_high()),
};
eprintln!("[spec-gate] demoted session to batched decode: {why} \
(model {}, committed {committed}, generated {})",
s.model, s.generated.len());
}
}
}
let tick_trace = tick_trace_enabled();
let t_spec = tick_trace.then(Instant::now);
let mut spec_calls = 0usize;
let mut spec_prev_end_ms = 0.0f32;
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());
}
let mut spec_order = spec_order.into_iter().peekable();
while let Some(i) = spec_order.next() {
if finished.contains(&i) { continue; }
let pair = spec_order.peek().copied().filter(|&j| {
!finished.contains(&j)
&& spec_pipe_pairable(&engine, &loaded, &active[i], &active[j])
});
if let Some(j) = pair {
let _ = spec_order.next();
let generated_i = active[i].generated.len();
let generated_j = active[j].generated.len();
let lane_i = active[i].lane;
let lane_j = active[j].lane;
let rounds_i = active[i].spec_rounds;
let rounds_j = active[j].spec_rounds;
let pair_started = Instant::now();
let step_result = if i < j {
let (left, right) = active.split_at_mut(j);
step_spec_pair(
&engine, &loaded, &mut left[i], &mut right[0], &mut spec_metrics,
)
} else {
let (left, right) = active.split_at_mut(i);
step_spec_pair(
&engine, &loaded, &mut right[0], &mut left[j], &mut spec_metrics,
)
};
let pair_ms = pair_started.elapsed().as_secs_f32() * 1000.0;
record_output_progress(
generated_i,
active[i].generated.len(),
lane_i,
pair_ms,
&mut n_tokens_out,
&mut lane_tokens,
&mut step_stats,
&mut last_interactive_decode,
);
record_output_progress(
generated_j,
active[j].generated.len(),
lane_j,
pair_ms,
&mut n_tokens_out,
&mut lane_tokens,
&mut step_stats,
&mut last_interactive_decode,
);
if tick_trace {
spec_calls += 2;
let end_ms = t_spec.unwrap().elapsed().as_secs_f32() * 1000.0;
eprintln!(
"[tick-spec-pipe] seq={}-{} slots={i},{j} gap_ms={:.3} wall_ms={pair_ms:.3} \
generated={generated_i}->{},{}->{} rounds={},{}",
spec_calls - 1,
spec_calls,
end_ms - pair_ms - spec_prev_end_ms,
active[i].generated.len(),
generated_j,
active[j].generated.len(),
active[i].spec_rounds.saturating_sub(rounds_i),
active[j].spec_rounds.saturating_sub(rounds_j),
);
spec_prev_end_ms = end_ms;
}
match step_result {
Ok((keep_i, keep_j)) => {
if !keep_i { finished.push(i); }
if !keep_j { finished.push(j); }
}
Err(err) => {
let message = format!("step error: {err}");
let _ = active[i].tx.send(Event::Error(EngineError::engine(message.clone())));
let _ = active[j].tx.send(Event::Error(EngineError::engine(message)));
finished.push(i);
finished.push(j);
}
}
continue;
}
let generated_before = active[i].generated.len();
let lane = active[i].lane;
let step_started = Instant::now();
let trace_before = t_spec.map(|phase_start| {
(
phase_start.elapsed().as_secs_f32() * 1000.0,
active[i].generated.len(),
active[i].spec_rounds,
active[i].spec_drafted,
active[i].spec_accepted,
active[i].spec_k,
active[i].trace_id.clone().unwrap_or_else(|| "-".into()),
Instant::now(),
)
});
let step_result = step_session(&engine, &loaded, &mut active[i], &mut spec_metrics);
let step_elapsed_ms = step_started.elapsed().as_secs_f32() * 1000.0;
if let Some((start_ms, generated0, rounds0, drafted0, accepted0, k, trace_id,
call_start)) = trace_before
{
let end_ms = t_spec.unwrap().elapsed().as_secs_f32() * 1000.0;
spec_calls += 1;
eprintln!(
"[tick-spec] seq={} slot={i} trace={} start_ms={start_ms:.3} \
gap_ms={:.3} wall_ms={:.3} generated={generated0}->{} rounds={} \
drafted={} accepted={} k={k}",
spec_calls,
trace_id,
start_ms - spec_prev_end_ms,
call_start.elapsed().as_secs_f32() * 1000.0,
active[i].generated.len(),
active[i].spec_rounds.saturating_sub(rounds0),
active[i].spec_drafted.saturating_sub(drafted0),
active[i].spec_accepted.saturating_sub(accepted0),
);
spec_prev_end_ms = end_ms;
}
record_output_progress(
generated_before,
active[i].generated.len(),
lane,
step_elapsed_ms,
&mut n_tokens_out,
&mut lane_tokens,
&mut step_stats,
&mut last_interactive_decode,
);
match step_result {
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) => {
n_step_oom_parks += 1;
requeue_oom.push_back(req);
finished.push(i);
}
None => {
let _ = s.tx.send(Event::Error(EngineError::engine(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(EngineError::engine(format!("step error: {err}"))));
finished.push(i);
}
}
}
let spec_ms = t_spec
.map(|started| started.elapsed().as_secs_f32() * 1000.0)
.unwrap_or(0.0);
let t_prefill = tick_trace.then(Instant::now);
let mut prefill_single_calls = 0usize;
let mut prefill_single_tokens = 0usize;
let mut prefill_batch_calls = 0usize;
let mut prefill_batch_tokens = 0usize;
let budgets = policy.prefill_budget;
let pb_hold_ms: u64 = std::env::var("MEMRA_PRIME_BATCH_HOLD_MS").ok()
.and_then(|v| v.parse().ok()).unwrap_or(4);
let dedup_advanced = dedup_interactive_prefixes(
&engine,
&loaded,
&eager_only,
&mut px,
&mut active,
&mut finished,
budgets[0],
&mut n_cached_in,
&mut ns_tokens,
);
let dedup_waiting: std::collections::HashSet<usize> =
if prefix_dedup_enabled() && !confidence_trace_enabled() && pb_hold_ms > 0 {
active.iter().enumerate()
.filter(|(i, s)| {
!finished.contains(i)
&& prefix_fanout_eligible(s, &eager_only)
&& s.t0.elapsed().as_millis() < pb_hold_ms as u128
})
.map(|(i, _)| i)
.collect()
} else {
Default::default()
};
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; }
if dedup_waiting.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)
&& !eager_only.contains(&s.model)
&& s.snapshot_at.is_none()
&& s.ckpt_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 mut held = false;
if cand.len() == 1 && pb_hold_ms > 0 {
let s = &active[cand[0]];
if s.t0.elapsed().as_millis() < pb_hold_ms as u128 {
held = true;
}
}
let mut fired = false;
if cand.len() >= 2 {
for &i in &cand {
if let Some(trace) = active[i].ttft.as_ref() {
trace.mark_prime_start();
}
}
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();
prefill_batch_calls += 1;
prefill_batch_tokens += toks;
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;
if let Some(trace) = s.ttft.as_ref() {
trace.mark_prime_end();
}
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);
};
let sole_unfinished = queue.is_empty()
&& requeue_oom.is_empty()
&& active.iter().enumerate()
.filter(|(i, _)| !finished.contains(i))
.count() == 1;
for i in 0..active.len() {
if finished.contains(&i) { continue; }
if dedup_advanced.contains(&i) { continue; }
if dedup_waiting.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; }
let fresh = s.fed.is_empty()
&& s.cache.as_ref().is_some_and(|c| c.pos == 0)
&& s.snapshot_at.is_none()
&& s.ckpt_at.is_none();
let budget = interactive_prefill_budget(
budgets[0],
prefill_tick_explicit,
sole_unfinished,
fresh,
s.prefill_queue.len(),
);
match prefill_tick(&engine, &loaded, &mut px, s, budget) {
Ok(consumed) => {
if consumed > 0 {
prefill_single_calls += 1;
prefill_single_tokens += consumed;
}
}
Err(err) => {
let _ = s.tx.send(Event::Error(EngineError::engine(format!("prefill error: {err}"))));
finished.push(i);
}
}
}
let prefill_ms = t_prefill
.map(|started| started.elapsed().as_secs_f32() * 1000.0)
.unwrap_or(0.0);
let t_decode = Instant::now();
for i in 0..active.len() {
if finished.contains(&i) { continue; }
if !eager_only.contains(&active[i].model) { continue; }
if active[i].spec.is_some() || !active[i].prefill_done
|| active[i].cache.is_none() { continue; }
let generated_before = active[i].generated.len();
let lane = active[i].lane;
let step_result = step_session(&engine, &loaded, &mut active[i], &mut spec_metrics);
let emitted = record_output_tokens(
generated_before,
active[i].generated.len(),
lane,
&mut n_tokens_out,
&mut lane_tokens,
);
if emitted > 0 && lane == crate::lanes::Lane::Interactive {
last_interactive_decode = Instant::now();
}
match step_result {
Ok(true) => {}
Ok(false) => finished.push(i),
Err(err) => {
let _ = active[i].tx.send(Event::Error(
EngineError::engine(format!("step error: {err}"))));
finished.push(i);
}
}
}
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()
&& !eager_only.contains(&active[i].model))
.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 generated_before = active[i].generated.len();
let lane = active[i].lane;
let (cont, next) = advance_sample_emit(&loaded, &mut active[i]);
let emitted = record_output_tokens(
generated_before,
active[i].generated.len(),
lane,
&mut n_tokens_out,
&mut lane_tokens,
);
had_interactive |= emitted > 0 && lane == crate::lanes::Lane::Interactive;
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(
EngineError::engine(format!("constraint mask: {err}"))));
finished.push(i);
continue;
}
ready.push((i, t));
}
(true, None) => {} }
}
for scheduled in group_chunks(&active, &ready, &chunk_policies) {
let ScheduledDecodeChunk { rows: chunk, wave_mid } = scheduled;
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();
match wave_mid {
Some(mid) => lm.model.decode_step_batch_sampled_lean_masked_scheduled(
&engine, &toks, &mut caches, &samp, &masks, serve_leanlogits(), mid,
),
None => 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]);
}
}
Err(err) => {
for &i in &idxs {
let _ = active[i].tx.send(Event::Error(EngineError::engine(format!("batch step: {err}"))));
finished.push(i);
}
}
}
}
if had_interactive {
last_interactive_decode = Instant::now();
}
last_batch = ready.len();
if tick_trace {
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();
let n_spec = active.iter().filter(|s| s.spec.is_some()).count();
eprintln!("[tick] act={} int={} priming={} ready={} spec={} demoted={} \
spec_calls={} spec_ms={:.1} \
prefill_single_calls={} prefill_single_tokens={} \
prefill_batch_calls={} prefill_batch_tokens={} \
prefill_ms={:.1} decode_ms={:.1}",
active.len(), n_int, n_pref, ready.len(), n_spec, n_demoted,
spec_calls, spec_ms,
prefill_single_calls, prefill_single_tokens,
prefill_batch_calls, prefill_batch_tokens, prefill_ms,
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.ckpt_at.is_some()
|| !s.cache.as_ref().is_some_and(|c| c.pos == s.fed.len()) { continue; }
if eager_only.contains(&s.model) { 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(EngineError::engine(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 mut s = active.remove(i);
retire_prefix_pin(&mut px, &mut s.prefix_pin);
let pool_key = s.pool_key(); n_completed += 1;
if panic_injection_due(n_completed) {
panic!("MEMRA_PANIC_AFTER={} fault injection: \
deliberate worker panic after {n_completed} completed request(s)",
panic_after().unwrap_or(0));
}
lane_completed[s.lane.idx()] += 1;
if s.spec_rounds > 0 { spec_telem_dirty = true; } if s.spec_drafted > 0 {
let tenant = crate::auth::meter_key(&s.cache_ns);
if let Some(event) = adsd_detector.observe(
&s.model,
tenant,
s.spec_accepted as u64,
s.spec_drafted as u64,
) {
eprintln!(
"[adsd-suspect] tenant={:?} model={:?} window_acceptance={:.3} \
baseline_acceptance={:.3} z={:.2} drafted={} detection_only=true",
event.tenant,
event.model,
event.tenant_rate,
event.baseline_rate,
event.z_score,
event.drafted,
);
}
}
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);
if prepare_park(
ParkedPool::Spec, &pool_key, &mut reuse, &mut spec_reuse,
&mut reuse_metrics, reuse_pool_per_namespace(), reuse_pool_global_cap(),
) {
spec_reuse.entry(pool_key).or_default().push(SpecReuseEntry {
sess, committed_text, affinity: s.affinity, fingerprint,
parked_at: Instant::now(),
});
}
}
} 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 tok = &loaded[&s.model].tok;
let ckpt = s.ckpt_snap;
let affinity = s.affinity;
let fingerprint = if ckpt.is_some() {
conversation_fingerprint(&s.fed, &|t| tok.token_is_control(t), false)
} else {
Vec::new()
};
if let Some(id) = affinity.as_deref() {
if let Some(sp) = spec_reuse.get_mut(&pool_key) {
let before = sp.len();
sp.retain(|e| e.affinity.as_deref() != Some(id));
reuse_metrics.spec_evictions += (before - sp.len()) as u64;
}
}
let cap = cache.max_ctx;
if prepare_park(
ParkedPool::Plain, &pool_key, &mut reuse, &mut spec_reuse,
&mut reuse_metrics, reuse_pool_per_namespace(), reuse_pool_global_cap(),
) {
reuse.entry(pool_key).or_default().push(ReuseEntry {
fed: s.fed, cache, last_logits, cap,
ckpt, affinity, fingerprint,
parked_at: Instant::now(),
});
}
}
}
}
}
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 || !finished.is_empty() {
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.prefix_misses = px.misses;
m.prefix_inserts = px.inserts;
m.prefix_evictions = px.evictions;
m.prefix_hit_tokens = px.hit_tokens;
m.admission_session_defers = n_session_defers;
m.admission_vram_defers = n_vram_defers;
m.step_oom_parks = n_step_oom_parks;
m.continuation_pool_hits = reuse_metrics.continuation_hits;
m.continuation_pool_evictions = reuse_metrics.continuation_evictions;
m.plain_affinity_rewinds = reuse_metrics.plain_affinity_rewinds;
m.spec_pool_hits = reuse_metrics.spec_hits;
m.spec_pool_misses = reuse_metrics.spec_misses;
m.spec_pool_affinity_rewinds = reuse_metrics.spec_affinity_rewinds;
m.spec_pool_evictions = reuse_metrics.spec_evictions;
m.active_sessions = active.len() as u64;
m.queued_requests = queue.len() as u64;
m.continuation_pool_entries =
reuse.values().map(|pool| pool.len() as u64).sum();
m.spec_pool_entries =
spec_reuse.values().map(|pool| pool.len() as u64).sum();
m.cuda_driver_free_bytes = engine.ctx().mem_get_info()
.map(|(free, _)| free as u64).unwrap_or(0);
let (pool_reserved, pool_used) = engine.pool_reserved_used();
m.cuda_pool_reserved_bytes = pool_reserved as u64;
m.cuda_pool_used_bytes = pool_used as u64;
m.cuda_pool_cached_bytes = engine.pool_cached_bytes() as u64;
m.lcp_hist = px.lcp_hist;
m.ns_tokens = ns_tokens.clone();
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_metrics.lifetime.clone();
m.spec_window = spec_metrics.window_snapshots();
m.adsd_suspect_total = adsd_detector.suspect_total.clone();
} }
if !finished.is_empty() && std::env::var("MEMRA_SPILL_STATS").as_deref() == Ok("1") {
let config_fallbacks = engine.spill_config_fallbacks();
if let Some((reads, bytes, errors, short, fallbacks, waits, ring_full)) = engine
.moe_pread_stats()
.or_else(|| (config_fallbacks != 0).then_some((0, 0, 0, 0, 0, 0, 0)))
{
eprintln!("[spill-pread] snapshot reads={reads} bytes={bytes} errors={errors} \
short_reads={short} config_fallbacks={config_fallbacks} \
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 fail_request(mut req: Box<Request>, error: EngineError) {
if let Some(ready) = req.constraint_ready.take() {
let _ = ready.send(Err(error));
} else {
let _ = req.tx.send(Event::Error(error));
}
}
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 error = EngineError::model_not_found(format!(
"unknown model {:?}; loaded: {:?}", req.model, order));
fail_request(req, error);
return;
}
queue.push_back(req);
}
}
}
pub(crate) fn constraint_timeout_error() -> EngineError {
EngineError::overloaded(format!(
"response_format compilation did not finish within {} ms; retry with a smaller schema",
crate::constrained::CONSTRAINT_COMPILE_TIMEOUT.as_millis(),
))
}
fn constraint_worker_limit_error() -> EngineError {
EngineError::overloaded(format!(
"response_format compiler is temporarily saturated after {} compile workers exceeded \
their deadline; retry shortly",
crate::constrained::CONSTRAINT_ABANDONED_WORKER_CAP,
))
}
fn resolve_constraint_compiles(
result_rx: &std::sync::mpsc::Receiver<crate::constrained::ConstraintCompileResult>,
pending: &mut HashMap<u64, PendingConstraintCompile>,
queue: &mut VecDeque<Box<Request>>,
) {
while let Ok(done) = result_rx.try_recv() {
let Some(mut pending_compile) = pending.remove(&done.id) else {
continue; };
if pending_compile.request.tx.is_closed() {
continue;
}
if done.finished_at > pending_compile.deadline {
fail_request(pending_compile.request, constraint_timeout_error());
continue;
}
pending_compile.request.grammar = Some(done.spec);
match done.result {
Ok(constraint) => {
pending_compile.request.prepared_constraint = Some(constraint);
if let Some(ready) = pending_compile.request.constraint_ready.take() {
if ready.send(Ok(())).is_err() {
continue; }
}
queue.push_back(pending_compile.request);
}
Err(crate::constrained::ConstraintCompileFailure::Invalid(err)) => {
fail_request(
pending_compile.request,
EngineError::invalid_param(err, "response_format"),
);
}
Err(crate::constrained::ConstraintCompileFailure::Internal(err)) => {
fail_request(pending_compile.request, EngineError::engine(err));
}
Err(crate::constrained::ConstraintCompileFailure::TimedOut) => {
fail_request(pending_compile.request, constraint_timeout_error());
}
Err(crate::constrained::ConstraintCompileFailure::AbandonedWorkerLimit) => {
fail_request(pending_compile.request, constraint_worker_limit_error());
}
}
}
}
fn expire_constraint_compiles(
pending: &mut HashMap<u64, PendingConstraintCompile>,
now: Instant,
) {
let expired: Vec<u64> = pending.iter()
.filter_map(|(id, compile)| {
(compile.request.tx.is_closed() || now >= compile.deadline).then_some(*id)
})
.collect();
for id in expired {
let Some(compile) = pending.remove(&id) else { continue };
if !compile.request.tx.is_closed() {
fail_request(compile.request, constraint_timeout_error());
}
}
}
fn constraint_poll_wait(
pending: &HashMap<u64, PendingConstraintCompile>,
now: Instant,
) -> Duration {
pending.values()
.map(|compile| compile.deadline.saturating_duration_since(now))
.min()
.unwrap_or(CONSTRAINT_RESULT_POLL)
.min(CONSTRAINT_RESULT_POLL)
}
fn request_ctx_cap(
server_ctx: usize,
model_ctx: usize,
prompt_len: usize,
max_ctx: Option<usize>,
max_new: usize,
) -> usize {
match (max_ctx, max_new) {
(Some(c), _) => c,
(None, MAX_NEW_CTX_BOUNDED) => {
let mut cap = server_ctx;
if prompt_len.saturating_add(16) > cap {
cap = prompt_len.saturating_add(server_ctx);
}
if model_ctx > 0 {
cap = cap.min(model_ctx);
}
cap
}
(None, max_new) => prompt_len
.saturating_add(max_new)
.saturating_add(8),
}
}
fn enforce_prompt_limit(
prompt_len: usize,
max_prompt_tokens: Option<usize>,
) -> Result<(), EngineError> {
if let Some(limit) = max_prompt_tokens
&& prompt_len > limit
{
return Err(EngineError::context_length(format!(
"prompt ({prompt_len} tok) exceeds configured model maximum ({limit})"
)));
}
Ok(())
}
fn prepare_request(
loaded: &HashMap<String, LoadedModel>,
req: &mut Request,
) -> Result<RequestShape, EngineError> {
let lm = &loaded[&req.model];
if req.prepared_prompt.is_none() {
if let Some(trace) = req.ttft.as_ref() {
trace.mark_tokenize_start();
}
let prompt = 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.reasoning_effort.is_none()
&& 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 {
lm.tok.apply_chat_template_tools(
&req.chat_turns,
true,
&req.tools_json,
req.think,
req.reasoning_effort.as_deref(),
).map_err(|err| EngineError::invalid_param(
format!("chat template: {err}"), "messages"))?
};
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(EngineError::invalid_param(
"empty prompt after tokenization", "prompt"));
}
if let Some(trace) = req.ttft.as_ref() {
trace.mark_tokenize_end(prompt.len());
}
req.prepared_prompt = Some(prompt);
}
let prompt_len = req.prepared_prompt.as_ref().unwrap().len();
enforce_prompt_limit(prompt_len, req.max_prompt_tokens)?;
let server_ctx = std::env::var("MEMRA_CTX")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(8192);
let ctx_cap = request_ctx_cap(
server_ctx,
lm.model.cfg.context_length as usize,
prompt_len,
req.params.max_ctx,
req.params.max_new,
);
if prompt_len >= ctx_cap {
return Err(EngineError::context_length(format!(
"prompt ({prompt_len} tok) >= context cap ({ctx_cap})")));
}
let budget = req.params.max_new.min(ctx_cap - prompt_len);
let need = prompt_len.saturating_add(budget).saturating_add(SPEC_SHRINK_SLACK);
Ok(RequestShape { ctx_cap, budget, need })
}
fn admission_request_may_spec(
lm: &LoadedModel,
req: &Request,
projected_active: usize,
prompt_len: usize,
) -> bool {
if confidence_trace_enabled() || lm.model.mtp.is_none() || !serve_spec_enabled() {
return false;
}
req.spec_k_replay.unwrap_or_else(|| {
choose_spec_k(
spec_k_pin(),
spec_gate_on(),
*spec_gate_thresholds(),
projected_active,
prompt_len,
0,
).k
}) > 0
}
fn effective_free_bytes(engine: &Engine) -> Option<(usize, usize)> {
engine.ctx().mem_get_info().ok().map(|(driver_free, _)| {
let pool_cached = engine.pool_cached_bytes();
(driver_free.saturating_add(pool_cached), pool_cached)
})
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct DualPpDeviceHeadroom {
requirement: DualPpDeviceRequirement,
free_bytes: usize,
pool_cached_bytes: usize,
pool_reserved_bytes: usize,
pool_used_bytes: usize,
}
#[derive(Debug)]
enum AdmissionHeadroom {
Primary {
free_bytes: usize,
pool_cached_bytes: usize,
},
Dual(Vec<DualPpDeviceHeadroom>),
}
impl AdmissionHeadroom {
fn sufficient(&self, primary_required: usize) -> bool {
match self {
Self::Primary { free_bytes, .. } => *free_bytes >= primary_required,
Self::Dual(devices) => devices
.iter()
.all(|device| device.free_bytes >= device.requirement.required()),
}
}
fn limiting_free_bytes(&self) -> usize {
match self {
Self::Primary { free_bytes, .. } => *free_bytes,
Self::Dual(devices) => devices
.iter()
.map(|device| device.free_bytes)
.min()
.unwrap_or(0),
}
}
fn primary_pool_cached_bytes(&self) -> usize {
match self {
Self::Primary { pool_cached_bytes, .. } => *pool_cached_bytes,
Self::Dual(_) => 0,
}
}
}
fn admission_headroom(
engine: &Engine,
dual_stages: Option<[DualPpStageAdmission; 2]>,
) -> Option<AdmissionHeadroom> {
let Some(stages) = dual_stages else {
return effective_free_bytes(engine).map(|(free_bytes, pool_cached_bytes)| {
AdmissionHeadroom::Primary { free_bytes, pool_cached_bytes }
});
};
let rt = memra_engine::pp::PpNRt::get(engine).ok()?;
if rt.n_stages() != 2 {
return None;
}
let stage_engines = [rt.engine(0, engine), rt.engine(1, engine)];
let requirements = dual_pp_device_requirements(
[stage_engines[0].ctx().ordinal(), stage_engines[1].ctx().ordinal()],
stages,
);
let mut devices = Vec::with_capacity(requirements.len());
for requirement in requirements {
let stage_engine = stage_engines
.iter()
.copied()
.find(|stage_engine| stage_engine.ctx().ordinal() == requirement.device)?;
let (free_bytes, pool_cached_bytes) = effective_free_bytes(stage_engine)?;
let (pool_reserved_bytes, pool_used_bytes) = stage_engine.pool_reserved_used();
devices.push(DualPpDeviceHeadroom {
requirement,
free_bytes,
pool_cached_bytes,
pool_reserved_bytes,
pool_used_bytes,
});
}
Some(AdmissionHeadroom::Dual(devices))
}
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,
reasoning_effort: p.reasoning_effort.clone(),
params: p.params.clone(),
sampler_cfg: p.sampler_cfg.clone(),
stop_strings: s.stop_strings.clone(),
trace_id: s.trace_id.clone(),
max_prompt_tokens: p.max_prompt_tokens,
cache_ns: s.cache_ns.clone(),
affinity: s.affinity.clone(),
lane: s.lane,
grammar: p.grammar.clone(),
prepared_constraint: None,
constraint_ready: None,
oom_retries: s.oom_retries,
spec_k_replay: Some(s.spec_k),
prepared_prompt: None,
ttft: s.ttft.clone(),
tx: s.tx.clone(),
}))
}
#[allow(clippy::too_many_arguments)]
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,
reuse_metrics: &mut ReuseMetrics,
px: &mut PrefixCache,
n_active: usize,
mut req: Request,
shape: RequestShape,
) -> Result<Session, (tokio::sync::mpsc::UnboundedSender<Event>, EngineError)> {
let lm = &loaded[&req.model];
let prompt = req.prepared_prompt.take()
.expect("admit requires a prompt prepared by the admission gate");
let RequestShape { ctx_cap, budget, need } = shape;
let pool_key: PoolKey = (req.model.clone(), req.cache_ns.clone());
let grammar = req.grammar.take();
let constrained = grammar.is_some();
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,
reasoning_effort: req.reasoning_effort.clone(),
params: req.params.clone(),
sampler_cfg: req.sampler_cfg.clone(),
grammar,
max_prompt_tokens: req.max_prompt_tokens,
});
let req_oom_retries = req.oom_retries;
let req_spec_k_replay = req.spec_k_replay;
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));
}
}
if reused.is_some() {
reuse_metrics.continuation_hits += 1;
}
if reuse_on && reused.is_none() && affinity_enabled() {
let req_fp = conversation_fingerprint(&prompt, &|t| lm.tok.token_is_control(t), true);
let candidate = if let Some(pool) = reuse.get(&pool_key) {
let mut why: String = "empty pool".into();
let cand = pool.iter().enumerate().rev().find_map(|(i, e)| {
let Some(ckpt) = e.ckpt.as_ref() else {
why = "no checkpoint retained".into(); return None;
};
let pos = ckpt.pos;
if !e.cache.can_rollback(&ckpt.snap, 0) {
why = format!("SWA ring lapped checkpoint {pos}");
return None;
}
let identity_matches = 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,
};
match affinity_resume_target(
&prompt, &e.fed, pos, e.cap, need, identity_matches,
) {
Ok(target_cap) => Some((i, target_cap)),
Err(reason) => { why = reason; None }
}
});
if cand.is_none() && !pool.is_empty() {
eprintln!("[worker] plain-affinity: declined ({why}; {} parked, {} prompt \
tokens; model {})", pool.len(), prompt.len(), req.model);
}
cand
} else {
None
};
if let Some((idx, target_cap)) = candidate {
let mut e = reuse.get_mut(&pool_key)
.expect("plain-affinity candidate pool vanished")
.remove(idx);
let ckpt = e.ckpt.take().expect("nominated candidate must carry a checkpoint");
let pos = ckpt.pos;
let old_cap = e.cap;
let restored: Result<(), Box<dyn std::error::Error>> = if target_cap > old_cap {
match alloc_with_single_reclaim_retry(
|| memra_engine::pp::new_cache(engine, &lm.model.cfg, target_cap),
|err| {
let evicted_prefix = px.evict_all();
if evicted_prefix > 0 {
eprintln!("[prefix-cache] evicted {evicted_prefix} entries after \
plain-affinity grow alloc failure; retrying");
}
let evicted_parked = if is_cuda_oom(&err.to_string()) {
evict_oldest_parked(reuse, spec_reuse, reuse_metrics)
} else {
None
};
if let Some(pool) = evicted_parked {
eprintln!("[admit-oom] plain-affinity grow: evicted oldest {} parked \
session (global LRU); retrying cache alloc once",
match pool {
ParkedPool::Plain => "plain",
ParkedPool::Spec => "spec",
});
}
evicted_prefix > 0 || evicted_parked.is_some()
},
) {
Ok(mut grown) => match memra_engine::pp::restore_cache_checkpoint(
engine,
&lm.model.cfg,
Some(&e.cache),
&mut grown,
&ckpt.snap,
) {
Ok(()) => {
e.cache = grown;
e.cap = target_cap;
Ok(())
}
Err(err) => Err(err),
},
Err(err) => Err(err),
}
} else {
memra_engine::pp::restore_cache_checkpoint(
engine,
&lm.model.cfg,
None,
&mut e.cache,
&ckpt.snap,
)
};
match restored {
Ok(()) => {
debug_assert_eq!(e.cache.pos, pos, "plain rewind landed off the checkpoint");
debug_assert_eq!(e.cache.max_ctx, e.cap, "plain cache cap metadata drift");
e.fed.truncate(pos);
e.last_logits = ckpt.last_logits;
reuse_metrics.continuation_hits += 1;
reuse_metrics.plain_affinity_rewinds += 1;
if target_cap > old_cap {
eprintln!("[worker] plain-affinity: grew parked cache {old_cap} -> \
{target_cap} rows (request-owned need)");
}
eprintln!("[worker] plain-affinity: rewound to {pos} of {} prompt tokens \
(priming {} suffix; model {})",
prompt.len(), prompt.len() - pos, req.model);
reused = Some(e);
}
Err(err) => eprintln!("[worker] plain-affinity resume failed ({err}); \
dropping session, full prime"),
}
}
}
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 (constrained, req.prepared_constraint.take()) {
(false, None) => None,
(true, Some(constraint)) => Some(constraint),
(true, None) => return Err((req.tx, EngineError::engine(
"response_format reached admission before off-tick compilation completed",
))),
(false, Some(_)) => return Err((req.tx, EngineError::engine(
"compiled response_format has no grammar specification",
))),
};
let prefix_requested = reuse_on && serve_batching() && prefix_cache_budget_bytes() > 0;
let ring_prefix_excluded = memra_engine::cache::swa_ring_on()
&& lm.model.cfg.arch.is_step35();
if prefix_requested && ring_prefix_excluded {
eprintln!("[prefix-cache] refused for MEMRA_SWA_RING=1 Step35 session (flat-history \
snapshots/restores are excluded)");
}
let prefix_on = prefix_requested && !ring_prefix_excluded;
let policy_lcp = if prefix_on { px.best_lcp(&pool_key, &prompt) } else { 0 };
let mut spec_k_decision = match req_spec_k_replay {
Some(k) => SpecKDecision { k, reason: SpecKReason::Replay },
None => choose_spec_k(
spec_k_pin(),
spec_gate_on(),
*spec_gate_thresholds(),
n_active + 1,
prompt.len(),
0,
),
};
let spec_eligible = serve_spec
&& spec_k_decision.k > 0
&& (constraint.is_none() || (sampler.is_greedy() && !constrain_host()))
&& (sampler.is_greedy() || sampler.temperature() > 0.0)
&& !greedy_penalized
&& lm.model.mtp.is_some();
let mut prefix_hit = false;
let mut prefix_pin = None;
let mut snapshot_at: Option<usize> = None;
let mut prefix_miss_lcp: 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 memra_engine::pp::new_cache(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,
ckpt: None,
affinity: None,
fingerprint: Vec::new(),
parked_at: Instant::now(),
}),
Err(err) => Err(format!("restore failed: {err}")),
},
Err(err) => Err(format!("session cache alloc failed: {err}")),
}
};
match restored {
Ok(entry) => {
prefix_pin = px.pin(&pool_key, i);
debug_assert!(prefix_pin.is_some(), "lookup entry vanished before pin");
px.hits += 1;
px.hit_tokens += entry.fed.len() as u64;
px.record_lcp(entry.fed.len()); 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);
px.record_lcp(l); prefix_miss_lcp = Some(l);
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;
for id in lm.tok.eog_ids() {
if !params.eos.contains(&id) { params.eos.push(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;
enum SpecResumeProbe {
Exact(usize),
Text { index: usize, suffix: Vec<u32> },
Affinity {
index: usize,
explicit: bool,
old_cap: usize,
target_cap: usize,
},
}
let resumed = if constraint.is_some() {
None
} else {
let mut probe = None;
if let Some(pool) = spec_reuse.get(&pool_key) {
if let Some(index) = pool.iter().rposition(|e|
e.sess.cache_max_ctx() >= ctx_cap
&& prompt.len() >= e.sess.committed.len()
&& prompt.starts_with(&e.sess.committed))
{
probe = Some(SpecResumeProbe::Exact(index));
} else if !req.prompt_text.is_empty() {
if let Some(index) = 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 rem = &req.prompt_text[pool[index].committed_text.len()..];
probe = Some(SpecResumeProbe::Text {
index,
suffix: lm.tok.encode(rem, false),
});
}
}
if probe.is_none() && affinity_enabled() {
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_map(|(index, e)| {
let Some(pos) = e.sess.rewind_pos() else {
why = "no turn checkpoint retained".into();
return None;
};
if !e.sess.rewind_is_resident() {
why = format!("SWA ring lapped checkpoint {pos}");
return None;
}
let identity_matches = 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,
};
match affinity_resume_target(
&prompt,
&e.sess.committed,
pos,
e.sess.cache_max_ctx(),
need,
identity_matches,
) {
Ok(target_cap) => Some(SpecResumeProbe::Affinity {
index,
explicit: e.affinity.is_some(),
old_cap: e.sess.cache_max_ctx(),
target_cap,
}),
Err(reason) => {
why = reason;
None
}
}
});
if cand.is_none() && !pool.is_empty() {
eprintln!("[worker] spec-affinity: declined ({why}; {} parked, {} prompt \
tokens; model {})", pool.len(), prompt.len(), req.model);
}
probe = cand;
}
}
match probe {
Some(SpecResumeProbe::Exact(index)) => Some(
spec_reuse.get_mut(&pool_key)
.expect("spec exact candidate pool vanished")
.remove(index)
.sess,
),
Some(SpecResumeProbe::Text { index, suffix }) => {
text_suffix = Some(suffix);
Some(
spec_reuse.get_mut(&pool_key)
.expect("spec text candidate pool vanished")
.remove(index)
.sess,
)
}
Some(SpecResumeProbe::Affinity {
index,
explicit,
old_cap,
target_cap,
}) => {
let mut entry = spec_reuse.get_mut(&pool_key)
.expect("spec affinity candidate pool vanished")
.remove(index);
let rewound = if target_cap > old_cap {
alloc_with_single_reclaim_retry(
|| lm.model.spec_grow_and_rewind_to_checkpoint(
engine,
&mut entry.sess,
target_cap,
),
|err| {
if !is_cuda_oom(&err.to_string()) {
return false;
}
let evicted_prefix = px.evict_all();
if evicted_prefix > 0 {
eprintln!("[prefix-cache] evicted {evicted_prefix} entries \
after spec-affinity grow OOM; retrying");
}
let evicted_parked =
evict_oldest_parked(reuse, spec_reuse, reuse_metrics);
if let Some(pool) = evicted_parked {
eprintln!("[admit-oom] spec-affinity grow: evicted oldest {} \
parked session (global LRU); retrying once",
match pool {
ParkedPool::Plain => "plain",
ParkedPool::Spec => "spec",
});
}
evicted_prefix > 0 || evicted_parked.is_some()
},
)
} else {
lm.model.spec_rewind_to_checkpoint(engine, &mut entry.sess)
};
match rewound {
Ok(Some(pos)) => {
if target_cap > old_cap {
eprintln!("[worker] spec-affinity: grew parked session {old_cap} \
-> {target_cap} rows (request-owned need)");
}
affinity_rewound = Some((
pos,
if explicit { "explicit" } else { "fingerprint" },
));
Some(entry.sess)
}
Ok(None) => {
eprintln!("[worker] affinity rewind failed (checkpoint vanished); \
dropping session, full prime");
None
}
Err(err) => {
eprintln!("[worker] affinity rewind failed ({err}); \
dropping session, full prime");
None
}
}
}
None => None,
}
};
match resumed {
Some(mut sess) => {
reuse_metrics.spec_hits += 1;
if affinity_rewound.is_some() {
reuse_metrics.spec_affinity_rewinds += 1;
}
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 => {
reuse_metrics.spec_misses += 1;
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)
{
reuse_metrics.spec_evictions += n as u64;
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 {
reuse_metrics.spec_evictions += evicted as u64;
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 alloc_with_single_reclaim_retry(
|| memra_engine::pp::new_cache(engine, &lm.model.cfg, ctx_cap),
|err| {
let evicted_prefix = px.evict_all();
if evicted_prefix > 0 {
eprintln!("[prefix-cache] evicted {evicted_prefix} entries after cache alloc \
failure; retrying");
}
let evicted_parked = if is_cuda_oom(&err.to_string()) {
evict_oldest_parked(reuse, spec_reuse, reuse_metrics)
} else {
None
};
if let Some(pool) = evicted_parked {
eprintln!("[admit-oom] reclaim-on-alloc-oom: evicted oldest {} parked \
session (global LRU); retrying cache alloc once",
match pool {
ParkedPool::Plain => "plain",
ParkedPool::Spec => "spec",
});
}
evicted_prefix > 0 || evicted_parked.is_some()
},
) {
Ok(c) => Some(c),
Err(err) => return Err((req.tx,
EngineError::engine(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())
};
if spec.is_some() && req_spec_k_replay.is_none() {
spec_k_decision = choose_spec_k(
spec_k_pin(),
spec_gate_on(),
*spec_gate_thresholds(),
n_active + 1,
n_prompt,
n_cached,
);
}
let spec_k = if spec.is_some() { spec_k_decision.k } else { 0 };
if lm.model.mtp.is_some() {
let source = if spec_k == spec_k_decision.k {
spec_k_decision.reason.as_str()
} else {
"eligibility-fallback"
};
let placement = if spec_gate_pp2_placement() {
"pp2-cross-device"
} else {
"single-or-non-pp2"
};
eprintln!(
"[spec-k] model={:?} tenant={:?} K={spec_k} source={source} \
prompt={n_prompt} cached={n_cached} lcp={policy_lcp} \
active={} placement={placement}",
req.model,
crate::auth::meter_key(&req.cache_ns),
n_active + 1,
);
}
let ckpt_at = if affinity_enabled()
&& spec.is_none()
&& !eager_only_model(lm)
&& !confidence_trace_enabled()
&& (req.affinity.is_some()
|| plain_ckpt_nominatable(&prompt, &|t| lm.tok.token_is_control(t)))
{
plain_checkpoint_boundary(&prompt, &|t| lm.tok.token_is_control(t))
.filter(|&b| b > seed_fed.len())
} else {
None
};
Ok(Session {
model: req.model,
spec_k,
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(),
tokens_emitted: 0,
params,
stop_strings: req.stop_strings,
trace_id: req.trace_id,
emitted_bytes: 0,
budget,
n_prompt,
n_cached,
snapshot_at,
ckpt_at,
ckpt_snap: None,
prefix_miss_lcp,
seed_prefix,
prefix_pin,
tx: req.tx,
ttft: req.ttft,
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 send_token_event(s: &mut Session, id: u32, text: String) -> bool {
if s.tx.send(Event::Token { id, text }).is_err() {
return false;
}
s.tokens_emitted += 1;
true
}
fn spec_visible_len(tokens: &[u32], request_room: usize, eos_ids: &[u32]) -> usize {
let capped = tokens.len().min(request_room);
tokens[..capped]
.iter()
.position(|id| eos_ids.contains(id))
.map_or(capped, |i| i + 1)
}
#[derive(Debug, Default, PartialEq, Eq)]
struct SpecEmitResult {
sent: usize,
send_ok: bool,
}
fn emit_spec_token_events<D, S>(
tokens: &[u32],
remaining: &mut usize,
decoded: &mut Vec<u8>,
cursor: &mut usize,
eos_ids: &[u32],
eos_seen: &mut bool,
mut decode: D,
mut send: S,
) -> SpecEmitResult
where
D: FnMut(u32) -> Vec<u8>,
S: FnMut(Event) -> bool,
{
let mut result = SpecEmitResult { send_ok: true, ..Default::default() };
for &id in tokens {
if *remaining == 0 || *eos_seen {
break;
}
*remaining -= 1;
let text = if eos_ids.contains(&id) {
*eos_seen = true;
String::new()
} else {
decoded.extend_from_slice(&decode(id));
utf8_delta(decoded, cursor)
};
if !send(Event::Token { id, text }) {
result.send_ok = false;
break;
}
result.sent += 1;
}
result
}
#[allow(clippy::too_many_arguments)]
fn record_output_progress(
generated_before: usize,
generated_after: usize,
lane: Lane,
elapsed_ms: f32,
n_tokens_out: &mut u64,
lane_tokens: &mut [u64; 3],
step_stats: &mut StepStats,
last_interactive_decode: &mut Instant,
) -> usize {
let emitted = record_output_tokens(
generated_before,
generated_after,
lane,
n_tokens_out,
lane_tokens,
);
if emitted == 0 {
return 0;
}
if lane == Lane::Interactive {
let per_token_ms = elapsed_ms / emitted as f32;
for _ in 0..emitted {
step_stats.record(per_token_ms);
}
*last_interactive_decode = Instant::now();
}
emitted
}
fn record_output_tokens(
generated_before: usize,
generated_after: usize,
lane: Lane,
n_tokens_out: &mut u64,
lane_tokens: &mut [u64; 3],
) -> usize {
let emitted = generated_after.saturating_sub(generated_before);
*n_tokens_out += emitted as u64;
lane_tokens[lane.idx()] += emitted as u64;
emitted
}
fn prefix_fanout_eligible(
s: &Session,
eager_only: &std::collections::HashSet<String>,
) -> bool {
s.spec.is_none()
&& s.graph.is_none()
&& !s.prefill_done
&& s.lane == crate::lanes::Lane::Interactive
&& s.fed.is_empty()
&& s.n_cached == 0
&& s.prefix_pin.is_none()
&& s.prefix_miss_lcp.is_some()
&& s.snapshot_at.is_none()
&& s.ckpt_at.is_none()
&& s.cache.as_ref().is_some_and(|c| c.pos == 0 && !c.has_swa_ring())
&& !eager_only.contains(&s.model)
&& s.prefill_queue.len() >= PREFIX_CACHE_MIN_TOKENS
}
#[allow(clippy::too_many_arguments)]
fn dedup_interactive_prefixes(
engine: &Engine,
loaded: &HashMap<String, LoadedModel>,
eager_only: &std::collections::HashSet<String>,
px: &mut PrefixCache,
active: &mut [Session],
finished: &mut Vec<usize>,
prefix_cap: usize,
n_cached_in: &mut u64,
ns_tokens: &mut HashMap<String, [u64; 2]>,
) -> std::collections::HashSet<usize> {
let mut advanced = std::collections::HashSet::new();
if !prefix_dedup_enabled()
|| prefix_cap < PREFIX_CACHE_MIN_TOKENS
|| confidence_trace_enabled()
{
return advanced;
}
let candidates: Vec<PrefixFanoutCandidate> = active.iter().enumerate()
.filter(|(i, s)| {
!finished.contains(i) && prefix_fanout_eligible(s, eager_only)
})
.map(|(active_idx, s)| PrefixFanoutCandidate {
active_idx,
key: s.pool_key(),
prompt: s.prefill_queue.iter().copied().collect(),
})
.collect();
for group in prefix_fanout_groups(&candidates, prefix_cap) {
let leader_i = group.members[0];
if finished.contains(&leader_i) {
continue;
}
let prefix: Vec<u32> = active[leader_i].prefill_queue.iter()
.take(group.prefix_len)
.copied()
.collect();
let queued_after = active[leader_i].prefill_queue.len() - group.prefix_len;
let model = active[leader_i].model.clone();
let key = active[leader_i].pool_key();
let t0 = Instant::now();
let leader_out = {
let s = &mut active[leader_i];
loaded[&model].model.prime_cache(
engine,
&prefix,
s.cache.as_mut().unwrap(),
queued_after,
)
};
let (leader_logits, _h, _x) = match leader_out {
Ok(out) => out,
Err(err) => {
let _ = active[leader_i].tx.send(Event::Error(EngineError::engine(
format!("prefix fanout prime failed: {err}"),
)));
finished.push(leader_i);
eprintln!("[prefix-dedup] leader prime FAILED (model {model}): {err}");
continue;
}
};
advanced.insert(leader_i);
let snapshot = prefix_snapshot(
engine,
active[leader_i].cache.as_ref().unwrap(),
&prefix,
&leader_logits,
);
{
let s = &mut active[leader_i];
s.last_logits = leader_logits;
s.prefill_queue.drain(..group.prefix_len);
for &tok in &prefix {
s.fed.push(tok);
s.sampler.accept(tok);
}
s.prefill_done = s.prefill_queue.is_empty();
s.prefix_miss_lcp = None;
}
let entry = match snapshot {
Ok(entry) => entry,
Err(err) => {
eprintln!("[prefix-dedup] snapshot failed ({err}); siblings prime cold");
continue;
}
};
let mut participants = vec![leader_i];
for &i in group.members.iter().skip(1) {
if finished.contains(&i) {
continue;
}
let restored = prefix_restore(
engine,
active[i].cache.as_mut().unwrap(),
&entry,
);
if let Err(err) = restored {
let _ = active[i].tx.send(Event::Error(EngineError::engine(
format!("prefix fanout restore failed: {err}"),
)));
finished.push(i);
eprintln!("[prefix-dedup] sibling restore FAILED (model {model}): {err}");
continue;
}
let miss_lcp = active[i].prefix_miss_lcp.take()
.expect("prefix fanout sibling must carry its admission miss");
{
let s = &mut active[i];
s.last_logits = entry.last_logits.clone();
s.prefill_queue.drain(..group.prefix_len);
for &tok in &prefix {
s.fed.push(tok);
s.sampler.accept(tok);
}
s.prefill_done = s.prefill_queue.is_empty();
s.n_cached += group.prefix_len;
s.seed_prefix = false;
}
px.promote_miss_to_hit(miss_lcp, group.prefix_len);
*n_cached_in += group.prefix_len as u64;
meter_cached_credit(
ns_tokens,
&active[i].cache_ns,
group.prefix_len as u64,
);
participants.push(i);
advanced.insert(i);
}
let pin = px.insert_pinned(&key, entry, "in-batch fanout", participants.len());
for &i in &participants {
active[i].seed_prefix = false;
if let Some(pin) = &pin {
debug_assert!(active[i].prefix_pin.is_none());
active[i].prefix_pin = Some(pin.clone());
}
}
eprintln!(
"[prefix-dedup] B={} prefix={} saved={} hash={:016x} in {:.1}ms \
retained={} (model {}{})",
participants.len(),
group.prefix_len,
group.prefix_len * participants.len().saturating_sub(1),
fnv1a(0xcbf29ce484222325, &prefix),
t0.elapsed().as_secs_f64() * 1e3,
pin.is_some(),
model,
ns_suffix(&key.1),
);
}
advanced
}
fn prefill_tick(
engine: &Engine,
loaded: &HashMap<String, LoadedModel>,
px: &mut PrefixCache,
s: &mut Session,
budget: usize,
) -> Result<usize, Box<dyn std::error::Error>> {
if let Some(trace) = s.ttft.as_ref() {
trace.mark_prime_start();
}
let lm = &loaded[&s.model];
let q = s.prefill_queue.len();
if q == 0 {
s.prefill_done = true;
maybe_prefix_seed(engine, px, s);
if let Some(trace) = s.ttft.as_ref() {
trace.mark_prime_end();
}
return Ok(0);
}
let mut consumed = 0usize;
let eager_mono = eager_only_model(lm);
let carried = s.cache.as_ref().is_some_and(|c| c.pos > 0);
if eager_mono {
s.snapshot_at = None;
s.ckpt_at = None;
}
let fed_len = s.fed.len();
let bound_rem = [s.snapshot_at, s.ckpt_at].into_iter()
.flatten()
.filter(|&b| b > fed_len)
.map(|b| b - fed_len)
.min();
if !confidence_trace_enabled()
&& q >= memra_engine::hybrid_forward::PRIME_MIN_T.max(2)
&& budget >= memra_engine::hybrid_forward::PRIME_MIN_T
&& !(eager_mono && carried)
&& bound_rem.is_none_or(|r| r >= memra_engine::hybrid_forward::PRIME_MIN_T)
{
let mut take = if eager_mono { q } else { 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.prefill_queue.len())?;
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");
}
maybe_plain_checkpoint(engine, s);
if s.prefill_queue.is_empty() {
s.prefill_done = true;
maybe_prefix_seed(engine, px, s);
if let Some(trace) = s.ttft.as_ref() {
trace.mark_prime_end();
}
}
Ok(consumed)
}
fn eager_only_model(lm: &LoadedModel) -> bool {
lm.model.cfg.gemma4.is_some() || lm.model.is_gemma4_e4b()
}
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(EngineError::engine(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 let Some(trace) = s.ttft.as_ref() {
trace.mark_first_decode();
}
if s.params.eos.contains(&next) {
if !send_token_event(s, next, String::new()) {
abort_log(s);
return (false, None);
}
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(EngineError::engine(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 !send_token_event(s, next, delta) {
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 let Some(trace) = s.ttft.as_ref() {
trace.mark_first_decode();
}
if s.params.eos.contains(&tok) {
if !send_token_event(s, tok, String::new()) {
abort_log(s);
return (false, ());
}
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(EngineError::engine(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 !send_token_event(s, tok, delta) {
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))
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
struct DecodeChunkPolicy {
wave_cap: usize,
dual: bool,
}
impl DecodeChunkPolicy {
const fn serial(wave_cap: usize) -> Self {
Self { wave_cap, dual: false }
}
fn tick_cap(self) -> usize {
if self.dual { self.wave_cap.saturating_mul(2) } else { self.wave_cap }
}
fn wave_mid(self, width: usize) -> Option<usize> {
self.dual.then(|| memra_engine::pp::dual_pp_wave_mid(width)).flatten()
}
}
#[derive(Debug, Eq, PartialEq)]
struct ScheduledDecodeChunk {
rows: Vec<(usize, u32)>,
wave_mid: Option<usize>,
}
fn schedule_decode_chunk(
rows: Vec<(usize, u32)>,
policy: DecodeChunkPolicy,
) -> ScheduledDecodeChunk {
let wave_mid = policy.wave_mid(rows.len());
ScheduledDecodeChunk { rows, wave_mid }
}
fn resolve_decode_chunk_policy(
wave_cap: usize,
dual_requested: bool,
overlap_requested: bool,
pp2_ready: bool,
host_bounce_active: bool,
) -> DecodeChunkPolicy {
DecodeChunkPolicy {
wave_cap,
dual: dual_requested && overlap_requested && pp2_ready && !host_bounce_active,
}
}
fn chunk_cap_for(lm: &LoadedModel) -> usize {
if lm.model.cfg.step35.is_some() {
if !HybridModel::step35_batch_on() {
return 1;
}
let cap = std::env::var("MEMRA_DECODE_BATCH_CAP").ok()
.and_then(|v| v.parse().ok()).unwrap_or(8usize);
return cap.clamp(1, 8);
}
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 decode_chunk_policy(lm: &LoadedModel) -> DecodeChunkPolicy {
let wave_cap = chunk_cap_for(lm);
let pp2_ready = memra_engine::pp::batch_pp_on()
&& !memra_engine::pp::pp2_streams_off()
&& memra_engine::pp::pp_cuts(lm.model.layers.len())
.is_some_and(|fence| fence.len() == 3);
resolve_decode_chunk_policy(
wave_cap,
memra_engine::pp::dual_pp_on(),
memra_engine::pp::pp2_overlap(),
pp2_ready,
memra_engine::pp::pp_host_bounce_active(),
)
}
fn group_chunks(
active: &[Session],
ready: &[(usize, u32)],
policies: &HashMap<String, DecodeChunkPolicy>,
) -> Vec<ScheduledDecodeChunk> {
let mut chunks: Vec<(Vec<(usize, u32)>, DecodeChunkPolicy)> = Vec::new();
for &(i, t) in ready {
let model = &active[i].model;
let policy = policies.get(model).copied().unwrap_or(DecodeChunkPolicy::serial(8));
let cap = policy.tick_cap();
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)], policy)),
}
}
chunks.into_iter().map(|(rows, policy)| schedule_decode_chunk(rows, policy)).collect()
}
fn spec_pipe_pairable(
engine: &Engine,
loaded: &HashMap<String, LoadedModel>,
a: &Session,
b: &Session,
) -> bool {
if std::env::var("MEMRA_SPEC_PIPE").as_deref() != Ok("1")
|| a.model != b.model
|| a.spec_k == 0
|| a.spec_k != b.spec_k
|| a.spec.is_none()
|| b.spec.is_none()
|| !a.prefill_done
|| !b.prefill_done
|| !a.prefill_queue.is_empty()
|| !b.prefill_queue.is_empty()
|| a.generated.is_empty()
|| b.generated.is_empty()
|| !a.sampler.is_greedy()
|| !b.sampler.is_greedy()
|| a.constraint.is_some()
|| b.constraint.is_some()
|| a.generated.len() >= a.budget
|| b.generated.len() >= b.budget
{
return false;
}
for sess in [a.spec.as_ref().unwrap(), b.spec.as_ref().unwrap()] {
if sess.committed_len() == 0 || (sess.next_pred.is_none() && !sess.has_pending()) {
return false;
}
}
loaded[&a.model].model.spec_pipe_available(engine)
}
fn finish_pipelined_spec_burst(
lm: &LoadedModel,
s: &mut Session,
burst: Vec<u32>,
drafted: usize,
accepted: usize,
telem_before: memra_engine::spec::SpecTelemetry,
request_room: usize,
spec_metrics: &mut SpecMetricState,
) -> Result<bool, Box<dyn std::error::Error>> {
if let Some(trace) = s.ttft.as_ref() {
trace.mark_prime_end();
if !burst.is_empty() {
trace.mark_first_decode();
}
}
let spec = s.spec.as_ref().expect("paired spec session disappeared");
let telem_delta = spec.telemetry().delta_since(&telem_before);
spec_metrics.record(&s.model, telem_delta);
s.spec_rounds += telem_delta.rounds;
s.spec_drafted += drafted;
s.spec_accepted += accepted;
if drafted > 0 {
eprintln!("[spec-acc] ctx={} burst={}/{} cum={}/{}={:.3}",
s.fed.len(), accepted, drafted, s.spec_accepted, s.spec_drafted,
s.spec_accepted as f64 / s.spec_drafted.max(1) as f64);
}
let tok_ref = &lm.tok;
let eos_ids = s.params.eos.clone();
let public_len = spec_visible_len(&burst, request_room, &eos_ids);
let public_burst = &burst[..public_len];
let mut decoded_visible = tok_ref.decode_bytes_special(&s.generated, true);
let mut cursor = s.emitted_bytes;
let mut emit_remaining = request_room;
let mut eos_seen = false;
let tx = s.tx.clone();
let emitted = emit_spec_token_events(
public_burst,
&mut emit_remaining,
&mut decoded_visible,
&mut cursor,
&eos_ids,
&mut eos_seen,
|id| tok_ref.decode_bytes_special(&[id], true),
|event| tx.send(event).is_ok(),
);
if emitted.send_ok {
debug_assert_eq!(emitted.sent, public_len, "one token event per public spec token");
}
let mut stop: Option<StopReason> = None;
for &tok in public_burst {
s.sampler.accept(tok);
s.generated.push(tok);
s.fed.push(tok);
if eos_ids.contains(&tok) {
stop = Some(StopReason::Eos);
break;
}
}
s.tokens_emitted += emitted.sent;
s.emitted_bytes = cursor;
if !emitted.send_ok {
abort_log(s);
return Ok(false);
}
let full = String::from_utf8_lossy(&decoded_visible);
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);
}
let context_full = s.spec.as_ref().is_some_and(|spec| {
spec.committed.len() + s.spec_k + 3 >= spec.cache_max_ctx()
});
if stop.is_none() && context_full {
stop = Some(StopReason::ContextFull);
}
if let Some(reason) = stop {
finish(s, reason);
return Ok(false);
}
Ok(true)
}
fn step_spec_pair(
engine: &Engine,
loaded: &HashMap<String, LoadedModel>,
a: &mut Session,
b: &mut Session,
spec_metrics: &mut SpecMetricState,
) -> Result<(bool, bool), Box<dyn std::error::Error>> {
debug_assert!(spec_pipe_pairable(engine, loaded, a, b));
let lm = &loaded[&a.model];
for s in [&mut *a, &mut *b] {
if let Some(trace) = s.ttft.as_ref() {
trace.mark_prime_start();
}
}
let burst_t: usize = std::env::var("MEMRA_SPEC_BURST")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(32);
let room_a = a.budget.saturating_sub(a.generated.len());
let room_b = b.budget.saturating_sub(b.generated.len());
let target_a = room_a.min(burst_t);
let target_b = room_b.min(burst_t);
let telem_a = a.spec.as_ref().unwrap().telemetry();
let telem_b = b.spec.as_ref().unwrap().telemetry();
let (burst_a, burst_b) = lm.model.generate_spec_session_pair(
engine,
a.spec.as_mut().unwrap(),
target_a,
a.spec_k,
b.spec.as_mut().unwrap(),
target_b,
b.spec_k,
)?;
let (out_a, drafted_a, accepted_a) = burst_a;
let (out_b, drafted_b, accepted_b) = burst_b;
let keep_a = finish_pipelined_spec_burst(
lm, a, out_a, drafted_a, accepted_a, telem_a, room_a, spec_metrics,
)?;
let keep_b = finish_pipelined_spec_burst(
lm, b, out_b, drafted_b, accepted_b, telem_b, room_b, spec_metrics,
)?;
Ok((keep_a, keep_b))
}
fn step_session(
engine: &Engine,
loaded: &HashMap<String, LoadedModel>,
s: &mut Session,
spec_metrics: &mut SpecMetricState,
) -> Result<bool, Box<dyn std::error::Error>> {
let lm = &loaded[&s.model];
if let Some(spec) = s.spec.as_mut() {
if let Some(trace) = s.ttft.as_ref() {
trace.mark_prime_start();
}
let burst_t: usize = std::env::var("MEMRA_SPEC_BURST").ok()
.and_then(|v| v.parse().ok()).unwrap_or(32);
let k = s.spec_k;
debug_assert!(k > 0, "plain sessions must not enter the spec round");
let request_room = s.budget.saturating_sub(s.generated.len());
if request_room == 0 { finish(s, StopReason::MaxNew); return Ok(false); }
let burst_target = request_room.min(burst_t);
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.telemetry();
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 tok_ref = &lm.tok;
let mut decoded_visible = tok_ref.decode_bytes_special(&s.generated, true);
let mut cursor = s.emitted_bytes;
let mut emit_remaining = request_room;
let mut eos_seen = false;
let mut send_ok = true;
let mut token_events = 0usize;
let flush_tx = s.tx.clone();
let eos_ids = s.params.eos.clone();
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 || emit_remaining == 0 || slice.is_empty() {
return keep;
}
if let Some(trace) = s.ttft.as_ref() {
trace.mark_prime_end();
trace.mark_first_decode();
}
if !send_ok {
return keep; }
let emitted = emit_spec_token_events(
slice,
&mut emit_remaining,
&mut decoded_visible,
&mut cursor,
&eos_ids,
&mut eos_seen,
|id| tok_ref.decode_bytes_special(&[id], true),
|event| flush_tx.send(event).is_ok(),
);
token_events += emitted.sent;
send_ok = emitted.send_ok;
keep
};
let on_commit: Option<&mut dyn FnMut(&[u32]) -> bool> =
if per_burst_emit && !admit_yield { None } else { Some(&mut flush_cb) };
let prime_split = if spec.committed.is_empty()
&& !suffix.is_empty()
&& affinity_enabled()
&& (s.affinity.is_some()
|| plain_ckpt_nominatable(&suffix, &|t| lm.tok.token_is_control(t)))
{
plain_checkpoint_boundary(&suffix, &|t| lm.tok.token_is_control(t))
} else {
None
};
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_prime_split(
engine, spec, &suffix, burst_target, k, sampling, Some(&mut g),
prime_split, on_commit)?
}
None => lm.model.generate_spec_session_sampled_prime_split(
engine, spec, &suffix, burst_target, k, sampling, prime_split, on_commit)?,
};
drop(flush_cb);
if let Some(trace) = s.ttft.as_ref() {
trace.mark_prime_end();
if !burst.is_empty() {
trace.mark_first_decode();
}
}
let telem_delta = spec.telemetry().delta_since(&telem_before);
spec_metrics.record(&s.model, 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 public_len = spec_visible_len(&burst, request_room, &eos_ids);
let public_burst = &burst[..public_len];
if per_burst_emit {
let emitted = emit_spec_token_events(
public_burst,
&mut emit_remaining,
&mut decoded_visible,
&mut cursor,
&eos_ids,
&mut eos_seen,
|id| tok_ref.decode_bytes_special(&[id], true),
|event| s.tx.send(event).is_ok(),
);
token_events += emitted.sent;
send_ok = emitted.send_ok;
}
if send_ok {
debug_assert_eq!(token_events, public_len, "one token event per public spec token");
}
let mut stop: Option<StopReason> = None;
for &tok in public_burst {
s.sampler.accept(tok);
s.generated.push(tok);
s.fed.push(tok);
if s.params.eos.contains(&tok) { stop = Some(StopReason::Eos); break; }
}
s.tokens_emitted += token_events;
s.emitted_bytes = cursor;
if !send_ok {
abort_log(s);
return Ok(false);
}
let full = String::from_utf8_lossy(&decoded_visible);
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 {
if let Some(trace) = s.ttft.as_ref() {
trace.mark_prime_start();
}
let q = s.prefill_queue.len();
let eager_mono = eager_only_model(lm);
let carried = s.cache.as_ref().is_some_and(|c| c.pos > 0);
if !confidence_trace_enabled()
&& q >= memra_engine::hybrid_forward::PRIME_MIN_T.max(2)
&& !(eager_mono && carried)
{
let mut take = if eager_mono { q } else { 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.prefill_queue.len())?;
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;
if let Some(trace) = s.ttft.as_ref() {
trace.mark_prime_end();
}
}
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 let Some(trace) = s.ttft.as_ref() {
trace.mark_first_decode();
}
if s.params.eos.contains(&next) {
if !send_token_event(s, next, String::new()) {
abort_log(s);
return Ok(false);
}
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 !send_token_event(s, next, delta) {
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();
assert_eq!(
s.tokens_emitted,
s.generated.len(),
"terminal token receipt mismatch: Event::Token count != generated count",
);
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::TokenSnapshot(s.generated.clone()));
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,
});
}
fn panic_after() -> Option<u64> {
static P: std::sync::OnceLock<Option<u64>> = std::sync::OnceLock::new();
*P.get_or_init(|| std::env::var("MEMRA_PANIC_AFTER").ok().and_then(|v| v.parse().ok()))
}
static PANIC_INJECTED: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(false);
fn panic_injection_due(n_completed: u64) -> bool {
match panic_after() {
Some(n) if n_completed >= n =>
!PANIC_INJECTED.swap(true, std::sync::atomic::Ordering::SeqCst),
_ => false,
}
}
const WORKER_RESPAWN_MAX: u32 = 1;
pub(crate) const WORKER_RESPAWN_BACKOFF_BASE_S: u64 = 2;
fn worker_respawn_max() -> u32 {
static R: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
*R.get_or_init(|| std::env::var("MEMRA_WORKER_RESPAWN").ok()
.and_then(|v| v.parse().ok()).unwrap_or(WORKER_RESPAWN_MAX))
}
const EXIT_WORKER_UNRECOVERABLE: i32 = 70;
#[allow(clippy::type_complexity)]
pub fn spawn(models: Vec<(String, String, Option<String>)>, health: crate::health::SharedHealth)
-> Result<(
Sender<Cmd>,
Arc<Vec<String>>,
Arc<HashMap<String, ModelCaps>>,
SharedMetrics,
std::thread::JoinHandle<()>,
), 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();
let h2 = health.clone();
let worker_thread = std::thread::Builder::new()
.name("memra-gpu-worker".into())
.spawn(move || {
let rx = cmd_rx;
let mut ready_tx = Some(ready_tx);
let mut attempt: u32 = 0;
loop {
let (models, m, h) = (models.clone(), m2.clone(), h2.clone());
let (rtx, rrx) = std::sync::mpsc::channel();
let caller = ready_tx.take();
let load_failed = Arc::new(std::sync::atomic::AtomicBool::new(false));
let (lf, hr) = (load_failed.clone(), h2.clone());
let relay = std::thread::Builder::new()
.name("memra-worker-ready".into())
.spawn(move || {
let verdict = rrx.recv()
.unwrap_or_else(|_| Err("worker died during init".into()));
if let Err(why) = &verdict {
lf.store(true, std::sync::atomic::Ordering::SeqCst);
hr.mark_dead(format!("model load failed: {why}"));
}
if let Some(tx) = caller {
let _ = tx.send(verdict);
}
});
let outcome = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
run(models, &rx, rtx, m, h)
}));
if let Ok(t) = relay { let _ = t.join(); }
match outcome {
Ok(()) if load_failed.load(std::sync::atomic::Ordering::SeqCst) => {
if attempt == 0 {
return;
}
eprintln!("[worker] FATAL: respawn attempt {attempt} could not reload \
the models — exiting the process so the supervisor can \
restart it whole");
crate::health::sd_notify("STATUS=respawn load failed; exiting");
std::io::stderr().flush().ok();
std::process::exit(EXIT_WORKER_UNRECOVERABLE);
}
Ok(()) => {
h2.set_phase(crate::health::PHASE_DEAD);
return;
}
Err(payload) => {
let why = payload.downcast_ref::<String>().cloned()
.or_else(|| payload.downcast_ref::<&str>().map(|s| s.to_string()))
.unwrap_or_else(|| "non-string panic payload".into());
attempt += 1;
h2.mark_dead(format!("worker thread panicked: {why}"));
eprintln!("[worker] PANIC in the GPU worker thread: {why}");
if attempt > worker_respawn_max() {
eprintln!("[worker] FATAL: worker unrecoverable after {} respawn \
attempt(s) — exiting the process so the supervisor can \
restart it whole (CUDA errors are sticky per process; a \
live HTTP listener with a dead worker serves nothing)",
attempt - 1);
crate::health::sd_notify("STATUS=worker unrecoverable; exiting");
std::io::stderr().flush().ok();
std::process::exit(EXIT_WORKER_UNRECOVERABLE);
}
let backoff = std::time::Duration::from_secs(
WORKER_RESPAWN_BACKOFF_BASE_S * attempt as u64);
eprintln!("[worker] respawn attempt {attempt}/{} in {:?} \
(reloading weights)", worker_respawn_max(), backoff);
std::thread::sleep(backoff);
h2.mark_respawning();
}
}
}
})
.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,
worker_thread,
)),
Ok(Err(err)) => Err(err),
Err(_) => Err("worker died during init".into()),
}
}
#[cfg(test)]
mod tests {
use super::{
emit_spec_token_events, interactive_prefill_budget, record_output_progress,
record_output_tokens,
spec_visible_len, summarize_confidence, utf8_delta, Event,
};
use super::{
PoolKey, PrefixCache, PrefixEntry, PrefixFanoutCandidate, PrefixFanoutGroup,
PREFIX_CACHE_MIN_TOKENS, prefix_fanout_groups, retire_prefix_pin,
};
use super::{meter_account, meter_cached_credit, HashMap, METER_TENANT_CAP};
use super::{
AdsdDetector, ADSD_MIN_RATE_DROP, ADSD_SUSTAINED_OBSERVATIONS,
ADSD_TENANT_WINDOW, ADSD_Z_THRESHOLD,
};
use super::{draft_verdict, draft_verdict_message, DraftVerdict};
use super::SpecTelemetryWindow;
use super::{
admission_required, admission_reserve, alloc_with_single_reclaim_retry, context_cache_bytes,
dual_pp_boundary_slot_bytes, dual_pp_device_requirements, dual_pp_stage_admission,
enforce_prompt_limit, is_cuda_oom, oldest_parked_candidate, parked_entry_count,
prepare_park, request_ctx_cap,
AdmissionCostModel, AdmissionHeadroom, DualPpDeviceHeadroom, ParkedCandidate, ParkedPool,
ReuseMetrics, MAX_NEW_CTX_BOUNDED, SPEC_SHRINK_RESERVE,
};
use super::{
choose_spec_k, parse_spec_k_pin, resolve_spec_gate_thresholds, spec_gate_defaults,
SpecKDecision, SpecKReason, SPEC_K_CACHED_LONG, SPEC_K_COLD_LONG, SPEC_K_COLD_SHORT,
SPEC_K_LONG_CACHE_MIN, SPEC_K_LONG_PROMPT_MIN,
};
use super::{resolve_decode_chunk_policy, schedule_decode_chunk, DecodeChunkPolicy};
use super::{optipipe_controller_threshold, worker_device};
use crate::lanes::{Lane, StepStats};
#[derive(Debug)]
struct TestParkedEntry {
id: u32,
parked_at: std::time::Instant,
}
impl super::ParkedEntryAge for TestParkedEntry {
fn parked_at(&self) -> std::time::Instant { self.parked_at }
}
fn ascii_decode(id: u32) -> Vec<u8> {
vec![b'a' + id as u8]
}
#[test]
fn spec_telemetry_window_aggregates_and_evicts_by_age() {
let start = std::time::Instant::now();
let mut window = SpecTelemetryWindow::new(30.0);
let mut first = memra_engine::spec::SpecTelemetry {
rounds: 2,
drafted: 6,
accepted: 3,
..Default::default()
};
first.pos_drafted[..3].copy_from_slice(&[2, 2, 2]);
first.pos_accepted[..3].copy_from_slice(&[2, 1, 0]);
let mut second = memra_engine::spec::SpecTelemetry {
rounds: 1,
drafted: 3,
accepted: 3,
..Default::default()
};
second.pos_drafted[..3].copy_from_slice(&[1, 1, 1]);
second.pos_accepted[..3].copy_from_slice(&[1, 1, 1]);
window.push_at(start, first);
window.push_at(start + std::time::Duration::from_secs(20), second);
let combined = window.snapshot_at(start + std::time::Duration::from_secs(20));
assert_eq!((combined.rounds, combined.drafted, combined.accepted), (3, 9, 6));
assert_eq!(&combined.pos_drafted[..3], &[3, 3, 3]);
assert_eq!(&combined.pos_accepted[..3], &[3, 2, 1]);
assert_eq!(combined.tau(), 2.0);
let recent = window.snapshot_at(start + std::time::Duration::from_secs(31));
assert_eq!((recent.rounds, recent.drafted, recent.accepted), (1, 3, 3));
assert_eq!(recent.tau(), 3.0);
let empty = window.snapshot_at(start + std::time::Duration::from_secs(51));
assert_eq!(empty.rounds, 0);
}
#[test]
fn slow_constraint_compile_times_out_while_normal_decode_and_heartbeat_progress() {
let (result_tx, result_rx) = std::sync::mpsc::channel();
let (started_tx, started_rx) = std::sync::mpsc::channel();
let (release_tx, release_rx) = std::sync::mpsc::channel();
let release_rx = std::sync::Arc::new(std::sync::Mutex::new(release_rx));
let compiler = crate::constrained::ConstraintCompiler::spawn_for_test(
result_tx,
move || {
let started_tx = started_tx.clone();
let release_rx = std::sync::Arc::clone(&release_rx);
move |_| {
let _ = started_tx.send(());
let _ = release_rx.lock().unwrap().recv();
Err("deliberately slow test compile".into())
}
},
);
let deadline = std::time::Instant::now() + std::time::Duration::from_millis(100);
let mut pathological = serde_json::json!({"type": "string"});
for _ in 0..24 {
pathological = serde_json::json!({"allOf": [pathological]});
}
compiler.try_submit(
7,
crate::constrained::GrammarSpec::JsonSchema(pathological),
deadline,
).unwrap();
started_rx.recv_timeout(std::time::Duration::from_millis(50))
.expect("test compiler did not start");
let (bad_tx, mut bad_rx) = tokio::sync::mpsc::unbounded_channel();
let (ready_tx, ready_rx) = tokio::sync::oneshot::channel();
let request = Box::new(super::Request {
model: "m".into(),
prompt_ids: vec![1],
prompt_text: String::new(),
chat: false,
chat_turns: Vec::new(),
tools_json: Vec::new(),
think: memra_tokenizer::chat::ThinkMode::Default,
reasoning_effort: None,
params: memra_engine::decode::GenParams::default(),
sampler_cfg: memra_engine::sampler::SamplerConfig::default(),
stop_strings: Vec::new(),
trace_id: None,
max_prompt_tokens: None,
cache_ns: String::new(),
affinity: None,
lane: crate::lanes::Lane::Interactive,
oom_retries: 0,
spec_k_replay: None,
grammar: None, prepared_constraint: None,
constraint_ready: Some(ready_tx),
prepared_prompt: None,
ttft: None,
tx: bad_tx,
});
let mut pending = super::HashMap::new();
pending.insert(7, super::PendingConstraintCompile { request, deadline });
let mut queue = std::collections::VecDeque::new();
let health = crate::health::WorkerHealth::with_stall_ms(50);
let (normal_tx, mut normal_rx) = tokio::sync::mpsc::unbounded_channel();
let mut normal_steps = 0u32;
while !pending.is_empty() {
health.beat_busy();
normal_steps += 1;
normal_tx.send(super::Event::Token {
id: normal_steps,
text: "x".into(),
}).unwrap();
super::resolve_constraint_compiles(&result_rx, &mut pending, &mut queue);
super::expire_constraint_compiles(&mut pending, std::time::Instant::now());
std::thread::sleep(std::time::Duration::from_millis(2));
}
let error = ready_rx.blocking_recv()
.expect("constraint-ready sender dropped")
.expect_err("timed-out request must fail preflight");
assert_eq!(error.class, super::ErrClass::Overloaded);
assert!(error.message.contains("did not finish"), "{}", error.message);
assert!(normal_steps >= 10, "normal decode stopped at {normal_steps} steps");
let mut received = 0u32;
while matches!(normal_rx.try_recv(), Ok(super::Event::Token { .. })) {
received += 1;
}
assert_eq!(received, normal_steps, "normal decode events stalled");
assert!(health.live().is_ok(), "heartbeat declared stalled: {:?}", health.live());
assert!(health.snapshot().beat_age_ms < health.snapshot().stall_threshold_ms);
release_tx.send(()).unwrap();
assert!(bad_rx.try_recv().is_err(), "compile failure leaked after preflight response");
}
#[test]
fn flip_naked_default_schedules_dual_on_pp2_and_one_flag_restores_serial() {
use memra_engine::pp::{dual_pp_mode_resolve, pp2_overlap_resolve, DualPpMode};
let policy_for = |dual_env: Option<&str>, overlap_env: Option<&str>, pp2_ready: bool,
host_bounce: bool| {
let mode = dual_pp_mode_resolve(dual_env);
resolve_decode_chunk_policy(
8,
mode != DualPpMode::Off,
pp2_overlap_resolve(overlap_env, mode),
pp2_ready,
host_bounce,
)
};
assert_eq!(policy_for(None, None, true, false).tick_cap(), 16);
assert!(policy_for(None, None, true, false).dual);
assert_eq!(policy_for(Some("0"), None, true, false), DecodeChunkPolicy::serial(8));
assert_eq!(policy_for(None, None, false, false), DecodeChunkPolicy::serial(8));
assert_eq!(policy_for(None, None, true, true), DecodeChunkPolicy::serial(8));
assert_eq!(policy_for(None, Some("0"), true, false), DecodeChunkPolicy::serial(8));
}
#[test]
fn dual_pp_scheduler_balances_every_live_width_within_two_wave_caps() {
let policy = resolve_decode_chunk_policy(8, true, true, true, false);
assert_eq!(policy.tick_cap(), 16);
assert_eq!(policy.wave_mid(1), None);
for width in 2..=16 {
let mid = policy.wave_mid(width).expect("width >=2 needs two waves");
assert_eq!(mid, (width + 1) / 2);
assert!(mid <= policy.wave_cap, "wave A exceeds exact cap at c={width}");
assert!(width - mid <= policy.wave_cap,
"wave B exceeds exact cap at c={width}");
}
assert_eq!(DecodeChunkPolicy::serial(8).tick_cap(), 8);
assert_eq!(resolve_decode_chunk_policy(8, false, true, true, false).tick_cap(), 8);
assert_eq!(resolve_decode_chunk_policy(8, true, false, true, false).tick_cap(), 8);
assert_eq!(resolve_decode_chunk_policy(8, true, true, false, false).tick_cap(), 8);
}
#[test]
fn dual_pp_scheduler_resolves_host_bounce_to_serial_before_dispatch() {
let policy = resolve_decode_chunk_policy(8, true, true, true, true);
let ordered = vec![(10, 100), (11, 101), (12, 102)];
let scheduled = schedule_decode_chunk(ordered.clone(), policy);
assert_eq!(policy, DecodeChunkPolicy::serial(8));
assert_eq!(policy.tick_cap(), 8);
assert_eq!(scheduled.wave_mid, None);
assert_eq!(scheduled.rows, ordered);
}
#[test]
fn dual_pp_scheduler_keeps_priority_order_across_odd_wave_boundary() {
let ordered = vec![(10, 100), (11, 101), (20, 200), (21, 201), (30, 300)];
let scheduled = schedule_decode_chunk(
ordered.clone(),
DecodeChunkPolicy { wave_cap: 8, dual: true },
);
assert_eq!(scheduled.wave_mid, Some(3));
assert_eq!(scheduled.rows, ordered);
let mid = scheduled.wave_mid.unwrap();
assert_eq!(&scheduled.rows[..mid], &[(10, 100), (11, 101), (20, 200)]);
assert_eq!(&scheduled.rows[mid..], &[(21, 201), (30, 300)]);
}
#[test]
fn admission_cost_scales_with_each_requests_context() {
let cost = AdmissionCostModel {
plain_bytes_per_token: 12_288,
spec_bytes_per_token: 16_384,
plain_ring_bytes_per_token: 0,
spec_ring_bytes_per_token: 0,
ring_rows: 0,
activation_bytes: 64 << 20,
last_logged: None,
};
let cost_128k = cost.estimate(131_072, false);
let cost_256k = cost.estimate(262_144, false);
assert_ne!(cost_128k, cost_256k, "128k must not inherit a 256k scalar");
assert_eq!(
cost_256k - cost_128k,
cost.plain_bytes_per_token * 131_072,
);
assert_eq!(
cost.estimate(131_072, true) - cost_128k,
(cost.spec_bytes_per_token - cost.plain_bytes_per_token) * 131_072,
"the spec scratch coefficient is charged only on the spec-shaped path",
);
}
#[test]
fn admission_caps_only_the_step35_swa_byte_class() {
let ctx = 262_144;
let rows = 512 + 4096 + 31;
let full = 83_520usize;
let swa = 61_248usize;
let got = context_cache_bytes(full, swa, rows, ctx);
let expected = (full - swa) * ctx + swa * rows;
assert_eq!(got, expected);
assert!(got * 3 < full * ctx, "the physical cache must deliver the ~3.5x geometry");
}
#[test]
fn plain_reserve_is_capped_at_the_measured_transient_floor() {
let small_cost = 192 << 20;
assert_eq!(admission_reserve(false, small_cost, None), small_cost);
let big_cost = 21_894 << 20;
assert_eq!(admission_reserve(false, big_cost, None), SPEC_SHRINK_RESERVE);
assert!(SPEC_SHRINK_RESERVE < big_cost);
}
#[test]
fn spec_reserve_keeps_the_full_transient_floor() {
assert_eq!(admission_reserve(true, 64 << 20, None), SPEC_SHRINK_RESERVE);
assert_eq!(admission_reserve(true, 21_894 << 20, None), SPEC_SHRINK_RESERVE);
}
#[test]
fn reserve_override_door_binds_both_paths() {
let forced = 16 << 20;
assert_eq!(admission_reserve(true, 21_894 << 20, Some(forced)), forced);
assert_eq!(admission_reserve(false, 21_894 << 20, Some(forced)), forced);
assert_eq!(admission_reserve(false, 8 << 20, Some(forced)), 8 << 20);
}
#[test]
fn dual_pp_admission_checks_both_devices_and_both_receiver_slots() {
let wave_cap = 16;
let n_embd = 7_168;
let slot_bytes = dual_pp_boundary_slot_bytes(wave_cap, n_embd);
let reserve = 1_500 << 20;
let activation = 32 << 20;
let stages = dual_pp_stage_admission(
[320 << 20, 360 << 20],
activation,
reserve,
slot_bytes,
);
let devices = dual_pp_device_requirements([0, 1], stages);
assert_eq!(devices.len(), 2);
assert_eq!(devices[0].device, 0);
assert_eq!(devices[0].session_bytes, (320 << 20) + activation);
assert_eq!(devices[0].reserve_bytes, reserve);
assert_eq!(devices[0].boundary_bytes, 0);
assert_eq!(devices[1].device, 1);
assert_eq!(devices[1].session_bytes, (360 << 20) + activation);
assert_eq!(devices[1].reserve_bytes, reserve);
assert_eq!(devices[1].boundary_bytes, slot_bytes * 2);
let headroom = AdmissionHeadroom::Dual(vec![
DualPpDeviceHeadroom {
requirement: devices[0],
free_bytes: devices[0].required(),
pool_cached_bytes: 0,
pool_reserved_bytes: 0,
pool_used_bytes: 0,
},
DualPpDeviceHeadroom {
requirement: devices[1],
free_bytes: devices[1].required() - 1,
pool_cached_bytes: 0,
pool_reserved_bytes: 0,
pool_used_bytes: 0,
},
]);
assert!(!headroom.sufficient(0), "one tight PP device must defer the admit");
}
#[test]
fn dual_pp_admission_aggregates_two_stages_on_one_device() {
let stages = dual_pp_stage_admission([10, 20], 3, 5, 7);
let devices = dual_pp_device_requirements([4, 4], stages);
assert_eq!(devices.len(), 1);
assert_eq!(devices[0].device, 4);
assert_eq!(devices[0].session_bytes, 36);
assert_eq!(devices[0].reserve_bytes, 10);
assert_eq!(devices[0].boundary_bytes, 14);
assert_eq!(devices[0].required(), 60);
}
#[test]
fn dual_pp_admission_arithmetic_saturates_and_serial_is_unchanged() {
let stages = dual_pp_stage_admission([usize::MAX, 1], 1, 1, usize::MAX);
let devices = dual_pp_device_requirements([0, 1], stages);
assert_eq!(devices[0].required(), usize::MAX);
assert_eq!(devices[1].required(), usize::MAX);
assert_eq!(admission_required(usize::MAX, 1), usize::MAX);
assert_eq!(
admission_required(680 << 20, 680 << 20),
(680 << 20) * 2,
"the serial rollback keeps the previous cost + reserve equation",
);
}
#[test]
fn admission_activation_residual_is_a_high_water_not_a_new_scalar() {
let mut cost = AdmissionCostModel {
plain_bytes_per_token: 4_096,
spec_bytes_per_token: 6_144,
plain_ring_bytes_per_token: 0,
spec_ring_bytes_per_token: 0,
ring_rows: 0,
activation_bytes: 0,
last_logged: None,
};
let ctx_8k = 8_192;
let linear_8k = cost.plain_bytes_per_token * ctx_8k;
assert_eq!(cost.observe(linear_8k + 32_000_000, ctx_8k, false), Some(32_000_000));
assert_eq!(cost.observe(linear_8k + 8_000_000, ctx_8k, false), None);
assert_eq!(cost.activation_bytes, 32_000_000, "the residual never moves down");
assert_eq!(
cost.estimate(131_072, false),
cost.plain_bytes_per_token * 131_072 + 32_000_000,
"an 8k observation contributes only its fixed residual to a later 128k request",
);
}
#[test]
fn admission_reclaim_selects_the_global_oldest_across_both_pools() {
let now = std::time::Instant::now();
let plain_key: PoolKey = ("model".into(), "plain-ns".into());
let spec_key: PoolKey = ("model".into(), "spec-ns".into());
let oldest = oldest_parked_candidate([
ParkedCandidate {
pool: ParkedPool::Plain,
key: plain_key,
index: 0,
parked_at: now - std::time::Duration::from_secs(2),
},
ParkedCandidate {
pool: ParkedPool::Spec,
key: spec_key.clone(),
index: 1,
parked_at: now - std::time::Duration::from_secs(3),
},
ParkedCandidate {
pool: ParkedPool::Plain,
key: ("model".into(), "newer-ns".into()),
index: 2,
parked_at: now - std::time::Duration::from_secs(1),
},
]).expect("a parked candidate exists");
assert_eq!(oldest.pool, ParkedPool::Spec);
assert_eq!(oldest.key, spec_key);
assert_eq!(oldest.index, 1);
assert!(oldest_parked_candidate(Vec::new()).is_none());
}
#[test]
fn parked_entry_ceiling_bounds_salt_fanout_and_evicts_global_oldest() {
const PER_NAMESPACE_CAP: usize = 2;
const GLOBAL_CAP: usize = 5;
const NAMESPACES: usize = 4;
fn park(
target: ParkedPool,
key: PoolKey,
entry: TestParkedEntry,
reuse: &mut HashMap<PoolKey, Vec<TestParkedEntry>>,
spec_reuse: &mut HashMap<PoolKey, Vec<TestParkedEntry>>,
metrics: &mut ReuseMetrics,
) {
assert!(prepare_park(
target, &key, reuse, spec_reuse, metrics,
PER_NAMESPACE_CAP, GLOBAL_CAP,
));
match target {
ParkedPool::Plain => reuse.entry(key).or_default().push(entry),
ParkedPool::Spec => spec_reuse.entry(key).or_default().push(entry),
}
assert!(parked_entry_count(reuse, spec_reuse) <= GLOBAL_CAP);
}
let now = std::time::Instant::now();
let mut reuse = HashMap::new();
let mut spec_reuse = HashMap::new();
let mut metrics = ReuseMetrics::default();
let mut next_id = 0u32;
for namespace in 0..NAMESPACES {
let target = if namespace % 2 == 0 {
ParkedPool::Plain
} else {
ParkedPool::Spec
};
let key: PoolKey = ("model".into(), format!("salt-{namespace}"));
for _ in 0..PER_NAMESPACE_CAP {
park(
target, key.clone(),
TestParkedEntry {
id: next_id,
parked_at: now + std::time::Duration::from_millis(next_id as u64),
},
&mut reuse, &mut spec_reuse, &mut metrics,
);
next_id += 1;
}
}
let mut live_ids: Vec<u32> = reuse.values().chain(spec_reuse.values())
.flat_map(|pool| pool.iter().map(|entry| entry.id))
.collect();
live_ids.sort_unstable();
assert_eq!(live_ids, vec![3, 4, 5, 6, 7]);
assert_eq!(parked_entry_count(&reuse, &spec_reuse), GLOBAL_CAP);
assert!(reuse.values().chain(spec_reuse.values())
.all(|pool| pool.len() <= PER_NAMESPACE_CAP));
assert_eq!(metrics.continuation_evictions, 2);
assert_eq!(metrics.spec_evictions, 1);
}
#[test]
fn cache_alloc_oom_reclaim_retries_exactly_once() {
let attempts = std::cell::Cell::new(0usize);
let reclaims = std::cell::Cell::new(0usize);
let result: Result<(), &'static str> = alloc_with_single_reclaim_retry(
|| {
attempts.set(attempts.get() + 1);
Err("DriverError(CUDA_ERROR_OUT_OF_MEMORY, out of memory)")
},
|err| {
assert!(is_cuda_oom(err));
reclaims.set(reclaims.get() + 1);
true },
);
assert!(result.is_err());
assert_eq!(attempts.get(), 2, "one initial allocation plus one retry");
assert_eq!(reclaims.get(), 1, "reclaim runs only after the first failure");
let non_oom_attempts = std::cell::Cell::new(0usize);
let result: Result<(), &'static str> = alloc_with_single_reclaim_retry(
|| {
non_oom_attempts.set(non_oom_attempts.get() + 1);
Err("CUDA_ERROR_INVALID_VALUE")
},
|err| is_cuda_oom(err),
);
assert!(result.is_err());
assert_eq!(
non_oom_attempts.get(),
1,
"a failure without reclaimed state is not retried",
);
}
#[test]
fn admission_request_context_uses_the_requests_own_bound() {
assert_eq!(
request_ctx_cap(262_144, 262_144, 128, Some(131_072), 64),
131_072,
"an explicit 128k request must not inherit the 262k server default",
);
assert_eq!(request_ctx_cap(8_192, 262_144, 128, Some(4_096), 64), 4_096);
assert_eq!(
request_ctx_cap(8_192, 262_144, 260_000, None, MAX_NEW_CTX_BOUNDED),
262_144,
"omitted max_tokens uses the server default and remains model-capped",
);
}
#[test]
fn admission_finite_request_does_not_inherit_large_server_default() {
assert_eq!(
request_ctx_cap(262_144, 262_144, 8_120, None, 64),
8_192,
"a finite 8k request must be charged from prompt + output + margin",
);
assert_eq!(request_ctx_cap(8_192, 262_144, 128, None, 64), 200);
let shape = super::RequestShape {
ctx_cap: 45_466,
budget: 32_768,
need: 45_522,
};
assert_eq!(
shape.admission_cap(),
45_522,
"the bounded affinity growth margin is charged to this request",
);
}
#[test]
fn provider_prompt_limit_is_inclusive_and_rejects_before_admission() {
assert!(enforce_prompt_limit(7_680, Some(7_680)).is_ok());
let err = enforce_prompt_limit(7_681, Some(7_680)).unwrap_err();
assert_eq!(err.class, super::ErrClass::ContextLength);
assert!(err.message.contains("7681 tok"));
assert!(enforce_prompt_limit(usize::MAX, None).is_ok());
}
#[test]
fn spec_emission_keeps_intermediate_scheduler_surplus_public() {
let requested_max = 64usize;
let prior_generated: [u32; 0] = [];
let burst_target = 32usize;
let burst: Vec<u32> = (0..=burst_target as u32).collect();
let request_room = requested_max - prior_generated.len();
let public_len = spec_visible_len(&burst, request_room, &[]);
let mut decoded = Vec::new();
let mut cursor = 0usize;
let mut remaining = request_room;
let mut eos_seen = false;
let mut events = Vec::new();
let result = emit_spec_token_events(
&burst,
&mut remaining,
&mut decoded,
&mut cursor,
&[],
&mut eos_seen,
ascii_decode,
|event| {
if let Event::Token { id, text } = event {
events.push((id, text));
}
true
},
);
assert_eq!(burst.len(), burst_target + 1, "engine crossed its scheduler target");
assert_eq!(public_len, burst.len(), "surplus still fits the request budget");
assert_eq!(result.sent, burst.len());
assert_eq!(events.len(), burst.len());
assert_eq!(remaining, requested_max - burst.len());
}
#[test]
fn spec_emission_clamps_engine_overshoot_to_the_request_budget() {
let requested_max = 5usize;
let prior_generated = [7, 8];
let mut decoded: Vec<u8> = prior_generated.iter()
.flat_map(|&id| ascii_decode(id)).collect();
let mut cursor = decoded.len();
let burst = [0, 1, 2, 3, 4]; let request_room = requested_max - prior_generated.len();
let public_len = spec_visible_len(&burst, request_room, &[]);
let mut remaining = request_room;
let mut eos_seen = false;
let mut events = Vec::new();
let result = emit_spec_token_events(
&burst,
&mut remaining,
&mut decoded,
&mut cursor,
&[],
&mut eos_seen,
ascii_decode,
|event| {
if let Event::Token { id, text } = event {
events.push((id, text));
}
true
},
);
assert_eq!(burst.len(), 5, "engine commit remains untouched");
assert_eq!(public_len, 3);
assert_eq!(prior_generated.len() + public_len, requested_max);
assert_eq!(events.len(), request_room);
assert_eq!(result.sent, request_room);
assert_eq!(remaining, 0);
assert_eq!(&burst[public_len..], &[3, 4], "surplus is not public output");
}
#[test]
fn spec_emission_publishes_one_event_per_visible_token_id() {
let burst = [0, 1, 9, 2];
let mut remaining = burst.len();
let mut decoded = Vec::new();
let mut cursor = 0usize;
let mut eos_seen = false;
let mut events = Vec::new();
let result = emit_spec_token_events(
&burst,
&mut remaining,
&mut decoded,
&mut cursor,
&[9],
&mut eos_seen,
ascii_decode,
|event| {
if let Event::Token { id, text } = event {
events.push((id, text));
}
true
},
);
assert_eq!(events, vec![(0, "a".into()), (1, "b".into()), (9, "".into())]);
assert_eq!(result.sent, spec_visible_len(&burst, burst.len(), &[9]));
assert!(result.send_ok);
assert!(eos_seen);
}
#[test]
fn spec_path_accounts_tokens_and_timing_per_emitted_token() {
let mut total = 0u64;
let mut lanes = [0u64; 3];
let mut stats = StepStats::new(30.0);
let old_decode = std::time::Instant::now() - std::time::Duration::from_secs(1);
let mut last_decode = old_decode;
let emitted = record_output_progress(
10,
14,
Lane::Interactive,
20.0,
&mut total,
&mut lanes,
&mut stats,
&mut last_decode,
);
assert_eq!(emitted, 4, "a four-token spec commit is not one output token");
assert_eq!(total, 4);
assert_eq!(lanes, [4, 0, 0]);
assert_eq!(stats.p(50.0), Some(5.0), "20 ms / 4 emitted tokens");
assert!(last_decode > old_decode);
}
#[test]
fn batched_scheduler_counts_terminal_tokens_before_row_retirement() {
let mut total = 0u64;
let mut lanes = [0u64; 3];
for finish_path in ["Eos", "Callback", "ContextFull"] {
let emitted = record_output_tokens(
7,
8,
Lane::Interactive,
&mut total,
&mut lanes,
);
assert_eq!(emitted, 1, "{finish_path} terminal token was lost");
}
assert_eq!(
record_output_tokens(8, 8, Lane::Interactive, &mut total, &mut lanes),
0,
);
assert_eq!(total, 3);
assert_eq!(lanes, [3, 0, 0]);
}
#[test]
fn legacy_round_robin_accounts_decode_but_not_prefill_steps() {
let mut total = 0u64;
let mut lanes = [0u64; 3];
let mut stats = StepStats::new(30.0);
let mut last_decode = std::time::Instant::now();
assert_eq!(
record_output_progress(
0, 0, Lane::Interactive, 12.0, &mut total, &mut lanes, &mut stats,
&mut last_decode,
),
0,
);
assert_eq!(stats.p(50.0), None, "prefill-only calls are not output steps");
assert_eq!(
record_output_progress(
0, 1, Lane::Interactive, 7.0, &mut total, &mut lanes, &mut stats,
&mut last_decode,
),
1,
);
assert_eq!(total, 1);
assert_eq!(lanes, [1, 0, 0]);
assert_eq!(stats.p(50.0), Some(7.0));
}
#[test]
fn naked_solo_fresh_prefill_uses_one_bounded_outer_call() {
assert_eq!(interactive_prefill_budget(1024, false, true, true, 4107), 4107);
assert_eq!(interactive_prefill_budget(1024, false, true, true, 20_000), 8192);
assert_eq!(interactive_prefill_budget(1024, false, true, true, 8200), 8200);
}
#[test]
fn solo_prefill_widening_preserves_operator_and_fairness_caps() {
assert_eq!(interactive_prefill_budget(1024, true, true, true, 4107), 1024);
assert_eq!(interactive_prefill_budget(1024, false, false, true, 4107), 1024);
assert_eq!(interactive_prefill_budget(1024, false, true, false, 4107), 1024);
}
#[test]
fn spec_gate_defaults_follow_placement() {
assert_eq!(spec_gate_defaults(false), (2, 4));
assert_eq!(spec_gate_defaults(true), (0, 1));
}
#[test]
fn spec_gate_threshold_overrides_remain_explicit_and_clamped() {
let pp2_c1 = resolve_spec_gate_thresholds(true, Some(1), Some(2));
assert_eq!((pp2_c1.low, pp2_c1.high), (1, 2));
assert!(pp2_c1.low_overridden);
assert!(pp2_c1.high_overridden);
assert!(!pp2_c1.high_clamped);
let bad = resolve_spec_gate_thresholds(false, Some(4), Some(4));
assert_eq!((bad.low, bad.raw_high, bad.high), (4, 4, 5));
assert!(bad.high_clamped);
}
#[test]
fn spec_k_pin_parsing_is_explicit_and_supports_plain() {
assert_eq!(parse_spec_k_pin(None), Ok(None));
assert_eq!(parse_spec_k_pin(Some("0")), Ok(Some(0)));
assert_eq!(parse_spec_k_pin(Some("5")), Ok(Some(5)));
assert!(parse_spec_k_pin(Some("all")).unwrap_err().contains("non-negative integer"));
assert!(parse_spec_k_pin(Some("-1")).is_err());
}
#[test]
fn spec_k_operator_pin_wins_over_every_policy_row() {
let pp2 = resolve_spec_gate_thresholds(true, None, None);
assert_eq!(
choose_spec_k(Some(5), true, pp2, 8, 16, 0),
SpecKDecision { k: 5, reason: SpecKReason::OperatorPin },
);
let single = resolve_spec_gate_thresholds(false, None, None);
assert_eq!(
choose_spec_k(Some(0), true, single, 1, 8192, 8192),
SpecKDecision { k: 0, reason: SpecKReason::OperatorPin },
);
}
#[test]
fn spec_k_policy_maps_placement_and_concurrency_to_plain() {
let pp2 = resolve_spec_gate_thresholds(true, None, None);
assert_eq!(
choose_spec_k(None, true, pp2, 1, 28, 0),
SpecKDecision { k: 0, reason: SpecKReason::Placement },
);
let single = resolve_spec_gate_thresholds(false, None, None);
assert_eq!(
choose_spec_k(None, true, single, 3, 28, 0),
SpecKDecision { k: 0, reason: SpecKReason::Concurrency },
);
let pp2_c1 = resolve_spec_gate_thresholds(true, Some(1), Some(2));
assert_eq!(
choose_spec_k(None, true, pp2_c1, 1, 28, 0),
SpecKDecision { k: SPEC_K_COLD_SHORT, reason: SpecKReason::ColdShort },
);
assert_eq!(
choose_spec_k(None, true, pp2_c1, 2, 28, 0),
SpecKDecision { k: 0, reason: SpecKReason::Concurrency },
);
}
#[test]
fn spec_k_prompt_cache_table_has_exact_boundaries() {
assert_eq!(
(SPEC_K_COLD_SHORT, SPEC_K_COLD_LONG, SPEC_K_CACHED_LONG),
(3, 3, 2),
);
let single = resolve_spec_gate_thresholds(false, None, None);
assert_eq!(
choose_spec_k(None, true, single, 1, SPEC_K_LONG_PROMPT_MIN - 1, 9999),
SpecKDecision { k: SPEC_K_COLD_SHORT, reason: SpecKReason::ColdShort },
);
assert_eq!(
choose_spec_k(
None, true, single, 1, SPEC_K_LONG_PROMPT_MIN, SPEC_K_LONG_CACHE_MIN - 1,
),
SpecKDecision { k: SPEC_K_COLD_LONG, reason: SpecKReason::ColdLong },
);
assert_eq!(
choose_spec_k(
None, true, single, 1, SPEC_K_LONG_PROMPT_MIN, SPEC_K_LONG_CACHE_MIN,
),
SpecKDecision { k: SPEC_K_CACHED_LONG, reason: SpecKReason::CachedLong },
);
assert_eq!(
choose_spec_k(None, false, resolve_spec_gate_thresholds(true, None, None),
8, SPEC_K_LONG_PROMPT_MIN, SPEC_K_LONG_CACHE_MIN),
SpecKDecision { k: SPEC_K_CACHED_LONG, reason: SpecKReason::CachedLong },
);
}
#[test]
fn worker_device_defaults_to_cuda_visible_zero_and_follows_the_pp_head_stage() {
assert_eq!(worker_device(None), Ok(0));
assert_eq!(worker_device(Some("")), Ok(0));
assert_eq!(worker_device(Some("1,0")), Ok(0));
assert_eq!(worker_device(Some("0,1")), Ok(1));
assert_eq!(worker_device(Some(" 3 , 4 ")), Ok(4));
}
#[test]
fn worker_device_rejects_an_invalid_pp_device() {
let err = worker_device(Some("gpu0,1")).unwrap_err();
assert!(err.contains("invalid device"), "{err}");
assert!(err.contains("gpu0"), "{err}");
let err = worker_device(Some("1,gpu0")).unwrap_err();
assert!(err.contains("gpu0"), "{err}");
}
#[test]
fn optipipe_controller_door_is_absent_by_default_and_bounds_thresholds() {
assert_eq!(optipipe_controller_threshold(None), Ok(None));
assert_eq!(optipipe_controller_threshold(Some("0")), Ok(Some(0.0)));
assert_eq!(optipipe_controller_threshold(Some("0.7")), Ok(Some(0.7)));
assert_eq!(optipipe_controller_threshold(Some("1")), Ok(Some(1.0)));
for invalid in ["-0.01", "1.01", "NaN", "not-a-number"] {
let err = optipipe_controller_threshold(Some(invalid)).unwrap_err();
assert!(err.contains("MEMRA_OPTI_CONTROLLER_Q"), "unexpected error: {err}");
}
}
#[test]
fn step35_without_drafter_warns_and_names_the_attach_spelling() {
let v = draft_verdict(false, true);
assert_eq!(v, DraftVerdict::NoDrafterExternalMtpArch);
let msg = draft_verdict_message(&v, "step", "/m/Step-3.7-flash-IQ4_XS-00001-of-00003.gguf")
.expect("a step35 model with no drafter MUST produce a line");
assert!(msg.contains("no MTP drafter attached"), "{msg}");
assert!(msg.contains("plain decode"), "{msg}");
assert!(msg.contains("MEMRA_MODELS"), "{msg}");
assert!(msg.contains("+/path/to/"), "the '+draft' convention must be spelled: {msg}");
assert!(msg.contains("SEPARATE GGUF"), "{msg}");
assert!(msg.contains("does NOT mean"), "{msg}");
}
#[test]
fn attached_drafter_is_quiet_and_so_is_a_non_step35_model_without_one() {
let v = draft_verdict(true, true);
assert_eq!(v, DraftVerdict::Attached);
assert!(draft_verdict_message(&v, "step", "/m.gguf").is_none());
let v = draft_verdict(false, false);
assert_eq!(v, DraftVerdict::NoDrafterQuiet);
assert!(draft_verdict_message(&v, "q27", "/m.gguf").is_none());
}
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,
pins: 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_fanout_groups_only_inside_exact_model_tenant_and_salt() {
let mut a = toks(96);
a.extend([10_001, 10_002]);
let mut b = toks(96);
b.extend([20_001, 20_002, 20_003]);
let mut other_prefix = toks(96);
other_prefix[63] = u32::MAX;
let acme_s1 = crate::auth::scope_namespace("acme", "s1");
let candidates = vec![
PrefixFanoutCandidate {
active_idx: 3,
key: ("m".into(), acme_s1.clone()),
prompt: a,
},
PrefixFanoutCandidate {
active_idx: 7,
key: ("m".into(), acme_s1.clone()),
prompt: b,
},
PrefixFanoutCandidate {
active_idx: 8,
key: ("m".into(), crate::auth::scope_namespace("acme", "s2")),
prompt: toks(96),
},
PrefixFanoutCandidate {
active_idx: 9,
key: ("m".into(), crate::auth::scope_namespace("blue", "s1")),
prompt: toks(96),
},
PrefixFanoutCandidate {
active_idx: 10,
key: ("other-model".into(), acme_s1.clone()),
prompt: toks(96),
},
PrefixFanoutCandidate {
active_idx: 11,
key: ("m".into(), acme_s1),
prompt: other_prefix,
},
];
assert_eq!(
prefix_fanout_groups(&candidates, 80),
vec![PrefixFanoutGroup {
members: vec![3, 7],
prefix_len: 80,
}],
);
assert!(prefix_fanout_groups(&candidates, PREFIX_CACHE_MIN_TOKENS - 1).is_empty());
}
#[test]
fn prefix_fanout_rewrites_one_provisional_miss_exactly_once() {
let mut px = PrefixCache::default();
px.misses = 2;
px.record_lcp(0);
px.record_lcp(0);
px.promote_miss_to_hit(0, 256);
assert_eq!(px.misses, 1);
assert_eq!(px.hits, 1);
assert_eq!(px.hit_tokens, 256);
assert_eq!(px.lcp_hist[PrefixCache::lcp_bucket(0)], 1);
assert_eq!(px.lcp_hist[PrefixCache::lcp_bucket(256)], 1);
assert_eq!(px.lcp_hist.iter().sum::<u64>(), 2);
}
#[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 lcp_histogram_buckets_are_lower_edge_and_record_samples() {
assert_eq!(PrefixCache::lcp_bucket(0), 0);
assert_eq!(PrefixCache::lcp_bucket(1), 1);
assert_eq!(PrefixCache::lcp_bucket(15), 1);
assert_eq!(PrefixCache::lcp_bucket(16), 2);
assert_eq!(PrefixCache::lcp_bucket(63), 3);
assert_eq!(PrefixCache::lcp_bucket(64), 4); assert_eq!(PrefixCache::lcp_bucket(127), 4);
assert_eq!(PrefixCache::lcp_bucket(128), 5);
assert_eq!(PrefixCache::lcp_bucket(256), 6);
assert_eq!(PrefixCache::lcp_bucket(511), 6); assert_eq!(PrefixCache::lcp_bucket(512), 7);
assert_eq!(PrefixCache::lcp_bucket(4095), 9);
assert_eq!(PrefixCache::lcp_bucket(4096), 10);
assert_eq!(PrefixCache::lcp_bucket(1 << 20), 10); let mut px = PrefixCache::default();
px.record_lcp(0);
px.record_lcp(100);
px.record_lcp(100);
assert_eq!(px.lcp_hist[0], 1);
assert_eq!(px.lcp_hist[4], 2);
assert_eq!(px.lcp_hist.iter().sum::<u64>(), 3);
}
#[test]
fn meter_account_keys_by_tenant_and_bounds_rows() {
let mut m: HashMap<String, [u64; 2]> = HashMap::new();
meter_account(&mut m, &crate::auth::scope_namespace("acme", "u1"), 100, 40);
meter_account(&mut m, &crate::auth::scope_namespace("acme", "u2"), 50, 10);
meter_account(&mut m, &crate::auth::scope_namespace("blue", ""), 30, 0);
meter_cached_credit(&mut m, &crate::auth::scope_namespace("acme", "u3"), 25);
assert_eq!(m["t:acme"], [150, 75]);
assert_eq!(m["t:blue"], [30, 0]);
meter_account(&mut m, "session-7", 20, 20);
meter_account(&mut m, "", 10, 5);
assert_eq!(m["session-7"], [20, 20]);
assert_eq!(m[""], [10, 5]);
let mut m: HashMap<String, [u64; 2]> = HashMap::new();
for i in 0..METER_TENANT_CAP {
meter_account(&mut m, &format!("s{i}"), 1, 0);
}
meter_account(&mut m, "one-too-many", 7, 3);
meter_account(&mut m, "s0", 2, 1);
assert_eq!(m.len(), METER_TENANT_CAP + 1);
assert_eq!(m["(other)"], [7, 3]);
assert_eq!(m["s0"], [3, 1]);
let total: u64 = m.values().map(|r| r[0]).sum();
assert_eq!(total, METER_TENANT_CAP as u64 + 7 + 2);
}
fn seed_adsd_baseline(detector: &mut AdsdDetector) {
for i in 0..24u64 {
let accepted = [70, 73, 71, 74][i as usize % 4];
assert!(detector.observe("model-a", "t:baseline", accepted, 100).is_none());
}
}
#[test]
fn adsd_detector_fires_once_on_sustained_acceptance_collapse() {
let mut detector = AdsdDetector::default();
seed_adsd_baseline(&mut detector);
let mut events = Vec::new();
for _ in 0..(ADSD_TENANT_WINDOW + ADSD_SUSTAINED_OBSERVATIONS as usize + 4) {
if let Some(event) = detector.observe("model-a", "t:attacker", 8, 100) {
events.push(event);
}
}
assert_eq!(events.len(), 1, "one sustained incident must not count every request");
let event = &events[0];
assert_eq!(event.tenant, "t:attacker");
assert!(event.baseline_rate - event.tenant_rate >= ADSD_MIN_RATE_DROP);
assert!(event.z_score <= ADSD_Z_THRESHOLD);
let baseline_accepted = 24.0 * 72.0;
let baseline_drafted = 24.0 * 100.0;
let tenant_accepted = ADSD_TENANT_WINDOW as f64 * 8.0;
let tenant_drafted = ADSD_TENANT_WINDOW as f64 * 100.0;
let pooled_rate = (baseline_accepted + tenant_accepted)
/ (baseline_drafted + tenant_drafted);
let expected_z = (tenant_accepted / tenant_drafted
- baseline_accepted / baseline_drafted)
/ (pooled_rate * (1.0 - pooled_rate)
* (1.0 / baseline_drafted + 1.0 / tenant_drafted)).sqrt();
assert!((event.z_score - expected_z).abs() < 1e-12);
assert_eq!(detector.suspect_total["t:attacker"], 1);
}
#[test]
fn adsd_detector_fires_on_single_tenant_historical_collapse() {
let mut detector = AdsdDetector::default();
for i in 0..24u64 {
let accepted = [70, 73, 71, 74][i as usize % 4];
assert!(detector.observe("model-a", "t:solo", accepted, 100).is_none());
}
let mut events = Vec::new();
for _ in 0..(ADSD_TENANT_WINDOW + ADSD_SUSTAINED_OBSERVATIONS as usize + 4) {
if let Some(event) = detector.observe("model-a", "t:solo", 8, 100) {
events.push(event);
}
}
assert_eq!(events.len(), 1, "a single-tenant collapse must emit one incident");
assert_eq!(events[0].tenant, "t:solo");
assert!(events[0].baseline_rate - events[0].tenant_rate >= ADSD_MIN_RATE_DROP);
assert!(events[0].z_score <= ADSD_Z_THRESHOLD);
assert_eq!(detector.suspect_total["t:solo"], 1);
}
#[test]
fn adsd_detector_stays_latched_during_boiling_frog_collapse() {
let mut detector = AdsdDetector::default();
seed_adsd_baseline(&mut detector);
let mut events = 0;
for _ in 0..(ADSD_TENANT_WINDOW + ADSD_SUSTAINED_OBSERVATIONS as usize) {
events += detector.observe("model-a", "t:attacker", 8, 100).is_some() as usize;
}
assert_eq!(events, 1, "the initial collapse must latch one incident");
for accepted in (0..8u64).rev() {
for _ in 0..8 {
events += detector.observe("model-a", "t:attacker", accepted, 100)
.is_some() as usize;
}
}
let key = ("model-a".to_string(), "t:attacker".to_string());
assert!(
detector.tenant_windows[&key].incident_latched,
"the incident must not rearm from the suspect tenant diluting its own baseline",
);
assert_eq!(events, 1, "one active incident must still emit exactly once");
}
#[test]
fn adsd_detector_does_not_fire_on_normal_acceptance_noise() {
let mut detector = AdsdDetector::default();
seed_adsd_baseline(&mut detector);
for i in 0..64usize {
let accepted = [65, 76, 69, 74, 71, 78, 67, 73][i % 8];
assert!(
detector.observe("model-a", "t:noisy", accepted, 100).is_none(),
"ordinary acceptance variation must not become an ADSD incident",
);
}
assert!(!detector.suspect_total.contains_key("t:noisy"));
}
#[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,
pins: 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_pin_refcount_blocks_eviction_until_last_release() {
let k = key("");
let mut px = PrefixCache::default();
px.insert_with_budget(&k, entry_b(0, 4), "test", 8);
px.insert_with_budget(&k, entry_b(1, 4), "test", 8);
let idx = px.entries[&k].iter().position(|e| e.toks[0] == 0).unwrap();
let pin = px.pin_n(&k, idx, 2).unwrap();
px.insert_with_budget(&k, entry_b(2, 4), "test", 8);
assert_eq!(px_survivors(&px), vec![0, 2]);
assert_eq!(px.entries[&k].iter().find(|e| e.id == pin.id).unwrap().pins, 2);
assert!(px.unpin(&pin));
px.insert_with_budget(&k, entry_b(3, 4), "test", 8);
assert_eq!(px_survivors(&px), vec![0, 3]);
assert!(px.unpin(&pin));
px.insert_with_budget(&k, entry_b(4, 4), "test", 8);
assert_eq!(px_survivors(&px), vec![0, 4]);
px.insert_with_budget(&k, entry_b(5, 4), "test", 8);
assert_eq!(px_survivors(&px), vec![4, 5]);
assert!(!px.unpin(&pin), "an evicted lease id must not release another entry");
}
#[test]
fn retiring_session_releases_prefix_pin_in_release() {
let k = key("");
let mut px = PrefixCache::default();
px.insert_with_budget(&k, entry_b(0, 4), "test", 8);
let idx = px.entries[&k].iter().position(|e| e.toks[0] == 0).unwrap();
let pin = px.pin(&k, idx).unwrap();
let pin_id = pin.id;
let mut session_pin = Some(pin);
retire_prefix_pin(&mut px, &mut session_pin);
retire_prefix_pin(&mut px, &mut session_pin);
assert!(session_pin.is_none(), "retirement consumes the session's one lease");
let entry = px.entries[&k].iter().find(|e| e.id == pin_id).unwrap();
assert_eq!(entry.pins, 0, "retirement must release the cache pin in release builds");
}
#[test]
fn prefix_cache_emergency_flush_preserves_inflight_pins() {
let k = key("");
let mut px = PrefixCache::default();
px.insert_with_budget(&k, entry_b(0, 4), "test", 8);
px.insert_with_budget(&k, entry_b(1, 4), "test", 8);
let idx = px.entries[&k].iter().position(|e| e.toks[0] == 0).unwrap();
let pin = px.pin(&k, idx).unwrap();
assert_eq!(px.evict_all(), 1);
assert_eq!(px_survivors(&px), vec![0]);
assert_eq!(px.total_bytes, 4);
assert!(px.unpin(&pin));
assert_eq!(px.evict_all(), 1);
assert!(px_survivors(&px).is_empty());
assert_eq!(px.total_bytes, 0);
}
#[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_preserves_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");
let committed = toks(11_500);
let prompt = toks(prompt_len);
assert_eq!(
super::affinity_resume_target(&prompt, &committed, 11_000, laddered, need, true),
Ok(laddered),
"a sufficient ladder landing must not be inflated back to ctx_cap",
);
assert_eq!(
super::affinity_resume_target(&prompt, &committed, 11_000, 8_192, need, true),
Ok(need),
"an undersized landing grows only to this request's need",
);
}
#[test]
fn plain_checkpoint_boundary_lands_before_the_live_generation_header() {
use super::plain_checkpoint_boundary;
let sys = body(1, 40);
let user1 = body(2, 40);
let asst1 = body(3, 40);
let user2 = body(4, 40);
let prompt = {
let mut v = convo(&[&sys, &user1, &asst1, &user2]);
v.push(IM);
v.extend_from_slice(&body(9, 4));
v
};
let b = plain_checkpoint_boundary(&prompt, &is_marker)
.expect("a multi-turn chat prompt has a locatable boundary");
let last_marker = prompt.iter().rposition(|&t| t == IM).unwrap();
assert_eq!(b, last_marker, "checkpoint sits at the start of the live header segment");
assert!(b < prompt.len() && b > super::REUSE_MIN_PREFIX);
assert!(b < prompt.len() - 1);
}
#[test]
fn plain_checkpoint_boundary_uses_a_guard_window_for_raw_prompts() {
use super::{plain_checkpoint_boundary, PLAIN_CKPT_RAW_GUARD};
let prompt = toks(512);
let b = plain_checkpoint_boundary(&prompt, &is_marker).expect("long raw prompt has one");
assert_eq!(b, 512 - PLAIN_CKPT_RAW_GUARD);
}
#[test]
fn plain_ckpt_capture_requires_a_nominatable_identity() {
use super::plain_ckpt_nominatable;
assert!(!plain_ckpt_nominatable(&toks(512), &is_marker),
"markerless raw prompt must NOT arm an implicit-tier capture");
let chat = convo(&[&body(1, 24), &body(2, 24), &body(3, 24), &body(9, 8)]);
assert!(plain_ckpt_nominatable(&chat, &is_marker),
"multi-turn chat traffic keeps the implicit-tier capture");
let opener = convo(&[&body(1, 24)]);
assert!(!plain_ckpt_nominatable(&opener, &is_marker));
}
#[test]
fn plain_checkpoint_boundary_declines_short_prompts() {
use super::plain_checkpoint_boundary;
assert!(plain_checkpoint_boundary(&toks(super::REUSE_MIN_PREFIX + 4), &is_marker).is_none());
assert!(plain_checkpoint_boundary(&toks(8), &is_marker).is_none());
}
#[test]
fn plain_affinity_resume_decision_is_bytes_over_identity() {
use super::{affinity_match, AffinityMatch, fingerprint_affinity};
let (sys, user1, user2) = (body(1, 40), body(2, 40), body(4, 40));
let mut asst1 = body(3, 60);
let committed = convo(&[&sys, &user1, &asst1]);
let pos = committed.len();
let parked_fp = fp_parked(&committed);
asst1.drain(super::FP_WINDOW..asst1.len() - super::FP_WINDOW);
let request = {
let mut v = convo(&[&sys, &user1, &asst1, &user2]);
v.push(IM);
v.extend_from_slice(&body(9, 4));
v
};
assert!(!request.starts_with(&committed), "exact-extension would miss (rewritten history)");
assert!(fingerprint_affinity(&fp(&request), &parked_fp) >= super::FP_MIN_SEGMENTS);
let early_pos = convo(&[&sys, &user1]).len();
assert!(early_pos < pos);
match affinity_match(&request, &committed[..early_pos]) {
AffinityMatch::Exact { suffix_from } => assert_eq!(suffix_from, early_pos),
other => panic!("expected exact match at the pre-generation boundary, got {other:?}"),
}
assert!(request.len() > early_pos, "non-empty suffix to prime");
}
#[test]
fn plain_affinity_declines_a_divergence_below_the_checkpoint() {
use super::{affinity_match, AffinityMatch};
let committed = toks(80);
let mut request = toks(120);
request[30] = 7777; match affinity_match(&request, &committed[..50]) {
AffinityMatch::Diverged { at } => assert_eq!(at, 30),
other => panic!("a below-boundary rewrite must decline, got {other:?}"),
}
}
#[test]
fn plain_affinity_fingerprint_collision_cannot_force_a_wrong_resume() {
use super::{affinity_match, AffinityMatch, fingerprint_affinity};
let (sys, user1) = (body(1, 40), body(2, 40));
let committed = convo(&[&sys, &user1, &body(3, 40)]);
let mut request = committed.clone();
request[5] = 4242; request.extend_from_slice(&body(9, 8));
let _nominated = fingerprint_affinity(&fp(&request), &fp_parked(&committed));
match affinity_match(&request, &committed) {
AffinityMatch::Diverged { at } => assert_eq!(at, 5, "the exact diff catches it"),
other => panic!("a byte divergence must never resume, got {other:?}"),
}
}
#[test]
fn plain_affinity_pi_shape_grows_instead_of_declining() {
use super::affinity_resume_target;
let committed = toks(12_640);
let prompt = toks(12_690);
let checkpoint_pos = 12_000;
let parked_cap = 45_064;
let incoming_request_need = 45_522;
assert_eq!(
affinity_resume_target(
&prompt,
&committed,
checkpoint_pos,
parked_cap,
incoming_request_need,
true,
),
Ok(incoming_request_need),
"a nominated exact checkpoint must grow to the next request instead of declining",
);
}
#[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);
}
}