package diceexpr import ( "crypto/rand" "errors" "fmt" "math/big" "sort" "strconv" "strings" "unicode" ) type Limits struct { MaxExprLen int MaxRepeat int MaxDice int MaxSides int MaxAstDepth int MaxExplosionDepth int } func DefaultLimits() Limits { return Limits{ MaxExprLen: 200, MaxRepeat: 25, MaxDice: 300, MaxSides: 1000000, MaxAstDepth: 64, MaxExplosionDepth: 100, } } type EvalResult struct { Total int Trace *Trace } type Trace struct { // Text is a normalized display of the (non-repeat) expression. Text string // Rolls includes all dice rolls (including dropped ones) in encounter order. Rolls []RollGroup } type RollGroup struct { Count int Sides int // Dice is the per-die roll chain; each die may have multiple faces due to explosions. Dice []Die // KeepDrop is the modifier used (kh/kl/dh/dl). Empty if none. KeepDrop string KeepDropN int } type Die struct { Faces []int Kept bool } type MultiResult struct { Repeat int Items []EvalResult } func EvalMany(input string, limits Limits) (MultiResult, error) { input = strings.TrimSpace(input) if input == "" { return MultiResult{}, errors.New("empty") } if limits.MaxExprLen > 0 && len(input) > limits.MaxExprLen { return MultiResult{}, errors.New("too long") } repeat, expr := splitRepeat(input) if repeat < 1 { repeat = 1 } if limits.MaxRepeat > 0 && repeat > limits.MaxRepeat { return MultiResult{}, fmt.Errorf("repeat too large") } // Parse once, evaluate many times. p := newParser(expr, limits) ast, err := p.parse() if err != nil { return MultiResult{}, err } norm := formatExpr(ast) out := MultiResult{Repeat: repeat, Items: make([]EvalResult, 0, repeat)} for range repeat { e := evaluator{limits: limits} total, trace, err := e.eval(ast) if err != nil { return MultiResult{}, err } trace.Text = norm out.Items = append(out.Items, EvalResult{Total: total, Trace: trace}) } return out, nil } func splitRepeat(s string) (int, string) { // Accept "R#expr" where R is positive integer. // We only treat the first '#' as repeat separator if it occurs before any dice operator. // This keeps "d#" or other oddities from being interpreted as repeat. i := strings.IndexByte(s, '#') if i <= 0 { return 1, s } prefix := strings.TrimSpace(s[:i]) if prefix == "" { return 1, s } for _, r := range prefix { if !unicode.IsDigit(r) { return 1, s } } n, err := strconv.Atoi(prefix) if err != nil || n <= 0 { return 1, s } return n, strings.TrimSpace(s[i+1:]) } // ---- Formatting ---- type FormatOptions struct { BoldMinMax bool } func DefaultFormatOptions() FormatOptions { return FormatOptions{BoldMinMax: true} } func FormatResult(r MultiResult, opts FormatOptions) string { lines := make([]string, 0, len(r.Items)) for _, item := range r.Items { lines = append(lines, formatOne(item, opts)) } return strings.Join(lines, "\n") } func formatOne(item EvalResult, opts FormatOptions) string { // If there were multiple roll groups, we still show them in sequence, e.g. "[...][...] expr". parts := make([]string, 0, len(item.Trace.Rolls)) for _, g := range item.Trace.Rolls { parts = append(parts, formatRollGroup(g, opts)) } breakdown := strings.Join(parts, " ") if breakdown == "" { breakdown = "[]" } return fmt.Sprintf("%d \u2190 %s %s", item.Total, breakdown, item.Trace.Text) } func formatRollGroup(g RollGroup, opts FormatOptions) string { items := make([]string, 0, len(g.Dice)) for _, d := range g.Dice { items = append(items, formatDie(d, g.Sides, opts)) } return "[" + strings.Join(items, ", ") + "]" } func formatDie(d Die, sides int, opts FormatOptions) string { // Chain faces with '+' if exploded. faceParts := make([]string, 0, len(d.Faces)) for _, v := range d.Faces { faceParts = append(faceParts, formatFace(v, sides, opts)) } txt := strings.Join(faceParts, "+") if len(d.Faces) > 1 { // denote explosion occurred (kept compact) txt = txt + "!" } if !d.Kept { return "~~" + txt + "~~" } return txt } func formatFace(v, sides int, opts FormatOptions) string { if !opts.BoldMinMax { return strconv.Itoa(v) } if v == 1 || v == sides { return "**" + strconv.Itoa(v) + "**" } return strconv.Itoa(v) } // ---- AST / Parser / Evaluator ---- type nodeKind int const ( kindNumber nodeKind = iota kindUnary kindBinary kindDice ) type node struct { kind nodeKind // number num int // unary unOp byte unA *node // binary binOp byte binL *node binR *node // dice diceCount *node diceSides *node explode bool keepDrop string // kh/kl/dh/dl keepN *node } type parser struct { s string i int limits Limits depth int } func newParser(s string, limits Limits) *parser { return &parser{s: strings.TrimSpace(s), limits: limits} } func (p *parser) parse() (*node, error) { p.skipSpaces() n, err := p.parseExpr() if err != nil { return nil, err } p.skipSpaces() if p.i != len(p.s) { return nil, errors.New("trailing input") } return n, nil } func (p *parser) parseExpr() (*node, error) { // expr = term {( + | - ) term} left, err := p.parseTerm() if err != nil { return nil, err } for { p.skipSpaces() if p.peek() != '+' && p.peek() != '-' { return left, nil } op := p.next() right, err := p.parseTerm() if err != nil { return nil, err } left = &node{kind: kindBinary, binOp: op, binL: left, binR: right} } } func (p *parser) parseTerm() (*node, error) { // term = factor {( * | / ) factor} left, err := p.parseFactor() if err != nil { return nil, err } for { p.skipSpaces() if p.peek() != '*' && p.peek() != '/' { return left, nil } op := p.next() right, err := p.parseFactor() if err != nil { return nil, err } left = &node{kind: kindBinary, binOp: op, binL: left, binR: right} } } func (p *parser) parseFactor() (*node, error) { // factor = unary return p.parseUnary() } func (p *parser) parseUnary() (*node, error) { p.skipSpaces() if p.peek() == '+' || p.peek() == '-' { op := p.next() a, err := p.parseUnary() if err != nil { return nil, err } return &node{kind: kindUnary, unOp: op, unA: a}, nil } return p.parsePrimaryOrDice() } func (p *parser) parsePrimaryOrDice() (*node, error) { p.skipSpaces() if p.peek() == '(' { p.next() p.depth++ if p.limits.MaxAstDepth > 0 && p.depth > p.limits.MaxAstDepth { return nil, errors.New("expression too deep") } n, err := p.parseExpr() if err != nil { return nil, err } p.skipSpaces() if p.peek() != ')' { return nil, errors.New("missing )") } p.next() p.depth-- return n, nil } // Try parse leading number or dice with omitted count. start := p.i num, hasNum, err := p.tryParseInt() if err != nil { return nil, err } p.skipSpaces() if p.peekLower() == 'd' { // dice: [count] d sides [keep/drop] [explode] var countNode *node if hasNum { countNode = &node{kind: kindNumber, num: num} } else { countNode = &node{kind: kindNumber, num: 1} } p.next() // d sidesNode, err := p.parseSides() if err != nil { return nil, err } keepDrop, keepN, err := p.parseKeepDrop() if err != nil { return nil, err } explode := false p.skipSpaces() if p.peek() == '!' { explode = true p.next() } return &node{ kind: kindDice, diceCount: countNode, diceSides: sidesNode, explode: explode, keepDrop: keepDrop, keepN: keepN, }, nil } // Not a dice; if we had a number, return number; else error. if hasNum { return &node{kind: kindNumber, num: num}, nil } p.i = start return nil, errors.New("expected number, dice, or (") } func (p *parser) parseSides() (*node, error) { p.skipSpaces() if p.peek() == '%' { p.next() return &node{kind: kindNumber, num: 100}, nil } n, ok, err := p.tryParseInt() if err != nil { return nil, err } if !ok { return nil, errors.New("expected sides") } return &node{kind: kindNumber, num: n}, nil } func (p *parser) parseKeepDrop() (string, *node, error) { p.skipSpaces() // khN klN dhN dlN if p.peekLower() != 'k' && p.peekLower() != 'd' { return "", nil, nil } c1 := p.peekLower() if c1 != 'k' && c1 != 'd' { return "", nil, nil } if p.i+1 >= len(p.s) { return "", nil, nil } c2 := byte(unicode.ToLower(rune(p.s[p.i+1]))) if c2 != 'h' && c2 != 'l' { return "", nil, nil } op := string([]byte{c1, c2}) p.i += 2 n, ok, err := p.tryParseInt() if err != nil { return "", nil, err } if !ok || n <= 0 { return "", nil, errors.New("expected keep/drop count") } return op, &node{kind: kindNumber, num: n}, nil } func (p *parser) tryParseInt() (int, bool, error) { p.skipSpaces() if p.i >= len(p.s) || !unicode.IsDigit(rune(p.s[p.i])) { return 0, false, nil } start := p.i for p.i < len(p.s) && unicode.IsDigit(rune(p.s[p.i])) { p.i++ } v, err := strconv.Atoi(p.s[start:p.i]) if err != nil { return 0, false, err } return v, true, nil } func (p *parser) skipSpaces() { for p.i < len(p.s) && unicode.IsSpace(rune(p.s[p.i])) { p.i++ } } func (p *parser) peek() byte { if p.i >= len(p.s) { return 0 } return p.s[p.i] } func (p *parser) peekLower() byte { if p.i >= len(p.s) { return 0 } return byte(unicode.ToLower(rune(p.s[p.i]))) } func (p *parser) next() byte { if p.i >= len(p.s) { return 0 } b := p.s[p.i] p.i++ return b } func normalizeSpaces(s string) string { // Keep compact: remove spaces entirely, but preserve leading repeat already stripped outside. var b strings.Builder b.Grow(len(s)) for _, r := range s { if !unicode.IsSpace(r) { b.WriteRune(r) } } return b.String() } // formatExpr renders a canonical expression with explicit dice counts and // spaces around binary operators. func formatExpr(n *node) string { return formatExprPrec(n, 0) } func formatExprPrec(n *node, parentPrec int) string { switch n.kind { case kindNumber: return strconv.Itoa(n.num) case kindUnary: inner := formatExprPrec(n.unA, 3) if n.unOp == '+' { return inner } return string(n.unOp) + inner case kindDice: // count and sides are currently numbers, but we keep it generic. count := formatExprPrec(n.diceCount, 4) sides := formatExprPrec(n.diceSides, 4) out := count + "d" + sides if n.keepDrop != "" && n.keepN != nil { out += n.keepDrop + formatExprPrec(n.keepN, 4) } if n.explode { out += "!" } return out case kindBinary: prec := binPrec(n.binOp) l := formatExprPrec(n.binL, prec) r := formatExprPrec(n.binR, prec+1) out := l + " " + string(n.binOp) + " " + r if prec < parentPrec { return "(" + out + ")" } return out default: return "" } } func binPrec(op byte) int { switch op { case '+', '-': return 1 case '*', '/': return 2 default: return 0 } } type evaluator struct { limits Limits dice int } func (e *evaluator) eval(n *node) (int, *Trace, error) { t := &Trace{} v, err := e.evalNode(n, t) if err != nil { return 0, nil, err } return v, t, nil } func (e *evaluator) evalNode(n *node, t *Trace) (int, error) { switch n.kind { case kindNumber: return n.num, nil case kindUnary: v, err := e.evalNode(n.unA, t) if err != nil { return 0, err } switch n.unOp { case '+': return v, nil case '-': return -v, nil default: return 0, errors.New("bad unary op") } case kindBinary: l, err := e.evalNode(n.binL, t) if err != nil { return 0, err } r, err := e.evalNode(n.binR, t) if err != nil { return 0, err } switch n.binOp { case '+': return l + r, nil case '-': return l - r, nil case '*': return l * r, nil case '/': if r == 0 { return 0, errors.New("division by zero") } return l / r, nil default: return 0, errors.New("bad binary op") } case kindDice: return e.evalDice(n, t) default: return 0, errors.New("unknown node") } } func (e *evaluator) evalDice(n *node, t *Trace) (int, error) { count, err := e.evalNode(n.diceCount, t) if err != nil { return 0, err } sides, err := e.evalNode(n.diceSides, t) if err != nil { return 0, err } if count <= 0 || sides <= 0 { return 0, errors.New("invalid dice") } if e.limits.MaxDice > 0 && e.dice+count > e.limits.MaxDice { return 0, errors.New("too many dice") } if e.limits.MaxSides > 0 && sides > e.limits.MaxSides { return 0, errors.New("sides too large") } e.dice += count keepDrop := n.keepDrop keepN := 0 if n.keepN != nil { keepN, err = e.evalNode(n.keepN, t) if err != nil { return 0, err } } if keepDrop != "" && keepN <= 0 { return 0, errors.New("invalid keep/drop") } if keepDrop != "" && keepN > count { keepN = count } g := RollGroup{ Count: count, Sides: sides, Dice: make([]Die, 0, count), KeepDrop: keepDrop, KeepDropN: keepN, } type dieScore struct { idx int score int } scores := make([]dieScore, 0, count) // Roll base dice. for i := 0; i < count; i++ { faces, err := rollExploding(sides, n.explode, e.limits.MaxExplosionDepth) if err != nil { return 0, err } sum := 0 for _, v := range faces { sum += v } g.Dice = append(g.Dice, Die{Faces: faces, Kept: true}) scores = append(scores, dieScore{idx: i, score: sum}) } // Apply keep/drop to base dice by comparing chain sums. if keepDrop != "" { sort.SliceStable(scores, func(i, j int) bool { return scores[i].score < scores[j].score }) kept := make(map[int]bool, count) switch keepDrop { case "kh": for i := len(scores) - keepN; i < len(scores); i++ { if i >= 0 && i < len(scores) { kept[scores[i].idx] = true } } case "kl": for i := 0; i < keepN && i < len(scores); i++ { kept[scores[i].idx] = true } case "dh": // drop highest N => keep all except highest N for i := 0; i < len(scores)-keepN; i++ { if i >= 0 && i < len(scores) { kept[scores[i].idx] = true } } case "dl": // drop lowest N => keep highest count-N for i := keepN; i < len(scores); i++ { kept[scores[i].idx] = true } default: return 0, errors.New("unknown keep/drop") } for i := range g.Dice { g.Dice[i].Kept = kept[i] } } // Sum kept dice. total := 0 for _, d := range g.Dice { if !d.Kept { continue } for _, v := range d.Faces { total += v } } t.Rolls = append(t.Rolls, g) return total, nil } func rollExploding(sides int, explode bool, maxDepth int) ([]int, error) { v, err := roll1(sides) if err != nil { return nil, err } out := []int{v} if !explode { return out, nil } depth := 0 for v == sides { depth++ if maxDepth > 0 && depth > maxDepth { return nil, errors.New("explosion depth limit") } v, err = roll1(sides) if err != nil { return nil, err } out = append(out, v) } return out, nil } func roll1(sides int) (int, error) { if sides <= 0 { return 0, errors.New("bad sides") } n, err := rand.Int(rand.Reader, big.NewInt(int64(sides))) if err != nil { return 0, err } return int(n.Int64()) + 1, nil }