use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use ferrox_core::cache::{KvBlockPool, KvCache, KvPoolExhausted as CacheKvPoolExhausted};
use ferrox_models::sampling::{Sampler, SamplingParams};
use ferrox_models::tokenizer::{prepend_bos, StopTokens};
use ferrox_models::{Ceiling, Decoder, Engine, KvElem, KvShape, PrefixCache, TextTokenizer};
use crate::budget::ContextCeiling;
use crate::json_mode::mask_logits_for_json;
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,
},
}
impl DecodeError {
pub fn retry_after_secs(&self) -> Option<u64> {
match self {
DecodeError::TokenOutOfVocab { .. } => None,
DecodeError::KvBudgetExceeded { .. } => None,
DecodeError::KvPoolExhausted | DecodeError::QueueFull { .. } => Some(1),
}
}
}
#[derive(Clone)]
pub struct KvPoolConfig {
pub pool: Arc<Mutex<KvBlockPool>>,
pub queue_wait: Duration,
}
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, Copy, PartialEq, Eq, serde::Serialize)]
#[serde(rename_all = "snake_case")]
pub enum FinishReason {
Stop,
Length,
Cancelled,
}
impl FinishReason {
pub fn as_str(&self) -> &'static str {
match self {
FinishReason::Stop => "stop",
FinishReason::Length => "length",
FinishReason::Cancelled => "cancelled",
}
}
}
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 cancel: Option<crate::cancel::CancelToken>,
}
impl GenerationParams {
pub(crate) fn is_cancelled(&self) -> bool {
self.cancel.as_ref().is_some_and(|c| c.is_cancelled())
}
}
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],
) -> Vec<f32> {
if let Some(chunk) = chunked_prefill_tokens() {
let mut pos = start_pos;
let mut last = Vec::new();
for part in tokens.chunks(chunk) {
last = decoder.forward_batch_last(part, pos, caches);
pos += part.len();
}
last
} else {
decoder.forward_batch_last(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>,
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 params.sampling.temperature <= 0.0 && ferrox_models::metal_greedy_gpu_enabled() {
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 max_seq_len = tokens.len() + params.max_tokens;
if let Some(ceiling) = ceiling {
if let Some(err) = ceiling.refusal(max_seq_len) {
return Err(err);
}
}
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 caches: Vec<KvCache>;
if let Some(m) = restored {
caches = m
.kv_caches
.expect("matched_len > 0 always carries kv_caches");
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, &mut caches);
pos += 1;
}
} else {
pos = m.matched_len;
let mut l = Vec::new();
for &tok in suffix {
l = decoder.forward_token(tok, pos, &mut caches);
pos += 1;
}
logits = l;
}
} else {
caches = match kv_pool {
Some(config) => acquire_pooled_caches(decoder, config, max_seq_len)
.map_err(|_| DecodeError::KvPoolExhausted)?,
None => 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 = decoder.forward_token(0, pos, &mut caches);
pos += 1;
l
} else {
pos = tokens.len();
forward_prompt_batch(decoder, &tokens, 0, &mut caches)
};
}
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());
}
let l = decoder.forward_token(next, pos, &mut caches);
#[cfg(feature = "metal")]
if kv_offload {
decoder.sync_metal_attn_kv_to_host(&mut caches);
}
l
},
&mut emit,
if params.json_object {
Some(&decode_token as &dyn Fn(usize) -> String)
} else {
None
},
);
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 kv_pool.is_none() {
if let Some(pc) = prefix_cache {
if logits.len() == vocab_size {
#[cfg(feature = "metal")]
decoder.sync_metal_attn_kv_to_host(&mut caches);
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: Option<&dyn Fn(usize) -> String>,
) -> (FinishReason, Vec<usize>, Vec<f32>) {
let mut matcher = crate::stop::StopMatcher::new(¶ms.stop, ¶ms.stop_token_ids);
let mut sampler = Sampler::new(params.seed);
let mut generated_ids: Vec<usize> = Vec::with_capacity(params.max_tokens);
let mut finish = FinishReason::Length;
for _ in 0..params.max_tokens {
if params.is_cancelled() {
finish = FinishReason::Cancelled;
break;
}
let next = if params.json_object {
if let Some(decode_token) = decode_token {
let mut mask_fn = |scores: &mut [f32]| {
mask_logits_for_json(scores, decode_token);
};
sampler.sample_with_mask(
&logits,
¶ms.sampling,
&generated_ids,
Some(&mut mask_fn),
)
} else {
sampler.sample(&logits, ¶ms.sampling, &generated_ids)
}
} else {
sampler.sample(&logits, ¶ms.sampling, &generated_ids)
};
if 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) => {
if !text.is_empty() {
emit(&text);
}
finish = FinishReason::Stop;
break;
}
}
}
let tail = matcher.flush();
if !tail.is_empty() {
emit(&tail);
}
(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,
None,
);
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(text: &str, stops: &[String]) -> Option<usize> {
stops
.iter()
.filter(|s| !s.is_empty())
.filter_map(|s| text.find(s.as_str()))
.min()
}
pub(crate) fn floor_char_boundary(s: &str, idx: usize) -> usize {
let mut i = idx.min(s.len());
while i > 0 && !s.is_char_boundary(i) {
i -= 1;
}
i
}
#[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,
cancel: None,
}
}
#[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,
|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,
|_| {},
);
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,
|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,
|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,
|_| {},
)
.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,
|_| {},
)
.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,
|_| {},
)
.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,
|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,
cancel: 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>) {
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,
&StopTokens::from_eos(None),
params,
|ids| ids.iter().copied().map(&render).collect::<String>(),
|_tok, _pos| logits_for(take()),
|chunk| chunks.push(chunk.to_string()),
None,
);
(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,
cancel: None,
}
}
#[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::Stop);
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,
|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],
stop_token_ids: Vec::new(),
json_object: false,
cancel: None,
},
None,
None,
None,
|s| truncated.push_str(s),
)
.unwrap();
assert_eq!(finish, FinishReason::Stop);
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,
|_| {},
)
.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,
|_| {},
)
.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,
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,
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,
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,
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,
|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,
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,
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,
|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),
Some(&pc),
None,
|_| {},
)
.unwrap();
generate(
&decoder,
&ServerTokenizer::Byte,
&StopTokens::default(),
None,
&prompt,
&greedy_params(5),
Some(&config),
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 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,
|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,
|_| {},
)
.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,
|_| {},
);
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,
|_| {},
);
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,
|_| {},
)
.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,
|_| {},
);
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_request_past_the_context_ceiling_is_refused_before_any_kv_is_acquired() {
let decoder = small_decoder();
let prompt = String::from_utf8(vec![1u8, 2]).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,
Some(&ceiling),
|_| panic!("no token may be emitted by a refused request"),
)
.expect_err("7 positions must not be admitted under a 4-position ceiling");
match &err {
DecodeError::KvBudgetExceeded {
binding,
positions,
positions_limit,
estimated_bytes,
limit_bytes,
..
} => {
assert_eq!(*binding, ferrox_models::Ceiling::ContextLength.code());
assert_eq!(*positions, 7);
assert_eq!(*positions_limit, 4);
assert_eq!(*estimated_bytes, shape.kv_bytes_for_tokens(7));
assert_eq!(*limit_bytes, shape.kv_bytes_for_tokens(4));
}
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_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,
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,
|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,
|_| {},
)
.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)
);
assert_eq!(
earliest_stop_match("hello world", &["nope".to_string()]),
None
);
}
}