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#[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 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 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 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 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 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 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 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 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 pub fn create<T>(&self) -> Result<T>
212 where
213 T: FromModuleRef,
214 {
215 T::from_module_ref(self)
216 }
217
218 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}