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