qualia_core_db/tensor/
kv_provenance.rs1use std::sync::{OnceLock, RwLock};
4
5pub const MAX_KV_PROVENANCE: usize = 1024;
6pub const KV_SLOT_UNMAPPED: u32 = u32::MAX;
7
8#[derive(Clone, Copy, Debug)]
9pub struct KvSlotInfo {
10 pub kv_slot: u32,
11 pub parent_page_id: u64,
12}
13
14impl KvSlotInfo {
15 pub const UNMAPPED: Self = Self {
16 kv_slot: KV_SLOT_UNMAPPED,
17 parent_page_id: 0,
18 };
19}
20
21pub struct KvProvenanceMap {
22 tensor_to_kv: [KvSlotInfo; MAX_KV_PROVENANCE],
23 generation: u32,
24}
25
26impl KvProvenanceMap {
27 pub fn new() -> Self {
28 Self {
29 tensor_to_kv: [KvSlotInfo::UNMAPPED; MAX_KV_PROVENANCE],
30 generation: 0,
31 }
32 }
33
34 #[inline]
35 pub fn generation(&self) -> u32 {
36 self.generation
37 }
38
39 pub fn record(&mut self, tensor_index: u32, kv_slot: u32, parent_page_id: u64) {
41 let ti = tensor_index as usize;
42 if ti < MAX_KV_PROVENANCE && kv_slot < MAX_KV_PROVENANCE as u32 {
43 self.tensor_to_kv[ti] = KvSlotInfo {
44 kv_slot,
45 parent_page_id,
46 };
47 self.generation = self.generation.wrapping_add(1);
48 }
49 }
50
51 pub fn build_prompt_alignment(
53 &mut self,
54 prompt_token_count: u32,
55 tensor_node_count: u32,
56 root_page_id: u64,
57 ) {
58 let n = prompt_token_count
59 .min(tensor_node_count)
60 .min(MAX_KV_PROVENANCE as u32);
61 for i in 0..n {
62 self.tensor_to_kv[i as usize] = KvSlotInfo {
63 kv_slot: i,
64 parent_page_id: root_page_id,
65 };
66 }
67 for i in n as usize..MAX_KV_PROVENANCE {
68 self.tensor_to_kv[i] = KvSlotInfo::UNMAPPED;
69 }
70 self.generation = self.generation.wrapping_add(1);
71 }
72
73 #[inline]
74 pub fn kv_slot_for_tensor(&self, tensor_index: u32) -> Option<u32> {
75 let ti = tensor_index as usize;
76 if ti >= MAX_KV_PROVENANCE {
77 return None;
78 }
79 let info = self.tensor_to_kv[ti];
80 if info.kv_slot == KV_SLOT_UNMAPPED {
81 None
82 } else {
83 Some(info.kv_slot)
84 }
85 }
86
87 #[inline]
88 pub fn page_id_for_tensor(&self, tensor_index: u32) -> Option<u64> {
89 let ti = tensor_index as usize;
90 if ti >= MAX_KV_PROVENANCE {
91 return None;
92 }
93 let info = self.tensor_to_kv[ti];
94 if info.kv_slot == KV_SLOT_UNMAPPED {
95 None
96 } else {
97 Some(info.parent_page_id)
98 }
99 }
100}
101
102static KV_PROVENANCE: OnceLock<RwLock<KvProvenanceMap>> = OnceLock::new();
103
104fn kv_lock() -> &'static RwLock<KvProvenanceMap> {
105 KV_PROVENANCE.get_or_init(|| RwLock::new(KvProvenanceMap::new()))
106}
107
108#[inline]
109pub fn global_kv_provenance() -> std::sync::RwLockReadGuard<'static, KvProvenanceMap> {
110 kv_lock().read().expect("kv provenance poisoned")
111}
112
113pub fn rebuild_prompt_provenance(
115 prompt_token_count: u32,
116 tensor_node_count: u32,
117 root_page_id: u64,
118) {
119 kv_lock()
120 .write()
121 .expect("kv provenance poisoned")
122 .build_prompt_alignment(prompt_token_count, tensor_node_count, root_page_id);
123}
124
125#[inline]
126pub fn record_kv_provenance(tensor_index: u32, kv_slot: u32, parent_page_id: u64) {
127 kv_lock().write().expect("kv provenance poisoned").record(
128 tensor_index,
129 kv_slot,
130 parent_page_id,
131 );
132}
133
134#[cfg(test)]
135mod tests {
136 use super::*;
137
138 #[test]
139 fn prompt_alignment_maps_tensor_to_kv() {
140 let mut map = KvProvenanceMap::new();
141 map.build_prompt_alignment(8, 10, 42);
142 assert_eq!(map.kv_slot_for_tensor(3), Some(3));
143 assert_eq!(map.page_id_for_tensor(3), Some(42));
144
145 assert_eq!(map.kv_slot_for_tensor(9), None);
146 }
147}