1 use super::*;
2 use crate::error::Result;
3 use crate::protection_profile::*;
4 
5 use rtcp::payload_feedbacks::*;
6 use util::conn::conn_pipe::*;
7 
8 use bytes::{Bytes, BytesMut};
9 use std::sync::Arc;
10 use tokio::sync::{mpsc, Mutex};
11 
12 async fn build_session_srtcp_pair() -> Result<(Session, Session)> {
13     let (ua, ub) = pipe();
14 
15     let ca = Config {
16         profile: ProtectionProfile::Aes128CmHmacSha1_80,
17         keys: SessionKeys {
18             local_master_key: vec![
19                 0xE1, 0xF9, 0x7A, 0x0D, 0x3E, 0x01, 0x8B, 0xE0, 0xD6, 0x4F, 0xA3, 0x2C, 0x06, 0xDE,
20                 0x41, 0x39,
21             ],
22             local_master_salt: vec![
23                 0x0E, 0xC6, 0x75, 0xAD, 0x49, 0x8A, 0xFE, 0xEB, 0xB6, 0x96, 0x0B, 0x3A, 0xAB, 0xE6,
24             ],
25             remote_master_key: vec![
26                 0xE1, 0xF9, 0x7A, 0x0D, 0x3E, 0x01, 0x8B, 0xE0, 0xD6, 0x4F, 0xA3, 0x2C, 0x06, 0xDE,
27                 0x41, 0x39,
28             ],
29             remote_master_salt: vec![
30                 0x0E, 0xC6, 0x75, 0xAD, 0x49, 0x8A, 0xFE, 0xEB, 0xB6, 0x96, 0x0B, 0x3A, 0xAB, 0xE6,
31             ],
32         },
33 
34         local_rtp_options: None,
35         remote_rtp_options: None,
36 
37         local_rtcp_options: None,
38         remote_rtcp_options: None,
39     };
40 
41     let cb = Config {
42         profile: ProtectionProfile::Aes128CmHmacSha1_80,
43         keys: SessionKeys {
44             local_master_key: vec![
45                 0xE1, 0xF9, 0x7A, 0x0D, 0x3E, 0x01, 0x8B, 0xE0, 0xD6, 0x4F, 0xA3, 0x2C, 0x06, 0xDE,
46                 0x41, 0x39,
47             ],
48             local_master_salt: vec![
49                 0x0E, 0xC6, 0x75, 0xAD, 0x49, 0x8A, 0xFE, 0xEB, 0xB6, 0x96, 0x0B, 0x3A, 0xAB, 0xE6,
50             ],
51             remote_master_key: vec![
52                 0xE1, 0xF9, 0x7A, 0x0D, 0x3E, 0x01, 0x8B, 0xE0, 0xD6, 0x4F, 0xA3, 0x2C, 0x06, 0xDE,
53                 0x41, 0x39,
54             ],
55             remote_master_salt: vec![
56                 0x0E, 0xC6, 0x75, 0xAD, 0x49, 0x8A, 0xFE, 0xEB, 0xB6, 0x96, 0x0B, 0x3A, 0xAB, 0xE6,
57             ],
58         },
59 
60         local_rtp_options: None,
61         remote_rtp_options: None,
62 
63         local_rtcp_options: None,
64         remote_rtcp_options: None,
65     };
66 
67     let sa = Session::new(Arc::new(ua), ca, false).await?;
68     let sb = Session::new(Arc::new(ub), cb, false).await?;
69 
70     Ok((sa, sb))
71 }
72 
73 const TEST_SSRC: u32 = 5000;
74 
75 #[tokio::test]
76 async fn test_session_srtcp_accept() -> Result<()> {
77     let (sa, sb) = build_session_srtcp_pair().await?;
78 
79     let rtcp_packet = picture_loss_indication::PictureLossIndication {
80         media_ssrc: TEST_SSRC,
81         ..Default::default()
82     };
83 
84     let test_payload = rtcp_packet.marshal()?;
85     sa.write_rtcp(&rtcp_packet).await?;
86 
87     let read_stream = sb.accept().await?;
88     let ssrc = read_stream.get_ssrc();
89     assert_eq!(
90         ssrc, TEST_SSRC,
91         "SSRC mismatch during accept exp({}) actual({})",
92         TEST_SSRC, ssrc
93     );
94 
95     let mut read_buffer = BytesMut::with_capacity(test_payload.len());
96     read_buffer.resize(test_payload.len(), 0u8);
97     read_stream.read(&mut read_buffer).await?;
98 
99     assert_eq!(
100         &test_payload[..],
101         &read_buffer[..],
102         "Sent buffer does not match the one received exp({:?}) actual({:?})",
103         &test_payload[..],
104         &read_buffer[..]
105     );
106 
107     sa.close().await?;
108     sb.close().await?;
109 
110     Ok(())
111 }
112 
113 #[tokio::test]
114 async fn test_session_srtcp_listen() -> Result<()> {
115     let (sa, sb) = build_session_srtcp_pair().await?;
116 
117     let rtcp_packet = picture_loss_indication::PictureLossIndication {
118         media_ssrc: TEST_SSRC,
119         ..Default::default()
120     };
121 
122     let test_payload = rtcp_packet.marshal()?;
123     let read_stream = sb.open(TEST_SSRC).await;
124 
125     sa.write_rtcp(&rtcp_packet).await?;
126 
127     let mut read_buffer = BytesMut::with_capacity(test_payload.len());
128     read_buffer.resize(test_payload.len(), 0u8);
129     read_stream.read(&mut read_buffer).await?;
130 
131     assert_eq!(
132         &test_payload[..],
133         &read_buffer[..],
134         "Sent buffer does not match the one received exp({:?}) actual({:?})",
135         &test_payload[..],
136         &read_buffer[..]
137     );
138 
139     sa.close().await?;
140     sb.close().await?;
141 
142     Ok(())
143 }
144 
145 fn encrypt_srtcp(
146     context: &mut Context,
147     pkt: &(dyn rtcp::packet::Packet + Send + Sync),
148 ) -> Result<Bytes> {
149     let decrypted = pkt.marshal()?;
150     let encrypted = context.encrypt_rtcp(&decrypted)?;
151     Ok(encrypted)
152 }
153 
154 const PLI_PACKET_SIZE: usize = 8;
155 
156 async fn get_sender_ssrc(read_stream: &Arc<Stream>) -> Result<u32> {
157     let auth_tag_size = ProtectionProfile::Aes128CmHmacSha1_80.auth_tag_len();
158 
159     let mut read_buffer = BytesMut::with_capacity(PLI_PACKET_SIZE + auth_tag_size);
160     read_buffer.resize(PLI_PACKET_SIZE + auth_tag_size, 0u8);
161 
162     let (n, _) = read_stream.read_rtcp(&mut read_buffer).await?;
163     let mut reader = &read_buffer[0..n];
164     let pli = picture_loss_indication::PictureLossIndication::unmarshal(&mut reader)?;
165 
166     Ok(pli.sender_ssrc)
167 }
168 
169 #[tokio::test]
170 async fn test_session_srtcp_replay_protection() -> Result<()> {
171     let (sa, sb) = build_session_srtcp_pair().await?;
172 
173     let read_stream = sb.open(TEST_SSRC).await;
174 
175     // Generate test packets
176     let mut packets = vec![];
177     let mut expected_ssrc = vec![];
178     {
179         let mut local_context = sa.local_context.lock().await;
180         for i in 0..0x10u32 {
181             expected_ssrc.push(i);
182 
183             let packet = picture_loss_indication::PictureLossIndication {
184                 media_ssrc: TEST_SSRC,
185                 sender_ssrc: i,
186             };
187 
188             let encrypted = encrypt_srtcp(&mut local_context, &packet)?;
189 
190             packets.push(encrypted);
191         }
192     }
193 
194     let (done_tx, mut done_rx) = mpsc::channel::<()>(1);
195 
196     let received_ssrc = Arc::new(Mutex::new(vec![]));
197     let cloned_received_ssrc = Arc::clone(&received_ssrc);
198     let count = expected_ssrc.len();
199 
200     tokio::spawn(async move {
201         let mut i = 0;
202         while i < count {
203             match get_sender_ssrc(&read_stream).await {
204                 Ok(ssrc) => {
205                     let mut r = cloned_received_ssrc.lock().await;
206                     r.push(ssrc);
207 
208                     i += 1;
209                 }
210                 Err(_) => break,
211             }
212         }
213 
214         drop(done_tx);
215     });
216 
217     // Write with replay attack
218     for packet in &packets {
219         sa.udp_tx.send(packet).await?;
220 
221         // Immediately replay
222         sa.udp_tx.send(packet).await?;
223     }
224     for packet in &packets {
225         // Delayed replay
226         sa.udp_tx.send(packet).await?;
227     }
228 
229     done_rx.recv().await;
230 
231     sa.close().await?;
232     sb.close().await?;
233 
234     {
235         let received_ssrc = received_ssrc.lock().await;
236         assert_eq!(&expected_ssrc[..], &received_ssrc[..]);
237     }
238 
239     Ok(())
240 }
241