use crate::cache::KvCache;
use crate::kv_swa::{BlockLayout, BlockLayoutError};
pub const BLOCK_FORMAT_VERSION: u32 = 2;
pub const READABLE_FORMAT_VERSIONS: &[u32] = &[2];
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum KvDtype {
F32,
}
impl KvDtype {
pub fn as_str(self) -> &'static str {
match self {
KvDtype::F32 => "f32",
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct CacheSignature {
pub format_version: u32,
pub model: String,
pub n_layers: usize,
pub n_kv_heads: usize,
pub head_dim: usize,
pub dtype: KvDtype,
pub tokens: usize,
pub layout: BlockLayout,
}
impl CacheSignature {
pub fn from_payload(
model: &str,
layout: BlockLayout,
layers: &[KvCache],
) -> Result<Self, SignatureError> {
let first = layers.first().ok_or(SignatureError::EmptyPayload)?;
let n_kv_heads = first.n_kv_heads;
let head_dim = first.head_dim;
if n_kv_heads == 0 || head_dim == 0 {
return Err(SignatureError::DegenerateLayer {
layer: 0,
n_kv_heads,
head_dim,
});
}
let per_token = n_kv_heads * head_dim;
let tokens = measure_layer(0, first, per_token)?;
for (index, layer) in layers.iter().enumerate().skip(1) {
if layer.n_kv_heads != n_kv_heads || layer.head_dim != head_dim {
return Err(SignatureError::RaggedPayload {
layer: index,
field: "layer shape",
expected: format!("{n_kv_heads}x{head_dim}"),
found: format!("{}x{}", layer.n_kv_heads, layer.head_dim),
});
}
let layer_tokens = measure_layer(index, layer, per_token)?;
if layer_tokens != tokens {
return Err(SignatureError::RaggedPayload {
layer: index,
field: "token count",
expected: tokens.to_string(),
found: layer_tokens.to_string(),
});
}
}
if tokens != layout.block_size() {
return Err(SignatureError::BlockSizeMismatch {
block_size: layout.block_size(),
tokens,
});
}
Ok(CacheSignature {
format_version: BLOCK_FORMAT_VERSION,
model: model.to_string(),
n_layers: layers.len(),
n_kv_heads,
head_dim,
dtype: KvDtype::F32,
tokens,
layout,
})
}
pub fn expected(
model: &str,
layout: BlockLayout,
n_layers: usize,
n_kv_heads: usize,
head_dim: usize,
tokens: usize,
) -> Self {
CacheSignature {
format_version: BLOCK_FORMAT_VERSION,
model: model.to_string(),
n_layers,
n_kv_heads,
head_dim,
dtype: KvDtype::F32,
tokens,
layout,
}
}
fn compare(
&self,
other: &CacheSignature,
mismatch: fn(&'static str, String, String) -> SignatureError,
) -> Result<(), SignatureError> {
if self.format_version != other.format_version {
return Err(mismatch(
"format_version",
self.format_version.to_string(),
other.format_version.to_string(),
));
}
if self.model != other.model {
return Err(mismatch("model", self.model.clone(), other.model.clone()));
}
if self.n_layers != other.n_layers {
return Err(mismatch(
"n_layers",
self.n_layers.to_string(),
other.n_layers.to_string(),
));
}
if self.n_kv_heads != other.n_kv_heads {
return Err(mismatch(
"n_kv_heads",
self.n_kv_heads.to_string(),
other.n_kv_heads.to_string(),
));
}
if self.head_dim != other.head_dim {
return Err(mismatch(
"head_dim",
self.head_dim.to_string(),
other.head_dim.to_string(),
));
}
if self.dtype != other.dtype {
return Err(mismatch(
"dtype",
self.dtype.as_str().to_string(),
other.dtype.as_str().to_string(),
));
}
if self.tokens != other.tokens {
return Err(mismatch(
"tokens",
self.tokens.to_string(),
other.tokens.to_string(),
));
}
if self.layout.block_size() != other.layout.block_size() {
return Err(mismatch(
"block_size",
self.layout.block_size().to_string(),
other.layout.block_size().to_string(),
));
}
if self.layout.sliding_window() != other.layout.sliding_window() {
return Err(mismatch(
"sliding_window",
describe_window(self.layout.sliding_window()),
describe_window(other.layout.sliding_window()),
));
}
Ok(())
}
}
fn describe_window(window: Option<usize>) -> String {
match window {
Some(w) => w.to_string(),
None => "none (full causal)".to_string(),
}
}
fn measure_layer(index: usize, layer: &KvCache, per_token: usize) -> Result<usize, SignatureError> {
if !layer.k.len().is_multiple_of(per_token) {
return Err(SignatureError::RaggedPayload {
layer: index,
field: "k length",
expected: format!("a multiple of {per_token}"),
found: layer.k.len().to_string(),
});
}
if layer.v.len() != layer.k.len() {
return Err(SignatureError::RaggedPayload {
layer: index,
field: "v length",
expected: layer.k.len().to_string(),
found: layer.v.len().to_string(),
});
}
let tokens = layer.k.len() / per_token;
if layer.seq_len != tokens {
return Err(SignatureError::RaggedPayload {
layer: index,
field: "seq_len",
expected: tokens.to_string(),
found: layer.seq_len.to_string(),
});
}
Ok(tokens)
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum SignatureError {
Unmarked,
EmptyPayload,
DegenerateLayer {
layer: usize,
n_kv_heads: usize,
head_dim: usize,
},
RaggedPayload {
layer: usize,
field: &'static str,
expected: String,
found: String,
},
PayloadMismatch {
field: &'static str,
recorded: String,
actual: String,
},
Incompatible {
field: &'static str,
expected: String,
found: String,
},
BlockSizeMismatch { block_size: usize, tokens: usize },
BadLayout(BlockLayoutError),
UnsupportedFormat {
found: u32,
readable: &'static [u32],
},
}
impl From<BlockLayoutError> for SignatureError {
fn from(err: BlockLayoutError) -> Self {
SignatureError::BadLayout(err)
}
}
impl std::fmt::Display for SignatureError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
SignatureError::Unmarked => write!(
f,
"KV block carries no cache signature; refusing to trust an unmarked block"
),
SignatureError::EmptyPayload => {
write!(f, "KV block has no layers; nothing to verify")
}
SignatureError::DegenerateLayer {
layer,
n_kv_heads,
head_dim,
} => write!(
f,
"KV block layer {layer} is degenerate: {n_kv_heads} kv heads x {head_dim} head dim"
),
SignatureError::RaggedPayload {
layer,
field,
expected,
found,
} => write!(
f,
"KV block payload is inconsistent at layer {layer}: {field} is {found}, expected {expected}"
),
SignatureError::PayloadMismatch {
field,
recorded,
actual,
} => write!(
f,
"KV block signature vouches for {field}={recorded} but its payload has {field}={actual}"
),
SignatureError::Incompatible {
field,
expected,
found,
} => write!(
f,
"KV block is incompatible: {field} is {found}, this server needs {expected}"
),
SignatureError::BlockSizeMismatch { block_size, tokens } => write!(
f,
"KV block signature declares a block size of {block_size} but its payload holds \
{tokens} token positions; a stored block is exactly one whole block"
),
SignatureError::BadLayout(err) => write!(f, "KV block layout is unusable: {err}"),
SignatureError::UnsupportedFormat { found, readable } => write!(
f,
"KV block format version {found} is not readable by this build (readable: {readable:?})"
),
}
}
}
impl std::error::Error for SignatureError {}
pub struct KvBlock {
signature: CacheSignature,
layers: Vec<KvCache>,
}
impl std::fmt::Debug for KvBlock {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("KvBlock")
.field("signature", &self.signature)
.field("layers", &self.layers.len())
.finish()
}
}
impl std::fmt::Debug for UnverifiedBlock {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("UnverifiedBlock")
.field("signature", &self.signature)
.field("layers", &self.layers.len())
.finish()
}
}
impl KvBlock {
pub fn stamp(
model: &str,
layout: BlockLayout,
layers: Vec<KvCache>,
) -> Result<Self, SignatureError> {
let signature = CacheSignature::from_payload(model, layout, &layers)?;
Ok(KvBlock { signature, layers })
}
pub fn layout(&self) -> BlockLayout {
self.signature.layout
}
pub fn signature(&self) -> &CacheSignature {
&self.signature
}
pub fn tokens(&self) -> usize {
self.signature.tokens
}
pub fn layers(&self) -> &[KvCache] {
&self.layers
}
pub fn into_layers(self) -> Vec<KvCache> {
self.layers
}
}
pub struct UnverifiedBlock {
pub signature: Option<CacheSignature>,
pub layers: Vec<KvCache>,
}
impl UnverifiedBlock {
pub fn new(signature: Option<CacheSignature>, layers: Vec<KvCache>) -> Self {
UnverifiedBlock { signature, layers }
}
pub fn verify(self, expected: &CacheSignature) -> Result<KvBlock, SignatureError> {
let recorded = self.signature.ok_or(SignatureError::Unmarked)?;
if !READABLE_FORMAT_VERSIONS.contains(&recorded.format_version) {
return Err(SignatureError::UnsupportedFormat {
found: recorded.format_version,
readable: READABLE_FORMAT_VERSIONS,
});
}
let actual = CacheSignature::from_payload(&recorded.model, recorded.layout, &self.layers)?;
recorded.compare(&actual, |field, recorded, actual| {
SignatureError::PayloadMismatch {
field,
recorded,
actual,
}
})?;
expected.compare(&actual, |field, expected, found| {
SignatureError::Incompatible {
field,
expected,
found,
}
})?;
Ok(KvBlock {
signature: actual,
layers: self.layers,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
fn layer(n_kv_heads: usize, head_dim: usize, tokens: usize) -> KvCache {
let mut cache = KvCache::new(n_kv_heads, head_dim);
let step = vec![0.5f32; n_kv_heads * head_dim];
for _ in 0..tokens {
cache.push(&step, &step).expect("unpooled push cannot fail");
}
cache
}
fn payload(n_layers: usize, n_kv_heads: usize, head_dim: usize, tokens: usize) -> Vec<KvCache> {
(0..n_layers)
.map(|_| layer(n_kv_heads, head_dim, tokens))
.collect()
}
fn flat(block_size: usize) -> BlockLayout {
BlockLayout::full_attention(block_size).expect("positive block size")
}
#[test]
fn signature_is_measured_from_the_payload() {
let block = KvBlock::stamp("model-a", flat(4), payload(3, 2, 8, 4)).expect("stamp");
let sig = block.signature();
assert_eq!(sig.n_layers, 3);
assert_eq!(sig.n_kv_heads, 2);
assert_eq!(sig.head_dim, 8);
assert_eq!(sig.tokens, 4);
assert_eq!(sig.dtype, KvDtype::F32);
assert_eq!(sig.format_version, BLOCK_FORMAT_VERSION);
assert_eq!(block.tokens(), 4);
assert_eq!(block.layers().len(), 3);
}
#[test]
fn a_stamped_block_round_trips_through_verification() {
let layers = payload(3, 2, 8, 4);
let signature =
CacheSignature::from_payload("model-a", flat(4), &layers).expect("signature");
let expected = CacheSignature::expected("model-a", flat(4), 3, 2, 8, 4);
let block = UnverifiedBlock::new(Some(signature), layers)
.verify(&expected)
.expect("a block that is what it says it is must verify");
assert_eq!(block.layers().len(), 3);
assert_eq!(block.into_layers().len(), 3);
}
#[test]
fn an_unmarked_block_is_rejected_not_trusted() {
let expected = CacheSignature::expected("model-a", flat(4), 3, 2, 8, 4);
let err = UnverifiedBlock::new(None, payload(3, 2, 8, 4))
.verify(&expected)
.expect_err("an unmarked block must be refused");
assert_eq!(err, SignatureError::Unmarked);
}
#[test]
fn a_signature_that_overstates_its_payload_is_rejected() {
let expected = CacheSignature::expected("model-a", flat(4), 3, 2, 16, 4);
let mut lying = expected.clone();
assert_eq!(lying.head_dim, 16);
let err = UnverifiedBlock::new(Some(lying.clone()), payload(3, 2, 8, 4))
.verify(&expected)
.expect_err("stamp claims head_dim 16 over an 8-wide payload");
assert_eq!(
err,
SignatureError::PayloadMismatch {
field: "head_dim",
recorded: "16".into(),
actual: "8".into(),
}
);
lying.head_dim = 8;
lying.tokens = 8;
let expected = CacheSignature::expected("model-a", flat(8), 3, 2, 8, 8);
let err = UnverifiedBlock::new(Some(lying.clone()), payload(3, 2, 8, 4))
.verify(&expected)
.expect_err("stamp claims 8 tokens over a 4-token payload");
assert_eq!(
err,
SignatureError::PayloadMismatch {
field: "tokens",
recorded: "8".into(),
actual: "4".into(),
}
);
lying.tokens = 4;
lying.n_layers = 4;
let expected = CacheSignature::expected("model-a", flat(4), 4, 2, 8, 4);
let err = UnverifiedBlock::new(Some(lying), payload(3, 2, 8, 4))
.verify(&expected)
.expect_err("stamp claims 4 layers over a 3-layer payload");
assert_eq!(
err,
SignatureError::PayloadMismatch {
field: "n_layers",
recorded: "4".into(),
actual: "3".into(),
}
);
}
#[test]
fn a_layer_whose_seq_len_contradicts_its_buffers_is_rejected() {
let mut layers = payload(2, 2, 8, 4);
layers[1].seq_len = 7;
let err = CacheSignature::from_payload("model-a", flat(4), &layers)
.expect_err("seq_len must be verified, not believed");
assert_eq!(
err,
SignatureError::RaggedPayload {
layer: 1,
field: "seq_len",
expected: "4".into(),
found: "7".into(),
}
);
}
#[test]
fn a_ragged_payload_is_rejected() {
let mut layers = payload(3, 2, 8, 4);
layers[2] = layer(2, 4, 4);
let err = CacheSignature::from_payload("model-a", flat(4), &layers)
.expect_err("shape disagreement");
assert!(matches!(
err,
SignatureError::RaggedPayload {
layer: 2,
field: "layer shape",
..
}
));
let mut layers = payload(3, 2, 8, 4);
layers[1] = layer(2, 8, 3);
let err = CacheSignature::from_payload("model-a", flat(4), &layers)
.expect_err("depth disagreement");
assert!(matches!(
err,
SignatureError::RaggedPayload {
layer: 1,
field: "token count",
..
}
));
let mut layers = payload(2, 2, 8, 4);
layers[0].v.truncate(8);
let err = CacheSignature::from_payload("model-a", flat(4), &layers)
.expect_err("k/v disagreement");
assert!(matches!(
err,
SignatureError::RaggedPayload {
layer: 0,
field: "v length",
..
}
));
}
#[test]
fn an_empty_payload_is_rejected() {
assert_eq!(
CacheSignature::from_payload("model-a", flat(4), &[])
.expect_err("nothing to vouch for"),
SignatureError::EmptyPayload
);
}
#[test]
fn an_honest_block_from_a_different_config_is_incompatible() {
let layers = payload(3, 2, 8, 4);
let signature =
CacheSignature::from_payload("model-a", flat(4), &layers).expect("signature");
let err = UnverifiedBlock::new(Some(signature.clone()), layers)
.verify(&CacheSignature::expected("model-b", flat(4), 3, 2, 8, 4))
.expect_err("a different model must not share KV state");
assert_eq!(
err,
SignatureError::Incompatible {
field: "model",
expected: "model-b".into(),
found: "model-a".into(),
}
);
let layers = payload(3, 2, 8, 4);
let err = UnverifiedBlock::new(Some(signature), layers)
.verify(&CacheSignature::expected("model-a", flat(4), 3, 4, 8, 4))
.expect_err("a different KV head count must not be reused");
assert_eq!(
err,
SignatureError::Incompatible {
field: "n_kv_heads",
expected: "4".into(),
found: "2".into(),
}
);
}
#[test]
fn an_unreadable_format_version_is_rejected() {
let layers = payload(2, 2, 8, 4);
let mut signature =
CacheSignature::from_payload("model-a", flat(4), &layers).expect("signature");
signature.format_version = 99;
let err = UnverifiedBlock::new(Some(signature), layers)
.verify(&CacheSignature::expected("model-a", flat(4), 2, 2, 8, 4))
.expect_err("an unknown layout must not be guessed at");
assert_eq!(
err,
SignatureError::UnsupportedFormat {
found: 99,
readable: READABLE_FORMAT_VERSIONS,
}
);
}
#[test]
fn a_block_written_under_a_different_window_is_refused_not_reused() {
let layout_128 = BlockLayout::new(4, Some(128)).expect("4 divides 128");
let layout_256 = BlockLayout::new(4, Some(256)).expect("4 divides 256");
let layers = payload(3, 2, 8, 4);
let signature =
CacheSignature::from_payload("model-a", layout_128, &layers).expect("signature");
let err = UnverifiedBlock::new(Some(signature.clone()), layers)
.verify(&CacheSignature::expected("model-a", layout_256, 3, 2, 8, 4))
.expect_err("a window change must invalidate the block, not be ignored");
assert_eq!(
err,
SignatureError::Incompatible {
field: "sliding_window",
expected: "256".into(),
found: "128".into(),
}
);
let layers = payload(3, 2, 8, 4);
UnverifiedBlock::new(Some(signature), layers)
.verify(&CacheSignature::expected("model-a", layout_128, 3, 2, 8, 4))
.expect("unchanged config must still hit");
}
#[test]
fn a_full_causal_reader_will_not_take_a_sliding_window_block() {
let sliding = BlockLayout::new(4, Some(128)).expect("aligned");
let layers = payload(2, 2, 8, 4);
let signature =
CacheSignature::from_payload("model-a", sliding, &layers).expect("signature");
let err = UnverifiedBlock::new(Some(signature), layers)
.verify(&CacheSignature::expected("model-a", flat(4), 2, 2, 8, 4))
.expect_err("no window and a 128 window are different configurations");
assert_eq!(
err,
SignatureError::Incompatible {
field: "sliding_window",
expected: "none (full causal)".into(),
found: "128".into(),
}
);
}
#[test]
fn a_block_cut_at_a_different_block_size_is_incompatible() {
let layers = payload(2, 2, 8, 4);
let signature =
CacheSignature::from_payload("model-a", flat(4), &layers).expect("signature");
let err = UnverifiedBlock::new(Some(signature), layers)
.verify(&CacheSignature::expected("model-a", flat(2), 2, 2, 8, 4))
.expect_err("a 4-token block is not a 2-token block");
assert_eq!(
err,
SignatureError::Incompatible {
field: "block_size",
expected: "2".into(),
found: "4".into(),
}
);
}
#[test]
fn a_stamp_may_not_claim_a_block_size_the_payload_lacks() {
let err = KvBlock::stamp("model-a", flat(8), payload(2, 2, 8, 4))
.expect_err("8-token blocks over a 4-token payload");
assert_eq!(
err,
SignatureError::BlockSizeMismatch {
block_size: 8,
tokens: 4,
}
);
let honest =
CacheSignature::from_payload("model-a", flat(4), &payload(2, 2, 8, 4)).expect("sig");
let mut lying = honest.clone();
lying.layout = flat(8);
lying.tokens = 8;
let err = UnverifiedBlock::new(Some(lying), payload(2, 2, 8, 4))
.verify(&CacheSignature::expected("model-a", flat(8), 2, 2, 8, 8))
.expect_err("the payload settles the block size, not the stamp");
assert_eq!(
err,
SignatureError::BlockSizeMismatch {
block_size: 8,
tokens: 4,
}
);
}
#[test]
fn blocks_from_the_pre_layout_format_are_not_readable() {
assert!(!READABLE_FORMAT_VERSIONS.contains(&1));
let layers = payload(2, 2, 8, 4);
let mut signature =
CacheSignature::from_payload("model-a", flat(4), &layers).expect("signature");
signature.format_version = 1;
let err = UnverifiedBlock::new(Some(signature), layers)
.verify(&CacheSignature::expected("model-a", flat(4), 2, 2, 8, 4))
.expect_err("a v1 block cannot say what layout it was cut under");
assert_eq!(
err,
SignatureError::UnsupportedFormat {
found: 1,
readable: READABLE_FORMAT_VERSIONS,
}
);
}
#[test]
fn errors_name_the_field_that_changed() {
let text = SignatureError::Incompatible {
field: "head_dim",
expected: "128".into(),
found: "64".into(),
}
.to_string();
assert!(text.contains("head_dim"), "{text}");
assert!(text.contains("64"), "{text}");
assert!(text.contains("128"), "{text}");
assert!(SignatureError::Unmarked.to_string().contains("unmarked"));
}
}