Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -0,0 +1,94 @@
import os
import tempfile
import unittest

from vcache.vcache_core.cache.embedding_store.embedding_metadata_storage import (
SQLiteEmbeddingMetadataStorage,
)
from vcache.vcache_core.cache.embedding_store.embedding_metadata_storage.embedding_metadata_obj import (
EmbeddingMetadataObj,
)


class TestSQLiteEmbeddingMetadataStorage(unittest.TestCase):
def setUp(self):
# ignore_cleanup_errors: on Windows, sqlite3 keeps the file handle
# open for the lifetime of the connection, which otherwise makes
# tearDown's rmtree fail with a PermissionError.
self.tmp_dir = tempfile.TemporaryDirectory(ignore_cleanup_errors=True)
self.db_path = os.path.join(self.tmp_dir.name, "metadata.sqlite3")

def tearDown(self):
self.tmp_dir.cleanup()

def test_sqlite_strategy(self):
embedding_metadata_storage = SQLiteEmbeddingMetadataStorage(self.db_path)

initial_obj = EmbeddingMetadataObj(embedding_id=0, response="test")
embedding_id = embedding_metadata_storage.add_metadata(
embedding_id=0, metadata=initial_obj
)
assert embedding_id == 0
assert embedding_metadata_storage.get_metadata(embedding_id=0) == initial_obj

updated_obj = EmbeddingMetadataObj(embedding_id=0, response="test2")
embedding_metadata_storage.update_metadata(embedding_id=0, metadata=updated_obj)
assert embedding_metadata_storage.get_metadata(embedding_id=0) == updated_obj

embedding_metadata_storage.flush()
with self.assertRaises(ValueError):
embedding_metadata_storage.get_metadata(embedding_id=0)

def test_get_missing_metadata_raises(self):
embedding_metadata_storage = SQLiteEmbeddingMetadataStorage(self.db_path)
with self.assertRaises(ValueError):
embedding_metadata_storage.get_metadata(embedding_id=42)

def test_update_missing_metadata_raises(self):
embedding_metadata_storage = SQLiteEmbeddingMetadataStorage(self.db_path)
with self.assertRaises(ValueError):
embedding_metadata_storage.update_metadata(
embedding_id=42,
metadata=EmbeddingMetadataObj(embedding_id=42, response="test"),
)

def test_remove_metadata(self):
embedding_metadata_storage = SQLiteEmbeddingMetadataStorage(self.db_path)
embedding_metadata_storage.add_metadata(
embedding_id=1,
metadata=EmbeddingMetadataObj(embedding_id=1, response="test"),
)
assert embedding_metadata_storage.remove_metadata(embedding_id=1) is True
assert embedding_metadata_storage.remove_metadata(embedding_id=1) is False

def test_get_all_embedding_metadata_objects(self):
embedding_metadata_storage = SQLiteEmbeddingMetadataStorage(self.db_path)
for i in range(3):
embedding_metadata_storage.add_metadata(
embedding_id=i,
metadata=EmbeddingMetadataObj(embedding_id=i, response=f"test{i}"),
)
all_metadata = embedding_metadata_storage.get_all_embedding_metadata_objects()
assert len(all_metadata) == 3
assert {meta.response for meta in all_metadata} == {"test0", "test1", "test2"}

def test_survives_simulated_restart(self):
"""Data written by one instance should be visible to a fresh instance
pointed at the same file, simulating a process restart."""
first_instance = SQLiteEmbeddingMetadataStorage(self.db_path)
first_instance.add_metadata(
embedding_id=7,
metadata=EmbeddingMetadataObj(
embedding_id=7, response="persisted", id_set=3
),
)
del first_instance

second_instance = SQLiteEmbeddingMetadataStorage(self.db_path)
restored = second_instance.get_metadata(embedding_id=7)
assert restored.response == "persisted"
assert restored.id_set == 3


if __name__ == "__main__":
unittest.main()
89 changes: 89 additions & 0 deletions tests/unit/VectorDBStrategy/test_persistent_hnsw_lib.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,89 @@
import os
import tempfile
import unittest

from vcache.vcache_core.cache.embedding_store.vector_db import (
PersistentHNSWLibVectorDB,
SimilarityMetricType,
)


class TestPersistentHNSWLibVectorDB(unittest.TestCase):
def setUp(self):
self.tmp_dir = tempfile.TemporaryDirectory()
self.persist_path = os.path.join(self.tmp_dir.name, "index.hnsw")

def tearDown(self):
self.tmp_dir.cleanup()

def test_add_and_get_knn(self):
vector_db = PersistentHNSWLibVectorDB(persist_path=self.persist_path)

