#![cfg(unix)]
use std::ptr;
use std::sync::Mutex;
use std::sync::atomic::{AtomicU8, AtomicU32, AtomicU64, Ordering, fence};
use bytes::Bytes;
use crate::NodeId;
use crate::id::NetId64;
use crate::ring::{Frame, RingSpec, RingTopology};
use crate::shm::{self, ShmRegion};
const MAGIC: u32 = 0x4F524254; const VERSION: u32 = 3;
const SLOT_ALIGNMENT: usize = 64;
#[repr(C, align(64))]
struct ShmRingHeader {
version_counter: AtomicU64,
capacity: u64,
magic: u32,
version: u32,
payload_capacity: u32,
slot_stride: u32,
kind: u8,
topology: u8,
lane_count: u16,
notification_generation: AtomicU32,
notification_waiters: AtomicU32,
_reserved: [u8; 20]
}
#[repr(C, align(64))]
struct ShmLaneHeader {
write_pos: AtomicU64,
_reserved: [u8; 56]
}
#[repr(C)]
struct ShmSlotHeader {
seq: AtomicU64,
id: AtomicU64,
ver: AtomicU64,
payload_len: AtomicU32,
kind: AtomicU8,
_reserved: [u8; 3]
}
const SLOT_HEADER_SIZE: usize = std::mem::size_of::<ShmSlotHeader>();
const HEADER_SIZE: usize = std::mem::size_of::<ShmRingHeader>();
const LANE_HEADER_SIZE: usize = std::mem::size_of::<ShmLaneHeader>();
const _: () = assert!(HEADER_SIZE == 64);
const _: () = assert!(LANE_HEADER_SIZE == 64);
const _: () = assert!(SLOT_HEADER_SIZE == 32);
fn invalid_input(message: impl Into<String>) -> std::io::Error {
std::io::Error::new(std::io::ErrorKind::InvalidInput, message.into())
}
fn lane_count_for(
spec: RingSpec,
fleet_capacity: u16
) -> std::io::Result<usize> {
if fleet_capacity == 0 {
return Err(invalid_input("ShmRing fleet capacity must be > 0"));
}
Ok(match spec.topology {
RingTopology::Shared | RingTopology::SharedOrdered => 1,
RingTopology::PerNode => usize::from(fleet_capacity)
})
}
fn checked_layout(
spec: RingSpec,
fleet_capacity: u16
) -> std::io::Result<(usize, usize, usize)> {
if spec.capacity == 0 {
return Err(invalid_input("ShmRing capacity must be > 0"));
}
if !spec.capacity.is_power_of_two() {
return Err(invalid_input("ShmRing capacity must be a power of two"));
}
if spec.payload_capacity > u32::MAX as usize {
return Err(invalid_input("ShmRing payload capacity must fit in u32"));
}
let unaligned = SLOT_HEADER_SIZE
.checked_add(spec.payload_capacity)
.ok_or_else(|| invalid_input("ShmRing slot size overflow"))?;
let slot_stride = unaligned
.checked_add(SLOT_ALIGNMENT - 1)
.map(|value| value & !(SLOT_ALIGNMENT - 1))
.ok_or_else(|| invalid_input("ShmRing slot stride overflow"))?;
if slot_stride > u32::MAX as usize {
return Err(invalid_input("ShmRing slot stride must fit in u32"));
}
let lane_count = lane_count_for(spec, fleet_capacity)?;
let lane_headers_size = lane_count
.checked_mul(LANE_HEADER_SIZE)
.ok_or_else(|| invalid_input("ShmRing lane header size overflow"))?;
let slots_offset = HEADER_SIZE
.checked_add(lane_headers_size)
.ok_or_else(|| invalid_input("ShmRing slots offset overflow"))?;
let slots_per_lane = spec
.capacity
.checked_mul(slot_stride)
.ok_or_else(|| invalid_input("ShmRing lane size overflow"))?;
let slots_size = lane_count
.checked_mul(slots_per_lane)
.ok_or_else(|| invalid_input("ShmRing slots size overflow"))?;
let segment_size = slots_offset
.checked_add(slots_size)
.ok_or_else(|| invalid_input("ShmRing segment size overflow"))?;
Ok((slot_stride, slots_offset, segment_size))
}
pub fn segment_size_for_spec(spec: RingSpec) -> std::io::Result<usize> {
segment_size_for_spec_and_fleet(spec, 1)
}
pub fn segment_size_for_spec_and_fleet(
spec: RingSpec,
fleet_capacity: u16
) -> std::io::Result<usize> {
checked_layout(spec, fleet_capacity).map(|(_, _, segment_size)| segment_size)
}
pub struct ShmRing {
region: ShmRegion,
kind: u8,
capacity: usize,
payload_capacity: usize,
topology: RingTopology,
lane_count: usize,
slot_stride: usize,
slots_offset: usize,
lane_stride: usize,
write_locks: Vec<Mutex<()>>
}
impl ShmRing {
pub fn open_or_create_with_policy(
fleet_name: &str,
kind: u8,
spec: RingSpec,
policy: shm::ShmAccessPolicy
) -> std::io::Result<Self> {
Self::open_or_create_for_fleet_with_policy(fleet_name, kind, spec, 1, policy)
}
pub fn open_or_create(
fleet_name: &str,
kind: u8,
spec: RingSpec
) -> std::io::Result<Self> {
Self::open_or_create_for_fleet(fleet_name, kind, spec, 1)
}
pub fn open_or_create_for_fleet(
fleet_name: &str,
kind: u8,
spec: RingSpec,
fleet_capacity: u16
) -> std::io::Result<Self> {
Self::open_or_create_for_fleet_with_policy(
fleet_name,
kind,
spec,
fleet_capacity,
shm::ShmAccessPolicy::default()
)
}
pub fn open_or_create_for_fleet_with_policy(
fleet_name: &str,
kind: u8,
spec: RingSpec,
fleet_capacity: u16,
policy: shm::ShmAccessPolicy
) -> std::io::Result<Self> {
let lane_count = lane_count_for(spec, fleet_capacity)?;
let (slot_stride, slots_offset, size) = checked_layout(spec, fleet_capacity)?;
let lane_stride = spec
.capacity
.checked_mul(slot_stride)
.ok_or_else(|| invalid_input("ShmRing lane stride overflow"))?;
let name = shm::ring_segment_name(fleet_name, kind);
let (region, _initialization_lock) =
ShmRegion::open_or_create_locked_with_policy(&name, size, policy)?;
if region.created() {
unsafe {
let header_ptr = region.as_ptr() as *mut ShmRingHeader;
ptr::write(
header_ptr,
ShmRingHeader {
version_counter: AtomicU64::new(0),
capacity: spec.capacity as u64,
magic: MAGIC,
version: VERSION,
payload_capacity: spec.payload_capacity as u32,
slot_stride: slot_stride as u32,
kind,
topology: spec.topology as u8,
lane_count: lane_count as u16,
notification_generation: AtomicU32::new(0),
notification_waiters: AtomicU32::new(0),
_reserved: [0; 20]
}
);
for lane in 0..lane_count {
let lane_ptr = region.as_ptr().add(HEADER_SIZE + lane * LANE_HEADER_SIZE)
as *mut ShmLaneHeader;
ptr::write(
lane_ptr,
ShmLaneHeader { write_pos: AtomicU64::new(0), _reserved: [0; 56] }
);
}
let slots_ptr = region.as_ptr().add(slots_offset);
ptr::write_bytes(slots_ptr, 0, lane_count * lane_stride);
}
} else {
let header = unsafe { &*(region.as_ptr() as *const ShmRingHeader) };
if header.magic != MAGIC {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!(
"SHM segment {} has wrong magic 0x{:08X} (expected 0x{:08X})",
name, header.magic, MAGIC
)
));
}
if header.version != VERSION {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("SHM segment {} version {} != local {}", name, header.version, VERSION)
));
}
if header.kind != kind {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("SHM segment {} kind {} != requested {}", name, header.kind, kind)
));
}
if header.topology != spec.topology as u8 {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!(
"SHM segment {} topology {} != requested {}",
name, header.topology, spec.topology as u8
)
));
}
if header.lane_count as usize != lane_count {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!(
"SHM segment {} lane count {} != requested {}",
name, header.lane_count, lane_count
)
));
}
if header.capacity as usize != spec.capacity {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!(
"SHM segment {} capacity {} != requested {}",
name, header.capacity, spec.capacity
)
));
}
if header.payload_capacity as usize != spec.payload_capacity {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!(
"SHM segment {} payload capacity {} != requested {}",
name, header.payload_capacity, spec.payload_capacity
)
));
}
if header.slot_stride as usize != slot_stride {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!(
"SHM segment {} slot stride {} != requested {}",
name, header.slot_stride, slot_stride
)
));
}
}
Ok(Self {
region,
kind,
capacity: spec.capacity,
payload_capacity: spec.payload_capacity,
topology: spec.topology,
lane_count,
slot_stride,
slots_offset,
lane_stride,
write_locks: (0..lane_count).map(|_| Mutex::new(())).collect()
})
}
pub fn created(&self) -> bool {
self.region.created()
}
pub fn unlink(&self) -> std::io::Result<()> {
self.region.unlink()
}
pub fn kind(&self) -> u8 {
self.kind
}
pub fn capacity(&self) -> usize {
self.capacity
}
pub fn lane_count(&self) -> usize {
self.lane_count
}
pub fn payload_capacity(&self) -> usize {
self.payload_capacity
}
pub fn spec(&self) -> RingSpec {
RingSpec {
capacity: self.capacity,
payload_capacity: self.payload_capacity,
topology: self.topology
}
}
fn lane_header(
&self,
lane: usize
) -> &ShmLaneHeader {
debug_assert!(lane < self.lane_count);
unsafe {
&*(self.region.as_ptr().add(HEADER_SIZE + lane * LANE_HEADER_SIZE)
as *const ShmLaneHeader)
}
}
#[cfg(any(target_os = "linux", target_os = "freebsd", target_os = "macos"))]
pub(crate) fn notification_generation(&self) -> &AtomicU32 {
unsafe { &(*(self.region.as_ptr() as *const ShmRingHeader)).notification_generation }
}
#[cfg(any(target_os = "linux", target_os = "freebsd", target_os = "macos"))]
pub(crate) fn notification_waiters(&self) -> &AtomicU32 {
unsafe { &(*(self.region.as_ptr() as *const ShmRingHeader)).notification_waiters }
}
fn slot_ptr(
&self,
lane: usize,
idx: usize
) -> *mut ShmSlotHeader {
debug_assert!(lane < self.lane_count);
debug_assert!(idx < self.capacity);
unsafe {
let base = self.region.as_ptr().add(self.slots_offset + lane * self.lane_stride);
base.add(idx * self.slot_stride).cast::<ShmSlotHeader>()
}
}
unsafe fn payload_ptr(slot_ptr: *mut ShmSlotHeader) -> *mut AtomicU8 {
unsafe { slot_ptr.cast::<u8>().add(SLOT_HEADER_SIZE).cast::<AtomicU8>() }
}
pub fn head(&self) -> u64 {
self.lane_header(0).write_pos.load(Ordering::Acquire)
}
pub fn lane_head(
&self,
node_id: NodeId
) -> u64 {
let lane = self.lane_index(node_id);
self.lane_header(lane).write_pos.load(Ordering::Acquire)
}
pub fn next_version(&self) -> u64 {
let header = unsafe { &*(self.region.as_ptr() as *const ShmRingHeader) };
header
.version_counter
.fetch_add(1, Ordering::AcqRel)
.checked_add(1)
.expect("SHM ring semantic version exhausted")
}
pub fn current_version(&self) -> u64 {
let header = unsafe { &*(self.region.as_ptr() as *const ShmRingHeader) };
header.version_counter.load(Ordering::Acquire)
}
pub fn write(
&self,
node_id: NodeId,
frame_kind: u8,
ver: u64,
payload: Bytes
) -> std::io::Result<NetId64> {
if payload.len() > self.payload_capacity {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
format!(
"payload {} > ring payload capacity {}",
payload.len(),
self.payload_capacity
)
));
}
let lane = self.lane_index(node_id);
match self.topology {
RingTopology::Shared => {
let counter = self.lane_header(lane).write_pos.fetch_add(1, Ordering::AcqRel);
Ok(self.write_slot(lane, node_id, counter, frame_kind, ver, &payload))
}
RingTopology::PerNode => {
let _write =
self.write_locks[lane].lock().unwrap_or_else(|error| error.into_inner());
let lane_header = self.lane_header(lane);
let counter = lane_header.write_pos.load(Ordering::Relaxed);
let id = self.write_slot(lane, node_id, counter, frame_kind, ver, &payload);
lane_header.write_pos.store(counter.wrapping_add(1), Ordering::Release);
Ok(id)
}
RingTopology::SharedOrdered => {
let _write =
self.write_locks[lane].lock().unwrap_or_else(|error| error.into_inner());
let _cross_process = self.region.lock_exclusive()?;
let lane_header = self.lane_header(lane);
let counter = lane_header.write_pos.load(Ordering::Relaxed);
let id = self.write_slot(lane, node_id, counter, frame_kind, ver, &payload);
lane_header.write_pos.store(counter.wrapping_add(1), Ordering::Release);
Ok(id)
}
}
}
pub fn write_batch(
&self,
node_id: NodeId,
frame_kind: u8,
ver: u64,
payloads: Vec<Bytes>
) -> std::io::Result<Vec<NetId64>> {
if payloads.len() > self.capacity {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
format!("batch {} > ring capacity {}", payloads.len(), self.capacity)
));
}
if let Some(payload) = payloads.iter().find(|payload| payload.len() > self.payload_capacity)
{
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
format!(
"payload {} > ring payload capacity {}",
payload.len(),
self.payload_capacity
)
));
}
if payloads.is_empty() {
return Ok(Vec::new());
}
let lane = self.lane_index(node_id);
let write_slots = |start: u64, payloads: Vec<Bytes>| {
payloads
.into_iter()
.enumerate()
.map(|(offset, payload)| {
self.write_slot(
lane,
node_id,
start.wrapping_add(offset as u64),
frame_kind,
ver,
&payload
)
})
.collect::<Vec<_>>()
};
match self.topology {
RingTopology::Shared => {
let start = self
.lane_header(lane)
.write_pos
.fetch_add(payloads.len() as u64, Ordering::AcqRel);
Ok(write_slots(start, payloads))
}
RingTopology::PerNode => {
let _write =
self.write_locks[lane].lock().unwrap_or_else(|error| error.into_inner());
let lane_header = self.lane_header(lane);
let start = lane_header.write_pos.load(Ordering::Relaxed);
let ids = write_slots(start, payloads);
lane_header
.write_pos
.store(start.wrapping_add(ids.len() as u64), Ordering::Release);
Ok(ids)
}
RingTopology::SharedOrdered => {
let _write =
self.write_locks[lane].lock().unwrap_or_else(|error| error.into_inner());
let _cross_process = self.region.lock_exclusive()?;
let lane_header = self.lane_header(lane);
let start = lane_header.write_pos.load(Ordering::Relaxed);
let ids = write_slots(start, payloads);
lane_header
.write_pos
.store(start.wrapping_add(ids.len() as u64), Ordering::Release);
Ok(ids)
}
}
}
pub fn read(
&self,
id: NetId64
) -> Option<Frame> {
if id.kind() != self.kind {
return None;
}
let lane = self.lane_index_for_frame(id)?;
let counter = id.counter();
let slot_idx = (counter as usize) & (self.capacity - 1);
let slot_ptr = self.slot_ptr(lane, slot_idx);
for _ in 0..3 {
let Some(frame) = (unsafe { read_committed_frame(slot_ptr, self.payload_capacity) })
else {
continue;
};
if frame.id.counter() == counter {
return Some(frame);
} else {
return None;
}
}
None
}
pub fn read_head(&self) -> Option<Frame> {
let head = self.head();
if head == 0 {
return None;
}
let counter = head - 1;
let slot_idx = (counter as usize) & (self.capacity - 1);
let slot_ptr = self.slot_ptr(0, slot_idx);
for _ in 0..3 {
if let Some(frame) = unsafe { read_committed_frame(slot_ptr, self.payload_capacity) } {
return Some(frame);
}
}
None
}
pub fn read_at(
&self,
counter: u64
) -> Option<Frame> {
self.read_lane_index_at(0, counter)
}
pub fn read_lane_at(
&self,
node_id: NodeId,
counter: u64
) -> Option<Frame> {
let lane = self.lane_index(node_id);
self.read_lane_index_at(lane, counter)
}
fn read_lane_index_at(
&self,
lane: usize,
counter: u64
) -> Option<Frame> {
let slot_idx = (counter as usize) & (self.capacity - 1);
let slot_ptr = self.slot_ptr(lane, slot_idx);
for _ in 0..3 {
if let Some(frame) = unsafe { read_committed_frame(slot_ptr, self.payload_capacity) } {
return Some(frame);
}
}
None
}
pub(crate) fn read_state_at(
&self,
counter: u64
) -> crate::ring::cursor::RingRead {
self.read_lane_index_state_at(0, counter)
}
pub(crate) fn read_lane_state_at(
&self,
node_id: NodeId,
counter: u64
) -> crate::ring::cursor::RingRead {
let lane = self.lane_index(node_id);
self.read_lane_index_state_at(lane, counter)
}
fn read_lane_index_state_at(
&self,
lane: usize,
counter: u64
) -> crate::ring::cursor::RingRead {
use crate::ring::cursor::RingRead;
let slot_idx = (counter as usize) & (self.capacity - 1);
let slot_ptr = self.slot_ptr(lane, slot_idx);
let expected_committed =
counter.checked_mul(2).and_then(|value| value.checked_add(2)).expect("seq overflow");
for _ in 0..3 {
let sequence = unsafe { &*slot_ptr }.seq.load(Ordering::Acquire);
if sequence < expected_committed {
return if self.topology != RingTopology::Shared {
RingRead::Unavailable
} else {
RingRead::Pending
};
}
if sequence > expected_committed {
return RingRead::Unavailable;
}
if let Some(frame) = unsafe { read_committed_frame(slot_ptr, self.payload_capacity) } {
return if frame.id.counter() == counter {
RingRead::Ready(frame)
} else {
RingRead::Unavailable
};
}
}
let sequence = unsafe { &*slot_ptr }.seq.load(Ordering::Acquire);
if sequence < expected_committed {
if self.topology != RingTopology::Shared {
RingRead::Unavailable
} else {
RingRead::Pending
}
} else {
RingRead::Unavailable
}
}
pub fn reset(&self) {
unsafe {
let slots_ptr = self.region.as_ptr().add(self.slots_offset);
ptr::write_bytes(slots_ptr, 0, self.lane_count * self.lane_stride);
}
for lane in 0..self.lane_count {
self.lane_header(lane).write_pos.store(0, Ordering::Release);
}
let header = unsafe { &*(self.region.as_ptr() as *const ShmRingHeader) };
header.version_counter.store(0, Ordering::Release);
}
fn lane_index(
&self,
node_id: NodeId
) -> usize {
let lane = match self.topology {
RingTopology::Shared | RingTopology::SharedOrdered => 0,
RingTopology::PerNode => usize::from(node_id.get())
};
assert!(
lane < self.lane_count,
"node {} is outside SHM ring lane count {}",
node_id.get(),
self.lane_count
);
lane
}
fn lane_index_for_frame(
&self,
id: NetId64
) -> Option<usize> {
let lane = match self.topology {
RingTopology::Shared | RingTopology::SharedOrdered => 0,
RingTopology::PerNode => usize::from(id.node())
};
(lane < self.lane_count).then_some(lane)
}
fn write_slot(
&self,
lane: usize,
node_id: NodeId,
counter: u64,
frame_kind: u8,
ver: u64,
payload: &[u8]
) -> NetId64 {
let id = NetId64::make(self.kind, node_id.get(), counter);
let slot_idx = (counter as usize) & (self.capacity - 1);
let slot_ptr = self.slot_ptr(lane, slot_idx);
unsafe {
let slot = &*slot_ptr;
let mid_seq = counter
.checked_mul(2)
.and_then(|value| value.checked_add(1))
.expect("seq overflow");
let final_seq = mid_seq.wrapping_add(1);
slot.seq.store(mid_seq, Ordering::Relaxed);
fence(Ordering::Release);
slot.id.store(id.raw(), Ordering::Relaxed);
slot.ver.store(ver, Ordering::Relaxed);
slot.payload_len.store(payload.len() as u32, Ordering::Relaxed);
slot.kind.store(frame_kind, Ordering::Relaxed);
let bytes = Self::payload_ptr(slot_ptr);
for (index, byte) in payload.iter().enumerate() {
(*bytes.add(index)).store(*byte, Ordering::Relaxed);
}
slot.seq.store(final_seq, Ordering::Release);
}
id
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ShmRingMetadata {
pub name: String,
pub segment_size: usize,
pub format_version: u32,
pub kind: u8,
pub spec: RingSpec,
pub lane_count: usize
}
pub struct ShmRingView {
region: ShmRegion,
metadata: ShmRingMetadata,
slot_stride: usize,
slots_offset: usize,
lane_stride: usize
}
impl std::fmt::Debug for ShmRingView {
fn fmt(
&self,
formatter: &mut std::fmt::Formatter<'_>
) -> std::fmt::Result {
formatter.debug_struct("ShmRingView").field("metadata", &self.metadata).finish()
}
}
impl ShmRingView {
pub fn attach_existing_with_policy(
fleet_name: &str,
kind: u8,
policy: shm::ShmAccessPolicy
) -> std::io::Result<Self> {
Self::attach_existing_for_uid_with_policy(
fleet_name,
kind,
unsafe { libc::geteuid() },
policy
)
}
pub fn attach_existing(
fleet_name: &str,
kind: u8
) -> std::io::Result<Self> {
let uid = unsafe { libc::geteuid() };
Self::attach_existing_for_uid(fleet_name, kind, uid)
}
pub fn attach_existing_for_uid(
fleet_name: &str,
kind: u8,
uid: u32
) -> std::io::Result<Self> {
Self::attach_existing_for_uid_with_policy(
fleet_name,
kind,
uid,
shm::ShmAccessPolicy::default()
)
}
pub fn attach_existing_for_uid_with_policy(
fleet_name: &str,
kind: u8,
uid: u32,
policy: shm::ShmAccessPolicy
) -> std::io::Result<Self> {
let name = shm::ring_segment_name_for_uid(fleet_name, kind, uid);
let region = ShmRegion::open_existing_read_only(&name, HEADER_SIZE, uid, policy)?;
let header = unsafe { &*(region.as_ptr() as *const ShmRingHeader) };
if header.magic != MAGIC {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!(
"SHM segment {} has wrong magic 0x{:08X} (expected 0x{:08X})",
name, header.magic, MAGIC
)
));
}
if header.version != VERSION {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("SHM segment {} version {} != local {}", name, header.version, VERSION)
));
}
if header.kind != kind {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("SHM segment {} kind {} != requested {}", name, header.kind, kind)
));
}
let topology = match header.topology {
value if value == RingTopology::Shared as u8 => RingTopology::Shared,
value if value == RingTopology::PerNode as u8 => RingTopology::PerNode,
value if value == RingTopology::SharedOrdered as u8 => RingTopology::SharedOrdered,
value => {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("SHM segment {name} has unknown topology {value}")
));
}
};
let capacity = usize::try_from(header.capacity).map_err(|_| {
std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!(
"SHM segment {name} capacity {} does not fit this platform",
header.capacity
)
)
})?;
let lane_count = usize::from(header.lane_count);
if lane_count == 0 {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("SHM segment {name} has zero lanes")
));
}
if topology != RingTopology::PerNode && lane_count != 1 {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!(
"SHM segment {name} topology {topology:?} requires one lane, found {lane_count}"
)
));
}
let spec =
RingSpec { capacity, payload_capacity: header.payload_capacity as usize, topology };
let fleet_capacity = if topology == RingTopology::PerNode { header.lane_count } else { 1 };
let (slot_stride, slots_offset, expected_size) = checked_layout(spec, fleet_capacity)
.map_err(|error| {
std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("SHM segment {name} has invalid geometry: {error}")
)
})?;
if header.slot_stride as usize != slot_stride {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!(
"SHM segment {} slot stride {} != geometry {}",
name, header.slot_stride, slot_stride
)
));
}
if region.len() < expected_size {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!(
"SHM segment {} size {} is smaller than persisted geometry {}",
name,
region.len(),
expected_size
)
));
}
let lane_stride = capacity
.checked_mul(slot_stride)
.ok_or_else(|| invalid_input(format!("SHM segment {name} lane stride overflow")))?;
Ok(Self {
metadata: ShmRingMetadata {
name,
segment_size: region.len(),
format_version: header.version,
kind,
spec,
lane_count
},
region,
slot_stride,
slots_offset,
lane_stride
})
}
pub fn attach_typed<T: crate::OrbitTyped>(fleet_name: &str) -> std::io::Result<Self> {
let view = Self::attach_existing(fleet_name, T::KIND)?;
if view.metadata.spec != T::RING_SPEC {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!(
"OrbitTyped KIND {} declares {:?}; existing spec is {:?}",
T::KIND,
T::RING_SPEC,
view.metadata.spec
)
));
}
Ok(view)
}
pub fn metadata(&self) -> &ShmRingMetadata {
&self.metadata
}
pub fn current_version(&self) -> u64 {
self.header().version_counter.load(Ordering::Acquire)
}
pub fn notification_generation(&self) -> u32 {
self.header().notification_generation.load(Ordering::Acquire)
}
pub fn lane(
&self,
lane: usize
) -> std::io::Result<ShmRingLaneView<'_>> {
if lane >= self.metadata.lane_count {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
format!("lane {lane} is outside SHM ring lane count {}", self.metadata.lane_count)
));
}
Ok(ShmRingLaneView { ring: self, lane })
}
fn header(&self) -> &ShmRingHeader {
unsafe { &*(self.region.as_ptr() as *const ShmRingHeader) }
}
fn lane_header(
&self,
lane: usize
) -> &ShmLaneHeader {
debug_assert!(lane < self.metadata.lane_count);
unsafe {
&*(self.region.as_ptr().add(HEADER_SIZE + lane * LANE_HEADER_SIZE)
as *const ShmLaneHeader)
}
}
fn slot_ptr(
&self,
lane: usize,
counter: u64
) -> *mut ShmSlotHeader {
debug_assert!(lane < self.metadata.lane_count);
let slot = (counter as usize) & (self.metadata.spec.capacity - 1);
unsafe {
self.region
.as_ptr()
.add(self.slots_offset + lane * self.lane_stride + slot * self.slot_stride)
.cast::<ShmSlotHeader>()
}
}
fn read_lane_state_at(
&self,
lane: usize,
counter: u64
) -> crate::ring::cursor::RingRead {
use crate::ring::cursor::RingRead;
let slot_ptr = self.slot_ptr(lane, counter);
let Some(expected_committed) =
counter.checked_mul(2).and_then(|value| value.checked_add(2))
else {
return RingRead::Unavailable;
};
for _ in 0..3 {
let sequence = unsafe { &*slot_ptr }.seq.load(Ordering::Acquire);
if sequence < expected_committed {
return if self.metadata.spec.topology == RingTopology::Shared {
RingRead::Pending
} else {
RingRead::Unavailable
};
}
if sequence > expected_committed {
return RingRead::Unavailable;
}
if let Some(frame) =
unsafe { read_committed_frame(slot_ptr, self.metadata.spec.payload_capacity) }
{
let lane_matches = self.metadata.spec.topology != RingTopology::PerNode
|| usize::from(frame.id.node()) == lane;
return if frame.id.kind() == self.metadata.kind
&& frame.id.counter() == counter
&& lane_matches
{
RingRead::Ready(frame)
} else {
RingRead::Unavailable
};
}
}
RingRead::Unavailable
}
}
#[derive(Clone, Copy, Debug)]
pub struct ShmRingLaneView<'a> {
ring: &'a ShmRingView,
lane: usize
}
impl ShmRingLaneView<'_> {
pub fn index(self) -> usize {
self.lane
}
pub fn head(self) -> u64 {
self.ring.lane_header(self.lane).write_pos.load(Ordering::Acquire)
}
pub fn retained_range(self) -> std::ops::Range<u64> {
let head = self.head();
head.saturating_sub(self.ring.metadata.spec.capacity as u64)..head
}
pub fn read_at(
self,
counter: u64
) -> Option<Frame> {
match self.ring.read_lane_state_at(self.lane, counter) {
crate::ring::cursor::RingRead::Ready(frame) => Some(frame),
crate::ring::cursor::RingRead::Pending | crate::ring::cursor::RingRead::Unavailable => {
None
}
}
}
pub fn read_head(self) -> Option<Frame> {
let head = self.head();
head.checked_sub(1).and_then(|counter| self.read_at(counter))
}
}
impl crate::ring::cursor::RingFrameSource for ShmRingLaneView<'_> {
fn kind(&self) -> u8 {
self.ring.metadata.kind
}
fn head(&self) -> u64 {
(*self).head()
}
fn capacity(&self) -> usize {
self.ring.metadata.spec.capacity
}
fn read_at(
&self,
counter: u64
) -> Option<Frame> {
(*self).read_at(counter)
}
fn read_state_at(
&self,
counter: u64
) -> crate::ring::cursor::RingRead {
self.ring.read_lane_state_at(self.lane, counter)
}
}
pub struct ShmRingRegistry {
fleet_name: String,
fleet_capacity: u16,
policies: std::collections::BTreeMap<u8, shm::ShmAccessPolicy>,
rings: dashmap::DashMap<u8, std::sync::Arc<ShmRing>>
}
impl ShmRingRegistry {
pub fn new(
fleet_name: impl Into<String>,
fleet_capacity: u16
) -> Self {
Self::with_policies(fleet_name, fleet_capacity, [])
}
pub fn with_policies(
fleet_name: impl Into<String>,
fleet_capacity: u16,
policies: impl IntoIterator<Item = (u8, shm::ShmAccessPolicy)>
) -> Self {
Self {
fleet_name: fleet_name.into(),
fleet_capacity,
policies: policies.into_iter().collect(),
rings: dashmap::DashMap::new()
}
}
pub fn get_or_create_for<T: crate::OrbitTyped>(
&self
) -> std::io::Result<std::sync::Arc<ShmRing>> {
if let Some(entry) = self.rings.get(&T::KIND) {
if entry.spec() != T::RING_SPEC {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!(
"OrbitTyped KIND {} was reused with ring spec {:?}; existing spec is {:?}",
T::KIND,
T::RING_SPEC,
entry.spec()
)
));
}
return Ok(entry.clone());
}
let ring = std::sync::Arc::new(ShmRing::open_or_create_for_fleet_with_policy(
&self.fleet_name,
T::KIND,
T::RING_SPEC,
self.fleet_capacity,
self.policies.get(&T::KIND).copied().unwrap_or_default()
)?);
let entry = self.rings.entry(T::KIND).or_insert_with(|| ring.clone());
if entry.spec() != T::RING_SPEC {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!(
"OrbitTyped KIND {} raced with ring spec {:?}; installed spec is {:?}",
T::KIND,
T::RING_SPEC,
entry.spec()
)
));
}
Ok(entry.clone())
}
pub fn lookup(
&self,
kind: u8
) -> Option<std::sync::Arc<ShmRing>> {
self.rings.get(&kind).map(|e| e.clone())
}
}
unsafe fn read_committed_frame(
slot_ptr: *mut ShmSlotHeader,
payload_capacity: usize
) -> Option<Frame> {
let slot = unsafe { &*slot_ptr };
let seq_pre = slot.seq.load(Ordering::Acquire);
if seq_pre == 0 {
return None;
}
if seq_pre & 1 == 1 {
return None;
}
let id = NetId64::from_raw(slot.id.load(Ordering::Relaxed));
let kind = slot.kind.load(Ordering::Relaxed);
let ver = slot.ver.load(Ordering::Relaxed);
let payload_len = slot.payload_len.load(Ordering::Relaxed) as usize;
if payload_len > payload_capacity {
return None;
}
let payload_src = unsafe { ShmRing::payload_ptr(slot_ptr) };
let mut payload_buf = Vec::with_capacity(payload_len);
for index in 0..payload_len {
payload_buf.push(unsafe { (*payload_src.add(index)).load(Ordering::Relaxed) });
}
fence(Ordering::Acquire);
let seq_post = slot.seq.load(Ordering::Relaxed);
if seq_pre != seq_post {
return None;
}
Some(Frame { id, kind, ver, payload: Bytes::from(payload_buf) })
}
#[cfg(test)]
mod tests {
use std::sync::atomic::{AtomicU64, Ordering};
use nix::sys::wait::{WaitStatus, waitpid};
use nix::unistd::{ForkResult, fork};
use super::*;
use crate::ring::cursor::{RingCursor, poll_ring};
#[test]
fn cursor_retries_a_claimed_slot_after_it_commits() {
static TEST_ID: AtomicU64 = AtomicU64::new(0);
let test_id = TEST_ID.fetch_add(1, Ordering::Relaxed);
let fleet_name = format!("p{:x}{test_id:x}", std::process::id());
let ring = ShmRing::open_or_create(&fleet_name, 199, RingSpec::new(4, 16))
.expect("create test ring");
ring.reset();
let slot_ptr = ring.slot_ptr(0, 0);
ring.lane_header(0).write_pos.store(1, Ordering::Release);
unsafe { &*slot_ptr }.seq.store(1, Ordering::Release);
let mut cursor = RingCursor::from_start();
let pending = poll_ring(&ring, &mut cursor);
assert!(pending.is_empty());
assert_eq!(cursor.next_counter(), 0);
let payload = b"ready";
unsafe {
let slot = &*slot_ptr;
slot.id.store(NetId64::make(199, NodeId::ZERO.get(), 0).raw(), Ordering::Relaxed);
slot.ver.store(7, Ordering::Relaxed);
slot.payload_len.store(payload.len() as u32, Ordering::Relaxed);
slot.kind.store(1, Ordering::Relaxed);
let bytes = ShmRing::payload_ptr(slot_ptr);
for (index, byte) in payload.iter().enumerate() {
(*bytes.add(index)).store(*byte, Ordering::Relaxed);
}
}
unsafe { &*slot_ptr }.seq.store(2, Ordering::Release);
let committed = poll_ring(&ring, &mut cursor);
assert_eq!(committed.frames.len(), 1);
assert_eq!(&committed.frames[0].payload[..], payload);
assert_eq!(cursor.next_counter(), 1);
ring.unlink().expect("unlink test ring");
}
#[test]
fn per_node_head_ignores_a_writer_that_dies_before_commit() {
static TEST_ID: AtomicU64 = AtomicU64::new(0);
let test_id = TEST_ID.fetch_add(1, Ordering::Relaxed);
let fleet_name = format!("d{:x}{test_id:x}", std::process::id());
let spec = RingSpec::per_node(4, 16);
let abandoned = ShmRing::open_or_create_for_fleet(&fleet_name, 198, spec, 2)
.expect("create per-node test ring");
abandoned.reset();
let slot_ptr = abandoned.slot_ptr(1, 0);
unsafe { &*slot_ptr }.seq.store(1, Ordering::Release);
assert_eq!(abandoned.lane_head(NodeId::new(1)), 0);
let replacement = ShmRing::open_or_create_for_fleet(&fleet_name, 198, spec, 2)
.expect("replacement attaches");
let id = replacement
.write(NodeId::new(1), 1, 9, Bytes::from_static(b"recovered"))
.expect("replacement commits");
assert_eq!(id.counter(), 0);
assert_eq!(replacement.lane_head(NodeId::new(1)), 1);
assert_eq!(&replacement.read(id).expect("frame visible").payload[..], b"recovered");
replacement.unlink().expect("unlink test ring");
}
#[test]
fn per_node_lane_serializes_concurrent_local_publishers() {
static TEST_ID: AtomicU64 = AtomicU64::new(0);
let test_id = TEST_ID.fetch_add(1, Ordering::Relaxed);
let fleet_name = format!("c{:x}{test_id:x}", std::process::id());
let ring = std::sync::Arc::new(
ShmRing::open_or_create_for_fleet(&fleet_name, 197, RingSpec::per_node(512, 0), 2)
.expect("create concurrent writer ring")
);
ring.reset();
let mut writers = Vec::new();
for _ in 0..4 {
let ring = ring.clone();
writers.push(std::thread::spawn(move || {
(0..64)
.map(|_| {
ring.write(NodeId::new(1), 1, 0, Bytes::new()).expect("publish").counter()
})
.collect::<Vec<_>>()
}));
}
let mut counters = writers
.into_iter()
.flat_map(|writer| writer.join().expect("writer joins"))
.collect::<Vec<_>>();
counters.sort_unstable();
assert_eq!(counters, (0..256).collect::<Vec<_>>());
assert_eq!(ring.lane_head(NodeId::new(1)), 256);
ring.unlink().expect("unlink test ring");
}
#[test]
fn shared_ordered_recovers_after_a_locked_writer_process_dies() {
static TEST_ID: AtomicU64 = AtomicU64::new(0);
let test_id = TEST_ID.fetch_add(1, Ordering::Relaxed);
let fleet_name = format!("o{:x}{test_id:x}", std::process::id());
let spec = RingSpec::shared_ordered(4, 16);
let ring = ShmRing::open_or_create_for_fleet(&fleet_name, 196, spec, 2)
.expect("create shared-ordered test ring");
ring.reset();
match unsafe { fork() }.expect("fork test writer") {
ForkResult::Child => {
let _lock = ring.region.lock_exclusive().expect("child locks ring");
let slot_ptr = ring.slot_ptr(0, 0);
unsafe { &*slot_ptr }.seq.store(1, Ordering::Release);
unsafe { libc::_exit(0) };
}
ForkResult::Parent { child } => {
let status = waitpid(child, None).expect("wait for abandoned writer");
assert!(matches!(status, WaitStatus::Exited(_, 0)));
let replacement = ShmRing::open_or_create_for_fleet(&fleet_name, 196, spec, 2)
.expect("replacement attaches");
let recovered_lock = replacement
.region
.try_lock_exclusive()
.expect("kernel released dead writer lock");
drop(recovered_lock);
let id = replacement
.write(NodeId::new(1), 1, 9, Bytes::from_static(b"recovered"))
.expect("replacement commits counter zero");
assert_eq!(id.counter(), 0);
assert_eq!(replacement.head(), 1);
assert_eq!(
&replacement.read(id).expect("recovered frame").payload[..],
b"recovered"
);
replacement.unlink().expect("unlink shared-ordered test ring");
}
}
}
}