use crate::decider::Decider;
use std::{
future::Future,
marker::PhantomData,
pin::Pin,
task::{Context, Poll},
};
use tower::{Layer, Service};
#[derive(Clone, Debug)]
pub struct ErrorLayer<'a, D, G> {
decider: D,
generator: G,
_phantom: PhantomData<&'a ()>,
}
impl<'a> ErrorLayer<'a, (), ()> {
pub fn builder() -> Self {
Self {
decider: (),
generator: (),
_phantom: PhantomData,
}
}
}
impl<'a, D, G> ErrorLayer<'a, D, G> {
pub fn new(decider: D, generator: G) -> Self {
Self {
decider,
generator,
_phantom: PhantomData,
}
}
pub fn with_decider<ND>(self, decider: ND) -> ErrorLayer<'a, ND, G> {
ErrorLayer {
decider,
generator: self.generator,
_phantom: PhantomData,
}
}
pub fn with_generator<NG>(self, generator: NG) -> ErrorLayer<'a, D, NG> {
ErrorLayer {
decider: self.decider,
generator,
_phantom: PhantomData,
}
}
}
impl<'a, D, G, S> Layer<S> for ErrorLayer<'a, D, G>
where
D: Clone,
G: Clone,
{
type Service = ErrorService<'a, D, G, S>;
fn layer(&self, inner: S) -> Self::Service {
ErrorService {
inner,
decider: self.decider.clone(),
generator: self.generator.clone(),
_phantom: PhantomData,
}
}
}
#[derive(Clone, Debug)]
pub struct ErrorService<'a, D, G, S> {
inner: S,
decider: D,
generator: G,
_phantom: PhantomData<&'a ()>,
}
impl<'a, D, G, S, R> Service<R> for ErrorService<'a, D, G, S>
where
D: Decider<R> + Clone,
G: Fn(&R) -> S::Error + Clone,
S: Service<R> + Send,
S::Future: Send + 'a,
S::Error: Send + 'a,
{
type Response = S::Response;
type Error = S::Error;
type Future = ErrorFuture<'a, R, S>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, request: R) -> Self::Future {
if self.decider.decide(&request) {
let error = (self.generator)(&request);
return Box::pin(async move { Err(error) });
}
Box::pin(self.inner.call(request))
}
}
type ErrorFuture<'a, R, S> = Pin<
Box<
dyn Future<Output = Result<<S as Service<R>>::Response, <S as Service<R>>::Error>>
+ Send
+ 'a,
>,
>;
#[cfg(test)]
mod tests {
use super::*;
use crate::test_utils::*;
#[tokio::test]
async fn error_success() {
let layer = ErrorLayer::new(0.0, |_: &()| String::from("error"));
let mut service = layer.layer(DummyService);
for _ in 0..1000 {
let res = service.call(()).await;
assert_eq!(res.unwrap(), String::from("ok"));
}
}
#[tokio::test]
async fn error_fail() {
let layer = ErrorLayer::new(1.0, |_: &()| String::from("error"));
let mut service = layer.layer(DummyService);
for _ in 0..1000 {
let res = service.call(()).await;
assert_eq!(res.unwrap_err(), String::from("error"));
}
}
}