use std::sync::Arc;
use std::sync::OnceLock;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::task::{Poll, ready};
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(Clone)]
pub(crate) struct FrameBuf(Arc<FrameBufInner>);
struct FrameBufInner {
capacity: usize,
written: AtomicUsize,
storage: OnceLock<FrameStorage>,
}
enum FrameStorage {
Shared(Bytes),
Mutable(MutableFrameBuf),
}
struct MutableFrameBuf {
data: *mut u8,
capacity: usize,
}
unsafe impl Send for MutableFrameBuf {}
unsafe impl Sync for MutableFrameBuf {}
impl Drop for MutableFrameBuf {
fn drop(&mut self) {
unsafe {
let slice = std::ptr::slice_from_raw_parts_mut(self.data, self.capacity);
drop(Box::from_raw(slice));
}
}
}
impl MutableFrameBuf {
fn new(size: usize) -> Self {
let boxed: Box<[u8]> = vec![0u8; size].into_boxed_slice();
let capacity = boxed.len();
let data = Box::into_raw(boxed) as *mut u8;
Self { data, capacity }
}
}
impl FrameBuf {
pub(crate) fn new(size: usize) -> Self {
Self(Arc::new(FrameBufInner {
capacity: size,
written: AtomicUsize::new(0),
storage: OnceLock::new(),
}))
}
pub(crate) fn capacity(&self) -> usize {
self.0.capacity
}
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.capacity() || self.written(Ordering::Acquire) != 0 {
return Err(bytes);
}
self.0
.storage
.set(FrameStorage::Shared(bytes))
.map_err(|storage| match storage {
FrameStorage::Shared(bytes) => bytes,
FrameStorage::Mutable(_) => unreachable!("try_set_bytes only installs shared storage"),
})
}
fn mutable(&self) -> Option<&MutableFrameBuf> {
match self
.0
.storage
.get_or_init(|| FrameStorage::Mutable(MutableFrameBuf::new(self.capacity())))
{
FrameStorage::Shared(_) => None,
FrameStorage::Mutable(buf) => Some(buf),
}
}
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 Some(buf) = self.mutable() else {
return;
};
unsafe {
std::ptr::copy_nonoverlapping(src.as_ptr(), buf.data.add(prev), src.len());
self.store_written(prev + src.len());
}
}
fn freeze(&self, size: usize) -> Bytes {
match self.0.storage.get() {
Some(FrameStorage::Shared(bytes)) => bytes.clone(),
_ => self.slice(0, size),
}
}
fn slice(&self, start: usize, end: usize) -> Bytes {
Bytes::from_owner(self.clone()).slice(start..end)
}
}
impl AsRef<[u8]> for FrameBuf {
fn as_ref(&self) -> &[u8] {
let written = self.0.written.load(Ordering::Acquire);
match self.0.storage.get() {
Some(FrameStorage::Shared(bytes)) => &bytes[..written],
Some(FrameStorage::Mutable(buf)) => {
unsafe { std::slice::from_raw_parts(buf.data, written) }
}
None => &[],
}
}
}
pub struct Producer<'a> {
group: &'a mut group::Producer,
buf: FrameBuf,
info: Info,
done: bool,
stats: stats::Meter,
}
impl std::ops::Deref for Producer<'_> {
type Target = Info;
fn deref(&self) -> &Self::Target {
&self.info
}
}
impl<'a> Producer<'a> {
pub(crate) fn new(group: &'a mut group::Producer, buf: FrameBuf, info: Info) -> Self {
Self {
group,
buf,
info,
done: false,
stats: stats::Meter::default(),
}
}
pub(crate) fn with_meter(mut self, meter: stats::Meter) -> Self {
self.stats = meter;
self
}
pub fn group(&self) -> group::Info {
self.group.info()
}
pub fn remaining(&self) -> usize {
self.buf.capacity() - self.buf.written(Ordering::Acquire)
}
pub 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.capacity() && self.buf.written(Ordering::Acquire) == 0 {
match self.buf.try_set_bytes(chunk.into_bytes()) {
Ok(()) => {
let cap = self.buf.capacity();
unsafe { self.buf.store_written(cap) };
}
Err(chunk) => self.buf.append(&chunk),
}
} else {
self.buf.append(chunk.as_ref());
}
self.group.frame_notify();
Ok(())
}
pub fn finish(mut self) -> Result<()> {
if self.buf.written(Ordering::Acquire) != self.buf.capacity() {
return Err(Error::WrongSize);
}
let payload = self.buf.freeze(self.buf.capacity());
self.group.frame_commit(Frame {
timestamp: self.info.timestamp,
payload,
})?;
self.done = true;
Ok(())
}
pub fn abort(mut self, err: Error) -> Result<()> {
self.group.frame_abort(err);
self.done = true;
Ok(())
}
}
impl Drop for Producer<'_> {
fn drop(&mut self) {
if !self.done {
tracing::warn!(
group = self.group.info().sequence,
"frame::Producer dropped before writing all bytes"
);
self.group.frame_abort(Error::Dropped);
}
}
}
#[derive(Clone)]
pub(crate) enum Source {
Complete(Bytes),
Partial(FrameBuf),
}
#[derive(Clone)]
pub struct Consumer {
state: kio::Consumer<GroupState>,
info: Info,
source: Source,
read_idx: usize,
stats: stats::Meter,
}
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(),
}
}
pub(crate) fn with_meter(mut self, meter: stats::Meter) -> Self {
self.stats = meter;
self
}
pub fn poll_read_chunk(&mut self, waiter: &kio::Waiter) -> Poll<Result<Option<Bytes>>> {
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);
Poll::Ready(Ok(Some(out)))
}
Source::Partial(buf) => {
let 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));
}
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>> {
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);
Poll::Ready(Ok(out))
}
Source::Partial(buf) => {
let buf = buf.clone();
let size = self.info.size as usize;
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()));
}
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)),
})
}