use std::io::Write as _;
use rama::{
Layer as _, Service,
error::extra::OpaqueError,
futures::StreamExt as _,
http::{
Body, BodyExtractExt, Method, Request, Response, StatusCode, Version,
body::util::BodyExt,
client::EasyHttpWebClient,
header::{ACCEPT_ENCODING, CONTENT_ENCODING},
headers::{ContentLength, HeaderMapExt, encoding::AcceptEncoding},
layer::decompression::DecompressionLayer,
service::client::{HttpClientExt, multipart},
},
rt::Executor,
service::BoxService,
tls::client::{ServerVerifyMode, TlsClientConfig},
utils::octets::{kib, mib},
utils::str::any_submatch_ignore_ascii_case,
};
use super::utils;
use flate2::{Compression, write::GzEncoder};
#[ignore]
#[tokio::test]
async fn test_http_tests() {
utils::init_tracing();
let _guard = utils::RamaService::serve_http_test(63133, false);
run_http_tests("http://127.0.0.1:63133").await;
}
#[ignore]
#[tokio::test]
async fn test_http_tests_over_tls() {
utils::init_tracing();
let _guard = utils::RamaService::serve_http_test(63134, true);
run_http_tests("https://127.0.0.1:63134").await;
}
async fn run_http_tests(base_uri: &'static str) {
let client = EasyHttpWebClient::connector_builder()
.with_default_transport_connector()
.with_default_dns_connector()
.without_tls_proxy_support()
.without_proxy_support()
.with_tls_support_using_boringssl(
TlsClientConfig::default_http().with_server_verify(ServerVerifyMode::Disable),
)
.with_default_http_connector(Executor::default())
.build_client()
.boxed();
for http_version in [Version::HTTP_10, Version::HTTP_11, Version::HTTP_2] {
run_http_test_endpoint_method(client.clone(), base_uri, http_version).await;
run_http_test_endpoint_request_compression(client.clone(), base_uri, http_version).await;
run_http_test_endpoint_response_compression(client.clone(), base_uri, http_version).await;
run_http_test_endpoint_response_stream(client.clone(), base_uri, http_version).await;
run_http_test_endpoint_response_stream_compression(client.clone(), base_uri, http_version)
.await;
run_http_test_endpoint_sse(client.clone(), base_uri, http_version).await;
run_http_test_endpoint_octet_stream(client.clone(), base_uri, http_version).await;
run_http_test_endpoint_multipart(client.clone(), base_uri, http_version).await;
run_http_test_endpoint_bytes(client.clone(), base_uri, http_version).await;
run_http_test_endpoint_sink(client.clone(), base_uri, http_version).await;
}
}
async fn run_http_test_endpoint_method(
client: BoxService<Request, Response, OpaqueError>,
base_uri: &'static str,
http_version: Version,
) {
for method in [
Method::GET,
Method::POST,
Method::PUT,
Method::TRACE,
Method::DELETE,
Method::PATCH,
Method::QUERY,
Method::from_bytes(b"COFFEE").unwrap(),
] {
let resp = client
.request(method.clone(), format!("{base_uri}/method"))
.version(http_version)
.send()
.await
.unwrap();
assert_eq!(StatusCode::OK, resp.status());
let ContentLength(content_length) = resp.headers().typed_get().unwrap();
let expected_payload = method.to_string();
assert_eq!(expected_payload.len(), content_length as usize);
assert_eq!(expected_payload, resp.try_into_string().await.unwrap());
}
assert!(
client
.connect(format!("{base_uri}/method"))
.version(http_version)
.send()
.await
.unwrap()
.status()
.is_client_error()
);
}
async fn run_http_test_endpoint_request_compression(
client: BoxService<Request, Response, OpaqueError>,
base_uri: &'static str,
http_version: Version,
) {
for method in [Method::POST, Method::PUT, Method::PATCH] {
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
encoder.write_all(b"Hello?").unwrap();
let body = encoder.finish().unwrap();
let req = Request::builder()
.uri(format!("{base_uri}/request-compression"))
.version(http_version)
.method(method)
.header(CONTENT_ENCODING, "gzip")
.body(Body::from(body))
.unwrap();
let resp = client.serve(req).await.unwrap();
assert_eq!(StatusCode::OK, resp.status());
let ContentLength(content_length) = resp.headers().typed_get().unwrap();
assert_eq!(6, content_length);
assert_eq!("Hello?", resp.try_into_string().await.unwrap());
}
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
encoder.write_all(&vec![0u8; 9 * 1024 * 1024]).unwrap();
let bomb = encoder.finish().unwrap();
let req = Request::builder()
.uri(format!("{base_uri}/request-compression"))
.version(http_version)
.method(Method::POST)
.header(CONTENT_ENCODING, "gzip")
.body(Body::from(bomb))
.unwrap();
let resp = client.serve(req).await.unwrap();
assert!(
resp.status().is_client_error(),
"decompression bomb must be rejected, got status {}",
resp.status(),
);
}
async fn run_http_test_endpoint_response_compression(
client: BoxService<Request, Response, OpaqueError>,
base_uri: &'static str,
http_version: Version,
) {
let client = DecompressionLayer::new().into_layer(client);
for maybe_accept_encoding in [
None,
Some(AcceptEncoding::new_deflate()),
Some(AcceptEncoding::new_deflate().with_br(true)),
Some(AcceptEncoding::new_gzip()),
Some(AcceptEncoding::new_zstd()),
Some(AcceptEncoding::new_zstd().with_gzip(true)),
Some(AcceptEncoding::new_br()),
Some(AcceptEncoding::default()),
] {
let req = client
.get(format!("{base_uri}/response-compression"))
.version(http_version);
let req = if let Some(accept_encoding) =
maybe_accept_encoding.and_then(|ae| ae.maybe_to_header_value())
{
req.header(ACCEPT_ENCODING, accept_encoding)
} else {
req
};
let resp = req.send().await.unwrap();
assert_eq!(StatusCode::OK, resp.status());
let payload = resp.try_into_string().await.unwrap();
assert!(payload.starts_with("# Ethical principles of hacking"));
assert!(payload.contains("All information should be free"));
assert!(payload.ends_with(
"the Chaos Computer Club (CCC).
"
));
}
}
async fn run_http_test_endpoint_response_stream(
client: BoxService<Request, Response, OpaqueError>,
base_uri: &'static str,
http_version: Version,
) {
let resp = client
.get(format!("{base_uri}/response-stream"))
.version(http_version)
.send()
.await
.unwrap();
assert_eq!(StatusCode::OK, resp.status());
assert!(!resp.headers().contains_key("content-length"));
let payload = resp.try_into_string().await.unwrap();
assert!(payload.contains("<title>Chunked transfer encoding test</title>"));
assert!(payload.contains("This is a chunked response after 100 ms"));
assert!(payload.contains("all chunks are sent to a client.</h5></body></html>"));
}
async fn run_http_test_endpoint_response_stream_compression(
client: BoxService<Request, Response, OpaqueError>,
base_uri: &'static str,
http_version: Version,
) {
let client = DecompressionLayer::new().into_layer(client);
for maybe_accept_encoding in [
None,
Some(AcceptEncoding::new_deflate()),
Some(AcceptEncoding::new_deflate().with_br(true)),
Some(AcceptEncoding::new_gzip()),
Some(AcceptEncoding::new_zstd()),
Some(AcceptEncoding::new_zstd().with_gzip(true)),
Some(AcceptEncoding::new_br()),
Some(AcceptEncoding::default()),
] {
let req = client
.get(format!("{base_uri}/response-stream-compression"))
.version(http_version);
let req = if let Some(accept_encoding) =
maybe_accept_encoding.and_then(|ae| ae.maybe_to_header_value())
{
req.header(ACCEPT_ENCODING, accept_encoding)
} else {
req
};
let resp = req.send().await.unwrap();
assert_eq!(StatusCode::OK, resp.status());
assert!(!resp.headers().contains_key("content-length"));
let payload = resp.try_into_string().await.unwrap_or_else(|err| {
panic!("decompression faile for {maybe_accept_encoding:?}: {err}")
});
assert!(payload.contains("<title>Chunked transfer encoding test</title>"));
assert!(payload.contains("This is a chunked response after 100 ms"));
assert!(payload.contains("all chunks are sent to a client.</h5></body></html>"));
}
}
async fn run_http_test_endpoint_sse(
client: BoxService<Request, Response, OpaqueError>,
base_uri: &'static str,
http_version: Version,
) {
let resp = client
.get(format!("{base_uri}/sse"))
.version(http_version)
.send()
.await
.unwrap();
assert_eq!(StatusCode::OK, resp.status());
let mut stream = resp.into_body().into_string_data_event_stream();
for expected_event in [
"Wake up slowly, enjoy morning light",
"Make loose plans, feel excited",
"Do one thing, celebrate it",
"Go to bed, feeling okay",
] {
let event = stream.next().await.unwrap().unwrap();
let data = event.into_data().unwrap();
assert_eq!(expected_event, data);
}
assert!(stream.next().await.is_none());
}
async fn run_http_test_endpoint_octet_stream(
client: BoxService<Request, Response, OpaqueError>,
base_uri: &'static str,
http_version: Version,
) {
let payload: Vec<u8> = (0u8..=200).collect();
let resp = client
.post(format!("{base_uri}/octet-stream"))
.version(http_version)
.octet_stream(payload.clone())
.send()
.await
.unwrap();
assert_eq!(StatusCode::OK, resp.status());
assert_eq!(
Some("application/octet-stream"),
resp.headers()
.get(rama::http::header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok()),
);
let echoed = resp.into_body().collect().await.unwrap().to_bytes();
assert_eq!(payload.as_slice(), echoed.as_ref());
}
async fn run_http_test_endpoint_multipart(
client: BoxService<Request, Response, OpaqueError>,
base_uri: &'static str,
http_version: Version,
) {
let form = multipart::Form::new().text("username", "glen").part(
"attachment",
multipart::Part::bytes(b"hello rama".as_slice())
.with_file_name("note.txt")
.try_with_mime_str("text/plain")
.unwrap(),
);
let resp = client
.post(format!("{base_uri}/multipart"))
.version(http_version)
.multipart(form)
.send()
.await
.unwrap();
assert_eq!(StatusCode::OK, resp.status());
let body: serde_json::Value = resp.try_into_json().await.unwrap();
let parts = body.get("parts").and_then(|v| v.as_array()).unwrap();
assert_eq!(parts.len(), 2);
let username = &parts[0];
assert_eq!(username["name"].as_str(), Some("username"));
assert!(username["filename"].is_null());
assert_eq!(username["size"].as_u64(), Some(4));
assert_eq!(username["text"].as_str(), Some("glen"));
let attachment = &parts[1];
assert_eq!(attachment["name"].as_str(), Some("attachment"));
assert_eq!(attachment["filename"].as_str(), Some("note.txt"));
assert_eq!(attachment["content_type"].as_str(), Some("text/plain"));
assert_eq!(attachment["size"].as_u64(), Some(10));
assert_eq!(attachment["text"].as_str(), Some("hello rama"));
let big = vec![b'x'; kib(300)];
let oversized = multipart::Form::new().part("blob", multipart::Part::bytes(big));
match client
.post(format!("{base_uri}/multipart"))
.version(http_version)
.multipart(oversized)
.send()
.await
{
Ok(resp) => assert_eq!(StatusCode::PAYLOAD_TOO_LARGE, resp.status()),
Err(err) => {
let msg = format!("{err:#}");
assert!(
any_submatch_ignore_ascii_case(
&msg,
[
"broken pipe",
"connection reset",
"connection aborted",
"connection closed",
],
),
"expected 413 response or transport-close error, got: {msg}",
);
}
}
}
async fn run_http_test_endpoint_bytes(
client: BoxService<Request, Response, OpaqueError>,
base_uri: &'static str,
http_version: Version,
) {
let resp = client
.get(format!("{base_uri}/bytes"))
.version(http_version)
.send()
.await
.unwrap();
assert_eq!(StatusCode::OK, resp.status());
let body = resp.into_body().collect().await.unwrap().to_bytes();
assert_eq!(1024, body.len());
assert!(body.iter().all(|&b| b == 0));
let resp = client
.get(format!("{base_uri}/bytes?size=4096"))
.version(http_version)
.send()
.await
.unwrap();
assert_eq!(StatusCode::OK, resp.status());
let body = resp.into_body().collect().await.unwrap().to_bytes();
assert_eq!(4096, body.len());
let resp = client
.get(format!("{base_uri}/bytes?size=4096&chunk=512"))
.version(http_version)
.send()
.await
.unwrap();
assert_eq!(StatusCode::OK, resp.status());
let body = resp.into_body().collect().await.unwrap().to_bytes();
assert_eq!(4096, body.len());
let resp = client
.get(format!("{base_uri}/bytes?size=0"))
.version(http_version)
.send()
.await
.unwrap();
assert_eq!(StatusCode::OK, resp.status());
let body = resp.into_body().collect().await.unwrap().to_bytes();
assert_eq!(0, body.len());
for bad_url in [
format!("{base_uri}/bytes?size=abc"),
format!("{base_uri}/bytes?chunk=0"),
format!("{base_uri}/bytes?delay_ms=99999999"),
] {
let resp = client
.get(bad_url)
.version(http_version)
.send()
.await
.unwrap();
assert_eq!(StatusCode::BAD_REQUEST, resp.status());
}
}
async fn run_http_test_endpoint_sink(
client: BoxService<Request, Response, OpaqueError>,
base_uri: &'static str,
http_version: Version,
) {
let resp = client
.post(format!("{base_uri}/sink"))
.version(http_version)
.send()
.await
.unwrap();
assert_eq!(StatusCode::OK, resp.status());
let body: serde_json::Value = resp.try_into_json().await.unwrap();
assert_eq!(body["bytes"].as_u64(), Some(0));
let payload = b"hello rama sink";
let resp = client
.post(format!("{base_uri}/sink"))
.version(http_version)
.octet_stream(payload.as_slice())
.send()
.await
.unwrap();
assert_eq!(StatusCode::OK, resp.status());
let body: serde_json::Value = resp.try_into_json().await.unwrap();
assert_eq!(body["bytes"].as_u64(), Some(payload.len() as u64));
let big = vec![0u8; mib(1)];
let resp = client
.post(format!("{base_uri}/sink"))
.version(http_version)
.octet_stream(big.clone())
.send()
.await
.unwrap();
assert_eq!(StatusCode::OK, resp.status());
let body: serde_json::Value = resp.try_into_json().await.unwrap();
assert_eq!(body["bytes"].as_u64(), Some(big.len() as u64));
}