xref: /webrtc/interceptor/src/twcc/receiver/mod.rs (revision 603f4064)
1 mod receiver_stream;
2 #[cfg(test)]
3 mod receiver_test;
4 
5 use crate::twcc::sender::TRANSPORT_CC_URI;
6 use crate::twcc::Recorder;
7 use crate::*;
8 use receiver_stream::ReceiverStream;
9 
10 use rtp::extension::transport_cc_extension::TransportCcExtension;
11 use std::time::{Duration, SystemTime};
12 use tokio::sync::{mpsc, Mutex};
13 use tokio::time::MissedTickBehavior;
14 use util::Unmarshal;
15 use waitgroup::WaitGroup;
16 
17 /// ReceiverBuilder is a InterceptorBuilder for a SenderInterceptor
18 #[derive(Default)]
19 pub struct ReceiverBuilder {
20     interval: Option<Duration>,
21 }
22 
23 impl ReceiverBuilder {
24     /// with_interval sets send interval for the interceptor.
25     pub fn with_interval(mut self, interval: Duration) -> ReceiverBuilder {
26         self.interval = Some(interval);
27         self
28     }
29 }
30 
31 impl InterceptorBuilder for ReceiverBuilder {
32     fn build(&self, _id: &str) -> Result<Arc<dyn Interceptor + Send + Sync>> {
33         let (close_tx, close_rx) = mpsc::channel(1);
34         let (packet_chan_tx, packet_chan_rx) = mpsc::channel(1);
35         Ok(Arc::new(Receiver {
36             internal: Arc::new(ReceiverInternal {
37                 interval: if let Some(interval) = &self.interval {
38                     *interval
39                 } else {
40                     Duration::from_millis(100)
41                 },
42                 recorder: Mutex::new(Recorder::default()),
43                 packet_chan_rx: Mutex::new(Some(packet_chan_rx)),
44                 streams: Mutex::new(HashMap::new()),
45                 close_rx: Mutex::new(Some(close_rx)),
46             }),
47             start_time: SystemTime::now(),
48             packet_chan_tx,
49             wg: Mutex::new(Some(WaitGroup::new())),
50             close_tx: Mutex::new(Some(close_tx)),
51         }))
52     }
53 }
54 
55 struct Packet {
56     hdr: rtp::header::Header,
57     sequence_number: u16,
58     arrival_time: i64,
59     ssrc: u32,
60 }
61 
62 struct ReceiverInternal {
63     interval: Duration,
64     recorder: Mutex<Recorder>,
65     packet_chan_rx: Mutex<Option<mpsc::Receiver<Packet>>>,
66     streams: Mutex<HashMap<u32, Arc<ReceiverStream>>>,
67     close_rx: Mutex<Option<mpsc::Receiver<()>>>,
68 }
69 
70 /// Receiver sends transport wide congestion control reports as specified in:
71 /// https://datatracker.ietf.org/doc/html/draft-holmer-rmcat-transport-wide-cc-extensions-01
72 pub struct Receiver {
73     internal: Arc<ReceiverInternal>,
74 
75     start_time: SystemTime,
76     packet_chan_tx: mpsc::Sender<Packet>,
77 
78     wg: Mutex<Option<WaitGroup>>,
79     close_tx: Mutex<Option<mpsc::Sender<()>>>,
80 }
81 
82 impl Receiver {
83     /// builder returns a new ReceiverBuilder.
84     pub fn builder() -> ReceiverBuilder {
85         ReceiverBuilder::default()
86     }
87 
88     async fn is_closed(&self) -> bool {
89         let close_tx = self.close_tx.lock().await;
90         close_tx.is_none()
91     }
92 
93     async fn run(
94         rtcp_writer: Arc<dyn RTCPWriter + Send + Sync>,
95         internal: Arc<ReceiverInternal>,
96     ) -> Result<()> {
97         let mut close_rx = {
98             let mut close_rx = internal.close_rx.lock().await;
99             if let Some(close_rx) = close_rx.take() {
100                 close_rx
101             } else {
102                 return Err(Error::ErrInvalidCloseRx);
103             }
104         };
105         let mut packet_chan_rx = {
106             let mut packet_chan_rx = internal.packet_chan_rx.lock().await;
107             if let Some(packet_chan_rx) = packet_chan_rx.take() {
108                 packet_chan_rx
109             } else {
110                 return Err(Error::ErrInvalidPacketRx);
111             }
112         };
113 
114         let a = Attributes::new();
115         let mut ticker = tokio::time::interval(internal.interval);
116         ticker.set_missed_tick_behavior(MissedTickBehavior::Skip);
117         loop {
118             tokio::select! {
119                 _ = close_rx.recv() =>{
120                     return Ok(());
121                 }
122                 p = packet_chan_rx.recv() => {
123                     if let Some(p) = p {
124                         let mut recorder = internal.recorder.lock().await;
125                         recorder.record(p.ssrc, p.sequence_number, p.arrival_time);
126                     }
127                 }
128                 _ = ticker.tick() =>{
129                     // build and send twcc
130                     let pkts = {
131                         let mut recorder = internal.recorder.lock().await;
132                         recorder.build_feedback_packet()
133                     };
134 
135                     if let Err(err) = rtcp_writer.write(&pkts, &a).await{
136                         log::error!("rtcp_writer.write got err: {}", err);
137                     }
138                 }
139             }
140         }
141     }
142 }
143 
144 #[async_trait]
145 impl Interceptor for Receiver {
146     /// bind_rtcp_reader lets you modify any incoming RTCP packets. It is called once per sender/receiver, however this might
147     /// change in the future. The returned method will be called once per packet batch.
148     async fn bind_rtcp_reader(
149         &self,
150         reader: Arc<dyn RTCPReader + Send + Sync>,
151     ) -> Arc<dyn RTCPReader + Send + Sync> {
152         reader
153     }
154 
155     /// bind_rtcp_writer lets you modify any outgoing RTCP packets. It is called once per PeerConnection. The returned method
156     /// will be called once per packet batch.
157     async fn bind_rtcp_writer(
158         &self,
159         writer: Arc<dyn RTCPWriter + Send + Sync>,
160     ) -> Arc<dyn RTCPWriter + Send + Sync> {
161         if self.is_closed().await {
162             return writer;
163         }
164 
165         {
166             let mut recorder = self.internal.recorder.lock().await;
167             *recorder = Recorder::new(rand::random::<u32>());
168         }
169 
170         let mut w = {
171             let wait_group = self.wg.lock().await;
172             wait_group.as_ref().map(|wg| wg.worker())
173         };
174         let writer2 = Arc::clone(&writer);
175         let internal = Arc::clone(&self.internal);
176         tokio::spawn(async move {
177             let _d = w.take();
178             if let Err(err) = Receiver::run(writer2, internal).await {
179                 log::warn!("bind_rtcp_writer TWCC Sender::run got error: {}", err);
180             }
181         });
182 
183         writer
184     }
185 
186     /// bind_local_stream lets you modify any outgoing RTP packets. It is called once for per LocalStream. The returned method
187     /// will be called once per rtp packet.
188     async fn bind_local_stream(
189         &self,
190         _info: &StreamInfo,
191         writer: Arc<dyn RTPWriter + Send + Sync>,
192     ) -> Arc<dyn RTPWriter + Send + Sync> {
193         writer
194     }
195 
196     /// unbind_local_stream is called when the Stream is removed. It can be used to clean up any data related to that track.
197     async fn unbind_local_stream(&self, _info: &StreamInfo) {}
198 
199     /// bind_remote_stream lets you modify any incoming RTP packets. It is called once for per RemoteStream. The returned method
200     /// will be called once per rtp packet.
201     async fn bind_remote_stream(
202         &self,
203         info: &StreamInfo,
204         reader: Arc<dyn RTPReader + Send + Sync>,
205     ) -> Arc<dyn RTPReader + Send + Sync> {
206         let mut hdr_ext_id = 0u8;
207         for e in &info.rtp_header_extensions {
208             if e.uri == TRANSPORT_CC_URI {
209                 hdr_ext_id = e.id as u8;
210                 break;
211             }
212         }
213         if hdr_ext_id == 0 {
214             // Don't try to read header extension if ID is 0, because 0 is an invalid extension ID
215             return reader;
216         }
217 
218         let stream = Arc::new(ReceiverStream::new(
219             reader,
220             hdr_ext_id,
221             info.ssrc,
222             self.packet_chan_tx.clone(),
223             self.start_time,
224         ));
225 
226         {
227             let mut streams = self.internal.streams.lock().await;
228             streams.insert(info.ssrc, Arc::clone(&stream));
229         }
230 
231         stream
232     }
233 
234     /// unbind_remote_stream is called when the Stream is removed. It can be used to clean up any data related to that track.
235     async fn unbind_remote_stream(&self, info: &StreamInfo) {
236         let mut streams = self.internal.streams.lock().await;
237         streams.remove(&info.ssrc);
238     }
239 
240     /// close closes the Interceptor, cleaning up any data if necessary.
241     async fn close(&self) -> Result<()> {
242         {
243             let mut close_tx = self.close_tx.lock().await;
244             close_tx.take();
245         }
246 
247         {
248             let mut wait_group = self.wg.lock().await;
249             if let Some(wg) = wait_group.take() {
250                 wg.wait().await;
251             }
252         }
253 
254         Ok(())
255     }
256 }
257