use super::body::GuardedBody;
use super::disconnect::ConnectionLiveness;
use super::handle::{ConnCtx, handle_request};
use super::router::ServerDispatch;
use super::server_lifecycle::{ConnectionLifecycle, ServerControl, wait_shutdown_control};
use crate::net::accept;
use crate::{RuntimeError, net};
use std::sync::Arc;
const CONNECTION_DRAIN_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(15);
pub(super) struct ConnectionState {
router: Arc<ServerDispatch>,
ctx: Arc<ConnCtx>,
lifecycle: ConnectionLifecycle,
keepalive_timeout: std::time::Duration,
remote_addr: Option<std::net::IpAddr>,
}
impl ConnectionState {
pub(super) fn new(
router: Arc<ServerDispatch>,
ctx: Arc<ConnCtx>,
lifecycle: ConnectionLifecycle,
keepalive_timeout: std::time::Duration,
remote_addr: Option<std::net::IpAddr>,
) -> Self {
Self {
router,
ctx,
lifecycle,
keepalive_timeout,
remote_addr,
}
}
}
fn connection_service(
state: ConnectionState,
liveness: Arc<ConnectionLiveness>,
) -> impl hyper::service::Service<
hyper::Request<hyper::body::Incoming>,
Response = hyper::Response<GuardedBody>,
Error = std::convert::Infallible,
Future: Send + 'static,
> + use<> {
let ConnectionState {
router,
ctx,
lifecycle,
remote_addr,
..
} = state;
let lifecycle = Arc::new(lifecycle);
hyper::service::service_fn(move |request| {
let router = Arc::clone(&router);
let ctx = Arc::clone(&ctx);
let lifecycle = Arc::clone(&lifecycle);
let liveness = Arc::clone(&liveness);
async move { serve_request(request, &router, &ctx, remote_addr, &lifecycle, liveness).await }
})
}
pub(super) async fn accept_loop(
listener: &net::Listener,
router: Arc<ServerDispatch>,
ctx: Arc<ConnCtx>,
shutdown: crate::runtime_state::ShutdownSignal,
keepalive_timeout: std::time::Duration,
tls_acceptor: Option<tokio_rustls::TlsAcceptor>,
conn_limit: Option<Arc<tokio::sync::Semaphore>>,
) -> Result<(), RuntimeError> {
match &listener.inner {
net::ListenerInner::Tcp(tcp) => {
accept_tcp(
tcp,
router,
ctx,
shutdown,
keepalive_timeout,
tls_acceptor,
conn_limit,
)
.await
}
net::ListenerInner::Unix(unix, _) => {
accept_unix(unix, router, ctx, shutdown, keepalive_timeout, conn_limit).await
}
}
}
pub(super) async fn accept_tcp(
listener: &tokio::net::TcpListener,
router: Arc<ServerDispatch>,
ctx: Arc<ConnCtx>,
shutdown: crate::runtime_state::ShutdownSignal,
keepalive_timeout: std::time::Duration,
tls_acceptor: Option<tokio_rustls::TlsAcceptor>,
conn_limit: Option<Arc<tokio::sync::Semaphore>>,
) -> Result<(), RuntimeError> {
let script = listener
.local_addr()
.ok()
.and_then(super::mock::lifecycle_script);
accept::accept_loop_with_permit(
listener,
&shutdown,
conn_limit.as_ref(),
script.as_ref(),
|(stream, addr), permit| {
let state =
synchronous_state(&router, &ctx, permit, keepalive_timeout, Some(addr.ip()));
let shutdown = shutdown.clone();
let acceptor = tls_acceptor.clone();
async move {
match acceptor {
Some(a) => serve_tls_connection(stream, a, state, shutdown).await,
None => serve_stream(stream, state, shutdown).await,
}
}
},
)
.await
}
async fn accept_unix(
listener: &tokio::net::UnixListener,
router: Arc<ServerDispatch>,
ctx: Arc<ConnCtx>,
shutdown: crate::runtime_state::ShutdownSignal,
keepalive_timeout: std::time::Duration,
conn_limit: Option<Arc<tokio::sync::Semaphore>>,
) -> Result<(), RuntimeError> {
accept::accept_loop_with_permit(
listener,
&shutdown,
conn_limit.as_ref(),
None,
|stream, permit| {
let state = synchronous_state(&router, &ctx, permit, keepalive_timeout, None);
let shutdown = shutdown.clone();
async move {
serve_stream(stream, state, shutdown).await;
}
},
)
.await
}
fn synchronous_state(
router: &Arc<ServerDispatch>,
ctx: &Arc<ConnCtx>,
permit: Option<tokio::sync::OwnedSemaphorePermit>,
keepalive_timeout: std::time::Duration,
remote_addr: Option<std::net::IpAddr>,
) -> ConnectionState {
ConnectionState::new(
Arc::clone(router),
Arc::clone(ctx),
ConnectionLifecycle::synchronous(permit),
keepalive_timeout,
remote_addr,
)
}
async fn serve_stream<S>(
stream: S,
state: ConnectionState,
shutdown: crate::runtime_state::ShutdownSignal,
) where
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
{
let liveness = ConnectionLiveness::latched(shutdown.flag());
let io = hyper_util::rt::TokioIo::new(liveness.wrap(stream));
serve_io(io, state, shutdown, liveness).await;
}
async fn serve_tls_connection(
stream: tokio::net::TcpStream,
acceptor: tokio_rustls::TlsAcceptor,
state: ConnectionState,
shutdown: crate::runtime_state::ShutdownSignal,
) {
let tls_stream = match accept::tls_handshake(stream, &acceptor).await {
Some(s) => s,
None => return,
};
serve_stream(tls_stream, state, shutdown).await;
}
async fn serve_io<I>(
io: hyper_util::rt::TokioIo<I>,
state: ConnectionState,
shutdown: crate::runtime_state::ShutdownSignal,
liveness: Arc<ConnectionLiveness>,
) where
I: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
{
let connection = build_connection(io, state, liveness);
tokio::pin!(connection);
drive_connection_until_shutdown(
connection.as_mut(),
async {
shutdown.wait().await;
ConnectionShutdown::Graceful
},
DrainPolicy::Bounded(CONNECTION_DRAIN_TIMEOUT),
)
.await;
}
pub(super) async fn serve_owned_connection(
stream: tokio::net::TcpStream,
tls_acceptor: Option<tokio_rustls::TlsAcceptor>,
state: ConnectionState,
control: tokio::sync::watch::Receiver<ServerControl>,
) {
let liveness = ConnectionLiveness::controlled(control.clone());
match tls_acceptor {
Some(acceptor) => serve_owned_tls(stream, acceptor, state, liveness, control).await,
None => serve_owned_stream(stream, state, liveness, control).await,
}
}
async fn serve_owned_tls(
stream: tokio::net::TcpStream,
acceptor: tokio_rustls::TlsAcceptor,
state: ConnectionState,
liveness: Arc<ConnectionLiveness>,
mut control: tokio::sync::watch::Receiver<ServerControl>,
) {
let handshake = accept::tls_handshake(stream, &acceptor);
tokio::pin!(handshake);
let tls_stream = tokio::select! {
biased;
_ = wait_connection_shutdown(&mut control) => return,
stream = &mut handshake => match stream {
Some(stream) => stream,
None => return,
},
};
serve_owned_stream(tls_stream, state, liveness, control).await;
}
#[cfg(feature = "ws")]
const OWNED_TRANSPORT_BUFFER_SIZE: usize = 8 * 1024;
#[cfg(feature = "ws")]
struct TransportStream<S> {
reader: Option<tokio::io::ReadHalf<S>>,
writer: tokio::io::WriteHalf<S>,
activation: Option<tokio::sync::oneshot::Sender<tokio::io::ReadHalf<S>>>,
incoming: tokio::sync::mpsc::Receiver<Result<bytes::Bytes, std::io::Error>>,
pending: Option<bytes::Bytes>,
reader_abort: tokio::task::AbortHandle,
}
#[cfg(feature = "ws")]
impl<S> Drop for TransportStream<S> {
fn drop(&mut self) {
self.reader_abort.abort();
}
}
#[cfg(feature = "ws")]
impl<S> tokio::io::AsyncRead for TransportStream<S>
where
S: tokio::io::AsyncRead + Unpin,
{
fn poll_read(
mut self: std::pin::Pin<&mut Self>,
context: &mut std::task::Context<'_>,
buffer: &mut tokio::io::ReadBuf<'_>,
) -> std::task::Poll<Result<(), std::io::Error>> {
if let Some(reader) = self.reader.as_mut() {
let filled = buffer.filled().len();
let result = std::pin::Pin::new(reader).poll_read(context, buffer);
let received_bytes =
matches!(result, std::task::Poll::Ready(Ok(()))) && buffer.filled().len() > filled;
self.activate_after_read(received_bytes);
return result;
}
if self.copy_pending(buffer) {
return std::task::Poll::Ready(Ok(()));
}
match self.incoming.poll_recv(context) {
std::task::Poll::Ready(Some(Ok(bytes))) => {
self.pending = Some(bytes);
self.copy_pending(buffer);
std::task::Poll::Ready(Ok(()))
}
std::task::Poll::Ready(Some(Err(error))) => std::task::Poll::Ready(Err(error)),
std::task::Poll::Ready(None) => std::task::Poll::Ready(Ok(())),
std::task::Poll::Pending => std::task::Poll::Pending,
}
}
}
#[cfg(feature = "ws")]
impl<S> TransportStream<S> {
fn activate_after_read(&mut self, received_bytes: bool) {
match received_bytes {
true => self.activate_reader(),
false => {}
}
}
fn activate_reader(&mut self) {
match (self.reader.take(), self.activation.take()) {
(Some(reader), Some(activation)) => hand_off_reader(activation, reader),
_ => {}
}
}
fn copy_pending(&mut self, buffer: &mut tokio::io::ReadBuf<'_>) -> bool {
let mut bytes = match self.pending.take() {
Some(bytes) => bytes,
None => return false,
};
let count = bytes.len().min(buffer.remaining());
buffer.put_slice(&bytes[..count]);
bytes::Buf::advance(&mut bytes, count);
if !bytes.is_empty() {
self.pending = Some(bytes);
}
true
}
}
#[cfg(feature = "ws")]
fn hand_off_reader<S>(
activation: tokio::sync::oneshot::Sender<tokio::io::ReadHalf<S>>,
reader: tokio::io::ReadHalf<S>,
) {
match activation.send(reader) {
Ok(()) => {}
Err(_) => tracing::debug!("upgrade transport reader is gone; connection read half dropped"),
}
}
#[cfg(feature = "ws")]
impl<S> tokio::io::AsyncWrite for TransportStream<S>
where
S: tokio::io::AsyncWrite + Unpin,
{
fn poll_write(
mut self: std::pin::Pin<&mut Self>,
context: &mut std::task::Context<'_>,
buffer: &[u8],
) -> std::task::Poll<Result<usize, std::io::Error>> {
std::pin::Pin::new(&mut self.writer).poll_write(context, buffer)
}
fn poll_flush(
mut self: std::pin::Pin<&mut Self>,
context: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), std::io::Error>> {
std::pin::Pin::new(&mut self.writer).poll_flush(context)
}
fn poll_shutdown(
mut self: std::pin::Pin<&mut Self>,
context: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), std::io::Error>> {
std::pin::Pin::new(&mut self.writer).poll_shutdown(context)
}
fn is_write_vectored(&self) -> bool {
self.writer.is_write_vectored()
}
fn poll_write_vectored(
mut self: std::pin::Pin<&mut Self>,
context: &mut std::task::Context<'_>,
buffers: &[std::io::IoSlice<'_>],
) -> std::task::Poll<Result<usize, std::io::Error>> {
std::pin::Pin::new(&mut self.writer).poll_write_vectored(context, buffers)
}
}
#[cfg(feature = "ws")]
struct OwnedTransport {
handle: Option<tokio::task::JoinHandle<()>>,
peer_closed: Option<tokio::sync::oneshot::Receiver<()>>,
barrier: tokio::sync::mpsc::Sender<tokio::sync::oneshot::Sender<()>>,
}
#[cfg(feature = "ws")]
impl OwnedTransport {
fn new<S>(
stream: S,
script: Option<Arc<super::mock::LifecycleScript>>,
) -> (TransportStream<S>, Self)
where
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
{
let (reader, writer) = tokio::io::split(stream);
let (activation, activated) = tokio::sync::oneshot::channel();
let (incoming_sender, incoming) = tokio::sync::mpsc::channel(1);
let (peer_closed_sender, peer_closed) = tokio::sync::oneshot::channel();
let (barrier, barriers) = tokio::sync::mpsc::channel(1);
let handle = tokio::spawn(drive_owned_reader(
activated,
incoming_sender,
peer_closed_sender,
barriers,
script,
));
let reader_abort = handle.abort_handle();
(
TransportStream {
reader: Some(reader),
writer,
activation: Some(activation),
incoming,
pending: None,
reader_abort,
},
Self {
handle: Some(handle),
peer_closed: Some(peer_closed),
barrier,
},
)
}
async fn peer_closed(&mut self) {
wait_for_peer_close(&mut self.peer_closed).await;
}
async fn join(&mut self) {
match self.handle.take() {
None => {}
Some(handle) => log_reader_join(handle.await),
}
}
async fn close(&mut self) {
if let Some(handle) = self.handle.as_ref() {
handle.abort();
}
self.join().await;
}
async fn peer_remains_open(&mut self) -> bool {
let (acknowledgement, acknowledged) = tokio::sync::oneshot::channel();
if self.barrier.send(acknowledgement).await.is_err() {
return false;
}
tokio::select! {
biased;
() = wait_for_peer_close(&mut self.peer_closed) => false,
result = acknowledged => result.is_ok(),
}
}
}
#[cfg(feature = "ws")]
fn log_reader_join(outcome: Result<(), tokio::task::JoinError>) {
match outcome {
Err(error) if error.is_panic() => {
tracing::warn!("upgrade transport reader panicked: {error}");
}
Ok(()) | Err(_) => {}
}
}
#[cfg(feature = "ws")]
fn report_unsent_read_error(
outcome: Result<(), tokio::sync::mpsc::error::SendError<Result<bytes::Bytes, std::io::Error>>>,
) {
match outcome {
Err(tokio::sync::mpsc::error::SendError(Err(error))) => {
tracing::debug!("upgrade transport reader is gone; connection read failed: {error}");
}
Ok(()) | Err(_) => {}
}
}
#[cfg(feature = "ws")]
async fn wait_for_peer_close(peer_closed: &mut Option<tokio::sync::oneshot::Receiver<()>>) {
match peer_closed.as_mut() {
Some(receiver) => {
let _ = receiver.await;
*peer_closed = None;
}
None => std::future::pending().await,
}
}
#[cfg(feature = "ws")]
impl Drop for OwnedTransport {
fn drop(&mut self) {
if let Some(handle) = self.handle.as_ref() {
handle.abort();
}
}
}
#[cfg(feature = "ws")]
async fn drive_owned_reader<S>(
activation: tokio::sync::oneshot::Receiver<tokio::io::ReadHalf<S>>,
incoming: tokio::sync::mpsc::Sender<Result<bytes::Bytes, std::io::Error>>,
peer_closed: tokio::sync::oneshot::Sender<()>,
mut barriers: tokio::sync::mpsc::Receiver<tokio::sync::oneshot::Sender<()>>,
script: Option<Arc<super::mock::LifecycleScript>>,
) where
S: tokio::io::AsyncRead + Unpin,
{
use tokio::io::AsyncReadExt;
let mut reader = match activation.await {
Ok(reader) => reader,
Err(_) => return,
};
let mut buffer = bytes::BytesMut::with_capacity(OWNED_TRANSPORT_BUFFER_SIZE);
loop {
buffer.reserve(OWNED_TRANSPORT_BUFFER_SIZE);
let result = tokio::select! {
biased;
result = reader.read_buf(&mut buffer) => result,
barrier = barriers.recv() => match barrier {
Some(barrier) => {
let _ = barrier.send(());
continue;
}
None => break,
},
};
let count = match result {
Ok(count) => count,
Err(error) => {
report_unsent_read_error(incoming.send(Err(error)).await);
break;
}
};
if count == 0 {
break;
}
if incoming.send(Ok(buffer.split().freeze())).await.is_err() {
break;
}
}
let _ = peer_closed.send(());
super::mock::LifecycleScript::pause_at(
script.as_deref(),
super::mock::LifecycleCheckpoint::UpgradePeerClosed,
)
.await;
}
async fn serve_owned_stream<S>(
stream: S,
state: ConnectionState,
liveness: Arc<ConnectionLiveness>,
control: tokio::sync::watch::Receiver<ServerControl>,
) where
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
{
let stream = liveness.wrap(stream);
#[cfg(feature = "ws")]
let (stream, transport) = OwnedTransport::new(stream, state.lifecycle.script());
let io = hyper_util::rt::TokioIo::new(stream);
serve_owned_io(
io,
#[cfg(feature = "ws")]
transport,
state,
liveness,
control,
)
.await;
}
async fn serve_owned_io<I>(
io: hyper_util::rt::TokioIo<I>,
#[cfg(feature = "ws")] mut transport: OwnedTransport,
state: ConnectionState,
liveness: Arc<ConnectionLiveness>,
mut control: tokio::sync::watch::Receiver<ServerControl>,
) where
I: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
{
#[cfg(feature = "ws")]
let mut state = state;
#[cfg(feature = "ws")]
let mut upgrade_transport = state.lifecycle.bind_upgrade_transport();
let connection = build_connection(io, state, liveness);
tokio::pin!(connection);
#[cfg(feature = "ws")]
let mut peer_closed = false;
#[cfg(feature = "ws")]
loop {
let event = next_owned_connection_event(
connection.as_mut(),
&mut control,
&mut upgrade_transport,
&mut transport,
)
.await;
let flow = match event {
OwnedConnectionEvent::Complete(result) => {
finish_owned_connection(result, &mut upgrade_transport, &mut transport).await;
ConnectionFlow::Finished
}
OwnedConnectionEvent::Shutdown(mode) => {
shutdown_owned_connection(
mode,
connection.as_mut(),
&mut upgrade_transport,
&mut transport,
)
.await;
ConnectionFlow::Finished
}
OwnedConnectionEvent::Registration(Some(registration)) => {
serve_upgrade_registration(
registration,
peer_closed,
connection.as_mut(),
&mut control,
&mut upgrade_transport,
&mut transport,
)
.await
}
OwnedConnectionEvent::Registration(None) => ConnectionFlow::Serving,
OwnedConnectionEvent::PeerClosed => {
peer_closed = true;
ConnectionFlow::Serving
}
};
match flow {
ConnectionFlow::Finished => return,
ConnectionFlow::Serving => {}
}
}
#[cfg(not(feature = "ws"))]
drive_connection_until_shutdown(
connection.as_mut(),
wait_connection_shutdown(&mut control),
DrainPolicy::Supervised,
)
.await;
}
type ConnectionResult = Result<(), Box<dyn std::error::Error + Send + Sync>>;
trait HyperConnection: std::future::Future<Output = ConnectionResult> {
fn begin_graceful_shutdown(self: std::pin::Pin<&mut Self>);
}
impl<I, S, B, E> HyperConnection
for hyper_util::server::conn::auto::UpgradeableConnection<'static, I, S, E>
where
S: hyper::service::Service<
hyper::Request<hyper::body::Incoming>,
Response = hyper::Response<B>,
>,
S::Future: 'static,
S::Error: Into<Box<dyn std::error::Error + Send + Sync>>,
B: hyper::body::Body + 'static,
B::Error: Into<Box<dyn std::error::Error + Send + Sync>>,
I: hyper::rt::Read + hyper::rt::Write + Unpin + Send + 'static,
E: hyper_util::server::conn::auto::HttpServerConnExec<S::Future, B>,
{
fn begin_graceful_shutdown(self: std::pin::Pin<&mut Self>) {
self.graceful_shutdown();
}
}
fn build_connection<I>(
io: hyper_util::rt::TokioIo<I>,
state: ConnectionState,
liveness: Arc<ConnectionLiveness>,
) -> impl HyperConnection + Send + use<I>
where
I: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
{
let builder = connection_builder(state.keepalive_timeout);
let service = connection_service(state, liveness);
builder
.serve_connection_with_upgrades(io, service)
.into_owned()
}
enum DrainPolicy {
Bounded(std::time::Duration),
Supervised,
}
enum ConnectionShutdown {
Graceful,
Abort,
}
enum ConnectionPhase {
Serving,
Draining,
}
async fn wait_connection_shutdown(
control: &mut tokio::sync::watch::Receiver<ServerControl>,
) -> ConnectionShutdown {
match wait_shutdown_control(control).await {
ServerControl::Graceful => ConnectionShutdown::Graceful,
ServerControl::Abort => ConnectionShutdown::Abort,
ServerControl::Running => {
tracing::error!("shutdown control reported a running server; connection kept serving");
std::future::pending().await
}
}
}
#[cfg(feature = "ws")]
enum ConnectionFlow {
Serving,
Finished,
}
#[cfg(feature = "ws")]
async fn serve_upgrade_registration<C>(
mut registration: super::server_lifecycle::TransportRegistration,
peer_closed: bool,
mut connection: std::pin::Pin<&mut C>,
control: &mut tokio::sync::watch::Receiver<ServerControl>,
upgrade_transport: &mut super::server_lifecycle::UpgradeTransportOwner,
transport: &mut OwnedTransport,
) -> ConnectionFlow
where
C: HyperConnection,
{
let cancellation = registration.prepare();
cancel_closed_registration(peer_closed, cancellation.as_ref());
let registration = registration.register();
tokio::pin!(registration);
let event = await_upgrade_registration(
connection.as_mut(),
control,
registration.as_mut(),
transport,
)
.await;
let settled = finish_interrupted_registration(
event,
connection.as_mut(),
registration.as_mut(),
cancellation.as_ref(),
upgrade_transport,
transport,
)
.await;
let settled = match settled {
Some(outcome) => outcome,
None => return ConnectionFlow::Finished,
};
let retained =
retain_open_upgrade(settled, connection.as_mut(), upgrade_transport, transport).await;
let retained = match retained {
Some(outcome) => outcome,
None => return ConnectionFlow::Finished,
};
let commitment = match retained.complete() {
Some(commitment) => commitment,
None => return ConnectionFlow::Serving,
};
let event = await_upgrade_commitment(connection.as_mut(), control, transport).await;
finish_upgrade_commitment(
event,
commitment,
connection.as_mut(),
upgrade_transport,
transport,
)
.await;
transport.join().await;
ConnectionFlow::Finished
}
#[cfg(feature = "ws")]
async fn finish_owned_connection(
result: ConnectionResult,
upgrade_transport: &mut super::server_lifecycle::UpgradeTransportOwner,
transport: &mut OwnedTransport,
) {
upgrade_transport.cancel();
upgrade_transport.abort_pending().await;
log_connection_result(result, ConnectionPhase::Serving);
transport.close().await;
}
async fn shutdown_hyper_connection<C>(
mode: ConnectionShutdown,
mut connection: std::pin::Pin<&mut C>,
drain: DrainPolicy,
) where
C: HyperConnection,
{
match mode {
ConnectionShutdown::Graceful | ConnectionShutdown::Abort => {
connection.as_mut().begin_graceful_shutdown();
drain_connection(connection, drain).await;
}
}
}
async fn drive_connection_until_shutdown<C, F>(
mut connection: std::pin::Pin<&mut C>,
shutdown: F,
drain: DrainPolicy,
) where
C: HyperConnection,
F: std::future::Future<Output = ConnectionShutdown>,
{
tokio::select! {
biased;
mode = shutdown => shutdown_hyper_connection(mode, connection.as_mut(), drain).await,
result = connection.as_mut() => log_connection_result(result, ConnectionPhase::Serving),
}
}
async fn drain_connection<C>(connection: std::pin::Pin<&mut C>, drain: DrainPolicy)
where
C: HyperConnection,
{
match drain {
DrainPolicy::Supervised => {
log_connection_result(connection.await, ConnectionPhase::Draining)
}
DrainPolicy::Bounded(budget) => drain_within_budget(connection, budget).await,
}
}
async fn drain_within_budget<C>(connection: std::pin::Pin<&mut C>, budget: std::time::Duration)
where
C: HyperConnection,
{
match tokio::time::timeout(budget, connection).await {
Ok(result) => log_connection_result(result, ConnectionPhase::Draining),
Err(_) => tracing::debug!("connection timed out during graceful shutdown"),
}
}
#[cfg(feature = "ws")]
async fn shutdown_owned_connection<C>(
mode: ConnectionShutdown,
connection: std::pin::Pin<&mut C>,
upgrade_transport: &mut super::server_lifecycle::UpgradeTransportOwner,
transport: &mut OwnedTransport,
) where
C: HyperConnection,
{
upgrade_transport.cancel();
upgrade_transport.abort_pending().await;
shutdown_hyper_connection(mode, connection, DrainPolicy::Supervised).await;
transport.close().await;
}
#[cfg(feature = "ws")]
fn cancel_prepared_upgrade(cancellation: Option<&super::server_lifecycle::UpgradeCancellation>) {
match cancellation {
Some(cancellation) => cancellation.cancel(),
None => {}
}
}
#[cfg(feature = "ws")]
fn cancel_closed_registration(
peer_closed: bool,
cancellation: Option<&super::server_lifecycle::UpgradeCancellation>,
) {
match (peer_closed, cancellation) {
(true, Some(cancellation)) => cancellation.cancel(),
_ => {}
}
}
#[cfg(feature = "ws")]
async fn settle_interrupted_registration<R>(
registration: std::pin::Pin<&mut R>,
cancellation: Option<&super::server_lifecycle::UpgradeCancellation>,
upgrade_transport: &mut super::server_lifecycle::UpgradeTransportOwner,
) where
R: std::future::Future<Output = super::server_lifecycle::TransportRegistrationOutcome>,
{
upgrade_transport.cancel();
cancel_prepared_upgrade(cancellation);
let outcome = registration.await;
drop(outcome.complete());
upgrade_transport.abort_pending().await;
}
#[cfg(feature = "ws")]
enum RegistrationInterruption {
Ended(ConnectionResult),
Shutdown(ConnectionShutdown),
PeerClosed,
}
#[cfg(feature = "ws")]
async fn finish_interrupted_registration<C, R>(
event: UpgradeRegistrationEvent,
connection: std::pin::Pin<&mut C>,
registration: std::pin::Pin<&mut R>,
cancellation: Option<&super::server_lifecycle::UpgradeCancellation>,
upgrade_transport: &mut super::server_lifecycle::UpgradeTransportOwner,
transport: &mut OwnedTransport,
) -> Option<super::server_lifecycle::TransportRegistrationOutcome>
where
C: HyperConnection,
R: std::future::Future<Output = super::server_lifecycle::TransportRegistrationOutcome>,
{
let interruption = match event {
UpgradeRegistrationEvent::Registered(outcome) => return Some(outcome),
UpgradeRegistrationEvent::Complete(result) => RegistrationInterruption::Ended(result),
UpgradeRegistrationEvent::Shutdown(mode) => RegistrationInterruption::Shutdown(mode),
UpgradeRegistrationEvent::PeerClosed => RegistrationInterruption::PeerClosed,
};
settle_interrupted_registration(registration, cancellation, upgrade_transport).await;
match interruption {
RegistrationInterruption::Ended(result) => {
log_connection_result(result, ConnectionPhase::Serving)
}
RegistrationInterruption::Shutdown(mode) => {
shutdown_hyper_connection(mode, connection, DrainPolicy::Supervised).await;
}
RegistrationInterruption::PeerClosed => {
log_connection_result(connection.await, ConnectionPhase::Serving)
}
}
transport.close().await;
None
}
#[cfg(feature = "ws")]
async fn cancel_admitted_upgrade<C>(
outcome: super::server_lifecycle::TransportRegistrationOutcome,
connection: std::pin::Pin<&mut C>,
upgrade_transport: &mut super::server_lifecycle::UpgradeTransportOwner,
transport: &mut OwnedTransport,
) where
C: HyperConnection,
{
upgrade_transport.cancel();
outcome.cancel();
upgrade_transport.abort_pending().await;
log_connection_result(connection.await, ConnectionPhase::Serving);
transport.close().await;
}
#[cfg(feature = "ws")]
async fn retain_open_upgrade<C>(
outcome: super::server_lifecycle::TransportRegistrationOutcome,
connection: std::pin::Pin<&mut C>,
upgrade_transport: &mut super::server_lifecycle::UpgradeTransportOwner,
transport: &mut OwnedTransport,
) -> Option<super::server_lifecycle::TransportRegistrationOutcome>
where
C: HyperConnection,
{
let peer_open = match outcome.admitted() {
true => transport.peer_remains_open().await,
false => true,
};
match peer_open {
true => Some(outcome),
false => {
cancel_admitted_upgrade(outcome, connection, upgrade_transport, transport).await;
None
}
}
}
#[cfg(feature = "ws")]
async fn commit_open_transport(
commitment: super::server_lifecycle::UpgradeCommitment,
upgrade_transport: &super::server_lifecycle::UpgradeTransportOwner,
transport: &mut OwnedTransport,
) {
match transport.peer_remains_open().await {
true => {
commitment.commit();
upgrade_transport.commit();
}
false => {
upgrade_transport.cancel();
drop(commitment);
}
}
}
#[cfg(feature = "ws")]
enum AbandonedCommitment {
Ended(ConnectionResult),
Shutdown(ConnectionShutdown),
PeerClosed,
}
#[cfg(feature = "ws")]
async fn finish_upgrade_commitment<C>(
event: UpgradeCommitmentEvent,
commitment: super::server_lifecycle::UpgradeCommitment,
connection: std::pin::Pin<&mut C>,
upgrade_transport: &super::server_lifecycle::UpgradeTransportOwner,
transport: &mut OwnedTransport,
) where
C: HyperConnection,
{
let abandoned = match event {
UpgradeCommitmentEvent::Complete(result) if result.is_ok() => {
commit_open_transport(commitment, upgrade_transport, transport).await;
log_connection_result(result, ConnectionPhase::Serving);
return;
}
UpgradeCommitmentEvent::Complete(result) => AbandonedCommitment::Ended(result),
UpgradeCommitmentEvent::Shutdown(mode) => AbandonedCommitment::Shutdown(mode),
UpgradeCommitmentEvent::PeerClosed => AbandonedCommitment::PeerClosed,
};
upgrade_transport.cancel();
drop(commitment);
match abandoned {
AbandonedCommitment::Ended(result) => {
log_connection_result(result, ConnectionPhase::Serving)
}
AbandonedCommitment::Shutdown(mode) => {
shutdown_hyper_connection(mode, connection, DrainPolicy::Supervised).await;
}
AbandonedCommitment::PeerClosed => {
log_connection_result(connection.await, ConnectionPhase::Serving)
}
}
}
#[cfg(feature = "ws")]
enum OwnedConnectionEvent {
Complete(ConnectionResult),
Shutdown(ConnectionShutdown),
Registration(Option<super::server_lifecycle::TransportRegistration>),
PeerClosed,
}
#[cfg(feature = "ws")]
async fn next_owned_connection_event<C>(
mut connection: std::pin::Pin<&mut C>,
control: &mut tokio::sync::watch::Receiver<ServerControl>,
transport: &mut super::server_lifecycle::UpgradeTransportOwner,
owned_transport: &mut OwnedTransport,
) -> OwnedConnectionEvent
where
C: HyperConnection,
{
tokio::select! {
biased;
() = owned_transport.peer_closed() => OwnedConnectionEvent::PeerClosed,
result = connection.as_mut() => OwnedConnectionEvent::Complete(result),
mode = wait_connection_shutdown(control) => OwnedConnectionEvent::Shutdown(mode),
registration = transport.next_registration() => {
OwnedConnectionEvent::Registration(registration)
}
}
}
#[cfg(feature = "ws")]
enum UpgradeRegistrationEvent {
Complete(ConnectionResult),
Shutdown(ConnectionShutdown),
Registered(super::server_lifecycle::TransportRegistrationOutcome),
PeerClosed,
}
#[cfg(feature = "ws")]
async fn await_upgrade_registration<C, R>(
mut connection: std::pin::Pin<&mut C>,
control: &mut tokio::sync::watch::Receiver<ServerControl>,
registration: std::pin::Pin<&mut R>,
transport: &mut OwnedTransport,
) -> UpgradeRegistrationEvent
where
C: HyperConnection,
R: std::future::Future<Output = super::server_lifecycle::TransportRegistrationOutcome>,
{
tokio::select! {
biased;
() = transport.peer_closed() => UpgradeRegistrationEvent::PeerClosed,
result = connection.as_mut() => UpgradeRegistrationEvent::Complete(result),
mode = wait_connection_shutdown(control) => UpgradeRegistrationEvent::Shutdown(mode),
outcome = registration => UpgradeRegistrationEvent::Registered(outcome),
}
}
#[cfg(feature = "ws")]
enum UpgradeCommitmentEvent {
Complete(ConnectionResult),
Shutdown(ConnectionShutdown),
PeerClosed,
}
#[cfg(feature = "ws")]
async fn await_upgrade_commitment<C>(
mut connection: std::pin::Pin<&mut C>,
control: &mut tokio::sync::watch::Receiver<ServerControl>,
transport: &mut OwnedTransport,
) -> UpgradeCommitmentEvent
where
C: HyperConnection,
{
tokio::select! {
biased;
() = transport.peer_closed() => UpgradeCommitmentEvent::PeerClosed,
result = connection.as_mut() => UpgradeCommitmentEvent::Complete(result),
mode = wait_connection_shutdown(control) => UpgradeCommitmentEvent::Shutdown(mode),
}
}
async fn serve_request(
request: hyper::Request<hyper::body::Incoming>,
router: &ServerDispatch,
ctx: &ConnCtx,
remote_addr: Option<std::net::IpAddr>,
lifecycle: &ConnectionLifecycle,
liveness: Arc<ConnectionLiveness>,
) -> Result<hyper::Response<GuardedBody>, std::convert::Infallible> {
let bodyless_request = request.method() == hyper::Method::HEAD;
let (signal, guard) = liveness.begin_response();
let response = handle_request(request, router, ctx, remote_addr, lifecycle, signal).await?;
Ok(GuardedBody::attach(response, guard, bodyless_request))
}
fn connection_builder(
keepalive_timeout: std::time::Duration,
) -> hyper_util::server::conn::auto::Builder<hyper_util::rt::TokioExecutor> {
let mut builder =
hyper_util::server::conn::auto::Builder::new(hyper_util::rt::TokioExecutor::new());
builder
.http1()
.keep_alive(true)
.timer(hyper_util::rt::TokioTimer::new())
.header_read_timeout(Some(keepalive_timeout));
builder
}
fn log_connection_result(result: ConnectionResult, phase: ConnectionPhase) {
match (result, phase) {
(Ok(()), _) => {}
(Err(ref error), _) if is_benign_hyper_error(&**error) => {}
(Err(error), ConnectionPhase::Draining) => {
tracing::warn!("connection error during shutdown: {error}");
}
(Err(error), ConnectionPhase::Serving) => tracing::warn!("connection error: {error}"),
}
}
fn is_benign_hyper_error(err: &(dyn std::error::Error + 'static)) -> bool {
let mut source: Option<&(dyn std::error::Error + 'static)> = Some(err);
while let Some(e) = source {
match e.downcast_ref::<std::io::Error>() {
Some(io_err) => return crate::error::is_benign_io(io_err),
None => source = e.source(),
}
}
false
}