tcp: listener produces RST response when conns exhausted

This commit is contained in:
Patricio Whittingslow
2026-02-23 15:07:33 -03:00
parent 5221dc8248
commit 58a9bf57e5
2 changed files with 119 additions and 0 deletions
+81
View File
@@ -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()
+38
View File
@@ -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)