goosedump 0.12.43

Browse, search, compact, and learn from coding-agent sessions
// SPDX-License-Identifier: LGPL-2.1-or-later
// Copyright (C) Jarkko Sakkinen 2026

use anyhow::Context as _;
use rusqlite::{Connection, params};

use crate::engine::model::{EMBEDDING_MODEL_ID, Embedder, Embedding};

use super::super::types::MemoryError;
use super::recall::{MemoryEmbedding, store_embeddings};
use super::util::EMBEDDING_BACKFILL_LIMIT;

pub(super) trait EmbeddingProvider {
    fn embed_batch(&self, texts: &[&str]) -> anyhow::Result<Vec<Embedding>>;
}

impl EmbeddingProvider for Embedder {
    fn embed_batch(&self, texts: &[&str]) -> anyhow::Result<Vec<Embedding>> {
        Embedder::embed_batch(self, texts)
    }
}

pub(super) fn embed_query(
    embedder: &impl EmbeddingProvider,
    query: &str,
) -> anyhow::Result<Embedding> {
    let mut embeddings = embedder.embed_batch(&[query])?;
    if embeddings.len() != 1 {
        return Err(MemoryError::EmbeddingBatchMismatch.into());
    }
    embeddings.pop().context("embedding batch omitted query")
}

pub(super) fn backfill_project_embeddings(
    conn: &Connection,
    project: &str,
    embedder: &impl EmbeddingProvider,
) -> anyhow::Result<usize> {
    let mut stmt = conn.prepare(
        "SELECT claims.id, claims.statement
         FROM claims
         JOIN projects ON projects.id = claims.project_id
         LEFT JOIN claim_embeddings
           ON claim_embeddings.claim_id = claims.id
          AND claim_embeddings.embedding_model = ?1
         WHERE claim_embeddings.claim_id IS NULL
           AND projects.path = ?2
           AND claims.status = 'active'
         ORDER BY claims.created_at, claims.id
         LIMIT ?3",
    )?;
    let rows = stmt
        .query_map(
            params![
                EMBEDDING_MODEL_ID,
                project,
                i64::try_from(EMBEDDING_BACKFILL_LIMIT)?
            ],
            |row| Ok((row.get::<_, String>(0)?, row.get::<_, String>(1)?)),
        )?
        .collect::<rusqlite::Result<Vec<_>>>()?;
    drop(stmt);
    if rows.is_empty() {
        return Ok(0);
    }

    let texts = rows
        .iter()
        .map(|(_, statement)| statement.as_str())
        .collect::<Vec<_>>();
    let vectors = embedder.embed_batch(&texts)?;
    if vectors.len() != rows.len() {
        return Err(MemoryError::EmbeddingBatchMismatch.into());
    }
    let embeddings = rows
        .into_iter()
        .zip(vectors)
        .map(|((claim_id, _), embedding)| MemoryEmbedding {
            claim_id,
            embedding,
        })
        .collect::<Vec<_>>();
    let stored = embeddings.len();
    store_embeddings(conn, &embeddings)?;
    Ok(stored)
}