use std::sync::Arc;
use std::time::Duration;
use bytes::Bytes;
use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
use super::commands::{
MuxCommand,
StreamRegistration,
};
use super::{
FRAME_WORKERS,
GOAWAY_STATUS_OK,
WRITE_BATCH_CAP,
};
use crate::codec::SpdyCodec;
use crate::error::Error;
use crate::transport::WsFrameWriter;
macro_rules! write_with_timeout {
($writer:expr, $payload:expr, $timeout:expr, $bytes_counter:ident += $len:expr, $err_msg:literal, $timeout_msg:literal) => {
match tokio::time::timeout($timeout, $writer.write_binary(Bytes::from($payload))).await {
Ok(Ok(())) => {
$bytes_counter += $len;
}
Ok(Err(e)) => {
tracing::warn!($err_msg, e);
return (true, $bytes_counter);
}
Err(_) => {
tracing::warn!($timeout_msg, $timeout);
return (true, $bytes_counter);
}
}
};
}
#[derive(Clone, Copy)]
pub(super) struct WriterConfig {
pub initial_window_size: u32,
pub max_concurrent_streams: u32,
pub ping_interval: Duration,
pub write_timeout: Duration,
pub max_frame_size: u32,
}
pub(super) struct WriterParts<W: WsFrameWriter> {
pub writer: W,
pub cmd_rx: mpsc::Receiver<MuxCommand>,
pub control_rx: mpsc::Receiver<MuxCommand>,
pub window_rx: mpsc::Receiver<MuxCommand>,
pub close_reg_txs: Arc<[mpsc::Sender<StreamRegistration>; FRAME_WORKERS]>,
pub cancel: CancellationToken,
}
pub(super) fn encode_command(cmd: MuxCommand, codec: &SpdyCodec) -> Result<Bytes, Error> {
match cmd {
MuxCommand::OpenStreamPairAndWrite { .. } => {
unreachable!("OpenStreamPairAndWrite must be handled inline in run_writer")
}
MuxCommand::SendWsPong { .. } => {
unreachable!("SendWsPong must be handled inline in run_writer")
}
MuxCommand::SendData {
stream_id,
payload,
fin,
} => {
let frame_bytes = codec.encode_data(stream_id, &payload, fin);
Ok(Bytes::from(frame_bytes))
}
MuxCommand::SendRawFrame { frame } => Ok(frame),
MuxCommand::CloseStream { .. } => {
unreachable!("CloseStream must be handled inline in run_writer")
}
MuxCommand::EncodePing { id } => {
let frame_bytes = codec.encode_ping(id);
Ok(Bytes::from(frame_bytes))
}
MuxCommand::EncodeWindowUpdate { stream_id, delta } => {
let frame_bytes = codec.encode_window_update(stream_id, delta);
Ok(Bytes::from(frame_bytes))
}
MuxCommand::GoAway {
last_good_stream_id,
} => {
let frame_bytes = codec.encode_goaway(last_good_stream_id, GOAWAY_STATUS_OK);
tracing::info!(
last_good_stream_id,
"SPDY writer: sending GOAWAY (graceful shutdown)"
);
Ok(Bytes::from(frame_bytes))
}
}
}
pub(super) async fn run_writer<W: WsFrameWriter>(parts: WriterParts<W>, config: WriterConfig) {
let WriterParts {
mut writer,
mut cmd_rx,
mut control_rx,
mut window_rx,
close_reg_txs,
cancel,
} = parts;
let mut codec = SpdyCodec::with_max_frame_size(config.max_frame_size);
{
let ping_frame = codec.encode_ping(1);
tracing::debug!("SPDY writer: sending initial PING");
let write_timeout = config.write_timeout;
match tokio::time::timeout(write_timeout, writer.write_binary(Bytes::from(ping_frame)))
.await
{
Ok(Ok(())) => {}
Ok(Err(e)) => {
tracing::warn!("SPDY writer: initial PING failed: {e}");
cancel.cancel();
return;
}
Err(_) => {
tracing::warn!("SPDY writer: initial PING timed out after {write_timeout:?}");
cancel.cancel();
return;
}
}
match tokio::time::timeout(write_timeout, writer.flush()).await {
Ok(Ok(())) => {}
Ok(Err(e)) => {
tracing::warn!("SPDY writer: initial PING flush failed: {e}");
cancel.cancel();
return;
}
Err(_) => {
tracing::warn!("SPDY writer: initial PING flush timed out after {write_timeout:?}");
cancel.cancel();
return;
}
}
}
{
let settings_entries = vec![
(7, config.initial_window_size), (4, config.max_concurrent_streams), ];
let settings_frame = codec.encode_settings(&settings_entries);
tracing::debug!(
initial_window_size = config.initial_window_size,
max_concurrent_streams = config.max_concurrent_streams,
"SPDY writer: sending our SETTINGS"
);
let write_timeout = config.write_timeout;
match tokio::time::timeout(
write_timeout,
writer.write_binary(Bytes::from(settings_frame)),
)
.await
{
Ok(Ok(())) => {}
Ok(Err(e)) => {
tracing::warn!("SPDY writer: SETTINGS send failed: {e}");
cancel.cancel();
return;
}
Err(_) => {
tracing::warn!("SPDY writer: SETTINGS send timed out after {write_timeout:?}");
cancel.cancel();
return;
}
}
match tokio::time::timeout(write_timeout, writer.flush()).await {
Ok(Ok(())) => {}
Ok(Err(e)) => {
tracing::warn!("SPDY writer: SETTINGS flush failed: {e}");
cancel.cancel();
return;
}
Err(_) => {
tracing::warn!("SPDY writer: SETTINGS flush timed out after {write_timeout:?}");
cancel.cancel();
return;
}
}
}
let mut ping_interval = tokio::time::interval(config.ping_interval);
ping_interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
ping_interval.tick().await;
let mut last_writer_progress = tokio::time::Instant::now();
let mut stall_check = tokio::time::interval(Duration::from_secs(5));
stall_check.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
stall_check.tick().await;
let write_timeout = config.write_timeout;
let mut ping_id: u32 = 3;
let mut had_error = false;
let mut pending_closes: Vec<u32> = Vec::new();
const OPPORTUNISTIC_CAP: usize = WRITE_BATCH_CAP - 1;
let mut bytes_since_flush: usize = 0;
let mut batch_start: Option<tokio::time::Instant> = None;
loop {
tokio::select! {
biased;
() = cancel.cancelled() => break,
win = window_rx.recv() => {
let Some(cmd) = win else {
tracing::debug!("SPDY writer: window channel closed");
break;
};
let (err, n) = process_writer_command(
cmd, &mut writer, &mut codec, &mut pending_closes, write_timeout,
).await;
bytes_since_flush += n;
if batch_start.is_none() && n > 0 { batch_start = Some(tokio::time::Instant::now()); }
if err { had_error = true; }
if !had_error {
while let Ok(c) = window_rx.try_recv() {
let (err, n) = process_writer_command(
c, &mut writer, &mut codec, &mut pending_closes, write_timeout,
).await;
bytes_since_flush += n;
if err {
had_error = true;
break;
}
}
}
}
ctrl = control_rx.recv() => {
let Some(cmd) = ctrl else {
tracing::debug!("SPDY writer: control channel closed");
break;
};
let (err, n) = process_writer_command(
cmd, &mut writer, &mut codec, &mut pending_closes, write_timeout,
).await;
bytes_since_flush += n;
if batch_start.is_none() && n > 0 { batch_start = Some(tokio::time::Instant::now()); }
if err { had_error = true; }
if !had_error {
for _ in 0..OPPORTUNISTIC_CAP {
match control_rx.try_recv() {
Ok(c) => {
let (err, n) = process_writer_command(
c, &mut writer, &mut codec, &mut pending_closes, write_timeout,
).await;
bytes_since_flush += n;
if err {
had_error = true;
break;
}
}
Err(_) => break,
}
}
}
}
cmd = cmd_rx.recv() => {
let Some(cmd) = cmd else {
tracing::debug!("SPDY writer: command channel closed");
break;
};
while let Ok(c) = control_rx.try_recv() {
let (err, n) = process_writer_command(
c, &mut writer, &mut codec, &mut pending_closes, write_timeout,
).await;
bytes_since_flush += n;
if batch_start.is_none() && n > 0 { batch_start = Some(tokio::time::Instant::now()); }
if err {
had_error = true;
break;
}
}
if !had_error {
let (err, n) = process_writer_command(
cmd, &mut writer, &mut codec, &mut pending_closes, write_timeout,
).await;
bytes_since_flush += n;
if batch_start.is_none() && n > 0 { batch_start = Some(tokio::time::Instant::now()); }
if err { had_error = true; }
}
if !had_error {
for _ in 0..OPPORTUNISTIC_CAP {
match cmd_rx.try_recv() {
Ok(c) => {
let (err, n) = process_writer_command(
c, &mut writer, &mut codec, &mut pending_closes, write_timeout,
).await;
bytes_since_flush += n;
if err {
had_error = true;
break;
}
}
Err(_) => break,
}
}
}
}
_ = ping_interval.tick() => {
let ping = codec.encode_ping(ping_id);
ping_id = ping_id.wrapping_add(2);
match tokio::time::timeout(write_timeout, writer.write_binary(Bytes::from(ping))).await {
Ok(Ok(())) => {}
Ok(Err(e)) => {
tracing::warn!("SPDY writer: keepalive PING feed failed: {e}");
had_error = true;
}
Err(_) => {
tracing::warn!("SPDY writer: keepalive PING timed out after {write_timeout:?}");
had_error = true;
}
}
}
_ = stall_check.tick() => {
let stall_threshold = write_timeout.saturating_mul(2);
let queued = control_rx.len() + cmd_rx.len();
if last_writer_progress.elapsed() > stall_threshold && queued > 0 {
tracing::warn!(
stalled_for = ?last_writer_progress.elapsed(),
queued_commands = queued,
"SPDY writer: stall watchdog triggered, killing handle"
);
cancel.cancel();
break;
}
continue;
}
}
if had_error {
break;
}
if bytes_since_flush > 0 {
match tokio::time::timeout(write_timeout, writer.flush()).await {
Ok(Ok(())) => {
last_writer_progress = tokio::time::Instant::now();
}
Ok(Err(e)) => {
tracing::warn!("SPDY writer: flush error: {e}");
break;
}
Err(_) => {
tracing::warn!(
"SPDY writer: flush timed out after {write_timeout:?}, peer stalled"
);
break;
}
}
bytes_since_flush = 0;
batch_start = None;
}
let mut requeued: Vec<u32> = Vec::new();
for stream_id in std::mem::take(&mut pending_closes) {
match close_reg_txs[(stream_id % FRAME_WORKERS as u32) as usize]
.try_send(StreamRegistration::Close { stream_id })
{
Ok(()) => {}
Err(mpsc::error::TrySendError::Full(_)) => {
requeued.push(stream_id);
}
Err(mpsc::error::TrySendError::Closed(_)) => {
had_error = true;
break;
}
}
}
pending_closes.extend(requeued);
if had_error {
break;
}
}
cancel.cancel();
let _ = writer.close().await;
}
pub(super) async fn process_writer_command<W: WsFrameWriter>(
cmd: MuxCommand, writer: &mut W, codec: &mut SpdyCodec, pending_closes: &mut Vec<u32>,
write_timeout: Duration,
) -> (bool, usize) {
let mut bytes_written: usize = 0;
if let MuxCommand::OpenStreamPairAndWrite {
error_id,
data_id,
error_headers,
data_headers,
first_payload,
} = cmd
{
match codec.encode_syn_stream(error_id, &error_headers, false) {
Ok(frame_bytes) => {
let len = frame_bytes.len();
tracing::debug!(
stream_id = error_id,
fin = false,
len,
"SPDY writer: sending SYN_STREAM (error)"
);
write_with_timeout!(
writer,
frame_bytes,
write_timeout,
bytes_written += len,
"SPDY writer: SYN_STREAM (error) feed error: {}",
"SPDY writer: SYN_STREAM (error) timed out after {:?}"
);
}
Err(e) => {
tracing::warn!("SPDY writer: SYN_STREAM (error) encode error: {e}");
return (true, bytes_written);
}
}
match codec.encode_syn_stream(data_id, &data_headers, false) {
Ok(frame_bytes) => {
let len = frame_bytes.len();
tracing::debug!(
stream_id = data_id,
fin = false,
len,
"SPDY writer: sending SYN_STREAM (data)"
);
write_with_timeout!(
writer,
frame_bytes,
write_timeout,
bytes_written += len,
"SPDY writer: SYN_STREAM (data) feed error: {}",
"SPDY writer: SYN_STREAM (data) timed out after {:?}"
);
}
Err(e) => {
tracing::warn!("SPDY writer: SYN_STREAM (data) encode error: {e}");
return (true, bytes_written);
}
}
{
let frame_bytes = codec.encode_data(error_id, &[], true);
let len = frame_bytes.len();
tracing::debug!(
stream_id = error_id,
fin = true,
len,
"SPDY writer: half-closing error stream with empty DATA+FIN"
);
write_with_timeout!(
writer,
frame_bytes,
write_timeout,
bytes_written += len,
"SPDY writer: error-stream DATA+FIN feed error: {}",
"SPDY writer: error-stream DATA+FIN timed out after {:?}"
);
}
if !first_payload.is_empty() {
let frame_bytes = codec.encode_data(data_id, &first_payload, false);
let len = frame_bytes.len();
tracing::debug!(
stream_id = data_id,
fin = false,
len,
"SPDY writer: sending first DATA after lazy open"
);
write_with_timeout!(
writer,
frame_bytes,
write_timeout,
bytes_written += len,
"SPDY writer: lazy-open first DATA feed error: {}",
"SPDY writer: lazy-open first DATA timed out after {:?}"
);
}
return (false, bytes_written);
}
if let MuxCommand::SendWsPong { payload } = cmd {
match tokio::time::timeout(write_timeout, writer.write_pong(payload)).await {
Ok(Ok(())) => {}
Ok(Err(e)) => {
tracing::warn!("SPDY writer: pong feed error: {e}");
return (true, 0);
}
Err(_) => {
tracing::warn!("SPDY writer: pong timed out after {write_timeout:?}");
return (true, 0);
}
}
return (false, 0);
}
if let MuxCommand::CloseStream { stream_id, status } = cmd {
let rst_bytes = codec.encode_rst_stream(stream_id, status);
let len = rst_bytes.len();
write_with_timeout!(
writer,
rst_bytes,
write_timeout,
bytes_written += len,
"SPDY writer: RST_STREAM feed error: {}",
"SPDY writer: RST_STREAM timed out after {:?}"
);
pending_closes.push(stream_id);
return (false, bytes_written);
}
match encode_command(cmd, codec) {
Ok(payload) => {
let len = payload.len();
write_with_timeout!(
writer,
payload,
write_timeout,
bytes_written += len,
"SPDY writer: feed error: {}",
"SPDY writer: write_binary timed out after {:?}"
);
(false, bytes_written)
}
Err(e) => {
tracing::warn!("SPDY writer: encode error: {e}");
(true, 0)
}
}
}