use crate::error::*;
use std::any::Any;
use std::any::TypeId;
use std::cell::RefCell;
use std::{collections::HashMap, marker::PhantomData, rc::Rc};
thread_local! {
static THREAD_LOCALS: RefCell<AmbientMap> = RefCell::new(AmbientMap::new());
}
struct AmbientMap {
data: HashMap<TypeId, Vec<Slot>>,
id_counter: usize,
}
#[derive(Debug)]
struct Slot {
ptr: Rc<dyn Any + 'static>,
id: usize,
}
impl AmbientMap {
pub fn new() -> AmbientMap {
AmbientMap {
data: HashMap::new(),
id_counter: 0,
}
}
pub fn peek<T: 'static>(&self) -> Option<Rc<T>> {
self.data.get(&TypeId::of::<T>()).and_then(|stack| {
if stack.len() > 0 {
stack[stack.len() - 1].ptr.clone().downcast::<T>().ok()
} else {
None
}
})
}
pub fn has<T: 'static>(&self) -> bool {
self.data
.get(&TypeId::of::<T>())
.map(|stack| !stack.is_empty())
.unwrap_or(false)
}
pub fn remove(&mut self, type_id: &TypeId, id: &usize) {
self.data
.get_mut(type_id)
.expect("tried to remove empty ambient data stack")
.retain(|slot| &slot.id != id);
}
pub fn push<T: 'static>(&mut self, new_val: Rc<T>) -> AmbientGuard<T> {
let type_id = TypeId::of::<T>();
let new_id = self.id_counter;
self.id_counter += 1;
self.data
.entry(type_id)
.or_insert_with(|| Vec::new())
.push(Slot {
ptr: new_val,
id: new_id,
});
AmbientGuard {
phantom: PhantomData::<*mut T>::default(),
id: new_id,
}
}
}
pub struct AmbientGuard<T>
where
T: 'static,
{
phantom: PhantomData<*mut T>,
id: usize,
}
impl<T> Drop for AmbientGuard<T>
where
T: 'static,
{
fn drop(&mut self) {
unset(&TypeId::of::<T>(), &self.id);
}
}
pub fn get<T: 'static>() -> Result<Rc<T>> {
THREAD_LOCALS.with(|frame_opt| {
frame_opt
.borrow()
.peek::<T>()
.ok_or_else(|| Error::ThreadAmbientUndefined(std::any::type_name::<T>()))
})
}
pub fn has<T: 'static>() -> bool {
THREAD_LOCALS.with(|frame_opt| frame_opt.borrow().has::<T>())
}
fn unset(type_id: &TypeId, id: &usize) {
THREAD_LOCALS.with(|frame_opt| {
let mut storage = frame_opt.borrow_mut();
storage.remove(type_id, id)
})
}
#[must_use]
pub fn set<T: 'static>(new_val: T) -> AmbientGuard<T> {
THREAD_LOCALS.with(|frame_opt| {
let mut storage = frame_opt.borrow_mut();
storage.push(Rc::new(new_val))
})
}
pub fn set_rc<T: 'static>(new_val: Rc<T>) -> AmbientGuard<T> {
THREAD_LOCALS.with(|frame_opt| {
let mut storage = frame_opt.borrow_mut();
storage.push(new_val)
})
}
fn _data_must_be_static() {}
fn _guard_is_send() {}
fn _guard_is_sync() {}
#[cfg(test)]
mod tests {
use crate::error::*;
use crate::thread as ambience;
use std::any::Any;
use std::rc::Rc;
#[test]
fn simple() -> Result<()> {
let one = 1u64;
{
let _frame_guard = ambience::set(one);
assert_eq!(*ambience::get::<u64>().unwrap(), one);
assert!(ambience::has::<u64>());
}
assert!(ambience::get::<u64>().is_err());
assert!(!ambience::has::<u64>());
Ok(())
}
#[test]
fn nested() -> Result<()> {
let one = 1u64;
let two = 2u64;
{
let _frame_guard = ambience::set(one.clone());
assert_eq!(*ambience::get::<u64>().unwrap(), one);
{
let _frame_guard = ambience::set(two.clone());
assert_eq!(*ambience::get::<u64>().unwrap(), two);
}
assert_eq!(*ambience::get::<u64>().unwrap(), one);
}
assert!(ambience::get::<u64>().is_err());
Ok(())
}
#[test]
fn types_are_independent() -> Result<()> {
let u_64 = 64u64;
let u_32 = 32u32;
{
let _frame_guard = ambience::set(u_32);
let _frame_guard = ambience::set(u_64);
assert_eq!(*ambience::get::<u64>().unwrap(), u_64);
assert_eq!(*ambience::get::<u32>().unwrap(), u_32);
}
assert!(ambience::get::<u64>().is_err());
assert!(ambience::get::<u32>().is_err());
Ok(())
}
#[test]
fn external_rc_can_be_reused() -> Result<()> {
let one = Rc::new(1u64);
{
let _frame_guard = ambience::set_rc(one.clone());
assert_eq!(*ambience::get::<u64>().unwrap(), *one);
}
assert!(ambience::get::<u64>().is_err());
Ok(())
}
#[test]
fn manually_dropped() -> Result<()> {
let one = 1u64;
let two = 2u64;
{
let outer_frame_guard = ambience::set(one);
assert_eq!(*ambience::get::<u64>().unwrap(), one);
{
let _frame_guard = ambience::set(two);
drop(outer_frame_guard);
assert_eq!(*ambience::get::<u64>().unwrap(), two);
}
assert!(ambience::get::<u64>().is_err());
}
assert!(ambience::get::<u64>().is_err());
Ok(())
}
#[test]
fn downcast_expectations() -> Result<()> {
let arc_any: Rc<dyn Any + 'static> = Rc::new(1u64);
assert!(arc_any.clone().downcast::<u64>().is_ok());
assert!(arc_any.clone().downcast::<Rc<u64>>().is_err());
Ok(())
}
}