use std::{
any::Any,
collections::{BTreeMap, BTreeSet, HashMap},
sync::Arc,
};
#[cfg(feature = "bound-upstream-request-body")]
use bytes::Bytes;
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct AuthenticatedIdentity {
subject_id: String,
roles: BTreeSet<String>,
teams: BTreeSet<String>,
custom_claims: BTreeMap<String, String>,
}
impl AuthenticatedIdentity {
#[cfg(any(feature = "basic-auth-filter", feature = "policy-engine", test))]
pub(crate) fn new(
subject_id: String,
roles: impl IntoIterator<Item = String>,
teams: impl IntoIterator<Item = String>,
custom_claims: impl IntoIterator<Item = (String, String)>,
) -> Option<Self> {
(!subject_id.is_empty()).then(|| Self {
subject_id,
roles: roles.into_iter().collect(),
teams: teams.into_iter().collect(),
custom_claims: custom_claims.into_iter().collect(),
})
}
pub fn subject_id(&self) -> &str {
&self.subject_id
}
pub fn roles(&self) -> &BTreeSet<String> {
&self.roles
}
pub fn teams(&self) -> &BTreeSet<String> {
&self.teams
}
pub fn custom_claims(&self) -> &BTreeMap<String, String> {
&self.custom_claims
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub(crate) struct SelectedClusterApplication {
protocol: Option<Arc<str>>,
provider: Option<Arc<str>>,
}
impl SelectedClusterApplication {
pub(crate) fn new(protocol: Option<Arc<str>>, provider: Option<Arc<str>>) -> Option<Self> {
(protocol.is_some() || provider.is_some()).then_some(Self { protocol, provider })
}
pub(crate) fn protocol(&self) -> Option<&str> {
self.protocol.as_deref()
}
pub(crate) fn provider(&self) -> Option<&str> {
self.provider.as_deref()
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub(crate) struct BoundUpstream {
cluster: Arc<str>,
application_protocol: Option<Arc<str>>,
application_provider: Option<Arc<str>>,
}
#[cfg(feature = "bound-upstream-request-body")]
#[derive(Clone, Debug)]
pub(crate) struct BoundRequestBodyRewrite(pub(crate) Bytes);
#[cfg(feature = "upstream-binding")]
#[derive(Clone, Copy, Debug)]
pub(crate) struct BoundUpstreamFrozen;
impl BoundUpstream {
#[cfg(feature = "upstream-binding")]
pub(crate) fn new(
cluster: Arc<str>,
application_protocol: Option<Arc<str>>,
application_provider: Option<Arc<str>>,
) -> Self {
Self {
cluster,
application_protocol,
application_provider,
}
}
pub(crate) fn cluster(&self) -> &str {
&self.cluster
}
pub(crate) fn application_protocol(&self) -> Option<&str> {
self.application_protocol.as_deref()
}
pub(crate) fn application_provider(&self) -> Option<&str> {
self.application_provider.as_deref()
}
}
#[derive(Default)]
pub struct RequestExtensions(HashMap<std::any::TypeId, Box<dyn Any + Send + Sync>>);
impl RequestExtensions {
pub fn new() -> Self {
Self::default()
}
pub fn insert<T: Send + Sync + 'static>(&mut self, val: T) {
self.0.insert(std::any::TypeId::of::<T>(), Box::new(val));
}
pub fn get<T: Send + Sync + 'static>(&self) -> Option<&T> {
self.0
.get(&std::any::TypeId::of::<T>())
.and_then(|boxed| boxed.downcast_ref())
}
pub fn get_mut<T: Send + Sync + 'static>(&mut self) -> Option<&mut T> {
self.0
.get_mut(&std::any::TypeId::of::<T>())
.and_then(|boxed| boxed.downcast_mut())
}
pub fn get_or_insert_with<T: Send + Sync + 'static>(&mut self, f: impl FnOnce() -> T) -> &mut T {
#[expect(clippy::expect_used, reason = "downcast cannot fail after typed insert")]
self.0
.entry(std::any::TypeId::of::<T>())
.or_insert_with(|| Box::new(f()))
.downcast_mut()
.expect("type mismatch after insert")
}
pub fn remove<T: Send + Sync + 'static>(&mut self) -> Option<T> {
self.0
.remove(&std::any::TypeId::of::<T>())
.and_then(|boxed| boxed.downcast().ok())
.map(|boxed| *boxed)
}
}
#[cfg(test)]
#[expect(clippy::allow_attributes, reason = "blanket test suppressions")]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic, reason = "tests")]
mod tests {
use super::*;
#[test]
fn authenticated_identity_exposes_stable_collections_and_getters() {
let identity = AuthenticatedIdentity::new(
"alice".to_owned(),
["writer".to_owned(), "admin".to_owned(), "admin".to_owned()],
["platform".to_owned()],
[("tenant".to_owned(), "acme".to_owned())],
)
.expect("non-empty subject");
assert_eq!(identity.subject_id(), "alice");
assert_eq!(
identity.roles().iter().map(String::as_str).collect::<Vec<_>>(),
["admin", "writer"],
);
assert_eq!(
identity.teams().iter().map(String::as_str).collect::<Vec<_>>(),
["platform"]
);
assert_eq!(identity.custom_claims().get("tenant").map(String::as_str), Some("acme"));
}
#[test]
fn authenticated_identity_rejects_empty_subject() {
assert!(
AuthenticatedIdentity::new(
String::new(),
std::iter::empty(),
std::iter::empty(),
std::iter::empty(),
)
.is_none(),
);
}
#[test]
fn default_is_empty() {
let ext = RequestExtensions::default();
assert!(ext.get::<String>().is_none(), "default should contain no values");
}
#[test]
fn selected_cluster_application_absent_when_both_fields_absent() {
assert!(
SelectedClusterApplication::new(None, None).is_none(),
"an untagged cluster must not produce a metadata value"
);
}
#[test]
fn selected_cluster_application_present_with_protocol_only() {
let app = SelectedClusterApplication::new(Some(Arc::from("openai_chat_completions")), None)
.expect("a protocol-only cluster should produce a value");
assert_eq!(
app.protocol(),
Some("openai_chat_completions"),
"protocol should read back"
);
assert!(app.provider().is_none(), "an absent provider should read back as None");
}
#[test]
fn selected_cluster_application_present_with_provider_only() {
let app = SelectedClusterApplication::new(None, Some(Arc::from("vllm")))
.expect("a provider-only cluster should produce a value");
assert!(app.protocol().is_none(), "an absent protocol should read back as None");
assert_eq!(app.provider(), Some("vllm"), "provider should read back");
}
#[test]
fn selected_cluster_application_exposes_both_fields() {
let app = SelectedClusterApplication::new(Some(Arc::from("openai_responses")), Some(Arc::from("openai")))
.expect("a fully tagged cluster should produce a value");
assert_eq!(app.protocol(), Some("openai_responses"), "protocol should read back");
assert_eq!(app.provider(), Some("openai"), "provider should read back");
}
#[cfg(feature = "upstream-binding")]
#[test]
fn bound_upstream_exposes_cluster_and_metadata() {
let bound = BoundUpstream::new(
Arc::from("inference-backend"),
Some(Arc::from("openai_responses")),
Some(Arc::from("openai")),
);
assert_eq!(bound.cluster(), "inference-backend", "cluster name should read back");
assert_eq!(
bound.application_protocol(),
Some("openai_responses"),
"protocol should read back"
);
assert_eq!(
bound.application_provider(),
Some("openai"),
"provider should read back"
);
}
#[cfg(feature = "upstream-binding")]
#[test]
fn bound_upstream_allows_untagged_cluster() {
let bound = BoundUpstream::new(Arc::from("backend"), None, None);
assert_eq!(
bound.cluster(),
"backend",
"an untagged cluster still produces a binding (the cluster name is mandatory)"
);
assert!(
bound.application_protocol().is_none(),
"absent protocol reads back as None"
);
assert!(
bound.application_provider().is_none(),
"absent provider reads back as None"
);
}
#[test]
fn insert_and_get() {
let mut ext = RequestExtensions::new();
ext.insert(42_u32);
assert_eq!(ext.get::<u32>(), Some(&42), "should retrieve inserted value");
}
#[test]
fn insert_and_get_mut() {
let mut ext = RequestExtensions::new();
ext.insert("hello".to_owned());
if let Some(val) = ext.get_mut::<String>() {
val.push_str(" world");
}
assert_eq!(
ext.get::<String>().map(String::as_str),
Some("hello world"),
"get_mut should allow mutation"
);
}
#[test]
fn multiple_types_coexist() {
let mut ext = RequestExtensions::new();
ext.insert(1_u32);
ext.insert("text".to_owned());
ext.insert(1.5_f64);
assert_eq!(ext.get::<u32>(), Some(&1), "u32 should be present");
assert_eq!(
ext.get::<String>().map(String::as_str),
Some("text"),
"String should be present"
);
assert_eq!(ext.get::<f64>(), Some(&1.5), "f64 should be present");
}
#[test]
fn insert_same_type_overwrites() {
let mut ext = RequestExtensions::new();
ext.insert(1_u32);
ext.insert(2_u32);
assert_eq!(ext.get::<u32>(), Some(&2), "second insert should overwrite first");
}
#[test]
fn remove_returns_owned_value() {
let mut ext = RequestExtensions::new();
ext.insert(99_u32);
let removed = ext.remove::<u32>();
assert_eq!(removed, Some(99), "remove should return the stored value");
assert!(ext.get::<u32>().is_none(), "value should be gone after remove");
}
#[test]
fn remove_absent_returns_none() {
let mut ext = RequestExtensions::new();
assert!(ext.remove::<u32>().is_none(), "removing absent type should return None");
}
#[test]
fn get_or_insert_with_creates_when_absent() {
let mut ext = RequestExtensions::new();
let val = ext.get_or_insert_with(|| 42_u32);
assert_eq!(*val, 42, "should create value when absent");
}
#[test]
fn get_or_insert_with_returns_existing() {
let mut ext = RequestExtensions::new();
ext.insert(10_u32);
let val = ext.get_or_insert_with(|| 42_u32);
assert_eq!(*val, 10, "should return existing value without calling factory");
}
#[test]
fn get_wrong_type_returns_none() {
let mut ext = RequestExtensions::new();
ext.insert(42_u32);
assert!(ext.get::<String>().is_none(), "wrong type should return None");
}
#[test]
fn newtypes_are_independent() {
struct FilterAState(u32);
struct FilterBState(u32);
let mut ext = RequestExtensions::new();
ext.insert(FilterAState(1));
ext.insert(FilterBState(2));
assert_eq!(
ext.get::<FilterAState>().map(|s| s.0),
Some(1),
"FilterAState should be 1"
);
assert_eq!(
ext.get::<FilterBState>().map(|s| s.0),
Some(2),
"FilterBState should be 2"
);
}
}