use crate::std::sync::Arc;
use super::owned::OwnedUriRef;
use super::parser::MAX_URI_LEN;
use super::{Uri, UriInner};
use rama_core::bytes::BytesMut;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ResolveError {
BaseHasNoScheme,
AsteriskNotResolvable,
ResultTooLong { len: usize },
DotSegmentTraversalPastRoot,
}
impl core::fmt::Display for ResolveError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::BaseHasNoScheme => f.write_str("base URI has no scheme"),
Self::AsteriskNotResolvable => {
f.write_str("asterisk-form URI cannot be used as a base or reference")
}
Self::ResultTooLong { len } => write!(f, "resolved URI is {len} bytes — exceeds cap"),
Self::DotSegmentTraversalPastRoot => {
f.write_str("`..` segment would traverse past path root (strict mode)")
}
}
}
}
impl core::error::Error for ResolveError {}
#[derive(Clone, Copy, PartialEq, Eq)]
pub(super) enum ResolveMode {
Graceful,
Strict,
}
pub(super) fn resolve(base: &Uri, reference: &Uri, mode: ResolveMode) -> Result<Uri, ResolveError> {
if matches!(base.inner, UriInner::Asterisk) || matches!(reference.inner, UriInner::Asterisk) {
return Err(ResolveError::AsteriskNotResolvable);
}
if base.scheme().is_none() {
return Err(ResolveError::BaseHasNoScheme);
}
let OwnedUriRef {
scheme: b_scheme,
authority: b_authority,
path: b_path,
query: b_query,
fragment: _,
} = base.as_owned_components();
let OwnedUriRef {
scheme: r_scheme,
authority: r_authority,
path: r_path,
query: r_query,
fragment: r_fragment,
} = reference.as_owned_components();
let r_has_effective_scheme = match (mode, &r_scheme, &b_scheme) {
(ResolveMode::Graceful, Some(r_s), Some(b_s)) if r_s == b_s => false,
_ => r_scheme.is_some(),
};
let (t_scheme, t_authority, t_path, t_query) = if r_has_effective_scheme {
(
r_scheme,
r_authority,
remove_dot_segments(&r_path, mode)?,
r_query.map(|q| q.bytes),
)
} else if r_authority.is_some() {
(
b_scheme,
r_authority,
remove_dot_segments(&r_path, mode)?,
r_query.map(|q| q.bytes),
)
} else if r_path.is_empty() {
let query = r_query
.map(|q| q.bytes)
.or_else(|| b_query.map(|q| q.bytes));
(b_scheme, b_authority, b_path, query)
} else {
let raw_path = if r_path.starts_with(b"/") {
r_path
} else {
merge_paths(b_authority.is_some(), &b_path, &r_path)
};
(
b_scheme,
b_authority,
remove_dot_segments(&raw_path, mode)?,
r_query.map(|q| q.bytes),
)
};
let t_fragment = r_fragment.map(|f| f.bytes);
let owned = OwnedUriRef {
scheme: t_scheme,
authority: t_authority,
path: t_path,
query: t_query.map(|bytes| super::Query { bytes }),
fragment: t_fragment.map(|bytes| super::Fragment { bytes }),
};
let total = serialized_len(&owned);
if total > MAX_URI_LEN {
return Err(ResolveError::ResultTooLong { len: total });
}
Ok(Uri {
inner: UriInner::Owned(Arc::new(owned)),
})
}
fn merge_paths(b_has_authority: bool, b_path: &[u8], r_path: &[u8]) -> BytesMut {
if b_has_authority && b_path.is_empty() {
let mut out = BytesMut::with_capacity(1 + r_path.len());
out.extend_from_slice(b"/");
out.extend_from_slice(r_path);
return out;
}
let cutoff = memchr::memrchr(b'/', b_path).map_or(0, |i| i + 1);
let mut out = BytesMut::with_capacity(cutoff + r_path.len());
out.extend_from_slice(&b_path[..cutoff]);
out.extend_from_slice(r_path);
out
}
fn remove_dot_segments(input: &[u8], mode: ResolveMode) -> Result<BytesMut, ResolveError> {
let mut output = BytesMut::with_capacity(input.len());
let mut i = 0;
while i < input.len() {
let rest = &input[i..];
if rest.starts_with(b"../") {
i += 3;
continue;
}
if rest.starts_with(b"./") {
i += 2;
continue;
}
if rest.starts_with(b"/./") {
i += 2;
continue;
}
if rest == b"/." {
output.extend_from_slice(b"/");
break;
}
if rest.starts_with(b"/../") {
pop_last_segment(&mut output, mode)?;
i += 3;
continue;
}
if rest == b"/.." {
pop_last_segment(&mut output, mode)?;
output.extend_from_slice(b"/");
break;
}
if rest == b"." || rest == b".." {
break;
}
let seg_end = if rest[0] == b'/' {
memchr::memchr(b'/', &rest[1..]).map_or(rest.len(), |p| p + 1)
} else {
memchr::memchr(b'/', rest).unwrap_or(rest.len())
};
output.extend_from_slice(&rest[..seg_end]);
i += seg_end;
}
Ok(output)
}
fn pop_last_segment(output: &mut BytesMut, mode: ResolveMode) -> Result<(), ResolveError> {
if let Some(last_slash) = memchr::memrchr(b'/', output) {
output.truncate(last_slash);
return Ok(());
}
if !output.is_empty() {
output.clear();
return Ok(());
}
match mode {
ResolveMode::Strict => Err(ResolveError::DotSegmentTraversalPastRoot),
ResolveMode::Graceful => Ok(()),
}
}
pub(super) fn remove_dot_segments_graceful(input: &[u8]) -> BytesMut {
remove_dot_segments(input, ResolveMode::Graceful).unwrap_or_else(|_| {
debug_assert!(false, "graceful remove_dot_segments must not error");
BytesMut::from(input)
})
}
fn serialized_len(owned: &OwnedUriRef) -> usize {
use core::fmt::Write as _;
let mut n = 0;
if let Some(scheme) = &owned.scheme {
n += scheme.as_str().len() + 1; }
if let Some(auth) = &owned.authority {
n += 2; if let Some(ui) = &auth.user_info {
n += ui.as_bytes().len() + 1; }
let mut counter = FmtLenCounter(0);
#[expect(
clippy::let_underscore_must_use,
reason = "FmtLenCounter::write_str is infallible by construction"
)]
let _ = write!(&mut counter, "{}", auth.address.host);
n += counter.0;
match auth.address.port {
crate::address::OptPort::Unset => {}
crate::address::OptPort::Empty => {
n += 1; }
crate::address::OptPort::Set(port) => {
n += 1 + port_decimal_len(port); }
}
}
n += owned.path.len();
if let Some(q) = &owned.query {
n += 1 + q.bytes.len(); }
if let Some(f) = &owned.fragment {
n += 1 + f.bytes.len(); }
n
}
struct FmtLenCounter(usize);
impl core::fmt::Write for FmtLenCounter {
fn write_str(&mut self, s: &str) -> core::fmt::Result {
self.0 += s.len();
Ok(())
}
}
#[inline]
fn port_decimal_len(port: u16) -> usize {
match port {
0..=9 => 1,
10..=99 => 2,
100..=999 => 3,
1000..=9999 => 4,
_ => 5,
}
}