diff --git a/bindings/c/CMakeLists.txt b/bindings/c/CMakeLists.txt index ca7f4d0ca..334b84706 100644 --- a/bindings/c/CMakeLists.txt +++ b/bindings/c/CMakeLists.txt @@ -40,6 +40,7 @@ set(SVS_C_API_SOURCES src/index_builder.hpp src/leanvec_training_data.hpp src/storage.hpp + src/stream.hpp src/threadpool.hpp src/types_support.hpp @@ -144,7 +145,7 @@ if (SVS_RUNTIME_ENABLE_LVQ_LEANVEC) else() # Links to LTO-enabled static library, requires GCC/G++ 11.2 if(CMAKE_CXX_COMPILER_ID STREQUAL "GNU" AND CMAKE_CXX_COMPILER_VERSION VERSION_GREATER_EQUAL "11.2" AND CMAKE_CXX_COMPILER_VERSION VERSION_LESS "11.3") - set(SVS_URL "https://github.com/intel/ScalableVectorSearch/releases/download/nightly/svs-shared-library-lto-nightly-2026-10-06-1613.tar.gz" + set(SVS_URL "https://github.com/intel/ScalableVectorSearch/releases/download/nightly/svs-shared-library-lto-nightly-2026-10-08-1629.tar.gz" CACHE STRING "URL to download SVS shared library") else() # The fallback is correct but slower, so nothing downstream fails and CI diff --git a/bindings/c/README.md b/bindings/c/README.md index 72af1fe29..f4d9fce3b 100644 --- a/bindings/c/README.md +++ b/bindings/c/README.md @@ -22,8 +22,8 @@ C applications and any language with C FFI support. The API is built around a small set of opaque handles and a builder pattern: configure an *algorithm*, optional *storage* and *thread pool*, hand them to an *index builder*, then use the resulting *index* to run TopK searches (with -optional ID filtering), save/load the index, and — for dynamic indices — add or -delete points at runtime. +optional ID filtering), save/load the index to disk or a caller-supplied stream, +and — for dynamic indices — add or delete points at runtime. For the design rationale, naming conventions, and full API reference see [docs/C_API_Design.md](docs/C_API_Design.md). @@ -181,16 +181,16 @@ cleanup: ## Samples -Runnable sample applications live in [samples/](samples/): +Runnable sample applications live in [`examples/c/`](../../examples/c/): -- [`simple.c`](samples/simple.c) – minimal static index build + search with a +- [`simple.c`](../../examples/c/simple.c) – minimal static index build + search with a custom thread pool -- [`dynamic.c`](samples/dynamic.c) – dynamic index with add / delete / +- [`dynamic.c`](../../examples/c/dynamic.c) – dynamic index with add / delete / consolidate -- [`save_load.c`](samples/save_load.c) – persisting and reloading indices from +- [`save_load.c`](../../examples/c/save_load.c) – persisting and reloading indices from disk - -Additional integration examples: [`examples/c/`](../../examples/c/). +- [`save_load_stream.c`](../../examples/c/save_load_stream.c) – stream-based index save + and load ## Further Reading diff --git a/bindings/c/docs/C_API_Design.md b/bindings/c/docs/C_API_Design.md index 597c120e9..b5b6950c5 100644 --- a/bindings/c/docs/C_API_Design.md +++ b/bindings/c/docs/C_API_Design.md @@ -44,6 +44,7 @@ - [6. Allocator Configuration](#6-allocator-configuration) - [7. Search Parameters](#7-search-parameters) - [8. ID Filter (optional)](#8-id-filter-optional) + - [9. Stream Interface](#9-stream-interface) - [API Overview](#api-overview) - [Headers](#headers) - [Types](#types) @@ -198,6 +199,10 @@ one of the two optional slots is used per name: (e.g. `svs_algorithm_create_vamana` creates a Vamana algorithm; `svs_index_build_dynamic` builds a dynamic index). +When both specializations apply (e.g. a save or load variant), `_stream` and +`_stream_dynamic` are stacked at the end: `svs_index_save_stream`, +`svs_index_load_stream_dynamic`. + **Examples:** | Function | Breakdown | Description | @@ -209,6 +214,8 @@ one of the two optional slots is used per name: | `svs_index_build_dynamic()` | `svs` + `index` + `build` + `dynamic` | Build a dynamic index | | `svs_index_dynamic_add_points()` | `svs` + `index` + `dynamic` + `add_points` | Add points to a dynamic index | | `svs_index_builder_set_threadpool()` | `svs` + `index_builder` + `set_threadpool` | Configure builder thread pool | +| `svs_index_save_stream()` | `svs` + `index` + `save` + `stream` | Save index to a caller-supplied stream | +| `svs_index_load_stream_dynamic()` | `svs` + `index` + `load` + `stream_dynamic` | Load a dynamic index from a stream | ### Examples by Category @@ -225,6 +232,7 @@ typedef enum svs_error_code svs_error_code_t; // Interface pointer types typedef struct svs_threadpool_interface* svs_threadpool_i; typedef struct svs_id_filter_interface* svs_id_filter_i; +typedef struct svs_stream_interface* svs_stream_i; ``` ## Core Components @@ -500,6 +508,64 @@ Providing a non-zero `filter_rate` lets the search account for the expected selectivity; if the observed hit rate ends up lower than the reported estimate the function returns an empty result set for that query. +### 9. Stream Interface + +Enables caller-supplied byte streams for index save and load operations, eliminating +the need for intermediate disk storage. Like the thread pool and allocator, the stream +interface is a versioned ops table plus an opaque `self` pointer. + +```c +struct svs_stream_interface_ops { + uint32_t version; // Set by SVS_INIT_STREAM_OPS + size_t struct_size; // Set by SVS_INIT_STREAM_OPS + size_t (*read)(void* self, void* buf, size_t n, svs_error_h out_err); + bool (*write)(void* self, const void* buf, size_t n, svs_error_h out_err); +}; + +struct svs_stream_interface { + struct svs_stream_interface_ops* ops; + void* self; // User-defined state +}; +typedef struct svs_stream_interface* svs_stream_i; + +// Initialisation macros +static svs_stream_ops_t my_stream_ops = + SVS_INIT_STREAM_OPS(my_read_func, my_write_func); +static svs_stream_t my_stream = SVS_MAKE_INTERFACE(user_state, my_stream_ops); +``` + +**Callback contracts:** + +- **`read`** — Read at most `n` bytes into `buf`. Return the number of bytes read; return + 0 to signal end of stream. Short reads are not errors; the library will call again as + needed. To report a read error, set `out_err` via `svs_error_set()` and return 0 — this + aborts the load with that error code, unlike returning 0 with no error set, which is a + clean end of stream. This differs from `write`, which signals failure through its return + value. Required for load operations; may be NULL for write-only streams. +- **`write`** — Write exactly `n` bytes from `buf`. Return `true` on success. A partial + write must be reported as failure. Required for save operations; may be NULL for + read-only streams. + +**Threading:** Both callbacks are invoked serially from the thread that called the streaming +save or load function. No synchronization between concurrent stream operations is required. + +**Lifetime:** The operations table is copied by value; `self` is retained as a bare pointer +and only needs to remain valid until the streaming function returns. Data is copied out of +the stream during load, so the stream buffer need not persist after the call completes. + +**Encodings:** SVS supports two mutually exclusive stream encodings identified by an 8-byte magic +at offset zero: the native stream encoding (`"SVS_STRM"`) and a tar-like directory archive. The load +functions accept both transparently; detection is handled by SVS internals. `svs_index_save_stream` +produces only the native encoding by design, so the save and load halves are deliberately asymmetric. +An index written to disk via `svs_index_save` cannot be streamed, because the C layer does not +expose the machinery to pack a directory archive into stream form. Streaming is therefore +self-sufficient only for indexes that were themselves stream-saved. + +**Read-ahead behavior:** The native stream encoding carries no total length. During load, the +library may read up to 64 KiB past the logical end of the index data, so callers embedding an +index inside a larger stream must frame the payload themselves (e.g. with a length prefix) and +bound the stream reads to that frame. + ## API Overview A concise map of the public surface. See [svs/c/svs_c.h](../include/svs/c/svs_c.h) @@ -518,8 +584,8 @@ for full signatures, parameters, and Doxygen documentation. - **Enums** (`_t`): `svs_error_code_t`, `svs_distance_metric_t`, `svs_algorithm_type_t`, `svs_data_type_t`, `svs_storage_kind_t`, `svs_threadpool_kind_t`, `svs_allocator_kind_t` -- **Custom interfaces**: `svs_threadpool_i`, `svs_allocator_i`, and `svs_id_filter_i` - (versioned ops-table + `self` pointer; build with `SVS_INIT_*_OPS()` / +- **Custom interfaces**: `svs_threadpool_i`, `svs_allocator_i`, `svs_id_filter_i`, and + `svs_stream_i` (versioned ops-table + `self` pointer; build with `SVS_INIT_*_OPS()` / `SVS_MAKE_INTERFACE()`) - **Value structs**: `svs_search_results_t` (CSR result buffer), `svs_memory_breakdown_t` @@ -535,7 +601,7 @@ for full signatures, parameters, and Doxygen documentation. | **Search params** | `svs_search_params_create_vamana`, `svs_search_params_free` | | **Builder** | `svs_index_builder_create`, `svs_index_builder_set_{storage,threadpool,threadpool_custom,allocator,allocator_custom}`, `svs_index_builder_free` | | **Memory estimation** | `svs_index_builder_estimate_memory`, `svs_index_builder_estimate_memory_dynamic`, `svs_index_builder_estimate_search_memory`, `svs_index_builder_estimate_search_memory_dynamic`, `svs_index_builder_get_default_blocksize_bytes` | -| **Index lifecycle** | `svs_index_build`, `svs_index_build_dynamic`, `svs_index_load`, `svs_index_load_dynamic`, `svs_index_save`, `svs_index_free` | +| **Index lifecycle** | `svs_index_build`, `svs_index_build_dynamic`, `svs_index_load`, `svs_index_load_dynamic`, `svs_index_load_stream`, `svs_index_load_stream_dynamic`, `svs_index_save`, `svs_index_save_stream`, `svs_index_free` | | **Dynamic ops** | `svs_index_dynamic_{add_points,delete_points,has_id,consolidate,compact}` | | **Introspection** | `svs_index_get_num_threads` / `set_num_threads`, `svs_index_get_distance`, `svs_index_reconstruct`, `svs_index_get_memory_usage`, `svs_index_get_memory_breakdown` | | **Search** | `svs_index_search_topk` (+ deprecated `svs_index_search`), `svs_search_results_free` | @@ -558,8 +624,8 @@ for full signatures, parameters, and Doxygen documentation. - See the top-level [../README.md](../README.md) for a quick start, build/consume instructions, and a complete end-to-end usage example. -- See [../samples/](../samples/) for runnable sample applications: - - `simple.c` – minimal static index build + search with a custom thread pool - - `dynamic.c` – dynamic index with add / delete / consolidate - - `save_load.c` – persisting and reloading indices from disk -- See [examples/c/](../../../examples/c/) for additional usage examples +- See [examples/c/](../../../examples/c/) for runnable sample applications: + - [`simple.c`](../../../examples/c/simple.c) – minimal static index build + search with a custom thread pool + - [`dynamic.c`](../../../examples/c/dynamic.c) – dynamic index with add / delete / consolidate + - [`save_load.c`](../../../examples/c/save_load.c) – persisting and reloading indices from disk + - [`save_load_stream.c`](../../../examples/c/save_load_stream.c) – stream-based index save and load diff --git a/bindings/c/include/svs/c/svs_c.h b/bindings/c/include/svs/c/svs_c.h index 224bf167c..f9c23f54e 100644 --- a/bindings/c/include/svs/c/svs_c.h +++ b/bindings/c/include/svs/c/svs_c.h @@ -294,6 +294,65 @@ struct svs_id_filter_interface { void* self; }; +/// @brief Operations table for a caller-supplied byte stream. +/// @remarks Access is strictly sequential: the library never repositions the stream. Both +/// callbacks are invoked serially from the thread that called the streaming save or load +/// function, so no synchronization is required +/// @remarks Exactly one direction is required per operation: @ref svs_index_save_stream +/// needs @p write, the load functions need @p read. The unused callback may be NULL. +/// @var svs_stream_interface_ops::version +/// Interface version, set by @ref SVS_INIT_STREAM_OPS. +/// @var svs_stream_interface_ops::struct_size +/// Size of this structure, set by @ref SVS_INIT_STREAM_OPS. +/// @var svs_stream_interface_ops::read +/// Reads at most @p n bytes into @p buf. +/// @param self Pointer to the stream instance. +/// @param buf Destination buffer. +/// @param n Maximum number of bytes to read. +/// @param out_err Handle to capture any error that occurs during the read. Returning 0 +/// with an error set on @p out_err via svs_error_set() reports a failed read and aborts +/// the load with that error code; returning 0 without setting one is a clean end of +/// stream. This differs from @p write, which signals failure through its return value. +/// @return The number of bytes read; 0 signals end of stream unless @p out_err carries an +/// error. A short read is not an error and the library will call again. +/// @var svs_stream_interface_ops::write +/// Writes exactly @p n bytes from @p buf. +/// @param self Pointer to the stream instance. +/// @param buf Source buffer. +/// @param n Number of bytes to write. +/// @param out_err Handle to capture any error that occurs during the write. User code may +/// call svs_error_set() to set the error code and message if an error occurs. +/// @return True on success. A partial write must be reported as failure. +struct svs_stream_interface_ops { + uint32_t version; + size_t struct_size; + size_t (*read)(void* self, void* buf, size_t n, svs_error_h out_err); + bool (*write)(void* self, const void* buf, size_t n, svs_error_h out_err); +}; + +/// @brief Macro to create a user-defined stream interface operations structure +/// @param read_func Function pointer that reads at most @p n bytes into @p buf, or NULL for +/// a write-only stream +/// @param write_func Function pointer that writes exactly @p n bytes from @p buf, or NULL +/// for a read-only stream +#define SVS_INIT_STREAM_OPS(read_func, write_func) \ + { \ + .version = SVS_C_API_VERSION, \ + .struct_size = sizeof(struct svs_stream_interface_ops), .read = (read_func), \ + .write = (write_func) \ + } + +/// @brief Structure representing a caller-supplied byte stream +/// @var svs_stream_interface::ops +/// Function pointers for the stream operations. +/// @var svs_stream_interface::self +/// Pointer to the user-defined stream instance. This pointer is passed to the function +/// pointers in @p ops when they are called. +struct svs_stream_interface { + struct svs_stream_interface_ops* ops; + void* self; +}; + /// @brief Macro to create a user-defined interface implementation structure /// @param user_ptr Pointer to the user-defined object /// @param vtable Function pointers for the interface operations @@ -498,6 +557,10 @@ typedef struct svs_id_filter_interface_ops svs_id_filter_ops_t; typedef struct svs_id_filter_interface svs_id_filter_t; typedef struct svs_id_filter_interface* svs_id_filter_i; +typedef struct svs_stream_interface_ops svs_stream_ops_t; +typedef struct svs_stream_interface svs_stream_t; +typedef struct svs_stream_interface* svs_stream_i; + typedef struct svs_search_results svs_search_results_t; typedef struct svs_memory_breakdown svs_memory_breakdown_t; @@ -996,6 +1059,44 @@ SVS_API svs_index_h svs_index_load_dynamic( svs_error_h out_err /*=NULL*/ ); +/// @brief Load an index from a caller-supplied stream +/// @param builder The index builder handle (used for configuration) +/// @param stream The stream interface to read the index from +/// @param out_err An optional error handle to capture errors +/// @return A handle to the loaded index +/// @remarks The operations table is copied, but @p stream->self is retained as-is. It only +/// needs to remain valid until this function returns, because index data is copied out of +/// the stream rather than referenced. +/// @remarks Accepts both the native stream encoding produced by @ref svs_index_save_stream +/// and a packed directory archive. The encoding is detected from the stream itself. +/// @remarks The native encoding carries no total length, so a load may read up to 64 KiB +/// past the end of the index. Callers embedding it in a larger stream must frame the +/// payload (e.g. with a length prefix) and bound reads to that frame. +SVS_API svs_index_h svs_index_load_stream( + svs_index_builder_h builder, svs_stream_i stream, svs_error_h out_err /*=NULL*/ +); + +/// @brief Load a dynamic index from a caller-supplied stream +/// @param builder The index builder handle (used for configuration) +/// @param stream The stream interface to read the index from +/// @param blocksize_bytes The block size in bytes for dynamic index loading (0 for default) +/// @param out_err An optional error handle to capture errors +/// @return A handle to the loaded dynamic index +/// @remarks The operations table is copied, but @p stream->self is retained as-is. It only +/// needs to remain valid until this function returns, because index data is copied out of +/// the stream rather than referenced. +/// @remarks Accepts both the native stream encoding produced by @ref svs_index_save_stream +/// and a packed directory archive. The encoding is detected from the stream itself. +/// @remarks The native encoding carries no total length, so a load may read up to 64 KiB +/// past the end of the index. Callers embedding it in a larger stream must frame the +/// payload (e.g. with a length prefix) and bound reads to that frame. +SVS_API svs_index_h svs_index_load_stream_dynamic( + svs_index_builder_h builder, + svs_stream_i stream, + size_t blocksize_bytes /*=0*/, + svs_error_h out_err /*=NULL*/ +); + /// @brief Convert an index using new builder configuration /// @param builder The index builder handle (used for configuration) /// @param src_index The source index handle to convert from @@ -1111,6 +1212,19 @@ static inline bool svs_index_search( SVS_API bool svs_index_save(svs_index_h index, const char* directory, svs_error_h out_err /*=NULL*/); +/// @brief Save the index to a caller-supplied stream +/// @param index The index handle +/// @param stream The stream interface to write the index to +/// @param out_err An optional error handle to capture errors +/// @return true on success, false on failure +/// @remarks The operations table is copied, but @p stream->self is retained as-is. It only +/// needs to remain valid until this function returns. +/// @remarks Produces the native stream encoding only. An index previously written with +/// @ref svs_index_save cannot be converted to a stream through this API. +SVS_API bool svs_index_save_stream( + svs_index_h index, svs_stream_i stream, svs_error_h out_err /*=NULL*/ +); + /// @brief Add points to a dynamic index /// @param index The dynamic index handle /// @param new_points Pointer to the new vector data (float array) diff --git a/bindings/c/src/dispatcher_dynamic_vamana.cpp b/bindings/c/src/dispatcher_dynamic_vamana.cpp index 3fa616ad4..993c751e9 100644 --- a/bindings/c/src/dispatcher_dynamic_vamana.cpp +++ b/bindings/c/src/dispatcher_dynamic_vamana.cpp @@ -31,6 +31,7 @@ #include #include +#include #include #include #include @@ -104,6 +105,36 @@ svs::DynamicVamana load_dynamic_vamana_index( ); } +template +svs::DynamicVamana load_stream_dynamic_vamana_index( + const svs::index::vamana::VamanaBuildParameters& SVS_UNUSED(build_params), + std::unique_ptr stream, + DataLoader SVS_UNUSED(loader), + Distance distance, + svs::threads::ThreadPoolHandle pool, + const AllocatorBuilder& allocator_builder, + size_t blocksize_bytes +) { + svs::data::BlockingParameters block_params; + if (blocksize_bytes != 0) { + block_params.blocksize_bytes = svs::lib::prevpow2(blocksize_bytes); + } + using allocator_type = typename DataLoader::allocator_type; + using value_type = typename allocator_type::value_type; + using data_type = typename DataLoader::data_type; + auto data_allocator_handle = allocator_builder.build(); + auto allocator = allocator_type{block_params, data_allocator_handle}; + + auto graph_allocator_handle = allocator_builder.build_for_graph(); + auto graph_allocator = svs::data::Blocked{block_params, graph_allocator_handle}; + + // svs_c.h lets the caller drop the stream once loading returns. That holds only while + // assemble copies data out; a view allocator here would leave the index dangling. + return svs::DynamicVamana::assemble( + *stream, distance, std::move(pool), allocator, graph_allocator + ); +} + template void register_dynamic_vamana_index_specializations(Dispatcher& dispatcher) { auto build_closure = [&dispatcher]() { @@ -112,20 +143,30 @@ void register_dynamic_vamana_index_specializations(Dispatcher& dispatcher) { auto load_closure = [&dispatcher]() { dispatcher.register_target(&load_dynamic_vamana_index); }; + auto load_stream_closure = [&dispatcher]() { + dispatcher.register_target(&load_stream_dynamic_vamana_index); + }; for_simple_specializations(build_closure); for_simple_specializations(load_closure); + for_simple_specializations(load_stream_closure); for_leanvec_specializations(build_closure); for_leanvec_specializations(load_closure); + for_leanvec_specializations(load_stream_closure); for_lvq_specializations(build_closure); for_lvq_specializations(load_closure); + for_lvq_specializations(load_stream_closure); for_sq_specializations(build_closure); for_sq_specializations(load_closure); + for_sq_specializations(load_stream_closure); } +// Stream load alternative, matched via generic variant DispatchConverter like the +// existing build and directory-load alternatives. using DynamicVamanaSource = std::variant< std::pair, std::span>, - std::filesystem::path>; + std::filesystem::path, + std::unique_ptr>; using BuildDynamicIndexDispatcher = svs::lib::Dispatcher< svs::DynamicVamana, @@ -334,6 +375,26 @@ svs::DynamicVamana dispatch_dynamic_vamana_index_load( ); } +svs::DynamicVamana dispatch_dynamic_vamana_index_load_stream( + const svs::index::vamana::VamanaBuildParameters& build_params, + std::unique_ptr stream, + const Storage* storage, + svs::DistanceType distance_type, + svs::threads::ThreadPoolHandle pool, + const AllocatorBuilder& allocator_builder, + size_t blocksize_bytes +) { + return build_dynamic_vamana_index_dispatcher().invoke( + build_params, + DynamicVamanaSource{std::move(stream)}, + storage, + distance_type, + std::move(pool), + allocator_builder, + blocksize_bytes + ); +} + svs::DynamicVamana dispatch_dynamic_vamana_index_copy( const svs::index::vamana::VamanaBuildParameters& build_params, const svs::DynamicVamana& src_index, diff --git a/bindings/c/src/dispatcher_dynamic_vamana.hpp b/bindings/c/src/dispatcher_dynamic_vamana.hpp index 5a96c62d9..afaeb366d 100644 --- a/bindings/c/src/dispatcher_dynamic_vamana.hpp +++ b/bindings/c/src/dispatcher_dynamic_vamana.hpp @@ -24,6 +24,8 @@ #include #include +#include +#include #include #include #include @@ -51,6 +53,16 @@ svs::DynamicVamana dispatch_dynamic_vamana_index_load( size_t blocksize_bytes ); +svs::DynamicVamana dispatch_dynamic_vamana_index_load_stream( + const svs::index::vamana::VamanaBuildParameters& build_params, + std::unique_ptr stream, + const Storage* storage, + svs::DistanceType distance_type, + svs::threads::ThreadPoolHandle pool, + const AllocatorBuilder& allocator_builder, + size_t blocksize_bytes +); + svs::DynamicVamana dispatch_dynamic_vamana_index_copy( const svs::index::vamana::VamanaBuildParameters& build_params, const svs::DynamicVamana& src_index, diff --git a/bindings/c/src/dispatcher_vamana.cpp b/bindings/c/src/dispatcher_vamana.cpp index 35f070718..e2414d95e 100644 --- a/bindings/c/src/dispatcher_vamana.cpp +++ b/bindings/c/src/dispatcher_vamana.cpp @@ -30,6 +30,7 @@ #include #include +#include #include #include #include @@ -80,6 +81,28 @@ svs::Vamana load_vamana_index( ); } +template +svs::Vamana load_stream_vamana_index( + const svs::index::vamana::VamanaBuildParameters& SVS_UNUSED(build_params), + std::unique_ptr stream, + DataLoader SVS_UNUSED(loader), + Distance distance, + svs::threads::ThreadPoolHandle pool, + const AllocatorBuilder& allocator_builder +) { + using value_type = typename DataLoader::allocator_type::value_type; + using data_type = typename DataLoader::data_type; + // svs_c.h lets the caller drop the stream once loading returns. That holds only while + // assemble copies data out; a view allocator here would leave the index dangling. + return svs::Vamana::assemble( + *stream, + distance, + std::move(pool), + allocator_builder.build(), + allocator_builder.build_for_graph() + ); +} + template void register_vamana_index_specializations(Dispatcher& dispatcher) { auto build_closure = [&dispatcher]() { @@ -88,19 +111,30 @@ void register_vamana_index_specializations(Dispatcher& dispatcher) { auto load_closure = [&dispatcher]() { dispatcher.register_target(&load_vamana_index); }; + auto load_stream_closure = [&dispatcher]() { + dispatcher.register_target(&load_stream_vamana_index); + }; for_simple_specializations(build_closure); for_simple_specializations(load_closure); + for_simple_specializations(load_stream_closure); for_leanvec_specializations(build_closure); for_leanvec_specializations(load_closure); + for_leanvec_specializations(load_stream_closure); for_lvq_specializations(build_closure); for_lvq_specializations(load_closure); + for_lvq_specializations(load_stream_closure); for_sq_specializations(build_closure); for_sq_specializations(load_closure); + for_sq_specializations(load_stream_closure); } -using VamanaSource = - std::variant, std::filesystem::path>; +// Stream load alternative, matched via generic variant DispatchConverter like the +// existing build and directory-load alternatives. +using VamanaSource = std::variant< + svs::data::ConstSimpleDataView, + std::filesystem::path, + std::unique_ptr>; using BuildIndexDispatcher = svs::lib::Dispatcher< svs::Vamana, @@ -280,6 +314,24 @@ svs::Vamana dispatch_vamana_index_load( ); } +svs::Vamana dispatch_vamana_index_load_stream( + const svs::index::vamana::VamanaBuildParameters& build_params, + std::unique_ptr stream, + const Storage* storage, + svs::DistanceType distance_type, + svs::threads::ThreadPoolHandle pool, + const AllocatorBuilder& allocator_builder +) { + return build_vamana_index_dispatcher().invoke( + build_params, + VamanaSource{std::move(stream)}, + storage, + distance_type, + std::move(pool), + allocator_builder + ); +} + svs::Vamana dispatch_vamana_index_copy( const svs::index::vamana::VamanaBuildParameters& build_params, const svs::Vamana& src_index, diff --git a/bindings/c/src/dispatcher_vamana.hpp b/bindings/c/src/dispatcher_vamana.hpp index 133de017d..d3f497ecb 100644 --- a/bindings/c/src/dispatcher_vamana.hpp +++ b/bindings/c/src/dispatcher_vamana.hpp @@ -27,6 +27,8 @@ #include #include +#include +#include namespace svs::c_runtime { svs::Vamana dispatch_vamana_index_build( @@ -47,6 +49,15 @@ svs::Vamana dispatch_vamana_index_load( const AllocatorBuilder& allocator_builder ); +svs::Vamana dispatch_vamana_index_load_stream( + const svs::index::vamana::VamanaBuildParameters& build_params, + std::unique_ptr stream, + const Storage* storage, + svs::DistanceType distance_type, + svs::threads::ThreadPoolHandle pool, + const AllocatorBuilder& allocator_builder +); + svs::Vamana dispatch_vamana_index_copy( const svs::index::vamana::VamanaBuildParameters& build_params, const svs::Vamana& src_index, diff --git a/bindings/c/src/error.hpp b/bindings/c/src/error.hpp index f1508678c..67fb50f89 100644 --- a/bindings/c/src/error.hpp +++ b/bindings/c/src/error.hpp @@ -102,6 +102,20 @@ class out_of_memory : public std::runtime_error { using std::runtime_error::runtime_error; }; +// Carries the code a callback (e.g. a stream interface) reported, so wrap_exceptions +// can surface it verbatim instead of collapsing it to SVS_ERROR_RUNTIME. +class coded_error : public std::runtime_error { + public: + coded_error(svs_error_code_t code, const std::string& msg) + : std::runtime_error(msg) + , code_{code} {} + + svs_error_code_t code() const noexcept { return code_; } + + private: + svs_error_code_t code_; +}; + // A helper to wrap C++ exceptions and convert them to C error codes/messages. template > Result wrap_exceptions(Callable&& func, svs_error_h err, Result err_res = {}) noexcept { @@ -129,6 +143,10 @@ Result wrap_exceptions(Callable&& func, svs_error_h err, Result err_res = {}) no } catch (const std::bad_alloc& ex) { SET_ERROR(err, SVS_ERROR_OUT_OF_MEMORY, ex.what()); return err_res; + } catch (const svs::c_runtime::coded_error& ex) { + // Must precede std::runtime_error, its base class, or this clause is unreachable. + SET_ERROR(err, ex.code(), ex.what()); + return err_res; } catch (const std::runtime_error& ex) { SET_ERROR(err, SVS_ERROR_RUNTIME, ex.what()); return err_res; diff --git a/bindings/c/src/index.hpp b/bindings/c/src/index.hpp index c81e3cdea..dc5066db4 100644 --- a/bindings/c/src/index.hpp +++ b/bindings/c/src/index.hpp @@ -28,6 +28,7 @@ #include #include +#include #include #include #include @@ -48,6 +49,7 @@ struct Index { const IDFilterInterface* id_filter = nullptr ) = 0; virtual void save(const std::filesystem::path& directory) = 0; + virtual void save(std::ostream& stream) = 0; virtual size_t dimensions() const = 0; virtual float get_distance(size_t id, std::span query) const = 0; virtual void diff --git a/bindings/c/src/index_builder.cpp b/bindings/c/src/index_builder.cpp index fb8c4fee0..d543f8d92 100644 --- a/bindings/c/src/index_builder.cpp +++ b/bindings/c/src/index_builder.cpp @@ -38,6 +38,7 @@ #include #include +#include #include #include @@ -87,6 +88,27 @@ std::shared_ptr IndexBuilder::load(const std::filesystem::path& directory return nullptr; } +std::shared_ptr IndexBuilder::load_stream(std::unique_ptr&& stream) { + if (algorithm->type == SVS_ALGORITHM_TYPE_VAMANA) { + auto vamana_algorithm = static_cast(algorithm.get()); + + auto index = std::make_shared( + *this, + dispatch_vamana_index_load_stream( + vamana_algorithm->build_parameters(), + std::move(stream), + storage.get(), + to_distance_type(distance_metric), + pool_builder.build(), + allocator_builder + ) + ); + + return index; + } + return nullptr; +} + namespace { // Helper function to validate that two IndexBuilder instances are compatible void validate_builder_compatibility( @@ -249,6 +271,30 @@ IndexBuilder::load_dynamic(const std::filesystem::path& directory, size_t blocks return nullptr; } +std::shared_ptr IndexBuilder::load_stream_dynamic( + std::unique_ptr&& stream, size_t blocksize_bytes +) { + if (algorithm->type == SVS_ALGORITHM_TYPE_VAMANA) { + auto vamana_algorithm = static_cast(algorithm.get()); + + auto index = std::make_shared( + *this, + dispatch_dynamic_vamana_index_load_stream( + vamana_algorithm->build_parameters(), + std::move(stream), + storage.get(), + to_distance_type(distance_metric), + pool_builder.build(), + allocator_builder, + blocksize_bytes + ) + ); + + return index; + } + return nullptr; +} + svs::index::vamana::MemoryBreakdown IndexBuilder::estimate_memory_breakdown(size_t num_vectors) const { NOT_IMPLEMENTED_IF( diff --git a/bindings/c/src/index_builder.hpp b/bindings/c/src/index_builder.hpp index a6aedbfb4..c98ecca89 100644 --- a/bindings/c/src/index_builder.hpp +++ b/bindings/c/src/index_builder.hpp @@ -29,6 +29,7 @@ #include #include +#include #include #include #include @@ -98,6 +99,8 @@ struct IndexBuilder { std::shared_ptr load(const std::filesystem::path& directory); + std::shared_ptr load_stream(std::unique_ptr&& stream); + std::shared_ptr copy(const std::shared_ptr& src_index); std::shared_ptr build_dynamic( @@ -109,6 +112,9 @@ struct IndexBuilder { std::shared_ptr load_dynamic(const std::filesystem::path& directory, size_t blocksize_bytes); + std::shared_ptr + load_stream_dynamic(std::unique_ptr&& stream, size_t blocksize_bytes); + std::shared_ptr copy_dynamic(const std::shared_ptr& src_index, size_t blocksize_bytes); diff --git a/bindings/c/src/index_vamana.hpp b/bindings/c/src/index_vamana.hpp index 7bd97c743..c7c436cd3 100644 --- a/bindings/c/src/index_vamana.hpp +++ b/bindings/c/src/index_vamana.hpp @@ -26,6 +26,7 @@ #include #include +#include #include #include #include @@ -48,6 +49,8 @@ struct IndexVamana : public Index { index.save(directory / "config", directory / "graph", directory / "data"); } + void save(std::ostream& stream) override { index.save(stream); } + size_t dimensions() const override { return index.dimensions(); } float get_distance(size_t id, std::span query) const override { @@ -89,6 +92,8 @@ struct DynamicIndexVamana : public DynamicIndex { index.save(directory / "config", directory / "graph", directory / "data"); } + void save(std::ostream& stream) override { index.save(stream); } + size_t dimensions() const override { return index.dimensions(); } size_t add_points( diff --git a/bindings/c/src/stream.hpp b/bindings/c/src/stream.hpp new file mode 100644 index 000000000..bb6ca2f0d --- /dev/null +++ b/bindings/c/src/stream.hpp @@ -0,0 +1,198 @@ +/* + * Copyright 2026 Intel Corporation + * + * Licensed 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. + */ +#pragma once + +#include "svs/c/svs_c.h" + +#include "error.hpp" + +#include +#include +#include +#include +#include +#include +#include + +namespace svs::c_runtime { + +// Bridges svs_stream_i's read/write callbacks to std::streambuf. Buffered at 64 KiB +class StreamBuf : public std::streambuf { + public: + static constexpr size_t buffer_size = 64 * 1024; + enum class Direction { read, write }; + + static void validate(svs_stream_i stream, bool need_write) { + if (stream == nullptr) { + throw std::invalid_argument("Stream pointer cannot be null."); + } + if (stream->ops == nullptr) { + throw std::invalid_argument("Stream interface is not initialized."); + } + if (stream->ops->version > svs_get_version()) { + throw std::invalid_argument("Stream interface version is not supported."); + } + if (stream->ops->struct_size < sizeof(svs_stream_ops_t)) { + throw std::invalid_argument("Incompatible stream interface struct size."); + } + if (need_write) { + if (stream->ops->write == nullptr) { + throw std::invalid_argument("Stream interface has no write callback."); + } + } else { + if (stream->ops->read == nullptr) { + throw std::invalid_argument("Stream interface has no read callback."); + } + } + } + + // Holds a value copy of the user's ops table; `self` is referenced and only needs to + // remain valid until the streaming save/load call returns. + StreamBuf(const svs_stream_ops_t& ops, void* self, Direction direction) + : ops_(ops) + , self_(self) + , read_buf_(direction == Direction::read ? buffer_size : 0) + , write_buf_(direction == Direction::write ? buffer_size : 0) { + if (direction == Direction::write) { + setp(write_buf_.data(), write_buf_.data() + write_buf_.size()); + } + } + + StreamBuf(const StreamBuf&) = delete; + StreamBuf& operator=(const StreamBuf&) = delete; + StreamBuf(StreamBuf&&) = delete; + StreamBuf& operator=(StreamBuf&&) = delete; + + // Unflushed bytes are dropped on destruction, never delivered: a flush during unwinding + // would hand the callback a truncated chunk. Writers must flush explicitly. + ~StreamBuf() override = default; + + protected: + int_type overflow(int_type ch) override { + flush_write_buffer(); + if (!traits_type::eq_int_type(ch, traits_type::eof())) { + *pptr() = traits_type::to_char_type(ch); + pbump(1); + } + return traits_type::not_eof(ch); + } + + int sync() override { + flush_write_buffer(); + return 0; + } + + int_type underflow() override { + if (gptr() < egptr()) { + return traits_type::to_int_type(*gptr()); + } + // Default-initialized to SVS_OK: a 0-byte read is legitimate EOF unless the + // callback explicitly reported an error. + svs_error_desc impl_error{}; + size_t n = ops_.read(self_, read_buf_.data(), read_buf_.size(), &impl_error); + if (n > read_buf_.size()) { + throw std::invalid_argument( + "Stream read callback returned more bytes than the buffer it was given." + ); + } + if (impl_error.code != SVS_OK) { + throw coded_error( + impl_error.code, + "Stream read callback failed: (" + std::to_string(impl_error.code) + ") " + + impl_error.message + ); + } + if (n == 0) { + return traits_type::eof(); + } + setg(read_buf_.data(), read_buf_.data(), read_buf_.data() + n); + return traits_type::to_int_type(*gptr()); + } + + pos_type seekoff( + off_type off, std::ios_base::seekdir way, std::ios_base::openmode which + ) override { + if (off == 0 && way == std::ios_base::cur && which == std::ios_base::out) { + // tellp() must count bytes handed to the streambuf, not bytes flushed to the + // callback, or the format's cache-line padding misaligns silently. + return pos_type(static_cast(written_ + (pptr() - pbase()))); + } + return pos_type(off_type(-1)); + } + + private: + void flush_write_buffer() { + auto n = static_cast(pptr() - pbase()); + // Reset the put area before the callback: a throw then leaves it empty, so a + // retried flush cannot redeliver the same bytes twice. + setp(write_buf_.data(), write_buf_.data() + write_buf_.size()); + if (n > 0) { + svs_error_desc impl_error{ + SVS_ERROR_UNKNOWN, "Unknown error in stream write callback"}; + if (!ops_.write(self_, write_buf_.data(), n, &impl_error)) { + if (impl_error.code == SVS_OK) { + impl_error.code = SVS_ERROR_UNKNOWN; + } + throw coded_error( + impl_error.code, + "Stream write callback failed: (" + std::to_string(impl_error.code) + + ") " + impl_error.message + ); + } + written_ += n; + } + } + + svs_stream_ops_t ops_; + void* self_; + std::vector read_buf_; + std::vector write_buf_; + size_t written_ = 0; +}; + +namespace detail { +// Base ordering trick: a base class initializes before other bases declared after it, so +// this guarantees `buf` exists before std::istream/std::ostream stores its address. +struct StreamBufHolder { + StreamBuf buf; + StreamBufHolder(const svs_stream_ops_t& ops, void* self, StreamBuf::Direction direction) + : buf(ops, self, direction) {} +}; +} // namespace detail + +class InputStream : private detail::StreamBufHolder, public std::istream { + public: + InputStream(const svs_stream_ops_t& ops, void* self) + : detail::StreamBufHolder(ops, self, StreamBuf::Direction::read) + , std::istream(&buf) { + // badbit alone lets a read ending inside a payload fail silently; failbit joins the + // mask, with eofbit excluded so plain EOF is not an exception. + exceptions(std::ios_base::badbit | std::ios_base::failbit); + } +}; + +class OutputStream : private detail::StreamBufHolder, public std::ostream { + public: + OutputStream(const svs_stream_ops_t& ops, void* self) + : detail::StreamBufHolder(ops, self, StreamBuf::Direction::write) + , std::ostream(&buf) { + // A write-callback throw surfaces as badbit from the sentry and as failbit from + // the streambuf inserter in write_table; a missing bit loses the callback's code. + exceptions(std::ios_base::badbit | std::ios_base::failbit); + } +}; + +} // namespace svs::c_runtime diff --git a/bindings/c/src/svs_c.cpp b/bindings/c/src/svs_c.cpp index d242e73ba..1a63d6e94 100644 --- a/bindings/c/src/svs_c.cpp +++ b/bindings/c/src/svs_c.cpp @@ -23,6 +23,7 @@ #include "index_builder.hpp" #include "leanvec_training_data.hpp" #include "storage.hpp" +#include "stream.hpp" #include "threadpool.hpp" #include "types_support.hpp" @@ -882,6 +883,65 @@ svs_index_load(svs_index_builder_h builder, const char* directory, svs_error_h o ); } +extern "C" svs_index_h svs_index_load_stream( + svs_index_builder_h builder, svs_stream_i stream, svs_error_h out_err +) { + using namespace svs::c_runtime; + return wrap_exceptions( + [&]() { + EXPECT_ARG_NOT_NULL(builder); + EXPECT_ARG_NOT_NULL(stream); + NOT_IMPLEMENTED_IF( + (builder->impl->algorithm->type != SVS_ALGORITHM_TYPE_VAMANA), + "Only Vamana algorithm is currently supported for index loading" + ); + StreamBuf::validate(stream, /*need_write=*/false); + auto index = builder->impl->load_stream( + std::make_unique(*stream->ops, stream->self) + ); + if (index == nullptr) { + SET_ERROR(out_err, SVS_ERROR_RUNTIME, "Index load failed"); + return svs_index_h{nullptr}; + } + auto result = new svs_index; + result->impl = index; + return result; + }, + out_err + ); +} + +extern "C" svs_index_h svs_index_load_stream_dynamic( + svs_index_builder_h builder, + svs_stream_i stream, + size_t blocksize_bytes, + svs_error_h out_err +) { + using namespace svs::c_runtime; + return wrap_exceptions( + [&]() { + EXPECT_ARG_NOT_NULL(builder); + EXPECT_ARG_NOT_NULL(stream); + NOT_IMPLEMENTED_IF( + (builder->impl->algorithm->type != SVS_ALGORITHM_TYPE_VAMANA), + "Only Vamana algorithm is currently supported for dynamic index loading" + ); + StreamBuf::validate(stream, /*need_write=*/false); + auto index = builder->impl->load_stream_dynamic( + std::make_unique(*stream->ops, stream->self), blocksize_bytes + ); + if (index == nullptr) { + SET_ERROR(out_err, SVS_ERROR_RUNTIME, "Dynamic index load failed"); + return svs_index_h{nullptr}; + } + auto result = new svs_index; + result->impl = index; + return result; + }, + out_err + ); +} + extern "C" svs_index_h svs_index_convert(svs_index_builder_h builder, svs_index_h src_index, svs_error_h out_err) { using namespace svs::c_runtime; @@ -1109,6 +1169,25 @@ svs_index_save(svs_index_h index, const char* directory, svs_error_h out_err) { ); } +extern "C" bool +svs_index_save_stream(svs_index_h index, svs_stream_i stream, svs_error_h out_err) { + using namespace svs::c_runtime; + return wrap_exceptions( + [&]() { + EXPECT_ARG_NOT_NULL(index); + EXPECT_ARG_NOT_NULL(stream); + StreamBuf::validate(stream, /*need_write=*/true); + OutputStream os(*stream->ops, stream->self); + index->impl->save(os); + // The core never flushes and ~StreamBuf drops unflushed bytes, so without this + // the final partial buffer would be lost while the save reports success. + os.flush(); + return true; + }, + out_err + ); +} + extern "C" bool svs_index_dynamic_add_points( svs_index_h index, const float* new_points, diff --git a/bindings/c/tests/CMakeLists.txt b/bindings/c/tests/CMakeLists.txt index a7c351ccf..82e25b8fc 100644 --- a/bindings/c/tests/CMakeLists.txt +++ b/bindings/c/tests/CMakeLists.txt @@ -51,6 +51,7 @@ set(C_API_TEST_SOURCES c_api_index.cpp c_api_index_convert.cpp c_api_dynamic_index.cpp + c_api_stream.cpp ) # Create test executable diff --git a/bindings/c/tests/README.md b/bindings/c/tests/README.md index 4a6ba567f..9e76f7191 100644 --- a/bindings/c/tests/README.md +++ b/bindings/c/tests/README.md @@ -29,6 +29,7 @@ The tests are organized into separate files by functionality: - **c_api_index_builder.cpp**: Tests for index builder creation and configuration - **c_api_index.cpp**: Tests for index building, searching, and basic operations - **c_api_dynamic_index.cpp**: Tests for dynamic index operations (add, delete, consolidate, compact) +- **c_api_stream.cpp**: Tests for stream-based save and load operations, error handling, and round-trip validation Note: The main() function is provided by Catch2::Catch2WithMain automatically. diff --git a/bindings/c/tests/c_api_stream.cpp b/bindings/c/tests/c_api_stream.cpp new file mode 100644 index 000000000..5fc000393 --- /dev/null +++ b/bindings/c/tests/c_api_stream.cpp @@ -0,0 +1,1117 @@ +/* + * Copyright 2026 Intel Corporation + * + * Licensed 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. + */ + +// C API +#include "svs/c/svs_c.h" + +// catch2 +#include "catch2/catch_test_macros.hpp" + +// Test utilities +#include "c_api_test_utils.h" + +// Standard library +#include +#include +#include +#include +#include +#include +#include + +namespace { + +// The streambuf adapter's buffer size (bindings/c/src/stream.hpp); not part of the public +// API, so tests that depend on it hardcode the value. +constexpr size_t STREAM_BUFFER_SIZE = 64 * 1024; + +// In-memory sink/source backing the stream interface tests, capped at `max_read` per call +// so tests can force short reads. +struct MemoryStream { + std::vector bytes; + size_t pos = 0; + size_t max_read = std::numeric_limits::max(); +}; + +size_t memory_stream_read(void* self, void* buf, size_t n, svs_error_h /*out_err*/) { + auto* stream = static_cast(self); + size_t remaining = stream->bytes.size() - stream->pos; + size_t to_copy = std::min({n, remaining, stream->max_read}); + std::memcpy(buf, stream->bytes.data() + stream->pos, to_copy); + stream->pos += to_copy; + return to_copy; +} + +bool memory_stream_write(void* self, const void* buf, size_t n, svs_error_h /*out_err*/) { + auto* stream = static_cast(self); + const auto* src = static_cast(buf); + stream->bytes.insert(stream->bytes.end(), src, src + n); + return true; +} + +bool always_fail_write( + void* /*self*/, const void* /*buf*/, size_t /*n*/, svs_error_h /*out_err*/ +) { + return false; +} + +bool oom_write(void* /*self*/, const void* /*buf*/, size_t /*n*/, svs_error_h out_err) { + svs_error_set(out_err, SVS_ERROR_OUT_OF_MEMORY, "simulated allocator exhaustion"); + return false; +} + +size_t oom_read(void* /*self*/, void* /*buf*/, size_t /*n*/, svs_error_h out_err) { + svs_error_set(out_err, SVS_ERROR_OUT_OF_MEMORY, "simulated allocator exhaustion"); + return 0; +} + +size_t over_report_read(void* /*self*/, void* /*buf*/, size_t n, svs_error_h /*out_err*/) { + return n + 1; +} + +size_t error_with_bytes_read(void* self, void* buf, size_t n, svs_error_h out_err) { + size_t copied = memory_stream_read(self, buf, n, out_err); + svs_error_set(out_err, SVS_ERROR_RUNTIME, "error reported alongside returned bytes"); + return copied; +} + +bool fail_write_with_ok_error( + void* /*self*/, const void* /*buf*/, size_t /*n*/, svs_error_h out_err +) { + svs_error_set(out_err, SVS_OK, "no error, yet still failing"); + return false; +} + +// Sized so the saved payload (data + graph) provably exceeds two StreamBuf write buffers: +// data alone is 200 * 700 * 4 = 560'000 bytes, well over 2 * STREAM_BUFFER_SIZE = 131'072. +constexpr size_t MULTIBUFFER_NUM_VECTORS = 200; +constexpr size_t MULTIBUFFER_DIMENSION = 700; + +// Owns the algorithm/builder/index triple built over the oversized data set above, shared +// by the sections that need a multi-buffer payload instead of duplicating build setup. +struct MultiBufferIndex { + svs_algorithm_h algorithm = nullptr; + svs_index_builder_h builder = nullptr; + svs_index_h index = nullptr; +}; + +MultiBufferIndex build_multibuffer_index(std::vector& data, svs_error_h error) { + MultiBufferIndex result; + result.algorithm = svs_algorithm_create_vamana(16, 32, 50, error); + result.builder = svs_index_builder_create( + SVS_DISTANCE_METRIC_EUCLIDEAN, MULTIBUFFER_DIMENSION, result.algorithm, error + ); + svs_index_builder_set_threadpool( + result.builder, SVS_THREADPOOL_KIND_SINGLE_THREAD, 1, error + ); + generate_test_data(data, MULTIBUFFER_NUM_VECTORS, MULTIBUFFER_DIMENSION); + result.index = + svs_index_build(result.builder, data.data(), MULTIBUFFER_NUM_VECTORS, error); + return result; +} + +// Rejects a write smaller than the adapter's buffer and counts the full-size writes that +// succeeded before it, so a test can assert the failure was the *last* of several writes. +struct CountingFailSink { + MemoryStream stream; + size_t full_write_count = 0; +}; + +bool fail_partial_write_counted( + void* self, const void* buf, size_t n, svs_error_h out_err +) { + auto* sink = static_cast(self); + if (n < STREAM_BUFFER_SIZE) { + svs_error_set(out_err, SVS_ERROR_RUNTIME, "refusing partial write"); + return false; + } + sink->full_write_count++; + return memory_stream_write(&sink->stream, buf, n, out_err); +} + +// Fails once, on the fail_at_invocation'th call, then records any later call: a +// redelivery after failure means a rejected chunk reached the callback again. +struct RecordingFailSink { + size_t fail_at_invocation = 0; + size_t invocation_count = 0; + size_t invocations_after_failure = 0; + bool has_failed = false; +}; + +bool fail_after_n_write(void* self, const void* /*buf*/, size_t n, svs_error_h out_err) { + auto* sink = static_cast(self); + if (sink->has_failed) { + sink->invocations_after_failure++; + return false; + } + sink->invocation_count++; + if (sink->invocation_count == sink->fail_at_invocation) { + sink->has_failed = true; + svs_error_set(out_err, SVS_ERROR_RUNTIME, "simulated failure mid-stream"); + return false; + } + return true; +} + +enum class StorageKind { Float16, ScalarQuantization, Lvq, LeanVec }; + +constexpr std::array STORAGE_KINDS = { + StorageKind::Float16, + StorageKind::ScalarQuantization, + StorageKind::Lvq, + StorageKind::LeanVec, +}; + +const char* storage_kind_name(StorageKind kind) { + switch (kind) { + case StorageKind::Float16: + return "float16"; + case StorageKind::ScalarQuantization: + return "scalar quantization (int8)"; + case StorageKind::Lvq: + return "LVQ (int4 primary, int8 residual)"; + case StorageKind::LeanVec: + return "LeanVec (int4 primary, int8 secondary)"; + } + return "unknown"; +} + +svs_storage_h create_storage_case(StorageKind kind, size_t dimension, svs_error_h error) { + switch (kind) { + case StorageKind::Float16: + return svs_storage_create_simple(SVS_DATA_TYPE_FLOAT16, error); + case StorageKind::ScalarQuantization: + return svs_storage_create_sq(SVS_DATA_TYPE_INT8, error); + case StorageKind::Lvq: + return svs_storage_create_lvq(SVS_DATA_TYPE_INT4, SVS_DATA_TYPE_INT8, error); + case StorageKind::LeanVec: + return svs_storage_create_leanvec( + dimension / 2, SVS_DATA_TYPE_INT4, SVS_DATA_TYPE_INT8, error + ); + } + return nullptr; +} + +// Compressed storages need not reproduce pre-save distances bit for bit. +bool distance_within_tolerance(float reloaded, float original) { + float scale = std::max({std::fabs(reloaded), std::fabs(original), 1.0f}); + return std::fabs(reloaded - original) <= 1e-3f * scale; +} + +} // namespace + +CATCH_TEST_CASE("C API Stream Save and Load", "[c_api][index][stream]") { + const size_t NUM_VECTORS = 100; + const size_t DIMENSION = 32; + const size_t K = 5; + + std::vector data; + std::vector queries; + generate_test_data(data, NUM_VECTORS, DIMENSION); + generate_test_data(queries, 3, DIMENSION); + + svs_error_h error = svs_error_create(); + + svs_algorithm_h algorithm = svs_algorithm_create_vamana(16, 32, 50, error); + CATCH_REQUIRE(algorithm != nullptr); + + svs_index_builder_h builder = svs_index_builder_create( + SVS_DISTANCE_METRIC_EUCLIDEAN, DIMENSION, algorithm, error + ); + CATCH_REQUIRE(builder != nullptr); + + // Single-threaded so a greedy search visits the same path on both indexes. + bool success = svs_index_builder_set_threadpool( + builder, SVS_THREADPOOL_KIND_SINGLE_THREAD, 1, error + ); + CATCH_REQUIRE(success); + CATCH_REQUIRE(svs_error_ok(error)); + + CATCH_SECTION("Static round-trip through an in-memory stream") { + svs_index_h index = svs_index_build(builder, data.data(), NUM_VECTORS, error); + CATCH_REQUIRE(index != nullptr); + + svs_search_results_t before = SVS_INIT_SEARCH_RESULTS(); + CATCH_REQUIRE(svs_index_search_topk( + index, queries.data(), 3, K, &before, nullptr, nullptr, error + )); + CATCH_REQUIRE(svs_error_ok(error)); + + MemoryStream stream; + svs_stream_interface_ops write_ops = + SVS_INIT_STREAM_OPS(nullptr, memory_stream_write); + svs_stream_interface out_stream = SVS_MAKE_INTERFACE(&stream, write_ops); + CATCH_REQUIRE(svs_index_save_stream(index, &out_stream, error)); + CATCH_REQUIRE(svs_error_ok(error)); + + svs_stream_interface_ops read_ops = + SVS_INIT_STREAM_OPS(memory_stream_read, nullptr); + svs_stream_interface in_stream = SVS_MAKE_INTERFACE(&stream, read_ops); + svs_index_h loaded = svs_index_load_stream(builder, &in_stream, error); + CATCH_REQUIRE(loaded != nullptr); + CATCH_REQUIRE(svs_error_ok(error)); + + svs_search_results_t after = SVS_INIT_SEARCH_RESULTS(); + CATCH_REQUIRE(svs_index_search_topk( + loaded, queries.data(), 3, K, &after, nullptr, nullptr, error + )); + CATCH_REQUIRE(svs_error_ok(error)); + CATCH_REQUIRE(after.num_queries == before.num_queries); + for (size_t i = 0; i < before.num_queries * K; ++i) { + CATCH_REQUIRE(after.indices[i] == before.indices[i]); + CATCH_REQUIRE(after.distances[i] == before.distances[i]); + } + + svs_search_results_free(&before); + svs_search_results_free(&after); + svs_index_free(loaded); + svs_index_free(index); + } + + CATCH_SECTION("Round-trip through an in-memory stream spanning multiple write buffers" + ) { + std::vector big_data; + MultiBufferIndex mb = build_multibuffer_index(big_data, error); + CATCH_REQUIRE(mb.algorithm != nullptr); + CATCH_REQUIRE(mb.builder != nullptr); + CATCH_REQUIRE(mb.index != nullptr); + CATCH_REQUIRE(svs_error_ok(error)); + + std::vector big_queries; + generate_test_data(big_queries, 3, MULTIBUFFER_DIMENSION); + + svs_search_results_t before = SVS_INIT_SEARCH_RESULTS(); + CATCH_REQUIRE(svs_index_search_topk( + mb.index, big_queries.data(), 3, K, &before, nullptr, nullptr, error + )); + CATCH_REQUIRE(svs_error_ok(error)); + + MemoryStream stream; + svs_stream_interface_ops write_ops = + SVS_INIT_STREAM_OPS(nullptr, memory_stream_write); + svs_stream_interface out_stream = SVS_MAKE_INTERFACE(&stream, write_ops); + CATCH_REQUIRE(svs_index_save_stream(mb.index, &out_stream, error)); + CATCH_REQUIRE(svs_error_ok(error)); + // Proves the write path flushed several full buffers and a trailing partial one, + // and the read path below refills the get area several times. + CATCH_REQUIRE(stream.bytes.size() > 2 * STREAM_BUFFER_SIZE); + + svs_stream_interface_ops read_ops = + SVS_INIT_STREAM_OPS(memory_stream_read, nullptr); + svs_stream_interface in_stream = SVS_MAKE_INTERFACE(&stream, read_ops); + svs_index_h loaded = svs_index_load_stream(mb.builder, &in_stream, error); + CATCH_REQUIRE(loaded != nullptr); + CATCH_REQUIRE(svs_error_ok(error)); + + svs_search_results_t after = SVS_INIT_SEARCH_RESULTS(); + CATCH_REQUIRE(svs_index_search_topk( + loaded, big_queries.data(), 3, K, &after, nullptr, nullptr, error + )); + CATCH_REQUIRE(svs_error_ok(error)); + CATCH_REQUIRE(after.num_queries == before.num_queries); + for (size_t i = 0; i < before.num_queries * K; ++i) { + CATCH_REQUIRE(after.indices[i] == before.indices[i]); + CATCH_REQUIRE(after.distances[i] == before.distances[i]); + } + + svs_search_results_free(&before); + svs_search_results_free(&after); + svs_index_free(loaded); + svs_index_free(mb.index); + svs_index_builder_free(mb.builder); + svs_algorithm_free(mb.algorithm); + } + + CATCH_SECTION("Save fails when only the final partial flush fails") { + std::vector big_data; + MultiBufferIndex mb = build_multibuffer_index(big_data, error); + CATCH_REQUIRE(mb.index != nullptr); + CATCH_REQUIRE(svs_error_ok(error)); + + MemoryStream probe; + svs_stream_interface_ops probe_ops = + SVS_INIT_STREAM_OPS(nullptr, memory_stream_write); + svs_stream_interface probe_stream = SVS_MAKE_INTERFACE(&probe, probe_ops); + CATCH_REQUIRE(svs_index_save_stream(mb.index, &probe_stream, error)); + // Needs two full buffers plus a partial one; otherwise the only flush failing + // looks the same as the final one failing. + CATCH_REQUIRE(probe.bytes.size() > 2 * STREAM_BUFFER_SIZE); + // The sink rejects only writes shorter than the buffer; a payload ending exactly + // on a buffer boundary leaves nothing for it to catch. + CATCH_REQUIRE(probe.bytes.size() % STREAM_BUFFER_SIZE != 0); + + CountingFailSink sink; + svs_stream_interface_ops fail_ops = + SVS_INIT_STREAM_OPS(nullptr, fail_partial_write_counted); + svs_stream_interface fail_stream = SVS_MAKE_INTERFACE(&sink, fail_ops); + CATCH_REQUIRE_FALSE(svs_index_save_stream(mb.index, &fail_stream, error)); + CATCH_REQUIRE_FALSE(svs_error_ok(error)); + // At least two full-size writes must have succeeded before the failing partial one, + // or this is once again just "the single flush failed". + CATCH_REQUIRE(sink.full_write_count >= 2); + + svs_index_free(mb.index); + svs_index_builder_free(mb.builder); + svs_algorithm_free(mb.algorithm); + } + + CATCH_SECTION("Write failure never redelivers the same bytes during destructor unwind" + ) { + std::vector big_data; + MultiBufferIndex mb = build_multibuffer_index(big_data, error); + CATCH_REQUIRE(mb.index != nullptr); + CATCH_REQUIRE(svs_error_ok(error)); + + // Fails on the 3rd of many calls, so bytes remain that a double-delivery bug + // would hand to the callback again. + RecordingFailSink sink; + sink.fail_at_invocation = 3; + svs_stream_interface_ops fail_ops = + SVS_INIT_STREAM_OPS(nullptr, fail_after_n_write); + svs_stream_interface fail_stream = SVS_MAKE_INTERFACE(&sink, fail_ops); + CATCH_REQUIRE_FALSE(svs_index_save_stream(mb.index, &fail_stream, error)); + CATCH_REQUIRE_FALSE(svs_error_ok(error)); + CATCH_REQUIRE(sink.invocation_count == 3); + CATCH_REQUIRE(sink.invocations_after_failure == 0); + + svs_index_free(mb.index); + svs_index_builder_free(mb.builder); + svs_algorithm_free(mb.algorithm); + } + + CATCH_SECTION("Write callback failure aborts save") { + svs_index_h index = svs_index_build(builder, data.data(), NUM_VECTORS, error); + CATCH_REQUIRE(index != nullptr); + + svs_stream_interface_ops fail_ops = SVS_INIT_STREAM_OPS(nullptr, always_fail_write); + svs_stream_interface fail_stream = SVS_MAKE_INTERFACE(nullptr, fail_ops); + CATCH_REQUIRE_FALSE(svs_index_save_stream(index, &fail_stream, error)); + CATCH_REQUIRE_FALSE(svs_error_ok(error)); + + svs_index_free(index); + } + + CATCH_SECTION("Write callback reports a specific error code") { + svs_index_h index = svs_index_build(builder, data.data(), NUM_VECTORS, error); + CATCH_REQUIRE(index != nullptr); + + svs_stream_interface_ops fail_ops = SVS_INIT_STREAM_OPS(nullptr, oom_write); + svs_stream_interface fail_stream = SVS_MAKE_INTERFACE(nullptr, fail_ops); + CATCH_REQUIRE_FALSE(svs_index_save_stream(index, &fail_stream, error)); + CATCH_REQUIRE(svs_error_get_code(error) == SVS_ERROR_OUT_OF_MEMORY); + + svs_index_free(index); + } + + CATCH_SECTION("Read callback reports a specific error code") { + svs_stream_interface_ops fail_ops = SVS_INIT_STREAM_OPS(oom_read, nullptr); + svs_stream_interface fail_stream = SVS_MAKE_INTERFACE(nullptr, fail_ops); + svs_index_h loaded = svs_index_load_stream(builder, &fail_stream, error); + CATCH_REQUIRE(loaded == nullptr); + CATCH_REQUIRE(svs_error_get_code(error) == SVS_ERROR_OUT_OF_MEMORY); + } + + CATCH_SECTION("Short reads and EOF on load") { + svs_index_h index = svs_index_build(builder, data.data(), NUM_VECTORS, error); + CATCH_REQUIRE(index != nullptr); + + MemoryStream stream; + svs_stream_interface_ops write_ops = + SVS_INIT_STREAM_OPS(nullptr, memory_stream_write); + svs_stream_interface out_stream = SVS_MAKE_INTERFACE(&stream, write_ops); + CATCH_REQUIRE(svs_index_save_stream(index, &out_stream, error)); + + // Hand back at most 3 bytes per call, forcing many short reads before the final + // 0-byte EOF. + stream.max_read = 3; + svs_stream_interface_ops read_ops = + SVS_INIT_STREAM_OPS(memory_stream_read, nullptr); + svs_stream_interface in_stream = SVS_MAKE_INTERFACE(&stream, read_ops); + svs_index_h loaded = svs_index_load_stream(builder, &in_stream, error); + CATCH_REQUIRE(loaded != nullptr); + CATCH_REQUIRE(svs_error_ok(error)); + + svs_search_results_t results = SVS_INIT_SEARCH_RESULTS(); + CATCH_REQUIRE(svs_index_search_topk( + loaded, queries.data(), 3, K, &results, nullptr, nullptr, error + )); + CATCH_REQUIRE(svs_error_ok(error)); + CATCH_REQUIRE(results.num_queries == 3); + + svs_search_results_free(&results); + svs_index_free(loaded); + svs_index_free(index); + } + + CATCH_SECTION("Load fails on an empty stream instead of hanging or crashing") { + MemoryStream stream; + svs_stream_interface_ops read_ops = + SVS_INIT_STREAM_OPS(memory_stream_read, nullptr); + svs_stream_interface in_stream = SVS_MAKE_INTERFACE(&stream, read_ops); + svs_index_h loaded = svs_index_load_stream(builder, &in_stream, error); + CATCH_REQUIRE(loaded == nullptr); + CATCH_REQUIRE_FALSE(svs_error_ok(error)); + } + + CATCH_SECTION("Load fails on a stream truncated partway through a valid payload") { + svs_index_h index = svs_index_build(builder, data.data(), NUM_VECTORS, error); + CATCH_REQUIRE(index != nullptr); + + MemoryStream stream; + svs_stream_interface_ops write_ops = + SVS_INIT_STREAM_OPS(nullptr, memory_stream_write); + svs_stream_interface out_stream = SVS_MAKE_INTERFACE(&stream, write_ops); + CATCH_REQUIRE(svs_index_save_stream(index, &out_stream, error)); + CATCH_REQUIRE(svs_error_ok(error)); + + // Cut the valid payload in half: the read callback hands back real bytes for a + // while, then reports EOF (0 bytes) before the format is fully consumed. + CATCH_REQUIRE(stream.bytes.size() > 1); + stream.bytes.resize(stream.bytes.size() / 2); + stream.pos = 0; + + svs_stream_interface_ops read_ops = + SVS_INIT_STREAM_OPS(memory_stream_read, nullptr); + svs_stream_interface in_stream = SVS_MAKE_INTERFACE(&stream, read_ops); + svs_index_h loaded = svs_index_load_stream(builder, &in_stream, error); + CATCH_REQUIRE(loaded == nullptr); + CATCH_REQUIRE_FALSE(svs_error_ok(error)); + + svs_index_free(index); + } + + CATCH_SECTION("Load fails when the last 16 bytes of a valid payload are missing") { + svs_index_h index = svs_index_build(builder, data.data(), NUM_VECTORS, error); + CATCH_REQUIRE(index != nullptr); + + MemoryStream stream; + svs_stream_interface_ops write_ops = + SVS_INIT_STREAM_OPS(nullptr, memory_stream_write); + svs_stream_interface out_stream = SVS_MAKE_INTERFACE(&stream, write_ops); + CATCH_REQUIRE(svs_index_save_stream(index, &out_stream, error)); + CATCH_REQUIRE(svs_error_ok(error)); + CATCH_REQUIRE(stream.bytes.size() > 16); + stream.bytes.resize(stream.bytes.size() - 16); + stream.pos = 0; + + svs_stream_interface_ops read_ops = + SVS_INIT_STREAM_OPS(memory_stream_read, nullptr); + svs_stream_interface in_stream = SVS_MAKE_INTERFACE(&stream, read_ops); + svs_index_h loaded = svs_index_load_stream(builder, &in_stream, error); + CATCH_REQUIRE(loaded == nullptr); + CATCH_REQUIRE_FALSE(svs_error_ok(error)); + + svs_index_free(index); + } + + CATCH_SECTION("Dynamic load fails when the last 16 bytes of a valid payload are missing" + ) { + std::vector ids(NUM_VECTORS); + std::iota(ids.begin(), ids.end(), size_t{0}); + const size_t BLOCK_SIZE = 1024 * 1024; + svs_index_h index = svs_index_build_dynamic( + builder, data.data(), ids.data(), NUM_VECTORS, BLOCK_SIZE, error + ); + CATCH_REQUIRE(index != nullptr); + + MemoryStream stream; + svs_stream_interface_ops write_ops = + SVS_INIT_STREAM_OPS(nullptr, memory_stream_write); + svs_stream_interface out_stream = SVS_MAKE_INTERFACE(&stream, write_ops); + CATCH_REQUIRE(svs_index_save_stream(index, &out_stream, error)); + CATCH_REQUIRE(svs_error_ok(error)); + CATCH_REQUIRE(stream.bytes.size() > 16); + stream.bytes.resize(stream.bytes.size() - 16); + stream.pos = 0; + + svs_stream_interface_ops read_ops = + SVS_INIT_STREAM_OPS(memory_stream_read, nullptr); + svs_stream_interface in_stream = SVS_MAKE_INTERFACE(&stream, read_ops); + svs_index_h loaded = + svs_index_load_stream_dynamic(builder, &in_stream, BLOCK_SIZE, error); + CATCH_REQUIRE(loaded == nullptr); + CATCH_REQUIRE_FALSE(svs_error_ok(error)); + + svs_index_free(index); + } + + CATCH_SECTION("Read callback returning more bytes than requested fails the load") { + svs_stream_interface_ops fail_ops = SVS_INIT_STREAM_OPS(over_report_read, nullptr); + svs_stream_interface fail_stream = SVS_MAKE_INTERFACE(nullptr, fail_ops); + svs_index_h loaded = svs_index_load_stream(builder, &fail_stream, error); + CATCH_REQUIRE(loaded == nullptr); + CATCH_REQUIRE_FALSE(svs_error_ok(error)); + } + + CATCH_SECTION("Read callback error is honored even when it also returns bytes") { + svs_index_h index = svs_index_build(builder, data.data(), NUM_VECTORS, error); + CATCH_REQUIRE(index != nullptr); + + MemoryStream stream; + svs_stream_interface_ops write_ops = + SVS_INIT_STREAM_OPS(nullptr, memory_stream_write); + svs_stream_interface out_stream = SVS_MAKE_INTERFACE(&stream, write_ops); + CATCH_REQUIRE(svs_index_save_stream(index, &out_stream, error)); + CATCH_REQUIRE(svs_error_ok(error)); + + svs_stream_interface_ops read_ops = + SVS_INIT_STREAM_OPS(error_with_bytes_read, nullptr); + svs_stream_interface in_stream = SVS_MAKE_INTERFACE(&stream, read_ops); + svs_index_h loaded = svs_index_load_stream(builder, &in_stream, error); + CATCH_REQUIRE(loaded == nullptr); + CATCH_REQUIRE(svs_error_get_code(error) == SVS_ERROR_RUNTIME); + + svs_index_free(index); + } + + CATCH_SECTION("Write callback failure with SVS_OK error is not reported as success") { + svs_index_h index = svs_index_build(builder, data.data(), NUM_VECTORS, error); + CATCH_REQUIRE(index != nullptr); + + svs_stream_interface_ops fail_ops = + SVS_INIT_STREAM_OPS(nullptr, fail_write_with_ok_error); + svs_stream_interface fail_stream = SVS_MAKE_INTERFACE(nullptr, fail_ops); + CATCH_REQUIRE_FALSE(svs_index_save_stream(index, &fail_stream, error)); + CATCH_REQUIRE_FALSE(svs_error_ok(error)); + CATCH_REQUIRE(svs_error_get_code(error) == SVS_ERROR_UNKNOWN); + + svs_index_free(index); + } + + CATCH_SECTION("Dynamic round-trip, add_points, and search") { + std::vector ids(NUM_VECTORS); + std::iota(ids.begin(), ids.end(), size_t{0}); + const size_t BLOCK_SIZE = 1024 * 1024; + svs_index_h index = svs_index_build_dynamic( + builder, data.data(), ids.data(), NUM_VECTORS, BLOCK_SIZE, error + ); + CATCH_REQUIRE(index != nullptr); + + svs_search_results_t before = SVS_INIT_SEARCH_RESULTS(); + CATCH_REQUIRE(svs_index_search_topk( + index, queries.data(), 3, K, &before, nullptr, nullptr, error + )); + CATCH_REQUIRE(svs_error_ok(error)); + + MemoryStream stream; + svs_stream_interface_ops write_ops = + SVS_INIT_STREAM_OPS(nullptr, memory_stream_write); + svs_stream_interface out_stream = SVS_MAKE_INTERFACE(&stream, write_ops); + CATCH_REQUIRE(svs_index_save_stream(index, &out_stream, error)); + CATCH_REQUIRE(svs_error_ok(error)); + + svs_stream_interface_ops read_ops = + SVS_INIT_STREAM_OPS(memory_stream_read, nullptr); + svs_stream_interface in_stream = SVS_MAKE_INTERFACE(&stream, read_ops); + svs_index_h loaded = + svs_index_load_stream_dynamic(builder, &in_stream, BLOCK_SIZE, error); + CATCH_REQUIRE(loaded != nullptr); + CATCH_REQUIRE(svs_error_ok(error)); + + // Compares element-by-element against the pre-save index, so a load that + // corrupts data or graph cannot pass just by returning. + svs_search_results_t after = SVS_INIT_SEARCH_RESULTS(); + CATCH_REQUIRE(svs_index_search_topk( + loaded, queries.data(), 3, K, &after, nullptr, nullptr, error + )); + CATCH_REQUIRE(svs_error_ok(error)); + CATCH_REQUIRE(after.num_queries == before.num_queries); + for (size_t i = 0; i < before.num_queries * K; ++i) { + CATCH_REQUIRE(after.indices[i] == before.indices[i]); + CATCH_REQUIRE(after.distances[i] == before.distances[i]); + } + svs_search_results_free(&before); + svs_search_results_free(&after); + + std::vector new_data; + std::vector new_ids = {NUM_VECTORS, NUM_VECTORS + 1}; + generate_test_data(new_data, 2, DIMENSION); + size_t added_count = 0; + CATCH_REQUIRE(svs_index_dynamic_add_points( + loaded, new_data.data(), new_ids.data(), 2, &added_count, error + )); + CATCH_REQUIRE(added_count == 2); + CATCH_REQUIRE(svs_error_ok(error)); + + svs_search_results_t results = SVS_INIT_SEARCH_RESULTS(); + CATCH_REQUIRE(svs_index_search_topk( + loaded, queries.data(), 3, K, &results, nullptr, nullptr, error + )); + CATCH_REQUIRE(svs_error_ok(error)); + CATCH_REQUIRE(results.num_queries == 3); + + // Queries with a newly added point's exact vector and requires its id in the + // result, so a load that dropped or corrupted the graph cannot pass on count alone. + std::vector probe_query( + new_data.begin(), new_data.begin() + static_cast(DIMENSION) + ); + svs_search_results_t probe_results = SVS_INIT_SEARCH_RESULTS(); + CATCH_REQUIRE(svs_index_search_topk( + loaded, probe_query.data(), 1, K, &probe_results, nullptr, nullptr, error + )); + CATCH_REQUIRE(svs_error_ok(error)); + bool found_new_id = std::any_of( + probe_results.indices, + probe_results.indices + K, + [new_id = new_ids[0]](size_t idx) { return idx == new_id; } + ); + CATCH_REQUIRE(found_new_id); + + svs_search_results_free(&results); + svs_search_results_free(&probe_results); + svs_index_free(loaded); + svs_index_free(index); + } + + CATCH_SECTION("Dynamic Stream Load Uses Custom Allocator For Graph") { + // Assert same number of bytes are allocated during original construction and after + // a streaming I/O loop + std::vector ids(NUM_VECTORS); + std::iota(ids.begin(), ids.end(), size_t{0}); + const size_t BLOCK_SIZE = 1024 * 1024; + + TrackingAllocator build_tracker; + svs_allocator_interface_ops build_alloc_ops = SVS_INIT_ALLOCATOR_OPS( + tracking_allocator_allocate, tracking_allocator_deallocate + ); + svs_allocator_interface build_allocator = + SVS_MAKE_INTERFACE(&build_tracker, build_alloc_ops); + CATCH_REQUIRE( + svs_index_builder_set_allocator_custom(builder, &build_allocator, error) + ); + CATCH_REQUIRE(svs_error_ok(error)); + + svs_index_h index = svs_index_build_dynamic( + builder, data.data(), ids.data(), NUM_VECTORS, BLOCK_SIZE, error + ); + CATCH_REQUIRE(index != nullptr); + CATCH_REQUIRE(svs_error_ok(error)); + size_t built_bytes = build_tracker.live_bytes; + + MemoryStream stream; + svs_stream_interface_ops write_ops = + SVS_INIT_STREAM_OPS(nullptr, memory_stream_write); + svs_stream_interface out_stream = SVS_MAKE_INTERFACE(&stream, write_ops); + CATCH_REQUIRE(svs_index_save_stream(index, &out_stream, error)); + CATCH_REQUIRE(svs_error_ok(error)); + + svs_index_builder_h load_builder = svs_index_builder_create( + SVS_DISTANCE_METRIC_EUCLIDEAN, DIMENSION, algorithm, error + ); + CATCH_REQUIRE(load_builder != nullptr); + CATCH_REQUIRE(svs_index_builder_set_threadpool( + load_builder, SVS_THREADPOOL_KIND_SINGLE_THREAD, 1, error + )); + CATCH_REQUIRE(svs_error_ok(error)); + + TrackingAllocator load_tracker; + svs_allocator_interface_ops load_alloc_ops = SVS_INIT_ALLOCATOR_OPS( + tracking_allocator_allocate, tracking_allocator_deallocate + ); + svs_allocator_interface load_allocator = + SVS_MAKE_INTERFACE(&load_tracker, load_alloc_ops); + CATCH_REQUIRE( + svs_index_builder_set_allocator_custom(load_builder, &load_allocator, error) + ); + CATCH_REQUIRE(svs_error_ok(error)); + + svs_stream_interface_ops read_ops = + SVS_INIT_STREAM_OPS(memory_stream_read, nullptr); + svs_stream_interface in_stream = SVS_MAKE_INTERFACE(&stream, read_ops); + svs_index_h loaded = + svs_index_load_stream_dynamic(load_builder, &in_stream, BLOCK_SIZE, error); + CATCH_REQUIRE(loaded != nullptr); + CATCH_REQUIRE(svs_error_ok(error)); + + size_t loaded_bytes = load_tracker.live_bytes; + CATCH_REQUIRE(loaded_bytes == built_bytes); + CATCH_REQUIRE(load_tracker.alloc_count > 0); + + svs_index_free(loaded); + svs_index_free(index); + svs_index_builder_free(load_builder); + } + + CATCH_SECTION("Static Stream Load Uses Custom Allocator For Graph") { + // Assert same number of bytes are allocated during original construction and after + // a streaming I/O loop + TrackingAllocator build_tracker; + svs_allocator_interface_ops build_alloc_ops = SVS_INIT_ALLOCATOR_OPS( + tracking_allocator_allocate, tracking_allocator_deallocate + ); + svs_allocator_interface build_allocator = + SVS_MAKE_INTERFACE(&build_tracker, build_alloc_ops); + CATCH_REQUIRE( + svs_index_builder_set_allocator_custom(builder, &build_allocator, error) + ); + CATCH_REQUIRE(svs_error_ok(error)); + + svs_index_h index = svs_index_build(builder, data.data(), NUM_VECTORS, error); + CATCH_REQUIRE(index != nullptr); + CATCH_REQUIRE(svs_error_ok(error)); + size_t built_bytes = build_tracker.live_bytes; + + MemoryStream stream; + svs_stream_interface_ops write_ops = + SVS_INIT_STREAM_OPS(nullptr, memory_stream_write); + svs_stream_interface out_stream = SVS_MAKE_INTERFACE(&stream, write_ops); + CATCH_REQUIRE(svs_index_save_stream(index, &out_stream, error)); + CATCH_REQUIRE(svs_error_ok(error)); + + svs_index_builder_h load_builder = svs_index_builder_create( + SVS_DISTANCE_METRIC_EUCLIDEAN, DIMENSION, algorithm, error + ); + CATCH_REQUIRE(load_builder != nullptr); + CATCH_REQUIRE(svs_index_builder_set_threadpool( + load_builder, SVS_THREADPOOL_KIND_SINGLE_THREAD, 1, error + )); + CATCH_REQUIRE(svs_error_ok(error)); + + TrackingAllocator load_tracker; + svs_allocator_interface_ops load_alloc_ops = SVS_INIT_ALLOCATOR_OPS( + tracking_allocator_allocate, tracking_allocator_deallocate + ); + svs_allocator_interface load_allocator = + SVS_MAKE_INTERFACE(&load_tracker, load_alloc_ops); + CATCH_REQUIRE( + svs_index_builder_set_allocator_custom(load_builder, &load_allocator, error) + ); + CATCH_REQUIRE(svs_error_ok(error)); + + svs_stream_interface_ops read_ops = + SVS_INIT_STREAM_OPS(memory_stream_read, nullptr); + svs_stream_interface in_stream = SVS_MAKE_INTERFACE(&stream, read_ops); + svs_index_h loaded = svs_index_load_stream(load_builder, &in_stream, error); + CATCH_REQUIRE(loaded != nullptr); + CATCH_REQUIRE(svs_error_ok(error)); + + size_t loaded_bytes = load_tracker.live_bytes; + CATCH_REQUIRE(loaded_bytes == built_bytes); + CATCH_REQUIRE(load_tracker.alloc_count > 0); + + svs_index_free(loaded); + svs_index_free(index); + svs_index_builder_free(load_builder); + } + + svs_index_builder_free(builder); + svs_algorithm_free(algorithm); + svs_error_free(error); +} + +CATCH_TEST_CASE("C API Stream Interface Validation", "[c_api][index][stream][error]") { + const size_t NUM_VECTORS = 20; + const size_t DIMENSION = 8; + + std::vector data; + generate_test_data(data, NUM_VECTORS, DIMENSION); + + svs_error_h error = svs_error_create(); + + svs_algorithm_h algorithm = svs_algorithm_create_vamana(16, 32, 50, error); + CATCH_REQUIRE(algorithm != nullptr); + + svs_index_builder_h builder = svs_index_builder_create( + SVS_DISTANCE_METRIC_EUCLIDEAN, DIMENSION, algorithm, error + ); + CATCH_REQUIRE(builder != nullptr); + + bool success = svs_index_builder_set_threadpool( + builder, SVS_THREADPOOL_KIND_SINGLE_THREAD, 1, error + ); + CATCH_REQUIRE(success); + + svs_index_h index = svs_index_build(builder, data.data(), NUM_VECTORS, error); + CATCH_REQUIRE(index != nullptr); + + CATCH_SECTION("Null interface pointer is rejected") { + CATCH_REQUIRE_FALSE(svs_index_save_stream(index, nullptr, error)); + CATCH_REQUIRE(svs_error_get_code(error) == SVS_ERROR_INVALID_ARGUMENT); + + CATCH_REQUIRE(svs_index_load_stream(builder, nullptr, error) == nullptr); + CATCH_REQUIRE(svs_error_get_code(error) == SVS_ERROR_INVALID_ARGUMENT); + } + + CATCH_SECTION("Null ops pointer is rejected") { + svs_stream_interface stream{nullptr, nullptr}; + CATCH_REQUIRE_FALSE(svs_index_save_stream(index, &stream, error)); + CATCH_REQUIRE(svs_error_get_code(error) == SVS_ERROR_INVALID_ARGUMENT); + + CATCH_REQUIRE(svs_index_load_stream(builder, &stream, error) == nullptr); + CATCH_REQUIRE(svs_error_get_code(error) == SVS_ERROR_INVALID_ARGUMENT); + } + + CATCH_SECTION("struct_size smaller than expected is rejected") { + svs_stream_interface_ops ops = + SVS_INIT_STREAM_OPS(memory_stream_read, memory_stream_write); + ops.struct_size = sizeof(uint32_t) + sizeof(size_t); + svs_stream_interface stream = SVS_MAKE_INTERFACE(nullptr, ops); + CATCH_REQUIRE_FALSE(svs_index_save_stream(index, &stream, error)); + CATCH_REQUIRE(svs_error_get_code(error) == SVS_ERROR_INVALID_ARGUMENT); + + CATCH_REQUIRE(svs_index_load_stream(builder, &stream, error) == nullptr); + CATCH_REQUIRE(svs_error_get_code(error) == SVS_ERROR_INVALID_ARGUMENT); + } + + CATCH_SECTION("version above svs_get_version() is rejected") { + svs_stream_interface_ops ops = + SVS_INIT_STREAM_OPS(memory_stream_read, memory_stream_write); + ops.version = svs_get_version() + 1; + svs_stream_interface stream = SVS_MAKE_INTERFACE(nullptr, ops); + CATCH_REQUIRE_FALSE(svs_index_save_stream(index, &stream, error)); + CATCH_REQUIRE(svs_error_get_code(error) == SVS_ERROR_INVALID_ARGUMENT); + + CATCH_REQUIRE(svs_index_load_stream(builder, &stream, error) == nullptr); + CATCH_REQUIRE(svs_error_get_code(error) == SVS_ERROR_INVALID_ARGUMENT); + } + + CATCH_SECTION("NULL write on save is rejected") { + svs_stream_interface_ops ops = SVS_INIT_STREAM_OPS(memory_stream_read, nullptr); + svs_stream_interface stream = SVS_MAKE_INTERFACE(nullptr, ops); + CATCH_REQUIRE_FALSE(svs_index_save_stream(index, &stream, error)); + CATCH_REQUIRE(svs_error_get_code(error) == SVS_ERROR_INVALID_ARGUMENT); + } + + CATCH_SECTION("NULL read on load is rejected") { + svs_stream_interface_ops ops = SVS_INIT_STREAM_OPS(nullptr, memory_stream_write); + svs_stream_interface stream = SVS_MAKE_INTERFACE(nullptr, ops); + CATCH_REQUIRE(svs_index_load_stream(builder, &stream, error) == nullptr); + CATCH_REQUIRE(svs_error_get_code(error) == SVS_ERROR_INVALID_ARGUMENT); + } + + svs_index_free(index); + svs_index_builder_free(builder); + svs_algorithm_free(algorithm); + svs_error_free(error); +} + +CATCH_TEST_CASE("C API Stream Storage Round Trips", "[c_api][index][stream][storage]") { + const size_t NUM_VECTORS = 100; + const size_t DIMENSION = 32; + const size_t K = 5; + + std::vector data; + std::vector queries; + generate_test_data(data, NUM_VECTORS, DIMENSION); + generate_test_data(queries, 3, DIMENSION); + + CATCH_SECTION("Static round trip per storage kind") { + for (StorageKind kind : STORAGE_KINDS) { + CATCH_INFO("storage kind: " << storage_kind_name(kind)); + svs_error_h error = svs_error_create(); + + svs_storage_h build_storage = create_storage_case(kind, DIMENSION, error); + CATCH_REQUIRE(check_storage_support(build_storage, error)); + if (!storage_usable(build_storage)) { + svs_storage_free(build_storage); + svs_error_free(error); + continue; + } + + svs_algorithm_h algorithm = svs_algorithm_create_vamana(16, 32, 50, error); + CATCH_REQUIRE(algorithm != nullptr); + svs_index_builder_h builder = svs_index_builder_create( + SVS_DISTANCE_METRIC_EUCLIDEAN, DIMENSION, algorithm, error + ); + CATCH_REQUIRE(builder != nullptr); + CATCH_REQUIRE(svs_index_builder_set_threadpool( + builder, SVS_THREADPOOL_KIND_SINGLE_THREAD, 1, error + )); + CATCH_REQUIRE(svs_index_builder_set_storage(builder, build_storage, error)); + svs_storage_free(build_storage); + + svs_index_h index = svs_index_build(builder, data.data(), NUM_VECTORS, error); + CATCH_REQUIRE(index != nullptr); + CATCH_REQUIRE(svs_error_ok(error)); + + svs_search_results_t before = SVS_INIT_SEARCH_RESULTS(); + CATCH_REQUIRE(svs_index_search_topk( + index, queries.data(), 3, K, &before, nullptr, nullptr, error + )); + CATCH_REQUIRE(svs_error_ok(error)); + + MemoryStream stream; + svs_stream_interface_ops write_ops = + SVS_INIT_STREAM_OPS(nullptr, memory_stream_write); + svs_stream_interface out_stream = SVS_MAKE_INTERFACE(&stream, write_ops); + CATCH_REQUIRE(svs_index_save_stream(index, &out_stream, error)); + CATCH_REQUIRE(svs_error_ok(error)); + + // Load dispatch follows the builder's storage, not the stream's metadata, so + // the load builder must use the kind that saved it. + svs_storage_h load_storage = create_storage_case(kind, DIMENSION, error); + CATCH_REQUIRE(storage_usable(load_storage)); + svs_index_builder_h load_builder = svs_index_builder_create( + SVS_DISTANCE_METRIC_EUCLIDEAN, DIMENSION, algorithm, error + ); + CATCH_REQUIRE(load_builder != nullptr); + CATCH_REQUIRE(svs_index_builder_set_threadpool( + load_builder, SVS_THREADPOOL_KIND_SINGLE_THREAD, 1, error + )); + CATCH_REQUIRE(svs_index_builder_set_storage(load_builder, load_storage, error)); + svs_storage_free(load_storage); + + svs_stream_interface_ops read_ops = + SVS_INIT_STREAM_OPS(memory_stream_read, nullptr); + svs_stream_interface in_stream = SVS_MAKE_INTERFACE(&stream, read_ops); + svs_index_h loaded = svs_index_load_stream(load_builder, &in_stream, error); + CATCH_REQUIRE(loaded != nullptr); + CATCH_REQUIRE(svs_error_ok(error)); + + svs_search_results_t after = SVS_INIT_SEARCH_RESULTS(); + CATCH_REQUIRE(svs_index_search_topk( + loaded, queries.data(), 3, K, &after, nullptr, nullptr, error + )); + CATCH_REQUIRE(svs_error_ok(error)); + CATCH_REQUIRE(after.num_queries == before.num_queries); + for (size_t i = 0; i < before.num_queries * K; ++i) { + CATCH_REQUIRE(after.indices[i] == before.indices[i]); + CATCH_REQUIRE( + distance_within_tolerance(after.distances[i], before.distances[i]) + ); + } + + svs_search_results_free(&before); + svs_search_results_free(&after); + svs_index_free(loaded); + svs_index_free(index); + svs_index_builder_free(load_builder); + svs_index_builder_free(builder); + svs_algorithm_free(algorithm); + svs_error_free(error); + } + } + + CATCH_SECTION("Dynamic round trip, add_points, and search per storage kind") { + std::vector ids(NUM_VECTORS); + std::iota(ids.begin(), ids.end(), size_t{0}); + const size_t BLOCK_SIZE = 1024 * 1024; + + for (StorageKind kind : STORAGE_KINDS) { + CATCH_INFO("storage kind: " << storage_kind_name(kind)); + svs_error_h error = svs_error_create(); + + svs_storage_h build_storage = create_storage_case(kind, DIMENSION, error); + CATCH_REQUIRE(check_storage_support(build_storage, error)); + if (!storage_usable(build_storage)) { + svs_storage_free(build_storage); + svs_error_free(error); + continue; + } + + svs_algorithm_h algorithm = svs_algorithm_create_vamana(16, 32, 50, error); + CATCH_REQUIRE(algorithm != nullptr); + svs_index_builder_h builder = svs_index_builder_create( + SVS_DISTANCE_METRIC_EUCLIDEAN, DIMENSION, algorithm, error + ); + CATCH_REQUIRE(builder != nullptr); + CATCH_REQUIRE(svs_index_builder_set_threadpool( + builder, SVS_THREADPOOL_KIND_SINGLE_THREAD, 1, error + )); + CATCH_REQUIRE(svs_index_builder_set_storage(builder, build_storage, error)); + svs_storage_free(build_storage); + + svs_index_h index = svs_index_build_dynamic( + builder, data.data(), ids.data(), NUM_VECTORS, BLOCK_SIZE, error + ); + CATCH_REQUIRE(index != nullptr); + CATCH_REQUIRE(svs_error_ok(error)); + + svs_search_results_t before = SVS_INIT_SEARCH_RESULTS(); + CATCH_REQUIRE(svs_index_search_topk( + index, queries.data(), 3, K, &before, nullptr, nullptr, error + )); + CATCH_REQUIRE(svs_error_ok(error)); + + MemoryStream stream; + svs_stream_interface_ops write_ops = + SVS_INIT_STREAM_OPS(nullptr, memory_stream_write); + svs_stream_interface out_stream = SVS_MAKE_INTERFACE(&stream, write_ops); + CATCH_REQUIRE(svs_index_save_stream(index, &out_stream, error)); + CATCH_REQUIRE(svs_error_ok(error)); + + svs_storage_h load_storage = create_storage_case(kind, DIMENSION, error); + CATCH_REQUIRE(storage_usable(load_storage)); + svs_index_builder_h load_builder = svs_index_builder_create( + SVS_DISTANCE_METRIC_EUCLIDEAN, DIMENSION, algorithm, error + ); + CATCH_REQUIRE(load_builder != nullptr); + CATCH_REQUIRE(svs_index_builder_set_threadpool( + load_builder, SVS_THREADPOOL_KIND_SINGLE_THREAD, 1, error + )); + CATCH_REQUIRE(svs_index_builder_set_storage(load_builder, load_storage, error)); + svs_storage_free(load_storage); + + svs_stream_interface_ops read_ops = + SVS_INIT_STREAM_OPS(memory_stream_read, nullptr); + svs_stream_interface in_stream = SVS_MAKE_INTERFACE(&stream, read_ops); + svs_index_h loaded = + svs_index_load_stream_dynamic(load_builder, &in_stream, BLOCK_SIZE, error); + CATCH_REQUIRE(loaded != nullptr); + CATCH_REQUIRE(svs_error_ok(error)); + + svs_search_results_t after = SVS_INIT_SEARCH_RESULTS(); + CATCH_REQUIRE(svs_index_search_topk( + loaded, queries.data(), 3, K, &after, nullptr, nullptr, error + )); + CATCH_REQUIRE(svs_error_ok(error)); + CATCH_REQUIRE(after.num_queries == before.num_queries); + for (size_t i = 0; i < before.num_queries * K; ++i) { + CATCH_REQUIRE(after.indices[i] == before.indices[i]); + CATCH_REQUIRE( + distance_within_tolerance(after.distances[i], before.distances[i]) + ); + } + svs_search_results_free(&before); + svs_search_results_free(&after); + + std::vector new_data; + std::vector new_ids = {NUM_VECTORS, NUM_VECTORS + 1}; + generate_test_data(new_data, 2, DIMENSION); + size_t added_count = 0; + CATCH_REQUIRE(svs_index_dynamic_add_points( + loaded, new_data.data(), new_ids.data(), 2, &added_count, error + )); + CATCH_REQUIRE(added_count == 2); + CATCH_REQUIRE(svs_error_ok(error)); + + // An exact-vector probe must return the new id; a corrupted graph would still + // pass a count check. + std::vector probe_query( + new_data.begin(), new_data.begin() + static_cast(DIMENSION) + ); + svs_search_results_t probe_results = SVS_INIT_SEARCH_RESULTS(); + CATCH_REQUIRE(svs_index_search_topk( + loaded, probe_query.data(), 1, K, &probe_results, nullptr, nullptr, error + )); + CATCH_REQUIRE(svs_error_ok(error)); + bool found_new_id = std::any_of( + probe_results.indices, + probe_results.indices + K, + [new_id = new_ids[0]](size_t idx) { return idx == new_id; } + ); + CATCH_REQUIRE(found_new_id); + + svs_search_results_free(&probe_results); + svs_index_free(loaded); + svs_index_free(index); + svs_index_builder_free(load_builder); + svs_index_builder_free(builder); + svs_algorithm_free(algorithm); + svs_error_free(error); + } + } +} diff --git a/examples/c/CMakeLists.txt b/examples/c/CMakeLists.txt index a1d0ebacb..1315bac98 100644 --- a/examples/c/CMakeLists.txt +++ b/examples/c/CMakeLists.txt @@ -29,7 +29,7 @@ else() set(EXAMPLE_LINK_TARGET svs_c_api) endif() -foreach(EXAMPLE_NAME simple save_load dynamic) +foreach(EXAMPLE_NAME simple save_load save_load_stream dynamic) set(EXAMPLE_TARGET c_api_${EXAMPLE_NAME}) list(APPEND EXAMPLE_TARGETS ${EXAMPLE_TARGET}) add_executable(${EXAMPLE_TARGET} ${EXAMPLE_NAME}.c) diff --git a/examples/c/save_load_stream.c b/examples/c/save_load_stream.c new file mode 100644 index 000000000..079e180b7 --- /dev/null +++ b/examples/c/save_load_stream.c @@ -0,0 +1,309 @@ +/* + * Copyright 2026 Intel Corporation + * + * Licensed 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. + */ + +#include "svs/c/svs_c.h" +#include +#include +#include +#include + +#define NUM_VECTORS 10000 +#define NUM_QUERIES 1 +#define DIMENSION 128 +#define K 10 + +// `size` is the high-water mark reached while writing; `pos` is the read/write cursor and +// gets rewound to 0 between the save pass and the load pass. +typedef struct { + unsigned char* data; + size_t capacity; + size_t size; + size_t pos; +} memory_stream_t; + +memory_stream_t* memory_stream_create(size_t initial_capacity) { + memory_stream_t* stream = (memory_stream_t*)malloc(sizeof(memory_stream_t)); + if (!stream) { + return NULL; + } + stream->data = (unsigned char*)malloc(initial_capacity); + if (!stream->data) { + free(stream); + return NULL; + } + stream->capacity = initial_capacity; + stream->size = 0; + stream->pos = 0; + return stream; +} + +void memory_stream_free(memory_stream_t* stream) { + if (stream) { + free(stream->data); + free(stream); + } +} + +// A short read is legal here and the library calls again for the remainder; this callback +// always fills the whole request, but a caller need not. +static size_t memory_stream_read(void* self, void* buf, size_t n, svs_error_h out_err) { + (void)out_err; + memory_stream_t* stream = (memory_stream_t*)self; + size_t available = stream->size - stream->pos; + size_t to_read = (n < available) ? n : available; + if (to_read > 0) { + memcpy(buf, stream->data + stream->pos, to_read); + stream->pos += to_read; + } + return to_read; +} + +// A partial write must be reported as failure; otherwise the saved stream is silently +// truncated and the failure surfaces only much later, as a load error. +static bool +memory_stream_write(void* self, const void* buf, size_t n, svs_error_h out_err) { + (void)out_err; + memory_stream_t* stream = (memory_stream_t*)self; + while (stream->pos + n > stream->capacity) { + size_t new_capacity = stream->capacity * 2; + unsigned char* new_data = (unsigned char*)realloc(stream->data, new_capacity); + if (!new_data) { + return false; + } + stream->data = new_data; + stream->capacity = new_capacity; + } + memcpy(stream->data + stream->pos, buf, n); + stream->pos += n; + if (stream->pos > stream->size) { + stream->size = stream->pos; + } + return true; +} + +void generate_random_data(float* data, size_t count, size_t dim) { + for (size_t i = 0; i < count * dim; i++) { + data[i] = (float)rand() / RAND_MAX; + } +} + +int main() { + int ret = 0; + srand(time(NULL)); + svs_error_h error = svs_error_create(); + + float* data = NULL; + float* queries = NULL; + svs_algorithm_h algorithm = NULL; + svs_storage_h storage = NULL; + svs_index_builder_h builder = NULL; + svs_index_h index = NULL; + svs_search_results_t results = SVS_INIT_SEARCH_RESULTS(); + memory_stream_t* stream = NULL; + svs_search_results_t loaded_results = SVS_INIT_SEARCH_RESULTS(); + + // Allocate random data + data = (float*)malloc(NUM_VECTORS * DIMENSION * sizeof(float)); + queries = (float*)malloc(NUM_QUERIES * DIMENSION * sizeof(float)); + + if (!data || !queries) { + fprintf(stderr, "Failed to allocate memory\n"); + ret = 1; + goto cleanup; + } + + generate_random_data(data, NUM_VECTORS, DIMENSION); + generate_random_data(queries, NUM_QUERIES, DIMENSION); + + // Create Vamana algorithm + algorithm = svs_algorithm_create_vamana(64, 128, 100, error); + if (!algorithm) { + fprintf(stderr, "Failed to create algorithm: %s\n", svs_error_get_message(error)); + ret = 1; + goto cleanup; + } + + // Create storage (simple float32) + storage = svs_storage_create_simple(SVS_DATA_TYPE_FLOAT32, error); + if (!storage) { + fprintf(stderr, "Failed to create storage: %s\n", svs_error_get_message(error)); + ret = 1; + goto cleanup; + } + + // Create index builder + builder = svs_index_builder_create( + SVS_DISTANCE_METRIC_EUCLIDEAN, DIMENSION, algorithm, error + ); + if (!builder) { + fprintf( + stderr, "Failed to create index builder: %s\n", svs_error_get_message(error) + ); + ret = 1; + goto cleanup; + } + + if (!svs_index_builder_set_storage(builder, storage, error)) { + fprintf(stderr, "Failed to set storage: %s\n", svs_error_get_message(error)); + ret = 1; + goto cleanup; + } + + // Build index + printf("Building index with %d vectors of dimension %d...\n", NUM_VECTORS, DIMENSION); + index = svs_index_build(builder, data, NUM_VECTORS, error); + if (!index) { + fprintf(stderr, "Failed to build index: %s\n", svs_error_get_message(error)); + ret = 1; + goto cleanup; + } + printf("Index built successfully!\n"); + + // Search + printf("Searching %d queries for top-%d neighbors...\n", NUM_QUERIES, K); + if (!svs_index_search_topk( + index, + queries, + NUM_QUERIES, + K, + &results, + NULL /* search_params */, + NULL /* id_filter */, + error + )) { + fprintf(stderr, "Failed to search index: %s\n", svs_error_get_message(error)); + ret = 1; + goto cleanup; + } + printf("Search completed successfully!\n"); + + // Create in-memory stream for saving + stream = memory_stream_create(1024 * 1024); + if (!stream) { + fprintf(stderr, "Failed to create stream\n"); + ret = 1; + goto cleanup; + } + + // Construct stream interface and save through it + static svs_stream_ops_t stream_ops = + SVS_INIT_STREAM_OPS(memory_stream_read, memory_stream_write); + svs_stream_t stream_iface = SVS_MAKE_INTERFACE(stream, stream_ops); + + printf("Saving index to in-memory stream...\n"); + if (!svs_index_save_stream(index, &stream_iface, error)) { + fprintf( + stderr, "Failed to save index to stream: %s\n", svs_error_get_message(error) + ); + ret = 1; + goto cleanup; + } + printf("Index saved successfully! Stream size: %zu bytes\n", stream->size); + + svs_index_free(index); + index = NULL; + + // Reset stream position for reading + stream->pos = 0; + + // Load the index from the stream + printf("Loading index from in-memory stream...\n"); + index = svs_index_load_stream(builder, &stream_iface, error); + if (!index) { + fprintf( + stderr, "Failed to load index from stream: %s\n", svs_error_get_message(error) + ); + ret = 1; + goto cleanup; + } + printf("Index loaded successfully!\n"); + + // Search the loaded index + printf( + "Searching loaded index for %d queries for top-%d neighbors...\n", NUM_QUERIES, K + ); + if (!svs_index_search_topk( + index, + queries, + NUM_QUERIES, + K, + &loaded_results, + NULL /* search_params */, + NULL /* id_filter */, + error + )) { + fprintf( + stderr, "Failed to search loaded index: %s\n", svs_error_get_message(error) + ); + ret = 1; + goto cleanup; + } + printf("Search on loaded index completed successfully!\n"); + + // Compare results + if (results.num_queries != loaded_results.num_queries) { + fprintf( + stderr, "Mismatch in number of queries between original and loaded results\n" + ); + ret = 1; + goto cleanup; + } + + size_t offset = 0; + for (size_t q = 0; q < results.num_queries; q++) { + size_t count = results.offsets[q + 1] - results.offsets[q]; + size_t loaded_count = loaded_results.offsets[q + 1] - loaded_results.offsets[q]; + if (count != loaded_count) { + fprintf(stderr, "Mismatch in number of results for query %zu\n", q); + ret = 1; + goto cleanup; + } + printf("Query %zu results:\n", q); + for (size_t i = 0; i < count; i++) { + if (results.indices[offset + i] != loaded_results.indices[offset + i]) { + fprintf( + stderr, "Mismatch in neighbor indices for query %zu, result %zu\n", q, i + ); + ret = 1; + goto cleanup; + } + printf( + " [%zu] id=%zu, distance=%.4f, diff=%.4f\n", + i, + results.indices[offset + i], + results.distances[offset + i], + results.distances[offset + i] - loaded_results.distances[offset + i] + ); + } + offset += count; + } + + printf("Done!\n"); + +cleanup: + svs_search_results_free(&results); + svs_search_results_free(&loaded_results); + svs_index_free(index); + svs_index_builder_free(builder); + svs_storage_free(storage); + svs_algorithm_free(algorithm); + free(data); + free(queries); + memory_stream_free(stream); + svs_error_free(error); + + return ret; +} diff --git a/include/svs/orchestrators/dynamic_vamana.h b/include/svs/orchestrators/dynamic_vamana.h index d2e99b0aa..803fcf94a 100644 --- a/include/svs/orchestrators/dynamic_vamana.h +++ b/include/svs/orchestrators/dynamic_vamana.h @@ -384,22 +384,42 @@ class DynamicVamana : public manager::IndexManager { } // Assembly from stream + /// + /// @brief Assemble a DynamicVamana index from a serialized stream. + /// + /// @tparam QueryTypes The set of query element types supported by the resulting + /// index. + /// @tparam Data The dataset type to load. + /// @tparam Distance Distance functor or ``svs::DistanceType`` enum. + /// @tparam ThreadPoolProto Thread pool type or size_t. + /// @tparam DataAllocator The type of allocator used for the dataset. + /// @tparam GraphAllocator The type of allocator used for the graph. Defaults to an + /// exact-size ``HugepageAllocator``; a blocked default commits a full block. + /// + /// @param stream Stream containing the serialized index. + /// @param distance Distance functor or enum. + /// @param threadpool_proto Thread pool or number of threads to use. + /// @param data_allocator Allocator instance to use for the dataset. + /// @param graph_allocator Allocator instance to use for the graph. + /// template < manager::QueryTypeDefinition QueryTypes, typename Data, typename Distance, typename ThreadPoolProto, - typename... DataLoaderArgs> + typename DataAllocator = typename Data::allocator_type, + typename GraphAllocator = HugepageAllocator> static DynamicVamana assemble( std::istream& stream, const Distance& distance, ThreadPoolProto threadpool_proto, - DataLoaderArgs&&... data_args + const DataAllocator& data_allocator = {}, + const GraphAllocator& graph_allocator = {} ) { auto deserializer = svs::lib::detail::Deserializer::build(stream); if (deserializer.is_native()) { auto threadpool = threads::as_threadpool(std::move(threadpool_proto)); - using GraphType = svs::GraphLoader<>::return_type; + using GraphType = graphs::SimpleGraph; if constexpr (std::is_same_v, DistanceType>) { auto dispatcher = DistanceDispatcher(distance); return dispatcher([&](auto distance_function) { @@ -407,12 +427,12 @@ class DynamicVamana : public manager::IndexManager { index::vamana::auto_dynamic_assemble( stream, // lazy graph loader - [&]() -> GraphType { return GraphType::load(stream); }, + [&]() -> GraphType { + return GraphType::load(stream, graph_allocator); + }, // lazy data loader [&]() -> Data { - return lib::load_from_stream( - stream, SVS_FWD(data_args)... - ); + return lib::load_from_stream(stream, data_allocator); }, distance_function, std::move(threadpool) @@ -424,12 +444,12 @@ class DynamicVamana : public manager::IndexManager { index::vamana::auto_dynamic_assemble( stream, // lazy graph loader - [&]() -> GraphType { return GraphType::load(stream); }, + [&]() -> GraphType { + return GraphType::load(stream, graph_allocator); + }, // lazy data loader [&]() -> Data { - return lib::load_from_stream( - stream, SVS_FWD(data_args)... - ); + return lib::load_from_stream(stream, data_allocator); }, distance, std::move(threadpool) @@ -460,8 +480,8 @@ class DynamicVamana : public manager::IndexManager { return assemble( config_path, - svs::GraphLoader{graph_path}, - lib::load_from_disk(data_path, SVS_FWD(data_args)...), + svs::GraphLoader{graph_path, graph_allocator}, + lib::load_from_disk(data_path, data_allocator), distance, threads::as_threadpool(std::move(threadpool_proto)), false diff --git a/include/svs/orchestrators/vamana.h b/include/svs/orchestrators/vamana.h index 6f27d6728..6618dfbff 100644 --- a/include/svs/orchestrators/vamana.h +++ b/include/svs/orchestrators/vamana.h @@ -480,62 +480,103 @@ class Vamana : public manager::IndexManager { } } - // Assembly from stream + /// + /// @brief Assemble a Vamana index from an in-memory, view-backed stream. + /// + /// @param stream The stream to load from. See ``svs::Vamana::save``. + /// @param distance The distance functor or ``svs::DistanceType`` enum to use for + /// similarity search computations. + /// @param threadpool_proto Precursor for the thread pool to use. Can either be an + /// acceptable thread pool instance or an integer specifying the number of + /// threads to use. + /// @param data_args Forwarded to the dataset loader. An allocator passed here must be + /// bound to ``stream``. + /// + /// The stream must be an in-memory stream in native format; the returned index views + /// its buffer directly, so the stream must outlive the index. + /// + /// @copydoc threadpool_requirements + /// + /// @sa save, build + /// template < manager::QueryTypeDefinition QueryTypes, typename Data, typename Distance, typename ThreadPoolProto, typename... DataLoaderArgs> + requires is_view_type_v static Vamana assemble( std::istream& stream, const Distance& distance, ThreadPoolProto threadpool_proto, DataLoaderArgs&&... data_args + ) { + auto deserializer = svs::lib::detail::Deserializer::build(stream); + if (!deserializer.is_native()) { + throw ANNEXCEPTION( + "Cannot load a view-backed Vamana index from a directory archive. " + "Directory archives are unpacked to a temporary directory and cannot " + "back a view; use the native stream format instead." + ); + } + + using Allocator = lib::rebind_allocator_t; + using GraphType = graphs::SimpleGraph; + auto load_graph = [&]() -> GraphType { return GraphType::load(stream); }; + auto load_data = [&]() -> Data { + return lib::load_from_stream(stream, SVS_FWD(data_args)...); + }; + + return assemble_native( + stream, load_graph, load_data, distance, std::move(threadpool_proto) + ); + } + + /// + /// @brief Assemble a Vamana index from a stream. + /// + /// @param stream The stream to load from. See ``svs::Vamana::save``. + /// @param distance The distance functor or ``svs::DistanceType`` enum to use for + /// similarity search computations. + /// @param threadpool_proto Precursor for the thread pool to use. Can either be an + /// acceptable thread pool instance or an integer specifying the number of + /// threads to use. + /// @param data_allocator Allocator to use for the loaded data. + /// @param graph_allocator Allocator to use for the loaded graph. + /// + /// @copydoc threadpool_requirements + /// + /// @sa save, build + /// + template < + manager::QueryTypeDefinition QueryTypes, + typename Data, + typename Distance, + typename ThreadPoolProto, + typename DataAllocator = typename Data::allocator_type, + typename GraphAllocator = HugepageAllocator> + requires(!is_view_type_v) + static Vamana assemble( + std::istream& stream, + const Distance& distance, + ThreadPoolProto threadpool_proto, + const DataAllocator& data_allocator = {}, + const GraphAllocator& graph_allocator = {} ) { auto deserializer = svs::lib::detail::Deserializer::build(stream); if (deserializer.is_native()) { - auto threadpool = threads::as_threadpool(std::move(threadpool_proto)); - - using GraphType = std::conditional_t< - is_view_type_v, - graphs::SimpleGraph< - uint32_t, - lib::rebind_allocator_t>, - GraphLoader<>::return_type>; - - if constexpr (std::is_same_v) { - auto dispatcher = DistanceDispatcher(distance); - return dispatcher([&](auto distance_function) { - return make_vamana>( - AssembleTag(), - stream, - // lazy-loader - [&]() -> GraphType { return GraphType::load(stream); }, - // lazy-loader - [&]() -> Data { - return lib::load_from_stream( - stream, SVS_FWD(data_args)... - ); - }, - distance_function, - std::move(threadpool) - ); - }); - } else { - return make_vamana>( - AssembleTag(), - stream, - // lazy-loader - [&]() -> GraphType { return GraphType::load(stream); }, - // lazy-loader - [&]() -> Data { - return lib::load_from_stream(stream, SVS_FWD(data_args)...); - }, - distance, - std::move(threadpool) - ); - } + using GraphType = graphs::SimpleGraph; + auto load_graph = [&]() -> GraphType { + return GraphType::load(stream, graph_allocator); + }; + auto load_data = [&]() -> Data { + return lib::load_from_stream(stream, data_allocator); + }; + + return assemble_native( + stream, load_graph, load_data, distance, std::move(threadpool_proto) + ); } else { namespace fs = std::filesystem; lib::UniqueTempDirectory tempdir{"svs_vamana_load"}; @@ -560,8 +601,8 @@ class Vamana : public manager::IndexManager { return assemble( config_path, - svs::GraphLoader{graph_path}, - lib::load_from_disk(data_path, SVS_FWD(data_args)...), + svs::GraphLoader{graph_path, graph_allocator}, + lib::load_from_disk(data_path, data_allocator), distance, threads::as_threadpool(std::move(threadpool_proto)) ); @@ -716,6 +757,45 @@ class Vamana : public manager::IndexManager { svs::index::vamana::VamanaIndexParameters parameters() const { return impl_->parameters(); } + + private: + template < + manager::QueryTypeDefinition QueryTypes, + typename GraphLoaderFn, + typename DataLoaderFn, + typename Distance, + typename ThreadPoolProto> + static Vamana assemble_native( + std::istream& stream, + const GraphLoaderFn& load_graph, + const DataLoaderFn& load_data, + const Distance& distance, + ThreadPoolProto threadpool_proto + ) { + auto threadpool = threads::as_threadpool(std::move(threadpool_proto)); + if constexpr (std::is_same_v) { + auto dispatcher = DistanceDispatcher(distance); + return dispatcher([&](auto distance_function) { + return make_vamana>( + AssembleTag(), + stream, + load_graph, + load_data, + distance_function, + std::move(threadpool) + ); + }); + } else { + return make_vamana>( + AssembleTag(), + stream, + load_graph, + load_data, + distance, + std::move(threadpool) + ); + } + } }; /// diff --git a/tests/svs/index/vamana/index.cpp b/tests/svs/index/vamana/index.cpp index 284bc68d2..d6c9b67ff 100644 --- a/tests/svs/index/vamana/index.cpp +++ b/tests/svs/index/vamana/index.cpp @@ -423,6 +423,53 @@ CATCH_TEST_CASE("Vamana Index Save and Load", "[vamana][index][saveload]") { CATCH_REQUIRE(modified_distance == Catch::Approx(0.0).epsilon(1e-5)); } + CATCH_SECTION("Load view with explicit stream-bound allocator") { + using ViewData_t = + svs::data::SimpleData>; + + // Save the full index to a stringstream. + auto ss = std::stringstream{}; + index.save(ss); + + // Load the Vamana index from the stream, passing the allocator explicitly. + ss.seekg(0); + auto loaded_index = svs::Vamana::assemble( + ss, + distance_function, + svs::threads::DefaultThreadPool(1), + svs::io::MemoryStreamAllocator{ss} + ); + + CATCH_REQUIRE(loaded_index.size() == index.size()); + CATCH_REQUIRE(loaded_index.dimensions() == index.dimensions()); + } + + CATCH_SECTION("Load view from directory archive throws") { + using ViewData_t = + svs::data::SimpleData>; + + std::stringstream ss; + { + svs::lib::UniqueTempDirectory tempdir{"svs_vamana_save"}; + const auto config_dir = tempdir.get() / "config"; + const auto graph_dir = tempdir.get() / "graph"; + const auto data_dir = tempdir.get() / "data"; + std::filesystem::create_directories(config_dir); + std::filesystem::create_directories(graph_dir); + std::filesystem::create_directories(data_dir); + index.save(config_dir, graph_dir, data_dir); + svs::lib::DirectoryArchiver::pack(tempdir, ss); + } + + // A directory archive cannot back a view-backed Data; assemble must throw. + CATCH_REQUIRE_THROWS_AS( + (svs::Vamana::assemble( + ss, distance_function, svs::threads::DefaultThreadPool(1) + )), + svs::ANNException + ); + } + CATCH_SECTION("Load with SimpleDataView pointing to memory mapped file") { // We will load the Vamana index's data as a SimpleDataView directly from the // stream, without copying.