1use std::collections::BTreeMap;
23
24use dotloom_constraints::{Expr, VarId};
25use dotloom_geometry::units::{Dim, Quantity};
26use thiserror::Error;
27
28#[derive(Debug, Clone, PartialEq, Eq, Error)]
30pub enum LangError {
31 #[error("syntax error at {pos}: {msg}")]
33 Syntax {
34 pos: usize,
36 msg: String,
38 },
39 #[error("unknown name `{0}`")]
41 UnknownName(String),
42 #[error("unsupported function `{0}`")]
44 UnknownFunction(String),
45 #[error("type error: {0}")]
47 Type(String),
48 #[error("expression too complex: {0}")]
50 TooComplex(String),
51}
52
53#[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 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#[derive(Debug, Clone, PartialEq)]
124pub enum Ast {
125 Num(f64, Dim),
127 Name(String),
129 Member(Box<Ast>, String),
131 Call(String, Vec<Ast>),
133 Neg(Box<Ast>),
135 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
260pub 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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
276pub enum Ty {
277 Scalar(Dim),
279 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#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
294pub enum Leaf {
295 Prop(String),
297 PropPoint(String, Axis),
299 RefAnchor(String, String, Axis),
302 RefParam(String, String),
304 AxisScale,
306 AxisOrigin,
308}
309
310#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
312pub enum Axis {
313 X,
315 Y,
317}
318
319#[derive(Debug, Clone, PartialEq)]
321pub struct Compiled {
322 pub ty: Ty,
324 pub parts: Vec<Expr>,
326 pub leaves: Vec<Leaf>,
328}
329
330impl Compiled {
331 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 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#[derive(Debug, Clone, PartialEq)]
358pub enum Binding {
359 NumberProp(Dim),
361 PointProp,
363 RefProp(Option<BTreeMap<String, Ty>>),
366 Value(Compiled),
368}
369
370#[derive(Debug, Clone, Default, PartialEq)]
372pub struct Scope {
373 pub names: BTreeMap<String, Binding>,
375 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 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
701pub 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
709pub 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 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}