embedding = [0.1, 0.2, 0.3]
id1 = vector_db.add(embedding=embedding)
knn = vector_db.get_knn(embedding=embedding, k=1)
assert len(knn) == 1
assert knn[0][1] == id1

vector_db.add(embedding=[0.2, 0.3, 0.4])
vector_db.add(embedding=[0.3, 0.4, 0.5])
knn = vector_db.get_knn(embedding=embedding, k=3)
assert len(knn) == 3

def test_remove(self):
vector_db = PersistentHNSWLibVectorDB(persist_path=self.persist_path)

id1 = vector_db.add(embedding=[0.1, 0.2, 0.3])
id2 = vector_db.add(embedding=[0.2, 0.3, 0.4])
vector_db.remove(embedding_id=id1)

knn = vector_db.get_knn(embedding=[0.1, 0.2, 0.3], k=2)
assert len(knn) == 1
assert knn[0][1] == id2

def test_persists_files_to_disk(self):
vector_db = PersistentHNSWLibVectorDB(persist_path=self.persist_path)
vector_db.add(embedding=[0.1, 0.2, 0.3])

assert os.path.exists(self.persist_path)
assert os.path.exists(self.persist_path + ".meta.json")

def test_survives_simulated_restart(self):
"""Embeddings added by one instance should be visible to a fresh
instance pointed at the same path, simulating a process restart."""
first_instance = PersistentHNSWLibVectorDB(persist_path=self.persist_path)
first_instance.add(embedding=[0.1, 0.2, 0.3])
first_instance.add(embedding=[0.2, 0.3, 0.4])
del first_instance

second_instance = PersistentHNSWLibVectorDB(persist_path=self.persist_path)
knn = second_instance.get_knn(embedding=[0.1, 0.2, 0.3], k=2)
assert len(knn) == 2

# New adds after reload should not collide with restored ids
new_id = second_instance.add(embedding=[0.3, 0.4, 0.5])
assert new_id == 2

def test_reset(self):
vector_db = PersistentHNSWLibVectorDB(persist_path=self.persist_path)
vector_db.add(embedding=[0.1, 0.2, 0.3])
vector_db.add(embedding=[0.2, 0.3, 0.4])

vector_db.reset()

knn = vector_db.get_knn(embedding=[0.1, 0.2, 0.3], k=3)
assert len(knn) == 0

def test_euclidean_metric(self):
vector_db = PersistentHNSWLibVectorDB(
persist_path=self.persist_path,
similarity_metric_type=SimilarityMetricType.EUCLIDEAN,
)
embedding = [0.1, 0.2, 0.3]
id1 = vector_db.add(embedding=embedding)
knn = vector_db.get_knn(embedding=embedding, k=1)
assert knn[0][1] == id1


if __name__ == "__main__":
unittest.main()
4 changes: 4 additions & 0 deletions vcache/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,13 +37,15 @@
from .vcache_core.cache.embedding_store.embedding_metadata_storage import (
InMemoryEmbeddingMetadataStorage,
LangchainMetadataStorage,
SQLiteEmbeddingMetadataStorage,
)

# Concrete Vector databases
from .vcache_core.cache.embedding_store.vector_db import (
ChromaVectorDB,
FAISSVectorDB,
HNSWLibVectorDB,
PersistentHNSWLibVectorDB,
SimilarityMetricType,
VectorDB,
)
Expand Down Expand Up @@ -112,6 +114,7 @@
# Concrete Vector databases
"FAISSVectorDB",
"HNSWLibVectorDB",
"PersistentHNSWLibVectorDB",
"ChromaVectorDB",
"SimilarityMetricType",
# Concrete Similarity evaluators
Expand All @@ -128,5 +131,6 @@
# Concrete Embedding metadata storage
"InMemoryEmbeddingMetadataStorage",
"LangchainMetadataStorage",
"SQLiteEmbeddingMetadataStorage",
"EmbeddingMetadataObj",
]
Original file line number Diff line number Diff line change
Expand Up @@ -7,9 +7,13 @@
from vcache.vcache_core.cache.embedding_store.embedding_metadata_storage.strategies.langchain import (
LangchainMetadataStorage,
)
from vcache.vcache_core.cache.embedding_store.embedding_metadata_storage.strategies.sqlite import (
SQLiteEmbeddingMetadataStorage,
)

__all__ = [
"EmbeddingMetadataStorage",
"InMemoryEmbeddingMetadataStorage",
"LangchainMetadataStorage",
"SQLiteEmbeddingMetadataStorage",
]
Original file line number Diff line number Diff line change
@@ -1,4 +1,9 @@
from .in_memory import InMemoryEmbeddingMetadataStorage
from .langchain import LangchainMetadataStorage
from .sqlite import SQLiteEmbeddingMetadataStorage

