mod message_fts;
mod message_index;
pub use message_fts::{MessageFtsHit, MessageFtsIndex};
pub use message_index::{MessageHit, MessageIndex};
use anyhow::Result;
use async_trait::async_trait;
#[async_trait]
pub trait SearchEmbedder: Send + Sync {
fn dims(&self) -> usize;
fn model_id(&self) -> &str;
fn is_local(&self) -> bool;
async fn embed(&self, text: &str) -> Result<Vec<f32>>;
}
use std::path::Path;
use anyhow::Context;
use rusqlite::Connection;
pub(crate) fn open_vec_connection(path: &Path) -> Result<Connection> {
use std::sync::Once;
static REGISTER: Once = Once::new();
REGISTER.call_once(|| {
unsafe {
rusqlite::ffi::sqlite3_auto_extension(Some(std::mem::transmute(
sqlite_vec::sqlite3_vec_init as *const (),
)));
}
});
let conn = if path == Path::new(":memory:") {
Connection::open_in_memory().context("opening in-memory search db")?
} else {
Connection::open(path).with_context(|| format!("opening search db {}", path.display()))?
};
Ok(conn)
}
pub(crate) fn encode_embedding(vec: &[f32]) -> Vec<u8> {
let mut bytes = Vec::with_capacity(vec.len() * 4);
for v in vec {
bytes.extend_from_slice(&v.to_le_bytes());
}
bytes
}
#[cfg(test)]
mod test_embedder {
use super::{Result, SearchEmbedder};
use async_trait::async_trait;
pub struct LocalHashingEmbedder {
dims: usize,
}
impl LocalHashingEmbedder {
pub fn new(dims: usize) -> Self {
Self { dims }
}
}
#[async_trait]
impl SearchEmbedder for LocalHashingEmbedder {
fn dims(&self) -> usize {
self.dims
}
fn model_id(&self) -> &str {
"local-hashing"
}
fn is_local(&self) -> bool {
true
}
async fn embed(&self, text: &str) -> Result<Vec<f32>> {
Ok(local_embed(text, self.dims))
}
}
fn local_embed(text: &str, dims: usize) -> Vec<f32> {
let mut vec = vec![0.0f32; dims];
for token in tokenize(text) {
let bucket = (fnv1a(&token) as usize) % dims;
vec[bucket] += 1.0;
}
l2_normalize(&mut vec);
vec
}
fn tokenize(text: &str) -> Vec<String> {
text.split(|c: char| !c.is_alphanumeric())
.filter(|s| !s.is_empty())
.map(|s| s.to_lowercase())
.collect()
}
fn fnv1a(s: &str) -> u64 {
const OFFSET: u64 = 0xcbf2_9ce4_8422_2325;
const PRIME: u64 = 0x0000_0100_0000_01b3;
let mut hash = OFFSET;
for byte in s.as_bytes() {
hash ^= u64::from(*byte);
hash = hash.wrapping_mul(PRIME);
}
hash
}
fn l2_normalize(vec: &mut [f32]) {
let norm: f32 = vec.iter().map(|v| v * v).sum::<f32>().sqrt();
if norm > f32::EPSILON {
for v in vec.iter_mut() {
*v /= norm;
}
}
}
}