use anyhow::{anyhow, ensure, Result};
use mlx_native::{DType, MlxBuffer, MlxDevice};
pub struct DrafterKvCache {
pub k_buf: MlxBuffer,
pub v_buf: MlxBuffer,
pub num_kv_heads: usize,
pub capacity: usize,
pub head_dim: usize,
len: usize,
}
impl DrafterKvCache {
pub fn new(
device: &MlxDevice,
num_kv_heads: usize,
capacity: usize,
head_dim: usize,
) -> Result<Self> {
ensure!(num_kv_heads > 0, "DrafterKvCache: num_kv_heads must be > 0");
ensure!(capacity > 0, "DrafterKvCache: capacity must be > 0");
ensure!(head_dim > 0, "DrafterKvCache: head_dim must be > 0");
let total_elems = num_kv_heads
.checked_mul(capacity)
.and_then(|v| v.checked_mul(head_dim))
.ok_or_else(|| {
anyhow!(
"DrafterKvCache: num_kv_heads ({}) * capacity ({}) * head_dim ({}) overflows usize",
num_kv_heads,
capacity,
head_dim
)
})?;
ensure!(
total_elems <= (u32::MAX as usize),
"DrafterKvCache: total elements ({}) exceeds u32::MAX",
total_elems
);
let total_bytes = total_elems
.checked_mul(std::mem::size_of::<f32>())
.ok_or_else(|| anyhow!("DrafterKvCache: byte size overflows usize"))?;
let k_buf = device
.alloc_buffer(
total_bytes,
DType::F32,
vec![num_kv_heads, capacity, head_dim],
)
.map_err(|e| anyhow!("alloc K cache: {e}"))?;
let v_buf = device
.alloc_buffer(
total_bytes,
DType::F32,
vec![num_kv_heads, capacity, head_dim],
)
.map_err(|e| anyhow!("alloc V cache: {e}"))?;
Ok(Self {
k_buf,
v_buf,
num_kv_heads,
capacity,
head_dim,
len: 0,
})
}
pub fn len(&self) -> usize {
self.len
}
pub fn is_empty(&self) -> bool {
self.len == 0
}
pub fn clear(&mut self) {
self.len = 0;
}
pub fn append(&mut self, k_row: &[f32], v_row: &[f32]) -> Result<()> {
ensure!(
self.len < self.capacity,
"DrafterKvCache::append: cache full (len={}, capacity={})",
self.len,
self.capacity
);
let expected = self.num_kv_heads * self.head_dim;
ensure!(
k_row.len() == expected,
"DrafterKvCache::append: k_row has {} elements, expected {} (num_kv_heads {} * head_dim {})",
k_row.len(),
expected,
self.num_kv_heads,
self.head_dim
);
ensure!(
v_row.len() == expected,
"DrafterKvCache::append: v_row has {} elements, expected {}",
v_row.len(),
expected
);
let k_slice = self
.k_buf
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("k_buf slice: {e}"))?;
let v_slice = self
.v_buf
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("v_buf slice: {e}"))?;
for h in 0..self.num_kv_heads {
let cache_offset = h * self.capacity * self.head_dim + self.len * self.head_dim;
let row_offset = h * self.head_dim;
k_slice[cache_offset..cache_offset + self.head_dim]
.copy_from_slice(&k_row[row_offset..row_offset + self.head_dim]);
v_slice[cache_offset..cache_offset + self.head_dim]
.copy_from_slice(&v_row[row_offset..row_offset + self.head_dim]);
}
self.len += 1;
Ok(())
}
pub fn rollback_to_accepted(&mut self, accepted: &[usize]) -> Result<()> {
ensure!(
!accepted.is_empty(),
"DrafterKvCache::rollback_to_accepted: accepted must be non-empty (root always included)"
);
ensure!(
accepted.len() <= self.capacity,
"DrafterKvCache::rollback_to_accepted: accepted len {} > capacity {}",
accepted.len(),
self.capacity
);
for (i, &idx) in accepted.iter().enumerate() {
ensure!(
idx < self.len,
"DrafterKvCache::rollback_to_accepted: accepted[{}] = {} >= current len {}",
i,
idx,
self.len
);
}
let mut seen = std::collections::HashSet::with_capacity(accepted.len());
for (i, &idx) in accepted.iter().enumerate() {
ensure!(
seen.insert(idx),
"DrafterKvCache::rollback_to_accepted: duplicate index {} at position {}",
idx,
i
);
}
let k_data = self
.k_buf
.as_slice::<f32>()
.map_err(|e| anyhow!("k_buf slice: {e}"))?
.to_vec();
let v_data = self
.v_buf
.as_slice::<f32>()
.map_err(|e| anyhow!("v_buf slice: {e}"))?
.to_vec();
let stride_per_head = self.capacity * self.head_dim;
let new_len = accepted.len();
let mut new_k = vec![0.0f32; self.num_kv_heads * stride_per_head];
let mut new_v = vec![0.0f32; self.num_kv_heads * stride_per_head];
for h in 0..self.num_kv_heads {
new_k[h * stride_per_head..(h + 1) * stride_per_head]
.copy_from_slice(&k_data[h * stride_per_head..(h + 1) * stride_per_head]);
new_v[h * stride_per_head..(h + 1) * stride_per_head]
.copy_from_slice(&v_data[h * stride_per_head..(h + 1) * stride_per_head]);
}
for (new_pos, &src_idx) in accepted.iter().enumerate() {
for h in 0..self.num_kv_heads {
let src_offset = h * stride_per_head + src_idx * self.head_dim;
let dst_offset = h * stride_per_head + new_pos * self.head_dim;
new_k[dst_offset..dst_offset + self.head_dim]
.copy_from_slice(&k_data[src_offset..src_offset + self.head_dim]);
new_v[dst_offset..dst_offset + self.head_dim]
.copy_from_slice(&v_data[src_offset..src_offset + self.head_dim]);
}
}
self.k_buf
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("k_buf mut slice: {e}"))?
.copy_from_slice(&new_k);
self.v_buf
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("v_buf mut slice: {e}"))?
.copy_from_slice(&new_v);
self.len = new_len;
Ok(())
}
}
pub struct MultiSeqDrafterKvCache {
pub n_seqs: u32,
pub k_buf: MlxBuffer,
pub v_buf: MlxBuffer,
pub num_kv_heads: usize,
pub capacity: usize,
pub head_dim: usize,
pub seq_lens: Vec<u32>,
}
impl MultiSeqDrafterKvCache {
pub const PADDING_SLOT: crate::serve::multi_seq_kv::SlotId =
crate::serve::multi_seq_kv::SlotId(u32::MAX);
pub fn reset_for_slot(
&mut self,
slot: crate::serve::multi_seq_kv::SlotId,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: self.n_seqs,
});
}
self.seq_lens[slot.0 as usize] = 0;
Ok(())
}
}
pub fn alloc_multi_seq_drafter_kv_for_layer(
device: &MlxDevice,
num_kv_heads: usize,
capacity: usize,
head_dim: usize,
n_seqs: u32,
) -> Result<MultiSeqDrafterKvCache> {
ensure!(
n_seqs > 0,
"alloc_multi_seq_drafter_kv_for_layer: n_seqs must be > 0"
);
ensure!(
n_seqs != u32::MAX,
"alloc_multi_seq_drafter_kv_for_layer: n_seqs must be < u32::MAX \
(reserved as PADDING_SLOT sentinel per vLLM/P-EAGLE convention)"
);
ensure!(
num_kv_heads > 0,
"alloc_multi_seq_drafter_kv_for_layer: num_kv_heads must be > 0"
);
ensure!(
capacity > 0,
"alloc_multi_seq_drafter_kv_for_layer: capacity must be > 0"
);
ensure!(
head_dim > 0,
"alloc_multi_seq_drafter_kv_for_layer: head_dim must be > 0"
);
let n = n_seqs as usize;
let total_elems = n
.checked_mul(num_kv_heads)
.and_then(|v| v.checked_mul(capacity))
.and_then(|v| v.checked_mul(head_dim))
.ok_or_else(|| {
anyhow!(
"alloc_multi_seq_drafter_kv_for_layer: n_seqs ({}) * num_kv_heads \
({}) * capacity ({}) * head_dim ({}) overflows usize",
n_seqs,
num_kv_heads,
capacity,
head_dim
)
})?;
ensure!(
total_elems <= (u32::MAX as usize),
"alloc_multi_seq_drafter_kv_for_layer: total elements ({}) exceeds u32::MAX \
(MlxBuffer shape bound)",
total_elems
);
let total_bytes = total_elems
.checked_mul(std::mem::size_of::<f32>())
.ok_or_else(|| {
anyhow!("alloc_multi_seq_drafter_kv_for_layer: byte size overflows usize")
})?;
let shape = vec![n, num_kv_heads, capacity, head_dim];
let mut k_buf = device
.alloc_buffer(total_bytes, DType::F32, shape.clone())
.map_err(|e| anyhow!("alloc multi-seq drafter K: {e}"))?;
let mut v_buf = device
.alloc_buffer(total_bytes, DType::F32, shape)
.map_err(|e| anyhow!("alloc multi-seq drafter V: {e}"))?;
if let Ok(s) = k_buf.as_mut_slice::<f32>() {
s.fill(0.0);
}
if let Ok(s) = v_buf.as_mut_slice::<f32>() {
s.fill(0.0);
}
Ok(MultiSeqDrafterKvCache {
n_seqs,
k_buf,
v_buf,
num_kv_heads,
capacity,
head_dim,
seq_lens: vec![0u32; n],
})
}
impl crate::serve::multi_seq_kv::MultiSeqKvCache for MultiSeqDrafterKvCache {
fn layout(&self) -> crate::serve::multi_seq_kv::MultiSeqLayout {
crate::serve::multi_seq_kv::MultiSeqLayout::SeparateSlots
}
fn slot_count(&self) -> u32 {
self.n_seqs
}
fn seq_len(
&self,
slot: crate::serve::multi_seq_kv::SlotId,
) -> Result<u32, crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: self.n_seqs,
});
}
Ok(self.seq_lens[slot.0 as usize])
}
fn append_for_seq(
&mut self,
slot: crate::serve::multi_seq_kv::SlotId,
n_tokens: u32,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: self.n_seqs,
});
}
let cur = &mut self.seq_lens[slot.0 as usize];
*cur = cur.saturating_add(n_tokens);
Ok(())
}
fn drop_seq(
&mut self,
slot: crate::serve::multi_seq_kv::SlotId,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: self.n_seqs,
});
}
self.seq_lens[slot.0 as usize] = 0;
Ok(())
}
fn fork_seq(
&mut self,
src: crate::serve::multi_seq_kv::SlotId,
dst: crate::serve::multi_seq_kv::SlotId,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if src.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot: src,
max_slots: self.n_seqs,
});
}
if dst.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot: dst,
max_slots: self.n_seqs,
});
}
if src == dst {
return Ok(());
}
let src_idx = src.0 as usize;
let dst_idx = dst.0 as usize;
let n_seqs = self.n_seqs as usize;
drafter_copy_buffer_slot_region(&mut self.k_buf, src_idx, dst_idx, n_seqs).map_err(
|e| crate::serve::multi_seq_kv::MultiSeqError::CapabilityUnsupported {
capability: drafter_leak_static_str(format!(
"fork_seq: MultiSeqDrafterKvCache k_buf copy failed ({e})"
)),
},
)?;
drafter_copy_buffer_slot_region(&mut self.v_buf, src_idx, dst_idx, n_seqs).map_err(
|e| crate::serve::multi_seq_kv::MultiSeqError::CapabilityUnsupported {
capability: drafter_leak_static_str(format!(
"fork_seq: MultiSeqDrafterKvCache v_buf copy failed ({e})"
)),
},
)?;
self.seq_lens[dst_idx] = self.seq_lens[src_idx];
Ok(())
}
}
#[inline]
fn drafter_leak_static_str(s: String) -> &'static str {
Box::leak(s.into_boxed_str())
}
fn drafter_copy_buffer_slot_region(
buf: &mut MlxBuffer,
src_idx: usize,
dst_idx: usize,
n_seqs: usize,
) -> Result<()> {
ensure!(n_seqs > 0, "fork_seq: n_seqs must be > 0");
let total_bytes = buf.byte_len();
ensure!(
total_bytes % n_seqs == 0,
"fork_seq: total_bytes={} not divisible by n_seqs={}",
total_bytes,
n_seqs
);
let per_slot_bytes = total_bytes / n_seqs;
ensure!(
src_idx < n_seqs && dst_idx < n_seqs,
"fork_seq: src/dst out of buffer range \
(src={src_idx}, dst={dst_idx}, n_seqs={n_seqs})"
);
if per_slot_bytes == 0 {
return Ok(());
}
let bytes = buf
.as_mut_slice::<u8>()
.map_err(|e| anyhow!("fork_seq: as_mut_slice<u8>: {e}"))?;
let src_off = src_idx * per_slot_bytes;
bytes.copy_within(src_off..src_off + per_slot_bytes, dst_idx * per_slot_bytes);
Ok(())
}
pub enum DrafterKvCacheVariant {
SingleSeq(DrafterKvCache),
MultiSeq(MultiSeqDrafterKvCache),
}
impl DrafterKvCacheVariant {
pub fn slot_count(&self) -> u32 {
match self {
Self::SingleSeq(_) => 1,
Self::MultiSeq(c) => c.n_seqs,
}
}
pub fn is_multi_seq(&self) -> bool {
matches!(self, Self::MultiSeq(_))
}
}
#[inline]
pub fn select_drafter_kv_variant_for_mode(max_slots: u32) -> DrafterKvCacheSelection {
if max_slots <= 1 {
DrafterKvCacheSelection::SingleSeq
} else {
DrafterKvCacheSelection::MultiSeq { n_seqs: max_slots }
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DrafterKvCacheSelection {
SingleSeq,
MultiSeq { n_seqs: u32 },
}
#[cfg(test)]
#[allow(clippy::expect_used, clippy::unwrap_used, clippy::panic)]
mod tests {
use super::*;
fn make_cache() -> Option<(MlxDevice, DrafterKvCache)> {
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return None,
};
let cache = DrafterKvCache::new(&device, 2, 8, 4).expect("alloc");
Some((device, cache))
}
fn sentinel_row(num_kv_heads: usize, head_dim: usize, tag: u32) -> Vec<f32> {
let mut out = vec![0.0f32; num_kv_heads * head_dim];
for h in 0..num_kv_heads {
for d in 0..head_dim {
out[h * head_dim + d] = (tag * 1000 + h as u32 * 100 + d as u32) as f32;
}
}
out
}
#[test]
fn adr_037_e5b_kv_cache_constructor_validates_dims_2026_05_22() {
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
assert!(DrafterKvCache::new(&device, 0, 8, 4).is_err());
assert!(DrafterKvCache::new(&device, 2, 0, 4).is_err());
assert!(DrafterKvCache::new(&device, 2, 8, 0).is_err());
}
#[test]
fn adr_037_e5b_kv_cache_initial_state_empty_2026_05_22() {
let (_dev, cache) = match make_cache() {
Some(t) => t,
None => return,
};
assert_eq!(cache.len(), 0);
assert!(cache.is_empty());
assert_eq!(cache.num_kv_heads, 2);
assert_eq!(cache.capacity, 8);
assert_eq!(cache.head_dim, 4);
}
#[test]
fn adr_037_e5b_kv_cache_append_grows_len_2026_05_22() {
let (_dev, mut cache) = match make_cache() {
Some(t) => t,
None => return,
};
for tag in 0..5 {
let k = sentinel_row(cache.num_kv_heads, cache.head_dim, tag);
let v = sentinel_row(cache.num_kv_heads, cache.head_dim, tag + 100);
cache.append(&k, &v).expect("append");
assert_eq!(cache.len(), (tag + 1) as usize);
}
}
#[test]
fn adr_037_e5b_kv_cache_append_rejects_full_2026_05_22() {
let (_dev, mut cache) = match make_cache() {
Some(t) => t,
None => return,
};
let k = sentinel_row(2, 4, 0);
let v = sentinel_row(2, 4, 0);
for _ in 0..8 {
cache.append(&k, &v).expect("append within capacity");
}
let err = cache.append(&k, &v).unwrap_err();
assert!(err.to_string().contains("cache full"), "got: {err}");
}
#[test]
fn adr_037_e5b_kv_cache_append_rejects_wrong_row_shape_2026_05_22() {
let (_dev, mut cache) = match make_cache() {
Some(t) => t,
None => return,
};
let bad_k = vec![0.0f32; 5]; let v = sentinel_row(2, 4, 0);
let err = cache.append(&bad_k, &v).unwrap_err();
assert!(err.to_string().contains("k_row has 5"), "got: {err}");
}
#[test]
fn adr_037_e5b_kv_cache_rollback_keeps_only_accepted_2026_05_22() {
let (_dev, mut cache) = match make_cache() {
Some(t) => t,
None => return,
};
for tag in 0..5 {
let k = sentinel_row(cache.num_kv_heads, cache.head_dim, tag);
let v = sentinel_row(cache.num_kv_heads, cache.head_dim, tag + 100);
cache.append(&k, &v).expect("append");
}
cache.rollback_to_accepted(&[0, 2, 4]).expect("rollback");
assert_eq!(cache.len(), 3);
let k_data: Vec<f32> = cache.k_buf.as_slice::<f32>().unwrap().to_vec();
let stride_per_head = cache.capacity * cache.head_dim;
for (new_pos, &src_tag) in [0_u32, 2, 4].iter().enumerate() {
for h in 0..cache.num_kv_heads {
for d in 0..cache.head_dim {
let offset = h * stride_per_head + new_pos * cache.head_dim + d;
let expected = (src_tag * 1000 + h as u32 * 100 + d as u32) as f32;
assert_eq!(
k_data[offset], expected,
"rollback pos {} head {} dim {} expected tag {} got {}",
new_pos, h, d, src_tag, k_data[offset]
);
}
}
}
}
#[test]
fn adr_037_e5b_kv_cache_rollback_rejects_empty_accepted_2026_05_22() {
let (_dev, mut cache) = match make_cache() {
Some(t) => t,
None => return,
};
let k = sentinel_row(2, 4, 0);
let v = sentinel_row(2, 4, 0);
cache.append(&k, &v).expect("append");
let err = cache.rollback_to_accepted(&[]).unwrap_err();
assert!(err.to_string().contains("must be non-empty"), "got: {err}");
}
#[test]
fn adr_037_e5b_kv_cache_rollback_rejects_out_of_range_idx_2026_05_22() {
let (_dev, mut cache) = match make_cache() {
Some(t) => t,
None => return,
};
let k = sentinel_row(2, 4, 0);
let v = sentinel_row(2, 4, 0);
cache.append(&k, &v).expect("append"); let err = cache.rollback_to_accepted(&[0, 5]).unwrap_err();
assert!(err.to_string().contains(">= current len"), "got: {err}");
}
#[test]
fn adr_037_e5b_kv_cache_rollback_rejects_duplicate_idx_2026_05_22() {
let (_dev, mut cache) = match make_cache() {
Some(t) => t,
None => return,
};
let k = sentinel_row(2, 4, 0);
let v = sentinel_row(2, 4, 0);
for _ in 0..3 {
cache.append(&k, &v).expect("append");
}
let err = cache.rollback_to_accepted(&[0, 1, 1]).unwrap_err();
assert!(err.to_string().contains("duplicate index"), "got: {err}");
}
#[test]
fn adr_037_e5b_kv_cache_rollback_to_root_only_2026_05_22() {
let (_dev, mut cache) = match make_cache() {
Some(t) => t,
None => return,
};
for tag in 0..4 {
let k = sentinel_row(cache.num_kv_heads, cache.head_dim, tag);
let v = sentinel_row(cache.num_kv_heads, cache.head_dim, tag + 100);
cache.append(&k, &v).expect("append");
}
cache.rollback_to_accepted(&[0]).expect("rollback to root");
assert_eq!(cache.len(), 1);
let k_data = cache.k_buf.as_slice::<f32>().unwrap();
let stride_per_head = cache.capacity * cache.head_dim;
for h in 0..cache.num_kv_heads {
for d in 0..cache.head_dim {
let offset = h * stride_per_head + d;
let expected = (0_u32 * 1000 + h as u32 * 100 + d as u32) as f32;
assert_eq!(k_data[offset], expected);
}
}
}
#[test]
fn adr_037_e5b_kv_cache_clear_resets_len_2026_05_22() {
let (_dev, mut cache) = match make_cache() {
Some(t) => t,
None => return,
};
let k = sentinel_row(2, 4, 0);
let v = sentinel_row(2, 4, 0);
cache.append(&k, &v).expect("append");
cache.append(&k, &v).expect("append");
assert_eq!(cache.len(), 2);
cache.clear();
assert_eq!(cache.len(), 0);
assert!(cache.is_empty());
}
#[test]
fn adr_037_e5b_kv_cache_integration_with_tree_walk_accept_2026_05_22() {
use crate::inference::spec_decode::eagle3::dynamic_tree::ExpandedTree;
use crate::inference::spec_decode::eagle3::tree_walk::walk_tree_accept;
let (_dev, mut cache) = match make_cache() {
Some(t) => t,
None => return,
};
for tag in 0..4 {
let k = sentinel_row(cache.num_kv_heads, cache.head_dim, tag);
let v = sentinel_row(cache.num_kv_heads, cache.head_dim, tag + 100);
cache.append(&k, &v).expect("append");
}
let tree = ExpandedTree {
tokens: vec![100, 1, 2, 3],
parents: vec![None, Some(0), Some(1), Some(0)],
depths: vec![0, 1, 2, 1],
cum_log_probs: vec![0.0, -0.1, -0.2, -0.5],
};
let argmax = vec![1_u32, 2, 0, 0];
let accepted = walk_tree_accept(&tree, &argmax).expect("walk");
assert_eq!(accepted, vec![0, 1, 2]);
cache.rollback_to_accepted(&accepted).expect("rollback");
assert_eq!(cache.len(), 3);
let k_data = cache.k_buf.as_slice::<f32>().unwrap();
let stride_per_head = cache.capacity * cache.head_dim;
for (new_pos, expected_tag) in [0_u32, 1, 2].iter().enumerate() {
for h in 0..cache.num_kv_heads {
let offset = h * stride_per_head + new_pos * cache.head_dim;
let expected = (expected_tag * 1000 + h as u32 * 100) as f32;
assert_eq!(k_data[offset], expected);
}
}
}
fn make_multi_seq_cache(n_seqs: u32) -> Option<(MlxDevice, MultiSeqDrafterKvCache)> {
let device = MlxDevice::new().ok()?;
let cache = alloc_multi_seq_drafter_kv_for_layer(&device, 2, 8, 4, n_seqs).expect("alloc");
Some((device, cache))
}
#[test]
fn h224_multi_seq_drafter_kv_cache_carries_dossier_shape_2026_05_30() {
let (_dev, cache) = match make_multi_seq_cache(3) {
Some(t) => t,
None => {
eprintln!(
"[skip] h224 — MlxDevice unavailable (CI without GPU); \
skip-mode structural pins for PADDING_SLOT + alloc helper \
run unconditionally elsewhere."
);
return;
}
};
assert_eq!(cache.n_seqs, 3, "H224: n_seqs round-trips alloc arg");
assert_eq!(cache.num_kv_heads, 2);
assert_eq!(cache.capacity, 8);
assert_eq!(cache.head_dim, 4);
assert_eq!(
cache.seq_lens.len(),
3,
"H224: seq_lens.len() == n_seqs by construction"
);
assert!(
cache.seq_lens.iter().all(|&l| l == 0),
"H224: all per-slot cursors start at 0"
);
assert_eq!(
cache.k_buf.byte_len(),
3 * 2 * 8 * 4 * std::mem::size_of::<f32>(),
"H224: K buffer total bytes match [n_seqs, nkv, cap, hd] F32"
);
assert_eq!(
cache.v_buf.byte_len(),
3 * 2 * 8 * 4 * std::mem::size_of::<f32>(),
"H224: V buffer total bytes match [n_seqs, nkv, cap, hd] F32"
);
}
#[test]
fn h225_alloc_multi_seq_drafter_kv_validates_dims_2026_05_30() {
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => {
eprintln!("[skip] h225 — MlxDevice unavailable");
return;
}
};
assert!(
alloc_multi_seq_drafter_kv_for_layer(&device, 2, 8, 4, 0).is_err(),
"H225: n_seqs == 0 must error"
);
let err = alloc_multi_seq_drafter_kv_for_layer(&device, 2, 8, 4, u32::MAX)
.err()
.expect("H225: n_seqs == u32::MAX must error");
assert!(
err.to_string().contains("PADDING_SLOT"),
"H225: PADDING_SLOT collision must be named in error. Got: {err}"
);
assert!(
alloc_multi_seq_drafter_kv_for_layer(&device, 0, 8, 4, 2).is_err(),
"H225: num_kv_heads == 0 must error"
);
assert!(
alloc_multi_seq_drafter_kv_for_layer(&device, 2, 0, 4, 2).is_err(),
"H225: capacity == 0 must error"
);
assert!(
alloc_multi_seq_drafter_kv_for_layer(&device, 2, 8, 0, 2).is_err(),
"H225: head_dim == 0 must error"
);
}
#[test]
fn h226_padding_slot_const_is_max_u32_per_vllm_p_eagle_2026_05_30() {
assert_eq!(
MultiSeqDrafterKvCache::PADDING_SLOT,
crate::serve::multi_seq_kv::SlotId(u32::MAX),
"H226: PADDING_SLOT MUST be SlotId(u32::MAX) per vLLM/P-EAGLE \
rejected-token convention (dossier §1.1 + §5)."
);
for n in [1u32, 2, 4, 8, 16, 1024, u32::MAX - 1] {
assert!(
MultiSeqDrafterKvCache::PADDING_SLOT.0 >= n,
"H226: PADDING_SLOT must be outside in-bounds range \
for every reachable n_seqs (got n={n})"
);
}
}
#[test]
fn h227_multi_seq_kv_cache_impl_bounds_first_and_fork_2026_05_30() {
use crate::serve::multi_seq_kv::{MultiSeqError, MultiSeqKvCache, SlotId};
let (_dev, mut cache) = match make_multi_seq_cache(3) {
Some(t) => t,
None => {
eprintln!("[skip] h227 — MlxDevice unavailable");
return;
}
};
assert_eq!(cache.slot_count(), 3, "H227: slot_count() == n_seqs");
let oor = cache.seq_len(SlotId(3)).unwrap_err();
assert_eq!(
oor,
MultiSeqError::SlotOutOfRange {
slot: SlotId(3),
max_slots: 3,
},
"H227: out-of-range slot must surface SlotOutOfRange"
);
let pad_err = cache
.seq_len(MultiSeqDrafterKvCache::PADDING_SLOT)
.unwrap_err();
assert!(
matches!(pad_err, MultiSeqError::SlotOutOfRange { .. }),
"H227: PADDING_SLOT against trait surface MUST be \
SlotOutOfRange. Got: {pad_err:?}"
);
cache.append_for_seq(SlotId(1), 7).expect("append");
assert_eq!(cache.seq_len(SlotId(0)).unwrap(), 0);
assert_eq!(cache.seq_len(SlotId(1)).unwrap(), 7);
assert_eq!(cache.seq_len(SlotId(2)).unwrap(), 0);
let slot_stride = cache.num_kv_heads * cache.capacity * cache.head_dim;
{
let k_slice = cache.k_buf.as_mut_slice::<f32>().expect("k slice");
for i in 0..slot_stride {
k_slice[slot_stride + i] = (i + 1) as f32; }
}
cache
.fork_seq(SlotId(1), SlotId(1))
.expect("same-slot fork");
cache.fork_seq(SlotId(1), SlotId(2)).expect("fork 1→2");
assert_eq!(
cache.seq_len(SlotId(2)).unwrap(),
7,
"H227: fork_seq copies the per-slot cursor"
);
let k_slice = cache.k_buf.as_slice::<f32>().expect("k slice");
for i in 0..slot_stride {
let expected = (i + 1) as f32;
assert_eq!(
k_slice[2 * slot_stride + i],
expected,
"H227: fork_seq must memcpy the K-buffer per-slot region \
(slot 2[{i}] = {expected})"
);
}
assert_eq!(cache.seq_len(SlotId(0)).unwrap(), 0);
cache.drop_seq(SlotId(2)).expect("drop slot 2");
assert_eq!(cache.seq_len(SlotId(2)).unwrap(), 0);
let k_after_drop = cache.k_buf.as_slice::<f32>().expect("k slice");
for i in 0..slot_stride {
assert_eq!(
k_after_drop[2 * slot_stride + i],
(i + 1) as f32,
"H227: drop_seq must NOT zero K/V bytes (preserves \
recurrent-content invariance)"
);
}
}
#[test]
fn h228_reset_for_slot_cursor_only_byte_preservation_2026_05_30() {
use crate::serve::multi_seq_kv::{MultiSeqError, SlotId};
let (_dev, mut cache) = match make_multi_seq_cache(2) {
Some(t) => t,
None => {
eprintln!("[skip] h228 — MlxDevice unavailable");
return;
}
};
let err = cache.reset_for_slot(SlotId(2)).unwrap_err();
assert!(
matches!(err, MultiSeqError::SlotOutOfRange { .. }),
"H228: bounds-first per A2b iter-1.5 cfa-finding-F5. Got: {err:?}"
);
let err_pad = cache
.reset_for_slot(MultiSeqDrafterKvCache::PADDING_SLOT)
.unwrap_err();
assert!(matches!(err_pad, MultiSeqError::SlotOutOfRange { .. }));
let slot_stride = cache.num_kv_heads * cache.capacity * cache.head_dim;
{
let k_slice = cache.k_buf.as_mut_slice::<f32>().expect("k slice");
for i in 0..(2 * slot_stride) {
k_slice[i] = (i + 1) as f32;
}
let v_slice = cache.v_buf.as_mut_slice::<f32>().expect("v slice");
for i in 0..(2 * slot_stride) {
v_slice[i] = (100 + i) as f32;
}
}
cache.seq_lens[0] = 5;
cache.seq_lens[1] = 6;
cache.reset_for_slot(SlotId(0)).expect("reset slot 0");
assert_eq!(cache.seq_lens[0], 0, "H228: cursor reset at slot 0");
assert_eq!(cache.seq_lens[1], 6, "H228: other-slot cursor untouched");
let k_slice = cache.k_buf.as_slice::<f32>().expect("k slice");
let v_slice = cache.v_buf.as_slice::<f32>().expect("v slice");
for i in 0..(2 * slot_stride) {
assert_eq!(
k_slice[i],
(i + 1) as f32,
"H228: reset_for_slot MUST NOT zero K bytes (cursor-masked)"
);
assert_eq!(
v_slice[i],
(100 + i) as f32,
"H228: reset_for_slot MUST NOT zero V bytes (cursor-masked)"
);
}
}
#[test]
fn h230_multi_seq_n_seqs_1_byte_equiv_to_legacy_drafter_kv_2026_05_30() {
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => {
eprintln!("[skip] h230 — MlxDevice unavailable");
return;
}
};
let legacy = DrafterKvCache::new(&device, 2, 8, 4).expect("legacy alloc");
let multi = alloc_multi_seq_drafter_kv_for_layer(&device, 2, 8, 4, 1).expect("multi alloc");
assert_eq!(
legacy.k_buf.byte_len(),
multi.k_buf.byte_len(),
"H230: legacy K byte count must equal multi-seq K byte count at n_seqs=1"
);
assert_eq!(
legacy.v_buf.byte_len(),
multi.v_buf.byte_len(),
"H230: legacy V byte count must equal multi-seq V byte count at n_seqs=1"
);
assert_eq!(multi.n_seqs, 1);
assert_eq!(multi.seq_lens, vec![0u32]);
}
#[test]
fn h231_legacy_drafter_kv_cache_surface_unchanged_2026_05_30() {
let _ctor: fn(&MlxDevice, usize, usize, usize) -> Result<DrafterKvCache> =
DrafterKvCache::new;
if let Ok(device) = MlxDevice::new() {
let cache = DrafterKvCache::new(&device, 2, 8, 4).expect("alloc");
assert_eq!(cache.len(), 0);
assert!(cache.is_empty());
assert_eq!(cache.num_kv_heads, 2);
assert_eq!(cache.capacity, 8);
assert_eq!(cache.head_dim, 4);
} else {
eprintln!(
"[skip] h231 — MlxDevice unavailable; signature pin still asserted at compile."
);
}
}
#[test]
fn h232_multi_seq_kv_trait_witness_cross_arch_unchanged_2026_05_30() {
fn assert_multi_seq_kv<T: crate::serve::multi_seq_kv::MultiSeqKvCache>() {}
assert_multi_seq_kv::<MultiSeqDrafterKvCache>();
assert_multi_seq_kv::<crate::inference::models::gemma4::kv_cache::MultiSeqHbKvBuffers>();
assert_multi_seq_kv::<crate::inference::models::gemma4::kv_cache::MultiSeqHybridKvBuffers>(
);
assert_multi_seq_kv::<crate::serve::multi_seq_kv::NoopMultiSeqKvCache>();
}
#[test]
fn h235a_drafter_kv_cache_selection_routes_by_max_slots_2026_05_30() {
assert_eq!(
select_drafter_kv_variant_for_mode(0),
DrafterKvCacheSelection::SingleSeq,
"H235a: max_slots == 0 MUST degrade to SingleSeq (defense-in-depth)"
);
assert_eq!(
select_drafter_kv_variant_for_mode(1),
DrafterKvCacheSelection::SingleSeq,
"H235a: max_slots == 1 MUST route to SingleSeq (byte-equivalent to \
MultiSeq at n_seqs=1; pre-A4 preserved)"
);
for n in [2u32, 3, 4, 8, 16, 1024] {
assert_eq!(
select_drafter_kv_variant_for_mode(n),
DrafterKvCacheSelection::MultiSeq { n_seqs: n },
"H235a: max_slots == {n} MUST route to MultiSeq with the requested n_seqs"
);
}
}
#[test]
fn h235b_drafter_kv_cache_variant_slot_count_pin_2026_05_30() {
let (_, single) = match make_cache() {
Some(t) => t,
None => {
eprintln!(
"[skip] h235b — MlxDevice unavailable; selection enum + \
is_multi_seq pinned in h235a / h235c."
);
return;
}
};
let variant = DrafterKvCacheVariant::SingleSeq(single);
assert_eq!(
variant.slot_count(),
1,
"H235b: SingleSeq arm slot_count == 1"
);
assert!(!variant.is_multi_seq(), "H235b: SingleSeq is NOT multi_seq");
}
#[test]
fn h235c_drafter_kv_cache_selection_copy_eq_witness_2026_05_30() {
let s = DrafterKvCacheSelection::MultiSeq { n_seqs: 4 };
let s2 = s; assert_eq!(s, s2);
let s3 = DrafterKvCacheSelection::SingleSeq;
assert_ne!(s, s3);
}
#[test]
fn h235d_drafter_dispatcher_cite_named_at_source_2026_05_30() {
let src = include_str!("kv_cache.rs");
assert!(
src.contains("iter-A4-cont-drafter-dispatcher"),
"H235d FALSIFIED: kv_cache.rs does NOT name \
`iter-A4-cont-drafter-dispatcher` at the dispatcher cite."
);
assert!(
src.contains("DrafterKvCacheVariant"),
"H235d FALSIFIED: DrafterKvCacheVariant variant carrier missing."
);
assert!(
src.contains("select_drafter_kv_variant_for_mode"),
"H235d FALSIFIED: select_drafter_kv_variant_for_mode routing \
helper missing."
);
assert!(
src.contains("SingleSeq") && src.contains("MultiSeq"),
"H235d FALSIFIED: SingleSeq / MultiSeq arm names missing."
);
}
}