use clap::crate_version;
use clap::error::ErrorKind;
use clap::{Args as ClapArgs, CommandFactory, Parser, Subcommand, ValueEnum};
use rusnel::cert;
use rusnel::common::quic::Congestion;
use rusnel::common::remote::RemoteRequest;
use rusnel::common::tls::{parse_fingerprint, ClientTlsConfig, ServerTlsConfig};
use rusnel::embedded::{self, Materialized};
use rusnel::{run_client, run_server, ClientConfig, ReconnectConfig, ServerConfig, ServerEndpoint};
#[derive(Debug, Clone, Copy, Default, ValueEnum)]
enum CongestionArg {
#[default]
Cubic,
Bbr,
}
impl From<CongestionArg> for Congestion {
fn from(c: CongestionArg) -> Self {
match c {
CongestionArg::Cubic => Congestion::Cubic,
CongestionArg::Bbr => Congestion::Bbr,
}
}
}
use std::net::{IpAddr, ToSocketAddrs};
use std::path::PathBuf;
use std::str::FromStr;
use std::time::Duration;
use tracing::{debug, info};
fn parse_duration_secs(s: &str) -> Result<Duration, String> {
let secs: u64 = s
.parse()
.map_err(|e| format!("invalid duration `{s}` (expected whole seconds): {e}"))?;
Ok(Duration::from_secs(secs))
}
fn max_retries_from_cli(raw: i64) -> Result<Option<u32>, String> {
if raw < 0 {
Ok(None)
} else {
u32::try_from(raw)
.map(Some)
.map_err(|_| format!("retry count `{raw}` is too large"))
}
}
fn parse_server_addr(s: &str) -> Result<ServerEndpoint, String> {
let host = if let Some(rest) = s.strip_prefix('[') {
let close = rest
.find(']')
.ok_or_else(|| format!("malformed IPv6 address in `{s}` (missing `]`)"))?;
rest[..close].to_string()
} else {
s.rsplit_once(':')
.ok_or_else(|| format!("expected host:port in `{s}`"))?
.0
.to_string()
};
let resolved: Vec<_> = s
.to_socket_addrs()
.map_err(|e| format!("failed to resolve server address `{s}`: {e}"))?
.collect();
if resolved.is_empty() {
return Err(format!("no addresses found for server `{s}`"));
}
let addrs = interleave_address_families(resolved);
Ok(ServerEndpoint { addrs, host })
}
fn interleave_address_families(resolved: Vec<std::net::SocketAddr>) -> Vec<std::net::SocketAddr> {
let mut v4 = Vec::new();
let mut v6 = Vec::new();
for a in resolved {
if a.is_ipv6() {
v6.push(a);
} else {
v4.push(a);
}
}
let (mut primary, mut alternate) = if !v6.is_empty() && !v4.is_empty() {
(v6.into_iter(), v4.into_iter())
} else {
(v4.into_iter(), v6.into_iter())
};
let mut out = Vec::new();
loop {
match (primary.next(), alternate.next()) {
(Some(a), Some(b)) => {
out.push(a);
out.push(b);
}
(Some(a), None) => out.push(a),
(None, Some(b)) => out.push(b),
(None, None) => break,
}
}
out
}
fn parse_remote(s: &str) -> Result<RemoteRequest, String> {
RemoteRequest::from_str(s).map_err(|e| format!("invalid remote `{s}`: {e}"))
}
#[derive(Parser)]
#[command(name = "Rusnel", version = crate_version!())]
#[command(about = "A fast tcp/udp tunnel", long_about = None)]
struct Args {
#[command(subcommand)]
mode: Mode,
}
#[derive(Debug, Subcommand)]
enum Mode {
#[allow(clippy::too_many_arguments)]
Server {
#[arg(long, default_value_t = IpAddr::V4(std::net::Ipv4Addr::new(0, 0, 0, 0)))]
host: IpAddr,
#[arg(long, short, default_value_t = 8080)]
port: u16,
#[arg(long, default_value_t = false)]
allow_reverse: bool,
#[arg(long, default_value_t = false)]
insecure: bool,
#[arg(long, default_value_t = false, conflicts_with_all = ["insecure", "tls_cert", "tls_key"])]
tls_self_signed: bool,
#[arg(long, value_name = "DIR", requires = "tls_self_signed")]
tls_state_dir: Option<PathBuf>,
#[arg(long, value_name = "PATH", requires = "tls_key", conflicts_with_all = ["insecure", "tls_self_signed"])]
tls_cert: Option<PathBuf>,
#[arg(long, value_name = "PATH", requires = "tls_cert")]
tls_key: Option<PathBuf>,
#[arg(long, value_name = "PATH", requires = "tls_cert", conflicts_with_all = ["insecure", "tls_self_signed"])]
tls_ca: Option<PathBuf>,
#[arg(long, value_enum, default_value_t = CongestionArg::Cubic)]
congestion: CongestionArg,
#[arg(long, value_name = "N", default_value_t = 0)]
max_connections: usize,
#[arg(short('v'), long("verbose"), default_value_t = false)]
is_verbose: bool,
#[arg(long("debug"), default_value_t = false)]
is_debug: bool,
},
#[allow(clippy::too_many_arguments)]
Client {
#[arg(value_parser = parse_server_addr)]
server: ServerEndpoint,
#[arg(name = "remote", required = true, value_parser = parse_remote, value_delimiter = ' ', num_args = 1.., help=r#"
<remote>s are remote connections tunneled through the server, each which come in the form:
<local-host>:<local-port>:<remote-host>:<remote-port>/<protocol>
■ local-host defaults to 0.0.0.0 (all interfaces).
■ local-port defaults to remote-port.
■ remote-port is required*.
■ remote-host defaults to 0.0.0.0 (server localhost).
■ protocol defaults to tcp.
which shares <remote-host>:<remote-port> from the server to the client as <local-host>:<local-port>, or:
R:<local-host>:<local-port>:<remote-host>:<remote-port>/<protocol>
which does reverse port forwarding,
sharing <remote-host>:<remote-port> from the client to the server\'s <local-host>:<local-port>.
example remotes
1337
example.com:1337
1337:google.com:80
192.168.1.14:5000:google.com:80
socks
5000:socks
R:2222:localhost:22
R:socks
R:5000:socks
1.1.1.1:53/udp
When the Rusnel server has --allow-reverse enabled, remotes can be prefixed with R to denote that they are reversed.
Remotes can specify "socks" in place of remote-host and remote-port.
The default local host and port for a "socks" remote is 127.0.0.1:1080.
"#)]
remotes: Vec<RemoteRequest>,
#[arg(long, default_value_t = false)]
insecure: bool,
#[arg(long, value_name = "SHA256", conflicts_with_all = ["insecure", "tls_ca"])]
tls_fingerprint: Option<String>,
#[arg(long, value_name = "PATH", conflicts_with = "insecure")]
tls_ca: Option<PathBuf>,
#[arg(long, value_name = "PATH", requires_all = ["tls_key", "tls_ca"])]
tls_cert: Option<PathBuf>,
#[arg(long, value_name = "PATH", requires_all = ["tls_cert", "tls_ca"])]
tls_key: Option<PathBuf>,
#[arg(long, value_name = "NAME")]
tls_server_name: Option<String>,
#[arg(long, value_enum, default_value_t = CongestionArg::Cubic)]
congestion: CongestionArg,
#[arg(long, value_name = "N", default_value_t = -1, allow_hyphen_values = true)]
max_retry_count: i64,
#[arg(long, value_name = "SECONDS", default_value = "300", value_parser = parse_duration_secs)]
max_retry_interval: Duration,
#[arg(short('v'), long("verbose"), default_value_t = false)]
is_verbose: bool,
#[arg(long("debug"), default_value_t = false)]
is_debug: bool,
},
Cert {
#[command(subcommand)]
action: CertAction,
},
}
#[derive(Debug, Subcommand)]
enum CertAction {
Ca(CaArgs),
Server(ServerCertArgs),
Client(ClientCertArgs),
Fingerprint {
cert: PathBuf,
},
}
#[derive(Debug, ClapArgs)]
struct CaArgs {
#[arg(long, value_name = "DIR", default_value = "./pki")]
out_dir: PathBuf,
#[arg(long, default_value = "rusnel-ca")]
common_name: String,
}
#[derive(Debug, ClapArgs)]
struct ServerCertArgs {
#[arg(long, value_name = "DIR", default_value = "./pki")]
out_dir: PathBuf,
#[arg(long, value_name = "PATH")]
ca: PathBuf,
#[arg(long, value_name = "PATH")]
ca_key: PathBuf,
#[arg(long)]
common_name: Option<String>,
#[arg(long = "name", value_name = "DNS")]
names: Vec<String>,
#[arg(long = "ip", value_name = "IP")]
ips: Vec<IpAddr>,
#[arg(long, default_value = "server")]
file_stem: String,
}
#[derive(Debug, ClapArgs)]
struct ClientCertArgs {
#[arg(long, value_name = "DIR", default_value = "./pki")]
out_dir: PathBuf,
#[arg(long, value_name = "PATH")]
ca: PathBuf,
#[arg(long, value_name = "PATH")]
ca_key: PathBuf,
#[arg(long, default_value = "rusnel-client")]
common_name: String,
#[arg(long)]
file_stem: Option<String>,
}
fn resolve_server_tls(
insecure: bool,
tls_self_signed: bool,
tls_state_dir: Option<PathBuf>,
tls_cert: Option<PathBuf>,
tls_key: Option<PathBuf>,
tls_ca: Option<PathBuf>,
embedded: &Materialized,
) -> Result<ServerTlsConfig, String> {
if insecure {
return Ok(ServerTlsConfig::Insecure);
}
if tls_self_signed {
let state_dir = match tls_state_dir {
Some(p) => p,
None => default_state_dir()?,
};
return Ok(ServerTlsConfig::SelfSigned { state_dir });
}
if let (Some(cert), Some(key)) = (tls_cert.clone(), tls_key.clone()) {
return Ok(match tls_ca {
Some(ca) => ServerTlsConfig::Mtls { cert, key, ca },
None => ServerTlsConfig::Provided { cert, key },
});
}
if let (Some(cert), Some(key)) = (embedded.server_cert.clone(), embedded.server_key.clone()) {
info!("using embedded server credentials baked in at build time");
return Ok(match embedded.ca.clone() {
Some(ca) => ServerTlsConfig::Mtls { cert, key, ca },
None => ServerTlsConfig::Provided { cert, key },
});
}
Err(
"no TLS mode specified. Pass one of --insecure, --tls-self-signed, \
--tls-cert + --tls-key (with optional --tls-ca for mTLS), or build \
with RUSNEL_EMBED_SERVER_CERT / RUSNEL_EMBED_SERVER_KEY."
.into(),
)
}
fn default_state_dir() -> Result<PathBuf, String> {
dirs::home_dir()
.map(|h| h.join(".rusnel"))
.ok_or_else(|| "could not determine home directory; pass --tls-state-dir explicitly".into())
}
fn resolve_client_tls(
insecure: bool,
tls_fingerprint: Option<String>,
tls_ca: Option<PathBuf>,
tls_cert: Option<PathBuf>,
tls_key: Option<PathBuf>,
tls_server_name: Option<String>,
embedded: &Materialized,
) -> Result<ClientTlsConfig, String> {
if insecure {
return Ok(ClientTlsConfig::Insecure);
}
let embedded_server_name = || embedded::EMBED_SERVER_NAME.map(|s| s.to_string());
if let Some(raw) = tls_fingerprint {
let sha256 = parse_fingerprint(&raw)
.map_err(|e| format!("invalid --tls-fingerprint value `{raw}`: {e}"))?;
return Ok(ClientTlsConfig::Fingerprint {
sha256,
server_name: tls_server_name.or_else(embedded_server_name),
});
}
if let Some(ca) = tls_ca {
return Ok(match (tls_cert, tls_key) {
(Some(cert), Some(key)) => ClientTlsConfig::Mtls {
ca,
cert,
key,
server_name: tls_server_name.or_else(embedded_server_name),
},
_ => ClientTlsConfig::Ca {
ca,
server_name: tls_server_name.or_else(embedded_server_name),
},
});
}
if let Some(ca) = embedded.ca.clone() {
info!("using embedded client credentials baked in at build time");
return Ok(
match (embedded.client_cert.clone(), embedded.client_key.clone()) {
(Some(cert), Some(key)) => ClientTlsConfig::Mtls {
ca,
cert,
key,
server_name: tls_server_name.or_else(embedded_server_name),
},
_ => ClientTlsConfig::Ca {
ca,
server_name: tls_server_name.or_else(embedded_server_name),
},
},
);
}
if let Some(fp) = embedded::EMBED_FINGERPRINT {
info!("using embedded server fingerprint baked in at build time");
let sha256 = parse_fingerprint(fp).map_err(|e| {
format!("invalid embedded fingerprint (RUSNEL_EMBED_FINGERPRINT) `{fp}`: {e}")
})?;
return Ok(ClientTlsConfig::Fingerprint {
sha256,
server_name: tls_server_name.or_else(embedded_server_name),
});
}
Err(
"no TLS mode specified. Pass one of --insecure, --tls-fingerprint, \
--tls-ca (with optional --tls-cert + --tls-key for mTLS), or build \
with RUSNEL_EMBED_CA / RUSNEL_EMBED_FINGERPRINT."
.into(),
)
}
fn main() {
if let Err(e) = rustls::crypto::ring::default_provider().install_default() {
Args::command()
.error(
ErrorKind::Io,
format!("failed to install rustls crypto provider: {e:?}"),
)
.exit();
}
let args = Args::parse();
match args.mode {
Mode::Server {
host,
port,
allow_reverse,
insecure,
tls_self_signed,
tls_state_dir,
tls_cert,
tls_key,
tls_ca,
congestion,
max_connections,
is_verbose,
is_debug,
} => {
set_log_level(is_verbose, is_debug);
let embedded = match embedded::materialize() {
Ok(m) => m,
Err(e) => Args::command()
.error(
ErrorKind::Io,
format!("failed to materialize embedded credentials: {e:#}"),
)
.exit(),
};
let tls = match resolve_server_tls(
insecure,
tls_self_signed,
tls_state_dir,
tls_cert,
tls_key,
tls_ca,
embedded,
) {
Ok(t) => t,
Err(msg) => Args::command().error(ErrorKind::InvalidValue, msg).exit(),
};
let server_config = ServerConfig {
host,
port,
allow_reverse,
tls,
congestion: congestion.into(),
max_connections: if max_connections == 0 {
None
} else {
Some(max_connections)
},
};
debug!("Initialized server with config: {:?}", server_config);
run_server(server_config);
}
Mode::Client {
server,
remotes,
insecure,
tls_fingerprint,
tls_ca,
tls_cert,
tls_key,
tls_server_name,
congestion,
max_retry_count,
max_retry_interval,
is_verbose,
is_debug,
} => {
set_log_level(is_verbose, is_debug);
let embedded = match embedded::materialize() {
Ok(m) => m,
Err(e) => Args::command()
.error(
ErrorKind::Io,
format!("failed to materialize embedded credentials: {e:#}"),
)
.exit(),
};
let tls = match resolve_client_tls(
insecure,
tls_fingerprint,
tls_ca,
tls_cert,
tls_key,
tls_server_name,
embedded,
) {
Ok(t) => t,
Err(msg) => Args::command().error(ErrorKind::InvalidValue, msg).exit(),
};
let max_retries = match max_retries_from_cli(max_retry_count) {
Ok(v) => v,
Err(msg) => Args::command().error(ErrorKind::InvalidValue, msg).exit(),
};
let reconnect = ReconnectConfig {
max_retries,
max_backoff: max_retry_interval,
..ReconnectConfig::default()
};
let client_config = ClientConfig {
server,
remotes,
tls,
congestion: congestion.into(),
reconnect,
};
debug!("Initialized client with config: {:?}", client_config);
run_client(client_config);
}
Mode::Cert { action } => {
tracing_subscriber::fmt()
.with_max_level(tracing::Level::INFO)
.with_target(false)
.without_time()
.init();
if let Err(e) = run_cert(action) {
Args::command()
.error(ErrorKind::Io, format!("{e:#}"))
.exit();
}
}
}
}
fn run_cert(action: CertAction) -> anyhow::Result<()> {
match action {
CertAction::Ca(a) => {
cert::generate_ca(&a.out_dir, &a.common_name)?;
}
CertAction::Server(a) => {
let cn = a
.common_name
.clone()
.or_else(|| a.names.first().cloned())
.unwrap_or_else(|| "rusnel-server".to_string());
cert::generate_server_cert(
&a.out_dir,
&a.ca,
&a.ca_key,
&cn,
&a.names,
&a.ips,
&a.file_stem,
)?;
}
CertAction::Client(a) => {
let stem = a.file_stem.clone().unwrap_or_else(|| a.common_name.clone());
cert::generate_client_cert(&a.out_dir, &a.ca, &a.ca_key, &a.common_name, &stem)?;
}
CertAction::Fingerprint { cert } => {
let fp = cert::print_fingerprint(&cert)?;
println!("{fp}");
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use std::net::{Ipv4Addr, Ipv6Addr, SocketAddr};
fn v4(p: u16) -> SocketAddr {
SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 0, 0, p as u8)), p)
}
fn v6(p: u16) -> SocketAddr {
SocketAddr::new(IpAddr::V6(Ipv6Addr::new(0, 0, 0, 0, 0, 0, 0, p as u16)), p)
}
#[test]
fn interleave_alternates_families_v6_first() {
let got = interleave_address_families(vec![v6(1), v6(2), v4(3), v4(4)]);
assert_eq!(got, vec![v6(1), v4(3), v6(2), v4(4)]);
}
#[test]
fn interleave_alternates_families_v4_first() {
let got = interleave_address_families(vec![v4(1), v4(2), v6(3), v6(4)]);
assert_eq!(got, vec![v6(3), v4(1), v6(4), v4(2)]);
}
#[test]
fn interleave_single_family_preserves_order() {
let got = interleave_address_families(vec![v4(1), v4(2), v4(3)]);
assert_eq!(got, vec![v4(1), v4(2), v4(3)]);
let got = interleave_address_families(vec![v6(1), v6(2)]);
assert_eq!(got, vec![v6(1), v6(2)]);
}
}
fn set_log_level(is_verbose: bool, is_debug: bool) {
let log_level = if is_debug {
tracing::Level::TRACE
} else if is_verbose {
tracing::Level::DEBUG
} else {
tracing::Level::INFO
};
tracing_subscriber::fmt().with_max_level(log_level).init();
debug!("log level: {}", log_level);
}