rtc_interceptor/nack/
generator.rs1use super::receive_log::ReceiveLog;
4use super::stream_supports_nack;
5use crate::Interceptor;
6use crate::stream_info::StreamInfo;
7use crate::{AttributedPacket, Packet, TaggedPacket};
8use sansio::Protocol;
9use shared::TransportContext;
10use shared::error::Error;
11use std::collections::{HashMap, VecDeque};
12use std::time::{Duration, Instant};
13
14pub struct NackGeneratorBuilder {
31 size: u16,
33 interval: Duration,
35 skip_last_n: u16,
37 max_nacks_per_packet: u16,
39}
40
41impl Default for NackGeneratorBuilder {
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 }
49 }
50}
51
52impl NackGeneratorBuilder {
53 pub fn new() -> Self {
55 Self::default()
56 }
57
58 pub fn with_size(mut self, size: u16) -> Self {
62 self.size = size;
63 self
64 }
65
66 pub fn with_interval(mut self, interval: Duration) -> Self {
68 self.interval = interval;
69 self
70 }
71
72 pub fn with_skip_last_n(mut self, skip_last_n: u16) -> Self {
77 self.skip_last_n = skip_last_n;
78 self
79 }
80
81 pub fn with_max_nacks_per_packet(mut self, max: u16) -> Self {
85 self.max_nacks_per_packet = max;
86 self
87 }
88
89 pub fn build(self) -> NackGeneratorInterceptor {
91 NackGeneratorInterceptor::new(
92 self.size,
93 self.interval,
94 self.skip_last_n,
95 self.max_nacks_per_packet,
96 )
97 }
98}
99
100pub struct NackGeneratorInterceptor {
106 size: u16,
108 interval: Duration,
109 skip_last_n: u16,
110 max_nacks_per_packet: u16,
111
112 next_timeout: Option<Instant>,
114
115 sender_ssrc: u32,
117
118 receive_logs: HashMap<u32, ReceiveLog>,
120
121 nack_counts: HashMap<u32, HashMap<u16, u16>>,
123
124 write_queue: VecDeque<TaggedPacket>,
126 read_queue: VecDeque<TaggedPacket>,
128}
129
130impl NackGeneratorInterceptor {
131 fn new(size: u16, interval: Duration, skip_last_n: u16, max_nacks_per_packet: u16) -> Self {
132 Self {
133 read_queue: VecDeque::new(),
134 size,
135 interval,
136 skip_last_n,
137 max_nacks_per_packet,
138 next_timeout: None,
139 sender_ssrc: rand::random(),
140 receive_logs: HashMap::new(),
141 nack_counts: HashMap::new(),
142 write_queue: VecDeque::new(),
143 }
144 }
145
146 fn generate_nacks(&mut self, now: Instant) {
148 for (&ssrc, receive_log) in &self.receive_logs {
149 let missing = receive_log.missing_seq_numbers(self.skip_last_n);
150 if missing.is_empty() {
151 self.nack_counts.remove(&ssrc);
153 continue;
154 }
155
156 let nack_count = self.nack_counts.entry(ssrc).or_default();
158
159 let filtered: Vec<u16> = if self.max_nacks_per_packet > 0 {
161 missing
162 .iter()
163 .filter(|&&seq| {
164 let count = nack_count.entry(seq).or_insert(0);
165 if *count < self.max_nacks_per_packet {
166 *count += 1;
167 true
168 } else {
169 false
170 }
171 })
172 .copied()
173 .collect()
174 } else {
175 missing.clone()
176 };
177
178 if filtered.is_empty() {
179 continue;
180 }
181
182 nack_count.retain(|seq, _| missing.contains(seq));
184
185 let nack = rtcp::transport_feedbacks::transport_layer_nack::TransportLayerNack {
187 sender_ssrc: self.sender_ssrc,
188 media_ssrc: ssrc,
189 nacks: rtcp::transport_feedbacks::transport_layer_nack::nack_pairs_from_sequence_numbers(
190 &filtered,
191 ),
192 };
193
194 self.write_queue.push_back(TaggedPacket {
195 now,
196 transport: TransportContext::default(),
197 message: AttributedPacket::new(Packet::Rtcp(vec![Box::new(nack)])),
198 });
199 }
200 }
201}
202
203impl Protocol<TaggedPacket, TaggedPacket, ()> for NackGeneratorInterceptor {
204 type Rout = TaggedPacket;
205 type Wout = TaggedPacket;
206 type Eout = ();
207 type Error = Error;
208 type Time = Instant;
209
210 fn handle_read(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
211 if let Packet::Rtp(ref rtp_packet) = msg.message.packet
213 && let Some(receive_log) = self.receive_logs.get_mut(&rtp_packet.header.ssrc)
214 {
215 receive_log.add(rtp_packet.header.sequence_number);
216
217 if self.next_timeout.is_none() {
220 self.next_timeout = Some(msg.now + self.interval);
221 }
222 }
223
224 self.read_queue.push_back(msg);
225
226 Ok(())
227 }
228
229 fn poll_read(&mut self) -> Option<Self::Rout> {
230 self.read_queue.pop_front()
231 }
232
233 fn handle_write(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
234 self.write_queue.push_back(msg);
235 Ok(())
236 }
237
238 fn poll_write(&mut self) -> Option<TaggedPacket> {
239 if let Some(pkt) = self.write_queue.pop_front() {
241 return Some(pkt);
242 }
243 None
244 }
245
246 fn handle_timeout(&mut self, now: Instant) -> Result<(), Error> {
247 if let Some(next_timeout) = self.next_timeout
248 && now >= next_timeout
249 {
250 self.next_timeout = Some(now + self.interval);
251 self.generate_nacks(now);
252 }
253 Ok(())
254 }
255
256 fn poll_timeout(&mut self) -> Option<Instant> {
257 self.next_timeout
258 }
259}
260
261impl Interceptor for NackGeneratorInterceptor {
262 fn bind_remote_stream(&mut self, info: &StreamInfo) {
263 if stream_supports_nack(info)
264 && let Some(receive_log) = ReceiveLog::new(self.size)
265 {
266 self.receive_logs.insert(info.ssrc, receive_log);
267 }
268 }
269
270 fn unbind_remote_stream(&mut self, info: &StreamInfo) {
271 self.receive_logs.remove(&info.ssrc);
272 self.nack_counts.remove(&info.ssrc);
273 }
274
275 fn bind_local_stream(&mut self, _info: &StreamInfo) {}
276
277 fn unbind_local_stream(&mut self, _info: &StreamInfo) {}
278}