#[cfg(unix)]
pub mod admin;
pub mod state;
#[cfg(unix)]
use std::path::PathBuf;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use anyhow::Result;
use quinn::{Connection, ConnectionError, VarInt};
use tokio::signal;
use tokio::sync::Semaphore;
use tokio::task::JoinSet;
use tracing::{debug, error, info, info_span, warn, Instrument};
use crate::common::quic::create_server_endpoint;
use crate::common::remote::{
Direction, DynamicTarget, RemoteKind, RemoteRequest, SessionHelloResponse,
};
use crate::common::socks::tunnel_socks_client;
use crate::common::tcp::{tunnel_tcp_client, tunnel_tcp_server};
use crate::common::tunnel::{
reply_open_conn, server_receive_session_hello, server_reply_session_hello,
};
use crate::common::udp::{tunnel_udp_client, tunnel_udp_server};
use crate::ServerConfig;
use self::state::{ServerState, TunnelEntry, TunnelHandle};
const CLOSE_CODE_SERVER_SHUTDOWN: u32 = 0;
pub async fn run_async(config: ServerConfig) -> Result<()> {
let endpoint =
create_server_endpoint(config.host, config.port, &config.tls, config.congestion)?;
let listen_addr = endpoint.local_addr()?;
info!(addr = %listen_addr, "server listening");
let client_counter = AtomicUsize::new(0);
let state = ServerState::new(listen_addr);
#[cfg(unix)]
let admin_handle: Option<tokio::task::JoinHandle<()>> = config
.admin_socket
.as_ref()
.map(|path| spawn_admin(state.clone(), path.clone()));
let connection_limiter: Option<Arc<Semaphore>> =
config.max_connections.map(|n| Arc::new(Semaphore::new(n)));
loop {
tokio::select! {
ctrl_c = signal::ctrl_c() => {
if let Err(e) = ctrl_c {
error!(error = %e, "failed to listen for ^C signal");
}
info!("shutdown signal received, notifying clients");
endpoint.close(VarInt::from_u32(CLOSE_CODE_SERVER_SHUTDOWN), b"server received ^C");
endpoint.wait_idle().await;
#[cfg(unix)]
{
if let Some(h) = admin_handle {
h.abort();
}
if let Some(path) = &config.admin_socket {
let _ = std::fs::remove_file(path);
}
}
info!("server stopped");
return Ok(());
}
maybe_conn = endpoint.accept() => {
let Some(conn) = maybe_conn else { break };
let client_id = (client_counter.fetch_add(1, Ordering::Relaxed) + 1) as u64;
let peer = conn.remote_address();
let span = info_span!("client", client_id = client_id, peer = %peer);
let permit = if let Some(limiter) = &connection_limiter {
match limiter.clone().try_acquire_owned() {
Ok(p) => Some(p),
Err(_) => {
warn!(peer = %peer, "rejected: max-connections cap reached");
conn.refuse();
continue;
}
}
} else {
None
};
let allow_reverse = config.allow_reverse;
let allow_socks = config.allow_socks;
let state_for_client = state.clone();
tokio::spawn(
async move {
info!("connected");
match handle_client_connection(
conn,
allow_reverse,
allow_socks,
client_id,
state_for_client,
)
.await
{
Ok(reason) => info!(reason = %reason, "disconnected"),
Err(e) => error!(error = %e, "session failed"),
}
drop(permit);
}
.instrument(span),
);
}
}
}
Ok(())
}
#[cfg(unix)]
fn spawn_admin(state: ServerState, path: PathBuf) -> tokio::task::JoinHandle<()> {
tokio::spawn(async move {
if let Err(e) = admin::serve(state, &path).await {
error!("admin API exited: {e:#}");
}
})
}
async fn handle_client_connection(
conn: quinn::Incoming,
allow_reverse: bool,
allow_socks: bool,
client_id: u64,
state: ServerState,
) -> Result<String> {
let connection = conn.await?;
let mut tunnels: JoinSet<()> = JoinSet::new();
let client_entry =
state.register_client(client_id, connection.remote_address(), connection.clone());
let hello_outcome = perform_session_hello(
&connection,
&state,
&client_entry,
allow_reverse,
allow_socks,
)
.await;
let registered_tunnels = match hello_outcome {
Ok(t) => t,
Err(e) => {
error!(error = %e, "session hello rejected");
state.deregister_client(client_id, format!("hello rejected: {e}"));
return Err(e);
}
};
info!(count = registered_tunnels.len(), "session established");
for tunnel in ®istered_tunnels {
let dir = match tunnel.direction {
Direction::Forward => "forward",
Direction::Reverse => "reverse",
};
info!(
tunnel_id = tunnel.id,
dir,
spec = %tunnel.spec,
"tunnel registered"
);
if matches!(tunnel.direction, Direction::Reverse) {
spawn_reverse_handler(
connection.clone(),
state.clone(),
tunnel.clone(),
&mut tunnels,
);
}
}
let outcome = loop {
let quic_connection = connection.clone();
let stream_result = tokio::select! {
r = quic_connection.accept_bi() => r,
Some(joined) = tunnels.join_next(), if !tunnels.is_empty() => {
if let Err(e) = joined {
if !e.is_cancelled() {
debug!("tunnel task panicked: {e}");
}
}
continue;
}
};
let stream = match stream_result {
Err(ConnectionError::ApplicationClosed(close)) => {
let reason = String::from_utf8_lossy(&close.reason);
break Ok(format!(
"client closed (code {}, {reason})",
close.error_code
));
}
Err(ConnectionError::ConnectionClosed(close)) => {
break Ok(format!("transport closed ({close})"));
}
Err(ConnectionError::LocallyClosed) => {
break Ok("locally closed".to_string());
}
Err(ConnectionError::TimedOut) => {
break Ok("idle timeout (peer went away)".to_string());
}
Err(ConnectionError::Reset) => {
break Ok("connection reset by peer".to_string());
}
Err(e) => {
error!(error = %e, "stream error");
break Err(e.into());
}
Ok(s) => s,
};
let state_for_conn = state.clone();
let fut = handle_open_conn(stream, state_for_conn);
tunnels.spawn(async move {
if let Err(e) = fut.await {
error!(error = %e, "conn failed");
}
});
};
let aborted = tunnels.len();
tunnels.shutdown().await;
if aborted > 0 {
debug!(aborted, "aborted in-flight tunnel tasks");
}
let reason = match &outcome {
Ok(r) => r.clone(),
Err(e) => format!("error: {e}"),
};
state.deregister_client(client_id, reason);
outcome
}
async fn perform_session_hello(
connection: &Connection,
state: &ServerState,
client: &Arc<state::ClientEntry>,
allow_reverse: bool,
allow_socks: bool,
) -> Result<Vec<Arc<TunnelEntry>>> {
let (mut send, mut recv) = connection.accept_bi().await?;
let hello = server_receive_session_hello(&mut recv).await?;
if let Err(reason) = validate_remotes(&hello.remotes, allow_reverse, allow_socks) {
let resp = SessionHelloResponse::Failed(reason.clone());
let _ = server_reply_session_hello(&mut send, &resp).await;
return Err(anyhow::anyhow!(reason));
}
let tunnels = state.register_tunnels(client, &hello.remotes);
let tunnel_ids: Vec<u64> = tunnels.iter().map(|t| t.id).collect();
server_reply_session_hello(&mut send, &SessionHelloResponse::Ok { tunnel_ids }).await?;
Ok(tunnels)
}
fn validate_remotes(
remotes: &[RemoteRequest],
allow_reverse: bool,
allow_socks: bool,
) -> Result<(), String> {
for r in remotes {
if r.is_reversed() && !allow_reverse {
return Err(format!("Reverse remotes are not allowed ({r})"));
}
if r.is_socks() && !allow_socks {
return Err(format!("SOCKS5 remotes are not allowed ({r})"));
}
}
Ok(())
}
fn spawn_reverse_handler(
connection: Connection,
state: ServerState,
tunnel: Arc<TunnelEntry>,
tasks: &mut JoinSet<()>,
) {
let handle = Arc::new(TunnelHandle::new(state, tunnel.clone()));
let span = info_span!(
"tunnel",
tunnel_id = tunnel.id,
dir = "reverse",
spec = %tunnel.spec,
);
tasks.spawn(
async move {
let request = RemoteRequest::new(tunnel.direction, tunnel.kind.clone());
let result = match &tunnel.kind {
RemoteKind::Tcp { .. } => {
tunnel_tcp_client(connection, request, Some(handle), tunnel.id).await
}
RemoteKind::Udp { .. } => {
tunnel_udp_client(connection, request, Some(handle), tunnel.id).await
}
RemoteKind::Socks5 { .. } => {
tunnel_socks_client(connection, request, Some(handle), tunnel.id).await
}
};
if let Err(e) = result {
error!(error = %e, "reverse handler failed");
}
}
.instrument(span),
);
}
async fn handle_open_conn(
(mut send, mut recv): (quinn::SendStream, quinn::RecvStream),
state: ServerState,
) -> Result<()> {
use crate::common::remote::OpenConnResponse;
let open = crate::common::tunnel::receive_open_conn(&mut recv).await?;
let tunnel = match state.tunnel(open.tunnel_id) {
Some(t) => t,
None => {
let _ = reply_open_conn(
&mut send,
&OpenConnResponse::Failed(format!("unknown tunnel id {}", open.tunnel_id)),
)
.await;
return Err(anyhow::anyhow!("unknown tunnel id {}", open.tunnel_id));
}
};
let dispatch = match resolve_dispatch(&tunnel, open.dynamic.as_ref()) {
Ok(d) => d,
Err(e) => {
let _ = reply_open_conn(&mut send, &OpenConnResponse::Failed(e.to_string())).await;
return Err(e);
}
};
reply_open_conn(&mut send, &OpenConnResponse::Ok).await?;
let peer = dispatch.peer_label();
let conn = state.register_conn(&tunnel, peer.clone());
let conn_id = conn.id();
let counters = conn.counters();
let span = info_span!(
"conn",
conn_id = conn_id,
tunnel_id = tunnel.id,
peer = peer.as_deref().unwrap_or("-"),
);
async move {
info!("conn opened");
let started = std::time::Instant::now();
let result = match dispatch {
ForwardDispatch::Tcp(req) => {
tunnel_tcp_server(recv, send, req, Some(counters.clone())).await
}
ForwardDispatch::Udp(req) => {
tunnel_udp_server(recv, send, req, Some(counters.clone())).await
}
};
let (bytes_in, bytes_out) = counters.snapshot();
let dur_ms = started.elapsed().as_millis() as u64;
match &result {
Ok(()) => info!(bytes_in, bytes_out, dur_ms, "conn closed"),
Err(e) => warn!(bytes_in, bytes_out, dur_ms, error = %e, "conn closed (error)"),
}
drop(conn);
result
}
.instrument(span)
.await
}
enum ForwardDispatch {
Tcp(RemoteRequest),
Udp(RemoteRequest),
}
impl ForwardDispatch {
fn peer_label(&self) -> Option<String> {
match self {
ForwardDispatch::Tcp(r) | ForwardDispatch::Udp(r) => r.remote_addr_string(),
}
}
}
fn resolve_dispatch(
tunnel: &TunnelEntry,
dynamic: Option<&DynamicTarget>,
) -> Result<ForwardDispatch> {
if !matches!(tunnel.direction, Direction::Forward) {
return Err(anyhow::anyhow!(
"OpenConn on reverse tunnel {} (server pushes reverse conns, not the client)",
tunnel.id
));
}
match (&tunnel.kind, dynamic) {
(RemoteKind::Tcp { local, remote }, None) => Ok(ForwardDispatch::Tcp(RemoteRequest::new(
Direction::Forward,
RemoteKind::Tcp {
local: *local,
remote: remote.clone(),
},
))),
(RemoteKind::Udp { local, remote }, None) => Ok(ForwardDispatch::Udp(RemoteRequest::new(
Direction::Forward,
RemoteKind::Udp {
local: *local,
remote: remote.clone(),
},
))),
(RemoteKind::Socks5 { local }, Some(DynamicTarget::Tcp(target))) => {
Ok(ForwardDispatch::Tcp(RemoteRequest::new(
Direction::Forward,
RemoteKind::Tcp {
local: *local,
remote: target.clone(),
},
)))
}
(RemoteKind::Socks5 { local }, Some(DynamicTarget::Udp(target))) => {
Ok(ForwardDispatch::Udp(RemoteRequest::new(
Direction::Forward,
RemoteKind::Udp {
local: *local,
remote: target.clone(),
},
)))
}
(RemoteKind::Socks5 { .. }, None) => Err(anyhow::anyhow!(
"OpenConn on SOCKS5 tunnel {} requires a `dynamic` target",
tunnel.id
)),
(_, Some(_)) => Err(anyhow::anyhow!(
"OpenConn on tunnel {} carried unexpected dynamic target",
tunnel.id
)),
}
}