use std::{net, sync::Arc};
use ana_gotatun::packet::PacketBufPool;
use reqwest_connect_rpc::token_source::TokenSource;
use sciparse::{address::ip_socket_addr::ScionSocketIpAddr, identifier::isd_asn::IsdAsn};
use snap_tun::client::{PACKET_BUF_POOL_SIZE, SnapTunEndpoint};
use socket2::{Domain, Protocol, Socket, Type};
use tokio::net::UdpSocket;
use url::Url;
use x25519_dalek::StaticSecret;
use crate::{
stack::{
BoundUnderlaySocket, DynUnderlayStack, InvalidBindAddressError, ScionSocketBindError,
SnapConnectionError, UnderlaySocket, builder::PreferredUnderlay,
},
underlays::{
discovery::{UnderlayDiscovery, UnderlayInfo},
udp::{OutboundIpResolver, UdpUnderlaySocket},
},
};
pub mod discovery;
pub mod snap;
pub mod udp;
pub(crate) struct SnapSocketConfig {
pub crpc_client: Option<reqwest::Client>,
pub snap_token_source: Option<Arc<dyn TokenSource>>,
}
pub(crate) struct UnderlayStack {
preferred_underlay: PreferredUnderlay,
underlay_discovery: Arc<dyn UnderlayDiscovery>,
outbound_ip_resolver: Arc<dyn OutboundIpResolver>,
snap_socket_config: SnapSocketConfig,
snap_tunnel_manager: Option<SnapTunEndpoint>,
pool: PacketBufPool<PACKET_BUF_POOL_SIZE>,
}
impl UnderlayStack {
pub fn new(
preferred_underlay: PreferredUnderlay,
underlay_discovery: Arc<dyn UnderlayDiscovery>,
outbound_ip_resolver: Arc<dyn OutboundIpResolver>,
static_identity: StaticSecret,
default_snap_socket_config: SnapSocketConfig,
) -> Self {
let snap_tunnel_manager = default_snap_socket_config
.snap_token_source
.as_ref()
.map(|token_source| SnapTunEndpoint::new(token_source.clone(), static_identity));
Self {
preferred_underlay,
underlay_discovery,
outbound_ip_resolver,
snap_socket_config: default_snap_socket_config,
snap_tunnel_manager,
pool: PacketBufPool::new(64),
}
}
fn select_underlay(&self, requested_isd_as: IsdAsn) -> Option<(IsdAsn, UnderlayInfo)> {
let underlays = self.underlay_discovery.underlays(requested_isd_as);
match self.preferred_underlay {
PreferredUnderlay::Snap => {
if let Some(underlay) = underlays
.iter()
.find(|(_, underlay)| matches!(underlay, UnderlayInfo::Snap(_)))
{
return Some(underlay.clone());
}
}
PreferredUnderlay::Udp => {
if let Some(underlay) = underlays
.iter()
.find(|(_, underlay)| matches!(underlay, UnderlayInfo::Udp(_)))
{
return Some(underlay.clone());
}
}
}
underlays.into_iter().next()
}
async fn bind_snap_socket(
&self,
requested_addr: Option<ScionSocketIpAddr>,
isd_as: IsdAsn,
cp_url: Url,
) -> Result<snap::SnapUnderlaySocket, ScionSocketBindError> {
let (Some(token_source), Some(snap_tunnel_manager)) = (
self.snap_socket_config.snap_token_source.as_ref(),
self.snap_tunnel_manager.as_ref(),
) else {
return Err(ScionSocketBindError::SnapConnectionError(
SnapConnectionError::SnapTokenSourceMissing,
))?;
};
let local_addr = match requested_addr {
Some(addr) => addr.socket_addr(),
None => {
if let Some(cp_addr) = cp_url
.socket_addrs(|| None)
.ok()
.and_then(|addrs| addrs.first().copied())
&& let Some(ip) = outbound_ip_towards(cp_addr).await
{
Ok(net::SocketAddr::new(ip, 0))
} else {
Err(ScionSocketBindError::InvalidBindAddress(
InvalidBindAddressError::NoLocalIpAddressFound,
))
}?
}
};
let bind_addr = ScionSocketIpAddr::new(isd_as, local_addr.ip(), local_addr.port());
let udp_socket = bind_udp_underlay_socket(local_addr)?;
let socket = snap::SnapUnderlaySocket::new(
bind_addr,
cp_url,
udp_socket,
snap_tunnel_manager,
token_source.clone(),
1024,
self.pool.clone(),
self.snap_socket_config.crpc_client.clone(),
)
.await?;
let assigned_addr = socket.local_addr();
if let Some(requested_addr) = requested_addr
&& requested_addr.isd_asn().matches(assigned_addr.isd_asn())
&& let requested_socket_addr = requested_addr.socket_addr()
&& let assigned_socket_addr = assigned_addr.socket_addr()
&& ((!requested_socket_addr.ip().is_unspecified() && assigned_socket_addr.ip() != requested_socket_addr.ip())
|| (requested_socket_addr.port() != 0 && assigned_socket_addr.port() != requested_socket_addr.port()))
{
return Err(crate::stack::ScionSocketBindError::InvalidBindAddress(
crate::stack::InvalidBindAddressError::AddressMismatch {
assigned_addr: ScionSocketIpAddr::new(
bind_addr.isd_asn(),
requested_socket_addr.ip(),
requested_socket_addr.port(),
),
bind_addr,
},
));
}
Ok(socket)
}
async fn resolve_udp_bind_addr(
&self,
isd_as: IsdAsn,
override_addr: Option<ScionSocketIpAddr>,
) -> Result<ScionSocketIpAddr, ScionSocketBindError> {
if let Some(addr) = override_addr {
return Ok(addr);
}
let local_address = *self
.outbound_ip_resolver
.outbound_ips()
.await
.first()
.ok_or(ScionSocketBindError::InvalidBindAddress(
InvalidBindAddressError::NoLocalIpAddressFound,
))?;
Ok(ScionSocketIpAddr::new(isd_as, local_address, 0))
}
async fn bind_udp_socket(
&self,
isd_as: IsdAsn,
bind_addr: Option<ScionSocketIpAddr>,
) -> Result<(ScionSocketIpAddr, UdpSocket), ScionSocketBindError> {
let bind_addr = self.resolve_udp_bind_addr(isd_as, bind_addr).await?;
let local_addr: net::SocketAddr = bind_addr.socket_addr();
let socket = bind_udp_underlay_socket(local_addr)?;
let local_addr = socket.local_addr().map_err(|e| {
ScionSocketBindError::Other(
anyhow::anyhow!("failed to get local address: {e}").into_boxed_dyn_error(),
)
})?;
let bind_addr =
ScionSocketIpAddr::new(bind_addr.isd_asn(), local_addr.ip(), local_addr.port());
Ok((bind_addr, socket))
}
}
impl DynUnderlayStack for UnderlayStack {
fn bind_socket(
&self,
_kind: crate::stack::SocketKind,
bind_addr: Option<ScionSocketIpAddr>,
) -> futures::future::BoxFuture<'_, Result<BoundUnderlaySocket, ScionSocketBindError>> {
Box::pin(async move {
let requested_isd_as = bind_addr.map_or(IsdAsn::WILDCARD, |addr| addr.isd_asn());
match self.select_underlay(requested_isd_as) {
Some((isd_as, UnderlayInfo::Snap(cp_url))) => {
let socket = self.bind_snap_socket(bind_addr, isd_as, cp_url).await?;
Ok(BoundUnderlaySocket {
local_addr: socket.local_addr(),
snap_data_plane: socket.snap_data_plane(),
socket: Box::new(socket) as Box<dyn UnderlaySocket>,
})
}
Some((isd_as, UnderlayInfo::Udp(_))) => {
let (bind_addr, socket) = self.bind_udp_socket(isd_as, bind_addr).await?;
Ok(BoundUnderlaySocket {
local_addr: bind_addr,
snap_data_plane: None,
socket: Box::new(UdpUnderlaySocket::new(
socket,
bind_addr,
self.underlay_discovery.clone(),
)) as Box<dyn UnderlaySocket>,
})
}
None => {
Err(ScionSocketBindError::NoUnderlayAvailable(
requested_isd_as.isd(),
))
}
}
})
}
fn local_ases(&self) -> Vec<IsdAsn> {
let mut isd_ases: Vec<IsdAsn> = self.underlay_discovery.isd_ases().into_iter().collect();
isd_ases.sort();
isd_ases
}
}
#[cfg(windows)]
fn set_exclusive_addr_use(sock: &Socket, enable: bool) -> std::io::Result<()> {
use std::{mem, os::windows::io::AsRawSocket};
use windows_sys::Win32::Networking::WinSock;
let val: u32 = if enable { 1 } else { 0 };
let rc = unsafe {
WinSock::setsockopt(
sock.as_raw_socket() as usize,
WinSock::SOL_SOCKET,
WinSock::SO_EXCLUSIVEADDRUSE,
&val as *const _ as *const _,
mem::size_of_val(&val) as _,
)
};
if rc == 0 {
Ok(())
} else {
Err(std::io::Error::last_os_error())
}
}
fn bind_udp_underlay_socket(
addr: net::SocketAddr,
) -> Result<tokio::net::UdpSocket, ScionSocketBindError> {
let socket = Socket::new(Domain::for_address(addr), Type::DGRAM, Some(Protocol::UDP))
.map_err(|e| ScionSocketBindError::Other(Box::new(e)))?;
socket
.set_nonblocking(true)
.map_err(|e| ScionSocketBindError::Other(Box::new(e)))?;
if addr.is_ipv6()
&& let Err(e) = socket.set_only_v6(false)
{
tracing::debug!(%e, "unable to make socket dual-stack");
}
#[cfg(windows)]
set_exclusive_addr_use(&socket, true).map_err(|e| ScionSocketBindError::Other(Box::new(e)))?;
socket.bind(&addr.into()).map_err(|e| {
match e.kind() {
std::io::ErrorKind::AddrInUse => ScionSocketBindError::PortAlreadyInUse(addr.port()),
std::io::ErrorKind::AddrNotAvailable | std::io::ErrorKind::InvalidInput => {
ScionSocketBindError::InvalidBindAddress(
InvalidBindAddressError::CannotBindToRequestedAddress(
ScionSocketIpAddr::new(IsdAsn::WILDCARD, addr.ip(), addr.port()),
format!("Failed to bind socket: {e:#}").into(),
),
)
}
#[cfg(windows)]
std::io::ErrorKind::PermissionDenied => {
ScionSocketBindError::PortAlreadyInUse(addr.port())
}
_ => ScionSocketBindError::Other(Box::new(e)),
}
})?;
tokio::net::UdpSocket::from_std(std::net::UdpSocket::from(socket))
.map_err(|e| ScionSocketBindError::Other(Box::new(e)))
}
pub(crate) async fn outbound_ip_towards(dst: net::SocketAddr) -> Option<net::IpAddr> {
let bind_addr = match dst.ip() {
net::IpAddr::V4(_) => net::Ipv4Addr::UNSPECIFIED.into(),
net::IpAddr::V6(_) => net::Ipv6Addr::UNSPECIFIED.into(),
};
if let Ok(socket) = tokio::net::UdpSocket::bind(net::SocketAddr::new(bind_addr, 0)).await
&& socket.connect(dst).await.is_ok()
&& let Ok(addr) = socket.local_addr()
{
return Some(addr.ip());
}
None
}