Skip to main content

nidus_core/provider/
mod.rs

1//! Provider registration primitives.
2
3use std::{
4    any::{Any, TypeId},
5    panic::{AssertUnwindSafe, catch_unwind},
6    sync::{Arc, Condvar, Mutex, MutexGuard, OnceLock},
7};
8
9use crate::{Container, NidusError, RequestScope, Result, resolution};
10
11/// Provider creation and reuse strategy.
12#[derive(Clone, Copy, Debug, Eq, PartialEq)]
13pub enum ProviderLifetime {
14    /// Create once and reuse for all resolutions.
15    Singleton,
16    /// Create a fresh value on every resolution.
17    Transient,
18    /// Create per request when request scopes are enabled.
19    Request,
20}
21
22/// Marker trait for injectable provider values.
23pub trait Provider: Send + Sync + 'static {}
24
25impl<T> Provider for T where T: Send + Sync + 'static {}
26
27type ErasedProvider = dyn Any + Send + Sync;
28type ProviderFactory = dyn Fn(&Container) -> Result<Arc<ErasedProvider>> + Send + Sync;
29type RequestProviderFactory =
30    dyn for<'scope> Fn(&RequestScope<'scope>) -> Result<Arc<ErasedProvider>> + Send + Sync;
31
32/// A typed provider registration stored by the container.
33pub struct ProviderEntry {
34    type_id: TypeId,
35    type_name: &'static str,
36    lifetime: ProviderLifetime,
37    factory: Arc<ProviderFactory>,
38    request_factory: Option<Arc<RequestProviderFactory>>,
39    singleton: Mutex<SingletonState>,
40    singleton_ready: Condvar,
41    // Lock-free read path for constructed singletons. Set exactly once, after a
42    // factory succeeds; failed or panicking factories leave it empty so the
43    // `singleton` state machine keeps its retry semantics.
44    singleton_cache: OnceLock<Arc<ErasedProvider>>,
45}
46
47enum SingletonState {
48    Empty,
49    Initializing,
50    Ready,
51}
52
53impl ProviderEntry {
54    /// Creates a provider entry from an erased factory.
55    pub fn new(
56        type_id: TypeId,
57        type_name: &'static str,
58        lifetime: ProviderLifetime,
59        factory: Arc<ProviderFactory>,
60    ) -> Self {
61        Self {
62            type_id,
63            type_name,
64            lifetime,
65            factory,
66            request_factory: None,
67            singleton: Mutex::new(SingletonState::Empty),
68            singleton_ready: Condvar::new(),
69            singleton_cache: OnceLock::new(),
70        }
71    }
72
73    /// Creates a request-scoped provider entry from an erased request-scope factory.
74    pub fn new_request_scoped(
75        type_id: TypeId,
76        type_name: &'static str,
77        factory: Arc<ProviderFactory>,
78        request_factory: Arc<RequestProviderFactory>,
79    ) -> Self {
80        Self {
81            type_id,
82            type_name,
83            lifetime: ProviderLifetime::Request,
84            factory,
85            request_factory: Some(request_factory),
86            singleton: Mutex::new(SingletonState::Empty),
87            singleton_ready: Condvar::new(),
88            singleton_cache: OnceLock::new(),
89        }
90    }
91
92    /// Returns the registered provider type name.
93    pub fn type_name(&self) -> &'static str {
94        self.type_name
95    }
96
97    /// Returns the configured provider lifetime.
98    pub fn lifetime(&self) -> ProviderLifetime {
99        self.lifetime
100    }
101
102    pub(crate) fn resolve_erased(&self, container: &Container) -> Result<Arc<ErasedProvider>> {
103        match self.lifetime {
104            ProviderLifetime::Singleton => self.resolve_singleton(container),
105            ProviderLifetime::Transient | ProviderLifetime::Request => {
106                self.create_erased(container)
107            }
108        }
109    }
110
111    pub(crate) fn resolve_erased_in_scope(
112        &self,
113        scope: &RequestScope<'_>,
114    ) -> Result<Arc<ErasedProvider>> {
115        match self.lifetime {
116            ProviderLifetime::Request => self.create_erased_in_scope(scope),
117            ProviderLifetime::Singleton | ProviderLifetime::Transient => {
118                self.resolve_erased(scope.container())
119            }
120        }
121    }
122
123    fn create_erased(&self, container: &Container) -> Result<Arc<ErasedProvider>> {
124        (self.factory)(container).map_err(|source| NidusError::ProviderFactory {
125            type_name: self.type_name,
126            source: Box::new(source),
127        })
128    }
129
130    fn resolve_singleton(&self, container: &Container) -> Result<Arc<ErasedProvider>> {
131        if let Some(instance) = self.singleton_cache.get() {
132            return Ok(Arc::clone(instance));
133        }
134        loop {
135            let mut singleton = lock_unpoisoned(&self.singleton);
136            match &*singleton {
137                SingletonState::Ready => {
138                    let instance = self
139                        .singleton_cache
140                        .get()
141                        .expect("ready singleton must be present in the cache");
142                    return Ok(Arc::clone(instance));
143                }
144                SingletonState::Initializing => {
145                    if resolution::is_active(self.type_id) {
146                        return Err(NidusError::CircularProviderResolution {
147                            type_name: self.type_name,
148                        });
149                    }
150                    drop(wait_unpoisoned(&self.singleton_ready, singleton));
151                }
152                SingletonState::Empty => {
153                    let _guard = resolution::enter(self.type_id, self.type_name)?;
154                    *singleton = SingletonState::Initializing;
155                    drop(singleton);
156
157                    let instance =
158                        match catch_unwind(AssertUnwindSafe(|| self.create_erased(container))) {
159                            Ok(outcome) => outcome,
160                            Err(panic_payload) => {
161                                let mut singleton = lock_unpoisoned(&self.singleton);
162                                *singleton = SingletonState::Empty;
163                                self.singleton_ready.notify_all();
164                                drop(singleton);
165                                std::panic::resume_unwind(panic_payload);
166                            }
167                        };
168                    let mut singleton = lock_unpoisoned(&self.singleton);
169                    match instance {
170                        Ok(instance) => {
171                            self.singleton_cache.get_or_init(|| Arc::clone(&instance));
172                            *singleton = SingletonState::Ready;
173                            self.singleton_ready.notify_all();
174                            return Ok(instance);
175                        }
176                        Err(error) => {
177                            *singleton = SingletonState::Empty;
178                            self.singleton_ready.notify_all();
179                            return Err(error);
180                        }
181                    }
182                }
183            }
184        }
185    }
186
187    fn create_erased_in_scope(&self, scope: &RequestScope<'_>) -> Result<Arc<ErasedProvider>> {
188        if let Some(factory) = &self.request_factory {
189            factory(scope).map_err(|source| NidusError::ProviderFactory {
190                type_name: self.type_name,
191                source: Box::new(source),
192            })
193        } else {
194            self.create_erased(scope.container())
195        }
196    }
197}
198
199fn lock_unpoisoned<T>(mutex: &Mutex<T>) -> MutexGuard<'_, T> {
200    mutex
201        .lock()
202        .unwrap_or_else(|poisoned| poisoned.into_inner())
203}
204
205fn wait_unpoisoned<'a, T>(condvar: &Condvar, guard: MutexGuard<'a, T>) -> MutexGuard<'a, T> {
206    condvar
207        .wait(guard)
208        .unwrap_or_else(|poisoned| poisoned.into_inner())
209}
210
211#[cfg(test)]
212mod tests {
213    use std::{
214        any::{Any, type_name},
215        sync::Arc,
216        thread,
217    };
218
219    use super::{ProviderEntry, ProviderLifetime};
220    use crate::Container;
221
222    #[test]
223    fn singleton_provider_reuses_the_constructed_instance() {
224        let provider = ProviderEntry::new(
225            std::any::TypeId::of::<String>(),
226            type_name::<String>(),
227            ProviderLifetime::Singleton,
228            Arc::new(|_container| Ok(Arc::new("ready".to_owned()) as Arc<dyn Any + Send + Sync>)),
229        );
230        let container = Container::new();
231
232        let first = provider.resolve_erased(&container).unwrap();
233        let second = provider.resolve_erased(&container).unwrap();
234        assert!(Arc::ptr_eq(&first, &second));
235        assert_eq!(
236            Arc::strong_count(&first),
237            3,
238            "the cache and two callers should be the only strong references"
239        );
240    }
241
242    #[test]
243    fn singleton_provider_retries_after_factory_error() {
244        use std::sync::atomic::{AtomicBool, Ordering};
245
246        let failed_once = Arc::new(AtomicBool::new(false));
247        let provider = ProviderEntry::new(
248            std::any::TypeId::of::<String>(),
249            type_name::<String>(),
250            ProviderLifetime::Singleton,
251            Arc::new({
252                let failed_once = Arc::clone(&failed_once);
253                move |_container| {
254                    if failed_once.swap(true, Ordering::SeqCst) {
255                        Ok(Arc::new("recovered".to_owned()) as Arc<dyn Any + Send + Sync>)
256                    } else {
257                        Err(crate::NidusError::MissingProvider {
258                            type_name: "transient failure",
259                        })
260                    }
261                }
262            }),
263        );
264        let container = Container::new();
265
266        assert!(provider.resolve_erased(&container).is_err());
267        let value = provider
268            .resolve_erased(&container)
269            .unwrap()
270            .downcast::<String>()
271            .unwrap();
272        assert_eq!(&*value, "recovered");
273    }
274
275    #[test]
276    fn singleton_provider_recovers_from_poisoned_cache() {
277        let provider = Arc::new(ProviderEntry::new(
278            std::any::TypeId::of::<String>(),
279            type_name::<String>(),
280            ProviderLifetime::Singleton,
281            Arc::new(|_container| Ok(Arc::new("ready".to_owned()) as Arc<dyn Any + Send + Sync>)),
282        ));
283        let poisoned_provider = Arc::clone(&provider);
284
285        let panic = thread::spawn(move || {
286            let _singleton = poisoned_provider.singleton.lock().unwrap();
287            panic!("poison singleton cache");
288        });
289        assert!(panic.join().is_err());
290
291        let value = provider
292            .resolve_erased(&Container::new())
293            .unwrap()
294            .downcast::<String>()
295            .unwrap();
296        assert_eq!(&*value, "ready");
297    }
298}