1use std::collections::HashMap;
29
30use spectrafit_types::CoreError;
31
32use crate::error::GraphError;
33
34#[derive(Debug, Clone, Copy, PartialEq, Eq)]
40pub enum BinOp {
41 Add,
43 Sub,
45 Mul,
47 Div,
49}
50
51#[derive(Debug, Clone, PartialEq)]
53pub enum Expr {
54 Num(f64),
56 Ref {
58 node: String,
60 param: String,
62 },
63 Binary {
65 op: BinOp,
67 lhs: Box<Expr>,
69 rhs: Box<Expr>,
71 },
72 Neg(Box<Expr>),
74}
75
76impl Expr {
77 pub fn references(&self) -> Vec<String> {
81 let mut out = Vec::new();
82 self.collect_refs(&mut out);
83 out
84 }
85
86 fn collect_refs(&self, out: &mut Vec<String>) {
87 match self {
88 Expr::Num(_) => {}
89 Expr::Ref { node, param } => out.push(format!("{}.{}", node, param)),
90 Expr::Binary { lhs, rhs, .. } => {
91 lhs.collect_refs(out);
92 rhs.collect_refs(out);
93 }
94 Expr::Neg(inner) => inner.collect_refs(out),
95 }
96 }
97
98 pub fn eval(&self, values: &HashMap<String, f64>) -> Result<f64, CoreError> {
104 match self {
105 Expr::Num(n) => Ok(*n),
106 Expr::Ref { node, param } => {
107 let key = format!("{}.{}", node, param);
108 values
109 .get(&key)
110 .copied()
111 .ok_or_else(|| GraphError::MissingParamKey(key).into())
112 }
113 Expr::Neg(inner) => Ok(-inner.eval(values)?),
114 Expr::Binary { op, lhs, rhs } => {
115 let l = lhs.eval(values)?;
116 let r = rhs.eval(values)?;
117 match op {
118 BinOp::Add => Ok(l + r),
119 BinOp::Sub => Ok(l - r),
120 BinOp::Mul => Ok(l * r),
121 BinOp::Div => {
122 if r == 0.0 {
123 Err(GraphError::DivisionByZero.into())
124 } else {
125 Ok(l / r)
126 }
127 }
128 }
129 }
130 }
131 }
132}
133
134#[derive(Debug, Clone, PartialEq)]
139enum Token {
140 Num(f64),
141 Ident(String),
142 Dot,
143 Plus,
144 Minus,
145 Star,
146 Slash,
147 LParen,
148 RParen,
149}
150
151fn tokenize(src: &str) -> Result<Vec<Token>, CoreError> {
152 let mut tokens = Vec::new();
153 let chars: Vec<char> = src.chars().collect();
154 let mut i = 0;
155 while i < chars.len() {
156 let c = chars[i];
157 match c {
158 ws if ws.is_whitespace() => {
159 i += 1;
160 }
161 '+' => {
162 tokens.push(Token::Plus);
163 i += 1;
164 }
165 '-' => {
166 tokens.push(Token::Minus);
167 i += 1;
168 }
169 '*' => {
170 tokens.push(Token::Star);
171 i += 1;
172 }
173 '/' => {
174 tokens.push(Token::Slash);
175 i += 1;
176 }
177 '(' => {
178 tokens.push(Token::LParen);
179 i += 1;
180 }
181 ')' => {
182 tokens.push(Token::RParen);
183 i += 1;
184 }
185 '.' => {
186 tokens.push(Token::Dot);
189 i += 1;
190 }
191 d if d.is_ascii_digit() => {
192 let start = i;
193 while i < chars.len() && (chars[i].is_ascii_digit() || chars[i] == '.') {
194 i += 1;
195 }
196 if i < chars.len() && (chars[i] == 'e' || chars[i] == 'E') {
198 i += 1;
199 if i < chars.len() && (chars[i] == '+' || chars[i] == '-') {
200 i += 1;
201 }
202 while i < chars.len() && chars[i].is_ascii_digit() {
203 i += 1;
204 }
205 }
206 let lexeme: String = chars[start..i].iter().collect();
207 let n = lexeme
208 .parse::<f64>()
209 .map_err(|_| CoreError::Eval(format!("invalid number literal '{}'", lexeme)))?;
210 tokens.push(Token::Num(n));
211 }
212 a if a.is_ascii_alphabetic() || a == '_' => {
213 let start = i;
214 while i < chars.len() && (chars[i].is_ascii_alphanumeric() || chars[i] == '_') {
215 i += 1;
216 }
217 let ident: String = chars[start..i].iter().collect();
218 tokens.push(Token::Ident(ident));
219 }
220 other => {
221 return Err(CoreError::Eval(format!(
222 "unexpected character '{}' in expression",
223 other
224 )));
225 }
226 }
227 }
228 Ok(tokens)
229}
230
231struct Parser {
236 tokens: Vec<Token>,
237 pos: usize,
238}
239
240impl Parser {
241 fn new(tokens: Vec<Token>) -> Self {
242 Parser { tokens, pos: 0 }
243 }
244
245 fn peek(&self) -> Option<&Token> {
246 self.tokens.get(self.pos)
247 }
248
249 fn next(&mut self) -> Option<Token> {
250 let t = self.tokens.get(self.pos).cloned();
251 if t.is_some() {
252 self.pos += 1;
253 }
254 t
255 }
256
257 fn expect(&mut self, want: &Token) -> Result<(), CoreError> {
258 match self.next() {
259 Some(ref got) if got == want => Ok(()),
260 Some(got) => Err(CoreError::Eval(format!(
261 "expected {:?}, found {:?}",
262 want, got
263 ))),
264 None => Err(CoreError::Eval(format!(
265 "expected {:?}, found end of input",
266 want
267 ))),
268 }
269 }
270
271 fn parse_expr(&mut self) -> Result<Expr, CoreError> {
273 let mut node = self.parse_term()?;
274 while let Some(tok) = self.peek() {
275 let op = match tok {
276 Token::Plus => BinOp::Add,
277 Token::Minus => BinOp::Sub,
278 _ => break,
279 };
280 self.next();
281 let rhs = self.parse_term()?;
282 node = Expr::Binary {
283 op,
284 lhs: Box::new(node),
285 rhs: Box::new(rhs),
286 };
287 }
288 Ok(node)
289 }
290
291 fn parse_term(&mut self) -> Result<Expr, CoreError> {
293 let mut node = self.parse_factor()?;
294 while let Some(tok) = self.peek() {
295 let op = match tok {
296 Token::Star => BinOp::Mul,
297 Token::Slash => BinOp::Div,
298 _ => break,
299 };
300 self.next();
301 let rhs = self.parse_factor()?;
302 node = Expr::Binary {
303 op,
304 lhs: Box::new(node),
305 rhs: Box::new(rhs),
306 };
307 }
308 Ok(node)
309 }
310
311 fn parse_factor(&mut self) -> Result<Expr, CoreError> {
313 match self.next() {
314 Some(Token::Num(n)) => Ok(Expr::Num(n)),
315 Some(Token::Minus) => Ok(Expr::Neg(Box::new(self.parse_factor()?))),
316 Some(Token::LParen) => {
317 let inner = self.parse_expr()?;
318 self.expect(&Token::RParen)?;
319 Ok(inner)
320 }
321 Some(Token::Ident(node)) => {
322 self.expect(&Token::Dot)?;
325 match self.next() {
326 Some(Token::Ident(param)) => Ok(Expr::Ref { node, param }),
327 other => Err(CoreError::Eval(format!(
328 "expected parameter name after '{}.', found {:?}",
329 node, other
330 ))),
331 }
332 }
333 other => Err(CoreError::Eval(format!(
334 "unexpected token {:?} while parsing expression",
335 other
336 ))),
337 }
338 }
339}
340
341pub fn parse(src: &str) -> Result<Expr, CoreError> {
347 let tokens = tokenize(src)?;
348 if tokens.is_empty() {
349 return Err(GraphError::EmptyExpression.into());
350 }
351 let mut parser = Parser::new(tokens);
352 let expr = parser.parse_expr()?;
353 if parser.pos != parser.tokens.len() {
354 return Err(GraphError::MalformedExpression(format!(
355 "trailing tokens after expression: {:?}",
356 &parser.tokens[parser.pos..]
357 ))
358 .into());
359 }
360 Ok(expr)
361}
362
363#[derive(Debug, Clone)]
369pub struct TiedParam {
370 pub target: String,
372 pub expr: Expr,
374}
375
376#[derive(Debug, Clone, Default)]
383pub struct TiedPlan {
384 pub order: Vec<TiedParam>,
386}
387
388impl TiedPlan {
389 pub fn build<'a, I>(edges: I) -> Result<Self, CoreError>
399 where
400 I: IntoIterator<Item = (&'a str, &'a str)>,
401 {
402 let mut parsed: Vec<TiedParam> = Vec::new();
404 let mut target_index: HashMap<String, usize> = HashMap::new();
405 for (target, src) in edges {
406 let expr = parse(src)?;
407 if target_index.contains_key(target) {
408 return Err(GraphError::DuplicateExprTarget(target.to_string()).into());
409 }
410 target_index.insert(target.to_string(), parsed.len());
411 parsed.push(TiedParam {
412 target: target.to_string(),
413 expr,
414 });
415 }
416
417 let order = topo_sort(&parsed, &target_index)?;
422 Ok(TiedPlan { order })
423 }
424
425 pub fn len(&self) -> usize {
427 self.order.len()
428 }
429
430 pub fn is_empty(&self) -> bool {
432 self.order.is_empty()
433 }
434
435 pub fn apply(&self, values: &mut HashMap<String, f64>) -> Result<(), CoreError> {
442 for tp in &self.order {
443 let v = tp.expr.eval(values)?;
444 values.insert(tp.target.clone(), v);
445 }
446 Ok(())
447 }
448}
449
450fn topo_sort(
453 parsed: &[TiedParam],
454 target_index: &HashMap<String, usize>,
455) -> Result<Vec<TiedParam>, CoreError> {
456 #[derive(Clone, Copy, PartialEq)]
457 enum Mark {
458 Unvisited,
459 InProgress,
460 Done,
461 }
462
463 let n = parsed.len();
464 let mut marks = vec![Mark::Unvisited; n];
465 let mut ordered: Vec<usize> = Vec::with_capacity(n);
466 let mut stack: Vec<(usize, usize)> = Vec::new();
469
470 for start in 0..n {
471 if marks[start] != Mark::Unvisited {
472 continue;
473 }
474 stack.push((start, 0));
475 marks[start] = Mark::InProgress;
476
477 while let Some(&(node, dep_cursor)) = stack.last() {
478 let deps: Vec<usize> = parsed[node]
480 .expr
481 .references()
482 .into_iter()
483 .filter_map(|r| target_index.get(&r).copied())
484 .collect();
485
486 if dep_cursor < deps.len() {
487 stack.last_mut().unwrap().1 += 1;
493 let dep = deps[dep_cursor];
494 match marks[dep] {
495 Mark::Done => {}
496 Mark::InProgress => {
497 return Err(GraphError::Cycle(parsed[dep].target.clone()).into());
498 }
499 Mark::Unvisited => {
500 marks[dep] = Mark::InProgress;
501 stack.push((dep, 0));
502 }
503 }
504 } else {
505 marks[node] = Mark::Done;
507 ordered.push(node);
508 stack.pop();
509 }
510 }
511 }
512
513 Ok(ordered.into_iter().map(|i| parsed[i].clone()).collect())
514}
515
516#[cfg(test)]
520mod tests {
521 use super::*;
522
523 #[test]
524 fn parse_scaled_reference() {
525 let ast = parse("0.5 * g1.amplitude").unwrap();
527 let expected = Expr::Binary {
528 op: BinOp::Mul,
529 lhs: Box::new(Expr::Num(0.5)),
530 rhs: Box::new(Expr::Ref {
531 node: "g1".to_string(),
532 param: "amplitude".to_string(),
533 }),
534 };
535 assert_eq!(ast, expected);
536 }
537
538 #[test]
539 fn parse_precedence_and_parens() {
540 let ast = parse("a.p + b.p * 2").unwrap();
542 match ast {
543 Expr::Binary {
544 op: BinOp::Add,
545 rhs,
546 ..
547 } => assert!(matches!(*rhs, Expr::Binary { op: BinOp::Mul, .. })),
548 other => panic!("unexpected AST: {:?}", other),
549 }
550
551 let ast2 = parse("(a.p + b.p) * 2").unwrap();
553 assert!(matches!(ast2, Expr::Binary { op: BinOp::Mul, .. }));
554 }
555
556 #[test]
557 fn eval_arithmetic() {
558 let ast = parse("0.5 * g1.amplitude + 1.0").unwrap();
559 let mut values = HashMap::new();
560 values.insert("g1.amplitude".to_string(), 4.0);
561 assert_eq!(ast.eval(&values).unwrap(), 3.0);
562 }
563
564 #[test]
565 fn references_collected() {
566 let ast = parse("a.x + b.y * c.z").unwrap();
567 let refs = ast.references();
568 assert_eq!(refs, vec!["a.x", "b.y", "c.z"]);
569 }
570
571 #[test]
572 fn parse_rejects_trailing_garbage() {
573 assert!(parse("1.0 +").is_err());
574 assert!(parse("g1.amplitude g2.amplitude").is_err());
575 assert!(parse("g1.").is_err());
576 assert!(parse("").is_err());
577 }
578
579 #[test]
580 fn plan_orders_dependencies() {
581 let plan = TiedPlan::build(vec![
583 ("c2.amp", "2.0 * c1.amp"),
584 ("c1.amp", "0.5 * g1.amplitude"),
585 ])
586 .unwrap();
587 assert_eq!(plan.len(), 2);
588 let order: Vec<&str> = plan.order.iter().map(|t| t.target.as_str()).collect();
589 let pos_c1 = order.iter().position(|t| *t == "c1.amp").unwrap();
590 let pos_c2 = order.iter().position(|t| *t == "c2.amp").unwrap();
591 assert!(pos_c1 < pos_c2, "c1.amp must be evaluated before c2.amp");
592 }
593
594 #[test]
595 fn plan_apply_resolves_chain() {
596 let plan = TiedPlan::build(vec![
597 ("c2.amp", "2.0 * c1.amp"),
598 ("c1.amp", "0.5 * g1.amplitude"),
599 ])
600 .unwrap();
601 let mut values = HashMap::new();
602 values.insert("g1.amplitude".to_string(), 8.0);
603 plan.apply(&mut values).unwrap();
604 assert_eq!(values["c1.amp"], 4.0);
605 assert_eq!(values["c2.amp"], 8.0);
606 }
607
608 #[test]
609 fn plan_detects_cycle_a_b_a() {
610 let res = TiedPlan::build(vec![("a.p", "b.p + 1.0"), ("b.p", "a.p * 2.0")]);
612 assert!(res.is_err(), "a→b→a cycle must be rejected");
613 let msg = format!("{}", res.unwrap_err());
614 assert!(msg.contains("cycle"), "error should mention cycle: {msg}");
615 }
616
617 #[test]
618 fn plan_detects_self_cycle() {
619 let res = TiedPlan::build(vec![("a.p", "a.p + 1.0")]);
620 assert!(res.is_err(), "self-reference cycle must be rejected");
621 }
622
623 #[test]
624 fn plan_rejects_duplicate_target() {
625 let res = TiedPlan::build(vec![("a.p", "1.0"), ("a.p", "2.0")]);
626 assert!(res.is_err());
627 }
628}