Skip to main content

dotloom_engine/
lang.rs

1//! The Dotloom expression language used by plugin model definitions.
2//!
3//! Expressions are parsed and type-checked (with physical dimensions) when a type
4//! definition is registered, then lowered to [`dotloom_constraints::Expr`] trees whose
5//! variables are *symbolic leaves* (properties, anchors of referenced entities,
6//! time-axis constants). The same compiled form is evaluated numerically for
7//! drawing and instantiated with solver variables for constraint solving, so the
8//! browser and native builds share one semantics. There is no `eval` and no
9//! callback into host code.
10//!
11//! ```text
12//! expr    := term (('+' | '-') term)*
13//! term    := unary (('*' | '/') unary)*
14//! unary   := '-' unary | postfix
15//! postfix := primary ('.' ident)*
16//! primary := number[unit] | ident | ident '(' args ')' | '(' expr ')'
17//! ```
18//!
19//! Units on literals: `mm cm m in ft deg rad ms s min h d`. Values: scalars with a
20//! dimension, and 2D vectors (lowered to component pairs).
21
22use std::collections::BTreeMap;
23
24use dotloom_constraints::{Expr, VarId};
25use dotloom_geometry::units::{Dim, Quantity};
26use thiserror::Error;
27
28/// Errors from parsing or type-checking expressions.
29#[derive(Debug, Clone, PartialEq, Eq, Error)]
30pub enum LangError {
31    /// Lexical/syntax error.
32    #[error("syntax error at {pos}: {msg}")]
33    Syntax {
34        /// Byte offset.
35        pos: usize,
36        /// Message.
37        msg: String,
38    },
39    /// Unknown identifier.
40    #[error("unknown name `{0}`")]
41    UnknownName(String),
42    /// Unknown function.
43    #[error("unsupported function `{0}`")]
44    UnknownFunction(String),
45    /// Type or dimension mismatch.
46    #[error("type error: {0}")]
47    Type(String),
48    /// Expression too large or nested too deeply.
49    #[error("expression too complex: {0}")]
50    TooComplex(String),
51}
52
53// ---------------------------------------------------------------------------
54// Lexer
55
56#[derive(Debug, Clone, PartialEq)]
57enum Tok {
58    Num(f64, Option<String>),
59    Ident(String),
60    Op(char),
61}
62
63fn lex(src: &str) -> Result<Vec<(usize, Tok)>, LangError> {
64    if src.len() > 4096 {
65        return Err(LangError::TooComplex("longer than 4096 bytes".into()));
66    }
67    let b = src.as_bytes();
68    let mut i = 0;
69    let mut out = Vec::new();
70    while i < b.len() {
71        let c = b[i] as char;
72        if c.is_ascii_whitespace() {
73            i += 1;
74        } else if c.is_ascii_digit() || (c == '.' && b.get(i + 1).is_some_and(u8::is_ascii_digit)) {
75            let start = i;
76            while i < b.len() && (b[i].is_ascii_digit() || b[i] == b'.') {
77                i += 1;
78            }
79            // Exponent: e/E followed by digits (optionally signed).
80            if i < b.len() && (b[i] == b'e' || b[i] == b'E') {
81                let mut j = i + 1;
82                if j < b.len() && (b[j] == b'+' || b[j] == b'-') {
83                    j += 1;
84                }
85                if j < b.len() && b[j].is_ascii_digit() {
86                    i = j;
87                    while i < b.len() && b[i].is_ascii_digit() {
88                        i += 1;
89                    }
90                }
91            }
92            let v: f64 =
93                src[start..i].parse().map_err(|_| LangError::Syntax { pos: start, msg: "bad number".into() })?;
94            if !v.is_finite() {
95                return Err(LangError::Syntax { pos: start, msg: "number out of range".into() });
96            }
97            let ustart = i;
98            while i < b.len() && b[i].is_ascii_alphabetic() {
99                i += 1;
100            }
101            let unit = (i > ustart).then(|| src[ustart..i].to_owned());
102            out.push((start, Tok::Num(v, unit)));
103        } else if c.is_ascii_alphabetic() || c == '_' {
104            let start = i;
105            while i < b.len() && (b[i].is_ascii_alphanumeric() || b[i] == b'_') {
106                i += 1;
107            }
108            out.push((start, Tok::Ident(src[start..i].to_owned())));
109        } else if "+-*/(),.".contains(c) {
110            out.push((i, Tok::Op(c)));
111            i += 1;
112        } else {
113            return Err(LangError::Syntax { pos: i, msg: format!("unexpected character `{c}`") });
114        }
115    }
116    Ok(out)
117}
118
119// ---------------------------------------------------------------------------
120// Parser
121
122/// Parsed expression.
123#[derive(Debug, Clone, PartialEq)]
124pub enum Ast {
125    /// Literal (canonical units).
126    Num(f64, Dim),
127    /// Name.
128    Name(String),
129    /// `a.b`.
130    Member(Box<Ast>, String),
131    /// Function call.
132    Call(String, Vec<Ast>),
133    /// Negation.
134    Neg(Box<Ast>),
135    /// Binary operation.
136    Bin(char, Box<Ast>, Box<Ast>),
137}
138
139struct Parser {
140    toks: Vec<(usize, Tok)>,
141    i: usize,
142    depth: usize,
143}
144
145const MAX_DEPTH: usize = 64;
146
147impl Parser {
148    fn peek(&self) -> Option<&Tok> {
149        self.toks.get(self.i).map(|t| &t.1)
150    }
151    fn pos(&self) -> usize {
152        self.toks.get(self.i).map_or(usize::MAX, |t| t.0)
153    }
154    fn err<T>(&self, msg: &str) -> Result<T, LangError> {
155        Err(LangError::Syntax { pos: self.pos(), msg: msg.to_owned() })
156    }
157    fn eat(&mut self, c: char) -> bool {
158        if self.peek() == Some(&Tok::Op(c)) {
159            self.i += 1;
160            true
161        } else {
162            false
163        }
164    }
165    fn enter(&mut self) -> Result<(), LangError> {
166        self.depth += 1;
167        if self.depth > MAX_DEPTH { Err(LangError::TooComplex("nesting deeper than 64".into())) } else { Ok(()) }
168    }
169    fn expr(&mut self) -> Result<Ast, LangError> {
170        self.enter()?;
171        let mut a = self.term()?;
172        loop {
173            if self.eat('+') {
174                a = Ast::Bin('+', Box::new(a), Box::new(self.term()?));
175            } else if self.eat('-') {
176                a = Ast::Bin('-', Box::new(a), Box::new(self.term()?));
177            } else {
178                break;
179            }
180        }
181        self.depth -= 1;
182        Ok(a)
183    }
184    fn term(&mut self) -> Result<Ast, LangError> {
185        let mut a = self.unary()?;
186        loop {
187            if self.eat('*') {
188                a = Ast::Bin('*', Box::new(a), Box::new(self.unary()?));
189            } else if self.eat('/') {
190                a = Ast::Bin('/', Box::new(a), Box::new(self.unary()?));
191            } else {
192                break;
193            }
194        }
195        Ok(a)
196    }
197    fn unary(&mut self) -> Result<Ast, LangError> {
198        if self.eat('-') {
199            self.enter()?;
200            let inner = self.unary()?;
201            self.depth -= 1;
202            return Ok(Ast::Neg(Box::new(inner)));
203        }
204        let mut a = self.primary()?;
205        while self.eat('.') {
206            match self.peek().cloned() {
207                Some(Tok::Ident(n)) => {
208                    self.i += 1;
209                    a = Ast::Member(Box::new(a), n);
210                }
211                _ => return self.err("expected member name after `.`"),
212            }
213        }
214        Ok(a)
215    }
216    fn primary(&mut self) -> Result<Ast, LangError> {
217        match self.peek().cloned() {
218            Some(Tok::Num(v, unit)) => {
219                self.i += 1;
220                let q = match unit {
221                    None => Quantity::scalar(v),
222                    Some(u) => Quantity::parse(&format!("{v} {u}"))
223                        .map_err(|_| LangError::Syntax { pos: self.pos(), msg: format!("unknown unit `{u}`") })?,
224                };
225                Ok(Ast::Num(q.value, q.dim))
226            }
227            Some(Tok::Ident(n)) => {
228                self.i += 1;
229                if self.eat('(') {
230                    let mut args = Vec::new();
231                    if !self.eat(')') {
232                        loop {
233                            args.push(self.expr()?);
234                            if self.eat(')') {
235                                break;
236                            }
237                            if !self.eat(',') {
238                                return self.err("expected `,` or `)`");
239                            }
240                        }
241                    }
242                    Ok(Ast::Call(n, args))
243                } else {
244                    Ok(Ast::Name(n))
245                }
246            }
247            Some(Tok::Op('(')) => {
248                self.i += 1;
249                let e = self.expr()?;
250                if !self.eat(')') {
251                    return self.err("expected `)`");
252                }
253                Ok(e)
254            }
255            _ => self.err("expected a value"),
256        }
257    }
258}
259
260/// Parse an expression.
261pub fn parse(src: &str) -> Result<Ast, LangError> {
262    let toks = lex(src)?;
263    let mut p = Parser { toks, i: 0, depth: 0 };
264    let e = p.expr()?;
265    if p.i != p.toks.len() {
266        return p.err("unexpected trailing input");
267    }
268    Ok(e)
269}
270
271// ---------------------------------------------------------------------------
272// Types, leaves and compiled values
273
274/// Value type.
275#[derive(Debug, Clone, Copy, PartialEq, Eq)]
276pub enum Ty {
277    /// Scalar with dimension.
278    Scalar(Dim),
279    /// 2D vector with dimension.
280    Vector(Dim),
281}
282
283impl Ty {
284    fn describe(self) -> String {
285        match self {
286            Self::Scalar(d) => format!("scalar[{d}]"),
287            Self::Vector(d) => format!("vector[{d}]"),
288        }
289    }
290}
291
292/// External inputs of compiled expressions.
293#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
294pub enum Leaf {
295    /// Numeric property of the entity itself.
296    Prop(String),
297    /// X/Y component of a point property of the entity itself.
298    PropPoint(String, Axis),
299    /// Anchor of a referenced entity (property `prop` holds the reference),
300    /// expressed in the referencing entity's local coordinates.
301    RefAnchor(String, String, Axis),
302    /// Numeric parameter of a referenced entity.
303    RefParam(String, String),
304    /// Document time axis: millimetres per second.
305    AxisScale,
306    /// Document time axis: origin in seconds.
307    AxisOrigin,
308}
309
310/// Vector component.
311#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
312pub enum Axis {
313    /// X.
314    X,
315    /// Y.
316    Y,
317}
318
319/// A compiled scalar or vector expression over [`Leaf`] inputs.
320#[derive(Debug, Clone, PartialEq)]
321pub struct Compiled {
322    /// Type.
323    pub ty: Ty,
324    /// One expression (scalar) or two (vector x, y). Variables index `leaves`.
325    pub parts: Vec<Expr>,
326    /// Leaf table.
327    pub leaves: Vec<Leaf>,
328}
329
330impl Compiled {
331    /// Evaluate with leaf values supplied by `f`.
332    pub fn eval(&self, f: &mut dyn FnMut(&Leaf) -> Option<f64>) -> Option<Vec<f64>> {
333        let mut vals = Vec::with_capacity(self.leaves.len());
334        for l in &self.leaves {
335            vals.push(f(l)?);
336        }
337        let out: Vec<f64> = self.parts.iter().map(|e| e.eval(&vals)).collect();
338        out.iter().all(|v| v.is_finite()).then_some(out)
339    }
340
341    /// Instantiate with leaf expressions supplied by `f` (solver variables or constants).
342    pub fn instantiate(&self, f: &mut dyn FnMut(&Leaf) -> Option<Expr>) -> Option<Vec<Expr>> {
343        let mut subs = Vec::with_capacity(self.leaves.len());
344        for l in &self.leaves {
345            subs.push(f(l)?);
346        }
347        Some(
348            self.parts
349                .iter()
350                .map(|e| e.substitute(&|v: VarId| subs.get(v.index()).cloned().unwrap_or(Expr::c(f64::NAN))))
351                .collect(),
352        )
353    }
354}
355
356/// What a name means while type-checking.
357#[derive(Debug, Clone, PartialEq)]
358pub enum Binding {
359    /// Numeric property with dimension.
360    NumberProp(Dim),
361    /// Point property.
362    PointProp,
363    /// Reference property; `target` lists the members available on the referenced
364    /// type (`None` = untyped reference, members are not allowed).
365    RefProp(Option<BTreeMap<String, Ty>>),
366    /// Already compiled helper (anchor or derived value).
367    Value(Compiled),
368}
369
370/// Names visible to an expression.
371#[derive(Debug, Clone, Default, PartialEq)]
372pub struct Scope {
373    /// Bindings.
374    pub names: BTreeMap<String, Binding>,
375    /// Whether `axis.mmPerSecond` / `axis.origin` are available.
376    pub has_axis: bool,
377}
378
379struct Lower<'a> {
380    scope: &'a Scope,
381    leaves: Vec<Leaf>,
382    nodes: usize,
383}
384
385type Val = (Ty, Vec<Expr>);
386
387fn s(d: Dim) -> Ty {
388    Ty::Scalar(d)
389}
390
391impl Lower<'_> {
392    fn leaf(&mut self, l: Leaf) -> Expr {
393        if let Some(i) = self.leaves.iter().position(|x| *x == l) {
394            return Expr::Var(VarId(u32::try_from(i).unwrap_or(u32::MAX)));
395        }
396        self.leaves.push(l);
397        Expr::Var(VarId(u32::try_from(self.leaves.len() - 1).unwrap_or(u32::MAX)))
398    }
399
400    fn import(&mut self, c: &Compiled) -> Val {
401        // Re-map the helper's leaves into this expression's leaf table.
402        let map: Vec<Expr> = c.leaves.iter().map(|l| self.leaf(l.clone())).collect();
403        let parts = c
404            .parts
405            .iter()
406            .map(|e| e.substitute(&|v: VarId| map.get(v.index()).cloned().unwrap_or(Expr::c(f64::NAN))))
407            .collect();
408        (c.ty, parts)
409    }
410
411    fn lower(&mut self, a: &Ast) -> Result<Val, LangError> {
412        self.nodes += 1;
413        if self.nodes > 10_000 {
414            return Err(LangError::TooComplex("more than 10000 nodes".into()));
415        }
416        match a {
417            Ast::Num(v, d) => Ok((s(*d), vec![Expr::c(*v)])),
418            Ast::Name(n) => self.name(n),
419            Ast::Member(base, member) => self.member(base, member),
420            Ast::Neg(x) => {
421                let (t, p) = self.lower(x)?;
422                Ok((t, p.into_iter().map(Expr::neg).collect()))
423            }
424            Ast::Bin(op, x, y) => {
425                let a = self.lower(x)?;
426                let b = self.lower(y)?;
427                self.binary(*op, a, b)
428            }
429            Ast::Call(f, args) => {
430                let vals = args.iter().map(|x| self.lower(x)).collect::<Result<Vec<_>, _>>()?;
431                self.call(f, vals)
432            }
433        }
434    }
435
436    fn name(&mut self, n: &str) -> Result<Val, LangError> {
437        if n == "pi" {
438            return Ok((s(Dim::ANGLE), vec![Expr::c(core::f64::consts::PI)]));
439        }
440        match self.scope.names.get(n) {
441            Some(Binding::NumberProp(d)) => {
442                let e = self.leaf(Leaf::Prop(n.to_owned()));
443                Ok((s(*d), vec![e]))
444            }
445            Some(Binding::PointProp) => {
446                let x = self.leaf(Leaf::PropPoint(n.to_owned(), Axis::X));
447                let y = self.leaf(Leaf::PropPoint(n.to_owned(), Axis::Y));
448                Ok((Ty::Vector(Dim::LENGTH), vec![x, y]))
449            }
450            Some(Binding::Value(c)) => {
451                let c = c.clone();
452                Ok(self.import(&c))
453            }
454            Some(Binding::RefProp(_)) => {
455                Err(LangError::Type(format!("reference `{n}` must be followed by a member, e.g. `{n}.start`")))
456            }
457            None => Err(LangError::UnknownName(n.to_owned())),
458        }
459    }
460
461    fn member(&mut self, base: &Ast, member: &str) -> Result<Val, LangError> {
462        let Ast::Name(b) = base else {
463            return Err(LangError::Type("member access is only allowed on references and `axis`".into()));
464        };
465        if b == "axis" {
466            if !self.scope.has_axis {
467                return Err(LangError::UnknownName("axis".into()));
468            }
469            return match member {
470                "mmPerSecond" => Ok((s(Dim::LENGTH.div(Dim::TIME)), vec![self.leaf(Leaf::AxisScale)])),
471                "origin" => Ok((s(Dim::TIME), vec![self.leaf(Leaf::AxisOrigin)])),
472                _ => Err(LangError::UnknownName(format!("axis.{member}"))),
473            };
474        }
475        match self.scope.names.get(b) {
476            Some(Binding::RefProp(Some(members))) => match members.get(member) {
477                Some(Ty::Vector(_)) => {
478                    let x = self.leaf(Leaf::RefAnchor(b.clone(), member.to_owned(), Axis::X));
479                    let y = self.leaf(Leaf::RefAnchor(b.clone(), member.to_owned(), Axis::Y));
480                    Ok((Ty::Vector(Dim::LENGTH), vec![x, y]))
481                }
482                Some(Ty::Scalar(d)) => {
483                    let d = *d;
484                    Ok((s(d), vec![self.leaf(Leaf::RefParam(b.clone(), member.to_owned()))]))
485                }
486                None => Err(LangError::UnknownName(format!("{b}.{member}"))),
487            },
488            Some(Binding::RefProp(None)) => {
489                Err(LangError::Type(format!("reference `{b}` has no declared target type")))
490            }
491            _ => Err(LangError::UnknownName(format!("{b}.{member}"))),
492        }
493    }
494
495    fn binary(&mut self, op: char, (ta, a): Val, (tb, b): Val) -> Result<Val, LangError> {
496        let mismatch = || LangError::Type(format!("cannot apply `{op}` to {} and {}", ta.describe(), tb.describe()));
497        match (op, ta, tb) {
498            ('+' | '-', Ty::Scalar(d1), Ty::Scalar(d2)) | ('+' | '-', Ty::Vector(d1), Ty::Vector(d2)) => {
499                if d1 != d2 {
500                    return Err(LangError::Type(format!(
501                        "cannot {} {d1} and {d2}",
502                        if op == '+' { "add" } else { "subtract" }
503                    )));
504                }
505                let parts = a
506                    .into_iter()
507                    .zip(b)
508                    .map(|(x, y)| if op == '+' { Expr::add(x, y) } else { Expr::sub(x, y) })
509                    .collect();
510                Ok((ta, parts))
511            }
512            ('*', Ty::Scalar(d1), Ty::Scalar(d2)) => Ok((s(d1.mul(d2)), vec![Expr::mul(one(a), one(b))])),
513            ('*', Ty::Vector(d1), Ty::Scalar(d2)) => {
514                let k = one(b);
515                Ok((Ty::Vector(d1.mul(d2)), a.into_iter().map(|x| Expr::mul(x, k.clone())).collect()))
516            }
517            ('*', Ty::Scalar(d1), Ty::Vector(d2)) => {
518                let k = one(a);
519                Ok((Ty::Vector(d1.mul(d2)), b.into_iter().map(|x| Expr::mul(k.clone(), x)).collect()))
520            }
521            ('/', Ty::Scalar(d1), Ty::Scalar(d2)) => Ok((s(d1.div(d2)), vec![Expr::div(one(a), one(b))])),
522            ('/', Ty::Vector(d1), Ty::Scalar(d2)) => {
523                let k = one(b);
524                Ok((Ty::Vector(d1.div(d2)), a.into_iter().map(|x| Expr::div(x, k.clone())).collect()))
525            }
526            _ => Err(mismatch()),
527        }
528    }
529
530    fn call(&mut self, f: &str, args: Vec<Val>) -> Result<Val, LangError> {
531        let arity = |n: usize| -> Result<(), LangError> {
532            if args.len() == n {
533                Ok(())
534            } else {
535                Err(LangError::Type(format!("`{f}` takes {n} argument(s), got {}", args.len())))
536            }
537        };
538        let scalar = |v: &Val| -> Result<(Dim, Expr), LangError> {
539            match v.0 {
540                Ty::Scalar(d) => Ok((d, v.1.first().cloned().unwrap_or(Expr::c(f64::NAN)))),
541                Ty::Vector(_) => Err(LangError::Type(format!("`{f}` expects a scalar"))),
542            }
543        };
544        let vector = |v: &Val| -> Result<(Dim, Expr, Expr), LangError> {
545            match (v.0, v.1.as_slice()) {
546                (Ty::Vector(d), [x, y]) => Ok((d, x.clone(), y.clone())),
547                _ => Err(LangError::Type(format!("`{f}` expects a vector"))),
548            }
549        };
550        let angle_like = |d: Dim| d == Dim::ANGLE || d == Dim::SCALAR;
551        match f {
552            "min" | "max" => {
553                arity(2)?;
554                let (a, b) = (&args[0], &args[1]);
555                if a.0 != b.0 || matches!(a.0, Ty::Vector(_)) {
556                    return Err(LangError::Type(format!("`{f}` needs two scalars of the same dimension")));
557                }
558                let (x, y) = (one(a.1.clone()), one(b.1.clone()));
559                Ok((a.0, vec![if f == "min" { Expr::min(x, y) } else { Expr::max(x, y) }]))
560            }
561            "clamp" => {
562                arity(3)?;
563                let ((d, x), (d1, lo), (d2, hi)) = (scalar(&args[0])?, scalar(&args[1])?, scalar(&args[2])?);
564                if d != d1 || d != d2 {
565                    return Err(LangError::Type("`clamp` arguments must share a dimension".into()));
566                }
567                Ok((s(d), vec![Expr::min(Expr::max(x, lo), hi)]))
568            }
569            "abs" => {
570                arity(1)?;
571                let (d, x) = scalar(&args[0])?;
572                Ok((s(d), vec![Expr::abs(x)]))
573            }
574            "sqrt" => {
575                arity(1)?;
576                let (d, x) = scalar(&args[0])?;
577                if d.length % 2 != 0 || d.angle % 2 != 0 || d.time % 2 != 0 {
578                    return Err(LangError::Type(format!("sqrt of {d} has no dimension")));
579                }
580                let half = Dim { length: d.length / 2, angle: d.angle / 2, time: d.time / 2 };
581                Ok((s(half), vec![Expr::sqrt(x)]))
582            }
583            "sin" | "cos" => {
584                arity(1)?;
585                let (d, x) = scalar(&args[0])?;
586                if !angle_like(d) {
587                    return Err(LangError::Type(format!("`{f}` expects an angle, got {d}")));
588                }
589                Ok((s(Dim::SCALAR), vec![if f == "sin" { Expr::sin(x) } else { Expr::cos(x) }]))
590            }
591            "atan2" => {
592                arity(2)?;
593                let ((d1, y), (d2, x)) = (scalar(&args[0])?, scalar(&args[1])?);
594                if d1 != d2 {
595                    return Err(LangError::Type("`atan2` arguments must share a dimension".into()));
596                }
597                Ok((s(Dim::ANGLE), vec![Expr::atan2(y, x)]))
598            }
599            "hypot" => {
600                arity(2)?;
601                let ((d1, x), (d2, y)) = (scalar(&args[0])?, scalar(&args[1])?);
602                if d1 != d2 {
603                    return Err(LangError::Type("`hypot` arguments must share a dimension".into()));
604                }
605                Ok((s(d1), vec![Expr::hypot(x, y)]))
606            }
607            "vec" => {
608                arity(2)?;
609                let ((d1, x), (d2, y)) = (scalar(&args[0])?, scalar(&args[1])?);
610                if d1 != d2 {
611                    return Err(LangError::Type("`vec` components must share a dimension".into()));
612                }
613                Ok((Ty::Vector(d1), vec![x, y]))
614            }
615            "x" | "y" => {
616                arity(1)?;
617                let (d, x, y) = vector(&args[0])?;
618                Ok((s(d), vec![if f == "x" { x } else { y }]))
619            }
620            "len" => {
621                arity(1)?;
622                let (d, x, y) = vector(&args[0])?;
623                Ok((s(d), vec![Expr::hypot(x, y)]))
624            }
625            "norm" => {
626                arity(1)?;
627                let (_, x, y) = vector(&args[0])?;
628                let l = Expr::hypot(x.clone(), y.clone());
629                Ok((Ty::Vector(Dim::SCALAR), vec![Expr::div(x, l.clone()), Expr::div(y, l)]))
630            }
631            "perp" => {
632                arity(1)?;
633                let (d, x, y) = vector(&args[0])?;
634                Ok((Ty::Vector(d), vec![Expr::neg(y), x]))
635            }
636            "dot" | "cross" => {
637                arity(2)?;
638                let ((d1, ax, ay), (d2, bx, by)) = (vector(&args[0])?, vector(&args[1])?);
639                let e = if f == "dot" {
640                    Expr::add(Expr::mul(ax, bx), Expr::mul(ay, by))
641                } else {
642                    Expr::sub(Expr::mul(ax, by), Expr::mul(ay, bx))
643                };
644                Ok((s(d1.mul(d2)), vec![e]))
645            }
646            "dist" => {
647                arity(2)?;
648                let ((d1, ax, ay), (d2, bx, by)) = (vector(&args[0])?, vector(&args[1])?);
649                if d1 != d2 {
650                    return Err(LangError::Type("`dist` arguments must share a dimension".into()));
651                }
652                Ok((s(d1), vec![Expr::hypot(Expr::sub(bx, ax), Expr::sub(by, ay))]))
653            }
654            "angle" => {
655                arity(1)?;
656                let (_, x, y) = vector(&args[0])?;
657                Ok((s(Dim::ANGLE), vec![Expr::atan2(y, x)]))
658            }
659            "rotate" => {
660                arity(2)?;
661                let (d, x, y) = vector(&args[0])?;
662                let (da, a) = scalar(&args[1])?;
663                if !angle_like(da) {
664                    return Err(LangError::Type("`rotate` expects an angle".into()));
665                }
666                let (c, sn) = (Expr::cos(a.clone()), Expr::sin(a));
667                Ok((
668                    Ty::Vector(d),
669                    vec![
670                        Expr::sub(Expr::mul(x.clone(), c.clone()), Expr::mul(y.clone(), sn.clone())),
671                        Expr::add(Expr::mul(x, sn), Expr::mul(y, c)),
672                    ],
673                ))
674            }
675            "lerp" => {
676                arity(3)?;
677                let (dt, t) = scalar(&args[2])?;
678                if dt != Dim::SCALAR {
679                    return Err(LangError::Type("`lerp` parameter must be dimensionless".into()));
680                }
681                let (a, b) = (&args[0], &args[1]);
682                if a.0 != b.0 {
683                    return Err(LangError::Type("`lerp` endpoints must share a type".into()));
684                }
685                let parts =
686                    a.1.iter()
687                        .zip(&b.1)
688                        .map(|(x, y)| Expr::add(x.clone(), Expr::mul(t.clone(), Expr::sub(y.clone(), x.clone()))))
689                        .collect();
690                Ok((a.0, parts))
691            }
692            other => Err(LangError::UnknownFunction(other.to_owned())),
693        }
694    }
695}
696
697fn one(v: Vec<Expr>) -> Expr {
698    v.into_iter().next().unwrap_or(Expr::c(f64::NAN))
699}
700
701/// Parse, type-check and lower `src` in `scope`.
702pub fn compile(src: &str, scope: &Scope) -> Result<Compiled, LangError> {
703    let ast = parse(src)?;
704    let mut l = Lower { scope, leaves: Vec::new(), nodes: 0 };
705    let (ty, parts) = l.lower(&ast)?;
706    Ok(Compiled { ty, parts, leaves: l.leaves })
707}
708
709/// Compile and require a type.
710pub fn compile_as(src: &str, scope: &Scope, want: Ty) -> Result<Compiled, LangError> {
711    let c = compile(src, scope)?;
712    if c.ty == want {
713        Ok(c)
714    } else {
715        Err(LangError::Type(format!("`{src}` is {} but {} is required", c.ty.describe(), want.describe())))
716    }
717}
718
719#[cfg(test)]
720mod tests {
721    use super::*;
722
723    fn scope() -> Scope {
724        let mut sc = Scope::default();
725        sc.names.insert("width".into(), Binding::NumberProp(Dim::LENGTH));
726        sc.names.insert("start".into(), Binding::PointProp);
727        sc.names.insert("end".into(), Binding::PointProp);
728        sc.names.insert("t0".into(), Binding::NumberProp(Dim::TIME));
729        let wall: BTreeMap<String, Ty> =
730            [("start".to_owned(), Ty::Vector(Dim::LENGTH)), ("thickness".to_owned(), Ty::Scalar(Dim::LENGTH))]
731                .into_iter()
732                .collect();
733        sc.names.insert("host".into(), Binding::RefProp(Some(wall)));
734        sc.has_axis = true;
735        sc
736    }
737
738    fn eval(src: &str) -> Vec<f64> {
739        let c = compile(src, &scope()).unwrap();
740        c.eval(&mut |l| {
741            Some(match l {
742                Leaf::Prop(n) if n == "width" => 600.0,
743                Leaf::Prop(n) if n == "t0" => 7200.0,
744                Leaf::PropPoint(n, a) => match (n.as_str(), a) {
745                    ("start", Axis::X) => 0.0,
746                    ("start", Axis::Y) => 0.0,
747                    ("end", Axis::X) => 3000.0,
748                    _ => 4000.0,
749                },
750                Leaf::RefAnchor(_, _, Axis::X) => 10.0,
751                Leaf::RefAnchor(_, _, Axis::Y) => 20.0,
752                Leaf::RefParam(_, _) => 200.0,
753                Leaf::AxisScale => 0.5,
754                Leaf::AxisOrigin => 3600.0,
755                Leaf::Prop(_) => return None,
756            })
757        })
758        .unwrap()
759    }
760
761    #[test]
762    fn arithmetic_units_and_vectors() {
763        assert_eq!(eval("width + 40cm"), vec![1000.0]);
764        assert_eq!(eval("dist(start, end)"), vec![5000.0]);
765        let n = eval("norm(end - start)");
766        assert!((n[0] - 0.6).abs() < 1e-12 && (n[1] - 0.8).abs() < 1e-12);
767        assert_eq!(eval("start + perp(norm(end - start)) * 100mm"), vec![-80.0, 60.0]);
768        assert_eq!(eval("host.start + vec(host.thickness, 0mm)"), vec![210.0, 20.0]);
769        assert_eq!(eval("(t0 - axis.origin) * axis.mmPerSecond"), vec![1800.0]);
770        let r = eval("rotate(vec(1m, 0m), 90deg)");
771        assert!(r[0].abs() < 1e-9 && (r[1] - 1000.0).abs() < 1e-9);
772        assert_eq!(eval("lerp(start, end, 0.5)"), vec![1500.0, 2000.0]);
773        assert_eq!(eval("max(width, 1m) - min(width, 1m)"), vec![400.0]);
774        assert_eq!(eval("sqrt(width * width)"), vec![600.0]);
775        assert_eq!(eval("1.5e3mm"), vec![1500.0]);
776    }
777
778    #[test]
779    fn dimension_errors() {
780        let sc = scope();
781        for bad in [
782            "width + t0",
783            "width + 90deg",
784            "start + width",
785            "sin(width)",
786            "sqrt(width)",
787            "lerp(start, end, width)",
788            "start * end",
789        ] {
790            assert!(matches!(compile(bad, &sc), Err(LangError::Type(_))), "{bad}");
791        }
792        assert!(matches!(compile("nope + 1", &sc), Err(LangError::UnknownName(_))));
793        assert!(matches!(compile("eval(1)", &sc), Err(LangError::UnknownFunction(_))));
794        assert!(matches!(compile("host", &sc), Err(LangError::Type(_))));
795        assert!(matches!(compile("host.missing", &sc), Err(LangError::UnknownName(_))));
796        assert!(matches!(compile("1 +", &sc), Err(LangError::Syntax { .. })));
797        assert!(matches!(compile("width $ 2", &sc), Err(LangError::Syntax { .. })));
798        assert!(matches!(compile("3 parsecs", &sc), Err(LangError::Syntax { .. })));
799        let deep = format!("{}1{}", "(".repeat(100), ")".repeat(100));
800        assert!(matches!(compile(&deep, &sc), Err(LangError::TooComplex(_))));
801        assert!(compile_as("width", &sc, Ty::Vector(Dim::LENGTH)).is_err());
802    }
803
804    #[test]
805    fn instantiate_matches_eval_and_has_exact_gradients() {
806        let c = compile("dist(start, host.start) * 2 + width", &scope()).unwrap();
807        // Leaves → solver variables 0..n.
808        let mut next = 0u32;
809        let mut map = BTreeMap::new();
810        let inst = c
811            .instantiate(&mut |l| {
812                let id = *map.entry(l.clone()).or_insert_with(|| {
813                    next += 1;
814                    next - 1
815                });
816                Some(Expr::Var(VarId(id)))
817            })
818            .unwrap();
819        let x: Vec<f64> = (0..next).map(|i| 1.0 + f64::from(i) * 3.0).collect();
820        let d = inst[0].eval_dual(&x);
821        assert!(d.v.is_finite());
822        assert!(!d.g.is_empty());
823    }
824}