1use alloc::vec::Vec;
33
34#[cfg(feature = "counters")]
35use core::cell::Cell;
36
37use crate::error::Error;
38use crate::id::FactId;
39
40const FACT_BYTES: usize = core::mem::size_of::<u32>();
42
43const SCALE_BYTES: usize = core::mem::size_of::<f32>();
45
46const HEAD: usize = FACT_BYTES + SCALE_BYTES;
48
49const SIG_WORD_BYTES: usize = core::mem::size_of::<u64>();
52
53#[derive(Debug, Default)]
56pub struct VecScratch {
57 cand: Vec<(u32, u32)>,
59 top: Vec<(f32, u32)>,
61 query: Vec<u8>,
63}
64
65impl VecScratch {
66 pub fn new() -> Self {
68 Self::default()
69 }
70}
71
72#[derive(Debug)]
83pub struct VecPool<'a> {
84 base: &'a [u8],
86 tail: Vec<u8>,
88 dim: usize,
89 max_bytes: usize,
90 #[cfg(feature = "counters")]
92 dots: Cell<u64>,
93}
94
95impl<'a> VecPool<'a> {
96 #[inline]
98 fn words(dim: usize) -> usize {
99 dim.div_ceil(64)
100 }
101
102 pub fn new(dim: usize, max_bytes: usize) -> Self {
105 Self {
106 base: &[],
107 tail: Vec::new(),
108 dim,
109 max_bytes,
110 #[cfg(feature = "counters")]
111 dots: Cell::new(0),
112 }
113 }
114
115 #[inline]
117 pub fn stride(&self) -> usize {
118 HEAD + Self::words(self.dim) * SIG_WORD_BYTES + self.dim
119 }
120
121 #[inline]
123 fn pool_len(&self) -> usize {
124 self.base.len() + self.tail.len()
125 }
126
127 #[inline]
132 pub(crate) fn slot_bytes(&self, i: usize) -> &[u8] {
133 let stride = self.stride();
134 let start = i * stride;
135 let base_len = self.base.len();
136 if start < base_len {
137 &self.base[start..start + stride]
138 } else {
139 let at = start - base_len;
140 &self.tail[at..at + stride]
141 }
142 }
143
144 #[inline]
146 pub fn len(&self) -> usize {
147 let pool_len = self.pool_len();
148 if pool_len == 0 {
149 0
150 } else {
151 pool_len / self.stride()
152 }
153 }
154
155 pub fn is_empty(&self) -> bool {
157 self.pool_len() == 0
158 }
159
160 pub fn pool_bytes(&self) -> usize {
162 self.pool_len()
163 }
164
165 #[inline]
167 pub fn slot_fact(&self, i: usize) -> u32 {
168 let slot = self.slot_bytes(i);
169 u32::from_le_bytes(slot[..FACT_BYTES].try_into().unwrap())
170 }
171
172 #[inline]
174 fn slot_scale(&self, i: usize) -> f32 {
175 let slot = self.slot_bytes(i);
176 f32::from_le_bytes(slot[FACT_BYTES..HEAD].try_into().unwrap())
177 }
178
179 #[inline]
183 pub(crate) fn quant(&self, i: usize) -> (f32, &[u8]) {
184 let stride = self.stride();
185 let q_off = HEAD + Self::words(self.dim) * SIG_WORD_BYTES;
186 let slot = self.slot_bytes(i);
187 let scale = f32::from_le_bytes(slot[FACT_BYTES..HEAD].try_into().unwrap());
188 (scale, &slot[q_off..stride])
189 }
190
191 #[inline]
195 pub(crate) fn sim(&self, a: u32, b: u32) -> f32 {
196 self.cosine_at(a as usize, b as usize)
197 }
198
199 fn encode_slot(&self, fact: u32, v: &[f32], out: &mut [u8]) -> Result<(), Error> {
207 if v.len() != self.dim {
208 return Err(Error::DimMismatch {
209 got: v.len(),
210 want: self.dim,
211 });
212 }
213 let mut norm_sq = 0.0f32;
214 for &x in v {
215 if !x.is_finite() {
216 return Err(Error::Invalid("vector must be finite"));
217 }
218 norm_sq += x * x;
219 }
220 let norm = libm::sqrtf(norm_sq);
223 if norm <= 0.0 {
224 return Err(Error::Invalid("vector must be nonzero"));
225 }
226 let inv_norm = 1.0 / norm;
227 let mut max_abs = 0.0f32;
228 for &x in v {
229 max_abs = max_abs.max(libm::fabsf(x * inv_norm));
230 }
231 let scale = max_abs / 127.0;
233 out[..FACT_BYTES].copy_from_slice(&fact.to_le_bytes());
234 out[FACT_BYTES..HEAD].copy_from_slice(&scale.to_le_bytes());
235 let words = Self::words(self.dim);
236 let q_off = HEAD + words * SIG_WORD_BYTES;
237 for (i, &x) in v.iter().enumerate() {
238 let qf = libm::roundf((x * inv_norm) / scale);
239 let qi = qf.clamp(-127.0, 127.0) as i32 as i8;
240 out[q_off + i] = qi as u8;
241 }
242 for w in 0..words {
244 let mut word = 0u64;
245 for b in 0..64 {
246 let i = w * 64 + b;
247 if i >= self.dim {
248 break;
249 }
250 if out[q_off + i] as i8 >= 0 {
251 word |= 1 << b;
252 }
253 }
254 out[HEAD + w * SIG_WORD_BYTES..HEAD + w * SIG_WORD_BYTES + SIG_WORD_BYTES]
255 .copy_from_slice(&word.to_le_bytes());
256 }
257 Ok(())
258 }
259
260 pub(crate) fn encode_slot_into(
265 &self,
266 fact: FactId,
267 v: &[f32],
268 out: &mut Vec<u8>,
269 ) -> Result<(), Error> {
270 out.clear();
271 out.resize(self.stride(), 0);
272 self.encode_slot(fact.0, v, out)
273 }
274
275 pub(crate) fn cosine_encoded_slot(&self, encoded: &[u8], slot: u32) -> f32 {
279 let stride = self.stride();
280 if encoded.len() != stride || slot as usize >= self.len() {
281 return 0.0;
282 }
283 let q_off = HEAD + Self::words(self.dim) * SIG_WORD_BYTES;
284 let stored = self.slot_bytes(slot as usize);
285 let query_scale = f32::from_le_bytes(encoded[FACT_BYTES..HEAD].try_into().unwrap());
286 let dot = dot_i8(&encoded[q_off..stride], &stored[q_off..stride]);
287 query_scale * self.slot_scale(slot as usize) * dot as f32
288 }
289
290 pub fn push(&mut self, fact: FactId, v: &[f32]) -> Result<u32, Error> {
298 let stride = self.stride();
299 let pool_len = self.pool_len();
300 if pool_len + stride > self.max_bytes {
301 return Err(Error::CapacityExceeded { what: "vectors" });
302 }
303 let index = u32::try_from(pool_len / stride).map_err(|_| Error::CapacityExceeded {
304 what: "vector slots",
305 })?;
306 let mut tail = core::mem::take(&mut self.tail);
310 let at = tail.len();
311 tail.resize(at + stride, 0);
312 let res = match self.encode_slot(fact.0, v, &mut tail[at..]) {
313 Ok(()) => Ok(index),
314 Err(e) => {
315 tail.truncate(at);
317 Err(e)
318 }
319 };
320 self.tail = tail;
321 res
322 }
323
324 pub(crate) fn quantized<'s>(&self, scratch: &'s VecScratch) -> (f32, &'s [u8]) {
328 let stride = self.stride();
329 let q_off = HEAD + Self::words(self.dim) * SIG_WORD_BYTES;
330 debug_assert_eq!(scratch.query.len(), stride);
331 (
332 f32::from_le_bytes(scratch.query[FACT_BYTES..HEAD].try_into().unwrap()),
333 &scratch.query[q_off..stride],
334 )
335 }
336
337 pub fn quantize_query(&self, v: &[f32], scratch: &mut VecScratch) -> Result<(), Error> {
340 let stride = self.stride();
341 scratch.query.clear();
342 scratch.query.resize(stride, 0);
343 let mut buf = core::mem::take(&mut scratch.query);
344 let res = self.encode_slot(0, v, &mut buf);
345 scratch.query = buf;
346 res
347 }
348
349 pub(crate) fn copy_slot(&mut self, src: &VecPool<'_>, i: u32) -> u32 {
354 debug_assert_eq!(self.dim, src.dim, "copy_slot across differing dims");
355 let stride = self.stride();
356 let index = (self.pool_len() / stride) as u32;
357 self.tail.extend_from_slice(src.slot_bytes(i as usize));
358 index
359 }
360
361 pub(crate) fn clone_slot_for_fact(&mut self, fact: FactId, source: u32) -> Result<u32, Error> {
365 let stride = self.stride();
366 let pool_len = self.pool_len();
367 if source as usize >= self.len() {
368 return Err(Error::Corrupt("retag vector slot is out of range"));
369 }
370 if pool_len + stride > self.max_bytes {
371 return Err(Error::CapacityExceeded { what: "vectors" });
372 }
373 let index = u32::try_from(pool_len / stride).map_err(|_| Error::CapacityExceeded {
374 what: "vector slots",
375 })?;
376 let source_start = source as usize * stride;
377 if source_start < self.base.len() {
378 self.tail
379 .extend_from_slice(&self.base[source_start..source_start + stride]);
380 } else {
381 let at = source_start - self.base.len();
382 let dst = self.tail.len();
383 self.tail.resize(dst + stride, 0);
384 self.tail.copy_within(at..at + stride, dst);
385 }
386 let at = self.tail.len() - stride;
387 self.tail[at..at + FACT_BYTES].copy_from_slice(&fact.0.to_le_bytes());
388 Ok(index)
389 }
390
391 fn cosine_at(&self, a: usize, b: usize) -> f32 {
394 let stride = self.stride();
395 let q_off = HEAD + Self::words(self.dim) * SIG_WORD_BYTES;
396 let (sa, sb) = (self.slot_bytes(a), self.slot_bytes(b));
397 let dot = dot_i8(&sa[q_off..stride], &sb[q_off..stride]);
398 self.slot_scale(a) * self.slot_scale(b) * dot as f32
399 }
400
401 pub fn cosine_slots(&self, a: u32, b: u32) -> f32 {
404 let n = self.len();
405 if a as usize >= n || b as usize >= n {
406 return 0.0;
407 }
408 self.cosine_at(a as usize, b as usize)
409 }
410
411 pub fn search(
416 &self,
417 query: &[f32],
418 k: usize,
419 admit: &mut dyn FnMut(FactId) -> bool,
420 scratch: &mut VecScratch,
421 out: &mut Vec<(FactId, f32)>,
422 ) -> Result<(), Error> {
423 out.clear();
424 let n = self.len();
425 if n == 0 || k == 0 {
426 return Ok(());
427 }
428 self.quantize_query(query, scratch)?;
429 let stride = self.stride();
430 let words = Self::words(self.dim);
431 let q_off = HEAD + words * SIG_WORD_BYTES;
432
433 let VecScratch {
448 cand, top, query, ..
449 } = scratch;
450 let q_sig = &query[HEAD..HEAD + words * SIG_WORD_BYTES];
451 let c = k.saturating_mul(4).max(64).min(n);
455 let cap = c * 2;
458 cand.clear();
459 cand.reserve(cap.min(n));
462 let mut limit: Option<(u32, u32)> = None;
463 let slots = self
464 .base
465 .chunks_exact(stride)
466 .chain(self.tail.chunks_exact(stride));
467 for (i, slot) in slots.enumerate() {
468 let s_sig = &slot[HEAD..HEAD + words * SIG_WORD_BYTES];
469 let mut ham = 0u32;
470 for (qw, sw) in q_sig
471 .chunks_exact(SIG_WORD_BYTES)
472 .zip(s_sig.chunks_exact(SIG_WORD_BYTES))
473 {
474 let a = u64::from_le_bytes(qw.try_into().unwrap());
475 let b = u64::from_le_bytes(sw.try_into().unwrap());
476 ham += (a ^ b).count_ones();
477 }
478 let entry = (ham, i as u32);
479 if limit.is_some_and(|worst| entry >= worst) {
480 continue;
481 }
482 cand.push(entry);
483 if cand.len() == cap {
484 cand.select_nth_unstable(c - 1);
485 cand.truncate(c);
486 limit = Some(cand[c - 1]);
487 }
488 }
489 if cand.len() > c {
490 cand.select_nth_unstable(c - 1);
491 }
492
493 let q_scale = f32::from_le_bytes(query[FACT_BYTES..HEAD].try_into().unwrap());
495 let q_q = &query[q_off..q_off + self.dim];
496 top.clear();
497 #[cfg(feature = "counters")]
498 let mut dots = 0u64;
499 for &(_, slot) in cand[..c].iter() {
500 let sb = self.slot_bytes(slot as usize);
501 let fact = FactId(u32::from_le_bytes(sb[..FACT_BYTES].try_into().unwrap()));
502 if !admit(fact) {
503 continue;
504 }
505 let s_scale = f32::from_le_bytes(sb[FACT_BYTES..HEAD].try_into().unwrap());
506 let dot = dot_i8(q_q, &sb[q_off..stride]);
507 top.push((q_scale * s_scale * dot as f32, fact.0));
508 #[cfg(feature = "counters")]
509 {
510 dots += 1;
511 }
512 }
513 #[cfg(feature = "counters")]
514 self.dots.set(self.dots.get() + dots);
515 let order = |a: &(f32, u32), b: &(f32, u32)| b.0.total_cmp(&a.0).then(a.1.cmp(&b.1));
520 let band = k.min(top.len());
521 if band < top.len() {
522 top.select_nth_unstable_by(band, order);
523 }
524 top[..band].sort_unstable_by(order);
525 for &(score, id) in top.iter().take(k) {
526 out.push((FactId(id), score));
527 }
528 Ok(())
529 }
530
531 #[cfg(test)]
536 pub(crate) fn dump(&self) -> Vec<u8> {
537 let mut out = Vec::with_capacity(self.pool_len());
538 out.extend_from_slice(self.base);
539 out.extend_from_slice(&self.tail);
540 out
541 }
542
543 pub(crate) fn pieces(&self) -> [&[u8]; 2] {
548 [self.base, &self.tail]
549 }
550
551 pub(crate) fn from_parts(dim: usize, max_bytes: usize, bytes: &[u8]) -> Result<Self, Error> {
556 Self::frame_check(dim, max_bytes, bytes.len())?;
557 let mut pool = Self::new(dim, max_bytes);
558 pool.tail = bytes.to_vec();
559 Ok(pool)
560 }
561
562 pub(crate) fn from_parts_borrowed(
568 dim: usize,
569 max_bytes: usize,
570 bytes: &'a [u8],
571 ) -> Result<Self, Error> {
572 Self::frame_check(dim, max_bytes, bytes.len())?;
573 let mut pool = Self::new(dim, max_bytes);
574 pool.base = bytes;
575 Ok(pool)
576 }
577
578 fn frame_check(dim: usize, max_bytes: usize, len: usize) -> Result<(), Error> {
582 if len > max_bytes {
583 return Err(Error::Corrupt("vector pool exceeds the configured ceiling"));
584 }
585 if dim == 0 {
586 if len != 0 {
587 return Err(Error::Corrupt("vector pool present with dim 0"));
588 }
589 return Ok(());
590 }
591 let stride = HEAD + Self::words(dim) * SIG_WORD_BYTES + dim;
592 if !len.is_multiple_of(stride) {
593 return Err(Error::Corrupt("vector pool is not a whole number of slots"));
594 }
595 Ok(())
596 }
597
598 pub(crate) fn validate(&self) -> Result<(), Error> {
604 if self.dim == 0 {
605 return Ok(());
606 }
607 let words = Self::words(self.dim);
608 let q_off = HEAD + words * SIG_WORD_BYTES;
609 for i in 0..self.len() {
610 let slot = self.slot_bytes(i);
611 let scale = f32::from_le_bytes(slot[FACT_BYTES..HEAD].try_into().unwrap());
612 if !scale.is_finite() || scale < 0.0 {
613 return Err(Error::Corrupt(
614 "vector slot scale is not finite and non-negative",
615 ));
616 }
617 for w in 0..words {
618 let stored = u64::from_le_bytes(
619 slot[HEAD + w * SIG_WORD_BYTES..HEAD + w * SIG_WORD_BYTES + SIG_WORD_BYTES]
620 .try_into()
621 .unwrap(),
622 );
623 let mut expect = 0u64;
624 for b in 0..64 {
625 let j = w * 64 + b;
626 if j >= self.dim {
627 break;
628 }
629 if slot[q_off + j] as i8 >= 0 {
630 expect |= 1 << b;
631 }
632 }
633 if stored != expect {
634 return Err(Error::Corrupt(
635 "vector slot signature disagrees with its components",
636 ));
637 }
638 }
639 }
640 Ok(())
641 }
642
643 #[cfg(feature = "counters")]
645 pub fn dots(&self) -> u64 {
646 self.dots.get()
647 }
648
649 #[cfg(feature = "counters")]
651 pub fn reset_dots(&self) {
652 self.dots.set(0);
653 }
654}
655
656const DOT_LANES: usize = 16;
662
663#[inline]
672pub(crate) fn dot_i8(a: &[u8], b: &[u8]) -> i32 {
673 let mut lanes = [0i32; DOT_LANES];
674 let mut a_chunks = a.chunks_exact(DOT_LANES);
675 let mut b_chunks = b.chunks_exact(DOT_LANES);
676 for (x, y) in a_chunks.by_ref().zip(b_chunks.by_ref()) {
677 for (lane, (&x, &y)) in lanes.iter_mut().zip(x.iter().zip(y.iter())) {
678 *lane += i32::from(x as i8) * i32::from(y as i8);
679 }
680 }
681 let mut acc: i32 = lanes.iter().sum();
682 for (&x, &y) in a_chunks.remainder().iter().zip(b_chunks.remainder().iter()) {
683 acc += i32::from(x as i8) * i32::from(y as i8);
684 }
685 acc
686}
687
688#[cfg(test)]
689mod tests {
690 use super::*;
691 use alloc::vec;
692
693 struct Lcg(u64);
696 impl Lcg {
697 fn next(&mut self) -> f32 {
698 self.0 = self
699 .0
700 .wrapping_mul(6_364_136_223_846_793_005)
701 .wrapping_add(1_442_695_040_888_963_407);
702 ((self.0 >> 40) as f32 / (1u64 << 24) as f32) * 2.0 - 1.0
703 }
704 fn vector(&mut self, dim: usize) -> Vec<f32> {
705 (0..dim).map(|_| self.next()).collect()
706 }
707 }
708
709 fn cosine_f32(a: &[f32], b: &[f32]) -> f32 {
711 let dot: f32 = a.iter().zip(b).map(|(x, y)| x * y).sum();
712 let na: f32 = libm::sqrtf(a.iter().map(|x| x * x).sum());
713 let nb: f32 = libm::sqrtf(b.iter().map(|x| x * x).sum());
714 dot / (na * nb)
715 }
716
717 #[test]
720 fn quantized_cosine_tracks_f32() {
721 let dim = 384;
722 let mut rng = Lcg(0x1234_5678);
723 let mut worst = 0.0f32;
724 for i in 0..200u32 {
725 let a = rng.vector(dim);
726 let b = rng.vector(dim);
727 let mut pool = VecPool::new(dim, usize::MAX);
728 pool.push(FactId(2 * i), &a).unwrap();
729 pool.push(FactId(2 * i + 1), &b).unwrap();
730 let q = pool.cosine_slots(0, 1);
731 let t = cosine_f32(&a, &b);
732 worst = worst.max(libm::fabsf(q - t));
733 }
734 assert!(
735 worst < 0.05,
736 "worst quantization error {worst} exceeds 0.05"
737 );
738 }
739
740 #[test]
744 fn golden_dim4() {
745 let dim = 4;
746 let mut pool = VecPool::new(dim, usize::MAX);
747 pool.push(FactId(0), &[1.0, 1.0, 0.0, 0.0]).unwrap();
748 pool.push(FactId(1), &[2.0, 2.0, 0.0, 0.0]).unwrap(); pool.push(FactId(2), &[0.0, 0.0, 1.0, 1.0]).unwrap(); assert!((pool.cosine_slots(0, 1) - 1.0).abs() < 1e-3);
752 assert!(pool.cosine_slots(0, 2).abs() < 1e-3);
754 pool.validate().unwrap();
756 let stride = pool.stride();
759 let sig = u64::from_le_bytes(pool.dump()[HEAD..HEAD + SIG_WORD_BYTES].try_into().unwrap());
760 assert_eq!(sig & 0b1111, 0b1111);
761 assert_eq!(pool.len(), 3);
762 assert_eq!(stride, HEAD + SIG_WORD_BYTES + 4);
764 }
765
766 #[test]
768 fn search_surfaces_the_nearest() {
769 let dim = 64;
770 let mut rng = Lcg(0xdead_beef);
771 let mut pool = VecPool::new(dim, usize::MAX);
772 let target = rng.vector(dim);
773 for i in 0..200u32 {
775 pool.push(FactId(i), &rng.vector(dim)).unwrap();
776 }
777 pool.push(FactId(500), &target).unwrap();
778 let mut scratch = VecScratch::new();
779 let mut out = Vec::new();
780 pool.search(&target, 5, &mut |_| true, &mut scratch, &mut out)
781 .unwrap();
782 assert_eq!(out[0].0, FactId(500), "exact match must rank first");
783 assert!(out[0].1 > 0.99, "self-cosine ≈ 1, got {}", out[0].1);
784 }
785
786 #[test]
788 fn degenerate_vectors_are_invalid() {
789 let mut pool = VecPool::new(3, usize::MAX);
790 assert_eq!(
791 pool.push(FactId(0), &[0.0, 0.0, 0.0]).unwrap_err(),
792 Error::Invalid("vector must be nonzero")
793 );
794 assert_eq!(
795 pool.push(FactId(0), &[1.0, f32::NAN, 0.0]).unwrap_err(),
796 Error::Invalid("vector must be finite")
797 );
798 assert!(matches!(
799 pool.push(FactId(0), &[1.0, 2.0]).unwrap_err(),
800 Error::DimMismatch { got: 2, want: 3 }
801 ));
802 assert_eq!(pool.len(), 0);
804 assert!(pool.is_empty());
805 }
806
807 #[test]
810 fn accessors_and_edges() {
811 let dim = 4;
812 let mut pool = VecPool::new(dim, usize::MAX);
813 assert!(pool.is_empty());
814 assert_eq!(pool.pool_bytes(), 0);
815 let mut scratch = VecScratch::new();
816 let mut out = vec![(FactId(9), 1.0)];
817 pool.search(&[1.0; 4], 5, &mut |_| true, &mut scratch, &mut out)
819 .unwrap();
820 assert!(out.is_empty());
821
822 pool.push(FactId(0), &[1.0, 0.0, 0.0, 0.0]).unwrap();
823 assert!(!pool.is_empty());
824 assert_eq!(pool.pool_bytes(), pool.stride());
825 pool.search(&[1.0; 4], 0, &mut |_| true, &mut scratch, &mut out)
827 .unwrap();
828 assert!(out.is_empty());
829 assert_eq!(pool.cosine_slots(0, 9), 0.0);
831
832 let mut tight = VecPool::new(dim, 4);
834 assert_eq!(
835 tight.push(FactId(0), &[1.0, 0.0, 0.0, 0.0]).unwrap_err(),
836 Error::CapacityExceeded { what: "vectors" }
837 );
838 }
839
840 #[test]
842 fn from_parts_frames_slots() {
843 let dim = 8;
844 let mut pool = VecPool::new(dim, usize::MAX);
845 pool.push(FactId(0), &vec![0.5; dim]).unwrap();
846 pool.push(FactId(1), &vec![-0.5; dim]).unwrap();
847 let bytes = pool.dump();
848 let rebuilt = VecPool::from_parts(dim, usize::MAX, &bytes).unwrap();
849 assert_eq!(rebuilt.len(), 2);
850 rebuilt.validate().unwrap();
851 assert!(VecPool::from_parts(dim, usize::MAX, &bytes[..bytes.len() - 1]).is_err());
853 assert!(VecPool::from_parts(0, usize::MAX, &bytes).is_err());
855 assert!(VecPool::from_parts(dim, bytes.len() - 1, &bytes).is_err());
857 }
858
859 #[test]
862 fn validate_rejects_malformed_slots() {
863 let dim = 8;
864 let mut pool = VecPool::new(dim, usize::MAX);
865 pool.push(FactId(0), &vec![0.5; dim]).unwrap();
866 let good = pool.dump();
867
868 let mut bad = good.clone();
870 bad[FACT_BYTES..HEAD].copy_from_slice(&f32::NAN.to_le_bytes());
871 assert!(
872 VecPool::from_parts(dim, usize::MAX, &bad)
873 .unwrap()
874 .validate()
875 .is_err()
876 );
877
878 let mut bad = good.clone();
881 let q_off = HEAD + VecPool::words(dim) * SIG_WORD_BYTES;
882 bad[q_off] = (-1i8) as u8; assert!(
884 VecPool::from_parts(dim, usize::MAX, &bad)
885 .unwrap()
886 .validate()
887 .is_err()
888 );
889 }
890
891 #[test]
896 fn overlay_appends_to_tail_and_reads_span_the_boundary() {
897 let dim = 16;
898 let mut rng = Lcg(0x0ace_1a75);
899 let (va, vb, vc) = (rng.vector(dim), rng.vector(dim), rng.vector(dim));
900
901 let mut owned = VecPool::new(dim, usize::MAX);
903 owned.push(FactId(10), &va).unwrap();
904 owned.push(FactId(11), &vb).unwrap();
905 owned.push(FactId(12), &vc).unwrap();
906
907 let mut seed = VecPool::new(dim, usize::MAX);
910 seed.push(FactId(10), &va).unwrap();
911 seed.push(FactId(11), &vb).unwrap();
912 let base = seed.dump();
913 let base_snapshot = base.clone();
914
915 let mut pool = VecPool::from_parts_borrowed(dim, usize::MAX, &base).unwrap();
918 assert_eq!(pool.len(), 2);
919 let idx = pool.push(FactId(12), &vc).unwrap();
920 assert_eq!(idx, 2);
921 assert_eq!(pool.len(), 3);
922
923 assert_eq!(pool.slot_fact(0), 10); assert_eq!(pool.slot_fact(2), 12); assert!((pool.cosine_slots(0, 2) - owned.cosine_slots(0, 2)).abs() < 1e-6);
929 pool.validate().unwrap();
930
931 let mut scratch = VecScratch::new();
933 let mut out = Vec::new();
934 pool.search(&vc, 1, &mut |_| true, &mut scratch, &mut out)
935 .unwrap();
936 assert_eq!(out[0].0, FactId(12));
937
938 assert_eq!(pool.dump(), owned.dump());
940 assert_eq!(base, base_snapshot);
941 }
942}