use crate::{storage::{self, KnowledgeBase}, text, types::*, Error, Result};
use parking_lot::{Condvar, Mutex};
use rusqlite::{params, params_from_iter, types::Value as SqlValue, Connection, OptionalExtension};
use serde::{Deserialize, Serialize};
use std::{cmp::Ordering, collections::{BTreeMap, BinaryHeap, HashMap, HashSet},
sync::{atomic::{AtomicBool, Ordering as AtomicOrdering}, Arc, Weak}, thread::JoinHandle};
fn text_version() -> u32 { 1 }
fn default_encoding() -> String { "sq8".into() }
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct EmbeddingSpace {
pub id: String,
pub model: String,
pub dimension: usize,
#[serde(default = "text_version")] pub text_version: u32,
#[serde(default = "default_encoding")] pub encoding: String,
}
#[derive(Debug, Clone)]
pub(crate) struct EmbeddingInput { pub key: RecordKey, pub text: String, pub fingerprint: String }
#[derive(Debug, Clone)]
pub(crate) struct EmbeddingWrite { pub key: RecordKey, pub fingerprint: String, pub values: Vec<f32> }
#[derive(Clone)]
pub struct EmbeddingStore(pub(crate) KnowledgeBase);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum EmbedErrorKind {
TooLarge,
RateLimited,
Other,
}
impl EmbedErrorKind {
pub fn code(self) -> &'static str {
match self { Self::TooLarge => "too_large", Self::RateLimited => "rate_limited", Self::Other => "other" }
}
pub fn from_code(code: &str) -> Option<Self> {
match code { "too_large" => Some(Self::TooLarge), "rate_limited" => Some(Self::RateLimited), "other" => Some(Self::Other), _ => None }
}
}
#[derive(Debug, Clone)]
pub struct EmbedCallbackError { pub kind: EmbedErrorKind, pub message: String }
impl EmbedCallbackError {
pub fn new(kind: EmbedErrorKind, message: impl Into<String>) -> Self { Self { kind, message: message.into() } }
pub fn too_large(message: impl Into<String>) -> Self { Self::new(EmbedErrorKind::TooLarge, message) }
pub fn rate_limited(message: impl Into<String>) -> Self { Self::new(EmbedErrorKind::RateLimited, message) }
pub fn other(message: impl Into<String>) -> Self { Self::new(EmbedErrorKind::Other, message) }
}
impl std::fmt::Display for EmbedCallbackError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { write!(f, "{}: {}", self.kind.code(), self.message) }
}
impl std::error::Error for EmbedCallbackError {}
pub trait Embedder: Send {
fn embed(&mut self, texts: &[String]) -> std::result::Result<Vec<Vec<f32>>, EmbedCallbackError>;
}
impl<F> Embedder for F
where F: FnMut(&[String]) -> std::result::Result<Vec<Vec<f32>>, EmbedCallbackError> + Send {
fn embed(&mut self, texts: &[String]) -> std::result::Result<Vec<Vec<f32>>, EmbedCallbackError> { self(texts) }
}
fn default_max_batch() -> usize { 32 }
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub struct EmbedderOptions {
#[serde(default = "default_max_batch")] pub max_batch: usize,
#[serde(default)] pub max_tokens_per_text: Option<usize>,
}
impl Default for EmbedderOptions {
fn default() -> Self { Self { max_batch: default_max_batch(), max_tokens_per_text: None } }
}
pub(crate) struct EmbedderEntry {
pub options: EmbedderOptions,
pub effective_batch: usize,
pub embedder: Box<dyn Embedder>,
}
impl EmbedderEntry {
pub(crate) fn embed(&mut self, texts: &[String]) -> std::result::Result<Vec<Vec<f32>>, EmbedCallbackError> {
match self.options.max_tokens_per_text {
Some(budget) => {
let budgeted: Vec<String> = texts.iter().map(|value| text::truncate_to_tokens(value, budget)).collect();
self.embedder.embed(&budgeted)
}
None => self.embedder.embed(texts),
}
}
}
#[derive(Default)]
pub(crate) struct EmbedderRegistry { entries: Mutex<HashMap<String, Arc<Mutex<EmbedderEntry>>>> }
impl EmbedderRegistry {
pub fn new() -> Self { Self::default() }
pub fn space_ids(&self) -> Vec<String> {
let mut ids: Vec<String> = self.entries.lock().keys().cloned().collect();
ids.sort();
ids
}
pub fn get(&self, space_id: &str) -> Option<Arc<Mutex<EmbedderEntry>>> { self.entries.lock().get(space_id).cloned() }
pub fn register(&self, space_id: String, entry: EmbedderEntry) { self.entries.lock().insert(space_id, Arc::new(Mutex::new(entry))); }
pub fn remove(&self, space_id: &str) -> bool { self.entries.lock().remove(space_id).is_some() }
}
const SAMPLE_TEXTS: [&str; 3] = ["样本一 sample", "样本二 sample", "样本三 sample"];
const RATE_LIMIT_ATTEMPTS: u32 = 3;
const RATE_LIMIT_BACKOFF_MS: u64 = 20;
pub(crate) fn get_space(conn: &Connection, id: &str) -> Result<EmbeddingSpace> {
conn.query_row("SELECT id,model,dimension,text_version,encoding FROM embedding_spaces WHERE id=?1", [id], |r|
Ok(EmbeddingSpace { id: r.get(0)?, model: r.get(1)?, dimension: r.get::<_, u32>(2)? as usize, text_version: r.get(3)?, encoding: r.get(4)? })).optional()?
.ok_or_else(|| Error::NotFound(format!("embedding space {id}")))
}
pub(crate) fn normalize(values: &[f32], dimension: usize) -> Result<Vec<f32>> {
if values.len() != dimension { return Err(Error::InvalidVector(format!("expected dimension {dimension}, received {}", values.len()))); }
if values.iter().any(|v| !v.is_finite()) { return Err(Error::InvalidVector("values must be finite".into())); }
let norm = values.iter().map(|v| f64::from(*v).powi(2)).sum::<f64>().sqrt();
if norm == 0.0 || !norm.is_finite() { return Err(Error::InvalidVector("zero or invalid vector norm".into())); }
Ok(values.iter().map(|v| (f64::from(*v) / norm) as f32).collect())
}
#[inline]
fn dot(a: &[f32], b: &[f32]) -> f32 {
let mut acc = [0f32; 8];
let chunks = a.len() / 8;
for c in 0..chunks {
let o = c * 8;
for k in 0..8 { acc[k] += a[o + k] * b[o + k]; }
}
let mut sum = acc.iter().sum::<f32>();
for i in chunks * 8..a.len() { sum += a[i] * b[i]; }
sum
}
fn encode_query_sq8(query: &[f32]) -> (Vec<i8>, f32) {
let max = query.iter().fold(0f32, |m, v| m.max(v.abs()));
let scale = if max == 0.0 { 1.0 } else { max / 127.0 };
let codes = query.iter().map(|v| (v / scale).round().clamp(-127.0, 127.0) as i8).collect();
(codes, scale)
}
#[inline]
fn dot_codes(left: &[i8], right: &[i8]) -> i32 {
#[cfg(target_arch = "x86_64")]
{
if std::is_x86_feature_detected!("avx2") { return unsafe { dot_codes_avx2(left, right) }; }
}
dot_codes_scalar(left, right)
}
fn dot_codes_scalar(left: &[i8], right: &[i8]) -> i32 {
left.iter().zip(right).map(|(a, b)| i32::from(*a) * i32::from(*b)).sum()
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn dot_codes_avx2(left: &[i8], right: &[i8]) -> i32 {
use std::arch::x86_64::*;
let mut acc = _mm256_setzero_si256();
let chunks = left.len() / 16;
for c in 0..chunks {
let o = c * 16;
let a = _mm_loadu_si128(left.as_ptr().add(o) as *const __m128i);
let b = _mm_loadu_si128(right.as_ptr().add(o) as *const __m128i);
acc = _mm256_add_epi32(acc, _mm256_madd_epi16(_mm256_cvtepi8_epi16(a), _mm256_cvtepi8_epi16(b)));
}
let mut total = {
let sum = _mm_add_epi32(_mm256_castsi256_si128(acc), _mm256_extracti128_si256(acc, 1));
let sum = _mm_add_epi32(sum, _mm_shuffle_epi32(sum, 0b01_00_11_10));
let sum = _mm_add_epi32(sum, _mm_shuffle_epi32(sum, 0b10_11_00_01));
_mm_cvtsi128_si32(sum)
};
for i in chunks * 16..left.len() { total += i32::from(left[i]) * i32::from(right[i]); }
total
}
fn encode_values(normalized: &[f32], encoding: &str) -> Result<Vec<u8>> {
match encoding {
"f32" => Ok(normalized.iter().flat_map(|v| v.to_le_bytes()).collect()),
"sq8" => {
let max = normalized.iter().fold(0f32, |m, v| m.max(v.abs()));
let scale = if max == 0.0 { 1.0 } else { max / 127.0 };
let mut out = Vec::with_capacity(4 + normalized.len());
out.extend_from_slice(&scale.to_le_bytes());
for v in normalized {
out.push((v / scale).round().clamp(-127.0, 127.0) as i8 as u8);
}
Ok(out)
}
other => Err(Error::Validation(format!("unknown encoding {other}"))),
}
}
fn decode_values(bytes: &[u8], dimension: usize, encoding: &str) -> Result<Vec<f32>> {
match encoding {
"f32" => {
if bytes.len() != dimension * 4 { return Err(Error::InvalidVector("stored vector dimension mismatch".into())); }
Ok(bytes.chunks_exact(4).map(|b| f32::from_le_bytes([b[0], b[1], b[2], b[3]])).collect())
}
"sq8" => {
let (scale, codes) = decode_sq8(bytes, dimension)?;
Ok(codes.into_iter().map(|c| f32::from(c) * scale).collect())
}
other => Err(Error::Validation(format!("unknown encoding {other}"))),
}
}
fn decode_sq8(bytes: &[u8], dimension: usize) -> Result<(f32, Vec<i8>)> {
if bytes.len() != dimension + 4 { return Err(Error::InvalidVector("stored vector dimension mismatch".into())); }
let scale = f32::from_le_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]);
Ok((scale, bytes[4..].iter().map(|b| *b as i8).collect()))
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum VectorizeTarget { Memory, Graph, Notes }
impl VectorizeTarget {
pub(crate) const ALL: [Self; 3] = [Self::Memory, Self::Graph, Self::Notes];
pub(crate) fn as_str(self) -> &'static str {
match self { Self::Memory => "memory", Self::Graph => "graph", Self::Notes => "notes" }
}
pub(crate) fn parse(value: &str) -> Result<Self> {
Self::ALL.into_iter().find(|target| target.as_str() == value)
.ok_or_else(|| Error::Validation(format!("vectorize target must be memory, graph or notes, got {value}")))
}
pub(crate) fn kinds(self) -> &'static [RecordKind] {
match self {
Self::Memory => &[RecordKind::Memory],
Self::Graph => &[RecordKind::Entity, RecordKind::Relation, RecordKind::Event],
Self::Notes => &[RecordKind::Chunk],
}
}
fn default_enabled(self) -> bool { !matches!(self, Self::Notes) }
}
fn vectorize_key(namespace: &str) -> String { format!("vectorize:{}", text::normalized_tag(namespace)) }
fn target_vectorize_key(namespace: &str, target: VectorizeTarget) -> String {
format!("{}:{}", vectorize_key(namespace), target.as_str())
}
pub(crate) fn namespace_vectorization(conn: &Connection, namespace: &str) -> Result<bool> {
let value: Option<i64> = conn.query_row("SELECT value FROM meta WHERE key=?1", [vectorize_key(namespace)], |r| r.get(0)).optional()?;
Ok(value != Some(0))
}
pub(crate) fn target_vectorization(conn: &Connection, namespace: &str, target: VectorizeTarget) -> Result<bool> {
let value: Option<i64> = conn.query_row("SELECT value FROM meta WHERE key=?1", [target_vectorize_key(namespace, target)], |r| r.get(0)).optional()?;
Ok(value.map_or_else(|| target.default_enabled(), |value| value != 0))
}
pub(crate) fn enabled_kinds(conn: &Connection, namespace: &str) -> Result<Vec<RecordKind>> {
if !namespace_vectorization(conn, namespace)? { return Ok(Vec::new()); }
let mut kinds = Vec::new();
for target in VectorizeTarget::ALL {
if target_vectorization(conn, namespace, target)? { kinds.extend_from_slice(target.kinds()); }
}
Ok(kinds)
}
pub(crate) fn ready_kinds(conn: &Connection, namespace: &str, space_id: &str) -> Result<Vec<RecordKind>> {
if !namespace_vectorization(conn, namespace)? { return Ok(Vec::new()); }
let mut kinds = Vec::new();
for target in VectorizeTarget::ALL {
if target_vectorization(conn, namespace, target)? && vector_ready(conn, namespace, space_id, target)? {
kinds.extend_from_slice(target.kinds());
}
}
Ok(kinds)
}
fn enabled_kinds_sql() -> String {
VectorizeTarget::ALL.iter().map(|target| target_enabled_sql(*target)).collect::<Vec<_>>().join(" OR ")
}
fn target_enabled_sql(target: VectorizeTarget) -> String {
let codes = target.kinds().iter().map(|kind| kind.code().to_string()).collect::<Vec<_>>().join(",");
format!("(r.kind IN ({codes}) AND COALESCE((SELECT value FROM meta WHERE key='vectorize:'||n.text||':{}'),{}) = 1)",
target.as_str(), i64::from(target.default_enabled()))
}
pub(crate) fn pending_candidates(conn: &Connection, index: &crate::index::TextIndex, space_id: &str, namespace: Option<&str>,
target: Option<VectorizeTarget>, limit: usize, after: Option<i64>, ids: Option<&[i64]>) -> Result<Vec<EmbeddingInput>> {
let mut sql = String::from("SELECT r.id,r.kind,r.payload_json,r.fingerprint FROM records r JOIN strings n ON n.id=r.namespace_id WHERE 1=1");
let mut values: Vec<SqlValue> = Vec::new();
if let Some(namespace) = namespace {
sql.push_str(" AND n.text=?");
values.push(SqlValue::Text(text::normalized_tag(namespace)));
}
sql.push_str(&format!(" AND ({})", match target { Some(target) => target_enabled_sql(target), None => enabled_kinds_sql() }));
sql.push_str(" AND COALESCE((SELECT value FROM meta WHERE key='vectorize:'||n.text),1)=1");
sql.push_str(" AND NOT EXISTS(SELECT 1 FROM embeddings e WHERE e.space_id=? AND e.record_id=r.id AND e.fingerprint=r.fingerprint)");
values.push(SqlValue::Text(space_id.into()));
if let Some(ids) = ids {
if ids.is_empty() { return Ok(Vec::new()); }
sql.push_str(&format!(" AND r.id IN ({})", vec!["?"; ids.len()].join(",")));
values.extend(ids.iter().map(|id| SqlValue::Integer(*id)));
}
if let Some(cursor) = after {
sql.push_str(" AND r.id>?");
values.push(SqlValue::Integer(cursor));
}
sql.push_str(" ORDER BY r.id LIMIT ?");
values.push(SqlValue::Integer(limit as i64));
let mut stmt = conn.prepare(&sql)?;
let mut items = Vec::new();
let mut candidates: Vec<(i64, i64, String, String)> = Vec::new();
for row in stmt.query_map(params_from_iter(values), |r| Ok((r.get::<_, i64>(0)?, r.get::<_, i64>(1)?, r.get::<_, String>(2)?, r.get::<_, String>(3)?)))? {
candidates.push(row?);
}
let chunk_ids: Vec<i64> = candidates.iter().filter(|(_, kind, _, _)| *kind == RecordKind::Chunk.code()).map(|(id, _, _, _)| *id).collect();
let bodies = if chunk_ids.is_empty() { BTreeMap::new() } else { index.bodies(&chunk_ids)? };
for (id, kind_code, payload_json, fingerprint) in candidates {
let kind = RecordKind::from_code(kind_code).ok_or_else(|| Error::Validation("invalid stored record kind".into()))?;
let payload: Option<serde_json::Value> = serde_json::from_str(&payload_json).ok();
let body = match kind {
RecordKind::Chunk => bodies.get(&id).cloned().unwrap_or_default(),
_ => payload.as_ref().map(|payload| storage::record_text(kind, payload)).unwrap_or_default(),
};
let name = match (kind, &payload) {
(RecordKind::Entity, Some(payload)) => payload.get("name").and_then(serde_json::Value::as_str).unwrap_or("").to_string(),
_ => String::new(),
};
let text = match (kind, record_tags(conn, id)?) {
(RecordKind::Memory, tags) if !tags.is_empty() => format!("{body}\n{}", tags.join(" ")),
(RecordKind::Entity, _) if !name.is_empty() => format!("{name}\n{body}"),
_ => body,
};
items.push(EmbeddingInput { key: RecordKey { id }, text, fingerprint });
}
Ok(items)
}
fn embeddable(input: &EmbeddingInput) -> bool { !input.text.trim().is_empty() }
const GAP_PROBE: usize = 256;
fn has_gap(conn: &Connection, index: &crate::index::TextIndex, space_id: &str, namespace: &str, target: VectorizeTarget) -> Result<bool> {
let mut cursor: Option<i64> = None;
loop {
let candidates = pending_candidates(conn, index, space_id, Some(namespace), Some(target), GAP_PROBE, cursor, None)?;
let Some(last) = candidates.last() else { return Ok(false) };
if candidates.iter().any(embeddable) { return Ok(true); }
cursor = Some(last.key.id);
}
}
fn vector_ready_key(namespace: &str, space_id: &str, target: VectorizeTarget) -> String {
format!("vector_ready:{}:{}:{}", text::normalized_tag(namespace), space_id, target.as_str())
}
pub(crate) fn vector_ready(conn: &Connection, namespace: &str, space_id: &str, target: VectorizeTarget) -> Result<bool> {
let value: Option<i64> = conn.query_row("SELECT value FROM meta WHERE key=?1", [vector_ready_key(namespace, space_id, target)], |r| r.get(0)).optional()?;
Ok(value == Some(1))
}
fn set_vector_ready(conn: &Connection, namespace: &str, space_id: &str, target: VectorizeTarget, ready: bool) -> Result<()> {
if ready {
conn.execute("INSERT INTO meta(key,value) VALUES (?1,1) ON CONFLICT(key) DO UPDATE SET value=1", [vector_ready_key(namespace, space_id, target)])?;
} else {
conn.execute("DELETE FROM meta WHERE key=?1", [vector_ready_key(namespace, space_id, target)])?;
}
Ok(())
}
pub(crate) fn clear_vector_ready(conn: &Connection, namespace: &str) -> Result<()> {
let prefix = format!("vector_ready:{}:", text::normalized_tag(namespace));
conn.execute("DELETE FROM meta WHERE substr(key,1,?1)=?2", params![prefix.chars().count() as i64, prefix])?;
Ok(())
}
fn record_tags(conn: &Connection, id: i64) -> Result<Vec<String>> {
let mut stmt = conn.prepare("SELECT t.text FROM record_tags rt JOIN strings t ON t.id=rt.tag_id WHERE rt.record_id=?1 ORDER BY t.text")?;
let mut tags = Vec::new();
for row in stmt.query_map([id], |r| r.get::<_, String>(0))? { tags.push(row?); }
Ok(tags)
}
enum EmbedOutcome { Vectors(Vec<Vec<f32>>), Shrunk, Failed(String) }
fn embed_with_retry(entry: &mut EmbedderEntry, texts: &[String]) -> EmbedOutcome {
let mut attempts = 0u32;
loop {
match entry.embed(texts) {
Ok(values) => return EmbedOutcome::Vectors(values),
Err(error) if error.kind == EmbedErrorKind::TooLarge => {
if entry.effective_batch <= 1 { return EmbedOutcome::Failed(format!("batch size 1 was still rejected: {}", error.message)); }
entry.effective_batch /= 2;
return EmbedOutcome::Shrunk;
}
Err(error) if error.kind == EmbedErrorKind::RateLimited && attempts < RATE_LIMIT_ATTEMPTS => {
attempts += 1;
std::thread::sleep(std::time::Duration::from_millis(RATE_LIMIT_BACKOFF_MS * u64::from(attempts)));
}
Err(error) => return EmbedOutcome::Failed(error.message),
}
}
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct SyncReport { pub scanned: usize, pub written: usize, pub batches: usize, pub interrupted: Option<String> }
impl EmbeddingStore {
pub fn register_space(&self, space: EmbeddingSpace) -> Result<WriteReceipt<EmbeddingSpace>> {
storage::validate_identity("space id", &space.id)?;
storage::validate_identity("model", &space.model)?;
if !(1..=65_536).contains(&space.dimension) || space.text_version != 1 { return Err(Error::Validation("dimension must be 1..65536 and text_version must be 1".into())); }
if space.encoding != "f32" && space.encoding != "sq8" { return Err(Error::Validation("encoding must be \"f32\" or \"sq8\"".into())); }
self.0.mutate_meta(|tx| {
match get_space(tx, &space.id) {
Ok(old) if old == space => return Ok(old),
Ok(_) => return Err(Error::Conflict("embedding space is immutable; register a new ID for a new model or dimension".into())),
Err(Error::NotFound(_)) => (),
Err(err) => return Err(err),
}
tx.execute("INSERT INTO embedding_spaces(id,model,dimension,text_version,encoding) VALUES (?1,?2,?3,?4,?5)", params![space.id, space.model, space.dimension as i64, space.text_version, space.encoding])?;
Ok(space)
})
}
pub fn spaces(&self) -> Result<Vec<EmbeddingSpace>> {
let state = self.0.read()?;
let mut stmt = state.conn().prepare("SELECT id,model,dimension,text_version,encoding FROM embedding_spaces ORDER BY id")?;
let rows = stmt.query_map([], |r| Ok(EmbeddingSpace { id: r.get(0)?, model: r.get(1)?, dimension: r.get::<_, u32>(2)? as usize, text_version: r.get(3)?, encoding: r.get(4)? }))?;
Ok(rows.collect::<std::result::Result<Vec<_>, _>>()?)
}
pub fn register_embedder<F: Embedder + 'static>(&self, space_id: &str, embedder: F) -> Result<()> {
self.register_embedder_with(space_id, embedder, EmbedderOptions::default())
}
pub fn register_embedder_with<F: Embedder + 'static>(&self, space_id: &str, embedder: F, options: EmbedderOptions) -> Result<()> {
storage::validate_identity("space id", space_id)?;
if !(1..=10_000).contains(&options.max_batch) { return Err(Error::Validation("max_batch must be between 1 and 10000".into())); }
if options.max_tokens_per_text == Some(0) { return Err(Error::Validation("max_tokens_per_text must be positive".into())); }
let space = { let state = self.0.read()?; get_space(state.conn(), space_id)? };
let mut entry = EmbedderEntry { options, effective_batch: options.max_batch, embedder: Box::new(embedder) };
let samples: Vec<String> = SAMPLE_TEXTS.iter().take(3.min(entry.effective_batch)).map(|sample| (*sample).to_string()).collect();
let produced = entry.embed(&samples)
.map_err(|error| Error::Validation(format!("embedder failed during registration ({}): {}", error.kind.code(), error.message)))?;
validate_vectors(&produced, samples.len(), &space)?;
self.0.engine.embedders.register(space_id.to_string(), entry);
if let Some(vectorizer) = self.0.engine.vectorizer.get() { vectorizer.notify_work(); }
Ok(())
}
pub fn embedder_space(&self, space_id: &str) -> Result<Option<EmbeddingSpace>> {
let state = self.0.read()?;
match get_space(state.conn(), space_id) { Ok(space) => Ok(Some(space)), Err(Error::NotFound(_)) => Ok(None), Err(error) => Err(error) }
}
pub fn unregister_embedder(&self, space_id: &str) -> Result<bool> { Ok(self.0.engine.embedders.remove(space_id)) }
pub fn namespace_vectorization(&self, namespace: &str) -> Result<bool> {
storage::validate_identity("namespace", namespace)?;
let state = self.0.read()?;
namespace_vectorization(state.conn(), namespace)
}
pub fn set_namespace_vectorization(&self, namespace: &str, enabled: bool) -> Result<WriteReceipt<bool>> {
storage::validate_identity("namespace", namespace)?;
let receipt = self.0.mutate_meta(|tx| {
tx.execute("INSERT INTO meta(key,value) VALUES (?1,?2) ON CONFLICT(key) DO UPDATE SET value=excluded.value",
params![vectorize_key(namespace), i64::from(enabled)])?;
clear_vector_ready(tx, namespace)?;
Ok(enabled)
})?;
if let Some(vectorizer) = self.0.engine.vectorizer.get() { vectorizer.notify_work(); }
Ok(receipt)
}
pub fn vectorization(&self, namespace: &str, target: &str) -> Result<bool> {
storage::validate_identity("namespace", namespace)?;
let target = VectorizeTarget::parse(target)?;
let state = self.0.read()?;
target_vectorization(state.conn(), namespace, target)
}
pub fn set_vectorization(&self, namespace: &str, target: &str, enabled: bool) -> Result<WriteReceipt<bool>> {
storage::validate_identity("namespace", namespace)?;
let target = VectorizeTarget::parse(target)?;
let receipt = self.0.mutate_meta(|tx| {
tx.execute("INSERT INTO meta(key,value) VALUES (?1,?2) ON CONFLICT(key) DO UPDATE SET value=excluded.value",
params![target_vectorize_key(namespace, target), i64::from(enabled)])?;
clear_vector_ready(tx, namespace)?;
Ok(enabled)
})?;
if let Some(vectorizer) = self.0.engine.vectorizer.get() { vectorizer.notify_work(); }
Ok(receipt)
}
pub fn vector_ready(&self, namespace: &str, space_id: &str, target: &str) -> Result<bool> {
storage::validate_identity("namespace", namespace)?;
storage::validate_identity("space id", space_id)?;
let target = VectorizeTarget::parse(target)?;
let state = self.0.read()?;
vector_ready(state.conn(), namespace, space_id, target)
}
pub fn sync(&self, space_id: &str, batch: usize) -> Result<WriteReceipt<SyncReport>> {
storage::validate_limit(batch)?;
if self.0.engine.embedders.get(space_id).is_none() {
return Err(Error::Validation(format!("no embedder registered for space {space_id}")));
}
let report = self.fill(space_id, batch, true)?;
let state = self.0.read()?;
Ok(WriteReceipt { value: report, revision: storage::current_revision(state.conn())? })
}
fn fill(&self, space_id: &str, batch: usize, blocking: bool) -> Result<SyncReport> {
let _filling = if blocking { None } else { self.0.engine.vectorizer.get().map(|vectorizer| vectorizer.begin_fill()) };
let Some(entry) = self.0.engine.embedders.get(space_id) else { return Ok(SyncReport::default()) };
self.0.catch_up_index()?;
let report = self.drain(space_id, &entry, batch, blocking)?;
if report.interrupted.is_some() { self.0.note_degrade(Degrade::EmbedFailed); }
self.verify_and_mark(space_id)?;
Ok(report)
}
fn verify_and_mark(&self, space_id: &str) -> Result<()> {
for _ in 0..VERIFY_ATTEMPTS {
let (revision, marks) = {
let state = self.0.read()?;
let index = self.0.index()?;
let conn = state.conn();
let pending: i64 = conn.query_row("SELECT COUNT(*) FROM index_updates", [], |r| r.get(0))?;
if pending > 0 { return Ok(()); }
let revision = storage::current_revision(conn)?;
let mut marks = Vec::new();
for namespace in storage::record_namespaces(conn)? {
for target in VectorizeTarget::ALL {
let enabled = target_vectorization(conn, &namespace, target)?;
let ready = enabled && !has_gap(conn, &index, space_id, &namespace, target)?;
marks.push((namespace.clone(), target, ready));
}
}
(revision, marks)
};
let settled = { let state = self.0.read()?; storage::current_revision(state.conn())? == revision };
if !settled { continue; }
self.0.mutate_meta(|tx| {
let pending: i64 = tx.query_row("SELECT COUNT(*) FROM index_updates", [], |r| r.get(0))?;
if pending > 0 { return Ok(()); }
for (namespace, target, ready) in &marks { set_vector_ready(tx, namespace, space_id, *target, *ready)?; }
Ok(())
})?;
return Ok(());
}
Ok(())
}
fn drain(&self, space_id: &str, entry: &Arc<Mutex<EmbedderEntry>>, batch: usize, blocking: bool) -> Result<SyncReport> {
let mut report = SyncReport::default();
let acquired = if blocking { Some(entry.lock()) } else { entry.try_lock() };
let Some(mut guard) = acquired else { return Ok(report) };
let mut cursor: Option<i64> = None;
loop {
let limit = batch.min(guard.effective_batch).max(1);
let candidates = {
let state = self.0.read()?;
let index = self.0.index()?;
pending_candidates(state.conn(), &index, space_id, None, None, limit, cursor, None)?
};
let Some(last) = candidates.last().map(|input| input.key.id) else { break };
let pending: Vec<EmbeddingInput> = candidates.into_iter().filter(embeddable).collect();
if pending.is_empty() { cursor = Some(last); continue; }
report.scanned += pending.len();
let texts: Vec<String> = pending.iter().map(|input| input.text.clone()).collect();
let values = match embed_with_retry(&mut guard, &texts) {
EmbedOutcome::Shrunk => continue,
EmbedOutcome::Failed(message) => { report.interrupted = Some(message); break; }
EmbedOutcome::Vectors(values) => values,
};
if values.len() != pending.len() {
report.interrupted = Some(format!("embedder returned {} vectors for {} inputs", values.len(), pending.len()));
break;
}
let writes: Vec<EmbeddingWrite> = pending.iter().zip(values).map(|(input, vector)|
EmbeddingWrite { key: input.key, fingerprint: input.fingerprint.clone(), values: vector }).collect();
match self.put(space_id, &writes) {
Ok(receipt) => { report.written += receipt.value; report.batches += 1; }
Err(Error::StaleRevision(_)) => {}
Err(error) => return Err(error),
}
cursor = Some(last);
}
Ok(report)
}
pub(crate) fn put(&self, space_id: &str, writes: &[EmbeddingWrite]) -> Result<WriteReceipt<usize>> {
self.0.mutate(|tx| {
let space = get_space(tx, space_id)?;
for write in writes {
let actual: Option<String> = tx.query_row("SELECT fingerprint FROM records WHERE id=?1", [write.key.id], |r| r.get(0)).optional()?;
let actual = actual.ok_or_else(|| Error::NotFound(write.key.id.to_string()))?;
if actual != write.fingerprint { return Err(Error::StaleRevision(write.key.id.to_string())); }
let normalized = normalize(&write.values, space.dimension)?;
let bytes = encode_values(&normalized, &space.encoding)?;
tx.execute("INSERT INTO embeddings(space_id,record_id,fingerprint,vector) VALUES (?1,?2,?3,?4)
ON CONFLICT(space_id,record_id) DO UPDATE SET fingerprint=excluded.fingerprint,vector=excluded.vector",
params![space_id, write.key.id, write.fingerprint, bytes])?;
}
for write in writes { storage::touch_record_namespace(tx, write.key.id)?; }
Ok(writes.len())
})
}
pub fn delete_space(&self, id: &str) -> Result<WriteReceipt<bool>> {
self.0.mutate_meta(|tx| Ok(tx.execute("DELETE FROM embedding_spaces WHERE id=?1", [id])? > 0))
}
}
fn validate_vectors(produced: &[Vec<f32>], expected: usize, space: &EmbeddingSpace) -> Result<()> {
if produced.len() != expected {
return Err(Error::InvalidVector(format!("embedder returned {} vectors for {expected} inputs", produced.len())));
}
for values in produced {
normalize(values, space.dimension)
.map_err(|error| Error::InvalidVector(format!("embedder output does not satisfy space {}: {error}", space.id)))?;
}
Ok(())
}
const SWEEP_BATCH: usize = 32;
const VERIFY_ATTEMPTS: usize = 3;
pub(crate) struct Vectorizer {
stopping: Mutex<bool>,
signal: Condvar,
pending: AtomicBool,
filling: AtomicBool,
handle: Mutex<Option<JoinHandle<()>>>,
}
pub(crate) struct FillGuard(Arc<Vectorizer>);
impl Drop for FillGuard {
fn drop(&mut self) { self.0.filling.store(false, AtomicOrdering::SeqCst); }
}
impl Vectorizer {
pub(crate) fn start(engine: &Arc<crate::storage::Engine>) -> Result<Arc<Self>> {
let vectorizer = Arc::new(Self { stopping: Mutex::new(false), signal: Condvar::new(),
pending: AtomicBool::new(true), filling: AtomicBool::new(false), handle: Mutex::new(None) });
let worker = Arc::clone(&vectorizer);
let engine = Arc::downgrade(engine);
let handle = std::thread::Builder::new().name("p-memory-vectorize".into())
.spawn(move || work_loop(&engine, &worker))?;
*vectorizer.handle.lock() = Some(handle);
Ok(vectorizer)
}
fn begin_fill(self: &Arc<Self>) -> FillGuard {
self.filling.store(true, AtomicOrdering::SeqCst);
FillGuard(Arc::clone(self))
}
pub(crate) fn is_filling(&self) -> bool { self.filling.load(AtomicOrdering::SeqCst) }
pub(crate) fn notify_work(&self) {
let _guard = self.stopping.lock();
self.pending.store(true, AtomicOrdering::SeqCst);
self.signal.notify_all();
}
pub(crate) fn stop(&self) {
{
let mut stopping = self.stopping.lock();
*stopping = true;
self.signal.notify_all();
}
if let Some(handle) = self.handle.lock().take() { let _ = handle.join(); }
}
}
fn work_loop(engine: &Weak<crate::storage::Engine>, vectorizer: &Vectorizer) {
loop {
{
let mut stopping = vectorizer.stopping.lock();
while !*stopping && !vectorizer.pending.load(AtomicOrdering::SeqCst) {
vectorizer.signal.wait(&mut stopping);
}
if *stopping { return; }
}
vectorizer.pending.store(false, AtomicOrdering::SeqCst);
let Some(engine) = engine.upgrade() else { return };
let kb = KnowledgeBase { engine };
loop {
let mut progressed = false;
for space_id in kb.engine.embedders.space_ids() {
if let Ok(report) = EmbeddingStore(kb.clone()).fill(&space_id, SWEEP_BATCH, false) {
if report.written > 0 { progressed = true; }
}
}
if !progressed { break; }
}
}
}
struct VectorRow { key: RecordKey, kind: RecordKind, tags: Vec<String>, note: i64 }
enum PartitionData {
F32(Vec<f32>),
Sq8 { codes: Vec<i8>, scales: Vec<f32> },
}
pub(crate) struct Partition { dimension: usize, rows: Vec<VectorRow>, data: PartitionData }
struct Candidate { score: f64, key: RecordKey }
impl PartialEq for Candidate { fn eq(&self, other: &Self) -> bool { self.score == other.score && self.key == other.key } }
impl Eq for Candidate {}
impl PartialOrd for Candidate { fn partial_cmp(&self, other: &Self) -> Option<Ordering> { Some(self.cmp(other)) } }
impl Ord for Candidate {
fn cmp(&self, other: &Self) -> Ordering { other.score.total_cmp(&self.score).then_with(|| self.key.cmp(&other.key)) }
}
impl Partition {
pub fn load(conn: &Connection, space: &EmbeddingSpace, namespace: &str, scope: &str) -> Result<Option<Self>> {
let namespace = text::normalized_tag(namespace);
let scope = text::normalized_tag(scope);
let mut stmt = conn.prepare("SELECT r.id,r.kind,e.vector,
(SELECT json_group_array(t.text) FROM record_tags rt JOIN strings t ON t.id=rt.tag_id WHERE rt.record_id=r.id),
COALESCE((SELECT c.note_id FROM chunks c WHERE c.record_id=r.id),0)
FROM embeddings e JOIN records r ON r.id=e.record_id AND r.fingerprint=e.fingerprint
WHERE e.space_id=?1 AND r.namespace_id=(SELECT id FROM strings WHERE text=?2)
AND r.scope_id=(SELECT id FROM strings WHERE text=?3)")?;
let sq8 = space.encoding == "sq8";
let mut partition = Self {
dimension: space.dimension,
rows: vec![],
data: if sq8 { PartitionData::Sq8 { codes: vec![], scales: vec![] } } else { PartitionData::F32(vec![]) },
};
let mut rows = stmt.query(params![space.id, namespace, scope])?;
while let Some(row) = rows.next()? {
let key = RecordKey { id: row.get(0)? };
let kind = RecordKind::from_code(row.get::<_, i64>(1)?).ok_or_else(|| Error::InvalidVector("invalid stored record kind".into()))?;
let bytes: Vec<u8> = row.get(2)?;
let tags: Vec<String> = serde_json::from_str(&row.get::<_, String>(3)?)?;
let note: i64 = row.get(4)?;
match &mut partition.data {
PartitionData::F32(values) => {
let decoded = decode_values(&bytes, space.dimension, "f32")?;
if decoded.iter().any(|v| !v.is_finite()) { return Err(Error::InvalidVector("stored vector contains nonfinite values".into())); }
values.extend(decoded);
}
PartitionData::Sq8 { codes, scales } => {
let (scale, decoded) = decode_sq8(&bytes, space.dimension)?;
if !scale.is_finite() { return Err(Error::InvalidVector("stored vector contains nonfinite values".into())); }
scales.push(scale);
codes.extend(decoded);
}
}
partition.rows.push(VectorRow { key, kind, tags, note });
}
Ok(if partition.rows.is_empty() { None } else { Some(partition) })
}
fn matches(&self, row: &VectorRow, kinds: &[RecordKind], tags: &[String], note_ids: &[i64], allowed: Option<&HashSet<i64>>) -> bool {
if allowed.is_some_and(|set| !set.contains(&row.key.id)) { return false; }
if !note_ids.is_empty() && !note_ids.contains(&row.note) { return false; }
(kinds.is_empty() || kinds.contains(&row.kind)) && tags.iter().all(|t| row.tags.contains(t))
}
fn retain(heap: &mut BinaryHeap<Candidate>, key: RecordKey, scored: f64, limit: usize) {
let score = scored.clamp(-1.0, 1.0);
if heap.len() < limit { heap.push(Candidate { score, key }); }
else if let Some(worst) = heap.peek() {
if score > worst.score || (score == worst.score && key < worst.key) {
heap.pop(); heap.push(Candidate { score, key });
}
}
}
pub fn search(&self, query: &[f32], kinds: &[RecordKind], tags: &[String], note_ids: &[i64], limit: usize, allowed: Option<&HashSet<i64>>) -> Result<Vec<(RecordKey, f64)>> {
let query = normalize(query, self.dimension)?;
let mut heap = BinaryHeap::<Candidate>::new();
match &self.data {
PartitionData::F32(values) => {
for (i, row) in self.rows.iter().enumerate() {
if !self.matches(row, kinds, tags, note_ids, allowed) { continue; }
let offset = i * self.dimension;
Self::retain(&mut heap, row.key, f64::from(dot(&query, &values[offset..offset + self.dimension])), limit);
}
}
PartitionData::Sq8 { codes, scales } => {
let (query_codes, query_scale) = encode_query_sq8(&query);
for (i, row) in self.rows.iter().enumerate() {
if !self.matches(row, kinds, tags, note_ids, allowed) { continue; }
let offset = i * self.dimension;
let raw = query_scale * scales[i] * dot_codes(&query_codes, &codes[offset..offset + self.dimension]) as f32;
Self::retain(&mut heap, row.key, f64::from(raw), limit);
}
}
}
let mut result: Vec<_> = heap.into_iter().map(|c| (c.key, c.score)).collect();
result.sort_by(|a,b| b.1.total_cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
Ok(result)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering as AtomicOrdering};
#[test]
fn integer_kernel_matches_scalar() {
for dimension in [1usize, 15, 16, 17, 250, 1024] {
let left: Vec<i8> = (0..dimension).map(|i| ((i * 37 % 255) as i32 - 127) as i8).collect();
let right: Vec<i8> = (0..dimension).map(|i| ((i * 91 % 255) as i32 - 127) as i8).collect();
assert_eq!(dot_codes(&left, &right), dot_codes_scalar(&left, &right), "dimension {dimension}");
}
}
#[test]
fn tied_scores_break_by_key_whatever_the_row_order() {
let dimension = 4usize;
let query = vec![1.0f32, 0.0, 0.0, 0.0];
let build = |order: &[i64]| Partition {
dimension,
rows: order.iter().map(|id| VectorRow { key: RecordKey { id: *id }, kind: RecordKind::Memory, tags: vec![], note: 0 }).collect(),
data: PartitionData::F32(order.iter().flat_map(|_| [1.0f32, 0.0, 0.0, 0.0]).collect()),
};
let top_two = |partition: &Partition| -> Vec<i64> {
partition.search(&query, &[], &[], &[], 2, None).unwrap().into_iter().map(|(key, _)| key.id).collect()
};
assert_eq!(top_two(&build(&[1, 2, 3, 4])), vec![1, 2], "并列时取 key 最小的两条");
assert_eq!(top_two(&build(&[4, 3, 2, 1])), vec![1, 2], "换个行序,结果必须一样");
}
#[test]
fn quantized_query_tracks_decoded_f32_kernel() {
let dimension = 256usize;
let stored_raw: Vec<f32> = (0..dimension).map(|i| (i as f32 * 0.37).sin() + 0.25).collect();
let query_raw: Vec<f32> = (0..dimension).map(|i| (i as f32 * 0.11).cos() - 0.1).collect();
let stored = normalize(&stored_raw, dimension).unwrap();
let query = normalize(&query_raw, dimension).unwrap();
let encoded = encode_values(&stored, "sq8").unwrap();
let (scale, codes) = decode_sq8(&encoded, dimension).unwrap();
let decoded = decode_values(&encoded, dimension, "sq8").unwrap();
let reference = dot(&query, &decoded);
let (query_codes, query_scale) = encode_query_sq8(&query);
let actual = query_scale * scale * dot_codes(&query_codes, &codes) as f32;
assert!((reference - actual).abs() < 1e-3, "reference {reference} vs actual {actual}");
}
#[test]
fn oversized_batches_shrink_and_persist() {
let lengths = std::sync::Arc::new(Mutex::new(Vec::<usize>::new()));
let observed = lengths.clone();
let mut entry = EmbedderEntry {
options: EmbedderOptions { max_batch: 8, max_tokens_per_text: None },
effective_batch: 8,
embedder: Box::new(move |texts: &[String]| {
observed.lock().push(texts.len());
if texts.len() > 4 { return Err(EmbedCallbackError::too_large("too many texts")); }
Ok(texts.iter().map(|_| vec![1.0f32, 0.0]).collect())
}),
};
let texts: Vec<String> = (0..8).map(|i| format!("文本 {i}")).collect();
assert!(matches!(embed_with_retry(&mut entry, &texts), EmbedOutcome::Shrunk));
assert_eq!(entry.effective_batch, 4, "减半值应当写在 entry 上并持久");
assert!(matches!(embed_with_retry(&mut entry, &texts[..4]), EmbedOutcome::Vectors(_)));
assert_eq!(*lengths.lock(), vec![8, 4]);
}
#[test]
fn callback_errors_are_classified() {
let mut broken = EmbedderEntry {
options: EmbedderOptions::default(), effective_batch: 1,
embedder: Box::new(|_: &[String]| Err(EmbedCallbackError::too_large("still too large"))),
};
assert!(matches!(embed_with_retry(&mut broken, &["a".to_string()]), EmbedOutcome::Failed(_)));
assert_eq!(broken.effective_batch, 1);
let attempts = std::sync::Arc::new(AtomicUsize::new(0));
let counter = attempts.clone();
let mut throttled = EmbedderEntry {
options: EmbedderOptions::default(), effective_batch: 4,
embedder: Box::new(move |_: &[String]| {
let attempt = counter.fetch_add(1, AtomicOrdering::SeqCst);
if attempt < 2 { Err(EmbedCallbackError::rate_limited("slow down")) } else { Ok(vec![vec![1.0f32, 0.0]]) }
}),
};
assert!(matches!(embed_with_retry(&mut throttled, &["a".to_string()]), EmbedOutcome::Vectors(_)));
assert_eq!(attempts.load(AtomicOrdering::SeqCst), 3, "限流应当退避重试后成功");
assert_eq!(throttled.effective_batch, 4, "限流不触发减半");
}
}