1use std::io::{self, BufRead as _, IoSlice, Read, Write};
2use std::ops::{Deref, DerefMut};
3use std::pin::Pin;
4use std::task::{Context, Poll};
5
6use rustls::{ConnectionCommon, SideData};
7use tokio::io::{AsyncBufRead, AsyncRead, AsyncWrite, ReadBuf};
8
9mod handshake;
10pub(crate) use handshake::{IoSession, MidHandshake};
11
12#[derive(Debug)]
13pub(crate) enum TlsState {
14 #[cfg(feature = "early-data")]
15 EarlyData(usize, Vec<u8>),
16 Stream,
17 ReadShutdown,
18 WriteShutdown,
19 FullyShutdown,
20}
21
22impl TlsState {
23 #[inline]
24 pub(crate) fn shutdown_read(&mut self) {
25 match *self {
26 Self::WriteShutdown | Self::FullyShutdown => *self = Self::FullyShutdown,
27 _ => *self = Self::ReadShutdown,
28 }
29 }
30
31 #[inline]
32 pub(crate) fn shutdown_write(&mut self) {
33 match *self {
34 Self::ReadShutdown | Self::FullyShutdown => *self = Self::FullyShutdown,
35 _ => *self = Self::WriteShutdown,
36 }
37 }
38
39 #[inline]
40 pub(crate) fn writeable(&self) -> bool {
41 !matches!(*self, Self::WriteShutdown | Self::FullyShutdown)
42 }
43
44 #[inline]
45 pub(crate) fn readable(&self) -> bool {
46 !matches!(*self, Self::ReadShutdown | Self::FullyShutdown)
47 }
48
49 #[inline]
50 #[cfg(feature = "early-data")]
51 pub(crate) fn is_early_data(&self) -> bool {
52 matches!(self, Self::EarlyData(..))
53 }
54
55 #[inline]
56 #[cfg(not(feature = "early-data"))]
57 pub(crate) const fn is_early_data(&self) -> bool {
58 false
59 }
60}
61
62pub(crate) struct Stream<'a, IO, C> {
63 pub(crate) io: &'a mut IO,
64 pub(crate) session: &'a mut C,
65 pub(crate) eof: bool,
66 pub(crate) need_flush: bool,
67}
68
69impl<'a, IO: AsyncRead + AsyncWrite + Unpin, C, SD> Stream<'a, IO, C>
70where
71 C: DerefMut + Deref<Target = ConnectionCommon<SD>>,
72 SD: SideData,
73{
74 pub(crate) fn new(io: &'a mut IO, session: &'a mut C) -> Self {
75 Stream {
76 io,
77 session,
78 eof: false,
81 need_flush: false,
83 }
84 }
85
86 pub(crate) fn set_eof(mut self, eof: bool) -> Self {
87 self.eof = eof;
88 self
89 }
90
91 pub(crate) fn set_need_flush(mut self, need_flush: bool) -> Self {
92 self.need_flush = need_flush;
93 self
94 }
95
96 pub(crate) fn as_mut_pin(&mut self) -> Pin<&mut Self> {
97 Pin::new(self)
98 }
99
100 pub(crate) fn read_io(&mut self, cx: &mut Context) -> Poll<io::Result<usize>> {
101 let mut reader = SyncReadAdapter { io: self.io, cx };
102
103 let n = match self.session.read_tls(&mut reader) {
104 Ok(n) => n,
105 Err(ref err) if err.kind() == io::ErrorKind::WouldBlock => return Poll::Pending,
106 Err(err) => return Poll::Ready(Err(err)),
107 };
108
109 self.session.process_new_packets().map_err(|err| {
110 let _ = self.write_io(cx);
114
115 io::Error::new(io::ErrorKind::InvalidData, err)
116 })?;
117
118 Poll::Ready(Ok(n))
119 }
120
121 pub(crate) fn write_io(&mut self, cx: &mut Context) -> Poll<io::Result<usize>> {
122 let mut writer = SyncWriteAdapter { io: self.io, cx };
123
124 match self.session.write_tls(&mut writer) {
125 Err(ref err) if err.kind() == io::ErrorKind::WouldBlock => Poll::Pending,
126 result => Poll::Ready(result),
127 }
128 }
129
130 pub(crate) fn handshake(&mut self, cx: &mut Context) -> Poll<io::Result<(usize, usize)>> {
131 let mut wrlen = 0;
132 let mut rdlen = 0;
133
134 loop {
135 let mut write_would_block = false;
136 let mut read_would_block = false;
137
138 while self.session.wants_write() {
139 match self.write_io(cx) {
140 Poll::Ready(Ok(0)) => return Poll::Ready(Err(io::ErrorKind::WriteZero.into())),
141 Poll::Ready(Ok(n)) => {
142 wrlen += n;
143 self.need_flush = true;
144 }
145 Poll::Pending => {
146 write_would_block = true;
147 break;
148 }
149 Poll::Ready(Err(err)) => return Poll::Ready(Err(err)),
150 }
151 }
152
153 if self.need_flush {
154 match Pin::new(&mut self.io).poll_flush(cx) {
155 Poll::Ready(Ok(())) => self.need_flush = false,
156 Poll::Ready(Err(err)) => return Poll::Ready(Err(err)),
157 Poll::Pending => write_would_block = true,
158 }
159 }
160
161 while !self.eof && self.session.wants_read() {
162 match self.read_io(cx) {
163 Poll::Ready(Ok(0)) => self.eof = true,
164 Poll::Ready(Ok(n)) => rdlen += n,
165 Poll::Pending => {
166 read_would_block = true;
167 break;
168 }
169 Poll::Ready(Err(err)) => return Poll::Ready(Err(err)),
170 }
171 }
172
173 return match (self.eof, self.session.is_handshaking()) {
174 (true, true) => {
175 let err = io::Error::new(io::ErrorKind::UnexpectedEof, "tls handshake eof");
176 Poll::Ready(Err(err))
177 }
178 (_, false) => Poll::Ready(Ok((rdlen, wrlen))),
179 (_, true) if write_would_block || read_would_block => {
180 if rdlen != 0 || wrlen != 0 {
181 Poll::Ready(Ok((rdlen, wrlen)))
182 } else {
183 Poll::Pending
184 }
185 }
186 (..) => continue,
187 };
188 }
189 }
190
191 pub(crate) fn poll_fill_buf(mut self, cx: &mut Context<'_>) -> Poll<io::Result<&'a [u8]>>
192 where
193 SD: 'a,
194 {
195 let mut io_pending = false;
196
197 while !self.eof && self.session.wants_read() {
199 match self.read_io(cx) {
200 Poll::Ready(Ok(0)) => {
201 break;
202 }
203 Poll::Ready(Ok(_)) => (),
204 Poll::Pending => {
205 io_pending = true;
206 break;
207 }
208 Poll::Ready(Err(err)) => return Poll::Ready(Err(err)),
209 }
210 }
211
212 match self.session.reader().into_first_chunk() {
213 Ok(buf) => {
214 Poll::Ready(Ok(buf))
217 }
218 Err(e) if e.kind() == io::ErrorKind::WouldBlock => {
219 if !io_pending {
220 cx.waker().wake_by_ref();
225 }
226
227 Poll::Pending
228 }
229 Err(e) => Poll::Ready(Err(e)),
230 }
231 }
232}
233
234impl<'a, IO: AsyncRead + AsyncWrite + Unpin, C, SD> AsyncRead for Stream<'a, IO, C>
235where
236 C: DerefMut + Deref<Target = ConnectionCommon<SD>>,
237 SD: SideData + 'a,
238{
239 fn poll_read(
240 mut self: Pin<&mut Self>,
241 cx: &mut Context<'_>,
242 buf: &mut ReadBuf<'_>,
243 ) -> Poll<io::Result<()>> {
244 let data = ready!(self.as_mut().poll_fill_buf(cx))?;
245 let amount = buf.remaining().min(data.len());
246 buf.put_slice(&data[..amount]);
247 self.session.reader().consume(amount);
248 Poll::Ready(Ok(()))
249 }
250}
251
252impl<'a, IO: AsyncRead + AsyncWrite + Unpin, C, SD> AsyncBufRead for Stream<'a, IO, C>
253where
254 C: DerefMut + Deref<Target = ConnectionCommon<SD>>,
255 SD: SideData + 'a,
256{
257 fn poll_fill_buf(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<&[u8]>> {
258 let this = self.get_mut();
259 Stream {
260 io: this.io,
262 session: this.session,
263 ..*this
264 }
265 .poll_fill_buf(cx)
266 }
267
268 fn consume(mut self: Pin<&mut Self>, amt: usize) {
269 self.session.reader().consume(amt);
270 }
271}
272
273impl<IO: AsyncRead + AsyncWrite + Unpin, C, SD> AsyncWrite for Stream<'_, IO, C>
274where
275 C: DerefMut + Deref<Target = ConnectionCommon<SD>>,
276 SD: SideData,
277{
278 fn poll_write(
279 mut self: Pin<&mut Self>,
280 cx: &mut Context,
281 buf: &[u8],
282 ) -> Poll<io::Result<usize>> {
283 let mut pos = 0;
284
285 while pos != buf.len() {
286 let mut would_block = false;
287
288 match self.session.writer().write(&buf[pos..]) {
289 Ok(n) => pos += n,
290 Err(err) => return Poll::Ready(Err(err)),
291 };
292
293 while self.session.wants_write() {
294 match self.write_io(cx) {
295 Poll::Ready(Ok(0)) | Poll::Pending => {
296 would_block = true;
297 break;
298 }
299 Poll::Ready(Ok(_)) => (),
300 Poll::Ready(Err(err)) => return Poll::Ready(Err(err)),
301 }
302 }
303
304 return match (pos, would_block) {
305 (0, true) => Poll::Pending,
306 (n, true) => Poll::Ready(Ok(n)),
307 (_, false) => continue,
308 };
309 }
310
311 Poll::Ready(Ok(pos))
312 }
313
314 fn poll_write_vectored(
315 mut self: Pin<&mut Self>,
316 cx: &mut Context<'_>,
317 bufs: &[IoSlice<'_>],
318 ) -> Poll<io::Result<usize>> {
319 if bufs.iter().all(|buf| buf.is_empty()) {
320 return Poll::Ready(Ok(0));
321 }
322
323 loop {
324 let mut would_block = false;
325 let written = self.session.writer().write_vectored(bufs)?;
326
327 while self.session.wants_write() {
328 match self.write_io(cx) {
329 Poll::Ready(Ok(0)) | Poll::Pending => {
330 would_block = true;
331 break;
332 }
333 Poll::Ready(Ok(_)) => (),
334 Poll::Ready(Err(err)) => return Poll::Ready(Err(err)),
335 }
336 }
337
338 return match (written, would_block) {
339 (0, true) => Poll::Pending,
340 (0, false) => continue,
341 (n, _) => Poll::Ready(Ok(n)),
342 };
343 }
344 }
345
346 #[inline]
347 fn is_write_vectored(&self) -> bool {
348 true
349 }
350
351 fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<io::Result<()>> {
352 self.session.writer().flush()?;
353 while self.session.wants_write() {
354 if ready!(self.write_io(cx))? == 0 {
355 return Poll::Ready(Err(io::ErrorKind::WriteZero.into()));
356 }
357 }
358 Pin::new(&mut self.io).poll_flush(cx)
359 }
360
361 fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
362 while self.session.wants_write() {
363 if ready!(self.write_io(cx))? == 0 {
364 return Poll::Ready(Err(io::ErrorKind::WriteZero.into()));
365 }
366 }
367
368 Poll::Ready(match ready!(Pin::new(&mut self.io).poll_shutdown(cx)) {
369 Ok(()) => Ok(()),
370 Err(err) if err.kind() == io::ErrorKind::NotConnected => Ok(()),
372 Err(err) => Err(err),
373 })
374 }
375}
376
377pub(crate) struct SyncReadAdapter<'a, 'b, T> {
382 pub(crate) io: &'a mut T,
383 pub(crate) cx: &'a mut Context<'b>,
384}
385
386impl<T: AsyncRead + Unpin> Read for SyncReadAdapter<'_, '_, T> {
387 #[inline]
388 fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
389 let mut buf = ReadBuf::new(buf);
390 match Pin::new(&mut self.io).poll_read(self.cx, &mut buf) {
391 Poll::Ready(Ok(())) => Ok(buf.filled().len()),
392 Poll::Ready(Err(err)) => Err(err),
393 Poll::Pending => Err(io::ErrorKind::WouldBlock.into()),
394 }
395 }
396}
397
398pub(crate) struct SyncWriteAdapter<'a, 'b, T> {
403 pub(crate) io: &'a mut T,
404 pub(crate) cx: &'a mut Context<'b>,
405}
406
407impl<T: Unpin> SyncWriteAdapter<'_, '_, T> {
408 #[inline]
409 fn poll_with<U>(
410 &mut self,
411 f: impl FnOnce(Pin<&mut T>, &mut Context<'_>) -> Poll<io::Result<U>>,
412 ) -> io::Result<U> {
413 match f(Pin::new(self.io), self.cx) {
414 Poll::Ready(result) => result,
415 Poll::Pending => Err(io::ErrorKind::WouldBlock.into()),
416 }
417 }
418}
419
420impl<T: AsyncWrite + Unpin> Write for SyncWriteAdapter<'_, '_, T> {
421 #[inline]
422 fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
423 self.poll_with(|io, cx| io.poll_write(cx, buf))
424 }
425
426 #[inline]
427 fn write_vectored(&mut self, bufs: &[IoSlice<'_>]) -> io::Result<usize> {
428 self.poll_with(|io, cx| io.poll_write_vectored(cx, bufs))
429 }
430
431 fn flush(&mut self) -> io::Result<()> {
432 self.poll_with(|io, cx| io.poll_flush(cx))
433 }
434}
435
436#[cfg(test)]
437mod test_stream;