use super::{accept, manager::WorkerError, LazyBoundStream};
use crate::{
crypto::{self, open::Application},
either::Either,
event::{self, EndpointPublisher, IntoEvent},
msg, packet,
path::secret::{self, map::Bidirectional},
stream::{
endpoint,
environment::tokio::{self as env, Environment},
recv, server, TransportFeatures,
},
uds::{self, sender::SendMsg},
};
use core::{
ops::ControlFlow,
pin::Pin,
task::{self, Poll},
time::Duration,
};
use nix::{sys::time::TimeValLike as _, time::ClockId};
use s2n_codec::{decoder, DecoderError, EncoderLenEstimator};
use s2n_quic_core::{
inet::SocketAddress,
ready,
time::{Clock, Timestamp},
};
use std::{future::Future, io, os::fd::OwnedFd, path::Path};
use tracing::debug;
pub struct Context<Sub>
where
Sub: event::Subscriber + Clone,
{
recv_buffer: msg::recv::Message,
env: Environment<Sub>,
secrets: secret::Map,
accept_flavor: accept::Flavor,
local_port: u16,
}
impl<Sub> Context<Sub>
where
Sub: event::Subscriber + Clone,
{
#[inline]
pub fn new<B: PollBehavior<Sub> + Clone>(acceptor: &super::Acceptor<Sub, B>) -> Self {
Self {
recv_buffer: msg::recv::Message::new(u16::MAX),
env: acceptor.env.clone(),
secrets: acceptor.secrets.clone(),
accept_flavor: acceptor.accept_flavor,
local_port: acceptor.socket.get_ref().local_addr().unwrap().port(),
}
}
}
pub struct Worker<Sub, B>
where
Sub: event::Subscriber + Clone,
B: PollBehavior<Sub>,
{
queue_time: Timestamp,
stream: Option<(LazyBoundStream, SocketAddress)>,
subscriber_ctx: Option<Sub::ConnectionContext>,
state: WorkerState,
poll_behavior: B,
}
impl<Sub, B> Worker<Sub, B>
where
Sub: event::Subscriber + Clone,
B: PollBehavior<Sub>,
{
#[inline]
pub fn new(now: Timestamp, poll_behavior: B) -> Self {
Self {
queue_time: now,
stream: None,
subscriber_ctx: None,
state: WorkerState::Init,
poll_behavior,
}
}
}
impl<Sub, B> super::manager::Worker for Worker<Sub, B>
where
Sub: event::Subscriber + Clone,
B: PollBehavior<Sub>,
{
type ConnectionContext = Sub::ConnectionContext;
type Stream = LazyBoundStream;
type Context = Context<Sub>;
#[inline]
fn replace<Pub, C>(
&mut self,
remote_address: SocketAddress,
stream: LazyBoundStream,
linger: Option<Duration>,
subscriber_ctx: Self::ConnectionContext,
publisher: &Pub,
clock: &C,
) where
Pub: EndpointPublisher,
C: Clock,
{
let _ = stream.set_nodelay(true);
if linger.is_some() {
let _ = stream.set_linger(linger);
}
let now = clock.get_time();
let prev_queue_time = core::mem::replace(&mut self.queue_time, now);
let prev_state = core::mem::replace(&mut self.state, WorkerState::Init);
let prev_stream = self.stream.replace((stream, remote_address));
let prev_ctx = self.subscriber_ctx.replace(subscriber_ctx);
if let Some(remote_address) = prev_stream.map(|(socket, remote_address)| {
if linger.is_none() || linger != Some(Duration::ZERO) {
let _ = socket.set_linger(Some(Duration::ZERO));
}
remote_address
}) {
let sojourn_time = now.saturating_duration_since(prev_queue_time);
let buffer_len = match prev_state {
WorkerState::Init => 0,
WorkerState::Buffering { buffer, .. } => buffer.payload_len(),
WorkerState::Erroring { .. } => 0,
WorkerState::Sending { .. } => 0,
};
publisher.on_acceptor_tcp_stream_replaced(event::builder::AcceptorTcpStreamReplaced {
remote_address: &remote_address,
sojourn_time,
buffer_len,
});
}
if let Some(ctx) = prev_ctx {
let _ = ctx;
}
}
#[inline]
fn poll<Pub, C>(
&mut self,
task_cx: &mut task::Context,
context: &mut Context<Sub>,
publisher: &Pub,
clock: &C,
) -> Poll<Result<ControlFlow<()>, WorkerError>>
where
Pub: EndpointPublisher,
C: Clock,
{
if self.stream.is_none() && !matches!(self.state, WorkerState::Sending { .. }) {
debug_assert!(
false,
"Worker::poll should only be called with an active socket"
);
return Poll::Ready(Ok(ControlFlow::Continue(())));
}
context.recv_buffer.clear();
let res = ready!(self.state.poll::<Sub, Pub, B>(
task_cx,
context,
&mut self.stream,
&mut self.subscriber_ctx,
self.queue_time,
clock.get_time(),
publisher,
&self.poll_behavior
));
self.state = WorkerState::Init;
self.stream = None;
if let Some(ctx) = self.subscriber_ctx.take() {
let _ = ctx;
}
Poll::Ready(res)
}
#[inline]
fn queue_time(&self) -> Timestamp {
self.queue_time
}
#[inline]
fn is_active(&self) -> bool {
if matches!(self.state, WorkerState::Sending { .. }) {
return true;
}
let is_active = self.stream.is_some();
if !is_active {
debug_assert!(matches!(self.state, WorkerState::Init));
debug_assert!(self.subscriber_ctx.is_none());
}
is_active
}
}
pub trait PollBehavior<Sub>
where
Sub: event::Subscriber + Clone,
{
fn poll<Pub>(
&self,
state: &mut WorkerState,
cx: &mut task::Context,
context: &mut Context<Sub>,
stream: &mut Option<(LazyBoundStream, SocketAddress)>,
subscriber_ctx: &mut Option<Sub::ConnectionContext>,
queue_time: Timestamp,
now: Timestamp,
publisher: &Pub,
) -> Poll<Result<ControlFlow<()>, WorkerError>>
where
Pub: EndpointPublisher,
Self: Sized;
}
pub enum WorkerState {
Init,
Buffering {
buffer: msg::recv::Message,
blocked_count: usize,
},
Erroring {
offset: usize,
buffer: Vec<u8>,
error: WorkerError,
},
Sending {
future: uds::sender::SendMsg,
event_data: SocketEventData,
},
}
impl WorkerState {
fn poll<Sub, Pub, B>(
&mut self,
cx: &mut task::Context,
context: &mut Context<Sub>,
stream: &mut Option<(LazyBoundStream, SocketAddress)>,
subscriber_ctx: &mut Option<Sub::ConnectionContext>,
queue_time: Timestamp,
now: Timestamp,
publisher: &Pub,
poll_behavior: &B,
) -> Poll<Result<ControlFlow<()>, WorkerError>>
where
Sub: event::Subscriber + Clone,
Pub: EndpointPublisher,
B: PollBehavior<Sub>,
{
poll_behavior.poll(
self,
cx,
context,
stream,
subscriber_ctx,
queue_time,
now,
publisher,
)
}
}
#[derive(Clone)]
pub struct DefaultBehavior<Sub>
where
Sub: event::Subscriber + Clone,
{
sender: accept::Sender<Sub>,
}
impl<Sub> DefaultBehavior<Sub>
where
Sub: event::Subscriber + Clone,
{
#[inline]
pub fn new(sender: &accept::Sender<Sub>) -> Self {
Self {
sender: sender.clone(),
}
}
}
impl<Sub> PollBehavior<Sub> for DefaultBehavior<Sub>
where
Sub: event::Subscriber + Clone,
{
fn poll<Pub>(
&self,
state: &mut WorkerState,
cx: &mut task::Context,
context: &mut Context<Sub>,
stream: &mut Option<(LazyBoundStream, SocketAddress)>,
subscriber_ctx: &mut Option<Sub::ConnectionContext>,
queue_time: Timestamp,
now: Timestamp,
publisher: &Pub,
) -> Poll<Result<ControlFlow<()>, WorkerError>>
where
Pub: EndpointPublisher,
{
let sojourn_time = now.saturating_duration_since(queue_time);
loop {
let (recv_buffer, blocked_count) = match state {
WorkerState::Init => (&mut context.recv_buffer, 0),
WorkerState::Buffering {
buffer,
blocked_count,
} => (buffer, *blocked_count),
WorkerState::Erroring { offset, buffer, .. } => {
let (stream, _remote_address) = stream.as_mut().unwrap();
let len = ready!(Pin::new(stream).poll_write(cx, &buffer[*offset..])).map_err(
|e| WorkerError {
error: e,
source: event::builder::AcceptorTcpIoErrorSource::Send,
},
)?;
*offset += len;
if *offset < buffer.len() {
continue;
}
let WorkerState::Erroring { error, .. } =
core::mem::replace(state, WorkerState::Init)
else {
unreachable!()
};
return Err(error).into();
}
WorkerState::Sending { .. } => unreachable!(),
};
let res = {
let (stream, remote_address) = stream.as_mut().unwrap();
WorkerState::poll_initial_packet(
cx,
stream,
remote_address,
recv_buffer,
sojourn_time,
publisher,
)
};
let Poll::Ready(res) = res else {
if blocked_count == 0 {
let buffer = recv_buffer.take();
*state = WorkerState::Buffering {
buffer,
blocked_count,
};
}
if let WorkerState::Buffering { blocked_count, .. } = state {
*blocked_count += 1;
}
return Poll::Pending;
};
let initial_packet = res?;
let subscriber_ctx = subscriber_ctx.take().unwrap();
let (socket, remote_address) = stream.take().unwrap();
let recv_buffer = recv::buffer::Local::new(recv_buffer.take(), None);
let recv_buffer = Either::A(recv_buffer);
let mut secret_control = vec![];
let (crypto, parameters) = match endpoint::derive_stream_credentials(
&initial_packet,
&context.secrets,
&TransportFeatures::TCP,
&mut secret_control,
) {
Ok(result) => result,
Err(error) => {
if !secret_control.is_empty() {
*stream = Some((socket, remote_address));
*state = WorkerState::Erroring {
offset: 0,
buffer: secret_control,
error: WorkerError {
error,
source: event::builder::AcceptorTcpIoErrorSource::Local,
},
};
continue;
} else {
let _ = socket.set_linger(Some(Duration::ZERO));
drop(socket);
}
return Err(WorkerError {
error,
source: event::builder::AcceptorTcpIoErrorSource::Local,
})
.into();
}
};
let peer = env::tcp::Reregistered {
socket,
peer_addr: remote_address,
local_port: context.local_port,
recv_buffer,
};
let stream_builder = match endpoint::accept_stream(
now,
&context.env,
peer,
&initial_packet,
&context.secrets,
subscriber_ctx,
None,
crypto,
parameters,
secret_control,
) {
Ok(stream) => stream,
Err(error) => {
return Err(WorkerError {
error: error.error,
source: event::builder::AcceptorTcpIoErrorSource::Local,
})
.into();
}
};
{
let remote_address: SocketAddress = stream_builder.shared.remote_addr();
let remote_address = &remote_address;
let creds = stream_builder.shared.credentials();
let credential_id = &*creds.id;
let stream_id = creds.key_id.as_u64();
publisher.on_acceptor_tcp_stream_enqueued(
event::builder::AcceptorTcpStreamEnqueued {
remote_address,
credential_id,
stream_id,
sojourn_time,
blocked_count,
},
);
}
let res = match context.accept_flavor {
accept::Flavor::Fifo => self.sender.send_back(stream_builder),
accept::Flavor::Lifo => self.sender.send_front(stream_builder),
};
return Poll::Ready(Ok(match res {
Ok(prev) => {
if let Some(stream) = prev {
stream.prune(
event::builder::AcceptorStreamPruneReason::AcceptQueueCapacityExceeded,
);
}
ControlFlow::Continue(())
}
Err(_err) => {
debug!("application accept queue dropped; shutting down");
ControlFlow::Break(())
}
}));
}
}
}
impl WorkerState {
#[inline]
fn poll_initial_packet<Pub>(
cx: &mut task::Context,
stream: &mut LazyBoundStream,
remote_address: &SocketAddress,
recv_buffer: &mut msg::recv::Message,
sojourn_time: Duration,
publisher: &Pub,
) -> Poll<Result<server::InitialPacket, WorkerError>>
where
Pub: EndpointPublisher,
{
loop {
if recv_buffer.payload_len() > 10_000 {
publisher.on_acceptor_tcp_packet_dropped(
event::builder::AcceptorTcpPacketDropped {
remote_address,
reason: DecoderError::UnexpectedBytes(recv_buffer.payload_len())
.into_event(),
sojourn_time,
},
);
let _ = stream.set_linger(Some(Duration::ZERO));
return Err(WorkerError {
source: event::builder::AcceptorTcpIoErrorSource::Remote,
error: io::Error::from(io::ErrorKind::FileTooLarge),
})
.into();
}
let res =
ready!(stream.poll_recv_buffer(cx, recv_buffer)).map_err(|error| WorkerError {
source: event::builder::AcceptorTcpIoErrorSource::Recv,
error,
})?;
match server::InitialPacket::peek(recv_buffer, 16) {
Ok(packet) => {
publisher.on_acceptor_tcp_packet_received(
event::builder::AcceptorTcpPacketReceived {
remote_address,
credential_id: &*packet.credentials.id,
stream_id: packet.stream_id.into_varint().as_u64(),
payload_len: packet.payload_len,
is_fin: packet.is_fin,
is_fin_known: packet.is_fin_known,
sojourn_time,
},
);
return Ok(packet).into();
}
Err(err) => {
if matches!(err, DecoderError::UnexpectedEof(_)) && res > 0 {
continue;
}
publisher.on_acceptor_tcp_packet_dropped(
event::builder::AcceptorTcpPacketDropped {
remote_address,
reason: err.into_event(),
sojourn_time,
},
);
let _ = stream.set_linger(Some(Duration::ZERO));
return Err(WorkerError {
source: event::builder::AcceptorTcpIoErrorSource::Remote,
error: io::Error::from(io::ErrorKind::InvalidData),
})
.into();
}
}
}
}
}
#[derive(Clone)]
pub struct SocketEventData {
blocked_count: usize,
payload_len: usize,
credential_id: Vec<u8>,
stream_id: u64,
sojourn_time: Duration,
}
#[derive(Clone)]
pub struct SocketBehavior {
sender: uds::sender::Sender,
}
impl SocketBehavior {
#[inline]
pub fn new(dest_path: &Path) -> Result<Self, std::io::Error> {
let sender = uds::sender::Sender::new(dest_path)?;
Ok(Self { sender })
}
fn poll_send<Pub>(
future: &mut SendMsg,
cx: &mut task::Context,
event_data: &SocketEventData,
publisher: &Pub,
) -> Poll<Result<ControlFlow<()>, WorkerError>>
where
Pub: EndpointPublisher,
{
match Pin::new(future).as_mut().poll(cx) {
Poll::Ready(res) => match res {
Ok(_) => {
publisher.on_acceptor_tcp_socket_sent(event::builder::AcceptorTcpSocketSent {
credential_id: &event_data.credential_id,
stream_id: event_data.stream_id,
payload_len: event_data.payload_len,
blocked_count: event_data.blocked_count,
sojourn_time: event_data.sojourn_time,
});
Poll::Ready(Ok(ControlFlow::Continue(())))
}
Err(err) => {
debug!("Error sending message to socket {:?}", err);
Err(WorkerError {
source: event::builder::AcceptorTcpIoErrorSource::UnixSend,
error: err,
})
.into()
}
},
Poll::Pending => Poll::Pending,
}
}
fn decrypt(keys: Bidirectional, recv_buffer: &mut [u8]) -> Result<(), WorkerError> {
let tag_len = keys.application.opener.tag_len();
let decoder = decoder::DecoderBufferMut::new(recv_buffer);
let (packet, _remaining) =
decoder
.decode_parameterized(tag_len)
.map_err(|error| WorkerError {
source: event::builder::AcceptorTcpIoErrorSource::Remote,
error: io::Error::new(
io::ErrorKind::InvalidData,
format!("Failed to decode stream packet: {:?}", error),
),
})?;
let packet::Packet::Stream(stream_packet) = packet else {
return Err(WorkerError {
source: event::builder::AcceptorTcpIoErrorSource::Remote,
error: io::Error::new(io::ErrorKind::InvalidData, "Expected stream packet"),
});
};
let mut payload_out = vec![0u8; stream_packet.payload().len()];
let payload_out = crypto::UninitSlice::new(&mut payload_out);
keys.application
.opener
.decrypt(
stream_packet.tag().key_phase(),
*stream_packet.packet_number(),
stream_packet.header(),
stream_packet.payload(),
stream_packet.auth_tag(),
payload_out,
)
.map_err(|error| WorkerError {
source: event::builder::AcceptorTcpIoErrorSource::Remote,
error: io::Error::new(
io::ErrorKind::InvalidData,
format!("Failed to decrypt stream packet: {:?}", error),
),
})?;
Ok(())
}
}
impl<Sub> PollBehavior<Sub> for SocketBehavior
where
Sub: event::Subscriber + Clone,
{
fn poll<Pub>(
&self,
state: &mut WorkerState,
cx: &mut task::Context,
context: &mut Context<Sub>,
stream: &mut Option<(LazyBoundStream, SocketAddress)>,
_subscriber_ctx: &mut Option<Sub::ConnectionContext>,
queue_time: Timestamp,
now: Timestamp,
publisher: &Pub,
) -> Poll<Result<ControlFlow<()>, WorkerError>>
where
Pub: EndpointPublisher,
{
let sojourn_time = now.saturating_duration_since(queue_time);
loop {
let (recv_buffer, blocked_count) = match state {
WorkerState::Init => (&mut context.recv_buffer, 0),
WorkerState::Buffering {
buffer,
blocked_count,
} => (buffer, *blocked_count),
WorkerState::Erroring { offset, buffer, .. } => {
let (stream, _remote_address) = stream.as_mut().unwrap();
let len = ready!(Pin::new(stream).poll_write(cx, &buffer[*offset..])).map_err(
|error| WorkerError {
source: event::builder::AcceptorTcpIoErrorSource::Send,
error,
},
)?;
*offset += len;
if *offset < buffer.len() {
continue;
}
let WorkerState::Erroring { error, .. } =
core::mem::replace(state, WorkerState::Init)
else {
unreachable!()
};
return Err(error).into();
}
WorkerState::Sending { future, event_data } => {
match Self::poll_send(future, cx, event_data, publisher) {
Poll::Ready(result) => return Poll::Ready(result),
Poll::Pending => {
event_data.blocked_count += 1;
return Poll::Pending;
}
}
}
};
let res = {
let (stream, remote_address) = stream.as_mut().unwrap();
WorkerState::poll_initial_packet(
cx,
stream,
remote_address,
recv_buffer,
sojourn_time,
publisher,
)
};
let Poll::Ready(res) = res else {
if blocked_count == 0 {
let buffer = recv_buffer.take();
*state = WorkerState::Buffering {
buffer,
blocked_count,
};
}
if let WorkerState::Buffering { blocked_count, .. } = state {
*blocked_count += 1;
}
return Poll::Pending;
};
let initial_packet = res?;
let (socket, remote_address) = stream.take().unwrap();
let recv_buffer = recv_buffer.make_contiguous();
let mut secret_control = vec![];
let credentials = &initial_packet.credentials;
let map = &context.secrets;
let Some((export_secret, ciphersuite, keys, application_params)) = map
.secret_for_credentials(
credentials,
initial_packet.source_queue_id,
&TransportFeatures::TCP,
&mut secret_control,
)
else {
let error = io::Error::new(
io::ErrorKind::NotFound,
format!("missing credentials for client: {credentials:?}"),
);
if !secret_control.is_empty() {
*stream = Some((socket, remote_address));
*state = WorkerState::Erroring {
offset: 0,
buffer: secret_control,
error: WorkerError {
source: event::builder::AcceptorTcpIoErrorSource::Local,
error,
},
};
continue;
} else {
let _ = socket.set_linger(Some(Duration::ZERO));
drop(socket);
}
return Err(WorkerError {
source: event::builder::AcceptorTcpIoErrorSource::Local,
error,
})
.into();
};
if let Err(err) = Self::decrypt(keys, recv_buffer) {
let _ = socket.set_linger(Some(Duration::ZERO));
drop(socket);
return Err(err).into();
};
#[cfg(target_os = "linux")]
let clock = ClockId::CLOCK_MONOTONIC_RAW;
#[cfg(not(target_os = "linux"))]
let clock = ClockId::CLOCK_MONOTONIC;
let now = clock.now().map_err(|errno| WorkerError {
source: event::builder::AcceptorTcpIoErrorSource::System,
error: io::Error::from(errno),
})?;
let encode_time = now.num_microseconds() as u64;
let mut estimator = EncoderLenEstimator::new(usize::MAX);
let size = packet::uds::encoder::encode(
&mut estimator,
&ciphersuite,
&export_secret,
&application_params,
encode_time,
recv_buffer,
);
let mut buffer = vec![0u8; size];
let mut encoder = s2n_codec::EncoderBuffer::new(&mut buffer);
packet::uds::encoder::encode(
&mut encoder,
&ciphersuite,
&export_secret,
&application_params,
encode_time,
recv_buffer,
);
let tcp_stream = socket.into_std().map_err(|error| WorkerError {
source: event::builder::AcceptorTcpIoErrorSource::System,
error,
})?;
let mut future = SendMsg::new(self.sender.clone(), buffer, OwnedFd::from(tcp_stream));
let mut event_data = SocketEventData {
credential_id: credentials.id.to_vec(),
stream_id: credentials.key_id.as_u64(),
payload_len: size,
blocked_count: 0,
sojourn_time,
};
match Self::poll_send(&mut future, cx, &event_data, publisher) {
Poll::Ready(result) => return Poll::Ready(result),
Poll::Pending => {
event_data.blocked_count += 1;
*state = WorkerState::Sending { future, event_data };
return Poll::Pending;
}
}
}
}
}