use crate::{
address::socket_addr::ScionSocketAddr,
core::encode::WireEncode,
identifier::isd_asn::IsdAsn,
packet::{
model::{ScionRawPacket, ScionScmpPacket, ScionUdpPacket},
view::{ScionPacketView, ScionRawPacketView, ScionScmpPacketView, ScionUdpPacketView},
},
};
#[derive(Debug, Clone, Copy)]
pub enum ClassifiedPacketView<'a> {
Udp(&'a ScionUdpPacketView),
Scmp(&'a ScionScmpPacketView),
Other(&'a ScionPacketView),
}
impl<'a> ClassifiedPacketView<'a> {
#[inline]
pub fn dst_socket_addr(&self) -> Option<ScionSocketAddr> {
let dst_port = self.dst_port()?;
let scion_addr = self.as_raw().dst_scion_addr().ok()?;
Some(ScionSocketAddr::new(
scion_addr.isd_asn(),
scion_addr.host(),
dst_port,
))
}
#[inline]
pub fn dst_port(&self) -> Option<u16> {
match self {
ClassifiedPacketView::Udp(udp) => Some(udp.udp().dst_port()),
ClassifiedPacketView::Scmp(scmp) => scmp.scmp().dst_port(),
ClassifiedPacketView::Other(_) => None,
}
}
#[inline]
pub const fn is_udp(&self) -> bool {
matches!(self, ClassifiedPacketView::Udp(_))
}
#[inline]
pub const fn is_scmp(&self) -> bool {
matches!(self, ClassifiedPacketView::Scmp(_))
}
#[inline]
pub const fn is_other(&self) -> bool {
matches!(self, ClassifiedPacketView::Other(_))
}
#[inline]
pub fn as_raw(&self) -> &'a ScionRawPacketView {
match self {
ClassifiedPacketView::Udp(udp) => udp.as_raw(),
ClassifiedPacketView::Scmp(scmp) => scmp.as_raw(),
ClassifiedPacketView::Other(other) => other,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub enum ClassifiedPacket {
Scmp(ScionScmpPacket),
Udp(ScionUdpPacket),
Other(ScionRawPacket),
}
impl ClassifiedPacket {
#[inline]
pub fn dst_socket_addr(&self) -> Option<ScionSocketAddr> {
let (header, dst_port) = match self {
ClassifiedPacket::Udp(packet) => (&packet.header, packet.payload.dst_port),
ClassifiedPacket::Scmp(packet) => (&packet.header, packet.payload.dst_port()?),
ClassifiedPacket::Other(_) => return None,
};
let isd_asn: IsdAsn = header.address.dst_ia;
let host = header.address.dst_host_addr.scion_host_addr().ok()?;
Some(ScionSocketAddr::new(isd_asn, host, dst_port))
}
#[inline]
pub fn into_raw(self) -> ScionRawPacket {
match self {
ClassifiedPacket::Udp(packet) => packet.into_raw(),
ClassifiedPacket::Scmp(packet) => packet.into_raw(),
ClassifiedPacket::Other(packet) => packet,
}
}
#[inline]
#[allow(clippy::result_large_err)]
pub fn try_into_udp(self) -> Result<ScionUdpPacket, Self> {
match self {
ClassifiedPacket::Udp(packet) => Ok(packet),
_ => Err(self),
}
}
#[inline]
#[allow(clippy::result_large_err)]
pub fn try_into_scmp(self) -> Result<ScionScmpPacket, Self> {
match self {
ClassifiedPacket::Scmp(packet) => Ok(packet),
_ => Err(self),
}
}
#[inline]
pub const fn header(&self) -> &crate::header::model::ScionPacketHeader {
match self {
ClassifiedPacket::Udp(packet) => &packet.header,
ClassifiedPacket::Scmp(packet) => &packet.header,
ClassifiedPacket::Other(packet) => &packet.header,
}
}
}
impl WireEncode for ClassifiedPacket {
#[inline]
fn wire_valid(&self) -> Result<(), crate::core::encode::InvalidStructureError> {
match self {
ClassifiedPacket::Udp(packet) => packet.wire_valid(),
ClassifiedPacket::Scmp(packet) => packet.wire_valid(),
ClassifiedPacket::Other(packet) => packet.wire_valid(),
}
}
#[inline]
unsafe fn encode_unchecked(&self, buf: &mut [u8]) -> usize {
unsafe {
match self {
ClassifiedPacket::Udp(packet) => packet.encode_unchecked(buf),
ClassifiedPacket::Scmp(packet) => packet.encode_unchecked(buf),
ClassifiedPacket::Other(packet) => packet.encode_unchecked(buf),
}
}
}
#[inline]
fn required_size(&self) -> usize {
match self {
ClassifiedPacket::Udp(packet) => packet.required_size(),
ClassifiedPacket::Scmp(packet) => packet.required_size(),
ClassifiedPacket::Other(packet) => packet.required_size(),
}
}
}
impl From<ClassifiedPacket> for ScionRawPacket {
#[inline]
fn from(packet: ClassifiedPacket) -> Self {
packet.into_raw()
}
}
impl TryFrom<ClassifiedPacket> for ScionUdpPacket {
type Error = ClassifiedPacket;
#[inline]
fn try_from(packet: ClassifiedPacket) -> Result<Self, Self::Error> {
packet.try_into_udp()
}
}
impl TryFrom<ClassifiedPacket> for ScionScmpPacket {
type Error = ClassifiedPacket;
#[inline]
fn try_from(packet: ClassifiedPacket) -> Result<Self, Self::Error> {
packet.try_into_scmp()
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum ClassifyError {
#[error("malformed UDP payload: {0}")]
MalformedUdp(crate::core::view::ViewConversionError),
#[error("malformed SCMP payload: {0}")]
MalformedScmp(crate::core::view::ViewConversionError),
}
#[cfg(feature = "proptest")]
pub mod ptest {
use proptest::prelude::*;
use super::*;
use crate::{
header::model::ScionPacketHeader,
packet::model::ScionPacket,
payload::{ProtocolNumber, scmp::model::ScmpMessage, udp::model::UdpDatagram},
};
#[derive(Debug, Clone)]
pub struct ArbitraryClassifiedPacketParams {
pub udp: u32,
pub scmp: u32,
pub other: u32,
pub header_params: <ScionPacketHeader as Arbitrary>::Parameters,
pub scmp_params: <ScmpMessage as Arbitrary>::Parameters,
}
impl Default for ArbitraryClassifiedPacketParams {
fn default() -> Self {
Self {
udp: 1,
scmp: 1,
other: 1,
header_params: Default::default(),
scmp_params: Default::default(),
}
}
}
impl Arbitrary for ClassifiedPacket {
type Parameters = ArbitraryClassifiedPacketParams;
type Strategy = BoxedStrategy<Self>;
fn arbitrary_with(params: Self::Parameters) -> Self::Strategy {
let hp = params.header_params;
prop_oneof![
params.udp => (
ScionPacketHeader::arbitrary_with(hp.clone()),
any::<UdpDatagram>(),
)
.prop_map(|(mut header, payload)| {
header.common.next_header = ProtocolNumber::Udp;
ClassifiedPacket::Udp(ScionPacket { header, payload })
}),
params.scmp => (
ScionPacketHeader::arbitrary_with(hp.clone()),
ScmpMessage::arbitrary_with(params.scmp_params),
)
.prop_map(|(mut header, payload)| {
header.common.next_header = ProtocolNumber::Scmp;
ClassifiedPacket::Scmp(ScionPacket { header, payload })
}),
params.other => (
ScionPacketHeader::arbitrary_with(hp),
proptest::collection::vec(any::<u8>(), 0..2048),
)
.prop_map(|(mut header, payload)| {
header.common.next_header = match header.common.next_header {
ProtocolNumber::Udp => ProtocolNumber::Other(255), ProtocolNumber::Scmp => ProtocolNumber::Other(255), v => v,
};
ClassifiedPacket::Other(ScionPacket { header, payload })
}),
]
.boxed()
}
}
}
#[cfg(test)]
mod tests {
use crate::{
address::socket_addr::ScionSocketAddr,
core::{encode::WireEncode, view::View},
dataplane_path::model::DpPath,
packet::{
classify::ClassifiedPacketView,
model::{ScionRawPacket, ScionScmpPacket, ScionUdpPacket},
view::ScionPacketView,
},
payload::{
ProtocolNumber,
scmp::model::{ScmpDestinationUnreachable, ScmpEchoReply},
},
};
#[test]
fn classify_udp_packet_succeeds() {
let buf = ScionUdpPacket::new(
"[1-ff00:0:110,10.0.0.1]:12345".parse().unwrap(),
"[1-ff00:0:111,10.0.0.2]:54321".parse().unwrap(),
DpPath::Empty,
b"payload".to_vec(),
)
.try_encode_to_vec()
.expect("failed to encode SCION UDP packet");
let (view, _) = ScionPacketView::try_from_slice(&buf).unwrap();
let classified = view.try_classify().unwrap();
match &classified {
ClassifiedPacketView::Udp { .. } => {}
_ => panic!("expected Udp variant"),
}
assert_eq!(
"[1-ff00:0:111,10.0.0.2]:54321"
.parse::<ScionSocketAddr>()
.unwrap(),
classified.dst_socket_addr().unwrap()
);
}
#[test]
fn classify_scmp_echo_reply_with_port() {
let scion_scmp_packet = ScionScmpPacket::new(
"1-ff00:0:110,10.0.0.1".parse().unwrap(),
"1-ff00:0:111,10.0.0.2".parse().unwrap(),
DpPath::Empty,
ScmpEchoReply::new(
54321, 1,
b"echo data".to_vec(),
)
.into(),
);
let buf = scion_scmp_packet
.try_encode_to_vec()
.expect("failed to encode SCION SCMP packet");
let (view, _) = ScionPacketView::try_from_slice(&buf).unwrap();
let classified = view.try_classify().unwrap();
match &classified {
ClassifiedPacketView::Scmp(scmp) => {
assert_eq!(Some(54321), scmp.scmp().dst_port());
}
_ => panic!("expected Scmp variant"),
}
assert_eq!(
"[1-ff00:0:111,10.0.0.2]:54321"
.parse::<ScionSocketAddr>()
.unwrap(),
classified.dst_socket_addr().unwrap()
);
}
#[test]
fn classify_scmp_destination_unreachable_with_parsable_payload() {
let quoted_udp = ScionUdpPacket::new(
"[1-ff00:0:111,10.0.0.2]:54321".parse().unwrap(),
"[1-ff00:0:110,10.0.0.1]:12345".parse().unwrap(),
DpPath::Empty,
b"quoted payload".to_vec(),
);
let quoted_udp_data = quoted_udp
.try_encode_to_vec()
.expect("failed to encode quoted UDP packet");
let scmp_packet = ScionScmpPacket::new(
"1-ff00:0:110,10.0.0.1".parse().unwrap(),
"1-ff00:0:111,10.0.0.2".parse().unwrap(),
DpPath::Empty,
ScmpDestinationUnreachable::new(
crate::payload::scmp::types::ScmpDestinationUnreachableCode::AddressUnreachable,
quoted_udp_data,
)
.into(),
);
let buf = scmp_packet.try_encode_to_vec().unwrap();
let (view, _) = ScionPacketView::try_from_slice(&buf).unwrap();
let classified = view.try_classify().unwrap();
match &classified {
ClassifiedPacketView::Scmp(scmp) => {
assert!(scmp.scmp().dst_port().is_some());
}
_ => panic!("expected Scmp variant"),
}
assert_eq!(
"[1-ff00:0:111,10.0.0.2]:54321"
.parse::<ScionSocketAddr>()
.unwrap(),
classified.dst_socket_addr().unwrap()
);
}
#[test]
fn classify_scmp_destination_unreachable_without_parsable_payload() {
let scmp_packet = ScionScmpPacket::new(
"1-ff00:0:110,10.0.0.1".parse().unwrap(),
"1-ff00:0:111,10.0.0.2".parse().unwrap(),
DpPath::Empty,
ScmpDestinationUnreachable::new(
crate::payload::scmp::types::ScmpDestinationUnreachableCode::AddressUnreachable,
b"not a valid quoted UDP packet".to_vec(),
)
.into(),
);
let buf = scmp_packet.try_encode_to_vec().unwrap();
let (view, _) = ScionPacketView::try_from_slice(&buf).unwrap();
let classified = view.try_classify().unwrap();
match &classified {
ClassifiedPacketView::Scmp(scmp) => {
assert!(scmp.scmp().dst_port().is_none());
}
_ => panic!("expected Scmp variant"),
}
assert!(classified.dst_socket_addr().is_none());
}
#[test]
fn classify_unknown_next_header_returns_other() {
let bytes = ScionRawPacket::new(
"1-ff00:0:110,10.0.0.1".parse().unwrap(),
"1-ff00:0:111,10.0.0.2".parse().unwrap(),
DpPath::Empty,
ProtocolNumber::Other(0xFE),
b"rawr".to_vec(),
)
.try_encode_to_vec()
.expect("failed to encode SCION raw packet");
let (view, _) = ScionPacketView::try_from_slice(&bytes).unwrap();
assert!(matches!(
view.try_classify().unwrap(),
ClassifiedPacketView::Other(_)
));
assert!(view.try_classify().unwrap().dst_socket_addr().is_none());
}
}