diff --git a/internal/commands/commands.go b/internal/commands/commands.go index bd272fc..2142325 100644 --- a/internal/commands/commands.go +++ b/internal/commands/commands.go @@ -10,6 +10,9 @@ import ( var AllCommands = []*discordgo.ApplicationCommand{ fun.Ping, + fun.Joke, + fun.Coinflip, + fun.Dice, help.Help, cmdlevel.Slvl, cmdlevel.Rank, @@ -18,8 +21,11 @@ var AllCommands = []*discordgo.ApplicationCommand{ } var handlers = map[string]func(s *discordgo.Session, i *discordgo.InteractionCreate){ - "ping": fun.PingHandler, - "help": help.HelpHandler, + "ping": fun.PingHandler, + "joke": fun.JokeHandler, + "coinflip": fun.CoinflipHandler, + "dice": fun.DiceHandler, + "help": help.HelpHandler, "slvl": cmdlevel.SlvlHandler, "rank": cmdlevel.RankHandler, "leaderboard": cmdlevel.LeaderboardHandler, diff --git a/internal/commands/fun/public/coinflip.go b/internal/commands/fun/public/coinflip.go new file mode 100644 index 0000000..c0f67f3 --- /dev/null +++ b/internal/commands/fun/public/coinflip.go @@ -0,0 +1,22 @@ +package public + +import ( + "math/rand" + "velox-bot/internal/commands/fun/shared" + + "github.com/bwmarrin/discordgo" +) + +var Coinflip = &discordgo.ApplicationCommand{ + Name: "coinflip", + Description: "Flip a coin", +} + +func CoinflipHandler(s *discordgo.Session, i *discordgo.InteractionCreate) { + num := rand.Intn(2) + if num == 0 { + shared.Respond(s, i, "Heads!") + } else { + shared.Respond(s, i, "Tails!") + } +} diff --git a/internal/commands/fun/public/dice.go b/internal/commands/fun/public/dice.go new file mode 100644 index 0000000..3015718 --- /dev/null +++ b/internal/commands/fun/public/dice.go @@ -0,0 +1,62 @@ +package public + +import ( + "strings" + + "velox-bot/internal/commands/fun/shared" + "velox-bot/internal/diceexpr" + + "github.com/bwmarrin/discordgo" +) + +var Dice = &discordgo.ApplicationCommand{ + Name: "dice", + Description: "Roll dice expressions (e.g. d20, 2d20+1, 2#d20+1)", + Options: []*discordgo.ApplicationCommandOption{ + { + Type: discordgo.ApplicationCommandOptionString, + Name: "expr", + Description: "Dice expression to roll", + Required: true, + }, + }, + Contexts: &[]discordgo.InteractionContextType{ + discordgo.InteractionContextGuild, + discordgo.InteractionContextBotDM, + discordgo.InteractionContextPrivateChannel, + }, +} + +func DiceHandler(s *discordgo.Session, i *discordgo.InteractionCreate) { + expr := "" + if opts := i.ApplicationCommandData().Options; len(opts) > 0 { + expr = opts[0].StringValue() + } + expr = strings.TrimSpace(expr) + if expr == "" { + respondDiceFailure(s, i) + return + } + + res, err := diceexpr.EvalMany(expr, diceexpr.DefaultLimits()) + if err != nil { + respondDiceFailure(s, i) + return + } + + out := diceexpr.FormatResult(res, diceexpr.DefaultFormatOptions()) + if len(out) > 1900 { + // Keep within Discord message limits; trim safely. + out = out[:1900] + "…" + } + shared.Respond(s, i, out) +} + +func respondDiceFailure(s *discordgo.Session, i *discordgo.InteractionCreate) { + if i.GuildID != "" { + shared.RespondEphemeral(s, i, "Failure!") + return + } + shared.Respond(s, i, "Failure!") +} + diff --git a/internal/commands/fun/public/joke.go b/internal/commands/fun/public/joke.go new file mode 100644 index 0000000..6bb9a47 --- /dev/null +++ b/internal/commands/fun/public/joke.go @@ -0,0 +1,153 @@ +package public + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "net/http" + "time" + + "velox-bot/internal/commands/fun/shared" + + "github.com/bwmarrin/discordgo" +) + +var Joke = &discordgo.ApplicationCommand{ + Name: "joke", + Description: "Get a random joke", + Contexts: &[]discordgo.InteractionContextType{ + discordgo.InteractionContextGuild, + discordgo.InteractionContextBotDM, + discordgo.InteractionContextPrivateChannel, + }, + Options: []*discordgo.ApplicationCommandOption{ + { + Type: discordgo.ApplicationCommandOptionString, + Name: "type", + Description: "The type of joke to get", + Required: false, + Choices: []*discordgo.ApplicationCommandOptionChoice{ + { + Name: "Programming", + Value: "Programming", + }, + { + Name: "Misc", + Value: "Miscellaneous", + }, + { + Name: "Dark", + Value: "Dark", + }, + { + Name: "Pun", + Value: "Pun", + }, + { + Name: "Spooky", + Value: "Spooky", + }, + { + Name: "Christmas", + Value: "Christmas", + }, + }, + }, + }, +} + +func JokeHandler(s *discordgo.Session, i *discordgo.InteractionCreate) { + jokeType := "Any" + if opts := i.ApplicationCommandData().Options; len(opts) > 0 { + if v := opts[0].StringValue(); v != "" { + jokeType = v + } + } + if !isValidJokeType(jokeType) { + jokeType = "Any" + } + + ctx, cancel := context.WithTimeout(context.Background(), 6*time.Second) + defer cancel() + + content, err := fetchJoke(ctx, jokeType) + if err != nil { + if i.GuildID != "" { + shared.RespondEphemeral(s, i, "Failed to get joke.") + } else { + shared.Respond(s, i, "Failed to get joke.") + } + return + } + + shared.Respond(s, i, content) +} + +func isValidJokeType(t string) bool { + switch t { + case "Any", "Programming", "Misc", "Dark", "Pun", "Spooky", "Christmas": + return true + default: + return false + } +} + +type jokeAPIResponse struct { + Error bool `json:"error"` + Message string `json:"message"` + Type string `json:"type"` + Joke string `json:"joke"` + Setup string `json:"setup"` + Delivery string `json:"delivery"` +} + +func fetchJoke(ctx context.Context, jokeType string) (string, error) { + const baseURL = "https://v2.jokeapi.dev/joke" + req, err := http.NewRequestWithContext(ctx, http.MethodGet, fmt.Sprintf("%s/%s", baseURL, jokeType), nil) + if err != nil { + return "", err + } + + resp, err := http.DefaultClient.Do(req) + if err != nil { + return "", err + } + defer resp.Body.Close() + + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + return "", fmt.Errorf("joke api status: %s", resp.Status) + } + + var data jokeAPIResponse + if err := json.NewDecoder(resp.Body).Decode(&data); err != nil { + return "", err + } + if data.Error { + if data.Message != "" { + return "", errors.New(data.Message) + } + return "", errors.New("joke api returned error") + } + + switch data.Type { + case "single": + if data.Joke == "" { + return "", errors.New("empty joke") + } + return data.Joke, nil + case "twopart": + if data.Setup == "" && data.Delivery == "" { + return "", errors.New("empty joke") + } + if data.Setup == "" { + return data.Delivery, nil + } + if data.Delivery == "" { + return data.Setup, nil + } + return data.Setup + "\n" + data.Delivery, nil + default: + return "", fmt.Errorf("unknown joke type: %q", data.Type) + } +} diff --git a/internal/commands/fun/ping.go b/internal/commands/fun/public/ping.go similarity index 96% rename from internal/commands/fun/ping.go rename to internal/commands/fun/public/ping.go index 2810ca5..f227c8d 100644 --- a/internal/commands/fun/ping.go +++ b/internal/commands/fun/public/ping.go @@ -1,4 +1,4 @@ -package fun +package public import "github.com/bwmarrin/discordgo" diff --git a/internal/commands/fun/registry.go b/internal/commands/fun/registry.go new file mode 100644 index 0000000..2808f4f --- /dev/null +++ b/internal/commands/fun/registry.go @@ -0,0 +1,23 @@ +package fun + +import ( + "velox-bot/internal/commands/fun/public" + + "github.com/bwmarrin/discordgo" +) + +// Commands +var ( + Ping *discordgo.ApplicationCommand = public.Ping + Joke *discordgo.ApplicationCommand = public.Joke + Coinflip *discordgo.ApplicationCommand = public.Coinflip + Dice *discordgo.ApplicationCommand = public.Dice +) + +// Handlers +func PingHandler(s *discordgo.Session, i *discordgo.InteractionCreate) { public.PingHandler(s, i) } +func JokeHandler(s *discordgo.Session, i *discordgo.InteractionCreate) { public.JokeHandler(s, i) } +func CoinflipHandler(s *discordgo.Session, i *discordgo.InteractionCreate) { + public.CoinflipHandler(s, i) +} +func DiceHandler(s *discordgo.Session, i *discordgo.InteractionCreate) { public.DiceHandler(s, i) } diff --git a/internal/commands/fun/shared/shared.go b/internal/commands/fun/shared/shared.go new file mode 100644 index 0000000..34be049 --- /dev/null +++ b/internal/commands/fun/shared/shared.go @@ -0,0 +1,24 @@ +package shared + +import ( + "github.com/bwmarrin/discordgo" +) + +func Respond(s *discordgo.Session, i *discordgo.InteractionCreate, msg string) { + _ = s.InteractionRespond(i.Interaction, &discordgo.InteractionResponse{ + Type: discordgo.InteractionResponseChannelMessageWithSource, + Data: &discordgo.InteractionResponseData{ + Content: msg, + }, + }) +} + +func RespondEphemeral(s *discordgo.Session, i *discordgo.InteractionCreate, msg string) { + _ = s.InteractionRespond(i.Interaction, &discordgo.InteractionResponse{ + Type: discordgo.InteractionResponseChannelMessageWithSource, + Data: &discordgo.InteractionResponseData{ + Content: msg, + Flags: discordgo.MessageFlagsEphemeral, + }, + }) +} diff --git a/internal/commands/level/registry.go b/internal/commands/level/registry.go index 7a2af42..1ba2da2 100644 --- a/internal/commands/level/registry.go +++ b/internal/commands/level/registry.go @@ -16,8 +16,11 @@ var ( ) // Handlers -func RankHandler(s *discordgo.Session, i *discordgo.InteractionCreate) { public.RankHandler(s, i) } -func LeaderboardHandler(s *discordgo.Session, i *discordgo.InteractionCreate) { public.LeaderboardHandler(s, i) } -func RewardsHandler(s *discordgo.Session, i *discordgo.InteractionCreate) { public.RewardsHandler(s, i) } -func SlvlHandler(s *discordgo.Session, i *discordgo.InteractionCreate) { config.SlvlHandler(s, i) } - +func RankHandler(s *discordgo.Session, i *discordgo.InteractionCreate) { public.RankHandler(s, i) } +func LeaderboardHandler(s *discordgo.Session, i *discordgo.InteractionCreate) { + public.LeaderboardHandler(s, i) +} +func RewardsHandler(s *discordgo.Session, i *discordgo.InteractionCreate) { + public.RewardsHandler(s, i) +} +func SlvlHandler(s *discordgo.Session, i *discordgo.InteractionCreate) { config.SlvlHandler(s, i) } diff --git a/internal/diceexpr/diceexpr.go b/internal/diceexpr/diceexpr.go new file mode 100644 index 0000000..3452520 --- /dev/null +++ b/internal/diceexpr/diceexpr.go @@ -0,0 +1,744 @@ +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 +} +