refactored dice output

This commit is contained in:
Lilith Schier 2026-08-11 22:04:00 +02:00
parent 7b6eecf40e
commit 8e2f75d624
2 changed files with 259 additions and 78 deletions

View file

@ -1,6 +1,9 @@
use anyhow::bail;
use rand::{Rng, RngExt};
use std::io::Write;
use std::convert::Infallible;
use std::error::Error;
use std::fmt::{Write, format};
use thiserror::Error;
pub type ExpBox = Box<Expression>;
@ -35,13 +38,13 @@ pub enum CompareFragment {
Le(ExpBox),
}
const DICE_POOL_LIMIT: usize = 100_000;
const DICE_POOL_LIMIT: usize = 10_000;
impl Expression {
pub fn collect_evaluation(&self, mut rng: impl Rng) -> anyhow::Result<String> {
let mut buffer = Vec::new();
let total = self.evaluate(&mut rng, &mut buffer, 100)?;
let result_text = String::from_utf8(buffer)?;
let mut witness = DiscordMdWitness::default();
let total = self.evaluate(&mut rng, &mut witness, 100)?;
let result_text = witness.buffer;
Ok(format!("**{total}** = {result_text}"))
}
@ -56,21 +59,25 @@ impl Expression {
Ok(average)
}
fn evaluate(
fn evaluate<W, E>(
&self,
rng: &mut impl Rng,
writer: &mut impl Write,
w: &mut W,
outer_precedence: i64,
) -> anyhow::Result<f64> {
) -> anyhow::Result<f64>
where
W: Witness<Error = E>,
E: Error + Send + Sync + 'static,
{
use Expression::*;
let precedence = self.precedence();
let needs_parens = precedence > outer_precedence;
if needs_parens {
write!(writer, "(")?;
write!(w, "(")?;
}
let result = match self {
Const(x) => {
write!(writer, "{}", x)?;
write!(w, "{}", x)?;
*x
}
Dice(DiceFormula {
@ -83,15 +90,15 @@ impl Expression {
x,
}) => {
let count = if let Some(count_node) = count_node {
count_node.evaluate(rng, writer, precedence)? as usize
count_node.evaluate(rng, w, precedence)? as usize
} else {
1
};
if count > DICE_POOL_LIMIT {
bail!("Too many dice.")
};
write!(writer, "d")?;
let size = size_node.evaluate(rng, writer, precedence)? as i64;
write!(w, "d")?;
let size = size_node.evaluate(rng, w, precedence)? as i64;
if size < 1 {
bail!("Invalid die size.")
}
@ -102,13 +109,13 @@ impl Expression {
// Exploding dice are early, they add rolls
if let Some(fragments) = x {
write!(writer, "x")?;
write!(w, "x")?;
if size < 2 {
bail!("Infinite explosion.")
}
let mut comparers = Vec::with_capacity(fragments.len());
for fragment in fragments {
comparers.push(fragment.comparer(rng, writer)?);
comparers.push(fragment.comparer(rng, w)?);
}
for i in 0..DICE_POOL_LIMIT {
if i >= rolls.len() {
@ -133,26 +140,26 @@ impl Expression {
// Advantage/disadvantage roll filters
let kh = if let Some(child) = kh {
write!(writer, "kh")?;
Some(child.evaluate(rng, writer, precedence)? as usize)
write!(w, "kh")?;
Some(child.evaluate(rng, w, precedence)? as usize)
} else {
None
};
let kl = if let Some(child) = kl {
write!(writer, "kl")?;
Some(child.evaluate(rng, writer, precedence)? as usize)
write!(w, "kl")?;
Some(child.evaluate(rng, w, precedence)? as usize)
} else {
None
};
let dh = if let Some(child) = dh {
write!(writer, "dh")?;
Some(child.evaluate(rng, writer, precedence)? as usize)
write!(w, "dh")?;
Some(child.evaluate(rng, w, precedence)? as usize)
} else {
None
};
let dl = if let Some(child) = dl {
write!(writer, "dl")?;
Some(child.evaluate(rng, writer, precedence)? as usize)
write!(w, "dl")?;
Some(child.evaluate(rng, w, precedence)? as usize)
} else {
None
};
@ -168,11 +175,9 @@ impl Expression {
}
let mut result_set = Vec::with_capacity(rolls.len());
let print_all_rolls = rolls.len() <= 100;
write!(w, " ")?;
let mut w_set = w.witness_set()?;
if print_all_rolls {
write!(writer, " ‹ ")?;
}
for idx in 0..rolls.len() {
let mut skip = false;
let roll = rolls[idx];
@ -200,58 +205,50 @@ impl Expression {
if !skip {
result_set.push(roll as f64);
}
if print_all_rolls {
if skip {
write!(writer, "~~`{}`~~", roll)?;
} else if size > 2 && (roll == 1 || roll == size) {
write!(writer, "**`{}`**", roll)?;
} else {
write!(writer, "`{}`", roll)?;
}
if idx != rolls.len() - 1 {
write!(writer, " ")?;
}
}
}
if print_all_rolls {
write!(writer, " ›")?;
w_set.witness_roll(DiceRoll {
size: size as usize,
roll: roll as usize,
is_admitted: !skip,
did_proliferate: false,
is_from_proliferation: idx >= count,
})?;
}
w_set.end()?;
result_set.iter().sum()
}
Neg(inner_node) => {
write!(writer, "-")?;
let inner = inner_node.evaluate(rng, writer, precedence)?;
write!(w, "-")?;
let inner = inner_node.evaluate(rng, w, precedence)?;
-inner
}
Add(lhs_node, rhs_node) => {
let lhs = lhs_node.evaluate(rng, writer, precedence)?;
write!(writer, " + ")?;
let rhs = rhs_node.evaluate(rng, writer, precedence)?;
let lhs = lhs_node.evaluate(rng, w, precedence)?;
write!(w, " + ")?;
let rhs = rhs_node.evaluate(rng, w, precedence)?;
lhs + rhs
}
Sub(lhs_node, rhs_node) => {
let lhs = lhs_node.evaluate(rng, writer, precedence)?;
write!(writer, " - ")?;
let rhs = rhs_node.evaluate(rng, writer, precedence)?;
let lhs = lhs_node.evaluate(rng, w, precedence)?;
write!(w, " - ")?;
let rhs = rhs_node.evaluate(rng, w, precedence)?;
lhs - rhs
}
Mul(lhs_node, rhs_node) => {
let lhs = lhs_node.evaluate(rng, writer, precedence)?;
write!(writer, " × ")?;
let rhs = rhs_node.evaluate(rng, writer, precedence)?;
let lhs = lhs_node.evaluate(rng, w, precedence)?;
write!(w, " × ")?;
let rhs = rhs_node.evaluate(rng, w, precedence)?;
lhs * rhs
}
Div(lhs_node, rhs_node) => {
let lhs = lhs_node.evaluate(rng, writer, precedence)?;
write!(writer, " ÷ ")?;
let rhs = rhs_node.evaluate(rng, writer, precedence)?;
let lhs = lhs_node.evaluate(rng, w, precedence)?;
write!(w, " ÷ ")?;
let rhs = rhs_node.evaluate(rng, w, precedence)?;
lhs / rhs
}
};
if needs_parens {
write!(writer, ")")?;
write!(w, ")")?;
}
Ok(result)
}
@ -315,35 +312,39 @@ impl Expression {
}
impl CompareFragment {
fn comparer(
fn comparer<W, E>(
&self,
rng: &mut impl Rng,
writer: &mut impl Write,
) -> anyhow::Result<Box<dyn Fn(f64) -> bool>> {
witness: &mut W,
) -> anyhow::Result<Box<dyn Fn(f64) -> bool>>
where
W: Witness<Error = E>,
E: Error + Send + Sync + 'static,
{
Ok(match self {
CompareFragment::Eq(inner) => {
write!(writer, "=")?;
let comp = inner.evaluate(rng, writer, 2)?;
write!(witness, "=")?;
let comp = inner.evaluate(rng, witness, 2)?;
Box::new(move |value| value == comp)
}
CompareFragment::Gt(inner) => {
write!(writer, ">")?;
let comp = inner.evaluate(rng, writer, 2)?;
write!(witness, ">")?;
let comp = inner.evaluate(rng, witness, 2)?;
Box::new(move |value| value > comp)
}
CompareFragment::Ge(inner) => {
write!(writer, ">=")?;
let comp = inner.evaluate(rng, writer, 2)?;
write!(witness, ">=")?;
let comp = inner.evaluate(rng, witness, 2)?;
Box::new(move |value| value >= comp)
}
CompareFragment::Lt(inner) => {
write!(writer, "<")?;
let comp = inner.evaluate(rng, writer, 2)?;
write!(witness, "<")?;
let comp = inner.evaluate(rng, witness, 2)?;
Box::new(move |value| value < comp)
}
CompareFragment::Le(inner) => {
write!(writer, "<=")?;
let comp = inner.evaluate(rng, writer, 2)?;
write!(witness, "<=")?;
let comp = inner.evaluate(rng, witness, 2)?;
Box::new(move |value| value <= comp)
}
})
@ -362,3 +363,173 @@ impl Expression {
}
}
}
trait Witness {
type Ok;
type Error: Error;
type WitnessSet<'a>: WitnessSet<Ok = Self::Ok, Error = Self::Error>
where
Self: 'a;
fn witness_source_text(&mut self, text: &str) -> Result<Self::Ok, Self::Error>;
fn witness_set(&mut self) -> Result<Self::WitnessSet<'_>, Self::Error>;
fn write_fmt(&mut self, args: std::fmt::Arguments<'_>) -> Result<Self::Ok, Self::Error>
where
Self: Sized,
{
self.witness_source_text(&format(args))
}
}
trait WitnessSet {
type Ok;
type Error: Error;
fn witness_roll(&mut self, dice: DiceRoll) -> Result<Self::Ok, Self::Error>;
fn end(self) -> Result<Self::Ok, Self::Error>;
}
#[derive(Debug)]
struct DiceRoll {
size: usize,
roll: usize,
is_admitted: bool,
did_proliferate: bool,
is_from_proliferation: bool,
}
impl Default for DiceRoll {
fn default() -> Self {
Self {
size: 1,
roll: 1,
is_admitted: true,
did_proliferate: false,
is_from_proliferation: false,
}
}
}
struct DiscordMdWitness {
buffer: String,
dice_written: usize,
}
struct DiscordMdWitnessSet<'a> {
parent: &'a mut DiscordMdWitness,
index: usize,
dice_elided: bool,
}
#[derive(Error, Debug)]
enum DiscordMdWitnessError {
#[error("formatting failed")]
Format(#[from] std::fmt::Error),
}
impl Default for DiscordMdWitness {
fn default() -> Self {
Self {
buffer: String::new(),
dice_written: 0,
}
}
}
impl Witness for DiscordMdWitness {
type Ok = ();
type Error = DiscordMdWitnessError;
type WitnessSet<'a> = DiscordMdWitnessSet<'a>;
fn witness_source_text(&mut self, text: &str) -> Result<Self::Ok, Self::Error> {
self.buffer.push_str(text);
Ok(())
}
fn witness_set(&mut self) -> Result<Self::WitnessSet<'_>, Self::Error> {
write!(self.buffer, "‹ ")?;
Ok(DiscordMdWitnessSet {
parent: self,
index: 0,
dice_elided: false,
})
}
}
impl<'a> WitnessSet for DiscordMdWitnessSet<'a> {
type Ok = ();
type Error = DiscordMdWitnessError;
fn witness_roll(&mut self, dice: DiceRoll) -> Result<Self::Ok, Self::Error> {
if self.parent.dice_written >= 100 {
if !self.dice_elided {
write!(self.parent.buffer, "…")?;
self.dice_elided = true;
}
return Ok(());
}
let is_extreme = dice.size >= 4 && (dice.roll == 1 || dice.roll == dice.size);
if self.index > 0 {
write!(self.parent.buffer, " ")?;
}
if !dice.is_admitted {
write!(self.parent.buffer, "~~")?;
}
if dice.is_from_proliferation {
write!(self.parent.buffer, "*")?;
}
if is_extreme || dice.did_proliferate {
write!(self.parent.buffer, "**")?;
}
write!(self.parent.buffer, "`{}`", dice.roll)?;
if is_extreme || dice.did_proliferate {
write!(self.parent.buffer, "**")?;
}
if dice.is_from_proliferation {
write!(self.parent.buffer, "*")?;
}
if !dice.is_admitted {
write!(self.parent.buffer, "~~")?;
}
self.index += 1;
self.parent.dice_written += 1;
Ok(())
}
fn end(self) -> Result<Self::Ok, Self::Error> {
write!(self.parent.buffer, " ›")?;
Ok(())
}
}
impl Witness for std::io::Sink {
type Ok = ();
type Error = Infallible;
type WitnessSet<'a>
= &'a mut Self
where
Self: 'a;
fn witness_source_text(&mut self, _text: &str) -> Result<Self::Ok, Self::Error> {
Ok(())
}
fn witness_set(&mut self) -> Result<Self::WitnessSet<'_>, Self::Error> {
Ok(self)
}
}
impl WitnessSet for &mut std::io::Sink {
type Ok = ();
type Error = Infallible;
fn witness_roll(&mut self, _dice: DiceRoll) -> Result<Self::Ok, Self::Error> {
Ok(())
}
fn end(self) -> Result<Self::Ok, Self::Error> {
Ok(())
}
}

View file

@ -14,7 +14,7 @@ use serenity::builder::{
};
use serenity::{async_trait, prelude::*};
use std::sync::Arc;
use tracing::{info, warn};
use tracing::{debug, info, warn};
#[tokio::main]
async fn main() -> anyhow::Result<()> {
@ -46,7 +46,7 @@ impl EventHandler for Handler {
match event {
FullEvent::InteractionCreate { interaction, .. } => {
if let Interaction::Command(command) = interaction {
info!("Received command interaction: {command:#?}");
debug!("Received command interaction: {command:#?}");
let content = match command.data.name.as_str() {
"roll" => {
@ -62,14 +62,12 @@ impl EventHandler for Handler {
let data = CreateInteractionResponseMessage::new().content(content);
let builder = CreateInteractionResponse::Message(data);
if let Err(why) = command.create_response(&ctx.http, builder).await {
info!("Cannot respond to slash command: {why}");
warn!("Cannot respond to slash command: {why}");
}
}
}
}
FullEvent::Ready { data_about_bot, .. } => {
info!("{} is connected!", data_about_bot.user.name);
let global_command = Command::create_global_command(
&ctx.http,
CreateCommand::new("roll")
@ -102,7 +100,7 @@ impl EventHandler for Handler {
)
.await;
info!("I created the following global slash command: {global_command:#?}");
debug!("I created the following global slash command: {global_command:#?}");
info!("{} is connected!", data_about_bot.user.name);
}
@ -200,8 +198,20 @@ pub async fn roll(ctx: &Context, command: &CommandInteraction) -> anyhow::Result
));
}
for _ in 0..repeat {
let result = match expression.collect_evaluation(&mut rng) {
Ok(res) => res,
Err(err) => {
command
.create_followup(
&ctx.http,
CreateInteractionResponseFollowup::new().content(format!("{}", err)),
)
.await?;
return Err(err);
}
};
text_components.push(CreateContainerComponent::TextDisplay(
CreateTextDisplay::new(format!("{}\n", expression.collect_evaluation(&mut rng)?)),
CreateTextDisplay::new(format!("{}\n", result)),
));
}
if fixed_seed.is_some() {