From 0f1a89002a52742ec37a5d8e9f407cfccf595c98 Mon Sep 17 00:00:00 2001 From: Mohammad Shoaib Ansari Date: Fri, 2 Oct 2026 03:04:20 +0000 Subject: [PATCH] Reject non-finite brute-force search distances --- CMakeLists.txt | 1 + README.md | 2 + hnswlib/bruteforce.h | 37 ++ python_bindings/bindings.cpp | 15 +- tests/cpp/bruteforce_numeric_range_test.cpp | 349 ++++++++++++++++++ .../python/bindings_test_bf_numeric_range.py | 153 ++++++++ 6 files changed, 553 insertions(+), 4 deletions(-) create mode 100644 tests/cpp/bruteforce_numeric_range_test.cpp create mode 100644 tests/python/bindings_test_bf_numeric_range.py diff --git a/CMakeLists.txt b/CMakeLists.txt index 75eb5d56..893bed0e 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -28,6 +28,7 @@ set(EXAMPLE_NAMES ) set(TEST_NAMES + bruteforce_numeric_range_test epsilon_search_test multiThread_replace_test multiThreadLoad_test diff --git a/README.md b/README.md index f8e04571..45aca317 100644 --- a/README.md +++ b/README.md @@ -71,6 +71,8 @@ Note that inner product is not an actual metric. An element can be closer to som For other spaces use the nmslib library https://github.com/nmslib/nmslib. +For brute-force search, `BFIndex.knn_query` raises `RuntimeError` if a distance to an allowed label is NaN or infinite, including when finite coordinates overflow during distance calculation. Use smaller magnitudes to keep squared L2 distances within float32 range. Labels excluded by `filter` do not cause this error, and `k=0` returns empty results without evaluating distances. The C++ `BruteforceSearch` and `BruteforceSearch` APIs return an error `Status` through `searchKnnNoExceptions`; their throwing wrappers raise `std::runtime_error` when exceptions are enabled. + #### API description * `hnswlib.get_simd()` returns the ISA selected at import (`sse`, `avx`, `avx512`, or `aarch64`). Set `HNSWLIB_SIMD` to force a lower-or-equal level (raises if the CPU cannot run it). * `hnswlib.Index(space, dim)` creates a non-initialized index an HNSW in space `space` with integer dimension `dim`. diff --git a/hnswlib/bruteforce.h b/hnswlib/bruteforce.h index e67b6cbd..994226b5 100644 --- a/hnswlib/bruteforce.h +++ b/hnswlib/bruteforce.h @@ -8,6 +8,9 @@ #include #include #include +#include +#include +#include namespace hnswlib { template @@ -120,9 +123,18 @@ class BruteforceSearch : public AlgorithmInterface { searchKnnNoExceptions(const void *query_data, size_t k, BaseFilterFunctor* isIdAllowed = nullptr) const override { assert(k <= cur_element_count); std::priority_queue> topResults; + if (k == 0) + return topResults; dist_t lastdist = std::numeric_limits::max(); for (int i = 0; i < cur_element_count; i++) { dist_t dist = fstdistfunc_(query_data, data_ + size_per_element_ * i, dist_func_param_); + if (!distanceIsFinite(dist)) { + labeltype label = getExternalLabel(i); + if (isIdAllowed && !(*isIdAllowed)(label)) + continue; + return Status("BruteforceSearch encountered a non-finite distance; " + "check input values or rescale to avoid overflow"); + } if (dist <= lastdist || topResults.size() < k) { labeltype label = getExternalLabel(i); if ((!isIdAllowed) || (*isIdAllowed)(label)) { @@ -196,5 +208,30 @@ class BruteforceSearch : public AlgorithmInterface { input.close(); } + + private: + template + static bool distanceIsFinite(const T&) { + return true; + } + + // Inspect the exponent bits: std::isfinite may be optimized away by fast-math. + static bool distanceIsFinite(float distance) { + static_assert(sizeof(float) == sizeof(std::uint32_t) && + std::numeric_limits::is_iec559, + "BruteforceSearch requires IEEE 754 binary32 floats"); + std::uint32_t bits; + std::memcpy(&bits, &distance, sizeof(bits)); + return (bits & 0x7f800000u) != 0x7f800000u; + } + + static bool distanceIsFinite(double distance) { + static_assert(sizeof(double) == sizeof(std::uint64_t) && + std::numeric_limits::is_iec559, + "BruteforceSearch requires IEEE 754 binary64 doubles"); + std::uint64_t bits; + std::memcpy(&bits, &distance, sizeof(bits)); + return (bits & 0x7ff0000000000000ULL) != 0x7ff0000000000000ULL; + } }; } // namespace hnswlib diff --git a/python_bindings/bindings.cpp b/python_bindings/bindings.cpp index 9d2c1082..221fd44d 100644 --- a/python_bindings/bindings.cpp +++ b/python_bindings/bindings.cpp @@ -6,6 +6,7 @@ #include "hnswlib.h" #include #include +#include #include #include @@ -856,6 +857,8 @@ class BFIndex { const std::function& filter = nullptr) { py::array_t < dist_t, py::array::c_style | py::array::forcecast > items(input); auto buffer = items.request(); + std::unique_ptr labels_owner; + std::unique_ptr distances_owner; hnswlib::labeltype *data_numpy_l; dist_t *data_numpy_d; size_t rows, features; @@ -867,8 +870,10 @@ class BFIndex { py::gil_scoped_release l; get_input_array_shapes(buffer, &rows, &features); - data_numpy_l = new hnswlib::labeltype[rows * k]; - data_numpy_d = new dist_t[rows * k]; + labels_owner.reset(new hnswlib::labeltype[rows * k]); + data_numpy_l = labels_owner.get(); + distances_owner.reset(new dist_t[rows * k]); + data_numpy_d = distances_owner.get(); CustomFilterFunctor idFilter(filter); CustomFilterFunctor* p_idFilter = filter ? &idFilter : nullptr; @@ -909,11 +914,13 @@ class BFIndex { } py::capsule free_when_done_l(data_numpy_l, [](void *f) { - delete[] f; + delete[] static_cast(f); }); + labels_owner.release(); py::capsule free_when_done_d(data_numpy_d, [](void *f) { - delete[] f; + delete[] static_cast(f); }); + distances_owner.release(); return py::make_tuple( diff --git a/tests/cpp/bruteforce_numeric_range_test.cpp b/tests/cpp/bruteforce_numeric_range_test.cpp new file mode 100644 index 00000000..686d0493 --- /dev/null +++ b/tests/cpp/bruteforce_numeric_range_test.cpp @@ -0,0 +1,349 @@ +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include "hnswlib/hnswlib.h" + +namespace { + +// std::isfinite can be optimized away under CMake's -Ofast flags. +bool finiteBits(float value) { + uint32_t bits; + std::memcpy(&bits, &value, sizeof(bits)); + return (bits & UINT32_C(0x7f800000)) != UINT32_C(0x7f800000); +} + +bool finiteBits(double value) { + uint64_t bits; + std::memcpy(&bits, &value, sizeof(bits)); + return (bits & UINT64_C(0x7ff0000000000000)) != UINT64_C(0x7ff0000000000000); +} + +void require(bool condition, const char* message) { + if (!condition) { + std::cerr << "Assertion failed: " << message << std::endl; + std::abort(); + } +} + +template +void expectNonFiniteError(const Result& result) { + require(!result.ok(), "eligible non-finite distance must return an error Status"); + require(std::strstr(result.status().message(), "non-finite distance") != nullptr, + "error Status must identify a non-finite distance"); +} + +#if defined(__EXCEPTIONS) || _HAS_EXCEPTIONS == 1 +template +void expectRuntimeError(Search search) { + bool caught = false; + try { + search(); + } catch (const std::runtime_error& error) { + caught = true; + require(std::strstr(error.what(), "non-finite distance") != nullptr, + "throwing wrapper must preserve the non-finite distance error"); + } + require(caught, "throwing wrapper must translate the error to std::runtime_error"); +} +#endif + +void expectNear(float actual, double expected) { + require(finiteBits(actual), "finite control must return a finite distance"); + const double difference = static_cast(actual) - expected; + require(difference >= -expected * 0.000002 && + difference <= expected * 0.000002, + "finite control must preserve its squared L2 distance"); +} + +class ExcludeLabel : public hnswlib::BaseFilterFunctor { + public: + explicit ExcludeLabel(hnswlib::labeltype excluded) : excluded_(excluded), calls(0) {} + + bool operator()(hnswlib::labeltype label) override { + ++calls; + return label != excluded_; + } + + private: + hnswlib::labeltype excluded_; + + public: + size_t calls; +}; + +void testReportedL2Case(size_t dim, float alpha, bool expect_error) { + std::vector p1(dim, 0.0f); + std::vector p2(dim, 0.0f); + std::vector query(dim, 0.0f); + p1[0] = alpha; + p2[1] = alpha; + query[0] = 0.6f * alpha; + query[1] = 0.8f * alpha; + for (size_t i = 0; i < dim; ++i) { + require(finiteBits(p1[i]) && finiteBits(p2[i]) && finiteBits(query[i]), + "reported overflow fixture must have only finite coordinates"); + } + + hnswlib::L2Space space(dim); + hnswlib::BruteforceSearch index(&space, 2); + require(index.addPointNoExceptions(p1.data(), 1).ok(), "p1 insertion must succeed"); + require(index.addPointNoExceptions(p2.data(), 2).ok(), "p2 insertion must succeed"); + + std::cout << "Reported L2 case: dim=" << dim << ", alpha=" << alpha << std::endl; + auto result = index.searchKnnNoExceptions(query.data(), 2); + if (expect_error) { + expectNonFiniteError(result); + expectNonFiniteError(index.searchKnnNoExceptions(query.data(), 1)); + expectNonFiniteError(index.searchKnnCloserFirstNoExceptions(query.data(), 2)); + return; + } + + require(result.ok(), "finite reported case must succeed"); + auto heap = std::move(result).value(); + require(heap.size() == 2, "finite reported case must return both labels"); + const double scale = static_cast(alpha) * static_cast(alpha); + require(heap.top().second == 1, "p1 must be the farther finite result"); + expectNear(heap.top().first, 0.8 * scale); + heap.pop(); + require(heap.top().second == 2, "p2 must be the nearest finite result"); + expectNear(heap.top().first, 0.4 * scale); + + auto closest = index.searchKnnCloserFirstNoExceptions(query.data(), 2); + require(closest.ok(), "finite closer-first search must succeed"); + require(closest.value().size() == 2 && + closest.value()[0].second == 2 && closest.value()[1].second == 1, + "finite closer-first search must preserve nearest-first label order"); +} + +void testL2OverflowShapes() { + const float zero[] = {0.0f, 0.0f}; + + // A single squared difference exceeds FLT_MAX. + const float individual_square[] = {3e19f}; + require(finiteBits(individual_square[0]), "individual-square coordinate must be finite"); + require(!finiteBits(individual_square[0] * individual_square[0]), + "individual-square fixture must overflow in one square"); + hnswlib::L2Space scalar_space(1); + hnswlib::BruteforceSearch scalar_index(&scalar_space, 1); + require(scalar_index.addPointNoExceptions(individual_square, 1).ok(), + "individual-square insertion must succeed"); + expectNonFiniteError(scalar_index.searchKnnNoExceptions(zero, 1)); + + // Each square is finite (2.25e38); their sum (4.5e38) exceeds FLT_MAX. + const float accumulation_only[] = {1.5e19f, 1.5e19f}; + for (size_t i = 0; i < 2; ++i) { + require(finiteBits(accumulation_only[i]), "accumulation coordinate must be finite"); + require(finiteBits(accumulation_only[i] * accumulation_only[i]), + "each accumulation-only square must remain finite"); + } + hnswlib::L2Space accumulation_space(2); + hnswlib::BruteforceSearch accumulation_index(&accumulation_space, 1); + require(accumulation_index.addPointNoExceptions(accumulation_only, 1).ok(), + "accumulation-only insertion must succeed"); + expectNonFiniteError(accumulation_index.searchKnnNoExceptions(zero, 1)); +} + +void testL2FilterIncludesOverflowAfterFiniteNeighbor() { + const float query[] = {0.0f, 0.0f}; + const float finite_neighbor[] = {1.0f, 0.0f}; + const float overflow_neighbor[] = {3e19f, 0.0f}; + hnswlib::L2Space space(2); + hnswlib::BruteforceSearch index(&space, 2); + require(index.addPointNoExceptions(finite_neighbor, 10).ok(), "finite insertion must succeed"); + require(index.addPointNoExceptions(overflow_neighbor, 20).ok(), "overflow insertion must succeed"); + + ExcludeLabel exclude_overflow(20); + auto finite_result = index.searchKnnNoExceptions(query, 1, &exclude_overflow); + require(finite_result.ok(), "filtered-out overflow must not cause an error"); + require(finite_result.value().size() == 1 && + finite_result.value().top().second == 10 && + finite_result.value().top().first == 1.0f, + "filter must retain the finite nearest neighbor"); + + ExcludeLabel include_both(0); + expectNonFiniteError(index.searchKnnNoExceptions(query, 1, &include_both)); +} + +// A custom numeric space isolates the search's handling of distance values, +// including double, from the built-in float-only L2 distance implementation. +template +class NumericDistanceSpace : public hnswlib::SpaceInterface { + public: + NumericDistanceSpace() : calls(0) {} + + size_t get_data_size() override { + return sizeof(Distance); + } + + hnswlib::DISTFUNC get_dist_func() override { + return distance; + } + + void* get_dist_func_param() override { + return &calls; + } + + size_t calls; + + private: + static Distance distance(const void*, const void* point, const void* parameter) { + ++*const_cast(static_cast(parameter)); + Distance result; + std::memcpy(&result, point, sizeof(result)); + return result; + } +}; + +// Build special values from IEEE bit patterns so Clang's fast-math warning +// does not reject compile-time infinity constants. +template +struct NonFiniteFixtures; + +template<> +struct NonFiniteFixtures { + float values[3]; + + NonFiniteFixtures() { + const uint32_t bits[] = { + UINT32_C(0x7f800000), // +infinity + UINT32_C(0xff800000), // -infinity + UINT32_C(0x7fc00000) // quiet NaN + }; + std::memcpy(values, bits, sizeof(values)); + } +}; + +template<> +struct NonFiniteFixtures { + double values[3]; + + NonFiniteFixtures() { + const uint64_t bits[] = { + UINT64_C(0x7ff0000000000000), // +infinity + UINT64_C(0xfff0000000000000), // -infinity + UINT64_C(0x7ff8000000000000) // quiet NaN + }; + std::memcpy(values, bits, sizeof(values)); + } +}; + +template +void testNonFiniteDistances() { + const NonFiniteFixtures fixtures; + const Distance* invalid = fixtures.values; + const Distance query = 0; + const Distance finite_neighbor = 1; + for (size_t i = 0; i < 3; ++i) { + require(!finiteBits(invalid[i]), "non-finite distance fixture must retain its bit pattern"); + NumericDistanceSpace space; + hnswlib::BruteforceSearch index(&space, 2); + require(index.addPointNoExceptions(&finite_neighbor, 10).ok(), + "custom finite insertion must succeed"); + require(index.addPointNoExceptions(&invalid[i], 20).ok(), + "custom non-finite insertion must succeed"); + + // Even after the heap is full, an eligible invalid candidate is an error. + expectNonFiniteError(index.searchKnnNoExceptions(&query, 1)); + expectNonFiniteError(index.searchKnnCloserFirstNoExceptions(&query, 1)); + + ExcludeLabel exclude_invalid(20); + auto result = index.searchKnnNoExceptions(&query, 1, &exclude_invalid); + require(result.ok(), "filtered-out non-finite distance must not cause an error"); + require(result.value().size() == 1 && result.value().top().second == 10 && + result.value().top().first == 1, + "filter must preserve the finite custom-distance result"); + +#if defined(__EXCEPTIONS) || _HAS_EXCEPTIONS == 1 + expectRuntimeError([&]() { index.searchKnn(&query, 1); }); + expectRuntimeError([&]() { index.searchKnnCloserFirst(&query, 1); }); +#endif + } +} + +template +void testZeroKDoesNotEvaluateDistancesOrFilters() { + NumericDistanceSpace space; + hnswlib::BruteforceSearch index(&space, 1); + const Distance query = 0; + const NonFiniteFixtures fixtures; + const Distance invalid = fixtures.values[2]; + ExcludeLabel include_all(0); + + auto empty = index.searchKnnNoExceptions(&query, 0, &include_all); + require(empty.ok() && empty.value().empty(), "empty index with k=0 must succeed"); + require(index.addPointNoExceptions(&invalid, 20).ok(), "k=0 fixture insertion must succeed"); + + auto heap = index.searchKnnNoExceptions(&query, 0, &include_all); + require(heap.ok() && heap.value().empty(), "k=0 must return an empty heap"); + auto closest = index.searchKnnCloserFirstNoExceptions(&query, 0, &include_all); + require(closest.ok() && closest.value().empty(), "k=0 must return an empty closer-first vector"); + require(index.searchKnn(&query, 0, &include_all).empty(), "throwing k=0 wrapper must return empty"); + require(index.searchKnnCloserFirst(&query, 0, &include_all).empty(), + "throwing closer-first k=0 wrapper must return empty"); + require(space.calls == 0, "k=0 must not evaluate any distance"); + require(include_all.calls == 0, "k=0 must not invoke the label filter"); +} + +void testFiniteCustomDoubleDistances() { + NumericDistanceSpace space; + hnswlib::BruteforceSearch index(&space, 3); + const double query = 0; + const double points[] = {-4.0, 1.0, 9.0}; + for (size_t i = 0; i < 3; ++i) { + require(index.addPointNoExceptions(&points[i], i + 1).ok(), + "finite double insertion must succeed"); + } + auto result = index.searchKnnCloserFirstNoExceptions(&query, 2); + require(result.ok(), "finite custom double distances must remain supported"); + require(result.value().size() == 2 && + result.value()[0].first == -4.0 && result.value()[0].second == 1 && + result.value()[1].first == 1.0 && result.value()[1].second == 2, + "finite custom double distances must preserve their values and order"); +} + +void testIntegerDistancesRemainSupported() { + hnswlib::L2SpaceI space(4); + hnswlib::BruteforceSearch index(&space, 2); + const unsigned char query[] = {0, 0, 0, 0}; + const unsigned char p1[] = {1, 2, 0, 0}; + const unsigned char p2[] = {3, 4, 0, 0}; + require(index.addPointNoExceptions(p1, 11).ok(), "integer p1 insertion must succeed"); + require(index.addPointNoExceptions(p2, 22).ok(), "integer p2 insertion must succeed"); + auto result = index.searchKnnCloserFirstNoExceptions(query, 2); + require(result.ok(), "integer distance searches must remain supported"); + require(result.value().size() == 2 && + result.value()[0].first == 5 && result.value()[0].second == 11 && + result.value()[1].first == 25 && result.value()[1].second == 22, + "integer L2 distances must preserve values and nearest-first order"); +} + +} // namespace + +int main() { + const size_t dims[] = {2, 4, 16, 17}; + testFiniteCustomDoubleDistances(); + testIntegerDistancesRemainSupported(); + for (size_t dim : dims) { + testReportedL2Case(dim, 1.0f, false); + testReportedL2Case(dim, 2e19f, false); + } + for (size_t dim : dims) { + testReportedL2Case(dim, 3e19f, true); + testReportedL2Case(dim, 2.5e19f, true); + } + testL2OverflowShapes(); + testL2FilterIncludesOverflowAfterFiniteNeighbor(); + testNonFiniteDistances(); + testNonFiniteDistances(); + testZeroKDoesNotEvaluateDistancesOrFilters(); + testZeroKDoesNotEvaluateDistancesOrFilters(); + std::cout << "Test ok" << std::endl; + return 0; +} diff --git a/tests/python/bindings_test_bf_numeric_range.py b/tests/python/bindings_test_bf_numeric_range.py new file mode 100644 index 00000000..a4915ba7 --- /dev/null +++ b/tests/python/bindings_test_bf_numeric_range.py @@ -0,0 +1,153 @@ +import unittest + +import hnswlib +import numpy as np + + +class BruteforceNumericRangeTestCase(unittest.TestCase): + def _index(self, data, space='l2'): + index = hnswlib.BFIndex(space=space, dim=data.shape[1]) + index.init_index(max_elements=len(data)) + index.add_items(data, np.arange(1, len(data) + 1)) + return index + + def test_l2_overflow_reports_error(self): + for dim in (2, 3, 4, 16, 17, 32, 128): + for alpha in (2.5e19, 3e19): + data = np.zeros((2, dim), dtype=np.float32) + data[0, 0] = alpha + data[1, 1] = alpha + query = np.zeros((1, dim), dtype=np.float32) + query[0, :2] = [0.6 * alpha, 0.8 * alpha] + reference = np.sum( + (data.astype(np.float64) - query.astype(np.float64)) ** 2, + axis=1) + self.assertTrue(np.isfinite(data).all()) + self.assertTrue(np.isfinite(query).all()) + self.assertLess(reference[1], reference[0]) + self.assertGreater(reference.max(), np.finfo(np.float32).max) + index = self._index(data) + for k in (1, 2): + with self.subTest(dim=dim, alpha=alpha, k=k): + with self.assertRaisesRegex( + RuntimeError, 'non-finite distance'): + index.knn_query(query, k=k, num_threads=1) + + def test_finite_l2_results_match_float64_reference(self): + for dim in (2, 3, 4, 16, 17, 128): + for alpha in (1.0, 1e10, 1e19): + with self.subTest(dim=dim, alpha=alpha): + data = np.zeros((2, dim), dtype=np.float32) + data[0, 0] = alpha + data[1, 1] = alpha + query = np.zeros((1, dim), dtype=np.float32) + query[0, :2] = [0.6 * alpha, 0.8 * alpha] + reference = np.sum( + (data.astype(np.float64) - + query.astype(np.float64)) ** 2, axis=1) + index = self._index(data) + for k in (1, 2): + labels, distances = index.knn_query( + query, k=k, num_threads=1) + np.testing.assert_array_equal( + labels[0], (np.argsort(reference) + 1)[:k]) + np.testing.assert_allclose( + distances[0], reference[labels[0] - 1], + rtol=2e-6, atol=0) + self.assertEqual(distances.dtype, np.dtype('float32')) + + def test_accumulation_overflow_reports_error(self): + for dim in (4, 16, 17): + with self.subTest(dim=dim): + data = np.full((1, dim), 1e19, dtype=np.float32) + squares = data.astype(np.float64) ** 2 + self.assertLess(squares.max(), np.finfo(np.float32).max) + self.assertGreater(squares.sum(), np.finfo(np.float32).max) + index = self._index(data) + with self.assertRaisesRegex(RuntimeError, 'non-finite distance'): + index.knn_query(np.zeros((1, dim), dtype=np.float32), + k=1, num_threads=1) + + def test_extreme_finite_coordinates_report_error(self): + largest = np.finfo(np.float32).max + data = np.array([[largest]], dtype=np.float32) + query = np.array([[-largest]], dtype=np.float32) + self.assertTrue(np.isfinite(data).all()) + self.assertTrue(np.isfinite(query).all()) + with self.assertRaisesRegex(RuntimeError, 'non-finite distance'): + self._index(data).knn_query(query, k=1, num_threads=1) + + def test_eligible_overflow_is_not_hidden_by_a_full_heap(self): + data = np.array([[0, 0], [1e20, 0]], dtype=np.float32) + with self.assertRaisesRegex(RuntimeError, 'non-finite distance'): + self._index(data).knn_query( + np.zeros((1, 2), dtype=np.float32), k=1, num_threads=1) + + def test_filter_can_exclude_an_overflowing_candidate(self): + data = np.array([[0, 0], [1e20, 0]], dtype=np.float32) + query = np.zeros((1, 2), dtype=np.float32) + index = self._index(data) + labels, distances = index.knn_query( + query, k=1, num_threads=1, filter=lambda label: label == 1) + np.testing.assert_array_equal(labels, [[1]]) + np.testing.assert_array_equal(distances, [[0]]) + with self.assertRaisesRegex(RuntimeError, 'non-finite distance'): + index.knn_query(query, k=1, num_threads=1, + filter=lambda label: label == 2) + + def test_error_propagates_from_parallel_queries_and_index_recovers(self): + data = np.array([[0, 0], [1e19, 0]], dtype=np.float32) + index = self._index(data) + queries = np.zeros((64, 2), dtype=np.float32) + queries[37, 0] = np.finfo(np.float32).max + for threads in (1, 4): + with self.subTest(threads=threads): + with self.assertRaisesRegex(RuntimeError, 'non-finite distance'): + index.knn_query(queries, k=2, num_threads=threads) + labels, distances = index.knn_query( + np.zeros((8, 2), dtype=np.float32), + k=2, num_threads=threads) + np.testing.assert_array_equal( + labels, np.tile([1, 2], (8, 1))) + self.assertTrue(np.isfinite(distances).all()) + + def test_cosine_parallel_error_and_recovery(self): + data = np.array([[1, 0], [0, 1]], dtype=np.float32) + index = self._index(data, space='cosine') + queries = np.tile(data[0], (64, 1)) + queries[37, 0] = np.nan + for threads in (1, 4): + with self.subTest(threads=threads): + with self.assertRaisesRegex(RuntimeError, 'non-finite distance'): + index.knn_query(queries, k=2, num_threads=threads) + labels, distances = index.knn_query( + np.tile(data[0], (8, 1)), k=2, num_threads=threads) + np.testing.assert_array_equal(labels, np.tile([1, 2], (8, 1))) + np.testing.assert_allclose(distances, np.tile([0, 1], (8, 1)), + rtol=0, atol=1e-6) + + def test_k_zero_returns_empty_arrays_without_calling_filter(self): + def unexpected_filter(label): + self.fail('k=0 should not evaluate the filter') + index = self._index(np.array([[1e20, 0]], dtype=np.float32)) + queries = np.zeros((8, 2), dtype=np.float32) + for threads in (1, 4): + labels, distances = index.knn_query( + queries, k=0, num_threads=threads, filter=unexpected_filter) + self.assertEqual(labels.shape, (8, 0)) + self.assertEqual(distances.shape, (8, 0)) + + def test_nonfinite_distance_results_report_error(self): + cases = ( + ('ip', 1e20, 1e20), # negative infinity distance + ('ip', -1e20, 1e20), # positive infinity distance + ('l2', np.nan, 0.0), # unordered distance + ) + for space, point, query in cases: + with self.subTest(space=space, point=point, query=query): + index = self._index(np.array([[point]], dtype=np.float32), + space=space) + with self.assertRaisesRegex(RuntimeError, 'non-finite distance'): + index.knn_query( + np.array([[query]], dtype=np.float32), + k=1, num_threads=1)