use crate::{
clock::Timer,
event::{self, ConnectionPublisher as _},
msg,
stream::{
recv, runtime,
shared::{AcceptState, ArcShared, ShutdownKind},
socket, Actor,
},
};
use core::{
fmt,
mem::ManuallyDrop,
pin::Pin,
task::{Context, Poll},
};
use s2n_quic_core::{
buffer::{self, writer::Storage as _},
ensure, ready,
stream::state,
task::waker,
time::{clock::Timer as _, timer::Provider as _, Timestamp},
};
use std::{io, net::SocketAddr};
mod builder;
pub use builder::Builder;
pub use crate::stream::recv::shared::AckMode;
#[derive(Clone, Copy, Debug, Default)]
pub enum ReadMode {
UntilFull,
#[default]
Once,
Drain,
}
pub struct Reader<Sub: event::Subscriber>(ManuallyDrop<Box<Inner<Sub>>>);
pub(crate) struct Inner<Sub>
where
Sub: event::Subscriber,
{
shared: ArcShared<Sub>,
sockets: socket::ArcApplication,
send_buffer: msg::send::Message,
read_mode: ReadMode,
ack_mode: AckMode,
timer: Option<Timer>,
local_state: LocalState,
runtime: runtime::ArcHandle<Sub>,
}
impl<Sub> fmt::Debug for Reader<Sub>
where
Sub: event::Subscriber,
{
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
let mut s = f.debug_struct("Reader");
for (name, addr) in [
("peer_addr", self.peer_addr()),
("local_addr", self.local_addr()),
] {
if let Ok(addr) = addr {
s.field(name, &addr);
}
}
s.finish()
}
}
#[derive(Clone, Debug, Default)]
enum LocalState {
#[default]
Ready,
Reading,
Drained,
Errored(recv::Error),
}
impl LocalState {
#[inline]
fn check(&self) -> Option<io::Result<()>> {
match self {
Self::Ready | Self::Reading => None,
Self::Drained => Some(Ok(())),
Self::Errored(err) => Some(Err((*err).into())),
}
}
#[inline]
fn on_read(&mut self) {
ensure!(matches!(self, Self::Ready));
*self = Self::Reading;
}
#[inline]
fn transition<Sub>(&mut self, target: Self, shared: &ArcShared<Sub>)
where
Sub: event::Subscriber,
{
ensure!(matches!(self, Self::Ready | Self::Reading));
*self = target;
shared
.common
.closed_halves
.fetch_add(1, core::sync::atomic::Ordering::Relaxed);
}
}
impl<Sub> Reader<Sub>
where
Sub: event::Subscriber,
{
#[inline]
pub fn peer_addr(&self) -> io::Result<SocketAddr> {
self.0.shared.common.ensure_open()?;
Ok(self.0.shared.remote_addr().into())
}
#[inline]
pub fn local_addr(&self) -> io::Result<SocketAddr> {
self.0.sockets.read_application().local_addr()
}
#[inline]
pub fn protocol(&self) -> socket::Protocol {
self.0.sockets.protocol()
}
#[inline]
pub fn set_read_mode(&mut self, read_mode: ReadMode) -> &mut Self {
self.0.read_mode = read_mode;
self
}
#[inline]
pub fn set_ack_mode(&mut self, ack_mode: AckMode) -> &mut Self {
self.0.ack_mode = ack_mode;
self
}
#[inline]
pub async fn read_into<S>(&mut self, out_buf: &mut S) -> io::Result<usize>
where
S: buffer::writer::Storage,
{
core::future::poll_fn(|cx| self.poll_read_into(cx, out_buf)).await
}
#[inline]
pub fn poll_read_into<S>(
&mut self,
cx: &mut Context,
out_buf: &mut S,
) -> Poll<io::Result<usize>>
where
S: buffer::writer::Storage,
{
let start_time = self.0.shared.clock.get_time();
let capacity = out_buf.remaining_capacity();
let result = waker::debug_assert_contract(cx, |cx| {
let mut out_buf = out_buf.track_write();
let res = self.0.poll_read_into(cx, &mut out_buf);
if res.is_pending() {
debug_assert_eq!(
out_buf.written_len(),
0,
"bytes should only be written on Ready(_)"
);
}
let res = ready!(res);
self.0.local_state.on_read();
res?;
Ok(out_buf.written_len()).into()
});
self.0.publish_read_events(capacity, start_time, &result);
result
}
}
impl<Sub> Inner<Sub>
where
Sub: event::Subscriber,
{
#[inline(always)]
fn poll_read_into<S>(
&mut self,
cx: &mut Context,
out_buf: &mut buffer::writer::storage::Tracked<S>,
) -> Poll<io::Result<()>>
where
S: buffer::writer::Storage,
{
if let Some(res) = self.local_state.check() {
return res.into();
}
let mut force_recv = !out_buf.has_remaining_capacity();
let shared = &self.shared;
let sockets = &self.sockets;
let transport_features = sockets.features();
let mut reader = shared.receiver.application_guard(
self.ack_mode,
&mut self.send_buffer,
shared,
sockets,
)?;
let reader = &mut *reader;
loop {
reader.process_recv_buffer(out_buf, shared, transport_features, AcceptState::Accepted);
if cfg!(debug_assertions) && out_buf.has_remaining_capacity() {
assert!(reader.reassembler.is_empty());
}
if let Err(err) = reader.receiver.check_error() {
self.local_state
.transition(LocalState::Errored(err), &self.shared);
if out_buf.written_len() > 0 {
break;
}
return Err(err.into()).into();
}
match reader.receiver.state() {
state::Receiver::Recv | state::Receiver::SizeKnown => {
}
state::Receiver::DataRecvd => {
ensure!(out_buf.has_remaining_capacity(), Ok(()).into());
continue;
}
state::Receiver::DataRead => {
self.local_state
.transition(LocalState::Drained, &self.shared);
break;
}
state::Receiver::ResetRecvd | state::Receiver::ResetRead => unreachable!(),
}
match self.read_mode {
_ if force_recv => {}
ReadMode::UntilFull if !out_buf.has_remaining_capacity() => break,
ReadMode::Once if out_buf.written_len() > 0 => break,
_ => {}
}
let recv = reader.poll_fill_recv_buffer(
cx,
Actor::Application,
self.sockets.read_application(),
&self.shared.clock,
&self.shared.subscriber,
);
let recv_len =
match Self::handle_socket_result(cx, &mut reader.receiver, &mut self.timer, recv) {
Poll::Ready(res) => res?,
Poll::Pending if out_buf.written_len() > 0 => break,
Poll::Pending => return Poll::Pending,
};
force_recv = false;
if recv_len == 0 {
if transport_features.is_stream() {
let publisher = self.shared.publisher();
reader.receiver.on_transport_close(&publisher);
continue;
} else {
debug_assert!(false, "datagram recv buffers should never be empty");
}
}
}
Ok(()).into()
}
#[inline]
fn handle_socket_result(
cx: &mut Context,
receiver: &mut recv::state::State,
timer: &mut Option<Timer>,
res: Poll<io::Result<usize>>,
) -> Poll<io::Result<usize>> {
if let Poll::Ready(res) = res {
return res.into();
}
let Some(timer) = timer.as_mut() else {
return Poll::Pending;
};
if let Some(target) = receiver.next_expiration() {
timer.update(target);
ready!(timer.poll_ready(cx));
Ok(1).into()
} else {
timer.cancel();
Poll::Pending
}
}
#[inline]
fn shutdown(mut self: Box<Self>) {
if let LocalState::Ready = self.local_state {
let mut storage = buffer::writer::storage::Empty;
let waker = s2n_quic_core::task::waker::noop();
let mut cx = core::task::Context::from_waker(&waker);
let _ = self.poll_read_into(&mut cx, &mut storage.track_write());
}
let background = matches!(self.local_state, LocalState::Ready);
self.shared
.publisher()
.on_stream_read_shutdown(event::builder::StreamReadShutdown { background });
if background {
tracing::debug!("spawning task to read server's response");
let runtime = self.runtime.clone();
let handle = Shutdown(self);
runtime.spawn_recv_shutdown(handle);
return;
}
self.local_state
.transition(LocalState::Drained, &self.shared);
let kind = if std::thread::panicking() {
ShutdownKind::Panicking
} else {
ShutdownKind::Normal
};
self.shared.receiver.shutdown(kind);
}
#[inline(always)]
fn publish_read_events(
&self,
capacity: usize,
start_time: Timestamp,
result: &Poll<io::Result<usize>>,
) {
let end_time = self.shared.clock.get_time();
let processing_duration = end_time.saturating_duration_since(start_time);
match result {
Poll::Ready(Ok(0)) if capacity > 0 => {
self.shared.common.publisher().on_stream_read_fin_flushed(
event::builder::StreamReadFinFlushed {
capacity,
processing_duration,
},
);
}
Poll::Ready(Ok(len)) => {
self.shared.common.publisher().on_stream_read_flushed(
event::builder::StreamReadFlushed {
capacity,
committed_len: *len,
processing_duration,
},
);
}
Poll::Ready(Err(error)) => {
let errno = error.raw_os_error();
self.shared.common.publisher().on_stream_read_errored(
event::builder::StreamReadErrored {
capacity,
processing_duration,
errno,
},
);
}
Poll::Pending => {
self.shared.common.publisher().on_stream_read_blocked(
event::builder::StreamReadBlocked {
capacity,
processing_duration,
},
);
}
};
}
}
#[cfg(feature = "tokio")]
impl<Sub> tokio::io::AsyncRead for Reader<Sub>
where
Sub: event::Subscriber,
{
#[inline]
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut tokio::io::ReadBuf<'_>,
) -> Poll<io::Result<()>> {
let mut buf = buffer::writer::storage::BufMut::new(buf);
ready!(self.poll_read_into(cx, &mut buf))?;
Ok(()).into()
}
}
impl<Sub> Drop for Reader<Sub>
where
Sub: event::Subscriber,
{
#[inline]
fn drop(&mut self) {
let inner = unsafe {
ManuallyDrop::take(&mut self.0)
};
inner.shutdown();
}
}
pub struct Shutdown<Sub: event::Subscriber>(Box<Inner<Sub>>);
impl<Sub> core::future::Future for Shutdown<Sub>
where
Sub: event::Subscriber,
{
type Output = ();
#[inline]
fn poll(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<()> {
let mut storage = buffer::writer::storage::Empty;
let _ = ready!(self.0.poll_read_into(cx, &mut storage.track_write()));
Poll::Ready(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[allow(dead_code)]
fn shutdown_traits_test<Sub>(shutdown: &Shutdown<Sub>)
where
Sub: event::Subscriber,
{
use crate::testing::*;
assert_send(shutdown);
assert_sync(shutdown);
assert_static(shutdown);
}
}