use std::path::Path;
use heed::types::{Bytes, Str};
use heed::{Database, DatabaseFlags, Env, EnvFlags, EnvOpenOptions};
use rkyv::rancor::Error as RkyvError;
use sha2::{Digest, Sha256};
use pylon_value::{ArchivedCachedEntry, CachedEntry, DecodedValue};
pub type Result<T> = std::result::Result<T, Box<dyn std::error::Error + Send + Sync>>;
const CACHE_FORMAT_VERSION: &[u8] = b"2";
pub fn cache_key(sql: &str, params: &[DecodedValue]) -> Result<String> {
let mut hasher = Sha256::new();
hasher.update(sql.as_bytes());
for param in params {
let bytes = rkyv::to_bytes::<RkyvError>(param)?;
hasher.update(&bytes);
}
Ok(hex::encode(hasher.finalize()))
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct CacheStats {
pub entry_count: u64,
pub used_bytes: u64,
}
pub struct Cache {
env: Env,
entries: Database<Str, Bytes>,
tags: Database<Str, Str>,
}
impl Cache {
pub fn open(path: &Path, max_size_mb: usize) -> Result<Self> {
Self::open_with_flags(path, max_size_mb, EnvFlags::NO_SYNC)
}
fn open_with_flags(path: &Path, max_size_mb: usize, flags: EnvFlags) -> Result<Self> {
std::fs::create_dir_all(path)?;
let env = unsafe {
EnvOpenOptions::new()
.map_size(max_size_mb * 1024 * 1024)
.max_dbs(3)
.flags(flags)
.open(path)?
};
let mut wtxn = env.write_txn()?;
let entries = env.create_database(&mut wtxn, Some("entries"))?;
let tags = env
.database_options()
.types::<Str, Str>()
.flags(DatabaseFlags::DUP_SORT)
.name("tags")
.create(&mut wtxn)?;
let meta: Database<Str, Bytes> = env.create_database(&mut wtxn, Some("meta"))?;
if meta.get(&wtxn, "format_version")?.map(<[u8]>::to_vec) != Some(CACHE_FORMAT_VERSION.to_vec()) {
entries.clear(&mut wtxn)?;
tags.clear(&mut wtxn)?;
meta.put(&mut wtxn, "format_version", CACHE_FORMAT_VERSION)?;
}
wtxn.commit()?;
Ok(Self { env, entries, tags })
}
pub fn get(&self, key: &str) -> Result<Option<CachedEntry>> {
let rtxn = self.env.read_txn()?;
let Some(bytes) = self.entries.get(&rtxn, key)? else {
return Ok(None);
};
let mut aligned = rkyv::util::AlignedVec::<16>::new();
aligned.extend_from_slice(bytes);
let archived = unsafe { rkyv::access_unchecked::<ArchivedCachedEntry>(&aligned) };
let entry: CachedEntry = rkyv::deserialize::<CachedEntry, RkyvError>(archived)?;
Ok(Some(entry))
}
pub fn put(&self, key: &str, rows: Vec<DecodedValue>, tags: Vec<String>) -> Result<()> {
let entry = CachedEntry {
rows,
tags: tags.clone(),
};
let bytes = rkyv::to_bytes::<RkyvError>(&entry)?;
let mut wtxn = self.env.write_txn()?;
self.entries.put(&mut wtxn, key, &bytes)?;
for tag in &tags {
self.tags.put(&mut wtxn, tag, key)?;
}
wtxn.commit()?;
Ok(())
}
pub fn flush(&self) -> Result<()> {
self.env.force_sync()?;
Ok(())
}
pub fn invalidate(&self, tags: &[String]) -> Result<()> {
let mut wtxn = self.env.write_txn()?;
let mut keys_to_remove: Vec<String> = Vec::new();
for tag in tags {
if let Some(iter) = self.tags.get_duplicates(&wtxn, tag.as_str())? {
for result in iter {
let (_, cache_key) = result?;
keys_to_remove.push(cache_key.to_owned());
}
}
}
for tag in tags {
self.tags.delete(&mut wtxn, tag.as_str())?;
}
for key in &keys_to_remove {
self.entries.delete(&mut wtxn, key.as_str())?;
}
wtxn.commit()?;
Ok(())
}
pub fn stat(&self) -> Result<CacheStats> {
let rtxn = self.env.read_txn()?;
let entry_count = self.entries.len(&rtxn)?;
drop(rtxn);
let used_bytes = self.env.non_free_pages_size()?;
Ok(CacheStats {
entry_count,
used_bytes,
})
}
pub fn clear(&self) -> Result<()> {
let mut wtxn = self.env.write_txn()?;
self.entries.clear(&mut wtxn)?;
self.tags.clear(&mut wtxn)?;
wtxn.commit()?;
Ok(())
}
}
impl Drop for Cache {
fn drop(&mut self) {
let _ = self.env.force_sync();
}
}
#[cfg(test)]
mod tests {
use super::*;
fn open_temp() -> (tempfile::TempDir, Cache) {
let dir = tempfile::tempdir().unwrap();
let cache = Cache::open(dir.path(), 10).unwrap();
(dir, cache)
}
#[test]
fn entries_survive_a_close_and_reopen() {
let dir = tempfile::tempdir().unwrap();
{
let cache = Cache::open(dir.path(), 10).unwrap();
for i in 0..32 {
cache
.put(
&format!("key{i}"),
vec![DecodedValue::I64(i)],
vec!["public.person".into()],
)
.unwrap();
}
}
let cache = Cache::open(dir.path(), 10).unwrap();
for i in 0..32 {
let entry = cache
.get(&format!("key{i}"))
.unwrap()
.unwrap_or_else(|| panic!("key{i} did not survive the reopen"));
assert_eq!(entry.rows, vec![DecodedValue::I64(i)]);
}
}
#[test]
fn an_explicit_flush_persists_without_closing() {
let dir = tempfile::tempdir().unwrap();
let cache = Cache::open(dir.path(), 10).unwrap();
cache
.put("key1", vec![DecodedValue::I64(7)], vec!["public.person".into()])
.unwrap();
cache.flush().unwrap();
assert!(cache.get("key1").unwrap().is_some());
}
#[test]
fn reopening_with_a_different_format_version_wipes_stale_entries() {
let dir = tempfile::tempdir().unwrap();
{
let cache = Cache::open(dir.path(), 10).unwrap();
cache
.put("key1", vec![DecodedValue::I64(1)], vec!["public.person".into()])
.unwrap();
assert!(cache.get("key1").unwrap().is_some());
}
{
let env = unsafe {
EnvOpenOptions::new()
.map_size(10 * 1024 * 1024)
.max_dbs(3)
.open(dir.path())
.unwrap()
};
let mut wtxn = env.write_txn().unwrap();
let meta: Database<Str, Bytes> = env.create_database(&mut wtxn, Some("meta")).unwrap();
meta.put(&mut wtxn, "format_version", b"a-different-version").unwrap();
wtxn.commit().unwrap();
}
let cache = Cache::open(dir.path(), 10).unwrap();
assert!(cache.get("key1").unwrap().is_none());
}
#[test]
fn reopening_with_the_same_format_version_preserves_entries() {
let dir = tempfile::tempdir().unwrap();
{
let cache = Cache::open(dir.path(), 10).unwrap();
cache
.put("key1", vec![DecodedValue::I64(1)], vec!["public.person".into()])
.unwrap();
}
let cache = Cache::open(dir.path(), 10).unwrap();
assert!(cache.get("key1").unwrap().is_some());
}
#[test]
fn put_then_get_round_trips() {
let (_dir, cache) = open_temp();
let rows = vec![DecodedValue::I64(1), DecodedValue::Str("hi".into())];
cache.put("key1", rows.clone(), vec!["public.person".into()]).unwrap();
let entry = cache.get("key1").unwrap().expect("entry present");
assert_eq!(entry.rows, rows);
assert_eq!(entry.tags, vec!["public.person".to_string()]);
}
#[test]
fn get_missing_key_returns_none() {
let (_dir, cache) = open_temp();
assert!(cache.get("nope").unwrap().is_none());
}
#[test]
fn invalidate_evicts_all_entries_sharing_a_tag() {
let (_dir, cache) = open_temp();
cache
.put("key1", vec![DecodedValue::I64(1)], vec!["public.person".into()])
.unwrap();
cache
.put(
"key2",
vec![DecodedValue::I64(2)],
vec!["public.person".into(), "public.pet".into()],
)
.unwrap();
cache
.put("key3", vec![DecodedValue::I64(3)], vec!["public.pet".into()])
.unwrap();
cache.invalidate(&["public.person".to_string()]).unwrap();
assert!(cache.get("key1").unwrap().is_none());
assert!(cache.get("key2").unwrap().is_none());
assert!(cache.get("key3").unwrap().is_some());
}
#[test]
fn invalidate_unknown_tag_is_a_no_op() {
let (_dir, cache) = open_temp();
cache
.put("key1", vec![DecodedValue::I64(1)], vec!["public.person".into()])
.unwrap();
cache.invalidate(&["public.nonexistent".to_string()]).unwrap();
assert!(cache.get("key1").unwrap().is_some());
}
#[test]
fn cache_key_is_stable_and_sensitive_to_params() {
let k1 = cache_key("select 1", &[DecodedValue::I64(1)]).unwrap();
let k2 = cache_key("select 1", &[DecodedValue::I64(1)]).unwrap();
let k3 = cache_key("select 1", &[DecodedValue::I64(2)]).unwrap();
assert_eq!(k1, k2);
assert_ne!(k1, k3);
}
#[test]
fn stat_on_empty_cache_reports_zero_entries() {
let (_dir, cache) = open_temp();
let stats = cache.stat().unwrap();
assert_eq!(stats.entry_count, 0);
}
#[test]
fn stat_reports_entry_count_and_nonzero_used_bytes_after_put() {
let (_dir, cache) = open_temp();
cache
.put("key1", vec![DecodedValue::I64(1)], vec!["public.person".into()])
.unwrap();
cache
.put("key2", vec![DecodedValue::I64(2)], vec!["public.pet".into()])
.unwrap();
let stats = cache.stat().unwrap();
assert_eq!(stats.entry_count, 2);
assert!(stats.used_bytes > 0);
}
#[test]
fn clear_removes_all_entries_and_tags() {
let (_dir, cache) = open_temp();
cache
.put("key1", vec![DecodedValue::I64(1)], vec!["public.person".into()])
.unwrap();
cache
.put("key2", vec![DecodedValue::I64(2)], vec!["public.pet".into()])
.unwrap();
cache.clear().unwrap();
assert!(cache.get("key1").unwrap().is_none());
assert!(cache.get("key2").unwrap().is_none());
assert_eq!(cache.stat().unwrap().entry_count, 0);
cache.invalidate(&["public.person".to_string()]).unwrap();
}
#[test]
fn clear_on_empty_cache_is_a_no_op() {
let (_dir, cache) = open_temp();
cache.clear().unwrap();
assert_eq!(cache.stat().unwrap().entry_count, 0);
}
fn bench_row(i: usize) -> DecodedValue {
DecodedValue::Composite(vec![
DecodedValue::Str("m::Article".to_string()),
DecodedValue::Str(format!("Title {i}")),
DecodedValue::Str(format!("slug-{i}")),
DecodedValue::Str("body text ".repeat(20)),
DecodedValue::Str("summary".into()),
DecodedValue::I64(i as i64 * 10),
DecodedValue::F64(4.5),
DecodedValue::Bool(true),
DecodedValue::Timestamptz(1_000_000_000),
DecodedValue::Timestamptz(1_000_000_001),
DecodedValue::Str("en".into()),
DecodedValue::Str("seed".into()),
DecodedValue::Str("abc".into()),
DecodedValue::I64(500),
DecodedValue::I64(3),
DecodedValue::Uuid([7; 16]),
DecodedValue::Composite(vec![
DecodedValue::Str("m::Author".into()),
DecodedValue::Str(format!("Author {i}")),
DecodedValue::Str(format!("a{i}@example.com")),
]),
DecodedValue::Array(vec![DecodedValue::Composite(vec![
DecodedValue::Str("m::Tag".into()),
DecodedValue::Str("red".into()),
])]),
])
}
fn bench_time(label: &str, iterations: u32, mut f: impl FnMut(u32)) -> f64 {
for i in 0..iterations.min(20) {
f(i);
}
let mut best = f64::INFINITY;
for round in 0..5 {
let start = std::time::Instant::now();
for i in 0..iterations {
f(round * iterations + i);
}
best = best.min(start.elapsed().as_secs_f64() / iterations as f64 * 1e6);
}
println!(" {label:<54} {best:9.2} µs");
best
}
#[test]
#[ignore = "benchmark, not a correctness test"]
fn write_cost_breakdown() {
let data: Vec<DecodedValue> = (0..49).map(bench_row).collect();
let tags = vec!["public.article".to_string()];
println!("\nSERIALIZATION ONLY");
let entry = CachedEntry {
rows: data.clone(),
tags: tags.clone(),
};
let encoded = rkyv::to_bytes::<RkyvError>(&entry).unwrap();
println!(" (entry encodes to {} bytes)", encoded.len());
bench_time("rkyv::to_bytes", 2_000, |_| {
std::hint::black_box(rkyv::to_bytes::<RkyvError>(&entry).unwrap());
});
println!("\nFULL PUT, DURABLE COMMIT");
let durable_dir = tempfile::tempdir().unwrap();
let durable = Cache::open_with_flags(durable_dir.path(), 64, EnvFlags::empty()).unwrap();
let t_durable = bench_time("Cache::put", 200, |i| {
durable.put(&format!("k{i}"), data.clone(), tags.clone()).unwrap();
});
println!("\nFULL PUT, DEFERRED SYNC (as shipped)");
let deferred_dir = tempfile::tempdir().unwrap();
let deferred = Cache::open(deferred_dir.path(), 64).unwrap();
let t_deferred = bench_time("Cache::put", 200, |i| {
deferred.put(&format!("k{i}"), data.clone(), tags.clone()).unwrap();
});
println!(
"\n Deferring the sync saves {:.0} µs per put ({:.1}x).\n",
t_durable - t_deferred,
t_durable / t_deferred
);
}
}