Skip to main content

tokio_rustls/
server.rs

1use std::future::Future;
2use std::io::{self, BufRead as _};
3#[cfg(unix)]
4use std::os::unix::io::{AsRawFd, RawFd};
5#[cfg(windows)]
6use std::os::windows::io::{AsRawSocket, RawSocket};
7use std::pin::Pin;
8use std::sync::Arc;
9use std::task::{Context, Poll};
10
11use rustls::server::AcceptedAlert;
12use rustls::{ServerConfig, ServerConnection};
13use tokio::io::{AsyncBufRead, AsyncRead, AsyncWrite, ReadBuf};
14
15use crate::common::{IoSession, MidHandshake, Stream, SyncReadAdapter, SyncWriteAdapter, TlsState};
16
17/// A wrapper around a `rustls::ServerConfig`, providing an async `accept` method.
18#[derive(Clone)]
19pub struct TlsAcceptor {
20    inner: Arc<ServerConfig>,
21}
22
23impl From<Arc<ServerConfig>> for TlsAcceptor {
24    fn from(inner: Arc<ServerConfig>) -> Self {
25        Self { inner }
26    }
27}
28
29impl TlsAcceptor {
30    /// Returns a future for completing a TLS handshake for a client using `stream`.
31    ///
32    /// You likely want to wrap this in a timeout (for example with [`tokio::time::timeout`][])
33    /// to bound the handshake time.
34    ///
35    /// [`tokio::time::timeout`]: https://docs.rs/tokio/latest/tokio/time/fn.timeout.html
36    #[inline]
37    pub fn accept<IO>(&self, stream: IO) -> Accept<IO>
38    where
39        IO: AsyncRead + AsyncWrite + Unpin,
40    {
41        self.accept_with(stream, |_| ())
42    }
43
44    /// Similar to [`Self::accept()`], but calls `f` before performing the handshake.
45    ///
46    /// As with [`Self::accept()`] you likely want to wrap this in a timeout to
47    /// bound the handshake time.
48    ///
49    /// The `f` handler is given a mutable reference to a [`ServerConnection`][] that can be used
50    /// to configure the connection before the handshake, for example, adjusting the buffer limit.
51    ///
52    /// Because no data has been read from `stream` yet when `f` is called ClientHello
53    /// dependent state (like early data) is not yet available.
54    ///
55    /// [`ServerConnection`]: https://docs.rs/rustls/latest/rustls/server/struct.ServerConnection.html
56    pub fn accept_with<IO, F>(&self, stream: IO, f: F) -> Accept<IO>
57    where
58        IO: AsyncRead + AsyncWrite + Unpin,
59        F: FnOnce(&mut ServerConnection),
60    {
61        let mut session = match ServerConnection::new(self.inner.clone()) {
62            Ok(session) => session,
63            Err(error) => {
64                return Accept(MidHandshake::Error {
65                    io: stream,
66                    // TODO(eliza): should this really return an `io::Error`?
67                    // Probably not...
68                    error: io::Error::new(io::ErrorKind::Other, error),
69                });
70            }
71        };
72        f(&mut session);
73
74        Accept(MidHandshake::Handshaking(TlsStream {
75            session,
76            io: stream,
77            state: TlsState::Stream,
78            need_flush: false,
79            error: None,
80        }))
81    }
82
83    /// Get a read-only reference to underlying config
84    pub fn config(&self) -> &Arc<ServerConfig> {
85        &self.inner
86    }
87}
88
89/// A future for reading a `ClientHello` from `io` without committing to a [`ServerConfig`][].
90///
91/// Awaiting it yields a [`StartHandshake`], which exposes the
92/// [`ClientHello`][] (for example, to choose a config based on SNI) and performs
93/// the rest of the handshake via [`StartHandshake::into_stream()`].
94///
95/// [`ServerConfig`]: https://docs.rs/rustls/latest/rustls/server/struct.ServerConfig.html
96/// [`ClientHello`]: https://docs.rs/rustls/latest/rustls/server/struct.ClientHello.html
97pub struct LazyConfigAcceptor<IO> {
98    acceptor: rustls::server::Acceptor,
99    io: Option<IO>,
100    alert: Option<(rustls::Error, AcceptedAlert)>,
101}
102
103impl<IO> LazyConfigAcceptor<IO>
104where
105    IO: AsyncRead + AsyncWrite + Unpin,
106{
107    /// Returns a new `LazyConfigAcceptor` that reads a `ClientHello` from `io`.
108    ///
109    /// You likely want to wrap awaiting the acceptor in a timeout to bound how long the
110    /// peer may take to send the `ClientHello`.
111    ///
112    /// Note that awaiting the acceptor is only the first half of the handshake and
113    /// [`StartHandshake::into_stream()`] performs the rest.
114    ///
115    /// To bound the time for the complete handshake, share one deadline across
116    /// both awaits (for example with [`tokio::time::timeout_at`][]) rather than giving each
117    /// its own timeout.
118    ///
119    /// If a timeout elapses before the `ClientHello` arrives, [`Self::take_io()`] can
120    /// recover the `io`, for example to answer the peer in plaintext before closing.
121    ///
122    /// [`tokio::time::timeout_at`]: https://docs.rs/tokio/latest/tokio/time/fn.timeout_at.html
123    #[inline]
124    pub fn new(acceptor: rustls::server::Acceptor, io: IO) -> Self {
125        Self {
126            acceptor,
127            io: Some(io),
128            alert: None,
129        }
130    }
131
132    /// Takes back the client connection. Will return `None` if called more than once or if the
133    /// connection has been accepted.
134    ///
135    /// # Example
136    ///
137    /// ```no_run
138    /// # fn choose_server_config(
139    /// #     _: rustls::server::ClientHello,
140    /// # ) -> std::sync::Arc<rustls::ServerConfig> {
141    /// #     unimplemented!();
142    /// # }
143    /// # #[allow(unused_variables)]
144    /// # async fn listen() {
145    /// use tokio::io::AsyncWriteExt;
146    /// let listener = tokio::net::TcpListener::bind("127.0.0.1:4443").await.unwrap();
147    /// let (stream, _) = listener.accept().await.unwrap();
148    ///
149    /// let acceptor = tokio_rustls::LazyConfigAcceptor::new(rustls::server::Acceptor::default(), stream);
150    /// tokio::pin!(acceptor);
151    ///
152    /// match acceptor.as_mut().await {
153    ///     Ok(start) => {
154    ///         let clientHello = start.client_hello();
155    ///         let config = choose_server_config(clientHello);
156    ///         let stream = start.into_stream(config).await.unwrap();
157    ///         // Proceed with handling the ServerConnection...
158    ///     }
159    ///     Err(err) => {
160    ///         if let Some(mut stream) = acceptor.take_io() {
161    ///             stream
162    ///                 .write_all(
163    ///                     format!("HTTP/1.1 400 Invalid Input\r\n\r\n\r\n{:?}\n", err)
164    ///                         .as_bytes()
165    ///                 )
166    ///                 .await
167    ///                 .unwrap();
168    ///         }
169    ///     }
170    /// }
171    /// # }
172    /// ```
173    pub fn take_io(&mut self) -> Option<IO> {
174        self.io.take()
175    }
176}
177
178impl<IO> Future for LazyConfigAcceptor<IO>
179where
180    IO: AsyncRead + AsyncWrite + Unpin,
181{
182    type Output = Result<StartHandshake<IO>, io::Error>;
183
184    fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
185        let this = self.get_mut();
186        loop {
187            let io = match this.io.as_mut() {
188                Some(io) => io,
189                None => {
190                    return Poll::Ready(Err(io::Error::new(
191                        io::ErrorKind::Other,
192                        "acceptor cannot be polled after acceptance",
193                    )));
194                }
195            };
196
197            if let Some((err, mut alert)) = this.alert.take() {
198                match alert.write(&mut SyncWriteAdapter { io, cx }) {
199                    Err(e) if e.kind() == io::ErrorKind::WouldBlock => {
200                        this.alert = Some((err, alert));
201                        return Poll::Pending;
202                    }
203                    Ok(0) | Err(_) => {
204                        return Poll::Ready(Err(io::Error::new(io::ErrorKind::InvalidData, err)));
205                    }
206                    Ok(_) => {
207                        this.alert = Some((err, alert));
208                        continue;
209                    }
210                };
211            }
212
213            let mut reader = SyncReadAdapter { io, cx };
214            match this.acceptor.read_tls(&mut reader) {
215                Ok(0) => return Err(io::ErrorKind::UnexpectedEof.into()).into(),
216                Ok(_) => {}
217                Err(e) if e.kind() == io::ErrorKind::WouldBlock => return Poll::Pending,
218                Err(e) => return Err(e).into(),
219            }
220
221            match this.acceptor.accept() {
222                Ok(Some(accepted)) => {
223                    let io = this.io.take().unwrap();
224                    return Poll::Ready(Ok(StartHandshake { accepted, io }));
225                }
226                Ok(None) => {}
227                Err((err, alert)) => {
228                    this.alert = Some((err, alert));
229                }
230            }
231        }
232    }
233}
234
235/// An incoming connection received through [`LazyConfigAcceptor`].
236///
237/// This contains the generic `IO` asynchronous transport,
238/// [`ClientHello`](rustls::server::ClientHello) data,
239/// and all the state required to continue the TLS handshake (e.g. via
240/// [`StartHandshake::into_stream`]).
241#[non_exhaustive]
242#[derive(Debug)]
243pub struct StartHandshake<IO> {
244    pub accepted: rustls::server::Accepted,
245    pub io: IO,
246}
247
248impl<IO> StartHandshake<IO>
249where
250    IO: AsyncRead + AsyncWrite + Unpin,
251{
252    /// Create a new object from an `IO` transport and prior TLS metadata.
253    pub fn from_parts(accepted: rustls::server::Accepted, transport: IO) -> Self {
254        Self {
255            accepted,
256            io: transport,
257        }
258    }
259
260    pub fn client_hello(&self) -> rustls::server::ClientHello<'_> {
261        self.accepted.client_hello()
262    }
263
264    /// Returns a future that performs the rest of the TLS handshake using `config`.
265    ///
266    /// You likely want to wrap this in a timeout to bound the handshake time. Ideally
267    /// with [`tokio::time::timeout_at`][], reusing the deadline that also bounded
268    /// awaiting the [`LazyConfigAcceptor`] so both halves of the handshake share
269    /// one budget. See [`LazyConfigAcceptor::new()`].
270    ///
271    /// [`tokio::time::timeout_at`]: https://docs.rs/tokio/latest/tokio/time/fn.timeout_at.html
272    pub fn into_stream(self, config: Arc<ServerConfig>) -> Accept<IO> {
273        self.into_stream_with(config, |_| ())
274    }
275
276    /// Similar to [`Self::into_stream()`], but calls `f` before performing the handshake.
277    ///
278    /// As with [`Self::into_stream()`] you likely want to wrap this in a timeout to
279    /// bound the handshake time.
280    ///
281    /// The `f` handler is given a mutable reference to a [`ServerConnection`][] that can be
282    /// used to configure the connection before the handshake.
283    ///
284    /// [`ServerConnection`]: https://docs.rs/rustls/latest/rustls/server/struct.ServerConnection.html
285    pub fn into_stream_with<F>(self, config: Arc<ServerConfig>, f: F) -> Accept<IO>
286    where
287        F: FnOnce(&mut ServerConnection),
288    {
289        let mut conn = match self.accepted.into_connection(config) {
290            Ok(conn) => conn,
291            Err((error, alert)) => {
292                return Accept(MidHandshake::SendAlert {
293                    io: self.io,
294                    alert,
295                    // TODO(eliza): should this really return an `io::Error`?
296                    // Probably not...
297                    error: io::Error::new(io::ErrorKind::InvalidData, error),
298                });
299            }
300        };
301        f(&mut conn);
302
303        Accept(MidHandshake::Handshaking(TlsStream {
304            session: conn,
305            io: self.io,
306            state: TlsState::Stream,
307            need_flush: false,
308            error: None,
309        }))
310    }
311}
312
313/// Future returned from `TlsAcceptor::accept` which will resolve
314/// once the accept handshake has finished.
315pub struct Accept<IO>(MidHandshake<TlsStream<IO>>);
316
317impl<IO> Accept<IO> {
318    #[inline]
319    pub fn into_fallible(self) -> FallibleAccept<IO> {
320        FallibleAccept(self.0)
321    }
322
323    pub fn get_ref(&self) -> Option<&IO> {
324        match &self.0 {
325            MidHandshake::Handshaking(sess) => Some(sess.get_ref().0),
326            MidHandshake::SendAlert { io, .. } => Some(io),
327            MidHandshake::Error { io, .. } => Some(io),
328            MidHandshake::End => None,
329        }
330    }
331
332    pub fn get_mut(&mut self) -> Option<&mut IO> {
333        match &mut self.0 {
334            MidHandshake::Handshaking(sess) => Some(sess.get_mut().0),
335            MidHandshake::SendAlert { io, .. } => Some(io),
336            MidHandshake::Error { io, .. } => Some(io),
337            MidHandshake::End => None,
338        }
339    }
340}
341
342impl<IO: AsyncRead + AsyncWrite + Unpin> Future for Accept<IO> {
343    type Output = io::Result<TlsStream<IO>>;
344
345    #[inline]
346    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
347        Pin::new(&mut self.0).poll(cx).map_err(|(err, _)| err)
348    }
349}
350
351/// Like [Accept], but returns `IO` on failure.
352pub struct FallibleAccept<IO>(MidHandshake<TlsStream<IO>>);
353
354impl<IO: AsyncRead + AsyncWrite + Unpin> Future for FallibleAccept<IO> {
355    type Output = Result<TlsStream<IO>, (io::Error, IO)>;
356
357    #[inline]
358    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
359        Pin::new(&mut self.0).poll(cx)
360    }
361}
362
363/// A wrapper around an underlying raw stream which implements the TLS or SSL
364/// protocol.
365#[derive(Debug)]
366pub struct TlsStream<IO> {
367    pub(crate) io: IO,
368    pub(crate) session: ServerConnection,
369    pub(crate) state: TlsState,
370    pub(crate) need_flush: bool,
371    /// Buffered error that occurred during batch reading
372    pub(crate) error: Option<io::Error>,
373}
374
375impl<IO> TlsStream<IO> {
376    #[inline]
377    pub fn get_ref(&self) -> (&IO, &ServerConnection) {
378        (&self.io, &self.session)
379    }
380
381    #[inline]
382    pub fn get_mut(&mut self) -> (&mut IO, &mut ServerConnection) {
383        (&mut self.io, &mut self.session)
384    }
385
386    #[inline]
387    pub fn into_inner(self) -> (IO, ServerConnection) {
388        (self.io, self.session)
389    }
390}
391
392impl<IO> IoSession for TlsStream<IO> {
393    type Io = IO;
394    type Session = ServerConnection;
395
396    #[inline]
397    fn skip_handshake(&self) -> bool {
398        false
399    }
400
401    #[inline]
402    fn get_mut(&mut self) -> (&mut TlsState, &mut Self::Io, &mut Self::Session, &mut bool) {
403        (
404            &mut self.state,
405            &mut self.io,
406            &mut self.session,
407            &mut self.need_flush,
408        )
409    }
410
411    #[inline]
412    fn into_io(self) -> Self::Io {
413        self.io
414    }
415}
416
417impl<IO> AsyncRead for TlsStream<IO>
418where
419    IO: AsyncRead + AsyncWrite + Unpin,
420{
421    fn poll_read(
422        mut self: Pin<&mut Self>,
423        cx: &mut Context<'_>,
424        buf: &mut ReadBuf<'_>,
425    ) -> Poll<io::Result<()>> {
426        if let Some(err) = self.error.take() {
427            return Poll::Ready(Err(err));
428        };
429        let data = ready!(self.as_mut().poll_fill_buf(cx))?;
430        let len = data.len().min(buf.remaining());
431        if len == 0 {
432            return Poll::Ready(Ok(()));
433        }
434        buf.put_slice(&data[..len]);
435        self.as_mut().consume(len);
436
437        while buf.remaining() > 0 {
438            let data = match self.as_mut().poll_fill_buf(cx) {
439                Poll::Ready(Ok([])) => break,
440                Poll::Ready(Ok(data)) => data,
441                Poll::Ready(Err(err)) => {
442                    self.error = Some(err);
443                    break;
444                }
445                Poll::Pending => break,
446            };
447            let len = Ord::min(data.len(), buf.remaining());
448            buf.put_slice(&data[..len]);
449            self.as_mut().consume(len);
450        }
451        Poll::Ready(Ok(()))
452    }
453}
454
455impl<IO> AsyncBufRead for TlsStream<IO>
456where
457    IO: AsyncRead + AsyncWrite + Unpin,
458{
459    fn poll_fill_buf(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<&[u8]>> {
460        match self.state {
461            TlsState::Stream | TlsState::WriteShutdown => {
462                let this = self.get_mut();
463                let stream =
464                    Stream::new(&mut this.io, &mut this.session).set_eof(!this.state.readable());
465
466                match stream.poll_fill_buf(cx) {
467                    Poll::Ready(Ok(buf)) => {
468                        if buf.is_empty() {
469                            this.state.shutdown_read();
470                        }
471
472                        Poll::Ready(Ok(buf))
473                    }
474                    Poll::Ready(Err(err)) if err.kind() == io::ErrorKind::ConnectionAborted => {
475                        this.state.shutdown_read();
476                        Poll::Ready(Err(err))
477                    }
478                    output => output,
479                }
480            }
481            TlsState::ReadShutdown | TlsState::FullyShutdown => Poll::Ready(Ok(&[])),
482            #[cfg(feature = "early-data")]
483            ref s => unreachable!("server TLS can not hit this state: {:?}", s),
484        }
485    }
486
487    fn consume(mut self: Pin<&mut Self>, amt: usize) {
488        self.session.reader().consume(amt);
489    }
490}
491
492impl<IO> AsyncWrite for TlsStream<IO>
493where
494    IO: AsyncRead + AsyncWrite + Unpin,
495{
496    /// Note: that it does not guarantee the final data to be sent.
497    /// To be cautious, you must manually call `flush`.
498    fn poll_write(
499        self: Pin<&mut Self>,
500        cx: &mut Context<'_>,
501        buf: &[u8],
502    ) -> Poll<io::Result<usize>> {
503        let this = self.get_mut();
504        let mut stream =
505            Stream::new(&mut this.io, &mut this.session).set_eof(!this.state.readable());
506        stream.as_mut_pin().poll_write(cx, buf)
507    }
508
509    /// Note: that it does not guarantee the final data to be sent.
510    /// To be cautious, you must manually call `flush`.
511    fn poll_write_vectored(
512        self: Pin<&mut Self>,
513        cx: &mut Context<'_>,
514        bufs: &[io::IoSlice<'_>],
515    ) -> Poll<io::Result<usize>> {
516        let this = self.get_mut();
517        let mut stream =
518            Stream::new(&mut this.io, &mut this.session).set_eof(!this.state.readable());
519        stream.as_mut_pin().poll_write_vectored(cx, bufs)
520    }
521
522    #[inline]
523    fn is_write_vectored(&self) -> bool {
524        true
525    }
526
527    fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
528        let this = self.get_mut();
529        let mut stream =
530            Stream::new(&mut this.io, &mut this.session).set_eof(!this.state.readable());
531        stream.as_mut_pin().poll_flush(cx)
532    }
533
534    fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
535        if self.state.writeable() {
536            self.session.send_close_notify();
537            self.state.shutdown_write();
538        }
539
540        let this = self.get_mut();
541        let mut stream =
542            Stream::new(&mut this.io, &mut this.session).set_eof(!this.state.readable());
543        stream.as_mut_pin().poll_shutdown(cx)
544    }
545}
546
547#[cfg(unix)]
548impl<IO> AsRawFd for TlsStream<IO>
549where
550    IO: AsRawFd,
551{
552    fn as_raw_fd(&self) -> RawFd {
553        self.get_ref().0.as_raw_fd()
554    }
555}
556
557#[cfg(windows)]
558impl<IO> AsRawSocket for TlsStream<IO>
559where
560    IO: AsRawSocket,
561{
562    fn as_raw_socket(&self) -> RawSocket {
563        self.get_ref().0.as_raw_socket()
564    }
565}