use std::sync::Arc;
use crate::context::Extensions;
pub(crate) type ExtensionBridge =
Arc<dyn Fn(&axum::http::Extensions, &mut Extensions) + Send + Sync>;
pub(crate) fn extension_bridge<T>() -> ExtensionBridge
where
T: Clone + Send + Sync + 'static,
{
Arc::new(|from, to| {
if let Some(value) = from.get::<T>() {
to.insert(value.clone());
}
})
}
pub(crate) fn apply_extension_bridges(
bridges: &[ExtensionBridge],
from: &axum::http::Extensions,
to: &mut Extensions,
) {
for bridge in bridges {
bridge(from, to);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Debug, Clone, PartialEq)]
struct Identity(&'static str);
#[derive(Debug, Clone, PartialEq)]
struct TraceId(u32);
#[derive(Debug, Clone, PartialEq)]
struct NeverRegistered;
#[test]
fn a_registered_type_crosses_and_an_unregistered_one_does_not() {
let mut from = axum::http::Extensions::new();
from.insert(Identity("agent-7"));
from.insert(NeverRegistered);
let mut to = Extensions::new();
apply_extension_bridges(&[extension_bridge::<Identity>()], &from, &mut to);
assert_eq!(to.get::<Identity>(), Some(&Identity("agent-7")));
assert!(
to.get::<NeverRegistered>().is_none(),
"only registered types cross the boundary"
);
}
#[test]
fn a_missing_type_is_skipped() {
let from = axum::http::Extensions::new();
let mut to = Extensions::new();
apply_extension_bridges(&[extension_bridge::<Identity>()], &from, &mut to);
assert!(to.get::<Identity>().is_none());
}
#[test]
fn every_registered_type_is_applied() {
let mut from = axum::http::Extensions::new();
from.insert(Identity("agent-7"));
from.insert(TraceId(42));
let mut to = Extensions::new();
apply_extension_bridges(
&[
extension_bridge::<Identity>(),
extension_bridge::<TraceId>(),
],
&from,
&mut to,
);
assert_eq!(to.get::<Identity>(), Some(&Identity("agent-7")));
assert_eq!(to.get::<TraceId>(), Some(&TraceId(42)));
}
}