use std::any::TypeId;
use std::ops::{Deref, DerefMut};
use crate::component::Component;
use crate::entity::Entity;
use crate::query::{Mut, QueryIter, QueryIterMut};
use crate::world::UnsafeWorldCell;
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum Access {
ResRead(TypeId),
ResWrite(TypeId),
CompRead(TypeId),
CompWrite(TypeId),
}
impl Access {
pub fn conflicts_with(&self, other: &Access) -> bool {
match (self, other) {
(Access::ResRead(a), Access::ResWrite(b))
| (Access::ResWrite(a), Access::ResRead(b))
| (Access::ResWrite(a), Access::ResWrite(b)) => a == b,
(Access::CompRead(a), Access::CompWrite(b))
| (Access::CompWrite(a), Access::CompRead(b))
| (Access::CompWrite(a), Access::CompWrite(b)) => a == b,
_ => false,
}
}
}
pub fn has_conflicts(a: &[Access], b: &[Access]) -> bool {
a.iter().any(|x| b.iter().any(|y| x.conflicts_with(y)))
}
pub unsafe trait SystemParam {
type Item<'w>;
fn access() -> Vec<Access>;
unsafe fn fetch<'w>(world: UnsafeWorldCell) -> Self::Item<'w>;
}
pub struct Res<'w, T: Send + 'static> {
value: &'w T,
}
impl<T: Send + 'static> Deref for Res<'_, T> {
type Target = T;
fn deref(&self) -> &T {
self.value
}
}
unsafe impl<T: Send + 'static> SystemParam for Res<'_, T> {
type Item<'w> = Res<'w, T>;
fn access() -> Vec<Access> {
vec![Access::ResRead(TypeId::of::<T>())]
}
unsafe fn fetch<'w>(world: UnsafeWorldCell) -> Res<'w, T> {
Res {
value: unsafe { world.get_resource::<T>() },
}
}
}
pub struct ResMut<'w, T: Send + 'static> {
value: &'w mut T,
}
impl<T: Send + 'static> Deref for ResMut<'_, T> {
type Target = T;
fn deref(&self) -> &T {
self.value
}
}
impl<T: Send + 'static> DerefMut for ResMut<'_, T> {
fn deref_mut(&mut self) -> &mut T {
self.value
}
}
unsafe impl<T: Send + 'static> SystemParam for ResMut<'_, T> {
type Item<'w> = ResMut<'w, T>;
fn access() -> Vec<Access> {
vec![Access::ResWrite(TypeId::of::<T>())]
}
unsafe fn fetch<'w>(world: UnsafeWorldCell) -> ResMut<'w, T> {
ResMut {
value: unsafe { world.get_resource_mut::<T>() },
}
}
}
pub struct Query<'w, T: Component> {
results: Vec<(Entity, &'w T)>,
}
impl<'w, T: Component> Query<'w, T> {
pub fn iter(&self) -> impl Iterator<Item = (Entity, &T)> {
self.results.iter().map(|&(e, v)| (e, v))
}
pub fn is_empty(&self) -> bool {
self.results.is_empty()
}
pub fn len(&self) -> usize {
self.results.len()
}
}
unsafe impl<T: Component> SystemParam for Query<'_, T> {
type Item<'w> = Query<'w, T>;
fn access() -> Vec<Access> {
vec![Access::CompRead(TypeId::of::<T>())]
}
unsafe fn fetch<'w>(world: UnsafeWorldCell) -> Query<'w, T> {
Query {
results: unsafe { QueryIter::<'w, &T>::new(world.archetypes()).collect() },
}
}
}
pub struct QueryMut<'w, T: Component> {
results: Vec<(Entity, Mut<'w, T>)>,
}
impl<'w, T: Component> QueryMut<'w, T> {
pub fn iter_mut(&mut self) -> impl Iterator<Item = (Entity, &mut Mut<'w, T>)> + '_ {
self.results.iter_mut().map(|(e, v)| (*e, v))
}
pub fn is_empty(&self) -> bool {
self.results.is_empty()
}
pub fn len(&self) -> usize {
self.results.len()
}
}
unsafe impl<T: Component> SystemParam for QueryMut<'_, T> {
type Item<'w> = QueryMut<'w, T>;
fn access() -> Vec<Access> {
vec![Access::CompWrite(TypeId::of::<T>())]
}
unsafe fn fetch<'w>(world: UnsafeWorldCell) -> QueryMut<'w, T> {
QueryMut {
results: unsafe {
let tick = world.change_tick();
QueryIterMut::<'w, &mut T>::new_from_ptr(world.archetypes_mut_ptr(), tick).collect()
},
}
}
}
unsafe impl SystemParam for () {
type Item<'w> = ();
fn access() -> Vec<Access> {
Vec::new()
}
unsafe fn fetch<'w>(_world: UnsafeWorldCell) -> Self::Item<'w> {}
}
macro_rules! impl_system_param_tuple {
($($P:ident),+) => {
unsafe impl<$($P: SystemParam),+> SystemParam for ($($P,)+) {
type Item<'w> = ($($P::Item<'w>,)+);
fn access() -> Vec<Access> {
let mut acc = Vec::new();
$(acc.extend($P::access());)+
acc
}
unsafe fn fetch<'w>(world: UnsafeWorldCell) -> Self::Item<'w> {
($(unsafe { $P::fetch(world) },)+)
}
}
};
}
impl_system_param_tuple!(P0);
impl_system_param_tuple!(P0, P1);
impl_system_param_tuple!(P0, P1, P2);
impl_system_param_tuple!(P0, P1, P2, P3);
impl_system_param_tuple!(P0, P1, P2, P3, P4);
impl_system_param_tuple!(P0, P1, P2, P3, P4, P5);
impl_system_param_tuple!(P0, P1, P2, P3, P4, P5, P6);
impl_system_param_tuple!(P0, P1, P2, P3, P4, P5, P6, P7);
#[cfg(test)]
mod tests {
use super::*;
use crate::component::Component;
use crate::world::World;
fn type_id<T: 'static>() -> TypeId {
TypeId::of::<T>()
}
#[derive(Debug, PartialEq)]
struct Pos {
x: f32,
}
impl Component for Pos {}
#[test]
fn read_read_same_type_no_conflict() {
let a = Access::ResRead(type_id::<u32>());
let b = Access::ResRead(type_id::<u32>());
assert!(!a.conflicts_with(&b));
}
#[test]
fn read_write_same_type_conflicts() {
let a = Access::ResRead(type_id::<u32>());
let b = Access::ResWrite(type_id::<u32>());
assert!(a.conflicts_with(&b));
assert!(b.conflicts_with(&a));
}
#[test]
fn write_write_same_type_conflicts() {
let a = Access::ResWrite(type_id::<u32>());
let b = Access::ResWrite(type_id::<u32>());
assert!(a.conflicts_with(&b));
}
#[test]
fn read_write_different_type_no_conflict() {
let a = Access::ResRead(type_id::<u32>());
let b = Access::ResWrite(type_id::<f32>());
assert!(!a.conflicts_with(&b));
}
#[test]
fn comp_read_write_conflicts() {
let a = Access::CompRead(type_id::<u64>());
let b = Access::CompWrite(type_id::<u64>());
assert!(a.conflicts_with(&b));
assert!(b.conflicts_with(&a));
}
#[test]
fn res_and_comp_same_type_no_conflict() {
let a = Access::ResWrite(type_id::<u32>());
let b = Access::CompWrite(type_id::<u32>());
assert!(!a.conflicts_with(&b));
}
#[test]
fn has_conflicts_finds_conflict_in_sets() {
let set_a = vec![
Access::ResRead(type_id::<u32>()),
Access::CompRead(type_id::<f32>()),
];
let set_b = vec![
Access::ResWrite(type_id::<u32>()),
Access::CompRead(type_id::<f32>()),
];
assert!(has_conflicts(&set_a, &set_b));
}
#[test]
fn has_conflicts_empty_sets_no_conflict() {
assert!(!has_conflicts(&[], &[]));
assert!(!has_conflicts(&[Access::ResRead(type_id::<u32>())], &[]));
assert!(!has_conflicts(&[], &[Access::ResWrite(type_id::<u32>())]));
}
#[test]
fn res_fetches_resource() {
let mut world = World::new();
world.insert_resource(42_i32);
let cell = unsafe { UnsafeWorldCell::new(&mut world as *mut World) };
unsafe {
let res: Res<'_, i32> = <Res<'_, i32> as SystemParam>::fetch(cell);
assert_eq!(*res, 42);
}
}
#[test]
fn res_access_is_read() {
let access = <Res<'_, i32> as SystemParam>::access();
assert_eq!(access, vec![Access::ResRead(TypeId::of::<i32>())]);
}
#[test]
fn res_mut_fetches_and_mutates() {
let mut world = World::new();
world.insert_resource(10_u32);
let cell = unsafe { UnsafeWorldCell::new(&mut world as *mut World) };
unsafe {
let mut res: ResMut<'_, u32> = <ResMut<'_, u32> as SystemParam>::fetch(cell);
*res = 20;
}
assert_eq!(*world.resource::<u32>(), 20);
}
#[test]
fn res_mut_access_is_write() {
let access = <ResMut<'_, u32> as SystemParam>::access();
assert_eq!(access, vec![Access::ResWrite(TypeId::of::<u32>())]);
}
#[test]
fn res_and_res_mut_different_types_no_conflict() {
let a = <Res<'_, i32> as SystemParam>::access();
let b = <ResMut<'_, u32> as SystemParam>::access();
assert!(!has_conflicts(&a, &b));
}
#[test]
fn res_and_res_mut_same_type_conflicts() {
let a = <Res<'_, i32> as SystemParam>::access();
let b = <ResMut<'_, i32> as SystemParam>::access();
assert!(has_conflicts(&a, &b));
}
#[test]
fn query_fetches_matching_entities() {
let mut world = World::new();
world.spawn((Pos { x: 1.0 },));
world.spawn((Pos { x: 2.0 },));
let cell = unsafe { UnsafeWorldCell::new(&mut world as *mut World) };
unsafe {
let q: Query<'_, Pos> = <Query<'_, Pos> as SystemParam>::fetch(cell);
assert_eq!(q.len(), 2);
}
}
#[test]
fn query_access_is_comp_read() {
let access = <Query<'_, Pos> as SystemParam>::access();
assert_eq!(access, vec![Access::CompRead(TypeId::of::<Pos>())]);
}
#[test]
fn query_mut_allows_mutation() {
let mut world = World::new();
let e = world.spawn((Pos { x: 5.0 },));
let cell = unsafe { UnsafeWorldCell::new(&mut world as *mut World) };
unsafe {
let mut q: QueryMut<'_, Pos> = <QueryMut<'_, Pos> as SystemParam>::fetch(cell);
for (_, pos) in q.iter_mut() {
pos.x += 10.0;
}
}
assert_eq!(world.get::<Pos>(e).unwrap().x, 15.0);
}
#[test]
fn query_mut_access_is_comp_write() {
let access = <QueryMut<'_, Pos> as SystemParam>::access();
assert_eq!(access, vec![Access::CompWrite(TypeId::of::<Pos>())]);
}
#[test]
fn query_empty_world() {
let mut world = World::new();
let cell = unsafe { UnsafeWorldCell::new(&mut world as *mut World) };
unsafe {
let q: Query<'_, Pos> = <Query<'_, Pos> as SystemParam>::fetch(cell);
assert!(q.is_empty());
}
}
#[test]
fn unit_tuple_has_no_access() {
let access = <() as SystemParam>::access();
assert!(access.is_empty());
}
#[test]
fn pair_tuple_aggregates_access() {
let access = <(Res<'_, i32>, ResMut<'_, u32>) as SystemParam>::access();
assert_eq!(access.len(), 2);
assert!(access.contains(&Access::ResRead(TypeId::of::<i32>())));
assert!(access.contains(&Access::ResWrite(TypeId::of::<u32>())));
}
#[test]
fn triple_tuple_aggregates_access() {
let access = <(Res<'_, i32>, ResMut<'_, u32>, Query<'_, Pos>) as SystemParam>::access();
assert_eq!(access.len(), 3);
}
#[test]
#[should_panic(expected = "resource not found")]
fn res_fetch_panics_on_missing_resource() {
let mut world = World::new();
let cell = unsafe { UnsafeWorldCell::new(&mut world as *mut World) };
unsafe {
let _: Res<'_, i32> = <Res<'_, i32> as SystemParam>::fetch(cell);
}
}
#[test]
fn query_mut_empty_world() {
let mut world = World::new();
let cell = unsafe { UnsafeWorldCell::new(&mut world as *mut World) };
unsafe {
let q: QueryMut<'_, Pos> = <QueryMut<'_, Pos> as SystemParam>::fetch(cell);
assert!(q.is_empty());
}
}
#[derive(Debug, PartialEq)]
struct Vel {
y: f32,
}
impl Component for Vel {}
struct TimeRes(f32);
struct GravRes(f32);
#[test]
fn four_arity_tuple_access_and_fetch() {
let mut world = World::new();
world.insert_resource(TimeRes(1.0));
world.insert_resource(GravRes(9.8));
world.spawn((Pos { x: 0.0 },));
world.spawn((Vel { y: 0.0 },));
let access = <(
Res<'_, TimeRes>,
Res<'_, GravRes>,
Query<'_, Pos>,
Query<'_, Vel>,
) as SystemParam>::access();
assert_eq!(access.len(), 4);
let cell = unsafe { UnsafeWorldCell::new(&mut world as *mut World) };
unsafe {
let (time, grav, positions, velocities) = <(
Res<'_, TimeRes>,
Res<'_, GravRes>,
Query<'_, Pos>,
Query<'_, Vel>,
) as SystemParam>::fetch(cell);
assert!((time.0 - 1.0).abs() < f32::EPSILON);
assert!((grav.0 - 9.8).abs() < f32::EPSILON);
assert_eq!(positions.len(), 1);
assert_eq!(velocities.len(), 1);
}
}
#[test]
fn query_read_and_query_mut_different_types() {
let mut world = World::new();
world.spawn((Pos { x: 1.0 }, Vel { y: 2.0 }));
world.spawn((Pos { x: 3.0 }, Vel { y: 4.0 }));
let a = <Query<'_, Pos> as SystemParam>::access();
let b = <QueryMut<'_, Vel> as SystemParam>::access();
assert!(!has_conflicts(&a, &b));
let cell = unsafe { UnsafeWorldCell::new(&mut world as *mut World) };
unsafe {
let positions: Query<'_, Pos> = <Query<'_, Pos> as SystemParam>::fetch(cell);
let mut velocities: QueryMut<'_, Vel> = <QueryMut<'_, Vel> as SystemParam>::fetch(cell);
assert_eq!(positions.len(), 2);
assert_eq!(velocities.len(), 2);
for (_, v) in velocities.iter_mut() {
v.y += 10.0;
}
}
let ys: Vec<f32> = world.query::<&Vel>().map(|(_, v)| v.y).collect();
assert!(ys.iter().all(|&y| y > 10.0));
}
}