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#[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 #[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 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 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 pub fn config(&self) -> &Arc<ServerConfig> {
85 &self.inner
86 }
87}
88
89pub 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 #[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 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#[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 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 pub fn into_stream(self, config: Arc<ServerConfig>) -> Accept<IO> {
273 self.into_stream_with(config, |_| ())
274 }
275
276 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 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
313pub 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
351pub 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#[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 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 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 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}