use std::any::{Any, TypeId};
use std::collections::HashMap;
use std::fmt;
use std::sync::Arc;
#[derive(Clone, Default)]
pub struct ThemeExtensions {
map: HashMap<TypeId, Arc<dyn Any + Send + Sync>>,
}
impl ThemeExtensions {
pub fn new() -> Self {
Self::default()
}
pub fn get<T: Any + Send + Sync>(&self) -> Option<&T> {
self.map
.get(&TypeId::of::<T>())
.and_then(|a| a.downcast_ref::<T>())
}
pub fn insert<T: Any + Send + Sync>(&mut self, value: T) {
self.map.insert(TypeId::of::<T>(), Arc::new(value));
}
pub fn remove<T: Any + Send + Sync>(&mut self) -> bool {
self.map.remove(&TypeId::of::<T>()).is_some()
}
pub fn contains<T: Any + Send + Sync>(&self) -> bool {
self.map.contains_key(&TypeId::of::<T>())
}
pub fn len(&self) -> usize {
self.map.len()
}
pub fn is_empty(&self) -> bool {
self.map.is_empty()
}
}
impl fmt::Debug for ThemeExtensions {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "ThemeExtensions({} entries)", self.map.len())
}
}
impl PartialEq for ThemeExtensions {
fn eq(&self, other: &Self) -> bool {
if self.map.len() != other.map.len() {
return false;
}
self.map.keys().all(|k| other.map.contains_key(k))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Debug, Clone, PartialEq)]
struct Marker(u32);
#[derive(Debug, Clone, PartialEq)]
struct Other(&'static str);
#[test]
fn insert_and_get_round_trip() {
let mut ext = ThemeExtensions::new();
ext.insert(Marker(7));
assert_eq!(ext.get::<Marker>(), Some(&Marker(7)));
assert!(ext.contains::<Marker>());
assert_eq!(ext.len(), 1);
}
#[test]
fn distinct_types_share_the_registry() {
let mut ext = ThemeExtensions::new();
ext.insert(Marker(1));
ext.insert(Other("hi"));
assert_eq!(ext.len(), 2);
assert_eq!(ext.get::<Marker>(), Some(&Marker(1)));
assert_eq!(ext.get::<Other>(), Some(&Other("hi")));
}
#[test]
fn remove_drops_the_slot() {
let mut ext = ThemeExtensions::new();
ext.insert(Marker(1));
assert!(ext.remove::<Marker>());
assert!(!ext.contains::<Marker>());
assert!(!ext.remove::<Marker>());
}
#[test]
fn debug_format_is_short() {
let mut ext = ThemeExtensions::new();
ext.insert(Marker(1));
let s = format!("{ext:?}");
assert!(s.contains("1 entries"), "got: {s}");
}
#[test]
fn equality_ignores_inner_values() {
let mut a = ThemeExtensions::new();
let mut b = ThemeExtensions::new();
a.insert(Marker(1));
b.insert(Marker(2));
assert_eq!(a, b);
}
}