use std::{
path::Path,
sync::{Arc, Mutex},
time::{Duration, SystemTime, UNIX_EPOCH},
};
use async_trait::async_trait;
use rusqlite::{Connection, OptionalExtension, params};
use super::{
backend::{CacheBackend, CacheError},
key::CacheKey,
};
const CREATE_TABLE: &str = "
CREATE TABLE IF NOT EXISTS cache_entries (
key_type TEXT NOT NULL,
key_id TEXT NOT NULL,
value TEXT NOT NULL,
expires_at INTEGER NOT NULL,
PRIMARY KEY (key_type, key_id)
);
CREATE INDEX IF NOT EXISTS idx_expires ON cache_entries(expires_at);
";
const CACHE_SCHEMA_VERSION: i64 = 1;
fn init_schema(conn: &Connection) -> Result<(), CacheError> {
conn.execute_batch(CREATE_TABLE)
.map_err(CacheError::backend)?;
let version: i64 = conn
.query_row("PRAGMA user_version", [], |row| row.get(0))
.map_err(CacheError::backend)?;
if version != CACHE_SCHEMA_VERSION {
conn.execute_batch("DELETE FROM cache_entries;")
.map_err(CacheError::backend)?;
conn.pragma_update(None, "user_version", CACHE_SCHEMA_VERSION)
.map_err(CacheError::backend)?;
}
Ok(())
}
type Clock = Arc<dyn Fn() -> u64 + Send + Sync>;
const MAX_ENTRIES: i64 = 50_000;
fn system_clock() -> Clock {
Arc::new(|| {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs()
})
}
#[derive(Clone)]
pub struct SqliteCache {
read_conn: Arc<Mutex<Connection>>,
write_conn: Arc<Mutex<Connection>>,
read_count: Arc<std::sync::atomic::AtomicU64>,
clock: Clock,
}
impl std::fmt::Debug for SqliteCache {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SqliteCache").finish_non_exhaustive()
}
}
impl SqliteCache {
pub fn new(path: impl AsRef<Path>) -> Result<Self, CacheError> {
let path = path.as_ref();
let write_conn = Connection::open(path).map_err(|e| CacheError::Backend(Box::new(e)))?;
write_conn
.execute_batch("PRAGMA journal_mode=WAL; PRAGMA synchronous=NORMAL;")
.map_err(|e| CacheError::Backend(Box::new(e)))?;
init_schema(&write_conn)?;
let read_conn = Connection::open(path).map_err(|e| CacheError::Backend(Box::new(e)))?;
Ok(Self {
read_conn: Arc::new(Mutex::new(read_conn)),
write_conn: Arc::new(Mutex::new(write_conn)),
read_count: Arc::new(std::sync::atomic::AtomicU64::new(0)),
clock: system_clock(),
})
}
pub fn in_memory() -> Result<Self, CacheError> {
let conn = Connection::open_in_memory().map_err(|e| CacheError::Backend(Box::new(e)))?;
init_schema(&conn)?;
let conn = Arc::new(Mutex::new(conn));
Ok(Self {
read_conn: Arc::clone(&conn),
write_conn: conn,
read_count: Arc::new(std::sync::atomic::AtomicU64::new(0)),
clock: system_clock(),
})
}
pub async fn purge_expired(&self) -> Result<(), CacheError> {
let conn = Arc::clone(&self.write_conn);
let now = self.now();
tokio::task::spawn_blocking(move || {
let conn = conn.lock().unwrap_or_else(|poison| poison.into_inner());
conn.execute(
"DELETE FROM cache_entries WHERE expires_at < ?1",
rusqlite::params![now as i64],
)
.map_err(|e| CacheError::Backend(Box::new(e)))?;
Ok(())
})
.await
.map_err(|e| CacheError::Backend(Box::from(format!("spawn_blocking failed: {e}"))))?
}
#[cfg(test)]
pub(crate) fn with_clock(mut self, clock: Clock) -> Self {
self.clock = clock;
self
}
fn now(&self) -> u64 {
(self.clock)()
}
fn enforce_limits(conn: &Connection, now: u64) {
let _ = conn.execute(
"DELETE FROM cache_entries WHERE expires_at < ?1",
params![now as i64],
);
let _ = conn.execute(
"DELETE FROM cache_entries WHERE rowid NOT IN \
(SELECT rowid FROM cache_entries ORDER BY expires_at DESC LIMIT ?1)",
params![MAX_ENTRIES],
);
}
fn maybe_purge_on_read(write_conn: Arc<Mutex<Connection>>, count: u64) {
if count != 0 && count.is_multiple_of(1000) {
let _purge_task = tokio::task::spawn_blocking(move || {
let conn = write_conn
.lock()
.unwrap_or_else(|poison| poison.into_inner());
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
let _ = conn.execute(
"DELETE FROM cache_entries WHERE expires_at < ?1",
params![now as i64],
);
});
}
}
}
#[async_trait]
impl CacheBackend for SqliteCache {
#[tracing::instrument(skip(self, key), fields(key_type = key.key_type(), key_id = %key.key_id()))]
async fn get(&self, key: &CacheKey) -> Result<Option<serde_json::Value>, CacheError> {
let conn = Arc::clone(&self.read_conn);
let write_conn = Arc::clone(&self.write_conn);
let read_count = Arc::clone(&self.read_count);
let key_type = key.key_type().to_string();
let key_id = key.key_id();
let now = self.now();
let result = tokio::task::spawn_blocking(move || {
let conn = conn.lock().unwrap_or_else(|poison| poison.into_inner());
let result = conn.query_row(
"SELECT value FROM cache_entries WHERE key_type = ?1 AND key_id = ?2 AND expires_at > ?3",
params![key_type, key_id, now as i64],
|row| row.get::<_, String>(0),
).optional();
match result.map_err(|e| CacheError::Backend(Box::new(e)))? {
Some(json_str) => {
let value = serde_json::from_str(&json_str)?;
Ok(Some(value))
}
None => Ok(None),
}
})
.await
.map_err(|e| CacheError::Backend(Box::from(format!("spawn_blocking failed: {e}"))))?;
let count = read_count.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
Self::maybe_purge_on_read(write_conn, count);
result
}
#[tracing::instrument(skip(self, key, value), fields(key_type = key.key_type(), key_id = %key.key_id(), ttl_secs = ttl.as_secs()))]
async fn set(
&self,
key: CacheKey,
value: serde_json::Value,
ttl: Duration,
) -> Result<(), CacheError> {
let conn = Arc::clone(&self.write_conn);
let key_type = key.key_type().to_string();
let key_id = key.key_id();
let json_str = serde_json::to_string(&value)?;
let now = self.now();
let expires_at = now + ttl.as_secs();
tokio::task::spawn_blocking(move || {
let conn = conn.lock().unwrap_or_else(|poison| poison.into_inner());
conn.execute(
"INSERT OR REPLACE INTO cache_entries (key_type, key_id, value, expires_at) VALUES (?1, ?2, ?3, ?4)",
params![key_type, key_id, json_str, expires_at as i64],
)
.map_err(|e| CacheError::Backend(Box::new(e)))?;
Self::enforce_limits(&conn, now);
Ok(())
})
.await
.map_err(|e| CacheError::Backend(Box::from(format!("spawn_blocking failed: {e}"))))?
}
#[tracing::instrument(skip(self, key), fields(key_type = key.key_type(), key_id = %key.key_id()))]
async fn invalidate(&self, key: &CacheKey) -> Result<(), CacheError> {
let conn = Arc::clone(&self.write_conn);
let key_type = key.key_type().to_string();
let key_id = key.key_id();
tokio::task::spawn_blocking(move || {
let conn = conn.lock().unwrap_or_else(|poison| poison.into_inner());
conn.execute(
"DELETE FROM cache_entries WHERE key_type = ?1 AND key_id = ?2",
params![key_type, key_id],
)
.map_err(|e| CacheError::Backend(Box::new(e)))?;
Ok(())
})
.await
.map_err(|e| CacheError::Backend(Box::from(format!("spawn_blocking failed: {e}"))))?
}
async fn clear(&self) -> Result<(), CacheError> {
let conn = Arc::clone(&self.write_conn);
tokio::task::spawn_blocking(move || {
let conn = conn.lock().unwrap_or_else(|poison| poison.into_inner());
conn.execute("DELETE FROM cache_entries", [])
.map_err(|e| CacheError::Backend(Box::new(e)))?;
Ok(())
})
.await
.map_err(|e| CacheError::Backend(Box::from(format!("spawn_blocking failed: {e}"))))?
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use super::*;
use crate::cache::{
CacheBackend,
key::{CacheKey, MediaType},
};
#[tokio::test]
async fn schema_version_bump_invalidates_old_entries() {
let path = std::env::temp_dir().join(format!("cameo_schema_ver_{}.db", std::process::id()));
let _ = std::fs::remove_file(&path);
let key = CacheKey::Detail {
media_type: MediaType::Movie,
provider_id: "tmdb:1".to_string(),
};
{
let cache = SqliteCache::new(&path).unwrap();
cache
.set(
key.clone(),
serde_json::json!({ "v": 1 }),
Duration::from_secs(3600),
)
.await
.unwrap();
assert!(cache.get(&key).await.unwrap().is_some());
}
{
let conn = Connection::open(&path).unwrap();
conn.pragma_update(None, "user_version", 999i64).unwrap();
}
let cache = SqliteCache::new(&path).unwrap();
assert!(cache.get(&key).await.unwrap().is_none());
let _ = std::fs::remove_file(&path);
}
fn mock(now: &Arc<std::sync::atomic::AtomicU64>) -> Clock {
let now = Arc::clone(now);
Arc::new(move || now.load(std::sync::atomic::Ordering::Relaxed))
}
#[tokio::test]
async fn ttl_expiry_is_deterministic_with_mock_clock() {
let now = Arc::new(std::sync::atomic::AtomicU64::new(1_000));
let cache = SqliteCache::in_memory().unwrap().with_clock(mock(&now));
let key = CacheKey::Detail {
media_type: MediaType::Movie,
provider_id: "tmdb:1".to_string(),
};
cache
.set(
key.clone(),
serde_json::json!({ "v": 1 }),
Duration::from_secs(60),
)
.await
.unwrap();
now.store(1_030, std::sync::atomic::Ordering::Relaxed); assert!(cache.get(&key).await.unwrap().is_some());
now.store(1_070, std::sync::atomic::Ordering::Relaxed); assert!(cache.get(&key).await.unwrap().is_none());
}
#[tokio::test]
async fn expired_rows_are_reclaimed_on_write() {
let now = Arc::new(std::sync::atomic::AtomicU64::new(1_000));
let cache = SqliteCache::in_memory().unwrap().with_clock(mock(&now));
let stale = CacheKey::Detail {
media_type: MediaType::Movie,
provider_id: "tmdb:1".to_string(),
};
cache
.set(
stale.clone(),
serde_json::json!({}),
Duration::from_secs(10),
)
.await
.unwrap();
now.store(2_000, std::sync::atomic::Ordering::Relaxed); let fresh = CacheKey::Detail {
media_type: MediaType::Movie,
provider_id: "tmdb:2".to_string(),
};
cache
.set(fresh, serde_json::json!({}), Duration::from_secs(3_600))
.await
.unwrap();
now.store(1_000, std::sync::atomic::Ordering::Relaxed);
assert!(cache.get(&stale).await.unwrap().is_none());
}
}