use std::sync::Arc;
use rustc_hash::FxHashSet;
use crate::{Expr, Pattern};
#[must_use]
pub fn free_vars(expr: &Expr) -> FxHashSet<Arc<str>> {
let mut vars = FxHashSet::default();
collect_free(expr, &mut FxHashSet::default(), &mut vars);
vars
}
fn collect_free(expr: &Expr, bound: &mut FxHashSet<Arc<str>>, free: &mut FxHashSet<Arc<str>>) {
match expr {
Expr::Var(name) => {
if !bound.contains(name) {
free.insert(Arc::clone(name));
}
}
Expr::Lam(param, body) => {
let newly_bound = bound.insert(Arc::clone(param));
collect_free(body, bound, free);
if newly_bound {
bound.remove(param);
}
}
Expr::App(func, arg) => {
collect_free(func, bound, free);
collect_free(arg, bound, free);
}
Expr::Lit(_) => {}
Expr::Record(fields) => {
for (_, v) in fields {
collect_free(v, bound, free);
}
}
Expr::List(items) => {
for item in items {
collect_free(item, bound, free);
}
}
Expr::Field(expr, _) => collect_free(expr, bound, free),
Expr::Index(expr, idx) => {
collect_free(expr, bound, free);
collect_free(idx, bound, free);
}
Expr::Match { scrutinee, arms } => {
collect_free(scrutinee, bound, free);
for (pat, body) in arms {
let pat_vars = pattern_vars(pat);
let mut inserted = Vec::new();
for v in &pat_vars {
if bound.insert(Arc::clone(v)) {
inserted.push(Arc::clone(v));
}
}
collect_free(body, bound, free);
for v in &inserted {
bound.remove(v);
}
}
}
Expr::Let { name, value, body } => {
collect_free(value, bound, free);
let newly_bound = bound.insert(Arc::clone(name));
collect_free(body, bound, free);
if newly_bound {
bound.remove(name);
}
}
Expr::Builtin(_, args) => {
for arg in args {
collect_free(arg, bound, free);
}
}
}
}
#[must_use]
pub fn pattern_vars(pat: &Pattern) -> Vec<Arc<str>> {
let mut vars = Vec::new();
collect_pattern_vars(pat, &mut vars);
vars
}
fn collect_pattern_vars(pat: &Pattern, vars: &mut Vec<Arc<str>>) {
match pat {
Pattern::Wildcard | Pattern::Lit(_) => {}
Pattern::Var(name) => vars.push(Arc::clone(name)),
Pattern::Record(fields) => {
for (_, p) in fields {
collect_pattern_vars(p, vars);
}
}
Pattern::List(items) => {
for p in items {
collect_pattern_vars(p, vars);
}
}
Pattern::Constructor(_, args) => {
for p in args {
collect_pattern_vars(p, vars);
}
}
}
}
fn rename_pattern_var(pat: &Pattern, from: &str, to: &Arc<str>) -> Pattern {
match pat {
Pattern::Wildcard | Pattern::Lit(_) => pat.clone(),
Pattern::Var(v) => {
if &**v == from {
Pattern::Var(Arc::clone(to))
} else {
pat.clone()
}
}
Pattern::Record(fields) => Pattern::Record(
fields
.iter()
.map(|(k, p)| (Arc::clone(k), rename_pattern_var(p, from, to)))
.collect(),
),
Pattern::List(items) => Pattern::List(
items
.iter()
.map(|p| rename_pattern_var(p, from, to))
.collect(),
),
Pattern::Constructor(ctor, args) => Pattern::Constructor(
Arc::clone(ctor),
args.iter()
.map(|p| rename_pattern_var(p, from, to))
.collect(),
),
}
}
fn rename_avoid_set(body: &Expr, name: &str, replacement: &Expr) -> FxHashSet<Arc<str>> {
let mut avoid = free_vars(replacement);
avoid.extend(free_vars(body));
avoid.insert(Arc::from(name));
avoid
}
#[must_use]
pub fn substitute(expr: &Expr, name: &str, replacement: &Expr) -> Expr {
match expr {
Expr::Var(v) => {
if &**v == name {
replacement.clone()
} else {
expr.clone()
}
}
Expr::Lam(param, body) => {
if &**param == name {
expr.clone()
} else if free_vars(replacement).contains(param) {
let fresh = fresh_name(param, &rename_avoid_set(body, name, replacement));
let renamed_body = substitute(body, param, &Expr::Var(Arc::clone(&fresh)));
Expr::Lam(
fresh,
Box::new(substitute(&renamed_body, name, replacement)),
)
} else {
Expr::Lam(
Arc::clone(param),
Box::new(substitute(body, name, replacement)),
)
}
}
Expr::App(func, arg) => Expr::App(
Box::new(substitute(func, name, replacement)),
Box::new(substitute(arg, name, replacement)),
),
Expr::Lit(_) => expr.clone(),
Expr::Record(fields) => Expr::Record(
fields
.iter()
.map(|(k, v)| (Arc::clone(k), substitute(v, name, replacement)))
.collect(),
),
Expr::List(items) => Expr::List(
items
.iter()
.map(|i| substitute(i, name, replacement))
.collect(),
),
Expr::Field(e, f) => Expr::Field(Box::new(substitute(e, name, replacement)), Arc::clone(f)),
Expr::Index(e, idx) => Expr::Index(
Box::new(substitute(e, name, replacement)),
Box::new(substitute(idx, name, replacement)),
),
Expr::Match { scrutinee, arms } => Expr::Match {
scrutinee: Box::new(substitute(scrutinee, name, replacement)),
arms: arms
.iter()
.map(|(pat, body)| substitute_match_arm(pat, body, name, replacement))
.collect(),
},
Expr::Let {
name: let_name,
value,
body,
} => {
let new_value = substitute(value, name, replacement);
if &**let_name == name {
Expr::Let {
name: Arc::clone(let_name),
value: Box::new(new_value),
body: body.clone(),
}
} else if free_vars(replacement).contains(let_name) {
let fresh = fresh_name(let_name, &rename_avoid_set(body, name, replacement));
let renamed_body = substitute(body, let_name, &Expr::Var(Arc::clone(&fresh)));
Expr::Let {
name: fresh,
value: Box::new(new_value),
body: Box::new(substitute(&renamed_body, name, replacement)),
}
} else {
Expr::Let {
name: Arc::clone(let_name),
value: Box::new(new_value),
body: Box::new(substitute(body, name, replacement)),
}
}
}
Expr::Builtin(op, args) => Expr::Builtin(
*op,
args.iter()
.map(|a| substitute(a, name, replacement))
.collect(),
),
}
}
fn substitute_match_arm(
pat: &Pattern,
body: &Expr,
name: &str,
replacement: &Expr,
) -> (Pattern, Expr) {
let pvars = pattern_vars(pat);
if pvars.iter().any(|v| &**v == name) {
return (pat.clone(), body.clone());
}
let replacement_free = free_vars(replacement);
let mut avoid = rename_avoid_set(body, name, replacement);
avoid.extend(pvars.iter().cloned());
let mut new_pat = pat.clone();
let mut new_body = body.clone();
for v in &pvars {
if !replacement_free.contains(v) {
continue;
}
let fresh = fresh_name(v, &avoid);
avoid.insert(Arc::clone(&fresh));
new_pat = rename_pattern_var(&new_pat, v, &fresh);
new_body = substitute(&new_body, v, &Expr::Var(Arc::clone(&fresh)));
}
(new_pat, substitute(&new_body, name, replacement))
}
fn fresh_name(base: &str, avoid: &FxHashSet<Arc<str>>) -> Arc<str> {
let mut candidate = format!("{base}'");
while avoid.contains(candidate.as_str()) {
candidate.push('\'');
}
Arc::from(candidate)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::eval::EvalConfig;
use crate::{Env, Literal};
#[test]
fn free_vars_simple() {
let expr = Expr::lam(
"x",
Expr::builtin(crate::BuiltinOp::Add, vec![Expr::var("x"), Expr::var("y")]),
);
let fv = free_vars(&expr);
assert!(fv.contains("y"));
assert!(!fv.contains("x"));
}
#[test]
fn substitute_simple() {
let expr = Expr::builtin(
crate::BuiltinOp::Add,
vec![Expr::var("x"), Expr::Lit(Literal::Int(1))],
);
let result = substitute(&expr, "x", &Expr::Lit(Literal::Int(42)));
assert_eq!(
result,
Expr::builtin(
crate::BuiltinOp::Add,
vec![Expr::Lit(Literal::Int(42)), Expr::Lit(Literal::Int(1))],
)
);
}
#[test]
fn substitute_avoids_capture() {
let expr = Expr::lam(
"y",
Expr::builtin(crate::BuiltinOp::Add, vec![Expr::var("x"), Expr::var("y")]),
);
let result = substitute(&expr, "x", &Expr::var("y"));
match &result {
Expr::Lam(param, _) => assert_ne!(&**param, "y"),
_ => panic!("expected Lam"),
}
}
#[test]
fn substitute_shadowed_by_let() {
let expr = Expr::let_in(
"x",
Expr::Lit(Literal::Int(1)),
Expr::builtin(crate::BuiltinOp::Add, vec![Expr::var("x"), Expr::var("y")]),
);
let result = substitute(&expr, "x", &Expr::Lit(Literal::Int(99)));
match &result {
Expr::Let { value, body, .. } => {
assert_eq!(**value, Expr::Lit(Literal::Int(1)));
assert!(
matches!(body.as_ref(), Expr::Builtin(_, args) if matches!(&args[0], Expr::Var(v) if &**v == "x"))
);
}
_ => panic!("expected Let"),
}
}
#[test]
fn free_vars_lambda_does_not_leak_into_siblings() {
let expr = Expr::Record(vec![
(Arc::from("f"), Expr::lam("x", Expr::var("x"))),
(Arc::from("g"), Expr::var("x")),
]);
let fv = free_vars(&expr);
assert!(
fv.contains("x"),
"sibling occurrence of x is free, got {fv:?}"
);
}
#[test]
fn free_vars_respects_shadowed_lambda_binder() {
let expr = Expr::lam(
"x",
Expr::app(Expr::lam("x", Expr::var("x")), Expr::var("x")),
);
let fv = free_vars(&expr);
assert!(fv.is_empty(), "expression is closed, got {fv:?}");
}
#[test]
fn free_vars_let_does_not_leak_into_siblings() {
let expr = Expr::Record(vec![
(
Arc::from("f"),
Expr::let_in("x", Expr::Lit(Literal::Int(1)), Expr::var("x")),
),
(Arc::from("g"), Expr::var("x")),
]);
let fv = free_vars(&expr);
assert!(
fv.contains("x"),
"sibling occurrence of x is free, got {fv:?}"
);
}
#[test]
fn free_vars_respects_shadowed_let_binder() {
let expr = Expr::lam(
"x",
Expr::let_in("x", Expr::Lit(Literal::Int(1)), Expr::var("x")),
);
let fv = free_vars(&expr);
assert!(fv.is_empty(), "expression is closed, got {fv:?}");
}
#[test]
fn substitute_avoids_capture_under_let() {
let expr = Expr::let_in("x", Expr::Lit(Literal::Int(1)), Expr::var("z"));
let result = substitute(&expr, "z", &Expr::var("x"));
let env = Env::new().extend(Arc::from("x"), Literal::Int(42));
let Ok(value) = crate::eval::eval(&result, &env, &EvalConfig::default()) else {
panic!("substituted expression must evaluate, got {result:?}");
};
assert_eq!(value, Literal::Int(42), "got {result:?}");
}
#[test]
fn substitute_avoids_capture_under_match_arm() {
let expr = Expr::Match {
scrutinee: Box::new(Expr::Lit(Literal::Int(1))),
arms: vec![(Pattern::Var(Arc::from("x")), Expr::var("z"))],
};
let result = substitute(&expr, "z", &Expr::var("x"));
let env = Env::new().extend(Arc::from("x"), Literal::Int(42));
let Ok(value) = crate::eval::eval(&result, &env, &EvalConfig::default()) else {
panic!("substituted expression must evaluate, got {result:?}");
};
assert_eq!(value, Literal::Int(42), "got {result:?}");
}
#[test]
fn substitute_fresh_name_avoids_body_free_variables() {
let expr = Expr::lam(
"y",
Expr::builtin(
crate::BuiltinOp::Add,
vec![
Expr::var("y'"),
Expr::builtin(crate::BuiltinOp::Add, vec![Expr::var("x"), Expr::var("y")]),
],
),
);
let result = substitute(&expr, "x", &Expr::var("y"));
match &result {
Expr::Lam(param, _) => {
assert_ne!(&**param, "y", "binder must be renamed");
assert_ne!(&**param, "y'", "binder must not capture the body's y'");
}
other => panic!("expected Lam, got {other:?}"),
}
}
}