#![cfg(feature = "slirp")]
use std::net::{IpAddr, Ipv4Addr};
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use pktkit::vtcp::segment::Segment;
use pktkit::vtcp::{Conn, ConnConfig, State};
use pktkit::{IpPrefix, L3Device, Packet, Protocol};
const CLIENT_IP: Ipv4Addr = Ipv4Addr::new(10, 0, 0, 5); const SERVER_IP: Ipv4Addr = Ipv4Addr::new(10, 0, 0, 1); const SERVER_PORT: u16 = 8080;
const CLIENT_PORT: u16 = 51000;
fn wrap(src: Ipv4Addr, dst: Ipv4Addr, seg: &[u8]) -> Vec<u8> {
let total = 20 + seg.len();
let mut ip = vec![0u8; total];
ip[0] = 0x45;
ip[2..4].copy_from_slice(&(total as u16).to_be_bytes());
ip[8] = 64;
ip[9] = Protocol::TCP.as_u8();
ip[12..16].copy_from_slice(&src.octets());
ip[16..20].copy_from_slice(&dst.octets());
let cs = pktkit::checksum(&ip[..20]);
ip[10..12].copy_from_slice(&cs.to_be_bytes());
ip[20..].copy_from_slice(seg);
ip
}
#[test]
fn inbound_accept_handshake_and_bidirectional_data() {
let stack = pktkit::slirp::Stack::new();
stack
.set_addr(IpPrefix::new(IpAddr::V4(SERVER_IP), 24))
.unwrap();
let listener = stack
.listen("tcp", &format!("{SERVER_IP}:{SERVER_PORT}"))
.unwrap();
let client = Arc::new(Mutex::new(Conn::new(ConnConfig {
local_port: CLIENT_PORT,
remote_port: SERVER_PORT,
mss: 1460,
..Default::default()
})));
let stack_for_handler = stack.clone();
let client_for_handler = client.clone();
stack.set_handler(Arc::new(move |pkt: &Packet| {
let bytes = pkt.as_bytes();
if bytes.len() < 40 || bytes[9] != Protocol::TCP.as_u8() {
return Ok(());
}
let seg = match Segment::parse(&bytes[20..]) {
Ok(s) => s,
Err(_) => return Ok(()),
};
let replies = {
let mut c = client_for_handler.lock().unwrap();
c.handle_segment(&seg)
};
for r in replies {
let ip = wrap(CLIENT_IP, SERVER_IP, &r);
let _ = stack_for_handler.send(Packet::from_slice(&ip));
}
Ok(())
}));
{
let segs = client.lock().unwrap().connect();
for s in segs {
let ip = wrap(CLIENT_IP, SERVER_IP, &s);
stack.send(Packet::from_slice(&ip)).unwrap();
}
}
{
let deadline = Instant::now() + Duration::from_secs(2);
loop {
if client.lock().unwrap().state() == State::Established {
break;
}
assert!(Instant::now() < deadline, "client never established");
std::thread::sleep(Duration::from_millis(10));
}
}
let server_conn = {
let l = listener.clone();
let (tx, rx) = std::sync::mpsc::channel();
std::thread::spawn(move || {
let _ = tx.send(l.accept());
});
rx.recv_timeout(Duration::from_secs(3))
.expect("accept timed out")
.expect("accept failed")
};
assert_eq!(server_conn.local_addr().port(), SERVER_PORT);
assert_eq!(server_conn.peer_addr().port(), CLIENT_PORT);
assert_eq!(server_conn.peer_addr().ip(), IpAddr::V4(CLIENT_IP));
server_conn.write(b"hello from server").unwrap();
{
let deadline = Instant::now() + Duration::from_secs(2);
let mut buf = [0u8; 64];
let got = loop {
let n = client.lock().unwrap().read(&mut buf);
if n > 0 {
break buf[..n].to_vec();
}
assert!(
Instant::now() < deadline,
"client never received server data"
);
std::thread::sleep(Duration::from_millis(10));
};
assert_eq!(&got, b"hello from server");
let acks = {
let mut c = client.lock().unwrap();
c.tick()
};
for a in acks {
let ip = wrap(CLIENT_IP, SERVER_IP, &a);
stack.send(Packet::from_slice(&ip)).unwrap();
}
}
{
let (n, segs) = client.lock().unwrap().write(b"hi from client");
assert_eq!(n, b"hi from client".len());
for s in segs {
let ip = wrap(CLIENT_IP, SERVER_IP, &s);
stack.send(Packet::from_slice(&ip)).unwrap();
}
}
server_conn.set_read_timeout(Some(Duration::from_secs(2)));
let mut buf = [0u8; 64];
let n = server_conn.read(&mut buf).expect("server read");
assert!(n > 0, "expected client data on server side");
assert_eq!(&buf[..n], b"hi from client");
let _ = server_conn.close();
let _ = listener.close();
let _ = stack.shutdown();
}