use rustc_hash::FxHashMap;
use std::collections::VecDeque;
use crate::error::{EcsError, Result};
use crate::system::{BoxedSystem, System, SystemAccess, SystemId};
#[derive(Debug, Clone)]
pub struct SystemNode {
pub id: SystemId,
pub accesses: SystemAccess,
}
#[derive(Debug, Clone)]
pub struct StageExecutionPlan {
pub name: String,
pub parallel_groups: Vec<crate::dependency::ExecutionStage>,
}
pub struct SystemGraph {
pub nodes: Vec<SystemNode>,
pub edges: FxHashMap<SystemId, Vec<SystemId>>,
pub reverse_edges: FxHashMap<SystemId, Vec<SystemId>>,
}
impl SystemGraph {
pub fn build(systems: &[BoxedSystem]) -> Self {
let mut nodes = Vec::with_capacity(systems.len());
let mut edges: FxHashMap<SystemId, Vec<SystemId>> = FxHashMap::default();
let mut reverse_edges: FxHashMap<SystemId, Vec<SystemId>> = FxHashMap::default();
for (i, system) in systems.iter().enumerate() {
let id = SystemId(i as u32);
let accesses = system.accesses();
nodes.push(SystemNode { id, accesses });
edges.insert(id, Vec::new());
reverse_edges.insert(id, Vec::new());
}
for i in 0..nodes.len() {
for j in (i + 1)..nodes.len() {
let id_a = nodes[i].id;
let id_b = nodes[j].id;
if nodes[i].accesses.conflicts_with(&nodes[j].accesses) {
edges.get_mut(&id_a).unwrap().push(id_b);
reverse_edges.get_mut(&id_b).unwrap().push(id_a);
}
}
}
Self {
nodes,
edges,
reverse_edges,
}
}
pub fn topological_sort(&self) -> Result<Vec<SystemId>> {
let mut in_degree: FxHashMap<SystemId, usize> = FxHashMap::default();
let mut queue = VecDeque::new();
let mut result = Vec::with_capacity(self.nodes.len());
for node in &self.nodes {
in_degree.insert(
node.id,
self.reverse_edges.get(&node.id).map_or(0, |v| v.len()),
);
}
for node in &self.nodes {
if in_degree[&node.id] == 0 {
queue.push_back(node.id);
}
}
while let Some(id) = queue.pop_front() {
result.push(id);
if let Some(neighbors) = self.edges.get(&id) {
for &neighbor in neighbors {
let degree = in_degree.get_mut(&neighbor).unwrap();
*degree -= 1;
if *degree == 0 {
queue.push_back(neighbor);
}
}
}
}
if result.len() != self.nodes.len() {
return Err(EcsError::SystemCycleDetected);
}
Ok(result)
}
}
#[derive(Debug, Clone)]
pub struct Stage {
pub name: String,
pub systems: Vec<SystemId>,
pub depends_on: Vec<String>, }
impl Stage {
pub fn new(name: &str) -> Self {
Self {
name: name.to_string(),
systems: Vec::new(),
depends_on: Vec::new(),
}
}
pub fn try_add(
&mut self,
system_id: SystemId,
access: &SystemAccess,
_graph: &SystemGraph,
) -> bool {
for &existing_id in &self.systems {
let existing_node = _graph.nodes.iter().find(|n| n.id == existing_id).unwrap();
if access.conflicts_with(&existing_node.accesses) {
return false;
}
}
self.systems.push(system_id);
true
}
}
impl Default for Stage {
fn default() -> Self {
Self::new("default")
}
}
#[derive(Debug, Clone)]
pub struct OrderingConstraint {
pub system_name: String,
pub before: Vec<String>,
pub after: Vec<String>,
}
pub struct Schedule {
pub(crate) systems: Vec<BoxedSystem>,
pub(crate) stages: Vec<Stage>,
pub(crate) graph: Option<SystemGraph>,
pub(crate) ordering_constraints: Vec<OrderingConstraint>,
pub(crate) parallel_plan: Vec<StageExecutionPlan>,
}
impl Default for Schedule {
fn default() -> Self {
Self::new()
}
}
impl Schedule {
pub fn from_systems(systems: Vec<BoxedSystem>) -> Result<Self> {
Self {
systems,
stages: Vec::new(),
graph: None,
ordering_constraints: Vec::new(),
parallel_plan: Vec::new(),
}
.build()
}
pub fn new() -> Self {
Self {
systems: Vec::new(),
stages: Vec::new(),
graph: None,
ordering_constraints: Vec::new(),
parallel_plan: Vec::new(),
}
}
pub fn add_stage(&mut self, name: &str) -> Result<()> {
if self.stages.iter().any(|s| s.name == name) {
return Err(EcsError::ScheduleError(format!(
"Stage '{name}' already exists"
)));
}
self.stages.push(Stage {
name: name.to_string(),
systems: Vec::new(),
depends_on: Vec::new(),
});
Ok(())
}
pub fn add_stage_dependency(&mut self, stage: &str, depends_on: &str) -> Result<()> {
if !self.stages.iter().any(|s| s.name == depends_on) {
return Err(EcsError::ScheduleError(format!(
"Dependency stage '{depends_on}' not found"
)));
}
let stage_def = self
.stages
.iter_mut()
.find(|s| s.name == stage)
.ok_or_else(|| EcsError::ScheduleError(format!("Stage '{stage}' not found")))?;
stage_def.depends_on.push(depends_on.to_string());
Ok(())
}
pub fn add_system_to_stage(&mut self, stage: &str, system: BoxedSystem) -> Result<()> {
let stage_def = self
.stages
.iter_mut()
.find(|s| s.name == stage)
.ok_or_else(|| EcsError::ScheduleError(format!("Stage '{stage}' not found")))?;
let system_id = SystemId(self.systems.len() as u32);
stage_def.systems.push(system_id);
self.systems.push(system);
Ok(())
}
pub fn validate_stages(&self) -> Result<()> {
let mut visited = std::collections::HashSet::new();
let mut temp = std::collections::HashSet::new();
fn visit(
stage_name: &str,
stages: &[Stage],
visited: &mut std::collections::HashSet<String>,
temp: &mut std::collections::HashSet<String>,
) -> Result<()> {
if temp.contains(stage_name) {
return Err(EcsError::ScheduleError(format!(
"Circular dependency detected involving stage '{stage_name}'"
)));
}
if visited.contains(stage_name) {
return Ok(());
}
temp.insert(stage_name.to_string());
if let Some(stage) = stages.iter().find(|s| s.name == stage_name) {
for dep in &stage.depends_on {
visit(dep, stages, visited, temp)?;
}
}
temp.remove(stage_name);
visited.insert(stage_name.to_string());
Ok(())
}
for stage in &self.stages {
visit(&stage.name, &self.stages, &mut visited, &mut temp)?;
}
Ok(())
}
pub fn with_system(mut self, system: BoxedSystem) -> Self {
self.add_system(system);
self
}
pub fn add_system(&mut self, system: BoxedSystem) {
self.systems.push(system);
self.invalidate();
}
pub fn add_system_before(&mut self, system: BoxedSystem, before: &str) {
let system_name = system.name().to_string();
self.systems.push(system);
if let Some(constraint) = self
.ordering_constraints
.iter_mut()
.find(|c| c.system_name == system_name)
{
constraint.before.push(before.to_string());
} else {
self.ordering_constraints.push(OrderingConstraint {
system_name: system_name.clone(),
before: vec![before.to_string()],
after: Vec::new(),
});
}
self.invalidate();
}
pub fn add_system_after(&mut self, system: BoxedSystem, after: &str) {
let system_name = system.name().to_string();
self.systems.push(system);
if let Some(constraint) = self
.ordering_constraints
.iter_mut()
.find(|c| c.system_name == system_name)
{
constraint.after.push(after.to_string());
} else {
self.ordering_constraints.push(OrderingConstraint {
system_name,
before: Vec::new(),
after: vec![after.to_string()],
});
}
self.invalidate();
}
fn invalidate(&mut self) {
self.graph = None;
self.parallel_plan.clear();
}
pub fn get_system_mut(&mut self, name: &str) -> Option<&mut (dyn System + 'static)> {
self.systems
.iter_mut()
.find(|sys| sys.name() == name)
.map(|sys| sys.as_mut())
}
pub fn build(mut self) -> Result<Self> {
self.rebuild()?;
Ok(self)
}
pub(crate) fn ensure_built(&mut self) -> Result<()> {
if self.graph.is_none() {
self.rebuild()?;
}
Ok(())
}
fn rebuild(&mut self) -> Result<()> {
if self.stages.is_empty() {
let mut default_stage = Stage::new("default");
for i in 0..self.systems.len() {
default_stage.systems.push(SystemId(i as u32));
}
self.stages.push(default_stage);
}
self.validate_stages()?;
let sorted_stages = self.topological_sort_stages_internal()?;
let mut parallel_plan = Vec::with_capacity(sorted_stages.len());
let all_accesses = self.get_accesses();
for stage in sorted_stages {
let stage_system_ids = &stage.systems;
if stage_system_ids.is_empty() {
continue;
}
let stage_accesses: Vec<SystemAccess> = stage_system_ids
.iter()
.map(|id| all_accesses[id.0 as usize].clone())
.collect();
let dep_graph = crate::dependency::DependencyGraph::new(stage_accesses);
let mut stage_parallel_groups = dep_graph.stages().to_vec();
for group in &mut stage_parallel_groups {
for idx in &mut group.system_indices {
*idx = stage_system_ids[*idx].0 as usize;
}
}
parallel_plan.push(StageExecutionPlan {
name: stage.name.clone(),
parallel_groups: stage_parallel_groups,
});
}
self.parallel_plan = parallel_plan;
self.graph = Some(SystemGraph::build(&self.systems));
Ok(())
}
fn topological_sort_stages_internal(&self) -> Result<Vec<Stage>> {
let mut sorted = Vec::new();
let mut visited = std::collections::HashSet::new();
let mut temp = std::collections::HashSet::new();
fn visit(
stage: &Stage,
all_stages: &[Stage],
sorted: &mut Vec<Stage>,
visited: &mut std::collections::HashSet<String>,
temp: &mut std::collections::HashSet<String>,
) -> Result<()> {
if temp.contains(&stage.name) {
return Err(EcsError::SystemCycleDetected);
}
if visited.contains(&stage.name) {
return Ok(());
}
temp.insert(stage.name.clone());
for dep_name in &stage.depends_on {
if let Some(dep_stage) = all_stages.iter().find(|s| s.name == *dep_name) {
visit(dep_stage, all_stages, sorted, visited, temp)?;
}
}
temp.remove(&stage.name);
visited.insert(stage.name.clone());
sorted.push(stage.clone());
Ok(())
}
for stage in &self.stages {
visit(stage, &self.stages, &mut sorted, &mut visited, &mut temp)?;
}
Ok(sorted)
}
pub fn stage_count(&self) -> usize {
self.stages.len()
}
pub fn stage_system_count(&self, stage_idx: usize) -> usize {
self.stages.get(stage_idx).map_or(0, |s| s.systems.len())
}
pub fn system_count(&self) -> usize {
self.systems.len()
}
pub(crate) fn system_mut_by_id(&mut self, id: SystemId) -> Option<&mut BoxedSystem> {
self.systems.get_mut(id.0 as usize)
}
pub(crate) fn stage_plan(&self) -> Vec<&[SystemId]> {
self.stages
.iter()
.map(|stage| stage.systems.as_slice())
.collect()
}
pub fn get_accesses(&self) -> Vec<SystemAccess> {
self.systems.iter().map(|s| s.accesses()).collect()
}
pub fn analyze_parallelization(&self) -> crate::dependency::DependencyGraph {
use crate::dependency::DependencyGraph;
DependencyGraph::new(self.get_accesses())
}
pub fn print_execution_plan(&self) {
let graph = self.analyze_parallelization();
graph.print_schedule();
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_stage_creation() {
let stage = Stage::new("test");
assert_eq!(stage.systems.len(), 0);
assert_eq!(stage.name, "test");
}
struct MockSystem;
impl crate::system::System for MockSystem {
fn run(
&mut self,
_world: &mut crate::World,
_commands: &mut crate::command::CommandBuffer,
) -> crate::error::Result<()> {
Ok(())
}
fn name(&self) -> &'static str {
"MockSystem"
}
fn accesses(&self) -> crate::system::SystemAccess {
crate::system::SystemAccess {
reads: vec![],
writes: vec![],
}
}
}
#[test]
fn test_lazy_rebuild() {
let mut schedule = Schedule::new();
schedule.add_system(Box::new(MockSystem));
assert!(
schedule.graph.is_none(),
"Graph should be None after add_system"
);
schedule.ensure_built().expect("Failed to build");
assert!(
schedule.graph.is_some(),
"Graph should be Some after ensure_built"
);
schedule.add_system(Box::new(MockSystem));
assert!(
schedule.graph.is_none(),
"Graph should be invalidated after adding new system"
);
}
}