use std::collections::BTreeSet;
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use crate::policy::anchor::{decode_slide, AnchorState, SlidingRequest, WindowPolicy};
use crate::policy::pool_budget::SWA_RETAIN_GAP;
use crate::policy::radix::{align_down, NodeId, RadixCache};
use ferrox_core::cache::{
KvBlockPool, KvCache, KvPoolExhausted as CacheKvPoolExhausted, PageGroup, PagedKvCache,
PagedStoreExhausted, SharedPagedKv,
};
use ferrox_models::sampling::SamplingParams;
use ferrox_models::tokenizer::{prepend_bos, StopTokens};
use ferrox_models::{Ceiling, Decoder, Engine, KvElem, KvShape, PrefixCache, TextTokenizer};
use crate::budget::ContextCeiling;
use crate::model::ServerTokenizer;
#[derive(Debug, thiserror::Error)]
pub enum DecodeError {
#[error("prompt encoded to token id {token}, which is outside this model's vocabulary of {vocab_size} (its tokenizer does not match this checkpoint)")]
TokenOutOfVocab { token: usize, vocab_size: usize },
#[error("server is at capacity: the shared KV cache block pool has no free blocks for a new request; retry shortly")]
KvPoolExhausted,
#[error("server is at capacity: {queued} requests are already queued for the batch scheduler (limit {cap}); retry shortly")]
QueueFull { queued: usize, cap: usize },
#[error("{binding}: {detail}")]
KvBudgetExceeded {
binding: &'static str,
estimated_bytes: u64,
limit_bytes: u64,
positions: usize,
positions_limit: usize,
detail: String,
},
#[error("grammar-constrained decoding stopped: {detail}")]
GrammarConstraint { detail: String },
}
impl DecodeError {
pub fn retry_after_secs(&self) -> Option<u64> {
match self {
DecodeError::TokenOutOfVocab { .. } => None,
DecodeError::GrammarConstraint { .. } => None,
DecodeError::KvBudgetExceeded { .. } => None,
DecodeError::KvPoolExhausted | DecodeError::QueueFull { .. } => Some(1),
}
}
}
#[derive(Clone)]
pub struct KvPoolConfig {
pub pool: Arc<Mutex<KvBlockPool>>,
pub queue_wait: Duration,
}
#[derive(Clone)]
pub struct PagedKvConfig {
pub store: Arc<SharedPagedKv>,
pub queue_wait: Duration,
pub radix: Option<Arc<Mutex<RadixCache>>>,
pub anchor_token: Option<u32>,
pub slide_interval: usize,
}
struct WindowSlide {
policy: WindowPolicy,
released: usize,
locked_prefix: usize,
decode_step: usize,
anchor: AnchorState,
anchor_token: Option<u32>,
}
fn slide_hold_bound(window: usize, policy: &WindowPolicy) -> usize {
2 * (window + SWA_RETAIN_GAP) + policy.eviction_interval + 2 * policy.page_size
}
pub struct PagedLease {
caches: Vec<PagedKvCache>,
store: Arc<SharedPagedKv>,
groups: Vec<Option<PageGroup>>,
spare: Vec<PageGroup>,
adopted: Option<(usize, NodeId)>,
radix: Option<Arc<Mutex<RadixCache>>>,
window: Option<WindowSlide>,
}
impl Drop for PagedLease {
fn drop(&mut self) {
if let (Some((_, node)), Some(radix)) = (self.adopted, self.radix.as_ref()) {
radix.lock().unwrap_or_else(|p| p.into_inner()).unlock(node);
}
for group in self.groups.drain(..).flatten() {
self.store.release_group(group);
}
for group in self.spare.drain(..) {
self.store.release_group(group);
}
}
}
impl PagedLease {
pub fn caches_mut(&mut self) -> &mut Vec<PagedKvCache> {
&mut self.caches
}
pub fn store(&self) -> &Arc<SharedPagedKv> {
&self.store
}
pub fn block_size(&self) -> usize {
self.store.read(0).block_size()
}
pub fn adopted_positions(&self, block_size: usize) -> usize {
self.adopted
.map(|(groups, _)| groups * block_size)
.unwrap_or(0)
}
pub fn has_slid(&self) -> bool {
self.window.as_ref().is_some_and(|w| w.released > 0)
}
pub fn observe_sampled(&mut self, token: usize, position: usize, finished: bool) {
if let Some(w) = self.window.as_mut() {
w.anchor
.observe(token as u32, w.anchor_token, position, finished);
}
}
pub fn before_step(&mut self, position: usize) {
if self.window.is_none() {
return;
}
let block_size = self.block_size();
self.slide(position, block_size);
self.extend_to(position, block_size);
#[cfg(debug_assertions)]
debug_assert!(
self.tables_match_groups(),
"a sliding lease must own every block its tables name"
);
}
fn slide(&mut self, position: usize, block_size: usize) {
let Some(w) = self.window.as_mut() else {
return;
};
w.decode_step += 1;
let request = SlidingRequest {
position,
already_released: w.released,
locked_prefix: w.locked_prefix,
decode_step: w.decode_step,
};
let Some(decision) =
decode_slide(&request, w.anchor.anchor_len(), &w.policy, w.decode_step)
else {
return;
};
if decision.drop_anchor {
w.anchor.clear();
}
if decision.frees_nothing() {
return;
}
let (from, to) = (
decision.free_from / block_size,
decision.free_to / block_size,
);
for slot in &mut self.groups[from..to] {
if let Some(group) = slot.take() {
self.spare.push(group);
}
}
w.released = decision.free_to;
}
fn extend_to(&mut self, position: usize, block_size: usize) {
let index = position / block_size;
while self.groups.len() <= index {
let Some(group) = self.spare.pop().or_else(|| self.store.acquire_group()) else {
return;
};
let blocks = self.store.group_blocks(group);
for (cache, &block) in self.caches.iter_mut().zip(&blocks) {
cache.append_block(block);
}
self.groups.push(Some(group));
}
}
#[cfg(debug_assertions)]
fn tables_match_groups(&self) -> bool {
self.caches
.iter()
.all(|c| c.block_table().len() == self.groups.len())
}
}
fn paged_groups_needed(
max_seq_len: usize,
prompt_len: usize,
block_size: usize,
window: Option<&WindowPolicy>,
) -> usize {
paged_hold_positions(max_seq_len, prompt_len, block_size, window)
.div_ceil(block_size)
.max(1)
}
pub(crate) fn paged_hold_positions(
max_seq_len: usize,
prompt_len: usize,
block_size: usize,
window: Option<&WindowPolicy>,
) -> usize {
let Some(policy) = window else {
return max_seq_len;
};
(prompt_len + slide_hold_bound(policy.sliding_window, policy) + block_size).min(max_seq_len)
}
pub(crate) fn paged_window_policy(
decoder: &Decoder,
config: &PagedKvConfig,
) -> Option<WindowPolicy> {
let block_size = config.store.read(0).block_size();
decoder
.config
.uniform_sliding_window()
.map(|w| WindowPolicy::new(w, block_size).with_eviction_interval(config.slide_interval))
}
pub(crate) fn acquire_paged_caches(
decoder: &Decoder,
config: &PagedKvConfig,
tokens: &[usize],
max_seq_len: usize,
) -> Result<PagedLease, PagedStoreExhausted> {
let block_size = config.store.read(0).block_size();
let deadline = Instant::now() + config.queue_wait;
let adopted = match config.radix.as_ref() {
Some(radix) => {
let ids: Vec<u32> = tokens.iter().map(|&t| t as u32).collect();
let mut tree = radix.lock().unwrap_or_else(|p| p.into_inner());
let m = tree.match_prefix(&ids);
let cap = align_down(ids.len().saturating_sub(1), block_size);
let cached_len = m.cached_len.min(cap);
if cached_len == 0 {
None
} else {
tree.lock(m.node);
let per_token = tree.matched_indices(m.node);
let groups: Vec<PageGroup> = per_token[..cached_len]
.iter()
.step_by(block_size)
.map(|&g| PageGroup(g))
.collect();
Some((cached_len, m.node, groups))
}
}
None => None,
};
if let Some((_, _, groups)) = adopted.as_ref() {
for &g in groups {
config.store.retain_group(g);
}
}
let (cached_len, node, adopted_groups) = match adopted {
Some((len, node, groups)) => (len, Some(node), groups),
None => (0, None, Vec::new()),
};
let policy = paged_window_policy(decoder, config);
let total_groups = paged_groups_needed(max_seq_len, tokens.len(), block_size, policy.as_ref());
let need = total_groups.saturating_sub(adopted_groups.len());
let make_window = || {
policy.map(|policy| WindowSlide {
policy,
released: 0,
locked_prefix: cached_len,
decode_step: 0,
anchor: AnchorState::new(),
anchor_token: config.anchor_token,
})
};
loop {
let mut fresh: Vec<PageGroup> = Vec::with_capacity(need);
while fresh.len() < need {
match config.store.acquire_group() {
Some(g) => fresh.push(g),
None => break,
}
}
if fresh.len() == need {
let mut groups = adopted_groups;
groups.extend(fresh);
let caches = seed_caches(decoder, config, &groups, cached_len, block_size);
return Ok(PagedLease {
caches,
store: Arc::clone(&config.store),
groups: groups.into_iter().map(Some).collect(),
spare: Vec::new(),
adopted: node.map(|n| (cached_len / block_size, n)),
radix: config.radix.clone(),
window: make_window(),
});
}
let short = need - fresh.len();
for g in fresh {
config.store.release_group(g);
}
let mut reclaimed = 0usize;
if let Some(radix) = config.radix.as_ref() {
let freed = {
let mut tree = radix.lock().unwrap_or_else(|p| p.into_inner());
let want = (short * block_size).min(tree.evictable_size());
let per_token = tree.evict(want);
per_token.into_iter().collect::<BTreeSet<u32>>()
};
for g in freed {
config.store.release_group(PageGroup(g));
reclaimed += 1;
}
}
if reclaimed > 0 {
continue;
}
let now = Instant::now();
if now >= deadline {
drop(PagedLease {
caches: Vec::new(),
store: Arc::clone(&config.store),
groups: adopted_groups.into_iter().map(Some).collect(),
spare: Vec::new(),
adopted: node.map(|n| (cached_len / block_size, n)),
radix: config.radix.clone(),
window: make_window(),
});
return Err(PagedStoreExhausted);
}
std::thread::sleep(Duration::from_millis(10).min(deadline - now));
}
}
fn seed_caches(
decoder: &Decoder,
config: &PagedKvConfig,
groups: &[PageGroup],
cached_len: usize,
block_size: usize,
) -> Vec<PagedKvCache> {
let per_group: Vec<Vec<usize>> = groups
.iter()
.map(|&g| config.store.group_blocks(g))
.collect();
(0..decoder.layers.len())
.map(|layer| {
let table: Vec<usize> = per_group.iter().map(|blocks| blocks[layer]).collect();
let mut cache = PagedKvCache::new();
cache.adopt_blocks(table, cached_len, block_size);
cache
})
.collect()
}
pub(crate) fn publish_to_radix(lease: &mut PagedLease, tokens: &[usize], block_size: usize) {
let Some(radix) = lease.radix.clone() else {
return;
};
if lease.has_slid() {
return;
}
let ids: Vec<u32> = tokens.iter().map(|&t| t as u32).collect();
let mut per_token: Vec<u32> = Vec::with_capacity(ids.len());
for (i, group) in lease.groups.iter().enumerate() {
let group = group.expect("a lease that has not slid holds every group it names");
let covered = block_size.min(ids.len().saturating_sub(i * block_size));
for _ in 0..covered {
per_token.push(group.0);
}
}
if per_token.len() < ids.len() {
return;
}
let result = {
let mut tree = radix.lock().unwrap_or_else(|p| p.into_inner());
tree.insert_prefix(&ids, &per_token[..ids.len()])
};
let kept = result.cached_len / block_size..result.inserted_len / block_size;
for i in kept {
if let Some(Some(g)) = lease.groups.get(i) {
lease.store.retain_group(*g);
}
}
}
enum Kv {
Contiguous(Vec<KvCache>),
Paged(PagedLease),
}
impl Kv {
fn prefill(&mut self, decoder: &Decoder, tokens: &[usize], host_kv: bool) -> Vec<f32> {
match self {
Kv::Contiguous(caches) => forward_prompt_batch(decoder, tokens, 0, caches, host_kv),
Kv::Paged(lease) => {
let done = lease.adopted_positions(lease.block_size());
decoder
.forward_batch_last_paged(
&tokens[done..],
done,
&mut lease.caches,
&lease.store,
)
.expect("the whole request's pages were reserved at admission")
}
}
}
fn step(&mut self, decoder: &Decoder, token: usize, pos: usize) -> Vec<f32> {
match self {
Kv::Contiguous(caches) => decoder.forward_token(token, pos, caches),
Kv::Paged(lease) => {
lease.before_step(pos);
decoder
.forward_token_paged(token, pos, &mut lease.caches, &lease.store)
.expect("admission reserved the prompt plus this request's window bound")
}
}
}
fn contiguous_mut(&mut self) -> Option<&mut Vec<KvCache>> {
match self {
Kv::Contiguous(caches) => Some(caches),
Kv::Paged(_) => None,
}
}
fn into_contiguous(self) -> Option<Vec<KvCache>> {
match self {
Kv::Contiguous(caches) => Some(caches),
Kv::Paged(_) => None,
}
}
}
fn pool_immovable_refusal(
decoder: &Decoder,
config: &KvPoolConfig,
max_seq_len: usize,
) -> Option<DecodeError> {
let (block_size, total_blocks) = {
let pool = config.pool.lock().unwrap_or_else(|p| p.into_inner());
(pool.block_size(), pool.total_blocks())
};
if block_size == 0 || decoder.layers.is_empty() {
return None;
}
let blocks_per_layer = max_seq_len.div_ceil(block_size).max(1);
let needed = blocks_per_layer.saturating_mul(decoder.layers.len());
if needed <= total_blocks {
return None;
}
let blocks_per_layer_limit = total_blocks / decoder.layers.len();
let positions_limit = blocks_per_layer_limit * block_size;
let shape = KvShape::from_config(&decoder.config, KvElem::F32, 1);
Some(DecodeError::KvBudgetExceeded {
binding: Ceiling::DeviceMemory.code(),
estimated_bytes: shape.kv_bytes_for_tokens(max_seq_len),
limit_bytes: shape.kv_bytes_for_tokens(positions_limit),
positions: max_seq_len,
positions_limit,
detail: format!(
"request needs {needed} KV pool blocks ({max_seq_len} token positions at \
{block_size} per block, across {} layers) but the whole pool is {total_blocks} \
blocks; an idle server would refuse it identically",
decoder.layers.len()
),
})
}
fn acquire_pooled_caches(
decoder: &Decoder,
config: &KvPoolConfig,
max_seq_len: usize,
) -> Result<Vec<KvCache>, CacheKvPoolExhausted> {
let deadline = Instant::now() + config.queue_wait;
loop {
let attempt: Result<Vec<KvCache>, CacheKvPoolExhausted> = decoder
.layers
.iter()
.map(|_| {
KvCache::with_pool(
decoder.config.n_kv_heads,
decoder.config.head_dim,
Arc::clone(&config.pool),
max_seq_len,
)
})
.collect();
let now = Instant::now();
if attempt.is_ok() || now >= deadline {
return attempt;
}
std::thread::sleep(Duration::from_millis(10).min(deadline - now));
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize)]
#[serde(rename_all = "snake_case")]
pub enum FinishReason {
Stop,
StopSequence(String),
Length,
Cancelled,
}
impl FinishReason {
pub fn as_str(&self) -> &'static str {
match self {
FinishReason::Stop | FinishReason::StopSequence(_) => "stop",
FinishReason::Length => "length",
FinishReason::Cancelled => "cancelled",
}
}
pub fn matched_stop(&self) -> Option<&str> {
match self {
FinishReason::StopSequence(stop) => Some(stop),
_ => None,
}
}
}
pub use ferrox_api::Usage;
#[derive(Clone)]
pub struct GenerationParams {
pub max_tokens: usize,
pub sampling: SamplingParams,
pub seed: u64,
pub stop: Vec<String>,
pub stop_token_ids: Vec<usize>,
pub json_object: bool,
pub grammar: Option<std::sync::Arc<ferrox_models::grammar::Grammar>>,
pub cancel: Option<crate::cancel::CancelToken>,
pub ignore_eos: bool,
}
impl GenerationParams {
pub(crate) fn is_cancelled(&self) -> bool {
self.cancel.as_ref().is_some_and(|c| c.is_cancelled())
}
pub(crate) fn needs_vocab_logits(&self) -> bool {
self.json_object || self.grammar.is_some()
}
}
#[cfg(any(feature = "metal", test))]
pub(crate) fn greedy_gpu_fold_allowed(params: &GenerationParams) -> bool {
params.sampling.temperature <= 0.0 && !params.needs_vocab_logits()
}
fn chunked_prefill_tokens() -> Option<usize> {
std::env::var("FERROX_CHUNKED_PREFILL")
.ok()
.and_then(|v| v.parse().ok())
.filter(|&n| n > 0)
}
#[cfg(feature = "metal")]
fn cpu_kv_offload_enabled() -> bool {
matches!(
std::env::var("FERROX_CPU_KV_OFFLOAD").ok().as_deref(),
Some("1")
)
}
fn forward_prompt_batch(
decoder: &Decoder,
tokens: &[usize],
start_pos: usize,
caches: &mut [KvCache],
host_kv: bool,
) -> Vec<f32> {
let run = |part: &[usize], pos: usize, caches: &mut [KvCache]| {
if host_kv {
decoder.forward_batch_last_host_kv(part, pos, caches)
} else {
decoder.forward_batch_last(part, pos, caches)
}
};
if let Some(chunk) = chunked_prefill_tokens() {
let mut pos = start_pos;
let mut last = Vec::new();
for part in tokens.chunks(chunk) {
last = run(part, pos, caches);
pos += part.len();
}
last
} else {
run(tokens, start_pos, caches)
}
}
#[allow(clippy::too_many_arguments)] pub fn generate(
decoder: &Decoder,
tokenizer: &ServerTokenizer,
stop_tokens: &StopTokens,
bos_id: Option<usize>,
prompt: &str,
params: &GenerationParams,
kv_pool: Option<&KvPoolConfig>,
paged_kv: Option<&PagedKvConfig>,
prefix_cache: Option<&Mutex<PrefixCache>>,
ceiling: Option<&ContextCeiling>,
mut emit: impl FnMut(&str),
) -> Result<(FinishReason, Usage), DecodeError> {
let vocab_size = decoder.config.vocab_size;
#[cfg(feature = "metal")]
let _metal_greedy_guard = {
struct Guard;
impl Drop for Guard {
fn drop(&mut self) {
ferrox_models::set_metal_greedy_argmax(false);
}
}
if greedy_gpu_fold_allowed(params) {
ferrox_models::set_metal_greedy_argmax(true);
Some(Guard)
} else {
None
}
};
let mut tokens = tokenizer.encode(prompt);
prepend_bos(&mut tokens, bos_id);
let prompt_tokens = tokens.len();
if let Some(&bad) = tokens.iter().find(|&&t| t >= vocab_size) {
return Err(DecodeError::TokenOutOfVocab {
token: bad,
vocab_size,
});
}
let clamped;
let params = match ceiling {
Some(ceiling) => {
if let Some(err) = ceiling.prompt_refusal(prompt_tokens) {
return Err(err);
}
if let Some(err) = ceiling.overflow_refusal(prompt_tokens, params.max_tokens) {
return Err(err);
}
match ceiling.limit() {
Some(limit) if prompt_tokens.saturating_add(params.max_tokens) > limit => {
let mut p = params.clone();
p.max_tokens = limit - prompt_tokens;
tracing::debug!(
"max_tokens clamped from {} to {} by the {limit}-position context ceiling",
params.max_tokens,
p.max_tokens
);
clamped = p;
&clamped
}
_ => params,
}
}
None => params,
};
let Some(max_seq_len) = prompt_tokens.checked_add(params.max_tokens) else {
return Err(DecodeError::KvBudgetExceeded {
binding: ferrox_models::Ceiling::ContextLength.code(),
estimated_bytes: 0,
limit_bytes: 0,
positions: usize::MAX,
positions_limit: usize::MAX,
detail: format!(
"prompt of {prompt_tokens} tokens plus max_tokens of {} overflows the position \
counter, so this request cannot be served by any deployment",
params.max_tokens
),
});
};
if let Some(config) = kv_pool {
if let Some(err) = pool_immovable_refusal(decoder, config, max_seq_len) {
return Err(err);
}
}
let restored = if kv_pool.is_none() {
prefix_cache.and_then(|pc| {
let m = pc
.lock()
.unwrap_or_else(|p| p.into_inner())
.find_longest_prefix(&tokens);
(m.matched_len > 0).then_some(m)
})
} else {
None
};
let cached_tokens = restored
.as_ref()
.map(|m| m.matched_len)
.or_else(|| (prefix_cache.is_some() && kv_pool.is_none()).then_some(0));
let prefill_start = std::time::Instant::now();
let mut pos;
let mut logits: Vec<f32>;
let mut kv: Kv;
if let Some(m) = restored {
kv = Kv::Contiguous(
m.kv_caches
.expect("matched_len > 0 always carries kv_caches"),
);
let caches = &mut *kv
.contiguous_mut()
.expect("a restored prefix is contiguous by construction");
let suffix = &tokens[m.matched_len..];
if suffix.is_empty() {
if let Some(pl) = m.pending_logits {
pos = m.matched_len;
logits = pl;
} else {
let back_to = m.matched_len - 1;
for c in caches.iter_mut() {
c.truncate(back_to);
}
pos = back_to;
logits = decoder.forward_token(tokens[back_to], pos, caches);
pos += 1;
}
} else {
pos = m.matched_len;
let mut l = Vec::new();
for &tok in suffix {
l = decoder.forward_token(tok, pos, caches);
pos += 1;
}
logits = l;
}
} else {
kv = match (paged_kv, kv_pool) {
(Some(config), _) => Kv::Paged(
acquire_paged_caches(decoder, config, &tokens, max_seq_len)
.map_err(|_| DecodeError::KvPoolExhausted)?,
),
(None, Some(config)) => Kv::Contiguous(
acquire_pooled_caches(decoder, config, max_seq_len)
.map_err(|_| DecodeError::KvPoolExhausted)?,
),
(None, None) => Kv::Contiguous(
decoder
.layers
.iter()
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect(),
),
};
pos = 0;
logits = if tokens.is_empty() {
let l = kv.step(decoder, 0, pos);
pos += 1;
l
} else {
pos = tokens.len();
kv.prefill(decoder, &tokens, prefix_cache.is_some())
};
}
let prefill_secs = prefill_start.elapsed().as_secs_f64();
let decode_start = std::time::Instant::now();
#[cfg(feature = "metal")]
let kv_offload = cpu_kv_offload_enabled();
let decode_token = |id: usize| tokenizer.decode(&[id]);
let mut first_token_at: Option<std::time::Instant> = None;
let (finish, generated_ids, final_logits) = sample_until_stop(
logits,
pos,
stop_tokens,
params,
|ids| tokenizer.decode(ids),
|next, pos| {
if first_token_at.is_none() {
first_token_at = Some(std::time::Instant::now());
}
if let Kv::Paged(lease) = &mut kv {
lease.observe_sampled(next, pos + 1, false);
}
let l = kv.step(decoder, next, pos);
#[cfg(feature = "metal")]
if kv_offload {
if let Some(caches) = kv.contiguous_mut() {
decoder.sync_metal_attn_kv_to_host(caches);
}
}
l
},
&mut emit,
&decode_token,
)?;
let decode_secs = decode_start.elapsed().as_secs_f64();
logits = final_logits;
let mut usage =
Usage::new(prompt_tokens, generated_ids.len()).with_timings(prefill_secs, decode_secs);
if let Some(at) = first_token_at {
usage = usage.with_ttft(at.duration_since(prefill_start).as_secs_f64());
}
if let Some(cached) = cached_tokens {
usage = usage.with_cached_tokens(cached);
}
if let Kv::Paged(lease) = &kv {
let adopted = lease.adopted_positions(lease.block_size());
if paged_kv.is_some_and(|c| c.radix.is_some()) {
usage = usage.with_cached_tokens(adopted);
}
}
if let Kv::Paged(lease) = &mut kv {
let mut full = tokens.clone();
full.extend(generated_ids.iter().copied());
let block_size = lease.block_size();
publish_to_radix(lease, &full, block_size);
}
if kv_pool.is_none() {
if let Some(pc) = prefix_cache {
if logits.len() == vocab_size {
#[cfg(feature = "metal")]
if let Some(caches) = kv.contiguous_mut() {
decoder.sync_metal_attn_kv_to_host(caches);
}
if let Some(caches) = kv.into_contiguous() {
tokens.extend(generated_ids);
pc.lock()
.unwrap_or_else(|p| p.into_inner())
.store(tokens, caches, logits);
}
}
}
}
Ok((finish, usage))
}
#[allow(clippy::too_many_arguments)] fn sample_until_stop(
mut logits: Vec<f32>,
mut pos: usize,
stop_tokens: &StopTokens,
params: &GenerationParams,
mut decode_one: impl FnMut(&[usize]) -> String,
mut step: impl FnMut(usize, usize) -> Vec<f32>,
mut emit: impl FnMut(&str),
decode_token: &dyn Fn(usize) -> String,
) -> Result<(FinishReason, Vec<usize>, Vec<f32>), DecodeError> {
let mut matcher = crate::stop::StopMatcher::new(¶ms.stop, ¶ms.stop_token_ids);
let mut state = crate::sample_step::SampleState::new(params.seed);
const PREALLOC_CAP: usize = 4096;
let mut generated_ids: Vec<usize> = Vec::with_capacity(params.max_tokens.min(PREALLOC_CAP));
let mut finish = FinishReason::Length;
for _ in 0..params.max_tokens {
if params.is_cancelled() {
finish = FinishReason::Cancelled;
break;
}
let next = match crate::sample_step::sample_next(
&mut state,
&logits,
params,
&generated_ids,
stop_tokens,
decode_token,
)? {
crate::sample_step::Step::Token(next) => next,
crate::sample_step::Step::GrammarComplete => {
finish = FinishReason::Stop;
break;
}
};
if !params.ignore_eos && stop_tokens.contains(next) {
finish = FinishReason::Stop;
break;
}
if matcher.is_stop_token(next) {
finish = FinishReason::Stop;
break;
}
generated_ids.push(next);
logits = step(next, pos);
pos += 1;
match matcher.push(&decode_one(&[next])) {
crate::stop::StopStep::Emit(text) => {
if !text.is_empty() {
emit(&text);
}
}
crate::stop::StopStep::Matched { text, stop } => {
if !text.is_empty() {
emit(&text);
}
finish = FinishReason::StopSequence(stop);
break;
}
}
}
let tail = matcher.flush();
if !tail.is_empty() {
emit(&tail);
}
Ok((finish, generated_ids, logits))
}
pub fn generate_engine<E: Engine, T: TextTokenizer>(
engine: &E,
tokenizer: &T,
stop_tokens: &StopTokens,
bos_id: Option<usize>,
prompt: &str,
params: &GenerationParams,
mut emit: impl FnMut(&str),
) -> Result<(FinishReason, Usage), DecodeError> {
let vocab_size = engine.vocab_size();
let mut tokens = tokenizer.encode(prompt);
prepend_bos(&mut tokens, bos_id);
let prompt_tokens = tokens.len();
if let Some(&bad) = tokens.iter().find(|&&t| t >= vocab_size) {
return Err(DecodeError::TokenOutOfVocab {
token: bad,
vocab_size,
});
}
let mut state = engine.new_state();
let mut pos = 0;
let prefill_start = std::time::Instant::now();
let logits = if tokens.is_empty() {
let l = engine.forward_token(0, pos, &mut state);
pos += 1;
l
} else {
let mut l = Vec::new();
for &tok in tokens.iter() {
l = engine.forward_token(tok, pos, &mut state);
pos += 1;
}
l
};
let prefill_secs = prefill_start.elapsed().as_secs_f64();
let decode_start = std::time::Instant::now();
let mut first_token_at: Option<std::time::Instant> = None;
let (finish, generated_ids, _final_logits) = sample_until_stop(
logits,
pos,
stop_tokens,
params,
|ids| tokenizer.decode(ids),
|next, pos| {
if first_token_at.is_none() {
first_token_at = Some(std::time::Instant::now());
}
engine.forward_token(next, pos, &mut state)
},
&mut emit,
&|id: usize| tokenizer.decode(&[id]),
)?;
let decode_secs = decode_start.elapsed().as_secs_f64();
let mut usage =
Usage::new(prompt_tokens, generated_ids.len()).with_timings(prefill_secs, decode_secs);
if let Some(at) = first_token_at {
usage = usage.with_ttft(at.duration_since(prefill_start).as_secs_f64());
}
Ok((finish, usage))
}
pub(crate) fn earliest_stop_match<'a>(text: &str, stops: &'a [String]) -> Option<(usize, &'a str)> {
stops
.iter()
.filter(|s| !s.is_empty())
.filter_map(|s| text.find(s.as_str()).map(|at| (at, s.as_str())))
.min_by_key(|(at, s)| (*at, std::cmp::Reverse(s.len())))
}
pub(crate) fn floor_char_boundary(s: &str, idx: usize) -> usize {
crate::policy::detokenize::floor_char_boundary(s, idx)
}
#[cfg(test)]
mod tests {
use super::*;
use ferrox_models::config::test_dense_fixture;
fn small_decoder() -> Decoder {
Decoder::new_random_small(test_dense_fixture(), 2, 256)
}
fn greedy_params(max_tokens: usize) -> GenerationParams {
GenerationParams {
max_tokens,
sampling: SamplingParams::default(),
seed: 1,
stop: Vec::new(),
stop_token_ids: Vec::new(),
json_object: false,
grammar: None,
cancel: None,
ignore_eos: false,
}
}
#[test]
fn prompt_processing_matches_forward_batch_ground_truth_with_no_duplicate_position() {
let decoder = small_decoder();
let tokens = vec![1usize, 2, 3, 4];
let mut fresh_caches: Vec<KvCache> = decoder
.layers
.iter()
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect();
let batch_logits = decoder.forward_batch(&tokens, 0, &mut fresh_caches);
let ground_truth_next_logits = batch_logits.last().unwrap().clone();
let mut caches: Vec<KvCache> = decoder
.layers
.iter()
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect();
let mut logits = Vec::new();
for (pos, &tok) in tokens.iter().enumerate() {
logits = decoder.forward_token(tok, pos, &mut caches);
}
assert_eq!(
caches[0].seq_len, fresh_caches[0].seq_len,
"must not push any position beyond the real prompt length"
);
assert_eq!(logits.len(), ground_truth_next_logits.len());
for (i, (a, b)) in logits.iter().zip(&ground_truth_next_logits).enumerate() {
assert!(
(a - b).abs() <= 1e-5 * a.abs().max(1.0),
"logit {i} predicting the first generated token: sequential {a} vs forward_batch {b}"
);
}
}
#[test]
fn generate_greedy_output_matches_independent_step_by_step_computation() {
let decoder = small_decoder();
let prompt_ids = vec![1usize, 2, 3];
let prompt = String::from_utf8(prompt_ids.iter().map(|&b| b as u8).collect()).unwrap();
let max_tokens = 8;
let mut caches: Vec<KvCache> = decoder
.layers
.iter()
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect();
let mut logits = decoder
.forward_batch(&prompt_ids, 0, &mut caches)
.pop()
.unwrap();
let mut expected_text = String::new();
for pos in (prompt_ids.len()..).take(max_tokens) {
let next = logits
.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.map(|(i, _)| i)
.unwrap();
expected_text.push_str(&ServerTokenizer::Byte.decode(&[next]));
logits = decoder.forward_token(next, pos, &mut caches);
}
let mut actual_text = String::new();
generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(max_tokens),
None,
None,
None,
None,
|s| actual_text.push_str(s),
)
.unwrap();
assert_eq!(actual_text, expected_text);
}
#[test]
fn rejects_out_of_vocab_prompt_tokens() {
let decoder = Decoder::new_random_small(test_dense_fixture(), 2, 32);
let result = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
"hello",
&greedy_params(4),
None,
None,
None,
None,
|_| {},
);
assert!(matches!(result, Err(DecodeError::TokenOutOfVocab { .. })));
}
#[test]
fn greedy_generation_hits_length_without_eos() {
let decoder = small_decoder();
let prompt = String::from_utf8(vec![1u8, 2, 3]).unwrap();
let mut chunks = String::new();
let (finish, _usage) = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(5),
None,
None,
None,
None,
|s| chunks.push_str(s),
)
.unwrap();
assert_eq!(finish, FinishReason::Length);
}
#[test]
fn a_cancelled_generation_stops_early_and_keeps_its_tokens() {
let decoder = small_decoder();
let prompt = String::from_utf8(vec![1u8, 2, 3]).unwrap();
let cancel = crate::cancel::CancelToken::new();
let mut params = greedy_params(200);
params.cancel = Some(cancel.clone());
let mut chunks = String::new();
let mut emitted = 0usize;
let (finish, usage) = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
¶ms,
None,
None,
None,
None,
|s| {
chunks.push_str(s);
emitted += 1;
if emitted == 3 {
cancel.cancel();
}
},
)
.unwrap();
assert_eq!(finish, FinishReason::Cancelled);
assert!(
usage.completion_tokens < 200,
"cancelling did not shorten the decode: {} tokens",
usage.completion_tokens
);
assert!(
!chunks.is_empty(),
"the tokens decoded before the cancel must survive it"
);
}
#[test]
fn an_uncancelled_generation_runs_to_its_normal_end() {
let decoder = small_decoder();
let prompt = String::from_utf8(vec![1u8, 2, 3]).unwrap();
let mut params = greedy_params(5);
params.cancel = Some(crate::cancel::CancelToken::new());
let (finish, usage) = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
¶ms,
None,
None,
None,
None,
|_| {},
)
.unwrap();
assert_eq!(finish, FinishReason::Length);
assert_eq!(usage.completion_tokens, 5);
}
fn greedy_next_token_after(decoder: &Decoder, prompt_ids: &[usize]) -> usize {
let mut caches: Vec<KvCache> = decoder
.layers
.iter()
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect();
let logits = decoder
.forward_batch(prompt_ids, 0, &mut caches)
.pop()
.unwrap();
logits
.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.map(|(i, _)| i)
.unwrap()
}
#[test]
fn eos_token_stops_generation_before_max_tokens() {
let decoder = small_decoder();
let prompt_ids = vec![1usize, 2];
let prompt = String::from_utf8(prompt_ids.iter().map(|&b| b as u8).collect()).unwrap();
let eos = greedy_next_token_after(&decoder, &prompt_ids);
let (finish, _usage) = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::from_eos(Some(eos)),
None,
&prompt,
&greedy_params(50),
None,
None,
None,
None,
|_| {},
)
.unwrap();
assert_eq!(
finish,
FinishReason::Stop,
"generation must stop as soon as the greedy-chosen token matches eos_id, not run to max_tokens"
);
}
#[test]
fn a_turn_ender_that_is_not_the_metadata_eos_still_stops_generation() {
let decoder = small_decoder();
let prompt_ids = vec![1usize, 2];
let prompt = String::from_utf8(prompt_ids.iter().map(|&b| b as u8).collect()).unwrap();
let turn_ender = greedy_next_token_after(&decoder, &prompt_ids);
let never_sampled = (turn_ender + 1) % decoder.config.vocab_size;
let stop = StopTokens::from_eos(Some(never_sampled)).with_id(Some(turn_ender));
let (finish, usage) = generate(
&decoder,
&ServerTokenizer::Byte,
&stop,
None,
&prompt,
&greedy_params(50),
None,
None,
None,
None,
|_| {},
)
.unwrap();
assert_eq!(finish, FinishReason::Stop);
assert_eq!(
usage.completion_tokens, 0,
"the very first sampled token was the turn ender"
);
}
#[test]
fn a_stop_sequence_that_never_matches_does_not_drop_any_generated_content() {
let decoder = small_decoder();
let prompt = String::from_utf8(vec![1u8]).unwrap();
let mut baseline = String::new();
let (baseline_finish, _usage) = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(20),
None,
None,
None,
None,
|s| baseline.push_str(s),
)
.unwrap();
let mut with_unmatchable_stop = String::new();
let (stop_finish, _usage2) = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&GenerationParams {
max_tokens: 20,
sampling: SamplingParams::default(),
seed: 1,
stop: vec!["ZZ_NEVER_MATCHES_ZZ".to_string()],
stop_token_ids: Vec::new(),
json_object: false,
grammar: None,
cancel: None,
ignore_eos: false,
},
None,
None,
None,
None,
|s| with_unmatchable_stop.push_str(s),
)
.unwrap();
assert_eq!(baseline_finish, FinishReason::Length);
assert_eq!(stop_finish, FinishReason::Length);
assert_eq!(with_unmatchable_stop, baseline);
}
fn run_scripted(
script: &[usize],
render: impl Fn(usize) -> String,
params: &GenerationParams,
) -> (FinishReason, Vec<usize>, Vec<String>) {
run_scripted_with_stops(script, render, params, StopTokens::from_eos(None))
}
fn run_scripted_with_stops(
script: &[usize],
render: impl Fn(usize) -> String,
params: &GenerationParams,
stop_tokens: StopTokens,
) -> (FinishReason, Vec<usize>, Vec<String>) {
try_run_scripted_with_stops(script, render, params, stop_tokens)
.expect("an unconstrained script cannot fail to decode")
}
fn try_run_scripted_with_stops(
script: &[usize],
render: impl Fn(usize) -> String,
params: &GenerationParams,
stop_tokens: StopTokens,
) -> Result<(FinishReason, Vec<usize>, Vec<String>), DecodeError> {
let vocab = script.iter().copied().max().unwrap_or(0) + 2;
let logits_for = |id: usize| {
let mut v = vec![0.0f32; vocab];
v[id] = 10.0;
v
};
let mut next = 0usize;
let mut take = || {
let id = script
.get(next)
.copied()
.unwrap_or(script[script.len() - 1]);
next += 1;
id
};
let first = logits_for(take());
let mut chunks: Vec<String> = Vec::new();
let (finish, ids, _) = sample_until_stop(
first,
0,
&stop_tokens,
params,
|ids| ids.iter().copied().map(&render).collect::<String>(),
|_tok, _pos| logits_for(take()),
|chunk| chunks.push(chunk.to_string()),
&render,
)?;
Ok((finish, ids, chunks))
}
fn scripted_params(max_tokens: usize) -> GenerationParams {
GenerationParams {
max_tokens,
sampling: SamplingParams {
temperature: 0.0,
..SamplingParams::default()
},
seed: 1,
stop: Vec::new(),
stop_token_ids: Vec::new(),
json_object: false,
grammar: None,
cancel: None,
ignore_eos: false,
}
}
#[test]
fn a_grammar_constrains_the_private_decode_loop() {
let render = |id: usize| char::from(b'a' + id as u8).to_string();
let script = [2usize, 2, 2, 2];
let unconstrained = run_scripted(&script, render, &scripted_params(2));
assert_eq!(
unconstrained.1,
vec![2, 2],
"the model wants \"cc\", so a grammar forbidding it has work to do"
);
let mut params = scripted_params(2);
params.grammar = Some(std::sync::Arc::new(
ferrox_models::grammar::Grammar::from_str_with_root(r#"root ::= "ab""#, "root")
.expect("test grammar parses"),
));
let (finish, ids, chunks) = run_scripted(&script, render, ¶ms);
assert_eq!(
ids,
vec![0, 1],
"the grammar was not applied token by token"
);
assert_eq!(chunks.concat(), "ab");
assert_eq!(finish, FinishReason::Length);
}
#[test]
fn a_completed_grammar_ends_the_private_decode_loop() {
let render = |id: usize| char::from(b'a' + id as u8).to_string();
let mut params = scripted_params(8);
params.grammar = Some(std::sync::Arc::new(
ferrox_models::grammar::Grammar::from_str_with_root(r#"root ::= "ab""#, "root")
.expect("test grammar parses"),
));
let (finish, ids, chunks) = run_scripted(&[2usize; 8], render, ¶ms);
assert_eq!(ids, vec![0, 1]);
assert_eq!(chunks.concat(), "ab");
assert_eq!(finish, FinishReason::Stop);
}
#[test]
fn a_grammar_the_vocabulary_cannot_spell_fails_the_generation() {
let render = |id: usize| char::from(b'a' + id as u8).to_string();
let mut params = scripted_params(4);
params.grammar = Some(std::sync::Arc::new(
ferrox_models::grammar::Grammar::from_str_with_root(r#"root ::= "z""#, "root")
.expect("test grammar parses"),
));
let err =
try_run_scripted_with_stops(&[2usize; 4], render, ¶ms, StopTokens::from_eos(None))
.expect_err("no token in this vocabulary renders as \"z\"");
assert!(
matches!(err, DecodeError::GrammarConstraint { .. }),
"{err}"
);
}
#[test]
fn a_token_level_stop_ends_generation_and_never_reaches_the_output() {
let render = |id: usize| char::from(b'a' + id as u8).to_string();
let script = [0usize, 0, 1, 0, 0, 1];
let (finish, ids, chunks) = run_scripted(&script, render, &scripted_params(6));
assert_eq!(finish, FinishReason::Length);
assert_eq!(ids, script.to_vec());
assert_eq!(chunks.concat(), "aabaab");
let (finish, ids, chunks) = run_scripted(
&script,
render,
&GenerationParams {
stop_token_ids: vec![1],
..scripted_params(6)
},
);
assert_eq!(finish, FinishReason::Stop);
assert_eq!(ids, vec![0, 0], "the stop token is not part of the answer");
assert_eq!(
chunks.concat(),
"aa",
"the stop token must not be rendered into the output"
);
}
#[test]
fn a_stop_token_that_renders_as_nothing_is_still_a_stop() {
let render = |id: usize| {
if id == 1 {
String::new()
} else {
char::from(b'a' + id as u8).to_string()
}
};
let script = [0usize, 1, 0, 0];
let (finish, _, chunks) = run_scripted(
&script,
render,
&GenerationParams {
stop: vec!["<|end|>".to_string()],
..scripted_params(4)
},
);
assert_eq!(finish, FinishReason::Length);
assert_eq!(chunks.concat(), "aaa");
let (finish, ids, chunks) = run_scripted(
&script,
render,
&GenerationParams {
stop: vec!["<|end|>".to_string()],
stop_token_ids: vec![1],
..scripted_params(4)
},
);
assert_eq!(finish, FinishReason::Stop);
assert_eq!(ids, vec![0]);
assert_eq!(chunks.concat(), "a");
}
#[test]
fn nothing_that_becomes_part_of_the_stop_is_ever_emitted() {
let render = |id: usize| char::from(b'a' + id as u8).to_string();
let script = [0usize, 1, 0, 1, 2, 0];
let params = GenerationParams {
stop: vec!["abc".to_string()],
..scripted_params(6)
};
let (finish, _, chunks) = run_scripted(&script, render, ¶ms);
assert_eq!(
finish,
FinishReason::StopSequence("abc".to_string()),
"the reason names the stop that fired, not merely that one did"
);
assert_eq!(
chunks.concat(),
"ab",
"the answer is everything before the stop, and nothing after it"
);
let mut seen = String::new();
for chunk in &chunks {
seen.push_str(chunk);
assert!(
"ab".starts_with(&seen),
"the stream ran ahead of the answer: {seen:?} (chunks: {chunks:?})"
);
}
}
#[test]
fn a_disproved_partial_is_released_by_the_token_that_disproves_it() {
let render = |id: usize| char::from(b'a' + id as u8).to_string();
let script = [0usize, 1, 3, 0];
let (_, _, chunks) = run_scripted(
&script,
render,
&GenerationParams {
stop: vec!["abc".to_string()],
..scripted_params(4)
},
);
assert_eq!(chunks.concat(), "abda", "no output is lost");
assert_eq!(
chunks.first().map(String::as_str),
Some("abd"),
"the whole disproved partial goes out at once: {chunks:?}"
);
}
#[test]
fn a_stop_sequence_that_does_match_truncates_output_before_it() {
let decoder = small_decoder();
let prompt = String::from_utf8(vec![1u8]).unwrap();
let mut baseline = String::new();
generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(20),
None,
None,
None,
None,
|s| baseline.push_str(s),
)
.unwrap();
let Some((cut, _)) = baseline.char_indices().nth(1) else {
return;
};
let stop_str = baseline[cut..].to_string();
if stop_str.is_empty() {
return;
}
let mut truncated = String::new();
let (finish, _usage) = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&GenerationParams {
max_tokens: 20,
sampling: SamplingParams::default(),
seed: 1,
stop: vec![stop_str.clone()],
stop_token_ids: Vec::new(),
json_object: false,
grammar: None,
cancel: None,
ignore_eos: false,
},
None,
None,
None,
None,
|s| truncated.push_str(s),
)
.unwrap();
assert_eq!(finish, FinishReason::StopSequence(stop_str));
assert_eq!(truncated, baseline[..cut]);
}
#[test]
fn usage_reports_both_phases_and_a_time_to_first_token() {
let decoder = small_decoder();
let prompt = String::from_utf8(vec![1u8, 2, 3]).unwrap();
let (_finish, usage) = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(5),
None,
None,
None,
None,
|_| {},
)
.unwrap();
assert_eq!(usage.prompt_tokens, 3);
assert_eq!(usage.completion_tokens, 5);
let prefill = usage.prompt_eval_duration_ms.expect("prefill timed");
let decode = usage.generation_duration_ms.expect("decode timed");
let ttft = usage.time_to_first_token_ms.expect("first token timed");
assert!(ttft >= prefill, "ttft {ttft} < prefill {prefill}");
assert!(
ttft <= prefill + decode + 1.0,
"ttft {ttft} exceeds the whole request"
);
}
#[test]
fn cached_tokens_distinguishes_a_miss_from_an_absent_prefix_cache() {
let decoder = small_decoder();
let prompt = String::from_utf8(vec![1u8, 2, 3]).unwrap();
let (_f, no_cache) = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(2),
None,
None,
None,
None,
|_| {},
)
.unwrap();
assert_eq!(no_cache.cached_tokens, None, "no prefix cache configured");
let pc = Mutex::new(PrefixCache::new(4));
let (_f, miss) = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(2),
None,
None,
Some(&pc),
None,
|_| {},
)
.unwrap();
assert_eq!(miss.cached_tokens, Some(0), "cache consulted, missed");
let longer = String::from_utf8(vec![1u8, 2, 3, 9]).unwrap();
let (_f, hit) = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&longer,
&greedy_params(2),
None,
None,
Some(&pc),
None,
|_| {},
)
.unwrap();
assert_eq!(hit.cached_tokens, Some(3));
}
#[test]
fn prefix_cache_reuses_a_shared_prefix_and_produces_the_same_output_as_a_fresh_run() {
let decoder = small_decoder();
let prompt1 = String::from_utf8(vec![1u8, 2, 3]).unwrap();
let pc = Mutex::new(PrefixCache::new(4));
let mut out1 = String::new();
generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt1,
&greedy_params(5),
None,
None,
Some(&pc),
None,
|s| out1.push_str(s),
)
.unwrap();
assert_eq!(pc.lock().unwrap().stats().misses, 1);
let prompt2 = String::from_utf8(vec![1u8, 2, 3, 9, 9]).unwrap();
let mut out2_with_cache = String::new();
generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt2,
&greedy_params(5),
None,
None,
Some(&pc),
None,
|s| out2_with_cache.push_str(s),
)
.unwrap();
let stats = pc.lock().unwrap().stats();
assert_eq!(stats.hits, 1, "prompt2 must hit the stored prompt1 entry");
assert_eq!(stats.total_positions_reused, 3);
let mut out2_fresh = String::new();
generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt2,
&greedy_params(5),
None,
None,
None,
None,
|s| out2_fresh.push_str(s),
)
.unwrap();
assert_eq!(
out2_with_cache, out2_fresh,
"restoring from the prefix cache must produce identical output to processing the whole prompt from scratch"
);
}
#[test]
fn prefix_cache_exact_repeat_skips_prompt_processing_via_pending_logits() {
let decoder = small_decoder();
let prompt = String::from_utf8(vec![1u8, 2, 3]).unwrap();
let pc = Mutex::new(PrefixCache::new(4));
let mut out1 = String::new();
generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(5),
None,
None,
Some(&pc),
None,
|s| out1.push_str(s),
)
.unwrap();
let mut out2_with_cache = String::new();
generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(5),
None,
None,
Some(&pc),
None,
|s| out2_with_cache.push_str(s),
)
.unwrap();
let mut out2_fresh = String::new();
generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(5),
None,
None,
None,
None,
|s| out2_fresh.push_str(s),
)
.unwrap();
assert_eq!(out2_with_cache, out2_fresh);
}
#[test]
fn prefix_cache_is_not_consulted_when_a_kv_pool_is_configured() {
let decoder = small_decoder(); let prompt = String::from_utf8(vec![1u8, 2, 3]).unwrap();
let pc = Mutex::new(PrefixCache::new(4));
let pool = Arc::new(Mutex::new(KvBlockPool::new(64, 2)));
let config = pool_config(pool, Duration::ZERO);
generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(5),
Some(&config),
None,
Some(&pc),
None,
|_| {},
)
.unwrap();
generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(5),
Some(&config),
None,
Some(&pc),
None,
|_| {},
)
.unwrap();
let stats = pc.lock().unwrap().stats();
assert_eq!(
stats.hits + stats.misses,
0,
"prefix cache must never be consulted while a KV pool is configured"
);
}
fn pool_config(pool: Arc<Mutex<KvBlockPool>>, queue_wait: Duration) -> KvPoolConfig {
KvPoolConfig { pool, queue_wait }
}
#[test]
fn published_prefixes_are_reclaimed_under_pressure() {
let decoder = small_decoder();
let block_size = 4;
let config = paged_config_with_radix(&decoder, block_size, 24, true);
let free_before = config.store.free_groups();
assert!(free_before > 0, "the fixture must start with free pages");
for round in 0..40u32 {
let tokens: Vec<usize> = (0..12).map(|i| (round * 100 + i) as usize).collect();
let mut lease = acquire_paged_caches(&decoder, &config, &tokens, tokens.len() + 8)
.expect("admission must keep succeeding once the tree can be evicted");
publish_to_radix(&mut lease, &tokens, block_size);
drop(lease);
}
let tree_holds = {
let tree = config
.radix
.as_ref()
.expect("configured with a tree")
.lock()
.unwrap_or_else(|p| p.into_inner());
tree.total_size().div_ceil(block_size)
};
assert_eq!(
config.store.free_groups() + tree_holds,
free_before,
"every page must be free or accounted for in the tree: \
free={} tree={} started={free_before}",
config.store.free_groups(),
tree_holds
);
}
fn paged_config(decoder: &Decoder, block_size: usize, blocks: usize) -> PagedKvConfig {
paged_config_with_radix(decoder, block_size, blocks, false)
}
fn paged_config_with_radix(
decoder: &Decoder,
block_size: usize,
blocks: usize,
share_prefixes: bool,
) -> PagedKvConfig {
PagedKvConfig {
store: Arc::new(SharedPagedKv::new(
decoder.layers.len(),
block_size,
blocks,
decoder.config.n_kv_heads,
decoder.config.head_dim,
)),
queue_wait: Duration::ZERO,
radix: share_prefixes.then(|| {
Arc::new(Mutex::new(crate::policy::radix::RadixCache::new(
block_size,
)))
}),
anchor_token: None,
slide_interval: crate::policy::pool_budget::DEFAULT_SWA_EVICTION_INTERVAL,
}
}
#[test]
fn a_paged_request_generates_the_same_text_as_a_contiguous_one() {
let decoder = small_decoder();
let prompt = String::from_utf8(vec![1u8, 2, 3]).unwrap();
let mut contiguous = String::new();
let (finish_a, usage_a) = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(6),
None,
None,
None,
None,
|s| contiguous.push_str(s),
)
.unwrap();
let config = paged_config(&decoder, 4, 64);
let mut paged = String::new();
let (finish_b, usage_b) = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(6),
None,
Some(&config),
None,
None,
|s| paged.push_str(s),
)
.unwrap();
assert_eq!(paged, contiguous, "paged serving changed the answer");
assert_eq!(finish_b, finish_a);
assert_eq!(usage_b.completion_tokens, usage_a.completion_tokens);
assert_eq!(usage_b.prompt_tokens, usage_a.prompt_tokens);
}
#[test]
fn a_finished_paged_request_returns_every_page_it_held() {
let decoder = small_decoder();
let prompt = String::from_utf8(vec![1u8, 2]).unwrap();
let config = paged_config(&decoder, 4, 64);
let before: Vec<usize> = (0..decoder.layers.len())
.map(|l| config.store.free_blocks(l))
.collect();
for _ in 0..3 {
let mut out = String::new();
generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(5),
None,
Some(&config),
None,
None,
|s| out.push_str(s),
)
.unwrap();
}
for (l, expected) in before.iter().enumerate() {
assert_eq!(
config.store.free_blocks(l),
*expected,
"layer {l} leaked pages across repeated requests"
);
}
}
fn windowed_decoder(window: usize) -> Decoder {
let mut cfg = test_dense_fixture();
cfg.sliding_window = Some(window);
cfg.swa_pattern = None;
Decoder::new_random_small(cfg, 2, 256)
}
fn run_to_completion(max_tokens: usize) -> GenerationParams {
GenerationParams {
ignore_eos: true,
..greedy_params(max_tokens)
}
}
#[test]
fn a_window_prices_a_request_at_its_prompt_plus_a_bound() {
let block_size = 4;
let policy = WindowPolicy::new(8, block_size);
let bound = slide_hold_bound(8, &policy);
assert_eq!(bound, 2 * (8 + SWA_RETAIN_GAP) + 128 + 2 * block_size);
assert_eq!(paged_hold_positions(10_000, 3, block_size, None), 10_000);
assert_eq!(paged_groups_needed(10_000, 3, block_size, None), 2_500);
let held = paged_hold_positions(10_000, 3, block_size, Some(&policy));
assert_eq!(held, 3 + bound + block_size);
assert_eq!(
paged_hold_positions(100_000, 3, block_size, Some(&policy)),
held,
"a longer generation must not cost more pages"
);
assert!(paged_hold_positions(10_000, 900, block_size, Some(&policy)) > held);
assert_eq!(paged_hold_positions(20, 3, block_size, Some(&policy)), 20);
assert_eq!(
paged_groups_needed(20, 3, block_size, Some(&policy)),
paged_groups_needed(20, 3, block_size, None)
);
}
#[test]
fn a_window_model_runs_on_a_store_too_small_for_its_whole_context() {
let window = 8;
let windowed = windowed_decoder(window);
assert_eq!(
windowed.config.uniform_sliding_window(),
Some(window),
"the fixture must be uniformly windowed or this test proves nothing"
);
let full = small_decoder();
assert_eq!(full.config.uniform_sliding_window(), None);
let block_size = 4;
let prompt = String::from_utf8(vec![1u8, 2, 3]).unwrap();
let max_tokens = 400;
let blocks = 60;
let params = run_to_completion(max_tokens);
let win_config = paged_config(&windowed, block_size, blocks);
let mut out = String::new();
let (_, usage) = generate(
&windowed,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
¶ms,
None,
Some(&win_config),
None,
None,
|s| out.push_str(s),
)
.expect("a window model must fit a store sized for its window");
assert_eq!(usage.completion_tokens, max_tokens);
let mut alternating_cfg = test_dense_fixture();
alternating_cfg.sliding_window = Some(window);
alternating_cfg.swa_pattern = Some(2);
let alternating = Decoder::new_random_small(alternating_cfg, 2, 256);
assert_eq!(alternating.config.kv_block_window(), Some(window));
assert_eq!(alternating.config.uniform_sliding_window(), None);
for (name, model) in [("full attention", &full), ("alternating", &alternating)] {
let config = paged_config(model, block_size, blocks);
let err = generate(
model,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
¶ms,
None,
Some(&config),
None,
None,
|_| {},
)
.unwrap_err();
assert!(matches!(err, DecodeError::KvPoolExhausted), "{name}: {err}");
}
}
#[test]
fn a_sliding_paged_request_says_the_same_thing_as_a_contiguous_one() {
let decoder = windowed_decoder(8);
let prompt = String::from_utf8(vec![1u8, 2, 3]).unwrap();
let params = run_to_completion(400);
let mut contiguous = String::new();
generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
¶ms,
None,
None,
None,
None,
|s| contiguous.push_str(s),
)
.unwrap();
let config = paged_config(&decoder, 4, 60);
let mut paged = String::new();
generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
¶ms,
None,
Some(&config),
None,
None,
|s| paged.push_str(s),
)
.unwrap();
assert_eq!(paged, contiguous, "the window slide changed the answer");
}
#[test]
fn a_long_windowed_generation_stops_taking_pages_from_the_store() {
let decoder = windowed_decoder(8);
let block_size = 4;
let config = paged_config(&decoder, block_size, 200);
let tokens: Vec<usize> = vec![1, 2, 3];
let mut lease = acquire_paged_caches(&decoder, &config, &tokens, tokens.len() + 4_000)
.expect("the store holds a window's worth");
let after_admission = config.store.free_groups();
let held = lease.groups.len();
for pos in tokens.len()..tokens.len() + 4_000 {
lease.before_step(pos);
assert_eq!(
config.store.free_groups(),
after_admission,
"position {pos} took a page from the store instead of recycling"
);
}
assert!(
lease.window.as_ref().unwrap().released > 0,
"4000 positions at a window of 8 must have slid"
);
assert_eq!(
lease.groups.iter().flatten().count() + lease.spare.len(),
held
);
}
#[test]
fn a_tool_call_anchor_holds_the_window_back_and_then_lets_go() {
let anchor_token = 77;
let block_size = 4;
let decoder = windowed_decoder(8);
let tokens: Vec<usize> = vec![1, 2, 3];
let released_at = |armed: bool, at: &[usize], check: usize| -> usize {
let mut config = paged_config(&decoder, block_size, 400);
config.slide_interval = 4;
config.anchor_token = armed.then_some(anchor_token as u32);
let mut lease = acquire_paged_caches(&decoder, &config, &tokens, tokens.len() + 1_000)
.expect("the store is large enough");
for pos in tokens.len()..check {
if at.contains(&(pos + 1)) {
lease.observe_sampled(anchor_token, pos + 1, false);
}
lease.before_step(pos);
}
lease.window.as_ref().unwrap().released
};
let first = 200;
assert!(
released_at(true, &[first], 208) < released_at(false, &[first], 208),
"the anchor did not hold the window back"
);
assert_eq!(
released_at(true, &[first], 600),
released_at(false, &[first], 600),
"the anchor was never dropped, so the hold is unbounded"
);
let second = 600;
assert!(
released_at(true, &[first, second], 610) < released_at(false, &[], 610),
"a dropped anchor left the request unable to take another"
);
}
#[test]
fn a_slid_request_returns_the_pages_it_recycled_as_well_as_the_ones_it_held() {
let decoder = windowed_decoder(8);
let prompt = String::from_utf8(vec![1u8, 2, 3]).unwrap();
let config = paged_config(&decoder, 4, 60);
let before: Vec<usize> = (0..decoder.layers.len())
.map(|l| config.store.free_blocks(l))
.collect();
for run in 0..3 {
let mut out = String::new();
generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&run_to_completion(400),
None,
Some(&config),
None,
None,
|s| out.push_str(s),
)
.unwrap_or_else(|e| panic!("run {run} was refused: {e}"));
}
for (l, expected) in before.iter().enumerate() {
assert_eq!(
config.store.free_blocks(l),
*expected,
"layer {l} kept the pages a slid request recycled"
);
}
}
#[test]
fn every_page_the_kernel_still_reads_is_one_the_lease_still_owns() {
let window = 8;
let block_size = 4;
let decoder = windowed_decoder(window);
let config = paged_config(&decoder, block_size, 200);
let tokens: Vec<usize> = vec![1, 2, 3];
let mut lease = acquire_paged_caches(&decoder, &config, &tokens, tokens.len() + 4_000)
.expect("the store holds a window's worth");
for pos in tokens.len()..tokens.len() + 4_000 {
lease.before_step(pos);
let seq_len = pos + 1;
let first = seq_len.saturating_sub(window) / block_size;
for i in first..=pos / block_size {
assert!(
lease.groups[i].is_some(),
"at position {pos} the window reaches page {i}, which was recycled"
);
}
}
assert!(
lease.window.as_ref().unwrap().released > 0,
"nothing was recycled, so this proved nothing"
);
}
#[test]
fn a_slid_request_never_recycles_the_prefix_the_tree_owns() {
let decoder = windowed_decoder(8);
let block_size = 4;
let config = paged_config_with_radix(&decoder, block_size, 400, true);
let shared: Vec<u8> = (1u8..=16).collect();
let prompt = String::from_utf8(shared.clone()).unwrap();
generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(2),
None,
Some(&config),
None,
None,
|_| {},
)
.unwrap();
let tokens: Vec<usize> = shared.iter().map(|&b| b as usize).collect();
let mut lease = acquire_paged_caches(&decoder, &config, &tokens, tokens.len() + 3_000)
.expect("the store is large enough");
let locked = lease.adopted_positions(block_size);
assert!(locked > 0, "this test needs a real prefix match");
for pos in tokens.len()..tokens.len() + 3_000 {
lease.before_step(pos);
}
let released = lease.window.as_ref().unwrap().released;
assert!(released >= locked, "the slide must have passed the prefix");
for i in 0..locked / block_size {
assert!(
lease.groups[i].is_some(),
"page {i} of the shared prefix was recycled out from under the tree"
);
}
}
#[test]
fn a_slid_sequence_is_not_published_to_the_tree() {
let decoder = windowed_decoder(8);
let block_size = 4;
let prompt = String::from_utf8((1u8..=16).collect::<Vec<u8>>()).unwrap();
let cached_after = |first: GenerationParams| -> usize {
let config = paged_config_with_radix(&decoder, block_size, 400, true);
let mut out = String::new();
generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&first,
None,
Some(&config),
None,
None,
|s| out.push_str(s),
)
.unwrap();
let mut probe = String::new();
let (_, usage) = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(2),
None,
Some(&config),
None,
None,
|s| probe.push_str(s),
)
.unwrap();
usage.cached_tokens.unwrap_or(0)
};
assert!(
cached_after(greedy_params(2)) > 0,
"a request that never slid must publish its pages"
);
assert_eq!(
cached_after(run_to_completion(600)),
0,
"a slid sequence published a prefix it no longer holds"
);
}
#[test]
fn two_prompts_sharing_a_prefix_hold_one_copy_of_it() {
let decoder = small_decoder();
let shared: Vec<u8> = (1u8..=8).collect();
let mut a = shared.clone();
a.push(40);
let mut b = shared.clone();
b.push(50);
let prompt_a = String::from_utf8(a).unwrap();
let prompt_b = String::from_utf8(b).unwrap();
let run = |config: &PagedKvConfig, prompt: &str| -> String {
let mut out = String::new();
generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
prompt,
&greedy_params(2),
None,
Some(config),
None,
None,
|s| out.push_str(s),
)
.unwrap();
out
};
let cfg = paged_config_with_radix(&decoder, 4, 64, true);
let start = cfg.store.free_groups();
let text_a = run(&cfg, &prompt_a);
let after_a = cfg.store.free_groups();
let text_b = run(&cfg, &prompt_b);
let after_b = cfg.store.free_groups();
let cost_a = start - after_a;
let cost_b = after_a - after_b;
assert!(cost_a > 0, "the first request must publish something");
assert!(
cost_b < cost_a,
"the second request shares A's prefix and must cost less \
(first {cost_a} groups, second {cost_b})"
);
let cfg2 = paged_config_with_radix(&decoder, 4, 64, true);
let mut sink = String::new();
let (_f, first_usage) = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt_a,
&greedy_params(2),
None,
Some(&cfg2),
None,
None,
|s| sink.push_str(s),
)
.unwrap();
assert_eq!(
first_usage.cached_tokens,
Some(0),
"a cold tree reuses nothing"
);
let (_f, second_usage) = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt_b,
&greedy_params(2),
None,
Some(&cfg2),
None,
None,
|s| sink.push_str(s),
)
.unwrap();
let reused = second_usage.cached_tokens.expect("a tree is configured");
assert!(
reused >= 8,
"the 8-token shared prefix must be reported as reused, got {reused}"
);
let plain = paged_config(&decoder, 4, 64);
assert_eq!(text_a, run(&plain, &prompt_a), "prompt A changed");
assert_eq!(text_b, run(&plain, &prompt_b), "prompt B changed");
}
#[test]
fn many_requests_off_one_prefix_neither_leak_nor_free_the_trees_pages() {
let decoder = small_decoder();
let config = paged_config_with_radix(&decoder, 4, 64, true);
let shared: Vec<u8> = (1u8..=8).collect();
let mut lows = Vec::new();
for suffix in 0..6u8 {
let mut p = shared.clone();
p.push(60 + suffix);
let prompt = String::from_utf8(p).unwrap();
let mut out = String::new();
generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(2),
None,
Some(&config),
None,
None,
|s| out.push_str(s),
)
.expect("the store is sized for many of these");
lows.push(config.store.free_groups());
}
assert_eq!(
lows[lows.len() - 1],
lows[lows.len() - 2],
"steady state expected once the shared prefix is published; \
free groups per run were {lows:?}"
);
assert!(
lows[lows.len() - 1] > 0,
"the store must not have been consumed: {lows:?}"
);
}
#[test]
fn a_paged_request_too_big_for_the_store_is_refused_before_emitting_anything() {
let decoder = small_decoder();
let prompt = String::from_utf8(vec![1u8, 2, 3, 4]).unwrap();
let config = paged_config(&decoder, 2, 1);
let result = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(8),
None,
Some(&config),
None,
None,
|_| panic!("a refused request must not emit"),
);
assert!(
matches!(result, Err(DecodeError::KvPoolExhausted)),
"expected a typed refusal, got {result:?}"
);
for l in 0..decoder.layers.len() {
assert_eq!(
config.store.free_blocks(l),
1,
"layer {l} must keep every page after a refusal"
);
}
}
#[test]
fn generate_succeeds_with_a_pool_that_has_enough_blocks() {
let decoder = small_decoder(); let prompt = String::from_utf8(vec![1u8, 2]).unwrap();
let pool = Arc::new(Mutex::new(KvBlockPool::new(64, 2)));
let config = pool_config(pool.clone(), Duration::ZERO);
let mut out = String::new();
let (finish, _usage) = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(5),
Some(&config),
None,
None,
None,
|s| out.push_str(s),
)
.unwrap();
assert_eq!(finish, FinishReason::Length);
assert_eq!(
pool.lock().unwrap().free_blocks(),
2,
"every acquired block must be released once the request finishes"
);
}
#[test]
fn generate_reserves_enough_blocks_up_front_for_a_sequence_spanning_multiple_blocks() {
let decoder = small_decoder(); let prompt = String::from_utf8(vec![1u8, 2]).unwrap(); let max_tokens = 10;
let block_size = 2;
let pool = Arc::new(Mutex::new(KvBlockPool::new(block_size, 12)));
let config = pool_config(pool.clone(), Duration::ZERO);
let (finish, _usage) = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(max_tokens),
Some(&config),
None,
None,
None,
|_| {},
)
.unwrap();
assert_eq!(finish, FinishReason::Length);
assert_eq!(pool.lock().unwrap().free_blocks(), 12);
}
#[test]
fn generate_fails_at_admission_not_mid_decode_when_the_pool_cannot_cover_the_worst_case() {
let decoder = small_decoder(); let prompt = String::from_utf8(vec![1u8, 2]).unwrap();
let max_tokens = 10;
let block_size = 2;
let pool = Arc::new(Mutex::new(KvBlockPool::new(block_size, 11)));
let config = pool_config(pool.clone(), Duration::ZERO);
let result = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(max_tokens),
Some(&config),
None,
None,
None,
|_| {},
);
let err = result.expect_err("11 blocks cannot cover a 12-block worst case");
assert!(
matches!(
&err,
DecodeError::KvBudgetExceeded { binding, positions, .. }
if *binding == ferrox_models::Ceiling::DeviceMemory.code()
&& *positions == 12
),
"expected an immovable device-memory refusal, got {err:?}"
);
assert_eq!(
err.retry_after_secs(),
None,
"no wait frees blocks that do not exist"
);
assert_eq!(
pool.lock().unwrap().free_blocks(),
11,
"a rejected request must leave the pool exactly as it found it"
);
}
#[test]
fn generate_rejects_the_request_without_leaking_blocks_when_the_pool_is_too_small() {
let decoder = small_decoder(); let prompt = String::from_utf8(vec![1u8, 2]).unwrap();
let pool = Arc::new(Mutex::new(KvBlockPool::new(64, 1)));
let config = pool_config(pool.clone(), Duration::ZERO);
let result = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(5),
Some(&config),
None,
None,
None,
|_| {},
);
let err = result.expect_err("one block cannot hold two layers' caches");
assert!(
matches!(
&err,
DecodeError::KvBudgetExceeded { binding, .. }
if *binding == ferrox_models::Ceiling::DeviceMemory.code()
),
"expected an immovable device-memory refusal, got {err:?}"
);
assert_eq!(
pool.lock().unwrap().free_blocks(),
1,
"a rejected request must leave the pool exactly as it found it"
);
}
#[test]
fn generate_releases_blocks_so_back_to_back_requests_do_not_starve_the_pool() {
let decoder = small_decoder(); let prompt = String::from_utf8(vec![1u8, 2]).unwrap();
let pool = Arc::new(Mutex::new(KvBlockPool::new(64, 2)));
let config = pool_config(pool.clone(), Duration::ZERO);
for _ in 0..3 {
let (finish, _usage) = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(5),
Some(&config),
None,
None,
None,
|_| {},
)
.unwrap();
assert_eq!(finish, FinishReason::Length);
}
assert_eq!(pool.lock().unwrap().free_blocks(), 2);
}
#[test]
fn generate_with_zero_queue_wait_rejects_immediately() {
let decoder = small_decoder(); let prompt = String::from_utf8(vec![1u8, 2]).unwrap();
let pool = Arc::new(Mutex::new(KvBlockPool::new(64, 2)));
let holder_pool = pool.clone();
let holder = std::thread::spawn(move || {
let mut held = KvCache::with_pool(1, 1, holder_pool, 0).unwrap();
held.push(&[0.0], &[0.0]).unwrap(); std::thread::sleep(Duration::from_millis(200));
drop(held);
});
std::thread::sleep(Duration::from_millis(15));
let config = pool_config(pool, Duration::ZERO);
let started = Instant::now();
let result = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(5),
Some(&config),
None,
None,
None,
|_| {},
);
assert!(
matches!(result, Err(DecodeError::KvPoolExhausted)),
"a pool that could serve this request once its holder lets go is momentary \
exhaustion, which is retryable"
);
assert!(
started.elapsed() < Duration::from_millis(50),
"queue_wait=0 must reject on the first attempt, not retry: took {:?}",
started.elapsed()
);
holder.join().unwrap();
}
#[test]
fn a_prompt_past_the_context_ceiling_is_refused_before_any_kv_is_acquired() {
let decoder = small_decoder();
let prompt = String::from_utf8(vec![1u8, 2, 3, 4, 5]).unwrap();
let pool = Arc::new(Mutex::new(KvBlockPool::new(64, 64)));
let config = pool_config(pool.clone(), Duration::ZERO);
let shape = KvShape::from_config(&decoder.config, KvElem::F32, 1);
let ceiling = ContextCeiling::new(Some(4), shape);
let err = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(5),
Some(&config),
None,
None,
Some(&ceiling),
|_| panic!("no token may be emitted by a refused request"),
)
.expect_err("a 5-token prompt must not be admitted under a 4-position ceiling");
match &err {
DecodeError::KvBudgetExceeded {
binding,
positions,
positions_limit,
detail,
..
} => {
assert_eq!(*binding, ferrox_models::Ceiling::ContextLength.code());
assert_eq!(*positions, 5);
assert_eq!(*positions_limit, 4);
assert_eq!(detail, "prompt is too long: 5 tokens > 4 maximum");
}
other => panic!("expected a context-length refusal, got {other:?}"),
}
assert_eq!(err.retry_after_secs(), None, "a 400, not a retryable 503");
assert_eq!(
pool.lock().unwrap().free_blocks(),
64,
"the refusal must land before any block is taken"
);
assert_eq!(ceiling.refused(), 1);
}
#[test]
fn a_prompt_that_fits_is_served_with_its_budget_clamped_rather_than_refused() {
let decoder = small_decoder();
let prompt = String::from_utf8(vec![1u8, 2]).unwrap();
let shape = KvShape::from_config(&decoder.config, KvElem::F32, 1);
let ceiling = ContextCeiling::new(Some(4), shape);
let mut emitted = 0usize;
let (finish, usage) = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(5),
None,
None,
None,
Some(&ceiling),
|_| emitted += 1,
)
.expect("a prompt that fits must be served");
assert_eq!(usage.completion_tokens, 2, "clamped to the room left");
assert_eq!(finish, FinishReason::Length);
assert_eq!(ceiling.refused(), 0, "a clamp is not a refusal");
assert!(emitted > 0);
}
#[test]
fn a_request_inside_the_ceiling_is_admitted_unchanged() {
let decoder = small_decoder();
let prompt = String::from_utf8(vec![1u8, 2]).unwrap();
let shape = KvShape::from_config(&decoder.config, KvElem::F32, 1);
let ceiling = ContextCeiling::new(Some(7), shape);
let mut with = String::new();
let (finish, _usage) = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(5),
None,
None,
None,
Some(&ceiling),
|s| with.push_str(s),
)
.expect("7 positions fits a 7-position ceiling exactly");
assert_eq!(finish, FinishReason::Length);
let mut without = String::new();
generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(5),
None,
None,
None,
None,
|s| without.push_str(s),
)
.unwrap();
assert_eq!(with, without, "an unbinding ceiling must change nothing");
assert_eq!(ceiling.refused(), 0);
}
#[test]
fn generate_with_a_queue_wait_succeeds_once_another_holder_releases_its_blocks() {
let decoder = small_decoder(); let prompt = String::from_utf8(vec![1u8, 2]).unwrap();
let pool = Arc::new(Mutex::new(KvBlockPool::new(64, 2)));
let holder_pool = pool.clone();
let holder = std::thread::spawn(move || {
let mut held = KvCache::with_pool(1, 1, holder_pool.clone(), 0).unwrap();
held.push(&[0.0], &[0.0]).unwrap(); std::thread::sleep(Duration::from_millis(80));
drop(held); });
std::thread::sleep(Duration::from_millis(15));
let config = pool_config(pool.clone(), Duration::from_millis(500));
let (finish, _usage) = generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(5),
Some(&config),
None,
None,
None,
|_| {},
)
.unwrap();
assert_eq!(
finish,
FinishReason::Length,
"a sufficiently long queue_wait must let the request succeed once the holder releases"
);
holder.join().unwrap();
assert_eq!(pool.lock().unwrap().free_blocks(), 2);
}
#[test]
fn earliest_stop_match_finds_the_leftmost_match_across_multiple_stops() {
assert_eq!(
earliest_stop_match("hello world", &["world".to_string(), "hello".to_string()]),
Some((0, "hello")),
"the leftmost match wins, not the caller's first entry"
);
assert_eq!(
earliest_stop_match("hello world", &["nope".to_string()]),
None
);
}
#[test]
fn ignore_eos_runs_a_request_out_to_its_full_budget() {
let render = |id: usize| char::from(b'a' + id as u8).to_string();
let script = [0usize, 1, 7, 2, 3, 4];
let eos = StopTokens::from_eos(Some(7));
let stops_early =
run_scripted_with_stops(&script, render, &scripted_params(6), eos.clone());
assert_eq!(stops_early.0, FinishReason::Stop);
assert_eq!(stops_early.1.len(), 2, "the model ended its own turn");
let runs_on = run_scripted_with_stops(
&script,
render,
&GenerationParams {
ignore_eos: true,
..scripted_params(6)
},
eos,
);
assert_eq!(runs_on.0, FinishReason::Length);
assert_eq!(
runs_on.1.len(),
6,
"exactly the budget, which is the whole point"
);
}
#[test]
fn ignore_eos_does_not_withdraw_the_callers_own_stop() {
let render = |id: usize| char::from(b'a' + id as u8).to_string();
let script = [0usize, 1, 2, 3, 4, 5];
let (finish, ids, _) = run_scripted_with_stops(
&script,
render,
&GenerationParams {
ignore_eos: true,
stop_token_ids: vec![2],
..scripted_params(6)
},
StopTokens::from_eos(Some(7)),
);
assert_eq!(finish, FinishReason::Stop);
assert_eq!(ids.len(), 2, "the caller's stop token still ends it");
}
}