summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorPaul Buetow <paul@buetow.org>2026-05-22 10:55:45 +0300
committerPaul Buetow <paul@buetow.org>2026-05-22 10:55:45 +0300
commite22223339ce55c11f57bae0f3f0426a57f14a1eb (patch)
tree7658f20adbf4c4a7a485d89215dd6a2e6cfef9a0
parentc97587eea0abc2e8d3e7d81e1f1c64cd1b7c0ae2 (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.go56
-rw-r--r--internal/rpn/operations.go102
-rw-r--r--internal/rpn/operations_metric.go162
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
+}