use std::ops::Range;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::{Mutex, OnceLock};
use std::task::{Poll, ready};
use arrayvec::ArrayVec;
use bytes::Bytes;
use crate::group::{self, GroupState};
use crate::{Error, IntoBytes, Result, Timestamp, stats};
#[derive(Clone, Copy, Debug)]
pub struct Info {
pub size: u64,
pub timestamp: Timestamp,
}
#[derive(Clone, Debug)]
pub struct Frame {
pub timestamp: Timestamp,
pub payload: Bytes,
}
#[derive(Debug, Default)]
pub struct Buffer<const N: usize = 8>(ArrayVec<Frame, N>);
impl<const N: usize> Buffer<N> {
pub fn new() -> Self {
Self(ArrayVec::new())
}
pub const fn capacity(&self) -> usize {
N
}
pub fn len(&self) -> usize {
self.0.len()
}
pub fn is_empty(&self) -> bool {
self.0.is_empty()
}
pub fn is_full(&self) -> bool {
self.0.is_full()
}
pub fn filled(&self) -> &[Frame] {
&self.0
}
pub fn filled_mut(&mut self) -> &mut [Frame] {
&mut self.0
}
pub fn push(&mut self, frame: Frame) -> std::result::Result<(), Frame> {
self.0.try_push(frame).map_err(|err| err.element())
}
pub fn drain(&mut self) -> impl ExactSizeIterator<Item = Frame> + '_ {
self.0.drain(..)
}
pub fn clear(&mut self) {
self.0.clear();
}
}
#[derive(Clone)]
pub(crate) struct Budget(Arc<AtomicUsize>);
impl Budget {
const DEFAULT: usize = 16 * 1024 * 1024;
pub(crate) fn new(bytes: usize) -> Self {
Self(Arc::new(AtomicUsize::new(bytes)))
}
pub(crate) fn reserve(&self, size: usize) -> Option<Reservation> {
self.0
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |left| left.checked_sub(size))
.ok()?;
Some(Reservation {
budget: self.clone(),
size,
})
}
}
impl Default for Budget {
fn default() -> Self {
Self::new(Self::DEFAULT)
}
}
pub(crate) struct Reservation {
budget: Budget,
size: usize,
}
impl Drop for Reservation {
fn drop(&mut self) {
self.budget.0.fetch_add(self.size, Ordering::Relaxed);
}
}
#[derive(Clone)]
pub(crate) struct FrameBuf(Arc<FrameBufInner>);
struct FrameBufInner {
size: usize,
grow: bool,
written: AtomicUsize,
storage: OnceLock<FrameStorage>,
}
enum FrameStorage {
Shared(Bytes),
Fixed(Arc<Segment>),
Growing(Mutex<Arc<Segment>>),
}
struct Segment {
data: *mut u8,
capacity: usize,
}
unsafe impl Send for Segment {}
unsafe impl Sync for Segment {}
impl Drop for Segment {
fn drop(&mut self) {
unsafe {
let slice = std::ptr::slice_from_raw_parts_mut(self.data, self.capacity);
drop(Box::from_raw(slice));
}
}
}
impl Segment {
fn new(capacity: usize) -> Self {
let boxed: Box<[u8]> = vec![0u8; capacity].into_boxed_slice();
let capacity = boxed.len();
let data = Box::into_raw(boxed) as *mut u8;
Self { data, capacity }
}
unsafe fn write(&self, offset: usize, src: &[u8]) {
debug_assert!(offset + src.len() <= self.capacity);
unsafe { std::ptr::copy_nonoverlapping(src.as_ptr(), self.data.add(offset), src.len()) };
}
}
struct Filled {
segment: Arc<Segment>,
len: usize,
}
impl AsRef<[u8]> for Filled {
fn as_ref(&self) -> &[u8] {
unsafe { std::slice::from_raw_parts(self.segment.data, self.len) }
}
}
impl FrameBuf {
pub(crate) fn new(size: usize) -> Self {
Self::with_growth(size, false)
}
pub(crate) fn growing(size: usize) -> Self {
Self::with_growth(size, true)
}
fn with_growth(size: usize, grow: bool) -> Self {
Self(Arc::new(FrameBufInner {
size,
grow,
written: AtomicUsize::new(0),
storage: OnceLock::new(),
}))
}
pub(crate) fn size(&self) -> usize {
self.0.size
}
pub(crate) fn written(&self, ord: Ordering) -> usize {
self.0.written.load(ord)
}
fn try_set_bytes(&self, bytes: Bytes) -> std::result::Result<(), Bytes> {
if bytes.len() != self.size() || self.written(Ordering::Acquire) != 0 {
return Err(bytes);
}
self.0
.storage
.set(FrameStorage::Shared(bytes))
.map_err(|storage| match storage {
FrameStorage::Shared(bytes) => bytes,
_ => unreachable!("try_set_bytes only installs shared storage"),
})
}
unsafe fn store_written(&self, new_written: usize) {
self.0.written.store(new_written, Ordering::Release);
}
fn append(&self, src: &[u8]) {
if src.is_empty() {
return;
}
let prev = self.written(Ordering::Relaxed);
let storage = self.0.storage.get_or_init(|| match self.0.grow {
true => FrameStorage::Growing(Mutex::new(Arc::new(Segment::new(0)))),
false => FrameStorage::Fixed(Arc::new(Segment::new(self.size()))),
});
match storage {
FrameStorage::Shared(_) => return,
FrameStorage::Fixed(segment) => unsafe { segment.write(prev, src) },
FrameStorage::Growing(current) => {
let mut current = current.lock().expect("mutex poisoned");
let needed = prev + src.len();
if current.capacity < needed {
let next = Segment::new(needed.max(current.capacity * 2).min(self.size()));
let filled = Filled {
segment: current.clone(),
len: prev,
};
unsafe { next.write(0, filled.as_ref()) };
*current = Arc::new(next);
}
unsafe { current.write(prev, src) };
}
}
unsafe { self.store_written(prev + src.len()) };
}
fn freeze(&self) -> Bytes {
self.slice(0, self.size())
}
fn slice(&self, start: usize, end: usize) -> Bytes {
let segment = match self.0.storage.get() {
Some(FrameStorage::Shared(bytes)) => return bytes.slice(start..end),
Some(FrameStorage::Fixed(segment)) => segment.clone(),
Some(FrameStorage::Growing(current)) => current.lock().expect("mutex poisoned").clone(),
None => return Bytes::new(),
};
debug_assert!(end <= segment.capacity);
Bytes::from_owner(Filled { segment, len: end }).slice(start..)
}
#[cfg(test)]
pub(crate) fn allocated(&self) -> usize {
match self.0.storage.get() {
Some(FrameStorage::Shared(bytes)) => bytes.len(),
Some(FrameStorage::Fixed(segment)) => segment.capacity,
Some(FrameStorage::Growing(current)) => current.lock().expect("mutex poisoned").capacity,
None => 0,
}
}
}
struct Raw<G: std::borrow::BorrowMut<group::Producer>> {
group: G,
buf: FrameBuf,
info: Info,
done: bool,
stats: stats::Meter,
}
impl<G: std::borrow::BorrowMut<group::Producer>> Raw<G> {
fn remaining(&self) -> usize {
self.buf.size() - self.buf.written(Ordering::Acquire)
}
fn write<B: IntoBytes>(&mut self, chunk: B) -> Result<()> {
let len = chunk.as_ref().len();
if len > self.remaining() {
return Err(Error::WrongSize);
}
self.stats.bytes(len as u64);
if len == self.buf.size() && self.buf.written(Ordering::Acquire) == 0 {
match self.buf.try_set_bytes(chunk.into_bytes()) {
Ok(()) => {
let size = self.buf.size();
unsafe { self.buf.store_written(size) };
}
Err(chunk) => self.buf.append(&chunk),
}
} else {
self.buf.append(chunk.as_ref());
}
Ok(())
}
fn finish(&mut self) -> Result<()> {
if self.buf.written(Ordering::Acquire) != self.buf.size() {
return Err(Error::WrongSize);
}
let payload = self.buf.freeze();
self.group.borrow_mut().frame_commit(Frame {
timestamp: self.info.timestamp,
payload,
})?;
self.done = true;
Ok(())
}
fn abort(&mut self, err: Error) -> Result<()> {
self.group.borrow_mut().frame_abort(err);
self.done = true;
Ok(())
}
}
impl<G: std::borrow::BorrowMut<group::Producer>> Drop for Raw<G> {
fn drop(&mut self) {
if !self.done {
let group = self.group.borrow_mut();
if !group.is_aborted() {
tracing::warn!(
group = group.info().sequence,
"frame::Producer dropped before writing all bytes"
);
}
group.frame_abort(Error::Dropped);
}
}
}
pub struct Producer<'a>(Raw<&'a mut group::Producer>);
impl std::ops::Deref for Producer<'_> {
type Target = Info;
fn deref(&self) -> &Self::Target {
&self.0.info
}
}
impl<'a> Producer<'a> {
pub(crate) fn new(group: &'a mut group::Producer, buf: FrameBuf, info: Info) -> Self {
Self(Raw {
group,
buf,
info,
done: false,
stats: stats::Meter::default(),
})
}
pub(crate) fn with_meter(mut self, meter: stats::Meter) -> Self {
self.0.stats = meter;
self
}
pub fn group(&self) -> group::Info {
self.0.group.info()
}
pub fn remaining(&self) -> usize {
self.0.remaining()
}
pub fn write<B: IntoBytes>(&mut self, chunk: B) -> Result<()> {
self.0.write(chunk)?;
self.0.group.frame_notify();
Ok(())
}
pub fn finish(mut self) -> Result<()> {
self.0.finish()
}
pub fn abort(mut self, err: Error) -> Result<()> {
self.0.abort(err)
}
}
pub(crate) struct ProducerOwned {
raw: Raw<group::Producer>,
_reserved: Option<Reservation>,
}
impl std::ops::Deref for ProducerOwned {
type Target = Info;
fn deref(&self) -> &Self::Target {
&self.raw.info
}
}
impl ProducerOwned {
pub(crate) fn new(group: group::Producer, buf: FrameBuf, info: Info, reserved: Option<Reservation>) -> Self {
Self {
raw: Raw {
group,
buf,
info,
done: false,
stats: stats::Meter::default(),
},
_reserved: reserved,
}
}
pub(crate) fn with_meter(mut self, meter: stats::Meter) -> Self {
self.raw.stats = meter;
self
}
pub fn remaining(&self) -> usize {
self.raw.remaining()
}
pub(crate) fn write<B: IntoBytes>(&mut self, chunk: B) -> Result<()> {
self.raw.write(chunk)
}
pub(crate) fn notify(&self) {
self.raw.group.frame_notify();
}
pub fn finish(mut self) -> Result<()> {
self.raw.finish()
}
pub fn abort(mut self, err: Error) -> Result<()> {
self.raw.abort(err)
}
#[cfg(test)]
pub(crate) fn allocated(&self) -> usize {
self.raw.buf.allocated()
}
}
#[derive(Clone)]
pub(crate) enum Source {
Complete(Bytes),
Partial(FrameBuf),
}
#[derive(Clone)]
pub(crate) struct Expiry {
policy: Arc<dyn group::Expiry>,
stale_stats: stats::Meter,
stale_counted: Arc<AtomicBool>,
tail: Range<usize>,
count_payload: bool,
}
impl Expiry {
pub(crate) fn new(
policy: Arc<dyn group::Expiry>,
stale_stats: stats::Meter,
stale_counted: Arc<AtomicBool>,
) -> Self {
Self {
policy,
stale_stats,
stale_counted,
tail: 0..0,
count_payload: false,
}
}
pub(crate) fn for_frame(mut self, tail: Range<usize>, count_payload: bool) -> Self {
self.tail = tail;
self.count_payload = count_payload;
self
}
}
#[derive(Clone)]
pub struct Consumer {
state: kio::Consumer<GroupState>,
info: Info,
source: Source,
read_idx: usize,
stats: stats::Meter,
expiry: Option<Expiry>,
expired: bool,
}
impl std::ops::Deref for Consumer {
type Target = Info;
fn deref(&self) -> &Self::Target {
&self.info
}
}
impl Consumer {
pub(crate) fn new(state: kio::Consumer<GroupState>, info: Info, source: Source) -> Self {
Self {
state,
info,
source,
read_idx: 0,
stats: stats::Meter::default(),
expiry: None,
expired: false,
}
}
pub(crate) fn with_meter(mut self, meter: stats::Meter) -> Self {
self.stats = meter;
self
}
pub(crate) fn with_expiry(mut self, expiry: Expiry) -> Self {
self.expiry = Some(expiry);
self
}
fn size(&self) -> usize {
match &self.source {
Source::Complete(bytes) => bytes.len(),
Source::Partial(_) => self.info.size as usize,
}
}
fn poll_expired(&mut self, waiter: &kio::Waiter) -> bool {
if self.expired || self.read_idx >= self.size() {
return self.expired;
}
let Some(expiry) = &self.expiry else {
return false;
};
if !expiry.policy.is_expired(waiter) {
return false;
}
self.expired = true;
if !expiry.stale_counted.swap(true, Ordering::Relaxed) {
let mut stale = self.state.read().content_range(expiry.tail.start, expiry.tail.end);
if expiry.count_payload {
stale.bytes += self.size().saturating_sub(self.read_idx) as u64;
}
expiry.stale_stats.stale(stale);
}
true
}
pub fn poll_read_chunk(&mut self, waiter: &kio::Waiter) -> Poll<Result<Option<Bytes>>> {
if self.expired {
return Poll::Ready(Err(Error::Old));
}
let buf = match &self.source {
Source::Complete(bytes) => {
if self.read_idx >= bytes.len() {
return Poll::Ready(Ok(None));
}
let out = bytes.slice(self.read_idx..);
self.read_idx = bytes.len();
self.stats.bytes(out.len() as u64);
return Poll::Ready(Ok(Some(out)));
}
Source::Partial(buf) => buf.clone(),
};
let size = self.info.size as usize;
loop {
let written = buf.written(Ordering::Acquire);
if written > self.read_idx {
let out = buf.slice(self.read_idx, written);
self.read_idx = written;
self.stats.bytes(out.len() as u64);
return Poll::Ready(Ok(Some(out)));
}
if written >= size {
return Poll::Ready(Ok(None));
}
if self.poll_expired(waiter) {
return Poll::Ready(Err(Error::Old));
}
let read_idx = self.read_idx;
ready!(poll_state(&self.state, waiter, |state| {
if let Some(err) = &state.abort {
return Poll::Ready(Err(err.clone()));
}
let w = buf.written(Ordering::Acquire);
if w > read_idx || w >= size {
Poll::Ready(Ok(()))
} else {
Poll::Pending
}
})?);
}
}
pub async fn read_chunk(&mut self) -> Result<Option<Bytes>> {
kio::wait(|waiter| self.poll_read_chunk(waiter)).await
}
pub fn poll_read_all(&mut self, waiter: &kio::Waiter) -> Poll<Result<Bytes>> {
if self.expired {
return Poll::Ready(Err(Error::Old));
}
let buf = match &self.source {
Source::Complete(bytes) => {
let out = bytes.slice(self.read_idx..);
self.read_idx = bytes.len();
self.stats.bytes(out.len() as u64);
return Poll::Ready(Ok(out));
}
Source::Partial(buf) => buf.clone(),
};
let size = self.info.size as usize;
let read_idx = self.read_idx;
if buf.written(Ordering::Acquire) < size && self.poll_expired(waiter) {
return Poll::Ready(Err(Error::Old));
}
ready!(poll_state(&self.state, waiter, |state| {
if let Some(err) = &state.abort {
return Poll::Ready(Err(err.clone()));
}
if buf.written(Ordering::Acquire) >= size {
Poll::Ready(Ok(()))
} else {
Poll::Pending
}
})?);
let out = buf.slice(read_idx, size);
self.read_idx = size;
self.stats.bytes(out.len() as u64);
Poll::Ready(Ok(out))
}
pub async fn read_all(&mut self) -> Result<Bytes> {
kio::wait(|waiter| self.poll_read_all(waiter)).await
}
}
fn poll_state<F, R>(state: &kio::Consumer<GroupState>, waiter: &kio::Waiter, f: F) -> Poll<Result<R>>
where
F: Fn(&kio::Ref<'_, GroupState>) -> Poll<Result<R>>,
{
Poll::Ready(match ready!(state.poll(waiter, f)) {
Ok(res) => res,
Err(state) => Err(state.abort.clone().unwrap_or(Error::Dropped)),
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn growing_buffer_keeps_handed_out_slices() {
let buf = FrameBuf::growing(10);
buf.append(b"abc");
assert_eq!(buf.allocated(), 3);
let early = buf.slice(0, 3);
buf.append(b"defg");
assert_eq!(buf.allocated(), 7);
buf.append(b"hij");
assert_eq!(buf.allocated(), 10, "doubling stops at the declared size");
assert_eq!(early, &b"abc"[..]);
assert_eq!(buf.slice(3, 7), &b"defg"[..]);
assert_eq!(buf.freeze(), &b"abcdefghij"[..]);
}
#[test]
fn growing_buffer_reads_across_threads() {
const SIZE: usize = 256 * 1024;
let buf = FrameBuf::growing(SIZE);
let reader = std::thread::spawn({
let buf = buf.clone();
move || {
let mut read = 0;
while read < SIZE {
let written = buf.written(Ordering::Acquire);
let chunk = buf.slice(read, written);
for (i, byte) in chunk.iter().enumerate() {
assert_eq!(*byte, ((read + i) % 251) as u8);
}
read = written;
std::thread::yield_now();
}
}
});
let payload: Vec<u8> = (0..SIZE).map(|i| (i % 251) as u8).collect();
for chunk in payload.chunks(1000) {
buf.append(chunk);
}
reader.join().unwrap();
assert_eq!(buf.freeze(), payload);
}
#[test]
fn budget_returns_reservations() {
let budget = Budget::new(100);
let a = budget.reserve(60).unwrap();
assert!(budget.reserve(41).is_none());
let b = budget.reserve(40).unwrap();
assert!(budget.reserve(1).is_none());
drop(a);
drop(b);
assert!(budget.reserve(100).is_some());
}
}