use crate::{CliError, CliErrorKind};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ArgLimits {
pub max_words: usize,
pub max_word_bytes: usize,
pub max_total_bytes: usize,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RawArgs {
words: Vec<String>,
}
impl Default for ArgLimits {
fn default() -> Self {
Self {
max_words: 1024,
max_word_bytes: 64 * 1024,
max_total_bytes: 1024 * 1024,
}
}
}
impl RawArgs {
pub fn from_env() -> Result<Self, CliError> {
Self::from_os_words(std::env::args_os().skip(1), ArgLimits::default())
}
fn from_os_words(
args: impl IntoIterator<Item = std::ffi::OsString>,
limits: ArgLimits,
) -> Result<Self, CliError> {
let mut words = Vec::new();
let mut total_bytes = 0usize;
for (index, word) in args.into_iter().enumerate() {
let word = word.into_string().map_err(|_| {
CliError::new(
CliErrorKind::InvalidInput,
format!("argument {} is not valid UTF-8", index + 1),
)
})?;
validate_raw_word(&word, words.len(), total_bytes, limits)?;
total_bytes = total_bytes.saturating_add(word.len());
words.push(word);
}
Ok(Self { words })
}
pub fn parse<I>(args: I) -> Result<Self, CliError>
where
I: IntoIterator,
I::Item: Into<String>,
{
Self::parse_with_limits(args, ArgLimits::default())
}
pub fn parse_with_limits<I>(args: I, limits: ArgLimits) -> Result<Self, CliError>
where
I: IntoIterator,
I::Item: Into<String>,
{
let mut words = Vec::new();
let mut total_bytes = 0usize;
for arg in args {
let value = arg.into();
validate_raw_word(&value, words.len(), total_bytes, limits)?;
total_bytes = total_bytes.saturating_add(value.len());
words.push(value);
}
Ok(Self { words })
}
pub fn words(&self) -> &[String] {
&self.words
}
}
fn validate_raw_word(
value: &str,
count: usize,
total: usize,
limits: ArgLimits,
) -> Result<(), CliError> {
if count >= limits.max_words {
return Err(CliError::new(
CliErrorKind::InvalidInput,
format!("argument count exceeds limit of {}", limits.max_words),
));
}
if value.len() > limits.max_word_bytes {
return Err(CliError::new(
CliErrorKind::InvalidInput,
format!(
"argument {} exceeds byte limit of {}",
count + 1,
limits.max_word_bytes
),
));
}
if total.saturating_add(value.len()) > limits.max_total_bytes {
return Err(CliError::new(
CliErrorKind::InvalidInput,
format!(
"total argument bytes exceed limit of {}",
limits.max_total_bytes
),
));
}
if value.contains('\0') {
return Err(CliError::new(
CliErrorKind::InvalidInput,
format!("argument {} contains a NUL byte", count + 1),
));
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use std::ffi::OsString;
#[test]
fn environment_ingestion_stops_consuming_at_first_limit_failure() {
let mut consumed = 0;
let args = ["first", "second", "never-read"].into_iter().map(|word| {
consumed += 1;
OsString::from(word)
});
let limits = ArgLimits {
max_words: 1,
..ArgLimits::default()
};
assert!(RawArgs::from_os_words(args, limits).is_err());
assert_eq!(consumed, 2);
}
#[test]
fn environment_ingestion_enforces_exact_aggregate_utf8_bytes() {
let limits = ArgLimits {
max_words: 2,
max_word_bytes: 2,
max_total_bytes: 3,
};
let accepted =
RawArgs::from_os_words([OsString::from("é"), OsString::from("a")], limits).unwrap();
assert_eq!(accepted.words(), &["é", "a"]);
assert!(
RawArgs::from_os_words([OsString::from("é"), OsString::from("ab")], limits).is_err()
);
}
#[cfg(unix)]
#[test]
fn invalid_environment_encoding_keeps_the_argument_index() {
use std::os::unix::ffi::OsStringExt;
let error = RawArgs::from_os_words(
[OsString::from("valid"), OsString::from_vec(vec![0xff])],
ArgLimits::default(),
)
.unwrap_err();
assert_eq!(error.to_string(), "argument 2 is not valid UTF-8");
}
}