diff --git a/internet/tcplistener_test.go b/internet/tcplistener_test.go index f1ff66f..f0a327c 100644 --- a/internet/tcplistener_test.go +++ b/internet/tcplistener_test.go @@ -309,6 +309,87 @@ func TestListener_MultiConn(t *testing.T) { } } +func TestListener_RSTOnPoolExhaustion(t *testing.T) { + rng := rand.New(rand.NewSource(1)) + var client1Stack, client2Stack, serverStack StackIP + var client1Conn, client2Conn, serverConn tcp.Conn + var listener tcp.Listener + + pool := newMockTCPPool(1, 3, 2048) // Pool size 1: will exhaust after first connection. + + setupClientServer(t, rng, &client1Stack, &serverStack, &client1Conn, &serverConn) + serverConn.Abort() + serverPort := uint16(80) + if err := listener.Reset(serverPort, pool); err != nil { + t.Fatal(err) + } + if err := serverStack.Register(&listener); err != nil { + t.Fatal(err) + } + + var buf [2048]byte + + // Complete full handshake for client1, exhausting the pool. + expectExchange(t, &client1Stack, &serverStack, buf[:]) // SYN + expectExchange(t, &serverStack, &client1Stack, buf[:]) // SYN-ACK + expectExchange(t, &client1Stack, &serverStack, buf[:]) // ACK + if pool.NumberOfAcquired() != 1 { + t.Fatalf("pool should have 1 acquired, got %d", pool.NumberOfAcquired()) + } + + // Setup client2 and send its SYN โ€” pool is full, server should queue RST. + const client2Port = uint16(1338) + setupClient(t, &client2Stack, &client2Conn, serverStack.Addr(), serverPort, client2Port) + + // Client2 sends SYN. + n, err := client2Stack.Encapsulate(buf[:], -1, 0) + if err != nil { + t.Fatal("client2 encapsulate:", err) + } else if n == 0 { + t.Fatal("client2 produced no SYN") + } + // Server receives SYN โ€” pool full, should return ErrPacketDrop but queue RST. + err = serverStack.Demux(buf[:n], 0) + if err == nil { + t.Fatal("expected error from server demux of rejected SYN") + } + + // Server encapsulates โ€” should produce RST (no connection data pending). + n, err = serverStack.Encapsulate(buf[:], -1, 0) + if err != nil { + t.Fatal("server encapsulate RST:", err) + } else if n == 0 { + t.Fatal("server produced no RST response") + } + + // Parse the IPv4+TCP frame to verify RST fields. + // IPv4 header is 20 bytes at offset 0, TCP starts at offset 20. + tfrm, err := tcp.NewFrame(buf[20:n]) + if err != nil { + t.Fatal("parse RST frame:", err) + } + _, flags := tfrm.OffsetAndFlags() + wantFlags := tcp.FlagRST | tcp.FlagACK + if flags != wantFlags { + t.Errorf("RST flags: got %s, want %s", flags, wantFlags) + } + if tfrm.SourcePort() != serverPort { + t.Errorf("RST source port: got %d, want %d", tfrm.SourcePort(), serverPort) + } + if tfrm.DestinationPort() != client2Port { + t.Errorf("RST dest port: got %d, want %d", tfrm.DestinationPort(), client2Port) + } + if tfrm.Seq() != 0 { + t.Errorf("RST SEQ: got %d, want 0", tfrm.Seq()) + } + // ACK should be client2's ISS + 1 (SYN occupies 1 sequence number). + // client2 was opened with ISS=100 (setupClient uses 100). + gotACK := tfrm.Ack() + if gotACK != 101 { + t.Errorf("RST ACK: got %d, want %d (client ISS+1)", gotACK, 101) + } +} + // tryExchange attempts an exchange but doesn't fail if no data to send. func tryExchange(t *testing.T, from, to *StackIP, buf []byte) { t.Helper() diff --git a/tcp/listener.go b/tcp/listener.go index d44feb5..5548613 100644 --- a/tcp/listener.go +++ b/tcp/listener.go @@ -28,6 +28,17 @@ type Listener struct { poolGet func() (*Conn, any, Value) poolReturn func(*Conn) logger + // rstQueue stores pending RST responses for SYNs rejected due to pool exhaustion. + // Per RFC 9293 ยง3.5.3: RST.SEQ=0, RST.ACK=SEG.SEQ+1, flags=RST|ACK. + rstQueue [4]rstEntry + rstQueueLen uint8 +} + +// rstEntry holds the minimum state needed to construct a stateless RST response. +type rstEntry struct { + remoteAddr [4]byte // IPv4 remote address. + remotePort uint16 + ackNum Value // SEG.SEQ + 1. } type handler struct { @@ -170,6 +181,26 @@ func (listener *Listener) Encapsulate(carrierData []byte, offsetToIP, offsetToFr err = nil } } + // Drain one RST entry if no connection data was sent. Lower priority than connection traffic. + if n == 0 && listener.rstQueueLen > 0 && offsetToIP >= 0 { + listener.rstQueueLen-- + entry := &listener.rstQueue[listener.rstQueueLen] + tfrm, err := NewFrame(carrierData[offsetToFrame:]) + if err == nil { + tfrm.SetSourcePort(listener.port) + tfrm.SetDestinationPort(entry.remotePort) + tfrm.SetSegment(Segment{ + SEQ: 0, + ACK: entry.ackNum, + Flags: FlagRST | FlagACK, + }, 5) + tfrm.SetUrgentPtr(0) + err = internal.SetIPAddrs(carrierData[offsetToIP:offsetToFrame], 0, nil, entry.remoteAddr[:]) + if err == nil { + return sizeHeaderTCP, nil + } + } + } if n == 0 { listener.maintainConns() } @@ -217,6 +248,13 @@ func (listener *Listener) Demux(carrierData []byte, tcpFrameOffset int) error { conn, userData, iss := listener.poolGet() if conn == nil { slog.Error("tcpListener:no-free-conn") + if len(srcaddr) == 4 && listener.rstQueueLen < uint8(len(listener.rstQueue)) { + entry := &listener.rstQueue[listener.rstQueueLen] + entry.remotePort = src + entry.ackNum = tfrm.Seq() + 1 + copy(entry.remoteAddr[:], srcaddr) + listener.rstQueueLen++ + } return lneto.ErrPacketDrop } err = conn.OpenListen(dst, iss)