diff options
| author | Paul Buetow <paul@buetow.org> | 2026-05-22 10:55:45 +0300 |
|---|---|---|
| committer | Paul Buetow <paul@buetow.org> | 2026-05-22 10:55:45 +0300 |
| commit | e22223339ce55c11f57bae0f3f0426a57f14a1eb (patch) | |
| tree | 7658f20adbf4c4a7a485d89215dd6a2e6cfef9a0 | |
| parent | c97587eea0abc2e8d3e7d81e1f1c64cd1b7c0ae2 (diff) | |
feat(rpn): implement metric-aware arithmetic operations
Add: compatible categories required; Cool absorbs; result uses left metric.
Sub: same as Add.
Mul: cross-category inference (DataRate×Time→bits, Speed×Time→meters).
Div: cross-category inference (DataSize/Time→bps, Distance/Time→mps).
Power/FastPower: drops metric (unitless ratio).
Mod: compatible categories required.
Operations convert to base units, compute, convert back to result metric.
Uses SI prefix mode (hardcoded; IEC toggle is a later phase).
New file: operations_metric.go with metric helper functions.
10 arithmetic test cases covering cross-category inference.
| -rw-r--r-- | internal/rpn/metric_test.go | 56 | ||||
| -rw-r--r-- | internal/rpn/operations.go | 102 | ||||
| -rw-r--r-- | internal/rpn/operations_metric.go | 162 |
3 files changed, 302 insertions, 18 deletions
diff --git a/internal/rpn/metric_test.go b/internal/rpn/metric_test.go index 9eb1cf7..ab9c215 100644 --- a/internal/rpn/metric_test.go +++ b/internal/rpn/metric_test.go @@ -5,6 +5,7 @@ package rpn import ( "fmt" + "strconv" "testing" ) @@ -843,3 +844,58 @@ func TestAtPrefixIntegration(t *testing.T) { t.Errorf("stack[0].Metric() = %v, want GB", m) } } + +func TestMetricAwareArithmetic(t *testing.T) { + reg := GetMetricRegistry() + tolerance := 0.001 + + tests := []struct { + expr string + wantNum float64 + wantMet string + wantErr bool + }{ + {"100Mbps 50Mbps +", 150, "Mbps", false}, + {"1 100hr +", 100, "hr", false}, + {"3 4 +", 7, "Cool", false}, + {"100Mbps 1hr *", 360000000000, "bits", false}, + {"1000000000bits 1s /", 1000000000, "bps", false}, + {"100kmh 1hr *", 100000, "m", false}, + {"1km 1s /", 1000, "mps", false}, + {"1Gbps 1000Mbps /", 1, "Cool", false}, + {"100Mbps 2 ^", 10000, "Cool", false}, + {"100Mbps 1hr +", 0, "", true}, + } + + for _, tt := range tests { + t.Run(tt.expr, func(t *testing.T) { + vars := NewVariables() + rpn := NewRPN(vars) + result, err := rpn.ParseAndEvaluate(tt.expr) + if tt.wantErr { + if err == nil { + t.Errorf("expected error, got %q", result) + } + return + } + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + resultVal, err := strconv.ParseFloat(result, 64) + if err != nil { + t.Fatalf("failed to parse result %q: %v", result, err) + } + if resultVal < tt.wantNum-tolerance || resultVal > tt.wantNum+tolerance { + t.Errorf("result = %g, want %g (tolerance %g)", resultVal, tt.wantNum, tolerance) + } + stack := rpn.GetCurrentStack() + if len(stack) > 0 { + m := stack[0].Metric() + expected, _ := reg.Find(tt.wantMet) + if m != expected { + t.Errorf("metric = %v, want %v", m, expected) + } + } + }) + } +} diff --git a/internal/rpn/operations.go b/internal/rpn/operations.go index d47557f..a615741 100644 --- a/internal/rpn/operations.go +++ b/internal/rpn/operations.go @@ -347,12 +347,24 @@ func (o *Operations) Add(stack *Stack) error { return err } - // Use the Number interface for arithmetic - result, err := a.Add(b) + aM, bM := resolveMetric(a), resolveMetric(b) + if !categoriesCompatible(aM, bM) { + return metricError("+", aM, bM) + } + + // Convert both to base units, add, convert back to result metric + aBase, err := convertToBase(a, SI) + if err != nil { + return buildError("addition", err) + } + bBase, err := convertToBase(b, SI) if err != nil { return buildError("addition", err) } - stack.Push(result) + resultMetric := compatibleMetric(aM, bM) + resultVal := convertFromBase(aBase+bBase, resultMetric, SI) + + stack.Push(NewNumber(resultVal, o.mode, resultMetric)) return nil } @@ -363,11 +375,23 @@ func (o *Operations) Subtract(stack *Stack) error { return err } - result, err := a.Sub(b) + aM, bM := resolveMetric(a), resolveMetric(b) + if !categoriesCompatible(aM, bM) { + return metricError("-", aM, bM) + } + + aBase, err := convertToBase(a, SI) if err != nil { return buildError("subtraction", err) } - stack.Push(result) + bBase, err := convertToBase(b, SI) + if err != nil { + return buildError("subtraction", err) + } + resultMetric := compatibleMetric(aM, bM) + resultVal := convertFromBase(aBase-bBase, resultMetric, SI) + + stack.Push(NewNumber(resultVal, o.mode, resultMetric)) return nil } @@ -378,11 +402,21 @@ func (o *Operations) Multiply(stack *Stack) error { return err } - result, err := a.Mul(b) + aM, bM := resolveMetric(a), resolveMetric(b) + + // Convert both to base units, multiply, convert back to result metric + aBase, err := convertToBase(a, SI) if err != nil { return buildError("multiplication", err) } - stack.Push(result) + bBase, err := convertToBase(b, SI) + if err != nil { + return buildError("multiplication", err) + } + resultMetric := resultMetricForMul(aM, bM) + resultVal := convertFromBase(aBase*bBase, resultMetric, SI) + + stack.Push(NewNumber(resultVal, o.mode, resultMetric)) return nil } @@ -402,26 +436,41 @@ func (o *Operations) Divide(stack *Stack) error { return err } - result, err := a.Div(b) + aM, bM := resolveMetric(a), resolveMetric(b) + + aBase, err := convertToBase(a, SI) + if err != nil { + return buildError("division", err) + } + bBase, err := convertToBase(b, SI) if err != nil { return buildError("division", err) } - stack.Push(result) + resultMetric := resultMetricForDiv(aM, bM) + resultVal := convertFromBase(aBase/bBase, resultMetric, SI) + + stack.Push(NewNumber(resultVal, o.mode, resultMetric)) return nil } // Power pops two values from stack, raises first to power of second (a ^ b), and pushes result. +// Result is unitless (Cool metric). func (o *Operations) Power(stack *Stack) error { a, b, err := popTwo(stack, "^") if err != nil { return err } - result, err := a.Pow(b) + aF, err := a.Float64() + if err != nil { + return buildError("power", err) + } + bF, err := b.Float64() if err != nil { return buildError("power", err) } - stack.Push(result) + + stack.Push(NewNumber(math.Pow(aF, bF), o.mode, GetCoolMetric())) return nil } @@ -432,7 +481,6 @@ func (o *Operations) Modulo(stack *Stack) error { return err } - // Check if operands are symbols (not supported for arithmetic) if sym, ok := a.(*Symbol); ok { return fmt.Errorf("symbol %s cannot be used with modulo operator", sym.Name()) } @@ -444,11 +492,23 @@ func (o *Operations) Modulo(stack *Stack) error { return buildError("%", fmt.Errorf("modulo by zero")) } - result, err := a.Mod(b) + aM, bM := resolveMetric(a), resolveMetric(b) + if !categoriesCompatible(aM, bM) { + return metricError("%", aM, bM) + } + + aBase, err := convertToBase(a, SI) + if err != nil { + return buildError("modulo", err) + } + bBase, err := convertToBase(b, SI) if err != nil { - return buildError("%", err) + return buildError("modulo", err) } - stack.Push(result) + resultMetric := compatibleMetric(aM, bM) + resultVal := convertFromBase(math.Mod(aBase, bBase), resultMetric, SI) + + stack.Push(NewNumber(resultVal, o.mode, resultMetric)) return nil } @@ -465,7 +525,6 @@ func (o *Operations) FastPower(stack *Stack) error { return err } - // Get the integer exponent from b bVal, err := b.Float64() if err != nil { return buildError("**", fmt.Errorf("exponent must be a number: %w", err)) @@ -476,11 +535,18 @@ func (o *Operations) FastPower(stack *Stack) error { return buildError("**", fmt.Errorf("exponent must be an integer, got %v", bVal)) } - result, err := a.PowInt(exp) + aF, err := a.Float64() if err != nil { return buildError("**", err) } - stack.Push(result) + + // Result is unitless (Cool metric) + if exp == 0 { + stack.Push(NewNumber(1, o.mode, GetCoolMetric())) + return nil + } + resultVal := binaryExponentiationFloat(aF, exp) + stack.Push(NewNumber(resultVal, o.mode, GetCoolMetric())) return nil } diff --git a/internal/rpn/operations_metric.go b/internal/rpn/operations_metric.go new file mode 100644 index 0000000..d15881c --- /dev/null +++ b/internal/rpn/operations_metric.go @@ -0,0 +1,162 @@ +// SPDX-License-Identifier: MIT +// Copyright (c) 2026 Paul Buetow + +package rpn + +import "fmt" + +// resolveMetric returns the metric for a Number, defaulting to Cool if nil. +func resolveMetric(n Number) *Metric { + m := n.Metric() + if m == nil { + return GetCoolMetric() + } + return m +} + +// categoriesCompatible checks if two metrics are compatible for arithmetic. +// Cool (Universal) is compatible with anything. Same category is compatible. +func categoriesCompatible(a, b *Metric) bool { + if a.Category == Universal || b.Category == Universal { + return true + } + return a.Category == b.Category +} + +// compatibleMetric returns the resulting metric for + and - operations. +// Cool absorbs: if either is Cool, result is the other's metric (or Cool if both). +// Same category: result uses left operand's metric. +func compatibleMetric(a, b *Metric) *Metric { + if a.Category == Universal && b.Category == Universal { + return a // Cool + } + if a.Category == Universal { + return b + } + if b.Category == Universal { + return a + } + // Same category: use left operand's metric + return a +} + +// convertToBase converts a Number's value to its metric's base unit. +// Returns the converted float64 value. +func convertToBase(n Number, mode PrefixMode) (float64, error) { + m := resolveMetric(n) + val, err := n.Float64() + if err != nil { + return 0, fmt.Errorf("convertToBase: %w", err) + } + return val * m.Factor(mode), nil +} + +// convertFromBase converts a base-unit value back to the given metric. +func convertFromBase(baseVal float64, m *Metric, mode PrefixMode) float64 { + return baseVal / m.Factor(mode) +} + +// resultMetricForMul computes the resulting metric for multiplication. +// Cross-category inference rules: +// - DataRate × Time → DataSize +// - Time × DataRate → DataSize +// - Speed × Time → Distance +// - Time × Speed → Distance +// - Universal × X → X (Cool absorbs) +// - Otherwise → Cool (result is unitless product) +func resultMetricForMul(a, b *Metric) *Metric { + if a.Category == Universal { + return b + } + if b.Category == Universal { + return a + } + + // Cross-category inference + switch { + case a.Category == DataRate && b.Category == Time: + // Rate × Time = Size — use base unit of DataSize + return dataSizeBaseMetric() + case a.Category == Time && b.Category == DataRate: + return dataSizeBaseMetric() + case a.Category == Speed && b.Category == Time: + // Speed × Time = Distance — use base unit of Distance + return distanceBaseMetric() + case a.Category == Time && b.Category == Speed: + return distanceBaseMetric() + default: + // Same category or unknown combination → Cool (unitless) + if a.Category == b.Category { + return GetCoolMetric() + } + return GetCoolMetric() + } +} + +// resultMetricForDiv computes the resulting metric for division. +// Cross-category inference rules: +// - DataSize / Time → DataRate (base unit) +// - Distance / Time → Speed (base unit) +// - DataRate / DataRate → Cool (ratio) +// - Speed / Speed → Cool (ratio) +// - Universal / X → X +// - X / Universal → X +// - Otherwise → Cool +func resultMetricForDiv(a, b *Metric) *Metric { + if a.Category == Universal && b.Category == Universal { + return a + } + if b.Category == Universal { + return a + } + if a.Category == Universal { + return b + } + + // Cross-category inference + switch { + case a.Category == DataSize && b.Category == Time: + return dataRateBaseMetric() + case a.Category == Distance && b.Category == Time: + return speedBaseMetric() + case a.Category == DataRate && b.Category == Time: + // Rate / Time = Size per time² → Cool + return GetCoolMetric() + default: + // Same category ratio → Cool + if a.Category == b.Category { + return GetCoolMetric() + } + return GetCoolMetric() + } +} + +// metricError returns a descriptive error for incompatible metric operations. +func metricError(op string, a, b *Metric) error { + return fmt.Errorf("%s: incompatible metrics %s (%s) and %s (%s)", + op, a.Name, a.Category, b.Name, b.Category) +} + +// dataSizeBaseMetric returns the DataSize base unit (bits). +func dataSizeBaseMetric() *Metric { + m, _ := GetMetricRegistry().Find("bits") + return m +} + +// dataRateBaseMetric returns the DataRate base unit (bps). +func dataRateBaseMetric() *Metric { + m, _ := GetMetricRegistry().Find("bps") + return m +} + +// distanceBaseMetric returns the Distance base unit (meters). +func distanceBaseMetric() *Metric { + m, _ := GetMetricRegistry().Find("m") + return m +} + +// speedBaseMetric returns the Speed base unit (mps). +func speedBaseMetric() *Metric { + m, _ := GetMetricRegistry().Find("mps") + return m +} |
