Documentation
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()) }
}