diff --git a/src/paperless_ai/embedding.py b/src/paperless_ai/embedding.py index 3fc3b6d18..0d11ea423 100644 --- a/src/paperless_ai/embedding.py +++ b/src/paperless_ai/embedding.py @@ -101,9 +101,12 @@ def get_configured_model_name(config: AIConfig) -> str: """Return the canonical name of the currently configured embedding model.""" # dict.get(key, default) overload resolution fails for TextChoices keys in some # type checkers; use `or` fallback to avoid the ambiguity. - default = _DEFAULT_MODEL_NAMES.get( - config.llm_embedding_backend, - ) or "sentence-transformers/all-MiniLM-L6-v2" + default = ( + _DEFAULT_MODEL_NAMES.get( + config.llm_embedding_backend, + ) + or "sentence-transformers/all-MiniLM-L6-v2" + ) return config.llm_embedding_model or default diff --git a/src/paperless_ai/tests/test_ai_indexing.py b/src/paperless_ai/tests/test_ai_indexing.py index 2628596d4..f245d9100 100644 --- a/src/paperless_ai/tests/test_ai_indexing.py +++ b/src/paperless_ai/tests/test_ai_indexing.py @@ -163,8 +163,6 @@ def test_update_llm_index( build_document_node.assert_called_once_with(real_document, chunk_size=512) - - @pytest.mark.django_db def test_update_llm_index_rebuilds_on_model_name_change( temp_llm_index_dir: Path, diff --git a/src/paperless_ai/tests/test_vector_store.py b/src/paperless_ai/tests/test_vector_store.py index fb066d1db..ff012a6b8 100644 --- a/src/paperless_ai/tests/test_vector_store.py +++ b/src/paperless_ai/tests/test_vector_store.py @@ -32,7 +32,12 @@ def store(tmp_path: Path) -> PaperlessSqliteVecVectorStore: return PaperlessSqliteVecVectorStore(uri=str(tmp_path)) -def _query(store: PaperlessSqliteVecVectorStore, embedding: list[float], top_k: int = 5, filters=None): +def _query( + store: PaperlessSqliteVecVectorStore, + embedding: list[float], + top_k: int = 5, + filters=None, +): from llama_index.core.vector_stores.types import VectorStoreQuery return store.query( @@ -52,7 +57,9 @@ def _in_filter(document_ids: list[str]): return MetadataFilters( filters=[ MetadataFilter( - key="document_id", operator=FilterOperator.IN, value=document_ids, + key="document_id", + operator=FilterOperator.IN, + value=document_ids, ), ], ) @@ -100,7 +107,10 @@ class TestCrud: [make_node(f"n{i}", str(i % 4), seed=float(i)) for i in range(12)], ) result = _query( - store, [0.0] * DIM, top_k=3, filters=_in_filter(["0", "1", "2", "3"]), + store, + [0.0] * DIM, + top_k=3, + filters=_in_filter(["0", "1", "2", "3"]), ) assert len(result.ids) == 3 assert result.similarities == sorted(result.similarities, reverse=True) @@ -151,7 +161,11 @@ class TestUpsert: store.upsert_document("1", []) assert _query(store, [0.0] * DIM, top_k=10).ids == ["b1"] - def test_upsert_is_atomic_for_concurrent_readers(self, store, tmp_path: Path) -> None: + def test_upsert_is_atomic_for_concurrent_readers( + self, + store, + tmp_path: Path, + ) -> None: """A second connection must never observe document 1 half-replaced.""" store.add([make_node("a1", "1"), make_node("a2", "1")]) reader = PaperlessSqliteVecVectorStore(uri=str(tmp_path)) @@ -171,13 +185,15 @@ class TestMetadataCoercion: class TestModelNameTracking: def test_stored_model_name_none_without_table(self, tmp_path: Path) -> None: store = PaperlessSqliteVecVectorStore( - uri=str(tmp_path), embed_model_name="model-a", + uri=str(tmp_path), + embed_model_name="model-a", ) assert store.stored_model_name() is None def test_model_name_stored_after_add_and_persists(self, tmp_path: Path) -> None: store = PaperlessSqliteVecVectorStore( - uri=str(tmp_path), embed_model_name="model-a", + uri=str(tmp_path), + embed_model_name="model-a", ) store.add([make_node("a1", "1")]) assert store.stored_model_name() == "model-a" @@ -186,7 +202,8 @@ class TestModelNameTracking: def test_config_mismatch_semantics(self, tmp_path: Path) -> None: store = PaperlessSqliteVecVectorStore( - uri=str(tmp_path), embed_model_name="model-a", + uri=str(tmp_path), + embed_model_name="model-a", ) assert not store.config_mismatch("anything") # no table yet store.add([make_node("a1", "1")]) @@ -194,7 +211,8 @@ class TestModelNameTracking: assert store.config_mismatch("model-b") def test_config_mismatch_false_when_table_predates_tracking( - self, tmp_path: Path, + self, + tmp_path: Path, ) -> None: store = PaperlessSqliteVecVectorStore(uri=str(tmp_path)) # no model name store.add([make_node("a1", "1")]) @@ -235,7 +253,8 @@ class TestCompact: def _churn(self, store, cycles: int) -> None: for i in range(cycles): store.upsert_document( - "1", [make_node(f"gen{i}-{j}", "1", seed=float(j)) for j in range(20)], + "1", + [make_node(f"gen{i}-{j}", "1", seed=float(j)) for j in range(20)], ) def test_compact_noop_below_threshold(self, store) -> None: @@ -247,11 +266,13 @@ class TestCompact: store.add([make_node("a1", "1"), make_node("b1", "2", seed=3.0)]) self._churn(store, 5) before = { - n.node_id: n.metadata for n in store.get_nodes(filters=_in_filter(["1", "2"])) + n.node_id: n.metadata + for n in store.get_nodes(filters=_in_filter(["1", "2"])) } store.compact(force=True) after = { - n.node_id: n.metadata for n in store.get_nodes(filters=_in_filter(["1", "2"])) + n.node_id: n.metadata + for n in store.get_nodes(filters=_in_filter(["1", "2"])) } assert after == before assert self._bloat_ratio(store) == pytest.approx(1.0) diff --git a/src/paperless_ai/vector_store.py b/src/paperless_ai/vector_store.py index 378c39e26..67e3f6cfc 100644 --- a/src/paperless_ai/vector_store.py +++ b/src/paperless_ai/vector_store.py @@ -353,7 +353,10 @@ class PaperlessSqliteVecVectorStore(BasePydanticVectorStore): # vec0 returns None distance when the query embedding is the zero vector # (no meaningful cosine angle); treat that as maximum distance (1.0) so # the row is included but ranked last. - sims = [1.0 - float(row["distance"] if row["distance"] is not None else 1.0) for row in rows] + sims = [ + 1.0 - float(row["distance"] if row["distance"] is not None else 1.0) + for row in rows + ] ids = [row["id"] for row in rows] return VectorStoreQueryResult(nodes=nodes, similarities=sims, ids=ids) @@ -442,7 +445,16 @@ class PaperlessSqliteVecVectorStore(BasePydanticVectorStore): f"INSERT INTO {self._table_name} " f"(id, document_id, modified, node_content, embedding) " f"VALUES (?, ?, ?, ?, ?)", - [(r["id"], r["document_id"], r["modified"], r["node_content"], bytes(r["embedding"])) for r in rows], + [ + ( + r["id"], + r["document_id"], + r["modified"], + r["node_content"], + bytes(r["embedding"]), + ) + for r in rows + ], ) # Reset the cumulative counter: after compact, total_inserts == live. new_conn.execute(