1use crate::GreenersError;
2
3#[derive(Debug, Clone)]
15pub struct Formula {
16 pub dependent: String,
18 pub independents: Vec<String>,
25 pub intercept: bool,
27}
28
29impl Formula {
30 pub fn parse(formula: &str) -> Result<Self, GreenersError> {
54 let formula = formula.trim();
55
56 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 let rhs_clean = rhs.replace("- 1", "").replace("-1", "");
77
78 let mut independents = Vec::new();
80 let mut intercept = true;
81
82 if rhs.contains("- 1") || rhs.contains("-1") {
84 intercept = false;
85 }
86
87 for term in rhs_clean.split('+') {
89 let term = term.trim();
90
91 if term.is_empty() {
92 continue;
93 }
94
95 if term == "1" {
97 intercept = true;
99 continue;
100 } else if term == "0" {
101 intercept = false;
103 continue;
104 }
105
106 if term.contains('*') {
108 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 independents.push(vars[0].to_string());
119 independents.push(vars[1].to_string());
120
121 independents.push(format!("{}:{}", vars[0], vars[1]));
123 } else if term.contains(':') {
124 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 independents.push(format!("{}:{}", vars[0], vars[1]));
135 }
136 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 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 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); }
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()); }
215
216 #[test]
217 fn test_full_interaction() {
218 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); }
225
226 #[test]
227 fn test_interaction_only() {
228 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); }
235
236 #[test]
237 fn test_mixed_interaction() {
238 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); }
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); }
254}