use std::collections::BTreeSet;
use std::fmt;
use unicode_normalization::UnicodeNormalization;
use crate::storage::error::FilenameError;
const MAX_INPUT_BYTES: usize = 4096;
pub const MAX_FILENAME_BYTES: usize = 255;
const MAX_EXTENSION_BYTES: usize = 16;
const RESERVED_DEVICE_NAMES: &[&str] = &[
"CON", "PRN", "AUX", "NUL", "CONIN$", "CONOUT$", "COM0", "COM1", "COM2", "COM3", "COM4",
"COM5", "COM6", "COM7", "COM8", "COM9", "LPT0", "LPT1", "LPT2", "LPT3", "LPT4", "LPT5", "LPT6",
"LPT7", "LPT8", "LPT9",
];
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct Extension(String);
impl Extension {
pub fn parse(extension: &str) -> Result<Self, FilenameError> {
if extension.is_empty() {
return Err(FilenameError::MissingExtension);
}
if extension.len() > MAX_EXTENSION_BYTES {
return Err(FilenameError::InvalidExtension);
}
if !extension.bytes().all(|b| b.is_ascii_alphanumeric()) {
return Err(FilenameError::InvalidExtension);
}
Ok(Self(extension.to_ascii_lowercase()))
}
#[must_use]
pub fn as_str(&self) -> &str {
&self.0
}
}
impl fmt::Display for Extension {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str(&self.0)
}
}
impl AsRef<str> for Extension {
fn as_ref(&self) -> &str {
&self.0
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct AllowedExtensions {
allowed: BTreeSet<Extension>,
}
impl AllowedExtensions {
pub fn new<I, S>(extensions: I) -> Result<Self, FilenameError>
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
let mut allowed = BTreeSet::new();
for extension in extensions {
allowed.insert(Extension::parse(extension.as_ref())?);
}
Ok(Self { allowed })
}
#[must_use]
pub fn images() -> Self {
Self::new(["jpg", "jpeg", "png", "gif", "webp"])
.expect("the built-in image extensions are valid")
}
#[must_use]
pub fn documents() -> Self {
Self::new(["pdf", "txt", "csv"]).expect("the built-in document extensions are valid")
}
#[must_use]
pub fn contains(&self, extension: &Extension) -> bool {
self.allowed.contains(extension)
}
pub fn iter(&self) -> impl Iterator<Item = &Extension> {
self.allowed.iter()
}
#[must_use]
pub fn len(&self) -> usize {
self.allowed.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.allowed.is_empty()
}
pub fn with(mut self, extension: &str) -> Result<Self, FilenameError> {
self.allowed.insert(Extension::parse(extension)?);
Ok(self)
}
}
impl<'a> IntoIterator for &'a AllowedExtensions {
type Item = &'a Extension;
type IntoIter = std::collections::btree_set::Iter<'a, Extension>;
fn into_iter(self) -> Self::IntoIter {
self.allowed.iter()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SafeFilename {
stem: String,
extension: Extension,
}
impl SafeFilename {
pub fn parse(filename: &str, allowed: &AllowedExtensions) -> Result<Self, FilenameError> {
if filename.is_empty() {
return Err(FilenameError::Empty);
}
if filename.len() > MAX_INPUT_BYTES {
return Err(FilenameError::TooLong);
}
if filename.chars().any(is_fatal_char) {
return Err(FilenameError::ControlChar);
}
let base = filename
.rsplit(['/', '\\'])
.next()
.unwrap_or_default()
.trim();
if base.is_empty() {
return Err(FilenameError::Empty);
}
if base == "." || base == ".." {
return Err(FilenameError::Traversal);
}
let base: String = base.nfc().collect();
let base = base.trim_end_matches(['.', ' ']).trim_start();
if base.is_empty() {
return Err(FilenameError::Empty);
}
let leading = base
.split('.')
.next()
.unwrap_or_default()
.trim_end_matches([' ', '.']);
if is_reserved_device_name(leading) {
return Err(FilenameError::ReservedName);
}
let (stem_raw, extension_raw) = base
.rsplit_once('.')
.ok_or(FilenameError::MissingExtension)?;
let extension = Extension::parse(extension_raw)?;
if !allowed.contains(&extension) {
return Err(FilenameError::ExtensionNotAllowed);
}
let stem: String = stem_raw
.chars()
.map(|ch| if is_repaired_char(ch) { '_' } else { ch })
.collect();
let stem = stem.trim().to_string();
if stem.is_empty() {
return Err(FilenameError::EmptyStem);
}
if is_reserved_device_name(&stem) {
return Err(FilenameError::ReservedName);
}
let budget = MAX_FILENAME_BYTES.saturating_sub(extension.as_str().len() + 1);
let stem = truncate_on_char_boundary(&stem, budget);
if stem.is_empty() {
return Err(FilenameError::EmptyStem);
}
Ok(Self {
stem: stem.to_string(),
extension,
})
}
#[must_use]
pub fn stem(&self) -> &str {
&self.stem
}
#[must_use]
pub fn extension(&self) -> &Extension {
&self.extension
}
}
impl fmt::Display for SafeFilename {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(formatter, "{}.{}", self.stem, self.extension)
}
}
fn is_fatal_char(ch: char) -> bool {
matches!(ch,
'\u{0}'..='\u{1F}'
| '\u{7F}'..='\u{9F}'
| '\u{00AD}'
| '\u{061C}'
| '\u{180E}'
| '\u{200B}'..='\u{200F}'
| '\u{2028}'..='\u{202E}'
| '\u{2060}'..='\u{2064}'
| '\u{2066}'..='\u{206F}'
| '\u{3164}'
| '\u{FEFF}'
| '\u{FFF9}'..='\u{FFFB}'
| '\u{E0000}'..='\u{E007F}')
}
fn is_repaired_char(ch: char) -> bool {
matches!(
ch,
'.' | '/' | '\\' | ':' | '*' | '?' | '"' | '<' | '>' | '|'
)
}
fn is_reserved_device_name(name: &str) -> bool {
RESERVED_DEVICE_NAMES
.iter()
.any(|reserved| reserved.eq_ignore_ascii_case(name))
}
fn truncate_on_char_boundary(text: &str, budget: usize) -> &str {
if text.len() <= budget {
return text;
}
let mut end = budget;
while end > 0 && !text.is_char_boundary(end) {
end -= 1;
}
&text[..end]
}
#[cfg(test)]
mod tests {
use super::*;
fn allowed() -> AllowedExtensions {
AllowedExtensions::new(["jpg", "jpeg", "png", "txt", "pdf"]).unwrap()
}
fn parse(filename: &str) -> Result<SafeFilename, FilenameError> {
SafeFilename::parse(filename, &allowed())
}
#[test]
fn accepts_an_ordinary_name() {
assert_eq!(parse("holiday.jpg").unwrap().to_string(), "holiday.jpg");
}
#[test]
fn lowercases_the_extension() {
assert_eq!(parse("HOLIDAY.JPG").unwrap().to_string(), "HOLIDAY.jpg");
}
#[test]
fn accepts_a_vietnamese_name_with_spaces_and_diacritics() {
let name = parse("Ảnh chụp màn hình.png").unwrap();
assert_eq!(name.to_string(), "Ảnh chụp màn hình.png");
}
#[test]
fn normalizes_to_nfc() {
let decomposed = parse("A\u{0309}nh.png").unwrap();
let composed = parse("\u{1EA2}nh.png").unwrap();
assert_eq!(decomposed, composed);
}
#[test]
fn strips_unix_directory_components() {
assert_eq!(
parse("/var/www/photo.jpg").unwrap().to_string(),
"photo.jpg"
);
}
#[test]
fn strips_windows_directory_components() {
assert_eq!(
parse(r"C:\Users\me\photo.jpg").unwrap().to_string(),
"photo.jpg"
);
}
#[test]
fn traversal_loses_its_directory_part() {
assert_eq!(
parse("../../etc/passwd"),
Err(FilenameError::MissingExtension)
);
assert_eq!(
parse("../../etc/passwd.txt").unwrap().to_string(),
"passwd.txt"
);
}
#[test]
fn rejects_bare_dot_segments() {
assert_eq!(parse(".."), Err(FilenameError::Traversal));
assert_eq!(parse("."), Err(FilenameError::Traversal));
}
#[test]
fn rejects_a_nul_byte() {
assert_eq!(parse("a\0b.jpg"), Err(FilenameError::ControlChar));
}
#[test]
fn rejects_a_nul_hidden_in_the_directory_part() {
assert_eq!(parse("../\0/photo.jpg"), Err(FilenameError::ControlChar));
}
#[test]
fn rejects_a_newline() {
assert_eq!(parse("a\r\nb.jpg"), Err(FilenameError::ControlChar));
}
#[test]
fn rejects_a_right_to_left_override() {
assert_eq!(
parse("invoice\u{202E}gpj.txt"),
Err(FilenameError::ControlChar)
);
}
#[test]
fn rejects_a_zero_width_joiner() {
assert_eq!(parse("pho\u{200D}to.jpg"), Err(FilenameError::ControlChar));
}
#[test]
fn rejects_a_unicode_tag_character() {
assert_eq!(parse("pho\u{E0041}to.jpg"), Err(FilenameError::ControlChar));
assert_eq!(parse("photo\u{E007F}.jpg"), Err(FilenameError::ControlChar));
}
#[test]
fn rejects_the_blank_rendering_separators() {
for blank in ['\u{00AD}', '\u{180E}', '\u{3164}'] {
let name = format!("pho{blank}to.jpg");
assert_eq!(
parse(&name),
Err(FilenameError::ControlChar),
"expected U+{:04X} to be fatal",
blank as u32
);
}
}
#[test]
fn keeps_a_variation_selector() {
let name = parse("heart\u{2764}\u{FE0F}.jpg").unwrap();
assert_eq!(name.to_string(), "heart\u{2764}\u{FE0F}.jpg");
}
#[test]
fn rejects_windows_device_names() {
for name in ["CON.txt", "con.txt", "NUL.txt", "AUX.txt", "COM1.txt"] {
assert_eq!(parse(name), Err(FilenameError::ReservedName), "{name}");
}
}
#[test]
fn rejects_a_device_name_behind_a_second_extension() {
assert_eq!(parse("CON.foo.txt"), Err(FilenameError::ReservedName));
}
#[test]
fn rejects_a_device_name_with_trailing_dots_and_spaces() {
assert_eq!(parse("CON.txt. "), Err(FilenameError::ReservedName));
}
#[test]
fn repairs_a_double_extension() {
assert_eq!(parse("shell.php.jpg").unwrap().to_string(), "shell_php.jpg");
}
#[test]
fn repairs_an_alternate_data_stream() {
assert_eq!(
parse("photo.jpg:payload.txt").unwrap().to_string(),
"photo_jpg_payload.txt"
);
}
#[test]
fn rejects_an_extension_outside_the_whitelist() {
for name in ["shell.php", "virus.exe", "page.phtml", "boot.svg"] {
assert_eq!(
parse(name),
Err(FilenameError::ExtensionNotAllowed),
"{name}"
);
}
}
#[test]
fn rejects_a_missing_extension() {
assert_eq!(parse("passwd"), Err(FilenameError::MissingExtension));
}
#[test]
fn rejects_a_dotfile_with_no_stem() {
assert_eq!(parse(".jpg"), Err(FilenameError::EmptyStem));
}
#[test]
fn rejects_an_empty_name() {
assert_eq!(parse(""), Err(FilenameError::Empty));
assert_eq!(parse(" "), Err(FilenameError::Empty));
assert_eq!(parse("/var/www/"), Err(FilenameError::Empty));
}
#[test]
fn rejects_an_over_long_input() {
let long = format!("{}.jpg", "a".repeat(MAX_INPUT_BYTES));
assert_eq!(
SafeFilename::parse(&long, &allowed()),
Err(FilenameError::TooLong)
);
}
#[test]
fn truncates_a_long_stem_on_a_character_boundary() {
let name = parse(&format!("{}.jpg", "ả".repeat(200))).unwrap();
assert!(name.to_string().len() <= MAX_FILENAME_BYTES);
assert!(name.stem().chars().all(|ch| ch == 'ả'));
}
#[test]
fn an_empty_whitelist_stores_nothing() {
let empty = AllowedExtensions::default();
assert!(empty.is_empty());
assert_eq!(
SafeFilename::parse("holiday.jpg", &empty),
Err(FilenameError::ExtensionNotAllowed)
);
}
#[test]
fn extension_parsing_is_strict() {
assert_eq!(Extension::parse(""), Err(FilenameError::MissingExtension));
let too_long = "a".repeat(MAX_EXTENSION_BYTES + 1);
for bad in ["jp g", "jpg.", "jpg/", "jpé", "j-pg", too_long.as_str()] {
assert_eq!(
Extension::parse(bad),
Err(FilenameError::InvalidExtension),
"{bad}"
);
}
assert_eq!(Extension::parse("mp4").unwrap().as_str(), "mp4");
}
#[test]
fn built_in_whitelists_are_what_they_say() {
let images = AllowedExtensions::images();
assert_eq!(images.len(), 5);
assert!(!images.contains(&Extension::parse("svg").unwrap()));
let documents = AllowedExtensions::documents();
assert!(documents.contains(&Extension::parse("pdf").unwrap()));
}
}