use crate::marker::InputMarker;
use alloc::vec::Vec;
use bevy_ecs::entity::{EntityMapper, MapEntities};
use bevy_ecs::query::QueryData;
use bevy_enhanced_input::action::ActionTime;
use bevy_enhanced_input::prelude::{ActionEvents, ActionValue, TriggerState};
use core::fmt::{Debug, Formatter};
use core::time::Duration;
use lightyear_core::prelude::Tick;
use lightyear_inputs::input_buffer::{Compressed, InputBuffer};
use lightyear_inputs::input_message::{ActionStateQueryData, ActionStateSequence, InputSnapshot};
use serde::{Deserialize, Serialize};
pub type BEIBuffer<C> = InputBuffer<ActionsSnapshot, C>;
#[derive(Serialize, Deserialize)]
pub struct BEIStateSequence<C> {
start_state: ActionsSnapshot,
diffs: Vec<Compressed<ActionsDiff>>,
marker: core::marker::PhantomData<C>,
}
impl<C> PartialEq for BEIStateSequence<C> {
fn eq(&self, other: &Self) -> bool {
self.start_state == other.start_state && self.diffs == other.diffs
}
}
impl<C> Debug for BEIStateSequence<C> {
fn fmt(&self, f: &mut Formatter<'_>) -> core::fmt::Result {
f.debug_struct("BEIStateSequence")
.field("start_state", &self.start_state)
.field("diffs", &self.diffs)
.finish()
}
}
impl<C> Clone for BEIStateSequence<C> {
fn clone(&self) -> Self {
Self {
start_state: self.start_state,
diffs: self.diffs.clone(),
marker: core::marker::PhantomData,
}
}
}
impl<C> MapEntities for BEIStateSequence<C> {
fn map_entities<E: EntityMapper>(&mut self, entity_mapper: &mut E) {}
}
#[derive(Serialize, Deserialize, Clone, Copy, PartialEq, Debug)]
pub struct ActionsSnapshot {
pub state: TriggerState,
pub value: ActionValue,
pub time: ActionTime,
pub events: ActionEvents,
}
impl Default for ActionsSnapshot {
fn default() -> Self {
Self {
state: TriggerState::default(),
value: ActionValue::Bool(false),
time: ActionTime::default(),
events: ActionEvents::empty(),
}
}
}
#[derive(Clone, Copy, PartialEq, Debug, Serialize, Deserialize)]
struct ActionsDiff {
state: TriggerState,
value: ActionValue,
}
impl InputSnapshot for ActionsSnapshot {
fn decay_tick(&mut self, tick_duration: Duration) {
self.events = ActionEvents::new(self.state, self.state);
let delta_secs = tick_duration.as_secs_f32();
self.time.update(delta_secs, self.state);
}
}
#[derive(QueryData, Debug)]
#[query_data(mutable)]
pub struct ActionData {
state: &'static mut TriggerState,
value: &'static mut ActionValue,
events: &'static mut ActionEvents,
time: &'static mut ActionTime,
}
#[derive(Debug)]
pub struct ActionDataInnerItem<'w> {
pub state: &'w mut TriggerState,
pub value: &'w mut ActionValue,
pub events: &'w mut ActionEvents,
pub time: &'w mut ActionTime,
}
impl ActionStateQueryData for ActionData {
type Mut = ActionData;
type MutItemInner<'w> = ActionDataInnerItem<'w>;
type Main = TriggerState;
type Bundle = (TriggerState, ActionValue, ActionEvents, ActionTime);
#[inline]
fn as_read_only<'a, 'w: 'a, 's>(
state: &'a <Self::Mut as QueryData>::Item<'w, 's>,
) -> <<Self::Mut as QueryData>::ReadOnly as QueryData>::Item<'a, 's> {
ActionDataReadOnlyItem {
state: &state.state,
value: &state.value,
events: &state.events,
time: &state.time,
}
}
#[inline]
fn into_inner<'w, 's>(
mut_item: <Self::Mut as QueryData>::Item<'w, 's>,
) -> Self::MutItemInner<'w> {
ActionDataInnerItem {
state: mut_item.state.into_inner(),
value: mut_item.value.into_inner(),
events: mut_item.events.into_inner(),
time: mut_item.time.into_inner(),
}
}
#[inline]
fn as_mut(bundle: &mut Self::Bundle) -> Self::MutItemInner<'_> {
let (state, value, events, time) = bundle;
ActionDataInnerItem {
state,
value,
events,
time,
}
}
#[inline]
fn base_value() -> Self::Bundle {
(
TriggerState::default(),
ActionValue::Bool(false),
ActionEvents::empty(),
ActionTime::default(),
)
}
}
impl<C: Send + Sync + 'static> ActionStateSequence for BEIStateSequence<C> {
type Action = C;
type Snapshot = ActionsSnapshot;
type State = ActionData;
type Marker = InputMarker<C>;
fn len(&self) -> usize {
self.diffs.len() + 1
}
fn get_snapshots_from_message(
self,
tick_duration: Duration,
) -> impl Iterator<Item = Compressed<Self::Snapshot>> {
let start_iter = core::iter::once(Compressed::Input(self.start_state));
let diffs_iter = self.diffs.into_iter().scan(
self.start_state,
move |state: &mut ActionsSnapshot, diff: Compressed<ActionsDiff>| {
let (new_state, new_value) = match diff {
Compressed::Absent => return Some(Compressed::Absent),
Compressed::SameAsPrecedent => (state.state, state.value),
Compressed::Input(diff) => (diff.state, diff.value),
};
state.events = ActionEvents::new(state.state, new_state);
let delta_secs = tick_duration.as_secs_f32();
state.time.update(delta_secs, state.state);
state.state = new_state;
state.value = new_value;
Some(Compressed::Input(*state))
},
);
start_iter.chain(diffs_iter)
}
fn build_from_input_buffer<'w, 's>(
input_buffer: &InputBuffer<Self::Snapshot, Self::Action>,
num_ticks: u32,
end_tick: Tick,
) -> Option<Self> {
let mut diffs = Vec::new();
let mut start_tick = end_tick - num_ticks + 1;
while start_tick <= end_tick {
if input_buffer.get(start_tick).is_some() {
break;
}
start_tick += 1;
}
if start_tick > end_tick {
return None;
}
let start_state = *input_buffer.get(start_tick).unwrap();
let mut tick = start_tick + 1;
let (mut cur_state, mut cur_value) = (start_state.state, start_state.value);
while tick <= end_tick {
let diff = match input_buffer.get_raw(tick) {
Compressed::Absent => Compressed::Absent,
Compressed::SameAsPrecedent => Compressed::SameAsPrecedent,
Compressed::Input(snapshot) => {
let diff = if snapshot.state == cur_state && snapshot.value == cur_value {
Compressed::SameAsPrecedent
} else {
Compressed::Input(ActionsDiff {
state: snapshot.state,
value: snapshot.value,
})
};
cur_state = snapshot.state;
cur_value = snapshot.value;
diff
}
};
diffs.push(diff);
tick += 1;
}
Some(Self {
start_state,
diffs,
marker: core::marker::PhantomData,
})
}
fn to_snapshot<'w, 's>(state: ActionDataReadOnlyItem) -> Self::Snapshot {
ActionsSnapshot {
state: *state.state,
value: *state.value,
events: *state.events,
time: *state.time,
}
}
fn from_snapshot<'w, 's>(state: ActionDataInnerItem, snapshot: &Self::Snapshot) {
*state.state = snapshot.state;
*state.value = snapshot.value;
*state.events = snapshot.events;
*state.time = snapshot.time;
}
}
#[cfg(test)]
mod tests {
use super::*;
use alloc::vec;
use core::time::Duration;
use bevy_enhanced_input::prelude::InputAction;
use bevy_reflect::Reflect;
use test_log::test;
use tracing::trace;
struct Context1;
#[derive(InputAction, Debug, Clone, PartialEq, Eq, Hash, Reflect)]
#[action_output(bool)]
struct Action1;
#[test]
fn test_create_message() {
let mut input_buffer = BEIBuffer::<Context1>::default();
let mut state = ActionsSnapshot::default();
input_buffer.set(Tick(2), state);
state.state = TriggerState::Fired;
state.value = ActionValue::Bool(true);
input_buffer.set(Tick(3), state);
state.state = TriggerState::None;
state.value = ActionValue::Bool(false);
input_buffer.set(Tick(7), state);
let sequence =
BEIStateSequence::<Context1>::build_from_input_buffer(&input_buffer, 9, Tick(10))
.unwrap();
assert_eq!(
sequence,
BEIStateSequence::<Context1> {
start_state: ActionsSnapshot {
state: TriggerState::None,
value: ActionValue::Bool(false),
events: ActionEvents::empty(),
time: ActionTime::default(),
},
diffs: vec![
Compressed::Input(ActionsDiff {
state: TriggerState::Fired,
value: ActionValue::Bool(true)
}),
Compressed::SameAsPrecedent,
Compressed::SameAsPrecedent,
Compressed::SameAsPrecedent,
Compressed::Input(ActionsDiff {
state: TriggerState::None,
value: ActionValue::Bool(false)
}),
Compressed::Absent,
Compressed::Absent,
Compressed::Absent,
],
marker: Default::default(),
}
);
}
#[test]
fn test_build_from_input_buffer_empty() {
let input_buffer: BEIBuffer<Context1> = InputBuffer::default();
let sequence =
BEIStateSequence::<Context1>::build_from_input_buffer(&input_buffer, 5, Tick(10));
assert!(sequence.is_none());
}
#[test]
fn test_build_from_input_buffer_partial_overlap() {
let mut input_buffer = BEIBuffer::<Context1>::default();
let mut state = ActionsSnapshot::default();
input_buffer.set(Tick(8), state);
state.state = TriggerState::Fired;
state.value = ActionValue::Bool(true);
input_buffer.set(Tick(10), state);
let sequence =
BEIStateSequence::<Context1>::build_from_input_buffer(&input_buffer, 5, Tick(12))
.unwrap();
assert_eq!(sequence.len(), 5);
}
#[test]
fn test_update_buffer_extends_left_and_right() {
let mut input_buffer = BEIBuffer::<Context1>::default();
let state = ActionsSnapshot::default();
let sequence = BEIStateSequence::<Context1> {
start_state: state,
diffs: vec![Compressed::SameAsPrecedent, Compressed::Absent],
marker: Default::default(),
};
sequence.update_buffer(&mut input_buffer, Tick(7), Duration::default());
assert!(input_buffer.get(Tick(5)).is_some());
assert!(input_buffer.get(Tick(6)).is_some());
assert!(input_buffer.get(Tick(7)).is_none());
}
#[test]
fn test_update_buffer_empty_buffer() {
let mut input_buffer = BEIBuffer::<Context1>::default();
let mut state = ActionsSnapshot::default();
state.state = TriggerState::Fired;
state.value = ActionValue::Bool(true);
let sequence = BEIStateSequence::<Context1> {
start_state: state,
diffs: vec![Compressed::SameAsPrecedent, Compressed::Absent],
marker: Default::default(),
};
let earliest_mismatch =
sequence.update_buffer(&mut input_buffer, Tick(7), Duration::default());
trace!("Input buffer after update: {:?}", input_buffer);
assert_eq!(earliest_mismatch, Some(Tick(5)));
assert_eq!(input_buffer.start_tick, Some(Tick(5)));
assert_eq!(input_buffer.get(Tick(5)), Some(&state));
state.events = ActionEvents::FIRE;
assert_eq!(input_buffer.get(Tick(6)), Some(&state));
assert_eq!(input_buffer.get(Tick(7)), None);
}
#[test]
fn test_update_buffer_last_action_absent_new_action_present() {
let mut input_buffer = BEIBuffer::<Context1>::default();
let mut state = ActionsSnapshot::default();
input_buffer.set_empty(Tick(5));
input_buffer.last_remote_tick = Some(Tick(5));
state.state = TriggerState::Fired;
state.value = ActionValue::Bool(true);
let sequence = BEIStateSequence::<Context1> {
start_state: state,
diffs: vec![Compressed::SameAsPrecedent],
marker: Default::default(),
};
let earliest_mismatch =
sequence.update_buffer(&mut input_buffer, Tick(8), Duration::default());
assert_eq!(earliest_mismatch, Some(Tick(7)));
assert_eq!(input_buffer.get_raw(Tick(6)), &Compressed::SameAsPrecedent);
assert_eq!(input_buffer.get(Tick(7)), Some(&state));
state.events = ActionEvents::FIRE;
assert_eq!(input_buffer.get(Tick(8)), Some(&state));
}
#[test]
fn test_update_buffer_action_mismatch() {
let mut input_buffer = BEIBuffer::<Context1>::default();
let mut state = ActionsSnapshot::default();
state.state = TriggerState::Fired;
state.value = ActionValue::Bool(true);
input_buffer.set(Tick(5), state);
input_buffer.last_remote_tick = Some(Tick(5));
state.state = TriggerState::Ongoing;
state.value = ActionValue::Bool(false);
let sequence = BEIStateSequence::<Context1> {
start_state: state,
diffs: vec![Compressed::SameAsPrecedent],
marker: Default::default(),
};
let earliest_mismatch =
sequence.update_buffer(&mut input_buffer, Tick(7), Duration::default());
assert_eq!(earliest_mismatch, Some(Tick(6)));
assert_eq!(input_buffer.get(Tick(6)), Some(&state));
state.events = ActionEvents::ONGOING;
assert_eq!(input_buffer.get(Tick(7)), Some(&state));
}
#[test]
fn test_update_buffer_no_mismatch_same_action() {
let mut input_buffer = BEIBuffer::<Context1>::default();
let mut state = ActionsSnapshot::default();
state.state = TriggerState::Fired;
state.value = ActionValue::Bool(true);
input_buffer.set(Tick(5), state);
input_buffer.last_remote_tick = Some(Tick(5));
let mut snapshot = state;
snapshot.decay_tick(Duration::default());
let sequence = BEIStateSequence::<Context1> {
start_state: snapshot,
diffs: vec![Compressed::SameAsPrecedent],
marker: Default::default(),
};
let earliest_mismatch =
sequence.update_buffer(&mut input_buffer, Tick(7), Duration::default());
assert_eq!(earliest_mismatch, None);
assert_eq!(
input_buffer.get_raw(Tick(6)),
&Compressed::Input(snapshot.clone())
);
snapshot.decay_tick(Duration::default());
assert_eq!(input_buffer.get(Tick(7)), Some(&snapshot));
assert_eq!(input_buffer.get_raw(Tick(8)), &Compressed::Absent);
}
#[test]
fn test_update_buffer_keeps_matching_overlap_before_previous_end() {
let mut input_buffer = BEIBuffer::<Context1>::default();
let mut state = ActionsSnapshot::default();
state.state = TriggerState::Fired;
state.value = ActionValue::Bool(true);
let first_state = state;
let mut first_state_decayed = first_state;
first_state_decayed.decay_tick(Duration::default());
input_buffer.set(Tick(5), first_state);
input_buffer.set(Tick(6), first_state_decayed);
input_buffer.last_remote_tick = Some(Tick(6));
state.state = TriggerState::Ongoing;
state.value = ActionValue::Bool(false);
let mut second_state = state;
let sequence = BEIStateSequence::<Context1> {
start_state: first_state,
diffs: vec![
Compressed::SameAsPrecedent, Compressed::Input(ActionsDiff {
state: second_state.state,
value: second_state.value,
}), Compressed::SameAsPrecedent, ],
marker: Default::default(),
};
let earliest_mismatch =
sequence.update_buffer(&mut input_buffer, Tick(8), Duration::default());
assert_eq!(earliest_mismatch, Some(Tick(7)));
assert_eq!(input_buffer.get(Tick(5)), Some(&first_state));
assert_eq!(input_buffer.get(Tick(6)), Some(&first_state_decayed));
second_state.events = ActionEvents::ONGOING;
assert_eq!(input_buffer.get(Tick(7)), Some(&second_state));
assert_eq!(input_buffer.get(Tick(8)), Some(&second_state));
}
#[test]
fn test_update_buffer_keeps_received_overlap_before_last_remote_tick() {
let mut input_buffer = BEIBuffer::<Context1>::default();
let neutral = ActionsSnapshot {
value: ActionValue::Bool(false),
..Default::default()
};
let fired = ActionsSnapshot {
state: TriggerState::Fired,
value: ActionValue::Bool(true),
events: ActionEvents::FIRE,
..Default::default()
};
for tick in 10..=14 {
input_buffer.set(Tick(tick), neutral);
}
input_buffer.last_remote_tick = Some(Tick(14));
let sequence = BEIStateSequence::<Context1> {
start_state: fired,
diffs: vec![
Compressed::SameAsPrecedent,
Compressed::SameAsPrecedent,
Compressed::SameAsPrecedent,
Compressed::SameAsPrecedent,
],
marker: Default::default(),
};
let earliest_mismatch =
sequence.update_buffer(&mut input_buffer, Tick(15), Duration::default());
assert_eq!(earliest_mismatch, Some(Tick(15)));
for tick in 11..=14 {
assert_eq!(input_buffer.get(Tick(tick)), Some(&neutral));
}
assert_eq!(input_buffer.get(Tick(15)), Some(&fired));
assert_eq!(input_buffer.last_remote_tick, Some(Tick(15)));
}
#[test]
fn test_update_buffer_last_remote_tick_before_end_tick() {
let mut input_buffer = BEIBuffer::default();
let mut state = ActionsSnapshot::default();
state.state = TriggerState::Fired;
state.value = ActionValue::Bool(true);
let first_state = state;
input_buffer.set(Tick(6), first_state);
let mut first_state_decay = first_state;
first_state_decay.decay_tick(Duration::from_millis(10));
input_buffer.set(Tick(7), first_state_decay);
input_buffer.last_remote_tick = Some(Tick(6));
trace!("Input buffer before update: {:?}", input_buffer);
let sequence = BEIStateSequence::<Context1> {
start_state: first_state,
diffs: vec![
Compressed::SameAsPrecedent, ],
marker: Default::default(),
};
let earliest_mismatch =
sequence.update_buffer(&mut input_buffer, Tick(7), Duration::default());
assert_eq!(earliest_mismatch, Some(Tick(7)));
}
}