#![cfg_attr(
not(test),
deny(
clippy::unwrap_used,
clippy::expect_used,
clippy::panic,
clippy::unreachable,
clippy::todo,
clippy::unimplemented,
clippy::indexing_slicing,
clippy::string_slice,
clippy::arithmetic_side_effects,
)
)]
use std::future::Future;
use std::pin::Pin;
use std::task::{Context, Poll};
use axum::body::Body;
use axum::http::Response;
use pin_project_lite::pin_project;
pin_project! {
#[project = ShortCircuitFutureProj]
pub enum ShortCircuitFuture<F> {
ShortCircuit { response: Option<Response<Body>> },
Forward {
#[pin]
inner: F,
},
}
}
impl<F> ShortCircuitFuture<F> {
pub(crate) const fn short_circuit(response: Response<Body>) -> Self {
Self::ShortCircuit {
response: Some(response),
}
}
pub(crate) const fn forward(inner: F) -> Self {
Self::Forward { inner }
}
}
impl<F, E> Future for ShortCircuitFuture<F>
where
F: Future<Output = Result<Response<Body>, E>>,
{
type Output = Result<Response<Body>, E>;
#[allow(
clippy::expect_used,
reason = "unreachable: future not polled after Ready"
)]
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
match self.project() {
ShortCircuitFutureProj::ShortCircuit { response } => Poll::Ready(Ok(response
.take()
.expect("ShortCircuitFuture polled after completion"))),
ShortCircuitFutureProj::Forward { inner } => inner.poll(cx),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use axum::http::StatusCode;
#[test]
fn the_future_is_named_not_boxed() {
let name = std::any::type_name::<ShortCircuitFuture<std::future::Ready<()>>>();
assert!(
!name.contains("Box"),
"ShortCircuitFuture must stay a named, unboxed future; got {name}"
);
}
type Never = std::future::Pending<Result<Response<Body>, std::convert::Infallible>>;
#[tokio::test]
async fn short_circuit_resolves_with_the_gate_response_without_the_inner_future() {
let response = Response::builder()
.status(StatusCode::BAD_REQUEST)
.body(Body::empty())
.expect("response builds");
let got = ShortCircuitFuture::<Never>::short_circuit(response)
.await
.expect("infallible");
assert_eq!(got.status(), StatusCode::BAD_REQUEST);
}
#[tokio::test]
async fn forward_resolves_with_the_inner_futures_output() {
let inner = std::future::ready(Ok::<_, std::convert::Infallible>(
Response::builder()
.status(StatusCode::IM_A_TEAPOT)
.body(Body::empty())
.expect("response builds"),
));
let got = ShortCircuitFuture::forward(inner)
.await
.expect("infallible");
assert_eq!(got.status(), StatusCode::IM_A_TEAPOT);
}
}