#![allow(clippy::pedantic)]
use std::collections::HashMap;
use std::sync::{Arc, RwLock};
use thiserror::Error;
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum ExtensionError {
#[error("extension `{name}` is already registered")]
AlreadyRegistered {
name: String,
},
#[error("extension `{name}` not found")]
NotFound {
name: String,
},
#[error("extension `{name}` has {strong_count} live references and cannot be removed")]
InUse {
name: String,
strong_count: usize,
},
#[error("extension registry lock poisoned")]
LockPoisoned,
}
pub struct ExtensionPoint<T: ?Sized> {
inner: Arc<ExtensionInner<T>>,
}
struct ExtensionInner<T: ?Sized> {
entries: RwLock<HashMap<String, Arc<T>>>,
}
impl<T: ?Sized> Default for ExtensionPoint<T> {
fn default() -> Self {
Self::new()
}
}
impl<T: ?Sized> Clone for ExtensionPoint<T> {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
}
}
}
impl<T: ?Sized> ExtensionPoint<T> {
#[must_use]
pub fn new() -> Self {
Self {
inner: Arc::new(ExtensionInner {
entries: RwLock::new(HashMap::new()),
}),
}
}
#[must_use]
pub fn len(&self) -> usize {
self.read().map(|m| m.len()).unwrap_or_default()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
#[must_use]
pub fn names(&self) -> Vec<String> {
let mut names = self
.read()
.map(|m| m.keys().cloned().collect::<Vec<_>>())
.unwrap_or_default();
names.sort_unstable();
names
}
#[must_use]
pub fn snapshot(&self) -> Vec<(String, Arc<T>)> {
self.read()
.map(|m| {
let mut entries: Vec<(String, Arc<T>)> =
m.iter().map(|(k, v)| (k.clone(), v.clone())).collect();
entries.sort_by(|a, b| a.0.cmp(&b.0));
entries
})
.unwrap_or_default()
}
fn read(
&self,
) -> Result<std::sync::RwLockReadGuard<'_, HashMap<String, Arc<T>>>, ExtensionError> {
self.inner
.entries
.read()
.map_err(|_| ExtensionError::LockPoisoned)
}
fn write(
&self,
) -> Result<std::sync::RwLockWriteGuard<'_, HashMap<String, Arc<T>>>, ExtensionError> {
self.inner
.entries
.write()
.map_err(|_| ExtensionError::LockPoisoned)
}
}
impl<T: ?Sized + Send + Sync + 'static> ExtensionPoint<T> {
pub fn register(&self, name: impl Into<String>, value: Arc<T>) -> Result<(), ExtensionError> {
let name = name.into();
let mut map = self.write()?;
if map.contains_key(&name) {
return Err(ExtensionError::AlreadyRegistered { name });
}
map.insert(name, value);
Ok(())
}
pub fn register_or_replace(&self, name: impl Into<String>, value: Arc<T>) -> Option<Arc<T>> {
let name = name.into();
self.write().ok().and_then(|mut m| m.insert(name, value))
}
pub fn unregister(&self, name: &str) -> Result<Option<Arc<T>>, ExtensionError> {
let map = self.read()?;
if let Some(existing) = map.get(name)
&& Arc::strong_count(existing) > 1
{
return Err(ExtensionError::InUse {
name: name.to_string(),
strong_count: Arc::strong_count(existing),
});
}
drop(map);
Ok(self.write()?.remove(name))
}
pub fn replace(&self, name: &str, new: Arc<T>) -> Result<Arc<T>, ExtensionError> {
let mut map = self.write()?;
let previous = map.remove(name).ok_or_else(|| ExtensionError::NotFound {
name: name.to_string(),
})?;
map.insert(name.to_string(), new);
Ok(previous)
}
#[must_use]
pub fn get(&self, name: &str) -> Option<Arc<T>> {
self.read().ok().and_then(|m| m.get(name).cloned())
}
pub fn get_required(&self, name: &str) -> Result<Arc<T>, ExtensionError> {
self.get(name).ok_or_else(|| ExtensionError::NotFound {
name: name.to_string(),
})
}
#[must_use]
pub fn contains(&self, name: &str) -> bool {
self.read().map(|m| m.contains_key(name)).unwrap_or(false)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn register_and_get_round_trip() {
let ep: ExtensionPoint<String> = ExtensionPoint::new();
ep.register("greeting", Arc::new("hello".to_string()))
.unwrap_or_else(|e| panic!("{e}"));
assert_eq!(
ep.get("greeting").map(|s| (*s).clone()),
Some("hello".to_string())
);
}
#[test]
fn register_rejects_duplicate_names() {
let ep: ExtensionPoint<String> = ExtensionPoint::new();
ep.register("a", Arc::new("first".to_string()))
.unwrap_or_else(|e| panic!("{e}"));
let err = match ep.register("a", Arc::new("second".to_string())) {
Ok(_) => panic!("expected Err, got Ok"),
Err(e) => e,
};
assert!(matches!(err, ExtensionError::AlreadyRegistered { .. }));
}
#[test]
fn register_or_replace_swallows_duplicates() {
let ep: ExtensionPoint<String> = ExtensionPoint::new();
let prev = ep.register_or_replace("a", Arc::new("first".to_string()));
assert!(prev.is_none());
let prev = ep.register_or_replace("a", Arc::new("second".to_string()));
assert_eq!(prev.map(|s| (*s).clone()), Some("first".to_string()));
assert_eq!(
ep.get("a").map(|s| (*s).clone()),
Some("second".to_string())
);
}
#[test]
fn replace_returns_previous_value() {
let ep: ExtensionPoint<String> = ExtensionPoint::new();
ep.register("a", Arc::new("v1".to_string()))
.unwrap_or_else(|e| panic!("{e}"));
let prev = ep
.replace("a", Arc::new("v2".to_string()))
.unwrap_or_else(|e| panic!("{e}"));
assert_eq!((*prev).clone(), "v1");
assert_eq!(ep.get("a").map(|s| (*s).clone()), Some("v2".to_string()));
}
#[test]
fn replace_missing_returns_not_found() {
let ep: ExtensionPoint<String> = ExtensionPoint::new();
let err = match ep.replace("missing", Arc::new("v".to_string())) {
Ok(_) => panic!("expected Err, got Ok"),
Err(e) => e,
};
assert!(matches!(err, ExtensionError::NotFound { .. }));
}
#[test]
fn unregister_refuses_in_use_entry() {
let ep: ExtensionPoint<String> = ExtensionPoint::new();
ep.register("a", Arc::new("v".to_string()))
.unwrap_or_else(|e| panic!("{e}"));
let _hold = match ep.get("a") {
Some(v) => v,
None => panic!("expected Some"),
};
let err = match ep.unregister("a") {
Ok(_) => panic!("expected Err, got Ok"),
Err(e) => e,
};
assert!(matches!(err, ExtensionError::InUse { .. }));
}
#[test]
fn unregister_drops_when_only_registry_holds() {
let ep: ExtensionPoint<String> = ExtensionPoint::new();
ep.register("a", Arc::new("v".to_string()))
.unwrap_or_else(|e| panic!("{e}"));
let removed = ep.unregister("a").unwrap_or_else(|e| panic!("{e}"));
assert_eq!(removed.map(|s| (*s).clone()), Some("v".to_string()));
assert!(ep.get("a").is_none());
}
#[test]
fn clone_shares_state() {
let a: ExtensionPoint<String> = ExtensionPoint::new();
let b = a.clone();
a.register("shared", Arc::new("x".to_string()))
.unwrap_or_else(|e| panic!("{e}"));
assert!(b.contains("shared"));
}
#[test]
fn snapshot_is_sorted_and_complete() {
let ep: ExtensionPoint<String> = ExtensionPoint::new();
ep.register("b", Arc::new("B".to_string()))
.unwrap_or_else(|e| panic!("{e}"));
ep.register("a", Arc::new("A".to_string()))
.unwrap_or_else(|e| panic!("{e}"));
let snap = ep.snapshot();
assert_eq!(snap.len(), 2);
assert_eq!(snap[0].0, "a");
assert_eq!(snap[1].0, "b");
}
}