Skip to main content

symplex/units/
inference.rs

1//! Runtime dimension inference for symbolic expressions.
2//!
3//! Given a mapping from variable names to physical dimensions, computes
4//! the dimension of an arbitrary expression by walking the expression tree.
5//!
6//! This is used for:
7//! - Validating `from_ex` assertions in debug builds
8//! - Debugging dimension errors interactively
9//! - Cross-checking symbolic derivations
10
11use std::collections::HashMap;
12
13use super::dim::ConstDim;
14use crate::prelude::Ex;
15
16// ---------------------------------------------------------------------------
17// DimMap — variable → dimension mapping
18// ---------------------------------------------------------------------------
19
20/// Maps symbolic variable names to their physical dimensions.
21///
22/// # Examples
23///
24/// ```
25/// use symplex::units::*;
26///
27/// let dims = DimMap::new()
28///     .with("m", ConstDim::MASS)
29///     .with("a", ConstDim::ACCELERATION)
30///     .with("g", ConstDim::ACCELERATION)
31///     .with("x", ConstDim::LENGTH);
32/// ```
33#[derive(Clone, Debug, Default)]
34pub struct DimMap {
35    map: HashMap<String, ConstDim>,
36}
37
38impl DimMap {
39    /// Create an empty dimension map.
40    pub fn new() -> Self {
41        Self::default()
42    }
43
44    /// Add a variable-dimension mapping. Chainable.
45    pub fn with(mut self, name: &str, dim: ConstDim) -> Self {
46        self.map.insert(name.to_string(), dim);
47        self
48    }
49
50    /// Add a mapping from an expression's display representation.
51    ///
52    /// Extracts the variable name from the `Ex`'s `Display` output.
53    pub fn with_var(mut self, var: &Ex, dim: ConstDim) -> Self {
54        let name = format!("{}", var);
55        self.map.insert(name, dim);
56        self
57    }
58
59    /// Look up a variable's dimension.
60    pub fn get(&self, name: &str) -> Option<&ConstDim> {
61        self.map.get(name)
62    }
63
64    /// Insert a mapping in-place (non-chaining).
65    pub fn insert(&mut self, name: &str, dim: ConstDim) {
66        self.map.insert(name.to_string(), dim);
67    }
68
69    /// Returns the number of entries.
70    pub fn len(&self) -> usize {
71        self.map.len()
72    }
73
74    /// Returns `true` if the map has no entries.
75    pub fn is_empty(&self) -> bool {
76        self.map.is_empty()
77    }
78}
79
80// ---------------------------------------------------------------------------
81// ConstDim extensions — pow and name
82// ---------------------------------------------------------------------------
83
84impl core::fmt::Display for ConstDim {
85    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
86        let n = self.name();
87        if n != "Unknown" {
88            return write!(f, "{}", n);
89        }
90        // Fall back to raw exponent notation
91        write!(
92            f,
93            "L^{} M^{} T^{} I^{} Θ^{} N^{} J^{}",
94            self.l, self.m, self.t, self.i, self.th, self.n, self.j
95        )
96    }
97}
98
99// ---------------------------------------------------------------------------
100// infer_dimension — walk an expression tree and compute its dimension
101// ---------------------------------------------------------------------------
102
103/// Infer the physical dimension of an expression.
104///
105/// Returns `Ok(dim)` if the expression is dimensionally consistent,
106/// or `Err(message)` describing the dimensional error.
107///
108/// # Dimension Rules
109///
110/// | Operation         | Rule                                                   |
111/// |-------------------|--------------------------------------------------------|
112/// | Number            | Dimensionless                                          |
113/// | Symbol            | Look up in `DimMap`                                    |
114/// | Constant (π, e)   | Dimensionless                                          |
115/// | Add(a, b)         | `a` and `b` must have same dimension                   |
116/// | Mul(a, b)         | Dimensions multiply (exponents add)                    |
117/// | Pow(base, exp)    | `exp` must be dimensionless; base dimension scaled     |
118/// | Neg(a)            | Same dimension as `a`                                  |
119/// | sin, cos, exp, ln | Argument must be dimensionless; result is dimensionless |
120/// | Derivative(f, x)  | `dim(f) / dim(x)`                                     |
121/// | Integral(f, x)    | `dim(f) × dim(x)`                                     |
122///
123/// # Examples
124///
125/// ```
126/// use symplex::prelude::*;
127/// use symplex::units::*;
128///
129/// let ctx = Context::new();
130/// let ctx = ctx.clone(); symplex::syms!(ctx; m, a);
131/// let dims = DimMap::new()
132///     .with("m", ConstDim::MASS)
133///     .with("a", ConstDim::ACCELERATION);
134///
135/// let d = infer_dimension(&(&m * &a), &dims).unwrap();
136/// assert!(d.eq(ConstDim::FORCE));
137/// ```
138pub fn infer_dimension(expr: &Ex, dims: &DimMap) -> Result<ConstDim, String> {
139    use crate::prelude::ExprType;
140
141    match expr.expr_type() {
142        ExprType::Number => Ok(ConstDim::DIMENSIONLESS),
143
144        ExprType::Symbol => {
145            let name = format!("{}", expr);
146            dims.get(&name)
147                .copied()
148                .ok_or_else(|| format!("Unknown variable '{}' — not in dimension map", name))
149        }
150
151        ExprType::Constant => {
152            // Physical constants (c, h, k_B, …) display as their symbol name.
153            // Check the DimMap first so callers can assign dimensions to them.
154            let name = format!("{}", expr);
155            if let Some(&dim) = dims.get(&name) {
156                return Ok(dim);
157            }
158            // π, e, i, ∞ — all dimensionless
159            Ok(ConstDim::DIMENSIONLESS)
160        }
161
162        ExprType::Add => {
163            let args = expr.args();
164            if args.is_empty() {
165                return Ok(ConstDim::DIMENSIONLESS);
166            }
167            let first_dim = infer_dimension(&args[0], dims)?;
168            for (i, arg) in args[1..].iter().enumerate() {
169                let arg_dim = infer_dimension(arg, dims)?;
170                if !first_dim.eq(arg_dim) {
171                    return Err(format!(
172                        "Dimension mismatch in addition: term 0 has dimension {} \
173                         but term {} has dimension {}",
174                        first_dim,
175                        i + 1,
176                        arg_dim,
177                    ));
178                }
179            }
180            Ok(first_dim)
181        }
182
183        ExprType::Mul => {
184            let args = expr.args();
185            let mut result = ConstDim::DIMENSIONLESS;
186            for arg in &args {
187                let arg_dim = infer_dimension(arg, dims)?;
188                result = result.mul(arg_dim);
189            }
190            Ok(result)
191        }
192
193        ExprType::Pow => {
194            let args = expr.args();
195            if args.len() != 2 {
196                return Err(format!(
197                    "Pow must have exactly 2 arguments, got {}",
198                    args.len()
199                ));
200            }
201            let base_dim = infer_dimension(&args[0], dims)?;
202            let exp_dim = infer_dimension(&args[1], dims)?;
203
204            // Exponent must be dimensionless
205            if !exp_dim.eq(ConstDim::DIMENSIONLESS) {
206                return Err(format!("Exponent must be dimensionless, got {}", exp_dim,));
207            }
208
209            // If the base is already dimensionless, short-circuit
210            if base_dim.eq(ConstDim::DIMENSIONLESS) {
211                return Ok(ConstDim::DIMENSIONLESS);
212            }
213
214            // Try to extract an integer exponent for dimension scaling
215            if let Ok(val) = args[1].eval_f64() {
216                let n = val.round() as i8;
217                if (val - f64::from(n)).abs() < 1e-10 {
218                    return Ok(base_dim.pow(n));
219                }
220            }
221
222            // Non-integer power of a dimensioned quantity — not allowed
223            Err(format!(
224                "Non-integer power of dimensioned quantity (base dimension: {})",
225                base_dim,
226            ))
227        }
228
229        ExprType::Neg => {
230            let args = expr.args();
231            if args.is_empty() {
232                return Ok(ConstDim::DIMENSIONLESS);
233            }
234            infer_dimension(&args[0], dims)
235        }
236
237        ExprType::Function => {
238            // sin, cos, tan, exp, ln, abs, etc.
239            // Argument(s) must be dimensionless; result is dimensionless.
240            let args = expr.args();
241            for (i, arg) in args.iter().enumerate() {
242                let dim = infer_dimension(arg, dims)?;
243                if !dim.eq(ConstDim::DIMENSIONLESS) {
244                    return Err(format!(
245                        "Function argument {} must be dimensionless, got {} \
246                         (in expression {})",
247                        i, dim, expr,
248                    ));
249                }
250            }
251            Ok(ConstDim::DIMENSIONLESS)
252        }
253
254        ExprType::Apply => {
255            // User-defined function application — treat like Function
256            let args = expr.args();
257            for (i, arg) in args.iter().enumerate() {
258                let dim = infer_dimension(arg, dims)?;
259                if !dim.eq(ConstDim::DIMENSIONLESS) {
260                    return Err(format!(
261                        "Applied function argument {} must be dimensionless, got {}",
262                        i, dim,
263                    ));
264                }
265            }
266            Ok(ConstDim::DIMENSIONLESS)
267        }
268
269        ExprType::Derivative => {
270            // d(body)/d(var) → dim(body) / dim(var)
271            let args = expr.args();
272            if args.len() >= 2 {
273                let body_dim = infer_dimension(&args[0], dims)?;
274                let var_dim = infer_dimension(&args[1], dims)?;
275                Ok(body_dim.div(var_dim))
276            } else {
277                Err("Derivative must have at least body and variable".to_string())
278            }
279        }
280
281        ExprType::Integral => {
282            // ∫ body d(var) → dim(body) × dim(var)
283            let args = expr.args();
284            if args.len() >= 2 {
285                let body_dim = infer_dimension(&args[0], dims)?;
286                let var_dim = infer_dimension(&args[1], dims)?;
287                Ok(body_dim.mul(var_dim))
288            } else {
289                Err("Integral must have at least body and variable".to_string())
290            }
291        }
292
293        // Set expressions and any future variants — treat as dimensionless
294        _ => Ok(ConstDim::DIMENSIONLESS),
295    }
296}
297
298// ---------------------------------------------------------------------------
299// Convenience wrappers
300// ---------------------------------------------------------------------------
301
302/// Check whether an expression is dimensionally consistent without
303/// returning the computed dimension.
304///
305/// Returns `Ok(())` on success, or `Err(message)` on dimensional mismatch.
306pub fn check_dimensions(expr: &Ex, dims: &DimMap) -> Result<(), String> {
307    infer_dimension(expr, dims).map(|_| ())
308}
309
310/// Infer the dimension and assert it matches the expected dimension.
311///
312/// Returns `Ok(())` if it matches, or `Err(message)` describing the mismatch.
313pub fn assert_dimension(expr: &Ex, dims: &DimMap, expected: ConstDim) -> Result<(), String> {
314    let actual = infer_dimension(expr, dims)?;
315    if actual.eq(expected) {
316        Ok(())
317    } else {
318        Err(format!(
319            "Expected dimension {} but expression has dimension {}",
320            expected, actual,
321        ))
322    }
323}
324
325// ---------------------------------------------------------------------------
326// Tests
327// ---------------------------------------------------------------------------
328
329#[cfg(test)]
330mod tests {
331    use super::*;
332
333    /// Build a standard dimension map for tests.
334    fn dims() -> DimMap {
335        DimMap::new()
336            .with("m", ConstDim::MASS)
337            .with("a", ConstDim::ACCELERATION)
338            .with("g", ConstDim::ACCELERATION)
339            .with("v", ConstDim::VELOCITY)
340            .with("t", ConstDim::TIME)
341            .with("x", ConstDim::LENGTH)
342            .with("k", ConstDim::STIFFNESS)
343            .with("F", ConstDim::FORCE)
344            .with("R", ConstDim::RESISTANCE)
345            .with("I", ConstDim::CURRENT)
346    }
347
348    // --- Basic inference ---
349
350    #[test]
351    fn infer_mass_times_accel_is_force() {
352        let ctx = crate::api::context::Context::new();
353        crate::syms!(ctx; m, a);
354        let expr = &m * &a;
355        let d = infer_dimension(&expr, &dims()).unwrap();
356        assert!(d.eq(ConstDim::FORCE), "Expected Force, got {}", d);
357    }
358
359    #[test]
360    fn infer_half_mv_squared_is_energy() {
361        let ctx = crate::api::context::Context::new();
362        crate::syms!(ctx; m, v);
363        // (1/2) * m * v^2
364        let half = ctx.int(1) / ctx.int(2);
365        let expr = &half * &m * v.powi(2);
366        let d = infer_dimension(&expr, &dims()).unwrap();
367        assert!(d.eq(ConstDim::ENERGY), "Expected Energy, got {}", d);
368    }
369
370    #[test]
371    fn infer_add_mismatch_is_error() {
372        let ctx = crate::api::context::Context::new();
373        crate::syms!(ctx; m, a);
374        let expr = &m + &a;
375        let result = infer_dimension(&expr, &dims());
376        assert!(
377            result.is_err(),
378            "Adding Mass + Acceleration should be an error"
379        );
380    }
381
382    #[test]
383    fn infer_voltage_is_current_times_resistance() {
384        let ctx = crate::api::context::Context::new();
385        let i_var = ctx.symbol("I");
386        let r_var = ctx.symbol("R");
387        let expr = &i_var * &r_var;
388        let d = infer_dimension(&expr, &dims()).unwrap();
389        assert!(d.eq(ConstDim::VOLTAGE), "Expected Voltage, got {}", d);
390    }
391
392    #[test]
393    fn infer_spring_force() {
394        let ctx = crate::api::context::Context::new();
395        crate::syms!(ctx; k, x);
396        let expr = &k * &x;
397        let d = infer_dimension(&expr, &dims()).unwrap();
398        assert!(d.eq(ConstDim::FORCE), "Expected Force, got {}", d);
399    }
400
401    #[test]
402    fn infer_pure_number_is_dimensionless() {
403        let ctx = crate::api::context::Context::new();
404        let d = infer_dimension(&ctx.int(42), &dims()).unwrap();
405        assert!(d.eq(ConstDim::DIMENSIONLESS));
406    }
407
408    #[test]
409    fn infer_unknown_variable_is_error() {
410        let ctx = crate::api::context::Context::new();
411        crate::syms!(ctx; unknown);
412        let result = infer_dimension(&unknown, &dims());
413        assert!(result.is_err(), "Unknown variable should produce an error");
414    }
415
416    // --- ConstDim::pow ---
417
418    #[test]
419    fn pow_length_squared_is_area() {
420        let d = ConstDim::LENGTH.pow(2);
421        assert!(d.eq(ConstDim::AREA));
422    }
423
424    #[test]
425    fn pow_length_cubed_is_volume() {
426        let d = ConstDim::LENGTH.pow(3);
427        assert!(d.eq(ConstDim::VOLUME));
428    }
429
430    #[test]
431    fn pow_zero_is_dimensionless() {
432        let d = ConstDim::FORCE.pow(0);
433        assert!(d.eq(ConstDim::DIMENSIONLESS));
434    }
435
436    #[test]
437    fn pow_one_is_identity() {
438        let d = ConstDim::VELOCITY.pow(1);
439        assert!(d.eq(ConstDim::VELOCITY));
440    }
441
442    // --- ConstDim::name ---
443
444    #[test]
445    fn name_known_dimensions() {
446        assert_eq!(ConstDim::FORCE.name(), "Force");
447        assert_eq!(ConstDim::ENERGY.name(), "Energy");
448        assert_eq!(ConstDim::VOLTAGE.name(), "Voltage");
449        assert_eq!(ConstDim::DIMENSIONLESS.name(), "Dimensionless");
450        assert_eq!(ConstDim::MASS.name(), "Mass");
451        assert_eq!(ConstDim::LENGTH.name(), "Length");
452        assert_eq!(ConstDim::TIME.name(), "Time");
453    }
454
455    #[test]
456    fn name_unknown_dimension() {
457        // An exotic dimension that doesn't match any named constant
458        let exotic = ConstDim::new(3, 2, -1, 0, 0, 0, 0);
459        assert_eq!(exotic.name(), "Unknown");
460    }
461
462    // --- DimMap ---
463
464    #[test]
465    fn dimmap_with_var() {
466        let ctx = crate::api::context::Context::new();
467        let x = ctx.symbol("x");
468        let dm = DimMap::new().with_var(&x, ConstDim::LENGTH);
469        assert_eq!(dm.get("x"), Some(&ConstDim::LENGTH));
470    }
471
472    #[test]
473    fn dimmap_len_and_empty() {
474        let dm = DimMap::new();
475        assert!(dm.is_empty());
476        assert_eq!(dm.len(), 0);
477
478        let dm = dm.with("x", ConstDim::LENGTH);
479        assert!(!dm.is_empty());
480        assert_eq!(dm.len(), 1);
481    }
482
483    // --- Addition dimensional consistency ---
484
485    #[test]
486    fn infer_add_consistent_is_ok() {
487        let ctx = crate::api::context::Context::new();
488        crate::syms!(ctx; m, a, g);
489        // m*a + m*g — both are Force
490        let expr = &m * &a + &m * &g;
491        let d = infer_dimension(&expr, &dims()).unwrap();
492        assert!(d.eq(ConstDim::FORCE), "Expected Force, got {}", d);
493    }
494
495    // --- Negation ---
496
497    #[test]
498    #[allow(non_snake_case)]
499    fn infer_negation_preserves_dimension() {
500        let ctx = crate::api::context::Context::new();
501        crate::syms!(ctx; F);
502        let expr = -&F;
503        let d = infer_dimension(&expr, &dims()).unwrap();
504        assert!(d.eq(ConstDim::FORCE), "Expected Force, got {}", d);
505    }
506
507    // --- assert_dimension helper ---
508
509    #[test]
510    fn assert_dimension_ok() {
511        let ctx = crate::api::context::Context::new();
512        crate::syms!(ctx; m, a);
513        let expr = &m * &a;
514        assert!(assert_dimension(&expr, &dims(), ConstDim::FORCE).is_ok());
515    }
516
517    #[test]
518    fn assert_dimension_mismatch() {
519        let ctx = crate::api::context::Context::new();
520        crate::syms!(ctx; m, a);
521        let expr = &m * &a;
522        assert!(assert_dimension(&expr, &dims(), ConstDim::ENERGY).is_err());
523    }
524}