Skip to main content

spectrafit_graph/
expr.rs

1//! Restricted-grammar expression parser & evaluator for tied parameters.
2//!
3//! This module implements the *scaffolding* for parameter tying via
4//! [`ExprEdge`](spectrafit_types::ExprEdge).  The supported grammar is
5//! deliberately small:
6//!
7//! ```text
8//! expr    := term (("+" | "-") term)*
9//! term    := factor (("*" | "/") factor)*
10//! factor  := NUMBER | REF | "(" expr ")" | "-" factor
11//! REF     := IDENT "." IDENT          # e.g. g1.amplitude
12//! ```
13//!
14//! Transcendental functions (`sin`, `exp`, …) and the inline
15//! [`ParameterSpec::expr`](spectrafit_types::ParameterSpec) string path are
16//! **out of scope** and deferred to a future iteration.
17//!
18//! # Status
19//!
20//! The parser, AST, evaluator, and dependency-ordering (topological sort with
21//! cycle detection) are fully implemented and unit-tested.  The per-iteration
22//! evaluation is wired into both solver front-ends via
23//! `spectrafit-solver::lm_problem::set_free_and_tied` (landed in M6); end-to-end
24//! tied-fit coverage lives in the solver crate
25//! (`dispatch::tests::test_tied_amplitude_fit_recovers_ratio`,
26//! `test_tied_fit_reduces_free_param_count`).
27
28use std::collections::HashMap;
29
30use spectrafit_types::CoreError;
31
32use crate::error::GraphError;
33
34// ---------------------------------------------------------------------------
35// AST
36// ---------------------------------------------------------------------------
37
38/// A binary arithmetic operator in the restricted grammar.
39#[derive(Debug, Clone, Copy, PartialEq, Eq)]
40pub enum BinOp {
41    /// Addition `+`.
42    Add,
43    /// Subtraction `-`.
44    Sub,
45    /// Multiplication `*`.
46    Mul,
47    /// Division `/`.
48    Div,
49}
50
51/// Parsed expression abstract syntax tree.
52#[derive(Debug, Clone, PartialEq)]
53pub enum Expr {
54    /// A numeric literal, e.g. `0.5`.
55    Num(f64),
56    /// A reference to another node's parameter, `node.param`.
57    Ref {
58        /// Source node id (left of the dot).
59        node: String,
60        /// Source parameter name (right of the dot).
61        param: String,
62    },
63    /// A binary operation `lhs op rhs`.
64    Binary {
65        /// Operator.
66        op: BinOp,
67        /// Left operand.
68        lhs: Box<Expr>,
69        /// Right operand.
70        rhs: Box<Expr>,
71    },
72    /// Unary negation `-operand`.
73    Neg(Box<Expr>),
74}
75
76impl Expr {
77    /// Collect every `node.param` reference appearing in this expression as a
78    /// flat `"node.param"` key.  Order follows a left-to-right tree walk;
79    /// duplicates are preserved (dedupe at the call site if required).
80    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    /// Evaluate this expression given a flat `"node.param" -> value` map.
99    ///
100    /// # Errors
101    /// Returns [`CoreError::Eval`] if a referenced key is missing or a
102    /// division by zero occurs.
103    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// ---------------------------------------------------------------------------
135// Lexer
136// ---------------------------------------------------------------------------
137
138#[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                // A leading-dot number (e.g. ".5") is handled in the digit arm;
187                // a bare dot here is the node.param separator.
188                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                // Optional exponent: 1e-3, 2.5E+4
197                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
231// ---------------------------------------------------------------------------
232// Recursive-descent parser
233// ---------------------------------------------------------------------------
234
235struct 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    // expr := term (("+"|"-") term)*
272    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    // term := factor (("*"|"/") factor)*
292    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    // factor := NUMBER | REF | "(" expr ")" | "-" factor
312    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                // A reference MUST be `node.param`.  A bare identifier is not a
323                // supported leaf (no free-standing symbols / functions yet).
324                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
341/// Parse a restricted-grammar expression string into an [`Expr`] AST.
342///
343/// # Errors
344/// Returns [`CoreError::Eval`] on any lex/parse error, including trailing
345/// tokens, unbalanced parentheses, or unsupported syntax.
346pub 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// ---------------------------------------------------------------------------
364// Dependency-ordered evaluation plan
365// ---------------------------------------------------------------------------
366
367/// A single tied-parameter assignment: `target = expr`.
368#[derive(Debug, Clone)]
369pub struct TiedParam {
370    /// Fully-qualified target key `"node.param"`.
371    pub target: String,
372    /// Parsed right-hand-side expression.
373    pub expr: Expr,
374}
375
376/// A dependency-ordered plan for evaluating tied parameters.
377///
378/// `order` lists [`TiedParam`]s such that every target appears *after* all of
379/// the tied targets it (transitively) references.  Evaluating the plan in
380/// order therefore guarantees each reference resolves to an already-updated
381/// value.
382#[derive(Debug, Clone, Default)]
383pub struct TiedPlan {
384    /// Tied assignments in dependency (topological) order.
385    pub order: Vec<TiedParam>,
386}
387
388impl TiedPlan {
389    /// Build a dependency-ordered plan from `(target, expression_src)` pairs.
390    ///
391    /// Performs a topological sort over the tied-target dependency graph and
392    /// detects cycles (e.g. `a → b → a`).
393    ///
394    /// # Errors
395    /// - [`CoreError::Eval`] if an expression fails to parse.
396    /// - [`CoreError::Eval`] if a target is assigned more than once.
397    /// - [`CoreError::Eval`] if a dependency cycle is detected.
398    pub fn build<'a, I>(edges: I) -> Result<Self, CoreError>
399    where
400        I: IntoIterator<Item = (&'a str, &'a str)>,
401    {
402        // Parse each edge, indexing by target key.
403        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        // Topological sort via DFS.  Edges point target → dependency (a tied
418        // target depends on the tied targets it references).  We only need to
419        // order *tied* targets among themselves; references to free/fixed
420        // params are leaves and impose no ordering constraint.
421        let order = topo_sort(&parsed, &target_index)?;
422        Ok(TiedPlan { order })
423    }
424
425    /// Number of tied parameters in the plan.
426    pub fn len(&self) -> usize {
427        self.order.len()
428    }
429
430    /// Whether the plan has no tied parameters.
431    pub fn is_empty(&self) -> bool {
432        self.order.is_empty()
433    }
434
435    /// Apply the plan in dependency order, mutating `values` in place so that
436    /// each tied target is set to its evaluated expression.
437    ///
438    /// # Errors
439    /// Returns [`CoreError::Eval`] if any expression references a key that is
440    /// not yet present in `values`.
441    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
450/// DFS topological sort over tied targets.  Returns the tied params ordered so
451/// that dependencies precede dependents.  Detects cycles.
452fn 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    // Explicit stack to avoid recursion-depth limits; each frame tracks the
467    // next dependency index to visit.
468    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            // Dependencies of `node` that are themselves tied targets.
479            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                // Advance the cursor on the current frame.
488                // INVARIANT: this branch is only reachable from within the
489                // `while let Some(&(node, dep_cursor)) = stack.last()` loop,
490                // which already confirmed `stack` is non-empty.  `last_mut`
491                // returns the same element that `last` just matched.
492                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                // All dependencies visited → finalize this node.
506                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// ---------------------------------------------------------------------------
517// Tests (real, passing — parser + topo/cycle behaviour)
518// ---------------------------------------------------------------------------
519#[cfg(test)]
520mod tests {
521    use super::*;
522
523    #[test]
524    fn parse_scaled_reference() {
525        // `0.5 * g1.amplitude` → Binary(Mul, Num(0.5), Ref(g1.amplitude))
526        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        // a.p + b.p * 2  →  add(ref, mul(ref, 2))
541        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        // (a.p + b.p) * 2  →  mul(add(...), 2)
552        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        // c2 depends on c1; provide them out of order — plan must put c1 first.
582        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        // a → b → a  (a references b, b references a)
611        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}