use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::Arc;
use std::time::Duration;
use anyhow::{anyhow, Result};
use futures::stream::{FuturesUnordered, StreamExt};
use quinn::{Connection, Endpoint, VarInt};
use tokio::sync::broadcast;
use tokio::{signal, task};
use tracing::{debug, error, info, info_span, warn, Instrument};
use crate::common::quic::{
client_server_name, create_client_endpoint, create_client_endpoint_via_proxy,
};
use crate::common::remote::{
Direction, DynamicTarget, OpenConnResponse, RemoteKind, RemoteRequest, SessionHello,
};
use crate::common::socks::tunnel_socks_client;
use crate::common::tcp::{tunnel_stdio_client, tunnel_tcp_client, tunnel_tcp_server};
use crate::common::tunnel::{client_send_session_hello, receive_open_conn, reply_open_conn};
use crate::common::udp::{tunnel_udp_client, tunnel_udp_server};
use crate::{ClientConfig, ReconnectConfig};
pub async fn run_async(config: ClientConfig) -> Result<()> {
let mut endpoints = match &config.proxy {
None => Some(EndpointPool::new(&config)?),
Some(p) => {
info!(proxy = %p, "routing QUIC through SOCKS5 proxy");
None
}
};
let server_name = client_server_name(&config.tls, &config.server.host);
let (shutdown_tx, _) = broadcast::channel::<()>(1);
let shutdown_tx_clone = shutdown_tx.clone();
tokio::spawn(async move {
if signal::ctrl_c().await.is_ok() {
info!("shutdown signal received");
let _ = shutdown_tx_clone.send(());
}
});
let result = run_with_reconnect(endpoints.as_mut(), &server_name, &config, &shutdown_tx).await;
if let Some(pool) = endpoints.as_ref() {
pool.wait_idle().await;
}
debug!("client run loop exited");
result
}
struct EndpointPool<'a> {
config: &'a ClientConfig,
v4: Option<Endpoint>,
v6: Option<Endpoint>,
}
impl<'a> EndpointPool<'a> {
fn new(config: &'a ClientConfig) -> Result<Self> {
let primary = config.server.primary();
let endpoint = create_client_endpoint(&config.tls, config.congestion, primary)?;
let mut pool = Self {
config,
v4: None,
v6: None,
};
if primary.is_ipv6() {
pool.v6 = Some(endpoint);
} else {
pool.v4 = Some(endpoint);
}
Ok(pool)
}
fn get_for(&mut self, addr: SocketAddr) -> Result<&Endpoint> {
let slot = if addr.is_ipv6() {
&mut self.v6
} else {
&mut self.v4
};
if slot.is_none() {
*slot = Some(create_client_endpoint(
&self.config.tls,
self.config.congestion,
addr,
)?);
}
Ok(slot.as_ref().expect("endpoint just inserted"))
}
async fn wait_idle(&self) {
if let Some(e) = &self.v4 {
e.wait_idle().await;
}
if let Some(e) = &self.v6 {
e.wait_idle().await;
}
}
}
const HAPPY_EYEBALLS_DELAY: Duration = Duration::from_millis(250);
async fn run_with_reconnect(
endpoints: Option<&mut EndpointPool<'_>>,
server_name: &str,
config: &ClientConfig,
shutdown_tx: &broadcast::Sender<()>,
) -> Result<()> {
let ReconnectConfig {
max_retries,
initial_backoff,
max_backoff,
} = config.reconnect.clone();
let mut backoff = initial_backoff;
let mut attempt: u32 = 0;
let mut endpoints = endpoints;
loop {
let mut shutdown_rx = shutdown_tx.subscribe();
info!(server = %config.server, sni = %server_name, "connecting");
let connect_outcome = tokio::select! {
res = async {
if let Some(proxy) = &config.proxy {
proxied_connect(proxy, config, server_name).await
} else {
let pool = endpoints
.as_deref_mut()
.expect("endpoint pool is built when no proxy is configured");
happy_eyeballs_connect(pool, &config.server.addrs, server_name).await
}
} => Some(res),
_ = shutdown_rx.recv() => return Ok(()),
};
match connect_outcome {
Some(Ok(connection)) => {
let peer = connection.remote_address();
info!(peer = %peer, "connected");
attempt = 0;
backoff = initial_backoff;
let session_span = info_span!("session", peer = %peer);
let outcome = run_connection(connection, config, shutdown_tx)
.instrument(session_span)
.await;
match outcome {
SessionOutcome::Shutdown => return Ok(()),
SessionOutcome::Disconnected(reason) => {
warn!(reason = %reason, "connection lost");
}
}
}
Some(Err(e)) => {
warn!(error = %e, "connect attempt failed");
}
None => unreachable!(),
}
attempt = attempt.saturating_add(1);
if let Some(max) = max_retries {
if attempt > max {
return Err(anyhow!(
"giving up after {} reconnect attempt(s)",
attempt - 1
));
}
}
let attempt_label = match max_retries {
Some(m) => format!("{attempt}/{m}"),
None => attempt.to_string(),
};
info!(
backoff_ms = backoff.as_millis() as u64,
attempt = %attempt_label,
"reconnecting"
);
tokio::select! {
_ = tokio::time::sleep(backoff) => {}
_ = shutdown_rx.recv() => return Ok(()),
}
backoff = next_backoff(backoff, max_backoff);
}
}
async fn happy_eyeballs_connect(
endpoints: &mut EndpointPool<'_>,
addrs: &[SocketAddr],
server_name: &str,
) -> Result<Connection> {
if addrs.is_empty() {
return Err(anyhow!("no candidate addresses to connect to"));
}
let mut races = FuturesUnordered::new();
let mut last_error: Option<String> = None;
for (idx, addr) in addrs.iter().enumerate() {
let endpoint = match endpoints.get_for(*addr) {
Ok(e) => e.clone(),
Err(e) => {
last_error = Some(format!("{addr}: failed to build endpoint: {e}"));
continue;
}
};
let connecting = match endpoint.connect(*addr, server_name) {
Ok(c) => c,
Err(e) => {
last_error = Some(format!("{addr}: {e}"));
continue;
}
};
let stagger = HAPPY_EYEBALLS_DELAY * (idx as u32);
let addr = *addr;
races.push(async move {
if !stagger.is_zero() {
tokio::time::sleep(stagger).await;
}
(addr, connecting.await)
});
}
while let Some((addr, res)) = races.next().await {
match res {
Ok(conn) => {
debug!(addr = %addr, "happy eyeballs winner");
return Ok(conn);
}
Err(e) => {
debug!(addr = %addr, error = %e, "happy eyeballs candidate failed");
last_error = Some(format!("{addr}: {e}"));
}
}
}
Err(anyhow!(
"all candidate addresses failed (last error: {})",
last_error.unwrap_or_else(|| "<none>".into())
))
}
async fn proxied_connect(
proxy: &crate::common::proxy::ProxyConfig,
config: &ClientConfig,
server_name: &str,
) -> Result<Connection> {
let server = config
.server
.addrs
.first()
.copied()
.ok_or_else(|| anyhow!("no candidate addresses to connect to"))?;
debug!(server = %server, proxy = %proxy, "opening SOCKS5 UDP ASSOCIATE for QUIC");
let endpoint =
create_client_endpoint_via_proxy(&config.tls, config.congestion, server, proxy).await?;
let connection = endpoint.connect(server, server_name)?.await?;
Ok(connection)
}
fn next_backoff(current: Duration, max_backoff: Duration) -> Duration {
current.saturating_mul(2).min(max_backoff)
}
enum SessionOutcome {
Shutdown,
Disconnected(String),
}
async fn run_connection(
connection: Connection,
config: &ClientConfig,
shutdown_tx: &broadcast::Sender<()>,
) -> SessionOutcome {
let tunnel_ids = match send_session_hello(&connection, &config.remotes).await {
Ok(ids) => ids,
Err(e) => {
return SessionOutcome::Disconnected(format!("session hello failed: {e}"));
}
};
info!(count = tunnel_ids.len(), "session established");
for (remote, tunnel_id) in config.remotes.iter().zip(tunnel_ids.iter().copied()) {
let dir = if matches!(remote.direction, Direction::Reverse) {
"reverse"
} else {
"forward"
};
info!(tunnel_id, dir, spec = %remote, "tunnel registered");
}
let remotes_by_id: Arc<HashMap<u64, RemoteRequest>> = Arc::new(
tunnel_ids
.iter()
.copied()
.zip(config.remotes.iter().cloned())
.collect(),
);
let mut tasks = Vec::new();
for (remote, tunnel_id) in config.remotes.iter().zip(tunnel_ids.iter().copied()) {
if matches!(remote.direction, Direction::Reverse) {
continue;
}
let remote = remote.clone();
let connection_clone = connection.clone();
let span = info_span!("tunnel", tunnel_id, dir = "forward", spec = %remote);
let shutdown_for_task = remote.is_stdio().then(|| shutdown_tx.clone());
let task = task::spawn(
async move {
if let Err(e) = handle_forward_tunnel(connection_clone, remote, tunnel_id).await {
error!(error = %e, "forward tunnel failed");
}
if let Some(tx) = shutdown_for_task {
let _ = tx.send(());
}
anyhow::Ok(())
}
.instrument(span),
);
tasks.push(task);
}
let connection_clone = connection.clone();
let remotes_for_accept = remotes_by_id.clone();
let accept_reverse_task = tokio::spawn(async move {
loop {
let quic_connection = connection_clone.clone();
let remotes = remotes_for_accept.clone();
if let Err(e) = client_accept_reverse_conn(quic_connection, remotes).await {
debug!(error = %e, "reverse-accept loop ended");
break;
}
}
anyhow::Ok(())
});
tasks.push(accept_reverse_task);
let mut shutdown_rx = shutdown_tx.subscribe();
let outcome = tokio::select! {
_ = shutdown_rx.recv() => {
info!("disconnecting and notifying server");
connection.close(VarInt::from_u32(130), b"client received ^C");
let _ = tokio::time::timeout(
Duration::from_millis(500),
connection.closed(),
)
.await;
SessionOutcome::Shutdown
}
reason = connection.closed() => {
SessionOutcome::Disconnected(reason.to_string())
}
};
for handle in tasks {
handle.abort();
}
outcome
}
async fn send_session_hello(
quic_connection: &Connection,
remotes: &[RemoteRequest],
) -> Result<Vec<u64>> {
let (mut send, mut recv) = quic_connection.open_bi().await?;
let hello = SessionHello {
remotes: remotes.to_vec(),
};
client_send_session_hello(&hello, &mut send, &mut recv).await
}
async fn handle_forward_tunnel(
quic_connection: Connection,
remote: RemoteRequest,
tunnel_id: u64,
) -> Result<()> {
if remote.is_stdio() {
return tunnel_stdio_client(quic_connection, tunnel_id).await;
}
match &remote.kind {
RemoteKind::Socks5 { .. } => {
tunnel_socks_client(quic_connection, remote, None, tunnel_id).await?
}
RemoteKind::Tcp { .. } => {
tunnel_tcp_client(quic_connection, remote, None, tunnel_id).await?
}
RemoteKind::Udp { .. } => {
tunnel_udp_client(quic_connection, remote, None, tunnel_id).await?
}
}
Ok(())
}
async fn client_accept_reverse_conn(
quic_connection: Connection,
remotes_by_id: Arc<HashMap<u64, RemoteRequest>>,
) -> Result<()> {
let (mut send, mut recv) = quic_connection.accept_bi().await?;
tokio::spawn(async move {
let open = match receive_open_conn(&mut recv).await {
Ok(o) => o,
Err(e) => {
error!(error = %e, "failed to read OpenConn frame");
return;
}
};
let parent = match remotes_by_id.get(&open.tunnel_id) {
Some(p) => p.clone(),
None => {
let _ = reply_open_conn(
&mut send,
&OpenConnResponse::Failed(format!("unknown tunnel id {}", open.tunnel_id)),
)
.await;
error!(
tunnel_id = open.tunnel_id,
"server pushed conn for unknown tunnel"
);
return;
}
};
let dispatch = match resolve_reverse_dispatch(&parent, open.dynamic) {
Ok(d) => d,
Err(e) => {
let _ = reply_open_conn(&mut send, &OpenConnResponse::Failed(e.to_string())).await;
error!(error = %e, "reverse OpenConn dispatch error");
return;
}
};
if let Err(e) = reply_open_conn(&mut send, &OpenConnResponse::Ok).await {
error!(error = %e, "failed to ack reverse OpenConn");
return;
}
let span =
info_span!("conn", tunnel_id = open.tunnel_id, dir = "reverse", target = %dispatch);
async move {
info!("conn opened");
let started = std::time::Instant::now();
let result = match dispatch {
ReverseDispatch::Tcp(req) => tunnel_tcp_server(recv, send, req, None).await,
ReverseDispatch::Udp(req) => tunnel_udp_server(recv, send, req, None).await,
};
let dur_ms = started.elapsed().as_millis() as u64;
match &result {
Ok(()) => info!(dur_ms, "conn closed"),
Err(e) => debug!(dur_ms, error = %e, "conn closed (error)"),
}
}
.instrument(span)
.await;
});
Ok(())
}
enum ReverseDispatch {
Tcp(RemoteRequest),
Udp(RemoteRequest),
}
impl std::fmt::Display for ReverseDispatch {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ReverseDispatch::Tcp(r) | ReverseDispatch::Udp(r) => write!(f, "{r}"),
}
}
}
fn resolve_reverse_dispatch(
parent: &RemoteRequest,
dynamic: Option<DynamicTarget>,
) -> Result<ReverseDispatch> {
if !matches!(parent.direction, Direction::Reverse) {
return Err(anyhow!(
"server pushed conn on a forward tunnel ({parent}) — protocol error"
));
}
match (&parent.kind, dynamic) {
(RemoteKind::Tcp { local, remote }, None) => Ok(ReverseDispatch::Tcp(RemoteRequest::new(
Direction::Reverse,
RemoteKind::Tcp {
local: *local,
remote: remote.clone(),
},
))),
(RemoteKind::Udp { local, remote }, None) => Ok(ReverseDispatch::Udp(RemoteRequest::new(
Direction::Reverse,
RemoteKind::Udp {
local: *local,
remote: remote.clone(),
},
))),
(RemoteKind::Socks5 { local }, Some(DynamicTarget::Tcp(target))) => {
Ok(ReverseDispatch::Tcp(RemoteRequest::new(
Direction::Reverse,
RemoteKind::Tcp {
local: *local,
remote: target,
},
)))
}
(RemoteKind::Socks5 { local }, Some(DynamicTarget::Udp(target))) => {
Ok(ReverseDispatch::Udp(RemoteRequest::new(
Direction::Reverse,
RemoteKind::Udp {
local: *local,
remote: target,
},
)))
}
(RemoteKind::Socks5 { .. }, None) => Err(anyhow!(
"server pushed reverse SOCKS5 conn without a dynamic target"
)),
(_, Some(_)) => Err(anyhow!(
"server pushed unexpected dynamic target on non-SOCKS reverse tunnel"
)),
}
}