1use crate::common::{PageID, Position};
2use crate::slice_reader::SliceReader;
3use std::io::{Cursor, Write};
4use umadb_dcb::{DcbError, DcbResult};
5
6#[derive(Debug, Clone, PartialEq, Eq)]
9pub struct TrackingLeafNode {
10 pub keys: Vec<String>,
11 pub values: Vec<Position>,
12}
13
14impl TrackingLeafNode {
15 pub fn new() -> Self {
16 Self {
17 keys: Vec::new(),
18 values: Vec::new(),
19 }
20 }
21
22 pub fn calc_serialized_size(&self) -> usize {
23 let mut size = 1 + 2; for k in &self.keys {
30 size += 1 + k.len();
31 }
32 size += 8 * self.values.len();
33 size
34 }
35
36 pub fn serialize_into(&self, buf: &mut [u8]) -> DcbResult<usize> {
37 let mut cursor = Cursor::new(buf);
38
39 cursor.write_all(&[1])?;
41
42 let klen = self.keys.len() as u16;
44 cursor.write_all(&klen.to_le_bytes())?;
45
46 for k in &self.keys {
48 let kb = k.as_bytes();
49
50 let kb_len = u8::try_from(kb.len()).map_err(|_| {
52 std::io::Error::new(
53 std::io::ErrorKind::InvalidInput,
54 "tracking key too long to serialize (len > 255)",
55 )
56 })?;
57
58 cursor.write_all(&[kb_len])?;
59 cursor.write_all(kb)?;
60 }
61
62 for v in &self.values {
64 cursor.write_all(&v.0.to_le_bytes())?;
65 }
66
67 Ok(cursor.position() as usize)
68 }
69
70 pub fn from_slice(slice: &[u8]) -> DcbResult<Self> {
71 let mut reader = SliceReader::new(slice);
72
73 let ver = reader.read_u8()?;
75
76 let count = if ver == 0 {
78 let cnt_u32 = reader.read_u32()?;
79 u16::try_from(cnt_u32).map_err(|_| {
80 DcbError::DeserializationError("v0 tracking leaf count exceeds u16".to_string())
81 })? as usize
82 } else {
83 reader.read_u16()? as usize
84 };
85
86 let mut keys = Vec::with_capacity(count);
88 for _ in 0..count {
89 let klen = if ver == 0 {
90 let klen_u32 = reader.read_u32()?;
91 u8::try_from(klen_u32).map_err(|_| {
92 DcbError::DeserializationError(
93 "v0 tracking leaf key length exceeds u8".to_string(),
94 )
95 })? as usize
96 } else {
97 reader.read_u8()? as usize
98 };
99
100 let k = reader.read_string(klen)?;
101 keys.push(k);
102 }
103
104 let mut values = Vec::with_capacity(count);
106 for _ in 0..count {
107 values.push(reader.read_position()?);
108 }
109
110 if ver == 0 {
112 let mut pairs: Vec<(String, Position)> = keys.into_iter().zip(values).collect();
114 pairs.sort_by(|a, b| a.0.cmp(&b.0));
115 let (new_keys, new_vals): (Vec<String>, Vec<Position>) = pairs.into_iter().unzip();
116 Ok(Self {
117 keys: new_keys,
118 values: new_vals,
119 })
120 } else {
121 Ok(Self { keys, values })
123 }
124 }
125
126 pub fn get(&self, source: &str) -> Option<Position> {
127 match self.keys.binary_search_by(|k| k.as_str().cmp(source)) {
128 Ok(i) => Some(self.values[i]),
129 Err(_) => None,
130 }
131 }
132
133 pub fn upsert_no_split(
136 &mut self,
137 source: &str,
138 pos: Position,
139 page_body_capacity: usize,
140 ) -> DcbResult<()> {
141 match self.keys.binary_search_by(|k| k.as_str().cmp(source)) {
142 Ok(idx) => {
143 self.values[idx] = pos;
145 Ok(())
146 }
147 Err(ins_idx) => {
148 let key_len = source.len();
150 let additional = key_len + 8;
152 let current = self.calc_serialized_size();
153 if current + additional > page_body_capacity {
154 return Err(DcbError::InternalError(
155 "not implemented: tracking split".to_string(),
156 ));
157 }
158 self.keys.insert(ins_idx, source.to_string());
159 self.values.insert(ins_idx, pos);
160 Ok(())
161 }
162 }
163 }
164}
165
166#[derive(Debug, Clone, PartialEq, Eq)]
167pub struct TrackingInternalNode {
168 pub keys: Vec<String>,
169 pub child_ids: Vec<PageID>,
170}
171
172impl TrackingInternalNode {
173 pub fn child_index_for_key(&self, key: &str) -> usize {
174 match self.keys.binary_search_by(|k| k.as_str().cmp(key)) {
175 Ok(idx) => idx + 1,
176 Err(idx) => idx,
177 }
178 }
179
180 pub fn replace_child_id_at(
181 &mut self,
182 idx: usize,
183 old_id: PageID,
184 new_id: PageID,
185 ) -> DcbResult<()> {
186 if idx >= self.child_ids.len() {
187 return Err(DcbError::DatabaseCorrupted(
188 "child index out of bounds".to_string(),
189 ));
190 }
191 if self.child_ids[idx] != old_id {
192 return Err(DcbError::DatabaseCorrupted("Child ID mismatch".to_string()));
193 }
194 self.child_ids[idx] = new_id;
195 Ok(())
196 }
197
198 pub fn insert_promoted_at(&mut self, idx: usize, key: String, right_child: PageID) {
199 self.keys.insert(idx, key);
200 self.child_ids.insert(idx + 1, right_child);
201 }
202
203 pub fn split_off(&mut self) -> DcbResult<(String, Vec<String>, Vec<PageID>)> {
205 if self.child_ids.len() < 4 || self.keys.len() + 1 != self.child_ids.len() {
206 return Err(DcbError::DatabaseCorrupted(
207 "Cannot split tracking internal with insufficient arity".to_string(),
208 ));
209 }
210 let mid = self.keys.len() / 2; let promoted_key = self.keys[mid].clone();
212 let right_keys: Vec<String> = self.keys[mid + 1..].to_vec();
213 let right_child_ids: Vec<PageID> = self.child_ids[mid + 1..].to_vec();
214 self.keys.truncate(mid);
216 self.child_ids.truncate(mid + 1);
217 Ok((promoted_key, right_keys, right_child_ids))
218 }
219 pub fn new() -> Self {
220 Self {
221 keys: Vec::new(),
222 child_ids: Vec::new(),
223 }
224 }
225
226 pub fn calc_serialized_size(&self) -> usize {
227 let mut size = 1 + 2; for k in &self.keys {
234 size += 1 + k.len();
235 }
236 size += 8 * (self.keys.len() + 1);
237 size
238 }
239
240 pub fn serialize_into(&self, buf: &mut [u8]) -> DcbResult<usize> {
241 let mut cursor = Cursor::new(buf);
242
243 cursor.write_all(&[1])?;
245
246 let klen = self.keys.len() as u16;
248 cursor.write_all(&klen.to_le_bytes())?;
249
250 for k in &self.keys {
252 let kb = k.as_bytes();
253
254 let kb_len = u8::try_from(kb.len()).map_err(|_| {
256 std::io::Error::new(
257 std::io::ErrorKind::InvalidInput,
258 "tracking internal key too long to serialize (len > 255)",
259 )
260 })?;
261
262 cursor.write_all(&[kb_len])?;
263 cursor.write_all(kb)?;
264 }
265
266 for id in &self.child_ids {
268 cursor.write_all(&id.0.to_le_bytes())?;
269 }
270
271 Ok(cursor.position() as usize)
272 }
273
274 pub fn from_slice(slice: &[u8]) -> DcbResult<Self> {
275 let mut reader = SliceReader::new(slice);
276
277 let ver = reader.read_u8()?;
279 if ver == 0 {
280 return Err(DcbError::DeserializationError(
281 "unsupported tracking internal version 0".to_string(),
282 ));
283 }
284
285 let count = reader.read_u16()? as usize;
287
288 let mut keys = Vec::with_capacity(count);
290 for _ in 0..count {
291 let klen = reader.read_u8()? as usize;
292 let k = reader.read_string(klen)?;
293 keys.push(k);
294 }
295
296 let child_count = keys.len() + 1;
298
299 let mut child_ids = Vec::with_capacity(child_count);
301 for _ in 0..child_count {
302 child_ids.push(reader.read_page_id()?);
303 }
304
305 Ok(Self { keys, child_ids })
306 }
307}
308
309#[cfg(test)]
310mod tests {
311 use super::*;
312 use byteorder::{ByteOrder, LittleEndian};
313
314 #[test]
315 fn test_tracking_leaf_roundtrip() {
316 let mut node = TrackingLeafNode::new();
317 node.keys = vec!["a".to_string(), "b".to_string()];
318 node.values = vec![Position(1), Position(2)];
319 let mut buf = vec![0u8; node.calc_serialized_size()];
320 let n = node.serialize_into(&mut buf).unwrap();
321 assert_eq!(n, buf.len());
322 let dec = TrackingLeafNode::from_slice(&buf).unwrap();
323 assert_eq!(node, dec);
324 assert_eq!(dec.get("a"), Some(Position(1)));
325 assert_eq!(dec.get("z"), None);
326 }
327
328 #[test]
329 fn test_deserialize_v0_unsorted_sorts_and_aligns() {
330 let keys = vec!["b", "c", "a"];
332 let values = vec![Position(2), Position(3), Position(1)];
333 let key_bytes: Vec<Vec<u8>> = keys.iter().map(|k| k.as_bytes().to_vec()).collect();
335 let mut size = 1 + 4;
336 for kb in &key_bytes {
337 size += 4 + kb.len();
338 }
339 size += 8 * values.len();
340 let mut buf = vec![0u8; size];
341 buf[0] = 0; LittleEndian::write_u32(&mut buf[1..5], keys.len() as u32);
343 let mut off = 5;
344 for kb in &key_bytes {
345 LittleEndian::write_u32(&mut buf[off..off + 4], kb.len() as u32);
346 off += 4;
347 buf[off..off + kb.len()].copy_from_slice(kb);
348 off += kb.len();
349 }
350 for v in &values {
351 LittleEndian::write_u64(&mut buf[off..off + 8], v.0);
352 off += 8;
353 }
354 let dec = TrackingLeafNode::from_slice(&buf).unwrap();
356 assert_eq!(dec.keys, vec!["a", "b", "c"]);
357 assert_eq!(dec.values, vec![Position(1), Position(2), Position(3)]);
358 assert_eq!(dec.get("a"), Some(Position(1)));
360 assert_eq!(dec.get("c"), Some(Position(3)));
361 assert_eq!(dec.get("z"), None);
362 }
363
364 #[test]
365 fn test_upsert_maintains_sorted_and_capacity_check() {
366 let mut node = TrackingLeafNode::new();
367 let capacity = 1 + 4 + (4 + 1) + (4 + 1) + (4 + 1) + 8 * 3; node.upsert_no_split("b", Position(2), capacity).unwrap();
370 node.upsert_no_split("a", Position(1), capacity).unwrap();
371 node.upsert_no_split("c", Position(3), capacity).unwrap();
372 assert_eq!(
373 node.keys,
374 vec!["a".to_string(), "b".to_string(), "c".to_string()]
375 );
376 assert_eq!(node.get("b"), Some(Position(2)));
377 let mut node2 = TrackingLeafNode::new();
379 let small_capacity = 1 + 4; let err = node2
381 .upsert_no_split("x", Position(9), small_capacity)
382 .unwrap_err();
383 match err {
384 DcbError::InternalError(s) => assert!(s.contains("tracking split")),
385 _ => panic!("unexpected error type"),
386 }
387 }
388
389 #[test]
390 fn test_tracking_internal_roundtrip() {
391 let node = TrackingInternalNode {
392 keys: vec!["alpha".into(), "beta".into(), "gamma".into()],
393 child_ids: vec![PageID(10), PageID(20), PageID(30), PageID(40)],
394 };
395 let mut buf = vec![0u8; node.calc_serialized_size()];
396 let n = node.serialize_into(&mut buf).unwrap();
397 assert_eq!(n, buf.len());
398 let dec = TrackingInternalNode::from_slice(&buf).unwrap();
399 assert_eq!(dec, node);
400 }
401
402 #[test]
403 fn test_tracking_internal_empty_keys_one_child_roundtrip() {
404 let node = TrackingInternalNode {
405 keys: vec![],
406 child_ids: vec![PageID(123)],
407 };
408 let mut buf = vec![0u8; node.calc_serialized_size()];
409 let n = node.serialize_into(&mut buf).unwrap();
410 assert_eq!(n, buf.len());
411 let dec = TrackingInternalNode::from_slice(&buf).unwrap();
412 assert_eq!(dec, node);
413 }
414
415 #[test]
416 fn test_tracking_internal_from_slice_truncated_children_err() {
417 let node = TrackingInternalNode {
419 keys: vec!["k1".into()],
420 child_ids: vec![PageID(1), PageID(2)],
421 };
422 let mut buf = vec![0u8; node.calc_serialized_size()];
423 let _ = node.serialize_into(&mut buf).unwrap();
424 buf.truncate(buf.len() - 4); let err = TrackingInternalNode::from_slice(&buf).unwrap_err();
426 match err {
427 DcbError::DeserializationError(_) => {}
428 _ => panic!("expected DeserializationError"),
429 }
430 }
431}