1use std::collections::HashMap;
12
13use super::dim::ConstDim;
14use crate::prelude::Ex;
15
16#[derive(Clone, Debug, Default)]
34pub struct DimMap {
35 map: HashMap<String, ConstDim>,
36}
37
38impl DimMap {
39 pub fn new() -> Self {
41 Self::default()
42 }
43
44 pub fn with(mut self, name: &str, dim: ConstDim) -> Self {
46 self.map.insert(name.to_string(), dim);
47 self
48 }
49
50 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 pub fn get(&self, name: &str) -> Option<&ConstDim> {
61 self.map.get(name)
62 }
63
64 pub fn insert(&mut self, name: &str, dim: ConstDim) {
66 self.map.insert(name.to_string(), dim);
67 }
68
69 pub fn len(&self) -> usize {
71 self.map.len()
72 }
73
74 pub fn is_empty(&self) -> bool {
76 self.map.is_empty()
77 }
78}
79
80impl 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 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
99pub 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 let name = format!("{}", expr);
155 if let Some(&dim) = dims.get(&name) {
156 return Ok(dim);
157 }
158 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 if !exp_dim.eq(ConstDim::DIMENSIONLESS) {
206 return Err(format!("Exponent must be dimensionless, got {}", exp_dim,));
207 }
208
209 if base_dim.eq(ConstDim::DIMENSIONLESS) {
211 return Ok(ConstDim::DIMENSIONLESS);
212 }
213
214 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 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 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 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 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 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 _ => Ok(ConstDim::DIMENSIONLESS),
295 }
296}
297
298pub fn check_dimensions(expr: &Ex, dims: &DimMap) -> Result<(), String> {
307 infer_dimension(expr, dims).map(|_| ())
308}
309
310pub 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#[cfg(test)]
330mod tests {
331 use super::*;
332
333 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 #[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 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 #[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 #[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 let exotic = ConstDim::new(3, 2, -1, 0, 0, 0, 0);
459 assert_eq!(exotic.name(), "Unknown");
460 }
461
462 #[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 #[test]
486 fn infer_add_consistent_is_ok() {
487 let ctx = crate::api::context::Context::new();
488 crate::syms!(ctx; m, a, g);
489 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 #[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 #[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}