use crate::{
clock,
credentials::Id,
event::{self, ConnectionPublisher},
msg,
stream::{
pacer, runtime,
send::{flow, queue},
shared::{ArcShared, ShutdownKind},
socket,
},
};
use core::{
fmt,
pin::Pin,
sync::atomic::Ordering,
task::{Context, Poll},
};
use s2n_quic_core::{buffer, ensure, ready, task::waker, time::Timestamp};
use std::{io, net::SocketAddr};
use tracing::trace;
mod builder;
pub mod state;
pub mod transmission;
use crate::stream::socket::Application;
pub use builder::Builder;
pub struct Writer<Sub: event::Subscriber>(Box<Inner<Sub>>);
struct Inner<Sub>
where
Sub: event::Subscriber,
{
shared: ArcShared<Sub>,
sockets: socket::ArcApplication,
queue: queue::Queue,
pacer: pacer::Naive,
status: Status,
runtime: runtime::ArcHandle<Sub>,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
enum Status {
#[default]
Open,
WroteFin,
Shutdown,
}
impl<Sub> fmt::Debug for Writer<Sub>
where
Sub: event::Subscriber,
{
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
let mut s = f.debug_struct("Writer");
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()
}
}
impl<Sub> Writer<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.write_application().local_addr()
}
#[inline]
pub fn path_secret_id(&self) -> &Id {
&self.0.shared.credentials().id
}
#[inline]
pub fn protocol(&self) -> socket::Protocol {
self.0.sockets.protocol()
}
#[inline]
pub async fn write_from<S>(&mut self, buf: &mut S) -> io::Result<usize>
where
S: buffer::reader::storage::Infallible,
{
core::future::poll_fn(|cx| self.poll_write_from(cx, buf, false)).await
}
#[inline]
pub async fn write_all_from<S>(&mut self, buf: &mut S) -> io::Result<usize>
where
S: buffer::reader::storage::Infallible,
{
let mut len = 0;
loop {
len += self.write_from(buf).await?;
if buf.buffer_is_empty() {
return Ok(len);
}
}
}
#[inline]
pub async fn write_from_fin<S>(&mut self, buf: &mut S) -> io::Result<usize>
where
S: buffer::reader::storage::Infallible,
{
core::future::poll_fn(|cx| self.poll_write_from(cx, buf, true)).await
}
#[inline]
pub async fn write_all_from_fin<S>(&mut self, buf: &mut S) -> io::Result<usize>
where
S: buffer::reader::storage::Infallible,
{
let mut len = 0;
loop {
len += self.write_from_fin(buf).await?;
if buf.buffer_is_empty() {
return Ok(len);
}
}
}
#[inline]
pub fn poll_write_from<S>(
&mut self,
cx: &mut Context,
buf: &mut S,
is_fin: bool,
) -> Poll<io::Result<usize>>
where
S: buffer::reader::storage::Infallible,
{
let start_time = self.0.shared.clock.get_time();
let provided_len = buf.buffered_len();
let res = waker::debug_assert_contract(cx, |cx| {
let res = ready!(self.0.poll_write_from(cx, buf, is_fin));
if res.is_err() {
let _ = self.0.shutdown(ShutdownType::Drop {
is_panicking: false,
});
}
res.into()
});
self.0
.publish_write_events(provided_len, is_fin, start_time, &res);
res
}
pub fn shutdown(&mut self) -> io::Result<()> {
self.0.shutdown(ShutdownType::Explicit)
}
pub fn query_event_context<C: 'static, R>(&self, query: impl FnOnce(&C) -> R) -> Option<R> {
let ctxt = &self.0.shared.common.subscriber.context;
let mut query = s2n_quic_core::query::Once::new(query);
Sub::query(ctxt, &mut query);
let res: Result<_, _> = query.into();
match res {
Ok(r) => Some(r),
Err(s2n_quic_core::query::Error::ConnectionLockPoisoned) => unreachable!(),
Err(s2n_quic_core::query::Error::ContextTypeMismatch) => None,
Err(_) => None,
}
}
pub fn peer_cert_chain(&self) -> Option<&crate::stream::tls::CertificateChain> {
self.0.shared.s2n_connection.as_ref()?.peer_cert_chain()
}
}
impl<Sub> Inner<Sub>
where
Sub: event::Subscriber,
{
#[inline(always)]
fn poll_write_from<S>(
&mut self,
cx: &mut Context,
buf: &mut S,
is_fin: bool,
) -> Poll<io::Result<usize>>
where
S: buffer::reader::storage::Infallible,
{
let flushed_len = ready!(self.poll_flush_buffer(cx, buf.buffered_len()))?;
ensure!(flushed_len == 0, Ok(flushed_len).into());
if !matches!(self.status, Status::Open) {
ensure!(
buf.buffer_is_empty() && is_fin,
Err(io::Error::from(io::ErrorKind::BrokenPipe)).into()
);
return Ok(0).into();
}
ensure!(self.queue.is_empty(), Ok(flushed_len).into());
let app = self.shared.application();
let max_header_len = app.max_header_len();
let max_segments = self.shared.gso.max_segments();
let initial_len = buf.buffered_len();
let mut request = flow::Request {
len: initial_len,
initial_len,
is_fin,
};
let path = self.shared.sender.path.load();
let features = self.sockets.features();
if !features.is_flow_controlled() {
request.clamp(path.max_flow_credits(max_header_len, max_segments));
}
let credits = ready!(self.shared.sender.flow.poll_acquire(cx, request, &features))?;
if credits.is_fin {
self.status = Status::WroteFin;
}
trace!(?credits);
let mut batch = if features.is_reliable() {
None
} else {
let batch = self
.shared
.sender
.application_transmission_queue
.alloc_batch(msg::segment::MAX_COUNT);
Some(batch)
};
let stream_id = self.shared.stream_id();
let local_queue_id = self.shared.local_queue_id();
self.queue.push_buffer(
buf,
&mut batch,
max_segments,
&self.shared.sender.segment_alloc,
|output, buf| {
self.shared.crypto.seal_with(
|sealer| {
app.transmit(
credits,
&path,
buf,
&self.shared.sender.packet_number,
sealer,
self.shared.credentials(),
&self.shared.s2n_connection,
&stream_id,
local_queue_id,
&clock::Cached::new(&self.shared.clock),
output,
&features,
&self.shared.publisher(),
)
},
|sealer| {
if features.is_reliable() {
sealer.update(&self.shared.clock, &self.shared.subscriber);
} else {
}
},
)
},
)?;
if let Some(batch) = batch {
self.shared.sender.push_to_worker(batch)?;
}
self.poll_flush_buffer(cx, usize::MAX)
}
#[inline]
fn poll_flush_buffer(
&mut self,
cx: &mut Context,
limit: usize,
) -> Poll<Result<usize, io::Error>> {
if !self.queue.is_empty() {
ready!(self.pacer.poll_pacing(cx, &self.shared.clock));
}
let len = ready!(self.queue.poll_flush(
cx,
limit,
self.sockets.write_application(),
&msg::addr::Addr::new(self.shared.remote_addr()),
&self.shared.sender.segment_alloc,
&self.shared.gso,
&self.shared.clock,
&self.shared.subscriber,
))?;
Ok(len).into()
}
#[inline]
fn shutdown(&mut self, ty: ShutdownType) -> io::Result<()> {
ensure!(
self.status != Status::Shutdown,
if cfg!(target_os = "macos") {
Err(io::ErrorKind::NotConnected.into())
} else {
Ok(())
}
);
if !matches!(ty, ShutdownType::Drop { is_panicking: true }) {
let waker = s2n_quic_core::task::waker::noop();
let mut cx = core::task::Context::from_waker(&waker);
let _ = self.poll_write_from(&mut cx, &mut buffer::reader::storage::Empty, true);
}
self.status = Status::Shutdown;
self.shared
.common
.closed_halves
.fetch_add(1, Ordering::Relaxed);
let queue = core::mem::take(&mut self.queue);
if matches!(ty, ShutdownType::Explicit) && queue.is_empty() {
self.sockets.write_application().send_finish()?;
}
let buffer_len = queue.accepted_len();
if !self.sockets.features().is_stream() {
self.shared
.publisher()
.on_stream_write_shutdown(event::builder::StreamWriteShutdown {
background: false,
buffer_len,
});
let is_panicking = matches!(ty, ShutdownType::Drop { is_panicking: true });
let shutdown_kind = if is_panicking {
ShutdownKind::Panicking
} else {
ShutdownKind::Normal
};
self.shared.sender.shutdown(queue, shutdown_kind);
return Ok(());
}
let background = !queue.is_empty();
self.shared
.publisher()
.on_stream_write_shutdown(event::builder::StreamWriteShutdown {
background,
buffer_len,
});
if background {
let shared = self.shared.clone();
let sockets = self.sockets.clone();
self.runtime.spawn_send_shutdown(Shutdown {
queue,
shared,
sockets,
ty,
});
}
Ok(())
}
#[inline(always)]
fn publish_write_events(
&self,
provided_len: usize,
is_fin: bool,
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(len)) if is_fin => {
self.shared.common.publisher().on_stream_write_fin_flushed(
event::builder::StreamWriteFinFlushed {
provided_len,
committed_len: *len,
processing_duration,
},
);
}
Poll::Ready(Ok(len)) => {
self.shared.common.publisher().on_stream_write_flushed(
event::builder::StreamWriteFlushed {
provided_len,
committed_len: *len,
processing_duration,
},
);
}
Poll::Ready(Err(error)) => {
let errno = error.raw_os_error();
self.shared.common.publisher().on_stream_write_errored(
event::builder::StreamWriteErrored {
provided_len,
is_fin,
processing_duration,
errno,
},
);
}
Poll::Pending => {
self.shared.common.publisher().on_stream_write_blocked(
event::builder::StreamWriteBlocked {
provided_len,
is_fin,
processing_duration,
},
);
}
};
}
}
#[cfg(feature = "tokio")]
impl<Sub> tokio::io::AsyncWrite for Writer<Sub>
where
Sub: event::Subscriber,
{
#[inline]
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
mut buf: &[u8],
) -> Poll<Result<usize, io::Error>> {
self.poll_write_from(cx, &mut buf, false)
}
#[inline]
fn poll_write_vectored(
mut self: Pin<&mut Self>,
cx: &mut Context,
buf: &[std::io::IoSlice],
) -> Poll<Result<usize, io::Error>> {
let mut buf = buffer::reader::storage::IoSlice::new(buf);
self.poll_write_from(cx, &mut buf, false)
}
#[inline]
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
Poll::Ready(Ok(()))
}
#[inline]
fn poll_shutdown(
mut self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Result<(), io::Error>> {
self.0.shutdown(ShutdownType::Explicit).into()
}
#[inline(always)]
fn is_write_vectored(&self) -> bool {
true
}
}
impl<Sub> Drop for Writer<Sub>
where
Sub: event::Subscriber,
{
#[inline]
fn drop(&mut self) {
let _ = self.0.shutdown(ShutdownType::Drop {
is_panicking: std::thread::panicking(),
});
}
}
#[derive(Clone, Copy, Debug)]
enum ShutdownType {
Explicit,
Drop { is_panicking: bool },
}
pub struct Shutdown<Sub>
where
Sub: event::Subscriber,
{
queue: queue::Queue,
shared: ArcShared<Sub>,
sockets: socket::ArcApplication,
ty: ShutdownType,
}
impl<Sub> core::future::Future for Shutdown<Sub>
where
Sub: event::Subscriber,
{
type Output = ();
#[inline]
fn poll(self: Pin<&mut Self>, cx: &mut Context) -> Poll<()> {
let Self {
queue,
sockets,
shared,
ty,
} = self.get_mut();
let _ = ready!(queue.poll_flush(
cx,
usize::MAX,
sockets.write_application(),
&msg::addr::Addr::new(shared.remote_addr()),
&shared.sender.segment_alloc,
&shared.gso,
&shared.clock,
&shared.subscriber,
));
if matches!(ty, ShutdownType::Explicit) {
let _ = sockets.write_application().send_finish();
}
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);
}
}