diff --git a/Cargo.lock b/Cargo.lock index 79ff5dd55b..fd12ad2e7e 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -605,7 +605,7 @@ dependencies = [ [[package]] name = "diskann-garnet" -version = "5.0.3" +version = "6.0.0" dependencies = [ "bytemuck", "crossbeam", @@ -617,6 +617,7 @@ dependencies = [ "diskann-vector", "foldhash", "rand", + "strum", "thiserror 2.0.17", "tokio", ] @@ -2311,6 +2312,27 @@ version = "0.11.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7da8b5736845d9f2fcb837ea5d9e2628564b3b043a70948a3f0b778838c5fb4f" +[[package]] +name = "strum" +version = "0.28.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9628de9b8791db39ceda2b119bbe13134770b56c138ec1d3af810d045c04f9bd" +dependencies = [ + "strum_macros", +] + +[[package]] +name = "strum_macros" +version = "0.28.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ab85eea0270ee17587ed4156089e10b9e6880ee688791d45a905f5b1ca36f664" +dependencies = [ + "heck", + "proc-macro2", + "quote", + "syn 2.0.117", +] + [[package]] name = "syn" version = "1.0.109" diff --git a/diskann-garnet/Cargo.toml b/diskann-garnet/Cargo.toml index be64165c85..ef05cbdd65 100644 --- a/diskann-garnet/Cargo.toml +++ b/diskann-garnet/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "diskann-garnet" -version = "5.0.3" +version = "6.0.0" edition = "2024" authors.workspace = true license.workspace = true @@ -23,6 +23,7 @@ diskann-utils.workspace = true diskann-vector.workspace = true foldhash = "0.2.0" rand.workspace = true +strum = { version = "0.28.0", features = ["derive"] } thiserror.workspace = true tokio = { workspace = true, features = ["sync"] } diff --git a/diskann-garnet/README.md b/diskann-garnet/README.md index 61e80992c7..7081b9b2f7 100644 --- a/diskann-garnet/README.md +++ b/diskann-garnet/README.md @@ -123,8 +123,8 @@ the one you want. ## Testing -Unit tests are run in the usual way with `cargo test`, but many are end-to-end -and run from the Garnet side. These two invocations will run the relevant tests: +Run the Rust tests from the workspace root with `cargo test -p diskann-garnet`. +Many end-to-end tests run from the Garnet side. These two invocations run them: ``` dotnet test test/standalone/Garnet.test.vectorset -f net10.0 -c Debug --filter RespVectorSetTests diff --git a/diskann-garnet/diskann-garnet.nuspec b/diskann-garnet/diskann-garnet.nuspec index 89e5748f90..8c53121e95 100644 --- a/diskann-garnet/diskann-garnet.nuspec +++ b/diskann-garnet/diskann-garnet.nuspec @@ -2,7 +2,7 @@ diskann-garnet - 5.0.3 + 6.0.0 docs/README.md Microsoft https://github.com/microsoft/DiskANN diff --git a/diskann-garnet/docs/data-design.md b/diskann-garnet/docs/data-design.md index c794c33370..c3a0d2c0a6 100644 --- a/diskann-garnet/docs/data-design.md +++ b/diskann-garnet/docs/data-design.md @@ -36,7 +36,7 @@ Diskann-garnet uses these bits to distinguish between differnet kinds of index d ### Key Data Prefixing -In order to reduce allocations in the data access path in Garnet, Garnet needs some place to scribble state into during operations. It uses a single byte immediately preceding the first key byte for this purpose. This means that any key pointer given to Garnet access methods must contain valid space preceding the real key. For this reason, key data pointers are `*mut u8` and not `*const u8` and care must be taken to ensure the memory preceding that pointer is valid. In diskann-garnet, we precede the key data with at least 4 bytes of scratch space. +Read callbacks receive keys with four-byte length prefixes. Write, delete, and RMW callbacks receive the key bytes and length separately. ## Term Types diff --git a/diskann-garnet/docs/ffi-design.rs b/diskann-garnet/docs/ffi-design.rs index d90856d03b..6ff12a7c4e 100644 --- a/diskann-garnet/docs/ffi-design.rs +++ b/diskann-garnet/docs/ffi-design.rs @@ -36,11 +36,13 @@ enum VectorQuantType { /// Status returned by `insert`, encoded as a `u8`. /// /// `SuccessStartTraining` signals that the insert crossed the threshold at which the quantizer -/// can be trained, and that Garnet should call `build_quant_table`. +/// can be trained, and that Garnet should call `build_quant_table`. Otherwise, `Success` is returned for a newly +/// inserted vector, and `SuccessUpdate` for an updated one. enum InsertResult { Fail = 0, Success = 1, SuccessStartTraining = 2, + SuccessUpdate = 3, } /// Read one or more keys from Garnet. @@ -155,8 +157,8 @@ extern "C" fn drop_index(context: u64, index_ptr: *const c_void); /// Insert a vector into an index. /// -/// Returns an `InsertResult` discriminant. `Fail` may result from the vector already being in -/// the index, or from writes failing. +/// Returns an `InsertResult` discriminant, distinguishing inserts from updates unless quantizer +/// training becomes ready. `Fail` may result from invalid input or writes failing. /// /// vector_len is a count of elements, not bytes; the element type follows from the index's /// `quant_type`. The pointer need not be aligned for that element type. diff --git a/diskann-garnet/src/dyn_index.rs b/diskann-garnet/src/dyn_index.rs index 0b60bad4f8..dd60d58956 100644 --- a/diskann-garnet/src/dyn_index.rs +++ b/diskann-garnet/src/dyn_index.rs @@ -116,26 +116,37 @@ impl DynIndex for DiskANNIndex> { /// /// The data slice here must be aligned to `T` or this will panic. fn insert(&self, context: &Context, id: &GarnetId, data: &[u8], attrs: &[u8]) -> ANNResult<()> { - self.insert( - &DynamicQuantization, - context, - id, - (bytemuck::cast_slice::(data), attrs), - ) + self.run(|_| async { + let _pending = self.inner.provider().reserve_external_id(id).await?; + self.inner + .insert( + &DynamicQuantization, + context, + id, + (bytemuck::cast_slice::(data), attrs), + ) + .await + }) } fn set_attributes(&self, context: &Context, id: &GarnetId, data: &[u8]) -> ANNResult<()> { - self.inner - .provider() - .set_attributes(context, id, data) - .map_err(|e| e.into()) + self.run(|_| async { + let provider = self.inner.provider(); + let _pending = provider.reserve_external_id(id).await?; + provider + .set_attributes(context, id, data) + .map_err(|error| error.into()) + }) } fn delete_attributes(&self, context: &Context, id: &GarnetId) -> ANNResult<()> { - self.inner - .provider() - .delete_attributes(context, id) - .map_err(|e| e.into()) + self.run(|_| async { + let provider = self.inner.provider(); + let _pending = provider.reserve_external_id(id).await?; + provider + .delete_attributes(context, id) + .map_err(|error| error.into()) + }) } fn search_vector( @@ -189,13 +200,18 @@ impl DynIndex for DiskANNIndex> { } fn remove(&self, context: &Context, id: &GarnetId) -> ANNResult<()> { - self.inplace_delete( - DynamicQuantization, - context, - id, - 3, - InplaceDeleteMethod::TwoHopAndOneHop, - ) + self.run(|_| async { + let _pending = self.inner.provider().reserve_external_id(id).await?; + self.inner + .inplace_delete( + DynamicQuantization, + context, + id, + 3, + InplaceDeleteMethod::TwoHopAndOneHop, + ) + .await + }) } fn approximate_count(&self) -> u64 { diff --git a/diskann-garnet/src/ffi_recall_tests.rs b/diskann-garnet/src/ffi_recall_tests.rs index 8d9b6a67c6..f2d006d96d 100644 --- a/diskann-garnet/src/ffi_recall_tests.rs +++ b/diskann-garnet/src/ffi_recall_tests.rs @@ -125,9 +125,9 @@ mod tests { let mut result_ids = Vec::with_capacity(count); let mut offset = 0; for _ in 0..count { - let mut id_len = 0u32; - bytemuck::bytes_of_mut(&mut id_len) - .copy_from_slice(&output_id_buffer[offset..offset + mem::size_of::()]); + let id_len = bytemuck::pod_read_unaligned::( + &output_id_buffer[offset..offset + mem::size_of::()], + ); offset += mem::size_of::(); let id_str = std::str::from_utf8(&output_id_buffer[offset..offset + id_len as usize]) .expect("id should be valid utf8"); diff --git a/diskann-garnet/src/ffi_tests.rs b/diskann-garnet/src/ffi_tests.rs index c3f054ab7b..e039b8303c 100644 --- a/diskann-garnet/src/ffi_tests.rs +++ b/diskann-garnet/src/ffi_tests.rs @@ -123,6 +123,64 @@ mod tests { } } + #[test] + fn insert_distinguishes_updates() { + for quant_type in [ + VectorQuantType::NoQuant, + VectorQuantType::Bin, + VectorQuantType::Q8, + ] { + let store = Store::new(); + let (index_ptr, ctx) = create_test_index(&store, quant_type); + + for (id, vector, expected, expected_internal_id) in [ + (42, [1.0, 2.0], 1, 1u32), + (42, [2.0, 1.0], 3, 1), + (43, [3.0, 4.0], 1, 2), + (43, [4.0, 3.0], 3, 2), + ] { + assert_eq!( + u8::from(insert_f32_vector(&ctx, index_ptr, id, &vector)), + expected, + "unexpected insert status for {quant_type:?} and ID {id}" + ); + let internal_id = store + .get(ctx.term(Term::IntMap).get(), bytemuck::bytes_of(&id)) + .unwrap(); + assert_eq!( + bytemuck::pod_read_unaligned::(&internal_id), + expected_internal_id + ); + assert_eq!( + store.get(ctx.term(Term::Vector).get(), &internal_id), + Some(bytemuck::cast_slice::(&vector).to_vec()) + ); + assert_eq!( + unsafe { card(ctx.get(), index_ptr) }, + u64::from(expected_internal_id) + ); + } + assert_eq!(store.int_map_reads(), 4); + + let (ids, distances) = do_search(&ctx, index_ptr, &[2.0, 1.0], 3, None); + assert_eq!(ids, [42, 43]); + assert_eq!(distances[0], 0.0); + + let id_bytes = bytemuck::bytes_of(&42u32); + assert!(unsafe { remove(ctx.get(), index_ptr, id_bytes.as_ptr(), id_bytes.len()) }); + store.clear_read_counts(); + assert_eq!( + insert_f32_vector(&ctx, index_ptr, 42, &[5.0, 6.0]), + InsertResult::Success + ); + assert_eq!(store.int_map_reads(), 1); + + unsafe { + drop_index(ctx.get(), index_ptr); + } + } + } + #[test] fn add_check_and_remove_vector() { let store = Store::new(); @@ -493,16 +551,16 @@ mod tests { let mut output_ids = vec![]; let mut offset = 0; for _ in 0..(count as usize) { - let mut id_len = 0u32; - bytemuck::bytes_of_mut(&mut id_len) - .copy_from_slice(&output_id_buffer[offset..offset + mem::size_of::()]); + let id_len = bytemuck::pod_read_unaligned::( + &output_id_buffer[offset..offset + mem::size_of::()], + ); offset += mem::size_of::(); assert_eq!(id_len, mem::size_of::() as u32); - let mut id = 0u64; - bytemuck::bytes_of_mut(&mut id) - .copy_from_slice(&output_id_buffer[offset..offset + mem::size_of::()]); + let id = bytemuck::pod_read_unaligned::( + &output_id_buffer[offset..offset + mem::size_of::()], + ); offset += mem::size_of::(); output_ids.push(id); @@ -603,16 +661,16 @@ mod tests { let mut output_ids = vec![]; let mut offset = 0; for _ in 0..(count as usize) { - let mut id_len = 0u32; - bytemuck::bytes_of_mut(&mut id_len) - .copy_from_slice(&output_id_buffer[offset..offset + mem::size_of::()]); + let id_len = bytemuck::pod_read_unaligned::( + &output_id_buffer[offset..offset + mem::size_of::()], + ); offset += mem::size_of::(); assert_eq!(id_len, mem::size_of::() as u32); - let mut id = 0u64; - bytemuck::bytes_of_mut(&mut id) - .copy_from_slice(&output_id_buffer[offset..offset + mem::size_of::()]); + let id = bytemuck::pod_read_unaligned::( + &output_id_buffer[offset..offset + mem::size_of::()], + ); offset += mem::size_of::(); output_ids.push(id); @@ -772,13 +830,13 @@ mod tests { let mut ids = vec![]; let mut offset = 0; for _ in 0..count { - let mut id_len = 0u32; - bytemuck::bytes_of_mut(&mut id_len) - .copy_from_slice(&output_id_buffer[offset..offset + mem::size_of::()]); + let id_len = bytemuck::pod_read_unaligned::( + &output_id_buffer[offset..offset + mem::size_of::()], + ); offset += mem::size_of::(); - let mut id = 0u32; - bytemuck::bytes_of_mut(&mut id) - .copy_from_slice(&output_id_buffer[offset..offset + id_len as usize]); + let id = bytemuck::pod_read_unaligned::( + &output_id_buffer[offset..offset + id_len as usize], + ); offset += id_len as usize; ids.push(id); } diff --git a/diskann-garnet/src/fsm.rs b/diskann-garnet/src/fsm.rs index 88bdf0a625..c39415826d 100644 --- a/diskann-garnet/src/fsm.rs +++ b/diskann-garnet/src/fsm.rs @@ -41,7 +41,7 @@ pub(crate) enum FsmError { IdOutOfRange(u32), } -/// Guard returned by `next_id()` to ensure correctness when reuse is enabled. +/// Guard protecting vector writes against quantization phase changes. pub(crate) struct ReuseGuard<'a> { id: u32, barrier: RwLockReadGuard<'a, Barrier>, @@ -59,6 +59,10 @@ impl<'a> ReuseGuard<'a> { pub(crate) fn should_quantize(&self) -> bool { self.barrier.quantization_enabled } + + pub(crate) fn max_id_for_backfill(&self) -> u32 { + self.barrier.max_id_for_backfill + } } struct Barrier { @@ -366,6 +370,11 @@ impl FreeSpaceMap { Ok(ReuseGuard::new(id, barrier)) } + /// Guard writes to an existing ID without changing its allocation state. + pub(crate) fn existing_id(&self, id: u32) -> ReuseGuard<'_> { + ReuseGuard::new(id, self.barrier.read().unwrap()) + } + /// Return the maximum ID that has been assigned to a vector. /// /// This ID may be free if that ID has been deleted since the ID was created. diff --git a/diskann-garnet/src/garnet.rs b/diskann-garnet/src/garnet.rs index d1547b1f4b..ed3b068f78 100644 --- a/diskann-garnet/src/garnet.rs +++ b/diskann-garnet/src/garnet.rs @@ -21,7 +21,11 @@ use thiserror::Error; /// Must have enough bits to represent all Term variants (max value is 6, needs 3 bits). pub(crate) const TERM_BITMASK: u64 = (1 << 3) - 1; -#[derive(Debug)] +#[derive(Debug, Error)] +#[error("Invalid term {0}")] +pub(crate) struct InvalidTerm(u32); + +#[derive(Copy, Clone, Debug, strum::VariantArray)] pub(crate) enum Term { Vector = 0, Neighbors = 1, @@ -32,17 +36,40 @@ pub(crate) enum Term { ExtMap = 6, } +impl TryFrom for Term { + type Error = InvalidTerm; + + fn try_from(value: u32) -> Result { + match value { + 0 => Ok(Term::Vector), + 1 => Ok(Term::Neighbors), + 2 => Ok(Term::Quantized), + 3 => Ok(Term::Attributes), + 4 => Ok(Term::Metadata), + 5 => Ok(Term::IntMap), + 6 => Ok(Term::ExtMap), + _ => Err(InvalidTerm(value)), + } + } +} + +#[derive(Debug, Default)] +struct ContextState { + quantizer_ready: AtomicBool, + insert_is_update: AtomicBool, +} + #[derive(Clone, Debug)] pub(crate) struct Context { inner: u64, - quantizer_ready: Arc, + state: Arc, } impl Context { pub(crate) fn new(inner: u64) -> Self { Self { inner, - quantizer_ready: Arc::new(AtomicBool::new(false)), + state: Arc::new(ContextState::default()), } } @@ -52,25 +79,27 @@ impl Context { } pub(crate) fn term(&self, kind: Term) -> Self { - let Context { - inner, - quantizer_ready, - } = self; + let Context { inner, state } = self; let inner = *inner | (kind as u64 & TERM_BITMASK); - let quantizer_ready = quantizer_ready.clone(); + let state = state.clone(); - Self { - inner, - quantizer_ready, - } + Self { inner, state } } pub(crate) fn quantizer_ready(&self) -> bool { - self.quantizer_ready.load(Ordering::Acquire) + self.state.quantizer_ready.load(Ordering::Acquire) } pub(crate) fn set_quantizer_ready(&self) { - self.quantizer_ready.store(true, Ordering::Release); + self.state.quantizer_ready.store(true, Ordering::Release); + } + + pub(crate) fn insert_is_update(&self) -> bool { + self.state.insert_is_update.load(Ordering::Acquire) + } + + pub(crate) fn set_insert_is_update(&self) { + self.state.insert_is_update.store(true, Ordering::Release); } } @@ -147,7 +176,6 @@ impl Callbacks { self.log_callback } - #[cfg(test)] pub(crate) fn exists_iid(&self, ctx: &Context, id: u32, length_hint: usize) -> bool { let key = [4, id]; // SAFETY: Key bytes are preceded by 4 bytes of space. @@ -162,19 +190,14 @@ impl Callbacks { unsafe { self.exists_raw(ctx, &key_bytes[4..], length_hint) } } - #[expect( - dead_code, - reason = "currently unused, but may be needed in the future" - )] pub(crate) fn exists_eid(&self, ctx: &Context, id: &GarnetId, length_hint: usize) -> bool { // SAFETY: GarnetId ensures there are 4 bytes preceding the key bytes. - unsafe { self.exists_raw(ctx, id, length_hint) } + unsafe { self.exists_raw(ctx, id.as_prefixed_key_bytes(), length_hint) } } /// Check for a key's existance in Garnet. /// - /// NOTE: The key bytes must be preceded by 4 valid bytes that Garnet can write into. - /// This invariant must be checked by the caller. + /// The key must be prefixed by a four byte length. unsafe fn exists_raw(&self, ctx: &Context, key: &[u8], length_hint: usize) -> bool { let mut called = false; let mut cb = |_, _: &[u8]| { @@ -253,8 +276,7 @@ impl Callbacks { /// Read a single key from Garnet. /// - /// NOTE: The key bytes must be preceded by 4 valid bytes that Garnet can write into. - /// This invariant must be checked by the caller. + /// The key must be prefixed by a four byte length. #[must_use] unsafe fn read_single_raw(&self, ctx: &Context, key: &[u8], value: &mut [u8]) -> bool { let length_hint = value.len() as u32; @@ -388,8 +410,7 @@ impl Callbacks { /// Write a value for a key in Garnet. /// - /// NOTE: The key bytes must be preceded by 4 valid bytes that Garnet can write into. - /// This invariant must be checked by the caller. + /// The key is passed without a length prefix. #[must_use] unsafe fn write_raw(&self, ctx: &Context, key: &[u8], value: &[u8]) -> bool { let value_ptr = value.as_ptr(); @@ -484,8 +505,7 @@ impl Callbacks { /// The provided function `f` will receive the current value, which it can then modify. If no /// value exists, zero-initialized value of length `write_len` will be passed in. /// - /// The key bytes must be preceded by 4 valid bytes that Garnet can write into. - /// This invariant must be checked by the caller. + /// The key is passed without a length prefix. /// /// `f` should not panic. #[must_use] @@ -595,15 +615,13 @@ pub(crate) enum GarnetError { /// A variable length byte string used as the vector ID in a Garnet vector set. /// -/// A wrapped type is used because the Garnet callbacks expect some padding bytes it can -/// use to avoid allocation, and this type ensures those bytes exist without interfering -/// with the "real" ID bytes. +/// This is cheap to clone as it uses `Arc` internally, and prefixes the data with a 4-byte length +/// appropriate for use with the read callbacks. /// -/// Dereferencing this type will return a slice to the actual ID bytes, without the padding, -/// which makes this interchangeable in most respects with using a raw `Box<[u8]>`. +/// Dereferencing returns only the ID bytes. #[derive(Clone, PartialEq)] pub(crate) struct GarnetId { - inner: Box<[u8]>, + inner: Arc<[u8]>, } impl GarnetId { @@ -620,11 +638,13 @@ impl fmt::Debug for GarnetId { impl From<&[u8]> for GarnetId { fn from(value: &[u8]) -> Self { - let mut id = Vec::with_capacity(value.len() + 4); + let mut inner = Arc::<[u8]>::new_uninit_slice(value.len() + 4); + let buffer = Arc::get_mut(&mut inner).unwrap(); let len = value.len() as u32; - id.extend_from_slice(bytemuck::bytes_of(&len)); - id.extend_from_slice(value); - let inner = id.into(); + buffer[..4].write_copy_of_slice(bytemuck::bytes_of(&len)); + buffer[4..].write_copy_of_slice(value); + // SAFETY: The prefix and ID copies initialize every byte of the allocation. + let inner = unsafe { inner.assume_init() }; Self { inner } } diff --git a/diskann-garnet/src/lib.rs b/diskann-garnet/src/lib.rs index 61484893b8..450e1998b3 100644 --- a/diskann-garnet/src/lib.rs +++ b/diskann-garnet/src/lib.rs @@ -491,14 +491,15 @@ fn interpret_vector<'a>( /// Return type for `insert()`. /// -/// `Fail` and `Success` are obvious. `SuccessStartTraining` is used when enough vectors have -/// been inserted to start training the quantizer. That return value signals to Garnet that -/// `build_quant_table` should be called. +/// `Success` indicates a new vector and `SuccessUpdate` indicates an existing vector. +/// `SuccessStartTraining` is returned when enough vectors have been inserted to start +/// training the quantizer, signaling to Garnet that `build_quant_table` should be called. #[derive(Debug, Clone, Copy, PartialEq)] enum InsertResult { Fail, Success, SuccessStartTraining, + SuccessUpdate, } impl From for u8 { @@ -507,6 +508,7 @@ impl From for u8 { InsertResult::Fail => 0, InsertResult::Success => 1, InsertResult::SuccessStartTraining => 2, + InsertResult::SuccessUpdate => 3, } } } @@ -517,6 +519,7 @@ impl From for InsertResult { match value { 1 => InsertResult::Success, 2 => InsertResult::SuccessStartTraining, + 3 => InsertResult::SuccessUpdate, _ => InsertResult::Fail, } } @@ -524,9 +527,8 @@ impl From for InsertResult { /// Insert a vector into the index. /// -/// Returns a status corresponding to the `InsertResult` enum. Aside from failure and success, -/// there is a third value that signals that the completed insert has reached the threshold to -/// begin quantization. +/// Returns a status corresponding to the `InsertResult` enum, distinguishing new inserts from +/// updates and when an insert allows quantization training to begin. /// /// # Safety /// @@ -573,6 +575,8 @@ pub unsafe extern "C" fn insert( let ready = ctx.quantizer_ready(); if !old_ready && ready { InsertResult::SuccessStartTraining.into() + } else if ctx.insert_is_update() { + InsertResult::SuccessUpdate.into() } else { InsertResult::Success.into() } @@ -770,9 +774,9 @@ impl Overflow { if id_index + prefix_len > ids.len() { break; } - let mut len = 0u32; - bytemuck::bytes_of_mut(&mut len) - .copy_from_slice(&self.id_buffer[self.id_index..self.id_index + prefix_len]); + let len = bytemuck::pod_read_unaligned::( + &self.id_buffer[self.id_index..self.id_index + prefix_len], + ); // We check there is room before advancing the indices if id_index + prefix_len + len as usize > ids.len() { @@ -1110,8 +1114,7 @@ pub unsafe extern "C" fn check_internal_id_valid( return false; } - let mut id: u32 = 0; - bytemuck::bytes_of_mut(&mut id).copy_from_slice(internal_id_bytes); + let id = bytemuck::pod_read_unaligned::(internal_id_bytes); index.inner.internal_id_exists(&ctx, id) } @@ -1282,14 +1285,12 @@ mod tests { let mut pos = 0usize; for (i, d) in dists.iter().enumerate() { - let mut size = 0u32; - bytemuck::bytes_of_mut(&mut size).copy_from_slice(&ids[pos..pos + 4]); + let size = bytemuck::pod_read_unaligned::(&ids[pos..pos + 4]); pos += 4; assert_eq!(size, 4); - let mut id = 0u32; - bytemuck::bytes_of_mut(&mut id).copy_from_slice(&ids[pos..pos + 4]); + let id = bytemuck::pod_read_unaligned::(&ids[pos..pos + 4]); pos += 4; assert_eq!(id, i as u32 + 1); diff --git a/diskann-garnet/src/provider.rs b/diskann-garnet/src/provider.rs index 556b6505e1..0e2ca5557f 100644 --- a/diskann-garnet/src/provider.rs +++ b/diskann-garnet/src/provider.rs @@ -18,8 +18,8 @@ use diskann::{ }, neighbor::Neighbor, provider::{ - DataProvider, Delete, ElementStatus, HasId, NeighborAccessor, NeighborAccessorMut, - NoopGuard, SetElement, + DataProvider, Delete, ElementStatus, Guard, HasId, NeighborAccessor, NeighborAccessorMut, + SetElement, }, utils::VectorRepr, }; @@ -35,15 +35,18 @@ use std::{ any::TypeId, collections::HashSet, future, + hash::BuildHasher, marker::PhantomData, mem, - ops::{Deref, DerefMut}, + ops::{Deref, DerefMut, Range}, sync::{ - Mutex, + Arc, Condvar, Mutex, atomic::{AtomicBool, AtomicU64, Ordering}, }, }; +use strum::VariantArray; use thiserror::Error; +use tokio::sync::watch; use crate::{ SearchResults, VectorQuantType, @@ -69,6 +72,9 @@ const RERANK_BUFFER_LENGTH: usize = 1024; /// length, so this is only an estimate used to size Garnet's read buffer. const ATTRIBUTE_LENGTH_HINT: usize = 1024; +/// Maximum number of reservation retries after waiting for another owner. +const RESERVATION_RETRY_LIMIT: usize = 5; + #[derive(Clone)] struct AdjList(AdjacencyList); @@ -112,11 +118,123 @@ pub(crate) enum GarnetProviderError { Quantizer(#[from] GarnetQuantizerError), #[error("Post processing error: {0}")] PostProcessing(Box), + #[error("External ID reservation retry limit reached")] + ReservationRetryLimit, } diskann::convert_error!(GarnetProviderError); diskann::always_escalate!(GarnetProviderError); +struct BackfillGuard { + ranges: Arc>>>, + notify: Arc, + range: Range, +} + +impl Drop for BackfillGuard { + fn drop(&mut self) { + let _ = self.ranges.lock().unwrap().remove(&self.range); + self.notify.notify_all(); + } +} + +pub(crate) struct ExternalIdGuard<'a> { + pending: &'a DashMap, foldhash::fast::RandomState>, + id_hash: u64, +} + +impl Drop for ExternalIdGuard<'_> { + fn drop(&mut self) { + self.pending.remove(&self.id_hash); + } +} + +pub(crate) struct InsertGuard { + callbacks: Callbacks, + context: Context, + external_id: GarnetId, + internal_id: u32, + fsm: Arc, + original: Option<[(Context, Option>); 4]>, + completed: bool, + _backfill: Option, +} + +impl std::fmt::Debug for InsertGuard { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("InsertGuard") + .field("internal_id", &self.internal_id) + .field("completed", &self.completed) + .finish_non_exhaustive() + } +} + +impl Guard for InsertGuard { + type Id = u32; + + async fn complete(mut self) { + self.completed = true; + } + + fn id(&self) -> u32 { + self.internal_id + } +} + +impl Drop for InsertGuard { + fn drop(&mut self) { + if self.completed { + return; + } + + let mut restored = true; + if let Some(original) = self.original.take() { + for (context, value) in original { + restored &= match value { + Some(value) => self.callbacks.write_iid(&context, self.internal_id, &value), + None => { + self.callbacks.delete_iid(&context, self.internal_id) + || !self.callbacks.exists_iid(&context, self.internal_id, 0) + } + }; + } + } else { + for &term in Term::VARIANTS { + match term { + Term::Vector + | Term::Quantized + | Term::Attributes + | Term::Neighbors + | Term::ExtMap => { + let context = self.context.term(term); + restored &= self.callbacks.delete_iid(&context, self.internal_id) + || !self.callbacks.exists_iid(&context, self.internal_id, 0); + } + Term::Metadata | Term::IntMap => {} + } + } + let context = self.context.term(Term::IntMap); + restored &= self.callbacks.delete_eid(&context, &self.external_id) + || !self + .callbacks + .exists_eid(&context, &self.external_id, mem::size_of::()); + if restored { + restored = self.fsm.mark_free(&self.context, self.internal_id).is_ok(); + } + } + if !restored { + self.callbacks.log( + &self.context, + &format!( + "Error: insert rollback failed for ID {:?}; stored terms may be inconsistent.", + self.external_id, + ), + ); + } + } +} + /// The Garnet DataProvider implementation. pub(crate) struct GarnetProvider { /// Dimension of the full precision vectors @@ -130,12 +248,17 @@ pub(crate) struct GarnetProvider { max_degree: usize, /// Garnet storage engine callbacks callbacks: Callbacks, + pending_external_ids: DashMap, foldhash::fast::RandomState>, /// The quantizer the index will use, or None if NOQUANT is used. quantizer: Option>, /// Tracks whether the index is ready to operate fully quantized. all_quantized: AtomicBool, /// Per job tracker for quantization backfill completion backfills_completed: AtomicU64, + /// Lock for active backfill and update ranges + backfill_lock: Arc>>>, + /// Signals released range reservations + backfill_notify: Arc, /// Lock to ensure training only happens once. training_lock: Mutex<()>, /// Pool of pre-allocated buffers to use for neighbor lists @@ -156,7 +279,7 @@ pub(crate) struct GarnetProvider { /// Small cache for the start points' quantized vector data start_point_quant_cache: DashMap, foldhash::fast::RandomState>, /// Free space map to track internal IDs - fsm: FreeSpaceMap, + fsm: Arc, _phantom: PhantomData, } @@ -302,9 +425,12 @@ impl GarnetProvider { metric_type, max_degree, callbacks, + pending_external_ids: DashMap::with_hasher(foldhash::fast::RandomState::default()), quantizer, all_quantized: AtomicBool::new(all_quantized), backfills_completed: AtomicU64::new(0), + backfill_lock: Arc::new(Mutex::new(HashSet::new())), + backfill_notify: Arc::new(Condvar::new()), training_lock: Mutex::new(()), id_buffer_pool, filtered_ids_pool, @@ -314,11 +440,56 @@ impl GarnetProvider { start_point_cache, start_point_quant_cache, neighbor_cache, - fsm, + fsm: Arc::new(fsm), _phantom: PhantomData, }) } + pub(crate) async fn reserve_external_id( + &self, + id: &GarnetId, + ) -> Result, GarnetProviderError> { + let id_hash = self.pending_external_ids.hasher().hash_one(&id[..]); + for retry in 0..=RESERVATION_RETRY_LIMIT { + let mut receiver = match self.pending_external_ids.entry(id_hash) { + dashmap::mapref::entry::Entry::Occupied(_) if retry == RESERVATION_RETRY_LIMIT => { + break; + } + dashmap::mapref::entry::Entry::Occupied(entry) => entry.get().subscribe(), + dashmap::mapref::entry::Entry::Vacant(entry) => { + let (sender, _) = watch::channel(()); + entry.insert(sender); + return Ok(ExternalIdGuard { + pending: &self.pending_external_ids, + id_hash, + }); + } + }; + let _ = receiver.changed().await; + } + Err(GarnetProviderError::ReservationRetryLimit) + } + + fn reserve_backfill_range(&self, range: Range) -> Option { + if range.is_empty() { + return None; + } + + let mut ranges = self.backfill_lock.lock().unwrap(); + while ranges + .iter() + .any(|active| active.start < range.end && range.start < active.end) + { + ranges = self.backfill_notify.wait(ranges).unwrap(); + } + let _ = ranges.insert(range.clone()); + Some(BackfillGuard { + ranges: self.backfill_lock.clone(), + notify: self.backfill_notify.clone(), + range, + }) + } + /// Called during `VADD` to ensure a start point exists. /// If there isn't a start point yet, the given point will be set as the start point; if there /// is a start point already, we ensure the caches are populated. @@ -640,6 +811,7 @@ impl GarnetProvider { let start_id = (work_count * task_idx) as u32; let end_id = (work_count * (task_idx + 1)).min(max_id + 1) as u32; + let _backfill_guard = self.reserve_backfill_range(start_id..end_id); let mut v = vec![T::default(); self.dim]; let mut f = vec![0f32; self.dim]; let mut q = vec![0u8; quantizer.bytes()]; @@ -983,7 +1155,7 @@ impl DataProvider for GarnetProvider { type InternalId = u32; type ExternalId = GarnetId; type Error = GarnetProviderError; - type Guard = NoopGuard; + type Guard = InsertGuard; fn to_internal_id( &self, @@ -1021,7 +1193,52 @@ impl SetElement<(&[T], &[u8])> for GarnetProvider { id: &Self::ExternalId, element: (&[T], &[u8]), ) -> Result { - let internal_id = self.fsm.next_id(context)?; + let (internal_id, is_update) = match self.to_internal_id(context, id) { + Ok(existing_id) => { + context.set_insert_is_update(); + (self.fsm.existing_id(existing_id), true) + } + Err(_) => (self.fsm.next_id(context)?, false), + }; + + let backfill_guard = if self.quantizer.is_some() + && !self.all_quantized.load(Ordering::Acquire) + && (!internal_id.should_quantize() + || internal_id.id() <= internal_id.max_id_for_backfill()) + { + let end_id = internal_id + .id() + .checked_add(1) + .ok_or(FsmError::IdOutOfRange(internal_id.id()))?; + self.reserve_backfill_range(internal_id.id()..end_id) + } else { + None + }; + + let guard = InsertGuard { + callbacks: self.callbacks, + context: context.clone(), + external_id: id.clone(), + internal_id: internal_id.id(), + fsm: self.fsm.clone(), + original: is_update.then(|| { + [ + Term::Vector, + Term::Quantized, + Term::Attributes, + Term::Neighbors, + ] + .map(|term| { + let context = context.term(term); + let value = self + .callbacks + .read_varsize_iid::(&context, internal_id.id()); + (context, value) + }) + }), + completed: false, + _backfill: backfill_guard, + }; // Set quantization readiness if let Some(quantizer) = &self.quantizer @@ -1032,7 +1249,8 @@ impl SetElement<(&[T], &[u8])> for GarnetProvider { context.set_quantizer_ready(); } - let insert = || -> Result<(), Self::SetError> { + let mut error_term = Term::Vector; + let mut insert = || -> Result<(), Self::SetError> { self.callbacks .write_iid(&context.term(Term::Vector), internal_id.id(), element.0) .then_some(()) @@ -1040,6 +1258,7 @@ impl SetElement<(&[T], &[u8])> for GarnetProvider { if let Some(quantizer) = &self.quantizer && internal_id.should_quantize() { + error_term = Term::Quantized; let mut quant = self .quant_buffer_pool .get_ref(Undef::new(quantizer.bytes())); @@ -1053,50 +1272,44 @@ impl SetElement<(&[T], &[u8])> for GarnetProvider { .ok_or(GarnetError::Write)?; } if !element.1.is_empty() { + error_term = Term::Attributes; self.callbacks .write_iid(&context.term(Term::Attributes), internal_id.id(), element.1) .then_some(()) .ok_or(GarnetError::Write)?; } - self.callbacks - .write_iid(&context.term(Term::ExtMap), internal_id.id(), id) - .then_some(()) - .ok_or(GarnetError::Write)?; - self.callbacks - .write_eid( - &context.term(Term::IntMap), - id, - bytemuck::bytes_of(&internal_id.id()), - ) - .then_some(()) - .ok_or(GarnetError::Write)?; + if !is_update { + error_term = Term::ExtMap; + self.callbacks + .write_iid(&context.term(Term::ExtMap), internal_id.id(), id) + .then_some(()) + .ok_or(GarnetError::Write)?; + error_term = Term::IntMap; + self.callbacks + .write_eid( + &context.term(Term::IntMap), + id, + bytemuck::bytes_of(&internal_id.id()), + ) + .then_some(()) + .ok_or(GarnetError::Write)?; + } Ok(()) }; match insert() { Ok(()) => (), - Err(e) => { - // Clean up any potential data we inserted, but ignore failures. - let _ = self - .callbacks - .delete_iid(&context.term(Term::Vector), internal_id.id()); - let _ = self - .callbacks - .delete_iid(&context.term(Term::Quantized), internal_id.id()); - let _ = self - .callbacks - .delete_iid(&context.term(Term::Attributes), internal_id.id()); - let _ = self - .callbacks - .delete_iid(&context.term(Term::ExtMap), internal_id.id()); - let _ = self.callbacks.delete_eid(&context.term(Term::IntMap), id); - - self.fsm.mark_free(context, internal_id.id())?; + Err(e) if is_update => { + self.callbacks.log( + &context.term(error_term), + &format!("Error: update failed for ID {id:?}, term {error_term:?}: {e}."), + ); return Err(e); } + Err(e) => return Err(e), } - Ok(NoopGuard::new(internal_id.id())) + Ok(guard) } } @@ -2068,26 +2281,40 @@ impl InplaceDeleteStrategy> for DynamicQuantiza #[cfg(test)] mod tests { - use std::mem; + use std::{ + collections::HashMap, + ffi::c_void, + hash::BuildHasher, + mem, + ops::Range, + sync::{Arc, Mutex, atomic::Ordering, mpsc}, + thread, + time::{Duration, Instant}, + }; + use dashmap::DashMap; use diskann::{ graph::{ config::{self, defaults::GRAPH_SLACK_FACTOR}, search, }, - provider::{Delete, SetElement}, + provider::{DataProvider, Delete, Guard, SetElement}, }; use diskann_providers::index::wrapped_async::DiskANNIndex; + use diskann_utils::views::rowmajor::{self, Matrix, MatrixMut}; use diskann_vector::distance::Metric; use rand::Rng; use crate::{ SearchResults, VectorQuantType, dyn_index::DynIndex, - garnet::{Context, GarnetId, Term}, - provider::{GarnetProvider, QUANT_STATE_KEY}, + garnet::{ + Callbacks, Context, GarnetError, GarnetId, ReadDataCallback, RmwDataCallback, + TERM_BITMASK, Term, WriteCallback, + }, + provider::{GarnetProvider, GarnetProviderError, QUANT_STATE_KEY, RESERVATION_RETRY_LIMIT}, quantization::{GarnetQuantizer, Spherical1Bit}, - test_utils::Store, + test_utils::{LOGS, Store}, }; #[tokio::test] @@ -2107,10 +2334,938 @@ mod tests { let id = GarnetId::from(bytemuck::bytes_of(&0)); let res = provider.set_element(&ctx, &id, (&[0f32, 0f32], &[])).await; - assert!(res.is_ok()); + res.unwrap().complete().await; let res = provider.delete(&ctx, &id).await; assert!(res.is_ok()); + + let guard = provider + .set_element(&ctx, &id, (&[0f32, 0f32], &[])) + .await + .unwrap(); + store.clear_read_counts(); + drop(guard); + assert_eq!(store.int_map_reads(), 0); + assert!(store.get(ctx.term(Term::IntMap).get(), &id).is_none()); + assert_eq!(provider.fsm.total_used(), 0); + LOGS.with(|logs| assert!(logs.lock().unwrap().is_empty())); + } + + fn assert_waits_for_range( + provider: &GarnetProvider, + range: Range, + should_wait: bool, + operation: impl FnOnce() + Send, + ) { + let (started_tx, started_rx) = mpsc::channel(); + let (finished_tx, finished_rx) = mpsc::channel(); + thread::scope(|scope| { + let reservation = provider.reserve_backfill_range(range); + scope.spawn(move || { + started_tx.send(()).unwrap(); + operation(); + finished_tx.send(()).unwrap(); + }); + + started_rx.recv_timeout(Duration::from_secs(5)).unwrap(); + let early_result = finished_rx.recv_timeout(if should_wait { + Duration::from_millis(50) + } else { + Duration::from_secs(5) + }); + drop(reservation); + if should_wait { + assert_eq!(early_result, Err(mpsc::RecvTimeoutError::Timeout)); + finished_rx.recv_timeout(Duration::from_secs(5)).unwrap(); + } else { + early_result.unwrap(); + } + }); + assert!(provider.backfill_lock.lock().unwrap().is_empty()); + } + + fn train_for_backfill(provider: &GarnetProvider, ctx: &Context) { + let quantizer = provider.quantizer.as_ref().unwrap(); + let mut data = rowmajor::Owned::from_element(quantizer.required_vectors(), 2, 0.0f32); + for row in 0..data.nrows() { + data.row_mut(row) + .copy_from_slice(&[(row + 1) as f32, (row % 7 + 1) as f32]); + } + quantizer.train(Metric::L2, data.as_view()).unwrap(); + let mut state = vec![0u8]; + state.extend_from_slice(&quantizer.serialize().unwrap()); + assert!( + provider + .callbacks + .write_iid(&ctx.term(Term::Metadata), QUANT_STATE_KEY, &state) + ); + provider.fsm.enable_quantization(); + } + + #[test] + fn backfill_ranges_exclude_overlaps() { + let store = Arc::new(DashMap::new()); + let state = ParallelContext::new(store); + let ctx = state.context(); + let index = create_2d_f32_index_with_callbacks( + VectorQuantType::NoQuant, + Metric::L2, + ParallelContext::callbacks(), + &ctx, + ); + let provider = index.inner.provider(); + + for (active, requested, should_wait) in [ + (0..10, 5..6, true), + (5..6, 0..10, true), + (0..10, 0..10, true), + (0..10, 9..20, true), + (0..10, 10..20, false), + (0..10, 5..5, false), + ] { + assert_waits_for_range(provider, active, should_wait, || { + let _reservation = provider.reserve_backfill_range(requested); + }); + } + } + + #[test] + fn updates_reserve_backfill_ranges_only_when_needed() { + for (quant_type, train, finish, above_boundary, should_wait) in [ + (VectorQuantType::NoQuant, false, false, false, false), + (VectorQuantType::Q8, false, false, false, false), + (VectorQuantType::Bin, false, false, false, true), + (VectorQuantType::Bin, true, false, false, true), + (VectorQuantType::Bin, true, false, true, false), + (VectorQuantType::Bin, true, true, false, false), + ] { + let store = Arc::new(DashMap::new()); + let state = ParallelContext::new(store.clone()); + let ctx = state.context(); + let index = create_2d_f32_index_with_callbacks( + quant_type, + Metric::L2, + ParallelContext::callbacks(), + &ctx, + ); + let provider = index.inner.provider(); + let original = [0.0f32, 1.0]; + let mut id = GarnetId::from(bytemuck::bytes_of(&42u32)); + let runtime = tokio::runtime::Builder::new_current_thread() + .build() + .unwrap(); + provider.maybe_set_start_point(&ctx, &original).unwrap(); + let guard = runtime + .block_on(provider.set_element(&ctx, &id, (&original, &[]))) + .unwrap(); + runtime.block_on(guard.complete()); + if train { + train_for_backfill(provider, &ctx); + } + if above_boundary { + id = GarnetId::from(bytemuck::bytes_of(&43u32)); + let guard = runtime + .block_on(provider.set_element(&ctx, &id, (&original, &[]))) + .unwrap(); + runtime.block_on(guard.complete()); + } + if finish { + assert!(provider.backfill_quant_vectors(&ctx, 0, 1)); + assert!(provider.all_quantized.load(Ordering::Acquire)); + } + let internal_id = parallel_get(&store, ctx.term(Term::IntMap).get(), &id).unwrap(); + let internal_id = bytemuck::pod_read_unaligned::(&internal_id); + let update_ctx = state.context(); + assert_waits_for_range(provider, internal_id..internal_id + 1, should_wait, || { + let updated = [1.0f32, 0.0]; + let guard = runtime + .block_on(provider.set_element(&update_ctx, &id, (&updated, &[]))) + .unwrap(); + runtime.block_on(guard.complete()); + assert!(update_ctx.insert_is_update()); + if let Some(quantizer) = &provider.quantizer + && quantizer.is_trained() + { + let mut expected = vec![0u8; quantizer.bytes()]; + quantizer.compress(&updated, &mut expected).unwrap(); + assert_eq!( + parallel_get( + &store, + ctx.term(Term::Quantized).get(), + bytemuck::bytes_of(&internal_id), + ), + Some(expected) + ); + } + }); + } + } + + #[test] + fn backfill_waits_for_overlapping_updates() { + for (range, should_wait) in [(1..2, true), (2..3, false)] { + let store = Arc::new(DashMap::new()); + let state = ParallelContext::new(store); + let ctx = state.context(); + let index = create_2d_f32_index_with_callbacks( + VectorQuantType::Bin, + Metric::L2, + ParallelContext::callbacks(), + &ctx, + ); + let provider = index.inner.provider(); + let original = [0.0f32, 1.0]; + let id = GarnetId::from(bytemuck::bytes_of(&42u32)); + provider.maybe_set_start_point(&ctx, &original).unwrap(); + DynIndex::insert(&index, &ctx, &id, bytemuck::cast_slice(&original), &[]).unwrap(); + train_for_backfill(provider, &ctx); + assert_waits_for_range(provider, range, should_wait, || { + assert!(provider.backfill_quant_vectors(&ctx, 0, 1)); + }); + assert!(provider.all_quantized.load(Ordering::Acquire)); + } + } + + /// Per-insert fault injection and synchronization for tests. + #[derive(Default)] + struct InsertControl { + /// One-shot failure: target term and matching operations to skip. + failure: Option<(u64, usize)>, + /// One-shot pause: target term, arrival sender, and resume receiver. + pause: Option<(u64, mpsc::Sender<()>, mpsc::Receiver<()>)>, + /// Store snapshot before the first vector write. + before_vector: Option, Vec>>, + } + + /// Parallel test callback state; the encoded pointer's low three bits hold the term tag. + struct ParallelContext { + /// Mock storage shared across threads. + store: Arc, Vec>>, + /// Per-insert fault, pause, and snapshot state. + control: Mutex, + } + + impl ParallelContext { + fn new(store: Arc, Vec>>) -> Box { + Box::new(Self { + store, + control: Mutex::new(InsertControl::default()), + }) + } + + fn context(&self) -> Context { + Context::new(std::ptr::from_ref(self).expose_provenance() as u64) + } + + unsafe fn from_context<'a>(context: u64) -> &'a Self { + let pointer = + std::ptr::with_exposed_provenance::((context & !TERM_BITMASK) as usize); + // SAFETY: Callers keep the boxed context alive until all callback operations finish. + unsafe { &*pointer } + } + + fn callbacks() -> Callbacks { + Callbacks::new( + parallel_read, + controlled_insert_write, + parallel_delete, + controlled_insert_rmw, + parallel_filter, + parallel_log, + ) + } + + fn fail_insert_operation(&self, context: u64) -> bool { + let mut control = self.control.lock().unwrap(); + let term = context & TERM_BITMASK; + if control + .pause + .as_ref() + .is_some_and(|(pause_term, _, _)| *pause_term == term) + { + let (_, entered, resume) = control.pause.take().unwrap(); + let _ = entered.send(()); + if resume.recv_timeout(Duration::from_secs(10)).is_err() { + return true; + } + } + if term == Term::Vector as u64 && control.before_vector.is_none() { + control.before_vector = Some(parallel_snapshot(&self.store)); + } + if let Some((failure_term, skip)) = control.failure.as_mut() + && *failure_term == term + { + if *skip == 0 { + control.failure = None; + return true; + } + *skip -= 1; + } + false + } + } + + fn parallel_key(context: u64, key: &[u8]) -> Vec { + let mut encoded = bytemuck::bytes_of(&(context & TERM_BITMASK)).to_vec(); + encoded.extend_from_slice(key); + encoded + } + + fn parallel_get( + store: &DashMap, Vec>, + context: u64, + key: &[u8], + ) -> Option> { + store + .get(¶llel_key(context, key)) + .map(|value| value.clone()) + } + + fn parallel_snapshot(store: &DashMap, Vec>) -> HashMap, Vec> { + store + .iter() + .map(|entry| (entry.key().clone(), entry.value().clone())) + .collect() + } + + unsafe extern "C" fn parallel_read( + context: u64, + count: u32, + _length_hint: u32, + keys: *const u8, + keys_len: usize, + callback: ReadDataCallback, + callback_context: *mut c_void, + ) { + let state = unsafe { ParallelContext::from_context(context) }; + let mut keys = unsafe { std::slice::from_raw_parts(keys, keys_len) }; + for index in 0..count { + let length = bytemuck::pod_read_unaligned::(&keys[..4]) as usize; + let key = &keys[4..4 + length]; + if let Some(value) = parallel_get(&state.store, context, key) { + unsafe { callback(index, callback_context, value.as_ptr(), value.len()) }; + } + keys = &keys[4 + length..]; + } + } + + unsafe extern "C" fn parallel_delete(context: u64, key: *const u8, key_len: usize) -> bool { + let state = unsafe { ParallelContext::from_context(context) }; + let key = unsafe { std::slice::from_raw_parts(key, key_len) }; + state.store.remove(¶llel_key(context, key)).is_some() + } + + unsafe extern "C" fn parallel_filter(_context: u64, _data: *const u8, _length: usize) -> bool { + true + } + + unsafe extern "C" fn parallel_log(_context: u64, _message: *const u8, _length: usize) {} + + unsafe extern "C" fn controlled_insert_write( + context: u64, + key: *const u8, + key_len: usize, + value: *const u8, + value_len: usize, + ) -> bool { + let state = unsafe { ParallelContext::from_context(context) }; + if state.fail_insert_operation(context) { + return false; + } + let key = unsafe { std::slice::from_raw_parts(key, key_len) }; + let value = unsafe { std::slice::from_raw_parts(value, value_len) }; + state + .store + .insert(parallel_key(context, key), value.to_vec()); + true + } + + unsafe extern "C" fn controlled_insert_rmw( + context: u64, + key: *const u8, + key_len: usize, + value_len: usize, + callback: RmwDataCallback, + callback_context: *mut c_void, + ) -> bool { + let state = unsafe { ParallelContext::from_context(context) }; + if state.fail_insert_operation(context) { + return false; + } + let key = unsafe { std::slice::from_raw_parts(key, key_len) }; + let mut value = state + .store + .entry(parallel_key(context, key)) + .or_insert_with(|| vec![0; value_len]); + unsafe { callback(callback_context, value.as_mut_ptr(), value.len()) }; + true + } + + fn wait_for_pending_receiver(provider: &GarnetProvider, id: &GarnetId) -> bool { + let id_hash = provider.pending_external_ids.hasher().hash_one(&id[..]); + let deadline = Instant::now() + Duration::from_secs(5); + loop { + if provider + .pending_external_ids + .get(&id_hash) + .is_some_and(|sender| sender.receiver_count() > 0) + { + return true; + } + if Instant::now() >= deadline { + return false; + } + thread::yield_now(); + } + } + + fn concurrent_insert_outcomes(first_fails: bool, second_fails: bool) { + for quant_type in [ + VectorQuantType::NoQuant, + VectorQuantType::Bin, + VectorQuantType::Q8, + ] { + // failure_term tells where to fail, and skip tells how many to succeed before we fail + for (failure_term, skip) in [ + (Term::Vector, 0), + (Term::Quantized, 0), + (Term::Attributes, 0), + (Term::ExtMap, 0), + (Term::IntMap, 0), + (Term::Neighbors, 0), + (Term::Neighbors, 1), + ] { + if matches!(failure_term, Term::Quantized) && quant_type != VectorQuantType::Q8 { + continue; + } + // Set up the index and capture its state before either insert. + let store = Arc::new(DashMap::new()); + let state = ParallelContext::new(store.clone()); + let ctx = state.context(); + let callbacks = ParallelContext::callbacks(); + let index = + create_2d_f32_index_with_callbacks(quant_type, Metric::L2, callbacks, &ctx); + let provider = index.inner.provider(); + let id = GarnetId::from(bytemuck::bytes_of(&42u32)); + let first_vector = [0.0f32, 1.0]; + let second_vector = [1.0f32, 0.0]; + provider.maybe_set_start_point(&ctx, &first_vector).unwrap(); + let initial_used = provider.fsm.total_used(); + let mut initial_snapshot = parallel_snapshot(&store); + initial_snapshot.retain(|key, _| { + bytemuck::pod_read_unaligned::(&key[..8]) != Term::Metadata as u64 + }); + let (entered_tx, entered_rx) = mpsc::channel(); + let (resume_tx, resume_rx) = mpsc::channel(); + let failure = (failure_term as u64, skip); + + // Pause the first insert while it owns the external ID, verify the second + // waits for that ID, then resume the first and collect both results. + let (first, second) = thread::scope(|scope| { + let first = scope.spawn(|| { + let state = ParallelContext::new(store.clone()); + *state.control.lock().unwrap() = InsertControl { + failure: first_fails.then_some(failure), + pause: Some((failure.0, entered_tx, resume_rx)), + before_vector: None, + }; + let context = state.context(); + let result = DynIndex::insert( + &index, + &context, + &id, + bytemuck::cast_slice(&first_vector), + b"first", + ); + (result.is_ok(), context.insert_is_update()) + }); + let entered = entered_rx.recv_timeout(Duration::from_secs(5)); + let second = scope.spawn(|| { + let state = ParallelContext::new(store.clone()); + let second_failure = if !first_fails + && (failure.0 == Term::ExtMap as u64 + || failure.0 == Term::IntMap as u64) + { + (Term::Attributes as u64, 0) + } else { + (failure.0, if first_fails { failure.1 } else { 0 }) + }; + *state.control.lock().unwrap() = InsertControl { + failure: second_fails.then_some(second_failure), + pause: None, + before_vector: None, + }; + let context = state.context(); + let result = DynIndex::insert( + &index, + &context, + &id, + bytemuck::cast_slice(&second_vector), + b"second", + ); + let snapshot = state.control.lock().unwrap().before_vector.take().unwrap(); + ((result.is_ok(), context.insert_is_update()), snapshot) + }); + let waiting = wait_for_pending_receiver(provider, &id); + let _ = resume_tx.send(()); + let results = (first.join().unwrap(), second.join().unwrap()); + assert!(entered.is_ok(), "first insert never acquired ownership"); + assert!(waiting, "second insert did not wait for ownership"); + results + }); + + // Verify success/update classification and rollback to the expected state. + assert_eq!( + first, + (!first_fails, false), + "first: {quant_type:?}, fault {failure:?}" + ); + assert_eq!( + second.0, + (!second_fails, !first_fails), + "second: {quant_type:?}, fault {failure:?}" + ); + if second_fails { + let mut current_snapshot = parallel_snapshot(&store); + if first_fails { + current_snapshot.retain(|key, _| { + bytemuck::pod_read_unaligned::(&key[..8]) != Term::Metadata as u64 + }); + assert_eq!(current_snapshot, initial_snapshot); + } else { + assert_eq!(current_snapshot, second.1); + } + } + // Check ownership cleanup, live member records, and freed IDs. + assert!(provider.pending_external_ids.is_empty()); + let has_member = !first_fails || !second_fails; + assert_eq!( + provider.fsm.total_used(), + initial_used + usize::from(has_member) + ); + let mapping = parallel_get(&store, ctx.term(Term::IntMap).get(), &id); + assert_eq!(mapping.is_some(), has_member); + let current_id = mapping.as_deref().map(bytemuck::pod_read_unaligned::); + for internal_id in 1..=provider.fsm.max_id() { + let key = bytemuck::bytes_of(&internal_id); + if Some(internal_id) == current_id { + let (vector, attrs) = if second_fails { + (&first_vector, &b"first"[..]) + } else { + (&second_vector, &b"second"[..]) + }; + assert_eq!(provider.get_full_vector(&ctx, internal_id).unwrap(), vector); + assert_eq!( + parallel_get(&store, ctx.term(Term::Attributes).get(), key), + Some(attrs.to_vec()) + ); + assert_eq!( + parallel_get(&store, ctx.term(Term::ExtMap).get(), key), + Some(id.to_vec()) + ); + if let Some(quantizer) = &provider.quantizer + && quantizer.is_trained() + { + let mut expected = vec![0; quantizer.bytes()]; + quantizer.compress(vector, &mut expected).unwrap(); + assert_eq!( + parallel_get(&store, ctx.term(Term::Quantized).get(), key), + Some(expected) + ); + } + } else { + assert!(provider.fsm.is_free(&ctx, internal_id).unwrap()); + for term in [ + Term::Vector, + Term::Quantized, + Term::Attributes, + Term::Neighbors, + Term::ExtMap, + ] { + assert!(parallel_get(&store, ctx.term(term).get(), key).is_none()); + } + } + } + + // Insert again, which will update or insert depending on the first two. + let retry_context = state.context(); + DynIndex::insert( + &index, + &retry_context, + &id, + bytemuck::cast_slice(&first_vector), + b"retry", + ) + .unwrap(); + assert_eq!(retry_context.insert_is_update(), has_member); + assert_eq!(provider.fsm.total_used(), initial_used + 1); + assert!(provider.pending_external_ids.is_empty()); + } + } + } + + #[test] + fn concurrent_inserts_both_succeed() { + concurrent_insert_outcomes(false, false); + } + + #[test] + fn concurrent_inserts_first_fails() { + concurrent_insert_outcomes(true, false); + } + + #[test] + fn concurrent_inserts_second_fails() { + concurrent_insert_outcomes(false, true); + } + + #[test] + fn concurrent_inserts_both_fail() { + concurrent_insert_outcomes(true, true); + } + + #[tokio::test] + async fn pending_external_ids_serialize_and_wake_waiters() { + use std::{future::Future, pin::pin, task::Poll}; + + let store = Store::new(); + let ctx = Context::new(0); + let provider = GarnetProvider::::new( + 2, + VectorQuantType::NoQuant, + Metric::L2, + 10, + store.callbacks(), + &ctx, + ) + .unwrap(); + let id = GarnetId::from(&b"same"[..]); + let other_id = GarnetId::from(&b"other"[..]); + let id_hash = provider.pending_external_ids.hasher().hash_one(&id[..]); + let receiver_count = || { + provider + .pending_external_ids + .get(&id_hash) + .unwrap() + .receiver_count() + }; + // Reserve one ID and register two waiters; a different ID remains available. + let owner = provider.reserve_external_id(&id).await.unwrap(); + assert_eq!(receiver_count(), 0); + let mut second = pin!(provider.reserve_external_id(&id)); + let mut third = pin!(provider.reserve_external_id(&id)); + let mut task = std::task::Context::from_waker(std::task::Waker::noop()); + assert!(second.as_mut().poll(&mut task).is_pending()); + assert!(third.as_mut().poll(&mut task).is_pending()); + assert_eq!(receiver_count(), 2); + drop(provider.reserve_external_id(&other_id).await.unwrap()); + + // Release owners in turn; only one waiter can hold the ID at a time. + drop(owner); + let Poll::Ready(Ok(second_owner)) = second.as_mut().poll(&mut task) else { + panic!("second waiter did not acquire the released ID"); + }; + assert!(third.as_mut().poll(&mut task).is_pending()); + assert_eq!(receiver_count(), 1); + drop(second_owner); + drop(third.await.unwrap()); + assert!(provider.pending_external_ids.is_empty()); + + // Cancelling a waiter removes its subscription without releasing the owner's ID. + let owner = provider.reserve_external_id(&id).await.unwrap(); + let mut cancelled = Box::pin(provider.reserve_external_id(&id)); + assert!(cancelled.as_mut().poll(&mut task).is_pending()); + assert_eq!(receiver_count(), 1); + drop(cancelled); + assert_eq!(receiver_count(), 0); + drop(owner); + drop(provider.reserve_external_id(&id).await.unwrap()); + assert!(provider.pending_external_ids.is_empty()); + } + + #[tokio::test] + async fn external_id_reservation_retry_limit() { + use std::{future::Future, pin::pin, task::Poll}; + + let store = Store::new(); + let ctx = Context::new(0); + let provider = GarnetProvider::::new( + 2, + VectorQuantType::NoQuant, + Metric::L2, + 10, + store.callbacks(), + &ctx, + ) + .unwrap(); + let id = GarnetId::from(&b"contended"[..]); + let id_hash = provider.pending_external_ids.hasher().hash_one(&id[..]); + let mut task = std::task::Context::from_waker(std::task::Waker::noop()); + + // Exercise both acquisition and continued contention on the final allowed retry. + for acquire_on_last_retry in [false, true] { + let mut owner = provider.reserve_external_id(&id).await.unwrap(); + let mut waiter = pin!(provider.reserve_external_id(&id)); + assert!(waiter.as_mut().poll(&mut task).is_pending()); + + // Reacquire the ID before polling the waiter to force repeated contention. + for _ in 1..RESERVATION_RETRY_LIMIT { + drop(owner); + owner = provider.reserve_external_id(&id).await.unwrap(); + assert!(waiter.as_mut().poll(&mut task).is_pending()); + } + + drop(owner); + if acquire_on_last_retry { + let Poll::Ready(Ok(guard)) = waiter.as_mut().poll(&mut task) else { + panic!("last reservation retry did not acquire the released ID"); + }; + drop(guard); + } else { + // Exhaustion removes the waiter but leaves the current owner's reservation. + let owner = provider.reserve_external_id(&id).await.unwrap(); + assert!(matches!( + waiter.as_mut().poll(&mut task), + Poll::Ready(Err(GarnetProviderError::ReservationRetryLimit)) + )); + assert_eq!( + provider + .pending_external_ids + .get(&id_hash) + .unwrap() + .receiver_count(), + 0 + ); + drop(owner); + } + assert!(provider.pending_external_ids.is_empty()); + } + } + + #[test] + fn member_mutations_wait_for_insert_owner() { + enum Mutation { + SetAttributes, + DeleteAttributes, + Remove, + } + + for mutation in [ + Mutation::SetAttributes, + Mutation::DeleteAttributes, + Mutation::Remove, + ] { + // Create an existing member with attributes for each mutation. + let store = Arc::new(DashMap::new()); + let state = ParallelContext::new(store.clone()); + let ctx = state.context(); + let index = create_2d_f32_index_with_callbacks( + VectorQuantType::NoQuant, + Metric::L2, + ParallelContext::callbacks(), + &ctx, + ); + let provider = index.inner.provider(); + let id = GarnetId::from(bytemuck::bytes_of(&42u32)); + let vector = [0.0f32, 1.0]; + provider.maybe_set_start_point(&ctx, &vector).unwrap(); + DynIndex::insert(&index, &ctx, &id, bytemuck::cast_slice(&vector), b"before").unwrap(); + let internal_id = provider.to_internal_id(&ctx, &id).unwrap(); + + // Hold the ID reservation and verify the mutation waits until it is released. + let owner = index.run(|_| provider.reserve_external_id(&id)).unwrap(); + thread::scope(|scope| { + let task = scope.spawn(|| match mutation { + Mutation::SetAttributes => { + DynIndex::set_attributes(&index, &ctx, &id, b"after") + } + Mutation::DeleteAttributes => DynIndex::delete_attributes(&index, &ctx, &id), + Mutation::Remove => DynIndex::remove(&index, &ctx, &id), + }); + let waiting = wait_for_pending_receiver(provider, &id); + drop(owner); + task.join().unwrap().unwrap(); + assert!(waiting, "member mutation did not wait for insert ownership"); + }); + + // Check the mutation's effect on attributes and membership, then reservation cleanup. + let attrs = parallel_get( + &store, + ctx.term(Term::Attributes).get(), + bytemuck::bytes_of(&internal_id), + ); + assert_eq!( + attrs, + match mutation { + Mutation::SetAttributes => Some(b"after".to_vec()), + _ => None, + } + ); + assert_eq!( + provider.to_internal_id(&ctx, &id).is_ok(), + !matches!(mutation, Mutation::Remove) + ); + assert!(provider.pending_external_ids.is_empty()); + } + } + + #[tokio::test] + async fn update_reuses_id_and_preserves_entry_on_write_failure() { + unsafe extern "C" fn fail_write( + _context: u64, + _key: *const u8, + _key_len: usize, + _value: *const u8, + _value_len: usize, + ) -> bool { + false + } + + unsafe extern "C" fn fail_after_vector_write( + context: u64, + key: *const u8, + key_len: usize, + value: *const u8, + value_len: usize, + ) -> bool { + context & TERM_BITMASK == Term::Vector as u64 + && unsafe { + (Store::attach().callbacks().write_callback())( + context, key, key_len, value, value_len, + ) + } + } + + for quant_type in [ + VectorQuantType::NoQuant, + VectorQuantType::Bin, + VectorQuantType::Q8, + ] { + let store = Store::new(); + let ctx = Context::new(0); + let mut provider = + GarnetProvider::::new(2, quant_type, Metric::L2, 10, store.callbacks(), &ctx) + .unwrap(); + let id = GarnetId::from(bytemuck::bytes_of(&42u32)); + let original = [0.0f32, 1.0]; + provider.maybe_set_start_point(&ctx, &original).unwrap(); + provider + .set_element(&ctx, &id, (&original, b"old")) + .await + .unwrap() + .complete() + .await; + let internal_id = store.get(ctx.term(Term::IntMap).get(), &id).unwrap(); + let max_id = provider.fsm.max_id(); + let total_used = provider.fsm.total_used(); + + let updated = [1.0f32, 0.0]; + provider + .set_element(&ctx, &id, (&updated, b"new")) + .await + .unwrap() + .complete() + .await; + assert!(ctx.insert_is_update()); + assert_eq!(provider.fsm.max_id(), max_id); + assert_eq!(provider.fsm.total_used(), total_used); + assert_eq!( + store.get(ctx.term(Term::IntMap).get(), &id), + Some(internal_id.clone()) + ); + assert_eq!( + store.get(ctx.term(Term::ExtMap).get(), &internal_id), + Some(id.to_vec()) + ); + assert_eq!( + store.get(ctx.term(Term::Vector).get(), &internal_id), + Some(bytemuck::cast_slice::(&updated).to_vec()) + ); + assert_eq!( + store.get(ctx.term(Term::Attributes).get(), &internal_id), + Some(b"new".to_vec()) + ); + if let Some(quantizer) = &provider.quantizer + && quantizer.is_trained() + { + let mut expected = vec![0u8; quantizer.bytes()]; + quantizer.compress(&updated, &mut expected).unwrap(); + assert_eq!( + store.get(ctx.term(Term::Quantized).get(), &internal_id), + Some(expected) + ); + } + + let quantized_before = store.get(ctx.term(Term::Quantized).get(), &internal_id); + let guard = provider + .set_element(&ctx, &id, (&original, b"discarded")) + .await + .unwrap(); + store.clear_read_counts(); + drop(guard); + assert_eq!(store.full_reads(), 0); + if quantized_before.is_some() { + assert_eq!(store.quant_reads(), 0); + } + assert_eq!( + store.get(ctx.term(Term::Vector).get(), &internal_id), + Some(bytemuck::cast_slice::(&updated).to_vec()) + ); + assert_eq!( + store.get(ctx.term(Term::Attributes).get(), &internal_id), + Some(b"new".to_vec()) + ); + assert_eq!( + store.get(ctx.term(Term::Quantized).get(), &internal_id), + quantized_before + ); + + for write_callback in [ + fail_write as WriteCallback, + fail_after_vector_write as WriteCallback, + ] { + let callbacks = store.callbacks(); + provider.callbacks = Callbacks::new( + callbacks.read_callback(), + write_callback, + callbacks.delete_callback(), + callbacks.rmw_callback(), + callbacks.filter_callback(), + callbacks.log_callback(), + ); + let error = provider + .set_element(&ctx, &id, (&original, b"failed")) + .await + .unwrap_err(); + assert!(matches!( + error, + GarnetProviderError::Garnet(GarnetError::Write) + )); + assert!(provider.backfill_lock.lock().unwrap().is_empty()); + assert_eq!(provider.fsm.max_id(), max_id); + assert_eq!(provider.fsm.total_used(), total_used); + assert_eq!( + store.get(ctx.term(Term::IntMap).get(), &id), + Some(internal_id.clone()) + ); + assert_eq!( + store.get(ctx.term(Term::ExtMap).get(), &internal_id), + Some(id.to_vec()) + ); + assert_eq!( + store.get(ctx.term(Term::Vector).get(), &internal_id), + Some(bytemuck::cast_slice::(&updated).to_vec()) + ); + assert_eq!( + store.get(ctx.term(Term::Quantized).get(), &internal_id), + quantized_before + ); + assert_eq!( + store.get(ctx.term(Term::Attributes).get(), &internal_id), + Some(b"new".to_vec()) + ); + } + } } fn create_2d_f32_index( @@ -2118,9 +3273,18 @@ mod tests { metric: Metric, store: &Store, ctx: &Context, + ) -> DiskANNIndex> { + create_2d_f32_index_with_callbacks(quant_type, metric, store.callbacks(), ctx) + } + + fn create_2d_f32_index_with_callbacks( + quant_type: VectorQuantType, + metric: Metric, + callbacks: Callbacks, + ctx: &Context, ) -> DiskANNIndex> { let provider = - GarnetProvider::::new(2, quant_type, metric, 10, store.callbacks(), ctx).unwrap(); + GarnetProvider::::new(2, quant_type, metric, 10, callbacks, ctx).unwrap(); let config = config::Builder::new( (10.0 / GRAPH_SLACK_FACTOR) as usize, diff --git a/diskann-garnet/src/test_utils.rs b/diskann-garnet/src/test_utils.rs index 827d637549..9b6fa47932 100644 --- a/diskann-garnet/src/test_utils.rs +++ b/diskann-garnet/src/test_utils.rs @@ -19,6 +19,7 @@ thread_local! { pub static LOGS: Mutex> = const { Mutex::new(Vec::new()) }; pub static FULL_READS: AtomicUsize = const { AtomicUsize::new(0) }; pub static QUANT_READS: AtomicUsize = const { AtomicUsize::new(0) }; + pub static INT_MAP_READS: AtomicUsize = const { AtomicUsize::new(0) }; } /// Mock storage for testing. @@ -58,6 +59,7 @@ impl Store { }); FULL_READS.with(|fr| fr.store(0, Ordering::Release)); QUANT_READS.with(|qr| qr.store(0, Ordering::Release)); + INT_MAP_READS.with(|reads| reads.store(0, Ordering::Release)); } pub fn set(&self, context: u64, key: &[u8], value: &[u8]) { @@ -102,6 +104,7 @@ impl Store { pub fn clear_read_counts(&self) { FULL_READS.with(|fr| fr.store(0, Ordering::Release)); QUANT_READS.with(|qr| qr.store(0, Ordering::Release)); + INT_MAP_READS.with(|reads| reads.store(0, Ordering::Release)); } pub fn full_reads(&self) -> usize { @@ -112,6 +115,10 @@ impl Store { QUANT_READS.with(|qr| qr.load(Ordering::Acquire)) } + pub fn int_map_reads(&self) -> usize { + INT_MAP_READS.with(|reads| reads.load(Ordering::Acquire)) + } + pub fn log(&self, context: u64, msg: &str) { LOGS.with(|l| { let mut guard = l.lock().unwrap(); @@ -129,13 +136,14 @@ unsafe extern "C" fn test_read( cb: ReadDataCallback, cb_ctx: *mut c_void, ) { + if ctx & TERM_BITMASK == Term::IntMap as u64 { + INT_MAP_READS.with(|reads| reads.fetch_add(1, Ordering::AcqRel)); + } let ids = unsafe { slice::from_raw_parts(id_bytes, id_len) }; let mut pos = 0usize; for idx in 0..count { - let mut len = 0u32; - let len_bytes = bytemuck::bytes_of_mut(&mut len); - len_bytes.copy_from_slice(&ids[pos..pos + 4]); + let len = bytemuck::pod_read_unaligned::(&ids[pos..pos + 4]); pos += 4;