Files
axum/src/tests.rs
T

386 lines
10 KiB
Rust
Raw Normal View History

use crate::{app, extract, handler::Handler};
2021-05-31 12:22:16 +02:00
use http::{Request, Response, StatusCode};
use hyper::{Body, Server};
use serde::Deserialize;
use serde_json::json;
use std::net::{SocketAddr, TcpListener};
use tower::{make::Shared, BoxError, Service};
#[tokio::test]
async fn hello_world() {
let app = app()
.at("/")
2021-05-31 22:54:21 +02:00
.get(|_: Request<Body>| async { "Hello, World!" })
2021-05-31 12:22:16 +02:00
.into_service();
let addr = run_in_background(app).await;
let res = reqwest::get(format!("http://{}", addr)).await.unwrap();
let body = res.text().await.unwrap();
assert_eq!(body, "Hello, World!");
}
#[tokio::test]
async fn consume_body() {
let app = app()
.at("/")
2021-05-31 22:54:21 +02:00
.get(|_: Request<Body>, body: String| async { body })
2021-05-31 12:22:16 +02:00
.into_service();
let addr = run_in_background(app).await;
let client = reqwest::Client::new();
let res = client
.get(format!("http://{}", addr))
.body("foo")
.send()
.await
.unwrap();
let body = res.text().await.unwrap();
assert_eq!(body, "foo");
}
#[tokio::test]
async fn deserialize_body() {
#[derive(Debug, Deserialize)]
struct Input {
foo: String,
}
let app = app()
.at("/")
2021-05-31 22:54:21 +02:00
.post(|_: Request<Body>, input: extract::Json<Input>| async { input.into_inner().foo })
2021-05-31 12:22:16 +02:00
.into_service();
let addr = run_in_background(app).await;
let client = reqwest::Client::new();
let res = client
.post(format!("http://{}", addr))
.json(&json!({ "foo": "bar" }))
.send()
.await
.unwrap();
let body = res.text().await.unwrap();
assert_eq!(body, "bar");
}
#[tokio::test]
async fn consume_body_to_json_requires_json_content_type() {
#[derive(Debug, Deserialize)]
struct Input {
foo: String,
}
let app = app()
.at("/")
.post(|_: Request<Body>, input: extract::Json<Input>| async {
let input = input.into_inner();
2021-05-31 22:54:21 +02:00
input.foo
2021-05-31 12:22:16 +02:00
})
.into_service();
let addr = run_in_background(app).await;
let client = reqwest::Client::new();
let res = client
.post(format!("http://{}", addr))
.body(r#"{ "foo": "bar" }"#)
.send()
.await
.unwrap();
2021-05-31 22:54:21 +02:00
let status = res.status();
dbg!(res.text().await.unwrap());
assert_eq!(status, StatusCode::BAD_REQUEST);
2021-05-31 12:22:16 +02:00
}
#[tokio::test]
async fn body_with_length_limit() {
use std::iter::repeat;
#[derive(Debug, Deserialize)]
struct Input {
foo: String,
}
const LIMIT: u64 = 8;
let app = app()
.at("/")
.post(
|req: Request<Body>, _body: extract::BytesMaxLength<LIMIT>| async move {
dbg!(&req);
},
)
.into_service();
let addr = run_in_background(app).await;
let client = reqwest::Client::new();
let res = client
.post(format!("http://{}", addr))
.body(repeat(0_u8).take((LIMIT - 1) as usize).collect::<Vec<_>>())
.send()
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let res = client
.post(format!("http://{}", addr))
.body(repeat(0_u8).take(LIMIT as usize).collect::<Vec<_>>())
.send()
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let res = client
.post(format!("http://{}", addr))
.body(repeat(0_u8).take((LIMIT + 1) as usize).collect::<Vec<_>>())
.send()
.await
.unwrap();
assert_eq!(res.status(), StatusCode::PAYLOAD_TOO_LARGE);
let res = client
.post(format!("http://{}", addr))
.body(reqwest::Body::wrap_stream(futures_util::stream::iter(
vec![Ok::<_, std::io::Error>(bytes::Bytes::new())],
)))
.send()
.await
.unwrap();
assert_eq!(res.status(), StatusCode::LENGTH_REQUIRED);
}
2021-05-31 14:04:05 +02:00
#[tokio::test]
async fn routing() {
let app = app()
.at("/users")
2021-05-31 22:54:21 +02:00
.get(|_: Request<Body>| async { "users#index" })
.post(|_: Request<Body>| async { "users#create" })
2021-05-31 14:04:05 +02:00
.at("/users/:id")
2021-05-31 22:54:21 +02:00
.get(|_: Request<Body>| async { "users#show" })
2021-05-31 14:04:05 +02:00
.at("/users/:id/action")
2021-05-31 22:54:21 +02:00
.get(|_: Request<Body>| async { "users#action" })
2021-05-31 14:04:05 +02:00
.into_service();
let addr = run_in_background(app).await;
let client = reqwest::Client::new();
let res = client.get(format!("http://{}", addr)).send().await.unwrap();
assert_eq!(res.status(), StatusCode::NOT_FOUND);
let res = client
.get(format!("http://{}/users", addr))
.send()
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
assert_eq!(res.text().await.unwrap(), "users#index");
let res = client
.post(format!("http://{}/users", addr))
.send()
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
assert_eq!(res.text().await.unwrap(), "users#create");
let res = client
.get(format!("http://{}/users/1", addr))
.send()
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
assert_eq!(res.text().await.unwrap(), "users#show");
let res = client
.get(format!("http://{}/users/1/action", addr))
.send()
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
assert_eq!(res.text().await.unwrap(), "users#action");
}
#[tokio::test]
async fn extracting_url_params() {
let app = app()
.at("/users/:id")
.get(
|_: Request<Body>, params: extract::UrlParams<(i32,)>| async move {
2021-05-31 22:54:21 +02:00
let id = params.into_inner();
2021-05-31 14:04:05 +02:00
assert_eq!(id, 42);
},
)
.post(
|_: Request<Body>, params_map: extract::UrlParamsMap| async move {
assert_eq!(params_map.get("id").unwrap(), "1337");
assert_eq!(params_map.get_typed::<i32>("id").unwrap(), 1337);
},
)
.into_service();
2021-05-31 12:22:16 +02:00
2021-05-31 14:04:05 +02:00
let addr = run_in_background(app).await;
let client = reqwest::Client::new();
let res = client
.get(format!("http://{}/users/42", addr))
.send()
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
let res = client
.post(format!("http://{}/users/1337", addr))
.send()
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
}
2021-05-31 12:22:16 +02:00
2021-05-31 16:28:26 +02:00
#[tokio::test]
async fn boxing() {
let app = app()
.at("/")
2021-05-31 22:54:21 +02:00
.get(|_: Request<Body>| async { "hi from GET" })
2021-05-31 16:28:26 +02:00
.boxed()
2021-05-31 22:54:21 +02:00
.post(|_: Request<Body>| async { "hi from POST" })
2021-05-31 16:28:26 +02:00
.into_service();
let addr = run_in_background(app).await;
let client = reqwest::Client::new();
2021-05-31 20:42:57 +02:00
let res = client.get(format!("http://{}", addr)).send().await.unwrap();
2021-05-31 16:28:26 +02:00
assert_eq!(res.status(), StatusCode::OK);
assert_eq!(res.text().await.unwrap(), "hi from GET");
let res = client
.post(format!("http://{}", addr))
.send()
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
assert_eq!(res.text().await.unwrap(), "hi from POST");
}
2021-05-31 12:22:16 +02:00
#[tokio::test]
async fn service_handlers() {
use crate::{body::BoxBody, ServiceExt as _};
use std::convert::Infallible;
use tower::service_fn;
use tower_http::services::ServeFile;
let app = app()
.at("/echo")
.post_service(service_fn(|req: Request<Body>| async move {
Ok::<_, Infallible>(Response::new(req.into_body()))
}))
// calling boxed isn't necessary here but done so
// we're sure it compiles
.boxed()
.at("/static/Cargo.toml")
.get_service(
ServeFile::new("Cargo.toml").handle_error(|error: std::io::Error| {
// `ServeFile` internally maps some errors to `404` so we don't have
// to handle those here
let body = BoxBody::from(error.to_string());
Response::builder()
.status(StatusCode::INTERNAL_SERVER_ERROR)
.body(body)
.unwrap()
}),
)
// calling boxed isn't necessary here but done so
// we're sure it compiles
.boxed()
.into_service();
let addr = run_in_background(app).await;
let client = reqwest::Client::new();
2021-06-01 00:34:09 +02:00
let res = client
.post(format!("http://{}/echo", addr))
.body("foobar")
.send()
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
assert_eq!(res.text().await.unwrap(), "foobar");
let res = client
.get(format!("http://{}/static/Cargo.toml", addr))
.body("foobar")
.send()
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
assert!(res.text().await.unwrap().contains("edition ="));
}
#[tokio::test]
async fn middleware_on_single_route() {
use tower::ServiceBuilder;
use tower_http::{compression::CompressionLayer, trace::TraceLayer};
async fn handle(_: Request<Body>) -> &'static str {
"Hello, World!"
}
let app = app()
.at("/")
.get(
handle.layer(
ServiceBuilder::new()
.layer(TraceLayer::new_for_http())
.layer(CompressionLayer::new())
.into_inner(),
),
)
.into_service();
let addr = run_in_background(app).await;
let res = reqwest::get(format!("http://{}", addr)).await.unwrap();
let body = res.text().await.unwrap();
assert_eq!(body, "Hello, World!");
}
2021-06-01 08:32:58 +02:00
2021-05-31 12:22:16 +02:00
/// Run a `tower::Service` in the background and get a URI for it.
pub async fn run_in_background<S, ResBody>(svc: S) -> SocketAddr
where
S: Service<Request<Body>, Response = Response<ResBody>> + Clone + Send + 'static,
ResBody: http_body::Body + Send + 'static,
ResBody::Data: Send,
ResBody::Error: Into<BoxError>,
S::Error: Into<BoxError>,
S::Future: Send,
{
let listener = TcpListener::bind("127.0.0.1:0").expect("Could not bind ephemeral socket");
let addr = listener.local_addr().unwrap();
println!("Listening on {}", addr);
let (tx, rx) = tokio::sync::oneshot::channel();
tokio::spawn(async move {
let server = Server::from_tcp(listener).unwrap().serve(Shared::new(svc));
tx.send(()).unwrap();
server.await.expect("server error");
});
rx.await.unwrap();
addr
}