use crate::backend::interface::{BackendKind, CacheSetItem};
use crate::backend::{CacheBackend, CacheConnector, CacheReader, CacheWriter};
use crate::error::OxCacheResult;
use hmac::{Hmac, KeyInit, Mac};
use sha2::Sha256;
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
type HmacSha256 = Hmac<Sha256>;
pub const HMAC_ENVELOPE_VERSION: u8 = 1;
const TAG_SIZE: usize = 32;
#[derive(Clone)]
pub struct HmacSigner {
key: Arc<[u8; TAG_SIZE]>,
}
impl std::fmt::Debug for HmacSigner {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("HmacSigner")
.field("key", &"<redacted 32-byte key>")
.finish()
}
}
impl HmacSigner {
pub fn new(key: [u8; TAG_SIZE]) -> Self {
Self { key: Arc::new(key) }
}
pub fn from_slice(key: &[u8]) -> crate::error::OxCacheResult<Self> {
let arr: [u8; TAG_SIZE] = key.try_into().map_err(|_| {
crate::error::OxCacheError::InvalidInput(format!(
"hmac signer key must be exactly {TAG_SIZE} bytes, got {}",
key.len()
))
})?;
Ok(Self::new(arr))
}
pub fn sign(&self, message: &[u8]) -> [u8; TAG_SIZE] {
let mut mac = HmacSha256::new_from_slice(self.key.as_ref())
.expect("HMAC accepts any key length; 32-byte key is valid");
mac.update(message);
mac.finalize().into_bytes().into()
}
pub fn verify(&self, message: &[u8], tag: &[u8]) -> bool {
let mut mac = HmacSha256::new_from_slice(self.key.as_ref())
.expect("HMAC accepts any key length; 32-byte key is valid");
mac.update(message);
mac.verify_slice(tag).is_ok()
}
}
pub struct IntegrityBackend {
inner: Arc<dyn CacheBackend>,
signer: HmacSigner,
}
impl std::fmt::Debug for IntegrityBackend {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("IntegrityBackend")
.field("inner", &self.inner.backend_kind())
.field("signer", &"<redacted>")
.finish()
}
}
impl IntegrityBackend {
pub fn new(inner: Arc<dyn CacheBackend>, key: [u8; TAG_SIZE]) -> Self {
Self {
inner,
signer: HmacSigner::new(key),
}
}
pub fn from_slice_key(
inner: Arc<dyn CacheBackend>,
key: &[u8],
) -> crate::error::OxCacheResult<Self> {
Ok(Self {
inner,
signer: HmacSigner::from_slice(key)?,
})
}
pub fn signer(&self) -> &HmacSigner {
&self.signer
}
fn record_integrity_failure(&self) -> Option<Vec<u8>> {
#[cfg(feature = "metrics")]
{
use crate::infra::metrics::{CacheOpResult, CacheOpType, CacheOperation};
crate::infra::GLOBAL_UNIFIED_METRICS.record_operation(CacheOperation {
layer: crate::core::CacheLayer::L1,
op_type: CacheOpType::Get,
result: CacheOpResult::Miss,
});
crate::infra::GLOBAL_UNIFIED_METRICS
.increment_counter("oxcache_integrity_failures_total", 1);
}
None
}
}
#[async_trait::async_trait]
impl CacheReader for IntegrityBackend {
async fn get(&self, key: &str) -> OxCacheResult<Option<Vec<u8>>> {
match self.inner.get(key).await? {
Some(envelope) => {
if envelope.is_empty() || envelope[0] != HMAC_ENVELOPE_VERSION {
return Ok(self.record_integrity_failure());
}
if envelope.len() < 1 + TAG_SIZE {
return Ok(self.record_integrity_failure());
}
let tag = &envelope[1..1 + TAG_SIZE];
let payload = &envelope[1 + TAG_SIZE..];
if self.signer.verify(payload, tag) {
Ok(Some(payload.to_vec()))
} else {
Ok(self.record_integrity_failure())
}
}
None => Ok(None),
}
}
async fn exists(&self, key: &str) -> OxCacheResult<bool> {
self.inner.exists(key).await
}
async fn ttl(&self, key: &str) -> OxCacheResult<Option<Duration>> {
self.inner.ttl(key).await
}
async fn len(&self) -> OxCacheResult<u64> {
self.inner.len().await
}
async fn capacity(&self) -> OxCacheResult<u64> {
self.inner.capacity().await
}
async fn stats(&self) -> OxCacheResult<HashMap<String, String>> {
self.inner.stats().await
}
async fn keys(&self, pattern: &str) -> OxCacheResult<Vec<String>> {
self.inner.keys(pattern).await
}
}
#[async_trait::async_trait]
impl CacheWriter for IntegrityBackend {
async fn set(
&self,
key: Arc<str>,
value: Arc<Vec<u8>>,
ttl: Option<Duration>,
) -> OxCacheResult<()> {
let tag = self.signer.sign(value.as_slice());
let mut envelope = Vec::with_capacity(1 + TAG_SIZE + value.len());
envelope.push(HMAC_ENVELOPE_VERSION);
envelope.extend_from_slice(&tag);
envelope.extend_from_slice(value.as_slice());
self.inner.set(key, Arc::new(envelope), ttl).await
}
async fn delete(&self, key: &str) -> OxCacheResult<()> {
self.inner.delete(key).await
}
async fn clear(&self) -> OxCacheResult<()> {
self.inner.clear().await
}
async fn expire(&self, key: &str, ttl: Duration) -> OxCacheResult<bool> {
self.inner.expire(key, ttl).await
}
async fn set_many(&self, items: &[CacheSetItem]) -> OxCacheResult<()> {
let mut signed = Vec::with_capacity(items.len());
for (key, value, ttl) in items {
let tag = self.signer.sign(value.as_slice());
let mut envelope = Vec::with_capacity(1 + TAG_SIZE + value.len());
envelope.push(HMAC_ENVELOPE_VERSION);
envelope.extend_from_slice(&tag);
envelope.extend_from_slice(value.as_slice());
signed.push((key.clone(), Arc::new(envelope), *ttl));
}
self.inner.set_many(&signed).await
}
async fn delete_many(&self, keys: &[String]) -> OxCacheResult<()> {
self.inner.delete_many(keys).await
}
}
#[async_trait::async_trait]
impl CacheConnector for IntegrityBackend {
async fn health_check(&self) -> OxCacheResult<()> {
self.inner.health_check().await
}
async fn shutdown(&self) {
self.inner.shutdown().await;
}
fn backend_kind(&self) -> BackendKind {
self.inner.backend_kind()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::backend::MockBackend;
use crate::backend::interface::{CacheReader, CacheWriter};
fn key32(seed: u8) -> [u8; TAG_SIZE] {
let mut key = [0u8; TAG_SIZE];
for (i, b) in key.iter_mut().enumerate() {
*b = seed.wrapping_add(i as u8);
}
key
}
fn integrity(seed: u8) -> IntegrityBackend {
IntegrityBackend::new(Arc::new(MockBackend::new("mock", 100, false)), key32(seed))
}
#[tokio::test]
async fn hmac_roundtrip_is_transparent() {
let backend = integrity(1);
backend
.set(Arc::from("k"), Arc::new(b"trusted-value".to_vec()), None)
.await
.unwrap();
assert_eq!(
backend.get("k").await.unwrap(),
Some(b"trusted-value".to_vec())
);
}
#[tokio::test]
async fn raw_storage_format_is_ver_tag_payload() {
let backend = integrity(1);
backend
.set(Arc::from("k"), Arc::new(b"payload!".to_vec()), None)
.await
.unwrap();
let raw = backend.inner.get("k").await.unwrap().unwrap();
assert_eq!(raw[0], HMAC_ENVELOPE_VERSION);
let tag = &raw[1..1 + TAG_SIZE];
assert_eq!(raw.len(), 1 + TAG_SIZE + 8);
assert!(backend.signer().verify(b"payload!", tag));
assert!(!backend.signer().verify(b"tampered!", tag));
}
#[tokio::test]
async fn tampered_payload_reads_as_miss() {
let backend = integrity(1);
backend
.set(Arc::from("k"), Arc::new(b"original".to_vec()), None)
.await
.unwrap();
let mut raw = backend.inner.get("k").await.unwrap().unwrap();
let payload_len = raw.len();
raw[payload_len - 1] ^= 0xFF; backend
.inner
.set(Arc::from("k"), Arc::new(raw), None)
.await
.unwrap();
assert_eq!(backend.get("k").await.unwrap(), None, "篡改后应视为 miss");
}
#[tokio::test]
async fn tampered_tag_reads_as_miss() {
let backend = integrity(1);
backend
.set(Arc::from("k"), Arc::new(b"original".to_vec()), None)
.await
.unwrap();
let mut raw = backend.inner.get("k").await.unwrap().unwrap();
raw[1] ^= 0x01; backend
.inner
.set(Arc::from("k"), Arc::new(raw), None)
.await
.unwrap();
assert_eq!(backend.get("k").await.unwrap(), None);
}
#[tokio::test]
async fn tampered_version_reads_as_miss() {
let backend = integrity(1);
backend
.set(Arc::from("k"), Arc::new(b"original".to_vec()), None)
.await
.unwrap();
let mut raw = backend.inner.get("k").await.unwrap().unwrap();
raw[0] = 0xFF; backend
.inner
.set(Arc::from("k"), Arc::new(raw), None)
.await
.unwrap();
assert_eq!(backend.get("k").await.unwrap(), None);
}
#[tokio::test]
#[serial_test::serial]
async fn integrity_failure_counts_miss_metric() {
let backend = integrity(1);
backend
.set(Arc::from("k"), Arc::new(b"original".to_vec()), None)
.await
.unwrap();
let mut raw = backend.inner.get("k").await.unwrap().unwrap();
let last = raw.len() - 1;
raw[last] ^= 0xFF;
backend
.inner
.set(Arc::from("k"), Arc::new(raw), None)
.await
.unwrap();
#[cfg(feature = "metrics")]
{
let before = crate::infra::GLOBAL_UNIFIED_METRICS
.get_counters()
.l1_misses;
let before_failures = crate::infra::GLOBAL_UNIFIED_METRICS
.get_dynamic_metrics()
.get("oxcache_integrity_failures_total")
.map(|v| match v {
crate::infra::metrics::MetricValue::Counter(c) => *c,
_ => 0,
})
.unwrap_or(0);
assert_eq!(backend.get("k").await.unwrap(), None);
let after = crate::infra::GLOBAL_UNIFIED_METRICS
.get_counters()
.l1_misses;
let after_failures = crate::infra::GLOBAL_UNIFIED_METRICS
.get_dynamic_metrics()
.get("oxcache_integrity_failures_total")
.map(|v| match v {
crate::infra::metrics::MetricValue::Counter(c) => *c,
_ => 0,
})
.unwrap_or(0);
assert_eq!(after, before + 1, "完整性失败应计 1 次 miss");
assert_eq!(after_failures, before_failures + 1, "完整性失败计数应递增");
}
#[cfg(not(feature = "metrics"))]
{
assert_eq!(backend.get("k").await.unwrap(), None);
}
}
#[tokio::test]
async fn wrong_key_reads_as_miss_not_error() {
let backend = integrity(1);
backend
.set(Arc::from("k"), Arc::new(b"original".to_vec()), None)
.await
.unwrap();
let other = IntegrityBackend::new(backend.inner.clone(), key32(99));
assert_eq!(other.get("k").await.unwrap(), None);
}
#[cfg(feature = "encrypt")]
#[tokio::test]
async fn composes_with_encryption_in_both_orders() {
use super::super::EncryptedBackend;
let mk_inner =
|| -> Arc<dyn CacheBackend> { Arc::new(MockBackend::new("mock", 100, false)) };
let a = IntegrityBackend::new(
Arc::new(EncryptedBackend::new(mk_inner(), key32(7))),
key32(8),
);
a.set(Arc::from("k"), Arc::new(b"both".to_vec()), None)
.await
.unwrap();
assert_eq!(a.get("k").await.unwrap(), Some(b"both".to_vec()));
let b = EncryptedBackend::new(
Arc::new(IntegrityBackend::new(mk_inner(), key32(8))),
key32(7),
);
b.set(Arc::from("k"), Arc::new(b"both".to_vec()), None)
.await
.unwrap();
assert_eq!(b.get("k").await.unwrap(), Some(b"both".to_vec()));
}
#[test]
fn hmac_key_length_validated_at_construction() {
let inner: Arc<dyn CacheBackend> = Arc::new(MockBackend::new("mock", 100, false));
let err = IntegrityBackend::from_slice_key(inner, b"tiny")
.expect_err("密钥长度错误必须在构造期报错");
assert!(matches!(err, crate::error::OxCacheError::InvalidInput(_)));
}
}