diff --git a/CMakeLists.txt b/CMakeLists.txt index d8987431..5059e503 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -104,7 +104,9 @@ function(add_example_or_test TARGET_NAME ...) endfunction() -option(HNSWLIB_ENABLE_EXCEPTIONS "Whether to enable exceptions in hnswlib" ON) +# Affects example/test executables in this repo only. The INTERFACE library +# does not export -fno-exceptions / /EHsc; consumers set their own flags. +option(HNSWLIB_ENABLE_EXCEPTIONS "Enable exceptions on hnswlib example/test targets" ON) if(HNSWLIB_ENABLE_EXCEPTIONS) message("Exceptions are enabled using HNSWLIB_ENABLE_EXCEPTIONS=ON (default)") else() @@ -145,7 +147,7 @@ endif() if (CMAKE_CXX_COMPILER_ID STREQUAL "MSVC") set(ENABLE_EXCEPTIONS_FLAGS /EHsc) - set(DISABLE_EXCEPTIONS_FLAGS /GR- /D_HAS_EXCEPTIONS=0) + set(DISABLE_EXCEPTIONS_FLAGS /EHs-c- /GR- /D_HAS_EXCEPTIONS=0) else() set(ENABLE_EXCEPTIONS_FLAGS -fexceptions) set(DISABLE_EXCEPTIONS_FLAGS -fno-exceptions) @@ -194,17 +196,6 @@ if(HNSWLIB_EXAMPLES) add_cxx_flags(-lrt) endif() elseif (CMAKE_CXX_COMPILER_ID STREQUAL "MSVC") - if (NOT HNSWLIB_ENABLE_EXCEPTIONS) - # Do not enable exceptions by default. We will enable them on a - # case by case basis when needed. - foreach(config IN ITEMS Debug Release RelWithDebInfo MinSizeRel) - string(TOUPPER ${config} config_upper) - set(FLAGS_VAR "CMAKE_CXX_FLAGS_${config_upper}") - string(REPLACE "/EHsc" "" ${FLAGS_VAR} "${${FLAGS_VAR}}") - set(${FLAGS_VAR} "${${FLAGS_VAR}}" CACHE STRING - "Flags for ${config} configuration." FORCE) - endforeach() - endif() add_cxx_flags(/O2 /W1 /openmp) endif() add_cxx_flags(-DHAVE_CXX0X) diff --git a/hnswlib/bruteforce.h b/hnswlib/bruteforce.h index 5a236a2c..6e618a90 100644 --- a/hnswlib/bruteforce.h +++ b/hnswlib/bruteforce.h @@ -64,6 +64,13 @@ class BruteforceSearch : public AlgorithmInterface { free(data_); } + // Labels sit at a packed offset; load via memcpy (nmslib/hnswlib#665). + inline labeltype getExternalLabel(size_t internal_id) const { + labeltype return_label; + memcpy(&return_label, data_ + internal_id * size_per_element_ + data_size_, sizeof(labeltype)); + return return_label; + } + Status addPointNoExceptions(const void *datapoint, labeltype label, bool replace_deleted = false) override { int idx; @@ -99,7 +106,7 @@ class BruteforceSearch : public AlgorithmInterface { dict_external_to_internal.erase(found); size_t cur_c = found->second; - labeltype label = *((labeltype*)(data_ + size_per_element_ * (cur_element_count-1) + data_size_)); + labeltype label = getExternalLabel(cur_element_count - 1); dict_external_to_internal[label] = cur_c; memcpy(data_ + size_per_element_ * cur_c, data_ + size_per_element_ * (cur_element_count-1), @@ -117,7 +124,7 @@ class BruteforceSearch : public AlgorithmInterface { for (int i = 0; i < cur_element_count; i++) { dist_t dist = fstdistfunc_(query_data, data_ + size_per_element_ * i, dist_func_param_); if (dist <= lastdist || topResults.size() < k) { - labeltype label = *((labeltype *) (data_ + size_per_element_ * i + data_size_)); + labeltype label = getExternalLabel(i); if ((!isIdAllowed) || (*isIdAllowed)(label)) { topResults.emplace(dist, label); if (topResults.size() > k) @@ -132,24 +139,23 @@ class BruteforceSearch : public AlgorithmInterface { Status saveIndexNoExceptions(std::ostream &output) { - StreamExceptionsOff guard(output); - return invokeWithoutStreamThrow([&]() -> Status { - if (!output) { - return Status("Cannot save index: output stream is not open or in a failed state"); - } - writeBinaryPOD(output, maxelements_); - writeBinaryPOD(output, size_per_element_); - writeBinaryPOD(output, cur_element_count); - if (!output.good()) { - return Status("Failed writing index metadata"); - } + // *NoExceptions I/O checks stream state. Callers must leave the default + // iostream exception mask (goodbit); enabling failbit/badbit can throw. + if (!output) { + return Status("Cannot save index: output stream is not open or in a failed state"); + } + writeBinaryPOD(output, maxelements_); + writeBinaryPOD(output, size_per_element_); + writeBinaryPOD(output, cur_element_count); + if (!output.good()) { + return Status("Failed writing index metadata"); + } - output.write(data_, maxelements_ * size_per_element_); - if (!output.good()) { - return Status("Failed writing vector data"); - } - return OkStatus(); - }); + output.write(data_, maxelements_ * size_per_element_); + if (!output.good()) { + return Status("Failed writing vector data"); + } + return OkStatus(); } @@ -162,32 +168,86 @@ class BruteforceSearch : public AlgorithmInterface { } - void loadIndex(std::istream &input, SpaceInterface *s) { + Status loadIndexNoExceptions(std::istream &input, SpaceInterface *s) { + if (!input) { + return Status("Cannot load index: input stream is not open or not readable"); + } + size_t file_maxelements = 0; + size_t file_size_per_element = 0; + size_t file_cur_count = 0; + readBinaryPOD(input, file_maxelements); + readBinaryPOD(input, file_size_per_element); + readBinaryPOD(input, file_cur_count); if (!input) { - HNSWLIB_THROW_RUNTIME_ERROR("Cannot load index: input stream is not open or not readable"); + return Status("Cannot load index: failed to read index header"); + } + if (file_cur_count > file_maxelements) { + return Status("Cannot load index: cur_element_count exceeds maxelements"); } - readBinaryPOD(input, maxelements_); - readBinaryPOD(input, size_per_element_); - readBinaryPOD(input, cur_element_count); - data_size_ = s->get_data_size(); + size_t data_size = s->get_data_size(); + size_t size_per_element = data_size + sizeof(labeltype); + if (file_size_per_element != size_per_element) { + return Status("Cannot load index: size_per_element does not match space"); + } + char *new_data = (char *) malloc(file_maxelements * size_per_element); + if (new_data == nullptr) + return Status("Not enough memory: loadIndex failed to allocate data"); + input.read(new_data, file_maxelements * size_per_element); + if (!input) { + free(new_data); + return Status("Cannot load index: failed to read vector data"); + } + + std::unordered_map new_dict; +#if defined(__EXCEPTIONS) || _HAS_EXCEPTIONS == 1 + try { +#endif + for (size_t i = 0; i < file_cur_count; i++) { + labeltype lab; + memcpy(&lab, new_data + i * size_per_element + data_size, sizeof(lab)); + new_dict[lab] = i; + } +#if defined(__EXCEPTIONS) || _HAS_EXCEPTIONS == 1 + } catch (const std::bad_alloc&) { + free(new_data); + return Status("Not enough memory: loadIndex failed to rebuild label map"); + } +#endif + + free(data_); + data_ = new_data; + maxelements_ = file_maxelements; + cur_element_count = file_cur_count; + data_size_ = data_size; + size_per_element_ = size_per_element; fstdistfunc_ = s->get_dist_func(); dist_func_param_ = s->get_dist_func_param(); - size_per_element_ = data_size_ + sizeof(labeltype); - data_ = (char *) malloc(maxelements_ * size_per_element_); - if (data_ == nullptr) - HNSWLIB_THROW_RUNTIME_ERROR("Not enough memory: loadIndex failed to allocate data"); - - input.read(data_, maxelements_ * size_per_element_); + dict_external_to_internal.swap(new_dict); + return OkStatus(); } - - void loadIndex(const std::string &location, SpaceInterface *s) { + Status loadIndexNoExceptions(const std::string &location, SpaceInterface *s) { std::ifstream input(location, std::ios::binary); + if (!input.is_open()) { + return Status("Cannot load index: input stream is not open or not readable"); + } + return loadIndexNoExceptions(input, s); + } + + void loadIndex(std::istream &input, SpaceInterface *s) { + Status status = loadIndexNoExceptions(input, s); + if (!status.ok()) { + HNSWLIB_THROW_RUNTIME_ERROR(status.message()); + } + } - loadIndex(input, s); - input.close(); + void loadIndex(const std::string &location, SpaceInterface *s) { + Status status = loadIndexNoExceptions(location, s); + if (!status.ok()) { + HNSWLIB_THROW_RUNTIME_ERROR(status.message()); + } } }; } // namespace hnswlib diff --git a/hnswlib/hnswalg.h b/hnswlib/hnswalg.h index 6d8f7374..5a2d6a52 100644 --- a/hnswlib/hnswalg.h +++ b/hnswlib/hnswalg.h @@ -161,16 +161,26 @@ class HierarchicalNSW : public AlgorithmInterface { void clear() { free(data_level0_memory_); data_level0_memory_ = nullptr; - for (tableint i = 0; i < cur_element_count; i++) { - if (element_levels_[i] > 0) - free(linkLists_[i]); - } if (linkLists_) { + const size_t n = cur_element_count; + const size_t n_levels = element_levels_.size(); + for (tableint i = 0; i < n && i < n_levels; i++) { + if (element_levels_[i] > 0) + free(linkLists_[i]); + } free(linkLists_); + linkLists_ = nullptr; } - linkLists_ = nullptr; cur_element_count = 0; + max_elements_ = 0; + num_deleted_ = 0; + label_lookup_.clear(); + deleted_elements.clear(); + element_levels_.clear(); + enterpoint_node_ = -1; + maxlevel_ = -1; visited_list_pool_.reset(nullptr); + std::vector().swap(link_list_locks_); } @@ -722,47 +732,46 @@ class HierarchicalNSW : public AlgorithmInterface { } Status saveIndexNoExceptions(std::ostream &output) { - StreamExceptionsOff guard(output); - return invokeWithoutStreamThrow([&]() -> Status { - if (!output) { - return Status("Cannot save index: output stream is not open or in a failed state"); - } - writeBinaryPOD(output, offsetLevel0_); - writeBinaryPOD(output, max_elements_); - writeBinaryPOD(output, cur_element_count); - writeBinaryPOD(output, size_data_per_element_); - writeBinaryPOD(output, label_offset_); - writeBinaryPOD(output, offsetData_); - writeBinaryPOD(output, maxlevel_); - writeBinaryPOD(output, enterpoint_node_); - writeBinaryPOD(output, maxM_); - - writeBinaryPOD(output, maxM0_); - writeBinaryPOD(output, M_); - writeBinaryPOD(output, mult_); - writeBinaryPOD(output, ef_construction_); + // *NoExceptions I/O checks stream state. Callers must leave the default + // iostream exception mask (goodbit); enabling failbit/badbit can throw. + if (!output) { + return Status("Cannot save index: output stream is not open or in a failed state"); + } + writeBinaryPOD(output, offsetLevel0_); + writeBinaryPOD(output, max_elements_); + writeBinaryPOD(output, cur_element_count); + writeBinaryPOD(output, size_data_per_element_); + writeBinaryPOD(output, label_offset_); + writeBinaryPOD(output, offsetData_); + writeBinaryPOD(output, maxlevel_); + writeBinaryPOD(output, enterpoint_node_); + writeBinaryPOD(output, maxM_); + + writeBinaryPOD(output, maxM0_); + writeBinaryPOD(output, M_); + writeBinaryPOD(output, mult_); + writeBinaryPOD(output, ef_construction_); + + if (!output.good()) { + return Status("Failed writing index metadata"); + } - if (!output.good()) { - return Status("Failed writing index metadata"); - } + output.write(data_level0_memory_, cur_element_count * size_data_per_element_); + if (!output.good()) { + return Status("Failed writing level 0 memory block"); + } - output.write(data_level0_memory_, cur_element_count * size_data_per_element_); - if (!output.good()) { - return Status("Failed writing level 0 memory block"); + for (size_t i = 0; i < cur_element_count; i++) { + unsigned int linkListSize = element_levels_[i] > 0 ? size_links_per_element_ * element_levels_[i] : 0; + writeBinaryPOD(output, linkListSize); + if (linkListSize) { + output.write(linkLists_[i], linkListSize); } - - for (size_t i = 0; i < cur_element_count; i++) { - unsigned int linkListSize = element_levels_[i] > 0 ? size_links_per_element_ * element_levels_[i] : 0; - writeBinaryPOD(output, linkListSize); - if (linkListSize) { - output.write(linkLists_[i], linkListSize); - } - if (!output.good()) { - return Status("Failed writing link list elements"); - } + if (!output.good()) { + return Status("Failed writing link list elements"); } - return OkStatus(); - }); + } + return OkStatus(); } Status saveIndexNoExceptions(const std::string &location) override { @@ -781,15 +790,11 @@ class HierarchicalNSW : public AlgorithmInterface { } Status loadIndexNoExceptions(std::istream &input, SpaceInterface *s, size_t max_elements_i = 0) { - // Must not destroy the live index until the stream is known to be readable. - // An unopened / failed stream used to seek to -1, skip the empty-index - // corruption loop, and return OkStatus() after clear(). + // Must not destroy the live index until the file layout is known to be valid. if (!input) { return Status("Cannot load index: input stream is not open or not readable"); } - StreamExceptionsOff guard(input); - return invokeWithoutStreamThrow([&]() -> Status { // Default-constructed ifstreams are often still good() on libc++. // If we cannot peek a byte, there is no index to load — leave the // live index untouched. @@ -797,7 +802,6 @@ class HierarchicalNSW : public AlgorithmInterface { return Status("Cannot load index: input stream is not open or not readable"); } - // get file size before mutating the in-memory index: input.seekg(0, input.end); std::streampos total_filesize = input.tellg(); input.seekg(0, input.beg); @@ -805,40 +809,51 @@ class HierarchicalNSW : public AlgorithmInterface { return Status("Cannot load index: failed to determine stream size"); } - clear(); - - readBinaryPOD(input, offsetLevel0_); - readBinaryPOD(input, max_elements_); - readBinaryPOD(input, cur_element_count); - - size_t max_elements = max_elements_i; - if (max_elements < cur_element_count) - max_elements = max_elements_; - max_elements_ = max_elements; - readBinaryPOD(input, size_data_per_element_); - readBinaryPOD(input, label_offset_); - readBinaryPOD(input, offsetData_); - readBinaryPOD(input, maxlevel_); - readBinaryPOD(input, enterpoint_node_); - - readBinaryPOD(input, maxM_); - readBinaryPOD(input, maxM0_); - readBinaryPOD(input, M_); - readBinaryPOD(input, mult_); - readBinaryPOD(input, ef_construction_); + size_t offsetLevel0 = 0; + size_t file_max_elements = 0; + size_t file_cur_count = 0; + size_t size_data_per_element = 0; + size_t label_offset = 0; + size_t offsetData = 0; + int maxlevel = -1; + tableint enterpoint_node = static_cast(-1); + size_t maxM = 0; + size_t maxM0 = 0; + size_t M = 0; + double mult = 0.0; + size_t ef_construction = 0; + + readBinaryPOD(input, offsetLevel0); + readBinaryPOD(input, file_max_elements); + readBinaryPOD(input, file_cur_count); + readBinaryPOD(input, size_data_per_element); + readBinaryPOD(input, label_offset); + readBinaryPOD(input, offsetData); + readBinaryPOD(input, maxlevel); + readBinaryPOD(input, enterpoint_node); + readBinaryPOD(input, maxM); + readBinaryPOD(input, maxM0); + readBinaryPOD(input, M); + readBinaryPOD(input, mult); + readBinaryPOD(input, ef_construction); if (!input) { return Status("Cannot load index: failed to read index header"); } + if (file_cur_count > file_max_elements) { + return Status("Cannot load index: cur_element_count exceeds max_elements"); + } + if (size_data_per_element == 0) { + return Status("Cannot load index: invalid size_data_per_element"); + } - data_size_ = s->get_data_size(); - fstdistfunc_ = s->get_dist_func(); - dist_func_param_ = s->get_dist_func_param(); + size_t max_elements = max_elements_i; + if (max_elements < file_cur_count) + max_elements = file_max_elements; auto pos = input.tellg(); - /// Optional - check if index is ok: - input.seekg(cur_element_count * size_data_per_element_, input.cur); - for (size_t i = 0; i < cur_element_count; i++) { + input.seekg(file_cur_count * size_data_per_element, input.cur); + for (size_t i = 0; i < file_cur_count; i++) { if (input.tellg() < 0 || input.tellg() >= total_filesize) { return Status("Index seems to be corrupted or unsupported"); } @@ -850,35 +865,61 @@ class HierarchicalNSW : public AlgorithmInterface { } } - // throw exception if it either corrupted or old index if (input.tellg() != total_filesize) return Status("Index seems to be corrupted or unsupported"); input.clear(); - /// Optional check end - input.seekg(pos, input.beg); + // Layout is valid. Replace the live index only now. + clear(); + + offsetLevel0_ = offsetLevel0; + max_elements_ = max_elements; + cur_element_count = file_cur_count; + size_data_per_element_ = size_data_per_element; + label_offset_ = label_offset; + offsetData_ = offsetData; + maxlevel_ = maxlevel; + enterpoint_node_ = enterpoint_node; + maxM_ = maxM; + maxM0_ = maxM0; + M_ = M; + mult_ = mult; + ef_construction_ = ef_construction; + + data_size_ = s->get_data_size(); + fstdistfunc_ = s->get_dist_func(); + dist_func_param_ = s->get_dist_func_param(); + data_level0_memory_ = (char *) malloc(max_elements * size_data_per_element_); - if (data_level0_memory_ == nullptr) + if (data_level0_memory_ == nullptr) { + clear(); return Status("Not enough memory: loadIndex failed to allocate level0"); - input.read(data_level0_memory_, cur_element_count * size_data_per_element_); + } + input.read(data_level0_memory_, file_cur_count * size_data_per_element_); size_links_per_element_ = maxM_ * sizeof(tableint) + sizeof(linklistsizeint); size_links_level0_ = maxM0_ * sizeof(tableint) + sizeof(linklistsizeint); +#if defined(__EXCEPTIONS) || _HAS_EXCEPTIONS == 1 + try { +#endif std::vector(max_elements).swap(link_list_locks_); std::vector(MAX_LABEL_OPERATION_LOCKS).swap(label_op_locks_); visited_list_pool_.reset(new VisitedListPool(1, max_elements)); linkLists_ = (char **) malloc(sizeof(void *) * max_elements); - if (linkLists_ == nullptr) + if (linkLists_ == nullptr) { + clear(); return Status("Not enough memory: loadIndex failed to allocate linklists"); + } + memset(linkLists_, 0, sizeof(void *) * max_elements); element_levels_ = std::vector(max_elements); revSize_ = 1.0 / mult_; ef_ = 10; - for (size_t i = 0; i < cur_element_count; i++) { + for (size_t i = 0; i < file_cur_count; i++) { label_lookup_[getExternalLabel(i)] = i; unsigned int linkListSize; readBinaryPOD(input, linkListSize); @@ -888,13 +929,15 @@ class HierarchicalNSW : public AlgorithmInterface { } else { element_levels_[i] = linkListSize / size_links_per_element_; linkLists_[i] = (char *) malloc(linkListSize); - if (linkLists_[i] == nullptr) + if (linkLists_[i] == nullptr) { + clear(); return Status("Not enough memory: loadIndex failed to allocate linklist"); + } input.read(linkLists_[i], linkListSize); } } - for (size_t i = 0; i < cur_element_count; i++) { + for (size_t i = 0; i < file_cur_count; i++) { if (isMarkedDeleted(i)) { num_deleted_ += 1; if (allow_replace_deleted_) deleted_elements.insert(i); @@ -902,7 +945,12 @@ class HierarchicalNSW : public AlgorithmInterface { } return OkStatus(); - }); +#if defined(__EXCEPTIONS) || _HAS_EXCEPTIONS == 1 + } catch (const std::bad_alloc&) { + clear(); + return Status("Not enough memory: loadIndex failed to allocate"); + } +#endif } @@ -1343,6 +1391,14 @@ class HierarchicalNSW : public AlgorithmInterface { if (curlevel) { linkLists_[cur_c] = (char *) malloc(size_links_per_element_ * curlevel + 1); if (linkLists_[cur_c] == nullptr) { + // Upper-level list failed. Never leave element_levels_ > 0 with + // a null linkLists_[cur_c] (saveIndex would write from nullptr). + element_levels_[cur_c] = 0; + std::unique_lock lock_table(label_lookup_lock); + if (cur_element_count == cur_c + 1) { + label_lookup_.erase(label); + cur_element_count--; + } return Status("Not enough memory: addPoint failed to allocate linklist"); } memset(linkLists_[cur_c], 0, size_links_per_element_ * curlevel + 1); diff --git a/hnswlib/hnswlib.h b/hnswlib/hnswlib.h index ff36772b..219dc3d1 100644 --- a/hnswlib/hnswlib.h +++ b/hnswlib/hnswlib.h @@ -125,8 +125,8 @@ static bool AVX512Capable() { #include #include #include -#include #include +#include #include #include #include @@ -152,23 +152,30 @@ static bool AVX512Capable() { namespace hnswlib { -// A lightweight Status class inspired by Abseil's Status class. +// Lightweight Status. Empty message is OK. The error text is copied, so +// callers may pass a stack buffer or a temporary std::string.c_str(). +// Copying may heap-allocate (typical messages exceed SSO). With +// -fno-exceptions that allocate is fatal; pass string literals. class HNSWLIB_NODISCARD Status { public: - Status() : message_(nullptr) {} + Status() {} - // Constructor with an error message (nullptr is interpreted as OK status). - Status(const char* message) : message_(message) {} + // nullptr is interpreted as OK status. + Status(const char* message) { + if (message != nullptr) { + message_.assign(message); + } + } + + Status(std::string message) : message_(std::move(message)) {} - // Returns true if the status is OK. - bool ok() const { return !message_; } + bool ok() const { return message_.empty(); } - // Returns the error message, or nullptr if OK. - const char* message() const { return message_; } + // nullptr if OK. Pointer is valid for the lifetime of *this. + const char* message() const { return ok() ? nullptr : message_.c_str(); } private: - // nullptr if OK, a message otherwise. - const char* message_; + std::string message_; }; inline Status OkStatus() { return Status(); } @@ -311,41 +318,6 @@ static void readBinaryPOD(std::istream &in, T &podRef) { in.read((char *) &podRef, sizeof(T)); } -// Temporarily disable iostream exceptions so *NoExceptions I/O reports -// failures via Status instead of throwing std::ios_base::failure. -class StreamExceptionsOff { - public: - explicit StreamExceptionsOff(std::ios& stream) - : stream_(stream), old_(stream.exceptions()) { - stream_.exceptions(std::ios::goodbit); - } - - ~StreamExceptionsOff() { - // Restoring a mask that includes failbit/badbit throws if those bits - // are already set. Only restore when that would be safe. - if (!stream_.fail()) { - stream_.exceptions(old_); - } - } - - private: - std::ios& stream_; - std::ios::iostate old_; -}; - -template -Status invokeWithoutStreamThrow(Fn&& fn) { -#if defined(__EXCEPTIONS) || _HAS_EXCEPTIONS == 1 - try { - return fn(); - } catch (const std::ios_base::failure&) { - return Status("Stream I/O failed"); - } -#else - return fn(); -#endif -} - template using DISTFUNC = MTYPE(*)(const void *, const void *, const void *); diff --git a/tests/cpp/no_exceptions_api_test.cpp b/tests/cpp/no_exceptions_api_test.cpp index dec3522a..4d0b7cda 100644 --- a/tests/cpp/no_exceptions_api_test.cpp +++ b/tests/cpp/no_exceptions_api_test.cpp @@ -1,8 +1,9 @@ #include #include +#include #include -#include #include +#include #include #include #include @@ -91,13 +92,6 @@ void testSaveIndexNoExceptionsDoesNotThrowOnWriteFailure() { FailingBuf buf; std::ostream out(&buf); -#if defined(__EXCEPTIONS) || _HAS_EXCEPTIONS == 1 - // The review case: write failure with iostream exceptions enabled must - // still return Status, not throw. Skipped when the TU is built with - // -fno-exceptions because enabling the mask would abort instead. - out.exceptions(std::ios::failbit | std::ios::badbit); -#endif - hnswlib::Status status = index.saveIndexNoExceptions(out); assert(!status.ok()); } @@ -118,6 +112,169 @@ void testAddPointIntegerLevelIsNotReplaceDeleted() { assert(index.getCurrentElementCount() == 2); } +void testReloadDropsStaleLabelsAndDeletedSet() { + const int dim = 4; + std::vector a(dim, 1.0f); + std::vector b(dim, 2.0f); + std::vector c(dim, 3.0f); + + hnswlib::L2Space space(dim); + hnswlib::HierarchicalNSW live( + &space, /*max_elements=*/8, /*M=*/16, /*ef_construction=*/16, + /*random_seed=*/100, /*allow_replace_deleted=*/true); + assert(live.addPointNoExceptions(a.data(), 1).ok()); + assert(live.addPointNoExceptions(b.data(), 2).ok()); + live.markDelete(1); + assert(live.getDeletedCount() == 1); + + hnswlib::HierarchicalNSW replacement(&space, 8); + assert(replacement.addPointNoExceptions(c.data(), 10).ok()); + std::ostringstream saved(std::ios::binary); + assert(replacement.saveIndexNoExceptions(saved).ok()); + + std::istringstream in(saved.str(), std::ios::binary); + assert(live.loadIndexNoExceptions(in, &space).ok()); + assert(live.getCurrentElementCount() == 1); + assert(live.getDeletedCount() == 0); + + auto missing = live.getDataByLabelNoExceptions(2); + assert(!missing.ok()); + std::vector got = live.getDataByLabel(10); + assert(got.size() == static_cast(dim)); + assert(got[0] == 3.0f); +} + +void testCorruptLoadDoesNotClearLiveIndex() { + const int dim = 4; + std::vector a(dim, 1.0f); + + hnswlib::L2Space space(dim); + hnswlib::HierarchicalNSW index(&space, 8); + assert(index.addPointNoExceptions(a.data(), 1).ok()); + + std::istringstream truncated("HNSW", std::ios::binary); + hnswlib::Status status = index.loadIndexNoExceptions(truncated, &space); + assert(!status.ok()); + assert(index.getCurrentElementCount() == 1); + std::vector restored = index.getDataByLabel(1); + assert(restored[0] == 1.0f); +} + +void testBruteforceLoadRebuildsLabelMap() { + const int dim = 4; + std::vector a(dim, 1.0f); + std::vector b(dim, 2.0f); + std::vector a2(dim, 9.0f); + + hnswlib::L2Space space(dim); + hnswlib::BruteforceSearch index(&space, 8); + assert(index.addPointNoExceptions(a.data(), 10).ok()); + assert(index.addPointNoExceptions(b.data(), 20).ok()); + assert(index.cur_element_count == 2); + + std::ostringstream saved(std::ios::binary); + assert(index.saveIndexNoExceptions(saved).ok()); + + hnswlib::BruteforceSearch loaded(&space, 8); + std::istringstream in(saved.str(), std::ios::binary); + assert(loaded.loadIndexNoExceptions(in, &space).ok()); + assert(loaded.cur_element_count == 2); + assert(loaded.addPointNoExceptions(a2.data(), 10).ok()); + assert(loaded.cur_element_count == 2); +} + +void testBruteforceLoadIndexNoExceptionsDoesNotMutateOnFailure() { + const int dim = 4; + std::vector a(dim, 1.0f); + + hnswlib::L2Space space(dim); + hnswlib::BruteforceSearch index(&space, 8); + assert(index.addPointNoExceptions(a.data(), 10).ok()); + assert(index.cur_element_count == 1); + + std::ifstream missing("hnswlib_rc_review_missing_bf.bin", std::ios::binary); + assert(!missing.is_open()); + assert(!index.loadIndexNoExceptions(missing, &space).ok()); + assert(index.cur_element_count == 1); + + std::ifstream never_opened; + assert(!index.loadIndexNoExceptions(never_opened, &space).ok()); + assert(index.cur_element_count == 1); + + std::istringstream truncated("HNSW", std::ios::binary); + assert(!index.loadIndexNoExceptions(truncated, &space).ok()); + assert(index.cur_element_count == 1); +} + +void testBruteforceLoadRejectsSizePerElementMismatch() { + const int dim = 4; + std::vector a(dim, 1.0f); + + hnswlib::L2Space space4(dim); + hnswlib::BruteforceSearch index(&space4, 8); + assert(index.addPointNoExceptions(a.data(), 10).ok()); + std::ostringstream saved(std::ios::binary); + assert(index.saveIndexNoExceptions(saved).ok()); + + hnswlib::L2Space space8(8); + hnswlib::BruteforceSearch loaded(&space8, 8); + std::istringstream in(saved.str(), std::ios::binary); + assert(!loaded.loadIndexNoExceptions(in, &space8).ok()); + assert(loaded.cur_element_count == 0); +} + +void testLoadRejectsCurCountGreaterThanMaxElements() { + const int dim = 4; + std::vector a(dim, 1.0f); + std::vector b(dim, 2.0f); + + hnswlib::L2Space space(dim); + hnswlib::HierarchicalNSW saved(&space, 8); + assert(saved.addPointNoExceptions(a.data(), 1).ok()); + assert(saved.addPointNoExceptions(b.data(), 2).ok()); + std::ostringstream out(std::ios::binary); + assert(saved.saveIndexNoExceptions(out).ok()); + std::string blob = out.str(); + assert(blob.size() >= 2 * sizeof(size_t)); + size_t patched_max = 1; + std::memcpy(&blob[sizeof(size_t)], &patched_max, sizeof(size_t)); + + hnswlib::HierarchicalNSW live(&space, 8); + assert(live.addPointNoExceptions(a.data(), 9).ok()); + std::istringstream in(blob, std::ios::binary); + assert(!live.loadIndexNoExceptions(in, &space).ok()); + assert(live.getCurrentElementCount() == 1); +} + +void testClearResetsCapacitySoAddPointDoesNotWriteNull() { + const int dim = 4; + std::vector a(dim, 1.0f); + + hnswlib::L2Space space(dim); + hnswlib::HierarchicalNSW index(&space, 8); + assert(index.addPointNoExceptions(a.data(), 1).ok()); + index.clear(); + hnswlib::Status status = index.addPointNoExceptions(a.data(), 1); + assert(!status.ok()); +} + +void testStatusCopiesStackMessage() { + char buf[32]; + std::snprintf(buf, sizeof(buf), "stack-%d", 7); + hnswlib::Status st(buf); + buf[0] = 'X'; + assert(!st.ok()); + assert(std::string(st.message()) == "stack-7"); + + hnswlib::Status ok; + assert(ok.ok()); + assert(ok.message() == nullptr); + + hnswlib::Status from_string(std::string("owned")); + assert(!from_string.ok()); + assert(std::string(from_string.message()) == "owned"); +} + void testSearchKnnCloserFirstIsConst() { const int dim = 4; std::vector a(dim, 1.0f); @@ -138,7 +295,15 @@ int main() { testBruteforceSaveIndexNoExceptionsDoesNotThrow(); testLoadIndexNoExceptionsDoesNotClearOnUnopenedStream(); testSaveIndexNoExceptionsDoesNotThrowOnWriteFailure(); + testReloadDropsStaleLabelsAndDeletedSet(); + testCorruptLoadDoesNotClearLiveIndex(); + testBruteforceLoadRebuildsLabelMap(); + testBruteforceLoadIndexNoExceptionsDoesNotMutateOnFailure(); testAddPointIntegerLevelIsNotReplaceDeleted(); + testBruteforceLoadRejectsSizePerElementMismatch(); + testLoadRejectsCurCountGreaterThanMaxElements(); + testClearResetsCapacitySoAddPointDoesNotWriteNull(); + testStatusCopiesStackMessage(); testSearchKnnCloserFirstIsConst(); return 0; }