1use std::path::Path;
31
32use heed::types::{Bytes, Str};
33use heed::{Database, DatabaseFlags, Env, EnvFlags, EnvOpenOptions};
34use rkyv::rancor::Error as RkyvError;
35use sha2::{Digest, Sha256};
36
37use pylon_value::{ArchivedCachedEntry, CachedEntry, DecodedValue};
38
39pub type Result<T> = std::result::Result<T, Box<dyn std::error::Error + Send + Sync>>;
40
41const CACHE_FORMAT_VERSION: &[u8] = b"2";
59
60pub fn cache_key(sql: &str, params: &[DecodedValue]) -> Result<String> {
65 let mut hasher = Sha256::new();
66 hasher.update(sql.as_bytes());
67 for param in params {
68 let bytes = rkyv::to_bytes::<RkyvError>(param)?;
69 hasher.update(&bytes);
70 }
71 Ok(hex::encode(hasher.finalize()))
72}
73
74#[derive(Debug, Clone, Copy, PartialEq, Eq)]
76pub struct CacheStats {
77 pub entry_count: u64,
79 pub used_bytes: u64,
83}
84
85pub struct Cache {
86 env: Env,
87 entries: Database<Str, Bytes>,
88 tags: Database<Str, Str>,
89}
90
91impl Cache {
92 pub fn open(path: &Path, max_size_mb: usize) -> Result<Self> {
112 Self::open_with_flags(path, max_size_mb, EnvFlags::NO_SYNC)
113 }
114
115 fn open_with_flags(path: &Path, max_size_mb: usize, flags: EnvFlags) -> Result<Self> {
116 std::fs::create_dir_all(path)?;
117 let env = unsafe {
121 EnvOpenOptions::new()
122 .map_size(max_size_mb * 1024 * 1024)
123 .max_dbs(3)
124 .flags(flags)
125 .open(path)?
126 };
127
128 let mut wtxn = env.write_txn()?;
129 let entries = env.create_database(&mut wtxn, Some("entries"))?;
130 let tags = env
131 .database_options()
132 .types::<Str, Str>()
133 .flags(DatabaseFlags::DUP_SORT)
134 .name("tags")
135 .create(&mut wtxn)?;
136 let meta: Database<Str, Bytes> = env.create_database(&mut wtxn, Some("meta"))?;
137
138 if meta.get(&wtxn, "format_version")?.map(<[u8]>::to_vec) != Some(CACHE_FORMAT_VERSION.to_vec()) {
143 entries.clear(&mut wtxn)?;
144 tags.clear(&mut wtxn)?;
145 meta.put(&mut wtxn, "format_version", CACHE_FORMAT_VERSION)?;
146 }
147 wtxn.commit()?;
148
149 Ok(Self { env, entries, tags })
150 }
151
152 pub fn get(&self, key: &str) -> Result<Option<CachedEntry>> {
153 let rtxn = self.env.read_txn()?;
154 let Some(bytes) = self.entries.get(&rtxn, key)? else {
155 return Ok(None);
156 };
157 let mut aligned = rkyv::util::AlignedVec::<16>::new();
161 aligned.extend_from_slice(bytes);
162 let archived = unsafe { rkyv::access_unchecked::<ArchivedCachedEntry>(&aligned) };
165 let entry: CachedEntry = rkyv::deserialize::<CachedEntry, RkyvError>(archived)?;
166 Ok(Some(entry))
167 }
168
169 pub fn put(&self, key: &str, rows: Vec<DecodedValue>, tags: Vec<String>) -> Result<()> {
170 let entry = CachedEntry {
171 rows,
172 tags: tags.clone(),
173 };
174 let bytes = rkyv::to_bytes::<RkyvError>(&entry)?;
175
176 let mut wtxn = self.env.write_txn()?;
177 self.entries.put(&mut wtxn, key, &bytes)?;
178 for tag in &tags {
179 self.tags.put(&mut wtxn, tag, key)?;
180 }
181 wtxn.commit()?;
182 Ok(())
183 }
184
185 pub fn flush(&self) -> Result<()> {
189 self.env.force_sync()?;
190 Ok(())
191 }
192
193 pub fn invalidate(&self, tags: &[String]) -> Result<()> {
196 let mut wtxn = self.env.write_txn()?;
197
198 let mut keys_to_remove: Vec<String> = Vec::new();
199 for tag in tags {
200 if let Some(iter) = self.tags.get_duplicates(&wtxn, tag.as_str())? {
201 for result in iter {
202 let (_, cache_key) = result?;
203 keys_to_remove.push(cache_key.to_owned());
204 }
205 }
206 }
207
208 for tag in tags {
209 self.tags.delete(&mut wtxn, tag.as_str())?;
213 }
214 for key in &keys_to_remove {
215 self.entries.delete(&mut wtxn, key.as_str())?;
216 }
217
218 wtxn.commit()?;
219 Ok(())
220 }
221
222 pub fn stat(&self) -> Result<CacheStats> {
224 let rtxn = self.env.read_txn()?;
225 let entry_count = self.entries.len(&rtxn)?;
226 drop(rtxn);
227 let used_bytes = self.env.non_free_pages_size()?;
228 Ok(CacheStats {
229 entry_count,
230 used_bytes,
231 })
232 }
233
234 pub fn clear(&self) -> Result<()> {
238 let mut wtxn = self.env.write_txn()?;
239 self.entries.clear(&mut wtxn)?;
240 self.tags.clear(&mut wtxn)?;
241 wtxn.commit()?;
242 Ok(())
243 }
244}
245
246impl Drop for Cache {
247 fn drop(&mut self) {
255 let _ = self.env.force_sync();
256 }
257}
258
259#[cfg(test)]
260mod tests {
261 use super::*;
262
263 fn open_temp() -> (tempfile::TempDir, Cache) {
264 let dir = tempfile::tempdir().unwrap();
265 let cache = Cache::open(dir.path(), 10).unwrap();
266 (dir, cache)
267 }
268
269 #[test]
270 fn entries_survive_a_close_and_reopen() {
271 let dir = tempfile::tempdir().unwrap();
276 {
277 let cache = Cache::open(dir.path(), 10).unwrap();
278 for i in 0..32 {
279 cache
280 .put(
281 &format!("key{i}"),
282 vec![DecodedValue::I64(i)],
283 vec!["public.person".into()],
284 )
285 .unwrap();
286 }
287 }
288 let cache = Cache::open(dir.path(), 10).unwrap();
289 for i in 0..32 {
290 let entry = cache
291 .get(&format!("key{i}"))
292 .unwrap()
293 .unwrap_or_else(|| panic!("key{i} did not survive the reopen"));
294 assert_eq!(entry.rows, vec![DecodedValue::I64(i)]);
295 }
296 }
297
298 #[test]
299 fn an_explicit_flush_persists_without_closing() {
300 let dir = tempfile::tempdir().unwrap();
301 let cache = Cache::open(dir.path(), 10).unwrap();
302 cache
303 .put("key1", vec![DecodedValue::I64(7)], vec!["public.person".into()])
304 .unwrap();
305 cache.flush().unwrap();
306 assert!(cache.get("key1").unwrap().is_some());
307 }
308
309 #[test]
310 fn reopening_with_a_different_format_version_wipes_stale_entries() {
311 let dir = tempfile::tempdir().unwrap();
316 {
317 let cache = Cache::open(dir.path(), 10).unwrap();
318 cache
319 .put("key1", vec![DecodedValue::I64(1)], vec!["public.person".into()])
320 .unwrap();
321 assert!(cache.get("key1").unwrap().is_some());
322 }
323 {
326 let env = unsafe {
327 EnvOpenOptions::new()
328 .map_size(10 * 1024 * 1024)
329 .max_dbs(3)
330 .open(dir.path())
331 .unwrap()
332 };
333 let mut wtxn = env.write_txn().unwrap();
334 let meta: Database<Str, Bytes> = env.create_database(&mut wtxn, Some("meta")).unwrap();
335 meta.put(&mut wtxn, "format_version", b"a-different-version").unwrap();
336 wtxn.commit().unwrap();
337 }
338 let cache = Cache::open(dir.path(), 10).unwrap();
339 assert!(cache.get("key1").unwrap().is_none());
340 }
341
342 #[test]
343 fn reopening_with_the_same_format_version_preserves_entries() {
344 let dir = tempfile::tempdir().unwrap();
345 {
346 let cache = Cache::open(dir.path(), 10).unwrap();
347 cache
348 .put("key1", vec![DecodedValue::I64(1)], vec!["public.person".into()])
349 .unwrap();
350 }
351 let cache = Cache::open(dir.path(), 10).unwrap();
352 assert!(cache.get("key1").unwrap().is_some());
353 }
354
355 #[test]
356 fn put_then_get_round_trips() {
357 let (_dir, cache) = open_temp();
358 let rows = vec![DecodedValue::I64(1), DecodedValue::Str("hi".into())];
359 cache.put("key1", rows.clone(), vec!["public.person".into()]).unwrap();
360
361 let entry = cache.get("key1").unwrap().expect("entry present");
362 assert_eq!(entry.rows, rows);
363 assert_eq!(entry.tags, vec!["public.person".to_string()]);
364 }
365
366 #[test]
367 fn get_missing_key_returns_none() {
368 let (_dir, cache) = open_temp();
369 assert!(cache.get("nope").unwrap().is_none());
370 }
371
372 #[test]
373 fn invalidate_evicts_all_entries_sharing_a_tag() {
374 let (_dir, cache) = open_temp();
375 cache
376 .put("key1", vec![DecodedValue::I64(1)], vec!["public.person".into()])
377 .unwrap();
378 cache
379 .put(
380 "key2",
381 vec![DecodedValue::I64(2)],
382 vec!["public.person".into(), "public.pet".into()],
383 )
384 .unwrap();
385 cache
386 .put("key3", vec![DecodedValue::I64(3)], vec!["public.pet".into()])
387 .unwrap();
388
389 cache.invalidate(&["public.person".to_string()]).unwrap();
390
391 assert!(cache.get("key1").unwrap().is_none());
392 assert!(cache.get("key2").unwrap().is_none());
393 assert!(cache.get("key3").unwrap().is_some());
394 }
395
396 #[test]
397 fn invalidate_unknown_tag_is_a_no_op() {
398 let (_dir, cache) = open_temp();
399 cache
400 .put("key1", vec![DecodedValue::I64(1)], vec!["public.person".into()])
401 .unwrap();
402 cache.invalidate(&["public.nonexistent".to_string()]).unwrap();
403 assert!(cache.get("key1").unwrap().is_some());
404 }
405
406 #[test]
407 fn cache_key_is_stable_and_sensitive_to_params() {
408 let k1 = cache_key("select 1", &[DecodedValue::I64(1)]).unwrap();
409 let k2 = cache_key("select 1", &[DecodedValue::I64(1)]).unwrap();
410 let k3 = cache_key("select 1", &[DecodedValue::I64(2)]).unwrap();
411 assert_eq!(k1, k2);
412 assert_ne!(k1, k3);
413 }
414
415 #[test]
416 fn stat_on_empty_cache_reports_zero_entries() {
417 let (_dir, cache) = open_temp();
418 let stats = cache.stat().unwrap();
419 assert_eq!(stats.entry_count, 0);
420 }
421
422 #[test]
423 fn stat_reports_entry_count_and_nonzero_used_bytes_after_put() {
424 let (_dir, cache) = open_temp();
425 cache
426 .put("key1", vec![DecodedValue::I64(1)], vec!["public.person".into()])
427 .unwrap();
428 cache
429 .put("key2", vec![DecodedValue::I64(2)], vec!["public.pet".into()])
430 .unwrap();
431
432 let stats = cache.stat().unwrap();
433 assert_eq!(stats.entry_count, 2);
434 assert!(stats.used_bytes > 0);
435 }
436
437 #[test]
438 fn clear_removes_all_entries_and_tags() {
439 let (_dir, cache) = open_temp();
440 cache
441 .put("key1", vec![DecodedValue::I64(1)], vec!["public.person".into()])
442 .unwrap();
443 cache
444 .put("key2", vec![DecodedValue::I64(2)], vec!["public.pet".into()])
445 .unwrap();
446
447 cache.clear().unwrap();
448
449 assert!(cache.get("key1").unwrap().is_none());
450 assert!(cache.get("key2").unwrap().is_none());
451 assert_eq!(cache.stat().unwrap().entry_count, 0);
452
453 cache.invalidate(&["public.person".to_string()]).unwrap();
456 }
457
458 #[test]
459 fn clear_on_empty_cache_is_a_no_op() {
460 let (_dir, cache) = open_temp();
461 cache.clear().unwrap();
462 assert_eq!(cache.stat().unwrap().entry_count, 0);
463 }
464
465 fn bench_row(i: usize) -> DecodedValue {
476 DecodedValue::Composite(vec![
477 DecodedValue::Str("m::Article".to_string()),
478 DecodedValue::Str(format!("Title {i}")),
479 DecodedValue::Str(format!("slug-{i}")),
480 DecodedValue::Str("body text ".repeat(20)),
481 DecodedValue::Str("summary".into()),
482 DecodedValue::I64(i as i64 * 10),
483 DecodedValue::F64(4.5),
484 DecodedValue::Bool(true),
485 DecodedValue::Timestamptz(1_000_000_000),
486 DecodedValue::Timestamptz(1_000_000_001),
487 DecodedValue::Str("en".into()),
488 DecodedValue::Str("seed".into()),
489 DecodedValue::Str("abc".into()),
490 DecodedValue::I64(500),
491 DecodedValue::I64(3),
492 DecodedValue::Uuid([7; 16]),
493 DecodedValue::Composite(vec![
494 DecodedValue::Str("m::Author".into()),
495 DecodedValue::Str(format!("Author {i}")),
496 DecodedValue::Str(format!("a{i}@example.com")),
497 ]),
498 DecodedValue::Array(vec![DecodedValue::Composite(vec![
499 DecodedValue::Str("m::Tag".into()),
500 DecodedValue::Str("red".into()),
501 ])]),
502 ])
503 }
504
505 fn bench_time(label: &str, iterations: u32, mut f: impl FnMut(u32)) -> f64 {
506 for i in 0..iterations.min(20) {
507 f(i);
508 }
509 let mut best = f64::INFINITY;
510 for round in 0..5 {
511 let start = std::time::Instant::now();
512 for i in 0..iterations {
513 f(round * iterations + i);
514 }
515 best = best.min(start.elapsed().as_secs_f64() / iterations as f64 * 1e6);
516 }
517 println!(" {label:<54} {best:9.2} µs");
518 best
519 }
520
521 #[test]
530 #[ignore = "benchmark, not a correctness test"]
531 fn write_cost_breakdown() {
532 let data: Vec<DecodedValue> = (0..49).map(bench_row).collect();
533 let tags = vec!["public.article".to_string()];
534
535 println!("\nSERIALIZATION ONLY");
536 let entry = CachedEntry {
537 rows: data.clone(),
538 tags: tags.clone(),
539 };
540 let encoded = rkyv::to_bytes::<RkyvError>(&entry).unwrap();
541 println!(" (entry encodes to {} bytes)", encoded.len());
542 bench_time("rkyv::to_bytes", 2_000, |_| {
543 std::hint::black_box(rkyv::to_bytes::<RkyvError>(&entry).unwrap());
544 });
545
546 println!("\nFULL PUT, DURABLE COMMIT");
547 let durable_dir = tempfile::tempdir().unwrap();
548 let durable = Cache::open_with_flags(durable_dir.path(), 64, EnvFlags::empty()).unwrap();
549 let t_durable = bench_time("Cache::put", 200, |i| {
550 durable.put(&format!("k{i}"), data.clone(), tags.clone()).unwrap();
551 });
552
553 println!("\nFULL PUT, DEFERRED SYNC (as shipped)");
554 let deferred_dir = tempfile::tempdir().unwrap();
555 let deferred = Cache::open(deferred_dir.path(), 64).unwrap();
556 let t_deferred = bench_time("Cache::put", 200, |i| {
557 deferred.put(&format!("k{i}"), data.clone(), tags.clone()).unwrap();
558 });
559
560 println!(
561 "\n Deferring the sync saves {:.0} µs per put ({:.1}x).\n",
562 t_durable - t_deferred,
563 t_durable / t_deferred
564 );
565 }
566}