use crate::backend::CacheBackend;
use crate::error::{OxCacheError, OxCacheResult};
use once_cell::sync::Lazy;
use std::collections::HashMap;
use std::sync::{Arc, RwLock};
#[derive(Debug, Clone, Default, PartialEq, Eq)]
#[cfg_attr(
any(feature = "serialization", feature = "full"),
derive(serde::Serialize, serde::Deserialize)
)]
pub struct BackendSpec {
pub kind: String,
#[cfg_attr(any(feature = "serialization", feature = "full"), serde(default))]
pub capacity: u64,
#[cfg_attr(any(feature = "serialization", feature = "full"), serde(default))]
pub default_ttl_ms: u64,
#[cfg_attr(any(feature = "serialization", feature = "full"), serde(default))]
pub url: Option<String>,
}
impl BackendSpec {
pub fn new(kind: impl Into<String>) -> Self {
Self {
kind: kind.into(),
..Default::default()
}
}
pub fn with_capacity(mut self, capacity: u64) -> Self {
self.capacity = capacity;
self
}
pub fn with_url(mut self, url: impl Into<String>) -> Self {
self.url = Some(url.into());
self
}
}
#[async_trait::async_trait]
pub trait BackendFactory: Send + Sync {
async fn build(&self, spec: &BackendSpec) -> OxCacheResult<Arc<dyn CacheBackend>>;
}
type BoxFutureBuild = std::pin::Pin<
Box<dyn std::future::Future<Output = OxCacheResult<Arc<dyn CacheBackend>>> + Send>,
>;
pub struct FnFactory<F>(pub F);
#[async_trait::async_trait]
impl<F> BackendFactory for FnFactory<F>
where
F: Fn(&BackendSpec) -> BoxFutureBuild + Send + Sync,
{
async fn build(&self, spec: &BackendSpec) -> OxCacheResult<Arc<dyn CacheBackend>> {
(self.0)(spec).await
}
}
#[derive(Default)]
pub struct BackendRegistry {
factories: RwLock<HashMap<String, Arc<dyn BackendFactory>>>,
}
impl BackendRegistry {
pub fn new() -> Self {
Self::default()
}
pub fn register(&self, kind: impl Into<String>, factory: Arc<dyn BackendFactory>) -> &Self {
self.factories
.write()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.insert(kind.into(), factory);
self
}
pub fn register_fn<F>(&self, kind: impl Into<String>, f: F) -> &Self
where
F: Fn(&BackendSpec) -> BoxFutureBuild + Send + Sync + 'static,
{
self.register(kind, Arc::new(FnFactory(f)))
}
pub fn registered(&self) -> Vec<String> {
let mut kinds: Vec<String> = self
.factories
.read()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.keys()
.cloned()
.collect();
kinds.sort();
kinds
}
pub async fn build(&self, spec: &BackendSpec) -> OxCacheResult<Arc<dyn CacheBackend>> {
let factory = match self.factories.read() {
Ok(map) => map.get(&spec.kind).cloned(),
Err(poisoned) => poisoned.into_inner().get(&spec.kind).cloned(),
};
match factory {
Some(factory) => factory.build(spec).await,
None => Err(OxCacheError::InvalidInput(format!(
"unknown backend kind '{}'; available: {}",
spec.kind,
self.registered().join(", ")
))),
}
}
fn with_builtins() -> Self {
let registry = Self::new();
#[cfg(feature = "memory")]
{
registry.register_fn("moka", |spec: &BackendSpec| {
let spec = spec.clone();
Box::pin(async move {
let mut builder =
crate::backend::MokaMemoryBackend::builder().capacity(spec.capacity.max(1));
if spec.default_ttl_ms > 0 {
builder =
builder.ttl(std::time::Duration::from_millis(spec.default_ttl_ms));
}
Ok(Arc::new(builder.build()) as Arc<dyn CacheBackend>)
})
});
registry.register_fn("dashmap", |spec: &BackendSpec| {
let spec = spec.clone();
Box::pin(async move {
let mut builder = crate::backend::DashMapMemoryBackend::builder();
if spec.capacity > 0 {
builder = builder.capacity(spec.capacity as usize);
}
if spec.default_ttl_ms > 0 {
builder = builder
.default_ttl(std::time::Duration::from_millis(spec.default_ttl_ms));
}
Ok(Arc::new(builder.build()) as Arc<dyn CacheBackend>)
})
});
registry.register_fn("memory", |spec: &BackendSpec| {
let spec = spec.clone();
Box::pin(async move {
let mut builder =
crate::backend::MokaMemoryBackend::builder().capacity(spec.capacity.max(1));
if spec.default_ttl_ms > 0 {
builder =
builder.ttl(std::time::Duration::from_millis(spec.default_ttl_ms));
}
Ok(Arc::new(builder.build()) as Arc<dyn CacheBackend>)
})
});
}
#[cfg(feature = "redis")]
{
registry.register_fn("redis", |spec: &BackendSpec| {
let spec = spec.clone();
Box::pin(async move {
let url = spec
.url
.clone()
.unwrap_or_else(|| "redis://127.0.0.1:6379".to_string());
Ok(Arc::new(crate::backend::RedisBackend::new(&url).await?)
as Arc<dyn CacheBackend>)
})
});
}
registry
}
}
pub static GLOBAL_BACKEND_REGISTRY: Lazy<BackendRegistry> =
Lazy::new(BackendRegistry::with_builtins);
impl BackendRegistry {
pub fn global() -> &'static BackendRegistry {
&GLOBAL_BACKEND_REGISTRY
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::backend::MockBackend;
#[tokio::test]
async fn builtin_memory_factories_build() {
let registry = BackendRegistry::global();
let spec = BackendSpec::new("moka").with_capacity(500);
let backend = registry.build(&spec).await.unwrap();
assert_eq!(backend.capacity().await.unwrap(), 500);
assert_eq!(backend.backend_kind(), crate::backend::BackendKind::Moka);
let backend = registry.build(&BackendSpec::new("memory")).await.unwrap();
assert_eq!(backend.backend_kind(), crate::backend::BackendKind::Moka);
let backend = registry
.build(&BackendSpec::new("dashmap").with_capacity(64))
.await
.unwrap();
assert_eq!(backend.backend_kind(), crate::backend::BackendKind::DashMap);
assert_eq!(backend.capacity().await.unwrap(), 64);
}
#[tokio::test]
async fn unknown_kind_error_lists_available() {
let registry = BackendRegistry::global();
let err = match registry.build(&BackendSpec::new("nosuch-backend")).await {
Err(e) => e,
Ok(_) => panic!("未知 kind 必须报错"),
};
let msg = err.to_string();
assert!(msg.contains("nosuch-backend"), "msg: {msg}");
assert!(msg.contains("available:"), "错误应附带可用列表: {msg}");
}
#[tokio::test]
async fn custom_factory_registration_and_lookup() {
let registry = BackendRegistry::new();
registry.register_fn("mock:test", |spec: &BackendSpec| {
let spec = spec.clone();
Box::pin(async move {
Ok(Arc::new(MockBackend::new(
"mock-test",
spec.capacity.min(255) as u8,
false,
)) as Arc<dyn CacheBackend>)
})
});
assert_eq!(registry.registered(), vec!["mock:test".to_string()]);
let backend = registry
.build(&BackendSpec::new("mock:test").with_capacity(7))
.await
.unwrap();
assert!(backend.exists("nothing").await.unwrap().eq(&false));
registry.register_fn("mock:test", |_spec| {
Box::pin(async move {
Ok(Arc::new(MockBackend::new("mock-2", 10, true)) as Arc<dyn CacheBackend>)
})
});
let backend = registry
.build(&BackendSpec::new("mock:test"))
.await
.unwrap();
assert_eq!(
backend
.stats()
.await
.unwrap()
.get("type")
.map(String::as_str),
Some("mock-2")
);
}
#[tokio::test]
async fn empty_registry_reports_empty_available_list() {
let registry = BackendRegistry::new();
let err = match registry.build(&BackendSpec::new("anything")).await {
Err(e) => e,
Ok(_) => panic!("空注册中心必须报错"),
};
assert!(err.to_string().contains("available: "), "got {err}");
}
#[test]
fn spec_serde_roundtrip() {
let spec = BackendSpec::new("moka").with_capacity(100);
let json = serde_json::to_string(&spec).unwrap();
let back: BackendSpec = serde_json::from_str(&json).unwrap();
assert_eq!(back, spec);
}
}