use std::convert::Infallible;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use axum::extract::{FromRequestParts, Request};
use axum::response::Response;
use http::request::Parts;
use tower::{Layer, Service};
use crate::event::{self, outcome_from_http_status, origin_from_headers, Actor, Event, EventHandle};
pub trait ActorResolver: Fn(&Parts) -> Actor + Send + Sync + 'static {}
impl<F: Fn(&Parts) -> Actor + Send + Sync + 'static> ActorResolver for F {}
#[derive(Clone)]
pub struct CurrentEvent(EventHandle);
impl CurrentEvent {
pub fn with<T>(&self, f: impl FnOnce(&mut Event) -> T) -> T {
self.0.with(f)
}
}
#[axum::async_trait]
impl<S: Send + Sync> FromRequestParts<S> for CurrentEvent {
type Rejection = Infallible;
async fn from_request_parts(_parts: &mut Parts, _state: &S) -> Result<Self, Infallible> {
Ok(CurrentEvent(event::current()))
}
}
pub struct EverscribeLayer<R, F> {
recorder: Arc<R>,
resolve: Arc<F>,
}
impl<R, F> Clone for EverscribeLayer<R, F> {
fn clone(&self) -> Self {
EverscribeLayer {
recorder: self.recorder.clone(),
resolve: self.resolve.clone(),
}
}
}
impl<R, F> EverscribeLayer<R, F> {
pub fn new(recorder: R, resolve: F) -> Self {
EverscribeLayer {
recorder: Arc::new(recorder),
resolve: Arc::new(resolve),
}
}
}
impl<S, R, F> Layer<S> for EverscribeLayer<R, F> {
type Service = EverscribeService<S, R, F>;
fn layer(&self, inner: S) -> Self::Service {
EverscribeService {
inner,
recorder: self.recorder.clone(),
resolve: self.resolve.clone(),
}
}
}
pub struct EverscribeService<S, R, F> {
inner: S,
recorder: Arc<R>,
resolve: Arc<F>,
}
impl<S: Clone, R, F> Clone for EverscribeService<S, R, F> {
fn clone(&self) -> Self {
EverscribeService {
inner: self.inner.clone(),
recorder: self.recorder.clone(),
resolve: self.resolve.clone(),
}
}
}
impl<S, R, F> Service<Request> for EverscribeService<S, R, F>
where
S: Service<Request, Response = Response> + Clone + Send + 'static,
S::Future: Send + 'static,
R: crate::recorder::Recorder + Send + Sync + 'static,
F: ActorResolver,
{
type Response = Response;
type Error = S::Error;
type Future = Pin<Box<dyn Future<Output = Result<Response, S::Error>> + Send>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, req: Request) -> Self::Future {
let clone = self.inner.clone();
let mut inner = std::mem::replace(&mut self.inner, clone);
let recorder = self.recorder.clone();
let resolve = self.resolve.clone();
Box::pin(async move {
let (parts, body) = req.into_parts();
let actor = (resolve)(&parts);
let origin = origin_from_headers(
|name| {
parts
.headers
.get(name)
.and_then(|v| v.to_str().ok())
.map(str::to_owned)
},
"",
);
let tmpl = Event {
actor,
origin,
..Event::default()
};
let (call_result, current) = event::scope(tmpl, None, async move {
inner.call(Request::from_parts(parts, body)).await
})
.await;
let resp = call_result?;
let outcome = outcome_from_http_status(resp.status().as_u16());
let recorder: &dyn event::Recorder = recorder.as_ref();
event::end(¤t, Some(&outcome), Some(recorder)).await;
Ok(resp)
})
}
}