diff --git a/src/paperless_ai/tables.py b/src/paperless_ai/tables.py new file mode 100644 index 000000000..1a25e1e0b --- /dev/null +++ b/src/paperless_ai/tables.py @@ -0,0 +1,221 @@ +"""Thin gateways over the plain relational side tables that sit alongside the +vec0 table. Each method takes the sqlite3.Connection to operate on +explicitly, rather than owning one -- the store swaps connections during +compact()/migration, and migrations always work across two connections +(src_conn, dst_conn) at once. +""" + +import sqlite3 +from collections.abc import Iterable +from typing import NamedTuple + + +class ChunkRow(NamedTuple): + chunk_id: str + document_id: int + + +class DocumentMetaRow(NamedTuple): + document_id: int + modified: str + + +class DocumentChunksTable: + """chunk_id -> document_id, indexed by document_id. Gives O(1) + per-document chunk lookup that vec0's own document_id metadata column + cannot (see PaperlessSqliteVecVectorStore._delete_chunks_by_document_id). + """ + + @staticmethod + def create(conn: sqlite3.Connection) -> None: + conn.execute( + "CREATE TABLE IF NOT EXISTS document_chunks " + "(chunk_id TEXT PRIMARY KEY, document_id INTEGER NOT NULL)", + ) + conn.execute( + "CREATE INDEX IF NOT EXISTS idx_document_chunks_document_id " + "ON document_chunks (document_id)", + ) + + @staticmethod + def insert_many(conn: sqlite3.Connection, rows: Iterable[ChunkRow]) -> None: + """rows must already be batch-bounded by the caller (e.g. vec0's own + fetchmany() loop) -- this never reads, so it can't itself introduce + an unbounded scan, but a whole-table iterable defeats the point.""" + conn.executemany( + "INSERT INTO document_chunks (chunk_id, document_id) VALUES (?, ?)", + rows, + ) + + @staticmethod + def chunk_ids_for_document( + conn: sqlite3.Connection, + document_id: int, + ) -> list[str]: + return [ + row["chunk_id"] + for row in conn.execute( + "SELECT chunk_id FROM document_chunks WHERE document_id = ?", + (document_id,), + ).fetchall() + ] + + @staticmethod + def delete_for_document(conn: sqlite3.Connection, document_id: int) -> None: + conn.execute( + "DELETE FROM document_chunks WHERE document_id = ?", + (document_id,), + ) + + @staticmethod + def delete_all(conn: sqlite3.Connection) -> None: + conn.execute("DELETE FROM document_chunks") + + @staticmethod + def count(conn: sqlite3.Connection) -> int: + """Cheap stand-in for vec0's own row count -- see compact().""" + return conn.execute("SELECT count(*) FROM document_chunks").fetchone()[0] + + +class DocumentMetaTable: + """document_id -> modified, one row per document. Lives outside vec0 + because vec0 only inlines TEXT metadata up to 12 bytes and `modified` + (an ISO timestamp) is always longer. + """ + + @staticmethod + def create(conn: sqlite3.Connection) -> None: + conn.execute( + "CREATE TABLE IF NOT EXISTS document_meta " + "(document_id INTEGER PRIMARY KEY, modified TEXT NOT NULL)", + ) + + @staticmethod + def upsert_many( + conn: sqlite3.Connection, + rows: Iterable[DocumentMetaRow], + ) -> None: + conn.executemany( + "INSERT INTO document_meta (document_id, modified) VALUES (?, ?) " + "ON CONFLICT(document_id) DO UPDATE SET modified = excluded.modified", + rows, + ) + + @staticmethod + def delete_for_document(conn: sqlite3.Connection, document_id: int) -> None: + conn.execute( + "DELETE FROM document_meta WHERE document_id = ?", + (document_id,), + ) + + @staticmethod + def delete_all(conn: sqlite3.Connection) -> None: + conn.execute("DELETE FROM document_meta") + + @staticmethod + def copy_all( + src_conn: sqlite3.Connection, + dst_conn: sqlite3.Connection, + batch_size: int, + ) -> None: + """Stream document_meta from src_conn into dst_conn in bounded + batches. The *only* sanctioned way to move this table across + connections (compact()/migrations) -- an unbounded fetchall here + would defeat the same OOM-avoidance the vec0 row copy already relies + on. batch_size has no default: forces the call site to think about + it (pass COMPACT_BATCH_SIZE).""" + cursor = src_conn.execute( + "SELECT document_id, modified FROM document_meta", + ) + while batch := cursor.fetchmany(batch_size): + DocumentMetaTable.upsert_many( + dst_conn, + (DocumentMetaRow(r["document_id"], r["modified"]) for r in batch), + ) + + @staticmethod + def all_modified_times(conn: sqlite3.Connection) -> dict[str, str]: + """Full document_id -> modified map, for get_modified_times()'s + public API only. One unbounded read by design (existing behavior). + Never use this for cross-connection copying; see copy_all().""" + return { + str(row["document_id"]): str(row["modified"] or "") + for row in conn.execute( + "SELECT document_id, modified FROM document_meta", + ) + } + + +class IndexMetaTable: + """Typed accessors over index_meta's key/value rows -- replaces + PaperlessSqliteVecVectorStore._meta_get_on/_meta_set_on, which returned + untyped str | None regardless of whether the key held an int (dim, + schema_version, total_inserts) or a string (embed_model). + """ + + @staticmethod + def create(conn: sqlite3.Connection) -> None: + conn.execute( + "CREATE TABLE IF NOT EXISTS index_meta (key TEXT PRIMARY KEY, value TEXT)", + ) + + @staticmethod + def _get(conn: sqlite3.Connection, key: str) -> str | None: + row = conn.execute( + "SELECT value FROM index_meta WHERE key = ?", + (key,), + ).fetchone() + return row["value"] if row else None + + @staticmethod + def _set(conn: sqlite3.Connection, key: str, value: str) -> None: + conn.execute( + "INSERT INTO index_meta (key, value) VALUES (?, ?) " + "ON CONFLICT(key) DO UPDATE SET value = excluded.value", + (key, value), + ) + + @staticmethod + def get_dim(conn: sqlite3.Connection) -> int | None: + value = IndexMetaTable._get(conn, "dim") + return int(value) if value is not None else None + + @staticmethod + def set_dim(conn: sqlite3.Connection, dim: int) -> None: + IndexMetaTable._set(conn, "dim", str(dim)) + + @staticmethod + def get_embed_model(conn: sqlite3.Connection) -> str | None: + return IndexMetaTable._get(conn, "embed_model") + + @staticmethod + def set_embed_model(conn: sqlite3.Connection, name: str) -> None: + IndexMetaTable._set(conn, "embed_model", name) + + @staticmethod + def get_schema_version(conn: sqlite3.Connection) -> int | None: + value = IndexMetaTable._get(conn, "schema_version") + return int(value) if value is not None else None + + @staticmethod + def set_schema_version(conn: sqlite3.Connection, version: int) -> None: + IndexMetaTable._set(conn, "schema_version", str(version)) + + @staticmethod + def get_total_inserts(conn: sqlite3.Connection) -> int: + value = IndexMetaTable._get(conn, "total_inserts") + return int(value) if value is not None else 0 + + @staticmethod + def increment_total_inserts(conn: sqlite3.Connection, count: int) -> None: + current = IndexMetaTable.get_total_inserts(conn) + IndexMetaTable._set(conn, "total_inserts", str(current + count)) + + @staticmethod + def reset_total_inserts(conn: sqlite3.Connection, count: int) -> None: + """Set total_inserts to an absolute value -- distinct from + increment_total_inserts(): used by compact()'s rebuild and by + m0001_v1_to_v2 after copying live rows into a fresh file, where + total_inserts must become exactly the live row count, not add to + whatever the source file's counter held.""" + IndexMetaTable._set(conn, "total_inserts", str(count)) diff --git a/src/paperless_ai/tests/test_tables.py b/src/paperless_ai/tests/test_tables.py new file mode 100644 index 000000000..7ac99a185 --- /dev/null +++ b/src/paperless_ai/tests/test_tables.py @@ -0,0 +1,318 @@ +import sqlite3 +from collections.abc import Generator + +import pytest + +from paperless_ai.tables import ChunkRow +from paperless_ai.tables import DocumentChunksTable +from paperless_ai.tables import DocumentMetaRow +from paperless_ai.tables import DocumentMetaTable +from paperless_ai.tables import IndexMetaTable + + +@pytest.fixture +def conn() -> Generator[sqlite3.Connection, None, None]: + connection = sqlite3.connect(":memory:") + connection.row_factory = sqlite3.Row + try: + yield connection + finally: + connection.close() + + +class TestDocumentChunksTable: + def test_create_is_idempotent(self, conn: sqlite3.Connection) -> None: + """ + GIVEN: + - A bare sqlite3 connection + WHEN: + - create() is called twice + THEN: + - No error is raised and the table exists + """ + DocumentChunksTable.create(conn) + DocumentChunksTable.create(conn) + assert DocumentChunksTable.count(conn) == 0 + + def test_insert_many_then_lookup_by_document_id( + self, + conn: sqlite3.Connection, + ) -> None: + """ + GIVEN: + - An empty document_chunks table + WHEN: + - Two chunks for document 1 and one for document 2 are inserted + THEN: + - chunk_ids_for_document returns exactly the matching chunk ids + """ + DocumentChunksTable.create(conn) + DocumentChunksTable.insert_many( + conn, + [ChunkRow("c1", 1), ChunkRow("c2", 1), ChunkRow("c3", 2)], + ) + assert sorted(DocumentChunksTable.chunk_ids_for_document(conn, 1)) == [ + "c1", + "c2", + ] + assert DocumentChunksTable.chunk_ids_for_document(conn, 2) == ["c3"] + assert DocumentChunksTable.chunk_ids_for_document(conn, 999) == [] + + def test_delete_for_document_removes_only_that_document( + self, + conn: sqlite3.Connection, + ) -> None: + """ + GIVEN: + - Chunks for two different documents + WHEN: + - delete_for_document() is called for one of them + THEN: + - Only that document's chunks are removed + """ + DocumentChunksTable.create(conn) + DocumentChunksTable.insert_many( + conn, + [ChunkRow("c1", 1), ChunkRow("c2", 2)], + ) + DocumentChunksTable.delete_for_document(conn, 1) + assert DocumentChunksTable.chunk_ids_for_document(conn, 1) == [] + assert DocumentChunksTable.chunk_ids_for_document(conn, 2) == ["c2"] + + def test_delete_all_clears_every_row(self, conn: sqlite3.Connection) -> None: + """ + GIVEN: + - Chunks for multiple documents + WHEN: + - delete_all() is called + THEN: + - count() returns 0 + """ + DocumentChunksTable.create(conn) + DocumentChunksTable.insert_many( + conn, + [ChunkRow("c1", 1), ChunkRow("c2", 2)], + ) + DocumentChunksTable.delete_all(conn) + assert DocumentChunksTable.count(conn) == 0 + + def test_count_reflects_live_rows(self, conn: sqlite3.Connection) -> None: + """ + GIVEN: + - An empty document_chunks table + WHEN: + - Rows are inserted then one document's rows are deleted + THEN: + - count() reflects the remaining row count + """ + DocumentChunksTable.create(conn) + DocumentChunksTable.insert_many( + conn, + [ChunkRow("c1", 1), ChunkRow("c2", 1), ChunkRow("c3", 2)], + ) + assert DocumentChunksTable.count(conn) == 3 + DocumentChunksTable.delete_for_document(conn, 1) + assert DocumentChunksTable.count(conn) == 1 + + +class TestDocumentMetaTable: + def test_upsert_many_then_all_modified_times( + self, + conn: sqlite3.Connection, + ) -> None: + """ + GIVEN: + - An empty document_meta table + WHEN: + - Two documents' modified timestamps are upserted + THEN: + - all_modified_times() returns both, keyed by str(document_id) + """ + DocumentMetaTable.create(conn) + DocumentMetaTable.upsert_many( + conn, + [ + DocumentMetaRow(1, "2026-01-01T00:00:00"), + DocumentMetaRow(2, "2026-02-02T00:00:00"), + ], + ) + assert DocumentMetaTable.all_modified_times(conn) == { + "1": "2026-01-01T00:00:00", + "2": "2026-02-02T00:00:00", + } + + def test_upsert_many_overwrites_existing_value( + self, + conn: sqlite3.Connection, + ) -> None: + """ + GIVEN: + - A document_meta row for document 1 + WHEN: + - upsert_many() is called again with a new modified value for + the same document_id + THEN: + - The stored value is replaced, not duplicated + """ + DocumentMetaTable.create(conn) + DocumentMetaTable.upsert_many(conn, [DocumentMetaRow(1, "old")]) + DocumentMetaTable.upsert_many(conn, [DocumentMetaRow(1, "new")]) + assert DocumentMetaTable.all_modified_times(conn) == {"1": "new"} + + def test_delete_for_document_removes_only_that_row( + self, + conn: sqlite3.Connection, + ) -> None: + """ + GIVEN: + - document_meta rows for two documents + WHEN: + - delete_for_document() is called for one of them + THEN: + - Only that document's row is removed + """ + DocumentMetaTable.create(conn) + DocumentMetaTable.upsert_many( + conn, + [DocumentMetaRow(1, "a"), DocumentMetaRow(2, "b")], + ) + DocumentMetaTable.delete_for_document(conn, 1) + assert DocumentMetaTable.all_modified_times(conn) == {"2": "b"} + + def test_delete_all_clears_every_row(self, conn: sqlite3.Connection) -> None: + """ + GIVEN: + - document_meta rows for multiple documents + WHEN: + - delete_all() is called + THEN: + - all_modified_times() returns an empty dict + """ + DocumentMetaTable.create(conn) + DocumentMetaTable.upsert_many( + conn, + [DocumentMetaRow(1, "a"), DocumentMetaRow(2, "b")], + ) + DocumentMetaTable.delete_all(conn) + assert DocumentMetaTable.all_modified_times(conn) == {} + + def test_copy_all_streams_every_row_to_destination( + self, + conn: sqlite3.Connection, + ) -> None: + """ + GIVEN: + - A source connection with document_meta rows for 5 documents + - A separate, empty destination connection + WHEN: + - copy_all() is called with a batch size smaller than the row + count, forcing multiple fetchmany() cycles + THEN: + - Every row is present on the destination connection + """ + DocumentMetaTable.create(conn) + DocumentMetaTable.upsert_many( + conn, + [DocumentMetaRow(i, f"modified-{i}") for i in range(5)], + ) + dst_conn = sqlite3.connect(":memory:") + dst_conn.row_factory = sqlite3.Row + try: + DocumentMetaTable.create(dst_conn) + DocumentMetaTable.copy_all(conn, dst_conn, batch_size=2) + assert DocumentMetaTable.all_modified_times(dst_conn) == { + str(i): f"modified-{i}" for i in range(5) + } + finally: + dst_conn.close() + + +class TestIndexMetaTable: + def test_dim_roundtrip(self, conn: sqlite3.Connection) -> None: + """ + GIVEN: + - An empty index_meta table + WHEN: + - set_dim() is called then get_dim() is read back + THEN: + - The same int value is returned + """ + IndexMetaTable.create(conn) + assert IndexMetaTable.get_dim(conn) is None + IndexMetaTable.set_dim(conn, 384) + assert IndexMetaTable.get_dim(conn) == 384 + + def test_embed_model_roundtrip(self, conn: sqlite3.Connection) -> None: + """ + GIVEN: + - An empty index_meta table + WHEN: + - set_embed_model() is called then get_embed_model() is read back + THEN: + - The same string value is returned + """ + IndexMetaTable.create(conn) + assert IndexMetaTable.get_embed_model(conn) is None + IndexMetaTable.set_embed_model(conn, "model-a") + assert IndexMetaTable.get_embed_model(conn) == "model-a" + + def test_schema_version_roundtrip(self, conn: sqlite3.Connection) -> None: + """ + GIVEN: + - An empty index_meta table + WHEN: + - set_schema_version() is called then get_schema_version() is + read back + THEN: + - The same int value is returned + """ + IndexMetaTable.create(conn) + assert IndexMetaTable.get_schema_version(conn) is None + IndexMetaTable.set_schema_version(conn, 2) + assert IndexMetaTable.get_schema_version(conn) == 2 + + def test_total_inserts_starts_at_zero(self, conn: sqlite3.Connection) -> None: + """ + GIVEN: + - An empty index_meta table + WHEN: + - get_total_inserts() is read before anything is set + THEN: + - 0 is returned + """ + IndexMetaTable.create(conn) + assert IndexMetaTable.get_total_inserts(conn) == 0 + + def test_increment_total_inserts_accumulates( + self, + conn: sqlite3.Connection, + ) -> None: + """ + GIVEN: + - An empty index_meta table + WHEN: + - increment_total_inserts() is called twice + THEN: + - get_total_inserts() returns the running sum + """ + IndexMetaTable.create(conn) + IndexMetaTable.increment_total_inserts(conn, 5) + IndexMetaTable.increment_total_inserts(conn, 3) + assert IndexMetaTable.get_total_inserts(conn) == 8 + + def test_reset_total_inserts_sets_absolute_value( + self, + conn: sqlite3.Connection, + ) -> None: + """ + GIVEN: + - A total_inserts counter already at a high value + WHEN: + - reset_total_inserts() is called with a lower value + THEN: + - get_total_inserts() returns exactly that value, not a sum + """ + IndexMetaTable.create(conn) + IndexMetaTable.increment_total_inserts(conn, 100) + IndexMetaTable.reset_total_inserts(conn, 7) + assert IndexMetaTable.get_total_inserts(conn) == 7