use core::net::IpAddr;
use super::owned::OwnedUriRef;
use super::resolve::remove_dot_segments_graceful;
use super::{Uri, UriInner};
use crate::address::{Domain, Host, UninterpretedHost, UninterpretedHostRef};
use crate::byte_sets::is_unreserved_byte;
use rama_core::bytes::{Bytes, BytesMut};
pub(super) fn canonicalize_uri(uri: Uri) -> Uri {
if matches!(uri.inner, UriInner::Asterisk) {
return uri;
}
let owned = uri.as_owned_components();
let canonical = canonicalize_owned(owned);
Uri {
inner: UriInner::Owned(crate::std::sync::Arc::new(canonical)),
}
}
fn canonicalize_owned(mut owned: OwnedUriRef) -> OwnedUriRef {
if let Some(scheme) = &owned.scheme
&& scheme.as_str().bytes().any(|b| b.is_ascii_uppercase())
{
let lower = scheme.as_str().to_ascii_lowercase();
if let Ok(lowered) = crate::Protocol::try_from(lower) {
owned.scheme = Some(lowered);
}
}
if let Some(authority) = &mut owned.authority {
if let Host::Uninterpreted(h) = &authority.address.host {
let h_ref = UninterpretedHostRef::from(h);
if let Ok(ip) = IpAddr::try_from(h_ref) {
authority.address.host = Host::Address(ip);
} else if let Ok(d) = Domain::try_from(h_ref) {
authority.address.host = Host::Name(d);
}
}
match &mut authority.address.host {
Host::Name(d) => normalize_host_domain(d),
Host::Uninterpreted(h) => normalize_host_uninterpreted(h),
Host::Address(_) => {}
}
if let Some(scheme) = &owned.scheme
&& let Some(default) = scheme.default_port()
&& authority.address.port == crate::address::OptPort::Set(default)
{
authority.address.port = crate::address::OptPort::Unset;
}
if authority.address.port == crate::address::OptPort::Empty {
authority.address.port = crate::address::OptPort::Unset;
}
}
normalize_pct(&mut owned.path);
if let Some(q) = &mut owned.query {
normalize_pct(&mut q.bytes);
}
if let Some(f) = &mut owned.fragment {
normalize_pct(&mut f.bytes);
}
owned.path = remove_dot_segments_graceful(&owned.path);
if owned.authority.is_some() && owned.path.is_empty() {
owned.path = BytesMut::from(&b"/"[..]);
}
owned
}
fn normalize_pct(buf: &mut BytesMut) {
if !buf.contains(&b'%') {
return;
}
let bytes = buf.as_mut();
let mut read = 0;
let mut write = 0;
while read < bytes.len() {
if bytes[read] == b'%' && read + 2 < bytes.len() {
let h1 = bytes[read + 1];
let h2 = bytes[read + 2];
if let Some(decoded) = rama_utils::hex::decode_pair(h1, h2) {
if is_unreserved_byte(decoded) {
bytes[write] = decoded;
write += 1;
read += 3;
continue;
}
bytes[write] = b'%';
bytes[write + 1] = h1.to_ascii_uppercase();
bytes[write + 2] = h2.to_ascii_uppercase();
write += 3;
read += 3;
continue;
}
}
if write != read {
bytes[write] = bytes[read];
}
write += 1;
read += 1;
}
buf.truncate(write);
}
fn normalize_host_domain(d: &mut Domain) {
let s = d.as_str();
if !s.bytes().any(|b| b.is_ascii_uppercase()) {
return;
}
let lower = Bytes::from(s.to_ascii_lowercase());
*d = unsafe { Domain::from_maybe_borrowed_unchecked(lower) };
}
fn normalize_host_uninterpreted(h: &mut UninterpretedHost) {
let bytes = h.as_bytes();
if !bytes.iter().any(|&b| b == b'%' || b.is_ascii_uppercase()) {
return;
}
let mut buf = BytesMut::from(bytes);
for b in buf.iter_mut() {
*b = b.to_ascii_lowercase();
}
normalize_pct(&mut buf);
*h = UninterpretedHost::from_validated_bytes(buf.freeze(), h.is_bracketed());
}
#[cfg(test)]
mod tests {
use super::*;
fn norm(input: &[u8]) -> Vec<u8> {
let mut b = BytesMut::from(input);
normalize_pct(&mut b);
b.to_vec()
}
#[test]
fn pct_decode_unreserved_alpha() {
assert_eq!(norm(b"exa%6Dple"), b"example");
assert_eq!(norm(b"exa%6dple"), b"example");
}
#[test]
fn pct_decode_unreserved_digit_and_dash() {
assert_eq!(norm(b"path%2D1%2E0"), b"path-1.0");
}
#[test]
fn pct_keeps_reserved_uppercased() {
assert_eq!(norm(b"foo%2fbar"), b"foo%2Fbar");
assert_eq!(norm(b"foo%2Fbar"), b"foo%2Fbar");
}
#[test]
fn pct_keeps_subdelim_uppercased() {
assert_eq!(norm(b"a%26b"), b"a%26b");
}
#[test]
fn pct_no_change_when_input_canonical() {
assert_eq!(norm(b"plain"), b"plain");
assert_eq!(norm(b""), b"");
assert_eq!(norm(b"a/b?c"), b"a/b?c");
}
#[test]
fn pct_mixed_decode_and_uppercase() {
assert_eq!(norm(b"exa%6dple%2fpath"), b"example%2Fpath");
}
#[test]
fn pct_truncated_passthrough() {
assert_eq!(norm(b"x%6"), b"x%6");
}
#[test]
fn host_domain_lowercases_ascii() {
let mut d = Domain::from_static("EXAMPLE.com");
normalize_host_domain(&mut d);
assert_eq!(d.as_str(), "example.com");
}
#[test]
fn host_domain_noop_when_already_lower() {
let mut d = Domain::from_static("example.com");
normalize_host_domain(&mut d);
assert_eq!(d.as_str(), "example.com");
}
fn make_uninterpreted(bytes: &'static [u8], bracketed: bool) -> UninterpretedHost {
UninterpretedHost::from_validated_bytes(Bytes::from_static(bytes), bracketed)
}
#[test]
fn host_uninterpreted_lowercases_ascii_alpha() {
let mut h = make_uninterpreted(b"TAG,WITH,COMMAS", false);
normalize_host_uninterpreted(&mut h);
assert_eq!(h.as_bytes(), b"tag,with,commas");
}
#[test]
fn host_uninterpreted_normalizes_pct_decode_unreserved() {
let mut h = make_uninterpreted(b"exa%6Dple.com", false);
normalize_host_uninterpreted(&mut h);
assert_eq!(h.as_bytes(), b"example.com");
}
#[test]
fn host_uninterpreted_normalizes_pct_keeps_reserved_uppercased() {
let mut h = make_uninterpreted(b"tag%21,more", false);
normalize_host_uninterpreted(&mut h);
assert_eq!(h.as_bytes(), b"tag%21,more");
}
#[test]
fn host_uninterpreted_preserves_bracketed_flag() {
let mut h = make_uninterpreted(b"v1.FE80::A", true);
normalize_host_uninterpreted(&mut h);
assert_eq!(h.as_bytes(), b"v1.fe80::a");
assert!(h.is_bracketed());
}
#[test]
fn host_uninterpreted_noop_when_already_canonical() {
let mut h = make_uninterpreted(b"tag,with,commas", false);
normalize_host_uninterpreted(&mut h);
assert_eq!(h.as_bytes(), b"tag,with,commas");
}
}