use crate::error::{CacheError, CacheResult};
use crate::traits::CacheStore;
use std::collections::HashSet;
use std::sync::Arc;
use std::time::Duration;
pub struct TaggedCache<C: CacheStore> {
cache: Arc<C>,
}
impl<C: CacheStore> TaggedCache<C> {
fn tag_set_key(tag: &str) -> String {
format!("__armature_tag__:{tag}")
}
fn key_tags_set_key(key: &str) -> String {
format!("__armature_keytags__:{key}")
}
const TAG_INDEX_KEY: &'static str = "__armature_tag_index__";
const RESERVED_KEY_PREFIX: &'static str = "__armature_";
fn validate_caller_key(key: &str) -> CacheResult<()> {
if key.starts_with(Self::RESERVED_KEY_PREFIX) {
Err(CacheError::Config(format!(
"cache key {key:?} is reserved for TaggedCache's internal tag index \
(the {:?} prefix is forbidden for caller-supplied keys)",
Self::RESERVED_KEY_PREFIX
)))
} else {
Ok(())
}
}
fn warn_on_partial_failure(op: &str, err: &CacheError) {
armature_log::warn!(
"TaggedCache::{op} failed partway through a multi-step tag-index update; \
the tag index may now be inconsistent with the cached value or with itself \
(no automatic rollback): {err}"
);
}
pub fn new(cache: Arc<C>) -> Self {
if !cache.supports_atomic_sets() {
armature_log::warn!(
"TaggedCache backing store does not support atomic set operations \
(SADD/SREM/SMEMBERS-equivalent); concurrent set_with_tags/invalidate_tag \
calls against the same tag from different instances can race and lose an \
update. Wrap a backend that overrides CacheStore::set_add/set_remove/\
set_members atomically (e.g. RedisCache) if you need that guarantee."
);
}
Self { cache }
}
pub async fn set_with_tags(
&self,
key: &str,
value: String,
tags: &[&str],
ttl: Option<Duration>,
) -> CacheResult<()> {
Self::validate_caller_key(key)?;
self.cache.set_json(key, value, ttl).await?;
let previous_tags = self.get_tags_for_key(key).await?;
let new_tags: HashSet<String> = tags.iter().map(|t| t.to_string()).collect();
let key_tags_key = Self::key_tags_set_key(key);
for old_tag in &previous_tags {
if !new_tags.contains(old_tag) {
self.cache
.set_remove(&Self::tag_set_key(old_tag), key)
.await
.inspect_err(|e| Self::warn_on_partial_failure("set_with_tags", e))?;
self.cache
.set_remove(&key_tags_key, old_tag)
.await
.inspect_err(|e| Self::warn_on_partial_failure("set_with_tags", e))?;
self.prune_tag_index_if_empty(old_tag).await?;
}
}
for tag in &new_tags {
self.cache
.set_add(&Self::tag_set_key(tag), key)
.await
.inspect_err(|e| Self::warn_on_partial_failure("set_with_tags", e))?;
self.cache
.set_add(&key_tags_key, tag)
.await
.inspect_err(|e| Self::warn_on_partial_failure("set_with_tags", e))?;
self.cache
.set_add(Self::TAG_INDEX_KEY, tag)
.await
.inspect_err(|e| Self::warn_on_partial_failure("set_with_tags", e))?;
}
Ok(())
}
pub async fn get(&self, key: &str) -> CacheResult<Option<String>> {
Self::validate_caller_key(key)?;
self.cache.get_json(key).await
}
pub async fn delete(&self, key: &str) -> CacheResult<()> {
Self::validate_caller_key(key)?;
self.cache.delete(key).await?;
let key_tags_key = Self::key_tags_set_key(key);
let tags = self.cache.set_members(&key_tags_key).await?;
for tag in &tags {
self.cache
.set_remove(&Self::tag_set_key(tag), key)
.await
.inspect_err(|e| Self::warn_on_partial_failure("delete", e))?;
self.prune_tag_index_if_empty(tag).await?;
}
if !tags.is_empty() {
self.cache
.delete(&key_tags_key)
.await
.inspect_err(|e| Self::warn_on_partial_failure("delete", e))?;
}
Ok(())
}
pub async fn invalidate_tag(&self, tag: &str) -> CacheResult<()> {
self.invalidate_tags(&[tag]).await
}
pub async fn invalidate_tags(&self, tags: &[&str]) -> CacheResult<()> {
let mut victims: HashSet<String> = HashSet::new();
for &tag in tags {
let members = self.cache.set_members(&Self::tag_set_key(tag)).await?;
victims.extend(members);
}
if !victims.is_empty() {
let key_refs: Vec<&str> = victims.iter().map(|s| s.as_str()).collect();
self.cache
.delete_many(&key_refs)
.await
.inspect_err(|e| Self::warn_on_partial_failure("invalidate_tags", e))?;
}
let removed: HashSet<&str> = tags.iter().copied().collect();
for key in &victims {
let key_tags_key = Self::key_tags_set_key(key);
let current_tags = self.cache.set_members(&key_tags_key).await?;
for tag in current_tags.iter().filter(|t| removed.contains(t.as_str())) {
self.cache
.set_remove(&key_tags_key, tag)
.await
.inspect_err(|e| Self::warn_on_partial_failure("invalidate_tags", e))?;
}
}
for &tag in tags {
self.cache
.delete(&Self::tag_set_key(tag))
.await
.inspect_err(|e| Self::warn_on_partial_failure("invalidate_tags", e))?;
self.cache
.set_remove(Self::TAG_INDEX_KEY, tag)
.await
.inspect_err(|e| Self::warn_on_partial_failure("invalidate_tags", e))?;
}
Ok(())
}
pub async fn get_keys_by_tag(&self, tag: &str) -> CacheResult<Vec<String>> {
self.cache.set_members(&Self::tag_set_key(tag)).await
}
pub async fn get_tags_for_key(&self, key: &str) -> CacheResult<Vec<String>> {
self.cache.set_members(&Self::key_tags_set_key(key)).await
}
pub async fn list_tags(&self) -> CacheResult<Vec<String>> {
self.cache.set_members(Self::TAG_INDEX_KEY).await
}
async fn prune_tag_index_if_empty(&self, tag: &str) -> CacheResult<()> {
let members = self.cache.set_members(&Self::tag_set_key(tag)).await?;
if members.is_empty() {
self.cache.set_remove(Self::TAG_INDEX_KEY, tag).await?;
}
Ok(())
}
}
impl<C: CacheStore> Clone for TaggedCache<C> {
fn clone(&self) -> Self {
Self {
cache: self.cache.clone(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::error::CacheResult;
use async_trait::async_trait;
use std::collections::HashMap;
use tokio::sync::RwLock;
use std::sync::atomic::{AtomicUsize, Ordering};
#[derive(Clone)]
struct MockCache {
data: Arc<RwLock<HashMap<String, String>>>,
mdel_calls: Arc<AtomicUsize>,
}
impl MockCache {
fn new() -> Self {
Self {
data: Arc::new(RwLock::new(HashMap::new())),
mdel_calls: Arc::new(AtomicUsize::new(0)),
}
}
fn mdel_calls(&self) -> usize {
self.mdel_calls.load(Ordering::Relaxed)
}
}
#[async_trait]
impl CacheStore for MockCache {
async fn get_json(&self, key: &str) -> CacheResult<Option<String>> {
Ok(self.data.read().await.get(key).cloned())
}
async fn set_json(
&self,
key: &str,
value: String,
_ttl: Option<Duration>,
) -> CacheResult<()> {
self.data.write().await.insert(key.to_string(), value);
Ok(())
}
async fn delete(&self, key: &str) -> CacheResult<()> {
self.data.write().await.remove(key);
Ok(())
}
async fn exists(&self, key: &str) -> CacheResult<bool> {
Ok(self.data.read().await.contains_key(key))
}
async fn clear(&self) -> CacheResult<()> {
self.data.write().await.clear();
Ok(())
}
async fn mdel(&self, keys: &[&str]) -> CacheResult<()> {
self.mdel_calls.fetch_add(1, Ordering::Relaxed);
let mut data = self.data.write().await;
for key in keys {
data.remove(*key);
}
Ok(())
}
async fn ttl(&self, _key: &str) -> CacheResult<Option<Duration>> {
Ok(None)
}
async fn expire(&self, _key: &str, _ttl: Duration) -> CacheResult<()> {
Ok(())
}
async fn increment(&self, _key: &str, _delta: i64) -> CacheResult<i64> {
Ok(0)
}
async fn decrement(&self, _key: &str, _delta: i64) -> CacheResult<i64> {
Ok(0)
}
}
#[tokio::test]
async fn test_tagged_cache() {
let cache = Arc::new(MockCache::new());
let tagged = TaggedCache::new(cache);
tagged
.set_with_tags("user:1", "Alice".to_string(), &["users", "active"], None)
.await
.unwrap();
tagged
.set_with_tags("user:2", "Bob".to_string(), &["users"], None)
.await
.unwrap();
let value = tagged.get("user:1").await.unwrap();
assert_eq!(value, Some("Alice".to_string()));
let user_keys = tagged.get_keys_by_tag("users").await.unwrap();
assert_eq!(user_keys.len(), 2);
tagged.invalidate_tag("users").await.unwrap();
let value = tagged.get("user:1").await.unwrap();
assert_eq!(value, None);
}
#[tokio::test]
async fn test_multiple_tags() {
let cache = Arc::new(MockCache::new());
let tagged = TaggedCache::new(cache);
tagged
.set_with_tags("key1", "value1".to_string(), &["tag1", "tag2"], None)
.await
.unwrap();
let tags = tagged.get_tags_for_key("key1").await.unwrap();
assert_eq!(tags.len(), 2);
tagged.invalidate_tag("tag1").await.unwrap();
let value = tagged.get("key1").await.unwrap();
assert_eq!(value, None);
}
#[tokio::test]
async fn test_invalidate_tags_coalesces_into_single_roundtrip() {
let cache = Arc::new(MockCache::new());
let tagged = TaggedCache::new(cache.clone());
tagged
.set_with_tags("k1", "a".to_string(), &["t1"], None)
.await
.unwrap();
tagged
.set_with_tags("k2", "b".to_string(), &["t2"], None)
.await
.unwrap();
tagged
.set_with_tags("shared", "c".to_string(), &["t1", "t3"], None)
.await
.unwrap();
tagged.invalidate_tags(&["t1", "t2", "t3"]).await.unwrap();
assert_eq!(tagged.get("k1").await.unwrap(), None);
assert_eq!(tagged.get("k2").await.unwrap(), None);
assert_eq!(tagged.get("shared").await.unwrap(), None);
assert_eq!(
cache.mdel_calls(),
1,
"invalidate_tags must issue a single coalesced batch delete"
);
assert!(tagged.list_tags().await.unwrap().is_empty());
assert!(tagged.get_tags_for_key("shared").await.unwrap().is_empty());
}
#[tokio::test]
async fn test_tag_index_visible_across_instances_sharing_backend() {
let shared_backend = Arc::new(MockCache::new());
let instance_a = TaggedCache::new(shared_backend.clone());
let instance_b = TaggedCache::new(shared_backend.clone());
instance_a
.set_with_tags("user:1", "Alice".to_string(), &["users"], None)
.await
.unwrap();
let keys = instance_b.get_keys_by_tag("users").await.unwrap();
assert_eq!(keys, vec!["user:1".to_string()]);
instance_b.invalidate_tag("users").await.unwrap();
assert_eq!(instance_a.get("user:1").await.unwrap(), None);
assert!(instance_a.list_tags().await.unwrap().is_empty());
}
#[tokio::test]
async fn test_set_with_tags_rejects_reserved_key_prefix() {
let cache = Arc::new(MockCache::new());
let tagged = TaggedCache::new(cache);
let err = tagged
.set_with_tags(
"__armature_tag__:users",
"corrupt".to_string(),
&["users"],
None,
)
.await
.unwrap_err();
assert!(
matches!(err, CacheError::Config(_)),
"expected CacheError::Config for a reserved-prefixed key, got: {err:?}"
);
assert!(tagged.get_keys_by_tag("users").await.unwrap().is_empty());
}
#[tokio::test]
async fn test_get_and_delete_reject_reserved_key_prefix() {
let cache = Arc::new(MockCache::new());
let tagged = TaggedCache::new(cache);
let get_err = tagged.get("__armature_keytags__:foo").await.unwrap_err();
assert!(matches!(get_err, CacheError::Config(_)));
let delete_err = tagged.delete("__armature_tag_index__").await.unwrap_err();
assert!(matches!(delete_err, CacheError::Config(_)));
}
#[tokio::test]
async fn test_key_containing_but_not_starting_with_reserved_prefix_is_allowed() {
let cache = Arc::new(MockCache::new());
let tagged = TaggedCache::new(cache);
tagged
.set_with_tags(
"user:__armature_tag__:not-a-prefix-collision",
"fine".to_string(),
&["users"],
None,
)
.await
.unwrap();
assert_eq!(
tagged
.get("user:__armature_tag__:not-a-prefix-collision")
.await
.unwrap(),
Some("fine".to_string())
);
}
#[tokio::test]
async fn test_new_does_not_fail_on_non_atomic_backend() {
let cache = Arc::new(MockCache::new());
assert!(!cache.supports_atomic_sets());
let tagged = TaggedCache::new(cache);
tagged
.set_with_tags("k", "v".to_string(), &["t"], None)
.await
.unwrap();
assert_eq!(tagged.get("k").await.unwrap(), Some("v".to_string()));
}
}