use crate::{Inject, error::Error};
use http::{Extensions, request::Parts};
use std::{
any::{Any, TypeId},
collections::HashMap,
fmt::Debug,
hash::{BuildHasherDefault, Hasher},
sync::{Arc, OnceLock},
};
pub use factory::GenericFactory;
pub mod factory;
#[inline]
fn make_resolver_fn<T, F, Args>(resolver: F) -> ResolverFn
where
T: Send + Sync + 'static,
F: GenericFactory<Args, Output = T>,
Args: Inject,
{
Arc::new(move |c: &Container| -> Result<ArcService, Error> {
let args = Args::inject(c)?;
resolver.call(args).map(|t| Arc::new(t) as ArcService)
})
}
#[inline]
fn make_inject_resolver_fn<T>() -> ResolverFn
where
T: Inject + 'static,
{
Arc::new(move |c: &Container| -> Result<ArcService, Error> {
T::inject(c).map(|t| Arc::new(t) as ArcService)
})
}
type ResolverFn = Arc<dyn Fn(&Container) -> Result<ArcService, Error> + Send + Sync>;
type ArcService = Arc<dyn Any + Send + Sync>;
pub(crate) enum ServiceEntry {
Singleton(ArcService),
Scoped(OnceLock<Result<ArcService, Error>>, ResolverFn),
Transient(ResolverFn),
}
impl Debug for ServiceEntry {
#[inline]
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("ServiceEntry(..)")
}
}
impl ServiceEntry {
#[inline(always)]
fn singleton<T: Send + Sync + 'static>(instance: T) -> Self {
Self::Singleton(Arc::new(instance))
}
#[inline(always)]
fn scoped(resolver: ResolverFn) -> Self {
Self::Scoped(OnceLock::new(), resolver)
}
#[inline(always)]
fn transient(resolver: ResolverFn) -> Self {
Self::Transient(resolver)
}
#[inline]
fn to_scope(&self) -> Self {
match self {
Self::Singleton(service) => Self::Singleton(service.clone()),
Self::Scoped(_, r) => Self::scoped(r.clone()),
Self::Transient(r) => Self::transient(r.clone()),
}
}
}
type ServiceMap = HashMap<TypeId, ServiceEntry, BuildHasherDefault<TypeIdHasher>>;
#[derive(Default)]
struct TypeIdHasher(u64);
impl Hasher for TypeIdHasher {
#[inline]
fn finish(&self) -> u64 {
self.0
}
#[cold]
fn write(&mut self, _: &[u8]) {
unreachable!("TypeId calls write_u64");
}
#[inline]
fn write_u64(&mut self, id: u64) {
self.0 = id;
}
}
#[derive(Debug)]
pub struct ContainerBuilder {
services: ServiceMap,
}
impl Default for ContainerBuilder {
#[inline]
fn default() -> Self {
Self::new()
}
}
impl ContainerBuilder {
#[inline]
pub fn new() -> Self {
Self {
services: ServiceMap::default(),
}
}
#[inline]
pub fn build(self) -> Container {
Container {
services: Arc::new(self.services),
}
}
pub fn register_singleton<T: Send + Sync + 'static>(&mut self, instance: T) {
self.services
.insert(TypeId::of::<T>(), ServiceEntry::singleton(instance));
}
pub fn register_scoped_factory<T, F, Args>(&mut self, factory: F)
where
T: Send + Sync + 'static,
F: GenericFactory<Args, Output = T>,
Args: Inject,
{
self.services.insert(
TypeId::of::<T>(),
ServiceEntry::scoped(make_resolver_fn(factory)),
);
}
pub fn register_scoped_default<T>(&mut self)
where
T: Default + Send + Sync + 'static,
{
self.register_scoped_factory(T::default);
}
pub fn register_scoped<T: Inject + 'static>(&mut self) {
self.services.insert(
TypeId::of::<T>(),
ServiceEntry::scoped(make_inject_resolver_fn::<T>()),
);
}
pub fn register_transient_factory<T, F, Args>(&mut self, factory: F)
where
T: Send + Sync + 'static,
F: GenericFactory<Args, Output = T>,
Args: Inject,
{
self.services.insert(
TypeId::of::<T>(),
ServiceEntry::transient(make_resolver_fn(factory)),
);
}
pub fn register_transient_default<T>(&mut self)
where
T: Default + Send + Sync + 'static,
{
self.register_transient_factory(T::default);
}
pub fn register_transient<T: Inject + 'static>(&mut self) {
self.services.insert(
TypeId::of::<T>(),
ServiceEntry::transient(make_inject_resolver_fn::<T>()),
);
}
}
#[derive(Debug, Clone)]
pub struct Container {
services: Arc<ServiceMap>,
}
impl Container {
#[inline]
pub fn create_scope(&self) -> Self {
let services = self
.services
.iter()
.map(|(key, value)| (*key, value.to_scope()))
.collect::<HashMap<_, _, _>>();
Self {
services: Arc::new(services),
}
}
#[inline]
pub fn resolve<T: Send + Sync + Clone + 'static>(&self) -> Result<T, Error> {
self.resolve_shared::<T>().map(|s| s.as_ref().clone())
}
#[inline]
pub fn resolve_shared<T: Send + Sync + 'static>(&self) -> Result<Arc<T>, Error> {
match self.get_service_entry::<T>()? {
ServiceEntry::Transient(r) => r(self).and_then(|s| Self::resolve_internal(&s)),
ServiceEntry::Scoped(cell, r) => self.resolve_scoped(cell, r),
ServiceEntry::Singleton(instance) => Self::resolve_internal(instance),
}
}
#[inline]
fn get_service_entry<T: Send + Sync + 'static>(&self) -> Result<&ServiceEntry, Error> {
let type_id = TypeId::of::<T>();
self.services
.get(&type_id)
.ok_or_else(|| Error::NotRegistered(std::any::type_name::<T>()))
}
#[inline]
fn resolve_scoped<T: Send + Sync + 'static>(
&self,
cell: &OnceLock<Result<ArcService, Error>>,
resolver_fn: &ResolverFn,
) -> Result<Arc<T>, Error> {
cell.get_or_init(|| resolver_fn(self))
.as_ref()
.map_err(|err| *err)
.and_then(Self::resolve_internal)
}
#[inline]
fn resolve_internal<T: Send + Sync + 'static>(instance: &ArcService) -> Result<Arc<T>, Error> {
instance
.clone()
.downcast::<T>()
.map_err(|_| Error::ResolveFailed(std::any::type_name::<T>()))
}
}
impl<'a> TryFrom<&'a Extensions> for &'a Container {
type Error = Error;
#[inline]
fn try_from(extensions: &'a Extensions) -> Result<Self, Self::Error> {
extensions.get::<Container>().ok_or(Error::ContainerMissing)
}
}
impl TryFrom<&Extensions> for Container {
type Error = Error;
#[inline]
fn try_from(extensions: &Extensions) -> Result<Self, Self::Error> {
let res: Result<&Container, Error> = extensions.try_into();
res.cloned()
}
}
impl TryFrom<&Parts> for Container {
type Error = Error;
#[inline]
fn try_from(parts: &Parts) -> Result<Self, Self::Error> {
Container::try_from(&parts.extensions)
}
}
#[cfg(test)]
mod tests {
use super::{Container, ContainerBuilder, Error, Inject};
use http::Request;
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
trait Cache: Send + Sync {
fn get(&self, key: &str) -> Option<String>;
fn set(&self, key: &str, value: &str);
}
#[derive(Clone, Default)]
struct InMemoryCache {
inner: Arc<Mutex<HashMap<String, String>>>,
}
impl Cache for InMemoryCache {
fn get(&self, key: &str) -> Option<String> {
self.inner.lock().unwrap().get(key).cloned()
}
fn set(&self, key: &str, value: &str) {
self.inner
.lock()
.unwrap()
.insert(key.to_string(), value.to_string());
}
}
#[derive(Clone)]
struct CacheWrapper {
inner: InMemoryCache,
}
impl Inject for CacheWrapper {
fn inject(container: &Container) -> Result<Self, Error> {
let inner = container.resolve::<InMemoryCache>()?;
Ok(Self { inner })
}
}
#[test]
fn it_registers_singleton() {
let mut container = ContainerBuilder::new();
container.register_singleton(InMemoryCache::default());
let container = container.build();
let cache = container.resolve::<InMemoryCache>().unwrap();
cache.set("key", "value");
let cache = container.resolve::<InMemoryCache>().unwrap();
let key = cache.get("key").unwrap();
assert_eq!(key, "value");
}
#[test]
fn it_registers_transient() {
let mut container = ContainerBuilder::new();
container.register_transient_default::<InMemoryCache>();
let container = container.build();
let cache = container.resolve::<InMemoryCache>().unwrap();
cache.set("key", "value");
let cache = container.resolve::<InMemoryCache>().unwrap();
let key = cache.get("key");
assert!(key.is_none());
}
#[test]
fn it_registers_scoped() {
let mut container = ContainerBuilder::new();
container.register_scoped_default::<InMemoryCache>();
let container = container.build();
let cache = container.resolve::<InMemoryCache>().unwrap();
cache.set("key", "value 1");
{
let scope = container.create_scope();
let cache = scope.resolve::<InMemoryCache>().unwrap();
cache.set("key", "value 2");
let cache = scope.resolve::<InMemoryCache>().unwrap();
let key = cache.get("key").unwrap();
assert_eq!(key, "value 2");
}
{
let scope = container.create_scope();
let cache = scope.resolve::<InMemoryCache>().unwrap();
let key = cache.get("key");
assert!(key.is_none());
}
let key = cache.get("key").unwrap();
assert_eq!(key, "value 1");
}
#[test]
fn it_resolves_inner_dependencies() {
let mut container = ContainerBuilder::new();
container.register_singleton(InMemoryCache::default());
container.register_scoped::<CacheWrapper>();
let container = container.build();
{
let scope = container.create_scope();
let cache = scope.resolve::<CacheWrapper>().unwrap();
cache.inner.set("key", "value 1");
}
let cache = container.resolve::<InMemoryCache>().unwrap();
let key = cache.get("key").unwrap();
assert_eq!(key, "value 1");
}
#[test]
fn inner_scope_does_not_affect_outer() {
let mut container = ContainerBuilder::new();
container.register_scoped_default::<InMemoryCache>();
container.register_scoped::<CacheWrapper>();
let container = container.build();
{
let scope = container.create_scope();
let cache = scope.resolve::<CacheWrapper>().unwrap();
cache.inner.set("key", "value 1");
let cache = scope.resolve::<CacheWrapper>().unwrap();
cache.inner.set("key", "value 2");
}
let cache = container.resolve::<InMemoryCache>().unwrap();
let key = cache.get("key");
assert!(key.is_none())
}
#[test]
fn it_resolves_inner_scoped_dependencies() {
let mut container = ContainerBuilder::new();
container.register_scoped_default::<InMemoryCache>();
container.register_scoped::<CacheWrapper>();
let container = container.build();
let scope = container.create_scope();
let cache = scope.resolve::<CacheWrapper>().unwrap();
cache.inner.set("key1", "value 1");
let cache = scope.resolve::<CacheWrapper>().unwrap();
cache.inner.set("key2", "value 2");
let cache = scope.resolve::<CacheWrapper>().unwrap();
assert_eq!(cache.inner.get("key1").unwrap(), "value 1");
assert_eq!(cache.inner.get("key2").unwrap(), "value 2");
}
#[test]
fn it_extracts_from_parts() {
let mut container = ContainerBuilder::new();
container.register_singleton(InMemoryCache::default());
let container = container.build();
let mut req = Request::get("/").body(()).unwrap();
req.extensions_mut().insert(container.create_scope());
let (parts, _) = req.into_parts();
let container = Container::try_from(&parts);
assert!(container.is_ok());
}
#[test]
fn it_returns_error_when_resolve_unregistered() {
let container = ContainerBuilder::new().build();
let cache = container.resolve::<CacheWrapper>();
assert!(cache.is_err());
}
#[test]
fn it_returns_error_when_resolve_unregistered_from_scope() {
let container = ContainerBuilder::new().build().create_scope();
let cache = container.resolve::<CacheWrapper>();
assert!(cache.is_err());
}
}