1use std::fmt;
4use std::sync::atomic::{AtomicBool, Ordering};
5use std::sync::{Arc, Mutex};
6use std::task::{Poll, ready};
7
8use bytes::Bytes;
9
10use crate::error::{Error, Result};
11use crate::key::TrackKey;
12use crate::limits::MAX_GROUPED_PAYLOAD;
13
14#[derive(Clone, Debug, PartialEq, Eq)]
16pub struct Frame {
17 pub timestamp: moq_net::Timestamp,
19 pub plaintext: Bytes,
21}
22
23pub struct Producer {
27 inner: moq_net::group::Producer,
28 key: Arc<Mutex<TrackKey>>,
29 next_frame: u32,
30}
31
32impl Producer {
33 pub(crate) fn new(inner: moq_net::group::Producer, key: Arc<Mutex<TrackKey>>) -> Self {
34 Self {
35 inner,
36 key,
37 next_frame: 0,
38 }
39 }
40
41 pub fn sequence(&self) -> u64 {
43 self.inner.sequence
44 }
45
46 pub fn next_frame(&self) -> u32 {
48 self.next_frame
49 }
50
51 pub fn write_frame(&mut self, timestamp: moq_net::Timestamp, plaintext: &[u8]) -> Result<()> {
62 let frame = self.next_frame;
63 if frame == u32::MAX {
64 return Err(Error::Identity);
65 }
66 timestamp
67 .convert(self.inner.timescale())
68 .map_err(|_| Error::Net(moq_net::Error::TimestampMismatch))?;
69 let payload = self.key.lock().expect("track key").protect(
70 self.inner.sequence,
71 u64::from(frame),
72 plaintext,
73 MAX_GROUPED_PAYLOAD,
74 )?;
75 self.next_frame = frame + 1;
76 self.inner.write_frame(timestamp, payload)?;
77 Ok(())
78 }
79
80 pub fn finish(self) -> Result<()> {
86 self.inner.finish()?;
87 Ok(())
88 }
89
90 pub fn abort(self) -> Result<()> {
96 self.inner.abort(moq_net::Error::Cancel)?;
97 Ok(())
98 }
99
100 #[cfg(test)]
101 pub(crate) fn invocations(&self) -> u64 {
102 self.key.lock().expect("track key").invocations()
103 }
104}
105
106impl fmt::Debug for Producer {
107 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
108 f.debug_struct("group::Producer")
109 .field("sequence", &self.sequence())
110 .field("next_frame", &self.next_frame)
111 .finish()
112 }
113}
114
115pub struct Consumer {
119 inner: moq_net::group::Consumer,
120 key: Arc<Mutex<TrackKey>>,
121 auth_failed: Arc<AtomicBool>,
122}
123
124impl Consumer {
125 pub(crate) fn new(
126 inner: moq_net::group::Consumer,
127 key: Arc<Mutex<TrackKey>>,
128 auth_failed: Arc<AtomicBool>,
129 ) -> Self {
130 Self {
131 inner,
132 key,
133 auth_failed,
134 }
135 }
136
137 pub fn sequence(&self) -> u64 {
139 self.inner.sequence
140 }
141
142 pub fn poll_read_frame(&mut self, waiter: &kio::Waiter) -> Poll<Result<Option<Frame>>> {
149 if self.auth_failed.load(Ordering::Acquire) {
150 return Poll::Ready(Err(Error::Authentication));
151 }
152 let Some(frame) = ready!(self.inner.poll_read_frame(waiter)?) else {
153 return Poll::Ready(Ok(None));
154 };
155 let index = self.inner.index().checked_sub(1).ok_or(Error::Identity)?;
158 let result =
159 self.key
160 .lock()
161 .expect("track key")
162 .open(self.inner.sequence, index, &frame.payload, MAX_GROUPED_PAYLOAD);
163 match result {
164 Ok(plaintext) => Poll::Ready(Ok(Some(Frame {
165 timestamp: frame.timestamp,
166 plaintext,
167 }))),
168 Err(err) => {
169 if matches!(err, Error::Authentication | Error::Identity) {
170 self.auth_failed.store(true, Ordering::Release);
171 }
172 Poll::Ready(Err(err))
173 }
174 }
175 }
176
177 pub async fn read_frame(&mut self) -> Result<Option<Frame>> {
183 kio::wait(|waiter| self.poll_read_frame(waiter)).await
184 }
185}
186
187impl fmt::Debug for Consumer {
188 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
189 f.debug_struct("group::Consumer")
190 .field("sequence", &self.sequence())
191 .field("index", &self.inner.index())
192 .finish()
193 }
194}