use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use anyhow::Result;
use quinn::{Connection, RecvStream, SendStream, VarInt};
use tokio::{
io::{AsyncRead, AsyncWriteExt, BufReader},
net::{TcpListener, TcpStream},
sync::Notify,
};
use tracing::{debug, debug_span, info, Instrument};
use crate::common::counted::{CountedReader, TunnelCounters};
use crate::common::remote::OpenConn;
use crate::common::tunnel::send_open_conn;
use crate::server::state::TunnelHandle;
use super::remote::RemoteRequest;
const TUNNEL_COPY_BUF: usize = 256 * 1024;
pub type Counters = Option<Arc<TunnelCounters>>;
pub type TunnelHandleOpt = Option<Arc<TunnelHandle>>;
fn count_in<R: AsyncRead + Unpin + Send + 'static>(
inner: R,
counters: &Counters,
) -> Box<dyn AsyncRead + Unpin + Send> {
match counters {
Some(c) => Box::new(CountedReader::new(inner, c.in_handle())),
None => Box::new(inner),
}
}
fn count_out<R: AsyncRead + Unpin + Send + 'static>(
inner: R,
counters: &Counters,
) -> Box<dyn AsyncRead + Unpin + Send> {
match counters {
Some(c) => Box::new(CountedReader::new(inner, c.out_handle())),
None => Box::new(inner),
}
}
pub async fn tunnel_tcp_stream(
tcp_stream: TcpStream,
mut send_channel: SendStream,
recv_channel: RecvStream,
counters: Counters,
) -> Result<()> {
if let Err(e) = tcp_stream.set_nodelay(true) {
debug!("set_nodelay failed: {}", e);
}
let (tcp_recv, mut tcp_send) = tcp_stream.into_split();
let tcp_recv = count_out(tcp_recv, &counters);
let quic_recv = count_in(recv_channel, &counters);
let mut tcp_recv = BufReader::with_capacity(TUNNEL_COPY_BUF, tcp_recv);
let mut quic_recv = BufReader::with_capacity(TUNNEL_COPY_BUF, quic_recv);
let abort: Arc<Notify> = Arc::new(Notify::new());
const ABORT_CODE: u32 = 0;
let client_to_server = {
let abort = abort.clone();
async move {
let outcome = tokio::select! {
biased;
_ = abort.notified() => CopyOutcome::Aborted,
r = tokio::io::copy_buf(&mut tcp_recv, &mut send_channel) => match r {
Ok(_) => CopyOutcome::Eof,
Err(e) => CopyOutcome::Errored(e),
},
};
match outcome {
CopyOutcome::Eof => {
let _ = send_channel.shutdown().await;
Ok(())
}
CopyOutcome::Errored(e) | CopyOutcome::AbortedWith(e) => {
let _ = send_channel.reset(VarInt::from_u32(ABORT_CODE));
abort.notify_waiters();
Err(anyhow::Error::from(e))
}
CopyOutcome::Aborted => {
let _ = send_channel.reset(VarInt::from_u32(ABORT_CODE));
Ok(())
}
}
}
};
let server_to_client = {
let abort = abort.clone();
async move {
let outcome = tokio::select! {
biased;
_ = abort.notified() => CopyOutcome::Aborted,
r = tokio::io::copy_buf(&mut quic_recv, &mut tcp_send) => match r {
Ok(_) => CopyOutcome::Eof,
Err(e) => CopyOutcome::Errored(e),
},
};
match outcome {
CopyOutcome::Eof => {
let _ = tcp_send.shutdown().await;
Ok(())
}
CopyOutcome::Errored(e) | CopyOutcome::AbortedWith(e) => {
let _ = tcp_send.shutdown().await;
abort.notify_waiters();
Err(anyhow::Error::from(e))
}
CopyOutcome::Aborted => {
let _ = tcp_send.shutdown().await;
Ok(())
}
}
}
};
let (c2s, s2c) = tokio::join!(client_to_server, server_to_client);
match (&c2s, &s2c) {
(Ok(_), Ok(_)) => debug!("closed"),
(Err(e), _) => debug!("client→server error: {}", e),
(_, Err(e)) => debug!("server→client error: {}", e),
}
Ok(())
}
enum CopyOutcome {
Eof,
Errored(std::io::Error),
Aborted,
#[allow(dead_code)]
AbortedWith(std::io::Error),
}
pub async fn tunnel_tcp_client(
quic_connection: Connection,
remote: RemoteRequest,
handle: TunnelHandleOpt,
tunnel_id: u64,
) -> Result<()> {
let local_addr = remote.local_socket_addr();
let listener = TcpListener::bind(local_addr).await?;
info!("listening on {}", local_addr);
let conn_counter = AtomicUsize::new(0);
loop {
let (local_socket, addr) = listener.accept().await?;
let conn_id = conn_counter.fetch_add(1, Ordering::Relaxed) + 1;
let span = debug_span!("conn", id = conn_id, peer = %addr);
let connection = quic_connection.clone();
let handle = handle.clone();
tokio::spawn(
async move {
debug!("open");
let (mut send, mut recv) = connection.open_bi().await?;
send_open_conn(
&OpenConn {
tunnel_id,
dynamic: None,
},
&mut send,
&mut recv,
)
.await?;
let _conn_guard = handle.as_ref().map(|h| h.open_conn(Some(addr.to_string())));
if let Some(g) = _conn_guard.as_ref() {
info!(conn_id = g.id(), peer = %addr, "conn opened");
}
let counters = _conn_guard.as_ref().map(|g| g.counters());
tunnel_tcp_stream(local_socket, send, recv, counters).await?;
Ok::<(), anyhow::Error>(())
}
.instrument(span),
);
}
}
pub async fn tunnel_tcp_server(
recv_channel: RecvStream,
send_channel: SendStream,
request: RemoteRequest,
counters: Counters,
) -> Result<()> {
let remote_addr = request
.remote_addr_string()
.ok_or_else(|| anyhow::anyhow!("TCP server tunnel requires a host:port remote"))?;
debug!("connecting to {}", remote_addr);
let tcp_stream = TcpStream::connect(&remote_addr).await?;
debug!("connected to {}", remote_addr);
tunnel_tcp_stream(tcp_stream, send_channel, recv_channel, counters).await?;
Ok(())
}