diff --git a/cpp/vector_engine_ffi.cpp b/cpp/vector_engine_ffi.cpp index 138bfeb..0574c47 100644 --- a/cpp/vector_engine_ffi.cpp +++ b/cpp/vector_engine_ffi.cpp @@ -1,6 +1,16 @@ #include "vector_engine_ffi.h" +#include "vector_engine.h" -struct NativeVectorEngine {}; +#include +#include + +struct NativeVectorEngine { + VectorEngine engine; +}; + +struct NativeSearchResults { + std::vector results; +}; extern "C" NativeVectorEngine* native_vector_engine_new() { return new NativeVectorEngine(); @@ -9,3 +19,33 @@ extern "C" NativeVectorEngine* native_vector_engine_new() { extern "C" void native_vector_engine_free(NativeVectorEngine* engine) { delete engine; } + +extern "C" size_t native_search_results_len(const NativeSearchResults* results) { + return results->results.size(); +} + +extern "C" const char* native_search_results_id_at(const NativeSearchResults* results, size_t index) { + return results->results[index].id.c_str(); +} + +extern "C" float native_search_results_score_at(const NativeSearchResults* results, size_t index) { + return results->results[index].score; +} + +extern "C" void native_search_results_free(NativeSearchResults* results) { + delete results; +} + +extern "C" void native_vector_engine_insert(NativeVectorEngine* engine, const char* id, const float* vector, size_t len) { + engine->engine.insert(std::string(id), vector, len); +} + +extern "C" bool native_vector_engine_delete(NativeVectorEngine* engine, const char* id) { + return engine->engine.erase(std::string(id)); +} + +extern "C" NativeSearchResults* native_vector_engine_search(const NativeVectorEngine* engine, const float* query, size_t len, size_t k) { + auto* results = new NativeSearchResults(); + results->results = engine->engine.search(query, len, k); + return results; +} diff --git a/cpp/vector_engine_ffi.h b/cpp/vector_engine_ffi.h index 4cb5d5a..ee00e1b 100644 --- a/cpp/vector_engine_ffi.h +++ b/cpp/vector_engine_ffi.h @@ -8,29 +8,13 @@ struct NativeSearchResults; extern "C" { - // subject to change NativeVectorEngine* native_vector_engine_new(); void native_vector_engine_free(NativeVectorEngine* engine); // insert / delete / search - void native_vector_engine_insert( - NativeVectorEngine* engine, - const char* id, - const float* vector, - size_t len - ); - - bool native_vector_engine_delete( - NativeVectorEngine* engine, - const char* id - ); - - NativeSearchResults* native_vector_engine_search( - const NativeVectorEngine* engine, - const float* query, - size_t len, - size_t k - ); + void native_vector_engine_insert(NativeVectorEngine* engine, const char* id, const float* vector, size_t len); + bool native_vector_engine_delete(NativeVectorEngine* engine, const char* id); + NativeSearchResults* native_vector_engine_search(NativeVectorEngine* engine, const float* query, size_t len, size_t k); // working with pointers to send info back to rust size_t native_search_results_len(const NativeSearchResults* results); diff --git a/crates/vdb-ffi/Cargo.lock b/crates/vdb-ffi/Cargo.lock index e826eb0..882e120 100644 --- a/crates/vdb-ffi/Cargo.lock +++ b/crates/vdb-ffi/Cargo.lock @@ -37,6 +37,16 @@ version = "2.11.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c4512299f36f043ab09a583e57bceb5a5aab7a73db1805848e8fef3c9e8c78b3" +[[package]] +name = "cc" +version = "1.2.62" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a1dce859f0832a7d088c4f1119888ab94ef4b5d6795d1ce05afb7fe159d79f98" +dependencies = [ + "find-msvc-tools", + "shlex", +] + [[package]] name = "cexpr" version = "0.6.0" @@ -69,6 +79,12 @@ version = "1.15.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "48c757948c5ede0e46177b7add2e67155f70e33c07fea8284df6576da70b3719" +[[package]] +name = "find-msvc-tools" +version = "0.1.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5baebc0774151f905a1a2cc41989300b1e6fbb29aff0ceffa1064fdd3088d582" + [[package]] name = "glob" version = "0.3.3" @@ -219,6 +235,7 @@ name = "vdb-ffi" version = "0.1.0" dependencies = [ "bindgen", + "cc", ] [[package]] diff --git a/crates/vdb-ffi/Cargo.toml b/crates/vdb-ffi/Cargo.toml index c0445f0..b9cfc95 100644 --- a/crates/vdb-ffi/Cargo.toml +++ b/crates/vdb-ffi/Cargo.toml @@ -8,3 +8,4 @@ path = "src/lib.rs" [build-dependencies] bindgen = "0.72" +cc = "1" diff --git a/crates/vdb-ffi/build.rs b/crates/vdb-ffi/build.rs index c87be6e..9ca089d 100644 --- a/crates/vdb-ffi/build.rs +++ b/crates/vdb-ffi/build.rs @@ -1,8 +1,19 @@ -use std::{env, path::PathBuf}; +use std::path::PathBuf; fn main() { - println!("cargo:rerun-if-changed=../../cpp/vector_engine_ffi.h"); + // linking the compiled cpp files + let cpp_dir = PathBuf::from("../../cpp"); + let mut native_build = cc::Build::new(); + native_build + .cpp(true) + .std("c++17") + .include(&cpp_dir) + .file(cpp_dir.join("vector_engine.cpp")) + .file(cpp_dir.join("vector_engine_ffi.cpp")); + native_build.compile("vector_engine_native"); + + // generates rust declarations based of cpp ffi let bindings = bindgen::Builder::default() .header("../../cpp/vector_engine_ffi.h") .clang_arg("-xc++") diff --git a/crates/vdb-ffi/src/engine.rs b/crates/vdb-ffi/src/engine.rs index 759d4a4..3f629c0 100644 --- a/crates/vdb-ffi/src/engine.rs +++ b/crates/vdb-ffi/src/engine.rs @@ -21,8 +21,7 @@ pub struct FfiSearchResults { handle: *mut NativeSearchResults, } -// The wrapper owns the native handle and is always accessed behind higher-level -// synchronization in the API layer. +// Send is built into Rust as a trait, allows FfiVectorEngine to be moved between threads safely. unsafe impl Send for FfiVectorEngine {} impl FfiVectorEngine {