use std::fmt;
pub const CMD_SWITCHES: &str = "/d /e:ON /v:OFF /s /c";
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum CmdLineError {
#[error(
"argument {index} cannot be passed to a .cmd or .bat file: it contains a line break, \
which terminates a cmd.exe command line"
)]
LineBreak {
index: usize,
},
#[error("argument {index} contains a NUL byte")]
Nul {
index: usize,
},
#[error("the script path contains a quote or ends with a backslash, so it cannot be quoted")]
ScriptPath,
}
mod unit {
pub(super) const QUOTE: u16 = b'"' as u16;
pub(super) const BACKSLASH: u16 = b'\\' as u16;
pub(super) const PERCENT: u16 = b'%' as u16;
pub(super) const CR: u16 = b'\r' as u16;
pub(super) const LF: u16 = b'\n' as u16;
pub(super) const NUL: u16 = 0;
pub(super) const SPACE: u16 = b' ' as u16;
}
const UNQUOTED_PUNCTUATION: &str = r"#$*+-./:?@\_";
const PERCENT_GUARD: &str = "%%cd:~,";
fn needs_quoting(arg: &[u16]) -> bool {
if arg.is_empty() {
return true;
}
if arg.last() == Some(&unit::BACKSLASH) {
return true;
}
arg.iter().any(|&unit| {
let Some(ch) = u8::try_from(unit).ok().filter(u8::is_ascii).map(char::from) else {
return char::from_u32(u32::from(unit)).is_some_and(char::is_control);
};
!(ch.is_ascii_alphanumeric() || UNQUOTED_PUNCTUATION.contains(ch))
})
}
pub fn append_argument(out: &mut Vec<u16>, arg: &[u16], index: usize) -> Result<(), CmdLineError> {
if arg.contains(&unit::CR) || arg.contains(&unit::LF) {
return Err(CmdLineError::LineBreak { index });
}
if arg.contains(&unit::NUL) {
return Err(CmdLineError::Nul { index });
}
let quote = needs_quoting(arg);
if quote {
out.push(unit::QUOTE);
}
let mut backslashes: usize = 0;
for &code in arg {
match code {
unit::BACKSLASH => backslashes += 1,
unit::QUOTE => {
out.extend(std::iter::repeat_n(unit::BACKSLASH, backslashes));
backslashes = 0;
out.push(unit::QUOTE);
}
unit::PERCENT => {
backslashes = 0;
out.extend(PERCENT_GUARD.encode_utf16());
}
_ => backslashes = 0,
}
out.push(code);
}
if quote {
out.extend(std::iter::repeat_n(unit::BACKSLASH, backslashes));
out.push(unit::QUOTE);
}
Ok(())
}
pub fn batch_command_line(script: &[u16], args: &[Vec<u16>]) -> Result<Vec<u16>, CmdLineError> {
if script.contains(&unit::QUOTE) || script.last() == Some(&unit::BACKSLASH) {
return Err(CmdLineError::ScriptPath);
}
if script.contains(&unit::NUL) {
return Err(CmdLineError::Nul { index: 0 });
}
let mut out: Vec<u16> = CMD_SWITCHES.encode_utf16().collect();
out.push(unit::SPACE);
out.push(unit::QUOTE);
out.push(unit::QUOTE);
out.extend_from_slice(script);
out.push(unit::QUOTE);
for (offset, arg) in args.iter().enumerate() {
out.push(unit::SPACE);
append_argument(&mut out, arg, offset + 1)?;
}
out.push(unit::QUOTE);
Ok(out)
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Utf16Display<'a>(pub &'a [u16]);
impl fmt::Display for Utf16Display<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
for ch in char::decode_utf16(self.0.iter().copied()) {
f.write_fmt(format_args!("{}", ch.unwrap_or(char::REPLACEMENT_CHARACTER)))?;
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
fn units(text: &str) -> Vec<u16> {
text.encode_utf16().collect()
}
fn escape(arg: &str) -> String {
let mut out = Vec::new();
append_argument(&mut out, &units(arg), 1).expect("representable");
Utf16Display(&out).to_string()
}
fn line(script: &str, args: &[&str]) -> String {
let args: Vec<Vec<u16>> = args.iter().map(|a| units(a)).collect();
let out = batch_command_line(&units(script), &args).expect("representable");
Utf16Display(&out).to_string()
}
#[test]
fn the_switches_disable_every_feature_that_could_run_something_else() {
assert!(CMD_SWITCHES.contains("/d"), "AutoRun must be disabled");
assert!(CMD_SWITCHES.contains("/v:OFF"), "delayed expansion must be disabled");
assert!(CMD_SWITCHES.contains("/s"), "the conditional quote rule must be pinned");
assert!(CMD_SWITCHES.contains("/e:ON"), "the percent guard needs command extensions");
}
#[test]
fn a_plain_word_is_left_alone() {
assert_eq!(escape("test"), "test");
assert_eq!(escape("build2"), "build2");
assert_eq!(escape("--flag"), "--flag");
assert_eq!(escape(r"C:\dir\file.txt"), r"C:\dir\file.txt");
}
#[test]
fn whitespace_forces_quoting() {
assert_eq!(escape("a b"), r#""a b""#);
assert_eq!(escape("a\tb"), "\"a\tb\"");
}
#[test]
fn an_empty_argument_is_quoted_so_it_survives() {
assert_eq!(escape(""), r#""""#);
}
#[test]
fn command_separators_are_quoted_rather_than_executed() {
assert_eq!(escape("a&b"), r#""a&b""#);
assert_eq!(escape("a|b"), r#""a|b""#);
assert_eq!(escape("a>b"), r#""a>b""#);
assert_eq!(escape("a<b"), r#""a<b""#);
assert_eq!(escape("a&&b"), r#""a&&b""#);
}
#[test]
fn a_caret_is_quoted_rather_than_treated_as_an_escape() {
assert_eq!(escape("a^b"), r#""a^b""#);
assert!(!escape("a^b").contains("^^"));
}
#[test]
fn a_variable_reference_is_defused_rather_than_expanded() {
let escaped = escape("%PATH%");
assert_eq!(escaped, "\"%%cd:~,%PATH%%cd:~,%\"");
assert_eq!(escaped.matches("%%cd:~,").count(), 2);
}
#[test]
fn a_lone_percent_is_guarded_too() {
assert_eq!(escape("100%"), "\"100%%cd:~,%\"");
assert_eq!(escape("a%b"), "\"a%%cd:~,%b\"");
}
#[test]
fn delayed_expansion_syntax_is_quoted_and_the_switch_disables_it() {
assert_eq!(escape("!DELAYED!"), r#""!DELAYED!""#);
assert!(CMD_SWITCHES.contains("/v:OFF"));
}
#[test]
fn an_inner_quote_is_doubled_not_backslash_escaped() {
assert_eq!(escape(r#"a"b"#), r#""a""b""#);
assert!(!escape(r#"a"b"#).contains(r#"\""#));
}
#[test]
fn backslashes_before_a_quote_are_doubled() {
assert_eq!(escape(r#"a\"b"#), r#""a\\""b""#);
assert_eq!(escape(r#"a\\"b"#), r#""a\\\\""b""#);
}
#[test]
fn a_trailing_backslash_is_doubled_against_the_closing_quote() {
assert_eq!(escape(r"a\"), r#""a\\""#);
assert_eq!(escape(r"C:\dir\"), r#""C:\dir\\""#);
}
#[test]
fn backslashes_not_adjacent_to_a_quote_are_left_alone() {
assert_eq!(escape(r"C:\a\b"), r"C:\a\b");
}
#[test]
fn a_line_break_is_refused_rather_than_truncated() {
let mut out = Vec::new();
assert_eq!(
append_argument(&mut out, &units("a\nb"), 3),
Err(CmdLineError::LineBreak { index: 3 })
);
assert_eq!(
append_argument(&mut out, &units("a\rb"), 1),
Err(CmdLineError::LineBreak { index: 1 })
);
}
#[test]
fn a_nul_is_refused() {
let mut out = Vec::new();
assert_eq!(
append_argument(&mut out, &[0x61, 0, 0x62], 2),
Err(CmdLineError::Nul { index: 2 })
);
}
#[test]
fn the_error_names_the_argument_position() {
let args = vec![units("ok"), units("bad\nvalue")];
let err = batch_command_line(&units(r"C:\n\npm.cmd"), &args).unwrap_err();
assert_eq!(err, CmdLineError::LineBreak { index: 2 });
assert!(err.to_string().contains("argument 2"));
}
#[test]
fn the_whole_command_is_wrapped_in_one_outer_quote_pair() {
let built = line(r"C:\Program Files\nodejs\npm.cmd", &["test"]);
assert_eq!(built, "/d /e:ON /v:OFF /s /c \"\"C:\\Program Files\\nodejs\\npm.cmd\" test\"");
assert!(built.ends_with('"'));
}
#[test]
fn an_unquotable_script_path_is_refused() {
assert_eq!(
batch_command_line(&units(r#"C:\we"ird.cmd"#), &[]),
Err(CmdLineError::ScriptPath)
);
assert_eq!(batch_command_line(&units(r"C:\dir\"), &[]), Err(CmdLineError::ScriptPath));
}
#[test]
fn the_full_adversarial_set_produces_a_balanced_command_line() {
let adversarial =
[r#"a"b"#, "a&b", "%PATH%", "!DELAYED!", "a b", "a^b", "", "a|b", r"a\", "100%"];
let built = line(r"C:\tools\shim.cmd", &adversarial);
assert_eq!(built.matches('"').count() % 2, 0, "unbalanced quotes in {built}");
assert!(!built.contains('\n') && !built.contains('\r'));
}
#[test]
fn non_ascii_arguments_survive_unquoted() {
assert_eq!(escape("café"), "café");
assert_eq!(escape("日本"), "日本");
}
#[test]
fn utf16_display_round_trips() {
assert_eq!(Utf16Display(&units("héllo")).to_string(), "héllo");
}
}