source_wire.rs (7993B)
1 //! Fixed-memory admission of decrypted ingress before WebSocket decoding. 2 3 use crate::source_ingress::IngressConnection; 4 use std::io; 5 use std::pin::Pin; 6 use std::task::{Context, Poll}; 7 use tokio::io::{AsyncRead, AsyncWrite, ReadBuf}; 8 9 const MAX_UPGRADE_BYTES: usize = 16 * 1024; 10 const READ_CHUNK_BYTES: usize = 4096; 11 12 /// Counts HTTP response and frame overhead conservatively in the same raw-byte 13 /// allowance as payloads. No message or payload is retained by this layer. 14 pub(crate) struct MeteredIo<S> { 15 inner: S, 16 ingress: IngressConnection, 17 parser: WireParser, 18 failed: bool, 19 } 20 21 impl<S> MeteredIo<S> { 22 pub(crate) fn new(inner: S, ingress: IngressConnection) -> Self { 23 Self { 24 inner, 25 ingress, 26 parser: WireParser::default(), 27 failed: false, 28 } 29 } 30 31 fn admitted(&self, context: &Context<'_>) -> io::Result<()> { 32 if self.failed || !self.ingress.admitted(context) { 33 Err(denied()) 34 } else { 35 Ok(()) 36 } 37 } 38 } 39 40 impl<S: AsyncRead + Unpin> AsyncRead for MeteredIo<S> { 41 fn poll_read( 42 self: Pin<&mut Self>, 43 context: &mut Context<'_>, 44 output: &mut ReadBuf<'_>, 45 ) -> Poll<io::Result<()>> { 46 let this = self.get_mut(); 47 if let Err(error) = this.admitted(context) { 48 return Poll::Ready(Err(error)); 49 } 50 if output.remaining() == 0 { 51 return Poll::Ready(Ok(())); 52 } 53 let mut scratch = [0; READ_CHUNK_BYTES]; 54 let length = output.remaining().min(scratch.len()); 55 let mut input = ReadBuf::new(&mut scratch[..length]); 56 match Pin::new(&mut this.inner).poll_read(context, &mut input) { 57 Poll::Pending => Poll::Pending, 58 Poll::Ready(Err(error)) => { 59 this.failed = true; 60 Poll::Ready(Err(error)) 61 } 62 Poll::Ready(Ok(())) => { 63 let bytes = input.filled(); 64 let result = if !this.ingress.charge(bytes.len(), 0, 0) { 65 Err(denied()) 66 } else if bytes.is_empty() { 67 this.parser.eof() 68 } else { 69 this.parser 70 .consume(bytes, |data| this.ingress.charge(0, 1, usize::from(data))) 71 }; 72 if let Err(error) = result { 73 this.failed = true; 74 return Poll::Ready(Err(error)); 75 } 76 // Cancellation may race a ready read. Revalidate before any 77 // plaintext becomes visible to the WebSocket implementation. 78 if let Err(error) = this.admitted(context) { 79 this.failed = true; 80 return Poll::Ready(Err(error)); 81 } 82 output.put_slice(bytes); 83 Poll::Ready(Ok(())) 84 } 85 } 86 } 87 } 88 89 impl<S: AsyncWrite + Unpin> AsyncWrite for MeteredIo<S> { 90 fn poll_write( 91 self: Pin<&mut Self>, 92 context: &mut Context<'_>, 93 bytes: &[u8], 94 ) -> Poll<io::Result<usize>> { 95 let this = self.get_mut(); 96 if this.failed || !this.ingress.admitted_write(context) { 97 return Poll::Ready(Err(denied())); 98 } 99 Pin::new(&mut this.inner).poll_write(context, bytes) 100 } 101 102 fn poll_flush(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<io::Result<()>> { 103 let this = self.get_mut(); 104 if this.failed || !this.ingress.admitted_write(context) { 105 return Poll::Ready(Err(denied())); 106 } 107 Pin::new(&mut this.inner).poll_flush(context) 108 } 109 110 fn poll_shutdown(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<io::Result<()>> { 111 // Closing the underlying transport remains possible after revocation. 112 Pin::new(&mut self.get_mut().inner).poll_shutdown(context) 113 } 114 } 115 116 #[derive(Default)] 117 struct WireParser { 118 upgrade_bytes: usize, 119 delimiter: u32, 120 upgraded: bool, 121 header: [u8; 10], 122 header_len: usize, 123 header_needed: usize, 124 payload_remaining: usize, 125 } 126 127 impl WireParser { 128 fn consume( 129 &mut self, 130 mut bytes: &[u8], 131 mut admit_frame: impl FnMut(bool) -> bool, 132 ) -> io::Result<()> { 133 while let Some((&byte, rest)) = bytes.split_first() { 134 if !self.upgraded { 135 self.upgrade_bytes += 1; 136 if self.upgrade_bytes > MAX_UPGRADE_BYTES { 137 return Err(invalid("WebSocket upgrade response exceeds its limit")); 138 } 139 self.delimiter = (self.delimiter << 8) | u32::from(byte); 140 self.upgraded = self.delimiter == u32::from_be_bytes(*b"\r\n\r\n"); 141 bytes = rest; 142 } else if self.payload_remaining != 0 { 143 let consumed = self.payload_remaining.min(bytes.len()); 144 self.payload_remaining -= consumed; 145 bytes = &bytes[consumed..]; 146 } else { 147 if self.header_len == 0 { 148 let opcode = byte & 0x0f; 149 if !admit_frame(matches!(opcode, 1 | 2)) { 150 return Err(denied()); 151 } 152 if byte & 0x70 != 0 || !matches!(opcode, 0 | 1 | 2 | 8 | 9 | 10) { 153 return Err(invalid("unsupported WebSocket frame flags or opcode")); 154 } 155 if opcode >= 8 && byte & 0x80 == 0 { 156 return Err(invalid("fragmented WebSocket control frame")); 157 } 158 self.header_needed = 2; 159 } 160 self.header[self.header_len] = byte; 161 self.header_len += 1; 162 bytes = rest; 163 if self.header_len == 2 { 164 if byte & 0x80 != 0 { 165 return Err(invalid("masked server WebSocket frame")); 166 } 167 self.header_needed = match byte { 168 126 => 4, 169 127 => 10, 170 _ => 2, 171 }; 172 } 173 if self.header_len == self.header_needed { 174 self.payload_remaining = self.payload_length()?; 175 self.header_len = 0; 176 } 177 } 178 } 179 Ok(()) 180 } 181 182 fn payload_length(&self) -> io::Result<usize> { 183 let (length, minimum) = match self.header[1] { 184 126 => ( 185 u64::from(u16::from_be_bytes([self.header[2], self.header[3]])), 186 126, 187 ), 188 127 => { 189 let mut extended = [0; 8]; 190 extended.copy_from_slice(&self.header[2..10]); 191 (u64::from_be_bytes(extended), 65536) 192 } 193 length => (u64::from(length), 0), 194 }; 195 if length < minimum || length > crate::relay::MAX_WIRE_MESSAGE_BYTES as u64 { 196 return Err(invalid( 197 "WebSocket frame length is noncanonical or exceeds its limit", 198 )); 199 } 200 if self.header[0] & 0x0f >= 8 && length > 125 { 201 return Err(invalid("WebSocket control frame exceeds its limit")); 202 } 203 usize::try_from(length).map_err(|_| invalid("WebSocket frame length cannot be represented")) 204 } 205 206 fn eof(&self) -> io::Result<()> { 207 if !self.upgraded || self.header_len != 0 || self.payload_remaining != 0 { 208 Err(io::Error::new( 209 io::ErrorKind::UnexpectedEof, 210 "truncated WebSocket ingress", 211 )) 212 } else { 213 Ok(()) 214 } 215 } 216 } 217 218 fn invalid(message: &'static str) -> io::Error { 219 io::Error::new(io::ErrorKind::InvalidData, message) 220 } 221 222 fn denied() -> io::Error { 223 io::Error::new( 224 io::ErrorKind::ConnectionAborted, 225 "relay ingress budget or generation revoked", 226 ) 227 } 228 229 #[cfg(test)] 230 #[path = "source_wire_tests.rs"] 231 mod tests;