use super::{Error, Result, be32, put_be32};
pub const PEER_ERROR_FILE_SYSTEM: u32 = 4000;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub struct KeepAlivePacket {
pub timestamp: u32,
pub dest_socket_id: u32,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub struct CongestionWarningPacket {
pub timestamp: u32,
pub dest_socket_id: u32,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub struct ShutdownPacket {
pub timestamp: u32,
pub dest_socket_id: u32,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub struct AckAckPacket {
pub ack_number: u32,
pub timestamp: u32,
pub dest_socket_id: u32,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub struct DropReqPacket {
pub message_number: u32,
pub timestamp: u32,
pub dest_socket_id: u32,
pub first_seq: u32,
pub last_seq: u32,
}
impl DropReqPacket {
pub(crate) fn parse_cif(
message_number: u32,
timestamp: u32,
dest_socket_id: u32,
cif: &[u8],
) -> Result<Self> {
if cif.len() != 8 {
return Err(Error::BufferTooShort {
need: 8,
have: cif.len(),
what: "drop request CIF",
});
}
Ok(DropReqPacket {
message_number,
timestamp,
dest_socket_id,
first_seq: be32(cif, 0),
last_seq: be32(cif, 4),
})
}
pub(crate) fn cif_len(&self) -> usize {
8
}
pub(crate) fn write_cif(&self, buf: &mut [u8]) {
put_be32(buf, 0, self.first_seq);
put_be32(buf, 4, self.last_seq);
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub struct PeerErrorPacket {
pub error_code: u32,
pub timestamp: u32,
pub dest_socket_id: u32,
}
#[cfg(test)]
mod tests {
use alloc::vec::Vec;
use super::super::control::ControlPacket;
use super::*;
#[test]
fn drop_req_round_trips() {
let d = DropReqPacket {
message_number: 7,
timestamp: 100,
dest_socket_id: 200,
first_seq: 10,
last_seq: 20,
};
let pkt = ControlPacket::DropReq(d);
let mut buf = [0u8; 24];
let n = pkt.serialize_into(&mut buf).unwrap();
assert_eq!(n, 24);
assert_eq!(&buf[4..8], &7u32.to_be_bytes()); assert_eq!(&buf[16..20], &10u32.to_be_bytes());
assert_eq!(&buf[20..24], &20u32.to_be_bytes());
let parsed = ControlPacket::parse(&buf).unwrap();
assert_eq!(parsed, pkt);
}
#[test]
fn keepalive_congestion_shutdown_ackack_peererror_round_trip() {
let cases: Vec<ControlPacket> = alloc::vec![
ControlPacket::KeepAlive(KeepAlivePacket {
timestamp: 1,
dest_socket_id: 2
}),
ControlPacket::CongestionWarning(CongestionWarningPacket {
timestamp: 3,
dest_socket_id: 4
}),
ControlPacket::Shutdown(ShutdownPacket {
timestamp: 5,
dest_socket_id: 6
}),
ControlPacket::AckAck(AckAckPacket {
ack_number: 9,
timestamp: 7,
dest_socket_id: 8
}),
ControlPacket::PeerError(PeerErrorPacket {
error_code: PEER_ERROR_FILE_SYSTEM,
timestamp: 11,
dest_socket_id: 12
}),
];
for pkt in cases {
let mut buf = [0u8; 16];
let n = pkt.serialize_into(&mut buf).unwrap();
assert_eq!(n, 16);
let parsed = ControlPacket::parse(&buf).unwrap();
assert_eq!(parsed, pkt);
}
}
#[test]
fn keepalive_rejects_trailing_bytes() {
let mut buf = [0u8; 17];
buf[0] = 0x80; buf[1] = 0x01; assert!(matches!(
ControlPacket::parse(&buf),
Err(Error::UnexpectedTrailingBytes { .. })
));
}
}