use super::error::ParseError;
use regex::Regex;
use std::sync::LazyLock;
static RE_SUPPRESS: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r"\((?:0|-1)(?:\+[^|]*)?\|[^)]*\)").unwrap());
static RE_NESTED: LazyLock<Regex> = LazyLock::new(|| {
Regex::new(r"\(1\|([_A-Za-z][_A-Za-z0-9]*)/([_A-Za-z][_A-Za-z0-9]*)\)").unwrap()
});
static RE_SLOPE: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r"\(1\+([^|]+?)\|([_A-Za-z][_A-Za-z0-9]*)\)").unwrap());
static RE_ISLOPE: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r"\(([_A-Za-z][^|]*?)\|([_A-Za-z][_A-Za-z0-9]*)\)").unwrap());
static RE_INT: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r"\(1\|([_A-Za-z][_A-Za-z0-9]*)\)").unwrap());
static RE_INTERACTION: LazyLock<Regex> = LazyLock::new(|| {
Regex::new(r"\(1\|([_A-Za-z][_A-Za-z0-9]*):([_A-Za-z][_A-Za-z0-9]*)\)").unwrap()
});
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ParsedFormula {
pub dependent: String,
pub predictors: Vec<String>,
pub terms: Vec<Term>,
pub random_effects: Vec<RandomEffect>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Term {
Main {
name: String,
},
Interaction {
vars: Vec<String>,
},
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RandomEffect {
Intercept {
group: String,
parent: Option<String>,
},
Slope {
group: String,
vars: Vec<String>,
},
}
pub fn parse(input: &str) -> Result<ParsedFormula, ParseError> {
let cleaned: String = input.chars().filter(|c| !c.is_whitespace()).collect();
if cleaned.is_empty() {
return Err(ParseError::EmptyFormula);
}
let (mut dep, rhs) = split_at_separator(&cleaned);
if dep.is_empty() {
dep = "explained_variable".to_string();
}
if rhs.is_empty() {
return Err(ParseError::EmptyFormula);
}
let (random_effects, rhs_stripped) = extract_random_effects(&rhs)?;
if find_term_removal(&rhs_stripped).is_some() {
return Err(ParseError::TermRemovalUnsupported);
}
let (predictors, terms) = if rhs_stripped.is_empty() {
(Vec::new(), Vec::new())
} else {
parse_rhs(&rhs_stripped)?
};
Ok(ParsedFormula {
dependent: dep,
predictors,
terms,
random_effects,
})
}
fn extract_stage(
work: &mut String,
seen: &mut std::collections::BTreeSet<String>,
effects: &mut Vec<RandomEffect>,
regex: &Regex,
mut make: impl FnMut(®ex::Captures) -> Result<Vec<(String, RandomEffect)>, ParseError>,
) -> Result<(), ParseError> {
loop {
let snapshot = work.clone();
let Some(m) = regex.captures(&snapshot) else {
break;
};
for (name, effect) in make(&m)? {
if !seen.insert(name.clone()) {
return Err(ParseError::DuplicateGroupingVar { name });
}
effects.push(effect);
}
*work = regex.replacen(&snapshot, 1, "").into_owned();
}
Ok(())
}
fn extract_random_effects(rhs: &str) -> Result<(Vec<RandomEffect>, String), ParseError> {
use std::collections::BTreeSet;
let mut seen: BTreeSet<String> = BTreeSet::new();
let mut effects: Vec<RandomEffect> = Vec::new();
let mut work = rhs.to_string();
if RE_SUPPRESS.is_match(rhs) {
return Err(ParseError::RandomInterceptSuppressionUnsupported);
}
extract_stage(&mut work, &mut seen, &mut effects, &RE_NESTED, |m| {
let parent_name = m.get(1).unwrap().as_str().to_string();
let child_name = m.get(2).unwrap().as_str().to_string();
let joined = format!("{parent_name}:{child_name}");
Ok(vec![
(
parent_name.clone(),
RandomEffect::Intercept {
group: parent_name.clone(),
parent: None,
},
),
(
joined.clone(),
RandomEffect::Intercept {
group: joined,
parent: Some(parent_name),
},
),
])
})?;
extract_stage(&mut work, &mut seen, &mut effects, &RE_SLOPE, |m| {
let var_list_raw = m.get(1).unwrap().as_str();
let group = m.get(2).unwrap().as_str().to_string();
let raw_tokens: Vec<&str> = var_list_raw
.split('+')
.map(str::trim)
.filter(|s| !s.is_empty())
.collect();
if raw_tokens.is_empty() {
return Err(ParseError::EmptySlopeTerm { group });
}
let vars: Vec<String> = raw_tokens
.iter()
.filter(|s| **s != "1")
.map(|s| String::from(*s))
.collect();
let effect = if vars.is_empty() {
RandomEffect::Intercept {
group: group.clone(),
parent: None,
}
} else {
RandomEffect::Slope {
group: group.clone(),
vars,
}
};
Ok(vec![(group, effect)])
})?;
extract_stage(&mut work, &mut seen, &mut effects, &RE_ISLOPE, |m| {
let var_list_raw = m.get(1).unwrap().as_str();
let group = m.get(2).unwrap().as_str().to_string();
let vars: Vec<String> = var_list_raw
.split('+')
.map(str::trim)
.filter(|s| !s.is_empty() && *s != "1")
.map(String::from)
.collect();
Ok(vec![(group.clone(), RandomEffect::Slope { group, vars })])
})?;
extract_stage(&mut work, &mut seen, &mut effects, &RE_INTERACTION, |m| {
let lhs = m.get(1).unwrap().as_str().to_string();
let rhs = m.get(2).unwrap().as_str().to_string();
let joined = format!("{lhs}:{rhs}");
Ok(vec![(
joined.clone(),
RandomEffect::Intercept {
group: joined,
parent: None,
},
)])
})?;
extract_stage(&mut work, &mut seen, &mut effects, &RE_INT, |m| {
let name = m.get(1).unwrap().as_str().to_string();
Ok(vec![(
name.clone(),
RandomEffect::Intercept {
group: name,
parent: None,
},
)])
})?;
let cleaned = clean_residual_plusses(&work);
Ok((effects, cleaned))
}
fn clean_residual_plusses(s: &str) -> String {
let mut out = String::with_capacity(s.len());
let mut prev_plus = false;
for ch in s.chars() {
if ch == '+' {
if !prev_plus && !out.is_empty() {
out.push('+');
prev_plus = true;
}
} else if !ch.is_whitespace() {
out.push(ch);
prev_plus = false;
}
}
while out.starts_with('+') {
out.remove(0);
}
while out.ends_with('+') {
out.pop();
}
out
}
fn split_at_separator(s: &str) -> (String, String) {
if let Some((l, r)) = s.split_once('~') {
(l.to_string(), r.to_string())
} else if let Some((l, r)) = s.split_once('=') {
(l.to_string(), r.to_string())
} else {
("explained_variable".to_string(), s.to_string())
}
}
fn find_term_removal(s: &str) -> Option<usize> {
let bytes = s.as_bytes();
let mut depth = 0i32;
for (i, &b) in bytes.iter().enumerate() {
match b {
b'(' => depth += 1,
b')' => depth -= 1,
b'-' if depth == 0 => {
let next = bytes.get(i + 1).copied().unwrap_or(b' ');
if !next.is_ascii_digit() {
return Some(i);
}
}
_ => {}
}
}
None
}
fn parse_rhs(rhs: &str) -> Result<(Vec<String>, Vec<Term>), ParseError> {
use std::collections::BTreeSet;
let mut predictors: Vec<String> = Vec::new();
let mut seen_pred: BTreeSet<String> = BTreeSet::new();
let mut terms: Vec<Term> = Vec::new();
let mut seen_term: BTreeSet<String> = BTreeSet::new();
for raw_term in rhs.split('+') {
let term = raw_term.trim();
if term.is_empty() {
continue;
}
if term.contains('*') {
let vars = parse_identifier_list(term, &['*'])?;
register_vars(&vars, &mut predictors, &mut seen_pred);
for v in &vars {
if seen_term.insert(v.clone()) {
terms.push(Term::Main { name: v.clone() });
}
}
for r in 2..=vars.len() {
for combo in combinations(&vars, r) {
let key = combo.join(":");
if seen_term.insert(key) {
terms.push(Term::Interaction { vars: combo });
}
}
}
} else if term.contains(':') {
let vars = parse_identifier_list(term, &[':'])?;
register_vars(&vars, &mut predictors, &mut seen_pred);
let key = vars.join(":");
if seen_term.insert(key) {
terms.push(Term::Interaction { vars });
}
} else {
let name = parse_single_identifier(term)?;
if seen_pred.insert(name.clone()) {
predictors.push(name.clone());
}
if seen_term.insert(name.clone()) {
terms.push(Term::Main { name });
}
}
}
Ok((predictors, terms))
}
fn parse_single_identifier(s: &str) -> Result<String, ParseError> {
if is_identifier(s) {
Ok(s.to_string())
} else {
Err(ParseError::Syntax {
pos: 0,
msg: format!("expected identifier, got '{s}'"),
})
}
}
fn parse_identifier_list(s: &str, seps: &[char]) -> Result<Vec<String>, ParseError> {
let mut parts: Vec<&str> = vec![s];
for sep in seps {
parts = parts.into_iter().flat_map(|p| p.split(*sep)).collect();
}
parts
.into_iter()
.map(str::trim)
.filter(|p| !p.is_empty())
.map(parse_single_identifier)
.collect()
}
fn is_identifier(s: &str) -> bool {
let mut chars = s.chars();
match chars.next() {
Some(c) if c.is_alphabetic() || c == '_' => {}
_ => return false,
}
chars.all(|c| c.is_alphanumeric() || c == '_')
}
fn register_vars(
vars: &[String],
predictors: &mut Vec<String>,
seen: &mut std::collections::BTreeSet<String>,
) {
for v in vars {
if seen.insert(v.clone()) {
predictors.push(v.clone());
}
}
}
fn combinations<T: Clone>(items: &[T], r: usize) -> Vec<Vec<T>> {
let n = items.len();
if r == 0 || r > n {
return vec![];
}
let mut idx: Vec<usize> = (0..r).collect();
let mut out: Vec<Vec<T>> = Vec::new();
loop {
out.push(idx.iter().map(|&i| items[i].clone()).collect());
let mut i = r;
while i > 0 && idx[i - 1] == n - r + (i - 1) {
i -= 1;
}
if i == 0 {
break;
}
idx[i - 1] += 1;
for j in i..r {
idx[j] = idx[j - 1] + 1;
}
}
out
}