use crate::error::{OxCacheError, OxCacheResult};
use async_trait::async_trait;
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use std::time::Duration;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct VersionedValue {
pub version: u64,
pub data: Vec<u8>,
}
#[async_trait]
pub trait VersionedStore: Send + Sync {
async fn version(&self, key: &str) -> OxCacheResult<Option<u64>>;
async fn compare_and_swap(
&self,
key: &str,
expect_version: u64,
new_value: Vec<u8>,
ttl: Option<Duration>,
) -> OxCacheResult<Option<u64>>;
async fn get_versioned(&self, key: &str) -> OxCacheResult<Option<VersionedValue>>;
async fn delete(&self, key: &str) -> OxCacheResult<()>;
}
#[derive(Default)]
pub struct MemoryVersionedCache {
entries: Mutex<HashMap<String, VersionedValue>>,
}
impl MemoryVersionedCache {
pub fn new() -> Self {
Self::default()
}
}
#[async_trait]
impl VersionedStore for MemoryVersionedCache {
async fn version(&self, key: &str) -> OxCacheResult<Option<u64>> {
Ok(self
.entries
.lock()
.map_err(|_| OxCacheError::Operation("version store poisoned".to_string()))?
.get(key)
.map(|v| v.version))
}
async fn compare_and_swap(
&self,
key: &str,
expect_version: u64,
new_value: Vec<u8>,
_ttl: Option<Duration>,
) -> OxCacheResult<Option<u64>> {
let mut map = self
.entries
.lock()
.map_err(|_| OxCacheError::Operation("version store poisoned".to_string()))?;
match map.get(key) {
Some(current) if current.version == expect_version => {
let new_version = current.version + 1;
map.insert(
key.to_string(),
VersionedValue {
version: new_version,
data: new_value,
},
);
Ok(Some(new_version))
}
Some(_) => Ok(None),
None if expect_version == 0 => {
map.insert(
key.to_string(),
VersionedValue {
version: 1,
data: new_value,
},
);
Ok(Some(1))
}
None => Ok(None),
}
}
async fn get_versioned(&self, key: &str) -> OxCacheResult<Option<VersionedValue>> {
Ok(self
.entries
.lock()
.map_err(|_| OxCacheError::Operation("version store poisoned".to_string()))?
.get(key)
.cloned())
}
async fn delete(&self, key: &str) -> OxCacheResult<()> {
self.entries
.lock()
.map_err(|_| OxCacheError::Operation("version store poisoned".to_string()))?
.remove(key);
Ok(())
}
}
#[cfg(feature = "redis")]
pub struct RedisVersionedCache {
backend: Arc<crate::backend::RedisBackend>,
}
#[cfg(feature = "redis")]
impl RedisVersionedCache {
pub fn new(backend: Arc<crate::backend::RedisBackend>) -> Self {
Self { backend }
}
fn encode(version: u64, data: &[u8]) -> Vec<u8> {
let mut out = Vec::with_capacity(8 + data.len());
out.extend_from_slice(&version.to_be_bytes());
out.extend_from_slice(data);
out
}
fn decode(raw: &[u8]) -> OxCacheResult<VersionedValue> {
if raw.len() < 8 {
return Err(OxCacheError::Operation(
"versioned envelope too short".to_string(),
));
}
let mut version_bytes = [0u8; 8];
version_bytes.copy_from_slice(&raw[..8]);
Ok(VersionedValue {
version: u64::from_be_bytes(version_bytes),
data: raw[8..].to_vec(),
})
}
}
#[cfg(feature = "redis")]
#[async_trait]
impl VersionedStore for RedisVersionedCache {
async fn version(&self, key: &str) -> OxCacheResult<Option<u64>> {
Ok(self.get_versioned(key).await?.map(|v| v.version))
}
async fn compare_and_swap(
&self,
key: &str,
expect_version: u64,
new_value: Vec<u8>,
ttl: Option<Duration>,
) -> OxCacheResult<Option<u64>> {
let mut conn = self.backend.conn();
redis::cmd("WATCH")
.arg(key)
.query_async::<()>(&mut conn)
.await
.map_err(|e| OxCacheError::Operation(format!("versioned WATCH failed: {e}")))?;
let current: Option<Vec<u8>> = redis::cmd("GET")
.arg(key)
.query_async(&mut conn)
.await
.map_err(|e| OxCacheError::Operation(format!("versioned GET failed: {e}")))?;
let current_version = match current.as_deref().map(Self::decode).transpose()? {
Some(v) => v.version,
None => 0,
};
if current_version != expect_version {
redis::cmd("UNWATCH")
.query_async::<()>(&mut conn)
.await
.ok();
return Ok(None);
}
let new_version = current_version + 1;
let envelope = Self::encode(new_version, &new_value);
redis::cmd("MULTI")
.query_async::<()>(&mut conn)
.await
.map_err(|e| OxCacheError::Operation(format!("versioned MULTI failed: {e}")))?;
let mut set_cmd = redis::cmd("SET");
set_cmd.arg(key).arg(envelope);
if let Some(ttl) = ttl {
set_cmd.arg("EX").arg(ttl.as_secs().max(1));
}
let _: Result<(), _> = set_cmd.query_async(&mut conn).await;
let exec: Result<Option<Vec<redis::Value>>, _> =
redis::cmd("EXEC").query_async(&mut conn).await;
match exec {
Ok(Some(_)) => Ok(Some(new_version)),
Ok(None) => Ok(None),
Err(e) => Err(OxCacheError::Operation(format!(
"versioned EXEC failed: {e}"
))),
}
}
async fn get_versioned(&self, key: &str) -> OxCacheResult<Option<VersionedValue>> {
let mut conn = self.backend.conn();
let raw: Option<Vec<u8>> = redis::cmd("GET")
.arg(key)
.query_async(&mut conn)
.await
.map_err(|e| OxCacheError::Operation(format!("versioned GET failed: {e}")))?;
match raw {
Some(bytes) => Ok(Some(Self::decode(&bytes)?)),
None => Ok(None),
}
}
async fn delete(&self, key: &str) -> OxCacheResult<()> {
let mut conn = self.backend.conn();
let _: i64 = redis::cmd("DEL")
.arg(key)
.query_async(&mut conn)
.await
.map_err(|e| OxCacheError::Operation(format!("versioned DEL failed: {e}")))?;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc as StdArc;
#[tokio::test]
async fn memory_create_with_version_zero() {
let store = MemoryVersionedCache::new();
assert_eq!(store.version("k").await.unwrap(), None);
let v1 = store
.compare_and_swap("k", 0, b"first".to_vec(), None)
.await
.unwrap()
.unwrap();
assert_eq!(v1, 1);
let value = store.get_versioned("k").await.unwrap().unwrap();
assert_eq!(value.version, 1);
assert_eq!(value.data, b"first".to_vec());
}
#[tokio::test]
async fn memory_cas_bumps_version_and_detects_mismatch() {
let store = MemoryVersionedCache::new();
let v1 = store
.compare_and_swap("k", 0, b"a".to_vec(), None)
.await
.unwrap()
.unwrap();
let v2 = store
.compare_and_swap("k", v1, b"b".to_vec(), None)
.await
.unwrap()
.unwrap();
assert_eq!(v2, 2);
assert_eq!(
store
.compare_and_swap("k", v1, b"stale".to_vec(), None)
.await
.unwrap(),
None,
"过期版本必须被拒绝"
);
assert_eq!(
store.get_versioned("k").await.unwrap().unwrap().data,
b"b".to_vec()
);
}
#[tokio::test]
async fn concurrent_writers_no_lost_update() {
let store = StdArc::new(MemoryVersionedCache::new());
store
.compare_and_swap("counter", 0, b"init".to_vec(), None)
.await
.unwrap()
.unwrap();
let a = store.clone();
let b = store.clone();
let (ra, rb) = tokio::join!(
a.compare_and_swap("counter", 1, b"writer-a".to_vec(), None),
b.compare_and_swap("counter", 1, b"writer-b".to_vec(), None),
);
let wins = [ra.unwrap(), rb.unwrap()].into_iter().flatten().count();
assert_eq!(wins, 1, "同版本并发写只能有一方成功");
assert_eq!(store.version("counter").await.unwrap(), Some(2));
}
#[tokio::test]
async fn delete_resets_version() {
let store = MemoryVersionedCache::new();
store
.compare_and_swap("k", 0, b"v".to_vec(), None)
.await
.unwrap()
.unwrap();
store.delete("k").await.unwrap();
assert_eq!(store.version("k").await.unwrap(), None);
assert_eq!(
store
.compare_and_swap("k", 0, b"new".to_vec(), None)
.await
.unwrap(),
Some(1)
);
}
#[tokio::test]
async fn missing_key_rejects_nonzero_expect() {
let store = MemoryVersionedCache::new();
assert_eq!(
store
.compare_and_swap("ghost", 3, b"x".to_vec(), None)
.await
.unwrap(),
None
);
assert_eq!(store.entries.lock().unwrap().len(), 0);
}
#[cfg(feature = "redis")]
#[tokio::test]
#[ignore = "needs live Redis at 127.0.0.1:6379"]
async fn redis_watch_cas_roundtrip() {
use crate::backend::RedisBackend;
let backend = Arc::new(RedisBackend::new("redis://127.0.0.1:6379").await.unwrap());
let store = RedisVersionedCache::new(backend);
let key = format!("oxcache:versioned:{}", uuid::Uuid::new_v4());
let v1 = store
.compare_and_swap(&key, 0, b"a".to_vec(), None)
.await
.unwrap()
.unwrap();
assert_eq!(v1, 1);
let v2 = store
.compare_and_swap(&key, v1, b"b".to_vec(), None)
.await
.unwrap()
.unwrap();
assert_eq!(v2, 2);
assert_eq!(
store
.compare_and_swap(&key, v1, b"stale".to_vec(), None)
.await
.unwrap(),
None
);
store.delete(&key).await.unwrap();
}
}