use std::fs::File;
use std::io::prelude::*;
use std::path::PathBuf;
use std::sync::RwLock;
use crate::crypto;
use crate::error::{Error, Result};
use crate::etree::ParseOps;
pub trait CasStore: Send + Sync {
fn save(&self, blob: &[u8], policy: &dyn crypto::CryptoPolicy) -> Result<String>;
fn load(&self, hash: &str, policy: &dyn crypto::CryptoPolicy) -> Result<Vec<u8>>;
fn contains(&self, hash: &str, policy: &dyn crypto::CryptoPolicy) -> Result<bool> {
Ok(self.load(hash, policy).is_ok())
}
}
pub struct LocalCas {
pub root: PathBuf,
pub verbose: bool,
}
impl LocalCas {
pub fn new(root: PathBuf) -> Self {
LocalCas {
root,
verbose: false,
}
}
fn path_for(&self, hash: &str) -> PathBuf {
self.root.join(hash)
}
}
impl CasStore for LocalCas {
fn save(&self, blob: &[u8], policy: &dyn crypto::CryptoPolicy) -> Result<String> {
let hexhash = crypto::hexdigest("sha3-256", blob, policy)?;
let path = self.path_for(&hexhash);
if path.is_file() {
if self.verbose {
eprintln!("cas::save(): {} already exists. Exiting.", path.display());
}
return Ok(hexhash);
}
let mut file_out = File::create(&path)
.map_err(|e| Error::Cas(format!("Failed to open {}: {}", path.display(), e)))?;
let bytes = file_out.write(blob).map_err(|e| {
Error::Cas(format!(
"Error writing {} bytes to {}: {}",
blob.len(),
path.display(),
e
))
})?;
if self.verbose {
eprintln!("cas::save(): {} bytes to {}", bytes, path.display());
}
Ok(hexhash)
}
fn load(&self, hash: &str, policy: &dyn crypto::CryptoPolicy) -> Result<Vec<u8>> {
hex::decode(hash).map_err(|_| Error::Cas(format!("Not a valid hex token: {}", hash)))?;
let path = self.path_for(hash);
let mut file_in = File::open(&path)
.map_err(|e| Error::Cas(format!("Failed to open {}: {}", path.display(), e)))?;
let mut blob = Vec::new();
let bytes = file_in
.read_to_end(&mut blob)
.map_err(|e| Error::Cas(format!("Error reading {}: {}", path.display(), e)))?;
if self.verbose {
eprintln!("cas::load(): {} bytes from {}", bytes, path.display());
}
let verify = crypto::hexdigest("sha3-256", &blob, policy)?;
if hash != verify {
return Err(Error::Cas(format!(
"CONTENT HASH MISMATCH!\ninput = {}\ncheck = {}",
hash, verify
)));
}
Ok(blob)
}
}
#[allow(dead_code)]
pub struct MemoryCas {
entries: RwLock<std::collections::BTreeMap<String, Vec<u8>>>,
}
impl MemoryCas {
#[allow(dead_code)]
pub fn new() -> Self {
MemoryCas {
entries: RwLock::new(std::collections::BTreeMap::new()),
}
}
}
impl Default for MemoryCas {
fn default() -> Self {
Self::new()
}
}
impl CasStore for MemoryCas {
fn save(&self, blob: &[u8], policy: &dyn crypto::CryptoPolicy) -> Result<String> {
let hexhash = crypto::hexdigest("sha3-256", blob, policy)?;
let mut map = self.entries.write().unwrap();
map.entry(hexhash.clone()).or_insert_with(|| blob.to_vec());
Ok(hexhash)
}
fn load(&self, hash: &str, policy: &dyn crypto::CryptoPolicy) -> Result<Vec<u8>> {
let map = self.entries.read().unwrap();
let blob = map
.get(hash)
.ok_or_else(|| Error::Cas(format!("hash not present in memory CAS: {}", hash)))?
.clone();
let verify = crypto::hexdigest("sha3-256", &blob, policy)?;
if hash != verify {
return Err(Error::Cas(format!(
"CONTENT HASH MISMATCH (memory)!\ninput = {}\ncheck = {}",
hash, verify
)));
}
Ok(blob)
}
fn contains(&self, _hash: &str, _policy: &dyn crypto::CryptoPolicy) -> Result<bool> {
Ok(self.entries.read().unwrap().contains_key(_hash))
}
}
pub fn load(hexhash: &str, paops: &mut ParseOps) -> Result<Vec<u8>> {
let policy: &dyn crypto::CryptoPolicy = &*paops.crypto.policy;
paops.io.cas.load(hexhash, policy)
}
pub fn save(blob: Vec<u8>, paops: &mut ParseOps) -> Result<String> {
let policy: &dyn crypto::CryptoPolicy = &*paops.crypto.policy;
paops.io.cas.save(&blob, policy)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::crypto::default_policy;
#[test]
fn local_cas_round_trip() {
let dir = tempfile::tempdir().unwrap();
let store = LocalCas::new(dir.path().to_path_buf());
let policy = default_policy();
let blob = b"hello, world";
let h = store.save(blob, &*policy).unwrap();
assert_eq!(h.len(), 64);
let loaded = store.load(&h, &*policy).unwrap();
assert_eq!(loaded, blob);
assert!(store.contains(&h, &*policy).unwrap());
assert!(!store.contains("0".repeat(64).as_str(), &*policy).unwrap());
}
#[test]
fn local_cas_idempotent() {
let dir = tempfile::tempdir().unwrap();
let store = LocalCas::new(dir.path().to_path_buf());
let policy = default_policy();
let blob = b"same content";
let h1 = store.save(blob, &*policy).unwrap();
let h2 = store.save(blob, &*policy).unwrap();
assert_eq!(h1, h2);
}
#[test]
fn local_cas_detects_hash_mismatch() {
let dir = tempfile::tempdir().unwrap();
let store = LocalCas::new(dir.path().to_path_buf());
let policy = default_policy();
let bad_hash = "0".repeat(64);
std::fs::write(dir.path().join(&bad_hash), b"unrelated content").unwrap();
let err = store.load(&bad_hash, &*policy);
assert!(err.is_err());
let msg = err.unwrap_err().to_string();
assert!(msg.contains("HASH MISMATCH"), "msg: {msg}");
}
#[test]
fn memory_cas_round_trip() {
let store = MemoryCas::new();
let policy = default_policy();
let blob = b"in-memory blob";
let h = store.save(blob, &*policy).unwrap();
let loaded = store.load(&h, &*policy).unwrap();
assert_eq!(loaded, blob);
assert!(store.contains(&h, &*policy).unwrap());
}
#[test]
fn memory_cas_idempotent() {
let store = MemoryCas::new();
let policy = default_policy();
let blob = b"x";
let h1 = store.save(blob, &*policy).unwrap();
let h2 = store.save(blob, &*policy).unwrap();
assert_eq!(h1, h2);
}
}