initial impl

This commit is contained in:
Lilith Schier 2026-08-11 08:01:33 +02:00
commit c24a777dcc
6 changed files with 3241 additions and 0 deletions

354
src/dice.rs Normal file
View file

@ -0,0 +1,354 @@
use anyhow::bail;
use id_arena::Arena;
use rand::{Rng, RngExt};
use std::fmt::Write;
#[derive(Debug)]
pub struct Expression {
pub(crate) arena: Arena<AstNode>,
pub(crate) root: AstNodeId,
}
pub type AstNodeId = id_arena::Id<AstNode>;
#[derive(Debug, PartialEq, Clone)]
pub enum AstNode {
Const(f64),
Dice(DiceFormula),
Neg(AstNodeId),
Add(AstNodeId, AstNodeId),
Sub(AstNodeId, AstNodeId),
Mul(AstNodeId, AstNodeId),
Div(AstNodeId, AstNodeId),
}
#[derive(Debug, PartialEq, Clone)]
pub struct DiceFormula {
pub(crate) count: Option<AstNodeId>,
pub(crate) size: AstNodeId,
pub(crate) kh: Option<AstNodeId>,
pub(crate) kl: Option<AstNodeId>,
pub(crate) dh: Option<AstNodeId>,
pub(crate) dl: Option<AstNodeId>,
}
impl DiceFormula {
pub fn new(size: AstNodeId) -> Self {
Self {
count: None,
size,
kh: None,
kl: None,
dh: None,
dl: None,
}
}
}
impl Expression {
pub fn evaluate(&self, mut rng: impl Rng) -> anyhow::Result<String> {
let mut result = String::new();
let total = self.node_evaluate(&mut rng, self.root, &mut result, 100)?;
Ok(format!("**{total}** = {result}"))
}
pub fn average(&self, rng: &mut impl Rng) -> anyhow::Result<f64> {
if let Ok(avg) = self.node_average(self.root) {
return Ok(avg);
}
let mut average = self.node_sample(rng, self.root)?;
for idx in 1..1000 {
average =
(average * idx as f64 + self.node_sample(rng, self.root)?) / (idx as f64 + 1f64);
}
Ok(average)
}
fn node_evaluate(
&self,
rng: &mut impl Rng,
node_id: AstNodeId,
output: &mut String,
outer_precedence: i64,
) -> anyhow::Result<f64> {
use AstNode::*;
let node = &self.arena[node_id];
let precedence = node.precedence();
let needs_parens = precedence > outer_precedence;
if needs_parens {
output.push('(')
}
let result = match node {
Const(x) => {
write!(output, "{}", x).unwrap();
*x
}
Dice(DiceFormula {
count: count_node,
size: size_node,
kh,
kl,
dh,
dl,
}) => {
let count = if let Some(count_node) = count_node {
self.node_evaluate(rng, *count_node, output, precedence)? as usize
} else {
1
};
if count > 1_000_000 {
anyhow::bail!("Too many dice.")
};
output.push_str("d");
let size = self.node_evaluate(rng, *size_node, output, precedence)? as i64;
let rolls = (0..count)
.map(|_| rng.random_range(1..=size))
.collect::<Vec<_>>();
let kh = if let Some(id) = kh {
output.push_str("kh");
Some(self.node_evaluate(rng, *id, output, precedence)? as usize)
} else {
None
};
let kl = if let Some(id) = kl {
output.push_str("kl");
Some(self.node_evaluate(rng, *id, output, precedence)? as usize)
} else {
None
};
let dh = if let Some(id) = dh {
output.push_str("dh");
Some(self.node_evaluate(rng, *id, output, precedence)? as usize)
} else {
None
};
let dl = if let Some(id) = dl {
output.push_str("dl");
Some(self.node_evaluate(rng, *id, output, precedence)? as usize)
} else {
None
};
let mut sorted_rolls = (0..count).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());
let print_all_rolls = count <= 100;
if print_all_rolls {
output.push_str(" [ ");
}
for idx in 0..count {
let mut skip = false;
let roll = rolls[idx];
let rank = ranks[idx];
if let Some(kh) = kh {
if count - 1 - rank >= kh {
skip = true;
}
}
if let Some(dh) = dh {
if count - 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 {
if print_all_rolls {
write!(output, "~~`{}`~~", roll)?;
}
continue;
}
result_set.push(roll as f64);
if print_all_rolls {
let bold = roll == 1 || roll == size;
if bold {
write!(output, "**`{}`**", roll)?;
} else {
write!(output, "`{}`", roll)?;
}
if idx != count - 1 {
output.push_str(", ");
}
}
}
if print_all_rolls {
output.push_str(" ]");
}
result_set.iter().sum()
}
Neg(inner_node) => {
output.push_str("-");
let inner = self.node_evaluate(rng, *inner_node, output, precedence)?;
-inner
}
Add(lhs_node, rhs_node) => {
let lhs = self.node_evaluate(rng, *lhs_node, output, precedence)?;
output.push_str(" + ");
let rhs = self.node_evaluate(rng, *rhs_node, output, precedence)?;
lhs + rhs
}
Sub(lhs_node, rhs_node) => {
let lhs = self.node_evaluate(rng, *lhs_node, output, precedence)?;
output.push_str(" - ");
let rhs = self.node_evaluate(rng, *rhs_node, output, precedence)?;
lhs - rhs
}
Mul(lhs_node, rhs_node) => {
let lhs = self.node_evaluate(rng, *lhs_node, output, precedence)?;
output.push_str(" × ");
let rhs = self.node_evaluate(rng, *rhs_node, output, precedence)?;
lhs * rhs
}
Div(lhs_node, rhs_node) => {
let lhs = self.node_evaluate(rng, *lhs_node, output, precedence)?;
output.push_str(" ÷ ");
let rhs = self.node_evaluate(rng, *rhs_node, output, precedence)?;
lhs / rhs
}
};
if needs_parens {
output.push(')')
}
Ok(result)
}
fn node_sample(&self, rng: &mut impl Rng, node_id: AstNodeId) -> anyhow::Result<f64> {
let node = &self.arena[node_id];
let result = match node {
AstNode::Const(x) => *x,
AstNode::Dice(DiceFormula {
count,
size,
kh,
kl,
dh,
dl,
}) => {
let count = if let Some(count) = count {
self.node_sample(rng, *count)? as usize
} else {
1
};
if count > 1_000_000 {
anyhow::bail!("Too many dice.")
};
let size = self.node_sample(rng, *size)? as i64;
let mut rolls = (0..count)
.map(|_| rng.random_range(1..=size))
.collect::<Vec<_>>();
rolls.sort();
let kh = if let Some(id) = kh {
self.node_sample(rng, *id)? as usize
} else {
count
};
let kl = if let Some(id) = kl {
self.node_sample(rng, *id)? as usize
} else {
count
};
let dh = if let Some(id) = dh {
self.node_sample(rng, *id)? as usize
} else {
0
};
let dl = if let Some(id) = dl {
self.node_sample(rng, *id)? as usize
} else {
0
};
rolls
.into_iter()
.skip(dl.max(count - kh))
.take(count - (dh.max(count - kl)))
.map(|x| x as f64)
.sum::<f64>()
}
AstNode::Neg(id) => -self.node_sample(rng, *id)?,
AstNode::Add(lhs, rhs) => self.node_sample(rng, *lhs)? + self.node_sample(rng, *rhs)?,
AstNode::Sub(lhs, rhs) => self.node_sample(rng, *lhs)? - self.node_sample(rng, *rhs)?,
AstNode::Mul(lhs, rhs) => self.node_sample(rng, *lhs)? * self.node_sample(rng, *rhs)?,
AstNode::Div(lhs, rhs) => self.node_sample(rng, *lhs)? / self.node_sample(rng, *rhs)?,
};
Ok(result)
}
fn node_average(&self, node_id: AstNodeId) -> anyhow::Result<f64> {
let node = &self.arena[node_id];
let result = match node {
AstNode::Const(x) => *x,
AstNode::Dice(DiceFormula {
count,
size,
kh,
kl,
dh,
dl,
}) => {
let count = if let Some(count) = count {
self.node_average(*count)? as usize
} else {
1
};
if count > 1_000_000 {
bail!("Too many dice.")
};
let size = self.node_average(*size)? as i64;
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
}
AstNode::Neg(id) => -self.node_average(*id)?,
AstNode::Add(lhs, rhs) => self.node_average(*lhs)? + self.node_average(*rhs)?,
AstNode::Sub(lhs, rhs) => self.node_average(*lhs)? - self.node_average(*rhs)?,
AstNode::Mul(lhs, rhs) => self.node_average(*lhs)? * self.node_average(*rhs)?,
AstNode::Div(lhs, rhs) => self.node_average(*lhs)? / self.node_average(*rhs)?,
};
Ok(result)
}
}
impl AstNode {
fn precedence(&self) -> i64 {
use AstNode::*;
match self {
Const(_) => 1,
Dice { .. } => 2,
Neg(_) => 5,
Mul(_, _) | Div(_, _) => 7,
Add(_, _) | Sub(_, _) => 8,
}
}
}