use crate::ip::{DataType, IpPacket, TcpFlags};
use std::io::{self, Read, Write};
use std::net::Ipv4Addr;
use std::num::NonZeroUsize;
#[derive(Debug)]
pub enum Action {
Pass,
Send(Option<Vec<u8>>),
Recv(Option<NonZeroUsize>),
Break,
}
pub trait Arbiter {
fn decide(&mut self, packet: &IpPacket) -> Action;
fn update(&mut self, data: &[u8], original: &[u8]);
}
pub struct TCPAddressArbiter {
client: Ipv4Addr,
host: Ipv4Addr,
}
impl TCPAddressArbiter {
#[must_use]
pub fn new(client: Ipv4Addr, host: Ipv4Addr) -> Self {
Self { client, host }
}
}
impl Arbiter for TCPAddressArbiter {
fn decide(&mut self, packet: &IpPacket) -> Action {
const CONN_END: TcpFlags = TcpFlags::RST.union(TcpFlags::FIN);
if let DataType::TCP(ref payload) = packet.payload {
if self.host == packet.dest && self.client == packet.source {
if payload.flags.contains(TcpFlags::PSH) {
return Action::Send(None);
}
if payload.flags.intersects(CONN_END) {
return Action::Break;
}
}
if self.host == packet.source && self.client == packet.dest {
if payload.flags.contains(TcpFlags::PSH) {
return Action::Recv(None);
}
if payload.flags.intersects(CONN_END) {
return Action::Break;
}
}
}
Action::Pass
}
fn update(&mut self, _data: &[u8], _original: &[u8]) {}
}
pub struct UDPAddressArbiter {
client: Ipv4Addr,
host: Ipv4Addr,
}
impl UDPAddressArbiter {
#[must_use]
pub fn new(client: Ipv4Addr, host: Ipv4Addr) -> Self {
Self { client, host }
}
}
impl Arbiter for UDPAddressArbiter {
fn decide(&mut self, packet: &IpPacket) -> Action {
if let DataType::UDP(_) = packet.payload {
if self.host == packet.dest && self.client == packet.source {
return Action::Send(None);
}
if self.host == packet.source && self.client == packet.dest {
return Action::Recv(None);
}
}
Action::Pass
}
fn update(&mut self, _data: &[u8], _original: &[u8]) {}
}
pub struct Replayer<C: Read + Write> {
socket: C,
packets: Vec<IpPacket>,
arbiter: Box<dyn Arbiter>,
}
impl<C: Read + Write> Replayer<C> {
#[must_use]
pub fn new(socket: C, packets: Vec<IpPacket>, arbiter: Box<dyn Arbiter>) -> Self {
Self {
socket,
packets,
arbiter,
}
}
pub fn replay(mut self) -> io::Result<C> {
let mut recv_buf = Box::new([0u8; 65536]);
for packet in &self.packets {
match self.arbiter.decide(packet) {
Action::Pass => {}
Action::Send(payload) => {
if let Some(data) = payload {
self.socket.write_all(&data)?;
} else {
self.socket.write_all(packet.payload.get_payload())?;
}
}
Action::Recv(read_size) => {
let read_bytes: usize;
if let Some(size) = read_size {
read_bytes = size.into();
self.socket.read_exact(&mut recv_buf[..read_bytes])?;
} else {
read_bytes = self.socket.read(&mut recv_buf[..])?;
}
self.arbiter
.update(&recv_buf[..read_bytes], packet.payload.get_payload());
}
Action::Break => break,
}
}
Ok(self.socket)
}
#[must_use]
pub fn into_socket(self) -> C {
self.socket
}
}