// Package calc evaluates the arithmetic people send the bot: decimal // numbers, + - * /, unary minus and parentheses. // // The expression is parsed by go/parser and computed by go/constant, // which does exact rational arithmetic: 5 * 5/2 is exactly 12.5, and // 0.1 + 0.2 is exactly 0.3, so a result carries no binary floating // point noise until the moment it is formatted. package calc import ( "errors" "go/ast" "go/constant" "go/parser" "go/token" "math" "regexp" "strconv" "strings" ) // MaxInputLength caps an expression, in bytes, so a message cannot make // the bot do unbounded work. Every operation's cost grows with the size // of its operands, and the operands can only grow with the input. const MaxInputLength = 256 // Results of magnitude plainUpper or more are written in exponent form // (1e+21 rather than twenty-two digits), and so are fractions smaller // than plainLower (1e-07 rather than 0.0000001). const ( plainUpper = 1e21 plainLower = 1e-6 ) // Errors returned by Evaluate. The bot turns each into a reply. var ( ErrTooLong = errors.New("expression too long") ErrNotArithmetic = errors.New("not an arithmetic expression") ErrDivisionByZero = errors.New("division by zero") ErrTooLarge = errors.New("result too large") ) // decimalLiteral is the only number syntax accepted. Go's own literal // syntax is wider, and parts of it are traps for someone typing // arithmetic: 010 is octal 8, and 0x10, 1_000 and 1i are not what a // calculator user means by a number. var decimalLiteral = regexp.MustCompile( `^([0-9]+\.?[0-9]*|\.[0-9]+)([eE][+-]?[0-9]+)?$`, ) // Evaluate computes an arithmetic expression and returns its result as // text: whole numbers without a decimal point, fractions in the // shortest form that reads back as the same float64. func Evaluate(input string) (string, error) { s := strings.TrimSpace(input) if len(s) > MaxInputLength { return "", ErrTooLong } if s == "" { return "", ErrNotArithmetic } expr, err := parser.ParseExpr(s) if err != nil { return "", ErrNotArithmetic } v, err := eval(expr) if err != nil { return "", err } return format(v) } // eval walks the syntax tree, allowing only the node types and // operators of arithmetic. Anything else — identifiers, calls, strings, // shifts, comparisons — is refused, not evaluated. func eval(e ast.Expr) (constant.Value, error) { switch n := e.(type) { case *ast.BasicLit: return literal(n) case *ast.ParenExpr: return eval(n.X) case *ast.UnaryExpr: if n.Op != token.ADD && n.Op != token.SUB { return nil, ErrNotArithmetic } x, err := eval(n.X) if err != nil { return nil, err } return constant.UnaryOp(n.Op, x, 0), nil case *ast.BinaryExpr: return binary(n) default: return nil, ErrNotArithmetic } } func binary(n *ast.BinaryExpr) (constant.Value, error) { switch n.Op { //nolint:exhaustive // every other operator is refused. case token.ADD, token.SUB, token.MUL, token.QUO: default: return nil, ErrNotArithmetic } x, err := eval(n.X) if err != nil { return nil, err } y, err := eval(n.Y) if err != nil { return nil, err } // constant.BinaryOp panics on a zero divisor. if n.Op == token.QUO && constant.Sign(y) == 0 { return nil, ErrDivisionByZero } // token.QUO divides exactly, integers included: 25/2 is 12.5. v := constant.BinaryOp(x, n.Op, y) // go/constant represents an overflow to infinity as Unknown. if v.Kind() == constant.Unknown { return nil, ErrTooLarge } return v, nil } func literal(n *ast.BasicLit) (constant.Value, error) { if n.Kind != token.INT && n.Kind != token.FLOAT { return nil, ErrNotArithmetic } if !decimalLiteral.MatchString(n.Value) { return nil, ErrNotArithmetic } // Read as FLOAT whatever the token says, which makes every literal // decimal: as INT, a leading zero would make it octal. v := constant.MakeFromLiteral(n.Value, token.FLOAT, 0) // The syntax was checked above, so Unknown here means the exponent // overflowed. if v.Kind() == constant.Unknown { return nil, ErrTooLarge } return v, nil } // format writes a result for a person to read. A whole number of // ordinary size is written exactly, digit for digit; anything else goes // through float64, whose shortest round-trip form is free of the noise // (0.30000000000000004) that printing a binary fraction to a fixed // precision produces. func format(v constant.Value) (string, error) { f, _ := constant.Float64Val(v) if math.IsInf(f, 0) || math.IsNaN(f) { return "", ErrTooLarge } abs := math.Abs(f) if i := constant.ToInt(v); i.Kind() == constant.Int && abs < plainUpper { return i.ExactString(), nil } if abs >= plainUpper || abs < plainLower { return strconv.FormatFloat(f, 'g', -1, 64), nil } return strconv.FormatFloat(f, 'f', -1, 64), nil }