use std::any::{type_name, TypeId};
use std::mem::size_of;
use std::sync::{Arc, RwLock};
use crate::engine::component::{component_id_of, ComponentRegistry, Signature};
use crate::engine::error::{
ECSError, ECSResult, ExecutionError, InternalViolation, InvalidAccessReason, RegistryError,
};
use crate::engine::systems::AccessSets;
use crate::engine::types::ComponentID;
#[derive(Clone, Copy, Debug, Default)]
pub struct QuerySignature {
pub read: Signature,
pub write: Signature,
pub without: Signature,
}
impl QuerySignature {
pub fn requires_all(&self, archetype_signature: &Signature) -> bool {
archetype_signature.contains_all(&self.read)
&& archetype_signature.contains_all(&self.write)
&& archetype_signature
.components
.iter()
.zip(self.without.components.iter())
.all(|(arch_word, without_word)| (arch_word & without_word) == 0)
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct QueryComponent {
component_id: ComponentID,
type_id: TypeId,
type_name: &'static str,
size: usize,
}
impl QueryComponent {
#[inline]
fn of<T: 'static>(component_id: ComponentID) -> Self {
Self {
component_id,
type_id: TypeId::of::<T>(),
type_name: type_name::<T>(),
size: size_of::<T>(),
}
}
#[inline]
pub fn from_desc(desc: &super::component::ComponentDesc) -> ECSResult<Self> {
let Some(component_id) = desc.component_id else {
return Err(ECSError::from(RegistryError::NotRegistered {
type_id: desc.type_id,
}));
};
Ok(Self {
component_id,
type_id: desc.type_id,
type_name: desc.name,
size: desc.size,
})
}
#[inline]
pub fn component_id(&self) -> ComponentID {
self.component_id
}
#[inline]
pub fn type_id(&self) -> TypeId {
self.type_id
}
#[inline]
pub fn type_name(&self) -> &'static str {
self.type_name
}
#[inline]
pub fn size(&self) -> usize {
self.size
}
}
#[derive(Clone, Debug)]
pub struct BuiltQuery {
signature: QuerySignature,
reads: Vec<QueryComponent>,
writes: Vec<QueryComponent>,
read_ids: Vec<ComponentID>,
write_ids: Vec<ComponentID>,
}
impl BuiltQuery {
#[inline]
pub(crate) fn signature(&self) -> &QuerySignature {
&self.signature
}
#[inline]
pub fn read_ids(&self) -> &[ComponentID] {
&self.read_ids
}
#[inline]
pub fn write_ids(&self) -> &[ComponentID] {
&self.write_ids
}
#[inline]
pub fn reads(&self) -> &[QueryComponent] {
&self.reads
}
#[inline]
pub fn writes(&self) -> &[QueryComponent] {
&self.writes
}
#[inline]
pub fn access_sets(&self) -> AccessSets {
AccessSets {
read: self.signature.read,
write: self.signature.write,
produces: Default::default(),
consumes: Default::default(),
}
}
#[inline]
pub(crate) fn validate_read_type<T: 'static>(
&self,
index: usize,
method: &'static str,
) -> ECSResult<()> {
self.validate_type::<T>(AccessKindForQuery::Read, index, method)
}
#[inline]
pub(crate) fn validate_write_type<T: 'static>(
&self,
index: usize,
method: &'static str,
) -> ECSResult<()> {
self.validate_type::<T>(AccessKindForQuery::Write, index, method)
}
fn validate_type<T: 'static>(
&self,
access: AccessKindForQuery,
index: usize,
method: &'static str,
) -> ECSResult<()> {
let columns = match access {
AccessKindForQuery::Read => &self.reads,
AccessKindForQuery::Write => &self.writes,
};
let Some(column) = columns.get(index) else {
return Err(ECSError::from(InternalViolation::QueryShapeMismatch {
method,
expected_reads: self.reads.len(),
expected_writes: self.writes.len(),
}));
};
if column.type_id != TypeId::of::<T>() {
return Err(ECSError::from(ExecutionError::QueryTypeMismatch {
method,
access: access.into(),
index,
component_id: column.component_id,
expected: column.type_name,
actual: type_name::<T>(),
}));
}
Ok(())
}
}
#[derive(Clone, Copy)]
enum AccessKindForQuery {
Read,
Write,
}
impl From<AccessKindForQuery> for crate::engine::error::AccessKind {
fn from(value: AccessKindForQuery) -> Self {
match value {
AccessKindForQuery::Read => Self::Read,
AccessKindForQuery::Write => Self::Write,
}
}
}
enum RegistrySource {
Global,
Instance(Arc<RwLock<ComponentRegistry>>),
}
impl RegistrySource {
fn resolve<T: 'static + Send + Sync>(&self) -> ECSResult<ComponentID> {
match self {
RegistrySource::Global => component_id_of::<T>(),
RegistrySource::Instance(registry) => {
let registry = registry.read().map_err(|_| RegistryError::PoisonedLock)?;
Ok(registry.require_id_of::<T>()?)
}
}
}
}
pub struct QueryBuilder {
signature: QuerySignature,
reads: Vec<QueryComponent>,
writes: Vec<QueryComponent>,
registry_source: RegistrySource,
}
impl Default for QueryBuilder {
fn default() -> Self {
Self::new()
}
}
impl QueryBuilder {
pub fn new() -> Self {
Self {
signature: QuerySignature::default(),
reads: vec![],
writes: vec![],
registry_source: RegistrySource::Global,
}
}
pub fn with_registry(registry: Arc<RwLock<ComponentRegistry>>) -> Self {
Self {
signature: QuerySignature::default(),
reads: vec![],
writes: vec![],
registry_source: RegistrySource::Instance(registry),
}
}
pub fn read<T: 'static + Send + Sync>(mut self) -> ECSResult<Self> {
let id = self.registry_source.resolve::<T>()?;
self.signature.read.set(id);
self.reads.push(QueryComponent::of::<T>(id));
Ok(self)
}
pub fn write<T: 'static + Send + Sync>(mut self) -> ECSResult<Self> {
let id = self.registry_source.resolve::<T>()?;
self.signature.write.set(id);
self.writes.push(QueryComponent::of::<T>(id));
Ok(self)
}
pub fn without<T: 'static + Send + Sync>(mut self) -> ECSResult<Self> {
let id = self.registry_source.resolve::<T>()?;
self.signature.without.set(id);
Ok(self)
}
pub fn read_id(mut self, desc: &super::component::ComponentDesc) -> ECSResult<Self> {
let component = QueryComponent::from_desc(desc)?;
self.signature.read.set(component.component_id());
self.reads.push(component);
Ok(self)
}
pub fn write_id(mut self, desc: &super::component::ComponentDesc) -> ECSResult<Self> {
let component = QueryComponent::from_desc(desc)?;
self.signature.write.set(component.component_id());
self.writes.push(component);
Ok(self)
}
pub fn build(self) -> ECSResult<BuiltQuery> {
let read_ids: Vec<ComponentID> = self
.reads
.iter()
.map(QueryComponent::component_id)
.collect();
let write_ids: Vec<ComponentID> = self
.writes
.iter()
.map(QueryComponent::component_id)
.collect();
let mut reads_sorted = read_ids.clone();
let mut writes_sorted = write_ids.clone();
reads_sorted.sort_unstable();
writes_sorted.sort_unstable();
let check_duplicates = |sorted: &[ComponentID]| -> ECSResult<()> {
for w in sorted.windows(2) {
if w[0] == w[1] {
return Err(ECSError::Execute(ExecutionError::InvalidQueryAccess {
component_id: w[0],
reason: InvalidAccessReason::DuplicateAccess,
}));
}
}
Ok(())
};
check_duplicates(&reads_sorted)?;
check_duplicates(&writes_sorted)?;
reads_sorted.dedup();
writes_sorted.dedup();
for component_id in &reads_sorted {
if writes_sorted.binary_search(component_id).is_ok() {
return Err(ECSError::Execute(ExecutionError::InvalidQueryAccess {
component_id: *component_id,
reason: InvalidAccessReason::ReadAndWrite,
}));
}
}
for (word_idx, (&w_word, &without_word)) in self
.signature
.write
.components
.iter()
.zip(self.signature.without.components.iter())
.enumerate()
{
let overlap = w_word & without_word;
if overlap != 0 {
let bit = overlap.trailing_zeros();
let component_id = (word_idx as u32) * 64 + bit;
return Err(ECSError::Execute(ExecutionError::InvalidQueryAccess {
component_id: component_id as ComponentID,
reason: InvalidAccessReason::WriteAndWithout,
}));
}
}
for (word_idx, (&r_word, &without_word)) in self
.signature
.read
.components
.iter()
.zip(self.signature.without.components.iter())
.enumerate()
{
let overlap = r_word & without_word;
if overlap != 0 {
let bit = overlap.trailing_zeros();
let component_id = (word_idx as u32) * 64 + bit;
return Err(ECSError::Execute(ExecutionError::InvalidQueryAccess {
component_id: component_id as ComponentID,
reason: InvalidAccessReason::ReadAndWithout,
}));
}
}
Ok(BuiltQuery {
signature: self.signature,
reads: self.reads,
writes: self.writes,
read_ids,
write_ids,
})
}
pub fn access_sets(&self) -> AccessSets {
AccessSets {
read: self.signature.read,
write: self.signature.write,
produces: Default::default(),
consumes: Default::default(),
}
}
}