a3s_boot/provider/
module_ref.rs1use super::{AnyProvider, ProviderDefinition, ProviderToken};
2use crate::{BootError, Result};
3use std::collections::BTreeMap;
4use std::fmt;
5use std::sync::{Arc, RwLock};
6
7#[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}