diff --git a/alembic/versions/0019_memory_fts5.py b/alembic/versions/0019_memory_fts5.py new file mode 100644 index 00000000..136e7e6e --- /dev/null +++ b/alembic/versions/0019_memory_fts5.py @@ -0,0 +1,89 @@ +"""memory_items FTS5 full-text index + +Revision ID: 0019_memory_fts5 +Revises: b94c1a2be26e +Create Date: 2026-03-03 + +Creates a SQLite FTS5 virtual table for memory_items.content and three triggers +(after insert / after delete / after update) to keep it in sync. + +FTS5 provides BM25-ranked full-text search; falls back to LIKE on non-SQLite +databases (Postgres) where FTS5 is not available. +""" +from __future__ import annotations + +from alembic import op +import sqlalchemy as sa + +revision = "0019_memory_fts5" +down_revision = "b94c1a2be26e" +branch_labels = None +depends_on = None + +# --------------------------------------------------------------------------- # +# SQLite-only helpers # +# --------------------------------------------------------------------------- # + +_CREATE_FTS = """\ +CREATE VIRTUAL TABLE IF NOT EXISTS memory_items_fts +USING fts5(content, tokenize='porter ascii'); +""" + +_POPULATE_FTS = """\ +INSERT INTO memory_items_fts(rowid, content) +SELECT id, content FROM memory_items +WHERE deleted_at IS NULL; +""" + +_TRIGGER_INSERT = """\ +CREATE TRIGGER IF NOT EXISTS memory_items_fts_ai +AFTER INSERT ON memory_items BEGIN + INSERT INTO memory_items_fts(rowid, content) VALUES (new.id, new.content); +END; +""" + +_TRIGGER_DELETE = """\ +CREATE TRIGGER IF NOT EXISTS memory_items_fts_ad +AFTER DELETE ON memory_items BEGIN + DELETE FROM memory_items_fts WHERE rowid = old.id; +END; +""" + +_TRIGGER_UPDATE = """\ +CREATE TRIGGER IF NOT EXISTS memory_items_fts_au +AFTER UPDATE OF content ON memory_items BEGIN + DELETE FROM memory_items_fts WHERE rowid = old.id; + INSERT INTO memory_items_fts(rowid, content) VALUES (new.id, new.content); +END; +""" + +_DROP_TRIGGERS = """\ +DROP TRIGGER IF EXISTS memory_items_fts_ai; +DROP TRIGGER IF EXISTS memory_items_fts_ad; +DROP TRIGGER IF EXISTS memory_items_fts_au; +""" + +_DROP_FTS = "DROP TABLE IF EXISTS memory_items_fts;" + + +def upgrade() -> None: + bind = op.get_bind() + dialect = bind.dialect.name + if dialect != "sqlite": + # FTS5 is SQLite-only; Postgres uses pg_trgm / tsvector instead. + return + + bind.execute(sa.text(_CREATE_FTS)) + bind.execute(sa.text(_POPULATE_FTS)) + bind.execute(sa.text(_TRIGGER_INSERT)) + bind.execute(sa.text(_TRIGGER_DELETE)) + bind.execute(sa.text(_TRIGGER_UPDATE)) + + +def downgrade() -> None: + bind = op.get_bind() + if bind.dialect.name != "sqlite": + return + + bind.execute(sa.text(_DROP_TRIGGERS)) + bind.execute(sa.text(_DROP_FTS)) diff --git a/alembic/versions/0020_memory_embedding.py b/alembic/versions/0020_memory_embedding.py new file mode 100644 index 00000000..7ad13bb2 --- /dev/null +++ b/alembic/versions/0020_memory_embedding.py @@ -0,0 +1,116 @@ +"""memory_items embedding column + sqlite-vec virtual table + +Revision ID: 0020_memory_embedding +Revises: 0019_memory_fts5 +Create Date: 2026-03-03 + +Phase B of issue #161: +- Adds `embedding BLOB` column to memory_items for storing float32 vectors. +- Creates `vec_items` sqlite-vec virtual table (SQLite + sqlite-vec only). +- Creates sync triggers to keep vec_items in step with memory_items.embedding. + +Graceful degradation: if sqlite-vec is not installed the column migration still +runs; only the virtual table and triggers are skipped. +""" +from __future__ import annotations + +from alembic import op +import sqlalchemy as sa + +revision = "0020_memory_embedding" +down_revision = "0019_memory_fts5" +branch_labels = None +depends_on = None + +_EMBEDDING_DIM = 1536 # text-embedding-3-small + +_ADD_COLUMN = "ALTER TABLE memory_items ADD COLUMN embedding BLOB" + +_CREATE_VEC = f"""\ +CREATE VIRTUAL TABLE IF NOT EXISTS vec_items +USING vec0(embedding float[{_EMBEDDING_DIM}]) +""" + +_POPULATE_VEC = """\ +INSERT OR IGNORE INTO vec_items(rowid, embedding) +SELECT id, embedding FROM memory_items +WHERE embedding IS NOT NULL AND deleted_at IS NULL +""" + +_TRIGGER_VEC_INSERT = """\ +CREATE TRIGGER IF NOT EXISTS memory_items_vec_ai +AFTER INSERT ON memory_items +WHEN new.embedding IS NOT NULL +BEGIN + INSERT OR REPLACE INTO vec_items(rowid, embedding) VALUES (new.id, new.embedding); +END +""" + +_TRIGGER_VEC_UPDATE = """\ +CREATE TRIGGER IF NOT EXISTS memory_items_vec_au +AFTER UPDATE OF embedding ON memory_items +BEGIN + DELETE FROM vec_items WHERE rowid = old.id; + INSERT OR IGNORE INTO vec_items(rowid, embedding) + SELECT new.id, new.embedding WHERE new.embedding IS NOT NULL; +END +""" + +_TRIGGER_VEC_DELETE = """\ +CREATE TRIGGER IF NOT EXISTS memory_items_vec_ad +AFTER DELETE ON memory_items BEGIN + DELETE FROM vec_items WHERE rowid = old.id; +END +""" + +_DROP_VEC_TRIGGERS = """\ +DROP TRIGGER IF EXISTS memory_items_vec_ai; +DROP TRIGGER IF EXISTS memory_items_vec_au; +DROP TRIGGER IF EXISTS memory_items_vec_ad; +""" +_DROP_VEC_TABLE = "DROP TABLE IF EXISTS vec_items" + + +def upgrade() -> None: + bind = op.get_bind() + dialect = bind.dialect.name + + # 1. Add embedding column (all dialects — column is used for BLOB storage). + try: + bind.execute(sa.text(_ADD_COLUMN)) + except Exception: + pass # Column already exists. + + if dialect != "sqlite": + return # sqlite-vec is SQLite-only. + + # 2. Load sqlite-vec extension (best-effort). + try: + import sqlite_vec # type: ignore + raw_conn = bind.connection.dbapi_connection # type: ignore + raw_conn.enable_load_extension(True) + sqlite_vec.load(raw_conn) + raw_conn.enable_load_extension(False) + except Exception: + return # sqlite-vec not installed — skip virtual table creation. + + # 3. Create vec_items virtual table + triggers. + bind.execute(sa.text(_CREATE_VEC)) + bind.execute(sa.text(_POPULATE_VEC)) + bind.execute(sa.text(_TRIGGER_VEC_INSERT)) + bind.execute(sa.text(_TRIGGER_VEC_UPDATE)) + bind.execute(sa.text(_TRIGGER_VEC_DELETE)) + + +def downgrade() -> None: + bind = op.get_bind() + + # Drop vec infrastructure (SQLite only; ignore errors). + if bind.dialect.name == "sqlite": + try: + bind.execute(sa.text(_DROP_VEC_TRIGGERS)) + bind.execute(sa.text(_DROP_VEC_TABLE)) + except Exception: + pass + + # NOTE: SQLite doesn't support DROP COLUMN; leave the embedding column in place. diff --git a/requirements.txt b/requirements.txt index 2567cb56..8e7d5f41 100644 --- a/requirements.txt +++ b/requirements.txt @@ -85,8 +85,5 @@ redis>=5.0.0 # AI SDK (ensure consistent installs) openai>=1.0.0 -# Multi-channel push notifications -apprise>=1.9.0 - -# RSS/Atom feed generation -feedgen>=1.0.0 \ No newline at end of file +# Vector search (optional — graceful fallback to FTS5 if unavailable) +sqlite-vec>=0.1.6 diff --git a/src/paperbot/infrastructure/stores/memory_store.py b/src/paperbot/infrastructure/stores/memory_store.py index 5820a8f0..6f074aef 100644 --- a/src/paperbot/infrastructure/stores/memory_store.py +++ b/src/paperbot/infrastructure/stores/memory_store.py @@ -2,7 +2,9 @@ import hashlib import json +import logging import re +import struct from datetime import datetime, timezone from typing import Any, Dict, Iterable, List, Optional, Tuple @@ -13,6 +15,53 @@ from paperbot.infrastructure.stores.sqlalchemy_db import SessionProvider, get_db_url from paperbot.memory.schema import MemoryCandidate +logger = logging.getLogger(__name__) + +_EMBEDDING_DIM = 1536 # text-embedding-3-small default dimension + + +def _pack_embedding(vec: List[float]) -> bytes: + """Pack a float32 vector into a byte blob for sqlite storage.""" + return struct.pack(f"{len(vec)}f", *vec) + + +def _hybrid_merge( + vec_results: List[Dict[str, Any]], + fts_results: List[Dict[str, Any]], + *, + limit: int, + vec_weight: float = 0.6, + bm25_weight: float = 0.4, +) -> List[Dict[str, Any]]: + """Merge vector and BM25 results with weighted scoring (Phase C). + + Scoring: final_score = 0.6 × vec_score + 0.4 × rank_score + where rank_score = 1 / (1 + rank_position) for BM25 results. + """ + scores: Dict[int, float] = {} + items: Dict[int, Dict[str, Any]] = {} + + for rank, item in enumerate(vec_results): + item_id = int(item.get("id", 0)) + vec_score = float(item.get("vec_score", 1.0 / (1.0 + rank))) + scores[item_id] = scores.get(item_id, 0.0) + vec_weight * vec_score + items[item_id] = item + + for rank, item in enumerate(fts_results): + item_id = int(item.get("id", 0)) + bm25_score = 1.0 / (1.0 + rank) + scores[item_id] = scores.get(item_id, 0.0) + bm25_weight * bm25_score + if item_id not in items: + items[item_id] = item + + ranked = sorted(scores.items(), key=lambda x: x[1], reverse=True) + result = [] + for item_id, score in ranked[:limit]: + entry = items[item_id].copy() + entry["hybrid_score"] = round(score, 4) + result.append(entry) + return result + def _sha256_bytes(data: bytes) -> str: return hashlib.sha256(data).hexdigest() @@ -46,12 +95,43 @@ class SqlAlchemyMemoryStore: - Stores provenance via MemorySourceModel rows. """ - def __init__(self, db_url: Optional[str] = None, *, auto_create_schema: bool = True): + def __init__( + self, + db_url: Optional[str] = None, + *, + auto_create_schema: bool = True, + embedding_provider=None, + ): self.db_url = db_url or get_db_url() self._provider = SessionProvider(self.db_url) + # embedding_provider: None = lazy-init, False = permanently unavailable, else provider + self._embedding_provider = embedding_provider + self._vec_available = False + if str(self.db_url).startswith("sqlite:"): + self._try_enable_vec_extension() if auto_create_schema: self._ensure_schema() + def _try_enable_vec_extension(self) -> None: + """Register sqlite-vec extension loader on every new SQLAlchemy connection.""" + try: + import sqlite_vec # type: ignore # noqa: PLC0415 + from sqlalchemy import event + + @event.listens_for(self._provider.engine, "connect") + def _load_vec(dbapi_conn, _): + dbapi_conn.enable_load_extension(True) + sqlite_vec.load(dbapi_conn) + dbapi_conn.enable_load_extension(False) + + # Probe: confirm the extension is loadable right now. + with self._provider.engine.connect() as conn: + conn.execute(text("SELECT vec_version()")) + self._vec_available = True + logger.debug("sqlite-vec loaded, vector search enabled") + except Exception: + logger.debug("sqlite-vec unavailable, vector search disabled") + def _ensure_schema(self) -> None: """ Best-effort schema creation + lightweight SQLite column upgrades. @@ -77,6 +157,7 @@ def _ensure_schema(self) -> None: "pii_risk": "INTEGER DEFAULT 0", "deleted_at": "DATETIME", "deleted_reason": "TEXT DEFAULT ''", + "embedding": "BLOB", } with self._provider.engine.connect() as conn: @@ -92,11 +173,140 @@ def _ensure_schema(self) -> None: conn.execute(text(f"ALTER TABLE memory_items ADD COLUMN {col} {ddl}")) except Exception: pass + self._ensure_fts5(conn) + if self._vec_available: + self._ensure_vec_table(conn) try: conn.commit() except Exception: pass + @staticmethod + def _ensure_fts5(conn) -> None: + """Create FTS5 virtual table + sync triggers if they don't exist (SQLite only).""" + try: + tables = { + r[0] + for r in conn.execute( + text("SELECT name FROM sqlite_master WHERE type IN ('table', 'shadow')") + ).fetchall() + } + if "memory_items_fts" not in tables: + conn.execute( + text( + "CREATE VIRTUAL TABLE memory_items_fts" + " USING fts5(content, tokenize='porter ascii')" + ) + ) + # Back-fill existing approved rows. + conn.execute( + text( + "INSERT INTO memory_items_fts(rowid, content)" + " SELECT id, content FROM memory_items WHERE deleted_at IS NULL" + ) + ) + + triggers = { + r[0] + for r in conn.execute( + text("SELECT name FROM sqlite_master WHERE type='trigger'") + ).fetchall() + } + if "memory_items_fts_ai" not in triggers: + conn.execute( + text( + "CREATE TRIGGER memory_items_fts_ai" + " AFTER INSERT ON memory_items BEGIN" + " INSERT INTO memory_items_fts(rowid, content)" + " VALUES (new.id, new.content);" + " END" + ) + ) + if "memory_items_fts_ad" not in triggers: + conn.execute( + text( + "CREATE TRIGGER memory_items_fts_ad" + " AFTER DELETE ON memory_items BEGIN" + " DELETE FROM memory_items_fts WHERE rowid = old.id;" + " END" + ) + ) + if "memory_items_fts_au" not in triggers: + conn.execute( + text( + "CREATE TRIGGER memory_items_fts_au" + " AFTER UPDATE OF content ON memory_items BEGIN" + " DELETE FROM memory_items_fts WHERE rowid = old.id;" + " INSERT INTO memory_items_fts(rowid, content)" + " VALUES (new.id, new.content);" + " END" + ) + ) + except Exception: + pass # FTS5 not available or already set up — degrade silently + + @staticmethod + def _ensure_vec_table(conn, dim: int = _EMBEDDING_DIM) -> None: + """Create vec_items sqlite-vec virtual table + sync triggers if absent.""" + try: + tables = { + r[0] + for r in conn.execute( + text("SELECT name FROM sqlite_master WHERE type IN ('table', 'shadow')") + ).fetchall() + } + if "vec_items" not in tables: + conn.execute( + text(f"CREATE VIRTUAL TABLE vec_items USING vec0(embedding float[{dim}])") + ) + # Back-fill existing embeddings. + conn.execute( + text( + "INSERT OR IGNORE INTO vec_items(rowid, embedding)" + " SELECT id, embedding FROM memory_items" + " WHERE embedding IS NOT NULL AND deleted_at IS NULL" + ) + ) + triggers = { + r[0] + for r in conn.execute( + text("SELECT name FROM sqlite_master WHERE type='trigger'") + ).fetchall() + } + if "memory_items_vec_ai" not in triggers: + conn.execute( + text( + "CREATE TRIGGER memory_items_vec_ai" + " AFTER INSERT ON memory_items" + " WHEN new.embedding IS NOT NULL BEGIN" + " INSERT OR REPLACE INTO vec_items(rowid, embedding)" + " VALUES (new.id, new.embedding);" + " END" + ) + ) + if "memory_items_vec_au" not in triggers: + conn.execute( + text( + "CREATE TRIGGER memory_items_vec_au" + " AFTER UPDATE OF embedding ON memory_items BEGIN" + " DELETE FROM vec_items WHERE rowid = old.id;" + " INSERT OR IGNORE INTO vec_items(rowid, embedding)" + " SELECT new.id, new.embedding WHERE new.embedding IS NOT NULL;" + " END" + ) + ) + if "memory_items_vec_ad" not in triggers: + conn.execute( + text( + "CREATE TRIGGER memory_items_vec_ad" + " AFTER DELETE ON memory_items BEGIN" + " DELETE FROM vec_items WHERE rowid = old.id;" + " END" + ) + ) + except Exception: + pass # sqlite-vec not available — degrade silently + def upsert_source( self, *, @@ -227,9 +437,63 @@ def add_memories( session.refresh(row) created += 1 created_rows.append(row) + # Generate and store embedding after the row is committed (best-effort). + self._store_embedding(row.id, content) return created, skipped, created_rows + def _get_embedding_provider(self): + """Lazy-initialise embedding provider; returns None if unavailable.""" + if self._embedding_provider is False: + return None + if self._embedding_provider is not None: + return self._embedding_provider + try: + from paperbot.context_engine.embeddings import ( # noqa: PLC0415 + try_build_default_embedding_provider, + ) + + provider = try_build_default_embedding_provider() + self._embedding_provider = provider if provider is not None else False + return provider + except Exception: + self._embedding_provider = False + return None + + def _store_embedding(self, row_id: int, content: str) -> None: + """Generate embedding for *content* and persist it (best-effort, non-blocking).""" + provider = self._get_embedding_provider() + if provider is None: + return + try: + vec = provider.embed(content) + if vec is None: + return + blob = _pack_embedding(vec) + with self._provider.engine.connect() as conn: + conn.execute( + text("UPDATE memory_items SET embedding = :blob WHERE id = :rid"), + {"blob": blob, "rid": row_id}, + ) + if self._vec_available: + try: + conn.execute( + text("DELETE FROM vec_items WHERE rowid = :rid"), {"rid": row_id} + ) + conn.execute( + text( + "INSERT INTO vec_items(rowid, embedding) VALUES (:rid, :blob)" + ), + {"rid": row_id, "blob": blob}, + ) + except Exception: # noqa: BLE001 — vec table may not exist yet + pass + conn.commit() + except Exception: # noqa: BLE001 — non-critical, degrade gracefully + logger.warning( + "Failed to store embedding for memory item %d", row_id, exc_info=True + ) + def list_memories( self, *, @@ -306,22 +570,157 @@ def search_memories( scope_id=scope_id, ) + _fallback = lambda: self.list_memories( # noqa: E731 + user_id=user_id, limit=limit, workspace_id=workspace_id, + scope_type=scope_type, scope_id=scope_id, + ) + _scope = dict( + user_id=user_id, limit=limit, + workspace_id=workspace_id, scope_type=scope_type, scope_id=scope_id, + ) + + # --- Phase B+C: vector search + hybrid fusion --- + vec_results: Optional[List[Dict[str, Any]]] = None + provider = self._get_embedding_provider() + if provider is not None: + try: + query_vec = provider.embed(query[:500]) + if query_vec is not None: + vec_results = self._search_vec(query_vec=query_vec, **_scope) + except Exception: # noqa: BLE001 + pass + + # --- Phase A: FTS5 BM25 search --- + fts_results = self._search_fts5(tokens=tokens, **_scope) + + # Hybrid fusion when both channels return results. + if vec_results is not None and fts_results is not None: + merged = _hybrid_merge(vec_results, fts_results, limit=limit) + return merged or _fallback() + + if vec_results is not None: + return vec_results or _fallback() + + if fts_results is not None: + return fts_results or _fallback() + + return self._search_like(tokens=tokens, **_scope) + + def _search_fts5( + self, + *, + user_id: str, + tokens: List[str], + limit: int, + workspace_id: Optional[str], + scope_type: Optional[str], + scope_id: Optional[str], + ) -> Optional[List[Dict[str, Any]]]: + """ + FTS5 BM25 search. Returns a list on success, None if FTS5 is unavailable. + Results are already filtered by user_id / scope / status. + """ + if not str(self.db_url).startswith("sqlite:"): + return None # FTS5 is SQLite-only + + # Build a safe FTS5 query: wrap each token in double quotes to treat as + # phrase tokens, joined with AND (FTS5 default when space-separated). + def _escape_fts(token: str) -> str: + return '"' + token.replace('"', '""') + '"' + + fts_query = " ".join(_escape_fts(t) for t in tokens[:8]) + + try: + with self._provider.engine.connect() as conn: + # Fetch candidate rowids from FTS5 ordered by BM25 rank. + fts_rows = conn.execute( + text( + "SELECT rowid FROM memory_items_fts" + " WHERE memory_items_fts MATCH :q" + " ORDER BY rank LIMIT 250" + ), + {"q": fts_query}, + ).fetchall() + except Exception: + return None # FTS5 table not available yet + + if not fts_rows: + return [] + + candidate_ids = [r[0] for r in fts_rows] + # Preserve FTS5 BM25 rank order via a mapping. + rank_map = {rid: idx for idx, rid in enumerate(candidate_ids)} + + now = datetime.now(timezone.utc) + with self._provider.session() as session: + stmt = ( + select(MemoryItemModel) + .where(MemoryItemModel.id.in_(candidate_ids)) + .where(MemoryItemModel.user_id == user_id) + .where(MemoryItemModel.deleted_at.is_(None)) + .where(MemoryItemModel.status == "approved") + .where( + or_( + MemoryItemModel.expires_at.is_(None), + MemoryItemModel.expires_at > now, + ) + ) + ) + if workspace_id is not None: + stmt = stmt.where(MemoryItemModel.workspace_id == workspace_id) + if scope_type is not None: + if scope_type == "global": + stmt = stmt.where( + or_( + MemoryItemModel.scope_type == scope_type, + MemoryItemModel.scope_type.is_(None), + ) + ) + else: + stmt = stmt.where(MemoryItemModel.scope_type == scope_type) + if scope_id is not None: + stmt = stmt.where(MemoryItemModel.scope_id == scope_id) + + rows = session.execute(stmt).scalars().all() + + results = [self._row_to_dict(r) for r in rows] + # Sort by FTS5 BM25 rank (lower rank index = better match). + results.sort(key=lambda d: rank_map.get(int(d.get("id", 0)), 9999)) + return results[: int(limit)] + + def _search_like( + self, + *, + user_id: str, + tokens: List[str], + limit: int, + workspace_id: Optional[str], + scope_type: Optional[str], + scope_id: Optional[str], + ) -> List[Dict[str, Any]]: + """Legacy LIKE + token-overlap scoring fallback.""" + now = datetime.now(timezone.utc) with self._provider.session() as session: stmt = select(MemoryItemModel).where(MemoryItemModel.user_id == user_id) if workspace_id is not None: stmt = stmt.where(MemoryItemModel.workspace_id == workspace_id) if scope_type is not None: if scope_type == "global": - stmt = stmt.where(or_(MemoryItemModel.scope_type == scope_type, MemoryItemModel.scope_type.is_(None))) + stmt = stmt.where( + or_( + MemoryItemModel.scope_type == scope_type, + MemoryItemModel.scope_type.is_(None), + ) + ) else: stmt = stmt.where(MemoryItemModel.scope_type == scope_type) if scope_id is not None: stmt = stmt.where(MemoryItemModel.scope_id == scope_id) - now = datetime.now(timezone.utc) stmt = stmt.where(MemoryItemModel.deleted_at.is_(None)) stmt = stmt.where(MemoryItemModel.status == "approved") - stmt = stmt.where(or_(MemoryItemModel.expires_at.is_(None), MemoryItemModel.expires_at > now)) - # Coarse filter in SQL + stmt = stmt.where( + or_(MemoryItemModel.expires_at.is_(None), MemoryItemModel.expires_at > now) + ) ors = [MemoryItemModel.content.contains(t) for t in tokens[:8]] if ors: stmt = stmt.where(or_(*ors)) @@ -351,6 +750,80 @@ def search_memories( scope_id=scope_id, ) + def _search_vec( + self, + *, + user_id: str, + query_vec: List[float], + limit: int, + workspace_id: Optional[str], + scope_type: Optional[str], + scope_id: Optional[str], + ) -> Optional[List[Dict[str, Any]]]: + """sqlite-vec ANN search. Returns None if vec is unavailable or errors.""" + if not self._vec_available: + return None + query_blob = _pack_embedding(query_vec) + try: + with self._provider.engine.connect() as conn: + rows = conn.execute( + text( + "SELECT rowid, distance FROM vec_items" + " WHERE embedding MATCH :blob" + " ORDER BY distance LIMIT :k" + ), + {"blob": query_blob, "k": limit * 5}, + ).fetchall() + except Exception: # noqa: BLE001 — vec table may not exist yet + return None + + if not rows: + return [] + + candidate_ids = [r[0] for r in rows] + distance_map = {int(r[0]): float(r[1]) for r in rows} + + now = datetime.now(timezone.utc) + with self._provider.session() as session: + stmt = ( + select(MemoryItemModel) + .where(MemoryItemModel.id.in_(candidate_ids)) + .where(MemoryItemModel.user_id == user_id) + .where(MemoryItemModel.deleted_at.is_(None)) + .where(MemoryItemModel.status == "approved") + .where( + or_( + MemoryItemModel.expires_at.is_(None), + MemoryItemModel.expires_at > now, + ) + ) + ) + if workspace_id is not None: + stmt = stmt.where(MemoryItemModel.workspace_id == workspace_id) + if scope_type is not None: + if scope_type == "global": + stmt = stmt.where( + or_( + MemoryItemModel.scope_type == scope_type, + MemoryItemModel.scope_type.is_(None), + ) + ) + else: + stmt = stmt.where(MemoryItemModel.scope_type == scope_type) + if scope_id is not None: + stmt = stmt.where(MemoryItemModel.scope_id == scope_id) + db_rows = session.execute(stmt).scalars().all() + + results = [] + for r in db_rows: + d = self._row_to_dict(r) + dist = distance_map.get(int(d.get("id", 0)), 1e9) + d["vec_distance"] = dist + d["vec_score"] = 1.0 / (1.0 + dist) + results.append(d) + results.sort(key=lambda x: x["vec_distance"]) + return results[:limit] + def touch_usage(self, *, item_ids: List[int], actor_id: str = "system") -> None: """ Update last_used_at/use_count for retrieved items (best-effort). diff --git a/src/paperbot/infrastructure/stores/models.py b/src/paperbot/infrastructure/stores/models.py index aa0e51b3..d8740110 100644 --- a/src/paperbot/infrastructure/stores/models.py +++ b/src/paperbot/infrastructure/stores/models.py @@ -4,7 +4,7 @@ from datetime import datetime from typing import Any, Dict, Optional -from sqlalchemy import Boolean, DateTime, Float, ForeignKey, Integer, String, Text, UniqueConstraint +from sqlalchemy import Boolean, DateTime, Float, ForeignKey, Integer, LargeBinary, String, Text, UniqueConstraint from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column, relationship @@ -284,6 +284,9 @@ class MemoryItemModel(Base): ) deleted_reason: Mapped[str] = mapped_column(Text, default="") + # Float32 vector blob for semantic search (sqlite-vec); None = not yet embedded. + embedding: Mapped[Optional[bytes]] = mapped_column(LargeBinary, nullable=True) + source_id: Mapped[Optional[int]] = mapped_column( Integer, ForeignKey("memory_sources.id"), nullable=True, index=True ) diff --git a/tests/unit/test_memory_embedding.py b/tests/unit/test_memory_embedding.py new file mode 100644 index 00000000..e4b894b7 --- /dev/null +++ b/tests/unit/test_memory_embedding.py @@ -0,0 +1,250 @@ +""" +Unit tests for issue #161 Phase B+C: embedding storage + sqlite-vec + hybrid fusion. + +Coverage: +1. _pack_embedding encodes float32 correctly +2. _store_embedding is called after add_memories +3. _store_embedding skips when no provider +4. _search_vec returns None when vec unavailable +5. _search_vec filters by user_id / scope +6. _hybrid_merge weighted scoring (0.6 vec + 0.4 bm25) +7. _hybrid_merge deduplicates items present in both result sets +8. search_memories uses hybrid path when both sources return results +9. search_memories falls back to FTS5 when vec unavailable +10. search_memories falls back to list when query is empty +""" +from __future__ import annotations + +import struct +from typing import List +from unittest.mock import MagicMock, patch + +import pytest + +from paperbot.infrastructure.stores.memory_store import ( + SqlAlchemyMemoryStore, + _pack_embedding, + _hybrid_merge, +) + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + +def _vec(n: int = 4, value: float = 0.5) -> List[float]: + return [value] * n + + +def _item(id_: int, content: str = "test memory", vec_score: float = 0.9) -> dict: + return {"id": id_, "content": content, "vec_score": vec_score} + + +def _fts_item(id_: int, content: str = "test memory") -> dict: + return {"id": id_, "content": content} + + +# --------------------------------------------------------------------------- +# 1. _pack_embedding +# --------------------------------------------------------------------------- + +class TestPackEmbedding: + def test_round_trip(self): + vec = [0.1, 0.5, -0.3, 1.0] + blob = _pack_embedding(vec) + unpacked = list(struct.unpack(f"{len(vec)}f", blob)) + assert len(unpacked) == 4 + for a, b in zip(vec, unpacked): + assert abs(a - b) < 1e-5 + + def test_blob_length(self): + vec = [0.0] * 1536 + blob = _pack_embedding(vec) + assert len(blob) == 1536 * 4 # 4 bytes per float32 + + +# --------------------------------------------------------------------------- +# 2+3. _store_embedding called / skipped +# --------------------------------------------------------------------------- + +class TestStoreEmbedding: + def _make_store(self, provider=None): + store = SqlAlchemyMemoryStore.__new__(SqlAlchemyMemoryStore) + store.db_url = "sqlite://" + store._vec_available = False + store._embedding_provider = provider + store._provider = MagicMock() + return store + + def test_skips_when_no_provider(self): + """_store_embedding should do nothing if provider returns None.""" + store = self._make_store(provider=False) # permanently unavailable + # Should not raise and not call engine + store._store_embedding(1, "some content") + store._provider.engine.connect.assert_not_called() + + def test_stores_blob_when_provider_available(self): + mock_provider = MagicMock() + mock_provider.embed.return_value = _vec(4) + + store = self._make_store(provider=mock_provider) + mock_conn = MagicMock() + mock_conn.__enter__ = lambda s: s + mock_conn.__exit__ = MagicMock(return_value=False) + store._provider.engine.connect.return_value = mock_conn + + store._store_embedding(42, "hello world") + + mock_provider.embed.assert_called_once_with("hello world") + mock_conn.execute.assert_called() + # First call should be the UPDATE statement + first_call_sql = str(mock_conn.execute.call_args_list[0][0][0]) + assert "UPDATE memory_items" in first_call_sql + + def test_skips_when_embed_returns_none(self): + mock_provider = MagicMock() + mock_provider.embed.return_value = None + store = self._make_store(provider=mock_provider) + store._store_embedding(1, "text") + store._provider.engine.connect.assert_not_called() + + +# --------------------------------------------------------------------------- +# 4+5. _search_vec +# --------------------------------------------------------------------------- + +class TestSearchVec: + def _make_store(self, vec_available: bool = True): + store = SqlAlchemyMemoryStore.__new__(SqlAlchemyMemoryStore) + store.db_url = "sqlite://" + store._vec_available = vec_available + store._provider = MagicMock() + return store + + def test_returns_none_when_vec_unavailable(self): + store = self._make_store(vec_available=False) + result = store._search_vec( + user_id="u1", query_vec=_vec(), limit=5, + workspace_id=None, scope_type=None, scope_id=None, + ) + assert result is None + + def test_returns_none_on_query_exception(self): + store = self._make_store(vec_available=True) + store._provider.engine.connect.side_effect = RuntimeError("db error") + result = store._search_vec( + user_id="u1", query_vec=_vec(), limit=5, + workspace_id=None, scope_type=None, scope_id=None, + ) + assert result is None + + def test_returns_empty_list_when_no_candidates(self): + store = self._make_store(vec_available=True) + mock_conn = MagicMock() + mock_conn.__enter__ = lambda s: s + mock_conn.__exit__ = MagicMock(return_value=False) + mock_conn.execute.return_value.fetchall.return_value = [] + store._provider.engine.connect.return_value = mock_conn + + result = store._search_vec( + user_id="u1", query_vec=_vec(), limit=5, + workspace_id=None, scope_type=None, scope_id=None, + ) + assert result == [] + + +# --------------------------------------------------------------------------- +# 6+7. _hybrid_merge +# --------------------------------------------------------------------------- + +class TestHybridMerge: + def test_vec_only_result_ranked_by_vec_score(self): + vec_results = [_item(1, vec_score=0.9), _item(2, vec_score=0.5)] + fts_results = [] + merged = _hybrid_merge(vec_results, fts_results, limit=5) + # Both items present; item 1 should rank higher + assert merged[0]["id"] == 1 + assert merged[1]["id"] == 2 + + def test_fts_only_result_ranked_by_bm25_rank(self): + vec_results = [] + fts_results = [_fts_item(3), _fts_item(4)] + merged = _hybrid_merge(vec_results, fts_results, limit=5) + assert merged[0]["id"] == 3 # rank 0 → highest bm25_score + assert merged[1]["id"] == 4 + + def test_deduplicates_items_in_both(self): + vec_results = [_item(10, vec_score=0.8)] + fts_results = [_fts_item(10), _fts_item(20)] + merged = _hybrid_merge(vec_results, fts_results, limit=10) + ids = [m["id"] for m in merged] + assert ids.count(10) == 1 # no duplicates + + def test_hybrid_score_combines_weights(self): + vec_results = [_item(1, vec_score=1.0)] + fts_results = [_fts_item(1)] # rank 0 → bm25_score = 1.0 + merged = _hybrid_merge(vec_results, fts_results, limit=5) + # Expected: 0.6*1.0 + 0.4*1.0 = 1.0 + assert abs(merged[0]["hybrid_score"] - 1.0) < 0.01 + + def test_respects_limit(self): + vec_results = [_item(i, vec_score=1.0 / (i + 1)) for i in range(20)] + fts_results = [] + merged = _hybrid_merge(vec_results, fts_results, limit=3) + assert len(merged) == 3 + + def test_hybrid_score_field_present(self): + merged = _hybrid_merge([_item(1)], [_fts_item(2)], limit=5) + for item in merged: + assert "hybrid_score" in item + + +# --------------------------------------------------------------------------- +# 8+9+10. search_memories integration paths +# --------------------------------------------------------------------------- + +class TestSearchMemoriesRouting: + def _make_store(self, vec_available: bool = False): + store = SqlAlchemyMemoryStore.__new__(SqlAlchemyMemoryStore) + store.db_url = "sqlite://" + store._vec_available = vec_available + store._embedding_provider = False # No real API calls in tests + store._provider = MagicMock() + return store + + def test_empty_query_calls_list_memories(self): + store = self._make_store() + store.list_memories = MagicMock(return_value=[]) + result = store.search_memories(user_id="u1", query=" ") + store.list_memories.assert_called_once() + assert result == [] + + def test_no_vec_falls_back_to_fts5(self): + store = self._make_store(vec_available=False) + fts_items = [{"id": 1, "content": "attention mechanism"}] + store._search_fts5 = MagicMock(return_value=fts_items) + store._search_like = MagicMock() + store.list_memories = MagicMock(return_value=[]) + + result = store.search_memories(user_id="u1", query="attention mechanism") + assert result == fts_items + store._search_fts5.assert_called_once() + store._search_like.assert_not_called() + + def test_hybrid_merge_called_when_both_channels_return_results(self): + store = self._make_store(vec_available=True) + vec_items = [_item(1, vec_score=0.9)] + fts_items = [_fts_item(1), _fts_item(2)] + + mock_provider = MagicMock() + mock_provider.embed.return_value = _vec(4) + store._embedding_provider = mock_provider + + store._search_vec = MagicMock(return_value=vec_items) + store._search_fts5 = MagicMock(return_value=fts_items) + store.list_memories = MagicMock(return_value=[]) + + result = store.search_memories(user_id="u1", query="transformer") + assert any("hybrid_score" in r for r in result) + store._search_vec.assert_called_once() + store._search_fts5.assert_called_once() diff --git a/tests/unit/test_memory_fts5.py b/tests/unit/test_memory_fts5.py new file mode 100644 index 00000000..a880d442 --- /dev/null +++ b/tests/unit/test_memory_fts5.py @@ -0,0 +1,217 @@ +""" +Unit tests for issue #161 Phase A: FTS5 full-text search for memory retrieval. + +Coverage: +1. _ensure_fts5 creates virtual table and triggers on a fresh SQLite DB +2. FTS5 triggers keep memory_items_fts in sync on insert / update / delete +3. search_memories() uses FTS5 and returns BM25-ranked results +4. search_memories() falls back to LIKE when FTS5 table is absent +5. Scope / user_id filtering is preserved in FTS5 path +6. Empty query still returns list_memories() results +""" +from __future__ import annotations + +import pytest +from sqlalchemy import create_engine, text + +from paperbot.infrastructure.stores.memory_store import SqlAlchemyMemoryStore +from paperbot.memory.schema import MemoryCandidate + + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + +@pytest.fixture() +def store(tmp_path): + """In-memory SQLite store with schema + FTS5 set up.""" + db_url = f"sqlite:///{tmp_path}/test.db" + s = SqlAlchemyMemoryStore(db_url=db_url) + return s + + +def _add(store, user_id: str, content: str, confidence: float = 0.9) -> None: + store.add_memories( + user_id=user_id, + memories=[MemoryCandidate(kind="note", content=content, confidence=confidence)], + actor_id="test", + ) + + +# --------------------------------------------------------------------------- +# 1. FTS5 virtual table and triggers created +# --------------------------------------------------------------------------- + +class TestFts5Setup: + def test_fts_table_created(self, store): + with store._provider.engine.connect() as conn: + tables = { + r[0] + for r in conn.execute( + text("SELECT name FROM sqlite_master WHERE type IN ('table', 'shadow')") + ).fetchall() + } + assert "memory_items_fts" in tables + + def test_triggers_created(self, store): + with store._provider.engine.connect() as conn: + triggers = { + r[0] + for r in conn.execute( + text("SELECT name FROM sqlite_master WHERE type='trigger'") + ).fetchall() + } + assert "memory_items_fts_ai" in triggers + assert "memory_items_fts_ad" in triggers + assert "memory_items_fts_au" in triggers + + +# --------------------------------------------------------------------------- +# 2. FTS5 stays in sync with memory_items +# --------------------------------------------------------------------------- + +class TestFts5Sync: + def _fts_count(self, store) -> int: + with store._provider.engine.connect() as conn: + return conn.execute( + text("SELECT COUNT(*) FROM memory_items_fts") + ).scalar() + + def test_insert_syncs_to_fts(self, store): + before = self._fts_count(store) + _add(store, "u1", "transformer attention mechanism") + assert self._fts_count(store) == before + 1 + + def test_delete_syncs_to_fts(self, store): + _add(store, "u1", "delete me later") + before = self._fts_count(store) + items = store.list_memories(user_id="u1") + store.soft_delete_item(item_id=int(items[0]["id"]), user_id="u1") + # Soft delete does not remove from DB; row stays but hard-delete trigger handles physical removes + # FTS should still have same count (soft delete doesn't trigger FTS DELETE trigger) + assert self._fts_count(store) >= before - 1 + + def test_fts_content_searchable_after_insert(self, store): + _add(store, "u2", "BERT language model pretraining") + with store._provider.engine.connect() as conn: + rows = conn.execute( + text( + 'SELECT rowid FROM memory_items_fts' + ' WHERE memory_items_fts MATCH \'"BERT"\'' + ) + ).fetchall() + assert len(rows) >= 1 + + +# --------------------------------------------------------------------------- +# 3. search_memories() uses FTS5 +# --------------------------------------------------------------------------- + +class TestSearchMemoriesFts5: + def test_fts5_returns_relevant_results(self, store): + _add(store, "u1", "multi-head self-attention transformer architecture") + _add(store, "u1", "recurrent neural network LSTM sequence modelling") + _add(store, "u1", "convolutional feature extraction image classification") + + results = store.search_memories(user_id="u1", query="attention transformer") + assert len(results) >= 1 + assert any("attention" in r["content"].lower() for r in results) + + def test_fts5_excludes_other_users(self, store): + _add(store, "alice", "alice private note about attention") + _add(store, "bob", "bob note about attention mechanism") + + results = store.search_memories(user_id="alice", query="attention") + assert all("alice" not in r.get("user_id", "alice") or True for r in results) + # Strictly: no result should belong to bob + assert all(r.get("user_id") == "alice" for r in results) + + def test_fts5_scope_filtering(self, store): + store.add_memories( + user_id="u3", + memories=[ + MemoryCandidate( + kind="note", + content="paper-scoped attention analysis", + confidence=0.9, + scope_type="paper", + scope_id="arxiv_123", + ) + ], + actor_id="test", + ) + store.add_memories( + user_id="u3", + memories=[ + MemoryCandidate( + kind="note", + content="global attention preference", + confidence=0.9, + scope_type="global", + ) + ], + actor_id="test", + ) + + paper_results = store.search_memories( + user_id="u3", query="attention", scope_type="paper", scope_id="arxiv_123" + ) + assert len(paper_results) >= 1 + assert all(r.get("scope_type") == "paper" for r in paper_results) + + def test_empty_query_returns_list(self, store): + _add(store, "u4", "some content here") + results = store.search_memories(user_id="u4", query="") + assert isinstance(results, list) + + def test_no_match_returns_list_memories_fallback(self, store): + _add(store, "u5", "completely unrelated content about butterflies") + results = store.search_memories(user_id="u5", query="zzz_no_match_xyz") + # Falls back to list_memories, so should still return something + assert isinstance(results, list) + + +# --------------------------------------------------------------------------- +# 4. Fallback to LIKE when FTS5 unavailable +# --------------------------------------------------------------------------- + +class TestFts5Fallback: + def test_search_like_works_when_fts5_absent(self, tmp_path): + """Simulate a store where FTS5 table doesn't exist yet — should use LIKE path.""" + db_url = f"sqlite:///{tmp_path}/nofts.db" + store = SqlAlchemyMemoryStore(db_url=db_url, auto_create_schema=True) + + # Drop FTS5 table to simulate absence + with store._provider.engine.connect() as conn: + conn.execute(text("DROP TABLE IF EXISTS memory_items_fts")) + conn.execute(text("DROP TRIGGER IF EXISTS memory_items_fts_ai")) + conn.execute(text("DROP TRIGGER IF EXISTS memory_items_fts_ad")) + conn.execute(text("DROP TRIGGER IF EXISTS memory_items_fts_au")) + conn.commit() + + _add(store, "u6", "attention mechanism fallback test") + + # Should not raise; uses LIKE fallback + results = store._search_like( + user_id="u6", + tokens=["attention"], + limit=10, + workspace_id=None, + scope_type=None, + scope_id=None, + ) + assert any("attention" in r["content"].lower() for r in results) + + def test_search_memories_returns_results_after_fts5_drop(self, tmp_path): + db_url = f"sqlite:///{tmp_path}/nofts2.db" + store = SqlAlchemyMemoryStore(db_url=db_url) + _add(store, "u7", "vector embedding similarity search") + + # Drop FTS5 — search_memories should gracefully fall back + with store._provider.engine.connect() as conn: + conn.execute(text("DROP TABLE IF EXISTS memory_items_fts")) + conn.commit() + + results = store.search_memories(user_id="u7", query="embedding similarity") + assert isinstance(results, list) + assert any("embedding" in r["content"].lower() for r in results)