use crate::error::{Error, Result};
use serde::{Deserialize, Serialize};
pub const CACHE_MAGIC: [u8; 4] = *b"CKIT";
pub const CURRENT_SCHEMA_VERSION: u32 = 1;
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
pub struct CacheEnvelope<T> {
pub magic: [u8; 4],
pub version: u32,
pub payload: T,
}
impl<T> CacheEnvelope<T> {
pub fn new(payload: T) -> Self {
Self {
magic: CACHE_MAGIC,
version: CURRENT_SCHEMA_VERSION,
payload,
}
}
}
pub fn serialize_for_cache<T: Serialize>(value: &T) -> Result<Vec<u8>> {
let envelope = CacheEnvelope::new(value);
postcard::to_allocvec(&envelope).map_err(|e| {
log::error!("Cache serialization failed: {}", e);
Error::SerializationError(e.to_string())
})
}
pub fn deserialize_from_cache<'de, T: Deserialize<'de>>(bytes: &'de [u8]) -> Result<T> {
let envelope: CacheEnvelope<T> = postcard::from_bytes(bytes).map_err(|e| {
log::error!("Cache deserialization failed: {}", e);
Error::DeserializationError(e.to_string())
})?;
if envelope.magic != CACHE_MAGIC {
log::warn!(
"Invalid cache entry: expected magic {:?}, got {:?}",
CACHE_MAGIC,
envelope.magic
);
return Err(Error::InvalidCacheEntry(format!(
"Invalid magic: expected {:?}, got {:?}",
CACHE_MAGIC, envelope.magic
)));
}
if envelope.version != CURRENT_SCHEMA_VERSION {
log::warn!(
"Cache version mismatch: expected {}, got {}",
CURRENT_SCHEMA_VERSION,
envelope.version
);
return Err(Error::VersionMismatch {
expected: CURRENT_SCHEMA_VERSION,
found: envelope.version,
});
}
Ok(envelope.payload)
}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Serialize, Deserialize, PartialEq, Debug, Clone)]
struct TestData {
id: u64,
name: String,
active: bool,
}
#[test]
fn test_roundtrip() {
let data = TestData {
id: 123,
name: "test".to_string(),
active: true,
};
let bytes = serialize_for_cache(&data).unwrap();
let deserialized: TestData = deserialize_from_cache(&bytes).unwrap();
assert_eq!(data, deserialized);
}
#[test]
fn test_envelope_structure() {
let data = TestData {
id: 123,
name: "test".to_string(),
active: true,
};
let bytes = serialize_for_cache(&data).unwrap();
let envelope: CacheEnvelope<TestData> = postcard::from_bytes(&bytes).unwrap();
assert_eq!(envelope.magic, CACHE_MAGIC);
assert_eq!(envelope.version, CURRENT_SCHEMA_VERSION);
assert_eq!(envelope.payload, data);
}
#[test]
fn test_envelope_new() {
let envelope = CacheEnvelope::new(42);
assert_eq!(envelope.magic, CACHE_MAGIC);
assert_eq!(envelope.version, CURRENT_SCHEMA_VERSION);
assert_eq!(envelope.payload, 42);
}
#[test]
fn test_invalid_magic_rejected() {
let mut bytes = vec![0u8; 100];
bytes[0..4].copy_from_slice(b"XXXX"); bytes[4..8].copy_from_slice(&1u32.to_le_bytes());
let result: Result<TestData> = deserialize_from_cache(&bytes);
assert!(result.is_err());
match result.unwrap_err() {
Error::InvalidCacheEntry(_) => {} e => panic!("Expected InvalidCacheEntry, got {:?}", e),
}
}
#[test]
fn test_version_mismatch_rejected() {
let data = TestData {
id: 123,
name: "test".to_string(),
active: true,
};
let mut envelope = CacheEnvelope::new(&data);
envelope.version = 999;
let bytes = postcard::to_allocvec(&envelope).unwrap();
let result: Result<TestData> = deserialize_from_cache(&bytes);
assert!(result.is_err());
match result.unwrap_err() {
Error::VersionMismatch { expected, found } => {
assert_eq!(expected, CURRENT_SCHEMA_VERSION);
assert_eq!(found, 999);
}
e => panic!("Expected VersionMismatch, got {:?}", e),
}
}
#[test]
fn test_deterministic_serialization() {
let data1 = TestData {
id: 123,
name: "test".to_string(),
active: true,
};
let data2 = data1.clone();
let bytes1 = serialize_for_cache(&data1).unwrap();
let bytes2 = serialize_for_cache(&data2).unwrap();
assert_eq!(bytes1, bytes2);
}
#[test]
fn test_corrupted_payload_rejected() {
let data = TestData {
id: 123,
name: "test".to_string(),
active: true,
};
let mut bytes = serialize_for_cache(&data).unwrap();
let original_len = bytes.len();
bytes.truncate(original_len / 2);
let result: Result<TestData> = deserialize_from_cache(&bytes);
assert!(result.is_err());
match result.unwrap_err() {
Error::DeserializationError(_) => {} e => panic!("Expected DeserializationError, got {:?}", e),
}
}
#[test]
fn test_empty_data_roundtrip() {
let data = TestData {
id: 0,
name: String::new(),
active: false,
};
let bytes = serialize_for_cache(&data).unwrap();
let deserialized: TestData = deserialize_from_cache(&bytes).unwrap();
assert_eq!(data, deserialized);
}
#[test]
fn test_large_data_roundtrip() {
let data = TestData {
id: u64::MAX,
name: "x".repeat(10000),
active: true,
};
let bytes = serialize_for_cache(&data).unwrap();
let deserialized: TestData = deserialize_from_cache(&bytes).unwrap();
assert_eq!(data, deserialized);
}
#[test]
fn test_postcard_smaller_than_json() {
let data = TestData {
id: 123,
name: "test".to_string(),
active: true,
};
let postcard_bytes = serialize_for_cache(&data).unwrap();
let json_bytes = serde_json::to_vec(&data).unwrap();
assert!(
postcard_bytes.len() < json_bytes.len(),
"Postcard ({} bytes) should be smaller than JSON ({} bytes)",
postcard_bytes.len(),
json_bytes.len()
);
}
}