Skip to main content

greeners_core/
formula.rs

1use crate::GreenersError;
2
3/// Represents a parsed formula in the form "y ~ x1 + x2 + ... + xn"
4///
5/// # Interaction Terms (v0.3.0)
6/// - `x1 * x2` : Full interaction (expands to x1 + x2 + x1:x2)
7/// - `x1 : x2` : Only the interaction term (x1 × x2)
8///
9/// # Categorical Variables (NEW in v0.4.0)
10/// - `C(var)` : Categorical encoding (creates dummies, drops first level)
11///
12/// # Polynomial Terms (NEW in v0.4.0)
13/// - `I(x^2)` : Polynomial terms (e.g., squared, cubed)
14#[derive(Debug, Clone)]
15pub struct Formula {
16    /// Name of the dependent variable (left-hand side)
17    pub dependent: String,
18    /// Names of independent variables (right-hand side)
19    /// May include:
20    /// - Regular variables: "x1"
21    /// - Interactions: "x1:x2"
22    /// - Categorical: "C(region)"
23    /// - Polynomials: "I(x^2)"
24    pub independents: Vec<String>,
25    /// Whether to include an intercept (default: true)
26    pub intercept: bool,
27}
28
29impl Formula {
30    /// Parse a formula string in the R/Python style: "y ~ x1 + x2 + x3"
31    ///
32    /// # Syntax
33    /// - Basic: "y ~ x1 + x2 + x3" (with intercept)
34    /// - No intercept: "y ~ x1 + x2 + x3 - 1" or "y ~ 0 + x1 + x2"
35    /// - Intercept only: "y ~ 1"
36    /// - Full interaction: "y ~ x1 * x2" (expands to x1 + x2 + x1:x2)
37    /// - Interaction only: "y ~ x1 : x2" (only the interaction term)
38    /// - Categorical: "y ~ C(region)" (creates dummies)
39    /// - Polynomial: "y ~ I(x^2)" or "y ~ I(x**2)" (power terms)
40    ///
41    /// # Examples
42    /// ```
43    /// use greeners_core::formula::Formula;
44    ///
45    /// let f = Formula::parse("fte ~ tratado + t + effect").unwrap();
46    /// assert_eq!(f.dependent, "fte");
47    /// assert_eq!(f.independents, vec!["tratado", "t", "effect"]);
48    /// assert_eq!(f.intercept, true);
49    ///
50    /// let f2 = Formula::parse("y ~ x1 + x2 - 1").unwrap();
51    /// assert_eq!(f2.intercept, false);
52    /// ```
53    pub fn parse(formula: &str) -> Result<Self, GreenersError> {
54        let formula = formula.trim();
55
56        // Split by ~ to get LHS and RHS
57        let parts: Vec<&str> = formula.split('~').collect();
58        if parts.len() != 2 {
59            return Err(GreenersError::FormulaError(format!(
60                "Invalid formula syntax. Expected 'y ~ x1 + x2', got: '{}'",
61                formula
62            )));
63        }
64
65        let dependent = parts[0].trim().to_string();
66        if dependent.is_empty() {
67            return Err(GreenersError::FormulaError(
68                "Dependent variable (LHS) cannot be empty".into(),
69            ));
70        }
71
72        let rhs = parts[1].trim();
73
74        // First, handle "- 1" or "- intercept" by removing it from the string
75        // Replace "- 1" or "-1" patterns before splitting
76        let rhs_clean = rhs.replace("- 1", "").replace("-1", "");
77
78        // Parse RHS: split by + and handle special cases
79        let mut independents = Vec::new();
80        let mut intercept = true;
81
82        // Check if the original had "- 1" to disable intercept
83        if rhs.contains("- 1") || rhs.contains("-1") {
84            intercept = false;
85        }
86
87        // Split by + and process each term
88        for term in rhs_clean.split('+') {
89            let term = term.trim();
90
91            if term.is_empty() {
92                continue;
93            }
94
95            // Check for intercept control
96            if term == "1" {
97                // Explicit intercept, already default
98                intercept = true;
99                continue;
100            } else if term == "0" {
101                // Remove intercept
102                intercept = false;
103                continue;
104            }
105
106            // Check for interaction terms
107            if term.contains('*') {
108                // Full interaction: x1 * x2 expands to x1 + x2 + x1:x2
109                let vars: Vec<&str> = term.split('*').map(|s| s.trim()).collect();
110                if vars.len() != 2 {
111                    return Err(GreenersError::FormulaError(format!(
112                        "Invalid interaction term '{}'. Expected 'var1 * var2'",
113                        term
114                    )));
115                }
116
117                // Add main effects
118                independents.push(vars[0].to_string());
119                independents.push(vars[1].to_string());
120
121                // Add interaction term (using : notation)
122                independents.push(format!("{}:{}", vars[0], vars[1]));
123            } else if term.contains(':') {
124                // Interaction only: x1:x2 (just the interaction term)
125                let vars: Vec<&str> = term.split(':').map(|s| s.trim()).collect();
126                if vars.len() != 2 {
127                    return Err(GreenersError::FormulaError(format!(
128                        "Invalid interaction term '{}'. Expected 'var1:var2'",
129                        term
130                    )));
131                }
132
133                // Add interaction term as-is
134                independents.push(format!("{}:{}", vars[0], vars[1]));
135            }
136            // Function transforms: log(), exp(), sqrt(), poly(), bs()
137            else if term.starts_with("log(")
138                || term.starts_with("exp(")
139                || term.starts_with("sqrt(")
140                || term.starts_with("poly(")
141                || term.starts_with("bs(")
142            {
143                independents.push(term.to_string());
144            } else {
145                // Regular term
146                independents.push(term.to_string());
147            }
148        }
149
150        let cleaned_independents = independents;
151
152        Ok(Formula {
153            dependent,
154            independents: cleaned_independents,
155            intercept,
156        })
157    }
158
159    /// Get the total number of columns in the design matrix (including intercept if present)
160    pub fn n_cols(&self) -> usize {
161        let base = self.independents.len();
162        if self.intercept {
163            base + 1
164        } else {
165            base
166        }
167    }
168}
169
170#[cfg(test)]
171mod tests {
172    use super::*;
173
174    #[test]
175    fn test_basic_formula() {
176        let f = Formula::parse("y ~ x1 + x2 + x3").unwrap();
177        assert_eq!(f.dependent, "y");
178        assert_eq!(f.independents, vec!["x1", "x2", "x3"]);
179        assert!(f.intercept);
180        assert_eq!(f.n_cols(), 4); // intercept + 3 vars
181    }
182
183    #[test]
184    fn test_formula_no_intercept() {
185        let f = Formula::parse("y ~ x1 + x2 - 1").unwrap();
186        assert_eq!(f.dependent, "y");
187        assert_eq!(f.independents, vec!["x1", "x2"]);
188        assert!(!f.intercept);
189        assert_eq!(f.n_cols(), 2);
190    }
191
192    #[test]
193    fn test_formula_zero_intercept() {
194        let f = Formula::parse("y ~ 0 + x1 + x2").unwrap();
195        assert_eq!(f.dependent, "y");
196        assert_eq!(f.independents, vec!["x1", "x2"]);
197        assert!(!f.intercept);
198    }
199
200    #[test]
201    fn test_intercept_only() {
202        let f = Formula::parse("y ~ 1").unwrap();
203        assert_eq!(f.dependent, "y");
204        assert_eq!(f.independents.len(), 0);
205        assert!(f.intercept);
206        assert_eq!(f.n_cols(), 1);
207    }
208
209    #[test]
210    fn test_invalid_formula() {
211        assert!(Formula::parse("invalid").is_err());
212        assert!(Formula::parse("~ x1 + x2").is_err());
213        assert!(Formula::parse("y ~").is_ok()); // empty RHS is technically ok
214    }
215
216    #[test]
217    fn test_full_interaction() {
218        // x1 * x2 should expand to x1 + x2 + x1:x2
219        let f = Formula::parse("y ~ x1 * x2").unwrap();
220        assert_eq!(f.dependent, "y");
221        assert_eq!(f.independents, vec!["x1", "x2", "x1:x2"]);
222        assert!(f.intercept);
223        assert_eq!(f.n_cols(), 4); // intercept + x1 + x2 + x1:x2
224    }
225
226    #[test]
227    fn test_interaction_only() {
228        // x1:x2 should only add the interaction term
229        let f = Formula::parse("y ~ x1 : x2").unwrap();
230        assert_eq!(f.dependent, "y");
231        assert_eq!(f.independents, vec!["x1:x2"]);
232        assert!(f.intercept);
233        assert_eq!(f.n_cols(), 2); // intercept + x1:x2
234    }
235
236    #[test]
237    fn test_mixed_interaction() {
238        // Combination of regular terms and interactions
239        let f = Formula::parse("y ~ x1 + x2 * x3 + x4").unwrap();
240        assert_eq!(f.dependent, "y");
241        assert_eq!(f.independents, vec!["x1", "x2", "x3", "x2:x3", "x4"]);
242        assert!(f.intercept);
243        assert_eq!(f.n_cols(), 6); // intercept + x1 + x2 + x3 + x2:x3 + x4
244    }
245
246    #[test]
247    fn test_interaction_no_intercept() {
248        let f = Formula::parse("y ~ x1 * x2 - 1").unwrap();
249        assert_eq!(f.dependent, "y");
250        assert_eq!(f.independents, vec!["x1", "x2", "x1:x2"]);
251        assert!(!f.intercept);
252        assert_eq!(f.n_cols(), 3); // x1 + x2 + x1:x2 (no intercept)
253    }
254}