use crate::ipc::{
BlobBackendKind, DeltaOp, IpcMessage, IpcValue, NodeState, ShmBlobArena, ShmBlobArenaError,
ShmBlobRef,
};
pub type BlobView<'a> = Option<&'a [u8]>;
pub trait BlobBackend {
fn kind(&self) -> BlobBackendKind;
fn write(&mut self, bytes: &[u8]) -> Result<ShmBlobRef, ShmBlobArenaError>;
fn read_view(&self, descriptor: &ShmBlobRef) -> BlobView<'_>;
fn advance_epoch(&mut self);
}
pub struct InProcessBackend {
arena: ShmBlobArena<Vec<u8>>,
epoch: u64,
}
pub const IN_PROCESS_DEFAULT_CAPACITY: usize = 1 << 20;
impl InProcessBackend {
pub fn new() -> Result<Self, ShmBlobArenaError> {
Self::with_capacity(IN_PROCESS_DEFAULT_CAPACITY)
}
pub fn with_capacity(capacity: usize) -> Result<Self, ShmBlobArenaError> {
Ok(Self {
arena: ShmBlobArena::with_capacity(capacity)?,
epoch: 0,
})
}
pub fn from_arena(arena: ShmBlobArena<Vec<u8>>) -> Self {
Self { arena, epoch: 0 }
}
pub fn arena(&self) -> &ShmBlobArena<Vec<u8>> {
&self.arena
}
pub fn epoch(&self) -> u64 {
self.epoch
}
}
impl Default for InProcessBackend {
fn default() -> Self {
Self::new().expect("IN_PROCESS_DEFAULT_CAPACITY >= SHM_BLOB_HEADER_LEN + 1")
}
}
impl BlobBackend for InProcessBackend {
fn kind(&self) -> BlobBackendKind {
BlobBackendKind::InProcess
}
fn write(&mut self, bytes: &[u8]) -> Result<ShmBlobRef, ShmBlobArenaError> {
let mut descriptor = self.arena.write_blob(self.epoch, bytes)?;
descriptor.backend = BlobBackendKind::InProcess;
Ok(descriptor)
}
fn read_view(&self, descriptor: &ShmBlobRef) -> BlobView<'_> {
if descriptor.epoch != self.epoch {
return None;
}
self.arena.read_blob(*descriptor).ok()
}
fn advance_epoch(&mut self) {
self.epoch = self.epoch.saturating_add(1);
}
}
pub struct ArrowBackend {
arena: ShmBlobArena<Vec<u8>>,
epoch: u64,
}
pub const ARROW_DEFAULT_CAPACITY: usize = 1 << 22;
impl ArrowBackend {
pub fn new() -> Result<Self, ShmBlobArenaError> {
Self::with_capacity(ARROW_DEFAULT_CAPACITY)
}
pub fn with_capacity(capacity: usize) -> Result<Self, ShmBlobArenaError> {
Ok(Self {
arena: ShmBlobArena::with_capacity(capacity)?,
epoch: 0,
})
}
pub fn epoch(&self) -> u64 {
self.epoch
}
}
impl Default for ArrowBackend {
fn default() -> Self {
Self::new().expect("ARROW_DEFAULT_CAPACITY >= SHM_BLOB_HEADER_LEN + 1")
}
}
impl BlobBackend for ArrowBackend {
fn kind(&self) -> BlobBackendKind {
BlobBackendKind::Arrow
}
fn write(&mut self, bytes: &[u8]) -> Result<ShmBlobRef, ShmBlobArenaError> {
let mut descriptor = self.arena.write_blob(self.epoch, bytes)?;
descriptor.backend = BlobBackendKind::Arrow;
Ok(descriptor)
}
fn read_view(&self, descriptor: &ShmBlobRef) -> BlobView<'_> {
if descriptor.epoch != self.epoch {
return None;
}
self.arena.read_blob(*descriptor).ok()
}
fn advance_epoch(&mut self) {
self.epoch = self.epoch.saturating_add(1);
}
}
#[cfg(all(unix, feature = "shm"))]
mod shm {
use super::{BlobBackend, BlobBackendKind, BlobView, ShmBlobArenaError, ShmBlobRef};
use std::io;
use std::sync::atomic::{AtomicU64, Ordering};
const SHM_MAGIC: u64 = 0x4c5a_5348_424c_4f42; const SLOT_HEADER_LEN: usize = 24;
#[repr(C)]
struct Header {
magic: AtomicU64,
capacity: u64,
bump: AtomicU64,
generation: AtomicU64,
epoch: AtomicU64,
}
#[repr(C)]
struct SlotHeader {
generation: u64,
len: u64,
checksum: u64,
}
const HEADER_LEN: usize = 40;
pub struct ShmBackend {
name: String,
fd: std::os::fd::RawFd,
base: *mut u8,
capacity: usize,
header: *mut Header,
}
unsafe impl Send for ShmBackend {}
unsafe impl Sync for ShmBackend {}
impl ShmBackend {
fn open_raw(name: &str, capacity: usize, create: bool) -> io::Result<Self> {
let c_name = ensure_leading_slash(name);
let flags = if create {
libc::O_RDWR | libc::O_CREAT
} else {
libc::O_RDWR
};
let fd = unsafe { libc::shm_open(c_name.as_ptr(), flags, 0o600) };
if fd < 0 {
return Err(io::Error::last_os_error());
}
if create && unsafe { libc::ftruncate(fd, capacity as libc::off_t) } != 0 {
let e = io::Error::last_os_error();
unsafe { libc::close(fd) };
return Err(e);
}
let base = unsafe {
libc::mmap(
std::ptr::null_mut(),
capacity,
libc::PROT_READ | libc::PROT_WRITE,
libc::MAP_SHARED,
fd,
0,
)
};
if base == libc::MAP_FAILED {
let e = io::Error::last_os_error();
unsafe { libc::close(fd) };
return Err(e);
}
let base = base as *mut u8;
let header = base as *mut Header;
if create {
unsafe {
(*header).magic.store(SHM_MAGIC, Ordering::Relaxed);
(*header).capacity = capacity as u64;
(*header).bump.store(HEADER_LEN as u64, Ordering::Relaxed);
(*header).generation.store(0, Ordering::Relaxed);
(*header).epoch.store(0, Ordering::Relaxed);
}
}
Ok(Self {
name: name.to_string(),
fd,
base,
capacity,
header,
})
}
pub fn create(name: &str, capacity: usize) -> Result<Self, ShmBlobArenaError> {
Self::open_raw(name, capacity, true).map_err(shm_io_err)
}
pub fn open(name: &str) -> Result<Self, ShmBlobArenaError> {
let probe = Self::open_raw(name, HEADER_LEN, false).map_err(shm_io_err)?;
let capacity = unsafe { (*probe.header).capacity } as usize;
let name_owned = probe.name.clone();
drop(probe);
Self::open_raw(&name_owned, capacity, false).map_err(shm_io_err)
}
pub fn unlink(name: &str) {
let c_name = ensure_leading_slash(name);
unsafe {
libc::shm_unlink(c_name.as_ptr());
}
}
pub fn epoch(&self) -> u64 {
unsafe { (*self.header).epoch.load(Ordering::Acquire) }
}
pub fn bump_offset(&self) -> u64 {
unsafe { (*self.header).bump.load(Ordering::Acquire) }
}
}
impl Drop for ShmBackend {
fn drop(&mut self) {
unsafe {
if !self.base.is_null() {
libc::munmap(self.base as *mut libc::c_void, self.capacity);
self.base = std::ptr::null_mut();
}
if self.fd >= 0 {
libc::close(self.fd);
self.fd = -1;
}
}
}
}
impl BlobBackend for ShmBackend {
fn kind(&self) -> BlobBackendKind {
BlobBackendKind::Shm
}
fn write(&mut self, bytes: &[u8]) -> Result<ShmBlobRef, ShmBlobArenaError> {
let need = SLOT_HEADER_LEN + bytes.len();
let header = unsafe { &*self.header };
let off = header.bump.fetch_add(need as u64, Ordering::AcqRel);
if off as usize + need > self.capacity {
return Err(ShmBlobArenaError::BlobTooLarge {
len: bytes.len(),
max_len: self.capacity.saturating_sub(SLOT_HEADER_LEN + HEADER_LEN),
});
}
let generation = header.generation.fetch_add(1, Ordering::AcqRel) + 1;
let ep = header.epoch.load(Ordering::Acquire);
let csum = fnv1a_64(bytes);
unsafe {
let slot = (self.base.add(off as usize)) as *mut SlotHeader;
std::ptr::write_unaligned(
slot,
SlotHeader {
generation,
len: bytes.len() as u64,
checksum: csum,
},
);
std::ptr::copy_nonoverlapping(
bytes.as_ptr(),
self.base.add(off as usize + SLOT_HEADER_LEN),
bytes.len(),
);
}
Ok(ShmBlobRef {
offset: off + SLOT_HEADER_LEN as u64,
len: bytes.len() as u64,
generation,
epoch: ep,
checksum: csum,
backend: BlobBackendKind::Shm,
})
}
fn read_view(&self, descriptor: &ShmBlobRef) -> BlobView<'_> {
let off = descriptor.offset as usize;
let slot_off = off.saturating_sub(SLOT_HEADER_LEN);
if slot_off + SLOT_HEADER_LEN > self.capacity {
return None;
}
let slot = unsafe { &*(self.base.add(slot_off) as *const SlotHeader) };
if slot.generation != descriptor.generation {
return None;
}
if slot.len != descriptor.len {
return None;
}
if slot.checksum != descriptor.checksum {
return None;
}
if unsafe { (*self.header).epoch.load(Ordering::Acquire) } != descriptor.epoch {
return None;
}
if off + descriptor.len as usize > self.capacity {
return None;
}
unsafe {
Some(std::slice::from_raw_parts(
self.base.add(off),
descriptor.len as usize,
))
}
}
fn advance_epoch(&mut self) {
unsafe {
(*self.header).epoch.fetch_add(1, Ordering::AcqRel);
}
}
}
fn ensure_leading_slash(name: &str) -> std::ffi::CString {
let prefixed = if name.starts_with('/') {
name.to_string()
} else {
format!("/{name}")
};
std::ffi::CString::new(prefixed).expect("shm name contains no NUL")
}
fn shm_io_err(e: io::Error) -> ShmBlobArenaError {
ShmBlobArenaError::BackendIo {
detail: e.to_string(),
}
}
fn fnv1a_64(bytes: &[u8]) -> u64 {
const FNV_OFFSET_BASIS: u64 = 0xcbf2_9ce4_8422_2325;
const FNV_PRIME: u64 = 0x0000_0100_0000_01b3;
bytes.iter().fold(FNV_OFFSET_BASIS, |hash, byte| {
(hash ^ u64::from(*byte)).wrapping_mul(FNV_PRIME)
})
}
}
#[cfg(all(unix, feature = "shm"))]
pub use shm::ShmBackend;
pub fn spill_value(value: &mut IpcValue, backend: &mut dyn BlobBackend, threshold: usize) -> usize {
if let IpcValue::Inline(bytes) = value
&& bytes.len() >= threshold
{
match backend.write(bytes) {
Ok(descriptor) => {
let spilled = bytes.len();
*value = IpcValue::SharedBlob(descriptor);
return spilled;
}
Err(_) => return 0,
}
}
0
}
fn spill_state(state: &mut NodeState, backend: &mut dyn BlobBackend, threshold: usize) -> usize {
if let NodeState::Payload(bytes) = state
&& bytes.len() >= threshold
{
match backend.write(bytes) {
Ok(descriptor) => {
let spilled = bytes.len();
*state = NodeState::SharedBlob(descriptor);
return spilled;
}
Err(_) => return 0,
}
}
0
}
pub fn spill_message(
message: &mut IpcMessage,
backend: &mut dyn BlobBackend,
threshold: usize,
) -> usize {
let mut total = 0;
match message {
IpcMessage::Snapshot(snap) => {
for node in &mut snap.nodes {
total += spill_state(&mut node.state, backend, threshold);
}
}
IpcMessage::Delta(delta) => {
for op in &mut delta.ops {
match op {
DeltaOp::CellSet { payload, .. } => {
total += spill_value(payload, backend, threshold);
}
DeltaOp::SlotValue { payload, .. } => {
total += spill_value(payload, backend, threshold);
}
DeltaOp::NodeAdd { state, .. } => {
total += spill_state(state, backend, threshold);
}
_ => {}
}
}
}
IpcMessage::CrdtSync(sync) => {
for op in &mut sync.ops {
total += spill_value(&mut op.state, backend, threshold);
}
}
}
total
}
pub fn resolve_value<'a>(value: &'a IpcValue, backend: &'a dyn BlobBackend) -> BlobView<'a> {
match value {
IpcValue::Inline(bytes) => Some(bytes.as_slice()),
IpcValue::SharedBlob(descriptor) => backend.read_view(descriptor),
}
}
pub struct BlobRouter<'a> {
backends: [Option<&'a dyn BlobBackend>; 3],
}
impl<'a> BlobRouter<'a> {
pub fn new() -> Self {
Self {
backends: [None, None, None],
}
}
pub fn register(&mut self, backend: &'a dyn BlobBackend) -> &mut Self {
self.backends[backend.kind() as usize] = Some(backend);
self
}
pub fn read_view(&self, descriptor: &ShmBlobRef) -> BlobView<'_> {
let idx = descriptor.backend as usize;
self.backends[idx].and_then(|b| b.read_view(descriptor))
}
pub fn resolve<'b>(&'b self, value: &'b IpcValue) -> BlobView<'b> {
match value {
IpcValue::Inline(bytes) => Some(bytes.as_slice()),
IpcValue::SharedBlob(descriptor) => self.read_view(descriptor),
}
}
}
impl<'a> Default for BlobRouter<'a> {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
fn bytes_eq(view: BlobView<'_>, expected: &[u8]) -> bool {
match view {
Some(b) => b == expected,
None => false,
}
}
#[test]
fn in_process_resolve_write() {
let mut backend = InProcessBackend::new().unwrap();
let payload = [1, 2, 3, 4, 5, 6, 7, 8];
let desc = backend.write(&payload).unwrap();
assert_eq!(desc.backend, BlobBackendKind::InProcess);
assert!(bytes_eq(backend.read_view(&desc), &payload));
}
#[test]
fn arrow_resolve_write() {
let mut backend = ArrowBackend::new().unwrap();
let payload = [10, 20, 30, 40];
let desc = backend.write(&payload).unwrap();
assert_eq!(desc.backend, BlobBackendKind::Arrow);
assert!(bytes_eq(backend.read_view(&desc), &payload));
}
#[test]
fn backend_isolation() {
let mut inproc = InProcessBackend::new().unwrap();
let desc = inproc.write(&[9, 9, 9]).unwrap();
let router = BlobRouter::new(); assert_eq!(router.read_view(&desc), None);
let mut router = BlobRouter::new();
router.register(&inproc);
assert!(router.read_view(&desc).is_some());
let mut shm_desc = desc;
shm_desc.backend = BlobBackendKind::Shm;
assert_eq!(router.read_view(&shm_desc), None);
}
#[test]
fn stale_generation_rejects() {
let mut backend = InProcessBackend::new().unwrap();
let desc = backend.write(&[1, 2, 3]).unwrap();
let mut stale = desc;
stale.generation += 1;
assert_eq!(backend.read_view(&stale), None);
}
#[test]
fn corrupt_checksum_rejects() {
let mut backend = InProcessBackend::new().unwrap();
let desc = backend.write(&[4, 5, 6]).unwrap();
let mut corrupt = desc;
corrupt.checksum = corrupt.checksum.wrapping_add(1);
assert_eq!(backend.read_view(&corrupt), None);
}
#[test]
fn epoch_advance_invalidates() {
let mut backend = InProcessBackend::new().unwrap();
let desc = backend.write(&[7, 8]).unwrap();
assert!(backend.read_view(&desc).is_some());
backend.advance_epoch();
assert_eq!(backend.read_view(&desc), None);
}
#[test]
fn spill_resolve_round_trip() {
use crate::{Delta, DeltaOp, NodeId};
let mut backend = InProcessBackend::new().unwrap();
let big = vec![0x5Au8; 500];
let mut msg = IpcMessage::Delta(Delta::next(
1,
vec![DeltaOp::slot_value(NodeId(7), big.clone())],
));
let spilled = spill_message(&mut msg, &mut backend, 64);
assert_eq!(spilled, big.len());
let router = BlobRouter::new();
let mut router = router;
router.register(&backend);
let IpcMessage::Delta(delta) = &msg else {
panic!("expected Delta");
};
let DeltaOp::SlotValue { payload, .. } = &delta.ops[0] else {
panic!("expected SlotValue");
};
assert!(matches!(payload, IpcValue::SharedBlob(_)));
assert!(bytes_eq(router.resolve(payload), &big));
}
#[test]
fn spill_snapshot_and_crdt() {
use crate::{CrdtOp, CrdtSync, NodeId, NodeSnapshot, Snapshot, WireStamp};
let mut backend = InProcessBackend::new().unwrap();
let big = vec![0xABu8; 300];
let mut msg = IpcMessage::Snapshot(Snapshot::new(
1,
vec![NodeSnapshot::payload(NodeId(1), "blob", big.clone())],
vec![],
vec![NodeId(1)],
));
let spilled = spill_message(&mut msg, &mut backend, 64);
assert_eq!(spilled, big.len());
let stamp = WireStamp {
wall_time: 1,
logical: 0,
peer: 1,
};
let mut crdt_msg = IpcMessage::CrdtSync(CrdtSync::new(
vec![(1, stamp)],
vec![CrdtOp::new(NodeId(1), stamp, big.clone())],
));
let spilled = spill_message(&mut crdt_msg, &mut backend, 64);
assert_eq!(spilled, big.len());
}
#[test]
fn sub_threshold_stays_inline() {
use crate::{Delta, DeltaOp, NodeId};
let mut backend = InProcessBackend::new().unwrap();
let mut msg = IpcMessage::Delta(Delta::next(
1,
vec![DeltaOp::slot_value(NodeId(1), vec![1, 2, 3])],
));
let spilled = spill_message(&mut msg, &mut backend, 64);
assert_eq!(spilled, 0);
let IpcMessage::Delta(delta) = &msg else {
panic!("expected Delta");
};
assert!(
matches!(delta.ops[0], DeltaOp::SlotValue { ref payload, .. } if matches!(payload, IpcValue::Inline(_)))
);
}
#[test]
fn multi_backend_routing() {
let mut inproc = InProcessBackend::new().unwrap();
let mut arrow = ArrowBackend::new().unwrap();
let inproc_desc = inproc.write(b"inproc bytes").unwrap();
let arrow_desc = arrow.write(b"arrow bytes").unwrap();
let mut router = BlobRouter::new();
router.register(&inproc).register(&arrow);
assert!(bytes_eq(router.read_view(&inproc_desc), b"inproc bytes"));
assert!(bytes_eq(router.read_view(&arrow_desc), b"arrow bytes"));
}
#[test]
fn arrow_ipc_stream_bytes() {
let mut arrow = ArrowBackend::new().unwrap();
let ipc_stream = [0x41, 0x52, 0x52, 0x4f, 0x57, 0x31, 0x00, 0x00];
let desc = arrow.write(&ipc_stream).unwrap();
assert_eq!(desc.backend, BlobBackendKind::Arrow);
assert!(bytes_eq(arrow.read_view(&desc), &ipc_stream));
}
#[cfg(all(unix, feature = "shm"))]
#[test]
fn shm_backend_round_trip() {
let name = format!("/lazily_shm_test_{}", std::process::id());
ShmBackend::unlink(&name);
let mut backend = ShmBackend::create(&name, 1 << 20).unwrap();
let payload: Vec<u8> = (0..1000).map(|i| (i * 7 + 1) as u8).collect();
let desc = backend.write(&payload).unwrap();
assert_eq!(desc.backend, BlobBackendKind::Shm);
assert!(bytes_eq(backend.read_view(&desc), &payload));
backend.advance_epoch();
assert_eq!(backend.read_view(&desc), None); ShmBackend::unlink(&name);
}
#[cfg(all(unix, feature = "shm"))]
#[test]
fn shm_backend_cross_process() {
let name = format!("/lazily_shm_xproc_{}", std::process::id());
ShmBackend::unlink(&name);
let payload: Vec<u8> = (0..1000).map(|i| (i * 7 + 1) as u8).collect();
let mut parent = ShmBackend::create(&name, 1 << 20).unwrap();
let desc = parent.write(&payload).unwrap();
assert_eq!(desc.backend, BlobBackendKind::Shm);
let pid = unsafe { libc::fork() };
if pid == 0 {
let child = ShmBackend::open(&name).unwrap();
let view = child.read_view(&desc);
let ok = matches!(view, Some(b) if b == payload.as_slice());
unsafe { libc::_exit(if ok { 0 } else { 1 }) };
}
let mut status = 0i32;
unsafe { libc::waitpid(pid, &mut status, 0) };
assert!(libc::WIFEXITED(status) && libc::WEXITSTATUS(status) == 0);
ShmBackend::unlink(&name);
}
}