1#![allow(clippy::pedantic)]
38use std::collections::HashMap;
39use std::sync::{Arc, RwLock};
40
41use thiserror::Error;
42
43#[derive(Debug, Error)]
45#[non_exhaustive]
46pub enum ExtensionError {
47 #[error("extension `{name}` is already registered")]
49 AlreadyRegistered {
50 name: String,
52 },
53 #[error("extension `{name}` not found")]
55 NotFound {
56 name: String,
58 },
59 #[error("extension `{name}` has {strong_count} live references and cannot be removed")]
62 InUse {
63 name: String,
65 strong_count: usize,
67 },
68 #[error("extension registry lock poisoned")]
70 LockPoisoned,
71}
72
73pub struct ExtensionPoint<T: ?Sized> {
77 inner: Arc<ExtensionInner<T>>,
78}
79
80struct ExtensionInner<T: ?Sized> {
81 entries: RwLock<HashMap<String, Arc<T>>>,
82}
83
84impl<T: ?Sized> Default for ExtensionPoint<T> {
85 fn default() -> Self {
86 Self::new()
87 }
88}
89
90impl<T: ?Sized> Clone for ExtensionPoint<T> {
91 fn clone(&self) -> Self {
92 Self {
93 inner: self.inner.clone(),
94 }
95 }
96}
97
98impl<T: ?Sized> ExtensionPoint<T> {
99 #[must_use]
101 pub fn new() -> Self {
102 Self {
103 inner: Arc::new(ExtensionInner {
104 entries: RwLock::new(HashMap::new()),
105 }),
106 }
107 }
108
109 #[must_use]
111 pub fn len(&self) -> usize {
112 self.read().map(|m| m.len()).unwrap_or_default()
113 }
114
115 #[must_use]
117 pub fn is_empty(&self) -> bool {
118 self.len() == 0
119 }
120
121 #[must_use]
124 pub fn names(&self) -> Vec<String> {
125 let mut names = self
126 .read()
127 .map(|m| m.keys().cloned().collect::<Vec<_>>())
128 .unwrap_or_default();
129 names.sort_unstable();
130 names
131 }
132
133 #[must_use]
135 pub fn snapshot(&self) -> Vec<(String, Arc<T>)> {
136 self.read()
137 .map(|m| {
138 let mut entries: Vec<(String, Arc<T>)> =
139 m.iter().map(|(k, v)| (k.clone(), v.clone())).collect();
140 entries.sort_by(|a, b| a.0.cmp(&b.0));
141 entries
142 })
143 .unwrap_or_default()
144 }
145
146 fn read(
147 &self,
148 ) -> Result<std::sync::RwLockReadGuard<'_, HashMap<String, Arc<T>>>, ExtensionError> {
149 self.inner
150 .entries
151 .read()
152 .map_err(|_| ExtensionError::LockPoisoned)
153 }
154
155 fn write(
156 &self,
157 ) -> Result<std::sync::RwLockWriteGuard<'_, HashMap<String, Arc<T>>>, ExtensionError> {
158 self.inner
159 .entries
160 .write()
161 .map_err(|_| ExtensionError::LockPoisoned)
162 }
163}
164
165impl<T: ?Sized + Send + Sync + 'static> ExtensionPoint<T> {
166 pub fn register(&self, name: impl Into<String>, value: Arc<T>) -> Result<(), ExtensionError> {
173 let name = name.into();
174 let mut map = self.write()?;
175 if map.contains_key(&name) {
176 return Err(ExtensionError::AlreadyRegistered { name });
177 }
178 map.insert(name, value);
179 Ok(())
180 }
181
182 pub fn register_or_replace(&self, name: impl Into<String>, value: Arc<T>) -> Option<Arc<T>> {
185 let name = name.into();
186 self.write().ok().and_then(|mut m| m.insert(name, value))
187 }
188
189 pub fn unregister(&self, name: &str) -> Result<Option<Arc<T>>, ExtensionError> {
195 let map = self.read()?;
196 if let Some(existing) = map.get(name)
197 && Arc::strong_count(existing) > 1
198 {
199 return Err(ExtensionError::InUse {
200 name: name.to_string(),
201 strong_count: Arc::strong_count(existing),
202 });
203 }
204 drop(map);
205 Ok(self.write()?.remove(name))
206 }
207
208 pub fn replace(&self, name: &str, new: Arc<T>) -> Result<Arc<T>, ExtensionError> {
218 let mut map = self.write()?;
219 let previous = map.remove(name).ok_or_else(|| ExtensionError::NotFound {
220 name: name.to_string(),
221 })?;
222 map.insert(name.to_string(), new);
223 Ok(previous)
224 }
225
226 #[must_use]
228 pub fn get(&self, name: &str) -> Option<Arc<T>> {
229 self.read().ok().and_then(|m| m.get(name).cloned())
230 }
231
232 pub fn get_required(&self, name: &str) -> Result<Arc<T>, ExtensionError> {
235 self.get(name).ok_or_else(|| ExtensionError::NotFound {
236 name: name.to_string(),
237 })
238 }
239
240 #[must_use]
242 pub fn contains(&self, name: &str) -> bool {
243 self.read().map(|m| m.contains_key(name)).unwrap_or(false)
244 }
245}
246
247#[cfg(test)]
248mod tests {
249 use super::*;
250
251 #[test]
252 fn register_and_get_round_trip() {
253 let ep: ExtensionPoint<String> = ExtensionPoint::new();
254 ep.register("greeting", Arc::new("hello".to_string()))
255 .unwrap_or_else(|e| panic!("{e}"));
256 assert_eq!(
257 ep.get("greeting").map(|s| (*s).clone()),
258 Some("hello".to_string())
259 );
260 }
261
262 #[test]
263 fn register_rejects_duplicate_names() {
264 let ep: ExtensionPoint<String> = ExtensionPoint::new();
265 ep.register("a", Arc::new("first".to_string()))
266 .unwrap_or_else(|e| panic!("{e}"));
267 let err = match ep.register("a", Arc::new("second".to_string())) {
268 Ok(_) => panic!("expected Err, got Ok"),
269 Err(e) => e,
270 };
271 assert!(matches!(err, ExtensionError::AlreadyRegistered { .. }));
272 }
273
274 #[test]
275 fn register_or_replace_swallows_duplicates() {
276 let ep: ExtensionPoint<String> = ExtensionPoint::new();
277 let prev = ep.register_or_replace("a", Arc::new("first".to_string()));
278 assert!(prev.is_none());
279 let prev = ep.register_or_replace("a", Arc::new("second".to_string()));
280 assert_eq!(prev.map(|s| (*s).clone()), Some("first".to_string()));
281 assert_eq!(
282 ep.get("a").map(|s| (*s).clone()),
283 Some("second".to_string())
284 );
285 }
286
287 #[test]
288 fn replace_returns_previous_value() {
289 let ep: ExtensionPoint<String> = ExtensionPoint::new();
290 ep.register("a", Arc::new("v1".to_string()))
291 .unwrap_or_else(|e| panic!("{e}"));
292 let prev = ep
293 .replace("a", Arc::new("v2".to_string()))
294 .unwrap_or_else(|e| panic!("{e}"));
295 assert_eq!((*prev).clone(), "v1");
296 assert_eq!(ep.get("a").map(|s| (*s).clone()), Some("v2".to_string()));
297 }
298
299 #[test]
300 fn replace_missing_returns_not_found() {
301 let ep: ExtensionPoint<String> = ExtensionPoint::new();
302 let err = match ep.replace("missing", Arc::new("v".to_string())) {
303 Ok(_) => panic!("expected Err, got Ok"),
304 Err(e) => e,
305 };
306 assert!(matches!(err, ExtensionError::NotFound { .. }));
307 }
308
309 #[test]
310 fn unregister_refuses_in_use_entry() {
311 let ep: ExtensionPoint<String> = ExtensionPoint::new();
312 ep.register("a", Arc::new("v".to_string()))
313 .unwrap_or_else(|e| panic!("{e}"));
314 let _hold = match ep.get("a") {
315 Some(v) => v,
316 None => panic!("expected Some"),
317 };
318 let err = match ep.unregister("a") {
319 Ok(_) => panic!("expected Err, got Ok"),
320 Err(e) => e,
321 };
322 assert!(matches!(err, ExtensionError::InUse { .. }));
323 }
324
325 #[test]
326 fn unregister_drops_when_only_registry_holds() {
327 let ep: ExtensionPoint<String> = ExtensionPoint::new();
328 ep.register("a", Arc::new("v".to_string()))
329 .unwrap_or_else(|e| panic!("{e}"));
330 let removed = ep.unregister("a").unwrap_or_else(|e| panic!("{e}"));
331 assert_eq!(removed.map(|s| (*s).clone()), Some("v".to_string()));
332 assert!(ep.get("a").is_none());
333 }
334
335 #[test]
336 fn clone_shares_state() {
337 let a: ExtensionPoint<String> = ExtensionPoint::new();
338 let b = a.clone();
339 a.register("shared", Arc::new("x".to_string()))
340 .unwrap_or_else(|e| panic!("{e}"));
341 assert!(b.contains("shared"));
342 }
343
344 #[test]
345 fn snapshot_is_sorted_and_complete() {
346 let ep: ExtensionPoint<String> = ExtensionPoint::new();
347 ep.register("b", Arc::new("B".to_string()))
348 .unwrap_or_else(|e| panic!("{e}"));
349 ep.register("a", Arc::new("A".to_string()))
350 .unwrap_or_else(|e| panic!("{e}"));
351 let snap = ep.snapshot();
352 assert_eq!(snap.len(), 2);
353 assert_eq!(snap[0].0, "a");
354 assert_eq!(snap[1].0, "b");
355 }
356}