use zenith_api::{CanonicalRequest, Method, Protocol, Transport};
use zenith_http1::{Http1Config, Http1Parser};
use zenith_http2::hpack::{HpackDecoder, HpackEncoder, HeaderField};
use zenith_http3::encoder::QpackEncoder;
use zenith_http3::qpack::QpackDecoder;
use zenith_web::{normalize_http1_request, normalize_http2_request, normalize_http3_request};
fn assert_semantically_equal(h1: &CanonicalRequest, h2: &CanonicalRequest, h3: &CanonicalRequest) {
assert_eq!(h1.method, h2.method, "method mismatch h1 vs h2");
assert_eq!(h1.method, h3.method, "method mismatch h1 vs h3");
assert_eq!(h1.path_str(), h2.path_str(), "path mismatch h1 vs h2");
assert_eq!(h1.path_str(), h3.path_str(), "path mismatch h1 vs h3");
assert_eq!(h1.protocol, Protocol::Http1);
assert_eq!(h2.protocol, Protocol::Http2);
assert_eq!(h3.protocol, Protocol::Http3);
let skip_headers = |name: &str| name.eq_ignore_ascii_case("host");
let count_regular = |c: &CanonicalRequest| {
c.headers_iter()
.iter()
.filter(|h| !skip_headers(h.name_str()))
.count()
};
assert_eq!(
count_regular(h1),
count_regular(h2),
"regular header count mismatch h1 vs h2"
);
assert_eq!(
count_regular(h1),
count_regular(h3),
"regular header count mismatch h1 vs h3"
);
for hdr in h1.headers_iter().iter() {
let name = hdr.name_str();
if skip_headers(name) {
continue;
}
let h2_hdr = h2
.find_header(name)
.unwrap_or_else(|| panic!("header '{name}' missing in h2"));
assert_eq!(
hdr.value_str(),
h2_hdr.value_str(),
"header '{name}' value mismatch h1 vs h2"
);
let h3_hdr = h3
.find_header(name)
.unwrap_or_else(|| panic!("header '{name}' missing in h3"));
assert_eq!(
hdr.value_str(),
h3_hdr.value_str(),
"header '{name}' value mismatch h1 vs h3"
);
}
let h1_host = h1
.find_header("host")
.map(|h| h.value_str().to_string())
.unwrap_or_default();
assert_eq!(
h2.authority_str(),
h1_host,
"authority mismatch: h1 host={h1_host} h2 authority={}",
h2.authority_str()
);
assert_eq!(
h3.authority_str(),
h1_host,
"authority mismatch: h1 host={h1_host} h3 authority={}",
h3.authority_str()
);
}
fn build_canonical_requests(
method: &str,
path: &str,
host: &str,
regular_headers: &[(&str, &str)],
scheme: &str,
) -> (CanonicalRequest, CanonicalRequest, CanonicalRequest) {
let mut raw_h1 = Vec::new();
raw_h1.extend_from_slice(method.as_bytes());
raw_h1.extend_from_slice(b" ");
raw_h1.extend_from_slice(path.as_bytes());
raw_h1.extend_from_slice(b" HTTP/1.1\r\n");
raw_h1.extend_from_slice(b"Host: ");
raw_h1.extend_from_slice(host.as_bytes());
raw_h1.extend_from_slice(b"\r\n");
for (name, value) in regular_headers {
raw_h1.extend_from_slice(name.as_bytes());
raw_h1.extend_from_slice(b": ");
raw_h1.extend_from_slice(value.as_bytes());
raw_h1.extend_from_slice(b"\r\n");
}
raw_h1.extend_from_slice(b"\r\n");
let mut parser = Http1Parser::new(Http1Config::default());
let (http1_req, consumed) = parser
.feed(&raw_h1)
.expect("HTTP/1.1 parse must succeed");
let http1_req = http1_req.expect("HTTP/1.1 request must be complete");
assert_eq!(consumed, raw_h1.len(), "HTTP/1.1 parser must consume all bytes");
let transport_h1 = if scheme == "https" {
Transport::Tls13
} else {
Transport::Plaintext
};
let canonical_h1 = normalize_http1_request(&http1_req, transport_h1)
.expect("HTTP/1.1 normalize must succeed");
let mut h2_headers: Vec<HeaderField> = Vec::with_capacity(4 + regular_headers.len());
h2_headers.push(HeaderField::new(":method", method));
h2_headers.push(HeaderField::new(":path", path));
h2_headers.push(HeaderField::new(":scheme", scheme));
h2_headers.push(HeaderField::new(":authority", host));
for (name, value) in regular_headers {
h2_headers.push(HeaderField::new(*name, *value));
}
let mut encoder_h2 = HpackEncoder::new(0);
let hpack_bytes = encoder_h2
.encode(&h2_headers)
.expect("HPACK encode must succeed");
let mut decoder_h2 = HpackDecoder::new(0);
let decoded_h2 = decoder_h2
.decode(&hpack_bytes)
.expect("HPACK decode must succeed");
assert_eq!(
decoded_h2.len(),
h2_headers.len(),
"HPACK roundtrip must preserve header count"
);
let canonical_h2 = normalize_http2_request(&decoded_h2, transport_h1)
.expect("HTTP/2 normalize must succeed");
let h3_headers: Vec<(Vec<u8>, Vec<u8>)> = {
let mut v: Vec<(Vec<u8>, Vec<u8>)> = Vec::with_capacity(4 + regular_headers.len());
v.push((b":method".to_vec(), method.as_bytes().to_vec()));
v.push((b":path".to_vec(), path.as_bytes().to_vec()));
v.push((b":scheme".to_vec(), scheme.as_bytes().to_vec()));
v.push((b":authority".to_vec(), host.as_bytes().to_vec()));
for (name, value) in regular_headers {
v.push((name.as_bytes().to_vec(), value.as_bytes().to_vec()));
}
v
};
let mut encoder_h3 = QpackEncoder::new(0).disable_auto_insert();
let (qpack_bytes, _ric, _delta_base) = encoder_h3
.encode_field_section(&h3_headers)
.expect("QPACK encode must succeed");
let decoder_h3 = QpackDecoder::new(0);
let decoded_h3 = decoder_h3
.decode_field_section(&qpack_bytes)
.expect("QPACK decode must succeed");
assert_eq!(
decoded_h3.len(),
h3_headers.len(),
"QPACK roundtrip must preserve header count"
);
let canonical_h3 = normalize_http3_request(&decoded_h3, transport_h1)
.expect("HTTP/3 normalize must succeed");
(canonical_h1, canonical_h2, canonical_h3)
}
#[test]
fn test_e2e_simple_get_all_protocols_consistent() {
let (h1, h2, h3) = build_canonical_requests(
"GET",
"/api/v1/users?active=true",
"example.com",
&[
("accept", "application/json"),
("user-agent", "Zenith/1.0"),
],
"https",
);
assert_semantically_equal(&h1, &h2, &h3);
assert_eq!(h1.method, Method::Get);
assert_eq!(h1.path_str(), "/api/v1/users");
assert_eq!(h1.query_str(), "active=true");
}
#[test]
fn test_e2e_post_request_all_protocols_consistent() {
let (h1, h2, h3) = build_canonical_requests(
"POST",
"/api/v1/orders",
"api.example.com",
&[
("content-type", "application/json"),
("accept", "application/json"),
],
"https",
);
assert_semantically_equal(&h1, &h2, &h3);
assert_eq!(h1.method, Method::Post);
}
#[test]
fn test_e2e_plaintext_http_scheme_consistent() {
let (h1, h2, h3) = build_canonical_requests(
"GET",
"/search?q=hello+world",
"example.com:80",
&[("accept", "text/html")],
"http",
);
assert_semantically_equal(&h1, &h2, &h3);
assert_eq!(h1.transport, Transport::Plaintext);
assert_eq!(h2.transport, Transport::Plaintext);
assert_eq!(h3.transport, Transport::Plaintext);
}
#[test]
fn test_e2e_delete_request_all_protocols_consistent() {
let (h1, h2, h3) = build_canonical_requests(
"DELETE",
"/api/v1/users/42",
"example.com",
&[("authorization", "Bearer token123")],
"https",
);
assert_semantically_equal(&h1, &h2, &h3);
assert_eq!(h1.method, Method::Delete);
}
#[test]
fn test_e2e_multiple_headers_preserved() {
let (h1, h2, h3) = build_canonical_requests(
"GET",
"/",
"example.com",
&[
("accept", "text/html,application/xhtml+xml"),
("accept-encoding", "gzip, deflate, br"),
("accept-language", "en-US,en;q=0.9"),
("cache-control", "no-cache"),
("x-request-id", "abc-123-def-456"),
],
"https",
);
assert_semantically_equal(&h1, &h2, &h3);
assert_eq!(h1.header_count(), 6); assert_eq!(h2.header_count(), 5); assert_eq!(h3.header_count(), 5); }
#[test]
fn test_h2_forbidden_connection_header_rejected() {
let h2_headers = vec![
HeaderField::new(":method", "GET"),
HeaderField::new(":path", "/"),
HeaderField::new(":scheme", "https"),
HeaderField::new(":authority", "example.com"),
HeaderField::new("connection", "keep-alive"),
];
let result = normalize_http2_request(&h2_headers, Transport::Tls13);
assert!(
result.is_err(),
"HTTP/2 must reject forbidden 'connection' header"
);
}
#[test]
fn test_h3_transfer_encoding_rejected() {
let h3_headers: Vec<(Vec<u8>, Vec<u8>)> = vec![
(b":method".to_vec(), b"GET".to_vec()),
(b":path".to_vec(), b"/".to_vec()),
(b":scheme".to_vec(), b"https".to_vec()),
(b":authority".to_vec(), b"example.com".to_vec()),
(b"transfer-encoding".to_vec(), b"chunked".to_vec()),
];
let result = normalize_http3_request(&h3_headers, Transport::Tls13);
assert!(
result.is_err(),
"HTTP/3 must reject forbidden 'transfer-encoding' header"
);
}
#[test]
fn test_h2_pseudo_after_regular_rejected() {
let h2_headers = vec![
HeaderField::new("content-type", "text/plain"),
HeaderField::new(":method", "GET"),
HeaderField::new(":path", "/"),
HeaderField::new(":scheme", "http"),
];
let result = normalize_http2_request(&h2_headers, Transport::Plaintext);
assert!(
result.is_err(),
"HTTP/2 must reject pseudo-header after regular header"
);
}
#[test]
fn test_scheme_transport_mismatch_rejected() {
let h2_headers = vec![
HeaderField::new(":method", "GET"),
HeaderField::new(":path", "/"),
HeaderField::new(":scheme", "https"),
HeaderField::new(":authority", "example.com"),
];
let result = normalize_http2_request(&h2_headers, Transport::Plaintext);
assert!(
result.is_err(),
"must reject scheme/transport mismatch (https over plaintext)"
);
}
#[test]
fn test_hpack_roundtrip_preserves_header_values() {
let original = vec![
HeaderField::new(":method", "POST"),
HeaderField::new(":path", "/api/v1/data?key=value&sort=desc"),
HeaderField::new(":scheme", "https"),
HeaderField::new(":authority", "api.example.com:8443"),
HeaderField::new("content-type", "application/json; charset=utf-8"),
HeaderField::new("authorization", "Bearer eyJhbGciOiJIUzI1NiJ9.payload.sig"),
HeaderField::new("x-custom-header", "值 with 中文 unicode"),
];
let mut encoder = HpackEncoder::new(0);
let encoded = encoder.encode(&original).expect("encode");
let mut decoder = HpackDecoder::new(0);
let decoded = decoder.decode(&encoded).expect("decode");
assert_eq!(decoded.len(), original.len());
for (i, (orig, dec)) in original.iter().zip(decoded.iter()).enumerate() {
assert_eq!(
orig.name, dec.name,
"header name mismatch at index {i}"
);
assert_eq!(
orig.value, dec.value,
"header value mismatch at index {i} (name={})",
orig.name
);
}
}
#[test]
fn test_qpack_roundtrip_preserves_header_values() {
let original: Vec<(Vec<u8>, Vec<u8>)> = vec![
(b":method".to_vec(), b"PUT".to_vec()),
(b":path".to_vec(), "/api/v2/resource/更新".as_bytes().to_vec()),
(b":scheme".to_vec(), b"https".to_vec()),
(b":authority".to_vec(), b"example.com".to_vec()),
(b"content-type".to_vec(), b"application/json".to_vec()),
(b"x-trace-id".to_vec(), b"550e8400-e29b-41d4-a716-446655440000".to_vec()),
];
let mut encoder = QpackEncoder::new(0).disable_auto_insert();
let (encoded, _ric, _delta) = encoder
.encode_field_section(&original)
.expect("encode");
let decoder = QpackDecoder::new(0);
let decoded = decoder
.decode_field_section(&encoded)
.expect("decode");
assert_eq!(decoded.len(), original.len());
for (i, (orig, dec)) in original.iter().zip(decoded.iter()).enumerate() {
assert_eq!(
orig.0, dec.0,
"QPACK header name mismatch at index {i}"
);
assert_eq!(
orig.1, dec.1,
"QPACK header value mismatch at index {i}"
);
}
}