use super::*;
use tokio::io::duplex;
async fn run(leftover: &[u8], input: &[u8], framing: Framing) -> std::io::Result<(u64, Vec<u8>)> {
let (mut src_tx, mut src_rx) = duplex(64 * 1024);
src_tx.write_all(input).await.unwrap();
drop(src_tx);
let mut out = Vec::new();
let mut buffered = Buffered::new(&mut src_rx, leftover.to_vec());
let n = forward(&mut buffered, &mut out, framing).await?;
Ok((n, out))
}
#[tokio::test]
async fn a_length_framed_body_forwards_exactly_its_declared_size() {
let (n, out) = run(b"", b"abcdefghij", Framing::Length(10)).await.unwrap();
assert_eq!(n, 10);
assert_eq!(out, b"abcdefghij");
}
#[tokio::test]
async fn bytes_already_read_with_the_head_are_forwarded_first() {
let (n, out) = run(b"abc", b"defghij", Framing::Length(10)).await.unwrap();
assert_eq!(n, 10);
assert_eq!(out, b"abcdefghij");
}
#[tokio::test]
async fn a_body_entirely_in_the_leftover_needs_no_read() {
let (n, out) = run(b"abcdefghij", b"", Framing::Length(10)).await.unwrap();
assert_eq!(n, 10);
assert_eq!(out, b"abcdefghij");
}
#[tokio::test]
async fn nothing_past_the_declared_length_is_forwarded() {
let (n, out) = run(b"", b"abcdeSURPLUS", Framing::Length(5)).await.unwrap();
assert_eq!(n, 5);
assert_eq!(out, b"abcde");
}
#[tokio::test]
async fn a_body_shorter_than_its_declared_length_is_an_error() {
let err = run(b"", b"abc", Framing::Length(10)).await.unwrap_err();
assert_eq!(err.kind(), std::io::ErrorKind::UnexpectedEof);
}
#[tokio::test]
async fn an_empty_framing_forwards_nothing() {
let (n, out) = run(b"ignored", b"also ignored", Framing::Empty)
.await
.unwrap();
assert_eq!(n, 0);
assert!(out.is_empty());
}
#[tokio::test]
async fn a_chunked_body_is_forwarded_with_its_framing_intact() {
let input = b"5\r\nhello\r\n6\r\n world\r\n0\r\n\r\n";
let (n, out) = run(b"", input, Framing::Chunked).await.unwrap();
assert_eq!(n, 11, "the payload bytes, not the framing overhead");
assert_eq!(out, input, "byte for byte, framing included");
}
#[tokio::test]
async fn a_chunked_body_with_no_chunks_is_forwarded() {
let (n, out) = run(b"", b"0\r\n\r\n", Framing::Chunked).await.unwrap();
assert_eq!(n, 0);
assert_eq!(out, b"0\r\n\r\n");
}
#[tokio::test]
async fn trailers_after_the_terminal_chunk_are_forwarded() {
let input = b"3\r\nabc\r\n0\r\nX-Checksum: 1\r\n\r\n";
let (_, out) = run(b"", input, Framing::Chunked).await.unwrap();
assert_eq!(out, input);
}
#[tokio::test]
async fn a_trailer_cannot_restate_a_header_the_head_strip_removes() {
let input = b"3\r\nabc\r\n0\r\n\
X-Forwarded-For: 203.0.113.9\r\n\
Connection: keep-alive\r\n\
Transfer-Encoding: chunked\r\n\
Content-Length: 99\r\n\
Host: elsewhere.example\r\n\
X-Checksum: 1\r\n\r\n";
let (_, out) = run(b"", input, Framing::Chunked).await.unwrap();
let text = String::from_utf8(out).expect("ascii");
assert!(
!text.contains("203.0.113.9"),
"a forwarding chain reached the far side through the trailers: {text}"
);
for stripped in [
"Connection:",
"Transfer-Encoding:",
"Content-Length:",
"Host:",
] {
assert!(!text.contains(stripped), "{stripped} survived: {text}");
}
assert!(text.contains("X-Checksum: 1"), "{text}");
assert!(text.ends_with("\r\n\r\n"), "the body still ends: {text:?}");
}
#[tokio::test]
async fn a_trailer_that_is_not_a_field_is_dropped() {
let input = b"0\r\nnot a header line\r\nX-Kept: 1\r\n\r\n";
let (_, out) = run(b"", input, Framing::Chunked).await.unwrap();
let text = String::from_utf8(out).expect("ascii");
assert!(!text.contains("not a header line"), "{text}");
assert!(text.contains("X-Kept: 1"), "{text}");
}
#[tokio::test]
async fn chunk_extensions_pass_through_uninterpreted() {
let input = b"3;name=value\r\nabc\r\n0\r\n\r\n";
let (n, out) = run(b"", input, Framing::Chunked).await.unwrap();
assert_eq!(n, 3);
assert_eq!(out, input);
}
#[tokio::test]
async fn a_chunk_size_that_is_not_hexadecimal_is_an_error() {
for input in [
&b"zz\r\nabc\r\n0\r\n\r\n"[..],
&b"-1\r\nabc\r\n0\r\n\r\n"[..],
&b" \r\n"[..],
] {
assert!(
run(b"", input, Framing::Chunked).await.is_err(),
"{input:?} is not a chunk size"
);
}
}
#[tokio::test]
async fn a_chunk_not_terminated_by_crlf_is_an_error() {
assert!(
run(b"", b"3\r\nabcXX\r\n0\r\n\r\n", Framing::Chunked)
.await
.is_err()
);
}
#[tokio::test]
async fn a_chunked_body_that_ends_before_its_terminal_chunk_is_an_error() {
let err = run(b"", b"5\r\nhello\r\n", Framing::Chunked)
.await
.unwrap_err();
assert_eq!(err.kind(), std::io::ErrorKind::UnexpectedEof);
}
#[tokio::test]
async fn an_over_long_chunk_size_line_is_refused() {
let input = vec![b'0'; MAX_CHUNK_LINE + 64];
assert!(run(b"", &input, Framing::Chunked).await.is_err());
}
#[tokio::test]
async fn a_close_framed_body_forwards_everything_until_the_source_ends() {
let (n, out) = run(b"lead", b"ing and trailing", Framing::UntilClose)
.await
.unwrap();
assert_eq!(n, 20);
assert_eq!(out, b"leading and trailing");
}
#[tokio::test]
async fn a_close_framed_body_may_be_empty() {
let (n, out) = run(b"", b"", Framing::UntilClose).await.unwrap();
assert_eq!(n, 0);
assert!(out.is_empty());
}
#[tokio::test]
async fn a_frame_reaches_the_destination_before_the_source_has_finished() {
let (mut src_tx, mut src_rx) = duplex(64 * 1024);
let (mut dst_tx, mut dst_rx) = duplex(64 * 1024);
let pump = tokio::spawn(async move {
let mut buffered = Buffered::new(&mut src_rx, Vec::new());
forward(&mut buffered, &mut dst_tx, Framing::UntilClose).await
});
src_tx.write_all(b"data: first\n\n").await.unwrap();
src_tx.flush().await.unwrap();
let mut seen = [0u8; 13];
tokio::time::timeout(
std::time::Duration::from_secs(5),
dst_rx.read_exact(&mut seen),
)
.await
.expect("the first frame must arrive before the source closes")
.expect("read");
assert_eq!(&seen, b"data: first\n\n");
src_tx.write_all(b"data: [DONE]\n\n").await.unwrap();
drop(src_tx);
let forwarded = pump.await.unwrap().unwrap();
assert_eq!(forwarded, 27);
}
#[tokio::test]
async fn a_trailer_cannot_smuggle_a_second_field_behind_a_bare_lf() {
let input = b"0\r\nX-Ok: 1\nContent-Length: 99\r\n\r\n";
let (_, out) = run(b"", input, Framing::Chunked).await.unwrap();
let text = String::from_utf8(out).expect("ascii");
assert!(
!text.contains("Content-Length"),
"a second field rode through behind a bare LF: {text:?}"
);
}
#[tokio::test]
async fn an_ordinary_trailer_still_survives_the_terminator_check() {
let input = b"3\r\nabc\r\n0\r\nX-Checksum: 1\r\nX-Other: 2\r\n\r\n";
let (_, out) = run(b"", input, Framing::Chunked).await.unwrap();
let text = String::from_utf8(out).expect("ascii");
assert!(text.contains("X-Checksum: 1"), "{text:?}");
assert!(text.contains("X-Other: 2"), "{text:?}");
}
#[tokio::test]
async fn a_chunk_size_that_is_not_plain_hexadecimal_is_refused() {
for bad in ["+a\r\nabcdefghij\r\n0\r\n\r\n", "-1\r\n\r\n0\r\n\r\n"] {
let err = run(b"", bad.as_bytes(), Framing::Chunked)
.await
.expect_err("must be refused");
assert!(
err.to_string().contains("hexadecimal"),
"{bad:?} gave {err}"
);
}
}
#[tokio::test]
async fn an_ordinary_chunk_size_is_still_hexadecimal() {
for (input, body) in [
("3\r\nabc\r\n0\r\n\r\n", "abc"),
("A\r\n0123456789\r\n0\r\n\r\n", "0123456789"),
("a\r\n0123456789\r\n0\r\n\r\n", "0123456789"),
] {
let (n, out) = run(b"", input.as_bytes(), Framing::Chunked)
.await
.unwrap_or_else(|e| panic!("{input:?} must frame: {e}"));
assert_eq!(n, body.len() as u64, "{input:?}");
assert!(String::from_utf8(out).unwrap().contains(body), "{input:?}");
}
}