use std::path::PathBuf;
#[derive(Debug, PartialEq, Eq)]
pub(super) enum Segment {
Literal(String),
Required { var: String },
Default { var: String, default: String },
RequiredMsg { var: String, message: String },
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum InterpolationError {
#[error(
"environment variable `{var}` is required at {path}:{key_path} \
but is unset or empty in the process environment"
)]
Required {
var: String,
path: PathBuf,
key_path: String,
},
#[error("environment variable `{var}` is required at {path}:{key_path}: {message}")]
RequiredWithMessage {
var: String,
message: String,
path: PathBuf,
key_path: String,
},
#[error(
"invalid interpolation syntax at {path}:{key_path}: {detail} \
(offending value: `{snippet}`)"
)]
Syntax {
detail: String,
snippet: String,
path: PathBuf,
key_path: String,
},
}
use nom::{
IResult, Parser,
branch::alt,
bytes::complete::{tag, take_while, take_while1},
character::complete::satisfy,
combinator::{cut, map, recognize},
multi::many0,
sequence::preceded,
};
pub(super) fn parse_template(input: &str) -> IResult<&str, Vec<Segment>> {
many0(alt((
map(tag("$$"), |_| Segment::Literal("$".to_string())),
parse_placeholder,
map(parse_literal_run, Segment::Literal),
map(tag("$"), |_| Segment::Literal("$".to_string())),
)))
.parse(input)
}
enum Modifier {
Default(String),
RequiredMsg(String),
}
fn parse_literal_run(input: &str) -> IResult<&str, String> {
let (rest, s) = take_while1(|c: char| c != '$')(input)?;
Ok((rest, s.to_string()))
}
fn parse_placeholder(input: &str) -> IResult<&str, Segment> {
let (input, _) = tag("${")(input)?;
cut(parse_placeholder_body).parse(input)
}
fn parse_placeholder_body(input: &str) -> IResult<&str, Segment> {
let (input, var_name) = parse_var_name(input)?;
let var = var_name.to_string();
let (input, modifier) = parse_modifier(input)?;
let (input, _) = tag("}")(input)?;
let segment = match modifier {
None => Segment::Required { var },
Some(Modifier::Default(default)) => Segment::Default { var, default },
Some(Modifier::RequiredMsg(message)) => Segment::RequiredMsg { var, message },
};
Ok((input, segment))
}
fn parse_var_name(input: &str) -> IResult<&str, &str> {
recognize((
satisfy(|c: char| c.is_ascii_alphabetic() || c == '_'),
take_while(|c: char| c.is_ascii_alphanumeric() || c == '_'),
))
.parse(input)
}
fn parse_modifier(input: &str) -> IResult<&str, Option<Modifier>> {
alt((
map(preceded(tag(":-"), parse_rest_until_brace), |s| {
Some(Modifier::Default(s))
}),
map(preceded(tag(":?"), parse_rest_until_brace), |s| {
Some(Modifier::RequiredMsg(s))
}),
parse_no_modifier,
))
.parse(input)
}
fn parse_no_modifier(input: &str) -> IResult<&str, Option<Modifier>> {
Ok((input, None))
}
fn parse_rest_until_brace(input: &str) -> IResult<&str, String> {
let mut out = String::new();
let mut chars = input;
loop {
if chars.is_empty() {
return Err(nom::Err::Error(nom::error::Error::new(
chars,
nom::error::ErrorKind::Eof,
)));
}
if let Some(rest) = chars.strip_prefix("$$") {
out.push('$');
chars = rest;
continue;
}
if chars.starts_with("${") {
return Err(nom::Err::Error(nom::error::Error::new(
chars,
nom::error::ErrorKind::Char,
)));
}
if chars.starts_with('}') {
return Ok((chars, out));
}
let mut iter = chars.chars();
let c = iter.next().unwrap();
out.push(c);
chars = iter.as_str();
}
}
pub(super) struct Interpolator<'a> {
lookup: &'a dyn Fn(&str) -> Option<String>,
}
impl<'a> Interpolator<'a> {
pub(super) fn new(lookup: &'a dyn Fn(&str) -> Option<String>) -> Self {
Self { lookup }
}
pub(super) fn interpolate_str(&self, input: &str) -> Result<String, InterpolationError> {
let (_rest, segments) = parse_template(input).map_err(|e| InterpolationError::Syntax {
detail: format!("{}", e),
snippet: input.to_string(),
path: PathBuf::new(),
key_path: String::new(),
})?;
let mut out = String::with_capacity(input.len());
for segment in segments {
match segment {
Segment::Literal(s) => out.push_str(&s),
Segment::Required { var } => match self.resolve(&var) {
Some(v) => out.push_str(&v),
None => {
return Err(InterpolationError::Required {
var,
path: PathBuf::new(),
key_path: String::new(),
});
}
},
Segment::Default { var, default } => match self.resolve(&var) {
Some(v) => out.push_str(&v),
None => out.push_str(&default),
},
Segment::RequiredMsg { var, message } => match self.resolve(&var) {
Some(v) => out.push_str(&v),
None => {
return Err(InterpolationError::RequiredWithMessage {
var,
message,
path: PathBuf::new(),
key_path: String::new(),
});
}
},
}
}
Ok(out)
}
fn resolve(&self, name: &str) -> Option<String> {
match (self.lookup)(name) {
Some(s) if !s.is_empty() => Some(s),
_ => None,
}
}
}
use std::path::Path;
#[derive(Debug, Clone)]
pub(super) enum KeyPathSegment {
Key(String),
Index(usize),
}
#[derive(Default, Clone)]
pub(super) struct KeyPath(Vec<KeyPathSegment>);
impl KeyPath {
pub(super) fn push_key(&mut self, key: &str) {
self.0.push(KeyPathSegment::Key(key.to_string()));
}
pub(super) fn push_index(&mut self, idx: usize) {
self.0.push(KeyPathSegment::Index(idx));
}
pub(super) fn pop(&mut self) {
self.0.pop();
}
}
impl std::fmt::Display for KeyPath {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
for (i, seg) in self.0.iter().enumerate() {
match seg {
KeyPathSegment::Key(k) if i == 0 => write!(f, "{}", k)?,
KeyPathSegment::Key(k) => write!(f, ".{}", k)?,
KeyPathSegment::Index(idx) => write!(f, "[{}]", idx)?,
}
}
Ok(())
}
}
impl<'a> Interpolator<'a> {
pub(super) fn interpolate_value(
&self,
value: &mut toml::Value,
toml_path: &Path,
) -> Result<(), InterpolationError> {
let mut key_path = KeyPath::default();
self.walk(value, toml_path, &mut key_path)
}
fn walk(
&self,
value: &mut toml::Value,
toml_path: &Path,
key_path: &mut KeyPath,
) -> Result<(), InterpolationError> {
match value {
toml::Value::String(s) => match self.interpolate_str(s) {
Ok(replaced) => {
*s = replaced;
Ok(())
}
Err(err) => Err(attach_context(err, toml_path, key_path)),
},
toml::Value::Table(table) => {
for (k, child) in table.iter_mut() {
key_path.push_key(k);
self.walk(child, toml_path, key_path)?;
key_path.pop();
}
Ok(())
}
toml::Value::Array(arr) => {
for (idx, child) in arr.iter_mut().enumerate() {
key_path.push_index(idx);
self.walk(child, toml_path, key_path)?;
key_path.pop();
}
Ok(())
}
_ => Ok(()),
}
}
}
fn attach_context(
err: InterpolationError,
toml_path: &Path,
key_path: &KeyPath,
) -> InterpolationError {
let path = toml_path.to_path_buf();
let key_path_str = key_path.to_string();
match err {
InterpolationError::Required { var, .. } => InterpolationError::Required {
var,
path,
key_path: key_path_str,
},
InterpolationError::RequiredWithMessage { var, message, .. } => {
InterpolationError::RequiredWithMessage {
var,
message,
path,
key_path: key_path_str,
}
}
InterpolationError::Syntax {
detail, snippet, ..
} => InterpolationError::Syntax {
detail,
snippet,
path,
key_path: key_path_str,
},
}
}
#[cfg(test)]
mod tests {
use super::*;
use rstest::rstest;
use std::path::PathBuf;
#[rstest]
fn interpolation_error_displays_context() {
let err = InterpolationError::Required {
var: "DB_HOST".into(),
path: PathBuf::from("local.toml"),
key_path: "database.host".into(),
};
let msg = err.to_string();
assert!(msg.contains("DB_HOST"));
assert!(msg.contains("local.toml"));
assert!(msg.contains("database.host"));
}
#[rstest]
#[case::empty("", vec![])]
#[case::literal_only("hello", vec![Segment::Literal("hello".into())])]
#[case::escape_only("$$", vec![Segment::Literal("$".into())])]
#[case::required("${VAR}", vec![Segment::Required { var: "VAR".into() }])]
#[case::default_empty(
"${VAR:-}",
vec![Segment::Default { var: "VAR".into(), default: "".into() }]
)]
#[case::default_with_value(
"${VAR:-localhost}",
vec![Segment::Default { var: "VAR".into(), default: "localhost".into() }]
)]
#[case::required_msg(
"${VAR:?Set me}",
vec![Segment::RequiredMsg { var: "VAR".into(), message: "Set me".into() }]
)]
#[case::two_placeholders(
"${A}${B}",
vec![
Segment::Required { var: "A".into() },
Segment::Required { var: "B".into() },
]
)]
#[case::placeholder_with_literal(
"a${A}-${B}b",
vec![
Segment::Literal("a".into()),
Segment::Required { var: "A".into() },
Segment::Literal("-".into()),
Segment::Required { var: "B".into() },
Segment::Literal("b".into()),
]
)]
#[case::escape_then_brace(
"$${VAR}",
vec![
Segment::Literal("$".into()),
Segment::Literal("{VAR}".into()),
]
)]
#[case::var_with_digits(
"${VAR_123}",
vec![Segment::Required { var: "VAR_123".into() }]
)]
#[case::dollar_then_letter(
"$abc",
vec![
Segment::Literal("$".into()),
Segment::Literal("abc".into()),
]
)]
#[case::escape_inside_default(
"${VAR:-foo$$bar}",
vec![Segment::Default { var: "VAR".into(), default: "foo$bar".into() }]
)]
#[case::lone_dollar_in_default(
"${VAR:-foo$bar}",
vec![Segment::Default { var: "VAR".into(), default: "foo$bar".into() }]
)]
fn parse_template_ok(#[case] input: &str, #[case] expected: Vec<Segment>) {
let result = parse_template(input);
let (rest, segments) = result.expect("parse should succeed");
assert_eq!(rest, "", "parser left unconsumed input");
assert_eq!(segments, expected);
}
#[rstest]
#[case::unclosed_brace("${VAR")]
#[case::empty_var_name("${}")]
#[case::var_starts_with_digit("${1VAR}")]
#[case::var_with_hyphen("${MY-VAR}")]
#[case::single_colon("${VAR:default}")]
#[case::dash_only("${VAR-default}")]
#[case::nested_placeholder("${A:-${B}}")]
#[case::question_only("${VAR?msg}")]
fn parse_template_err(#[case] input: &str) {
let result = parse_template(input);
assert!(
matches!(result, Err(nom::Err::Failure(_))),
"input `{}` should produce Err::Failure (cut-protected), got {:?}",
input,
result,
);
}
use std::collections::HashMap;
macro_rules! envmap {
( $( $k:expr => $v:expr ),* $(,)? ) => {{
#[allow(unused_mut)]
let mut m: HashMap<&'static str, &'static str> = HashMap::new();
$( m.insert($k, $v); )*
m
}};
}
#[rstest]
#[case::set_value(envmap!{"VAR" => "value"}, "${VAR}", "value")]
#[case::with_prefix_suffix(envmap!{"H" => "ex"}, "http://${H}/x", "http://ex/x")]
#[case::default_used(envmap!{}, "${VAR:-fallback}", "fallback")]
#[case::default_overridden(envmap!{"V" => "x"}, "${V:-fallback}", "x")]
#[case::empty_default_explicit(envmap!{}, "${VAR:-}", "")]
#[case::empty_env_treated_as_unset(envmap!{"V" => ""}, "${V:-fb}", "fb")]
#[case::escape_passthrough(envmap!{}, "$$50", "$50")]
fn interpolate_str_ok(
#[case] env: HashMap<&'static str, &'static str>,
#[case] tmpl: &str,
#[case] expected: &str,
) {
let lookup = |n: &str| env.get(n).map(|s| s.to_string());
let interp = Interpolator::new(&lookup);
let result = interp
.interpolate_str(tmpl)
.expect("interpolation should succeed");
assert_eq!(result, expected);
}
#[rstest]
fn interpolate_str_required_unset_returns_required_error() {
let env: HashMap<&'static str, &'static str> = HashMap::new();
let lookup = |n: &str| env.get(n).map(|s| s.to_string());
let interp = Interpolator::new(&lookup);
let err = interp.interpolate_str("${MISSING}").unwrap_err();
assert!(matches!(
&err,
InterpolationError::Required { var, .. } if var == "MISSING"
));
}
#[rstest]
fn interpolate_str_required_empty_returns_required_error() {
let env = envmap! {"V" => ""};
let lookup = |n: &str| env.get(n).map(|s| s.to_string());
let interp = Interpolator::new(&lookup);
let err = interp.interpolate_str("${V}").unwrap_err();
assert!(matches!(err, InterpolationError::Required { .. }));
}
#[rstest]
fn interpolate_str_required_msg_returns_message() {
let env: HashMap<&'static str, &'static str> = HashMap::new();
let lookup = |n: &str| env.get(n).map(|s| s.to_string());
let interp = Interpolator::new(&lookup);
let err = interp.interpolate_str("${P:?Set via direnv}").unwrap_err();
assert!(matches!(
&err,
InterpolationError::RequiredWithMessage { var, message, .. }
if var == "P" && message == "Set via direnv"
));
}
#[rstest]
fn interpolate_str_syntax_error_for_unclosed_brace() {
let env: HashMap<&'static str, &'static str> = HashMap::new();
let lookup = |n: &str| env.get(n).map(|s| s.to_string());
let interp = Interpolator::new(&lookup);
let err = interp.interpolate_str("${UNCLOSED").unwrap_err();
assert!(matches!(err, InterpolationError::Syntax { .. }));
}
#[rstest]
fn interpolate_value_walks_nested_table() {
let mut value: toml::Value = toml::from_str(
r#"
[database]
host = "${DB_HOST}"
port = 5432
"#,
)
.unwrap();
let env = envmap! {"DB_HOST" => "postgres"};
let lookup = |n: &str| env.get(n).map(|s| s.to_string());
let interp = Interpolator::new(&lookup);
interp
.interpolate_value(&mut value, Path::new("test.toml"))
.expect("walk should succeed");
assert_eq!(value["database"]["host"].as_str(), Some("postgres"));
assert_eq!(value["database"]["port"].as_integer(), Some(5432));
}
#[rstest]
fn interpolate_value_propagates_key_path_in_error() {
let mut value: toml::Value = toml::from_str(
r#"
[core.databases.default]
host = "${MISSING_VAR}"
"#,
)
.unwrap();
let env: HashMap<&'static str, &'static str> = HashMap::new();
let lookup = |n: &str| env.get(n).map(|s| s.to_string());
let interp = Interpolator::new(&lookup);
let err = interp
.interpolate_value(&mut value, Path::new("local.toml"))
.unwrap_err();
let msg = err.to_string();
assert!(msg.contains("local.toml"), "msg = {}", msg);
assert!(msg.contains("core.databases.default.host"), "msg = {}", msg);
assert!(msg.contains("MISSING_VAR"), "msg = {}", msg);
}
#[rstest]
fn interpolate_value_walks_array_with_index_in_path() {
let mut value: toml::Value = toml::from_str(
r#"
services = ["${SVC_A}", "${SVC_B:-default-b}"]
"#,
)
.unwrap();
let env = envmap! {"SVC_A" => "alpha"};
let lookup = |n: &str| env.get(n).map(|s| s.to_string());
let interp = Interpolator::new(&lookup);
interp
.interpolate_value(&mut value, Path::new("test.toml"))
.unwrap();
let arr = value["services"].as_array().unwrap();
assert_eq!(arr[0].as_str(), Some("alpha"));
assert_eq!(arr[1].as_str(), Some("default-b"));
}
#[rstest]
fn interpolate_value_array_error_includes_index_in_path() {
let mut value: toml::Value = toml::from_str(r#"services = ["${MISSING_SVC}"]"#).unwrap();
let env: HashMap<&'static str, &'static str> = HashMap::new();
let lookup = |n: &str| env.get(n).map(|s| s.to_string());
let interp = Interpolator::new(&lookup);
let err = interp
.interpolate_value(&mut value, Path::new("test.toml"))
.unwrap_err();
let msg = err.to_string();
assert!(msg.contains("services[0]"), "msg = {}", msg);
}
#[rstest]
fn interpolate_value_does_not_recurse_into_resolved_value() {
let mut value = toml::Value::String("${OUTER}".to_string());
let env = envmap! {"OUTER" => "${INNER}"};
let lookup = |n: &str| env.get(n).map(|s| s.to_string());
let interp = Interpolator::new(&lookup);
interp
.interpolate_value(&mut value, Path::new("x.toml"))
.unwrap();
assert_eq!(value.as_str(), Some("${INNER}"));
}
#[rstest]
fn interpolate_value_skips_non_string_types() {
let mut value: toml::Value = toml::from_str(
r#"
port = 5432
enabled = true
rate = 1.5
"#,
)
.unwrap();
let env: HashMap<&'static str, &'static str> = HashMap::new();
let lookup = |n: &str| env.get(n).map(|s| s.to_string());
let interp = Interpolator::new(&lookup);
interp
.interpolate_value(&mut value, Path::new("x.toml"))
.unwrap();
assert_eq!(value["port"].as_integer(), Some(5432));
assert_eq!(value["enabled"].as_bool(), Some(true));
assert_eq!(value["rate"].as_float(), Some(1.5));
}
}