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}