#![deny(missing_docs)]
#![deny(rustdoc::broken_intra_doc_links)]
pub mod region;
mod ring;
pub mod slot;
use std::marker::PhantomData;
use std::sync::Arc;
use crate::core::{
Encode, Error, Lifecycle, Message, MetricKey, MetricKind, MetricsCollector, NullCollector,
Result, SchemaId, Sink,
};
use region::Region;
use slot::{FLAG_SUSPECT, FLAG_VALID, SlotHeader, slot_size};
#[derive(Debug, Clone, Copy)]
pub enum MemoryMetricKey {
MessagesWritten,
RingFull,
Oversized,
}
impl MetricKey for MemoryMetricKey {
fn name(&self) -> &str {
match self {
Self::MessagesWritten => "memory.messages_written",
Self::RingFull => "memory.ring_full",
Self::Oversized => "memory.oversized",
}
}
}
pub const DEFAULT_SLOT_COUNT: usize = 1024;
pub const DEFAULT_MAX_PAYLOAD: usize = 96;
const STACK_ENCODE_CAP: usize = 256;
pub struct SharedMemorySink<M: Message + Encode> {
region: Region,
max_payload: usize,
written: u64,
metrics: Arc<dyn MetricsCollector>,
_marker: PhantomData<M>,
}
impl<M: Message + Encode> SharedMemorySink<M> {
pub fn new(slot_count: usize, max_payload: usize) -> Result<Self> {
if max_payload > u32::MAX as usize {
return Err(Error::config("max_payload exceeds u32::MAX"));
}
let ss = slot_size(max_payload)?;
let region = Region::anonymous(slot_count, ss)?;
Ok(Self {
region,
max_payload,
written: 0,
metrics: Arc::new(NullCollector),
_marker: PhantomData,
})
}
pub fn with_metrics(mut self, metrics: Arc<dyn MetricsCollector>) -> Self {
self.metrics = metrics;
self
}
pub fn written(&self) -> u64 {
self.written
}
pub fn len(&self) -> usize {
self.region.len()
}
pub fn is_empty(&self) -> bool {
self.region.is_empty()
}
pub fn is_full(&self) -> bool {
self.region.is_full()
}
pub fn pop(&mut self, consume: impl FnOnce(&[u8])) -> bool {
self.region.pop(consume)
}
fn reset_ring(&mut self) {
while self.region.pop(|_| {}) {}
self.written = 0;
}
}
impl<M: Message + Encode> Lifecycle for SharedMemorySink<M> {
fn init(&mut self) -> Result<()> {
self.reset_ring();
Ok(())
}
fn shutdown(&mut self) -> Result<()> {
Ok(())
}
}
impl<M: Message + Encode> Sink for SharedMemorySink<M> {
type Message = M;
fn write(&mut self, message: &M) -> Result<()> {
let payload_len = message.encoded_len();
if payload_len > self.max_payload {
self.metrics
.record(&MemoryMetricKey::Oversized, MetricKind::Counter, 1.0);
return Err(Error::sink(format!(
"encoded message ({} bytes) exceeds max_payload ({} bytes)",
payload_len, self.max_payload
)));
}
let payload_len_u32 = u32::try_from(payload_len)
.map_err(|_| Error::sink("encoded message length exceeds u32::MAX"))?;
let meta = message.metadata();
let ts = message.timestamp();
let schema_id = message.schema_id().id();
let flags = FLAG_VALID | if meta.suspect { FLAG_SUSPECT } else { 0 };
let header = SlotHeader::new(
schema_id,
flags,
meta.sequence,
ts.as_nanos(),
payload_len_u32,
);
let pushed = if payload_len <= STACK_ENCODE_CAP {
let mut stack = [0u8; STACK_ENCODE_CAP];
let n = message.encode_into(&mut stack[..payload_len])?;
if n != payload_len {
return Err(Error::encode(
"encode_into length does not match encoded_len",
));
}
self.region
.push(|buf| slot::encode(&header, &stack[..n], buf))?
} else {
let mut heap = vec![0u8; payload_len];
let n = message.encode_into(&mut heap)?;
if n != payload_len {
return Err(Error::encode(
"encode_into length does not match encoded_len",
));
}
self.region
.push(|buf| slot::encode(&header, &heap[..n], buf))?
};
if pushed {
self.written += 1;
self.metrics
.record_counter(&MemoryMetricKey::MessagesWritten, 1);
Ok(())
} else {
self.metrics.record_counter(&MemoryMetricKey::RingFull, 1);
Err(Error::back_pressure("ring buffer full"))
}
}
}
#[derive(Debug, Default, Clone, Copy)]
pub struct StubMessage {
pub seq: u64,
}
impl Message for StubMessage {
type Schema = crate::core::DefaultSchemaId;
fn schema_id(&self) -> crate::core::DefaultSchemaId {
crate::core::DefaultSchemaId(1)
}
fn timestamp(&self) -> crate::core::Timestamp {
crate::core::Timestamp::from_nanos(0)
}
fn metadata(&self) -> crate::core::Metadata {
crate::core::Metadata {
sequence: self.seq,
suspect: false,
}
}
}
impl Encode for StubMessage {
fn encoded_len(&self) -> usize {
8
}
fn encode_into(&self, dst: &mut [u8]) -> Result<usize> {
if dst.len() < 8 {
return Err(Error::encode("buffer too small for StubMessage"));
}
dst[..8].copy_from_slice(&self.seq.to_be_bytes());
Ok(8)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::{ErrorKind, Sink};
fn make_sink() -> SharedMemorySink<StubMessage> {
SharedMemorySink::new(16, 64).unwrap()
}
fn msg(seq: u64) -> StubMessage {
StubMessage { seq }
}
#[test]
fn write_and_count() {
let mut sink = make_sink();
sink.write(&msg(1)).unwrap();
sink.write(&msg(2)).unwrap();
assert_eq!(sink.written(), 2);
assert_eq!(sink.len(), 2);
}
#[test]
fn write_pop_roundtrip() {
let mut sink = make_sink();
sink.write(&msg(42)).unwrap();
let mut recovered = 0u64;
sink.pop(|buf| {
let (_hdr, payload) = slot::decode(buf).unwrap();
recovered = u64::from_be_bytes(payload.try_into().unwrap());
});
assert_eq!(recovered, 42);
}
#[test]
fn full_ring_returns_back_pressure() {
let mut sink = SharedMemorySink::new(4, 64).unwrap();
for i in 0..4 {
sink.write(&msg(i)).unwrap();
}
let err = sink.write(&msg(99)).unwrap_err();
assert_eq!(err.kind(), ErrorKind::BackPressure);
}
#[test]
fn oversized_payload_returns_error() {
let mut sink = SharedMemorySink::new(4, 4).unwrap();
assert!(sink.write(&msg(0)).is_err());
}
#[test]
fn sequence_monotonic_across_writes() {
let mut sink = make_sink();
for i in 0..8u64 {
sink.write(&msg(i)).unwrap();
}
let mut seqs = Vec::new();
while sink.pop(|buf| {
let (hdr, _) = slot::decode(buf).unwrap();
seqs.push(hdr.sequence);
}) {}
let expected: Vec<u64> = (0..8).collect();
assert_eq!(seqs, expected);
}
#[test]
fn reinit_clears_ring() {
let mut sink = make_sink();
sink.init().unwrap();
sink.write(&msg(1)).unwrap();
assert_eq!(sink.len(), 1);
sink.shutdown().unwrap();
sink.init().unwrap();
assert_eq!(sink.written(), 0);
assert!(sink.is_empty());
}
#[test]
fn failed_encode_does_not_advance_ring() {
struct BadMsg;
impl Message for BadMsg {
type Schema = crate::core::DefaultSchemaId;
fn schema_id(&self) -> crate::core::DefaultSchemaId {
crate::core::DefaultSchemaId(1)
}
fn timestamp(&self) -> crate::core::Timestamp {
crate::core::Timestamp::from_nanos(0)
}
fn metadata(&self) -> crate::core::Metadata {
crate::core::Metadata::default()
}
}
impl Encode for BadMsg {
fn encoded_len(&self) -> usize {
8
}
fn encode_into(&self, _dst: &mut [u8]) -> Result<usize> {
Err(Error::encode("boom"))
}
}
let mut sink: SharedMemorySink<BadMsg> = SharedMemorySink::new(4, 64).unwrap();
assert!(sink.write(&BadMsg).is_err());
assert!(sink.is_empty());
}
}