net.rsannotatednet.rssource282 lines · 9.8 KB · raw
1//! Non-blocking TCP for the executor, and the adapter hyper reads it
2//! through. Futures-io traits in the middle, so futures-rustls can sit
3//! between the socket and hyper without tokio.
4
5use std::io::{self, Read, Write};
6use std::mem::MaybeUninit;
7use std::net::{SocketAddr, TcpStream, UdpSocket};
8use std::os::fd::{AsRawFd, RawFd};
9use std::pin::Pin;
10use std::task::{Context, Poll};
11
12use futures_io::{AsyncRead, AsyncWrite};
13use socket2::{Domain, Protocol, Socket, Type};
14
15use crate::executor::{self, Interest};
16
17/// A connected, non-blocking socket polled through the executor. It is
18/// `Send + Sync` (a bare fd) because hickory requires it; the executor
19/// asserts it is only ever polled on the backend thread.
20pub struct TcpIo {
21    stream: TcpStream,
22}
23
24impl TcpIo {
25    pub async fn connect(addr: SocketAddr) -> io::Result<TcpIo> {
26        let socket = Socket::new(Domain::for_address(addr), Type::STREAM, Some(Protocol::TCP))?;
27        socket.set_nonblocking(true)?;
28        socket.set_tcp_nodelay(true)?;
29        match socket.connect(&addr.into()) {
30            Ok(()) => {}
31            // io::ErrorKind::InProgress is still unstable (rustc 1.98).
32            Err(e) if e.raw_os_error() == Some(libc::EINPROGRESS) => {}
33            Err(e) => return Err(e),
34        }
35        let io = TcpIo { stream: socket.into() };
36        // Writable means the handshake finished, one way or the other.
37        std::future::poll_fn(|cx| io.poll_connected(cx)).await?;
38        Ok(io)
39    }
40
41    fn poll_connected(&self, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
42        if let Some(e) = self.stream.take_error()? {
43            return Poll::Ready(Err(e));
44        }
45        match self.stream.peer_addr() {
46            Ok(_) => Poll::Ready(Ok(())),
47            Err(e) if e.kind() == io::ErrorKind::NotConnected => {
48                executor::register(self.fd(), Interest::Write, cx.waker());
49                Poll::Pending
50            }
51            Err(e) => Poll::Ready(Err(e)),
52        }
53    }
54
55    fn fd(&self) -> RawFd {
56        self.stream.as_raw_fd()
57    }
58
59    fn poll_io<T>(
60        &self,
61        interest: Interest,
62        cx: &mut Context<'_>,
63        op: impl FnOnce(&TcpStream) -> io::Result<T>,
64    ) -> Poll<io::Result<T>> {
65        match op(&self.stream) {
66            Err(e) if e.kind() == io::ErrorKind::WouldBlock => {
67                executor::register(self.fd(), interest, cx.waker());
68                Poll::Pending
69            }
70            other => Poll::Ready(other),
71        }
72    }
73}
74
75impl Drop for TcpIo {
76    fn drop(&mut self) {
77        executor::deregister(self.fd());
78    }
79}
80
81impl AsyncRead for TcpIo {
82    fn poll_read(self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut [u8]) -> Poll<io::Result<usize>> {
83        self.poll_io(Interest::Read, cx, |mut s| s.read(buf))
84    }
85}
86
87impl AsyncWrite for TcpIo {
88    fn poll_write(self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &[u8]) -> Poll<io::Result<usize>> {
89        self.poll_io(Interest::Write, cx, |mut s| s.write(buf))
90    }
91
92    fn poll_flush(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<io::Result<()>> {
93        Poll::Ready(Ok(()))
94    }
95
96    fn poll_close(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<io::Result<()>> {
97        Poll::Ready(self.stream.shutdown(std::net::Shutdown::Write))
98    }
99}
100
101/// hyper's I/O traits over any futures-io stream.
102pub struct HyperIo<T>(pub T);
103
104impl<T: AsyncRead + Unpin> hyper::rt::Read for HyperIo<T> {
105    fn poll_read(
106        mut self: Pin<&mut Self>,
107        cx: &mut Context<'_>,
108        mut buf: hyper::rt::ReadBufCursor<'_>,
109    ) -> Poll<io::Result<()>> {
110        // SAFETY: the unfilled part is zeroed before it is viewed as
111        // `[u8]`, and `advance` covers only the bytes the read wrote.
112        unsafe {
113            let unfilled = buf.as_mut();
114            unfilled.fill(MaybeUninit::new(0));
115            let bytes = &mut *(unfilled as *mut [MaybeUninit<u8>] as *mut [u8]);
116            match Pin::new(&mut self.0).poll_read(cx, bytes) {
117                Poll::Ready(Ok(n)) => {
118                    // `n` comes from a safe trait; never let it widen the
119                    // filled region past what exists.
120                    assert!(n <= bytes.len(), "AsyncRead reported {n} of {} bytes", bytes.len());
121                    buf.advance(n);
122                    Poll::Ready(Ok(()))
123                }
124                Poll::Ready(Err(e)) => Poll::Ready(Err(e)),
125                Poll::Pending => Poll::Pending,
126            }
127        }
128    }
129}
130
131impl<T: AsyncWrite + Unpin> hyper::rt::Write for HyperIo<T> {
132    fn poll_write(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &[u8]) -> Poll<io::Result<usize>> {
133        Pin::new(&mut self.0).poll_write(cx, buf)
134    }
135
136    fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
137        Pin::new(&mut self.0).poll_flush(cx)
138    }
139
140    fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
141        Pin::new(&mut self.0).poll_close(cx)
142    }
143}
144
145/// A non-blocking UDP socket polled through the executor, for DNS.
146pub struct UdpIo {
147    socket: UdpSocket,
148}
149
150impl UdpIo {
151    pub fn bind(local: SocketAddr) -> io::Result<UdpIo> {
152        let socket = UdpSocket::bind(local)?;
153        socket.set_nonblocking(true)?;
154        Ok(UdpIo { socket })
155    }
156
157    pub fn poll_recv_from(&self, cx: &mut Context<'_>, buf: &mut [u8]) -> Poll<io::Result<(usize, SocketAddr)>> {
158        match self.socket.recv_from(buf) {
159            Err(e) if e.kind() == io::ErrorKind::WouldBlock => {
160                executor::register(self.socket.as_raw_fd(), Interest::Read, cx.waker());
161                Poll::Pending
162            }
163            other => Poll::Ready(other),
164        }
165    }
166
167    pub fn poll_send_to(&self, cx: &mut Context<'_>, buf: &[u8], target: SocketAddr) -> Poll<io::Result<usize>> {
168        match self.socket.send_to(buf, target) {
169            Err(e) if e.kind() == io::ErrorKind::WouldBlock => {
170                executor::register(self.socket.as_raw_fd(), Interest::Write, cx.waker());
171                Poll::Pending
172            }
173            other => Poll::Ready(other),
174        }
175    }
176}
177
178impl Drop for UdpIo {
179    fn drop(&mut self) {
180        executor::deregister(self.socket.as_raw_fd());
181    }
182}
183
184/// Reads the peer's `SETTINGS_MAX_CONCURRENT_STREAMS` off the bytes hyper
185/// reads, because hyper keeps h2's view of it private. It walks the
186/// server's frames (9-byte header, then payload) and looks inside
187/// SETTINGS frames only; it never changes a byte.
188pub struct PeerSettings<T> {
189    io: T,
190    frames: FrameWalk,
191    /// Called with each advertised limit, in the order the peer sent them.
192    on_max_streams: Box<dyn FnMut(u32) + Send + Sync>,
193}
194
195impl<T> PeerSettings<T> {
196    pub fn new(io: T, on_max_streams: impl FnMut(u32) + Send + Sync + 'static) -> Self {
197        PeerSettings { io, frames: FrameWalk::default(), on_max_streams: Box::new(on_max_streams) }
198    }
199}
200
201impl<T: AsyncRead + Unpin> AsyncRead for PeerSettings<T> {
202    fn poll_read(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut [u8]) -> Poll<io::Result<usize>> {
203        let this = &mut *self;
204        let n = std::task::ready!(Pin::new(&mut this.io).poll_read(cx, buf))?;
205        this.frames.feed(&buf[..n.min(buf.len())], &mut *this.on_max_streams);
206        Poll::Ready(Ok(n))
207    }
208}
209
210impl<T: AsyncWrite + Unpin> AsyncWrite for PeerSettings<T> {
211    fn poll_write(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &[u8]) -> Poll<io::Result<usize>> {
212        Pin::new(&mut self.io).poll_write(cx, buf)
213    }
214
215    fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
216        Pin::new(&mut self.io).poll_flush(cx)
217    }
218
219    fn poll_close(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
220        Pin::new(&mut self.io).poll_close(cx)
221    }
222}
223
224const FRAME_SETTINGS: u8 = 0x4;
225const FLAG_ACK: u8 = 0x1;
226const SETTINGS_MAX_CONCURRENT_STREAMS: u16 = 0x3;
227
228/// Where the walk is in the server's frame stream (RFC 9113 §4.1).
229#[derive(Default)]
230struct FrameWalk {
231    header: [u8; 9],
232    have: usize,
233    /// Payload bytes of the current frame still to pass.
234    remaining: usize,
235    /// Inside a SETTINGS frame (not an ACK), collecting 6-byte entries.
236    settings: bool,
237    entry: [u8; 6],
238    entry_have: usize,
239}
240
241impl FrameWalk {
242    fn feed(&mut self, mut bytes: &[u8], on_max_streams: &mut dyn FnMut(u32)) {
243        while !bytes.is_empty() {
244            if self.remaining == 0 && self.have < 9 {
245                let take = (9 - self.have).min(bytes.len());
246                self.header[self.have..self.have + take].copy_from_slice(&bytes[..take]);
247                self.have += take;
248                bytes = &bytes[take..];
249                if self.have < 9 {
250                    return;
251                }
252                let h = self.header;
253                self.remaining = u32::from_be_bytes([0, h[0], h[1], h[2]]) as usize;
254                self.settings = h[3] == FRAME_SETTINGS && h[4] & FLAG_ACK == 0;
255                self.entry_have = 0;
256                if self.remaining == 0 {
257                    self.have = 0;
258                }
259                continue;
260            }
261            let take = self.remaining.min(bytes.len());
262            if self.settings {
263                for &b in &bytes[..take] {
264                    self.entry[self.entry_have] = b;
265                    self.entry_have += 1;
266                    if self.entry_have == 6 {
267                        self.entry_have = 0;
268                        let e = self.entry;
269                        if u16::from_be_bytes([e[0], e[1]]) == SETTINGS_MAX_CONCURRENT_STREAMS {
270                            on_max_streams(u32::from_be_bytes([e[2], e[3], e[4], e[5]]));
271                        }
272                    }
273                }
274            }
275            self.remaining -= take;
276            bytes = &bytes[take..];
277            if self.remaining == 0 {
278                self.have = 0;
279            }
280        }
281    }
282}