use std::{
collections::{BTreeMap, BTreeSet},
mem
};
use super::{
meter::{Halt, Meter, charge_of},
record::{Powers, Roll}
};
use crate::{Distribution, Weight};
#[derive(Debug, Clone, Default, PartialEq, Eq, Hash, PartialOrd, Ord)]
struct Setup
{
roll: Option<Roll>,
lowest: i32,
highest: i32
}
impl Setup
{
fn results(&self) -> i32 { self.roll.as_ref().map_or(0, Roll::results) }
fn total(&self) -> Weight
{
self.roll.as_ref().map_or(Weight::ONE, Roll::total)
}
fn dice(&self) -> u64 { self.roll.as_ref().map_or(0, Roll::dice) }
fn cells(&self) -> u64 { self.roll.as_ref().map_or(1, Roll::cells) }
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub(super) struct Mixture(BTreeMap<Setup, Weight>);
impl Default for Mixture
{
fn default() -> Self
{
Self(BTreeMap::from([(Setup::default(), Weight::ONE)]))
}
}
impl Mixture
{
pub(super) fn rolls(rolls: impl IntoIterator<Item = (Roll, Weight)>)
-> Self
{
let mut setups = BTreeMap::<Setup, Weight>::new();
for (roll, weight) in rolls
{
let setup = Setup {
roll: Some(roll),
..Setup::default()
};
*setups.entry(setup).or_default() += weight;
}
Self(setups)
}
pub(super) fn cells(&self) -> u64 { self.0.keys().map(Setup::cells).sum() }
pub(super) fn drop_lowest(
&mut self,
counts: &Distribution,
meter: &mut Meter
) -> Result<(), Halt>
{
self.drop(counts, |setup| &mut setup.lowest, meter)
}
pub(super) fn drop_highest(
&mut self,
counts: &Distribution,
meter: &mut Meter
) -> Result<(), Halt>
{
self.drop(counts, |setup| &mut setup.highest, meter)
}
fn drop(
&mut self,
counts: &Distribution,
dropped: impl Fn(&mut Setup) -> &mut i32,
meter: &mut Meter
) -> Result<(), Halt>
{
let outcomes = counts.len() as u128;
let (steps, cells) =
self.0.keys().fold((0u128, 0u128), |(steps, cells), setup| {
(
steps + outcomes * (1 + setup.dice() as u128),
cells + outcomes * setup.cells() as u128
)
});
meter.charge(charge_of(steps), charge_of(cells))?;
let mut setups = BTreeMap::<Setup, Weight>::new();
for (setup, weight) in mem::take(&mut self.0)
{
for (count, count_weight) in counts
{
let mut setup = setup.clone();
let results = setup.results();
let dropped = dropped(&mut setup);
*dropped =
dropped.saturating_add(count.max(0)).clamp(0, results);
*setups.entry(setup).or_default() += &weight * count_weight;
}
}
self.0 = setups;
Ok(())
}
fn denominator(&self, meter: &mut Meter) -> Result<Weight, Halt>
{
meter.charge(self.0.len() as u64, 0)?;
Ok(self
.0
.keys()
.fold(Weight::ONE, |lcm, setup| lcm.lcm(&setup.total())))
}
pub(super) fn total(&self, meter: &mut Meter) -> Result<Weight, Halt>
{
meter.charge(self.0.len() as u64, 0)?;
Ok(self.0.values().sum::<Weight>() * self.denominator(meter)?)
}
pub(super) fn sum(&self, meter: &mut Meter) -> Result<Distribution, Halt>
{
let mut powers = Powers::default();
if self.0.len() == 1
&& let Some((setup, weight)) = self.0.first_key_value()
&& *weight == Weight::ONE
{
return setup_sum(setup, &mut powers, meter)
}
let denominator = self.denominator(meter)?;
let mut sums = BTreeMap::<i32, Weight>::new();
let mut ranges = Ranges::default();
for (setup, weight) in &self.0
{
meter.charge(1, 0)?;
let scale = weight * denominator.exact_div(&setup.total());
match setup
{
&Setup {
roll: Some(Roll::Range { start, end }),
lowest: 0,
highest: 0
} if start <= end => ranges.add(start, end, scale, meter)?,
_ =>
{
let sum = setup_sum(setup, &mut powers, meter)?;
let len = sum.len() as u64;
let before = sums.len() as u64;
meter.charge(len, len)?;
for (outcome, weight) in sum
{
*sums.entry(outcome).or_default() += weight * &scale;
}
let grown = sums.len() as u64 - before;
meter.free(2 * len - grown);
}
}
}
ranges.sweep(&mut sums, meter)?;
let len = sums.len() as u64;
meter.charge(len, len)?;
let distribution =
Distribution::from_weights(sums).expect("a mixture has outcomes");
meter.free(2 * len - distribution.len() as u64);
Ok(distribution)
}
pub(super) fn outcomes(
&self,
meter: &mut Meter
) -> Result<Vec<(Mixture, Weight)>, Halt>
{
let denominator = self.denominator(meter)?;
let mut outcomes = BTreeMap::<Setup, Weight>::new();
for (setup, weight) in &self.0
{
meter.charge(1, 0)?;
let scale = weight * denominator.exact_div(&setup.total());
let rolls = match &setup.roll
{
Some(roll) => roll
.outcomes(meter)?
.into_iter()
.map(|(roll, branches)| (Some(roll), branches))
.collect(),
None =>
{
meter.charge(1, 1)?;
vec![(None, Weight::ONE)]
}
};
meter.charge(rolls.len() as u64, 0)?;
for (roll, branches) in rolls
{
let outcome = Setup {
roll,
..setup.clone()
};
*outcomes.entry(outcome).or_default() += branches * &scale;
}
}
Ok(outcomes
.into_iter()
.map(|(setup, weight)| {
(Self(BTreeMap::from([(setup, Weight::ONE)])), weight)
})
.collect())
}
}
fn setup_sum(
setup: &Setup,
powers: &mut Powers,
meter: &mut Meter
) -> Result<Distribution, Halt>
{
match &setup.roll
{
Some(roll) => roll.sum(setup.lowest, setup.highest, powers, meter),
None =>
{
meter.charge(1, 1)?;
Ok(Distribution::point(0))
}
}
}
#[derive(Debug, Default)]
struct Ranges
{
rises: BTreeMap<i64, Weight>,
falls: BTreeMap<i64, Weight>
}
impl Ranges
{
fn add(
&mut self,
start: i32,
end: i32,
weight: Weight,
meter: &mut Meter
) -> Result<(), Halt>
{
meter.charge(2, 2)?;
*self.falls.entry(end as i64 + 1).or_default() += &weight;
*self.rises.entry(start as i64).or_default() += weight;
Ok(())
}
fn sweep(
self,
sums: &mut BTreeMap<i32, Weight>,
meter: &mut Meter
) -> Result<(), Halt>
{
let (Some(first), Some(last)) =
(self.rises.first_key_value(), self.falls.last_key_value())
else
{
return Ok(())
};
let span = (last.0 - first.0) as u64;
let edges = (self.rises.len() + self.falls.len()) as u64;
meter.charge(span + edges, span + 2 * edges)?;
let before = sums.len() as u64;
let breaks = self
.rises
.keys()
.chain(self.falls.keys())
.copied()
.collect::<BTreeSet<_>>()
.into_iter()
.collect::<Vec<_>>();
let mut level = Weight::ZERO;
for pair in breaks.windows(2)
{
let (here, next) = (pair[0], pair[1]);
if let Some(rise) = self.rises.get(&here)
{
level += rise;
}
if let Some(fall) = self.falls.get(&here)
{
level = level
.checked_sub(fall)
.expect("a range ends only after it starts");
}
if !level.is_zero()
{
for outcome in here..next
{
*sums.entry(outcome as i32).or_default() += &level;
}
}
}
let grown = sums.len() as u64 - before;
meter.free(span + 2 * edges - grown);
Ok(())
}
}