diff --git a/src/framer_async.rs b/src/framer_async.rs index f63815c..8828501 100644 --- a/src/framer_async.rs +++ b/src/framer_async.rs @@ -46,6 +46,11 @@ where TWebSocketType: WebSocketType, { websocket: WebSocket, + // Small buffer for a websocket frame header, so we can progress + // reading data even when we received only a part of websocket header. + header_buf: [u8; 16], + // How much of header_buf contains data. + header_buf_len: usize, frame_cursor: usize, rx_remainder_len: usize, } @@ -137,6 +142,8 @@ where websocket, frame_cursor: 0, rx_remainder_len: 0, + header_buf: [0; 16], + header_buf_len: 0, } } @@ -196,10 +203,15 @@ where Ok(()) } + /// Return true if framer can make a progress without doing a read from a stream + pub fn read_ready(&self) -> bool { + self.rx_remainder_len != 0 + } + // NOTE: any unused bytes read from the stream but not decoded are stored at the end // of the buffer to be used next time this read function is called. This also applies to // any unused bytes read when the connect handshake was made. Therefore it is important that - // the caller does not clear this buffer between calls or use it for anthing other than reads. + // the caller does not clear this buffer between calls or use it for anything other than reads. pub async fn read<'a, B: Deref, E>( &mut self, stream: &mut (impl Stream> + Sink<&'a [u8], Error = E> + Unpin), @@ -211,15 +223,22 @@ where if self.rx_remainder_len == 0 { match stream.next().await { Some(Ok(input)) => { - if buffer.len() < input.len() { + if buffer.len() < input.len() + self.frame_cursor + self.header_buf_len { return Some(Err(FramerError::RxBufferTooSmall(input.len()))); } - let rx_start = buffer.len() - input.len(); + let rx_start = buffer.len() - input.len() - self.header_buf_len; // copy to end of buffer - buffer[rx_start..].copy_from_slice(&input); - self.rx_remainder_len = input.len() + + // copy previous part of frame header if any + buffer[rx_start..rx_start + self.header_buf_len] + .copy_from_slice(&self.header_buf[0..self.header_buf_len]); + // copy new data + buffer[rx_start + self.header_buf_len..].copy_from_slice(&input); + + self.rx_remainder_len = input.len() + self.header_buf_len; + self.header_buf_len = 0; } Some(Err(e)) => { return Some(Err(FramerError::Io(e))); @@ -230,9 +249,24 @@ where let rx_start = buffer.len() - self.rx_remainder_len; let (frame_buf, rx_buf) = buffer.split_at_mut(rx_start); + let frame_buf_remaining = &mut frame_buf[self.frame_cursor..]; - let ws_result = match self.websocket.read(rx_buf, frame_buf) { + let ws_result = match self.websocket.read(rx_buf, frame_buf_remaining) { Ok(ws_result) => ws_result, + // We might receive part of websocket header. Ex. receive data one byte at a time. + // Store enough data in a small buffer so we can at least read the header + Err(crate::Error::ReadFrameIncomplete) => { + if rx_buf.len() + self.header_buf_len > self.header_buf.len() { + return Some(Err(FramerError::WebSocket( + crate::Error::ReadFrameIncomplete, + ))); + } + self.header_buf[self.header_buf_len..self.header_buf_len + rx_buf.len()] + .copy_from_slice(&rx_buf); + self.header_buf_len += rx_buf.len(); + self.rx_remainder_len -= rx_buf.len(); + return None; + } Err(e) => return Some(Err(FramerError::WebSocket(e))), };