use crate::kv_cache::AdapterId;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct CrossTurnSlotId(pub u64);
impl CrossTurnSlotId {
pub const DEFAULT: Self = Self(0);
pub const fn new(value: u64) -> Self {
Self(value)
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct CrossTurnPrefixMetadata {
pub model_fingerprint: u64,
pub tokenizer_fingerprint: u64,
pub adapter_id: AdapterId,
pub vocab_size: usize,
pub max_cache_len: usize,
pub kv_f16: bool,
pub rope_theta_bits: u64,
pub partial_rotary_factor_bits: Option<u32>,
pub layer_pattern_hash: u64,
pub chat_template_version: u32,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct KvPrefixHandle {
pub represented_len: usize,
pub num_full_attention_layers: usize,
pub kv_dim: usize,
pub max_cache_len: usize,
pub kv_f16: bool,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct CrossTurnPrefixEntry {
pub slot_id: CrossTurnSlotId,
pub metadata: CrossTurnPrefixMetadata,
pub token_ids: Vec<u32>,
pub represented_len: usize,
pub kv: KvPrefixHandle,
pub gdn_snapshot_len: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum PrefixReuseMode {
FullRefill,
ExactAppend,
ReplayFromCheckpoint { checkpoint_len: usize },
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PrefixRestorePlan {
pub mode: PrefixReuseMode,
pub shared_token_prefix_len: usize,
pub reusable_len: usize,
pub suffix_start: usize,
pub suffix_len: usize,
pub old_represented_len: usize,
}
pub fn longest_common_token_prefix(a: &[u32], b: &[u32]) -> usize {
a.iter().zip(b.iter()).take_while(|(x, y)| x == y).count()
}
pub fn plan_prefix_reuse(
entry: Option<&CrossTurnPrefixEntry>,
metadata: &CrossTurnPrefixMetadata,
new_prompt_ids: &[u32],
sparse_checkpoint_lens: &[usize],
) -> PrefixRestorePlan {
let full_refill = |shared: usize, old_represented_len: usize| PrefixRestorePlan {
mode: PrefixReuseMode::FullRefill,
shared_token_prefix_len: shared,
reusable_len: 0,
suffix_start: 0,
suffix_len: new_prompt_ids.len(),
old_represented_len,
};
let Some(entry) = entry else {
return full_refill(0, 0);
};
if &entry.metadata != metadata || new_prompt_ids.is_empty() {
return full_refill(0, entry.represented_len);
}
let shared = longest_common_token_prefix(&entry.token_ids, new_prompt_ids);
if shared == 0 {
return full_refill(0, entry.represented_len);
}
if shared == entry.represented_len
&& entry.gdn_snapshot_len == entry.represented_len
&& new_prompt_ids.len() > shared
{
return PrefixRestorePlan {
mode: PrefixReuseMode::ExactAppend,
shared_token_prefix_len: shared,
reusable_len: shared,
suffix_start: shared,
suffix_len: new_prompt_ids.len() - shared,
old_represented_len: entry.represented_len,
};
}
if let Some(&checkpoint_len) = sparse_checkpoint_lens
.iter()
.filter(|&&len| len > 0 && len <= shared)
.max()
{
return PrefixRestorePlan {
mode: PrefixReuseMode::ReplayFromCheckpoint { checkpoint_len },
shared_token_prefix_len: shared,
reusable_len: checkpoint_len,
suffix_start: checkpoint_len,
suffix_len: new_prompt_ids.len() - checkpoint_len,
old_represented_len: entry.represented_len,
};
}
full_refill(shared, entry.represented_len)
}
pub fn checkpoint_survives_save(
len: usize,
common_prefix_len: usize,
new_represented_len: usize,
) -> bool {
len > 0 && len <= common_prefix_len && len < new_represented_len
}
#[cfg(test)]
mod tests {
use super::*;
fn metadata() -> CrossTurnPrefixMetadata {
CrossTurnPrefixMetadata {
model_fingerprint: 1,
tokenizer_fingerprint: 2,
adapter_id: AdapterId::BASE,
vocab_size: 1000,
max_cache_len: 4096,
kv_f16: false,
rope_theta_bits: 0,
partial_rotary_factor_bits: None,
layer_pattern_hash: 42,
chat_template_version: 1,
}
}
fn entry(
token_ids: Vec<u32>,
represented_len: usize,
gdn_snapshot_len: usize,
) -> CrossTurnPrefixEntry {
CrossTurnPrefixEntry {
slot_id: CrossTurnSlotId::DEFAULT,
metadata: metadata(),
token_ids,
represented_len,
kv: KvPrefixHandle {
represented_len,
num_full_attention_layers: 6,
kv_dim: 512,
max_cache_len: 4096,
kv_f16: false,
},
gdn_snapshot_len,
}
}
#[test]
fn longest_common_prefix_full_hit() {
assert_eq!(longest_common_token_prefix(&[1, 2, 3], &[1, 2, 3]), 3);
}
#[test]
fn longest_common_prefix_empty_hit() {
assert_eq!(longest_common_token_prefix(&[], &[1, 2, 3]), 0);
assert_eq!(longest_common_token_prefix(&[1, 2, 3], &[]), 0);
assert_eq!(longest_common_token_prefix(&[9], &[1, 2, 3]), 0);
}
#[test]
fn longest_common_prefix_first_divergence() {
assert_eq!(longest_common_token_prefix(&[1, 2, 3], &[1, 2, 9]), 2);
assert_eq!(longest_common_token_prefix(&[1, 9, 3], &[1, 2, 3]), 1);
}
#[test]
fn longest_common_prefix_shorter_old() {
assert_eq!(longest_common_token_prefix(&[1, 2], &[1, 2, 3, 4]), 2);
}
#[test]
fn longest_common_prefix_shorter_new() {
assert_eq!(longest_common_token_prefix(&[1, 2, 3, 4], &[1, 2]), 2);
}
#[test]
fn plan_no_entry_is_full_refill() {
let plan = plan_prefix_reuse(None, &metadata(), &[1, 2, 3], &[]);
assert_eq!(plan.mode, PrefixReuseMode::FullRefill);
assert_eq!(plan.suffix_start, 0);
assert_eq!(plan.suffix_len, 3);
}
#[test]
fn plan_metadata_mismatch_is_full_refill() {
let e = entry(vec![1, 2, 3], 3, 3);
let mut other = metadata();
other.model_fingerprint = 999;
let plan = plan_prefix_reuse(Some(&e), &other, &[1, 2, 3, 4], &[]);
assert_eq!(plan.mode, PrefixReuseMode::FullRefill);
assert_eq!(plan.reusable_len, 0);
}
#[test]
fn plan_empty_new_prompt_is_full_refill() {
let e = entry(vec![1, 2, 3], 3, 3);
let plan = plan_prefix_reuse(Some(&e), &metadata(), &[], &[]);
assert_eq!(plan.mode, PrefixReuseMode::FullRefill);
}
#[test]
fn plan_zero_shared_prefix_is_full_refill() {
let e = entry(vec![1, 2, 3], 3, 3);
let plan = plan_prefix_reuse(Some(&e), &metadata(), &[9, 8, 7], &[]);
assert_eq!(plan.mode, PrefixReuseMode::FullRefill);
assert_eq!(plan.shared_token_prefix_len, 0);
}
#[test]
fn plan_exact_append() {
let e = entry(vec![1, 2, 3], 3, 3);
let plan = plan_prefix_reuse(Some(&e), &metadata(), &[1, 2, 3, 4, 5], &[]);
assert_eq!(plan.mode, PrefixReuseMode::ExactAppend);
assert_eq!(plan.shared_token_prefix_len, 3);
assert_eq!(plan.reusable_len, 3);
assert_eq!(plan.suffix_start, 3);
assert_eq!(plan.suffix_len, 2);
assert_eq!(plan.old_represented_len, 3);
}
#[test]
fn plan_exact_equal_prompt_is_full_refill() {
let e = entry(vec![1, 2, 3], 3, 3);
let plan = plan_prefix_reuse(Some(&e), &metadata(), &[1, 2, 3], &[]);
assert_eq!(plan.mode, PrefixReuseMode::FullRefill);
assert_eq!(plan.shared_token_prefix_len, 3);
assert_eq!(plan.old_represented_len, 3);
assert_eq!(plan.suffix_len, 3);
}
#[test]
fn plan_exact_append_requires_gdn_snapshot_at_boundary() {
let e = entry(vec![1, 2, 3], 3, 2);
let plan = plan_prefix_reuse(Some(&e), &metadata(), &[1, 2, 3, 4], &[]);
assert_eq!(plan.mode, PrefixReuseMode::FullRefill);
}
#[test]
fn plan_mid_history_divergence_without_checkpoint_is_full_refill() {
let e = entry(vec![1, 2, 3], 3, 3);
let plan = plan_prefix_reuse(Some(&e), &metadata(), &[1, 2, 9, 9], &[]);
assert_eq!(plan.mode, PrefixReuseMode::FullRefill);
assert_eq!(plan.shared_token_prefix_len, 2);
assert_eq!(plan.old_represented_len, 3);
}
#[test]
fn plan_mid_history_divergence_shorter_new_prompt_is_full_refill() {
let e = entry(vec![1, 2, 3, 4, 5], 5, 5);
let plan = plan_prefix_reuse(Some(&e), &metadata(), &[1, 2, 3], &[]);
assert_eq!(plan.mode, PrefixReuseMode::FullRefill);
assert_eq!(plan.shared_token_prefix_len, 3);
}
#[test]
fn plan_sparse_checkpoint_replay_only_when_valid() {
let e = entry(vec![1, 2, 3, 4, 5], 5, 5);
let plan = plan_prefix_reuse(Some(&e), &metadata(), &[1, 2, 3, 9, 9], &[3]);
assert_eq!(
plan.mode,
PrefixReuseMode::ReplayFromCheckpoint { checkpoint_len: 3 }
);
assert_eq!(plan.reusable_len, 3);
assert_eq!(plan.suffix_start, 3);
assert_eq!(plan.suffix_len, 2);
}
#[test]
fn plan_sparse_checkpoint_beyond_shared_prefix_is_ignored() {
let e = entry(vec![1, 2, 3, 4, 5], 5, 5);
let plan = plan_prefix_reuse(Some(&e), &metadata(), &[1, 2, 3, 9, 9], &[4]);
assert_eq!(plan.mode, PrefixReuseMode::FullRefill);
}
#[test]
fn plan_picks_the_deepest_valid_checkpoint() {
let e = entry(vec![1, 2, 3, 4, 5], 5, 5);
let plan = plan_prefix_reuse(Some(&e), &metadata(), &[1, 2, 3, 9, 9], &[1, 3]);
assert_eq!(
plan.mode,
PrefixReuseMode::ReplayFromCheckpoint { checkpoint_len: 3 }
);
}
#[test]
fn checkpoint_survival_requires_all_three_conditions() {
assert!(checkpoint_survives_save(3, 5, 8));
assert!(checkpoint_survives_save(5, 5, 8));
assert!(!checkpoint_survives_save(0, 5, 8));
assert!(!checkpoint_survives_save(6, 5, 8));
assert!(!checkpoint_survives_save(8, 8, 8));
assert!(!checkpoint_survives_save(1, 5, 0));
}
}