1 use std::sync::Mutex;
2 
3 use super::*;
4 
5 use crate::nack::UINT16SIZE_HALF;
6 
7 use util::Unmarshal;
8 
9 struct GeneratorStreamInternal {
10     packets: Vec<u64>,
11     size: u16,
12     end: u16,
13     started: bool,
14     last_consecutive: u16,
15 }
16 
17 impl GeneratorStreamInternal {
18     fn new(log2_size_minus_6: u8) -> Self {
19         GeneratorStreamInternal {
20             packets: vec![0u64; 1 << log2_size_minus_6],
21             size: 1 << (log2_size_minus_6 + 6),
22             end: 0,
23             started: false,
24             last_consecutive: 0,
25         }
26     }
27 
28     fn add(&mut self, seq: u16) {
29         if !self.started {
30             self.set_received(seq);
31             self.end = seq;
32             self.started = true;
33             self.last_consecutive = seq;
34             return;
35         }
36 
37         let last_consecutive_plus1 = self.last_consecutive.wrapping_add(1);
38         let diff = seq.wrapping_sub(self.end);
39         if diff == 0 {
40             return;
41         } else if diff < UINT16SIZE_HALF {
42             // this means a positive diff, in other words seq > end (with counting for rollovers)
43             let mut i = self.end.wrapping_add(1);
44             while i != seq {
45                 // clear packets between end and seq (these may contain packets from a "size" ago)
46                 self.del_received(i);
47                 i = i.wrapping_add(1);
48             }
49             self.end = seq;
50 
51             let seq_sub_last_consecutive = seq.wrapping_sub(self.last_consecutive);
52             if last_consecutive_plus1 == seq {
53                 self.last_consecutive = seq;
54             } else if seq_sub_last_consecutive > self.size {
55                 let diff = seq.wrapping_sub(self.size);
56                 self.last_consecutive = diff;
57                 self.fix_last_consecutive(); // there might be valid packets at the beginning of the buffer now
58             }
59         } else if last_consecutive_plus1 == seq {
60             // negative diff, seq < end (with counting for rollovers)
61             self.last_consecutive = seq;
62             self.fix_last_consecutive(); // there might be other valid packets after seq
63         }
64 
65         self.set_received(seq);
66     }
67 
68     fn get(&self, seq: u16) -> bool {
69         let diff = self.end.wrapping_sub(seq);
70         if diff >= UINT16SIZE_HALF {
71             return false;
72         }
73 
74         if diff >= self.size {
75             return false;
76         }
77 
78         self.get_received(seq)
79     }
80 
81     fn missing_seq_numbers(&self, skip_last_n: u16) -> Vec<u16> {
82         let until = self.end.wrapping_sub(skip_last_n);
83         let diff = until.wrapping_sub(self.last_consecutive);
84         if diff >= UINT16SIZE_HALF {
85             // until < s.last_consecutive (counting for rollover)
86             return vec![];
87         }
88 
89         let mut missing_packet_seq_nums = vec![];
90         let mut i = self.last_consecutive.wrapping_add(1);
91         let util_plus1 = until.wrapping_add(1);
92         while i != util_plus1 {
93             if !self.get_received(i) {
94                 missing_packet_seq_nums.push(i);
95             }
96             i = i.wrapping_add(1);
97         }
98 
99         missing_packet_seq_nums
100     }
101 
102     fn set_received(&mut self, seq: u16) {
103         let pos = (seq % self.size) as usize;
104         self.packets[pos / 64] |= 1u64 << (pos % 64);
105     }
106 
107     fn del_received(&mut self, seq: u16) {
108         let pos = (seq % self.size) as usize;
109         self.packets[pos / 64] &= u64::MAX ^ (1u64 << (pos % 64));
110     }
111 
112     fn get_received(&self, seq: u16) -> bool {
113         let pos = (seq % self.size) as usize;
114         (self.packets[pos / 64] & (1u64 << (pos % 64))) != 0
115     }
116 
117     fn fix_last_consecutive(&mut self) {
118         let mut i = self.last_consecutive.wrapping_add(1);
119         while i != self.end.wrapping_add(1) && self.get_received(i) {
120             // find all consecutive packets
121             i = i.wrapping_add(1);
122         }
123         self.last_consecutive = i.wrapping_sub(1);
124     }
125 }
126 
127 pub(super) struct GeneratorStream {
128     parent_rtp_reader: Arc<dyn RTPReader + Send + Sync>,
129 
130     internal: Mutex<GeneratorStreamInternal>,
131 }
132 
133 impl GeneratorStream {
134     pub(super) fn new(log2_size_minus_6: u8, reader: Arc<dyn RTPReader + Send + Sync>) -> Self {
135         GeneratorStream {
136             parent_rtp_reader: reader,
137             internal: Mutex::new(GeneratorStreamInternal::new(log2_size_minus_6)),
138         }
139     }
140 
141     pub(super) fn missing_seq_numbers(&self, skip_last_n: u16) -> Vec<u16> {
142         let internal = self.internal.lock().unwrap();
143         internal.missing_seq_numbers(skip_last_n)
144     }
145 
146     pub(super) fn add(&self, seq: u16) {
147         let mut internal = self.internal.lock().unwrap();
148         internal.add(seq);
149     }
150 }
151 
152 /// RTPReader is used by Interceptor.bind_remote_stream.
153 #[async_trait]
154 impl RTPReader for GeneratorStream {
155     /// read a rtp packet
156     async fn read(&self, buf: &mut [u8], a: &Attributes) -> Result<(usize, Attributes)> {
157         let (n, attr) = self.parent_rtp_reader.read(buf, a).await?;
158 
159         let mut b = &buf[..n];
160         let pkt = rtp::packet::Packet::unmarshal(&mut b)?;
161         self.add(pkt.header.sequence_number);
162 
163         Ok((n, attr))
164     }
165 }
166 
167 #[cfg(test)]
168 mod test {
169     use super::*;
170 
171     #[test]
172     fn test_generator_stream() -> Result<()> {
173         let tests: Vec<u16> = vec![
174             0, 1, 127, 128, 129, 511, 512, 513, 32767, 32768, 32769, 65407, 65408, 65409, 65534,
175             65535,
176         ];
177         for start in tests {
178             let mut rl = GeneratorStreamInternal::new(1);
179 
180             let all = |min: u16, max: u16| -> Vec<u16> {
181                 let mut result = vec![];
182                 let mut i = min;
183                 let max_plus_1 = max.wrapping_add(1);
184                 while i != max_plus_1 {
185                     result.push(i);
186                     i = i.wrapping_add(1);
187                 }
188                 result
189             };
190 
191             let join = |parts: &[&[u16]]| -> Vec<u16> {
192                 let mut result = vec![];
193                 for p in parts {
194                     result.extend_from_slice(*p);
195                 }
196                 result
197             };
198 
199             let add = |rl: &mut GeneratorStreamInternal, nums: &[u16]| {
200                 for n in nums {
201                     let seq = start.wrapping_add(*n);
202                     rl.add(seq);
203                 }
204             };
205 
206             let assert_get = |rl: &GeneratorStreamInternal, nums: &[u16]| {
207                 for n in nums {
208                     let seq = start.wrapping_add(*n);
209                     assert!(rl.get(seq), "not found: {}", seq);
210                 }
211             };
212 
213             let assert_not_get = |rl: &GeneratorStreamInternal, nums: &[u16]| {
214                 for n in nums {
215                     let seq = start.wrapping_add(*n);
216                     assert!(
217                         !rl.get(seq),
218                         "packet found: start {}, n {}, seq {}",
219                         start,
220                         *n,
221                         seq
222                     );
223                 }
224             };
225 
226             let assert_missing = |rl: &GeneratorStreamInternal, skip_last_n: u16, nums: &[u16]| {
227                 let missing = rl.missing_seq_numbers(skip_last_n);
228                 let mut want = vec![];
229                 for n in nums {
230                     let seq = start.wrapping_add(*n);
231                     want.push(seq);
232                 }
233                 assert_eq!(want, missing, "missing want/got, ");
234             };
235 
236             let assert_last_consecutive = |rl: &GeneratorStreamInternal, last_consecutive: u16| {
237                 let want = last_consecutive.wrapping_add(start);
238                 assert_eq!(rl.last_consecutive, want, "invalid last_consecutive want");
239             };
240 
241             add(&mut rl, &[0]);
242             assert_get(&rl, &[0]);
243             assert_missing(&rl, 0, &[]);
244             assert_last_consecutive(&rl, 0); // first element added
245 
246             add(&mut rl, &all(1, 127));
247             assert_get(&rl, &all(1, 127));
248             assert_missing(&rl, 0, &[]);
249             assert_last_consecutive(&rl, 127);
250 
251             add(&mut rl, &[128]);
252             assert_get(&rl, &[128]);
253             assert_not_get(&rl, &[0]);
254             assert_missing(&rl, 0, &[]);
255             assert_last_consecutive(&rl, 128);
256 
257             add(&mut rl, &[130]);
258             assert_get(&rl, &[130]);
259             assert_not_get(&rl, &[1, 2, 129]);
260             assert_missing(&rl, 0, &[129]);
261             assert_last_consecutive(&rl, 128);
262 
263             add(&mut rl, &[333]);
264             assert_get(&rl, &[333]);
265             assert_not_get(&rl, &all(0, 332));
266             assert_missing(&rl, 0, &all(206, 332)); // all 127 elements missing before 333
267             assert_missing(&rl, 10, &all(206, 323)); // skip last 10 packets (324-333) from check
268             assert_last_consecutive(&rl, 205); // lastConsecutive is still out of the buffer
269 
270             add(&mut rl, &[329]);
271             assert_get(&rl, &[329]);
272             assert_missing(&rl, 0, &join(&[&all(206, 328), &all(330, 332)]));
273             assert_missing(&rl, 5, &join(&[&all(206, 328)])); // skip last 5 packets (329-333) from check
274             assert_last_consecutive(&rl, 205);
275 
276             add(&mut rl, &all(207, 320));
277             assert_get(&rl, &all(207, 320));
278             assert_missing(&rl, 0, &join(&[&[206], &all(321, 328), &all(330, 332)]));
279             assert_last_consecutive(&rl, 205);
280 
281             add(&mut rl, &[334]);
282             assert_get(&rl, &[334]);
283             assert_not_get(&rl, &[206]);
284             assert_missing(&rl, 0, &join(&[&all(321, 328), &all(330, 332)]));
285             assert_last_consecutive(&rl, 320); // head of buffer is full of consecutive packages
286 
287             add(&mut rl, &all(322, 328));
288             assert_get(&rl, &all(322, 328));
289             assert_missing(&rl, 0, &join(&[&[321], &all(330, 332)]));
290             assert_last_consecutive(&rl, 320);
291 
292             add(&mut rl, &[321]);
293             assert_get(&rl, &[321]);
294             assert_missing(&rl, 0, &all(330, 332));
295             assert_last_consecutive(&rl, 329); // after adding a single missing packet, lastConsecutive should jump forward
296         }
297 
298         Ok(())
299     }
300 
301     #[test]
302     fn test_generator_stream_rollover() {
303         let mut rl = GeneratorStreamInternal::new(1);
304         // Make sure it doesn't panic.
305         rl.add(65533);
306         rl.add(65535);
307         rl.add(65534);
308 
309         let mut rl = GeneratorStreamInternal::new(1);
310         // Make sure it doesn't panic.
311         rl.add(65534);
312         rl.add(0);
313         rl.add(65535);
314     }
315 }
316