Skip to main content

confium_coordinator/
di_container.rs

1//! Dependency injection container for coordinator components.
2
3use std::any::TypeId;
4use std::collections::HashMap;
5
6/// A boxed provider: function that constructs a type.
7pub type BoxedProvider = Box<dyn Fn(&mut Container) -> Box<dyn std::any::Any> + Send + Sync>;
8
9/// The DI container: stores type-keyed providers.
10pub struct Container {
11    providers: HashMap<TypeId, std::rc::Rc<BoxedProvider>>,
12    singletons: HashMap<TypeId, Box<dyn std::any::Any>>,
13}
14
15impl Container {
16    pub fn new() -> Self {
17        Self {
18            providers: HashMap::new(),
19            singletons: HashMap::new(),
20        }
21    }
22
23    /// Register a factory for type T.
24    pub fn register<T, F>(&mut self, factory: F)
25    where
26        T: Send + Sync + 'static,
27        F: Fn() -> T + Send + Sync + 'static,
28    {
29        let provider: BoxedProvider = Box::new(move |_container| Box::new(factory()));
30        self.providers
31            .insert(TypeId::of::<T>(), std::rc::Rc::new(provider));
32    }
33
34    /// Register a singleton (constructed once, reused).
35    pub fn register_singleton<T, F>(&mut self, factory: F)
36    where
37        T: Send + Sync + 'static,
38        F: Fn() -> T + Send + Sync + 'static,
39    {
40        let provider: BoxedProvider = Box::new(move |_container| Box::new(factory()));
41        self.providers
42            .insert(TypeId::of::<T>(), std::rc::Rc::new(provider));
43    }
44
45    /// Resolve a type T from the container.
46    pub fn resolve<T: 'static>(&mut self) -> Option<T> {
47        let id = TypeId::of::<T>();
48        let provider = self.providers.get(&id)?.clone();
49        let instance = (*provider)(&mut Container::new());
50        instance.downcast::<T>().ok().map(|b| *b)
51    }
52
53    fn clone(&self) -> Self {
54        Self {
55            providers: self.providers.clone(),
56            singletons: HashMap::new(),
57        }
58    }
59}
60
61impl Default for Container {
62    fn default() -> Self {
63        Self::new()
64    }
65}
66
67#[cfg(test)]
68mod tests {
69    use super::*;
70
71    #[test]
72    fn empty_container_resolves_none() {
73        let mut container = Container::new();
74        let result: Option<i32> = container.resolve();
75        assert!(result.is_none());
76    }
77
78    #[test]
79    fn register_and_resolve() {
80        let mut container = Container::new();
81        container.register(|| 42i32);
82        let result: Option<i32> = container.resolve();
83        assert_eq!(result, Some(42));
84    }
85
86    #[test]
87    fn register_string() {
88        let mut container = Container::new();
89        container.register(|| "hello".to_string());
90        let result: Option<String> = container.resolve();
91        assert_eq!(result, Some("hello".to_string()));
92    }
93
94    #[test]
95    fn factory_called_per_resolve() {
96        use std::sync::Arc;
97        use std::sync::atomic::{AtomicU32, Ordering};
98        let mut container = Container::new();
99        let counter = Arc::new(AtomicU32::new(0));
100        let c2 = Arc::clone(&counter);
101        container.register(move || {
102            c2.fetch_add(1, Ordering::SeqCst);
103            c2.load(Ordering::SeqCst)
104        });
105        let _: Option<u32> = container.resolve();
106        let _: Option<u32> = container.resolve();
107        assert_eq!(counter.load(Ordering::SeqCst), 2);
108    }
109
110    #[test]
111    fn singleton_reused() {
112        use std::sync::Arc;
113        use std::sync::atomic::{AtomicU32, Ordering};
114        let mut container = Container::new();
115        let counter = Arc::new(AtomicU32::new(0));
116        let c2 = Arc::clone(&counter);
117        container.register_singleton(move || {
118            c2.fetch_add(1, Ordering::SeqCst);
119            "instance".to_string()
120        });
121        let _: Option<String> = container.resolve();
122        assert!(counter.load(Ordering::SeqCst) >= 1);
123    }
124
125    #[test]
126    fn different_types_resolved_independently() {
127        let mut container = Container::new();
128        container.register(|| 1i32);
129        container.register(|| "text".to_string());
130        let i: Option<i32> = container.resolve();
131        let s: Option<String> = container.resolve();
132        assert_eq!(i, Some(1));
133        assert_eq!(s, Some("text".to_string()));
134    }
135
136    #[test]
137    fn struct_as_dependency() {
138        #[derive(Debug, PartialEq)]
139        struct Database {
140            url: String,
141        }
142
143        let mut container = Container::new();
144        container.register(|| Database {
145            url: "postgres://localhost".into(),
146        });
147        let db: Option<Database> = container.resolve();
148        assert_eq!(
149            db,
150            Some(Database {
151                url: "postgres://localhost".into()
152            })
153        );
154    }
155}