a69d7caefb
- chunk(): byte splits that landed mid multibyte char silently dropped the whole chunk (accented/CJK lore). Back the boundary off to a char boundary; regression test included. - add_document(): DELETE the source first — re-adding a file no longer doubles its chunks. - chunks now record their embed_model (ALTER TABLE migration for old DBs); search skips chunks from a different model; rag_list reports it; new rag_reindex re-embeds everything, surfaced in the Lore panel as a mismatch banner with a one-click Reindex.
316 lines
12 KiB
Rust
316 lines
12 KiB
Rust
use crate::llm::LlmConfig;
|
|
use rusqlite::{params, Connection};
|
|
use serde::Deserialize;
|
|
use std::sync::Mutex;
|
|
|
|
// ponytail: brute-force cosine over normalized vectors, not sqlite-vec.
|
|
// Ceiling: ~10k chunks stays sub-millisecond on one core. Upgrade path:
|
|
// swap search() for a sqlite-vec virtual table when chunk count grows past
|
|
// that and scan time shows up in profiles. Avoids native extension loading.
|
|
|
|
/// One embedding vector, stored as little-endian f32 bytes.
|
|
/// Max chunk size in BYTES (chunk() splits on byte boundaries safely).
|
|
const MAX_CHUNK_BYTES: usize = 1000;
|
|
|
|
/// Legacy chunks indexed before the embed_model column existed. Unknown
|
|
/// model — assumed compatible so search keeps working; a one-click
|
|
/// Reindex stamps everything with the current model.
|
|
const EMBED_MODEL_UNKNOWN: &str = "";
|
|
|
|
pub struct RagStore {
|
|
db: Mutex<Connection>,
|
|
}
|
|
|
|
#[derive(Debug, Deserialize)]
|
|
struct EmbedResponse {
|
|
embeddings: Vec<Vec<f32>>,
|
|
}
|
|
|
|
impl RagStore {
|
|
pub fn open(dir: &std::path::Path) -> anyhow::Result<Self> {
|
|
std::fs::create_dir_all(dir)?;
|
|
let path = dir.join("lore.db");
|
|
let db = Connection::open(path)?;
|
|
db.execute_batch(
|
|
"CREATE TABLE IF NOT EXISTS chunks (
|
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
|
source TEXT NOT NULL,
|
|
text TEXT NOT NULL,
|
|
emb BLOB NOT NULL,
|
|
embed_model TEXT NOT NULL DEFAULT ''
|
|
);
|
|
CREATE INDEX IF NOT EXISTS idx_chunks_source ON chunks(source);",
|
|
)?;
|
|
// Migration for pre-embed_model DBs: add the column if it's missing.
|
|
// Duplicate-column error is expected on already-migrated DBs — ignore.
|
|
let _ = db.execute(
|
|
"ALTER TABLE chunks ADD COLUMN embed_model TEXT NOT NULL DEFAULT ''",
|
|
[],
|
|
);
|
|
Ok(Self { db: Mutex::new(db) })
|
|
}
|
|
|
|
/// Chunk `text` on paragraph boundaries, capping each chunk's length.
|
|
/// ponytail: no overlap in v1 — fine for retrieval at this scale; add
|
|
/// a sliding window if recall on boundary-spanning facts drops.
|
|
fn chunk(text: &str) -> Vec<String> {
|
|
let mut out = Vec::new();
|
|
for para in text.split("\n\n") {
|
|
let para = para.trim();
|
|
if para.is_empty() {
|
|
continue;
|
|
}
|
|
if para.len() <= MAX_CHUNK_BYTES {
|
|
out.push(para.to_string());
|
|
} else {
|
|
// Hard-cap long paragraphs on char boundaries. A naive byte
|
|
// split can land mid multibyte char — back the boundary off
|
|
// so accented/CJK text is never silently dropped.
|
|
let mut start = 0;
|
|
while start < para.len() {
|
|
let mut end = (start + MAX_CHUNK_BYTES).min(para.len());
|
|
while end < para.len() && !para.is_char_boundary(end) {
|
|
end -= 1;
|
|
}
|
|
out.push(para[start..end].trim().to_string());
|
|
start = end;
|
|
}
|
|
}
|
|
}
|
|
out
|
|
}
|
|
|
|
/// Embed a batch of texts via Ollama `/api/embed`. Vectors come back
|
|
/// L2-normalized from the server, so cosine similarity = dot product.
|
|
async fn embed(client: &reqwest::Client, config: &LlmConfig, texts: &[String]) -> anyhow::Result<Vec<Vec<f32>>> {
|
|
let url = format!("{}/api/embed", config.api_url.trim_end_matches('/'));
|
|
let body = serde_json::json!({ "model": config.embed_model, "input": texts });
|
|
let res = client.post(&url).json(&body).send().await?;
|
|
if !res.status().is_success() {
|
|
let status = res.status();
|
|
let text = res.text().await.unwrap_or_default();
|
|
anyhow::bail!("embed error {status}: {text}");
|
|
}
|
|
let parsed: EmbedResponse = res.json().await?;
|
|
Ok(parsed.embeddings)
|
|
}
|
|
|
|
fn vec_to_blob(v: &[f32]) -> Vec<u8> {
|
|
let mut bytes = Vec::with_capacity(v.len() * 4);
|
|
for f in v {
|
|
bytes.extend_from_slice(&f.to_le_bytes());
|
|
}
|
|
bytes
|
|
}
|
|
|
|
fn blob_to_vec(b: &[u8]) -> Vec<f32> {
|
|
b.chunks_exact(4)
|
|
.map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
|
|
.collect()
|
|
}
|
|
|
|
/// Add a lore document: chunk it, embed, and store. Returns chunk count.
|
|
pub async fn add_document(
|
|
&self,
|
|
config: &LlmConfig,
|
|
source: &str,
|
|
text: &str,
|
|
) -> anyhow::Result<usize> {
|
|
let chunks = Self::chunk(text);
|
|
if chunks.is_empty() {
|
|
return Ok(0);
|
|
}
|
|
let client = reqwest::Client::new();
|
|
let embeddings = Self::embed(&client, config, &chunks).await?;
|
|
if embeddings.len() != chunks.len() {
|
|
anyhow::bail!("embedding count mismatch");
|
|
}
|
|
let db = self.db.lock().map_err(|e| anyhow::anyhow!("db lock: {e}"))?;
|
|
// Re-adding a source replaces it (no duplicate chunks skewing scores).
|
|
db.execute("DELETE FROM chunks WHERE source = ?", params![source])?;
|
|
let mut stmt = db.prepare(
|
|
"INSERT INTO chunks (source, text, emb, embed_model) VALUES (?, ?, ?, ?)",
|
|
)?;
|
|
for (text, emb) in chunks.iter().zip(embeddings.iter()) {
|
|
stmt.execute(params![source, text, Self::vec_to_blob(emb), config.embed_model])?;
|
|
}
|
|
Ok(chunks.len())
|
|
}
|
|
|
|
/// Brute-force cosine search. Returns (text, source, score) for top_k.
|
|
pub async fn search(
|
|
&self,
|
|
config: &LlmConfig,
|
|
query: &str,
|
|
top_k: usize,
|
|
) -> anyhow::Result<Vec<(String, String, f32)>> {
|
|
let client = reqwest::Client::new();
|
|
let q_emb = Self::embed(&client, config, &[query.to_string()])
|
|
.await?
|
|
.into_iter()
|
|
.next()
|
|
.ok_or_else(|| anyhow::anyhow!("no query embedding"))?;
|
|
|
|
let rows: Vec<(String, String, Vec<u8>, String)> = {
|
|
let db = self.db.lock().map_err(|e| anyhow::anyhow!("db lock: {e}"))?;
|
|
let mut stmt =
|
|
db.prepare("SELECT text, source, emb, embed_model FROM chunks")?;
|
|
let rows = stmt.query_map([], |r| {
|
|
Ok((
|
|
r.get::<_, String>(0)?,
|
|
r.get::<_, String>(1)?,
|
|
r.get::<_, Vec<u8>>(2)?,
|
|
r.get::<_, String>(3)?,
|
|
))
|
|
})?;
|
|
rows.filter_map(|r| r.ok()).collect()
|
|
};
|
|
if rows.is_empty() {
|
|
return Ok(Vec::new());
|
|
}
|
|
|
|
// ponytail: naive O(n) scan. Fine to ~10k chunks; see module note.
|
|
// Skip chunks embedded with a DIFFERENT model — their dot products
|
|
// against this query vector are garbage. '' rows are legacy/unknown
|
|
// (indexed before per-chunk stamping) and stay in play; Reindex
|
|
// re-stamps them.
|
|
let mut scored: Vec<(String, String, f32)> = rows
|
|
.into_iter()
|
|
.filter(|(_, _, _, model)| {
|
|
model.as_str() == EMBED_MODEL_UNKNOWN || model == &config.embed_model
|
|
})
|
|
.map(|(text, source, blob, _)| {
|
|
let v = Self::blob_to_vec(&blob);
|
|
let dot: f32 = v.iter().zip(q_emb.iter()).map(|(a, b)| a * b).sum();
|
|
(text, source, dot)
|
|
})
|
|
.collect();
|
|
scored.sort_by(|a, b| b.2.partial_cmp(&a.2).unwrap_or(std::cmp::Ordering::Equal));
|
|
scored.truncate(top_k);
|
|
Ok(scored)
|
|
}
|
|
|
|
/// (source, chunk_count, embed_model) for every distinct source.
|
|
pub fn list_sources(&self) -> anyhow::Result<Vec<(String, i64, String)>> {
|
|
let db = self.db.lock().map_err(|e| anyhow::anyhow!("db lock: {e}"))?;
|
|
let mut stmt = db.prepare(
|
|
"SELECT source, COUNT(*), MAX(embed_model) FROM chunks GROUP BY source ORDER BY source",
|
|
)?;
|
|
let rows = stmt.query_map([], |r| {
|
|
Ok((r.get::<_, String>(0)?, r.get::<_, i64>(1)?, r.get::<_, String>(2)?))
|
|
})?;
|
|
let mut out = Vec::new();
|
|
for r in rows {
|
|
out.push(r?);
|
|
}
|
|
Ok(out)
|
|
}
|
|
|
|
/// Re-embed every source with the CURRENT embed model: for each source,
|
|
/// join its stored chunk texts and run them through add_document (which
|
|
/// replaces the old rows and stamps the model). Also the fix-it button
|
|
/// when the DM changes embedding models. Returns (sources, chunks).
|
|
pub async fn reindex(&self, config: &LlmConfig) -> anyhow::Result<(usize, usize)> {
|
|
let sources: Vec<String> = {
|
|
let db = self.db.lock().map_err(|e| anyhow::anyhow!("db lock: {e}"))?;
|
|
let mut stmt = db.prepare("SELECT DISTINCT source FROM chunks ORDER BY source")?;
|
|
let rows = stmt.query_map([], |r| r.get::<_, String>(0))?;
|
|
rows.filter_map(|r| r.ok()).collect()
|
|
};
|
|
let mut chunks = 0;
|
|
for s in &sources {
|
|
let texts: Vec<String> = {
|
|
let db = self.db.lock().map_err(|e| anyhow::anyhow!("db lock: {e}"))?;
|
|
let mut stmt =
|
|
db.prepare("SELECT text FROM chunks WHERE source = ? ORDER BY id")?;
|
|
let rows = stmt.query_map(params![s], |r| r.get::<_, String>(0))?;
|
|
rows.filter_map(|r| r.ok()).collect()
|
|
};
|
|
chunks += self.add_document(config, &s, &texts.join("\n\n")).await?;
|
|
}
|
|
Ok((sources.len(), chunks))
|
|
}
|
|
|
|
/// Clear all chunks, or just one source if given.
|
|
pub fn clear(&self, source: Option<&str>) -> anyhow::Result<()> {
|
|
let db = self.db.lock().map_err(|e| anyhow::anyhow!("db lock: {e}"))?;
|
|
match source {
|
|
Some(s) => { db.execute("DELETE FROM chunks WHERE source = ?", params![s])?; }
|
|
None => { db.execute("DELETE FROM chunks", [])?; }
|
|
}
|
|
Ok(())
|
|
}
|
|
|
|
/// Preview chunks for a source: (id, first 200 chars of text). Used by the
|
|
/// Lore panel so a DM can see what actually got embedded without a full dump.
|
|
pub fn list_chunks(&self, source: &str) -> anyhow::Result<Vec<(i64, String)>> {
|
|
let db = self.db.lock().map_err(|e| anyhow::anyhow!("db lock: {e}"))?;
|
|
let mut stmt = db.prepare("SELECT id, substr(text, 1, 200) FROM chunks WHERE source = ? ORDER BY id")?;
|
|
let rows = stmt.query_map(params![source], |r| {
|
|
Ok((r.get::<_, i64>(0)?, r.get::<_, String>(1)?))
|
|
})?;
|
|
let mut out = Vec::new();
|
|
for r in rows {
|
|
out.push(r?);
|
|
}
|
|
Ok(out)
|
|
}
|
|
|
|
/// Read-only SQL pass-through for the advanced SQLite viewer.
|
|
pub fn raw_query(
|
|
&self,
|
|
sql: &str,
|
|
) -> anyhow::Result<(Vec<String>, Vec<Vec<serde_json::Value>>)> {
|
|
let db = self.db.lock().map_err(|e| anyhow::anyhow!("db lock: {e}"))?;
|
|
crate::sql_viewer::run_readonly_query(&db, sql)
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
|
|
#[test]
|
|
fn chunk_splits_paragraphs_and_caps_long_ones() {
|
|
let long = "a".repeat(MAX_CHUNK_BYTES * 2 + 50);
|
|
let text = format!("short para\n\n{long}\n\nanother");
|
|
let chunks = RagStore::chunk(&text);
|
|
// "short para", (2 or 3 long pieces), "another"
|
|
assert!(chunks.len() >= 4);
|
|
assert_eq!(chunks.first().unwrap(), "short para");
|
|
assert!(chunks.last().unwrap() == "another");
|
|
for c in &chunks {
|
|
assert!(c.len() <= MAX_CHUNK_BYTES);
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn chunk_never_drops_multibyte_text() {
|
|
// é is 2 bytes: a 1002-byte paragraph forces a split that used to land
|
|
// wherever the byte counter said — including mid-character, silently
|
|
// dropping the whole chunk. Every byte must survive now.
|
|
let para = format!("é{}", "a".repeat(MAX_CHUNK_BYTES));
|
|
let chunks = RagStore::chunk(¶);
|
|
let rejoined = chunks.concat();
|
|
assert_eq!(rejoined, para, "chunking must not drop multibyte text");
|
|
for c in &chunks {
|
|
assert!(c.len() <= MAX_CHUNK_BYTES);
|
|
}
|
|
// Force a mid-character boundary: 998 a's then 3-byte chars, so the
|
|
// 1000-byte cut lands inside the first 日.
|
|
let para2 = format!("{}{}", "a".repeat(MAX_CHUNK_BYTES - 2), "日日日");
|
|
let chunks2 = RagStore::chunk(¶2);
|
|
assert_eq!(chunks2.concat(), para2);
|
|
}
|
|
|
|
#[test]
|
|
fn blob_roundtrip() {
|
|
let v = vec![0.0, 1.5, -2.25, 3.33];
|
|
let blob = RagStore::vec_to_blob(&v);
|
|
assert_eq!(blob.len(), v.len() * 4);
|
|
let back = RagStore::blob_to_vec(&blob);
|
|
for (a, b) in v.iter().zip(back.iter()) {
|
|
assert!((a - b).abs() < 1e-6);
|
|
}
|
|
}
|
|
} |