feat(embeddings): per-chunk rows — schema, write path, version-aware backfill (#280 steps 2+3)
note_embeddings becomes one row per chunk: PK (note_id, chunk_index), plus chunk_text (what this vector actually encodes) and chunker_version. Migration 0077 clears the table — embeddings are derived (0067 precedent) and the old whole-document rows are indistinguishable from single-chunk notes, so the startup backfill regenerates the corpus at the new shape. The backfill is now version-aware: a future shape change is a CHUNKER_VERSION bump that re-embeds exactly the stale notes, not another wipe. upsert_note_embedding takes (title, body) and chunks internally — one path for the write path, the recurrence spawn and the backfill. The recurrence spawn's own embed call is deleted outright: create_note already embeds via embed_note (#2056), so the spawn was a second copy of the rule. An emptied record now CLEARS its stale vectors instead of leaving them findable. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01UaYUaouG9jjhATyuxCKrQs
This commit is contained in:
@@ -1,7 +1,7 @@
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from pgvector.sqlalchemy import Vector
|
||||
from sqlalchemy import DateTime, ForeignKey, Integer
|
||||
from sqlalchemy import DateTime, ForeignKey, Integer, Text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from scribe.models import Base
|
||||
@@ -14,7 +14,15 @@ EMBEDDING_DIM = 384
|
||||
|
||||
|
||||
class NoteEmbedding(Base):
|
||||
"""Stores the embedding vector for a note, used for semantic search."""
|
||||
"""One embedding vector per CHUNK of a note (#280, migration 0077).
|
||||
|
||||
The model reads at most 512 tokens, so a single whole-document vector
|
||||
permanently lost everything past ~400 words. A note now stores one row per
|
||||
chunk of `embeddings.chunk_document`, and a query matches the note if it
|
||||
matches ANY chunk — retrieval collapses rows to best-chunk-per-note.
|
||||
A short note has exactly one row (chunk_index 0) whose text is the
|
||||
historical `title\\nbody` shape.
|
||||
"""
|
||||
|
||||
__tablename__ = "note_embeddings"
|
||||
|
||||
@@ -23,8 +31,16 @@ class NoteEmbedding(Base):
|
||||
ForeignKey("notes.id", ondelete="CASCADE"),
|
||||
primary_key=True,
|
||||
)
|
||||
chunk_index: Mapped[int] = mapped_column(Integer, primary_key=True)
|
||||
user_id: Mapped[int] = mapped_column(Integer, nullable=False, index=True)
|
||||
embedding: Mapped[list] = mapped_column(Vector(EMBEDDING_DIM), nullable=False)
|
||||
# Exactly what this vector encodes — inspectable when a ranking surprises,
|
||||
# and the hook for surfacing WHICH section matched, later.
|
||||
chunk_text: Mapped[str] = mapped_column(Text, nullable=False)
|
||||
# embeddings.CHUNKER_VERSION at write time. The startup backfill re-embeds
|
||||
# any note whose rows carry a stale version — shape changes become a
|
||||
# version bump instead of a table wipe.
|
||||
chunker_version: Mapped[int] = mapped_column(Integer, nullable=False)
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
default=lambda: datetime.now(timezone.utc),
|
||||
|
||||
@@ -75,10 +75,19 @@ async def get_embedding(text: str) -> list[float]:
|
||||
Raises if the fastembed model fails to load. Callers should catch and
|
||||
degrade to keyword search.
|
||||
"""
|
||||
return (await get_embeddings([text]))[0]
|
||||
|
||||
|
||||
async def get_embeddings(texts: list[str]) -> list[list[float]]:
|
||||
"""Embed several texts in one model call (the chunked write path).
|
||||
|
||||
fastembed batches internally, so N chunks cost far less than N single
|
||||
calls. Raises like get_embedding; callers catch and degrade.
|
||||
"""
|
||||
embedder = await _get_model()
|
||||
# embed() is synchronous CPU work; offload so we don't block the event loop.
|
||||
vecs = await asyncio.to_thread(lambda: list(embedder.embed([text])))
|
||||
return vecs[0].tolist()
|
||||
vecs = await asyncio.to_thread(lambda: list(embedder.embed(texts)))
|
||||
return [v.tolist() for v in vecs]
|
||||
|
||||
|
||||
def _cosine_similarity(a: list[float], b: list[float]) -> float:
|
||||
@@ -322,12 +331,36 @@ def chunk_document(title: str | None, body: str | None) -> list[str]:
|
||||
return chunks
|
||||
|
||||
|
||||
async def upsert_note_embedding(note_id: int, user_id: int, text: str) -> None:
|
||||
"""Generate and persist an embedding for a note. Safe to fire-and-forget."""
|
||||
if not text or not text.strip():
|
||||
return
|
||||
async def upsert_note_embedding(
|
||||
note_id: int, user_id: int, title: str | None, body: str | None
|
||||
) -> None:
|
||||
"""Chunk, embed and persist a note's vectors. Safe to fire-and-forget.
|
||||
|
||||
Takes title/body rather than pre-built text so the chunking happens HERE —
|
||||
one path for the write path, the recurrence spawn and the startup backfill,
|
||||
which is the same single-definition discipline embedding_text existed for.
|
||||
|
||||
Replacement is atomic per note: old rows are deleted and the new chunk set
|
||||
inserted in one transaction, so a concurrent read sees the old shape or the
|
||||
new one, never a mixture.
|
||||
"""
|
||||
chunks = chunk_document(title, body)
|
||||
try:
|
||||
embedding = await get_embedding(text)
|
||||
if not chunks:
|
||||
# A record emptied of content should stop being findable by its
|
||||
# old content — clear stale vectors rather than leaving them.
|
||||
async with async_session() as session:
|
||||
await session.execute(
|
||||
delete(NoteEmbedding).where(NoteEmbedding.note_id == note_id)
|
||||
)
|
||||
await session.commit()
|
||||
return
|
||||
except Exception:
|
||||
logger.warning("Failed to clear embedding for note %d", note_id, exc_info=True)
|
||||
return
|
||||
|
||||
try:
|
||||
vectors = await get_embeddings(chunks)
|
||||
except Exception:
|
||||
logger.debug("Skipping embedding for note %d — embedder unavailable", note_id)
|
||||
return
|
||||
@@ -337,9 +370,19 @@ async def upsert_note_embedding(note_id: int, user_id: int, text: str) -> None:
|
||||
await session.execute(
|
||||
delete(NoteEmbedding).where(NoteEmbedding.note_id == note_id)
|
||||
)
|
||||
session.add(NoteEmbedding(note_id=note_id, user_id=user_id, embedding=embedding))
|
||||
for index, (chunk, vector) in enumerate(zip(chunks, vectors)):
|
||||
session.add(
|
||||
NoteEmbedding(
|
||||
note_id=note_id,
|
||||
chunk_index=index,
|
||||
user_id=user_id,
|
||||
embedding=vector,
|
||||
chunk_text=chunk,
|
||||
chunker_version=CHUNKER_VERSION,
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
logger.debug("Upserted embedding for note %d", note_id)
|
||||
logger.debug("Upserted %d chunk embedding(s) for note %d", len(chunks), note_id)
|
||||
except Exception:
|
||||
logger.warning("Failed to persist embedding for note %d", note_id, exc_info=True)
|
||||
|
||||
@@ -500,40 +543,50 @@ async def semantic_search_notes(
|
||||
|
||||
|
||||
async def backfill_note_embeddings() -> None:
|
||||
"""Generate embeddings for all notes that don't have one yet.
|
||||
"""(Re-)embed every note that is missing vectors OR whose stored vectors
|
||||
were produced by an older chunker.
|
||||
|
||||
Runs as a background task at startup. Adds a small sleep between notes
|
||||
so a large backfill doesn't peg CPU.
|
||||
Runs as a background task at startup. Version-awareness is what makes a
|
||||
document-shape change deployable: migration 0077 cleared the table once,
|
||||
and every later CHUNKER_VERSION bump re-embeds the stale notes here — a
|
||||
version comparison instead of another wipe. Adds a small sleep between
|
||||
notes so a large backfill doesn't peg CPU.
|
||||
"""
|
||||
try:
|
||||
async with async_session() as session:
|
||||
existing = {
|
||||
current = {
|
||||
row[0]
|
||||
for row in (
|
||||
await session.execute(select(NoteEmbedding.note_id))
|
||||
await session.execute(
|
||||
select(NoteEmbedding.note_id).where(
|
||||
NoteEmbedding.chunker_version == CHUNKER_VERSION
|
||||
)
|
||||
)
|
||||
).fetchall()
|
||||
}
|
||||
result = await session.execute(
|
||||
select(Note.id, Note.user_id, Note.title, Note.body)
|
||||
)
|
||||
notes_to_embed = [
|
||||
row for row in result.fetchall() if row[0] not in existing
|
||||
row for row in result.fetchall() if row[0] not in current
|
||||
]
|
||||
except Exception:
|
||||
logger.warning("Embedding backfill: failed to query notes", exc_info=True)
|
||||
return
|
||||
|
||||
if not notes_to_embed:
|
||||
logger.info("Embedding backfill: all notes already have embeddings")
|
||||
logger.info("Embedding backfill: all notes current at chunker v%d", CHUNKER_VERSION)
|
||||
return
|
||||
|
||||
logger.info("Embedding backfill: generating embeddings for %d notes", len(notes_to_embed))
|
||||
logger.info(
|
||||
"Embedding backfill: embedding %d notes at chunker v%d",
|
||||
len(notes_to_embed), CHUNKER_VERSION,
|
||||
)
|
||||
success = 0
|
||||
for note_id, user_id, title, body in notes_to_embed:
|
||||
text = embedding_text(title, body)
|
||||
if not text:
|
||||
if not chunk_document(title, body):
|
||||
continue
|
||||
await upsert_note_embedding(note_id, user_id, text)
|
||||
await upsert_note_embedding(note_id, user_id, title, body)
|
||||
success += 1
|
||||
await asyncio.sleep(0.05) # gentle pacing
|
||||
|
||||
|
||||
@@ -33,11 +33,12 @@ def embed_note(note) -> None:
|
||||
try:
|
||||
import asyncio
|
||||
|
||||
from scribe.services.embeddings import embedding_text, upsert_note_embedding
|
||||
text = embedding_text(note.title, note.body)
|
||||
if not text:
|
||||
return
|
||||
asyncio.create_task(upsert_note_embedding(note.id, note.user_id, text))
|
||||
from scribe.services.embeddings import upsert_note_embedding
|
||||
# Chunking and the empty-record gate live inside upsert_note_embedding —
|
||||
# one path for every writer (#280).
|
||||
asyncio.create_task(
|
||||
upsert_note_embedding(note.id, note.user_id, note.title, note.body)
|
||||
)
|
||||
except RuntimeError:
|
||||
pass # no running loop — a sync caller, not a failure
|
||||
except Exception: # noqa: BLE001 - never let indexing break a write
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
"""Recurring task rule validation, date calculation, and scheduled spawning."""
|
||||
import asyncio
|
||||
import calendar
|
||||
import logging
|
||||
from datetime import date, datetime, timedelta, timezone
|
||||
@@ -103,7 +102,6 @@ async def spawn_recurring_tasks() -> int:
|
||||
|
||||
Returns the number of tasks spawned.
|
||||
"""
|
||||
from scribe.services.embeddings import embedding_text, upsert_note_embedding
|
||||
from scribe.services.notes import create_note
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
@@ -139,9 +137,9 @@ async def spawn_recurring_tasks() -> int:
|
||||
milestone_id=task.milestone_id,
|
||||
recurrence_rule=task.recurrence_rule,
|
||||
)
|
||||
text = embedding_text(child.title, child.body)
|
||||
if text:
|
||||
asyncio.create_task(upsert_note_embedding(child.id, task.user_id, text))
|
||||
# No explicit embed here: create_note calls embed_note itself, so
|
||||
# the spawn path stopped being a second copy of the embedding rule
|
||||
# the moment that moved into the service (#2056, #280).
|
||||
except Exception:
|
||||
logger.exception("Failed to spawn recurring task %d", task.id)
|
||||
continue
|
||||
|
||||
Reference in New Issue
Block a user