use std::any::{Any, TypeId};
use std::collections::HashMap;
use std::sync::Arc;
use crate::{
Injectable, InjectableError, InjectableResult, PreDestruct, Provider, ProviderRegistry,
SingletonStore,
};
pub type SingletonCache = Arc<
tokio::sync::Mutex<HashMap<TypeId, Arc<tokio::sync::OnceCell<Arc<dyn Any + Send + Sync>>>>>,
>;
struct DestructorEntry {
instance: Arc<dyn PreDestruct>,
type_name: &'static str,
}
pub struct ResolveContext {
store: Arc<dyn SingletonStore>,
registry: Arc<ProviderRegistry>,
destructors: Arc<tokio::sync::Mutex<Vec<DestructorEntry>>>,
singleton_cache: SingletonCache,
}
impl ResolveContext {
pub fn new(store: Arc<dyn SingletonStore>, registry: Arc<ProviderRegistry>) -> Self {
Self {
store,
registry,
destructors: Arc::new(tokio::sync::Mutex::new(Vec::new())),
singleton_cache: Arc::new(tokio::sync::Mutex::new(HashMap::new())),
}
}
pub fn from_store(store: Arc<dyn SingletonStore>) -> Self {
Self {
store,
registry: Arc::new(ProviderRegistry::new()),
destructors: Arc::new(tokio::sync::Mutex::new(Vec::new())),
singleton_cache: Arc::new(tokio::sync::Mutex::new(HashMap::new())),
}
}
pub fn store(&self) -> &Arc<dyn SingletonStore> {
&self.store
}
pub fn registry(&self) -> &ProviderRegistry {
&self.registry
}
pub async fn extract<T>(&self) -> crate::InjectableResult<T>
where
T: crate::Extract + Send + Sync + 'static,
{
T::extract(self).await
}
pub async fn clone_from_singleton<T: Injectable + Clone>(&self) -> InjectableResult<T> {
Ok(Arc::unwrap_or_clone(
self.resolve_singleton_arc::<T>().await?,
))
}
#[allow(clippy::expect_used)]
pub(crate) async fn resolve_singleton_arc<T: Injectable>(&self) -> InjectableResult<Arc<T>> {
let type_id = TypeId::of::<T>();
let cell = {
let mut cache = self.singleton_cache.lock().await;
Arc::clone(
cache
.entry(type_id)
.or_insert_with(|| Arc::new(tokio::sync::OnceCell::new())),
)
};
let arc_any = cell
.get_or_try_init(|| async {
let value = T::Provider::provide(self).await?;
let arc_t: Arc<T> = Arc::new(value);
Ok(Arc::new(arc_t) as Arc<dyn Any + Send + Sync>)
})
.await?;
let inner: &Arc<T> = (**arc_any)
.downcast_ref::<Arc<T>>()
.expect("TypeId guarantees type correctness; downcast cannot fail");
Ok(Arc::clone(inner))
}
pub async fn resolve_external<T: Send + Sync + 'static>(&self) -> InjectableResult<T> {
self.resolve_external_with_token::<T>(crate::registry::DEFAULT_TOKEN)
.await
}
pub async fn resolve_external_with_token<T: Send + Sync + 'static>(
&self,
token: &str,
) -> InjectableResult<T> {
match self
.registry
.resolve_with_token::<T>(token, Arc::new(self.clone()))
.await
{
Some(result) => result,
None => Err(InjectableError::MissingDependency {
type_name: std::any::type_name::<T>(),
}),
}
}
pub async fn try_resolve_external<T: Send + Sync + 'static>(
&self,
) -> Option<InjectableResult<T>> {
self.registry
.resolve_with_token::<T>(crate::registry::DEFAULT_TOKEN, Arc::new(self.clone()))
.await
}
pub async fn try_resolve_external_with_token<T: Send + Sync + 'static>(
&self,
token: &str,
) -> Option<InjectableResult<T>> {
self.registry
.resolve_with_token::<T>(token, Arc::new(self.clone()))
.await
}
pub fn register_destructor(&self, instance: Arc<dyn PreDestruct>) {
if let Ok(mut destructors) = self.destructors.try_lock() {
destructors.push(DestructorEntry {
type_name: "",
instance,
});
}
}
pub fn register_destructor_with_name(
&self,
type_name: &'static str,
instance: Arc<dyn PreDestruct>,
) {
if let Ok(mut destructors) = self.destructors.try_lock() {
destructors.push(DestructorEntry {
type_name,
instance,
});
}
}
pub async fn run_destructors(&self) -> Result<(), Vec<crate::InjectableError>> {
let mut destructors = self.destructors.lock().await;
let mut errors = Vec::new();
while let Some(entry) = destructors.pop() {
match entry.instance.pre_destruct().await {
Ok(()) => {}
Err(e) => {
errors.push(crate::InjectableError::LifecycleHookFailed {
type_name: entry.type_name,
hook: "pre_destruct",
reason: e.to_string(),
});
}
}
}
if errors.is_empty() {
Ok(())
} else {
Err(errors)
}
}
pub async fn destructor_count(&self) -> usize {
self.destructors.lock().await.len()
}
}
impl Clone for ResolveContext {
fn clone(&self) -> Self {
Self {
store: Arc::clone(&self.store),
registry: Arc::clone(&self.registry),
destructors: Arc::clone(&self.destructors),
singleton_cache: Arc::clone(&self.singleton_cache),
}
}
}
impl std::fmt::Debug for ResolveContext {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ResolveContext")
.field("store", &"Arc<dyn SingletonStore>")
.field("registry", &self.registry)
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{DynProvider, EmptySingletonStore, HookResult, PreDestruct, ProviderRegistry};
use std::sync::Arc;
fn make_ctx() -> ResolveContext {
ResolveContext::new(
Arc::new(EmptySingletonStore),
Arc::new(ProviderRegistry::new()),
)
}
#[test]
fn from_store_creates_context() {
let ctx = ResolveContext::from_store(Arc::new(EmptySingletonStore));
assert!(ctx.registry().is_empty());
}
#[test]
fn store_and_registry_accessors() {
let ctx = make_ctx();
assert_eq!(ctx.store().len(), 0);
assert!(ctx.registry().is_empty());
}
#[test]
fn clone_shares_destructors() {
let ctx = make_ctx();
let ctx2 = ctx.clone();
assert!(Arc::ptr_eq(&ctx.destructors, &ctx2.destructors));
}
#[test]
fn debug_impl() {
let ctx = make_ctx();
let s = format!("{ctx:?}");
assert!(s.contains("ResolveContext"));
}
#[tokio::test]
async fn destructor_count_starts_zero() {
let ctx = make_ctx();
assert_eq!(ctx.destructor_count().await, 0);
}
#[tokio::test]
async fn run_destructors_empty_ok() {
let ctx = make_ctx();
assert!(ctx.run_destructors().await.is_ok());
}
#[tokio::test]
async fn register_destructor_increments_count() {
struct NoopDestructor;
#[async_trait::async_trait]
impl PreDestruct for NoopDestructor {
async fn pre_destruct(&self) -> HookResult {
Ok(())
}
}
let ctx = make_ctx();
ctx.register_destructor(Arc::new(NoopDestructor));
assert_eq!(ctx.destructor_count().await, 1);
}
#[tokio::test]
async fn register_destructor_with_name_increments_count() {
struct NoopDestructor;
#[async_trait::async_trait]
impl PreDestruct for NoopDestructor {
async fn pre_destruct(&self) -> HookResult {
Ok(())
}
}
let ctx = make_ctx();
ctx.register_destructor_with_name("TestType", Arc::new(NoopDestructor));
assert_eq!(ctx.destructor_count().await, 1);
}
#[tokio::test]
async fn run_destructors_calls_hooks() {
use std::sync::atomic::{AtomicBool, Ordering};
static CALLED: AtomicBool = AtomicBool::new(false);
struct FlagDestructor;
#[async_trait::async_trait]
impl PreDestruct for FlagDestructor {
async fn pre_destruct(&self) -> HookResult {
CALLED.store(true, Ordering::SeqCst);
Ok(())
}
}
CALLED.store(false, Ordering::SeqCst);
let ctx = make_ctx();
ctx.register_destructor(Arc::new(FlagDestructor));
ctx.run_destructors().await.unwrap();
assert!(CALLED.load(Ordering::SeqCst));
assert_eq!(ctx.destructor_count().await, 0);
}
#[tokio::test]
async fn run_destructors_collects_errors() {
struct FailingDestructor;
#[async_trait::async_trait]
impl PreDestruct for FailingDestructor {
async fn pre_destruct(&self) -> HookResult {
Err(Box::new(std::io::Error::new(
std::io::ErrorKind::Other,
"fail",
)))
}
}
let ctx = make_ctx();
ctx.register_destructor(Arc::new(FailingDestructor));
let result = ctx.run_destructors().await;
assert!(result.is_err());
assert_eq!(result.unwrap_err().len(), 1);
}
#[tokio::test]
async fn resolve_external_missing_returns_error() {
let ctx = make_ctx();
let result = ctx.resolve_external::<String>().await;
assert!(result.is_err());
}
#[tokio::test]
async fn resolve_external_registered_returns_value() {
let mut registry = ProviderRegistry::new();
registry.register("", DynProvider::from_value(42u32));
let ctx = ResolveContext::new(Arc::new(EmptySingletonStore), Arc::new(registry));
let v: u32 = ctx.resolve_external().await.unwrap();
assert_eq!(v, 42);
}
#[tokio::test]
async fn try_resolve_external_missing_returns_none() {
let ctx = make_ctx();
let result = ctx.try_resolve_external::<String>().await;
assert!(result.is_none());
}
#[tokio::test]
async fn try_resolve_external_registered_returns_some() {
let mut registry = ProviderRegistry::new();
registry.register("", DynProvider::from_value(99u32));
let ctx = ResolveContext::new(Arc::new(EmptySingletonStore), Arc::new(registry));
let result = ctx.try_resolve_external::<u32>().await.unwrap();
assert_eq!(result.unwrap(), 99u32);
}
#[tokio::test]
async fn resolve_external_with_token_resolves_named_provider() {
let mut registry = ProviderRegistry::new();
registry.register("primary", DynProvider::from_value(1u32));
registry.register("replica", DynProvider::from_value(2u32));
let ctx = ResolveContext::new(Arc::new(EmptySingletonStore), Arc::new(registry));
let primary: u32 = ctx.resolve_external_with_token("primary").await.unwrap();
let replica: u32 = ctx.resolve_external_with_token("replica").await.unwrap();
assert_eq!(primary, 1);
assert_eq!(replica, 2);
}
#[tokio::test]
async fn resolve_external_with_default_token_doesnt_find_named() {
let mut registry = ProviderRegistry::new();
registry.register("primary", DynProvider::from_value(1u32));
let ctx = ResolveContext::new(Arc::new(EmptySingletonStore), Arc::new(registry));
let result = ctx.resolve_external::<u32>().await;
assert!(
result.is_err(),
"default token should not resolve named provider"
);
}
#[tokio::test]
async fn try_resolve_external_with_token_returns_none_for_wrong_token() {
let mut registry = ProviderRegistry::new();
registry.register("primary", DynProvider::from_value(1u32));
let ctx = ResolveContext::new(Arc::new(EmptySingletonStore), Arc::new(registry));
let result = ctx.try_resolve_external_with_token::<u32>("replica").await;
assert!(result.is_none());
}
}