use std::io;
use std::pin::Pin;
use std::task::{Context, Poll};
use fastwebsockets::{Frame, OpCode, Payload, Role, WebSocket};
use h2ts_server::{bridge, bridge_with, BridgeConfig, CloseFrame};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, ReadBuf};
#[tokio::test]
async fn bridge_forwards_bytes_both_directions() {
let (client_io, server_io) = tokio::io::duplex(16 * 1024);
let (peer_for_bridge, mut peer_test) = tokio::io::duplex(16 * 1024);
tokio::spawn(async move {
let _ = bridge(server_io, peer_for_bridge).await;
});
let mut client_ws = WebSocket::after_handshake(client_io, Role::Client);
client_ws
.write_frame(Frame::binary(Payload::Owned(b"hello".to_vec())))
.await
.unwrap();
let mut buf = [0u8; 5];
peer_test.read_exact(&mut buf).await.unwrap();
assert_eq!(&buf, b"hello");
peer_test.write_all(b"world").await.unwrap();
let frame = client_ws.read_frame().await.unwrap();
assert_eq!(frame.payload.to_vec(), b"world".to_vec());
}
#[tokio::test]
async fn bridge_streams_a_large_payload() {
let (client_io, server_io) = tokio::io::duplex(16 * 1024);
let (peer_for_bridge, mut peer_test) = tokio::io::duplex(16 * 1024);
tokio::spawn(async move {
let _ = bridge(server_io, peer_for_bridge).await;
});
let payload: Vec<u8> = (0..512 * 1024).map(|i| (i % 251) as u8).collect();
let expected = payload.clone();
tokio::spawn(async move {
peer_test.write_all(&payload).await.unwrap();
});
let mut client_ws = WebSocket::after_handshake(client_io, Role::Client);
let mut got = Vec::new();
while got.len() < expected.len() {
let frame = client_ws.read_frame().await.unwrap();
got.extend_from_slice(&frame.payload);
}
assert_eq!(got, expected);
}
#[tokio::test]
async fn bridge_reassembles_a_large_inbound_frame() {
let (client_io, server_io) = tokio::io::duplex(16 * 1024);
let (peer_for_bridge, mut peer_test) = tokio::io::duplex(16 * 1024);
tokio::spawn(async move {
let _ = bridge(server_io, peer_for_bridge).await;
});
let payload: Vec<u8> = (0..256 * 1024).map(|i| (i % 251) as u8).collect();
let expected = payload.clone();
let mut client_ws = WebSocket::after_handshake(client_io, Role::Client);
tokio::spawn(async move {
client_ws
.write_frame(Frame::binary(Payload::Owned(payload)))
.await
.unwrap();
});
let mut got = vec![0u8; expected.len()];
peer_test.read_exact(&mut got).await.unwrap();
assert_eq!(
got, expected,
"every byte of the large frame must reach the peer, in order"
);
}
#[tokio::test]
async fn bridge_handles_many_frames_in_one_read() {
let (client_io, server_io) = tokio::io::duplex(64 * 1024);
let (peer_for_bridge, mut peer_test) = tokio::io::duplex(64 * 1024);
tokio::spawn(async move {
let _ = bridge(server_io, peer_for_bridge).await;
});
let mut expected = Vec::new();
let mut client_ws = WebSocket::after_handshake(client_io, Role::Client);
for i in 0..64u32 {
let msg = format!("frame-{i:04}-");
expected.extend_from_slice(msg.as_bytes());
client_ws
.write_frame(Frame::binary(Payload::Owned(msg.into_bytes())))
.await
.unwrap();
}
let mut got = vec![0u8; expected.len()];
peer_test.read_exact(&mut got).await.unwrap();
assert_eq!(got, expected, "all frames forwarded once, in order");
}
#[tokio::test]
async fn bridge_answers_ping_with_pong() {
let (client_io, server_io) = tokio::io::duplex(16 * 1024);
let (peer_for_bridge, _peer_test) = tokio::io::duplex(16 * 1024);
tokio::spawn(async move {
let _ = bridge(server_io, peer_for_bridge).await;
});
let mut client_ws = WebSocket::after_handshake(client_io, Role::Client);
client_ws
.write_frame(Frame::new(
true,
OpCode::Ping,
None,
Payload::Owned(b"ping-payload".to_vec()),
))
.await
.unwrap();
let frame = client_ws.read_frame().await.unwrap();
assert_eq!(frame.opcode, OpCode::Pong);
assert_eq!(frame.payload.to_vec(), b"ping-payload".to_vec());
}
#[tokio::test]
async fn bridge_propagates_client_close_to_peer() {
let (client_io, server_io) = tokio::io::duplex(16 * 1024);
let (peer_for_bridge, mut peer_test) = tokio::io::duplex(16 * 1024);
tokio::spawn(async move {
let _ = bridge(server_io, peer_for_bridge).await;
});
let mut client_ws = WebSocket::after_handshake(client_io, Role::Client);
client_ws
.write_frame(Frame::close(1000, b""))
.await
.unwrap();
let mut buf = [0u8; 16];
let n = peer_test.read(&mut buf).await.unwrap();
assert_eq!(n, 0, "peer should observe EOF after the client close");
}
#[tokio::test]
async fn bridge_reports_1006_on_transport_drop_without_close() {
let (client_io, server_io) = tokio::io::duplex(16 * 1024);
let (peer_for_bridge, _peer_test) = tokio::io::duplex(16 * 1024);
let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel::<CloseFrame>();
let config = BridgeConfig {
on_close: Some(Box::new(move |cf: &CloseFrame| {
let _ = tx.send(cf.clone());
})),
..Default::default()
};
tokio::spawn(async move {
let _ = bridge_with(server_io, peer_for_bridge, config).await;
});
drop(client_io);
let got = rx
.recv()
.await
.expect("on_close should fire on abnormal end");
assert_eq!(got.code, 1006, "abnormal closure code");
assert!(got.reason.is_empty(), "no reason on an abnormal close");
}
struct FailingPeer;
impl AsyncRead for FailingPeer {
fn poll_read(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
_buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
Poll::Pending
}
}
impl AsyncWrite for FailingPeer {
fn poll_write(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
_buf: &[u8],
) -> Poll<io::Result<usize>> {
Poll::Ready(Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"peer write failed",
)))
}
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<()>> {
Poll::Ready(Ok(()))
}
}
#[tokio::test]
async fn bridge_reports_error_close_on_peer_write_failure() {
let (client_io, server_io) = tokio::io::duplex(16 * 1024);
let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel::<CloseFrame>();
let config = BridgeConfig {
error_close: CloseFrame {
code: 1014,
reason: "bad gateway".to_string(),
},
on_close: Some(Box::new(move |cf: &CloseFrame| {
let _ = tx.send(cf.clone());
})),
..Default::default()
};
tokio::spawn(async move {
let _ = bridge_with(server_io, FailingPeer, config).await;
});
let mut client_ws = WebSocket::after_handshake(client_io, Role::Client);
client_ws
.write_frame(Frame::binary(Payload::Owned(b"trigger".to_vec())))
.await
.unwrap();
let got = rx
.recv()
.await
.expect("on_close should fire on write failure");
assert_eq!(got.code, 1014, "a write failure uses error_close, not a bare 1006");
assert_eq!(got.reason, "bad gateway");
}
struct ReadErrPeer;
impl AsyncRead for ReadErrPeer {
fn poll_read(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
_buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
Poll::Ready(Err(io::Error::new(
io::ErrorKind::ConnectionReset,
"upstream reset",
)))
}
}
impl AsyncWrite for ReadErrPeer {
fn poll_write(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
Poll::Ready(Ok(buf.len()))
}
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<()>> {
Poll::Ready(Ok(()))
}
}
#[tokio::test]
async fn bridge_reports_error_close_on_peer_read_error() {
let (client_io, server_io) = tokio::io::duplex(16 * 1024);
let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel::<CloseFrame>();
let config = BridgeConfig {
error_close: CloseFrame {
code: 1014,
reason: "bad gateway".to_string(),
},
keepalive: None,
on_close: Some(Box::new(move |cf: &CloseFrame| {
let _ = tx.send(cf.clone());
})),
..Default::default()
};
tokio::spawn(async move {
let _ = bridge_with(server_io, ReadErrPeer, config).await;
});
let _client_ws = WebSocket::after_handshake(client_io, Role::Client);
let got = rx.recv().await.expect("on_close should fire on read error");
assert_eq!(got.code, 1014, "a read error uses error_close, not the clean close");
assert_eq!(got.reason, "bad gateway");
}