From 3e961a167149570f785f15a793c00bbeda05ccd8 Mon Sep 17 00:00:00 2001 From: Junrui Lee Date: Mon, 21 Sep 2026 11:50:12 +0800 Subject: [PATCH 1/5] feat: add range search across language bindings --- .github/workflows/ci.yml | 14 + Cargo.lock | 1 + README.md | 2 + c/range_test_support.h | 286 +++++++ c/test_vindex.c | 296 +++++++ cpp/test_vindex.cpp | 232 ++++++ docs/api.html | 2 +- docs/range-search.html | 81 +- ffi/Cargo.toml | 3 + ffi/examples/range_search_fixture.rs | 206 +++++ ffi/src/lib.rs | 6 + ffi/src/range.rs | 425 ++++++++++ ffi/src/range_tests.rs | 483 ++++++++++++ include/paimon_vindex.hpp | 161 ++++ .../index/vector/VectorDistanceBand.java | 104 +++ .../index/vector/VectorIndexNative.java | 21 + .../index/vector/VectorIndexReader.java | 91 +++ .../index/vector/VectorRangeSearchParams.java | 43 + .../index/vector/VectorRangeSearchResult.java | 145 ++++ .../index/vector/VectorIndexJavaApiTest.java | 1 + .../VectorIndexNativeValidationTest.java | 1 + .../vector/VectorIndexRangeOracleTest.java | 190 +++++ .../vector/VectorIndexRangeSearchTest.java | 589 ++++++++++++++ jni/src/lib.rs | 1 + jni/src/range.rs | 403 ++++++++++ python/README.md | 102 +++ python/paimon_vindex/__init__.py | 308 +++++++- python/paimon_vindex/_ffi.py | 91 +++ python/tests/test_range_search.py | 739 ++++++++++++++++++ 29 files changed, 5022 insertions(+), 5 deletions(-) create mode 100644 c/range_test_support.h create mode 100644 ffi/examples/range_search_fixture.rs create mode 100644 ffi/src/range.rs create mode 100644 ffi/src/range_tests.rs create mode 100644 java/src/main/java/org/apache/paimon/index/vector/VectorDistanceBand.java create mode 100644 java/src/main/java/org/apache/paimon/index/vector/VectorRangeSearchParams.java create mode 100644 java/src/main/java/org/apache/paimon/index/vector/VectorRangeSearchResult.java create mode 100644 java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeOracleTest.java create mode 100644 java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeSearchTest.java create mode 100644 jni/src/range.rs create mode 100644 python/README.md create mode 100644 python/tests/test_range_search.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index c2063765..0870f10a 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -114,6 +114,8 @@ jobs: c-test: runs-on: ubuntu-latest + env: + PVI_RANGE_FIXTURES: ${{ github.workspace }}/target/range-fixtures steps: - uses: actions/checkout@v4 @@ -133,6 +135,7 @@ jobs: - name: Build C test run: | + cargo run --release -p paimon-vindex-ffi --example range_search_fixture -- "$PVI_RANGE_FIXTURES" cmake -S c -B c/build \ -DPAIMON_VINDEX_FFI_LIB=${{ github.workspace }}/target/release/libpaimon_vindex_ffi.so cmake --build c/build @@ -144,6 +147,8 @@ jobs: cpp-test: runs-on: ubuntu-latest + env: + PVI_RANGE_FIXTURES: ${{ github.workspace }}/target/range-fixtures steps: - uses: actions/checkout@v4 @@ -163,6 +168,7 @@ jobs: - name: Build C++ test run: | + cargo run --release -p paimon-vindex-ffi --example range_search_fixture -- "$PVI_RANGE_FIXTURES" cmake -S cpp -B cpp/build \ -DPAIMON_VINDEX_FFI_LIB=${{ github.workspace }}/target/release/libpaimon_vindex_ffi.so cmake --build cpp/build @@ -174,6 +180,8 @@ jobs: jni-build: runs-on: ubuntu-latest + env: + PVI_RANGE_FIXTURES: ${{ github.workspace }}/target/range-fixtures steps: - uses: actions/checkout@v4 @@ -197,6 +205,9 @@ jobs: - name: Build JNI library run: cargo build -p paimon-vindex-jni --release + - name: Generate core range-search oracle + run: cargo run --release -p paimon-vindex-ffi --example range_search_fixture -- "$PVI_RANGE_FIXTURES" + - name: Test Java API run: > mvn -f java/pom.xml test @@ -230,6 +241,8 @@ jobs: python-build: runs-on: ubuntu-latest + env: + PVI_RANGE_FIXTURES: ${{ github.workspace }}/target/range-fixtures steps: - uses: actions/checkout@v4 @@ -255,6 +268,7 @@ jobs: - name: Test Python API working-directory: python run: | + cargo run --manifest-path ../Cargo.toml --release -p paimon-vindex-ffi --example range_search_fixture -- "$PVI_RANGE_FIXTURES" pip install -e ".[test]" pytest -v env: diff --git a/Cargo.lock b/Cargo.lock index eaac3198..9817e08b 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -677,6 +677,7 @@ version = "0.6.0" dependencies = [ "cbindgen", "paimon-vindex-core", + "roaring", ] [[package]] diff --git a/README.md b/README.md index 953c53e3..a1e89510 100644 --- a/README.md +++ b/README.md @@ -45,6 +45,8 @@ retention is bounded and the shared cold cache is sharded for concurrent hits. DiskANN and configure build, search, local-SSD, and object-store parameters. - [API and language bindings](docs/api.html): lifecycle, query parameters, warm-up, Rust, C, C++, Java, Python, and metadata filter pushdown. +- [Distance range search](docs/range-search.html#bindings): distance bands, + variable-length results, ownership, and core-backed cross-language verification. - [Development and benchmarks](docs/development.html): workspace layout, build and test commands, ANN benchmarks, and storage compatibility checks. - [Storage format specification](core/STORAGE_FORMAT.md): normative v1 binary diff --git a/c/range_test_support.h b/c/range_test_support.h new file mode 100644 index 00000000..86522198 --- /dev/null +++ b/c/range_test_support.h @@ -0,0 +1,286 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +#ifndef PAIMON_VINDEX_RANGE_TEST_SUPPORT_H +#define PAIMON_VINDEX_RANGE_TEST_SUPPORT_H + +#include +#include +#include +#include +#include +#include + +enum { + RANGE_DIMENSION = 8, + RANGE_NLIST = 4, + RANGE_VECTOR_COUNT = 512, + RANGE_QUERY_COUNT = 3 +}; + +static const uint8_t range_filter[] = { + 1, 0, 0, 0, 0, 0, 0, 0, + 0, 1, 0, 0, + 58, 48, 0, 0, 1, 0, 0, 0, + 0, 0, 1, 0, 16, 0, 0, 0, + 1, 0, 3, 0 +}; +static const uint8_t range_empty_filter[8] = {0}; + +static uint32_t range_float_bits(float value) { + uint32_t bits; + memcpy(&bits, &value, sizeof(bits)); + return bits; +} + +static float range_float_from_bits(uint32_t bits) { + float value; + memcpy(&value, &bits, sizeof(value)); + return value; +} + +static void range_fill_data(float *data, int64_t *labels, float *queries) { + for (size_t row = 0; row < RANGE_VECTOR_COUNT; ++row) { + labels[row] = (INT64_C(1) << 40) + (int64_t)row; + for (size_t dimension = 0; dimension < RANGE_DIMENSION; ++dimension) { + data[row * RANGE_DIMENSION + dimension] = + (float)((int)((row * 17 + dimension * 13 + row * dimension) % 101) - 50) / 16.0f; + } + } + for (size_t query = 0; query < RANGE_QUERY_COUNT; ++query) { + memcpy(queries + query * RANGE_DIMENSION, + data + query * 71 * RANGE_DIMENSION, + RANGE_DIMENSION * sizeof(float)); + } +} + +static PaimonVindexRangeSearchParams range_all_params(uint32_t metric) { + PaimonVindexRangeSearchParams params = {{0, 0, 0.0f, 0, 0.0f}, 0}; + params.band.metric = metric; + params.band.lower_kind = PAIMON_VINDEX_BOUND_UNBOUNDED; + params.band.upper_kind = PAIMON_VINDEX_BOUND_UNBOUNDED; + params.nprobe = RANGE_NLIST; + return params; +} + +static void range_assert_shape(const PaimonVindexRangeSearchResultView *view) { + ASSERT_TRUE(view->lims != NULL); + ASSERT_TRUE(view->lims[0] == 0); + ASSERT_TRUE(view->lims[view->query_count] == view->hit_count); + ASSERT_TRUE(view->hit_count == 0 || (view->labels != NULL && view->distances != NULL)); + ASSERT_TRUE(view->query_count == 0 || view->stats != NULL); + for (size_t query = 0; query < view->query_count; ++query) { + ASSERT_TRUE(view->lims[query] <= view->lims[query + 1]); + ASSERT_TRUE(view->stats[query].rows_committed == view->lims[query + 1] - view->lims[query]); + ASSERT_TRUE(view->stats[query].rows_scanned >= view->stats[query].rows_committed); + ASSERT_TRUE(view->stats[query].early_abandoned <= view->stats[query].rows_scanned); + ASSERT_TRUE(view->stats[query].lists_probed <= RANGE_NLIST); + } + for (size_t hit = 0; hit < view->hit_count; ++hit) { + ASSERT_TRUE(view->labels[hit] >= (INT64_C(1) << 40)); + ASSERT_TRUE(isfinite(view->distances[hit])); + } +} + +static void range_assert_stats_equal( + const PaimonVindexRangeSearchStats *actual, + const PaimonVindexRangeSearchStats *expected) { + ASSERT_TRUE(actual->lists_probed == expected->lists_probed); + ASSERT_TRUE(actual->rows_scanned == expected->rows_scanned); + ASSERT_TRUE(actual->rows_committed == expected->rows_committed); + ASSERT_TRUE(actual->early_abandoned == expected->early_abandoned); +} + +struct RangeFixture { + size_t dimension; + PaimonVindexRangeSearchParams params; + size_t query_count; + size_t filter_len; + size_t hit_count; + size_t list_reads; + float *queries; + uint8_t *filter; + size_t *lims; + int64_t *labels; + uint32_t *distance_bits; + PaimonVindexRangeSearchStats *stats; + uint8_t *index_data; + size_t index_len; +}; + +static void *range_fixture_allocate(size_t count, size_t element_size) { + ASSERT_TRUE(element_size != 0 && count <= SIZE_MAX / element_size); + void *allocation = calloc(count == 0 ? 1 : count, element_size); + ASSERT_TRUE(allocation != NULL); + return allocation; +} + +static FILE *range_fixture_open(const char *directory, const char *name, const char *suffix) { + char filename[4096]; + int written = snprintf(filename, sizeof(filename), "%s/%s%s", directory, name, suffix); + ASSERT_TRUE(written >= 0 && (size_t)written < sizeof(filename)); + FILE *file = fopen(filename, "rb"); + if (file == NULL) { + fprintf(stderr, "Cannot read range fixture %s\n", filename); + abort(); + } + return file; +} + +static void range_fixture_load( + const char *directory, const char *name, const char *index_name, + struct RangeFixture *fixture) { + memset(fixture, 0, sizeof(*fixture)); + FILE *file = range_fixture_open(directory, name, ".expected"); + uint32_t lower_bits; + uint32_t upper_bits; + ASSERT_TRUE(fscanf(file, "%zu %" SCNu32 " %zu %zu %" SCNu32 " %" SCNu32 + " %" SCNu32 " %" SCNu32 " %zu %zu %zu", + &fixture->dimension, &fixture->params.band.metric, + &fixture->query_count, &fixture->params.nprobe, + &fixture->params.band.lower_kind, &lower_bits, + &fixture->params.band.upper_kind, &upper_bits, + &fixture->filter_len, &fixture->hit_count, &fixture->list_reads) == 11); + fixture->params.band.lower = range_float_from_bits(lower_bits); + fixture->params.band.upper = range_float_from_bits(upper_bits); + ASSERT_TRUE(fixture->dimension != 0); + ASSERT_TRUE(fixture->query_count < SIZE_MAX); + ASSERT_TRUE(fixture->query_count <= SIZE_MAX / fixture->dimension); + size_t query_len = fixture->dimension * fixture->query_count; + fixture->queries = (float *)range_fixture_allocate(query_len, sizeof(float)); + fixture->filter = (uint8_t *)range_fixture_allocate(fixture->filter_len, sizeof(uint8_t)); + fixture->lims = (size_t *)range_fixture_allocate(fixture->query_count + 1, sizeof(size_t)); + fixture->labels = (int64_t *)range_fixture_allocate(fixture->hit_count, sizeof(int64_t)); + fixture->distance_bits = (uint32_t *)range_fixture_allocate(fixture->hit_count, sizeof(uint32_t)); + fixture->stats = (PaimonVindexRangeSearchStats *)range_fixture_allocate( + fixture->query_count, sizeof(PaimonVindexRangeSearchStats)); + for (size_t element = 0; element < query_len; ++element) { + uint32_t bits; + ASSERT_TRUE(fscanf(file, "%" SCNu32, &bits) == 1); + fixture->queries[element] = range_float_from_bits(bits); + } + for (size_t element = 0; element < fixture->filter_len; ++element) { + unsigned int byte; + ASSERT_TRUE(fscanf(file, "%u", &byte) == 1 && byte <= UINT8_MAX); + fixture->filter[element] = (uint8_t)byte; + } + for (size_t element = 0; element <= fixture->query_count; ++element) { + ASSERT_TRUE(fscanf(file, "%zu", &fixture->lims[element]) == 1); + } + for (size_t element = 0; element < fixture->hit_count; ++element) { + ASSERT_TRUE(fscanf(file, "%" SCNd64, &fixture->labels[element]) == 1); + } + for (size_t element = 0; element < fixture->hit_count; ++element) { + ASSERT_TRUE(fscanf(file, "%" SCNu32, &fixture->distance_bits[element]) == 1); + } + for (size_t query = 0; query < fixture->query_count; ++query) { + PaimonVindexRangeSearchStats *stats = &fixture->stats[query]; + ASSERT_TRUE(fscanf(file, "%zu %zu %zu %zu", &stats->lists_probed, + &stats->rows_scanned, &stats->rows_committed, + &stats->early_abandoned) == 4); + } + char trailing; + ASSERT_TRUE(fscanf(file, " %c", &trailing) == EOF); + ASSERT_TRUE(fclose(file) == 0); + file = range_fixture_open(directory, index_name, ""); + ASSERT_TRUE(fseek(file, 0, SEEK_END) == 0); + long index_len = ftell(file); + ASSERT_TRUE(index_len > 0); + fixture->index_len = (size_t)index_len; + fixture->index_data = (uint8_t *)range_fixture_allocate(fixture->index_len, 1); + ASSERT_TRUE(fseek(file, 0, SEEK_SET) == 0); + ASSERT_TRUE(fread(fixture->index_data, 1, fixture->index_len, file) == fixture->index_len); + ASSERT_TRUE(fclose(file) == 0); +} + +struct RangeFixtureHit { + int64_t label; + uint32_t distance_bits; +}; + +static int range_fixture_compare_hits(const void *left, const void *right) { + const struct RangeFixtureHit *left_hit = (const struct RangeFixtureHit *)left; + const struct RangeFixtureHit *right_hit = (const struct RangeFixtureHit *)right; + if (left_hit->label != right_hit->label) return left_hit->label < right_hit->label ? -1 : 1; + if (left_hit->distance_bits != right_hit->distance_bits) { + return left_hit->distance_bits < right_hit->distance_bits ? -1 : 1; + } + return 0; +} + +static void range_fixture_assert( + const struct RangeFixture *fixture, const PaimonVindexRangeSearchResultView *view) { + ASSERT_TRUE(view->query_count == fixture->query_count); + ASSERT_TRUE(view->hit_count == fixture->hit_count); + ASSERT_TRUE(view->list_reads == fixture->list_reads); + for (size_t query = 0; query <= fixture->query_count; ++query) { + ASSERT_TRUE(view->lims[query] == fixture->lims[query]); + } + struct RangeFixtureHit *actual = (struct RangeFixtureHit *)range_fixture_allocate( + fixture->hit_count, sizeof(struct RangeFixtureHit)); + struct RangeFixtureHit *expected = (struct RangeFixtureHit *)range_fixture_allocate( + fixture->hit_count, sizeof(struct RangeFixtureHit)); + for (size_t hit = 0; hit < fixture->hit_count; ++hit) { + actual[hit].label = view->labels[hit]; + actual[hit].distance_bits = range_float_bits(view->distances[hit]); + expected[hit].label = fixture->labels[hit]; + expected[hit].distance_bits = fixture->distance_bits[hit]; + } + for (size_t query = 0; query < fixture->query_count; ++query) { + size_t begin = fixture->lims[query]; + size_t count = fixture->lims[query + 1] - begin; + qsort(actual + begin, count, sizeof(struct RangeFixtureHit), range_fixture_compare_hits); + qsort(expected + begin, count, sizeof(struct RangeFixtureHit), range_fixture_compare_hits); + for (size_t hit = begin; hit < begin + count; ++hit) { + ASSERT_TRUE(range_fixture_compare_hits(&actual[hit], &expected[hit]) == 0); + } + range_assert_stats_equal(&view->stats[query], &fixture->stats[query]); + } + free(actual); + free(expected); +} + +static void range_fixture_run_all(void (*consume)(const struct RangeFixture *)) { + const char *directory = getenv("PVI_RANGE_FIXTURES"); + if (directory == NULL) return; + FILE *manifest = range_fixture_open(directory, "manifest.txt", ""); + char name[256]; + char index_name[256]; + size_t case_count = 0; + int fields; + while ((fields = fscanf(manifest, "%255s %255s", name, index_name)) == 2) { + struct RangeFixture fixture; + printf("ORACLE %s\n", name); + range_fixture_load(directory, name, index_name, &fixture); + consume(&fixture); + free(fixture.queries); + free(fixture.filter); + free(fixture.lims); + free(fixture.labels); + free(fixture.distance_bits); + free(fixture.stats); + free(fixture.index_data); + ++case_count; + } + ASSERT_TRUE(fields == EOF && case_count != 0); + ASSERT_TRUE(fclose(manifest) == 0); + printf("PASS range_core_oracle (%zu cases)\n", case_count); +} + +#endif diff --git a/c/test_vindex.c b/c/test_vindex.c index 474aba7d..f2d6a5d2 100644 --- a/c/test_vindex.c +++ b/c/test_vindex.c @@ -42,6 +42,8 @@ } \ } while (0) +#include "range_test_support.h" + struct MemBuffer { uint8_t *data; size_t len; @@ -391,6 +393,16 @@ static void run_roundtrip( } assert_id_in_cluster(batch_ids[0], 0); assert_id_in_cluster(batch_ids[1], 1); + int range_supported = -1; + ASSERT_TRUE(paimon_vindex_reader_supports_range_search(reader, &range_supported) == 0); + ASSERT_TRUE(range_supported == (expected_index_type != PAIMON_VINDEX_INDEX_TYPE_DISKANN)); + if (!range_supported) { + PaimonVindexRangeSearchResult *unsupported = (PaimonVindexRangeSearchResult *)(uintptr_t)1; + PaimonVindexRangeSearchParams range_params = range_all_params(PAIMON_VINDEX_METRIC_L2); + ASSERT_TRUE(paimon_vindex_reader_range_search( + reader, query, ROUNDTRIP_DIMENSION, range_params, &unsupported) != 0); + ASSERT_TRUE(unsupported == NULL); + } paimon_vindex_reader_free(reader); free(buf.data); free(data); @@ -551,11 +563,295 @@ static void test_extensible_search_params_defaults(void) { printf("PASS extensible_search_params_defaults\n"); } +static PaimonVindexRangeSearchResultView range_view(PaimonVindexRangeSearchResult *result) { + PaimonVindexRangeSearchResultView view = {0}; + if (paimon_vindex_range_search_result_view(result, &view) != 0) fail_ffi("range view"); + return view; +} + +static void consume_range_fixture(const struct RangeFixture *fixture) { + struct MemBuffer buffer = {0}; + buffer.data = fixture->index_data; + buffer.len = fixture->index_len; + PaimonVindexInputFile input = {.ctx = &buffer, .read_ranges_fn = mem_read_ranges}; + PaimonVindexReaderHandle *reader = paimon_vindex_reader_open(input); + ASSERT_TRUE(reader != NULL); + int supported = 0; + ASSERT_TRUE(paimon_vindex_reader_supports_range_search(reader, &supported) == 0 && supported == 1); + PaimonVindexRangeSearchResult *result = NULL; + size_t query_len = fixture->dimension * fixture->query_count; + int status; + if (fixture->query_count == 1) { + status = fixture->filter_len == 0 + ? paimon_vindex_reader_range_search(reader, fixture->queries, query_len, fixture->params, &result) + : paimon_vindex_reader_range_search_with_roaring_filter( + reader, fixture->queries, query_len, fixture->params, + fixture->filter, fixture->filter_len, &result); + } else { + status = fixture->filter_len == 0 + ? paimon_vindex_reader_range_search_batch( + reader, fixture->queries, query_len, fixture->query_count, fixture->params, &result) + : paimon_vindex_reader_range_search_batch_with_roaring_filter( + reader, fixture->queries, query_len, fixture->query_count, fixture->params, + fixture->filter, fixture->filter_len, &result); + } + if (status != 0) fail_ffi("range oracle"); + paimon_vindex_reader_free(reader); + PaimonVindexRangeSearchResultView view = range_view(result); + range_fixture_assert(fixture, &view); + paimon_vindex_range_search_result_destroy(result); +} + +static void test_range_endpoints(void) { + const uint32_t metrics[] = { + PAIMON_VINDEX_METRIC_L2, PAIMON_VINDEX_METRIC_COSINE, PAIMON_VINDEX_METRIC_INNER_PRODUCT}; + const float candidates[] = {-1.0f, -0.5f, -0.0f, 0.0f, 0.25f, 0.5f, 1.0f, 4.0f}; + for (size_t metric_index = 0; metric_index < 3; ++metric_index) { + uint32_t metric = metrics[metric_index]; + PaimonVindexDistanceBand band; + ASSERT_TRUE(paimon_vindex_distance_band_from_endpoints(metric, NULL, NULL, &band) == 0); + ASSERT_TRUE(band.metric == metric); + ASSERT_TRUE(band.lower_kind == PAIMON_VINDEX_BOUND_UNBOUNDED); + ASSERT_TRUE(band.upper_kind == PAIMON_VINDEX_BOUND_UNBOUNDED); + for (uint32_t lower_op = PAIMON_VINDEX_CUT_GE; lower_op <= PAIMON_VINDEX_CUT_GT; ++lower_op) { + for (uint32_t upper_op = PAIMON_VINDEX_CUT_LE; upper_op <= PAIMON_VINDEX_CUT_LT; ++upper_op) { + PaimonVindexDistanceEndpoint lower = {0.5, lower_op}; + PaimonVindexDistanceEndpoint upper = {1.0, upper_op}; + ASSERT_TRUE(paimon_vindex_distance_band_from_endpoints(metric, &lower, &upper, &band) == 0); + for (size_t candidate = 0; candidate < sizeof(candidates) / sizeof(candidates[0]); ++candidate) { + float distance = candidates[candidate]; + if (metric == PAIMON_VINDEX_METRIC_L2 && distance < 0) continue; + double value = metric == PAIMON_VINDEX_METRIC_L2 ? (double)sqrtf(distance) + : metric == PAIMON_VINDEX_METRIC_INNER_PRODUCT ? -(double)distance : (double)distance; + int expected = (lower_op == PAIMON_VINDEX_CUT_GE ? value >= 0.5 : value > 0.5) && + (upper_op == PAIMON_VINDEX_CUT_LE ? value <= 1.0 : value < 1.0); + ASSERT_TRUE(expected == + ((band.lower_kind == PAIMON_VINDEX_BOUND_UNBOUNDED || distance >= band.lower) && + (band.upper_kind == PAIMON_VINDEX_BOUND_UNBOUNDED || distance < band.upper))); + } + } + } + PaimonVindexDistanceEndpoint precise = {nextafter(1.0, 2.0), PAIMON_VINDEX_CUT_GE}; + ASSERT_TRUE(paimon_vindex_distance_band_from_endpoints(metric, &precise, NULL, &band) == 0); + float boundary = metric == PAIMON_VINDEX_METRIC_INNER_PRODUCT ? -1.0f : 1.0f; + ASSERT_TRUE(!((band.lower_kind == PAIMON_VINDEX_BOUND_UNBOUNDED || boundary >= band.lower) && + (band.upper_kind == PAIMON_VINDEX_BOUND_UNBOUNDED || boundary < band.upper))); + } + PaimonVindexDistanceBand band; + PaimonVindexDistanceEndpoint endpoint = {1.0, PAIMON_VINDEX_CUT_LT}; + ASSERT_TRUE(paimon_vindex_distance_band_from_endpoints(PAIMON_VINDEX_METRIC_L2, &endpoint, NULL, &band) != 0); + endpoint.op = PAIMON_VINDEX_CUT_GE; + ASSERT_TRUE(paimon_vindex_distance_band_from_endpoints(PAIMON_VINDEX_METRIC_L2, NULL, &endpoint, &band) != 0); + endpoint.value = NAN; + ASSERT_TRUE(paimon_vindex_distance_band_from_endpoints(PAIMON_VINDEX_METRIC_L2, &endpoint, NULL, &band) != 0); + endpoint.value = INFINITY; + ASSERT_TRUE(paimon_vindex_distance_band_from_endpoints(PAIMON_VINDEX_METRIC_L2, &endpoint, NULL, &band) != 0); + endpoint.value = 1; + endpoint.op = UINT32_MAX; + ASSERT_TRUE(paimon_vindex_distance_band_from_endpoints(PAIMON_VINDEX_METRIC_L2, &endpoint, NULL, &band) != 0); + ASSERT_TRUE(paimon_vindex_distance_band_from_endpoints(UINT32_MAX, NULL, NULL, &band) != 0); + ASSERT_TRUE(paimon_vindex_distance_band_from_endpoints(PAIMON_VINDEX_METRIC_L2, NULL, NULL, NULL) != 0); + paimon_vindex_range_search_result_destroy(NULL); + printf("PASS range_endpoints\n"); +} + +#define ASSERT_RANGE_ERROR(expression) do { \ + result = (PaimonVindexRangeSearchResult *)(uintptr_t)1; \ + ASSERT_TRUE((expression) != 0); \ + ASSERT_TRUE(result == NULL); \ + ASSERT_TRUE(paimon_vindex_last_error() != NULL); \ +} while (0) + +static void test_range_errors(PaimonVindexReaderHandle *reader, const float *query, + PaimonVindexRangeSearchParams params) { + PaimonVindexRangeSearchResult *result; + ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search(NULL, query, RANGE_DIMENSION, params, &result)); + ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search(reader, NULL, RANGE_DIMENSION, params, &result)); + ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search(reader, query, RANGE_DIMENSION - 1, params, &result)); + ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search(reader, NULL, 0, params, &result)); + ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search_batch(reader, NULL, 0, 0, params, &result)); + ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search_batch(reader, query, RANGE_DIMENSION, 2, params, &result)); + ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search_batch( + reader, NULL, 0, SIZE_MAX / RANGE_DIMENSION + 1, params, &result)); + ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search_with_roaring_filter( + reader, query, RANGE_DIMENSION, params, NULL, 1, &result)); + ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search_with_roaring_filter( + reader, query, RANGE_DIMENSION, params, NULL, 0, &result)); + ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search_batch_with_roaring_filter( + reader, query, RANGE_DIMENSION, 1, params, range_filter, 1, &result)); + ASSERT_TRUE(paimon_vindex_reader_range_search(reader, query, RANGE_DIMENSION, params, NULL) != 0); + int supported; + ASSERT_TRUE(paimon_vindex_reader_supports_range_search(NULL, &supported) != 0); + ASSERT_TRUE(paimon_vindex_reader_supports_range_search(reader, NULL) != 0); + PaimonVindexRangeSearchResultView view; + ASSERT_TRUE(paimon_vindex_range_search_result_view(NULL, &view) != 0); + PaimonVindexRangeSearchParams invalid = params; + invalid.nprobe = 0; + ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search(reader, query, RANGE_DIMENSION, invalid, &result)); + invalid = params; + invalid.band.metric = (params.band.metric + 1) % 3; + ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search(reader, query, RANGE_DIMENSION, invalid, &result)); + invalid.band.metric = UINT32_MAX; + ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search(reader, query, RANGE_DIMENSION, invalid, &result)); + invalid = params; + invalid.band.lower_kind = UINT32_MAX; + ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search(reader, query, RANGE_DIMENSION, invalid, &result)); + invalid = params; + invalid.band.upper_kind = PAIMON_VINDEX_BOUND_FINITE; + invalid.band.upper = NAN; + ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search(reader, query, RANGE_DIMENSION, invalid, &result)); + invalid.band.upper = INFINITY; + ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search(reader, query, RANGE_DIMENSION, invalid, &result)); + invalid.band.upper = 0; + invalid.band.lower_kind = PAIMON_VINDEX_BOUND_FINITE; + invalid.band.lower = 1; + ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search(reader, query, RANGE_DIMENSION, invalid, &result)); + invalid.band.lower = 0; + invalid.nprobe = 0; + ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search(reader, query, RANGE_DIMENSION, invalid, &result)); + float bad_query[RANGE_DIMENSION]; + memcpy(bad_query, query, sizeof(bad_query)); + bad_query[0] = NAN; + ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search(reader, bad_query, RANGE_DIMENSION, params, &result)); +} + +static void test_range_matrix(void) { + const char *index_types[] = {"ivf_flat", "ivf_sq", "ivf_pq", "ivf_rq"}; + const char *metrics[] = {"l2", "inner_product", "cosine"}; + const uint32_t metric_codes[] = { + PAIMON_VINDEX_METRIC_L2, PAIMON_VINDEX_METRIC_INNER_PRODUCT, PAIMON_VINDEX_METRIC_COSINE}; + float data[RANGE_VECTOR_COUNT * RANGE_DIMENSION]; + int64_t labels[RANGE_VECTOR_COUNT]; + float queries[RANGE_QUERY_COUNT * RANGE_DIMENSION]; + range_fill_data(data, labels, queries); + for (size_t index_type = 0; index_type < 4; ++index_type) { + for (size_t metric = 0; metric < 3; ++metric) { + const char *keys[] = {"index.type", "dimension", "nlist", "metric"}; + const char *values[] = {index_types[index_type], "8", "4", metrics[metric]}; + PaimonVindexTrainerHandle *trainer = paimon_vindex_trainer_open(keys, values, 4); + if (trainer == NULL) fail_ffi("range trainer open"); + ASSERT_TRUE(paimon_vindex_trainer_add_training_vectors(trainer, data, RANGE_VECTOR_COUNT) == 0); + PaimonVindexTrainingHandle *training = paimon_vindex_trainer_finish(trainer); + ASSERT_TRUE(training != NULL); + paimon_vindex_trainer_free(trainer); + PaimonVindexWriterHandle *writer = paimon_vindex_writer_open(training); + ASSERT_TRUE(writer != NULL); + paimon_vindex_training_free(training); + ASSERT_TRUE(paimon_vindex_writer_add_vectors(writer, labels, data, RANGE_VECTOR_COUNT) == 0); + struct MemBuffer buffer = {0}; + PaimonVindexOutputFile output = { + .ctx = &buffer, .write_fn = mem_write, .flush_fn = mem_flush, .get_pos_fn = mem_pos}; + ASSERT_TRUE(paimon_vindex_writer_write_index(writer, output) == 0); + paimon_vindex_writer_free(writer); + PaimonVindexInputFile input = {.ctx = &buffer, .read_ranges_fn = mem_read_ranges}; + PaimonVindexReaderHandle *reader = paimon_vindex_reader_open(input); + ASSERT_TRUE(reader != NULL); + int supported = 0; + ASSERT_TRUE(paimon_vindex_reader_supports_range_search(reader, &supported) == 0 && supported == 1); + PaimonVindexRangeSearchParams params = range_all_params(metric_codes[metric]); + PaimonVindexRangeSearchResult *batch = NULL; + ASSERT_TRUE(paimon_vindex_reader_range_search_batch( + reader, queries, RANGE_QUERY_COUNT * RANGE_DIMENSION, RANGE_QUERY_COUNT, params, &batch) == 0); + PaimonVindexRangeSearchResultView batch_view = range_view(batch); + range_assert_shape(&batch_view); + ASSERT_TRUE(batch_view.query_count == RANGE_QUERY_COUNT); + ASSERT_TRUE(batch_view.hit_count == RANGE_VECTOR_COUNT * RANGE_QUERY_COUNT); + ASSERT_TRUE(batch_view.list_reads > 0); + for (size_t query = 0; query < RANGE_QUERY_COUNT; ++query) { + ASSERT_TRUE(batch_view.stats[query].lists_probed == RANGE_NLIST); + ASSERT_TRUE(batch_view.stats[query].rows_scanned == RANGE_VECTOR_COUNT); + ASSERT_TRUE(batch_view.stats[query].early_abandoned == 0); + } + PaimonVindexRangeSearchResult *single = NULL; + ASSERT_TRUE(paimon_vindex_reader_range_search(reader, queries, RANGE_DIMENSION, params, &single) == 0); + PaimonVindexRangeSearchResultView single_view = range_view(single); + ASSERT_TRUE(single_view.hit_count == RANGE_VECTOR_COUNT); + PaimonVindexRangeSearchResult *filtered = NULL; + ASSERT_TRUE(paimon_vindex_reader_range_search_batch_with_roaring_filter( + reader, queries, RANGE_QUERY_COUNT * RANGE_DIMENSION, RANGE_QUERY_COUNT, params, + range_filter, sizeof(range_filter), &filtered) == 0); + PaimonVindexRangeSearchResultView filtered_view = range_view(filtered); + range_assert_shape(&filtered_view); + ASSERT_TRUE(filtered_view.hit_count == RANGE_QUERY_COUNT * 2); + for (size_t hit = 0; hit < filtered_view.hit_count; ++hit) { + ASSERT_TRUE(filtered_view.labels[hit] == labels[1] || filtered_view.labels[hit] == labels[3]); + } + paimon_vindex_range_search_result_destroy(filtered); + ASSERT_TRUE(paimon_vindex_reader_range_search_with_roaring_filter( + reader, queries, RANGE_DIMENSION, params, range_filter, sizeof(range_filter), &filtered) == 0); + filtered_view = range_view(filtered); + ASSERT_TRUE(filtered_view.hit_count == 2); + paimon_vindex_range_search_result_destroy(filtered); + ASSERT_TRUE(paimon_vindex_reader_range_search_with_roaring_filter( + reader, queries, RANGE_DIMENSION, params, range_empty_filter, sizeof(range_empty_filter), &filtered) == 0); + filtered_view = range_view(filtered); + range_assert_shape(&filtered_view); + ASSERT_TRUE(filtered_view.hit_count == 0 && filtered_view.lims[1] == 0); + paimon_vindex_range_search_result_destroy(filtered); + PaimonVindexRangeSearchParams bounded = params; + float minimum = single_view.distances[0]; + float maximum = minimum; + for (size_t hit = 1; hit < single_view.hit_count; ++hit) { + minimum = fminf(minimum, single_view.distances[hit]); + maximum = fmaxf(maximum, single_view.distances[hit]); + } + bounded.band.lower_kind = PAIMON_VINDEX_BOUND_FINITE; + bounded.band.lower = metric_codes[metric] == PAIMON_VINDEX_METRIC_L2 + ? fmaxf(0, minimum) : minimum; + bounded.band.upper_kind = PAIMON_VINDEX_BOUND_FINITE; + bounded.band.upper = minimum + (maximum - minimum) / 2; + PaimonVindexRangeSearchResult *subset = NULL; + ASSERT_TRUE(paimon_vindex_reader_range_search_batch( + reader, queries, RANGE_QUERY_COUNT * RANGE_DIMENSION, RANGE_QUERY_COUNT, bounded, &subset) == 0); + PaimonVindexRangeSearchResultView subset_view = range_view(subset); + range_assert_shape(&subset_view); + ASSERT_TRUE(subset_view.hit_count > 0 && subset_view.hit_count < batch_view.hit_count); + for (size_t query = 0; query < RANGE_QUERY_COUNT; ++query) { + size_t expected = 0; + for (size_t hit = batch_view.lims[query]; hit < batch_view.lims[query + 1]; ++hit) { + if (batch_view.distances[hit] >= bounded.band.lower && batch_view.distances[hit] < bounded.band.upper) { + ++expected; + int found = 0; + for (size_t candidate = subset_view.lims[query]; candidate < subset_view.lims[query + 1]; ++candidate) { + if (subset_view.labels[candidate] == batch_view.labels[hit]) found = 1; + } + ASSERT_TRUE(found); + } + } + ASSERT_TRUE(subset_view.lims[query + 1] - subset_view.lims[query] == expected); + } + paimon_vindex_range_search_result_destroy(subset); + bounded.band.lower = bounded.band.upper; + ASSERT_TRUE(paimon_vindex_reader_range_search_batch( + reader, queries, RANGE_QUERY_COUNT * RANGE_DIMENSION, RANGE_QUERY_COUNT, bounded, &subset) == 0); + subset_view = range_view(subset); + range_assert_shape(&subset_view); + ASSERT_TRUE(subset_view.hit_count == 0 && subset_view.list_reads == 0); + paimon_vindex_range_search_result_destroy(subset); + test_range_errors(reader, queries, params); + ASSERT_TRUE(paimon_vindex_range_search_result_view(single, NULL) != 0); + paimon_vindex_reader_free(reader); + free(buffer.data); + PaimonVindexRangeSearchResultView retained_view = range_view(single); + ASSERT_TRUE(retained_view.labels == single_view.labels); + ASSERT_TRUE(retained_view.distances == single_view.distances); + range_assert_shape(&retained_view); + range_assert_shape(&batch_view); + paimon_vindex_range_search_result_destroy(single); + paimon_vindex_range_search_result_destroy(batch); + printf("PASS range_matrix %s %s\n", index_types[index_type], metrics[metric]); + } + } +} + int main(void) { test_extensible_search_params_defaults(); test_supported_index_roundtrips(); test_output_write_callback_error_propagates(); test_output_flush_callback_error_propagates(); test_input_read_ranges_callback_error_propagates(); + test_range_endpoints(); + test_range_matrix(); + range_fixture_run_all(consume_range_fixture); return 0; } diff --git a/cpp/test_vindex.cpp b/cpp/test_vindex.cpp index af7104a8..f504b87b 100644 --- a/cpp/test_vindex.cpp +++ b/cpp/test_vindex.cpp @@ -25,6 +25,7 @@ #include #include #include +#include #include #include @@ -42,6 +43,8 @@ } \ } while (0) +#include "../c/range_test_support.h" + struct MemBuffer { std::vector data; size_t pos = 0; @@ -233,6 +236,16 @@ static void run_roundtrip( ASSERT_EQ(batch.ids.size(), 2); assert_id_in_cluster(batch.ids[0], 0); assert_id_in_cluster(batch.ids[1], 1); + ASSERT_EQ(reader.supports_range_search(), expected_index_type != PAIMON_VINDEX_INDEX_TYPE_DISKANN); + if (expected_index_type == PAIMON_VINDEX_INDEX_TYPE_DISKANN) { + bool rejected = false; + try { + reader.range_search(query, paimon::vindex::RangeSearchParams{}); + } catch (const paimon::vindex::Error&) { + rejected = true; + } + ASSERT_TRUE(rejected); + } printf("PASS %s\n", name); } @@ -339,9 +352,228 @@ static void test_extensible_search_params_forward_query_tuning() { printf("PASS extensible_search_params_forward_query_tuning\n"); } +template +static void assert_range_error(Operation operation) { + bool rejected = false; + try { + operation(); + } catch (const paimon::vindex::Error& error) { + rejected = !std::string(error.what()).empty(); + } + ASSERT_TRUE(rejected); +} + +static PaimonVindexRangeSearchResultView range_view( + const paimon::vindex::RangeSearchResult& result) { + return {result.query_count, result.labels.size(), result.lims.data(), + result.labels.data(), result.distances.data(), result.stats.data(), + result.list_reads}; +} + +static paimon::vindex::RangeSearchParams range_cpp_params( + PaimonVindexRangeSearchParams raw) { + return {{raw.band.metric, raw.band.lower_kind, raw.band.lower, + raw.band.upper_kind, raw.band.upper}, raw.nprobe}; +} + +static void consume_range_fixture(const RangeFixture* fixture) { + MemBuffer buffer; + buffer.data.assign(fixture->index_data, fixture->index_data + fixture->index_len); + paimon::vindex::RangeSearchResult result; + { + paimon::vindex::Reader reader(make_input(buffer)); + ASSERT_TRUE(reader.supports_range_search()); + ASSERT_EQ(reader.metadata().dimension, fixture->dimension); + auto params = range_cpp_params(fixture->params); + const size_t query_len = fixture->dimension * fixture->query_count; + if (fixture->query_count == 1) { + result = fixture->filter_len == 0 + ? reader.range_search(fixture->queries, query_len, params) + : reader.range_search_with_roaring_filter( + fixture->queries, query_len, params, fixture->filter, fixture->filter_len); + } else { + result = fixture->filter_len == 0 + ? reader.range_search_batch(fixture->queries, query_len, fixture->query_count, params) + : reader.range_search_batch_with_roaring_filter( + fixture->queries, query_len, fixture->query_count, params, + fixture->filter, fixture->filter_len); + } + } + auto view = range_view(result); + range_fixture_assert(fixture, &view); +} + +static void test_range_endpoints() { + using namespace paimon::vindex; + for (uint32_t metric : {PAIMON_VINDEX_METRIC_L2, PAIMON_VINDEX_METRIC_COSINE, + PAIMON_VINDEX_METRIC_INNER_PRODUCT}) { + for (uint32_t lower_op : {PAIMON_VINDEX_CUT_GE, PAIMON_VINDEX_CUT_GT}) { + for (uint32_t upper_op : {PAIMON_VINDEX_CUT_LE, PAIMON_VINDEX_CUT_LT}) { + auto band = DistanceBand::from_endpoints( + metric, DistanceEndpoint{0.5, lower_op}, DistanceEndpoint{1.0, upper_op}); + for (float distance : {-1.0f, -0.5f, -0.0f, 0.0f, 0.25f, + std::nextafter(0.25f, 0.0f), 0.5f, 1.0f, + std::nextafter(1.0f, 2.0f), 4.0f}) { + if (metric == PAIMON_VINDEX_METRIC_L2 && distance < 0) continue; + double public_value = metric == PAIMON_VINDEX_METRIC_L2 + ? static_cast(std::sqrt(distance)) + : metric == PAIMON_VINDEX_METRIC_INNER_PRODUCT + ? -static_cast(distance) : static_cast(distance); + bool expected = (lower_op == PAIMON_VINDEX_CUT_GE + ? public_value >= 0.5 : public_value > 0.5) && + (upper_op == PAIMON_VINDEX_CUT_LE ? public_value <= 1.0 : public_value < 1.0); + ASSERT_EQ(expected, + (band.lower_kind == PAIMON_VINDEX_BOUND_UNBOUNDED || distance >= band.lower) && + (band.upper_kind == PAIMON_VINDEX_BOUND_UNBOUNDED || distance < band.upper)); + } + } + } + auto unbounded = DistanceBand::from_endpoints(metric); + ASSERT_EQ(unbounded.lower_kind, PAIMON_VINDEX_BOUND_UNBOUNDED); + ASSERT_EQ(unbounded.upper_kind, PAIMON_VINDEX_BOUND_UNBOUNDED); + auto precise = DistanceBand::from_endpoints( + metric, DistanceEndpoint{std::nextafter(1.0, 2.0), PAIMON_VINDEX_CUT_GE}); + float boundary = metric == PAIMON_VINDEX_METRIC_INNER_PRODUCT ? -1.0f : 1.0f; + ASSERT_TRUE(!((precise.lower_kind == PAIMON_VINDEX_BOUND_UNBOUNDED || boundary >= precise.lower) && + (precise.upper_kind == PAIMON_VINDEX_BOUND_UNBOUNDED || boundary < precise.upper))); + } + assert_range_error([] { + DistanceBand::from_endpoints(PAIMON_VINDEX_METRIC_L2, + DistanceEndpoint{1.0, PAIMON_VINDEX_CUT_LT}); + }); + assert_range_error([] { + DistanceBand::from_endpoints(PAIMON_VINDEX_METRIC_L2, std::nullopt, + DistanceEndpoint{std::numeric_limits::infinity(), PAIMON_VINDEX_CUT_LE}); + }); + printf("PASS range_endpoints\n"); +} + +static void test_range_matrix() { + using namespace paimon::vindex; + const std::pair metrics[] = { + {"l2", PAIMON_VINDEX_METRIC_L2}, {"cosine", PAIMON_VINDEX_METRIC_COSINE}, + {"inner_product", PAIMON_VINDEX_METRIC_INNER_PRODUCT}}; + std::vector data(RANGE_VECTOR_COUNT * RANGE_DIMENSION); + std::vector labels(RANGE_VECTOR_COUNT); + std::vector queries(RANGE_QUERY_COUNT * RANGE_DIMENSION); + range_fill_data(data.data(), labels.data(), queries.data()); + const std::vector query(queries.begin(), queries.begin() + RANGE_DIMENSION); + for (const char* index_type : {"ivf_flat", "ivf_sq", "ivf_pq", "ivf_rq"}) { + for (const auto& metric : metrics) { + Trainer trainer({{"index.type", index_type}, {"dimension", "8"}, + {"nlist", "4"}, {"metric", metric.first}}); + Writer writer(trainer.add_training_vectors(data.data(), RANGE_VECTOR_COUNT).finish_training()); + writer.add_vectors(labels.data(), data.data(), RANGE_VECTOR_COUNT); + MemBuffer buffer; + writer.write_index(make_output(buffer)); + RangeSearchResult retained; + { + Reader reader(make_input(buffer)); + ASSERT_TRUE(reader.supports_range_search()); + auto params = range_cpp_params(range_all_params(metric.second)); + auto batch = reader.range_search_batch(queries, RANGE_QUERY_COUNT, params); + auto batch_view = range_view(batch); + range_assert_shape(&batch_view); + ASSERT_EQ(batch.query_count, RANGE_QUERY_COUNT); + ASSERT_EQ(batch.labels.size(), RANGE_QUERY_COUNT * RANGE_VECTOR_COUNT); + ASSERT_TRUE(batch.list_reads > 0); + for (const auto& stats : batch.stats) { + ASSERT_EQ(stats.lists_probed, RANGE_NLIST); + ASSERT_EQ(stats.rows_scanned, RANGE_VECTOR_COUNT); + ASSERT_EQ(stats.rows_committed, RANGE_VECTOR_COUNT); + ASSERT_EQ(stats.early_abandoned, 0); + } + retained = reader.range_search(query, params); + ASSERT_EQ(retained.labels.size(), RANGE_VECTOR_COUNT); + auto topk = reader.search(query.data(), SearchParams{5, RANGE_NLIST}); + ASSERT_EQ(topk.ids.size(), 5); + for (int64_t label : topk.ids) { + ASSERT_TRUE(std::find(retained.labels.begin(), retained.labels.end(), label) != retained.labels.end()); + } + auto filtered = reader.range_search_batch_with_roaring_filter( + queries, RANGE_QUERY_COUNT, params, range_filter, sizeof(range_filter)); + auto filtered_view = range_view(filtered); + range_assert_shape(&filtered_view); + ASSERT_EQ(filtered.labels.size(), RANGE_QUERY_COUNT * 2); + for (int64_t label : filtered.labels) { + ASSERT_TRUE(label == labels[1] || label == labels[3]); + } + auto filtered_single = reader.range_search_with_roaring_filter( + query, params, range_filter, sizeof(range_filter)); + ASSERT_EQ(filtered_single.labels.size(), 2); + auto empty_filter = reader.range_search_with_roaring_filter( + query, params, range_empty_filter, sizeof(range_empty_filter)); + ASSERT_TRUE(empty_filter.labels.empty()); + ASSERT_TRUE(empty_filter.lims == std::vector({0, 0})); + auto bounded = params; + auto bounds = std::minmax_element(retained.distances.begin(), retained.distances.end()); + bounded.band.lower_kind = PAIMON_VINDEX_BOUND_FINITE; + bounded.band.lower = metric.second == PAIMON_VINDEX_METRIC_L2 + ? std::max(0.0f, *bounds.first) : *bounds.first; + bounded.band.upper_kind = PAIMON_VINDEX_BOUND_FINITE; + bounded.band.upper = *bounds.first + (*bounds.second - *bounds.first) / 2.0f; + auto subset = reader.range_search_batch(queries, RANGE_QUERY_COUNT, bounded); + auto subset_view = range_view(subset); + range_assert_shape(&subset_view); + ASSERT_TRUE(subset.labels.size() > 0 && subset.labels.size() < batch.labels.size()); + for (size_t query_index = 0; query_index < RANGE_QUERY_COUNT; ++query_index) { + size_t expected = 0; + for (size_t hit = batch.lims[query_index]; hit < batch.lims[query_index + 1]; ++hit) { + if (batch.distances[hit] >= bounded.band.lower && batch.distances[hit] < bounded.band.upper) { + ++expected; + ASSERT_TRUE(std::find(subset.labels.begin() + subset.lims[query_index], + subset.labels.begin() + subset.lims[query_index + 1], batch.labels[hit]) != + subset.labels.begin() + subset.lims[query_index + 1]); + } + } + ASSERT_EQ(subset.lims[query_index + 1] - subset.lims[query_index], expected); + } + bounded.band.lower = bounded.band.upper; + auto empty = reader.range_search_batch(queries, RANGE_QUERY_COUNT, bounded); + ASSERT_TRUE(empty.labels.empty() && empty.distances.empty()); + ASSERT_TRUE(empty.lims == std::vector({0, 0, 0, 0})); + ASSERT_EQ(empty.list_reads, 0); + assert_range_error([&] { reader.range_search(nullptr, RANGE_DIMENSION, params); }); + assert_range_error([&] { reader.range_search(std::vector{}, params); }); + assert_range_error([&] { reader.range_search_batch(queries, 2, params); }); + assert_range_error([&] { reader.range_search_batch(nullptr, 0, 0, params); }); + assert_range_error([&] { + reader.range_search_batch(nullptr, 0, SIZE_MAX / RANGE_DIMENSION + 1, params); + }); + assert_range_error([&] { reader.range_search_with_roaring_filter(query, params, nullptr, 1); }); + assert_range_error([&] { reader.range_search_with_roaring_filter(query, params, range_filter, 1); }); + auto invalid = params; + invalid.nprobe = 0; + assert_range_error([&] { reader.range_search(query, invalid); }); + invalid = bounded; + invalid.nprobe = 0; + assert_range_error([&] { reader.range_search(query, invalid); }); + invalid = params; + invalid.band.metric = (metric.second + 1) % 3; + assert_range_error([&] { reader.range_search(query, invalid); }); + auto nan_query = query; + nan_query[0] = std::numeric_limits::quiet_NaN(); + assert_range_error([&] { reader.range_search(nan_query, params); }); + Reader moved(std::move(reader)); + assert_range_error([&] { reader.supports_range_search(); }); + assert_range_error([&] { reader.range_search(query, params); }); + ASSERT_TRUE(moved.supports_range_search()); + ASSERT_EQ(moved.range_search(query, params).labels.size(), retained.labels.size()); + } + auto retained_view = range_view(retained); + range_assert_shape(&retained_view); + ASSERT_EQ(retained.labels.size(), RANGE_VECTOR_COUNT); + printf("PASS range_matrix %s %s\n", index_type, metric.first); + } + } +} + int main() { test_supported_index_roundtrips(); test_worker_callback_reentry_is_rejected(); test_extensible_search_params_forward_query_tuning(); + test_range_endpoints(); + test_range_matrix(); + range_fixture_run_all(consume_range_fixture); return 0; } diff --git a/docs/api.html b/docs/api.html index 56bd9f9a..97bcb129 100644 --- a/docs/api.html +++ b/docs/api.html @@ -76,7 +76,7 @@

Shared search parameters

Range search parameters and results

-

Rust range search returns every eligible probed row inside a half-open distance band instead of a fixed number of nearest neighbours. All four entry points support IVF-FLAT, IVF-SQ, IVF-PQ and IVF-RQ with L2, cosine and inner product. FLAT tests exact distances; SQ, PQ and RQ test quantized estimates. Query capability with IndexType::supports_range_search(metric) or reader.supports_range_search(). DiskANN reports Unsupported; exact, complete membership requires full-probe FLAT or an exhaustive scan, not top-K followed by filtering. See Range search for public endpoint conversion and the full contract. Range bindings are not included.

+

Range search returns every eligible probed row inside a half-open distance band instead of a fixed number of nearest neighbours. Rust, C, C++, JNI/Java and Python support IVF-FLAT, IVF-SQ, IVF-PQ and IVF-RQ with L2, cosine and inner product, for single/batch queries with or without a Roaring allow-list. FLAT tests exact distances; SQ, PQ and RQ test quantized estimates. Capability queries delegate to core. DiskANN remains unsupported; exact, complete membership requires full-probe FLAT or an exhaustive scan, not top-K followed by filtering. See Range search for the full contract and language APIs, ownership and examples. Existing top-K entry points and parameter layouts are unchanged.

diff --git a/docs/range-search.html b/docs/range-search.html index ca63ae9e..568d0aae 100644 --- a/docs/range-search.html +++ b/docs/range-search.html @@ -11,9 +11,9 @@
-

Every row inside a distance band

Range search

Return every eligible probed row whose family-specific distance falls inside a half-open band [lower, upper), with no result limit. IVF-FLAT computes exact distances; IVF-SQ, IVF-PQ and IVF-RQ compute estimates. Range search answers "which rows are within this distance", where Top-K answers "which rows are closest".

No limit, no capHalf-open intervalIVF-FLAT / SQ / PQ / RQL2 / cosine / inner product · Rust API
+

Every row inside a distance band

Range search

Return every eligible probed row whose family-specific distance falls inside a half-open band [lower, upper), with no result limit. IVF-FLAT computes exact distances; IVF-SQ, IVF-PQ and IVF-RQ compute estimates. Range search answers "which rows are within this distance", where Top-K answers "which rows are closest".

No limit, no capHalf-open intervalIVF-FLAT / SQ / PQ / RQL2 / cosine / inner product · All bindings
- +

Semantic contract

@@ -95,9 +95,84 @@

Usage

For IVF-RQ, lists_probed includes empty selected lists; rows_scanned counts filter-eligible rows evaluated; rows_committed counts returned rows; and early_abandoned is zero. Call-level list_reads counts unique non-empty lists, not query/list pairs or storage read rounds. These result-owned counters leave the last top-K statistics unchanged.

+
+

Language bindings and ownership

+

Every binding delegates to the same core entry points and endpoint conversion. All return raw internal distances and variable-length CSR results: lims has query_count + 1 offsets, and query i occupies [lims[i], lims[i + 1]) in the label and distance arrays. Results have no padding, sorting, top-K cap, or extra reranking. A batch must contain at least one query. One optional Roaring allow-list applies to every query; absent filters, valid serialized empty filters, and malformed zero-byte filters are distinct.

+
Entry pointSignature
Single queryrange_search(query, params)
Single query, filteredrange_search_with_roaring_filter(query, params, filter_bytes)
+ + + + +
BindingAPI and lifetime
C ABI / generated headerpaimon_vindex_reader_range_search, _range_search_batch, and both _with_roaring_filter variants return an owned opaque result through an output pointer. paimon_vindex_range_search_result_view borrows its buffers; paimon_vindex_range_search_result_destroy releases them. Query/filter lengths are explicit. The generated include/paimon_vindex.h is rebuilt by cbindgen, not maintained manually.
C++Reader::range_search, Reader::range_search_batch, and their _with_roaring_filter variants copy CSR buffers into owning vectors. A native-result RAII guard handles cleanup, including allocation exceptions. Existing top-K methods remain unchanged.
JNI / JavaVectorIndexReader.rangeSearch and rangeSearchBatch accept VectorRangeSearchParams and optional filter bytes. VectorRangeSearchResult owns defensively copied Java arrays. Native lengths are checked against JVM array limits; offsets, labels and counters use long.
PythonVectorIndexReader.range_search and range_search_batch accept RangeSearchParams and optional roaring_filter=. RangeSearchResult owns copied NumPy arrays (uintp offsets, int64 labels, float32 distances). Native results are freed in finally, even if conversion fails.
+

Results remain valid after the reader closes. In C, destroy each successful result exactly once, never free its individual arrays, and never access a view after destruction. Input queries and filter bytes are borrowed only for the call. Search sets a valid output pointer to NULL before doing work; an error returns -1, sets paimon_vindex_last_error(), and transfers no result. destroy(NULL) is safe. Callers must supply valid, aligned buffers and live handles; C callers must synchronize operations on a reader. C++, Java and Python retain their existing callback-aware handle locks.

+

Capability is available as paimon_vindex_reader_supports_range_search(reader, &supported), C++/Python reader.supports_range_search(), or Java reader.supportsRangeSearch(). DiskANN returns false. Invalid parameters, malformed filters, unsupported searches, I/O errors and non-finite evaluated distances fail rather than pretending there are no matches. Length multiplication, slice byte limits and language-specific array limits are checked instead of truncating offsets or counters.

+

All results expose per-query lists_probed, rows_scanned, rows_committed and early_abandoned, plus call-level list_reads (camelCase accessors in Java). These are core's logical counters, not binding-side estimates. In particular, list_reads is not the sum of a batch's per-query probe counts, and an IVF-SQ cache hit does not count as a payload read.

+
C · public L2 distance <= 2.0
PaimonVindexDistanceEndpoint upper = {2.0, PAIMON_VINDEX_CUT_LE};
+PaimonVindexDistanceBand band;
+if (paimon_vindex_distance_band_from_endpoints(
+        PAIMON_VINDEX_METRIC_L2, NULL, &upper, &band) != 0) {
+    return -1;
+}
+PaimonVindexRangeSearchParams params = {band, 8};
+PaimonVindexRangeSearchResult *result = NULL;
+int status = paimon_vindex_reader_range_search(
+        reader, query, dimension, params, &result);
+if (status == 0) {
+    PaimonVindexRangeSearchResultView view;
+    status = paimon_vindex_range_search_result_view(result, &view);
+    if (status == 0) {
+        consume_rows(view.labels, view.distances, view.hit_count);
+    }
+}
+paimon_vindex_range_search_result_destroy(result);
+
C++ · owning CSR vectors
using namespace paimon::vindex;
+auto band = DistanceBand::from_endpoints(
+        PAIMON_VINDEX_METRIC_L2, std::nullopt,
+        DistanceEndpoint{2.0, PAIMON_VINDEX_CUT_LE});
+auto result = reader.range_search_batch(queries, query_count, RangeSearchParams{band, 8});
+auto first_begin = result.lims[0];
+auto first_end = result.lims[1];
+
Java · shared filtered batch
VectorDistanceBand band = VectorDistanceBand.fromEndpoints(
+        "l2", null, null, 2.0, VectorDistanceBand.CutOperator.LE);
+VectorRangeSearchParams params = new VectorRangeSearchParams(band, 8);
+VectorRangeSearchResult result = reader.rangeSearchBatch(
+        queries, queryCount, params, roaringFilter);
+long[] labels = result.labelsForQuery(0);
+long listReads = result.listReads();
+
Python · shared filtered batch
from paimon_vindex import DistanceBand, DistanceEndpoint, DistanceEndpointOp, RangeSearchParams
+
+band = DistanceBand.from_endpoints(
+    "l2", upper=DistanceEndpoint(2.0, DistanceEndpointOp.LE))
+result = reader.range_search_batch(
+    queries, RangeSearchParams(band, nprobe=8), roaring_filter=roaring_filter)
+labels, distances = result.query(0)
+stats = result.stats[0]
+

Raw band constructors take float32 cuts in internal distance space. To express public predicate endpoints, use the native-backed conversion helpers instead: do not square L2 endpoints, round double literals to float32, or negate/reverse inner-product endpoints yourself. Missing endpoints are structural unboundedness, not infinities or sentinel numbers.

+

Python can still import and run existing top-K calls with a native library predating the range ABI. If any required range export is absent, capability returns false and new range calls or endpoint conversion request a native-library upgrade with a clear error. Range operations never silently fall back to top-K. NumPy query inputs are normalized for both contiguity and alignment.

+
+ +
+

Cross-language verification

+

The core-only ffi/examples/range_search_fixture.rs generator writes the same index bytes and 168 oracle cases for C, C++, Java and Python: four IVF families, three metrics, bounded/unbounded/empty bands, single/batch queries, and absent/non-empty/empty Roaring filters. Each consumer opens a fresh reader and checks per-query label/distance-bit multisets, CSR offsets, and all statistics against core. Comparison does not impose ordering that core does not promise. Tests also cover invalid endpoints, shape/length errors, unsupported indexes, ownership and existing top-K behavior. CI generates and consumes the oracle in every binding job.

+
Linux · from the repository root
cargo fmt --all -- --check
+cargo test --workspace
+cargo build -p paimon-vindex-ffi -p paimon-vindex-jni
+export PVI_RANGE_FIXTURES="$PWD/target/range-fixtures"
+cargo run -p paimon-vindex-ffi --example range_search_fixture -- "$PVI_RANGE_FIXTURES"
+export PAIMON_VINDEX_LIB_PATH="$PWD/target/debug"
+cmake -S c -B c/build -DPAIMON_VINDEX_FFI_LIB="$PWD/target/debug/libpaimon_vindex_ffi.so"
+cmake --build c/build && c/build/test_vindex
+cmake -S cpp -B cpp/build -DPAIMON_VINDEX_FFI_LIB="$PWD/target/debug/libpaimon_vindex_ffi.so"
+cmake --build cpp/build && cpp/build/test_vindex_cpp
+mvn -f java/pom.xml test -Dpaimon.vindex.native.path="$PWD/target/debug/libpaimon_vindex_jni.so"
+java -cp java/target/test-classes:java/target/classes org.apache.paimon.index.vector.VectorIndexNativeValidationTest "$PWD/target/debug/libpaimon_vindex_jni.so"
+PYTHONPATH=python python3 -m pytest python/tests
+

Use .dylib library paths on macOS. Standalone binding tests still run without PVI_RANGE_FIXTURES; setting it enables the shared core oracle as well. No algorithm, index storage format or performance-tuning change is part of these additive bindings.

+
+

Choosing an index type

-

Range search supports IVF-FLAT, IVF-SQ, IVF-PQ and IVF-RQ under L2, cosine and inner product, with fixed positive probe widths and Rust entry points. Query capability through IndexType::supports_range_search(metric) or reader.supports_range_search(). DiskANN remains unsupported; no storage-format, C/JNI binding or top-K behavior changes are included.

+

Range search supports IVF-FLAT, IVF-SQ, IVF-PQ and IVF-RQ under L2, cosine and inner product, with fixed positive probe widths in Rust and every language binding. Query capability through core or the binding's reader capability method. DiskANN remains unsupported; storage-format and top-K behavior are unchanged.

Choose according to the membership requirement.IVF-FLAT tests full-vector distances. IVF-RQ uses RaBitQ, with a one-bit estimate for one-bit files and the full multi-bit estimate otherwise, not Faiss's residual/additive quantizer. Its band predicate is precise relative to that estimate, not to the raw vector. Tests cover an independent estimated-distance oracle, single/batch equivalence, filters, statistics, parallel scans, and non-finite inputs/data; they do not establish exact-distance recall guarantees. See IVF-RQ range semantics.
IVF-SQ membership uses an estimate.IVF-FLAT computes exact distances from stored f32 vectors. IVF-SQ instead reuses top-K's blocked scalar-quantized estimator, reconstructing residuals with each list's stored bounds and centroid. The same estimated value determines band membership and is returned in distances; there is no original-vector reranking and no top-K fallback. Prefer IVF-FLAT if original-distance membership must be exact.
IVF-PQ uses complete floating-point ADC estimates.Both packed 4-bit and 8-bit codes, residual encoding and OPQ are supported. L2 sums direct squared subvector distances. Cosine uses half the ADC squared distance after query normalization, a unit-vector surrogate rather than exact cosine of a re-normalized reconstruction. Inner product uses negative estimated dot product. The range path does not reuse top-K's quantized FastScan tables, expanded L2 tables or cosine score scale, so scores need not be bit-identical to top-K. It never truncates or reranks raw vectors. Shared lists are read once, oversized lists stream in bounded chunks, and the allow-list is evaluated once per list row across the batch.
diff --git a/ffi/Cargo.toml b/ffi/Cargo.toml index dda13433..807672f9 100644 --- a/ffi/Cargo.toml +++ b/ffi/Cargo.toml @@ -33,3 +33,6 @@ paimon-vindex-core = { path = "../core", version = "0.6.0" } [build-dependencies] cbindgen = "0.27" + +[dev-dependencies] +roaring = "0.11" diff --git a/ffi/examples/range_search_fixture.rs b/ffi/examples/range_search_fixture.rs new file mode 100644 index 00000000..77eca7fa --- /dev/null +++ b/ffi/examples/range_search_fixture.rs @@ -0,0 +1,206 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +use paimon_vindex_core::distance::MetricType; +use paimon_vindex_core::index::{ + VectorIndexConfig, VectorIndexReader, VectorIndexTrainer, VectorIndexWriter, +}; +use paimon_vindex_core::io::PosWriter; +use paimon_vindex_core::range::{Bound, DistanceBand, VectorRangeSearchParams}; +use roaring::RoaringTreemap; +use std::collections::HashMap; +use std::fs::{self, File}; +use std::io::{self, Cursor, Write}; +use std::path::Path; + +const DIMENSION: usize = 16; +const ROWS: usize = 256; +const NPROBE: usize = 3; + +fn main() -> io::Result<()> { + let output = std::env::args_os() + .nth(1) + .ok_or_else(|| io::Error::other("usage: range_search_fixture OUTPUT_DIR"))?; + let output = Path::new(&output); + fs::create_dir_all(output)?; + let mut manifest = File::create(output.join("manifest.txt"))?; + let vectors = (0..ROWS) + .flat_map(|row| { + (0..DIMENSION) + .map(move |dimension| ((row * 13 + dimension * 7) % 67) as f32 / 31.0 - 1.0) + }) + .collect::>(); + let ids = (0..ROWS) + .map(|row| (1_i64 << 40) + row as i64) + .collect::>(); + let queries = [0, 17, 129] + .into_iter() + .flat_map(|row| { + vectors[row * DIMENSION..(row + 1) * DIMENSION] + .iter() + .copied() + }) + .collect::>(); + let allow_list = ids + .iter() + .enumerate() + .filter(|(row, _)| row % 2 == 1) + .map(|(_, &label)| label as u64) + .collect::(); + let mut filter = Vec::new(); + allow_list.serialize_into(&mut filter)?; + let mut empty_filter = Vec::new(); + RoaringTreemap::new().serialize_into(&mut empty_filter)?; + let mut case_count = 0; + + for family in ["ivf_flat", "ivf_sq", "ivf_pq", "ivf_rq"] { + for (metric_name, metric, metric_code, lower, upper) in [ + ("l2", MetricType::L2, 0, 0.15, 1.75), + ("cosine", MetricType::Cosine, 2, -0.15, 0.45), + ("inner_product", MetricType::InnerProduct, 1, -4.0, 0.25), + ] { + let options = HashMap::from([ + ("index.type".to_string(), family.to_string()), + ("dimension".to_string(), DIMENSION.to_string()), + ("metric".to_string(), metric_name.to_string()), + ("nlist".to_string(), NPROBE.to_string()), + ]); + let config = VectorIndexConfig::from_options(&options)?; + let training = VectorIndexTrainer::train(config, &vectors, ROWS)?; + let mut writer = VectorIndexWriter::new(training); + writer.add_vectors(&ids, &vectors, ROWS)?; + let mut bytes = Vec::new(); + writer.write(&mut PosWriter::new(&mut bytes))?; + let index_name = format!("{family}-{metric_name}.index"); + fs::write(output.join(&index_name), &bytes)?; + + for (band_name, low, high) in [ + ("bounded", Bound::Finite(lower), Bound::Finite(upper)), + ("unbounded", Bound::Unbounded, Bound::Unbounded), + ("empty", Bound::Finite(0.0), Bound::Finite(0.0)), + ] { + for query_count in [1, 3] { + for (filter_name, filter_bytes) in + [("all", None), ("filtered", Some(filter.as_slice()))] + { + let name = format!( + "{family}-{metric_name}-{band_name}-{query_count}-{filter_name}" + ); + write_case( + output, + &name, + &bytes, + metric, + metric_code, + low, + high, + &queries[..query_count * DIMENSION], + query_count, + filter_bytes, + )?; + writeln!(manifest, "{name} {index_name}")?; + case_count += 1; + } + } + } + for query_count in [1, 3] { + let name = format!("{family}-{metric_name}-empty-filter-{query_count}"); + write_case( + output, + &name, + &bytes, + metric, + metric_code, + Bound::Unbounded, + Bound::Unbounded, + &queries[..query_count * DIMENSION], + query_count, + Some(empty_filter.as_slice()), + )?; + writeln!(manifest, "{name} {index_name}")?; + case_count += 1; + } + } + } + println!( + "Generated {case_count} core range-search oracle cases in {}", + output.display() + ); + Ok(()) +} + +#[allow(clippy::too_many_arguments)] +fn write_case( + output: &Path, + name: &str, + index: &[u8], + metric: MetricType, + metric_code: u32, + lower: Bound, + upper: Bound, + queries: &[f32], + query_count: usize, + filter: Option<&[u8]>, +) -> io::Result<()> { + let mut reader = VectorIndexReader::open(Cursor::new(index.to_vec()))?; + assert!(reader.supports_range_search()); + let params = VectorRangeSearchParams::new(DistanceBand::new(lower, upper, metric)?, NPROBE); + let result = match (query_count == 1, filter) { + (true, None) => reader.range_search(queries, params), + (true, Some(filter)) => reader.range_search_with_roaring_filter(queries, params, filter), + (false, None) => reader.range_search_batch(queries, query_count, params), + (false, Some(filter)) => { + reader.range_search_batch_with_roaring_filter(queries, query_count, params, filter) + } + }?; + let encode = |bound| match bound { + Bound::Unbounded => (0, 0), + Bound::Finite(value) => (1, value.to_bits()), + }; + let (lower_kind, lower_bits) = encode(lower); + let (upper_kind, upper_bits) = encode(upper); + let filter = filter.unwrap_or(&[]); + let mut file = File::create(output.join(format!("{name}.expected")))?; + writeln!(file, "{DIMENSION} {metric_code} {query_count} {NPROBE} {lower_kind} {lower_bits} {upper_kind} {upper_bits} {} {} {}", filter.len(), result.labels().len(), result.call_stats().list_reads())?; + for query in queries { + writeln!(file, "{}", query.to_bits())?; + } + for byte in filter { + writeln!(file, "{byte}")?; + } + for offset in result.lims() { + writeln!(file, "{offset}")?; + } + for label in result.labels() { + writeln!(file, "{label}")?; + } + for distance in result.distances() { + writeln!(file, "{}", distance.to_bits())?; + } + for query in 0..query_count { + let stats = result.query(query).stats; + writeln!( + file, + "{} {} {} {}", + stats.lists_probed(), + stats.rows_scanned(), + stats.rows_committed(), + stats.early_abandoned() + )?; + } + Ok(()) +} diff --git a/ffi/src/lib.rs b/ffi/src/lib.rs index 90b99160..b6dc138d 100644 --- a/ffi/src/lib.rs +++ b/ffi/src/lib.rs @@ -17,6 +17,12 @@ #![allow(clippy::missing_safety_doc)] +mod range; +pub use range::*; + +#[cfg(test)] +mod range_tests; + use paimon_vindex_core::distance::MetricType; use paimon_vindex_core::index::{ IvfPqBatchTableReuseMode, SearchWidth, VectorIndexConfig, VectorIndexMetadata, diff --git a/ffi/src/range.rs b/ffi/src/range.rs new file mode 100644 index 00000000..c8928e08 --- /dev/null +++ b/ffi/src/range.rs @@ -0,0 +1,425 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +use super::*; +use paimon_vindex_core::range::{ + Bound, CutOperator, DistanceBand, DistanceEndpoint, RangeSearchResult, VectorRangeSearchParams, +}; + +pub const PAIMON_VINDEX_BOUND_UNBOUNDED: u32 = 0; +pub const PAIMON_VINDEX_BOUND_FINITE: u32 = 1; +pub const PAIMON_VINDEX_CUT_GE: u32 = 0; +pub const PAIMON_VINDEX_CUT_GT: u32 = 1; +pub const PAIMON_VINDEX_CUT_LE: u32 = 2; +pub const PAIMON_VINDEX_CUT_LT: u32 = 3; + +/// Internal half-open distance band: squared L2, 1-cosine, or negative inner product. +/// Unbounded sides ignore their value field; finite sides must be finite. +#[repr(C)] +#[derive(Clone, Copy)] +pub struct PaimonVindexDistanceBand { + pub metric: u32, + pub lower_kind: u32, + pub lower: f32, + pub upper_kind: u32, + pub upper: f32, +} + +/// Public-distance predicate endpoint. Values are passed to core without rounding. +/// Lower endpoints accept GE/GT; upper endpoints accept LE/LT. +#[repr(C)] +#[derive(Clone, Copy)] +pub struct PaimonVindexDistanceEndpoint { + pub value: f64, + pub op: u32, +} + +/// Fixed IVF probe count, greater than zero and clamped by core to nlist. +#[repr(C)] +#[derive(Clone, Copy)] +pub struct PaimonVindexRangeSearchParams { + pub band: PaimonVindexDistanceBand, + pub nprobe: usize, +} + +#[repr(C)] +#[derive(Clone, Copy)] +pub struct PaimonVindexRangeSearchStats { + pub lists_probed: usize, + pub rows_scanned: usize, + pub rows_committed: usize, + pub early_abandoned: usize, +} + +/// Borrowed CSR view; all pointers remain valid until the owning result is destroyed. +/// lims has query_count + 1 entries; labels/distances have hit_count entries; +/// stats has query_count entries. Empty arrays must not be dereferenced. +#[repr(C)] +pub struct PaimonVindexRangeSearchResultView { + pub query_count: usize, + pub hit_count: usize, + pub lims: *const usize, + pub labels: *const i64, + pub distances: *const f32, + pub stats: *const PaimonVindexRangeSearchStats, + pub list_reads: usize, +} + +/// Owns all range result buffers independently of the reader and input buffers. +pub struct PaimonVindexRangeSearchResult { + inner: RangeSearchResult, + stats: Vec, +} + +fn range_metric(code: u32) -> Result { + match code { + PAIMON_VINDEX_METRIC_L2 => Ok(MetricType::L2), + PAIMON_VINDEX_METRIC_COSINE => Ok(MetricType::Cosine), + PAIMON_VINDEX_METRIC_INNER_PRODUCT => Ok(MetricType::InnerProduct), + _ => Err(format!("invalid range metric: {code}")), + } +} + +fn range_bound(kind: u32, value: f32) -> Result { + match kind { + PAIMON_VINDEX_BOUND_UNBOUNDED => Ok(Bound::Unbounded), + PAIMON_VINDEX_BOUND_FINITE => Ok(Bound::Finite(value)), + _ => Err(format!("invalid bound kind: {kind}")), + } +} + +fn range_params(params: PaimonVindexRangeSearchParams) -> Result { + let band = DistanceBand::new( + range_bound(params.band.lower_kind, params.band.lower)?, + range_bound(params.band.upper_kind, params.band.upper)?, + range_metric(params.band.metric)?, + ) + .map_err(|error| format!("range band: {error}"))?; + Ok(VectorRangeSearchParams::new(band, params.nprobe)) +} + +unsafe fn range_endpoint( + endpoint: *const PaimonVindexDistanceEndpoint, +) -> Result, String> { + if endpoint.is_null() { + return Ok(None); + } + let endpoint = unsafe { &*endpoint }; + let op = match endpoint.op { + PAIMON_VINDEX_CUT_GE => CutOperator::Ge, + PAIMON_VINDEX_CUT_GT => CutOperator::Gt, + PAIMON_VINDEX_CUT_LE => CutOperator::Le, + PAIMON_VINDEX_CUT_LT => CutOperator::Lt, + code => return Err(format!("invalid endpoint operator: {code}")), + }; + Ok(Some(DistanceEndpoint { + value: endpoint.value, + op, + })) +} + +fn range_band_to_ffi(band: DistanceBand) -> PaimonVindexDistanceBand { + let encode = |bound| match bound { + Bound::Unbounded => (PAIMON_VINDEX_BOUND_UNBOUNDED, 0.0), + Bound::Finite(value) => (PAIMON_VINDEX_BOUND_FINITE, value), + }; + let (lower_kind, lower) = encode(band.lower()); + let (upper_kind, upper) = encode(band.upper()); + PaimonVindexDistanceBand { + metric: metric_code(band.metric()), + lower_kind, + lower, + upper_kind, + upper, + } +} + +/// Converts public-distance endpoints through core, including f64 boundary rounding +/// and inner-product side reversal. NULL endpoints mean structurally unbounded. +/// Returns zero on success or -1 with last_error set; out is unchanged on error. +#[no_mangle] +pub unsafe extern "C" fn paimon_vindex_distance_band_from_endpoints( + metric: u32, + lower: *const PaimonVindexDistanceEndpoint, + upper: *const PaimonVindexDistanceEndpoint, + out: *mut PaimonVindexDistanceBand, +) -> c_int { + ffi_status(|| { + if out.is_null() { + return Err("out band pointer is null".to_string()); + } + let band = DistanceBand::from_endpoints( + unsafe { range_endpoint(lower) }?, + unsafe { range_endpoint(upper) }?, + range_metric(metric)?, + ) + .map_err(|error| format!("range endpoints: {error}"))?; + unsafe { *out = range_band_to_ffi(band) }; + Ok(()) + }) +} + +/// Queries core capability without issuing a search; out receives zero or one. +#[no_mangle] +pub unsafe extern "C" fn paimon_vindex_reader_supports_range_search( + handle: *const PaimonVindexReaderHandle, + out: *mut c_int, +) -> c_int { + ffi_status(|| { + if out.is_null() { + return Err("out capability pointer is null".to_string()); + } + let handle = unsafe { reader_ref(handle) }?; + unsafe { *out = c_int::from(handle.inner.supports_range_search()) }; + Ok(()) + }) +} + +fn range_len(len: usize, name: &str) -> Result<(), String> { + if len + .checked_mul(size_of::()) + .is_none_or(|bytes| bytes > isize::MAX as usize) + { + return Err(format!("{name} byte length overflow")); + } + Ok(()) +} + +unsafe fn range_slice<'a, T>(data: *const T, len: usize, name: &str) -> Result<&'a [T], String> { + range_len::(len, name)?; + if len > 0 && !data.is_aligned() { + return Err(format!("{name} pointer is not aligned")); + } + unsafe { const_slice(data, len, name) } +} + +impl PaimonVindexRangeSearchResult { + fn new(inner: RangeSearchResult) -> Result { + let mut stats = Vec::new(); + stats + .try_reserve_exact(inner.query_count()) + .map_err(|error| format!("range result allocation: {error}"))?; + for query in 0..inner.query_count() { + let counters = inner.query(query).stats; + stats.push(PaimonVindexRangeSearchStats { + lists_probed: counters.lists_probed(), + rows_scanned: counters.rows_scanned(), + rows_committed: counters.rows_committed(), + early_abandoned: counters.early_abandoned(), + }); + } + Ok(Self { inner, stats }) + } +} + +#[allow(clippy::too_many_arguments)] +unsafe fn range_search( + handle: *mut PaimonVindexReaderHandle, + queries: *const f32, + queries_len: usize, + query_count: usize, + params: PaimonVindexRangeSearchParams, + filter: Option<(*const u8, usize)>, + batch: bool, + out: *mut *mut PaimonVindexRangeSearchResult, +) -> c_int { + ffi_status(|| { + if out.is_null() { + return Err("out result pointer is null".to_string()); + } + unsafe { *out = ptr::null_mut() }; + let handle = unsafe { reader_mut(handle) }?; + let expected_len = checked_len(query_count, handle.inner.dimension(), "range queries")?; + if queries_len != expected_len { + return Err(format!( + "range query length {queries_len} does not match required {expected_len}" + )); + } + let lims_len = query_count + .checked_add(1) + .ok_or_else(|| "range offsets length overflow".to_string())?; + range_len::(lims_len, "range offsets")?; + range_len::(query_count, "range statistics")?; + let queries = unsafe { range_slice(queries, queries_len, "range queries") }?; + let params = range_params(params)?; + let filter = filter + .map(|(data, len)| unsafe { range_slice(data, len, "range filter") }) + .transpose()?; + let result = match (batch, filter) { + (false, None) => handle.inner.range_search(queries, params), + (false, Some(filter)) => handle + .inner + .range_search_with_roaring_filter(queries, params, filter), + (true, None) => handle + .inner + .range_search_batch(queries, query_count, params), + (true, Some(filter)) => handle.inner.range_search_batch_with_roaring_filter( + queries, + query_count, + params, + filter, + ), + } + .map_err(|error| format!("range search: {error}"))?; + let result = Box::new(PaimonVindexRangeSearchResult::new(result)?); + unsafe { *out = Box::into_raw(result) }; + Ok(()) + }) +} + +/// Returns an owned variable-length result; no top-K limit or padding is applied. +/// Inputs are borrowed for this call only. A valid out pointer is set to NULL +/// before any work; on failure last_error is set and no result is transferred. +/// Query length is in floats and must equal the reader dimension. +#[no_mangle] +pub unsafe extern "C" fn paimon_vindex_reader_range_search( + handle: *mut PaimonVindexReaderHandle, + query: *const f32, + query_len: usize, + params: PaimonVindexRangeSearchParams, + out: *mut *mut PaimonVindexRangeSearchResult, +) -> c_int { + unsafe { range_search(handle, query, query_len, 1, params, None, false, out) } +} + +/// Filtered single-query variant. The filter must be a serialized Roaring +/// bitmap/treemap; zero bytes are malformed, not equivalent to no filter. +#[no_mangle] +pub unsafe extern "C" fn paimon_vindex_reader_range_search_with_roaring_filter( + handle: *mut PaimonVindexReaderHandle, + query: *const f32, + query_len: usize, + params: PaimonVindexRangeSearchParams, + filter: *const u8, + filter_len: usize, + out: *mut *mut PaimonVindexRangeSearchResult, +) -> c_int { + unsafe { + range_search( + handle, + query, + query_len, + 1, + params, + Some((filter, filter_len)), + false, + out, + ) + } +} + +/// Batched CSR variant. queries_len must equal query_count * dimension. +/// The query count must be positive, as required by core. +#[no_mangle] +pub unsafe extern "C" fn paimon_vindex_reader_range_search_batch( + handle: *mut PaimonVindexReaderHandle, + queries: *const f32, + queries_len: usize, + query_count: usize, + params: PaimonVindexRangeSearchParams, + out: *mut *mut PaimonVindexRangeSearchResult, +) -> c_int { + unsafe { + range_search( + handle, + queries, + queries_len, + query_count, + params, + None, + true, + out, + ) + } +} + +/// Batched CSR variant sharing one serialized Roaring allow-list across queries. +#[no_mangle] +pub unsafe extern "C" fn paimon_vindex_reader_range_search_batch_with_roaring_filter( + handle: *mut PaimonVindexReaderHandle, + queries: *const f32, + queries_len: usize, + query_count: usize, + params: PaimonVindexRangeSearchParams, + filter: *const u8, + filter_len: usize, + out: *mut *mut PaimonVindexRangeSearchResult, +) -> c_int { + unsafe { + range_search( + handle, + queries, + queries_len, + query_count, + params, + Some((filter, filter_len)), + true, + out, + ) + } +} + +/// Obtains a borrowed view without transferring ownership or copying buffers. +#[no_mangle] +pub unsafe extern "C" fn paimon_vindex_range_search_result_view( + result: *const PaimonVindexRangeSearchResult, + out: *mut PaimonVindexRangeSearchResultView, +) -> c_int { + ffi_status(|| { + if result.is_null() || out.is_null() { + return Err("range result or view pointer is null".to_string()); + } + let result = unsafe { &*result }; + unsafe { + *out = PaimonVindexRangeSearchResultView { + query_count: result.inner.query_count(), + hit_count: result.inner.labels().len(), + lims: result.inner.lims().as_ptr(), + labels: result.inner.labels().as_ptr(), + distances: result.inner.distances().as_ptr(), + stats: result.stats.as_ptr(), + list_reads: result.inner.call_stats().list_reads(), + } + }; + Ok(()) + }) +} + +/// Destroys a result exactly once, invalidating every borrowed view. NULL is a no-op. +#[no_mangle] +pub unsafe extern "C" fn paimon_vindex_range_search_result_destroy( + result: *mut PaimonVindexRangeSearchResult, +) { + if !result.is_null() { + unsafe { drop(Box::from_raw(result)) }; + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn range_query_byte_limits_are_checked_before_slice_construction() { + let maximum = isize::MAX as usize / size_of::(); + assert!(range_len::(maximum, "queries").is_ok()); + for len in [maximum + 1, usize::MAX] { + let error = unsafe { range_slice::(ptr::dangling(), len, "queries") }.unwrap_err(); + assert_eq!(error, "queries byte length overflow"); + } + } +} diff --git a/ffi/src/range_tests.rs b/ffi/src/range_tests.rs new file mode 100644 index 00000000..272f1bb5 --- /dev/null +++ b/ffi/src/range_tests.rs @@ -0,0 +1,483 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +use super::*; +use paimon_vindex_core::range::{Bound, CutOperator, DistanceBand, DistanceEndpoint}; + +#[test] +fn range_endpoints_match_core_for_every_metric_and_operator() { + for (metric_code, metric) in [ + (PAIMON_VINDEX_METRIC_L2, MetricType::L2), + (PAIMON_VINDEX_METRIC_COSINE, MetricType::Cosine), + (PAIMON_VINDEX_METRIC_INNER_PRODUCT, MetricType::InnerProduct), + ] { + for (lower_code, lower_op) in [ + (PAIMON_VINDEX_CUT_GE, CutOperator::Ge), + (PAIMON_VINDEX_CUT_GT, CutOperator::Gt), + ] { + for (upper_code, upper_op) in [ + (PAIMON_VINDEX_CUT_LE, CutOperator::Le), + (PAIMON_VINDEX_CUT_LT, CutOperator::Lt), + ] { + for (lower_value, upper_value) in [ + (0.1, 1.1), + (1.0, 2.0), + (1.0_f64.next_down(), 2.0_f64.next_up()), + (1.0_f64.next_up(), 2.0_f64.next_down()), + ] { + let lower = PaimonVindexDistanceEndpoint { + value: lower_value, + op: lower_code, + }; + let upper = PaimonVindexDistanceEndpoint { + value: upper_value, + op: upper_code, + }; + let expected = DistanceBand::from_endpoints( + Some(DistanceEndpoint { + value: lower.value, + op: lower_op, + }), + Some(DistanceEndpoint { + value: upper.value, + op: upper_op, + }), + metric, + ) + .unwrap(); + let mut actual = std::mem::MaybeUninit::uninit(); + assert_eq!( + unsafe { + paimon_vindex_distance_band_from_endpoints( + metric_code, + &lower, + &upper, + actual.as_mut_ptr(), + ) + }, + 0 + ); + let actual = unsafe { actual.assume_init() }; + assert_eq!(actual.metric, metric_code); + for (kind, value, expected) in [ + (actual.lower_kind, actual.lower, expected.lower()), + (actual.upper_kind, actual.upper, expected.upper()), + ] { + match expected { + Bound::Unbounded => assert_eq!(kind, PAIMON_VINDEX_BOUND_UNBOUNDED), + Bound::Finite(expected) => { + assert_eq!(kind, PAIMON_VINDEX_BOUND_FINITE); + assert_eq!(value.to_bits(), expected.to_bits()); + } + } + } + } + } + } + } +} + +#[test] +fn range_errors_leave_no_owned_result() { + let params = PaimonVindexRangeSearchParams { + band: PaimonVindexDistanceBand { + metric: 0, + lower_kind: 0, + lower: 0.0, + upper_kind: 0, + upper: 0.0, + }, + nprobe: 1, + }; + let mut result = ptr::dangling_mut::(); + assert_eq!( + unsafe { + paimon_vindex_reader_range_search(ptr::null_mut(), ptr::null(), 0, params, &mut result) + }, + -1 + ); + assert!(result.is_null()); + unsafe { paimon_vindex_range_search_result_destroy(result) }; +} + +struct RangeInput { + bytes: Vec, + fail_reads: bool, +} + +struct RangeReader { + handle: *mut PaimonVindexReaderHandle, + input: Box, +} + +impl RangeReader { + fn new() -> Self { + let config = VectorIndexConfig::from_options(&HashMap::from([ + ("index.type".to_string(), "ivf_flat".to_string()), + ("dimension".to_string(), "2".to_string()), + ("nlist".to_string(), "1".to_string()), + ("metric".to_string(), "l2".to_string()), + ])) + .unwrap(); + let vectors = [1.0, 0.0, 0.0, 1.0, -1.0, 0.0, 0.0, -1.0]; + let training = VectorIndexTrainer::train(config, &vectors, 4).unwrap(); + let mut writer = VectorIndexWriter::new(training); + writer + .add_vectors(&[1, 2, i64::MAX, -1099511627776], &vectors, 4) + .unwrap(); + let mut bytes = Vec::new(); + writer + .write(&mut paimon_vindex_core::io::PosWriter::new(&mut bytes)) + .unwrap(); + let mut input = Box::new(RangeInput { + bytes, + fail_reads: false, + }); + let handle = unsafe { + paimon_vindex_reader_open(PaimonVindexInputFile { + ctx: (&mut *input as *mut RangeInput).cast(), + read_ranges_fn: Some(range_read), + estimated_random_read_latency_nanos: 0, + preferred_window_bytes: 0, + max_ranges_per_read: 0, + }) + }; + assert!(!handle.is_null()); + Self { handle, input } + } +} + +impl Drop for RangeReader { + fn drop(&mut self) { + unsafe { paimon_vindex_reader_free(self.handle) }; + } +} + +unsafe extern "C" fn range_read( + ctx: *mut c_void, + requests: *mut PaimonVindexReadRequest, + count: usize, +) -> c_int { + let input = unsafe { &*ctx.cast::() }; + if input.fail_reads { + return -1; + } + for request in unsafe { slice::from_raw_parts_mut(requests, count) } { + let Ok(start) = usize::try_from(request.offset) else { + return -1; + }; + let Some(end) = start.checked_add(request.len) else { + return -1; + }; + let Some(source) = input.bytes.get(start..end) else { + return -1; + }; + unsafe { slice::from_raw_parts_mut(request.buf, request.len) }.copy_from_slice(source); + } + 0 +} + +fn unbounded_params() -> PaimonVindexRangeSearchParams { + PaimonVindexRangeSearchParams { + band: PaimonVindexDistanceBand { + metric: 0, + lower_kind: 0, + lower: f32::NAN, + upper_kind: 0, + upper: f32::NAN, + }, + nprobe: 1, + } +} + +#[test] +fn range_result_outlives_reader_and_preserves_signed_labels_and_statistics() { + let reader = RangeReader::new(); + let mut supported = 0; + assert_eq!( + unsafe { paimon_vindex_reader_supports_range_search(reader.handle, &mut supported) }, + 0 + ); + assert_eq!(supported, 1); + let queries = [1.0, 0.0, 0.0, 1.0]; + let mut result = ptr::null_mut(); + assert_eq!( + unsafe { + paimon_vindex_reader_range_search_batch( + reader.handle, + queries.as_ptr(), + queries.len(), + 2, + unbounded_params(), + &mut result, + ) + }, + 0 + ); + let mut expected_reader = + VectorIndexReader::open(std::io::Cursor::new(reader.input.bytes.clone())).unwrap(); + let expected = expected_reader + .range_search_batch( + &queries, + 2, + paimon_vindex_core::range::VectorRangeSearchParams::new( + DistanceBand::new(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(), + 1, + ), + ) + .unwrap(); + drop(reader); + let mut view = std::mem::MaybeUninit::uninit(); + assert_eq!( + unsafe { paimon_vindex_range_search_result_view(result, view.as_mut_ptr()) }, + 0 + ); + let view = unsafe { view.assume_init() }; + assert_eq!(view.query_count, 2); + assert_eq!(view.hit_count, 8); + assert_eq!( + unsafe { slice::from_raw_parts(view.lims, 3) }, + expected.lims() + ); + assert_eq!( + unsafe { slice::from_raw_parts(view.labels, view.hit_count) }, + expected.labels() + ); + assert_eq!( + unsafe { slice::from_raw_parts(view.distances, view.hit_count) }, + expected.distances() + ); + assert_eq!(view.list_reads, expected.call_stats().list_reads()); + for (query, actual) in unsafe { slice::from_raw_parts(view.stats, 2) } + .iter() + .enumerate() + { + let stats = expected.query(query).stats; + assert_eq!(actual.lists_probed, stats.lists_probed()); + assert_eq!(actual.rows_scanned, stats.rows_scanned()); + assert_eq!(actual.rows_committed, stats.rows_committed()); + assert_eq!(actual.early_abandoned, stats.early_abandoned()); + } + unsafe { paimon_vindex_range_search_result_destroy(result) }; +} + +#[test] +fn range_rejects_invalid_lengths_before_dereferencing_inputs() { + let reader = RangeReader::new(); + let mut result = ptr::null_mut(); + for (data, len, count) in [ + (ptr::null(), 2, 1), + (ptr::null(), 0, 0), + (ptr::dangling(), usize::MAX, usize::MAX), + (ptr::dangling(), usize::MAX - 1, usize::MAX / 2), + (ptr::dangling(), 1, 1), + (ptr::without_provenance(1), 2, 1), + ] { + assert_eq!( + unsafe { + paimon_vindex_reader_range_search_batch( + reader.handle, + data, + len, + count, + unbounded_params(), + &mut result, + ) + }, + -1 + ); + assert!(result.is_null()); + } + let query = [1.0, 0.0]; + assert_eq!( + unsafe { + paimon_vindex_reader_range_search_with_roaring_filter( + reader.handle, + query.as_ptr(), + 2, + unbounded_params(), + ptr::dangling(), + usize::MAX, + &mut result, + ) + }, + -1 + ); + assert!(result.is_null()); + assert_eq!( + unsafe { + paimon_vindex_reader_range_search( + reader.handle, + query.as_ptr(), + 2, + unbounded_params(), + ptr::null_mut(), + ) + }, + -1 + ); +} + +#[test] +fn range_validation_and_io_errors_do_not_transfer_results() { + let mut reader = RangeReader::new(); + let query = [1.0, 0.0]; + let mut result = ptr::null_mut(); + let mut params = unbounded_params(); + params.nprobe = 0; + assert_eq!( + unsafe { + paimon_vindex_reader_range_search(reader.handle, query.as_ptr(), 2, params, &mut result) + }, + -1 + ); + params = unbounded_params(); + for metric in [1, 2, u32::MAX] { + params.band.metric = metric; + assert_eq!( + unsafe { + paimon_vindex_reader_range_search( + reader.handle, + query.as_ptr(), + 2, + params, + &mut result, + ) + }, + -1 + ); + } + params = unbounded_params(); + for kind in [2, u32::MAX] { + params.band.lower_kind = kind; + assert_eq!( + unsafe { + paimon_vindex_reader_range_search( + reader.handle, + query.as_ptr(), + 2, + params, + &mut result, + ) + }, + -1 + ); + } + params = unbounded_params(); + params.band.lower_kind = PAIMON_VINDEX_BOUND_FINITE; + for value in [f32::NAN, f32::INFINITY, -1.0] { + params.band.lower = value; + assert_eq!( + unsafe { + paimon_vindex_reader_range_search( + reader.handle, + query.as_ptr(), + 2, + params, + &mut result, + ) + }, + -1 + ); + } + assert_eq!( + unsafe { + paimon_vindex_reader_range_search_with_roaring_filter( + reader.handle, + query.as_ptr(), + 2, + unbounded_params(), + ptr::null(), + 0, + &mut result, + ) + }, + -1 + ); + assert_eq!( + unsafe { + paimon_vindex_reader_range_search( + reader.handle, + [f32::NAN, 0.0].as_ptr(), + 2, + unbounded_params(), + &mut result, + ) + }, + -1 + ); + reader.input.fail_reads = true; + assert_eq!( + unsafe { + paimon_vindex_reader_range_search( + reader.handle, + query.as_ptr(), + 2, + unbounded_params(), + &mut result, + ) + }, + -1 + ); + assert!(result.is_null()); + assert!(!paimon_vindex_last_error().is_null()); +} + +#[test] +fn range_endpoint_errors_leave_output_unchanged() { + let mut output = unbounded_params().band; + output.metric = 77; + for endpoint in [ + PaimonVindexDistanceEndpoint { + value: f64::NAN, + op: PAIMON_VINDEX_CUT_GE, + }, + PaimonVindexDistanceEndpoint { + value: 1.0, + op: PAIMON_VINDEX_CUT_LE, + }, + PaimonVindexDistanceEndpoint { value: 1.0, op: 99 }, + PaimonVindexDistanceEndpoint { + value: 1.0e30, + op: PAIMON_VINDEX_CUT_GE, + }, + ] { + assert_eq!( + unsafe { + paimon_vindex_distance_band_from_endpoints(0, &endpoint, ptr::null(), &mut output) + }, + -1 + ); + assert_eq!(output.metric, 77); + } + assert_eq!( + unsafe { + paimon_vindex_distance_band_from_endpoints(77, ptr::null(), ptr::null(), &mut output) + }, + -1 + ); + assert_eq!( + unsafe { + paimon_vindex_distance_band_from_endpoints(0, ptr::null(), ptr::null(), ptr::null_mut()) + }, + -1 + ); + assert_eq!( + unsafe { paimon_vindex_range_search_result_view(ptr::null(), ptr::null_mut()) }, + -1 + ); +} diff --git a/include/paimon_vindex.hpp b/include/paimon_vindex.hpp index f3638bbd..9cac0ef0 100644 --- a/include/paimon_vindex.hpp +++ b/include/paimon_vindex.hpp @@ -25,8 +25,10 @@ extern "C" { #include #include +#include #include #include +#include #include #include #include @@ -224,6 +226,79 @@ struct SearchResult { std::vector distances; }; +using DistanceEndpoint = PaimonVindexDistanceEndpoint; +using RangeSearchStats = PaimonVindexRangeSearchStats; + +struct DistanceBand { + uint32_t metric = PAIMON_VINDEX_METRIC_L2; + uint32_t lower_kind = PAIMON_VINDEX_BOUND_UNBOUNDED; + float lower = 0.0f; + uint32_t upper_kind = PAIMON_VINDEX_BOUND_UNBOUNDED; + float upper = 0.0f; + + static DistanceBand from_endpoints( + uint32_t metric, + std::optional lower = std::nullopt, + std::optional upper = std::nullopt) { + PaimonVindexDistanceBand raw{}; + check(paimon_vindex_distance_band_from_endpoints( + metric, lower ? &*lower : nullptr, upper ? &*upper : nullptr, &raw)); + return {raw.metric, raw.lower_kind, raw.lower, raw.upper_kind, raw.upper}; + } + + PaimonVindexDistanceBand to_ffi() const { + return {metric, lower_kind, lower, upper_kind, upper}; + } +}; + +struct RangeSearchParams { + DistanceBand band; + size_t nprobe = 1; + + PaimonVindexRangeSearchParams to_ffi() const { + return {band.to_ffi(), nprobe}; + } +}; + +struct RangeSearchResult { + size_t query_count = 0; + std::vector lims; + std::vector labels; + std::vector distances; + std::vector stats; + size_t list_reads = 0; +}; + +namespace detail { + +inline RangeSearchResult copy_range_result(PaimonVindexRangeSearchResult* raw, int status) { + std::unique_ptr guard( + raw, &paimon_vindex_range_search_result_destroy); + check(status); + PaimonVindexRangeSearchResultView view{}; + check(paimon_vindex_range_search_result_view(guard.get(), &view)); + if (view.query_count == std::numeric_limits::max() || !view.lims || + (view.hit_count != 0 && (!view.labels || !view.distances)) || + (view.query_count != 0 && !view.stats)) { + throw Error("invalid native range result view"); + } + RangeSearchResult result; + result.query_count = view.query_count; + result.lims.assign(view.lims, view.lims + view.query_count + 1); + if (view.hit_count != 0) { + result.labels.assign(view.labels, view.labels + view.hit_count); + result.distances.assign(view.distances, view.distances + view.hit_count); + } + if (view.query_count != 0) { + result.stats.assign(view.stats, view.stats + view.query_count); + } + result.list_reads = view.list_reads; + return result; +} + +} + struct SearchParams { size_t top_k = 0; uint32_t search_width = PAIMON_VINDEX_SEARCH_WIDTH_AUTO; @@ -584,6 +659,77 @@ class Reader { return result; } + bool supports_range_search() const { + std::lock_guard lock(native_handle_mutex_); + int supported = 0; + check(paimon_vindex_reader_supports_range_search(require_open(), &supported)); + return supported != 0; + } + + RangeSearchResult range_search( + const float* query, size_t query_len, RangeSearchParams params) { + std::lock_guard lock(native_handle_mutex_); + validate_range_queries(query, query_len, 1); + PaimonVindexRangeSearchResult* raw = nullptr; + int status = paimon_vindex_reader_range_search( + require_open(), query, query_len, params.to_ffi(), &raw); + return detail::copy_range_result(raw, status); + } + + RangeSearchResult range_search(const std::vector& query, RangeSearchParams params) { + return range_search(query.data(), query.size(), params); + } + + RangeSearchResult range_search_with_roaring_filter( + const float* query, size_t query_len, RangeSearchParams params, + const uint8_t* filter, size_t filter_len) { + std::lock_guard lock(native_handle_mutex_); + validate_range_queries(query, query_len, 1); + PaimonVindexRangeSearchResult* raw = nullptr; + int status = paimon_vindex_reader_range_search_with_roaring_filter( + require_open(), query, query_len, params.to_ffi(), filter, filter_len, &raw); + return detail::copy_range_result(raw, status); + } + + RangeSearchResult range_search_with_roaring_filter( + const std::vector& query, RangeSearchParams params, + const uint8_t* filter, size_t filter_len) { + return range_search_with_roaring_filter(query.data(), query.size(), params, filter, filter_len); + } + + RangeSearchResult range_search_batch( + const float* queries, size_t queries_len, size_t query_count, RangeSearchParams params) { + std::lock_guard lock(native_handle_mutex_); + validate_range_queries(queries, queries_len, query_count); + PaimonVindexRangeSearchResult* raw = nullptr; + int status = paimon_vindex_reader_range_search_batch( + require_open(), queries, queries_len, query_count, params.to_ffi(), &raw); + return detail::copy_range_result(raw, status); + } + + RangeSearchResult range_search_batch( + const std::vector& queries, size_t query_count, RangeSearchParams params) { + return range_search_batch(queries.data(), queries.size(), query_count, params); + } + + RangeSearchResult range_search_batch_with_roaring_filter( + const float* queries, size_t queries_len, size_t query_count, RangeSearchParams params, + const uint8_t* filter, size_t filter_len) { + std::lock_guard lock(native_handle_mutex_); + validate_range_queries(queries, queries_len, query_count); + PaimonVindexRangeSearchResult* raw = nullptr; + int status = paimon_vindex_reader_range_search_batch_with_roaring_filter( + require_open(), queries, queries_len, query_count, params.to_ffi(), filter, filter_len, &raw); + return detail::copy_range_result(raw, status); + } + + RangeSearchResult range_search_batch_with_roaring_filter( + const std::vector& queries, size_t query_count, RangeSearchParams params, + const uint8_t* filter, size_t filter_len) { + return range_search_batch_with_roaring_filter( + queries.data(), queries.size(), query_count, params, filter, filter_len); + } + SearchResult search(const float* query, SearchParams params) { std::lock_guard lock(native_handle_mutex_); SearchResult result; @@ -669,6 +815,21 @@ class Reader { } private: + void validate_range_queries(const float* queries, size_t queries_len, size_t query_count) const { + PaimonVindexMetadata metadata{}; + check(paimon_vindex_reader_metadata(require_open(), &metadata)); + if (metadata.dimension != 0 && + query_count > std::numeric_limits::max() / metadata.dimension) { + throw Error("range query dimensions overflow"); + } + if (queries_len != query_count * metadata.dimension) { + throw Error("range query length does not match dimension and query count"); + } + if (queries_len != 0 && !queries) { + throw Error("range queries must not be null for a nonempty input"); + } + } + PaimonVindexReaderHandle* require_open() const { if (!handle_) throw Error("vector index reader is closed"); return handle_; diff --git a/java/src/main/java/org/apache/paimon/index/vector/VectorDistanceBand.java b/java/src/main/java/org/apache/paimon/index/vector/VectorDistanceBand.java new file mode 100644 index 00000000..ad3c4276 --- /dev/null +++ b/java/src/main/java/org/apache/paimon/index/vector/VectorDistanceBand.java @@ -0,0 +1,104 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package org.apache.paimon.index.vector; + +import java.util.Objects; + +/** + * A half-open band [lower, upper) in raw index distance space: squared L2, 1 - cosine, or negative + * inner product. A null cut is structurally unbounded, not a finite sentinel. + */ +public final class VectorDistanceBand { + + public enum CutOperator { + GE(0), + GT(1), + LE(2), + LT(3); + + private final int code; + + CutOperator(int code) { + this.code = code; + } + } + + private final String metric; + private final Float lower; + private final Float upper; + + public VectorDistanceBand(String metric, Float lower, Float upper) { + this.metric = Objects.requireNonNull(metric, "metric"); + if (!"l2".equals(metric) && !"cosine".equals(metric) && !"inner_product".equals(metric)) { + throw new IllegalArgumentException("unknown metric: " + metric); + } + validateCut(lower); + validateCut(upper); + if (lower != null && upper != null && lower > upper) { + throw new IllegalArgumentException("inverted distance band"); + } + this.lower = lower; + this.upper = upper; + } + + /** + * Converts public predicate endpoints through the core's exact f64-to-f32 cut conversion. + * Public L2 is the f32 square root of stored squared distance; cosine is 1 - cosine; inner + * product is the positive dot product. Null endpoint/operator pairs are unbounded. Lower + * accepts GE/GT and upper accepts LE/LT. Unrepresentable L2 cuts fail; out-of-domain linear + * endpoints produce empty or structurally unbounded bands as defined by core. + */ + public static VectorDistanceBand fromEndpoints( + String metric, + Double lower, + CutOperator lowerOperator, + Double upper, + CutOperator upperOperator) { + Objects.requireNonNull(metric, "metric"); + if ((lower == null) != (lowerOperator == null) + || (upper == null) != (upperOperator == null)) { + throw new IllegalArgumentException( + "endpoint value and operator must both be null or set"); + } + return VectorIndexNative.distanceBandFromEndpoints( + metric, + lower, + lowerOperator == null ? -1 : lowerOperator.code, + upper, + upperOperator == null ? -1 : upperOperator.code); + } + + public String metric() { + return metric; + } + + public Float lower() { + return lower; + } + + public Float upper() { + return upper; + } + + private void validateCut(Float cut) { + if (cut != null && (!Float.isFinite(cut) || ("l2".equals(metric) && cut < 0))) { + throw new IllegalArgumentException( + "cut must be finite and non-negative for squared L2"); + } + } +} diff --git a/java/src/main/java/org/apache/paimon/index/vector/VectorIndexNative.java b/java/src/main/java/org/apache/paimon/index/vector/VectorIndexNative.java index 7d48be37..ceffd290 100644 --- a/java/src/main/java/org/apache/paimon/index/vector/VectorIndexNative.java +++ b/java/src/main/java/org/apache/paimon/index/vector/VectorIndexNative.java @@ -77,4 +77,25 @@ static native VectorSearchBatchResult searchBatchWithRoaringFilter( byte[] roaringFilter); static native void freeReader(long ptr); + + static native VectorDistanceBand distanceBandFromEndpoints( + String metric, Double lower, int lowerOperator, Double upper, int upperOperator); + + static native boolean supportsRangeSearch(long ptr); + + static native VectorRangeSearchResult rangeSearch( + long ptr, float[] query, VectorRangeSearchParams params); + + static native VectorRangeSearchResult rangeSearchWithRoaringFilter( + long ptr, float[] query, VectorRangeSearchParams params, byte[] roaringFilter); + + static native VectorRangeSearchResult rangeSearchBatch( + long ptr, float[] queries, int queryCount, VectorRangeSearchParams params); + + static native VectorRangeSearchResult rangeSearchBatchWithRoaringFilter( + long ptr, + float[] queries, + int queryCount, + VectorRangeSearchParams params, + byte[] roaringFilter); } diff --git a/java/src/main/java/org/apache/paimon/index/vector/VectorIndexReader.java b/java/src/main/java/org/apache/paimon/index/vector/VectorIndexReader.java index 0dee5371..d7ba6c6d 100644 --- a/java/src/main/java/org/apache/paimon/index/vector/VectorIndexReader.java +++ b/java/src/main/java/org/apache/paimon/index/vector/VectorIndexReader.java @@ -212,6 +212,97 @@ public VectorSearchBatchResult searchBatch( } } + public boolean supportsRangeSearch() { + rejectCallbackReentry(); + synchronized (nativeHandleLock) { + enterNativeHandle(); + try { + return VectorIndexNative.supportsRangeSearch(requireOpen()); + } finally { + exitNativeHandle(); + } + } + } + + public VectorRangeSearchResult rangeSearch(float[] query, VectorRangeSearchParams params) { + validateQuery(query); + if (params == null) { + throw new NullPointerException("params"); + } + rejectCallbackReentry(); + synchronized (nativeHandleLock) { + enterNativeHandle(); + try { + return VectorIndexNative.rangeSearch(requireOpen(), query, params); + } finally { + exitNativeHandle(); + } + } + } + + public VectorRangeSearchResult rangeSearch( + float[] query, VectorRangeSearchParams params, byte[] roaringFilter) { + validateQuery(query); + if (params == null) { + throw new NullPointerException("params"); + } + if (roaringFilter == null) { + throw new NullPointerException("roaringFilter"); + } + rejectCallbackReentry(); + synchronized (nativeHandleLock) { + enterNativeHandle(); + try { + return VectorIndexNative.rangeSearchWithRoaringFilter( + requireOpen(), query, params, roaringFilter); + } finally { + exitNativeHandle(); + } + } + } + + public VectorRangeSearchResult rangeSearchBatch( + float[] queries, int queryCount, VectorRangeSearchParams params) { + if (queries == null) { + throw new NullPointerException("queries"); + } + if (params == null) { + throw new NullPointerException("params"); + } + rejectCallbackReentry(); + synchronized (nativeHandleLock) { + enterNativeHandle(); + try { + return VectorIndexNative.rangeSearchBatch(requireOpen(), queries, queryCount, params); + } finally { + exitNativeHandle(); + } + } + } + + public VectorRangeSearchResult rangeSearchBatch( + float[] queries, int queryCount, VectorRangeSearchParams params, byte[] roaringFilter) { + if (queries == null) { + throw new NullPointerException("queries"); + } + if (params == null) { + throw new NullPointerException("params"); + } + if (roaringFilter == null) { + throw new NullPointerException("roaringFilter"); + } + rejectCallbackReentry(); + synchronized (nativeHandleLock) { + enterNativeHandle(); + try { + return VectorIndexNative.rangeSearchBatchWithRoaringFilter( + requireOpen(), queries, queryCount, params, roaringFilter); + } finally { + exitNativeHandle(); + } + } + } + @Override public void close() { rejectCallbackReentry(); diff --git a/java/src/main/java/org/apache/paimon/index/vector/VectorRangeSearchParams.java b/java/src/main/java/org/apache/paimon/index/vector/VectorRangeSearchParams.java new file mode 100644 index 00000000..0c863a2b --- /dev/null +++ b/java/src/main/java/org/apache/paimon/index/vector/VectorRangeSearchParams.java @@ -0,0 +1,43 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package org.apache.paimon.index.vector; + +import java.util.Objects; + +/** Range-search parameters with an explicit positive IVF probe count and no result cap. */ +public final class VectorRangeSearchParams { + + private final VectorDistanceBand band; + private final int nprobe; + + public VectorRangeSearchParams(VectorDistanceBand band, int nprobe) { + this.band = Objects.requireNonNull(band, "band"); + if (nprobe <= 0) { + throw new IllegalArgumentException("nprobe must be greater than 0"); + } + this.nprobe = nprobe; + } + + public VectorDistanceBand band() { + return band; + } + + public int nprobe() { + return nprobe; + } +} diff --git a/java/src/main/java/org/apache/paimon/index/vector/VectorRangeSearchResult.java b/java/src/main/java/org/apache/paimon/index/vector/VectorRangeSearchResult.java new file mode 100644 index 00000000..d0902983 --- /dev/null +++ b/java/src/main/java/org/apache/paimon/index/vector/VectorRangeSearchResult.java @@ -0,0 +1,145 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package org.apache.paimon.index.vector; + +import java.util.Arrays; +import java.util.Objects; + +/** + * CSR range-search output in core scan order, with raw distances and no sorting or top-K cap. + * IVF-Flat distances are exact; SQ, PQ and RQ return their core distance estimates. Each query + * occupies [lims[query], lims[query + 1]); all arrays are defensively copied. + */ +public final class VectorRangeSearchResult { + + private final long[] labels; + private final float[] distances; + private final long[] lims; + private final long[] listsProbed; + private final long[] rowsScanned; + private final long[] rowsCommitted; + private final long[] earlyAbandoned; + private final long listReads; + + public VectorRangeSearchResult( + long[] labels, + float[] distances, + long[] lims, + long[] listsProbed, + long[] rowsScanned, + long[] rowsCommitted, + long[] earlyAbandoned, + long listReads) { + this.labels = Objects.requireNonNull(labels, "labels").clone(); + this.distances = Objects.requireNonNull(distances, "distances").clone(); + this.lims = Objects.requireNonNull(lims, "lims").clone(); + if (this.labels.length != this.distances.length + || this.lims.length == 0 + || this.lims[0] != 0 + || this.lims[this.lims.length - 1] != this.labels.length) { + throw new IllegalArgumentException("invalid CSR result shape"); + } + for (int offset = 1; offset < this.lims.length; offset++) { + if (this.lims[offset] < this.lims[offset - 1] + || this.lims[offset] > this.labels.length) { + throw new IllegalArgumentException("invalid CSR limits"); + } + } + this.listsProbed = copyCounters(listsProbed, "listsProbed"); + this.rowsScanned = copyCounters(rowsScanned, "rowsScanned"); + this.rowsCommitted = copyCounters(rowsCommitted, "rowsCommitted"); + this.earlyAbandoned = copyCounters(earlyAbandoned, "earlyAbandoned"); + if (listReads < 0) { + throw new IllegalArgumentException("listReads must be non-negative"); + } + this.listReads = listReads; + } + + public int queryCount() { + return lims.length - 1; + } + + public long[] labels() { + return labels.clone(); + } + + public float[] distances() { + return distances.clone(); + } + + public long[] lims() { + return lims.clone(); + } + + /** Logical IVF ranks probed per query, including empty lists. */ + public long[] listsProbed() { + return listsProbed.clone(); + } + + /** Allow-listed rows evaluated per query, including early-abandoned rows. */ + public long[] rowsScanned() { + return rowsScanned.clone(); + } + + public long[] rowsCommitted() { + return rowsCommitted.clone(); + } + + /** Core diagnostic count, not an arithmetic work measure. */ + public long[] earlyAbandoned() { + return earlyAbandoned.clone(); + } + + /** Non-empty unique list reads for the whole call; SQ cache hits do not count. */ + public long listReads() { + return listReads; + } + + public long[] labelsForQuery(int queryIndex) { + checkQueryIndex(queryIndex); + return Arrays.copyOfRange( + labels, Math.toIntExact(lims[queryIndex]), Math.toIntExact(lims[queryIndex + 1])); + } + + public float[] distancesForQuery(int queryIndex) { + checkQueryIndex(queryIndex); + return Arrays.copyOfRange( + distances, + Math.toIntExact(lims[queryIndex]), + Math.toIntExact(lims[queryIndex + 1])); + } + + private void checkQueryIndex(int queryIndex) { + if (queryIndex < 0 || queryIndex >= queryCount()) { + throw new IndexOutOfBoundsException("queryIndex " + queryIndex + " out of range"); + } + } + + private long[] copyCounters(long[] counters, String name) { + long[] copy = Objects.requireNonNull(counters, name).clone(); + if (copy.length != queryCount()) { + throw new IllegalArgumentException(name + " length must equal queryCount"); + } + for (long value : copy) { + if (value < 0) { + throw new IllegalArgumentException(name + " must be non-negative"); + } + } + return copy; + } +} diff --git a/java/src/test/java/org/apache/paimon/index/vector/VectorIndexJavaApiTest.java b/java/src/test/java/org/apache/paimon/index/vector/VectorIndexJavaApiTest.java index b9bc6a3d..fc3f3237 100644 --- a/java/src/test/java/org/apache/paimon/index/vector/VectorIndexJavaApiTest.java +++ b/java/src/test/java/org/apache/paimon/index/vector/VectorIndexJavaApiTest.java @@ -24,6 +24,7 @@ public class VectorIndexJavaApiTest { public static void main(String[] args) { + VectorIndexRangeSearchTest.testValueTypes(); testNativeLibraryResourcePaths(); testSingleResultCopiesArrays(); testBatchResultCopiesArraysAndSlicesRows(); diff --git a/java/src/test/java/org/apache/paimon/index/vector/VectorIndexNativeValidationTest.java b/java/src/test/java/org/apache/paimon/index/vector/VectorIndexNativeValidationTest.java index e33692d9..83050ae8 100644 --- a/java/src/test/java/org/apache/paimon/index/vector/VectorIndexNativeValidationTest.java +++ b/java/src/test/java/org/apache/paimon/index/vector/VectorIndexNativeValidationTest.java @@ -42,6 +42,7 @@ public static void main(String[] args) { testHighLevelTrainingPreservesExpectedVectorCount(); testSupportedIndexRoundtrips(); testDiskAnnInnerProductAndCosine(); + VectorIndexRangeSearchTest.testNative(); } private static void testReaderCapabilityFailuresArePropagatedBeforeOpen() { diff --git a/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeOracleTest.java b/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeOracleTest.java new file mode 100644 index 00000000..e4b9aca3 --- /dev/null +++ b/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeOracleTest.java @@ -0,0 +1,190 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package org.apache.paimon.index.vector; + +import java.io.IOException; +import java.nio.charset.StandardCharsets; +import java.nio.file.Files; +import java.nio.file.Path; +import java.nio.file.Paths; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Collections; +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import java.util.Scanner; + +public class VectorIndexRangeOracleTest { + + public static void main(String[] args) { + VectorIndexNativeLoaderSmokeTest.configureExternalLibrary(args); + runIfConfigured(); + } + + static void runIfConfigured() { + String directory = System.getenv("PVI_RANGE_FIXTURES"); + if (directory == null || directory.isEmpty()) { + return; + } + Path root = Paths.get(directory); + int cases = 0; + try { + for (String line : + Files.readAllLines(root.resolve("manifest.txt"), StandardCharsets.UTF_8)) { + if (line.trim().isEmpty()) { + continue; + } + String[] fields = line.trim().split("\\s+"); + require(fields.length == 2, "manifest entry: " + line); + try { + runCase(root.resolve(fields[0] + ".expected"), root.resolve(fields[1])); + } catch (Throwable error) { + throw new AssertionError("Core range oracle case " + fields[0], error); + } + cases++; + } + } catch (IOException error) { + throw new AssertionError("Cannot read core range oracle " + root, error); + } + require(cases > 0, "oracle manifest is empty"); + System.out.println("Core range oracle: " + cases + " cases passed"); + } + + private static void runCase(Path expected, Path index) throws IOException { + try (Scanner values = new Scanner(expected, StandardCharsets.UTF_8.name())) { + int dimension = values.nextInt(); + int metricCode = values.nextInt(); + int queryCount = values.nextInt(); + int nprobe = values.nextInt(); + Float lower = readBound(values); + Float upper = readBound(values); + int filterLength = values.nextInt(); + int hitCount = values.nextInt(); + long listReads = values.nextLong(); + float[] queries = new float[Math.multiplyExact(dimension, queryCount)]; + for (int offset = 0; offset < queries.length; offset++) { + queries[offset] = Float.intBitsToFloat(readBits(values)); + } + byte[] filter = new byte[filterLength]; + for (int offset = 0; offset < filter.length; offset++) { + int value = values.nextInt(); + require(value >= 0 && value <= 255, "filter byte"); + filter[offset] = (byte) value; + } + long[] lims = readLongs(values, Math.addExact(queryCount, 1)); + long[] labels = readLongs(values, hitCount); + int[] distances = new int[hitCount]; + for (int offset = 0; offset < hitCount; offset++) { + distances[offset] = readBits(values); + } + long[][] stats = new long[4][queryCount]; + for (int queryIndex = 0; queryIndex < queryCount; queryIndex++) { + for (int counter = 0; counter < 4; counter++) { + stats[counter][queryIndex] = values.nextLong(); + } + } + require(!values.hasNext(), "unexpected trailing oracle data"); + String[] metrics = {"l2", "inner_product", "cosine"}; + require(metricCode >= 0 && metricCode < metrics.length, "metric code"); + VectorRangeSearchParams params = + new VectorRangeSearchParams( + new VectorDistanceBand(metrics[metricCode], lower, upper), nprobe); + try (VectorIndexReader reader = + new VectorIndexReader( + new VectorIndexNativeValidationTest.ByteArraySeekableInputStream( + Files.readAllBytes(index)))) { + require(reader.supportsRangeSearch(), "range support"); + VectorRangeSearchResult actual; + if (queryCount == 1) { + actual = + filterLength == 0 + ? reader.rangeSearch(queries, params) + : reader.rangeSearch(queries, params, filter); + } else { + actual = + filterLength == 0 + ? reader.rangeSearchBatch(queries, queryCount, params) + : reader.rangeSearchBatch(queries, queryCount, params, filter); + } + require(actual.queryCount() == queryCount, "query count"); + require(Arrays.equals(lims, actual.lims()), "lims"); + float[] actualDistances = actual.distances(); + require(actualDistances.length == hitCount, "distance count"); + int[] actualBits = new int[hitCount]; + for (int offset = 0; offset < hitCount; offset++) { + actualBits[offset] = Float.floatToRawIntBits(actualDistances[offset]); + } + long[] actualLabels = actual.labels(); + for (int queryIndex = 0; queryIndex < queryCount; queryIndex++) { + int start = Math.toIntExact(lims[queryIndex]); + int end = Math.toIntExact(lims[queryIndex + 1]); + require( + rows(labels, distances, start, end) + .equals(rows(actualLabels, actualBits, start, end)), + "label/distance-bit multiset for query " + queryIndex); + } + require(Arrays.equals(stats[0], actual.listsProbed()), "listsProbed"); + require(Arrays.equals(stats[1], actual.rowsScanned()), "rowsScanned"); + require(Arrays.equals(stats[2], actual.rowsCommitted()), "rowsCommitted"); + require(Arrays.equals(stats[3], actual.earlyAbandoned()), "earlyAbandoned"); + require(listReads == actual.listReads(), "listReads"); + } + } + } + + private static Float readBound(Scanner values) { + int kind = values.nextInt(); + int bits = readBits(values); + require(kind == 0 || kind == 1, "bound kind"); + return kind == 0 ? null : Float.intBitsToFloat(bits); + } + + private static Map> rows( + long[] labels, int[] distances, int start, int end) { + Map> result = new HashMap>(); + for (int offset = start; offset < end; offset++) { + result.computeIfAbsent(labels[offset], label -> new ArrayList()) + .add(distances[offset]); + } + for (List values : result.values()) { + Collections.sort(values); + } + return result; + } + + private static int readBits(Scanner values) { + long value = values.nextLong(); + require(value >= 0 && value <= 0xffff_ffffL, "f32 bits"); + return (int) value; + } + + private static long[] readLongs(Scanner values, int count) { + long[] result = new long[count]; + for (int offset = 0; offset < count; offset++) { + result[offset] = values.nextLong(); + } + return result; + } + + private static void require(boolean condition, String description) { + if (!condition) { + throw new AssertionError(description); + } + } +} diff --git a/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeSearchTest.java b/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeSearchTest.java new file mode 100644 index 00000000..e106b741 --- /dev/null +++ b/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeSearchTest.java @@ -0,0 +1,589 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package org.apache.paimon.index.vector; + +import static org.apache.paimon.index.vector.VectorDistanceBand.CutOperator.GE; +import static org.apache.paimon.index.vector.VectorDistanceBand.CutOperator.GT; +import static org.apache.paimon.index.vector.VectorDistanceBand.CutOperator.LE; +import static org.apache.paimon.index.vector.VectorDistanceBand.CutOperator.LT; + +import java.nio.ByteBuffer; +import java.nio.ByteOrder; +import java.util.Arrays; +import java.util.HashMap; +import java.util.Map; + +public class VectorIndexRangeSearchTest { + + private static final long LABEL_BASE = 1L << 33; + private static final int VECTOR_COUNT = 256; + + public static void main(String[] args) { + VectorIndexNativeLoaderSmokeTest.configureExternalLibrary(args); + testValueTypes(); + testNative(); + System.out.println( + "Range search: value types, 12 IVF/metric combinations, endpoints, filters, validation and callbacks passed"); + } + + static void testValueTypes() { + VectorDistanceBand band = new VectorDistanceBand("l2", null, 4.0f); + check(band.lower() == null && band.upper() == 4.0f, "structural bounds"); + check("l2".equals(band.metric()), "metric"); + check(new VectorRangeSearchParams(band, 2).nprobe() == 2, "nprobe"); + expect(IllegalArgumentException.class, () -> new VectorRangeSearchParams(band, 0)); + expect(NullPointerException.class, () -> new VectorRangeSearchParams(null, 1)); + expect(IllegalArgumentException.class, () -> new VectorDistanceBand("other", null, null)); + expect(IllegalArgumentException.class, () -> new VectorDistanceBand("l2", -1.0f, null)); + expect(IllegalArgumentException.class, () -> new VectorDistanceBand("cosine", 2.0f, 1.0f)); + expect(IllegalArgumentException.class, () -> new VectorDistanceBand("l2", null, Float.NaN)); + expect( + IllegalArgumentException.class, + () -> new VectorDistanceBand("cosine", Float.NEGATIVE_INFINITY, null)); + + long[] labels = {LABEL_BASE, 7}; + float[] distances = {1, 3}; + long[] lims = {0, 0, 2}; + long[] counters = {0, 2}; + VectorRangeSearchResult result = + new VectorRangeSearchResult( + labels, distances, lims, counters, counters, counters, new long[2], 2); + labels[0] = -1; + distances[0] = -1; + lims[1] = 2; + counters[1] = -1; + check(result.queryCount() == 2 && result.listReads() == 2, "result shape"); + check(result.labelsForQuery(0).length == 0, "empty first row"); + check(result.labelsForQuery(1)[0] == LABEL_BASE, "64-bit label and copy"); + check(result.distancesForQuery(1)[0] == 1, "distance copy"); + result.labels()[0] = -2; + result.distances()[0] = -2; + result.lims()[1] = 2; + result.listsProbed()[1] = -2; + result.rowsScanned()[1] = -2; + result.rowsCommitted()[1] = -2; + result.earlyAbandoned()[1] = -2; + check(result.labels()[0] == LABEL_BASE && result.distances()[0] == 1, "defensive arrays"); + check(result.lims()[1] == 0 && result.listsProbed()[1] == 2, "defensive CSR and stats"); + check( + result.rowsScanned()[1] == 2 + && result.rowsCommitted()[1] == 2 + && result.earlyAbandoned()[1] == 0, + "defensive counters"); + expect(IndexOutOfBoundsException.class, () -> result.labelsForQuery(-1)); + expect(IndexOutOfBoundsException.class, () -> result.distancesForQuery(2)); + expect( + IllegalArgumentException.class, + () -> + new VectorRangeSearchResult( + new long[0], + new float[0], + new long[] {0, Long.MAX_VALUE}, + new long[1], + new long[1], + new long[1], + new long[1], + 0)); + expect( + IllegalArgumentException.class, + () -> + new VectorRangeSearchResult( + new long[0], + new float[0], + new long[] {0}, + new long[1], + new long[0], + new long[0], + new long[0], + 0)); + expect( + IllegalArgumentException.class, + () -> + new VectorRangeSearchResult( + new long[0], + new float[0], + new long[] {0}, + new long[0], + new long[0], + new long[0], + new long[0], + -1)); + + VectorIndexReader closed = VectorIndexReader.fromNativePointerForTesting(0); + VectorRangeSearchParams params = new VectorRangeSearchParams(band, 1); + expect(IllegalStateException.class, closed::supportsRangeSearch); + expect(IllegalStateException.class, () -> closed.rangeSearch(new float[] {0}, params)); + expect( + IllegalStateException.class, + () -> closed.rangeSearch(new float[] {0}, params, new byte[8])); + expect(IllegalStateException.class, () -> closed.rangeSearchBatch(new float[0], 0, params)); + expect( + IllegalStateException.class, + () -> closed.rangeSearchBatch(new float[0], 0, params, new byte[8])); + expect(NullPointerException.class, () -> closed.rangeSearch(null, params)); + expect(NullPointerException.class, () -> closed.rangeSearch(new float[1], null)); + expect(NullPointerException.class, () -> closed.rangeSearch(new float[1], params, null)); + } + + static void testNative() { + for (String indexType : new String[] {"ivf_flat", "ivf_sq", "ivf_pq", "ivf_rq"}) { + for (String metric : new String[] {"l2", "cosine", "inner_product"}) { + testIndex(indexType, metric); + } + } + testExactDistancesAndEndpoints(); + testNativeValidation(); + testCallbacks(); + testUnsupported(); + VectorIndexRangeOracleTest.runIfConfigured(); + } + + private static void testIndex(String indexType, String metric) { + float[] data = new float[VECTOR_COUNT * 8]; + for (int row = 0; row < VECTOR_COUNT; row++) { + for (int column = 0; column < 8; column++) { + data[row * 8 + column] = ((row * 17 + column * 11) % 67 - 33) / 16.0f; + } + } + float[] queries = Arrays.copyOf(data, 16); + VectorRangeSearchParams all = params(metric, null, null); + try (VectorIndexReader reader = open(build(indexType, metric, 8, data))) { + check(reader.supportsRangeSearch(), indexType + " " + metric + " support"); + VectorRangeSearchResult full = reader.rangeSearchBatch(queries, 2, all); + assertShape(full, 2); + check(full.labels().length == VECTOR_COUNT * 2, "unbounded includes every row"); + check(full.listReads() <= 2, "shared list reads"); + for (int queryIndex = 0; queryIndex < 2; queryIndex++) { + float[] query = Arrays.copyOfRange(queries, queryIndex * 8, (queryIndex + 1) * 8); + assertRows(full, queryIndex, reader.rangeSearch(query, all), 0); + check(full.listsProbed()[queryIndex] == 2, "lists probed"); + check(full.rowsScanned()[queryIndex] == VECTOR_COUNT, "rows scanned"); + } + float[] sorted = full.distancesForQuery(0); + Arrays.sort(sorted); + Float lower = sorted[VECTOR_COUNT / 4]; + Float upper = sorted[VECTOR_COUNT * 3 / 4]; + VectorRangeSearchParams bounded = params(metric, lower, upper); + VectorRangeSearchResult selected = reader.rangeSearchBatch(queries, 2, bounded); + assertSelected(full, selected, lower, upper, false); + VectorRangeSearchResult filtered = + reader.rangeSearchBatch(queries, 2, bounded, filter()); + assertSelected(full, filtered, lower, upper, true); + for (int queryIndex = 0; queryIndex < 2; queryIndex++) { + float[] query = Arrays.copyOfRange(queries, queryIndex * 8, (queryIndex + 1) * 8); + assertRows(selected, queryIndex, reader.rangeSearch(query, bounded), 0); + assertRows(filtered, queryIndex, reader.rangeSearch(query, bounded, filter()), 0); + } + assertSelected( + full, + reader.rangeSearchBatch(queries, 2, params(metric, null, upper)), + null, + upper, + false); + assertSelected( + full, + reader.rangeSearchBatch(queries, 2, params(metric, lower, null)), + lower, + null, + false); + expectMessage( + RuntimeException.class, + "query count must be greater than 0", + () -> reader.rangeSearchBatch(new float[0], 0, all)); + expectMessage( + RuntimeException.class, + "query count must be greater than 0", + () -> reader.rangeSearchBatch(new float[0], 0, all, filter())); + VectorRangeSearchResult empty = + reader.rangeSearchBatch(queries, 2, params(metric, 0.0f, 0.0f)); + check(empty.labels().length == 0 && empty.listReads() == 0, "empty band"); + check( + reader.rangeSearchBatch(queries, 2, all, new byte[8]).labels().length == 0, + "empty allow-list"); + expect( + RuntimeException.class, + () -> reader.rangeSearchBatch(queries, 2, all, new byte[] {1})); + } + } + + private static void testExactDistancesAndEndpoints() { + for (String metric : new String[] {"l2", "cosine", "inner_product"}) { + try (VectorIndexReader reader = + open(build("ivf_flat", metric, 2, new float[] {1, 0, 2, 0, 0, 1, -1, 0}))) { + float[] query = {1, 0}; + VectorRangeSearchResult raw = reader.rangeSearch(query, params(metric, null, null)); + Map rows = rows(raw, 0); + check( + rows.get(LABEL_BASE + 1) + == ("l2".equals(metric) + ? 1.0f + : "cosine".equals(metric) ? 0.0f : -2.0f), + "raw metric distance"); + double endpoint = + "inner_product".equals(metric) + ? 1.0 + : "l2".equals(metric) ? Math.sqrt(2.0f) : 1.0; + if ("l2".equals(metric)) { + endpoint = (double) (float) endpoint; + } + for (VectorDistanceBand.CutOperator operator : + VectorDistanceBand.CutOperator.values()) { + boolean lower = operator == GE || operator == GT; + for (double literal : + new double[] { + endpoint, Math.nextUp(endpoint), Math.nextDown(endpoint) + }) { + VectorDistanceBand band = + VectorDistanceBand.fromEndpoints( + metric, + lower ? literal : null, + lower ? operator : null, + lower ? null : literal, + lower ? null : operator); + Map selected = + rows( + reader.rangeSearch( + query, new VectorRangeSearchParams(band, 2)), + 0); + for (Map.Entry row : rows.entrySet()) { + double value = + "l2".equals(metric) + ? (double) (float) Math.sqrt(row.getValue()) + : "inner_product".equals(metric) + ? -row.getValue() + : row.getValue(); + boolean admitted = + operator == GE + ? value >= literal + : operator == GT + ? value > literal + : operator == LE + ? value <= literal + : value < literal; + check( + selected.containsKey(row.getKey()) == admitted, + "core endpoint conversion " + metric + " " + operator); + } + } + } + } + } + VectorDistanceBand unbounded = + VectorDistanceBand.fromEndpoints("inner_product", null, null, null, null); + check( + unbounded.lower() == null && unbounded.upper() == null, + "unbounded endpoint conversion"); + VectorDistanceBand outside = + VectorDistanceBand.fromEndpoints( + "cosine", -Double.MAX_VALUE, GE, Double.MAX_VALUE, LE); + check( + outside.lower() == -Float.MAX_VALUE && outside.upper() == null, + "linear out-of-domain endpoints retain core cuts and unbounded upper"); + VectorDistanceBand equal = VectorDistanceBand.fromEndpoints("l2", 1.0, GE, 1.0, LE); + check( + equal.lower() <= 1.0f && equal.upper() > 1.0f, + "inclusive equal endpoints retain equality bucket"); + expect( + RuntimeException.class, + () -> VectorDistanceBand.fromEndpoints("l2", 1.0, LT, null, null)); + expect( + RuntimeException.class, + () -> VectorDistanceBand.fromEndpoints("cosine", null, null, Double.NaN, LT)); + expect( + RuntimeException.class, + () -> VectorDistanceBand.fromEndpoints("l2", null, null, Double.MAX_VALUE, LE)); + expect( + IllegalArgumentException.class, + () -> VectorDistanceBand.fromEndpoints("l2", 1.0, null, null, null)); + expect( + RuntimeException.class, + () -> VectorIndexNative.distanceBandFromEndpoints("l2", 1.0, 4, null, -1)); + } + + private static void testNativeValidation() { + VectorRangeSearchParams all = params("l2", null, null); + expect(RuntimeException.class, () -> VectorIndexNative.supportsRangeSearch(0)); + expect(RuntimeException.class, () -> VectorIndexNative.rangeSearch(0, new float[2], all)); + try (VectorIndexReader reader = + open(build("ivf_flat", "l2", 2, new float[] {0, 0, 1, 1, 2, 2, 3, 3}))) { + expect(RuntimeException.class, () -> reader.rangeSearch(new float[1], all)); + expectMessage( + RuntimeException.class, + "query length", + () -> reader.rangeSearch(new float[1], all, new byte[] {1})); + expect( + RuntimeException.class, + () -> reader.rangeSearch(new float[] {Float.NaN, 0}, all)); + expect(RuntimeException.class, () -> reader.rangeSearchBatch(new float[2], 2, all)); + expect(RuntimeException.class, () -> reader.rangeSearchBatch(new float[0], -1, all)); + expect( + RuntimeException.class, + () -> reader.rangeSearchBatch(new float[0], Integer.MAX_VALUE, all)); + expect( + RuntimeException.class, + () -> reader.rangeSearch(new float[2], params("cosine", null, null))); + expect( + RuntimeException.class, + () -> reader.rangeSearchBatch(new float[0], 0, params("cosine", null, null))); + check( + reader.rangeSearch(new float[2], all).labels().length == 4, + "reader survives errors"); + check( + reader.rangeSearch( + new float[2], + new VectorRangeSearchParams( + all.band(), Integer.MAX_VALUE)) + .listsProbed()[0] + == 2, + "core clamps nprobe"); + } + } + + private static void testCallbacks() { + byte[] bytes = build("ivf_flat", "l2", 2, new float[] {0, 0, 1, 1, 2, 2, 3, 3}); + VectorIndexNativeValidationTest.ByteArraySeekableInputStream delegate = + new VectorIndexNativeValidationTest.ByteArraySeekableInputStream(bytes); + VectorIndexReader[] holder = new VectorIndexReader[1]; + int[] callbacks = {0}; + boolean[] fail = {false}; + IllegalStateException callbackFailure = new IllegalStateException("range input failure"); + VectorRangeSearchParams all = params("l2", null, null); + try (VectorIndexReader reader = + new VectorIndexReader( + (positions, buffers) -> { + if (holder[0] != null) { + callbacks[0]++; + expect(IllegalStateException.class, holder[0]::close); + expect(IllegalStateException.class, holder[0]::supportsRangeSearch); + expect( + IllegalStateException.class, + () -> holder[0].rangeSearch(new float[2], all)); + expect( + IllegalStateException.class, + () -> holder[0].rangeSearch(new float[2], all, filter())); + expect( + IllegalStateException.class, + () -> holder[0].rangeSearchBatch(new float[2], 1, all)); + expect( + IllegalStateException.class, + () -> + holder[0].rangeSearchBatch( + new float[2], 1, all, filter())); + } + if (fail[0]) { + throw callbackFailure; + } + delegate.pread(positions, buffers); + })) { + holder[0] = reader; + reader.rangeSearch(new float[2], all); + reader.rangeSearchBatch(new float[4], 2, all, filter()); + check(callbacks[0] > 0, "range callbacks exercised"); + fail[0] = true; + check( + expect(IllegalStateException.class, () -> reader.rangeSearch(new float[2], all)) + == callbackFailure, + "callback exception identity preserved"); + check( + expect( + IllegalStateException.class, + () -> reader.rangeSearchBatch(new float[4], 2, all, filter())) + == callbackFailure, + "batch callback exception identity preserved"); + fail[0] = false; + check( + reader.rangeSearch(new float[2], all).labels().length == 4, + "reader usable after callback failure"); + } + } + + private static void testUnsupported() { + float[] data = new float[64 * 8]; + for (int offset = 0; offset < data.length; offset++) { + data[offset] = (offset * 17 % 67) / 16.0f; + } + try (VectorIndexReader reader = open(build("diskann", "l2", 8, data))) { + check(!reader.supportsRangeSearch(), "DiskANN unsupported"); + VectorRangeSearchParams empty = params("l2", 0.0f, 0.0f); + expectMessage( + RuntimeException.class, + "not supported", + () -> reader.rangeSearch(new float[8], empty)); + expectMessage( + RuntimeException.class, + "not supported", + () -> reader.rangeSearch(new float[8], empty, filter())); + expectMessage( + RuntimeException.class, + "not supported", + () -> reader.rangeSearchBatch(new float[16], 2, empty)); + expectMessage( + RuntimeException.class, + "not supported", + () -> reader.rangeSearchBatch(new float[16], 2, empty, filter())); + } + } + + private static VectorRangeSearchParams params(String metric, Float lower, Float upper) { + return new VectorRangeSearchParams(new VectorDistanceBand(metric, lower, upper), 2); + } + + private static VectorIndexReader open(byte[] bytes) { + return new VectorIndexReader( + new VectorIndexNativeValidationTest.ByteArraySeekableInputStream(bytes)); + } + + private static byte[] build(String type, String metric, int dimension, float[] data) { + Map options = new HashMap(); + options.put("index.type", type); + options.put("metric", metric); + options.put("dimension", Integer.toString(dimension)); + if (!"diskann".equals(type)) { + options.put("nlist", "2"); + } + if ("ivf_pq".equals(type) || "diskann".equals(type)) { + options.put("pq.m", "2"); + } + if ("ivf_pq".equals(type)) { + options.put("use-opq", "false"); + } + if ("ivf_rq".equals(type)) { + options.put("rq.bits", "5"); + } + if ("diskann".equals(type)) { + options.put("pq.bits", "4"); + options.put("diskann.max-degree", "8"); + options.put("diskann.build-search-list-size", "16"); + } + int count = data.length / dimension; + long[] ids = new long[count]; + for (int row = 0; row < count; row++) { + ids[row] = LABEL_BASE + row; + } + VectorIndexNativeValidationTest.ByteArrayPositionOutputStream output = + new VectorIndexNativeValidationTest.ByteArrayPositionOutputStream(); + try (VectorIndexTraining training = VectorIndexTrainer.train(options, data, count); + VectorIndexWriter writer = new VectorIndexWriter(training)) { + writer.addVectors(ids, data, count); + writer.writeIndex(output); + } + return output.toByteArray(); + } + + private static byte[] filter() { + ByteBuffer bytes = ByteBuffer.allocate(36).order(ByteOrder.LITTLE_ENDIAN); + bytes.putLong(1).putInt(2).putInt(12346).putInt(1); + bytes.putShort((short) 0).putShort((short) 3).putInt(16); + bytes.putShort((short) 0).putShort((short) 1).putShort((short) 16).putShort((short) 63); + return bytes.array(); + } + + private static void assertShape(VectorRangeSearchResult result, int count) { + check(result.queryCount() == count && result.lims().length == count + 1, "CSR shape"); + check( + result.lims()[0] == 0 && result.lims()[count] == result.labels().length, + "CSR limits"); + check(result.labels().length == result.distances().length, "parallel results"); + for (int queryIndex = 0; queryIndex < count; queryIndex++) { + check( + result.rowsCommitted()[queryIndex] == result.labelsForQuery(queryIndex).length, + "committed counter"); + check( + result.rowsScanned()[queryIndex] >= result.rowsCommitted()[queryIndex], + "scanned counter"); + check( + result.earlyAbandoned()[queryIndex] <= result.rowsScanned()[queryIndex], + "abandon counter"); + } + } + + private static Map rows(VectorRangeSearchResult result, int queryIndex) { + Map rows = new HashMap(); + long[] labels = result.labelsForQuery(queryIndex); + float[] distances = result.distancesForQuery(queryIndex); + for (int row = 0; row < labels.length; row++) { + check(Float.isFinite(distances[row]), "finite raw distance"); + check(rows.put(labels[row], distances[row]) == null, "unique labels"); + } + return rows; + } + + private static void assertRows( + VectorRangeSearchResult expected, + int expectedQuery, + VectorRangeSearchResult actual, + int actualQuery) { + assertShape(actual, actual.queryCount()); + check( + rows(expected, expectedQuery).equals(rows(actual, actualQuery)), + "single/batch row equality"); + } + + private static void assertSelected( + VectorRangeSearchResult full, + VectorRangeSearchResult actual, + Float lower, + Float upper, + boolean filtered) { + assertShape(actual, full.queryCount()); + for (int queryIndex = 0; queryIndex < full.queryCount(); queryIndex++) { + Map expected = rows(full, queryIndex); + expected.entrySet() + .removeIf( + row -> + (lower != null && row.getValue() < lower) + || (upper != null && row.getValue() >= upper) + || (filtered + && row.getKey() != LABEL_BASE + && row.getKey() != LABEL_BASE + 1 + && row.getKey() != LABEL_BASE + 16 + && row.getKey() != LABEL_BASE + 63)); + check(expected.equals(rows(actual, queryIndex)), "half-open raw membership and filter"); + } + } + + private static void check(boolean condition, String message) { + if (!condition) { + throw new AssertionError(message); + } + } + + private static void expectMessage( + Class type, String message, Runnable action) { + Throwable error = expect(type, action); + check( + error.getMessage() != null && error.getMessage().contains(message), + "exception message: " + error.getMessage()); + } + + private static Throwable expect(Class type, Runnable action) { + try { + action.run(); + } catch (Throwable error) { + if (!type.isInstance(error)) { + throw new AssertionError("expected " + type.getName(), error); + } + if (error.getMessage() != null + && error.getMessage().contains("Rust panic in JNI call")) { + throw new AssertionError("validation must not panic", error); + } + return error; + } + throw new AssertionError("expected " + type.getName()); + } +} diff --git a/jni/src/lib.rs b/jni/src/lib.rs index e025cbd2..be5ddca6 100644 --- a/jni/src/lib.rs +++ b/jni/src/lib.rs @@ -16,6 +16,7 @@ // under the License. mod log_bridge; +mod range; mod stream; use jni::objects::{JByteArray, JClass, JFloatArray, JLongArray, JObject, JValue}; diff --git a/jni/src/range.rs b/jni/src/range.rs new file mode 100644 index 00000000..4fbcd0e0 --- /dev/null +++ b/jni/src/range.rs @@ -0,0 +1,403 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +use jni::objects::{ + AutoLocal, JByteArray, JClass, JFloatArray, JLongArray, JObject, JString, JValue, +}; +use jni::sys::{jboolean, jint, jlong, jobject}; +use jni::JNIEnv; +use paimon_vindex_core::distance::MetricType; +use paimon_vindex_core::range::{ + Bound, CutOperator, DistanceBand, DistanceEndpoint, RangeSearchResult, VectorRangeSearchParams, +}; + +use crate::{ + call_int_method, deref_reader, jni_call, read_byte_array, read_float_array, throw_and_return, +}; + +fn range_call(env: JNIEnv, action: impl FnOnce(&mut JNIEnv) -> Result) -> T { + jni_call(env, |env| match action(env) { + Ok(result) => result, + Err(error) => { + if env.exception_check().unwrap_or(true) { + T::default() + } else { + throw_and_return(env, &error) + } + } + }) +} + +fn checked_array_length(value: usize, name: &str) -> Result { + jint::try_from(value).map_err(|_| format!("{name} exceeds Java array length limit")) +} + +fn checked_counter(value: usize, name: &str) -> Result { + jlong::try_from(value).map_err(|_| format!("{name} exceeds Java long limit")) +} + +fn parse_metric(value: &str) -> Result { + match value { + "l2" => Ok(MetricType::L2), + "inner_product" => Ok(MetricType::InnerProduct), + "cosine" => Ok(MetricType::Cosine), + _ => Err(format!("unknown range metric: {value}")), + } +} + +fn read_metric(env: &mut JNIEnv, metric: JObject) -> Result { + if metric.is_null() { + return Err("metric is null".to_string()); + } + let metric = env.auto_local(JString::from(metric)); + let value: String = env + .get_string(&metric) + .map_err(|error| error.to_string())? + .into(); + parse_metric(&value) +} + +fn cut_operator(code: jint) -> Result { + match code { + 0 => Ok(CutOperator::Ge), + 1 => Ok(CutOperator::Gt), + 2 => Ok(CutOperator::Le), + 3 => Ok(CutOperator::Lt), + _ => Err(format!("invalid endpoint operator: {code}")), + } +} + +fn read_endpoint( + env: &mut JNIEnv, + value: JObject, + code: jint, +) -> Result, String> { + if value.is_null() { + if code != -1 { + return Err("unbounded endpoint must not have an operator".to_string()); + } + return Ok(None); + } + let op = cut_operator(code)?; + let value = env + .call_method(value, "doubleValue", "()D", &[]) + .and_then(|value| value.d()) + .map_err(|error| error.to_string())?; + Ok(Some(DistanceEndpoint { value, op })) +} + +fn read_bound(env: &mut JNIEnv, band: &JObject, name: &str) -> Result { + let value = env + .call_method(band, name, "()Ljava/lang/Float;", &[]) + .and_then(|value| value.l()) + .map_err(|error| error.to_string())?; + let value = env.auto_local(value); + if value.is_null() { + Ok(Bound::Unbounded) + } else { + env.call_method(&value, "floatValue", "()F", &[]) + .and_then(|value| value.f()) + .map(Bound::Finite) + .map_err(|error| error.to_string()) + } +} + +fn range_params(env: &mut JNIEnv, params: JObject) -> Result { + if params.is_null() { + return Err("params is null".to_string()); + } + let nprobe = call_int_method(env, ¶ms, "nprobe")?; + let nprobe = usize::try_from(nprobe).map_err(|_| "nprobe must be positive".to_string())?; + let band = env + .call_method( + ¶ms, + "band", + "()Lorg/apache/paimon/index/vector/VectorDistanceBand;", + &[], + ) + .and_then(|value| value.l()) + .map_err(|error| error.to_string())?; + let band = env.auto_local(band); + if band.is_null() { + return Err("band is null".to_string()); + } + let metric = env + .call_method(&band, "metric", "()Ljava/lang/String;", &[]) + .and_then(|value| value.l()) + .map_err(|error| error.to_string())?; + let metric = read_metric(env, metric)?; + let lower = read_bound(env, &band, "lower")?; + let upper = read_bound(env, &band, "upper")?; + let band = DistanceBand::new(lower, upper, metric).map_err(|error| error.to_string())?; + Ok(VectorRangeSearchParams::new(band, nprobe)) +} + +fn boxed_bound<'local>( + env: &mut JNIEnv<'local>, + bound: Bound, +) -> Result>, String> { + let object = match bound { + Bound::Unbounded => JObject::null(), + Bound::Finite(value) => env + .new_object("java/lang/Float", "(F)V", &[JValue::Float(value)]) + .map_err(|error| error.to_string())?, + }; + Ok(env.auto_local(object)) +} + +fn long_array<'local>( + env: &mut JNIEnv<'local>, + values: &[jlong], + name: &str, +) -> Result>, String> { + let array = env + .new_long_array(checked_array_length(values.len(), name)?) + .map_err(|error| error.to_string())?; + let array = env.auto_local(array); + env.set_long_array_region(&array, 0, values) + .map_err(|error| error.to_string())?; + Ok(array) +} + +fn build_range_result(env: &mut JNIEnv, result: RangeSearchResult) -> Result { + checked_array_length(result.query_count(), "query count")?; + checked_array_length(result.lims().len(), "lims")?; + checked_array_length(result.labels().len(), "labels")?; + let distance_count = checked_array_length(result.distances().len(), "distances")?; + let list_reads = checked_counter(result.call_stats().list_reads(), "listReads")?; + let lims = result + .lims() + .iter() + .map(|&value| checked_counter(value, "lims")) + .collect::, _>>()?; + let mut lists_probed = Vec::with_capacity(result.query_count()); + let mut rows_scanned = Vec::with_capacity(result.query_count()); + let mut rows_committed = Vec::with_capacity(result.query_count()); + let mut early_abandoned = Vec::with_capacity(result.query_count()); + for query_index in 0..result.query_count() { + let stats = result.query(query_index).stats; + lists_probed.push(checked_counter(stats.lists_probed(), "listsProbed")?); + rows_scanned.push(checked_counter(stats.rows_scanned(), "rowsScanned")?); + rows_committed.push(checked_counter(stats.rows_committed(), "rowsCommitted")?); + early_abandoned.push(checked_counter(stats.early_abandoned(), "earlyAbandoned")?); + } + let labels = long_array(env, result.labels(), "labels")?; + let distances = env + .new_float_array(distance_count) + .map_err(|error| error.to_string())?; + let distances = env.auto_local(distances); + env.set_float_array_region(&distances, 0, result.distances()) + .map_err(|error| error.to_string())?; + let lims = long_array(env, &lims, "lims")?; + let lists_probed = long_array(env, &lists_probed, "listsProbed")?; + let rows_scanned = long_array(env, &rows_scanned, "rowsScanned")?; + let rows_committed = long_array(env, &rows_committed, "rowsCommitted")?; + let early_abandoned = long_array(env, &early_abandoned, "earlyAbandoned")?; + env.new_object( + "org/apache/paimon/index/vector/VectorRangeSearchResult", + "([J[F[J[J[J[J[JJ)V", + &[ + JValue::Object(&labels), + JValue::Object(&distances), + JValue::Object(&lims), + JValue::Object(&lists_probed), + JValue::Object(&rows_scanned), + JValue::Object(&rows_committed), + JValue::Object(&early_abandoned), + JValue::Long(list_reads), + ], + ) + .map(|object| object.into_raw()) + .map_err(|error| error.to_string()) +} + +fn range_search( + env: &mut JNIEnv, + ptr: jlong, + queries: JFloatArray, + query_count: Option, + params: JObject, + filter: Option, +) -> Result { + let reader = deref_reader(ptr) + .ok_or_else(|| "null native pointer (reader already freed?)".to_string())?; + let query_count = query_count + .map(|value| usize::try_from(value).map_err(|_| format!("invalid query count: {value}"))) + .transpose()?; + let params = range_params(env, params)?; + let queries = read_float_array(env, &queries, "queries")?; + let filter = filter + .map(|filter| read_byte_array(env, filter)) + .transpose()?; + let result = match (query_count, filter) { + (None, None) => reader.range_search(&queries, params), + (None, Some(filter)) => reader.range_search_with_roaring_filter(&queries, params, &filter), + (Some(count), None) => reader.range_search_batch(&queries, count, params), + (Some(count), Some(filter)) => { + reader.range_search_batch_with_roaring_filter(&queries, count, params, &filter) + } + } + .map_err(|error| format!("range_search: {error}"))?; + build_range_result(env, result) +} + +#[no_mangle] +pub extern "system" fn Java_org_apache_paimon_index_vector_VectorIndexNative_distanceBandFromEndpoints( + env: JNIEnv, + _class: JClass, + metric: JObject, + lower: JObject, + lower_operator: jint, + upper: JObject, + upper_operator: jint, +) -> jobject { + range_call(env, |env| { + let metric = read_metric(env, metric)?; + let lower = read_endpoint(env, lower, lower_operator)?; + let upper = read_endpoint(env, upper, upper_operator)?; + let band = DistanceBand::from_endpoints(lower, upper, metric) + .map_err(|error| error.to_string())?; + let lower = boxed_bound(env, band.lower())?; + let upper = boxed_bound(env, band.upper())?; + let metric = env + .new_string(metric.as_str()) + .map_err(|error| error.to_string())?; + let metric = env.auto_local(metric); + env.new_object( + "org/apache/paimon/index/vector/VectorDistanceBand", + "(Ljava/lang/String;Ljava/lang/Float;Ljava/lang/Float;)V", + &[ + JValue::Object(&metric), + JValue::Object(&lower), + JValue::Object(&upper), + ], + ) + .map(|object| object.into_raw()) + .map_err(|error| error.to_string()) + }) +} + +#[no_mangle] +pub extern "system" fn Java_org_apache_paimon_index_vector_VectorIndexNative_supportsRangeSearch( + env: JNIEnv, + _class: JClass, + ptr: jlong, +) -> jboolean { + range_call(env, |_env| { + let reader = deref_reader(ptr) + .ok_or_else(|| "null native pointer (reader already freed?)".to_string())?; + Ok(jboolean::from(reader.supports_range_search())) + }) +} + +#[no_mangle] +pub extern "system" fn Java_org_apache_paimon_index_vector_VectorIndexNative_rangeSearch( + env: JNIEnv, + _class: JClass, + ptr: jlong, + query: JFloatArray, + params: JObject, +) -> jobject { + range_call(env, |env| range_search(env, ptr, query, None, params, None)) +} + +#[no_mangle] +pub extern "system" fn Java_org_apache_paimon_index_vector_VectorIndexNative_rangeSearchWithRoaringFilter( + env: JNIEnv, + _class: JClass, + ptr: jlong, + query: JFloatArray, + params: JObject, + filter: JByteArray, +) -> jobject { + range_call(env, |env| { + range_search(env, ptr, query, None, params, Some(filter)) + }) +} + +#[no_mangle] +pub extern "system" fn Java_org_apache_paimon_index_vector_VectorIndexNative_rangeSearchBatch( + env: JNIEnv, + _class: JClass, + ptr: jlong, + queries: JFloatArray, + query_count: jint, + params: JObject, +) -> jobject { + range_call(env, |env| { + range_search(env, ptr, queries, Some(query_count), params, None) + }) +} + +#[no_mangle] +pub extern "system" fn Java_org_apache_paimon_index_vector_VectorIndexNative_rangeSearchBatchWithRoaringFilter( + env: JNIEnv, + _class: JClass, + ptr: jlong, + queries: JFloatArray, + query_count: jint, + params: JObject, + filter: JByteArray, +) -> jobject { + range_call(env, |env| { + range_search(env, ptr, queries, Some(query_count), params, Some(filter)) + }) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn native_lengths_and_counters_are_checked() { + assert_eq!(checked_array_length(0, "labels").unwrap(), 0); + assert_eq!( + checked_array_length(jint::MAX as usize, "labels").unwrap(), + jint::MAX + ); + assert!(checked_array_length(jint::MAX as usize + 1, "labels").is_err()); + assert_eq!(checked_counter(123, "rowsScanned").unwrap(), 123); + if usize::BITS == 64 { + assert_eq!( + checked_counter(jlong::MAX as usize, "listReads").unwrap(), + jlong::MAX + ); + assert!(checked_counter(usize::MAX, "listReads").is_err()); + } + } + + #[test] + fn external_metric_and_operator_codes_are_explicit() { + assert_eq!(parse_metric("l2").unwrap(), MetricType::L2); + assert_eq!(parse_metric("cosine").unwrap(), MetricType::Cosine); + assert_eq!( + parse_metric("inner_product").unwrap(), + MetricType::InnerProduct + ); + assert!(parse_metric("unknown").is_err()); + for (code, expected) in [ + (0, CutOperator::Ge), + (1, CutOperator::Gt), + (2, CutOperator::Le), + (3, CutOperator::Lt), + ] { + assert_eq!(cut_operator(code).unwrap(), expected); + } + assert!(cut_operator(-1).is_err()); + assert!(cut_operator(4).is_err()); + } +} diff --git a/python/README.md b/python/README.md new file mode 100644 index 00000000..75e11f18 --- /dev/null +++ b/python/README.md @@ -0,0 +1,102 @@ + + +# Python distance range search + +`VectorIndexReader.range_search` and `range_search_batch` add variable-length +results alongside the unchanged top-K `search` and `search_batch` APIs. IVF-Flat, +IVF-PQ, IVF-RQ, and IVF-SQ support L2, inner product, and cosine. Query +`reader.supports_range_search()` before selecting this API; DiskANN is unsupported. +With an older native library lacking range exports, existing imports and top-K +calls still work: capability returns `False`, and range calls/endpoint conversion +raise a clear error asking for a native-library upgrade. A partially available +range ABI is treated as unavailable rather than risking an unfreeable result. + +```python +from paimon_vindex import ( + DistanceBand, + DistanceEndpoint, + DistanceEndpointOp, + RangeSearchParams, +) + +band = DistanceBand.from_endpoints( + "l2", upper=DistanceEndpoint(2.0, DistanceEndpointOp.LE) +) +result = reader.range_search_batch(queries, RangeSearchParams(band, nprobe=8)) +labels, distances = result.query(0) +stats = result.stats[0] +``` + +## Distance contracts + +- `DistanceBand(metric, lower=None, upper=None)` describes a half-open + **internal-distance** interval `[lower, upper)`. Metric names are `l2`, + `inner_product`, and `cosine`. Internal distances are squared L2, negative inner + product, and cosine distance, respectively. +- `None` means structurally unbounded, not a finite sentinel. Equal finite cuts + describe a valid empty band. Core validates finite, ordered, metric-compatible + cuts during search; invalid bands raise `RuntimeError`. +- `DistanceBand.from_endpoints` accepts optional `DistanceEndpoint(value, op)` + objects: lower uses `GE`/`GT`, upper uses `LE`/`LT`. Endpoint values are public + distances (L2 square root, inner product, or cosine distance). Double literals + pass directly to core, which handles rounding, operator strictness, and inner + product direction reversal. Python performs no endpoint conversion math. +- Results retain internal distances. Quantized IVF variants return their own + distance estimates, not an additional exact-vector reranking. + +## Queries, filters, and results + +Single queries have shape `(dimension,)`; batches have shape +`(query_count, dimension)`. Inputs are converted to aligned, contiguous float32 after +shape and buffer-size checks. Core rejects zero-query batches. `nprobe` is a +positive, platform-sized integer and is independent of top-K parameters. + +Both methods accept `roaring_filter=` containing serialized **RoaringTreemap** +bytes, bytearray, or memoryview. `None` selects the unfiltered API; a serialized +empty treemap selects no rows. Empty bytes are not a serialized empty treemap +and are passed to core for validation. + +`RangeSearchResult` contains owned NumPy arrays `lims` (`uintp`), `labels` +(`int64`), and `distances` (`float32`), plus an immutable tuple `stats` and the +call-level `list_reads`. `query_count`, `hit_count`, and `query(index)` expose +the CSR shape. The query accessor returns label/distance slices and rejects +negative or out-of-range indices. Arrays and statistics remain valid after +reader closure and native result destruction. + +Each `RangeSearchStats` contains `lists_probed`, `rows_scanned`, +`rows_committed`, and `early_abandoned`. These are core's logical counters; +`list_reads` is call-level and should not be summed from per-query counters. +The native result handle is destroyed in `finally`, including conversion and +allocation failures, while the reader's callback-aware lock remains held. + +## Verification + +```sh +PYTHONPATH=python PAIMON_VINDEX_LIB_PATH=/path/to/native/library \ + python3 -m pytest python/tests +``` + +The standalone matrix tests run without shared fixtures. To additionally run +the exact cross-language oracle, set `PVI_RANGE_FIXTURES` to the output directory +from `cargo run -p paimon-vindex-ffi --example range_search_fixture -- DIR`. +Oracle tests open a fresh reader for each manifest entry and compare per-query +`(label, float32 bits)` multisets, CSR limits, and every statistics field exactly. +Core does not guarantee row order across calls; the Python wrapper preserves +the order it receives without sorting. diff --git a/python/paimon_vindex/__init__.py b/python/paimon_vindex/__init__.py index bba1ef07..881e6196 100644 --- a/python/paimon_vindex/__init__.py +++ b/python/paimon_vindex/__init__.py @@ -20,7 +20,7 @@ import threading from dataclasses import dataclass from enum import IntEnum -from typing import Mapping, Optional +from typing import Mapping, Optional, Tuple import numpy as np @@ -128,6 +128,160 @@ class IvfPqBatchTableReuseMode(IntEnum): AUTO = 2 +class DistanceEndpointOp(IntEnum): + GE = 0 + GT = 1 + LE = 2 + LT = 3 + + +@dataclass(frozen=True) +class DistanceEndpoint: + """A public-distance predicate literal, preserved as a double for core.""" + + value: float + op: DistanceEndpointOp + + def __post_init__(self): + try: + endpoint_op = DistanceEndpointOp(self.op) + except (TypeError, ValueError) as exc: + raise ValueError("distance endpoint operator is invalid") from exc + object.__setattr__(self, "value", float(self.value)) + object.__setattr__(self, "op", endpoint_op) + + def to_ffi(self): + return _ffi.PaimonVindexDistanceEndpoint(self.value, int(self.op)) + + +def _metric_code(metric): + for code, name in METRICS.items(): + if metric == name: + return code + raise ValueError("metric must be l2, inner_product, or cosine") + + +@dataclass(frozen=True) +class DistanceBand: + """Half-open internal-distance band; None denotes an unbounded side. + + Cuts use squared L2, negative inner product, or cosine distance. Core + validates the band during search. Use from_endpoints for public-distance + predicates instead of converting or rounding their literals in Python. + """ + + metric: str + lower: Optional[float] = None + upper: Optional[float] = None + + def __post_init__(self): + _metric_code(self.metric) + for name in ("lower", "upper"): + value = getattr(self, name) + if value is not None: + object.__setattr__(self, name, float(value)) + + @classmethod + def from_endpoints( + cls, + metric: str, + lower: Optional[DistanceEndpoint] = None, + upper: Optional[DistanceEndpoint] = None, + ): + """Delegate GE/GT lower and LE/LT upper predicates to core. + + Literals are public distances: sqrtf(squared L2), inner product, + or cosine distance. No float32 rounding precedes the native call. + """ + _require_range_search() + metric_code = _metric_code(metric) + for endpoint in (lower, upper): + if endpoint is not None and not isinstance(endpoint, DistanceEndpoint): + raise TypeError("endpoints must be DistanceEndpoint or None") + raw_lower = lower.to_ffi() if lower is not None else None + raw_upper = upper.to_ffi() if upper is not None else None + band = _ffi.PaimonVindexDistanceBand() + rc = lib.paimon_vindex_distance_band_from_endpoints( + metric_code, + ctypes.byref(raw_lower) if raw_lower is not None else None, + ctypes.byref(raw_upper) if raw_upper is not None else None, + ctypes.byref(band), + ) + if rc != 0: + _check_error("distance endpoint conversion failed") + return cls( + metric, + band.lower if band.lower_kind else None, + band.upper if band.upper_kind else None, + ) + + def to_ffi(self): + return _ffi.PaimonVindexDistanceBand( + _metric_code(self.metric), + int(self.lower is not None), + self.lower if self.lower is not None else 0.0, + int(self.upper is not None), + self.upper if self.upper is not None else 0.0, + ) + + +@dataclass(frozen=True) +class RangeSearchParams: + band: DistanceBand + nprobe: int + + def __post_init__(self): + if not isinstance(self.band, DistanceBand): + raise TypeError("band must be a DistanceBand") + object.__setattr__( + self, "nprobe", _size_t(self.nprobe, "nprobe", allow_zero=False) + ) + + def to_ffi(self): + return _ffi.PaimonVindexRangeSearchParams(self.band.to_ffi(), self.nprobe) + + +@dataclass(frozen=True) +class RangeSearchStats: + """Per-query logical counters copied from the native result.""" + + lists_probed: int + rows_scanned: int + rows_committed: int + early_abandoned: int + + +@dataclass(frozen=True) +class RangeSearchResult: + """Owned CSR arrays and per-query statistics, independent of the reader. + + Distances retain the index's internal distance representation. list_reads + counts call-level list reads, not the sum of per-query lists_probed. + """ + + lims: np.ndarray + labels: np.ndarray + distances: np.ndarray + stats: Tuple[RangeSearchStats, ...] + list_reads: int + + @property + def query_count(self): + return len(self.lims) - 1 + + @property + def hit_count(self): + return len(self.labels) + + def query(self, index): + """Return label/distance slices for a nonnegative query index.""" + index = operator.index(index) + if not 0 <= index < self.query_count: + raise IndexError("range query index out of bounds") + start, end = int(self.lims[index]), int(self.lims[index + 1]) + return self.labels[start:end], self.distances[start:end] + + @dataclass(frozen=True) class VectorIndexMetadata: index_type: str @@ -284,6 +438,79 @@ def _float32_vector(value, name): return np.ascontiguousarray(array) +def _require_range_search(): + if not _ffi.RANGE_SEARCH_AVAILABLE: + raise RuntimeError( + "loaded native library does not support range search; " + "rebuild or upgrade the native library" + ) + + +def _range_buffer_length(length, itemsize, name): + length = _size_t(length, name, allow_zero=True) + if length > min(_SIZE_T_MAX, np.iinfo(np.intp).max) // itemsize: + raise ValueError(f"{name} is too large: buffer length overflow") + return length + + +def _range_queries(value, dimension, batch): + array = np.asarray(value) + name = "queries" if batch else "query" + ndim = 2 if batch else 1 + if array.ndim != ndim: + raise ValueError(f"{name} must be a {ndim}-dimensional float32 array") + if array.shape[-1] != dimension: + raise RuntimeError( + f"{name} dimension {array.shape[-1]} does not match " + f"index dimension {dimension}" + ) + _range_buffer_length(array.size, ctypes.sizeof(ctypes.c_float), name) + return np.require(array, dtype=np.float32, requirements=["C", "A"]) + + +def _range_array_copy(pointer, length, dtype, name): + _range_buffer_length(length, np.dtype(dtype).itemsize, name) + if not length: + return np.empty(0, dtype=dtype) + if not pointer: + raise RuntimeError(f"range result {name} pointer is null") + return np.ctypeslib.as_array(pointer, shape=(length,)).copy() + + +def _range_result_copy(handle, query_count): + view = _ffi.PaimonVindexRangeSearchResultView() + if lib.paimon_vindex_range_search_result_view(handle, ctypes.byref(view)) != 0: + _check_error("failed to view range result") + if view.query_count != query_count: + raise RuntimeError("range result query count does not match input") + _range_buffer_length( + query_count, ctypes.sizeof(_ffi.PaimonVindexRangeSearchStats), "stats" + ) + lims = _range_array_copy(view.lims, query_count + 1, np.uintp, "lims") + if ( + lims[0] != 0 + or lims[-1] != view.hit_count + or np.any(lims[1:] < lims[:-1]) + ): + raise RuntimeError("range result has invalid CSR limits") + labels = _range_array_copy(view.labels, view.hit_count, np.int64, "labels") + distances = _range_array_copy( + view.distances, view.hit_count, np.float32, "distances" + ) + if query_count and not view.stats: + raise RuntimeError("range result stats pointer is null") + stats = tuple( + RangeSearchStats( + view.stats[index].lists_probed, + view.stats[index].rows_scanned, + view.stats[index].rows_committed, + view.stats[index].early_abandoned, + ) + for index in range(query_count) + ) + return RangeSearchResult(lims, labels, distances, stats, view.list_reads) + + def _int64_vector(value, name): array = np.asarray(value, dtype=np.int64) if array.ndim != 1: @@ -772,6 +999,79 @@ def _filter_args(self, filter_bytes): return None, 0, None return _bytes_buffer(filter_bytes, "filter_bytes") + def supports_range_search(self): + """Ask core whether this reader supports IVF distance range search.""" + with self._native_handle_lock: + self._require_open() + if not _ffi.RANGE_SEARCH_AVAILABLE: + return False + supported = ctypes.c_int() + rc = lib.paimon_vindex_reader_supports_range_search( + self._handle, ctypes.byref(supported) + ) + if rc != 0: + _check_error("range search capability check failed") + return bool(supported.value) + + def range_search(self, query, params: RangeSearchParams, roaring_filter=None): + """Search one query, returning an owned one-query CSR result. + + roaring_filter is optional serialized RoaringTreemap bytes. An empty + byte string is still a supplied filter and is validated by core. + """ + return self._range_search(query, params, roaring_filter, batch=False) + + def range_search_batch( + self, queries, params: RangeSearchParams, roaring_filter=None + ): + """Search a query matrix; core rejects a zero-query batch.""" + return self._range_search(queries, params, roaring_filter, batch=True) + + def _range_search(self, value, params, roaring_filter, *, batch): + _require_range_search() + if not isinstance(params, RangeSearchParams): + raise TypeError("params must be RangeSearchParams") + queries = _range_queries(value, self._metadata.dimension, batch) + query_count = queries.shape[0] if batch else 1 + ffi_params = params.to_ffi() + with self._native_handle_lock: + self._require_open() + args = [ + self._handle, + queries.ctypes.data_as(ctypes.POINTER(ctypes.c_float)), + queries.size, + ] + if batch: + args.append(query_count) + args.append(ffi_params) + if roaring_filter is None: + search = ( + lib.paimon_vindex_reader_range_search_batch + if batch else lib.paimon_vindex_reader_range_search + ) + else: + filter_buf, filter_len, _ = _bytes_buffer( + roaring_filter, "roaring_filter" + ) + _range_buffer_length(filter_len, 1, "roaring_filter") + args.extend((filter_buf, filter_len)) + search = ( + lib.paimon_vindex_reader_range_search_batch_with_roaring_filter + if batch + else lib.paimon_vindex_reader_range_search_with_roaring_filter + ) + handle = ctypes.c_void_p() + try: + rc = search(*args, ctypes.byref(handle)) + if rc != 0: + _check_error("range search failed") + if not handle: + raise RuntimeError("range search returned a null result") + return _range_result_copy(handle, query_count) + finally: + if handle: + lib.paimon_vindex_range_search_result_destroy(handle) + def search(self, query, params: SearchParams, filter_bytes=None): query = _float32_vector(query, "query") if query.shape[0] != self._metadata.dimension: @@ -875,7 +1175,13 @@ def __del__(self): __all__ = [ + "DistanceBand", + "DistanceEndpoint", + "DistanceEndpointOp", "IvfPqBatchTableReuseMode", + "RangeSearchParams", + "RangeSearchResult", + "RangeSearchStats", "SearchParams", "VectorIndexMetadata", "VectorIndexReadPlan", diff --git a/python/paimon_vindex/_ffi.py b/python/paimon_vindex/_ffi.py index 02cd900c..3ec3626c 100644 --- a/python/paimon_vindex/_ffi.py +++ b/python/paimon_vindex/_ffi.py @@ -23,6 +23,7 @@ POINTER, Structure, c_char_p, + c_double, c_float, c_int, c_int64, @@ -156,6 +157,51 @@ class PaimonVindexSearchParamsV2(Structure): ] +class PaimonVindexDistanceBand(Structure): + _fields_ = [ + ("metric", c_uint32), + ("lower_kind", c_uint32), + ("lower", c_float), + ("upper_kind", c_uint32), + ("upper", c_float), + ] + + +class PaimonVindexDistanceEndpoint(Structure): + _fields_ = [ + ("value", c_double), + ("op", c_uint32), + ] + + +class PaimonVindexRangeSearchParams(Structure): + _fields_ = [ + ("band", PaimonVindexDistanceBand), + ("nprobe", c_size_t), + ] + + +class PaimonVindexRangeSearchStats(Structure): + _fields_ = [ + ("lists_probed", c_size_t), + ("rows_scanned", c_size_t), + ("rows_committed", c_size_t), + ("early_abandoned", c_size_t), + ] + + +class PaimonVindexRangeSearchResultView(Structure): + _fields_ = [ + ("query_count", c_size_t), + ("hit_count", c_size_t), + ("lims", POINTER(c_size_t)), + ("labels", POINTER(c_int64)), + ("distances", POINTER(c_float)), + ("stats", POINTER(PaimonVindexRangeSearchStats)), + ("list_reads", c_size_t), + ] + + class PaimonVindexReaderOptions(Structure): _fields_ = [ ("memory_budget_bytes", c_size_t), @@ -338,3 +384,48 @@ class PaimonVindexReadPlan(Structure): c_size_t, ] lib.paimon_vindex_reader_search_batch_with_roaring_filter_v2.restype = c_int + +def _configure_range_api(): + signatures = ( + ("paimon_vindex_distance_band_from_endpoints", [ + c_uint32, + POINTER(PaimonVindexDistanceEndpoint), + POINTER(PaimonVindexDistanceEndpoint), + POINTER(PaimonVindexDistanceBand), + ], c_int), + ("paimon_vindex_reader_supports_range_search", [ + c_void_p, POINTER(c_int), + ], c_int), + ("paimon_vindex_reader_range_search", [ + c_void_p, POINTER(c_float), c_size_t, + PaimonVindexRangeSearchParams, POINTER(c_void_p), + ], c_int), + ("paimon_vindex_reader_range_search_with_roaring_filter", [ + c_void_p, POINTER(c_float), c_size_t, PaimonVindexRangeSearchParams, + POINTER(c_uint8), c_size_t, POINTER(c_void_p), + ], c_int), + ("paimon_vindex_reader_range_search_batch", [ + c_void_p, POINTER(c_float), c_size_t, c_size_t, + PaimonVindexRangeSearchParams, POINTER(c_void_p), + ], c_int), + ("paimon_vindex_reader_range_search_batch_with_roaring_filter", [ + c_void_p, POINTER(c_float), c_size_t, c_size_t, + PaimonVindexRangeSearchParams, POINTER(c_uint8), c_size_t, + POINTER(c_void_p), + ], c_int), + ("paimon_vindex_range_search_result_view", [ + c_void_p, POINTER(PaimonVindexRangeSearchResultView), + ], c_int), + ("paimon_vindex_range_search_result_destroy", [c_void_p], None), + ) + try: + functions = [(getattr(lib, name), args, result) for name, args, result in signatures] + except AttributeError: + return False + for function, args, result in functions: + function.argtypes = args + function.restype = result + return True + + +RANGE_SEARCH_AVAILABLE = _configure_range_api() diff --git a/python/tests/test_range_search.py b/python/tests/test_range_search.py new file mode 100644 index 00000000..ae2cdb1b --- /dev/null +++ b/python/tests/test_range_search.py @@ -0,0 +1,739 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +import ast +import ctypes +import io +import os +import struct +import subprocess +import sys +import threading +import textwrap +from collections import Counter +from pathlib import Path + +import numpy as np +import pytest + + +def test_range_public_api_is_exported_without_loading_native_library(): + source = Path(__file__).parents[1] / "paimon_vindex" / "__init__.py" + module = ast.parse(source.read_text()) + exports = next( + ast.literal_eval(statement.value) + for statement in module.body + if isinstance(statement, ast.Assign) + and any( + isinstance(target, ast.Name) and target.id == "__all__" + for target in statement.targets + ) + ) + assert { + "DistanceBand", + "DistanceEndpoint", + "DistanceEndpointOp", + "RangeSearchParams", + "RangeSearchResult", + "RangeSearchStats", + } <= set(exports) + + +@pytest.fixture(scope="module") +def vindex(): + import paimon_vindex + + return paimon_vindex + + +class BytesInput: + def __init__(self, data): + self.data = data + + def pread_many(self, ranges): + return [self.data[offset : offset + length] for offset, length in ranges] + + +def make_index(vindex, index_type="ivf_flat", metric="l2"): + count = 512 if index_type == "ivf_pq" else 128 + data = np.random.default_rng(1729).normal(size=(count, 16)).astype(np.float32) + labels = np.arange(len(data), dtype=np.int64) + (1 << 40) + options = { + "index.type": index_type, + "dimension": "16", + "metric": metric, + } + if index_type == "diskann": + options.update({ + "pq.m": "4", + "pq.bits": "4", + "diskann.max-degree": "8", + "diskann.build-search-list-size": "16", + }) + else: + options["nlist"] = "4" + if index_type == "ivf_pq": + options.update({"pq.m": "4", "use-opq": "false"}) + output = io.BytesIO() + training = vindex.VectorIndexTrainer.train(options, data) + with vindex.VectorIndexWriter(training) as writer: + writer.add_vectors(labels, data) + writer.write(output) + return output.getvalue(), data, labels + + +def roaring_allowlist(labels): + if not len(labels): + return struct.pack("= result.lims[:-1]) + assert len(result.stats) == len(queries) + assert 0 <= result.list_reads <= 4 + for array in (result.lims, result.labels, result.distances): + assert array.flags.owndata + for query_index, query in enumerate(queries): + with vindex.VectorIndexReader(BytesInput(payload)) as reader: + reference = reader.range_search(query, params, roaring_filter=filter_bytes) + actual_labels, actual_distances = result.query(query_index) + actual_order = np.argsort(actual_labels) + reference_order = np.argsort(reference.labels) + np.testing.assert_array_equal( + actual_labels[actual_order], reference.labels[reference_order] + ) + np.testing.assert_array_equal( + actual_distances[actual_order].view(np.uint32), + reference.distances[reference_order].view(np.uint32), + ) + assert set(actual_labels) <= set(allowed) + assert result.stats[query_index] == reference.stats[0] + stats = result.stats[query_index] + assert stats.rows_committed == len(actual_labels) + assert stats.rows_scanned >= stats.rows_committed + assert stats.early_abandoned <= stats.rows_scanned + if metric != "l2" or index_type in ("ivf_pq", "ivf_rq"): + assert stats.early_abandoned == 0 + if band_kind == "unbounded": + assert set(actual_labels) == set(allowed) + if band_kind == "empty" or filter_kind == "empty": + assert len(actual_labels) == 0 + if band.lower is not None: + assert np.all(actual_distances >= np.float32(band.lower)) + if band.upper is not None: + assert np.all(actual_distances < np.float32(band.upper)) + + +@pytest.mark.parametrize("batch", [False, True]) +@pytest.mark.parametrize("filter_bytes", [b"", b"invalid roaring"]) +def test_invalid_filter(vindex, flat_index, batch, filter_bytes): + payload, data, _ = flat_index + params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + with vindex.VectorIndexReader(BytesInput(payload)) as reader: + method = reader.range_search_batch if batch else reader.range_search + with pytest.raises(RuntimeError, match="[Rr]oaring|filter"): + method( + data[:2] if batch else data[0], params, roaring_filter=filter_bytes + ) + + +@pytest.mark.parametrize("lower, upper", [ + (float("nan"), None), (None, float("inf")), (-1.0, None), (2.0, 1.0), +]) +def test_invalid_bands_are_rejected_by_core(vindex, flat_index, lower, upper): + payload, data, _ = flat_index + params = vindex.RangeSearchParams(vindex.DistanceBand("l2", lower, upper), 4) + with vindex.VectorIndexReader(BytesInput(payload)) as reader: + with pytest.raises(RuntimeError): + reader.range_search(data[0], params) + with pytest.raises(RuntimeError): + reader.range_search_batch(data[:0], params) + + +@pytest.mark.parametrize("nprobe", [-1, 0, 1.5, ctypes.c_size_t(-1).value + 1]) +def test_nprobe_rejects_invalid_and_wrapping_values(vindex, nprobe): + with pytest.raises(ValueError, match="nprobe"): + vindex.RangeSearchParams(vindex.DistanceBand("l2"), nprobe) + + +@pytest.mark.parametrize("metric", ["l2", "inner_product", "cosine"]) +@pytest.mark.parametrize("lower_op", [0, 1]) +@pytest.mark.parametrize("upper_op", [2, 3]) +def test_endpoints_match_direct_core_conversion(vindex, metric, lower_op, upper_op): + ffi = vindex._ffi + lower = vindex.DistanceEndpoint(0.10000000000000002, lower_op) + upper = vindex.DistanceEndpoint(1.1000000000000003, upper_op) + band = vindex.DistanceBand.from_endpoints(metric, lower, upper) + raw_lower = ffi.PaimonVindexDistanceEndpoint(lower.value, lower_op) + raw_upper = ffi.PaimonVindexDistanceEndpoint(upper.value, upper_op) + expected = ffi.PaimonVindexDistanceBand() + assert ffi.lib.paimon_vindex_distance_band_from_endpoints( + {"l2": 0, "inner_product": 1, "cosine": 2}[metric], + ctypes.byref(raw_lower), ctypes.byref(raw_upper), ctypes.byref(expected), + ) == 0 + actual = band.to_ffi() + assert bytes(actual) == bytes(expected) + + +def test_empty_batch_and_query_accessor(vindex, flat_index): + payload, data, _ = flat_index + with vindex.VectorIndexReader(BytesInput(payload)) as reader: + params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + with pytest.raises(RuntimeError, match="query count must be greater than 0"): + reader.range_search_batch(data[:0], params) + result = reader.range_search(data[0], params) + for index in (-1, 1, ctypes.c_size_t(-1).value + 1): + with pytest.raises(IndexError): + result.query(index) + with pytest.raises(TypeError): + result.query(0.5) + + +@pytest.mark.parametrize("batch,shape,error", [ + (False, (), ValueError), (False, (1, 16), ValueError), + (False, (15,), RuntimeError), (False, (0,), RuntimeError), + (True, (16,), ValueError), (True, (1, 1, 16), ValueError), + (True, (1, 15), RuntimeError), (True, (0, 15), RuntimeError), +]) +def test_query_shapes(vindex, flat_index, batch, shape, error): + payload, _, _ = flat_index + params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + with vindex.VectorIndexReader(BytesInput(payload)) as reader: + method = reader.range_search_batch if batch else reader.range_search + with pytest.raises(error): + method(np.zeros(shape), params) + + +def test_strided_queries_and_filter_buffer_types(vindex, flat_index): + payload, data, labels = flat_index + queries = np.asfortranarray(data[:4].astype(np.float64))[::2] + params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + serialized = roaring_allowlist(labels[::3]) + for filter_bytes in (bytearray(serialized), memoryview(serialized)): + with vindex.VectorIndexReader(BytesInput(payload)) as reader: + result = reader.range_search_batch( + queries, params, roaring_filter=filter_bytes + ) + assert result.query_count == 2 + for query_index in range(2): + assert set(result.query(query_index)[0]) == set(labels[::3]) + + +def test_query_buffer_overflow_before_copy(vindex, flat_index): + payload, _, _ = flat_index + params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + huge = np.lib.stride_tricks.as_strided( + np.zeros(1, dtype=np.uint8), + shape=(np.iinfo(np.intp).max // 32, 16), strides=(0, 0), + ) + with vindex.VectorIndexReader(BytesInput(payload)) as reader: + with pytest.raises(ValueError, match="overflow|too large"): + reader.range_search_batch(huge, params) + + +@pytest.mark.parametrize("operation", ["single", "batch", "capability"]) +def test_closed_reader_and_reentry(vindex, flat_index, operation): + payload, data, _ = flat_index + reader = vindex.VectorIndexReader(BytesInput(payload)) + params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + invoke = { + "single": lambda: reader.range_search(data[0], params), + "batch": lambda: reader.range_search_batch(data[:2], params), + "capability": reader.supports_range_search, + }[operation] + with reader._native_handle_lock: + with pytest.raises(RuntimeError, match="reentrant"): + invoke() + reader.close() + with pytest.raises(RuntimeError, match="closed"): + invoke() + + +def test_diskann_is_unsupported(vindex): + payload, data, _ = make_index(vindex, "diskann") + with vindex.VectorIndexReader(BytesInput(payload)) as reader: + assert reader.supports_range_search() is False + params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + for query in (data[0], data[:2]): + method = ( + reader.range_search if query.ndim == 1 else reader.range_search_batch + ) + with pytest.raises( + RuntimeError, match="[Uu]nsupported|[Dd]isk[Aa][Nn][Nn]|range" + ): + method(query, params) + + +@pytest.mark.parametrize("failure", [ + "view", "copy", "stats", "search", "null", "null_distances", + "null_stats", "null_lims", "length", "overflow", "lims", +]) +def test_result_destroyed_on_failure(vindex, flat_index, monkeypatch, failure): + payload, data, _ = flat_index + params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + ffi = vindex._ffi + destroyed = [] + destroy = ffi.lib.paimon_vindex_range_search_result_destroy + view_result = ffi.lib.paimon_vindex_range_search_result_view + search = ffi.lib.paimon_vindex_reader_range_search + + def track_destroy(handle): + destroyed.append(handle.value) + destroy(handle) + + def alter_view(handle, output): + status = view_result(handle, output) + raw = ctypes.cast( + output, ctypes.POINTER(ffi.PaimonVindexRangeSearchResultView) + ).contents + if failure == "view": + return -1 + if failure == "null": + raw.labels = ctypes.POINTER(ctypes.c_int64)() + if failure == "null_distances": + raw.distances = ctypes.POINTER(ctypes.c_float)() + if failure == "null_stats": + raw.stats = ctypes.POINTER(ffi.PaimonVindexRangeSearchStats)() + if failure == "null_lims": + raw.lims = ctypes.POINTER(ctypes.c_size_t)() + if failure == "length": + raw.query_count = 2 + if failure == "overflow": + raw.hit_count = ctypes.c_size_t(-1).value + raw.lims[raw.query_count] = raw.hit_count + if failure == "lims": + raw.lims[0] = 1 + return status + + def failed_search(*args): + assert search(*args) == 0 + return -1 + + def fail_copy(*args, **kwargs): + raise MemoryError("copy failed") + + monkeypatch.setattr( + ffi.lib, "paimon_vindex_range_search_result_destroy", track_destroy + ) + monkeypatch.setattr(ffi.lib, "paimon_vindex_range_search_result_view", alter_view) + if failure == "search": + monkeypatch.setattr(ffi.lib, "paimon_vindex_reader_range_search", failed_search) + if failure == "copy": + monkeypatch.setattr(np.ctypeslib, "as_array", fail_copy) + if failure == "stats": + monkeypatch.setattr(vindex, "RangeSearchStats", fail_copy) + with vindex.VectorIndexReader(BytesInput(payload)) as reader: + with pytest.raises((MemoryError, RuntimeError, ValueError)): + reader.range_search(data[0], params) + assert len(destroyed) == 1 and destroyed[0] + + +def test_result_copies_survive_native_destruction(vindex, flat_index, monkeypatch): + payload, data, _ = flat_index + params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + ffi = vindex._ffi + destroy = ffi.lib.paimon_vindex_range_search_result_destroy + destroyed = [] + + def poison_destroy(handle): + view = ffi.PaimonVindexRangeSearchResultView() + assert ffi.lib.paimon_vindex_range_search_result_view( + handle, ctypes.byref(view) + ) == 0 + ctypes.memset( + view.lims, 255, (view.query_count + 1) * ctypes.sizeof(ctypes.c_size_t) + ) + ctypes.memset(view.labels, 255, view.hit_count * ctypes.sizeof(ctypes.c_int64)) + ctypes.memset( + view.distances, 255, view.hit_count * ctypes.sizeof(ctypes.c_float) + ) + ctypes.memset( + view.stats, 255, + view.query_count * ctypes.sizeof(ffi.PaimonVindexRangeSearchStats), + ) + destroyed.append(handle.value) + destroy(handle) + + monkeypatch.setattr( + ffi.lib, "paimon_vindex_range_search_result_destroy", poison_destroy + ) + with vindex.VectorIndexReader(BytesInput(payload)) as reader: + result = reader.range_search_batch(data[:2], params) + assert len(destroyed) == 1 + np.testing.assert_array_equal(result.lims, [0, 128, 256]) + assert np.all(result.labels >= (1 << 40)) + assert np.all(np.isfinite(result.distances)) + assert [stats.rows_committed for stats in result.stats] == [128, 128] + + +@pytest.mark.parametrize("metric", ["l2", "inner_product", "cosine"]) +@pytest.mark.parametrize("side", ["lower", "upper"]) +@pytest.mark.parametrize("value", [float("nan"), float("inf"), -float("inf")]) +def test_nonfinite_endpoints_are_core_errors(vindex, metric, side, value): + endpoint = vindex.DistanceEndpoint(value, 0 if side == "lower" else 3) + with pytest.raises(RuntimeError, match="finite"): + vindex.DistanceBand.from_endpoints(metric, **{side: endpoint}) + + +@pytest.mark.parametrize("side,op", [ + ("lower", 2), ("lower", 3), ("upper", 0), ("upper", 1), +]) +def test_wrong_side_endpoint_operators_are_core_errors(vindex, side, op): + with pytest.raises(RuntimeError, match="endpoint"): + vindex.DistanceBand.from_endpoints( + "l2", **{side: vindex.DistanceEndpoint(1.0, op)} + ) + + +@pytest.mark.parametrize("op", [-1, 4, 1 << 32]) +def test_endpoint_operator_cannot_wrap(vindex, op): + with pytest.raises(ValueError, match="operator"): + vindex.DistanceEndpoint(1.0, op) + + +def test_nullable_and_extreme_endpoints(vindex): + for metric in ("l2", "inner_product", "cosine"): + assert vindex.DistanceBand.from_endpoints(metric) == vindex.DistanceBand(metric) + assert vindex.DistanceBand.from_endpoints( + metric, lower=vindex.DistanceEndpoint(0.1, vindex.DistanceEndpointOp.GE) + ).metric == metric + assert vindex.DistanceBand.from_endpoints( + metric, upper=vindex.DistanceEndpoint(1.1, vindex.DistanceEndpointOp.LE) + ).metric == metric + with pytest.raises(RuntimeError): + vindex.DistanceBand.from_endpoints( + "l2", lower=vindex.DistanceEndpoint(1e300, 0) + ) + with pytest.raises(TypeError, match="DistanceEndpoint"): + vindex.DistanceBand.from_endpoints("l2", lower=(1.0, 0)) + + +def test_metric_and_parameter_validation(vindex, flat_index): + for metric in ("unknown", -1, 1 << 32, None): + with pytest.raises(ValueError, match="metric"): + vindex.DistanceBand(metric) + with pytest.raises(TypeError, match="band"): + vindex.RangeSearchParams("l2", 4) + payload, data, _ = flat_index + with vindex.VectorIndexReader(BytesInput(payload)) as reader: + with pytest.raises(TypeError, match="RangeSearchParams"): + reader.range_search(data[0], vindex.SearchParams.ivf(4, 4)) + with pytest.raises(RuntimeError, match="metric"): + reader.range_search( + data[0], vindex.RangeSearchParams(vindex.DistanceBand("cosine"), 4) + ) + with pytest.raises(ValueError, match="bytes"): + reader.range_search( + data[0], vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4), + roaring_filter="not bytes", + ) + + +def test_range_callback_reentry(vindex, flat_index): + payload, data, _ = flat_index + params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + + class ReentrantInput(BytesInput): + operation = None + + def pread_many(self, ranges): + if self.operation is not None: + self.operation() + return super().pread_many(ranges) + + source = ReentrantInput(payload) + with vindex.VectorIndexReader(source) as reader: + errors = [] + + def reenter(): + for operation in ( + reader.supports_range_search, + lambda: reader.range_search(data[0], params), + lambda: reader.range_search_batch(data[:2], params), + reader.close, + ): + with pytest.raises(RuntimeError, match="reentrant"): + operation() + errors.append(True) + + source.operation = reenter + result = reader.range_search(data[0], params) + assert errors and result.hit_count == len(data) + + +def test_range_close_waits_through_result_copy(vindex, flat_index, monkeypatch): + payload, data, _ = flat_index + reader = vindex.VectorIndexReader(BytesInput(payload)) + params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + copy_entered = threading.Event() + release_copy = threading.Event() + close_entered = threading.Event() + close_done = threading.Event() + errors = [] + result_view = vindex._ffi.lib.paimon_vindex_range_search_result_view + + def blocking_view(*args): + status = result_view(*args) + copy_entered.set() + assert release_copy.wait(timeout=5) + return status + + def search(): + try: + reader.range_search(data[0], params) + except Exception as error: + errors.append(error) + + def close(): + close_entered.set() + reader.close() + close_done.set() + + monkeypatch.setattr( + vindex._ffi.lib, "paimon_vindex_range_search_result_view", blocking_view + ) + search_thread = threading.Thread(target=search) + close_thread = threading.Thread(target=close) + search_thread.start() + try: + assert copy_entered.wait(timeout=5) + close_thread.start() + assert close_entered.wait(timeout=5) + assert not close_done.wait(timeout=0.05) + finally: + release_copy.set() + search_thread.join(timeout=5) + if close_thread.ident is not None: + close_thread.join(timeout=5) + reader.close() + assert not search_thread.is_alive() and not close_thread.is_alive() + assert close_done.is_set() and not errors + + +def oracle_cases(): + root = os.environ.get("PVI_RANGE_FIXTURES") + if not root: + return [] + directory = Path(root) + return [ + (directory, *line.split()) + for line in (directory / "manifest.txt").read_text().splitlines() + if line.strip() + ] + + +@pytest.mark.parametrize( + "directory,case_name,index_filename", oracle_cases(), ids=str +) +def test_core_oracle(vindex, directory, case_name, index_filename): + tokens = iter((directory / f"{case_name}.expected").read_text().split()) + + def integers(count): + return [int(next(tokens)) for _ in range(count)] + + ( + dimension, metric, query_count, nprobe, lower_kind, lower_bits, + upper_kind, upper_bits, filter_len, hit_count, list_reads, + ) = integers(11) + queries = ( + np.array(integers(dimension * query_count), dtype=np.uint32) + .view(np.float32) + .reshape(query_count, dimension) + ) + filter_bytes = bytes(integers(filter_len)) if filter_len else None + expected_lims = np.array(integers(query_count + 1), dtype=np.uintp) + expected_labels = np.array(integers(hit_count), dtype=np.int64) + expected_distances = np.array(integers(hit_count), dtype=np.uint32) + expected_stats = [tuple(integers(4)) for _ in range(query_count)] + assert next(tokens, None) is None + assert lower_kind in (0, 1) and upper_kind in (0, 1) + lower = ( + struct.unpack(" Date: Mon, 21 Sep 2026 12:07:22 +0800 Subject: [PATCH 2/5] fix: link C binding tests with libm The range oracle tests use fminf/fmaxf, which require an explicit libm dependency when linking on Linux. Link the C test executable against m on UNIX without changing the binding or core behavior. Validated with CMake, -Wall -Wextra -Werror, all C tests including 168 core oracle cases, cargo fmt, and the ASF header check. Linux validation follows in PR CI. --- c/CMakeLists.txt | 3 +++ 1 file changed, 3 insertions(+) diff --git a/c/CMakeLists.txt b/c/CMakeLists.txt index 2e989e1d..beb81a1b 100644 --- a/c/CMakeLists.txt +++ b/c/CMakeLists.txt @@ -35,3 +35,6 @@ endif() add_executable(test_vindex test_vindex.c) target_include_directories(test_vindex PRIVATE ${CMAKE_SOURCE_DIR}/../include) target_link_libraries(test_vindex ${PAIMON_VINDEX_FFI_LIB}) +if(UNIX) + target_link_libraries(test_vindex m) +endif() From d4a2d3cb26d737f30b3f1f3ef56f161cb378ff2a Mon Sep 17 00:00:00 2001 From: Junrui Lee Date: Mon, 21 Sep 2026 14:30:13 +0800 Subject: [PATCH 3/5] fix: address range search binding review feedback --- cpp/test_vindex.cpp | 30 +++++++ include/paimon_vindex.hpp | 19 ---- .../index/vector/VectorRangeSearchResult.java | 74 ++++++++++++--- .../vector/VectorIndexRangeSearchTest.java | 89 +++++++++++++++++++ jni/src/range.rs | 6 +- 5 files changed, 184 insertions(+), 34 deletions(-) diff --git a/cpp/test_vindex.cpp b/cpp/test_vindex.cpp index f504b87b..8dde10f0 100644 --- a/cpp/test_vindex.cpp +++ b/cpp/test_vindex.cpp @@ -540,6 +540,36 @@ static void test_range_matrix() { assert_range_error([&] { reader.range_search_batch(nullptr, 0, SIZE_MAX / RANGE_DIMENSION + 1, params); }); + assert_range_error([&] { + reader.range_search_with_roaring_filter( + nullptr, RANGE_DIMENSION, params, range_filter, sizeof(range_filter)); + }); + assert_range_error([&] { + reader.range_search_with_roaring_filter( + std::vector{}, params, range_filter, sizeof(range_filter)); + }); + assert_range_error([&] { + reader.range_search_with_roaring_filter( + query.data(), SIZE_MAX, params, range_filter, sizeof(range_filter)); + }); + assert_range_error([&] { + reader.range_search_batch_with_roaring_filter( + nullptr, queries.size(), RANGE_QUERY_COUNT, params, + range_filter, sizeof(range_filter)); + }); + assert_range_error([&] { + reader.range_search_batch_with_roaring_filter( + queries, 2, params, range_filter, sizeof(range_filter)); + }); + assert_range_error([&] { + reader.range_search_batch_with_roaring_filter( + nullptr, 0, 0, params, range_filter, sizeof(range_filter)); + }); + assert_range_error([&] { + reader.range_search_batch_with_roaring_filter( + queries.data(), 0, SIZE_MAX / RANGE_DIMENSION + 1, params, + range_filter, sizeof(range_filter)); + }); assert_range_error([&] { reader.range_search_with_roaring_filter(query, params, nullptr, 1); }); assert_range_error([&] { reader.range_search_with_roaring_filter(query, params, range_filter, 1); }); auto invalid = params; diff --git a/include/paimon_vindex.hpp b/include/paimon_vindex.hpp index 9cac0ef0..72424073 100644 --- a/include/paimon_vindex.hpp +++ b/include/paimon_vindex.hpp @@ -669,7 +669,6 @@ class Reader { RangeSearchResult range_search( const float* query, size_t query_len, RangeSearchParams params) { std::lock_guard lock(native_handle_mutex_); - validate_range_queries(query, query_len, 1); PaimonVindexRangeSearchResult* raw = nullptr; int status = paimon_vindex_reader_range_search( require_open(), query, query_len, params.to_ffi(), &raw); @@ -684,7 +683,6 @@ class Reader { const float* query, size_t query_len, RangeSearchParams params, const uint8_t* filter, size_t filter_len) { std::lock_guard lock(native_handle_mutex_); - validate_range_queries(query, query_len, 1); PaimonVindexRangeSearchResult* raw = nullptr; int status = paimon_vindex_reader_range_search_with_roaring_filter( require_open(), query, query_len, params.to_ffi(), filter, filter_len, &raw); @@ -700,7 +698,6 @@ class Reader { RangeSearchResult range_search_batch( const float* queries, size_t queries_len, size_t query_count, RangeSearchParams params) { std::lock_guard lock(native_handle_mutex_); - validate_range_queries(queries, queries_len, query_count); PaimonVindexRangeSearchResult* raw = nullptr; int status = paimon_vindex_reader_range_search_batch( require_open(), queries, queries_len, query_count, params.to_ffi(), &raw); @@ -716,7 +713,6 @@ class Reader { const float* queries, size_t queries_len, size_t query_count, RangeSearchParams params, const uint8_t* filter, size_t filter_len) { std::lock_guard lock(native_handle_mutex_); - validate_range_queries(queries, queries_len, query_count); PaimonVindexRangeSearchResult* raw = nullptr; int status = paimon_vindex_reader_range_search_batch_with_roaring_filter( require_open(), queries, queries_len, query_count, params.to_ffi(), filter, filter_len, &raw); @@ -815,21 +811,6 @@ class Reader { } private: - void validate_range_queries(const float* queries, size_t queries_len, size_t query_count) const { - PaimonVindexMetadata metadata{}; - check(paimon_vindex_reader_metadata(require_open(), &metadata)); - if (metadata.dimension != 0 && - query_count > std::numeric_limits::max() / metadata.dimension) { - throw Error("range query dimensions overflow"); - } - if (queries_len != query_count * metadata.dimension) { - throw Error("range query length does not match dimension and query count"); - } - if (queries_len != 0 && !queries) { - throw Error("range queries must not be null for a nonempty input"); - } - } - PaimonVindexReaderHandle* require_open() const { if (!handle_) throw Error("vector index reader is closed"); return handle_; diff --git a/java/src/main/java/org/apache/paimon/index/vector/VectorRangeSearchResult.java b/java/src/main/java/org/apache/paimon/index/vector/VectorRangeSearchResult.java index d0902983..4aeb0eab 100644 --- a/java/src/main/java/org/apache/paimon/index/vector/VectorRangeSearchResult.java +++ b/java/src/main/java/org/apache/paimon/index/vector/VectorRangeSearchResult.java @@ -23,7 +23,8 @@ /** * CSR range-search output in core scan order, with raw distances and no sorting or top-K cap. * IVF-Flat distances are exact; SQ, PQ and RQ return their core distance estimates. Each query - * occupies [lims[query], lims[query + 1]); all arrays are defensively copied. + * occupies [lims[query], lims[query + 1]). Public construction and array access make defensive + * copies; native construction takes exclusive ownership of its arrays. */ public final class VectorRangeSearchResult { @@ -36,6 +37,27 @@ public final class VectorRangeSearchResult { private final long[] earlyAbandoned; private final long listReads; + static VectorRangeSearchResult fromNative( + long[] labels, + float[] distances, + long[] lims, + long[] listsProbed, + long[] rowsScanned, + long[] rowsCommitted, + long[] earlyAbandoned, + long listReads) { + return new VectorRangeSearchResult( + labels, + distances, + lims, + listsProbed, + rowsScanned, + rowsCommitted, + earlyAbandoned, + listReads, + false); + } + public VectorRangeSearchResult( long[] labels, float[] distances, @@ -45,9 +67,34 @@ public VectorRangeSearchResult( long[] rowsCommitted, long[] earlyAbandoned, long listReads) { - this.labels = Objects.requireNonNull(labels, "labels").clone(); - this.distances = Objects.requireNonNull(distances, "distances").clone(); - this.lims = Objects.requireNonNull(lims, "lims").clone(); + this( + labels, + distances, + lims, + listsProbed, + rowsScanned, + rowsCommitted, + earlyAbandoned, + listReads, + true); + } + + private VectorRangeSearchResult( + long[] labels, + float[] distances, + long[] lims, + long[] listsProbed, + long[] rowsScanned, + long[] rowsCommitted, + long[] earlyAbandoned, + long listReads, + boolean copyArrays) { + Objects.requireNonNull(labels, "labels"); + Objects.requireNonNull(distances, "distances"); + Objects.requireNonNull(lims, "lims"); + this.labels = copyArrays ? labels.clone() : labels; + this.distances = copyArrays ? distances.clone() : distances; + this.lims = copyArrays ? lims.clone() : lims; if (this.labels.length != this.distances.length || this.lims.length == 0 || this.lims[0] != 0 @@ -60,10 +107,10 @@ public VectorRangeSearchResult( throw new IllegalArgumentException("invalid CSR limits"); } } - this.listsProbed = copyCounters(listsProbed, "listsProbed"); - this.rowsScanned = copyCounters(rowsScanned, "rowsScanned"); - this.rowsCommitted = copyCounters(rowsCommitted, "rowsCommitted"); - this.earlyAbandoned = copyCounters(earlyAbandoned, "earlyAbandoned"); + this.listsProbed = validatedCounters(listsProbed, "listsProbed", copyArrays); + this.rowsScanned = validatedCounters(rowsScanned, "rowsScanned", copyArrays); + this.rowsCommitted = validatedCounters(rowsCommitted, "rowsCommitted", copyArrays); + this.earlyAbandoned = validatedCounters(earlyAbandoned, "earlyAbandoned", copyArrays); if (listReads < 0) { throw new IllegalArgumentException("listReads must be non-negative"); } @@ -130,16 +177,17 @@ private void checkQueryIndex(int queryIndex) { } } - private long[] copyCounters(long[] counters, String name) { - long[] copy = Objects.requireNonNull(counters, name).clone(); - if (copy.length != queryCount()) { + private long[] validatedCounters(long[] counters, String name, boolean copyArrays) { + Objects.requireNonNull(counters, name); + long[] values = copyArrays ? counters.clone() : counters; + if (values.length != queryCount()) { throw new IllegalArgumentException(name + " length must equal queryCount"); } - for (long value : copy) { + for (long value : values) { if (value < 0) { throw new IllegalArgumentException(name + " must be non-negative"); } } - return copy; + return values; } } diff --git a/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeSearchTest.java b/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeSearchTest.java index e106b741..544d8846 100644 --- a/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeSearchTest.java +++ b/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeSearchTest.java @@ -22,6 +22,7 @@ import static org.apache.paimon.index.vector.VectorDistanceBand.CutOperator.LE; import static org.apache.paimon.index.vector.VectorDistanceBand.CutOperator.LT; +import java.lang.reflect.Field; import java.nio.ByteBuffer; import java.nio.ByteOrder; import java.util.Arrays; @@ -42,6 +43,7 @@ public static void main(String[] args) { } static void testValueTypes() { + testNativeResultOwnership(); VectorDistanceBand band = new VectorDistanceBand("l2", null, 4.0f); check(band.lower() == null && band.upper() == 4.0f, "structural bounds"); check("l2".equals(band.metric()), "metric"); @@ -140,6 +142,93 @@ static void testValueTypes() { expect(NullPointerException.class, () -> closed.rangeSearch(new float[1], params, null)); } + private static void testNativeResultOwnership() { + long[] labels = {LABEL_BASE, 7}; + float[] distances = {1, 3}; + long[] lims = {0, 0, 2}; + long[] listsProbed = {0, 1}; + long[] rowsScanned = {0, 3}; + long[] rowsCommitted = {0, 2}; + long[] earlyAbandoned = {0, 1}; + VectorRangeSearchResult result = + VectorRangeSearchResult.fromNative( + labels, + distances, + lims, + listsProbed, + rowsScanned, + rowsCommitted, + earlyAbandoned, + 1); + checkOwnedArray(result, "labels", labels); + checkOwnedArray(result, "distances", distances); + checkOwnedArray(result, "lims", lims); + checkOwnedArray(result, "listsProbed", listsProbed); + checkOwnedArray(result, "rowsScanned", rowsScanned); + checkOwnedArray(result, "rowsCommitted", rowsCommitted); + checkOwnedArray(result, "earlyAbandoned", earlyAbandoned); + result.labels()[0] = -1; + result.distances()[0] = -1; + result.lims()[1] = 2; + result.listsProbed()[1] = -1; + result.rowsScanned()[1] = -1; + result.rowsCommitted()[1] = -1; + result.earlyAbandoned()[1] = -1; + result.labelsForQuery(1)[0] = -1; + result.distancesForQuery(1)[0] = -1; + check(result.queryCount() == 2 && result.listReads() == 1, "owned result shape"); + check(result.labelsForQuery(0).length == 0, "owned empty query"); + check(result.labelsForQuery(1)[0] == LABEL_BASE, "owned labels remain defensive"); + check(result.distancesForQuery(1)[0] == 1, "owned distances remain defensive"); + check(result.lims()[1] == 0 && result.listsProbed()[1] == 1, "owned CSR and stats"); + check( + result.rowsScanned()[1] == 3 + && result.rowsCommitted()[1] == 2 + && result.earlyAbandoned()[1] == 1, + "owned counters remain defensive"); + expect( + NullPointerException.class, + () -> + VectorRangeSearchResult.fromNative( + null, distances, lims, listsProbed, rowsScanned, + rowsCommitted, earlyAbandoned, 1)); + expect( + IllegalArgumentException.class, + () -> + VectorRangeSearchResult.fromNative( + labels, distances, new long[] {0, 3, 2}, listsProbed, + rowsScanned, rowsCommitted, earlyAbandoned, 1)); + expect( + IllegalArgumentException.class, + () -> + VectorRangeSearchResult.fromNative( + labels, distances, lims, new long[0], rowsScanned, + rowsCommitted, earlyAbandoned, 1)); + expect( + IllegalArgumentException.class, + () -> + VectorRangeSearchResult.fromNative( + labels, distances, lims, listsProbed, new long[] {0, -1}, + rowsCommitted, earlyAbandoned, 1)); + expect( + IllegalArgumentException.class, + () -> + VectorRangeSearchResult.fromNative( + labels, distances, lims, listsProbed, rowsScanned, + rowsCommitted, earlyAbandoned, -1)); + } + + private static void checkOwnedArray( + VectorRangeSearchResult result, String name, Object expected) { + try { + Field field = VectorRangeSearchResult.class.getDeclaredField(name); + field.setAccessible(true); + check(field.get(result) == expected, "native result must own " + name + " without copying"); + } catch (ReflectiveOperationException error) { + throw new AssertionError(error); + } + } + static void testNative() { for (String indexType : new String[] {"ivf_flat", "ivf_sq", "ivf_pq", "ivf_rq"}) { for (String metric : new String[] {"l2", "cosine", "inner_product"}) { diff --git a/jni/src/range.rs b/jni/src/range.rs index 4fbcd0e0..7e3d1b2e 100644 --- a/jni/src/range.rs +++ b/jni/src/range.rs @@ -207,9 +207,10 @@ fn build_range_result(env: &mut JNIEnv, result: RangeSearchResult) -> Result Result Date: Mon, 21 Sep 2026 18:19:03 +0800 Subject: [PATCH 4/5] fix: add zero-copy range result consumption Add indexed Java result access, bound JNI metadata staging, and cover memory and ownership regressions. --- .github/workflows/ci.yml | 2 + docs/range-search.html | 13 +- .../index/vector/VectorRangeSearchResult.java | 53 ++++- .../vector/VectorIndexRangeOracleTest.java | 38 ++-- .../vector/VectorIndexRangeSearchTest.java | 187 ++++++++++++++++++ jni/src/range.rs | 62 ++++-- python/tests/test_range_search.py | 86 +++++++- 7 files changed, 393 insertions(+), 48 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 0870f10a..0a5fbbd2 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -215,6 +215,8 @@ jobs: - name: Test JNI native behavior run: | + java -Xmx36m -cp java/target/test-classes:java/target/classes org.apache.paimon.index.vector.VectorIndexRangeSearchTest \ + --zero-copy-memory java -cp java/target/test-classes:java/target/classes org.apache.paimon.index.vector.VectorIndexNativeValidationTest \ "$(pwd)/target/release/libpaimon_vindex_jni.so" java -cp java/target/test-classes:java/target/classes org.apache.paimon.index.vector.VectorIndexNativePanicBoundaryTest \ diff --git a/docs/range-search.html b/docs/range-search.html index 568d0aae..b3d636e0 100644 --- a/docs/range-search.html +++ b/docs/range-search.html @@ -101,7 +101,7 @@

Language bindings and ownership

- +
BindingAPI and lifetime
C ABI / generated headerpaimon_vindex_reader_range_search, _range_search_batch, and both _with_roaring_filter variants return an owned opaque result through an output pointer. paimon_vindex_range_search_result_view borrows its buffers; paimon_vindex_range_search_result_destroy releases them. Query/filter lengths are explicit. The generated include/paimon_vindex.h is rebuilt by cbindgen, not maintained manually.
C++Reader::range_search, Reader::range_search_batch, and their _with_roaring_filter variants copy CSR buffers into owning vectors. A native-result RAII guard handles cleanup, including allocation exceptions. Existing top-K methods remain unchanged.
JNI / JavaVectorIndexReader.rangeSearch and rangeSearchBatch accept VectorRangeSearchParams and optional filter bytes. VectorRangeSearchResult owns defensively copied Java arrays. Native lengths are checked against JVM array limits; offsets, labels and counters use long.
JNI / JavaVectorIndexReader.rangeSearch and rangeSearchBatch accept VectorRangeSearchParams and optional filter bytes. VectorRangeSearchResult takes ownership of JNI-created Java arrays. Its public constructor and array-returning accessors still make defensive copies. hitCount(), labelAt(int), distanceAt(int), queryStart(int), queryEnd(int), and per-query counter overloads provide allocation-free consumption. Native lengths are checked against JVM array limits; labels, stored offsets and counters use long, while Java hit indices and half-open query bounds use int.
PythonVectorIndexReader.range_search and range_search_batch accept RangeSearchParams and optional roaring_filter=. RangeSearchResult owns copied NumPy arrays (uintp offsets, int64 labels, float32 distances). Native results are freed in finally, even if conversion fails.

Results remain valid after the reader closes. In C, destroy each successful result exactly once, never free its individual arrays, and never access a view after destruction. Input queries and filter bytes are borrowed only for the call. Search sets a valid output pointer to NULL before doing work; an error returns -1, sets paimon_vindex_last_error(), and transfers no result. destroy(NULL) is safe. Callers must supply valid, aligned buffers and live handles; C callers must synchronize operations on a reader. C++, Java and Python retain their existing callback-aware handle locks.

@@ -137,8 +137,17 @@

Language bindings and ownership

VectorRangeSearchParams params = new VectorRangeSearchParams(band, 8); VectorRangeSearchResult result = reader.rangeSearchBatch( queries, queryCount, params, roaringFilter); -long[] labels = result.labelsForQuery(0); +for (int queryIndex = 0; queryIndex < result.queryCount(); queryIndex++) { + for (int hitIndex = result.queryStart(queryIndex); + hitIndex < result.queryEnd(queryIndex); hitIndex++) { + long label = result.labelAt(hitIndex); + float distance = result.distanceAt(hitIndex); + consume(label, distance); + } + long rowsScanned = result.rowsScanned(queryIndex); +} long listReads = result.listReads(); +

Range results are uncapped and still require storage proportional to the number of hits and queries. In Java, prefer indexed access for large results: the array-returning and per-query slice methods intentionally allocate copies. C views borrow result-owned storage, C++ vectors are directly accessible, and Python query slices share the result's owned NumPy arrays. Zero-copy consumption avoids another payload allocation; it is not streaming search or a bound on native search memory.

Python · shared filtered batch
from paimon_vindex import DistanceBand, DistanceEndpoint, DistanceEndpointOp, RangeSearchParams
 
 band = DistanceBand.from_endpoints(
diff --git a/java/src/main/java/org/apache/paimon/index/vector/VectorRangeSearchResult.java b/java/src/main/java/org/apache/paimon/index/vector/VectorRangeSearchResult.java
index 4aeb0eab..402431b3 100644
--- a/java/src/main/java/org/apache/paimon/index/vector/VectorRangeSearchResult.java
+++ b/java/src/main/java/org/apache/paimon/index/vector/VectorRangeSearchResult.java
@@ -23,8 +23,15 @@
 /**
  * CSR range-search output in core scan order, with raw distances and no sorting or top-K cap.
  * IVF-Flat distances are exact; SQ, PQ and RQ return their core distance estimates. Each query
- * occupies [lims[query], lims[query + 1]). Public construction and array access make defensive
- * copies; native construction takes exclusive ownership of its arrays.
+ * occupies [lims[query], lims[query + 1]). Public construction and array-returning accessors make
+ * defensive copies; native construction takes exclusive ownership of its arrays.
+ *
+ * 

For allocation-free consumption, use {@link #hitCount()}, {@link #labelAt(int)}, + * {@link #distanceAt(int)}, and the half-open bounds {@link #queryStart(int)} and + * {@link #queryEnd(int)}. Counter accessors taking a query index also avoid copying arrays. + * Hit indices and query bounds fit in {@code int}, while labels and counters retain their full + * {@code long} range. Indexed accessors reject out-of-range indices with + * {@link IndexOutOfBoundsException}. */ public final class VectorRangeSearchResult { @@ -121,14 +128,36 @@ public int queryCount() { return lims.length - 1; } + public int hitCount() { + return labels.length; + } + + public int queryStart(int queryIndex) { + checkQueryIndex(queryIndex); + return Math.toIntExact(lims[queryIndex]); + } + + public int queryEnd(int queryIndex) { + checkQueryIndex(queryIndex); + return Math.toIntExact(lims[queryIndex + 1]); + } + public long[] labels() { return labels.clone(); } + public long labelAt(int hitIndex) { + return labels[hitIndex]; + } + public float[] distances() { return distances.clone(); } + public float distanceAt(int hitIndex) { + return distances[hitIndex]; + } + public long[] lims() { return lims.clone(); } @@ -138,20 +167,40 @@ public long[] listsProbed() { return listsProbed.clone(); } + public long listsProbed(int queryIndex) { + checkQueryIndex(queryIndex); + return listsProbed[queryIndex]; + } + /** Allow-listed rows evaluated per query, including early-abandoned rows. */ public long[] rowsScanned() { return rowsScanned.clone(); } + public long rowsScanned(int queryIndex) { + checkQueryIndex(queryIndex); + return rowsScanned[queryIndex]; + } + public long[] rowsCommitted() { return rowsCommitted.clone(); } + public long rowsCommitted(int queryIndex) { + checkQueryIndex(queryIndex); + return rowsCommitted[queryIndex]; + } + /** Core diagnostic count, not an arithmetic work measure. */ public long[] earlyAbandoned() { return earlyAbandoned.clone(); } + public long earlyAbandoned(int queryIndex) { + checkQueryIndex(queryIndex); + return earlyAbandoned[queryIndex]; + } + /** Non-empty unique list reads for the whole call; SQ cache hits do not count. */ public long listReads() { return listReads; diff --git a/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeOracleTest.java b/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeOracleTest.java index e4b9aca3..2493bf37 100644 --- a/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeOracleTest.java +++ b/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeOracleTest.java @@ -23,7 +23,6 @@ import java.nio.file.Path; import java.nio.file.Paths; import java.util.ArrayList; -import java.util.Arrays; import java.util.Collections; import java.util.HashMap; import java.util.List; @@ -123,26 +122,23 @@ private static void runCase(Path expected, Path index) throws IOException { : reader.rangeSearchBatch(queries, queryCount, params, filter); } require(actual.queryCount() == queryCount, "query count"); - require(Arrays.equals(lims, actual.lims()), "lims"); - float[] actualDistances = actual.distances(); - require(actualDistances.length == hitCount, "distance count"); - int[] actualBits = new int[hitCount]; - for (int offset = 0; offset < hitCount; offset++) { - actualBits[offset] = Float.floatToRawIntBits(actualDistances[offset]); - } - long[] actualLabels = actual.labels(); + require(actual.hitCount() == hitCount, "hit count"); for (int queryIndex = 0; queryIndex < queryCount; queryIndex++) { int start = Math.toIntExact(lims[queryIndex]); int end = Math.toIntExact(lims[queryIndex + 1]); + require(actual.queryStart(queryIndex) == start, "query start"); + require(actual.queryEnd(queryIndex) == end, "query end"); require( rows(labels, distances, start, end) - .equals(rows(actualLabels, actualBits, start, end)), + .equals(rows(actual, queryIndex)), "label/distance-bit multiset for query " + queryIndex); + require(stats[0][queryIndex] == actual.listsProbed(queryIndex), "listsProbed"); + require(stats[1][queryIndex] == actual.rowsScanned(queryIndex), "rowsScanned"); + require(stats[2][queryIndex] == actual.rowsCommitted(queryIndex), "rowsCommitted"); + require( + stats[3][queryIndex] == actual.earlyAbandoned(queryIndex), + "earlyAbandoned"); } - require(Arrays.equals(stats[0], actual.listsProbed()), "listsProbed"); - require(Arrays.equals(stats[1], actual.rowsScanned()), "rowsScanned"); - require(Arrays.equals(stats[2], actual.rowsCommitted()), "rowsCommitted"); - require(Arrays.equals(stats[3], actual.earlyAbandoned()), "earlyAbandoned"); require(listReads == actual.listReads(), "listReads"); } } @@ -168,6 +164,20 @@ private static Map> rows( return result; } + private static Map> rows(VectorRangeSearchResult result, int queryIndex) { + Map> rows = new HashMap>(); + for (int hitIndex = result.queryStart(queryIndex); + hitIndex < result.queryEnd(queryIndex); + hitIndex++) { + rows.computeIfAbsent(result.labelAt(hitIndex), label -> new ArrayList()) + .add(Float.floatToRawIntBits(result.distanceAt(hitIndex))); + } + for (List values : rows.values()) { + Collections.sort(values); + } + return rows; + } + private static int readBits(Scanner values) { long value = values.nextLong(); require(value >= 0 && value <= 0xffff_ffffL, "f32 bits"); diff --git a/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeSearchTest.java b/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeSearchTest.java index 544d8846..4d600781 100644 --- a/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeSearchTest.java +++ b/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeSearchTest.java @@ -35,6 +35,11 @@ public class VectorIndexRangeSearchTest { private static final int VECTOR_COUNT = 256; public static void main(String[] args) { + if (args.length == 1 && "--zero-copy-memory".equals(args[0])) { + testMemoryBoundedConsumption(); + System.out.println("Range result consumption passed with a bounded heap"); + return; + } VectorIndexNativeLoaderSmokeTest.configureExternalLibrary(args); testValueTypes(); testNative(); @@ -44,6 +49,7 @@ public static void main(String[] args) { static void testValueTypes() { testNativeResultOwnership(); + testIndexedResultAccess(); VectorDistanceBand band = new VectorDistanceBand("l2", null, 4.0f); check(band.lower() == null && band.upper() == 4.0f, "structural bounds"); check("l2".equals(band.metric()), "metric"); @@ -142,6 +148,125 @@ static void testValueTypes() { expect(NullPointerException.class, () -> closed.rangeSearch(new float[1], params, null)); } + private static void testIndexedResultAccess() { + long[] labels = {Long.MIN_VALUE, LABEL_BASE, Long.MAX_VALUE}; + float[] distances = {-0.0f, 1.5f, Float.POSITIVE_INFINITY}; + long[] lims = {0, 0, 2, 2, 3}; + long[] listsProbed = {0, 4, 0, 2}; + long[] rowsScanned = {0, Long.MAX_VALUE, 0, 3}; + long[] rowsCommitted = {0, 2, 0, 1}; + long[] earlyAbandoned = {0, 1, 0, 2}; + VectorRangeSearchResult[] results = { + new VectorRangeSearchResult( + labels, distances, lims, listsProbed, rowsScanned, + rowsCommitted, earlyAbandoned, 3), + VectorRangeSearchResult.fromNative( + labels, distances, lims, listsProbed, rowsScanned, + rowsCommitted, earlyAbandoned, 3) + }; + for (VectorRangeSearchResult result : results) { + check(result.hitCount() == labels.length, "indexed hit count"); + for (int queryIndex = 0; queryIndex < result.queryCount(); queryIndex++) { + check(result.queryStart(queryIndex) == lims[queryIndex], "query start"); + check(result.queryEnd(queryIndex) == lims[queryIndex + 1], "query end"); + check(result.listsProbed(queryIndex) == listsProbed[queryIndex], "indexed listsProbed"); + check(result.rowsScanned(queryIndex) == rowsScanned[queryIndex], "indexed rowsScanned"); + check( + result.rowsCommitted(queryIndex) == rowsCommitted[queryIndex], + "indexed rowsCommitted"); + check( + result.earlyAbandoned(queryIndex) == earlyAbandoned[queryIndex], + "indexed earlyAbandoned"); + for (int hitIndex = result.queryStart(queryIndex); + hitIndex < result.queryEnd(queryIndex); + hitIndex++) { + check(result.labelAt(hitIndex) == labels[hitIndex], "indexed label"); + check( + Float.floatToRawIntBits(result.distanceAt(hitIndex)) + == Float.floatToRawIntBits(distances[hitIndex]), + "indexed distance bits"); + } + } + result.labels()[0] = 0; + result.distances()[0] = 1; + result.lims()[1] = 1; + result.listsProbed()[1] = 0; + result.rowsScanned()[1] = 0; + result.rowsCommitted()[1] = 0; + result.earlyAbandoned()[1] = 0; + check(result.labelAt(0) == Long.MIN_VALUE, "indexed labels remain immutable"); + check(Float.floatToRawIntBits(result.distanceAt(0)) == 0x80000000, "signed zero retained"); + check(result.queryEnd(0) == 0, "indexed offsets remain immutable"); + check( + result.listsProbed(1) == 4 && result.rowsScanned(1) == Long.MAX_VALUE, + "indexed stats remain immutable"); + check( + result.rowsCommitted(1) == 2 && result.earlyAbandoned(1) == 1, + "indexed counters remain immutable"); + for (int hitIndex : new int[] {-1, result.hitCount(), Integer.MAX_VALUE}) { + expect(IndexOutOfBoundsException.class, () -> result.labelAt(hitIndex)); + expect(IndexOutOfBoundsException.class, () -> result.distanceAt(hitIndex)); + } + for (int queryIndex : new int[] {-1, result.queryCount(), Integer.MAX_VALUE}) { + checkInvalidQueryIndex(result, queryIndex); + } + } + for (long[] emptyLims : new long[][] {{0}, {0, 0, 0}}) { + long[] counters = new long[emptyLims.length - 1]; + VectorRangeSearchResult empty = + VectorRangeSearchResult.fromNative( + new long[0], new float[0], emptyLims, + counters, counters, counters, counters, 0); + check(empty.hitCount() == 0, "empty indexed result"); + for (int queryIndex = 0; queryIndex < empty.queryCount(); queryIndex++) { + check( + empty.queryStart(queryIndex) == 0 && empty.queryEnd(queryIndex) == 0, + "empty query bounds"); + } + expect(IndexOutOfBoundsException.class, () -> empty.labelAt(0)); + expect(IndexOutOfBoundsException.class, () -> empty.distanceAt(0)); + checkInvalidQueryIndex(empty, empty.queryCount()); + } + } + + private static void checkInvalidQueryIndex(VectorRangeSearchResult result, int queryIndex) { + expect(IndexOutOfBoundsException.class, () -> result.queryStart(queryIndex)); + expect(IndexOutOfBoundsException.class, () -> result.queryEnd(queryIndex)); + expect(IndexOutOfBoundsException.class, () -> result.listsProbed(queryIndex)); + expect(IndexOutOfBoundsException.class, () -> result.rowsScanned(queryIndex)); + expect(IndexOutOfBoundsException.class, () -> result.rowsCommitted(queryIndex)); + expect(IndexOutOfBoundsException.class, () -> result.earlyAbandoned(queryIndex)); + } + + private static void testMemoryBoundedConsumption() { + check(Runtime.getRuntime().maxMemory() <= 40L * 1024 * 1024, "run with -Xmx36m"); + int hitCount = 2_000_000; + long[] labels = new long[hitCount]; + float[] distances = new float[hitCount]; + labels[0] = LABEL_BASE; + labels[hitCount - 1] = LABEL_BASE + 1; + distances[0] = 1.25f; + distances[hitCount - 1] = -3.5f; + VectorRangeSearchResult result = + VectorRangeSearchResult.fromNative( + labels, distances, new long[] {0, 0, hitCount}, + new long[] {0, 1}, new long[] {0, hitCount}, + new long[] {0, hitCount}, new long[2], 1); + check(result.hitCount() == hitCount && result.queryCount() == 2, "large result shape"); + check(result.queryStart(0) == result.queryEnd(0), "large result empty query"); + long labelSum = 0; + double distanceSum = 0; + for (int hitIndex = result.queryStart(1); hitIndex < result.queryEnd(1); hitIndex++) { + labelSum += result.labelAt(hitIndex); + distanceSum += result.distanceAt(hitIndex); + } + check(labelSum == 2 * LABEL_BASE + 1 && distanceSum == -2.25, "large result consumption"); + check(result.listsProbed(1) == 1 && result.rowsScanned(1) == hitCount, "large result stats"); + check( + result.rowsCommitted(1) == hitCount && result.earlyAbandoned(1) == 0, + "large result counters"); + } + private static void testNativeResultOwnership() { long[] labels = {LABEL_BASE, 7}; float[] distances = {1, 3}; @@ -237,11 +362,73 @@ static void testNative() { } testExactDistancesAndEndpoints(); testNativeValidation(); + testBatchMetadataTransfer(); testCallbacks(); testUnsupported(); VectorIndexRangeOracleTest.runIfConfigured(); } + private static void testBatchMetadataTransfer() { + float[] data = new float[16]; + Arrays.fill(data, 8, 16, 1.0f); + VectorRangeSearchParams params = params("l2", null, 0.5f); + VectorRangeSearchResult[] singles = new VectorRangeSearchResult[3]; + int[] queryCounts = {1023, 1024, 1025, 2048, 2051}; + VectorRangeSearchResult[] batches = new VectorRangeSearchResult[queryCounts.length]; + try (VectorIndexReader reader = open(build("ivf_flat", "l2", 8, data))) { + for (int queryKind = 0; queryKind < singles.length; queryKind++) { + float[] query = new float[8]; + Arrays.fill(query, queryKind == 2 ? 3.0f : queryKind); + singles[queryKind] = reader.rangeSearch(query, params); + check( + singles[queryKind].hitCount() == (queryKind == 2 ? 0 : 1), + "metadata reference hits"); + } + for (int batchIndex = 0; batchIndex < queryCounts.length; batchIndex++) { + int queryCount = queryCounts[batchIndex]; + float[] queries = new float[queryCount * 8]; + for (int queryIndex = 0; queryIndex < queryCount; queryIndex++) { + int queryKind = queryIndex % singles.length; + Arrays.fill( + queries, queryIndex * 8, (queryIndex + 1) * 8, + queryKind == 2 ? 3.0f : queryKind); + } + batches[batchIndex] = reader.rangeSearchBatch(queries, queryCount, params); + } + } + for (int batchIndex = 0; batchIndex < queryCounts.length; batchIndex++) { + VectorRangeSearchResult batch = batches[batchIndex]; + check(batch.queryCount() == queryCounts[batchIndex], "metadata query count"); + int expectedStart = 0; + for (int queryIndex = 0; queryIndex < batch.queryCount(); queryIndex++) { + VectorRangeSearchResult single = singles[queryIndex % singles.length]; + check(batch.queryStart(queryIndex) == expectedStart, "metadata query start"); + check( + batch.queryEnd(queryIndex) == expectedStart + single.hitCount(), + "metadata query end"); + check(batch.listsProbed(queryIndex) == single.listsProbed(0), "metadata listsProbed"); + check(batch.rowsScanned(queryIndex) == single.rowsScanned(0), "metadata rowsScanned"); + check( + batch.rowsCommitted(queryIndex) == single.rowsCommitted(0), + "metadata rowsCommitted"); + check( + batch.earlyAbandoned(queryIndex) == single.earlyAbandoned(0), + "metadata earlyAbandoned"); + for (int hitIndex = 0; hitIndex < single.hitCount(); hitIndex++) { + check( + batch.labelAt(expectedStart + hitIndex) == single.labelAt(hitIndex), + "retained batch label"); + check( + Float.floatToRawIntBits(batch.distanceAt(expectedStart + hitIndex)) + == Float.floatToRawIntBits(single.distanceAt(hitIndex)), + "retained batch distance bits"); + } + expectedStart += single.hitCount(); + } + check(batch.hitCount() == expectedStart, "metadata total hits"); + } + } + private static void testIndex(String indexType, String metric) { float[] data = new float[VECTOR_COUNT * 8]; for (int row = 0; row < VECTOR_COUNT; row++) { diff --git a/jni/src/range.rs b/jni/src/range.rs index 7e3d1b2e..7c834305 100644 --- a/jni/src/range.rs +++ b/jni/src/range.rs @@ -173,28 +173,38 @@ fn long_array<'local>( Ok(array) } +fn counter_array<'local>( + env: &mut JNIEnv<'local>, + len: usize, + name: &str, + value_at: impl Fn(usize) -> usize, +) -> Result>, String> { + let array = env + .new_long_array(checked_array_length(len, name)?) + .map_err(|error| error.to_string())?; + let array = env.auto_local(array); + let mut scratch = [0; 1024]; + for offset in (0..len).step_by(scratch.len()) { + let count = (len - offset).min(scratch.len()); + for (index, value) in scratch[..count].iter_mut().enumerate() { + *value = checked_counter(value_at(offset + index), name)?; + } + env.set_long_array_region( + &array, + checked_array_length(offset, name)?, + &scratch[..count], + ) + .map_err(|error| error.to_string())?; + } + Ok(array) +} + fn build_range_result(env: &mut JNIEnv, result: RangeSearchResult) -> Result { checked_array_length(result.query_count(), "query count")?; checked_array_length(result.lims().len(), "lims")?; checked_array_length(result.labels().len(), "labels")?; let distance_count = checked_array_length(result.distances().len(), "distances")?; let list_reads = checked_counter(result.call_stats().list_reads(), "listReads")?; - let lims = result - .lims() - .iter() - .map(|&value| checked_counter(value, "lims")) - .collect::, _>>()?; - let mut lists_probed = Vec::with_capacity(result.query_count()); - let mut rows_scanned = Vec::with_capacity(result.query_count()); - let mut rows_committed = Vec::with_capacity(result.query_count()); - let mut early_abandoned = Vec::with_capacity(result.query_count()); - for query_index in 0..result.query_count() { - let stats = result.query(query_index).stats; - lists_probed.push(checked_counter(stats.lists_probed(), "listsProbed")?); - rows_scanned.push(checked_counter(stats.rows_scanned(), "rowsScanned")?); - rows_committed.push(checked_counter(stats.rows_committed(), "rowsCommitted")?); - early_abandoned.push(checked_counter(stats.early_abandoned(), "earlyAbandoned")?); - } let labels = long_array(env, result.labels(), "labels")?; let distances = env .new_float_array(distance_count) @@ -202,11 +212,21 @@ fn build_range_result(env: &mut JNIEnv, result: RangeSearchResult) -> Result Date: Tue, 22 Sep 2026 11:52:21 +0800 Subject: [PATCH 5/5] fix: make range distance semantics explicit Use raw-distance names across Rust, C, C++, Java, and Python, including indexed and per-query access. Make raw band construction explicit and keep public endpoint conversion as the normal entry point. Document the range API migration and add contract, endpoint, ABI-layout, and zero-copy regressions. Preserve result values, ownership, C ABI layout, and existing top-K APIs. --- c/range_test_support.h | 38 +-- c/test_vindex.c | 134 +++++++--- core/README.md | 21 +- core/src/collect.rs | 12 +- core/src/ivfpq.rs | 18 +- core/src/ivfrq_io.rs | 10 +- core/src/ivfsq_io.rs | 13 +- core/src/range.rs | 167 ++++++++----- core/tests/range_metrics.rs | 117 ++++++++- core/tests/range_search.rs | 130 +++++----- cpp/test_vindex.cpp | 154 ++++++++++-- docs/api.html | 4 +- docs/range-search.html | 49 ++-- ffi/examples/range_search_fixture.rs | 5 +- ffi/src/range.rs | 50 ++-- ffi/src/range_tests.rs | 92 +++++-- include/paimon_vindex.hpp | 42 ++-- .../index/vector/VectorDistanceBand.java | 37 +-- .../index/vector/VectorRangeSearchResult.java | 35 +-- .../vector/VectorIndexRangeOracleTest.java | 4 +- .../vector/VectorIndexRangeSearchTest.java | 233 +++++++++++++++--- jni/src/range.rs | 27 +- python/README.md | 49 ++-- python/paimon_vindex/__init__.py | 105 +++++--- python/paimon_vindex/_ffi.py | 16 +- python/tests/test_range_search.py | 231 +++++++++++++---- 26 files changed, 1299 insertions(+), 494 deletions(-) diff --git a/c/range_test_support.h b/c/range_test_support.h index 86522198..e24fc182 100644 --- a/c/range_test_support.h +++ b/c/range_test_support.h @@ -73,8 +73,8 @@ static void range_fill_data(float *data, int64_t *labels, float *queries) { static PaimonVindexRangeSearchParams range_all_params(uint32_t metric) { PaimonVindexRangeSearchParams params = {{0, 0, 0.0f, 0, 0.0f}, 0}; params.band.metric = metric; - params.band.lower_kind = PAIMON_VINDEX_BOUND_UNBOUNDED; - params.band.upper_kind = PAIMON_VINDEX_BOUND_UNBOUNDED; + params.band.raw_lower_kind = PAIMON_VINDEX_BOUND_UNBOUNDED; + params.band.raw_upper_kind = PAIMON_VINDEX_BOUND_UNBOUNDED; params.nprobe = RANGE_NLIST; return params; } @@ -83,7 +83,7 @@ static void range_assert_shape(const PaimonVindexRangeSearchResultView *view) { ASSERT_TRUE(view->lims != NULL); ASSERT_TRUE(view->lims[0] == 0); ASSERT_TRUE(view->lims[view->query_count] == view->hit_count); - ASSERT_TRUE(view->hit_count == 0 || (view->labels != NULL && view->distances != NULL)); + ASSERT_TRUE(view->hit_count == 0 || (view->labels != NULL && view->raw_distances != NULL)); ASSERT_TRUE(view->query_count == 0 || view->stats != NULL); for (size_t query = 0; query < view->query_count; ++query) { ASSERT_TRUE(view->lims[query] <= view->lims[query + 1]); @@ -94,7 +94,7 @@ static void range_assert_shape(const PaimonVindexRangeSearchResultView *view) { } for (size_t hit = 0; hit < view->hit_count; ++hit) { ASSERT_TRUE(view->labels[hit] >= (INT64_C(1) << 40)); - ASSERT_TRUE(isfinite(view->distances[hit])); + ASSERT_TRUE(isfinite(view->raw_distances[hit])); } } @@ -118,7 +118,7 @@ struct RangeFixture { uint8_t *filter; size_t *lims; int64_t *labels; - uint32_t *distance_bits; + uint32_t *raw_distance_bits; PaimonVindexRangeSearchStats *stats; uint8_t *index_data; size_t index_len; @@ -148,17 +148,17 @@ static void range_fixture_load( struct RangeFixture *fixture) { memset(fixture, 0, sizeof(*fixture)); FILE *file = range_fixture_open(directory, name, ".expected"); - uint32_t lower_bits; - uint32_t upper_bits; + uint32_t raw_lower_bits; + uint32_t raw_upper_bits; ASSERT_TRUE(fscanf(file, "%zu %" SCNu32 " %zu %zu %" SCNu32 " %" SCNu32 " %" SCNu32 " %" SCNu32 " %zu %zu %zu", &fixture->dimension, &fixture->params.band.metric, &fixture->query_count, &fixture->params.nprobe, - &fixture->params.band.lower_kind, &lower_bits, - &fixture->params.band.upper_kind, &upper_bits, + &fixture->params.band.raw_lower_kind, &raw_lower_bits, + &fixture->params.band.raw_upper_kind, &raw_upper_bits, &fixture->filter_len, &fixture->hit_count, &fixture->list_reads) == 11); - fixture->params.band.lower = range_float_from_bits(lower_bits); - fixture->params.band.upper = range_float_from_bits(upper_bits); + fixture->params.band.raw_lower = range_float_from_bits(raw_lower_bits); + fixture->params.band.raw_upper = range_float_from_bits(raw_upper_bits); ASSERT_TRUE(fixture->dimension != 0); ASSERT_TRUE(fixture->query_count < SIZE_MAX); ASSERT_TRUE(fixture->query_count <= SIZE_MAX / fixture->dimension); @@ -167,7 +167,7 @@ static void range_fixture_load( fixture->filter = (uint8_t *)range_fixture_allocate(fixture->filter_len, sizeof(uint8_t)); fixture->lims = (size_t *)range_fixture_allocate(fixture->query_count + 1, sizeof(size_t)); fixture->labels = (int64_t *)range_fixture_allocate(fixture->hit_count, sizeof(int64_t)); - fixture->distance_bits = (uint32_t *)range_fixture_allocate(fixture->hit_count, sizeof(uint32_t)); + fixture->raw_distance_bits = (uint32_t *)range_fixture_allocate(fixture->hit_count, sizeof(uint32_t)); fixture->stats = (PaimonVindexRangeSearchStats *)range_fixture_allocate( fixture->query_count, sizeof(PaimonVindexRangeSearchStats)); for (size_t element = 0; element < query_len; ++element) { @@ -187,7 +187,7 @@ static void range_fixture_load( ASSERT_TRUE(fscanf(file, "%" SCNd64, &fixture->labels[element]) == 1); } for (size_t element = 0; element < fixture->hit_count; ++element) { - ASSERT_TRUE(fscanf(file, "%" SCNu32, &fixture->distance_bits[element]) == 1); + ASSERT_TRUE(fscanf(file, "%" SCNu32, &fixture->raw_distance_bits[element]) == 1); } for (size_t query = 0; query < fixture->query_count; ++query) { PaimonVindexRangeSearchStats *stats = &fixture->stats[query]; @@ -211,15 +211,15 @@ static void range_fixture_load( struct RangeFixtureHit { int64_t label; - uint32_t distance_bits; + uint32_t raw_distance_bits; }; static int range_fixture_compare_hits(const void *left, const void *right) { const struct RangeFixtureHit *left_hit = (const struct RangeFixtureHit *)left; const struct RangeFixtureHit *right_hit = (const struct RangeFixtureHit *)right; if (left_hit->label != right_hit->label) return left_hit->label < right_hit->label ? -1 : 1; - if (left_hit->distance_bits != right_hit->distance_bits) { - return left_hit->distance_bits < right_hit->distance_bits ? -1 : 1; + if (left_hit->raw_distance_bits != right_hit->raw_distance_bits) { + return left_hit->raw_distance_bits < right_hit->raw_distance_bits ? -1 : 1; } return 0; } @@ -238,9 +238,9 @@ static void range_fixture_assert( fixture->hit_count, sizeof(struct RangeFixtureHit)); for (size_t hit = 0; hit < fixture->hit_count; ++hit) { actual[hit].label = view->labels[hit]; - actual[hit].distance_bits = range_float_bits(view->distances[hit]); + actual[hit].raw_distance_bits = range_float_bits(view->raw_distances[hit]); expected[hit].label = fixture->labels[hit]; - expected[hit].distance_bits = fixture->distance_bits[hit]; + expected[hit].raw_distance_bits = fixture->raw_distance_bits[hit]; } for (size_t query = 0; query < fixture->query_count; ++query) { size_t begin = fixture->lims[query]; @@ -273,7 +273,7 @@ static void range_fixture_run_all(void (*consume)(const struct RangeFixture *)) free(fixture.filter); free(fixture.lims); free(fixture.labels); - free(fixture.distance_bits); + free(fixture.raw_distance_bits); free(fixture.stats); free(fixture.index_data); ++case_count; diff --git a/c/test_vindex.c b/c/test_vindex.c index f2d6a5d2..ac0e9c80 100644 --- a/c/test_vindex.c +++ b/c/test_vindex.c @@ -44,6 +44,13 @@ #include "range_test_support.h" +_Static_assert(_Generic(((PaimonVindexRangeSearchParams *)0)->band, + PaimonVindexRawDistanceBand: 1, default: 0), + "range parameters must use an explicitly raw band"); +_Static_assert(_Generic(((PaimonVindexRangeSearchResultView *)0)->raw_distances, + const float *: 1, default: 0), + "range results must expose borrowed raw distances"); + struct MemBuffer { uint8_t *data; size_t len; @@ -602,42 +609,105 @@ static void consume_range_fixture(const struct RangeFixture *fixture) { paimon_vindex_range_search_result_destroy(result); } +static void test_range_endpoint_raw_results(void) { + const char *metrics[] = {"l2", "inner_product"}; + const uint32_t metric_codes[] = {PAIMON_VINDEX_METRIC_L2, PAIMON_VINDEX_METRIC_INNER_PRODUCT}; + const float coordinates[][3] = {{3.0f, 4.0f, 5.0f}, {3.0f, 2.5f, 2.0f}}; + const float expected_raw_distances[][2] = {{9.0f, 16.0f}, {-6.0f, -5.0f}}; + const int64_t labels[] = {INT64_C(1) << 40, (INT64_C(1) << 40) + 1, (INT64_C(1) << 40) + 2}; + for (size_t metric_index = 0; metric_index < 2; ++metric_index) { + float data[3 * RANGE_DIMENSION] = {0}; + float query[RANGE_DIMENSION] = {0}; + for (size_t row = 0; row < 3; ++row) { + data[row * RANGE_DIMENSION] = coordinates[metric_index][row]; + } + if (metric_codes[metric_index] == PAIMON_VINDEX_METRIC_INNER_PRODUCT) query[0] = 2.0f; + const char *keys[] = {"index.type", "dimension", "nlist", "metric"}; + const char *values[] = {"ivf_flat", "8", "1", metrics[metric_index]}; + PaimonVindexTrainerHandle *trainer = paimon_vindex_trainer_open(keys, values, 4); + ASSERT_TRUE(trainer != NULL); + ASSERT_TRUE(paimon_vindex_trainer_add_training_vectors(trainer, data, 3) == 0); + PaimonVindexTrainingHandle *training = paimon_vindex_trainer_finish(trainer); + ASSERT_TRUE(training != NULL); + paimon_vindex_trainer_free(trainer); + PaimonVindexWriterHandle *writer = paimon_vindex_writer_open(training); + ASSERT_TRUE(writer != NULL); + paimon_vindex_training_free(training); + ASSERT_TRUE(paimon_vindex_writer_add_vectors(writer, labels, data, 3) == 0); + struct MemBuffer buffer = {0}; + PaimonVindexOutputFile output = { + .ctx = &buffer, .write_fn = mem_write, .flush_fn = mem_flush, .get_pos_fn = mem_pos}; + ASSERT_TRUE(paimon_vindex_writer_write_index(writer, output) == 0); + paimon_vindex_writer_free(writer); + PaimonVindexInputFile input = {.ctx = &buffer, .read_ranges_fn = mem_read_ranges}; + PaimonVindexReaderHandle *reader = paimon_vindex_reader_open(input); + ASSERT_TRUE(reader != NULL); + PaimonVindexRangeSearchParams params = range_all_params(metric_codes[metric_index]); + params.nprobe = 1; + PaimonVindexDistanceEndpoint radius = {4.0, PAIMON_VINDEX_CUT_LE}; + PaimonVindexDistanceEndpoint similarity = {5.0, PAIMON_VINDEX_CUT_GE}; + ASSERT_TRUE(paimon_vindex_distance_band_from_endpoints( + metric_codes[metric_index], metric_index == 0 ? NULL : &similarity, + metric_index == 0 ? &radius : NULL, ¶ms.band) == 0); + PaimonVindexRangeSearchResult *result = NULL; + ASSERT_TRUE(paimon_vindex_reader_range_search(reader, query, RANGE_DIMENSION, params, &result) == 0); + paimon_vindex_reader_free(reader); + free(buffer.data); + PaimonVindexRangeSearchResultView view = range_view(result); + range_assert_shape(&view); + ASSERT_TRUE(view.query_count == 1 && view.hit_count == 2); + for (size_t expected = 0; expected < 2; ++expected) { + size_t matches = 0; + for (size_t hit = 0; hit < view.hit_count; ++hit) { + if (view.labels[hit] == labels[expected]) { + ASSERT_TRUE(range_float_bits(view.raw_distances[hit]) == + range_float_bits(expected_raw_distances[metric_index][expected])); + ++matches; + } + } + ASSERT_TRUE(matches == 1); + } + paimon_vindex_range_search_result_destroy(result); + printf("PASS range_endpoint_raw_results %s\n", metrics[metric_index]); + } +} + static void test_range_endpoints(void) { const uint32_t metrics[] = { PAIMON_VINDEX_METRIC_L2, PAIMON_VINDEX_METRIC_COSINE, PAIMON_VINDEX_METRIC_INNER_PRODUCT}; const float candidates[] = {-1.0f, -0.5f, -0.0f, 0.0f, 0.25f, 0.5f, 1.0f, 4.0f}; for (size_t metric_index = 0; metric_index < 3; ++metric_index) { uint32_t metric = metrics[metric_index]; - PaimonVindexDistanceBand band; + PaimonVindexRawDistanceBand band; ASSERT_TRUE(paimon_vindex_distance_band_from_endpoints(metric, NULL, NULL, &band) == 0); ASSERT_TRUE(band.metric == metric); - ASSERT_TRUE(band.lower_kind == PAIMON_VINDEX_BOUND_UNBOUNDED); - ASSERT_TRUE(band.upper_kind == PAIMON_VINDEX_BOUND_UNBOUNDED); + ASSERT_TRUE(band.raw_lower_kind == PAIMON_VINDEX_BOUND_UNBOUNDED); + ASSERT_TRUE(band.raw_upper_kind == PAIMON_VINDEX_BOUND_UNBOUNDED); for (uint32_t lower_op = PAIMON_VINDEX_CUT_GE; lower_op <= PAIMON_VINDEX_CUT_GT; ++lower_op) { for (uint32_t upper_op = PAIMON_VINDEX_CUT_LE; upper_op <= PAIMON_VINDEX_CUT_LT; ++upper_op) { PaimonVindexDistanceEndpoint lower = {0.5, lower_op}; PaimonVindexDistanceEndpoint upper = {1.0, upper_op}; ASSERT_TRUE(paimon_vindex_distance_band_from_endpoints(metric, &lower, &upper, &band) == 0); for (size_t candidate = 0; candidate < sizeof(candidates) / sizeof(candidates[0]); ++candidate) { - float distance = candidates[candidate]; - if (metric == PAIMON_VINDEX_METRIC_L2 && distance < 0) continue; - double value = metric == PAIMON_VINDEX_METRIC_L2 ? (double)sqrtf(distance) - : metric == PAIMON_VINDEX_METRIC_INNER_PRODUCT ? -(double)distance : (double)distance; + float raw_distance = candidates[candidate]; + if (metric == PAIMON_VINDEX_METRIC_L2 && raw_distance < 0) continue; + double value = metric == PAIMON_VINDEX_METRIC_L2 ? (double)sqrtf(raw_distance) + : metric == PAIMON_VINDEX_METRIC_INNER_PRODUCT ? -(double)raw_distance : (double)raw_distance; int expected = (lower_op == PAIMON_VINDEX_CUT_GE ? value >= 0.5 : value > 0.5) && (upper_op == PAIMON_VINDEX_CUT_LE ? value <= 1.0 : value < 1.0); ASSERT_TRUE(expected == - ((band.lower_kind == PAIMON_VINDEX_BOUND_UNBOUNDED || distance >= band.lower) && - (band.upper_kind == PAIMON_VINDEX_BOUND_UNBOUNDED || distance < band.upper))); + ((band.raw_lower_kind == PAIMON_VINDEX_BOUND_UNBOUNDED || raw_distance >= band.raw_lower) && + (band.raw_upper_kind == PAIMON_VINDEX_BOUND_UNBOUNDED || raw_distance < band.raw_upper))); } } } PaimonVindexDistanceEndpoint precise = {nextafter(1.0, 2.0), PAIMON_VINDEX_CUT_GE}; ASSERT_TRUE(paimon_vindex_distance_band_from_endpoints(metric, &precise, NULL, &band) == 0); - float boundary = metric == PAIMON_VINDEX_METRIC_INNER_PRODUCT ? -1.0f : 1.0f; - ASSERT_TRUE(!((band.lower_kind == PAIMON_VINDEX_BOUND_UNBOUNDED || boundary >= band.lower) && - (band.upper_kind == PAIMON_VINDEX_BOUND_UNBOUNDED || boundary < band.upper))); + float raw_boundary = metric == PAIMON_VINDEX_METRIC_INNER_PRODUCT ? -1.0f : 1.0f; + ASSERT_TRUE(!((band.raw_lower_kind == PAIMON_VINDEX_BOUND_UNBOUNDED || raw_boundary >= band.raw_lower) && + (band.raw_upper_kind == PAIMON_VINDEX_BOUND_UNBOUNDED || raw_boundary < band.raw_upper))); } - PaimonVindexDistanceBand band; + PaimonVindexRawDistanceBand band; PaimonVindexDistanceEndpoint endpoint = {1.0, PAIMON_VINDEX_CUT_LT}; ASSERT_TRUE(paimon_vindex_distance_band_from_endpoints(PAIMON_VINDEX_METRIC_L2, &endpoint, NULL, &band) != 0); endpoint.op = PAIMON_VINDEX_CUT_GE; @@ -694,19 +764,19 @@ static void test_range_errors(PaimonVindexReaderHandle *reader, const float *que invalid.band.metric = UINT32_MAX; ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search(reader, query, RANGE_DIMENSION, invalid, &result)); invalid = params; - invalid.band.lower_kind = UINT32_MAX; + invalid.band.raw_lower_kind = UINT32_MAX; ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search(reader, query, RANGE_DIMENSION, invalid, &result)); invalid = params; - invalid.band.upper_kind = PAIMON_VINDEX_BOUND_FINITE; - invalid.band.upper = NAN; + invalid.band.raw_upper_kind = PAIMON_VINDEX_BOUND_FINITE; + invalid.band.raw_upper = NAN; ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search(reader, query, RANGE_DIMENSION, invalid, &result)); - invalid.band.upper = INFINITY; + invalid.band.raw_upper = INFINITY; ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search(reader, query, RANGE_DIMENSION, invalid, &result)); - invalid.band.upper = 0; - invalid.band.lower_kind = PAIMON_VINDEX_BOUND_FINITE; - invalid.band.lower = 1; + invalid.band.raw_upper = 0; + invalid.band.raw_lower_kind = PAIMON_VINDEX_BOUND_FINITE; + invalid.band.raw_lower = 1; ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search(reader, query, RANGE_DIMENSION, invalid, &result)); - invalid.band.lower = 0; + invalid.band.raw_lower = 0; invalid.nprobe = 0; ASSERT_RANGE_ERROR(paimon_vindex_reader_range_search(reader, query, RANGE_DIMENSION, invalid, &result)); float bad_query[RANGE_DIMENSION]; @@ -789,17 +859,17 @@ static void test_range_matrix(void) { ASSERT_TRUE(filtered_view.hit_count == 0 && filtered_view.lims[1] == 0); paimon_vindex_range_search_result_destroy(filtered); PaimonVindexRangeSearchParams bounded = params; - float minimum = single_view.distances[0]; + float minimum = single_view.raw_distances[0]; float maximum = minimum; for (size_t hit = 1; hit < single_view.hit_count; ++hit) { - minimum = fminf(minimum, single_view.distances[hit]); - maximum = fmaxf(maximum, single_view.distances[hit]); + minimum = fminf(minimum, single_view.raw_distances[hit]); + maximum = fmaxf(maximum, single_view.raw_distances[hit]); } - bounded.band.lower_kind = PAIMON_VINDEX_BOUND_FINITE; - bounded.band.lower = metric_codes[metric] == PAIMON_VINDEX_METRIC_L2 + bounded.band.raw_lower_kind = PAIMON_VINDEX_BOUND_FINITE; + bounded.band.raw_lower = metric_codes[metric] == PAIMON_VINDEX_METRIC_L2 ? fmaxf(0, minimum) : minimum; - bounded.band.upper_kind = PAIMON_VINDEX_BOUND_FINITE; - bounded.band.upper = minimum + (maximum - minimum) / 2; + bounded.band.raw_upper_kind = PAIMON_VINDEX_BOUND_FINITE; + bounded.band.raw_upper = minimum + (maximum - minimum) / 2; PaimonVindexRangeSearchResult *subset = NULL; ASSERT_TRUE(paimon_vindex_reader_range_search_batch( reader, queries, RANGE_QUERY_COUNT * RANGE_DIMENSION, RANGE_QUERY_COUNT, bounded, &subset) == 0); @@ -809,7 +879,8 @@ static void test_range_matrix(void) { for (size_t query = 0; query < RANGE_QUERY_COUNT; ++query) { size_t expected = 0; for (size_t hit = batch_view.lims[query]; hit < batch_view.lims[query + 1]; ++hit) { - if (batch_view.distances[hit] >= bounded.band.lower && batch_view.distances[hit] < bounded.band.upper) { + if (batch_view.raw_distances[hit] >= bounded.band.raw_lower && + batch_view.raw_distances[hit] < bounded.band.raw_upper) { ++expected; int found = 0; for (size_t candidate = subset_view.lims[query]; candidate < subset_view.lims[query + 1]; ++candidate) { @@ -821,7 +892,7 @@ static void test_range_matrix(void) { ASSERT_TRUE(subset_view.lims[query + 1] - subset_view.lims[query] == expected); } paimon_vindex_range_search_result_destroy(subset); - bounded.band.lower = bounded.band.upper; + bounded.band.raw_lower = bounded.band.raw_upper; ASSERT_TRUE(paimon_vindex_reader_range_search_batch( reader, queries, RANGE_QUERY_COUNT * RANGE_DIMENSION, RANGE_QUERY_COUNT, bounded, &subset) == 0); subset_view = range_view(subset); @@ -834,7 +905,7 @@ static void test_range_matrix(void) { free(buffer.data); PaimonVindexRangeSearchResultView retained_view = range_view(single); ASSERT_TRUE(retained_view.labels == single_view.labels); - ASSERT_TRUE(retained_view.distances == single_view.distances); + ASSERT_TRUE(retained_view.raw_distances == single_view.raw_distances); range_assert_shape(&retained_view); range_assert_shape(&batch_view); paimon_vindex_range_search_result_destroy(single); @@ -851,6 +922,7 @@ int main(void) { test_output_flush_callback_error_propagates(); test_input_read_ranges_callback_error_propagates(); test_range_endpoints(); + test_range_endpoint_raw_results(); test_range_matrix(); range_fixture_run_all(consume_range_fixture); return 0; diff --git a/core/README.md b/core/README.md index 992900b2..5fef38d0 100644 --- a/core/README.md +++ b/core/README.md @@ -36,7 +36,7 @@ The range path does not change top-K search or the v1 storage format. See the [range search guide](../docs/range-search.html) for membership, -validation, filtering, and statistics. C/JNI range bindings are not included. +validation, filtering, statistics, and the C, C++, Java, and Python bindings. The DiskANN and Vamana code is an independent Apache-licensed implementation based on the published algorithms and this project's existing storage @@ -47,11 +47,15 @@ lower-is-better distance semantics as the IVF indexes. ## Distance semantics -`DistanceBand::new` takes internal, lower-is-better f32 scores: squared L2, +`DistanceBand::from_raw` takes internal, lower-is-better f32 scores: squared L2, cosine distance (`1 - cos`, not clamped), or negative inner product. `MetricType::public_distance` converts a score to its public f64 predicate value: L2 takes the f32 square root before widening, cosine widens unchanged, and inner -product negates. Returned `distances` always remain in internal units. +product negates. Returned `raw_distances` always remain in internal units. +Both `RangeSearchResult::raw_distances()` and `QueryResult::raw_distances` borrow +those raw values without converting or copying them. A public L2 radius of 4 +can return raw distance 9; an inner-product lower endpoint of 5 can return -6. +Do not compare raw results directly with public endpoint literals. Use `DistanceBand::from_endpoints` for public f64 predicates. `Ge`/`Gt` belong on the lower side and `Le`/`Lt` on the upper side. L2 resolves square-root rounding; @@ -61,10 +65,17 @@ Endpoints are never rounded to f32 first. Missing sides are structurally unbounded, so an inclusive endpoint at `f32::MAX` does not lose that value. Out-of-domain cosine/IP predicates become empty or unbounded bands; unrepresentable L2 cuts retain the existing `Unsupported` response. +The `raw_lower()` and `raw_upper()` accessors expose converted half-open cuts, +not the original endpoints. `admit_raw()` accepts internal values only. + +Range API callers must migrate `DistanceBand::new` to `from_raw` (or preferably +`from_endpoints`), `distances` to `raw_distances`, and raw cut/membership access +to the explicit names. This is a source API change, not a change to stored data, +returned values, or top-K APIs. `IndexType::supports_range_search(metric)` and `reader.supports_range_search()` -report capability without reading list payloads. DiskANN and C/JNI range APIs -remain unsupported. Bad queries, mismatched metrics, non-finite endpoints, +report capability without reading list payloads. DiskANN range search remains +unsupported. Bad queries, mismatched metrics, non-finite endpoints, malformed filters and zero `nprobe` are errors even for an empty band. Non-finite consumed distances or cosine norms return `InvalidData`, not partial results. Cosine queries use the existing normalization, including leaving zero diff --git a/core/src/collect.rs b/core/src/collect.rs index bd3c829a..6701b853 100644 --- a/core/src/collect.rs +++ b/core/src/collect.rs @@ -90,7 +90,7 @@ fn early_abandon_threshold(band: DistanceBand) -> f32 { if band.metric() != MetricType::L2 { return f32::INFINITY; } - match band.upper() { + match band.raw_upper() { Bound::Finite(upper) => upper, Bound::Unbounded => f32::INFINITY, } @@ -164,7 +164,7 @@ impl Collector for RangeCollector { format!("non-finite distance {value} computed for row {id}"), )); } - if self.band.admit(value) { + if self.band.admit_raw(value) { self.rows.push((id, value)); } Ok(()) @@ -178,7 +178,7 @@ mod tests { use crate::range::{Bound, DistanceBand}; fn l2_band(lower: f32, upper: f32) -> DistanceBand { - DistanceBand::new(Bound::Finite(lower), Bound::Finite(upper), MetricType::L2).unwrap() + DistanceBand::from_raw(Bound::Finite(lower), Bound::Finite(upper), MetricType::L2).unwrap() } #[test] @@ -221,7 +221,8 @@ mod tests { #[test] fn an_unbounded_upper_reports_no_cutoff() { - let band = DistanceBand::new(Bound::Finite(1.0), Bound::Unbounded, MetricType::L2).unwrap(); + let band = + DistanceBand::from_raw(Bound::Finite(1.0), Bound::Unbounded, MetricType::L2).unwrap(); assert_eq!(RangeCollector::new(band).cutoff(), f32::INFINITY); } @@ -230,7 +231,8 @@ mod tests { // A partial cosine or inner-product accumulation does not bound the full // value, so the cutoff must stay infinite no matter what the band says. for metric in [MetricType::Cosine, MetricType::InnerProduct] { - let band = DistanceBand::new(Bound::Finite(0.1), Bound::Finite(0.5), metric).unwrap(); + let band = + DistanceBand::from_raw(Bound::Finite(0.1), Bound::Finite(0.5), metric).unwrap(); assert_eq!( RangeCollector::new(band).cutoff(), f32::INFINITY, diff --git a/core/src/ivfpq.rs b/core/src/ivfpq.rs index 60236a46..4852bb3c 100644 --- a/core/src/ivfpq.rs +++ b/core/src/ivfpq.rs @@ -3474,7 +3474,7 @@ mod tests { let queries = (0..16) .map(|query_index| vec![query_index as f32 * 0.25; 128]) .collect::>(); - let band = crate::range::DistanceBand::new( + let band = crate::range::DistanceBand::from_raw( crate::range::Bound::Unbounded, crate::range::Bound::Unbounded, metric, @@ -3541,7 +3541,7 @@ mod tests { count: usize, metric: MetricType, ) -> Vec> { - let band = crate::range::DistanceBand::new( + let band = crate::range::DistanceBand::from_raw( crate::range::Bound::Unbounded, crate::range::Bound::Unbounded, metric, @@ -3662,7 +3662,7 @@ mod tests { (Bound::Finite(0.0), Bound::Finite(1.0)), (Bound::Finite(4.0), Bound::Finite(5.0)), ] { - let band = DistanceBand::new(lower, upper, MetricType::InnerProduct).unwrap(); + let band = DistanceBand::from_raw(lower, upper, MetricType::InnerProduct).unwrap(); let mut queries = [PqRangeQuery::new( 0, vec![1.0; 2], @@ -3696,7 +3696,7 @@ mod tests { ) .unwrap(); let distance = (16_777_216.0_f32 + offset) - 16_777_216.0; - if band.admit(distance) { + if band.admit_raw(distance) { expected.push((ids[0], distance.to_bits())); } assert_eq!(queries[0].table_list, Some(0)); @@ -3721,8 +3721,9 @@ mod tests { fn pq_range_residual_ip_preserves_zero_bits_with_and_without_cache() { use crate::range::{Bound, DistanceBand}; - let band = DistanceBand::new(Bound::Unbounded, Bound::Unbounded, MetricType::InnerProduct) - .unwrap(); + let band = + DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, MetricType::InnerProduct) + .unwrap(); for bits in [4, 8] { for subquantizers in [1, 2] { let mut pq = ProductQuantizer::with_nbits(subquantizers, subquantizers, bits); @@ -3769,8 +3770,9 @@ mod tests { fn pq_range_residual_ip_rejects_nonfinite_offsets_with_and_without_cache() { use crate::range::{Bound, DistanceBand}; - let band = DistanceBand::new(Bound::Unbounded, Bound::Unbounded, MetricType::InnerProduct) - .unwrap(); + let band = + DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, MetricType::InnerProduct) + .unwrap(); for bits in [4, 8] { let mut pq = ProductQuantizer::with_nbits(2, 2, bits); pq.set_centroids(vec![1.0; 2 * pq.ksub()]); diff --git a/core/src/ivfrq_io.rs b/core/src/ivfrq_io.rs index 4a0f1984..e88b0125 100644 --- a/core/src/ivfrq_io.rs +++ b/core/src/ivfrq_io.rs @@ -1715,7 +1715,7 @@ mod tests { fn all_range_params(nprobe: usize) -> VectorRangeSearchParams { use crate::range::{Bound, DistanceBand}; VectorRangeSearchParams::new( - DistanceBand::new(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(), + DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(), nprobe, ) } @@ -1725,7 +1725,7 @@ mod tests { let mut pairs = query .labels .iter() - .zip(query.distances) + .zip(query.raw_distances) .map(|(&id, &distance)| (id, distance.to_bits())) .collect::>(); pairs.sort_unstable(); @@ -1754,7 +1754,7 @@ mod tests { let queries = [vec![0.0; index.d], vec![0.25; index.d], vec![-2.0; index.d]].concat(); let params = VectorRangeSearchParams::new( - DistanceBand::new(Bound::Unbounded, Bound::Unbounded, metric).unwrap(), + DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, metric).unwrap(), 1, ); let result = reader.range_search_batch(&queries, 3, params).unwrap(); @@ -1802,7 +1802,7 @@ mod tests { let mut filter_bytes = Vec::new(); filter.serialize_into(&mut filter_bytes).unwrap(); let empty = VectorRangeSearchParams::new( - DistanceBand::new(Bound::Finite(1.0), Bound::Finite(1.0), MetricType::L2).unwrap(), + DistanceBand::from_raw(Bound::Finite(1.0), Bound::Finite(1.0), MetricType::L2).unwrap(), 3, ); let calls = stats.lock().unwrap().calls; @@ -1884,7 +1884,7 @@ mod tests { io::ErrorKind::InvalidInput ); let wrong_metric = VectorRangeSearchParams::new( - DistanceBand::new(Bound::Unbounded, Bound::Unbounded, MetricType::Cosine).unwrap(), + DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, MetricType::Cosine).unwrap(), 3, ); assert_eq!( diff --git a/core/src/ivfsq_io.rs b/core/src/ivfsq_io.rs index 4bd08830..a48a2ae7 100644 --- a/core/src/ivfsq_io.rs +++ b/core/src/ivfsq_io.rs @@ -1731,7 +1731,8 @@ mod tests { [MetricType::L2, MetricType::Cosine, MetricType::InnerProduct] .map(|metric| (selection, metric)) }) { - let band = DistanceBand::new(Bound::Finite(1.0), Bound::Finite(200.0), metric).unwrap(); + let band = + DistanceBand::from_raw(Bound::Finite(1.0), Bound::Finite(200.0), metric).unwrap(); let workers = AtomicU64::new(0); let query_indices = [14, 2, 12, 4, 10, 6, 8, 0]; let mut collectors = query_indices @@ -1841,7 +1842,8 @@ mod tests { let (index, data, ids) = build_index(37, 4, 1_024); let mut reader = IVFSQIndexReader::open(Cursor::new(serialized_index(&index))).unwrap(); let filter = CountingFilter(AtomicUsize::new(0)); - let band = DistanceBand::new(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(); + let band = + DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(); let result = reader .range_search_batch_with_filter( &data[..37 * 3], @@ -1901,8 +1903,9 @@ mod tests { Bound::Finite(0.0), Bound::Finite(full.distances[count / 2]), ] { - let band = DistanceBand::new(Bound::Finite(0.0), upper, MetricType::L2) - .unwrap(); + let band = + DistanceBand::from_raw(Bound::Finite(0.0), upper, MetricType::L2) + .unwrap(); let mut collector = RangeCollector::new(band); scan_sq_rows( &query, @@ -1922,7 +1925,7 @@ mod tests { .zip(full.distances.iter().copied()) .filter(|(id, distance)| { filter.map(|filter| filter.contains(*id)).unwrap_or(true) - && band.admit(*distance) + && band.admit_raw(*distance) }) .map(|(id, distance)| (id, distance.to_bits())) .collect::>(); diff --git a/core/src/range.rs b/core/src/range.rs index 0a75d2d9..0f4505f9 100644 --- a/core/src/range.rs +++ b/core/src/range.rs @@ -49,6 +49,16 @@ pub enum Bound { /// A validated internal interval `[lower, upper)` together with the metric it /// belongs to. +/// +/// Use [`Self::from_endpoints`] for public-distance predicates, or explicitly +/// opt into internal cuts with [`Self::from_raw`]. Ambiguous raw construction +/// is intentionally unavailable: +/// +/// ```compile_fail +/// use paimon_vindex_core::distance::MetricType; +/// use paimon_vindex_core::range::{Bound, DistanceBand}; +/// DistanceBand::new(Bound::Unbounded, Bound::Finite(4.0), MetricType::L2); +/// ``` #[derive(Debug, Clone, Copy, PartialEq)] pub struct DistanceBand { lower: Bound, @@ -61,14 +71,16 @@ fn invalid(message: impl Into) -> io::Error { } impl DistanceBand { - /// Validates and constructs a band, rejecting non-finite cuts, negative cuts - /// under squared L2, and inverted intervals. + /// Constructs a band from raw f32 cuts: squared L2, cosine distance, or + /// negative inner product. Rejects non-finite cuts, negative squared-L2 + /// cuts, and inverted intervals. Prefer [`Self::from_endpoints`] for public + /// endpoint literals; this method does not convert units or operators. /// /// This has to fail loud rather than "quietly return no rows". Any caller can /// pass an illegal value, and an empty result would be read upstream as /// "this bucket genuinely has no matches" -- indistinguishable from a /// correct answer, and so undetectable. - pub fn new(lower: Bound, upper: Bound, metric: MetricType) -> io::Result { + pub fn from_raw(lower: Bound, upper: Bound, metric: MetricType) -> io::Result { for (side, bound) in [("lower", lower), ("upper", upper)] { if let Bound::Finite(value) = bound { if !value.is_finite() { @@ -97,11 +109,13 @@ impl DistanceBand { self.metric } - pub fn lower(&self) -> Bound { + /// Inclusive raw lower cut, not the original public lower endpoint. + pub fn raw_lower(&self) -> Bound { self.lower } - pub fn upper(&self) -> Bound { + /// Exclusive raw upper cut, not the original public upper endpoint. + pub fn raw_upper(&self) -> Bound { self.upper } @@ -110,7 +124,7 @@ impl DistanceBand { matches!((self.lower, self.upper), (Bound::Finite(low), Bound::Finite(high)) if low >= high) } - /// Membership test. Left-closed, right-open. + /// Membership test for a raw distance. Left-closed, right-open. /// /// A non-finite value is never a member, whichever side is unbounded. /// Without this an unbounded side would admit `NaN` and the matching @@ -121,7 +135,7 @@ impl DistanceBand { /// rather than raising -- but neither should ever call such a value a match, /// and this is the public membership authority. #[inline] - pub fn admit(&self, value: f32) -> bool { + pub fn admit_raw(&self, value: f32) -> bool { if !value.is_finite() { return false; } @@ -414,9 +428,9 @@ impl DistanceBand { || upper_cut == Some(-f32::MAX) || matches!((lower_cut, upper_cut), (Some(Some(low)), Some(high)) if low > high) { - return Self::new(Bound::Finite(0.0), Bound::Finite(0.0), metric); + return Self::from_raw(Bound::Finite(0.0), Bound::Finite(0.0), metric); } - return Self::new( + return Self::from_raw( lower_cut.flatten().map_or(Bound::Unbounded, Bound::Finite), upper_cut.map_or(Bound::Unbounded, Bound::Finite), metric, @@ -432,7 +446,7 @@ impl DistanceBand { }; let lower_bound = resolve(lower_finder)?; let upper_bound = resolve(upper_finder)?; - Self::new(lower_bound, upper_bound, metric) + Self::from_raw(lower_bound, upper_bound, metric) } } @@ -498,19 +512,36 @@ impl RangeSearchCallStats { #[derive(Debug)] pub struct QueryResult<'a> { pub labels: &'a [i64], - pub distances: &'a [f32], + /// Raw squared-L2, cosine-distance, or negative-inner-product values. + pub raw_distances: &'a [f32], pub stats: &'a RangeSearchStats, } -/// The CSR-shaped batch result: the repository's existing flat -/// `(Vec, Vec)` convention plus the one thing variable-length results -/// require, `lims`. The `lims` / `labels` / `distances` names are aligned with -/// Faiss's `RangeSearchResult`; its `nq` is derivable here and so is not stored. +/// CSR-shaped batch result with offsets, labels, and explicitly raw distances. +/// Values remain in core units, including quantized estimates; public endpoint +/// conversion does not transform the output. There is no ambiguous distance +/// accessor: +/// +/// ```compile_fail +/// use paimon_vindex_core::range::RangeSearchResult; +/// fn ambiguous(result: &RangeSearchResult) { +/// let _ = result.distances(); +/// } +/// ``` +/// +/// Per-query views follow the same contract: +/// +/// ```compile_fail +/// use paimon_vindex_core::range::QueryResult; +/// fn ambiguous(query: QueryResult<'_>) { +/// let _ = query.distances; +/// } +/// ``` #[derive(Debug)] pub struct RangeSearchResult { lims: Vec, labels: Vec, - distances: Vec, + raw_distances: Vec, stats: Vec, call_stats: RangeSearchCallStats, } @@ -526,7 +557,7 @@ impl RangeSearchResult { let (start, end) = (self.lims[i], self.lims[i + 1]); QueryResult { labels: &self.labels[start..end], - distances: &self.distances[start..end], + raw_distances: &self.raw_distances[start..end], stats: &self.stats[i], } } @@ -543,8 +574,11 @@ impl RangeSearchResult { &self.labels } - pub fn distances(&self) -> &[f32] { - &self.distances + /// Borrows raw squared-L2, cosine-distance, or negative-inner-product values. + /// Use [`MetricType::public_distance`] for a public predicate value; do not + /// compare these values directly with public L2 or inner-product endpoints. + pub fn raw_distances(&self) -> &[f32] { + &self.raw_distances } } @@ -614,19 +648,19 @@ impl RangeResultBuilder { let total: usize = self.rows.iter().map(Vec::len).sum(); let mut lims = Vec::with_capacity(self.rows.len() + 1); let mut labels = Vec::with_capacity(total); - let mut distances = Vec::with_capacity(total); + let mut raw_distances = Vec::with_capacity(total); lims.push(0); for query_rows in &self.rows { for (id, distance) in query_rows { labels.push(*id); - distances.push(*distance); + raw_distances.push(*distance); } lims.push(labels.len()); } RangeSearchResult { lims, labels, - distances, + raw_distances, stats: self.stats, call_stats: self.call_stats, } @@ -809,33 +843,33 @@ mod tests { #[test] fn a_finite_l2_band_is_left_closed_right_open() { let band = - DistanceBand::new(Bound::Finite(1.0), Bound::Finite(3.0), MetricType::L2).unwrap(); - assert!(!band.admit(0.999)); - assert!(band.admit(1.0), "the lower end is closed"); - assert!(band.admit(2.999)); - assert!(!band.admit(3.0), "the upper end is open"); + DistanceBand::from_raw(Bound::Finite(1.0), Bound::Finite(3.0), MetricType::L2).unwrap(); + assert!(!band.admit_raw(0.999)); + assert!(band.admit_raw(1.0), "the lower end is closed"); + assert!(band.admit_raw(2.999)); + assert!(!band.admit_raw(3.0), "the upper end is open"); } #[test] fn unboundedness_is_structural_not_a_sentinel() { // Cosine's 1-cos can go slightly negative because distance.rs does not // clamp it, so a 0.0 sentinel would drop the most similar rows. - let band = - DistanceBand::new(Bound::Unbounded, Bound::Finite(0.5), MetricType::Cosine).unwrap(); + let band = DistanceBand::from_raw(Bound::Unbounded, Bound::Finite(0.5), MetricType::Cosine) + .unwrap(); assert!( - band.admit(-1.1920929e-7), + band.admit_raw(-1.1920929e-7), "an unbounded lower end must admit a slightly negative cosine distance" ); // An inner-product internal distance can be exactly f32::MAX, which a // half-open interval with that sentinel as its upper cut would exclude. - let band = DistanceBand::new( + let band = DistanceBand::from_raw( Bound::Finite(0.0), Bound::Unbounded, MetricType::InnerProduct, ) .unwrap(); assert!( - band.admit(f32::MAX), + band.admit_raw(f32::MAX), "an unbounded upper end must admit f32::MAX" ); } @@ -847,40 +881,48 @@ mod tests { // public membership authority; it answers "not a member" rather than // raising, which is not the same response as the collector's fail-loud // push, but neither may ever call such a value a match. - let whole = DistanceBand::new(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(); + let whole = + DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(); let lower_open = - DistanceBand::new(Bound::Unbounded, Bound::Finite(1.0), MetricType::L2).unwrap(); + DistanceBand::from_raw(Bound::Unbounded, Bound::Finite(1.0), MetricType::L2).unwrap(); let upper_open = - DistanceBand::new(Bound::Finite(0.0), Bound::Unbounded, MetricType::L2).unwrap(); + DistanceBand::from_raw(Bound::Finite(0.0), Bound::Unbounded, MetricType::L2).unwrap(); let closed = - DistanceBand::new(Bound::Finite(0.0), Bound::Finite(1.0), MetricType::L2).unwrap(); + DistanceBand::from_raw(Bound::Finite(0.0), Bound::Finite(1.0), MetricType::L2).unwrap(); for band in [whole, lower_open, upper_open, closed] { for value in [f32::NAN, f32::INFINITY, f32::NEG_INFINITY] { - assert!(!band.admit(value), "{band:?} must not admit {value}"); + assert!(!band.admit_raw(value), "{band:?} must not admit {value}"); } } } #[test] fn a_whole_space_band_admits_everything_finite() { - let band = DistanceBand::new(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(); - assert!(band.admit(0.0) && band.admit(f32::MAX) && !band.is_empty()); + let band = + DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(); + assert!(band.admit_raw(0.0) && band.admit_raw(f32::MAX) && !band.is_empty()); } #[test] fn an_inverted_or_non_finite_band_fails_loud() { - assert!(DistanceBand::new(Bound::Finite(5.0), Bound::Finite(3.0), MetricType::L2).is_err()); assert!( - DistanceBand::new(Bound::Finite(f32::NAN), Bound::Finite(3.0), MetricType::L2).is_err() + DistanceBand::from_raw(Bound::Finite(5.0), Bound::Finite(3.0), MetricType::L2).is_err() ); - assert!(DistanceBand::new( + assert!(DistanceBand::from_raw( + Bound::Finite(f32::NAN), + Bound::Finite(3.0), + MetricType::L2 + ) + .is_err()); + assert!(DistanceBand::from_raw( Bound::Finite(0.0), Bound::Finite(f32::INFINITY), MetricType::L2 ) .is_err()); assert!( - DistanceBand::new(Bound::Finite(-1.0), Bound::Finite(3.0), MetricType::L2).is_err(), + DistanceBand::from_raw(Bound::Finite(-1.0), Bound::Finite(3.0), MetricType::L2) + .is_err(), "L2 is a squared distance, so a negative cut is illegal" ); } @@ -888,15 +930,15 @@ mod tests { #[test] fn an_empty_band_is_legal_and_admits_nothing() { let band = - DistanceBand::new(Bound::Finite(2.0), Bound::Finite(2.0), MetricType::L2).unwrap(); + DistanceBand::from_raw(Bound::Finite(2.0), Bound::Finite(2.0), MetricType::L2).unwrap(); assert!(band.is_empty()); - assert!(!band.admit(2.0)); + assert!(!band.admit_raw(2.0)); } #[test] fn all_metrics_accept_positive_probe_widths() { for metric in [MetricType::L2, MetricType::Cosine, MetricType::InnerProduct] { - let band = DistanceBand::new(Bound::Unbounded, Bound::Unbounded, metric).unwrap(); + let band = DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, metric).unwrap(); assert_eq!( VectorRangeSearchParams::new(band, 3).validate(2).unwrap(), 2 @@ -954,10 +996,10 @@ mod tests { op: CutOperator::Lt, }; let band = DistanceBand::from_endpoints(None, Some(ep), MetricType::L2).unwrap(); - assert_eq!(band.lower(), Bound::Unbounded); + assert_eq!(band.raw_lower(), Bound::Unbounded); let band = DistanceBand::from_endpoints(None, None, MetricType::L2).unwrap(); assert_eq!( - (band.lower(), band.upper()), + (band.raw_lower(), band.raw_upper()), (Bound::Unbounded, Bound::Unbounded) ); } @@ -972,7 +1014,7 @@ mod tests { let band = DistanceBand::from_endpoints(Some(ep), None, metric).unwrap(); for distance in [-1.0, -0.5, 0.0, 0.5, 1.0] { assert_eq!( - band.admit(distance), + band.admit_raw(distance), metric.public_distance(distance) >= 0.5 ); } @@ -1023,7 +1065,7 @@ mod tests { MetricType::L2, ) .unwrap(); - let Bound::Finite(cut) = band.lower() else { + let Bound::Finite(cut) = band.raw_lower() else { panic!("lower must be finite") }; cut @@ -1075,7 +1117,7 @@ mod tests { MetricType::L2, ) .unwrap(); - let Bound::Finite(cut) = band.upper() else { + let Bound::Finite(cut) = band.raw_upper() else { panic!("upper must be finite") }; cut @@ -1104,14 +1146,14 @@ mod tests { op: CutOperator::Lt, }; let band = DistanceBand::from_endpoints(Some(lower), Some(upper), MetricType::L2).unwrap(); - let (Bound::Finite(lo), Bound::Finite(hi)) = (band.lower(), band.upper()) else { + let (Bound::Finite(lo), Bound::Finite(hi)) = (band.raw_lower(), band.raw_upper()) else { panic!("both cuts must be finite") }; // Just before the boundary, on it, and just before the upper cut. - assert!(!band.admit(f32::from_bits(lo.to_bits() - 1))); - assert!(band.admit(lo)); - assert!(band.admit(f32::from_bits(hi.to_bits() - 1))); - assert!(!band.admit(hi)); + assert!(!band.admit_raw(f32::from_bits(lo.to_bits() - 1))); + assert!(band.admit_raw(lo)); + assert!(band.admit_raw(f32::from_bits(hi.to_bits() - 1))); + assert!(!band.admit_raw(hi)); } #[test] @@ -1150,7 +1192,7 @@ mod tests { }; let band = DistanceBand::from_endpoints(Some(ep), None, MetricType::L2).unwrap(); assert_eq!( - band.lower(), + band.raw_lower(), Bound::Finite(0.0), "every distance is >= a negative endpoint" ); @@ -1167,15 +1209,15 @@ mod tests { builder.push_row(0, 11, 2.5); let result = builder.build(); assert_eq!(result.query(0).labels, &[10, 11]); - assert_eq!(result.query(0).distances, &[1.5, 2.5]); + assert_eq!(result.query(0).raw_distances, &[1.5, 2.5]); assert_eq!(result.query(1).labels, &[20]); - assert_eq!(result.query(1).distances, &[3.5]); + assert_eq!(result.query(1).raw_distances, &[3.5]); } #[test] fn a_query_with_no_hits_yields_empty_slices() { let result = RangeResultBuilder::new(1).build(); - assert!(result.query(0).labels.is_empty() && result.query(0).distances.is_empty()); + assert!(result.query(0).labels.is_empty() && result.query(0).raw_distances.is_empty()); assert_eq!(result.call_stats().list_reads(), 0); } @@ -1199,7 +1241,7 @@ mod tests { // --- Task 7: width and params ---------------------------------------- fn l2_band_for_params() -> DistanceBand { - DistanceBand::new(Bound::Finite(0.0), Bound::Finite(1.0), MetricType::L2).unwrap() + DistanceBand::from_raw(Bound::Finite(0.0), Bound::Finite(1.0), MetricType::L2).unwrap() } #[test] @@ -1224,7 +1266,8 @@ mod tests { #[test] fn a_zero_nprobe_is_invalid_for_every_metric() { for metric in [MetricType::Cosine, MetricType::InnerProduct] { - let band = DistanceBand::new(Bound::Finite(0.0), Bound::Finite(1.0), metric).unwrap(); + let band = + DistanceBand::from_raw(Bound::Finite(0.0), Bound::Finite(1.0), metric).unwrap(); let err = VectorRangeSearchParams::new(band, 0) .validate(1024) .unwrap_err(); diff --git a/core/tests/range_metrics.rs b/core/tests/range_metrics.rs index b5abb843..fb4c3851 100644 --- a/core/tests/range_metrics.rs +++ b/core/tests/range_metrics.rs @@ -403,12 +403,12 @@ fn result_pairs(result: &RangeSearchResult, query: usize) -> Vec<(i64, u32)> { .labels .iter() .copied() - .zip(result.query(query).distances.iter().copied()), + .zip(result.query(query).raw_distances.iter().copied()), ) } fn band(metric: MetricType) -> DistanceBand { - DistanceBand::new(Bound::Unbounded, Bound::Unbounded, metric).unwrap() + DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, metric).unwrap() } fn filter_bytes(ids: impl IntoIterator) -> Vec { @@ -447,7 +447,8 @@ fn all_families_match_independent_metric_oracles_and_four_entry_points() { let lower = distances[distances.len() / 4]; let upper = distances[distances.len() * 3 / 4]; let cuts = - DistanceBand::new(Bound::Finite(lower), Bound::Finite(upper), metric).unwrap(); + DistanceBand::from_raw(Bound::Finite(lower), Bound::Finite(upper), metric) + .unwrap(); for band in [band(metric), cuts] { let params = VectorRangeSearchParams::new(band, nprobe); let allowed = oracles @@ -478,7 +479,7 @@ fn all_families_match_independent_metric_oracles_and_four_entry_points() { for (query_index, query) in queries.chunks_exact(DIMENSION).enumerate() { let expected = pairs(oracles[query_index].iter().copied().filter( |(id, distance)| { - band.admit(*distance) + band.admit_raw(*distance) && (selected.is_none() || allowed.contains(id)) }, )); @@ -591,7 +592,7 @@ fn endpoints_match_public_predicates_at_ulps_zeros_and_extremes() { CutOperator::Lt => displayed < value, }; assert_eq!( - band.admit(candidate), + band.admit_raw(candidate), expected, "{metric:?} {op:?} {value} candidate={candidate}" ); @@ -617,7 +618,7 @@ fn endpoints_match_public_predicates_at_ulps_zeros_and_extremes() { .unwrap(); for &candidate in &values { assert_eq!( - singleton.admit(candidate), + singleton.admit_raw(candidate), metric.public_distance(candidate) == value, "{metric:?} singleton={value} candidate={candidate}" ); @@ -647,7 +648,7 @@ fn endpoints_match_public_predicates_at_ulps_zeros_and_extremes() { metric, ) .unwrap(); - assert!(values.iter().all(|&value| !empty.admit(value))); + assert!(values.iter().all(|&value| !empty.admit_raw(value))); } } } @@ -672,7 +673,7 @@ fn capability_validation_empty_bands_and_topk_are_preserved() { topk ); let empty = VectorRangeSearchParams::new( - DistanceBand::new(Bound::Finite(0.5), Bound::Finite(0.5), metric).unwrap(), + DistanceBand::from_raw(Bound::Finite(0.5), Bound::Finite(0.5), metric).unwrap(), LISTS, ); assert_eq!( @@ -837,7 +838,7 @@ fn cosine_zero_queries_still_validate_consumed_row_norms() { let result = direct(&bytes, family, &queries, count > 1, params, Some(&good)).unwrap(); for query in 0..count { assert_eq!(result.query(query).labels, &[1]); - assert_eq!(result.query(query).distances, &[1.0]); + assert_eq!(result.query(query).raw_distances, &[1.0]); assert_eq!(result.query(query).stats.rows_scanned(), 1); } } @@ -891,6 +892,98 @@ fn nonfinite_coarse_data_and_finite_query_overflow_fail_loud() { } } +#[test] +fn public_endpoints_return_explicit_raw_values_across_entry_points() { + for metric in METRICS { + let mut index = IVFFlatIndex::new(DIMENSION, 1, metric); + index.set_quantizer_centroids(vec![0.0; DIMENSION]); + index.ids[0] = vec![11, 12]; + let mut vectors = vec![0.0; 2 * DIMENSION]; + let mut query = [0.0; DIMENSION]; + let (lower, upper, expected_raw) = match metric { + MetricType::L2 => { + vectors[0] = 3.0; + vectors[DIMENSION] = 5.0; + ( + None, + Some(DistanceEndpoint { + value: 4.0, + op: CutOperator::Lt, + }), + 9.0, + ) + } + MetricType::InnerProduct => { + query[0] = 2.0; + vectors[0] = 3.0; + vectors[DIMENSION] = 1.0; + ( + Some(DistanceEndpoint { + value: 5.0, + op: CutOperator::Ge, + }), + None, + -6.0, + ) + } + MetricType::Cosine => { + query[0] = 1.0; + vectors[1] = 1.0; + vectors[DIMENSION] = -1.0; + ( + None, + Some(DistanceEndpoint { + value: 1.5, + op: CutOperator::Lt, + }), + 1.0, + ) + } + }; + index.vectors[0] = vectors; + let bytes = Source::Flat(index).bytes(); + let band = DistanceBand::from_endpoints(lower, upper, metric).unwrap(); + assert!(band.admit_raw(expected_raw)); + let explicit_raw = + DistanceBand::from_raw(band.raw_lower(), band.raw_upper(), metric).unwrap(); + assert_eq!(band, explicit_raw); + let filter = filter_bytes([11]); + for batch in [false, true] { + let query_count = if batch { 2 } else { 1 }; + let queries = query.repeat(query_count); + for selected in [None, Some(filter.as_slice())] { + let result = direct( + &bytes, + Family::Flat, + &queries, + batch, + VectorRangeSearchParams::new(band, 1), + selected, + ) + .unwrap(); + assert_eq!(result.labels(), vec![11; query_count]); + assert_eq!(result.raw_distances(), vec![expected_raw; query_count]); + for query_index in 0..query_count { + let query_result = result.query(query_index); + assert_eq!(query_result.raw_distances, &[expected_raw]); + assert_eq!( + query_result.raw_distances.as_ptr(), + result.raw_distances()[query_index..].as_ptr() + ); + } + let public_value = metric.public_distance(result.raw_distances()[0]); + if metric == MetricType::L2 { + assert_eq!(public_value, 3.0); + } else if metric == MetricType::InnerProduct { + assert_eq!(public_value, 6.0); + } else { + assert_eq!(public_value, 1.0); + } + } + } + } +} + #[test] fn flat_zero_vectors_and_inner_product_extrema_use_public_semantics() { for metric in [MetricType::Cosine, MetricType::InnerProduct] { @@ -914,7 +1007,7 @@ fn flat_zero_vectors_and_inner_product_extrema_use_public_semantics() { let params = VectorRangeSearchParams::new(band(metric), 1); let result = direct(&bytes, Family::Flat, &query, false, params, None).unwrap(); assert_eq!(result.labels().len(), values.len()); - for (&id, &distance) in result.labels().iter().zip(result.distances()) { + for (&id, &distance) in result.labels().iter().zip(result.raw_distances()) { let value = values[id as usize]; let expected = if metric == MetricType::Cosine { if value == 0.0 || first == 0.0 { @@ -942,7 +1035,7 @@ fn flat_zero_vectors_and_inner_product_extrema_use_public_semantics() { .labels() .iter() .copied() - .zip(result.distances().iter().copied()) + .zip(result.raw_distances().iter().copied()) .filter(|&(_, distance)| metric.public_distance(distance) >= public), ); let found = direct( @@ -1053,7 +1146,7 @@ fn pq_streaming_shares_reads_filters_and_propagates_errors() { assert_eq!(result.query(query_index).labels, &[0, count as i64 - 1]); assert_eq!(result.query(query_index).stats.rows_scanned(), 2); assert_eq!(result.query(query_index).stats.early_abandoned(), 0); - assert_eq!(result.query(query_index).distances, &[4.5, 4.5]); + assert_eq!(result.query(query_index).raw_distances, &[4.5, 4.5]); } let mut centroids = reader.pq.centroids().to_vec(); centroids[0] = f32::NAN; diff --git a/core/tests/range_search.rs b/core/tests/range_search.rs index 74486758..61741f6d 100644 --- a/core/tests/range_search.rs +++ b/core/tests/range_search.rs @@ -291,7 +291,7 @@ fn pq_rows(query: QueryResult<'_>) -> Vec<(i64, f32)> { .labels .iter() .copied() - .zip(query.distances.iter().copied()) + .zip(query.raw_distances.iter().copied()) .collect::>(); rows.sort_by_key(|&(id, _)| id); rows @@ -299,7 +299,7 @@ fn pq_rows(query: QueryResult<'_>) -> Vec<(i64, f32)> { fn pq_all_params(nprobe: usize) -> VectorRangeSearchParams { VectorRangeSearchParams::new( - DistanceBand::new(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(), + DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(), nprobe, ) } @@ -335,7 +335,7 @@ fn pq_range_matches_decoded_oracle_for_both_code_widths() { pq_rows(result.query(0)), expected .into_iter() - .filter(|&(_, value)| band.admit(value)) + .filter(|&(_, value)| band.admit_raw(value)) .collect::>() ); } @@ -389,7 +389,7 @@ fn check_pq_range_dense_opq(metric: MetricType) { })); let filter = serialize_roaring(&allowed); let params = VectorRangeSearchParams::new( - DistanceBand::new(Bound::Unbounded, Bound::Unbounded, metric).unwrap(), + DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, metric).unwrap(), index.nlist, ); let bytes = pq_bytes(&index); @@ -423,11 +423,11 @@ fn check_pq_range_dense_opq(metric: MetricType) { assert!(distances[0] < cut && cut < *distances.last().unwrap()); let bands = [ params.band(), - DistanceBand::new(Bound::Unbounded, Bound::Finite(cut), metric).unwrap(), - DistanceBand::new(Bound::Finite(cut), Bound::Unbounded, metric).unwrap(), - DistanceBand::new(Bound::Finite(cut), Bound::Finite(cut.next_up()), metric) + DistanceBand::from_raw(Bound::Unbounded, Bound::Finite(cut), metric).unwrap(), + DistanceBand::from_raw(Bound::Finite(cut), Bound::Unbounded, metric).unwrap(), + DistanceBand::from_raw(Bound::Finite(cut), Bound::Finite(cut.next_up()), metric) .unwrap(), - DistanceBand::new(Bound::Finite(cut.next_down()), Bound::Finite(cut), metric) + DistanceBand::from_raw(Bound::Finite(cut.next_down()), Bound::Finite(cut), metric) .unwrap(), ]; for pool in &pools { @@ -459,7 +459,7 @@ fn check_pq_range_dense_opq(metric: MetricType) { let expected = reference[query_index] .iter() .copied() - .filter(|&(_, value)| band.admit(value)) + .filter(|&(_, value)| band.admit_raw(value)) .collect::>(); let expected_filtered = expected .iter() @@ -530,7 +530,7 @@ fn check_pq_coarse_parallel(metric: MetricType) { .copied() .collect::>(); let params = VectorRangeSearchParams::new( - DistanceBand::new(Bound::Unbounded, Bound::Unbounded, metric).unwrap(), + DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, metric).unwrap(), 2, ); let snapshots = [1, 4].map(|threads| { @@ -644,8 +644,8 @@ fn pq_range_boundaries_partition_the_full_estimated_result() { let mut reader = pq_reader(&index); let bands = [ pq_all_params(4).band(), - DistanceBand::new(Bound::Unbounded, Bound::Finite(cut), MetricType::L2).unwrap(), - DistanceBand::new(Bound::Finite(cut), Bound::Unbounded, MetricType::L2).unwrap(), + DistanceBand::from_raw(Bound::Unbounded, Bound::Finite(cut), MetricType::L2).unwrap(), + DistanceBand::from_raw(Bound::Finite(cut), Bound::Unbounded, MetricType::L2).unwrap(), l2(cut, cut.next_up()), l2(cut.next_down(), cut), l2(cut, cut), @@ -659,7 +659,7 @@ fn pq_range_boundaries_partition_the_full_estimated_result() { expected .iter() .copied() - .filter(|&(_, value)| band.admit(value)) + .filter(|&(_, value)| band.admit_raw(value)) .collect::>() ); assert_eq!(result.query(0).stats.early_abandoned(), 0); @@ -873,7 +873,8 @@ fn pq_range_validates_metric_and_nprobe_even_for_empty_bands() { for band_metric in [MetricType::L2, MetricType::Cosine, MetricType::InnerProduct] { for nprobe in [0, 4] { let band = - DistanceBand::new(Bound::Finite(1.0), Bound::Finite(1.0), band_metric).unwrap(); + DistanceBand::from_raw(Bound::Finite(1.0), Bound::Finite(1.0), band_metric) + .unwrap(); let params = VectorRangeSearchParams::new(band, nprobe); for result in pq_entry_points(&mut reader, &[0.0; 8], params, &filter) .into_iter() @@ -966,7 +967,7 @@ fn pq_range_uses_estimated_not_original_distances() { let original = [0.25; 8]; index.add(&original, &[42], 1); let band = l2(0.0, 0.25); - assert!(!band.admit(fvec_l2sqr(&original, &[0.0; 8]))); + assert!(!band.admit_raw(fvec_l2sqr(&original, &[0.0; 8]))); let result = pq_reader(&index) .range_search(&[0.0; 8], VectorRangeSearchParams::new(band, 1)) .unwrap(); @@ -1144,7 +1145,8 @@ fn check_pq_range_streaming(metric: MetricType, residual: bool) { &queries, query_count, VectorRangeSearchParams::new( - DistanceBand::new(Bound::Unbounded, Bound::Unbounded, metric).unwrap(), + DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, metric) + .unwrap(), 1, ), &filter, @@ -1341,11 +1343,12 @@ fn rq_range_matches_estimated_distance_oracle() { let upper = values[values.len() * 3 / 4].max(lower); for band in [ l2(lower, upper), - DistanceBand::new(Bound::Unbounded, Bound::Finite(upper), MetricType::L2) + DistanceBand::from_raw(Bound::Unbounded, Bound::Finite(upper), MetricType::L2) + .unwrap(), + DistanceBand::from_raw(Bound::Finite(lower), Bound::Unbounded, MetricType::L2) .unwrap(), - DistanceBand::new(Bound::Finite(lower), Bound::Unbounded, MetricType::L2) + DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, MetricType::L2) .unwrap(), - DistanceBand::new(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(), ] { let result = rq_reader(&index) .range_search(&query, VectorRangeSearchParams::new(band, nprobe)) @@ -1353,7 +1356,7 @@ fn rq_range_matches_estimated_distance_oracle() { let expected = oracle .iter() .copied() - .filter(|row| band.admit(row.1)) + .filter(|row| band.admit_raw(row.1)) .collect(); assert_eq!( pairs_of(result.query(0)), @@ -1374,7 +1377,7 @@ fn rq_range_matches_estimated_distance_oracle() { } fn rq_all_distances() -> DistanceBand { - DistanceBand::new(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap() + DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap() } #[test] @@ -1404,7 +1407,7 @@ fn rq_range_bounded_probes_match_oracle_across_many_lists() { for (query_index, query) in queries.chunks_exact(index.d).enumerate() { let expected = rq_estimated_oracle(&index, query, nprobe) .into_iter() - .filter(|row| band.admit(row.1)) + .filter(|row| band.admit_raw(row.1)) .collect::>(); assert_eq!( pairs_of(batch.query(query_index)), @@ -1473,7 +1476,7 @@ fn rq_range_batch_single_filters_and_statistics_agree() { let expected = oracle .iter() .copied() - .filter(|row| allowed.contains(&row.0) && params.band().admit(row.1)) + .filter(|row| allowed.contains(&row.0) && params.band().admit_raw(row.1)) .collect(); assert_eq!(pairs_of(batch.query(query_index)), bits_of(expected)); let stats = batch.query(query_index).stats; @@ -1530,11 +1533,19 @@ fn rq_range_membership_is_estimated_not_exact() { let result = rq_reader(&index) .range_search(&query, VectorRangeSearchParams::new(band, 1)) .unwrap(); - assert_eq!(result.labels().contains(&witness), band.admit(estimated)); - assert_ne!(result.labels().contains(&witness), band.admit(exact)); + assert_eq!( + result.labels().contains(&witness), + band.admit_raw(estimated) + ); + assert_ne!(result.labels().contains(&witness), band.admit_raw(exact)); assert_eq!( pairs_of(result.query(0)), - bits_of(oracle.into_iter().filter(|row| band.admit(row.1)).collect()) + bits_of( + oracle + .into_iter() + .filter(|row| band.admit_raw(row.1)) + .collect() + ) ); } @@ -1820,8 +1831,8 @@ fn rq_range_unified_metric_capability_and_validation_precedence() { let mut reader = rq_reader(&index); let filter = serialize_roaring(&HashSet::from([1])); for band in [ - DistanceBand::new(Bound::Unbounded, Bound::Unbounded, metric).unwrap(), - DistanceBand::new(Bound::Finite(1.0), Bound::Finite(1.0), metric).unwrap(), + DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, metric).unwrap(), + DistanceBand::from_raw(Bound::Finite(1.0), Bound::Finite(1.0), metric).unwrap(), ] { let params = VectorRangeSearchParams::new(band, 1); for result in [ @@ -2099,7 +2110,7 @@ pub fn brute_force_band( .enumerate() .filter_map(|(row, &id)| { let distance = fvec_l2sqr(query, &vectors[row * d..(row + 1) * d]); - band.admit(distance).then_some((id, distance)) + band.admit_raw(distance).then_some((id, distance)) }) .collect() } @@ -2114,7 +2125,7 @@ pub fn pairs_of(result: QueryResult<'_>) -> Vec<(i64, u32)> { let mut pairs: Vec<(i64, u32)> = result .labels .iter() - .zip(result.distances) + .zip(result.raw_distances) .map(|(id, dist)| (*id, dist.to_bits())) .collect(); pairs.sort_unstable(); @@ -2131,7 +2142,7 @@ fn bits_of(rows: Vec<(i64, f32)>) -> Vec<(i64, u32)> { } fn l2(lower: f32, upper: f32) -> DistanceBand { - DistanceBand::new(Bound::Finite(lower), Bound::Finite(upper), MetricType::L2).unwrap() + DistanceBand::from_raw(Bound::Finite(lower), Bound::Finite(upper), MetricType::L2).unwrap() } /// Serializes an allow-list the way the reader's Roaring decoder expects it. @@ -2215,7 +2226,7 @@ fn flat_sq_non_l2_coarse_parallel_batches_preserve_results_stats_filters_and_rea let dimension = 64; for metric in [MetricType::Cosine, MetricType::InnerProduct] { let params = VectorRangeSearchParams::new( - DistanceBand::new(Bound::Unbounded, Bound::Unbounded, metric).unwrap(), + DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, metric).unwrap(), 2, ); for bytes in flat_sq_coarse_batch_fixtures(metric) { @@ -2328,7 +2339,7 @@ fn flat_sq_non_l2_coarse_parallel_errors_precede_payload_io() { let empty_filter = serialize_roaring(&HashSet::new()); for metric in [MetricType::Cosine, MetricType::InnerProduct] { let params = VectorRangeSearchParams::new( - DistanceBand::new(Bound::Unbounded, Bound::Unbounded, metric).unwrap(), + DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, metric).unwrap(), 1, ); for bytes in flat_sq_coarse_batch_fixtures(metric) { @@ -2425,7 +2436,7 @@ fn ivf_sq_range_matches_sq_estimates_for_all_entry_points() { let expected = bits_of( ids.into_iter() .zip(distances) - .filter(|(id, distance)| *id != -1 && band.admit(*distance)) + .filter(|(id, distance)| *id != -1 && band.admit_raw(*distance)) .collect(), ); assert!(!expected.is_empty()); @@ -2472,7 +2483,7 @@ fn ivf_sq_range_uses_estimated_not_original_distance() { labels .into_iter() .zip(distances) - .filter(|(_, distance)| band.admit(*distance)) + .filter(|(_, distance)| band.admit_raw(*distance)) .collect(), ); assert_eq!(pairs_of(result.query(0)), expected); @@ -2496,9 +2507,9 @@ fn ivf_sq_range_preserves_boundaries_and_has_no_top_k_cap() { l2(lower, upper), l2(lower, lower), l2(0.0, 0.0), - DistanceBand::new(Bound::Unbounded, Bound::Finite(upper), MetricType::L2).unwrap(), - DistanceBand::new(Bound::Finite(lower), Bound::Unbounded, MetricType::L2).unwrap(), - DistanceBand::new(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(), + DistanceBand::from_raw(Bound::Unbounded, Bound::Finite(upper), MetricType::L2).unwrap(), + DistanceBand::from_raw(Bound::Finite(lower), Bound::Unbounded, MetricType::L2).unwrap(), + DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(), ] { let result = reader .range_search(&query, VectorRangeSearchParams::new(band, 10)) @@ -2507,7 +2518,7 @@ fn ivf_sq_range_preserves_boundaries_and_has_no_top_k_cap() { ids.iter() .copied() .zip(distances.iter().copied()) - .filter(|(_, distance)| band.admit(*distance)) + .filter(|(_, distance)| band.admit_raw(*distance)) .collect(), ); assert_eq!(pairs_of(result.query(0)), expected); @@ -2532,7 +2543,7 @@ fn ivf_sq_range_parallel_scans_preserve_query_order_and_membership() { .flat_map(|value| vec![value; 65]) .collect(); let mut reader = VectorIndexReader::open(Cursor::new(bytes)).unwrap(); - let whole = DistanceBand::new(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(); + let whole = DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(); let all = reader .range_search_batch(&queries, 3, VectorRangeSearchParams::new(whole, 4)) .unwrap(); @@ -2545,8 +2556,8 @@ fn ivf_sq_range_parallel_scans_preserve_query_order_and_membership() { .labels .iter() .copied() - .zip(all.query(query_index).distances.iter().copied()) - .filter(|(_, distance)| band.admit(*distance)) + .zip(all.query(query_index).raw_distances.iter().copied()) + .filter(|(_, distance)| band.admit_raw(*distance)) .collect(), ) }) @@ -2631,7 +2642,7 @@ fn ivf_sq_range_reuses_shared_lists_and_cache_with_query_local_filters() { index.ids[3].clear(); index.codes[3].clear(); let bytes = serialize_sq(&index); - let whole = DistanceBand::new(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(); + let whole = DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(); let params = VectorRangeSearchParams::new(whole, 4); for budget in [0, 4 * 1024 * 1024] { let trace = Arc::new(Mutex::new(SqReadTrace::default())); @@ -2694,7 +2705,7 @@ fn ivf_sq_range_respects_bounded_reads_and_deduplicates_partial_probes() { trace: Arc::clone(&trace), }); let mut reader = IVFSQIndexReader::open(source).unwrap(); - let band = DistanceBand::new(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(); + let band = DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(); for (nprobe, reads, rows) in [(1, 2, 67), (4, 4, 268)] { *trace.lock().unwrap() = SqReadTrace::default(); let result = reader @@ -2791,14 +2802,15 @@ fn ivf_sq_range_batch_validates_inputs_before_empty_band_shortcuts() { let mut index = build_sq_index(33, 35, 2); index.metric = metric; let mut reader = IVFSQIndexReader::open(Cursor::new(serialize_sq(&index))).unwrap(); - let empty = DistanceBand::new(Bound::Finite(1.0), Bound::Finite(1.0), metric).unwrap(); + let empty = DistanceBand::from_raw(Bound::Finite(1.0), Bound::Finite(1.0), metric).unwrap(); let mismatched_metric = if metric == MetricType::L2 { MetricType::Cosine } else { MetricType::L2 }; let mismatched = - DistanceBand::new(Bound::Finite(1.0), Bound::Finite(1.0), mismatched_metric).unwrap(); + DistanceBand::from_raw(Bound::Finite(1.0), Bound::Finite(1.0), mismatched_metric) + .unwrap(); let query = [0.0; 33]; for (case, queries, query_count, band, nprobe) in [ ("dimension", &query[..32], 1, empty, 2), @@ -2856,7 +2868,7 @@ fn ivf_sq_range_wrappers_delegate_input_validation() { fn ivf_sq_range_propagates_nonfinite_estimates_and_payload_errors() { let index = build_sq_index(33, 35, 2); let bytes = serialize_sq(&index); - let whole = DistanceBand::new(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(); + let whole = DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(); let mut reader = VectorIndexReader::open(Cursor::new(bytes.clone())).unwrap(); for band in [whole, l2(0.0, 1.0)] { let params = VectorRangeSearchParams::new(band, 2); @@ -2930,8 +2942,8 @@ fn a_band_whose_metric_disagrees_with_the_index_is_rejected() { for index_metric in [MetricType::L2, MetricType::Cosine, MetricType::InnerProduct] { let (mut reader, ..) = build_flat_fixture_with_metric(128, 8, 4, index_metric); for band_metric in [MetricType::L2, MetricType::Cosine, MetricType::InnerProduct] { - let band = - DistanceBand::new(Bound::Finite(0.0), Bound::Finite(1.0), band_metric).unwrap(); + let band = DistanceBand::from_raw(Bound::Finite(0.0), Bound::Finite(1.0), band_metric) + .unwrap(); let outcome = reader.range_search(&[0.1; 8], VectorRangeSearchParams::new(band, 4)); if band_metric != index_metric { assert_eq!( @@ -2972,7 +2984,7 @@ fn an_unsupported_index_type_fails_loud_for_every_band() { fn a_whole_space_band_returns_every_probed_row() { let (mut reader, vectors, ids, d) = build_flat_fixture(256, 16, 8); let query = vectors[0..d].to_vec(); - let band = DistanceBand::new(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(); + let band = DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(); let result = reader .range_search(&query, VectorRangeSearchParams::new(band, 8)) .unwrap(); @@ -3169,7 +3181,7 @@ fn a_band_splits_into_adjacent_sub_bands_without_losing_rows() { // version of this test fell into. let seam = *whole .query(0) - .distances + .raw_distances .iter() .find(|d| **d > 0.0) .expect("the parent band must contain a row at a positive distance"); @@ -3183,7 +3195,7 @@ fn a_band_splits_into_adjacent_sub_bands_without_losing_rows() { .query(0) .labels .iter() - .zip(whole.query(0).distances) + .zip(whole.query(0).raw_distances) .find(|(_, d)| **d == seam) .map(|(id, _)| *id) .expect("the seam distance came from this result"); @@ -3371,7 +3383,8 @@ fn the_l2_cutoff_does_not_change_which_rows_are_returned() { let (mut reader, vectors, _ids, d) = build_flat_fixture(512, 16, 8); let query = vectors[0..d].to_vec(); let bounded = l2(0.0, 2.0); - let unbounded = DistanceBand::new(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(); + let unbounded = + DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(); let with_cutoff = reader .range_search(&query, VectorRangeSearchParams::new(bounded, 8)) .unwrap(); @@ -3383,7 +3396,7 @@ fn the_l2_cutoff_does_not_change_which_rows_are_returned() { .query(0) .labels .iter() - .zip(without.query(0).distances) + .zip(without.query(0).raw_distances) .filter(|(_, dist)| **dist < 2.0) .map(|(id, _)| *id) .collect(); @@ -3444,7 +3457,7 @@ fn early_abandoned_counts_exactly_the_rows_above_the_upper_cut() { // non-empty and a row sits exactly on it. let lower = sorted[ids.len() / 8]; let upper = sorted[ids.len() / 2]; - let band = DistanceBand::new(Bound::Finite(lower), Bound::Finite(upper), MetricType::L2) + let band = DistanceBand::from_raw(Bound::Finite(lower), Bound::Finite(upper), MetricType::L2) .expect("a well-ordered band"); // nprobe = nlist, so every row is probed and the expected counts are over @@ -3729,7 +3742,8 @@ fn range_probe_selection_is_invariant_to_batch_size() { let query = vec![BASE]; // nprobe = 1 so only the single best-ranked list is scanned; the band is wide // enough that whichever list is chosen contributes its row. - let band = DistanceBand::new(Bound::Finite(0.0), Bound::Unbounded, MetricType::L2).unwrap(); + let band = + DistanceBand::from_raw(Bound::Finite(0.0), Bound::Unbounded, MetricType::L2).unwrap(); let params = VectorRangeSearchParams::new(band, 1); let alone = reader.range_search(&query, params).unwrap(); @@ -3808,7 +3822,7 @@ fn a_caller_bug_outranks_an_unsupported_family() { // A band whose metric disagrees with the index is also a caller bug, and was // the field the first version of this hoist forgot. let cosine_band = - DistanceBand::new(Bound::Finite(0.0), Bound::Finite(1.0), MetricType::Cosine).unwrap(); + DistanceBand::from_raw(Bound::Finite(0.0), Bound::Finite(1.0), MetricType::Cosine).unwrap(); let err = reader .range_search(&[0.0; 8], VectorRangeSearchParams::new(cosine_band, 4)) .unwrap_err(); @@ -3832,7 +3846,7 @@ fn a_zero_nprobe_outranks_the_metric_capability_gap() { // over a zero nprobe if the metric gap were reported first. let (mut cosine_reader, ..) = build_flat_fixture_with_metric(64, 8, 4, MetricType::Cosine); let cosine_band = - DistanceBand::new(Bound::Finite(0.0), Bound::Finite(1.0), MetricType::Cosine).unwrap(); + DistanceBand::from_raw(Bound::Finite(0.0), Bound::Finite(1.0), MetricType::Cosine).unwrap(); let err = cosine_reader .range_search(&[0.0; 8], VectorRangeSearchParams::new(cosine_band, 0)) .unwrap_err(); diff --git a/cpp/test_vindex.cpp b/cpp/test_vindex.cpp index 8dde10f0..32fa9bc0 100644 --- a/cpp/test_vindex.cpp +++ b/cpp/test_vindex.cpp @@ -27,6 +27,7 @@ #include #include #include +#include #include #define ASSERT_EQ(a, b) do { \ @@ -45,6 +46,24 @@ #include "../c/range_test_support.h" +static_assert(!std::is_aggregate_v); +static_assert(!std::is_default_constructible_v); +static_assert(!std::is_constructible_v); +static_assert(!std::is_constructible_v); +static_assert(std::is_same_v>); +static_assert(std::is_same_v>); + +template +struct HasDistances : std::false_type {}; + +template +struct HasDistances().distances)>> : std::true_type {}; + +static_assert(!HasDistances::value); +static_assert(HasDistances::value); + struct MemBuffer { std::vector data; size_t pos = 0; @@ -366,14 +385,15 @@ static void assert_range_error(Operation operation) { static PaimonVindexRangeSearchResultView range_view( const paimon::vindex::RangeSearchResult& result) { return {result.query_count, result.labels.size(), result.lims.data(), - result.labels.data(), result.distances.data(), result.stats.data(), + result.labels.data(), result.raw_distances.data(), result.stats.data(), result.list_reads}; } static paimon::vindex::RangeSearchParams range_cpp_params( PaimonVindexRangeSearchParams raw) { - return {{raw.band.metric, raw.band.lower_kind, raw.band.lower, - raw.band.upper_kind, raw.band.upper}, raw.nprobe}; + return {paimon::vindex::DistanceBand::from_raw( + raw.band.metric, raw.band.raw_lower_kind, raw.band.raw_lower, + raw.band.raw_upper_kind, raw.band.raw_upper), raw.nprobe}; } static void consume_range_fixture(const RangeFixture* fixture) { @@ -403,6 +423,86 @@ static void consume_range_fixture(const RangeFixture* fixture) { range_fixture_assert(fixture, &view); } +static void test_range_raw_factory() { + using namespace paimon::vindex; + for (uint32_t metric : {PAIMON_VINDEX_METRIC_L2, PAIMON_VINDEX_METRIC_INNER_PRODUCT}) { + const float raw_lower = metric == PAIMON_VINDEX_METRIC_L2 ? 4.0f : -6.0f; + const float raw_upper = metric == PAIMON_VINDEX_METRIC_L2 ? 9.0f : -5.0f; + const auto band = DistanceBand::from_raw( + metric, PAIMON_VINDEX_BOUND_FINITE, raw_lower, PAIMON_VINDEX_BOUND_FINITE, raw_upper); + ASSERT_EQ(band.metric(), metric); + ASSERT_EQ(band.raw_lower_kind(), PAIMON_VINDEX_BOUND_FINITE); + ASSERT_EQ(band.raw_lower(), raw_lower); + ASSERT_EQ(band.raw_upper_kind(), PAIMON_VINDEX_BOUND_FINITE); + ASSERT_EQ(band.raw_upper(), raw_upper); + auto raw = band.to_ffi(); + ASSERT_EQ(raw.metric, band.metric()); + ASSERT_EQ(raw.raw_lower_kind, band.raw_lower_kind()); + ASSERT_EQ(range_float_bits(raw.raw_lower), range_float_bits(raw_lower)); + ASSERT_EQ(raw.raw_upper_kind, band.raw_upper_kind()); + ASSERT_EQ(range_float_bits(raw.raw_upper), range_float_bits(raw_upper)); + } + const RangeSearchParams defaults; + ASSERT_EQ(defaults.band.metric(), PAIMON_VINDEX_METRIC_L2); + ASSERT_EQ(defaults.band.raw_lower_kind(), PAIMON_VINDEX_BOUND_UNBOUNDED); + ASSERT_EQ(defaults.band.raw_upper_kind(), PAIMON_VINDEX_BOUND_UNBOUNDED); + ASSERT_EQ(defaults.nprobe, 1); + printf("PASS range_raw_factory\n"); +} + +static void test_range_endpoint_raw_results() { + using namespace paimon::vindex; + const char* metrics[] = {"l2", "inner_product"}; + const uint32_t metric_codes[] = {PAIMON_VINDEX_METRIC_L2, PAIMON_VINDEX_METRIC_INNER_PRODUCT}; + const float coordinates[][3] = {{3.0f, 4.0f, 5.0f}, {3.0f, 2.5f, 2.0f}}; + const float expected_raw_distances[][2] = {{9.0f, 16.0f}, {-6.0f, -5.0f}}; + const std::vector labels = { + INT64_C(1) << 40, (INT64_C(1) << 40) + 1, (INT64_C(1) << 40) + 2}; + for (size_t metric_index = 0; metric_index < 2; ++metric_index) { + RangeSearchResult result; + { + std::vector data(3 * RANGE_DIMENSION); + std::vector query(RANGE_DIMENSION); + for (size_t row = 0; row < 3; ++row) { + data[row * RANGE_DIMENSION] = coordinates[metric_index][row]; + } + if (metric_codes[metric_index] == PAIMON_VINDEX_METRIC_INNER_PRODUCT) query[0] = 2.0f; + Trainer trainer({{"index.type", "ivf_flat"}, {"dimension", "8"}, + {"nlist", "1"}, {"metric", metrics[metric_index]}}); + Writer writer(trainer.add_training_vectors(data.data(), 3).finish_training()); + writer.add_vectors(labels.data(), data.data(), 3); + MemBuffer buffer; + writer.write_index(make_output(buffer)); + Reader reader(make_input(buffer)); + const auto band = metric_index == 0 + ? DistanceBand::from_endpoints(metric_codes[metric_index], std::nullopt, + DistanceEndpoint{4.0, PAIMON_VINDEX_CUT_LE}) + : DistanceBand::from_endpoints(metric_codes[metric_index], + DistanceEndpoint{5.0, PAIMON_VINDEX_CUT_GE}); + result = reader.range_search(query, RangeSearchParams{band, 1}); + const auto raw_band = DistanceBand::from_raw( + band.metric(), band.raw_lower_kind(), band.raw_lower(), + band.raw_upper_kind(), band.raw_upper()); + auto raw_result = reader.range_search(query, RangeSearchParams{raw_band, 1}); + ASSERT_TRUE(result.labels == raw_result.labels); + ASSERT_TRUE(result.raw_distances == raw_result.raw_distances); + } + const auto view = range_view(result); + range_assert_shape(&view); + ASSERT_EQ(result.query_count, 1); + ASSERT_EQ(result.labels.size(), 2); + ASSERT_EQ(result.raw_distances.size(), 2); + for (size_t expected = 0; expected < 2; ++expected) { + ASSERT_EQ(std::count(result.labels.begin(), result.labels.end(), labels[expected]), 1); + auto found = std::find(result.labels.begin(), result.labels.end(), labels[expected]); + size_t hit = static_cast(found - result.labels.begin()); + ASSERT_EQ(range_float_bits(result.raw_distances[hit]), + range_float_bits(expected_raw_distances[metric_index][expected])); + } + printf("PASS range_endpoint_raw_results %s\n", metrics[metric_index]); + } +} + static void test_range_endpoints() { using namespace paimon::vindex; for (uint32_t metric : {PAIMON_VINDEX_METRIC_L2, PAIMON_VINDEX_METRIC_COSINE, @@ -411,31 +511,31 @@ static void test_range_endpoints() { for (uint32_t upper_op : {PAIMON_VINDEX_CUT_LE, PAIMON_VINDEX_CUT_LT}) { auto band = DistanceBand::from_endpoints( metric, DistanceEndpoint{0.5, lower_op}, DistanceEndpoint{1.0, upper_op}); - for (float distance : {-1.0f, -0.5f, -0.0f, 0.0f, 0.25f, + for (float raw_distance : {-1.0f, -0.5f, -0.0f, 0.0f, 0.25f, std::nextafter(0.25f, 0.0f), 0.5f, 1.0f, std::nextafter(1.0f, 2.0f), 4.0f}) { - if (metric == PAIMON_VINDEX_METRIC_L2 && distance < 0) continue; + if (metric == PAIMON_VINDEX_METRIC_L2 && raw_distance < 0) continue; double public_value = metric == PAIMON_VINDEX_METRIC_L2 - ? static_cast(std::sqrt(distance)) + ? static_cast(std::sqrt(raw_distance)) : metric == PAIMON_VINDEX_METRIC_INNER_PRODUCT - ? -static_cast(distance) : static_cast(distance); + ? -static_cast(raw_distance) : static_cast(raw_distance); bool expected = (lower_op == PAIMON_VINDEX_CUT_GE ? public_value >= 0.5 : public_value > 0.5) && (upper_op == PAIMON_VINDEX_CUT_LE ? public_value <= 1.0 : public_value < 1.0); ASSERT_EQ(expected, - (band.lower_kind == PAIMON_VINDEX_BOUND_UNBOUNDED || distance >= band.lower) && - (band.upper_kind == PAIMON_VINDEX_BOUND_UNBOUNDED || distance < band.upper)); + (band.raw_lower_kind() == PAIMON_VINDEX_BOUND_UNBOUNDED || raw_distance >= band.raw_lower()) && + (band.raw_upper_kind() == PAIMON_VINDEX_BOUND_UNBOUNDED || raw_distance < band.raw_upper())); } } } auto unbounded = DistanceBand::from_endpoints(metric); - ASSERT_EQ(unbounded.lower_kind, PAIMON_VINDEX_BOUND_UNBOUNDED); - ASSERT_EQ(unbounded.upper_kind, PAIMON_VINDEX_BOUND_UNBOUNDED); + ASSERT_EQ(unbounded.raw_lower_kind(), PAIMON_VINDEX_BOUND_UNBOUNDED); + ASSERT_EQ(unbounded.raw_upper_kind(), PAIMON_VINDEX_BOUND_UNBOUNDED); auto precise = DistanceBand::from_endpoints( metric, DistanceEndpoint{std::nextafter(1.0, 2.0), PAIMON_VINDEX_CUT_GE}); - float boundary = metric == PAIMON_VINDEX_METRIC_INNER_PRODUCT ? -1.0f : 1.0f; - ASSERT_TRUE(!((precise.lower_kind == PAIMON_VINDEX_BOUND_UNBOUNDED || boundary >= precise.lower) && - (precise.upper_kind == PAIMON_VINDEX_BOUND_UNBOUNDED || boundary < precise.upper))); + float raw_boundary = metric == PAIMON_VINDEX_METRIC_INNER_PRODUCT ? -1.0f : 1.0f; + ASSERT_TRUE(!((precise.raw_lower_kind() == PAIMON_VINDEX_BOUND_UNBOUNDED || raw_boundary >= precise.raw_lower()) && + (precise.raw_upper_kind() == PAIMON_VINDEX_BOUND_UNBOUNDED || raw_boundary < precise.raw_upper()))); } assert_range_error([] { DistanceBand::from_endpoints(PAIMON_VINDEX_METRIC_L2, @@ -506,12 +606,11 @@ static void test_range_matrix() { ASSERT_TRUE(empty_filter.labels.empty()); ASSERT_TRUE(empty_filter.lims == std::vector({0, 0})); auto bounded = params; - auto bounds = std::minmax_element(retained.distances.begin(), retained.distances.end()); - bounded.band.lower_kind = PAIMON_VINDEX_BOUND_FINITE; - bounded.band.lower = metric.second == PAIMON_VINDEX_METRIC_L2 - ? std::max(0.0f, *bounds.first) : *bounds.first; - bounded.band.upper_kind = PAIMON_VINDEX_BOUND_FINITE; - bounded.band.upper = *bounds.first + (*bounds.second - *bounds.first) / 2.0f; + auto bounds = std::minmax_element(retained.raw_distances.begin(), retained.raw_distances.end()); + bounded.band = DistanceBand::from_raw( + metric.second, PAIMON_VINDEX_BOUND_FINITE, + metric.second == PAIMON_VINDEX_METRIC_L2 ? std::max(0.0f, *bounds.first) : *bounds.first, + PAIMON_VINDEX_BOUND_FINITE, *bounds.first + (*bounds.second - *bounds.first) / 2.0f); auto subset = reader.range_search_batch(queries, RANGE_QUERY_COUNT, bounded); auto subset_view = range_view(subset); range_assert_shape(&subset_view); @@ -519,7 +618,8 @@ static void test_range_matrix() { for (size_t query_index = 0; query_index < RANGE_QUERY_COUNT; ++query_index) { size_t expected = 0; for (size_t hit = batch.lims[query_index]; hit < batch.lims[query_index + 1]; ++hit) { - if (batch.distances[hit] >= bounded.band.lower && batch.distances[hit] < bounded.band.upper) { + if (batch.raw_distances[hit] >= bounded.band.raw_lower() && + batch.raw_distances[hit] < bounded.band.raw_upper()) { ++expected; ASSERT_TRUE(std::find(subset.labels.begin() + subset.lims[query_index], subset.labels.begin() + subset.lims[query_index + 1], batch.labels[hit]) != @@ -528,9 +628,11 @@ static void test_range_matrix() { } ASSERT_EQ(subset.lims[query_index + 1] - subset.lims[query_index], expected); } - bounded.band.lower = bounded.band.upper; + bounded.band = DistanceBand::from_raw( + metric.second, PAIMON_VINDEX_BOUND_FINITE, bounded.band.raw_upper(), + PAIMON_VINDEX_BOUND_FINITE, bounded.band.raw_upper()); auto empty = reader.range_search_batch(queries, RANGE_QUERY_COUNT, bounded); - ASSERT_TRUE(empty.labels.empty() && empty.distances.empty()); + ASSERT_TRUE(empty.labels.empty() && empty.raw_distances.empty()); ASSERT_TRUE(empty.lims == std::vector({0, 0, 0, 0})); ASSERT_EQ(empty.list_reads, 0); assert_range_error([&] { reader.range_search(nullptr, RANGE_DIMENSION, params); }); @@ -579,7 +681,9 @@ static void test_range_matrix() { invalid.nprobe = 0; assert_range_error([&] { reader.range_search(query, invalid); }); invalid = params; - invalid.band.metric = (metric.second + 1) % 3; + invalid.band = DistanceBand::from_raw( + (metric.second + 1) % 3, params.band.raw_lower_kind(), params.band.raw_lower(), + params.band.raw_upper_kind(), params.band.raw_upper()); assert_range_error([&] { reader.range_search(query, invalid); }); auto nan_query = query; nan_query[0] = std::numeric_limits::quiet_NaN(); @@ -603,6 +707,8 @@ int main() { test_worker_callback_reentry_is_rejected(); test_extensible_search_params_forward_query_tuning(); test_range_endpoints(); + test_range_raw_factory(); + test_range_endpoint_raw_results(); test_range_matrix(); range_fixture_run_all(consume_range_fixture); return 0; diff --git a/docs/api.html b/docs/api.html index 97bcb129..7874ac74 100644 --- a/docs/api.html +++ b/docs/api.html @@ -84,7 +84,7 @@

Range search parameters and results

Batch, filteredrange_search_batch_with_roaring_filter(queries, query_count, params, filter_bytes)
- +
ParameterDescription
DistanceBandA validated half-open interval [lower, upper) carrying its metric. Each side is either Finite(f32) or Unbounded; unboundedness is a distinct state rather than a sentinel number. Build it from cuts with new, or from a SQL predicate's endpoints with from_endpoints. An empty band where lower == upper is legal and returns zero rows.
DistanceBandA validated half-open interval [raw_lower, raw_upper) carrying its metric. Each side is either Finite(f32) or Unbounded; unboundedness is a distinct state rather than a sentinel number. Prefer from_endpoints for public predicate endpoints; from_raw explicitly accepts internal f32 cuts without converting them. The raw_lower()/raw_upper() accessors expose converted cuts, not the original endpoints. An empty band where the raw cuts are equal is legal and returns zero rows.
DistanceEndpointOne side of a predicate: the already-folded literal plus its CutOperator. lower accepts Ge/Gt and upper accepts Le/Lt; a mismatch is rejected rather than reinterpreted.
VectorRangeSearchParamsCarries the band and the probe width. Construct with new(band, nprobe). nprobe is clamped to nlist exactly as the Top-K path clamps its width, and 0 is rejected. There is no automatically widening mode: one would be public surface with no working use until it is implemented, and the private fields let it arrive later without breaking new.
@@ -92,7 +92,7 @@

Range search parameters and results

- +
FieldLengthDescription
limsquery_count + 1Offsets: query i spans lims[i]..lims[i+1].
labelsTotal rowsRow IDs, in no promised order.
distancesTotal rowsDistances in the index's own space, so squared for l2.
raw_distancesTotal rowsValues in the index's own space: squared L2, cosine distance, or negative inner product, including family-specific estimates. Public endpoint conversion does not convert returned values; do not compare raw L2/IP results directly with the original endpoints.
query(i).statsPer querylists_probed, rows_scanned (every row read and at least partially evaluated, including rejected ones), rows_committed, and early_abandoned. early_abandoned counts the rows the scan rejected against its abandon cutoff, which under l2 means a distance above the upper cut; it is a diagnostic rather than a work measure, since a row is counted whether the kernel stopped at its first term or its last.
call_stats()One per calllist_reads: the number of non-empty unique lists read, counting an oversized list's several chunks once. It is call-level rather than per query because a batch reads a shared list once and fans it out to every query that selected it, so the count belongs to no single query.
diff --git a/docs/range-search.html b/docs/range-search.html index b3d636e0..8b3908e0 100644 --- a/docs/range-search.html +++ b/docs/range-search.html @@ -28,7 +28,7 @@

Units and ordering

Cuts are expressed in the index's own distance space, not in whatever unit the caller happens to think in.

For IVF-RQ under L2, this is the raw estimate in squared-L2 units. Negative estimates have no real-valued Euclidean radius. The L2 endpoint examples describe non-negative squared distances; endpoint conversion does not turn estimated membership into an exact predicate over the original vectors.

MetricInternal valuePublic predicate value
l2Squared Euclidean distance or estimatef32 square root, then widened to f64
cosine1 - cos or family estimate; no clampingInternal value widened to f64
inner_product-inner_product or family estimateNegated internal value widened to f64
-

MetricType::public_distance defines this conversion. Returned distances and DistanceBand::new always use internal units; DistanceBand::from_endpoints takes public f64 endpoints. Cosine queries are normalized before both probe selection and scanning; zero queries remain zero. FLAT and SQ return cosine distance 1 when either vector has zero norm. PQ and RQ retain their unit-vector estimators even for zero queries, rather than claiming exact cosine membership.

+

MetricType::public_distance defines this conversion. Returned raw_distances and DistanceBand::from_raw always use internal units; DistanceBand::from_endpoints takes public f64 endpoints. Cosine queries are normalized before both probe selection and scanning; zero queries remain zero. FLAT and SQ return cosine distance 1 when either vector has zero norm. PQ and RQ retain their unit-vector estimators even for zero queries, rather than claiming exact cosine membership.

Row order is not part of the contract. Within one list rows come back in physical order, but no order across lists is specified or promised, and two runs of the same query may differ. Do not depend on any observed order: sorting is the caller's job, and in SQL it is ORDER BY's. Results are neither padded nor sorted, which is how range search differs from Top-K.

A row within a few ULP of the upper cutUnder L2, IVF-FLAT can abandon a row when its partially accumulated squared distance passes the upper cut, using the same accumulation for pruning and collection. IVF-SQ similarly prunes blocked squared estimates at the exclusive upper cut. Cosine and inner product never use partial-sum pruning. PQ and RQ always evaluate complete f32 estimates, without top-K's FastScan or coarse-bound pruning. All four families use the supplied cuts without a margin.
@@ -53,20 +53,19 @@

Endpoints and cuts

Usage

Rust · a two-sided band
use paimon_vindex_core::distance::MetricType;
-use paimon_vindex_core::range::{Bound, DistanceBand, VectorRangeSearchParams};
+use paimon_vindex_core::range::{CutOperator, DistanceBand, DistanceEndpoint, VectorRangeSearchParams};
 
-// Squared-L2 cuts, half-open: 0.5 is returned, 1.5 is not.
-let band = DistanceBand::new(
-    Bound::Finite(0.5),
-    Bound::Finite(1.5),
+let band = DistanceBand::from_endpoints(
+    Some(DistanceEndpoint { value: 0.5, op: CutOperator::Ge }),
+    Some(DistanceEndpoint { value: 1.5, op: CutOperator::Lt }),
     MetricType::L2,
 )?;
 let result = reader.range_search(&query, VectorRangeSearchParams::new(band, 16))?;
 
 let rows = result.query(0);
-for (id, distance) in rows.labels.iter().zip(rows.distances) {
-    // `distance` is a squared L2 value; take sqrt for a Euclidean radius.
-    println!("{id} {distance}");
+for (id, raw_distance) in rows.labels.iter().zip(rows.raw_distances) {
+    let public_distance = MetricType::L2.public_distance(*raw_distance);
+    println!("{id} raw={raw_distance} public={public_distance}");
 }
Rust · from a SQL predicate, and an unbounded side
use paimon_vindex_core::range::{CutOperator, DistanceEndpoint};
 
@@ -78,7 +77,11 @@ 

Usage

)?; // A one-sided band: everything at or beyond 2.0, with no upper end. -let tail = DistanceBand::new(Bound::Finite(2.0), Bound::Unbounded, MetricType::L2)?;
+let tail = DistanceBand::from_endpoints( + Some(DistanceEndpoint { value: 2.0, op: CutOperator::Ge }), + None, + MetricType::L2, +)?;
Rust · cosine distance and inner-product similarity
let near_cosine = DistanceBand::from_endpoints(
     None,
     Some(DistanceEndpoint { value: 0.2, op: CutOperator::Le }),
@@ -90,7 +93,7 @@ 

Usage

MetricType::InnerProduct, )?;

The inner-product predicate is a lower bound on public similarity, not on the returned negative-dot score. Endpoint conversion reverses the internal cut automatically.

-

Results use a CSR layout, so a batch of queries shares three contiguous buffers. lims holds query_count + 1 offsets, and query i owns labels[lims[i]..lims[i+1]] together with the matching slice of distances. Per-query counters are available through query(i).stats, and counters covering the whole call through call_stats().

+

Results use a CSR layout, so a batch of queries shares three contiguous buffers. lims holds query_count + 1 offsets, and query i owns labels[lims[i]..lims[i+1]] together with the matching slice of raw_distances. Per-query counters are available through query(i).stats, and counters covering the whole call through call_stats().

All four families expose range_search, range_search_batch, and their _with_roaring_filter variants through both typed readers and VectorIndexReader. The filter is an allow-list and does not widen the fixed nprobe. Roaring filters admit only non-negative row IDs; a direct RowIdFilter may admit signed IDs. A query has the same label/distance multiset alone or in a batch; order remains unspecified. Unique non-empty lists are read at most once per call and shared across queries; IVF-SQ cache hits require no payload read.

For IVF-RQ, lists_probed includes empty selected lists; rows_scanned counts filter-eligible rows evaluated; rows_committed counts returned rows; and early_abandoned is zero. Call-level list_reads counts unique non-empty lists, not query/list pairs or storage read rounds. These result-owned counters leave the last top-K statistics unchanged.

@@ -101,14 +104,14 @@

Language bindings and ownership

- +
BindingAPI and lifetime
C ABI / generated headerpaimon_vindex_reader_range_search, _range_search_batch, and both _with_roaring_filter variants return an owned opaque result through an output pointer. paimon_vindex_range_search_result_view borrows its buffers; paimon_vindex_range_search_result_destroy releases them. Query/filter lengths are explicit. The generated include/paimon_vindex.h is rebuilt by cbindgen, not maintained manually.
C++Reader::range_search, Reader::range_search_batch, and their _with_roaring_filter variants copy CSR buffers into owning vectors. A native-result RAII guard handles cleanup, including allocation exceptions. Existing top-K methods remain unchanged.
JNI / JavaVectorIndexReader.rangeSearch and rangeSearchBatch accept VectorRangeSearchParams and optional filter bytes. VectorRangeSearchResult takes ownership of JNI-created Java arrays. Its public constructor and array-returning accessors still make defensive copies. hitCount(), labelAt(int), distanceAt(int), queryStart(int), queryEnd(int), and per-query counter overloads provide allocation-free consumption. Native lengths are checked against JVM array limits; labels, stored offsets and counters use long, while Java hit indices and half-open query bounds use int.
JNI / JavaVectorIndexReader.rangeSearch and rangeSearchBatch accept VectorRangeSearchParams and optional filter bytes. VectorRangeSearchResult takes ownership of JNI-created Java arrays. Its public constructor and array-returning accessors still make defensive copies. hitCount(), labelAt(int), rawDistanceAt(int), queryStart(int), queryEnd(int), and per-query counter overloads provide allocation-free consumption. Native lengths are checked against JVM array limits; labels, stored offsets and counters use long, while Java hit indices and half-open query bounds use int.
PythonVectorIndexReader.range_search and range_search_batch accept RangeSearchParams and optional roaring_filter=. RangeSearchResult owns copied NumPy arrays (uintp offsets, int64 labels, float32 distances). Native results are freed in finally, even if conversion fails.

Results remain valid after the reader closes. In C, destroy each successful result exactly once, never free its individual arrays, and never access a view after destruction. Input queries and filter bytes are borrowed only for the call. Search sets a valid output pointer to NULL before doing work; an error returns -1, sets paimon_vindex_last_error(), and transfers no result. destroy(NULL) is safe. Callers must supply valid, aligned buffers and live handles; C callers must synchronize operations on a reader. C++, Java and Python retain their existing callback-aware handle locks.

Capability is available as paimon_vindex_reader_supports_range_search(reader, &supported), C++/Python reader.supports_range_search(), or Java reader.supportsRangeSearch(). DiskANN returns false. Invalid parameters, malformed filters, unsupported searches, I/O errors and non-finite evaluated distances fail rather than pretending there are no matches. Length multiplication, slice byte limits and language-specific array limits are checked instead of truncating offsets or counters.

All results expose per-query lists_probed, rows_scanned, rows_committed and early_abandoned, plus call-level list_reads (camelCase accessors in Java). These are core's logical counters, not binding-side estimates. In particular, list_reads is not the sum of a batch's per-query probe counts, and an IVF-SQ cache hit does not count as a payload read.

C · public L2 distance <= 2.0
PaimonVindexDistanceEndpoint upper = {2.0, PAIMON_VINDEX_CUT_LE};
-PaimonVindexDistanceBand band;
+PaimonVindexRawDistanceBand band;
 if (paimon_vindex_distance_band_from_endpoints(
         PAIMON_VINDEX_METRIC_L2, NULL, &upper, &band) != 0) {
     return -1;
@@ -121,7 +124,7 @@ 

Language bindings and ownership

PaimonVindexRangeSearchResultView view; status = paimon_vindex_range_search_result_view(result, &view); if (status == 0) { - consume_rows(view.labels, view.distances, view.hit_count); + consume_rows(view.labels, view.raw_distances, view.hit_count); } } paimon_vindex_range_search_result_destroy(result);
@@ -131,7 +134,10 @@

Language bindings and ownership

DistanceEndpoint{2.0, PAIMON_VINDEX_CUT_LE}); auto result = reader.range_search_batch(queries, query_count, RangeSearchParams{band, 8}); auto first_begin = result.lims[0]; -auto first_end = result.lims[1]; +auto first_end = result.lims[1]; +for (size_t hit_index = first_begin; hit_index < first_end; ++hit_index) { + consume(result.labels[hit_index], result.raw_distances[hit_index]); +}
Java · shared filtered batch
VectorDistanceBand band = VectorDistanceBand.fromEndpoints(
         "l2", null, null, 2.0, VectorDistanceBand.CutOperator.LE);
 VectorRangeSearchParams params = new VectorRangeSearchParams(band, 8);
@@ -141,8 +147,8 @@ 

Language bindings and ownership

for (int hitIndex = result.queryStart(queryIndex); hitIndex < result.queryEnd(queryIndex); hitIndex++) { long label = result.labelAt(hitIndex); - float distance = result.distanceAt(hitIndex); - consume(label, distance); + float rawDistance = result.rawDistanceAt(hitIndex); + consume(label, rawDistance); } long rowsScanned = result.rowsScanned(queryIndex); } @@ -154,9 +160,12 @@

Language bindings and ownership

"l2", upper=DistanceEndpoint(2.0, DistanceEndpointOp.LE)) result = reader.range_search_batch( queries, RangeSearchParams(band, nprobe=8), roaring_filter=roaring_filter) -labels, distances = result.query(0) +query_result = result.query(0) +labels, raw_distances = query_result.labels, query_result.raw_distances stats = result.stats[0]
-

Raw band constructors take float32 cuts in internal distance space. To express public predicate endpoints, use the native-backed conversion helpers instead: do not square L2 endpoints, round double literals to float32, or negate/reverse inner-product endpoints yourself. Missing endpoints are structural unboundedness, not infinities or sentinel numbers.

+

Use from_endpoints (fromEndpoints in Java) for public predicate endpoints. Explicit raw construction uses DistanceBand::from_raw in Rust/C++, DistanceBand.from_raw in Python, VectorDistanceBand.fromRaw in Java, or the C PaimonVindexRawDistanceBand type. Raw constructors take float32 cuts in internal distance space; they do not convert units or endpoint operators. Do not square L2 endpoints, round double literals to float32, or negate/reverse inner-product endpoints yourself. Missing endpoints are structural unboundedness, not infinities or sentinel numbers.

+
Public endpoints are not raw results.For L2, a public radius of 4 can return a raw_distance of 9, whose public distance is 3. An inner-product predicate of at least 5 can return a raw value of -6, whose public similarity is 6. The raw_lower/raw_upper cuts (rawLower()/rawUpper() in Java) are transformed half-open boundaries, not the original endpoint values; inner product also reverses their sides. Cosine uses cosine distance, not cosine similarity.
+

Range API migration: replace distances with raw_distances; Java uses rawDistances(), rawDistanceAt(), and rawDistancesForQuery(). Replace ambiguous raw-band construction with an explicit raw factory or, preferably, the public-endpoint factory. Python query views expose labels and raw_distances while retaining tuple unpacking. These are source API changes, including Rust; Java method renames also require recompiling callers. C field/type renames preserve the native layout and exported function symbols. Existing top-K names, result values, index formats, and ownership rules are unchanged.

Python can still import and run existing top-K calls with a native library predating the range ABI. If any required range export is absent, capability returns false and new range calls or endpoint conversion request a native-library upgrade with a clear error. Range operations never silently fall back to top-K. NumPy query inputs are normalized for both contiguity and alignment.

@@ -183,7 +192,7 @@

Cross-language verification

Choosing an index type

Range search supports IVF-FLAT, IVF-SQ, IVF-PQ and IVF-RQ under L2, cosine and inner product, with fixed positive probe widths in Rust and every language binding. Query capability through core or the binding's reader capability method. DiskANN remains unsupported; storage-format and top-K behavior are unchanged.

Choose according to the membership requirement.IVF-FLAT tests full-vector distances. IVF-RQ uses RaBitQ, with a one-bit estimate for one-bit files and the full multi-bit estimate otherwise, not Faiss's residual/additive quantizer. Its band predicate is precise relative to that estimate, not to the raw vector. Tests cover an independent estimated-distance oracle, single/batch equivalence, filters, statistics, parallel scans, and non-finite inputs/data; they do not establish exact-distance recall guarantees. See IVF-RQ range semantics.
-
IVF-SQ membership uses an estimate.IVF-FLAT computes exact distances from stored f32 vectors. IVF-SQ instead reuses top-K's blocked scalar-quantized estimator, reconstructing residuals with each list's stored bounds and centroid. The same estimated value determines band membership and is returned in distances; there is no original-vector reranking and no top-K fallback. Prefer IVF-FLAT if original-distance membership must be exact.
+
IVF-SQ membership uses an estimate.IVF-FLAT computes exact distances from stored f32 vectors. IVF-SQ instead reuses top-K's blocked scalar-quantized estimator, reconstructing residuals with each list's stored bounds and centroid. The same estimated value determines band membership and is returned in raw_distances; there is no original-vector reranking and no top-K fallback. Prefer IVF-FLAT if original-distance membership must be exact.
IVF-PQ uses complete floating-point ADC estimates.Both packed 4-bit and 8-bit codes, residual encoding and OPQ are supported. L2 sums direct squared subvector distances. Cosine uses half the ADC squared distance after query normalization, a unit-vector surrogate rather than exact cosine of a re-normalized reconstruction. Inner product uses negative estimated dot product. The range path does not reuse top-K's quantized FastScan tables, expanded L2 tables or cosine score scale, so scores need not be bit-identical to top-K. It never truncates or reranks raw vectors. Shared lists are read once, oversized lists stream in bounded chunks, and the allow-list is evaluated once per list row across the batch.

A reproducible boundary example is in core/tests/range_search.rs: with a one-dimensional centroid of zero and SQ bounds [0, 255], inputs 0.49 and 0.51 quantize to 0 and 1. For query zero, band [0, 0.1) includes the first estimate despite its true squared distance being outside; band [0.2, 0.3) misses both although both true squared distances lie inside. This demonstrates the membership gap, not a general recall estimate.

Performance and memory: Under L2, IVF-SQ keeps the existing SIMD block layout and uses a finite upper cut to abandon a block once all partial squared distances reach that exclusive cut. A lower cut alone cannot prune a partial sum; cosine and IP do not prune partial sums. Batch queries read each unique list once, reuse cached partitions, and keep query-owned collectors instead of materializing a list-by-query result matrix. Large single queries can scan lists in parallel and merge once per list, not per row. Oversized lists stream in bounded chunks, with reusable scan scratch. Output memory still grows with all admitted rows; there is no result cap.

diff --git a/ffi/examples/range_search_fixture.rs b/ffi/examples/range_search_fixture.rs index 77eca7fa..15f6dc90 100644 --- a/ffi/examples/range_search_fixture.rs +++ b/ffi/examples/range_search_fixture.rs @@ -158,7 +158,8 @@ fn write_case( ) -> io::Result<()> { let mut reader = VectorIndexReader::open(Cursor::new(index.to_vec()))?; assert!(reader.supports_range_search()); - let params = VectorRangeSearchParams::new(DistanceBand::new(lower, upper, metric)?, NPROBE); + let params = + VectorRangeSearchParams::new(DistanceBand::from_raw(lower, upper, metric)?, NPROBE); let result = match (query_count == 1, filter) { (true, None) => reader.range_search(queries, params), (true, Some(filter)) => reader.range_search_with_roaring_filter(queries, params, filter), @@ -188,7 +189,7 @@ fn write_case( for label in result.labels() { writeln!(file, "{label}")?; } - for distance in result.distances() { + for distance in result.raw_distances() { writeln!(file, "{}", distance.to_bits())?; } for query in 0..query_count { diff --git a/ffi/src/range.rs b/ffi/src/range.rs index c8928e08..32d1a92e 100644 --- a/ffi/src/range.rs +++ b/ffi/src/range.rs @@ -27,16 +27,18 @@ pub const PAIMON_VINDEX_CUT_GT: u32 = 1; pub const PAIMON_VINDEX_CUT_LE: u32 = 2; pub const PAIMON_VINDEX_CUT_LT: u32 = 3; -/// Internal half-open distance band: squared L2, 1-cosine, or negative inner product. -/// Unbounded sides ignore their value field; finite sides must be finite. +/// Explicitly raw half-open distance band: squared L2, 1-cosine, or negative +/// inner product. Prefer paimon_vindex_distance_band_from_endpoints for public +/// predicates. Raw cuts are not the original endpoints; inner product reverses +/// their sides. Unbounded sides ignore their value; finite sides must be finite. #[repr(C)] #[derive(Clone, Copy)] -pub struct PaimonVindexDistanceBand { +pub struct PaimonVindexRawDistanceBand { pub metric: u32, - pub lower_kind: u32, - pub lower: f32, - pub upper_kind: u32, - pub upper: f32, + pub raw_lower_kind: u32, + pub raw_lower: f32, + pub raw_upper_kind: u32, + pub raw_upper: f32, } /// Public-distance predicate endpoint. Values are passed to core without rounding. @@ -52,7 +54,7 @@ pub struct PaimonVindexDistanceEndpoint { #[repr(C)] #[derive(Clone, Copy)] pub struct PaimonVindexRangeSearchParams { - pub band: PaimonVindexDistanceBand, + pub band: PaimonVindexRawDistanceBand, pub nprobe: usize, } @@ -66,15 +68,17 @@ pub struct PaimonVindexRangeSearchStats { } /// Borrowed CSR view; all pointers remain valid until the owning result is destroyed. -/// lims has query_count + 1 entries; labels/distances have hit_count entries; +/// lims has query_count + 1 entries; labels/raw_distances have hit_count entries; /// stats has query_count entries. Empty arrays must not be dereferenced. +/// raw_distances retain squared-L2, cosine-distance, or negative-inner-product +/// units; they cannot be compared directly with public L2 or IP endpoints. #[repr(C)] pub struct PaimonVindexRangeSearchResultView { pub query_count: usize, pub hit_count: usize, pub lims: *const usize, pub labels: *const i64, - pub distances: *const f32, + pub raw_distances: *const f32, pub stats: *const PaimonVindexRangeSearchStats, pub list_reads: usize, } @@ -103,9 +107,9 @@ fn range_bound(kind: u32, value: f32) -> Result { } fn range_params(params: PaimonVindexRangeSearchParams) -> Result { - let band = DistanceBand::new( - range_bound(params.band.lower_kind, params.band.lower)?, - range_bound(params.band.upper_kind, params.band.upper)?, + let band = DistanceBand::from_raw( + range_bound(params.band.raw_lower_kind, params.band.raw_lower)?, + range_bound(params.band.raw_upper_kind, params.band.raw_upper)?, range_metric(params.band.metric)?, ) .map_err(|error| format!("range band: {error}"))?; @@ -132,19 +136,19 @@ unsafe fn range_endpoint( })) } -fn range_band_to_ffi(band: DistanceBand) -> PaimonVindexDistanceBand { +fn range_band_to_ffi(band: DistanceBand) -> PaimonVindexRawDistanceBand { let encode = |bound| match bound { Bound::Unbounded => (PAIMON_VINDEX_BOUND_UNBOUNDED, 0.0), Bound::Finite(value) => (PAIMON_VINDEX_BOUND_FINITE, value), }; - let (lower_kind, lower) = encode(band.lower()); - let (upper_kind, upper) = encode(band.upper()); - PaimonVindexDistanceBand { + let (raw_lower_kind, raw_lower) = encode(band.raw_lower()); + let (raw_upper_kind, raw_upper) = encode(band.raw_upper()); + PaimonVindexRawDistanceBand { metric: metric_code(band.metric()), - lower_kind, - lower, - upper_kind, - upper, + raw_lower_kind, + raw_lower, + raw_upper_kind, + raw_upper, } } @@ -156,7 +160,7 @@ pub unsafe extern "C" fn paimon_vindex_distance_band_from_endpoints( metric: u32, lower: *const PaimonVindexDistanceEndpoint, upper: *const PaimonVindexDistanceEndpoint, - out: *mut PaimonVindexDistanceBand, + out: *mut PaimonVindexRawDistanceBand, ) -> c_int { ffi_status(|| { if out.is_null() { @@ -390,7 +394,7 @@ pub unsafe extern "C" fn paimon_vindex_range_search_result_view( hit_count: result.inner.labels().len(), lims: result.inner.lims().as_ptr(), labels: result.inner.labels().as_ptr(), - distances: result.inner.distances().as_ptr(), + raw_distances: result.inner.raw_distances().as_ptr(), stats: result.stats.as_ptr(), list_reads: result.inner.call_stats().list_reads(), } diff --git a/ffi/src/range_tests.rs b/ffi/src/range_tests.rs index 272f1bb5..7abbbf83 100644 --- a/ffi/src/range_tests.rs +++ b/ffi/src/range_tests.rs @@ -18,6 +18,54 @@ use super::*; use paimon_vindex_core::range::{Bound, CutOperator, DistanceBand, DistanceEndpoint}; +#[test] +fn raw_names_preserve_c_abi_layout() { + use std::mem::{align_of, offset_of, size_of}; + + assert_eq!(size_of::(), 20); + assert_eq!(align_of::(), 4); + assert_eq!(offset_of!(PaimonVindexRawDistanceBand, metric), 0); + assert_eq!(offset_of!(PaimonVindexRawDistanceBand, raw_lower_kind), 4); + assert_eq!(offset_of!(PaimonVindexRawDistanceBand, raw_lower), 8); + assert_eq!(offset_of!(PaimonVindexRawDistanceBand, raw_upper_kind), 12); + assert_eq!(offset_of!(PaimonVindexRawDistanceBand, raw_upper), 16); + + let word = size_of::(); + assert_eq!(size_of::(), 7 * word); + assert_eq!( + align_of::(), + align_of::() + ); + assert_eq!( + offset_of!(PaimonVindexRangeSearchResultView, query_count), + 0 + ); + assert_eq!( + offset_of!(PaimonVindexRangeSearchResultView, hit_count), + word + ); + assert_eq!( + offset_of!(PaimonVindexRangeSearchResultView, lims), + 2 * word + ); + assert_eq!( + offset_of!(PaimonVindexRangeSearchResultView, labels), + 3 * word + ); + assert_eq!( + offset_of!(PaimonVindexRangeSearchResultView, raw_distances), + 4 * word + ); + assert_eq!( + offset_of!(PaimonVindexRangeSearchResultView, stats), + 5 * word + ); + assert_eq!( + offset_of!(PaimonVindexRangeSearchResultView, list_reads), + 6 * word + ); +} + #[test] fn range_endpoints_match_core_for_every_metric_and_operator() { for (metric_code, metric) in [ @@ -74,8 +122,16 @@ fn range_endpoints_match_core_for_every_metric_and_operator() { let actual = unsafe { actual.assume_init() }; assert_eq!(actual.metric, metric_code); for (kind, value, expected) in [ - (actual.lower_kind, actual.lower, expected.lower()), - (actual.upper_kind, actual.upper, expected.upper()), + ( + actual.raw_lower_kind, + actual.raw_lower, + expected.raw_lower(), + ), + ( + actual.raw_upper_kind, + actual.raw_upper, + expected.raw_upper(), + ), ] { match expected { Bound::Unbounded => assert_eq!(kind, PAIMON_VINDEX_BOUND_UNBOUNDED), @@ -94,12 +150,12 @@ fn range_endpoints_match_core_for_every_metric_and_operator() { #[test] fn range_errors_leave_no_owned_result() { let params = PaimonVindexRangeSearchParams { - band: PaimonVindexDistanceBand { + band: PaimonVindexRawDistanceBand { metric: 0, - lower_kind: 0, - lower: 0.0, - upper_kind: 0, - upper: 0.0, + raw_lower_kind: 0, + raw_lower: 0.0, + raw_upper_kind: 0, + raw_upper: 0.0, }, nprobe: 1, }; @@ -193,12 +249,12 @@ unsafe extern "C" fn range_read( fn unbounded_params() -> PaimonVindexRangeSearchParams { PaimonVindexRangeSearchParams { - band: PaimonVindexDistanceBand { + band: PaimonVindexRawDistanceBand { metric: 0, - lower_kind: 0, - lower: f32::NAN, - upper_kind: 0, - upper: f32::NAN, + raw_lower_kind: 0, + raw_lower: f32::NAN, + raw_upper_kind: 0, + raw_upper: f32::NAN, }, nprobe: 1, } @@ -235,7 +291,7 @@ fn range_result_outlives_reader_and_preserves_signed_labels_and_statistics() { &queries, 2, paimon_vindex_core::range::VectorRangeSearchParams::new( - DistanceBand::new(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(), + DistanceBand::from_raw(Bound::Unbounded, Bound::Unbounded, MetricType::L2).unwrap(), 1, ), ) @@ -258,8 +314,8 @@ fn range_result_outlives_reader_and_preserves_signed_labels_and_statistics() { expected.labels() ); assert_eq!( - unsafe { slice::from_raw_parts(view.distances, view.hit_count) }, - expected.distances() + unsafe { slice::from_raw_parts(view.raw_distances, view.hit_count) }, + expected.raw_distances() ); assert_eq!(view.list_reads, expected.call_stats().list_reads()); for (query, actual) in unsafe { slice::from_raw_parts(view.stats, 2) } @@ -363,7 +419,7 @@ fn range_validation_and_io_errors_do_not_transfer_results() { } params = unbounded_params(); for kind in [2, u32::MAX] { - params.band.lower_kind = kind; + params.band.raw_lower_kind = kind; assert_eq!( unsafe { paimon_vindex_reader_range_search( @@ -378,9 +434,9 @@ fn range_validation_and_io_errors_do_not_transfer_results() { ); } params = unbounded_params(); - params.band.lower_kind = PAIMON_VINDEX_BOUND_FINITE; + params.band.raw_lower_kind = PAIMON_VINDEX_BOUND_FINITE; for value in [f32::NAN, f32::INFINITY, -1.0] { - params.band.lower = value; + params.band.raw_lower = value; assert_eq!( unsafe { paimon_vindex_reader_range_search( diff --git a/include/paimon_vindex.hpp b/include/paimon_vindex.hpp index 72424073..c6e8308b 100644 --- a/include/paimon_vindex.hpp +++ b/include/paimon_vindex.hpp @@ -229,30 +229,42 @@ struct SearchResult { using DistanceEndpoint = PaimonVindexDistanceEndpoint; using RangeSearchStats = PaimonVindexRangeSearchStats; -struct DistanceBand { - uint32_t metric = PAIMON_VINDEX_METRIC_L2; - uint32_t lower_kind = PAIMON_VINDEX_BOUND_UNBOUNDED; - float lower = 0.0f; - uint32_t upper_kind = PAIMON_VINDEX_BOUND_UNBOUNDED; - float upper = 0.0f; - +class DistanceBand { +public: static DistanceBand from_endpoints( uint32_t metric, std::optional lower = std::nullopt, std::optional upper = std::nullopt) { - PaimonVindexDistanceBand raw{}; + PaimonVindexRawDistanceBand raw{}; check(paimon_vindex_distance_band_from_endpoints( metric, lower ? &*lower : nullptr, upper ? &*upper : nullptr, &raw)); - return {raw.metric, raw.lower_kind, raw.lower, raw.upper_kind, raw.upper}; + return DistanceBand(raw); } - PaimonVindexDistanceBand to_ffi() const { - return {metric, lower_kind, lower, upper_kind, upper}; + static DistanceBand from_raw( + uint32_t metric, uint32_t raw_lower_kind, float raw_lower, + uint32_t raw_upper_kind, float raw_upper) { + return DistanceBand({metric, raw_lower_kind, raw_lower, raw_upper_kind, raw_upper}); } + + uint32_t metric() const { return raw_.metric; } + uint32_t raw_lower_kind() const { return raw_.raw_lower_kind; } + float raw_lower() const { return raw_.raw_lower; } + uint32_t raw_upper_kind() const { return raw_.raw_upper_kind; } + float raw_upper() const { return raw_.raw_upper; } + + PaimonVindexRawDistanceBand to_ffi() const { return raw_; } + +private: + explicit DistanceBand(PaimonVindexRawDistanceBand raw) : raw_(raw) {} + + PaimonVindexRawDistanceBand raw_; }; struct RangeSearchParams { - DistanceBand band; + DistanceBand band = DistanceBand::from_raw( + PAIMON_VINDEX_METRIC_L2, PAIMON_VINDEX_BOUND_UNBOUNDED, 0.0f, + PAIMON_VINDEX_BOUND_UNBOUNDED, 0.0f); size_t nprobe = 1; PaimonVindexRangeSearchParams to_ffi() const { @@ -264,7 +276,7 @@ struct RangeSearchResult { size_t query_count = 0; std::vector lims; std::vector labels; - std::vector distances; + std::vector raw_distances; std::vector stats; size_t list_reads = 0; }; @@ -279,7 +291,7 @@ inline RangeSearchResult copy_range_result(PaimonVindexRangeSearchResult* raw, i PaimonVindexRangeSearchResultView view{}; check(paimon_vindex_range_search_result_view(guard.get(), &view)); if (view.query_count == std::numeric_limits::max() || !view.lims || - (view.hit_count != 0 && (!view.labels || !view.distances)) || + (view.hit_count != 0 && (!view.labels || !view.raw_distances)) || (view.query_count != 0 && !view.stats)) { throw Error("invalid native range result view"); } @@ -288,7 +300,7 @@ inline RangeSearchResult copy_range_result(PaimonVindexRangeSearchResult* raw, i result.lims.assign(view.lims, view.lims + view.query_count + 1); if (view.hit_count != 0) { result.labels.assign(view.labels, view.labels + view.hit_count); - result.distances.assign(view.distances, view.distances + view.hit_count); + result.raw_distances.assign(view.raw_distances, view.raw_distances + view.hit_count); } if (view.query_count != 0) { result.stats.assign(view.stats, view.stats + view.query_count); diff --git a/java/src/main/java/org/apache/paimon/index/vector/VectorDistanceBand.java b/java/src/main/java/org/apache/paimon/index/vector/VectorDistanceBand.java index ad3c4276..e539f46a 100644 --- a/java/src/main/java/org/apache/paimon/index/vector/VectorDistanceBand.java +++ b/java/src/main/java/org/apache/paimon/index/vector/VectorDistanceBand.java @@ -20,8 +20,8 @@ import java.util.Objects; /** - * A half-open band [lower, upper) in raw index distance space: squared L2, 1 - cosine, or negative - * inner product. A null cut is structurally unbounded, not a finite sentinel. + * A half-open band [rawLower, rawUpper) in raw index distance space: squared L2, 1 - cosine, or + * negative inner product. A null cut is structurally unbounded, not a finite sentinel. */ public final class VectorDistanceBand { @@ -39,21 +39,26 @@ public enum CutOperator { } private final String metric; - private final Float lower; - private final Float upper; + private final Float rawLower; + private final Float rawUpper; - public VectorDistanceBand(String metric, Float lower, Float upper) { + private VectorDistanceBand(String metric, Float rawLower, Float rawUpper) { this.metric = Objects.requireNonNull(metric, "metric"); if (!"l2".equals(metric) && !"cosine".equals(metric) && !"inner_product".equals(metric)) { throw new IllegalArgumentException("unknown metric: " + metric); } - validateCut(lower); - validateCut(upper); - if (lower != null && upper != null && lower > upper) { + validateRawCut(rawLower); + validateRawCut(rawUpper); + if (rawLower != null && rawUpper != null && rawLower > rawUpper) { throw new IllegalArgumentException("inverted distance band"); } - this.lower = lower; - this.upper = upper; + this.rawLower = rawLower; + this.rawUpper = rawUpper; + } + + /** Creates raw half-open cuts without converting endpoint distances or similarities. */ + public static VectorDistanceBand fromRaw(String metric, Float rawLower, Float rawUpper) { + return new VectorDistanceBand(metric, rawLower, rawUpper); } /** @@ -87,16 +92,16 @@ public String metric() { return metric; } - public Float lower() { - return lower; + public Float rawLower() { + return rawLower; } - public Float upper() { - return upper; + public Float rawUpper() { + return rawUpper; } - private void validateCut(Float cut) { - if (cut != null && (!Float.isFinite(cut) || ("l2".equals(metric) && cut < 0))) { + private void validateRawCut(Float rawCut) { + if (rawCut != null && (!Float.isFinite(rawCut) || ("l2".equals(metric) && rawCut < 0))) { throw new IllegalArgumentException( "cut must be finite and non-negative for squared L2"); } diff --git a/java/src/main/java/org/apache/paimon/index/vector/VectorRangeSearchResult.java b/java/src/main/java/org/apache/paimon/index/vector/VectorRangeSearchResult.java index 402431b3..8ac55448 100644 --- a/java/src/main/java/org/apache/paimon/index/vector/VectorRangeSearchResult.java +++ b/java/src/main/java/org/apache/paimon/index/vector/VectorRangeSearchResult.java @@ -21,13 +21,14 @@ import java.util.Objects; /** - * CSR range-search output in core scan order, with raw distances and no sorting or top-K cap. + * CSR range-search output in core scan order, with raw distances (squared L2, 1 - cosine, or + * negative inner product) and no sorting or top-K cap. * IVF-Flat distances are exact; SQ, PQ and RQ return their core distance estimates. Each query * occupies [lims[query], lims[query + 1]). Public construction and array-returning accessors make * defensive copies; native construction takes exclusive ownership of its arrays. * *

For allocation-free consumption, use {@link #hitCount()}, {@link #labelAt(int)}, - * {@link #distanceAt(int)}, and the half-open bounds {@link #queryStart(int)} and + * {@link #rawDistanceAt(int)}, and the half-open bounds {@link #queryStart(int)} and * {@link #queryEnd(int)}. Counter accessors taking a query index also avoid copying arrays. * Hit indices and query bounds fit in {@code int}, while labels and counters retain their full * {@code long} range. Indexed accessors reject out-of-range indices with @@ -36,7 +37,7 @@ public final class VectorRangeSearchResult { private final long[] labels; - private final float[] distances; + private final float[] rawDistances; private final long[] lims; private final long[] listsProbed; private final long[] rowsScanned; @@ -46,7 +47,7 @@ public final class VectorRangeSearchResult { static VectorRangeSearchResult fromNative( long[] labels, - float[] distances, + float[] rawDistances, long[] lims, long[] listsProbed, long[] rowsScanned, @@ -55,7 +56,7 @@ static VectorRangeSearchResult fromNative( long listReads) { return new VectorRangeSearchResult( labels, - distances, + rawDistances, lims, listsProbed, rowsScanned, @@ -67,7 +68,7 @@ static VectorRangeSearchResult fromNative( public VectorRangeSearchResult( long[] labels, - float[] distances, + float[] rawDistances, long[] lims, long[] listsProbed, long[] rowsScanned, @@ -76,7 +77,7 @@ public VectorRangeSearchResult( long listReads) { this( labels, - distances, + rawDistances, lims, listsProbed, rowsScanned, @@ -88,7 +89,7 @@ public VectorRangeSearchResult( private VectorRangeSearchResult( long[] labels, - float[] distances, + float[] rawDistances, long[] lims, long[] listsProbed, long[] rowsScanned, @@ -97,12 +98,12 @@ private VectorRangeSearchResult( long listReads, boolean copyArrays) { Objects.requireNonNull(labels, "labels"); - Objects.requireNonNull(distances, "distances"); + Objects.requireNonNull(rawDistances, "rawDistances"); Objects.requireNonNull(lims, "lims"); this.labels = copyArrays ? labels.clone() : labels; - this.distances = copyArrays ? distances.clone() : distances; + this.rawDistances = copyArrays ? rawDistances.clone() : rawDistances; this.lims = copyArrays ? lims.clone() : lims; - if (this.labels.length != this.distances.length + if (this.labels.length != this.rawDistances.length || this.lims.length == 0 || this.lims[0] != 0 || this.lims[this.lims.length - 1] != this.labels.length) { @@ -150,12 +151,12 @@ public long labelAt(int hitIndex) { return labels[hitIndex]; } - public float[] distances() { - return distances.clone(); + public float[] rawDistances() { + return rawDistances.clone(); } - public float distanceAt(int hitIndex) { - return distances[hitIndex]; + public float rawDistanceAt(int hitIndex) { + return rawDistances[hitIndex]; } public long[] lims() { @@ -212,10 +213,10 @@ public long[] labelsForQuery(int queryIndex) { labels, Math.toIntExact(lims[queryIndex]), Math.toIntExact(lims[queryIndex + 1])); } - public float[] distancesForQuery(int queryIndex) { + public float[] rawDistancesForQuery(int queryIndex) { checkQueryIndex(queryIndex); return Arrays.copyOfRange( - distances, + rawDistances, Math.toIntExact(lims[queryIndex]), Math.toIntExact(lims[queryIndex + 1])); } diff --git a/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeOracleTest.java b/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeOracleTest.java index 2493bf37..8a39a163 100644 --- a/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeOracleTest.java +++ b/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeOracleTest.java @@ -103,7 +103,7 @@ private static void runCase(Path expected, Path index) throws IOException { require(metricCode >= 0 && metricCode < metrics.length, "metric code"); VectorRangeSearchParams params = new VectorRangeSearchParams( - new VectorDistanceBand(metrics[metricCode], lower, upper), nprobe); + VectorDistanceBand.fromRaw(metrics[metricCode], lower, upper), nprobe); try (VectorIndexReader reader = new VectorIndexReader( new VectorIndexNativeValidationTest.ByteArraySeekableInputStream( @@ -170,7 +170,7 @@ private static Map> rows(VectorRangeSearchResult result, int hitIndex < result.queryEnd(queryIndex); hitIndex++) { rows.computeIfAbsent(result.labelAt(hitIndex), label -> new ArrayList()) - .add(Float.floatToRawIntBits(result.distanceAt(hitIndex))); + .add(Float.floatToRawIntBits(result.rawDistanceAt(hitIndex))); } for (List values : rows.values()) { Collections.sort(values); diff --git a/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeSearchTest.java b/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeSearchTest.java index 4d600781..96a17c56 100644 --- a/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeSearchTest.java +++ b/java/src/test/java/org/apache/paimon/index/vector/VectorIndexRangeSearchTest.java @@ -23,6 +23,8 @@ import static org.apache.paimon.index.vector.VectorDistanceBand.CutOperator.LT; import java.lang.reflect.Field; +import java.lang.reflect.Method; +import java.lang.reflect.Modifier; import java.nio.ByteBuffer; import java.nio.ByteOrder; import java.util.Arrays; @@ -48,21 +50,43 @@ public static void main(String[] args) { } static void testValueTypes() { + testRawApiContract(); testNativeResultOwnership(); testIndexedResultAccess(); - VectorDistanceBand band = new VectorDistanceBand("l2", null, 4.0f); - check(band.lower() == null && band.upper() == 4.0f, "structural bounds"); + VectorDistanceBand band = VectorDistanceBand.fromRaw("l2", null, 4.0f); + check(band.rawLower() == null && band.rawUpper() == 4.0f, "structural bounds"); check("l2".equals(band.metric()), "metric"); + VectorDistanceBand rawL2 = VectorDistanceBand.fromRaw("l2", -0.0f, 16.0f); + check( + Float.floatToRawIntBits(rawL2.rawLower()) == 0x80000000 + && rawL2.rawUpper() == 16.0f, + "raw L2 cuts retain signed zero and squared distances"); + VectorDistanceBand rawIp = VectorDistanceBand.fromRaw("inner_product", -6.0f, -5.0f); + check( + rawIp.rawLower() == -6.0f && rawIp.rawUpper() == -5.0f, + "raw inner product cuts are not negated"); + VectorDistanceBand rawCosine = VectorDistanceBand.fromRaw("cosine", 0.25f, 0.75f); + check( + rawCosine.rawLower() == 0.25f && rawCosine.rawUpper() == 0.75f, + "raw cosine cuts are unchanged"); + VectorDistanceBand rawUnbounded = VectorDistanceBand.fromRaw("l2", null, null); + check( + rawUnbounded.rawLower() == null && rawUnbounded.rawUpper() == null, + "raw unbounded cuts remain structural"); check(new VectorRangeSearchParams(band, 2).nprobe() == 2, "nprobe"); expect(IllegalArgumentException.class, () -> new VectorRangeSearchParams(band, 0)); expect(NullPointerException.class, () -> new VectorRangeSearchParams(null, 1)); - expect(IllegalArgumentException.class, () -> new VectorDistanceBand("other", null, null)); - expect(IllegalArgumentException.class, () -> new VectorDistanceBand("l2", -1.0f, null)); - expect(IllegalArgumentException.class, () -> new VectorDistanceBand("cosine", 2.0f, 1.0f)); - expect(IllegalArgumentException.class, () -> new VectorDistanceBand("l2", null, Float.NaN)); + expect(NullPointerException.class, () -> VectorDistanceBand.fromRaw(null, null, null)); + expect(IllegalArgumentException.class, () -> VectorDistanceBand.fromRaw("other", null, null)); + expect(IllegalArgumentException.class, () -> VectorDistanceBand.fromRaw("l2", -1.0f, null)); + expect(IllegalArgumentException.class, () -> VectorDistanceBand.fromRaw("cosine", 2.0f, 1.0f)); + expect(IllegalArgumentException.class, () -> VectorDistanceBand.fromRaw("l2", null, Float.NaN)); + expect( + IllegalArgumentException.class, + () -> VectorDistanceBand.fromRaw("cosine", Float.NEGATIVE_INFINITY, null)); expect( IllegalArgumentException.class, - () -> new VectorDistanceBand("cosine", Float.NEGATIVE_INFINITY, null)); + () -> VectorDistanceBand.fromRaw("inner_product", null, Float.POSITIVE_INFINITY)); long[] labels = {LABEL_BASE, 7}; float[] distances = {1, 3}; @@ -78,15 +102,15 @@ static void testValueTypes() { check(result.queryCount() == 2 && result.listReads() == 2, "result shape"); check(result.labelsForQuery(0).length == 0, "empty first row"); check(result.labelsForQuery(1)[0] == LABEL_BASE, "64-bit label and copy"); - check(result.distancesForQuery(1)[0] == 1, "distance copy"); + check(result.rawDistancesForQuery(1)[0] == 1, "distance copy"); result.labels()[0] = -2; - result.distances()[0] = -2; + result.rawDistances()[0] = -2; result.lims()[1] = 2; result.listsProbed()[1] = -2; result.rowsScanned()[1] = -2; result.rowsCommitted()[1] = -2; result.earlyAbandoned()[1] = -2; - check(result.labels()[0] == LABEL_BASE && result.distances()[0] == 1, "defensive arrays"); + check(result.labels()[0] == LABEL_BASE && result.rawDistances()[0] == 1, "defensive arrays"); check(result.lims()[1] == 0 && result.listsProbed()[1] == 2, "defensive CSR and stats"); check( result.rowsScanned()[1] == 2 @@ -94,7 +118,7 @@ static void testValueTypes() { && result.earlyAbandoned()[1] == 0, "defensive counters"); expect(IndexOutOfBoundsException.class, () -> result.labelsForQuery(-1)); - expect(IndexOutOfBoundsException.class, () -> result.distancesForQuery(2)); + expect(IndexOutOfBoundsException.class, () -> result.rawDistancesForQuery(2)); expect( IllegalArgumentException.class, () -> @@ -148,6 +172,63 @@ static void testValueTypes() { expect(NullPointerException.class, () -> closed.rangeSearch(new float[1], params, null)); } + private static void testRawApiContract() { + for (Method method : VectorRangeSearchResult.class.getMethods()) { + check( + !"distances".equals(method.getName()) + && !"distanceAt".equals(method.getName()) + && !"distancesForQuery".equals(method.getName()), + "ambiguous public range result method: " + method.getName()); + } + check( + VectorDistanceBand.class.getConstructors().length == 0, + "raw distance band constructor must not be public"); + for (Method method : VectorDistanceBand.class.getMethods()) { + check( + !"lower".equals(method.getName()) && !"upper".equals(method.getName()), + "ambiguous public distance band method: " + method.getName()); + } + try { + check( + VectorRangeSearchResult.class.getMethod("rawDistances").getReturnType() + == float[].class, + "raw distance array accessor"); + check( + VectorRangeSearchResult.class.getMethod("rawDistanceAt", int.class) + .getReturnType() + == float.class, + "raw distance indexed accessor"); + check( + VectorRangeSearchResult.class.getMethod("rawDistancesForQuery", int.class) + .getReturnType() + == float[].class, + "raw distance query accessor"); + check( + Modifier.isPrivate( + VectorRangeSearchResult.class.getDeclaredField("rawDistances") + .getModifiers()), + "raw distance storage is private"); + Method factory = + VectorDistanceBand.class.getMethod( + "fromRaw", String.class, Float.class, Float.class); + check( + Modifier.isStatic(factory.getModifiers()) + && factory.getReturnType() == VectorDistanceBand.class, + "explicit raw band factory"); + for (String bound : new String[] {"rawLower", "rawUpper"}) { + check( + VectorDistanceBand.class.getMethod(bound).getReturnType() == Float.class, + "raw bound accessor: " + bound); + check( + Modifier.isPrivate( + VectorDistanceBand.class.getDeclaredField(bound).getModifiers()), + "raw bound storage is private: " + bound); + } + } catch (ReflectiveOperationException error) { + throw new AssertionError(error); + } + } + private static void testIndexedResultAccess() { long[] labels = {Long.MIN_VALUE, LABEL_BASE, Long.MAX_VALUE}; float[] distances = {-0.0f, 1.5f, Float.POSITIVE_INFINITY}; @@ -182,20 +263,20 @@ private static void testIndexedResultAccess() { hitIndex++) { check(result.labelAt(hitIndex) == labels[hitIndex], "indexed label"); check( - Float.floatToRawIntBits(result.distanceAt(hitIndex)) + Float.floatToRawIntBits(result.rawDistanceAt(hitIndex)) == Float.floatToRawIntBits(distances[hitIndex]), "indexed distance bits"); } } result.labels()[0] = 0; - result.distances()[0] = 1; + result.rawDistances()[0] = 1; result.lims()[1] = 1; result.listsProbed()[1] = 0; result.rowsScanned()[1] = 0; result.rowsCommitted()[1] = 0; result.earlyAbandoned()[1] = 0; check(result.labelAt(0) == Long.MIN_VALUE, "indexed labels remain immutable"); - check(Float.floatToRawIntBits(result.distanceAt(0)) == 0x80000000, "signed zero retained"); + check(Float.floatToRawIntBits(result.rawDistanceAt(0)) == 0x80000000, "signed zero retained"); check(result.queryEnd(0) == 0, "indexed offsets remain immutable"); check( result.listsProbed(1) == 4 && result.rowsScanned(1) == Long.MAX_VALUE, @@ -205,7 +286,7 @@ private static void testIndexedResultAccess() { "indexed counters remain immutable"); for (int hitIndex : new int[] {-1, result.hitCount(), Integer.MAX_VALUE}) { expect(IndexOutOfBoundsException.class, () -> result.labelAt(hitIndex)); - expect(IndexOutOfBoundsException.class, () -> result.distanceAt(hitIndex)); + expect(IndexOutOfBoundsException.class, () -> result.rawDistanceAt(hitIndex)); } for (int queryIndex : new int[] {-1, result.queryCount(), Integer.MAX_VALUE}) { checkInvalidQueryIndex(result, queryIndex); @@ -224,12 +305,13 @@ private static void testIndexedResultAccess() { "empty query bounds"); } expect(IndexOutOfBoundsException.class, () -> empty.labelAt(0)); - expect(IndexOutOfBoundsException.class, () -> empty.distanceAt(0)); + expect(IndexOutOfBoundsException.class, () -> empty.rawDistanceAt(0)); checkInvalidQueryIndex(empty, empty.queryCount()); } } private static void checkInvalidQueryIndex(VectorRangeSearchResult result, int queryIndex) { + expect(IndexOutOfBoundsException.class, () -> result.rawDistancesForQuery(queryIndex)); expect(IndexOutOfBoundsException.class, () -> result.queryStart(queryIndex)); expect(IndexOutOfBoundsException.class, () -> result.queryEnd(queryIndex)); expect(IndexOutOfBoundsException.class, () -> result.listsProbed(queryIndex)); @@ -258,7 +340,7 @@ private static void testMemoryBoundedConsumption() { double distanceSum = 0; for (int hitIndex = result.queryStart(1); hitIndex < result.queryEnd(1); hitIndex++) { labelSum += result.labelAt(hitIndex); - distanceSum += result.distanceAt(hitIndex); + distanceSum += result.rawDistanceAt(hitIndex); } check(labelSum == 2 * LABEL_BASE + 1 && distanceSum == -2.25, "large result consumption"); check(result.listsProbed(1) == 1 && result.rowsScanned(1) == hitCount, "large result stats"); @@ -286,25 +368,25 @@ private static void testNativeResultOwnership() { earlyAbandoned, 1); checkOwnedArray(result, "labels", labels); - checkOwnedArray(result, "distances", distances); + checkOwnedArray(result, "rawDistances", distances); checkOwnedArray(result, "lims", lims); checkOwnedArray(result, "listsProbed", listsProbed); checkOwnedArray(result, "rowsScanned", rowsScanned); checkOwnedArray(result, "rowsCommitted", rowsCommitted); checkOwnedArray(result, "earlyAbandoned", earlyAbandoned); result.labels()[0] = -1; - result.distances()[0] = -1; + result.rawDistances()[0] = -1; result.lims()[1] = 2; result.listsProbed()[1] = -1; result.rowsScanned()[1] = -1; result.rowsCommitted()[1] = -1; result.earlyAbandoned()[1] = -1; result.labelsForQuery(1)[0] = -1; - result.distancesForQuery(1)[0] = -1; + result.rawDistancesForQuery(1)[0] = -1; check(result.queryCount() == 2 && result.listReads() == 1, "owned result shape"); check(result.labelsForQuery(0).length == 0, "owned empty query"); check(result.labelsForQuery(1)[0] == LABEL_BASE, "owned labels remain defensive"); - check(result.distancesForQuery(1)[0] == 1, "owned distances remain defensive"); + check(result.rawDistancesForQuery(1)[0] == 1, "owned distances remain defensive"); check(result.lims()[1] == 0 && result.listsProbed()[1] == 1, "owned CSR and stats"); check( result.rowsScanned()[1] == 3 @@ -361,6 +443,7 @@ static void testNative() { } } testExactDistancesAndEndpoints(); + testEndpointSearchReturnsRawDistances(); testNativeValidation(); testBatchMetadataTransfer(); testCallbacks(); @@ -419,8 +502,8 @@ private static void testBatchMetadataTransfer() { batch.labelAt(expectedStart + hitIndex) == single.labelAt(hitIndex), "retained batch label"); check( - Float.floatToRawIntBits(batch.distanceAt(expectedStart + hitIndex)) - == Float.floatToRawIntBits(single.distanceAt(hitIndex)), + Float.floatToRawIntBits(batch.rawDistanceAt(expectedStart + hitIndex)) + == Float.floatToRawIntBits(single.rawDistanceAt(hitIndex)), "retained batch distance bits"); } expectedStart += single.hitCount(); @@ -450,7 +533,7 @@ private static void testIndex(String indexType, String metric) { check(full.listsProbed()[queryIndex] == 2, "lists probed"); check(full.rowsScanned()[queryIndex] == VECTOR_COUNT, "rows scanned"); } - float[] sorted = full.distancesForQuery(0); + float[] sorted = full.rawDistancesForQuery(0); Arrays.sort(sorted); Float lower = sorted[VECTOR_COUNT / 4]; Float upper = sorted[VECTOR_COUNT * 3 / 4]; @@ -562,17 +645,17 @@ query, new VectorRangeSearchParams(band, 2)), VectorDistanceBand unbounded = VectorDistanceBand.fromEndpoints("inner_product", null, null, null, null); check( - unbounded.lower() == null && unbounded.upper() == null, + unbounded.rawLower() == null && unbounded.rawUpper() == null, "unbounded endpoint conversion"); VectorDistanceBand outside = VectorDistanceBand.fromEndpoints( "cosine", -Double.MAX_VALUE, GE, Double.MAX_VALUE, LE); check( - outside.lower() == -Float.MAX_VALUE && outside.upper() == null, + outside.rawLower() == -Float.MAX_VALUE && outside.rawUpper() == null, "linear out-of-domain endpoints retain core cuts and unbounded upper"); VectorDistanceBand equal = VectorDistanceBand.fromEndpoints("l2", 1.0, GE, 1.0, LE); check( - equal.lower() <= 1.0f && equal.upper() > 1.0f, + equal.rawLower() <= 1.0f && equal.rawUpper() > 1.0f, "inclusive equal endpoints retain equality bucket"); expect( RuntimeException.class, @@ -591,6 +674,96 @@ query, new VectorRangeSearchParams(band, 2)), () -> VectorIndexNative.distanceBandFromEndpoints("l2", 1.0, 4, null, -1)); } + private static void testEndpointSearchReturnsRawDistances() { + assertEndpointSearchReturnsRawDistances( + "l2", + new float[] {3, 0, 4, 0, 0, 3, 5, 0}, + new float[] {0, 0}, + VectorDistanceBand.fromEndpoints("l2", null, null, 4.0, LT), + new long[] {LABEL_BASE, LABEL_BASE + 2}, + new float[] {9.0f, 9.0f}); + assertEndpointSearchReturnsRawDistances( + "inner_product", + new float[] {3, 0, 2, 0, 4, 0, 2.5f, 0}, + new float[] {2, 0}, + VectorDistanceBand.fromEndpoints("inner_product", 5.0, GE, null, null), + new long[] {LABEL_BASE, LABEL_BASE + 2, LABEL_BASE + 3}, + new float[] {-6.0f, -8.0f, -5.0f}); + assertEndpointSearchReturnsRawDistances( + "cosine", + new float[] {1, 0, 0, 1, 1, 1, -1, 0}, + new float[] {1, 0}, + VectorDistanceBand.fromEndpoints("cosine", null, null, 1.0, LT), + new long[] {LABEL_BASE, LABEL_BASE + 2}, + new float[] {0.0f, (float) (1.0 - 1.0 / Math.sqrt(2.0))}); + } + + private static void assertEndpointSearchReturnsRawDistances( + String metric, + float[] data, + float[] query, + VectorDistanceBand endpointBand, + long[] expectedLabels, + float[] expectedRawDistances) { + VectorDistanceBand rawBand = + VectorDistanceBand.fromRaw( + metric, endpointBand.rawLower(), endpointBand.rawUpper()); + float[] queries = new float[query.length * 2]; + System.arraycopy(query, 0, queries, 0, query.length); + System.arraycopy(query, 0, queries, query.length, query.length); + try (VectorIndexReader reader = open(build("ivf_flat", metric, 2, data))) { + for (VectorDistanceBand band : new VectorDistanceBand[] {endpointBand, rawBand}) { + VectorRangeSearchParams params = new VectorRangeSearchParams(band, 2); + for (boolean filtered : new boolean[] {false, true}) { + Map expected = new HashMap(); + for (int hitIndex = 0; hitIndex < expectedLabels.length; hitIndex++) { + long label = expectedLabels[hitIndex]; + if (!filtered || label == LABEL_BASE || label == LABEL_BASE + 1) { + expected.put(label, expectedRawDistances[hitIndex]); + } + } + VectorRangeSearchResult single = + filtered + ? reader.rangeSearch(query, params, filter()) + : reader.rangeSearch(query, params); + VectorRangeSearchResult batch = + filtered + ? reader.rangeSearchBatch(queries, 2, params, filter()) + : reader.rangeSearchBatch(queries, 2, params); + assertShape(single, 1); + assertShape(batch, 2); + assertRawValues(single, 0, expected); + for (int queryIndex = 0; queryIndex < batch.queryCount(); queryIndex++) { + assertRows(single, 0, batch, queryIndex); + assertRawValues(batch, queryIndex, expected); + } + } + } + } + } + + private static void assertRawValues( + VectorRangeSearchResult result, int queryIndex, Map expected) { + check(rows(result, queryIndex).keySet().equals(expected.keySet()), "endpoint membership"); + float[] rawDistances = result.rawDistances(); + float[] queryRawDistances = result.rawDistancesForQuery(queryIndex); + for (int hitIndex = result.queryStart(queryIndex); + hitIndex < result.queryEnd(queryIndex); + hitIndex++) { + float rawDistance = result.rawDistanceAt(hitIndex); + check( + Math.abs(rawDistance - expected.get(result.labelAt(hitIndex))) <= 1e-6f, + "endpoint search returns raw distance, not endpoint value"); + check( + Float.floatToRawIntBits(rawDistance) + == Float.floatToRawIntBits(rawDistances[hitIndex]) + && Float.floatToRawIntBits(rawDistance) + == Float.floatToRawIntBits( + queryRawDistances[hitIndex - result.queryStart(queryIndex)]), + "raw accessors retain identical distance bits"); + } + } + private static void testNativeValidation() { VectorRangeSearchParams all = params("l2", null, null); expect(RuntimeException.class, () -> VectorIndexNative.supportsRangeSearch(0)); @@ -716,7 +889,7 @@ private static void testUnsupported() { } private static VectorRangeSearchParams params(String metric, Float lower, Float upper) { - return new VectorRangeSearchParams(new VectorDistanceBand(metric, lower, upper), 2); + return new VectorRangeSearchParams(VectorDistanceBand.fromRaw(metric, lower, upper), 2); } private static VectorIndexReader open(byte[] bytes) { @@ -774,7 +947,7 @@ private static void assertShape(VectorRangeSearchResult result, int count) { check( result.lims()[0] == 0 && result.lims()[count] == result.labels().length, "CSR limits"); - check(result.labels().length == result.distances().length, "parallel results"); + check(result.labels().length == result.rawDistances().length, "parallel results"); for (int queryIndex = 0; queryIndex < count; queryIndex++) { check( result.rowsCommitted()[queryIndex] == result.labelsForQuery(queryIndex).length, @@ -791,7 +964,7 @@ private static void assertShape(VectorRangeSearchResult result, int count) { private static Map rows(VectorRangeSearchResult result, int queryIndex) { Map rows = new HashMap(); long[] labels = result.labelsForQuery(queryIndex); - float[] distances = result.distancesForQuery(queryIndex); + float[] distances = result.rawDistancesForQuery(queryIndex); for (int row = 0; row < labels.length; row++) { check(Float.isFinite(distances[row]), "finite raw distance"); check(rows.put(labels[row], distances[row]) == null, "unique labels"); diff --git a/jni/src/range.rs b/jni/src/range.rs index 7c834305..9b01708a 100644 --- a/jni/src/range.rs +++ b/jni/src/range.rs @@ -140,9 +140,10 @@ fn range_params(env: &mut JNIEnv, params: JObject) -> Result Result Result RangeSearchQueryResult: + """Return unpackable labels/raw_distances views, without payload copies. + + The mutable views share the result arrays and keep them alive even + after this result is released. The query index must be nonnegative. + """ index = operator.index(index) if not 0 <= index < self.query_count: raise IndexError("range query index out of bounds") start, end = int(self.lims[index]), int(self.lims[index + 1]) - return self.labels[start:end], self.distances[start:end] + return RangeSearchQueryResult( + self.labels[start:end], self.raw_distances[start:end] + ) @dataclass(frozen=True) @@ -494,8 +530,8 @@ def _range_result_copy(handle, query_count): ): raise RuntimeError("range result has invalid CSR limits") labels = _range_array_copy(view.labels, view.hit_count, np.int64, "labels") - distances = _range_array_copy( - view.distances, view.hit_count, np.float32, "distances" + raw_distances = _range_array_copy( + view.raw_distances, view.hit_count, np.float32, "raw_distances" ) if query_count and not view.stats: raise RuntimeError("range result stats pointer is null") @@ -508,7 +544,7 @@ def _range_result_copy(handle, query_count): ) for index in range(query_count) ) - return RangeSearchResult(lims, labels, distances, stats, view.list_reads) + return RangeSearchResult(lims, labels, raw_distances, stats, view.list_reads) def _int64_vector(value, name): @@ -1014,7 +1050,7 @@ def supports_range_search(self): return bool(supported.value) def range_search(self, query, params: RangeSearchParams, roaring_filter=None): - """Search one query, returning an owned one-query CSR result. + """Search one query, returning owned CSR arrays with raw_distances. roaring_filter is optional serialized RoaringTreemap bytes. An empty byte string is still a supplied filter and is validated by core. @@ -1024,7 +1060,7 @@ def range_search(self, query, params: RangeSearchParams, roaring_filter=None): def range_search_batch( self, queries, params: RangeSearchParams, roaring_filter=None ): - """Search a query matrix; core rejects a zero-query batch.""" + """Return raw_distances for a query matrix; core rejects empty batches.""" return self._range_search(queries, params, roaring_filter, batch=True) def _range_search(self, value, params, roaring_filter, *, batch): @@ -1180,6 +1216,7 @@ def __del__(self): "DistanceEndpointOp", "IvfPqBatchTableReuseMode", "RangeSearchParams", + "RangeSearchQueryResult", "RangeSearchResult", "RangeSearchStats", "SearchParams", diff --git a/python/paimon_vindex/_ffi.py b/python/paimon_vindex/_ffi.py index 3ec3626c..83cd0d31 100644 --- a/python/paimon_vindex/_ffi.py +++ b/python/paimon_vindex/_ffi.py @@ -157,13 +157,13 @@ class PaimonVindexSearchParamsV2(Structure): ] -class PaimonVindexDistanceBand(Structure): +class PaimonVindexRawDistanceBand(Structure): _fields_ = [ ("metric", c_uint32), - ("lower_kind", c_uint32), - ("lower", c_float), - ("upper_kind", c_uint32), - ("upper", c_float), + ("raw_lower_kind", c_uint32), + ("raw_lower", c_float), + ("raw_upper_kind", c_uint32), + ("raw_upper", c_float), ] @@ -176,7 +176,7 @@ class PaimonVindexDistanceEndpoint(Structure): class PaimonVindexRangeSearchParams(Structure): _fields_ = [ - ("band", PaimonVindexDistanceBand), + ("band", PaimonVindexRawDistanceBand), ("nprobe", c_size_t), ] @@ -196,7 +196,7 @@ class PaimonVindexRangeSearchResultView(Structure): ("hit_count", c_size_t), ("lims", POINTER(c_size_t)), ("labels", POINTER(c_int64)), - ("distances", POINTER(c_float)), + ("raw_distances", POINTER(c_float)), ("stats", POINTER(PaimonVindexRangeSearchStats)), ("list_reads", c_size_t), ] @@ -391,7 +391,7 @@ def _configure_range_api(): c_uint32, POINTER(PaimonVindexDistanceEndpoint), POINTER(PaimonVindexDistanceEndpoint), - POINTER(PaimonVindexDistanceBand), + POINTER(PaimonVindexRawDistanceBand), ], c_int), ("paimon_vindex_reader_supports_range_search", [ c_void_p, POINTER(c_int), diff --git a/python/tests/test_range_search.py b/python/tests/test_range_search.py index 023852c6..22a6ecfd 100644 --- a/python/tests/test_range_search.py +++ b/python/tests/test_range_search.py @@ -25,6 +25,7 @@ import threading import textwrap from collections import Counter +from dataclasses import FrozenInstanceError from pathlib import Path import numpy as np @@ -48,6 +49,7 @@ def test_range_public_api_is_exported_without_loading_native_library(): "DistanceEndpoint", "DistanceEndpointOp", "RangeSearchParams", + "RangeSearchQueryResult", "RangeSearchResult", "RangeSearchStats", } <= set(exports) @@ -121,6 +123,144 @@ def flat_index(vindex): return make_index(vindex) +def test_raw_distance_api_names(vindex): + assert set(vindex.RangeSearchResult.__dataclass_fields__) == { + "lims", "labels", "raw_distances", "stats", "list_reads", + } + assert set(vindex.DistanceBand.__dataclass_fields__) == { + "metric", "raw_lower", "raw_upper", + } + ffi = vindex._ffi + assert not hasattr(ffi, "PaimonVindexDistanceBand") + assert ffi.PaimonVindexRawDistanceBand._fields_ == [ + ("metric", ctypes.c_uint32), + ("raw_lower_kind", ctypes.c_uint32), + ("raw_lower", ctypes.c_float), + ("raw_upper_kind", ctypes.c_uint32), + ("raw_upper", ctypes.c_float), + ] + assert ctypes.sizeof(ffi.PaimonVindexRawDistanceBand) == 20 + assert [ + getattr(ffi.PaimonVindexRawDistanceBand, name).offset + for name, _ in ffi.PaimonVindexRawDistanceBand._fields_ + ] == [0, 4, 8, 12, 16] + assert ffi.PaimonVindexRangeSearchParams._fields_[0] == ( + "band", ffi.PaimonVindexRawDistanceBand + ) + assert [name for name, _ in ffi.PaimonVindexRangeSearchResultView._fields_] == [ + "query_count", "hit_count", "lims", "labels", "raw_distances", "stats", + "list_reads", + ] + + +@pytest.mark.parametrize("args,kwargs", [ + (("l2",), {}), + (("l2", 0.0, 4.0), {}), + ((), {"metric": "l2", "lower": 0.0, "upper": 4.0}), + ((), {"metric": "l2", "raw_lower": 0.0, "raw_upper": 4.0}), +]) +def test_ambiguous_band_construction_is_rejected(vindex, args, kwargs): + with pytest.raises(TypeError, match="from_endpoints.*from_raw"): + vindex.DistanceBand(*args, **kwargs) + + +def test_explicit_raw_band_factory(vindex): + band = vindex.DistanceBand.from_raw("inner_product", raw_lower=-6, raw_upper="2") + assert band.metric == "inner_product" + assert band.raw_lower == -6.0 and type(band.raw_lower) is float + assert band.raw_upper == 2.0 and type(band.raw_upper) is float + assert not hasattr(band, "lower") and not hasattr(band, "upper") + assert band == vindex.DistanceBand.from_raw("inner_product", -6.0, 2.0) + assert band.to_ffi().raw_lower == -6.0 + assert band.to_ffi().raw_upper == 2.0 + unbounded = vindex.DistanceBand.from_raw("l2") + assert unbounded.raw_lower is None and unbounded.raw_upper is None + assert unbounded.to_ffi().raw_lower_kind == 0 + assert unbounded.to_ffi().raw_upper_kind == 0 + with pytest.raises(FrozenInstanceError): + band.raw_lower = 0.0 + with pytest.raises(TypeError): + vindex.DistanceBand.from_raw("l2", lower=0.0, upper=4.0) + with pytest.raises(ValueError): + vindex.DistanceBand.from_raw("l2", raw_lower="invalid") + + +@pytest.fixture(scope="module", params=["l2", "inner_product", "cosine"]) +def endpoint_index(request, vindex): + metric = request.param + data = np.zeros((128, 16), dtype=np.float32) + query = np.zeros(16, dtype=np.float32) + if metric == "l2": + data[:, 0] = 5.0 + data[0, 0], data[1, 0] = 3.0, 4.0 + endpoints = {"upper": vindex.DistanceEndpoint(4.0, vindex.DistanceEndpointOp.LE)} + expected = [9.0, 16.0] + elif metric == "inner_product": + query[0] = 2.0 + data[:, 0] = 2.0 + data[0, 0], data[1, 0] = 3.0, 2.5 + endpoints = {"lower": vindex.DistanceEndpoint(5.0, vindex.DistanceEndpointOp.GE)} + expected = [-6.0, -5.0] + else: + query[0] = 1.0 + data[:, 1] = 1.0 + data[0, 0] = 1.0 + data[1] = query + endpoints = {"upper": vindex.DistanceEndpoint(0.5, vindex.DistanceEndpointOp.LE)} + expected = [1.0 - 1.0 / np.sqrt(2.0), 0.0] + labels = np.arange(len(data), dtype=np.int64) + (1 << 40) + options = { + "index.type": "ivf_flat", "dimension": "16", "metric": metric, + "nlist": "1", + } + output = io.BytesIO() + training = vindex.VectorIndexTrainer.train(options, data) + with vindex.VectorIndexWriter(training) as writer: + writer.add_vectors(labels, data) + writer.write(output) + return metric, output.getvalue(), query, labels, endpoints, expected + + +@pytest.mark.parametrize("batch", [False, True]) +@pytest.mark.parametrize("filter_kind", ["none", "subset", "empty"]) +def test_endpoint_search_returns_raw_distances(vindex, endpoint_index, batch, filter_kind): + metric, payload, query, labels, endpoints, expected = endpoint_index + band = vindex.DistanceBand.from_endpoints(metric, **endpoints) + params = vindex.RangeSearchParams(band, 1) + allowed = {"none": labels, "subset": labels[::3], "empty": labels[:0]}[filter_kind] + filter_bytes = None if filter_kind == "none" else roaring_allowlist(allowed) + with vindex.VectorIndexReader(BytesInput(payload)) as reader: + method = reader.range_search_batch if batch else reader.range_search + result = method( + np.stack([query, query]) if batch else query, params, + roaring_filter=filter_bytes, + ) + assert not hasattr(result, "distances") + assert result.raw_distances.dtype == np.dtype(np.float32) + assert result.raw_distances.flags.owndata and result.raw_distances.flags.writeable + assert result.query_count == (2 if batch else 1) + expected_hits = { + label: raw_distance + for label, raw_distance in zip(labels[:2], expected) + if label in allowed + } + for query_index in range(result.query_count): + hits = result.query(query_index) + assert isinstance(hits, tuple) + assert hits._fields == ("labels", "raw_distances") + assert not hasattr(hits, "distances") + hit_labels, raw_distances = hits + assert hit_labels is hits.labels and raw_distances is hits.raw_distances + assert hit_labels.base is result.labels + assert raw_distances.base is result.raw_distances + assert set(hit_labels) == set(expected_hits) + for label, raw_distance in zip(hit_labels, raw_distances): + assert raw_distance == pytest.approx(expected_hits[label], abs=1e-6) + if expected_hits: + assert np.shares_memory(hit_labels, result.labels) + assert np.shares_memory(raw_distances, result.raw_distances) + + @pytest.mark.parametrize("batch", [False, True]) def test_range_accepts_unaligned_contiguous_queries(vindex, flat_index, batch): payload, data, _ = flat_index @@ -131,7 +271,7 @@ def test_range_accepts_unaligned_contiguous_queries(vindex, flat_index, batch): ) queries[:] = source assert queries.flags.c_contiguous and not queries.flags.aligned - params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + params = vindex.RangeSearchParams(vindex.DistanceBand.from_raw("l2"), 4) with vindex.VectorIndexReader(BytesInput(payload)) as reader: method = reader.range_search_batch if batch else reader.range_search result = method(queries, params) @@ -180,7 +320,7 @@ def __getattr__(self, name): expected = reader.search(data[0], params) assert expected[0][0] == labels[0] assert not reader.supports_range_search() - range_params = api.RangeSearchParams(api.DistanceBand("l2"), 4) + range_params = api.RangeSearchParams(api.DistanceBand.from_raw("l2"), 4) for call in ( lambda: reader.range_search(data[0], range_params), lambda: reader.range_search_batch(data[:3], range_params), @@ -208,16 +348,16 @@ def test_range_matrix(vindex, index_case, batch, filter_kind, band_kind): index_type, metric, (payload, data, labels) = index_case queries = data[:3] if batch else data[:1] if band_kind == "unbounded": - band = vindex.DistanceBand(metric) + band = vindex.DistanceBand.from_raw(metric) elif band_kind == "empty": - band = vindex.DistanceBand(metric, 0.0, 0.0) + band = vindex.DistanceBand.from_raw(metric, 0.0, 0.0) else: lower, upper = { "l2": (0.0, 24.0), "inner_product": (-4.0, 2.0), "cosine": (0.0, 0.95), }[metric] - band = vindex.DistanceBand(metric, lower, upper) + band = vindex.DistanceBand.from_raw(metric, lower, upper) allowed = { "none": labels, "subset": labels[::3], @@ -232,15 +372,15 @@ def test_range_matrix(vindex, index_case, batch, filter_kind, band_kind): queries if batch else queries[0], params, roaring_filter=filter_bytes ) assert result.query_count == len(queries) - assert result.hit_count == len(result.labels) == len(result.distances) + assert result.hit_count == len(result.labels) == len(result.raw_distances) assert result.lims.dtype == np.dtype(np.uintp) assert result.labels.dtype == np.dtype(np.int64) - assert result.distances.dtype == np.dtype(np.float32) + assert result.raw_distances.dtype == np.dtype(np.float32) assert result.lims[0] == 0 and result.lims[-1] == result.hit_count assert np.all(result.lims[1:] >= result.lims[:-1]) assert len(result.stats) == len(queries) assert 0 <= result.list_reads <= 4 - for array in (result.lims, result.labels, result.distances): + for array in (result.lims, result.labels, result.raw_distances): assert array.flags.owndata for query_index, query in enumerate(queries): with vindex.VectorIndexReader(BytesInput(payload)) as reader: @@ -253,7 +393,7 @@ def test_range_matrix(vindex, index_case, batch, filter_kind, band_kind): ) np.testing.assert_array_equal( actual_distances[actual_order].view(np.uint32), - reference.distances[reference_order].view(np.uint32), + reference.raw_distances[reference_order].view(np.uint32), ) assert set(actual_labels) <= set(allowed) assert result.stats[query_index] == reference.stats[0] @@ -267,17 +407,17 @@ def test_range_matrix(vindex, index_case, batch, filter_kind, band_kind): assert set(actual_labels) == set(allowed) if band_kind == "empty" or filter_kind == "empty": assert len(actual_labels) == 0 - if band.lower is not None: - assert np.all(actual_distances >= np.float32(band.lower)) - if band.upper is not None: - assert np.all(actual_distances < np.float32(band.upper)) + if band.raw_lower is not None: + assert np.all(actual_distances >= np.float32(band.raw_lower)) + if band.raw_upper is not None: + assert np.all(actual_distances < np.float32(band.raw_upper)) @pytest.mark.parametrize("batch", [False, True]) @pytest.mark.parametrize("filter_bytes", [b"", b"invalid roaring"]) def test_invalid_filter(vindex, flat_index, batch, filter_bytes): payload, data, _ = flat_index - params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + params = vindex.RangeSearchParams(vindex.DistanceBand.from_raw("l2"), 4) with vindex.VectorIndexReader(BytesInput(payload)) as reader: method = reader.range_search_batch if batch else reader.range_search with pytest.raises(RuntimeError, match="[Rr]oaring|filter"): @@ -291,7 +431,7 @@ def test_invalid_filter(vindex, flat_index, batch, filter_bytes): ]) def test_invalid_bands_are_rejected_by_core(vindex, flat_index, lower, upper): payload, data, _ = flat_index - params = vindex.RangeSearchParams(vindex.DistanceBand("l2", lower, upper), 4) + params = vindex.RangeSearchParams(vindex.DistanceBand.from_raw("l2", lower, upper), 4) with vindex.VectorIndexReader(BytesInput(payload)) as reader: with pytest.raises(RuntimeError): reader.range_search(data[0], params) @@ -302,7 +442,7 @@ def test_invalid_bands_are_rejected_by_core(vindex, flat_index, lower, upper): @pytest.mark.parametrize("nprobe", [-1, 0, 1.5, ctypes.c_size_t(-1).value + 1]) def test_nprobe_rejects_invalid_and_wrapping_values(vindex, nprobe): with pytest.raises(ValueError, match="nprobe"): - vindex.RangeSearchParams(vindex.DistanceBand("l2"), nprobe) + vindex.RangeSearchParams(vindex.DistanceBand.from_raw("l2"), nprobe) @pytest.mark.parametrize("metric", ["l2", "inner_product", "cosine"]) @@ -315,7 +455,7 @@ def test_endpoints_match_direct_core_conversion(vindex, metric, lower_op, upper_ band = vindex.DistanceBand.from_endpoints(metric, lower, upper) raw_lower = ffi.PaimonVindexDistanceEndpoint(lower.value, lower_op) raw_upper = ffi.PaimonVindexDistanceEndpoint(upper.value, upper_op) - expected = ffi.PaimonVindexDistanceBand() + expected = ffi.PaimonVindexRawDistanceBand() assert ffi.lib.paimon_vindex_distance_band_from_endpoints( {"l2": 0, "inner_product": 1, "cosine": 2}[metric], ctypes.byref(raw_lower), ctypes.byref(raw_upper), ctypes.byref(expected), @@ -327,7 +467,7 @@ def test_endpoints_match_direct_core_conversion(vindex, metric, lower_op, upper_ def test_empty_batch_and_query_accessor(vindex, flat_index): payload, data, _ = flat_index with vindex.VectorIndexReader(BytesInput(payload)) as reader: - params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + params = vindex.RangeSearchParams(vindex.DistanceBand.from_raw("l2"), 4) with pytest.raises(RuntimeError, match="query count must be greater than 0"): reader.range_search_batch(data[:0], params) result = reader.range_search(data[0], params) @@ -342,7 +482,7 @@ def test_empty_batch_and_query_accessor(vindex, flat_index): @pytest.mark.parametrize("filter_kind", ["none", "subset", "empty"]) def test_query_access_shares_owned_payload(vindex, flat_index, batch, filter_kind): payload, data, labels = flat_index - params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + params = vindex.RangeSearchParams(vindex.DistanceBand.from_raw("l2"), 4) filter_bytes = { "none": None, "subset": roaring_allowlist(labels[::3]), @@ -356,10 +496,13 @@ def test_query_access_shares_owned_payload(vindex, flat_index, batch, filter_kin retained_views = [] for query_index in range(result.query_count): start, end = int(result.lims[query_index]), int(result.lims[query_index + 1]) + hits = result.query(query_index) + repeated_hits = result.query(query_index) + assert hits._fields == ("labels", "raw_distances") for owned, view, repeated in zip( - (result.labels, result.distances), - result.query(query_index), - result.query(query_index), + (result.labels, result.raw_distances), + hits, + repeated_hits, ): assert view.base is owned and repeated.base is owned assert view.flags.writeable and repeated.flags.writeable @@ -371,7 +514,7 @@ def test_query_access_shares_owned_payload(vindex, flat_index, batch, filter_kin owned[end - 1] = -2 assert view[-1] == repeated[-1] == -2 retained_views.append(view) - del result, owned, view, repeated + del result, hits, repeated_hits, owned, view, repeated for view in retained_views: if len(view): assert view[-1] == -2 @@ -385,7 +528,7 @@ def test_query_access_shares_owned_payload(vindex, flat_index, batch, filter_kin ]) def test_query_shapes(vindex, flat_index, batch, shape, error): payload, _, _ = flat_index - params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + params = vindex.RangeSearchParams(vindex.DistanceBand.from_raw("l2"), 4) with vindex.VectorIndexReader(BytesInput(payload)) as reader: method = reader.range_search_batch if batch else reader.range_search with pytest.raises(error): @@ -395,7 +538,7 @@ def test_query_shapes(vindex, flat_index, batch, shape, error): def test_strided_queries_and_filter_buffer_types(vindex, flat_index): payload, data, labels = flat_index queries = np.asfortranarray(data[:4].astype(np.float64))[::2] - params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + params = vindex.RangeSearchParams(vindex.DistanceBand.from_raw("l2"), 4) serialized = roaring_allowlist(labels[::3]) for filter_bytes in (bytearray(serialized), memoryview(serialized)): with vindex.VectorIndexReader(BytesInput(payload)) as reader: @@ -409,7 +552,7 @@ def test_strided_queries_and_filter_buffer_types(vindex, flat_index): def test_query_buffer_overflow_before_copy(vindex, flat_index): payload, _, _ = flat_index - params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + params = vindex.RangeSearchParams(vindex.DistanceBand.from_raw("l2"), 4) huge = np.lib.stride_tricks.as_strided( np.zeros(1, dtype=np.uint8), shape=(np.iinfo(np.intp).max // 32, 16), strides=(0, 0), @@ -423,7 +566,7 @@ def test_query_buffer_overflow_before_copy(vindex, flat_index): def test_closed_reader_and_reentry(vindex, flat_index, operation): payload, data, _ = flat_index reader = vindex.VectorIndexReader(BytesInput(payload)) - params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + params = vindex.RangeSearchParams(vindex.DistanceBand.from_raw("l2"), 4) invoke = { "single": lambda: reader.range_search(data[0], params), "batch": lambda: reader.range_search_batch(data[:2], params), @@ -441,7 +584,7 @@ def test_diskann_is_unsupported(vindex): payload, data, _ = make_index(vindex, "diskann") with vindex.VectorIndexReader(BytesInput(payload)) as reader: assert reader.supports_range_search() is False - params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + params = vindex.RangeSearchParams(vindex.DistanceBand.from_raw("l2"), 4) for query in (data[0], data[:2]): method = ( reader.range_search if query.ndim == 1 else reader.range_search_batch @@ -455,15 +598,15 @@ def test_diskann_is_unsupported(vindex): @pytest.mark.parametrize("batch", [False, True]) @pytest.mark.parametrize("filtered", [False, True]) @pytest.mark.parametrize("failure", [ - "view", "copy_lims", "copy_labels", "copy_distances", "stats", "result", - "search", "null", "null_distances", + "view", "copy_lims", "copy_labels", "copy_raw_distances", "stats", "result", + "search", "null", "null_raw_distances", "null_stats", "null_lims", "length", "overflow", "lims", ]) def test_result_destroyed_on_failure( vindex, flat_index, monkeypatch, failure, batch, filtered ): payload, data, labels = flat_index - params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + params = vindex.RangeSearchParams(vindex.DistanceBand.from_raw("l2"), 4) ffi = vindex._ffi destroyed = [] destroy = ffi.lib.paimon_vindex_range_search_result_destroy @@ -476,7 +619,7 @@ def test_result_destroyed_on_failure( search = getattr(ffi.lib, search_name) as_array = np.ctypeslib.as_array array_count = 0 - copy_failure_at = {"copy_lims": 1, "copy_labels": 2, "copy_distances": 3}.get( + copy_failure_at = {"copy_lims": 1, "copy_labels": 2, "copy_raw_distances": 3}.get( failure ) @@ -493,8 +636,8 @@ def alter_view(handle, output): return -1 if failure == "null": raw.labels = ctypes.POINTER(ctypes.c_int64)() - if failure == "null_distances": - raw.distances = ctypes.POINTER(ctypes.c_float)() + if failure == "null_raw_distances": + raw.raw_distances = ctypes.POINTER(ctypes.c_float)() if failure == "null_stats": raw.stats = ctypes.POINTER(ffi.PaimonVindexRangeSearchStats)() if failure == "null_lims": @@ -547,7 +690,7 @@ def fail_array(*args, **kwargs): def test_result_copies_survive_native_destruction(vindex, flat_index, monkeypatch): payload, data, _ = flat_index - params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + params = vindex.RangeSearchParams(vindex.DistanceBand.from_raw("l2"), 4) ffi = vindex._ffi destroy = ffi.lib.paimon_vindex_range_search_result_destroy destroyed = [] @@ -562,7 +705,7 @@ def poison_destroy(handle): ) ctypes.memset(view.labels, 255, view.hit_count * ctypes.sizeof(ctypes.c_int64)) ctypes.memset( - view.distances, 255, view.hit_count * ctypes.sizeof(ctypes.c_float) + view.raw_distances, 255, view.hit_count * ctypes.sizeof(ctypes.c_float) ) ctypes.memset( view.stats, 255, @@ -579,7 +722,7 @@ def poison_destroy(handle): assert len(destroyed) == 1 np.testing.assert_array_equal(result.lims, [0, 128, 256]) assert np.all(result.labels >= (1 << 40)) - assert np.all(np.isfinite(result.distances)) + assert np.all(np.isfinite(result.raw_distances)) assert [stats.rows_committed for stats in result.stats] == [128, 128] @@ -610,7 +753,7 @@ def test_endpoint_operator_cannot_wrap(vindex, op): def test_nullable_and_extreme_endpoints(vindex): for metric in ("l2", "inner_product", "cosine"): - assert vindex.DistanceBand.from_endpoints(metric) == vindex.DistanceBand(metric) + assert vindex.DistanceBand.from_endpoints(metric) == vindex.DistanceBand.from_raw(metric) assert vindex.DistanceBand.from_endpoints( metric, lower=vindex.DistanceEndpoint(0.1, vindex.DistanceEndpointOp.GE) ).metric == metric @@ -628,7 +771,7 @@ def test_nullable_and_extreme_endpoints(vindex): def test_metric_and_parameter_validation(vindex, flat_index): for metric in ("unknown", -1, 1 << 32, None): with pytest.raises(ValueError, match="metric"): - vindex.DistanceBand(metric) + vindex.DistanceBand.from_raw(metric) with pytest.raises(TypeError, match="band"): vindex.RangeSearchParams("l2", 4) payload, data, _ = flat_index @@ -637,18 +780,18 @@ def test_metric_and_parameter_validation(vindex, flat_index): reader.range_search(data[0], vindex.SearchParams.ivf(4, 4)) with pytest.raises(RuntimeError, match="metric"): reader.range_search( - data[0], vindex.RangeSearchParams(vindex.DistanceBand("cosine"), 4) + data[0], vindex.RangeSearchParams(vindex.DistanceBand.from_raw("cosine"), 4) ) with pytest.raises(ValueError, match="bytes"): reader.range_search( - data[0], vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4), + data[0], vindex.RangeSearchParams(vindex.DistanceBand.from_raw("l2"), 4), roaring_filter="not bytes", ) def test_range_callback_reentry(vindex, flat_index): payload, data, _ = flat_index - params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + params = vindex.RangeSearchParams(vindex.DistanceBand.from_raw("l2"), 4) class ReentrantInput(BytesInput): operation = None @@ -681,7 +824,7 @@ def reenter(): def test_range_close_waits_through_result_copy(vindex, flat_index, monkeypatch): payload, data, _ = flat_index reader = vindex.VectorIndexReader(BytesInput(payload)) - params = vindex.RangeSearchParams(vindex.DistanceBand("l2"), 4) + params = vindex.RangeSearchParams(vindex.DistanceBand.from_raw("l2"), 4) copy_entered = threading.Event() release_copy = threading.Event() close_entered = threading.Event() @@ -770,7 +913,7 @@ def integers(count): upper = ( struct.unpack("