use crate::error::{CacheError, CacheResult};
use crate::traits::CacheStore;
use futures::future::join_all;
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}"
);
}
fn first_error(op: &str, results: Vec<CacheResult<()>>) -> CacheResult<()> {
for result in results {
result.inspect_err(|e| Self::warn_on_partial_failure(op, e))?;
}
Ok(())
}
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);
let stale_tags: Vec<&str> = previous_tags
.iter()
.filter(|t| !new_tags.contains(*t))
.map(|t| t.as_str())
.collect();
if !stale_tags.is_empty() {
let stale_set_keys: Vec<String> =
stale_tags.iter().map(|t| Self::tag_set_key(t)).collect();
let removals = join_all(
stale_set_keys
.iter()
.map(|set_key| self.cache.set_remove(set_key, key)),
)
.await;
Self::first_error("set_with_tags", removals)?;
self.cache
.set_remove_many(&key_tags_key, &stale_tags)
.await
.inspect_err(|e| Self::warn_on_partial_failure("set_with_tags", e))?;
self.prune_tag_index_where_empty(&stale_tags).await?;
}
let added_tags: Vec<&str> = new_tags.iter().map(|t| t.as_str()).collect();
if !added_tags.is_empty() {
let added_set_keys: Vec<String> =
added_tags.iter().map(|t| Self::tag_set_key(t)).collect();
let additions = join_all(
added_set_keys
.iter()
.map(|set_key| self.cache.set_add(set_key, key)),
)
.await;
Self::first_error("set_with_tags", additions)?;
self.cache
.set_add_many(&key_tags_key, &added_tags)
.await
.inspect_err(|e| Self::warn_on_partial_failure("set_with_tags", e))?;
self.cache
.set_add_many(Self::TAG_INDEX_KEY, &added_tags)
.await
.inspect_err(|e| Self::warn_on_partial_failure("set_with_tags", e))?;
if let Some(ttl) = ttl
&& let Err(e) = self.cache.expire(&key_tags_key, ttl).await
{
armature_log::warn!(
"TaggedCache::set_with_tags could not mirror the value TTL onto the \
reverse tag index for key {key:?}; the index entry may outlive the \
value (stale members are still reconciled on read): {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?;
if !tags.is_empty() {
let tag_set_keys: Vec<String> = tags.iter().map(|t| Self::tag_set_key(t)).collect();
let removals = join_all(
tag_set_keys
.iter()
.map(|set_key| self.cache.set_remove(set_key, key)),
)
.await;
Self::first_error("delete", removals)?;
let tag_refs: Vec<&str> = tags.iter().map(|t| t.as_str()).collect();
self.prune_tag_index_where_empty(&tag_refs).await?;
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 tag_set_keys: Vec<String> = tags.iter().map(|t| Self::tag_set_key(t)).collect();
let member_lists = join_all(
tag_set_keys
.iter()
.map(|set_key| self.cache.set_members(set_key)),
)
.await;
let mut victims: HashSet<String> = HashSet::new();
for members in member_lists {
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();
let reverse_updates = join_all(victims.iter().map(|key| {
let removed = &removed;
async move {
let key_tags_key = Self::key_tags_set_key(key.as_str());
let current_tags = self.cache.set_members(&key_tags_key).await?;
let doomed: Vec<&str> = current_tags
.iter()
.map(|t| t.as_str())
.filter(|t| removed.contains(t))
.collect();
if doomed.is_empty() {
return Ok(());
}
self.cache.set_remove_many(&key_tags_key, &doomed).await
}
}))
.await;
Self::first_error("invalidate_tags", reverse_updates)?;
let tag_set_key_refs: Vec<&str> = tag_set_keys.iter().map(|k| k.as_str()).collect();
if !tag_set_key_refs.is_empty() {
self.cache
.delete_many(&tag_set_key_refs)
.await
.inspect_err(|e| Self::warn_on_partial_failure("invalidate_tags", e))?;
self.cache
.set_remove_many(Self::TAG_INDEX_KEY, tags)
.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>> {
let tag_key = Self::tag_set_key(tag);
let members = self.cache.set_members(&tag_key).await?;
if members.is_empty() {
return Ok(Vec::new());
}
let member_refs: Vec<&str> = members.iter().map(|m| m.as_str()).collect();
let present = self.cache.exists_many(&member_refs).await?;
let mut live: Vec<String> = Vec::with_capacity(members.len());
let mut stale: Vec<&str> = Vec::new();
for (member, exists) in members.iter().zip(present) {
if exists {
live.push(member.clone());
} else {
stale.push(member.as_str());
}
}
if !stale.is_empty() {
match self.cache.set_remove_many(&tag_key, &stale).await {
Ok(()) => {
if live.is_empty()
&& let Err(e) = self.prune_tag_index_where_empty(&[tag]).await
{
armature_log::warn!(
"TaggedCache::get_keys_by_tag could not prune the now-empty tag \
{tag:?} from the tag index: {e}"
);
}
}
Err(e) => {
let stale_count = stale.len();
armature_log::warn!(
"TaggedCache::get_keys_by_tag could not prune {stale_count} expired \
member(s) from tag {tag:?}; they are excluded from this result and \
the prune will be retried on the next read: {e}"
);
}
}
}
Ok(live)
}
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_where_empty(&self, tags: &[&str]) -> CacheResult<()> {
if tags.is_empty() {
return Ok(());
}
let tag_set_keys: Vec<String> = tags.iter().map(|t| Self::tag_set_key(t)).collect();
let member_lists = join_all(
tag_set_keys
.iter()
.map(|set_key| self.cache.set_members(set_key)),
)
.await;
let mut empty: Vec<&str> = Vec::new();
for (tag, members) in tags.iter().zip(member_lists) {
if members?.is_empty() {
empty.push(*tag);
}
}
if !empty.is_empty() {
self.cache
.set_remove_many(Self::TAG_INDEX_KEY, &empty)
.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 crate::tiered::InMemoryCache;
#[derive(Clone)]
struct MockCache {
data: Arc<RwLock<HashMap<String, String>>>,
mdel_batches: Arc<RwLock<Vec<Vec<String>>>>,
}
impl MockCache {
fn new() -> Self {
Self {
data: Arc::new(RwLock::new(HashMap::new())),
mdel_batches: Arc::new(RwLock::new(Vec::new())),
}
}
async fn mdel_batches(&self) -> Vec<Vec<String>> {
self.mdel_batches
.read()
.await
.iter()
.map(|batch| {
let mut batch = batch.clone();
batch.sort();
batch
})
.collect()
}
}
#[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_batches
.write()
.await
.push(keys.iter().map(|k| k.to_string()).collect());
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);
let batches = cache.mdel_batches().await;
assert_eq!(
batches.len(),
2,
"invalidate_tags must coalesce its deletes, got batches: {batches:?}"
);
assert_eq!(batches[0], vec!["k1", "k2", "shared"]);
assert_eq!(
batches[1],
vec![
"__armature_tag__:t1",
"__armature_tag__:t2",
"__armature_tag__:t3"
]
);
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()));
}
#[tokio::test(start_paused = true)]
async fn test_set_with_tags_mirrors_ttl_onto_reverse_index() {
let cache = Arc::new(InMemoryCache::new());
let tagged = TaggedCache::new(cache.clone());
tagged
.set_with_tags(
"user:1",
"Alice".to_string(),
&["users"],
Some(Duration::from_secs(60)),
)
.await
.unwrap();
let key_tags_key = TaggedCache::<InMemoryCache>::key_tags_set_key("user:1");
let index_ttl = cache
.ttl(&key_tags_key)
.await
.unwrap()
.expect("the reverse tag index must inherit the value's TTL");
assert!(index_ttl > Duration::from_secs(0));
assert!(index_ttl <= Duration::from_secs(60));
tokio::time::advance(Duration::from_secs(61)).await;
assert_eq!(tagged.get("user:1").await.unwrap(), None);
assert!(tagged.get_tags_for_key("user:1").await.unwrap().is_empty());
}
#[tokio::test(start_paused = true)]
async fn test_set_with_tags_without_ttl_leaves_index_unexpiring() {
let cache = Arc::new(InMemoryCache::new());
let tagged = TaggedCache::new(cache.clone());
tagged
.set_with_tags("user:1", "Alice".to_string(), &["users"], None)
.await
.unwrap();
let key_tags_key = TaggedCache::<InMemoryCache>::key_tags_set_key("user:1");
assert_eq!(cache.ttl(&key_tags_key).await.unwrap(), None);
tokio::time::advance(Duration::from_secs(3600)).await;
assert_eq!(
tagged.get_keys_by_tag("users").await.unwrap(),
vec!["user:1".to_string()],
"a key with no TTL must stay a live member of its tags"
);
}
#[tokio::test(start_paused = true)]
async fn test_get_keys_by_tag_reconciles_expired_members() {
let cache = Arc::new(InMemoryCache::new());
let tagged = TaggedCache::new(cache.clone());
tagged
.set_with_tags(
"short",
"gone-soon".to_string(),
&["users"],
Some(Duration::from_secs(1)),
)
.await
.unwrap();
tagged
.set_with_tags("forever", "stays".to_string(), &["users"], None)
.await
.unwrap();
let mut keys = tagged.get_keys_by_tag("users").await.unwrap();
keys.sort();
assert_eq!(keys, vec!["forever".to_string(), "short".to_string()]);
tokio::time::advance(Duration::from_secs(2)).await;
assert_eq!(
tagged.get_keys_by_tag("users").await.unwrap(),
vec!["forever".to_string()]
);
let tag_key = TaggedCache::<InMemoryCache>::tag_set_key("users");
assert_eq!(
cache.set_members(&tag_key).await.unwrap(),
vec!["forever".to_string()]
);
}
#[tokio::test(start_paused = true)]
async fn test_reconciliation_prunes_emptied_tag_from_index() {
let cache = Arc::new(InMemoryCache::new());
let tagged = TaggedCache::new(cache.clone());
tagged
.set_with_tags(
"short",
"gone-soon".to_string(),
&["ephemeral"],
Some(Duration::from_secs(1)),
)
.await
.unwrap();
assert_eq!(
tagged.list_tags().await.unwrap(),
vec!["ephemeral".to_string()]
);
tokio::time::advance(Duration::from_secs(2)).await;
assert!(
tagged
.get_keys_by_tag("ephemeral")
.await
.unwrap()
.is_empty()
);
assert!(
tagged.list_tags().await.unwrap().is_empty(),
"an emptied tag must be pruned from the global tag index"
);
}
#[tokio::test]
async fn test_retagging_removes_key_from_dropped_tags() {
let cache = Arc::new(MockCache::new());
let tagged = TaggedCache::new(cache);
tagged
.set_with_tags("k", "v1".to_string(), &["a", "b", "c"], None)
.await
.unwrap();
tagged
.set_with_tags("k", "v2".to_string(), &["c", "d"], None)
.await
.unwrap();
assert!(tagged.get_keys_by_tag("a").await.unwrap().is_empty());
assert!(tagged.get_keys_by_tag("b").await.unwrap().is_empty());
assert_eq!(
tagged.get_keys_by_tag("c").await.unwrap(),
vec!["k".to_string()]
);
assert_eq!(
tagged.get_keys_by_tag("d").await.unwrap(),
vec!["k".to_string()]
);
let mut tags = tagged.get_tags_for_key("k").await.unwrap();
tags.sort();
assert_eq!(tags, vec!["c".to_string(), "d".to_string()]);
let mut listed = tagged.list_tags().await.unwrap();
listed.sort();
assert_eq!(listed, vec!["c".to_string(), "d".to_string()]);
}
}