use core::{
cell::{Cell, UnsafeCell},
marker::PhantomData,
};
use crate::session::types::{Lane, SessionId};
mod fault;
#[cfg(kani)]
mod kani;
mod storage;
pub(crate) use fault::SessionFaultKind;
const ENTRY_COUNT_BITS: u32 = 9;
const ENTRY_COUNT_MASK: u16 = (1u16 << ENTRY_COUNT_BITS) - 1;
const ENTRY_FAULT_SHIFT: u32 = ENTRY_COUNT_BITS;
const ENTRY_COUNT_MAX: u16 = u8::MAX as u16 + 1;
#[inline]
fn next_attachment_count(current: u16) -> Option<u16> {
if current == 0 || current > ENTRY_COUNT_MAX {
crate::invariant();
}
if current == ENTRY_COUNT_MAX {
None
} else {
Some(current + 1)
}
}
pub(super) struct AssocTable {
lane_base: Cell<u32>,
lane_slots: Cell<u16>,
assoc_slots: Cell<u16>,
entry_sids: UnsafeCell<*mut SessionId>,
entry_lanes: UnsafeCell<*mut u8>,
entry_states: UnsafeCell<*mut u16>,
_no_send_sync: PhantomData<*mut ()>,
}
impl AssocTable {
#[inline]
const fn entry_count(raw: u16) -> u16 {
raw & ENTRY_COUNT_MASK
}
#[inline]
const fn entry_fault_code(raw: u16) -> u8 {
(raw >> ENTRY_FAULT_SHIFT) as u8
}
#[inline]
const fn entry_fault(raw: u16) -> Option<SessionFaultKind> {
SessionFaultKind::decode(Self::entry_fault_code(raw))
}
#[inline]
const fn entry_state(count: u16, fault: u8) -> u16 {
if count > ENTRY_COUNT_MAX || fault as u16 > (u16::MAX >> ENTRY_FAULT_SHIFT) {
crate::invariant();
}
((fault as u16) << ENTRY_FAULT_SHIFT) | count
}
const EMPTY_ENTRY_STATE: u16 = Self::entry_state(0, SessionFaultKind::ABSENT_CODE);
const LIVE_ENTRY_STATE: u16 = Self::entry_state(1, SessionFaultKind::ABSENT_CODE);
#[inline]
fn lane_slots(&self) -> usize {
self.lane_slots.get() as usize
}
#[inline]
pub(super) fn assoc_slots(&self) -> usize {
self.assoc_slots.get() as usize
}
#[inline]
fn entry_sids_ptr(&self) -> *mut SessionId {
unsafe { *self.entry_sids.get() }
}
#[inline]
fn entry_lanes_ptr(&self) -> *mut u8 {
unsafe { *self.entry_lanes.get() }
}
#[inline]
fn entry_states_ptr(&self) -> *mut u16 {
unsafe { *self.entry_states.get() }
}
#[inline]
fn lane_offset(&self, lane: Lane) -> Option<u8> {
let lane_raw = lane.raw();
if lane_raw < self.lane_base.get() {
return None;
}
let offset = lane_raw - self.lane_base.get();
if (offset as usize) >= self.lane_slots() || offset > u8::MAX as u32 {
return None;
}
Some(offset as u8)
}
#[inline]
fn find_entry_by_offset(&self, lane_offset: u8, sid: SessionId) -> Option<usize> {
unsafe {
let lanes = self.entry_lanes_ptr();
let sids = self.entry_sids_ptr();
let states = self.entry_states_ptr();
let mut idx = 0usize;
while idx < self.assoc_slots() {
if Self::entry_count(*states.add(idx)) != 0
&& *lanes.add(idx) == lane_offset
&& *sids.add(idx) == sid
{
return Some(idx);
}
idx += 1;
}
}
None
}
#[inline]
pub(super) fn active_entry_count(&self) -> usize {
unsafe {
let states = self.entry_states_ptr();
let mut idx = 0usize;
let mut live = 0usize;
while idx < self.assoc_slots() {
if Self::entry_count(*states.add(idx)) != 0 {
live += 1;
}
idx += 1;
}
live
}
}
#[inline]
pub(super) fn active_lane_slots(&self) -> usize {
unsafe {
let lanes = self.entry_lanes_ptr();
let states = self.entry_states_ptr();
let mut idx = 0usize;
let mut required = 0usize;
while idx < self.assoc_slots() {
if Self::entry_count(*states.add(idx)) != 0 {
required = required.max(*lanes.add(idx) as usize + 1);
}
idx += 1;
}
required
}
}
#[inline]
pub(super) fn shrink_lane_slots(&self, required_lane_slots: usize) {
if required_lane_slots > self.lane_slots()
|| required_lane_slots < self.active_lane_slots()
|| required_lane_slots > usize::from(u16::MAX)
{
crate::invariant();
}
self.lane_slots.set(required_lane_slots as u16);
}
#[inline]
pub(super) fn has_entry(&self, lane: Lane, sid: SessionId) -> bool {
let Some(lane_offset) = self.lane_offset(lane) else {
return false;
};
self.find_entry_by_offset(lane_offset, sid).is_some()
}
#[inline]
pub(super) fn register(&self, lane: Lane, sid: SessionId) -> bool {
let Some(lane_offset) = self.lane_offset(lane) else {
crate::invariant();
};
if self.find_entry_by_offset(lane_offset, sid).is_some() {
crate::invariant();
}
unsafe {
let lanes = self.entry_lanes_ptr();
let sids = self.entry_sids_ptr();
let states = self.entry_states_ptr();
let mut idx = 0usize;
while idx < self.assoc_slots() {
if Self::entry_count(*states.add(idx)) == 0 {
lanes.add(idx).write(lane_offset);
sids.add(idx).write(sid);
states.add(idx).write(Self::LIVE_ENTRY_STATE);
return true;
}
idx += 1;
}
}
false
}
#[inline]
pub(super) fn increment(&self, lane: Lane, sid: SessionId) -> Option<u16> {
let lane_offset = crate::invariant_some(self.lane_offset(lane));
let idx = crate::invariant_some(self.find_entry_by_offset(lane_offset, sid));
unsafe {
let states = self.entry_states_ptr();
let raw = *states.add(idx);
let current = Self::entry_count(raw);
let next = next_attachment_count(current)?;
states
.add(idx)
.write(Self::entry_state(next, Self::entry_fault_code(raw)));
Some(next)
}
}
#[inline]
pub(super) fn decrement(&self, lane: Lane, sid: SessionId) -> u16 {
let lane_offset = crate::invariant_some(self.lane_offset(lane));
let idx = crate::invariant_some(self.find_entry_by_offset(lane_offset, sid));
unsafe {
let states = self.entry_states_ptr();
let raw = *states.add(idx);
let current = Self::entry_count(raw);
if current == 0 {
crate::invariant();
}
let next = current - 1;
if next == 0 {
self.remove_entry(idx);
} else {
states
.add(idx)
.write(Self::entry_state(next, Self::entry_fault_code(raw)));
}
next
}
}
#[inline]
pub(super) fn session_fault(&self, sid: SessionId) -> Option<SessionFaultKind> {
unsafe {
let sids = self.entry_sids_ptr();
let states = self.entry_states_ptr();
let mut idx = 0usize;
while idx < self.assoc_slots() {
let raw = *states.add(idx);
if Self::entry_count(raw) != 0
&& *sids.add(idx) == sid
&& let Some(kind) = Self::entry_fault(raw)
{
return Some(kind);
}
idx += 1;
}
None
}
}
#[inline]
pub(super) fn poison_session(
&self,
sid: SessionId,
cause: SessionFaultKind,
) -> SessionFaultKind {
if let Some(existing) = self.session_fault(sid) {
return existing;
}
unsafe {
let sids = self.entry_sids_ptr();
let states = self.entry_states_ptr();
let encoded = cause.encode();
let mut idx = 0usize;
while idx < self.assoc_slots() {
let raw = *states.add(idx);
let count = Self::entry_count(raw);
if count != 0 && *sids.add(idx) == sid {
states.add(idx).write(Self::entry_state(count, encoded));
}
idx += 1;
}
}
cause
}
}