use std::{
future::Future,
pin::Pin,
sync::Arc,
task::{Context, Poll},
};
use super::config::ClerkAuthLayerConfig;
use super::verification::{VerifiedRequest, Verifier, verify_request};
use axum::body::Body;
use axum::http::{Request, Response};
use tower::{Layer, Service};
#[cfg(not(target_arch = "wasm32"))]
trait MaybeSend: Send {}
#[cfg(not(target_arch = "wasm32"))]
impl<T: Send + ?Sized> MaybeSend for T {}
#[cfg(target_arch = "wasm32")]
trait MaybeSend {}
#[cfg(target_arch = "wasm32")]
impl<T: ?Sized> MaybeSend for T {}
#[cfg(not(target_arch = "wasm32"))]
type BoxServiceFuture<E> = Pin<Box<dyn Future<Output = Result<Response<Body>, E>> + Send>>;
#[cfg(target_arch = "wasm32")]
type BoxServiceFuture<E> = Pin<Box<dyn Future<Output = Result<Response<Body>, E>>>>;
#[cfg(all(target_arch = "wasm32", feature = "worker"))]
type ServiceFuture<E> = send_wrapper::SendWrapper<BoxServiceFuture<E>>;
#[cfg(not(all(target_arch = "wasm32", feature = "worker")))]
type ServiceFuture<E> = BoxServiceFuture<E>;
#[cfg(all(target_arch = "wasm32", feature = "worker"))]
fn service_future<E>(future: BoxServiceFuture<E>) -> ServiceFuture<E> {
send_wrapper::SendWrapper::new(future)
}
#[cfg(not(all(target_arch = "wasm32", feature = "worker")))]
fn service_future<E>(future: BoxServiceFuture<E>) -> ServiceFuture<E> {
future
}
#[derive(Clone)]
pub struct ClerkAuthLayer {
inner: Arc<Inner>,
}
impl std::fmt::Debug for ClerkAuthLayer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ClerkAuthLayer").finish_non_exhaustive()
}
}
struct Inner {
verifier: Verifier,
}
impl ClerkAuthLayer {
pub fn new(secret_key: impl Into<String>) -> Result<Self, crate::core::ClerkError> {
Self::from_config(ClerkAuthLayerConfig::new(secret_key))
}
pub fn from_env() -> Result<Self, crate::core::ClerkError> {
Self::from_config(ClerkAuthLayerConfig::from_env()?)
}
pub fn from_config(config: ClerkAuthLayerConfig) -> Result<Self, crate::core::ClerkError> {
let verifier = Verifier::new(config)?;
Ok(Self {
inner: Arc::new(Inner { verifier }),
})
}
}
impl<S> Layer<S> for ClerkAuthLayer {
type Service = ClerkAuthService<S>;
fn layer(&self, inner: S) -> Self::Service {
ClerkAuthService {
inner,
layer: self.clone(),
}
}
}
#[derive(Clone)]
pub struct ClerkAuthService<S> {
inner: S,
layer: ClerkAuthLayer,
}
impl<S: std::fmt::Debug> std::fmt::Debug for ClerkAuthService<S> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ClerkAuthService")
.field("inner", &self.inner)
.finish_non_exhaustive()
}
}
impl<S> Service<Request<Body>> for ClerkAuthService<S>
where
S: Service<Request<Body>, Response = Response<Body>> + Clone + MaybeSend + 'static,
S::Future: MaybeSend + 'static,
{
type Response = Response<Body>;
type Error = S::Error;
type Future = ServiceFuture<S::Error>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, req: Request<Body>) -> Self::Future {
let layer = self.layer.clone();
let clone = self.inner.clone();
let mut inner = std::mem::replace(&mut self.inner, clone);
service_future(Box::pin(async move {
match verify_request(req, layer.inner.verifier.clone()).await {
VerifiedRequest::Forward(req) => inner.call(req).await,
VerifiedRequest::Unavailable(response) => Ok(response),
}
}))
}
}