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!("routing QUIC through proxy: {p}");
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. Broadcasting shutdown...");
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!("Run function completed");
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!(
"connecting to server at: {} (sni: {})",
config.server, server_name
);
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)) => {
info!("Connected successfully");
attempt = 0;
backoff = initial_backoff;
let outcome = run_connection(connection, config, shutdown_tx).await;
match outcome {
SessionOutcome::Shutdown => return Ok(()),
SessionOutcome::Disconnected(reason) => {
warn!("connection lost: {}", reason);
}
}
}
Some(Err(e)) => {
warn!("connection attempt failed: {}", e);
}
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
));
}
}
info!(
"reconnecting in {:?} (attempt {}{})",
backoff,
attempt,
max_retries
.map(|m| format!("/{m}"))
.unwrap_or_else(|| "".to_string()),
);
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, "happy eyeballs winner");
return Ok(conn);
}
Err(e) => {
debug!(%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!(
target = %server,
proxy = %proxy,
"establishing SOCKS5 UDP ASSOCIATE for QUIC connection",
);
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!("session established with {} tunnel(s)", tunnel_ids.len());
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 = tunnel_id, remote = %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!("failed: {}", e)
}
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!("reverse conn accept loop ended: {}", e);
break;
}
}
anyhow::Ok(())
});
tasks.push(accept_reverse_task);
let mut shutdown_rx = shutdown_tx.subscribe();
let outcome = tokio::select! {
_ = shutdown_rx.recv() => {
info!("Shutting down client 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!("failed to read OpenConn frame: {e}");
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!(
"server pushed conn for unknown tunnel id {}",
open.tunnel_id
);
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!("reverse OpenConn dispatch error: {e}");
return;
}
};
if let Err(e) = reply_open_conn(&mut send, &OpenConnResponse::Ok).await {
error!("failed to ack reverse OpenConn: {e}");
return;
}
let span = info_span!("conn", tunnel_id = open.tunnel_id, remote = %dispatch);
async move {
info!("reverse conn opened");
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,
};
if let Err(e) = result {
debug!("reverse conn ended: {e}");
}
}
.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"
)),
}
}