use std::path::PathBuf;
use std::sync::OnceLock;
use anyhow::{Context, Result};
use cute_sqlite_kv::BlobStore;
use sha2::{Digest, Sha256};
use tracing::{info, warn};
use crate::problem::parse::PuzzleParse;
use crate::problem::util::exec::ProgramRunner;
const CACHE_VERSION: &str = "parse-v1";
const ZSTD_LEVEL: i32 = 6;
fn cache_db_path() -> Option<PathBuf> {
let dir = match std::env::var("DEMYSTIFY_PARSE_CACHE") {
Ok(v) => {
let t = v.trim();
if t.is_empty() || t.eq_ignore_ascii_case("off") || t == "0" || t == "false" {
return None;
}
PathBuf::from(t)
}
Err(_) => std::env::temp_dir().join("demystify-parse-cache"),
};
Some(dir.join("parse.sqlite"))
}
fn tool_versions() -> Result<&'static (String, String)> {
static VERSIONS: OnceLock<(String, String)> = OnceLock::new();
if VERSIONS.get().is_none() {
let sr = ProgramRunner::get_savilerow_version()
.map_err(|e| anyhow::anyhow!("parse cache: reading Savile Row version: {e}"))?;
let conjure = ProgramRunner::get_conjure_version()
.map_err(|e| anyhow::anyhow!("parse cache: reading Conjure version: {e}"))?;
let _ = VERSIONS.set((sr, conjure));
}
Ok(VERSIONS.get().expect("just initialised"))
}
fn hash_key(
src_hash: &str,
sr_ver: &str,
conjure_ver: &str,
model_ext: &str,
model_bytes: &[u8],
param_bytes: &[u8],
) -> String {
let mut h = Sha256::new();
let mut field = |label: &str, bytes: &[u8]| {
h.update(label.as_bytes());
h.update((bytes.len() as u64).to_le_bytes());
h.update(bytes);
};
field("src", src_hash.as_bytes());
field("savilerow", sr_ver.as_bytes());
field("conjure", conjure_ver.as_bytes());
field("ext", model_ext.as_bytes());
field("model", model_bytes);
field("param", param_bytes);
format!("{CACHE_VERSION}-{:x}", h.finalize())
}
pub fn cache_key(model_bytes: &[u8], param_bytes: &[u8], model_ext: &str) -> Result<String> {
let (sr_ver, conjure_ver) = tool_versions()?;
Ok(hash_key(
env!("DEMYSTIFY_SRC_HASH"),
sr_ver,
conjure_ver,
model_ext,
model_bytes,
param_bytes,
))
}
fn open_store() -> Result<Option<BlobStore>> {
let Some(path) = cache_db_path() else {
return Ok(None);
};
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)
.with_context(|| format!("creating parse cache directory {parent:?}"))?;
}
let store = BlobStore::new_from_file(&path)
.with_context(|| format!("opening parse cache at {path:?}"))?;
Ok(Some(store))
}
pub fn try_load(key: &str) -> Result<Option<PuzzleParse>> {
let Some(store) = open_store()? else {
return Ok(None);
};
let Some(compressed) = store.get(key) else {
return Ok(None);
};
let json = zstd::decode_all(compressed.as_slice())
.context("parse cache: decompressing cached parse (cache may be corrupt)")?;
let puzzle = PuzzleParse::from_json_bytes(&json)
.context("parse cache: deserializing cached parse (cache may be corrupt)")?;
Ok(Some(puzzle))
}
pub fn store(key: &str, puzzle: &PuzzleParse) -> Result<()> {
let Some(store) = open_store()? else {
return Ok(());
};
let json = puzzle.to_json_bytes()?;
let compressed =
zstd::encode_all(json.as_slice(), ZSTD_LEVEL).context("parse cache: compressing parse")?;
store.insert(key, &compressed);
info!(target: "progress", "stored parse in cache ({} bytes compressed)", compressed.len());
Ok(())
}
pub fn log_status() {
match cache_db_path() {
Some(p) => info!(target: "progress", "parse cache enabled at {p:?}"),
None => warn!(target: "progress", "parse cache disabled (DEMYSTIFY_PARSE_CACHE=off)"),
}
}
#[cfg(test)]
mod tests {
use super::hash_key;
fn key(ext: &str, model: &[u8], param: &[u8]) -> String {
hash_key("srchash", "sr-1.10", "conjure-2.5", ext, model, param)
}
#[test]
fn key_is_deterministic() {
assert_eq!(
key("eprime", b"model", b"param"),
key("eprime", b"model", b"param")
);
}
#[test]
fn key_changes_with_every_field() {
let base = key("eprime", b"model", b"param");
assert_ne!(
base,
hash_key(
"OTHER",
"sr-1.10",
"conjure-2.5",
"eprime",
b"model",
b"param"
)
);
assert_ne!(
base,
hash_key(
"srchash",
"sr-9.99",
"conjure-2.5",
"eprime",
b"model",
b"param"
)
);
assert_ne!(
base,
hash_key(
"srchash",
"sr-1.10",
"conjure-9.9",
"eprime",
b"model",
b"param"
)
);
assert_ne!(base, key("essence", b"model", b"param"));
assert_ne!(base, key("eprime", b"MODEL", b"param"));
assert_ne!(base, key("eprime", b"model", b"PARAM"));
}
#[test]
fn key_fields_are_length_delimited() {
assert_ne!(key("eprime", b"ab", b"c"), key("eprime", b"a", b"bc"));
}
}