use crate::{
allocator::Allocator,
clock,
credentials::Credentials,
crypto::{self, UninitSlice},
event,
packet::{control, stream},
stream::{
recv::{
ack,
error::{self, Error},
packet,
},
shared::AcceptState,
TransportFeatures, DEFAULT_IDLE_TIMEOUT,
},
};
use core::{task::Poll, time::Duration};
use s2n_codec::{EncoderBuffer, EncoderValue};
use s2n_quic_core::{
buffer::{self, reader::storage::Infallible as _},
dc::ApplicationParams,
endpoint::Location,
ensure,
frame::{self, ack::EcnCounts},
inet::ExplicitCongestionNotification,
packet::number::PacketNumberSpace,
ready,
stream::state::Receiver,
time::{
timer::{self, Provider as _},
Clock, Timer, Timestamp,
},
varint::VarInt,
};
#[derive(Clone, Copy, Debug)]
struct ErrorState {
error: Error,
source: Location,
}
#[derive(Debug)]
pub struct State {
ecn_counts: EcnCounts,
control_packet_number: u64,
stream_ack: ack::Space,
recovery_ack: ack::Space,
pub(crate) state: Receiver,
idle_timer: Timer,
idle_timeout: Duration,
tick_timer: Timer,
_should_transmit: bool,
is_reliable: bool,
max_data: VarInt,
max_data_window: VarInt,
error: Option<ErrorState>,
fin_ack_packet_number: Option<VarInt>,
features: TransportFeatures,
}
impl State {
#[inline]
pub fn new<C>(
stream_id: stream::Id,
params: &ApplicationParams,
features: TransportFeatures,
clock: &C,
) -> Self
where
C: Clock + ?Sized,
{
let (initial_max_data, max_data_window) = if features.is_flow_controlled() {
(VarInt::MAX, VarInt::MAX)
} else {
let initial_max_data = params.remote_max_data;
let data_window = params.local_recv_max_data;
(initial_max_data, data_window)
};
let now = clock.get_time();
let idle_timeout = params.max_idle_timeout().unwrap_or(DEFAULT_IDLE_TIMEOUT);
let mut idle_timer = Timer::default();
idle_timer.set(now + idle_timeout);
let tick_timer = idle_timer.clone();
Self {
is_reliable: stream_id.is_reliable,
ecn_counts: Default::default(),
control_packet_number: Default::default(),
stream_ack: Default::default(),
recovery_ack: Default::default(),
state: Default::default(),
idle_timer,
idle_timeout,
tick_timer,
_should_transmit: false,
max_data: initial_max_data,
max_data_window,
error: None,
fin_ack_packet_number: None,
features,
}
}
#[inline]
pub fn state(&self) -> &Receiver {
&self.state
}
#[inline]
pub fn timer(&self) -> Option<Timestamp> {
self.next_expiration()
}
#[inline]
pub fn is_open(&self) -> bool {
!self.state.is_terminal()
}
#[inline]
pub fn is_finished(&self) -> bool {
ensure!(self.state.is_terminal(), false);
ensure!(self.timer().is_none(), false);
true
}
#[inline]
pub fn stop_sending<Pub>(&mut self, error: s2n_quic_core::application::Error, publisher: &Pub)
where
Pub: event::ConnectionPublisher,
{
ensure!(matches!(self.state, Receiver::Recv | Receiver::SizeKnown));
self.on_error(
error::Kind::ApplicationError { error },
Location::Local,
publisher,
);
}
#[inline]
pub fn on_read_buffer<B, C, Clk>(
&mut self,
out_buf: &mut B,
chunk: &mut C,
accept_state: AcceptState,
_clock: &Clk,
) where
B: buffer::Duplex<Error = core::convert::Infallible>,
C: buffer::writer::Storage,
Clk: Clock + ?Sized,
{
if chunk.has_remaining_capacity() && !out_buf.buffer_is_empty() {
out_buf.infallible_copy_into(chunk);
}
if matches!(accept_state, AcceptState::Accepted) {
let new_max_data = out_buf
.current_offset()
.saturating_add(self.max_data_window);
if new_max_data > self.max_data {
self.max_data = new_max_data;
self.needs_transmission("new_max_data");
}
}
if out_buf.final_offset().is_some() {
let _ = self.state.on_receive_fin();
}
if out_buf.has_buffered_fin() && self.state.on_receive_all_data().is_ok() {
self.needs_transmission("receive_all_data");
}
if out_buf.is_consumed() && self.state.on_app_read_all_data().is_ok() {
self.needs_transmission("app_read_all_data");
}
}
#[inline]
pub fn precheck_stream_packet<Pub>(
&mut self,
credentials: &Credentials,
packet: &stream::decoder::Packet,
publisher: &Pub,
) -> Result<(), Error>
where
Pub: event::ConnectionPublisher,
{
match self.precheck_stream_packet_impl(credentials, packet) {
Ok(()) => Ok(()),
Err(err) => {
if err.is_fatal(&self.features) {
self.on_error(err, Location::Local, publisher);
} else {
tracing::debug!(non_fatal_error = %err, ?packet);
}
Err(err)
}
}
}
#[inline]
fn precheck_stream_packet_impl(
&mut self,
credentials: &Credentials,
packet: &stream::decoder::Packet,
) -> Result<(), Error> {
ensure!(
packet.credentials() == credentials,
Err(error::Kind::CredentialMismatch {
expected: *credentials,
actual: *packet.credentials(),
}
.err())
);
if self.features.is_stream() {
let expected_pn = self
.stream_ack
.packets
.max_value()
.map_or(0, |v| v.as_u64() + 1);
let actual_pn = packet.packet_number().as_u64();
ensure!(
expected_pn == actual_pn,
Err(error::Kind::OutOfOrder {
expected: expected_pn,
actual: actual_pn,
}
.err())
);
}
if self.features.is_reliable() {
ensure!(
!packet.is_retransmission(),
Err(error::Kind::UnexpectedRetransmission.err())
);
}
Ok(())
}
#[inline]
pub fn on_stream_packet<D, C, B, Clk, Pub>(
&mut self,
opener: &D,
control: &C,
credentials: &Credentials,
packet: &mut stream::decoder::Packet,
ecn: ExplicitCongestionNotification,
accept_state: AcceptState,
clock: &Clk,
out_buf: &mut B,
publisher: &Pub,
) -> Result<(), Error>
where
D: crypto::open::Application,
C: crypto::open::control::Stream,
Clk: Clock + ?Sized,
B: buffer::Duplex<Error = core::convert::Infallible>,
Pub: event::ConnectionPublisher,
{
publisher.on_stream_packet_received(event::builder::StreamPacketReceived {
packet_len: packet.total_len(),
packet_number: packet.packet_number().as_u64(),
stream_offset: packet.stream_offset().as_u64(),
payload_len: packet.payload().len(),
is_fin: packet.is_fin(),
is_retransmission: packet.is_retransmission(),
});
match self.on_stream_packet_impl(
opener,
control,
credentials,
packet,
ecn,
accept_state,
clock,
out_buf,
publisher,
) {
Ok(()) => Ok(()),
Err(err) => {
if err.is_fatal(&self.features) {
self.on_error(err, Location::Local, publisher);
} else {
tracing::debug!(non_fatal_error = %err, ?packet);
}
Err(err)
}
}
}
#[inline]
fn on_stream_packet_impl<D, C, B, Clk, Pub>(
&mut self,
opener: &D,
control: &C,
credentials: &Credentials,
packet: &mut stream::decoder::Packet,
ecn: ExplicitCongestionNotification,
accept_state: AcceptState,
clock: &Clk,
out_buf: &mut B,
publisher: &Pub,
) -> Result<(), Error>
where
D: crypto::open::Application,
C: crypto::open::control::Stream,
Clk: Clock + ?Sized,
B: buffer::Duplex<Error = core::convert::Infallible>,
Pub: event::ConnectionPublisher,
{
use buffer::reader::Storage as _;
self.precheck_stream_packet_impl(credentials, packet)?;
let is_max_data_ok = self.ensure_max_data(packet);
let mut packet = packet::Packet {
packet: &mut *packet,
payload_cursor: 0,
is_decrypted_in_place: false,
ecn,
clock,
opener,
control,
receiver: self,
copied_len: 0,
publisher,
};
if !is_max_data_ok {
let _ = packet.read_chunk(usize::MAX)?;
tracing::error!(
message = "max data exceeded",
allowed = packet.receiver.max_data.as_u64(),
requested = packet
.packet
.stream_offset()
.as_u64()
.saturating_add(packet.packet.payload().len() as u64),
);
let error = error::Kind::MaxDataExceeded.err();
self.on_error(error, Location::Local, publisher);
return Err(error);
}
let initial = out_buf.buffered_len();
out_buf.read_from(&mut packet)?;
let new = out_buf.buffered_len();
if !packet.packet.payload().is_empty() {
if packet.copied_len == 0 {
publisher.on_stream_packet_spuriously_retransmitted(
event::builder::StreamPacketSpuriouslyRetransmitted {
packet_len: packet.packet.total_len(),
packet_number: packet.packet.packet_number().as_u64(),
stream_offset: packet.packet.stream_offset().as_u64(),
payload_len: packet.packet.payload().len(),
is_fin: packet.packet.is_fin(),
is_retransmission: packet.packet.is_retransmission(),
},
);
}
publisher.on_stream_decrypt_packet(event::builder::StreamDecryptPacket {
decrypted_in_place: packet.is_decrypted_in_place,
forced_copy: if packet.is_decrypted_in_place {
packet.packet.payload().len()
} else {
new.saturating_sub(initial)
},
required_application_buffer: packet.packet.payload().len(),
});
}
let mut chunk = buffer::writer::storage::Empty;
self.on_read_buffer(out_buf, &mut chunk, accept_state, clock);
Ok(())
}
#[inline]
pub(super) fn on_stream_packet_in_place<D, C, Clk, Pub>(
&mut self,
crypto: &D,
control: &C,
packet: &mut stream::decoder::Packet,
ecn: ExplicitCongestionNotification,
clock: &Clk,
publisher: &Pub,
) -> Result<(), Error>
where
D: crypto::open::Application,
C: crypto::open::control::Stream,
Clk: Clock + ?Sized,
Pub: event::ConnectionPublisher,
{
let res = packet.decrypt_in_place(crypto, control);
res?;
self.on_cleartext_stream_packet(packet, ecn, clock, publisher)
}
#[inline]
pub(super) fn on_stream_packet_copy<D, C, Clk, Pub>(
&mut self,
crypto: &D,
control: &C,
packet: &mut stream::decoder::Packet,
ecn: ExplicitCongestionNotification,
payload_out: &mut UninitSlice,
clock: &Clk,
publisher: &Pub,
) -> Result<(), Error>
where
D: crypto::open::Application,
C: crypto::open::control::Stream,
Clk: Clock + ?Sized,
Pub: event::ConnectionPublisher,
{
let res = packet.decrypt(crypto, control, payload_out);
res?;
self.on_cleartext_stream_packet(packet, ecn, clock, publisher)
}
#[inline]
fn ensure_max_data(&self, packet: &stream::decoder::Packet) -> bool {
ensure!(!self.features.is_flow_controlled(), true);
self.max_data
.as_u64()
.checked_sub(packet.payload().len() as u64)
.and_then(|v| v.checked_sub(packet.stream_offset().as_u64()))
.is_some()
}
#[inline]
fn on_cleartext_stream_packet<Clk, Pub>(
&mut self,
packet: &mut stream::decoder::Packet,
ecn: ExplicitCongestionNotification,
clock: &Clk,
publisher: &Pub,
) -> Result<(), Error>
where
Clk: Clock + ?Sized,
Pub: event::ConnectionPublisher,
{
tracing::trace!(
stream_id = %packet.stream_id(),
stream_offset = packet.stream_offset().as_u64(),
payload_len = packet.payload().len(),
final_offset = ?packet.final_offset().map(|v| v.as_u64()),
);
let space = match packet.tag().packet_space() {
stream::PacketSpace::Stream => &mut self.stream_ack,
stream::PacketSpace::Recovery => &mut self.recovery_ack,
};
ensure!(
space.filter.on_packet(packet).is_ok(),
Err(error::Kind::Duplicate.err())
);
let packet_number = PacketNumberSpace::Initial.new_packet_number(packet.packet_number());
if let Err(err) = space.packets.insert_packet_number(packet_number) {
tracing::debug!("could not record packet number {packet_number} with error {err:?}");
}
self.needs_transmission("new_packet");
if matches!(self.state, Receiver::Recv | Receiver::SizeKnown)
|| packet.stream_offset() == VarInt::ZERO
{
self.update_idle_timer(clock);
}
for frame in packet.control_frames_mut() {
let Ok(frame) = frame else {
return Err(error::Kind::Decode.err());
};
match frame {
frame::Frame::ConnectionClose(close) => {
let error = if close.frame_type.is_some() {
error::Kind::TransportError {
code: close.error_code,
}
} else {
error::Kind::ApplicationError {
error: close.error_code.into(),
}
}
.err();
self.on_error(error, Location::Remote, publisher);
return Err(error);
}
_ => {
}
}
}
self.ecn_counts.increment(ecn);
if !self.is_reliable {
}
self.on_next_expected_control_packet(packet.next_expected_control_packet());
Ok(())
}
#[inline]
pub fn should_transmit(&self) -> bool {
self._should_transmit
}
#[inline]
pub fn on_transport_close<Pub>(&mut self, publisher: &Pub)
where
Pub: event::ConnectionPublisher,
{
ensure!(self.features.is_stream());
ensure!(matches!(self.state, Receiver::Recv | Receiver::SizeKnown));
self.on_error(error::Kind::TruncatedTransport, Location::Local, publisher);
}
#[inline]
fn needs_transmission(&mut self, reason: &str) {
if self.error.is_none() {
if self.features.is_reliable() && self.features.is_flow_controlled() {
tracing::trace!(skipping_transmission = reason);
return;
}
}
if !self._should_transmit {
tracing::trace!(needs_transmission = reason);
}
self._should_transmit = true;
}
#[inline]
fn on_next_expected_control_packet(&mut self, next_expected_control_packet: VarInt) {
if let Some(largest_delivered_control_packet) =
next_expected_control_packet.checked_sub(VarInt::from_u8(1))
{
self.stream_ack
.on_largest_delivered_packet(largest_delivered_control_packet);
self.recovery_ack
.on_largest_delivered_packet(largest_delivered_control_packet);
if let Some(fin_ack_packet_number) = self.fin_ack_packet_number {
if largest_delivered_control_packet >= fin_ack_packet_number {
self.silent_shutdown();
}
}
}
}
#[inline]
fn update_idle_timer<Clk: Clock + ?Sized>(&mut self, clock: &Clk) {
let target = clock.get_time() + self.idle_timeout;
self.idle_timer.set(target);
if !self.tick_timer.is_armed() {
self.tick_timer.set(target);
}
}
#[inline]
fn mtu(&self) -> u16 {
1200
}
#[inline]
fn ecn(&self) -> ExplicitCongestionNotification {
ExplicitCongestionNotification::Ect0
}
#[inline]
#[track_caller]
pub fn on_error<E, Pub>(&mut self, error: E, source: Location, publisher: &Pub)
where
Error: From<E>,
Pub: event::ConnectionPublisher,
{
let error = Error::from(error);
debug_assert!(error.is_fatal(&self.features));
let _ = self.state.on_reset();
self.stream_ack.clear();
self.recovery_ack.clear();
ensure!(self.error.is_none());
self.error = Some(ErrorState { error, source });
publisher
.on_stream_receiver_errored(event::builder::StreamReceiverErrored { error, source });
if matches!(source, Location::Local) {
self.needs_transmission("on_error");
} else {
let _ = self.state.on_app_read_reset();
self.silent_shutdown();
}
}
#[inline]
pub fn check_error(&self) -> Result<(), Error> {
ensure!(
!matches!(self.state, Receiver::DataRead | Receiver::DataRecvd),
Ok(())
);
if let Some(err) = self.error {
Err(err.error)
} else {
Ok(())
}
}
#[inline]
pub fn on_timeout<Clk, Ld, Pub>(&mut self, clock: &Clk, load_last_activity: Ld, publisher: &Pub)
where
Clk: Clock + ?Sized,
Ld: FnOnce() -> Timestamp,
Pub: event::ConnectionPublisher,
{
let now = clock.get_time();
if self.poll_idle_timer(clock, load_last_activity).is_ready() {
self.silent_shutdown();
ensure!(matches!(self.state, Receiver::Recv | Receiver::SizeKnown));
let mut did_transition = false;
did_transition |= self.state.on_reset().is_ok();
did_transition |= self.state.on_app_read_reset().is_ok();
if did_transition {
self.on_error(error::Kind::IdleTimeout, Location::Local, publisher);
self._should_transmit = false;
}
return;
}
if self.tick_timer.poll_expiration(now).is_ready() {
self.tick_timer = self.idle_timer.clone();
}
}
#[inline]
fn poll_idle_timer<Clk, Ld>(&mut self, clock: &Clk, load_last_activity: Ld) -> Poll<()>
where
Clk: Clock + ?Sized,
Ld: FnOnce() -> Timestamp,
{
let now = clock.get_time();
ready!(self.idle_timer.poll_expiration(now));
let last_peer_activity = load_last_activity();
self.update_idle_timer(&last_peer_activity);
ready!(self.idle_timer.poll_expiration(now));
Poll::Ready(())
}
#[inline]
fn silent_shutdown(&mut self) {
self._should_transmit = false;
self.idle_timer.cancel();
self.tick_timer.cancel();
self.stream_ack.clear();
self.recovery_ack.clear();
tracing::trace!("silent_shutdown");
}
#[inline]
pub fn on_transmit<K, A, Clk, Pub>(
&mut self,
key: &K,
credentials: &Credentials,
stream_id: stream::Id,
source_queue_id: Option<VarInt>,
output: &mut A,
clock: &Clk,
publisher: &Pub,
) where
K: crypto::seal::control::Stream,
A: Allocator,
Clk: Clock + ?Sized,
Pub: event::ConnectionPublisher,
{
(if self.error.is_none() {
Self::on_transmit_ack
} else {
Self::on_transmit_error
})(
self,
key,
credentials,
stream_id,
source_queue_id,
output,
&clock::Cached::new(clock),
publisher,
)
}
#[inline]
fn on_transmit_ack<K, A, Clk, Pub>(
&mut self,
key: &K,
credentials: &Credentials,
stream_id: stream::Id,
source_queue_id: Option<VarInt>,
output: &mut A,
_clock: &Clk,
publisher: &Pub,
) where
K: crypto::seal::control::Stream,
A: Allocator,
Clk: Clock + ?Sized,
Pub: event::ConnectionPublisher,
{
ensure!(self.should_transmit());
let mtu = self.mtu();
output.set_ecn(self.ecn());
let packet_number = self.next_pn();
ensure!(let Some(segment) = output.alloc());
let buffer = output.get_mut(&segment);
buffer.resize(mtu as _, 0);
let encoder = EncoderBuffer::new(buffer);
let ack_delay = VarInt::ZERO;
let max_data = frame::MaxData {
maximum_data: self.max_data,
};
let max_data_encoding_size: VarInt = max_data.encoding_size().try_into().unwrap();
let (recovery_ack, max_data_encoding_size) =
self.recovery_ack
.encoding(max_data_encoding_size, ack_delay, None, mtu);
let (stream_ack, max_data_encoding_size) = self.stream_ack.encoding(
max_data_encoding_size,
ack_delay,
Some(self.ecn_counts),
mtu,
);
let encoding_size = max_data_encoding_size;
tracing::trace!(?stream_ack, ?recovery_ack, ?max_data);
let frame = ((max_data, stream_ack), recovery_ack);
let result = control::encoder::encode(
encoder,
source_queue_id,
Some(stream_id),
packet_number,
VarInt::ZERO,
&mut &[][..],
encoding_size,
&frame,
key,
credentials,
);
match result {
0 => {
output.free(segment);
return;
}
packet_len => {
buffer.truncate(packet_len);
let intervals = self.stream_ack.packets.interval_len()
+ self.recovery_ack.packets.interval_len();
let mut duplicate_threshold = 20;
let mut should_duplicate = false;
should_duplicate |= intervals > duplicate_threshold;
should_duplicate |= !self.recovery_ack.packets.is_empty();
let duplicate = if should_duplicate {
Some(buffer.clone())
} else {
None
};
output.push(segment);
if let Some(buffer) = duplicate {
loop {
let Some(segment) = output.alloc() else {
break;
};
let buf = output.get_mut(&segment);
if intervals > duplicate_threshold {
buf.extend_from_slice(&buffer);
output.push(segment);
duplicate_threshold *= 2;
continue;
} else {
*buf = buffer;
output.push(segment);
break;
}
}
}
}
}
ensure!(!output.is_empty());
publisher.on_stream_control_packet_transmitted(
event::builder::StreamControlPacketTransmitted {
packet_len: result,
control_data_len: encoding_size.as_u64() as usize,
packet_number: packet_number.as_u64(),
},
);
self.stream_ack.on_transmit(packet_number);
self.recovery_ack.on_transmit(packet_number);
self.on_packet_sent(packet_number);
}
#[inline]
fn on_transmit_error<K, A, Clk, Pub>(
&mut self,
control_key: &K,
credentials: &Credentials,
stream_id: stream::Id,
source_queue_id: Option<VarInt>,
output: &mut A,
_clock: &Clk,
publisher: &Pub,
) where
K: crypto::seal::control::Stream,
A: Allocator,
Clk: Clock + ?Sized,
Pub: event::ConnectionPublisher,
{
ensure!(self.should_transmit());
let Some(error) = self.error else {
return;
};
ensure!(matches!(error.source, Location::Local));
let mtu = self.mtu() as usize;
output.set_ecn(self.ecn());
let packet_number = self.next_pn();
ensure!(let Some(segment) = output.alloc());
let buffer = output.get_mut(&segment);
buffer.resize(mtu, 0);
let encoder = EncoderBuffer::new(buffer);
let frame = error
.error
.connection_close()
.unwrap_or_else(|| s2n_quic_core::transport::Error::NO_ERROR.into());
let encoding_size = frame.encoding_size().try_into().unwrap();
let result = control::encoder::encode(
encoder,
source_queue_id,
Some(stream_id),
packet_number,
VarInt::ZERO,
&mut &[][..],
encoding_size,
&frame,
control_key,
credentials,
);
match result {
0 => {
output.free(segment);
return;
}
packet_len => {
buffer.truncate(packet_len);
output.push(segment);
}
}
tracing::debug!(connection_close = ?frame);
publisher.on_stream_control_packet_transmitted(
event::builder::StreamControlPacketTransmitted {
packet_len: result,
control_data_len: encoding_size.as_u64() as usize,
packet_number: packet_number.as_u64(),
},
);
self.stream_ack.clear();
self.recovery_ack.clear();
self.on_packet_sent(packet_number);
}
#[inline]
fn next_pn(&mut self) -> VarInt {
VarInt::new(self.control_packet_number).expect("2^62 is a lot of packets")
}
#[inline]
fn on_packet_sent(&mut self, packet_number: VarInt) {
if !matches!(self.state, Receiver::Recv | Receiver::SizeKnown)
&& self.fin_ack_packet_number.is_none()
{
self.fin_ack_packet_number = Some(packet_number);
}
self.control_packet_number += 1;
self._should_transmit = false;
}
}
impl timer::Provider for State {
#[inline]
fn timers<Q: timer::Query>(&self, query: &mut Q) -> timer::Result {
self.idle_timer.timers(query)?;
self.tick_timer.timers(query)?;
Ok(())
}
}