1use super::send_buffer::SendBuffer;
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::Instant;
12
13pub struct NackResponderBuilder<P> {
27 size: u16,
29 _phantom: PhantomData<P>,
30}
31
32impl<P> Default for NackResponderBuilder<P> {
33 fn default() -> Self {
34 Self {
35 size: 1024,
36 _phantom: PhantomData,
37 }
38 }
39}
40
41impl<P> NackResponderBuilder<P> {
42 pub fn new() -> Self {
44 Self::default()
45 }
46
47 pub fn with_size(mut self, size: u16) -> Self {
52 self.size = size;
53 self
54 }
55
56 pub fn build(self) -> impl FnOnce(P) -> NackResponderInterceptor<P> {
58 move |inner| NackResponderInterceptor::new(inner, self.size)
59 }
60}
61
62struct LocalStream {
64 send_buffer: SendBuffer,
66 ssrc_rtx: Option<u32>,
68 payload_type_rtx: Option<u8>,
70 rtx_sequence_number: u16,
72}
73
74#[derive(Interceptor)]
79pub struct NackResponderInterceptor<P> {
80 #[next]
81 inner: P,
82
83 size: u16,
85
86 streams: HashMap<u32, LocalStream>,
88
89 write_queue: VecDeque<TaggedPacket>,
91}
92
93impl<P> NackResponderInterceptor<P> {
94 fn new(inner: P, size: u16) -> Self {
95 Self {
96 inner,
97 size,
98 streams: HashMap::new(),
99 write_queue: VecDeque::new(),
100 }
101 }
102
103 fn handle_nack(
105 &mut self,
106 now: Instant,
107 nack: &rtcp::transport_feedbacks::transport_layer_nack::TransportLayerNack,
108 ) {
109 let mut seqs_to_retransmit = Vec::new();
111
112 for nack_pair in &nack.nacks {
113 seqs_to_retransmit.push(nack_pair.packet_id);
115
116 for i in 0..16 {
118 if nack_pair.lost_packets & (1 << i) != 0 {
119 let seq = nack_pair.packet_id.wrapping_add(i + 1);
120 seqs_to_retransmit.push(seq);
121 }
122 }
123 }
124
125 let Some(stream) = self.streams.get_mut(&nack.media_ssrc) else {
126 return;
127 };
128
129 for seq in seqs_to_retransmit {
131 let Some(original_packet) = stream.send_buffer.get(seq) else {
132 continue;
133 };
134
135 let packet = if let (Some(ssrc_rtx), Some(pt_rtx)) =
136 (stream.ssrc_rtx, stream.payload_type_rtx)
137 {
138 let original_seq = original_packet.header.sequence_number;
143 let mut rtx_payload = Vec::with_capacity(2 + original_packet.payload.len());
144 rtx_payload.extend_from_slice(&original_seq.to_be_bytes());
145 rtx_payload.extend_from_slice(&original_packet.payload);
146
147 let rtx_seq = stream.rtx_sequence_number;
148 stream.rtx_sequence_number = stream.rtx_sequence_number.wrapping_add(1);
149
150 rtp::Packet {
151 header: rtp::header::Header {
152 version: 2,
158 ssrc: ssrc_rtx,
159 payload_type: pt_rtx,
160 sequence_number: rtx_seq,
161 timestamp: original_packet.header.timestamp,
162 marker: original_packet.header.marker,
163 ..Default::default()
164 },
165 payload: rtx_payload.into(),
166 }
167 } else {
168 original_packet.clone()
170 };
171
172 self.write_queue.push_back(TaggedPacket {
173 now,
174 transport: TransportContext::default(),
175 message: Packet::Rtp(packet),
176 });
177 }
178 }
179}
180
181#[interceptor]
182impl<P: Interceptor> NackResponderInterceptor<P> {
183 #[overrides]
184 fn handle_read(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
185 if let Packet::Rtcp(ref rtcp_packets) = msg.message {
187 for rtcp_packet in rtcp_packets {
188 if let Some(nack) = rtcp_packet
189 .as_any()
190 .downcast_ref::<rtcp::transport_feedbacks::transport_layer_nack::TransportLayerNack>()
191 {
192 self.handle_nack(msg.now, nack);
193 }
194 }
195 }
196
197 self.inner.handle_read(msg)
198 }
199
200 #[overrides]
201 fn handle_write(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
202 if let Packet::Rtp(ref rtp_packet) = msg.message
204 && let Some(stream) = self.streams.get_mut(&rtp_packet.header.ssrc)
205 {
206 stream.send_buffer.add(rtp_packet.clone());
207 }
208
209 self.inner.handle_write(msg)
210 }
211
212 #[overrides]
213 fn poll_write(&mut self) -> Option<Self::Wout> {
214 if let Some(pkt) = self.write_queue.pop_front() {
216 return Some(pkt);
217 }
218 self.inner.poll_write()
219 }
220
221 #[overrides]
222 fn bind_local_stream(&mut self, info: &StreamInfo) {
223 if stream_supports_nack(info)
224 && let Some(send_buffer) = SendBuffer::new(self.size)
225 {
226 self.streams.insert(
227 info.ssrc,
228 LocalStream {
229 send_buffer,
230 ssrc_rtx: info.ssrc_rtx,
231 payload_type_rtx: info.payload_type_rtx,
232 rtx_sequence_number: 0,
233 },
234 );
235 }
236 self.inner.bind_local_stream(info);
237 }
238
239 #[overrides]
240 fn unbind_local_stream(&mut self, info: &StreamInfo) {
241 self.streams.remove(&info.ssrc);
242 self.inner.unbind_local_stream(info);
243 }
244}
245
246#[cfg(test)]
247mod tests {
248 use super::*;
249 use crate::Registry;
250 use crate::stream_info::RTCPFeedback;
251 use sansio::Protocol;
252
253 fn make_rtp_packet(ssrc: u32, seq: u16, payload: &[u8]) -> TaggedPacket {
254 TaggedPacket {
255 now: Instant::now(),
256 transport: Default::default(),
257 message: Packet::Rtp(rtp::Packet {
258 header: rtp::header::Header {
259 ssrc,
260 sequence_number: seq,
261 ..Default::default()
262 },
263 payload: payload.to_vec().into(),
264 }),
265 }
266 }
267
268 fn make_nack_packet(sender_ssrc: u32, media_ssrc: u32, nacks: Vec<(u16, u16)>) -> TaggedPacket {
269 let nack_pairs: Vec<rtcp::transport_feedbacks::transport_layer_nack::NackPair> = nacks
270 .into_iter()
271 .map(|(packet_id, lost_packets)| {
272 rtcp::transport_feedbacks::transport_layer_nack::NackPair {
273 packet_id,
274 lost_packets,
275 }
276 })
277 .collect();
278
279 TaggedPacket {
280 now: Instant::now(),
281 transport: Default::default(),
282 message: Packet::Rtcp(vec![Box::new(
283 rtcp::transport_feedbacks::transport_layer_nack::TransportLayerNack {
284 sender_ssrc,
285 media_ssrc,
286 nacks: nack_pairs,
287 },
288 )]),
289 }
290 }
291
292 #[test]
293 fn test_nack_responder_builder_defaults() {
294 let chain = Registry::new()
295 .with(NackResponderBuilder::default().build())
296 .build();
297
298 assert_eq!(chain.size, 1024);
299 }
300
301 #[test]
302 fn test_nack_responder_builder_custom() {
303 let chain = Registry::new()
304 .with(NackResponderBuilder::new().with_size(2048).build())
305 .build();
306
307 assert_eq!(chain.size, 2048);
308 }
309
310 #[test]
311 fn test_nack_responder_retransmits_packet() {
312 let mut chain = Registry::new()
313 .with(NackResponderBuilder::new().with_size(8).build())
314 .build();
315
316 let info = StreamInfo {
318 ssrc: 12345,
319 clock_rate: 90000,
320 rtcp_feedback: vec![RTCPFeedback {
321 typ: "nack".to_string(),
322 parameter: "".to_string(),
323 }],
324 ..Default::default()
325 };
326 chain.bind_local_stream(&info);
327
328 let now = Instant::now();
329
330 for seq in [10u16, 11, 12, 14, 15] {
332 let mut pkt = make_rtp_packet(12345, seq, &[seq as u8]);
333 pkt.now = now;
334 chain.handle_write(pkt).unwrap();
335 chain.poll_write(); }
337
338 let mut nack = make_nack_packet(999, 12345, vec![(11, 0b1011)]);
341 nack.now = now;
342 chain.handle_read(nack).unwrap();
343
344 let mut retransmitted = Vec::new();
346 while let Some(pkt) = chain.poll_write() {
347 if let Packet::Rtp(rtp) = pkt.message {
348 retransmitted.push(rtp.header.sequence_number);
349 }
350 }
351
352 assert!(retransmitted.contains(&11));
353 assert!(retransmitted.contains(&12));
354 assert!(!retransmitted.contains(&13)); assert!(retransmitted.contains(&15));
356 }
357
358 #[test]
359 fn test_nack_responder_no_retransmit_without_binding() {
360 let mut chain = Registry::new()
361 .with(NackResponderBuilder::new().with_size(8).build())
362 .build();
363
364 let now = Instant::now();
365
366 for seq in [10u16, 11, 12] {
368 let mut pkt = make_rtp_packet(12345, seq, &[seq as u8]);
369 pkt.now = now;
370 chain.handle_write(pkt).unwrap();
371 chain.poll_write();
372 }
373
374 let mut nack = make_nack_packet(999, 12345, vec![(11, 0)]);
376 nack.now = now;
377 chain.handle_read(nack).unwrap();
378
379 assert!(chain.poll_write().is_none());
381 }
382
383 #[test]
384 fn test_nack_responder_no_retransmit_expired_packet() {
385 let mut chain = Registry::new()
386 .with(NackResponderBuilder::new().with_size(8).build())
387 .build();
388
389 let info = StreamInfo {
390 ssrc: 12345,
391 clock_rate: 90000,
392 rtcp_feedback: vec![RTCPFeedback {
393 typ: "nack".to_string(),
394 parameter: "".to_string(),
395 }],
396 ..Default::default()
397 };
398 chain.bind_local_stream(&info);
399
400 let now = Instant::now();
401
402 for seq in 0..16u16 {
404 let mut pkt = make_rtp_packet(12345, seq, &[seq as u8]);
405 pkt.now = now;
406 chain.handle_write(pkt).unwrap();
407 chain.poll_write();
408 }
409
410 let mut nack = make_nack_packet(999, 12345, vec![(0, 0)]);
412 nack.now = now;
413 chain.handle_read(nack).unwrap();
414
415 assert!(chain.poll_write().is_none());
417
418 let mut nack = make_nack_packet(999, 12345, vec![(10, 0)]);
420 nack.now = now;
421 chain.handle_read(nack).unwrap();
422
423 let pkt = chain.poll_write();
424 assert!(pkt.is_some());
425 if let Some(tagged) = pkt
426 && let Packet::Rtp(rtp) = tagged.message
427 {
428 assert_eq!(rtp.header.sequence_number, 10);
429 }
430 }
431
432 #[test]
433 fn test_nack_responder_unbind_removes_stream() {
434 let mut chain = Registry::new()
435 .with(NackResponderBuilder::new().with_size(8).build())
436 .build();
437
438 let info = StreamInfo {
439 ssrc: 12345,
440 clock_rate: 90000,
441 rtcp_feedback: vec![RTCPFeedback {
442 typ: "nack".to_string(),
443 parameter: "".to_string(),
444 }],
445 ..Default::default()
446 };
447
448 chain.bind_local_stream(&info);
449 assert!(chain.streams.contains_key(&12345));
450
451 chain.unbind_local_stream(&info);
452 assert!(!chain.streams.contains_key(&12345));
453 }
454
455 #[test]
456 fn test_nack_responder_no_nack_support() {
457 let mut chain = Registry::new()
458 .with(NackResponderBuilder::new().with_size(8).build())
459 .build();
460
461 let info = StreamInfo {
463 ssrc: 12345,
464 clock_rate: 90000,
465 rtcp_feedback: vec![], ..Default::default()
467 };
468 chain.bind_local_stream(&info);
469
470 assert!(!chain.streams.contains_key(&12345));
472 }
473
474 #[test]
475 fn test_nack_responder_passthrough() {
476 let mut chain = Registry::new()
477 .with(NackResponderBuilder::new().with_size(8).build())
478 .build();
479
480 let now = Instant::now();
481
482 let mut pkt = make_rtp_packet(12345, 1, &[1]);
484 pkt.now = now;
485 chain.handle_write(pkt).unwrap();
486 let out = chain.poll_write();
487 assert!(out.is_some());
488
489 let mut nack = make_nack_packet(999, 12345, vec![(1, 0)]);
491 nack.now = now;
492 chain.handle_read(nack).unwrap();
493 let out = chain.poll_read();
494 assert!(out.is_none());
495 }
496
497 #[test]
498 fn test_nack_responder_rfc4588_rtx() {
499 let mut chain = Registry::new()
500 .with(NackResponderBuilder::new().with_size(8).build())
501 .build();
502
503 let info = StreamInfo {
505 ssrc: 1,
506 ssrc_rtx: Some(2), payload_type: 96,
508 payload_type_rtx: Some(97), clock_rate: 90000,
510 rtcp_feedback: vec![RTCPFeedback {
511 typ: "nack".to_string(),
512 parameter: "".to_string(),
513 }],
514 ..Default::default()
515 };
516 chain.bind_local_stream(&info);
517
518 let now = Instant::now();
519
520 for seq in [10u16, 11, 12, 14, 15] {
522 let mut pkt = make_rtp_packet(1, seq, &[seq as u8]);
523 pkt.now = now;
524 chain.handle_write(pkt).unwrap();
525 chain.poll_write(); }
527
528 let mut nack = make_nack_packet(999, 1, vec![(11, 0b1011)]);
531 nack.now = now;
532 chain.handle_read(nack).unwrap();
533
534 let mut rtx_seq = 0u16;
536 for expected_original_seq in [11u16, 12, 15] {
537 let pkt = chain.poll_write();
538 assert!(
539 pkt.is_some(),
540 "Expected RTX packet for seq {}",
541 expected_original_seq
542 );
543
544 if let Some(tagged) = pkt {
545 if let Packet::Rtp(rtp) = tagged.message {
546 assert_eq!(rtp.header.ssrc, 2, "RTX packet should use RTX SSRC");
548 assert_eq!(
550 rtp.header.payload_type, 97,
551 "RTX packet should use RTX payload type"
552 );
553 assert_eq!(
555 rtp.header.sequence_number, rtx_seq,
556 "RTX seq should be {}",
557 rtx_seq
558 );
559 rtx_seq += 1;
560
561 assert!(
563 rtp.payload.len() >= 2,
564 "RTX payload should have at least 2 bytes"
565 );
566 let original_seq_from_payload =
567 u16::from_be_bytes([rtp.payload[0], rtp.payload[1]]);
568 assert_eq!(
569 original_seq_from_payload, expected_original_seq,
570 "RTX payload should contain original seq"
571 );
572
573 assert_eq!(
575 rtp.payload[2..],
576 [expected_original_seq as u8],
577 "Original payload should follow seq number"
578 );
579 } else {
580 panic!("Expected RTP packet");
581 }
582 }
583 }
584
585 assert!(chain.poll_write().is_none());
587 }
588}