use std::{
io,
net::{self},
sync::Arc,
};
use async_trait::async_trait;
use sciparse::{
address::ip_socket_addr::ScionSocketIpAddr,
core::view::View,
dataplane_path::view::ScionDpPathViewExt,
identifier::isd_asn::IsdAsn,
packet::view::{ScionPacketView, ScionRawPacketView},
path::metadata::path_interface::PathInterface,
};
use tokio::net::UdpSocket;
use crate::{
stack::{ScionSocketReceiveError, ScionSocketSendError, UnderlaySocket},
underlays::{discovery::UnderlayDiscovery, outbound_ip_towards},
};
#[async_trait::async_trait]
pub trait OutboundIpResolver: Send + Sync {
async fn outbound_ips(&self) -> Vec<net::IpAddr>;
}
#[async_trait::async_trait]
impl OutboundIpResolver for Vec<net::IpAddr> {
async fn outbound_ips(&self) -> Vec<net::IpAddr> {
self.clone()
}
}
pub struct TargetAddrOutboundIpResolver {
url: url::Url,
host_overrides: Vec<(String, net::IpAddr)>,
}
impl TargetAddrOutboundIpResolver {
pub fn new(url: url::Url, host_overrides: Vec<(String, net::IpAddr)>) -> Self {
Self {
url,
host_overrides,
}
}
}
#[async_trait::async_trait]
impl OutboundIpResolver for TargetAddrOutboundIpResolver {
async fn outbound_ips(&self) -> Vec<net::IpAddr> {
let url = &self.url;
let host = match url.host_str() {
Some(h) => h,
None => return vec![],
};
let port = url.port_or_known_default().unwrap_or(0);
let target = if let Some((_, ip)) = self.host_overrides.iter().find(|(h, _)| h == host) {
net::SocketAddr::new(*ip, port)
} else if let Ok(ip) = host.parse::<net::IpAddr>() {
net::SocketAddr::new(ip, port)
} else {
match url
.socket_addrs(|| None)
.ok()
.and_then(|addrs| addrs.into_iter().next())
{
Some(addr) => addr,
None => {
tracing::warn!(url = %url, "failed to resolve API URL for local IP detection");
return vec![];
}
}
};
match outbound_ip_towards(target).await {
Some(ip) => vec![ip],
None => vec![],
}
}
}
pub(crate) struct UdpUnderlaySocket {
pub(crate) socket: UdpSocket,
pub(crate) bind_addr: ScionSocketIpAddr,
pub(crate) underlay_discovery: Arc<dyn UnderlayDiscovery>,
}
impl UdpUnderlaySocket {
pub(crate) fn new(
socket: UdpSocket,
bind_addr: ScionSocketIpAddr,
underlay_discovery: Arc<dyn UnderlayDiscovery>,
) -> Self {
Self {
socket,
bind_addr,
underlay_discovery,
}
}
fn resolve_local_dispatch_addr(
&self,
packet: &ScionPacketView,
) -> Result<net::SocketAddr, ScionSocketSendError> {
let dst_addr = packet
.header()
.dst_host_addr()
.map_err(|e| {
crate::stack::ScionSocketSendError::InvalidPacket(
format!("Packet for local dispatch had invalid destination address: {e}")
.into(),
)
})?
.ip()
.ok_or(crate::stack::ScionSocketSendError::InvalidPacket(
"Packet for local dispatch dst was not an IP address".into(),
))?;
let classified = packet.try_classify().map_err(|e| {
crate::stack::ScionSocketSendError::InvalidPacket(
format!("Failed to classify packet for local dispatch: {e}").into(),
)
})?;
let dst_port =
classified
.dst_port()
.ok_or(crate::stack::ScionSocketSendError::InvalidPacket(
"Coulldn't determine destination port in packet for local dispatch".into(),
))?;
Ok(net::SocketAddr::new(dst_addr, dst_port))
}
fn try_dispatch_local(&self, packet: &ScionPacketView) -> Result<(), ScionSocketSendError> {
let dst_addr = self.resolve_local_dispatch_addr(packet)?;
self.socket
.try_send_to(packet.as_slice(), dst_addr)
.map_err(|e| Self::map_send_io_error(e, packet.header().src_ia(), 0, dst_addr))?;
Ok(())
}
fn map_send_io_error(
e: io::Error,
src: IsdAsn,
interface_id: u16,
next_hop: net::SocketAddr,
) -> ScionSocketSendError {
use std::io::ErrorKind::{
BrokenPipe, ConnectionAborted, ConnectionReset, HostUnreachable, NetworkUnreachable,
};
match e.kind() {
HostUnreachable | NetworkUnreachable => {
ScionSocketSendError::UnderlayNextHopUnreachable {
isd_as: src,
interface_id,
address: Some(next_hop),
msg: e.to_string(),
}
}
ConnectionAborted | ConnectionReset | BrokenPipe => ScionSocketSendError::Closed,
_ => ScionSocketSendError::IoError(e),
}
}
}
#[async_trait]
impl UnderlaySocket for UdpUnderlaySocket {
fn try_send(&self, packet: &ScionRawPacketView) -> Result<(), ScionSocketSendError> {
let source_ia = packet.header().src_ia();
if packet.header().dst_ia() == source_ia {
return self.try_dispatch_local(packet);
}
let Some(egress_if) = packet.header().path().first_egress_interface() else {
return Err(ScionSocketSendError::InvalidPacket(
"Can't determine egress interface for packet.".into(),
));
};
let next_hop = match self
.underlay_discovery
.resolve_udp_underlay_next_hop(PathInterface {
isd_asn: source_ia,
id: egress_if,
}) {
Some(next_hop) => next_hop,
None => {
return Err(ScionSocketSendError::UnderlayNextHopUnreachable {
isd_as: source_ia,
interface_id: egress_if,
address: None,
msg: "next hop not found".to_string(),
});
}
};
self.socket
.try_send_to(packet.as_slice(), next_hop)
.map_err(|e| Self::map_send_io_error(e, source_ia, egress_if, next_hop))?;
Ok(())
}
async fn writeable(&self) {
let _ = self.socket.writable().await;
}
fn try_recv(&self, buf: &mut [u8]) -> Result<usize, ScionSocketReceiveError> {
loop {
let n = match self.socket.try_recv_from(buf) {
Ok((n, _)) => n,
Err(e) => return Err(ScionSocketReceiveError::IoError(e)),
};
match ScionRawPacketView::try_from_slice(&buf[..n]) {
Ok((packet, _rest)) => {
match packet.dst_scion_addr() {
Ok(dst) if dst == self.bind_addr.scion_addr() => return Ok(n),
Ok(_) => {
tracing::debug!(
"Packet destination does not match assigned address, skipping"
);
}
Err(e) => {
tracing::error!(error = %e, "Failed to read destination address from packet");
}
}
}
Err(e) => {
tracing::error!(error = %e, "Failed to decode SCION packet");
}
}
}
}
async fn readable(&self) {
let _ = self.socket.readable().await;
}
}