use bytes::{Buf, BufMut, Bytes, BytesMut};
use std::collections::HashMap;
use std::io;
use std::pin::Pin;
use std::sync::{
atomic::{AtomicBool, AtomicU32, AtomicU8, Ordering},
Arc, Mutex,
};
use std::task::{Context, Poll};
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use tokio::sync::{mpsc, Notify};
use tokio::time::interval;
use tokio_util::sync::PollSender;
use tracing::{debug, trace, warn};
use crate::errors::{Error, Result};
use crate::session::Session;
const VERSION: u8 = 1;
const CMD_SYN: u8 = 0;
const CMD_FIN: u8 = 1;
const CMD_PSH: u8 = 2;
const CMD_NOP: u8 = 3;
const HEADER_SIZE: usize = 8;
const MAX_PAYLOAD: usize = 32 * 1024;
const FRAME_CHANNEL_CAP: usize = 256;
const STREAM_CHANNEL_CAP: usize = 64;
const KEEPALIVE_INTERVAL: std::time::Duration = std::time::Duration::from_secs(10);
fn encode_frame(cmd: u8, stream_id: u32, payload: &[u8]) -> Bytes {
let mut buf = BytesMut::with_capacity(HEADER_SIZE + payload.len());
buf.put_u8(VERSION);
buf.put_u8(cmd);
buf.put_u16_le(payload.len() as u16);
buf.put_u32_le(stream_id);
buf.put_slice(payload);
buf.freeze()
}
#[inline]
fn encode_ctrl(cmd: u8, stream_id: u32) -> Bytes {
encode_frame(cmd, stream_id, &[])
}
fn decode_frame(buf: &mut BytesMut) -> Option<(u8, u32, Bytes)> {
loop {
if buf.len() < HEADER_SIZE {
return None;
}
let version = buf[0];
let length = u16::from_le_bytes([buf[2], buf[3]]) as usize;
if buf.len() < HEADER_SIZE + length {
return None;
}
if version != VERSION {
warn!(version, "Unexpected smux version byte — discarding frame");
buf.advance(HEADER_SIZE + length);
continue;
}
buf.advance(1); let cmd = buf.get_u8();
buf.advance(2); let stream_id = buf.get_u32_le();
let payload = buf.split_to(length).freeze();
return Some((cmd, stream_id, payload));
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum StreamCloseReason {
Clean = 0,
SlowConsumer = 1,
}
const REASON_OPEN: u8 = 255;
type StreamEntry = (mpsc::Sender<Bytes>, Arc<AtomicU8>);
struct Inner {
streams: Mutex<HashMap<u32, StreamEntry>>,
frame_tx: mpsc::Sender<Bytes>,
next_id: AtomicU32,
closed: AtomicBool,
die: Notify,
}
impl Inner {
fn close(&self) {
if !self.closed.swap(true, Ordering::SeqCst) {
self.die.notify_waiters();
}
let mut streams = self.streams.lock().unwrap();
for (_, reason_tag) in streams.values() {
reason_tag
.compare_exchange(
REASON_OPEN,
StreamCloseReason::Clean as u8,
Ordering::SeqCst,
Ordering::SeqCst,
)
.ok();
}
streams.clear();
}
fn is_closed(&self) -> bool {
self.closed.load(Ordering::SeqCst)
}
fn route_psh(&self, stream_id: u32, data: Bytes) {
let entry = self
.streams
.lock()
.unwrap()
.get(&stream_id)
.map(|(tx, reason)| (tx.clone(), Arc::clone(reason)));
if entry.is_none() {
debug!(
stream_id,
"smux route_psh: no stream registered for id — dropping frame"
);
}
if let Some((tx, reason_tag)) = entry {
trace!(
stream_id,
bytes = data.len(),
"smux route_psh: routing to stream"
);
match tx.try_send(data) {
Ok(()) => {}
Err(mpsc::error::TrySendError::Full(_)) => {
warn!(
stream_id,
"Per-stream receive buffer full — evicting slow consumer"
);
reason_tag.store(StreamCloseReason::SlowConsumer as u8, Ordering::SeqCst);
self.streams.lock().unwrap().remove(&stream_id);
let _ = self.frame_tx.try_send(encode_ctrl(CMD_FIN, stream_id));
}
Err(mpsc::error::TrySendError::Closed(_)) => {
self.streams.lock().unwrap().remove(&stream_id);
}
}
}
}
fn route_fin(&self, stream_id: u32) {
if let Some((_, reason_tag)) = self.streams.lock().unwrap().remove(&stream_id) {
reason_tag
.compare_exchange(
REASON_OPEN,
StreamCloseReason::Clean as u8,
Ordering::SeqCst,
Ordering::SeqCst,
)
.ok();
}
}
}
async fn recv_task(mut output_rx: mpsc::Receiver<Bytes>, inner: Arc<Inner>) {
let mut buf = BytesMut::new();
loop {
tokio::select! {
biased;
_ = inner.die.notified() => break,
chunk = output_rx.recv() => {
match chunk {
Some(bytes) => {
trace!(len = bytes.len(), "smux recv_task: chunk received");
buf.extend_from_slice(&bytes);
if !dispatch_frames(&mut buf, &inner) {
break; }
}
None => {
debug!("smux recv_task: output_rx closed (no more data)");
break; }
}
}
}
}
dispatch_frames(&mut buf, &inner);
inner.close();
}
fn dispatch_frames(buf: &mut BytesMut, inner: &Inner) -> bool {
loop {
if buf.len() >= HEADER_SIZE {
let version = buf[0];
let length = u16::from_le_bytes([buf[2], buf[3]]) as usize;
if version == VERSION && length > MAX_PAYLOAD {
warn!(
length,
MAX_PAYLOAD,
"smux frame length exceeds MAX_PAYLOAD — protocol violation, tearing down mux"
);
buf.clear();
inner.close();
return false;
}
}
match decode_frame(buf) {
None => return true, Some((cmd, stream_id, data)) => {
trace!(
cmd,
stream_id,
payload_len = data.len(),
"smux dispatch_frames: decoded frame"
);
match cmd {
CMD_PSH => inner.route_psh(stream_id, data),
CMD_FIN => {
debug!(stream_id, "smux dispatch_frames: FIN received");
inner.route_fin(stream_id);
}
CMD_NOP | CMD_SYN => {}
_ => warn!(cmd, stream_id, "Unknown smux command – ignoring"),
}
}
}
}
}
async fn send_task(session: Arc<Session>, mut frame_rx: mpsc::Receiver<Bytes>, inner: Arc<Inner>) {
loop {
let die = inner.die.notified();
tokio::pin!(die);
if inner.is_closed() {
break;
}
tokio::select! {
biased;
_ = &mut die => break,
frame = frame_rx.recv() => {
match frame {
Some(f) => {
let send_die = inner.die.notified();
tokio::pin!(send_die);
tokio::select! {
biased;
_ = &mut send_die => break,
result = session.send(f) => {
if let Err(e) = result {
warn!(error = ?e, "smux send task: session error");
break;
}
}
}
}
None => break,
}
}
}
}
inner.close();
}
async fn keepalive_task(inner: Arc<Inner>) {
let mut ticker = interval(KEEPALIVE_INTERVAL);
ticker.tick().await;
loop {
tokio::select! {
biased;
_ = inner.die.notified() => break,
_ = ticker.tick() => {
if inner.is_closed() {
break;
}
let nop = encode_ctrl(CMD_NOP, 0);
match inner.frame_tx.try_send(nop) {
Ok(()) => {}
Err(mpsc::error::TrySendError::Full(_)) => {}
Err(mpsc::error::TrySendError::Closed(_)) => break,
}
}
}
}
}
#[derive(Debug, Clone, Default)]
pub struct SmuxConfig {
pub keepalive: bool,
}
pub struct SmuxSession {
inner: Arc<Inner>,
}
impl SmuxSession {
pub fn new(session: Arc<Session>, config: SmuxConfig) -> Self {
let (frame_tx, frame_rx) = mpsc::channel::<Bytes>(FRAME_CHANNEL_CAP);
let inner = Arc::new(Inner {
streams: Mutex::new(HashMap::new()),
frame_tx,
next_id: AtomicU32::new(1), closed: AtomicBool::new(false),
die: Notify::new(),
});
let output_rx = session.subscribe_output();
debug!("Subscribed to session output (direct_subs)");
tokio::spawn(recv_task(output_rx, Arc::clone(&inner)));
tokio::spawn(send_task(
Arc::clone(&session),
frame_rx,
Arc::clone(&inner),
));
if config.keepalive {
tokio::spawn(keepalive_task(Arc::clone(&inner)));
}
Self { inner }
}
pub fn open_stream(&self) -> Result<SmuxStream> {
if self.inner.is_closed() {
return Err(Error::InvalidState("smux session is closed".to_string()));
}
let stream_id = self.inner.next_id.fetch_add(2, Ordering::SeqCst);
let (data_tx, data_rx) = mpsc::channel::<Bytes>(STREAM_CHANNEL_CAP);
let close_reason_tag = Arc::new(AtomicU8::new(REASON_OPEN));
self.inner
.streams
.lock()
.unwrap()
.insert(stream_id, (data_tx, Arc::clone(&close_reason_tag)));
match self
.inner
.frame_tx
.try_send(encode_ctrl(CMD_SYN, stream_id))
{
Ok(()) => {}
Err(mpsc::error::TrySendError::Closed(_)) => {
self.inner.streams.lock().unwrap().remove(&stream_id);
return Err(Error::InvalidState("smux send task has exited".to_string()));
}
Err(mpsc::error::TrySendError::Full(_)) => {
self.inner.streams.lock().unwrap().remove(&stream_id);
return Err(Error::InvalidState(
"smux frame queue is full; retry open_stream".to_string(),
));
}
}
debug!(stream_id, "Opened smux stream");
let frame_sink = PollSender::new(self.inner.frame_tx.clone());
Ok(SmuxStream {
stream_id,
inner: Arc::clone(&self.inner),
data_rx,
current_chunk: None,
read_closed: false,
write_closed: false,
close_reason_tag,
frame_sink,
})
}
#[allow(dead_code)]
pub fn is_closed(&self) -> bool {
self.inner.is_closed()
}
}
impl Drop for SmuxSession {
fn drop(&mut self) {
self.inner.close();
}
}
pub struct SmuxStream {
stream_id: u32,
inner: Arc<Inner>,
data_rx: mpsc::Receiver<Bytes>,
current_chunk: Option<Bytes>,
read_closed: bool,
write_closed: bool,
close_reason_tag: Arc<AtomicU8>,
frame_sink: PollSender<Bytes>,
}
impl SmuxStream {
pub fn close_reason(&self) -> Option<StreamCloseReason> {
if !self.read_closed {
return None;
}
match self.close_reason_tag.load(Ordering::SeqCst) {
x if x == StreamCloseReason::SlowConsumer as u8 => {
Some(StreamCloseReason::SlowConsumer)
}
_ => Some(StreamCloseReason::Clean),
}
}
}
impl AsyncRead for SmuxStream {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
let this = self.get_mut();
if this.read_closed {
return Poll::Ready(Ok(())); }
if let Some(ref mut chunk) = this.current_chunk {
let n = chunk.len().min(buf.remaining());
buf.put_slice(&chunk[..n]);
chunk.advance(n);
if chunk.is_empty() {
this.current_chunk = None;
}
return Poll::Ready(Ok(()));
}
match this.data_rx.poll_recv(cx) {
Poll::Ready(Some(mut chunk)) => {
let n = chunk.len().min(buf.remaining());
buf.put_slice(&chunk[..n]);
chunk.advance(n);
if !chunk.is_empty() {
this.current_chunk = Some(chunk);
}
Poll::Ready(Ok(()))
}
Poll::Ready(None) => {
this.read_closed = true;
Poll::Ready(Ok(()))
}
Poll::Pending => {
if this.inner.is_closed() {
this.read_closed = true;
return Poll::Ready(Ok(()));
}
Poll::Pending
}
}
}
}
impl AsyncWrite for SmuxStream {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
let this = self.get_mut();
if this.write_closed {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"stream is write-closed",
)));
}
if this.inner.is_closed() {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"smux session closed",
)));
}
let n = buf.len().min(MAX_PAYLOAD);
match this.frame_sink.poll_reserve(cx) {
Poll::Pending => Poll::Pending,
Poll::Ready(Err(_)) => Poll::Ready(Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"smux session closed",
))),
Poll::Ready(Ok(())) => {
let frame = encode_frame(CMD_PSH, this.stream_id, &buf[..n]);
this.frame_sink.send_item(frame).map_err(|_| {
io::Error::new(io::ErrorKind::BrokenPipe, "smux session closed")
})?;
Poll::Ready(Ok(n))
}
}
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(())) }
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
let this = self.get_mut();
if this.write_closed {
return Poll::Ready(Ok(()));
}
match this.frame_sink.poll_reserve(cx) {
Poll::Pending => return Poll::Pending,
Poll::Ready(Err(_)) => {
this.write_closed = true;
return Poll::Ready(Ok(()));
}
Poll::Ready(Ok(())) => {}
}
this.write_closed = true;
this.frame_sink
.send_item(encode_ctrl(CMD_FIN, this.stream_id))
.ok();
debug!(stream_id = this.stream_id, "Sent FIN");
Poll::Ready(Ok(()))
}
}
impl Drop for SmuxStream {
fn drop(&mut self) {
if !self.write_closed {
self.write_closed = true;
self.inner
.frame_tx
.try_send(encode_ctrl(CMD_FIN, self.stream_id))
.ok();
}
self.inner.streams.lock().unwrap().remove(&self.stream_id);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_encode_decode_frame() {
let payload = b"hello";
let frame = encode_frame(CMD_PSH, 0x0000_0003, payload);
assert_eq!(frame.len(), HEADER_SIZE + payload.len());
assert_eq!(frame[0], VERSION);
assert_eq!(frame[1], CMD_PSH);
assert_eq!(&frame[2..4], &[5u8, 0]);
assert_eq!(&frame[4..8], &[3u8, 0, 0, 0]);
assert_eq!(&frame[8..], payload);
}
#[test]
fn test_roundtrip() {
let original = b"smux test payload";
let stream_id = 7u32;
let encoded = encode_frame(CMD_PSH, stream_id, original);
let mut buf = BytesMut::from(&encoded[..]);
let (cmd, sid, data) = decode_frame(&mut buf).expect("should decode");
assert_eq!(cmd, CMD_PSH);
assert_eq!(sid, stream_id);
assert_eq!(&data[..], original);
assert!(buf.is_empty());
}
#[test]
fn test_partial_frame_returns_none() {
let frame = encode_frame(CMD_PSH, 1, b"data");
let mut buf = BytesMut::from(&frame[..HEADER_SIZE]);
assert!(decode_frame(&mut buf).is_none());
}
#[test]
fn test_ctrl_frame() {
let frame = encode_ctrl(CMD_SYN, 5);
assert_eq!(frame.len(), HEADER_SIZE);
let mut buf = BytesMut::from(&frame[..]);
let (cmd, sid, data) = decode_frame(&mut buf).expect("should decode");
assert_eq!(cmd, CMD_SYN);
assert_eq!(sid, 5);
assert!(data.is_empty());
}
#[test]
fn test_unknown_version_discarded() {
let mut frame = BytesMut::from(&encode_frame(CMD_PSH, 3, b"payload")[..]);
frame[0] = 0x02;
let mut buf = frame;
assert!(
decode_frame(&mut buf).is_none(),
"Unknown version must be discarded"
);
assert!(buf.is_empty(), "Buffer must be drained after discard");
}
}