use std::collections::{BTreeMap, HashMap};
use super::meter::{Halt, Meter, charge_of};
use crate::{Distribution, Weight};
#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub(super) enum Roll
{
Range
{
start: i32,
end: i32
},
Standard
{
count: i32,
faces: i32
},
Custom
{
count: i32,
faces: Vec<(i32, Weight)>
},
Results(Vec<(i32, i32)>)
}
pub(super) fn distinct_faces(faces: &[i32]) -> Vec<(i32, Weight)>
{
let mut distinct = BTreeMap::<i32, Weight>::new();
for &face in faces
{
*distinct.entry(face).or_default() += Weight::ONE;
}
distinct.into_iter().collect()
}
impl Roll
{
pub(super) fn results(&self) -> i32
{
match self
{
Roll::Range { .. } => 1,
Roll::Standard { count, .. } | Roll::Custom { count, .. } =>
{
(*count).max(0)
},
Roll::Results(results) =>
{
results.iter().map(|(_, copies)| copies).sum()
},
}
}
pub(super) fn dice(&self) -> u64
{
match self
{
Roll::Standard { count, .. } | Roll::Custom { count, .. } =>
{
(*count).max(0) as u64
},
Roll::Range { .. } | Roll::Results(_) => 0
}
}
pub(super) fn cells(&self) -> u64
{
1 + match self
{
Roll::Custom { faces, .. } => faces.len() as u64,
Roll::Results(results) => results.len() as u64,
Roll::Range { .. } | Roll::Standard { .. } => 0
}
}
fn faces(&self) -> u64
{
match self
{
Roll::Standard { faces, .. } => (*faces).max(0) as u64,
Roll::Custom { faces, .. } => faces.len() as u64,
Roll::Range { .. } | Roll::Results(_) => 0
}
}
pub(super) fn total(&self) -> Weight
{
match self
{
Roll::Range { start, end } if end < start => Weight::ONE,
Roll::Range { start, end } =>
{
Weight::from((*end as i64 - *start as i64 + 1) as u64)
},
Roll::Standard { count, faces } if *count > 0 && *faces > 0 =>
{
Weight::from(*faces as u32).pow(*count as u32)
},
Roll::Custom { count, faces }
if *count > 0 && !faces.is_empty() =>
{
faces
.iter()
.map(|(_, weight)| weight)
.sum::<Weight>()
.pow(*count as u32)
},
_ => Weight::ONE
}
}
fn die(&self) -> Vec<(i32, Weight)>
{
match self
{
Roll::Standard { faces, .. } =>
{
(1..=*faces).map(|face| (face, Weight::ONE)).collect()
},
Roll::Custom { faces, .. } => faces.clone(),
Roll::Range { .. } => unreachable!("a range has no die"),
Roll::Results(_) => unreachable!("fixed results have no die")
}
}
pub(super) fn outcomes(
&self,
meter: &mut Meter
) -> Result<Vec<(Roll, Weight)>, Halt>
{
let die = match self
{
Roll::Range { start, end } if end < start =>
{
meter.charge(1, 2)?;
return Ok(vec![(Roll::Results(vec![(0, 1)]), Weight::ONE)])
},
&Roll::Range { start, end } =>
{
let width = (end as i64 - start as i64 + 1) as u64;
meter.charge(width, 2 * width)?;
return Ok((start..=end)
.map(|result| {
(Roll::Results(vec![(result, 1)]), Weight::ONE)
})
.collect())
},
Roll::Results(_) =>
{
meter.charge(1, self.cells())?;
return Ok(vec![(self.clone(), Weight::ONE)])
},
_ if self.results() == 0 =>
{
meter.charge(1, 1)?;
return Ok(vec![(Roll::Results(Vec::new()), Weight::ONE)])
},
Roll::Standard { .. } | Roll::Custom { .. } =>
{
let faces = self.faces();
meter.charge(faces, faces)?;
self.die()
}
};
let n = self.results() as usize;
if die.is_empty()
{
meter.charge(1, 2)?;
return Ok(vec![(Roll::Results(vec![(0, n as i32)]), Weight::ONE)])
}
meter.charge(1, 1)?;
let mut held = 1;
let mut partial = vec![(Vec::<(i32, i32)>::new(), 0usize, Weight::ONE)];
for (index, (face, weight)) in die.iter().enumerate()
{
let greatest = index == die.len() - 1;
let (steps, grown) = partial.iter().fold(
(0u128, 0u128),
|(steps, grown), (_, j, _)| match greatest
{
true => (steps + *j as u128 + 1, grown + 1),
false =>
{
let m = (n - j + 1) as u128;
(steps + m, grown + m)
}
}
);
let cells = charge_of(grown * (index as u128 + 2));
meter.charge(charge_of(steps), cells)?;
let mut next = Vec::new();
for (results, j, w) in partial
{
if greatest
{
let m = n - j;
let binomial = (1..=j).fold(Weight::ONE, |binomial, i| {
(binomial * Weight::from(n - j + i))
.exact_div(&Weight::from(i))
});
let power = weight.pow(m as u32);
let mut results = results;
if m > 0
{
results.push((*face, m as i32));
}
next.push((results, n, w * binomial * power));
continue
}
let mut binomial = Weight::ONE;
let mut power = Weight::ONE;
for m in 0..=n - j
{
let mut results = results.clone();
if m > 0
{
binomial = (binomial * Weight::from(j + m))
.exact_div(&Weight::from(m));
power *= weight;
results.push((*face, m as i32));
}
next.push((results, j + m, &w * &binomial * &power));
}
}
let occupied = next
.iter()
.map(|(results, _, _)| results.len() as u64 + 1)
.sum::<u64>();
meter.free(cells - occupied + held);
held = occupied;
partial = next;
}
meter.free(die.len() as u64);
Ok(partial
.into_iter()
.map(|(results, _, weight)| (Roll::Results(results), weight))
.collect())
}
}
impl Roll
{
#[cfg_attr(doc, aquamarine::aquamarine)]
pub(super) fn sum(
&self,
lowest: i32,
highest: i32,
powers: &mut Powers,
meter: &mut Meter
) -> Result<Distribution, Halt>
{
match self
{
Roll::Results(results) =>
{
meter.charge(results.len() as u64, 1)?;
Ok(Distribution::point(kept_sum(results, lowest, highest)))
},
Roll::Range { start, end } if end < start =>
{
meter.charge(1, 1)?;
Ok(Distribution::point(0))
},
Roll::Range { .. } if lowest > 0 || highest > 0 =>
{
meter.charge(1, 1)?;
Ok(Distribution::from_weights([(0, self.total())])
.expect("a roll has a nonzero total"))
},
&Roll::Range { start, end } =>
{
let width = (end as i64 - start as i64 + 1) as u64;
meter.charge(width, width)?;
Ok(Distribution::from_weights(
(start..=end).map(|outcome| (outcome, Weight::ONE))
)
.expect("a range that is not empty has outcomes"))
},
Roll::Standard { count, .. } | Roll::Custom { count, .. } =>
{
let count = *count;
let faces = self.faces();
if count <= 0 || faces == 0
{
meter.charge(1, 1)?;
return Ok(Distribution::point(0))
}
if lowest as i64 + highest as i64 >= count as i64
{
meter.charge(1, 1)?;
return Ok(Distribution::from_weights([(0, self.total())])
.expect("a roll has a nonzero total"))
}
meter.charge(faces, faces)?;
let die = self.die();
if lowest == 0 && highest == 0 && folds_to_clamp(&die, count)
{
let die = die
.into_iter()
.map(|(face, weight)| (face as i64, weight))
.collect::<Vec<_>>();
clamp(powers.power(die, count as u32, meter)?, meter)
}
else
{
let sum =
order_statistics(&die, count, lowest, highest, meter)?;
meter.free(faces);
Ok(sum)
}
}
}
}
}
fn kept_sum(results: &[(i32, i32)], lowest: i32, highest: i32) -> i32
{
let count = results
.iter()
.map(|(_, copies)| *copies as i64)
.sum::<i64>();
let kept = lowest as i64..count - highest as i64;
let mut j = 0i64;
let mut sum = 0i32;
for &(result, copies) in results
{
let placed = kept.start.max(j)..kept.end.min(j + copies as i64);
let added = (placed.end - placed.start).max(0) * result as i64;
sum =
(sum as i64 + added).clamp(i32::MIN as i64, i32::MAX as i64) as i32;
j += copies as i64;
}
sum
}
fn folds_to_clamp(die: &[(i32, Weight)], count: i32) -> bool
{
let least = die[0].0 as i64;
let greatest = die[die.len() - 1].0 as i64;
least >= 0 || greatest <= 0 || count as i64 * least >= i32::MIN as i64
}
pub(super) fn clamp(
wide: Vec<(i64, Weight)>,
meter: &mut Meter
) -> Result<Distribution, Halt>
{
let len = wide.len() as u64;
meter.charge(len, len)?;
let clamped = Distribution::from_weights(wide.into_iter().map(
|(outcome, weight)| {
(
outcome.clamp(i32::MIN as i64, i32::MAX as i64) as i32,
weight
)
}
))
.expect("a convolution of distributions has outcomes");
meter.free(2 * len - clamped.len() as u64);
Ok(clamped)
}
pub(super) fn convolve(
xs: &[(i64, Weight)],
ys: &[(i64, Weight)],
meter: &mut Meter
) -> Result<Vec<(i64, Weight)>, Halt>
{
let least = xs[0].0 + ys[0].0;
let greatest = xs[xs.len() - 1].0 + ys[ys.len() - 1].0;
let span = (greatest as i128 - least as i128 + 1) as u128;
let pairs = xs.len() as u128 * ys.len() as u128;
let sums = charge_of(span.min(pairs));
meter.charge(charge_of(pairs), sums.saturating_mul(2))?;
let convolution = convolve_charged(xs, ys, least, span, pairs);
meter.free(sums.saturating_mul(2) - convolution.len() as u64);
Ok(convolution)
}
fn convolve_charged(
xs: &[(i64, Weight)],
ys: &[(i64, Weight)],
least: i64,
span: u128,
pairs: u128
) -> Vec<(i64, Weight)>
{
if span <= pairs
{
convolve_packed(xs, ys, least, span)
.unwrap_or_else(|| convolve_dense(xs, ys, least, span))
}
else
{
let mut sums = BTreeMap::<i64, Weight>::new();
for (x, wx) in xs
{
for (y, wy) in ys
{
*sums.entry(x + y).or_default() += wx * wy;
}
}
sums.into_iter().collect()
}
}
pub(crate) fn convolve_dense(
xs: &[(i64, Weight)],
ys: &[(i64, Weight)],
least: i64,
span: u128
) -> Vec<(i64, Weight)>
{
let mut sums = vec![Weight::ZERO; span as usize];
for (x, wx) in xs
{
for (y, wy) in ys
{
sums[(x + y - least) as usize] += wx * wy;
}
}
sums.into_iter()
.zip(least..)
.filter(|(weight, _)| !weight.is_zero())
.map(|(weight, outcome)| (outcome, weight))
.collect()
}
pub(crate) fn convolve_packed(
xs: &[(i64, Weight)],
ys: &[(i64, Weight)],
least: i64,
span: u128
) -> Option<Vec<(i64, Weight)>>
{
let bits = |ws: &[(i64, Weight)]| {
ws.iter().map(|(_, w)| w.bits()).max().unwrap_or(0)
};
let (bx, by) = (bits(xs), bits(ys));
if bx + by <= u128::BITS as u64
{
return None
}
let terms = xs.len().min(ys.len()) as u64;
let slot = (bx + by + (u64::BITS - terms.leading_zeros()) as u64)
.div_ceil(u32::BITS as u64) as usize;
let pack = |ws: &[(i64, Weight)]| {
let origin = ws[0].0;
let mut digits =
vec![0u32; (ws[ws.len() - 1].0 - origin + 1) as usize * slot];
for (outcome, weight) in ws
{
let at = (outcome - origin) as usize * slot;
weight.write_u32_digits(&mut digits[at..at + slot]);
}
Weight::from_u32_digits(&digits)
};
let digits = (&pack(xs) * &pack(ys)).to_u32_digits();
Some(
(0..span as usize)
.zip(least..)
.filter_map(|(k, outcome)| {
let start = (k * slot).min(digits.len());
let end = ((k + 1) * slot).min(digits.len());
let weight = Weight::from_u32_digits(&digits[start..end]);
(!weight.is_zero()).then_some((outcome, weight))
})
.collect()
)
}
#[derive(Debug, Default)]
pub(super) struct Powers(HashMap<Wide, (u32, Wide)>);
type Wide = Vec<(i64, Weight)>;
impl Powers
{
fn power(
&mut self,
die: Vec<(i64, Weight)>,
n: u32,
meter: &mut Meter
) -> Result<Vec<(i64, Weight)>, Halt>
{
let faces = die.len() as u64;
let power = match self.0.get(&die)
{
Some((k, power)) if *k == n =>
{
let len = power.len() as u64;
meter.charge(len, len)?;
meter.free(faces);
return Ok(power.clone())
},
Some((k, power)) if *k < n =>
{
meter.charge(faces, faces)?;
let step = convolution_power(die.clone(), n - k, meter)?;
let power = convolve(power, &step, meter)?;
meter.free(step.len() as u64);
power
},
Some(_) => return convolution_power(die, n, meter),
None =>
{
meter.charge(faces, faces)?;
convolution_power(die.clone(), n, meter)?
}
};
let len = power.len() as u64;
meter.charge(len, len)?;
if let Some((_, old)) = self.0.insert(die, (n, power.clone()))
{
meter.free(old.len() as u64 + faces);
}
Ok(power)
}
}
fn convolution_power(
die: Vec<(i64, Weight)>,
mut n: u32,
meter: &mut Meter
) -> Result<Vec<(i64, Weight)>, Halt>
{
debug_assert!(n > 0, "at least one die");
let mut power = None::<Vec<(i64, Weight)>>;
let mut square = die;
loop
{
if n & 1 == 1
{
power = Some(match power
{
None =>
{
let len = square.len() as u64;
meter.charge(len, len)?;
square.clone()
},
Some(power) =>
{
let next = convolve(&power, &square, meter)?;
meter.free(power.len() as u64);
next
}
});
}
n >>= 1;
if n == 0
{
meter.free(square.len() as u64);
return Ok(power.expect("at least one die"))
}
let next = convolve(&square, &square, meter)?;
meter.free(square.len() as u64);
square = next;
}
}
fn order_statistics(
die: &[(i32, Weight)],
count: i32,
lowest: i32,
highest: i32,
meter: &mut Meter
) -> Result<Distribution, Halt>
{
let n = count as usize;
let kept = lowest as usize..n - highest as usize;
let slots = n as u64 + 1;
meter.charge(slots, slots + 1)?;
let mut held = slots + 1;
let mut states = vec![BTreeMap::<i32, Weight>::new(); n + 1];
states[0].insert(0, Weight::ONE);
for (index, (face, weight)) in die.iter().enumerate()
{
let greatest = index == die.len() - 1;
let (steps, moved) = states.iter().enumerate().fold(
(slots as u128, 0u128),
|(steps, moved), (j, sums)| {
let placements = (n - j + 1) as u128;
let sums = sums.len() as u128;
let moves = match greatest
{
true => sums,
false => sums * placements
};
(steps + placements + moves, moved + moves)
}
);
let cells = charge_of(moved).saturating_add(slots);
meter.charge(charge_of(steps), cells)?;
let mut next = vec![BTreeMap::<i32, Weight>::new(); n + 1];
for (j, sums) in states.iter().enumerate()
{
if sums.is_empty()
{
continue
}
let mut binomial = Weight::ONE;
let mut power = Weight::ONE;
for m in 0..=n - j
{
if m > 0
{
binomial = (binomial * Weight::from(j + m))
.exact_div(&Weight::from(m));
power *= weight;
}
if greatest && m < n - j
{
continue
}
let placed = kept.start.max(j)..kept.end.min(j + m);
let added = placed.len() as i64 * *face as i64;
let factor = &binomial * &power;
let target = &mut next[j + m];
for (&sum, w) in sums
{
let sum = (sum as i64 + added)
.clamp(i32::MIN as i64, i32::MAX as i64)
as i32;
*target.entry(sum).or_default() += w * &factor;
}
}
}
let occupied =
slots + next.iter().map(|sums| sums.len() as u64).sum::<u64>();
meter.free(cells - occupied + held);
held = occupied;
states = next;
}
let sums = states.pop().expect("n + 1 states");
let len = sums.len() as u64;
meter.charge(len, len)?;
let distribution =
Distribution::from_weights(sums).expect("a roll has outcomes");
meter.free(held + len - distribution.len() as u64);
Ok(distribution)
}