use std::collections::{BTreeMap, HashMap};
use std::fs::OpenOptions;
use std::io::{self, Read, Write};
use std::net::{
Ipv4Addr, Shutdown, SocketAddr, SocketAddrV4, TcpListener, TcpStream, ToSocketAddrs, UdpSocket,
};
use std::os::unix::io::{AsRawFd, FromRawFd};
use std::path::{Path, PathBuf};
use std::ptr;
use std::sync::atomic::{AtomicU32, AtomicU64, Ordering, fence};
use std::sync::{Arc, Mutex};
use std::thread::{self, JoinHandle};
use std::time::{Duration, Instant};
use memmap2::{Mmap, MmapMut};
use crate::arena::{ArenaHeader, BlobRef, PayloadArena};
use crate::error::{Result, RingfireError};
use crate::header::{
FLAG_MODE_SPMC, FLAG_POLICY_LATEST_WINS, FLAG_SPARSE, FLAG_WITH_ARENA, FLAG_WITH_REGISTRY,
ReaderSlot, RingHeader, RingLayout, RingView, SLOT_WRITING, validate_ring,
};
use crate::registry::ReaderRegistry;
use crate::shm::create_backing_file;
use crate::spmc::oldest_retained;
use crate::wait::wake_futex;
pub const REPLICATION_MAGIC: u64 = 0x5249_4E47_4D49_5252;
pub const REPLICATION_VERSION: u32 = 1;
pub const HELLO_LATEST: u64 = 0;
pub const HELLO_OLDEST: u64 = u64::MAX;
pub const DEFAULT_MTU: usize = 1472;
const FRAME_HEADER_LEN: usize = 16;
const HELLO_LEN: usize = 16;
const GEOMETRY_LEN: usize = 48;
const NAK_LEN: usize = 8;
const BLOB_REF_LEN: usize = 16;
const UDP_MAX_PAYLOAD: usize = 65_507;
const TCP_FRAME_BYTES: usize = 1 << 20;
const MULTICAST_LEN: usize = 16;
const KIND_HELLO: u8 = 1;
const KIND_GEOMETRY: u8 = 2;
const KIND_DATA: u8 = 3;
const KIND_GAP: u8 = 4;
const KIND_HEARTBEAT: u8 = 5;
const KIND_NAK: u8 = 6;
const KIND_MULTICAST: u8 = 7;
const KIND_PUNCH: u8 = 8;
const HELLO_UDP: u8 = 0x01;
const HELLO_PREFER_UNICAST: u8 = 0x02;
const PUNCH_INTERVAL: Duration = Duration::from_millis(200);
const PUNCH_KEEPALIVE: Duration = Duration::from_secs(5);
const GEOMETRY_RESET: u8 = 0x01;
const GEOMETRY_MULTICAST: u8 = 0x02;
const HEARTBEAT_INTERVAL: Duration = Duration::from_millis(200);
const WRITE_TIMEOUT: Duration = Duration::from_secs(5);
const DEFAULT_BATCH: usize = 256;
const MAX_DATAGRAM: usize = 65_536;
const NAK_MAX: u64 = u16::MAX as u64;
const PENDING_MAX: usize = 8192;
const POLL_SLICE: Duration = Duration::from_millis(10);
const _: () = assert!(
cfg!(target_endian = "little"),
"the replication protocol carries native slot bytes and assumes little-endian hosts"
);
const DATAGRAMS_PER_STEP: usize = 32;
const ADAPTIVE_PACE: Duration = Duration::from_micros(50);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Geometry {
pub capacity: u64,
pub element_size: u32,
pub flags: u32,
pub schema_sig: u64,
pub registry_count: u32,
pub slots_offset: u32,
pub arena_offset: u64,
pub arena_size: u64,
}
impl Geometry {
fn from_header(header: &RingHeader, view: &RingView) -> Self {
Self {
capacity: view.capacity,
element_size: view.slot_size as u32,
flags: header.flags,
schema_sig: header.schema_sig,
registry_count: header.reader_registry_count,
slots_offset: view.slots_offset as u32,
arena_offset: header.arena_offset,
arena_size: header.arena_size,
}
}
pub fn payload_len(&self) -> usize {
self.element_size as usize - 8
}
pub fn has_arena(&self) -> bool {
self.arena_offset != 0
}
fn same_layout(&self, other: &Geometry) -> bool {
self.capacity == other.capacity
&& self.element_size == other.element_size
&& self.schema_sig == other.schema_sig
&& self.registry_count == other.registry_count
&& self.slots_offset == other.slots_offset
&& self.arena_offset == other.arena_offset
&& self.arena_size == other.arena_size
}
fn mirror_flags(&self) -> u32 {
FLAG_MODE_SPMC
| FLAG_POLICY_LATEST_WINS
| FLAG_SPARSE
| if self.registry_count > 0 {
FLAG_WITH_REGISTRY
} else {
0
}
| if self.has_arena() { FLAG_WITH_ARENA } else { 0 }
}
fn total_len(&self) -> usize {
if self.has_arena() {
self.arena_offset as usize
+ std::mem::size_of::<ArenaHeader>()
+ self.arena_size as usize
} else {
self.slots_offset as usize + self.capacity as usize * self.element_size as usize
}
}
fn encode(&self, out: &mut [u8; GEOMETRY_LEN]) {
out[0..8].copy_from_slice(&self.capacity.to_le_bytes());
out[8..12].copy_from_slice(&self.element_size.to_le_bytes());
out[12..16].copy_from_slice(&self.flags.to_le_bytes());
out[16..24].copy_from_slice(&self.schema_sig.to_le_bytes());
out[24..28].copy_from_slice(&self.registry_count.to_le_bytes());
out[28..32].copy_from_slice(&self.slots_offset.to_le_bytes());
out[32..40].copy_from_slice(&self.arena_offset.to_le_bytes());
out[40..48].copy_from_slice(&self.arena_size.to_le_bytes());
}
fn decode(b: &[u8; GEOMETRY_LEN]) -> Result<Self> {
let g = Self {
capacity: u64::from_le_bytes(b[0..8].try_into().unwrap()),
element_size: u32::from_le_bytes(b[8..12].try_into().unwrap()),
flags: u32::from_le_bytes(b[12..16].try_into().unwrap()),
schema_sig: u64::from_le_bytes(b[16..24].try_into().unwrap()),
registry_count: u32::from_le_bytes(b[24..28].try_into().unwrap()),
slots_offset: u32::from_le_bytes(b[28..32].try_into().unwrap()),
arena_offset: u64::from_le_bytes(b[32..40].try_into().unwrap()),
arena_size: u64::from_le_bytes(b[40..48].try_into().unwrap()),
};
if !g.capacity.is_power_of_two() {
return Err(RingfireError::Protocol(
"geometry capacity is not a power of two",
));
}
if g.element_size < 8 || !g.element_size.is_multiple_of(8) {
return Err(RingfireError::Protocol(
"geometry element size is not a multiple of 8",
));
}
let min_offset = std::mem::size_of::<RingHeader>()
+ g.registry_count as usize * std::mem::size_of::<ReaderSlot>();
if (g.slots_offset as usize) < min_offset || !g.slots_offset.is_multiple_of(64) {
return Err(RingfireError::Protocol(
"geometry slots offset is misplaced",
));
}
let Some(slots_end) = (g.capacity as usize)
.checked_mul(g.element_size as usize)
.and_then(|b| b.checked_add(g.slots_offset as usize))
else {
return Err(RingfireError::Protocol("geometry does not fit in memory"));
};
if g.has_arena() {
if (g.arena_offset as usize) < slots_end || !g.arena_offset.is_multiple_of(64) {
return Err(RingfireError::Protocol(
"geometry arena offset is misplaced",
));
}
if !g.arena_size.is_power_of_two() || g.arena_size < 64 {
return Err(RingfireError::Protocol(
"geometry arena size is not a power of two",
));
}
if g.payload_len() < BLOB_REF_LEN {
return Err(RingfireError::Protocol(
"geometry arena ring has no blob descriptor",
));
}
if (g.arena_offset as usize)
.checked_add(std::mem::size_of::<ArenaHeader>())
.and_then(|b| b.checked_add(g.arena_size as usize))
.is_none()
{
return Err(RingfireError::Protocol("geometry does not fit in memory"));
}
} else if g.arena_size != 0 {
return Err(RingfireError::Protocol(
"geometry arena size without an arena",
));
}
Ok(g)
}
}
#[derive(Debug, Clone, Copy)]
pub struct MulticastConfig {
pub group: Ipv4Addr,
pub port: u16,
pub interface: Ipv4Addr,
pub mtu: usize,
pub ttl: u8,
pub heartbeat: Duration,
#[doc(hidden)]
pub drop_every: u64,
#[doc(hidden)]
pub swap_every: u64,
}
impl MulticastConfig {
pub fn new(group: Ipv4Addr, port: u16) -> Self {
Self {
group,
port,
interface: Ipv4Addr::UNSPECIFIED,
mtu: DEFAULT_MTU,
ttl: 1,
heartbeat: Duration::from_millis(1),
drop_every: 0,
swap_every: 0,
}
}
pub fn interface(mut self, interface: Ipv4Addr) -> Self {
self.interface = interface;
self
}
pub fn mtu(mut self, mtu: usize) -> Self {
self.mtu = mtu.clamp(FRAME_HEADER_LEN + 8, MAX_DATAGRAM);
self
}
pub fn ttl(mut self, ttl: u8) -> Self {
self.ttl = ttl;
self
}
pub fn heartbeat(mut self, heartbeat: Duration) -> Self {
self.heartbeat = heartbeat;
self
}
#[doc(hidden)]
pub fn drop_every(mut self, n: u64) -> Self {
self.drop_every = n;
self
}
#[doc(hidden)]
pub fn swap_every(mut self, n: u64) -> Self {
self.swap_every = n;
self
}
fn encode(&self, session: u8, out: &mut [u8; MULTICAST_LEN]) {
encode_udp_info(self.group, self.port, self.mtu, self.ttl, session, 0, out);
}
}
fn encode_udp_info(
group: Ipv4Addr,
port: u16,
mtu: usize,
ttl: u8,
session: u8,
token: u32,
out: &mut [u8; MULTICAST_LEN],
) {
out[0..4].copy_from_slice(&group.octets());
out[4..6].copy_from_slice(&port.to_le_bytes());
out[6..8].copy_from_slice(&(mtu.min(u16::MAX as usize) as u16).to_le_bytes());
out[8] = ttl;
out[9] = session;
out[10..14].copy_from_slice(&token.to_le_bytes());
out[14..16].fill(0);
}
#[derive(Debug, Clone, Copy)]
struct MulticastInfo {
group: Ipv4Addr,
port: u16,
session: u8,
token: u32,
}
impl MulticastInfo {
fn decode(b: &[u8; MULTICAST_LEN]) -> Result<Self> {
let group = Ipv4Addr::new(b[0], b[1], b[2], b[3]);
if !group.is_multicast() && !group.is_unspecified() {
return Err(RingfireError::Protocol(
"MULTICAST group is not a multicast address",
));
}
Ok(Self {
group,
port: u16::from_le_bytes(b[4..6].try_into().unwrap()),
session: b[9],
token: u32::from_le_bytes(b[10..14].try_into().unwrap()),
})
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct Frame {
kind: u8,
flags: u8,
count: u16,
len: u32,
seq: u64,
}
impl Frame {
fn control(kind: u8, seq: u64) -> Self {
Self {
kind,
flags: 0,
count: 0,
len: 0,
seq,
}
}
fn encode(&self) -> [u8; FRAME_HEADER_LEN] {
let mut b = [0u8; FRAME_HEADER_LEN];
b[0] = self.kind;
b[1] = self.flags;
b[2..4].copy_from_slice(&self.count.to_le_bytes());
b[4..8].copy_from_slice(&self.len.to_le_bytes());
b[8..16].copy_from_slice(&self.seq.to_le_bytes());
b
}
fn decode(b: &[u8; FRAME_HEADER_LEN]) -> Self {
Self {
kind: b[0],
flags: b[1],
count: u16::from_le_bytes(b[2..4].try_into().unwrap()),
len: u32::from_le_bytes(b[4..8].try_into().unwrap()),
seq: u64::from_le_bytes(b[8..16].try_into().unwrap()),
}
}
}
fn protocol(what: &'static str) -> io::Error {
io::Error::new(io::ErrorKind::InvalidData, what)
}
fn peer_gone(e: &io::Error) -> bool {
matches!(
e.kind(),
io::ErrorKind::UnexpectedEof
| io::ErrorKind::ConnectionReset
| io::ErrorKind::ConnectionAborted
| io::ErrorKind::BrokenPipe
)
}
fn wait_readable(fds: &[i32], timeout: Duration) {
let mut polls: Vec<libc::pollfd> = fds
.iter()
.map(|&fd| libc::pollfd {
fd,
events: libc::POLLIN,
revents: 0,
})
.collect();
let ms = timeout.as_millis().min(i32::MAX as u128) as i32;
unsafe { libc::poll(polls.as_mut_ptr(), polls.len() as libc::nfds_t, ms) };
}
fn read_full(stream: &mut TcpStream, buf: &mut [u8], spin: bool) -> io::Result<()> {
let mut filled = 0;
while filled < buf.len() {
match stream.read(&mut buf[filled..]) {
Ok(0) => return Err(io::Error::from(io::ErrorKind::UnexpectedEof)),
Ok(n) => filled += n,
Err(e) if e.kind() == io::ErrorKind::WouldBlock => {
if spin {
core::hint::spin_loop();
} else {
wait_readable(&[stream.as_raw_fd()], POLL_SLICE);
}
}
Err(e) if e.kind() == io::ErrorKind::Interrupted => {}
Err(e) => return Err(e),
}
}
Ok(())
}
fn write_full(stream: &mut TcpStream, mut buf: &[u8]) -> io::Result<()> {
while !buf.is_empty() {
match stream.write(buf) {
Ok(0) => return Err(io::Error::from(io::ErrorKind::WriteZero)),
Ok(n) => buf = &buf[n..],
Err(e) if e.kind() == io::ErrorKind::WouldBlock => {
let mut poll = libc::pollfd {
fd: stream.as_raw_fd(),
events: libc::POLLOUT,
revents: 0,
};
unsafe { libc::poll(&mut poll, 1, POLL_SLICE.as_millis() as i32) };
}
Err(e) if e.kind() == io::ErrorKind::Interrupted => {}
Err(e) => return Err(e),
}
}
Ok(())
}
fn set_sockopt<T>(fd: i32, level: i32, name: i32, value: &T) -> io::Result<()> {
let rc = unsafe {
libc::setsockopt(
fd,
level,
name,
(value as *const T).cast(),
std::mem::size_of::<T>() as libc::socklen_t,
)
};
if rc == 0 {
Ok(())
} else {
Err(io::Error::last_os_error())
}
}
fn udp_sender(port: u16, multicast: Option<&MulticastConfig>) -> io::Result<UdpSocket> {
let sock = UdpSocket::bind(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, port))?;
if let Some(cfg) = multicast {
sock.set_multicast_ttl_v4(u32::from(cfg.ttl))?;
sock.set_multicast_loop_v4(true)?;
if !cfg.interface.is_unspecified() {
let addr = libc::in_addr {
s_addr: u32::from(cfg.interface).to_be(),
};
set_sockopt(
sock.as_raw_fd(),
libc::IPPROTO_IP,
libc::IP_MULTICAST_IF,
&addr,
)?;
}
}
sock.set_nonblocking(true)?;
Ok(sock)
}
fn multicast_receiver(
group: Ipv4Addr,
port: u16,
interface: Ipv4Addr,
rcvbuf: usize,
) -> io::Result<UdpSocket> {
let fd = unsafe { libc::socket(libc::AF_INET, libc::SOCK_DGRAM, 0) };
if fd < 0 {
return Err(io::Error::last_os_error());
}
let sock = unsafe { UdpSocket::from_raw_fd(fd) };
unsafe { libc::fcntl(fd, libc::F_SETFD, libc::FD_CLOEXEC) };
set_sockopt(fd, libc::SOL_SOCKET, libc::SO_REUSEADDR, &1i32)?;
#[cfg(not(target_os = "linux"))]
set_sockopt(fd, libc::SOL_SOCKET, libc::SO_REUSEPORT, &1i32)?;
if rcvbuf > 0 {
let bytes = rcvbuf.min(i32::MAX as usize) as i32;
set_sockopt(fd, libc::SOL_SOCKET, libc::SO_RCVBUF, &bytes)?;
}
let mut addr: libc::sockaddr_in = unsafe { std::mem::zeroed() };
addr.sin_family = libc::AF_INET as libc::sa_family_t;
addr.sin_port = port.to_be();
addr.sin_addr.s_addr = u32::from(Ipv4Addr::UNSPECIFIED).to_be();
let rc = unsafe {
libc::bind(
fd,
(&addr as *const libc::sockaddr_in).cast(),
std::mem::size_of::<libc::sockaddr_in>() as libc::socklen_t,
)
};
if rc != 0 {
return Err(io::Error::last_os_error());
}
if group.is_multicast() {
sock.join_multicast_v4(&group, &interface)?;
}
sock.set_nonblocking(true)?;
Ok(sock)
}
fn send_datagram(sock: &UdpSocket, dest: SocketAddrV4, bytes: &[u8]) -> io::Result<()> {
loop {
match sock.send_to(bytes, dest) {
Ok(_) => return Ok(()),
Err(e) if e.kind() == io::ErrorKind::WouldBlock => core::hint::spin_loop(),
Err(e) if e.kind() == io::ErrorKind::Interrupted => {}
Err(e) => return Err(e),
}
}
}
enum RawRead {
Item,
Pending,
Overwritten(u64),
Lost,
}
struct Collected {
count: usize,
lapped: Option<u64>,
lost: bool,
}
fn blob_ref_at(payload: &[u8]) -> BlobRef {
let b = &payload[payload.len() - BLOB_REF_LEN..];
BlobRef {
offset: u64::from_le_bytes(b[0..8].try_into().unwrap()),
len: u32::from_le_bytes(b[8..12].try_into().unwrap()),
flags: u32::from_le_bytes(b[12..16].try_into().unwrap()),
}
}
fn put_blob_ref(payload: &mut [u8], r: BlobRef) {
let n = payload.len();
let b = &mut payload[n - BLOB_REF_LEN..];
b[0..8].copy_from_slice(&r.offset.to_le_bytes());
b[8..12].copy_from_slice(&r.len.to_le_bytes());
b[12..16].copy_from_slice(&r.flags.to_le_bytes());
}
struct ArenaView {
header: *const ArenaHeader,
data: *const u8,
capacity: u64,
mask: u64,
}
impl ArenaView {
unsafe fn open(base: *const u8, geometry: &Geometry) -> Result<Self> {
let header = unsafe { base.add(geometry.arena_offset as usize) }.cast::<ArenaHeader>();
let (capacity, mask) = unsafe { ((*header).capacity, (*header).mask) };
if capacity != geometry.arena_size || mask != capacity - 1 {
return Err(RingfireError::CorruptLayout(
"arena header disagrees with ring header",
));
}
Ok(Self {
header,
data: unsafe { header.cast::<u8>().add(std::mem::size_of::<ArenaHeader>()) },
capacity,
mask,
})
}
fn in_bounds(&self, r: BlobRef) -> bool {
(r.offset & self.mask) + r.len as u64 <= self.capacity
}
fn copy_into(&self, r: BlobRef, out: &mut Vec<u8>) {
let len = r.len as usize;
out.reserve(len);
unsafe {
ptr::copy_nonoverlapping(
self.data.add((r.offset & self.mask) as usize),
out.as_mut_ptr().add(out.len()),
len,
);
out.set_len(out.len() + len);
}
}
fn is_lapped(&self, r: BlobRef) -> bool {
fence(Ordering::Acquire);
let reserved = unsafe { (*self.header).reserved.load(Ordering::Relaxed) };
reserved.saturating_sub(r.offset) > self.capacity
}
}
struct SourceRing {
base: *const u8,
view: RingView,
geometry: Geometry,
arena: Option<ArenaView>,
_mmap: Mmap,
}
unsafe impl Send for SourceRing {}
impl SourceRing {
fn open(path: &Path) -> Result<Self> {
let file = OpenOptions::new().read(true).open(path)?;
let mmap = unsafe { Mmap::map(&file)? };
let view = unsafe { validate_ring(mmap.as_ptr(), mmap.len(), None, 8)? };
let header = unsafe { &*mmap.as_ptr().cast::<RingHeader>() };
let geometry = Geometry::from_header(header, &view);
let arena = if header.flags & FLAG_WITH_ARENA != 0 {
if geometry.payload_len() < BLOB_REF_LEN {
return Err(RingfireError::Unsupported(
"arena ring without a blob descriptor",
));
}
Some(unsafe { ArenaView::open(mmap.as_ptr(), &geometry)? })
} else {
None
};
Ok(Self {
base: mmap.as_ptr(),
view,
geometry,
arena,
_mmap: mmap,
})
}
fn header(&self) -> &RingHeader {
unsafe { &*self.base.cast::<RingHeader>() }
}
fn write_seq(&self) -> u64 {
self.header().write_seq.load(Ordering::Acquire)
}
fn read(&self, want: u64, out: &mut [u8]) -> RawRead {
let slot = unsafe {
self.base.add(
self.view.slots_offset + (want & self.view.mask) as usize * self.view.slot_size,
)
};
let seq = unsafe { &*slot.cast::<AtomicU64>() };
let s1 = seq.load(Ordering::Acquire);
if s1 == want {
unsafe { ptr::copy_nonoverlapping(slot.add(8), out.as_mut_ptr(), out.len()) };
fence(Ordering::Acquire);
let s2 = seq.load(Ordering::Relaxed);
if s2 == want {
RawRead::Item
} else {
RawRead::Overwritten(if s2 == SLOT_WRITING { 0 } else { s2 })
}
} else if s1 == SLOT_WRITING || s1 < want {
RawRead::Pending
} else {
RawRead::Overwritten(s1)
}
}
fn read_record(&self, want: u64, out: &mut Vec<u8>) -> RawRead {
let payload_len = self.geometry.payload_len();
let start = out.len();
out.resize(start + payload_len, 0);
match self.read(want, &mut out[start..]) {
RawRead::Item => {}
other => {
out.truncate(start);
return other;
}
}
if let Some(arena) = &self.arena {
let r = blob_ref_at(&out[start..]);
if r.len > 0 {
if !arena.in_bounds(r) {
out.truncate(start);
return RawRead::Lost;
}
arena.copy_into(r, out);
if arena.is_lapped(r) {
out.truncate(start);
return RawRead::Lost;
}
}
}
RawRead::Item
}
fn collect(
&self,
cursor: u64,
max_records: usize,
max_bytes: usize,
out: &mut Vec<u8>,
) -> Collected {
let mut count = 0usize;
while count < max_records {
let before = out.len();
match self.read_record(cursor + count as u64, out) {
RawRead::Item => {
if count > 0 && out.len() > max_bytes {
out.truncate(before);
break;
}
count += 1;
if out.len() >= max_bytes {
break;
}
}
RawRead::Pending => break,
RawRead::Overwritten(seen) => {
return Collected {
count,
lapped: Some(seen),
lost: false,
};
}
RawRead::Lost => {
return Collected {
count,
lapped: None,
lost: true,
};
}
}
}
Collected {
count,
lapped: None,
lost: false,
}
}
fn resync(&self, seen: u64) -> u64 {
oldest_retained(self.write_seq().max(seen), self.view.capacity)
}
fn collect_lingering(
&self,
cursor: u64,
max_records: usize,
max_bytes: usize,
out: &mut Vec<u8>,
linger: Duration,
) -> Collected {
let mut got = self.collect(cursor, max_records, max_bytes, out);
let full = |got: &Collected, out: &Vec<u8>| {
got.count >= max_records || out.len() >= max_bytes || got.lapped.is_some() || got.lost
};
if got.count == 0 || full(&got, out) || linger.is_zero() {
return got;
}
let until = Instant::now() + linger;
while !full(&got, out) && Instant::now() < until {
let more = self.collect(
cursor + got.count as u64,
max_records - got.count,
max_bytes,
out,
);
got.count += more.count;
got.lapped = more.lapped;
got.lost = more.lost;
if more.count == 0 {
core::hint::spin_loop();
}
}
got
}
}
fn data_frame(out: &mut [u8], flags: u8, count: usize, seq: u64) -> &[u8] {
let frame = Frame {
kind: KIND_DATA,
flags,
count: count as u16,
len: (out.len() - FRAME_HEADER_LEN) as u32,
seq,
};
out[..FRAME_HEADER_LEN].copy_from_slice(&frame.encode());
out
}
struct MirrorRing {
base: *mut u8,
view: RingView,
geometry: Geometry,
next_seq: u64,
arena: Option<PayloadArena>,
_mmap: MmapMut,
_file: std::fs::File,
}
unsafe impl Send for MirrorRing {}
impl MirrorRing {
fn create(path: &Path, geometry: &Geometry, mode: u32) -> Result<Self> {
let header_size = std::mem::size_of::<RingHeader>();
let registry_count = geometry.registry_count as usize;
let slot_size = geometry.element_size as usize;
let slots_offset = geometry.slots_offset as usize;
let capacity = geometry.capacity;
let total = geometry.total_len();
let file = create_backing_file(path, mode, true, total as u64)?;
let mut mmap = unsafe { MmapMut::map_mut(&file)? };
let base = mmap.as_mut_ptr();
let arena = unsafe {
RingHeader::initialize(
base.cast(),
&RingLayout {
capacity,
slot_size,
flags: geometry.mirror_flags(),
schema_sig: geometry.schema_sig,
claim_seq: 0,
read_seq: 0,
registry_offset: if registry_count > 0 { header_size } else { 0 },
registry_count,
slots_offset,
arena_offset: geometry.arena_offset as usize,
arena_size: geometry.arena_size as usize,
},
);
if registry_count > 0 {
let _ = ReaderRegistry::init(base.add(header_size), registry_count);
}
for i in 0..capacity as usize {
let seq = &*base.add(slots_offset + i * slot_size).cast::<AtomicU64>();
seq.store(0, Ordering::Relaxed);
}
let arena = if geometry.has_arena() {
Some(PayloadArena::init(
base.add(geometry.arena_offset as usize),
geometry.arena_size as usize,
)?)
} else {
None
};
RingHeader::publish(base.cast());
arena
};
Ok(Self {
base,
view: RingView {
capacity,
mask: capacity - 1,
slot_size,
slots_offset,
},
geometry: *geometry,
next_seq: 1,
arena,
_mmap: mmap,
_file: file,
})
}
fn adopt(path: &Path) -> Result<Option<Self>> {
let file = match OpenOptions::new().read(true).write(true).open(path) {
Ok(file) => file,
Err(e) if e.kind() == io::ErrorKind::NotFound => return Ok(None),
Err(e) => return Err(e.into()),
};
if unsafe { libc::flock(file.as_raw_fd(), libc::LOCK_EX | libc::LOCK_NB) } != 0 {
return Err(RingfireError::ProducerAlreadyExists);
}
let Ok(mut mmap) = (unsafe { MmapMut::map_mut(&file) }) else {
return Ok(None);
};
let base = mmap.as_mut_ptr();
let Ok(view) = (unsafe { validate_ring(base, mmap.len(), None, 8) }) else {
return Ok(None);
};
let header = unsafe { &*base.cast::<RingHeader>() };
let geometry = Geometry::from_header(header, &view);
let arena = if geometry.has_arena() {
match unsafe { PayloadArena::from_ptr(base.add(geometry.arena_offset as usize)) } {
Ok(arena) if arena.capacity() as u64 == geometry.arena_size => Some(arena),
_ => return Ok(None),
}
} else {
None
};
let next_seq = header.write_seq.load(Ordering::Acquire) + 1;
Ok(Some(Self {
base,
view,
geometry,
next_seq,
arena,
_mmap: mmap,
_file: file,
}))
}
fn header(&self) -> &RingHeader {
unsafe { &*self.base.cast::<RingHeader>() }
}
fn slot(&self, seq: u64) -> *mut u8 {
unsafe {
self.base
.add(self.view.slots_offset + (seq & self.view.mask) as usize * self.view.slot_size)
}
}
fn write(&mut self, seq: u64, payload: &[u8]) {
debug_assert_eq!(seq, self.next_seq);
debug_assert_eq!(payload.len(), self.view.slot_size - 8);
let slot = self.slot(seq);
let word = unsafe { &*slot.cast::<AtomicU64>() };
word.store(SLOT_WRITING, Ordering::Relaxed);
fence(Ordering::Release);
unsafe { ptr::copy_nonoverlapping(payload.as_ptr(), slot.add(8), payload.len()) };
word.store(seq, Ordering::Release);
self.next_seq = seq + 1;
}
fn write_blob(&mut self, seq: u64, descriptor: &[u8], blob: &[u8]) -> Result<()> {
debug_assert_eq!(seq, self.next_seq);
debug_assert_eq!(descriptor.len(), self.view.slot_size - 8);
let arena = self.arena.as_ref().expect("arena ring");
let source = blob_ref_at(descriptor);
let placed = arena.write_blob(blob, source.flags)?;
let slot = self.slot(seq);
let word = unsafe { &*slot.cast::<AtomicU64>() };
word.store(SLOT_WRITING, Ordering::Relaxed);
fence(Ordering::Release);
unsafe {
ptr::copy_nonoverlapping(descriptor.as_ptr(), slot.add(8), descriptor.len());
put_blob_ref(
std::slice::from_raw_parts_mut(slot.add(8), descriptor.len()),
placed,
);
}
word.store(seq, Ordering::Release);
self.next_seq = seq + 1;
Ok(())
}
fn publish(&self) {
let header = self.header();
header.write_seq.store(self.next_seq - 1, Ordering::Release);
wake_futex(header, i32::MAX);
}
fn skip_to(&mut self, seq: u64) {
if seq > self.next_seq {
self.next_seq = seq;
}
}
}
pub struct ReplicaServer {
ring_path: PathBuf,
listener: TcpListener,
batch: usize,
spin: bool,
linger: Option<Duration>,
multicast: Option<MulticastConfig>,
unicast: Option<(u16, usize)>,
duplicate: u8,
session: u8,
peers: Arc<Mutex<HashMap<u32, Option<SocketAddrV4>>>>,
next_token: Arc<AtomicU32>,
}
struct UdpDelivery {
port: u16,
mtu: usize,
multicast: Option<MulticastConfig>,
peers: Arc<Mutex<HashMap<u32, Option<SocketAddrV4>>>>,
duplicate: u8,
heartbeat: Duration,
}
struct UdpOffer {
multicast: Option<MulticastConfig>,
unicast: Option<(u16, usize)>,
peers: Arc<Mutex<HashMap<u32, Option<SocketAddrV4>>>>,
next_token: Arc<AtomicU32>,
}
#[derive(Clone, Copy)]
struct Linger {
fixed: Option<Duration>,
last_send: Instant,
}
impl Linger {
fn new(fixed: Option<Duration>) -> Self {
Self {
fixed,
last_send: Instant::now() - ADAPTIVE_PACE,
}
}
fn current(&self) -> Duration {
match self.fixed {
Some(fixed) => fixed,
None => ADAPTIVE_PACE.saturating_sub(self.last_send.elapsed()),
}
}
fn sent(&mut self) {
self.last_send = Instant::now();
}
}
impl ReplicaServer {
pub fn bind<P: AsRef<Path>, A: ToSocketAddrs>(ring_path: P, addr: A) -> Result<Self> {
let ring_path = ring_path.as_ref().to_path_buf();
SourceRing::open(&ring_path)?;
let listener = TcpListener::bind(addr)?;
let session = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| (d.subsec_nanos() ^ d.as_secs() as u32) as u8)
.unwrap_or(1);
Ok(Self {
ring_path,
listener,
batch: DEFAULT_BATCH,
spin: false,
linger: None,
multicast: None,
unicast: None,
duplicate: 1,
session,
peers: Arc::new(Mutex::new(HashMap::new())),
next_token: Arc::new(AtomicU32::new(1)),
})
}
pub fn unicast(mut self, port: u16, mtu: usize) -> Self {
self.unicast = Some((port, mtu.clamp(FRAME_HEADER_LEN + 8, MAX_DATAGRAM)));
self
}
pub fn duplicate(mut self, copies: u8) -> Self {
self.duplicate = copies.max(1);
self
}
pub fn linger(mut self, linger: Option<Duration>) -> Self {
self.linger = linger;
self
}
pub fn batch(mut self, batch: usize) -> Self {
self.batch = batch.clamp(1, u16::MAX as usize);
self
}
pub fn spin(mut self, spin: bool) -> Self {
self.spin = spin;
self
}
pub fn multicast(mut self, config: MulticastConfig) -> Self {
self.multicast = Some(config);
self
}
pub fn local_addr(&self) -> io::Result<SocketAddr> {
self.listener.local_addr()
}
pub fn ring_path(&self) -> &Path {
&self.ring_path
}
fn start_udp(&self) -> Result<()> {
if self.multicast.is_none() && self.unicast.is_none() {
return Ok(());
}
let ring = SourceRing::open(&self.ring_path)?;
let (batch, spin, session, linger) = (self.batch, self.spin, self.session, self.linger);
let delivery = UdpDelivery {
port: self.unicast.map_or(0, |(port, _)| port),
mtu: match (self.multicast, self.unicast) {
(Some(cfg), Some((_, mtu))) => cfg.mtu.min(mtu),
(Some(cfg), None) => cfg.mtu,
(None, Some((_, mtu))) => mtu,
(None, None) => DEFAULT_MTU,
},
multicast: self.multicast,
peers: self.peers.clone(),
duplicate: self.duplicate,
heartbeat: self
.multicast
.map_or(Duration::from_millis(1), |cfg| cfg.heartbeat),
};
thread::Builder::new()
.name("ringfire-udp".into())
.spawn(move || {
if let Err(e) = udp_loop(ring, delivery, session, batch, spin, linger) {
eprintln!("ringfire udp sender stopped: {}", e);
}
})?;
Ok(())
}
pub fn run(&self) -> Result<()> {
self.start_udp()?;
loop {
let (stream, peer) = self.listener.accept()?;
let ring = SourceRing::open(&self.ring_path)?;
let (batch, spin, session, linger) = (self.batch, self.spin, self.session, self.linger);
let offer = UdpOffer {
multicast: self.multicast,
unicast: self.unicast,
peers: self.peers.clone(),
next_token: self.next_token.clone(),
};
thread::Builder::new()
.name(format!("ringfire-serve-{}", peer))
.spawn(move || {
if let Err(e) = serve_client(ring, stream, batch, spin, linger, offer, session)
&& e.kind() != io::ErrorKind::UnexpectedEof
{
eprintln!("ringfire serve: mirror {} dropped: {}", peer, e);
}
})?;
}
}
pub fn serve_one(&self) -> Result<()> {
let (stream, _) = self.listener.accept()?;
let ring = SourceRing::open(&self.ring_path)?;
let offer = UdpOffer {
multicast: None,
unicast: None,
peers: self.peers.clone(),
next_token: self.next_token.clone(),
};
serve_client(
ring,
stream,
self.batch,
self.spin,
self.linger,
offer,
self.session,
)?;
Ok(())
}
pub fn spawn(self) -> io::Result<JoinHandle<Result<()>>> {
thread::Builder::new()
.name("ringfire-serve".into())
.spawn(move || self.run())
}
}
fn udp_loop(
ring: SourceRing,
delivery: UdpDelivery,
session: u8,
batch: usize,
spin: bool,
linger: Option<Duration>,
) -> io::Result<()> {
let mut linger = Linger::new(linger);
let sock = udp_sender(delivery.port, delivery.multicast.as_ref())?;
let group = delivery
.multicast
.map(|cfg| SocketAddrV4::new(cfg.group, cfg.port));
let batch = batch.clamp(1, u16::MAX as usize);
let mut max_bytes = delivery.mtu.max(FRAME_HEADER_LEN + 1);
let mut wire: Vec<u8> = Vec::with_capacity(MAX_DATAGRAM);
let mut punch = [0u8; FRAME_HEADER_LEN];
let mut cursor = ring.write_seq() + 1;
let mut sent = 0u64;
let mut idle = 0u32;
let mut last_beat = Instant::now();
let mut held: Option<Vec<u8>> = None;
let (drop_every, swap_every) = delivery
.multicast
.map_or((0, 0), |cfg| (cfg.drop_every, cfg.swap_every));
let deliver = |sock: &UdpSocket, bytes: &[u8]| -> io::Result<()> {
for _ in 0..delivery.duplicate {
if let Some(group) = group {
send_datagram(sock, group, bytes)?;
}
let peers = delivery.peers.lock().unwrap();
for addr in peers.values().flatten() {
send_datagram(sock, *addr, bytes)?;
}
}
Ok(())
};
loop {
loop {
match sock.recv_from(&mut punch) {
Ok((n, SocketAddr::V4(from))) if n == FRAME_HEADER_LEN => {
let frame = Frame::decode(&punch);
if frame.kind == KIND_PUNCH && frame.flags == session {
let token = frame.seq as u32;
let mut peers = delivery.peers.lock().unwrap();
if let Some(slot) = peers.get_mut(&token) {
*slot = Some(from);
}
}
}
Ok(_) => {}
Err(e) if e.kind() == io::ErrorKind::WouldBlock => break,
Err(e) if e.kind() == io::ErrorKind::Interrupted => {}
Err(e) => return Err(e),
}
}
wire.clear();
wire.resize(FRAME_HEADER_LEN, 0);
let got = ring.collect_lingering(cursor, batch, max_bytes, &mut wire, linger.current());
if got.count > 0 {
if wire.len() > UDP_MAX_PAYLOAD {
cursor += got.count as u64;
continue;
}
sent += 1;
let frame = data_frame(&mut wire, session, got.count, cursor);
if swap_every != 0 && sent.is_multiple_of(swap_every) && held.is_none() {
held = Some(frame.to_vec());
} else if drop_every == 0 || !sent.is_multiple_of(drop_every) {
match deliver(&sock, frame) {
Ok(()) => {
if let Some(late) = held.take() {
deliver(&sock, &late)?;
}
}
Err(e) if e.raw_os_error() == Some(libc::EMSGSIZE) => {
max_bytes = max_bytes
.min(wire.len().saturating_sub(1))
.max(FRAME_HEADER_LEN + 1);
}
Err(e) => return Err(e),
}
}
linger.sent();
cursor += got.count as u64;
idle = 0;
last_beat = Instant::now();
}
if let Some(seen) = got.lapped {
let resync = ring.resync(seen);
if resync > cursor {
cursor = resync;
}
continue;
}
if got.lost {
cursor += 1;
continue;
}
if got.count == 0 {
if spin {
core::hint::spin_loop();
} else {
idle = idle.saturating_add(1);
if idle < 2000 {
core::hint::spin_loop();
} else if idle < 4000 {
thread::yield_now();
} else {
thread::sleep(Duration::from_micros(50));
}
}
if last_beat.elapsed() >= delivery.heartbeat {
let mut beat = Frame::control(KIND_HEARTBEAT, cursor - 1);
beat.flags = session;
deliver(&sock, &beat.encode())?;
last_beat = Instant::now();
}
}
}
}
fn serve_client(
ring: SourceRing,
mut stream: TcpStream,
batch: usize,
spin: bool,
linger: Option<Duration>,
offer: UdpOffer,
session: u8,
) -> io::Result<()> {
let mut linger = Linger::new(linger);
stream.set_nodelay(true)?;
stream.set_write_timeout(Some(WRITE_TIMEOUT))?;
let mut hdr = [0u8; FRAME_HEADER_LEN];
stream.read_exact(&mut hdr)?;
let hello = Frame::decode(&hdr);
if hello.kind != KIND_HELLO || hello.len as usize != HELLO_LEN {
return Err(protocol("expected HELLO"));
}
let mut body = [0u8; HELLO_LEN];
stream.read_exact(&mut body)?;
let magic = u64::from_le_bytes(body[0..8].try_into().unwrap());
let version = u32::from_le_bytes(body[8..12].try_into().unwrap());
if magic != REPLICATION_MAGIC {
return Err(protocol("bad HELLO magic"));
}
if version != REPLICATION_VERSION {
return Err(protocol("unsupported replication protocol version"));
}
let hello_flags = body[12];
let wants_udp = hello_flags & HELLO_UDP != 0;
let prefers_unicast = hello_flags & HELLO_PREFER_UNICAST != 0;
let unicast = match offer.unicast {
Some(unicast) if wants_udp && (prefers_unicast || offer.multicast.is_none()) => {
Some(unicast)
}
_ => None,
};
let multicast = if unicast.is_none() {
offer.multicast
} else {
None
};
let token = unicast.map(|_| {
let token = offer.next_token.fetch_add(1, Ordering::Relaxed);
offer.peers.lock().unwrap().insert(token, None);
token
});
struct Unregister(Arc<Mutex<HashMap<u32, Option<SocketAddrV4>>>>, Option<u32>);
impl Drop for Unregister {
fn drop(&mut self) {
if let Some(token) = self.1 {
self.0.lock().unwrap().remove(&token);
}
}
}
let _unregister = Unregister(offer.peers.clone(), token);
let capacity = ring.view.capacity;
let write_seq = ring.write_seq();
let oldest = oldest_retained(write_seq, capacity);
let (mut cursor, mut flags) = match hello.seq {
HELLO_LATEST => (write_seq + 1, 0),
HELLO_OLDEST => (oldest, 0),
wanted if wanted > write_seq + 1 => (write_seq + 1, GEOMETRY_RESET),
wanted => (wanted.max(oldest), 0),
};
if multicast.is_some() || unicast.is_some() {
flags |= GEOMETRY_MULTICAST;
}
let mut geometry = [0u8; GEOMETRY_LEN];
ring.geometry.encode(&mut geometry);
let mut out =
Vec::with_capacity(FRAME_HEADER_LEN + GEOMETRY_LEN + FRAME_HEADER_LEN + MULTICAST_LEN);
out.extend_from_slice(
&Frame {
kind: KIND_GEOMETRY,
flags,
count: 0,
len: GEOMETRY_LEN as u32,
seq: cursor,
}
.encode(),
);
out.extend_from_slice(&geometry);
let udp_info = match (multicast, unicast) {
(Some(cfg), _) => {
let mut info = [0u8; MULTICAST_LEN];
cfg.encode(session, &mut info);
Some(info)
}
(None, Some((port, mtu))) => {
let mut info = [0u8; MULTICAST_LEN];
encode_udp_info(
Ipv4Addr::UNSPECIFIED,
port,
mtu,
0,
session,
token.unwrap_or(0),
&mut info,
);
Some(info)
}
(None, None) => None,
};
if let Some(info) = udp_info {
out.extend_from_slice(
&Frame {
kind: KIND_MULTICAST,
flags: 0,
count: 0,
len: MULTICAST_LEN as u32,
seq: 0,
}
.encode(),
);
out.extend_from_slice(&info);
}
stream.write_all(&out)?;
let batch = batch.clamp(1, u16::MAX as usize);
let mut wire: Vec<u8> = Vec::with_capacity(1 << 16);
if multicast.is_some() || unicast.is_some() {
loop {
stream.read_exact(&mut hdr)?;
let frame = Frame::decode(&hdr);
match frame.kind {
KIND_NAK if frame.len as usize == NAK_LEN => {
let mut to = [0u8; NAK_LEN];
stream.read_exact(&mut to)?;
let to = u64::from_le_bytes(to);
serve_range(&ring, &mut stream, frame.seq, to, batch, &mut wire)?;
}
_ => return Err(protocol("expected NAK")),
}
}
}
let mut idle = 0u32;
let mut last_beat = Instant::now();
loop {
wire.clear();
wire.resize(FRAME_HEADER_LEN, 0);
let got =
ring.collect_lingering(cursor, batch, TCP_FRAME_BYTES, &mut wire, linger.current());
if got.count > 0 {
stream.write_all(data_frame(&mut wire, 0, got.count, cursor))?;
linger.sent();
cursor += got.count as u64;
idle = 0;
last_beat = Instant::now();
}
if let Some(seen) = got.lapped {
let resync = ring.resync(seen);
if resync > cursor {
stream.write_all(&Frame::control(KIND_GAP, resync).encode())?;
cursor = resync;
}
continue;
}
if got.lost {
cursor += 1;
stream.write_all(&Frame::control(KIND_GAP, cursor).encode())?;
continue;
}
if got.count == 0 {
if spin {
core::hint::spin_loop();
} else {
idle = idle.saturating_add(1);
if idle < 2000 {
core::hint::spin_loop();
} else if idle < 4000 {
thread::yield_now();
} else {
thread::sleep(Duration::from_micros(50));
}
}
if last_beat.elapsed() >= HEARTBEAT_INTERVAL {
stream.write_all(&Frame::control(KIND_HEARTBEAT, ring.write_seq()).encode())?;
last_beat = Instant::now();
}
}
}
}
fn serve_range(
ring: &SourceRing,
stream: &mut TcpStream,
from: u64,
to: u64,
batch: usize,
wire: &mut Vec<u8>,
) -> io::Result<()> {
let mut cursor = from.max(1);
let oldest = oldest_retained(ring.write_seq(), ring.view.capacity);
if cursor < oldest {
stream.write_all(&Frame::control(KIND_GAP, oldest).encode())?;
cursor = oldest;
}
while cursor <= to {
let max = batch.min((to - cursor + 1) as usize);
wire.clear();
wire.resize(FRAME_HEADER_LEN, 0);
let got = ring.collect(cursor, max, TCP_FRAME_BYTES, wire);
if got.count > 0 {
stream.write_all(data_frame(wire, 0, got.count, cursor))?;
cursor += got.count as u64;
}
if let Some(seen) = got.lapped {
let resync = ring.resync(seen);
if resync > cursor {
stream.write_all(&Frame::control(KIND_GAP, resync).encode())?;
cursor = resync;
}
continue;
}
if got.lost {
cursor += 1;
stream.write_all(&Frame::control(KIND_GAP, cursor).encode())?;
continue;
}
if got.count == 0 {
break;
}
}
Ok(())
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum MirrorStart {
#[default]
Latest,
Oldest,
Resume,
Sequence(u64),
}
#[derive(Debug, Clone, Copy)]
pub struct MirrorBuilder {
start: MirrorStart,
spin: bool,
file_mode: u32,
interface: Ipv4Addr,
rcvbuf: usize,
nak_timeout: Duration,
prefer_unicast: bool,
}
impl Default for MirrorBuilder {
fn default() -> Self {
Self::new()
}
}
impl MirrorBuilder {
pub fn new() -> Self {
Self {
start: MirrorStart::Latest,
spin: false,
file_mode: 0o660,
interface: Ipv4Addr::UNSPECIFIED,
rcvbuf: 8 << 20,
nak_timeout: Duration::from_millis(20),
prefer_unicast: false,
}
}
pub fn unicast(mut self, prefer: bool) -> Self {
self.prefer_unicast = prefer;
self
}
pub fn start(mut self, start: MirrorStart) -> Self {
self.start = start;
self
}
pub fn spin(mut self, spin: bool) -> Self {
self.spin = spin;
self
}
pub fn file_mode(mut self, mode: u32) -> Self {
self.file_mode = mode;
self
}
pub fn interface(mut self, interface: Ipv4Addr) -> Self {
self.interface = interface;
self
}
pub fn rcvbuf(mut self, bytes: usize) -> Self {
self.rcvbuf = bytes;
self
}
pub fn nak_timeout(mut self, timeout: Duration) -> Self {
self.nak_timeout = timeout;
self
}
pub fn connect<A: ToSocketAddrs, P: AsRef<Path>>(
self,
source: A,
ring_path: P,
) -> Result<Mirror> {
let ring_path = ring_path.as_ref().to_path_buf();
let adopted = match self.start {
MirrorStart::Resume => MirrorRing::adopt(&ring_path)?,
_ => None,
};
let wanted = match self.start {
MirrorStart::Latest => HELLO_LATEST,
MirrorStart::Oldest => HELLO_OLDEST,
MirrorStart::Sequence(seq) => seq.max(1),
MirrorStart::Resume => adopted.as_ref().map_or(HELLO_OLDEST, |ring| ring.next_seq),
};
let mut stream = TcpStream::connect(source)?;
stream.set_nodelay(true)?;
let mut hello = Vec::with_capacity(FRAME_HEADER_LEN + HELLO_LEN);
hello.extend_from_slice(
&Frame {
kind: KIND_HELLO,
flags: 0,
count: 0,
len: HELLO_LEN as u32,
seq: wanted,
}
.encode(),
);
hello.extend_from_slice(&REPLICATION_MAGIC.to_le_bytes());
hello.extend_from_slice(&REPLICATION_VERSION.to_le_bytes());
hello.push(
HELLO_UDP
| if self.prefer_unicast {
HELLO_PREFER_UNICAST
} else {
0
},
);
hello.extend_from_slice(&[0u8; 3]);
stream.write_all(&hello)?;
let mut hdr = [0u8; FRAME_HEADER_LEN];
stream.read_exact(&mut hdr)?;
let frame = Frame::decode(&hdr);
if frame.kind != KIND_GEOMETRY || frame.len as usize != GEOMETRY_LEN {
return Err(RingfireError::Protocol("expected GEOMETRY"));
}
let mut body = [0u8; GEOMETRY_LEN];
stream.read_exact(&mut body)?;
let geometry = Geometry::decode(&body)?;
let reset = frame.flags & GEOMETRY_RESET != 0;
let mut udp = None;
let mut session = 0u8;
let mut punch = None;
if frame.flags & GEOMETRY_MULTICAST != 0 {
stream.read_exact(&mut hdr)?;
let mc = Frame::decode(&hdr);
if mc.kind != KIND_MULTICAST || mc.len as usize != MULTICAST_LEN {
return Err(RingfireError::Protocol("expected MULTICAST"));
}
let mut info = [0u8; MULTICAST_LEN];
stream.read_exact(&mut info)?;
let info = MulticastInfo::decode(&info)?;
session = info.session;
if info.group.is_unspecified() {
let SocketAddr::V4(peer) = stream.peer_addr()? else {
return Err(RingfireError::Unsupported(
"unicast delivery needs an IPv4 source",
));
};
udp = Some(multicast_receiver(
info.group,
0,
self.interface,
self.rcvbuf,
)?);
punch = Some(Punch {
to: SocketAddrV4::new(*peer.ip(), info.port),
token: info.token,
last: Instant::now() - PUNCH_INTERVAL,
heard: false,
});
} else {
udp = Some(multicast_receiver(
info.group,
info.port,
self.interface,
self.rcvbuf,
)?);
}
}
let (ring, resumed) = match adopted {
Some(ring) if !reset && ring.geometry.same_layout(&geometry) => (ring, true),
other => {
drop(other);
let mut ring = MirrorRing::create(&ring_path, &geometry, self.file_mode)?;
ring.skip_to(frame.seq);
(ring, false)
}
};
if self.spin || udp.is_some() {
stream.set_nonblocking(true)?;
}
Ok(Mirror {
stream,
udp,
punch,
session,
ring: Some(ring),
ring_path,
geometry,
spin: self.spin,
file_mode: self.file_mode,
nak_timeout: self.nak_timeout,
first_seq: frame.seq,
resumed,
source_seq: frame.seq.saturating_sub(1),
frames: 0,
gaps: 0,
datagrams: 0,
naks: 0,
retransmitted: 0,
nak: None,
last_nak: None,
pending: BTreeMap::new(),
buf: Vec::new(),
dgram: vec![0u8; MAX_DATAGRAM],
})
}
}
#[derive(Debug, Clone, Copy)]
struct Nak {
from: u64,
to: u64,
sent: Instant,
}
#[derive(Debug, Clone, Copy)]
struct Punch {
to: SocketAddrV4,
token: u32,
last: Instant,
heard: bool,
}
pub struct Mirror {
stream: TcpStream,
udp: Option<UdpSocket>,
punch: Option<Punch>,
session: u8,
ring: Option<MirrorRing>,
ring_path: PathBuf,
geometry: Geometry,
spin: bool,
file_mode: u32,
nak_timeout: Duration,
first_seq: u64,
resumed: bool,
source_seq: u64,
frames: u64,
gaps: u64,
datagrams: u64,
naks: u64,
retransmitted: u64,
nak: Option<Nak>,
last_nak: Option<(u64, u64)>,
pending: BTreeMap<u64, (usize, Vec<u8>)>,
buf: Vec<u8>,
dgram: Vec<u8>,
}
#[derive(Debug)]
pub struct MirrorHandle {
stream: TcpStream,
}
impl MirrorHandle {
pub fn shutdown(&self) -> io::Result<()> {
self.stream.shutdown(Shutdown::Both)
}
}
impl Mirror {
pub fn builder() -> MirrorBuilder {
MirrorBuilder::new()
}
pub fn connect<A: ToSocketAddrs, P: AsRef<Path>>(source: A, ring_path: P) -> Result<Self> {
MirrorBuilder::new().connect(source, ring_path)
}
pub fn handle(&self) -> io::Result<MirrorHandle> {
Ok(MirrorHandle {
stream: self.stream.try_clone()?,
})
}
pub fn geometry(&self) -> Geometry {
self.geometry
}
pub fn path(&self) -> &Path {
&self.ring_path
}
pub fn peer_addr(&self) -> io::Result<SocketAddr> {
self.stream.peer_addr()
}
pub fn is_multicast(&self) -> bool {
self.udp.is_some()
}
pub fn is_unicast(&self) -> bool {
self.punch.is_some()
}
#[doc(hidden)]
pub fn session(&self) -> u8 {
self.session
}
#[doc(hidden)]
pub fn last_nak(&self) -> Option<(u64, u64)> {
self.last_nak
}
pub fn first_sequence(&self) -> u64 {
self.first_seq
}
pub fn resumed(&self) -> bool {
self.resumed
}
pub fn sequence(&self) -> u64 {
self.ring().next_seq - 1
}
pub fn source_sequence(&self) -> u64 {
self.source_seq
}
pub fn gaps(&self) -> u64 {
self.gaps
}
pub fn frames(&self) -> u64 {
self.frames
}
pub fn datagrams(&self) -> u64 {
self.datagrams
}
pub fn naks(&self) -> u64 {
self.naks
}
pub fn retransmitted(&self) -> u64 {
self.retransmitted
}
fn ring(&self) -> &MirrorRing {
self.ring.as_ref().expect("mirror ring")
}
fn ring_mut(&mut self) -> &mut MirrorRing {
self.ring.as_mut().expect("mirror ring")
}
pub fn run(&mut self) -> Result<()> {
while self.step()? {}
Ok(())
}
pub fn step(&mut self) -> Result<bool> {
if self.udp.is_some() {
return self.step_multicast();
}
let mut hdr = [0u8; FRAME_HEADER_LEN];
match read_full(&mut self.stream, &mut hdr, self.spin) {
Ok(()) => {}
Err(e) if peer_gone(&e) => return Ok(false),
Err(e) => return Err(e.into()),
}
let frame = Frame::decode(&hdr);
self.frames += 1;
match frame.kind {
KIND_DATA => {
let (count, len) = match self.read_data_payload(frame) {
Ok(sizes) => sizes,
Err(RingfireError::Io(e)) if peer_gone(&e) => return Ok(false),
Err(e) => return Err(e),
};
let next = self.ring().next_seq;
if frame.seq < next {
self.ring = None;
let mut ring =
MirrorRing::create(&self.ring_path, &self.geometry, self.file_mode)?;
ring.skip_to(frame.seq);
self.ring = Some(ring);
self.gaps += 1;
} else if frame.seq > next {
self.gaps += 1;
self.ring_mut().skip_to(frame.seq);
}
let bytes = std::mem::take(&mut self.buf);
let result = self.write_records(frame.seq, count, &bytes[..len]);
self.buf = bytes;
result?;
}
KIND_GAP => {
if frame.seq > self.ring().next_seq {
self.gaps += 1;
self.ring_mut().skip_to(frame.seq);
}
}
KIND_HEARTBEAT => {
self.source_seq = self.source_seq.max(frame.seq);
}
_ => return Err(RingfireError::Protocol("unexpected frame kind")),
}
Ok(true)
}
fn read_data_payload(&mut self, frame: Frame) -> Result<(usize, usize)> {
let count = frame.count as usize;
let len = frame.len as usize;
if !self.frame_sizes_ok(count, len) {
return Err(RingfireError::Protocol(
"DATA length does not match its record count",
));
}
if self.buf.len() < len {
self.buf.resize(len, 0);
}
read_full(&mut self.stream, &mut self.buf[..len], self.spin)?;
Ok((count, len))
}
fn frame_sizes_ok(&self, count: usize, len: usize) -> bool {
let descriptors = count * self.geometry.payload_len();
count > 0
&& if self.geometry.has_arena() {
len >= descriptors
} else {
len == descriptors
}
}
fn write_records(&mut self, seq: u64, count: usize, bytes: &[u8]) -> Result<()> {
let payload_len = self.geometry.payload_len();
let arena = self.geometry.has_arena();
let ring = self.ring.as_mut().expect("mirror ring");
let next = ring.next_seq;
let mut pos = 0usize;
for i in 0..count {
let s = seq + i as u64;
if pos + payload_len > bytes.len() {
return Err(RingfireError::Protocol(
"DATA frame shorter than its records",
));
}
let descriptor = &bytes[pos..pos + payload_len];
pos += payload_len;
if arena {
let len = blob_ref_at(descriptor).len as usize;
if pos + len > bytes.len() {
return Err(RingfireError::Protocol("DATA frame shorter than its blobs"));
}
let blob = &bytes[pos..pos + len];
pos += len;
if s >= next {
ring.write_blob(s, descriptor, blob)?;
}
} else if s >= next {
ring.write(s, descriptor);
}
}
if pos != bytes.len() {
return Err(RingfireError::Protocol(
"DATA frame longer than its records",
));
}
ring.publish();
self.source_seq = self.source_seq.max(seq + count as u64 - 1);
Ok(())
}
fn step_multicast(&mut self) -> Result<bool> {
let mut progressed = false;
let mut drained = 0;
while drained < DATAGRAMS_PER_STEP {
drained += 1;
let received = {
let udp = self.udp.as_ref().expect("multicast socket");
udp.recv(&mut self.dgram)
};
match received {
Ok(n) => {
progressed = true;
self.datagrams += 1;
let dgram = std::mem::take(&mut self.dgram);
let result = self.handle_datagram(&dgram[..n]);
self.dgram = dgram;
match result {
Ok(()) => {}
Err(RingfireError::Io(e)) if peer_gone(&e) => return Ok(false),
Err(e) => return Err(e),
}
}
Err(e) if e.kind() == io::ErrorKind::WouldBlock => break,
Err(e) if e.kind() == io::ErrorKind::Interrupted => {}
Err(e) => return Err(e.into()),
}
}
match self.try_read_frame() {
Ok(Some(frame)) => {
progressed = true;
self.frames += 1;
match self.handle_control(frame) {
Ok(()) => {}
Err(RingfireError::Io(e)) if peer_gone(&e) => return Ok(false),
Err(e) => return Err(e),
}
}
Ok(None) => {}
Err(e) if peer_gone(&e) => return Ok(false),
Err(e) => return Err(e.into()),
}
if let Some(nak) = self.nak
&& nak.sent.elapsed() >= self.nak_timeout
{
match self.send_nak(nak.from, nak.to) {
Ok(()) => {}
Err(RingfireError::Io(e)) if peer_gone(&e) => return Ok(false),
Err(e) => return Err(e),
}
progressed = true;
}
if let Some(punch) = self.punch.as_mut() {
let due = if punch.heard {
PUNCH_KEEPALIVE
} else {
PUNCH_INTERVAL
};
if punch.last.elapsed() >= due {
let mut frame = Frame::control(KIND_PUNCH, u64::from(punch.token));
frame.flags = self.session;
let udp = self.udp.as_ref().expect("unicast socket");
match udp.send_to(&frame.encode(), punch.to) {
Ok(_) => {}
Err(e) if e.kind() == io::ErrorKind::WouldBlock => {}
Err(e) => return Err(e.into()),
}
punch.last = Instant::now();
progressed = true;
}
}
if !progressed {
if self.spin {
core::hint::spin_loop();
} else {
let udp_fd = self.udp.as_ref().expect("multicast socket").as_raw_fd();
wait_readable(
&[udp_fd, self.stream.as_raw_fd()],
self.nak_timeout.min(POLL_SLICE),
);
}
}
Ok(true)
}
fn try_read_frame(&mut self) -> io::Result<Option<Frame>> {
let mut hdr = [0u8; FRAME_HEADER_LEN];
let n = match self.stream.read(&mut hdr) {
Ok(0) => return Err(io::Error::from(io::ErrorKind::UnexpectedEof)),
Ok(n) => n,
Err(e)
if e.kind() == io::ErrorKind::WouldBlock
|| e.kind() == io::ErrorKind::Interrupted =>
{
return Ok(None);
}
Err(e) => return Err(e),
};
if n < FRAME_HEADER_LEN {
read_full(&mut self.stream, &mut hdr[n..], self.spin)?;
}
Ok(Some(Frame::decode(&hdr)))
}
fn handle_control(&mut self, frame: Frame) -> Result<()> {
match frame.kind {
KIND_DATA => {
let (count, len) = self.read_data_payload(frame)?;
self.retransmitted += count as u64;
let payload = std::mem::take(&mut self.buf);
let result = self.apply(frame.seq, count, &payload[..len]);
self.buf = payload;
result
}
KIND_GAP => {
if frame.seq > self.ring().next_seq {
self.gaps += 1;
self.ring_mut().skip_to(frame.seq);
self.after_advance()
} else {
Ok(())
}
}
KIND_HEARTBEAT => self.on_source_seq(frame.seq),
KIND_MULTICAST => Ok(()),
_ => Err(RingfireError::Protocol("unexpected frame kind")),
}
}
fn handle_datagram(&mut self, bytes: &[u8]) -> Result<()> {
if bytes.len() < FRAME_HEADER_LEN {
return Ok(());
}
let frame = Frame::decode(bytes[..FRAME_HEADER_LEN].try_into().unwrap());
if frame.flags != self.session {
return Ok(());
}
if let Some(punch) = self.punch.as_mut() {
punch.heard = true;
}
match frame.kind {
KIND_DATA => {
let count = frame.count as usize;
let len = frame.len as usize;
if !self.frame_sizes_ok(count, len) || bytes.len() < FRAME_HEADER_LEN + len {
return Ok(());
}
self.apply(
frame.seq,
count,
&bytes[FRAME_HEADER_LEN..FRAME_HEADER_LEN + len],
)
}
KIND_HEARTBEAT => self.on_source_seq(frame.seq),
_ => Ok(()),
}
}
fn apply(&mut self, seq: u64, count: usize, payload: &[u8]) -> Result<()> {
let next = self.ring().next_seq;
let end = seq + count as u64;
if end <= next {
return Ok(());
}
if seq > next {
if self.pending.len() < PENDING_MAX {
self.pending
.entry(seq)
.or_insert_with(|| (count, payload.to_vec()));
}
return self.request(next, seq - 1);
}
self.write_records(seq, count, payload)?;
self.after_advance()
}
fn after_advance(&mut self) -> Result<()> {
loop {
let next = self.ring().next_seq;
if let Some(nak) = self.nak
&& next > nak.to
{
self.nak = None;
}
let Some((&seq, _)) = self.pending.first_key_value() else {
return Ok(());
};
let (count, payload) = self.pending.remove(&seq).expect("first pending entry");
let end = seq + count as u64;
if end <= next {
continue;
}
if seq > next {
self.pending.insert(seq, (count, payload));
return self.request(next, seq - 1);
}
self.write_records(seq, count, &payload)?;
}
}
fn on_source_seq(&mut self, seq: u64) -> Result<()> {
self.source_seq = self.source_seq.max(seq);
let next = self.ring().next_seq;
if seq >= next {
self.request(next, seq)?;
}
Ok(())
}
fn request(&mut self, from: u64, to: u64) -> Result<()> {
let to = to.min(from + NAK_MAX - 1);
if let Some(nak) = self.nak
&& nak.from == from
&& nak.to >= to
&& nak.sent.elapsed() < self.nak_timeout
{
return Ok(());
}
self.send_nak(from, to)
}
fn send_nak(&mut self, from: u64, to: u64) -> Result<()> {
let mut frame = Vec::with_capacity(FRAME_HEADER_LEN + NAK_LEN);
frame.extend_from_slice(
&Frame {
kind: KIND_NAK,
flags: 0,
count: 0,
len: NAK_LEN as u32,
seq: from,
}
.encode(),
);
frame.extend_from_slice(&to.to_le_bytes());
write_full(&mut self.stream, &frame)?;
self.nak = Some(Nak {
from,
to,
sent: Instant::now(),
});
self.last_nak = Some((from, to));
self.naks += 1;
Ok(())
}
}
impl std::fmt::Debug for Mirror {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Mirror")
.field("path", &self.ring_path)
.field("multicast", &self.udp.is_some())
.field("sequence", &self.sequence())
.field("source_sequence", &self.source_seq)
.field("gaps", &self.gaps)
.field("naks", &self.naks)
.finish()
}
}
#[cfg(test)]
#[path = "../tests/unit/replication.rs"]
mod tests;