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()
});
static CBIND_LHS: LazyLock<Regex> = LazyLock::new(|| {
Regex::new(r"^cbind\(([_A-Za-z][_A-Za-z0-9]*),([_A-Za-z][_A-Za-z0-9]*)\)$").unwrap()
});
static OFFSET_TERM: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r"(^|\+)offset\(((?:[^()]|\([^()]*\))*)\)").unwrap());
static NO_INTERCEPT_MINUS: LazyLock<Regex> = LazyLock::new(|| Regex::new(r"-1(\+|$)").unwrap());
static NO_INTERCEPT_ZERO: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r"(?:^|\+)0(\+|$)").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>,
pub has_intercept: bool,
pub offset: Option<String>,
pub cbind: Option<(String, String)>,
}
#[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);
let cbind = if let Some(m) = CBIND_LHS.captures(&dep) {
Some((
m.get(1).unwrap().as_str().to_string(),
m.get(2).unwrap().as_str().to_string(),
))
} else if dep.starts_with("cbind(") {
return Err(ParseError::Syntax {
pos: 0,
msg: format!("cbind() takes exactly two column names, got '{dep}'"),
});
} else if !dep.is_empty() && !is_identifier(&dep) {
return Err(ParseError::Syntax {
pos: 0,
msg: format!("expected a response column name, got '{dep}'"),
});
} else {
None
};
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)?;
let mut offset = None;
let mut rhs_stripped = rhs_stripped;
let offsets: Vec<String> = OFFSET_TERM
.captures_iter(&rhs_stripped)
.map(|m| m.get(2).unwrap().as_str().to_string())
.collect();
if offsets.len() > 1 {
return Err(ParseError::Syntax {
pos: 0,
msg: "at most one offset() term is allowed".into(),
});
}
if let Some(expr) = offsets.into_iter().next() {
if !(is_identifier(&expr) || parse_transform(&expr).is_some()) {
return Err(ParseError::Syntax {
pos: 0,
msg: format!(
"offset() takes a column name or a whitelisted transform of one, got '{expr}'"
),
});
}
offset = Some(expr);
rhs_stripped = clean_residual_plusses(&OFFSET_TERM.replace(&rhs_stripped, "+"));
}
let mut has_intercept = true;
if NO_INTERCEPT_MINUS.is_match(&rhs_stripped) || NO_INTERCEPT_ZERO.is_match(&rhs_stripped) {
has_intercept = false;
rhs_stripped = NO_INTERCEPT_MINUS
.replace_all(&rhs_stripped, "$1")
.into_owned();
rhs_stripped = NO_INTERCEPT_ZERO
.replace_all(&rhs_stripped, "+$1")
.into_owned();
rhs_stripped = clean_residual_plusses(&rhs_stripped);
}
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,
has_intercept,
offset,
cbind,
})
}
fn extract_stage(
rhs: &str,
rank: u8,
work: &mut String,
seen: &mut std::collections::BTreeSet<String>,
effects: &mut Vec<((u8, usize), 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;
};
let pos = rhs.find(m.get(0).unwrap().as_str()).unwrap_or(usize::MAX);
for (name, effect) in make(&m)? {
if !seen.insert(name.clone()) {
return Err(ParseError::DuplicateGroupingVar { name });
}
effects.push(((rank, pos), 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<((u8, usize), RandomEffect)> = Vec::new();
let mut work = rhs.to_string();
if RE_SUPPRESS.is_match(rhs) {
return Err(ParseError::RandomInterceptSuppressionUnsupported);
}
extract_stage(
rhs,
0,
&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(rhs, 1, &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(
rhs,
1,
&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(
rhs,
1,
&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(rhs, 1, &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,
},
)])
})?;
effects.sort_by_key(|(key, _)| *key);
let effects = effects.into_iter().map(|(_, e)| e).collect();
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 => 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) || parse_transform(s).is_some() {
Ok(s.to_string())
} else {
Err(ParseError::Syntax {
pos: 0,
msg: format!("expected identifier, got '{s}'"),
})
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) enum Transform {
Log,
Sqrt,
Exp,
Pow(u32),
}
static TRANSFORM_CALL: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r"^(log|sqrt|exp)\(([_A-Za-z][_A-Za-z0-9]*)\)$").unwrap());
static TRANSFORM_POW: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r"^I\(([_A-Za-z][_A-Za-z0-9]*)\^([2-9]|[1-9][0-9])\)$").unwrap());
pub(super) fn parse_transform(s: &str) -> Option<(Transform, &str)> {
if let Some(m) = TRANSFORM_CALL.captures(s) {
let t = match m.get(1).unwrap().as_str() {
"log" => Transform::Log,
"sqrt" => Transform::Sqrt,
_ => Transform::Exp,
};
return Some((t, m.get(2).unwrap().as_str()));
}
let m = TRANSFORM_POW.captures(s)?;
let k: u32 = m.get(2).unwrap().as_str().parse().ok()?;
Some((Transform::Pow(k), m.get(1).unwrap().as_str()))
}
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
}