Skip to main content

remem/retrieval/vector/
backfill.rs

1use std::time::Instant;
2
3use anyhow::{Context, Result};
4use rusqlite::{params, Connection};
5
6use crate::retrieval::embedding::{
7    EmbeddingBackfillTarget, EmbeddingConfig, EmbeddingFallbackCache, EmbeddingProviderStatus,
8};
9
10use super::reindex::{
11    prepare_memory_embedding_batch, select_memory_embedding_reindex_candidates,
12    PreparedMemoryEmbedding,
13};
14
15#[derive(Debug, Clone, PartialEq)]
16pub struct EmbeddingReindexReport {
17    pub selected: usize,
18    pub processed: usize,
19    pub model: String,
20    pub dimensions: usize,
21    pub timings: Vec<crate::perf::PhaseTiming>,
22}
23
24#[derive(Debug)]
25pub(crate) struct EmbeddingBackfillSession {
26    target: EmbeddingBackfillTarget,
27    fallback_cache: EmbeddingFallbackCache,
28    pinned_config: EmbeddingConfig,
29    pinned_status: EmbeddingProviderStatus,
30    prune_block_reason: Option<String>,
31    disabled: bool,
32}
33
34impl EmbeddingBackfillSession {
35    pub(crate) fn start() -> Result<Self> {
36        let config_before = crate::retrieval::embedding::resolve_embedding_config()?;
37        let status_before = crate::retrieval::embedding::embedding_provider_status_without_probe()?;
38        if status_before.disabled {
39            if let Some(error) =
40                crate::retrieval::embedding::disabled_provider_status_error(&status_before)
41            {
42                if !crate::retrieval::embedding::is_embedding_provider_off_error(&error) {
43                    return Err(error);
44                }
45            }
46            return Ok(Self {
47                target: EmbeddingBackfillTarget {
48                    model: "off".to_string(),
49                    dimensions: 0,
50                },
51                fallback_cache: EmbeddingFallbackCache::default(),
52                pinned_config: config_before,
53                prune_block_reason: status_before
54                    .degradation_reason
55                    .clone()
56                    .or_else(|| Some("embedding provider is off".to_string())),
57                pinned_status: status_before,
58                disabled: true,
59            });
60        }
61
62        let mut fallback_cache = EmbeddingFallbackCache::default();
63        let target = crate::retrieval::embedding::configured_backfill_target_with_fallback_cache(
64            &mut fallback_cache,
65        )?;
66        let config_after = crate::retrieval::embedding::resolve_embedding_config()?;
67        let status_after = crate::retrieval::embedding::embedding_provider_status_without_probe()?;
68        if config_after != config_before || status_after != status_before {
69            anyhow::bail!(
70                "embedding configuration or active profile changed while pinning backfill target; refusing to start"
71            );
72        }
73
74        let fallback_target = fallback_cache.call_failure_fallback_target();
75        if let Some(fallback_target) = fallback_target.as_ref() {
76            if fallback_target != &target {
77                anyhow::bail!(
78                    "pinned embedding profile changed while starting backfill: expected model={} dimensions={}, got fallback model={} dimensions={}",
79                    target.model,
80                    target.dimensions,
81                    fallback_target.model,
82                    fallback_target.dimensions
83                );
84            }
85        }
86        let prune_block_reason = status_after
87            .degraded
88            .then(|| {
89                status_after
90                    .degradation_reason
91                    .clone()
92                    .unwrap_or_else(|| "embedding provider is degraded".to_string())
93            })
94            .or_else(|| {
95                fallback_target.map(|fallback_target| {
96                    format!(
97                        "typed provider fallback selected model={} dimensions={}",
98                        fallback_target.model, fallback_target.dimensions
99                    )
100                })
101            });
102        Ok(Self {
103            target,
104            fallback_cache,
105            pinned_config: config_after,
106            pinned_status: status_after,
107            prune_block_reason,
108            disabled: false,
109        })
110    }
111
112    pub(crate) fn target(&self) -> &EmbeddingBackfillTarget {
113        &self.target
114    }
115
116    fn is_disabled(&self) -> bool {
117        self.disabled
118    }
119
120    pub(crate) fn ensure_environment_unchanged(&mut self, phase: &str) -> Result<()> {
121        let current_config = crate::retrieval::embedding::resolve_embedding_config()?;
122        if current_config != self.pinned_config {
123            let reason = format!(
124                "embedding configuration changed {phase} for pinned model={} dimensions={}",
125                self.target.model, self.target.dimensions
126            );
127            self.prune_block_reason.get_or_insert(reason.clone());
128            anyhow::bail!("{reason}");
129        }
130        let current_status =
131            crate::retrieval::embedding::embedding_provider_status_without_probe()?;
132        if current_status != self.pinned_status {
133            let current_dimensions = current_status
134                .active_dimensions
135                .map(|dimensions| dimensions.to_string())
136                .unwrap_or_else(|| "<none>".to_string());
137            let reason = format!(
138                "active embedding provider/profile changed {phase}: pinned model={} dimensions={}, current model={} dimensions={}",
139                self.target.model,
140                self.target.dimensions,
141                current_status.active_model_id.as_deref().unwrap_or("<none>"),
142                current_dimensions
143            );
144            self.prune_block_reason.get_or_insert(reason.clone());
145            anyhow::bail!("{reason}");
146        }
147        Ok(())
148    }
149
150    fn observe_call_failure_fallback(&mut self) -> Result<()> {
151        let Some(fallback_target) = self.fallback_cache.call_failure_fallback_target() else {
152            return Ok(());
153        };
154        self.prune_block_reason.get_or_insert_with(|| {
155            format!(
156                "typed provider fallback selected model={} dimensions={}",
157                fallback_target.model, fallback_target.dimensions
158            )
159        });
160        if fallback_target != self.target {
161            anyhow::bail!(
162                "pinned embedding profile changed after typed fallback: expected model={} dimensions={}, got model={} dimensions={}",
163                self.target.model,
164                self.target.dimensions,
165                fallback_target.model,
166                fallback_target.dimensions
167            );
168        }
169        Ok(())
170    }
171
172    pub(crate) fn ensure_prune_preconditions(&mut self) -> Result<()> {
173        self.ensure_environment_unchanged("before pruning")?;
174        if self.disabled {
175            anyhow::bail!("cannot prune embedding profiles while embedding provider is off");
176        }
177        if let Some(reason) = self.prune_block_reason.as_deref() {
178            anyhow::bail!(
179                "refusing to prune embedding profiles after an unsafe provider transition: {reason}"
180            );
181        }
182        Ok(())
183    }
184}
185
186pub fn backfill_missing_memory_embeddings(conn: &Connection, limit: i64) -> Result<usize> {
187    reindex_memory_embeddings(conn, limit)
188}
189
190pub fn reindex_memory_embeddings(conn: &Connection, limit: i64) -> Result<usize> {
191    let mut remaining_limit = limit.max(0);
192    if remaining_limit == 0
193        || !super::table_exists(conn, "memories")?
194        || !super::table_exists(conn, "memory_embeddings")?
195    {
196        return Ok(0);
197    }
198    let mut session = EmbeddingBackfillSession::start()?;
199    let mut processed = 0usize;
200    while remaining_limit > 0 {
201        let batch_limit = remaining_limit.min(super::EMBEDDING_REINDEX_WRITE_BATCH_SIZE as i64);
202        let report =
203            reindex_memory_embeddings_with_session_report(conn, batch_limit, &mut session)?;
204        if report.processed == 0 {
205            break;
206        }
207        processed += report.processed;
208        remaining_limit -= report.processed as i64;
209        if report.processed < batch_limit as usize {
210            break;
211        }
212    }
213    session.ensure_environment_unchanged("before finalizing backfill")?;
214    Ok(processed)
215}
216
217pub fn reindex_memory_embeddings_with_report(
218    conn: &Connection,
219    limit: i64,
220) -> Result<EmbeddingReindexReport> {
221    let total_start = Instant::now();
222    let mut timings = vec![];
223    if crate::retrieval::embedding::provider_disabled_or_error()? {
224        crate::perf::push_elapsed(&mut timings, "total", total_start);
225        return Ok(empty_report("off", 0, timings));
226    }
227    if !super::table_exists(conn, "memories")?
228        || !super::table_exists(conn, "memory_embeddings")?
229        || limit.max(0) == 0
230    {
231        crate::perf::push_elapsed(&mut timings, "total", total_start);
232        return Ok(empty_report("", 0, timings));
233    }
234    let mut session = EmbeddingBackfillSession::start()?;
235    reindex_memory_embeddings_with_session_report(conn, limit, &mut session)
236}
237
238pub(crate) fn reindex_memory_embeddings_with_session_report(
239    conn: &Connection,
240    limit: i64,
241    session: &mut EmbeddingBackfillSession,
242) -> Result<EmbeddingReindexReport> {
243    let total_start = Instant::now();
244    let mut timings = vec![];
245    if session.is_disabled() {
246        crate::perf::push_elapsed(&mut timings, "total", total_start);
247        return Ok(empty_report("off", 0, timings));
248    }
249    if !super::table_exists(conn, "memories")?
250        || !super::table_exists(conn, "memory_embeddings")?
251        || limit.max(0) == 0
252    {
253        crate::perf::push_elapsed(&mut timings, "total", total_start);
254        return Ok(empty_report("", 0, timings));
255    }
256
257    let profile_start = Instant::now();
258    session.ensure_environment_unchanged("before selecting a backfill batch")?;
259    crate::perf::push_elapsed(&mut timings, "profile_probe", profile_start);
260    let target = session.target.clone();
261
262    let select_start = Instant::now();
263    let pending = select_memory_embedding_reindex_candidates(conn, &target, limit)?;
264    crate::perf::push_elapsed(&mut timings, "select_pending", select_start);
265    let selected = pending.len();
266    if pending.is_empty() {
267        session.ensure_environment_unchanged("before completing an empty backfill batch")?;
268        crate::perf::push_elapsed(&mut timings, "total", total_start);
269        return Ok(EmbeddingReindexReport {
270            selected,
271            processed: 0,
272            model: target.model,
273            dimensions: target.dimensions,
274            timings,
275        });
276    }
277
278    let prepared =
279        prepare_memory_embedding_batch(&pending, &mut timings, &mut session.fallback_cache)?;
280    session.observe_call_failure_fallback()?;
281    validate_prepared_embedding_profiles(&target, &prepared)?;
282    session.ensure_environment_unchanged("before writing a backfill batch")?;
283
284    let processed = upsert_prepared_memory_embedding_batch(conn, &prepared, &mut timings)?;
285    session.ensure_environment_unchanged("after writing a backfill batch")?;
286    crate::perf::push_elapsed(&mut timings, "total", total_start);
287    Ok(EmbeddingReindexReport {
288        selected,
289        processed,
290        model: target.model,
291        dimensions: target.dimensions,
292        timings,
293    })
294}
295
296fn validate_prepared_embedding_profiles(
297    target: &EmbeddingBackfillTarget,
298    prepared: &[PreparedMemoryEmbedding],
299) -> Result<()> {
300    for embedding in prepared {
301        let actual_dimensions = embedding.values.len();
302        if embedding.model != target.model || actual_dimensions != target.dimensions {
303            anyhow::bail!(
304                "pinned embedding profile changed while preparing memory id={}: expected model={} dimensions={}, got model={} dimensions={}; refusing to write mixed backfill batch",
305                embedding.memory_id,
306                target.model,
307                target.dimensions,
308                embedding.model,
309                actual_dimensions
310            );
311        }
312    }
313    Ok(())
314}
315
316fn empty_report(
317    model: &str,
318    dimensions: usize,
319    timings: Vec<crate::perf::PhaseTiming>,
320) -> EmbeddingReindexReport {
321    EmbeddingReindexReport {
322        selected: 0,
323        processed: 0,
324        model: model.to_string(),
325        dimensions,
326        timings,
327    }
328}
329
330pub fn pending_memory_embedding_count(conn: &Connection) -> Result<i64> {
331    pending_memory_embedding_reindex_count(conn)
332}
333
334pub fn pending_memory_embedding_reindex_count(conn: &Connection) -> Result<i64> {
335    if crate::retrieval::embedding::provider_disabled_or_error()? {
336        return Ok(0);
337    }
338    if !super::table_exists(conn, "memories")? || !super::table_exists(conn, "memory_embeddings")? {
339        return Ok(0);
340    }
341    let target = match crate::retrieval::embedding::configured_backfill_target() {
342        Ok(target) => target,
343        Err(error) if crate::retrieval::embedding::is_embedding_provider_off_error(&error) => {
344            return Ok(0);
345        }
346        Err(error) => return Err(error),
347    };
348    pending_memory_embedding_reindex_count_for_target(conn, &target)
349}
350
351pub fn pending_memory_embedding_reindex_count_for_target(
352    conn: &Connection,
353    target: &EmbeddingBackfillTarget,
354) -> Result<i64> {
355    if target.dimensions == 0
356        || !super::table_exists(conn, "memories")?
357        || !super::table_exists(conn, "memory_embeddings")?
358    {
359        return Ok(0);
360    }
361    Ok(conn.query_row(
362        "SELECT COUNT(*)
363         FROM memories m
364         LEFT JOIN memory_embeddings e
365           ON e.memory_id = m.id
366          AND e.model = ?1
367          AND e.dimensions = ?2
368         WHERE (e.memory_id IS NULL
369                OR e.updated_at_epoch < m.updated_at_epoch)
370           AND m.status IN ('active', 'stale', 'archived')",
371        params![target.model.as_str(), target.dimensions as i64],
372        |row| row.get(0),
373    )?)
374}
375
376fn upsert_prepared_memory_embedding_batch(
377    conn: &Connection,
378    prepared: &[PreparedMemoryEmbedding],
379    timings: &mut Vec<crate::perf::PhaseTiming>,
380) -> Result<usize> {
381    if prepared.is_empty() {
382        return Ok(0);
383    }
384    let prepared_count = prepared.len();
385    conn.execute_batch("SAVEPOINT remem_embedding_reindex_batch")
386        .context("start memory embedding reindex savepoint")?;
387    let result = (|| -> Result<()> {
388        let upsert_start = Instant::now();
389        {
390            let mut stmt = conn.prepare(super::UPSERT_EMBEDDING_SQL)?;
391            for embedding in prepared {
392                super::execute_embedding_upsert(
393                    &mut stmt,
394                    embedding.memory_id,
395                    &embedding.model,
396                    &embedding.content_hash,
397                    &embedding.values,
398                    embedding.updated_at_epoch,
399                )
400                .with_context(|| {
401                    format!(
402                        "memory embedding upsert failed for memory id={}",
403                        embedding.memory_id
404                    )
405                })?;
406            }
407        }
408        let mut by_dimensions: std::collections::BTreeMap<usize, Vec<i64>> =
409            std::collections::BTreeMap::new();
410        for embedding in prepared {
411            by_dimensions
412                .entry(embedding.values.len())
413                .or_default()
414                .push(embedding.memory_id);
415        }
416        for (dimensions, memory_ids) in by_dimensions {
417            super::vec_index::sync_vec_upsert_batch(conn, dimensions, &memory_ids)?;
418        }
419        crate::perf::push_elapsed(timings, "upsert_embeddings", upsert_start);
420        Ok(())
421    })();
422
423    match result {
424        Ok(()) => {
425            let commit_start = Instant::now();
426            conn.execute_batch("RELEASE SAVEPOINT remem_embedding_reindex_batch")
427                .context("release memory embedding reindex savepoint")?;
428            crate::perf::push_elapsed(timings, "commit", commit_start);
429            Ok(prepared_count)
430        }
431        Err(error) => {
432            let rollback_result = conn.execute_batch(
433                "ROLLBACK TO SAVEPOINT remem_embedding_reindex_batch;
434                 RELEASE SAVEPOINT remem_embedding_reindex_batch",
435            );
436            match rollback_result {
437                Ok(()) => Err(error),
438                Err(rollback_error) => Err(error).context(format!(
439                    "memory embedding reindex failed and rollback failed: {rollback_error}"
440                )),
441            }
442        }
443    }
444}