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#[derive(Clone)]
21pub struct TlsConnector {
22 inner: Arc<ClientConfig>,
23 #[cfg(feature = "early-data")]
24 early_data: bool,
25}
26
27impl TlsConnector {
28 #[cfg(feature = "early-data")]
33 pub fn early_data(mut self, flag: bool) -> Self {
34 self.early_data = flag;
35 self
36 }
37
38 #[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 #[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 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 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 #[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 #[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
182pub 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
229pub struct FallibleConnect<IO>(MidHandshake<TlsStream<IO>>);
231
232#[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 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 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 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 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 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 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 while stream.session.is_handshaking() {
577 ready!(stream.handshake(cx))?;
578 }
579
580 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 *state = TlsState::Stream;
590
591 if let Some(waker) = early_waker.take() {
592 waker.wake();
593 }
594 }
595
596 Poll::Ready(Ok(0))
597}