Files
goca/evaluator.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
}