use crate::comm::ExecuteContext;
use crate::errors::Result;
use crate::interceptor::shared::InterceptorBase;
use crate::interceptor::{InterceptorType, OperationType};
use std::collections::HashMap;
use std::sync::{Arc, RwLock};
use std::time::{Duration, Instant};
struct CacheEntry {
data: Vec<u8>,
created_at: Instant,
ttl: Duration,
}
impl CacheEntry {
fn is_expired(&self) -> bool {
self.created_at.elapsed() > self.ttl
}
}
pub trait CacheProvider: Send + Sync {
fn get(&self, key: &str) -> Option<Vec<u8>>;
fn set(&self, key: &str, value: Vec<u8>, ttl: Duration);
fn delete(&self, key: &str);
fn clear(&self);
fn size(&self) -> usize;
}
pub struct MemoryCacheProvider {
cache: Arc<RwLock<HashMap<String, CacheEntry>>>,
max_size: usize,
}
impl MemoryCacheProvider {
pub fn new() -> Self {
Self {
cache: Arc::new(RwLock::new(HashMap::new())),
max_size: 1000,
}
}
pub fn with_max_size(mut self, max_size: usize) -> Self {
self.max_size = max_size;
self
}
fn cleanup(&self) {
if let Ok(mut cache) = self.cache.write() {
cache.retain(|_, entry| !entry.is_expired());
}
}
}
impl Default for MemoryCacheProvider {
fn default() -> Self {
Self::new()
}
}
impl CacheProvider for MemoryCacheProvider {
fn get(&self, key: &str) -> Option<Vec<u8>> {
self.cleanup();
if let Ok(cache) = self.cache.read() {
cache.get(key).map(|entry| entry.data.clone())
} else {
None
}
}
fn set(&self, key: &str, value: Vec<u8>, ttl: Duration) {
self.cleanup();
if let Ok(mut cache) = self.cache.write() {
if cache.len() >= self.max_size {
if let Some(oldest_key) = cache
.iter()
.min_by_key(|(_, entry)| entry.created_at)
.map(|(k, _)| k.clone())
{
cache.remove(&oldest_key);
}
}
cache.insert(
key.to_string(),
CacheEntry {
data: value,
created_at: Instant::now(),
ttl,
},
);
}
}
fn delete(&self, key: &str) {
if let Ok(mut cache) = self.cache.write() {
cache.remove(key);
}
}
fn clear(&self) {
if let Ok(mut cache) = self.cache.write() {
cache.clear();
}
}
fn size(&self) -> usize {
if let Ok(cache) = self.cache.read() {
cache.len()
} else {
0
}
}
}
pub struct NoopCacheProvider;
impl NoopCacheProvider {
pub fn new() -> Self {
Self
}
}
impl Default for NoopCacheProvider {
fn default() -> Self {
Self::new()
}
}
impl CacheProvider for NoopCacheProvider {
fn get(&self, _key: &str) -> Option<Vec<u8>> {
None
}
fn set(&self, _key: &str, _value: Vec<u8>, _ttl: Duration) {}
fn delete(&self, _key: &str) {}
fn clear(&self) {}
fn size(&self) -> usize {
0
}
}
pub struct CacheInterceptor {
provider: Box<dyn CacheProvider>,
default_ttl: Duration,
pub(crate) enabled: bool,
cache_prefix: String,
}
impl CacheInterceptor {
pub fn new(provider: Box<dyn CacheProvider>) -> Self {
Self {
provider,
default_ttl: Duration::from_secs(300), enabled: true,
cache_prefix: "akita_cache:".to_string(),
}
}
pub fn with_default_ttl(mut self, ttl: Duration) -> Self {
self.default_ttl = ttl;
self
}
pub fn with_enabled(mut self, enabled: bool) -> Self {
self.enabled = enabled;
self
}
pub fn with_cache_prefix(mut self, prefix: &str) -> Self {
self.cache_prefix = prefix.to_string();
self
}
pub fn generate_cache_key(&self, sql: &str) -> String {
use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher};
let mut hasher = DefaultHasher::new();
sql.hash(&mut hasher);
let hash = hasher.finish();
format!("{}{:016x}", self.cache_prefix, hash)
}
pub fn get_from_cache(&self, key: &str) -> Option<Vec<u8>> {
self.provider.get(key)
}
pub fn set_in_cache(&self, key: &str, value: Vec<u8>) {
self.provider.set(key, value, self.default_ttl);
}
pub fn invalidate_pattern(&self, pattern: &str) {
if pattern.contains('*') || pattern.contains('%') {
self.provider.clear();
} else {
self.provider.delete(pattern);
}
}
}
impl InterceptorBase for CacheInterceptor {
fn name(&self) -> &'static str {
"cache"
}
fn interceptor_type(&self) -> InterceptorType {
InterceptorType::Cache
}
fn order(&self) -> i32 {
5 }
fn supports_operation(&self, _operation: &OperationType) -> bool {
true
}
}
#[cfg(any(
feature = "mysql-sync",
feature = "postgres-sync",
feature = "sqlite-sync",
feature = "oracle-sync",
feature = "mssql-sync"
))]
impl crate::interceptor::blocking::AkitaInterceptor for CacheInterceptor {
fn before_execute(&self, ctx: &mut ExecuteContext) -> Result<()> {
if !self.enabled {
return Ok(());
}
let sql = ctx.final_sql();
let key = self.generate_cache_key(sql);
if matches!(ctx.operation_type(), OperationType::Select) {
if let Some(cached) = self.get_from_cache(&key) {
ctx.set_metadata("cache_hit".to_string(), "true".to_string());
ctx.set_metadata("cache_key".to_string(), key);
ctx.set_metadata(
"cached_data".to_string(),
String::from_utf8_lossy(&cached).to_string(),
);
} else {
ctx.set_metadata("cache_hit".to_string(), "false".to_string());
ctx.set_metadata("cache_key".to_string(), key);
}
} else {
ctx.set_metadata("cache_invalidate".to_string(), "true".to_string());
}
Ok(())
}
fn after_execute(
&self,
ctx: &mut ExecuteContext,
result: &mut std::result::Result<crate::comm::ExecuteResult, crate::errors::AkitaError>,
) -> Result<()> {
if !self.enabled {
return Ok(());
}
if matches!(ctx.operation_type(), OperationType::Select) {
let cache_hit = ctx
.get_metadata("cache_hit")
.map(|v| v.to_string())
.unwrap_or_default();
if cache_hit == "false" {
let key = ctx
.get_metadata("cache_key")
.map(|v| v.to_string())
.unwrap_or_default();
if !key.is_empty() {
self.set_in_cache(&key, b"cached".to_vec());
}
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_memory_cache_provider() {
let provider = MemoryCacheProvider::new();
provider.set("key1", b"value1".to_vec(), Duration::from_secs(60));
assert_eq!(provider.get("key1"), Some(b"value1".to_vec()));
provider.delete("key1");
assert_eq!(provider.get("key1"), None);
provider.set("key2", b"value2".to_vec(), Duration::from_secs(60));
provider.set("key3", b"value3".to_vec(), Duration::from_secs(60));
assert_eq!(provider.size(), 2);
provider.clear();
assert_eq!(provider.size(), 0);
}
#[test]
fn test_memory_cache_provider_max_size() {
let provider = MemoryCacheProvider::new().with_max_size(2);
provider.set("key1", b"value1".to_vec(), Duration::from_secs(60));
provider.set("key2", b"value2".to_vec(), Duration::from_secs(60));
provider.set("key3", b"value3".to_vec(), Duration::from_secs(60));
assert!(provider.size() <= 2);
}
#[test]
fn test_noop_cache_provider() {
let provider = NoopCacheProvider;
provider.set("key1", b"value1".to_vec(), Duration::from_secs(60));
assert_eq!(provider.get("key1"), None);
assert_eq!(provider.size(), 0);
}
#[test]
fn test_cache_interceptor() {
let provider = MemoryCacheProvider::new();
let interceptor = CacheInterceptor::new(Box::new(provider))
.with_default_ttl(Duration::from_secs(60))
.with_cache_prefix("test:");
assert_eq!(interceptor.name(), "cache");
assert_eq!(interceptor.interceptor_type(), InterceptorType::Cache);
assert_eq!(interceptor.order(), 5);
}
#[test]
fn test_generate_cache_key() {
let provider = MemoryCacheProvider::new();
let interceptor = CacheInterceptor::new(Box::new(provider)).with_cache_prefix("test:");
let key1 = interceptor.generate_cache_key("SELECT * FROM users");
let key2 = interceptor.generate_cache_key("SELECT * FROM users");
let key3 = interceptor.generate_cache_key("SELECT * FROM orders");
assert_eq!(key1, key2);
assert_ne!(key1, key3);
assert!(key1.starts_with("test:"));
}
}