569 lines
17 KiB
Go
569 lines
17 KiB
Go
package main
|
|
|
|
import (
|
|
"fmt"
|
|
"math"
|
|
"regexp"
|
|
"strconv"
|
|
"strings"
|
|
|
|
"github.com/shopspring/decimal"
|
|
)
|
|
|
|
var variables = make(map[string]Result)
|
|
var userFuncs = make(map[string]UserFunc)
|
|
var stringTable []string
|
|
|
|
// evaluate evaluates a full command string, potentially handling "to/in" conversions.
|
|
func evaluate(exprStr string) (Result, string, error) {
|
|
re := regexp.MustCompile(`(?i)\s+(to|in)\s+`)
|
|
loc := re.FindStringIndex(exprStr)
|
|
|
|
var targetUnit string
|
|
var exprToEval string
|
|
|
|
if loc != nil {
|
|
exprToEval = exprStr[:loc[0]]
|
|
targetUnit = strings.ToUpper(strings.TrimSpace(exprStr[loc[1]:]))
|
|
} else {
|
|
exprToEval = exprStr
|
|
}
|
|
|
|
p := NewParser(exprToEval)
|
|
ast, err := p.Parse()
|
|
if err != nil {
|
|
return Result{}, "", fmt.Errorf("parse error: %w", err)
|
|
}
|
|
|
|
res, err := evalExpr(ast, nil)
|
|
if err != nil {
|
|
return Result{}, "", fmt.Errorf("eval error: %w", err)
|
|
}
|
|
|
|
if targetUnit != "" {
|
|
resolvedUnit := targetUnit
|
|
switch targetUnit {
|
|
case "$":
|
|
resolvedUnit = "USD"
|
|
case "€":
|
|
resolvedUnit = "EUR"
|
|
case "£":
|
|
resolvedUnit = "GBP"
|
|
case "¥":
|
|
resolvedUnit = "JPY"
|
|
}
|
|
u, ok := unitRegistry[resolvedUnit]
|
|
if !ok {
|
|
return Result{}, "", fmt.Errorf("unknown unit: %s", targetUnit)
|
|
}
|
|
if res.Dim != u.Dimension && res.Dim != DimNone {
|
|
return Result{}, "", fmt.Errorf("cannot convert %s to %s", getDimName(res.Dim), targetUnit)
|
|
}
|
|
if res.Dim == DimTemp {
|
|
var convertedVal decimal.Decimal
|
|
switch targetUnit {
|
|
case "C", "CELSIUS":
|
|
convertedVal = res.Value.Sub(d(273.15))
|
|
case "F", "FAHRENHEIT":
|
|
convertedVal = res.Value.Sub(d(273.15)).Mul(d(1.8)).Add(d(32))
|
|
case "K", "KELVIN":
|
|
convertedVal = res.Value
|
|
default:
|
|
return Result{}, "", fmt.Errorf("unknown temperature unit: %s", targetUnit)
|
|
}
|
|
return Result{convertedVal, u.Dimension}, targetUnit, nil
|
|
}
|
|
return Result{res.Value.Div(u.Factor), u.Dimension}, targetUnit, nil
|
|
}
|
|
|
|
return res, "", nil
|
|
}
|
|
|
|
func parseIP(s string) (int64, error) {
|
|
parts := strings.Split(s, ".")
|
|
if len(parts) != 4 {
|
|
return 0, fmt.Errorf("invalid IP format")
|
|
}
|
|
var ip int64
|
|
for i := 0; i < 4; i++ {
|
|
val, err := strconv.ParseInt(parts[i], 10, 64)
|
|
if err != nil || val < 0 || val > 255 {
|
|
return 0, fmt.Errorf("invalid IP format")
|
|
}
|
|
ip = (ip << 8) | val
|
|
}
|
|
return ip, nil
|
|
}
|
|
|
|
func parseCIDR(s string) (int64, error) {
|
|
parts := strings.Split(s, "/")
|
|
if len(parts) != 2 {
|
|
return 0, fmt.Errorf("invalid CIDR format")
|
|
}
|
|
ip, err := parseIP(parts[0])
|
|
if err != nil {
|
|
return 0, err
|
|
}
|
|
mask, err := strconv.ParseInt(parts[1], 10, 64)
|
|
if err != nil || mask < 0 || mask > 32 {
|
|
return 0, fmt.Errorf("invalid CIDR mask")
|
|
}
|
|
return (ip << 8) | mask, nil
|
|
}
|
|
|
|
func evalExpr(node Expr, locals map[string]Result) (Result, error) {
|
|
switch n := node.(type) {
|
|
case *StringExpr:
|
|
stringTable = append(stringTable, n.Value)
|
|
idx := len(stringTable) - 1
|
|
return Result{decimal.NewFromInt(int64(idx)), DimString}, nil
|
|
|
|
case *NumberExpr:
|
|
text := n.Value
|
|
if strings.HasPrefix(strings.ToLower(text), "0x") || strings.HasPrefix(strings.ToLower(text), "0b") || strings.HasPrefix(strings.ToLower(text), "0o") {
|
|
val, err := strconv.ParseInt(text, 0, 64)
|
|
if err != nil {
|
|
return Result{}, fmt.Errorf("invalid integer format: %w", err)
|
|
}
|
|
return Result{decimal.NewFromInt(val), DimNone}, nil
|
|
}
|
|
val, err := decimal.NewFromString(text)
|
|
if err != nil {
|
|
return Result{}, fmt.Errorf("invalid number %s: %w", text, err)
|
|
}
|
|
return Result{val, DimNone}, nil
|
|
|
|
case *IdentExpr:
|
|
upper := strings.ToUpper(n.Name)
|
|
switch upper {
|
|
case "PI":
|
|
return Result{d(math.Pi), DimNone}, nil
|
|
case "E":
|
|
return Result{d(math.E), DimNone}, nil
|
|
}
|
|
if locals != nil {
|
|
if v, ok := locals[n.Name]; ok {
|
|
return v, nil
|
|
}
|
|
}
|
|
if v, ok := variables[n.Name]; ok { // variables are case-sensitive
|
|
return v, nil
|
|
}
|
|
// Currencies shorthand
|
|
switch n.Name {
|
|
case "$":
|
|
upper = "USD"
|
|
case "€":
|
|
upper = "EUR"
|
|
case "£":
|
|
upper = "GBP"
|
|
case "¥":
|
|
upper = "JPY"
|
|
}
|
|
u, ok := unitRegistry[upper]
|
|
if !ok {
|
|
return Result{}, fmt.Errorf("unknown unit or identifier: %s", n.Name)
|
|
}
|
|
return Result{u.Factor, u.Dimension}, nil
|
|
|
|
case *UnaryExpr:
|
|
res, err := evalExpr(n.Expr, locals)
|
|
if err != nil {
|
|
return Result{}, err
|
|
}
|
|
if n.Op == "-" {
|
|
return Result{res.Value.Neg(), res.Dim}, nil
|
|
} else if n.Op == "~" {
|
|
val := ^res.Value.IntPart()
|
|
return Result{decimal.NewFromInt(val), DimNone}, nil
|
|
} else if n.Op == "!" {
|
|
if res.Value.IsZero() {
|
|
return Result{decimal.NewFromInt(1), DimNone}, nil
|
|
}
|
|
return Result{decimal.Zero, DimNone}, nil
|
|
}
|
|
return Result{}, fmt.Errorf("unknown unary op: %s", n.Op)
|
|
|
|
case *BinaryExpr:
|
|
left, err := evalExpr(n.Left, locals)
|
|
if err != nil {
|
|
return Result{}, err
|
|
}
|
|
|
|
// Handle short-circuiting or specific ops
|
|
// Right is evaluated eagerly here for most ops
|
|
right, err := evalExpr(n.Right, locals)
|
|
if err != nil {
|
|
return Result{}, err
|
|
}
|
|
|
|
if n.Op == "implicit_mult" {
|
|
// Check if right is a temperature unit
|
|
isRightTemp := false
|
|
if right.Dim == DimTemp {
|
|
if id, ok := n.Right.(*IdentExpr); ok {
|
|
upper := strings.ToUpper(id.Name)
|
|
switch upper {
|
|
case "C", "CELSIUS":
|
|
val := left.Value.Add(d(273.15))
|
|
return Result{val, DimTemp}, nil
|
|
case "F", "FAHRENHEIT":
|
|
val := left.Value.Sub(d(32)).Div(d(1.8)).Add(d(273.15))
|
|
return Result{val, DimTemp}, nil
|
|
}
|
|
isRightTemp = true
|
|
}
|
|
}
|
|
if isRightTemp {
|
|
return Result{left.Value.Mul(right.Value), DimTemp}, nil // K
|
|
}
|
|
|
|
if left.Dim != DimNone && right.Dim != DimNone {
|
|
return Result{}, fmt.Errorf("cannot multiply two units")
|
|
}
|
|
dim := left.Dim
|
|
if dim == DimNone {
|
|
dim = right.Dim
|
|
}
|
|
return Result{left.Value.Mul(right.Value), dim}, nil
|
|
}
|
|
|
|
switch n.Op {
|
|
case "==", "!=", "<", ">", "<=", ">=":
|
|
if left.Dim != right.Dim && left.Dim != DimNone && right.Dim != DimNone {
|
|
return Result{}, fmt.Errorf("dimension mismatch in comparison")
|
|
}
|
|
cmp := left.Value.Cmp(right.Value)
|
|
var isTrue bool
|
|
switch n.Op {
|
|
case "==": isTrue = cmp == 0
|
|
case "!=": isTrue = cmp != 0
|
|
case "<": isTrue = cmp < 0
|
|
case ">": isTrue = cmp > 0
|
|
case "<=": isTrue = cmp <= 0
|
|
case ">=": isTrue = cmp >= 0
|
|
}
|
|
if isTrue {
|
|
return Result{decimal.NewFromInt(1), DimNone}, nil
|
|
}
|
|
return Result{decimal.Zero, DimNone}, nil
|
|
case "+", "-":
|
|
if left.Dim != right.Dim && left.Dim != DimNone && right.Dim != DimNone {
|
|
return Result{}, fmt.Errorf("dimension mismatch: cannot add %s and %s", getDimName(left.Dim), getDimName(right.Dim))
|
|
}
|
|
dim := left.Dim
|
|
if dim == DimNone {
|
|
dim = right.Dim
|
|
}
|
|
if n.Op == "+" {
|
|
return Result{left.Value.Add(right.Value), dim}, nil
|
|
}
|
|
return Result{left.Value.Sub(right.Value), dim}, nil
|
|
case "*":
|
|
if left.Dim != DimNone && right.Dim != DimNone {
|
|
return Result{}, fmt.Errorf("cannot multiply two units")
|
|
}
|
|
dim := left.Dim
|
|
if dim == DimNone {
|
|
dim = right.Dim
|
|
}
|
|
return Result{left.Value.Mul(right.Value), dim}, nil
|
|
case "/":
|
|
if right.Value.IsZero() {
|
|
return Result{}, fmt.Errorf("division by zero")
|
|
}
|
|
if right.Dim != DimNone {
|
|
return Result{}, fmt.Errorf("cannot divide by a unit")
|
|
}
|
|
return Result{left.Value.Div(right.Value), left.Dim}, nil
|
|
case "%":
|
|
if right.Value.IsZero() {
|
|
return Result{}, fmt.Errorf("division by zero")
|
|
}
|
|
if left.Dim != DimNone || right.Dim != DimNone {
|
|
return Result{}, fmt.Errorf("modulo operator not supported on units")
|
|
}
|
|
return Result{left.Value.Mod(right.Value), DimNone}, nil
|
|
case "**":
|
|
f1, _ := left.Value.Float64()
|
|
f2, _ := right.Value.Float64()
|
|
return Result{d(math.Pow(f1, f2)), DimNone}, nil
|
|
case "|":
|
|
return Result{decimal.NewFromInt(left.Value.IntPart() | right.Value.IntPart()), DimNone}, nil
|
|
case "^":
|
|
return Result{decimal.NewFromInt(left.Value.IntPart() ^ right.Value.IntPart()), DimNone}, nil
|
|
case "&":
|
|
return Result{decimal.NewFromInt(left.Value.IntPart() & right.Value.IntPart()), DimNone}, nil
|
|
case "<<":
|
|
shift := right.Value.IntPart()
|
|
if shift < 0 {
|
|
return Result{}, fmt.Errorf("negative shift count")
|
|
}
|
|
if shift >= 64 {
|
|
return Result{decimal.Zero, DimNone}, nil
|
|
}
|
|
return Result{decimal.NewFromInt(left.Value.IntPart() << uint64(shift)), DimNone}, nil
|
|
case ">>":
|
|
shift := right.Value.IntPart()
|
|
if shift < 0 {
|
|
return Result{}, fmt.Errorf("negative shift count")
|
|
}
|
|
if shift >= 64 {
|
|
return Result{decimal.Zero, DimNone}, nil
|
|
}
|
|
return Result{decimal.NewFromInt(left.Value.IntPart() >> uint64(shift)), DimNone}, nil
|
|
}
|
|
|
|
case *CallExpr:
|
|
fName := strings.ToLower(n.Func)
|
|
if fName == "if" {
|
|
if len(n.Args) != 3 {
|
|
return Result{}, fmt.Errorf("if expects 3 arguments: condition, true_val, false_val")
|
|
}
|
|
cond, err := evalExpr(n.Args[0], locals)
|
|
if err != nil {
|
|
return Result{}, err
|
|
}
|
|
if !cond.Value.IsZero() {
|
|
return evalExpr(n.Args[1], locals)
|
|
}
|
|
return evalExpr(n.Args[2], locals)
|
|
}
|
|
|
|
args := make([]Result, len(n.Args))
|
|
for i, a := range n.Args {
|
|
res, err := evalExpr(a, locals)
|
|
if err != nil {
|
|
return Result{}, err
|
|
}
|
|
args[i] = res
|
|
}
|
|
|
|
// Helper to require exactly N arguments
|
|
requireArgs := func(count int) error {
|
|
if len(args) != count {
|
|
return fmt.Errorf("%s expects %d argument(s)", fName, count)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
switch fName {
|
|
case "ip":
|
|
if err := requireArgs(1); err != nil { return Result{}, err }
|
|
if args[0].Dim != DimString { return Result{}, fmt.Errorf("ip expects a string") }
|
|
str := stringTable[args[0].Value.IntPart()]
|
|
val, err := parseIP(str)
|
|
if err != nil { return Result{}, err }
|
|
return Result{decimal.NewFromInt(val), DimIP}, nil
|
|
|
|
case "cidr":
|
|
if err := requireArgs(1); err != nil { return Result{}, err }
|
|
if args[0].Dim != DimString { return Result{}, fmt.Errorf("cidr expects a string") }
|
|
str := stringTable[args[0].Value.IntPart()]
|
|
val, err := parseCIDR(str)
|
|
if err != nil { return Result{}, err }
|
|
return Result{decimal.NewFromInt(val), DimCIDR}, nil
|
|
|
|
case "network":
|
|
if err := requireArgs(1); err != nil { return Result{}, err }
|
|
if args[0].Dim != DimCIDR { return Result{}, fmt.Errorf("network expects a CIDR") }
|
|
val := args[0].Value.IntPart()
|
|
ip := val >> 8
|
|
mask := val & 0xFF
|
|
shift := 32 - mask
|
|
network := (ip >> shift) << shift
|
|
return Result{decimal.NewFromInt(network), DimIP}, nil
|
|
|
|
case "broadcast":
|
|
if err := requireArgs(1); err != nil { return Result{}, err }
|
|
if args[0].Dim != DimCIDR { return Result{}, fmt.Errorf("broadcast expects a CIDR") }
|
|
val := args[0].Value.IntPart()
|
|
ip := val >> 8
|
|
mask := val & 0xFF
|
|
shift := 32 - mask
|
|
broadcast := ip | ((1 << shift) - 1)
|
|
return Result{decimal.NewFromInt(broadcast), DimIP}, nil
|
|
|
|
case "mask":
|
|
if err := requireArgs(1); err != nil { return Result{}, err }
|
|
if args[0].Dim != DimCIDR { return Result{}, fmt.Errorf("mask expects a CIDR") }
|
|
val := args[0].Value.IntPart()
|
|
mask := val & 0xFF
|
|
shift := 32 - mask
|
|
maskIp := ((1 << mask) - 1) << shift
|
|
return Result{decimal.NewFromInt(int64(maskIp)), DimIP}, nil
|
|
|
|
case "hosts":
|
|
if err := requireArgs(1); err != nil { return Result{}, err }
|
|
if args[0].Dim != DimCIDR { return Result{}, fmt.Errorf("hosts expects a CIDR") }
|
|
val := args[0].Value.IntPart()
|
|
mask := val & 0xFF
|
|
if mask >= 31 {
|
|
return Result{decimal.Zero, DimNone}, nil
|
|
}
|
|
hosts := (1 << (32 - mask)) - 2
|
|
return Result{decimal.NewFromInt(int64(hosts)), DimNone}, nil
|
|
|
|
case "range":
|
|
if err := requireArgs(1); err != nil { return Result{}, err }
|
|
if args[0].Dim != DimCIDR { return Result{}, fmt.Errorf("range expects a CIDR") }
|
|
val := args[0].Value.IntPart()
|
|
ip := val >> 8
|
|
mask := val & 0xFF
|
|
shift := 32 - mask
|
|
network := (ip >> shift) << shift
|
|
broadcast := ip | ((1 << shift) - 1)
|
|
|
|
var first, last int64
|
|
if mask >= 31 {
|
|
first = network
|
|
last = broadcast
|
|
} else {
|
|
first = network + 1
|
|
last = broadcast - 1
|
|
}
|
|
|
|
formatIP := func(i int64) string {
|
|
return fmt.Sprintf("%d.%d.%d.%d", (i>>24)&0xFF, (i>>16)&0xFF, (i>>8)&0xFF, i&0xFF)
|
|
}
|
|
|
|
str := fmt.Sprintf("%s - %s", formatIP(first), formatIP(last))
|
|
stringTable = append(stringTable, str)
|
|
return Result{decimal.NewFromInt(int64(len(stringTable) - 1)), DimString}, nil
|
|
|
|
case "sin", "cos", "tan":
|
|
if err := requireArgs(1); err != nil {
|
|
return Result{}, err
|
|
}
|
|
if args[0].Dim != DimNone && args[0].Dim != DimAngle {
|
|
return Result{}, fmt.Errorf("%s expects a dimensionless number or angle", fName)
|
|
}
|
|
val, _ := args[0].Value.Float64()
|
|
var res float64
|
|
if fName == "sin" { res = math.Sin(val) }
|
|
if fName == "cos" { res = math.Cos(val) }
|
|
if fName == "tan" { res = math.Tan(val) }
|
|
return Result{d(res), DimNone}, nil
|
|
|
|
case "asin", "acos", "atan":
|
|
if err := requireArgs(1); err != nil {
|
|
return Result{}, err
|
|
}
|
|
if args[0].Dim != DimNone {
|
|
return Result{}, fmt.Errorf("%s expects a dimensionless number", fName)
|
|
}
|
|
val, _ := args[0].Value.Float64()
|
|
var res float64
|
|
if fName == "asin" { res = math.Asin(val) }
|
|
if fName == "acos" { res = math.Acos(val) }
|
|
if fName == "atan" { res = math.Atan(val) }
|
|
return Result{d(res), DimAngle}, nil
|
|
|
|
case "sinh", "cosh", "tanh":
|
|
if err := requireArgs(1); err != nil {
|
|
return Result{}, err
|
|
}
|
|
if args[0].Dim != DimNone {
|
|
return Result{}, fmt.Errorf("%s expects a dimensionless number", fName)
|
|
}
|
|
val, _ := args[0].Value.Float64()
|
|
var res float64
|
|
if fName == "sinh" { res = math.Sinh(val) }
|
|
if fName == "cosh" { res = math.Cosh(val) }
|
|
if fName == "tanh" { res = math.Tanh(val) }
|
|
return Result{d(res), DimNone}, nil
|
|
|
|
case "sqrt":
|
|
if err := requireArgs(1); err != nil { return Result{}, err }
|
|
f, _ := args[0].Value.Float64()
|
|
return Result{d(math.Sqrt(f)), DimNone}, nil
|
|
case "log":
|
|
if err := requireArgs(1); err != nil { return Result{}, err }
|
|
f, _ := args[0].Value.Float64()
|
|
return Result{d(math.Log10(f)), DimNone}, nil
|
|
case "log2":
|
|
if err := requireArgs(1); err != nil { return Result{}, err }
|
|
f, _ := args[0].Value.Float64()
|
|
return Result{d(math.Log2(f)), DimNone}, nil
|
|
case "ln":
|
|
if err := requireArgs(1); err != nil { return Result{}, err }
|
|
f, _ := args[0].Value.Float64()
|
|
return Result{d(math.Log(f)), DimNone}, nil
|
|
case "exp":
|
|
if err := requireArgs(1); err != nil { return Result{}, err }
|
|
f, _ := args[0].Value.Float64()
|
|
return Result{d(math.Exp(f)), DimNone}, nil
|
|
case "abs":
|
|
if err := requireArgs(1); err != nil { return Result{}, err }
|
|
return Result{args[0].Value.Abs(), DimNone}, nil
|
|
case "ceil":
|
|
if err := requireArgs(1); err != nil { return Result{}, err }
|
|
return Result{args[0].Value.Ceil(), DimNone}, nil
|
|
case "floor":
|
|
if err := requireArgs(1); err != nil { return Result{}, err }
|
|
return Result{args[0].Value.Floor(), DimNone}, nil
|
|
case "round":
|
|
if err := requireArgs(1); err != nil { return Result{}, err }
|
|
return Result{args[0].Value.Round(0), DimNone}, nil
|
|
case "pow":
|
|
if err := requireArgs(2); err != nil { return Result{}, err }
|
|
f1, _ := args[0].Value.Float64()
|
|
f2, _ := args[1].Value.Float64()
|
|
return Result{d(math.Pow(f1, f2)), DimNone}, nil
|
|
case "min":
|
|
if len(args) < 1 { return Result{}, fmt.Errorf("min expects at least 1 argument") }
|
|
minVal := args[0].Value
|
|
for _, arg := range args[1:] {
|
|
if arg.Value.LessThan(minVal) { minVal = arg.Value }
|
|
}
|
|
return Result{minVal, DimNone}, nil
|
|
case "max":
|
|
if len(args) < 1 { return Result{}, fmt.Errorf("max expects at least 1 argument") }
|
|
maxVal := args[0].Value
|
|
for _, arg := range args[1:] {
|
|
if arg.Value.GreaterThan(maxVal) { maxVal = arg.Value }
|
|
}
|
|
return Result{maxVal, DimNone}, nil
|
|
case "mod":
|
|
if err := requireArgs(2); err != nil { return Result{}, err }
|
|
return Result{args[0].Value.Mod(args[1].Value), DimNone}, nil
|
|
case "fact":
|
|
if err := requireArgs(1); err != nil { return Result{}, err }
|
|
f, _ := args[0].Value.Float64()
|
|
if f < 0 || f != math.Trunc(f) {
|
|
return Result{}, fmt.Errorf("fact expects a non-negative integer")
|
|
}
|
|
fact := decimal.NewFromInt(1)
|
|
valInt := args[0].Value.IntPart()
|
|
for i := int64(1); i <= valInt; i++ {
|
|
fact = fact.Mul(decimal.NewFromInt(i))
|
|
}
|
|
return Result{fact, DimNone}, nil
|
|
default:
|
|
if uf, ok := userFuncs[n.Func]; ok {
|
|
if len(args) != len(uf.Args) {
|
|
return Result{}, fmt.Errorf("%s expects %d argument(s)", n.Func, len(uf.Args))
|
|
}
|
|
newLocals := make(map[string]Result)
|
|
for i, argName := range uf.Args {
|
|
newLocals[argName] = args[i]
|
|
}
|
|
p := NewParser(uf.Expr)
|
|
ast, err := p.Parse()
|
|
if err != nil {
|
|
return Result{}, fmt.Errorf("error in user function %s: %w", n.Func, err)
|
|
}
|
|
return evalExpr(ast, newLocals)
|
|
}
|
|
return Result{}, fmt.Errorf("unknown function: %s", fName)
|
|
}
|
|
}
|
|
return Result{}, fmt.Errorf("unknown AST node")
|
|
}
|
|
|
|
// evaluateExpr is a helper wrapper for tests that directly evaluate simple expressions
|
|
func evaluateExpr(expr string) (Result, error) {
|
|
res, _, err := evaluate(expr)
|
|
return res, err
|
|
}
|