use bevy::{platform::collections::HashSet, prelude::*, reflect::Reflect};
use bevy_ecs::component::Mutable;
use bevy_ecs::{component::StorageType};
use bevy_ecs::entity::MapEntities;
use crate::{active::{Active, Inactive}, guards::Guards, history::{History, HistoryState}};
pub mod active;
pub mod guards;
pub mod history;
pub mod prelude;
pub mod state_component;
pub mod transitions;
pub struct GearboxPlugin;
impl Plugin for GearboxPlugin {
fn build(&self, app: &mut App) {
app.add_observer(transition_observer)
.add_observer(active::add_active)
.add_observer(active::add_inactive)
.add_observer(initialize_state_machine);
app.register_type::<StateMachineRoot>();
app.register_type::<Parallel>();
app.register_type::<InitialState>();
app.register_type::<CurrentState>();
app.register_type::<History>();
app.register_type::<HistoryState>();
app.register_type::<StateChildren>();
app.register_type::<StateChildOf>();
app.register_type::<Guards>();
app.register_type::<Active>();
app.register_type::<Inactive>();
app.register_type::<EnterState>();
app.register_type::<ExitState>();
app.register_type::<TransitionActions>();
app.register_type::<OnAdd>();
app.register_type::<transitions::Source>();
app.register_type::<transitions::Transitions>();
app.register_type::<transitions::Target>();
app.register_type::<transitions::AlwaysEdge>();
app.register_type::<transitions::TransitionKind>();
app.add_observer(transitions::transition_always);
app.add_observer(transitions::start_after_on_enter);
app.add_observer(transitions::cancel_after_on_exit);
app.add_systems(Update, transitions::check_always_on_guards_changed);
app.add_systems(Update, transitions::tick_after_system);
}
}
#[derive(Component, Reflect, Default)]
#[reflect(Component)]
#[require(CurrentState)]
pub struct StateMachineRoot;
#[derive(Reflect, Component)]
#[reflect(Component)]
#[relationship_target(relationship = StateChildOf, linked_spawn)]
pub struct StateChildren(Vec<Entity>);
impl MapEntities for StateChildren {
fn map_entities<E: EntityMapper>(&mut self, entity_mapper: &mut E) {
for child in &mut self.0 {
*child = entity_mapper.get_mapped(*child);
}
}
}
impl<'a> IntoIterator for &'a StateChildren {
type Item = <Self::IntoIter as Iterator>::Item;
type IntoIter = std::slice::Iter<'a, Entity>;
#[inline(always)]
fn into_iter(self) -> Self::IntoIter {
self.0.iter()
}
}
#[derive(Reflect, Component)]
#[reflect(Component)]
#[relationship(relationship_target = StateChildren)]
pub struct StateChildOf(pub Entity);
impl MapEntities for StateChildOf {
fn map_entities<E: EntityMapper>(&mut self, entity_mapper: &mut E) {
self.0 = entity_mapper.get_mapped(self.0);
}
}
#[derive(Event)]
pub struct Transition {
pub source: Entity,
pub edge: Entity,
}
#[derive(Event, Reflect)]
pub struct TransitionActions {
pub source: Entity,
pub edge: Entity,
pub target: Entity,
}
#[derive(Component, Reflect, Default)]
#[reflect(Component)]
pub struct Parallel;
#[derive(Reflect)]
#[reflect(Component)]
pub struct InitialState(pub Entity);
impl Component for InitialState {
const STORAGE_TYPE: StorageType = StorageType::Table;
type Mutability = Mutable;
fn map_entities<E: EntityMapper>(this: &mut Self, entity_mapper: &mut E) {
this.0 = entity_mapper.get_mapped(this.0);
}
}
#[derive(Component, Reflect, Default)]
#[reflect(Component)]
pub struct CurrentState(pub HashSet<Entity>);
#[derive(Event, Reflect, Default)]
pub struct EnterState;
#[derive(Event, Reflect, Default)]
pub struct ExitState;
pub fn transition_observer(
trigger: Trigger<Transition>,
mut machine_query: Query<&mut CurrentState>,
parallel_query: Query<&Parallel>,
children_query: Query<&StateChildren>,
initial_state_query: Query<&InitialState>,
history_query: Query<&History>,
mut history_state_query: Query<&mut HistoryState>,
child_of_query: Query<&StateChildOf>,
edge_target_query: Query<&transitions::Target>,
kind_query: Query<&transitions::TransitionKind>,
mut commands: Commands,
) {
let machine_entity = trigger.target();
let source_state = trigger.event().source;
let new_super_state = match edge_target_query.get(trigger.event().edge) {
Ok(edge_target) => edge_target.0,
Err(_) => trigger.event().edge,
};
let Ok(mut current_state) = machine_query.get_mut(machine_entity) else {
return;
};
if current_state.0.is_empty() {
commands.trigger_targets(EnterState, machine_entity);
let mut path_to_target: Vec<Entity> = vec![new_super_state];
path_to_target.extend(
child_of_query
.iter_ancestors(new_super_state)
.take_while(|&ancestor| ancestor != machine_entity),
);
for entity in path_to_target.iter().rev() {
commands.trigger_targets(EnterState, *entity);
}
let new_leaf_states = get_all_leaf_states(
new_super_state,
&initial_state_query,
&children_query,
¶llel_query,
&history_query,
&history_state_query,
&child_of_query,
&mut commands,
);
current_state.0.extend(new_leaf_states);
return;
}
let Some(exiting_leaf_state) = current_state.0.iter().find(|leaf| {
**leaf == source_state
|| child_of_query
.iter_ancestors(**leaf)
.any(|ancestor| ancestor == source_state)
}).copied() else {
return;
};
let exit_path = get_path_to_root(exiting_leaf_state, &child_of_query);
let enter_path = get_path_to_root(new_super_state, &child_of_query);
let mut lca_depth = exit_path
.iter()
.rev()
.zip(enter_path.iter().rev())
.take_while(|(a, b)| a == b)
.count();
let lca_entity = if lca_depth > 0 {
Some(exit_path[exit_path.len() - lca_depth])
} else {
None
};
let is_internal = matches!(kind_query.get(trigger.event().edge), Ok(transitions::TransitionKind::Internal));
if !is_internal {
if new_super_state == exiting_leaf_state {
lca_depth = lca_depth.saturating_sub(1);
}
else if lca_entity == Some(source_state) {
lca_depth = lca_depth.saturating_sub(1);
}
}
let states_to_exit = &exit_path[..exit_path.len() - lca_depth];
let states_to_enter = &enter_path[..enter_path.len() - lca_depth];
for entity in states_to_exit {
if let Ok(history) = history_query.get(*entity) {
let states_to_save = match history {
History::Shallow => {
current_state.0.iter()
.filter(|&&state| {
if let Ok(parent) = child_of_query.get(state).map(|child_of| child_of.0) {
parent == *entity
} else {
false
}
})
.copied()
.collect()
}
History::Deep => {
current_state.0.iter()
.filter(|&&state| {
state == *entity || child_of_query
.iter_ancestors(state)
.any(|ancestor| ancestor == *entity)
})
.copied()
.collect()
}
};
if let Ok(mut existing_history) = history_state_query.get_mut(*entity) {
existing_history.0 = states_to_save;
} else {
commands.entity(*entity).insert(HistoryState(states_to_save));
}
}
commands.trigger_targets(ExitState, *entity);
}
current_state.0.remove(&exiting_leaf_state);
commands.trigger_targets(
TransitionActions {
source: source_state,
edge: trigger.event().edge,
target: new_super_state,
},
machine_entity,
);
for entity in states_to_enter.iter().rev() {
commands.trigger_targets(EnterState, *entity);
}
let new_leaf_states = get_all_leaf_states(
new_super_state,
&initial_state_query,
&children_query,
¶llel_query,
&history_query,
&history_state_query,
&child_of_query,
&mut commands,
);
current_state.0.extend(new_leaf_states);
}
fn get_path_to_root(start_entity: Entity, child_of_query: &Query<&StateChildOf>) -> Vec<Entity> {
let mut path = vec![start_entity];
path.extend(child_of_query.iter_ancestors(start_entity));
path
}
pub fn get_all_leaf_states(
start_node: Entity,
initial_state_query: &Query<&InitialState>,
children_query: &Query<&StateChildren>,
parallel_query: &Query<&Parallel>,
history_query: &Query<&History>,
history_state_query: &Query<&mut HistoryState>,
child_of_query: &Query<&StateChildOf>,
commands: &mut Commands,
) -> HashSet<Entity> {
let mut leaves = HashSet::new();
let mut stack = vec![start_node];
while let Some(entity) = stack.pop() {
let mut found_next = false;
if parallel_query.get(entity).is_ok() {
if let Ok(children) = children_query.get(entity) {
found_next = true;
for &child in children {
commands.trigger_targets(EnterState, child);
stack.push(child);
}
}
}
else if let (Ok(history), Ok(history_state)) = (history_query.get(entity), history_state_query.get(entity)) {
found_next = true;
match history {
History::Shallow => {
for &saved_state in &history_state.0 {
commands.trigger_targets(EnterState, saved_state);
stack.push(saved_state);
}
}
History::Deep => {
for &saved_state in &history_state.0 {
if saved_state != entity {
commands.trigger_targets(EnterState, saved_state);
}
leaves.insert(saved_state);
}
continue;
}
}
}
else if let Ok(initial_state) = initial_state_query.get(entity) {
found_next = true;
let mut path_to_substate = vec![initial_state.0];
path_to_substate.extend(
child_of_query
.iter_ancestors(initial_state.0)
.take_while(|&ancestor| ancestor != entity),
);
for e in path_to_substate.iter().rev() {
commands.trigger_targets(EnterState, *e);
}
stack.push(initial_state.0);
}
if !found_next {
leaves.insert(entity);
}
}
leaves
}
fn initialize_state_machine(
trigger: Trigger<OnAdd, StateMachineRoot>,
initial_state_query: Query<&InitialState>,
mut commands: Commands,
) {
let target = trigger.target();
let Ok(_initial_state) = initial_state_query.get(target) else {
return;
};
commands.trigger_targets(Transition { source: target, edge: target }, target);
}