use std::collections::HashMap;
use std::io::Write;
use std::path::{Path, PathBuf};
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::mpsc;
use std::sync::{Arc, Condvar, Mutex};
use std::time::Duration;
use crc::{Crc, CRC_32_ISCSI};
use crate::encryption::{AesCipher, Cipher};
use crate::engine::{LookupMetrics, LookupMetricsSnapshot};
pub const FRAME_MAGIC: [u8; 4] = *b"MLCP";
pub const FRAME_FORMAT_VERSION: u16 = 1;
pub(crate) const CLEAR_KEY: u64 = u64::MAX;
const FRAME_PREFIX_LEN: usize = 4 + 2 + 2 + 8 + 8 + 8 + 8 + 8;
const FRAME_TRAILER_LEN: usize = 4;
const FRAME_HEADER_LEN: usize = FRAME_PREFIX_LEN + 4;
const FRAME_FULL_OVERHEAD: usize = FRAME_HEADER_LEN + FRAME_TRAILER_LEN;
const NONCE_LEN: usize = 12;
const CRC32C: Crc<u32> = Crc::<u32>::new(&CRC_32_ISCSI);
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PersistedHeader {
pub format_version: u16,
pub table_id: u64,
pub schema_id: u64,
pub run_generation: u64,
pub cache_key: u64,
pub entry_generation: u64,
pub payload_len: u32,
}
#[derive(Debug, Clone)]
pub struct PersistedFrame {
pub header: PersistedHeader,
pub payload: Vec<u8>,
}
pub fn encode_frame(frame: &PersistedFrame) -> Vec<u8> {
let mut out = Vec::with_capacity(FRAME_FULL_OVERHEAD + frame.payload.len());
out.extend_from_slice(&FRAME_MAGIC);
out.extend_from_slice(&frame.header.format_version.to_le_bytes());
out.extend_from_slice(&0u16.to_le_bytes()); out.extend_from_slice(&frame.header.table_id.to_le_bytes());
out.extend_from_slice(&frame.header.schema_id.to_le_bytes());
out.extend_from_slice(&frame.header.run_generation.to_le_bytes());
out.extend_from_slice(&frame.header.cache_key.to_le_bytes());
out.extend_from_slice(&frame.header.entry_generation.to_le_bytes());
out.extend_from_slice(&frame.header.payload_len.to_le_bytes());
out.extend_from_slice(&frame.payload);
let crc = CRC32C.checksum(&out[0..FRAME_PREFIX_LEN + 4 + frame.payload.len()]);
out.extend_from_slice(&crc.to_le_bytes());
out
}
pub fn decode_frame(bytes: &[u8]) -> Option<PersistedFrame> {
if bytes.len() < FRAME_FULL_OVERHEAD {
return None;
}
let body_len = bytes.len() - FRAME_TRAILER_LEN;
let stored_crc = u32::from_le_bytes(bytes[body_len..].try_into().ok()?);
let actual_crc = CRC32C.checksum(&bytes[..body_len]);
if actual_crc != stored_crc {
return None;
}
let mut p = &bytes[..body_len];
let mut magic = [0u8; 4];
magic.copy_from_slice(&p[..4]);
if magic != FRAME_MAGIC {
return None;
}
p = &p[4..];
let format_version = u16::from_le_bytes(p[..2].try_into().ok()?);
p = &p[2..];
let _reserved = u16::from_le_bytes(p[..2].try_into().ok()?);
p = &p[2..];
let table_id = u64::from_le_bytes(p[..8].try_into().ok()?);
p = &p[8..];
let schema_id = u64::from_le_bytes(p[..8].try_into().ok()?);
p = &p[8..];
let run_generation = u64::from_le_bytes(p[..8].try_into().ok()?);
p = &p[8..];
let cache_key = u64::from_le_bytes(p[..8].try_into().ok()?);
p = &p[8..];
let entry_generation = u64::from_le_bytes(p[..8].try_into().ok()?);
p = &p[8..];
let payload_len = u32::from_le_bytes(p[..4].try_into().ok()?) as usize;
p = &p[4..];
if p.len() < payload_len {
return None;
}
let payload = p[..payload_len].to_vec();
Some(PersistedFrame {
header: PersistedHeader {
format_version,
table_id,
schema_id,
run_generation,
cache_key,
entry_generation,
payload_len: payload_len as u32,
},
payload,
})
}
pub fn read_header_only(bytes: &[u8]) -> Option<PersistedHeader> {
decode_frame(bytes).map(|f| f.header)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct PersistentCacheIdentity {
pub table_id: u64,
pub schema_id: u64,
pub logical_generation: u64,
}
#[derive(Debug, Clone, Copy)]
pub struct PersistContext {
pub identity: PersistentCacheIdentity,
pub key: u64,
pub entry_generation: u64,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PersistenceDisabledReason {
NoDirectory,
WorkerSpawnFailed,
WorkerShutdown,
QueueUnavailable,
}
impl PersistenceDisabledReason {
pub fn label(&self) -> &'static str {
match self {
Self::NoDirectory => "no_directory",
Self::WorkerSpawnFailed => "worker_spawn_failed",
Self::WorkerShutdown => "worker_shutdown",
Self::QueueUnavailable => "queue_unavailable",
}
}
}
impl std::fmt::Display for PersistenceDisabledReason {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.label())
}
}
pub enum PersistentPublicationState {
Async(Arc<PersistentResultCacheWriter>),
Disabled(PersistenceDisabledReason),
}
impl std::fmt::Debug for PersistentPublicationState {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Async(_) => f.write_str("PersistentPublicationState::Async(..)"),
Self::Disabled(reason) => write!(f, "PersistentPublicationState::Disabled({reason:?})"),
}
}
}
#[derive(Debug)]
pub struct CachePersistError {
pub message: String,
}
impl std::fmt::Display for CachePersistError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(&self.message)
}
}
impl std::error::Error for CachePersistError {}
#[derive(Debug, Clone, PartialEq, Eq)]
#[allow(dead_code)] pub enum CacheLoadRejection {
LegacyUnframed,
Truncated,
BadCrc,
UnsupportedVersion(u16),
TableIdMismatch { expected: u64, found: u64 },
SchemaIdMismatch { expected: u64, found: u64 },
GenerationMismatch { expected: u64, found: u64 },
KeyMismatch { expected: u64, found: u64 },
MissingCipher,
DecryptionFailed,
PayloadInvalid,
}
impl std::fmt::Display for CacheLoadRejection {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::LegacyUnframed => f.write_str("legacy unframed data"),
Self::Truncated => f.write_str("truncated frame"),
Self::BadCrc => f.write_str("bad crc"),
Self::UnsupportedVersion(v) => write!(f, "unsupported format version {v}"),
Self::TableIdMismatch { expected, found } => {
write!(f, "table id mismatch: expected {expected}, found {found}")
}
Self::SchemaIdMismatch { expected, found } => {
write!(f, "schema id mismatch: expected {expected}, found {found}")
}
Self::GenerationMismatch { expected, found } => write!(
f,
"logical generation mismatch: expected {expected}, found {found}"
),
Self::KeyMismatch { expected, found } => {
write!(f, "cache key mismatch: expected {expected}, found {found}")
}
Self::MissingCipher => f.write_str("cipher required for encrypted frame"),
Self::DecryptionFailed => f.write_str("decryption failed"),
Self::PayloadInvalid => f.write_str("payload deserialize failed"),
}
}
}
impl std::error::Error for CacheLoadRejection {}
pub fn encode_persisted_entry(
context: PersistContext,
payload: &[u8],
cipher: Option<&AesCipher>,
) -> Result<Vec<u8>, CachePersistError> {
let inner = encrypt_payload(cipher, payload).map_err(|e| CachePersistError {
message: format!("encrypt_persisted_entry: {:?}", e),
})?;
let frame = PersistedFrame {
header: PersistedHeader {
format_version: FRAME_FORMAT_VERSION,
table_id: context.identity.table_id,
schema_id: context.identity.schema_id,
run_generation: context.identity.logical_generation,
cache_key: context.key,
entry_generation: context.entry_generation,
payload_len: inner.len() as u32,
},
payload: inner,
};
Ok(encode_frame(&frame))
}
pub fn decode_persisted_entry(
expected: PersistentCacheIdentity,
expected_key: u64,
bytes: &[u8],
cipher: Option<&AesCipher>,
) -> Result<Vec<u8>, CacheLoadRejection> {
if bytes.len() < FRAME_MAGIC.len() || bytes[..FRAME_MAGIC.len()] != FRAME_MAGIC {
return Err(CacheLoadRejection::LegacyUnframed);
}
let frame = decode_frame(bytes).ok_or(CacheLoadRejection::BadCrc)?;
if frame.header.format_version != FRAME_FORMAT_VERSION {
return Err(CacheLoadRejection::UnsupportedVersion(
frame.header.format_version,
));
}
if frame.header.table_id != expected.table_id {
return Err(CacheLoadRejection::TableIdMismatch {
expected: expected.table_id,
found: frame.header.table_id,
});
}
if frame.header.schema_id != expected.schema_id {
return Err(CacheLoadRejection::SchemaIdMismatch {
expected: expected.schema_id,
found: frame.header.schema_id,
});
}
if frame.header.run_generation != expected.logical_generation {
return Err(CacheLoadRejection::GenerationMismatch {
expected: expected.logical_generation,
found: frame.header.run_generation,
});
}
if frame.header.cache_key != expected_key {
return Err(CacheLoadRejection::KeyMismatch {
expected: expected_key,
found: frame.header.cache_key,
});
}
let plaintext = decrypt_payload(cipher, &frame.payload).map_err(|_e| {
CacheLoadRejection::DecryptionFailed
})?;
plaintext.ok_or(CacheLoadRejection::DecryptionFailed)
}
#[derive(Debug)]
pub enum PendingCacheOp {
Store(PersistableEntry),
Remove,
Clear,
}
pub struct PersistableEntry {
pub key: u64,
pub table_id: u64,
pub schema_id: u64,
pub run_generation: u64,
pub entry_generation: u64,
pub bytes: usize,
pub payload_factory: Box<dyn FnOnce() -> Option<Vec<u8>> + Send + 'static>,
}
impl std::fmt::Debug for PersistableEntry {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PersistableEntry")
.field("key", &self.key)
.field("table_id", &self.table_id)
.field("schema_id", &self.schema_id)
.field("run_generation", &self.run_generation)
.field("entry_generation", &self.entry_generation)
.field("bytes", &self.bytes)
.finish_non_exhaustive()
}
}
#[derive(Debug)]
pub struct PendingCacheState {
pub clear_generation: u64,
pub next_key_generation: u64,
pub key_generations: HashMap<u64, u64>,
pub operations: HashMap<u64, PendingCacheOp>,
pub approx_bytes: usize,
}
impl Default for PendingCacheState {
fn default() -> Self {
Self {
clear_generation: 0,
next_key_generation: 1,
key_generations: HashMap::new(),
operations: HashMap::new(),
approx_bytes: 0,
}
}
}
#[derive(Debug, Clone)]
pub struct WriterLimits {
pub max_pending_keys: usize,
pub max_pending_bytes: usize,
}
impl Default for WriterLimits {
fn default() -> Self {
Self {
max_pending_keys: 16_384,
max_pending_bytes: 64 * 1024 * 1024,
}
}
}
struct Inner {
state: PendingCacheState,
limits: WriterLimits,
shutdown: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DrainOutcome {
StorePublished,
StoreStale,
StoreErrored,
RemoveApplied,
RemoveErrored,
ClearApplied,
ClearErrored,
Abandoned,
}
#[derive(Debug)]
pub struct DrainedOp {
pub key: u64,
pub op: PendingCacheOp,
pub key_generation: u64,
pub clear_generation: u64,
}
pub struct PersistentResultCacheWriter {
inner: Mutex<Inner>,
cond: Condvar,
metrics: LookupMetrics,
enqueued_total: AtomicU64,
coalesced_total: AtomicU64,
dropped_total: AtomicU64,
remove_total: AtomicU64,
stale_total: AtomicU64,
errors_total: AtomicU64,
abandoned_total: AtomicU64,
queue_depth: AtomicU64,
writes_in_flight: AtomicU64,
}
impl PersistentResultCacheWriter {
pub fn new(metrics: LookupMetrics, limits: WriterLimits) -> Self {
Self {
inner: Mutex::new(Inner {
state: PendingCacheState::default(),
limits,
shutdown: false,
}),
cond: Condvar::new(),
metrics,
enqueued_total: AtomicU64::new(0),
coalesced_total: AtomicU64::new(0),
dropped_total: AtomicU64::new(0),
remove_total: AtomicU64::new(0),
stale_total: AtomicU64::new(0),
errors_total: AtomicU64::new(0),
abandoned_total: AtomicU64::new(0),
queue_depth: AtomicU64::new(0),
writes_in_flight: AtomicU64::new(0),
}
}
pub fn enqueue_store(&self, entry: PersistableEntry) {
let mut guard = self.inner.lock().expect("writer mutex poisoned");
let key = entry.key;
let bytes = entry.bytes;
let already_present = guard.state.operations.contains_key(&key);
let would_grow = !already_present;
let key_cap = guard.limits.max_pending_keys;
let byte_cap = guard.limits.max_pending_bytes;
if (would_grow && guard.state.operations.len() >= key_cap)
|| guard.state.approx_bytes.saturating_add(bytes) > byte_cap
{
self.dropped_total.fetch_add(1, Ordering::Relaxed);
self.metrics
.result_cache_persist_dropped_store_total
.fetch_add(1, Ordering::Relaxed);
return;
}
if let Some(PendingCacheOp::Remove) = guard.state.operations.get(&key) {
self.dropped_total.fetch_add(1, Ordering::Relaxed);
self.metrics
.result_cache_persist_dropped_store_total
.fetch_add(1, Ordering::Relaxed);
return;
}
if let Some(PendingCacheOp::Store(prev)) = guard.state.operations.get(&key) {
guard.state.approx_bytes = guard.state.approx_bytes.saturating_sub(prev.bytes);
self.coalesced_total.fetch_add(1, Ordering::Relaxed);
self.metrics
.result_cache_persist_coalesced_total
.fetch_add(1, Ordering::Relaxed);
} else {
let gen = guard.state.next_key_generation;
guard.state.next_key_generation = guard.state.next_key_generation.wrapping_add(1);
guard.state.key_generations.insert(key, gen);
}
guard
.state
.operations
.insert(key, PendingCacheOp::Store(entry));
guard.state.approx_bytes = guard.state.approx_bytes.saturating_add(bytes);
let depth = guard.state.operations.len() as u64;
drop(guard);
self.queue_depth.store(depth, Ordering::Relaxed);
self.enqueued_total.fetch_add(1, Ordering::Relaxed);
self.metrics
.result_cache_persist_enqueued_total
.fetch_add(1, Ordering::Relaxed);
self.cond.notify_one();
}
pub fn enqueue_remove(&self, key: u64) {
let mut guard = self.inner.lock().expect("writer mutex poisoned");
let new_gen = guard.state.next_key_generation;
guard.state.next_key_generation = guard.state.next_key_generation.wrapping_add(1);
guard.state.key_generations.insert(key, new_gen);
if let Some(PendingCacheOp::Store(prev)) = guard.state.operations.get(&key) {
guard.state.approx_bytes = guard.state.approx_bytes.saturating_sub(prev.bytes);
}
guard.state.operations.insert(key, PendingCacheOp::Remove);
let depth = guard.state.operations.len() as u64;
drop(guard);
self.remove_total.fetch_add(1, Ordering::Relaxed);
self.metrics
.result_cache_persist_remove_total
.fetch_add(1, Ordering::Relaxed);
self.queue_depth.store(depth, Ordering::Relaxed);
self.cond.notify_one();
}
pub fn enqueue_clear(&self) {
let mut guard = self.inner.lock().expect("writer mutex poisoned");
guard.state.clear_generation = guard.state.next_key_generation;
guard.state.next_key_generation = guard.state.next_key_generation.wrapping_add(1);
let mut to_sub = 0usize;
for entry in guard.state.operations.values() {
if let PendingCacheOp::Store(s) = entry {
to_sub = to_sub.saturating_add(s.bytes);
}
}
guard.state.approx_bytes = guard.state.approx_bytes.saturating_sub(to_sub);
guard.state.operations.clear();
guard.state.key_generations.clear();
guard
.state
.operations
.insert(CLEAR_KEY, PendingCacheOp::Clear);
let depth = guard.state.operations.len() as u64;
drop(guard);
self.queue_depth.store(depth, Ordering::Relaxed);
self.cond.notify_all();
}
pub fn drain_one(&self) -> Option<DrainedOp> {
let mut guard = self.inner.lock().expect("writer mutex poisoned");
loop {
if let Some(next_key) = guard.state.operations.keys().next().copied() {
let op = guard
.state
.operations
.remove(&next_key)
.expect("just observed");
let key_gen = guard
.state
.key_generations
.get(&next_key)
.copied()
.unwrap_or(0);
if let PendingCacheOp::Store(s) = &op {
guard.state.approx_bytes = guard.state.approx_bytes.saturating_sub(s.bytes);
}
guard.state.key_generations.remove(&next_key);
let clear_gen = guard.state.clear_generation;
let depth = guard.state.operations.len() as u64;
drop(guard);
self.queue_depth.store(depth, Ordering::Relaxed);
return Some(DrainedOp {
key: next_key,
op,
key_generation: key_gen,
clear_generation: clear_gen,
});
}
if guard.shutdown {
return None;
}
guard = self.cond.wait(guard).expect("condvar wait poisoned");
}
}
pub fn try_drain_one(&self) -> Option<DrainedOp> {
let mut guard = self.inner.lock().expect("writer mutex poisoned");
let next_key = {
let mut iter = guard.state.operations.iter();
let (k, _) = iter.next()?;
*k
};
let op = guard
.state
.operations
.remove(&next_key)
.expect("just observed");
let key_gen = guard
.state
.key_generations
.get(&next_key)
.copied()
.unwrap_or(0);
if let PendingCacheOp::Store(s) = &op {
guard.state.approx_bytes = guard.state.approx_bytes.saturating_sub(s.bytes);
}
guard.state.key_generations.remove(&next_key);
let clear_gen = guard.state.clear_generation;
let depth = guard.state.operations.len() as u64;
drop(guard);
self.queue_depth.store(depth, Ordering::Relaxed);
Some(DrainedOp {
key: next_key,
op,
key_generation: key_gen,
clear_generation: clear_gen,
})
}
pub fn shutdown(&self) {
let mut guard = self.inner.lock().expect("writer mutex poisoned");
guard.shutdown = true;
drop(guard);
self.cond.notify_all();
}
pub fn pending_count(&self) -> u64 {
self.queue_depth.load(Ordering::Relaxed)
}
pub fn drain_all_as_abandoned(&self) -> u64 {
let mut guard = self.inner.lock().expect("writer mutex poisoned");
let mut n = 0u64;
let mut bytes = 0usize;
for (_, op) in guard.state.operations.drain() {
if let PendingCacheOp::Store(s) = op {
bytes = bytes.saturating_add(s.bytes);
}
n += 1;
}
guard.state.key_generations.clear();
guard.state.approx_bytes = 0;
drop(guard);
if n > 0 {
self.abandoned_total.fetch_add(n, Ordering::Relaxed);
self.metrics
.result_cache_persist_shutdown_abandoned_total
.fetch_add(n, Ordering::Relaxed);
let _ = bytes; }
self.queue_depth.store(0, Ordering::Relaxed);
self.cond.notify_all();
n
}
pub fn record_outcome(&self, outcome: DrainOutcome) {
match outcome {
DrainOutcome::StorePublished => {}
DrainOutcome::RemoveApplied => {}
DrainOutcome::ClearApplied => {}
DrainOutcome::StoreStale => {
self.stale_total.fetch_add(1, Ordering::Relaxed);
self.metrics
.result_cache_persist_stale_store_skipped_total
.fetch_add(1, Ordering::Relaxed);
}
DrainOutcome::StoreErrored
| DrainOutcome::RemoveErrored
| DrainOutcome::ClearErrored => {
self.errors_total.fetch_add(1, Ordering::Relaxed);
self.metrics
.result_cache_persist_errors_total
.fetch_add(1, Ordering::Relaxed);
}
DrainOutcome::Abandoned => {
self.abandoned_total.fetch_add(1, Ordering::Relaxed);
self.metrics
.result_cache_persist_shutdown_abandoned_total
.fetch_add(1, Ordering::Relaxed);
}
}
}
pub fn enqueued_total(&self) -> u64 {
self.enqueued_total.load(Ordering::Relaxed)
}
pub fn coalesced_total(&self) -> u64 {
self.coalesced_total.load(Ordering::Relaxed)
}
pub fn dropped_total(&self) -> u64 {
self.dropped_total.load(Ordering::Relaxed)
}
pub fn remove_total(&self) -> u64 {
self.remove_total.load(Ordering::Relaxed)
}
pub fn stale_total(&self) -> u64 {
self.stale_total.load(Ordering::Relaxed)
}
pub fn errors_total(&self) -> u64 {
self.errors_total.load(Ordering::Relaxed)
}
pub fn abandoned_total(&self) -> u64 {
self.abandoned_total.load(Ordering::Relaxed)
}
pub fn queue_depth(&self) -> usize {
self.queue_depth.load(Ordering::Relaxed) as usize
}
pub fn writes_in_flight(&self) -> u64 {
self.writes_in_flight.load(Ordering::Relaxed)
}
pub fn persist_snapshot(&self) -> LookupMetricsSnapshot {
LookupMetricsSnapshot {
result_cache_persist_enqueued_total: self.enqueued_total.load(Ordering::Relaxed),
result_cache_persist_coalesced_total: self.coalesced_total.load(Ordering::Relaxed),
result_cache_persist_dropped_store_total: self.dropped_total.load(Ordering::Relaxed),
result_cache_persist_remove_total: self.remove_total.load(Ordering::Relaxed),
result_cache_persist_stale_store_skipped_total: self
.stale_total
.load(Ordering::Relaxed),
result_cache_persist_errors_total: self.errors_total.load(Ordering::Relaxed),
result_cache_persist_shutdown_abandoned_total: self
.abandoned_total
.load(Ordering::Relaxed),
result_cache_persist_queue_depth: self.queue_depth.load(Ordering::Relaxed),
..LookupMetricsSnapshot::default()
}
}
pub fn for_test(limits: WriterLimits) -> Self {
Self::new(LookupMetrics::default(), limits)
}
pub fn bump_persist_generation(&self, key: u64) -> u64 {
let mut guard = self.inner.lock().expect("writer mutex poisoned");
let gen = guard.state.next_key_generation;
guard.state.next_key_generation = guard.state.next_key_generation.wrapping_add(1);
guard.state.key_generations.insert(key, gen);
gen
}
}
pub trait PersistentCacheIo: Send + Sync {
fn write_atomic(&self, key: u64, frame: &[u8]) -> Result<(), IoError>;
fn remove(&self, key: u64) -> Result<(), IoError>;
fn clear(&self) -> Result<(), IoError>;
fn load(&self, key: u64) -> Result<Option<Vec<u8>>, IoError>;
fn exists(&self, key: u64) -> bool;
}
#[derive(Debug, Clone)]
pub struct IoError {
pub kind: IoErrorKind,
pub message: String,
}
impl IoError {
pub fn new(kind: IoErrorKind, message: impl Into<String>) -> Self {
Self {
kind,
message: message.into(),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum IoErrorKind {
Other,
}
pub struct RealPersistentCacheIo {
dir: PathBuf,
}
impl RealPersistentCacheIo {
pub fn new(dir: PathBuf) -> std::io::Result<Self> {
std::fs::create_dir_all(&dir)?;
Ok(Self { dir })
}
fn final_path(&self, key: u64) -> PathBuf {
self.dir.join(format!("{key:016x}.bin"))
}
fn temp_path(&self, key: u64) -> PathBuf {
self.dir.join(format!("{key:016x}.bin.tmp"))
}
}
impl PersistentCacheIo for RealPersistentCacheIo {
fn write_atomic(&self, key: u64, frame: &[u8]) -> Result<(), IoError> {
let final_path = self.final_path(key);
let tmp_path = self.temp_path(key);
let _ = std::fs::remove_file(&tmp_path);
{
let mut f = std::fs::File::create(&tmp_path).map_err(|e| {
IoError::new(
IoErrorKind::Other,
format!("create {}: {e}", tmp_path.display()),
)
})?;
f.write_all(frame)
.map_err(|e| IoError::new(IoErrorKind::Other, format!("write tmp: {e}")))?;
f.flush()
.map_err(|e| IoError::new(IoErrorKind::Other, format!("flush tmp: {e}")))?;
f.sync_all()
.map_err(|e| IoError::new(IoErrorKind::Other, format!("fsync tmp: {e}")))?;
}
std::fs::rename(&tmp_path, &final_path).map_err(|e| {
let _ = std::fs::remove_file(&tmp_path);
IoError::new(
IoErrorKind::Other,
format!(
"rename {} -> {}: {e}",
tmp_path.display(),
final_path.display()
),
)
})?;
if let Ok(dir) = std::fs::File::open(&self.dir) {
let _ = dir.sync_all();
}
Ok(())
}
fn remove(&self, key: u64) -> Result<(), IoError> {
let path = self.final_path(key);
match std::fs::remove_file(&path) {
Ok(()) => Ok(()),
Err(e) if e.kind() == std::io::ErrorKind::NotFound => Ok(()),
Err(e) => Err(IoError::new(
IoErrorKind::Other,
format!("remove {}: {e}", path.display()),
)),
}
}
fn clear(&self) -> Result<(), IoError> {
let entries = match std::fs::read_dir(&self.dir) {
Ok(e) => e,
Err(e) if e.kind() == std::io::ErrorKind::NotFound => return Ok(()),
Err(e) => {
return Err(IoError::new(
IoErrorKind::Other,
format!("read_dir {}: {e}", self.dir.display()),
))
}
};
for entry in entries.flatten() {
let path = entry.path();
if path.extension().and_then(|s| s.to_str()) == Some("bin") {
if let Err(e) = std::fs::remove_file(&path) {
if e.kind() != std::io::ErrorKind::NotFound {
return Err(IoError::new(
IoErrorKind::Other,
format!("remove {}: {e}", path.display()),
));
}
}
}
}
Ok(())
}
fn load(&self, key: u64) -> Result<Option<Vec<u8>>, IoError> {
let path = self.final_path(key);
match std::fs::read(&path) {
Ok(b) => Ok(Some(b)),
Err(e) if e.kind() == std::io::ErrorKind::NotFound => Ok(None),
Err(e) => Err(IoError::new(
IoErrorKind::Other,
format!("read {}: {e}", path.display()),
)),
}
}
fn exists(&self, key: u64) -> bool {
self.final_path(key).exists()
}
}
pub fn encrypt_payload(cipher: Option<&AesCipher>, plaintext: &[u8]) -> Result<Vec<u8>, IoError> {
let Some(cipher) = cipher else {
return Ok(plaintext.to_vec());
};
let mut nonce = [0u8; NONCE_LEN];
crate::encryption::fill_random(&mut nonce)
.map_err(|e| IoError::new(IoErrorKind::Other, format!("fill_random nonce: {e}")))?;
let ct = cipher
.encrypt_page(&nonce, plaintext)
.map_err(|e| IoError::new(IoErrorKind::Other, format!("aes encrypt: {e}")))?;
let mut out = Vec::with_capacity(NONCE_LEN + ct.len());
out.extend_from_slice(&nonce);
out.extend_from_slice(&ct);
Ok(out)
}
pub fn decrypt_payload(
cipher: Option<&AesCipher>,
bytes: &[u8],
) -> Result<Option<Vec<u8>>, IoError> {
let Some(cipher) = cipher else {
return Ok(Some(bytes.to_vec()));
};
if bytes.len() < NONCE_LEN {
return Ok(None);
}
let nonce: [u8; NONCE_LEN] = bytes[..NONCE_LEN].try_into().expect("checked above");
match cipher.decrypt_page(&nonce, &bytes[NONCE_LEN..]) {
Ok(plaintext) => Ok(Some(plaintext)),
Err(_) => Ok(None),
}
}
#[cfg(test)]
#[allow(clippy::items_after_test_module)]
mod tests {
use super::*;
fn entry(key: u64, bytes: usize) -> PersistableEntry {
PersistableEntry {
key,
table_id: 0,
schema_id: 0,
run_generation: 0,
entry_generation: 0,
bytes,
payload_factory: Box::new(|| Some(Vec::new())),
}
}
#[test]
fn frame_round_trip() {
let frame = PersistedFrame {
header: PersistedHeader {
format_version: FRAME_FORMAT_VERSION,
table_id: 7,
schema_id: 9,
run_generation: 11,
cache_key: 13,
entry_generation: 17,
payload_len: 5,
},
payload: b"hello".to_vec(),
};
let encoded = encode_frame(&frame);
let decoded = decode_frame(&encoded).expect("round-trip");
assert_eq!(decoded.header, frame.header);
assert_eq!(decoded.payload, frame.payload);
}
#[test]
fn frame_rejects_bad_crc() {
let frame = PersistedFrame {
header: PersistedHeader {
format_version: FRAME_FORMAT_VERSION,
table_id: 1,
schema_id: 1,
run_generation: 1,
cache_key: 1,
entry_generation: 1,
payload_len: 4,
},
payload: b"test".to_vec(),
};
let mut encoded = encode_frame(&frame);
let last = encoded.len() - 1;
encoded[last] ^= 0xFF;
assert!(decode_frame(&encoded).is_none());
}
#[test]
fn coalescing_replaces_earlier_store() {
let m = LookupMetrics::default();
let w = PersistentResultCacheWriter::new(m, WriterLimits::default());
w.enqueue_store(entry(1, 100));
w.enqueue_store(entry(1, 200));
assert_eq!(w.coalesced_total(), 1);
assert_eq!(w.queue_depth(), 1);
}
#[test]
fn coalescing_does_not_grow_key_count() {
let m = LookupMetrics::default();
let w = PersistentResultCacheWriter::new(
m,
WriterLimits {
max_pending_keys: 1,
max_pending_bytes: 1024 * 1024,
},
);
w.enqueue_store(entry(7, 100));
w.enqueue_store(entry(7, 200));
assert_eq!(w.coalesced_total(), 1);
assert_eq!(w.dropped_total(), 0);
assert_eq!(w.queue_depth(), 1);
}
#[test]
fn remove_supersedes_pending_store() {
let m = LookupMetrics::default();
let w = PersistentResultCacheWriter::new(m, WriterLimits::default());
w.enqueue_store(entry(1, 100));
w.enqueue_remove(1);
w.enqueue_store(entry(1, 100));
assert_eq!(w.dropped_total(), 1);
assert_eq!(w.queue_depth(), 1);
}
#[test]
fn clear_invalidates_every_queued_op() {
let m = LookupMetrics::default();
let w = PersistentResultCacheWriter::new(m, WriterLimits::default());
w.enqueue_store(entry(1, 100));
w.enqueue_store(entry(2, 100));
w.enqueue_clear();
assert_eq!(w.queue_depth(), 1);
let drained = w.drain_one().expect("drain clear");
assert!(matches!(drained.op, PendingCacheOp::Clear));
assert_eq!(w.queue_depth(), 0);
}
#[test]
fn capacity_overflow_drops_new_key() {
let m = LookupMetrics::default();
let w = PersistentResultCacheWriter::new(
m,
WriterLimits {
max_pending_keys: 1,
max_pending_bytes: 1024,
},
);
w.enqueue_store(entry(1, 100));
w.enqueue_store(entry(2, 100));
assert_eq!(w.dropped_total(), 1);
assert_eq!(w.queue_depth(), 1);
}
#[test]
fn byte_accounting_exact() {
let m = LookupMetrics::default();
let w = PersistentResultCacheWriter::new(
m,
WriterLimits {
max_pending_keys: 64,
max_pending_bytes: 1000,
},
);
w.enqueue_store(entry(1, 400));
w.enqueue_store(entry(2, 400));
w.enqueue_store(entry(3, 400));
assert_eq!(w.dropped_total(), 1);
assert_eq!(w.queue_depth(), 2);
w.enqueue_store(entry(1, 200));
w.enqueue_store(entry(3, 400));
assert_eq!(w.dropped_total(), 1);
let m2 = LookupMetrics::default();
let w2 = PersistentResultCacheWriter::new(
m2,
WriterLimits {
max_pending_keys: 64,
max_pending_bytes: 1000,
},
);
w2.enqueue_store(entry(1, 400));
w2.enqueue_store(entry(2, 400));
let _ = w2.drain_one();
w2.enqueue_store(entry(3, 400));
assert_eq!(w2.dropped_total(), 0);
}
#[test]
fn shutdown_drains_pending_ops() {
let m = LookupMetrics::default();
let w = PersistentResultCacheWriter::new(m, WriterLimits::default());
w.enqueue_store(entry(1, 100));
w.enqueue_store(entry(2, 100));
w.shutdown();
assert!(w.drain_one().is_some());
assert!(w.drain_one().is_some());
assert!(w.drain_one().is_none());
}
#[test]
fn drain_all_as_abandoned_clears_queue() {
let m = LookupMetrics::default();
let w = PersistentResultCacheWriter::new(m, WriterLimits::default());
w.enqueue_store(entry(1, 100));
w.enqueue_store(entry(2, 100));
w.enqueue_store(entry(3, 100));
let n = w.drain_all_as_abandoned();
assert_eq!(n, 3);
assert_eq!(w.queue_depth(), 0);
assert_eq!(w.abandoned_total(), 3);
}
fn generation_context(table_id: u64, schema_id: u64, generation: u64) -> PersistContext {
PersistContext {
identity: PersistentCacheIdentity {
table_id,
schema_id,
logical_generation: generation,
},
key: 7,
entry_generation: 1,
}
}
fn assert_exact_generation_identity(cipher: Option<&AesCipher>) {
let expected = PersistentCacheIdentity {
table_id: 1,
schema_id: 1,
logical_generation: 50,
};
let payload = b"cached-rows-payload";
let bytes =
encode_persisted_entry(generation_context(1, 1, 50), payload, cipher).expect("encode");
let decoded = decode_persisted_entry(expected, 7, &bytes, cipher).expect("equal accepted");
assert_eq!(decoded, payload);
let bytes =
encode_persisted_entry(generation_context(1, 1, 49), payload, cipher).expect("encode");
let err = decode_persisted_entry(expected, 7, &bytes, cipher)
.expect_err("older generation rejected");
assert_eq!(
err,
CacheLoadRejection::GenerationMismatch {
expected: 50,
found: 49
}
);
let bytes =
encode_persisted_entry(generation_context(1, 1, 51), payload, cipher).expect("encode");
let err = decode_persisted_entry(expected, 7, &bytes, cipher)
.expect_err("future generation rejected");
assert_eq!(
err,
CacheLoadRejection::GenerationMismatch {
expected: 50,
found: 51
}
);
}
#[test]
fn loader_requires_exact_logical_generation_plaintext() {
assert_exact_generation_identity(None);
}
#[test]
fn loader_requires_exact_logical_generation_encrypted() {
let cipher = AesCipher::new(&[0x42u8; 32]).expect("32-byte key");
assert_exact_generation_identity(Some(&cipher));
}
}
pub fn real_io_final_path(dir: &Path, key: u64) -> PathBuf {
dir.join(format!("{key:016x}.bin"))
}
pub struct WorkerConfig {
pub writer: Arc<PersistentResultCacheWriter>,
pub io: Arc<dyn PersistentCacheIo>,
pub cipher: Option<Arc<AesCipher>>,
pub staleness: Arc<dyn StalenessGuard>,
pub max_staleness_retries: u32,
pub completion: Option<mpsc::Sender<()>>,
}
pub trait StalenessGuard: Send + Sync {
fn is_current(&self, key: u64, key_generation: u64, clear_generation: u64) -> bool;
}
pub struct WriterStalenessGuard {
writer: Arc<PersistentResultCacheWriter>,
}
impl WriterStalenessGuard {
pub fn new(writer: Arc<PersistentResultCacheWriter>) -> Self {
Self { writer }
}
}
impl StalenessGuard for WriterStalenessGuard {
fn is_current(&self, key: u64, key_generation: u64, clear_generation: u64) -> bool {
let guard = self.writer.inner.lock().expect("writer mutex poisoned");
if guard.state.clear_generation != clear_generation {
return false;
}
match guard.state.key_generations.get(&key) {
Some(g) => *g == key_generation,
None => true, }
}
}
pub fn spawn_persistent_cache_worker(config: WorkerConfig) -> std::thread::JoinHandle<()> {
std::thread::spawn(move || run_persistent_cache_worker(config))
}
fn run_persistent_cache_worker(config: WorkerConfig) {
let WorkerConfig {
writer,
io,
cipher,
staleness,
max_staleness_retries,
completion,
} = config;
loop {
let drained = match writer.drain_one() {
Some(d) => d,
None => {
if let Some(tx) = completion {
let _ = tx.send(());
}
return;
}
};
let mut attempts = 0u32;
let is_current = loop {
if staleness.is_current(
drained.key,
drained.key_generation,
drained.clear_generation,
) {
break true;
}
if attempts >= max_staleness_retries {
break false;
}
attempts += 1;
std::thread::yield_now();
};
if !is_current {
if matches!(drained.op, PendingCacheOp::Store(_)) {
writer.record_outcome(DrainOutcome::StoreStale);
}
continue;
}
if !staleness.is_current(
drained.key,
drained.key_generation,
drained.clear_generation,
) {
if matches!(drained.op, PendingCacheOp::Store(_)) {
writer.record_outcome(DrainOutcome::StoreStale);
}
continue;
}
match drained.op {
PendingCacheOp::Store(entry) => {
writer.writes_in_flight.fetch_add(1, Ordering::Relaxed);
let payload = match (entry.payload_factory)() {
Some(p) => p,
None => {
writer.record_outcome(DrainOutcome::StoreErrored);
writer.writes_in_flight.fetch_sub(1, Ordering::Relaxed);
continue;
}
};
let context = PersistContext {
identity: PersistentCacheIdentity {
table_id: entry.table_id,
schema_id: entry.schema_id,
logical_generation: entry.run_generation,
},
key: entry.key,
entry_generation: entry.entry_generation,
};
let bytes = match encode_persisted_entry(context, &payload, cipher.as_deref()) {
Ok(b) => b,
Err(_) => {
writer.record_outcome(DrainOutcome::StoreErrored);
writer.writes_in_flight.fetch_sub(1, Ordering::Relaxed);
continue;
}
};
let outcome = match io.write_atomic(entry.key, &bytes) {
Ok(()) => DrainOutcome::StorePublished,
Err(_) => DrainOutcome::StoreErrored,
};
writer.record_outcome(outcome);
writer.writes_in_flight.fetch_sub(1, Ordering::Relaxed);
}
PendingCacheOp::Remove => {
writer.writes_in_flight.fetch_add(1, Ordering::Relaxed);
let outcome = match io.remove(drained.key) {
Ok(()) => DrainOutcome::RemoveApplied,
Err(_) => DrainOutcome::RemoveErrored,
};
writer.record_outcome(outcome);
writer.writes_in_flight.fetch_sub(1, Ordering::Relaxed);
}
PendingCacheOp::Clear => {
writer.writes_in_flight.fetch_add(1, Ordering::Relaxed);
let outcome = match io.clear() {
Ok(()) => DrainOutcome::ClearApplied,
Err(_) => DrainOutcome::ClearErrored,
};
writer.record_outcome(outcome);
writer.writes_in_flight.fetch_sub(1, Ordering::Relaxed);
}
}
}
}
pub fn recv_completion_with_timeout(rx: &mpsc::Receiver<()>, timeout: Duration) -> bool {
matches!(rx.recv_timeout(timeout), Ok(()))
}