use std::fmt;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
use std::task::{Poll, ready};
use bytes::Bytes;
use crate::error::{Error, Result};
use crate::key::TrackKey;
use crate::limits::MAX_GROUPED_PAYLOAD;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Frame {
pub timestamp: moq_net::Timestamp,
pub plaintext: Bytes,
}
pub struct Producer {
inner: moq_net::group::Producer,
key: Arc<Mutex<TrackKey>>,
next_frame: u32,
}
impl Producer {
pub(crate) fn new(inner: moq_net::group::Producer, key: Arc<Mutex<TrackKey>>) -> Self {
Self {
inner,
key,
next_frame: 0,
}
}
pub fn sequence(&self) -> u64 {
self.inner.sequence
}
pub fn next_frame(&self) -> u32 {
self.next_frame
}
pub fn write_frame(&mut self, timestamp: moq_net::Timestamp, plaintext: &[u8]) -> Result<()> {
let frame = self.next_frame;
if frame == u32::MAX {
return Err(Error::Identity);
}
timestamp
.convert(self.inner.timescale())
.map_err(|_| Error::Net(moq_net::Error::TimestampMismatch))?;
let payload = self.key.lock().expect("track key").protect(
self.inner.sequence,
u64::from(frame),
plaintext,
MAX_GROUPED_PAYLOAD,
)?;
self.next_frame = frame + 1;
self.inner.write_frame(timestamp, payload)?;
Ok(())
}
pub fn finish(self) -> Result<()> {
self.inner.finish()?;
Ok(())
}
pub fn abort(self) -> Result<()> {
self.inner.abort(moq_net::Error::Cancel)?;
Ok(())
}
#[cfg(test)]
pub(crate) fn invocations(&self) -> u64 {
self.key.lock().expect("track key").invocations()
}
}
impl fmt::Debug for Producer {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("group::Producer")
.field("sequence", &self.sequence())
.field("next_frame", &self.next_frame)
.finish()
}
}
pub struct Consumer {
inner: moq_net::group::Consumer,
key: Arc<Mutex<TrackKey>>,
auth_failed: Arc<AtomicBool>,
}
impl Consumer {
pub(crate) fn new(
inner: moq_net::group::Consumer,
key: Arc<Mutex<TrackKey>>,
auth_failed: Arc<AtomicBool>,
) -> Self {
Self {
inner,
key,
auth_failed,
}
}
pub fn sequence(&self) -> u64 {
self.inner.sequence
}
pub fn poll_read_frame(&mut self, waiter: &kio::Waiter) -> Poll<Result<Option<Frame>>> {
if self.auth_failed.load(Ordering::Acquire) {
return Poll::Ready(Err(Error::Authentication));
}
let Some(frame) = ready!(self.inner.poll_read_frame(waiter)?) else {
return Poll::Ready(Ok(None));
};
let index = self.inner.index().checked_sub(1).ok_or(Error::Identity)?;
let result =
self.key
.lock()
.expect("track key")
.open(self.inner.sequence, index, &frame.payload, MAX_GROUPED_PAYLOAD);
match result {
Ok(plaintext) => Poll::Ready(Ok(Some(Frame {
timestamp: frame.timestamp,
plaintext,
}))),
Err(err) => {
if matches!(err, Error::Authentication | Error::Identity) {
self.auth_failed.store(true, Ordering::Release);
}
Poll::Ready(Err(err))
}
}
}
pub async fn read_frame(&mut self) -> Result<Option<Frame>> {
kio::wait(|waiter| self.poll_read_frame(waiter)).await
}
}
impl fmt::Debug for Consumer {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("group::Consumer")
.field("sequence", &self.sequence())
.field("index", &self.inner.index())
.finish()
}
}