use core::net::{Ipv4Addr, Ipv6Addr};
use core::str::FromStr;
use crate::error::ParseError;
use crate::util::IpAddr;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct Uri<'a> {
scheme: &'a str,
authority: Option<Authority<'a>>,
path: &'a str,
query: Option<&'a str>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct Authority<'a> {
host: Host<'a>,
port: Option<u16>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum Host<'a> {
IpAddr(IpAddr),
RegName(&'a str),
}
impl<'a> Uri<'a> {
pub fn parse(input: &'a str) -> Result<Self, ParseError> {
Parser::new(input).parse_uri()
}
#[must_use]
pub const fn scheme(&self) -> &'a str {
self.scheme
}
#[must_use]
pub const fn authority(&self) -> Option<&Authority<'a>> {
self.authority.as_ref()
}
#[must_use]
pub fn port_or_default(&self) -> u16 {
self
.authority()
.and_then(Authority::port)
.unwrap_or_else(|| {
if self.scheme.eq_ignore_ascii_case("https") {
443
} else {
80
}
})
}
#[must_use]
pub const fn path(&self) -> &'a str {
self.path
}
#[must_use]
pub const fn query(&self) -> Option<&'a str> {
self.query
}
#[must_use]
pub fn to_path_and_query(&self) -> alloc::string::String {
let path = if self.path().is_empty() {
"/"
} else {
self.path()
};
self.query.map_or_else(
|| alloc::string::String::from(path),
|query| alloc::format!("{path}?{query}"),
)
}
pub fn resolve_relative(
&self,
location: &str,
) -> Result<alloc::string::String, ParseError> {
if let Some((scheme, rest)) = location.split_once(':')
&& (scheme.eq_ignore_ascii_case("http") || scheme.eq_ignore_ascii_case("https"))
{
let scheme_lower = scheme.to_ascii_lowercase();
return Ok(alloc::format!("{scheme_lower}:{rest}"));
}
if let Some(rest) = location.strip_prefix("//") {
return Ok(alloc::format!("{}://{rest}", self.scheme));
}
if location.starts_with('?') {
let path = if self.path.is_empty() {
"/"
} else {
self.path
};
return self.recompose_with_path(&alloc::format!("{path}{location}"));
}
let path = if location.starts_with('/') {
alloc::string::String::from(location)
} else {
if self.path.is_empty() {
alloc::format!("/{location}")
} else {
let dir_end = self.path.rfind('/').map_or(0, |i| i.saturating_add(1));
let prefix = self.path.get(..dir_end).unwrap_or("");
alloc::format!("{prefix}{location}")
}
};
self.recompose_with_path(&path)
}
fn recompose_with_path(
&self,
path: &str,
) -> Result<alloc::string::String, ParseError> {
let authority = self.authority.as_ref().ok_or(ParseError::InvalidUri)?;
let port = self.port_or_default();
let host_str = match &authority.host {
Host::RegName(name) => alloc::string::String::from(*name),
Host::IpAddr(addr) => crate::util::format_ip_for_host(*addr),
};
if (self.scheme.eq_ignore_ascii_case("http") && port == 80)
|| (self.scheme.eq_ignore_ascii_case("https") && port == 443)
{
Ok(alloc::format!(
"{scheme}://{host}{path}",
scheme = self.scheme,
host = host_str
))
} else {
Ok(alloc::format!(
"{scheme}://{host}:{port}{path}",
scheme = self.scheme,
host = host_str
))
}
}
}
impl<'a> Authority<'a> {
#[must_use]
pub const fn host(&self) -> &Host<'a> {
&self.host
}
#[must_use]
pub const fn port(&self) -> Option<u16> {
self.port
}
}
struct Parser<'a> {
input: &'a str,
pos: usize,
}
impl<'a> Parser<'a> {
const fn new(input: &'a str) -> Self {
Self { input, pos: 0 }
}
fn peek(&self) -> Option<u8> {
self.input.as_bytes().get(self.pos).copied()
}
fn peek_at(
&self,
offset: usize,
) -> Option<u8> {
let idx = self.pos.saturating_add(offset);
self.input.as_bytes().get(idx).copied()
}
const fn advance(&mut self) {
if self.pos < self.input.len() {
self.pos = self.pos.saturating_add(1);
}
}
fn advance_by(
&mut self,
n: usize,
) {
self.pos = self.pos.saturating_add(n).min(self.input.len());
}
fn slice_from(
&self,
start: usize,
) -> &'a str {
&self.input[start..self.pos]
}
fn parse_uri(mut self) -> Result<Uri<'a>, ParseError> {
let scheme = self.parse_scheme()?;
if self.peek() != Some(b':') || self.peek_at(1) != Some(b'/') || self.peek_at(2) != Some(b'/') {
return Err(ParseError::InvalidUri);
}
self.advance_by(3);
let authority = self.parse_authority()?;
let path = self.parse_path_abempty();
let query = if self.peek() == Some(b'?') {
self.advance();
Some(self.parse_query()?)
} else {
None
};
if self.pos != self.input.len() {
return Err(ParseError::InvalidUri);
}
Ok(Uri {
scheme,
authority: Some(authority),
path,
query,
})
}
fn parse_scheme(&mut self) -> Result<&'a str, ParseError> {
let start = self.pos;
let rest = &self.input[start..];
let scheme = if rest
.get(..5)
.is_some_and(|s| s.eq_ignore_ascii_case("https"))
{
self.advance_by(5);
self.slice_from(start)
} else if rest
.get(..4)
.is_some_and(|s| s.eq_ignore_ascii_case("http"))
{
self.advance_by(4);
self.slice_from(start)
} else {
return Err(ParseError::InvalidUri);
};
Ok(scheme)
}
fn parse_authority(&mut self) -> Result<Authority<'a>, ParseError> {
if self.find_char_in_authority(b'@') {
return Err(ParseError::InvalidUri);
}
let host = self.parse_host()?;
let port = if self.peek() == Some(b':') {
self.advance();
Some(self.parse_port()?)
} else {
None
};
Ok(Authority { host, port })
}
fn find_char_in_authority(
&self,
target: u8,
) -> bool {
let mut pos = self.pos;
let bytes = self.input.as_bytes();
while let Some(&ch) = bytes.get(pos) {
match ch {
b'/' | b'?' | b'#' => return false,
_ if ch == target => return true,
_ => pos = pos.saturating_add(1),
}
}
false
}
fn parse_host(&mut self) -> Result<Host<'a>, ParseError> {
if self.peek() == Some(b'[') {
return self.parse_ip_literal();
}
let start = self.pos;
while let Some(ch) = self.peek() {
match ch {
b':' | b'/' | b'?' | b'#' => break,
_ if is_reg_name_char(ch) => self.advance(),
_ => break,
}
}
let host_str = self.slice_from(start);
if host_str.is_empty() {
return Err(ParseError::InvalidUri);
}
if looks_like_ipv4(host_str)
&& let Ok(v4) = Ipv4Addr::from_str(host_str)
{
return Ok(Host::IpAddr(IpAddr::V4(v4)));
}
Ok(Host::RegName(host_str))
}
fn parse_ip_literal(&mut self) -> Result<Host<'a>, ParseError> {
if self.peek() != Some(b'[') {
return Err(ParseError::InvalidUri);
}
self.advance();
let start = self.pos;
while let Some(ch) = self.peek() {
if ch == b']' {
break;
}
self.advance();
}
if self.peek() != Some(b']') {
return Err(ParseError::InvalidUri);
}
let addr_str = self.slice_from(start);
self.advance();
let v6 = parse_ipv6(addr_str)?;
Ok(Host::IpAddr(IpAddr::V6(v6)))
}
fn parse_port(&mut self) -> Result<u16, ParseError> {
let start = self.pos;
while let Some(b'0'..=b'9') = self.peek() {
self.advance();
}
if start == self.pos {
return Ok(0);
}
let port_str = self.slice_from(start);
port_str.parse::<u16>().map_err(|_| ParseError::InvalidUri)
}
fn parse_path_abempty(&mut self) -> &'a str {
let start = self.pos;
while self.peek() == Some(b'/') {
self.advance();
while let Some(ch) = self.peek() {
if is_pchar(ch) {
self.advance();
} else {
break;
}
}
}
self.slice_from(start)
}
fn parse_query(&mut self) -> Result<&'a str, ParseError> {
let start = self.pos;
while let Some(ch) = self.peek() {
match ch {
b'#' => break,
_ if is_pchar(ch) || ch == b'/' || ch == b'?' => {
self.advance();
},
_ => return Err(ParseError::InvalidUri),
}
}
Ok(self.slice_from(start))
}
}
const fn is_unreserved(ch: u8) -> bool {
ch.is_ascii_alphanumeric() || matches!(ch, b'-' | b'.' | b'_' | b'~')
}
const fn is_sub_delim(ch: u8) -> bool {
matches!(
ch,
b'!' | b'$' | b'&' | b'\'' | b'(' | b')' | b'*' | b'+' | b',' | b';' | b'='
)
}
const fn is_pchar(ch: u8) -> bool {
is_unreserved(ch) || is_sub_delim(ch) || ch == b':' || ch == b'@' || ch == b'%'
}
const fn is_reg_name_char(ch: u8) -> bool {
is_unreserved(ch) || is_sub_delim(ch) || ch == b'%'
}
#[inline]
fn looks_like_ipv4(host: &str) -> bool {
let bytes = host.as_bytes();
if bytes.is_empty() {
return false;
}
let mut has_dot = false;
for &b in bytes {
match b {
b'0'..=b'9' => {},
b'.' => has_dot = true,
_ => return false,
}
}
has_dot
}
fn parse_ipv6(s: &str) -> Result<Ipv6Addr, ParseError> {
let bytes = s.as_bytes();
let mut pos = 0usize;
let mut head = [0u16; 8];
let (head_size, head_ipv4) = read_ipv6_groups(bytes, &mut pos, &mut head)?;
if head_size == 8 {
if pos != bytes.len() {
return Err(ParseError::InvalidUri);
}
return Ok(Ipv6Addr::from(head));
}
if head_ipv4 {
return Err(ParseError::InvalidUri);
}
if bytes.get(pos).copied() != Some(b':') || bytes.get(pos.saturating_add(1)).copied() != Some(b':') {
return Err(ParseError::InvalidUri);
}
pos = pos.saturating_add(2);
let mut tail = [0u16; 7];
let limit = 8usize.saturating_sub(head_size.saturating_add(1));
let tail_slot = tail.get_mut(..limit).ok_or(ParseError::InvalidUri)?;
let (tail_size, _) = read_ipv6_groups(bytes, &mut pos, tail_slot)?;
if pos != bytes.len() {
return Err(ParseError::InvalidUri);
}
let fill_at = 8usize.saturating_sub(tail_size);
let dest = head.get_mut(fill_at..8).ok_or(ParseError::InvalidUri)?;
let src = tail.get(..tail_size).ok_or(ParseError::InvalidUri)?;
dest.copy_from_slice(src);
Ok(Ipv6Addr::from(head))
}
fn read_ipv6_groups(
bytes: &[u8],
pos: &mut usize,
groups: &mut [u16],
) -> Result<(usize, bool), ParseError> {
let limit = groups.len();
let mut i = 0usize;
while i < limit {
if i < limit.saturating_sub(1) {
let save = *pos;
if i > 0 {
if bytes.get(*pos).copied() == Some(b':') {
*pos = pos.saturating_add(1);
} else {
*pos = save;
return Ok((i, false));
}
}
match read_embedded_ipv4(bytes, pos) {
Some(v4) => {
let oct = v4.octets();
let Some(slot0) = groups.get_mut(i) else {
return Err(ParseError::InvalidUri);
};
*slot0 = u16::from_be_bytes([oct[0], oct[1]]);
let Some(slot1) = groups.get_mut(i.saturating_add(1)) else {
return Err(ParseError::InvalidUri);
};
*slot1 = u16::from_be_bytes([oct[2], oct[3]]);
return Ok((i.saturating_add(2), true));
},
None => {
*pos = save;
},
}
}
let save = *pos;
if i > 0 {
if bytes.get(*pos).copied() == Some(b':') {
*pos = pos.saturating_add(1);
} else {
*pos = save;
return Ok((i, false));
}
}
if let Some(g) = read_hextet(bytes, pos) {
let Some(slot) = groups.get_mut(i) else {
return Err(ParseError::InvalidUri);
};
*slot = g;
i = i.saturating_add(1);
} else {
*pos = save;
return Ok((i, false));
}
}
Ok((limit, false))
}
fn read_hextet(
bytes: &[u8],
pos: &mut usize,
) -> Option<u16> {
let start = *pos;
let mut value: u16 = 0;
let mut digits = 0u32;
while let Some(&b) = bytes.get(*pos) {
let digit = match b {
b'0'..=b'9' => u16::from(b - b'0'),
b'a'..=b'f' => u16::from(b - b'a') + 10,
b'A'..=b'F' => u16::from(b - b'A') + 10,
_ => break,
};
if digits >= 4 {
*pos = start;
return None;
}
value = (value << 4) | digit;
digits = digits.saturating_add(1);
*pos = pos.saturating_add(1);
}
if digits == 0 {
*pos = start;
None
} else {
Some(value)
}
}
fn read_embedded_ipv4(
bytes: &[u8],
pos: &mut usize,
) -> Option<Ipv4Addr> {
let start = *pos;
let mut octets = [0u8; 4];
for (i, slot) in octets.iter_mut().enumerate() {
if i > 0 {
if bytes.get(*pos).copied() != Some(b'.') {
*pos = start;
return None;
}
*pos = pos.saturating_add(1);
}
if let Some(o) = read_decimal_octet(bytes, pos) {
*slot = o;
} else {
*pos = start;
return None;
}
}
Some(Ipv4Addr::from(octets))
}
fn read_decimal_octet(
bytes: &[u8],
pos: &mut usize,
) -> Option<u8> {
let first = bytes.get(*pos).copied()?;
if !first.is_ascii_digit() {
return None;
}
*pos = pos.saturating_add(1);
let mut value = u32::from(first - b'0');
let mut digits = 1u32;
while let Some(&b) = bytes.get(*pos) {
if !b.is_ascii_digit() {
break;
}
if digits >= 3 {
return None;
}
value = value.saturating_mul(10).saturating_add(u32::from(b - b'0'));
digits = digits.saturating_add(1);
*pos = pos.saturating_add(1);
}
if first == b'0' && digits > 1 {
return None;
}
u8::try_from(value).ok()
}
impl core::fmt::Display for Uri<'_> {
fn fmt(
&self,
f: &mut core::fmt::Formatter<'_>,
) -> core::fmt::Result {
f.write_str(self.scheme)?;
f.write_str("://")?;
if let Some(auth) = &self.authority {
match &auth.host {
Host::RegName(name) => f.write_str(name)?,
Host::IpAddr(IpAddr::V4(v4)) => write!(f, "{v4}")?,
Host::IpAddr(IpAddr::V6(v6)) => write!(f, "[{v6}]")?,
}
if let Some(port) = auth.port {
write!(f, ":{port}")?;
}
}
f.write_str(self.path)?;
if let Some(query) = self.query {
write!(f, "?{query}")?;
}
Ok(())
}
}
impl<'a> core::convert::TryFrom<&'a str> for Uri<'a> {
type Error = ParseError;
fn try_from(s: &'a str) -> Result<Self, Self::Error> {
Self::parse(s)
}
}
#[cfg(kani)]
mod kani_uri_proofs {
use super::{Host, Uri};
#[kani::proof]
fn reg_name_preserves_case() {
let uri = Uri::parse("http://Example.COM/path").unwrap();
match uri.authority().unwrap().host() {
Host::RegName(name) => assert_eq!(*name, "Example.COM"),
Host::IpAddr(_) => panic!("expected RegName"),
}
}
#[kani::proof]
fn https_default_port_case_insensitive() {
assert_eq!(
Uri::parse("HTTPS://example.com/")
.unwrap()
.port_or_default(),
443
);
assert_eq!(Uri::parse("Http://example.com/").unwrap().port_or_default(), 80);
}
#[kani::proof]
#[kani::unwind(16)]
fn ascii_lowercase_idempotent_label() {
let mut raw = [0u8; 4];
for b in &mut raw {
*b = kani::any();
kani::assume(b.is_ascii_alphanumeric());
}
let s = core::str::from_utf8(&raw).unwrap();
let once = s.to_ascii_lowercase();
let twice = once.to_ascii_lowercase();
assert_eq!(once, twice);
}
}