Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions .gitattributes
Original file line number Diff line number Diff line change
@@ -0,0 +1,2 @@
*.sh text eol=lf
*.pkt text eol=lf
85 changes: 72 additions & 13 deletions src/tcp/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -69,7 +69,14 @@ pub struct TcpListener {
ip_stack: IpStack,
packet_receiver: Receiver<TransportPacket>,
local_addr: Option<SocketAddr>,
tcb_map: HashMap<NetworkTuple, (Tcb, Instant)>,
tcb_map: HashMap<NetworkTuple, HalfOpenTcb>,
}

struct HalfOpenTcb {
tcb: Tcb,
created_at: Instant,
retransmit_at: Instant,
rto: Duration,
}

/// A TCP stream between a local and a remote socket.
Expand Down Expand Up @@ -171,23 +178,37 @@ impl TcpListener {
}
pub async fn accept(&mut self) -> io::Result<(TcpStream, SocketAddr)> {
loop {
if let Some(packet) = self.packet_receiver.recv().await {
let packet = if let Some(deadline) = self.tcb_map.values().map(|entry| entry.retransmit_at).min() {
tokio::select! {
packet = self.packet_receiver.recv() => packet,
_ = tokio::time::sleep_until(deadline.into()) => {
self.retransmit_half_open().await?;
continue;
}
}
} else {
self.packet_receiver.recv().await
};
if let Some(packet) = packet {
let network_tuple = &packet.network_tuple;
if let Some(v) = self.ip_stack.inner.tcp_stream_map.get(network_tuple).as_deref().cloned() {
// If a TCP stream has already been generated, hand it over to the corresponding stream
_ = v.send(packet).await;
continue;
}
let Some(tcp_packet) = pnet_packet::tcp::TcpPacket::new(&packet.buf) else {
return Err(Error::new(io::ErrorKind::InvalidInput, "not tcp"));
continue;
};
if tcb::validated_tcp_header_len(&tcp_packet, &packet.buf).is_none() {
continue;
}
let acknowledgement = tcp_packet.get_acknowledgement();
let sequence = tcp_packet.get_sequence();
let local_addr = network_tuple.dst;
let peer_addr = network_tuple.src;
if tcp_packet.get_flags() & RST == RST {
if let Some((tcb, _)) = self.tcb_map.get(network_tuple) {
if tcb.rst_acceptable(&tcp_packet) {
if let Some(entry) = self.tcb_map.get(network_tuple) {
if entry.tcb.rst_acceptable(&tcp_packet) {
self.tcb_map.remove(network_tuple);
self.ip_stack.remove_tcp_half_open(network_tuple);
}
Expand All @@ -200,6 +221,14 @@ impl TcpListener {
log::debug!("drop tcp syn: half-open connection limit reached");
continue;
}
if let Some(entry) = self.tcb_map.get_mut(network_tuple) {
if let Some(relay_packet) = entry.tcb.try_syn_received(&tcp_packet) {
entry.retransmit_at = Instant::now() + entry.rto;
self.ip_stack.send_packet(relay_packet).await?;
continue;
}
}

// LISTEN -> SYN_RECEIVED
let mut tcp_config = self.ip_stack.config.tcp_config();
if tcp_config.mss.is_none() {
Expand All @@ -210,20 +239,30 @@ impl TcpListener {
}
let mut tcb = Tcb::new_listen(local_addr, peer_addr, tcp_config);
if let Some(relay_packet) = tcb.try_syn_received(&tcp_packet) {
let now = Instant::now();
let rto = tcb.rto();
self.ip_stack.add_tcp_half_open(*network_tuple);
self.tcb_map.insert(*network_tuple, (tcb, Instant::now()));
self.tcb_map.insert(
*network_tuple,
HalfOpenTcb {
tcb,
created_at: now,
retransmit_at: now + rto,
rto,
},
);
self.ip_stack.send_packet(relay_packet).await?;
continue;
}
} else if let Some((tcb, _)) = self.tcb_map.get_mut(network_tuple) {
} else if let Some(entry) = self.tcb_map.get_mut(network_tuple) {
// SYN_RECEIVED -> ESTABLISHED
if tcb.try_syn_received_to_established(packet.buf) {
let (tcb, _) = self.tcb_map.remove(network_tuple).unwrap();
let stream = TcpStream::new(self.ip_stack.clone(), tcb);
if entry.tcb.try_syn_received_to_established(packet.buf) {
let entry = self.tcb_map.remove(network_tuple).unwrap();
let stream = TcpStream::new(self.ip_stack.clone(), entry.tcb);
self.ip_stack.remove_tcp_half_open(network_tuple);
return Ok((stream?, peer_addr));
}
if tcb.is_close() {
if entry.tcb.is_close() {
self.tcb_map.remove(network_tuple).unwrap();
self.ip_stack.remove_tcp_half_open(network_tuple);
}
Expand All @@ -247,14 +286,34 @@ impl TcpListener {
let now = Instant::now();
let timeout = Duration::from_secs(10);
let ip_stack = self.ip_stack.clone();
self.tcb_map.retain(|network_tuple, (_, created_at)| {
let keep = *created_at + timeout > now;
self.tcb_map.retain(|network_tuple, entry| {
let keep = entry.created_at + timeout > now;
if !keep {
ip_stack.remove_tcp_half_open(network_tuple);
}
keep
});
}

async fn retransmit_half_open(&mut self) -> io::Result<()> {
self.expire_half_open();
let now = Instant::now();
let mut packets = Vec::new();
for entry in self.tcb_map.values_mut() {
if entry.retransmit_at <= now {
entry.tcb.note_syn_retransmission();
if let Some(packet) = entry.tcb.syn_ack_packet() {
packets.push(packet);
}
entry.rto = (entry.rto * 2).min(Duration::from_secs(3));
entry.retransmit_at = now + entry.rto;
}
}
for packet in packets {
self.ip_stack.send_packet(packet).await?;
}
Ok(())
}
}
#[cfg(feature = "global-ip-stack")]
impl TcpStream {
Expand Down
15 changes: 3 additions & 12 deletions src/tcp/sys.rs
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,6 @@ const MAX_INBOUND_BATCH: usize = 64;
#[derive(Debug)]
pub struct TcpStreamTask {
_bind_addr: Option<BindAddr>,
quick_end: bool,
tcb: Tcb,
ip_stack: IpStack,
application_layer_receiver: Receiver<BytesMut>,
Expand Down Expand Up @@ -80,7 +79,6 @@ impl TcpStreamTask {
) -> Self {
Self {
_bind_addr,
quick_end: ip_stack.config.tcp_config().quick_end,
tcb,
ip_stack,
application_layer_receiver,
Expand Down Expand Up @@ -109,16 +107,6 @@ impl TcpStreamTask {
if self.tcb.is_close() {
return Ok(());
}
if self.quick_end
&& self.read_half_closed()
&& self.write_half_closed
&& self.tcb.no_inflight_packet()
&& self.tcb.fin_acknowledged()
{
// Both halves are closed and the peer has acknowledged everything:
// it is safe to skip the remaining teardown states
return Ok(());
}
if !self.write_half_closed && !self.retransmission {
self.flush().await?;
}
Expand Down Expand Up @@ -429,6 +417,9 @@ impl TcpStreamTask {
let mut attempts = 0;
let mut time = self.tcb.rto();
while attempts < 50 {
if attempts > 0 {
self.tcb.note_syn_retransmission();
}
let Some(packet) = self.tcb.try_syn_sent() else {
return if self.tcb.is_close() {
Err(io::Error::from(io::ErrorKind::ConnectionRefused))
Expand Down
Loading
Loading