use std::cell::UnsafeCell;
use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, LazyLock, Mutex, Weak};
use super::aligned_buffer::AlignedBuffer;
use super::error::MessagingError;
use super::message::{Message, MessageTypeID};
use crate::engine::error::{ECSError, ECSResult};
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub(crate) struct MessageRuntimeID(u64);
static NEXT_RUNTIME_ID: AtomicU64 = AtomicU64::new(1);
pub(crate) fn alloc_runtime_id() -> MessageRuntimeID {
MessageRuntimeID(NEXT_RUNTIME_ID.fetch_add(1, Ordering::Relaxed))
}
pub(crate) struct WorkerEmitSlots {
slots: UnsafeCell<Vec<Option<AlignedBuffer>>>,
worker_id: u32,
}
unsafe impl Send for WorkerEmitSlots {}
unsafe impl Sync for WorkerEmitSlots {}
impl WorkerEmitSlots {
fn new(num_message_types: usize, worker_id: u32) -> Self {
let slots: Vec<Option<AlignedBuffer>> = (0..num_message_types).map(|_| None).collect();
WorkerEmitSlots {
slots: UnsafeCell::new(slots),
worker_id,
}
}
}
static GLOBAL_EMIT_REGISTRY: LazyLock<Mutex<HashMap<MessageRuntimeID, Vec<Arc<WorkerEmitSlots>>>>> =
LazyLock::new(|| Mutex::new(HashMap::new()));
thread_local! {
static THIS_WORKER: std::cell::RefCell<HashMap<MessageRuntimeID, Weak<WorkerEmitSlots>>> =
std::cell::RefCell::new(HashMap::new());
}
pub(crate) fn ensure_worker_registered_fallible(
runtime_id: MessageRuntimeID,
num_message_types: usize,
) -> ECSResult<Arc<WorkerEmitSlots>> {
THIS_WORKER.with(|cell| {
let existing = {
let workers = cell.borrow();
workers.get(&runtime_id).and_then(Weak::upgrade)
};
if let Some(existing) = existing {
return Ok(existing);
}
let worker_id = crate::engine::workers::worker_id();
let slots = Arc::new(WorkerEmitSlots::new(num_message_types, worker_id));
cell.borrow_mut().insert(runtime_id, Arc::downgrade(&slots));
let mut registry = GLOBAL_EMIT_REGISTRY
.lock()
.map_err(|_| ECSError::from(MessagingError::LockPoisoned("global emit registry")))?;
let workers = registry.entry(runtime_id).or_default();
let at = workers.partition_point(|w| w.worker_id < worker_id);
workers.insert(at, Arc::clone(&slots));
drop(registry);
Ok(slots)
})
}
pub(crate) fn deregister_runtime(runtime_id: MessageRuntimeID) {
if let Ok(mut registry) = GLOBAL_EMIT_REGISTRY.lock() {
registry.remove(&runtime_id);
}
THIS_WORKER.with(|cell| {
cell.borrow_mut().remove(&runtime_id);
});
}
#[cfg(test)]
pub(crate) fn registered_worker_count_for_test(runtime_id: MessageRuntimeID) -> usize {
GLOBAL_EMIT_REGISTRY
.lock()
.ok()
.and_then(|registry| registry.get(&runtime_id).map(Vec::len))
.unwrap_or(0)
}
#[cfg(test)]
pub(crate) fn current_thread_has_worker_for_test(runtime_id: MessageRuntimeID) -> bool {
THIS_WORKER.with(|cell| cell.borrow().contains_key(&runtime_id))
}
pub(crate) fn emit<M: Message>(
runtime_id: MessageRuntimeID,
num_message_types: usize,
mtid: MessageTypeID,
item_size: usize,
item_align: usize,
capacity: usize,
msg: M,
) -> ECSResult<()> {
let worker = ensure_worker_registered_fallible(runtime_id, num_message_types)?;
let slots = unsafe { &mut *worker.slots.get() };
if slots[mtid.index()].is_none() {
slots[mtid.index()] = Some(AlignedBuffer::with_capacity(
item_size, item_align, capacity,
));
}
let buf = slots[mtid.index()].as_mut().unwrap();
unsafe { buf.push(msg) };
Ok(())
}
type EmitterMarker<M> = std::marker::PhantomData<(*mut (), fn() -> M)>;
pub struct MessageEmitter<'a, M: Message> {
worker: Arc<WorkerEmitSlots>,
slot_index: usize,
item_size: usize,
item_align: usize,
initial_capacity: usize,
_owner: std::marker::PhantomData<&'a ()>,
_marker: EmitterMarker<M>,
}
impl<'a, M: Message> MessageEmitter<'a, M> {
pub(crate) fn new(
runtime_id: MessageRuntimeID,
num_message_types: usize,
mtid: MessageTypeID,
item_size: usize,
item_align: usize,
initial_capacity: usize,
) -> ECSResult<Self> {
let worker = ensure_worker_registered_fallible(runtime_id, num_message_types)?;
Ok(Self {
worker,
slot_index: mtid.index(),
item_size,
item_align,
initial_capacity,
_owner: std::marker::PhantomData,
_marker: std::marker::PhantomData,
})
}
#[inline]
pub fn emit(&self, msg: M) {
let slots = unsafe { &mut *self.worker.slots.get() };
if slots[self.slot_index].is_none() {
slots[self.slot_index] = Some(AlignedBuffer::with_capacity(
self.item_size,
self.item_align,
self.initial_capacity,
));
}
let buffer = slots[self.slot_index]
.as_mut()
.expect("slot initialised above");
unsafe { buffer.push(msg) };
}
}
pub(crate) fn drain_into(
runtime_id: MessageRuntimeID,
mtid: MessageTypeID,
out: &mut AlignedBuffer,
) -> ECSResult<()> {
let registry = GLOBAL_EMIT_REGISTRY
.lock()
.map_err(|_| ECSError::from(MessagingError::LockPoisoned("global emit registry")))?;
let Some(workers) = registry.get(&runtime_id) else {
return Ok(());
};
for worker in workers {
let slots = unsafe { &mut *worker.slots.get() };
if let Some(ref buf) = slots[mtid.index()] {
out.extend_from(buf);
}
}
Ok(())
}
pub(crate) fn clear_for_tick(runtime_id: MessageRuntimeID, mtid: MessageTypeID) -> ECSResult<()> {
let registry = GLOBAL_EMIT_REGISTRY
.lock()
.map_err(|_| ECSError::from(MessagingError::LockPoisoned("global emit registry")))?;
let Some(workers) = registry.get(&runtime_id) else {
return Ok(());
};
for worker in workers {
let slots = unsafe { &mut *worker.slots.get() };
if let Some(ref mut buf) = slots[mtid.index()] {
buf.clear();
}
}
Ok(())
}