use std::net::SocketAddrV4;
use boringtun::noise::{Tunn, TunnResult};
use crate::config::PeerConfig;
use crate::framing::{build_ipv4_udp, parse_ipv4_udp};
const SCRATCH_LEN: usize = 65535;
#[derive(Default)]
pub(crate) struct EngineOutput {
pub to_stack: Vec<Vec<u8>>,
pub to_network: Vec<Vec<u8>>,
}
fn make_tunn(peer: &PeerConfig, index: u32) -> Tunn {
Tunn::new(
peer.client_private_key.inner().clone(),
peer.gateway_public_key.inner(),
peer.preshared_key,
None,
index,
None,
)
}
fn encap(tunn: &mut Tunn, scratch: &mut [u8], plaintext: &[u8]) -> Option<Vec<u8>> {
match tunn.encapsulate(plaintext, scratch) {
TunnResult::WriteToNetwork(p) => Some(p.to_vec()),
TunnResult::Err(e) => {
tracing::warn!("wireguard encapsulate error: {e:?}");
None
}
_ => None,
}
}
fn decap(tunn: &mut Tunn, scratch: &mut [u8], datagram: &[u8]) -> (Vec<Vec<u8>>, Vec<Vec<u8>>) {
let mut inner = Vec::new();
let mut net = Vec::new();
let queued = match tunn.decapsulate(None, datagram, scratch) {
TunnResult::WriteToTunnelV4(p, _) | TunnResult::WriteToTunnelV6(p, _) => {
inner.push(p.to_vec());
false
}
TunnResult::WriteToNetwork(p) => {
net.push(p.to_vec());
true
}
TunnResult::Err(e) => {
tracing::warn!("wireguard decapsulate error: {e:?}");
false
}
_ => false,
};
if queued {
while let TunnResult::WriteToNetwork(p) = tunn.decapsulate(None, &[], scratch) {
net.push(p.to_vec());
}
}
(inner, net)
}
fn timer(tunn: &mut Tunn, scratch: &mut [u8]) -> Option<Vec<u8>> {
match tunn.update_timers(scratch) {
TunnResult::WriteToNetwork(p) => Some(p.to_vec()),
_ => None,
}
}
fn handshake_init(tunn: &mut Tunn, scratch: &mut [u8]) -> Option<Vec<u8>> {
match tunn.format_handshake_initiation(scratch, false) {
TunnResult::WriteToNetwork(p) => Some(p.to_vec()),
_ => None,
}
}
fn wrap_for_exit(
entry: &mut Tunn,
scratch: &mut [u8],
tunnel_src: SocketAddrV4,
exit_endpoint: SocketAddrV4,
exit_packet: &[u8],
) -> Option<Vec<u8>> {
let carrier = build_ipv4_udp(tunnel_src, exit_endpoint, exit_packet);
encap(entry, scratch, &carrier)
}
#[allow(clippy::large_enum_variant)]
enum Inner {
SingleHop {
tunn: Tunn,
},
TwoHop {
entry: Tunn,
exit: Tunn,
tunnel_src: SocketAddrV4,
exit_endpoint: SocketAddrV4,
},
}
pub(crate) struct WgEngine {
inner: Inner,
scratch: Box<[u8]>,
saw_inbound: bool,
entry_established: bool,
exit_established: bool,
}
impl WgEngine {
pub(crate) fn single_hop(peer: &PeerConfig) -> Self {
WgEngine {
inner: Inner::SingleHop {
tunn: make_tunn(peer, 0),
},
scratch: vec![0u8; SCRATCH_LEN].into_boxed_slice(),
saw_inbound: false,
entry_established: false,
exit_established: false,
}
}
pub(crate) fn two_hop(
entry: &PeerConfig,
exit: &PeerConfig,
tunnel_src: SocketAddrV4,
exit_endpoint: SocketAddrV4,
) -> Self {
WgEngine {
inner: Inner::TwoHop {
entry: make_tunn(entry, 0),
exit: make_tunn(exit, 1),
tunnel_src,
exit_endpoint,
},
scratch: vec![0u8; SCRATCH_LEN].into_boxed_slice(),
saw_inbound: false,
entry_established: false,
exit_established: false,
}
}
pub(crate) fn establishment(&self) -> (bool, Option<bool>) {
match &self.inner {
Inner::SingleHop { .. } => (self.entry_established, None),
Inner::TwoHop { .. } => (self.entry_established, Some(self.exit_established)),
}
}
fn note_progress(&mut self) {
match &self.inner {
Inner::SingleHop { tunn } => {
if !self.entry_established && tunn.stats().0.is_some() {
self.entry_established = true;
tracing::info!("wireguard session established");
}
}
Inner::TwoHop { entry, exit, .. } => {
if !self.entry_established && entry.stats().0.is_some() {
self.entry_established = true;
tracing::info!("entry-hop wireguard session established");
}
if !self.exit_established && exit.stats().0.is_some() {
self.exit_established = true;
tracing::info!("exit-hop wireguard session established");
}
}
}
}
pub(crate) fn encapsulate_app(&mut self, app: &[u8]) -> EngineOutput {
let mut out = EngineOutput::default();
let WgEngine { inner, scratch, .. } = self;
match inner {
Inner::SingleHop { tunn } => {
if let Some(p) = encap(tunn, scratch, app) {
out.to_network.push(p);
}
}
Inner::TwoHop {
entry,
exit,
tunnel_src,
exit_endpoint,
} => {
if let Some(c_exit) = encap(exit, scratch, app) {
if let Some(c_entry) =
wrap_for_exit(entry, scratch, *tunnel_src, *exit_endpoint, &c_exit)
{
out.to_network.push(c_entry);
}
}
}
}
out
}
pub(crate) fn decapsulate_incoming(&mut self, wg: &[u8]) -> EngineOutput {
if !self.saw_inbound {
self.saw_inbound = true;
tracing::info!("first datagram received from the entry transport");
}
let mut out = EngineOutput::default();
let WgEngine { inner, scratch, .. } = self;
match inner {
Inner::SingleHop { tunn } => {
let (inner_pkts, net) = decap(tunn, scratch, wg);
out.to_stack = inner_pkts;
out.to_network = net;
}
Inner::TwoHop {
entry,
exit,
tunnel_src,
exit_endpoint,
} => {
let (carriers, entry_net) = decap(entry, scratch, wg);
out.to_network.extend(entry_net); for carrier in carriers {
let Some(parsed) = parse_ipv4_udp(&carrier) else {
continue;
};
if parsed.src != *exit_endpoint {
continue;
}
let (app_pkts, exit_net) = decap(exit, scratch, &parsed.payload);
out.to_stack.extend(app_pkts);
for c_exit in exit_net {
if let Some(c_entry) =
wrap_for_exit(entry, scratch, *tunnel_src, *exit_endpoint, &c_exit)
{
out.to_network.push(c_entry);
}
}
}
}
}
self.note_progress();
out
}
pub(crate) fn update_timers(&mut self) -> EngineOutput {
let mut out = EngineOutput::default();
let WgEngine { inner, scratch, .. } = self;
match inner {
Inner::SingleHop { tunn } => {
if let Some(p) = timer(tunn, scratch) {
out.to_network.push(p);
}
}
Inner::TwoHop {
entry,
exit,
tunnel_src,
exit_endpoint,
} => {
if let Some(p) = timer(entry, scratch) {
out.to_network.push(p);
}
if let Some(c_exit) = timer(exit, scratch) {
if let Some(c_entry) =
wrap_for_exit(entry, scratch, *tunnel_src, *exit_endpoint, &c_exit)
{
out.to_network.push(c_entry);
}
}
}
}
out
}
pub(crate) fn initiate_handshakes(&mut self) -> EngineOutput {
let mut out = EngineOutput::default();
let WgEngine { inner, scratch, .. } = self;
match inner {
Inner::SingleHop { tunn } => {
if let Some(p) = handshake_init(tunn, scratch) {
out.to_network.push(p);
}
}
Inner::TwoHop {
entry,
exit,
tunnel_src,
exit_endpoint,
} => {
if let Some(p) = handshake_init(entry, scratch) {
out.to_network.push(p);
}
if let Some(c_exit) = handshake_init(exit, scratch) {
if let Some(c_entry) =
wrap_for_exit(entry, scratch, *tunnel_src, *exit_endpoint, &c_exit)
{
out.to_network.push(c_entry);
}
}
}
}
out
}
}