| 1 | //! A stream wrapper that replays bytes we already consumed. |
| 2 | //! |
| 3 | //! We need to look at the HTTP request head *before* handing the stream to the |
| 4 | //! WebSocket handshake, so we can serve a human-readable page to someone who |
| 5 | //! navigated to https://127.0.0.1:7681 to accept the certificate. TLS streams |
| 6 | //! can't be peeked like a TcpStream, so we read, decide, and rewind. |
| 7 | |
| 8 | use std::pin::Pin; |
| 9 | use std::task::{Context, Poll}; |
| 10 | |
| 11 | use tokio::io::{AsyncRead, AsyncWrite, ReadBuf}; |
| 12 | |
| 13 | pub struct Rewind<S> { |
| 14 | prefix: Vec<u8>, |
| 15 | pos: usize, |
| 16 | inner: S, |
| 17 | } |
| 18 | |
| 19 | impl<S> Rewind<S> { |
| 20 | pub fn new(prefix: Vec<u8>, inner: S) -> Self { |
| 21 | Self { |
| 22 | prefix, |
| 23 | pos: 0, |
| 24 | inner, |
| 25 | } |
| 26 | } |
| 27 | } |
| 28 | |
| 29 | impl<S: AsyncRead + Unpin> AsyncRead for Rewind<S> { |
| 30 | fn poll_read( |
| 31 | mut self: Pin<&mut Self>, |
| 32 | cx: &mut Context<'_>, |
| 33 | buf: &mut ReadBuf<'_>, |
| 34 | ) -> Poll<std::io::Result<()>> { |
| 35 | if self.pos < self.prefix.len() { |
| 36 | let remaining = &self.prefix[self.pos..]; |
| 37 | let n = remaining.len().min(buf.remaining()); |
| 38 | buf.put_slice(&remaining[..n]); |
| 39 | self.pos += n; |
| 40 | return Poll::Ready(Ok(())); |
| 41 | } |
| 42 | Pin::new(&mut self.inner).poll_read(cx, buf) |
| 43 | } |
| 44 | } |
| 45 | |
| 46 | impl<S: AsyncWrite + Unpin> AsyncWrite for Rewind<S> { |
| 47 | fn poll_write( |
| 48 | mut self: Pin<&mut Self>, |
| 49 | cx: &mut Context<'_>, |
| 50 | buf: &[u8], |
| 51 | ) -> Poll<std::io::Result<usize>> { |
| 52 | Pin::new(&mut self.inner).poll_write(cx, buf) |
| 53 | } |
| 54 | fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> { |
| 55 | Pin::new(&mut self.inner).poll_flush(cx) |
| 56 | } |
| 57 | fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> { |
| 58 | Pin::new(&mut self.inner).poll_shutdown(cx) |
| 59 | } |
| 60 | } |
| 61 | |
| 62 | /// Read the request head (up to the blank line) without consuming the body. |
| 63 | /// Returns the bytes read so they can be replayed. |
| 64 | pub async fn read_head<S: AsyncRead + Unpin>( |
| 65 | stream: &mut S, |
| 66 | limit: usize, |
| 67 | ) -> std::io::Result<Vec<u8>> { |
| 68 | use tokio::io::AsyncReadExt; |
| 69 | let mut buf = Vec::with_capacity(1024); |
| 70 | let mut byte = [0u8; 1]; |
| 71 | while buf.len() < limit { |
| 72 | let n = stream.read(&mut byte).await?; |
| 73 | if n == 0 { |
| 74 | break; |
| 75 | } |
| 76 | buf.push(byte[0]); |
| 77 | if buf.ends_with(b"\r\n\r\n") { |
| 78 | break; |
| 79 | } |
| 80 | } |
| 81 | Ok(buf) |
| 82 | } |
| 83 | |
| 84 | pub fn is_websocket_upgrade(head: &[u8]) -> bool { |
| 85 | let text = String::from_utf8_lossy(head).to_ascii_lowercase(); |
| 86 | text.contains("upgrade: websocket") || text.contains("upgrade:websocket") |
| 87 | } |