diff --git a/src/openhuman/memory_queue/handlers/mod.rs b/src/openhuman/memory_queue/handlers/mod.rs index 4fabaaaec..e39f6dbc4 100644 --- a/src/openhuman/memory_queue/handlers/mod.rs +++ b/src/openhuman/memory_queue/handlers/mod.rs @@ -23,7 +23,7 @@ use crate::openhuman::memory_store::content::{ self as content_store, read as content_read, tags as content_tags, }; use crate::openhuman::memory_tree::score; -use crate::openhuman::memory_tree::score::embed::{build_write_embedder, pack_checked}; +use crate::openhuman::memory_tree::score::embed::{build_write_embedder, pack_checked, Embedder}; use crate::openhuman::memory_tree::score::store as score_store; use crate::openhuman::memory_tree::tree::store as summary_store; use crate::openhuman::memory_tree::tree::{LeafRef, TreeFactory}; @@ -703,6 +703,89 @@ fn try_mark_summary_reembed_skipped( } } +/// Read each row's stored source text, embed the readable bodies in **one +/// batched call**, and classify per position exactly as the legacy per-row +/// loop did. Generic over the chunk vs summary read/tombstone functions so +/// both Phase-2 passes share one implementation. +/// +/// Failure semantics are preserved verbatim from #1574 §6: +/// - body read failure → log + persistent tombstone (`"body read failed: {e}"`) +/// - embed wrong dim → log + tombstone (`"embed wrong dim"`) +/// - embed error → log + tombstone (`"embed failed: {e}"`) +/// +/// The single difference vs the old code is the embed call shape: one +/// `embed_batch` (which collapses N provider round-trips into one for batch- +/// capable providers, falling back to per-text internally) instead of N +/// sequential `embed` awaits. Ordering is irrelevant — results are zipped +/// back to their ids by position and folded into order-independent state. +async fn reembed_collect( + config: &Config, + embedder: &dyn Embedder, + active_sig: &str, + ids: &[String], + label: &str, + read_body: impl Fn(&Config, &str) -> Result, + mark_skipped: impl Fn(&Config, &str, &str, &str), +) -> Result)>> { + // Phase A: read bodies; persistently tombstone read failures so an + // unreadable row is attempted at most once per signature. + let mut readable: Vec<(&String, String)> = Vec::with_capacity(ids.len()); + for id in ids { + match read_body(config, id) { + Ok(body) => readable.push((id, body)), + Err(e) => { + log::warn!( + "[memory::jobs] reembed_backfill: {label} {id} body read failed: {e}; skipping (sig={active_sig})" + ); + mark_skipped(config, id, active_sig, &format!("body read failed: {e}")); + } + } + } + if readable.is_empty() { + return Ok(Vec::new()); + } + + // Phase B: one batched embed call. Scope `texts` so its borrow on + // `readable` ends before we consume `readable` below. + let results = { + let texts: Vec<&str> = readable.iter().map(|(_, body)| body.as_str()).collect(); + embedder.embed_batch(&texts).await + }; + if results.len() != readable.len() { + // `embed_batch`'s contract is one result per input position. A + // violation means we can't attribute results to ids. Returning an + // empty Vec here would make the handler write nothing yet still + // `Defer`, re-selecting the same ids forever — a non-converging + // chain. Surface it as an error so the chain terminates instead. + anyhow::bail!( + "reembed_backfill: {label} embed_batch returned {} results for {} texts (sig={active_sig})", + results.len(), + readable.len() + ); + } + + // Phase C: classify per position exactly as the legacy loop did. + let mut out: Vec<(String, Vec)> = Vec::with_capacity(readable.len()); + for ((id, _body), result) in readable.into_iter().zip(results) { + match result { + Ok(v) if pack_checked(&v).is_ok() => out.push((id.clone(), v)), + Ok(_) => { + log::warn!( + "[memory::jobs] reembed_backfill: {label} {id} embed wrong dim, skipping (sig={active_sig})" + ); + mark_skipped(config, id, active_sig, "embed wrong dim"); + } + Err(e) => { + log::warn!( + "[memory::jobs] reembed_backfill: {label} {id} embed failed: {e}; skipping (sig={active_sig})" + ); + mark_skipped(config, id, active_sig, &format!("embed failed: {e}")); + } + } + } + Ok(out) +} + async fn handle_reembed_backfill(config: &Config, job: &Job) -> Result { let payload: ReembedBackfillPayload = serde_json::from_str(&job.payload_json).context("parse ReembedBackfill payload")?; @@ -820,78 +903,26 @@ async fn handle_reembed_backfill(config: &Config, job: &Job) -> Result)> = Vec::new(); - for id in &chunk_ids { - match content_read::read_chunk_body(config, id) { - Ok(body) => match embedder.embed(&body).await { - Ok(v) if pack_checked(&v).is_ok() => chunk_vecs.push((id.clone(), v)), - Ok(_) => { - log::warn!( - "[memory::jobs] reembed_backfill: chunk {id} embed wrong dim, skipping (sig={active_sig})" - ); - try_mark_chunk_reembed_skipped(config, id, &active_sig, "embed wrong dim"); - } - Err(e) => { - log::warn!( - "[memory::jobs] reembed_backfill: chunk {id} embed failed: {e}; skipping (sig={active_sig})" - ); - try_mark_chunk_reembed_skipped( - config, - id, - &active_sig, - &format!("embed failed: {e}"), - ); - } - }, - Err(e) => { - log::warn!( - "[memory::jobs] reembed_backfill: chunk {id} body read failed: {e}; skipping (sig={active_sig})" - ); - try_mark_chunk_reembed_skipped( - config, - id, - &active_sig, - &format!("body read failed: {e}"), - ); - } - } - } - let mut summary_vecs: Vec<(String, Vec)> = Vec::new(); - for id in &summary_ids { - match content_read::read_summary_body(config, id) { - Ok(body) => match embedder.embed(&body).await { - Ok(v) if pack_checked(&v).is_ok() => summary_vecs.push((id.clone(), v)), - Ok(_) => { - log::warn!( - "[memory::jobs] reembed_backfill: summary {id} embed wrong dim, skipping (sig={active_sig})" - ); - try_mark_summary_reembed_skipped(config, id, &active_sig, "embed wrong dim"); - } - Err(e) => { - log::warn!( - "[memory::jobs] reembed_backfill: summary {id} embed failed: {e}; skipping (sig={active_sig})" - ); - try_mark_summary_reembed_skipped( - config, - id, - &active_sig, - &format!("embed failed: {e}"), - ); - } - }, - Err(e) => { - log::warn!( - "[memory::jobs] reembed_backfill: summary {id} body read failed: {e}; skipping (sig={active_sig})" - ); - try_mark_summary_reembed_skipped( - config, - id, - &active_sig, - &format!("body read failed: {e}"), - ); - } - } - } + let chunk_vecs = reembed_collect( + config, + embedder.as_ref(), + &active_sig, + &chunk_ids, + "chunk", + content_read::read_chunk_body, + try_mark_chunk_reembed_skipped, + ) + .await?; + let summary_vecs = reembed_collect( + config, + embedder.as_ref(), + &active_sig, + &summary_ids, + "summary", + content_read::read_summary_body, + try_mark_summary_reembed_skipped, + ) + .await?; // Phase 3 (one short tx): persist all collected vectors to the sidecar. chunk_store::with_connection(config, |conn| { diff --git a/src/openhuman/memory_tree/score/embed/cloud.rs b/src/openhuman/memory_tree/score/embed/cloud.rs index 682c1eeed..c0787cb9b 100644 --- a/src/openhuman/memory_tree/score/embed/cloud.rs +++ b/src/openhuman/memory_tree/score/embed/cloud.rs @@ -76,6 +76,14 @@ impl Embedder for CloudEmbedder { } Ok(v) } + + /// Collapse N per-text round-trips into a single batched Voyage request by + /// delegating to the inner provider's native batch `embed`. Falls back to + /// per-text embedding (preserving per-position error attribution) on a + /// whole-batch failure or a length mismatch. + async fn embed_batch(&self, texts: &[&str]) -> Vec>> { + super::embed_batch_via_provider(&self.inner, "cloud", texts).await + } } #[cfg(test)] diff --git a/src/openhuman/memory_tree/score/embed/mod.rs b/src/openhuman/memory_tree/score/embed/mod.rs index 7f678857b..00a783085 100644 --- a/src/openhuman/memory_tree/score/embed/mod.rs +++ b/src/openhuman/memory_tree/score/embed/mod.rs @@ -61,6 +61,120 @@ pub trait Embedder: Send + Sync { /// [`EMBEDDING_DIM`]. Hard failure — ingest / seal treat `Err` as /// "don't persist the row" so retries stay idempotent on `chunk_id`. async fn embed(&self, text: &str) -> Result>; + + /// Embed many texts, returning **one [`Result`] per input position** + /// aligned by index. A single failing text does not strand the rest of + /// the batch — its slot carries the `Err` while the others succeed — + /// which lets bulk callers (e.g. the re-embed backfill) attribute and + /// skip individual rows exactly as a per-text loop would. + /// + /// The default implementation issues one sequential [`Embedder::embed`] + /// call per text: correct for any provider, but with no batching win. + /// Providers whose backend accepts many texts in a single request + /// (cloud / OpenAI-compatible) override this to collapse N network + /// round-trips into one — see [`embed_batch_via_provider`]. + /// + /// The returned vector always has `texts.len()` elements. + async fn embed_batch(&self, texts: &[&str]) -> Vec>> { + log::debug!( + "[memory_tree::embed::{}] embed_batch:enter sequential texts={}", + self.name(), + texts.len() + ); + let mut out = Vec::with_capacity(texts.len()); + for text in texts { + out.push(self.embed(text).await); + } + out + } +} + +/// Validate that a freshly-produced embedding has exactly [`EMBEDDING_DIM`] +/// floats, returning a labelled error otherwise. Shared by the per-text and +/// batched provider adapters so the "wrong dims" diagnostic is identical +/// regardless of path. +pub(crate) fn check_embed_dim(v: Vec, label: &str) -> Result> { + if v.len() != EMBEDDING_DIM { + anyhow::bail!( + "{label} embedder returned {} dims, expected {}", + v.len(), + EMBEDDING_DIM + ); + } + Ok(v) +} + +/// Batch-embed `texts` through a unified [`EmbeddingProvider`] in a single +/// request and adapt the result to the [`Embedder::embed_batch`] contract: +/// one dimension-checked [`Result`] per input position. +/// +/// On a wholesale batch failure (network / auth) **or** a length-contract +/// violation from the provider, this falls back to per-text +/// [`EmbeddingProvider::embed_one`] so a single transient blip cannot fail — +/// and, in the backfill, *tombstone* — every row in the batch. That preserves +/// the resilience of the original per-item loop while keeping the happy-path +/// single-round-trip win. +pub(crate) async fn embed_batch_via_provider( + inner: &dyn crate::openhuman::embeddings::EmbeddingProvider, + label: &str, + texts: &[&str], +) -> Vec>> { + if texts.is_empty() { + return Vec::new(); + } + log::debug!( + "[memory_tree::embed::{label}] embed_batch:enter texts={}", + texts.len() + ); + match inner.embed(texts).await { + Ok(vectors) if vectors.len() == texts.len() => { + log::debug!( + "[memory_tree::embed::{label}] embed_batch:success collapsed {} texts into one provider call", + texts.len() + ); + vectors + .into_iter() + .map(|v| check_embed_dim(v, label)) + .collect() + } + Ok(vectors) => { + log::warn!( + "[memory_tree::embed::{label}] embed_batch:fallback batch returned {} vectors for {} texts; \ + falling back to per-text embedding", + vectors.len(), + texts.len() + ); + embed_each_via_provider(inner, label, texts).await + } + Err(e) => { + log::warn!( + "[memory_tree::embed::{label}] embed_batch:fallback batch embed failed ({e:#}); \ + falling back to per-text embedding" + ); + embed_each_via_provider(inner, label, texts).await + } + } +} + +/// Sequential per-text fallback used when a provider's native batch call is +/// unavailable or fails wholesale. Each slot is dimension-checked so the +/// result is interchangeable with the happy-path mapping in +/// [`embed_batch_via_provider`]. +async fn embed_each_via_provider( + inner: &dyn crate::openhuman::embeddings::EmbeddingProvider, + label: &str, + texts: &[&str], +) -> Vec>> { + let mut out = Vec::with_capacity(texts.len()); + for text in texts { + let result = inner + .embed_one(text) + .await + .with_context(|| format!("{label} embeddings failed")) + .and_then(|v| check_embed_dim(v, label)); + out.push(result); + } + out } /// Cosine similarity between two equal-length vectors. @@ -234,4 +348,193 @@ mod tests { let correct = vec![0.0_f32; EMBEDDING_DIM]; assert!(pack_checked(&correct).is_ok()); } + + // --- batch-embedding (variant B) scaffolding + tests --- + + use crate::openhuman::embeddings::EmbeddingProvider; + use std::sync::atomic::{AtomicUsize, Ordering}; + use std::sync::Arc; + + fn ok_vec() -> Vec { + vec![0.5_f32; EMBEDDING_DIM] + } + + #[derive(Clone)] + enum ProviderMode { + /// One correct-dim vector per text (single batch call succeeds). + Ok, + /// Batch (`len > 1`) call errors, per-text (`len == 1`) succeeds — + /// exercises the whole-batch-error fallback path. + BatchFailsPerTextOk, + /// Batch (`len > 1`) returns one extra vector, per-text is fine — + /// exercises the length-mismatch fallback path. + WrongCount, + /// Returns `len` vectors but the one at `idx` has the wrong dim — + /// length matches so no fallback; that position must map to `Err`. + OneWrongDim(usize), + } + + struct FakeProvider { + calls: Arc, + mode: ProviderMode, + } + + #[async_trait::async_trait] + impl EmbeddingProvider for FakeProvider { + fn name(&self) -> &str { + "fake" + } + fn model_id(&self) -> &str { + "fake-model" + } + fn dimensions(&self) -> usize { + EMBEDDING_DIM + } + async fn embed(&self, texts: &[&str]) -> anyhow::Result>> { + self.calls.fetch_add(1, Ordering::SeqCst); + match self.mode { + ProviderMode::Ok => Ok(texts.iter().map(|_| ok_vec()).collect()), + ProviderMode::BatchFailsPerTextOk => { + if texts.len() > 1 { + anyhow::bail!("simulated batch endpoint failure") + } else { + Ok(texts.iter().map(|_| ok_vec()).collect()) + } + } + ProviderMode::WrongCount => { + if texts.len() > 1 { + Ok((0..texts.len() + 1).map(|_| ok_vec()).collect()) + } else { + Ok(texts.iter().map(|_| ok_vec()).collect()) + } + } + ProviderMode::OneWrongDim(idx) => Ok(texts + .iter() + .enumerate() + .map(|(i, _)| if i == idx { vec![0.0_f32; 3] } else { ok_vec() }) + .collect()), + } + } + } + + #[tokio::test] + async fn embed_batch_via_provider_happy_is_single_call() { + let calls = Arc::new(AtomicUsize::new(0)); + let p = FakeProvider { + calls: calls.clone(), + mode: ProviderMode::Ok, + }; + let out = embed_batch_via_provider(&p, "test", &["a", "b", "c"]).await; + assert_eq!(out.len(), 3); + assert!(out.iter().all(|r| r.is_ok())); + assert_eq!( + calls.load(Ordering::SeqCst), + 1, + "happy path must collapse to exactly one batch call" + ); + } + + #[tokio::test] + async fn embed_batch_via_provider_empty_makes_no_call() { + let calls = Arc::new(AtomicUsize::new(0)); + let p = FakeProvider { + calls: calls.clone(), + mode: ProviderMode::Ok, + }; + let texts: [&str; 0] = []; + let out = embed_batch_via_provider(&p, "test", &texts).await; + assert!(out.is_empty()); + assert_eq!(calls.load(Ordering::SeqCst), 0); + } + + #[tokio::test] + async fn embed_batch_via_provider_falls_back_on_batch_error() { + let calls = Arc::new(AtomicUsize::new(0)); + let p = FakeProvider { + calls: calls.clone(), + mode: ProviderMode::BatchFailsPerTextOk, + }; + let out = embed_batch_via_provider(&p, "test", &["a", "b", "c"]).await; + assert_eq!(out.len(), 3); + assert!( + out.iter().all(|r| r.is_ok()), + "per-text fallback should still produce all vectors" + ); + // 1 failed batch call + 3 per-text calls. + assert_eq!(calls.load(Ordering::SeqCst), 4); + } + + #[tokio::test] + async fn embed_batch_via_provider_falls_back_on_length_mismatch() { + let calls = Arc::new(AtomicUsize::new(0)); + let p = FakeProvider { + calls: calls.clone(), + mode: ProviderMode::WrongCount, + }; + let out = embed_batch_via_provider(&p, "test", &["a", "b"]).await; + assert_eq!(out.len(), 2); + assert!(out.iter().all(|r| r.is_ok())); + // 1 mismatched batch call + 2 per-text calls. + assert_eq!(calls.load(Ordering::SeqCst), 3); + } + + #[tokio::test] + async fn embed_batch_via_provider_maps_wrong_dim_per_position() { + let calls = Arc::new(AtomicUsize::new(0)); + let p = FakeProvider { + calls: calls.clone(), + mode: ProviderMode::OneWrongDim(1), + }; + let out = embed_batch_via_provider(&p, "test", &["a", "b", "c"]).await; + assert_eq!(out.len(), 3); + assert!(out[0].is_ok()); + assert!(out[1].is_err(), "wrong-dim vector maps to Err at its slot"); + assert!(out[2].is_ok()); + // Length matched, so no fallback — a single batch call. + assert_eq!(calls.load(Ordering::SeqCst), 1); + } + + struct SeqEmbedder { + calls: Arc, + } + + #[async_trait::async_trait] + impl Embedder for SeqEmbedder { + fn name(&self) -> &'static str { + "seq" + } + async fn embed(&self, text: &str) -> Result> { + self.calls.fetch_add(1, Ordering::SeqCst); + if text == "bad" { + anyhow::bail!("simulated per-text failure") + } + Ok(ok_vec()) + } + // Uses the default `embed_batch`. + } + + #[tokio::test] + async fn default_embed_batch_calls_embed_per_text() { + let calls = Arc::new(AtomicUsize::new(0)); + let e = SeqEmbedder { + calls: calls.clone(), + }; + let out = e.embed_batch(&["a", "b", "c"]).await; + assert_eq!(out.len(), 3); + assert!(out.iter().all(|r| r.is_ok())); + assert_eq!(calls.load(Ordering::SeqCst), 3); + } + + #[tokio::test] + async fn default_embed_batch_preserves_per_position_errors() { + let calls = Arc::new(AtomicUsize::new(0)); + let e = SeqEmbedder { + calls: calls.clone(), + }; + let out = e.embed_batch(&["ok", "bad", "ok"]).await; + assert_eq!(out.len(), 3); + assert!(out[0].is_ok()); + assert!(out[1].is_err()); + assert!(out[2].is_ok()); + } } diff --git a/src/openhuman/memory_tree/score/embed/openai_compat.rs b/src/openhuman/memory_tree/score/embed/openai_compat.rs index 9e18cd735..8dd8bfb92 100644 --- a/src/openhuman/memory_tree/score/embed/openai_compat.rs +++ b/src/openhuman/memory_tree/score/embed/openai_compat.rs @@ -107,6 +107,14 @@ impl Embedder for OpenAiCompatEmbedder { } Ok(v) } + + /// Collapse N per-text round-trips into a single batched request by + /// delegating to the inner provider's native batch `embed`. Falls back to + /// per-text embedding (preserving per-position error attribution) on a + /// whole-batch failure or a length mismatch. + async fn embed_batch(&self, texts: &[&str]) -> Vec>> { + super::embed_batch_via_provider(self.inner.as_ref(), self.label, texts).await + } } #[cfg(test)]