1#![deny(missing_docs, missing_debug_implementations, unsafe_code)]
7#![warn(unreachable_pub, unused_qualifications, unused_lifetimes)]
8#![warn(
9 clippy::must_use_candidate,
10 clippy::unwrap_in_result,
11 clippy::panic_in_result_fn
12)]
13
14use futures_io::{AsyncRead, AsyncWrite};
15use openssl::{
16 error::ErrorStack,
17 ssl::{self, ErrorCode, ShutdownResult, Ssl, SslRef},
18};
19use std::{
20 fmt, future,
21 io::{self, Read, Write},
22 pin::Pin,
23 sync::Arc,
24 task::{Context, Poll, Wake, Waker},
25};
26
27#[cfg(test)]
28mod test;
29
30struct StreamWrapper<S: Unpin> {
31 stream: S,
32 waker: Option<Waker>,
33}
34
35impl<S> fmt::Debug for StreamWrapper<S>
36where
37 S: fmt::Debug + Unpin,
38{
39 fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
40 self.stream.fmt(fmt)
41 }
42}
43
44impl<S: Unpin> StreamWrapper<S> {
45 fn parts(&mut self) -> (Pin<&mut S>, Context<'_>) {
46 let stream = Pin::new(&mut self.stream);
47 let context = Context::from_waker(self.waker.as_ref().unwrap_or(Waker::noop()));
50 (stream, context)
51 }
52}
53
54impl<S> Read for StreamWrapper<S>
55where
56 S: AsyncRead + Unpin,
57{
58 fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
59 let (stream, mut cx) = self.parts();
60 match stream.poll_read(&mut cx, buf)? {
61 Poll::Ready(nread) => Ok(nread),
62 Poll::Pending => Err(io::Error::from(io::ErrorKind::WouldBlock)),
63 }
64 }
65}
66
67impl<S> Write for StreamWrapper<S>
68where
69 S: AsyncWrite + Unpin,
70{
71 fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
72 let (stream, mut cx) = self.parts();
73 match stream.poll_write(&mut cx, buf) {
74 Poll::Ready(r) => r,
75 Poll::Pending => Err(io::Error::from(io::ErrorKind::WouldBlock)),
76 }
77 }
78
79 fn flush(&mut self) -> io::Result<()> {
80 let (stream, mut cx) = self.parts();
81 match stream.poll_flush(&mut cx) {
82 Poll::Ready(r) => r,
83 Poll::Pending => Err(io::Error::from(io::ErrorKind::WouldBlock)),
84 }
85 }
86}
87
88fn cvt<T>(r: io::Result<T>) -> Poll<io::Result<T>> {
89 match r {
90 Ok(v) => Poll::Ready(Ok(v)),
91 Err(ref e) if e.kind() == io::ErrorKind::WouldBlock => Poll::Pending,
92 Err(e) => Poll::Ready(Err(e)),
93 }
94}
95
96fn cvt_ossl<T>(r: Result<T, ssl::Error>) -> Poll<Result<T, ssl::Error>> {
97 match r {
98 Ok(v) => Poll::Ready(Ok(v)),
99 Err(e) => match e.code() {
100 ErrorCode::WANT_READ | ErrorCode::WANT_WRITE
101 if e.io_error()
102 .is_none_or(|io_error| io_error.kind() == io::ErrorKind::WouldBlock) =>
103 {
104 Poll::Pending
105 }
106 _ => Poll::Ready(Err(e)),
107 },
108 }
109}
110
111fn ssl_error_to_io(e: ssl::Error) -> io::Error {
112 e.into_io_error().unwrap_or_else(io::Error::other)
113}
114
115const WRITE_CHUNK_SIZE: usize = 16 * 1024;
116
117struct WakeBoth(Waker, Waker);
118
119impl Wake for WakeBoth {
120 fn wake(self: Arc<Self>) {
121 self.wake_by_ref();
122 }
123
124 fn wake_by_ref(self: &Arc<Self>) {
125 self.0.wake_by_ref();
126 self.1.wake_by_ref();
127 }
128}
129
130struct PendingWrite {
131 bytes: Vec<u8>,
132 caller_addr: usize,
133 caller_len: usize,
134 operation_id: Option<u64>,
135 accepted: usize,
136 waker: Option<Waker>,
137 other_waker: Option<Waker>,
138}
139
140impl PendingWrite {
141 fn new(buf: &[u8], operation_id: Option<u64>, mut bytes: Vec<u8>) -> Self {
142 bytes.clear();
143 bytes.extend_from_slice(&buf[..buf.len().min(WRITE_CHUNK_SIZE)]);
144 Self {
145 bytes,
146 caller_addr: buf.as_ptr() as usize,
147 caller_len: buf.len(),
148 operation_id,
149 accepted: 0,
150 waker: None,
151 other_waker: None,
152 }
153 }
154
155 fn set_waker(&mut self, waker: &Waker) {
156 match &mut self.waker {
157 Some(current) => current.clone_from(waker),
158 slot @ None => *slot = Some(waker.clone()),
159 }
160 }
161
162 fn set_other_waker(&mut self, waker: &Waker) {
163 if self
164 .waker
165 .as_ref()
166 .is_some_and(|writer| writer.will_wake(waker))
167 {
168 return;
169 }
170 match &mut self.other_waker {
171 Some(current) => current.clone_from(waker),
172 slot @ None => *slot = Some(waker.clone()),
173 }
174 }
175
176 fn wake_waiters(&mut self) {
177 if let Some(waker) = self.waker.take() {
178 waker.wake();
179 }
180 if let Some(waker) = self.other_waker.take() {
181 waker.wake();
182 }
183 }
184
185 fn combined_waker(&self) -> Option<Waker> {
186 Some(Waker::from(Arc::new(WakeBoth(
187 self.waker.as_ref()?.clone(),
188 self.other_waker.as_ref()?.clone(),
189 ))))
190 }
191
192 fn is_same_write(&self, buf: &[u8], operation_id: Option<u64>) -> bool {
193 match operation_id {
194 Some(id) => self.operation_id == Some(id),
195 None => {
196 self.operation_id.is_none()
199 && self.caller_addr == buf.as_ptr() as usize
200 && self.caller_len == buf.len()
201 && self.bytes.as_slice() == &buf[..self.bytes.len()]
202 }
203 }
204 }
205
206 fn staged(&self) -> &[u8] {
207 &self.bytes
208 }
209}
210
211#[derive(Default)]
212struct WriteState {
213 pending: Option<PendingWrite>,
214 completed: Option<PendingWrite>,
216 buffer: Vec<u8>,
217}
218
219impl WriteState {
220 fn register_writer(&mut self, buf: &[u8], operation_id: Option<u64>, waker: &Waker) {
221 if let Some(pending) = &mut self.pending {
222 if pending.is_same_write(buf, operation_id) {
223 pending.set_waker(waker);
224 } else {
225 pending.set_other_waker(waker);
226 }
227 }
228 }
229
230 fn take_completion(&mut self, buf: &[u8], operation_id: Option<u64>) -> Option<usize> {
231 let completed = self.completed.take()?;
232 let accepted = completed
233 .is_same_write(buf, operation_id)
234 .then_some(completed.accepted);
235 self.buffer = completed.bytes;
236 accepted
237 }
238
239 fn recycle_completion(&mut self) {
240 if let Some(completed) = self.completed.take() {
241 self.buffer = completed.bytes;
242 }
243 }
244
245 fn finish_initial(
246 &mut self,
247 mut pending: PendingWrite,
248 waker: &Waker,
249 result: Poll<Result<usize, ssl::Error>>,
250 ) -> Poll<Result<usize, ssl::Error>> {
251 match result {
252 Poll::Pending => {
253 pending.set_waker(waker);
254 self.pending = Some(pending);
255 Poll::Pending
256 }
257 ready => {
258 self.buffer = pending.bytes;
259 ready
260 }
261 }
262 }
263
264 fn finish_retry(
265 &mut self,
266 mut pending: PendingWrite,
267 result: Poll<Result<usize, ssl::Error>>,
268 zero_is_success: bool,
269 ) -> Poll<Result<(), ssl::Error>> {
270 match result {
271 Poll::Pending => {
272 self.pending = Some(pending);
273 Poll::Pending
274 }
275 Poll::Ready(Ok(n)) => {
276 pending.wake_waiters();
277 if n == 0 && !(zero_is_success && pending.bytes.is_empty()) {
278 self.buffer = pending.bytes;
279 Poll::Ready(Err(ErrorStack::get().into()))
280 } else {
281 pending.accepted = n;
283 self.completed = Some(pending);
284 Poll::Ready(Ok(()))
285 }
286 }
287 Poll::Ready(Err(error)) => {
288 pending.wake_waiters();
289 self.buffer = pending.bytes;
290 Poll::Ready(Err(error))
291 }
292 }
293 }
294}
295
296#[derive(Clone, Copy)]
297enum WriteKind {
298 Normal,
299 #[cfg(ossl111)]
300 Early,
301}
302
303impl WriteKind {
304 fn zero_is_success(self) -> bool {
305 match self {
306 Self::Normal => false,
307 #[cfg(ossl111)]
308 Self::Early => true,
309 }
310 }
311}
312
313pub struct SslStream<S: Unpin> {
324 inner: ssl::SslStream<StreamWrapper<S>>,
325 write: WriteState,
327 #[cfg(ossl111)]
328 early_write: WriteState,
329 next_write_id: u64,
330 close_notify_sent: bool,
334}
335
336impl<S> fmt::Debug for SslStream<S>
337where
338 S: fmt::Debug + Unpin,
339{
340 fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
341 fmt.debug_tuple("SslStream").field(&self.inner).finish()
342 }
343}
344
345impl<S> SslStream<S>
346where
347 S: AsyncRead + AsyncWrite + Unpin,
348{
349 pub fn new(ssl: Ssl, stream: S) -> Result<Self, ErrorStack> {
351 ssl::SslStream::new(
352 ssl,
353 StreamWrapper {
354 stream,
355 waker: None,
356 },
357 )
358 .map(|inner| SslStream {
359 inner,
360 write: WriteState::default(),
361 #[cfg(ossl111)]
362 early_write: WriteState::default(),
363 next_write_id: 0,
364 close_notify_sent: false,
365 })
366 }
367
368 pub async fn write_cancellable(mut self: Pin<&mut Self>, buf: &[u8]) -> io::Result<usize> {
374 let id = self.as_mut().get_mut().allocate_write_id();
375 future::poll_fn(|cx| self.as_mut().poll_write_inner(cx, buf, Some(id))).await
376 }
377
378 pub fn poll_connect(
380 mut self: Pin<&mut Self>,
381 cx: &mut Context<'_>,
382 ) -> Poll<Result<(), ssl::Error>> {
383 std::task::ready!(self.as_mut().poll_finish_pending_writes(cx))?;
384 self.as_mut().with_context(cx, |s| cvt_ossl(s.connect()))
385 }
386
387 pub async fn connect(mut self: Pin<&mut Self>) -> Result<(), ssl::Error> {
389 future::poll_fn(|cx| self.as_mut().poll_connect(cx)).await
390 }
391
392 pub fn poll_accept(
394 mut self: Pin<&mut Self>,
395 cx: &mut Context<'_>,
396 ) -> Poll<Result<(), ssl::Error>> {
397 std::task::ready!(self.as_mut().poll_finish_pending_writes(cx))?;
398 self.as_mut().with_context(cx, |s| cvt_ossl(s.accept()))
399 }
400
401 pub async fn accept(mut self: Pin<&mut Self>) -> Result<(), ssl::Error> {
403 future::poll_fn(|cx| self.as_mut().poll_accept(cx)).await
404 }
405
406 pub fn poll_do_handshake(
408 mut self: Pin<&mut Self>,
409 cx: &mut Context<'_>,
410 ) -> Poll<Result<(), ssl::Error>> {
411 std::task::ready!(self.as_mut().poll_finish_pending_writes(cx))?;
412 self.as_mut()
413 .with_context(cx, |s| cvt_ossl(s.do_handshake()))
414 }
415
416 pub async fn do_handshake(mut self: Pin<&mut Self>) -> Result<(), ssl::Error> {
418 future::poll_fn(|cx| self.as_mut().poll_do_handshake(cx)).await
419 }
420
421 pub fn poll_peek(
423 mut self: Pin<&mut Self>,
424 cx: &mut Context<'_>,
425 buf: &mut [u8],
426 ) -> Poll<Result<usize, ssl::Error>> {
427 if buf.is_empty() {
431 return Poll::Ready(Ok(0));
432 }
433 std::task::ready!(self.as_mut().poll_finish_pending_writes(cx))?;
434 self.as_mut()
435 .with_context(cx, |s| cvt_ossl(s.ssl_peek(buf)))
436 }
437
438 pub async fn peek(mut self: Pin<&mut Self>, buf: &mut [u8]) -> Result<usize, ssl::Error> {
440 future::poll_fn(|cx| self.as_mut().poll_peek(cx, buf)).await
441 }
442
443 #[cfg(ossl111)]
445 pub fn poll_read_early_data(
446 mut self: Pin<&mut Self>,
447 cx: &mut Context<'_>,
448 buf: &mut [u8],
449 ) -> Poll<Result<usize, ssl::Error>> {
450 std::task::ready!(self.as_mut().poll_finish_pending_writes(cx))?;
451 self.with_context(cx, |s| cvt_ossl(s.read_early_data(buf)))
452 }
453
454 #[cfg(ossl111)]
456 pub async fn read_early_data(
457 mut self: Pin<&mut Self>,
458 buf: &mut [u8],
459 ) -> Result<usize, ssl::Error> {
460 future::poll_fn(|cx| self.as_mut().poll_read_early_data(cx, buf)).await
461 }
462
463 #[cfg(ossl111)]
469 pub fn poll_write_early_data(
470 self: Pin<&mut Self>,
471 cx: &mut Context<'_>,
472 buf: &[u8],
473 ) -> Poll<Result<usize, ssl::Error>> {
474 self.poll_write_operation(cx, buf, None, WriteKind::Early)
475 }
476
477 #[cfg(ossl111)]
483 pub async fn write_early_data(
484 mut self: Pin<&mut Self>,
485 buf: &[u8],
486 ) -> Result<usize, ssl::Error> {
487 let id = self.as_mut().get_mut().allocate_write_id();
488 future::poll_fn(|cx| {
489 self.as_mut()
490 .poll_write_operation(cx, buf, Some(id), WriteKind::Early)
491 })
492 .await
493 }
494}
495
496impl<S: Unpin> SslStream<S> {
497 #[must_use]
499 pub fn ssl(&self) -> &SslRef {
500 self.inner.ssl()
501 }
502
503 #[must_use]
505 pub fn get_ref(&self) -> &S {
506 &self.inner.get_ref().stream
507 }
508
509 pub fn get_mut(&mut self) -> &mut S {
515 &mut self.inner.get_mut().stream
516 }
517
518 #[must_use]
524 pub fn get_pin_mut(self: Pin<&mut Self>) -> Pin<&mut S> {
525 Pin::new(&mut self.get_mut().inner.get_mut().stream)
526 }
527
528 fn with_context<F, R>(self: Pin<&mut Self>, ctx: &mut Context<'_>, f: F) -> R
529 where
530 F: FnOnce(&mut ssl::SslStream<StreamWrapper<S>>) -> R,
531 {
532 let this = self.get_mut();
533 match &mut this.inner.get_mut().waker {
534 Some(waker) => waker.clone_from(ctx.waker()),
537 waker @ None => *waker = Some(ctx.waker().clone()),
538 }
539 f(&mut this.inner)
540 }
541}
542
543impl<S> AsyncRead for SslStream<S>
544where
545 S: AsyncRead + AsyncWrite + Unpin,
546{
547 fn poll_read(
548 mut self: Pin<&mut Self>,
549 ctx: &mut Context<'_>,
550 buf: &mut [u8],
551 ) -> Poll<io::Result<usize>> {
552 if buf.is_empty() {
553 return Poll::Ready(Ok(0));
554 }
555 std::task::ready!(self.as_mut().poll_finish_pending_writes(ctx))
556 .map_err(ssl_error_to_io)?;
557 self.as_mut().with_context(ctx, |s| cvt(s.read(buf)))
558 }
559}
560
561impl<S> AsyncWrite for SslStream<S>
562where
563 S: AsyncRead + AsyncWrite + Unpin,
564{
565 fn poll_write(self: Pin<&mut Self>, ctx: &mut Context, buf: &[u8]) -> Poll<io::Result<usize>> {
566 self.poll_write_inner(ctx, buf, None)
567 }
568
569 fn poll_flush(mut self: Pin<&mut Self>, ctx: &mut Context) -> Poll<io::Result<()>> {
570 std::task::ready!(self.as_mut().poll_finish_pending_writes(ctx))
571 .map_err(ssl_error_to_io)?;
572 self.as_mut().get_mut().recycle_completed_writes();
573 self.with_context(ctx, |s| cvt(s.flush()))
574 }
575
576 fn poll_close(mut self: Pin<&mut Self>, ctx: &mut Context) -> Poll<io::Result<()>> {
577 std::task::ready!(self.as_mut().poll_finish_pending_writes(ctx))
578 .map_err(ssl_error_to_io)?;
579 self.as_mut().get_mut().recycle_completed_writes();
580 if !self.close_notify_sent {
590 match self.as_mut().with_context(ctx, |s| s.shutdown()) {
591 Ok(ShutdownResult::Sent | ShutdownResult::Received) => {}
592 Err(ref e) if e.code() == ErrorCode::ZERO_RETURN => {}
593 Err(ref e)
594 if e.code() == ErrorCode::WANT_READ || e.code() == ErrorCode::WANT_WRITE =>
595 {
596 return Poll::Pending;
597 }
598 Err(e) => {
599 return Poll::Ready(Err(ssl_error_to_io(e)));
600 }
601 }
602 self.as_mut().get_mut().close_notify_sent = true;
603 }
604
605 self.get_pin_mut().poll_close(ctx)
606 }
607}
608
609impl<S> SslStream<S>
610where
611 S: AsyncRead + AsyncWrite + Unpin,
612{
613 fn poll_ssl_write_kind(
614 self: Pin<&mut Self>,
615 ctx: &mut Context<'_>,
616 buf: &[u8],
617 kind: WriteKind,
618 ) -> Poll<Result<usize, ssl::Error>> {
619 match kind {
620 WriteKind::Normal => self.with_context(ctx, |s| {
621 loop {
622 match s.ssl_write(buf) {
623 Err(ref e)
626 if e.code() == ErrorCode::WANT_READ && e.io_error().is_none() => {}
627 result => break cvt_ossl(result),
628 }
629 }
630 }),
631 #[cfg(ossl111)]
632 WriteKind::Early => self.with_context(ctx, |s| cvt_ossl(s.write_early_data(buf))),
633 }
634 }
635
636 fn allocate_write_id(&mut self) -> u64 {
637 let id = self.next_write_id;
638 self.next_write_id = id.wrapping_add(1);
639 id
640 }
641
642 fn write_state_mut(&mut self, kind: WriteKind) -> &mut WriteState {
643 match kind {
644 WriteKind::Normal => &mut self.write,
645 #[cfg(ossl111)]
646 WriteKind::Early => &mut self.early_write,
647 }
648 }
649
650 fn recycle_completed_writes(&mut self) {
651 self.write.recycle_completion();
652 #[cfg(ossl111)]
653 self.early_write.recycle_completion();
654 }
655
656 fn poll_write_inner(
657 self: Pin<&mut Self>,
658 ctx: &mut Context<'_>,
659 buf: &[u8],
660 operation_id: Option<u64>,
661 ) -> Poll<io::Result<usize>> {
662 self.poll_write_operation(ctx, buf, operation_id, WriteKind::Normal)
663 .map(|result| result.map_err(ssl_error_to_io))
664 }
665
666 fn poll_write_operation(
667 mut self: Pin<&mut Self>,
668 ctx: &mut Context<'_>,
669 buf: &[u8],
670 operation_id: Option<u64>,
671 kind: WriteKind,
672 ) -> Poll<Result<usize, ssl::Error>> {
673 if matches!(kind, WriteKind::Normal) && buf.is_empty() {
674 return Poll::Ready(Ok(0));
675 }
676
677 #[cfg(ossl111)]
679 {
680 let other = match kind {
681 WriteKind::Normal => WriteKind::Early,
682 WriteKind::Early => WriteKind::Normal,
683 };
684 std::task::ready!(self.as_mut().poll_finish_pending_write_kind(ctx, other))?;
685 }
686
687 self.as_mut()
688 .get_mut()
689 .write_state_mut(kind)
690 .register_writer(buf, operation_id, ctx.waker());
691 std::task::ready!(self.as_mut().poll_finish_pending_write_kind(ctx, kind))?;
692 if let Some(written) = self
693 .as_mut()
694 .get_mut()
695 .write_state_mut(kind)
696 .take_completion(buf, operation_id)
697 {
698 return Poll::Ready(Ok(written));
699 }
700 let bytes = std::mem::take(&mut self.as_mut().get_mut().write_state_mut(kind).buffer);
703 let pending = PendingWrite::new(buf, operation_id, bytes);
704 let result = self
705 .as_mut()
706 .poll_ssl_write_kind(ctx, pending.staged(), kind);
707 self.get_mut()
708 .write_state_mut(kind)
709 .finish_initial(pending, ctx.waker(), result)
710 }
711
712 fn poll_finish_pending_writes(
715 mut self: Pin<&mut Self>,
716 ctx: &mut Context<'_>,
717 ) -> Poll<Result<(), ssl::Error>> {
718 #[cfg(ossl111)]
719 std::task::ready!(
720 self.as_mut()
721 .poll_finish_pending_write_kind(ctx, WriteKind::Early)
722 )?;
723 std::task::ready!(
724 self.as_mut()
725 .poll_finish_pending_write_kind(ctx, WriteKind::Normal)
726 )?;
727 Poll::Ready(Ok(()))
728 }
729
730 fn poll_finish_pending_write_kind(
731 mut self: Pin<&mut Self>,
732 ctx: &mut Context<'_>,
733 kind: WriteKind,
734 ) -> Poll<Result<(), ssl::Error>> {
735 let Some(mut pending) = self.as_mut().get_mut().write_state_mut(kind).pending.take() else {
736 return Poll::Ready(Ok(()));
737 };
738 pending.set_other_waker(ctx.waker());
739 let result = if let Some(waker) = pending.combined_waker() {
740 let mut combined_context = Context::from_waker(&waker);
741 self.as_mut()
742 .poll_ssl_write_kind(&mut combined_context, pending.staged(), kind)
743 } else {
744 self.as_mut()
745 .poll_ssl_write_kind(ctx, pending.staged(), kind)
746 };
747 self.get_mut()
748 .write_state_mut(kind)
749 .finish_retry(pending, result, kind.zero_is_success())
750 }
751}