use std::any::{Any, TypeId};
use std::cell::RefCell;
use std::collections::HashMap;
use std::future::Future;
type Values = HashMap<TypeId, Box<dyn Any + Send + Sync>>;
tokio::task_local! {
static CONTEXT: RefCell<Values>;
}
pub fn set<T: Clone + Send + Sync + 'static>(value: T) -> bool {
CONTEXT
.try_with(|values| {
values
.borrow_mut()
.insert(TypeId::of::<T>(), Box::new(value));
})
.is_ok()
}
pub fn get<T: Clone + Send + Sync + 'static>() -> Option<T> {
CONTEXT
.try_with(|values| {
values
.borrow()
.get(&TypeId::of::<T>())
.and_then(|value| value.downcast_ref::<T>())
.cloned()
})
.ok()
.flatten()
}
pub fn remove<T: Send + Sync + 'static>() {
let _ = CONTEXT.try_with(|values| values.borrow_mut().remove(&TypeId::of::<T>()));
}
pub async fn scope<F: Future>(fut: F) -> F::Output {
CONTEXT.scope(RefCell::new(Values::new()), fut).await
}
#[derive(Debug, Clone)]
pub struct Current<T>(pub T);
impl<T, S> axum::extract::FromRequestParts<S> for Current<T>
where
T: Clone + Send + Sync + 'static,
S: Send + Sync,
{
type Rejection = crate::Error;
async fn from_request_parts(
_: &mut axum::http::request::Parts,
_: &S,
) -> Result<Self, crate::Error> {
get::<T>().map(Current).ok_or_else(|| {
anyhow::anyhow!(
"no `{}` in the request's context: set it in a middleware with renox::context::set",
std::any::type_name::<T>()
)
.into()
})
}
}
impl<T, S> axum::extract::OptionalFromRequestParts<S> for Current<T>
where
T: Clone + Send + Sync + 'static,
S: Send + Sync,
{
type Rejection = std::convert::Infallible;
async fn from_request_parts(
_: &mut axum::http::request::Parts,
_: &S,
) -> Result<Option<Self>, std::convert::Infallible> {
Ok(get::<T>().map(Current))
}
}
pub fn app() -> Option<crate::AppState> {
get::<crate::AppState>()
}
pub(crate) async fn scope_app<F: Future>(state: crate::AppState, fut: F) -> F::Output {
scope(async move {
set(state);
fut.await
})
.await
}
pub(crate) async fn middleware(
axum::extract::State(state): axum::extract::State<crate::AppState>,
req: axum::extract::Request,
next: axum::middleware::Next,
) -> axum::response::Response {
let info = RequestInfo {
method: req.method().to_string(),
path: req.uri().path().to_owned(),
id: req
.extensions()
.get::<crate::RequestId>()
.map(|id| id.0.clone())
.unwrap_or_default(),
ip: crate::ClientIp::of(&req).map(|ip| ip.to_string()),
};
scope_app(state, async move {
set(info);
next.run(req).await
})
.await
}
#[derive(Debug, Clone)]
pub(crate) struct RequestInfo {
pub method: String,
pub path: String,
pub id: String,
pub ip: Option<String>,
}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Clone, Debug, PartialEq)]
struct Team(i64);
#[tokio::test]
async fn values_live_in_their_scope() {
assert!(!set(Team(1)), "no context outside a scope");
assert_eq!(get::<Team>(), None);
scope(async {
assert!(set(Team(1)));
set(Team(2));
assert_eq!(get::<Team>(), Some(Team(2)));
assert_eq!(get::<String>(), None);
scope(async { assert_eq!(get::<Team>(), None, "a new scope starts empty") }).await;
remove::<Team>();
assert_eq!(get::<Team>(), None);
})
.await;
}
}