working through the plan + UI/ UX

This commit is contained in:
itsamejms
2026-07-12 22:13:43 +01:00
parent 5b256242be
commit a24f3615e0
38 changed files with 5295 additions and 328 deletions
+212
View File
@@ -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);
}
}
}