rust: bssl-tls: Async polling I/Os and Datagrams
As we prepare `bssl-tls-tokio`, we are ready to publish the functions
for performing `async` I/O as public APIs.
Update-Note: DTLS I/O APIs have been not functional fully, but if it
has worked so far, please migrate to the new {a,}sync_{send, recv} APIs
for they are semantically different from stream sockets.
Bug: 532601068
Signed-off-by: Xiangfei Ding <xfding@google.com>
Change-Id: I3f2d7b21f513acbb6f09c4167997058c6a6a6964
Reviewed-on: https://boringssl-review.googlesource.com/c/boringssl/+/100847
Presubmit-BoringSSL-Verified: boringssl-scoped@luci-project-accounts.iam.gserviceaccount.com <boringssl-scoped@luci-project-accounts.iam.gserviceaccount.com>
Reviewed-by: Adam Langley <agl@google.com>
diff --git a/rust/Cargo.lock b/rust/Cargo.lock
index c5e4524..4a8843e 100644
--- a/rust/Cargo.lock
+++ b/rust/Cargo.lock
@@ -66,6 +66,7 @@
"futures",
"hyper",
"hyper-util",
+ "libc",
"tokio",
"tower",
]
diff --git a/rust/bssl-tls-tokio/Cargo.toml b/rust/bssl-tls-tokio/Cargo.toml
index 93dfac3..33b03bc 100644
--- a/rust/bssl-tls-tokio/Cargo.toml
+++ b/rust/bssl-tls-tokio/Cargo.toml
@@ -23,6 +23,9 @@
version = "0.4"
optional = true
+[dependencies.libc]
+version = "0.2"
+
[dev-dependencies]
tokio = { version = "1.0", features = ["full"] }
futures = "0.3"
diff --git a/rust/bssl-tls-tokio/src/lib.rs b/rust/bssl-tls-tokio/src/lib.rs
index c1fe4c2..f8065a7 100644
--- a/rust/bssl-tls-tokio/src/lib.rs
+++ b/rust/bssl-tls-tokio/src/lib.rs
@@ -141,6 +141,7 @@
UseFd, //
};
use bssl_tls::{
+ ReceiveBuffer,
connection::{
Client,
Server,
@@ -156,7 +157,7 @@
io::{
AbstractReader, AbstractSocket, AbstractSocketResult, AbstractWriter, IoStatus,
NoAsyncContext, stdio::PollFor,
- }, //
+ },
};
#[cfg(test)]
@@ -337,6 +338,18 @@
/// Wrapper for datagram sockets to satisfy orphan rule.
pub struct TokioDatagramIo<T>(pub T);
+#[inline]
+fn os_has_no_resource(err: &io::Error) -> bool {
+ #[cfg(unix)]
+ {
+ matches!(err.raw_os_error(), Some(libc::ENOBUFS | libc::ENOMEM))
+ }
+ #[cfg(not(unix))]
+ {
+ false
+ }
+}
+
macro_rules! gen_impl_datagram {
($ty:ty) => {
impl AbstractReader for TokioDatagramIo<$ty> {
@@ -369,7 +382,13 @@
match self.0.poll_send(cx, buf) {
Poll::Pending => AbstractSocketResult::Retry,
Poll::Ready(Ok(len)) => AbstractSocketResult::Ok(len),
- Poll::Ready(Err(e)) => translate_stdio_err(e),
+ Poll::Ready(Err(e)) => {
+ if os_has_no_resource(&e) {
+ AbstractSocketResult::Ok(buf.len())
+ } else {
+ translate_stdio_err(e)
+ }
+ }
}
}
@@ -423,18 +442,21 @@
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
- // Note: This assumes `aread_inner` is made public in `bssl-tls`.
- let status = match self
- .inner
- .as_pin_mut()
- .aread_inner(buf.initialize_unfilled(), cx)
- {
+ let mut recv_buf = ReceiveBuffer::new_uninit(unsafe {
+ // Safety: we will only ever advance the cursor.
+ buf.unfilled_mut()
+ });
+ let status = match self.inner.as_pin_mut().async_poll_read(&mut recv_buf, cx) {
Ok(Some(status)) => status,
Ok(None) => return Poll::Pending,
Err(e) => return Poll::Ready(Err(io::Error::new(io::ErrorKind::Other, e))),
};
match status {
IoStatus::Ok(bytes) => {
+ unsafe {
+ // Safety: BoringSSL filled `bytes` bytes.
+ buf.assume_init(bytes);
+ }
buf.advance(bytes);
Poll::Ready(Ok(()))
}
@@ -453,8 +475,7 @@
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
- // Note: This assumes `awrite_inner` is made public in `bssl-tls`.
- let status = match self.inner.as_pin_mut().awrite_inner(buf, cx) {
+ let status = match self.inner.as_pin_mut().async_poll_write(buf, cx) {
Ok(Some(status)) => status,
Ok(None) => return Poll::Pending,
Err(e) => return Poll::Ready(Err(io::Error::new(io::ErrorKind::Other, e))),
@@ -470,8 +491,7 @@
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
- // Note: This assumes `aflush_inner` is made public in `bssl-tls`.
- let status = match self.inner.as_pin_mut().aflush_inner(cx) {
+ let status = match self.inner.as_pin_mut().async_poll_flush(cx) {
Ok(Some(status)) => status,
Ok(None) => return Poll::Pending,
Err(e) => return Poll::Ready(Err(io::Error::new(io::ErrorKind::Other, e))),
@@ -487,8 +507,7 @@
}
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
- // Note: This assumes `ashutdown_inner` is made public in `bssl-tls`.
- match self.inner.as_pin_mut().ashutdown_inner(cx) {
+ match self.inner.as_pin_mut().async_poll_shutdown(cx) {
Ok(Some(ShutdownStatus::CloseNotifyReceived)) => Poll::Ready(Ok(())),
Ok(Some(ShutdownStatus::RemainingApplicationData)) => Poll::Ready(Err(io::Error::new(
io::ErrorKind::Other,
diff --git a/rust/bssl-tls-tokio/src/tests/datagram.rs b/rust/bssl-tls-tokio/src/tests/datagram.rs
index ed81644..2d54bb4 100644
--- a/rust/bssl-tls-tokio/src/tests/datagram.rs
+++ b/rust/bssl-tls-tokio/src/tests/datagram.rs
@@ -14,7 +14,11 @@
#![cfg(unix)]
+use std::mem::MaybeUninit;
+use std::time::Duration;
+
use bssl_tls::{
+ ReceiveBuffer,
connection::{
Client,
Server,
@@ -24,7 +28,12 @@
DtlsMode,
TlsContextBuilder, //
},
- errors::Error, //
+ errors::Error,
+ io::IoStatus, //
+};
+use tokio::{
+ select,
+ time::sleep, //
};
use crate::{
@@ -40,90 +49,106 @@
server_ctx_builder
.with_credential(super::server_credential())
.unwrap();
- let server_conn = server_ctx_builder.build().new_server_connection().build();
+ let mut server_conn = server_ctx_builder.build().new_server_connection();
+ server_conn.with_mtu(500).unwrap();
+ let server_conn = server_conn.build();
let mut client_ctx_builder = TlsContextBuilder::new_dtls();
client_ctx_builder.with_certificate_store(&super::client_cert_store());
- let client_conn = client_ctx_builder.build().new_client_connection().build();
+ let mut client_conn = client_ctx_builder.build().new_client_connection();
+ client_conn.with_mtu(500).unwrap();
+ let client_conn = client_conn.build();
(server_conn, client_conn)
}
-async fn async_ping_pong(
+async fn drive_async_dtls_handshake<R: Send + 'static>(
+ conn: &mut TlsConnection<R, DtlsMode>,
+) -> Result<(), Error> {
+ loop {
+ let timeout = conn.dtlsv1_get_timeout().unwrap_or(Duration::from_secs(5));
+ select! {
+ biased;
+ res = conn.async_handshake() => match res? {
+ None => break Ok(()),
+ Some(reason) => panic!("unexpected retry reason {reason:?}"),
+ },
+ _ = sleep(timeout) => {
+ conn.dtlsv1_handle_timeout()?;
+ }
+ }
+ }
+}
+
+async fn async_dtls_recv<R: Send + 'static>(
+ conn: &mut TlsConnection<R, DtlsMode>,
+ buf: &mut ReceiveBuffer<'_>,
+) -> Result<IoStatus, Error> {
+ conn.as_pin_mut().async_recv(buf).await
+}
+
+async fn async_dtls_send<R: Send + 'static>(
+ conn: &mut TlsConnection<R, DtlsMode>,
+ data: &[u8],
+) -> Result<IoStatus, Error> {
+ conn.as_pin_mut().async_send(data).await
+}
+
+async fn async_dtls_ping_pong(
mut server_conn: TlsConnection<Server, DtlsMode>,
mut client_conn: TlsConnection<Client, DtlsMode>,
) -> Result<(), Error> {
- use bssl_tls::io::IoStatus;
- use std::time::Duration;
-
let task = tokio::spawn(async move {
- server_conn.async_handshake().await?;
+ drive_async_dtls_handshake(&mut server_conn).await?;
- let mut message = [0; 21];
+ let mut buf = [MaybeUninit::uninit(); 21];
+ let mut message = ReceiveBuffer::new_uninit(&mut buf);
let mut read_bytes = 0;
while read_bytes < 21 {
- match server_conn
- .as_pin_mut()
- .async_read(&mut message[read_bytes..])
- .await?
- {
+ match async_dtls_recv(&mut server_conn, &mut message).await? {
IoStatus::Ok(n) => read_bytes += n,
IoStatus::EndOfStream => break,
_ => {}
}
}
- assert_eq!(&message, b"BoringSSL is awesome!");
- tokio::time::sleep(Duration::from_secs(2)).await;
- server_conn
- .as_pin_mut()
- .async_write(b"Oh yeah definitely!")
- .await?;
- server_conn.as_pin_mut().async_shutdown().await?;
+ assert_eq!(message.filled(), b"BoringSSL is awesome!");
+ async_dtls_send(&mut server_conn, b"Oh yeah definitely!").await?;
Ok::<_, Error>(())
});
- client_conn.async_handshake().await?;
- client_conn
- .as_pin_mut()
- .async_write(b"BoringSSL is awesome!")
- .await?;
- let mut message = [0; 19];
- let mut read_bytes = 0;
- while read_bytes < 19 {
- match client_conn
- .as_pin_mut()
- .async_read(&mut message[read_bytes..])
- .await?
- {
- IoStatus::Ok(n) => read_bytes += n,
+ drive_async_dtls_handshake(&mut client_conn).await?;
+ async_dtls_send(&mut client_conn, b"BoringSSL is awesome!").await?;
+ let mut buf = [MaybeUninit::uninit(); 19];
+ let mut message = ReceiveBuffer::new_uninit(&mut buf);
+ while message.remaining() > 0 {
+ match async_dtls_recv(&mut client_conn, &mut message).await? {
+ IoStatus::Ok(_) => {}
IoStatus::EndOfStream => break,
_ => {}
}
}
- assert_eq!(&message, b"Oh yeah definitely!");
- assert!(matches!(
- client_conn.as_pin_mut().async_shutdown().await,
- Ok(_) | Err(Error::Io(bssl_tls::errors::IoError::EndOfStream))
- ));
+ assert_eq!(message.filled(), b"Oh yeah definitely!");
task.await.unwrap()?;
Ok(())
}
#[cfg(unix)]
#[tokio::test]
-#[ignore = "https://crbug.com/532601068"]
async fn async_dtls() -> Result<(), Error> {
let (mut server_conn, mut client_conn) = dumb_dtls_server_client();
let (server_sock, client_sock) = tokio::net::UnixDatagram::pair().unwrap();
- server_conn.set_io(TokioDatagramIo(server_sock)).unwrap();
- client_conn.set_io(TokioDatagramIo(client_sock)).unwrap();
+ server_conn
+ .set_datagram_socket(TokioDatagramIo(server_sock))
+ .unwrap();
+ client_conn
+ .set_datagram_socket(TokioDatagramIo(client_sock))
+ .unwrap();
- async_ping_pong(server_conn, client_conn).await
+ async_dtls_ping_pong(server_conn, client_conn).await
}
#[cfg(unix)]
#[tokio::test]
-#[ignore = "https://crbug.com/532601068"]
async fn async_dtls_over_fd() -> Result<(), Error> {
let (mut server_conn, mut client_conn) = dumb_dtls_server_client();
let (server_sock, client_sock) = std::os::unix::net::UnixDatagram::pair().unwrap();
@@ -131,8 +156,8 @@
client_sock.set_nonblocking(true).unwrap();
let server_sock = new_std_datagram_with_tokio(server_sock).unwrap();
let client_sock = new_std_datagram_with_tokio(client_sock).unwrap();
- server_conn.set_io(server_sock).unwrap();
- client_conn.set_io(client_sock).unwrap();
+ server_conn.set_datagram_socket(server_sock).unwrap();
+ client_conn.set_datagram_socket(client_sock).unwrap();
- async_ping_pong(server_conn, client_conn).await
+ async_dtls_ping_pong(server_conn, client_conn).await
}
diff --git a/rust/bssl-tls/src/connection/io.rs b/rust/bssl-tls/src/connection/io.rs
index b156cab..070b55b 100644
--- a/rust/bssl-tls/src/connection/io.rs
+++ b/rust/bssl-tls/src/connection/io.rs
@@ -33,8 +33,9 @@
methods::HasTlsConnectionMethod, //
},
context::{
- HasBasicIo,
- TlsMode, //
+ HasDatagramIo,
+ HasShutdown,
+ HasStreamIo, //
},
errors::{
Error,
@@ -57,7 +58,7 @@
}
}
- fn take_io_err(&mut self) -> Option<Box<dyn core::error::Error + Send + Sync>> {
+ pub(crate) fn take_io_err(&mut self) -> Option<Box<dyn core::error::Error + Send + Sync>> {
let bio = self.get_connection_methods().bio.as_mut()?;
bio.as_mut().take_io_err()
}
@@ -75,10 +76,7 @@
}
}
- /// Read data from the socket.
- ///
- /// This method reads up to `buffer.len()` bytes from `buffer`.
- pub fn sync_read(&mut self, buffer: &mut ReceiveBuffer<'_>) -> Result<IoStatus, Error> {
+ fn read_inner(&mut self, buffer: &mut ReceiveBuffer<'_>) -> Result<IoStatus, Error> {
let buf = unsafe {
// Safety:
// - the use of this pointer is outlived by this function callframe.
@@ -104,7 +102,75 @@
}
}
- /// Peek `buffer.len()` bytes of application data into the `buffer`.
+ fn write_inner(&mut self, buffer: &[u8]) -> Result<IoStatus, Error> {
+ let (ptr, len) = slice_into_ffi_raw_parts(buffer);
+ let num = c_int::try_from(len).unwrap_or(c_int::MAX);
+ let rc = unsafe {
+ // Safety: the validity of the handle `self.ptr()` is witnessed by `self`
+ bssl_sys::SSL_write(self.ptr(), ptr as _, num)
+ };
+ if rc > 0 {
+ Ok(IoStatus::Ok(rc as usize))
+ } else {
+ self.translate_io_error(rc)
+ }
+ }
+}
+
+impl<R, M> TlsConnection<R, M>
+where
+ M: HasTlsConnectionMethod,
+{
+ /// For `async` operations, obtain a pinned mutable reference.
+ pub fn as_pin_mut(&mut self) -> Pin<&mut Self> {
+ Pin::new(self)
+ }
+
+ /// For `async` operations, obtain a pinned immutable reference.
+ pub fn as_pin(&self) -> Pin<&Self> {
+ Pin::new(self)
+ }
+}
+
+impl<R, M> TlsConnection<R, M>
+where
+ M: HasTlsConnectionMethod,
+{
+ fn do_async_io(
+ mut self: Pin<&mut Self>,
+ cx: &mut Context<'_>,
+ sync_op: impl FnOnce(&mut TlsConnection<R, M>) -> Result<IoStatus, Error>,
+ ) -> Result<Option<IoStatus>, Error> {
+ self.set_waker(cx.waker());
+
+ let reason = match sync_op(&mut *self) {
+ Ok(
+ status @ (IoStatus::Ok(..)
+ | IoStatus::EndOfStream
+ | IoStatus::Empty
+ | IoStatus::Err),
+ ) => return Ok(Some(status)),
+ Err(e) => return Err(e),
+ Ok(IoStatus::Retry(reason)) => reason,
+ };
+ self.get_connection_methods().set_pending_reason(reason);
+ Ok(None)
+ }
+}
+
+/// I/O for stream sockets
+impl<R, M> TlsConnection<R, M>
+where
+ M: HasTlsConnectionMethod + HasStreamIo,
+{
+ /// Read data from the socket.
+ ///
+ /// This method reads up to `buffer.remaining()` bytes from `buffer`.
+ pub fn sync_read(&mut self, buffer: &mut ReceiveBuffer<'_>) -> Result<IoStatus, Error> {
+ self.read_inner(buffer)
+ }
+
+ /// Peek `buffer.remaining()` bytes of application data into the `buffer`.
pub fn peek(&mut self, buffer: &mut ReceiveBuffer<'_>) -> Result<IoStatus, Error> {
let buf = unsafe {
// Safety:
@@ -135,17 +201,7 @@
///
/// This method writes up to `buffer.len()` bytes from `buffer`.
pub fn sync_write(&mut self, buffer: &[u8]) -> Result<IoStatus, Error> {
- let (ptr, len) = slice_into_ffi_raw_parts(buffer);
- let num = c_int::try_from(len).unwrap_or(c_int::MAX);
- let rc = unsafe {
- // Safety: the validity of the handle `self.ptr()` is witnessed by `self`
- bssl_sys::SSL_write(self.ptr(), ptr as _, num)
- };
- if rc > 0 {
- Ok(IoStatus::Ok(rc as usize))
- } else {
- self.translate_io_error(rc)
- }
+ self.write_inner(buffer)
}
/// Flush the data on the **transport**.
@@ -181,69 +237,36 @@
Ok(IoStatus::Ok(0))
}
}
-}
-/// Async I/O
-impl<R, M> TlsConnection<R, M>
-where
- M: HasTlsConnectionMethod,
-{
- /// For `async` operations, obtain a pinned mutable reference.
- pub fn as_pin_mut(&mut self) -> Pin<&mut Self> {
- Pin::new(self)
- }
-
- /// For `async` operations, obtain a pinned immutable reference.
- pub fn as_pin(&self) -> Pin<&Self> {
- Pin::new(self)
- }
-}
-
-impl<R, M> TlsConnection<R, M>
-where
- M: HasTlsConnectionMethod + HasBasicIo,
-{
- fn do_async_io(
- mut self: Pin<&mut Self>,
- cx: &mut Context<'_>,
- sync_op: impl FnOnce(&mut TlsConnection<R, M>) -> Result<IoStatus, Error>,
- ) -> Result<Option<IoStatus>, Error> {
- self.set_waker(cx.waker());
-
- let reason = match sync_op(&mut *self) {
- Ok(
- status @ (IoStatus::Ok(..)
- | IoStatus::EndOfStream
- | IoStatus::Empty
- | IoStatus::Err),
- ) => return Ok(Some(status)),
- Err(e) => return Err(e),
- Ok(IoStatus::Retry(reason)) => reason,
- };
- self.get_connection_methods().set_pending_reason(reason);
- Ok(None)
- }
- #[doc(hidden)]
- pub fn aread_inner(
+ /// Poll from the connection once for receiving application data.
+ ///
+ /// When the transport is not ready or the handshake has pending resolution,
+ /// this function will register a waker and return [`None`].
+ pub fn async_poll_read(
self: Pin<&mut Self>,
- buffer: &mut [u8],
+ buffer: &mut ReceiveBuffer<'_>,
cx: &mut Context<'_>,
) -> Result<Option<IoStatus>, Error> {
- let mut buffer = ReceiveBuffer::new(buffer);
- self.do_async_io(cx, move |this| this.sync_read(&mut buffer))
+ self.do_async_io(cx, move |this| this.read_inner(buffer))
}
- #[doc(hidden)]
- pub fn awrite_inner(
+ /// Poll from the connection once for writing application data.
+ ///
+ /// When the transport is not ready or the handshake has pending resolution,
+ /// this function will register a waker and return [`None`].
+ pub fn async_poll_write(
self: Pin<&mut Self>,
buffer: &[u8],
cx: &mut Context<'_>,
) -> Result<Option<IoStatus>, Error> {
- self.do_async_io(cx, move |this| this.sync_write(buffer))
+ self.do_async_io(cx, move |this| this.write_inner(buffer))
}
- #[doc(hidden)]
- pub fn aflush_inner(
+ /// Poll from the connection once for flushing pending writes.
+ ///
+ /// When the transport is not ready or the handshake has pending resolution,
+ /// this function will register a waker and return [`None`].
+ pub fn async_poll_flush(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Result<Option<IoStatus>, Error> {
@@ -256,9 +279,9 @@
/// The reason can be inspected by invoking [`Self::take_pending_reason`].
pub fn async_read<'a>(
mut self: Pin<&'a mut Self>,
- buffer: &'a mut [u8],
+ buffer: &'a mut ReceiveBuffer<'_>,
) -> impl 'a + Send + Future<Output = Result<IoStatus, Error>> {
- poll_fn(move |cx| match self.as_mut().aread_inner(buffer, cx) {
+ poll_fn(move |cx| match self.as_mut().async_poll_read(buffer, cx) {
Ok(Some(status)) => Poll::Ready(Ok(status)),
Ok(None) => Poll::Pending,
Err(e) => Poll::Ready(Err(e)),
@@ -273,7 +296,7 @@
mut self: Pin<&'a mut Self>,
buffer: &'a [u8],
) -> impl 'a + Send + Future<Output = Result<IoStatus, Error>> {
- poll_fn(move |cx| match self.as_mut().awrite_inner(buffer, cx) {
+ poll_fn(move |cx| match self.as_mut().async_poll_write(buffer, cx) {
Ok(Some(status)) => Poll::Ready(Ok(status)),
Ok(None) => Poll::Pending,
Err(e) => Poll::Ready(Err(e)),
@@ -287,15 +310,113 @@
pub fn async_flush<'a>(
mut self: Pin<&'a mut Self>,
) -> impl 'a + Send + Future<Output = Result<IoStatus, Error>> {
- poll_fn(move |cx| match self.as_mut().aflush_inner(cx) {
+ poll_fn(move |cx| match self.as_mut().async_poll_flush(cx) {
+ Ok(Some(status)) => Poll::Ready(Ok(status)),
+ Ok(None) => Poll::Pending,
+ Err(e) => Poll::Ready(Err(e)),
+ })
+ }
+}
+
+/// I/O for datagram sockets.
+impl<R, M> TlsConnection<R, M>
+where
+ M: HasTlsConnectionMethod + HasDatagramIo,
+{
+ /// Receive an application datagram from the socket.
+ ///
+ /// This method reads up to `buffer.len()` bytes from `buffer`.
+ /// Be sure to use [`Self::dtlsv1_get_timeout`] to arm a timer and
+ /// notify the connection about deadline with [`Self::dtlsv1_handle_timeout`].
+ pub fn sync_recv(&mut self, buffer: &mut ReceiveBuffer<'_>) -> Result<IoStatus, Error> {
+ self.read_inner(buffer)
+ }
+
+ /// Send an application datagram down the socket.
+ ///
+ /// This method writes up to `buffer.len()` bytes from `buffer`.
+ /// Be sure to use [`Self::dtlsv1_get_timeout`] to arm a timer and
+ /// notify the connection about deadline with [`Self::dtlsv1_handle_timeout`].
+ pub fn sync_send(&mut self, buffer: &[u8]) -> Result<IoStatus, Error> {
+ self.write_inner(buffer)
+ }
+
+ /// Poll from the connection once for receiving application datagrams.
+ ///
+ /// When the transport is not ready or the handshake has pending resolution,
+ /// this function will register a waker and return [`None`].
+ ///
+ /// Be sure to use [`Self::dtlsv1_get_timeout`] to arm a timer and
+ /// notify the connection about deadline with [`Self::dtlsv1_handle_timeout`].
+ pub fn async_poll_recv(
+ self: Pin<&mut Self>,
+ buffer: &mut ReceiveBuffer<'_>,
+ cx: &mut Context<'_>,
+ ) -> Result<Option<IoStatus>, Error> {
+ self.do_async_io(cx, move |this| this.read_inner(buffer))
+ }
+
+ /// Poll from the connection once for sending application datagrams.
+ ///
+ /// When the transport is not ready or the handshake has pending resolution,
+ /// this function will register a waker and return [`None`].
+ ///
+ /// Be sure to use [`Self::dtlsv1_get_timeout`] to arm a timer and
+ /// notify the connection about deadline with [`Self::dtlsv1_handle_timeout`].
+ pub fn async_poll_send(
+ self: Pin<&mut Self>,
+ buffer: &[u8],
+ cx: &mut Context<'_>,
+ ) -> Result<Option<IoStatus>, Error> {
+ self.do_async_io(cx, move |this| this.write_inner(buffer))
+ }
+
+ /// Asynchronously receive an application datagram.
+ ///
+ /// This method will intercept [`IoStatus::Retry`] and suspend the future.
+ /// The reason can be inspected by invoking [`Self::take_pending_reason`].
+ ///
+ /// Be sure to use [`Self::dtlsv1_get_timeout`] to arm a timer and
+ /// notify the connection about deadline with [`Self::dtlsv1_handle_timeout`].
+ pub fn async_recv<'a>(
+ mut self: Pin<&'a mut Self>,
+ buffer: &'a mut ReceiveBuffer<'_>,
+ ) -> impl 'a + Send + Future<Output = Result<IoStatus, Error>> {
+ poll_fn(move |cx| match self.as_mut().async_poll_recv(buffer, cx) {
Ok(Some(status)) => Poll::Ready(Ok(status)),
Ok(None) => Poll::Pending,
Err(e) => Poll::Ready(Err(e)),
})
}
- #[doc(hidden)]
- pub fn ashutdown_inner(
+ /// Asynchronously send an application datagram.
+ ///
+ /// This method will intercept [`IoStatus::Retry`] and suspend the future.
+ /// The reason can be inspected by invoking [`Self::take_pending_reason`].
+ ///
+ /// Be sure to use [`Self::dtlsv1_get_timeout`] to arm a timer and
+ /// notify the connection about deadline with [`Self::dtlsv1_handle_timeout`].
+ pub fn async_send<'a>(
+ mut self: Pin<&'a mut Self>,
+ buffer: &'a [u8],
+ ) -> impl 'a + Send + Future<Output = Result<IoStatus, Error>> {
+ poll_fn(move |cx| match self.as_mut().async_poll_send(buffer, cx) {
+ Ok(Some(status)) => Poll::Ready(Ok(status)),
+ Ok(None) => Poll::Pending,
+ Err(e) => Poll::Ready(Err(e)),
+ })
+ }
+}
+
+impl<R, M> TlsConnection<R, M>
+where
+ M: HasTlsConnectionMethod + HasShutdown,
+{
+ /// Poll from the connection once for shutting down the connection.
+ ///
+ /// When the transport is not ready or the shutdown has pending resolution,
+ /// this function will register a waker and return [`None`].
+ pub fn async_poll_shutdown(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Result<Option<ShutdownStatus>, Error> {
@@ -315,7 +436,7 @@
pub fn async_shutdown<'a>(
mut self: Pin<&'a mut Self>,
) -> impl 'a + Send + Future<Output = Result<(), Error>> {
- poll_fn(move |cx| match self.as_mut().ashutdown_inner(cx) {
+ poll_fn(move |cx| match self.as_mut().async_poll_shutdown(cx) {
Ok(Some(ShutdownStatus::CloseNotifyReceived)) => Poll::Ready(Ok(())),
Ok(Some(ShutdownStatus::EndOfStream)) => {
Poll::Ready(Err(Error::Io(IoError::EndOfStream)))
diff --git a/rust/bssl-tls/src/connection/io/stdio.rs b/rust/bssl-tls/src/connection/io/stdio.rs
index 95cde2d..3406818 100644
--- a/rust/bssl-tls/src/connection/io/stdio.rs
+++ b/rust/bssl-tls/src/connection/io/stdio.rs
@@ -14,14 +14,14 @@
use std::io;
-use super::{
- Error,
- TlsMode, //
-};
+use super::Error;
use crate::{
ReceiveBuffer,
connection::TlsConnection,
- context::DtlsMode,
+ context::{
+ DtlsMode, //
+ TlsMode,
+ },
errors::{
IoError,
TlsRetryReason, //
@@ -97,11 +97,11 @@
impl<R> DatagramSocket for TlsConnection<R, DtlsMode> {
fn send(&mut self, datagram: &[u8]) -> AbstractSocketResult {
- translate_result_for_datagram(self.sync_write(datagram))
+ translate_result_for_datagram(self.sync_send(datagram))
}
fn recv(&mut self, datagram: &mut [u8]) -> AbstractSocketResult {
let mut datagram = ReceiveBuffer::new(datagram);
- translate_result_for_datagram(self.sync_read(&mut datagram))
+ translate_result_for_datagram(self.sync_recv(&mut datagram))
}
}
diff --git a/rust/bssl-tls/src/connection/lifecycle.rs b/rust/bssl-tls/src/connection/lifecycle.rs
index 64be5e2..6043267 100644
--- a/rust/bssl-tls/src/connection/lifecycle.rs
+++ b/rust/bssl-tls/src/connection/lifecycle.rs
@@ -43,13 +43,14 @@
methods::HasTlsConnectionMethod, //
},
context::{
- HasBasicIo,
+ HasShutdown,
SupportedMode,
TlsMode, //
},
credentials::TlsCredential,
errors::{
Error,
+ IoError,
TlsErrorReason,
TlsRetryReason, //
},
@@ -139,10 +140,14 @@
&mut self,
alert: AlertDescription,
) -> Result<Option<TlsRetryReason>, Error> {
- Ok(check_tls_error!(self.ptr(), {
+ let ret = check_tls_error!(self.ptr(), {
// Safety: `self.0` is still a valid handle and `alert` is valid by construction.
bssl_sys::SSL_send_fatal_alert(self.ptr(), alert as u8)
- }))
+ });
+ if let Some(err) = self.take_io_err() {
+ return Err(Error::Io(IoError::Transport(err)));
+ }
+ Ok(ret)
}
/// Send fatal alert asynchronously.
@@ -184,7 +189,7 @@
/// # Handshake
impl<R, M> TlsConnection<R, M>
where
- M: HasTlsConnectionMethod,
+ M: SupportedMode,
{
/// Drive the handshake.
///
@@ -196,13 +201,17 @@
/// before this method can make progress again.
pub fn do_handshake(&mut self) -> Result<Option<TlsRetryReason>, Error> {
let conn = self.ptr();
- Ok(check_tls_error!(conn, bssl_sys::SSL_do_handshake(conn)))
+ let ret = check_tls_error!(conn, bssl_sys::SSL_do_handshake(conn));
+ if let Some(err) = self.take_io_err() {
+ return Err(Error::Io(IoError::Transport(err)));
+ }
+ Ok(ret)
}
}
impl<M> TlsConnection<Server, M>
where
- M: HasTlsConnectionMethod,
+ M: SupportedMode,
{
/// Accept a connection by responding to `ClientHello` with `ServerHello`.
///
@@ -216,7 +225,7 @@
impl<M> TlsConnection<Client, M>
where
- M: HasTlsConnectionMethod,
+ M: SupportedMode,
{
/// Initiate a connection by sending a `ClientHello`.
///
@@ -248,7 +257,7 @@
impl<R, M> EstablishedTlsConnection<'_, R, M>
where
- M: HasTlsConnectionMethod + HasBasicIo,
+ M: HasTlsConnectionMethod + HasShutdown,
{
/// Perform synchronising shutdown.
///
@@ -275,9 +284,6 @@
// Safety: we have exclusive access to the connection state.
bssl_sys::SSL_shutdown(self.ptr())
};
- if self.is_write_closed() {
- return Ok(Some(ShutdownStatus::EndOfStream));
- }
match rc {
0 => Ok(Some(ShutdownStatus::CloseNotifyPosted)),
1 => Ok(Some(ShutdownStatus::CloseNotifyReceived)),
@@ -289,6 +295,9 @@
Ok(IoStatus::Retry(TlsRetryReason::WantRead | TlsRetryReason::WantWrite)) => {
Ok(None)
}
+ Ok(IoStatus::Retry(TlsRetryReason::Syscall)) => {
+ Ok(Some(ShutdownStatus::EndOfStream))
+ }
Ok(IoStatus::Retry(reason)) => panic!("unexpected retry reason {reason:?}"),
Err(Error::TlsReason(TlsErrorReason::ApplicationDataOnShutdown)) => {
Ok(Some(ShutdownStatus::RemainingApplicationData))
diff --git a/rust/bssl-tls/src/connection/transport.rs b/rust/bssl-tls/src/connection/transport.rs
index 30d44dd..5bfaceb 100644
--- a/rust/bssl-tls/src/connection/transport.rs
+++ b/rust/bssl-tls/src/connection/transport.rs
@@ -15,23 +15,31 @@
//! TLS Connection transport settings
//!
-use core::mem::{
- MaybeUninit,
- transmute, //
+use core::{
+ mem::{
+ MaybeUninit,
+ transmute, //
+ },
+ time::Duration, //
};
use crate::{
check_lib_error,
- check_tls_error,
config::ConfigurationError,
connection::{
TlsConnection,
TlsConnectionBuilder,
methods::HasTlsConnectionMethod, //
},
- context::DtlsMode,
- context::HasBasicIo,
- errors::Error,
+ context::{
+ DtlsMode,
+ HasDatagramIo,
+ HasStreamIo, //
+ },
+ errors::{
+ Error,
+ IoError, //
+ },
io::{
AbstractReader,
AbstractSocket,
@@ -45,10 +53,10 @@
/// These are the methods to configure the underlying IO drivers and transport configurations.
impl<R, M> TlsConnection<R, M>
where
- M: HasBasicIo + HasTlsConnectionMethod,
+ M: HasTlsConnectionMethod,
{
/// Set up underlying transport driver.
- pub fn set_io<S: 'static + AbstractSocket>(&mut self, socket: S) -> Result<&mut Self, Error> {
+ fn set_io_inner<S: 'static + AbstractSocket>(&mut self, socket: S) -> Result<&mut Self, Error> {
let bio = RustBio::new_duplex(socket)?;
unsafe {
// Safety: the additional ref-count is to compensate for `SSL` taking ownership.
@@ -59,6 +67,35 @@
self.get_connection_methods().bio = Some(bio);
Ok(self)
}
+}
+
+/// # Transport configurations
+///
+/// These are the methods to configure the underlying IO drivers and transport configurations.
+impl<R, M> TlsConnection<R, M>
+where
+ M: HasDatagramIo + HasTlsConnectionMethod,
+{
+ /// Set up datagram socket driver.
+ pub fn set_datagram_socket<S: 'static + AbstractSocket>(
+ &mut self,
+ socket: S,
+ ) -> Result<&mut Self, Error> {
+ self.set_io_inner(socket)
+ }
+}
+
+/// # Transport configurations
+///
+/// These are the methods to configure the underlying IO drivers and transport configurations.
+impl<R, M> TlsConnection<R, M>
+where
+ M: HasStreamIo + HasTlsConnectionMethod,
+{
+ /// Set up underlying transport driver.
+ pub fn set_io<S: 'static + AbstractSocket>(&mut self, socket: S) -> Result<&mut Self, Error> {
+ self.set_io_inner(socket)
+ }
/// Set up underlying transport driver, with a pair of read and write ends.
pub fn set_split_io<Reader, Writer>(
@@ -151,15 +188,23 @@
if rc == 0 {
return Ok(false);
}
- let _ = check_tls_error!(conn, rc);
+ // Clear error queue first.
+ let lib_err = Error::extract_lib_err();
+ if let Some(err) = self.take_io_err() {
+ return Err(Error::Io(IoError::Transport(err)));
+ }
+ if rc < 0 {
+ return Err(lib_err);
+ }
Ok(true)
}
- /// Get connection's remaining timeout.
+ /// Get connection's remaining DTLS timer timeout.
///
- /// If a timeout is in effect, this method call returns the remaining seconds,
- /// followed by the remaining microseconds.
- pub fn dtlsv1_get_timeout(&self) -> Option<(i64, i64)> {
+ /// If a timeout is in effect, this method returns the remaining [`Duration`].
+ ///
+ /// This function returns [`None`] when TLS does not have any pending flights.
+ pub fn dtlsv1_get_timeout(&self) -> Option<Duration> {
#[cfg(windows)]
#[repr(C)]
struct timeval {
@@ -183,7 +228,9 @@
// Safety: timeval is now valid as per BoringSSL specification.
timeval.assume_init()
};
- Some((timeval.tv_sec as i64, timeval.tv_usec as i64))
+ let secs = u64::try_from(timeval.tv_sec).unwrap_or(0);
+ let usecs = u64::try_from(timeval.tv_usec).unwrap_or(0);
+ Some(Duration::from_secs(secs) + Duration::from_micros(usecs))
}
0 => None,
rc => {
diff --git a/rust/bssl-tls/src/context.rs b/rust/bssl-tls/src/context.rs
index 84f7650..0d80c93 100644
--- a/rust/bssl-tls/src/context.rs
+++ b/rust/bssl-tls/src/context.rs
@@ -76,16 +76,14 @@
/// [`TlsContextBuilder::with_certificate_verifier`].
pub enum DtlsExternalVerifierMode {}
-pub(crate) trait HasBasicIo {}
-
/// A marker trait for modes that have built-in X.509 support.
-pub trait UseBuiltinX509 {}
+pub(crate) trait UseBuiltinX509: SupportedMode {}
impl UseBuiltinX509 for TlsMode {}
impl UseBuiltinX509 for DtlsMode {}
/// A collection of supported mode of operations.
-pub trait SupportedMode:
+pub(crate) trait SupportedMode:
HasTlsContextMethod + HasTlsConnectionMethod + HasPrivateKeyMethods
{
}
@@ -96,10 +94,22 @@
impl SupportedMode for TlsExternalVerifierMode {}
impl SupportedMode for DtlsExternalVerifierMode {}
-impl HasBasicIo for TlsMode {}
-impl HasBasicIo for DtlsMode {}
-impl HasBasicIo for TlsExternalVerifierMode {}
-impl HasBasicIo for DtlsExternalVerifierMode {}
+pub(crate) trait HasStreamIo: SupportedMode {}
+
+impl HasStreamIo for TlsMode {}
+impl HasStreamIo for TlsExternalVerifierMode {}
+
+pub(crate) trait HasDatagramIo: SupportedMode {}
+
+impl HasDatagramIo for DtlsMode {}
+impl HasDatagramIo for DtlsExternalVerifierMode {}
+
+pub(crate) trait HasShutdown: SupportedMode {}
+
+impl HasShutdown for TlsMode {}
+impl HasShutdown for DtlsMode {}
+impl HasShutdown for TlsExternalVerifierMode {}
+impl HasShutdown for DtlsExternalVerifierMode {}
/// General TLS configuration
///
diff --git a/rust/bssl-tls/src/credentials/tests.rs b/rust/bssl-tls/src/credentials/tests.rs
index 15f1586..9ff4dac 100644
--- a/rust/bssl-tls/src/credentials/tests.rs
+++ b/rust/bssl-tls/src/credentials/tests.rs
@@ -44,6 +44,7 @@
TlsMode, //
},
errors::Error,
+ ffi::ReceiveBuffer,
tests::{
P256_SERVER_CERT,
P256_SERVER_CERT_DER,
@@ -584,14 +585,16 @@
client_conn.as_pin_mut().async_write(b"hello").await?;
let mut server_buf = [0u8; 5];
- let read_len = server_conn.as_pin_mut().async_read(&mut server_buf).await?;
+ let mut recv_buf = ReceiveBuffer::new(&mut server_buf);
+ let read_len = server_conn.as_pin_mut().async_read(&mut recv_buf).await?;
assert!(matches!(read_len, IoStatus::Ok(5)));
assert_eq!(&server_buf, b"hello");
server_conn.as_pin_mut().async_write(b"world").await?;
let mut client_buf = [0u8; 5];
- let read_len = client_conn.as_pin_mut().async_read(&mut client_buf).await?;
+ let mut recv_buf = ReceiveBuffer::new(&mut client_buf);
+ let read_len = client_conn.as_pin_mut().async_read(&mut recv_buf).await?;
assert!(matches!(read_len, IoStatus::Ok(5)));
assert_eq!(&client_buf, b"world");
diff --git a/rust/bssl-tls/src/ffi.rs b/rust/bssl-tls/src/ffi.rs
index d27ce62..ae6397a 100644
--- a/rust/bssl-tls/src/ffi.rs
+++ b/rust/bssl-tls/src/ffi.rs
@@ -136,6 +136,9 @@
_p: PhantomData<&'a mut [u8]>,
}
+// Safety: by construction `ReceiveBuffer` owns the buffer region for exclusive access.
+unsafe impl Send for ReceiveBuffer<'_> {}
+
impl<'a> ReceiveBuffer<'a> {
/// Create a new receiver buffer, with uninitialised bytes.
pub fn new_uninit(buffer: &'a mut [MaybeUninit<u8>]) -> Self {
diff --git a/rust/bssl-tls/src/io/unix.rs b/rust/bssl-tls/src/io/unix.rs
index 5dd4a4b..54ada24 100644
--- a/rust/bssl-tls/src/io/unix.rs
+++ b/rust/bssl-tls/src/io/unix.rs
@@ -32,14 +32,15 @@
}, //
};
-#[cfg(feature = "libc")]
-use crate::ffi::{
- mut_slice_into_ffi_raw_parts,
- slice_into_ffi_raw_parts, //
-};
-use crate::io::stdio::{
- DatagramSocket,
- PollFor, //
+use crate::{
+ ffi::{
+ mut_slice_into_ffi_raw_parts,
+ slice_into_ffi_raw_parts, //
+ },
+ io::stdio::{
+ DatagramSocket,
+ PollFor, //
+ }, //
};
use super::{
@@ -119,11 +120,26 @@
{
}
+/// Check if an I/O error corresponds to `ENOBUFS` (kernel socket buffer exhaustion).
+#[inline]
+fn os_has_no_resource(err: &io::Error) -> bool {
+ #[cfg(unix)]
+ {
+ matches!(err.raw_os_error(), Some(libc::ENOBUFS | libc::ENOMEM))
+ }
+ #[cfg(not(unix))]
+ {
+ false
+ }
+}
+
impl DatagramSocket for UnixDatagram {
fn send(&mut self, datagram: &[u8]) -> AbstractSocketResult {
loop {
return match UnixDatagram::send(self, datagram) {
Ok(bytes) => AbstractSocketResult::Ok(bytes),
+ Err(e) if matches!(e.kind(), io::ErrorKind::Interrupted) => continue,
+ Err(e) if os_has_no_resource(&e) => AbstractSocketResult::Ok(datagram.len()),
Err(e) => crate::retry_on_interrupt!(e),
};
}
@@ -199,19 +215,20 @@
target_os = "ios",
))]
let flag = 0;
- #[cfg(any(windows, target_os = "none"))]
- let flag = 0;
loop {
let rc = unsafe {
// Safety: the socket file descriptor is exclusively owned.
libc::send(self.as_raw_fd(), buf as _, len, flag)
};
- return if rc < 0 {
+ if rc < 0 {
let err = io::Error::last_os_error();
- crate::retry_on_interrupt!(err)
+ if os_has_no_resource(&err) {
+ return AbstractSocketResult::Ok(datagram.len());
+ }
+ return crate::retry_on_interrupt!(err);
} else {
- AbstractSocketResult::Ok(rc as usize)
- };
+ return AbstractSocketResult::Ok(rc as usize);
+ }
}
}
diff --git a/rust/bssl-tls/src/lib.rs b/rust/bssl-tls/src/lib.rs
index fa1719b..67eb1d7 100644
--- a/rust/bssl-tls/src/lib.rs
+++ b/rust/bssl-tls/src/lib.rs
@@ -33,6 +33,7 @@
extern crate alloc;
extern crate core;
+use alloc::boxed::Box;
use core::panic::AssertUnwindSafe;
pub mod alerts;
diff --git a/rust/bssl-tls/src/tests.rs b/rust/bssl-tls/src/tests.rs
index 5a4636f..3386ca6 100644
--- a/rust/bssl-tls/src/tests.rs
+++ b/rust/bssl-tls/src/tests.rs
@@ -106,7 +106,8 @@
fn sync_ping_pong<
M: crate::connection::methods::HasTlsConnectionMethod
+ crate::context::SupportedMode
- + crate::context::HasBasicIo
+ + crate::context::HasStreamIo
+ + crate::context::HasShutdown
+ 'static,
>(
mut server_conn: TlsConnection<Server, M>,
@@ -447,11 +448,12 @@
let server_data = async move {
let mut buf = [0u8; TEST_DATA.len()];
+ let mut message = ReceiveBuffer::new(&mut buf);
let mut read_bytes = 0;
while read_bytes < TEST_DATA.len() {
match server_conn
.as_pin_mut()
- .async_read(&mut buf[read_bytes..])
+ .async_read(&mut message)
.await
.unwrap()
{
diff --git a/rust/bssl-tls/src/tests/credentials.rs b/rust/bssl-tls/src/tests/credentials.rs
index d74c210..47e407a 100644
--- a/rust/bssl-tls/src/tests/credentials.rs
+++ b/rust/bssl-tls/src/tests/credentials.rs
@@ -163,8 +163,9 @@
}
let mut message = [0; 21];
+ let mut recv_buf = crate::ffi::ReceiveBuffer::new(&mut message);
assert!(matches!(
- server_conn.as_pin_mut().async_read(&mut message).await?,
+ server_conn.as_pin_mut().async_read(&mut recv_buf).await?,
IoStatus::Ok(21)
));
assert_eq!(message, *b"BoringSSL is awesome!");
diff --git a/rust/bssl-tls/src/tests/datagram.rs b/rust/bssl-tls/src/tests/datagram.rs
index 68ca698..f185c01 100644
--- a/rust/bssl-tls/src/tests/datagram.rs
+++ b/rust/bssl-tls/src/tests/datagram.rs
@@ -12,15 +12,35 @@
// See the License for the specific language governing permissions and
// limitations under the License.
+use std::{
+ mem::MaybeUninit,
+ thread::sleep, //
+};
+
use bssl_x509::{
- certificates::X509Certificate, keys::PrivateKey, params::Trust, store::X509StoreBuilder,
+ certificates::X509Certificate,
+ keys::PrivateKey,
+ params::Trust,
+ store::X509StoreBuilder, //
};
use crate::{
- connection::{Client, Server, TlsConnection},
- context::{DtlsMode, TlsContextBuilder},
- credentials::{Certificate, TlsCredentialBuilder},
+ connection::{
+ Client,
+ Server,
+ TlsConnection, //
+ },
+ context::{
+ DtlsMode,
+ TlsContextBuilder, //
+ },
+ credentials::{
+ Certificate,
+ TlsCredentialBuilder, //
+ },
errors::Error,
+ ffi::ReceiveBuffer,
+ io::IoStatus, //
};
// TODO(@xfding): this function will come useful for Windows tests.
@@ -46,7 +66,9 @@
};
server_ctx_builder.with_credential(server_cred.unwrap())?;
let server_ctx = server_ctx_builder.build();
- let server_conn = server_ctx.new_server_connection().build();
+ let mut server_conn = server_ctx.new_server_connection();
+ server_conn.with_mtu(500)?;
+ let server_conn = server_conn.build();
let mut client_ctx_builder = TlsContextBuilder::new_dtls();
let ca = X509Certificate::parse_one_from_pem(super::CA)?;
@@ -55,27 +77,147 @@
let cert_store = cert_store.build();
client_ctx_builder.with_certificate_store(&cert_store);
let client_ctx = client_ctx_builder.build();
- let client_conn = client_ctx.new_client_connection().build();
+ let mut client_conn = client_ctx.new_client_connection();
+ client_conn.with_mtu(500)?;
+ let client_conn = client_conn.build();
Ok((server_conn, client_conn))
}
+use std::time::Duration;
+
+use crate::connection::lifecycle::ShutdownStatus;
+use crate::errors::TlsRetryReason;
+
+fn handle_sync_dtls_timeout<R>(conn: &mut TlsConnection<R, DtlsMode>) -> Result<(), Error> {
+ if let Some(timeout) = conn.dtlsv1_get_timeout() {
+ sleep(timeout.min(Duration::from_millis(10)));
+ conn.dtlsv1_handle_timeout()?;
+ } else {
+ sleep(Duration::from_millis(5));
+ }
+ Ok(())
+}
+
+fn drive_dtls_handshake<R>(conn: &mut TlsConnection<R, DtlsMode>) -> Result<(), Error> {
+ loop {
+ match conn.do_handshake() {
+ Ok(None) => break Ok(()),
+ Ok(Some(TlsRetryReason::WantRead | TlsRetryReason::WantWrite)) => {
+ handle_sync_dtls_timeout(conn)?;
+ }
+ Ok(Some(reason)) => panic!("unexpected retry reason {reason:?}"),
+ Err(e) => break Err(e),
+ }
+ }
+}
+
+fn dtls_sync_recv<R>(
+ conn: &mut TlsConnection<R, DtlsMode>,
+ buf: &mut ReceiveBuffer<'_>,
+) -> Result<usize, Error> {
+ loop {
+ match conn.sync_recv(buf) {
+ Ok(IoStatus::Ok(n)) => break Ok(n),
+ Ok(IoStatus::Retry(TlsRetryReason::WantRead | TlsRetryReason::WantWrite)) => {
+ handle_sync_dtls_timeout(conn)?;
+ }
+ Ok(IoStatus::EndOfStream) => break Ok(0),
+ Ok(status) => panic!("unexpected status {status:?}"),
+ Err(e) => break Err(e),
+ }
+ }
+}
+
+fn dtls_sync_send<R>(conn: &mut TlsConnection<R, DtlsMode>, data: &[u8]) -> Result<usize, Error> {
+ loop {
+ match conn.sync_send(data) {
+ Ok(IoStatus::Ok(n)) => break Ok(n),
+ Ok(IoStatus::Retry(TlsRetryReason::WantRead | TlsRetryReason::WantWrite)) => {
+ handle_sync_dtls_timeout(conn)?;
+ }
+ Ok(status) => panic!("unexpected status {status:?}"),
+ Err(e) => break Err(e),
+ }
+ }
+}
+
+fn dtls_sync_shutdown<R>(conn: &mut TlsConnection<R, DtlsMode>) -> Result<(), Error> {
+ loop {
+ let Some(mut established) = conn.established() else {
+ break Ok(());
+ };
+ match established.sync_shutdown() {
+ Ok(Some(ShutdownStatus::CloseNotifyReceived | ShutdownStatus::EndOfStream)) => {
+ break Ok(());
+ }
+ Ok(Some(ShutdownStatus::CloseNotifyPosted)) => break Ok(()),
+ Ok(Some(ShutdownStatus::RemainingApplicationData)) => {
+ let mut discard = [MaybeUninit::uninit(); 128];
+ let mut discard_buf = ReceiveBuffer::new_uninit(&mut discard);
+ let _ = conn.sync_recv(&mut discard_buf);
+ }
+ Ok(None) => {
+ handle_sync_dtls_timeout(conn)?;
+ }
+ Err(e) => break Err(e),
+ }
+ }
+}
+
+fn sync_ping_pong_datagram(
+ mut server_conn: TlsConnection<Server, DtlsMode>,
+ mut client_conn: TlsConnection<Client, DtlsMode>,
+) -> Result<(), Error> {
+ let thread = std::thread::spawn(move || {
+ drive_dtls_handshake(&mut server_conn)?;
+ assert!(!server_conn.is_in_handshake());
+ let mut message = [MaybeUninit::uninit(); 21];
+ let mut message = ReceiveBuffer::new_uninit(&mut message);
+ let n = dtls_sync_recv(&mut server_conn, &mut message)?;
+ assert_eq!(n, 21);
+ assert_eq!(*message, *b"BoringSSL is awesome!");
+ dtls_sync_send(&mut server_conn, b"Oh yeah definitely!")?;
+ dtls_sync_shutdown(&mut server_conn)?;
+ // Second shutdown poll.
+ let _ = dtls_sync_shutdown(&mut server_conn);
+ Ok::<_, Error>(())
+ });
+
+ drive_dtls_handshake(&mut client_conn)?;
+ assert!(!client_conn.is_in_handshake());
+ dtls_sync_send(&mut client_conn, b"BoringSSL is awesome!")?;
+ let mut message = [MaybeUninit::uninit(); 19];
+ let mut message = ReceiveBuffer::new_uninit(&mut message);
+ let n = dtls_sync_recv(&mut client_conn, &mut message)?;
+ assert_eq!(n, 19);
+ assert_eq!(*message, *b"Oh yeah definitely!");
+ dtls_sync_shutdown(&mut client_conn)?;
+ thread.join().unwrap()?;
+
+ Ok(())
+}
+
#[cfg(unix)]
-#[ignore = "https://crbug.com/532601068"]
#[test]
fn dtls() {
- use crate::{io::sync_io::NoAsync, io::unix::StdDatagram, tests::sync_ping_pong};
+ use crate::{io::sync_io::NoAsync, io::unix::StdDatagram};
let (mut server_conn, mut client_conn) = dumb_dtls_server_client().unwrap();
let (server_sock, client_sock) = std::os::unix::net::UnixDatagram::pair().unwrap();
+ server_sock
+ .set_read_timeout(Some(Duration::from_millis(50)))
+ .unwrap();
+ client_sock
+ .set_read_timeout(Some(Duration::from_millis(50)))
+ .unwrap();
let server_sock = StdDatagram::new(server_sock, NoAsync);
let client_sock = StdDatagram::new(client_sock, NoAsync);
- server_conn.set_io(server_sock).unwrap();
- client_conn.set_io(client_sock).unwrap();
- sync_ping_pong(server_conn, client_conn).unwrap();
+ server_conn.set_datagram_socket(server_sock).unwrap();
+ client_conn.set_datagram_socket(client_sock).unwrap();
+ sync_ping_pong_datagram(server_conn, client_conn).unwrap();
}
-#[ignore = "https://crbug.com/532601068"]
#[test]
fn test_async_dtls() -> Result<(), Error> {
use crate::io::IoStatus;
@@ -85,8 +227,8 @@
let (client_socket, server_socket, mut executor) = create_mock_datagram();
- server_conn.set_io(server_socket)?;
- client_conn.set_io(client_socket)?;
+ server_conn.set_datagram_socket(server_socket)?;
+ client_conn.set_datagram_socket(client_socket)?;
let test_future = async {
futures::future::try_join(server_conn.async_handshake(), client_conn.async_handshake())
@@ -94,13 +236,10 @@
let server_data = async {
let mut buf = [0u8; TEST_DATA.len()];
+ let mut message = ReceiveBuffer::new(&mut buf);
let mut read_bytes = 0;
while read_bytes < TEST_DATA.len() {
- match server_conn
- .as_pin_mut()
- .async_read(&mut buf[read_bytes..])
- .await?
- {
+ match server_conn.as_pin_mut().async_recv(&mut message).await? {
IoStatus::Ok(n) => read_bytes += n,
IoStatus::EndOfStream => break,
_ => {}
@@ -111,7 +250,7 @@
};
let client_data = async {
- client_conn.as_pin_mut().async_write(TEST_DATA).await?;
+ client_conn.as_pin_mut().async_send(TEST_DATA).await?;
Ok::<(), Error>(())
};
diff --git a/rust/bssl-x509/src/ffi.rs b/rust/bssl-x509/src/ffi.rs
index 11f186f..a267716 100644
--- a/rust/bssl-x509/src/ffi.rs
+++ b/rust/bssl-x509/src/ffi.rs
@@ -17,7 +17,7 @@
ptr::{
NonNull,
null, //
- },//
+ }, //
};
use bssl_crypto::FfiSlice;