use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::time::Duration;
use async_trait::async_trait;
use dashmap::DashMap;
use tokio::time::Instant;
use crate::connector::CacheConnectorConfig;
use crate::errors::OrionError;
#[async_trait]
pub trait CacheBackend: Send + Sync {
async fn get(&self, key: &str) -> Result<Option<String>, OrionError>;
async fn set(&self, key: &str, value: &str) -> Result<(), OrionError>;
async fn set_ex(&self, key: &str, value: &str, ttl_secs: u64) -> Result<(), OrionError>;
async fn remove(&self, key: &str) -> Result<(), OrionError>;
async fn claim_dedup_key(
&self,
key: &str,
owner: &str,
window_secs: u64,
) -> Result<Option<String>, OrionError>;
}
struct MemoryEntry {
value: String,
expires_at: Option<Instant>,
last_access: AtomicU64,
}
const EVICTION_BATCH_DIVISOR: usize = 10;
pub struct MemoryCacheBackend {
entries: DashMap<String, MemoryEntry>,
max_entries: usize,
clock: AtomicU64,
evicting: AtomicBool,
}
impl MemoryCacheBackend {
pub fn new(cleanup_interval_secs: u64, max_entries: usize) -> Arc<Self> {
let store = Arc::new(Self {
entries: DashMap::new(),
max_entries,
clock: AtomicU64::new(0),
evicting: AtomicBool::new(false),
});
let weak = Arc::downgrade(&store);
tokio::spawn(async move {
let interval = Duration::from_secs(cleanup_interval_secs.max(1));
loop {
tokio::time::sleep(interval).await;
let Some(store) = weak.upgrade() else {
break;
};
store.purge_expired();
}
});
store
}
fn purge_expired(&self) {
let now = Instant::now();
self.entries
.retain(|_, entry| entry.expires_at.is_none_or(|exp| exp > now));
}
fn tick(&self) -> u64 {
self.clock.fetch_add(1, Ordering::Relaxed)
}
fn new_entry(&self, value: &str, expires_at: Option<Instant>) -> MemoryEntry {
MemoryEntry {
value: value.to_string(),
expires_at,
last_access: AtomicU64::new(self.tick()),
}
}
fn enforce_bound(&self) {
if self.max_entries == 0 || self.entries.len() <= self.max_entries {
return;
}
if self
.evicting
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
.is_err()
{
return;
}
self.purge_expired();
let len = self.entries.len();
if len > self.max_entries {
let batch = (self.max_entries / EVICTION_BATCH_DIVISOR).max(1);
let target = (len - self.max_entries + batch).min(len);
let mut ticks: Vec<u64> = self
.entries
.iter()
.map(|e| e.last_access.load(Ordering::Relaxed))
.collect();
if target >= 1 && target <= ticks.len() {
let (_, nth, _) = ticks.select_nth_unstable(target - 1);
let threshold = *nth;
self.entries
.retain(|_, entry| entry.last_access.load(Ordering::Relaxed) > threshold);
tracing::warn!(
max_entries = self.max_entries,
evicted = len.saturating_sub(self.entries.len()),
"In-memory cache at capacity, evicted least-recently-used entries"
);
}
}
self.evicting.store(false, Ordering::Release);
}
}
#[async_trait]
impl CacheBackend for MemoryCacheBackend {
async fn get(&self, key: &str) -> Result<Option<String>, OrionError> {
let Some(entry) = self.entries.get(key) else {
return Ok(None);
};
if let Some(exp) = entry.expires_at
&& Instant::now() >= exp
{
drop(entry); self.entries.remove(key);
return Ok(None);
}
entry.last_access.store(self.tick(), Ordering::Relaxed);
Ok(Some(entry.value.clone()))
}
async fn set(&self, key: &str, value: &str) -> Result<(), OrionError> {
self.entries
.insert(key.to_string(), self.new_entry(value, None));
self.enforce_bound();
Ok(())
}
async fn set_ex(&self, key: &str, value: &str, ttl_secs: u64) -> Result<(), OrionError> {
self.entries.insert(
key.to_string(),
self.new_entry(value, Some(Instant::now() + Duration::from_secs(ttl_secs))),
);
self.enforce_bound();
Ok(())
}
async fn remove(&self, key: &str) -> Result<(), OrionError> {
self.entries.remove(key);
Ok(())
}
async fn claim_dedup_key(
&self,
key: &str,
owner: &str,
window_secs: u64,
) -> Result<Option<String>, OrionError> {
use dashmap::mapref::entry::Entry;
let now = Instant::now();
let expires_at = now + Duration::from_secs(window_secs);
let holder = match self.entries.entry(key.to_string()) {
Entry::Vacant(vacant) => {
vacant.insert(self.new_entry(owner, Some(expires_at)));
None }
Entry::Occupied(mut occupied) => {
if let Some(exp) = occupied.get().expires_at
&& now >= exp
{
occupied.insert(self.new_entry(owner, Some(expires_at)));
None
} else {
Some(occupied.get().value.clone())
}
}
};
if holder.is_none() {
self.enforce_bound();
}
Ok(holder)
}
}
pub struct RedisCacheBackend {
conn: redis::aio::ConnectionManager,
}
impl RedisCacheBackend {
pub fn new(conn: redis::aio::ConnectionManager) -> Self {
Self { conn }
}
}
#[async_trait]
impl CacheBackend for RedisCacheBackend {
async fn get(&self, key: &str) -> Result<Option<String>, OrionError> {
use redis::AsyncCommands;
let mut conn = self.conn.clone();
conn.get(key).await.map_err(|e| OrionError::Internal {
context: format!("Redis GET failed for key '{key}'"),
source: Some(Box::new(e)),
})
}
async fn set(&self, key: &str, value: &str) -> Result<(), OrionError> {
use redis::AsyncCommands;
let mut conn = self.conn.clone();
conn.set::<_, _, ()>(key, value)
.await
.map_err(|e| OrionError::Internal {
context: format!("Redis SET failed for key '{key}'"),
source: Some(Box::new(e)),
})
}
async fn set_ex(&self, key: &str, value: &str, ttl_secs: u64) -> Result<(), OrionError> {
use redis::AsyncCommands;
let mut conn = self.conn.clone();
conn.set_ex::<_, _, ()>(key, value, ttl_secs)
.await
.map_err(|e| OrionError::Internal {
context: format!("Redis SETEX failed for key '{key}'"),
source: Some(Box::new(e)),
})
}
async fn remove(&self, key: &str) -> Result<(), OrionError> {
use redis::AsyncCommands;
let mut conn = self.conn.clone();
conn.del::<_, ()>(key)
.await
.map_err(|e| OrionError::Internal {
context: format!("Redis DEL failed for key '{key}'"),
source: Some(Box::new(e)),
})
}
async fn claim_dedup_key(
&self,
key: &str,
owner: &str,
window_secs: u64,
) -> Result<Option<String>, OrionError> {
use redis::AsyncCommands;
let mut conn = self.conn.clone();
for _ in 0..2 {
let claimed: Option<String> = redis::cmd("SET")
.arg(key)
.arg(owner)
.arg("NX")
.arg("EX")
.arg(window_secs)
.query_async(&mut conn)
.await
.map_err(|e| OrionError::Internal {
context: format!("Redis SET NX EX failed for key '{key}'"),
source: Some(Box::new(e)),
})?;
if claimed.is_some() {
return Ok(None);
}
let holder: Option<String> = conn.get(key).await.map_err(|e| OrionError::Internal {
context: format!("Redis GET failed for key '{key}'"),
source: Some(Box::new(e)),
})?;
if let Some(holder) = holder {
return Ok(Some(holder));
}
}
tracing::warn!(
key = %key,
"Deduplication key expired between claim attempts twice; treating the message as new"
);
Ok(None)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum CachePurpose {
Workflow,
Dedup,
ResponseCache,
}
impl CachePurpose {
fn as_str(self) -> &'static str {
match self {
CachePurpose::Workflow => "workflow",
CachePurpose::Dedup => "dedup",
CachePurpose::ResponseCache => "response_cache",
}
}
}
impl std::fmt::Display for CachePurpose {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(match self {
CachePurpose::Workflow => "workflow cache",
CachePurpose::Dedup => "deduplication",
CachePurpose::ResponseCache => "response cache",
})
}
}
pub struct CachePool {
memory: DashMap<String, Arc<MemoryCacheBackend>>,
cleanup_interval_secs: u64,
max_memory_cache_entries: usize,
redis: Arc<super::redis_pool::RedisPoolCache>,
}
impl CachePool {
pub fn new(
max_redis_pool_entries: usize,
cleanup_interval_secs: u64,
max_memory_cache_entries: usize,
) -> Self {
Self {
memory: DashMap::new(),
cleanup_interval_secs,
max_memory_cache_entries,
redis: Arc::new(super::redis_pool::RedisPoolCache::new(
max_redis_pool_entries,
)),
}
}
fn memory_namespace(&self, namespace: &str) -> Arc<dyn CacheBackend> {
if let Some(existing) = self.memory.get(namespace) {
return existing.value().clone() as Arc<dyn CacheBackend>;
}
self.memory
.entry(namespace.to_string())
.or_insert_with(|| {
MemoryCacheBackend::new(self.cleanup_interval_secs, self.max_memory_cache_entries)
})
.clone() as Arc<dyn CacheBackend>
}
pub async fn get_backend(
&self,
purpose: CachePurpose,
connector_name: &str,
config: &CacheConnectorConfig,
) -> Result<Arc<dyn CacheBackend>, OrionError> {
match config.backend.as_str() {
"memory" => {
Ok(self.memory_namespace(&format!("{}:{connector_name}", purpose.as_str())))
}
"redis" => {
let conn = self.redis.get_conn(connector_name, config).await?;
Ok(Arc::new(RedisCacheBackend::new(conn)))
}
other => Err(OrionError::validation(format!(
"Unknown cache backend '{other}'. Must be 'redis' or 'memory'"
))),
}
}
pub fn default_memory(&self, purpose: CachePurpose) -> Arc<dyn CacheBackend> {
self.memory_namespace(purpose.as_str())
}
pub async fn evict_pool(&self, connector_name: &str) {
self.redis.evict(connector_name).await;
}
pub async fn evict_all_pools(&self) {
self.redis.evict_all().await;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_memory_get_set() {
let backend = MemoryCacheBackend::new(60, 0);
assert!(backend.get("k1").await.expect("test").is_none());
backend.set("k1", "v1").await.expect("test");
assert_eq!(
backend.get("k1").await.expect("test"),
Some("v1".to_string())
);
}
#[tokio::test(start_paused = true)]
async fn test_memory_set_ex_expires() {
let backend = MemoryCacheBackend::new(60, 0);
backend.set_ex("k1", "v1", 1).await.expect("test");
assert_eq!(
backend.get("k1").await.expect("test"),
Some("v1".to_string())
);
tokio::time::advance(Duration::from_secs(2)).await;
assert!(backend.get("k1").await.expect("test").is_none());
}
#[tokio::test]
async fn test_memory_claim_dedup_key_new() {
let backend = MemoryCacheBackend::new(60, 0);
assert_eq!(
backend
.claim_dedup_key("dedup-1", "owner-a", 300)
.await
.expect("test"),
None
);
}
#[tokio::test]
async fn test_memory_claim_dedup_key_reports_the_holder() {
let backend = MemoryCacheBackend::new(60, 0);
assert_eq!(
backend
.claim_dedup_key("dedup-1", "owner-a", 300)
.await
.expect("test"),
None
);
assert_eq!(
backend
.claim_dedup_key("dedup-1", "owner-b", 300)
.await
.expect("test"),
Some("owner-a".to_string())
);
}
#[tokio::test]
async fn test_memory_remove_frees_a_claim() {
let backend = MemoryCacheBackend::new(60, 0);
backend
.claim_dedup_key("dedup-1", "owner-a", 300)
.await
.expect("test");
backend.remove("dedup-1").await.expect("test");
assert_eq!(
backend
.claim_dedup_key("dedup-1", "owner-b", 300)
.await
.expect("test"),
None,
"a released key must be claimable again"
);
}
#[tokio::test(start_paused = true)]
async fn test_memory_claim_dedup_key_expired() {
let backend = MemoryCacheBackend::new(60, 0);
assert_eq!(
backend
.claim_dedup_key("k", "owner-a", 1)
.await
.expect("test"),
None
);
tokio::time::advance(Duration::from_secs(2)).await;
assert_eq!(
backend
.claim_dedup_key("k", "owner-b", 1)
.await
.expect("test"),
None
);
}
#[tokio::test(start_paused = true)]
async fn test_memory_purge_expired() {
let backend = MemoryCacheBackend::new(60, 0);
backend.set_ex("keep", "val", 3600).await.expect("test");
backend.set_ex("expire", "val", 1).await.expect("test");
tokio::time::advance(Duration::from_secs(2)).await;
backend.purge_expired();
assert!(backend.get("keep").await.expect("test").is_some());
assert!(backend.get("expire").await.expect("test").is_none());
}
#[tokio::test]
async fn test_memory_set_overwrites() {
let backend = MemoryCacheBackend::new(60, 0);
backend.set("k", "v1").await.expect("test");
backend.set("k", "v2").await.expect("test");
assert_eq!(
backend.get("k").await.expect("test"),
Some("v2".to_string())
);
}
#[tokio::test]
async fn test_memory_set_without_ttl_is_bounded() {
let backend = MemoryCacheBackend::new(60, 100);
for i in 0..10_000 {
backend.set(&format!("k{i}"), "v").await.expect("test");
}
assert!(
backend.entries.len() <= 100,
"unbounded growth: {} entries",
backend.entries.len()
);
}
#[tokio::test]
async fn test_memory_set_ex_is_bounded() {
let backend = MemoryCacheBackend::new(60, 50);
for i in 0..2_000 {
backend
.set_ex(&format!("k{i}"), "v", 3600)
.await
.expect("test");
}
assert!(backend.entries.len() <= 50);
}
#[tokio::test]
async fn test_memory_claim_dedup_key_is_bounded() {
let backend = MemoryCacheBackend::new(60, 64);
for i in 0..5_000 {
backend
.claim_dedup_key(&format!("dedup-{i}"), "owner", 3600)
.await
.expect("test");
}
assert!(backend.entries.len() <= 64);
}
#[tokio::test]
async fn test_memory_lru_evicts_coldest_first() {
let backend = MemoryCacheBackend::new(60, 10);
for i in 0..10 {
backend.set(&format!("k{i}"), "v").await.expect("test");
}
assert!(backend.get("k0").await.expect("test").is_some());
backend.set("overflow", "v").await.expect("test");
assert!(
backend.get("k0").await.expect("test").is_some(),
"a key read since insertion must outlive colder ones"
);
assert!(
backend.get("k1").await.expect("test").is_none(),
"the coldest key must be the one evicted"
);
assert!(backend.get("overflow").await.expect("test").is_some());
}
#[tokio::test(start_paused = true)]
async fn test_memory_expired_entries_evicted_before_live_ones() {
let backend = MemoryCacheBackend::new(3600, 4);
for i in 0..4 {
backend
.set_ex(&format!("gone{i}"), "v", 1)
.await
.expect("test");
}
tokio::time::advance(Duration::from_secs(2)).await;
backend.set("live", "v").await.expect("test");
assert_eq!(
backend.get("live").await.expect("test"),
Some("v".to_string()),
"a fresh entry must not be evicted while expired ones remain"
);
assert!(backend.entries.len() <= 4);
}
#[tokio::test]
async fn test_memory_zero_max_entries_is_unbounded() {
let backend = MemoryCacheBackend::new(3600, 0);
for i in 0..500 {
backend.set(&format!("k{i}"), "v").await.expect("test");
}
assert_eq!(backend.entries.len(), 500);
}
fn memory_connector() -> CacheConnectorConfig {
CacheConnectorConfig {
backend: "memory".to_string(),
url: None,
allow_private_urls: false,
operations: Default::default(),
}
}
#[tokio::test]
async fn workflow_writes_cannot_poison_the_dedup_store() {
let pool = CachePool::new(4, 60, 1000);
let workflow = pool
.get_backend(CachePurpose::Workflow, "wf-cache", &memory_connector())
.await
.expect("test");
workflow
.set("dedup:orders:token-1", "\"1\"")
.await
.expect("test");
let dedup = pool.default_memory(CachePurpose::Dedup);
assert_eq!(
dedup
.claim_dedup_key("dedup:orders:token-1", "owner", 300)
.await
.expect("test"),
None,
"a workflow-written key must not read as a duplicate in the dedup store"
);
}
#[tokio::test]
async fn workflow_writes_cannot_forge_a_cached_response() {
let pool = CachePool::new(4, 60, 1000);
let workflow = pool
.get_backend(CachePurpose::Workflow, "wf-cache", &memory_connector())
.await
.expect("test");
workflow
.set("cache:orders:00000000deadbeef", r#"{"forged":true}"#)
.await
.expect("test");
let response_cache = pool.default_memory(CachePurpose::ResponseCache);
assert!(
response_cache
.get("cache:orders:00000000deadbeef")
.await
.expect("test")
.is_none(),
"a workflow-written key must not surface as a cached response"
);
}
#[tokio::test]
async fn two_memory_connectors_have_distinct_keyspaces() {
let pool = CachePool::new(4, 60, 1000);
let a = pool
.get_backend(CachePurpose::Workflow, "conn-a", &memory_connector())
.await
.expect("test");
let b = pool
.get_backend(CachePurpose::Workflow, "conn-b", &memory_connector())
.await
.expect("test");
a.set("shared-key", "\"from-a\"").await.expect("test");
assert!(
b.get("shared-key").await.expect("test").is_none(),
"connector 'conn-b' must not see keys written through 'conn-a'"
);
let a2 = pool
.get_backend(CachePurpose::Workflow, "conn-a", &memory_connector())
.await
.expect("test");
assert_eq!(
a2.get("shared-key").await.expect("test"),
Some("\"from-a\"".to_string())
);
}
#[tokio::test]
async fn lru_budgets_are_per_namespace() {
let pool = CachePool::new(4, 60, 10);
let dedup = pool.default_memory(CachePurpose::Dedup);
assert_eq!(
dedup
.claim_dedup_key("dedup:ch:tok", "owner", 300)
.await
.expect("test"),
None
);
let workflow = pool
.get_backend(CachePurpose::Workflow, "hot", &memory_connector())
.await
.expect("test");
for i in 0..1_000 {
workflow.set(&format!("k{i}"), "v").await.expect("test");
}
assert_eq!(
dedup
.claim_dedup_key("dedup:ch:tok", "owner", 300)
.await
.expect("test"),
Some("owner".to_string()),
"a hot workflow cache must not evict dedup entries"
);
}
#[tokio::test]
async fn test_memory_bound_holds_under_concurrent_writers() {
let backend = MemoryCacheBackend::new(3600, 100);
let mut handles = Vec::new();
for w in 0..8 {
let backend = backend.clone();
handles.push(tokio::spawn(async move {
for i in 0..500 {
backend.set(&format!("w{w}-k{i}"), "v").await.expect("test");
}
}));
}
for h in handles {
h.await.expect("test");
}
assert!(
backend.entries.len() <= 400,
"bound not holding: {} entries",
backend.entries.len()
);
}
}