use s2n_quic_core::{
buffer::{reader::Incremental, writer::Storage as _, Writer as _},
event::IntoEvent,
inet::ExplicitCongestionNotification,
time::{Clock, Timestamp},
varint::VarInt,
};
use std::{
cell::UnsafeCell,
io,
sync::{Arc, Mutex},
task::Poll,
time::Duration,
};
use tokio::net::TcpStream;
use crate::{
msg,
stream::{
environment::{tokio::Environment, Environment as _},
recv,
socket::{application::Single, Application},
},
};
mod cert_chain;
pub use cert_chain::CertificateChain;
pub struct S2nTlsConnection {
socket: Arc<Single<TcpStream>>,
connection: Mutex<(Conn, ReadState)>,
cert_chain: Option<CertificateChain>,
}
struct ReadState {
reader: Incremental,
buffer: bytes::BytesMut,
}
pub type Conn = Box<dyn AsMut<s2n_tls::connection::Connection> + Send>;
pub trait ConnectionBuilder: Send + Sync {
fn build_connection(&self, mode: s2n_tls::enums::Mode) -> Result<Conn, s2n_tls::error::Error>;
}
impl<B> ConnectionBuilder for B
where
B: s2n_tls::connection::Builder + Send + Sync,
B::Output: Send + 'static,
{
fn build_connection(&self, mode: s2n_tls::enums::Mode) -> Result<Conn, s2n_tls::error::Error> {
Ok(Box::new(self.build_connection(mode)?))
}
}
impl S2nTlsConnection {
pub fn from_connection(
socket: Arc<Single<TcpStream>>,
mut connection: Conn,
) -> io::Result<Self> {
(*connection)
.as_mut()
.set_blinding(s2n_tls::enums::Blinding::SelfService)?;
Ok(S2nTlsConnection {
socket,
connection: Mutex::new((
connection,
ReadState {
reader: Incremental::new(VarInt::ZERO),
buffer: bytes::BytesMut::with_capacity(8192),
},
)),
cert_chain: None,
})
}
pub(crate) async fn negotiate(
&mut self,
mut initial_read_buffer: Option<crate::msg::recv::Message>,
) -> io::Result<()> {
std::future::poll_fn(|cx| -> Poll<io::Result<()>> {
let s2n_connection = &mut self.connection.get_mut().unwrap().0;
let context = NegotiateContext {
socket: &self.socket,
waker: cx.waker(),
initial_read_buffer: initial_read_buffer.as_mut(),
};
let mut connection = CallbackResetGuard {
conn: (**s2n_connection).as_mut(),
reset_write: true,
reset_read: true,
};
connection.set_receive_callback(Some(recv_direct_cb))?;
connection.set_send_callback(Some(send_direct_cb))?;
connection.set_waker(Some(cx.waker()))?;
let mut connection = connection.set_context(&context);
let res = match connection.poll_negotiate() {
Poll::Ready(Ok(_)) => {
drop(connection);
if let Some(buffer) = &mut initial_read_buffer {
if !buffer.is_empty() {
let e = io::Error::new(
std::io::ErrorKind::InvalidData,
"received data pre-handshake",
);
return Poll::Ready(Err(e));
}
}
self.cert_chain = Some(CertificateChain::new(
(**s2n_connection).as_mut().peer_cert_chain()?,
)?);
Poll::Ready(Ok(()))
}
Poll::Ready(Err(e)) => Poll::Ready(Err(e.into())),
Poll::Pending => Poll::Pending,
};
res
})
.await
}
pub(crate) fn write<M, R>(
&self,
message: &mut M,
reader: &mut R,
is_fin: bool,
) -> Result<(), crate::stream::send::Error>
where
M: super::send::application::state::Message,
R: s2n_quic_core::buffer::reader::storage::Infallible,
{
let mut guard = self.connection.lock().unwrap();
let conn = CallbackResetGuard {
conn: (*guard.0).as_mut(),
reset_write: true,
reset_read: false,
};
let mut conn = conn.set_context_mut(message);
conn.set_send_callback(Some(send_io_cb::<M>))
.expect("infallible");
while !reader.buffer_is_empty() {
let Ok(chunk) = reader.read_chunk(usize::MAX);
let mut consumed = 0;
while consumed < chunk.len() {
match conn.poll_send(&chunk) {
Poll::Ready(Ok(l)) => consumed += l,
Poll::Ready(Err(e)) => {
tracing::warn!("s2n_tls::poll_send() = Err({:?})", &e);
return Err(crate::stream::send::Error::new(
crate::stream::send::ErrorKind::FatalError,
));
}
Poll::Pending => unreachable!(
"TODO: verify, but s2n-tls shouldn't block when the network doesn't"
),
}
}
}
if is_fin {
match conn.poll_shutdown_send() {
Poll::Ready(Ok(_)) => {}
Poll::Ready(Err(e)) => {
tracing::warn!("s2n_tls::poll_shutdown_send() = Err({:?})", &e);
return Err(crate::stream::send::Error::new(
crate::stream::send::ErrorKind::FatalError,
));
}
Poll::Pending => unreachable!(
"TODO: verify, but s2n-tls shouldn't block when the network doesn't"
),
}
}
Ok(())
}
pub(crate) fn read(
&self,
input: &mut super::recv::shared::RecvBuffer,
output: &mut s2n_quic_core::buffer::duplex::Interposer<
'_,
impl s2n_quic_core::buffer::writer::Storage,
s2n_quic_core::buffer::Reassembler,
>,
) -> Result<(), super::recv::Error> {
let mut guard = self.connection.lock().unwrap();
let (conn, read_state) = &mut *guard;
let conn = CallbackResetGuard {
conn: (**conn).as_mut(),
reset_write: false,
reset_read: true,
};
let mut conn = conn.set_context_mut(input);
conn.set_receive_callback(Some(recv_io_cb)).unwrap();
read_state.buffer.reserve(8192);
match conn.poll_recv_uninitialized(read_state.buffer.spare_capacity_mut()) {
Poll::Ready(Ok(len)) => {
unsafe {
let original = read_state.buffer.len();
read_state.buffer.set_len(
original
.checked_add(len)
.expect("single buffer cannot exceed isize::MAX, so cannot overflow"),
);
}
let is_fin = len == 0;
let mut reader = match read_state
.reader
.with_storage(&mut read_state.buffer, is_fin)
{
Ok(r) => r,
Err(s2n_quic_core::buffer::Error::OutOfRange) => {
return Err(super::recv::Error::new(
super::recv::ErrorKind::MaxDataExceeded,
))
}
Err(s2n_quic_core::buffer::Error::InvalidFin) => {
return Err(super::recv::Error::new(super::recv::ErrorKind::InvalidFin))
}
};
match output.read_from(&mut reader) {
Ok(()) => {}
Err(s2n_quic_core::buffer::Error::OutOfRange) => {
return Err(super::recv::Error::new(
super::recv::ErrorKind::MaxDataExceeded,
))
}
Err(s2n_quic_core::buffer::Error::InvalidFin) => {
return Err(super::recv::Error::new(super::recv::ErrorKind::InvalidFin))
}
}
}
Poll::Ready(Err(e)) => {
tracing::warn!("s2n_tls::poll_recv() = Err({:?})", &e);
return Err(super::recv::Error::new(super::recv::ErrorKind::Decode));
}
Poll::Pending => {
}
}
Ok(())
}
pub(crate) fn peer_cert_chain(&self) -> Option<&CertificateChain> {
self.cert_chain.as_ref()
}
}
struct NegotiateContext<'a> {
socket: &'a Single<TcpStream>,
waker: &'a std::task::Waker,
initial_read_buffer: Option<&'a mut crate::msg::recv::Message>,
}
#[allow(clippy::extra_unused_lifetimes)]
unsafe extern "C" fn recv_direct_cb<'a>(
ctx: *mut core::ffi::c_void,
buf: *mut u8,
len: u32,
) -> i32 {
let ctx = ctx.cast::<NegotiateContext<'a>>().as_mut::<'a>().unwrap();
let mut cx = std::task::Context::from_waker(ctx.waker);
let buf = std::slice::from_raw_parts_mut(buf, len as usize);
if let Some(initial_read_buffer) = ctx.initial_read_buffer.as_mut() {
let peeked = initial_read_buffer.peek();
let consumed = std::cmp::min(buf.len(), peeked.len());
buf[..consumed].copy_from_slice(&peeked[..consumed]);
initial_read_buffer.consume(consumed);
if consumed > 0 {
return consumed as i32;
}
}
let buf = std::io::IoSliceMut::new(buf);
let mut addr = Default::default();
let mut cmsg = Default::default();
match ctx
.socket
.read_application()
.poll_recv(&mut cx, &mut addr, &mut cmsg, &mut [buf])
{
Poll::Ready(Ok(r)) => r as i32,
Poll::Ready(Err(e)) => {
nix::errno::Errno::try_from(e)
.unwrap_or(nix::errno::Errno::EIO)
.set();
-1
}
Poll::Pending => {
nix::errno::Errno::EWOULDBLOCK.set();
-1
}
}
}
#[allow(clippy::extra_unused_lifetimes)]
unsafe extern "C" fn send_direct_cb<'a>(
ctx: *mut core::ffi::c_void,
buf: *const u8,
len: u32,
) -> i32 {
let ctx = ctx.cast::<NegotiateContext<'a>>().as_ref::<'a>().unwrap();
let mut cx = std::task::Context::from_waker(ctx.waker);
let buf = std::slice::from_raw_parts(buf, len as usize);
let buf = std::io::IoSlice::new(buf);
let addr = Default::default();
let ecn = Default::default();
match ctx
.socket
.write_application()
.poll_send(&mut cx, &addr, ecn, &[buf])
{
Poll::Ready(Ok(r)) => r as i32,
Poll::Ready(Err(e)) => {
nix::errno::Errno::try_from(e)
.unwrap_or(nix::errno::Errno::EIO)
.set();
-1
}
Poll::Pending => {
nix::errno::Errno::EWOULDBLOCK.set();
-1
}
}
}
unsafe extern "C" fn send_io_cb<'a, M>(ctx: *mut core::ffi::c_void, buf: *const u8, len: u32) -> i32
where
M: 'a + super::send::application::state::Message,
{
let message = ctx.cast::<M>().as_mut::<'a>().unwrap();
let mut buf = std::slice::from_raw_parts(buf, len as usize);
while !buf.is_empty() {
let part = buf
.split_off(..buf.len().clamp(0, u16::MAX as usize))
.unwrap();
message.push(part.len(), |mut b| {
b.put_slice(part);
crate::stream::send::application::transmission::Event {
packet_number: VarInt::ZERO,
info: crate::stream::send::application::transmission::Info {
packet_len: part.len() as u16,
retransmission: None,
stream_offset: VarInt::ZERO,
payload_len: 0,
included_fin: Default::default(),
time_sent: unsafe { Timestamp::from_duration(Duration::from_millis(1)) },
ecn: Default::default(),
},
has_more_app_data: false,
}
});
}
i32::try_from(len).unwrap()
}
#[allow(clippy::extra_unused_lifetimes)]
unsafe extern "C" fn recv_io_cb<'a>(ctx: *mut core::ffi::c_void, buf: *mut u8, len: u32) -> i32 {
let mut ctx = ctx
.cast::<super::recv::shared::RecvBuffer>()
.as_mut::<'a>()
.unwrap();
let crate::either::Either::A(a) = &mut ctx else {
unreachable!("only local buffer for TLS stream");
};
let output = bytes::buf::UninitSlice::from_raw_parts_mut(buf, len as usize);
let written = a.copy_into(output);
if written == 0 {
if a.saw_fin() {
0
} else {
nix::errno::Errno::EWOULDBLOCK.set();
-1
}
} else {
written as i32
}
}
unsafe extern "C" fn unreachable_recv_io_cb(_: *mut core::ffi::c_void, _: *mut u8, _: u32) -> i32 {
unreachable!(
"s2n-tls should not call I/O callbacks outside of application controlled send/receive"
);
}
unsafe extern "C" fn unreachable_send_io_cb(
_: *mut core::ffi::c_void,
_: *const u8,
_: u32,
) -> i32 {
unreachable!(
"s2n-tls should not call I/O callbacks outside of application controlled send/receive"
);
}
struct CallbackResetGuard<'a> {
conn: &'a mut s2n_tls::connection::Connection,
reset_write: bool,
reset_read: bool,
}
impl<'a> CallbackResetGuard<'a> {
fn set_context<T>(self, context: &'a T) -> Self {
unsafe {
if self.reset_write {
self.conn
.set_send_context(context as *const T as *mut std::ffi::c_void)
.expect("infallible");
}
if self.reset_read {
self.conn
.set_receive_context(context as *const T as *mut std::ffi::c_void)
.expect("infallible");
}
self
}
}
fn set_context_mut<T>(self, context: &'a mut T) -> Self {
unsafe {
if self.reset_write {
self.conn
.set_send_context(context as *mut _ as *mut std::ffi::c_void)
.expect("infallible");
}
if self.reset_read {
self.conn
.set_receive_context(context as *mut _ as *mut std::ffi::c_void)
.expect("infallible");
}
self
}
}
}
impl std::ops::Deref for CallbackResetGuard<'_> {
type Target = s2n_tls::connection::Connection;
fn deref(&self) -> &Self::Target {
self.conn
}
}
impl std::ops::DerefMut for CallbackResetGuard<'_> {
fn deref_mut(&mut self) -> &mut Self::Target {
self.conn
}
}
impl Drop for CallbackResetGuard<'_> {
fn drop(&mut self) {
unsafe {
if self.reset_write {
self.conn
.set_send_context(std::ptr::null_mut())
.unwrap_or_else(|_| std::process::abort());
self.conn
.set_send_callback(Some(unreachable_send_io_cb))
.unwrap_or_else(|_| std::process::abort());
}
if self.reset_read {
self.conn
.set_receive_context(std::ptr::null_mut())
.unwrap_or_else(|_| std::process::abort());
self.conn
.set_receive_callback(Some(unreachable_recv_io_cb))
.unwrap_or_else(|_| std::process::abort());
}
}
}
}
pub(crate) fn build_stream<Sub>(
kernel_start_time: Timestamp,
addr: std::net::SocketAddr,
socket: Arc<Single<TcpStream>>,
s2n_connection: crate::stream::tls::S2nTlsConnection,
env: &Environment<Sub>,
map: &crate::path::secret::Map,
endpoint_type: s2n_quic_core::endpoint::Type,
) -> io::Result<crate::stream::application::Builder<Sub>>
where
Sub: crate::event::Subscriber + Clone,
{
let peer_addr = if addr.ip().is_unspecified() {
socket.0.peer_addr()?
} else {
addr
};
let stream_id = crate::packet::stream::Id {
queue_id: VarInt::ZERO,
is_reliable: true,
is_bidirectional: true,
};
let params = s2n_quic_core::dc::ApplicationParams::new(
1 << 14,
&Default::default(),
&Default::default(),
);
let meta = crate::event::api::ConnectionMeta {
id: 0, timestamp: env.clock().get_time().into_event(),
};
let info = crate::event::api::ConnectionInfo {};
let subscriber = env.subscriber().clone();
let subscriber_ctx = subscriber.create_connection_context(&meta, &info);
let mut secret = [0; 32];
aws_lc_rs::rand::fill(&mut secret).unwrap();
let secret = crate::path::secret::schedule::Secret::new(
crate::path::secret::schedule::Ciphersuite::AES_GCM_128_SHA256,
s2n_quic_core::dc::SUPPORTED_VERSIONS[0],
endpoint_type,
&secret,
);
let common = {
let application = crate::stream::send::application::state::State { is_reliable: true };
let fixed = crate::stream::shared::FixedValues {
remote_ip: UnsafeCell::new(peer_addr.ip().into()),
application: UnsafeCell::new(application),
credentials: UnsafeCell::new(crate::credentials::Credentials {
id: crate::credentials::Id::from([1; 16]),
key_id: VarInt::ZERO,
}),
};
crate::stream::shared::Common {
clock: env.clock().clone(),
gso: env.gso(),
remote_port: peer_addr.port().into(),
remote_queue_id: stream_id.queue_id.as_u64().into(),
local_queue_id: u64::MAX.into(),
last_peer_activity: Default::default(),
fixed,
closed_halves: 0u8.into(),
subscriber: crate::stream::shared::Subscriber {
subscriber,
context: subscriber_ctx,
},
s2n_connection: Some(s2n_connection),
}
};
let pair = crate::path::secret::map::ApplicationPair::new(
&secret,
VarInt::ZERO,
crate::path::secret::schedule::Initiator::Local,
crate::path::secret::map::Dedup::disabled(),
);
let shared = Arc::new(crate::stream::shared::Shared {
receiver: crate::stream::recv::shared::State::new(
stream_id,
¶ms,
crate::stream::TransportFeatures::TCP,
crate::stream::recv::shared::RecvBuffer::A(recv::buffer::Local::new(
msg::recv::Message::new(9000),
None,
)),
endpoint_type,
&env.clock(),
),
sender: crate::stream::send::shared::State::new(
crate::stream::send::flow::non_blocking::State::new(VarInt::MAX),
crate::stream::send::path::Info {
max_datagram_size: params.max_datagram_size(),
send_quantum: 10,
ecn: ExplicitCongestionNotification::Ect0,
next_expected_control_packet: VarInt::ZERO,
},
None,
),
crypto: crate::stream::shared::Crypto::new(pair.sealer, pair.opener, None, map),
common,
});
let read = crate::stream::recv::application::Builder::new(endpoint_type, env.reader_rt());
let write = crate::stream::send::application::Builder::new(env.writer_rt());
Ok(crate::stream::application::Builder {
read,
write,
shared,
sockets: Box::new(socket),
kernel_start_time,
app_queue_time: None,
})
}
#[cfg(test)]
mod test;