use super::ast::Comprehension;
use super::metadata::{IndexFn, cycle_length};
use super::strategy::StrategyName;
pub mod antidiagonal;
pub mod diagonal;
pub mod extrema;
pub mod halton;
pub mod lex;
pub mod lhs;
pub mod prng;
pub mod reverse_lex;
pub mod shells;
pub mod shuffle;
pub mod sobol;
pub type MultiIndex = Vec<u64>;
#[derive(Debug, Clone, PartialEq)]
pub struct Tuple {
pub bindings: Vec<(String, TupleValue)>,
}
#[derive(Debug, Clone, PartialEq)]
pub enum TupleValue {
U64(u64),
I64(i64),
F64(f64),
Str(String),
Bool(bool),
}
impl Tuple {
pub fn new() -> Self {
Self {
bindings: Vec::new(),
}
}
pub fn with<K: Into<String>>(mut self, key: K, value: TupleValue) -> Self {
self.bindings.push((key.into(), value));
self
}
}
impl Default for Tuple {
fn default() -> Self {
Self::new()
}
}
pub struct EvaluatedInput {
pub tuples: Vec<Tuple>,
pub cardinality: u64,
pub index_fn: IndexFn,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Selection {
Prefix(u64),
Reverse {
total: u64,
len: u64,
},
Positions(Vec<u64>),
}
impl Selection {
pub fn len(&self) -> u64 {
match self {
Selection::Prefix(n) => *n,
Selection::Reverse { len, .. } => *len,
Selection::Positions(p) => p.len() as u64,
}
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn get(&self, i: u64) -> Option<u64> {
match self {
Selection::Prefix(n) => (i < *n).then_some(i),
Selection::Reverse { total, len } => (i < *len).then(|| total - 1 - i),
Selection::Positions(p) => usize::try_from(i).ok().and_then(|i| p.get(i).copied()),
}
}
pub fn iter(&self) -> impl Iterator<Item = u64> + '_ {
(0..self.len()).filter_map(|i| self.get(i))
}
pub(crate) fn from_multi_indices(
idx: &IndexFn,
multi_indices: Vec<MultiIndex>,
cardinality: u64,
) -> Self {
Selection::Positions(
multi_indices
.into_iter()
.filter_map(|mi| multi_index_to_flat(idx, &mi))
.map(|flat| flat as u64)
.filter(|p| *p < cardinality)
.collect(),
)
}
}
pub trait Strategy {
fn name(&self) -> StrategyName;
fn selects_from_shape(&self) -> bool;
fn accepts_input(&self, idx: Option<&IndexFn>) -> bool;
fn has_closed_form_for(&self, idx: &IndexFn) -> bool;
fn select(
&self,
index_fn: &IndexFn,
cardinality: u64,
truncation: Option<u64>,
seed: Option<u64>,
) -> Selection;
fn select_surviving(
&self,
index_fn: &IndexFn,
cardinality: u64,
truncation: Option<u64>,
seed: Option<u64>,
survivors: &[u64],
) -> Selection {
surviving_in_rank(
&|count| self.select(index_fn, cardinality, count, seed),
cardinality,
truncation,
survivors,
)
}
fn apply(&self, input: &EvaluatedInput, truncation: Option<u64>) -> Vec<Tuple> {
self.apply_seeded(input, truncation, None)
}
fn apply_seeded(
&self,
input: &EvaluatedInput,
truncation: Option<u64>,
seed: Option<u64>,
) -> Vec<Tuple> {
self.select(&input.index_fn, input.tuples.len() as u64, truncation, seed)
.iter()
.filter_map(|p| input.tuples.get(p as usize).cloned())
.collect()
}
}
pub(crate) fn surviving_in_rank(
select: &dyn Fn(Option<u64>) -> Selection,
cardinality: u64,
truncation: Option<u64>,
survivors: &[u64],
) -> Selection {
let want = capped(truncation, survivors.len() as u64) as usize;
if want == 0 {
return Selection::Positions(Vec::new());
}
let mut count = truncation.map(|t| t.min(cardinality));
loop {
let whole = count.is_none_or(|k| k >= cardinality);
let selected = select(count);
let reached = selected
.iter()
.filter(|p| survivors.binary_search(p).is_ok());
let rest = survivors.iter().copied().filter(|_| whole);
let mut taken = std::collections::HashSet::with_capacity(want);
let mut out = Vec::with_capacity(want);
for p in reached.chain(rest) {
if out.len() == want {
break;
}
if taken.insert(p) {
out.push(p);
}
}
if out.len() == want || whole {
return Selection::Positions(out);
}
count = count.map(|k| k.saturating_mul(2).min(cardinality));
}
}
pub(crate) fn capped(truncation: Option<u64>, total: u64) -> u64 {
truncation.map_or(total, |t| t.min(total))
}
pub fn for_name(name: StrategyName) -> Box<dyn Strategy + Send + Sync> {
match name {
StrategyName::Lex => Box::new(lex::Lex),
StrategyName::ReverseLex => Box::new(reverse_lex::ReverseLex),
StrategyName::Shuffle => Box::new(shuffle::Shuffle),
StrategyName::Halton => Box::new(halton::Halton),
StrategyName::Sobol => Box::new(sobol::Sobol),
StrategyName::Lhs => Box::new(lhs::Lhs),
StrategyName::Extrema => Box::new(extrema::Extrema),
StrategyName::Shells => Box::new(shells::Shells),
StrategyName::Diagonal => Box::new(diagonal::Diagonal),
StrategyName::Antidiagonal => Box::new(antidiagonal::Antidiagonal),
}
}
pub fn shape_input(child: &Comprehension, strategy: StrategyName) -> &Comprehension {
if !for_name(strategy).selects_from_shape() {
return child;
}
let mut input = child;
while let Comprehension::Order {
child,
truncation: None,
..
} = input
{
input = child;
}
input
}
pub fn ranked_filter(
child: &Comprehension,
strategy: StrategyName,
) -> Option<(&Comprehension, &str)> {
if strategy == StrategyName::Lex {
return None;
}
match shape_input(child, strategy) {
Comprehension::Filter { child, predicate } => {
Some((shape_input(child, strategy), predicate.as_str()))
}
_ => None,
}
}
pub fn multi_index_to_flat(idx: &IndexFn, mi: &MultiIndex) -> Option<usize> {
match idx {
IndexFn::Lattice { axis_sizes } => {
if mi.len() != axis_sizes.len() {
return None;
}
let mut flat: u64 = 0;
let mut stride: u64 = 1;
for i in (0..axis_sizes.len()).rev() {
let pos = mi[i];
let size = axis_sizes[i];
if pos >= size {
return None;
}
flat = flat.checked_add(pos.checked_mul(stride)?)?;
stride = stride.checked_mul(size)?;
}
Some(flat as usize)
}
IndexFn::Lockstep { length } => {
if mi.len() != 1 || mi[0] >= *length {
return None;
}
Some(mi[0] as usize)
}
IndexFn::Modular { axis_sizes } => {
if mi.len() != 1 || mi[0] >= cycle_length(axis_sizes) {
return None;
}
Some(mi[0] as usize)
}
IndexFn::Concatenation { segment_sizes } => {
let total: u64 = segment_sizes.iter().copied().sum();
if mi.len() != 1 || mi[0] >= total {
return None;
}
Some(mi[0] as usize)
}
IndexFn::Continuous { .. } | IndexFn::Hybrid { .. } => None,
}
}
pub fn index_fn_supports_lookup(idx: &IndexFn) -> bool {
!matches!(idx, IndexFn::Continuous { .. } | IndexFn::Hybrid { .. })
}
pub(crate) fn index_fn_size(idx: &IndexFn) -> u64 {
match idx {
IndexFn::Lattice { axis_sizes } => axis_sizes
.iter()
.copied()
.fold(1u64, |a, b| a.saturating_mul(b)),
IndexFn::Lockstep { length } => *length,
IndexFn::Modular { axis_sizes } => cycle_length(axis_sizes),
IndexFn::Concatenation { segment_sizes } => segment_sizes
.iter()
.copied()
.fold(0u64, |a, b| a.saturating_add(b)),
IndexFn::Continuous { .. } | IndexFn::Hybrid { .. } => 0,
}
}
pub(crate) fn index_fn_dim(idx: &IndexFn) -> usize {
match idx {
IndexFn::Lattice { axis_sizes } => axis_sizes.len(),
IndexFn::Continuous { intervals, .. } => intervals.len(),
IndexFn::Hybrid {
discrete_axes,
continuous_axes,
..
} => discrete_axes.len() + continuous_axes.len(),
IndexFn::Lockstep { .. } | IndexFn::Modular { .. } | IndexFn::Concatenation { .. } => 1,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn for_name_dispatches_to_correct_strategy() {
assert_eq!(for_name(StrategyName::Lex).name(), StrategyName::Lex);
assert_eq!(for_name(StrategyName::Halton).name(), StrategyName::Halton);
assert_eq!(
for_name(StrategyName::Extrema).name(),
StrategyName::Extrema
);
}
#[test]
fn index_fn_size_lattice() {
let idx = IndexFn::Lattice {
axis_sizes: vec![3, 4, 5],
};
assert_eq!(index_fn_size(&idx), 60);
}
#[test]
fn index_fn_size_concatenation() {
let idx = IndexFn::Concatenation {
segment_sizes: vec![10, 20, 30],
};
assert_eq!(index_fn_size(&idx), 60);
}
#[test]
fn index_fn_dim_classifies_correctly() {
assert_eq!(
index_fn_dim(&IndexFn::Lattice {
axis_sizes: vec![3, 4]
}),
2
);
assert_eq!(index_fn_dim(&IndexFn::Lockstep { length: 10 }), 1);
assert_eq!(
index_fn_dim(&IndexFn::Concatenation {
segment_sizes: vec![1, 2, 3]
}),
1
);
}
#[test]
fn one_axis_inputs_select_within_their_length() {
let inputs = [
IndexFn::Modular {
axis_sizes: vec![2, 7, 3],
},
IndexFn::Concatenation {
segment_sizes: vec![2, 3, 4],
},
IndexFn::Lockstep { length: 9 },
];
for idx in &inputs {
let total = index_fn_size(idx);
for name in [
StrategyName::Lex,
StrategyName::ReverseLex,
StrategyName::Diagonal,
StrategyName::Antidiagonal,
StrategyName::Extrema,
StrategyName::Shells,
StrategyName::Halton,
StrategyName::Sobol,
StrategyName::Lhs,
StrategyName::Shuffle,
] {
let full: Vec<u64> = for_name(name)
.select(idx, total, None, None)
.iter()
.collect();
let mut sorted = full.clone();
sorted.sort_unstable();
sorted.dedup();
assert!(
full.iter().all(|p| *p < total),
"{name:?} over {idx:?}: {full:?}"
);
if !matches!(name, StrategyName::Halton | StrategyName::Sobol) {
assert_eq!(
sorted.len() as u64,
total,
"{name:?} over {idx:?} reaches every position: {full:?}"
);
}
}
}
}
}