Skip to main content

tokio_rustls/
client.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;
9#[cfg(feature = "early-data")]
10use std::task::Waker;
11use std::task::{Context, Poll};
12
13use rustls::pki_types::ServerName;
14use rustls::{ClientConfig, ClientConnection};
15use tokio::io::{AsyncBufRead, AsyncRead, AsyncWrite, ReadBuf};
16
17use crate::common::{IoSession, MidHandshake, Stream, TlsState};
18
19/// A wrapper around a `rustls::ClientConfig`, providing an async `connect` method.
20#[derive(Clone)]
21pub struct TlsConnector {
22    inner: Arc<ClientConfig>,
23    #[cfg(feature = "early-data")]
24    early_data: bool,
25}
26
27impl TlsConnector {
28    /// Enable 0-RTT.
29    ///
30    /// If you want to use 0-RTT,
31    /// You must also set `ClientConfig.enable_early_data` to `true`.
32    #[cfg(feature = "early-data")]
33    pub fn early_data(mut self, flag: bool) -> Self {
34        self.early_data = flag;
35        self
36    }
37
38    /// Returns a future that performs a TLS handshake with `domain` using the `stream`.
39    ///
40    /// You likely want to wrap this in a timeout (for example with [`tokio::time::timeout`][])
41    /// to bound the handshake time.
42    ///
43    /// [`tokio::time::timeout`]: https://docs.rs/tokio/latest/tokio/time/fn.timeout.html
44    #[inline]
45    pub fn connect<IO>(&self, domain: ServerName<'static>, stream: IO) -> Connect<IO>
46    where
47        IO: AsyncRead + AsyncWrite + Unpin,
48    {
49        self.connect_impl(domain, stream, None, |_| ())
50    }
51
52    /// Similar to [`Self::connect()`], but calls `f` before performing the handshake.
53    ///
54    /// As with [`Self::connect()`] you likely want to wrap this in a timeout to
55    /// bound the handshake time.
56    ///
57    /// The `f` handler is given a mutable reference to a [`ClientConnection`][] that can
58    /// be used for tasks like writing early data.
59    ///
60    /// [`ClientConnection`]: https://docs.rs/rustls/latest/rustls/client/struct.ClientConnection.html
61    #[inline]
62    pub fn connect_with<IO, F>(&self, domain: ServerName<'static>, stream: IO, f: F) -> Connect<IO>
63    where
64        IO: AsyncRead + AsyncWrite + Unpin,
65        F: FnOnce(&mut ClientConnection),
66    {
67        self.connect_impl(domain, stream, None, f)
68    }
69
70    fn connect_impl<IO, F>(
71        &self,
72        domain: ServerName<'static>,
73        stream: IO,
74        alpn_protocols: Option<Vec<Vec<u8>>>,
75        f: F,
76    ) -> Connect<IO>
77    where
78        IO: AsyncRead + AsyncWrite + Unpin,
79        F: FnOnce(&mut ClientConnection),
80    {
81        let alpn = alpn_protocols.unwrap_or_else(|| self.inner.alpn_protocols.clone());
82        let mut session = match ClientConnection::new_with_alpn(self.inner.clone(), domain, alpn) {
83            Ok(session) => session,
84            Err(error) => {
85                return Connect(MidHandshake::Error {
86                    io: stream,
87                    // TODO(eliza): should this really return an `io::Error`?
88                    // Probably not...
89                    error: io::Error::new(io::ErrorKind::Other, error),
90                });
91            }
92        };
93        f(&mut session);
94
95        Connect(MidHandshake::Handshaking(TlsStream {
96            io: stream,
97
98            #[cfg(not(feature = "early-data"))]
99            state: TlsState::Stream,
100
101            #[cfg(feature = "early-data")]
102            state: if self.early_data && session.early_data().is_some() {
103                TlsState::EarlyData(0, Vec::new())
104            } else {
105                TlsState::Stream
106            },
107
108            need_flush: false,
109            error: None,
110
111            #[cfg(feature = "early-data")]
112            early_waker: None,
113
114            session,
115        }))
116    }
117
118    pub fn with_alpn(&self, alpn_protocols: Vec<Vec<u8>>) -> TlsConnectorWithAlpn<'_> {
119        TlsConnectorWithAlpn {
120            inner: self,
121            alpn_protocols,
122        }
123    }
124
125    /// Get a read-only reference to underlying config
126    pub fn config(&self) -> &Arc<ClientConfig> {
127        &self.inner
128    }
129}
130
131impl From<Arc<ClientConfig>> for TlsConnector {
132    fn from(inner: Arc<ClientConfig>) -> Self {
133        Self {
134            inner,
135            #[cfg(feature = "early-data")]
136            early_data: false,
137        }
138    }
139}
140
141pub struct TlsConnectorWithAlpn<'c> {
142    inner: &'c TlsConnector,
143    alpn_protocols: Vec<Vec<u8>>,
144}
145
146impl TlsConnectorWithAlpn<'_> {
147    /// Returns a future that performs a TLS handshake with `domain` using the `stream`.
148    ///
149    /// You likely want to wrap this in a timeout (for example with [`tokio::time::timeout`][])
150    /// to bound the handshake time.
151    ///
152    /// [`tokio::time::timeout`]: https://docs.rs/tokio/latest/tokio/time/fn.timeout.html
153    #[inline]
154    pub fn connect<IO>(self, domain: ServerName<'static>, stream: IO) -> Connect<IO>
155    where
156        IO: AsyncRead + AsyncWrite + Unpin,
157    {
158        self.inner
159            .connect_impl(domain, stream, Some(self.alpn_protocols), |_| ())
160    }
161
162    /// Similar to [`Self::connect()`], but calls `f` before performing the handshake.
163    ///
164    /// As with [`Self::connect()`] you likely want to wrap this in a timeout to
165    /// bound the handshake time.
166    ///
167    /// The `f` handler is given a mutable reference to a [`ClientConnection`][] that can
168    /// be used for tasks like writing early data.
169    ///
170    /// [`ClientConnection`]: https://docs.rs/rustls/latest/rustls/client/struct.ClientConnection.html
171    #[inline]
172    pub fn connect_with<IO, F>(self, domain: ServerName<'static>, stream: IO, f: F) -> Connect<IO>
173    where
174        IO: AsyncRead + AsyncWrite + Unpin,
175        F: FnOnce(&mut ClientConnection),
176    {
177        self.inner
178            .connect_impl(domain, stream, Some(self.alpn_protocols), f)
179    }
180}
181
182/// Future returned from `TlsConnector::connect` which will resolve
183/// once the connection handshake has finished.
184pub struct Connect<IO>(MidHandshake<TlsStream<IO>>);
185
186impl<IO> Connect<IO> {
187    #[inline]
188    pub fn into_fallible(self) -> FallibleConnect<IO> {
189        FallibleConnect(self.0)
190    }
191
192    pub fn get_ref(&self) -> Option<&IO> {
193        match &self.0 {
194            MidHandshake::Handshaking(sess) => Some(sess.get_ref().0),
195            MidHandshake::SendAlert { io, .. } => Some(io),
196            MidHandshake::Error { io, .. } => Some(io),
197            MidHandshake::End => None,
198        }
199    }
200
201    pub fn get_mut(&mut self) -> Option<&mut IO> {
202        match &mut self.0 {
203            MidHandshake::Handshaking(sess) => Some(sess.get_mut().0),
204            MidHandshake::SendAlert { io, .. } => Some(io),
205            MidHandshake::Error { io, .. } => Some(io),
206            MidHandshake::End => None,
207        }
208    }
209}
210
211impl<IO: AsyncRead + AsyncWrite + Unpin> Future for Connect<IO> {
212    type Output = io::Result<TlsStream<IO>>;
213
214    #[inline]
215    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
216        Pin::new(&mut self.0).poll(cx).map_err(|(err, _)| err)
217    }
218}
219
220impl<IO: AsyncRead + AsyncWrite + Unpin> Future for FallibleConnect<IO> {
221    type Output = Result<TlsStream<IO>, (io::Error, IO)>;
222
223    #[inline]
224    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
225        Pin::new(&mut self.0).poll(cx)
226    }
227}
228
229/// Like [Connect], but returns `IO` on failure.
230pub struct FallibleConnect<IO>(MidHandshake<TlsStream<IO>>);
231
232/// A wrapper around an underlying raw stream which implements the TLS or SSL
233/// protocol.
234#[derive(Debug)]
235pub struct TlsStream<IO> {
236    pub(crate) io: IO,
237    pub(crate) session: ClientConnection,
238    pub(crate) state: TlsState,
239    pub(crate) need_flush: bool,
240    /// Buffered error that occurred during batch reading
241    pub(crate) error: Option<io::Error>,
242
243    #[cfg(feature = "early-data")]
244    pub(crate) early_waker: Option<Waker>,
245}
246
247impl<IO> TlsStream<IO> {
248    #[inline]
249    pub fn get_ref(&self) -> (&IO, &ClientConnection) {
250        (&self.io, &self.session)
251    }
252
253    #[inline]
254    pub fn get_mut(&mut self) -> (&mut IO, &mut ClientConnection) {
255        (&mut self.io, &mut self.session)
256    }
257
258    #[inline]
259    pub fn into_inner(self) -> (IO, ClientConnection) {
260        (self.io, self.session)
261    }
262}
263
264#[cfg(unix)]
265impl<S> AsRawFd for TlsStream<S>
266where
267    S: AsRawFd,
268{
269    fn as_raw_fd(&self) -> RawFd {
270        self.get_ref().0.as_raw_fd()
271    }
272}
273
274#[cfg(windows)]
275impl<S> AsRawSocket for TlsStream<S>
276where
277    S: AsRawSocket,
278{
279    fn as_raw_socket(&self) -> RawSocket {
280        self.get_ref().0.as_raw_socket()
281    }
282}
283
284impl<IO> IoSession for TlsStream<IO> {
285    type Io = IO;
286    type Session = ClientConnection;
287
288    #[inline]
289    fn skip_handshake(&self) -> bool {
290        self.state.is_early_data()
291    }
292
293    #[inline]
294    fn get_mut(&mut self) -> (&mut TlsState, &mut Self::Io, &mut Self::Session, &mut bool) {
295        (
296            &mut self.state,
297            &mut self.io,
298            &mut self.session,
299            &mut self.need_flush,
300        )
301    }
302
303    #[inline]
304    fn into_io(self) -> Self::Io {
305        self.io
306    }
307}
308
309#[cfg(feature = "early-data")]
310impl<IO> TlsStream<IO>
311where
312    IO: AsyncRead + AsyncWrite + Unpin,
313{
314    fn poll_early_data(&mut self, cx: &mut Context<'_>) {
315        // In the EarlyData state, we have not really established a Tls connection.
316        // Before writing data through `AsyncWrite` and completing the tls handshake,
317        // we ignore read readiness and return to pending.
318        //
319        // In order to avoid event loss,
320        // we need to register a waker and wake it up after tls is connected.
321        if self
322            .early_waker
323            .as_ref()
324            .filter(|waker| cx.waker().will_wake(waker))
325            .is_none()
326        {
327            self.early_waker = Some(cx.waker().clone());
328        }
329    }
330}
331
332impl<IO> AsyncRead for TlsStream<IO>
333where
334    IO: AsyncRead + AsyncWrite + Unpin,
335{
336    fn poll_read(
337        mut self: Pin<&mut Self>,
338        cx: &mut Context<'_>,
339        buf: &mut ReadBuf<'_>,
340    ) -> Poll<io::Result<()>> {
341        if let Some(err) = self.error.take() {
342            return Poll::Ready(Err(err));
343        };
344        let data = ready!(self.as_mut().poll_fill_buf(cx))?;
345        let len = data.len().min(buf.remaining());
346        if len == 0 {
347            return Poll::Ready(Ok(()));
348        }
349        buf.put_slice(&data[..len]);
350        self.as_mut().consume(len);
351
352        while buf.remaining() > 0 {
353            let data = match self.as_mut().poll_fill_buf(cx) {
354                Poll::Ready(Ok([])) => break,
355                Poll::Ready(Ok(data)) => data,
356                Poll::Ready(Err(err)) => {
357                    self.error = Some(err);
358                    break;
359                }
360                Poll::Pending => break,
361            };
362            let len = Ord::min(data.len(), buf.remaining());
363            buf.put_slice(&data[..len]);
364            self.as_mut().consume(len);
365        }
366        Poll::Ready(Ok(()))
367    }
368}
369
370impl<IO> AsyncBufRead for TlsStream<IO>
371where
372    IO: AsyncRead + AsyncWrite + Unpin,
373{
374    fn poll_fill_buf(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<&[u8]>> {
375        match self.state {
376            #[cfg(feature = "early-data")]
377            TlsState::EarlyData(..) => {
378                self.get_mut().poll_early_data(cx);
379                Poll::Pending
380            }
381            TlsState::Stream | TlsState::WriteShutdown => {
382                let this = self.get_mut();
383                let stream =
384                    Stream::new(&mut this.io, &mut this.session).set_eof(!this.state.readable());
385
386                match stream.poll_fill_buf(cx) {
387                    Poll::Ready(Ok(buf)) => {
388                        if buf.is_empty() {
389                            this.state.shutdown_read();
390                        }
391
392                        Poll::Ready(Ok(buf))
393                    }
394                    Poll::Ready(Err(err)) if err.kind() == io::ErrorKind::ConnectionAborted => {
395                        this.state.shutdown_read();
396                        Poll::Ready(Err(err))
397                    }
398                    output => output,
399                }
400            }
401            TlsState::ReadShutdown | TlsState::FullyShutdown => Poll::Ready(Ok(&[])),
402        }
403    }
404
405    fn consume(mut self: Pin<&mut Self>, amt: usize) {
406        self.session.reader().consume(amt);
407    }
408}
409
410impl<IO> AsyncWrite for TlsStream<IO>
411where
412    IO: AsyncRead + AsyncWrite + Unpin,
413{
414    /// Note: that it does not guarantee the final data to be sent.
415    /// To be cautious, you must manually call `flush`.
416    fn poll_write(
417        self: Pin<&mut Self>,
418        cx: &mut Context<'_>,
419        buf: &[u8],
420    ) -> Poll<io::Result<usize>> {
421        let this = self.get_mut();
422        let mut stream = Stream::new(&mut this.io, &mut this.session)
423            .set_eof(!this.state.readable())
424            .set_need_flush(this.need_flush);
425
426        #[cfg(feature = "early-data")]
427        {
428            let bufs = [io::IoSlice::new(buf)];
429            let written = poll_handle_early_data(
430                &mut this.state,
431                &mut stream,
432                &mut this.early_waker,
433                cx,
434                &bufs,
435            )?;
436            match written {
437                Poll::Ready(0) => {}
438                Poll::Ready(written) => return Poll::Ready(Ok(written)),
439                Poll::Pending => {
440                    this.need_flush = stream.need_flush;
441                    return Poll::Pending;
442                }
443            }
444        }
445
446        stream.as_mut_pin().poll_write(cx, buf)
447    }
448
449    /// Note: that it does not guarantee the final data to be sent.
450    /// To be cautious, you must manually call `flush`.
451    fn poll_write_vectored(
452        self: Pin<&mut Self>,
453        cx: &mut Context<'_>,
454        bufs: &[io::IoSlice<'_>],
455    ) -> Poll<io::Result<usize>> {
456        let this = self.get_mut();
457        let mut stream = Stream::new(&mut this.io, &mut this.session)
458            .set_eof(!this.state.readable())
459            .set_need_flush(this.need_flush);
460
461        #[cfg(feature = "early-data")]
462        {
463            let written = poll_handle_early_data(
464                &mut this.state,
465                &mut stream,
466                &mut this.early_waker,
467                cx,
468                bufs,
469            )?;
470            match written {
471                Poll::Ready(0) => {}
472                Poll::Ready(written) => return Poll::Ready(Ok(written)),
473                Poll::Pending => {
474                    this.need_flush = stream.need_flush;
475                    return Poll::Pending;
476                }
477            }
478        }
479
480        stream.as_mut_pin().poll_write_vectored(cx, bufs)
481    }
482
483    #[inline]
484    fn is_write_vectored(&self) -> bool {
485        true
486    }
487
488    fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
489        let this = self.get_mut();
490        let mut stream = Stream::new(&mut this.io, &mut this.session)
491            .set_eof(!this.state.readable())
492            .set_need_flush(this.need_flush);
493
494        #[cfg(feature = "early-data")]
495        {
496            let written = poll_handle_early_data(
497                &mut this.state,
498                &mut stream,
499                &mut this.early_waker,
500                cx,
501                &[],
502            )?;
503            if written.is_pending() {
504                this.need_flush = stream.need_flush;
505                return Poll::Pending;
506            }
507        }
508
509        stream.as_mut_pin().poll_flush(cx)
510    }
511
512    fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
513        #[cfg(feature = "early-data")]
514        {
515            // complete handshake
516            if matches!(self.state, TlsState::EarlyData(..)) {
517                ready!(self.as_mut().poll_flush(cx))?;
518            }
519        }
520
521        if self.state.writeable() {
522            self.session.send_close_notify();
523            self.state.shutdown_write();
524        }
525
526        let this = self.get_mut();
527        let mut stream =
528            Stream::new(&mut this.io, &mut this.session).set_eof(!this.state.readable());
529        stream.as_mut_pin().poll_shutdown(cx)
530    }
531}
532
533#[cfg(feature = "early-data")]
534fn poll_handle_early_data<IO>(
535    state: &mut TlsState,
536    stream: &mut Stream<IO, ClientConnection>,
537    early_waker: &mut Option<Waker>,
538    cx: &mut Context<'_>,
539    bufs: &[io::IoSlice<'_>],
540) -> Poll<io::Result<usize>>
541where
542    IO: AsyncRead + AsyncWrite + Unpin,
543{
544    if let TlsState::EarlyData(pos, data) = state {
545        use std::io::Write;
546
547        // write early data
548        if let Some(mut early_data) = stream.session.early_data() {
549            let mut written = 0;
550
551            for buf in bufs {
552                if buf.is_empty() {
553                    continue;
554                }
555
556                let len = match early_data.write(buf) {
557                    Ok(0) => break,
558                    Ok(n) => n,
559                    Err(err) => return Poll::Ready(Err(err)),
560                };
561
562                written += len;
563                data.extend_from_slice(&buf[..len]);
564
565                if len < buf.len() {
566                    break;
567                }
568            }
569
570            if written != 0 {
571                return Poll::Ready(Ok(written));
572            }
573        }
574
575        // complete handshake
576        while stream.session.is_handshaking() {
577            ready!(stream.handshake(cx))?;
578        }
579
580        // write early data (fallback)
581        if !stream.session.is_early_data_accepted() {
582            while *pos < data.len() {
583                let len = ready!(stream.as_mut_pin().poll_write(cx, &data[*pos..]))?;
584                *pos += len;
585            }
586        }
587
588        // end
589        *state = TlsState::Stream;
590
591        if let Some(waker) = early_waker.take() {
592            waker.wake();
593        }
594    }
595
596    Poll::Ready(Ok(0))
597}