use anyhow::bail; use rand::{Rng, RngExt}; use std::convert::Infallible; use std::error::Error; use std::fmt::{Write, format}; use thiserror::Error; pub type ExpBox = Box; #[derive(Debug, PartialEq, Clone)] pub enum Expression { Const(f64), Dice(DiceFormula), Neg(ExpBox), Add(ExpBox, ExpBox), Sub(ExpBox, ExpBox), Mul(ExpBox, ExpBox), Div(ExpBox, ExpBox), } #[derive(Debug, PartialEq, Clone)] pub struct DiceFormula { pub(crate) count: Option, pub(crate) size: ExpBox, pub(crate) kh: Option, pub(crate) kl: Option, pub(crate) dh: Option, pub(crate) dl: Option, pub(crate) x: Option>, } #[derive(Debug, PartialEq, Clone)] pub enum CompareFragment { Eq(ExpBox), Gt(ExpBox), Ge(ExpBox), Lt(ExpBox), Le(ExpBox), } const DICE_POOL_LIMIT: usize = 10_000; impl Expression { pub fn collect_evaluation(&self, mut rng: impl Rng) -> anyhow::Result { let mut witness = DiscordMdWitness::default(); let total = self.evaluate(&mut rng, &mut witness, 100)?; let result_text = witness.buffer; Ok(format!("**{total}** = {result_text}")) } pub fn average(&self, rng: &mut impl Rng) -> anyhow::Result { if let Ok(avg) = self.avg() { return Ok(avg); } let mut average = self.sample(rng)?; for idx in 1..3000 { average = (average * idx as f64 + self.sample(rng)?) / (idx as f64 + 1f64); } Ok(average) } fn evaluate( &self, rng: &mut impl Rng, w: &mut W, outer_precedence: i64, ) -> anyhow::Result where W: Witness, E: Error + Send + Sync + 'static, { use Expression::*; let precedence = self.precedence(); let needs_parens = precedence > outer_precedence; if needs_parens { write!(w, "(")?; } let result = match self { Const(x) => { write!(w, "{}", x)?; *x } Dice(DiceFormula { count: count_node, size: size_node, kh, kl, dh, dl, x, }) => { let count = if let Some(count_node) = count_node { count_node.evaluate(rng, w, precedence)? as usize } else { 1 }; if count > DICE_POOL_LIMIT { bail!("Too many dice.") }; write!(w, "d")?; let size = size_node.evaluate(rng, w, precedence)? as i64; if size < 1 { bail!("Invalid die size.") } let mut rolls = (0..count) .map(|_| rng.random_range(1..=size)) .collect::>(); // Exploding dice are early, they add rolls if let Some(fragments) = 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, w)?); } for i in 0..DICE_POOL_LIMIT { if i >= rolls.len() { break; } if comparers.len() == 0 { // Explode on max size dice if rolls[i] == size { rolls.push(rng.random_range(1..=size)); } } else { // Explode based on comparers provided for comparer in comparers.iter() { if comparer(rolls[i] as f64) { rolls.push(rng.random_range(1..=size)); break; } } } } } // Advantage/disadvantage roll filters let kh = if let Some(child) = kh { write!(w, "kh")?; Some(child.evaluate(rng, w, precedence)? as usize) } else { None }; let kl = if let Some(child) = kl { write!(w, "kl")?; Some(child.evaluate(rng, w, precedence)? as usize) } else { None }; let dh = if let Some(child) = dh { write!(w, "dh")?; Some(child.evaluate(rng, w, precedence)? as usize) } else { None }; let dl = if let Some(child) = dl { write!(w, "dl")?; Some(child.evaluate(rng, w, precedence)? as usize) } else { None }; // Don't sort the rolls directly, to preserve their original order for other logic // instead, `ranks` contains the indices where they would be in the array, // if it was sorted. let mut sorted_rolls = (0..rolls.len()).collect::>(); sorted_rolls.sort_by_key(|idx| rolls[*idx]); let mut ranks = vec![0usize; sorted_rolls.len()]; for (rank, roll_idx) in sorted_rolls.into_iter().enumerate() { ranks[roll_idx] = rank; } let mut result_set = Vec::with_capacity(rolls.len()); write!(w, " ")?; let mut w_set = w.witness_set()?; for idx in 0..rolls.len() { let mut skip = false; let roll = rolls[idx]; let rank = ranks[idx]; if let Some(kh) = kh { if rolls.len() - 1 - rank >= kh { skip = true; } } if let Some(dh) = dh { if rolls.len() - 1 - rank < dh { skip = true; } } if let Some(kl) = kl { if rank >= kl { skip = true; } } if let Some(dl) = dl { if rank < dl { skip = true; } } if !skip { result_set.push(roll as f64); } 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!(w, "-")?; let inner = inner_node.evaluate(rng, w, precedence)?; -inner } Add(lhs_node, rhs_node) => { 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, 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, 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, w, precedence)?; write!(w, " ÷ ")?; let rhs = rhs_node.evaluate(rng, w, precedence)?; lhs / rhs } }; if needs_parens { write!(w, ")")?; } Ok(result) } fn sample(&self, rng: &mut impl Rng) -> anyhow::Result { self.evaluate(rng, &mut std::io::sink(), 100) } fn avg(&self) -> anyhow::Result { let node = self; let result = match node { Expression::Const(x) => *x, Expression::Dice(DiceFormula { count, size, kh, kl, dh, dl, x, }) => { let count = if let Some(count) = count { count.avg()? as usize } else { 1 }; if count > 1_000_000 { bail!("Too many dice.") }; let size = size.avg()? as i64; if size < 1 { bail!("Invalid die size.") } if let Some(_) = x { bail!("Not implemented yet"); } if let Some(_) = kh { bail!("Not implemented yet"); }; if let Some(_) = kl { bail!("Not implemented yet"); }; if let Some(_) = dh { bail!("Not implemented yet"); }; if let Some(_) = dl { bail!("Not implemented yet"); }; count as f64 * (size as f64 + 1f64) * 0.5 } Expression::Neg(child) => -child.avg()?, Expression::Add(lhs, rhs) => lhs.avg()? + rhs.avg()?, Expression::Sub(lhs, rhs) => lhs.avg()? - rhs.avg()?, Expression::Mul(lhs, rhs) => lhs.avg()? * rhs.avg()?, Expression::Div(lhs, rhs) => lhs.avg()? / rhs.avg()?, }; Ok(result) } } impl CompareFragment { fn comparer( &self, rng: &mut impl Rng, witness: &mut W, ) -> anyhow::Result bool>> where W: Witness, E: Error + Send + Sync + 'static, { Ok(match self { CompareFragment::Eq(inner) => { write!(witness, "=")?; let comp = inner.evaluate(rng, witness, 2)?; Box::new(move |value| value == comp) } CompareFragment::Gt(inner) => { write!(witness, ">")?; let comp = inner.evaluate(rng, witness, 2)?; Box::new(move |value| value > comp) } CompareFragment::Ge(inner) => { write!(witness, ">=")?; let comp = inner.evaluate(rng, witness, 2)?; Box::new(move |value| value >= comp) } CompareFragment::Lt(inner) => { write!(witness, "<")?; let comp = inner.evaluate(rng, witness, 2)?; Box::new(move |value| value < comp) } CompareFragment::Le(inner) => { write!(witness, "<=")?; let comp = inner.evaluate(rng, witness, 2)?; Box::new(move |value| value <= comp) } }) } } impl Expression { fn precedence(&self) -> i64 { use Expression::*; match self { Const(_) => 1, Dice { .. } => 2, Neg(_) => 5, Mul(_, _) | Div(_, _) => 7, Add(_, _) | Sub(_, _) => 8, } } } trait Witness { type Ok; type Error: Error; type WitnessSet<'a>: WitnessSet where Self: 'a; fn witness_source_text(&mut self, text: &str) -> Result; fn witness_set(&mut self) -> Result, Self::Error>; fn write_fmt(&mut self, args: std::fmt::Arguments<'_>) -> Result where Self: Sized, { self.witness_source_text(&format(args)) } } trait WitnessSet { type Ok; type Error: Error; fn witness_roll(&mut self, dice: DiceRoll) -> Result; fn end(self) -> Result; } #[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.buffer.push_str(text); Ok(()) } fn witness_set(&mut self) -> Result, 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 { 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 { 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 { Ok(()) } fn witness_set(&mut self) -> Result, Self::Error> { Ok(self) } } impl WitnessSet for &mut std::io::Sink { type Ok = (); type Error = Infallible; fn witness_roll(&mut self, _dice: DiceRoll) -> Result { Ok(()) } fn end(self) -> Result { Ok(()) } }