1use 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#[derive(Clone, Copy, Debug, Eq, PartialEq)]
13pub enum ProviderLifetime {
14 Singleton,
16 Transient,
18 Request,
20}
21
22pub 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
32pub 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 singleton_cache: OnceLock<Arc<ErasedProvider>>,
45}
46
47enum SingletonState {
48 Empty,
49 Initializing,
50 Ready,
51}
52
53impl ProviderEntry {
54 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 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 pub fn type_name(&self) -> &'static str {
94 self.type_name
95 }
96
97 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}