use anyhow::{Context, Result};
use tokio::io::{copy_bidirectional_with_sizes, AsyncRead, AsyncWrite};
use tracing::info;
use crate::metrics_helper::MetricsCounter;
pub async fn forward_bidirectional<StreamA, StreamB, StreamName, CounterA, CounterB>(
a: &mut StreamA,
b: &mut StreamB,
id: StreamName,
stream_a_counter: CounterA,
stream_b_counter: CounterB,
) -> Result<()>
where
StreamA: AsyncRead + AsyncWrite + Unpin,
StreamB: AsyncRead + AsyncWrite + Unpin,
StreamName: std::fmt::Display,
CounterA: MetricsCounter,
CounterB: MetricsCounter,
{
let buf_size = crate::PortRedirectProtocol::QUIC_STREAM_READ_BUFFER_SIZE;
let result = copy_bidirectional_with_sizes(a, b, buf_size, buf_size).await;
match result {
Ok((bytes_a, bytes_b)) => {
stream_a_counter.inc_by(bytes_a);
stream_b_counter.inc_by(bytes_b);
info!(
"Stream (id={}): forwarded (A:B) ({}:{}) bytes",
id, bytes_a, bytes_b
);
Ok(())
}
Err(err) => {
if err.to_string().contains("error 0") {
info!("Stream (id={}): graceful shutdown detected: {}", id, err);
Ok(())
} else {
Err(err).context("Bidirectional copy failed")
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::metrics_helper::DummyCounter;
use anyhow::Result;
use std::pin::Pin;
use std::task::{Context, Poll};
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
#[derive(Debug)]
struct TestStream {
read_data: Vec<u8>,
write_data: Vec<u8>,
pos: usize,
}
impl TestStream {
fn new(initial_data: &[u8]) -> Self {
Self {
read_data: initial_data.to_vec(),
write_data: Vec::new(),
pos: 0,
}
}
}
impl AsyncRead for TestStream {
fn poll_read(
mut self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<Result<(), std::io::Error>> {
let remaining = &self.read_data[self.pos..];
if remaining.is_empty() {
return Poll::Ready(Ok(()));
}
let to_copy = std::cmp::min(remaining.len(), buf.remaining());
buf.put_slice(&remaining[..to_copy]);
self.pos += to_copy;
Poll::Ready(Ok(()))
}
}
impl AsyncWrite for TestStream {
fn poll_write(
mut self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<Result<usize, std::io::Error>> {
self.write_data.extend_from_slice(buf);
Poll::Ready(Ok(buf.len()))
}
fn poll_flush(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Result<(), std::io::Error>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Result<(), std::io::Error>> {
Poll::Ready(Ok(()))
}
}
struct FailingStream;
impl AsyncRead for FailingStream {
fn poll_read(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
_buf: &mut ReadBuf<'_>,
) -> Poll<Result<(), std::io::Error>> {
Poll::Ready(Err(std::io::Error::new(
std::io::ErrorKind::Other,
"read failure",
)))
}
}
impl AsyncWrite for FailingStream {
fn poll_write(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
_buf: &[u8],
) -> Poll<Result<usize, std::io::Error>> {
Poll::Ready(Err(std::io::Error::new(
std::io::ErrorKind::Other,
"write failure",
)))
}
fn poll_flush(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Result<(), std::io::Error>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Result<(), std::io::Error>> {
Poll::Ready(Ok(()))
}
}
#[tokio::test]
async fn test_forward_bidirectional_success() -> Result<()> {
let mut stream_a = TestStream::new(b"hello");
let mut stream_b = TestStream::new(b"world");
let dummy_counter_a = DummyCounter::new();
let dummy_counter_b = DummyCounter::new();
forward_bidirectional(
&mut stream_a,
&mut stream_b,
"A",
&dummy_counter_a,
&dummy_counter_b,
)
.await?;
assert_eq!(stream_a.write_data, b"world");
assert_eq!(stream_b.write_data, b"hello");
Ok(())
}
#[tokio::test]
async fn test_forward_bidirectional_failure() {
let mut normal_stream = TestStream::new(b"data");
let mut failing_stream = FailingStream;
let dummy_counter_a = DummyCounter::new();
let dummy_counter_b = DummyCounter::new();
let result = forward_bidirectional(
&mut failing_stream,
&mut normal_stream,
"fail",
&dummy_counter_a,
&dummy_counter_b,
)
.await;
assert!(result.is_err());
let err_msg = format!("{:?}", result.err().unwrap());
assert!(
err_msg.contains("Bidirectional copy failed"),
"Error message did not contain expected context, got: {}",
err_msg
);
}
struct GracefulStream;
impl AsyncRead for GracefulStream {
fn poll_read(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
_buf: &mut ReadBuf<'_>,
) -> Poll<Result<(), std::io::Error>> {
Poll::Ready(Err(std::io::Error::new(
std::io::ErrorKind::Other,
"sending stopped by peer: error 0",
)))
}
}
impl AsyncWrite for GracefulStream {
fn poll_write(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
_buf: &[u8],
) -> Poll<Result<usize, std::io::Error>> {
Poll::Ready(Err(std::io::Error::new(
std::io::ErrorKind::Other,
"sending stopped by peer: error 0",
)))
}
fn poll_flush(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Result<(), std::io::Error>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Result<(), std::io::Error>> {
Poll::Ready(Ok(()))
}
}
#[tokio::test]
async fn test_forward_bidirectional_graceful_shutdown() -> Result<()> {
let mut normal_stream = TestStream::new(b"normal");
let mut graceful_stream = GracefulStream;
let dummy_counter_a = DummyCounter::new();
let dummy_counter_b = DummyCounter::new();
forward_bidirectional(
&mut normal_stream,
&mut graceful_stream,
"normal",
&dummy_counter_a,
&dummy_counter_b,
)
.await?;
Ok(())
}
}