use crate::constants::{
DEFAULT_MAX_FILE_LEN, DOT, DOTDOT, REPLACEMENT_CHAR, REPLACEMENT_UNDERSCORE,
};
use crate::error::{SafeNameError, SafeNameResult};
use crate::lookup::{BLOCKED_ALWAYS, BLOCKED_FINAL, BLOCKED_INITIAL};
#[derive(Debug, Clone)]
pub struct FileSanitizationOptions<'a> {
pub replacement: &'a [u8],
pub max_len: usize,
}
impl Default for FileSanitizationOptions<'_> {
fn default() -> Self {
Self {
replacement: REPLACEMENT_UNDERSCORE,
max_len: DEFAULT_MAX_FILE_LEN,
}
}
}
impl FileSanitizationOptions<'_> {
pub fn with_unicode_replacement_char() -> Self {
Self {
replacement: REPLACEMENT_CHAR,
..Default::default()
}
}
}
pub fn sanitize_file(filename: impl AsRef<[u8]>) -> SafeNameResult<Vec<u8>> {
sanitize_file_with_options(filename, &FileSanitizationOptions::default())
}
pub fn sanitize_file_with_options(
filename: impl AsRef<[u8]>,
options: &FileSanitizationOptions<'_>,
) -> SafeNameResult<Vec<u8>> {
let filename = filename.as_ref();
if filename.is_empty() {
return Ok(options.replacement.to_vec());
}
if filename == DOT || filename == DOTDOT {
return Ok(filename.to_vec());
}
let mut result = Vec::with_capacity(filename.len());
let len = filename.len();
let last_idx = len - 1;
for (idx, &byte) in filename.iter().enumerate() {
let byte_index = byte as usize;
let is_blocked = BLOCKED_ALWAYS[byte_index]
|| (idx == 0 && BLOCKED_INITIAL[byte_index])
|| (idx == last_idx && BLOCKED_FINAL[byte_index]);
#[cfg(feature = "require-ascii")]
let is_blocked = is_blocked || byte >= 0x80;
if is_blocked {
result.extend_from_slice(options.replacement);
} else {
result.push(byte);
}
}
#[cfg(feature = "require-utf8")]
{
if std::str::from_utf8(&result).is_err() {
result = replace_invalid_utf8(&result, options.replacement);
}
}
if result.len() > options.max_len {
return Err(SafeNameError::InvalidLength {
len: result.len(),
max: options.max_len,
});
}
Ok(result)
}
#[cfg(feature = "require-utf8")]
fn replace_invalid_utf8(bytes: &[u8], replacement: &[u8]) -> Vec<u8> {
let mut result = Vec::with_capacity(bytes.len());
let mut idx = 0;
while idx < bytes.len() {
let byte = bytes[idx];
let seq_len = match byte {
0x00..=0x7F => 1,
0xC2..=0xDF => 2,
0xE0..=0xEF => 3,
0xF0..=0xF4 => 4,
_ => {
result.extend_from_slice(replacement);
idx += 1;
continue;
}
};
if idx + seq_len > bytes.len() {
result.extend_from_slice(replacement);
idx += 1;
continue;
}
let seq = &bytes[idx..idx + seq_len];
if std::str::from_utf8(seq).is_ok() {
result.extend_from_slice(seq);
idx += seq_len;
} else {
result.extend_from_slice(replacement);
idx += 1;
}
}
result
}
#[cfg(test)]
mod tests {
use super::*;
use crate::name::validation::validate_file;
use proptest::prelude::*;
#[test]
#[cfg(feature = "require-utf8")]
fn test_sanitize_invalid_utf8_uses_custom_replacement() {
let result = sanitize_file(b"file\x80name").unwrap();
assert_eq!(result, b"file_name");
let opts = FileSanitizationOptions {
replacement: b"X",
..Default::default()
};
let result = sanitize_file_with_options(b"file\x80name", &opts).unwrap();
assert_eq!(result, b"fileXname");
}
#[test]
#[cfg(feature = "require-ascii")]
fn test_sanitize_non_ascii() {
let result = sanitize_file(b"file\x80name").unwrap();
assert_eq!(result, b"file_name");
let result = sanitize_file("fileαname".as_bytes()).unwrap();
assert_eq!(result, b"file__name");
let opts = FileSanitizationOptions {
replacement: b"X",
..Default::default()
};
let result = sanitize_file_with_options(b"file\x80name", &opts).unwrap();
assert_eq!(result, b"fileXname");
}
#[test]
fn test_sanitize_empty() {
assert_eq!(sanitize_file(b"").unwrap(), b"_");
}
#[test]
fn test_sanitize_valid_unchanged() {
assert_eq!(sanitize_file(b"normal.txt").unwrap(), b"normal.txt");
assert_eq!(sanitize_file(b"file-name").unwrap(), b"file-name");
assert_eq!(sanitize_file(b".hidden").unwrap(), b".hidden");
}
#[test]
fn test_sanitize_replaces_slash() {
assert_eq!(sanitize_file(b"/hello/a.txt").unwrap(), b"_hello_a.txt");
assert_eq!(sanitize_file(b"foo/bar").unwrap(), b"foo_bar");
}
#[test]
fn test_sanitize_dot_dotdot() {
assert_eq!(sanitize_file(b".").unwrap(), b".");
assert_eq!(sanitize_file(b"..").unwrap(), b"..");
}
#[test]
fn test_sanitize_leading_dash() {
assert_eq!(sanitize_file(b"-rf").unwrap(), b"_rf");
assert!(validate_file(sanitize_file(b"-rf").unwrap()).is_ok());
}
#[test]
fn test_sanitize_leading_tilde() {
assert_eq!(sanitize_file(b"~user").unwrap(), b"_user");
assert!(validate_file(sanitize_file(b"~user").unwrap()).is_ok());
}
#[test]
fn test_sanitize_leading_space() {
assert_eq!(sanitize_file(b" file").unwrap(), b"_file");
assert!(validate_file(sanitize_file(b" file").unwrap()).is_ok());
}
#[test]
fn test_sanitize_trailing_space() {
assert_eq!(sanitize_file(b"file ").unwrap(), b"file_");
assert!(validate_file(sanitize_file(b"file ").unwrap()).is_ok());
}
#[test]
fn test_sanitize_control_char() {
assert_eq!(sanitize_file(b"file\x00name").unwrap(), b"file_name");
assert!(validate_file(sanitize_file(b"file\x00name").unwrap()).is_ok());
}
#[test]
fn test_sanitize_with_replacement_char() {
let opts = FileSanitizationOptions::with_unicode_replacement_char();
assert_eq!(
sanitize_file_with_options(b"-rf", &opts).unwrap(),
b"\xEF\xBF\xBDrf"
);
assert_eq!(
sanitize_file_with_options(b"file ", &opts).unwrap(),
b"file\xEF\xBF\xBD"
);
}
#[test]
fn test_sanitize_too_long_returns_error() {
let long_name = [b'a'; 300];
let result = sanitize_file(long_name);
assert_eq!(
result,
Err(SafeNameError::InvalidLength { len: 300, max: 255 })
);
}
#[test]
fn test_sanitize_result_validates() {
let test_cases: &[&[u8]] = &[
b"",
b"-rf",
b"~user",
b" leading",
b"trailing ",
b"file\x00name",
b"\x7Fdel",
b"file\xFFbad",
];
for &input in test_cases {
if let Ok(result) = sanitize_file(input) {
assert!(
validate_file(&result).is_ok(),
"sanitize_file({:?}) = {:?} failed validation",
input,
result
);
}
}
}
#[test]
#[cfg(feature = "require-utf8")]
fn test_sanitize_multibyte_utf8_sequences() {
let result = sanitize_file("señor".as_bytes()).unwrap();
assert_eq!(result, "señor".as_bytes());
let result = sanitize_file("100€".as_bytes()).unwrap();
assert_eq!(result, "100€".as_bytes());
let result = sanitize_file("hi😀".as_bytes()).unwrap();
assert_eq!(result, "hi😀".as_bytes());
}
#[test]
#[cfg(feature = "require-utf8")]
fn test_sanitize_truncated_utf8_sequences() {
let result = sanitize_file(b"file\xC3").unwrap();
assert_eq!(result, b"file_");
let result = sanitize_file(b"file\xE2\x82").unwrap();
assert_eq!(result, b"file__");
let result = sanitize_file(b"file\xF0\x9F\x98").unwrap();
assert_eq!(result, b"file___");
}
#[test]
#[cfg(feature = "require-utf8")]
fn test_sanitize_invalid_continuation_bytes() {
let result = sanitize_file(b"file\xC3\x00name").unwrap();
assert_eq!(result, b"file__name");
let result = sanitize_file(b"a\xE2\x82\x00b").unwrap();
assert_eq!(result, b"a___b");
let result = sanitize_file(b"a\xF0\x9F\x98Xb").unwrap();
assert_eq!(result, b"a___Xb");
}
#[test]
#[cfg(feature = "require-utf8")]
fn test_sanitize_overlong_utf8_encodings() {
let result = sanitize_file(b"file\xC0\xAFname").unwrap();
assert_eq!(result, b"file__name");
let result = sanitize_file(b"file\xC1\x80name").unwrap();
assert_eq!(result, b"file__name");
}
proptest! {
#[test]
fn prop_sanitized_output_always_valid(input in prop::collection::vec(any::<u8>(), 0..256)) {
if let Ok(sanitized) = sanitize_file(&input) {
prop_assert!(
validate_file(&sanitized).is_ok(),
"sanitize_file({:?}) = {:?} failed validation",
input, sanitized
);
}
}
#[test]
fn prop_valid_filenames_unchanged(input in "[a-zA-Z0-9._][a-zA-Z0-9._-]{0,50}") {
let bytes = input.as_bytes();
if validate_file(bytes).is_ok() {
let sanitized = sanitize_file(bytes).unwrap();
prop_assert_eq!(
sanitized.as_slice(), bytes,
"Valid filename was modified by sanitization"
);
}
}
}
}