use std::collections::HashMap;
use std::net::{IpAddr, Ipv4Addr, SocketAddr, ToSocketAddrs};
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use anyhow::{Context, Result};
use clap::Parser;
use log::{debug, error, info, warn};
use tokio::net::UdpSocket;
use tokio::sync::mpsc;
use shadowvpn::config::{ServerArgs, ServerConfig};
use shadowvpn::crypto::{decrypt_packet, encrypt_packet};
use shadowvpn::mesh::{self, Control, RouteApproval, RoutePush, SubnetTable};
use shadowvpn::nat::{Ingress, Nat};
use shadowvpn::obfs::{self, Obfuscator};
use shadowvpn::protocol::{max_datagram_size, MAX_IP_PACKET};
use shadowvpn::tun_device::TunDevice;
#[derive(Default)]
struct Learned {
clients: HashMap<IpAddr, SocketAddr>,
subnets: SubnetTable,
}
impl Learned {
fn learn(&mut self, src: IpAddr, peer: SocketAddr, via: &str) {
if self.clients.insert(src, peer) != Some(peer) {
info!("client {src} reachable via {peer}{via}");
}
}
fn lookup(&self, dst: IpAddr) -> Option<SocketAddr> {
self.clients
.get(&dst)
.copied()
.or_else(|| self.subnets.lookup(dst))
}
}
enum Routing {
Learn(Learned),
Nat(Nat),
}
type Shared = Arc<Mutex<Routing>>;
const CHANNEL_DEPTH: usize = 1024;
#[tokio::main]
async fn main() -> Result<()> {
env_logger::Builder::from_env(env_logger::Env::default().default_filter_or("info")).init();
let cfg = ServerArgs::parse()
.resolve()
.context("failed to resolve server configuration")?;
if let Err(err) = run(cfg).await {
error!("server exited with error: {err:#}");
return Err(err);
}
Ok(())
}
async fn run(cfg: ServerConfig) -> Result<()> {
let listen_addr = cfg
.listen
.to_socket_addrs()
.with_context(|| format!("resolving listen address {}", cfg.listen))?
.next()
.with_context(|| format!("no address resolved for {}", cfg.listen))?;
let socket = shadowvpn::net::bind_udp(listen_addr)
.with_context(|| format!("failed to bind UDP socket on {}", cfg.listen))?;
let socket = Arc::new(socket);
let tun = TunDevice::create(&cfg.tun)
.context("failed to create TUN device (TUN setup needs root / elevated privileges)")?;
let tun = Arc::new(tun);
let tun_name = tun.name().unwrap_or_else(|_| {
cfg.tun
.name
.clone()
.unwrap_or_else(|| "<unknown>".to_string())
});
print_banner(&cfg, &tun_name);
let routing: Shared = Arc::new(Mutex::new(if cfg.nat {
let nat = Nat::new(cfg.tun.ip, cfg.tun.netmask, cfg.lease_ttl);
info!(
" NAT : ENABLED ({} clients max, idle TTL {}s)",
nat.capacity(),
cfg.lease_ttl.as_secs()
);
Routing::Nat(nat)
} else {
Routing::Learn(Learned::default())
}));
let obfuscator: Option<Arc<Obfuscator>> = cfg
.obfs
.as_deref()
.and_then(Obfuscator::from_name)
.map(Arc::new);
if let Some(name) = cfg.obfs.as_deref() {
info!(" obfuscation : {name} datagram shaping ENABLED");
}
let nat_enabled = cfg.nat;
let lease_ttl = cfg.lease_ttl;
let cfg = Arc::new(cfg);
let a = {
let socket = Arc::clone(&socket);
let tun = Arc::clone(&tun);
let routing = Arc::clone(&routing);
let cfg = Arc::clone(&cfg);
let obfs = obfuscator.clone();
tokio::spawn(async move { udp_to_tun(socket, tun, routing, cfg, obfs).await })
};
let b = {
let socket = Arc::clone(&socket);
let tun = Arc::clone(&tun);
let routing = Arc::clone(&routing);
let cfg = Arc::clone(&cfg);
let obfs = obfuscator.clone();
tokio::spawn(async move { tun_to_udp(socket, tun, routing, cfg, obfs).await })
};
let _sweeper = {
let routing = Arc::clone(&routing);
let interval = (lease_ttl / 2).max(Duration::from_secs(5));
let _ = nat_enabled;
tokio::spawn(async move {
let mut tick = tokio::time::interval(interval);
tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
loop {
tick.tick().await;
match &mut *routing.lock().unwrap() {
Routing::Nat(nat) => {
nat.reap(Instant::now());
}
Routing::Learn(learned) => {
for net in learned.subnets.expire(lease_ttl, Instant::now()) {
info!("subnet route {net} expired (advertiser went quiet)");
}
}
}
}
})
};
tokio::select! {
res = a => res.context("UDP→TUN task panicked")?,
res = b => res.context("TUN→UDP task panicked")?,
}
}
async fn udp_to_tun(
socket: Arc<UdpSocket>,
tun: Arc<TunDevice>,
routing: Shared,
cfg: Arc<ServerConfig>,
obfuscator: Option<Arc<Obfuscator>>,
) -> Result<()> {
let cipher = cfg.cipher;
let (tx, mut rx) = mpsc::channel::<(SocketAddr, Vec<u8>)>(CHANNEL_DEPTH);
let socket_out = Arc::clone(&socket);
let reader = tokio::spawn(async move {
let mut buf = vec![0u8; max_datagram_size(cipher) + obfs::MAX_HEADER];
loop {
let (n, peer) = socket
.recv_from(&mut buf)
.await
.context("UDP recv_from failed")?;
if tx.send((peer, buf[..n].to_vec())).await.is_err() {
return Ok(()); }
}
});
let processor = tokio::spawn(async move {
while let Some((peer, pkt)) = rx.recv().await {
let n = pkt.len();
let decoded;
let datagram: &[u8] = match obfuscator {
Some(ref o) => match o.unwrap(&pkt) {
Some(inner) => {
decoded = inner;
&decoded
}
None => {
debug!("dropping {n}-byte non-obfs datagram from {peer}");
continue;
}
},
None => &pkt,
};
let mut plaintext = match decrypt_packet(cipher, &cfg.master_key, datagram) {
Ok(pt) => pt,
Err(err) => {
debug!("dropping {n}-byte datagram from {peer}: decrypt failed: {err}");
continue;
}
};
let now = Instant::now();
if mesh::is_control(&plaintext) {
let reply = handle_control(&routing, &cfg.route_approval, peer, &plaintext, now);
if let Some(push) = reply {
send_ciphered(
&socket_out,
cipher,
&cfg.master_key,
&obfuscator,
&push.encode(),
peer,
)
.await;
}
continue;
}
if plaintext.len() < 20 {
debug!(
"dropping {}-byte sub-IP-header payload from {peer}",
plaintext.len()
);
continue;
}
enum Action {
Tun,
Relay(SocketAddr),
Bounce,
}
let action = {
let mut guard = routing.lock().unwrap();
match &mut *guard {
Routing::Learn(learned) => {
if let Some(src) = ip_src(&plaintext) {
learned.learn(src, peer, "");
} else {
debug!("datagram from {peer} is not a parseable IP packet; forwarding");
}
match ip_dst(&plaintext).and_then(|dst| learned.lookup(dst)) {
Some(next) if next != peer => Action::Relay(next),
Some(_) => Action::Bounce,
None => Action::Tun,
}
}
Routing::Nat(nat) => match nat.ingress(peer, &mut plaintext, now) {
Ingress::Rewritten(_) => Action::Tun,
Ingress::Exhausted => {
warn!("NAT address pool exhausted; dropping packet from {peer}");
continue;
}
Ingress::Invalid => {
debug!("unparseable IPv4 packet from {peer}; dropping");
continue;
}
},
}
};
match action {
Action::Tun => {
tun.send(&plaintext)
.await
.context("failed to write packet to TUN")?;
}
Action::Relay(next) => {
send_ciphered(
&socket_out,
cipher,
&cfg.master_key,
&obfuscator,
&plaintext,
next,
)
.await;
}
Action::Bounce => {
debug!("dropping {n}-byte packet from {peer}: destination routes back to its sender");
}
}
}
Ok(())
});
let mut reader = reader;
let mut processor = processor;
tokio::select! {
r = &mut reader => { processor.abort(); r.context("UDP→TUN reader task panicked")? }
r = &mut processor => { reader.abort(); r.context("UDP→TUN processor task panicked")? }
}
}
async fn tun_to_udp(
socket: Arc<UdpSocket>,
tun: Arc<TunDevice>,
routing: Shared,
cfg: Arc<ServerConfig>,
obfuscator: Option<Arc<Obfuscator>>,
) -> Result<()> {
let cipher = cfg.cipher;
let (tx, mut rx) = mpsc::channel::<Vec<u8>>(CHANNEL_DEPTH);
let reader = tokio::spawn(async move {
let mut buf = vec![0u8; MAX_IP_PACKET];
loop {
let n = tun
.recv(&mut buf)
.await
.context("failed to read from TUN")?;
if tx.send(buf[..n].to_vec()).await.is_err() {
return Ok(());
}
}
});
let processor = tokio::spawn(async move {
while let Some(mut pkt) = rx.recv().await {
let n = pkt.len();
let now = Instant::now();
let peer = {
let mut guard = routing.lock().unwrap();
match &mut *guard {
Routing::Learn(learned) => ip_dst(&pkt).and_then(|dst| learned.lookup(dst)),
Routing::Nat(nat) => nat.egress(&mut pkt, now),
}
};
let peer = match peer {
Some(peer) => peer,
None => {
debug!("dropping {n}-byte TUN packet: no known client for its destination");
continue;
}
};
let datagram = match encrypt_packet(cipher, &cfg.master_key, &pkt) {
Ok(d) => d,
Err(err) => {
warn!("failed to encrypt packet for {peer}: {err}");
continue;
}
};
let datagram = match obfuscator {
Some(ref o) => o.wrap(&datagram),
None => datagram,
};
if let Err(err) = socket.send_to(&datagram, peer).await {
warn!("failed to send datagram to {peer}: {err}");
}
}
Ok(())
});
let mut reader = reader;
let mut processor = processor;
tokio::select! {
r = &mut reader => { processor.abort(); r.context("TUN→UDP reader task panicked")? }
r = &mut processor => { reader.abort(); r.context("TUN→UDP processor task panicked")? }
}
}
fn handle_control(
routing: &Shared,
approval: &RouteApproval,
peer: SocketAddr,
payload: &[u8],
now: Instant,
) -> Option<RoutePush> {
let control = match mesh::parse_control(payload) {
Some(control) => control,
None => {
debug!(
"dropping malformed {}-byte control payload from {peer}",
payload.len()
);
return None;
}
};
let mut guard = routing.lock().unwrap();
match (&mut *guard, control) {
(Routing::Nat(nat), control) => {
nat.touch(peer, now);
if !matches!(control, Control::Keepalive(_)) {
debug!("ignoring mesh control from {peer}: NAT mode has no subnet routing");
}
None
}
(Routing::Learn(learned), Control::Keepalive(src)) => {
if let Some(src) = src {
learned.learn(IpAddr::V4(src), peer, " (keepalive)");
}
None
}
(Routing::Learn(learned), Control::RouteAdvert(advert)) => {
learned.learn(IpAddr::V4(advert.tunnel_ip), peer, " (advert)");
if let Some(ip6) = advert.tunnel_ip6 {
learned.learn(IpAddr::V6(ip6), peer, " (advert)");
}
let outcome = learned
.subnets
.advertise(peer, &advert.routes, approval, now);
let who = advert.tunnel_ip;
for net in &outcome.approved {
info!("subnet route {net} via client {who} approved");
}
for net in &outcome.awaiting {
warn!(
"subnet route {net} from client {who} is awaiting approval \
(add it to approve_routes, or set auto_approve_routes)"
);
}
for net in &outcome.moved {
info!("subnet route {net} moved to client {who} ({peer})");
}
for net in &outcome.withdrawn {
info!("subnet route {net} withdrawn by client {who}");
}
advert.accept_routes.then(|| RoutePush {
routes: learned.subnets.routes_for(peer),
})
}
(Routing::Learn(_), Control::RoutePush(_)) => {
debug!("ignoring route push from {peer}: pushes only flow server→client");
None
}
}
}
async fn send_ciphered(
socket: &UdpSocket,
cipher: shadowvpn::crypto::Cipher,
master_key: &[u8],
obfuscator: &Option<Arc<Obfuscator>>,
plaintext: &[u8],
peer: SocketAddr,
) {
let datagram = match encrypt_packet(cipher, master_key, plaintext) {
Ok(d) => d,
Err(err) => {
warn!(
"failed to encrypt {}-byte payload for {peer}: {err}",
plaintext.len()
);
return;
}
};
let datagram = match obfuscator {
Some(o) => o.wrap(&datagram),
None => datagram,
};
if let Err(err) = socket.send_to(&datagram, peer).await {
warn!("failed to send datagram to {peer}: {err}");
}
}
fn ip_src(packet: &[u8]) -> Option<IpAddr> {
match packet.first()? >> 4 {
4 if packet.len() >= 20 => Some(IpAddr::V4(Ipv4Addr::new(
packet[12], packet[13], packet[14], packet[15],
))),
6 if packet.len() >= 40 => Some(IpAddr::V6(std::net::Ipv6Addr::from(
<[u8; 16]>::try_from(&packet[8..24]).expect("16 bytes"),
))),
_ => None,
}
}
fn ip_dst(packet: &[u8]) -> Option<IpAddr> {
match packet.first()? >> 4 {
4 if packet.len() >= 20 => Some(IpAddr::V4(Ipv4Addr::new(
packet[16], packet[17], packet[18], packet[19],
))),
6 if packet.len() >= 40 => Some(IpAddr::V6(std::net::Ipv6Addr::from(
<[u8; 16]>::try_from(&packet[24..40]).expect("16 bytes"),
))),
_ => None,
}
}
fn print_banner(cfg: &ServerConfig, tun_name: &str) {
info!("ShadowVPN server starting");
info!(" listen (UDP) : {}", cfg.listen);
info!(" cipher : {}", cfg.cipher.name());
info!(
" TUN interface : {tun_name} ip={} netmask={} peer={} mtu={}",
cfg.tun.ip, cfg.tun.netmask, cfg.tun.peer_ip, cfg.tun.mtu
);
if let Some(ip6) = cfg.tun.ip6 {
info!(" TUN IPv6 : {ip6}");
}
info!(" routing : learn inner src IP -> UDP addr; route by inner dst IP");
if cfg.route_approval.auto {
info!(" mesh routes : auto-approving every advertised subnet route");
} else if !cfg.route_approval.allowlist.is_empty() {
info!(
" mesh routes : approving advertised routes within {:?}",
cfg.route_approval
.allowlist
.iter()
.map(ToString::to_string)
.collect::<Vec<_>>()
);
}
info!("To route client traffic beyond this host, enable forwarding + NAT:");
#[cfg(target_os = "linux")]
{
info!(" Linux: sysctl -w net.ipv4.ip_forward=1");
info!(
" Linux: iptables -t nat -A POSTROUTING -s {}/{} -o <wan-if> -j MASQUERADE",
cfg.tun.ip, cfg.tun.netmask
);
}
#[cfg(target_os = "macos")]
{
info!(" macOS: sysctl -w net.inet.ip.forwarding=1");
info!(
" macOS: configure pf NAT (nat on <wan-if> from {} -> (<wan-if>))",
cfg.tun.ip
);
}
}
#[cfg(test)]
mod tests {
use super::*;
fn ipv4_header(src: [u8; 4], dst: [u8; 4]) -> Vec<u8> {
let mut p = vec![0u8; 20];
p[0] = 0x45; p[12..16].copy_from_slice(&src);
p[16..20].copy_from_slice(&dst);
p
}
fn ipv6_header(src: std::net::Ipv6Addr, dst: std::net::Ipv6Addr) -> Vec<u8> {
let mut p = vec![0u8; 40];
p[0] = 0x60; p[8..24].copy_from_slice(&src.octets());
p[24..40].copy_from_slice(&dst.octets());
p
}
#[test]
fn parses_v4_src_and_dst() {
let p = ipv4_header([10, 7, 0, 2], [10, 7, 0, 1]);
assert_eq!(ip_src(&p), Some("10.7.0.2".parse().unwrap()));
assert_eq!(ip_dst(&p), Some("10.7.0.1".parse().unwrap()));
}
#[test]
fn parses_v6_src_and_dst() {
let src: std::net::Ipv6Addr = "fd07:7::2".parse().unwrap();
let dst: std::net::Ipv6Addr = "fd42:cafe::1".parse().unwrap();
let p = ipv6_header(src, dst);
assert_eq!(ip_src(&p), Some(IpAddr::V6(src)));
assert_eq!(ip_dst(&p), Some(IpAddr::V6(dst)));
}
#[test]
fn rejects_too_short() {
let p = vec![0x45u8; 10];
assert_eq!(ip_src(&p), None);
assert_eq!(ip_dst(&p), None);
let p = vec![0x60u8; 20];
assert_eq!(ip_src(&p), None);
assert_eq!(ip_dst(&p), None);
}
#[test]
fn rejects_unknown_version() {
let mut p = ipv4_header([1, 2, 3, 4], [5, 6, 7, 8]);
p[0] = 0x50; assert_eq!(ip_src(&p), None);
assert_eq!(ip_dst(&p), None);
}
#[test]
fn learned_lookup_prefers_exact_client_over_subnet() {
use shadowvpn::mesh::RouteApproval;
let mut learned = Learned::default();
let peer_a: SocketAddr = "198.51.100.1:1000".parse().unwrap();
let peer_b: SocketAddr = "198.51.100.2:2000".parse().unwrap();
learned.learn("10.77.0.2".parse().unwrap(), peer_a, "");
learned.subnets.advertise(
peer_b,
&["10.77.0.0/16".parse().unwrap()],
&RouteApproval {
auto: true,
allowlist: vec![],
},
Instant::now(),
);
assert_eq!(learned.lookup("10.77.0.2".parse().unwrap()), Some(peer_a));
assert_eq!(learned.lookup("10.77.9.9".parse().unwrap()), Some(peer_b));
assert_eq!(learned.lookup("192.0.2.1".parse().unwrap()), None);
}
#[test]
fn control_handling_learns_and_replies_to_accepting_clients() {
use shadowvpn::mesh::{RouteAdvert, RouteApproval};
let routing: Shared = Arc::new(Mutex::new(Routing::Learn(Learned::default())));
let approval = RouteApproval {
auto: false,
allowlist: vec!["192.168.200.0/24".parse().unwrap()],
};
let now = Instant::now();
let router_peer: SocketAddr = "198.51.100.1:1000".parse().unwrap();
let client_peer: SocketAddr = "198.51.100.2:2000".parse().unwrap();
let advert = RouteAdvert {
tunnel_ip: "10.77.0.2".parse().unwrap(),
tunnel_ip6: Some("fd07:7::2".parse().unwrap()),
accept_routes: false,
routes: vec![
"192.168.200.0/24".parse().unwrap(),
"10.99.0.0/16".parse().unwrap(),
],
};
assert_eq!(
handle_control(&routing, &approval, router_peer, &advert.encode(), now),
None
);
let advert = RouteAdvert {
tunnel_ip: "10.77.0.3".parse().unwrap(),
tunnel_ip6: None,
accept_routes: true,
routes: vec![],
};
let push = handle_control(&routing, &approval, client_peer, &advert.encode(), now)
.expect("accepting client gets a push");
assert_eq!(push.routes, vec!["192.168.200.0/24".parse().unwrap()]);
let guard = routing.lock().unwrap();
let Routing::Learn(learned) = &*guard else {
panic!("learning mode")
};
assert_eq!(
learned.lookup("10.77.0.2".parse().unwrap()),
Some(router_peer)
);
assert_eq!(
learned.lookup("fd07:7::2".parse().unwrap()),
Some(router_peer)
);
assert_eq!(
learned.lookup("192.168.200.7".parse().unwrap()),
Some(router_peer)
);
assert_eq!(learned.lookup("10.99.1.1".parse().unwrap()), None);
}
}