use std::net::{Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4};
use std::sync::{Arc, RwLock};
use std::time::Duration;
use futures::channel::mpsc;
use futures::StreamExt;
use nym_smol_core::{ChannelDevice, DnsConfig, Stack, StackConfig, TcpStream, UdpSocket};
use crate::connectors::TunnelConnector;
use crate::topup::{
event_channel, run_topup, BandwidthCredentialSource, BandwidthEvent, TopupConfig,
};
use tokio::sync::{broadcast, watch};
use tokio::task::JoinHandle;
use tokio_util::sync::CancellationToken;
use tracing::{debug, info, warn};
use crate::bridge::{self, BridgeParams};
use crate::config::{DnsMode, MtuConfig, PeerConfig, TunnelConfig};
use crate::engine::WgEngine;
use crate::error::{DvpnError, Result};
use crate::transport::{direct_transport, SocketProtector, WgReceiver, WgSender};
const TIMER_INTERVAL: Duration = Duration::from_millis(250);
const TRANSPORT_ERROR_BACKOFF: Duration = Duration::from_millis(100);
type SharedStack = Arc<RwLock<Arc<Stack>>>;
type StackChannels = (
mpsc::UnboundedReceiver<Vec<u8>>,
mpsc::UnboundedSender<Vec<u8>>,
);
fn build_stack(
assigned: Ipv4Addr,
ipv6: Option<Ipv6Addr>,
dns: DnsMode,
interface_mtu: usize,
) -> (Stack, StackChannels) {
let (stack_out_tx, stack_out_rx) = mpsc::unbounded::<Vec<u8>>();
let (stack_in_tx, stack_in_rx) = mpsc::unbounded::<Vec<u8>>();
let device = ChannelDevice::new(stack_in_rx, stack_out_tx, Some(interface_mtu));
let mut stack_config = StackConfig::new(assigned);
if let Some(v6) = ipv6 {
stack_config = stack_config.with_ipv6(v6);
}
let mut stack = Stack::new(device, stack_config);
if let DnsMode::InTunnelServer(server) = dns {
stack = stack.with_dns_config(DnsConfig {
server,
..DnsConfig::default()
});
}
(stack, (stack_out_rx, stack_in_tx))
}
#[derive(Clone, Debug)]
enum TransportChoice {
Direct,
Quic(BridgeParams),
}
struct TopupSpec {
config: TopupConfig,
source: Option<Arc<dyn BandwidthCredentialSource>>,
}
pub struct TunnelBuilder {
entry: PeerConfig,
exit: Option<PeerConfig>,
config: TunnelConfig,
transport: TransportChoice,
protector: Option<SocketProtector>,
cancel: Option<CancellationToken>,
topup: Option<TopupSpec>,
}
impl TunnelBuilder {
pub fn single_hop(gateway: PeerConfig) -> Self {
Self {
entry: gateway,
exit: None,
config: TunnelConfig::default(),
transport: TransportChoice::Direct,
protector: None,
cancel: None,
topup: None,
}
}
pub fn two_hop(entry: PeerConfig, exit: PeerConfig) -> Self {
Self {
entry,
exit: Some(exit),
config: TunnelConfig::default(),
transport: TransportChoice::Direct,
protector: None,
cancel: None,
topup: None,
}
}
pub fn config(mut self, config: TunnelConfig) -> Self {
self.config = config;
self
}
pub fn quic_bridge(mut self, params: BridgeParams) -> Self {
self.transport = TransportChoice::Quic(params);
self
}
pub fn socket_protector(mut self, protector: SocketProtector) -> Self {
self.protector = Some(protector);
self
}
pub fn cancellation_token(mut self, token: CancellationToken) -> Self {
self.cancel = Some(token);
self
}
pub fn bandwidth_topup(
mut self,
config: TopupConfig,
source: Arc<dyn BandwidthCredentialSource>,
) -> Self {
self.topup = Some(TopupSpec {
config,
source: Some(source),
});
self
}
pub fn bandwidth_monitor(mut self, config: TopupConfig) -> Self {
self.topup = Some(TopupSpec {
config,
source: None,
});
self
}
pub async fn connect(self) -> Result<Tunnel> {
Tunnel::connect(self).await
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct NotEstablished {
pub entry: bool,
pub exit: Option<bool>,
}
impl std::fmt::Display for NotEstablished {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self.exit {
None => write!(f, "wireguard session not established"),
Some(exit) => write!(
f,
"wireguard session(s) not established (entry: {}, exit: {})",
if self.entry { "up" } else { "down" },
if exit { "up" } else { "down" },
),
}
}
}
impl std::error::Error for NotEstablished {}
pub struct Tunnel {
stack: SharedStack,
cancel: CancellationToken,
task: Option<JoinHandle<()>>,
topup_task: Option<JoinHandle<()>>,
assigned: Ipv4Addr,
ipv6: Option<Ipv6Addr>,
dns: DnsMode,
two_hop: bool,
mtu: RwLock<MtuConfig>,
swap_tx: mpsc::UnboundedSender<StackChannels>,
swap_lock: std::sync::Mutex<()>,
events_tx: broadcast::Sender<BandwidthEvent>,
established_rx: watch::Receiver<(bool, Option<bool>)>,
}
impl Tunnel {
async fn connect(builder: TunnelBuilder) -> Result<Tunnel> {
let cancel = builder.cancel.unwrap_or_default();
let two_hop = builder.exit.is_some();
if matches!(builder.transport, TransportChoice::Quic(_)) && !two_hop {
return Err(DvpnError::QuicRequiresTwoHop);
}
let (assigned, ipv6, interface_mtu) = if let Some(exit) = &builder.exit {
(
exit.assigned_ipv4,
exit.assigned_ipv6,
builder.config.mtu.exit,
)
} else {
(
builder.entry.assigned_ipv4,
builder.entry.assigned_ipv6,
builder.config.mtu.entry,
)
};
if cancel.is_cancelled() {
return Err(DvpnError::Cancelled);
}
let (stack, (stack_out_rx, stack_in_tx)) =
build_stack(assigned, ipv6, builder.config.dns, interface_mtu);
let (swap_tx, swap_rx) = mpsc::unbounded::<StackChannels>();
let engine = if let Some(exit) = &builder.exit {
let exit_endpoint = as_socket_addr_v4(exit.endpoint, "exit gateway endpoint")?;
let tunnel_src =
SocketAddrV4::new(builder.entry.assigned_ipv4, builder.config.exit_client_port);
WgEngine::two_hop(&builder.entry, exit, tunnel_src, exit_endpoint)
} else {
WgEngine::single_hop(&builder.entry)
};
let (sender, receiver) = match &builder.transport {
TransportChoice::Direct => {
direct_transport(builder.entry.endpoint, builder.protector.as_ref()).await?
}
TransportChoice::Quic(params) => {
let (s, r) = bridge::connect(params, &cancel).await?;
(WgSender::Quic(s), WgReceiver::Quic(r))
}
};
info!(
two_hop,
quic = matches!(builder.transport, TransportChoice::Quic(_)),
assigned = %assigned,
mtu = interface_mtu,
entry_endpoint = %builder.entry.endpoint,
exit_endpoint = builder.exit.as_ref().map(|e| e.endpoint.to_string()),
"dVPN tunnel datapath starting"
);
let (established_tx, established_rx) = watch::channel((false, two_hop.then_some(false)));
let task = tokio::spawn(datapath(
engine,
sender,
receiver,
(stack_out_rx, stack_in_tx),
swap_rx,
established_tx,
cancel.clone(),
));
let shared_stack: SharedStack = Arc::new(RwLock::new(Arc::new(stack)));
let events_tx = event_channel();
let topup_task = builder.topup.map(|spec| {
let connector = TunnelConnector::new(shared_stack.clone());
tokio::spawn(run_topup(
spec.config,
spec.source,
connector,
events_tx.clone(),
cancel.clone(),
))
});
Ok(Tunnel {
stack: shared_stack,
cancel,
task: Some(task),
topup_task,
assigned,
ipv6,
dns: builder.config.dns,
two_hop,
mtu: RwLock::new(builder.config.mtu),
swap_tx,
swap_lock: std::sync::Mutex::new(()),
events_tx,
established_rx,
})
}
pub async fn await_established(
&self,
timeout: Duration,
) -> std::result::Result<(), NotEstablished> {
let mut rx = self.established_rx.clone();
let established = |(entry, exit): &(bool, Option<bool>)| *entry && exit.unwrap_or(true);
let done = matches!(
tokio::time::timeout(timeout, rx.wait_for(established)).await,
Ok(Ok(_))
);
if done {
Ok(())
} else {
let (entry, exit) = *rx.borrow();
Err(NotEstablished { entry, exit })
}
}
pub fn mtu(&self) -> MtuConfig {
*self.mtu.read().expect("mtu lock poisoned")
}
pub fn set_mtu(&self, mtu: MtuConfig) -> Result<()> {
let interface_mtu = if self.two_hop { mtu.exit } else { mtu.entry };
let (stack, channels) = build_stack(self.assigned, self.ipv6, self.dns, interface_mtu);
let _swap = self.swap_lock.lock().expect("swap lock poisoned");
self.swap_tx
.unbounded_send(channels)
.map_err(|_| DvpnError::Transport("datapath has stopped".into()))?;
*self.stack.write().expect("stack lock poisoned") = Arc::new(stack);
*self.mtu.write().expect("mtu lock poisoned") = mtu;
info!(mtu = interface_mtu, "tunnel MTU updated at runtime");
Ok(())
}
pub fn bandwidth_events(&self) -> broadcast::Receiver<BandwidthEvent> {
self.events_tx.subscribe()
}
pub async fn tcp_connect(&self, addr: SocketAddr) -> Result<TcpStream> {
Ok(self.stack().tcp_connect(addr).await?)
}
pub async fn udp_socket(&self) -> Result<UdpSocket> {
Ok(self.stack().udp_socket().await?)
}
pub async fn resolve(&self, host: &str) -> Result<Vec<std::net::IpAddr>> {
Ok(self.stack().resolve(host).await?)
}
pub async fn tcp_connect_host(&self, host: &str, port: u16) -> Result<TcpStream> {
Ok(self.stack().tcp_connect_host(host, port).await?)
}
pub fn stack(&self) -> Arc<Stack> {
self.stack.read().expect("stack lock poisoned").clone()
}
pub fn connector(&self) -> TunnelConnector {
TunnelConnector::new(self.stack.clone())
}
pub async fn shutdown(mut self) {
self.cancel.cancel();
stop_task(self.task.take()).await;
stop_task(self.topup_task.take()).await;
info!("dVPN tunnel shut down");
}
}
const SHUTDOWN_GRACE: Duration = Duration::from_secs(2);
async fn stop_task(task: Option<JoinHandle<()>>) {
let Some(mut handle) = task else { return };
if tokio::time::timeout(SHUTDOWN_GRACE, &mut handle)
.await
.is_err()
{
handle.abort();
let _ = handle.await;
}
}
impl Drop for Tunnel {
fn drop(&mut self) {
self.cancel.cancel();
}
}
fn as_socket_addr_v4(addr: SocketAddr, what: &str) -> Result<SocketAddrV4> {
match addr {
SocketAddr::V4(v4) => Ok(v4),
SocketAddr::V6(_) => Err(DvpnError::Config(format!(
"{what} must be IPv4 for two-hop inner framing"
))),
}
}
async fn datapath(
mut engine: WgEngine,
mut sender: WgSender,
mut receiver: WgReceiver,
stack_channels: StackChannels,
mut swap_rx: mpsc::UnboundedReceiver<StackChannels>,
established_tx: watch::Sender<(bool, Option<bool>)>,
cancel: CancellationToken,
) {
let (mut stack_out_rx, mut stack_in_tx) = stack_channels;
let init = engine.initiate_handshakes();
for pkt in init.to_network {
if let Err(e) = sender.send(&pkt).await {
warn!("failed to send handshake init: {e}");
}
}
let mut ticker = tokio::time::interval(TIMER_INTERVAL);
ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
loop {
tokio::select! {
biased;
maybe_swap = swap_rx.next() => {
match maybe_swap {
Some((new_out_rx, new_in_tx)) => {
debug!("datapath swapping stack channels (runtime MTU change)");
stack_out_rx = new_out_rx;
stack_in_tx = new_in_tx;
}
None => break,
}
}
_ = cancel.cancelled() => {
debug!("datapath cancelled");
break;
}
maybe_app = stack_out_rx.next() => {
match maybe_app {
Some(app) => {
let out = engine.encapsulate_app(&app);
send_all(&mut sender, out.to_network).await;
}
None => break, }
}
incoming = receiver.recv() => {
match incoming {
Ok(wg) => {
let out = engine.decapsulate_incoming(&wg);
let est = engine.establishment();
established_tx.send_if_modified(|v| {
let changed = *v != est;
if changed { *v = est; }
changed
});
for pkt in out.to_stack {
if stack_in_tx.unbounded_send(pkt).is_err() {
debug!("stack channel closed");
return;
}
}
send_all(&mut sender, out.to_network).await;
}
Err(e) => {
if receiver.is_bridge() {
warn!("bridge transport failed, stopping datapath: {e}");
break;
}
debug!("transport recv error (transient): {e}");
tokio::time::sleep(TRANSPORT_ERROR_BACKOFF).await;
}
}
}
_ = ticker.tick() => {
let out = engine.update_timers();
send_all(&mut sender, out.to_network).await;
}
}
}
}
async fn send_all(sender: &mut WgSender, packets: Vec<Vec<u8>>) {
for pkt in packets {
if let Err(e) = sender.send(&pkt).await {
warn!("transport send error: {e}");
}
}
}