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 }