feat(memory): add per-model embedding tables (#2010)

Co-authored-by: honor2030 <19909783+honor2030@users.noreply.github.com>
This commit is contained in:
이민재
2026-05-17 18:09:37 -07:00
committed by GitHub
co-authored by honor2030
parent db99318ea4
commit 3b7c4cdc38
8 changed files with 528 additions and 1 deletions
+84 -1
View File
@@ -7,7 +7,7 @@
//! - Tier B (local LLM): runs on segment close if local AI is enabled.
use parking_lot::Mutex;
use rusqlite::{params, Connection};
use rusqlite::{params, Connection, OptionalExtension};
use serde::{Deserialize, Serialize};
use std::sync::Arc;
@@ -62,6 +62,21 @@ CREATE TRIGGER IF NOT EXISTS event_au AFTER UPDATE ON event_log BEGIN
INSERT INTO event_fts(rowid, content, subject, event_type)
VALUES (new.rowid, new.content, new.subject, new.event_type);
END;
-- Per-(event, embedding model) vectors (#1574). The legacy event_log.embedding
-- column stays available during the dual-write migration; this table records
-- vector-space provenance for safe provider/model switches.
CREATE TABLE IF NOT EXISTS event_embeddings (
event_id TEXT NOT NULL REFERENCES event_log(event_id) ON DELETE CASCADE,
model_signature TEXT NOT NULL,
vector BLOB NOT NULL,
dim INTEGER NOT NULL,
created_at REAL NOT NULL,
PRIMARY KEY (event_id, model_signature)
);
CREATE INDEX IF NOT EXISTS idx_event_embeddings_model
ON event_embeddings(model_signature);
"#;
/// Event types extracted from conversations.
@@ -153,6 +168,54 @@ pub fn event_insert(conn: &Arc<Mutex<Connection>>, event: &EventRecord) -> anyho
Ok(())
}
/// Store an event embedding for a specific provider/model/dimension signature.
///
/// This writes only the per-model table introduced for #1574. The legacy
/// `event_log.embedding` column remains available for dual-read fallback.
pub fn event_embedding_upsert(
conn: &Arc<Mutex<Connection>>,
event_id: &str,
model_signature: &str,
embedding: &[f32],
created_at: f64,
) -> anyhow::Result<()> {
let bytes = vec_to_bytes(embedding);
let dim = i64::try_from(embedding.len())?;
let conn = conn.lock();
conn.execute(
"INSERT INTO event_embeddings (event_id, model_signature, vector, dim, created_at)
VALUES (?1, ?2, ?3, ?4, ?5)
ON CONFLICT(event_id, model_signature) DO UPDATE SET
vector = excluded.vector,
dim = excluded.dim,
created_at = excluded.created_at",
params![event_id, model_signature, bytes, dim, created_at],
)?;
Ok(())
}
/// Fetch an event embedding for exactly one provider/model/dimension signature.
pub fn event_embedding_get(
conn: &Arc<Mutex<Connection>>,
event_id: &str,
model_signature: &str,
) -> anyhow::Result<Option<Vec<f32>>> {
let conn = conn.lock();
let row: Option<(Vec<u8>, i64)> = conn
.query_row(
"SELECT vector, dim
FROM event_embeddings
WHERE event_id = ?1 AND model_signature = ?2",
params![event_id, model_signature],
|r| Ok((r.get(0)?, r.get(1)?)),
)
.optional()?;
match row {
None => Ok(None),
Some((bytes, dim)) => decode_embedding_row(&bytes, dim),
}
}
/// Search events via FTS5, scoped to a namespace.
pub fn event_search_fts(
conn: &Arc<Mutex<Connection>>,
@@ -398,6 +461,26 @@ fn bytes_to_vec(bytes: &[u8]) -> Vec<f32> {
.collect()
}
fn decode_embedding_row(bytes: &[u8], dim: i64) -> anyhow::Result<Option<Vec<f32>>> {
if dim < 0 {
anyhow::bail!("event embedding has negative dimension {dim}");
}
if !bytes.len().is_multiple_of(4) {
anyhow::bail!(
"event embedding blob length {} not a multiple of 4",
bytes.len()
);
}
let vector = bytes_to_vec(bytes);
if vector.len() != dim as usize {
anyhow::bail!(
"event embedding dimension mismatch: dim column says {dim}, blob contains {} floats",
vector.len()
);
}
Ok(Some(vector))
}
#[cfg(test)]
#[path = "events_tests.rs"]
mod tests;
@@ -210,3 +210,52 @@ fn event_fts_matches_subject_field() {
assert_eq!(by_subject.len(), 1, "FTS should match on subject field");
assert_eq!(by_subject[0].event_id, "evt-subj");
}
#[test]
fn event_embeddings_are_scoped_by_model_signature() {
let conn = setup_db();
let event = EventRecord {
event_id: "evt-embed".into(),
segment_id: "seg-1".into(),
session_id: "s1".into(),
namespace: "global".into(),
event_type: EventType::Fact,
content: "The user prefers Korean summaries".into(),
subject: Some("language preference".into()),
timestamp_ref: None,
confidence: 0.9,
embedding: None,
source_turn_ids: None,
created_at: 1000.0,
};
event_insert(&conn, &event).unwrap();
event_embedding_upsert(
&conn,
"evt-embed",
"openai/text-embedding-3-small@1536",
&[0.1, 0.2],
1001.0,
)
.unwrap();
event_embedding_upsert(
&conn,
"evt-embed",
"local/bge-small@384",
&[0.3, 0.4, 0.5],
1002.0,
)
.unwrap();
assert_eq!(
event_embedding_get(&conn, "evt-embed", "openai/text-embedding-3-small@1536").unwrap(),
Some(vec![0.1, 0.2])
);
assert_eq!(
event_embedding_get(&conn, "evt-embed", "local/bge-small@384").unwrap(),
Some(vec![0.3, 0.4, 0.5])
);
assert!(event_embedding_get(&conn, "evt-embed", "missing/model@1")
.unwrap()
.is_none());
}
@@ -37,6 +37,22 @@ CREATE INDEX IF NOT EXISTS idx_segments_namespace
CREATE INDEX IF NOT EXISTS idx_segments_status
ON conversation_segments(status, session_id);
-- Per-model segment embeddings for #1574. The legacy
-- `conversation_segments.embedding` column stays in place during staged
-- migration; this table lets provider/model switches become query-time
-- filters instead of destructive rewrites.
CREATE TABLE IF NOT EXISTS segment_embeddings (
segment_id TEXT NOT NULL REFERENCES conversation_segments(segment_id) ON DELETE CASCADE,
model_signature TEXT NOT NULL,
vector BLOB NOT NULL,
dim INTEGER NOT NULL,
created_at REAL NOT NULL,
PRIMARY KEY (segment_id, model_signature)
);
CREATE INDEX IF NOT EXISTS idx_segment_embeddings_model
ON segment_embeddings(model_signature);
"#;
/// Segment status lifecycle: open → closed → summarised.
@@ -258,6 +274,55 @@ pub fn segment_set_embedding(
Ok(())
}
/// Store a segment embedding for a specific provider/model/dimension signature.
///
/// This writes only the per-model table introduced for #1574. The legacy
/// `conversation_segments.embedding` column remains available for dual-read
/// fallback while query paths migrate.
pub fn segment_embedding_upsert(
conn: &Arc<Mutex<Connection>>,
segment_id: &str,
model_signature: &str,
embedding: &[f32],
created_at: f64,
) -> anyhow::Result<()> {
let bytes = vec_to_bytes(embedding);
let dim = i64::try_from(embedding.len())?;
let conn = conn.lock();
conn.execute(
"INSERT INTO segment_embeddings (segment_id, model_signature, vector, dim, created_at)
VALUES (?1, ?2, ?3, ?4, ?5)
ON CONFLICT(segment_id, model_signature) DO UPDATE SET
vector = excluded.vector,
dim = excluded.dim,
created_at = excluded.created_at",
params![segment_id, model_signature, bytes, dim, created_at],
)?;
Ok(())
}
/// Fetch a segment embedding for exactly one provider/model/dimension signature.
pub fn segment_embedding_get(
conn: &Arc<Mutex<Connection>>,
segment_id: &str,
model_signature: &str,
) -> anyhow::Result<Option<Vec<f32>>> {
let conn = conn.lock();
let row: Option<(Vec<u8>, i64)> = conn
.query_row(
"SELECT vector, dim
FROM segment_embeddings
WHERE segment_id = ?1 AND model_signature = ?2",
params![segment_id, model_signature],
|r| Ok((r.get(0)?, r.get(1)?)),
)
.optional()?;
match row {
None => Ok(None),
Some((bytes, dim)) => decode_embedding_row(&bytes, dim),
}
}
/// Store topic keywords for the segment.
pub fn segment_set_keywords(
conn: &Arc<Mutex<Connection>>,
@@ -502,6 +567,26 @@ fn bytes_to_vec(bytes: &[u8]) -> Vec<f32> {
.collect()
}
fn decode_embedding_row(bytes: &[u8], dim: i64) -> anyhow::Result<Option<Vec<f32>>> {
if dim < 0 {
anyhow::bail!("segment embedding has negative dimension {dim}");
}
if !bytes.len().is_multiple_of(4) {
anyhow::bail!(
"segment embedding blob length {} not a multiple of 4",
bytes.len()
);
}
let floats = bytes_to_vec(bytes);
if floats.len() != dim as usize {
anyhow::bail!(
"segment embedding dimension mismatch: dim column says {dim}, blob contains {} floats",
floats.len()
);
}
Ok(Some(floats))
}
#[cfg(test)]
#[path = "segments_tests.rs"]
mod tests;
@@ -21,6 +21,44 @@ fn create_and_get_segment() {
assert_eq!(seg.status, SegmentStatus::Open);
}
#[test]
fn segment_embeddings_are_scoped_by_model_signature() {
let conn = setup_db();
segment_create(&conn, "seg-embed", "s1", "global", 1, 1000.0, 1000.0).unwrap();
segment_embedding_upsert(
&conn,
"seg-embed",
"openai/text-embedding-3-small@1536",
&[0.1, 0.2],
1001.0,
)
.unwrap();
segment_embedding_upsert(
&conn,
"seg-embed",
"local/bge-small@384",
&[0.3, 0.4, 0.5],
1002.0,
)
.unwrap();
assert_eq!(
segment_embedding_get(&conn, "seg-embed", "openai/text-embedding-3-small@1536").unwrap(),
Some(vec![0.1, 0.2])
);
assert_eq!(
segment_embedding_get(&conn, "seg-embed", "local/bge-small@384").unwrap(),
Some(vec![0.3, 0.4, 0.5])
);
assert!(segment_embedding_get(&conn, "seg-embed", "missing/model@1")
.unwrap()
.is_none());
let legacy_segment = segment_get(&conn, "seg-embed").unwrap().unwrap();
assert!(legacy_segment.embedding.is_none());
}
#[test]
fn append_and_close_segment() {
let conn = setup_db();
+106
View File
@@ -71,6 +71,21 @@ CREATE INDEX IF NOT EXISTS idx_mem_tree_chunks_owner
CREATE INDEX IF NOT EXISTS idx_mem_tree_chunks_source_seq
ON mem_tree_chunks(source_kind, source_id, seq_in_source);
-- Per-(chunk, embedding model) vectors (#1574). The legacy
-- mem_tree_chunks.embedding column remains in place during the dual-write
-- migration; this table lets multiple vector spaces coexist safely.
CREATE TABLE IF NOT EXISTS mem_tree_chunk_embeddings (
chunk_id TEXT NOT NULL REFERENCES mem_tree_chunks(id) ON DELETE CASCADE,
model_signature TEXT NOT NULL,
vector BLOB NOT NULL,
dim INTEGER NOT NULL,
created_at REAL NOT NULL,
PRIMARY KEY (chunk_id, model_signature)
);
CREATE INDEX IF NOT EXISTS idx_mem_tree_chunk_embeddings_model
ON mem_tree_chunk_embeddings(model_signature);
-- Phase 2 (#708): per-chunk score rationale for admission debugging.
CREATE TABLE IF NOT EXISTS mem_tree_score (
chunk_id TEXT PRIMARY KEY,
@@ -163,6 +178,21 @@ CREATE INDEX IF NOT EXISTS idx_mem_tree_summaries_sealed_at
CREATE INDEX IF NOT EXISTS idx_mem_tree_summaries_deleted
ON mem_tree_summaries(deleted);
-- Per-(summary, embedding model) vectors (#1574). Kept separate from the
-- legacy mem_tree_summaries.embedding column so provider/model switches can
-- be query-time filters instead of destructive rewrites.
CREATE TABLE IF NOT EXISTS mem_tree_summary_embeddings (
summary_id TEXT NOT NULL REFERENCES mem_tree_summaries(id) ON DELETE CASCADE,
model_signature TEXT NOT NULL,
vector BLOB NOT NULL,
dim INTEGER NOT NULL,
created_at REAL NOT NULL,
PRIMARY KEY (summary_id, model_signature)
);
CREATE INDEX IF NOT EXISTS idx_mem_tree_summary_embeddings_model
ON mem_tree_summary_embeddings(model_signature);
-- `mem_tree_buffers` holds the unsealed frontier per (tree, level). One row
-- per active level per tree; deleted when the buffer seals (clears) in the
-- same transaction as the new summary node row.
@@ -933,6 +963,58 @@ pub fn set_chunk_embedding(config: &Config, chunk_id: &str, embedding: &[f32]) -
})
}
/// Store a chunk embedding for a specific provider/model/dimension signature.
///
/// This is the Stage-1 per-model table write path for #1574. The legacy
/// `mem_tree_chunks.embedding` column is intentionally left untouched by this
/// helper so callers can dual-write while query paths migrate.
pub fn set_chunk_embedding_for_signature(
config: &Config,
chunk_id: &str,
model_signature: &str,
embedding: &[f32],
) -> Result<()> {
let bytes = embedding_to_blob(embedding);
let dim = i64::try_from(embedding.len()).context("embedding dimension does not fit i64")?;
let created_at = Utc::now().timestamp_millis() as f64 / 1000.0;
with_connection(config, |conn| {
conn.execute(
"INSERT INTO mem_tree_chunk_embeddings
(chunk_id, model_signature, vector, dim, created_at)
VALUES (?1, ?2, ?3, ?4, ?5)
ON CONFLICT(chunk_id, model_signature) DO UPDATE SET
vector = excluded.vector,
dim = excluded.dim,
created_at = excluded.created_at",
rusqlite::params![chunk_id, model_signature, bytes, dim, created_at],
)?;
Ok(())
})
}
/// Fetch a chunk embedding for exactly one provider/model/dimension signature.
pub fn get_chunk_embedding_for_signature(
config: &Config,
chunk_id: &str,
model_signature: &str,
) -> Result<Option<Vec<f32>>> {
with_connection(config, |conn| {
let row: Option<(Vec<u8>, i64)> = conn
.query_row(
"SELECT vector, dim
FROM mem_tree_chunk_embeddings
WHERE chunk_id = ?1 AND model_signature = ?2",
rusqlite::params![chunk_id, model_signature],
|r| Ok((r.get(0)?, r.get(1)?)),
)
.optional()?;
match row {
None => Ok(None),
Some((bytes, dim)) => embedding_from_blob(&bytes, dim, "chunk embedding"),
}
})
}
/// Fetch a chunk's embedding, decoding the stored little-endian `f32` blob.
///
/// Returns `Ok(None)` if the chunk doesn't exist or has no embedding stored.
@@ -961,6 +1043,30 @@ pub fn get_chunk_embedding(config: &Config, chunk_id: &str) -> Result<Option<Vec
})
}
fn embedding_to_blob(embedding: &[f32]) -> Vec<u8> {
embedding.iter().flat_map(|f| f.to_le_bytes()).collect()
}
fn embedding_from_blob(bytes: &[u8], dim: i64, label: &str) -> Result<Option<Vec<f32>>> {
if dim < 0 {
anyhow::bail!("{label} has negative dimension {dim}");
}
if !bytes.len().is_multiple_of(4) {
anyhow::bail!("{label} blob length {} not a multiple of 4", bytes.len());
}
let floats: Vec<f32> = bytes
.chunks_exact(4)
.map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect();
if floats.len() != dim as usize {
anyhow::bail!(
"{label} dimension mismatch: dim column says {dim}, blob contains {} floats",
floats.len()
);
}
Ok(Some(floats))
}
#[cfg(test)]
#[path = "store_tests.rs"]
mod tests;
+33
View File
@@ -71,6 +71,39 @@ fn reingest_preserves_existing_embedding() {
assert_eq!(got.token_count, 99);
}
#[test]
fn chunk_embeddings_are_scoped_by_model_signature() {
let (_tmp, cfg) = test_config();
let c = sample_chunk("slack:#eng", 0, 1_700_000_000_000);
upsert_chunks(&cfg, &[c.clone()]).unwrap();
set_chunk_embedding_for_signature(
&cfg,
&c.id,
"openai/text-embedding-3-small@1536",
&[0.1, 0.2],
)
.unwrap();
set_chunk_embedding_for_signature(&cfg, &c.id, "local/bge-small@384", &[0.3, 0.4, 0.5])
.unwrap();
assert_eq!(
get_chunk_embedding_for_signature(&cfg, &c.id, "openai/text-embedding-3-small@1536")
.unwrap(),
Some(vec![0.1, 0.2])
);
assert_eq!(
get_chunk_embedding_for_signature(&cfg, &c.id, "local/bge-small@384").unwrap(),
Some(vec![0.3, 0.4, 0.5])
);
assert!(
get_chunk_embedding_for_signature(&cfg, &c.id, "missing/model@1")
.unwrap()
.is_none()
);
assert!(get_chunk_embedding(&cfg, &c.id).unwrap().is_none());
}
#[test]
fn list_filters_by_source_kind() {
let (_tmp, cfg) = test_config();
@@ -290,6 +290,68 @@ pub fn get_summary_embedding(config: &Config, summary_id: &str) -> Result<Option
})
}
/// Store a summary embedding for a specific provider/model/dimension signature.
///
/// This writes the #1574 per-model table only; the legacy
/// `mem_tree_summaries.embedding` column remains available for dual-read
/// fallback while query paths migrate.
pub fn set_summary_embedding_for_signature(
config: &Config,
summary_id: &str,
model_signature: &str,
embedding: &[f32],
) -> Result<()> {
let blob = pack_embedding_blob(embedding);
let dim = i64::try_from(embedding.len()).context("embedding dimension does not fit i64")?;
let created_at = Utc::now().timestamp_millis() as f64 / 1000.0;
with_connection(config, |conn| {
conn.execute(
"INSERT INTO mem_tree_summary_embeddings
(summary_id, model_signature, vector, dim, created_at)
VALUES (?1, ?2, ?3, ?4, ?5)
ON CONFLICT(summary_id, model_signature) DO UPDATE SET
vector = excluded.vector,
dim = excluded.dim,
created_at = excluded.created_at",
params![summary_id, model_signature, blob, dim, created_at],
)?;
Ok(())
})
}
/// Fetch a summary embedding for exactly one provider/model/dimension signature.
pub fn get_summary_embedding_for_signature(
config: &Config,
summary_id: &str,
model_signature: &str,
) -> Result<Option<Vec<f32>>> {
with_connection(config, |conn| {
let row: Option<(Option<Vec<u8>>, i64)> = conn
.query_row(
"SELECT vector, dim
FROM mem_tree_summary_embeddings
WHERE summary_id = ?1 AND model_signature = ?2",
params![summary_id, model_signature],
|r| Ok((Some(r.get(0)?), r.get(1)?)),
)
.optional()?;
match row {
None => Ok(None),
Some((blob, dim)) => {
let decoded =
decode_signature_blob(blob, dim, &format!("summary_id={summary_id}"))?;
if decoded.as_ref().is_some_and(|v| v.len() != dim as usize) {
anyhow::bail!(
"summary embedding dimension mismatch: dim column says {dim}, blob contains {} floats",
decoded.as_ref().map_or(0, Vec::len)
);
}
Ok(decoded)
}
}
})
}
/// Fetch one summary by id. Soft-deleted rows are returned with
/// `deleted = true` so callers can decide filtering policy.
pub fn get_summary(config: &Config, id: &str) -> Result<Option<SummaryNode>> {
@@ -517,6 +579,33 @@ fn row_to_buffer(row: &rusqlite::Row<'_>) -> rusqlite::Result<Buffer> {
})
}
fn pack_embedding_blob(embedding: &[f32]) -> Vec<u8> {
embedding.iter().flat_map(|f| f.to_le_bytes()).collect()
}
fn decode_signature_blob(blob: Option<Vec<u8>>, dim: i64, label: &str) -> Result<Option<Vec<f32>>> {
let Some(bytes) = blob else {
return Ok(None);
};
if dim < 0 {
anyhow::bail!("{label} has negative dimension {dim}");
}
if !bytes.len().is_multiple_of(4) {
anyhow::bail!("{label} blob length {} not a multiple of 4", bytes.len());
}
let floats: Vec<f32> = bytes
.chunks_exact(4)
.map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect();
if floats.len() != dim as usize {
anyhow::bail!(
"summary embedding dimension mismatch: dim column says {dim}, blob contains {} floats",
floats.len()
);
}
Ok(Some(floats))
}
#[cfg(test)]
#[path = "store_tests.rs"]
mod tests;
@@ -102,6 +102,50 @@ fn summary_insert_is_idempotent_on_id() {
assert_eq!(count_summaries(&cfg, "tree-1").unwrap(), 1);
}
#[test]
fn summary_embeddings_are_scoped_by_model_signature() {
let (_tmp, cfg) = test_config();
insert_tree(&cfg, &sample_tree("tree-1", "slack:#eng")).unwrap();
let node = sample_summary("sum-embed", "tree-1", 1);
with_connection(&cfg, |conn| {
let tx = conn.unchecked_transaction()?;
insert_summary_tx(&tx, &node, None)?;
tx.commit()?;
Ok(())
})
.unwrap();
set_summary_embedding_for_signature(
&cfg,
"sum-embed",
"openai/text-embedding-3-small@1536",
&[0.1, 0.2],
)
.unwrap();
set_summary_embedding_for_signature(&cfg, "sum-embed", "local/bge-small@384", &[0.3, 0.4, 0.5])
.unwrap();
assert_eq!(
get_summary_embedding_for_signature(
&cfg,
"sum-embed",
"openai/text-embedding-3-small@1536",
)
.unwrap(),
Some(vec![0.1, 0.2])
);
assert_eq!(
get_summary_embedding_for_signature(&cfg, "sum-embed", "local/bge-small@384").unwrap(),
Some(vec![0.3, 0.4, 0.5])
);
assert!(
get_summary_embedding_for_signature(&cfg, "sum-embed", "missing/model@1")
.unwrap()
.is_none()
);
assert!(get_summary_embedding(&cfg, "sum-embed").unwrap().is_none());
}
#[test]
fn buffer_upsert_and_clear() {
let (_tmp, cfg) = test_config();