Skip to main content

a3s_boot/provider/
module_ref.rs

1use super::cache::{new_provider_cache, ProviderCache};
2use super::entry::ProviderEntry;
3use super::provider_ref::ProviderRef;
4use super::resolution::{new_resolution_stack, ProviderResolutionStack};
5use super::{AnyProvider, FromModuleRef, ProviderDefinition, ProviderScope, ProviderToken};
6use crate::{BootError, Result};
7use std::collections::BTreeMap;
8use std::fmt;
9use std::sync::{Arc, RwLock};
10
11/// Runtime provider container. This is Boot's Rust equivalent of Nest's ModuleRef.
12#[derive(Clone, Default)]
13pub struct ModuleRef {
14    providers: Arc<RwLock<BTreeMap<ProviderToken, ProviderEntry>>>,
15    provider_order: Arc<RwLock<Vec<ProviderToken>>>,
16    visible_scopes: Arc<RwLock<Vec<ModuleRef>>>,
17    request_cache: Option<ProviderCache>,
18    resolution_stack: Option<ProviderResolutionStack>,
19}
20
21impl fmt::Debug for ModuleRef {
22    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
23        let len = self
24            .providers
25            .read()
26            .map(|providers| providers.len())
27            .unwrap_or(0);
28        let visible = self
29            .visible_scopes
30            .read()
31            .map(|scopes| scopes.len())
32            .unwrap_or(0);
33        f.debug_struct("ModuleRef")
34            .field("providers", &len)
35            .field("visible_scopes", &visible)
36            .finish()
37    }
38}
39
40impl ModuleRef {
41    pub fn new() -> Self {
42        Self::default()
43    }
44
45    pub fn request_scope(&self) -> Self {
46        self.with_request_cache(new_provider_cache())
47    }
48
49    pub(crate) fn with_request_cache(&self, request_cache: ProviderCache) -> Self {
50        Self {
51            providers: Arc::clone(&self.providers),
52            provider_order: Arc::clone(&self.provider_order),
53            visible_scopes: Arc::clone(&self.visible_scopes),
54            request_cache: Some(request_cache),
55            resolution_stack: self.resolution_stack.clone(),
56        }
57    }
58
59    pub(crate) fn with_resolution_stack(&self, resolution_stack: ProviderResolutionStack) -> Self {
60        Self {
61            providers: Arc::clone(&self.providers),
62            provider_order: Arc::clone(&self.provider_order),
63            visible_scopes: Arc::clone(&self.visible_scopes),
64            request_cache: self.request_cache.clone(),
65            resolution_stack: Some(resolution_stack),
66        }
67    }
68
69    pub fn register(&self, definition: ProviderDefinition) -> Result<()> {
70        let token = definition.token().clone();
71        self.validate_registration(&token, &definition)?;
72        if definition.is_async_factory() {
73            return Err(BootError::Internal(format!(
74                "async provider factory requires async registration: {token}"
75            )));
76        }
77
78        let entry = ProviderEntry::new(definition);
79        self.insert_entry(token, entry)
80    }
81
82    pub async fn register_async(&self, definition: ProviderDefinition) -> Result<()> {
83        let token = definition.token().clone();
84        self.validate_registration(&token, &definition)?;
85
86        let entry = ProviderEntry::new(definition);
87        self.insert_entry(token, entry)
88    }
89
90    pub fn insert<T>(&self, value: T) -> Result<()>
91    where
92        T: Send + Sync + 'static,
93    {
94        self.insert_arc(Arc::new(value))
95    }
96
97    pub fn insert_arc<T>(&self, value: Arc<T>) -> Result<()>
98    where
99        T: Send + Sync + 'static,
100    {
101        let token = ProviderToken::of::<T>();
102        let entry = ProviderEntry::new(ProviderDefinition::from_arc(value));
103        self.insert_entry(token, entry)
104    }
105
106    pub fn get<T>(&self) -> Result<Arc<T>>
107    where
108        T: Send + Sync + 'static,
109    {
110        self.get_token::<T>(&ProviderToken::of::<T>())
111    }
112
113    pub fn get_named<T>(&self, token: &str) -> Result<Arc<T>>
114    where
115        T: Send + Sync + 'static,
116    {
117        self.get_token::<T>(&ProviderToken::named(token))
118    }
119
120    pub fn get_optional<T>(&self) -> Result<Option<Arc<T>>>
121    where
122        T: Send + Sync + 'static,
123    {
124        self.get_optional_token::<T>(&ProviderToken::of::<T>())
125    }
126
127    pub fn get_optional_named<T>(&self, token: &str) -> Result<Option<Arc<T>>>
128    where
129        T: Send + Sync + 'static,
130    {
131        self.get_optional_token::<T>(&ProviderToken::named(token))
132    }
133
134    /// Create a lazy provider handle for a typed dependency.
135    pub fn provider_ref<T>(&self) -> ProviderRef<T>
136    where
137        T: Send + Sync + 'static,
138    {
139        ProviderRef::new(self.clone(), ProviderToken::of::<T>())
140    }
141
142    /// Create a lazy provider handle for a named dependency.
143    pub fn named_provider_ref<T>(&self, token: &str) -> ProviderRef<T>
144    where
145        T: Send + Sync + 'static,
146    {
147        ProviderRef::new(self.clone(), ProviderToken::named(token))
148    }
149
150    /// Create a lazy provider handle only when the typed dependency is visible.
151    pub fn optional_provider_ref<T>(&self) -> Result<Option<ProviderRef<T>>>
152    where
153        T: Send + Sync + 'static,
154    {
155        if self.contains_provider::<T>()? {
156            Ok(Some(self.provider_ref::<T>()))
157        } else {
158            Ok(None)
159        }
160    }
161
162    /// Create a lazy provider handle only when the named dependency is visible.
163    pub fn optional_named_provider_ref<T>(&self, token: &str) -> Result<Option<ProviderRef<T>>>
164    where
165        T: Send + Sync + 'static,
166    {
167        if self.contains_named(token)? {
168            Ok(Some(self.named_provider_ref::<T>(token)))
169        } else {
170            Ok(None)
171        }
172    }
173
174    /// Resolve a typed provider in a fresh resolution context.
175    ///
176    /// This mirrors Nest's `ModuleRef.resolve(...)`: singleton providers reuse
177    /// their application instance, while request-scoped dependencies share one
178    /// temporary context for this resolution and transient providers are rebuilt.
179    pub fn resolve<T>(&self) -> Result<Arc<T>>
180    where
181        T: Send + Sync + 'static,
182    {
183        self.resolve_token::<T>(&ProviderToken::of::<T>())
184    }
185
186    /// Resolve a named provider in a fresh resolution context.
187    pub fn resolve_named<T>(&self, token: &str) -> Result<Arc<T>>
188    where
189        T: Send + Sync + 'static,
190    {
191        self.resolve_token::<T>(&ProviderToken::named(token))
192    }
193
194    /// Resolve a typed provider in a fresh resolution context when it exists.
195    pub fn resolve_optional<T>(&self) -> Result<Option<Arc<T>>>
196    where
197        T: Send + Sync + 'static,
198    {
199        self.resolve_optional_token::<T>(&ProviderToken::of::<T>())
200    }
201
202    /// Resolve a named provider in a fresh resolution context when it exists.
203    pub fn resolve_optional_named<T>(&self, token: &str) -> Result<Option<Arc<T>>>
204    where
205        T: Send + Sync + 'static,
206    {
207        self.resolve_optional_token::<T>(&ProviderToken::named(token))
208    }
209
210    /// Create an injectable value without registering it in the provider graph.
211    pub fn create<T>(&self) -> Result<T>
212    where
213        T: FromModuleRef,
214    {
215        T::from_module_ref(self)
216    }
217
218    /// Create an injectable `Arc<T>` without registering it in the provider graph.
219    pub fn create_arc<T>(&self) -> Result<Arc<T>>
220    where
221        T: FromModuleRef,
222    {
223        Ok(Arc::new(self.create::<T>()?))
224    }
225
226    pub fn contains(&self, token: &ProviderToken) -> Result<bool> {
227        Ok(self.get_entry(token)?.is_some())
228    }
229
230    pub fn contains_provider<T>(&self) -> Result<bool>
231    where
232        T: Send + Sync + 'static,
233    {
234        self.contains(&ProviderToken::of::<T>())
235    }
236
237    pub fn contains_named(&self, token: &str) -> Result<bool> {
238        self.contains(&ProviderToken::named(token))
239    }
240
241    pub fn tokens(&self) -> Result<Vec<ProviderToken>> {
242        let mut tokens = BTreeMap::new();
243        self.collect_tokens(&mut tokens)?;
244        Ok(tokens.into_keys().collect())
245    }
246
247    fn insert_entry(&self, token: ProviderToken, entry: ProviderEntry) -> Result<()> {
248        let mut provider_order = self.write_provider_order()?;
249        let mut providers = self.write_providers()?;
250        if providers.contains_key(&token) {
251            return Err(BootError::DuplicateProvider(token.to_string()));
252        }
253        providers.insert(token.clone(), entry);
254        provider_order.push(token);
255        Ok(())
256    }
257
258    fn validate_registration(
259        &self,
260        token: &ProviderToken,
261        definition: &ProviderDefinition,
262    ) -> Result<()> {
263        if self.contains_local(token)? {
264            return Err(BootError::DuplicateProvider(token.to_string()));
265        }
266        if definition.is_async_factory() && definition.scope() != ProviderScope::Singleton {
267            return Err(BootError::Internal(format!(
268                "async provider factories require singleton scope: {token}"
269            )));
270        }
271        if definition.lifecycle().has_hooks() && definition.scope() != ProviderScope::Singleton {
272            return Err(BootError::Internal(format!(
273                "provider lifecycle hooks require singleton scope: {token}"
274            )));
275        }
276        if definition.lifecycle().has_hooks() && definition.is_alias() {
277            return Err(BootError::Internal(format!(
278                "provider aliases cannot define lifecycle hooks: {token}"
279            )));
280        }
281        Ok(())
282    }
283
284    pub(crate) fn add_visible_scope(&self, module_ref: ModuleRef) -> Result<()> {
285        self.write_visible_scopes()?.push(module_ref);
286        Ok(())
287    }
288
289    pub(crate) fn export_from(&self, module_ref: &ModuleRef, token: &ProviderToken) -> Result<()> {
290        let entry = module_ref
291            .get_entry(token)?
292            .ok_or_else(|| BootError::MissingProvider(token.to_string()))?;
293        self.insert_entry(token.clone(), entry.with_owner(module_ref.clone()))
294    }
295
296    pub(crate) fn local_tokens(&self) -> Result<Vec<ProviderToken>> {
297        Ok(self.read_provider_order()?.clone())
298    }
299
300    pub(crate) fn initialize_local_singletons(&self) -> Result<()> {
301        for entry in self.local_entries()? {
302            if entry.is_local_singleton() {
303                let resolution_stack = new_resolution_stack();
304                entry.resolve_singleton(self, &resolution_stack)?;
305            }
306        }
307        Ok(())
308    }
309
310    pub(crate) async fn initialize_local_singletons_async(&self) -> Result<()> {
311        for entry in self.local_entries()? {
312            if entry.is_local_singleton() && entry.is_async_factory() {
313                entry.seed_singleton_async(self.clone()).await?;
314            }
315        }
316
317        self.initialize_local_singletons()
318    }
319
320    pub(crate) fn initialize_local_providers(&self) -> Result<()> {
321        for entry in self.local_entries()? {
322            entry.on_module_init(self)?;
323        }
324        Ok(())
325    }
326
327    pub(crate) async fn bootstrap_local_providers(&self) -> Result<()> {
328        for entry in self.local_entries()? {
329            entry.on_application_bootstrap(self.clone()).await?;
330        }
331        Ok(())
332    }
333
334    pub(crate) async fn destroy_local_providers(&self, signal: Option<String>) -> Result<()> {
335        let mut entries = self.local_entries()?;
336        entries.reverse();
337        for entry in entries {
338            entry
339                .on_module_destroy(self.clone(), signal.clone())
340                .await?;
341        }
342        Ok(())
343    }
344
345    pub(crate) async fn before_application_shutdown_local_providers(
346        &self,
347        signal: Option<String>,
348    ) -> Result<()> {
349        let mut entries = self.local_entries()?;
350        entries.reverse();
351        for entry in entries {
352            entry
353                .before_application_shutdown(self.clone(), signal.clone())
354                .await?;
355        }
356        Ok(())
357    }
358
359    pub(crate) async fn shutdown_local_providers(&self, signal: Option<String>) -> Result<()> {
360        let mut entries = self.local_entries()?;
361        entries.reverse();
362        for entry in entries {
363            entry
364                .on_application_shutdown(self.clone(), signal.clone())
365                .await?;
366        }
367        Ok(())
368    }
369
370    pub(crate) fn get_token<T>(&self, token: &ProviderToken) -> Result<Arc<T>>
371    where
372        T: Send + Sync + 'static,
373    {
374        let value = self
375            .get_any(token)?
376            .ok_or_else(|| BootError::MissingProvider(token.to_string()))?;
377
378        Arc::downcast::<T>(value).map_err(|_| BootError::ProviderTypeMismatch(token.to_string()))
379    }
380
381    pub(crate) fn get_optional_token<T>(&self, token: &ProviderToken) -> Result<Option<Arc<T>>>
382    where
383        T: Send + Sync + 'static,
384    {
385        let value = self.get_any(token)?;
386        match value {
387            Some(value) => Arc::downcast::<T>(value)
388                .map(Some)
389                .map_err(|_| BootError::ProviderTypeMismatch(token.to_string())),
390            None => Ok(None),
391        }
392    }
393
394    pub(crate) fn resolve_token<T>(&self, token: &ProviderToken) -> Result<Arc<T>>
395    where
396        T: Send + Sync + 'static,
397    {
398        self.request_scope().get_token(token)
399    }
400
401    pub(crate) fn resolve_optional_token<T>(&self, token: &ProviderToken) -> Result<Option<Arc<T>>>
402    where
403        T: Send + Sync + 'static,
404    {
405        self.request_scope().get_optional_token(token)
406    }
407
408    fn get_any(&self, token: &ProviderToken) -> Result<Option<Arc<AnyProvider>>> {
409        let mut alias_path = Vec::new();
410        let resolution_stack = self
411            .resolution_stack
412            .clone()
413            .unwrap_or_else(new_resolution_stack);
414        self.get_any_with_request_cache_inner(
415            token,
416            self.request_cache.clone(),
417            &resolution_stack,
418            &mut alias_path,
419        )
420    }
421
422    pub(crate) fn get_any_with_request_cache_inner(
423        &self,
424        token: &ProviderToken,
425        request_cache: Option<ProviderCache>,
426        resolution_stack: &ProviderResolutionStack,
427        alias_path: &mut Vec<ProviderToken>,
428    ) -> Result<Option<Arc<AnyProvider>>> {
429        if let Some(entry) = self.read_providers()?.get(token).cloned() {
430            return entry
431                .resolve(self, request_cache, resolution_stack, alias_path)
432                .map(Some);
433        }
434
435        for scope in self.visible_scopes()? {
436            if let Some(value) = scope.get_any_with_request_cache_inner(
437                token,
438                request_cache.clone(),
439                resolution_stack,
440                alias_path,
441            )? {
442                return Ok(Some(value));
443            }
444        }
445
446        Ok(None)
447    }
448
449    fn get_entry(&self, token: &ProviderToken) -> Result<Option<ProviderEntry>> {
450        if let Some(entry) = self.read_providers()?.get(token).cloned() {
451            return Ok(Some(entry));
452        }
453
454        for scope in self.visible_scopes()? {
455            if let Some(entry) = scope.get_entry(token)? {
456                return Ok(Some(entry));
457            }
458        }
459
460        Ok(None)
461    }
462
463    fn contains_local(&self, token: &ProviderToken) -> Result<bool> {
464        Ok(self.read_providers()?.contains_key(token))
465    }
466
467    fn collect_tokens(&self, tokens: &mut BTreeMap<ProviderToken, ()>) -> Result<()> {
468        for token in self.read_providers()?.keys() {
469            tokens.insert(token.clone(), ());
470        }
471        for scope in self.visible_scopes()? {
472            scope.collect_tokens(tokens)?;
473        }
474        Ok(())
475    }
476
477    fn local_entries(&self) -> Result<Vec<ProviderEntry>> {
478        let provider_order = self.read_provider_order()?.clone();
479        let providers = self.read_providers()?;
480        let mut entries = Vec::with_capacity(provider_order.len());
481        for token in provider_order {
482            if let Some(entry) = providers.get(&token) {
483                entries.push(entry.clone());
484            }
485        }
486        Ok(entries)
487    }
488
489    fn visible_scopes(&self) -> Result<Vec<ModuleRef>> {
490        Ok(self.read_visible_scopes()?.clone())
491    }
492
493    fn read_providers(
494        &self,
495    ) -> Result<std::sync::RwLockReadGuard<'_, BTreeMap<ProviderToken, ProviderEntry>>> {
496        self.providers
497            .read()
498            .map_err(|_| BootError::Internal("provider registry lock is poisoned".to_string()))
499    }
500
501    fn write_providers(
502        &self,
503    ) -> Result<std::sync::RwLockWriteGuard<'_, BTreeMap<ProviderToken, ProviderEntry>>> {
504        self.providers
505            .write()
506            .map_err(|_| BootError::Internal("provider registry lock is poisoned".to_string()))
507    }
508
509    fn read_provider_order(&self) -> Result<std::sync::RwLockReadGuard<'_, Vec<ProviderToken>>> {
510        self.provider_order
511            .read()
512            .map_err(|_| BootError::Internal("provider order lock is poisoned".to_string()))
513    }
514
515    fn write_provider_order(&self) -> Result<std::sync::RwLockWriteGuard<'_, Vec<ProviderToken>>> {
516        self.provider_order
517            .write()
518            .map_err(|_| BootError::Internal("provider order lock is poisoned".to_string()))
519    }
520
521    fn read_visible_scopes(&self) -> Result<std::sync::RwLockReadGuard<'_, Vec<ModuleRef>>> {
522        self.visible_scopes
523            .read()
524            .map_err(|_| BootError::Internal("provider registry lock is poisoned".to_string()))
525    }
526
527    fn write_visible_scopes(&self) -> Result<std::sync::RwLockWriteGuard<'_, Vec<ModuleRef>>> {
528        self.visible_scopes
529            .write()
530            .map_err(|_| BootError::Internal("provider registry lock is poisoned".to_string()))
531    }
532}