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