__all__ = ["InMemoryEmbeddingMetadataStorage", "LangchainMetadataStorage"]
__all__ = [
"InMemoryEmbeddingMetadataStorage",
"LangchainMetadataStorage",
"SQLiteEmbeddingMetadataStorage",
]
Original file line number Diff line number Diff line change
@@ -0,0 +1,153 @@
import pickle
import sqlite3
import threading
from typing import List

from vcache.vcache_core.cache.embedding_store.embedding_metadata_storage.embedding_metadata_obj import (
EmbeddingMetadataObj,
)
from vcache.vcache_core.cache.embedding_store.embedding_metadata_storage.embedding_metadata_storage import (
EmbeddingMetadataStorage,
)


class SQLiteEmbeddingMetadataStorage(EmbeddingMetadataStorage):
"""
SQLite-backed implementation of embedding metadata storage.

Unlike `InMemoryEmbeddingMetadataStorage`, this implementation persists
metadata to disk, so it survives process restarts. Each metadata object
is stored as a pickled blob keyed by `embedding_id`, rather than mapped to
individual columns, since `EmbeddingMetadataObj` mixes datetimes, optional
floats, and tuples, and gains fields over time.
"""

def __init__(self, db_path: str):
"""
Initialize SQLite-backed embedding metadata storage.

Args:
db_path: Path to the SQLite database file. Created if it does
not yet exist.
"""
self.db_path = db_path
self._lock = threading.Lock()
self._connection = sqlite3.connect(db_path, check_same_thread=False)
with self._lock:
self._connection.execute(
"CREATE TABLE IF NOT EXISTS embedding_metadata ("
"embedding_id INTEGER PRIMARY KEY, data BLOB NOT NULL)"
)
self._connection.commit()

def add_metadata(self, embedding_id: int, metadata: EmbeddingMetadataObj) -> int:
"""
Add metadata for a specific embedding.

Args:
embedding_id: The id of the embedding to add the metadata for.
metadata: The metadata to add to the embedding.

Returns:
The id of the embedding.
"""
with self._lock:
self._connection.execute(
"INSERT OR REPLACE INTO embedding_metadata (embedding_id, data) "
"VALUES (?, ?)",
(embedding_id, pickle.dumps(metadata)),
)
self._connection.commit()
return embedding_id

def get_metadata(self, embedding_id: int) -> EmbeddingMetadataObj:
"""
Get metadata for a specific embedding.

Args:
embedding_id: The id of the embedding to get the metadata for.

Returns:
The metadata of the embedding.

Raises:
ValueError: If embedding metadata is not found.
"""
with self._lock:
row = self._connection.execute(
"SELECT data FROM embedding_metadata WHERE embedding_id = ?",
(embedding_id,),
).fetchone()
if row is None:
raise ValueError(
f"Embedding metadata for embedding id {embedding_id} not found"
)
return pickle.loads(row[0])

def update_metadata(
self, embedding_id: int, metadata: EmbeddingMetadataObj
) -> EmbeddingMetadataObj:
"""
Update metadata for a specific embedding.

Args:
embedding_id: The id of the embedding to update the metadata for.
metadata: The metadata to update the embedding with.

Returns:
The updated metadata of the embedding.

Raises:
ValueError: If embedding metadata is not found.
"""
with self._lock:
cursor = self._connection.execute(
"UPDATE embedding_metadata SET data = ? WHERE embedding_id = ?",
(pickle.dumps(metadata), embedding_id),
)
self._connection.commit()
not_found = cursor.rowcount == 0
if not_found:
raise ValueError(
f"Embedding metadata for embedding id {embedding_id} not found"
)
return metadata

def remove_metadata(self, embedding_id: int) -> bool:
"""
Remove metadata for a specific embedding.

Args:
embedding_id: The id of the embedding to remove metadata for.

Returns:
True if metadata was removed, False if not found.
"""
with self._lock:
cursor = self._connection.execute(
"DELETE FROM embedding_metadata WHERE embedding_id = ?",
(embedding_id,),
)
self._connection.commit()
return cursor.rowcount > 0

def flush(self) -> None:
"""
Flush all metadata from storage.
"""
with self._lock:
self._connection.execute("DELETE FROM embedding_metadata")
self._connection.commit()

def get_all_embedding_metadata_objects(self) -> List[EmbeddingMetadataObj]:
"""
Get all embedding metadata objects in storage.

Returns:
A list of all the embedding metadata objects in the storage.
"""
with self._lock:
rows = self._connection.execute(
"SELECT data FROM embedding_metadata"
).fetchall()
return [pickle.loads(row[0]) for row in rows]
Loading
Loading