use anyhow::Context;
use hyper::StatusCode;
use percent_encoding::percent_decode_str;
use std::path::{Component, Path, PathBuf};
use crate::Result;
pub(crate) trait PathExt {
fn is_hidden(&self) -> bool;
fn contains_symlink(&self, base: &Path) -> Result<bool>;
}
impl PathExt for Path {
fn is_hidden(&self) -> bool {
self.components()
.filter_map(|cmp| match cmp {
Component::Normal(s) => s.to_str(),
_ => None,
})
.any(|s| s.starts_with('.'))
}
fn contains_symlink(&self, base: &Path) -> Result<bool> {
let mut current = base.to_path_buf();
current.reserve(self.as_os_str().len());
for component in self.components() {
match component {
Component::Normal(c) => {
current.push(c);
let meta = std::fs::symlink_metadata(¤t).with_context(|| {
format!("unable to get metadata for path '{}'", current.display())
})?;
if meta.file_type().is_symlink() {
return Ok(true);
}
}
Component::CurDir => {}
Component::Prefix(_) | Component::RootDir | Component::ParentDir => {
tracing::debug!(
"dir: skipping segment containing invalid prefix, dots or backslashes"
);
}
}
}
Ok(false)
}
}
#[cfg(unix)]
fn path_from_bytes(bytes: &[u8]) -> PathBuf {
use std::ffi::OsStr;
use std::os::unix::ffi::OsStrExt;
OsStr::from_bytes(bytes).into()
}
#[cfg(windows)]
fn path_from_bytes(bytes: &[u8]) -> PathBuf {
String::from_utf8_lossy(bytes).into_owned().into()
}
fn decode_tail_path(tail: &str) -> PathBuf {
let bytes = percent_decode_str(tail.trim_start_matches('/')).collect::<Vec<_>>();
path_from_bytes(&bytes)
}
#[doc(hidden)]
pub fn sanitize_path(base: &Path, tail: &str) -> Result<PathBuf, StatusCode> {
let path_decoded = decode_tail_path(tail);
let mut full_path = base.to_path_buf();
tracing::trace!("dir: base={:?}, route={:?}", full_path, path_decoded);
for component in path_decoded.components() {
match component {
Component::Normal(comp) => {
if Path::new(&comp)
.components()
.all(|c| matches!(c, Component::Normal(_)))
{
full_path.push(comp)
} else {
tracing::debug!("dir: skipping segment with invalid prefix");
}
}
Component::CurDir => {}
Component::Prefix(_) | Component::RootDir | Component::ParentDir => {
tracing::debug!(
"dir: skipping segment containing invalid prefix, dots or backslashes"
);
}
}
}
Ok(full_path)
}
#[cfg(test)]
mod tests {
use super::{PathExt, sanitize_path};
use std::path::PathBuf;
fn root_dir() -> PathBuf {
PathBuf::from("docker/public/")
}
#[test]
fn test_sanitize_path() {
let base_dir = &PathBuf::from("docker/public");
assert_eq!(
sanitize_path(base_dir, "/index.html").unwrap(),
root_dir().join("index.html")
);
assert_eq!(
sanitize_path(base_dir, "/../foo.html").unwrap(),
root_dir().join("foo.html"),
);
assert_eq!(
sanitize_path(base_dir, "/../W�foo.html").unwrap(),
root_dir().join("W�foo.html"),
);
assert_eq!(
sanitize_path(base_dir, "/%EF%BF%BD/../bar.html").unwrap(),
root_dir().join("�/bar.html"),
);
assert_eq!(
sanitize_path(base_dir, "àí/é%20/öüñ").unwrap(),
root_dir().join("àí/é /öüñ"),
);
#[cfg(unix)]
let expected_path = root_dir().join("C:\\/foo.html");
#[cfg(windows)]
let expected_path = PathBuf::from("docker/public/\\foo.html");
assert_eq!(
sanitize_path(base_dir, "/C:\\/foo.html").unwrap(),
expected_path
);
}
#[test]
fn test_contains_symlink_returns_false() {
let base = PathBuf::from("tests/fixtures/public");
let user_path = PathBuf::from("./index.htm");
match user_path.contains_symlink(&base) {
Ok(contains) => assert!(!contains),
Err(err) => panic!("unexpected error when checking for symlinks: {err}"),
}
}
#[test]
fn test_contains_symlink_returns_true() {
let base = PathBuf::from("tests/fixtures/public");
let user_path = PathBuf::from("./readme.md");
match user_path.contains_symlink(&base) {
Ok(contains) => assert!(contains),
Err(err) => panic!("unexpected error when checking for symlinks: {err}"),
}
}
#[test]
fn test_contains_symlink_returns_error() {
let base = PathBuf::from("tests/fixtures/public");
let user_path = PathBuf::from("./unknown_file.txt");
match user_path.contains_symlink(&base) {
Ok(_) => panic!("expected error when checking for symlinks on non-existent path"),
Err(err) => assert!(
err.to_string().contains("unable to get metadata for path"),
"unexpected error message: {err}"
),
}
}
use proptest::prelude::*;
use std::path::Component;
fn assert_under_base(base: &std::path::Path, full: &std::path::Path) {
let extra = full
.strip_prefix(base)
.expect("sanitized path must extend base");
for comp in extra.components() {
match comp {
Component::Normal(_) | Component::CurDir => {}
Component::ParentDir | Component::RootDir | Component::Prefix(_) => {
panic!("sanitize_path leaked traversal component {comp:?} in {full:?}");
}
}
}
}
proptest! {
#![proptest_config(ProptestConfig {
cases: 256, ..ProptestConfig::default()
})]
#[test]
fn prop_sanitize_path_never_escapes_base(tail in "\\PC{0,128}") {
let base = PathBuf::from("docker/public");
if let Ok(full) = sanitize_path(&base, &tail) {
assert_under_base(&base, &full);
}
}
#[test]
fn prop_sanitize_path_resists_dot_dot_floods(
tail in proptest::collection::vec(
prop_oneof![
Just("../".to_string()),
Just("..\\".to_string()),
Just("%2e%2e/".to_string()),
Just("./".to_string()),
"[a-zA-Z0-9_.-]{1,8}".prop_map(String::from),
],
0..32,
)
) {
let base = PathBuf::from("docker/public");
let tail = tail.concat();
if let Ok(full) = sanitize_path(&base, &tail) {
assert_under_base(&base, &full);
}
}
}
}