def embed_pending_messages(
connection: sqlite3.Connection,
embedder: Embedder,
*,
batch_size: int = 64,
) -> int:
if batch_size < 1:
raise ValueError("embedding batch size must be positive")
rows = connection.execute(
"SELECT m.id, m.text, m.content_hash FROM messages m "
"LEFT JOIN embeddings e ON e.message_id = m.id "
"WHERE m.role IN ('user', 'assistant') AND length(trim(m.text)) >= 3 "
"AND length(m.text) <= ? "
"AND (e.message_id IS NULL OR e.model_id != ? OR e.dimension != ? "
"OR e.content_hash != m.content_hash) ORDER BY m.rowid",
(MAX_EMBEDDING_CHARS, embedder.model_id, embedder.dimension),
).fetchall()
embedded = 0
for offset in range(0, len(rows), batch_size):
batch = rows[offset : offset + batch_size]
vectors = embedder.encode([row["text"] for row in batch])
if vectors.shape != (len(batch), embedder.dimension):
raise ValueError("embedder returned an unexpected batch shape")
with connection:
for row, vector in zip(batch, vectors, strict=True):
connection.execute(
"INSERT INTO embeddings(message_id, model_id, dimension, content_hash, vector) "
"VALUES (?, ?, ?, ?, ?) ON CONFLICT(message_id) DO UPDATE SET "
"model_id=excluded.model_id, dimension=excluded.dimension, "
"content_hash=excluded.content_hash, vector=excluded.vector",
(
row["id"],
embedder.model_id,
embedder.dimension,
row["content_hash"],
encode_vector(vector),
),
)
embedded += len(batch)
return embedded