#![doc = include_str!("../docs/tower.md")]
use std::{
borrow::Cow,
convert::Infallible,
fmt::{self, Display},
future::Future,
pin::{Pin, pin},
sync::Arc,
task::{Context, Poll},
};
use bytes::Bytes;
use tokio::sync::{mpsc, oneshot};
use topcoat_core::{
context::{Cx, try_request_context},
error::{Error, Result},
};
use tower::ServiceExt;
use crate::{
Body, BoxError, IntoPath, Layer, LayerFuture, Methods, Next, OwnedMethods, Path, Route,
RouteFuture, RouteId, Router,
request::{Request, parts},
response::Response,
};
pub struct TowerRoute<S> {
id: RouteId,
methods: OwnedMethods,
path: Cow<'static, Path>,
service: S,
}
impl<S> TowerRoute<S> {
#[must_use]
#[track_caller]
pub fn new(methods: impl Into<OwnedMethods>, path: impl IntoPath, service: S) -> Self {
Self {
id: RouteId::new(),
methods: methods.into(),
path: path.into_path(),
service,
}
}
#[must_use]
#[track_caller]
pub fn any(path: impl IntoPath, service: S) -> Self {
Self::new(Methods::Any, path, service)
}
}
impl<S, ResBody> Route for TowerRoute<S>
where
S: tower::Service<Request, Response = http::Response<ResBody>> + Clone + Send + Sync + 'static,
S::Error: Into<BoxError> + Send,
S::Future: Send,
ResBody: http_body::Body<Data = Bytes> + Send + 'static,
ResBody::Error: Into<BoxError>,
{
fn id(&self) -> RouteId {
self.id
}
fn methods(&self) -> Methods<'_> {
self.methods.as_methods()
}
fn path(&self) -> &Path {
&self.path
}
fn handle<'cx>(&'cx self, cx: &'cx Cx, body: Body) -> RouteFuture<'cx> {
let service = self.service.clone();
Box::pin(async move {
let request = Request::from_parts(parts(cx).clone(), body);
match service.oneshot(request).await {
Ok(response) => Ok(response.map(Body::new)),
Err(error) => Err(TowerServiceError(error.into()).into()),
}
})
}
}
pub struct TowerLayer<S> {
path: Option<Cow<'static, Path>>,
service: S,
}
impl<S> TowerLayer<S> {
#[must_use]
pub fn new<L>(layer: L) -> Self
where
L: tower::Layer<TowerNext, Service = S>,
{
Self {
path: None,
service: layer.layer(TowerNext::new()),
}
}
#[must_use]
#[track_caller]
pub fn at(mut self, path: impl IntoPath) -> Self {
self.path = Some(path.into_path());
self
}
}
impl<S, ResBody> Layer for TowerLayer<S>
where
S: tower::Service<Request, Response = http::Response<ResBody>> + Clone + Send + Sync + 'static,
S::Error: Into<BoxError> + Send,
S::Future: Send,
ResBody: http_body::Body<Data = Bytes> + Send + 'static,
ResBody::Error: Into<BoxError>,
{
fn path(&self) -> Option<&Path> {
self.path.as_deref()
}
fn handle<'a>(&'a self, cx: &'a Cx, body: Body, next: Next<'a>) -> LayerFuture<'a> {
let service = self.service.clone();
Box::pin(async move {
let parts = try_request_context::<http::request::Parts>(cx)
.expect("router context contains parts")
.clone();
let mut request = Request::from_parts(parts, body);
let (relay, mut chain_calls) = relay_channel();
request.extensions_mut().insert(relay);
let mut middleware = pin!(service.oneshot(request));
let (request, respond_to) = tokio::select! {
result = &mut middleware => return finish(result),
called = chain_calls.recv() => match called {
Some(call) => call,
None => return finish(middleware.await),
},
};
let chain = async move {
let (mut parts, body) = request.into_parts();
parts.extensions.remove::<Relay>();
let cx = cx.with(parts);
let result = next.run(&cx, body).await.map_err(TowerNextError::tunneled);
let _ = respond_to.send(result);
while let Some((_, respond_to)) = chain_calls.recv().await {
let _ = respond_to.send(Err(TowerNextError::consumed()));
}
};
let mut chain = pin!(chain);
tokio::select! {
result = &mut middleware => return finish(result),
() = &mut chain => {}
}
finish(middleware.await)
})
}
}
#[derive(Clone, Debug)]
pub struct TowerNext {
_priv: (),
}
impl TowerNext {
fn new() -> Self {
Self { _priv: () }
}
}
impl tower::Service<Request> for TowerNext {
type Response = Response;
type Error = TowerNextError;
type Future = Pin<Box<dyn Future<Output = Result<Response, TowerNextError>> + Send>>;
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, mut request: Request) -> Self::Future {
let relay = request.extensions_mut().remove::<Relay>();
Box::pin(async move {
let Some(Relay(relay)) = relay else {
return Err(TowerNextError::detached());
};
let (respond_to, response) = oneshot::channel();
if relay.send((request, respond_to)).await.is_err() {
return Err(TowerNextError::cancelled());
}
response
.await
.unwrap_or_else(|_| Err(TowerNextError::cancelled()))
})
}
}
#[derive(Debug)]
pub struct TowerNextError {
repr: Repr,
}
#[derive(Debug)]
enum Repr {
Tunneled(Error),
Consumed,
Detached,
Cancelled,
}
impl TowerNextError {
fn tunneled(error: Error) -> Self {
Self {
repr: Repr::Tunneled(error),
}
}
fn consumed() -> Self {
Self {
repr: Repr::Consumed,
}
}
fn detached() -> Self {
Self {
repr: Repr::Detached,
}
}
fn cancelled() -> Self {
Self {
repr: Repr::Cancelled,
}
}
}
impl Display for TowerNextError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match &self.repr {
Repr::Tunneled(error) => Display::fmt(error, f),
Repr::Consumed => f.write_str("the chain wrapped by this TowerLayer has already run"),
Repr::Detached => {
f.write_str("the request no longer carries the relay to its TowerLayer")
}
Repr::Cancelled => {
f.write_str("the TowerLayer was dropped before the chain produced a response")
}
}
}
}
impl std::error::Error for TowerNextError {}
#[derive(Debug)]
pub struct TowerServiceError(BoxError);
impl TowerServiceError {
#[must_use]
pub fn get_ref(&self) -> &BoxError {
&self.0
}
#[must_use]
pub fn into_inner(self) -> BoxError {
self.0
}
}
impl Display for TowerServiceError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("tower service error")
}
}
impl std::error::Error for TowerServiceError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
Some(self.0.as_ref())
}
}
type ChainCall = (Request, oneshot::Sender<Result<Response, TowerNextError>>);
#[derive(Clone)]
struct Relay(mpsc::Sender<ChainCall>);
fn relay_channel() -> (Relay, mpsc::Receiver<ChainCall>) {
let (sender, receiver) = mpsc::channel(1);
(Relay(sender), receiver)
}
fn finish<ResBody, E>(result: Result<http::Response<ResBody>, E>) -> Result<Response>
where
ResBody: http_body::Body<Data = Bytes> + Send + 'static,
ResBody::Error: Into<BoxError>,
E: Into<BoxError>,
{
match result {
Ok(response) => Ok(response.map(Body::new)),
Err(error) => Err(recover(error.into())),
}
}
fn recover(error: BoxError) -> Error {
match error.downcast::<TowerNextError>() {
Ok(error) => match error.repr {
Repr::Tunneled(error) => error,
repr => TowerNextError { repr }.into(),
},
Err(error) => TowerServiceError(error).into(),
}
}
#[derive(Clone)]
pub struct TowerService {
router: Arc<Router>,
}
impl TowerService {
#[must_use]
pub fn new(router: Router) -> Self {
Self {
router: Arc::new(router),
}
}
}
impl<B> tower::Service<Request<B>> for TowerService
where
B: http_body::Body<Data = Bytes> + Send + 'static,
B::Error: Into<BoxError>,
{
type Response = Response;
type Error = Infallible;
type Future = Pin<Box<dyn Future<Output = Result<Response, Infallible>> + Send>>;
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, request: Request<B>) -> Self::Future {
let router = Arc::clone(&self.router);
Box::pin(async move { Ok(router.handle(request.map(Body::new)).await) })
}
}
#[cfg(test)]
mod tests {
use std::{
borrow::Cow,
convert::Infallible,
sync::{
Arc,
atomic::{AtomicUsize, Ordering},
},
time::Duration,
};
use http::{HeaderValue, StatusCode, request::Parts};
use topcoat_core::context::{Cx, request_context, try_request_context};
use super::*;
use crate::{
Method, RouteFn, RouteFuture, Router, Terminal,
error::{NotFoundError, not_found},
request::Bytes,
response::IntoResponse,
to_bytes,
};
fn block_on<F: Future>(future: F) -> F::Output {
tokio::runtime::Builder::new_current_thread()
.enable_time()
.build()
.unwrap()
.block_on(future)
}
fn path(s: &'static str) -> Cow<'static, Path> {
Cow::Borrowed(Path::new(s))
}
fn cx_for(uri: &str) -> Cx {
let (parts, ()) = http::Request::builder()
.uri(uri)
.body(())
.unwrap()
.into_parts();
Cx::default().with(parts)
}
fn run(layer: &dyn Layer, cx: &Cx, route: &RouteFn) -> Result<Response> {
let next = Next::new(&[], Terminal::Route(route));
block_on(layer.handle(cx, Body::empty(), next))
}
fn body_bytes(response: Response) -> Bytes {
let (_, body) = response.into_parts();
block_on(to_bytes(body, usize::MAX)).unwrap()
}
fn say_route(cx: &Cx, _body: Body) -> RouteFuture<'_> {
Box::pin(async move { "route".into_response(cx) })
}
fn echo_header(cx: &Cx, _body: Body) -> RouteFuture<'_> {
Box::pin(async move {
let value = crate::request::headers(cx)
.get("x-tower")
.and_then(|value| value.to_str().ok())
.unwrap_or("missing")
.to_owned();
value.into_response(cx)
})
}
fn hang(_cx: &Cx, _body: Body) -> RouteFuture<'_> {
Box::pin(std::future::pending())
}
fn not_found_route(_cx: &Cx, _body: Body) -> RouteFuture<'_> {
Box::pin(async move { Err(not_found().into()) })
}
fn long_route(cx: &Cx, _body: Body) -> RouteFuture<'_> {
Box::pin(async move { "route ".repeat(64).into_response(cx) })
}
async fn echo_service(request: Request) -> Result<Response, Infallible> {
let (parts, body) = request.into_parts();
let bytes = to_bytes(body, usize::MAX).await.unwrap();
let reply = format!(
"{} {} {}",
parts.method,
parts.uri,
String::from_utf8_lossy(&bytes)
);
Ok(Response::new(Body::from(reply)))
}
fn send(router: &Router, uri: &str) -> Response {
let request = http::Request::builder()
.uri(uri)
.body(Body::empty())
.unwrap();
block_on(router.handle(request))
}
struct MarkRequestLayer;
impl<S> tower::Layer<S> for MarkRequestLayer {
type Service = MarkRequest<S>;
fn layer(&self, inner: S) -> Self::Service {
MarkRequest { inner }
}
}
#[derive(Clone)]
struct MarkRequest<S> {
inner: S,
}
impl<S> tower::Service<Request> for MarkRequest<S>
where
S: tower::Service<Request>,
{
type Response = S::Response;
type Error = S::Error;
type Future = S::Future;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, mut request: Request) -> Self::Future {
request
.headers_mut()
.insert("x-tower", HeaderValue::from_static("marked"));
self.inner.call(request)
}
}
struct MarkResponseLayer;
impl<S> tower::Layer<S> for MarkResponseLayer {
type Service = MarkResponse<S>;
fn layer(&self, inner: S) -> Self::Service {
MarkResponse { inner }
}
}
#[derive(Clone)]
struct MarkResponse<S> {
inner: S,
}
impl<S> tower::Service<Request> for MarkResponse<S>
where
S: tower::Service<Request, Response = Response> + Clone + Send + 'static,
S::Future: Send,
{
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, request: Request) -> Self::Future {
let mut inner = self.inner.clone();
Box::pin(async move {
let mut response = inner.call(request).await?;
response
.headers_mut()
.insert("x-tower", HeaderValue::from_static("marked"));
Ok(response)
})
}
}
struct ShortCircuitLayer;
impl<S> tower::Layer<S> for ShortCircuitLayer {
type Service = ShortCircuit;
fn layer(&self, _inner: S) -> Self::Service {
ShortCircuit
}
}
#[derive(Clone)]
struct ShortCircuit;
impl tower::Service<Request> for ShortCircuit {
type Response = Response;
type Error = Infallible;
type Future = std::future::Ready<Result<Response, Infallible>>;
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, request: Request) -> Self::Future {
drop(request);
std::future::ready(Ok(Response::new(Body::from("short"))))
}
}
struct CountingLayer {
builds: Arc<AtomicUsize>,
requests: Arc<AtomicUsize>,
}
impl<S> tower::Layer<S> for CountingLayer {
type Service = Counting<S>;
fn layer(&self, inner: S) -> Self::Service {
self.builds.fetch_add(1, Ordering::SeqCst);
Counting {
requests: self.requests.clone(),
inner,
}
}
}
#[derive(Clone)]
struct Counting<S> {
requests: Arc<AtomicUsize>,
inner: S,
}
impl<S> tower::Service<Request> for Counting<S>
where
S: tower::Service<Request>,
{
type Response = S::Response;
type Error = S::Error;
type Future = S::Future;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, request: Request) -> Self::Future {
self.requests.fetch_add(1, Ordering::SeqCst);
self.inner.call(request)
}
}
struct CallTwiceLayer;
impl<S> tower::Layer<S> for CallTwiceLayer {
type Service = CallTwice<S>;
fn layer(&self, inner: S) -> Self::Service {
CallTwice { inner }
}
}
#[derive(Clone)]
struct CallTwice<S> {
inner: S,
}
impl<S> tower::Service<Request> for CallTwice<S>
where
S: tower::Service<Request, Response = Response> + Clone + Send + 'static,
S::Future: Send,
{
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, request: Request) -> Self::Future {
let mut inner = self.inner.clone();
Box::pin(async move {
let extensions = request.extensions().clone();
inner.call(request).await?;
let mut retry = Request::new(Body::empty());
*retry.extensions_mut() = extensions;
inner.call(retry).await
})
}
}
struct DetachLayer;
impl<S> tower::Layer<S> for DetachLayer {
type Service = Detach<S>;
fn layer(&self, inner: S) -> Self::Service {
Detach { inner }
}
}
#[derive(Clone)]
struct Detach<S> {
inner: S,
}
impl<S> tower::Service<Request> for Detach<S>
where
S: tower::Service<Request>,
{
type Response = S::Response;
type Error = S::Error;
type Future = S::Future;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, request: Request) -> Self::Future {
drop(request);
self.inner.call(Request::new(Body::empty()))
}
}
#[test]
fn tower_layer_exposes_its_path() {
let layer = TowerLayer::new(tower::layer::util::Identity::new()).at("/admin");
assert_eq!(layer.path(), Some(Path::new("/admin")));
}
#[test]
fn passes_the_request_through_to_the_route() {
let layer = TowerLayer::new(tower::layer::util::Identity::new());
let route = RouteFn::new(Method::GET, path("/x"), say_route);
let cx = cx_for("/x");
let response = run(&layer, &cx, &route).unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(&body_bytes(response)[..], b"route");
}
#[test]
fn request_edits_reach_the_route_but_not_the_caller() {
let layer = TowerLayer::new(MarkRequestLayer);
let route = RouteFn::new(Method::GET, path("/x"), echo_header);
let cx = cx_for("/x");
let response = run(&layer, &cx, &route).unwrap();
assert_eq!(&body_bytes(response)[..], b"marked");
assert!(
!request_context::<Parts>(&cx)
.headers
.contains_key("x-tower")
);
}
#[test]
fn response_edits_reach_the_caller() {
let layer = TowerLayer::new(MarkResponseLayer);
let route = RouteFn::new(Method::GET, path("/x"), say_route);
let cx = cx_for("/x");
let response = run(&layer, &cx, &route).unwrap();
assert_eq!(response.headers().get("x-tower").unwrap(), "marked");
assert_eq!(&body_bytes(response)[..], b"route");
}
#[test]
fn middleware_can_short_circuit_without_calling_the_chain() {
let layer = TowerLayer::new(ShortCircuitLayer);
let route = RouteFn::new(Method::GET, path("/x"), say_route);
let cx = cx_for("/x");
let response = run(&layer, &cx, &route).unwrap();
assert_eq!(&body_bytes(response)[..], b"short");
assert!(try_request_context::<Parts>(&cx).is_some());
}
#[test]
fn chain_errors_tunnel_through_unchanged() {
let layer = TowerLayer::new(tower::layer::util::Identity::new());
let route = RouteFn::new(Method::GET, path("/missing"), not_found_route);
let cx = cx_for("/missing");
let next = Next::new(&[], Terminal::Route(&route));
let result = block_on(layer.handle(&cx, Body::empty(), next));
assert!(
result
.unwrap_err()
.downcast_ref::<NotFoundError>()
.is_some()
);
}
#[test]
fn chain_errors_tunnel_through_an_error_boxing_middleware() {
let layer = TowerLayer::new(tower::timeout::TimeoutLayer::new(Duration::from_mins(1)));
let route = RouteFn::new(Method::GET, path("/missing"), not_found_route);
let cx = cx_for("/missing");
let next = Next::new(&[], Terminal::Route(&route));
let result = block_on(layer.handle(&cx, Body::empty(), next));
assert!(
result
.unwrap_err()
.downcast_ref::<NotFoundError>()
.is_some()
);
}
#[test]
fn middleware_is_built_once_and_shared_across_requests() {
let builds = Arc::new(AtomicUsize::new(0));
let requests = Arc::new(AtomicUsize::new(0));
let layer = TowerLayer::new(CountingLayer {
builds: builds.clone(),
requests: requests.clone(),
});
assert_eq!(builds.load(Ordering::SeqCst), 1);
let route = RouteFn::new(Method::GET, path("/x"), say_route);
for _ in 0..2 {
let cx = cx_for("/x");
run(&layer, &cx, &route).unwrap();
}
assert_eq!(builds.load(Ordering::SeqCst), 1);
assert_eq!(requests.load(Ordering::SeqCst), 2);
}
#[test]
fn timeout_middleware_cancels_a_hung_route() {
let layer = TowerLayer::new(tower::timeout::TimeoutLayer::new(Duration::from_millis(10)));
let route = RouteFn::new(Method::GET, path("/x"), hang);
let cx = cx_for("/x");
let error = run(&layer, &cx, &route).unwrap_err();
let middleware = error.downcast_ref::<TowerServiceError>().unwrap();
assert!(middleware.get_ref().is::<tower::timeout::error::Elapsed>());
}
#[test]
fn calling_the_chain_twice_errors() {
let layer = TowerLayer::new(CallTwiceLayer);
let route = RouteFn::new(Method::GET, path("/x"), say_route);
let cx = cx_for("/x");
let error = run(&layer, &cx, &route).unwrap_err();
assert!(error.downcast_ref::<TowerNextError>().is_some());
}
#[test]
fn calling_the_chain_without_the_relay_errors() {
let layer = TowerLayer::new(DetachLayer);
let route = RouteFn::new(Method::GET, path("/x"), say_route);
let cx = cx_for("/x");
let error = run(&layer, &cx, &route).unwrap_err();
assert!(error.downcast_ref::<TowerNextError>().is_some());
}
#[test]
fn works_with_tower_concurrency_limit() {
let router = Router::builder()
.route(RouteFn::new(Method::GET, path("/x"), say_route))
.layer(TowerLayer::new(tower::limit::ConcurrencyLimitLayer::new(1)))
.build();
for _ in 0..2 {
let response = send(&router, "/x");
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(&body_bytes(response)[..], b"route");
}
}
#[test]
fn works_with_tower_buffer_and_rate_limit() {
block_on(async {
let router = Router::builder()
.route(RouteFn::new(Method::GET, path("/x"), say_route))
.layer(TowerLayer::new(
tower::ServiceBuilder::new()
.buffer::<Request>(8)
.rate_limit(100, Duration::from_secs(1))
.into_inner(),
))
.build();
for _ in 0..2 {
let request = http::Request::builder()
.uri("/x")
.body(Body::empty())
.unwrap();
let response = router.handle(request).await;
assert_eq!(response.status(), StatusCode::OK);
}
});
}
#[test]
fn works_with_tower_http_set_response_header() {
let router = Router::builder()
.route(RouteFn::new(Method::GET, path("/admin/x"), say_route))
.route(RouteFn::new(Method::GET, path("/public"), say_route))
.layer(
TowerLayer::new(
tower_http::set_header::SetResponseHeaderLayer::if_not_present(
http::header::HeaderName::from_static("x-tower"),
HeaderValue::from_static("marked"),
),
)
.at("/admin"),
)
.build();
let response = send(&router, "/admin/x");
assert_eq!(response.headers().get("x-tower").unwrap(), "marked");
let response = send(&router, "/public");
assert!(!response.headers().contains_key("x-tower"));
}
#[test]
fn works_with_tower_http_cors() {
let router = Router::builder()
.route(RouteFn::new(Method::GET, path("/x"), say_route))
.layer(TowerLayer::new(tower_http::cors::CorsLayer::permissive()))
.build();
let request = http::Request::builder()
.method(Method::OPTIONS)
.uri("/x")
.header(http::header::ORIGIN, "https://example.com")
.header(http::header::ACCESS_CONTROL_REQUEST_METHOD, "GET")
.body(Body::empty())
.unwrap();
let response = block_on(router.handle(request));
assert_eq!(response.status(), StatusCode::OK);
assert!(
response
.headers()
.contains_key(http::header::ACCESS_CONTROL_ALLOW_ORIGIN)
);
let response = send(&router, "/x");
assert!(
response
.headers()
.contains_key(http::header::ACCESS_CONTROL_ALLOW_ORIGIN)
);
assert_eq!(&body_bytes(response)[..], b"route");
}
#[test]
fn works_with_tower_http_compression() {
let router = Router::builder()
.route(RouteFn::new(Method::GET, path("/x"), long_route))
.layer(TowerLayer::new(
tower_http::compression::CompressionLayer::new(),
))
.build();
let request = http::Request::builder()
.uri("/x")
.header(http::header::ACCEPT_ENCODING, "gzip")
.body(Body::empty())
.unwrap();
let response = block_on(router.handle(request));
assert_eq!(
response
.headers()
.get(http::header::CONTENT_ENCODING)
.unwrap(),
"gzip"
);
let compressed = body_bytes(response);
assert!(!compressed.is_empty());
assert!(compressed.len() < "route ".repeat(64).len());
}
#[test]
fn works_with_tower_http_trace() {
let router = Router::builder()
.route(RouteFn::new(Method::GET, path("/x"), say_route))
.layer(TowerLayer::new(
tower_http::trace::TraceLayer::new_for_http(),
))
.build();
let response = send(&router, "/x");
assert_eq!(response.status(), StatusCode::OK);
let response = send(&router, "/missing");
assert_eq!(response.status(), StatusCode::NOT_FOUND);
}
#[test]
fn works_with_tower_http_timeout() {
let router = Router::builder()
.route(RouteFn::new(Method::GET, path("/x"), hang))
.layer(TowerLayer::new(
tower_http::timeout::TimeoutLayer::with_status_code(
StatusCode::REQUEST_TIMEOUT,
Duration::from_millis(10),
),
))
.build();
let response = send(&router, "/x");
assert_eq!(response.status(), StatusCode::REQUEST_TIMEOUT);
}
#[test]
fn tower_route_exposes_its_methods_and_path() {
let route = TowerRoute::new(
Method::POST,
Path::new("/legacy"),
tower::service_fn(echo_service),
);
assert_eq!(route.methods(), Methods::Only(&[Method::POST]));
assert_eq!(route.path(), Path::new("/legacy"));
let route = TowerRoute::new(
Methods::Any,
Path::new("/legacy"),
tower::service_fn(echo_service),
);
assert_eq!(route.methods(), Methods::Any);
}
#[test]
fn an_any_route_responds_to_every_method() {
let route = TowerRoute::any(Path::new("/legacy"), tower::service_fn(echo_service));
assert_eq!(route.methods(), Methods::Any);
assert_eq!(route.path(), Path::new("/legacy"));
}
#[test]
fn mounts_a_service_at_a_catch_all_path() {
let router = Router::builder()
.route(TowerRoute::any(
Path::new("/legacy/{*rest}"),
tower::service_fn(echo_service),
))
.build();
let request = http::Request::builder()
.method(Method::POST)
.uri("/legacy/users/7?page=2")
.body(Body::from("payload"))
.unwrap();
let response = block_on(router.handle(request));
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
&body_bytes(response)[..],
b"POST /legacy/users/7?page=2 payload"
);
}
#[test]
fn a_strip_prefix_layer_rewrites_the_uri_a_tower_route_sees() {
let router = Router::builder()
.route(TowerRoute::new(
Methods::Any,
Path::new("/legacy/{*rest}"),
tower::service_fn(echo_service),
))
.layer(crate::StripPrefixLayer::new("/legacy"))
.build();
let request = http::Request::builder()
.method(Method::POST)
.uri("/legacy/users/7?page=2")
.body(Body::from("payload"))
.unwrap();
let response = block_on(router.handle(request));
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(&body_bytes(response)[..], b"POST /users/7?page=2 payload");
}
#[test]
fn a_tower_route_serves_only_its_declared_methods() {
let router = Router::builder()
.route(TowerRoute::new(
Method::POST,
Path::new("/legacy"),
tower::service_fn(echo_service),
))
.build();
let request = http::Request::builder()
.method(Method::POST)
.uri("/legacy")
.body(Body::empty())
.unwrap();
assert_eq!(block_on(router.handle(request)).status(), StatusCode::OK);
let response = send(&router, "/legacy");
assert_eq!(response.status(), StatusCode::METHOD_NOT_ALLOWED);
}
#[test]
fn layers_wrap_a_mounted_service() {
let router = Router::builder()
.route(TowerRoute::any(
Path::new("/legacy/{*rest}"),
tower::service_fn(echo_service),
))
.layer(
TowerLayer::new(
tower_http::set_header::SetResponseHeaderLayer::if_not_present(
http::header::HeaderName::from_static("x-tower"),
HeaderValue::from_static("marked"),
),
)
.at("/legacy"),
)
.build();
let response = send(&router, "/legacy/x");
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(response.headers().get("x-tower").unwrap(), "marked");
}
#[test]
fn a_mounted_service_error_surfaces_as_a_tower_service_error() {
let failing = tower::service_fn(|_request: Request| async {
Err::<Response, _>(std::io::Error::other("legacy failure"))
});
let route = TowerRoute::any(Path::new("/legacy"), failing);
let cx = cx_for("/legacy");
let error = block_on(route.handle(&cx, Body::empty())).unwrap_err();
let route_error = error.downcast_ref::<TowerServiceError>().unwrap();
assert!(route_error.get_ref().is::<std::io::Error>());
}
fn echo_body(cx: &Cx, body: Body) -> RouteFuture<'_> {
Box::pin(async move {
let bytes = to_bytes(body, usize::MAX).await.unwrap();
String::from_utf8_lossy(&bytes)
.into_owned()
.into_response(cx)
})
}
fn panic_route(_cx: &Cx, _body: Body) -> RouteFuture<'_> {
Box::pin(async move { panic!("handler panicked") })
}
fn assert_server_bounds<S>(service: S) -> S
where
S: tower::Service<Request, Error = Infallible> + Clone + Send + Sync + 'static,
{
service
}
#[test]
fn tower_service_dispatches_to_the_router() {
let router = Router::builder()
.route(RouteFn::new(Method::GET, path("/x"), say_route))
.build();
let service = assert_server_bounds(TowerService::new(router));
let request = http::Request::builder()
.uri("/x")
.body(Body::empty())
.unwrap();
let response = block_on(service.clone().oneshot(request)).unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(&body_bytes(response)[..], b"route");
let request = http::Request::builder()
.uri("/missing")
.body(Body::empty())
.unwrap();
let response = block_on(service.oneshot(request)).unwrap();
assert_eq!(response.status(), StatusCode::NOT_FOUND);
}
#[test]
fn tower_service_accepts_foreign_request_bodies() {
let router = Router::builder()
.route(RouteFn::new(Method::POST, path("/echo"), echo_body))
.build();
let service = TowerService::new(router);
let request = http::Request::builder()
.method(Method::POST)
.uri("/echo")
.body(http_body_util::Full::new(Bytes::from("payload")))
.unwrap();
let response = block_on(service.oneshot(request)).unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(&body_bytes(response)[..], b"payload");
}
#[test]
fn tower_service_renders_a_panic_as_a_response() {
let router = Router::builder()
.route(RouteFn::new(Method::GET, path("/panic"), panic_route))
.build();
let service = TowerService::new(router);
let request = http::Request::builder()
.uri("/panic")
.body(Body::empty())
.unwrap();
let response = block_on(service.oneshot(request)).unwrap();
assert_eq!(response.status(), StatusCode::INTERNAL_SERVER_ERROR);
}
}