use bevy_ecs::archetype::ArchetypeEntity;
use bevy_ecs::component::{ComponentId, StorageType};
use bevy_ecs::ptr::{Ptr, PtrMut};
use bevy_ecs::query::{IterQueryData, QueryFilter, QueryItem};
use bevy_ecs::storage::TableId;
use bevy_ecs::system::Query;
use bevy_ecs::world::unsafe_world_cell::UnsafeWorldCell;
#[macro_export]
macro_rules! adaptive_for_each_mut {
($query:expr, |$item:pat_param| $body:expr $(,)?) => {
$crate::ecs::AdaptiveQueryIterMut::new(&mut $query, 1).for_each(|$item| $body)
};
($query:expr, $serial_threshold:expr, |$item:pat_param| $body:expr $(,)?) => {
$crate::ecs::AdaptiveQueryIterMut::new(&mut $query, $serial_threshold)
.for_each(|$item| $body)
};
($query:expr $(,)?) => {
$crate::ecs::AdaptiveQueryIterMut::new(&mut $query, 1)
};
}
#[doc(hidden)]
pub struct AdaptiveQueryIterMut<'query, 'world, 'state, D, F>
where
D: IterQueryData,
F: QueryFilter,
{
query: &'query mut Query<'world, 'state, D, F>,
serial_threshold: usize,
}
impl<'query, 'world, 'state, D, F> AdaptiveQueryIterMut<'query, 'world, 'state, D, F>
where
D: IterQueryData,
F: QueryFilter,
{
#[doc(hidden)]
pub fn new(query: &'query mut Query<'world, 'state, D, F>, serial_threshold: usize) -> Self {
Self {
query,
serial_threshold,
}
}
pub fn for_each<Func>(self, func: Func)
where
Func: for<'item> Fn(QueryItem<'item, 'state, D>) + Send + Sync + Clone,
{
if self.query.iter().nth(self.serial_threshold).is_some() {
self.query.par_iter_mut().for_each(func);
} else {
self.query.iter_mut().for_each(func);
}
}
}
pub unsafe fn get_component_unchecked_mut<'w>(
unsafe_world_cell: UnsafeWorldCell<'w>,
entity: &'w ArchetypeEntity,
table_id: TableId,
storage: StorageType,
component_id: ComponentId,
) -> PtrMut<'w> {
let storages = unsafe { unsafe_world_cell.storages() };
match storage {
StorageType::Table => unsafe {
let table = storages.tables.get(table_id).unwrap_unchecked();
table
.get_component(component_id, entity.table_row())
.unwrap_unchecked()
.assert_unique()
},
StorageType::SparseSet => unsafe {
let sparse_set = storages.sparse_sets.get(component_id).unwrap_unchecked();
sparse_set
.get(entity.id())
.unwrap_unchecked()
.assert_unique()
},
}
}
pub unsafe fn get_component_unchecked<'w>(
unsafe_world_cell: UnsafeWorldCell<'w>,
entity: &'w ArchetypeEntity,
table_id: TableId,
storage: StorageType,
component_id: ComponentId,
) -> Ptr<'w> {
let storages = unsafe { unsafe_world_cell.storages() };
match storage {
StorageType::Table => unsafe {
let table = storages.tables.get(table_id).unwrap_unchecked();
table
.get_component(component_id, entity.table_row())
.unwrap_unchecked()
},
StorageType::SparseSet => unsafe {
let sparse_set = storages.sparse_sets.get(component_id).unwrap_unchecked();
sparse_set.get(entity.id()).unwrap_unchecked()
},
}
}
#[cfg(test)]
mod tests {
use bevy_app::{App, TaskPoolPlugin};
use bevy_ecs::prelude::*;
#[derive(Component)]
struct Value(u32);
#[test]
fn adaptive_iteration_supports_default_and_custom_thresholds() {
let mut app = App::new();
app.add_plugins(TaskPoolPlugin::default());
let world = app.world_mut();
let first = world.spawn(Value(0)).id();
let mut query_state = world.query::<&mut Value>();
{
let mut query = query_state.query_mut(world);
crate::adaptive_for_each_mut!(query, |mut value| value.0 += 1);
}
let second = world.spawn(Value(0)).id();
{
let mut query = query_state.query_mut(world);
crate::adaptive_for_each_mut!(query, 2, |mut value| value.0 += 1);
}
{
let mut query = query_state.query_mut(world);
crate::adaptive_for_each_mut!(query, |mut value| value.0 += 1);
}
assert_eq!(world.get::<Value>(first).unwrap().0, 3);
assert_eq!(world.get::<Value>(second).unwrap().0, 2);
}
}