use crate::analytic::Constraints;
use crate::error::RuntimeError;
use crate::value::Value;
use crate::weight::Weight;
use probl_sema::SlotSet;
use probl_sema::ir::SlotId;
use rustc_hash::{FxHashMap, FxHasher};
use std::collections::BTreeSet;
use std::collections::hash_map::Entry;
use std::hash::{Hash, Hasher};
use std::sync::Arc;
pub type Returned = (Value, Weight, Constraints);
#[derive(Clone, Debug)]
pub struct World {
pub slots: Vec<Value>,
pub constraints: Constraints,
pub inherited: Arc<BTreeSet<u64>>,
pub weight: Weight,
pub run: u32,
}
impl World {
pub fn scaled(mut self, factor: f64) -> World {
self.weight = self.weight.scale(factor);
self
}
}
pub fn total_weight(worlds: &[World]) -> Weight {
Weight::sum(worlds.iter().map(|w| w.weight))
}
#[derive(Debug, Default)]
pub struct Flow {
pub next: Vec<World>,
pub broke: Vec<World>,
pub continued: Vec<World>,
pub returned: Vec<Returned>,
pub faulted: Vec<(World, RuntimeError)>,
}
impl Flow {
pub fn next(worlds: Vec<World>) -> Flow {
Flow {
next: worlds,
..Flow::default()
}
}
pub fn join(&mut self, other: Flow) {
self.next.extend(other.next);
self.broke.extend(other.broke);
self.continued.extend(other.continued);
self.returned.extend(other.returned);
self.faulted.extend(other.faulted);
}
}
pub fn clear(worlds: &mut [World], slots: &[SlotId]) {
if slots.is_empty() {
return;
}
for w in worlds {
for &s in slots {
w.slots[s as usize] = Value::Dead;
}
}
}
pub fn clear_dead(worlds: &mut [World], live: &SlotSet) {
let Some(n) = worlds.first().map(|w| w.slots.len()) else {
return;
};
let dead: Vec<SlotId> = live.iter_missing(n).collect();
clear(worlds, &dead);
}
pub fn live_slots(live: &SlotSet, n: usize) -> Vec<usize> {
live.iter().map(|s| s as usize).filter(|&s| s < n).collect()
}
pub fn state_hash(w: &World, live: &[usize]) -> u64 {
let mut hasher = FxHasher::default();
for &i in live {
w.slots[i].hash(&mut hasher);
}
w.constraints.hash(&mut hasher);
hasher.finish()
}
pub fn merge(mut worlds: Vec<World>, live: &SlotSet, enabled: bool) -> Vec<World> {
for w in &mut worlds {
if w.constraints.keys().all(|id| w.inherited.contains(id)) {
continue;
}
let mut ids = (*w.inherited).clone();
for slot in live.iter() {
crate::analytic::collect_ids(&w.slots[slot as usize], &mut ids);
}
Arc::make_mut(&mut w.constraints).retain(|id, _| ids.contains(id));
}
if !enabled || worlds.len() < 2 {
return worlds;
}
let n = worlds[0].slots.len();
let live = live_slots(live, n);
let mut out: Vec<World> = Vec::with_capacity(worlds.len());
let mut first: FxHashMap<u64, usize> = FxHashMap::with_capacity_and_hasher(worlds.len(), Default::default());
let mut next: Vec<usize> = Vec::with_capacity(worlds.len());
const NONE: usize = usize::MAX;
for w in worlds {
let same =
|kept: &World| kept.constraints == w.constraints && live.iter().all(|&i| kept.slots[i] == w.slots[i]);
let found = match first.entry(state_hash(&w, &live)) {
Entry::Vacant(entry) => {
entry.insert(out.len());
None
}
Entry::Occupied(entry) => {
let mut i = *entry.get();
loop {
if same(&out[i]) {
break Some(i);
}
if next[i] == NONE {
next[i] = out.len();
break None;
}
i = next[i];
}
}
};
match found {
Some(i) => out[i].weight += w.weight,
None => {
next.push(NONE);
out.push(w);
}
}
}
out
}
pub fn merge_values(pairs: Vec<Returned>) -> Vec<Returned> {
let mut out: Vec<Returned> = Vec::with_capacity(pairs.len());
let mut index: FxHashMap<(Value, Constraints), usize> = FxHashMap::default();
for (v, w, constraints) in pairs {
match index.entry((v.clone(), constraints.clone())) {
Entry::Occupied(entry) => out[*entry.get()].1 += w,
Entry::Vacant(entry) => {
entry.insert(out.len());
out.push((v, w, constraints));
}
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
fn world(slots: Vec<Value>) -> World {
World {
slots,
constraints: Default::default(),
inherited: Default::default(),
weight: Weight::ONE,
run: 0,
}
}
#[test]
fn worlds_merge_on_their_live_slots() {
let mut live = SlotSet::with_capacity(3);
live.insert(0);
let worlds = vec![
world(vec![Value::Int(1.into()), Value::Int(7.into()), Value::Dead]),
world(vec![Value::Int(2.into()), Value::Int(7.into()), Value::Dead]),
world(vec![
Value::Int(1.into()),
Value::list(vec![Value::Int(8.into())]),
Value::Dead,
]),
];
let mut merged = merge(worlds, &live, true);
assert_eq!(merged.len(), 2);
assert_eq!(merged[0].slots[0], Value::Int(1.into()));
assert_eq!(merged[0].weight.to_f64(), 2.0);
assert_eq!(merged[1].weight.to_f64(), 1.0);
clear_dead(&mut merged, &live);
assert!(merged.iter().all(|w| w.slots[1] == Value::Dead));
assert_eq!(merged[1].slots[0], Value::Int(2.into()));
}
}