use pylon_core::query::CompiledQuery;
use pylon_value::DecodedValue;
use crate::error::{Error, Result};
fn map_err<E: std::fmt::Display>(e: E) -> Error {
Error::Cache(e.to_string())
}
#[derive(Clone, Copy, Default)]
pub(crate) struct CacheAccess<'a> {
cache: Option<&'a pylon_cache::Cache>,
read_through: bool,
}
impl<'a> CacheAccess<'a> {
pub(crate) fn read_write(cache: Option<&'a pylon_cache::Cache>) -> Self {
Self {
cache,
read_through: true,
}
}
pub(crate) fn evict_only(cache: Option<&'a pylon_cache::Cache>) -> Self {
Self {
cache,
read_through: false,
}
}
fn readable(&self) -> Option<&'a pylon_cache::Cache> {
if self.read_through { self.cache } else { None }
}
}
fn cache_key(kind: &str, sql: &str, params: &[DecodedValue]) -> Result<String> {
pylon_cache::cache_key(&format!("{kind}\0{sql}"), params).map_err(map_err)
}
fn is_cacheable(compiled: &CompiledQuery) -> bool {
!compiled.mutates && !compiled.tags.is_empty()
}
pub(crate) fn invalidate_for(access: CacheAccess<'_>, compiled: &CompiledQuery) -> Result<()> {
let Some(cache) = access.cache else { return Ok(()) };
if !compiled.mutates || compiled.tags.is_empty() {
return Ok(());
}
cache.invalidate(&compiled.tags).map_err(map_err)
}
pub(crate) fn get_rows(
access: CacheAccess<'_>,
compiled: &CompiledQuery,
params: &[DecodedValue],
) -> Result<Option<Vec<DecodedValue>>> {
let Some(cache) = access.readable() else {
return Ok(None);
};
if !is_cacheable(compiled) {
return Ok(None);
}
let key = cache_key("rows", &compiled.sql, params)?;
Ok(cache.get(&key).map_err(map_err)?.map(|entry| entry.rows))
}
pub(crate) fn put_rows(
access: CacheAccess<'_>,
compiled: &CompiledQuery,
params: &[DecodedValue],
rows: &[DecodedValue],
) -> Result<()> {
let Some(cache) = access.readable() else { return Ok(()) };
if !is_cacheable(compiled) {
return Ok(());
}
let key = cache_key("rows", &compiled.sql, params)?;
cache.put(&key, rows.to_vec(), compiled.tags.clone()).map_err(map_err)
}
pub(crate) fn get_json(
access: CacheAccess<'_>,
kind: &str,
compiled: &CompiledQuery,
params: &[DecodedValue],
) -> Result<Option<Option<String>>> {
let Some(cache) = access.readable() else {
return Ok(None);
};
if !is_cacheable(compiled) {
return Ok(None);
}
let key = cache_key(kind, &compiled.sql, params)?;
let Some(entry) = cache.get(&key).map_err(map_err)? else {
return Ok(None);
};
Ok(Some(match entry.rows.into_iter().next() {
Some(DecodedValue::Str(s)) => Some(s),
_ => None,
}))
}
pub(crate) fn put_json(
access: CacheAccess<'_>,
kind: &str,
compiled: &CompiledQuery,
params: &[DecodedValue],
value: Option<&str>,
) -> Result<()> {
let Some(cache) = access.readable() else { return Ok(()) };
if !is_cacheable(compiled) {
return Ok(());
}
let key = cache_key(kind, &compiled.sql, params)?;
let rows = value
.map(|v| vec![DecodedValue::Str(v.to_string())])
.unwrap_or_default();
cache.put(&key, rows, compiled.tags.clone()).map_err(map_err)
}
#[cfg(test)]
mod tests {
use super::*;
fn rw(cache: &pylon_cache::Cache) -> CacheAccess<'_> {
CacheAccess::read_write(Some(cache))
}
fn open_temp() -> (tempfile::TempDir, pylon_cache::Cache) {
let dir = tempfile::tempdir().unwrap();
let cache = pylon_cache::Cache::open(dir.path(), 10).unwrap();
(dir, cache)
}
fn compiled_with_tags(sql: &str, tags: &[&str]) -> CompiledQuery {
compiled(sql, tags, false)
}
fn compiled(sql: &str, tags: &[&str], mutates: bool) -> CompiledQuery {
let shape = pylon_core::query::ShapeDescriptor {
root: pylon_core::query::ShapeNode::RawScalar,
};
let shape_id = pylon_core::query::derive_shape_id(sql, &shape);
CompiledQuery {
sql: sql.to_string(),
param_names: vec![],
params: vec![],
shape,
warnings: vec![],
inference_plan: None,
tags: tags.iter().map(|t| t.to_string()).collect(),
mutates,
analyze_paths: None,
shape_id,
}
}
#[test]
fn a_mutating_statement_is_never_cached() {
let (_dir, cache) = open_temp();
let insert = compiled("insert person", &["public.person"], true);
put_rows(rw(&cache), &insert, &[], &[DecodedValue::Str("row".into())]).unwrap();
assert_eq!(
get_rows(rw(&cache), &insert, &[]).unwrap(),
None,
"a write must not be served from cache"
);
put_json(rw(&cache), "json_all", &insert, &[], Some("[]")).unwrap();
assert_eq!(get_json(rw(&cache), "json_all", &insert, &[]).unwrap(), None);
}
#[test]
fn the_same_tags_are_still_cacheable_for_a_read() {
let (_dir, cache) = open_temp();
let read = compiled("select person", &["public.person"], false);
put_rows(rw(&cache), &read, &[], &[DecodedValue::Str("row".into())]).unwrap();
assert!(get_rows(rw(&cache), &read, &[]).unwrap().is_some());
}
#[test]
fn a_write_evicts_a_read_sharing_its_tag() {
let (_dir, cache) = open_temp();
let read = compiled("select person", &["public.person"], false);
put_rows(rw(&cache), &read, &[], &[DecodedValue::Str("before".into())]).unwrap();
assert!(get_rows(rw(&cache), &read, &[]).unwrap().is_some());
let write = compiled("update person", &["public.person"], true);
invalidate_for(rw(&cache), &write).unwrap();
assert_eq!(
get_rows(rw(&cache), &read, &[]).unwrap(),
None,
"the read must be evicted by a write to the same set"
);
}
#[test]
fn a_write_leaves_an_unrelated_tag_alone() {
let (_dir, cache) = open_temp();
let other = compiled("select company", &["public.company"], false);
put_rows(rw(&cache), &other, &[], &[DecodedValue::Str("row".into())]).unwrap();
invalidate_for(rw(&cache), &compiled("update person", &["public.person"], true)).unwrap();
assert!(get_rows(rw(&cache), &other, &[]).unwrap().is_some());
}
#[test]
fn transaction_access_evicts_without_populating() {
let (_dir, cache) = open_temp();
let tx = CacheAccess::evict_only(Some(&cache));
let read = compiled("select person", &["public.person"], false);
put_rows(tx, &read, &[], &[DecodedValue::Str("uncommitted".into())]).unwrap();
assert_eq!(get_rows(rw(&cache), &read, &[]).unwrap(), None);
put_rows(rw(&cache), &read, &[], &[DecodedValue::Str("committed".into())]).unwrap();
assert_eq!(get_rows(tx, &read, &[]).unwrap(), None);
invalidate_for(tx, &compiled("insert person", &["public.person"], true)).unwrap();
assert_eq!(
get_rows(rw(&cache), &read, &[]).unwrap(),
None,
"a transactional write must evict"
);
}
#[test]
fn a_read_never_evicts() {
let (_dir, cache) = open_temp();
let read = compiled("select person", &["public.person"], false);
put_rows(rw(&cache), &read, &[], &[DecodedValue::Str("row".into())]).unwrap();
invalidate_for(rw(&cache), &read).unwrap();
assert!(get_rows(rw(&cache), &read, &[]).unwrap().is_some());
}
#[test]
fn rows_miss_then_hit() {
let (_dir, cache) = open_temp();
let compiled = compiled_with_tags("select 1", &["public.person"]);
assert_eq!(get_rows(rw(&cache), &compiled, &[]).unwrap(), None);
put_rows(rw(&cache), &compiled, &[], &[DecodedValue::I64(1)]).unwrap();
assert_eq!(
get_rows(rw(&cache), &compiled, &[]).unwrap(),
Some(vec![DecodedValue::I64(1)])
);
}
#[test]
fn no_tags_means_put_is_a_no_op() {
let (_dir, cache) = open_temp();
let compiled = compiled_with_tags("select 1", &[]);
put_rows(rw(&cache), &compiled, &[], &[DecodedValue::I64(1)]).unwrap();
assert_eq!(get_rows(rw(&cache), &compiled, &[]).unwrap(), None);
}
#[test]
fn rows_and_json_kinds_do_not_collide_on_the_same_sql() {
let (_dir, cache) = open_temp();
let compiled = compiled_with_tags("select 1", &["public.person"]);
put_rows(rw(&cache), &compiled, &[], &[DecodedValue::I64(1)]).unwrap();
put_json(rw(&cache), "json_all", &compiled, &[], Some("[1]")).unwrap();
assert_eq!(
get_rows(rw(&cache), &compiled, &[]).unwrap(),
Some(vec![DecodedValue::I64(1)])
);
assert_eq!(
get_json(rw(&cache), "json_all", &compiled, &[]).unwrap(),
Some(Some("[1]".to_string()))
);
assert_eq!(get_json(rw(&cache), "json_single", &compiled, &[]).unwrap(), None);
}
#[test]
fn json_single_caches_a_legitimately_empty_result_distinct_from_a_miss() {
let (_dir, cache) = open_temp();
let compiled = compiled_with_tags("select Person filter false", &["public.person"]);
assert_eq!(get_json(rw(&cache), "json_single", &compiled, &[]).unwrap(), None);
put_json(rw(&cache), "json_single", &compiled, &[], None).unwrap();
assert_eq!(get_json(rw(&cache), "json_single", &compiled, &[]).unwrap(), Some(None));
}
}