use std::collections::HashMap;
pub const ENV_INFRA_VARS: &[&str] = &[
"PATH",
"HOME",
"USER",
"LOGNAME",
"SHELL",
"LANG",
"LC_ALL",
"LC_CTYPE",
"LC_MESSAGES",
"TERM",
"TMPDIR",
"TZ",
"PWD",
"SystemRoot",
"TEMP",
"TMP",
"USERPROFILE",
"USERNAME",
"PATHEXT",
"COMSPEC",
"APPDATA",
"LOCALAPPDATA",
"PROGRAMDATA",
"NUMBER_OF_PROCESSORS",
"PROCESSOR_ARCHITECTURE",
"OS",
];
pub fn build_base_env(parent: &HashMap<String, String>) -> HashMap<String, String> {
let mut out = HashMap::new();
for name in ENV_INFRA_VARS {
if let Some(val) = parent.get(*name) {
out.insert((*name).to_string(), val.clone());
}
}
out
}
pub fn expand_env_ref(value: &str, env: &HashMap<String, String>) -> String {
fn noop(_: &str) {}
expand_env_ref_with_allowlist(value, env, None, noop)
}
pub fn expand_env_ref_with_allowlist<F: FnMut(&str)>(
value: &str,
env: &HashMap<String, String>,
allowlist: Option<&[String]>,
mut on_blocked: F,
) -> String {
let bytes = value.as_bytes();
let mut out = String::with_capacity(value.len());
let mut i = 0;
while i < bytes.len() {
if bytes[i] == b'$' && i + 1 < bytes.len() && bytes[i + 1] == b'{' {
if let Some(end) = bytes[i + 2..].iter().position(|&b| b == b'}') {
let name = &value[i + 2..i + 2 + end];
if is_valid_env_ident(name) {
let allowed = allowlist.map(|a| a.iter().any(|n| n == name)).unwrap_or(true);
if allowed {
if let Some(v) = env.get(name) {
out.push_str(v);
}
} else {
on_blocked(&format!(
"env_allowlist blocked ${{{name}}}; expanded to empty"
));
}
i = i + 2 + end + 1;
continue;
}
}
}
out.push(bytes[i] as char);
i += 1;
}
out
}
fn is_valid_env_ident(name: &str) -> bool {
if name.is_empty() {
return false;
}
let mut chars = name.chars();
let first = chars.next().unwrap();
if !(first.is_ascii_alphabetic() || first == '_') {
return false;
}
chars.all(|c| c.is_ascii_alphanumeric() || c == '_')
}
fn validate_safe_placeholder_value(template: &str, field: &str, value: &str) -> Result<(), PlaceholderError> {
let token = match field {
"message" => "{{MESSAGE}}",
"session_id" => "{{SESSION_ID}}",
"session_name" => "{{SESSION_NAME}}",
_ => return Ok(()),
};
if !template.contains(token) {
return Ok(());
}
for b in value.bytes() {
if b == 0 || b == b'\n' || b == b'\r' {
return Err(PlaceholderError::UnsafeValue { field: field.to_string() });
}
}
Ok(())
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum PlaceholderError {
#[error("placeholder `{field}` value contains an unsafe byte (NUL/newline/CR) — refusing to inject into command/args")]
UnsafeValue { field: String },
}
pub fn substitute(value: &str, ctx: &SubstCtx) -> Result<String, PlaceholderError> {
validate_safe_placeholder_value(value, "message", &ctx.message)?;
validate_safe_placeholder_value(value, "session_id", &ctx.session_id)?;
validate_safe_placeholder_value(value, "session_name", &ctx.session_name)?;
Ok(value
.replace("{{MESSAGE}}", &ctx.message)
.replace("{{SESSION_ID}}", &ctx.session_id)
.replace("{{SESSION_NAME}}", &ctx.session_name)
.replace("{{PROFILE_DIR}}", &ctx.profile_dir))
}
#[derive(Debug, Clone, Default)]
pub struct SubstCtx {
pub message: String,
pub session_id: String,
pub session_name: String,
pub profile_dir: String,
}
#[cfg(test)]
mod tests {
use super::*;
fn env_of(pairs: &[(&str, &str)]) -> HashMap<String, String> {
pairs.iter().map(|(k, v)| (k.to_string(), v.to_string())).collect()
}
#[test]
fn expand_known_var() {
let env = env_of(&[("FOO", "bar")]);
assert_eq!(expand_env_ref("x${FOO}y", &env), "xbary");
}
#[test]
fn expand_unknown_to_empty() {
let env = HashMap::new();
assert_eq!(expand_env_ref("x${NOPE}y", &env), "xy");
}
#[test]
fn expand_with_allowlist_blocks() {
let env = env_of(&[("ALLOWED", "yes"), ("SECRET", "shh")]);
let allow = vec!["ALLOWED".to_string()];
let mut warnings = Vec::new();
let out = expand_env_ref_with_allowlist(
"${ALLOWED}-${SECRET}",
&env,
Some(&allow),
|m| warnings.push(m.to_string()),
);
assert_eq!(out, "yes-");
assert_eq!(warnings.len(), 1);
assert!(warnings[0].contains("SECRET"));
}
#[test]
fn expand_invalid_ident_left_untouched() {
let env = HashMap::new();
assert_eq!(expand_env_ref("${1BAD}", &env), "${1BAD}");
}
#[test]
fn build_base_env_copies_only_infra() {
let parent = env_of(&[
("PATH", "/bin"),
("HOME", "/u"),
("SECRET", "shh"),
]);
let base = build_base_env(&parent);
assert_eq!(base.get("PATH").map(|s| s.as_str()), Some("/bin"));
assert_eq!(base.get("HOME").map(|s| s.as_str()), Some("/u"));
assert!(!base.contains_key("SECRET"));
}
#[test]
fn substitute_placeholders() {
let ctx = SubstCtx {
message: "hi".into(),
session_id: "s1".into(),
session_name: "feat".into(),
profile_dir: "/p".into(),
};
assert_eq!(
substitute("{{MESSAGE}}|{{SESSION_ID}}|{{SESSION_NAME}}|{{PROFILE_DIR}}", &ctx).unwrap(),
"hi|s1|feat|/p"
);
}
#[test]
fn substitute_rejects_unsafe_message_newline() {
let ctx = SubstCtx {
message: "hello\nworld".into(),
..Default::default()
};
let err = substitute("echo {{MESSAGE}}", &ctx).unwrap_err();
assert!(matches!(err, PlaceholderError::UnsafeValue { field } if field == "message"));
}
#[test]
fn substitute_rejects_unsafe_message_nul_and_cr() {
let ctx_nul = SubstCtx { message: "a\0b".into(), ..Default::default() };
assert!(substitute("{{MESSAGE}}", &ctx_nul).is_err());
let ctx_cr = SubstCtx { message: "a\rb".into(), ..Default::default() };
assert!(substitute("{{MESSAGE}}", &ctx_cr).is_err());
}
#[test]
fn substitute_allows_newlines_when_no_placeholder() {
let ctx = SubstCtx { message: "multi\nline".into(), ..Default::default() };
assert_eq!(substitute("no placeholders here", &ctx).unwrap(), "no placeholders here");
}
#[test]
fn substitute_rejects_unsafe_session_fields() {
let ctx = SubstCtx {
session_id: "s\n1".into(),
session_name: "feat\rname".into(),
..Default::default()
};
assert!(substitute("--session {{SESSION_ID}}", &ctx).is_err());
assert!(substitute("--name {{SESSION_NAME}}", &ctx).is_err());
}
#[test]
fn substitute_profile_dir_not_validated() {
let ctx = SubstCtx { profile_dir: "/a\nb".into(), ..Default::default() };
assert_eq!(substitute("{{PROFILE_DIR}}/script", &ctx).unwrap(), "/a\nb/script");
}
}