source_wire_tests.rs (10543B)
1 use super::*; 2 use crate::source::budget::{FetchBudget, MAX_FETCH_BYTES}; 3 use crate::source_ingress::{FetchRegistration, IngressRegistry}; 4 use futures::task::{ArcWake, waker}; 5 use std::sync::{ 6 Arc, 7 atomic::{AtomicUsize, Ordering}, 8 }; 9 10 const UPGRADE: &[u8] = b"HTTP/1.1 101 Switching Protocols\r\n\r\n"; 11 const RELAY: &str = "ws://127.0.0.1:1234"; 12 13 fn upgraded() -> WireParser { 14 let mut parser = WireParser::default(); 15 parser 16 .consume(UPGRADE, |_| panic!("HTTP is not a frame")) 17 .expect("upgrade"); 18 parser 19 } 20 21 fn header(length: u64) -> Vec<u8> { 22 let mut bytes = vec![0x82, 127]; 23 bytes.extend(length.to_be_bytes()); 24 bytes 25 } 26 27 #[test] 28 fn maximum_frame_length_is_admitted_and_next_byte_or_overflow_is_denied() { 29 let maximum = crate::relay::MAX_WIRE_MESSAGE_BYTES as u64; 30 let mut parser = upgraded(); 31 parser 32 .consume(&header(maximum), |_| true) 33 .expect("exact limit"); 34 assert_eq!(parser.payload_remaining, maximum as usize); 35 for length in [maximum + 1, 1 << 63, u64::MAX] { 36 assert!(upgraded().consume(&header(length), |_| true).is_err()); 37 } 38 } 39 40 #[test] 41 fn headers_payloads_and_upgrade_can_arrive_one_byte_at_a_time() { 42 let mut parser = WireParser::default(); 43 let mut bytes = UPGRADE.to_vec(); 44 bytes.extend([0x01, 126, 0, 126]); 45 bytes.extend([b'x'; 126]); 46 bytes.extend([0x89, 0, 0x80, 2, b'y', b'z', 0x82, 0]); 47 let mut frames = 0; 48 let mut data = 0; 49 for byte in bytes { 50 parser 51 .consume(&[byte], |is_data| { 52 frames += 1; 53 data += usize::from(is_data); 54 true 55 }) 56 .expect("chopped ingress"); 57 } 58 assert_eq!((frames, data), (4, 2)); 59 parser.eof().expect("complete frames"); 60 } 61 62 #[test] 63 fn zero_length_continuations_and_controls_each_require_admission() { 64 let mut parser = upgraded(); 65 let mut admissions = 0; 66 let error = parser 67 .consume(&[0x01, 0, 0x00, 0, 0x89, 0, 0x80, 0], |is_data| { 68 admissions += 1; 69 assert_eq!(is_data, admissions == 1); 70 admissions <= 3 71 }) 72 .expect_err("fourth frame is denied"); 73 assert_eq!(error.kind(), io::ErrorKind::ConnectionAborted); 74 assert_eq!(admissions, 4); 75 } 76 77 #[test] 78 fn invalid_flags_masks_lengths_and_controls_fail_closed() { 79 for bytes in [ 80 vec![0xc1, 0], 81 vec![0x83, 0], 82 vec![0x09, 0], 83 vec![0x81, 0x80], 84 vec![0x82, 126, 0, 125], 85 header(65535), 86 vec![0x89, 126, 0, 126], 87 ] { 88 assert!( 89 upgraded().consume(&bytes, |_| true).is_err(), 90 "accepted {bytes:?}" 91 ); 92 } 93 } 94 95 #[test] 96 fn upgrade_prefix_has_an_exact_fixed_bound() { 97 let mut exact = vec![b'x'; MAX_UPGRADE_BYTES - 4]; 98 exact.extend(b"\r\n\r\n"); 99 let mut parser = WireParser::default(); 100 parser.consume(&exact, |_| true).expect("exact maximum"); 101 parser.eof().expect("upgrade boundary"); 102 let mut excessive = vec![b'x']; 103 excessive.extend(exact); 104 assert!(WireParser::default().consume(&excessive, |_| true).is_err()); 105 } 106 107 #[test] 108 fn eof_rejects_partial_upgrade_header_and_payload() { 109 assert!(WireParser::default().eof().is_err()); 110 for bytes in [vec![0x81], vec![0x82, 126, 0], vec![0x81, 2, b'x']] { 111 let mut parser = upgraded(); 112 parser.consume(&bytes, |_| true).expect("partial input"); 113 assert_eq!( 114 parser.eof().expect_err("truncation").kind(), 115 io::ErrorKind::UnexpectedEof 116 ); 117 } 118 upgraded().eof().expect("frame boundary"); 119 } 120 121 #[derive(Default)] 122 struct Calls { 123 reads: AtomicUsize, 124 writes: AtomicUsize, 125 wakes: AtomicUsize, 126 } 127 128 impl ArcWake for Calls { 129 fn wake_by_ref(arc_self: &Arc<Self>) { 130 arc_self.wakes.fetch_add(1, Ordering::SeqCst); 131 } 132 } 133 134 struct TestIo { 135 bytes: Vec<u8>, 136 offset: usize, 137 chunk: usize, 138 pending: bool, 139 calls: Arc<Calls>, 140 } 141 142 impl AsyncRead for TestIo { 143 fn poll_read( 144 mut self: Pin<&mut Self>, 145 _: &mut Context<'_>, 146 output: &mut ReadBuf<'_>, 147 ) -> Poll<io::Result<()>> { 148 self.calls.reads.fetch_add(1, Ordering::SeqCst); 149 if self.pending { 150 return Poll::Pending; 151 } 152 let count = (self.bytes.len() - self.offset) 153 .min(output.remaining()) 154 .min(self.chunk); 155 output.put_slice(&self.bytes[self.offset..self.offset + count]); 156 self.offset += count; 157 Poll::Ready(Ok(())) 158 } 159 } 160 161 impl AsyncWrite for TestIo { 162 fn poll_write( 163 self: Pin<&mut Self>, 164 _: &mut Context<'_>, 165 bytes: &[u8], 166 ) -> Poll<io::Result<usize>> { 167 self.calls.writes.fetch_add(1, Ordering::SeqCst); 168 Poll::Ready(Ok(bytes.len())) 169 } 170 fn poll_flush(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<io::Result<()>> { 171 self.calls.writes.fetch_add(1, Ordering::SeqCst); 172 Poll::Ready(Ok(())) 173 } 174 fn poll_shutdown(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<io::Result<()>> { 175 Poll::Ready(Ok(())) 176 } 177 } 178 179 fn fixture( 180 bytes: Vec<u8>, 181 ) -> ( 182 MeteredIo<TestIo>, 183 Arc<FetchBudget>, 184 FetchRegistration, 185 Arc<Calls>, 186 ) { 187 let registry = IngressRegistry::new([RELAY.to_owned()].into_iter()); 188 let budget = Arc::new(FetchBudget::default()); 189 let registration = registry 190 .register([RELAY.to_owned()].into_iter(), budget.clone()) 191 .expect("registration"); 192 let connection = registry.connection(RELAY).expect("connection"); 193 let calls = Arc::new(Calls::default()); 194 let inner = TestIo { 195 bytes, 196 offset: 0, 197 chunk: READ_CHUNK_BYTES, 198 pending: false, 199 calls: calls.clone(), 200 }; 201 ( 202 MeteredIo::new(inner, connection), 203 budget, 204 registration, 205 calls, 206 ) 207 } 208 209 fn read(meter: &mut MeteredIo<TestIo>, calls: &Arc<Calls>) -> (Poll<io::Result<()>>, usize) { 210 let waker = waker(calls.clone()); 211 let mut context = Context::from_waker(&waker); 212 let mut output = [0; READ_CHUNK_BYTES]; 213 let mut buffer = ReadBuf::new(&mut output); 214 let result = Pin::new(meter).poll_read(&mut context, &mut buffer); 215 (result, buffer.filled().len()) 216 } 217 218 #[test] 219 fn raw_upgrade_and_frame_overhead_count_before_decode() { 220 let mut bytes = UPGRADE.to_vec(); 221 bytes.extend([0x81, 0]); 222 let length = bytes.len(); 223 let (mut meter, budget, _registration, calls) = fixture(bytes); 224 assert!(budget.wire(MAX_FETCH_BYTES - length, 0, 0)); 225 assert!(matches!(read(&mut meter, &calls), (Poll::Ready(Ok(())), count) if count == length)); 226 assert!(!budget.wire(1, 0, 0)); 227 } 228 229 #[test] 230 fn excessive_read_is_not_exposed_and_subsequent_io_is_rejected() { 231 let (mut meter, budget, _registration, calls) = fixture(UPGRADE.to_vec()); 232 assert!(budget.wire(MAX_FETCH_BYTES, 0, 0)); 233 assert!(matches!(read(&mut meter, &calls), (Poll::Ready(Err(_)), 0))); 234 assert!(budget.exhausted()); 235 let reads = calls.reads.load(Ordering::SeqCst); 236 assert!(matches!(read(&mut meter, &calls), (Poll::Ready(Err(_)), 0))); 237 assert_eq!(calls.reads.load(Ordering::SeqCst), reads); 238 let waker = waker(calls.clone()); 239 let mut context = Context::from_waker(&waker); 240 assert!(matches!( 241 Pin::new(&mut meter).poll_write(&mut context, b"REQ"), 242 Poll::Ready(Err(_)) 243 )); 244 assert!(matches!( 245 Pin::new(&mut meter).poll_flush(&mut context), 246 Poll::Ready(Err(_)) 247 )); 248 assert_eq!(calls.writes.load(Ordering::SeqCst), 0); 249 } 250 251 #[test] 252 fn cancellation_wakes_pending_read_and_revokes_both_io_directions() { 253 let (mut meter, _budget, registration, calls) = fixture(Vec::new()); 254 meter.inner.pending = true; 255 assert!(matches!(read(&mut meter, &calls), (Poll::Pending, 0))); 256 let writer_calls = Arc::new(Calls::default()); 257 let writer_waker = waker(writer_calls.clone()); 258 let mut writer_context = Context::from_waker(&writer_waker); 259 assert!(matches!( 260 Pin::new(&mut meter).poll_write(&mut writer_context, b"REQ"), 261 Poll::Ready(Ok(3)) 262 )); 263 assert!(matches!( 264 Pin::new(&mut meter).poll_flush(&mut writer_context), 265 Poll::Ready(Ok(())) 266 )); 267 drop(registration); 268 assert!(calls.wakes.load(Ordering::SeqCst) > 0); 269 assert!(writer_calls.wakes.load(Ordering::SeqCst) > 0); 270 assert!(matches!(read(&mut meter, &calls), (Poll::Ready(Err(_)), 0))); 271 let waker = waker(calls.clone()); 272 let mut context = Context::from_waker(&waker); 273 assert!(matches!( 274 Pin::new(&mut meter).poll_write(&mut context, b"REQ"), 275 Poll::Ready(Err(_)) 276 )); 277 assert!(matches!( 278 Pin::new(&mut meter).poll_shutdown(&mut context), 279 Poll::Ready(Ok(())) 280 )); 281 assert_eq!(calls.reads.load(Ordering::SeqCst), 1); 282 assert_eq!(calls.writes.load(Ordering::SeqCst), 2); 283 } 284 285 #[test] 286 fn continuation_flood_exhausts_frame_budget_without_assembled_message() { 287 let mut bytes = UPGRADE.to_vec(); 288 bytes.extend([0x01, 0]); 289 for _ in 0..8192 { 290 bytes.extend([0x00, 0]); 291 } 292 let (mut meter, budget, _registration, calls) = fixture(bytes); 293 loop { 294 match read(&mut meter, &calls) { 295 (Poll::Ready(Ok(())), count) => assert!(count > 0, "unexpected EOF"), 296 (Poll::Ready(Err(_)), 0) => break, 297 other => panic!("unexpected result {other:?}"), 298 } 299 } 300 assert!(budget.exhausted()); 301 } 302 303 #[test] 304 fn empty_data_frames_exhaust_attempt_budget() { 305 let mut bytes = UPGRADE.to_vec(); 306 for _ in 0..4097 { 307 bytes.extend([0x81, 0]); 308 } 309 let (mut meter, budget, _registration, calls) = fixture(bytes); 310 loop { 311 match read(&mut meter, &calls) { 312 (Poll::Ready(Ok(())), count) => assert!(count > 0, "unexpected EOF"), 313 (Poll::Ready(Err(_)), 0) => break, 314 other => panic!("unexpected result {other:?}"), 315 } 316 } 317 assert!(budget.exhausted()); 318 } 319 320 #[test] 321 fn read_chunks_remain_bounded_and_truncation_is_terminal() { 322 let mut bytes = UPGRADE.to_vec(); 323 bytes.extend([0x81, 126, 0x20, 0]); 324 bytes.extend([b'x'; 5000]); 325 let (mut meter, _budget, _registration, calls) = fixture(bytes); 326 assert!(matches!( 327 read(&mut meter, &calls), 328 (Poll::Ready(Ok(())), READ_CHUNK_BYTES) 329 )); 330 assert!(matches!(read(&mut meter, &calls), (Poll::Ready(Ok(())), count) if count > 0)); 331 assert!( 332 matches!(read(&mut meter, &calls), (Poll::Ready(Err(error)), 0) if error.kind() == io::ErrorKind::UnexpectedEof) 333 ); 334 }