use std::any::Any;
use std::sync::Arc;
use std::time::Duration;
use autumn_web::cache::{Cache, RawCacheBytes};
use rocksdb::{ColumnFamily, WriteOptions};
use tokio::runtime::RuntimeFlavor;
use crate::client::RocksDb;
use crate::envelope;
use crate::error::RocksDbError;
use crate::open::{CACHE_CF, Database};
#[derive(Clone)]
pub struct RocksCache {
db: RocksDb,
}
impl RocksCache {
pub fn new(db: RocksDb) -> Result<Self, RocksDbError> {
if db.config().is_read_only() {
return Err(RocksDbError::ReadOnly);
}
if !db.column_families().iter().any(|name| name == CACHE_CF) {
return Err(RocksDbError::UnknownColumnFamily {
name: CACHE_CF.to_owned(),
});
}
Ok(Self { db })
}
fn with_cf<T>(
&self,
work: impl FnOnce(&Database, &ColumnFamily) -> Result<T, rocksdb::Error>,
) -> Option<T> {
let (_slot, db) = self.db.try_slot()?;
let cf = db.cf_handle(CACHE_CF)?;
match blocking(|| work(&db, cf)) {
Ok(value) => Some(value),
Err(err) => {
tracing::warn!(kind = ?err.kind(), "a RocksDB cache call failed");
None
}
}
}
fn fits(&self, key: &str, value: &[u8]) -> bool {
let config = self.db.config();
key.len() <= config.max_key_bytes && value.len() <= config.max_value_bytes
}
fn write_options(&self) -> WriteOptions {
let mut options = WriteOptions::default();
options.set_no_slowdown(true);
options.set_sync(self.db.config().sync_writes);
options
}
fn store(&self, key: &str, bytes: &[u8], ttl: Option<Duration>) {
if !self.fits(key, bytes) {
return;
}
let longest = Duration::from_secs(self.db.config().cache_ttl_secs);
let ttl = ttl.map_or(longest, |ttl| ttl.min(longest));
let value = envelope::encode(bytes, Some(envelope::expiry(envelope::now_ms(), ttl)));
let options = self.write_options();
self.with_cf(|db, cf| db.put_cf_opt(cf, key, &value, &options));
}
fn load(&self, key: &str) -> Option<Vec<u8>> {
if key.len() > self.db.config().max_key_bytes {
return None;
}
let stored = self.with_cf(|db, cf| db.get_cf(cf, key)).flatten()?;
let envelope = envelope::decode(&stored).ok()?;
if envelope.is_expired(envelope::now_ms()) {
return None;
}
Some(envelope.payload.to_vec())
}
}
impl Cache for RocksCache {
fn get_value(&self, key: &str) -> Option<Arc<dyn Any + Send + Sync>> {
let found = self.load(key);
if found.is_some() {
self.db.metrics().cache_hit();
} else {
self.db.metrics().cache_miss();
}
found.map(|bytes| Arc::new(RawCacheBytes(bytes)) as Arc<dyn Any + Send + Sync>)
}
fn insert_value(&self, key: &str, value: Arc<dyn Any + Send + Sync>) {
let bytes = value.downcast_ref::<RawCacheBytes>().map_or_else(
|| {
json(value.downcast_ref::<String>())
.or_else(|| json(value.downcast_ref::<i64>()))
.or_else(|| json(value.downcast_ref::<i32>()))
},
|raw| Some(raw.0.clone()),
);
if let Some(bytes) = bytes {
self.store(key, &bytes, None);
}
}
fn insert_raw_bytes(&self, key: &str, bytes: Vec<u8>, ttl: Option<Duration>) {
self.store(key, &bytes, ttl);
}
fn invalidate(&self, key: &str) {
if key.len() <= self.db.config().max_key_bytes {
let options = self.write_options();
self.with_cf(|db, cf| db.delete_cf_opt(cf, key, &options));
}
}
fn clear(&self) {
let options = self.write_options();
self.with_cf(|db, cf| {
db.delete_range_cf_opt(cf, [].as_slice(), [0xFF].as_slice(), &options)
});
}
}
fn json<T: serde::Serialize>(value: Option<&T>) -> Option<Vec<u8>> {
serde_json::to_vec(value?).ok()
}
fn blocking<T>(work: impl FnOnce() -> T) -> T {
match tokio::runtime::Handle::try_current() {
Ok(handle) if handle.runtime_flavor() == RuntimeFlavor::MultiThread => {
tokio::task::block_in_place(work)
}
_ => work(),
}
}
impl std::fmt::Debug for RocksCache {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RocksCache").finish_non_exhaustive()
}
}
#[cfg(test)]
mod tests;