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, info, info_span, 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!(error = %e, "set_nodelay failed");
}
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!("stream closed"),
(Err(e), _) => debug!(direction = "tx", error = %e, "stream copy error"),
(_, Err(e)) => debug!(direction = "rx", error = %e, "stream copy error"),
}
Ok(())
}
enum CopyOutcome {
Eof,
Errored(std::io::Error),
Aborted,
#[allow(dead_code)]
AbortedWith(std::io::Error),
}
pub async fn tunnel_stdio_client(quic_connection: Connection, tunnel_id: u64) -> Result<()> {
use crate::common::tunnel::send_open_conn;
debug!("opening stdio tunnel");
let (mut send, mut recv) = quic_connection.open_bi().await?;
send_open_conn(
&OpenConn {
tunnel_id,
dynamic: None,
},
&mut send,
&mut recv,
)
.await?;
info!("stdio tunnel ready");
let mut stdin = BufReader::with_capacity(TUNNEL_COPY_BUF, tokio::io::stdin());
let mut stdout = tokio::io::stdout();
let mut quic_recv = BufReader::with_capacity(TUNNEL_COPY_BUF, recv);
let abort: Arc<Notify> = Arc::new(Notify::new());
const ABORT_CODE: u32 = 0;
let stdin_to_quic = {
let abort = abort.clone();
async move {
let outcome = tokio::select! {
biased;
_ = abort.notified() => CopyOutcome::Aborted,
r = tokio::io::copy_buf(&mut stdin, &mut send) => match r {
Ok(_) => CopyOutcome::Eof,
Err(e) => CopyOutcome::Errored(e),
},
};
match outcome {
CopyOutcome::Eof => {
let _ = send.shutdown().await;
Ok(())
}
CopyOutcome::Errored(e) | CopyOutcome::AbortedWith(e) => {
let _ = send.reset(VarInt::from_u32(ABORT_CODE));
abort.notify_waiters();
Err(anyhow::Error::from(e))
}
CopyOutcome::Aborted => {
let _ = send.reset(VarInt::from_u32(ABORT_CODE));
Ok(())
}
}
}
};
let quic_to_stdout = {
let abort = abort.clone();
async move {
let outcome = tokio::select! {
biased;
_ = abort.notified() => CopyOutcome::Aborted,
r = tokio::io::copy_buf(&mut quic_recv, &mut stdout) => match r {
Ok(_) => CopyOutcome::Eof,
Err(e) => CopyOutcome::Errored(e),
},
};
let _ = stdout.flush().await;
match outcome {
CopyOutcome::Eof => Ok(()),
CopyOutcome::Errored(e) | CopyOutcome::AbortedWith(e) => {
abort.notify_waiters();
Err(anyhow::Error::from(e))
}
CopyOutcome::Aborted => Ok(()),
}
}
};
let (s2q, q2s) = tokio::join!(stdin_to_quic, quic_to_stdout);
match (&s2q, &q2s) {
(Ok(_), Ok(_)) => debug!("stdio tunnel closed"),
(Err(e), _) => debug!(direction = "tx", error = %e, "stdio copy error"),
(_, Err(e)) => debug!(direction = "rx", error = %e, "stdio copy error"),
}
Ok(())
}
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!(addr = %local_addr, "listening");
let local_counter = AtomicUsize::new(0);
loop {
let (local_socket, peer) = listener.accept().await?;
let connection = quic_connection.clone();
let handle = handle.clone();
let local_id = (local_counter.fetch_add(1, Ordering::Relaxed) + 1) as u64;
tokio::spawn(async move {
let (mut send, mut recv) = match connection.open_bi().await {
Ok(s) => s,
Err(e) => {
let span = info_span!("conn", conn_id = local_id, tunnel_id, peer = %peer);
let _g = span.enter();
debug!(error = %e, "open_bi failed");
return Err::<(), anyhow::Error>(e.into());
}
};
if let Err(e) = send_open_conn(
&OpenConn {
tunnel_id,
dynamic: None,
},
&mut send,
&mut recv,
)
.await
{
let span = info_span!("conn", conn_id = local_id, tunnel_id, peer = %peer);
let _g = span.enter();
debug!(error = %e, "OpenConn rejected");
return Err(e);
}
let conn_guard = handle.as_ref().map(|h| h.open_conn(Some(peer.to_string())));
let conn_id = conn_guard.as_ref().map(|g| g.id()).unwrap_or(local_id);
let counters = conn_guard.as_ref().map(|g| g.counters());
let span = info_span!("conn", conn_id, tunnel_id, peer = %peer);
async move {
info!("conn opened");
let started = std::time::Instant::now();
let result = tunnel_tcp_stream(local_socket, send, recv, counters.clone()).await;
let dur_ms = started.elapsed().as_millis() as u64;
let snap = counters.as_ref().map(|c| c.snapshot());
match (&result, snap) {
(Ok(()), Some((bytes_in, bytes_out))) => {
info!(bytes_in, bytes_out, dur_ms, "conn closed")
}
(Ok(()), None) => info!(dur_ms, "conn closed"),
(Err(e), Some((bytes_in, bytes_out))) => {
debug!(bytes_in, bytes_out, dur_ms, error = %e, "conn closed (error)")
}
(Err(e), None) => debug!(dur_ms, error = %e, "conn closed (error)"),
}
drop(conn_guard);
result
}
.instrument(span)
.await
});
}
}
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!(target = %remote_addr, "dialing");
let tcp_stream = TcpStream::connect(&remote_addr).await?;
debug!(target = %remote_addr, "dialed");
tunnel_tcp_stream(tcp_stream, send_channel, recv_channel, counters).await?;
Ok(())
}