perf(memory_tree): batch chunk + score fetch in fetch_leaves (#2941)

Co-authored-by: Steven Enamakel <enamakel@tinyhumans.ai>
This commit is contained in:
mysma-9403
2026-05-30 09:20:03 -07:00
committed by GitHub
co-authored by Steven Enamakel
parent b37bab7134
commit c29cefd01d
5 changed files with 319 additions and 19 deletions
+73 -1
View File
@@ -25,7 +25,7 @@
use anyhow::{Context, Result};
use chrono::{DateTime, TimeZone, Utc};
use rusqlite::{params, Connection, OptionalExtension, Transaction};
use std::collections::HashSet;
use std::collections::{HashMap, HashSet};
#[cfg(test)]
use std::sync::Arc;
use std::time::Duration;
@@ -538,6 +538,78 @@ pub fn get_chunk(config: &Config, id: &str) -> Result<Option<Chunk>> {
})
}
/// Defensive cap for batched `IN (?,?,…)` reads.
///
/// SQLite's compile-time limit on bound parameters in a single statement
/// (`SQLITE_MAX_VARIABLE_NUMBER`) has been **32 766** since 3.32 (2020),
/// so 500 leaves a ~65× safety margin. The current call-site
/// (`memory_tree::retrieval::fetch::fetch_leaves`) is capped at 20 ids,
/// so the chunked loop runs exactly once today. The window exists so
/// future call-sites passing larger id lists do not blow up against a
/// host with a lower compile-time SQLite cap (older builds, custom
/// embeddings, etc.).
///
/// Volume is **not** reduced: all input ids in → all matching rows out.
/// The loop only splits the SQL; the merged `HashMap` is byte-identical
/// to what one giant query would return.
const MAX_FETCH_BATCH: usize = 500;
/// Batched read of full chunk rows by id.
///
/// Contract mirror of looping [`get_chunk`] per id, but in
/// `O(ceil(n / MAX_FETCH_BATCH))` SQLite round-trips instead of `O(n)`.
/// The returned map contains only ids that exist in `mem_tree_chunks`;
/// missing ids are silently absent (same as `get_chunk` returning
/// `Ok(None)`). Callers that depend on input order must iterate their
/// own id slice and look each id up in the map.
///
/// Reuses [`row_to_chunk`] so decoding stays bit-identical to the
/// per-row helper — no risk of decoder drift.
pub fn get_chunks_batch(config: &Config, chunk_ids: &[String]) -> Result<HashMap<String, Chunk>> {
if chunk_ids.is_empty() {
return Ok(HashMap::new());
}
log::debug!(
"[memory::chunk_store] get_chunks_batch: n={} windows={}",
chunk_ids.len(),
chunk_ids.len().div_ceil(MAX_FETCH_BATCH)
);
with_connection(config, |conn| {
let mut out: HashMap<String, Chunk> = HashMap::with_capacity(chunk_ids.len());
for window in chunk_ids.chunks(MAX_FETCH_BATCH) {
// Build the placeholder list `?1, ?2, …, ?n` matching the
// window length; rusqlite assigns positional binds 1..n in
// the order the values are passed.
let placeholders = (1..=window.len())
.map(|i| format!("?{i}"))
.collect::<Vec<_>>()
.join(",");
let sql = format!(
"SELECT id, source_kind, source_id, source_ref, owner,
timestamp_ms, time_range_start_ms, time_range_end_ms,
tags_json, content, token_count, seq_in_source, created_at_ms
FROM mem_tree_chunks WHERE id IN ({placeholders})"
);
let mut stmt = conn.prepare(&sql).context("prepare get_chunks_batch")?;
let params: Vec<&dyn rusqlite::ToSql> =
window.iter().map(|id| id as &dyn rusqlite::ToSql).collect();
let rows = stmt
.query_map(params.as_slice(), row_to_chunk)
.context("query get_chunks_batch")?;
for row in rows {
let chunk = row.context("decode get_chunks_batch row")?;
out.insert(chunk.id.clone(), chunk);
}
}
log::debug!(
"[memory::chunk_store] get_chunks_batch: matched {}/{} ids",
out.len(),
chunk_ids.len()
);
Ok(out)
})
}
/// Query parameters for [`list_chunks`]. All fields are optional filters —
/// callers pass `ListChunksQuery::default()` to get recent-across-everything.
#[derive(Debug, Default, Clone)]
@@ -963,6 +963,53 @@ fn validate_reembed_skip_key_rejects_empty_and_oversized() {
);
}
// ---------- get_chunks_batch ----------
//
// Contract: equivalent to looping `get_chunk` per id but in
// `O(ceil(n / MAX_FETCH_BATCH))` SQLite round-trips. The map carries
// only ids that exist; missing ids are silently absent (same as the
// per-row helper returning Ok(None)).
#[test]
fn get_chunks_batch_returns_present_ids_in_map() {
let (_tmp, cfg) = test_config();
let c1 = sample_chunk("slack:#eng", 0, 1_700_000_000_000);
let c2 = sample_chunk("slack:#eng", 1, 1_700_000_000_000);
let c3 = sample_chunk("slack:#ops", 0, 1_700_000_000_000);
upsert_chunks(&cfg, &[c1.clone(), c2.clone(), c3.clone()]).unwrap();
let ids = vec![c1.id.clone(), c2.id.clone(), c3.id.clone()];
let map = get_chunks_batch(&cfg, &ids).unwrap();
assert_eq!(map.len(), 3);
assert_eq!(map.get(&c1.id), Some(&c1));
assert_eq!(map.get(&c2.id), Some(&c2));
assert_eq!(map.get(&c3.id), Some(&c3));
}
#[test]
fn get_chunks_batch_empty_input_and_missing_ids() {
// Empty input: empty map (no SQL issued).
let (_tmp, cfg) = test_config();
let empty = get_chunks_batch(&cfg, &[]).unwrap();
assert!(empty.is_empty());
// Missing ids: silently absent (mirrors per-row Ok(None)).
// `fetch_leaves` relies on this so partial-result detection
// (`hits.len() < ids.len()`) keeps working unchanged.
let c = sample_chunk("slack:#eng", 0, 1_700_000_000_000);
upsert_chunks(&cfg, &[c.clone()]).unwrap();
let ids = vec![
c.id.clone(),
"ghost:no-such-1".into(),
"ghost:no-such-2".into(),
];
let map = get_chunks_batch(&cfg, &ids).unwrap();
assert_eq!(map.len(), 1);
assert_eq!(map.get(&c.id), Some(&c));
assert!(map.get("ghost:no-such-1").is_none());
assert!(map.get("ghost:no-such-2").is_none());
}
// ---------- get_chunk_embeddings_for_signature_batch ----------
//
// Contract: equivalent to looping `get_chunk_embedding_for_signature`
+95 -17
View File
@@ -13,10 +13,10 @@
use anyhow::Result;
use crate::openhuman::config::Config;
use crate::openhuman::memory_store::chunks::store::get_chunk;
use crate::openhuman::memory_store::chunks::store::get_chunks_batch;
use crate::openhuman::memory_store::content::read as content_read;
use crate::openhuman::memory_tree::retrieval::types::{hit_from_chunk, RetrievalHit};
use crate::openhuman::memory_tree::score::store::get_score;
use crate::openhuman::memory_tree::score::store::get_scores_batch;
/// Max batch size. Callers that pass more than this get truncated with a
/// warn log — no error surface so the LLM sees a partial result.
@@ -47,22 +47,33 @@ pub async fn fetch_leaves(config: &Config, chunk_ids: &[String]) -> Result<Vec<R
let config_owned = config.clone();
let hits = tokio::task::spawn_blocking(move || -> Result<Vec<RetrievalHit>> {
// Two batched SQLite reads up front instead of 2N per-id queries
// inside the loop. With the `MAX_BATCH = 20` cap above, this turns
// 40 round-trips into 2. Per-row decoders are reused inside both
// helpers so the returned `Chunk` and `score.total` values are
// byte-identical to the old per-id path. Missing ids are absent
// from the maps (same contract as `get_chunk` / `get_score`
// returning `Ok(None)`).
let chunk_by_id = get_chunks_batch(&config_owned, &ids)?;
let score_by_id = get_scores_batch(&config_owned, &ids)?;
// Walk the input ids in order so the response preserves caller
// ordering. Missing ids are dropped exactly as before — callers
// detect partial results via `hits.len() < ids.len()`. File I/O
// (`read_chunk_body`) stays per-id: each MD body lives in its
// own on-disk file, so batching there would mean concurrent file
// opens, not a single round-trip — left untouched.
let mut out: Vec<RetrievalHit> = Vec::with_capacity(ids.len());
for id in &ids {
let chunk = match get_chunk(&config_owned, id)? {
Some(c) => c,
None => {
log::debug!(
"[retrieval::fetch] chunk not found — skipping (1 of {} requested)",
ids.len()
);
continue;
}
};
let score = match get_score(&config_owned, id)? {
Some(s) => s.total,
None => 0.0,
for (idx, id) in ids.iter().enumerate() {
let Some(chunk) = chunk_by_id.get(id) else {
log::debug!(
"[retrieval::fetch] chunk not found at index {}/{} — skipping",
idx + 1,
ids.len()
);
continue;
};
let score = score_by_id.get(id).copied().unwrap_or(0.0);
// Leaves are not attached to a materialised tree id via the
// chunk row. `scope` falls back to the chunk's own source_id so
// consumers still see provenance (e.g. "slack:#eng").
@@ -71,7 +82,7 @@ pub async fn fetch_leaves(config: &Config, chunk_ids: &[String]) -> Result<Vec<R
// The `content` column in SQLite holds a ≤500-char preview after
// the MD-on-disk migration; the retrieval API must return the
// complete chunk text so the LLM sees untruncated content.
let mut chunk_with_body = chunk;
let mut chunk_with_body = chunk.clone();
match content_read::read_chunk_body(&config_owned, id) {
Ok(body) => chunk_with_body.content = body,
Err(e) => {
@@ -213,4 +224,71 @@ mod tests {
assert_eq!(out[0].source_ref.as_deref(), Some("slack://slack:#eng/0"));
assert_eq!(out[0].tree_scope, "slack:#eng");
}
/// After the batch refactor, ordering and score propagation rely on
/// walking the input slice and looking each id up in two HashMaps
/// (chunk + score). This test pins both invariants: interleaved
/// present/missing ids keep their input order, and each kept hit
/// carries the score from its own `mem_tree_score` row (not the
/// row of a neighbour, not the 0.0 fallback when a row exists).
#[tokio::test]
async fn fetch_leaves_preserves_input_order_and_propagates_scores() {
use crate::openhuman::memory_tree::score::signals::ScoreSignals;
use crate::openhuman::memory_tree::score::store::{upsert_score, ScoreRow};
let (_tmp, cfg) = test_config();
let c1 = sample_chunk("slack:#eng", 0);
let c2 = sample_chunk("slack:#eng", 1);
let c3 = sample_chunk("slack:#eng", 2);
upsert_chunks(&cfg, &[c1.clone(), c2.clone(), c3.clone()]).unwrap();
stage_test_chunks(&cfg, &[c1.clone(), c2.clone(), c3.clone()]);
// Distinct totals so we can tell which chunk a hit's score came
// from — proves the per-id HashMap lookup keys by chunk_id and
// not by iteration index.
let mk_row = |id: &str, total: f32| ScoreRow {
chunk_id: id.to_string(),
total,
signals: ScoreSignals {
token_count: 0.0,
unique_words: 0.0,
metadata_weight: 0.0,
source_weight: 0.0,
interaction: 0.0,
entity_density: 0.0,
llm_importance: 0.0,
},
llm_importance_reason: None,
dropped: false,
reason: None,
computed_at_ms: 0,
};
upsert_score(&cfg, &mk_row(&c1.id, 0.1)).unwrap();
upsert_score(&cfg, &mk_row(&c2.id, 0.2)).unwrap();
// c3 intentionally has NO score row so we also pin the 0.0
// fallback after the get_scores_batch contract.
// Request order: c2, ghost, c3, c1 — none in natural id order.
let out = fetch_leaves(
&cfg,
&[
c2.id.clone(),
"ghost:no-such".into(),
c3.id.clone(),
c1.id.clone(),
],
)
.await
.unwrap();
assert_eq!(out.len(), 3, "ghost dropped, 3 real chunks returned");
assert_eq!(out[0].node_id, c2.id);
assert_eq!(out[1].node_id, c3.id);
assert_eq!(out[2].node_id, c1.id);
assert!((out[0].score - 0.2).abs() < 1e-6, "c2 score");
assert!(
out[1].score.abs() < 1e-6,
"c3 has no score row → 0.0 fallback"
);
assert!((out[2].score - 0.1).abs() < 1e-6, "c1 score");
}
}
+63 -1
View File
@@ -8,7 +8,9 @@
//! Schema is declared in `memory/tree/store.rs::SCHEMA`; this file only
//! owns the CRUD operations.
use anyhow::Result;
use std::collections::HashMap;
use anyhow::{Context, Result};
use rusqlite::{params, Connection, OptionalExtension, Transaction};
use serde::{Deserialize, Serialize};
@@ -169,6 +171,66 @@ pub fn get_score(config: &Config, chunk_id: &str) -> Result<Option<ScoreRow>> {
})
}
/// Defensive cap for batched `IN (?,?,…)` reads.
///
/// See identical rationale on [`crate::openhuman::memory_store::chunks::store`]'s
/// `MAX_FETCH_BATCH`: SQLite's `SQLITE_MAX_VARIABLE_NUMBER` has been
/// 32 766 since 3.32, so 500 leaves a ~65× safety margin. The
/// `fetch_leaves` call-site is capped at 20 ids so the chunked loop
/// runs exactly once today — the window exists purely so future
/// call-sites passing larger id lists do not blow up against a host
/// with a lower compile-time SQLite cap.
const MAX_FETCH_BATCH: usize = 500;
/// Batched read of just the `total` field for many chunk ids.
///
/// Narrow on purpose — `fetch_leaves` only needs the float, not the
/// full [`ScoreRow`] with all signals. The returned map contains only
/// `chunk_id`s that have a score row; missing ids are silently absent,
/// matching the per-row [`get_score`] contract (callers then fall back
/// to the documented 0.0 neutral).
pub fn get_scores_batch(config: &Config, chunk_ids: &[String]) -> Result<HashMap<String, f32>> {
if chunk_ids.is_empty() {
return Ok(HashMap::new());
}
log::debug!(
"[memory_tree::score] get_scores_batch: n={} windows={}",
chunk_ids.len(),
chunk_ids.len().div_ceil(MAX_FETCH_BATCH)
);
with_connection(config, |conn| {
let mut out: HashMap<String, f32> = HashMap::with_capacity(chunk_ids.len());
for window in chunk_ids.chunks(MAX_FETCH_BATCH) {
let placeholders = (1..=window.len())
.map(|i| format!("?{i}"))
.collect::<Vec<_>>()
.join(",");
let sql = format!(
"SELECT chunk_id, total FROM mem_tree_score
WHERE chunk_id IN ({placeholders})"
);
let mut stmt = conn.prepare(&sql).context("prepare get_scores_batch")?;
let params: Vec<&dyn rusqlite::ToSql> =
window.iter().map(|id| id as &dyn rusqlite::ToSql).collect();
let rows = stmt
.query_map(params.as_slice(), |row| {
Ok((row.get::<_, String>(0)?, row.get::<_, f32>(1)?))
})
.context("query get_scores_batch")?;
for row in rows {
let (chunk_id, total) = row.context("decode get_scores_batch row")?;
out.insert(chunk_id, total);
}
}
log::debug!(
"[memory_tree::score] get_scores_batch: matched {}/{} ids",
out.len(),
chunk_ids.len()
);
Ok(out)
})
}
/// Index one (entity, chunk) association.
///
/// Idempotent on the composite primary key `(entity_id, node_id)` so
@@ -205,3 +205,44 @@ fn summary_entity_index_kind_is_parseable() {
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].entity_kind, EntityKind::Email);
}
// ---------- get_scores_batch ----------
//
// Narrow on purpose: returns only chunk_id -> total. `fetch_leaves` is
// the call-site and only needs the float. Missing chunk_ids are absent
// from the map (mirrors per-row `get_score` Ok(None) → caller fallback
// to 0.0 neutral).
#[test]
fn get_scores_batch_returns_present_chunk_ids() {
let (_tmp, cfg) = test_config();
let r1 = sample_row("c1", false);
let mut r2 = sample_row("c2", false);
r2.total = 0.3;
upsert_score(&cfg, &r1).unwrap();
upsert_score(&cfg, &r2).unwrap();
let ids = vec!["c1".to_string(), "c2".to_string()];
let map = get_scores_batch(&cfg, &ids).unwrap();
assert_eq!(map.len(), 2);
assert!((map.get("c1").copied().unwrap() - 0.7).abs() < 1e-6);
assert!((map.get("c2").copied().unwrap() - 0.3).abs() < 1e-6);
}
#[test]
fn get_scores_batch_empty_input_and_missing_chunk_ids() {
// Empty input: empty map (no SQL issued).
let (_tmp, cfg) = test_config();
let empty = get_scores_batch(&cfg, &[]).unwrap();
assert!(empty.is_empty());
// Missing ids: silently absent so `fetch_leaves` can fall back to
// its documented 0.0 neutral without ambient errors.
let r = sample_row("c1", false);
upsert_score(&cfg, &r).unwrap();
let ids = vec!["c1".to_string(), "ghost:no-such".to_string()];
let map = get_scores_batch(&cfg, &ids).unwrap();
assert_eq!(map.len(), 1);
assert!((map.get("c1").copied().unwrap() - 0.7).abs() < 1e-6);
assert!(map.get("ghost:no-such").is_none());
}