xref: /webrtc/sctp/src/stream/stream_test.rs (revision 7ceeeeb0)
1 use super::*;
2 use std::sync::atomic::{AtomicU32, Ordering};
3 use std::sync::Arc;
4 use tokio::io::AsyncReadExt;
5 use tokio::io::AsyncWriteExt;
6 
7 #[test]
8 fn test_stream_buffered_amount() -> Result<()> {
9     let s = Stream::default();
10 
11     assert_eq!(0, s.buffered_amount());
12     assert_eq!(0, s.buffered_amount_low_threshold());
13 
14     s.buffered_amount.store(8192, Ordering::SeqCst);
15     s.set_buffered_amount_low_threshold(2048);
16     assert_eq!(8192, s.buffered_amount(), "unexpected bufferedAmount");
17     assert_eq!(
18         2048,
19         s.buffered_amount_low_threshold(),
20         "unexpected threshold"
21     );
22 
23     Ok(())
24 }
25 
26 #[tokio::test]
27 async fn test_stream_amount_on_buffered_amount_low() -> Result<()> {
28     let s = Stream::default();
29 
30     s.buffered_amount.store(4096, Ordering::SeqCst);
31     s.set_buffered_amount_low_threshold(2048);
32 
33     let n_cbs = Arc::new(AtomicU32::new(0));
34     let n_cbs2 = n_cbs.clone();
35 
36     s.on_buffered_amount_low(Box::new(move || {
37         n_cbs2.fetch_add(1, Ordering::SeqCst);
38         Box::pin(async {})
39     }))
40     .await;
41 
42     // Negative value should be ignored (by design)
43     s.on_buffer_released(-32).await; // bufferedAmount = 3072
44     assert_eq!(4096, s.buffered_amount(), "unexpected bufferedAmount");
45     assert_eq!(0, n_cbs.load(Ordering::SeqCst), "callback count mismatch");
46 
47     // Above to above, no callback
48     s.on_buffer_released(1024).await; // bufferedAmount = 3072
49     assert_eq!(3072, s.buffered_amount(), "unexpected bufferedAmount");
50     assert_eq!(0, n_cbs.load(Ordering::SeqCst), "callback count mismatch");
51 
52     // Above to equal, callback should be made
53     s.on_buffer_released(1024).await; // bufferedAmount = 2048
54     assert_eq!(2048, s.buffered_amount(), "unexpected bufferedAmount");
55     assert_eq!(1, n_cbs.load(Ordering::SeqCst), "callback count mismatch");
56 
57     // Eaual to below, no callback
58     s.on_buffer_released(1024).await; // bufferedAmount = 1024
59     assert_eq!(1024, s.buffered_amount(), "unexpected bufferedAmount");
60     assert_eq!(1, n_cbs.load(Ordering::SeqCst), "callback count mismatch");
61 
62     // Blow to below, no callback
63     s.on_buffer_released(1024).await; // bufferedAmount = 0
64     assert_eq!(0, s.buffered_amount(), "unexpected bufferedAmount");
65     assert_eq!(1, n_cbs.load(Ordering::SeqCst), "callback count mismatch");
66 
67     // Capped at 0, no callback
68     s.on_buffer_released(1024).await; // bufferedAmount = 0
69     assert_eq!(0, s.buffered_amount(), "unexpected bufferedAmount");
70     assert_eq!(1, n_cbs.load(Ordering::SeqCst), "callback count mismatch");
71 
72     Ok(())
73 }
74 
75 #[tokio::test]
76 async fn test_stream() -> std::result::Result<(), io::Error> {
77     let s = Stream::new(
78         "test_poll_stream".to_owned(),
79         0,
80         4096,
81         Arc::new(AtomicU32::new(4096)),
82         Arc::new(AtomicU8::new(AssociationState::Established as u8)),
83         None,
84         Arc::new(PendingQueue::new()),
85     );
86 
87     // getters
88     assert_eq!(0, s.stream_identifier());
89     assert_eq!(0, s.buffered_amount());
90     assert_eq!(0, s.buffered_amount_low_threshold());
91     assert_eq!(0, s.get_num_bytes_in_reassembly_queue().await);
92 
93     // setters
94     s.set_default_payload_type(PayloadProtocolIdentifier::Binary);
95     s.set_reliability_params(true, ReliabilityType::Reliable, 0);
96 
97     // write
98     let n = s.write(&Bytes::from("Hello ")).await?;
99     assert_eq!(6, n);
100     assert_eq!(6, s.buffered_amount());
101     let n = s
102         .write_sctp(&Bytes::from("world"), PayloadProtocolIdentifier::Binary)
103         .await?;
104     assert_eq!(5, n);
105     assert_eq!(11, s.buffered_amount());
106 
107     // async read
108     //  1. pretend that we've received a chunk
109     s.handle_data(ChunkPayloadData {
110         unordered: true,
111         beginning_fragment: true,
112         ending_fragment: true,
113         user_data: Bytes::from_static(&[0, 1, 2, 3, 4]),
114         payload_type: PayloadProtocolIdentifier::Binary,
115         ..Default::default()
116     })
117     .await;
118     //  2. read it
119     let mut buf = [0; 5];
120     s.read(&mut buf).await?;
121     assert_eq!(buf, [0, 1, 2, 3, 4]);
122 
123     // shutdown write
124     s.shutdown(Shutdown::Write).await?;
125     // write must fail
126     assert!(s.write(&Bytes::from("error")).await.is_err());
127     // read should continue working
128     s.handle_data(ChunkPayloadData {
129         unordered: true,
130         beginning_fragment: true,
131         ending_fragment: true,
132         user_data: Bytes::from_static(&[5, 6, 7, 8, 9]),
133         payload_type: PayloadProtocolIdentifier::Binary,
134         ..Default::default()
135     })
136     .await;
137     let mut buf = [0; 5];
138     s.read(&mut buf).await?;
139     assert_eq!(buf, [5, 6, 7, 8, 9]);
140 
141     // shutdown read
142     s.shutdown(Shutdown::Read).await?;
143     // read must return 0
144     assert_eq!(Ok(0), s.read(&mut buf).await);
145 
146     Ok(())
147 }
148 
149 #[tokio::test]
150 async fn test_poll_stream() -> std::result::Result<(), io::Error> {
151     let s = Arc::new(Stream::new(
152         "test_poll_stream".to_owned(),
153         0,
154         4096,
155         Arc::new(AtomicU32::new(4096)),
156         Arc::new(AtomicU8::new(AssociationState::Established as u8)),
157         None,
158         Arc::new(PendingQueue::new()),
159     ));
160     let mut poll_stream = PollStream::new(s.clone());
161 
162     // getters
163     assert_eq!(0, poll_stream.stream_identifier());
164     assert_eq!(0, poll_stream.buffered_amount());
165     assert_eq!(0, poll_stream.buffered_amount_low_threshold());
166     assert_eq!(0, poll_stream.get_num_bytes_in_reassembly_queue().await);
167 
168     // async write
169     let n = poll_stream.write(&[1, 2, 3]).await?;
170     assert_eq!(3, n);
171     poll_stream.flush().await?;
172     assert_eq!(3, poll_stream.buffered_amount());
173 
174     // async read
175     //  1. pretend that we've received a chunk
176     let sc = s.clone();
177     sc.handle_data(ChunkPayloadData {
178         unordered: true,
179         beginning_fragment: true,
180         ending_fragment: true,
181         user_data: Bytes::from_static(&[0, 1, 2, 3, 4]),
182         payload_type: PayloadProtocolIdentifier::Binary,
183         ..Default::default()
184     })
185     .await;
186     //  2. read it
187     let mut buf = [0; 5];
188     poll_stream.read(&mut buf).await?;
189     assert_eq!(buf, [0, 1, 2, 3, 4]);
190 
191     // shutdown write
192     poll_stream.shutdown().await?;
193     // write must fail
194     assert!(poll_stream.write(&[1, 2, 3]).await.is_err());
195     // read should continue working
196     sc.handle_data(ChunkPayloadData {
197         unordered: true,
198         beginning_fragment: true,
199         ending_fragment: true,
200         user_data: Bytes::from_static(&[5, 6, 7, 8, 9]),
201         payload_type: PayloadProtocolIdentifier::Binary,
202         ..Default::default()
203     })
204     .await;
205     let mut buf = [0; 5];
206     poll_stream.read(&mut buf).await?;
207     assert_eq!(buf, [5, 6, 7, 8, 9]);
208 
209     // misc.
210     let clone = poll_stream.clone();
211     assert_eq!(clone.stream_identifier(), poll_stream.stream_identifier());
212 
213     Ok(())
214 }
215