use crate::{
allocator::Allocator,
clock, msg,
packet::{stream, Packet},
stream::{
recv,
server::handshake,
shared::{ArcShared, Half},
socket::{self, Socket},
TransportFeatures,
},
task::waker::worker::Waker as WorkerWaker,
};
use core::{
mem::ManuallyDrop,
ops,
task::{Context, Poll},
};
use s2n_codec::{DecoderBufferMut, DecoderError};
use s2n_quic_core::{buffer, dc, ensure, ready, stream::state, time::Clock};
use std::{
io,
sync::{
atomic::{AtomicU64, AtomicU8, Ordering},
Mutex, MutexGuard,
},
};
#[derive(Clone, Copy, Debug, Default)]
pub enum AckMode {
#[default]
Application,
Worker,
}
pub enum ApplicationState {
Open,
Closed { is_panicking: bool },
}
impl ApplicationState {
const IS_CLOSED_MASK: u8 = 1;
const IS_PANICKING_MASK: u8 = 1 << 1;
#[inline]
fn load(shared: &AtomicU8) -> Self {
let value = shared.load(Ordering::Acquire);
if value == 0 {
return Self::Open;
}
let is_panicking = value & Self::IS_PANICKING_MASK != 0;
Self::Closed { is_panicking }
}
#[inline]
fn close(shared: &AtomicU8, is_panicking: bool) {
let mut value = Self::IS_CLOSED_MASK;
if is_panicking {
value |= Self::IS_PANICKING_MASK;
}
shared.store(value, Ordering::Release);
}
}
#[derive(Debug)]
pub struct State {
inner: Mutex<Inner>,
application_epoch: AtomicU64,
application_state: AtomicU8,
pub worker_waker: WorkerWaker,
}
impl State {
#[inline]
pub fn new(
stream_id: stream::Id,
params: &dc::ApplicationParams,
handshake: Option<handshake::Receiver>,
features: TransportFeatures,
recv_buffer: Option<&mut msg::recv::Message>,
) -> Self {
let recv_buffer = match recv_buffer {
Some(prev) => prev.take(),
None => msg::recv::Message::new(9000u16),
};
let receiver = recv::state::State::new(stream_id, params, features);
let reassembler = Default::default();
let inner = Inner {
receiver,
reassembler,
handshake,
recv_buffer,
};
let inner = Mutex::new(inner);
Self {
inner,
application_epoch: AtomicU64::new(0),
application_state: AtomicU8::new(0),
worker_waker: Default::default(),
}
}
#[inline]
pub fn application_state(&self) -> ApplicationState {
ApplicationState::load(&self.application_state)
}
#[inline]
pub fn application_epoch(&self) -> u64 {
self.application_epoch.load(Ordering::Acquire)
}
#[inline]
pub fn application_guard<'a>(
&'a self,
ack_mode: AckMode,
send_buffer: &'a mut msg::send::Message,
shared: &'a ArcShared,
sockets: &'a dyn socket::Application,
) -> io::Result<AppGuard<'a>> {
self.application_epoch.fetch_add(1, Ordering::AcqRel);
let inner = self.inner.lock().map_err(|_| {
io::Error::new(io::ErrorKind::Other, "shared recv state has been poisoned")
})?;
let initial_state = inner.receiver.state().clone();
let inner = ManuallyDrop::new(inner);
Ok(AppGuard {
inner,
ack_mode,
send_buffer,
shared,
sockets,
initial_state,
})
}
#[inline]
pub fn shutdown(&self, is_panicking: bool) {
ApplicationState::close(&self.application_state, is_panicking);
self.worker_waker.wake();
}
#[inline]
pub fn worker_try_lock(&self) -> io::Result<Option<MutexGuard<Inner>>> {
match self.inner.try_lock() {
Ok(lock) => Ok(Some(lock)),
Err(std::sync::TryLockError::WouldBlock) => Ok(None),
Err(_) => Err(io::Error::new(
io::ErrorKind::Other,
"shared recv state has been poisoned",
)),
}
}
}
pub struct AppGuard<'a> {
inner: ManuallyDrop<MutexGuard<'a, Inner>>,
ack_mode: AckMode,
send_buffer: &'a mut msg::send::Message,
shared: &'a ArcShared,
sockets: &'a dyn socket::Application,
initial_state: state::Receiver,
}
impl<'a> AppGuard<'a> {
#[inline]
fn send_ack(&mut self) -> bool {
ensure!(
!self.sockets.read_application().features().is_reliable(),
false
);
match self.ack_mode {
AckMode::Application => {
self.inner
.fill_transmit_queue(self.shared, self.send_buffer);
ensure!(!self.send_buffer.is_empty(), false);
let did_send = self
.sockets
.read_application()
.try_send_buffer(self.send_buffer)
.is_ok();
let _ = self.send_buffer.drain();
!did_send
}
AckMode::Worker => {
self.inner.receiver.should_transmit()
}
}
}
}
impl<'a> ops::Deref for AppGuard<'a> {
type Target = Inner;
#[inline]
fn deref(&self) -> &Self::Target {
&self.inner
}
}
impl<'a> ops::DerefMut for AppGuard<'a> {
#[inline]
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.inner
}
}
impl<'a> Drop for AppGuard<'a> {
#[inline]
fn drop(&mut self) {
let wake_worker_for_ack = self.send_ack();
let current_state = self.inner.receiver.state().clone();
unsafe {
ManuallyDrop::drop(&mut self.inner);
}
if wake_worker_for_ack && !current_state.is_terminal() {
}
ensure!(self.initial_state != current_state);
if current_state.is_terminal() {
self.shared.receiver.shutdown(false);
}
}
}
#[derive(Debug)]
pub struct Inner {
pub receiver: recv::state::State,
pub reassembler: buffer::Reassembler,
pub recv_buffer: msg::recv::Message,
pub handshake: Option<handshake::Receiver>,
}
impl Inner {
#[inline]
pub fn fill_transmit_queue(
&mut self,
shared: &ArcShared,
send_buffer: &mut msg::send::Message,
) {
let source_control_port = shared.source_control_port();
shared.crypto.seal_with(|sealer| {
self.receiver
.on_transmit(sealer, source_control_port, send_buffer, &shared.clock)
});
ensure!(!send_buffer.is_empty());
send_buffer.set_remote_address(shared.read_remote_addr());
}
#[inline]
pub fn poll_fill_recv_buffer<S>(&mut self, cx: &mut Context, socket: &S) -> Poll<io::Result<()>>
where
S: ?Sized + Socket,
{
loop {
if let Some(chan) = self.handshake.as_mut() {
match chan.poll_recv(cx) {
Poll::Ready(Some(recv_buffer)) => {
debug_assert!(!recv_buffer.is_empty());
ensure!(!recv_buffer.is_empty(), continue);
self.recv_buffer = recv_buffer;
return Ok(()).into();
}
Poll::Ready(None) => {
self.handshake = None;
}
Poll::Pending => {
}
}
}
ready!(socket.poll_recv_buffer(cx, &mut self.recv_buffer))?;
return Ok(()).into();
}
}
#[inline]
pub fn process_recv_buffer(
&mut self,
out_buf: &mut impl buffer::writer::Storage,
shared: &ArcShared,
features: TransportFeatures,
) -> bool {
let clock = clock::Cached::new(&shared.clock);
let clock = &clock;
self.receiver
.on_read_buffer(&mut self.reassembler, out_buf, clock);
if !self.recv_buffer.is_empty() {
if features.is_stream() {
self.dispatch_buffer_stream(out_buf, shared, clock, features)
} else {
self.dispatch_buffer_datagram(out_buf, shared, clock, features)
}
self.receiver
.on_read_buffer(&mut self.reassembler, out_buf, clock);
}
if !features.is_reliable() {
self.receiver
.on_timeout(clock, || shared.last_peer_activity());
}
self.receiver.should_transmit()
}
#[inline]
fn dispatch_buffer_stream<C: Clock + ?Sized>(
&mut self,
out_buf: &mut impl buffer::writer::Storage,
shared: &ArcShared,
clock: &C,
features: TransportFeatures,
) {
let msg = &mut self.recv_buffer;
let remote_addr = msg.remote_address();
let ecn = msg.ecn();
let tag_len = shared.crypto.tag_len();
let mut any_valid_packets = false;
let mut did_complete_handshake = false;
let mut prev_packet_len = None;
let mut out_buf = buffer::duplex::Interposer::new(out_buf, &mut self.reassembler);
loop {
if let Some(packet_len) = prev_packet_len.take() {
msg.consume(packet_len);
}
let segment = msg.peek();
ensure!(!segment.is_empty(), break);
let initial_len = segment.len();
let decoder = DecoderBufferMut::new(segment);
let mut packet = match decoder.decode_parameterized(tag_len) {
Ok((packet, remaining)) => {
prev_packet_len = Some(initial_len - remaining.len());
packet
}
Err(decoder_error) => {
if let DecoderError::UnexpectedEof(len) = decoder_error {
if msg.make_contiguous().len() > initial_len {
continue;
}
let max_datagram_size =
shared.sender.path.load().max_datagram_size as usize;
if msg.payload_len() > max_datagram_size {
tracing::error!(
unconsumed = msg.payload_len(),
remaining_capacity = msg.remaining_capacity()
);
msg.clear();
self.receiver.on_error(recv::Error::Decode);
return;
}
tracing::trace!(
protocol_features = ?features,
unexpected_eof = len,
buffer_len = initial_len
);
break;
}
tracing::error!(
protocol_features = ?features,
fatal_error = %decoder_error,
payload_len = msg.payload_len()
);
msg.clear();
self.receiver.on_error(recv::Error::Decode);
return;
}
};
tracing::trace!(?packet);
match &mut packet {
Packet::Stream(packet) => {
debug_assert_eq!(Some(packet.total_len()), prev_packet_len);
if self.receiver.precheck_stream_packet(packet).is_err() {
if self.receiver.check_error().is_err() {
msg.clear();
return;
} else {
continue;
}
}
let _ = shared.crypto.open_with(|opener| {
self.receiver
.on_stream_packet(opener, packet, ecn, clock, &mut out_buf)?;
any_valid_packets = true;
did_complete_handshake |=
packet.next_expected_control_packet().as_u64() > 0;
<Result<_, recv::Error>>::Ok(())
});
if self.receiver.check_error().is_err() {
msg.clear();
return;
}
}
other => {
let kind = other.kind();
shared.crypto.map().handle_unexpected_packet(other);
msg.clear();
self.receiver
.on_error(recv::Error::UnexpectedPacket { packet: kind });
return;
}
}
}
if let Some(len) = prev_packet_len.take() {
msg.consume(len);
}
if any_valid_packets {
shared.on_valid_packet(&remote_addr, Half::Read, did_complete_handshake);
}
}
#[inline]
fn dispatch_buffer_datagram<C: Clock + ?Sized>(
&mut self,
out_buf: &mut impl buffer::writer::Storage,
shared: &ArcShared,
clock: &C,
features: TransportFeatures,
) {
let msg = &mut self.recv_buffer;
let remote_addr = msg.remote_address();
let ecn = msg.ecn();
let tag_len = shared.crypto.tag_len();
let mut any_valid_packets = false;
let mut did_complete_handshake = false;
let mut out_buf = buffer::duplex::Interposer::new(out_buf, &mut self.reassembler);
for segment in msg.segments() {
let segment_len = segment.len();
let mut decoder = DecoderBufferMut::new(segment);
'segment: while !decoder.is_empty() {
let packet = match decoder.decode_parameterized(tag_len) {
Ok((packet, remaining)) => {
decoder = remaining;
packet
}
Err(decoder_error) => {
tracing::warn!(
protocol_features = ?features,
%decoder_error,
segment_len
);
break 'segment;
}
};
match packet {
Packet::Stream(mut packet) => {
ensure!(
self.receiver.precheck_stream_packet(&packet).is_ok(),
continue
);
let _ = shared.crypto.open_with(|opener| {
self.receiver.on_stream_packet(
opener,
&mut packet,
ecn,
clock,
&mut out_buf,
)?;
any_valid_packets = true;
did_complete_handshake |=
packet.next_expected_control_packet().as_u64() > 0;
<Result<_, recv::Error>>::Ok(())
});
}
other => {
shared.crypto.map().handle_unexpected_packet(&other);
}
}
}
}
if any_valid_packets {
shared.on_valid_packet(&remote_addr, Half::Read, did_complete_handshake);
}
}
}