anvilsign in

collin/browser-terminal-extension

main / daemon / src / rewind.rs
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
8use std::pin::Pin;
9use std::task::{Context, Poll};
10
11use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
12
13pub struct Rewind<S> {
14 prefix: Vec<u8>,
15 pos: usize,
16 inner: S,
17}
18
19impl<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
29impl<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
46impl<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.
64pub 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
84pub 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}