use std::sync::Arc;
use rusqlite::Connection;
use serde::{de::DeserializeOwned, Serialize};
use thiserror::Error;
use tokio::sync::Mutex;
mod bm25;
mod hnsw;
mod shadow;
#[cfg(test)]
mod tests;
pub use shadow::{ShadowMetrics, ShadowValidationConfig, ShadowValidationResult};
#[derive(Error, Debug)]
pub enum PersistError {
#[error("SQLite error: {0}")]
Sqlite(#[from] rusqlite::Error),
#[error("Serialization error: {0}")]
Serialize(String),
#[error("Deserialization error: {0}")]
Deserialize(String),
#[error("Task join error: {0}")]
TaskJoin(String),
#[error("Snapshot verification failed: {0}")]
SnapshotVerification(String),
#[error("Validation error: {0}")]
Validation(String),
#[error("Blocking task failed: {0}")]
BlockingJoin(String),
#[error("Tokio join error: {0}")]
Join(#[from] tokio::task::JoinError),
#[error("Internal error: {0}")]
Internal(String),
#[error("Embedding error: {0}")]
Embedding(String),
#[error("Retrieval error: {0}")]
Retrieval(String),
}
pub struct RetrievalPersistence {
pub(crate) conn: Arc<Mutex<Connection>>,
pub(crate) namespace: Arc<str>,
}
impl RetrievalPersistence {
pub fn new(conn: Arc<Mutex<Connection>>, namespace: impl Into<String>) -> Self {
Self {
conn,
namespace: Arc::from(namespace.into()),
}
}
pub async fn init_schema(&self) -> Result<(), PersistError> {
let conn = self.conn.clone();
tokio::task::spawn_blocking(move || {
let conn = conn.blocking_lock();
conn.execute_batch(
r#"
CREATE TABLE IF NOT EXISTS retrieval_snapshots (
namespace TEXT NOT NULL,
index_type TEXT NOT NULL,
snapshot BLOB NOT NULL,
created_at INTEGER NOT NULL,
PRIMARY KEY (namespace, index_type)
);
CREATE INDEX IF NOT EXISTS idx_retrieval_snapshots_namespace
ON retrieval_snapshots(namespace);
"#,
)?;
Ok(())
})
.await
.map_err(|e| PersistError::TaskJoin(e.to_string()))?
}
pub(crate) async fn persist_snapshot<T: Serialize + Send + Sync>(
&self,
index_type: &str,
snapshot: &T,
) -> Result<(), PersistError> {
let data =
serde_json::to_vec(snapshot).map_err(|e| PersistError::Serialize(e.to_string()))?;
let conn = self.conn.clone();
let namespace = self.namespace.clone();
let index_type = index_type.to_string();
tokio::task::spawn_blocking(move || {
let conn = conn.blocking_lock();
conn.execute(
r#"
INSERT OR REPLACE INTO retrieval_snapshots
(namespace, index_type, snapshot, created_at)
VALUES
(?1, ?2, ?3, ?4)
"#,
rusqlite::params![
&*namespace,
index_type,
data,
chrono::Utc::now().timestamp_micros()
],
)?;
Ok(())
})
.await
.map_err(|e| PersistError::TaskJoin(e.to_string()))?
}
pub(crate) async fn load_snapshot<T: DeserializeOwned + Send + 'static>(
&self,
index_type: &str,
) -> Result<Option<T>, PersistError> {
let conn = self.conn.clone();
let namespace = self.namespace.clone();
let index_type = index_type.to_string();
tokio::task::spawn_blocking(move || {
let conn = conn.blocking_lock();
let mut stmt = conn.prepare(
r#"
SELECT snapshot FROM retrieval_snapshots
WHERE namespace = ?1 AND index_type = ?2
"#,
)?;
let result: Option<Vec<u8>> = match stmt
.query_row(rusqlite::params![&*namespace, index_type], |row| row.get(0))
{
Ok(data) => Some(data),
Err(rusqlite::Error::QueryReturnedNoRows) => None,
Err(e) => return Err(PersistError::Sqlite(e)),
};
match result {
Some(data) => {
let snapshot: T = serde_json::from_slice(&data)
.map_err(|e| PersistError::Deserialize(e.to_string()))?;
Ok(Some(snapshot))
}
None => Ok(None),
}
})
.await
.map_err(|e| PersistError::TaskJoin(e.to_string()))?
}
pub async fn clear(&self) -> Result<(), PersistError> {
let conn = self.conn.clone();
let namespace = self.namespace.clone();
tokio::task::spawn_blocking(move || {
let conn = conn.blocking_lock();
conn.execute(
"DELETE FROM retrieval_snapshots WHERE namespace = ?1",
rusqlite::params![&*namespace],
)?;
Ok(())
})
.await
.map_err(|e| PersistError::TaskJoin(e.to_string()))?
}
pub async fn stats(&self) -> Result<PersistenceStats, PersistError> {
let conn = self.conn.clone();
let namespace = self.namespace.clone();
tokio::task::spawn_blocking(move || {
let conn = conn.blocking_lock();
let mut stmt = conn.prepare(
r#"
SELECT index_type, length(snapshot), created_at
FROM retrieval_snapshots
WHERE namespace = ?1
"#,
)?;
let mut stats = PersistenceStats::default();
let mut rows = stmt.query(rusqlite::params![&*namespace])?;
while let Some(row) = rows.next()? {
let index_type: String = row.get(0)?;
let size: i64 = row.get(1)?;
let created_at: i64 = row.get(2)?;
match index_type.as_str() {
"hnsw" => {
stats.hnsw_snapshot_size = size as usize;
stats.hnsw_snapshot_at = Some(created_at);
}
"bm25" => {
stats.bm25_snapshot_size = size as usize;
stats.bm25_snapshot_at = Some(created_at);
}
_ => {}
}
}
Ok(stats)
})
.await
.map_err(|e| PersistError::TaskJoin(e.to_string()))?
}
}
#[derive(Debug, Default, Clone)]
pub struct PersistenceStats {
pub hnsw_snapshot_size: usize,
pub hnsw_snapshot_at: Option<i64>,
pub bm25_snapshot_size: usize,
pub bm25_snapshot_at: Option<i64>,
}