From d5ea2efdf483a81f0a33402682c513db9dae9779 Mon Sep 17 00:00:00 2001 From: TuteMthCD <43007973+TuteMthCD@users.noreply.github.com> Date: Tue, 23 Jun 2026 00:19:16 -0300 Subject: [PATCH] Retransmit active-open SYN after RTO (#130) * added: sync packet retransmission * test: sync packet retransmission * refactor: remove packet time from tcppool * fix: return comment * refactor: move retrying to stackRetrying * added: retries & timeout to stackGoConfig * test: test retries stack * refactor: move StackRetrying to stackBlocking * fix: line gofmt --------- Co-authored-by: Pat Whittingslow --- tcp/conn.go | 16 +++++++ tcp/handler.go | 39 ++++++++++++++-- tcp/handler_test.go | 97 ++++++++++++++++++++++++++++++++++++++++ x/xnet/stack-blocking.go | 10 ++++- x/xnet/stack-go.go | 29 +++++++++--- x/xnet/stack-retrying.go | 27 ++++++++--- x/xnet/xnet_test.go | 53 ++++++++++++++++++++++ 7 files changed, 256 insertions(+), 15 deletions(-) diff --git a/tcp/conn.go b/tcp/conn.go index cf3d5b8..5196982 100644 --- a/tcp/conn.go +++ b/tcp/conn.go @@ -108,6 +108,22 @@ func (conn *Conn) RemotePort() uint16 { return conn.h.RemotePort() } +// IsAwaitingControl reports whether the connection is waiting for a response to +// a control segment that can be retransmitted to advance connection state. +func (conn *Conn) IsAwaitingControl() bool { + conn.mu.Lock() + defer conn.mu.Unlock() + return conn.h.IsAwaitingControl() +} + +// RequeueControl asks the next packet emission to retransmit the outstanding +// control segment, if the connection is waiting for one. +func (conn *Conn) RequeueControl() { + conn.mu.Lock() + defer conn.mu.Unlock() + conn.h.RequeueControl() +} + // RemoteAddr returns the address of the peer Conn is exchanging data with. func (conn *Conn) RemoteAddr() []byte { conn.mu.Lock() diff --git a/tcp/handler.go b/tcp/handler.go index ae0510f..822a8db 100644 --- a/tcp/handler.go +++ b/tcp/handler.go @@ -32,7 +32,8 @@ type Handler struct { closing bool shutdownRx bool // nRetransmit stores the number of times the oldest packet was retransmit. - nRetransmit uint8 + nRetransmit uint8 + requeueControl bool } // SetLoggers sets the [slog.Logger] for the Handler and internal [ControlBlock]. @@ -271,6 +272,7 @@ func (h *Handler) Send(b []byte) (int, error) { return 0, net.ErrClosed } awaitingSyn := h.AwaitingSynSend() + requeueControl := h.requeueControl buffered := h.bufTx.BufferedUnsent() if h.scb.State() == StateCloseWait && !h.closing && buffered == 0 && !h.scb.HasPending() { // Remote closed with no application data left to send: initiate our own close. @@ -278,7 +280,7 @@ func (h *Handler) Send(b []byte) (int, error) { // before Send is called, implementing the half-close per RFC 9293 ยง3.5. h.closing = true } - if !awaitingSyn && buffered == 0 && !h.closing && !h.scb.HasPending() { + if !awaitingSyn && !requeueControl && buffered == 0 && !h.closing && !h.scb.HasPending() { // Early nop short circuit. return 0, nil } @@ -301,11 +303,27 @@ func (h *Handler) Send(b []byte) (int, error) { offset := uint8(5) mss := uint16(len(b) - sizeHeaderTCP) var segment Segment - if awaitingSyn { + if awaitingSyn || requeueControl && h.scb.State() == StateSynSent { // Handling init syn segment. segment = ClientSynSegment(h.bufTx.iss, Size(h.bufRx.Size())) h.optcodec.PutOption16(b[sizeHeaderTCP:], OptMaxSegmentSize, mss) offset++ + if requeueControl { + h.info("tcp.Handler:requeue-syn", slog.Uint64("port", uint64(h.localPort)), slog.Uint64("rport", uint64(h.remotePort))) + } + } else if requeueControl && h.scb.State() == StateSynRcvd { + segment = Segment{ + SEQ: h.scb.snd.UNA, + ACK: h.scb.rcv.NXT, + WND: Size(h.bufRx.Free()), + Flags: synack, + } + h.optcodec.PutOption16(b[sizeHeaderTCP:], OptMaxSegmentSize, mss) + offset++ + h.info("tcp.Handler:requeue-synack", slog.Uint64("port", uint64(h.localPort)), slog.Uint64("rport", uint64(h.remotePort))) + } else if requeueControl { + h.requeueControl = false + return 0, nil } else { var ok bool maxPayload := len(b) - sizeHeaderTCP @@ -335,6 +353,7 @@ func (h *Handler) Send(b []byte) (int, error) { } else if prevState != h.scb.State() && h.logenabled(slog.LevelInfo) { h.info("tcp.Handler:tx-statechange", slog.Uint64("port", uint64(h.localPort)), slog.String("oldState", prevState.String()), slog.String("newState", h.scb.State().String()), slog.String("txflags", segment.Flags.String())) } + h.requeueControl = false tfrm.SetSourcePort(h.localPort) tfrm.SetDestinationPort(h.remotePort) tfrm.SetSegment(segment, offset) @@ -443,6 +462,20 @@ func (h *Handler) AwaitingSynResponse() bool { return h.remotePort != 0 && h.scb.State() == StateSynSent } +// IsAwaitingControl reports whether the connection is waiting for a response to +// a control segment that can be retransmitted to advance connection state. +func (h *Handler) IsAwaitingControl() bool { + return h.AwaitingSynResponse() || h.scb.State() == StateSynRcvd +} + +// RequeueControl asks the next Send call to retransmit the outstanding control +// segment, if the connection is waiting for one. +func (h *Handler) RequeueControl() { + if h.IsAwaitingControl() { + h.requeueControl = true + } +} + // AwaitingSynAck returns true if the Handler is a passive server opened with [Handler.OpenListen] and not yet received a valid SYN remote packet. func (h *Handler) AwaitingSynAck() bool { return h.remotePort == 0 && h.scb.State() == StateListen diff --git a/tcp/handler_test.go b/tcp/handler_test.go index c411cbf..8722106 100644 --- a/tcp/handler_test.go +++ b/tcp/handler_test.go @@ -20,6 +20,103 @@ func TestHandler(t *testing.T) { sendDataFull(t, client, server, []byte("hello"), rawbuf[:]) } +func TestHandler_RequeueControlRetransmitsSYN(t *testing.T) { + const mtu = ethernet.MaxMTU + const maxpackets = 3 + rng := rand.New(rand.NewSource(1)) + client, server := newHandler(t, mtu, maxpackets), newHandler(t, mtu, maxpackets) + setupClientServer(t, rng, client, server) + + var rawbuf [mtu]byte + n, err := client.Send(rawbuf[:]) + if err != nil { + t.Fatal("client sending initial SYN:", err) + } else if n < sizeHeaderTCP { + t.Fatalf("initial SYN size=%d, want at least %d", n, sizeHeaderTCP) + } + initial := mustSegment(t, rawbuf[:n], 0) + if initial.Flags != FlagSYN { + t.Fatalf("initial flags=%s, want SYN", initial.Flags) + } + if client.State() != StateSynSent { + t.Fatalf("client state=%s, want SYN-SENT", client.State()) + } + + clear(rawbuf[:]) + n, err = client.Send(rawbuf[:]) + if err != nil { + t.Fatal("client sending before RequeueControl:", err) + } else if n != 0 { + t.Fatalf("Send before RequeueControl wrote %d bytes, want 0", n) + } + + client.RequeueControl() + clear(rawbuf[:]) + n, err = client.Send(rawbuf[:]) + if err != nil { + t.Fatal("client retransmitting SYN:", err) + } else if n < sizeHeaderTCP { + t.Fatalf("retransmitted SYN size=%d, want at least %d", n, sizeHeaderTCP) + } + retransmit := mustSegment(t, rawbuf[:n], 0) + if retransmit.Flags != FlagSYN { + t.Fatalf("retransmit flags=%s, want SYN", retransmit.Flags) + } + if retransmit.SEQ != initial.SEQ { + t.Fatalf("retransmit SEQ=%d, want initial SEQ=%d", retransmit.SEQ, initial.SEQ) + } + + if err := server.Recv(rawbuf[:n]); err != nil { + t.Fatal("server receiving retransmitted SYN:", err) + } + clear(rawbuf[:]) + n, err = server.Send(rawbuf[:]) + if err != nil { + t.Fatal("server sending SYN-ACK:", err) + } + segSynAck := mustSegment(t, rawbuf[:n], 0) + if segSynAck.Flags != synack { + t.Fatalf("server flags=%s, want SYN-ACK", segSynAck.Flags) + } + if err := client.Recv(rawbuf[:n]); err != nil { + t.Fatal("client receiving SYN-ACK:", err) + } + if client.State() != StateEstablished { + t.Fatalf("client state=%s, want ESTABLISHED", client.State()) + } + + clear(rawbuf[:]) + n, err = client.Send(rawbuf[:]) // final ACK. + if err != nil { + t.Fatal("client sending final ACK:", err) + } + if ack := mustSegment(t, rawbuf[:n], 0); ack.Flags != FlagACK { + t.Fatalf("client final flags=%s, want ACK", ack.Flags) + } + if err := server.Recv(rawbuf[:n]); err != nil { + t.Fatal("server receiving final ACK:", err) + } + + client.RequeueControl() + clear(rawbuf[:]) + n, err = client.Send(rawbuf[:]) + if err != nil { + t.Fatal("client sending after establishment:", err) + } else if n != 0 { + seg := mustSegment(t, rawbuf[:n], 0) + t.Fatalf("Send after establishment wrote %d bytes (%s), want 0", n, seg.Flags) + } +} + +func mustSegment(t *testing.T, b []byte, payloadLen int) Segment { + t.Helper() + frame, err := NewFrame(b) + if err != nil { + t.Fatal("parse TCP frame:", err) + } + return frame.Segment(payloadLen) +} + func sendDataFull(t *testing.T, client, server *Handler, data, packetBuf []byte) { n, err := client.Write(data) if err != nil { diff --git a/x/xnet/stack-blocking.go b/x/xnet/stack-blocking.go index 314b3be..575b8eb 100644 --- a/x/xnet/stack-blocking.go +++ b/x/xnet/stack-blocking.go @@ -181,6 +181,14 @@ func (s StackBlocking) DoDialTCP(conn *tcp.Conn, localPort uint16, addrp netip.A if err != nil { return err } + err = s.waitDialTCP(conn, timeout) + if err != nil { + conn.Abort() + } + return err +} + +func (s StackBlocking) waitDialTCP(conn *tcp.Conn, timeout time.Duration) (err error) { deadline := time.Now().Add(timeout) var backoffs uint for range maxIter { @@ -189,12 +197,10 @@ func (s StackBlocking) DoDialTCP(conn *tcp.Conn, localPort uint16, addrp netip.A return nil } else if state == tcp.StateSynSent || state == tcp.StateSynRcvd || conn.AwaitingSynSend() { if err = s.checkDeadline(deadline); err != nil { - conn.Abort() return err } } else { // Unexpected state, abort and terminate connection. - conn.Abort() return errTCPFailedToConnect } s.backoff(backoffs) diff --git a/x/xnet/stack-go.go b/x/xnet/stack-go.go index 34861ae..5d8fda6 100644 --- a/x/xnet/stack-go.go +++ b/x/xnet/stack-go.go @@ -18,10 +18,15 @@ import ( const ( sockSTREAM = 0x1 sockDGRAM = 0x2 + + defaultTCPDialTimeout = 2 * time.Second + defaultTCPDialRetries = 1 ) type StackGoConfig struct { ListenerPoolConfig TCPPoolConfig + TCPDialTimeout time.Duration + TCPDialRetries int } func (s *StackAsync) StackGo(stackProtoBackoff lneto.BackoffStrategy, cfg StackGoConfig) StackGo { @@ -32,16 +37,30 @@ func (s *StackAsync) StackGo(stackProtoBackoff lneto.BackoffStrategy, cfg StackG } func (s StackBlocking) StackGo(cfg StackGoConfig) StackGo { + tcpDialTimeout := cfg.TCPDialTimeout + tcpDialRetries := cfg.TCPDialRetries + // Defaults + if tcpDialTimeout <= 0 { + tcpDialTimeout = defaultTCPDialTimeout + } + if tcpDialRetries <= 0 { + tcpDialRetries = defaultTCPDialRetries + } + sg := StackGo{ - blk: s, - plcfg: cfg.ListenerPoolConfig, + blk: s, + plcfg: cfg.ListenerPoolConfig, + tcpDialTimeout: tcpDialTimeout, + tcpDialRetries: tcpDialRetries, } return sg } type StackGo struct { - blk StackBlocking - plcfg TCPPoolConfig + blk StackBlocking + plcfg TCPPoolConfig + tcpDialTimeout time.Duration + tcpDialRetries int } func (s StackGo) Socket(ctx context.Context, network string, family, sotype int, laddr, raddr net.Addr) (c any, err error) { @@ -171,7 +190,7 @@ func (s StackGo) SocketNetip(ctx context.Context, network string, family, sotype if err != nil { return nil, err } - err = s.blk.async.DialTCP(&conn, laddr.Port(), raddr) + err = s.blk.StackRetrying().DoDialTCP(&conn, laddr.Port(), raddr, s.tcpDialTimeout, s.tcpDialRetries) if err != nil { return nil, err } diff --git a/x/xnet/stack-retrying.go b/x/xnet/stack-retrying.go index dc1bdaa..dafd6f3 100644 --- a/x/xnet/stack-retrying.go +++ b/x/xnet/stack-retrying.go @@ -13,9 +13,11 @@ func (s *StackAsync) StackRetrying(stackProtoBackoff lneto.BackoffStrategy) Stac if stackProtoBackoff == nil { panic("nil backoff to StackRetrying") } - return StackRetrying{ - block: s.StackBlocking(stackProtoBackoff), - } + return s.StackBlocking(stackProtoBackoff).StackRetrying() +} + +func (s StackBlocking) StackRetrying() StackRetrying { + return StackRetrying{block: s} } var ( @@ -93,13 +95,28 @@ func (s StackRetrying) DoResolveHardwareAddress6(addr netip.Addr, timeout time.D func (s StackRetrying) DoDialTCP(conn *tcp.Conn, localPort uint16, addrp netip.AddrPort, timeout time.Duration, retries int) (err error) { expectEnd := time.Now().Add(timeout * time.Duration(retries)) var firstErr error - for range retries { - err = s.block.DoDialTCP(conn, localPort, addrp, timeout) + for i := range retries { + if i == 0 || conn.State().IsClosed() { + err = s.block.async.DialTCP(conn, localPort, addrp) + if err != nil { + if firstErr == nil { + firstErr = err + } + continue + } + } else if conn.IsAwaitingControl() { + conn.RequeueControl() + } + + err = s.block.waitDialTCP(conn, timeout) if err == nil { return nil } else if firstErr == nil { firstErr = err } + if !conn.IsAwaitingControl() { + conn.Abort() + } } if time.Now().Before(expectEnd) { if err != firstErr { diff --git a/x/xnet/xnet_test.go b/x/xnet/xnet_test.go index 8d38e2e..bbca978 100644 --- a/x/xnet/xnet_test.go +++ b/x/xnet/xnet_test.go @@ -2,10 +2,12 @@ package xnet import ( "bytes" + "context" "errors" "math/rand" "net/netip" "sync" + "syscall" "testing" "time" @@ -105,6 +107,57 @@ func TestTCPConn_ReadBlocksUntilDataAvailable(t *testing.T) { } } +func TestStackGoTCPDialRetriesPendingControl(t *testing.T) { + const seed = 5678 + const MTU = ethernet.MaxMTU + + client, sv, _, _ := newTCPStacks(t, seed, MTU) + sg := client.StackBlocking(backoffYield).StackGo(StackGoConfig{ + ListenerPoolConfig: TCPPoolConfig{ + QueueSize: 4, + TxBufSize: MTU, + RxBufSize: MTU, + NewBackoff: func() lneto.BackoffStrategy { + return backoffYield + }, + }, + TCPDialTimeout: 10 * time.Millisecond, + TCPDialRetries: 2, + }) + + laddr := netip.AddrPortFrom(netip.AddrFrom4(client.Addr4()), 1234) + raddr := netip.AddrPortFrom(netip.AddrFrom4(sv.Addr4()), 22) + done := make(chan error, 1) + go func() { + _, err := sg.SocketNetip(context.Background(), "tcp", syscall.AF_INET, sockSTREAM, laddr, raddr) + done <- err + }() + + var buf [ethernet.MaxMTU + ethernet.MaxOverheadSize]byte + waitForEgress := func() { + t.Helper() + deadline := time.Now().Add(100 * time.Millisecond) + for time.Now().Before(deadline) { + n, err := client.EgressEthernet(buf[:]) + if err != nil { + t.Fatal(err) + } + if n > 0 { + return + } + time.Sleep(time.Millisecond) + } + t.Fatal("timed out waiting for TCP dial egress packet") + } + waitForEgress() + waitForEgress() + + err := <-done + if err == nil { + t.Fatal("expected TCP dial to fail after retries without peer response") + } +} + func TestStackAsyncTCP_multipacket(t *testing.T) { const seed = 1234 const MTU = 512