1use yo_common::{Code, Error, Result};
65use yo_shape::Metric;
66
67use crate::db::Handle;
68
69pub use yo_vector::Match;
70
71#[derive(Clone)]
76pub struct Vectors {
77 pub(crate) db: Handle,
78 pub(crate) at: usize,
79 pub(crate) dim: usize,
80 pub(crate) metric: Metric,
81}
82
83impl core::fmt::Debug for Vectors {
84 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
85 f.debug_struct("Vectors")
86 .field("dim", &self.dim)
87 .field("metric", &self.metric)
88 .finish_non_exhaustive()
89 }
90}
91
92impl Vectors {
93 #[must_use]
95 pub fn dim(&self) -> usize {
96 self.dim
97 }
98
99 #[must_use]
101 pub fn metric(&self) -> Metric {
102 self.metric
103 }
104
105 pub fn put(&self, key: impl AsRef<[u8]>, v: &[f32]) -> Result<bool> {
119 let key = key.as_ref();
120 self.db
121 .write(|inner| inner.collections[self.at].data.vectors_mut().put(key, v))
122 }
123
124 pub fn get(&self, key: impl AsRef<[u8]>) -> Result<Option<Vec<f32>>> {
134 self.with(key, <[f32]>::to_vec)
135 }
136
137 pub fn with<R>(&self, key: impl AsRef<[u8]>, f: impl FnOnce(&[f32]) -> R) -> Result<Option<R>> {
147 let key = key.as_ref();
148 self.db
149 .read(|inner| Ok(inner.collections[self.at].data.vectors().get(key).map(f)))
150 }
151
152 pub fn contains(&self, key: impl AsRef<[u8]>) -> Result<bool> {
158 let key = key.as_ref();
159 self.db
160 .read(|inner| Ok(inner.collections[self.at].data.vectors().contains(key)))
161 }
162
163 pub fn remove(&self, key: impl AsRef<[u8]>) -> Result<bool> {
174 let key = key.as_ref();
175 self.db
176 .write(|inner| Ok(inner.collections[self.at].data.vectors_mut().remove(key)))
177 }
178
179 pub fn len(&self) -> Result<usize> {
185 self.db
186 .read(|inner| Ok(inner.collections[self.at].data.vectors().len()))
187 }
188
189 pub fn is_empty(&self) -> Result<bool> {
195 self.len().map(|n| n == 0)
196 }
197
198 pub fn search(&self, q: &[f32], k: usize) -> Result<Vec<Match>> {
222 self.db
223 .read(|inner| inner.collections[self.at].data.vectors().search(q, k, None))
224 }
225
226 pub fn near(&self, key: impl AsRef<[u8]>, k: usize) -> Result<Vec<Match>> {
237 let key = key.as_ref();
238 self.db.read(|inner| {
239 let store = inner.collections[self.at].data.vectors();
240 let Some(q) = store.get(key) else {
241 return Err(Error::fmt(
242 Code::NotFound,
243 format_args!(
244 "this collection has no vector under that key, so there is nothing to be near to"
245 ),
246 ));
247 };
248 store.search(q, k, Some(key))
249 })
250 }
251}
252
253#[cfg(test)]
254mod tests {
255 use yo_vector::collection::MAX_DIM;
256
257 use crate::db::{MEMORY, open};
258
259 use super::*;
260
261 fn axes(v: &Vectors) {
263 v.put("x", &[1.0, 0.0, 0.0]).unwrap();
264 v.put("y", &[0.0, 1.0, 0.0]).unwrap();
265 v.put("z", &[0.0, 0.0, 1.0]).unwrap();
266 }
267
268 #[test]
269 fn a_vector_comes_back_the_way_it_went_in() {
270 let db = open(MEMORY).unwrap();
271 let v = db.vectors("e", 3).unwrap();
272 assert!(v.is_empty().unwrap());
273
274 assert!(v.put("x", &[1.0, 2.0, 3.0]).unwrap(), "the key is new");
275 assert!(!v.put("x", &[1.0, 2.0, 3.0]).unwrap(), "and then it is not");
276 assert_eq!(v.get("x").unwrap(), Some(vec![1.0, 2.0, 3.0]));
277 assert_eq!(v.len().unwrap(), 1);
278 assert!(v.contains("x").unwrap());
279 assert_eq!(v.get("nobody").unwrap(), None);
280 assert_eq!(v.with("x", |x| x.len()).unwrap(), Some(3));
281 }
282
283 #[test]
284 fn the_nearest_answer_is_the_nearest_vector() {
285 let db = open(MEMORY).unwrap();
286 let v = db.vectors("e", 3).unwrap();
287 axes(&v);
288
289 let hits = v.search(&[0.9, 0.2, 0.1], 3).unwrap();
290 let keys: Vec<&[u8]> = hits.iter().map(|h| h.key.as_slice()).collect();
291 assert_eq!(keys, vec![&b"x"[..], &b"y"[..], &b"z"[..]]);
292 assert!(hits[0].distance < hits[1].distance);
293 let want = (0.01f32 + 0.04 + 0.01).sqrt();
296 assert!((hits[0].distance - want).abs() < 1e-6, "{hits:?}");
297 }
298
299 #[test]
300 fn asking_for_more_than_there_is_gets_what_there_is() {
301 let db = open(MEMORY).unwrap();
302 let v = db.vectors("e", 3).unwrap();
303 assert!(v.search(&[1.0, 0.0, 0.0], 4).unwrap().is_empty());
304 axes(&v);
305 assert_eq!(v.search(&[1.0, 0.0, 0.0], 10).unwrap().len(), 3);
306 assert!(v.search(&[1.0, 0.0, 0.0], 0).unwrap().is_empty());
307 }
308
309 #[test]
310 fn a_removed_vector_is_not_an_answer_and_its_slot_comes_back() {
311 let db = open(MEMORY).unwrap();
312 let v = db.vectors("e", 3).unwrap();
313 axes(&v);
314
315 assert!(v.remove("x").unwrap());
316 assert!(!v.remove("x").unwrap(), "twice is not there twice");
317 assert_eq!(v.len().unwrap(), 2);
318 assert!(!v.contains("x").unwrap());
319
320 let hits = v.search(&[1.0, 0.0, 0.0], 3).unwrap();
321 assert_eq!(hits.len(), 2);
322 assert!(hits.iter().all(|h| h.key != b"x".to_vec()));
323
324 v.put("w", &[1.0, 0.0, 0.0]).unwrap();
325 let hits = v.search(&[1.0, 0.0, 0.0], 1).unwrap();
326 assert_eq!(hits[0].key, b"w".to_vec(), "the reused slot answers as w");
327 }
328
329 #[test]
333 fn a_replaced_vector_is_searched_at_its_new_place() {
334 let db = open(MEMORY).unwrap();
335 let v = db.vectors("e", 3).unwrap();
336 axes(&v);
337
338 v.put("x", &[0.0, 0.0, 1.0]).unwrap();
339 assert_eq!(v.len().unwrap(), 3, "a replacement is not a second key");
340
341 let hits = v.search(&[1.0, 0.0, 0.0], 1).unwrap();
342 assert_eq!(hits[0].key, b"y".to_vec(), "x moved away from that corner");
343 let hits = v.search(&[0.0, 0.0, 1.0], 2).unwrap();
344 let keys: Vec<Vec<u8>> = hits.into_iter().map(|h| h.key).collect();
345 assert!(keys.contains(&b"x".to_vec()) && keys.contains(&b"z".to_vec()));
346 }
347
348 #[test]
349 fn more_like_this_leaves_the_thing_itself_out() {
350 let db = open(MEMORY).unwrap();
351 let v = db.vectors("e", 3).unwrap();
352 axes(&v);
353 v.put("x2", &[0.9, 0.1, 0.0]).unwrap();
354
355 let hits = v.near("x", 2).unwrap();
356 assert_eq!(hits.len(), 2);
357 assert_eq!(hits[0].key, b"x2".to_vec());
358 assert!(hits.iter().all(|h| h.key != b"x".to_vec()));
359
360 let e = v.near("nobody", 2).expect_err("no such key");
361 assert_eq!(e.code(), Code::NotFound);
362 }
363
364 #[test]
365 fn a_cosine_collection_stores_the_direction_and_reports_the_angle() {
366 let db = open(MEMORY).unwrap();
367 let v = db.vectors_with("e", 2, Metric::Cosine).unwrap();
368
369 v.put("east", &[7.0, 0.0]).unwrap();
370 v.put("north", &[0.0, 3.0]).unwrap();
371 v.put("west", &[-2.0, 0.0]).unwrap();
372 assert_eq!(v.get("east").unwrap(), Some(vec![1.0, 0.0]));
373
374 let hits = v.search(&[100.0, 0.0], 3).unwrap();
377 assert_eq!(hits[0].key, b"east".to_vec());
378 assert!(hits[0].distance.abs() < 1e-6, "{hits:?}");
379 assert!(
380 (hits[1].distance - 1.0).abs() < 1e-6,
381 "north is a right angle"
382 );
383 assert!(
384 (hits[2].distance - 2.0).abs() < 1e-6,
385 "west is the opposite"
386 );
387
388 let e = v.put("nowhere", &[0.0, 0.0]).expect_err("no direction");
389 assert_eq!(e.code(), Code::Invalid);
390 }
391
392 #[test]
393 fn a_vector_of_the_wrong_length_or_shape_is_refused() {
394 let db = open(MEMORY).unwrap();
395 let v = db.vectors("e", 3).unwrap();
396
397 let e = v.put("x", &[1.0, 2.0]).expect_err("two is not three");
398 assert_eq!(e.code(), Code::Invalid);
399 assert!(e.message().contains("3 dimensional"), "{e}");
400
401 let e = v.put("x", &[1.0, f32::NAN, 2.0]).expect_err("not a number");
402 assert_eq!(e.code(), Code::Invalid);
403 assert!(e.message().contains("coordinate 1"), "{e}");
404
405 let e = v.search(&[1.0], 1).expect_err("one is not three");
406 assert_eq!(e.code(), Code::Invalid);
407 }
408
409 #[test]
412 fn recall_holds_once_the_index_has_split() {
413 let db = open(MEMORY).unwrap();
414 let v = db.vectors("e", 8).unwrap();
415
416 let mut seed = 0x2026u64;
417 let mut next = move || {
418 seed = seed.wrapping_mul(6364136223846793005).wrapping_add(1);
419 ((seed >> 33) as f32 / (1u64 << 31) as f32) - 0.5
420 };
421 let mut all: Vec<Vec<f32>> = Vec::new();
422 for i in 0..2000usize {
423 let x: Vec<f32> = (0..8).map(|_| next()).collect();
424 v.put(format!("k{i}"), &x).unwrap();
425 all.push(x);
426 }
427
428 let mut found = 0;
429 for (i, q) in all.iter().enumerate().step_by(50) {
430 let hits = v.search(q, 1).unwrap();
431 if hits[0].key == format!("k{i}").into_bytes() {
432 found += 1;
433 }
434 }
435 assert!(found >= 39, "{found} of 40 queries found their own vector");
436 assert_eq!(v.len().unwrap(), 2000);
437 assert!(db.memory_bytes().unwrap() > 2000 * 8 * 4);
438 }
439
440 #[test]
441 fn a_dimension_or_a_metric_the_build_cannot_hold_is_refused_at_open() {
442 let db = open(MEMORY).unwrap();
443
444 let e = db.vectors("e", 0).expect_err("zero dimensions is nothing");
445 assert_eq!(e.code(), Code::Invalid);
446 let e = db.vectors("e", MAX_DIM + 1).expect_err("past the limit");
447 assert_eq!(e.code(), Code::Invalid);
448
449 let e = db
450 .vectors_with("e", 8, Metric::Ip)
451 .expect_err("not a distance");
452 assert_eq!(e.code(), Code::Unsupported);
453 assert!(e.message().contains("cosine"), "{e}");
454 let e = db
455 .vectors_with("e", 8, Metric::Hamming)
456 .expect_err("not floats");
457 assert_eq!(e.code(), Code::Unsupported);
458 }
459
460 #[test]
464 fn the_dimension_and_the_metric_are_part_of_the_shape() {
465 let db = open(MEMORY).unwrap();
466 let v = db.vectors("e", 3).unwrap();
467 v.put("x", &[1.0, 0.0, 0.0]).unwrap();
468
469 let same = db.vectors("e", 3).unwrap();
470 assert_eq!(same.len().unwrap(), 1, "the same name is the same store");
471
472 let e = db.vectors("e", 4).expect_err("that is another collection");
473 assert_eq!(e.code(), Code::ShapeMismatch);
474 let e = db
475 .vectors_with("e", 3, Metric::Cosine)
476 .expect_err("and so is that");
477 assert_eq!(e.code(), Code::ShapeMismatch);
478
479 let e = db.map::<String, u64>("e").expect_err("and so is a map");
480 assert_eq!(e.code(), Code::ShapeMismatch);
481 }
482}