use alloy::primitives::Address;
use moka::future::Cache;
use std::time::Duration;
#[derive(Clone, Debug)]
pub struct PolicyContractCacheConfig {
pub version_ttl_secs: u64,
pub version_max_capacity: u64,
}
impl Default for PolicyContractCacheConfig {
fn default() -> Self {
Self {
version_ttl_secs: 600,
version_max_capacity: 512,
}
}
}
pub struct PolicyContractCache {
policy_version: Cache<(u64, Address), String>,
}
impl std::fmt::Debug for PolicyContractCache {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PolicyContractCache")
.field("policy_version_entry_count", &self.policy_version.entry_count())
.finish()
}
}
impl Default for PolicyContractCache {
fn default() -> Self {
Self::new(PolicyContractCacheConfig::default())
}
}
impl PolicyContractCache {
pub fn new(config: PolicyContractCacheConfig) -> Self {
Self {
policy_version: Cache::builder()
.time_to_live(Duration::from_secs(config.version_ttl_secs))
.max_capacity(config.version_max_capacity)
.build(),
}
}
pub async fn get_policy_version(&self, chain_id: u64, policy_address: Address) -> Option<String> {
self.policy_version.get(&(chain_id, policy_address)).await
}
pub async fn insert_policy_version(&self, chain_id: u64, policy_address: Address, version: String) {
self.policy_version.insert((chain_id, policy_address), version).await;
}
pub async fn get_or_try_insert_policy_version<F, E>(
&self,
chain_id: u64,
policy_address: Address,
init: F,
) -> Result<String, std::sync::Arc<E>>
where
F: std::future::Future<Output = Result<String, E>>,
E: Send + Sync + 'static,
{
self.policy_version.try_get_with((chain_id, policy_address), init).await
}
pub async fn invalidate_policy_version(&self, chain_id: u64, policy_address: Address) {
self.policy_version.invalidate(&(chain_id, policy_address)).await;
}
}
#[cfg(test)]
mod tests {
use super::*;
use alloy::primitives::Address;
use std::time::Duration;
fn test_config() -> PolicyContractCacheConfig {
PolicyContractCacheConfig {
version_ttl_secs: 1,
version_max_capacity: 10,
}
}
#[tokio::test]
async fn test_version_cache_miss_then_hit() {
let cache = PolicyContractCache::new(test_config());
let chain_id = 1u64;
let addr = Address::repeat_byte(0x01);
assert!(cache.get_policy_version(chain_id, addr).await.is_none());
cache.insert_policy_version(chain_id, addr, "0.3.0".to_string()).await;
let cached = cache.get_policy_version(chain_id, addr).await.unwrap();
assert_eq!(cached, "0.3.0");
}
#[tokio::test]
async fn test_version_ttl_expiry() {
let cache = PolicyContractCache::new(test_config());
let chain_id = 1u64;
let addr = Address::repeat_byte(0x01);
cache.insert_policy_version(chain_id, addr, "0.3.0".to_string()).await;
assert!(cache.get_policy_version(chain_id, addr).await.is_some());
tokio::time::sleep(Duration::from_millis(1100)).await;
assert!(cache.get_policy_version(chain_id, addr).await.is_none());
}
#[tokio::test]
async fn test_default_config() {
let config = PolicyContractCacheConfig::default();
assert_eq!(config.version_ttl_secs, 600);
assert_eq!(config.version_max_capacity, 512);
}
}