diff --git a/README.md b/README.md index 0ffa98f0e..1a9de9593 100644 --- a/README.md +++ b/README.md @@ -75,6 +75,7 @@ All the database client supported | pgvecto.rs | `pip install vectordb-bench[pgvecto_rs]` | | redis | `pip install vectordb-bench[redis]` | | memorydb | `pip install vectordb-bench[memorydb]` | +| kividb | `pip install vectordb-bench[kividb]` | | chromadb | `pip install vectordb-bench[chromadb]` | | cockroachdb | `pip install vectordb-bench[cockroachdb]` | | awsopensearch | `pip install vectordb-bench[opensearch]` | diff --git a/pyproject.toml b/pyproject.toml index 9fbbb7748..5d4d3da34 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -68,6 +68,7 @@ pgvector = [ "psycopg", "psycopg-binary", "pgvector" ] pgvecto_rs = [ "pgvecto_rs[psycopg3]>=0.2.2" ] redis = [ "redis" ] memorydb = [ "memorydb" ] +kividb = [ "redis" ] chromadb = [ "chromadb" ] opensearch = [ "opensearch-py", "boto3", "requests-aws4auth" ] aliyun_opensearch = [ "alibabacloud_ha3engine_vector" ] diff --git a/tests/test_kividb.py b/tests/test_kividb.py new file mode 100644 index 000000000..1658ca7ad --- /dev/null +++ b/tests/test_kividb.py @@ -0,0 +1,156 @@ +"""Tests for the KiviDB client. + +The unit tests need nothing. The `integration` tests need a KiviDB server +(>= 1.0.5) and are skipped when none is reachable: + + docker run -d -p 6380:6380 quay.io/kividbio/kividb:v1.0.5-full + KIVIDB_HOST=localhost KIVIDB_PORT=6380 pytest tests/test_kividb.py -v + +Recall is checked against brute-force ground truth over the exact vectors +inserted, so a filter that is dropped or mis-built fails here rather than +producing a plausible-looking number. +""" + +import os +from collections.abc import Callable + +import numpy as np +import pytest +import redis + +from vectordb_bench.backend.clients import DB, IndexType, MetricType +from vectordb_bench.backend.clients.kividb.config import KiviDBFLATConfig, KiviDBHNSWConfig +from vectordb_bench.backend.clients.kividb.kividb import escape_tag_value +from vectordb_bench.backend.filter import FilterOp, IntFilter, LabelFilter, non_filter + +HOST = os.environ.get("KIVIDB_HOST", "localhost") +PORT = int(os.environ.get("KIVIDB_PORT", "6380")) +DIM, COUNT, NQ, K = 32, 2000, 20, 10 +LABELS = ["label_1p", "label_5p", "label_50p"] + + +# ---------------------------------------------------------------- unit tests +def test_registered_with_the_benchmark(): + assert DB.KiviDB.value == "KiviDB" + assert DB.KiviDB.init_cls.name == "KiviDB" + assert DB.KiviDB.config_cls.__name__ == "KiviDBConfig" + assert DB.KiviDB.case_config_cls(IndexType.HNSW) is KiviDBHNSWConfig + assert DB.KiviDB.case_config_cls(IndexType.Flat) is KiviDBFLATConfig + + +def test_supported_filters(): + cls = DB.KiviDB.init_cls + assert set(cls.supported_filter_types) == {FilterOp.NonFilter, FilterOp.NumGE, FilterOp.StrEqual} + + +def test_metric_mapping(): + assert KiviDBHNSWConfig(metric_type=MetricType.L2).index_param()["metric"] == "L2" + assert KiviDBHNSWConfig(metric_type=MetricType.IP).index_param()["metric"] == "IP" + assert KiviDBHNSWConfig(metric_type=MetricType.COSINE).index_param()["metric"] == "COSINE" + + +def test_tag_values_are_escaped(): + assert escape_tag_value("label_5p") == "label_5p" + assert escape_tag_value("a-b c|d") == "a\\-b\\ c\\|d" + + +# ---------------------------------------------------------- integration tests +def _server_available() -> bool: + try: + return "kividb_version" in redis.Redis(host=HOST, port=PORT, socket_timeout=2).info("server") + except Exception: + return False + + +live = pytest.mark.skipif(not _server_available(), reason=f"no KiviDB server at {HOST}:{PORT}") + + +def _data(): + rng = np.random.default_rng(7) + vectors = rng.random((COUNT, DIM), dtype=np.float32) + queries = rng.random((NQ, DIM), dtype=np.float32) + labels = [LABELS[i % len(LABELS)] for i in range(COUNT)] + return vectors, queries, labels + + +def _brute_force(vectors: np.ndarray, query: np.ndarray, candidates: list[int], k: int) -> list[int]: + v = vectors[candidates] + sims = (v @ query) / (np.linalg.norm(v, axis=1) * np.linalg.norm(query)) + return [candidates[i] for i in np.argsort(-sims)[:k]] + + +def _client(with_scalar_labels: bool = False): + cfg = KiviDBHNSWConfig(metric_type=MetricType.COSINE, M=16, ef_construction=200, ef_runtime=200) + return DB.KiviDB.init_cls( + dim=DIM, + db_config={"host": HOST, "port": PORT, "password": None, "ssl": False}, + db_case_config=cfg, + collection_name="vdbbench_kividb_test", + drop_old=True, + with_scalar_labels=with_scalar_labels, + ) + + +def _recall( + db: object, + vectors: np.ndarray, + queries: np.ndarray, + candidates: list[int], + check: Callable[[int], bool] | None = None, +) -> float: + hits = 0 + for q in queries: + got = db.search_embedding(q.tolist(), k=K) + if check: + assert all(check(i) for i in got), f"result outside the filter: {got}" + hits += len(set(got) & set(_brute_force(vectors, q, candidates, K))) + return hits / (len(queries) * K) + + +@pytest.mark.integration +@live +@pytest.mark.parametrize( + ("filters", "candidates_of", "check_of"), + [ + (non_filter, lambda _labels: list(range(COUNT)), None), + ( + IntFilter(int_value=1500, filter_rate=0.75), + lambda _labels: list(range(1500, COUNT)), + lambda _labels: lambda i: i >= 1500, + ), + ( + LabelFilter(label_percentage=0.05), + lambda labels: [i for i in range(COUNT) if labels[i] == "label_5p"], + lambda labels: lambda i: labels[i] == "label_5p", + ), + ], + ids=["no-filter", "int-ge", "label-eq"], +) +def test_insert_search_recall( + filters: object, + candidates_of: Callable[[list[str]], list[int]], + check_of: Callable[[list[str]], Callable[[int], bool]] | None, +): + vectors, queries, labels = _data() + db = _client(with_scalar_labels=filters.type == FilterOp.StrEqual) + with db.init(): + n, err = db.insert_embeddings(vectors.tolist(), list(range(COUNT)), labels_data=labels) + assert (n, err) == (COUNT, None) + db.optimize(data_size=COUNT) + with db.init(): + db.prepare_filter(filters) + check = check_of(labels) if check_of else None + recall = _recall(db, vectors, queries, candidates_of(labels), check) + assert recall >= 0.95, f"recall {recall:.3f} < 0.95 for {filters.type}" + + +@pytest.mark.integration +@live +def test_drop_old_removes_previous_corpus(): + vectors, _, labels = _data() + db = _client() + with db.init(): + db.insert_embeddings(vectors.tolist(), list(range(COUNT)), labels_data=labels) + _client() # drop_old=True again + conn = redis.Redis(host=HOST, port=PORT) + assert not list(conn.scan_iter(match="vdbbench_kividb_test:*", count=1000)) diff --git a/vectordb_bench/backend/clients/__init__.py b/vectordb_bench/backend/clients/__init__.py index dbd83bb66..b50ceb53a 100644 --- a/vectordb_bench/backend/clients/__init__.py +++ b/vectordb_bench/backend/clients/__init__.py @@ -37,6 +37,7 @@ class DB(Enum): LakebaseVector = "LakebaseVector" Redis = "Redis" MemoryDB = "MemoryDB" + KiviDB = "KiviDB" Chroma = "Chroma" AWSOpenSearch = "OpenSearch" OSSOpenSearch = "OSSOpenSearch" @@ -135,6 +136,11 @@ def init_cls(self) -> type[VectorDB]: # noqa: PLR0911, PLR0912, C901, PLR0915 return MemoryDB + if self == DB.KiviDB: + from .kividb.kividb import KiviDB + + return KiviDB + if self == DB.Chroma: from .chroma.chroma import ChromaClient @@ -358,6 +364,11 @@ def config_cls(self) -> type[DBConfig]: # noqa: PLR0911, PLR0912, C901, PLR0915 return MemoryDBConfig + if self == DB.KiviDB: + from .kividb.config import KiviDBConfig + + return KiviDBConfig + if self == DB.Chroma: from .chroma.config import ChromaConfig @@ -735,6 +746,11 @@ def case_config_cls( # noqa: C901, PLR0911, PLR0912, PLR0915 return AdbpgIndexConfig + if self == DB.KiviDB: + from .kividb.config import _kividb_case_config + + return _kividb_case_config.get(index_type) + # DB.Pinecone, DB.Redis return EmptyDBCaseConfig diff --git a/vectordb_bench/backend/clients/kividb/cli.py b/vectordb_bench/backend/clients/kividb/cli.py new file mode 100644 index 000000000..d561bbf81 --- /dev/null +++ b/vectordb_bench/backend/clients/kividb/cli.py @@ -0,0 +1,58 @@ +from typing import Annotated, TypedDict, Unpack + +import click +from pydantic import SecretStr + +from ....cli.cli import ( + CommonTypedDict, + HNSWFlavor2, + cli, + click_parameter_decorators_from_typed_dict, + run, +) +from .. import DB +from .config import KiviDBHNSWConfig + + +class KiviDBTypedDict(TypedDict): + host: Annotated[str, click.option("--host", type=str, help="KiviDB host", required=True)] + password: Annotated[str, click.option("--password", type=str, help="KiviDB password")] + port: Annotated[int, click.option("--port", type=int, default=6380, show_default=True, help="KiviDB port")] + ssl: Annotated[ + bool, + click.option( + "--ssl/--no-ssl", + is_flag=True, + show_default=True, + default=False, + help="Connect over TLS (needs a -tls or -full KiviDB build)", + ), + ] + + +class KiviDBHNSWTypedDict(CommonTypedDict, KiviDBTypedDict, HNSWFlavor2): ... + + +@cli.command() +@click_parameter_decorators_from_typed_dict(KiviDBHNSWTypedDict) +def KiviDB(**parameters: Unpack[KiviDBHNSWTypedDict]): + from .config import KiviDBConfig + + case_config = {"ef_runtime": parameters["ef_runtime"]} + if parameters["m"] is not None: + case_config["M"] = parameters["m"] + if parameters["ef_construction"] is not None: + case_config["ef_construction"] = parameters["ef_construction"] + + run( + db=DB.KiviDB, + db_config=KiviDBConfig( + db_label=parameters["db_label"], + host=SecretStr(parameters["host"]), + port=parameters["port"], + password=SecretStr(parameters["password"]) if parameters["password"] else None, + ssl=parameters["ssl"], + ), + db_case_config=KiviDBHNSWConfig(**case_config), + **parameters, + ) diff --git a/vectordb_bench/backend/clients/kividb/config.py b/vectordb_bench/backend/clients/kividb/config.py new file mode 100644 index 000000000..3114a4c05 --- /dev/null +++ b/vectordb_bench/backend/clients/kividb/config.py @@ -0,0 +1,63 @@ +from pydantic import BaseModel, SecretStr + +from ..api import DBCaseConfig, DBConfig, IndexType, MetricType + + +class KiviDBConfig(DBConfig): + host: SecretStr + port: int = 6380 + password: SecretStr | None = None + ssl: bool = False + + def to_dict(self) -> dict: + return { + "host": self.host.get_secret_value(), + "port": self.port, + "password": self.password.get_secret_value() if self.password else None, + "ssl": self.ssl, + } + + +class KiviDBIndexConfig(BaseModel, DBCaseConfig): + metric_type: MetricType | None = None + + def parse_metric(self) -> str: + if self.metric_type == MetricType.L2: + return "L2" + if self.metric_type == MetricType.IP: + return "IP" + return "COSINE" + + +class KiviDBHNSWConfig(KiviDBIndexConfig): + M: int = 16 + ef_construction: int = 200 + ef_runtime: int | None = None + index: IndexType = IndexType.HNSW + + def index_param(self) -> dict: + return { + "index_type": self.index.value, + "metric": self.parse_metric(), + "m": self.M, + "ef_construction": self.ef_construction, + } + + def search_param(self) -> dict: + return {"ef_runtime": self.ef_runtime} + + +class KiviDBFLATConfig(KiviDBIndexConfig): + index: IndexType = IndexType.Flat + + def index_param(self) -> dict: + return {"index_type": self.index.value, "metric": self.parse_metric()} + + def search_param(self) -> dict: + return {"ef_runtime": None} + + +_kividb_case_config = { + IndexType.HNSW: KiviDBHNSWConfig, + IndexType.Flat: KiviDBFLATConfig, +} diff --git a/vectordb_bench/backend/clients/kividb/kividb.py b/vectordb_bench/backend/clients/kividb/kividb.py new file mode 100644 index 000000000..8a3fc4348 --- /dev/null +++ b/vectordb_bench/backend/clients/kividb/kividb.py @@ -0,0 +1,214 @@ +"""KiviDB client: a multi-threaded, Redis-compatible store with RediSearch-style +`FT.*` vector search built in (no module to load). + +Speaks RESP2 with raw `FT.*` commands so reply parsing does not depend on the +redis-py search helpers or on the negotiated protocol version. +""" + +import logging +import re +import time +from collections.abc import Generator +from contextlib import contextmanager +from typing import Any + +import numpy as np +import redis + +from ...filter import Filter, FilterOp +from ..api import VectorDB +from .config import KiviDBIndexConfig + +log = logging.getLogger(__name__) + +DEFAULT_INDEX_NAME = "vdbbench_kividb" +ID_FIELD = "id" +LABEL_FIELD = "labels" +VECTOR_FIELD = "vector" +PIPELINE_BATCH_SIZE = 1000 +INDEX_WAIT_TIMEOUT_S = 3600 + +# RediSearch TAG syntax: anything but letters, digits and `_` must be escaped. +_TAG_ESCAPE = re.compile(r"([^A-Za-z0-9_])") + + +def escape_tag_value(value: str) -> str: + return _TAG_ESCAPE.sub(r"\\\1", value) + + +class KiviDB(VectorDB): + supported_filter_types: list[FilterOp] = [ + FilterOp.NonFilter, + FilterOp.NumGE, + FilterOp.StrEqual, + ] + name = "KiviDB" + + def __init__( + self, + dim: int, + db_config: dict, + db_case_config: KiviDBIndexConfig, + collection_name: str = DEFAULT_INDEX_NAME, + drop_old: bool = False, + with_scalar_labels: bool = False, + **kwargs, + ): + self.db_config = db_config + self.case_config = db_case_config + self.index_name = collection_name or DEFAULT_INDEX_NAME + self.key_prefix = f"{self.index_name}:" + self.with_scalar_labels = with_scalar_labels + self.filter_expr = "*" + self.conn: redis.Redis | None = None + + conn = self._connect() + info = conn.info("server") + log.info( + f"Connected to KiviDB {info.get('kividb_version', 'unknown')} " + f"(redis_version {info.get('redis_version', 'unknown')})" + ) + if drop_old: + self._drop(conn) + self._create_index(dim, conn) + conn.close() + + def _connect(self) -> redis.Redis: + return redis.Redis( + host=self.db_config["host"], + port=self.db_config["port"], + password=self.db_config["password"], + ssl=self.db_config.get("ssl", False), + protocol=2, + ) + + def _index_exists(self, conn: redis.Redis) -> bool: + try: + conn.execute_command("FT.INFO", self.index_name) + except redis.exceptions.ResponseError: + return False + return True + + def _drop(self, conn: redis.Redis): + if self._index_exists(conn): + conn.execute_command("FT.DROPINDEX", self.index_name) + deleted = 0 + batch: list[bytes] = [] + for key in conn.scan_iter(match=f"{self.key_prefix}*", count=10_000): + batch.append(key) + if len(batch) >= 10_000: + deleted += conn.unlink(*batch) + batch = [] + if batch: + deleted += conn.unlink(*batch) + log.info(f"Dropped KiviDB index {self.index_name} and {deleted} keys") + + def _create_index(self, dim: int, conn: redis.Redis): + if self._index_exists(conn): + return + index_param = self.case_config.index_param() + vector_attrs = ["TYPE", "FLOAT32", "DIM", dim, "DISTANCE_METRIC", index_param["metric"]] + if index_param["index_type"] == "HNSW": + vector_attrs += ["M", index_param["m"], "EF_CONSTRUCTION", index_param["ef_construction"]] + schema = [ID_FIELD, "NUMERIC"] + if self.with_scalar_labels: + schema += [LABEL_FIELD, "TAG"] + schema += [VECTOR_FIELD, "VECTOR", index_param["index_type"], len(vector_attrs), *vector_attrs] + conn.execute_command( + "FT.CREATE", self.index_name, "ON", "HASH", "PREFIX", "1", self.key_prefix, "SCHEMA", *schema + ) + + @contextmanager + def init(self) -> Generator[None, None, None]: + self.conn = self._connect() + ef_runtime = self.case_config.search_param()["ef_runtime"] + self.ef_clause = f" EF_RUNTIME {ef_runtime}" if ef_runtime else "" + yield + self.conn.close() + self.conn = None + + def insert_embeddings( + self, + embeddings: list[list[float]], + metadata: list[int], + labels_data: list[str] | None = None, + **kwargs: Any, + ) -> tuple[int, Exception | None]: + assert self.conn is not None, "call init() first" + if self.with_scalar_labels and labels_data is None: + return 0, ValueError("labels_data is required when with_scalar_labels is set") + try: + with self.conn.pipeline(transaction=False) as pipe: + for i, embedding in enumerate(embeddings): + mapping = { + ID_FIELD: metadata[i], + VECTOR_FIELD: np.asarray(embedding, dtype=np.float32).tobytes(), + } + if self.with_scalar_labels: + mapping[LABEL_FIELD] = labels_data[i] + pipe.hset(f"{self.key_prefix}{metadata[i]}", mapping=mapping) + if (i + 1) % PIPELINE_BATCH_SIZE == 0: + pipe.execute() + pipe.execute() + except Exception as e: + log.warning(f"KiviDB insert failed: {e}") + return 0, e + return len(embeddings), None + + def _num_docs(self, conn: redis.Redis) -> int: + info = conn.execute_command("FT.INFO", self.index_name) + fields = dict(zip(info[0::2], info[1::2], strict=False)) + return int(fields.get(b"num_docs", 0)) + + def optimize(self, data_size: int | None = None): + """KiviDB indexes each vector inside the HSET that stores it; this only + confirms the index holds the whole corpus before search begins.""" + if not data_size: + return + conn = self._connect() + deadline = time.monotonic() + INDEX_WAIT_TIMEOUT_S + while (indexed := self._num_docs(conn)) < data_size: + if time.monotonic() > deadline: + conn.close() + msg = f"KiviDB index {self.index_name} has {indexed} of {data_size} docs after {INDEX_WAIT_TIMEOUT_S}s" + raise TimeoutError(msg) + time.sleep(1) + conn.close() + + def prepare_filter(self, filters: Filter): + if filters.type == FilterOp.NonFilter: + self.filter_expr = "*" + elif filters.type == FilterOp.NumGE: + self.filter_expr = f"(@{ID_FIELD}:[{int(filters.int_value)} +inf])" + elif filters.type == FilterOp.StrEqual: + self.filter_expr = f"(@{LABEL_FIELD}:{{{escape_tag_value(filters.label_value)}}})" + else: + msg = f"Unsupported filter for KiviDB: {filters}" + raise ValueError(msg) + + def search_embedding( + self, + query: list[float], + k: int = 100, + **kwargs: Any, + ) -> list[int]: + assert self.conn is not None, "call init() first" + res = self.conn.execute_command( + "FT.SEARCH", + self.index_name, + f"{self.filter_expr}=>[KNN {k} @{VECTOR_FIELD} $vec{self.ef_clause} AS score]", + "PARAMS", + "2", + "vec", + np.asarray(query, dtype=np.float32).tobytes(), + "SORTBY", + "score", + "NOCONTENT", + "LIMIT", + "0", + str(k), + "DIALECT", + "2", + ) + # RESP2 reply with NOCONTENT: [total, key1, key2, ...] + return [int(key.rsplit(b":", 1)[1]) for key in res[1:]] diff --git a/vectordb_bench/cli/vectordbbench.py b/vectordb_bench/cli/vectordbbench.py index 1bbc462ef..6d4d5e132 100644 --- a/vectordb_bench/cli/vectordbbench.py +++ b/vectordb_bench/cli/vectordbbench.py @@ -24,6 +24,7 @@ LanceDBIVFHNSWSQ, LanceDBIVFPQ, ) +from ..backend.clients.kividb.cli import KiviDB from ..backend.clients.lindorm.cli import LindormHNSW, LindormIVFBQ, LindormIVFPQ from ..backend.clients.mariadb.cli import MariaDBHNSW from ..backend.clients.memorydb.cli import MemoryDB @@ -65,6 +66,7 @@ cli.add_command(PgVectoRSIVFFlat) cli.add_command(Redis) cli.add_command(MemoryDB) +cli.add_command(KiviDB) cli.add_command(Weaviate) cli.add_command(Test) cli.add_command(ZillizAutoIndex)