1use super::receive_log::ReceiveLog;
4use super::stream_supports_nack;
5use crate::stream_info::StreamInfo;
6use crate::{Interceptor, Packet, TaggedPacket, interceptor};
7use shared::TransportContext;
8use shared::error::Error;
9use std::collections::{HashMap, VecDeque};
10use std::marker::PhantomData;
11use std::time::{Duration, Instant};
12
13pub struct NackGeneratorBuilder<P> {
30 size: u16,
32 interval: Duration,
34 skip_last_n: u16,
36 max_nacks_per_packet: u16,
38 _phantom: PhantomData<P>,
39}
40
41impl<P> Default for NackGeneratorBuilder<P> {
42 fn default() -> Self {
43 Self {
44 size: 512,
45 interval: Duration::from_millis(100),
46 skip_last_n: 0,
47 max_nacks_per_packet: 0,
48 _phantom: PhantomData,
49 }
50 }
51}
52
53impl<P> NackGeneratorBuilder<P> {
54 pub fn new() -> Self {
56 Self::default()
57 }
58
59 pub fn with_size(mut self, size: u16) -> Self {
63 self.size = size;
64 self
65 }
66
67 pub fn with_interval(mut self, interval: Duration) -> Self {
69 self.interval = interval;
70 self
71 }
72
73 pub fn with_skip_last_n(mut self, skip_last_n: u16) -> Self {
78 self.skip_last_n = skip_last_n;
79 self
80 }
81
82 pub fn with_max_nacks_per_packet(mut self, max: u16) -> Self {
86 self.max_nacks_per_packet = max;
87 self
88 }
89
90 pub fn build(self) -> impl FnOnce(P) -> NackGeneratorInterceptor<P> {
92 move |inner| {
93 NackGeneratorInterceptor::new(
94 inner,
95 self.size,
96 self.interval,
97 self.skip_last_n,
98 self.max_nacks_per_packet,
99 )
100 }
101 }
102}
103
104#[derive(Interceptor)]
110pub struct NackGeneratorInterceptor<P> {
111 #[next]
112 inner: P,
113
114 size: u16,
116 interval: Duration,
117 skip_last_n: u16,
118 max_nacks_per_packet: u16,
119
120 next_timeout: Option<Instant>,
122
123 sender_ssrc: u32,
125
126 receive_logs: HashMap<u32, ReceiveLog>,
128
129 nack_counts: HashMap<u32, HashMap<u16, u16>>,
131
132 write_queue: VecDeque<TaggedPacket>,
134}
135
136impl<P> NackGeneratorInterceptor<P> {
137 fn new(
138 inner: P,
139 size: u16,
140 interval: Duration,
141 skip_last_n: u16,
142 max_nacks_per_packet: u16,
143 ) -> Self {
144 Self {
145 inner,
146 size,
147 interval,
148 skip_last_n,
149 max_nacks_per_packet,
150 next_timeout: None,
151 sender_ssrc: rand::random(),
152 receive_logs: HashMap::new(),
153 nack_counts: HashMap::new(),
154 write_queue: VecDeque::new(),
155 }
156 }
157
158 fn generate_nacks(&mut self, now: Instant) {
160 for (&ssrc, receive_log) in &self.receive_logs {
161 let missing = receive_log.missing_seq_numbers(self.skip_last_n);
162 if missing.is_empty() {
163 self.nack_counts.remove(&ssrc);
165 continue;
166 }
167
168 let nack_count = self.nack_counts.entry(ssrc).or_default();
170
171 let filtered: Vec<u16> = if self.max_nacks_per_packet > 0 {
173 missing
174 .iter()
175 .filter(|&&seq| {
176 let count = nack_count.entry(seq).or_insert(0);
177 if *count < self.max_nacks_per_packet {
178 *count += 1;
179 true
180 } else {
181 false
182 }
183 })
184 .copied()
185 .collect()
186 } else {
187 missing.clone()
188 };
189
190 if filtered.is_empty() {
191 continue;
192 }
193
194 nack_count.retain(|seq, _| missing.contains(seq));
196
197 let nack = rtcp::transport_feedbacks::transport_layer_nack::TransportLayerNack {
199 sender_ssrc: self.sender_ssrc,
200 media_ssrc: ssrc,
201 nacks: rtcp::transport_feedbacks::transport_layer_nack::nack_pairs_from_sequence_numbers(
202 &filtered,
203 ),
204 };
205
206 self.write_queue.push_back(TaggedPacket {
207 now,
208 transport: TransportContext::default(),
209 message: Packet::Rtcp(vec![Box::new(nack)]),
210 });
211 }
212 }
213}
214
215#[interceptor]
216impl<P: Interceptor> NackGeneratorInterceptor<P> {
217 #[overrides]
218 fn handle_read(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
219 if let Packet::Rtp(ref rtp_packet) = msg.message
221 && let Some(receive_log) = self.receive_logs.get_mut(&rtp_packet.header.ssrc)
222 {
223 receive_log.add(rtp_packet.header.sequence_number);
224
225 if self.next_timeout.is_none() {
228 self.next_timeout = Some(msg.now + self.interval);
229 }
230 }
231
232 self.inner.handle_read(msg)
233 }
234
235 #[overrides]
236 fn poll_write(&mut self) -> Option<Self::Wout> {
237 if let Some(pkt) = self.write_queue.pop_front() {
239 return Some(pkt);
240 }
241 self.inner.poll_write()
242 }
243
244 #[overrides]
245 fn handle_timeout(&mut self, now: Self::Time) -> Result<(), Self::Error> {
246 if let Some(next_timeout) = self.next_timeout
247 && now >= next_timeout
248 {
249 self.next_timeout = Some(now + self.interval);
250 self.generate_nacks(now);
251 }
252
253 self.inner.handle_timeout(now)
254 }
255
256 #[overrides]
257 fn poll_timeout(&mut self) -> Option<Self::Time> {
258 match (self.next_timeout, self.inner.poll_timeout()) {
259 (Some(a), Some(b)) => Some(a.min(b)),
260 (Some(a), None) => Some(a),
261 (None, Some(b)) => Some(b),
262 (None, None) => None,
263 }
264 }
265
266 #[overrides]
267 fn bind_remote_stream(&mut self, info: &StreamInfo) {
268 if stream_supports_nack(info)
269 && let Some(receive_log) = ReceiveLog::new(self.size)
270 {
271 self.receive_logs.insert(info.ssrc, receive_log);
272 }
273 self.inner.bind_remote_stream(info);
274 }
275
276 #[overrides]
277 fn unbind_remote_stream(&mut self, info: &StreamInfo) {
278 self.receive_logs.remove(&info.ssrc);
279 self.nack_counts.remove(&info.ssrc);
280 self.inner.unbind_remote_stream(info);
281 }
282}
283
284#[cfg(test)]
285mod tests {
286 use super::*;
287 use crate::Registry;
288 use crate::stream_info::RTCPFeedback;
289 use sansio::Protocol;
290
291 fn make_rtp_packet(ssrc: u32, seq: u16) -> TaggedPacket {
292 TaggedPacket {
293 now: Instant::now(),
294 transport: Default::default(),
295 message: Packet::Rtp(rtp::Packet {
296 header: rtp::header::Header {
297 ssrc,
298 sequence_number: seq,
299 ..Default::default()
300 },
301 ..Default::default()
302 }),
303 }
304 }
305
306 #[test]
307 fn test_nack_generator_builder_defaults() {
308 let chain = Registry::new()
309 .with(NackGeneratorBuilder::default().build())
310 .build();
311
312 assert_eq!(chain.size, 512);
313 assert_eq!(chain.interval, Duration::from_millis(100));
314 assert_eq!(chain.skip_last_n, 0);
315 assert_eq!(chain.max_nacks_per_packet, 0);
316 }
317
318 #[test]
319 fn test_nack_generator_builder_custom() {
320 let chain = Registry::new()
321 .with(
322 NackGeneratorBuilder::new()
323 .with_size(1024)
324 .with_interval(Duration::from_millis(50))
325 .with_skip_last_n(3)
326 .with_max_nacks_per_packet(5)
327 .build(),
328 )
329 .build();
330
331 assert_eq!(chain.size, 1024);
332 assert_eq!(chain.interval, Duration::from_millis(50));
333 assert_eq!(chain.skip_last_n, 3);
334 assert_eq!(chain.max_nacks_per_packet, 5);
335 }
336
337 #[test]
338 fn test_nack_generator_no_nack_without_binding() {
339 let mut chain = Registry::new()
340 .with(
341 NackGeneratorBuilder::new()
342 .with_interval(Duration::from_millis(100))
343 .build(),
344 )
345 .build();
346
347 let now = Instant::now();
348
349 chain.handle_read(make_rtp_packet(12345, 0)).unwrap();
351 chain.handle_read(make_rtp_packet(12345, 2)).unwrap(); let later = now + Duration::from_millis(200);
355 chain.handle_timeout(later).unwrap();
356
357 assert!(chain.poll_write().is_none());
359 }
360
361 #[test]
362 fn test_nack_generator_generates_nack() {
363 let mut chain = Registry::new()
364 .with(
365 NackGeneratorBuilder::new()
366 .with_size(64)
367 .with_interval(Duration::from_millis(100))
368 .build(),
369 )
370 .build();
371
372 let info = StreamInfo {
374 ssrc: 12345,
375 clock_rate: 90000,
376 rtcp_feedback: vec![RTCPFeedback {
377 typ: "nack".to_string(),
378 parameter: "".to_string(),
379 }],
380 ..Default::default()
381 };
382 chain.bind_remote_stream(&info);
383
384 let base_time = Instant::now();
385
386 let mut pkt = make_rtp_packet(12345, 10);
388 pkt.now = base_time;
389 chain.handle_read(pkt).unwrap();
390
391 let mut pkt = make_rtp_packet(12345, 12); pkt.now = base_time;
393 chain.handle_read(pkt).unwrap();
394
395 chain.poll_read();
396
397 let later = base_time + Duration::from_millis(200);
399 chain.handle_timeout(later).unwrap();
400
401 let nack_pkt = chain.poll_write();
403 assert!(nack_pkt.is_some());
404
405 if let Some(tagged) = nack_pkt {
406 if let Packet::Rtcp(rtcp_packets) = tagged.message {
407 assert_eq!(rtcp_packets.len(), 1);
408 let nack = rtcp_packets[0]
409 .as_any()
410 .downcast_ref::<rtcp::transport_feedbacks::transport_layer_nack::TransportLayerNack>()
411 .expect("Expected TransportLayerNack");
412 assert_eq!(nack.media_ssrc, 12345);
413 assert!(!nack.nacks.is_empty());
414 } else {
415 panic!("Expected RTCP packet");
416 }
417 }
418 }
419
420 #[test]
421 fn test_nack_generator_skip_last_n() {
422 let mut chain = Registry::new()
423 .with(
424 NackGeneratorBuilder::new()
425 .with_size(64)
426 .with_interval(Duration::from_millis(100))
427 .with_skip_last_n(2)
428 .build(),
429 )
430 .build();
431
432 let info = StreamInfo {
433 ssrc: 12345,
434 clock_rate: 90000,
435 rtcp_feedback: vec![RTCPFeedback {
436 typ: "nack".to_string(),
437 parameter: "".to_string(),
438 }],
439 ..Default::default()
440 };
441 chain.bind_remote_stream(&info);
442
443 let base_time = Instant::now();
444
445 for seq in [10u16, 11, 12, 14, 16, 18] {
447 let mut pkt = make_rtp_packet(12345, seq);
448 pkt.now = base_time;
449 chain.handle_read(pkt).unwrap();
450 }
451
452 let later = base_time + Duration::from_millis(200);
454 chain.handle_timeout(later).unwrap();
455
456 let nack_pkt = chain.poll_write();
458 assert!(nack_pkt.is_some());
459
460 if let Some(tagged) = nack_pkt
461 && let Packet::Rtcp(rtcp_packets) = tagged.message
462 {
463 let nack = rtcp_packets[0]
464 .as_any()
465 .downcast_ref::<rtcp::transport_feedbacks::transport_layer_nack::TransportLayerNack>()
466 .expect("Expected TransportLayerNack");
467
468 let mut nacked_seqs = Vec::new();
470 for nack_pair in &nack.nacks {
471 nacked_seqs.push(nack_pair.packet_id);
472 for i in 0..16 {
473 if nack_pair.lost_packets & (1 << i) != 0 {
474 nacked_seqs.push(nack_pair.packet_id.wrapping_add(i + 1));
475 }
476 }
477 }
478
479 assert!(nacked_seqs.contains(&13));
481 assert!(nacked_seqs.contains(&15));
482 assert!(!nacked_seqs.contains(&17));
483 }
484 }
485
486 #[test]
487 fn test_nack_generator_unbind_removes_stream() {
488 let mut chain = Registry::new()
489 .with(
490 NackGeneratorBuilder::new()
491 .with_size(64)
492 .with_interval(Duration::from_millis(100))
493 .build(),
494 )
495 .build();
496
497 let info = StreamInfo {
498 ssrc: 12345,
499 clock_rate: 90000,
500 rtcp_feedback: vec![RTCPFeedback {
501 typ: "nack".to_string(),
502 parameter: "".to_string(),
503 }],
504 ..Default::default()
505 };
506
507 chain.bind_remote_stream(&info);
508 assert!(chain.receive_logs.contains_key(&12345));
509
510 chain.unbind_remote_stream(&info);
511 assert!(!chain.receive_logs.contains_key(&12345));
512 assert!(!chain.nack_counts.contains_key(&12345));
513 }
514
515 #[test]
516 fn test_nack_generator_no_nack_support() {
517 let mut chain = Registry::new()
518 .with(
519 NackGeneratorBuilder::new()
520 .with_size(64)
521 .with_interval(Duration::from_millis(100))
522 .build(),
523 )
524 .build();
525
526 let info = StreamInfo {
528 ssrc: 12345,
529 clock_rate: 90000,
530 rtcp_feedback: vec![], ..Default::default()
532 };
533 chain.bind_remote_stream(&info);
534
535 assert!(!chain.receive_logs.contains_key(&12345));
537 }
538}