535 lines
16 KiB
Rust
535 lines
16 KiB
Rust
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<Expression>;
|
||
|
||
#[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<ExpBox>,
|
||
pub(crate) size: ExpBox,
|
||
pub(crate) kh: Option<ExpBox>,
|
||
pub(crate) kl: Option<ExpBox>,
|
||
pub(crate) dh: Option<ExpBox>,
|
||
pub(crate) dl: Option<ExpBox>,
|
||
pub(crate) x: Option<Vec<CompareFragment>>,
|
||
}
|
||
|
||
#[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<String> {
|
||
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<f64> {
|
||
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<W, E>(
|
||
&self,
|
||
rng: &mut impl Rng,
|
||
w: &mut W,
|
||
outer_precedence: i64,
|
||
) -> 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!(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::<Vec<_>>();
|
||
|
||
// 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::<Vec<_>>();
|
||
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<f64> {
|
||
self.evaluate(rng, &mut std::io::sink(), 100)
|
||
}
|
||
|
||
fn avg(&self) -> anyhow::Result<f64> {
|
||
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<W, E>(
|
||
&self,
|
||
rng: &mut impl Rng,
|
||
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!(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<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(())
|
||
}
|
||
}
|