use std::mem::MaybeUninit;
use std::net::{Ipv4Addr, SocketAddr, UdpSocket};
use std::time::{Duration, Instant};
pub const TCP_FIN: u8 = 0x01;
pub const TCP_SYN: u8 = 0x02;
pub const TCP_RST: u8 = 0x04;
pub const TCP_ACK: u8 = 0x10;
const TCP_HEADER_LEN: usize = 20;
const IPV4_HEADER_LEN: usize = 20;
const IPPROTO_TCP: u8 = 6;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SynOutcome {
Open,
Closed,
Filtered,
}
pub fn classify_flags(flags: u8) -> SynOutcome {
if flags & TCP_RST != 0 {
SynOutcome::Closed
} else if flags & TCP_SYN != 0 && flags & TCP_ACK != 0 {
SynOutcome::Open
} else {
SynOutcome::Filtered
}
}
fn checksum(data: &[u8]) -> u16 {
let mut sum: u32 = 0;
let mut chunks = data.chunks_exact(2);
for pair in &mut chunks {
sum += u32::from(u16::from_be_bytes([pair[0], pair[1]]));
}
if let [last] = chunks.remainder() {
sum += u32::from(u16::from_be_bytes([*last, 0]));
}
while sum >> 16 != 0 {
sum = (sum & 0xffff) + (sum >> 16);
}
!(sum as u16)
}
pub fn build_tcp_syn(
src: Ipv4Addr,
dst: Ipv4Addr,
src_port: u16,
dst_port: u16,
seq: u32,
) -> Vec<u8> {
let mut seg = vec![0u8; TCP_HEADER_LEN];
seg[0..2].copy_from_slice(&src_port.to_be_bytes());
seg[2..4].copy_from_slice(&dst_port.to_be_bytes());
seg[4..8].copy_from_slice(&seq.to_be_bytes());
seg[12] = 0x50; seg[13] = TCP_SYN;
seg[14..16].copy_from_slice(&1024u16.to_be_bytes());
let mut pseudo = Vec::with_capacity(12 + TCP_HEADER_LEN);
pseudo.extend_from_slice(&src.octets());
pseudo.extend_from_slice(&dst.octets());
pseudo.push(0);
pseudo.push(IPPROTO_TCP);
pseudo.extend_from_slice(&(TCP_HEADER_LEN as u16).to_be_bytes());
pseudo.extend_from_slice(&seg);
let sum = checksum(&pseudo);
seg[16..18].copy_from_slice(&sum.to_be_bytes());
seg
}
pub fn build_ipv4_syn(
src: Ipv4Addr,
dst: Ipv4Addr,
src_port: u16,
dst_port: u16,
seq: u32,
) -> Vec<u8> {
let tcp = build_tcp_syn(src, dst, src_port, dst_port, seq);
let total_len = (IPV4_HEADER_LEN + tcp.len()) as u16;
let mut ip = vec![0u8; IPV4_HEADER_LEN];
ip[0] = 0x45; ip[2..4].copy_from_slice(&total_len.to_be_bytes());
ip[6..8].copy_from_slice(&0x4000u16.to_be_bytes()); ip[8] = 64; ip[9] = IPPROTO_TCP;
ip[12..16].copy_from_slice(&src.octets());
ip[16..20].copy_from_slice(&dst.octets());
let sum = checksum(&ip);
ip[10..12].copy_from_slice(&sum.to_be_bytes());
ip.extend_from_slice(&tcp);
ip
}
pub fn parse_ipv4_tcp(packet: &[u8]) -> Option<(u16, u16, u8)> {
if packet.len() < IPV4_HEADER_LEN {
return None;
}
if packet[0] >> 4 != 4 {
return None;
}
let ihl = (packet[0] & 0x0f) as usize * 4;
if ihl < IPV4_HEADER_LEN || packet.len() < ihl {
return None;
}
if packet[9] != IPPROTO_TCP {
return None;
}
let tcp = &packet[ihl..];
if tcp.len() < TCP_HEADER_LEN {
return None;
}
let src_port = u16::from_be_bytes([tcp[0], tcp[1]]);
let dst_port = u16::from_be_bytes([tcp[2], tcp[3]]);
let flags = tcp[13];
Some((src_port, dst_port, flags))
}
pub fn raw_socket_available() -> bool {
use socket2::{Domain, Protocol, Socket, Type};
Socket::new(
Domain::IPV4,
Type::RAW,
Some(Protocol::from(i32::from(IPPROTO_TCP))),
)
.is_ok()
}
fn local_ipv4_for(dst: Ipv4Addr) -> Option<Ipv4Addr> {
let socket = UdpSocket::bind("0.0.0.0:0").ok()?;
socket.connect(SocketAddr::new(dst.into(), 80)).ok()?;
match socket.local_addr().ok()?.ip() {
std::net::IpAddr::V4(v4) => Some(v4),
std::net::IpAddr::V6(_) => None,
}
}
pub fn syn_scan_port(dst: Ipv4Addr, dst_port: u16, timeout: Duration) -> Option<SynOutcome> {
use socket2::{Domain, Protocol, Socket, Type};
crate::rate::gate();
let src = local_ipv4_for(dst)?;
let src_port = 40000u16.wrapping_add(dst_port);
let segment = build_tcp_syn(src, dst, src_port, dst_port, 0x1234_5678);
let socket = Socket::new(
Domain::IPV4,
Type::RAW,
Some(Protocol::from(i32::from(IPPROTO_TCP))),
)
.ok()?;
socket.set_read_timeout(Some(timeout)).ok()?;
let destination: socket2::SockAddr = SocketAddr::new(dst.into(), dst_port).into();
socket.send_to(&segment, &destination).ok()?;
let deadline = Instant::now() + timeout;
let mut buf = [MaybeUninit::<u8>::uninit(); 1500];
while Instant::now() < deadline {
let Ok(n) = socket.recv(&mut buf) else {
break;
};
let packet = unsafe { &*(&buf[..n] as *const [MaybeUninit<u8>] as *const [u8]) };
if let Some((reply_src, reply_dst, flags)) = parse_ipv4_tcp(packet)
&& reply_src == dst_port
&& reply_dst == src_port
{
return Some(classify_flags(flags));
}
}
Some(SynOutcome::Filtered)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn classify_covers_open_closed_filtered() {
assert_eq!(classify_flags(TCP_SYN | TCP_ACK), SynOutcome::Open);
assert_eq!(classify_flags(TCP_RST | TCP_ACK), SynOutcome::Closed);
assert_eq!(classify_flags(TCP_RST), SynOutcome::Closed);
assert_eq!(classify_flags(TCP_ACK), SynOutcome::Filtered);
assert_eq!(classify_flags(0), SynOutcome::Filtered);
}
#[test]
fn checksum_of_a_valid_block_verifies_to_zero() {
let ip = build_ipv4_syn(
"10.0.0.1".parse().unwrap(),
"10.0.0.2".parse().unwrap(),
40000,
80,
0x11223344,
);
assert_eq!(checksum(&ip[..IPV4_HEADER_LEN]), 0);
}
#[test]
fn tcp_syn_has_expected_shape_and_valid_checksum() {
let src: Ipv4Addr = "192.168.1.10".parse().unwrap();
let dst: Ipv4Addr = "192.168.1.20".parse().unwrap();
let seg = build_tcp_syn(src, dst, 50000, 443, 0xdeadbeef);
assert_eq!(seg.len(), TCP_HEADER_LEN);
assert_eq!(u16::from_be_bytes([seg[0], seg[1]]), 50000);
assert_eq!(u16::from_be_bytes([seg[2], seg[3]]), 443);
assert_eq!(seg[12] >> 4, 5, "data offset should be 5 words");
assert_eq!(seg[13], TCP_SYN, "only the SYN flag should be set");
let mut pseudo = Vec::new();
pseudo.extend_from_slice(&src.octets());
pseudo.extend_from_slice(&dst.octets());
pseudo.push(0);
pseudo.push(IPPROTO_TCP);
pseudo.extend_from_slice(&(TCP_HEADER_LEN as u16).to_be_bytes());
pseudo.extend_from_slice(&seg);
assert_eq!(checksum(&pseudo), 0);
}
#[test]
fn ipv4_syn_total_length_and_protocol_are_correct() {
let ip = build_ipv4_syn(
"1.2.3.4".parse().unwrap(),
"5.6.7.8".parse().unwrap(),
33333,
22,
1,
);
assert_eq!(ip.len(), IPV4_HEADER_LEN + TCP_HEADER_LEN);
assert_eq!(ip[0] >> 4, 4, "version 4");
assert_eq!(
u16::from_be_bytes([ip[2], ip[3]]) as usize,
IPV4_HEADER_LEN + TCP_HEADER_LEN
);
assert_eq!(ip[9], IPPROTO_TCP);
}
#[test]
fn build_then_parse_round_trips_ports_and_flags() {
let ip = build_ipv4_syn(
"10.1.1.1".parse().unwrap(),
"10.1.1.2".parse().unwrap(),
44444,
8080,
42,
);
let (src_port, dst_port, flags) = parse_ipv4_tcp(&ip).expect("should parse");
assert_eq!(src_port, 44444);
assert_eq!(dst_port, 8080);
assert_eq!(flags, TCP_SYN);
}
#[test]
fn parse_rejects_non_ipv4_and_non_tcp() {
assert!(parse_ipv4_tcp(&[]).is_none());
assert!(parse_ipv4_tcp(&[0x60; 40]).is_none());
let mut pkt = build_ipv4_syn(
"1.1.1.1".parse().unwrap(),
"2.2.2.2".parse().unwrap(),
1,
2,
3,
);
pkt[9] = 17;
assert!(parse_ipv4_tcp(&pkt).is_none());
}
#[test]
fn parse_handles_ip_options() {
let mut pkt = build_ipv4_syn(
"9.9.9.9".parse().unwrap(),
"8.8.8.8".parse().unwrap(),
12345,
53,
7,
);
let tcp = pkt.split_off(IPV4_HEADER_LEN);
pkt.extend_from_slice(&[0u8; 4]); pkt.extend_from_slice(&tcp);
pkt[0] = 0x46; let new_len = pkt.len() as u16;
pkt[2..4].copy_from_slice(&new_len.to_be_bytes());
let (src_port, dst_port, flags) = parse_ipv4_tcp(&pkt).expect("should parse with options");
assert_eq!(src_port, 12345);
assert_eq!(dst_port, 53);
assert_eq!(flags, TCP_SYN);
}
}