nat1/src/dice.rs

535 lines
16 KiB
Rust
Raw Blame History

This file contains invisible Unicode characters

This file contains invisible Unicode characters that are indistinguishable to humans but may be processed differently by a computer. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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(())
}
}