use std::sync::Arc;
use rusqlite::Connection;
use serde::{de::DeserializeOwned, Serialize};
use thiserror::Error;
use tokio::sync::Mutex;
#[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(include_str!("../../sql/retrieval_snapshots_create.sql"))?;
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(
include_str!("../../sql/retrieval_snapshots_upsert.sql"),
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(include_str!(
"../../sql/retrieval_snapshots_select_snapshot.sql"
))?;
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(
include_str!("../../sql/retrieval_snapshots_delete_namespace.sql"),
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(include_str!(
"../../sql/retrieval_snapshots_list_namespace.sql"
))?;
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>,
}