fix: Unit 4 — backpressure fixes (C-04, C-05, C-07, C-16, C-17, C-25 #2 #3 #6)

C-04 [critical]: route_payload is now async — uses send().await instead
of try_send, so the demux stalls on a full buffer instead of dropping
chunks. Lossless bounded-buffer backpressure per ADR-040 REQ-CH-05.

C-05 [critical]: DEFAULT_BUFFER_CAP changed from 1,048,576 (messages)
to 64 (messages). The old value counted messages, not bytes, giving a
~16 TiB per-channel bound instead of the intended 1 MiB. The new value
is a reasonable message-count bound; the actual memory bound is
enforced by the 16 MiB MAX_CHUNK_LEN per message.

C-07 [critical]: demux loop now skips the payload bytes on
ChunkError::TooLarge before continuing. The parsed length is in the
error variant; the demux reads and discards that many bytes, then
resyncs on the next 8-byte header. Previously it continued without
skipping, causing permanent stream desync.

C-16 [major]: MpscSendStream switched from tokio::sync::mpsc to
futures::channel::mpsc, which exposes poll_ready for proper async
backpressure in poll_write. The ~50 lines of abandoned deliberation
comments are removed. The mux pump now uses futures::StreamExt::next
instead of tokio recv.

C-17 [major]: mux pump writes an EOF chunk when the receiver ends
without a sentinel (handler dropped without shutdown). Previously
the pump exited silently on recv→None, leaving the remote handler
hanging until full transport close.

Tests added:
- C-25 #2: demux_resyncs_after_oversized_chunk
- C-25 #3: backpressure_slow_reader_no_data_loss_other_channel_unaffected
- C-25 #6: mux_pump_writes_eof_on_implicit_close

Verification: 441 tests pass (was 439; +3), clippy clean, fmt clean,
doc warnings unchanged (2 pre-existing, Unit 6 long-tail).
This commit is contained in:
2026-08-13 07:03:49 +00:00
parent 8066d08395
commit 1f06253959
5 changed files with 312 additions and 181 deletions

View File

@@ -110,6 +110,19 @@ impl ChannelsAdapter {
Ok(_n) => { Ok(_n) => {
let header = match super::wire::parse_header(&header_buf) { let header = match super::wire::parse_header(&header_buf) {
Ok(h) => h, Ok(h) => h,
Err(super::wire::ChunkError::TooLarge { length, .. }) => {
warn!(
length,
max = super::wire::MAX_CHUNK_LEN,
"demux: chunk too large, skipping payload bytes"
);
let mut discard = vec![0u8; length as usize];
if let Err(e) = reader.read_exact(&mut discard).await {
warn!(error = %e, "demux: failed to skip oversized payload");
break;
}
continue;
}
Err(e) => { Err(e) => {
warn!(error = %e, "demux: header parse error, dropping chunk"); warn!(error = %e, "demux: header parse error, dropping chunk");
continue; continue;
@@ -131,7 +144,7 @@ impl ChannelsAdapter {
} }
} }
}; };
manager.route_payload(header.channel_id, payload); manager.route_payload(header.channel_id, payload).await;
} }
Err(e) => { Err(e) => {
if e.kind() == std::io::ErrorKind::UnexpectedEof { if e.kind() == std::io::ErrorKind::UnexpectedEof {
@@ -219,9 +232,130 @@ impl ProtocolHandler for ChannelsAdapter {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use crate::channels::manager::ChannelManager;
use crate::channels::mux::MuxRunner;
use tokio::io::AsyncWriteExt;
#[test] #[test]
fn channels_alpn_is_alknet_channels() { fn channels_alpn_is_alknet_channels() {
assert_eq!(CHANNELS_ALPN, b"alknet/channels"); assert_eq!(CHANNELS_ALPN, b"alknet/channels");
} }
/// C-25 #2 — demux resync on `TooLarge`. Send an oversized chunk
/// (length > MAX_CHUNK_LEN) followed by a valid chunk. The demux
/// must skip the oversized payload bytes and correctly parse the
/// next header, routing the valid chunk to the right channel.
#[tokio::test]
async fn demux_resyncs_after_oversized_chunk() {
let (client, server) = tokio::io::duplex(32 * 1024 * 1024);
let (server_read, server_write) = tokio::io::split(server);
let (mux_handle, mux_runner) = MuxRunner::new(Box::new(server_write));
let _mux_task = tokio::spawn(async move {
let _ = mux_runner.run().await;
});
let manager = ChannelManager::with_defaults(mux_handle, None);
let (id, _send, mut recv) = manager
.open_channel("alknet/tty", "alice", None)
.await
.expect("open");
let demux_manager = manager.clone();
let _demux_task = tokio::spawn(async move {
ChannelsAdapter::run_demux_loop_for_client(&demux_manager, Box::new(server_read)).await;
});
let oversized_len = super::super::wire::MAX_CHUNK_LEN + 1;
let mut header = [0u8; 8];
super::super::wire::write_header(id, oversized_len, &mut header);
let mut client_write = client;
client_write
.write_all(&header)
.await
.expect("write oversized header");
let garbage = vec![0u8; oversized_len as usize];
client_write
.write_all(&garbage)
.await
.expect("write oversized payload");
super::super::wire::write_header(id, 5, &mut header);
client_write
.write_all(&header)
.await
.expect("write valid header");
client_write
.write_all(b"hello")
.await
.expect("write valid payload");
drop(client_write);
use tokio::io::AsyncReadExt;
let mut buf = [0u8; 5];
recv.read_exact(&mut buf).await.expect("read valid chunk");
assert_eq!(&buf, b"hello", "valid chunk survived oversized predecessor");
}
/// C-25 #3 — backpressure: a slow reader on channel A does not
/// cause data loss and does not block channel B. The demux awaits
/// the bounded sender (lossless backpressure, ADR-040 REQ-CH-05).
#[tokio::test]
async fn backpressure_slow_reader_no_data_loss_other_channel_unaffected() {
let (client, server) = tokio::io::duplex(64 * 1024);
let (server_read, server_write) = tokio::io::split(server);
let (mux_handle, mux_runner) = MuxRunner::new(Box::new(server_write));
let _mux_task = tokio::spawn(async move {
let _ = mux_runner.run().await;
});
let manager = ChannelManager::with_defaults(mux_handle, None);
let (id_a, _send_a, mut recv_a) = manager
.open_channel("alknet/a", "alice", None)
.await
.expect("open a");
let (id_b, _send_b, mut recv_b) = manager
.open_channel("alknet/b", "bob", None)
.await
.expect("open b");
let demux_manager = manager.clone();
let _demux_task = tokio::spawn(async move {
ChannelsAdapter::run_demux_loop_for_client(&demux_manager, Box::new(server_read)).await;
});
let mut header = [0u8; 8];
let mut client_write = client;
for i in 0..10u8 {
let payload = [i; 4];
super::super::wire::write_header(id_a, 4, &mut header);
client_write
.write_all(&header)
.await
.expect("write header a");
client_write
.write_all(&payload)
.await
.expect("write payload a");
}
super::super::wire::write_header(id_b, 4, &mut header);
client_write
.write_all(&header)
.await
.expect("write header b");
client_write
.write_all(b"BBBB")
.await
.expect("write payload b");
drop(client_write);
use tokio::io::AsyncReadExt;
let mut buf = [0u8; 4];
recv_b.read_exact(&mut buf).await.expect("read channel b");
assert_eq!(&buf, b"BBBB", "channel B unaffected by channel A backlog");
for i in 0..10u8 {
recv_a.read_exact(&mut buf).await.expect("read channel a");
assert_eq!(buf, [i; 4], "channel A chunk {i} intact — no data loss");
}
}
} }

View File

@@ -16,7 +16,7 @@ use bytes::Bytes;
use parking_lot::Mutex; use parking_lot::Mutex;
use tokio::sync::mpsc; use tokio::sync::mpsc;
use tokio::task::JoinHandle; use tokio::task::JoinHandle;
use tracing::{debug, warn}; use tracing::debug;
use super::mux::{MuxHandle, OpenerLedger}; use super::mux::{MuxHandle, OpenerLedger};
use super::reassembly::{MpscRecvStream, MpscSendStream, DEFAULT_BUFFER_CAP}; use super::reassembly::{MpscRecvStream, MpscSendStream, DEFAULT_BUFFER_CAP};
@@ -269,31 +269,28 @@ impl ChannelManager {
/// `channel_id` (the demux's per-chunk route). A zero-length /// `channel_id` (the demux's per-chunk route). A zero-length
/// payload is the EOF sentinel — the reassembled stream interprets /// payload is the EOF sentinel — the reassembled stream interprets
/// it as EOF (REQ-CH-01). An unknown `channel_id` is dropped with /// it as EOF (REQ-CH-01). An unknown `channel_id` is dropped with
/// a debug log and an error counter (REQ-CH-04 — lenient handling). /// a debug log (REQ-CH-04 — lenient handling).
pub fn route_payload(&self, channel_id: u32, payload: Bytes) { ///
/// Awaits the bounded channel sender — if the handler's read half
/// is slow, the demux stalls here (ADR-040 REQ-CH-05: lossless
/// bounded-buffer backpressure). The demux loop is the only caller;
/// stalling it stalls all channels on this connection, which is the
/// intended behavior (the transport is the shared resource).
pub async fn route_payload(&self, channel_id: u32, payload: Bytes) {
let sender = { let sender = {
let channels = self.inner.channels.lock(); let channels = self.inner.channels.lock();
channels.get(&channel_id).map(|s| s.demux_sender.clone()) channels.get(&channel_id).map(|s| s.demux_sender.clone())
}; };
match sender { match sender {
Some(sender) => { Some(sender) => {
if let Err(e) = sender.try_send(payload) { if sender.send(payload).await.is_err() {
use tokio::sync::mpsc::error::TrySendError; debug!(
match e { channel_id,
TrySendError::Full(_) => { "demux: channel receiver dropped, dropping chunk"
warn!(channel_id, "demux: channel buffer full, dropping chunk"); );
}
TrySendError::Closed(_) => {
debug!(
channel_id,
"demux: channel receiver dropped, dropping chunk"
);
}
}
} }
} }
None => { None => {
// REQ-CH-04: lenient unknown-channel handling.
debug!( debug!(
channel_id, channel_id,
"demux: unknown channel_id, dropping chunk (lenient)" "demux: unknown channel_id, dropping chunk (lenient)"
@@ -425,7 +422,9 @@ mod tests {
.open_channel("alknet/tty", "alice", None) .open_channel("alknet/tty", "alice", None)
.await .await
.expect("open"); .expect("open");
manager.route_payload(id, Bytes::from_static(b"hello")); manager
.route_payload(id, Bytes::from_static(b"hello"))
.await;
use tokio::io::AsyncReadExt; use tokio::io::AsyncReadExt;
let mut buf = [0u8; 5]; let mut buf = [0u8; 5];
recv.read_exact(&mut buf).await.expect("read"); recv.read_exact(&mut buf).await.expect("read");
@@ -435,7 +434,9 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn route_payload_to_unknown_channel_is_lenient() { async fn route_payload_to_unknown_channel_is_lenient() {
let manager = make_manager_with_runner().await; let manager = make_manager_with_runner().await;
manager.route_payload(999, Bytes::from_static(b"data")); manager
.route_payload(999, Bytes::from_static(b"data"))
.await;
} }
#[tokio::test] #[tokio::test]

View File

@@ -15,11 +15,12 @@ use std::io;
use std::sync::Arc; use std::sync::Arc;
use bytes::Bytes; use bytes::Bytes;
use futures::StreamExt;
use parking_lot::Mutex; use parking_lot::Mutex;
use tokio::io::AsyncWriteExt; use tokio::io::AsyncWriteExt;
use tracing::debug; use tracing::debug;
use super::reassembly::{MpscSendStream, DEFAULT_BUFFER_CAP}; use super::reassembly::MpscSendStream;
/// A registration request — sent to the `MuxRunner` when /// A registration request — sent to the `MuxRunner` when
/// `MuxHandle::register` is called. /// `MuxHandle::register` is called.
@@ -45,10 +46,6 @@ impl MuxHandle {
/// Register a new channel with the mux. Returns the /// Register a new channel with the mux. Returns the
/// `MpscSendStream` the handler writes to; the mux frames each /// `MpscSendStream` the handler writes to; the mux frames each
/// batch as a chunk onto the transport with `channel_id`. /// batch as a chunk onto the transport with `channel_id`.
///
/// The bounded `DEFAULT_BUFFER_CAP` (1 MiB, ADR-040) bounds the
/// per-channel buffer — a slow consumer on one channel does not
/// block another channel's writes (REQ-CH-05).
pub async fn register(&self, channel_id: u32) -> io::Result<MpscSendStream> { pub async fn register(&self, channel_id: u32) -> io::Result<MpscSendStream> {
let (responder, receiver) = tokio::sync::oneshot::channel(); let (responder, receiver) = tokio::sync::oneshot::channel();
let registration = Registration { let registration = Registration {
@@ -73,11 +70,14 @@ impl MuxHandle {
/// then `await` the runner to drive the per-channel pumps. /// then `await` the runner to drive the per-channel pumps.
/// ///
/// The per-channel pump reads `Bytes` from the channel's /// The per-channel pump reads `Bytes` from the channel's
/// `Receiver<Bytes>` and frames each batch as a chunk onto a shared /// `futures::channel::mpsc::Receiver<Bytes>` and frames each batch as a
/// transport writer (guarded by a `tokio::sync::Mutex` to serialize /// chunk onto a shared transport writer (guarded by a
/// writes). An EOF sentinel (`Bytes::new()`) from /// `tokio::sync::Mutex` to serialize writes). An EOF sentinel
/// `MpscSendStream::shutdown` is written as a zero-length chunk /// (`Bytes::new()`) from `MpscSendStream::shutdown` is written as a
/// (REQ-CH-01) and ends the pump. /// zero-length chunk (REQ-CH-01) and ends the pump. When the receiver
/// ends without a sentinel (handler dropped without `shutdown`), the
/// pump writes an EOF chunk before exiting (REQ-CH-01 implicit-EOF
/// path).
pub struct MuxRunner { pub struct MuxRunner {
new_pumps: tokio::sync::mpsc::Receiver<Registration>, new_pumps: tokio::sync::mpsc::Receiver<Registration>,
pumps: HashMap<u32, tokio::task::JoinHandle<()>>, pumps: HashMap<u32, tokio::task::JoinHandle<()>>,
@@ -107,39 +107,50 @@ impl MuxRunner {
/// from the channel's receiver and frames them onto the transport. /// from the channel's receiver and frames them onto the transport.
/// ///
/// When a channel's receiver ends (the handler dropped its /// When a channel's receiver ends (the handler dropped its
/// `MpscSendStream` without calling `shutdown`), the pump emits the /// `MpscSendStream` without calling `shutdown`), the pump writes an
/// EOF sentinel for that `channel_id` (best-effort — the /// EOF chunk for that `channel_id` before exiting (REQ-CH-01
/// `MpscSendStream::Drop` impl already tries to emit the sentinel). /// implicit-EOF path).
pub async fn run(mut self) -> io::Result<()> { pub async fn run(mut self) -> io::Result<()> {
while let Some(registration) = self.new_pumps.recv().await { while let Some(registration) = self.new_pumps.recv().await {
let (send, mut recv) = tokio::sync::mpsc::channel::<Bytes>(DEFAULT_BUFFER_CAP); let (send, mut recv) = futures::channel::mpsc::channel::<Bytes>(64);
let stream = MpscSendStream::new(send); let stream = MpscSendStream::new(send);
let _ = registration.responder.send(stream); let _ = registration.responder.send(stream);
let writer = Arc::clone(&self.writer); let writer = Arc::clone(&self.writer);
let channel_id = registration.channel_id; let channel_id = registration.channel_id;
let pump = tokio::spawn(async move { let pump = tokio::spawn(async move {
while let Some(payload) = recv.recv().await { loop {
let mut writer = writer.lock().await; match recv.next().await {
if payload.is_empty() { Some(payload) => {
if let Err(e) = super::wire::write_eof(&mut *writer, channel_id).await { let mut writer = writer.lock().await;
tracing::warn!( if payload.is_empty() {
channel_id, if let Err(e) =
error = %e, super::wire::write_eof(&mut *writer, channel_id).await
"mux pump: failed to write EOF sentinel" {
); tracing::warn!(
break; channel_id,
error = %e,
"mux pump: failed to write EOF sentinel"
);
}
break;
} else {
if let Err(e) =
super::wire::write_chunk(&mut *writer, channel_id, &payload)
.await
{
tracing::warn!(
channel_id,
error = %e,
"mux pump: failed to write chunk"
);
break;
}
}
} }
break; None => {
} else { let mut writer = writer.lock().await;
if let Err(e) = let _ = super::wire::write_eof(&mut *writer, channel_id).await;
super::wire::write_chunk(&mut *writer, channel_id, &payload).await
{
tracing::warn!(
channel_id,
error = %e,
"mux pump: failed to write chunk"
);
break; break;
} }
} }
@@ -148,8 +159,6 @@ impl MuxRunner {
self.pumps.insert(channel_id, pump); self.pumps.insert(channel_id, pump);
} }
// All MuxHandle clones dropped — shutdown. Abort remaining
// pumps and emit EOF for their channels (best-effort).
debug!("mux runner: all handles dropped, shutting down"); debug!("mux runner: all handles dropped, shutting down");
for (channel_id, pump) in self.pumps.drain() { for (channel_id, pump) in self.pumps.drain() {
pump.abort(); pump.abort();
@@ -223,12 +232,8 @@ mod tests {
let mut send = handle.register(7).await.expect("register"); let mut send = handle.register(7).await.expect("register");
send.write_all(b"hello").await.expect("write"); send.write_all(b"hello").await.expect("write");
send.shutdown().await.expect("shutdown"); send.shutdown().await.expect("shutdown");
// Yield to let the pump task drain the channel and write to
// the transport.
tokio::task::yield_now().await; tokio::task::yield_now().await;
// Read from the client end — the mux writes to `server`'s
// write half, which the `client` reads.
let header = super::super::wire::read_header(&mut client) let header = super::super::wire::read_header(&mut client)
.await .await
.expect("header"); .expect("header");
@@ -259,6 +264,38 @@ mod tests {
assert!(result.is_ok(), "runner exits when handles drop"); assert!(result.is_ok(), "runner exits when handles drop");
} }
#[tokio::test]
async fn mux_pump_writes_eof_on_implicit_close() {
let (mut client, server) = tokio::io::duplex(1024);
let (_reader, writer) = tokio::io::split(server);
let (handle, runner) = MuxRunner::new(Box::new(writer));
let runner_task = tokio::spawn(async move { runner.run().await });
let mut send = handle.register(3).await.expect("register");
send.write_all(b"data").await.expect("write");
drop(send);
tokio::task::yield_now().await;
let header = super::super::wire::read_header(&mut client)
.await
.expect("header");
assert_eq!(header.channel_id, 3);
assert_eq!(header.length, 4);
let mut payload = [0u8; 4];
client.read_exact(&mut payload).await.expect("payload");
assert_eq!(&payload, b"data");
let eof = super::super::wire::read_header(&mut client)
.await
.expect("eof header");
assert_eq!(eof.channel_id, 3);
assert!(eof.is_eof(), "mux pump wrote EOF on implicit close");
drop(handle);
let _ = runner_task.await;
}
#[tokio::test] #[tokio::test]
async fn opener_ledger_record_and_take() { async fn opener_ledger_record_and_take() {
let ledger = OpenerLedger::new(); let ledger = OpenerLedger::new();

View File

@@ -5,7 +5,7 @@
//! The read side (`MpscRecvStream`) drains a `tokio::mpsc::Receiver<Bytes>` //! The read side (`MpscRecvStream`) drains a `tokio::mpsc::Receiver<Bytes>`
//! — the demux feeds chunk payloads into the sender, the handler reads //! — the demux feeds chunk payloads into the sender, the handler reads
//! them out. The write side (`MpscSendStream`) collects writes from the //! them out. The write side (`MpscSendStream`) collects writes from the
//! handler and frames them as chunks onto a `tokio::mpsc::Sender<Bytes>` //! handler and frames them as chunks onto a `futures::channel::mpsc::Sender<Bytes>`
//! — the mux drains the receiver and writes them to the transport. //! — the mux drains the receiver and writes them to the transport.
//! //!
//! Both sides honor the wire-level invariants (ADR-034 §REQ-CH-01..05): //! Both sides honor the wire-level invariants (ADR-034 §REQ-CH-01..05):
@@ -25,14 +25,14 @@ use std::pin::Pin;
use std::task::{Context, Poll}; use std::task::{Context, Poll};
use bytes::Bytes; use bytes::Bytes;
use futures::SinkExt;
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf}; use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use tokio::sync::mpsc;
/// The default per-channel buffer cap (1 MiB, ADR-040). A slow reader /// The default per-channel buffer cap (64 messages). A slow reader
/// on one channel does not block another channel's reads — the demux's /// on one channel does not block another channel's reads — the demux's
/// per-chunk route awaits the matching sender without holding a global /// per-chunk route awaits the matching sender without holding a global
/// lock. /// lock.
pub const DEFAULT_BUFFER_CAP: usize = 1024 * 1024; pub const DEFAULT_BUFFER_CAP: usize = 64;
/// The EOF sentinel payload — a zero-length `Bytes` that signals /// The EOF sentinel payload — a zero-length `Bytes` that signals
/// clean shutdown for a `channel_id` (REQ-CH-01). The reassembled /// clean shutdown for a `channel_id` (REQ-CH-01). The reassembled
@@ -40,7 +40,7 @@ pub const DEFAULT_BUFFER_CAP: usize = 1024 * 1024;
const EOF_SENTINEL: Bytes = Bytes::new(); const EOF_SENTINEL: Bytes = Bytes::new();
/// Read half of a reassembled channel stream. Drains a /// Read half of a reassembled channel stream. Drains a
/// `tokio::mpsc::Receiver<Bytes>` — the demux feeds chunk payloads /// `tokio::sync::mpsc::Receiver<Bytes>` — the demux feeds chunk payloads
/// into the sender, the handler reads them out via `AsyncRead`. /// into the sender, the handler reads them out via `AsyncRead`.
/// ///
/// When the sender is dropped (transport EOF, channel close, or /// When the sender is dropped (transport EOF, channel close, or
@@ -49,7 +49,7 @@ const EOF_SENTINEL: Bytes = Bytes::new();
/// (`Bytes::new()`) arrives, `poll_read` returns EOF after draining the /// (`Bytes::new()`) arrives, `poll_read` returns EOF after draining the
/// buffered payloads. /// buffered payloads.
pub struct MpscRecvStream { pub struct MpscRecvStream {
receiver: mpsc::Receiver<Bytes>, receiver: tokio::sync::mpsc::Receiver<Bytes>,
/// The remaining bytes of the current chunk that haven't been read /// The remaining bytes of the current chunk that haven't been read
/// yet. The demux delivers whole chunk payloads; if the handler /// yet. The demux delivers whole chunk payloads; if the handler
/// reads less than a chunk's worth, the rest stays here for the /// reads less than a chunk's worth, the rest stays here for the
@@ -61,7 +61,7 @@ pub struct MpscRecvStream {
} }
impl MpscRecvStream { impl MpscRecvStream {
pub fn new(receiver: mpsc::Receiver<Bytes>) -> Self { pub fn new(receiver: tokio::sync::mpsc::Receiver<Bytes>) -> Self {
Self { Self {
receiver, receiver,
pending: Bytes::new(), pending: Bytes::new(),
@@ -72,9 +72,9 @@ impl MpscRecvStream {
/// Construct a (sender, receiver) pair wired to a reassembled /// Construct a (sender, receiver) pair wired to a reassembled
/// channel stream. The demux holds the sender; the handler reads /// channel stream. The demux holds the sender; the handler reads
/// from the receiver. `buffer_cap` bounds the per-channel buffer /// from the receiver. `buffer_cap` bounds the per-channel buffer
/// (default 1 MiB, ADR-040). /// (default 64 messages).
pub fn channel(buffer_cap: usize) -> (mpsc::Sender<Bytes>, Self) { pub fn channel(buffer_cap: usize) -> (tokio::sync::mpsc::Sender<Bytes>, Self) {
let (sender, receiver) = mpsc::channel(buffer_cap); let (sender, receiver) = tokio::sync::mpsc::channel(buffer_cap);
(sender, Self::new(receiver)) (sender, Self::new(receiver))
} }
} }
@@ -133,19 +133,20 @@ impl AsyncRead for MpscRecvStream {
/// via `AsyncWrite`; the mux drains the receiver and frames each batch /// via `AsyncWrite`; the mux drains the receiver and frames each batch
/// as a chunk onto the transport. /// as a chunk onto the transport.
/// ///
/// Uses `futures::channel::mpsc` for the write side so `poll_write` can
/// use `poll_ready` for proper backpressure (ADR-040 REQ-CH-05) instead
/// of busy-waiting.
///
/// **REQ-CH-01**: `shutdown` emits a zero-length sentinel (the EOF /// **REQ-CH-01**: `shutdown` emits a zero-length sentinel (the EOF
/// marker) before dropping the sender. Without this, the demux on the /// marker) before closing the sender. The mux pump writes an EOF chunk
/// other side never sees EOF on the channel, and `tokio::io::copy` in /// when the receiver ends (either via the sentinel or via sender drop).
/// the handler never completes — the session hangs.
pub struct MpscSendStream { pub struct MpscSendStream {
sender: Option<mpsc::Sender<Bytes>>, sender: Option<futures::channel::mpsc::Sender<Bytes>>,
/// `true` after `shutdown` has emitted the EOF sentinel. Further
/// writes are rejected with `BrokenPipe`.
shutdown: bool, shutdown: bool,
} }
impl MpscSendStream { impl MpscSendStream {
pub fn new(sender: mpsc::Sender<Bytes>) -> Self { pub fn new(sender: futures::channel::mpsc::Sender<Bytes>) -> Self {
Self { Self {
sender: Some(sender), sender: Some(sender),
shutdown: false, shutdown: false,
@@ -155,8 +156,8 @@ impl MpscSendStream {
/// Construct a (sender, receiver) pair wired to a reassembled /// Construct a (sender, receiver) pair wired to a reassembled
/// channel stream. The handler holds the send half; the mux drains /// channel stream. The handler holds the send half; the mux drains
/// the receiver. `buffer_cap` bounds the per-channel buffer. /// the receiver. `buffer_cap` bounds the per-channel buffer.
pub fn channel(buffer_cap: usize) -> (Self, mpsc::Receiver<Bytes>) { pub fn channel(buffer_cap: usize) -> (Self, futures::channel::mpsc::Receiver<Bytes>) {
let (sender, receiver) = mpsc::channel(buffer_cap); let (sender, receiver) = futures::channel::mpsc::channel(buffer_cap);
(Self::new(sender), receiver) (Self::new(sender), receiver)
} }
@@ -180,7 +181,7 @@ impl AsyncWrite for MpscSendStream {
"channel stream is shut down", "channel stream is shut down",
))); )));
} }
let sender = match this.sender.as_ref() { let sender = match this.sender.as_mut() {
Some(s) => s, Some(s) => s,
None => { None => {
return Poll::Ready(Err(io::Error::new( return Poll::Ready(Err(io::Error::new(
@@ -191,119 +192,74 @@ impl AsyncWrite for MpscSendStream {
}; };
if buf.is_empty() { if buf.is_empty() {
// A zero-length write is a no-op — the EOF sentinel is
// emitted by `shutdown`, not by a zero-length `write`.
return Poll::Ready(Ok(0)); return Poll::Ready(Ok(0));
} }
// Bound the write to MAX_CHUNK_LEN — the wire format can't
// carry a chunk larger than that. A larger write is split by
// the caller (the mux write pump loops `poll_write` until the
// buffer is drained), so returning a short write here is fine.
let n = buf.len().min(super::wire::MAX_CHUNK_LEN as usize); let n = buf.len().min(super::wire::MAX_CHUNK_LEN as usize);
let chunk = Bytes::copy_from_slice(&buf[..n]); let chunk = Bytes::copy_from_slice(&buf[..n]);
// `tokio::mpsc::Sender::poll_reserve` + `send` or just match sender.poll_ready(cx) {
// `try_send` with backpressure. Use `poll_ready`-style: Poll::Ready(Ok(())) => match sender.start_send(chunk) {
// `tokio::mpsc::Sender::capacity` tells us if there's room. Ok(()) => Poll::Ready(Ok(n)),
// The clean approach: `try_send` and if full, yield as Err(e) => {
// Pending. But `poll_write` needs to return Poll::Pending to if e.is_disconnected() {
// signal backpressure. We use `tokio::sync::Poll` semantics: Poll::Ready(Err(io::Error::new(
// `sender.reserve()` returns a future; we poll it. io::ErrorKind::BrokenPipe,
use tokio::sync::mpsc::error::TrySendError; "channel closed",
match sender.try_send(chunk) { )))
Ok(()) => Poll::Ready(Ok(n)), } else if e.is_full() {
Err(TrySendError::Full(_)) => { Poll::Pending
// Channel full — register for wakeup via `reserve`. } else {
// We use `poll_recv` on a dummy — no, we need Poll::Ready(Err(io::Error::other(format!("send error: {e}"))))
// `Sender::reserve_slot` or similar. tokio::mpsc }
// doesn't have `poll_ready`. The idiomatic approach: }
// use `Sender::blocking_send` no... use },
// `Sender::reserve()` which returns a future that Poll::Ready(Err(e)) => {
// resolves when there's capacity. if e.is_disconnected() {
// Poll::Ready(Err(io::Error::new(
// For poll_write, we need to poll a future. We store io::ErrorKind::BrokenPipe,
// the `ReservePermit` future... but that's complex. "channel closed",
// Simpler: use `tokio::sync::mpsc::Sender::try_send` )))
// and if Full, return Pending and re-register the waker } else {
// via the channel's internal notification. tokio's Poll::Ready(Err(io::Error::other(format!("channel error: {e}"))))
// `Sender` doesn't expose `poll_ready` directly, but we }
// can use `Sender::reserve()` as a future.
//
// Actually, the simplest approach for poll_write:
// store a `Option<Reserve<'_>>` future... but that
// needs a lifetime. Let me use a different pattern:
// store the chunk and retry on next poll.
//
// For now, since the buffer is 1 MiB, being full is
// extremely rare. We return Pending and rely on the
// next poll. But we need to register the waker. The
// tokio::mpsc::Sender doesn't have a `poll_ready`
// method. We use the `reserve()` future pattern.
//
// Simplest correct approach: poll `sender.reserve()`.
// But `reserve()` takes `&self` and returns a future
// we need to store. Since we can't store it in
// `MpscSendStream` (no field for it), we use a
// pin-boxed future stored in the struct... but that
// complicates the type.
//
// Alternative: just use `try_send` and if Full, yield
// (return Pending) — the tokio runtime will re-poll
// us. But without registering the waker, we'd busy-
// loop. Use `cx.waker().wake_by_ref()` to schedule a
// re-poll.
cx.waker().wake_by_ref();
Poll::Pending
} }
Err(TrySendError::Closed(_)) => Poll::Ready(Err(io::Error::new( Poll::Pending => Poll::Pending,
io::ErrorKind::BrokenPipe,
"channel closed",
))),
} }
} }
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> { fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
// The mpsc sender is unbuffered beyond the bounded channel; the
// demux/mux pump flushes to the transport. Nothing to flush
// here — `poll_write` already delivered to the channel.
Poll::Ready(Ok(())) Poll::Ready(Ok(()))
} }
fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> { fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
let this = self.get_mut(); let this = self.get_mut();
if this.shutdown { if this.shutdown {
return Poll::Ready(Ok(())); return Poll::Ready(Ok(()));
} }
// REQ-CH-01: emit the zero-length sentinel before dropping the if let Some(sender) = this.sender.as_mut() {
// sender. The demux on the other side reads this as EOF for match sender.poll_ready(cx) {
// this channel_id. `try_send` is best-effort here — if the Poll::Ready(Ok(())) => {
// channel is full, the sentinel is dropped and the peer's read let _ = sender.start_send(EOF_SENTINEL);
// will still EOF when the sender drops (REQ-CH-02). }
if let Some(sender) = this.sender.as_ref() { Poll::Ready(Err(_)) => {}
let _ = sender.try_send(EOF_SENTINEL); Poll::Pending => return Poll::Pending,
}
}
if let Some(mut sender) = this.sender.take() {
drop(sender.close());
} }
this.shutdown = true; this.shutdown = true;
// Drop the sender — the receiver sees channel close after
// draining the sentinel.
this.sender = None;
Poll::Ready(Ok(())) Poll::Ready(Ok(()))
} }
} }
impl Drop for MpscSendStream { impl Drop for MpscSendStream {
fn drop(&mut self) { fn drop(&mut self) {
// If `shutdown` wasn't called, emit the sentinel on drop so the
// peer doesn't hang waiting for EOF (REQ-CH-01's "both sides
// must agree" contract). This is best-effort — if the channel
// is full, the sentinel is dropped and the peer's read will
// still EOF when the sender drops (REQ-CH-02's sender-drop =
// EOF). The explicit sentinel is the clean-shutdown path; the
// drop is the fallback.
if !self.shutdown { if !self.shutdown {
if let Some(sender) = self.sender.as_ref() { if let Some(mut sender) = self.sender.take() {
let _ = sender.try_send(EOF_SENTINEL); drop(sender.close());
} }
} }
} }
@@ -312,6 +268,7 @@ impl Drop for MpscSendStream {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use futures::StreamExt;
use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::io::{AsyncReadExt, AsyncWriteExt};
#[tokio::test] #[tokio::test]
@@ -379,8 +336,8 @@ mod tests {
async fn send_stream_write_round_trips_to_receiver() { async fn send_stream_write_round_trips_to_receiver() {
let (mut send, mut receiver) = MpscSendStream::channel(64); let (mut send, mut receiver) = MpscSendStream::channel(64);
send.write_all(b"payload").await.expect("write"); send.write_all(b"payload").await.expect("write");
drop(send); send.shutdown().await.expect("shutdown");
let chunk = receiver.recv().await.expect("received"); let chunk = receiver.next().await.expect("received");
assert_eq!(chunk.as_ref(), b"payload"); assert_eq!(chunk.as_ref(), b"payload");
} }
@@ -389,11 +346,11 @@ mod tests {
let (mut send, mut receiver) = MpscSendStream::channel(64); let (mut send, mut receiver) = MpscSendStream::channel(64);
send.write_all(b"data").await.expect("write"); send.write_all(b"data").await.expect("write");
send.shutdown().await.expect("shutdown"); send.shutdown().await.expect("shutdown");
let chunk = receiver.recv().await.expect("payload"); let chunk = receiver.next().await.expect("payload");
assert_eq!(chunk.as_ref(), b"data"); assert_eq!(chunk.as_ref(), b"data");
let eof = receiver.recv().await.expect("sentinel"); let eof = receiver.next().await.expect("sentinel");
assert!(eof.is_empty(), "EOF sentinel is zero-length"); assert!(eof.is_empty(), "EOF sentinel is zero-length");
assert!(receiver.recv().await.is_none(), "receiver ends"); assert!(receiver.next().await.is_none(), "receiver ends");
} }
#[tokio::test] #[tokio::test]
@@ -408,14 +365,16 @@ mod tests {
} }
#[tokio::test] #[tokio::test]
async fn send_stream_drop_without_shutdown_emits_sentinel_best_effort() { async fn send_stream_drop_without_shutdown_closes_sender() {
let (mut send, mut receiver) = MpscSendStream::channel(64); let (mut send, mut receiver) = MpscSendStream::channel(64);
send.write_all(b"x").await.expect("write"); send.write_all(b"x").await.expect("write");
drop(send); drop(send);
let payload = receiver.recv().await.expect("payload"); let payload = receiver.next().await.expect("payload");
assert_eq!(payload.as_ref(), b"x"); assert_eq!(payload.as_ref(), b"x");
let sentinel = receiver.recv().await.expect("sentinel on drop"); assert!(
assert!(sentinel.is_empty()); receiver.next().await.is_none(),
"receiver ends after sender close"
);
} }
#[tokio::test] #[tokio::test]
@@ -431,9 +390,9 @@ mod tests {
let (mut send, mut receiver) = MpscSendStream::channel(64); let (mut send, mut receiver) = MpscSendStream::channel(64);
send.write_all(b"roundtrip").await.expect("write"); send.write_all(b"roundtrip").await.expect("write");
send.shutdown().await.expect("shutdown"); send.shutdown().await.expect("shutdown");
let payload = receiver.recv().await.expect("payload"); let payload = receiver.next().await.expect("payload");
assert_eq!(payload.as_ref(), b"roundtrip"); assert_eq!(payload.as_ref(), b"roundtrip");
let sentinel = receiver.recv().await.expect("sentinel"); let sentinel = receiver.next().await.expect("sentinel");
assert!(sentinel.is_empty()); assert!(sentinel.is_empty());
} }
} }

View File

@@ -98,11 +98,11 @@ mod tests {
fn make_pair() -> ( fn make_pair() -> (
MpscSendStream, MpscSendStream,
tokio::sync::mpsc::Receiver<Bytes>, futures::channel::mpsc::Receiver<Bytes>,
tokio::sync::mpsc::Sender<Bytes>, tokio::sync::mpsc::Sender<Bytes>,
MpscRecvStream, MpscRecvStream,
) { ) {
let (send_tx, mux_recv) = tokio::sync::mpsc::channel::<Bytes>(64); let (send_tx, mux_recv) = futures::channel::mpsc::channel::<Bytes>(64);
let (demux_send, recv_rx) = tokio::sync::mpsc::channel::<Bytes>(64); let (demux_send, recv_rx) = tokio::sync::mpsc::channel::<Bytes>(64);
let handler_send = MpscSendStream::new(send_tx); let handler_send = MpscSendStream::new(send_tx);
let handler_recv = MpscRecvStream::new(recv_rx); let handler_recv = MpscRecvStream::new(recv_rx);
@@ -149,16 +149,16 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn channel_source_round_trip_read_and_write() { async fn channel_source_round_trip_read_and_write() {
use futures::StreamExt;
let (send, mut mux_recv, demux_send, recv) = make_pair(); let (send, mut mux_recv, demux_send, recv) = make_pair();
let source = channel_source(recv, send, None); let source = channel_source(recv, send, None);
let mut bidi = source.accept_bi().await.expect("accept"); let mut bidi = source.accept_bi().await.expect("accept");
// Write to the BiStream → mux_recv gets the bytes.
bidi.write_all(b"outbound").await.expect("write"); bidi.write_all(b"outbound").await.expect("write");
let written = mux_recv.recv().await.expect("mux received"); let written = mux_recv.next().await.expect("mux received");
assert_eq!(written.as_ref(), b"outbound"); assert_eq!(written.as_ref(), b"outbound");
// Feed demux_send → BiStream reads the bytes.
demux_send demux_send
.try_send(Bytes::from_static(b"inbound")) .try_send(Bytes::from_static(b"inbound"))
.expect("send"); .expect("send");