diff --git a/Cargo.lock b/Cargo.lock index fbbc655..3471623 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -135,6 +135,35 @@ dependencies = [ "tracing", ] +[[package]] +name = "axum-test" +version = "18.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0ce2a8627e8d8851f894696b39f2b67807d6375c177361d376173ace306a21e2" +dependencies = [ + "anyhow", + "axum", + "bytes", + "bytesize", + "cookie", + "expect-json", + "http 1.4.0", + "http-body-util", + "hyper", + "hyper-util", + "mime", + "pretty_assertions", + "reserve-port", + "rust-multipart-rfc7578_2", + "serde", + "serde_json", + "serde_urlencoded", + "smallvec", + "tokio", + "tower", + "url", +] + [[package]] name = "backtrace" version = "0.3.76" @@ -237,6 +266,12 @@ version = "1.11.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1e748733b7cbc798e1434b6ac524f0c1ff2ab456fe201501e6497c8417a4fc33" +[[package]] +name = "bytesize" +version = "2.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6bd91ee7b2422bcb158d90ef4d14f75ef67f340943fc4149891dcce8f8b972a3" + [[package]] name = "bzip2-sys" version = "0.1.13+1.0.8" @@ -354,6 +389,16 @@ dependencies = [ "static_assertions", ] +[[package]] +name = "cookie" +version = "0.18.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4ddef33a339a91ea89fb53151bd0a4689cfce27055c291dfa69945475d22c747" +dependencies = [ + "time", + "version_check", +] + [[package]] name = "core-foundation" version = "0.9.4" @@ -443,10 +488,27 @@ checksum = "d7a1e2f27636f116493b8b860f5546edb47c8d8f8ea73e1d2a20be88e28d1fea" name = "defs" version = "0.1.0" dependencies = [ + "axum", "serde", + "snafu", "uuid", ] +[[package]] +name = "deranged" +version = "0.5.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7cd812cc2bc1d69d4764bd80df88b4317eaef9e773c75226407d9bc0876b211c" +dependencies = [ + "powerfmt", +] + +[[package]] +name = "diff" +version = "0.1.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "56254986775e3233ffa9c4d7d3faaf6d36a2c09d30b20687e9f88bc8bafc16c8" + [[package]] name = "digest" version = "0.10.7" @@ -480,6 +542,15 @@ version = "1.15.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "48c757948c5ede0e46177b7add2e67155f70e33c07fea8284df6576da70b3719" +[[package]] +name = "email_address" +version = "0.2.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e079f19b08ca6239f47f8ba8509c11cf3ea30095831f7fed61441475edd8c449" +dependencies = [ + "serde", +] + [[package]] name = "encoding_rs" version = "0.8.35" @@ -495,6 +566,17 @@ version = "1.0.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" +[[package]] +name = "erased-serde" +version = "0.4.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d2add8a07dd6a8d93ff627029c51de145e12686fbc36ecb298ac22e74cf02dec" +dependencies = [ + "serde", + "serde_core", + "typeid", +] + [[package]] name = "errno" version = "0.3.14" @@ -505,6 +587,35 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "expect-json" +version = "1.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "869f97f4abe8e78fc812a94ad6b721d72c4fb5532877c79610f2c238d7ccf6c4" +dependencies = [ + "chrono", + "email_address", + "expect-json-macros", + "num", + "regex", + "serde", + "serde_json", + "thiserror", + "typetag", + "uuid", +] + +[[package]] +name = "expect-json-macros" +version = "1.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c0637949cd816934f3b7aab44ff98e7ec1fb903c379e07dcb9eac943ec33499e" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "eyre" version = "0.6.12" @@ -777,11 +888,13 @@ version = "0.1.0" dependencies = [ "api", "axum", + "axum-test", "defs", "index", "serde", "serde_json", "storage", + "tempfile", "tokio", "tracing", ] @@ -1087,6 +1200,15 @@ dependencies = [ "serde_core", ] +[[package]] +name = "inventory" +version = "0.3.24" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a4f0c30c76f2f4ccee3fe55a2435f691ca00c0e4bd87abe4f4a851b1d4dac39b" +dependencies = [ + "rustversion", +] + [[package]] name = "ipnet" version = "2.12.0" @@ -1386,6 +1508,76 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "num" +version = "0.4.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "35bd024e8b2ff75562e5f34e7f4905839deb4b22955ef5e73d2fea1b9813cb23" +dependencies = [ + "num-bigint", + "num-complex", + "num-integer", + "num-iter", + "num-rational", + "num-traits", +] + +[[package]] +name = "num-bigint" +version = "0.4.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a5e44f723f1133c9deac646763579fdb3ac745e418f2a7af9cd0c431da1f20b9" +dependencies = [ + "num-integer", + "num-traits", +] + +[[package]] +name = "num-complex" +version = "0.4.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "73f88a1307638156682bada9d7604135552957b7818057dcef22705b4d509495" +dependencies = [ + "num-traits", +] + +[[package]] +name = "num-conv" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "521739c6d2bac4aa25192232afe6841231376b2b26d4d9fae5ecf8ca5772e441" + +[[package]] +name = "num-integer" +version = "0.1.46" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7969661fd2958a5cb096e56c8e1ad0444ac2bbcd0061bd28660485a44879858f" +dependencies = [ + "num-traits", +] + +[[package]] +name = "num-iter" +version = "0.1.45" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1429034a0490724d0075ebb2bc9e875d6503c3cf69e235a8941aa757d83ef5bf" +dependencies = [ + "autocfg", + "num-integer", + "num-traits", +] + +[[package]] +name = "num-rational" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f83d14da390562dca69fc84082e73e548e1ad308d24accdedd2720017cb37824" +dependencies = [ + "num-bigint", + "num-integer", + "num-traits", +] + [[package]] name = "num-traits" version = "0.2.19" @@ -1565,6 +1757,12 @@ dependencies = [ "zerovec", ] +[[package]] +name = "powerfmt" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "439ee305def115ba05938db6eb1644ff94165c5ab5e9420d1c1bcedbba909391" + [[package]] name = "ppv-lite86" version = "0.2.21" @@ -1574,6 +1772,16 @@ dependencies = [ "zerocopy", ] +[[package]] +name = "pretty_assertions" +version = "1.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3ae130e2f271fbc2ac3a40fb1d07180839cdbbe443c7a27e1e3c13c5cac0116d" +dependencies = [ + "diff", + "yansi", +] + [[package]] name = "prettyplease" version = "0.2.37" @@ -1826,6 +2034,15 @@ dependencies = [ "web-sys", ] +[[package]] +name = "reserve-port" +version = "2.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "94070964579245eb2f76e62a7668fe87bd9969ed6c41256f3bf614e3323dd3cc" +dependencies = [ + "thiserror", +] + [[package]] name = "ring" version = "0.17.14" @@ -1850,6 +2067,21 @@ dependencies = [ "librocksdb-sys", ] +[[package]] +name = "rust-multipart-rfc7578_2" +version = "0.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c839d037155ebc06a571e305af66ff9fd9063a6e662447051737e1ac75beea41" +dependencies = [ + "bytes", + "futures-core", + "futures-util", + "http 1.4.0", + "mime", + "rand", + "thiserror", +] + [[package]] name = "rustc-demangle" version = "0.1.27" @@ -2327,6 +2559,26 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "thiserror" +version = "2.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4288b5bcbc7920c07a1149a35cf9590a2aa808e0bc1eafaade0b80947865fbc4" +dependencies = [ + "thiserror-impl", +] + +[[package]] +name = "thiserror-impl" +version = "2.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ebc4ee7f67670e9b64d05fa4253e753e016c6c95ff35b89b7941d6b856dec1d5" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "thread_local" version = "1.1.9" @@ -2336,6 +2588,37 @@ dependencies = [ "cfg-if", ] +[[package]] +name = "time" +version = "0.3.47" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "743bd48c283afc0388f9b8827b976905fb217ad9e647fae3a379a9283c4def2c" +dependencies = [ + "deranged", + "itoa", + "num-conv", + "powerfmt", + "serde_core", + "time-core", + "time-macros", +] + +[[package]] +name = "time-core" +version = "0.1.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7694e1cfe791f8d31026952abf09c69ca6f6fa4e1a1229e18988f06a04a12dca" + +[[package]] +name = "time-macros" +version = "0.2.27" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2e70e4c5a0e0a8a4823ad65dfe1a6930e4f4d756dcd9dd7939022b5e8c501215" +dependencies = [ + "num-conv", + "time-core", +] + [[package]] name = "tinystr" version = "0.8.2" @@ -2628,12 +2911,42 @@ dependencies = [ "uuid", ] +[[package]] +name = "typeid" +version = "1.0.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bc7d623258602320d5c55d1bc22793b57daff0ec7efc270ea7d55ce1d5f5471c" + [[package]] name = "typenum" version = "1.19.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "562d481066bde0658276a35467c4af00bdc6ee726305698a55b86e61d7ad82bb" +[[package]] +name = "typetag" +version = "0.2.22" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c5a897b12c6c1151ad0b138b8db50252dc301f93bc3b027db05eec82aeed298c" +dependencies = [ + "erased-serde", + "inventory", + "once_cell", + "serde", + "typetag-impl", +] + +[[package]] +name = "typetag-impl" +version = "0.2.22" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cf808357c6ed7e13ba0f3277ec8d8f21b2d501274895104263985330c726c1c5" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "unicase" version = "2.9.0" @@ -3205,6 +3518,12 @@ dependencies = [ "rustix", ] +[[package]] +name = "yansi" +version = "1.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cfe53a6657fd280eaa890a3bc59152892ffa3e30101319d168b781ed6529b049" + [[package]] name = "yoke" version = "0.8.1" diff --git a/client/python/USAGE.md b/client/python/USAGE.md index 082bfd0..be2e5d5 100644 --- a/client/python/USAGE.md +++ b/client/python/USAGE.md @@ -39,6 +39,30 @@ The client supports usage as a context manager, which automatically closes the u Example available in: ```examples/context_manager_usage.py``` +### Async Client Support + +For async applications, use `AsyncVortexDB`. It mirrors the synchronous client API and uses `grpc.aio` under the hood, including full support for `batch_insert` and `batch_search`. + +Examples available in: +```examples/async_usage.py``` & ```examples/async_batch_usage.py``` + +```python +async with AsyncVortexDB( + grpc_url="localhost:50051", + api_key="your-api-key", +) as db: + point_id = await db.insert( + vector=DenseVector([0.1, 0.2, 0.3]), + payload=Payload.text("hello async vortex"), + ) +``` + +### Batch Insertion and Search Support + +Both `VortexDB` and `AsyncVortexDB` support batch insertion and batch search queries. +Methods of usage and examples available in: +```examples/batch_insert_usage.py``` & ```examples/search_query_usage.py``` & ```examples/async_batch_usage.py``` + --- ## Client API @@ -47,6 +71,27 @@ Example available in: Main client class for interacting with the VortexDB gRPC server. +### `AsyncVortexDB` + +Async client class for I/O-heavy applications. It has the same constructor and method names as `VortexDB`, but methods are awaitable: + +``` +await db.insert(...) +await db.batch_insert(...) +await db.get(...) +await db.search(...) +await db.batch_search(...) +await db.delete(...) +await db.close() +``` + +It also supports async context manager usage: + +``` +async with AsyncVortexDB(...) as db: + ... +``` + #### **Constructor** ``` @@ -78,6 +123,22 @@ Raises --- +#### **Batch Insert** + +Insert multiple vectors with payloads in a single request +``` +batch_insert(*, items: list[tuple[DenseVector, Payload]]) -> list[str] +``` + +Returns +- List of `point_id` (UUID string) + +Raises +- `TypeError` if input structure is invalid +- gRPC-mapped errors (see Error Handling) + +--- + #### **Get** Fetch a point by its ID @@ -112,6 +173,32 @@ Raises --- +#### **Batch Search** + +Search for nearest neighbours for multiple queries in a single request +``` +batch_search( + *, + queries, + similarity: Similarity | None = None, + limit: int | None = None, +) -> list[list[str]] +``` + +Returns +- `TypeError` for invalid query formats +- `ValueError` if required parameters are missing + +Supported Input Formats: +The `queries` parameter is flexible and supports multiple formats: +- List of `SearchQuery` objects +- List of `(DenseVector, Similarity, Limit)` tuples +- List of `(DenseVector, Similarity)` tuples with a global `Limit` +- List of `(DenseVector, Limit)` tuples with a global `Similarity` +- List of `DenseVector` with global `Similarity` and `Limit` + +--- + #### **Delete** Delete a point by its ID @@ -177,6 +264,19 @@ All fields are directly accessible: --- +### `SearchQuery` + +``` +SearchQuery( + vector: DenseVector, + similarity: Similarity, + limit: int, +) +``` +Structured representation of a search request + +--- + ### `Similarity` Enum representing distance functions: @@ -274,4 +374,4 @@ python -m grpc_tools.protoc \ After running this: - `vector_db_pb2_grpc.py` and `vector_db_pb2.py` will be updated -- No other client code should need changes +- No other client code should need changes \ No newline at end of file diff --git a/client/python/examples/all.py b/client/python/examples/all.py new file mode 100644 index 0000000..05deb55 --- /dev/null +++ b/client/python/examples/all.py @@ -0,0 +1,28 @@ +# This file is like a master test. Runs all the examples +# Not exactly the purpose of the examples dir, +# but helps in checking if any code updates haven't broken the API + +from pathlib import Path +import subprocess +import pytest + +# Didn't know I could do this with pytest, so cool +# Just run: pytest ./all.py -v + +EXAMPLES_DIR = Path(__file__).parent +example_files = sorted(EXAMPLES_DIR.glob("*_usage.py")) + + +@pytest.mark.parametrize( + "script_path", + example_files, + ids=lambda p: p.stem, +) +def test(script_path): + """Run all example scripts to check if they crash or not""" + result = subprocess.run( + ["python3", str(script_path)], capture_output=True, text=True + ) + assert result.returncode == 0, ( + f"Script {script_path} failed with stderr:\n{result.stderr}" + ) diff --git a/client/python/examples/async_batch_usage.py b/client/python/examples/async_batch_usage.py new file mode 100644 index 0000000..72d2493 --- /dev/null +++ b/client/python/examples/async_batch_usage.py @@ -0,0 +1,85 @@ +import asyncio + +from vortexdb import AsyncVortexDB +from vortexdb import Payload, Similarity, SearchQuery, to_dense_vectors + + +async def main(): + async with AsyncVortexDB( + grpc_url="localhost:50051", + api_key="my-secret-password", + ) as db: + raw_vectors = [ + [0.1, 0.2, 0.3], + [0.4, 0.5, 0.6], + [0.7, 0.8, 0.9], + ] + vectors = to_dense_vectors(raw_vectors) + + p1 = Payload.text("hello world") + p2 = Payload.image("/img/a.png") + p3 = Payload.text("foo bar") + + items = [ + (vectors[0], p1), + (vectors[1], p2), + (vectors[2], p3), + ] + + # Batch Insert + point_ids = await db.batch_insert(items=items) + print("Inserted ids:\n", point_ids) + + q = SearchQuery( + vector=vectors[0], + similarity=Similarity.COSINE, + limit=3, + ) + res = await db.search(query=q) + print("\nSingle SearchQuery:\n", res) + + # List of SearchQuery + queries = [ + SearchQuery(vectors[0], Similarity.HAMMING, 3), + SearchQuery(vectors[1], Similarity.EUCLIDEAN, 2), + q, + ] + res = await db.batch_search(queries=queries) + print("\nBatch SearchQuery:\n", res) + + # List of vectors with global Similarity and Limit + res = await db.batch_search( + queries=vectors, + similarity=Similarity.COSINE, + limit=3, + ) + print("\nList of DenseVectors:\n", res) + + # List of tuple (DenseVector, Similarity) with global Limit + queries = [ + (vectors[0], Similarity.COSINE), + (vectors[1], Similarity.MANHATTAN), + ] + res = await db.batch_search( + queries=queries, + limit=3, + ) + print("\nList of (DenseVector, Similarity):\n", res) + + # List of tuple (DenseVector, Limit) with global Similarity + queries = [ + (vectors[0], 2), + (vectors[1], 4), + ] + res = await db.batch_search( + queries=queries, + similarity=Similarity.COSINE, + ) + print("\nList of (DenseVector, Limit):\n", res) + + for pid in point_ids: + await db.delete(point_id=pid) + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/client/python/examples/async_usage.py b/client/python/examples/async_usage.py new file mode 100644 index 0000000..882a926 --- /dev/null +++ b/client/python/examples/async_usage.py @@ -0,0 +1,31 @@ +import asyncio + +from vortexdb import AsyncVortexDB, DenseVector, Payload, Similarity + + +async def main(): + async with AsyncVortexDB( + grpc_url="localhost:50051", + api_key="my-secret-password", + ) as db: + point_id = await db.insert( + vector=DenseVector([0.1, 0.2, 0.3]), + payload=Payload.text("hello async vortex"), + ) + + point = await db.get(point_id=point_id) + if point is not None: + print(point.pretty()) + + results = await db.search( + vector=DenseVector([0.1, 0.2, 0.3]), + similarity=Similarity.COSINE, + limit=5, + ) + print(results) + + await db.delete(point_id=point_id) + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/client/python/examples/basic_usage.py b/client/python/examples/basic_usage.py index 4288c02..f58d8c5 100644 --- a/client/python/examples/basic_usage.py +++ b/client/python/examples/basic_usage.py @@ -1,5 +1,6 @@ from vortexdb import VortexDB -from vortexdb import DenseVector, Payload, Similarity # from vortexdb.models +from vortexdb import DenseVector, Payload, Similarity # from vortexdb.models + def main(): # Initialize client @@ -32,5 +33,6 @@ def main(): # Close connection db.close() + if __name__ == "__main__": main() diff --git a/client/python/examples/batch_insert_usage.py b/client/python/examples/batch_insert_usage.py new file mode 100644 index 0000000..e021b1a --- /dev/null +++ b/client/python/examples/batch_insert_usage.py @@ -0,0 +1,39 @@ +from vortexdb import VortexDB +from vortexdb import Payload, to_dense_vectors + + +def main(): + db = VortexDB( + grpc_url="localhost:50051", + api_key="my-secret-password", + ) + + raw_vectors = [ + [0.1, 0.2, 0.3], + [0.4, 0.5, 0.6], + [0.7, 0.8, 0.9], + ] + vectors = to_dense_vectors(raw_vectors) + + p1 = Payload.text("hello world") + p2 = Payload.image("/img/a.png") + p3 = Payload.text("foo bar") + + items = [ + (vectors[0], p1), + (vectors[1], p2), + (vectors[2], p3), + ] + + # Batch Insert + point_ids = db.batch_insert(items=items) + print("Inserted ids:\n", point_ids) + + for pid in point_ids: + db.delete(point_id=pid) + + db.close() + + +if __name__ == "__main__": + main() diff --git a/client/python/examples/context_manager_usage.py b/client/python/examples/context_manager_usage.py index 0f86b99..d2f1de2 100644 --- a/client/python/examples/context_manager_usage.py +++ b/client/python/examples/context_manager_usage.py @@ -1,11 +1,11 @@ from vortexdb import VortexDB, DenseVector, Payload, Similarity + def main(): with VortexDB( grpc_url="localhost:50051", api_key="my-secret-password", ) as db: - # Insert a vector point_id = db.insert( vector=DenseVector([0.1, 0.2, 0.3]), @@ -30,5 +30,6 @@ def main(): # At this point, the gRPC channel is closed automatically print("Connection closed") + if __name__ == "__main__": main() diff --git a/client/python/examples/search_query_usage.py b/client/python/examples/search_query_usage.py new file mode 100644 index 0000000..3044056 --- /dev/null +++ b/client/python/examples/search_query_usage.py @@ -0,0 +1,69 @@ +from vortexdb import VortexDB +from vortexdb import Similarity, SearchQuery, to_dense_vectors + + +def main(): + db = VortexDB( + grpc_url="localhost:50051", + api_key="my-secret-password", + ) + + raw_vectors = [ + [0.1, 0.2, 0.3], + [0.4, 0.5, 0.6], + [0.7, 0.8, 0.9], + ] + vectors = to_dense_vectors(raw_vectors) + + q = SearchQuery( + vector=vectors[0], + similarity=Similarity.COSINE, + limit=3, + ) + res = db.search(query=q) + print("Single SearchQuery:\n", res) + + # List of SearchQuery + queries = [ + SearchQuery(vectors[0], Similarity.HAMMING, 3), + SearchQuery(vectors[1], Similarity.EUCLIDEAN, 2), + q, + ] + res = db.batch_search(queries=queries) + print("\nBatch SearchQuery:\n", res) + + # List of vectors with global Similarity and Limit + res = db.batch_search( + queries=vectors, + similarity=Similarity.COSINE, + limit=3, + ) + print("\nList of DenseVectors:\n", res) + + # List of tuple (DenseVector, Similarity) with global Limit + queries = [ + (vectors[0], Similarity.COSINE), + (vectors[1], Similarity.MANHATTAN), + ] + res = db.batch_search( + queries=queries, + limit=3, + ) + print("\nList of (DenseVector, Similarity):\n", res) + + # List of tuple (DenseVector, Limit) with global Similarity + queries = [ + (vectors[0], 2), + (vectors[1], 4), + ] + res = db.batch_search( + queries=queries, + similarity=Similarity.COSINE, + ) + print("\nList of (DenseVector, Limit):\n", res) + + db.close() + + +if __name__ == "__main__": + main() diff --git a/client/python/pyproject.toml b/client/python/pyproject.toml index 8528f91..7f4f031 100644 --- a/client/python/pyproject.toml +++ b/client/python/pyproject.toml @@ -24,13 +24,13 @@ classifiers = [ ] dependencies = [ - "grpcio>=1.60", - "protobuf>=4.25", + "grpcio>=1.81.1", + "protobuf>=6.33.5,<7.0.0", ] [project.optional-dependencies] dev = [ - "grpcio-tools>=1.60", + "grpcio-tools>=1.81.1,<2.0.0", "pytest>=7.0", ] @@ -42,3 +42,6 @@ Repository = "https://github.com/sdslabs/VortexDB" [tool.pytest.ini_options] testpaths = ["tests"] + +[tool.ruff] +exclude = ["vortexdb/grpc"] diff --git a/client/python/tests/test_api_parity.py b/client/python/tests/test_api_parity.py new file mode 100644 index 0000000..ee8e611 --- /dev/null +++ b/client/python/tests/test_api_parity.py @@ -0,0 +1,40 @@ +# VortexDB and AsyncVortexDB should expose the same public methods +# with the same signatures +import inspect + +from vortexdb.client import VortexDB +from vortexdb.async_client import AsyncVortexDB + + +def public_methods(cls): + return { + name: value + for name, value in vars(cls).items() + if callable(value) and not name.startswith("_") + } + + +def test_sync_async_client_api_parity(): + sync_methods = public_methods(VortexDB) + async_methods = public_methods(AsyncVortexDB) + + sync_names = set(sync_methods) + async_names = set(async_methods) + + assert sync_names == async_names, ( + "Sync and async clients expose different methods. " + f"Only sync: {sorted(sync_names - async_names)}. " + f"Only async: {sorted(async_names - sync_names)}." + ) + + mismatches = [ + f" {name}: sync{inspect.signature(sync_methods[name])}" + f" != async{inspect.signature(async_methods[name])}" + for name in sync_names + if inspect.signature(sync_methods[name]) + != inspect.signature(async_methods[name]) + ] + + assert not mismatches, ( + "Sync and async methods have mismatched signatures:\n" + "\n".join(mismatches) + ) diff --git a/client/python/tests/test_async_client.py b/client/python/tests/test_async_client.py new file mode 100644 index 0000000..d2af9fd --- /dev/null +++ b/client/python/tests/test_async_client.py @@ -0,0 +1,385 @@ +import asyncio +from unittest.mock import AsyncMock, Mock + +import pytest + +from vortexdb.async_client import AsyncVortexDB +from vortexdb.async_connection import AsyncGRPCConnection +from vortexdb.models import ( + ContentType, + DenseVector, + Payload, + Point, + Similarity, + SearchQuery, +) + + +@pytest.fixture +def mock_connection(monkeypatch): + """ + Replace AsyncGRPCConnection with a mock instance. + """ + conn = Mock(spec=AsyncGRPCConnection) + conn.stub = Mock() + conn.call = AsyncMock() + conn.close = AsyncMock() + monkeypatch.setattr("vortexdb.async_client.AsyncGRPCConnection", lambda _: conn) + return conn + + +@pytest.fixture +def client(mock_connection): + return AsyncVortexDB( + grpc_url="localhost:50051", + api_key="secret", + ) + + +def test_async_insert_success(client, mock_connection): + async def run(): + response = Mock() + response.id = Mock() + response.id.value = "point-123" + + mock_connection.call.return_value = response + + point_id = await client.insert( + vector=DenseVector([1, 2, 3]), + payload=Payload.text("hello"), + ) + + assert point_id == "point-123" + + asyncio.run(run()) + + +def test_async_insert_rejects_invalid_vector(client): + async def run(): + with pytest.raises(TypeError): + await client.insert( + vector=[1, 2, 3], + payload=Payload.text("hello"), + ) + + asyncio.run(run()) + + +# Batch Insert + + +def test_async_batch_insert_success(client, mock_connection): + async def run(): + response = Mock() + response.ids = [ + Mock(id=Mock(value="p1")), + Mock(id=Mock(value="p2")), + ] + mock_connection.call.return_value = response + + items = [ + (DenseVector([1, 2, 3]), Payload.text("a")), + (DenseVector([4, 5, 6]), Payload.text("b")), + ] + result = await client.batch_insert(items=items) + assert result == ["p1", "p2"] + + asyncio.run(run()) + + +def test_async_batch_insert_invalid_items_type(client): + async def run(): + with pytest.raises(TypeError): + await client.batch_insert(items="not-a-list") + + asyncio.run(run()) + + +def test_async_batch_insert_invalid_tuple_structure(client): + async def run(): + items = [ + (DenseVector([1, 2, 3]),), # only one element + ] + with pytest.raises(TypeError): + await client.batch_insert(items=items) + + asyncio.run(run()) + + +def test_async_batch_insert_invalid_vector(client): + async def run(): + items = [ + ([1, 2, 3], Payload.text("a")), # not DenseVector + ] + with pytest.raises(TypeError): + await client.batch_insert(items=items) + + asyncio.run(run()) + + +# Get + + +def test_async_get_point_success(client, mock_connection): + async def run(): + proto_point = Mock() + proto_point.id.id.value = "point-123" + proto_point.vector.values = [1, 2, 3] + proto_point.payload.content_type = ContentType.TEXT.to_proto() + proto_point.payload.content = "hello" + + mock_connection.call.return_value = proto_point + + point = await client.get(point_id="point-123") + + assert isinstance(point, Point) + assert point.id == "point-123" + assert point.payload.content == "hello" + + asyncio.run(run()) + + +def test_async_get_point_not_found(client, mock_connection): + async def run(): + mock_connection.call.return_value = None + + result = await client.get(point_id="missing") + + assert result is None + + asyncio.run(run()) + + +# Delete + + +def test_async_delete_success(client, mock_connection): + async def run(): + mock_connection.call.return_value = None + + await client.delete(point_id="point-123") + + mock_connection.call.assert_awaited_once() + + asyncio.run(run()) + + +# Search + + +def test_async_search_success(client, mock_connection): + async def run(): + mock_connection.call.return_value = Mock( + result_point_ids=[ + Mock(id=Mock(value="p1")), + Mock(id=Mock(value="p2")), + ] + ) + + results = await client.search( + vector=DenseVector([1, 2, 3]), + similarity=Similarity.COSINE, + limit=2, + ) + + assert results == ["p1", "p2"] + + asyncio.run(run()) + + +def test_async_search_with_query_object(client, mock_connection): + async def run(): + mock_connection.call.return_value = Mock( + result_point_ids=[ + Mock(id=Mock(value="p1")), + ] + ) + + q = SearchQuery(DenseVector([1, 2, 3]), Similarity.COSINE, 2) + results = await client.search(query=q) + + assert results == ["p1"] + + asyncio.run(run()) + + +def test_async_search_accepts_ef(client, mock_connection): + async def run(): + mock_connection.call.return_value = Mock(result_point_ids=[]) + + await client.search( + vector=DenseVector([1, 2, 3]), + similarity=Similarity.COSINE, + limit=2, + ef=128, + ) + + request = mock_connection.call.call_args.args[1] + assert request.ef == 128 + + asyncio.run(run()) + + +def test_async_search_invalid_vector(client): + async def run(): + with pytest.raises(TypeError): + await client.search( + vector=[1, 2, 3], + similarity=Similarity.COSINE, + limit=2, + ) + + asyncio.run(run()) + + +# Batch Search + + +def test_async_batch_search_full_tuple(client, mock_connection): + async def run(): + mock_connection.call.return_value = Mock( + results=[ + Mock(result_point_ids=[Mock(id=Mock(value="p1"))]), + Mock(result_point_ids=[Mock(id=Mock(value="p2"))]), + ] + ) + queries = [ + (DenseVector([1, 2, 3]), Similarity.COSINE, 2), + (DenseVector([4, 5, 6]), Similarity.EUCLIDEAN, 1), + ] + result = await client.batch_search(queries=queries) + assert result == [["p1"], ["p2"]] + + asyncio.run(run()) + + +def test_async_batch_search_searchquery_objects(client, mock_connection): + async def run(): + mock_connection.call.return_value = Mock( + results=[ + Mock(result_point_ids=[Mock(id=Mock(value="p1"))]), + ] + ) + queries = [ + SearchQuery(DenseVector([1, 2, 3]), Similarity.COSINE, 2), + ] + result = await client.batch_search(queries=queries) + assert result == [["p1"]] + + asyncio.run(run()) + + +def test_async_batch_search_vectors_with_global_params(client, mock_connection): + async def run(): + mock_connection.call.return_value = Mock( + results=[ + Mock(result_point_ids=[Mock(id=Mock(value="p1"))]), + ] + ) + queries = [DenseVector([1, 2, 3])] + result = await client.batch_search( + queries=queries, + similarity=Similarity.MANHATTAN, + limit=2, + ) + assert result == [["p1"]] + + asyncio.run(run()) + + +def test_async_batch_search_vector_similarity_with_global_limit( + client, mock_connection +): + async def run(): + mock_connection.call.return_value = Mock( + results=[ + Mock(result_point_ids=[Mock(id=Mock(value="p1"))]), + ] + ) + queries = [ + (DenseVector([1, 2, 3]), Similarity.COSINE), + ] + result = await client.batch_search( + queries=queries, + limit=2, + ) + assert result == [["p1"]] + + asyncio.run(run()) + + +def test_async_batch_search_accepts_ef(client, mock_connection): + async def run(): + mock_connection.call.return_value = Mock(results=[]) + + await client.batch_search( + queries=[ + (DenseVector([1, 2, 3]), Similarity.COSINE, 2), + (DenseVector([4, 5, 6]), Similarity.COSINE, 1), + ], + ef=256, + ) + + request = mock_connection.call.call_args.args[1] + assert [query.ef for query in request.queries] == [256, 256] + + asyncio.run(run()) + + +def test_async_batch_search_missing_globals_for_vector(client): + async def run(): + queries = [DenseVector([1, 2, 3])] + with pytest.raises(ValueError): + await client.batch_search(queries=queries) + + asyncio.run(run()) + + +def test_async_batch_search_missing_limit(client): + async def run(): + queries = [ + (DenseVector([1, 2, 3]), Similarity.COSINE), + ] + with pytest.raises(ValueError): + await client.batch_search(queries=queries) + + asyncio.run(run()) + + +def test_async_batch_search_invalid_format(client): + async def run(): + queries = ["invalid"] + with pytest.raises(TypeError): + await client.batch_search(queries=queries) + + asyncio.run(run()) + + +# Close / Context Manager + + +def test_async_close_closes_connection(client, mock_connection): + async def run(): + await client.close() + mock_connection.close.assert_awaited_once() + + asyncio.run(run()) + + +def test_async_context_manager_closes_connection(monkeypatch): + async def run(): + conn = Mock(spec=AsyncGRPCConnection) + conn.stub = Mock() + conn.call = AsyncMock() + conn.close = AsyncMock() + monkeypatch.setattr("vortexdb.async_client.AsyncGRPCConnection", lambda _: conn) + + async with AsyncVortexDB( + grpc_url="localhost:50051", + api_key="secret", + ) as db: + assert db is not None + + conn.close.assert_awaited_once() + + asyncio.run(run()) diff --git a/client/python/tests/test_async_connection.py b/client/python/tests/test_async_connection.py new file mode 100644 index 0000000..62f1043 --- /dev/null +++ b/client/python/tests/test_async_connection.py @@ -0,0 +1,106 @@ +import asyncio +from unittest.mock import AsyncMock, Mock, patch + +import grpc + +from vortexdb._grpc_common import map_grpc_error +from vortexdb.async_connection import AsyncGRPCConnection +from vortexdb.config import VortexDBConfig +from vortexdb.exceptions import ( + AuthenticationError, + InternalServerError, + InvalidArgumentError, + NotFoundError, + ServiceUnavailableError, + TimeoutError, +) + + +class FakeAioRpcError: + """ + Minimal AioRpcError-compatible object for mapping tests. + """ + + def __init__(self, status_code: grpc.StatusCode, details: str): + self._status_code = status_code + self._details = details + + def code(self): + return self._status_code + + def details(self): + return self._details + + +def make_config() -> VortexDBConfig: + return VortexDBConfig( + grpc_url="localhost:50051", + api_key="secret", + timeout=3.0, + ) + + +def test_async_channel_created_with_correct_url(): + with patch("grpc.aio.insecure_channel") as mock_channel: + mock_channel.return_value = Mock() + AsyncGRPCConnection(make_config()) + mock_channel.assert_called_once_with("localhost:50051") + + +def test_async_metadata_is_attached(): + with patch("grpc.aio.insecure_channel") as mock_channel: + mock_channel.return_value = Mock() + connection = AsyncGRPCConnection(make_config()) + + assert ("authorization", "Bearer secret") in connection._metadata + + +def test_successful_async_rpc_call(): + async def run(): + with patch("grpc.aio.insecure_channel") as mock_channel: + mock_channel.return_value = Mock() + connection = AsyncGRPCConnection(make_config()) + + fake_rpc = AsyncMock(return_value="ok") + + result = await connection.call(fake_rpc, request="req") + + fake_rpc.assert_awaited_once_with( + "req", + timeout=3.0, + metadata=connection._metadata, + ) + assert result == "ok" + + asyncio.run(run()) + + +def test_async_grpc_error_mapping(): + cases = [ + (grpc.StatusCode.UNAUTHENTICATED, AuthenticationError), + (grpc.StatusCode.NOT_FOUND, NotFoundError), + (grpc.StatusCode.INVALID_ARGUMENT, InvalidArgumentError), + (grpc.StatusCode.DEADLINE_EXCEEDED, TimeoutError), + (grpc.StatusCode.UNAVAILABLE, ServiceUnavailableError), + (grpc.StatusCode.UNKNOWN, InternalServerError), + ] + + for status_code, expected_exception in cases: + error = FakeAioRpcError(status_code, "boom") + mapped = map_grpc_error(error) + assert isinstance(mapped, expected_exception) + + +def test_async_close_closes_channel(): + async def run(): + with patch("grpc.aio.insecure_channel") as mock_channel: + channel = Mock() + channel.close = AsyncMock() + mock_channel.return_value = channel + connection = AsyncGRPCConnection(make_config()) + + await connection.close() + + channel.close.assert_awaited_once() + + asyncio.run(run()) diff --git a/client/python/tests/test_client.py b/client/python/tests/test_client.py index a752320..b639da5 100644 --- a/client/python/tests/test_client.py +++ b/client/python/tests/test_client.py @@ -4,12 +4,12 @@ from vortexdb.client import VortexDB from vortexdb.connection import GRPCConnection from vortexdb.models import DenseVector, Payload, Similarity, ContentType, Point -from vortexdb.exceptions import InvalidArgumentError - +from vortexdb.models import SearchQuery # Fixtures for a mock connection and client layer + @pytest.fixture def mock_connection(monkeypatch): """ @@ -30,6 +30,7 @@ def client(mock_connection): # Insert + def test_insert_success(client, mock_connection): response = Mock() response.id = Mock() @@ -45,7 +46,6 @@ def test_insert_success(client, mock_connection): assert point_id == "point-123" - def test_insert_rejects_invalid_vector(client): with pytest.raises(TypeError): client.insert( @@ -54,8 +54,48 @@ def test_insert_rejects_invalid_vector(client): ) +# Batch Insert + + +def test_batch_insert_success(client, mock_connection): + response = Mock() + response.ids = [ + Mock(id=Mock(value="p1")), + Mock(id=Mock(value="p2")), + ] + mock_connection.call.return_value = response + items = [ + (DenseVector([1, 2, 3]), Payload.text("a")), + (DenseVector([4, 5, 6]), Payload.text("b")), + ] + result = client.batch_insert(items=items) + assert result == ["p1", "p2"] + + +def test_batch_insert_invalid_items_type(client): + with pytest.raises(TypeError): + client.batch_insert(items="not-a-list") + + +def test_batch_insert_invalid_tuple_structure(client): + items = [ + (DenseVector([1, 2, 3]),), # only one element + ] + with pytest.raises(TypeError): + client.batch_insert(items=items) + + +def test_batch_insert_invalid_vector(client): + items = [ + ([1, 2, 3], Payload.text("a")), # not DenseVector + ] + with pytest.raises(TypeError): + client.batch_insert(items=items) + + # Get + def test_get_point_success(client, mock_connection): proto_point = Mock() proto_point.id.id.value = "point-123" @@ -82,6 +122,7 @@ def test_get_point_not_found(client, mock_connection): # Delete + def test_delete_success(client, mock_connection): mock_connection.call.return_value = None @@ -92,6 +133,7 @@ def test_delete_success(client, mock_connection): # Search + def test_search_success(client, mock_connection): mock_connection.call.return_value = Mock( result_point_ids=[ @@ -109,6 +151,20 @@ def test_search_success(client, mock_connection): assert results == ["p1", "p2"] +def test_search_accepts_ef(client, mock_connection): + mock_connection.call.return_value = Mock(result_point_ids=[]) + + client.search( + vector=DenseVector([1, 2, 3]), + similarity=Similarity.COSINE, + limit=2, + ef=128, + ) + + request = mock_connection.call.call_args.args[1] + assert request.ef == 128 + + def test_search_invalid_vector(client): with pytest.raises(TypeError): client.search( @@ -118,12 +174,111 @@ def test_search_invalid_vector(client): ) +def test_batch_search_accepts_ef(client, mock_connection): + mock_connection.call.return_value = Mock(results=[]) + + client.batch_search( + queries=[ + (DenseVector([1, 2, 3]), Similarity.COSINE, 2), + (DenseVector([4, 5, 6]), Similarity.COSINE, 1), + ], + ef=256, + ) + + request = mock_connection.call.call_args.args[1] + assert [query.ef for query in request.queries] == [256, 256] + + +# Batch Search + + +def test_batch_search_full_tuple(client, mock_connection): + mock_connection.call.return_value = Mock( + results=[ + Mock(result_point_ids=[Mock(id=Mock(value="p1"))]), + Mock(result_point_ids=[Mock(id=Mock(value="p2"))]), + ] + ) + queries = [ + (DenseVector([1, 2, 3]), Similarity.COSINE, 2), + (DenseVector([4, 5, 6]), Similarity.EUCLIDEAN, 1), + ] + result = client.batch_search(queries=queries) + assert result == [["p1"], ["p2"]] + + +def test_batch_search_vectors_with_global_params(client, mock_connection): + mock_connection.call.return_value = Mock( + results=[ + Mock(result_point_ids=[Mock(id=Mock(value="p1"))]), + ] + ) + queries = [DenseVector([1, 2, 3])] + result = client.batch_search( + queries=queries, + similarity=Similarity.MANHATTAN, + limit=2, + ) + assert result == [["p1"]] + + +def test_batch_search_vector_similarity_with_global_limit(client, mock_connection): + mock_connection.call.return_value = Mock( + results=[ + Mock(result_point_ids=[Mock(id=Mock(value="p1"))]), + ] + ) + queries = [ + (DenseVector([1, 2, 3]), Similarity.COSINE), + ] + result = client.batch_search( + queries=queries, + limit=2, + ) + assert result == [["p1"]] + + +def test_batch_search_searchquery_objects(client, mock_connection): + mock_connection.call.return_value = Mock( + results=[ + Mock(result_point_ids=[Mock(id=Mock(value="p1"))]), + ] + ) + queries = [ + SearchQuery(DenseVector([1, 2, 3]), Similarity.COSINE, 2), + ] + result = client.batch_search(queries=queries) + assert result == [["p1"]] + + +def test_batch_search_missing_globals_for_vector(client): + queries = [DenseVector([1, 2, 3])] + with pytest.raises(ValueError): + client.batch_search(queries=queries) + + +def test_batch_search_missing_limit(client): + queries = [ + (DenseVector([1, 2, 3]), Similarity.COSINE), + ] + with pytest.raises(ValueError): + client.batch_search(queries=queries) + + +def test_batch_search_invalid_format(client): + queries = ["invalid"] + with pytest.raises(TypeError): + client.batch_search(queries=queries) + + # Close + def test_close_closes_connection(client, mock_connection): client.close() mock_connection.close.assert_called_once() + def test_context_manager_closes_connection(monkeypatch): conn = Mock(spec=GRPCConnection) monkeypatch.setattr("vortexdb.client.GRPCConnection", lambda _: conn) diff --git a/client/python/tests/test_config.py b/client/python/tests/test_config.py index 1017250..f23ab2b 100644 --- a/client/python/tests/test_config.py +++ b/client/python/tests/test_config.py @@ -16,6 +16,7 @@ def clean_env(monkeypatch): # Checking from_env + def test_config_requires_api_key(clean_env): with pytest.raises(ConfigurationError): VortexDBConfig.from_env() @@ -32,8 +33,10 @@ def test_config_from_explicit_args(clean_env): assert cfg.api_key == "secret" assert cfg.timeout == 10.0 + # Env vars fallback + def test_config_from_env_vars(clean_env, monkeypatch): monkeypatch.setenv("VORTEXDB_GRPC_URL", "127.0.0.1:1234") monkeypatch.setenv("VORTEXDB_API_KEY", "env-secret") @@ -48,6 +51,7 @@ def test_config_from_env_vars(clean_env, monkeypatch): # Defaults + def test_config_default_grpc_url(clean_env, monkeypatch): monkeypatch.setenv("VORTEXDB_API_KEY", "secret") @@ -66,6 +70,7 @@ def test_config_default_timeout(clean_env, monkeypatch): # Invalid Timeout + def test_config_invalid_timeout(clean_env, monkeypatch): monkeypatch.setenv("VORTEXDB_API_KEY", "secret") monkeypatch.setenv("VORTEXDB_TIMEOUT", "not-a-number") diff --git a/client/python/tests/test_connection.py b/client/python/tests/test_connection.py index ee82152..6d5c108 100644 --- a/client/python/tests/test_connection.py +++ b/client/python/tests/test_connection.py @@ -2,6 +2,7 @@ import pytest from unittest.mock import Mock, patch +from vortexdb._grpc_common import map_grpc_error from vortexdb.connection import GRPCConnection from vortexdb.config import VortexDBConfig from vortexdb.exceptions import ( @@ -15,6 +16,7 @@ # Fake gRPC error, required for testing + class FakeRpcError(grpc.RpcError): """ RpcError implementation for unit testing. @@ -34,6 +36,7 @@ def details(self): # Pytest fixtures for config and channel + @pytest.fixture def config(): return VortexDBConfig( @@ -52,6 +55,7 @@ def connection(config): # Basic connection testing + def test_channel_created_with_correct_url(config): with patch("grpc.insecure_channel") as mock_channel: GRPCConnection(config) @@ -77,6 +81,7 @@ def test_successful_rpc_call(connection): # Error mapping test + @pytest.mark.parametrize( "status_code,expected_exception", [ @@ -89,22 +94,22 @@ def test_successful_rpc_call(connection): ) def test_grpc_error_mapping(status_code, expected_exception, connection): error = FakeRpcError(status_code, "boom") - fake_rpc = Mock(side_effect=error) - with pytest.raises(expected_exception): - connection.call(fake_rpc, request="req") + mapped = map_grpc_error(error) + + assert isinstance(mapped, expected_exception) def test_unknown_grpc_error_maps_to_internal_error(connection): error = FakeRpcError(grpc.StatusCode.UNKNOWN, "unknown") - fake_rpc = Mock(side_effect=error) + mapped = map_grpc_error(error) - with pytest.raises(InternalServerError): - connection.call(fake_rpc, request="req") + assert isinstance(mapped, InternalServerError) # Clean connection closure test + def test_close_closes_channel(config): with patch("grpc.insecure_channel") as mock_channel: mock_channel.return_value = Mock() diff --git a/client/python/tests/test_models.py b/client/python/tests/test_models.py index e78d686..796631f 100644 --- a/client/python/tests/test_models.py +++ b/client/python/tests/test_models.py @@ -12,6 +12,7 @@ # DenseVector Tests + def test_dense_vector_valid(): a = [1, 2.5, 3] v = DenseVector(a) @@ -47,6 +48,7 @@ def test_dense_vector_to_proto(): # Similarity Test + def test_similarity_to_proto(): assert Similarity.EUCLIDEAN.to_proto() == vector_db_pb2.Euclidean assert Similarity.MANHATTAN.to_proto() == vector_db_pb2.Manhattan @@ -56,6 +58,7 @@ def test_similarity_to_proto(): # ContentType Tests + def test_content_type_to_proto(): assert ContentType.TEXT.to_proto() == vector_db_pb2.Text assert ContentType.IMAGE.to_proto() == vector_db_pb2.Image @@ -73,6 +76,7 @@ def test_content_type_from_proto_invalid(): # Payload Tests + def test_payload_text_factory(): p = Payload.text("hello") assert p.content_type == ContentType.TEXT @@ -91,24 +95,20 @@ def test_payload_to_proto(): assert proto.content == "hello" assert proto.content_type == vector_db_pb2.Text + def test_payload_rejects_invalid_content_type(): with pytest.raises(TypeError): Payload("text", "hello") - # Point Test + def test_point_from_proto(): proto = vector_db_pb2.Point( - id=vector_db_pb2.PointID( - id=vector_db_pb2.UUID(value="point-123") - ), + id=vector_db_pb2.PointID(id=vector_db_pb2.UUID(value="point-123")), vector=vector_db_pb2.DenseVector(values=[1, 2, 3]), - payload=vector_db_pb2.Payload( - content_type=vector_db_pb2.Text, - content="hello" - ) + payload=vector_db_pb2.Payload(content_type=vector_db_pb2.Text, content="hello"), ) point = Point.from_proto(proto) @@ -118,12 +118,11 @@ def test_point_from_proto(): assert point.payload.content_type == ContentType.TEXT assert point.payload.content == "hello" + def test_point_from_proto_without_payload(): proto = vector_db_pb2.Point( - id=vector_db_pb2.PointID( - id=vector_db_pb2.UUID(value="p1") - ), - vector=vector_db_pb2.DenseVector(values=[1,2,3]), + id=vector_db_pb2.PointID(id=vector_db_pb2.UUID(value="p1")), + vector=vector_db_pb2.DenseVector(values=[1, 2, 3]), payload=None, ) diff --git a/client/python/vortexdb/__init__.py b/client/python/vortexdb/__init__.py index 62c100f..9cc9dd4 100644 --- a/client/python/vortexdb/__init__.py +++ b/client/python/vortexdb/__init__.py @@ -1,11 +1,14 @@ # vortexdb/__init__.py from vortexdb.client import VortexDB +from vortexdb.async_client import AsyncVortexDB from vortexdb.models import ( DenseVector, Payload, Point, Similarity, + SearchQuery, + to_dense_vectors, ) from vortexdb.exceptions import ( VortexDBError, @@ -19,10 +22,13 @@ __all__ = [ "VortexDB", + "AsyncVortexDB", "DenseVector", "Payload", "Point", "Similarity", + "SearchQuery", + "to_dense_vectors", "VortexDBError", "AuthenticationError", "NotFoundError", diff --git a/client/python/vortexdb/_grpc_common.py b/client/python/vortexdb/_grpc_common.py new file mode 100644 index 0000000..2257f3d --- /dev/null +++ b/client/python/vortexdb/_grpc_common.py @@ -0,0 +1,38 @@ +from typing import Any + +import grpc + +from vortexdb.exceptions import ( + AuthenticationError, + InternalServerError, + InvalidArgumentError, + NotFoundError, + ServiceUnavailableError, + TimeoutError, + VortexDBError, +) + + +def build_auth_metadata(api_key: str) -> tuple[tuple[str, str], ...]: + return (("authorization", f"Bearer {api_key}"),) + + +def map_grpc_error(error: Any) -> VortexDBError: + code = error.code() + + if code == grpc.StatusCode.UNAUTHENTICATED: + return AuthenticationError(error.details()) + + if code == grpc.StatusCode.NOT_FOUND: + return NotFoundError(error.details()) + + if code == grpc.StatusCode.INVALID_ARGUMENT: + return InvalidArgumentError(error.details()) + + if code == grpc.StatusCode.DEADLINE_EXCEEDED: + return TimeoutError(error.details()) + + if code == grpc.StatusCode.UNAVAILABLE: + return ServiceUnavailableError(error.details()) + + return InternalServerError(error.details()) diff --git a/client/python/vortexdb/async_client.py b/client/python/vortexdb/async_client.py new file mode 100644 index 0000000..3b7c2f5 --- /dev/null +++ b/client/python/vortexdb/async_client.py @@ -0,0 +1,211 @@ +from typing import List + +from vortexdb import protoutils as proto +from vortexdb.async_connection import AsyncGRPCConnection +from vortexdb.config import VortexDBConfig +from vortexdb.models import DenseVector, Payload, Point, Similarity, SearchQuery + + +class AsyncVortexDB: + """High-level async Python client for VortexDB.""" + + def __init__( + self, + *, + grpc_url: str | None = None, + api_key: str | None = None, + timeout: float | None = None, + ): + # Config order followed - args -> env vars -> defaults + self._config = VortexDBConfig.from_env( + grpc_url=grpc_url, + api_key=api_key, + timeout=timeout, + ) + + self._conn = AsyncGRPCConnection(self._config) + + async def insert(self, *, vector: DenseVector, payload: Payload) -> str: + """ + Insert a vector with payload. + Returns: point_id (str) + """ + if not isinstance(vector, DenseVector): + raise TypeError( + "vector must be a DenseVector. Use: DenseVector([1.0, 2.0, 3.0])" + ) + + request = proto.build_insert_request( + vector=vector, + payload=payload, + ) + + response = await self._conn.call( + self._conn.stub.InsertVector, + request, + ) + + return response.id.value + + async def batch_insert( + self, *, items: list[tuple[DenseVector, Payload]] + ) -> list[str]: + """ + Insert multiple vectors. + Returns: list of point_id (str) + """ + request = proto.build_batch_insert_request(items=items) + + response = await self._conn.call( + self._conn.stub.InsertVectorsBatch, + request, + ) + + return [pid.id.value for pid in response.ids] + + async def get(self, *, point_id: str) -> Point | None: + """ + Retrieve a point by ID. + """ + request = proto.build_point_id_request(point_id) + + response = await self._conn.call( + self._conn.stub.GetPoint, + request, + ) + + if response is None: + return None + + return Point.from_proto(response) + + async def delete(self, *, point_id: str) -> None: + """ + Delete a point by ID. + """ + request = proto.build_point_id_request(point_id) + + await self._conn.call( + self._conn.stub.DeletePoint, + request, + ) + + async def search( + self, + *, + vector: DenseVector | None = None, + similarity: Similarity | None = None, + limit: int | None = None, + query: SearchQuery | None = None, + ef: int | None = None, + ) -> List[str]: + """ + Search for nearest neighbors. + Returns: List of point IDs + """ + if query is not None: + if not isinstance(query, SearchQuery): + raise TypeError("query must be a SearchQuery") + vector = query.vector + similarity = query.similarity + limit = query.limit + else: + if not isinstance(vector, DenseVector): + raise TypeError( + "vector must be a DenseVector. Use: DenseVector([1.0, 2.0, 3.0])" + ) + if not isinstance(similarity, Similarity): + raise TypeError("similarity must be a Similarity enum") + if not isinstance(limit, int): + raise TypeError("limit must be an int") + + request = proto.build_search_request( + vector=vector, + similarity=similarity, + limit=limit, + ef=ef, + ) + + response = await self._conn.call( + self._conn.stub.SearchPoints, + request, + ) + + return [pid.id.value for pid in response.result_point_ids] + + async def batch_search( + self, + *, + queries, + similarity: Similarity | None = None, + limit: int | None = None, + ef: int | None = None, + ) -> List[List[str]]: + """ + Flexible batch search. + + Accepts: + - List[SearchQuery] + - List[(DenseVector, Similarity, int)] + - List[(DenseVector, Similarity)] + global limit + - List[(DenseVector, int)] + global similarity + - List[DenseVector] + global similarity + limit + """ + normalized = [] + + for i, q in enumerate(queries): + if ( + hasattr(q, "vector") + and hasattr(q, "similarity") + and hasattr(q, "limit") + ): + normalized.append((q.vector, q.similarity, q.limit)) + continue + + if isinstance(q, DenseVector): + if similarity is None or limit is None: + raise ValueError( + f"queries[{i}] requires global similarity and limit" + ) + normalized.append((q, similarity, limit)) + continue + + if isinstance(q, (list, tuple)): + if len(q) == 3: + normalized.append(q) + continue + if len(q) == 2: + a, b = q + + if isinstance(a, DenseVector) and isinstance(b, Similarity): + if limit is None: + raise ValueError(f"queries[{i}] missing global limit") + normalized.append((a, b, limit)) + continue + + if isinstance(a, DenseVector) and isinstance(b, int): + if similarity is None: + raise ValueError(f"queries[{i}] missing global similarity") + normalized.append((a, similarity, b)) + continue + + raise TypeError(f"Invalid query format at index {i}") + + request = proto.build_batch_search_request(queries=normalized, ef=ef) + response = await self._conn.call(self._conn.stub.SearchPointsBatch, request) + return [ + [pid.id.value for pid in result.result_point_ids] + for result in response.results + ] + + async def close(self) -> None: + """ + Close the async gRPC connection. + """ + await self._conn.close() + + async def __aenter__(self) -> "AsyncVortexDB": + return self + + async def __aexit__(self, exc_type, exc, tb) -> None: + await self.close() diff --git a/client/python/vortexdb/async_connection.py b/client/python/vortexdb/async_connection.py new file mode 100644 index 0000000..69dc9b6 --- /dev/null +++ b/client/python/vortexdb/async_connection.py @@ -0,0 +1,41 @@ +from typing import Any, Callable + +import grpc + +from vortexdb._grpc_common import build_auth_metadata, map_grpc_error +from vortexdb.config import VortexDBConfig +from vortexdb.grpc.vector_db_pb2_grpc import VectorDBStub + + +class AsyncGRPCConnection: + """Async gRPC connection wrapper for VortexDB.""" + + def __init__(self, config: VortexDBConfig): + self._config = config + self._channel = grpc.aio.insecure_channel(config.grpc_url) + self._stub = VectorDBStub(self._channel) + self._metadata = build_auth_metadata(config.api_key) + + @property + def stub(self) -> VectorDBStub: + return self._stub + + async def call( + self, + rpc: Callable[..., Any], + request: Any, + ) -> Any: + """Execute an async gRPC call with standard error handling.""" + try: + return await rpc( + request, + timeout=self._config.timeout, + metadata=self._metadata, + ) + + except grpc.aio.AioRpcError as e: + raise map_grpc_error(e) from e + + async def close(self) -> None: + """Close the underlying async gRPC channel.""" + await self._channel.close() diff --git a/client/python/vortexdb/client.py b/client/python/vortexdb/client.py index 38e0553..9984279 100644 --- a/client/python/vortexdb/client.py +++ b/client/python/vortexdb/client.py @@ -7,13 +7,14 @@ Payload, Point, Similarity, + SearchQuery, ) from vortexdb import protoutils as proto class VortexDB: - """ High-level Python client for VortexDB """ + """High-level Python client for VortexDB""" def __init__( self, @@ -22,7 +23,7 @@ def __init__( api_key: str | None = None, timeout: float | None = None, ): - # Config order followed - args -> env vars -> defaults + # Config order followed - args -> env vars -> defaults self._config = VortexDBConfig.from_env( grpc_url=grpc_url, api_key=api_key, @@ -31,18 +32,14 @@ def __init__( self._conn = GRPCConnection(self._config) -# The basic operations + # The basic operations def insert(self, *, vector: DenseVector, payload: Payload) -> str: """ Insert a vector with payload. Returns: point_id (str) """ - if not isinstance(vector, DenseVector): - raise TypeError( - "vector must be a DenseVector. " - "Use: DenseVector([1.0, 2.0, 3.0])" - ) + self._validate_dense_vector(vector) request = proto.build_insert_request( vector=vector, @@ -56,6 +53,19 @@ def insert(self, *, vector: DenseVector, payload: Payload) -> str: return response.id.value + def batch_insert(self, *, items: list[tuple[DenseVector, Payload]]) -> list[str]: + """ + Insert multiple vectors. + Returns: list of point_id (str) + """ + request = proto.build_batch_insert_request(items=items) + + response = self._conn.call( + self._conn.stub.InsertVectorsBatch, + request, + ) + return [pid.id.value for pid in response.ids] + def get(self, *, point_id: str) -> Point | None: """ Retrieve a point by ID. @@ -72,7 +82,6 @@ def get(self, *, point_id: str) -> Point | None: return Point.from_proto(response) - def delete(self, *, point_id: str) -> None: """ Delete a point by ID. @@ -87,33 +96,113 @@ def delete(self, *, point_id: str) -> None: def search( self, *, - vector: DenseVector, - similarity: Similarity, - limit: int, + vector: DenseVector | None = None, + similarity: Similarity | None = None, + limit: int | None = None, + query: SearchQuery | None = None, + ef: int | None = None, ) -> List[str]: """ Search for nearest neighbors. Returns: List of point IDs """ - if not isinstance(vector, DenseVector): - raise TypeError( - "vector must be a DenseVector. " - "Use: DenseVector([1.0, 2.0, 3.0])" - ) + if query is not None: + if not isinstance(query, SearchQuery): + raise TypeError("query must be a SearchQuery") + vector = query.vector + similarity = query.similarity + limit = query.limit + else: + self._validate_dense_vector(vector) + if not isinstance(similarity, Similarity): + raise TypeError("similarity must be a Similarity enum") + if not isinstance(limit, int): + raise TypeError("limit must be an int") request = proto.build_search_request( vector=vector, similarity=similarity, limit=limit, + ef=ef, ) - response = self._conn.call( self._conn.stub.SearchPoints, request, ) - return [pid.id.value for pid in response.result_point_ids] + def batch_search( + self, + *, + queries, + similarity: Similarity | None = None, + limit: int | None = None, + ef: int | None = None, + ) -> List[List[str]]: + """ + Flexible batch search. + + Accepts: + - List[SearchQuery] + - List[(DenseVector, Similarity, int)] + - List[(DenseVector, Similarity)] + global limit + - List[(DenseVector, int)] + global similarity + - List[DenseVector] + global similarity + limit + """ + normalized = [] + + for i, q in enumerate(queries): + if ( + hasattr(q, "vector") + and hasattr(q, "similarity") + and hasattr(q, "limit") + ): + normalized.append((q.vector, q.similarity, q.limit)) + continue + + if isinstance(q, DenseVector): + if similarity is None or limit is None: + raise ValueError( + f"queries[{i}] requires global similarity and limit" + ) + normalized.append((q, similarity, limit)) + continue + + if isinstance(q, (list, tuple)): + if len(q) == 3: + normalized.append(q) + continue + if len(q) == 2: + a, b = q + + if isinstance(a, DenseVector) and isinstance(b, Similarity): + if limit is None: + raise ValueError(f"queries[{i}] missing global limit") + normalized.append((a, b, limit)) + continue + + if isinstance(a, DenseVector) and isinstance(b, int): + if similarity is None: + raise ValueError(f"queries[{i}] missing global similarity") + normalized.append((a, similarity, b)) + continue + + raise TypeError(f"Invalid query format at index {i}") + + request = proto.build_batch_search_request(queries=normalized, ef=ef) + response = self._conn.call(self._conn.stub.SearchPointsBatch, request) + return [ + [pid.id.value for pid in result.result_point_ids] + for result in response.results + ] + + @staticmethod + def _validate_dense_vector(vector: DenseVector) -> None: + if not isinstance(vector, DenseVector): + raise TypeError( + "vector must be a DenseVector. Use: DenseVector([1.0, 2.0, 3.0])" + ) + def close(self) -> None: """ Close the gRPC connection. diff --git a/client/python/vortexdb/config.py b/client/python/vortexdb/config.py index 46bd508..d85be84 100644 --- a/client/python/vortexdb/config.py +++ b/client/python/vortexdb/config.py @@ -5,6 +5,8 @@ DEFAULT_GRPC_HOST = "localhost" DEFAULT_GRPC_PORT = 50051 DEFAULT_TIMEOUT = 5.0 + + @dataclass(frozen=True) class VortexDBConfig: """Configuration for the VortexDB Python client""" @@ -20,7 +22,7 @@ def from_env( api_key: str | None = None, timeout: float | None = None, ) -> "VortexDBConfig": - """ Load configuration from explicit arguments with environment variable fallback """ + """Load configuration from explicit arguments with environment variable fallback""" resolved_grpc_url = ( grpc_url diff --git a/client/python/vortexdb/connection.py b/client/python/vortexdb/connection.py index d7a61a5..bd348ac 100644 --- a/client/python/vortexdb/connection.py +++ b/client/python/vortexdb/connection.py @@ -1,31 +1,21 @@ import grpc from typing import Any, Callable +from vortexdb._grpc_common import build_auth_metadata, map_grpc_error from vortexdb.config import VortexDBConfig -from vortexdb.exceptions import ( - AuthenticationError, - NotFoundError, - InvalidArgumentError, - TimeoutError, - ServiceUnavailableError, - InternalServerError, - VortexDBError, -) from vortexdb.grpc.vector_db_pb2_grpc import VectorDBStub class GRPCConnection: - """ gRPC connection wrapper for VortexDB""" + """gRPC connection wrapper for VortexDB""" def __init__(self, config: VortexDBConfig): self._config = config self._channel = grpc.insecure_channel(config.grpc_url) self._stub = VectorDBStub(self._channel) # Because this is required in every request - self._metadata = ( - ("authorization", f"Bearer {config.api_key}"), - ) + self._metadata = build_auth_metadata(config.api_key) @property def stub(self) -> VectorDBStub: @@ -36,7 +26,7 @@ def call( rpc: Callable[..., Any], request: Any, ) -> Any: - """ Execute a gRPC call with standard error handling """ + """Execute a gRPC call with standard error handling""" try: return rpc( request, @@ -45,29 +35,8 @@ def call( ) except grpc.RpcError as e: - raise self._map_grpc_error(e) from e + raise map_grpc_error(e) from e def close(self) -> None: - """ Close the underlying gRPC channel """ + """Close the underlying gRPC channel""" self._channel.close() - - @staticmethod - def _map_grpc_error(error: grpc.RpcError) -> VortexDBError: - code = error.code() - - if code == grpc.StatusCode.UNAUTHENTICATED: - return AuthenticationError(error.details()) - - if code == grpc.StatusCode.NOT_FOUND: - return NotFoundError(error.details()) - - if code == grpc.StatusCode.INVALID_ARGUMENT: - return InvalidArgumentError(error.details()) - - if code == grpc.StatusCode.DEADLINE_EXCEEDED: - return TimeoutError(error.details()) - - if code == grpc.StatusCode.UNAVAILABLE: - return ServiceUnavailableError(error.details()) - - return InternalServerError(error.details()) diff --git a/client/python/vortexdb/exceptions.py b/client/python/vortexdb/exceptions.py index 7a4498b..07d914a 100644 --- a/client/python/vortexdb/exceptions.py +++ b/client/python/vortexdb/exceptions.py @@ -25,6 +25,6 @@ class ServiceUnavailableError(VortexDBError): class InternalServerError(VortexDBError): """Internal error in the server""" + class ConfigurationError(VortexDBError): """Invalid or missing client configuration.""" - diff --git a/client/python/vortexdb/grpc/vector_db_pb2.py b/client/python/vortexdb/grpc/vector_db_pb2.py index 2b8cbb8..faa4212 100644 --- a/client/python/vortexdb/grpc/vector_db_pb2.py +++ b/client/python/vortexdb/grpc/vector_db_pb2.py @@ -2,7 +2,7 @@ # Generated by the protocol buffer compiler. DO NOT EDIT! # NO CHECKED-IN PROTOBUF GENCODE # source: vector-db.proto -# Protobuf Python Version: 6.31.1 +# Protobuf Python Version: 6.33.5 """Generated protocol buffer code.""" from google.protobuf import descriptor as _descriptor from google.protobuf import descriptor_pool as _descriptor_pool @@ -12,8 +12,8 @@ _runtime_version.ValidateProtobufRuntimeVersion( _runtime_version.Domain.PUBLIC, 6, - 31, - 1, + 33, + 5, '', 'vector-db.proto' ) @@ -25,33 +25,41 @@ from google.protobuf import empty_pb2 as google_dot_protobuf_dot_empty__pb2 -DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x0fvector-db.proto\x12\x08vectordb\x1a\x1bgoogle/protobuf/empty.proto\"\x15\n\x04UUID\x12\r\n\x05value\x18\x01 \x01(\t\"`\n\x13InsertVectorRequest\x12%\n\x06vector\x18\x01 \x01(\x0b\x32\x15.vectordb.DenseVector\x12\"\n\x07payload\x18\x02 \x01(\x0b\x32\x11.vectordb.Payload\"u\n\rSearchRequest\x12+\n\x0cquery_vector\x18\x01 \x01(\x0b\x32\x15.vectordb.DenseVector\x12(\n\nsimilarity\x18\x02 \x01(\x0e\x32\x14.vectordb.Similarity\x12\r\n\x05limit\x18\x03 \x01(\x04\"=\n\x0eSearchResponse\x12+\n\x10result_point_ids\x18\x01 \x03(\x0b\x32\x11.vectordb.PointID\"\x1d\n\x0b\x44\x65nseVector\x12\x0e\n\x06values\x18\x01 \x03(\x02\"q\n\x05Point\x12\x1d\n\x02id\x18\x01 \x01(\x0b\x32\x11.vectordb.PointID\x12\"\n\x07payload\x18\x02 \x01(\x0b\x32\x11.vectordb.Payload\x12%\n\x06vector\x18\x03 \x01(\x0b\x32\x15.vectordb.DenseVector\"%\n\x07PointID\x12\x1a\n\x02id\x18\x01 \x01(\x0b\x32\x0e.vectordb.UUID\"G\n\x07Payload\x12+\n\x0c\x63ontent_type\x18\x01 \x01(\x0e\x32\x15.vectordb.ContentType\x12\x0f\n\x07\x63ontent\x18\x02 \x01(\t*C\n\nSimilarity\x12\r\n\tEuclidean\x10\x00\x12\r\n\tManhattan\x10\x01\x12\x0b\n\x07Hamming\x10\x02\x12\n\n\x06\x43osine\x10\x03*\"\n\x0b\x43ontentType\x12\t\n\x05Image\x10\x00\x12\x08\n\x04Text\x10\x01\x32\x81\x02\n\x08VectorDB\x12\x42\n\x0cInsertVector\x12\x1d.vectordb.InsertVectorRequest\x1a\x11.vectordb.PointID\"\x00\x12:\n\x0b\x44\x65letePoint\x12\x11.vectordb.PointID\x1a\x16.google.protobuf.Empty\"\x00\x12\x30\n\x08GetPoint\x12\x11.vectordb.PointID\x1a\x0f.vectordb.Point\"\x00\x12\x43\n\x0cSearchPoints\x12\x17.vectordb.SearchRequest\x1a\x18.vectordb.SearchResponse\"\x00\x62\x06proto3') +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x0fvector-db.proto\x12\x08vectordb\x1a\x1bgoogle/protobuf/empty.proto\"\x15\n\x04UUID\x12\r\n\x05value\x18\x01 \x01(\t\"`\n\x13InsertVectorRequest\x12%\n\x06vector\x18\x01 \x01(\x0b\x32\x15.vectordb.DenseVector\x12\"\n\x07payload\x18\x02 \x01(\x0b\x32\x11.vectordb.Payload\"\x81\x01\n\rSearchRequest\x12+\n\x0cquery_vector\x18\x01 \x01(\x0b\x32\x15.vectordb.DenseVector\x12(\n\nsimilarity\x18\x02 \x01(\x0e\x32\x14.vectordb.Similarity\x12\r\n\x05limit\x18\x03 \x01(\x04\x12\n\n\x02\x65\x66\x18\x04 \x01(\x04\"=\n\x0eSearchResponse\x12+\n\x10result_point_ids\x18\x01 \x03(\x0b\x32\x11.vectordb.PointID\"\x1d\n\x0b\x44\x65nseVector\x12\x0e\n\x06values\x18\x01 \x03(\x02\"q\n\x05Point\x12\x1d\n\x02id\x18\x01 \x01(\x0b\x32\x11.vectordb.PointID\x12\"\n\x07payload\x18\x02 \x01(\x0b\x32\x11.vectordb.Payload\x12%\n\x06vector\x18\x03 \x01(\x0b\x32\x15.vectordb.DenseVector\"%\n\x07PointID\x12\x1a\n\x02id\x18\x01 \x01(\x0b\x32\x0e.vectordb.UUID\"G\n\x07Payload\x12+\n\x0c\x63ontent_type\x18\x01 \x01(\x0e\x32\x15.vectordb.ContentType\x12\x0f\n\x07\x63ontent\x18\x02 \x01(\t\"K\n\x19InsertVectorsBatchRequest\x12.\n\x07vectors\x18\x01 \x03(\x0b\x32\x1d.vectordb.InsertVectorRequest\"<\n\x1aInsertVectorsBatchResponse\x12\x1e\n\x03ids\x18\x01 \x03(\x0b\x32\x11.vectordb.PointID\"D\n\x18SearchPointsBatchRequest\x12(\n\x07queries\x18\x01 \x03(\x0b\x32\x17.vectordb.SearchRequest\"F\n\x19SearchPointsBatchResponse\x12)\n\x07results\x18\x01 \x03(\x0b\x32\x18.vectordb.SearchResponse*C\n\nSimilarity\x12\r\n\tEuclidean\x10\x00\x12\r\n\tManhattan\x10\x01\x12\x0b\n\x07Hamming\x10\x02\x12\n\n\x06\x43osine\x10\x03*\"\n\x0b\x43ontentType\x12\t\n\x05Image\x10\x00\x12\x08\n\x04Text\x10\x01\x32\xc4\x03\n\x08VectorDB\x12\x42\n\x0cInsertVector\x12\x1d.vectordb.InsertVectorRequest\x1a\x11.vectordb.PointID\"\x00\x12:\n\x0b\x44\x65letePoint\x12\x11.vectordb.PointID\x1a\x16.google.protobuf.Empty\"\x00\x12\x30\n\x08GetPoint\x12\x11.vectordb.PointID\x1a\x0f.vectordb.Point\"\x00\x12\x43\n\x0cSearchPoints\x12\x17.vectordb.SearchRequest\x1a\x18.vectordb.SearchResponse\"\x00\x12\x61\n\x12InsertVectorsBatch\x12#.vectordb.InsertVectorsBatchRequest\x1a$.vectordb.InsertVectorsBatchResponse\"\x00\x12^\n\x11SearchPointsBatch\x12\".vectordb.SearchPointsBatchRequest\x1a#.vectordb.SearchPointsBatchResponse\"\x00\x62\x06proto3') _globals = globals() _builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) _builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, 'vector_db_pb2', _globals) if not _descriptor._USE_C_DESCRIPTORS: DESCRIPTOR._loaded_options = None - _globals['_SIMILARITY']._serialized_start=619 - _globals['_SIMILARITY']._serialized_end=686 - _globals['_CONTENTTYPE']._serialized_start=688 - _globals['_CONTENTTYPE']._serialized_end=722 + _globals['_SIMILARITY']._serialized_start=913 + _globals['_SIMILARITY']._serialized_end=980 + _globals['_CONTENTTYPE']._serialized_start=982 + _globals['_CONTENTTYPE']._serialized_end=1016 _globals['_UUID']._serialized_start=58 _globals['_UUID']._serialized_end=79 _globals['_INSERTVECTORREQUEST']._serialized_start=81 _globals['_INSERTVECTORREQUEST']._serialized_end=177 - _globals['_SEARCHREQUEST']._serialized_start=179 - _globals['_SEARCHREQUEST']._serialized_end=296 - _globals['_SEARCHRESPONSE']._serialized_start=298 - _globals['_SEARCHRESPONSE']._serialized_end=359 - _globals['_DENSEVECTOR']._serialized_start=361 - _globals['_DENSEVECTOR']._serialized_end=390 - _globals['_POINT']._serialized_start=392 - _globals['_POINT']._serialized_end=505 - _globals['_POINTID']._serialized_start=507 - _globals['_POINTID']._serialized_end=544 - _globals['_PAYLOAD']._serialized_start=546 - _globals['_PAYLOAD']._serialized_end=617 - _globals['_VECTORDB']._serialized_start=725 - _globals['_VECTORDB']._serialized_end=982 + _globals['_SEARCHREQUEST']._serialized_start=180 + _globals['_SEARCHREQUEST']._serialized_end=309 + _globals['_SEARCHRESPONSE']._serialized_start=311 + _globals['_SEARCHRESPONSE']._serialized_end=372 + _globals['_DENSEVECTOR']._serialized_start=374 + _globals['_DENSEVECTOR']._serialized_end=403 + _globals['_POINT']._serialized_start=405 + _globals['_POINT']._serialized_end=518 + _globals['_POINTID']._serialized_start=520 + _globals['_POINTID']._serialized_end=557 + _globals['_PAYLOAD']._serialized_start=559 + _globals['_PAYLOAD']._serialized_end=630 + _globals['_INSERTVECTORSBATCHREQUEST']._serialized_start=632 + _globals['_INSERTVECTORSBATCHREQUEST']._serialized_end=707 + _globals['_INSERTVECTORSBATCHRESPONSE']._serialized_start=709 + _globals['_INSERTVECTORSBATCHRESPONSE']._serialized_end=769 + _globals['_SEARCHPOINTSBATCHREQUEST']._serialized_start=771 + _globals['_SEARCHPOINTSBATCHREQUEST']._serialized_end=839 + _globals['_SEARCHPOINTSBATCHRESPONSE']._serialized_start=841 + _globals['_SEARCHPOINTSBATCHRESPONSE']._serialized_end=911 + _globals['_VECTORDB']._serialized_start=1019 + _globals['_VECTORDB']._serialized_end=1471 # @@protoc_insertion_point(module_scope) diff --git a/client/python/vortexdb/grpc/vector_db_pb2_grpc.py b/client/python/vortexdb/grpc/vector_db_pb2_grpc.py index edc3c8f..1442f69 100644 --- a/client/python/vortexdb/grpc/vector_db_pb2_grpc.py +++ b/client/python/vortexdb/grpc/vector_db_pb2_grpc.py @@ -4,9 +4,9 @@ import warnings from google.protobuf import empty_pb2 as google_dot_protobuf_dot_empty__pb2 -from vortexdb.grpc import vector_db_pb2 as vector__db__pb2 +from . import vector_db_pb2 as vector__db__pb2 -GRPC_GENERATED_VERSION = '1.76.0' +GRPC_GENERATED_VERSION = '1.81.1' GRPC_VERSION = grpc.__version__ _version_not_supported = False @@ -26,7 +26,7 @@ ) -class VectorDBStub(object): +class VectorDBStub: """Missing associated documentation comment in .proto file.""" def __init__(self, channel): @@ -55,9 +55,19 @@ def __init__(self, channel): request_serializer=vector__db__pb2.SearchRequest.SerializeToString, response_deserializer=vector__db__pb2.SearchResponse.FromString, _registered_method=True) + self.InsertVectorsBatch = channel.unary_unary( + '/vectordb.VectorDB/InsertVectorsBatch', + request_serializer=vector__db__pb2.InsertVectorsBatchRequest.SerializeToString, + response_deserializer=vector__db__pb2.InsertVectorsBatchResponse.FromString, + _registered_method=True) + self.SearchPointsBatch = channel.unary_unary( + '/vectordb.VectorDB/SearchPointsBatch', + request_serializer=vector__db__pb2.SearchPointsBatchRequest.SerializeToString, + response_deserializer=vector__db__pb2.SearchPointsBatchResponse.FromString, + _registered_method=True) -class VectorDBServicer(object): +class VectorDBServicer: """Missing associated documentation comment in .proto file.""" def InsertVector(self, request, context): @@ -88,6 +98,18 @@ def SearchPoints(self, request, context): context.set_details('Method not implemented!') raise NotImplementedError('Method not implemented!') + def InsertVectorsBatch(self, request, context): + """Missing associated documentation comment in .proto file.""" + context.set_code(grpc.StatusCode.UNIMPLEMENTED) + context.set_details('Method not implemented!') + raise NotImplementedError('Method not implemented!') + + def SearchPointsBatch(self, request, context): + """Missing associated documentation comment in .proto file.""" + context.set_code(grpc.StatusCode.UNIMPLEMENTED) + context.set_details('Method not implemented!') + raise NotImplementedError('Method not implemented!') + def add_VectorDBServicer_to_server(servicer, server): rpc_method_handlers = { @@ -111,6 +133,16 @@ def add_VectorDBServicer_to_server(servicer, server): request_deserializer=vector__db__pb2.SearchRequest.FromString, response_serializer=vector__db__pb2.SearchResponse.SerializeToString, ), + 'InsertVectorsBatch': grpc.unary_unary_rpc_method_handler( + servicer.InsertVectorsBatch, + request_deserializer=vector__db__pb2.InsertVectorsBatchRequest.FromString, + response_serializer=vector__db__pb2.InsertVectorsBatchResponse.SerializeToString, + ), + 'SearchPointsBatch': grpc.unary_unary_rpc_method_handler( + servicer.SearchPointsBatch, + request_deserializer=vector__db__pb2.SearchPointsBatchRequest.FromString, + response_serializer=vector__db__pb2.SearchPointsBatchResponse.SerializeToString, + ), } generic_handler = grpc.method_handlers_generic_handler( 'vectordb.VectorDB', rpc_method_handlers) @@ -119,7 +151,7 @@ def add_VectorDBServicer_to_server(servicer, server): # This class is part of an EXPERIMENTAL API. -class VectorDB(object): +class VectorDB: """Missing associated documentation comment in .proto file.""" @staticmethod @@ -229,3 +261,57 @@ def SearchPoints(request, timeout, metadata, _registered_method=True) + + @staticmethod + def InsertVectorsBatch(request, + target, + options=(), + channel_credentials=None, + call_credentials=None, + insecure=False, + compression=None, + wait_for_ready=None, + timeout=None, + metadata=None): + return grpc.experimental.unary_unary( + request, + target, + '/vectordb.VectorDB/InsertVectorsBatch', + vector__db__pb2.InsertVectorsBatchRequest.SerializeToString, + vector__db__pb2.InsertVectorsBatchResponse.FromString, + options, + channel_credentials, + insecure, + call_credentials, + compression, + wait_for_ready, + timeout, + metadata, + _registered_method=True) + + @staticmethod + def SearchPointsBatch(request, + target, + options=(), + channel_credentials=None, + call_credentials=None, + insecure=False, + compression=None, + wait_for_ready=None, + timeout=None, + metadata=None): + return grpc.experimental.unary_unary( + request, + target, + '/vectordb.VectorDB/SearchPointsBatch', + vector__db__pb2.SearchPointsBatchRequest.SerializeToString, + vector__db__pb2.SearchPointsBatchResponse.FromString, + options, + channel_credentials, + insecure, + call_credentials, + compression, + wait_for_ready, + timeout, + metadata, + _registered_method=True) diff --git a/client/python/vortexdb/models.py b/client/python/vortexdb/models.py index f2cbe19..0ba1e6c 100644 --- a/client/python/vortexdb/models.py +++ b/client/python/vortexdb/models.py @@ -4,11 +4,12 @@ from vortexdb.grpc import vector_db_pb2 -# I found this to be a good idea, because +# I found this to be a good idea, because # 1. readability # 2. will help in HTTP client # 3. transport conversion at the very end, won't break if proto enum changes + class Similarity(Enum): EUCLIDEAN = "euclidean" MANHATTAN = "manhattan" @@ -57,20 +58,21 @@ def __post_init__(self): for v in self.values: if not isinstance(v, (int, float)): - raise TypeError( - "DenseVector values must be numeric (int or float)" - ) + raise TypeError("DenseVector values must be numeric (int or float)") # force float normalization object.__setattr__(self, "values", [float(v) for v in self.values]) - + def to_proto(self) -> vector_db_pb2.DenseVector: return vector_db_pb2.DenseVector(values=self.values) - + def to_list(self) -> list[float]: return list(self.values) +# & Helper Function for Batch of DenseVectors +def to_dense_vectors(arr): + return [DenseVector(x) for x in arr] @dataclass(frozen=True) @@ -90,7 +92,6 @@ def __post_init__(self): if not isinstance(self.content_type, ContentType): raise TypeError("content_type must be ContentType enum") - def to_proto(self) -> vector_db_pb2.Payload: return vector_db_pb2.Payload( content_type=self.content_type.to_proto(), @@ -120,7 +121,7 @@ def from_proto(proto: vector_db_pb2.Point) -> "Point": vector=DenseVector(list(proto.vector.values)), payload=payload_obj, ) - + def pretty(self) -> str: return ( f"\nPoint:\n id = {self.id},\n" @@ -129,3 +130,18 @@ def pretty(self) -> str: f" payload_type = {self.payload.content_type.name},\n" f" payload = '{self.payload.content}'" ) + + +# I added this because using tuples will get messy if we increase fields in a search query +@dataclass(frozen=True) +class SearchQuery: + vector: DenseVector + similarity: Similarity + limit: int + + def to_proto(self) -> vector_db_pb2.SearchRequest: + return vector_db_pb2.SearchRequest( + query_vector=self.vector.to_proto(), + similarity=self.similarity.to_proto(), + limit=self.limit, + ) diff --git a/client/python/vortexdb/protoutils.py b/client/python/vortexdb/protoutils.py index 8562b18..c6aae69 100644 --- a/client/python/vortexdb/protoutils.py +++ b/client/python/vortexdb/protoutils.py @@ -1,6 +1,7 @@ from vortexdb.grpc import vector_db_pb2 from vortexdb.models import DenseVector, Payload, Similarity + def build_insert_request( *, vector: DenseVector, @@ -11,19 +12,95 @@ def build_insert_request( payload=payload.to_proto(), ) + +def build_batch_insert_request( + *, + items: list[tuple[DenseVector, Payload]], +) -> vector_db_pb2.InsertVectorsBatchRequest: + if not isinstance(items, (list, tuple)): + raise TypeError("Items must be a list of (DenseVector, Payload) tuples") + + if not items: + raise ValueError("Items cannot be empty") + + requests = [] + + for i, pair in enumerate(items): + if not isinstance(pair, (list, tuple)) or len(pair) != 2: + raise TypeError(f"items[{i}] must be a tuple of (DenseVector, Payload)") + + vector, payload = pair + if not isinstance(vector, DenseVector): + raise TypeError( + f"items[{i}][0] must be a DenseVectorUse: DenseVector([1.0, 2.0, 3.0])" + ) + + if not isinstance(payload, Payload): + raise TypeError(f"items[{i}][1] must be Payload") + + requests.append(build_insert_request(vector=vector, payload=payload)) + + return vector_db_pb2.InsertVectorsBatchRequest(vectors=requests) + + def build_point_id_request(point_id: str) -> vector_db_pb2.PointID: - return vector_db_pb2.PointID( - id=vector_db_pb2.UUID(value=point_id) - ) + return vector_db_pb2.PointID(id=vector_db_pb2.UUID(value=point_id)) + def build_search_request( *, vector: DenseVector, similarity: Similarity, limit: int, + ef: int | None = None, ) -> vector_db_pb2.SearchRequest: return vector_db_pb2.SearchRequest( query_vector=vector.to_proto(), similarity=similarity.to_proto(), limit=limit, + ef=ef or 0, ) + + +def build_batch_search_request( + *, + queries: list[tuple[DenseVector, Similarity, int]], + ef: int | None = None, +) -> vector_db_pb2.SearchPointsBatchRequest: + if not isinstance(queries, (list, tuple)): + raise TypeError( + "Queries must be a list of (DenseVector, Similarity, Limit (int)) tuples" + ) + + if not queries: + raise ValueError("Queries cannot be empty") + + requests = [] + + for i, trio in enumerate(queries): + if not isinstance(trio, (list, tuple)) or len(trio) != 3: + raise TypeError( + f"queries[{i}] must be a tuple of (DenseVector, Similarity, Limit(int))" + ) + + vector, similarity, limit = trio + if not isinstance(vector, DenseVector): + raise TypeError( + f"queries[{i}][0] must be a DenseVector" + "Use: DenseVector([1.0, 2.0, 3.0])" + ) + if not isinstance(similarity, Similarity): + raise TypeError(f"queries[{i}][1] must be Similarity") + if not isinstance(limit, int): + raise TypeError(f"queries[{i}][2] must be an integer value") + + requests.append( + vector_db_pb2.SearchRequest( + query_vector=vector.to_proto(), + similarity=similarity.to_proto(), + limit=limit, + ef=ef, + ) + ) + + return vector_db_pb2.SearchPointsBatchRequest(queries=requests) diff --git a/crates/api/src/lib.rs b/crates/api/src/lib.rs index de028c0..08f89ea 100644 --- a/crates/api/src/lib.rs +++ b/crates/api/src/lib.rs @@ -1,7 +1,7 @@ -use defs::{DbError, Dimension, IndexedVector, Similarity, SnapshottableDb}; -use defs::{DenseVector, Payload, Point, PointId}; -use index::hnsw::HnswIndex; -use index::kd_tree::index::KDTree; +use defs::{DbError, Dimension, IndexedVector, SearchQueryInput, Similarity, SnapshottableDb}; +use defs::{DenseVector, Payload, Point, PointId, PointInput}; +use index::hnsw::{HnswConfig, HnswIndex}; +use index::kd_tree::{KDTree, KDTreeConfig}; use std::path::{Path, PathBuf}; use tempfile::tempdir; // use std::sync::atomic::{AtomicU64, Ordering}; @@ -10,8 +10,7 @@ use std::sync::{Arc, RwLock}; use index::flat::index::FlatIndex; use index::{IndexType, VectorIndex}; use snapshot::Snapshot; -use storage::rocks_db::RocksDbStorage; -use storage::{StorageEngine, StorageType, VectorPage}; +use storage::{StorageEngine, StorageType, VectorPage, create_storage_engine}; use uuid::Uuid; @@ -70,6 +69,37 @@ impl VectorDb { Ok(point_id) } + pub fn insert_batch(&self, points: Vec) -> Result> { + let mut ids = Vec::with_capacity(points.len()); + + for point in points { + let id = point.id.unwrap_or_else(Uuid::new_v4); + let vector = point.vector; + let payload = point.payload; + + if let Some(ref v) = vector + && v.len() != self.dimension + { + return Err(ApiError::DimensionMismatch { + expected: self.dimension, + got: v.len(), + }); + } + + self.storage.insert_point(id, vector.clone(), payload)?; + + if let Some(v) = vector { + let indexed = IndexedVector { id, vector: v }; + let mut index = self.index.write().map_err(|_| ApiError::LockError)?; + index.insert(indexed)?; + } + + ids.push(id); + } + + Ok(ids) + } + //TODO: Make this an atomic operation pub fn delete(&self, id: PointId) -> Result { // Remove from storage @@ -95,22 +125,17 @@ impl VectorDb { } } - pub fn search( - &self, - query: DenseVector, - similarity: Similarity, - limit: usize, - ) -> Result> { + pub fn search(&self, query: SearchQueryInput) -> Result> { // Validate search limit - if limit == 0 { - return Err(ApiError::InvalidSearchLimit { limit }); + if query.limit == 0 { + return Err(ApiError::InvalidSearchLimit { limit: query.limit }); } // Validate query dimension - if query.len() != self.dimension { + if query.vector.len() != self.dimension { return Err(ApiError::DimensionMismatch { expected: self.dimension, - got: query.len(), + got: query.vector.len(), }); } @@ -118,11 +143,25 @@ impl VectorDb { let index = self.index.read().map_err(|_| ApiError::LockError)?; //TODO: Add feat of returning similarity scores in the search - let vectors = index.search(query, similarity, limit)?; + let vectors = + index.search_with_ef(query.vector, query.similarity, query.limit, query.ef)?; Ok(vectors) } + pub fn search_batch(&self, queries: Vec) -> Result>> { + let mut results = Vec::with_capacity(queries.len()); + let index = self.index.read().unwrap(); + + for query in queries { + let found = + index.search_with_ef(query.vector, query.similarity, query.limit, query.ef)?; + results.push(found); + } + + Ok(results) + } + pub fn list(&self, offset: PointId, limit: usize) -> Result> { let page = self.storage.list_vectors(offset, limit)?; Ok(page) @@ -188,6 +227,8 @@ pub struct DbConfig { pub data_path: PathBuf, pub dimension: Dimension, pub similarity: Similarity, + pub hnsw_config: HnswConfig, + pub kd_tree_config: KDTreeConfig, } #[derive(Debug)] @@ -214,18 +255,19 @@ pub fn restore_from_snapshot(config: &DbRestoreConfig) -> Result Result { // Initialize the storage engine - let storage = match config.storage_type { - StorageType::RocksDb => Arc::new(RocksDbStorage::new(config.data_path)?), - _ => Arc::new(RocksDbStorage::new(config.data_path)?), - }; + let storage = create_storage_engine(config.storage_type, config.data_path)?; // Initialize the vector index let index: Arc> = match config.index_type { IndexType::Flat => Arc::new(RwLock::new(FlatIndex::new())), - IndexType::KDTree => Arc::new(RwLock::new(KDTree::build_empty(config.dimension))), - IndexType::HNSW => Arc::new(RwLock::new(HnswIndex::new( + IndexType::KDTree => Arc::new(RwLock::new(KDTree::build_empty_with_config( + config.dimension, + config.kd_tree_config, + ))), + IndexType::HNSW => Arc::new(RwLock::new(HnswIndex::with_config( config.similarity, config.dimension, + config.hnsw_config, ))), }; @@ -252,17 +294,30 @@ mod tests { // Helper function to create a test database fn create_test_db() -> (VectorDb, TempDir) { + create_test_db_with_storage(StorageType::RocksDb) + } + + fn create_test_db_with_storage(storage_type: StorageType) -> (VectorDb, TempDir) { let temp_dir = tempdir().unwrap(); let config = DbConfig { - storage_type: StorageType::RocksDb, + storage_type, index_type: IndexType::Flat, data_path: temp_dir.path().to_path_buf(), dimension: 3, similarity: Similarity::Cosine, + hnsw_config: HnswConfig::default(), + kd_tree_config: KDTreeConfig::default(), }; (init_api(config).unwrap(), temp_dir) } + fn test_payload(content: &str) -> Payload { + Payload { + content_type: ContentType::Text, + content: content.to_string(), + } + } + #[test] fn test_insert_and_get() { let (db, _temp_dir) = create_test_db(); @@ -288,6 +343,20 @@ mod tests { assert_eq!(point.payload.as_ref().unwrap().content, "Test content"); } + #[test] + fn test_insert_and_get_with_in_memory_storage() { + let (db, _temp_dir) = create_test_db_with_storage(StorageType::InMemory); + let vector = vec![1.0, 2.0, 3.0]; + let payload = test_payload("Test content"); + + let id = db.insert(vector.clone(), payload.clone()).unwrap(); + let point = db.get(id).unwrap().unwrap(); + + assert_eq!(point.id, id); + assert_eq!(point.vector, Some(vector)); + assert_eq!(point.payload, Some(payload)); + } + #[test] fn test_dimension_mismatch() { let (db, _temp_dir) = create_test_db(); @@ -359,7 +428,14 @@ mod tests { // Search for the closest vector to [1.0, 0.1, 0.1] let query = vec![1.0, 0.1, 0.1]; - let results = db.search(query, Similarity::Cosine, 1).unwrap(); + let results = db + .search(SearchQueryInput { + vector: query, + similarity: Similarity::Cosine, + limit: 1, + ef: None, + }) + .unwrap(); assert_eq!(results.len(), 1); assert_eq!(results[0], ids[0]); // The first vector should be closest @@ -387,7 +463,14 @@ mod tests { // Search with limit 3 let query = vec![0.0, 0.0, 0.0]; - let results = db.search(query, Similarity::Euclidean, 3).unwrap(); + let results = db + .search(SearchQueryInput { + vector: query, + similarity: Similarity::Euclidean, + limit: 3, + ef: None, + }) + .unwrap(); assert_eq!(results.len(), 3); } @@ -397,7 +480,12 @@ mod tests { let (db, _temp_dir) = create_test_db(); let query = vec![1.0, 2.0, 3.0]; - let result = db.search(query, Similarity::Cosine, 0); + let result = db.search(SearchQueryInput { + vector: query, + similarity: Similarity::Cosine, + limit: 0, + ef: None, + }); assert!(result.is_err()); match result.unwrap_err() { @@ -416,7 +504,14 @@ mod tests { assert!(db.get(Uuid::new_v4()).unwrap().is_none()); let query = vec![1.0, 2.0, 3.0]; - let results = db.search(query, Similarity::Cosine, 10).unwrap(); + let results = db + .search(SearchQueryInput { + vector: query, + similarity: Similarity::Cosine, + limit: 10, + ef: None, + }) + .unwrap(); assert_eq!(results.len(), 0); } @@ -533,6 +628,34 @@ mod tests { assert!(loaded_db.get(id2).unwrap().unwrap().vector.unwrap() == v2); } + #[test] + fn test_create_and_load_snapshot_with_in_memory_storage() { + let (old_db, temp_dir) = create_test_db_with_storage(StorageType::InMemory); + + let v1 = vec![0.0, 1.0, 2.0]; + let v2 = vec![3.0, 4.0, 5.0]; + let v3 = vec![6.0, 7.0, 8.0]; + + let id1 = old_db.insert(v1.clone(), test_payload("one")).unwrap(); + let id2 = old_db.insert(v2.clone(), test_payload("two")).unwrap(); + + let temp_snapshot_dir = tempdir().unwrap(); + let snapshot_path = old_db.create_snapshot(temp_snapshot_dir.path()).unwrap(); + + let id3 = old_db.insert(v3, test_payload("three")).unwrap(); + + let reload_config = DbRestoreConfig { + data_path: temp_dir.path().to_path_buf(), + snapshot_path, + }; + + let loaded_db = restore_from_snapshot(&reload_config).unwrap(); + + assert_eq!(loaded_db.get(id1).unwrap().unwrap().vector, Some(v1)); + assert_eq!(loaded_db.get(id2).unwrap().unwrap().vector, Some(v2)); + assert!(loaded_db.get(id3).unwrap().is_none()); + } + #[test] fn test_snapshot_engine() { let (_db, _temp_dir) = create_test_db(); diff --git a/crates/defs/Cargo.toml b/crates/defs/Cargo.toml index 600b80c..81d8dd9 100644 --- a/crates/defs/Cargo.toml +++ b/crates/defs/Cargo.toml @@ -7,5 +7,7 @@ edition.workspace = true license.workspace = true [dependencies] +axum.workspace = true serde.workspace = true +snafu.workspace = true uuid.workspace = true diff --git a/crates/defs/src/error.rs b/crates/defs/src/error.rs index 8354768..fa0e852 100644 --- a/crates/defs/src/error.rs +++ b/crates/defs/src/error.rs @@ -24,16 +24,32 @@ pub enum DbError { PointNotFound { id: PointId }, } -#[derive(Debug)] +use axum::{http::StatusCode, response::IntoResponse}; +use snafu::Snafu; + +#[derive(Debug, Snafu)] pub enum ServerError { - Bind(io::Error), - Serve(io::Error), + #[snafu(display("Failed to bind: {source}"))] + Bind { source: io::Error }, + + #[snafu(display("Failed to serve: {source}"))] + Serve { source: io::Error }, } #[derive(Debug)] pub enum AppError { ServerError(ServerError), + Api(String), } +impl IntoResponse for AppError { + fn into_response(self) -> axum::response::Response { + let (status, message) = match self { + AppError::Api(msg) => (StatusCode::BAD_REQUEST, msg), + AppError::ServerError(e) => (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()), + }; + (status, message).into_response() + } +} // Error type for server pub type BoxError = Box; diff --git a/crates/defs/src/types.rs b/crates/defs/src/types.rs index 525df26..1507653 100644 --- a/crates/defs/src/types.rs +++ b/crates/defs/src/types.rs @@ -21,18 +21,26 @@ pub enum StoredVector { Dense(DenseVector), } -#[derive(Serialize, Deserialize, Clone, Copy, Debug, PartialEq)] +#[derive(Serialize, Deserialize, Clone, Copy, Debug, Default, PartialEq)] pub enum ContentType { + #[default] Text, Image, } -#[derive(Serialize, Deserialize, Clone, Debug, PartialEq)] +#[derive(Serialize, Deserialize, Clone, Debug, Default, PartialEq)] pub struct Payload { pub content_type: ContentType, pub content: String, } +#[derive(Serialize, Deserialize, Clone, Debug, PartialEq)] +pub struct PointInput { + pub id: Option, + pub vector: Option, + pub payload: Option, +} + #[derive(Serialize, Deserialize, Clone, Debug, PartialEq)] pub struct Point { pub id: PointId, @@ -55,6 +63,42 @@ pub enum Similarity { Cosine, } +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct BatchInsertRequest { + pub points: Vec, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct BatchInsertResponse { + pub inserted: usize, + pub ids: Vec, +} + +// For batch search +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct BatchSearchRequest { + pub queries: Vec, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct SearchQueryInput { + pub vector: DenseVector, + pub similarity: Similarity, + pub limit: usize, + #[serde(default)] + pub ef: Option, +} + +#[derive(Clone, Serialize, Deserialize, Debug)] +pub struct SearchResponse { + pub results: Vec, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct BatchSearchResponse { + pub results: Vec, +} + // Struct which stores the distance between a vector and query vector and implements ordering traits #[derive(Copy, Clone)] pub struct DistanceOrderedVector<'q> { diff --git a/crates/grpc/proto/vector-db.proto b/crates/grpc/proto/vector-db.proto index b3834e9..0b5dfe1 100644 --- a/crates/grpc/proto/vector-db.proto +++ b/crates/grpc/proto/vector-db.proto @@ -20,6 +20,9 @@ service VectorDB { //Search for the k nearest vectors to a target vector given a distance function rpc SearchPoints(SearchRequest) returns (SearchResponse) {} + + rpc InsertVectorsBatch(InsertVectorsBatchRequest) returns (InsertVectorsBatchResponse) {} +rpc SearchPointsBatch(SearchPointsBatchRequest) returns (SearchPointsBatchResponse) {} } @@ -33,6 +36,7 @@ message SearchRequest { DenseVector query_vector = 1; Similarity similarity = 2; uint64 limit = 3; + uint64 ef = 4; } @@ -71,3 +75,18 @@ message Payload { string content = 2; } +message InsertVectorsBatchRequest { + repeated InsertVectorRequest vectors = 1; +} + +message InsertVectorsBatchResponse { + repeated PointID ids = 1; +} + +message SearchPointsBatchRequest { + repeated SearchRequest queries = 1; +} + +message SearchPointsBatchResponse { + repeated SearchResponse results = 1; +} diff --git a/crates/grpc/src/error.rs b/crates/grpc/src/error.rs index c8bd1fd..faf02b4 100644 --- a/crates/grpc/src/error.rs +++ b/crates/grpc/src/error.rs @@ -135,6 +135,15 @@ impl From for GrpcError { StorageError::RocksDbFlush { source: _ } => GrpcError::Internal { message: "flush error".to_string(), }, + StorageError::InMemoryLock {} => GrpcError::Internal { + message: "failed to lock in-memory storage".to_string(), + }, + StorageError::InMemoryCheckpoint { msg } => GrpcError::Internal { + message: format!("in-memory checkpoint error: {}", msg), + }, + StorageError::InMemoryCheckpointIo { msg, source: _ } => GrpcError::Internal { + message: format!("in-memory checkpoint io error: {}", msg), + }, } } } diff --git a/crates/grpc/src/service.rs b/crates/grpc/src/service.rs index 6fcb0ad..6e0d2c1 100644 --- a/crates/grpc/src/service.rs +++ b/crates/grpc/src/service.rs @@ -1,15 +1,18 @@ use std::str::FromStr; use std::sync::Arc; +use crate::error::GrpcError; use crate::interceptors; use crate::service::vectordb::{ContentType, Uuid}; use crate::utils::log_rpc; use crate::{constants::SIMILARITY_PROTOBUFF_MAP, utils::ServerEndpoint}; +use defs::SearchQueryInput; use tonic::{Request, Response, Status, service::InterceptorLayer, transport::Server}; use tracing::{Level, event}; use uuid::Uuid as UuidCrate; use vectordb::{ - DenseVector, InsertVectorRequest, Point, PointId, SearchRequest, SearchResponse, + DenseVector, InsertVectorRequest, InsertVectorsBatchRequest, InsertVectorsBatchResponse, Point, + PointId, SearchPointsBatchRequest, SearchPointsBatchResponse, SearchRequest, SearchResponse, vector_db_server::{VectorDb, VectorDbServer}, }; @@ -129,7 +132,12 @@ impl VectorDb for VectorDBService { let result_point_ids = self .vector_db - .search(query_vect.values, *similarity, limit as usize) + .search(SearchQueryInput { + vector: query_vect.values, + similarity: *similarity, + limit: limit as usize, + ef: (search_request.ef > 0).then_some(search_request.ef as usize), + }) .map_err(|e| Status::from(crate::error::GrpcError::from(e)))?; // create a mapped vector of PointIds @@ -166,8 +174,83 @@ impl VectorDb for VectorDBService { Err(e) => Err(Status::from(crate::error::GrpcError::from(e))), } } -} + async fn insert_vectors_batch( + &self, + request: tonic::Request, + ) -> Result, tonic::Status> { + let req = request.into_inner(); + let mut ids = Vec::with_capacity(req.vectors.len()); + + for vec in req.vectors { + let payload = vec.payload.map(|p| defs::Payload { + content_type: match ContentType::try_from(p.content_type) + .unwrap_or(ContentType::Text) + { + ContentType::Text => defs::ContentType::Text, + ContentType::Image => defs::ContentType::Image, + }, + content: p.content, + }); + + let id = self + .vector_db + .insert( + vec.vector.unwrap_or_default().values, + payload.unwrap_or_default(), + ) + .map_err(|e| tonic::Status::internal(e.to_string()))?; + + ids.push(PointId { + id: Some(Uuid { + value: id.to_string(), + }), + }); + } + + Ok(tonic::Response::new(InsertVectorsBatchResponse { ids })) + } + + async fn search_points_batch( + &self, + request: tonic::Request, + ) -> Result, tonic::Status> { + let req = request.into_inner(); + let mut results = Vec::with_capacity(req.queries.len()); + + for query in req.queries { + let similarity = SIMILARITY_PROTOBUFF_MAP + .get(query.similarity as usize) + .ok_or(tonic::Status::invalid_argument("Invalid similarity"))?; + + let ids = self + .vector_db + .search(SearchQueryInput { + vector: query + .query_vector + .ok_or(tonic::Status::invalid_argument("missing query_vector"))? + .values, + similarity: *similarity, + limit: query.limit as usize, + ef: (query.ef > 0).then_some(query.ef as usize), + }) + .map_err(|e| tonic::Status::from(GrpcError::from(e)))?; + + results.push(SearchResponse { + result_point_ids: ids + .into_iter() + .map(|id| PointId { + id: Some(Uuid { + value: id.to_string(), + }), + }) + .collect(), + }); + } + + Ok(tonic::Response::new(SearchPointsBatchResponse { results })) + } +} pub async fn run_server( vector_db_service: VectorDBService, endpoint: ServerEndpoint, diff --git a/crates/grpc/src/tests.rs b/crates/grpc/src/tests.rs index ff04820..704dd8f 100644 --- a/crates/grpc/src/tests.rs +++ b/crates/grpc/src/tests.rs @@ -7,7 +7,7 @@ use crate::service::{VectorDBService, run_server}; use crate::utils::ServerEndpoint; use api::DbConfig; use defs::Similarity; -use index::IndexType; +use index::{IndexType, hnsw::HnswConfig, kd_tree::KDTreeConfig}; use std::net::SocketAddr; use std::sync::Arc; use storage::StorageType; @@ -35,6 +35,8 @@ async fn start_test_server() -> Result<(SocketAddr, TempDir), Box, + Json(request): Json, +) -> Result, AppError> { + let ids = state + .db + .insert_batch(request.points) + .map_err(|e| AppError::Api(e.to_string()))?; + + Ok(Json(BatchInsertResponse { + inserted: ids.len(), + ids, + })) +} + pub async fn get_point_handler( Path(point_id): Path, State(app_state): State, @@ -74,26 +92,11 @@ pub async fn delete_point_handler( } } -#[derive(Deserialize)] -pub struct SearchRequest { - pub vector: DenseVector, - pub similarity: Similarity, - pub limit: usize, -} - -#[derive(Serialize, Deserialize, Debug)] -pub struct SearchResponse { - pub results: Vec, -} - pub async fn search_points_handler( State(app_state): State, - Json(request): Json, + Json(request): Json, ) -> Result, (StatusCode, String)> { - match app_state - .db - .search(request.vector, request.similarity, request.limit) - { + match app_state.db.search(request) { Ok(results) => { let response = SearchResponse { results }; Ok(Json(response)) @@ -105,6 +108,21 @@ pub async fn search_points_handler( } } +pub async fn batch_search_handler( + State(state): State, + Json(request): Json, +) -> Result, AppError> { + let results = state + .db + .search_batch(request.queries) + .map_err(|e| AppError::Api(e.to_string()))? + .into_iter() + .map(|ids| SearchResponse { results: ids }) + .collect(); + + Ok(Json(BatchSearchResponse { results })) +} + /// Map `ApiError` into an HTTP `(StatusCode, String)` response. fn api_error_to_response(err: &ApiError) -> (StatusCode, String) { match err { @@ -131,7 +149,10 @@ fn api_error_to_response(err: &ApiError) -> (StatusCode, String) { | StorageError::RocksDbFlush { .. } | StorageError::RocksDbInitialization { .. } | StorageError::RocksDbCheckpointMsg { .. } - | StorageError::RocksDbCheckpointIo { .. } => { + | StorageError::RocksDbCheckpointIo { .. } + | StorageError::InMemoryLock { .. } + | StorageError::InMemoryCheckpoint { .. } + | StorageError::InMemoryCheckpointIo { .. } => { (StatusCode::INTERNAL_SERVER_ERROR, source.to_string()) } }, diff --git a/crates/http/src/lib.rs b/crates/http/src/lib.rs index 8554157..7843aff 100644 --- a/crates/http/src/lib.rs +++ b/crates/http/src/lib.rs @@ -12,8 +12,8 @@ use tokio::net::TcpListener; use tracing::info; use handler::{ - delete_point_handler, get_point_handler, health_handler, insert_point_handler, root_handler, - search_points_handler, + batch_insert_handler, batch_search_handler, delete_point_handler, get_point_handler, + health_handler, insert_point_handler, root_handler, search_points_handler, }; #[derive(Clone)] @@ -33,6 +33,8 @@ pub fn create_router(db: Arc) -> Router { get(get_point_handler).delete(delete_point_handler), ) .route("/points/search", post(search_points_handler)) + .route("/points/batch", post(batch_insert_handler)) + .route("/points/search/batch", post(batch_search_handler)) .with_state(app_state) } @@ -44,3 +46,69 @@ pub async fn run_http_server(db: Arc, addr: SocketAddr) -> Result<(), axum::serve(listener, app.into_make_service()).await?; Ok(()) } + +#[cfg(test)] +mod tests { + use super::*; + use api::DbConfig; + use axum::http::StatusCode; + use axum_test::TestServer; + use defs::Similarity; + use index::{IndexType, hnsw::HnswConfig, kd_tree::KDTreeConfig}; + use serde_json::json; + use storage::StorageType; + + #[tokio::test] + async fn in_memory_storage_http_smoke_test() { + let temp_dir = tempfile::tempdir().unwrap(); + let db = api::init_api(DbConfig { + storage_type: StorageType::InMemory, + index_type: IndexType::Flat, + data_path: temp_dir.path().to_path_buf(), + dimension: 3, + similarity: Similarity::Cosine, + hnsw_config: HnswConfig::default(), + kd_tree_config: KDTreeConfig::default(), + }) + .unwrap(); + let server = TestServer::new(create_router(Arc::new(db))).unwrap(); + + let insert_response = server + .post("/points") + .json(&json!({ + "vector": [1.0, 0.0, 0.0], + "payload": { + "content_type": "Text", + "content": "smoke-test" + } + })) + .await; + insert_response.assert_status(StatusCode::CREATED); + let insert_body: serde_json::Value = insert_response.json(); + let point_id = insert_body["point_id"].as_str().unwrap(); + + let get_response = server.get(&format!("/points/{point_id}")).await; + get_response.assert_status_ok(); + let point_body: serde_json::Value = get_response.json(); + assert_eq!(point_body["payload"]["content"], "smoke-test"); + assert_eq!(point_body["vector"], json!([1.0, 0.0, 0.0])); + + let search_response = server + .post("/points/search") + .json(&json!({ + "vector": [1.0, 0.0, 0.0], + "similarity": "Cosine", + "limit": 1 + })) + .await; + search_response.assert_status_ok(); + let search_body: serde_json::Value = search_response.json(); + assert_eq!(search_body["results"], json!([point_id])); + + let delete_response = server.delete(&format!("/points/{point_id}")).await; + delete_response.assert_status(StatusCode::NO_CONTENT); + + let missing_response = server.get(&format!("/points/{point_id}")).await; + missing_response.assert_status(StatusCode::NOT_FOUND); + } +} diff --git a/crates/index/src/hnsw/index.rs b/crates/index/src/hnsw/index.rs index 8bf91d8..f730198 100644 --- a/crates/index/src/hnsw/index.rs +++ b/crates/index/src/hnsw/index.rs @@ -4,11 +4,33 @@ use defs::{DenseVector, Dimension, IndexedVector, PointId, Similarity}; use uuid::Uuid; use crate::VectorIndex; -use crate::{IndexError, Result}; +use crate::{IndexError, Result, distance}; use super::types::{HnswStats, LevelGenerator, Node, PointIndexation}; use std::cmp::{max, min}; +#[derive(Debug, Clone, Copy)] +pub struct HnswConfig { + pub max_connections: usize, + pub max_connections_0: usize, + pub max_layer: usize, + pub ef_construction: usize, + pub ef: usize, +} + +impl Default for HnswConfig { + fn default() -> Self { + let max_connections = 16; + Self { + max_connections, + max_connections_0: 2 * max_connections, + max_layer: 16, + ef_construction: 200, + ef: 100, + } + } +} + pub struct HnswIndex { // Construction/search parameters pub ef_construction: usize, @@ -26,11 +48,19 @@ pub struct HnswIndex { impl HnswIndex { pub fn new(similarity: Similarity, data_dimension: Dimension) -> Self { - let max_connections = 16; - let max_connections_0 = 32; // M0 = 2 * M (common default) - let max_layer = 16; - let ef_construction = 200; - let ef = 100; + Self::with_config(similarity, data_dimension, HnswConfig::default()) + } + + pub fn with_config( + similarity: Similarity, + data_dimension: Dimension, + config: HnswConfig, + ) -> Self { + let max_connections = config.max_connections.max(2); + let max_connections_0 = config.max_connections_0.max(max_connections); + let max_layer = config.max_layer.max(1); + let ef_construction = config.ef_construction.max(1); + let ef = config.ef.max(1); let level_generator = LevelGenerator::from_m(max_connections); let index = PointIndexation { @@ -83,7 +113,7 @@ impl VectorIndex for HnswIndex { let new_id: PointId = vector.id; - let mut query_vec = vector.vector.clone(); + let mut query_vec = vector.vector; self.normalize_if_cosine(&mut query_vec); self.cache.insert(new_id, query_vec.clone()); @@ -173,11 +203,16 @@ impl VectorIndex for HnswIndex { /// - greedy descend from the top layer to level 1 /// - run ef-best-first at level 0 with ef0 = max(ef, k) /// - return up to k ids by ascending distance - fn search( + fn search(&self, query: DenseVector, similarity: Similarity, k: usize) -> Result> { + self.search_with_ef(query, similarity, k, None) + } + + fn search_with_ef( &self, mut query: DenseVector, _similarity: Similarity, k: usize, + ef: Option, ) -> Result> { if k == 0 { return Ok(Vec::new()); @@ -209,7 +244,7 @@ impl VectorIndex for HnswIndex { ep = self.greedy_search_layer(ep, level, &query)?; } } - let ef0 = max(self.ef, k); + let ef0 = max(ef.unwrap_or(self.ef), k); let mut w = self.search_layer_for_insert(ep, 0, &query, ef0)?; w.truncate(k); let result: Vec = w.into_iter().map(|(id, _)| id).collect(); @@ -293,4 +328,23 @@ impl HnswIndex { } } } + + pub(super) fn distance(&self, a: &[f32], b: &[f32]) -> f32 { + debug_assert_eq!(a.len(), b.len()); + match self.similarity { + Similarity::Euclidean => a + .iter() + .zip(b.iter()) + .map(|(&x, &y)| { + let d = x - y; + d * d + }) + .sum(), + Similarity::Cosine => { + let dot = a.iter().zip(b.iter()).map(|(&x, &y)| x * y).sum::(); + 1.0 - dot + } + Similarity::Manhattan | Similarity::Hamming => distance(a, b, self.similarity), + } + } } diff --git a/crates/index/src/hnsw/mod.rs b/crates/index/src/hnsw/mod.rs index 0cedc17..129985a 100644 --- a/crates/index/src/hnsw/mod.rs +++ b/crates/index/src/hnsw/mod.rs @@ -6,7 +6,7 @@ pub mod search; pub mod serialize; pub mod types; use defs::Magic; -pub use index::HnswIndex; +pub use index::{HnswConfig, HnswIndex}; pub const HNSW_MAGIC_BYTES: Magic = [0x02, 0x01, 0x03, 0x00]; diff --git a/crates/index/src/hnsw/search.rs b/crates/index/src/hnsw/search.rs index bdc6640..8be71db 100644 --- a/crates/index/src/hnsw/search.rs +++ b/crates/index/src/hnsw/search.rs @@ -5,7 +5,6 @@ use std::collections::HashSet; use defs::{OrdF32, PointId}; use crate::Result; -use crate::distance; use super::index::HnswIndex; @@ -23,7 +22,7 @@ impl HnswIndex { let mut current = ep; loop { let cur_vec = self.get_vec(current)?; - let mut best_score = distance(query, cur_vec, self.similarity); + let mut best_score = self.distance(query, cur_vec); let mut best_id = current; let empty: &[PointId] = &[]; @@ -46,7 +45,7 @@ impl HnswIndex { continue; } let n_vec = self.get_vec(n)?; - let score = distance(query, n_vec, self.similarity); + let score = self.distance(query, n_vec); if score < best_score { best_score = score; best_id = n; @@ -91,7 +90,7 @@ impl HnswIndex { .unwrap_or(ep), }; - let ep_score = distance(query, self.get_vec(seed)?, self.similarity); + let ep_score = self.distance(query, self.get_vec(seed)?); candidates.push((Reverse(OrdF32::new(ep_score)), seed)); w_heap.push((OrdF32::new(ep_score), seed)); visited.insert(seed); @@ -114,7 +113,7 @@ impl HnswIndex { .unwrap_or(empty); for &n in neighbors { - if visited.contains(&n) { + if !visited.insert(n) { continue; } // Skip deleted neighbors @@ -124,8 +123,7 @@ impl HnswIndex { continue; } - visited.insert(n); - let score = distance(query, self.get_vec(n)?, self.similarity); + let score = self.distance(query, self.get_vec(n)?); let score = OrdF32::new(score); candidates.push((Reverse(score), n)); if w_heap.len() < ef_construction { @@ -168,7 +166,7 @@ impl HnswIndex { let cand_vec = self.get_vec(cand_id)?; for &r_id in &result { let r_vec = self.get_vec(r_id)?; - let cand_to_r = distance(cand_vec, r_vec, self.similarity); + let cand_to_r = self.distance(cand_vec, r_vec); if cand_to_r < cand_dist_to_q { continue 'outer; } @@ -273,7 +271,7 @@ impl HnswIndex { let mut scored: Vec<(PointId, f32)> = Vec::with_capacity(merged.len()); for nid in merged { - let d = distance(center_vec, self.get_vec(nid)?, self.similarity); + let d = self.distance(center_vec, self.get_vec(nid)?); scored.push((nid, d)); } scored.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap()); diff --git a/crates/index/src/kd_tree/helpers.rs b/crates/index/src/kd_tree/helpers.rs index 8934ae7..e7ffe11 100644 --- a/crates/index/src/kd_tree/helpers.rs +++ b/crates/index/src/kd_tree/helpers.rs @@ -4,16 +4,13 @@ use super::types::KDTreeNode; use defs::IndexedVector; impl KDTree { - pub const BALANCE_THRESHOLD: f32 = 0.7; - pub const DELETE_REBUILD_RATIO: f32 = 0.25; - /// Checks if a node is unbalanced based on the balance threshold - pub fn is_unbalanced(node: &KDTreeNode) -> bool { + pub fn is_unbalanced(&self, node: &KDTreeNode) -> bool { let left_size = node.left.as_ref().map_or(0, |n| n.subtree_size); let right_size = node.right.as_ref().map_or(0, |n| n.subtree_size); let max_child = left_size.max(right_size); - max_child as f32 > Self::BALANCE_THRESHOLD * node.subtree_size as f32 + max_child as f32 > self.config.balance_threshold * node.subtree_size as f32 } /// Recursively collects non-deleted vectors from the tree @@ -37,7 +34,9 @@ impl KDTree { } /// Checks if the tree should be globally rebuilt based on deletion ratio - pub fn should_rebuild_global(total_nodes: usize, deleted_count: usize) -> bool { - total_nodes > 0 && (deleted_count as f32 / total_nodes as f32) > Self::DELETE_REBUILD_RATIO + pub fn should_rebuild_global(&self) -> bool { + self.total_nodes > 0 + && (self.deleted_count as f32 / self.total_nodes as f32) + > self.config.delete_rebuild_ratio } } diff --git a/crates/index/src/kd_tree/index.rs b/crates/index/src/kd_tree/index.rs index bf4ed7a..31b6c6e 100644 --- a/crates/index/src/kd_tree/index.rs +++ b/crates/index/src/kd_tree/index.rs @@ -9,6 +9,21 @@ use std::{ }; use uuid::Uuid; +#[derive(Debug, Clone, Copy)] +pub struct KDTreeConfig { + pub balance_threshold: f32, + pub delete_rebuild_ratio: f32, +} + +impl Default for KDTreeConfig { + fn default() -> Self { + Self { + balance_threshold: 0.7, + delete_rebuild_ratio: 0.25, + } + } +} + pub struct KDTree { pub dim: usize, pub root: Option>, @@ -17,22 +32,35 @@ pub struct KDTree { // Rebuild tracking pub total_nodes: usize, pub deleted_count: usize, + pub config: KDTreeConfig, } impl KDTree { // Build an empty index with no points pub fn build_empty(dim: usize) -> Self { + Self::build_empty_with_config(dim, KDTreeConfig::default()) + } + + pub fn build_empty_with_config(dim: usize, config: KDTreeConfig) -> Self { KDTree { dim, root: None, point_ids: HashSet::new(), total_nodes: 0, deleted_count: 0, + config: Self::sanitize_config(config), } } // Builds the vector index from provided vectors, there should atleast be single vector for dim calculation - pub fn build(mut vectors: Vec) -> Result { + pub fn build(vectors: Vec) -> Result { + Self::build_with_config(vectors, KDTreeConfig::default()) + } + + pub fn build_with_config( + mut vectors: Vec, + config: KDTreeConfig, + ) -> Result { if vectors.is_empty() { Err(IndexError::NotInitialized) } else { @@ -50,10 +78,18 @@ impl KDTree { point_ids, total_nodes: vectors.len(), deleted_count: 0, + config: Self::sanitize_config(config), }) } } + fn sanitize_config(config: KDTreeConfig) -> KDTreeConfig { + KDTreeConfig { + balance_threshold: config.balance_threshold.clamp(0.5, 1.0), + delete_rebuild_ratio: config.delete_rebuild_ratio.clamp(0.0, 1.0), + } + } + // Builds the tree recursively with given vectors and returns the pointer of the root node pub fn build_recursive( vectors: &mut [IndexedVector], @@ -246,7 +282,7 @@ impl KDTree { // Check root first (depth 0) if let Some(node) = current - && Self::is_unbalanced(node) + && self.is_unbalanced(node) { unbalanced_depth = Some(0); } @@ -267,7 +303,7 @@ impl KDTree { // Check the child node we just moved to (at depth idx + 1) if let Some(child) = current - && Self::is_unbalanced(child) + && self.is_unbalanced(child) { unbalanced_depth = Some(idx + 1); break; @@ -289,7 +325,7 @@ impl KDTree { self.point_ids.remove(point_id); } - if Self::should_rebuild_global(self.total_nodes, self.deleted_count) + if self.should_rebuild_global() && let Some(root) = self.root.take() { let mut vectors = Self::collect_active_vectors(*root); diff --git a/crates/index/src/kd_tree/mod.rs b/crates/index/src/kd_tree/mod.rs index 4a29b61..61233f4 100644 --- a/crates/index/src/kd_tree/mod.rs +++ b/crates/index/src/kd_tree/mod.rs @@ -9,3 +9,5 @@ pub mod types; mod tests; pub const KD_TREE_MAGIC_BYTES: Magic = [0x00, 0x01, 0x02, 0x00]; + +pub use index::{KDTree, KDTreeConfig}; diff --git a/crates/index/src/kd_tree/serialize.rs b/crates/index/src/kd_tree/serialize.rs index a651d10..397d629 100644 --- a/crates/index/src/kd_tree/serialize.rs +++ b/crates/index/src/kd_tree/serialize.rs @@ -101,6 +101,7 @@ impl KDTree { point_ids: non_deleted, total_nodes: metadata.total_nodes, deleted_count: metadata.deleted_count, + config: Default::default(), }) } } diff --git a/crates/index/src/lib.rs b/crates/index/src/lib.rs index 13202d8..8b80cd9 100644 --- a/crates/index/src/lib.rs +++ b/crates/index/src/lib.rs @@ -21,6 +21,23 @@ pub trait VectorIndex: Send + Sync + SerializableIndex { similarity: Similarity, k: usize, ) -> Result>; // Return a Vec of ids of closest vectors (length max k) + + fn search_with_ef( + &self, + query_vector: DenseVector, + similarity: Similarity, + k: usize, + _ef: Option, + ) -> Result> { + self.search(query_vector, similarity, k) + } + + fn insert_batch(&mut self, vectors: Vec) -> Result<()> { + for v in vectors { + self.insert(v)?; + } + Ok(()) + } } /// Distance function to get the distance between two vectors (taken from old version) diff --git a/crates/server/src/config.rs b/crates/server/src/config.rs index e593afa..dfcd0f4 100644 --- a/crates/server/src/config.rs +++ b/crates/server/src/config.rs @@ -1,7 +1,7 @@ use api::DbConfig; use defs::Similarity; use dotenv::dotenv; -use index::IndexType; +use index::{IndexType, hnsw::HnswConfig, kd_tree::KDTreeConfig}; use snafu::prelude::*; use std::env; use std::fs; @@ -12,6 +12,12 @@ use tracing::{Level, event}; const DEFAULT_HTTP_PORT: &str = "3000"; const DEFAULT_GRPC_PORT: &str = "50051"; +const DEFAULT_HNSW_M: usize = 16; +const DEFAULT_HNSW_MAX_LAYER: usize = 16; +const DEFAULT_HNSW_EF_CONSTRUCTION: usize = 200; +const DEFAULT_HNSW_EF: usize = 100; +const DEFAULT_KD_TREE_BALANCE_THRESHOLD: f32 = 0.7; +const DEFAULT_KD_TREE_DELETE_REBUILD_RATIO: f32 = 0.25; #[derive(Debug)] pub struct ServerConfig { @@ -197,12 +203,33 @@ impl ServerConfig { } }; + let hnsw_m = load_usize_env("HNSW_M", DEFAULT_HNSW_M); + let hnsw_config = HnswConfig { + max_connections: hnsw_m, + max_connections_0: load_usize_env("HNSW_M0", 2 * hnsw_m), + max_layer: load_usize_env("HNSW_MAX_LAYER", DEFAULT_HNSW_MAX_LAYER), + ef_construction: load_usize_env("HNSW_EF_CONSTRUCTION", DEFAULT_HNSW_EF_CONSTRUCTION), + ef: load_usize_env("HNSW_EF", DEFAULT_HNSW_EF), + }; + let kd_tree_config = KDTreeConfig { + balance_threshold: load_f32_env( + "KD_TREE_BALANCE_THRESHOLD", + DEFAULT_KD_TREE_BALANCE_THRESHOLD, + ), + delete_rebuild_ratio: load_f32_env( + "KD_TREE_DELETE_REBUILD_RATIO", + DEFAULT_KD_TREE_DELETE_REBUILD_RATIO, + ), + }; + let db_config = DbConfig { storage_type, index_type, data_path, dimension, similarity, + hnsw_config, + kd_tree_config, }; Ok(ServerConfig { @@ -215,3 +242,35 @@ impl ServerConfig { }) } } + +fn load_usize_env(name: &str, default: usize) -> usize { + match env::var(name) { + Ok(value) => value.parse().unwrap_or_else(|_| { + event!( + Level::WARN, + "{}='{}' is invalid, defaulting to {}", + name, + value, + default + ); + default + }), + Err(_) => default, + } +} + +fn load_f32_env(name: &str, default: f32) -> f32 { + match env::var(name) { + Ok(value) => value.parse().unwrap_or_else(|_| { + event!( + Level::WARN, + "{}='{}' is invalid, defaulting to {}", + name, + value, + default + ); + default + }), + Err(_) => default, + } +} diff --git a/crates/snapshot/src/lib.rs b/crates/snapshot/src/lib.rs index 51b26fc..7324a75 100644 --- a/crates/snapshot/src/lib.rs +++ b/crates/snapshot/src/lib.rs @@ -26,7 +26,8 @@ use std::{ time::SystemTime, }; use storage::{ - StorageEngine, StorageType, checkpoint::StorageCheckpoint, rocks_db::RocksDbStorage, + StorageEngine, StorageType, checkpoint::StorageCheckpoint, in_memory::MemoryStorage, + rocks_db::RocksDbStorage, }; use tar::Archive; use tempfile::tempdir; @@ -201,17 +202,12 @@ impl Snapshot { )); } - // only rocksdb is supported for snapshots as of now let mut storage_engine: Box = match manifest.storage_type { + StorageType::InMemory => Box::new(MemoryStorage::new()), StorageType::RocksDb => Box::new( RocksDbStorage::new(storage_data_path) .map_err(|e| DbError::StorageError(format!("Could not open storage: {e}")))?, ), - _ => { - return Err(DbError::SnapshotError( - "Unsupported storage type".to_string(), - )); - } }; let id = manifest.id; diff --git a/crates/storage/src/checkpoint.rs b/crates/storage/src/checkpoint.rs index 09827fd..fc4be08 100644 --- a/crates/storage/src/checkpoint.rs +++ b/crates/storage/src/checkpoint.rs @@ -35,6 +35,7 @@ impl StorageCheckpoint { .0; let storage_type = match marker { + INMEMORY_CHECKPOINT_FILENAME_MARKER => StorageType::InMemory, ROCKSDB_CHECKPOINT_FILENAME_MARKER => StorageType::RocksDb, _ => { return Err(DbError::StorageCheckpointError( diff --git a/crates/storage/src/error.rs b/crates/storage/src/error.rs index d54f67e..70971a7 100644 --- a/crates/storage/src/error.rs +++ b/crates/storage/src/error.rs @@ -37,6 +37,15 @@ pub enum StorageError { #[snafu(display("Failed to iterate over storage: {source}"))] RocksDbIteration { source: rocksdb::Error }, + #[snafu(display("Failed to lock in-memory storage"))] + InMemoryLock {}, + + #[snafu(display("In-memory checkpoint error: {}", msg))] + InMemoryCheckpoint { msg: String }, + + #[snafu(display("{} : {}", msg, source))] + InMemoryCheckpointIo { msg: String, source: std::io::Error }, + #[snafu(display("Failed to serialize point {id}: {source}"))] Serialization { id: PointId, source: bincode::Error }, diff --git a/crates/storage/src/in_memory.rs b/crates/storage/src/in_memory.rs index 096d9ed..1e8ca60 100644 --- a/crates/storage/src/in_memory.rs +++ b/crates/storage/src/in_memory.rs @@ -1,18 +1,29 @@ use crate::StorageType; use crate::error::StorageError; use crate::{StorageEngine, VectorPage, checkpoint::StorageCheckpoint}; -use defs::{DenseVector, Payload, PointId}; -use std::path::{Path, PathBuf}; +use bincode::{deserialize_from, serialize_into}; +use defs::{DenseVector, Payload, Point, PointId}; +use std::collections::BTreeMap; +use std::fs::File; +use std::io::{BufReader, BufWriter, Read, Write}; +use std::ops::Bound::{Excluded, Unbounded}; +use std::path::Path; +use std::sync::RwLock; pub const INMEMORY_CHECKPOINT_FILENAME_MARKER: &str = "inmemory"; +const INMEMORY_CHECKPOINT_EXTENSION: &str = "bin"; +const INMEMORY_CHECKPOINT_MAGIC: &[u8; 8] = b"VDBIMCP\0"; +const INMEMORY_CHECKPOINT_VERSION: u16 = 1; pub struct MemoryStorage { - // define here how MemoryStorage will be defined + points: RwLock>, } impl MemoryStorage { pub fn new() -> Self { - MemoryStorage {} + MemoryStorage { + points: RwLock::new(BTreeMap::new()), + } } } @@ -25,40 +36,401 @@ impl Default for MemoryStorage { impl StorageEngine for MemoryStorage { fn insert_point( &self, - _id: PointId, - _vector: Option, - _payload: Option, + id: PointId, + vector: Option, + payload: Option, ) -> Result<(), StorageError> { + let mut points = self + .points + .write() + .map_err(|_| StorageError::InMemoryLock {})?; + points.insert( + id, + Point { + id, + vector, + payload, + }, + ); Ok(()) } - fn contains_point(&self, _id: PointId) -> Result { - Ok(true) + fn contains_point(&self, id: PointId) -> Result { + let points = self + .points + .read() + .map_err(|_| StorageError::InMemoryLock {})?; + Ok(points.contains_key(&id)) } - fn delete_point(&self, _id: PointId) -> Result<(), StorageError> { + fn delete_point(&self, id: PointId) -> Result<(), StorageError> { + let mut points = self + .points + .write() + .map_err(|_| StorageError::InMemoryLock {})?; + points.remove(&id); Ok(()) } - fn get_payload(&self, _id: PointId) -> Result, StorageError> { - Ok(None) + fn get_payload(&self, id: PointId) -> Result, StorageError> { + let points = self + .points + .read() + .map_err(|_| StorageError::InMemoryLock {})?; + Ok(points.get(&id).and_then(|point| point.payload.clone())) } - fn get_vector(&self, _id: PointId) -> Result, StorageError> { - Ok(None) + fn get_vector(&self, id: PointId) -> Result, StorageError> { + let points = self + .points + .read() + .map_err(|_| StorageError::InMemoryLock {})?; + Ok(points.get(&id).and_then(|point| point.vector.clone())) } fn list_vectors( &self, - _offset: PointId, - _limit: usize, + offset: PointId, + limit: usize, ) -> Result, StorageError> { - Ok(None) + if limit < 1 { + return Ok(None); + } + + let points = self + .points + .read() + .map_err(|_| StorageError::InMemoryLock {})?; + let mut result = Vec::with_capacity(limit); + let mut last_id = offset; + + for (id, point) in points.range((Excluded(offset), Unbounded)) { + if let Some(vector) = &point.vector { + last_id = *id; + result.push((*id, vector.clone())); + if result.len() == limit { + break; + } + } + } + + Ok(Some((result, last_id))) } - fn checkpoint_at(&self, _path: &Path) -> Result { + fn checkpoint_at(&self, path: &Path) -> Result { + let checkpoint_filename = format!( + "{}-{}.{}", + INMEMORY_CHECKPOINT_FILENAME_MARKER, + uuid::Uuid::new_v4(), + INMEMORY_CHECKPOINT_EXTENSION + ); + let checkpoint_path = path.join(checkpoint_filename); + let file = File::create(&checkpoint_path).map_err(|source| { + StorageError::InMemoryCheckpointIo { + msg: "Couldn't create in-memory checkpoint".to_string(), + source, + } + })?; + let point_snapshot: Vec = { + let points = self + .points + .read() + .map_err(|_| StorageError::InMemoryLock {})?; + points.values().cloned().collect() + }; + + let mut writer = BufWriter::new(file); + writer + .write_all(INMEMORY_CHECKPOINT_MAGIC) + .and_then(|_| writer.write_all(&INMEMORY_CHECKPOINT_VERSION.to_le_bytes())) + .and_then(|_| writer.write_all(&(point_snapshot.len() as u64).to_le_bytes())) + .map_err(|source| StorageError::InMemoryCheckpointIo { + msg: "Couldn't write in-memory checkpoint header".to_string(), + source, + })?; + + for point in &point_snapshot { + serialize_into(&mut writer, point).map_err(|source| StorageError::Serialization { + id: point.id, + source, + })?; + } + + writer + .flush() + .map_err(|source| StorageError::InMemoryCheckpointIo { + msg: "Couldn't flush in-memory checkpoint".to_string(), + source, + })?; + Ok(StorageCheckpoint { - path: PathBuf::default(), + path: checkpoint_path, storage_type: StorageType::InMemory, }) } - fn restore_checkpoint(&mut self, _checkpoint: &StorageCheckpoint) -> Result<(), StorageError> { + fn restore_checkpoint(&mut self, checkpoint: &StorageCheckpoint) -> Result<(), StorageError> { + if checkpoint.storage_type != StorageType::InMemory { + return Err(StorageError::InMemoryCheckpoint { + msg: "Invalid storage type".to_string(), + }); + } + + let checkpoint_filename = checkpoint + .path + .file_name() + .ok_or_else(|| StorageError::InMemoryCheckpoint { + msg: "Could not read checkpoint filename".to_string(), + })? + .to_str() + .ok_or_else(|| StorageError::InMemoryCheckpoint { + msg: "Checkpoint filename is not valid UTF-8".to_string(), + })?; + if !checkpoint_filename.starts_with(INMEMORY_CHECKPOINT_FILENAME_MARKER) + || checkpoint.path.extension().and_then(|ext| ext.to_str()) + != Some(INMEMORY_CHECKPOINT_EXTENSION) + { + return Err(StorageError::InMemoryCheckpoint { + msg: "Invalid file name".to_string(), + }); + } + + let file = + File::open(&checkpoint.path).map_err(|source| StorageError::InMemoryCheckpointIo { + msg: "Couldn't open in-memory checkpoint".to_string(), + source, + })?; + let mut reader = BufReader::new(file); + let mut magic = [0u8; INMEMORY_CHECKPOINT_MAGIC.len()]; + reader + .read_exact(&mut magic) + .map_err(|source| StorageError::InMemoryCheckpointIo { + msg: "Couldn't read in-memory checkpoint magic".to_string(), + source, + })?; + if &magic != INMEMORY_CHECKPOINT_MAGIC { + return Err(StorageError::InMemoryCheckpoint { + msg: "Invalid checkpoint magic".to_string(), + }); + } + + let mut version_bytes = [0u8; size_of::()]; + reader.read_exact(&mut version_bytes).map_err(|source| { + StorageError::InMemoryCheckpointIo { + msg: "Couldn't read in-memory checkpoint version".to_string(), + source, + } + })?; + let version = u16::from_le_bytes(version_bytes); + if version != INMEMORY_CHECKPOINT_VERSION { + return Err(StorageError::InMemoryCheckpoint { + msg: format!("Unsupported checkpoint version: {version}"), + }); + } + + let mut count_bytes = [0u8; size_of::()]; + reader.read_exact(&mut count_bytes).map_err(|source| { + StorageError::InMemoryCheckpointIo { + msg: "Couldn't read in-memory checkpoint point count".to_string(), + source, + } + })?; + let point_count = u64::from_le_bytes(count_bytes); + let mut restored_points = BTreeMap::new(); + for _ in 0..point_count { + let point: Point = + deserialize_from(&mut reader).map_err(|source| StorageError::Deserialization { + id: PointId::nil(), + source, + })?; + restored_points.insert(point.id, point); + } + + let mut points = self + .points + .write() + .map_err(|_| StorageError::InMemoryLock {})?; + *points = restored_points; Ok(()) } } + +#[cfg(test)] +mod tests { + use super::*; + use defs::ContentType; + use std::io::{Read, Write}; + use tempfile::{TempDir, tempdir}; + use uuid::Uuid; + + fn create_test_storage() -> MemoryStorage { + MemoryStorage::new() + } + + fn test_payload(content: &str) -> Payload { + Payload { + content_type: ContentType::Text, + content: content.to_string(), + } + } + + #[test] + fn test_insert_and_get_vector() { + let storage = create_test_storage(); + let id = Uuid::new_v4(); + let vector = Some(vec![0.1, 0.2, 0.3]); + let payload = Some(test_payload("Test")); + + storage.insert_point(id, vector.clone(), payload).unwrap(); + + assert_eq!(storage.get_vector(id).unwrap(), vector); + } + + #[test] + fn test_insert_and_get_payload() { + let storage = create_test_storage(); + let id = Uuid::new_v4(); + let payload = Some(test_payload("Test")); + + storage.insert_point(id, None, payload.clone()).unwrap(); + + assert_eq!(storage.get_payload(id).unwrap(), payload); + } + + #[test] + fn test_contains_and_delete_point() { + let storage = create_test_storage(); + let id = Uuid::new_v4(); + + assert!(!storage.contains_point(id).unwrap()); + + storage + .insert_point(id, Some(vec![0.4, 0.5, 0.6]), Some(test_payload("Test"))) + .unwrap(); + assert!(storage.contains_point(id).unwrap()); + + storage.delete_point(id).unwrap(); + assert!(!storage.contains_point(id).unwrap()); + assert_eq!(storage.get_vector(id).unwrap(), None); + assert_eq!(storage.get_payload(id).unwrap(), None); + } + + #[test] + fn test_list_vectors_respects_offset_limit_and_skips_payload_only_points() { + let storage = create_test_storage(); + let ids = [ + Uuid::from_u128(1), + Uuid::from_u128(2), + Uuid::from_u128(3), + Uuid::from_u128(4), + ]; + + storage + .insert_point(ids[0], Some(vec![1.0, 1.1]), Some(test_payload("one"))) + .unwrap(); + storage + .insert_point(ids[1], None, Some(test_payload("payload-only"))) + .unwrap(); + storage + .insert_point(ids[2], Some(vec![3.0, 3.1]), Some(test_payload("three"))) + .unwrap(); + storage + .insert_point(ids[3], Some(vec![4.0, 4.1]), Some(test_payload("four"))) + .unwrap(); + + let (first_page, next_offset) = storage.list_vectors(Uuid::nil(), 2).unwrap().unwrap(); + assert_eq!( + first_page, + vec![(ids[0], vec![1.0, 1.1]), (ids[2], vec![3.0, 3.1])] + ); + assert_eq!(next_offset, ids[2]); + + let (second_page, next_offset) = storage.list_vectors(next_offset, 2).unwrap().unwrap(); + assert_eq!(second_page, vec![(ids[3], vec![4.0, 4.1])]); + assert_eq!(next_offset, ids[3]); + } + + #[test] + fn test_list_vectors_with_zero_limit_returns_none() { + let storage = create_test_storage(); + + assert_eq!(storage.list_vectors(Uuid::nil(), 0).unwrap(), None); + } + + #[test] + fn test_create_and_restore_checkpoint() { + let mut storage = create_test_storage(); + let temp_dir: TempDir = tempdir().unwrap(); + let id_before_checkpoint = Uuid::new_v4(); + let id_after_checkpoint = Uuid::new_v4(); + + storage + .insert_point( + id_before_checkpoint, + Some(vec![0.1, 0.2, 0.3]), + Some(test_payload("before")), + ) + .unwrap(); + let checkpoint = storage.checkpoint_at(temp_dir.path()).unwrap(); + + storage + .insert_point( + id_after_checkpoint, + Some(vec![0.4, 0.5, 0.6]), + Some(test_payload("after")), + ) + .unwrap(); + + storage.restore_checkpoint(&checkpoint).unwrap(); + + assert!(storage.contains_point(id_before_checkpoint).unwrap()); + assert!(!storage.contains_point(id_after_checkpoint).unwrap()); + assert_eq!( + storage.get_payload(id_before_checkpoint).unwrap(), + Some(test_payload("before")) + ); + } + + #[test] + fn test_checkpoint_writes_header() { + let storage = create_test_storage(); + let temp_dir = tempdir().unwrap(); + let id = Uuid::new_v4(); + + storage + .insert_point(id, Some(vec![0.1, 0.2, 0.3]), Some(test_payload("point"))) + .unwrap(); + let checkpoint = storage.checkpoint_at(temp_dir.path()).unwrap(); + + let mut file = File::open(checkpoint.path).unwrap(); + let mut magic = [0u8; INMEMORY_CHECKPOINT_MAGIC.len()]; + file.read_exact(&mut magic).unwrap(); + assert_eq!(&magic, INMEMORY_CHECKPOINT_MAGIC); + + let mut version_bytes = [0u8; size_of::()]; + file.read_exact(&mut version_bytes).unwrap(); + assert_eq!( + u16::from_le_bytes(version_bytes), + INMEMORY_CHECKPOINT_VERSION + ); + + let mut count_bytes = [0u8; size_of::()]; + file.read_exact(&mut count_bytes).unwrap(); + assert_eq!(u64::from_le_bytes(count_bytes), 1); + } + + #[test] + fn test_restore_rejects_invalid_checkpoint_magic() { + let mut storage = create_test_storage(); + let temp_dir = tempdir().unwrap(); + let checkpoint_path = temp_dir.path().join("inmemory-invalid.bin"); + let mut file = File::create(&checkpoint_path).unwrap(); + file.write_all(b"BADMAGIC").unwrap(); + file.write_all(&INMEMORY_CHECKPOINT_VERSION.to_le_bytes()) + .unwrap(); + file.write_all(&0u64.to_le_bytes()).unwrap(); + + let checkpoint = StorageCheckpoint { + path: checkpoint_path, + storage_type: StorageType::InMemory, + }; + + let error = storage.restore_checkpoint(&checkpoint).unwrap_err(); + assert!(matches!(error, StorageError::InMemoryCheckpoint { .. })); + } +} diff --git a/crates/storage/src/rocks_db.rs b/crates/storage/src/rocks_db.rs index 6d7e4db..166aefc 100644 --- a/crates/storage/src/rocks_db.rs +++ b/crates/storage/src/rocks_db.rs @@ -181,13 +181,15 @@ impl StorageEngine for RocksDbStorage { let point: Point = deserialize(&v).context(error::DeserializationSnafu { id: offset })?; - if point.id <= offset { + let id = point.id; + + if id <= offset { continue; } if let Some(vec) = point.vector { - last_id = point.id; - result.push((point.id, vec)); + last_id = id; + result.push((id, vec)); if result.len() == limit { break; } diff --git a/crates/tui/README.md b/crates/tui/README.md new file mode 100644 index 0000000..39ac88b --- /dev/null +++ b/crates/tui/README.md @@ -0,0 +1,35 @@ +# VortexDB TUI + +A terminal user interface for managing VortexDB databases locally. + +## Important Note + +> **This is currently a local admin/development tool**, not a network client. The TUI operates directly on database files using embedded storage — it does not connect to a running VortexDB server. This will change once the server gets the multiple-database support. +> +> For remote server access, use the [Python client](../../client/python/). + +## Usage + +```bash +cargo run -p tui +``` + +Databases are stored in `./databases/` by default. + +### Environment Variables + +For embedding features, set these in your `.env` file or environment: + +```bash +TEXT_EMBEDDING_URL=http://localhost:8000/embed/text +IMAGE_EMBEDDING_URL=http://localhost:8000/embed/image +``` + +These point to an external embedding service that generates vectors from text/images. + +## Roadmap + +This TUI will evolve into a full client that connects to VortexDB servers over gRPC. Planned changes: + +- Add remote connection mode via gRPC (like the Python client) +- Move to `client/tui/` once server multi-database support lands diff --git a/crates/tui/src/app/database.rs b/crates/tui/src/app/database.rs index ee68e4d..db67df7 100644 --- a/crates/tui/src/app/database.rs +++ b/crates/tui/src/app/database.rs @@ -1,6 +1,6 @@ use api::{DbConfig, VectorDb, init_api}; use defs::Similarity; -use index::IndexType; +use index::{IndexType, hnsw::HnswConfig, kd_tree::KDTreeConfig}; use std::io; use std::path::PathBuf; use std::sync::Arc; @@ -49,6 +49,8 @@ impl DatabaseManager { data_path: path.clone(), dimension: 512, similarity: Similarity::Cosine, + hnsw_config: HnswConfig::default(), + kd_tree_config: KDTreeConfig::default(), }; match init_api(cfg) { @@ -82,6 +84,8 @@ impl DatabaseManager { data_path: path.clone(), dimension: 512, similarity: Similarity::Cosine, + hnsw_config: HnswConfig::default(), + kd_tree_config: KDTreeConfig::default(), }; match init_api(cfg) { diff --git a/crates/tui/src/app/events.rs b/crates/tui/src/app/events.rs index b2736cb..5f4668f 100644 --- a/crates/tui/src/app/events.rs +++ b/crates/tui/src/app/events.rs @@ -1,6 +1,6 @@ use super::{App, AppState, ModalType, VectorListItem}; use crossterm::event::{Event, KeyCode, KeyEvent}; -use defs::{ContentType, Payload, Similarity}; +use defs::{ContentType, Payload, SearchQueryInput, Similarity}; use std::io; use std::path::PathBuf; use uuid::Uuid; @@ -313,7 +313,14 @@ fn execute_modal_action(app: &mut App) -> io::Result<()> { } }; - let ids = db.search(query, Similarity::Cosine, k).map_err(to_io)?; + let ids = db + .search(SearchQueryInput { + vector: query, + similarity: Similarity::Cosine, + limit: k, + ef: None, + }) + .map_err(to_io)?; app.vector_list_items.clear(); app.vector_detail = None; diff --git a/crates/tui/src/ui/components.rs b/crates/tui/src/ui/components.rs index 4fffc6f..3cd5e7c 100644 --- a/crates/tui/src/ui/components.rs +++ b/crates/tui/src/ui/components.rs @@ -17,7 +17,7 @@ impl PageTitle { .block( Block::default() .borders(Borders::ALL) - .title("Vector DB") + .title("VortexDB") .title_alignment(Alignment::Center) .border_style(Style::default().fg(self.color)), ) diff --git a/crates/tui/src/ui/dashboard.rs b/crates/tui/src/ui/dashboard.rs index 71956bc..fc6a705 100644 --- a/crates/tui/src/ui/dashboard.rs +++ b/crates/tui/src/ui/dashboard.rs @@ -21,74 +21,37 @@ pub fn render_dashboard(f: &mut Frame, app: &App) { Line::from(""), Line::from(""), Line::from(vec![Span::styled( - "██╗ ██╗███████╗ ██████╗████████╗ ██████╗ ██████╗ ", + "██╗ ██╗ ██████╗ ██████╗ ████████╗███████╗██╗ ██╗██████╗ ██████╗ ", Style::default() .fg(Color::Cyan) .add_modifier(ratatui::style::Modifier::BOLD), )]), Line::from(vec![Span::styled( - "██║ ██║██╔════╝██╔════╝╚══██╔══╝██╔═══██╗██╔══██╗", + "██║ ██║██╔═══██╗██╔══██╗╚══██╔══╝██╔════╝╚██╗██╔╝██╔══██╗██╔══██╗", Style::default() .fg(Color::Cyan) .add_modifier(ratatui::style::Modifier::BOLD), )]), Line::from(vec![Span::styled( - "██║ ██║█████╗ ██║ ██║ ██║ ██║██████╔╝", + "██║ ██║██║ ██║██████╔╝ ██║ █████╗ ╚███╔╝ ██║ ██║██████╔╝", Style::default() .fg(Color::Cyan) .add_modifier(ratatui::style::Modifier::BOLD), )]), Line::from(vec![Span::styled( - "╚██╗ ██╔╝██╔══╝ ██║ ██║ ██║ ██║██╔══██╗", + "╚██╗ ██╔╝██║ ██║██╔══██╗ ██║ ██╔══╝ ██╔██╗ ██║ ██║██╔══██╗", Style::default() .fg(Color::Cyan) .add_modifier(ratatui::style::Modifier::BOLD), )]), Line::from(vec![Span::styled( - " ╚████╔╝ ███████╗╚██████╗ ██║ ╚██████╔╝██║ ██║", + " ╚████╔╝ ╚██████╔╝██║ ██║ ██║ ███████╗██╔╝ ██╗██████╔╝██████╔╝", Style::default() .fg(Color::Cyan) .add_modifier(ratatui::style::Modifier::BOLD), )]), Line::from(vec![Span::styled( - " ╚═══╝ ╚══════╝ ╚═════╝ ╚═╝ ╚═════╝ ╚═╝ ╚═╝", - Style::default() - .fg(Color::Cyan) - .add_modifier(ratatui::style::Modifier::BOLD), - )]), - Line::from(""), - Line::from(vec![Span::styled( - "██████╗ ██████╗ ", - Style::default() - .fg(Color::Cyan) - .add_modifier(ratatui::style::Modifier::BOLD), - )]), - Line::from(vec![Span::styled( - "██╔══██╗██╔══██╗", - Style::default() - .fg(Color::Cyan) - .add_modifier(ratatui::style::Modifier::BOLD), - )]), - Line::from(vec![Span::styled( - "██║ ██║██████╔╝", - Style::default() - .fg(Color::Cyan) - .add_modifier(ratatui::style::Modifier::BOLD), - )]), - Line::from(vec![Span::styled( - "██║ ██║██╔══██╗", - Style::default() - .fg(Color::Cyan) - .add_modifier(ratatui::style::Modifier::BOLD), - )]), - Line::from(vec![Span::styled( - "██████╔╝██████╔╝", - Style::default() - .fg(Color::Cyan) - .add_modifier(ratatui::style::Modifier::BOLD), - )]), - Line::from(vec![Span::styled( - "╚═════╝ ╚═════╝ ", + " ╚═══╝ ╚═════╝ ╚═╝ ╚═╝ ╚═╝ ╚══════╝╚═╝ ╚═╝╚═════╝ ╚═════╝ ", Style::default() .fg(Color::Cyan) .add_modifier(ratatui::style::Modifier::BOLD),