Skip to main content

a3s_boot/provider/
module_ref.rs

1use super::{AnyProvider, ProviderDefinition, ProviderToken};
2use crate::{BootError, Result};
3use std::collections::BTreeMap;
4use std::fmt;
5use std::sync::{Arc, RwLock};
6
7/// Runtime provider container. This is Boot's Rust equivalent of Nest's ModuleRef.
8#[derive(Clone, Default)]
9pub struct ModuleRef {
10    providers: Arc<RwLock<BTreeMap<ProviderToken, Arc<AnyProvider>>>>,
11}
12
13impl fmt::Debug for ModuleRef {
14    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
15        let len = self
16            .providers
17            .read()
18            .map(|providers| providers.len())
19            .unwrap_or(0);
20        f.debug_struct("ModuleRef")
21            .field("providers", &len)
22            .finish()
23    }
24}
25
26impl ModuleRef {
27    pub fn new() -> Self {
28        Self::default()
29    }
30
31    pub fn register(&self, definition: ProviderDefinition) -> Result<()> {
32        let token = definition.token().clone();
33        if self.read_providers()?.contains_key(&token) {
34            return Err(BootError::DuplicateProvider(token.to_string()));
35        }
36        let value = definition.build(self)?;
37        self.insert_any(token, value)
38    }
39
40    pub fn insert<T>(&self, value: T) -> Result<()>
41    where
42        T: Send + Sync + 'static,
43    {
44        self.insert_arc(Arc::new(value))
45    }
46
47    pub fn insert_arc<T>(&self, value: Arc<T>) -> Result<()>
48    where
49        T: Send + Sync + 'static,
50    {
51        self.insert_any(ProviderToken::of::<T>(), value as Arc<AnyProvider>)
52    }
53
54    pub fn get<T>(&self) -> Result<Arc<T>>
55    where
56        T: Send + Sync + 'static,
57    {
58        self.get_token::<T>(&ProviderToken::of::<T>())
59    }
60
61    pub fn get_named<T>(&self, token: &str) -> Result<Arc<T>>
62    where
63        T: Send + Sync + 'static,
64    {
65        self.get_token::<T>(&ProviderToken::named(token))
66    }
67
68    pub fn get_optional<T>(&self) -> Result<Option<Arc<T>>>
69    where
70        T: Send + Sync + 'static,
71    {
72        self.get_optional_token::<T>(&ProviderToken::of::<T>())
73    }
74
75    pub fn get_optional_named<T>(&self, token: &str) -> Result<Option<Arc<T>>>
76    where
77        T: Send + Sync + 'static,
78    {
79        self.get_optional_token::<T>(&ProviderToken::named(token))
80    }
81
82    pub fn contains(&self, token: &ProviderToken) -> Result<bool> {
83        Ok(self.read_providers()?.contains_key(token))
84    }
85
86    pub fn contains_provider<T>(&self) -> Result<bool>
87    where
88        T: Send + Sync + 'static,
89    {
90        self.contains(&ProviderToken::of::<T>())
91    }
92
93    pub fn contains_named(&self, token: &str) -> Result<bool> {
94        self.contains(&ProviderToken::named(token))
95    }
96
97    pub fn tokens(&self) -> Result<Vec<ProviderToken>> {
98        Ok(self.read_providers()?.keys().cloned().collect())
99    }
100
101    fn insert_any(&self, token: ProviderToken, value: Arc<AnyProvider>) -> Result<()> {
102        let mut providers = self.write_providers()?;
103        if providers.contains_key(&token) {
104            return Err(BootError::DuplicateProvider(token.to_string()));
105        }
106        providers.insert(token, value);
107        Ok(())
108    }
109
110    fn get_token<T>(&self, token: &ProviderToken) -> Result<Arc<T>>
111    where
112        T: Send + Sync + 'static,
113    {
114        let value = self
115            .read_providers()?
116            .get(token)
117            .cloned()
118            .ok_or_else(|| BootError::MissingProvider(token.to_string()))?;
119
120        Arc::downcast::<T>(value).map_err(|_| BootError::ProviderTypeMismatch(token.to_string()))
121    }
122
123    fn get_optional_token<T>(&self, token: &ProviderToken) -> Result<Option<Arc<T>>>
124    where
125        T: Send + Sync + 'static,
126    {
127        let value = self.read_providers()?.get(token).cloned();
128        match value {
129            Some(value) => Arc::downcast::<T>(value)
130                .map(Some)
131                .map_err(|_| BootError::ProviderTypeMismatch(token.to_string())),
132            None => Ok(None),
133        }
134    }
135
136    fn read_providers(
137        &self,
138    ) -> Result<std::sync::RwLockReadGuard<'_, BTreeMap<ProviderToken, Arc<AnyProvider>>>> {
139        self.providers
140            .read()
141            .map_err(|_| BootError::Internal("provider registry lock is poisoned".to_string()))
142    }
143
144    fn write_providers(
145        &self,
146    ) -> Result<std::sync::RwLockWriteGuard<'_, BTreeMap<ProviderToken, Arc<AnyProvider>>>> {
147        self.providers
148            .write()
149            .map_err(|_| BootError::Internal("provider registry lock is poisoned".to_string()))
150    }
151}