use std::any::{Any, TypeId};
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use parking_lot::RwLock;
type FactoryFn = Arc<dyn Fn(&Container) -> Box<dyn Any + Send + Sync> + Send + Sync>;
enum Entry {
Factory(FactoryFn),
SingletonFactory(FactoryFn),
Instance(Box<dyn Any + Send + Sync>),
}
pub struct Container {
entries: RwLock<HashMap<TypeId, Entry>>,
frozen: AtomicBool,
}
impl Container {
pub fn new() -> Self {
Self {
entries: RwLock::new(HashMap::new()),
frozen: AtomicBool::new(false),
}
}
pub fn freeze(&self) {
self.frozen.store(true, Ordering::SeqCst);
}
pub fn is_frozen(&self) -> bool {
self.frozen.load(Ordering::SeqCst)
}
fn assert_not_frozen(&self) {
assert!(
!self.frozen.load(Ordering::SeqCst),
"Container is frozen — cannot mutate after boot()"
);
}
pub fn singleton<T, F>(&self, factory: F)
where
T: Send + Sync + 'static,
F: Fn(&Container) -> T + Send + Sync + 'static,
{
self.assert_not_frozen();
let key = TypeId::of::<T>();
let mut entries = self.entries.write();
assert!(
!entries.contains_key(&key),
"Container: a binding for `{}` already exists",
std::any::type_name::<T>()
);
let f: FactoryFn = Arc::new(move |c| Box::new(factory(c)));
entries.insert(key, Entry::SingletonFactory(f));
}
pub fn singleton_or_replace<T, F>(&self, factory: F)
where
T: Send + Sync + 'static,
F: Fn(&Container) -> T + Send + Sync + 'static,
{
self.assert_not_frozen();
let f: FactoryFn = Arc::new(move |c| Box::new(factory(c)));
self.entries
.write()
.insert(TypeId::of::<T>(), Entry::SingletonFactory(f));
}
pub fn instance<T>(&self, value: T)
where
T: Send + Sync + 'static,
{
self.assert_not_frozen();
let key = TypeId::of::<T>();
let mut entries = self.entries.write();
assert!(
!entries.contains_key(&key),
"Container: a binding for `{}` already exists",
std::any::type_name::<T>()
);
entries.insert(key, Entry::Instance(Box::new(Arc::new(value))));
}
pub fn instance_or_replace<T>(&self, value: T)
where
T: Send + Sync + 'static,
{
self.assert_not_frozen();
self.entries.write().insert(
TypeId::of::<T>(),
Entry::Instance(Box::new(Arc::new(value))),
);
}
pub fn bind<T, F>(&self, factory: F)
where
T: Send + Sync + 'static,
F: Fn(&Container) -> T + Send + Sync + 'static,
{
self.assert_not_frozen();
let key = TypeId::of::<T>();
let mut entries = self.entries.write();
assert!(
!entries.contains_key(&key),
"Container: a binding for `{}` already exists",
std::any::type_name::<T>()
);
let f: FactoryFn = Arc::new(move |c| Box::new(factory(c)));
entries.insert(key, Entry::Factory(f));
}
pub fn bind_or_replace<T, F>(&self, factory: F)
where
T: Send + Sync + 'static,
F: Fn(&Container) -> T + Send + Sync + 'static,
{
self.assert_not_frozen();
let f: FactoryFn = Arc::new(move |c| Box::new(factory(c)));
self.entries
.write()
.insert(TypeId::of::<T>(), Entry::Factory(f));
}
pub fn resolve<T: Send + Sync + 'static>(&self) -> Option<Arc<T>> {
let key = TypeId::of::<T>();
let needs_materialise = {
let entries = self.entries.read();
matches!(entries.get(&key), Some(Entry::SingletonFactory(_)))
};
if needs_materialise {
let factory = {
let mut entries = self.entries.write();
match entries.remove(&key)? {
Entry::SingletonFactory(f) => f,
other => {
entries.insert(key, other);
return self.read_instance::<T>();
}
}
};
let raw: Box<dyn Any + Send + Sync> = factory(self);
let mut entries = self.entries.write();
match raw.downcast::<T>() {
Ok(typed) => {
let arc: Arc<T> = Arc::from(*typed);
entries.insert(key, Entry::Instance(Box::new(arc)));
}
Err(_) => return None, }
}
self.read_instance::<T>()
}
fn read_instance<T: Send + Sync + 'static>(&self) -> Option<Arc<T>> {
let entries = self.entries.read();
match entries.get(&TypeId::of::<T>())? {
Entry::Instance(boxed) => {
let arc_ref: &Arc<T> = boxed.downcast_ref::<Arc<T>>()?;
Some(Arc::clone(arc_ref))
}
_ => None,
}
}
pub fn resolve_fresh<T: Send + Sync + 'static>(&self) -> Option<T> {
let key = TypeId::of::<T>();
let f: FactoryFn = {
let entries = self.entries.read();
match entries.get(&key)? {
Entry::Factory(f) => Arc::clone(f),
_ => return None,
}
};
let instance = f(self);
instance.downcast::<T>().ok().map(|b| *b)
}
pub fn has<T: 'static>(&self) -> bool {
self.entries.read().contains_key(&TypeId::of::<T>())
}
pub fn forget<T: 'static>(&self) {
self.assert_not_frozen();
self.entries.write().remove(&TypeId::of::<T>());
}
pub fn flush(&self) {
self.assert_not_frozen();
self.entries.write().clear();
}
pub fn len(&self) -> usize {
self.entries.read().len()
}
pub fn is_empty(&self) -> bool {
self.entries.read().is_empty()
}
pub fn bind_trait<Key: 'static, Trait: ?Sized + Send + Sync + 'static>(
&self,
value: Arc<Trait>,
) {
self.assert_not_frozen();
self.entries
.write()
.insert(TypeId::of::<Key>(), Entry::Instance(Box::new(value)));
}
pub fn resolve_trait<Key: 'static, Trait: ?Sized + Send + Sync + 'static>(
&self,
) -> Option<Arc<Trait>> {
self.entries
.read()
.get(&TypeId::of::<Key>())
.and_then(|entry| match entry {
Entry::Instance(boxed) => boxed.downcast_ref::<Arc<Trait>>().cloned(),
_ => None,
})
}
}
impl Default for Container {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Mutex;
#[derive(Debug, PartialEq)]
struct Greeter {
prefix: String,
}
#[test]
fn singleton_returns_same_instance() {
let c = Container::new();
c.singleton(|_| Greeter {
prefix: "Hi".into(),
});
let a = c.resolve::<Greeter>().unwrap();
let b = c.resolve::<Greeter>().unwrap();
assert_eq!(Arc::as_ptr(&a), Arc::as_ptr(&b));
}
#[test]
fn singleton_factory_receives_container() {
let c = Container::new();
c.instance("base:".to_string());
c.singleton(|cx: &Container| {
let base: Arc<String> = cx.resolve().unwrap();
Greeter {
prefix: format!("{base}hello"),
}
});
let g = c.resolve::<Greeter>().unwrap();
assert_eq!(g.prefix, "base:hello");
}
#[test]
fn prebuilt_instance() {
let c = Container::new();
c.instance(42u32);
assert_eq!(*c.resolve::<u32>().unwrap(), 42);
}
#[test]
fn transient_factory_new_each_time() {
let c = Container::new();
let counter = Arc::new(Mutex::new(0u32));
let cnt = counter.clone();
c.bind(move |_| {
let mut n = cnt.lock().unwrap();
*n += 1;
*n
});
assert_eq!(c.resolve_fresh::<u32>().unwrap(), 1);
assert_eq!(c.resolve_fresh::<u32>().unwrap(), 2);
assert_eq!(c.resolve_fresh::<u32>().unwrap(), 3);
assert!(c.resolve::<u32>().is_none());
}
#[test]
fn transient_factory_receives_container() {
let c = Container::new();
c.instance(100u32);
c.bind(move |cx: &Container| {
let base: Arc<u32> = cx.resolve().unwrap();
format!("value={base}")
});
assert_eq!(c.resolve_fresh::<String>().unwrap(), "value=100");
}
#[test]
fn trait_object_binding() {
trait Calc: Send + Sync {
fn double(&self, x: i32) -> i32;
}
struct Doubler;
impl Calc for Doubler {
fn double(&self, x: i32) -> i32 {
x * 2
}
}
struct CalcService;
let c = Container::new();
c.bind_trait::<CalcService, dyn Calc>(Arc::new(Doubler));
let calc = c.resolve_trait::<CalcService, dyn Calc>().unwrap();
assert_eq!(calc.double(5), 10);
let calc2 = c.resolve_trait::<CalcService, dyn Calc>().unwrap();
assert_eq!(Arc::as_ptr(&calc), Arc::as_ptr(&calc2));
}
#[test]
fn has_and_forget() {
let c = Container::new();
assert!(!c.has::<u32>());
c.instance(1u32);
assert!(c.has::<u32>());
c.forget::<u32>();
assert!(!c.has::<u32>());
}
#[test]
fn flush_removes_all() {
let c = Container::new();
c.instance(1u32);
c.instance("hello");
assert_eq!(c.len(), 2);
c.flush();
assert!(c.is_empty());
}
#[test]
fn freeze_prevents_writes() {
let c = Container::new();
c.instance(42u32);
c.freeze();
assert!(c.is_frozen());
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
c.instance(99u32);
}));
assert!(result.is_err());
}
#[test]
fn resolve_works_after_freeze() {
let c = Container::new();
c.singleton(|_| Greeter {
prefix: "Hi".into(),
});
c.freeze();
let g = c.resolve::<Greeter>().unwrap();
assert_eq!(g.prefix, "Hi");
}
#[test]
fn concurrent_resolve_after_freeze() {
use std::thread;
let c = Arc::new(Container::new());
c.singleton(|_| Greeter {
prefix: "Threaded".into(),
});
c.freeze();
let handles: Vec<_> = (0..4)
.map(|_| {
let c = Arc::clone(&c);
thread::spawn(move || {
let g = c.resolve::<Greeter>().unwrap();
assert_eq!(g.prefix, "Threaded");
})
})
.collect();
for h in handles {
h.join().unwrap();
}
}
}