diff --git a/README.hi-IN.md b/README.hi-IN.md index df3ecd6cf..baa63674d 100644 --- a/README.hi-IN.md +++ b/README.hi-IN.md @@ -525,6 +525,7 @@ pip install "code-review-graph[all]" # All optional dependencies | `CRG_EMBEDDING_MODEL` | स्थानीय वेक्टर एम्बेडिंग का डिफ़ॉल्ट मॉडल | `all-MiniLM-L6-v2` | | `CRG_ACCEPT_CLOUD_EMBEDDINGS` | `1` करने पर क्लाउड एम्बेडिंग की बाहर-भेजने वाली चेतावनी दबती है | - | | `CRG_ALLOW_REMOTE_CODE` | `trust_remote_code=True` माँगने वाले HuggingFace मॉडल की अनुमति | `0` | +| `CRG_VECTOR_CACHE` | `1` करने पर डिकोड किए गए एम्बेडिंग वेक्टर सिमैंटिक खोजों के बीच मेमोरी में रहते हैं, हर क्वेरी पर SQLite से सब दोबारा नहीं पढ़े जाते। बदले में RAM (हर कैश किए गए इंडेक्स के लिए वेक्टर × आयाम × 4 बाइट) और हर कैश किए गए डेटाबेस के लिए प्रोसेस चलने तक खुला एक रीड कनेक्शन लगता है: Windows पर खुले रहने तक `graph.db` को हटाया या बदला नहीं जा सकता | - | | `CRG_MAX_IMPACT_NODES` | प्रभाव विश्लेषण में अधिकतम नोड | `500` | | `CRG_MAX_IMPACT_DEPTH` | प्रभाव क्षेत्र विश्लेषण की खोज गहराई | `2` | | `CRG_MAX_BFS_DEPTH` | ग्राफ़ ट्रैवर्सल की अधिकतम गहराई | `15` | diff --git a/README.ja-JP.md b/README.ja-JP.md index c3444568b..6ab757b07 100644 --- a/README.ja-JP.md +++ b/README.ja-JP.md @@ -525,6 +525,7 @@ pip install "code-review-graph[all]" # All optional dependencies | `CRG_EMBEDDING_MODEL` | ローカルのベクトル埋め込みの既定モデル | `all-MiniLM-L6-v2` | | `CRG_ACCEPT_CLOUD_EMBEDDINGS` | `1` にするとクラウド埋め込みの送信警告を抑制 | - | | `CRG_ALLOW_REMOTE_CODE` | `trust_remote_code=True` を要する HuggingFace モデルを許可 | `0` | +| `CRG_VECTOR_CACHE` | `1` にすると、デコード済みの埋め込みベクトルをセマンティック検索の合間もメモリに保持し、クエリごとに SQLite からすべて読み直さない。代わりに RAM(キャッシュするインデックスごとに ベクトル数 × 次元数 × 4 バイト)と、キャッシュ対象のデータベースごとにプロセス終了まで開いたままの読み取り接続 1 本を使う。Windows では開いている間 `graph.db` を削除・置換できない | - | | `CRG_MAX_IMPACT_NODES` | 影響分析に含める最大ノード数 | `500` | | `CRG_MAX_IMPACT_DEPTH` | 影響範囲分析の探索の深さ | `2` | | `CRG_MAX_BFS_DEPTH` | グラフ探索の最大の深さ | `15` | diff --git a/README.ko-KR.md b/README.ko-KR.md index 239986ff6..c3a23dff8 100644 --- a/README.ko-KR.md +++ b/README.ko-KR.md @@ -525,6 +525,7 @@ pip install "code-review-graph[all]" # All optional dependencies | `CRG_EMBEDDING_MODEL` | 로컬 벡터 임베딩의 기본 모델 | `all-MiniLM-L6-v2` | | `CRG_ACCEPT_CLOUD_EMBEDDINGS` | `1`로 두면 클라우드 임베딩 전송 경고를 숨김 | - | | `CRG_ALLOW_REMOTE_CODE` | `trust_remote_code=True`가 필요한 HuggingFace 모델 허용 | `0` | +| `CRG_VECTOR_CACHE` | `1`로 두면 디코딩한 임베딩 벡터를 시맨틱 검색 사이에도 메모리에 유지해, 쿼리마다 SQLite에서 전부 다시 읽지 않음. 대신 RAM(캐시한 인덱스마다 벡터 수 × 차원 수 × 4바이트)과, 캐시한 데이터베이스마다 프로세스가 끝날 때까지 열려 있는 읽기 연결 1개를 사용함. Windows에서는 열려 있는 동안 `graph.db`를 삭제하거나 교체할 수 없음 | - | | `CRG_MAX_IMPACT_NODES` | 영향 분석에 넣는 최대 노드 수 | `500` | | `CRG_MAX_IMPACT_DEPTH` | 영향 범위 분석의 탐색 깊이 | `2` | | `CRG_MAX_BFS_DEPTH` | 그래프 순회의 최대 깊이 | `15` | diff --git a/README.md b/README.md index a16038f85..671c9bbde 100644 --- a/README.md +++ b/README.md @@ -525,6 +525,7 @@ pip install "code-review-graph[all]" # All optional dependencies | `CRG_EMBEDDING_MODEL` | Default model for local vector embeddings | `all-MiniLM-L6-v2` | | `CRG_ACCEPT_CLOUD_EMBEDDINGS` | Set to `1` to suppress the cloud embedding egress warning | - | | `CRG_ALLOW_REMOTE_CODE` | Allow HuggingFace models that require `trust_remote_code=True` | `0` | +| `CRG_VECTOR_CACHE` | Set to `1` to keep the decoded embedding vectors in memory between semantic searches instead of reading them all back from SQLite on every query. Costs RAM (vectors × dimensions × 4 bytes per cached index) and one open read connection per cached database for the life of the process: on Windows, `graph.db` cannot be deleted or replaced while it is open | - | | `CRG_MAX_IMPACT_NODES` | Maximum nodes in impact analysis | `500` | | `CRG_MAX_IMPACT_DEPTH` | Search depth for blast-radius analysis | `2` | | `CRG_MAX_BFS_DEPTH` | Maximum depth for graph traversal | `15` | diff --git a/README.zh-CN.md b/README.zh-CN.md index b64777ce9..6e746c721 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -525,6 +525,7 @@ pip install "code-review-graph[all]" # All optional dependencies | `CRG_EMBEDDING_MODEL` | 本地向量嵌入的默认模型 | `all-MiniLM-L6-v2` | | `CRG_ACCEPT_CLOUD_EMBEDDINGS` | 设为 `1` 可抑制云端嵌入的出网警告 | - | | `CRG_ALLOW_REMOTE_CODE` | 允许需要 `trust_remote_code=True` 的 HuggingFace 模型 | `0` | +| `CRG_VECTOR_CACHE` | 设为 `1` 时,在多次语义搜索之间把解码后的嵌入向量保留在内存中,不再每次查询都从 SQLite 全部重新读取。代价是内存(每个缓存的索引占 向量数 × 维度 × 4 字节),以及每个缓存的数据库在进程存续期间保持打开的一个读取连接:在 Windows 上,连接打开期间无法删除或替换 `graph.db` | - | | `CRG_MAX_IMPACT_NODES` | 影响分析中的最大节点数 | `500` | | `CRG_MAX_IMPACT_DEPTH` | 影响半径分析的搜索深度 | `2` | | `CRG_MAX_BFS_DEPTH` | 图谱遍历的最大深度 | `15` | diff --git a/code_review_graph/embeddings.py b/code_review_graph/embeddings.py index 9e60645e6..2ef1e83c5 100644 --- a/code_review_graph/embeddings.py +++ b/code_review_graph/embeddings.py @@ -1069,6 +1069,39 @@ def _cosine_similarity(a: list[float], b: list[float]) -> float: return dot / (norm_a * norm_b) +# Rows scored per numpy step: the fetchmany() size of the search, and the +# slice of a cached matrix converted to float64 at a time. +_SEARCH_CHUNK_ROWS = 500 + + +def _decode_rows(blobs: list[bytes], dims: int) -> Any: + """Stack vector blobs of *dims* components into a float32 matrix (needs numpy).""" + import numpy as np + + return np.frombuffer(b"".join(blobs), dtype=np.float32).reshape(len(blobs), dims) + + +def _row_norms(matrix: Any) -> Any: + """Each row's norm in float64, with ``inf`` for a zero row (needs numpy). + + ``inf`` turns the division in :func:`_cosine_scores` into 0.0 for a zero + vector, which is what :func:`_cosine_similarity` returns for it. + """ + import numpy as np + + rows = matrix.astype(np.float64) + norms = np.sqrt(np.einsum("ij,ij->i", rows, rows)) + norms[norms == 0.0] = np.inf + return norms + + +def _cosine_scores(matrix: Any, norms: Any, query: Any, query_norm: float) -> Any: + """Cosine of each float32 row of *matrix* with *query*, in float64.""" + import numpy as np + + return (matrix.astype(np.float64) @ query) / (norms * query_norm) + + _IDENTIFIER_SPLIT_RE = re.compile(r"([a-z])([A-Z])|[_./\-]+") _MAX_EMBEDDED_DOCSTRING_CHARS = 400 @@ -1245,13 +1278,211 @@ def embed_nodes(self, nodes: list[GraphNode], batch_size: int = 64) -> int: return embedded def search(self, query: str, limit: int = 20) -> list[tuple[str, float]]: - """Search for nodes by semantic similarity.""" + """Search for nodes by semantic similarity. + + Every stored vector of the current provider is scored against the + query. With numpy installed (the ``embeddings`` extra pulls it in), + each chunk of rows is scored with one matrix-vector product instead of + a Python loop over every component of every row, which took tens of + seconds per search over tens of thousands of high-dimensional vectors. + Without numpy, :meth:`_search_pure_python` runs that loop instead. + + Both return the same ranking: a vector whose dimensionality differs + from the query's, or whose norm is zero, scores 0.0, and equal scores + keep the order the rows were read in. + + With ``CRG_VECTOR_CACHE=1`` the decoded vectors stay in memory between + searches instead of being read back from SQLite every time; see + :meth:`_cached_vectors`. The ranking is the same either way. + """ if not self.provider: return [] provider_name = self.provider.name query_vec = self.provider.embed_query(query) + try: + import numpy as np + except ImportError: + return self._search_pure_python(query_vec, provider_name, limit) + + # float64, like the loop's Python floats: the stored components are + # float32, so every product is exact and only the summation order + # differs between the two paths. + query_arr = np.asarray(query_vec, dtype=np.float64) + dims = int(query_arr.shape[0]) + query_norm = float(np.sqrt(query_arr @ query_arr)) + + cached = self._cached_vectors(provider_name, dims) + if cached is None: + names, all_scores = self._stream_scores(provider_name, query_arr, query_norm) + else: + names, rows, matrix, norms = cached + all_scores = np.zeros(len(names)) + if query_norm: + for start in range(0, len(rows), _SEARCH_CHUNK_ROWS): + end = start + _SEARCH_CHUNK_ROWS + all_scores[rows[start:end]] = _cosine_scores( + matrix[start:end], norms[start:end], query_arr, query_norm, + ) + + if not names: + return [] + # Stable, like list.sort(reverse=True): equal scores keep read order. + order = np.argsort(-all_scores, kind="stable")[:limit] + return [(names[i], float(all_scores[i])) for i in order] + + def _stream_scores( + self, provider_name: str, query_arr: Any, query_norm: float, + ) -> tuple[list[str], Any]: + """Score the stored vectors chunk by chunk, keeping none of them.""" + import numpy as np + + dims = int(query_arr.shape[0]) + blob_size = dims * 4 # float32 components, see _encode_vector + names: list[str] = [] + chunk_scores: list[Any] = [] + cursor = self._conn.execute( + "SELECT qualified_name, vector FROM embeddings WHERE provider = ?", + (provider_name,), + ) + while True: + rows = cursor.fetchmany(_SEARCH_CHUNK_ROWS) + if not rows: + break + scores = np.zeros(len(rows)) + same_dims = [i for i, row in enumerate(rows) if len(row["vector"]) == blob_size] + if same_dims and query_norm: + matrix = _decode_rows([rows[i]["vector"] for i in same_dims], dims) + scores[same_dims] = _cosine_scores( + matrix, _row_norms(matrix), query_arr, query_norm, + ) + names.extend(row["qualified_name"] for row in rows) + chunk_scores.append(scores) + return names, np.concatenate(chunk_scores) if chunk_scores else np.zeros(0) + + # Opt-in (CRG_VECTOR_CACHE=1) cache of decoded vectors, shared by every + # store of the process: _embedding_search() opens and closes a store per + # query, so a cache kept on the instance would never be hit twice. + # (database, provider, dims) -> (token, names, rows, matrix, norms), with + # the least recently used entry first. + _vector_cache: dict[tuple[str, str, int], tuple[Any, ...]] = {} + # database -> (file identity, connection that only reads data_version) + _vector_cache_watchers: dict[str, tuple[tuple[int, int], sqlite3.Connection]] = {} + _vector_cache_lock = threading.Lock() + _VECTOR_CACHE_MAX_ENTRIES = 4 + + def _cached_vectors(self, provider_name: str, dims: int) -> tuple[Any, ...] | None: + """``(names, rows, matrix, norms)`` from the process-wide cache, or None. + + *names* lists every stored vector of *provider_name* in read order, + *rows* the positions in *names* of those with *dims* components, + *matrix* their float32 components and *norms* their norms (``inf`` for + a zero vector). + + None, and the caller reads SQLite instead, unless ``CRG_VECTOR_CACHE`` + is ``1``. Also None for a database that is not a file, and while this + connection holds an uncommitted transaction whose rows only it can + see. + + An entry is valid while ``PRAGMA data_version`` on a long-lived + watcher connection stays put. That value moves whenever any other + connection commits, this store's included. It is read before the rows, + so a commit that races the load costs one extra reload, never a stale + answer. The watcher stays open while its database has a cached entry. + On Windows that means ``graph.db`` cannot be deleted or replaced + while it is open; :meth:`clear_vector_cache` releases it. + """ + if os.environ.get("CRG_VECTOR_CACHE", "").strip() != "1": + return None + if self._conn.in_transaction: + return None + path = str(self.db_path) + if path == ":memory:" or not os.path.isfile(path): + return None + db = os.path.normcase(os.path.realpath(path)) + + cls = EmbeddingStore + with cls._vector_cache_lock: + try: + token = cls._vector_cache_token(db) + except (sqlite3.Error, OSError) as exc: + logger.warning("Vector cache unavailable for %s: %s", db, exc) + return None + + key = (db, provider_name, dims) + entry = cls._vector_cache.pop(key, None) + if entry is not None and entry[0] != token: + entry = None # release the stale matrix before decoding a new one + if entry is None: + entry = (token, *self._load_vectors(provider_name, dims)) + cls._vector_cache[key] = entry # most recently used goes last + + while len(cls._vector_cache) > cls._VECTOR_CACHE_MAX_ENTRIES: + evicted = next(iter(cls._vector_cache)) + del cls._vector_cache[evicted] + if all(other[0] != evicted[0] for other in cls._vector_cache): + cls._close_vector_cache_watcher(evicted[0]) + return entry[1:] + + def _load_vectors(self, provider_name: str, dims: int) -> tuple[Any, ...]: + """Decode every stored vector of *provider_name* for the cache.""" + import numpy as np + + blob_size = dims * 4 + names: list[str] = [] + rows: list[int] = [] + blobs: list[bytes] = [] + for name, blob in self._conn.execute( + "SELECT qualified_name, vector FROM embeddings WHERE provider = ?", + (provider_name,), + ): + if len(blob) == blob_size: + rows.append(len(names)) + blobs.append(blob) + names.append(name) + + matrix = _decode_rows(blobs, dims) + norms = np.concatenate([ + _row_norms(matrix[start:start + _SEARCH_CHUNK_ROWS]) + for start in range(0, len(blobs), _SEARCH_CHUNK_ROWS) + ]) if blobs else np.zeros(0) + return names, np.asarray(rows, dtype=np.intp), matrix, norms + + @classmethod + def _vector_cache_token(cls, db: str) -> tuple[tuple[int, int], int]: + """(file identity, data_version) of *db*. The caller holds the lock.""" + stat = os.stat(db) + file_id = (stat.st_dev, stat.st_ino) + watcher = cls._vector_cache_watchers.get(db) + if watcher is None or watcher[0] != file_id: + # A replaced file gets a fresh watcher: data_version only compares + # commits seen by one connection on one file. + cls._close_vector_cache_watcher(db) + watcher = (file_id, sqlite3.connect(db, timeout=30, check_same_thread=False)) + cls._vector_cache_watchers[db] = watcher + # fetchall() runs the statement to completion: one left open would pin + # a read snapshot and freeze data_version on this connection. + return file_id, watcher[1].execute("PRAGMA data_version").fetchall()[0][0] + + @classmethod + def _close_vector_cache_watcher(cls, db: str) -> None: + watcher = cls._vector_cache_watchers.pop(db, None) + if watcher is not None: + watcher[1].close() + + @classmethod + def clear_vector_cache(cls) -> None: + """Drop every cached matrix and close the connections that watch them.""" + with cls._vector_cache_lock: + cls._vector_cache.clear() + for db in list(cls._vector_cache_watchers): + cls._close_vector_cache_watcher(db) + + def _search_pure_python( + self, query_vec: list[float], provider_name: str, limit: int, + ) -> list[tuple[str, float]]: + """Score every stored vector one component at a time (no numpy).""" # Process in chunks, only matching current provider scored: list[tuple[str, float]] = [] cursor = self._conn.execute( diff --git a/pyproject.toml b/pyproject.toml index d527ef10a..d3c5d6148 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -85,6 +85,9 @@ enrichment = [ ] dev = [ "mypy>=1.10,<3", + # Optional at runtime (the "embeddings" extra pulls it in), but + # EmbeddingStore.search() has a numpy path the test job has to exercise. + "numpy>=1.26,<3", "pytest>=8.0,<9", "pytest-asyncio>=0.23,<2", "pytest-cov>=4.0,<8", diff --git a/tests/test_embeddings.py b/tests/test_embeddings.py index 079217a2c..a39158dac 100644 --- a/tests/test_embeddings.py +++ b/tests/test_embeddings.py @@ -2,6 +2,9 @@ import json import os +import random +import sqlite3 +import sys import threading import time from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer @@ -231,6 +234,283 @@ def dimension(self): store.close() +class TestEmbeddingStoreSearch: + """``search`` scores with numpy when it can, and ranks like the Python loop.""" + + PROVIDER = "test:provider" + + @pytest.fixture(autouse=True) + def _vector_cache_env(self, monkeypatch): + # The uncached path, even when the developer exports the variable. + monkeypatch.delenv("CRG_VECTOR_CACHE", raising=False) + + class _Provider: + name = "test:provider" + + def __init__(self, query_vec): + self.query_vec = query_vec + + def embed(self, texts): + raise AssertionError("search must not embed stored vectors") + + def embed_query(self, text): + return list(self.query_vec) + + @property + def dimension(self): + return len(self.query_vec) + + def _store(self, tmp_path, query_vec, name="embeddings.db"): + provider = self._Provider(query_vec) + with patch("code_review_graph.embeddings.get_provider", return_value=provider): + return EmbeddingStore(tmp_path / name) + + @staticmethod + def _insert(store, name, vec, provider=PROVIDER): + store._conn.execute( + "INSERT INTO embeddings (qualified_name, vector, text_hash, provider) " + "VALUES (?, ?, ?, ?)", + (name, _encode_vector(vec), "hash", provider), + ) + + def test_matches_the_python_loop(self, tmp_path): + pytest.importorskip("numpy") + rng = random.Random(7) + query = [rng.uniform(-1.0, 1.0) for _ in range(24)] + store = self._store(tmp_path, query) + try: + # 1234 rows cross the 500-row chunk boundary twice. + for i in range(1234): + vec = [rng.uniform(-1.0, 1.0) for _ in range(24)] + self._insert(store, f"file.py::func_{i}", vec) + + expected = store._search_pure_python(query, self.PROVIDER, 1234) + results = store.search("query", limit=1234) + + assert [name for name, _ in results] == [name for name, _ in expected] + for (_, score), (_, want) in zip(results, expected): + assert score == pytest.approx(want, abs=1e-12) + finally: + store.close() + + def test_zero_norm_and_other_dimensions_score_zero_in_read_order(self, tmp_path): + pytest.importorskip("numpy") + store = self._store(tmp_path, [1.0, 0.0, 0.0]) + try: + self._insert(store, "zero_norm", [0.0, 0.0, 0.0]) + self._insert(store, "two_dims", [1.0, 0.0]) + self._insert(store, "orthogonal", [0.0, 1.0, 0.0]) + self._insert(store, "aligned", [2.0, 0.0, 0.0]) + self._insert(store, "other_provider", [1.0, 0.0, 0.0], provider="other:x") + + results = store.search("query", limit=10) + + assert results == [ + ("aligned", 1.0), + ("zero_norm", 0.0), + ("two_dims", 0.0), + ("orthogonal", 0.0), + ] + assert results == store._search_pure_python([1.0, 0.0, 0.0], self.PROVIDER, 10) + finally: + store.close() + + def test_zero_query_scores_every_row_zero_and_limit_applies(self, tmp_path): + pytest.importorskip("numpy") + store = self._store(tmp_path, [0.0, 0.0]) + try: + for name in ("a", "b", "c"): + self._insert(store, name, [1.0, 2.0]) + + assert store.search("query", limit=2) == [("a", 0.0), ("b", 0.0)] + assert store.search("query", limit=0) == [] + finally: + store.close() + + def test_empty_index_returns_nothing(self, tmp_path): + store = self._store(tmp_path, [1.0, 0.0]) + try: + assert store.search("query") == [] + finally: + store.close() + + def test_falls_back_to_the_python_loop_without_numpy(self, tmp_path): + store = self._store(tmp_path, [1.0, 0.0]) + try: + self._insert(store, "a", [0.0, 1.0]) + self._insert(store, "b", [1.0, 1.0]) + + with patch.dict(sys.modules, {"numpy": None}), patch.object( + store, "_search_pure_python", wraps=store._search_pure_python, + ) as loop: + results = store.search("query", limit=5) + + loop.assert_called_once_with([1.0, 0.0], self.PROVIDER, 5) + assert [name for name, _ in results] == ["b", "a"] + assert results[0][1] == pytest.approx(2 ** -0.5) + assert results[1][1] == 0.0 + finally: + store.close() + + +class TestEmbeddingStoreVectorCache(TestEmbeddingStoreSearch): + """``CRG_VECTOR_CACHE=1``: every search test above, plus the cache's own. + + The inherited tests run with the cache on, so they pin that it ranks + exactly like the uncached path and the Python loop. + """ + + @pytest.fixture(autouse=True) + def _vector_cache_env(self, monkeypatch): + monkeypatch.setenv("CRG_VECTOR_CACHE", "1") + EmbeddingStore.clear_vector_cache() + yield + # Also closes the watcher connections, which on Windows would keep + # tmp_path's databases from being deleted. + EmbeddingStore.clear_vector_cache() + + @staticmethod + def _entry(): + (entry,) = EmbeddingStore._vector_cache.values() + return entry + + def test_off_without_the_variable(self, tmp_path, monkeypatch): + pytest.importorskip("numpy") + monkeypatch.delenv("CRG_VECTOR_CACHE") + store = self._store(tmp_path, [1.0, 0.0]) + try: + self._insert(store, "a", [1.0, 0.0]) + assert [name for name, _ in store.search("query")] == ["a"] + assert EmbeddingStore._vector_cache == {} + assert EmbeddingStore._vector_cache_watchers == {} + finally: + store.close() + + def test_a_second_store_reuses_the_decoded_vectors(self, tmp_path): + pytest.importorskip("numpy") + first = self._store(tmp_path, [1.0, 0.0]) + try: + self._insert(first, "a", [1.0, 0.0]) + first.search("query") + entry = self._entry() + finally: + first.close() + + second = self._store(tmp_path, [1.0, 0.0]) + try: + assert second.search("query") == [("a", 1.0)] + assert self._entry() is entry + finally: + second.close() + + def test_a_commit_from_another_connection_refreshes_the_cache(self, tmp_path): + pytest.importorskip("numpy") + store = self._store(tmp_path, [1.0, 0.0]) + try: + self._insert(store, "a", [0.0, 1.0]) + assert [name for name, _ in store.search("query")] == ["a"] + + other = sqlite3.connect(tmp_path / "embeddings.db") + other.execute( + "INSERT INTO embeddings VALUES ('b', ?, 'hash', ?)", + (_encode_vector([1.0, 0.0]), self.PROVIDER), + ) + other.commit() + other.close() + + assert [name for name, _ in store.search("query")] == ["b", "a"] + finally: + store.close() + + def test_uncommitted_rows_bypass_the_cache(self, tmp_path): + pytest.importorskip("numpy") + store = self._store(tmp_path, [1.0, 0.0]) + try: + self._insert(store, "a", [0.0, 1.0]) + store.search("query") + entry = self._entry() + + store._conn.execute("BEGIN") + self._insert(store, "b", [1.0, 0.0]) + assert [name for name, _ in store.search("query")] == ["b", "a"] + store._conn.execute("ROLLBACK") + + assert [name for name, _ in store.search("query")] == ["a"] + assert self._entry() is entry + finally: + store.close() + + def test_evicting_a_database_closes_its_watcher(self, tmp_path, monkeypatch): + pytest.importorskip("numpy") + monkeypatch.setattr(EmbeddingStore, "_VECTOR_CACHE_MAX_ENTRIES", 1) + one = self._store(tmp_path, [1.0, 0.0], name="one.db") + two = self._store(tmp_path, [1.0, 0.0], name="two.db") + try: + self._insert(one, "a", [1.0, 0.0]) + self._insert(two, "b", [1.0, 0.0]) + one.search("query") + (watcher,) = [conn for _, conn in EmbeddingStore._vector_cache_watchers.values()] + + two.search("query") + + assert len(EmbeddingStore._vector_cache) == 1 + assert list(EmbeddingStore._vector_cache_watchers) == [ + os.path.normcase(os.path.realpath(tmp_path / "two.db")), + ] + with pytest.raises(sqlite3.ProgrammingError): + watcher.execute("SELECT 1") + finally: + one.close() + two.close() + + def test_a_new_file_identity_gets_a_fresh_watcher(self, tmp_path): + # Stands in for graph.db being replaced, say restored from a backup. + # Windows refuses the real replacement while the watcher holds the + # file open, so the new inode is simulated. + pytest.importorskip("numpy") + store = self._store(tmp_path, [1.0, 0.0]) + try: + self._insert(store, "a", [1.0, 0.0]) + store.search("query") + entry = self._entry() + ((old_id, old_watcher),) = EmbeddingStore._vector_cache_watchers.values() + + real_stat = os.stat + + def replaced(path, *args, **kwargs): + st = real_stat(path, *args, **kwargs) + return os.stat_result((st.st_mode, st.st_ino + 1) + tuple(st)[2:]) + + with patch.object(os, "stat", replaced): + assert store.search("query") == [("a", 1.0)] + + ((new_id, new_watcher),) = EmbeddingStore._vector_cache_watchers.values() + assert new_id != old_id + assert new_watcher is not old_watcher + with pytest.raises(sqlite3.ProgrammingError): + old_watcher.execute("SELECT 1") + assert self._entry() is not entry + finally: + store.close() + + def test_clear_vector_cache_closes_the_watchers(self, tmp_path): + pytest.importorskip("numpy") + store = self._store(tmp_path, [1.0, 0.0]) + try: + self._insert(store, "a", [1.0, 0.0]) + store.search("query") + (watcher,) = [conn for _, conn in EmbeddingStore._vector_cache_watchers.values()] + + EmbeddingStore.clear_vector_cache() + + assert EmbeddingStore._vector_cache == {} + assert EmbeddingStore._vector_cache_watchers == {} + with pytest.raises(sqlite3.ProgrammingError): + watcher.execute("SELECT 1") + finally: + store.close() + + class TestLocalEmbeddingProviderModelName: """Tests for configurable model name on LocalEmbeddingProvider.""" diff --git a/uv.lock b/uv.lock index 80677c1d2..0c458bf9a 100644 --- a/uv.lock +++ b/uv.lock @@ -474,6 +474,9 @@ communities = [ ] dev = [ { name = "mypy" }, + { name = "numpy", version = "2.2.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, + { name = "numpy", version = "2.4.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version == '3.11.*'" }, + { name = "numpy", version = "2.5.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, { name = "pytest" }, { name = "pytest-asyncio" }, { name = "pytest-cov" }, @@ -524,6 +527,7 @@ requires-dist = [ { name = "mcp", specifier = ">=1.0.0,<3" }, { name = "mypy", marker = "extra == 'dev'", specifier = ">=1.10,<3" }, { name = "networkx", specifier = ">=3.2,<4" }, + { name = "numpy", marker = "extra == 'dev'", specifier = ">=1.26,<3" }, { name = "numpy", marker = "extra == 'embeddings'", specifier = ">=1.26,<3" }, { name = "ollama", marker = "extra == 'wiki'", specifier = ">=0.1.0" }, { name = "playwright", marker = "extra == 'browser-test'", specifier = ">=1.45,<2" },