use std::any::{Any, TypeId};
use std::collections::HashMap;
use std::sync::Arc;
use crate::{DynProvider, InjectableError, InjectableResult, ResolveContext};
pub const DEFAULT_TOKEN: &str = "";
pub type ErasedProviderPinnedFuture<'a> = std::pin::Pin<
Box<dyn std::future::Future<Output = InjectableResult<Box<dyn Any + Send>>> + Send + 'a>,
>;
trait ErasedProvider: Send + Sync + 'static {
fn provide_as_any(&self, ctx: Arc<ResolveContext>) -> ErasedProviderPinnedFuture<'_>;
}
impl<T: Send + Sync + 'static> ErasedProvider for DynProvider<T> {
fn provide_as_any(&self, ctx: Arc<ResolveContext>) -> ErasedProviderPinnedFuture<'_> {
Box::pin(async move {
let value = self.provide(ctx).await?;
Ok(Box::new(value) as Box<dyn Any + Send>)
})
}
}
type RegistryKey = (TypeId, String);
pub struct ProviderRegistry {
providers: HashMap<RegistryKey, Box<dyn ErasedProvider>>,
duplicates: Vec<String>,
}
impl ProviderRegistry {
pub fn new() -> Self {
Self {
providers: HashMap::new(),
duplicates: Vec::new(),
}
}
fn make_key<T: 'static>(token: &str) -> RegistryKey {
(TypeId::of::<T>(), token.to_string())
}
fn duplicate_label<T: 'static>(token: &str) -> String {
if token.is_empty() {
std::any::type_name::<T>().to_string()
} else {
format!("{}[{}]", std::any::type_name::<T>(), token)
}
}
pub fn register<T: Send + Sync + 'static>(
&mut self,
token: impl Into<String>,
provider: DynProvider<T>,
) {
let token = token.into();
let key = Self::make_key::<T>(&token);
if self.providers.contains_key(&key) {
self.duplicates.push(Self::duplicate_label::<T>(&token));
}
self.providers.insert(key, Box::new(provider));
}
pub fn register_or_replace<T: Send + Sync + 'static>(
&mut self,
token: impl Into<String>,
provider: DynProvider<T>,
) {
let key = (TypeId::of::<T>(), token.into());
self.providers.insert(key, Box::new(provider));
}
pub fn duplicates(&self) -> &[String] {
&self.duplicates
}
pub fn has<T: 'static>(&self) -> bool {
self.has_with_token::<T>(DEFAULT_TOKEN)
}
pub fn has_with_token<T: 'static>(&self, token: &str) -> bool {
self.providers.contains_key(&Self::make_key::<T>(token))
}
pub(crate) async fn resolve_with_token<T: Send + Sync + 'static>(
&self,
token: &str,
ctx: Arc<ResolveContext>,
) -> Option<InjectableResult<T>> {
let key = Self::make_key::<T>(token);
if let Some(provider) = self.providers.get(&key) {
let result = provider.provide_as_any(Arc::clone(&ctx)).await;
return Some(
result.and_then(|boxed| match boxed.downcast::<T>() {
Ok(t) => Ok(*t),
Err(_) => Err(InjectableError::ConstructionFailed {
type_name: std::any::type_name::<T>(),
reason: "downcast failed (this should never happen with correct TypeId)"
.to_string(),
}),
}),
);
}
if token == DEFAULT_TOKEN {
let target_id = TypeId::of::<T>();
for factory in inventory::iter::<InjectableArcFactory>() {
if factory.type_id() == target_id {
let result = factory.provide(ctx).await;
return Some(result.and_then(|boxed| match boxed.downcast::<T>() {
Ok(t) => Ok(*t),
Err(_) => Err(InjectableError::ConstructionFailed {
type_name: std::any::type_name::<T>(),
reason: "InjectableArcFactory downcast failed".to_string(),
}),
}));
}
}
}
None
}
pub fn len(&self) -> usize {
self.providers.len()
}
pub fn is_empty(&self) -> bool {
self.providers.is_empty()
}
}
impl Default for ProviderRegistry {
fn default() -> Self {
Self::new()
}
}
impl std::fmt::Debug for ProviderRegistry {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ProviderRegistry")
.field("count", &self.providers.len())
.finish()
}
}
pub type InjectableProvideFnPtr = fn(
std::sync::Arc<ResolveContext>,
) -> std::pin::Pin<
Box<
dyn std::future::Future<Output = InjectableResult<Box<dyn std::any::Any + Send>>>
+ Send
+ 'static,
>,
>;
pub struct InjectableArcFactory {
pub type_name: &'static str,
type_id_fn: fn() -> std::any::TypeId,
provide_fn: InjectableProvideFnPtr,
}
impl InjectableArcFactory {
pub const fn new_const(
type_name: &'static str,
type_id_fn: fn() -> std::any::TypeId,
provide_fn: InjectableProvideFnPtr,
) -> Self {
Self {
type_name,
type_id_fn,
provide_fn,
}
}
pub fn type_id(&self) -> std::any::TypeId {
(self.type_id_fn)()
}
pub fn provide(&self, ctx: std::sync::Arc<ResolveContext>) -> ErasedProviderPinnedFuture<'_> {
(self.provide_fn)(ctx)
}
}
inventory::collect!(InjectableArcFactory);
pub type PostConstructFnPtr = fn(
std::sync::Arc<dyn std::any::Any + std::marker::Send + std::marker::Sync>,
) -> std::pin::Pin<
Box<dyn std::future::Future<Output = crate::HookResult> + std::marker::Send + 'static>,
>;
pub type MakePreDestructFnPtr = fn(
std::sync::Arc<dyn std::any::Any + std::marker::Send + std::marker::Sync>,
) -> std::sync::Arc<dyn crate::PreDestruct>;
pub struct InjectableHooksEntry {
type_id_fn: fn() -> std::any::TypeId,
post_construct_fn: Option<PostConstructFnPtr>,
make_pre_destruct_fn: Option<MakePreDestructFnPtr>,
}
impl InjectableHooksEntry {
pub const fn new_const(
type_id_fn: fn() -> std::any::TypeId,
post_construct_fn: Option<PostConstructFnPtr>,
make_pre_destruct_fn: Option<MakePreDestructFnPtr>,
) -> Self {
Self {
type_id_fn,
post_construct_fn,
make_pre_destruct_fn,
}
}
pub fn type_id(&self) -> std::any::TypeId {
(self.type_id_fn)()
}
pub fn post_construct_fn(&self) -> Option<PostConstructFnPtr> {
self.post_construct_fn
}
pub fn make_pre_destruct_fn(&self) -> Option<MakePreDestructFnPtr> {
self.make_pre_destruct_fn
}
}
inventory::collect!(InjectableHooksEntry);
#[cfg(test)]
mod tests {
use super::*;
use crate::DynProvider;
#[test]
fn new_registry_is_empty() {
let r = ProviderRegistry::new();
assert!(r.is_empty());
assert_eq!(r.len(), 0);
}
#[test]
fn has_returns_false_for_unregistered() {
let r = ProviderRegistry::new();
assert!(!r.has::<u32>());
assert!(!r.has_with_token::<u32>("primary"));
}
#[test]
fn has_returns_true_after_register_default() {
let mut r = ProviderRegistry::new();
r.register("", DynProvider::from_value(42u32));
assert!(r.has::<u32>());
assert!(r.has_with_token::<u32>(""));
assert!(!r.has_with_token::<u32>("other"));
assert_eq!(r.len(), 1);
assert!(!r.is_empty());
}
#[test]
fn has_returns_true_after_register_named() {
let mut r = ProviderRegistry::new();
r.register("primary", DynProvider::from_value(42u32));
assert!(!r.has::<u32>(), "default token should be absent");
assert!(r.has_with_token::<u32>("primary"));
assert!(!r.has_with_token::<u32>("replica"));
}
#[test]
fn multiple_tokens_same_type_coexist() {
let mut r = ProviderRegistry::new();
r.register("primary", DynProvider::from_value(1u32));
r.register("replica", DynProvider::from_value(2u32));
r.register("", DynProvider::from_value(0u32));
assert_eq!(r.len(), 3);
assert!(r.has::<u32>());
assert!(r.has_with_token::<u32>("primary"));
assert!(r.has_with_token::<u32>("replica"));
}
#[test]
fn duplicate_same_token_is_recorded() {
let mut r = ProviderRegistry::new();
r.register("", DynProvider::from_value(1u32));
r.register("", DynProvider::from_value(2u32));
assert_eq!(r.len(), 1); assert_eq!(r.duplicates().len(), 1);
}
#[test]
fn duplicate_different_tokens_not_recorded() {
let mut r = ProviderRegistry::new();
r.register("primary", DynProvider::from_value(1u32));
r.register("replica", DynProvider::from_value(2u32));
assert_eq!(
r.duplicates().len(),
0,
"different tokens are not duplicates"
);
}
#[test]
fn register_or_replace_does_not_record_duplicate() {
let mut r = ProviderRegistry::new();
r.register("", DynProvider::from_value(1u32));
r.register_or_replace("", DynProvider::from_value(2u32));
assert_eq!(r.duplicates().len(), 0);
}
#[test]
fn debug_shows_count() {
let mut r = ProviderRegistry::new();
r.register("", DynProvider::from_value(0u8));
let s = format!("{r:?}");
assert!(s.contains("ProviderRegistry"));
assert!(s.contains('1'));
}
#[test]
fn default_creates_empty() {
let r = ProviderRegistry::default();
assert!(r.is_empty());
}
#[test]
fn duplicate_label_includes_token_for_named() {
let label = ProviderRegistry::duplicate_label::<u32>("primary");
assert!(label.contains("primary"));
assert!(label.contains("u32"));
}
#[test]
fn duplicate_label_no_token_suffix_for_default() {
let label = ProviderRegistry::duplicate_label::<u32>("");
assert!(!label.contains('['), "default token adds no bracket suffix");
}
}