#![allow(clippy::disallowed_types)]
use core::any::{Any, TypeId};
use core::fmt;
use core::pin::Pin;
use crate::std::{boxed::Box, sync::Arc, vec::Vec};
pub use rama_macros::{Extension, FromExtensions};
use rama_utils::collections::AppendOnlyVec;
use rama_utils::macros::impl_deref;
#[derive(Clone, Default)]
pub struct Extensions {
extensions: Arc<AppendOnlyVec<TypeErasedExtension, 12, 3>>,
parent: Option<Box<Self>>,
}
impl Extensions {
#[inline(always)]
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn fork(&self) -> Self {
Self {
extensions: Arc::new(AppendOnlyVec::new()),
parent: Some(Box::new(self.clone())),
}
}
#[inline(always)]
#[must_use]
pub fn parent(&self) -> Option<&Self> {
self.parent.as_deref()
}
#[must_use]
pub fn with_base(&self, base: &Self) -> Self {
let parent = match &self.parent {
Some(parent) => parent.with_base(base),
None => base.clone(),
};
Self {
extensions: Arc::clone(&self.extensions),
parent: Some(Box::new(parent)),
}
}
pub fn insert<T: Extension>(&self, val: T) -> &T {
let extension = TypeErasedExtension::new(val);
let idx = self.extensions.push(extension);
#[expect(
clippy::unwrap_used,
reason = "`downcast_ref` can only be none if TypeId doesn't match, but we just inserted this type"
)]
self.extensions[idx].downcast_ref::<T>().unwrap()
}
pub fn insert_arc<T: Extension>(&self, val: Arc<T>) -> Arc<T> {
let extension = TypeErasedExtension::new_arc(val);
let idx = self.extensions.push(extension);
#[expect(
clippy::unwrap_used,
reason = "`cloned_downcast` can only be none if TypeId doesn't match, but we just inserted this type"
)]
self.extensions[idx].cloned_downcast::<T>().unwrap()
}
pub fn extend(&self, other: &Self) {
for ext in other.extensions.iter() {
self.extensions.push(ext.clone());
}
}
#[must_use]
pub fn contains<T: Extension>(&self) -> bool {
self.get_ref::<T>().is_some()
}
#[must_use]
pub fn self_contains<T: Extension>(&self) -> bool {
let type_id = TypeId::of::<T>();
self.extensions
.iter()
.rev()
.any(|item| item.type_id == type_id)
}
#[must_use]
pub fn get_ref<T: Extension>(&self) -> Option<&T> {
let target = TypeId::of::<T>();
let egress_id = TypeId::of::<Egress<Self>>();
let ingress_id = TypeId::of::<Ingress<Self>>();
for ext in self.extensions.iter().rev() {
if ext.type_id == target {
if let Some(v) = ext.downcast_ref::<T>() {
return Some(v);
}
} else if ext.type_id == egress_id
&& let Some(eg) = ext.downcast_ref::<Egress<Self>>()
&& let Some(v) = eg.0.get_ref::<T>()
{
return Some(v);
} else if ext.type_id == ingress_id
&& let Some(ig) = ext.downcast_ref::<Ingress<Self>>()
&& let Some(v) = ig.0.get_ref::<T>()
{
return Some(v);
}
}
self.parent().and_then(|p| p.get_ref::<T>())
}
#[must_use]
pub fn self_get_ref<T: Extension>(&self) -> Option<&T> {
let type_id = TypeId::of::<T>();
self.extensions
.iter()
.rev()
.find(|item| item.type_id == type_id)
.and_then(|ext| ext.downcast_ref())
}
#[must_use]
pub fn get_arc<T: Extension>(&self) -> Option<Arc<T>> {
let target = TypeId::of::<T>();
let egress_id = TypeId::of::<Egress<Self>>();
let ingress_id = TypeId::of::<Ingress<Self>>();
for ext in self.extensions.iter().rev() {
if ext.type_id == target {
if let Some(v) = ext.cloned_downcast::<T>() {
return Some(v);
}
} else if ext.type_id == egress_id
&& let Some(eg) = ext.downcast_ref::<Egress<Self>>()
&& let Some(v) = eg.0.get_arc::<T>()
{
return Some(v);
} else if ext.type_id == ingress_id
&& let Some(ig) = ext.downcast_ref::<Ingress<Self>>()
&& let Some(v) = ig.0.get_arc::<T>()
{
return Some(v);
}
}
self.parent().and_then(|p| p.get_arc::<T>())
}
#[must_use]
pub fn self_get_arc<T: Extension>(&self) -> Option<Arc<T>> {
let type_id = TypeId::of::<T>();
self.extensions
.iter()
.rev()
.find(|item| item.type_id == type_id)
.and_then(|ext| ext.cloned_downcast())
}
pub fn get_many_ref<'a, T: GetManyRef<'a>>(&'a self) -> T::Output {
T::get_many_ref(self)
}
pub fn get_many_arc<T: GetManyArc>(&self) -> T::Output {
T::get_many_arc(self)
}
#[doc(hidden)]
pub fn get_many_erased<'a, const N: usize>(
&'a self,
targets: &[TypeId; N],
out: &mut [Option<(&'a TypeErasedExtension, usize)>; N],
) {
let mut rank = 0;
self.get_many_erased_ranked(targets, out, &mut rank);
}
fn get_many_erased_ranked<'a, const N: usize>(
&'a self,
targets: &[TypeId; N],
out: &mut [Option<(&'a TypeErasedExtension, usize)>; N],
rank: &mut usize,
) {
let egress_id = TypeId::of::<Egress<Self>>();
let ingress_id = TypeId::of::<Ingress<Self>>();
let mut remaining = out.iter().filter(|slot| slot.is_none()).count();
if remaining == 0 {
return;
}
for ext in self.extensions.iter().rev() {
let current = *rank;
*rank += 1;
for (i, &tid) in targets.iter().enumerate() {
if out[i].is_none() && ext.type_id == tid {
out[i] = Some((ext, current));
remaining -= 1;
if remaining == 0 {
return;
}
}
}
if ext.type_id == egress_id
&& let Some(eg) = ext.downcast_ref::<Egress<Self>>()
{
eg.0.get_many_erased_ranked(targets, out, rank);
remaining = out.iter().filter(|slot| slot.is_none()).count();
if remaining == 0 {
return;
}
} else if ext.type_id == ingress_id
&& let Some(ig) = ext.downcast_ref::<Ingress<Self>>()
{
ig.0.get_many_erased_ranked(targets, out, rank);
remaining = out.iter().filter(|slot| slot.is_none()).count();
if remaining == 0 {
return;
}
}
}
if remaining != 0
&& let Some(parent) = self.parent()
{
parent.get_many_erased_ranked(targets, out, rank);
}
}
pub fn get_ref_or_insert<T, F>(&self, create_fn: F) -> &T
where
T: Extension,
F: FnOnce() -> T,
{
self.get_ref().unwrap_or_else(|| self.insert(create_fn()))
}
pub fn get_arc_or_insert<T, F>(&self, create_fn: F) -> Arc<T>
where
T: Extension,
F: FnOnce() -> Arc<T>,
{
self.get_arc()
.unwrap_or_else(|| self.insert_arc(create_fn()))
}
pub fn self_get_ref_or_insert<T, F>(&self, create_fn: F) -> &T
where
T: Extension,
F: FnOnce() -> T,
{
self.self_get_ref()
.unwrap_or_else(|| self.insert(create_fn()))
}
pub fn self_get_arc_or_insert<T, F>(&self, create_fn: F) -> Arc<T>
where
T: Extension,
F: FnOnce() -> Arc<T>,
{
self.self_get_arc()
.unwrap_or_else(|| self.insert_arc(create_fn()))
}
#[must_use]
pub fn self_first_ref<T: Extension>(&self) -> Option<&T> {
let type_id = TypeId::of::<T>();
self.extensions
.iter()
.find(|item| item.type_id == type_id)
.and_then(|ext| ext.downcast_ref())
}
#[must_use]
pub fn self_first_arc<T: Extension>(&self) -> Option<Arc<T>> {
let type_id = TypeId::of::<T>();
self.extensions
.iter()
.find(|item| item.type_id == type_id)
.and_then(|ext| ext.cloned_downcast())
}
pub fn self_iter_ref<T: Extension>(&self) -> impl Iterator<Item = &T> {
let type_id = TypeId::of::<T>();
self.extensions
.iter()
.rev()
.filter(move |item| item.type_id == type_id)
.filter_map(TypeErasedExtension::downcast_ref::<T>)
}
pub fn self_iter_arc<T: Extension>(&self) -> impl Iterator<Item = Arc<T>> {
let type_id = TypeId::of::<T>();
self.extensions
.iter()
.rev()
.filter(move |item| item.type_id == type_id)
.filter_map(TypeErasedExtension::cloned_downcast::<T>)
}
pub fn self_iter_all(&self) -> impl Iterator<Item = &TypeErasedExtension> {
self.extensions.iter()
}
pub fn iter_ref<T: Extension>(&self) -> impl Iterator<Item = &T> + '_ {
self.iter_ref_inner::<T>()
}
pub fn iter_arc<T: Extension>(&self) -> impl Iterator<Item = Arc<T>> + '_ {
self.iter_arc_inner::<T>()
}
fn iter_ref_inner<T: Extension>(&self) -> Box<dyn Iterator<Item = &T> + '_> {
let target = TypeId::of::<T>();
let egress_id = TypeId::of::<Egress<Self>>();
let ingress_id = TypeId::of::<Ingress<Self>>();
let local = self.extensions.iter().rev().flat_map(
move |ext| -> Box<dyn Iterator<Item = &T> + '_> {
if ext.type_id == target {
match ext.downcast_ref::<T>() {
Some(v) => Box::new(core::iter::once(v)),
None => Box::new(core::iter::empty()),
}
} else if ext.type_id == egress_id {
match ext.downcast_ref::<Egress<Self>>() {
Some(e) => e.0.iter_ref_inner::<T>(),
None => Box::new(core::iter::empty()),
}
} else if ext.type_id == ingress_id {
match ext.downcast_ref::<Ingress<Self>>() {
Some(i) => i.0.iter_ref_inner::<T>(),
None => Box::new(core::iter::empty()),
}
} else {
Box::new(core::iter::empty())
}
},
);
let parent: Box<dyn Iterator<Item = &T>> = match self.parent() {
Some(p) => p.iter_ref_inner::<T>(),
None => Box::new(core::iter::empty()),
};
Box::new(local.chain(parent))
}
fn iter_arc_inner<T: Extension>(&self) -> Box<dyn Iterator<Item = Arc<T>> + '_> {
let target = TypeId::of::<T>();
let egress_id = TypeId::of::<Egress<Self>>();
let ingress_id = TypeId::of::<Ingress<Self>>();
let local = self.extensions.iter().rev().flat_map(
move |ext| -> Box<dyn Iterator<Item = Arc<T>> + '_> {
if ext.type_id == target {
match ext.cloned_downcast::<T>() {
Some(v) => Box::new(core::iter::once(v)),
None => Box::new(core::iter::empty()),
}
} else if ext.type_id == egress_id {
match ext.downcast_ref::<Egress<Self>>() {
Some(e) => e.0.iter_arc_inner::<T>(),
None => Box::new(core::iter::empty()),
}
} else if ext.type_id == ingress_id {
match ext.downcast_ref::<Ingress<Self>>() {
Some(i) => i.0.iter_arc_inner::<T>(),
None => Box::new(core::iter::empty()),
}
} else {
Box::new(core::iter::empty())
}
},
);
let parent: Box<dyn Iterator<Item = Arc<T>>> = match self.parent() {
Some(p) => p.iter_arc_inner::<T>(),
None => Box::new(core::iter::empty()),
};
Box::new(local.chain(parent))
}
#[must_use]
pub fn ingress(&self) -> Option<&Ingress<Self>> {
self.get_ref::<Ingress<Self>>()
}
#[must_use]
pub fn egress(&self) -> Option<&Egress<Self>> {
self.get_ref::<Egress<Self>>()
}
pub fn clone_to<T: Extension>(&self, target: &Self) -> Option<Arc<T>> {
let item = self.get_arc();
if let Some(item) = item.clone() {
target.insert_arc(item);
};
item
}
pub fn clone_to_if_absent<T: Extension>(&self, target: &Self) -> Option<Arc<T>> {
let item = target.get_arc::<T>();
if item.is_some() {
return item;
}
self.clone_to(target)
}
}
#[doc(hidden)]
pub trait FromExtensionsGroup<'a>: Sized {
const TARGETS: usize;
fn from_ext_targets(targets: &mut [TypeId], offset: usize);
fn from_ext_slots(
out: &[Option<(&'a TypeErasedExtension, usize)>],
offset: usize,
) -> Option<Self>;
}
pub trait GetManyRef<'a>: Sized {
type Output;
#[doc(hidden)]
const SEALED: seal::Seal;
#[doc(hidden)]
fn get_many_ref(ext: &'a Extensions) -> Self::Output;
}
pub trait GetManyArc: Sized {
type Output;
#[doc(hidden)]
const SEALED: seal::Seal;
#[doc(hidden)]
fn get_many_arc(ext: &Extensions) -> Self::Output;
}
macro_rules! impl_get_many {
($n:literal; $($T:ident => $idx:tt),+ $(,)?) => {
impl<'a, $($T: Extension),+> GetManyRef<'a> for ($($T,)+) {
type Output = ($(Option<&'a $T>,)+);
const SEALED: seal::Seal = seal::Seal;
fn get_many_ref(ext: &'a Extensions) -> Self::Output {
let targets = [$(TypeId::of::<$T>()),+];
let mut out: [Option<(&'a TypeErasedExtension, usize)>; $n] = [None; $n];
ext.get_many_erased(&targets, &mut out);
($(out[$idx].and_then(|(e, _)| e.downcast_ref::<$T>()),)+)
}
}
impl<$($T: Extension),+> GetManyArc for ($($T,)+) {
type Output = ($(Option<Arc<$T>>,)+);
const SEALED: seal::Seal = seal::Seal;
fn get_many_arc(ext: &Extensions) -> Self::Output {
let targets = [$(TypeId::of::<$T>()),+];
let mut out: [Option<(&TypeErasedExtension, usize)>; $n] = [None; $n];
ext.get_many_erased(&targets, &mut out);
($(out[$idx].and_then(|(e, _)| e.cloned_downcast::<$T>()),)+)
}
}
};
}
impl_get_many!(1; A => 0);
impl_get_many!(2; A => 0, B => 1);
impl_get_many!(3; A => 0, B => 1, C => 2);
impl_get_many!(4; A => 0, B => 1, C => 2, D => 3);
impl_get_many!(5; A => 0, B => 1, C => 2, D => 3, E => 4);
impl_get_many!(6; A => 0, B => 1, C => 2, D => 3, E => 4, F => 5);
impl_get_many!(7; A => 0, B => 1, C => 2, D => 3, E => 4, F => 5, G => 6);
impl_get_many!(8; A => 0, B => 1, C => 2, D => 3, E => 4, F => 5, G => 6, H => 7);
impl_get_many!(9; A => 0, B => 1, C => 2, D => 3, E => 4, F => 5, G => 6, H => 7, I => 8);
impl_get_many!(10; A => 0, B => 1, C => 2, D => 3, E => 4, F => 5, G => 6, H => 7, I => 8, J => 9);
impl_get_many!(11; A => 0, B => 1, C => 2, D => 3, E => 4, F => 5, G => 6, H => 7, I => 8, J => 9, K => 10);
impl_get_many!(12; A => 0, B => 1, C => 2, D => 3, E => 4, F => 5, G => 6, H => 7, I => 8, J => 9, K => 10, L => 11);
impl fmt::Debug for Extensions {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let mut s = f.debug_struct("Extensions");
if let Some(parent) = self.parent() {
s.field("parent", parent);
}
s.field(
"entries",
&self.extensions.iter().map(|e| &e.value).collect::<Vec<_>>(),
);
s.finish()
}
}
#[derive(Clone, Debug)]
pub struct TypeErasedExtension {
type_id: TypeId,
value: Arc<dyn Extension>,
}
impl TypeErasedExtension {
pub fn new<T: Extension>(value: T) -> Self {
Self {
type_id: TypeId::of::<T>(),
value: Arc::new(value),
}
}
pub fn new_arc<T: Extension>(value: Arc<T>) -> Self {
Self {
type_id: TypeId::of::<T>(),
value,
}
}
pub fn type_id(&self) -> TypeId {
self.type_id
}
pub fn cloned_downcast<T: Extension>(&self) -> Option<Arc<T>> {
let any = self.value.clone() as Arc<dyn Any + Send + Sync>;
any.downcast::<T>().ok()
}
pub fn downcast_ref<T: Extension>(&self) -> Option<&T> {
let inner_any = self.value.as_ref() as &dyn Any;
(inner_any).downcast_ref::<T>()
}
}
#[derive(Debug, Clone, Extension)]
pub struct Ingress<T>(pub T);
impl_deref!(Ingress);
#[derive(Debug, Clone, Extension)]
pub struct Egress<T>(pub T);
impl_deref!(Egress);
pub trait Extension: Any + Send + Sync + core::fmt::Debug + 'static {}
pub trait TlsExtension: Extension {}
pub trait HttpExtension: Extension {}
pub trait NetExtension: Extension {}
pub trait UaExtension: Extension {}
pub trait ProxyExtension: Extension {}
pub trait WsExtension: Extension {}
pub trait DnsExtension: Extension {}
pub trait GrpcExtension: Extension {}
pub trait ExtensionsRef {
fn extensions(&self) -> &Extensions;
}
impl ExtensionsRef for Extensions {
fn extensions(&self) -> &Extensions {
self
}
}
impl<T> ExtensionsRef for &T
where
T: ExtensionsRef,
{
#[inline(always)]
fn extensions(&self) -> &Extensions {
(**self).extensions()
}
}
impl<T> ExtensionsRef for &mut T
where
T: ExtensionsRef,
{
#[inline(always)]
fn extensions(&self) -> &Extensions {
(**self).extensions()
}
}
impl<T> ExtensionsRef for Box<T>
where
T: ExtensionsRef,
{
fn extensions(&self) -> &Extensions {
(**self).extensions()
}
}
impl<T> ExtensionsRef for Pin<Box<T>>
where
T: ExtensionsRef,
{
fn extensions(&self) -> &Extensions {
(**self).extensions()
}
}
impl<T> ExtensionsRef for Arc<T>
where
T: ExtensionsRef,
{
fn extensions(&self) -> &Extensions {
(**self).extensions()
}
}
macro_rules! impl_extensions_either {
($id:ident, $($param:ident),+ $(,)?) => {
impl<$($param),+,> ExtensionsRef for crate::combinators::$id<$($param),+>
where
$($param: ExtensionsRef,)+
{
fn extensions(&self) -> &Extensions {
match self {
$(crate::combinators::$id::$param(s) => s.extensions(),)+
}
}
}
};
}
crate::combinators::impl_either!(impl_extensions_either);
mod seal {
pub struct Seal;
}
#[cfg(test)]
mod tests {
use super::*;
use core::any::TypeId;
use core::pin::Pin;
use core::sync::atomic::{AtomicUsize, Ordering};
#[derive(Debug, Clone, PartialEq, Eq, Extension)]
struct TraceNote(String);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Extension)]
struct RetryBudget(u32);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Extension)]
struct ConnectionTimeoutMs(u64);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Extension)]
struct WorkerId(i32);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Extension)]
struct HealthSignal(u8);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Extension)]
struct FeatureToggle(bool);
#[test]
fn get_ref_returns_last_inserted() {
let ext = Extensions::new();
ext.insert(TraceNote("first".to_owned()));
ext.insert(TraceNote("second".to_owned()));
ext.insert(TraceNote("third".to_owned()));
assert_eq!(
ext.get_ref::<TraceNote>(),
Some(&TraceNote("third".to_owned()))
);
}
#[test]
fn clone_shares_backing_store() {
let ext = Extensions::new();
ext.insert(TraceNote("first".to_owned()));
let clone = ext.clone();
clone.insert(TraceNote("second".to_owned()));
assert_eq!(
ext.get_ref::<TraceNote>(),
Some(&TraceNote("second".to_owned()))
);
assert_eq!(
clone.get_ref::<TraceNote>(),
Some(&TraceNote("second".to_owned()))
);
}
#[test]
fn get_ref_none_when_absent() {
let ext = Extensions::new();
assert_eq!(ext.get_ref::<TraceNote>(), None);
}
#[test]
fn get_arc_none_when_absent() {
let ext = Extensions::new();
assert!(ext.get_arc::<TraceNote>().is_none());
}
#[test]
fn first_ref_none_when_absent() {
let ext = Extensions::new();
assert_eq!(ext.self_first_ref::<TraceNote>(), None);
}
#[test]
fn first_arc_none_when_absent() {
let ext = Extensions::new();
assert!(ext.self_first_arc::<TraceNote>().is_none());
}
#[test]
fn first_ref_returns_first_inserted() {
let ext = Extensions::new();
ext.insert(TraceNote("first".to_owned()));
ext.insert(TraceNote("second".to_owned()));
assert_eq!(
ext.self_first_ref::<TraceNote>(),
Some(&TraceNote("first".to_owned()))
);
}
#[test]
fn extend_appends_other_extensions() {
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Extension)]
struct DerivedMetric(i32);
let source = Extensions::new();
source.insert(WorkerId(5));
source.insert(DerivedMetric(10));
let target = Extensions::new();
target.extend(&source);
assert_eq!(target.get_ref::<WorkerId>(), Some(&WorkerId(5)));
assert_eq!(target.get_ref::<DerivedMetric>(), Some(&DerivedMetric(10)));
}
#[test]
fn insert_arc_can_be_retrieved_via_get_arc() {
let ext = Extensions::new();
let inserted = ext.insert_arc(Arc::new(TraceNote(String::from("hello"))));
let retrieved = ext.get_arc::<TraceNote>();
assert_eq!(inserted.0.as_str(), "hello");
assert_eq!(retrieved.as_deref().map(|it| it.0.as_str()), Some("hello"));
}
#[test]
fn insert_arc_can_be_retrieved_via_get_ref() {
let ext = Extensions::new();
ext.insert_arc(Arc::new(WorkerId(99)));
assert_eq!(ext.get_ref::<WorkerId>(), Some(&WorkerId(99)));
}
#[test]
fn contains_reports_presence_and_absence() {
let ext = Extensions::new();
assert!(!ext.contains::<RetryBudget>());
ext.insert(RetryBudget(1));
assert!(ext.contains::<RetryBudget>());
assert!(!ext.contains::<ConnectionTimeoutMs>());
}
#[test]
fn get_arc_and_first_arc_report_latest_and_oldest() {
let ext = Extensions::new();
ext.insert_arc(Arc::new(TraceNote(String::from("first"))));
ext.insert_arc(Arc::new(TraceNote(String::from("second"))));
assert_eq!(
ext.self_first_arc::<TraceNote>()
.as_deref()
.map(|it| it.0.as_str()),
Some("first")
);
assert_eq!(
ext.get_arc::<TraceNote>()
.as_deref()
.map(|it| it.0.as_str()),
Some("second")
);
}
#[test]
fn get_ref_or_insert_uses_existing_or_inserts_once() {
let ext = Extensions::new();
ext.insert(RetryBudget(5));
let calls = AtomicUsize::new(0);
let existing = ext.self_get_ref_or_insert(|| {
calls.fetch_add(1, Ordering::SeqCst);
RetryBudget(6)
});
assert_eq!(existing.0, 5u32);
assert_eq!(calls.load(Ordering::SeqCst), 0);
let missing = ext.self_get_ref_or_insert(|| {
calls.fetch_add(1, Ordering::SeqCst);
ConnectionTimeoutMs(7)
});
assert_eq!(missing.0, 7u64);
assert_eq!(calls.load(Ordering::SeqCst), 1);
}
#[test]
fn get_arc_or_insert_uses_existing_or_inserts_once() {
let ext = Extensions::new();
ext.insert_arc(Arc::new(TraceNote(String::from("stored"))));
let calls = AtomicUsize::new(0);
let existing = ext.self_get_arc_or_insert(|| {
calls.fetch_add(1, Ordering::SeqCst);
Arc::new(TraceNote(String::from("new")))
});
assert_eq!(existing.0.as_str(), "stored");
assert_eq!(calls.load(Ordering::SeqCst), 0);
let missing = ext.self_get_arc_or_insert(|| {
calls.fetch_add(1, Ordering::SeqCst);
Arc::new(RetryBudget(11))
});
assert_eq!(missing.0, 11u32);
assert_eq!(calls.load(Ordering::SeqCst), 1);
}
#[test]
fn iter_all_exposes_all_items_in_insert_order() {
let ext = Extensions::new();
ext.insert(HealthSignal(1));
ext.insert(FeatureToggle(true));
ext.insert(HealthSignal(2));
let type_ids: Vec<TypeId> = ext
.self_iter_all()
.map(TypeErasedExtension::type_id)
.collect();
assert_eq!(
type_ids,
vec![
TypeId::of::<HealthSignal>(),
TypeId::of::<FeatureToggle>(),
TypeId::of::<HealthSignal>()
]
);
}
#[test]
fn iter_for_missing_type_is_empty() {
let ext = Extensions::new();
ext.insert(HealthSignal(1));
assert_eq!(ext.self_iter_ref::<TraceNote>().count(), 0);
assert_eq!(ext.self_iter_arc::<TraceNote>().count(), 0);
}
#[test]
fn iter_ref_returns_items_for_present_type_in_newest_to_oldest_order() {
let ext = Extensions::new();
ext.insert(TraceNote(String::from("first")));
ext.insert(HealthSignal(9));
ext.insert(TraceNote(String::from("second")));
let output: Vec<&str> = ext
.self_iter_ref::<TraceNote>()
.map(|it| it.0.as_str())
.collect();
assert_eq!(output, vec!["second", "first"]);
}
#[test]
fn iter_arc_returns_items_for_present_type_in_newest_to_oldest_order() {
let ext = Extensions::new();
ext.insert(TraceNote(String::from("first")));
ext.insert(HealthSignal(9));
ext.insert(TraceNote(String::from("second")));
let output: Vec<String> = ext
.self_iter_arc::<TraceNote>()
.map(|arc| arc.0.clone())
.collect();
assert_eq!(output, vec!["second".to_owned(), "first".to_owned()]);
}
#[test]
fn type_erased_new_supports_downcast_ref_and_cloned_downcast() {
let ext = TypeErasedExtension::new(TraceNote(String::from("hello")));
assert_eq!(ext.type_id(), TypeId::of::<TraceNote>());
assert_eq!(
ext.downcast_ref::<TraceNote>().map(|it| it.0.as_str()),
Some("hello")
);
assert_eq!(
ext.cloned_downcast::<TraceNote>()
.as_deref()
.map(|it| it.0.as_str()),
Some("hello")
);
assert!(ext.downcast_ref::<RetryBudget>().is_none());
assert!(ext.cloned_downcast::<RetryBudget>().is_none());
}
#[test]
fn type_erased_new_arc_supports_all_downcasts() {
let ext = TypeErasedExtension::new_arc(Arc::new(TraceNote(String::from("hello"))));
assert_eq!(ext.type_id(), TypeId::of::<TraceNote>());
assert_eq!(
ext.downcast_ref::<TraceNote>().map(|it| it.0.as_str()),
Some("hello")
);
assert_eq!(
ext.cloned_downcast::<TraceNote>()
.as_deref()
.map(|it| it.0.as_str()),
Some("hello")
);
assert!(ext.downcast_ref::<RetryBudget>().is_none());
assert!(ext.cloned_downcast::<RetryBudget>().is_none());
}
#[test]
fn extensions_ref_blanket_impls_forward_to_underlying_extensions() {
let base = Extensions::new();
base.insert(RetryBudget(7));
let by_ref: &Extensions = &base;
assert_eq!(
by_ref.extensions().get_ref::<RetryBudget>(),
Some(&RetryBudget(7))
);
let mut base_for_mut = base.clone();
let by_mut_ref: &mut Extensions = &mut base_for_mut;
assert_eq!(
by_mut_ref.extensions().get_ref::<RetryBudget>(),
Some(&RetryBudget(7))
);
let boxed = Box::new(base.clone());
assert_eq!(
boxed.extensions().get_ref::<RetryBudget>(),
Some(&RetryBudget(7))
);
let pinned = Pin::new(Box::new(base.clone()));
assert_eq!(
pinned.extensions().get_ref::<RetryBudget>(),
Some(&RetryBudget(7))
);
let arced = Arc::new(base);
assert_eq!(
arced.extensions().get_ref::<RetryBudget>(),
Some(&RetryBudget(7))
);
}
#[derive(Debug, Clone, PartialEq, Eq, Extension)]
struct ConnSocketInfo(&'static str);
#[derive(Debug, Clone, PartialEq, Eq, Extension)]
struct RequestId(u64);
#[test]
fn get_finds_local() {
let req = Extensions::new();
req.insert(RequestId(42));
assert_eq!(req.get_ref::<RequestId>(), Some(&RequestId(42)));
}
#[test]
fn get_walks_parent_chain() {
let req = Extensions::new();
req.insert(RequestId(7));
let resp = req.fork();
assert_eq!(resp.get_ref::<RequestId>(), Some(&RequestId(7)));
}
#[test]
fn local_shadows_parent() {
let req = Extensions::new();
req.insert(RequestId(7));
let attempt = req.fork();
attempt.insert(RequestId(99));
assert_eq!(attempt.get_ref::<RequestId>(), Some(&RequestId(99)));
}
#[test]
fn fork_isolates_writes() {
let req = Extensions::new();
req.insert(RequestId(1));
let attempt = req.fork();
attempt.insert(RequestId(2));
assert_eq!(req.get_ref::<RequestId>(), Some(&RequestId(1)));
}
#[test]
fn ingress_view_walks_parent() {
let conn_ext = Extensions::new();
conn_ext.insert(ConnSocketInfo("client-in"));
let req = Extensions::new();
req.insert(Ingress(conn_ext));
assert_eq!(
req.ingress().and_then(|i| i.get_ref::<ConnSocketInfo>()),
Some(&ConnSocketInfo("client-in"))
);
}
#[test]
fn egress_view_walks_parent() {
let conn_ext = Extensions::new();
conn_ext.insert(ConnSocketInfo("egress-side"));
let req = Extensions::new();
req.insert(Egress(conn_ext));
assert_eq!(
req.egress().and_then(|e| e.get_ref::<ConnSocketInfo>()),
Some(&ConnSocketInfo("egress-side"))
);
}
#[test]
fn ingress_egress_disambiguate_in_mitm() {
let in_conn = Extensions::new();
in_conn.insert(ConnSocketInfo("in"));
let out_conn = Extensions::new();
out_conn.insert(ConnSocketInfo("out"));
let req = Extensions::new();
req.insert(Ingress(in_conn));
req.insert(Egress(out_conn));
assert_eq!(
req.ingress().and_then(|i| i.get_ref::<ConnSocketInfo>()),
Some(&ConnSocketInfo("in"))
);
assert_eq!(
req.egress().and_then(|e| e.get_ref::<ConnSocketInfo>()),
Some(&ConnSocketInfo("out"))
);
}
#[test]
fn egress_view_walks_through_parent_to_find_wrapper() {
let conn_ext = Extensions::new();
conn_ext.insert(ConnSocketInfo("inside-parent"));
let req = Extensions::new();
req.insert(Egress(conn_ext));
let resp = req.fork();
assert_eq!(
resp.egress().and_then(|e| e.get_ref::<ConnSocketInfo>()),
Some(&ConnSocketInfo("inside-parent"))
);
}
#[test]
fn ingress_egress_return_none_when_absent() {
let req = Extensions::new();
assert!(req.ingress().is_none());
assert!(req.egress().is_none());
}
#[test]
fn iter_ref_yields_local_then_parent_newest_to_oldest() {
let parent = Extensions::new();
parent.insert(RequestId(1));
parent.insert(RequestId(2));
let child = parent.fork();
child.insert(RequestId(3));
child.insert(RequestId(4));
let ids: Vec<_> = child.iter_ref::<RequestId>().map(|r| r.0).collect();
assert_eq!(ids, vec![4, 3, 2, 1]);
}
#[test]
fn iter_ref_walks_egress_and_ingress_wrappers_inline() {
let conn_in = Extensions::new();
conn_in.insert(RequestId(10));
conn_in.insert(RequestId(11));
let conn_out = Extensions::new();
conn_out.insert(RequestId(20));
let req = Extensions::new();
req.insert(RequestId(1));
req.insert(Ingress(conn_in));
req.insert(Egress(conn_out));
let ids: Vec<_> = req.iter_ref::<RequestId>().map(|r| r.0).collect();
assert_eq!(ids, vec![20, 11, 10, 1]);
}
#[test]
fn local_direct_after_wrapper_shadows_wrapper() {
let conn = Extensions::new();
conn.insert(RequestId(99));
let req = Extensions::new();
req.insert(Ingress(conn));
req.insert(RequestId(1));
assert_eq!(req.get_ref::<RequestId>(), Some(&RequestId(1)));
}
#[test]
fn wrapper_after_local_direct_shadows_direct() {
let conn = Extensions::new();
conn.insert(RequestId(99));
let req = Extensions::new();
req.insert(RequestId(1));
req.insert(Ingress(conn));
assert_eq!(req.get_ref::<RequestId>(), Some(&RequestId(99)));
}
#[test]
fn iter_ref_first_matches_get_ref() {
let parent = Extensions::new();
parent.insert(RequestId(1));
let child = parent.fork();
child.insert(RequestId(2));
assert_eq!(
child.iter_ref::<RequestId>().next(),
child.get_ref::<RequestId>()
);
}
#[test]
fn get_many_present_and_absent() {
let ext = Extensions::new();
ext.insert(RequestId(7));
ext.insert(ConnSocketInfo("a"));
let (id, sock, toggle) = ext.get_many_ref::<(RequestId, ConnSocketInfo, FeatureToggle)>();
assert_eq!(id, Some(&RequestId(7)));
assert_eq!(sock, Some(&ConnSocketInfo("a")));
assert_eq!(toggle, None);
}
#[test]
fn get_many_each_slot_matches_get_ref() {
let ext = Extensions::new();
ext.insert(RequestId(1));
ext.insert(RequestId(2)); ext.insert(ConnSocketInfo("x"));
let (id, sock) = ext.get_many_ref::<(RequestId, ConnSocketInfo)>();
assert_eq!(id, ext.get_ref::<RequestId>());
assert_eq!(sock, ext.get_ref::<ConnSocketInfo>());
assert_eq!(id, Some(&RequestId(2)));
}
#[test]
fn get_many_walks_parent_chain() {
let parent = Extensions::new();
parent.insert(RequestId(1));
let child = parent.fork();
child.insert(ConnSocketInfo("x"));
let (id, sock) = child.get_many_ref::<(RequestId, ConnSocketInfo)>();
assert_eq!(id, Some(&RequestId(1)));
assert_eq!(sock, Some(&ConnSocketInfo("x")));
}
#[test]
fn with_base_and_chain() {
let base = Extensions::new();
base.insert(RequestId(1));
base.insert(ConnSocketInfo("base"));
let req = Extensions::new();
req.insert(RequestId(2));
req.insert(TraceNote("mid".to_owned()));
let req_retry = req.fork();
req_retry.insert(RequestId(3));
let combined = req_retry.with_base(&base);
assert_eq!(combined.get_ref::<RequestId>(), Some(&RequestId(3)));
assert_eq!(
combined.get_ref::<TraceNote>(),
Some(&TraceNote("mid".to_owned())),
);
assert_eq!(
combined.get_ref::<ConnSocketInfo>(),
Some(&ConnSocketInfo("base")),
);
assert_eq!(base.get_ref::<RequestId>(), Some(&RequestId(1)));
assert!(req_retry.get_ref::<ConnSocketInfo>().is_none());
}
#[test]
fn get_many_walks_wrappers() {
let conn = Extensions::new();
conn.insert(ConnSocketInfo("in"));
let req = Extensions::new();
req.insert(RequestId(7));
req.insert(Ingress(conn));
let (id, sock) = req.get_many_ref::<(RequestId, ConnSocketInfo)>();
assert_eq!(id, Some(&RequestId(7)));
assert_eq!(sock, Some(&ConnSocketInfo("in")));
}
#[derive(FromExtensions)]
struct GatherView<'a> {
id: Option<&'a RequestId>,
sock: Option<&'a ConnSocketInfo>,
toggle: Option<&'a FeatureToggle>,
}
#[test]
fn derive_from_extensions_gathers_pieces() {
let ext = Extensions::new();
ext.insert(RequestId(7));
ext.insert(ConnSocketInfo("a"));
let view = GatherView::from_extensions(&ext);
assert_eq!(view.id, Some(&RequestId(7)));
assert_eq!(view.sock, Some(&ConnSocketInfo("a")));
assert_eq!(view.toggle, None);
assert_eq!(view.id, ext.get_ref::<RequestId>());
}
#[test]
fn derive_from_extensions_walks_parent() {
let parent = Extensions::new();
parent.insert(RequestId(1));
let child = parent.fork();
child.insert(ConnSocketInfo("x"));
let view = GatherView::from_extensions(&child);
assert_eq!(view.id, Some(&RequestId(1)));
assert_eq!(view.sock, Some(&ConnSocketInfo("x")));
assert_eq!(view.toggle, None);
}
#[test]
fn get_many_arc_returns_owned_arcs() {
let ext = Extensions::new();
ext.insert(RequestId(7));
let (id, sock) = ext.get_many_arc::<(RequestId, ConnSocketInfo)>();
assert_eq!(id.as_deref(), Some(&RequestId(7)));
assert_eq!(sock, None);
}
#[derive(FromExtensions)]
struct MixedView<'a> {
id_ref: Option<&'a RequestId>,
sock_arc: Option<Arc<ConnSocketInfo>>,
}
#[test]
fn derive_from_extensions_mixed_ref_and_arc() {
let ext = Extensions::new();
ext.insert(RequestId(7));
ext.insert(ConnSocketInfo("a"));
let view = MixedView::from_extensions(&ext);
assert_eq!(view.id_ref, Some(&RequestId(7)));
assert_eq!(view.sock_arc.as_deref(), Some(&ConnSocketInfo("a")));
}
#[derive(FromExtensions)]
struct AllArc {
id_ref: Option<Arc<RequestId>>,
sock_arc: Option<Arc<ConnSocketInfo>>,
}
#[test]
fn derive_from_extensions_all_arc() {
let ext = Extensions::new();
ext.insert(RequestId(7));
ext.insert(ConnSocketInfo("a"));
let view = AllArc::from_extensions(&ext);
assert_eq!(view.id_ref.as_deref(), Some(&RequestId(7)));
assert_eq!(view.sock_arc.as_deref(), Some(&ConnSocketInfo("a")));
}
#[derive(FromExtensions)]
struct RankedView<'a> {
id: Option<(&'a RequestId, usize)>,
sock: Option<(&'a ConnSocketInfo, usize)>,
toggle: Option<(&'a FeatureToggle, usize)>,
}
#[test]
fn derive_from_extensions_captures_rank() {
let ext = Extensions::new();
ext.insert(RequestId(7)); ext.insert(ConnSocketInfo("a"));
let view = RankedView::from_extensions(&ext);
assert_eq!(view.sock, Some((&ConnSocketInfo("a"), 0)));
assert_eq!(view.id, Some((&RequestId(7), 1)));
assert_eq!(view.toggle, None);
assert!(view.sock.unwrap().1 < view.id.unwrap().1);
}
#[test]
fn derive_from_extensions_rank_arc_variant() {
#[derive(FromExtensions)]
struct RankedArc {
id: Option<(Arc<RequestId>, usize)>,
}
let ext = Extensions::new();
ext.insert(ConnSocketInfo("a"));
ext.insert(RequestId(7));
let view = RankedArc::from_extensions(&ext);
let (id, rank) = view.id.expect("present");
assert_eq!(&*id, &RequestId(7));
assert_eq!(rank, 0);
}
#[derive(Debug, PartialEq, Eq, FromExtensions)]
enum AnyOf<'a> {
Req(&'a RequestId),
Sock(&'a ConnSocketInfo),
}
#[test]
fn derive_from_extensions_enum_newest_wins() {
let ext = Extensions::new();
ext.insert(ConnSocketInfo("a"));
ext.insert(RequestId(7));
assert_eq!(
AnyOf::from_extensions(&ext),
Some(AnyOf::Req(&RequestId(7)))
);
let ext = Extensions::new();
ext.insert(RequestId(7));
ext.insert(ConnSocketInfo("a"));
assert_eq!(
AnyOf::from_extensions(&ext),
Some(AnyOf::Sock(&ConnSocketInfo("a")))
);
let ext = Extensions::new();
ext.insert(RequestId(7));
assert_eq!(
AnyOf::from_extensions(&ext),
Some(AnyOf::Req(&RequestId(7)))
);
let ext = Extensions::new();
ext.insert(FeatureToggle(true));
assert_eq!(AnyOf::from_extensions(&ext), None);
}
#[derive(FromExtensions)]
struct ConfigView<'a> {
toggle: Option<&'a FeatureToggle>,
either: Option<AnyOf<'a>>,
}
#[derive(Debug, PartialEq, Eq, FromExtensions)]
enum SameType<'a> {
First(&'a RequestId),
Second(&'a RequestId),
}
#[test]
fn derive_from_extensions_enum_same_type_ties_to_earlier_variant() {
let ext = Extensions::new();
ext.insert(RequestId(7));
assert_eq!(
SameType::from_extensions(&ext),
Some(SameType::First(&RequestId(7)))
);
}
#[test]
fn derive_from_extensions_nested_group_field() {
let ext = Extensions::new();
ext.insert(FeatureToggle(true));
ext.insert(RequestId(7));
ext.insert(ConnSocketInfo("a"));
let view = ConfigView::from_extensions(&ext);
assert_eq!(view.toggle, Some(&FeatureToggle(true)));
assert_eq!(view.either, Some(AnyOf::Sock(&ConnSocketInfo("a"))));
let ext = Extensions::new();
ext.insert(FeatureToggle(false));
let view = ConfigView::from_extensions(&ext);
assert_eq!(view.toggle, Some(&FeatureToggle(false)));
assert_eq!(view.either, None);
}
#[test]
fn derive_from_extensions_enum_newest_wins_across_parent() {
let parent = Extensions::new();
parent.insert(RequestId(7));
let child = parent.fork();
child.insert(ConnSocketInfo("a"));
assert_eq!(
AnyOf::from_extensions(&child),
Some(AnyOf::Sock(&ConnSocketInfo("a")))
);
let parent = Extensions::new();
parent.insert(ConnSocketInfo("a"));
let child = parent.fork();
child.insert(RequestId(7));
assert_eq!(
AnyOf::from_extensions(&child),
Some(AnyOf::Req(&RequestId(7)))
);
}
}