#[cfg(test)]
mod tests;
use std::fmt::Write as _;
use std::path::Path;
use globset::GlobBuilder;
use regex::bytes::{Regex as BytesRegex, RegexBuilder as BytesRegexBuilder};
use regex::{Regex, RegexBuilder};
use regex_syntax::ast::parse::Parser;
use regex_syntax::ast::{self, Ast};
use serde::de::Error as _;
pub(super) const MAX_SHELL_ALLOWLIST_RULES: usize = 32;
pub(super) const MAX_SHELL_ALLOWLIST_PATTERN_BYTES: usize = 2 * 1024;
pub(super) const MAX_SHELL_ALLOWLIST_DESCRIPTION_BYTES: usize = 1024;
pub(super) const MAX_SHELL_ALLOWLIST_COMPILE_BYTES: usize = 256 * 1024;
pub(super) fn deserialize_shell_allowlist<'de, D>(
deserializer: D,
) -> Result<Option<Vec<ShellAllowRule>>, D::Error>
where
D: serde::Deserializer<'de>,
{
struct AllowlistVisitor;
impl<'de> serde::de::Visitor<'de> for AllowlistVisitor {
type Value = Vec<ShellAllowRule>;
fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("a sequence of shell allowlist rules")
}
fn visit_seq<A>(self, mut sequence: A) -> Result<Self::Value, A::Error>
where
A: serde::de::SeqAccess<'de>,
{
if sequence
.size_hint()
.is_some_and(|count| MAX_SHELL_ALLOWLIST_RULES < count)
{
return Err(A::Error::custom(format!(
"shell allowlist permits at most {MAX_SHELL_ALLOWLIST_RULES} rules"
)));
}
let mut rules = Vec::with_capacity(MAX_SHELL_ALLOWLIST_RULES);
while let Some(raw_rule) = sequence.next_element::<RawShellAllowRule>()? {
if MAX_SHELL_ALLOWLIST_RULES <= rules.len() {
return Err(A::Error::custom(format!(
"shell allowlist permits at most {MAX_SHELL_ALLOWLIST_RULES} rules"
)));
}
rules.push(ShellAllowRule::try_from(raw_rule).map_err(A::Error::custom)?);
}
Ok(rules)
}
}
deserializer.deserialize_seq(AllowlistVisitor).map(Some)
}
#[derive(Clone, Debug)]
pub(super) struct ShellAllowRule {
workdir: String,
workdir_matcher: BytesRegex,
command_matcher: ShellCommandMatcher,
description: Option<String>,
}
impl ShellAllowRule {
pub(super) fn matches(&self, canonical_cwd: &str, command: &str) -> bool {
self.workdir_matcher.is_match(canonical_cwd.as_bytes())
&& self.command_matcher.is_match(command)
}
pub(super) fn append_diagnostic(&self, message: &mut String) {
let command = serde_json::to_string(self.command_matcher.pattern())
.expect("serializing a string to JSON cannot fail");
let workdir =
serde_json::to_string(&self.workdir).expect("serializing a string to JSON cannot fail");
let _ = write!(
message,
"\n- {}: {command}\n workdir: {workdir}",
self.command_matcher.field_name()
);
if let Some(description) = &self.description {
let description = serde_json::to_string(description)
.expect("serializing a string to JSON cannot fail");
let _ = write!(message, "\n description: {description}");
}
}
pub(super) fn prompt_selector(&self) -> String {
let command = prompt_json_string(self.command_matcher.pattern());
let workdir = prompt_json_string(&self.workdir);
let mut selector = format!(
"{}: {command}; workdir: {workdir}",
self.command_matcher.field_name()
);
if let Some(description) = &self.description {
let description = prompt_json_string(description);
let _ = write!(selector, "; description: {description}");
}
selector
}
}
fn prompt_json_string(value: &str) -> String {
serde_json::to_string(value)
.expect("serializing a string to JSON cannot fail")
.replace('{', r"\u007b")
.replace('}', r"\u007d")
}
#[derive(serde::Deserialize)]
#[serde(deny_unknown_fields)]
struct RawShellAllowRule {
workdir: String,
command: Option<String>,
command_regex: Option<String>,
description: Option<String>,
}
#[derive(Clone, Debug)]
enum ShellCommandMatcher {
Glob {
pattern: String,
matcher: BytesRegex,
},
Regex {
pattern: String,
matcher: Regex,
},
}
impl ShellCommandMatcher {
fn is_match(&self, command: &str) -> bool {
match self {
Self::Glob { matcher, .. } => matcher.is_match(command.as_bytes()),
Self::Regex { matcher, .. } => matcher.is_match(command),
}
}
fn field_name(&self) -> &'static str {
match self {
Self::Glob { .. } => "command_glob",
Self::Regex { .. } => "command_regex",
}
}
fn pattern(&self) -> &str {
match self {
Self::Glob { pattern, .. } | Self::Regex { pattern, .. } => pattern,
}
}
}
impl TryFrom<RawShellAllowRule> for ShellAllowRule {
type Error = String;
fn try_from(raw: RawShellAllowRule) -> Result<Self, Self::Error> {
if raw
.description
.as_ref()
.is_some_and(|description| MAX_SHELL_ALLOWLIST_DESCRIPTION_BYTES < description.len())
{
return Err(format!(
"shell allowlist description must not exceed {MAX_SHELL_ALLOWLIST_DESCRIPTION_BYTES} authored UTF-8 bytes"
));
}
if !Path::new(&raw.workdir).is_absolute() {
return Err("shell allowlist workdir glob must be absolute".to_owned());
}
require_pattern_limit("workdir", &raw.workdir)?;
let workdir_matcher = compile_workdir_glob(&raw.workdir)?;
let command_matcher = match (raw.command, raw.command_regex) {
(Some(command), None) => {
require_pattern_limit("command", &command)?;
ShellCommandMatcher::Glob {
matcher: compile_command_glob(&command)?,
pattern: command,
}
}
(None, Some(command_regex)) => {
require_pattern_limit("command_regex", &command_regex)?;
ShellCommandMatcher::Regex {
matcher: compile_command_regex(&command_regex)?,
pattern: command_regex,
}
}
(None, None) | (Some(_), Some(_)) => {
return Err(
"shell allowlist rule requires exactly one of `command` or `command_regex`"
.to_owned(),
);
}
};
Ok(Self {
workdir: raw.workdir,
workdir_matcher,
command_matcher,
description: raw.description,
})
}
}
fn require_pattern_limit(field: &str, pattern: &str) -> Result<(), String> {
if MAX_SHELL_ALLOWLIST_PATTERN_BYTES < pattern.len() {
return Err(format!(
"shell allowlist `{field}` must not exceed {MAX_SHELL_ALLOWLIST_PATTERN_BYTES} authored UTF-8 bytes"
));
}
Ok(())
}
fn compile_workdir_glob(pattern: &str) -> Result<BytesRegex, String> {
compile_glob(pattern, GlobMatcherRole::Workdir)
}
fn compile_command_glob(pattern: &str) -> Result<BytesRegex, String> {
compile_glob(pattern, GlobMatcherRole::Command)
}
enum GlobMatcherRole {
Workdir,
Command,
}
impl GlobMatcherRole {
fn field_name(&self) -> &'static str {
match self {
Self::Workdir => "workdir",
Self::Command => "command",
}
}
fn literal_separator(&self) -> bool {
match self {
Self::Workdir => true,
Self::Command => false,
}
}
}
fn compile_glob(pattern: &str, role: GlobMatcherRole) -> Result<BytesRegex, String> {
let glob = GlobBuilder::new(pattern)
.literal_separator(role.literal_separator())
.backslash_escape(true)
.build()
.map_err(|_| format!("invalid shell allowlist {} glob", role.field_name()))?;
BytesRegexBuilder::new(glob.regex())
.dot_matches_new_line(true)
.size_limit(MAX_SHELL_ALLOWLIST_COMPILE_BYTES)
.build()
.map_err(|error| regex_compile_error(role.field_name(), "glob", error))
}
fn compile_command_regex(pattern: &str) -> Result<Regex, String> {
let ast = Parser::new()
.parse(pattern)
.map_err(|_| "invalid shell allowlist command regex".to_owned())?;
if ast_enables_case_insensitive_matching(&ast) {
return Err("shell allowlist command regex must remain case-sensitive".to_owned());
}
let wrapped = format!(r"\A(?:{pattern})\z");
RegexBuilder::new(&wrapped)
.size_limit(MAX_SHELL_ALLOWLIST_COMPILE_BYTES)
.build()
.map_err(|error| regex_compile_error("command", "regex", error))
}
fn regex_compile_error(field: &str, matcher_type: &str, error: regex::Error) -> String {
match error {
regex::Error::CompiledTooBig(_) => format!(
"shell allowlist {field} {matcher_type} compilation must not exceed {MAX_SHELL_ALLOWLIST_COMPILE_BYTES} bytes"
),
_ => format!("invalid shell allowlist {field} {matcher_type}"),
}
}
fn ast_enables_case_insensitive_matching(ast: &Ast) -> bool {
let mut pending = vec![ast];
while let Some(ast) = pending.pop() {
match ast {
Ast::Flags(flags)
if flags.flags.flag_state(ast::Flag::CaseInsensitive) == Some(true) =>
{
return true;
}
Ast::Group(group) => {
if group
.flags()
.is_some_and(|flags| flags.flag_state(ast::Flag::CaseInsensitive) == Some(true))
{
return true;
}
pending.push(&group.ast);
}
Ast::Repetition(repetition) => pending.push(&repetition.ast),
Ast::Alternation(alternation) => pending.extend(&alternation.asts),
Ast::Concat(concat) => pending.extend(&concat.asts),
Ast::Empty(_)
| Ast::Literal(_)
| Ast::Dot(_)
| Ast::Assertion(_)
| Ast::ClassUnicode(_)
| Ast::ClassPerl(_)
| Ast::ClassBracketed(_)
| Ast::Flags(_) => {}
}
}
false
}