use std::sync::Arc;
use axum::{
extract::{FromRequestParts, Request},
middleware::Next,
response::Response,
};
use crate::{
DI, DIK,
error::ServiceNotFoundError,
scope::{Scope, ScopeProvider},
};
impl<T: ?Sized + Send + Sync + 'static, S: Send + Sync> FromRequestParts<S> for DI<T> {
type Rejection = &'static str;
async fn from_request_parts(
parts: &mut axum::http::request::Parts,
_: &S,
) -> Result<Self, Self::Rejection> {
let Some(sp) = crate::spo() else {
return Err("Failed to get service provider in DI extractor");
};
let Some(scope) = parts.extensions.get::<Arc<Scope>>() else {
return Err("Failed to extract scope");
};
let Ok(service) = sp.dis(Arc::new(Arc::new(AxumExtensionsScopeProvider(
scope.clone(),
)))) else {
return Err("Failed to get service");
};
Ok(DI(service))
}
}
impl<T: ?Sized + Send + Sync + 'static, K: Send + Sync + 'static, S: Send + Sync>
FromRequestParts<S> for DIK<T, K>
{
type Rejection = &'static str;
async fn from_request_parts(
parts: &mut axum::http::request::Parts,
_: &S,
) -> Result<Self, Self::Rejection> {
let Some(sp) = crate::spo() else {
return Err("Failed to get service provider in DI extractor");
};
let Some(scope) = parts.extensions.get::<Arc<Scope>>() else {
return Err("Failed to extract scope");
};
let Ok(service) = sp.disk::<T, K>(Arc::new(Arc::new(AxumExtensionsScopeProvider(
scope.clone(),
)))) else {
return Err("Failed to get service");
};
Ok(DIK(service, std::marker::PhantomData))
}
}
pub async fn scope_middleware(mut req: Request, next: Next) -> Response {
req.extensions_mut().insert(Arc::new(Scope::new()));
next.run(req).await
}
pub struct AxumExtensionsScopeProvider(pub Arc<Scope>);
impl<'a> ScopeProvider for AxumExtensionsScopeProvider {
fn provide_scope(&self) -> Result<Arc<Scope>, ServiceNotFoundError> { Ok(self.0.clone()) }
}