From 24f453b294735b049fa5dfeb0355032aab5b07e7 Mon Sep 17 00:00:00 2001 From: Gabriel Date: Wed, 30 Sep 2026 16:14:06 +0800 Subject: [PATCH 1/2] feat: expose distance-bounded nearest search in C and C++ --- README.md | 40 ++++++ include/lance/lance.h | 25 ++++ include/lance/lance.hpp | 14 ++ src/scanner.rs | 69 ++++++++++ tests/c_api_test.rs | 268 +++++++++++++++++++++++++++++++++++++ tests/cpp/test_c_api.c | 39 ++++++ tests/cpp/test_cpp_api.cpp | 37 +++++ tests/multivector_test.rs | 29 ++++ 8 files changed, 521 insertions(+) diff --git a/README.md b/README.md index 9ceccdc..f4dccb2 100644 --- a/README.md +++ b/README.md @@ -69,6 +69,46 @@ Based on the [liblance RFC](https://github.com/lance-format/lance/discussions/60 | [x] | Dataset metadata | `lance_dataset_version()`, `lance_dataset_count_rows()`, `lance_dataset_latest_version()` | | [x] | Filter pushdown | `lance_scanner_set_substrait_filter()` accepts a serialized Substrait `ExtendedExpression`; `lance_scanner_additional_sql_filter()` adds SQL predicates with AND before scanning starts | +## Distance-bounded vector search + +After configuring a single-vector nearest-neighbor query, use +`lance_scanner_set_distance_range` or C++ `Scanner::distance_range` to restrict +results to `lower_bound <= _distance < upper_bound`: + +```cpp +const float query[] = {1.0f, 0.0f}; +auto scanner = dataset.scan(); +scanner.nearest("embedding", query, 2, 100) + .metric(LANCE_METRIC_L2) + .distance_range(std::nullopt, 0.5f); +``` + +This returns **at most 100** neighbors with distance below `0.5`. It is a +range-constrained Top-K search, not an unbounded enumeration of all matches. +For an annulus, use `.distance_range(0.2f, 0.5f)`; for a lower bound only, use +`.distance_range(0.2f)`; `.distance_range()` clears both bounds. In C, pass +pointers to bounds and `NULL` for an unbounded side: + +```c +float upper_bound = 0.5f; +int32_t status = lance_scanner_set_distance_range(scanner, NULL, &upper_bound); +/* Check status and lance_last_error_* before starting the scan. */ +``` + +Bounds are copied and must be finite. When both are set, the lower bound must +be strictly smaller than the upper bound. Negative bounds are allowed (for +example, Dot distances can be negative). Distances use the selected metric's +units: L2 reports **squared Euclidean distance**, so a geometric radius `r` +corresponds to an upper bound of `r * r`, with the boundary excluded. Index-based +search retains Lance's approximate candidate selection; the range does not +guarantee exhaustive recall. Use `.use_index(false)` for an exact scan, still +subject to `k`. + +Set bounds after `nearest` and before starting the scan. Replacing the nearest +query clears the bounds; an invalid range leaves the previous range unchanged. +Multi-vector queries are not supported by this setter because their scores +aggregate distances across subvectors. + ## Multi-vector search Use `lance_scanner_nearest_multivector` or the C++ `Scanner::nearest_multivector` diff --git a/include/lance/lance.h b/include/lance/lance.h index 7630913..d29fc0b 100644 --- a/include/lance/lance.h +++ b/include/lance/lance.h @@ -2125,6 +2125,31 @@ int32_t lance_scanner_nearest( uint32_t k ); +/** + * Restrict a single-vector nearest query to lower_bound <= _distance < upper_bound. + * + * Call after lance_scanner_nearest and before starting the scan. Both bounds + * are copied before returning; NULL means unbounded on that side. Passing NULL + * for both clears the range. A successful replacement nearest query also clears + * the range. Multi-vector queries and scans without nearest are rejected. + * + * Bounds must be finite; negative distances are allowed. When both are present, + * lower_bound must be strictly smaller than upper_bound. Invalid calls leave the + * previous range unchanged. Distances use the configured metric's units (L2 is + * squared Euclidean distance). Results still have the nearest query's k cap; + * index-based search remains approximate, not an exhaustive range enumeration. + * + * @param scanner Scanner with a single-vector nearest query. + * @param lower_bound Inclusive lower bound, or NULL. + * @param upper_bound Exclusive upper bound, or NULL. + * @return 0 on success, -1 on error (check lance_last_error_*). + */ +int32_t lance_scanner_set_distance_range( + LanceScanner* scanner, + const float* lower_bound, + const float* upper_bound +); + /** * Set one multi-vector query on a List> column. * Inner vectors must be non-nullable and contain no null elements; the outer list may be nullable. diff --git a/include/lance/lance.hpp b/include/lance/lance.hpp index 286724e..882059a 100644 --- a/include/lance/lance.hpp +++ b/include/lance/lance.hpp @@ -1622,6 +1622,20 @@ class Scanner { return *this; } + /// Restrict single-vector nearest results to [lower_bound, upper_bound). + /// Call after nearest() and before scanning. Omitted bounds are unbounded; + /// distance_range() clears both. Bounds must be finite and lower < upper + /// when both are present. Uses metric distance units (squared L2 for L2). + /// The nearest query's k cap and ANN candidate selection still apply. + Scanner& distance_range(std::optional lower_bound = std::nullopt, + std::optional upper_bound = std::nullopt) { + if (lance_scanner_set_distance_range(handle_.get(), + lower_bound ? &*lower_bound : nullptr, + upper_bound ? &*upper_bound : nullptr) != 0) + check_error(); + return *this; + } + /// One multi-vector query, copied from dimension * num_vectors row-major elements. Scanner& nearest_multivector(const std::string& column, const void* query_data, size_t dimension, size_t num_vectors, diff --git a/src/scanner.rs b/src/scanner.rs index 71d2dd2..61717bf 100644 --- a/src/scanner.rs +++ b/src/scanner.rs @@ -148,6 +148,8 @@ struct NearestQuery { column: String, query: arrow_array::ArrayRef, k: u32, + lower_bound: Option, + upper_bound: Option, } /// The effective adaptive partition-search range shared by all three nprobes @@ -444,6 +446,7 @@ impl LanceScanner { } if let Some(n) = &self.nearest { scanner.nearest(&n.column, n.query.as_ref(), n.k as usize)?; + scanner.distance_range(n.lower_bound, n.upper_bound); if let Some(minimum_nprobes) = self.nprobes.minimum { scanner.minimum_nprobes(minimum_nprobes as usize); } @@ -2502,6 +2505,68 @@ macro_rules! scanner_set_u32 { scanner_set_u32!(lance_scanner_set_refine_factor, refine_factor); scanner_set_u32!(lance_scanner_set_ef, ef); +/// Set inclusive lower and exclusive upper distance bounds on a single-vector query. +/// +/// NULL bounds are unbounded. Values are copied and must be finite, with lower < upper +/// when both are present. Call after nearest and before scanning; k still caps results. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn lance_scanner_set_distance_range( + scanner: *mut LanceScanner, + lower_bound: *const f32, + upper_bound: *const f32, +) -> i32 { + scanner_poison_check!(scanner, -1); + scanner_ffi_try!(scanner, unsafe { + scanner_set_distance_range_inner(scanner, lower_bound, upper_bound) + }) +} + +unsafe fn scanner_set_distance_range_inner( + scanner: *mut LanceScanner, + lower_bound: *const f32, + upper_bound: *const f32, +) -> Result { + let invalid = |message: String| lance_core::Error::invalid_input_source(message.into()); + if scanner.is_null() { + return Err(invalid("scanner is NULL".into())); + } + let scanner = unsafe { &mut *scanner }; + scanner.ensure_scan_not_started("distance_range")?; + let nearest = scanner + .nearest + .as_mut() + .ok_or_else(|| invalid("distance_range requires nearest() to be configured".into()))?; + // Multi-vector scores aggregate subvectors; filtering each subvector's + // candidates would not enforce the requested range on the final score. + if matches!( + nearest.query.data_type(), + arrow_schema::DataType::FixedSizeList(_, _) + ) { + return Err(invalid( + "distance_range does not support multi-vector queries".into(), + )); + } + let lower_bound = unsafe { lower_bound.as_ref() }.copied(); + let upper_bound = unsafe { upper_bound.as_ref() }.copied(); + for (name, bound) in [("lower_bound", lower_bound), ("upper_bound", upper_bound)] { + if let Some(value) = bound + && !value.is_finite() + { + return Err(invalid(format!("{name} must be finite, got {value}"))); + } + } + if let (Some(lower), Some(upper)) = (lower_bound, upper_bound) + && lower >= upper + { + return Err(invalid(format!( + "lower_bound ({lower}) must be less than upper_bound ({upper})" + ))); + } + nearest.lower_bound = lower_bound; + nearest.upper_bound = upper_bound; + Ok(0) +} + /// Set both vector-index partition-search bounds to the same value. /// /// This replaces any values previously configured through @@ -2851,6 +2916,8 @@ unsafe fn scanner_nearest_inner( column: column_str.to_string(), query, k, + lower_bound: None, + upper_bound: None, }); Ok(0) } @@ -3005,6 +3072,8 @@ unsafe fn nearest_multivector_inner( column: column.to_string(), query: Arc::new(query), k, + lower_bound: None, + upper_bound: None, }); Ok(0) } diff --git a/tests/c_api_test.rs b/tests/c_api_test.rs index 1a25e34..34ff063 100644 --- a/tests/c_api_test.rs +++ b/tests/c_api_test.rs @@ -7306,6 +7306,274 @@ fn test_create_index_replace_false_conflicts() { // Vector search (k-NN) tests (Phase 2) // --------------------------------------------------------------------------- +fn range_scanner(dataset: *const LanceDataset, k: u32) -> *mut LanceScanner { + let scanner = unsafe { lance_scanner_new(dataset, ptr::null(), ptr::null()) }; + assert!(!scanner.is_null()); + let query: Vec = (0..8).map(|i| 31.0 + i as f32 / 8.0).collect(); + assert_eq!( + unsafe { + lance_scanner_nearest( + scanner, + c_str("embedding").as_ptr(), + query.as_ptr().cast(), + query.len(), + LanceDataType::Float32 as i32, + k, + ) + }, + 0 + ); + scanner +} + +fn range_results(scanner: *mut LanceScanner) -> Vec<(i32, f32)> { + let mut stream = FFI_ArrowArrayStream::empty(); + assert_eq!( + unsafe { lance_scanner_to_arrow_stream(scanner, &mut stream) }, + 0, + "{}", + take_last_error_message() + ); + let reader = unsafe { ArrowArrayStreamReader::from_raw(&mut stream) }.unwrap(); + reader + .flat_map(|batch| { + let batch = batch.unwrap(); + let ids = batch + .column_by_name("id") + .unwrap() + .as_any() + .downcast_ref::() + .unwrap(); + let distances = batch + .column_by_name("_distance") + .unwrap() + .as_any() + .downcast_ref::() + .unwrap(); + (0..batch.num_rows()) + .map(|i| (ids.value(i), distances.value(i))) + .collect::>() + }) + .collect() +} + +#[test] +fn test_scanner_distance_range_flat_and_indexed_boundaries() { + let (_tmp, uri) = create_multi_fragment_vector_dataset(2, 32, 8, false); + let dataset = unsafe { lance_dataset_open(c_str(&uri).as_ptr(), ptr::null(), 0) }; + assert!(!dataset.is_null()); + let params = LanceVectorIndexParams { + index_type: LanceVectorIndexType::IvfFlat, + metric: LanceMetricType::L2, + num_partitions: 2, + num_sub_vectors: 0, + num_bits: 0, + max_iterations: 0, + hnsw_m: 0, + hnsw_ef_construction: 0, + sample_rate: 0, + }; + assert_eq!( + unsafe { + lance_dataset_create_vector_index( + dataset, + c_str("embedding").as_ptr(), + ptr::null(), + ¶ms, + false, + ) + }, + 0, + "{}", + take_last_error_message() + ); + for use_index in [false, true] { + for (lower, upper, k, expected) in [ + (None, Some(32.0), 64, vec![30, 31, 32]), + (Some(8.0), Some(72.0), 64, vec![29, 30, 32, 33]), + (Some(8.0), Some(72.0), 2, vec![30, 32]), + (Some(32.0), Some(72.0), 64, vec![29, 33]), + (Some(7688.0), None, 64, vec![0, 62, 63]), + (None, Some(0.0), 64, vec![]), + ] { + let scanner = range_scanner(dataset, k); + assert_eq!( + unsafe { lance_scanner_set_use_index(scanner, use_index) }, + 0 + ); + assert_eq!(unsafe { lance_scanner_set_nprobes(scanner, 2) }, 0); + assert_eq!( + unsafe { + lance_scanner_set_distance_range( + scanner, + lower.as_ref().map_or(ptr::null(), |v| v), + upper.as_ref().map_or(ptr::null(), |v| v), + ) + }, + 0 + ); + let result = range_results(scanner); + assert!(result.windows(2).all(|w| w[0].1 <= w[1].1)); + assert!( + result + .iter() + .all(|(_, d)| lower.is_none_or(|l| *d >= l) && upper.is_none_or(|u| *d < u)) + ); + let mut ids: Vec<_> = result.into_iter().map(|(id, _)| id).collect(); + ids.sort_unstable(); + assert_eq!( + ids, expected, + "use_index={use_index}, lower={lower:?}, upper={upper:?}, k={k}" + ); + unsafe { lance_scanner_close(scanner) }; + } + } + unsafe { lance_dataset_close(dataset) }; +} + +#[test] +fn test_scanner_distance_range_validation_is_atomic_and_copies_bounds() { + let (_tmp, uri) = create_multi_fragment_vector_dataset(2, 32, 8, false); + let dataset = unsafe { lance_dataset_open(c_str(&uri).as_ptr(), ptr::null(), 0) }; + assert_eq!( + unsafe { lance_scanner_set_distance_range(ptr::null_mut(), ptr::null(), ptr::null()) }, + -1 + ); + let plain = unsafe { lance_scanner_new(dataset, ptr::null(), ptr::null()) }; + assert_eq!( + unsafe { lance_scanner_set_distance_range(plain, ptr::null(), &32.0) }, + -1 + ); + assert_eq!(lance_last_error_code(), LanceErrorCode::InvalidArgument); + unsafe { lance_scanner_close(plain) }; + let scanner = range_scanner(dataset, 64); + let mut upper = 32.0; + assert_eq!( + unsafe { lance_scanner_set_distance_range(scanner, ptr::null(), &upper) }, + 0 + ); + upper = 0.0; + assert_eq!(upper, 0.0); + for (lower, upper) in [ + (f32::NAN, 32.0), + (0.0, f32::NAN), + (f32::NEG_INFINITY, 32.0), + (0.0, f32::INFINITY), + (32.0, 8.0), + (32.0, 32.0), + ] { + assert_eq!( + unsafe { lance_scanner_set_distance_range(scanner, &lower, &upper) }, + -1 + ); + assert_eq!(lance_last_error_code(), LanceErrorCode::InvalidArgument); + } + let mut ids: Vec<_> = range_results(scanner) + .into_iter() + .map(|(id, _)| id) + .collect(); + ids.sort_unstable(); + assert_eq!(ids, vec![30, 31, 32]); + assert_eq!( + unsafe { lance_scanner_set_distance_range(scanner, ptr::null(), ptr::null()) }, + -1 + ); + assert_eq!(lance_last_error_code(), LanceErrorCode::InvalidArgument); + unsafe { + lance_scanner_close(scanner); + lance_dataset_close(dataset); + } +} + +#[test] +fn test_scanner_distance_range_can_clear_and_nearest_replacement_resets_bounds() { + let (_tmp, uri) = create_multi_fragment_vector_dataset(2, 32, 8, false); + let dataset = unsafe { lance_dataset_open(c_str(&uri).as_ptr(), ptr::null(), 0) }; + for replace_nearest in [false, true] { + let scanner = range_scanner(dataset, 5); + assert_eq!( + unsafe { lance_scanner_set_distance_range(scanner, ptr::null(), &0.0) }, + 0 + ); + if replace_nearest { + let query = [0.0f32; 8]; + assert_eq!( + unsafe { + lance_scanner_nearest( + scanner, + c_str("embedding").as_ptr(), + query.as_ptr().cast(), + 8, + LanceDataType::Float32 as i32, + 5, + ) + }, + 0 + ); + } else { + assert_eq!( + unsafe { lance_scanner_set_distance_range(scanner, ptr::null(), ptr::null()) }, + 0 + ); + } + assert_eq!(range_results(scanner).len(), 5); + unsafe { lance_scanner_close(scanner) }; + } + unsafe { lance_dataset_close(dataset) }; +} + +#[test] +fn test_scanner_distance_range_combines_with_prefilter_and_result_window() { + let (_tmp, uri) = create_multi_fragment_vector_dataset(2, 32, 8, false); + unsafe { + let dataset = lance_dataset_open(c_str(&uri).as_ptr(), ptr::null(), 0); + let scanner = range_scanner(dataset, 64); + assert_eq!(lance_scanner_set_use_index(scanner, false), 0); + assert_eq!(lance_scanner_set_prefilter(scanner, true), 0); + assert_eq!( + lance_scanner_additional_sql_filter(scanner, c_str("id >= 31").as_ptr()), + 0 + ); + assert_eq!(lance_scanner_set_distance_range(scanner, &8.0, &72.0), 0); + assert_eq!(lance_scanner_set_offset(scanner, 1), 0); + assert_eq!(lance_scanner_set_limit(scanner, 1), 0); + assert_eq!(range_results(scanner), vec![(33, 32.0)]); + lance_scanner_close(scanner); + lance_dataset_close(dataset); + } +} + +#[test] +fn test_scanner_distance_range_accepts_negative_dot_distances() { + let (_tmp, uri) = create_multi_fragment_vector_dataset(2, 32, 8, false); + unsafe { + let dataset = lance_dataset_open(c_str(&uri).as_ptr(), ptr::null(), 0); + let scanner = lance_scanner_new(dataset, ptr::null(), ptr::null()); + let query = [1.0f32; 8]; + assert_eq!( + lance_scanner_nearest( + scanner, + c_str("embedding").as_ptr(), + query.as_ptr().cast(), + 8, + LanceDataType::Float32 as i32, + 64 + ), + 0 + ); + assert_eq!( + lance_scanner_set_metric(scanner, LanceMetricType::Dot as i32), + 0 + ); + assert_eq!(lance_scanner_set_use_index(scanner, false), 0); + // Dot distance is 1 - dot(query, row), which can be negative. + assert_eq!(lance_scanner_set_distance_range(scanner, &-18.5, &-2.5), 0); + assert_eq!(range_results(scanner), vec![(2, -18.5), (1, -10.5)]); + lance_scanner_close(scanner); + lance_dataset_close(dataset); + } +} + #[test] fn test_scanner_nearest_brute_force() { let (_tmp, uri) = create_vector_dataset(64, 8); diff --git a/tests/cpp/test_c_api.c b/tests/cpp/test_c_api.c index 5df9cde..4610ced 100644 --- a/tests/cpp/test_c_api.c +++ b/tests/cpp/test_c_api.c @@ -1310,6 +1310,44 @@ static void test_delete(const char *write_uri) { printf("deleted=%llu... OK\n", (unsigned long long)deleted); } +static void test_distance_range(const char *uri) { + printf(" test_distance_range... "); + LanceDataset *ds = lance_dataset_open(uri, NULL, 0); + ASSERT(ds != NULL, "open failed"); + float query[8]; + for (int i = 0; i < 8; ++i) query[i] = 0.1f + (float)i; + const float lower = 0.02f, upper = 0.125f; + const int64_t expected[] = {2, 1, 20}; + for (int mode = 0; mode < 3; ++mode) { + LanceScanner *scanner = lance_scanner_new(ds, NULL, NULL); + ASSERT(scanner != NULL, "create scanner failed"); + ASSERT(lance_scanner_nearest(scanner, "embedding", query, 8, + LANCE_DTYPE_FLOAT32, 20) == 0, "nearest failed"); + ASSERT(lance_scanner_set_use_index(scanner, false) == 0, "use_index failed"); + ASSERT(lance_scanner_set_distance_range(scanner, mode == 1 ? &lower : NULL, + &upper) == 0, "distance_range failed"); + if (mode == 2) { + ASSERT(lance_scanner_set_distance_range(scanner, NULL, NULL) == 0, + "clearing distance_range failed"); + } + struct ArrowArrayStream stream = {0}; + ASSERT(lance_scanner_to_arrow_stream(scanner, &stream) == 0, "stream failed"); + int64_t rows = 0; + while (1) { + struct ArrowArray array = {0}; + ASSERT(stream.get_next(&stream, &array) == 0, "get_next failed"); + if (!array.release) break; + rows += array.length; + array.release(&array); + } + ASSERT(rows == expected[mode], "distance_range row count mismatch"); + stream.release(&stream); + lance_scanner_close(scanner); + } + lance_dataset_close(ds); + printf("OK\n"); +} + int main(int argc, char **argv) { if (argc < 4) { fprintf(stderr, "Usage: %s \n", argv[0]); @@ -1324,6 +1362,7 @@ int main(int argc, char **argv) { test_open_and_metadata(uri); test_shared_session(uri); test_scan(uri); + test_distance_range(uri); test_scan_with_limit(uri); test_scanner_blob_handling(blob_uri); test_take_blobs(blob_uri); diff --git a/tests/cpp/test_cpp_api.cpp b/tests/cpp/test_cpp_api.cpp index 0332d1a..e3cde45 100644 --- a/tests/cpp/test_cpp_api.cpp +++ b/tests/cpp/test_cpp_api.cpp @@ -1198,6 +1198,42 @@ static void test_delete_rows(const std::string& dst_uri) { PASS(); } +static void test_distance_range(const std::string& uri) { + TEST(test_distance_range); + auto ds = lance::Dataset::open(uri); + float query[8]; + for (int i = 0; i < 8; ++i) query[i] = 0.1f + static_cast(i); + const int64_t expected[] = {2, 1, 20}; + for (int mode = 0; mode < 3; ++mode) { + auto scanner = ds.scan(); + scanner.nearest("embedding", query, 8, 20).use_index(false); + if (mode == 1) scanner.distance_range(0.02f, 0.125f); + else scanner.distance_range(std::nullopt, 0.125f); + if (mode == 2) scanner.distance_range(); + bool rejected = false; + try { + scanner.distance_range(1.0f, 0.0f); + } catch (const lance::Error& e) { + rejected = true; + assert(e.code == LANCE_ERR_INVALID_ARGUMENT); + } + assert(rejected); + ArrowArrayStream stream{}; + scanner.to_arrow_stream(&stream); + int64_t rows = 0; + while (true) { + ArrowArray array{}; + assert(stream.get_next(&stream, &array) == 0); + if (!array.release) break; + rows += array.length; + array.release(&array); + } + assert(rows == expected[mode]); + stream.release(&stream); + } + PASS(); +} + int main(int argc, char** argv) { if (argc < 4) { fprintf(stderr, "Usage: %s \n", argv[0]); @@ -1213,6 +1249,7 @@ int main(int argc, char** argv) { test_shared_session(uri); test_dataset_schema(uri); test_scanner_fluent(uri); + test_distance_range(uri); test_scanner_async_stream_ownership(uri); test_scanner_blob_handling(blob_uri); test_take_blobs(blob_uri); diff --git a/tests/multivector_test.rs b/tests/multivector_test.rs index f02e850..4223502 100644 --- a/tests/multivector_test.rs +++ b/tests/multivector_test.rs @@ -963,3 +963,32 @@ fn strict_batches_apply_after_multivector_offset_and_limit() { lance_dataset_close(ds); } } + +#[test] +fn multivector_distance_range_is_rejected_instead_of_filtering_subvectors() { + let (_dir, uri) = fixture(); + unsafe { + let ds = lance_dataset_open(uri.as_ptr(), ptr::null(), 0); + let scan = lance_scanner_new(ds, ptr::null(), ptr::null()); + let query = [1.0f32, 0.0]; + assert_eq!( + lance_scanner_nearest_multivector( + scan, + CString::new("vectors").unwrap().as_ptr(), + query.as_ptr().cast(), + 2, + 1, + LanceDataType::Float32 as i32, + 3 + ), + 0 + ); + assert_eq!( + lance_scanner_set_distance_range(scan, ptr::null(), &1.0), + -1 + ); + assert_eq!(lance_last_error_code(), LanceErrorCode::InvalidArgument); + lance_scanner_close(scan); + lance_dataset_close(ds); + } +} From 3dc61c41c430df4c409d9e4ec3e7305f0c87ddd1 Mon Sep 17 00:00:00 2001 From: Gabriel Date: Wed, 30 Sep 2026 17:29:43 +0800 Subject: [PATCH 2/2] test: strengthen distance range binding and validation coverage --- tests/c_api_test.rs | 24 +++++++++---- tests/cpp/test_c_api.c | 73 ++++++++++++++++++++++++++++++-------- tests/cpp/test_cpp_api.cpp | 67 ++++++++++++++++++++++++++-------- 3 files changed, 127 insertions(+), 37 deletions(-) diff --git a/tests/c_api_test.rs b/tests/c_api_test.rs index 34ff063..2a00f5e 100644 --- a/tests/c_api_test.rs +++ b/tests/c_api_test.rs @@ -7454,16 +7454,26 @@ fn test_scanner_distance_range_validation_is_atomic_and_copies_bounds() { ); upper = 0.0; assert_eq!(upper, 0.0); + // Leave the other side unbounded so ordering validation cannot mask a + // regression in the finite-value check, especially for +inf lower/-inf upper. for (lower, upper) in [ - (f32::NAN, 32.0), - (0.0, f32::NAN), - (f32::NEG_INFINITY, 32.0), - (0.0, f32::INFINITY), - (32.0, 8.0), - (32.0, 32.0), + (Some(f32::NAN), None), + (None, Some(f32::NAN)), + (Some(f32::NEG_INFINITY), None), + (Some(f32::INFINITY), None), + (None, Some(f32::NEG_INFINITY)), + (None, Some(f32::INFINITY)), + (Some(32.0), Some(8.0)), + (Some(32.0), Some(32.0)), ] { assert_eq!( - unsafe { lance_scanner_set_distance_range(scanner, &lower, &upper) }, + unsafe { + lance_scanner_set_distance_range( + scanner, + lower.as_ref().map_or(ptr::null(), |v| v), + upper.as_ref().map_or(ptr::null(), |v| v), + ) + }, -1 ); assert_eq!(lance_last_error_code(), LanceErrorCode::InvalidArgument); diff --git a/tests/cpp/test_c_api.c b/tests/cpp/test_c_api.c index 4610ced..ce59a40 100644 --- a/tests/cpp/test_c_api.c +++ b/tests/cpp/test_c_api.c @@ -1310,38 +1310,81 @@ static void test_delete(const char *write_uri) { printf("deleted=%llu... OK\n", (unsigned long long)deleted); } +static void check_distance_range_results(struct ArrowArrayStream *stream, + const float *lower, const float *upper, + uint32_t expected_ids, int64_t expected_count) { + struct ArrowSchema schema = {0}; + ASSERT(stream->get_schema(stream, &schema) == 0, "get_schema failed"); + int id_column = -1, distance_column = -1; + for (int64_t i = 0; i < schema.n_children; ++i) { + if (strcmp(schema.children[i]->name, "id") == 0) id_column = (int)i; + if (strcmp(schema.children[i]->name, "_distance") == 0) distance_column = (int)i; + } + ASSERT(id_column >= 0 && distance_column >= 0, "missing range result columns"); + ASSERT(strcmp(schema.children[id_column]->format, "i") == 0 && + strcmp(schema.children[distance_column]->format, "f") == 0, "unexpected range result types"); + schema.release(&schema); + int64_t rows = 0; + uint32_t seen = 0; + while (1) { + struct ArrowArray array = {0}; + ASSERT(stream->get_next(stream, &array) == 0, "get_next failed"); + if (!array.release) break; + const struct ArrowArray *ids = array.children[id_column]; + const struct ArrowArray *distances = array.children[distance_column]; + const int32_t *id_values = (const int32_t *)ids->buffers[1]; + const float *distance_values = (const float *)distances->buffers[1]; + for (int64_t i = 0; i < array.length; ++i) { + int32_t id = id_values[ids->offset + i]; + float distance = distance_values[distances->offset + i]; + ASSERT(id >= 1 && id <= 20, "unexpected range result id"); + uint32_t bit = 1u << (id - 1); + ASSERT((seen & bit) == 0, "duplicate range result id"); + seen |= bit; + ASSERT(distance >= 0.0f && (!lower || distance >= *lower) && (!upper || distance < *upper), "distance outside half-open range"); + } + rows += array.length; + array.release(&array); + } + ASSERT(rows == expected_count && seen == expected_ids, "range result ids/count mismatch"); + stream->release(stream); +} + static void test_distance_range(const char *uri) { printf(" test_distance_range... "); LanceDataset *ds = lance_dataset_open(uri, NULL, 0); ASSERT(ds != NULL, "open failed"); float query[8]; for (int i = 0; i < 8; ++i) query[i] = 0.1f + (float)i; - const float lower = 0.02f, upper = 0.125f; - const int64_t expected[] = {2, 1, 20}; - for (int mode = 0; mode < 3; ++mode) { + const float positive_lower = 0.02f, positive_upper = 0.125f, zero = 0.0f; + const int64_t expected[] = {2, 1, 20, 20, 0}; + const uint32_t expected_ids[] = {3u, 2u, (1u << 20) - 1, (1u << 20) - 1, 0u}; + for (int mode = 0; mode < 5; ++mode) { LanceScanner *scanner = lance_scanner_new(ds, NULL, NULL); ASSERT(scanner != NULL, "create scanner failed"); ASSERT(lance_scanner_nearest(scanner, "embedding", query, 8, LANCE_DTYPE_FLOAT32, 20) == 0, "nearest failed"); ASSERT(lance_scanner_set_use_index(scanner, false) == 0, "use_index failed"); - ASSERT(lance_scanner_set_distance_range(scanner, mode == 1 ? &lower : NULL, - &upper) == 0, "distance_range failed"); + const float *lower = mode == 1 ? &positive_lower : NULL; + const float *upper = &positive_upper; + /* Row 1 lies on zero: distinguish inclusive lower and exclusive upper bounds. */ + if (mode == 3) { lower = &zero; upper = NULL; } + if (mode == 4) upper = &zero; + ASSERT(lance_scanner_set_distance_range(scanner, lower, upper) == 0, + "distance_range failed"); if (mode == 2) { ASSERT(lance_scanner_set_distance_range(scanner, NULL, NULL) == 0, "clearing distance_range failed"); + lower = NULL; + upper = NULL; } + ASSERT(lance_scanner_set_distance_range(scanner, &positive_upper, &positive_lower) == -1, + "reversed distance range must fail"); + ASSERT(lance_last_error_code() == LANCE_ERR_INVALID_ARGUMENT, + "expected INVALID_ARGUMENT"); struct ArrowArrayStream stream = {0}; ASSERT(lance_scanner_to_arrow_stream(scanner, &stream) == 0, "stream failed"); - int64_t rows = 0; - while (1) { - struct ArrowArray array = {0}; - ASSERT(stream.get_next(&stream, &array) == 0, "get_next failed"); - if (!array.release) break; - rows += array.length; - array.release(&array); - } - ASSERT(rows == expected[mode], "distance_range row count mismatch"); - stream.release(&stream); + check_distance_range_results(&stream, lower, upper, expected_ids[mode], expected[mode]); lance_scanner_close(scanner); } lance_dataset_close(ds); diff --git a/tests/cpp/test_cpp_api.cpp b/tests/cpp/test_cpp_api.cpp index e3cde45..e82582d 100644 --- a/tests/cpp/test_cpp_api.cpp +++ b/tests/cpp/test_cpp_api.cpp @@ -1198,18 +1198,63 @@ static void test_delete_rows(const std::string& dst_uri) { PASS(); } +static void check_distance_range_results(struct ArrowArrayStream *stream, + const float *lower, const float *upper, + uint32_t expected_ids, int64_t expected_count) { + struct ArrowSchema schema = {}; + assert(stream->get_schema(stream, &schema) == 0); + int id_column = -1, distance_column = -1; + for (int64_t i = 0; i < schema.n_children; ++i) { + if (strcmp(schema.children[i]->name, "id") == 0) id_column = (int)i; + if (strcmp(schema.children[i]->name, "_distance") == 0) distance_column = (int)i; + } + assert(id_column >= 0 && distance_column >= 0); + assert(strcmp(schema.children[id_column]->format, "i") == 0 && + strcmp(schema.children[distance_column]->format, "f") == 0); + schema.release(&schema); + int64_t rows = 0; + uint32_t seen = 0; + while (1) { + struct ArrowArray array = {}; + assert(stream->get_next(stream, &array) == 0); + if (!array.release) break; + const struct ArrowArray *ids = array.children[id_column]; + const struct ArrowArray *distances = array.children[distance_column]; + const int32_t *id_values = (const int32_t *)ids->buffers[1]; + const float *distance_values = (const float *)distances->buffers[1]; + for (int64_t i = 0; i < array.length; ++i) { + int32_t id = id_values[ids->offset + i]; + float distance = distance_values[distances->offset + i]; + assert(id >= 1 && id <= 20); + uint32_t bit = 1u << (id - 1); + assert((seen & bit) == 0); + seen |= bit; + assert(distance >= 0.0f && (!lower || distance >= *lower) && (!upper || distance < *upper)); + } + rows += array.length; + array.release(&array); + } + assert(rows == expected_count && seen == expected_ids); + stream->release(stream); +} + static void test_distance_range(const std::string& uri) { TEST(test_distance_range); auto ds = lance::Dataset::open(uri); float query[8]; for (int i = 0; i < 8; ++i) query[i] = 0.1f + static_cast(i); - const int64_t expected[] = {2, 1, 20}; - for (int mode = 0; mode < 3; ++mode) { + const int64_t expected[] = {2, 1, 20, 20, 0}; + const uint32_t expected_ids[] = {3u, 2u, (1u << 20) - 1, (1u << 20) - 1, 0u}; + for (int mode = 0; mode < 5; ++mode) { auto scanner = ds.scan(); scanner.nearest("embedding", query, 8, 20).use_index(false); - if (mode == 1) scanner.distance_range(0.02f, 0.125f); - else scanner.distance_range(std::nullopt, 0.125f); - if (mode == 2) scanner.distance_range(); + std::optional lower = mode == 1 ? std::optional(0.02f) : std::nullopt; + std::optional upper = 0.125f; + // Row 1 has exactly zero distance: modes 3/4 distinguish >= from > and < from <=. + if (mode == 3) { lower = 0.0f; upper = std::nullopt; } + if (mode == 4) upper = 0.0f; + scanner.distance_range(lower, upper); + if (mode == 2) { scanner.distance_range(); lower.reset(); upper.reset(); } bool rejected = false; try { scanner.distance_range(1.0f, 0.0f); @@ -1220,16 +1265,8 @@ static void test_distance_range(const std::string& uri) { assert(rejected); ArrowArrayStream stream{}; scanner.to_arrow_stream(&stream); - int64_t rows = 0; - while (true) { - ArrowArray array{}; - assert(stream.get_next(&stream, &array) == 0); - if (!array.release) break; - rows += array.length; - array.release(&array); - } - assert(rows == expected[mode]); - stream.release(&stream); + check_distance_range_results(&stream, lower ? &*lower : nullptr, + upper ? &*upper : nullptr, expected_ids[mode], expected[mode]); } PASS(); }