working through the plan + UI/ UX
This commit is contained in:
@@ -0,0 +1,212 @@
|
||||
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.
|
||||
const MAX_CHUNK_CHARS: usize = 1000;
|
||||
|
||||
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
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_chunks_source ON chunks(source);",
|
||||
)?;
|
||||
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_CHARS {
|
||||
out.push(para.to_string());
|
||||
} else {
|
||||
// Hard-cap long paragraphs on char boundaries.
|
||||
for chunk in para.as_bytes().chunks(MAX_CHUNK_CHARS) {
|
||||
if let Ok(s) = std::str::from_utf8(chunk) {
|
||||
out.push(s.trim().to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
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}"))?;
|
||||
let mut stmt = db.prepare("INSERT INTO chunks (source, text, emb) VALUES (?, ?, ?)")?;
|
||||
for (text, emb) in chunks.iter().zip(embeddings.iter()) {
|
||||
stmt.execute(params![source, text, Self::vec_to_blob(emb)])?;
|
||||
}
|
||||
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>)> = {
|
||||
let db = self.db.lock().map_err(|e| anyhow::anyhow!("db lock: {e}"))?;
|
||||
let mut stmt = db.prepare("SELECT text, source, emb FROM chunks")?;
|
||||
let rows = stmt.query_map([], |r| {
|
||||
Ok((
|
||||
r.get::<_, String>(0)?,
|
||||
r.get::<_, String>(1)?,
|
||||
r.get::<_, Vec<u8>>(2)?,
|
||||
))
|
||||
})?;
|
||||
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.
|
||||
let mut scored: Vec<(String, String, f32)> = rows
|
||||
.into_iter()
|
||||
.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) for every distinct source.
|
||||
pub fn list_sources(&self) -> anyhow::Result<Vec<(String, i64)>> {
|
||||
let db = self.db.lock().map_err(|e| anyhow::anyhow!("db lock: {e}"))?;
|
||||
let mut stmt = db.prepare("SELECT source, COUNT(*) FROM chunks GROUP BY source ORDER BY source")?;
|
||||
let rows = stmt.query_map([], |r| Ok((r.get::<_, String>(0)?, r.get::<_, i64>(1)?)))?;
|
||||
let mut out = Vec::new();
|
||||
for r in rows {
|
||||
out.push(r?);
|
||||
}
|
||||
Ok(out)
|
||||
}
|
||||
|
||||
/// 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(())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn chunk_splits_paragraphs_and_caps_long_ones() {
|
||||
let long = "a".repeat(MAX_CHUNK_CHARS * 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_CHARS);
|
||||
}
|
||||
}
|
||||
|
||||
#[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);
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user