use std::any::{Any, TypeId};
use rustc_hash::FxHashMap;
use crate::correlate::TimeBucketedCounter;
#[derive(Default)]
pub struct StateMap {
by_type: FxHashMap<TypeId, Box<dyn Any + Send>>,
}
impl StateMap {
pub fn get_or_init_mut<T: Default + Send + 'static>(&mut self) -> &mut T {
let id = TypeId::of::<T>();
self.by_type
.entry(id)
.or_insert_with(|| Box::<T>::default())
.downcast_mut::<T>()
.expect("StateMap invariant: TypeId keys to its own T")
}
pub fn get_or_init_with<T, F>(&mut self, factory: F) -> &mut T
where
T: Send + 'static,
F: FnOnce() -> T,
{
let id = TypeId::of::<T>();
self.by_type
.entry(id)
.or_insert_with(|| Box::new(factory()))
.downcast_mut::<T>()
.expect("StateMap invariant: TypeId keys to its own T")
}
pub fn insert<T: Send + 'static>(&mut self, value: T) {
self.by_type.insert(TypeId::of::<T>(), Box::new(value));
}
pub fn get_mut<T: 'static>(&mut self) -> Option<&mut T> {
self.by_type
.get_mut(&TypeId::of::<T>())
.and_then(|b| b.downcast_mut::<T>())
}
pub fn get<T: 'static>(&self) -> Option<&T> {
self.by_type
.get(&TypeId::of::<T>())
.and_then(|b| b.downcast_ref::<T>())
}
pub fn take_dyn(&mut self, type_id: TypeId) -> Option<Box<dyn Any + Send>> {
self.by_type.remove(&type_id)
}
pub fn len(&self) -> usize {
self.by_type.len()
}
pub fn is_empty(&self) -> bool {
self.by_type.is_empty()
}
}
#[derive(Default)]
pub struct FlowStateRegistry {
by_type: FxHashMap<TypeId, Box<dyn Any + Send>>,
}
impl FlowStateRegistry {
pub fn register<T>(&mut self, idle_timeout: std::time::Duration)
where
T: Default + Send + 'static,
{
let map: flowscope::correlate::FlowStateMap<T, flowscope::extract::FiveTupleKey> =
flowscope::correlate::FlowStateMap::new(idle_timeout);
self.by_type.insert(TypeId::of::<T>(), Box::new(map));
}
pub fn get_mut<T>(
&mut self,
) -> Option<&mut flowscope::correlate::FlowStateMap<T, flowscope::extract::FiveTupleKey>>
where
T: Default + Send + 'static,
{
self.by_type
.get_mut(&TypeId::of::<T>())
.and_then(|b| b.downcast_mut())
}
pub fn is_empty(&self) -> bool {
self.by_type.is_empty()
}
}
#[derive(Default)]
pub struct CounterRegistry {
by_type: FxHashMap<TypeId, Box<dyn Any + Send>>,
registered_type_names: Vec<&'static str>,
}
impl CounterRegistry {
pub fn register<K>(&mut self, counter: TimeBucketedCounter<K>)
where
K: std::hash::Hash + Eq + Clone + Send + 'static,
{
let id = TypeId::of::<K>();
let name = std::any::type_name::<K>();
if !self.registered_type_names.contains(&name) {
self.registered_type_names.push(name);
}
self.by_type.insert(id, Box::new(counter));
}
pub fn registered_type_names(&self) -> &[&'static str] {
&self.registered_type_names
}
pub fn get<K>(&self) -> Option<&TimeBucketedCounter<K>>
where
K: std::hash::Hash + Eq + Clone + Send + 'static,
{
self.by_type
.get(&TypeId::of::<K>())
.and_then(|b| b.downcast_ref::<TimeBucketedCounter<K>>())
}
pub fn get_mut<K>(&mut self) -> &mut TimeBucketedCounter<K>
where
K: std::hash::Hash + Eq + Clone + Send + 'static,
{
let id = TypeId::of::<K>();
self.by_type
.get_mut(&id)
.expect("counter::<K> not registered โ call .counter::<K>(...) on the builder")
.downcast_mut::<TimeBucketedCounter<K>>()
.expect("CounterRegistry invariant: TypeId keys to its own counter")
}
pub fn len(&self) -> usize {
self.by_type.len()
}
pub fn is_empty(&self) -> bool {
self.by_type.is_empty()
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use super::*;
#[derive(Default)]
struct Counter1 {
n: u64,
}
#[derive(Default)]
struct Counter2 {
m: u32,
}
#[test]
fn state_map_lazy_creates_then_returns_same() {
let mut m = StateMap::default();
m.get_or_init_mut::<Counter1>().n = 7;
assert_eq!(m.get_or_init_mut::<Counter1>().n, 7);
assert_eq!(m.len(), 1);
}
#[test]
fn state_map_segregates_by_type() {
let mut m = StateMap::default();
m.get_or_init_mut::<Counter1>().n = 1;
m.get_or_init_mut::<Counter2>().m = 2;
assert_eq!(m.get_or_init_mut::<Counter1>().n, 1);
assert_eq!(m.get_or_init_mut::<Counter2>().m, 2);
assert_eq!(m.len(), 2);
}
#[test]
fn state_map_is_empty_initially() {
let m = StateMap::default();
assert!(m.is_empty());
assert_eq!(m.len(), 0);
}
#[test]
fn take_dyn_removes_slot_and_downcasts() {
let mut m = StateMap::default();
m.get_or_init_mut::<Counter1>().n = 42;
let taken = m.take_dyn(std::any::TypeId::of::<Counter1>());
assert_eq!(taken.unwrap().downcast::<Counter1>().unwrap().n, 42);
assert!(m.is_empty()); assert!(m.take_dyn(std::any::TypeId::of::<Counter1>()).is_none());
}
#[test]
fn counter_registry_returns_registered() {
let mut r = CounterRegistry::default();
r.register::<u32>(TimeBucketedCounter::<u32>::new_unbounded(
Duration::from_secs(10),
Duration::from_secs(1),
));
let c = r.get_mut::<u32>();
c.bump(1u32, flowscope::Timestamp::new(0, 0));
assert!(!r.is_empty());
assert_eq!(r.len(), 1);
}
#[test]
#[should_panic(expected = "counter::<K> not registered")]
fn counter_registry_panics_on_missing_key() {
let mut r = CounterRegistry::default();
let _ = r.get_mut::<u64>();
}
}