use std::any::{Any, TypeId};
use std::collections::{HashMap, HashSet};
use std::sync::RwLock;
use crate::engine::types::ChannelID;
use super::error::{EnvironmentError, EnvironmentResult};
use super::handle::EnvKey;
struct EntrySlot {
channel_id: ChannelID,
value: RwLock<EntryValue>,
}
struct EntryValue {
value: Box<dyn Any + Send + Sync>,
type_id: TypeId,
type_name: &'static str,
}
pub struct Environment {
entries: HashMap<String, EntrySlot>,
dirty_channels: RwLock<HashSet<ChannelID>>,
}
const _: fn() = || {
fn assert_send_sync<T: Send + Sync>() {}
assert_send_sync::<Environment>();
};
impl Environment {
pub(super) fn from_schema(
schema: Vec<(String, Box<dyn Any + Send + Sync>, TypeId, &'static str)>,
channel_ids: Vec<ChannelID>,
) -> Self {
debug_assert_eq!(
schema.len(),
channel_ids.len(),
"schema and channel_ids must be the same length"
);
let mut entries = HashMap::with_capacity(schema.len());
for ((key, value, type_id, type_name), channel_id) in schema.into_iter().zip(channel_ids) {
entries.insert(
key,
EntrySlot {
channel_id,
value: RwLock::new(EntryValue {
value,
type_id,
type_name,
}),
},
);
}
Self {
entries,
dirty_channels: RwLock::new(HashSet::new()),
}
}
#[inline]
pub fn len(&self) -> usize {
self.entries.len()
}
#[inline]
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
#[inline]
pub fn contains_key(&self, key: &str) -> bool {
self.entries.contains_key(key)
}
#[inline]
pub fn channel_of(&self, key: &str) -> Option<ChannelID> {
self.entries.get(key).map(|slot| slot.channel_id)
}
pub fn all_channel_ids(&self) -> Vec<ChannelID> {
let mut ids: Vec<ChannelID> = self.entries.values().map(|s| s.channel_id).collect();
ids.sort_unstable();
ids
}
#[inline]
pub fn env_key<T: Any + Clone + Send + Sync>(&self, key: &'static str) -> Option<EnvKey<T>> {
let channel_id = self.channel_of(key)?;
Some(EnvKey::new(key, channel_id))
}
pub fn get<T: Any + Clone + Send + Sync>(&self, key: &str) -> EnvironmentResult<T> {
let slot = self
.entries
.get(key)
.ok_or_else(|| EnvironmentError::KeyNotFound(key.to_owned()))?;
let entry = slot
.value
.read()
.map_err(|_| EnvironmentError::LockPoisoned {
what: "environment entry",
})?;
let requested_type_id = TypeId::of::<T>();
if entry.type_id != requested_type_id {
return Err(EnvironmentError::TypeMismatch {
key: key.to_owned(),
expected: entry.type_name,
actual: std::any::type_name::<T>(),
});
}
let value = entry
.value
.downcast_ref::<T>()
.expect("TypeId matched but downcast failed - this is a bug");
Ok(value.clone())
}
pub fn set<T: Any + Clone + Send + Sync>(&self, key: &str, value: T) -> EnvironmentResult<()> {
let slot = self
.entries
.get(key)
.ok_or_else(|| EnvironmentError::KeyNotFound(key.to_owned()))?;
let channel_id = slot.channel_id;
{
let mut entry = slot
.value
.write()
.map_err(|_| EnvironmentError::LockPoisoned {
what: "environment entry",
})?;
let requested_type_id = TypeId::of::<T>();
if entry.type_id != requested_type_id {
return Err(EnvironmentError::TypeMismatch {
key: key.to_owned(),
expected: entry.type_name,
actual: std::any::type_name::<T>(),
});
}
*entry
.value
.downcast_mut::<T>()
.expect("TypeId matched but downcast_mut failed - this is a bug") = value;
}
self.dirty_channels
.write()
.map_err(|_| EnvironmentError::LockPoisoned {
what: "environment dirty channels",
})?
.insert(channel_id);
Ok(())
}
#[inline]
pub(crate) fn has_any_dirty_channels(
&self,
channels: impl Iterator<Item = ChannelID>,
) -> EnvironmentResult<bool> {
let dirty = self
.dirty_channels
.read()
.map_err(|_| EnvironmentError::LockPoisoned {
what: "environment dirty channels",
})?;
Ok(channels.into_iter().any(|id| dirty.contains(&id)))
}
#[cfg(feature = "gpu")]
#[inline]
pub(crate) fn is_channel_dirty(&self, id: ChannelID) -> EnvironmentResult<bool> {
Ok(self
.dirty_channels
.read()
.map_err(|_| EnvironmentError::LockPoisoned {
what: "environment dirty channels",
})?
.contains(&id))
}
#[inline]
pub(crate) fn dirty_channel_ids(&self) -> EnvironmentResult<HashSet<ChannelID>> {
Ok(self
.dirty_channels
.read()
.map_err(|_| EnvironmentError::LockPoisoned {
what: "environment dirty channels",
})?
.clone())
}
#[inline]
pub(crate) fn clear_dirty_for_channels(&self, channels: &[ChannelID]) -> EnvironmentResult<()> {
let mut dirty =
self.dirty_channels
.write()
.map_err(|_| EnvironmentError::LockPoisoned {
what: "environment dirty channels",
})?;
for id in channels {
dirty.remove(id);
}
Ok(())
}
#[inline]
pub(crate) fn clear_dirty(&self) -> EnvironmentResult<()> {
self.dirty_channels
.write()
.map_err(|_| EnvironmentError::LockPoisoned {
what: "environment dirty channels",
})?
.clear();
Ok(())
}
#[cfg(test)]
pub(crate) fn poison_dirty_channels_for_test(&self) {
let _guard = self.dirty_channels.write().unwrap();
panic!("poison environment dirty-channel lock");
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::environment::builder::EnvironmentBuilder;
use std::sync::Arc;
fn build_env() -> Arc<Environment> {
EnvironmentBuilder::new()
.register::<f32>("interest_rate", 0.05)
.unwrap()
.register::<u32>("world_width", 100)
.unwrap()
.register::<bool>("verbose", false)
.unwrap()
.build()
.unwrap()
}
#[test]
fn get_registered_value() {
let env = build_env();
let v: f32 = env.get("interest_rate").unwrap();
assert!((v - 0.05f32).abs() < f32::EPSILON);
}
#[test]
fn get_key_not_found() {
let env = build_env();
let err = env.get::<f32>("missing_key").unwrap_err();
assert!(matches!(err, EnvironmentError::KeyNotFound(_)));
}
#[test]
fn get_type_mismatch() {
let env = build_env();
let err = env.get::<f64>("interest_rate").unwrap_err();
assert!(matches!(err, EnvironmentError::TypeMismatch { .. }));
}
#[test]
fn set_updates_value() {
let env = build_env();
env.set::<f32>("interest_rate", 0.10).unwrap();
let v: f32 = env.get("interest_rate").unwrap();
assert!((v - 0.10f32).abs() < f32::EPSILON);
}
#[test]
fn set_marks_channel_dirty() {
let env = build_env();
let id = env.channel_of("interest_rate").unwrap();
env.set::<f32>("interest_rate", 0.10).unwrap();
let dirty = env.dirty_channel_ids().unwrap();
assert!(dirty.contains(&id));
}
#[test]
fn clear_dirty_empties_set() {
let env = build_env();
env.set::<f32>("interest_rate", 0.10).unwrap();
env.clear_dirty().unwrap();
assert!(env.dirty_channel_ids().unwrap().is_empty());
}
#[test]
fn set_type_mismatch_returns_error() {
let env = build_env();
let err = env.set::<f64>("interest_rate", 0.10f64).unwrap_err();
assert!(matches!(err, EnvironmentError::TypeMismatch { .. }));
}
#[test]
fn set_missing_key_returns_error() {
let env = build_env();
let err = env.set::<f32>("nonexistent", 1.0).unwrap_err();
assert!(matches!(err, EnvironmentError::KeyNotFound(_)));
}
#[test]
fn contains_key() {
let env = build_env();
assert!(env.contains_key("world_width"));
assert!(!env.contains_key("planet_radius"));
}
#[test]
fn bool_roundtrip() {
let env = build_env();
let v: bool = env.get("verbose").unwrap();
assert!(!v);
env.set("verbose", true).unwrap();
assert!(env.get::<bool>("verbose").unwrap());
}
#[test]
fn len_and_is_empty() {
let env = build_env();
assert_eq!(env.len(), 3);
assert!(!env.is_empty());
let empty = EnvironmentBuilder::new().build().unwrap();
assert!(empty.is_empty());
}
#[test]
fn dirty_channels_are_deduplicated() {
let env = build_env();
let id = env.channel_of("interest_rate").unwrap();
for _ in 0..100 {
env.set::<f32>("interest_rate", 0.10).unwrap();
}
let dirty = env.dirty_channel_ids().unwrap();
assert_eq!(dirty.len(), 1);
assert!(dirty.contains(&id));
}
#[test]
fn dirty_tracks_multiple_distinct_channels() {
let env = build_env();
let id_rate = env.channel_of("interest_rate").unwrap();
let id_width = env.channel_of("world_width").unwrap();
env.set::<f32>("interest_rate", 0.10).unwrap();
env.set::<u32>("world_width", 200).unwrap();
let dirty = env.dirty_channel_ids().unwrap();
assert_eq!(dirty.len(), 2);
assert!(dirty.contains(&id_rate));
assert!(dirty.contains(&id_width));
}
#[test]
fn has_any_dirty_channels_returns_true_for_dirty() {
let env = build_env();
let id = env.channel_of("interest_rate").unwrap();
env.set::<f32>("interest_rate", 0.10).unwrap();
assert!(env.has_any_dirty_channels([id].into_iter()).unwrap());
}
#[test]
fn has_any_dirty_channels_returns_false_for_clean() {
let env = build_env();
let id = env.channel_of("interest_rate").unwrap();
assert!(!env.has_any_dirty_channels([id].into_iter()).unwrap());
}
#[test]
fn has_any_dirty_channels_ignores_unrelated() {
let env = build_env();
let id_rate = env.channel_of("interest_rate").unwrap();
let id_width = env.channel_of("world_width").unwrap();
env.set::<u32>("world_width", 200).unwrap();
assert!(!env.has_any_dirty_channels([id_rate].into_iter()).unwrap());
assert!(env
.has_any_dirty_channels([id_rate, id_width].into_iter())
.unwrap());
}
#[test]
fn clear_dirty_for_channels_is_selective() {
let env = build_env();
let id_rate = env.channel_of("interest_rate").unwrap();
let id_width = env.channel_of("world_width").unwrap();
env.set::<f32>("interest_rate", 0.10).unwrap();
env.set::<u32>("world_width", 200).unwrap();
env.clear_dirty_for_channels(&[id_rate]).unwrap();
let dirty = env.dirty_channel_ids().unwrap();
assert!(!dirty.contains(&id_rate));
assert!(dirty.contains(&id_width));
}
#[test]
fn channel_of_returns_none_for_unknown_key() {
let env = build_env();
assert!(env.channel_of("nonexistent").is_none());
}
#[test]
fn channel_of_returns_distinct_ids_for_distinct_keys() {
let env = build_env();
let id_rate = env.channel_of("interest_rate").unwrap();
let id_width = env.channel_of("world_width").unwrap();
let id_verbose = env.channel_of("verbose").unwrap();
assert_ne!(id_rate, id_width);
assert_ne!(id_rate, id_verbose);
assert_ne!(id_width, id_verbose);
}
#[test]
fn env_key_roundtrip() {
let env = build_env();
let key = env.env_key::<f32>("interest_rate").unwrap();
assert_eq!(key.name(), "interest_rate");
assert_eq!(key.channel_id(), env.channel_of("interest_rate").unwrap());
}
#[test]
fn env_key_returns_none_for_unknown_key() {
let env = build_env();
assert!(env.env_key::<f32>("nonexistent").is_none());
}
#[test]
fn concurrent_reads_do_not_block() {
use std::thread;
let env = build_env();
let env2 = Arc::clone(&env);
let handle = thread::spawn(move || {
for _ in 0..1000 {
let _ = env2.get::<f32>("interest_rate").unwrap();
}
});
for _ in 0..1000 {
let _ = env.get::<u32>("world_width").unwrap();
}
handle.join().unwrap();
}
#[test]
fn channel_of_is_lock_free_against_writers() {
use std::thread;
use std::time::{Duration, Instant};
let env = build_env();
let env_writer = Arc::clone(&env);
let stop_at = Instant::now() + Duration::from_millis(50);
let handle = thread::spawn(move || {
while Instant::now() < stop_at {
env_writer.set::<f32>("interest_rate", 0.07).unwrap();
}
});
let start = Instant::now();
for _ in 0..100_000 {
let _ = env.channel_of("interest_rate").unwrap();
}
let elapsed = start.elapsed();
handle.join().unwrap();
assert!(
elapsed < Duration::from_millis(500),
"channel_of appears to serialise against writers: {:?}",
elapsed
);
}
#[test]
fn poisoned_entry_lock_returns_structured_error() {
use std::thread;
let env = build_env();
let env_for_thread = Arc::clone(&env);
let _ = thread::spawn(move || {
let slot = env_for_thread.entries.get("interest_rate").unwrap();
let _guard = slot.value.write().unwrap();
panic!("poison environment entry");
})
.join();
let err = env.get::<f32>("interest_rate").unwrap_err();
assert!(matches!(
err,
EnvironmentError::LockPoisoned {
what: "environment entry"
}
));
}
#[test]
fn poisoned_dirty_channel_helpers_return_structured_errors() {
use std::thread;
let env = build_env();
let id = env.channel_of("interest_rate").unwrap();
let env_for_thread = Arc::clone(&env);
let _ = thread::spawn(move || env_for_thread.poison_dirty_channels_for_test()).join();
assert!(matches!(
env.has_any_dirty_channels([id].into_iter()),
Err(EnvironmentError::LockPoisoned {
what: "environment dirty channels"
})
));
assert!(matches!(
env.dirty_channel_ids(),
Err(EnvironmentError::LockPoisoned {
what: "environment dirty channels"
})
));
assert!(matches!(
env.clear_dirty_for_channels(&[id]),
Err(EnvironmentError::LockPoisoned {
what: "environment dirty channels"
})
));
assert!(matches!(
env.clear_dirty(),
Err(EnvironmentError::LockPoisoned {
what: "environment dirty channels"
})
));
}
}