use alloc::vec;
use alloc::vec::Vec;
use crate::ecs::{Access, BuiltSystem, SystemEntry};
pub(crate) struct ExecSchedule {
waves: Vec<Vec<usize>>,
accesses: Vec<Access>,
}
impl ExecSchedule {
#[cfg_attr(
not(test),
expect(
dead_code,
reason = "read by this module's tests until `World::step` walks waves"
)
)]
pub(crate) fn waves(&self) -> &[Vec<usize>] {
&self.waves
}
pub(crate) fn access(&self, system: usize) -> Access {
self.accesses[system]
}
pub(crate) fn len(&self) -> usize {
self.accesses.len()
}
pub(crate) fn is_empty(&self) -> bool {
self.accesses.is_empty()
}
}
pub(crate) fn build(systems: &[BuiltSystem], entries: &[SystemEntry]) -> ExecSchedule {
let names: Vec<&'static str> = systems.iter().map(|s| s.name()).collect();
let accesses: Vec<Access> = systems.iter().map(|s| s.access()).collect();
build_from(&names, &accesses, |name| {
entries
.iter()
.find(|e| e.name == name)
.map(|e| (e.after, e.before))
.unwrap_or((&[], &[]))
})
}
fn build_from(
names: &[&'static str],
accesses: &[Access],
edges_of: impl Fn(&str) -> (&'static [&'static str], &'static [&'static str]),
) -> ExecSchedule {
let position = |name: &str| names.iter().position(|n| *n == name);
let mut preds: Vec<Vec<usize>> = vec![Vec::new(); names.len()];
for (i, name) in names.iter().enumerate() {
let (after, before) = edges_of(name);
for a in after {
if let Some(j) = position(a) {
assert!(
j < i,
"schedule edge violated: {a} is declared before {name} but the table runs it later",
);
preds[i].push(j);
}
}
for b in before {
if let Some(j) = position(b) {
assert!(
i < j,
"schedule edge violated: {name} is declared before {b} but the table runs it later",
);
preds[j].push(i);
}
}
}
for i in 0..names.len() {
for j in (i + 1)..names.len() {
if accesses[i].conflicts_with(accesses[j]) {
preds[j].push(i);
}
}
}
let mut level = vec![0usize; names.len()];
for i in 0..names.len() {
level[i] = preds[i].iter().map(|&p| level[p] + 1).max().unwrap_or(0);
}
let wave_count = level.iter().map(|&l| l + 1).max().unwrap_or(0);
let mut waves: Vec<Vec<usize>> = vec![Vec::new(); wave_count];
for (i, &l) in level.iter().enumerate() {
waves[l].push(i);
}
ExecSchedule {
waves,
accesses: accesses.to_vec(),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ecs::{ComponentId, ComponentMask};
use alloc::vec;
fn mask(ids: &[u8]) -> ComponentMask {
let mut m = ComponentMask::EMPTY;
for &id in ids {
m.insert(ComponentId::new(id));
}
m
}
fn no_edges(_: &str) -> (&'static [&'static str], &'static [&'static str]) {
(&[], &[])
}
#[test]
fn exclusive_systems_are_singleton_waves_in_table_order() {
let names = ["A", "B", "C"];
let accesses = vec![Access::new().exclusive(); 3];
let s = build_from(&names, &accesses, no_edges);
let flat: Vec<usize> = s.waves().iter().flatten().copied().collect();
assert_eq!(flat, vec![0, 1, 2]);
assert!(s.waves().iter().all(|w| w.len() == 1));
}
#[test]
fn non_conflicting_systems_share_a_wave() {
let names = ["A", "B", "C"];
let accesses = vec![
Access::new().writes_components(mask(&[1])),
Access::new().writes_components(mask(&[2])),
Access::new().reads_components(mask(&[1])),
];
let s = build_from(&names, &accesses, no_edges);
assert_eq!(s.waves(), &[vec![0, 1], vec![2]]);
}
#[test]
fn declared_edges_order_non_conflicting_systems() {
let names = ["A", "B"];
let accesses = vec![Access::new(); 2];
let s = build_from(&names, &accesses, |n| {
if n == "B" {
(&["A"][..], &[][..])
} else {
(&[], &[])
}
});
assert_eq!(s.waves(), &[vec![0], vec![1]]);
}
#[test]
#[should_panic(expected = "schedule edge violated")]
fn edge_contradicting_table_order_panics() {
let names = ["A", "B"];
let accesses = vec![Access::new(); 2];
let _ = build_from(&names, &accesses, |n| {
if n == "A" {
(&["B"][..], &[][..])
} else {
(&[], &[])
}
});
}
#[test]
fn absent_edge_targets_are_ignored() {
let names = ["B"];
let accesses = vec![Access::new()];
let s = build_from(&names, &accesses, |_| (&["A"][..], &["Z"][..]));
assert_eq!(s.waves(), &[vec![0]]);
}
#[test]
fn chains_stack_levels() {
let names = ["A", "B", "C", "D"];
let w = |id| Access::new().writes_components(mask(&[id]));
let accesses = vec![
w(1),
Access::new()
.reads_components(mask(&[1]))
.writes_components(mask(&[2])),
Access::new().reads_components(mask(&[2])),
w(9),
];
let s = build_from(&names, &accesses, no_edges);
assert_eq!(s.waves(), &[vec![0, 3], vec![1], vec![2]]);
}
}