From 07afcfd9247b185e83cfdd3190c65a9d21bdc8d8 Mon Sep 17 00:00:00 2001 From: Pat Whittingslow Date: Mon, 7 Sep 2026 10:46:24 -0300 Subject: [PATCH] tcp: simplified Policy implementation based on @MDr164 (#190) * begin prepping policy refactor manually * implement tcp.Policy and refactor rto to use it * fix CI * chatting with claude gave me idea to reformulate Policy * tcp.Policy: add newTransmitLimit output * merge with main and fix failing tests * remove fix.patch * answer my own comments --- tcp/conn.go | 20 +- tcp/control.go | 32 ++- tcp/handler.go | 143 ++++++----- tcp/loss.go | 83 ------ tcp/loss_test.go | 262 ------------------- tcp/policy.go | 30 +++ tcp/policy_test.go | 449 +++++++++++++++++++++++++++++++++ tcp/rto/integration_test.go | 222 ++++++++++++++++ tcp/{rto.go => rto/timer.go} | 162 ++++++++---- tcp/rto/timer_test.go | 311 +++++++++++++++++++++++ tcp/rto_test.go | 206 --------------- tcp/rtointegration_test.go | 120 --------- tcp/txqueue.go | 69 ++++- tcp/txqueue_retransmit_test.go | 215 ++++++++++++++++ 14 files changed, 1518 insertions(+), 806 deletions(-) delete mode 100644 tcp/loss.go delete mode 100644 tcp/loss_test.go create mode 100644 tcp/policy.go create mode 100644 tcp/policy_test.go create mode 100644 tcp/rto/integration_test.go rename tcp/{rto.go => rto/timer.go} (50%) create mode 100644 tcp/rto/timer_test.go delete mode 100644 tcp/rto_test.go delete mode 100644 tcp/rtointegration_test.go create mode 100644 tcp/txqueue_retransmit_test.go diff --git a/tcp/conn.go b/tcp/conn.go index 547563d..fdb4c7c 100644 --- a/tcp/conn.go +++ b/tcp/conn.go @@ -76,16 +76,10 @@ type ConnConfig struct { // Logger sets the [Conn] logger. // Lower level logging available at [Handler.SetLoggers] via [Conn.InternalHandler]. Logger *slog.Logger - // LossRecovery is the optional packet-loss recovery algorithm (RTO, - // congestion control, ...) for the connection. If set, Nanotime must also be - // set (else Configure returns an error). Leaving it nil disables loss - // recovery. See [LossRecovery]. - LossRecovery LossRecovery - // Nanotime is the monotonic time source in nanoseconds (the func() int64 - // convention used across lneto) that drives LossRecovery. It is required when - // LossRecovery is set and unused otherwise. The tcp package reads it only to - // stamp the loss-recovery hooks; it holds no clock itself. - Nanotime func() int64 + // Policy is the optional transmit-steering algorithm (RTO, congestion + // control, ...) for the connection. nil disables it. A Policy needing time + // carries its own clock. See [Policy]. + Policy Policy } // Configure should be called on any newly created connection before usage. See [ConnConfig]. @@ -93,10 +87,6 @@ func (conn *Conn) Configure(config ConnConfig) (err error) { if config.RWBackoff == nil { return lneto.ErrMissingHALConfig } - if config.LossRecovery != nil && config.Nanotime == nil { - // The tcp package holds no clock: a loss-recovery algorithm cannot run without it. - return lneto.ErrInvalidConfig - } conn.mu.Lock() defer conn.mu.Unlock() err = conn.h.SetBuffers(config.TxBuf, config.RxBuf, config.TxPacketQueueSize) @@ -105,7 +95,7 @@ func (conn *Conn) Configure(config ConnConfig) (err error) { } conn._backoff = config.RWBackoff conn.logger.log = config.Logger - conn.h.SetLossRecovery(config.LossRecovery, config.Nanotime) + conn.h.SetPolicy(config.Policy) return nil } diff --git a/tcp/control.go b/tcp/control.go index 31b91da..ba258ab 100644 --- a/tcp/control.go +++ b/tcp/control.go @@ -95,6 +95,12 @@ func (tcb *ControlBlock) RecvWindow() Size { return tcb.rcv.WND } // ISS returns the initial sequence number of the connection that was defined on a call to Open by user. func (tcb *ControlBlock) ISS() Value { return tcb.snd.ISS } +// SendUNA returns snd.UNA, the oldest sequence number not yet acked by the remote. +func (tcb *ControlBlock) SendUNA() Value { return tcb.snd.UNA } + +// SendNext returns snd.NXT, one past the highest sequence number sent. +func (tcb *ControlBlock) SendNext() Value { return tcb.snd.NXT } + // MaxInFlightData returns the maximum size of a segment that can be sent by taking into account // the send window size and the unacked data. Returns 0 before StateSynRcvd. func (tcb *ControlBlock) MaxInFlightData() Size { @@ -257,14 +263,34 @@ func (tcb *ControlBlock) HasPendingRetransmit() bool { return tcb._state.TxDataOpen() && tcb.dupack >= retransmitAfterDupacks && tcb.nRetransmit <= tcb.dupack-retransmitAfterDupacks } +// RetransmitFrom rewinds snd.NXT back to newNxt so the next PendingSegment and +// Send calls retransmit unacknowledged data from that sequence number onwards. +// It must be paired with ringTx.RetransmitFrom to rewind the transmit buffer to +// the same point. Implements RFC 9293 §3.10.8 (RETRANSMISSION TIMEOUT). +// +// It reports false and changes nothing when newNxt falls outside the +// unacknowledged range [snd.UNA, snd.NXT] or the connection cannot send data, so +// a misbehaving [Policy] cannot corrupt the send sequence space. +func (tcb *ControlBlock) RetransmitFrom(newNxt Value) bool { + if !tcb._state.txQueuedDataOpen() { + // Matches [State.TxDataOpen] and other states that may have data queued to make progress. + // Matches [ControlBlock.PendingSegment] gate (RFC 9293 §3.10.8). + return false + } else if newNxt.LessThan(tcb.snd.UNA) || tcb.snd.NXT.LessThan(newNxt) { + return false + } + tcb.snd.NXT = newNxt + tcb.dupack = 0 + tcb.nRetransmit = 0 + return true +} + // RetransmitAll rewinds snd.NXT back to snd.UNA so the next PendingSegment and // Send calls retransmit all unacknowledged data from the oldest sequence number // (go-back-N). It must be paired with ringTx.RetransmitFromUNA to rewind the // transmit buffer. Implements RFC 9293 §3.10.8 (RETRANSMISSION TIMEOUT). func (tcb *ControlBlock) RetransmitAll() { - tcb.snd.NXT = tcb.snd.UNA - tcb.dupack = 0 - tcb.nRetransmit = 0 + tcb.RetransmitFrom(tcb.snd.UNA) } // PendingSegment calculates a suitable next segment to send from a payload length. diff --git a/tcp/handler.go b/tcp/handler.go index 9ff7a29..09eca89 100644 --- a/tcp/handler.go +++ b/tcp/handler.go @@ -31,14 +31,8 @@ type Handler struct { optcodec OptionCodec // reasm tracks out-of-order segments staged in bufRx's free region. Always // enabled once buffers are set (see [Handler.SetBuffers]). - reasm reassembly - // loss is the optional packet-loss recovery algorithm (RTO, congestion - // control, ...) driven from the rx/tx hooks. nil disables loss recovery, in - // which case the connection behaves as if no timing existed. nanotime is the - // monotonic time source (nanoseconds) passed to those hooks; it is non-nil - // whenever loss is non-nil (enforced by [Conn.Configure]). See [LossRecovery]. - loss LossRecovery - nanotime func() int64 + reasm reassembly + policy Policy closing bool shutdownRx bool @@ -79,27 +73,16 @@ func (h *Handler) SetBuffers(txbuf, rxbuf []byte, packets int) error { return h.bufTx.ResetOrReuse(txbuf, packets, 0) } -// SetLossRecovery installs the packet-loss recovery algorithm and the monotonic -// time source (nanoseconds, the func() int64 convention used across lneto) that -// drives it. The tcp package keeps no clock of its own; nanotime is read only to -// stamp the rx/tx hooks (see [LossRecovery]). Passing loss == nil disables loss -// recovery. It should be set before the connection is opened. -func (h *Handler) SetLossRecovery(loss LossRecovery, nanotime func() int64) { - h.loss = loss - h.nanotime = nanotime +// SetPolicy installs the transmit-steering algorithm. nil disables it. +// It should be set before the connection is opened. See [Policy]. +func (h *Handler) SetPolicy(policy Policy) { + h.policy = policy } +func (h *Handler) policyEnabled() bool { return h.policy != nil } -func (h *Handler) lossEnabled() bool { return h.loss != nil } - -// NextDeadline returns the monotonic-nanosecond instant at which the connection -// must next be serviced by a transmit attempt (e.g. an RTO expiry), or 0 when -// there is no deadline or no loss recovery is configured. See [LossRecovery]. -func (h *Handler) NextDeadline() int64 { - if h.loss == nil { - return 0 - } - return h.loss.NextDeadline() -} +// ControlBlock returns the state machine underlying the Handler, mainly so a +// [Policy] can read the sequence spaces. Not for modification. +func (h *Handler) ControlBlock() *ControlBlock { return &h.scb } // LocalPort returns the local port of the connection. Returns 0 if the connection is closed and uninitialized. func (h *Handler) LocalPort() uint16 { @@ -165,16 +148,15 @@ func (h *Handler) reset(localPort, remotePort uint16, iss Value) { shutdownRx: false, // Persist configuration across reopen: validator: h.validator, - loss: h.loss, - nanotime: h.nanotime, + policy: h.policy, logger: h.logger, // persist memory across repoen: bufTx: h.bufTx, bufRx: h.bufRx, reasm: h.reasm, } - if h.lossEnabled() { - h.loss.Reset() + if h.policyEnabled() { + h.policy.Reset() } h.reasm.clear() // preserve metadata capacity across reopen, drop held segments. h.bufTx.ResetOrReuse(nil, 0, iss) @@ -212,9 +194,7 @@ func (h *Handler) Recv(incomingPacket []byte) error { return nil } - // Notify loss recovery of the received segment (RTT sampling, timer - // management) and let it drop the segment before processing if it asks to. - if h.lossEnabled() && !h.loss.PreRx(segIncoming, h.nanotime()).Keep { + if h.policyEnabled() && !h.policy.PreRx(h, tfrm) { return nil } @@ -245,6 +225,9 @@ func (h *Handler) Recv(incomingPacket []byte) error { if prevState != h.scb.State() { h.info("tcp.Handler:rx-statechange", slog.Uint64("port", uint64(h.localPort)), slog.String("old", prevState.String()), slog.String("new", h.scb.State().String()), slog.String("rxflags", segIncoming.Flags.String())) } + if h.policyEnabled() { + h.policy.PostRx(h, prevState, tfrm) + } if segIncoming.DATALEN != 0 && h.shutdownRx && (h.scb.State() == StateFinWait1 || h.scb.State() == StateFinWait2) { // soypat/lneto#50: the application is done in both directions — read side // shut down (CloseRead) and our FIN sent (Close) — so inbound data has no @@ -380,16 +363,31 @@ func (h *Handler) Send(b []byte) (int, error) { if h.IsTxOver() { return 0, net.ErrClosed } - var now int64 - if h.lossEnabled() { - now = h.nanotime() - if h.loss.PreTx(now).RetransmitAll { - // Go-back-N retransmission directed by loss recovery: rewind the - // send sequence and transmit buffer so unacknowledged data is resent - // from snd.UNA. Done before the early short-circuit below so an - // expired RTO retransmits even with no new data queued. - h.scb.RetransmitAll() - h.bufTx.RetransmitFromUNA() + tfrm, err := NewFrame(b) + if err != nil { + return 0, err + } + offset := uint8(5) + txLimit := TransmitUnlimited + if h.policyEnabled() { + // Hand the Policy a defined frame: zeroed header at the minimum offset. + // It may append options and raise the offset, which is read back below. + tfrm.ClearHeader() + tfrm.SetOffsetAndFlags(offset, 0) + limit, rtxFrom, doRtx := h.policy.PreTx(h, tfrm) + txLimit = limit + if limit == 0 { + h.info("tcp.Policy:newTxLimit=0") // Can cause headaches for users. + } + if doRtx && h.scb.RetransmitFrom(rtxFrom) { + // Retransmission directed by the Policy: rewind the transmit buffer + // to match the send sequence so unacknowledged data is resent. Done + // before the early short-circuit below so an expired RTO + // retransmits even with no new data queued. + h.bufTx.RetransmitFrom(rtxFrom) + } + if o, _ := tfrm.OffsetAndFlags(); o > offset && int(o)*4 < len(b) { + offset = o } } awaitingSyn := h.AwaitingSynSend() @@ -405,29 +403,27 @@ func (h *Handler) Send(b []byte) (int, error) { // Early nop short circuit. return 0, nil } - tfrm, err := NewFrame(b) - if err != nil { - return 0, err - } if buffered == 0 && h.closing && (h.scb.State() != StateCloseWait || !h.scb.HasPending()) { // If Close called and no more data to be sent, terminate connection. // In CLOSE-WAIT: wait until the pending ACK is sent first, since scb.Close() // overwrites pending with [FIN|ACK] (unlike ESTABLISHED which merges via bitmask). h.closing = false - err = h.scb.Close() + err := h.scb.Close() if err != nil { h.logerr("tcp.Handler.Close", slog.String("err", errstr(err)), slog.String("state", h.State().String())) h.Abort() return 0, io.EOF } } - offset := uint8(5) - mss := uint16(len(b) - sizeHeaderTCP) + // optHead is where the Handler's own options begin: after the fixed header + // and after any options the Policy already wrote, so neither clobbers the other. + optHead := int(offset) * 4 + mss := uint16(len(b) - optHead) var segment Segment 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) + h.optcodec.PutOption16(b[optHead:], OptMaxSegmentSize, mss) offset++ if requeueControl { h.info("tcp.Handler:requeue-syn", slog.Uint64("port", uint64(h.localPort)), slog.Uint64("rport", uint64(h.remotePort))) @@ -439,7 +435,7 @@ func (h *Handler) Send(b []byte) (int, error) { WND: Size(h.bufRx.Free()), Flags: synack, } - h.optcodec.PutOption16(b[sizeHeaderTCP:], OptMaxSegmentSize, mss) + h.optcodec.PutOption16(b[optHead:], OptMaxSegmentSize, mss) offset++ h.info("tcp.Handler:requeue-synack", slog.Uint64("port", uint64(h.localPort)), slog.Uint64("rport", uint64(h.remotePort))) } else if requeueControl { @@ -447,17 +443,21 @@ func (h *Handler) Send(b []byte) (int, error) { return 0, nil } else { var ok bool - maxPayload := len(b) - sizeHeaderTCP + maxPayload := len(b) - optHead + if txLimit < Size(maxPayload) && !h.nextSegmentIsRetransmit() { + // Policy clamped new data. + maxPayload = int(txLimit) + } segment, ok = h.scb.PendingSegment(maxPayload) segment.WND = h.recvWindow() if !ok { // No pending control segment or data to send. Yield. return 0, nil } else if segment.Flags == synack { - h.optcodec.PutOption16(b[sizeHeaderTCP:], OptMaxSegmentSize, mss) + h.optcodec.PutOption16(b[optHead:], OptMaxSegmentSize, mss) offset++ } else if segment.DATALEN > 0 { - n, err := h.bufTx.MakePacket(b[sizeHeaderTCP:sizeHeaderTCP+segment.DATALEN], segment.SEQ) + n, err := h.bufTx.MakePacket(b[optHead:optHead+int(segment.DATALEN)], segment.SEQ) if err != nil { return 0, err } @@ -474,15 +474,19 @@ 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())) } - if h.lossEnabled() { - h.loss.PostTx(segment, now) - } h.requeueControl = false tfrm.SetSourcePort(h.localPort) tfrm.SetDestinationPort(h.remotePort) tfrm.SetSegment(segment, offset) tfrm.SetUrgentPtr(0) datalen := int(offset)*4 + int(segment.DATALEN) + if h.policyEnabled() { + // Frame trimmed to what is actually emitted so the Policy's Payload() + // is the segment data and nothing more. + if sent, err := NewFrame(b[:datalen]); err == nil { + h.policy.PostTx(h, sent) + } + } closedSuccess := prevState == StateTimeWait && segment.Flags.HasAny(FlagACK) if closedSuccess { h.reset(0, 0, 0) @@ -494,6 +498,27 @@ func (h *Handler) Send(b []byte) (int, error) { return datalen, nil } +// nextSegmentIsRetransmit reports whether the next data segment would resend +// already-transmitted bytes rather than open new sequence space. Used to let a +// retransmission through while a [Policy] holds new data back. +func (h *Handler) nextSegmentIsRetransmit() bool { + endSeq, hasSent := h.bufTx.sentEndSeq() + return hasSent && h.scb.snd.NXT.LessThan(endSeq) +} + +// NextSegmentSYN returns syn=true if next outgoing segment is a handshake SYN. +// This method is exported for use by [Policy] implementations to decide handshake-only options (window scale, SACK-permitted, timestamps). +func (h *Handler) NextSegmentSYN() (syn, ack bool) { + state := h.scb.State() + if h.AwaitingSynSend() || h.requeueControl && state == StateSynSent { + return true, false // SYN initial/requeue. + } else if h.requeueControl && state == StateSynRcvd { + return true, true // SYNACK requeue. + } + pending := h.scb.pending[0] + return pending.HasAny(FlagSYN), pending.HasAny(FlagACK) +} + // Write implements [io.Writer] by copying b to a internal buffer to be sent over the network on the next // [Handler.Send] call that can send data to remote peer. Use [Handler.Free] to know the maximum length the argument slice can be before erroring. func (h *Handler) Write(b []byte) (int, error) { diff --git a/tcp/loss.go b/tcp/loss.go deleted file mode 100644 index 98cfe5c..0000000 --- a/tcp/loss.go +++ /dev/null @@ -1,83 +0,0 @@ -package tcp - -// LossRecovery abstracts TCP packet-loss recovery: RTO, congestion control and -// any similar algorithm that observes segment traffic and steers the -// connection's transmit behaviour. As far as the tcp package is concerned these -// are all the same thing — packet-loss recovery algorithms — so they share one -// interface (see discussion #157). -// -// The tcp package stays free of any time source: the current monotonic time in -// nanoseconds (the func() int64 convention used across lneto) is passed in at -// each hook boundary. It originates from [ConnConfig.Nanotime] and satisfies the -// "WHEN was this segment rx/tx'd" requirement without a clock living inside the -// state machine, which also keeps implementations deterministic for testing -// (see issue #140). -// -// The interface is intentionally free of errors: an implementation handles or -// reports its own errors rather than propagating them into lneto internals. -// -// Introspection (smoothed RTT, current window, ...) is deliberately left off the -// interface; expose it on the concrete implementation the caller constructs and -// hands to [ConnConfig]. -type LossRecovery interface { - // Reset returns the implementation to its initial, pre-connection state. It - // is invoked whenever the connection is (re)opened or aborted so a single - // LossRecovery value can be reused across the lifetime of connection reuse - // (see discussion #115). - Reset() - - // NextDeadline returns the monotonic-nanosecond instant at which the - // connection must next be serviced by a transmit attempt — typically the RTO - // expiry. A return of 0 means there is no pending deadline. It replaces a - // poll/atomic-flag scheme with a deadline the caller's event loop can - // schedule against. - NextDeadline() int64 - - // PreRx is called for every segment received on the TCP port before the - // state machine processes it, with the monotonic time the segment arrived. It - // returns whether the segment should be kept (processed) or dropped. - PreRx(incoming Segment, now int64) RxDirective - - // PreTx is called on entering the transmit path (Encapsulate), before a - // segment is built, with the current monotonic time. Its directive tells the - // connection whether to retransmit unacknowledged data, rewind the send - // pointer, or hold back new data. - PreTx(now int64) TxDirective - - // PostTx is called on leaving the transmit path with the segment that was - // actually emitted and the monotonic time it was sent. This is where segment - // timing (for RTT sampling and the retransmission timer) is recorded. - PostTx(outgoing Segment, now int64) -} - -// TxDirective is returned by [LossRecovery.PreTx] to steer the transmit path. -// The zero value directs the connection to proceed normally (send new data if -// available, no retransmission). -type TxDirective struct { - // RewindNXT is the number of sequence-space octets to rewind snd.NXT by - // before transmitting, for partial (e.g. selective) retransmission. Zero - // means no rewind. It is independent of Retransmit, which rewinds fully to - // snd.UNA. - // RewindNXT uint32 - - // RetransmitAll requests go-back-N retransmission: the connection rewinds - // snd.NXT to snd.UNA and resends unacknowledged data from the oldest - // sequence number. - RetransmitAll bool - // HoldNew pauses transmission of new data (for example when the congestion - // window is exhausted). Retransmissions already directed by this same - // directive still proceed. - // HoldNew bool -} - -// RxDirective is returned by [LossRecovery.PreRx]. -// -// NOTE: its shape is the minimum viable contract — it mirrors the original -// PreRx "keep" boolean from discussion #157 — and is the one element of the -// interface not yet fully settled there. It is a struct (rather than a bare -// bool) so fields can be added without breaking implementations. -type RxDirective struct { - // Keep reports whether the received segment should be handed to the state - // machine. A false value drops the segment before it is processed. - Keep bool -} diff --git a/tcp/loss_test.go b/tcp/loss_test.go deleted file mode 100644 index 18dd13a..0000000 --- a/tcp/loss_test.go +++ /dev/null @@ -1,262 +0,0 @@ -package tcp - -import ( - "math/rand" - "testing" - - "github.com/soypat/lneto/ethernet" -) - -// recordingLoss is a test LossRecovery that records every hook invocation and -// lets the test steer the directives returned to the Handler. It is the -// interface counterpart driven by the Handler under test. -type recordingLoss struct { - resets int - preRx []hookCall - preTx []int64 - postTx []hookCall - deadline int64 // value NextDeadline reports back. - - // Directives handed back to the Handler. - keep bool // PreRx result. Default true (see newRecordingLoss). - tx TxDirective // PreTx result. -} - -type hookCall struct { - seg Segment - now int64 -} - -func newRecordingLoss() *recordingLoss { return &recordingLoss{keep: true} } - -var _ LossRecovery = (*recordingLoss)(nil) - -func (l *recordingLoss) Reset() { l.resets++ } -func (l *recordingLoss) NextDeadline() int64 { return l.deadline } - -func (l *recordingLoss) PreRx(incoming Segment, now int64) RxDirective { - l.preRx = append(l.preRx, hookCall{seg: incoming, now: now}) - return RxDirective{Keep: l.keep} -} - -func (l *recordingLoss) PreTx(now int64) TxDirective { - l.preTx = append(l.preTx, now) - return l.tx -} - -func (l *recordingLoss) PostTx(outgoing Segment, now int64) { - l.postTx = append(l.postTx, hookCall{seg: outgoing, now: now}) -} - -// TestLossRecovery_DisabledByDefault verifies the Handler runs normally with no -// loss recovery installed: NextDeadline reports no deadline and the transmit/ -// receive paths never touch a nil LossRecovery. -func TestLossRecovery_DisabledByDefault(t *testing.T) { - const mtu = ethernet.MaxMTU - rng := rand.New(rand.NewSource(1)) - client, server := newHandler(t, mtu, 3), newHandler(t, mtu, 3) - setupClientServer(t, rng, client, server) - - if d := client.NextDeadline(); d != 0 { - t.Fatalf("NextDeadline with no loss recovery = %d, want 0", d) - } - var buf [mtu]byte - establish(t, client, server, buf[:]) // must not panic on nil loss recovery. -} - -// TestLossRecovery_HooksInvoked verifies the Handler drives the full hook -// contract across a handshake: Reset on open, PreTx+PostTx on every transmit, -// PreRx on every receive, each stamped with the configured monotonic clock. -func TestLossRecovery_HooksInvoked(t *testing.T) { - const mtu = ethernet.MaxMTU - rng := rand.New(rand.NewSource(2)) - client, server := newHandler(t, mtu, 3), newHandler(t, mtu, 3) - - loss := newRecordingLoss() - const clockNow = 1_000_000 - client.SetLossRecovery(loss, func() int64 { return clockNow }) - - setupClientServer(t, rng, client, server) // OpenActive → reset → Reset(). - if loss.resets == 0 { - t.Fatal("Reset not called on open") - } - - var buf [mtu]byte - establish(t, client, server, buf[:]) - - // Client emitted SYN and the final ACK: both paths must have hit PreTx/PostTx. - if len(loss.preTx) == 0 { - t.Fatal("PreTx never called on transmit") - } - if len(loss.postTx) == 0 { - t.Fatal("PostTx never called on transmit") - } - if len(loss.preTx) != len(loss.postTx) { - t.Fatalf("PreTx calls=%d, PostTx calls=%d, want equal", len(loss.preTx), len(loss.postTx)) - } - // Client received the SYN-ACK: PreRx must have seen it. - if len(loss.preRx) == 0 { - t.Fatal("PreRx never called on receive") - } - - // The Handler holds no clock: every hook must be stamped from the supplied - // nanotime source. - for i, c := range loss.postTx { - if c.now != clockNow { - t.Fatalf("PostTx[%d].now = %d, want clock %d", i, c.now, clockNow) - } - } - for i, now := range loss.preTx { - if now != clockNow { - t.Fatalf("PreTx[%d].now = %d, want clock %d", i, now, clockNow) - } - } - for i, c := range loss.preRx { - if c.now != clockNow { - t.Fatalf("PreRx[%d].now = %d, want clock %d", i, c.now, clockNow) - } - } - - // PostTx receives the segment actually emitted: the first is the SYN. - if !loss.postTx[0].seg.Flags.HasAny(FlagSYN) { - t.Fatalf("first PostTx segment flags=%s, want SYN", loss.postTx[0].seg.Flags) - } -} - -// TestLossRecovery_NextDeadlineDelegates verifies NextDeadline is forwarded to -// the installed LossRecovery unchanged. -func TestLossRecovery_NextDeadlineDelegates(t *testing.T) { - const mtu = ethernet.MaxMTU - rng := rand.New(rand.NewSource(3)) - client, server := newHandler(t, mtu, 3), newHandler(t, mtu, 3) - - loss := newRecordingLoss() - loss.deadline = 4242 - client.SetLossRecovery(loss, func() int64 { return 1 }) - setupClientServer(t, rng, client, server) - - if d := client.NextDeadline(); d != 4242 { - t.Fatalf("NextDeadline = %d, want delegated 4242", d) - } -} - -// TestLossRecovery_PreRxDropsSegment verifies a PreRx directive of Keep=false -// drops the segment before the state machine sees it: the payload is not -// buffered and connection state is untouched. -func TestLossRecovery_PreRxDropsSegment(t *testing.T) { - const mtu = ethernet.MaxMTU - rng := rand.New(rand.NewSource(4)) - client, server := newHandler(t, mtu, 3), newHandler(t, mtu, 3) - - loss := newRecordingLoss() - server.SetLossRecovery(loss, func() int64 { return 1 }) - setupClientServer(t, rng, client, server) - var buf [mtu]byte - establish(t, client, server, buf[:]) // keep=true so handshake completes. - - // Now start dropping everything the server receives. - loss.keep = false - preRxBefore := len(loss.preRx) - - data := []byte("dropme") - if _, err := client.Write(data); err != nil { - t.Fatal("client write:", err) - } - clear(buf[:]) - n, err := client.Send(buf[:]) - if err != nil { - t.Fatal("client send:", err) - } - - if err := server.Recv(buf[:n]); err != nil { - t.Fatalf("dropped segment must return nil, got %v", err) - } - if len(loss.preRx) != preRxBefore+1 { - t.Fatalf("PreRx calls=%d, want %d (segment must reach PreRx)", len(loss.preRx), preRxBefore+1) - } - if server.BufferedInput() != 0 { - t.Fatalf("dropped segment must not be buffered, got %d bytes", server.BufferedInput()) - } - if server.State() != StateEstablished { - t.Fatalf("dropped segment must not change state, got %s", server.State()) - } -} - -// TestLossRecovery_PreTxRetransmitAll verifies a PreTx directive of -// RetransmitAll drives go-back-N: the Handler rewinds and re-emits already-sent, -// unacknowledged data from snd.UNA on the next transmit. -func TestLossRecovery_PreTxRetransmitAll(t *testing.T) { - const mtu = ethernet.MaxMTU - rng := rand.New(rand.NewSource(5)) - client, server := newHandler(t, mtu, 3), newHandler(t, mtu, 3) - - loss := newRecordingLoss() - client.SetLossRecovery(loss, func() int64 { return 1 }) - setupClientServer(t, rng, client, server) - var buf [mtu]byte - establish(t, client, server, buf[:]) - - // Emit one data segment; server never ACKs, so it stays unacknowledged. - data := []byte("payload") - if _, err := client.Write(data); err != nil { - t.Fatal("client write:", err) - } - clear(buf[:]) - n, err := client.Send(buf[:]) - if err != nil { - t.Fatal("client send data:", err) - } - if n <= sizeHeaderTCP { - t.Fatal("expected data segment") - } - firstSeg := mustSegment(t, buf[:n], n-sizeHeaderTCP) - - // Direct go-back-N on the next transmit. - loss.tx = TxDirective{RetransmitAll: true} - clear(buf[:]) - n, err = client.Send(buf[:]) - if err != nil { - t.Fatal("client send retransmit:", err) - } - if n <= sizeHeaderTCP { - t.Fatal("expected retransmitted data segment") - } - rtSeg := mustSegment(t, buf[:n], n-sizeHeaderTCP) - - if rtSeg.SEQ != firstSeg.SEQ { - t.Fatalf("retransmit SEQ=%d, want original UNA SEQ=%d (go-back-N)", rtSeg.SEQ, firstSeg.SEQ) - } - if rtSeg.DATALEN != firstSeg.DATALEN { - t.Fatalf("retransmit DATALEN=%d, want %d", rtSeg.DATALEN, firstSeg.DATALEN) - } -} - -// TestLossRecovery_ResetOnReopen verifies Reset fires on every (re)open and on -// Abort, so a single LossRecovery value can be reused across connection reuse. -func TestLossRecovery_ResetOnReopen(t *testing.T) { - const mtu = ethernet.MaxMTU - client := newHandler(t, mtu, 3) - loss := newRecordingLoss() - client.SetLossRecovery(loss, func() int64 { return 1 }) - - if err := client.OpenActive(1234, 5678, 0); err != nil { - t.Fatal("open 1:", err) - } - afterOpen := loss.resets - if afterOpen == 0 { - t.Fatal("Reset not called on first open") - } - - client.Abort() - if loss.resets <= afterOpen { - t.Fatalf("Reset not called on Abort: resets=%d, want >%d", loss.resets, afterOpen) - } - afterAbort := loss.resets - - if err := client.OpenActive(1234, 5678, 0); err != nil { - t.Fatal("open 2:", err) - } - if loss.resets <= afterAbort { - t.Fatalf("Reset not called on reopen: resets=%d, want >%d", loss.resets, afterAbort) - } -} diff --git a/tcp/policy.go b/tcp/policy.go new file mode 100644 index 0000000..43742d3 --- /dev/null +++ b/tcp/policy.go @@ -0,0 +1,30 @@ +package tcp + +// TransmitUnlimited size returned by [Policy.PreTx] to signal no new data transmit limit (no congestion control). +const TransmitUnlimited = ^Size(0) + +// Policy observes segment traffic and steers transmit behaviour: RTO, +// congestion control and the like (discussion #157). The tcp package holds no +// clock, so a Policy needing time carries its own (issue #140). +type Policy interface { + // Reset returns the Policy to its pre-connection state. Should be called on every + // Open/Listen on connection creation. Configuration like clock setting and fine tuning + // should persist throughout the Policy lifetime after Reset calls. + Reset() + + // PreTx is called before writing to a frame. + // The outgoing frame options can be set by the Policy and will be respected if Frame offset >5. + // retransmitFrom is ignored unless within [snd.UNA, snd.NXT] and returned retransmit==true. + // newTransmitLimit sets the maximum number of new bytes to send over the wire (congestion control). + // If not implementing congestion control then newTransmitLimit=[TransmitUnlimited]. + PreTx(h *Handler, outgoingOpts Frame) (newTransmitLimit Size, retransmitFrom Value, retransmit bool) + // PostTx called on leaving the transmit path with the fully written frame. + PostTx(h *Handler, outgoing Frame) + + // PreRx is called by [Handler] on every incoming segment. + // PreRx can choose to drop segment if it returns keep=false. + PreRx(h *Handler, incoming Frame) (keep bool) + // PostRx is called by [Handler] after accepting an incoming segment. + // To access [ControlBlock.SendUNA] before incoming frame was processed save UNA in PreRx. + PostRx(h *Handler, prevState State, accepted Frame) +} diff --git a/tcp/policy_test.go b/tcp/policy_test.go new file mode 100644 index 0000000..283c20a --- /dev/null +++ b/tcp/policy_test.go @@ -0,0 +1,449 @@ +package tcp + +import ( + "math/rand" + "testing" + + "github.com/soypat/lneto/ethernet" +) + +// recordingPolicy records every hook invocation and lets the test steer what is +// returned to the Handler. It is the [Policy] counterpart driven by the Handler +// under test. +type recordingPolicy struct { + resets int + preRx []Segment + preTx int + postRx []Segment + postTx []txRecord + + // Values handed back to the Handler. + keep bool // PreRx result. Default true (see newRecordingPolicy). + rtxFrom Value + retransmit bool + txLimit Size // PreTx new-data limit. Default TransmitUnlimited (see newRecordingPolicy). + // writeOpts, when non-empty, is appended as TCP options by PreTx. + writeOpts []byte +} + +// txRecord is what PostTx observed on the emitted frame. +type txRecord struct { + seg Segment + offset uint8 + sport uint16 + dport uint16 +} + +func newRecordingPolicy() *recordingPolicy { + return &recordingPolicy{keep: true, txLimit: TransmitUnlimited} +} + +var _ Policy = (*recordingPolicy)(nil) + +func (p *recordingPolicy) Reset() { p.resets++ } + +func (p *recordingPolicy) PreRx(h *Handler, incoming Frame) bool { + p.preRx = append(p.preRx, incoming.Segment(len(incoming.Payload()))) + return p.keep +} + +func (p *recordingPolicy) PostRx(h *Handler, prevState State, accepted Frame) { + p.postRx = append(p.postRx, accepted.Segment(len(accepted.Payload()))) +} + +func (p *recordingPolicy) PreTx(h *Handler, outgoingOpts Frame) (Size, Value, bool) { + p.preTx++ + if len(p.writeOpts) > 0 { + // Raise the offset first: Options() is sized from it. + words := uint8(5 + (len(p.writeOpts)+3)/4) + outgoingOpts.SetOffsetAndFlags(words, 0) + copy(outgoingOpts.Options(), p.writeOpts) + } + return p.txLimit, p.rtxFrom, p.retransmit +} + +func (p *recordingPolicy) PostTx(h *Handler, outgoing Frame) { + offset, _ := outgoing.OffsetAndFlags() + p.postTx = append(p.postTx, txRecord{ + seg: outgoing.Segment(len(outgoing.Payload())), + offset: offset, + sport: outgoing.SourcePort(), + dport: outgoing.DestinationPort(), + }) +} + +// TestPolicy_DisabledByDefault verifies the Handler runs normally with no Policy +// installed: the transmit and receive paths never touch a nil Policy. +func TestPolicy_DisabledByDefault(t *testing.T) { + const mtu = ethernet.MaxMTU + rng := rand.New(rand.NewSource(1)) + client, server := newHandler(t, mtu, 3), newHandler(t, mtu, 3) + setupClientServer(t, rng, client, server) + + var buf [mtu]byte + establish(t, client, server, buf[:]) // must not panic on nil Policy. +} + +// TestPolicy_HooksInvoked verifies the Handler drives the full hook contract +// across a handshake: Reset on open, PreTx+PostTx on transmit, PreRx+PostRx on +// receive. +func TestPolicy_HooksInvoked(t *testing.T) { + const mtu = ethernet.MaxMTU + rng := rand.New(rand.NewSource(2)) + client, server := newHandler(t, mtu, 3), newHandler(t, mtu, 3) + + pol := newRecordingPolicy() + client.SetPolicy(pol) + + setupClientServer(t, rng, client, server) // OpenActive → reset → Reset(). + if pol.resets == 0 { + t.Fatal("Reset not called on open") + } + + var buf [mtu]byte + establish(t, client, server, buf[:]) + + if pol.preTx == 0 { + t.Fatal("PreTx never called on transmit") + } + if len(pol.postTx) == 0 { + t.Fatal("PostTx never called on transmit") + } + if pol.preTx < len(pol.postTx) { + t.Fatalf("PreTx calls=%d < PostTx calls=%d: PostTx must never fire without PreTx", pol.preTx, len(pol.postTx)) + } + // Client received the SYN-ACK and accepted it. + if len(pol.preRx) == 0 { + t.Fatal("PreRx never called on receive") + } + if len(pol.postRx) == 0 { + t.Fatal("PostRx never called on accepted receive") + } + // PostTx receives the segment actually emitted: the first is the SYN. + if !pol.postTx[0].seg.Flags.HasAny(FlagSYN) { + t.Fatalf("first PostTx segment flags=%s, want SYN", pol.postTx[0].seg.Flags) + } +} + +// TestPolicy_PostTxSeesWrittenFrame verifies PostTx observes the fully populated +// frame — ports, sequence numbers and payload length as emitted — and not the +// frame as it stood before the segment was written into it. +func TestPolicy_PostTxSeesWrittenFrame(t *testing.T) { + const mtu = ethernet.MaxMTU + rng := rand.New(rand.NewSource(6)) + client, server := newHandler(t, mtu, 3), newHandler(t, mtu, 3) + + pol := newRecordingPolicy() + client.SetPolicy(pol) + setupClientServer(t, rng, client, server) + var buf [mtu]byte + establish(t, client, server, buf[:]) + + data := []byte("payload") + if _, err := client.Write(data); err != nil { + t.Fatal("client write:", err) + } + clear(buf[:]) + n, err := client.Send(buf[:]) + if err != nil { + t.Fatal("client send:", err) + } + last := pol.postTx[len(pol.postTx)-1] + wantSeg := mustSegment(t, buf[:n], n-int(last.offset)*4) + if last.seg != wantSeg { + t.Fatalf("PostTx segment=%+v, want emitted %+v", last.seg, wantSeg) + } + if int(last.seg.DATALEN) != len(data) { + t.Fatalf("PostTx DATALEN=%d, want %d", last.seg.DATALEN, len(data)) + } + if last.sport != client.LocalPort() || last.dport != client.RemotePort() { + t.Fatalf("PostTx ports=%d→%d, want %d→%d", last.sport, last.dport, client.LocalPort(), client.RemotePort()) + } +} + +// TestPolicy_NoPostTxWithoutSegment verifies a transmit attempt that emits +// nothing still runs PreTx but never PostTx, so a Policy cannot mistake a +// no-op Send for a segment on the wire. +func TestPolicy_NoPostTxWithoutSegment(t *testing.T) { + const mtu = ethernet.MaxMTU + rng := rand.New(rand.NewSource(7)) + client, server := newHandler(t, mtu, 3), newHandler(t, mtu, 3) + + pol := newRecordingPolicy() + client.SetPolicy(pol) + setupClientServer(t, rng, client, server) + var buf [mtu]byte + establish(t, client, server, buf[:]) + + preTxBefore, postTxBefore := pol.preTx, len(pol.postTx) + n, err := client.Send(buf[:]) // Nothing queued: no segment. + if err != nil { + t.Fatal("client send:", err) + } + if n != 0 { + t.Fatalf("expected no segment, got %d bytes", n) + } + if pol.preTx != preTxBefore+1 { + t.Fatalf("PreTx calls=%d, want %d: PreTx must run on every attempt", pol.preTx, preTxBefore+1) + } + if len(pol.postTx) != postTxBefore { + t.Fatalf("PostTx calls=%d, want %d: no segment was emitted", len(pol.postTx), postTxBefore) + } +} + +// TestPolicy_PreTxOptions verifies options written by PreTx survive to the wire: +// the data offset accounts for them and the payload starts after them. +func TestPolicy_PreTxOptions(t *testing.T) { + const mtu = ethernet.MaxMTU + rng := rand.New(rand.NewSource(8)) + client, server := newHandler(t, mtu, 3), newHandler(t, mtu, 3) + + pol := newRecordingPolicy() + client.SetPolicy(pol) + setupClientServer(t, rng, client, server) + var buf [mtu]byte + establish(t, client, server, buf[:]) + + // One 4-byte option word: NOP,NOP,NOP,EOL. + opts := []byte{1, 1, 1, 0} + pol.writeOpts = opts + + data := []byte("payload") + if _, err := client.Write(data); err != nil { + t.Fatal("client write:", err) + } + clear(buf[:]) + n, err := client.Send(buf[:]) + if err != nil { + t.Fatal("client send:", err) + } + frm, err := NewFrame(buf[:n]) + if err != nil { + t.Fatal("frame:", err) + } + offset, _ := frm.OffsetAndFlags() + if offset != 6 { + t.Fatalf("data offset=%d, want 6 (header + one option word)", offset) + } + if got := frm.Options(); string(got) != string(opts) { + t.Fatalf("options=%v, want %v", got, opts) + } + if got := frm.Payload(); string(got) != string(data) { + t.Fatalf("payload=%q, want %q: options must not overlap data", got, data) + } + if n != int(offset)*4+len(data) { + t.Fatalf("frame length=%d, want %d", n, int(offset)*4+len(data)) + } +} + +// TestPolicy_PreRxDropsSegment verifies keep=false drops the segment before the +// state machine sees it: the payload is not buffered, connection state is +// untouched and PostRx never fires. +func TestPolicy_PreRxDropsSegment(t *testing.T) { + const mtu = ethernet.MaxMTU + rng := rand.New(rand.NewSource(4)) + client, server := newHandler(t, mtu, 3), newHandler(t, mtu, 3) + + pol := newRecordingPolicy() + server.SetPolicy(pol) + setupClientServer(t, rng, client, server) + var buf [mtu]byte + establish(t, client, server, buf[:]) // keep=true so handshake completes. + + // Now start dropping everything the server receives. + pol.keep = false + preRxBefore, postRxBefore := len(pol.preRx), len(pol.postRx) + + data := []byte("dropme") + if _, err := client.Write(data); err != nil { + t.Fatal("client write:", err) + } + clear(buf[:]) + n, err := client.Send(buf[:]) + if err != nil { + t.Fatal("client send:", err) + } + + if err := server.Recv(buf[:n]); err != nil { + t.Fatalf("dropped segment must return nil, got %v", err) + } + if len(pol.preRx) != preRxBefore+1 { + t.Fatalf("PreRx calls=%d, want %d (segment must reach PreRx)", len(pol.preRx), preRxBefore+1) + } + if len(pol.postRx) != postRxBefore { + t.Fatalf("PostRx calls=%d, want %d: a dropped segment was never accepted", len(pol.postRx), postRxBefore) + } + if server.BufferedInput() != 0 { + t.Fatalf("dropped segment must not be buffered, got %d bytes", server.BufferedInput()) + } + if server.State() != StateEstablished { + t.Fatalf("dropped segment must not change state, got %s", server.State()) + } +} + +// TestPolicy_PreTxRetransmit verifies a PreTx retransmit directive drives +// go-back-N: the Handler rewinds the send sequence and the transmit buffer +// together and re-emits already-sent, unacknowledged data from snd.UNA. +func TestPolicy_PreTxRetransmit(t *testing.T) { + const mtu = ethernet.MaxMTU + rng := rand.New(rand.NewSource(5)) + client, server := newHandler(t, mtu, 3), newHandler(t, mtu, 3) + + pol := newRecordingPolicy() + client.SetPolicy(pol) + setupClientServer(t, rng, client, server) + var buf [mtu]byte + establish(t, client, server, buf[:]) + + // Emit one data segment; server never ACKs, so it stays unacknowledged. + data := []byte("payload") + if _, err := client.Write(data); err != nil { + t.Fatal("client write:", err) + } + clear(buf[:]) + n, err := client.Send(buf[:]) + if err != nil { + t.Fatal("client send data:", err) + } + if n <= sizeHeaderTCP { + t.Fatal("expected data segment") + } + firstSeg := mustSegment(t, buf[:n], n-sizeHeaderTCP) + firstData := append([]byte(nil), buf[sizeHeaderTCP:n]...) + + // Direct go-back-N on the next transmit. + pol.rtxFrom, pol.retransmit = client.ControlBlock().SendUNA(), true + clear(buf[:]) + n, err = client.Send(buf[:]) + if err != nil { + t.Fatal("client send retransmit:", err) + } + if n <= sizeHeaderTCP { + t.Fatal("expected retransmitted data segment") + } + rtSeg := mustSegment(t, buf[:n], n-sizeHeaderTCP) + + if rtSeg.SEQ != firstSeg.SEQ { + t.Fatalf("retransmit SEQ=%d, want original UNA SEQ=%d (go-back-N)", rtSeg.SEQ, firstSeg.SEQ) + } + if rtSeg.DATALEN != firstSeg.DATALEN { + t.Fatalf("retransmit DATALEN=%d, want %d", rtSeg.DATALEN, firstSeg.DATALEN) + } + if got := buf[sizeHeaderTCP:n]; string(got) != string(firstData) { + t.Fatalf("retransmit payload=%q, want %q", got, firstData) + } +} + +// TestPolicy_PreTxRetransmitOutOfRange verifies an out-of-range rtxFrom is +// refused, leaving the send sequence and transmit buffer untouched. +func TestPolicy_PreTxRetransmitOutOfRange(t *testing.T) { + const mtu = ethernet.MaxMTU + rng := rand.New(rand.NewSource(9)) + client, server := newHandler(t, mtu, 3), newHandler(t, mtu, 3) + + pol := newRecordingPolicy() + client.SetPolicy(pol) + setupClientServer(t, rng, client, server) + var buf [mtu]byte + establish(t, client, server, buf[:]) + + if _, err := client.Write([]byte("payload")); err != nil { + t.Fatal("client write:", err) + } + clear(buf[:]) + if _, err := client.Send(buf[:]); err != nil { + t.Fatal("client send data:", err) + } + nxtBefore := client.ControlBlock().SendNext() + + // Well beyond snd.NXT: must be refused. + pol.rtxFrom, pol.retransmit = nxtBefore+1000, true + clear(buf[:]) + if _, err := client.Send(buf[:]); err != nil { + t.Fatal("client send:", err) + } + if got := client.ControlBlock().SendNext(); got != nxtBefore { + t.Fatalf("snd.NXT=%d, want unchanged %d: out-of-range rtxFrom must be refused", got, nxtBefore) + } +} + +// TestPolicy_TransmitLimit verifies the PreTx new-data limit caps the payload +// sent while leaving control segments free to go out: a zero limit suppresses +// data entirely, a partial limit truncates the segment. +func TestPolicy_TransmitLimit(t *testing.T) { + const mtu = ethernet.MaxMTU + rng := rand.New(rand.NewSource(10)) + client, server := newHandler(t, mtu, 3), newHandler(t, mtu, 3) + + pol := newRecordingPolicy() + client.SetPolicy(pol) + setupClientServer(t, rng, client, server) + var buf [mtu]byte + establish(t, client, server, buf[:]) + + const payload = "payload" + pol.txLimit = 0 + if _, err := client.Write([]byte(payload)); err != nil { + t.Fatal("client write:", err) + } + clear(buf[:]) + n, err := client.Send(buf[:]) + if err != nil { + t.Fatal("client send:", err) + } + if n > sizeHeaderTCP { + t.Fatalf("a zero limit must suppress new data, got %d payload bytes", n-sizeHeaderTCP) + } + + // A partial limit lets only that many bytes out. + pol.txLimit = 3 + clear(buf[:]) + n, err = client.Send(buf[:]) + if err != nil { + t.Fatal("client send under partial limit:", err) + } + if got := n - sizeHeaderTCP; got != int(pol.txLimit) { + t.Fatalf("got %d payload bytes, want the limit of %d", got, pol.txLimit) + } + + // Releasing the limit lets the rest of the data out. + pol.txLimit = TransmitUnlimited + clear(buf[:]) + n, err = client.Send(buf[:]) + if err != nil { + t.Fatal("client send after limit lifted:", err) + } + if got := n - sizeHeaderTCP; got != len(payload)-3 { + t.Fatalf("got %d payload bytes, want the remaining %d once unlimited", got, len(payload)-3) + } +} + +// TestPolicy_ResetOnReopen verifies Reset fires on every (re)open and on Abort, +// so a single Policy value can be reused across connection reuse. +func TestPolicy_ResetOnReopen(t *testing.T) { + const mtu = ethernet.MaxMTU + client := newHandler(t, mtu, 3) + pol := newRecordingPolicy() + client.SetPolicy(pol) + + if err := client.OpenActive(1234, 5678, 0); err != nil { + t.Fatal("open 1:", err) + } + afterOpen := pol.resets + if afterOpen == 0 { + t.Fatal("Reset not called on first open") + } + + client.Abort() + if pol.resets <= afterOpen { + t.Fatalf("Reset not called on Abort: resets=%d, want >%d", pol.resets, afterOpen) + } + afterAbort := pol.resets + + if err := client.OpenActive(1234, 5678, 0); err != nil { + t.Fatal("open 2:", err) + } + if pol.resets <= afterAbort { + t.Fatalf("Reset not called on reopen: resets=%d, want >%d", pol.resets, afterAbort) + } +} diff --git a/tcp/rto/integration_test.go b/tcp/rto/integration_test.go new file mode 100644 index 0000000..128e410 --- /dev/null +++ b/tcp/rto/integration_test.go @@ -0,0 +1,222 @@ +package rto + +import ( + "math/rand" + "testing" + "time" + + "github.com/soypat/lneto/ethernet" + "github.com/soypat/lneto/tcp" +) + +// sizeHeaderTCP is the fixed TCP header length. The tcp package's own constant +// is unexported and these tests live outside it. +const sizeHeaderTCP = 20 + +// TestRTO_HandlerRetransmitsAfterTimeout covers the seam between a Handler and +// its Policy, which the Timer unit tests do not: a lost data segment must be +// resent once the timer expires, with nothing arriving to prompt it. +func TestRTO_HandlerRetransmitsAfterTimeout(t *testing.T) { + const mtu = ethernet.MaxMTU + const maxpackets = 4 + rng := rand.New(rand.NewSource(5)) + client, server := newHandler(t, mtu, maxpackets), newHandler(t, mtu, maxpackets) + + var now int64 // injected monotonic clock, in nanoseconds + client.SetPolicy(newTimer(t, func() int64 { return now })) + + setupClientServer(t, rng, client, server) + var rawbuf [mtu]byte + establish(t, client, server, rawbuf[:]) + + data := []byte("hello") + if n, err := client.Write(data); err != nil || n != len(data) { + t.Fatal("client write:", n, err) + } + clear(rawbuf[:]) + n, err := client.Send(rawbuf[:]) + if err != nil || n == 0 { + t.Fatal("client send:", n, err) + } + // That frame is lost: it is never handed to the server. + + // Nothing may come back before the timer expires. + var probe [mtu]byte + if n, err := client.Send(probe[:]); err != nil || n != 0 { + t.Fatalf("client sent %d bytes before the RTO expired (err %v)", n, err) + } + + now += int64(3 * time.Second) // past the initial RTO and one backoff + + clear(probe[:]) + n, err = client.Send(probe[:]) + if err != nil { + t.Fatal("client send after RTO:", err) + } + if n == 0 { + t.Fatal("no retransmission after the RTO expired: the Policy directive is never applied") + } + if err := server.Recv(probe[:n]); err != nil { + t.Fatal("server refused the retransmission:", err) + } + got := make([]byte, 16) + nr, err := server.Read(got) + if err != nil || string(got[:nr]) != string(data) { + t.Fatalf("server read %q (%v), want %q", got[:nr], err, data) + } +} + +// TestRTO_HandlerRetransmitsAfterCloseWithUnackedData is the write-then-close +// case every server performs. With the last data segment lost, the FIN behind it +// sits above a gap the peer cannot cross, so FIN-WAIT-1 must still retransmit +// that data or both sides wait forever. +func TestRTO_HandlerRetransmitsAfterCloseWithUnackedData(t *testing.T) { + const mtu = ethernet.MaxMTU + const maxpackets = 4 + rng := rand.New(rand.NewSource(9)) + client, server := newHandler(t, mtu, maxpackets), newHandler(t, mtu, maxpackets) + + var now int64 + client.SetPolicy(newTimer(t, func() int64 { return now })) + + setupClientServer(t, rng, client, server) + var rawbuf [mtu]byte + establish(t, client, server, rawbuf[:]) + + data := []byte("last response bytes") + if n, err := client.Write(data); err != nil || n != len(data) { + t.Fatal("client write:", n, err) + } + clear(rawbuf[:]) + n, err := client.Send(rawbuf[:]) // this frame is lost in transit + if err != nil || n == 0 { + t.Fatal("client send:", n, err) + } + + // The application closes right after writing. + if err := client.Close(); err != nil { + t.Fatal("client close:", err) + } + var finbuf [mtu]byte + nfin, err := client.Send(finbuf[:]) // FIN (also lost, or simply unacked) + if err != nil { + t.Fatal("client send FIN:", err) + } + t.Logf("state after close: %s (FIN frame %d bytes)", client.State(), nfin) + + now += int64(3 * time.Second) // past the RTO + + var probe [mtu]byte + n, err = client.Send(probe[:]) + if err != nil { + t.Fatal("client send after RTO:", err) + } + if n == 0 { + t.Fatalf("no retransmission in %s: unacknowledged data is stranded by the close", client.State()) + } + if err := server.Recv(probe[:n]); err != nil { + t.Fatal("server refused the retransmission:", err) + } + got := make([]byte, 32) + nr, err := server.Read(got) + if err != nil || string(got[:nr]) != string(data) { + t.Fatalf("server read %q (%v), want %q", got[:nr], err, data) + } +} + +// newTimer returns a Timer driven by nanotime, ready to install as a [tcp.Policy]. +func newTimer(t *testing.T, nanotime func() int64) *Timer { + t.Helper() + r := new(Timer) + err := r.Configure(nanotime) + if err != nil { + t.Fatal(err) + } + return r +} + +// The handshake helpers below mirror those in the tcp package's own tests, which +// are unexported and so unavailable here. They drive two Handlers against each +// other over a single packet buffer, with no network in between. + +func newHandler(t *testing.T, mtu, minpackets int) *tcp.Handler { + t.Helper() + h := new(tcp.Handler) + err := h.SetBuffers(make([]byte, mtu), make([]byte, mtu), minpackets) + if err != nil { + t.Fatal(err) + } + return h +} + +func setupClientServer(t *testing.T, rng *rand.Rand, client, server *tcp.Handler) { + t.Helper() + err := server.OpenListen(uint16(rng.Uint32()), 0) + if err != nil { + t.Fatal(err) + } + err = client.OpenActive(uint16(rng.Uint32()), server.LocalPort(), 0) + if err != nil { + t.Fatal(err) + } + if !client.AwaitingSynSend() { + t.Fatal("client in wrong state") + } + if !server.AwaitingSynAck() { + t.Fatal("server in wrong state") + } +} + +func establish(t *testing.T, client, server *tcp.Handler, packetBuf []byte) { + t.Helper() + if client.State() != tcp.StateClosed { + t.Fatal("client in wrong state") + } else if server.State() != tcp.StateListen { + t.Fatal("server in wrong state") + } + clear(packetBuf) + + // Commence 3-way handshake: client sends SYN, server sends SYN-ACK, client sends ACK. + n, err := client.Send(packetBuf) + if err != nil { + t.Fatal("client sending:", err) + } else if n < sizeHeaderTCP { + t.Fatal("expected client to send SYN packet") + } else if client.State() != tcp.StateSynSent { + t.Fatal("client did not transition to SynSent state:", client.State().String()) + } + err = server.Recv(packetBuf[:n]) // Server receives SYN. + if err != nil { + t.Fatal(err) + } else if server.State() != tcp.StateSynRcvd { + t.Fatal("server did not transition to SynReceived state:", server.State().String()) + } + + clear(packetBuf) + n, err = server.Send(packetBuf) // Server sends SYNACK. + if err != nil { + t.Fatal("server sending:", err) + } else if n < sizeHeaderTCP { + t.Fatal("expected server to send SYNACK packet") + } + err = client.Recv(packetBuf[:n]) // Client receives SYNACK, is established but must send ACK. + if err != nil { + t.Fatal(err) + } else if client.State() != tcp.StateEstablished { + t.Fatal("client did not transition to Established state:", client.State().String()) + } + + clear(packetBuf) + n, err = client.Send(packetBuf) // Client sends ACK. + if err != nil { + t.Fatal("client sending ACK:", err) + } else if n < sizeHeaderTCP { + t.Fatal("expected client to send ACK packet") + } + err = server.Recv(packetBuf[:n]) // Server receives ACK. + if err != nil { + t.Fatal(err) + } else if server.State() != tcp.StateEstablished { + t.Fatal("server did not transition to Established state on ACK receive:", server.State().String()) + } +} diff --git a/tcp/rto.go b/tcp/rto/timer.go similarity index 50% rename from tcp/rto.go rename to tcp/rto/timer.go index 180d518..9ea4a07 100644 --- a/tcp/rto.go +++ b/tcp/rto/timer.go @@ -1,6 +1,11 @@ -package tcp +package rto -import "time" +import ( + "time" + + "github.com/soypat/lneto" + "github.com/soypat/lneto/tcp" +) // RFC 6298 retransmission-timeout (RTO) parameters. The algorithm keeps a // single retransmission timer per connection (RFC 6298 §5): the timer is @@ -30,61 +35,80 @@ const ( backoffMax = 12 ) -// RTO implements the RFC 6298 round-trip-time estimator and the single -// retransmission timer as a [LossRecovery]. Construct it with new(RTO) and hand -// it to [ConnConfig.LossRecovery]; the connection calls [RTO.Reset] on open, so -// the zero value is ready to use. +// Timer implements the RFC 6298 round-trip-time estimator and the single +// retransmission timer as a [tcp.Policy]. Construct it with [NewTimer] and hand +// it to [tcp.ConnConfig.Policy]. // -// RTO is a pure, reactive state machine: it observes the segments a connection -// sends and receives (via the LossRecovery hooks) and the monotonic time handed -// in at each hook, and from those alone derives RTT estimates and retransmission -// decisions. It holds no clock and allocates nothing, which keeps it -// deterministic for unit testing (see issue #140). +// Timer is a pure, reactive state machine: it observes the segments a connection +// sends and receives (via the tcp.Policy hooks) and from those alone derives RTT +// estimates and retransmission decisions. The tcp package holds no clock, so the +// Timer carries its own; injecting it keeps the estimator deterministic for unit +// testing (see issue #140). // -// RTO tracks its own shadow of the send sequence space purely from the segments -// it observes: [RTO.PostTx] advances the highest sequence sent and [RTO.PreRx] +// Timer tracks its own shadow of the send sequence space purely from the segments +// it observes: [Timer.PostTx] advances the highest sequence sent and [Timer.PreRx] // advances the highest sequence acknowledged. This is what lets it manage the // timer (RFC 6298 §5.2/§5.3) without reaching into the tcp state machine, and it // is also how retransmissions are distinguished for Karn's algorithm — a segment // whose sequence space is not beyond the shadow snd.NXT is a retransmission and // is never RTT-sampled. -type RTO struct { +type Timer struct { + // nanotime is the monotonic time source in nanoseconds. Preserved by Reset. + nanotime func() int64 + srtt time.Duration // smoothed round-trip time (SRTT). rttvar time.Duration // round-trip-time variation (RTTVAR). rto time.Duration // current retransmission timeout. haveRTT bool // false until the first RTT sample is taken. // Shadow of the send sequence space, derived from observed segments. - haveSeq bool // false until the first data segment is observed. - sndUNA Value // highest acknowledged sequence number seen on the wire. - sndNXT Value // one past the highest sequence number sent. + haveSeq bool // false until the first data segment is observed. + sndUNA tcp.Value // highest acknowledged sequence number seen on the wire. + sndNXT tcp.Value // one past the highest sequence number sent. // RTT sampling state (Karn's algorithm, RFC 6298 §3): at most one segment is // timed at a time and retransmitted segments are never sampled. timing bool - timedSeq Value // ACK at or beyond this value completes the sample. - timedAt int64 // send time (monotonic ns) of the timed segment. + timedSeq tcp.Value // ACK at or beyond this value completes the sample. + timedAt int64 // send time (monotonic ns) of the timed segment. // Retransmission timer state. running bool deadline int64 // time (monotonic ns) at which the timer expires. backoff uint8 // consecutive timeouts, for exponential backoff. + + // expirations counts timeouts since Reset. It exists so a policy sharing this + // timer can notice a timeout it did not itself drive: a congestion controller + // must collapse its window on one, and a policy that composes the timer as a + // peer never sees the timer's own directive. + expirations uint32 } -var _ LossRecovery = (*RTO)(nil) +var _ tcp.Policy = (*Timer)(nil) -// Reset returns the estimator to its pre-connection state with the initial RTO. -// It implements [LossRecovery] and is called when the connection opens or aborts -// so the estimator can be reused across connection reuse. -func (r *RTO) Reset() { *r = RTO{rto: rtoInitial} } +// Configure prepares the Timer for use with nanotime, the monotonic time source +// in nanoseconds (the func() int64 convention used across lneto). It must be +// called before the connection is opened. +func (r *Timer) Configure(nanotime func() int64) error { + if nanotime == nil { + return lneto.ErrMissingHALConfig // The estimator cannot run without a clock. + } + *r = Timer{rto: rtoInitial, nanotime: nanotime} + return nil +} + +// Reset returns the estimator to its pre-connection state with the initial RTO, +// preserving the configured clock. It implements [tcp.Policy] and is called when +// the connection opens or aborts so the estimator survives connection reuse. +func (r *Timer) Reset() { *r = Timer{rto: rtoInitial, nanotime: r.nanotime} } // SmoothedRTT returns the current smoothed round-trip time (SRTT), or zero // before the first RTT measurement. It is concrete-type introspection and is -// intentionally not part of [LossRecovery]. -func (r *RTO) SmoothedRTT() time.Duration { return r.srtt } +// intentionally not part of [tcp.Policy]. +func (r *Timer) SmoothedRTT() time.Duration { return r.srtt } // CurrentRTO returns the timeout currently in effect, clamped to [rtoMin, rtoMax]. -func (r *RTO) CurrentRTO() time.Duration { +func (r *Timer) CurrentRTO() time.Duration { rto := r.rto if rto < rtoMin { rto = rtoMin @@ -95,23 +119,46 @@ func (r *RTO) CurrentRTO() time.Duration { } // Running reports whether the retransmission timer is currently armed. -func (r *RTO) Running() bool { return r.running } +func (r *Timer) Running() bool { return r.running } + +// Expirations returns how many times the retransmission timer has expired since +// [Timer.Reset]. A policy that shares this timer rather than driving it watches +// this for a change to learn that a timeout happened, since it never sees the +// timer's own directive. It is concrete-type introspection and is intentionally +// not part of [tcp.Policy]. +func (r *Timer) Expirations() uint32 { return r.expirations } // NextDeadline returns the monotonic-nanosecond instant at which the timer -// expires, or 0 when it is not armed. It implements [LossRecovery]. -func (r *RTO) NextDeadline() int64 { +// expires, or 0 when it is not armed. It is concrete-type introspection, not +// part of [tcp.Policy]: an event loop that wants to schedule against the RTO +// holds the Timer it configured and reads this. +func (r *Timer) NextDeadline() int64 { if !r.running { return 0 } return r.deadline } -// PreRx samples the RTT and manages the retransmission timer from a received -// segment (RFC 6298 §5.2/§5.3). It implements [LossRecovery] and always keeps -// the segment (the estimator never drops traffic). -func (r *RTO) PreRx(incoming Segment, now int64) RxDirective { - if !r.haveSeq || !incoming.Flags.HasAny(FlagACK) { - return RxDirective{Keep: true} +// PreRx keeps every segment: the estimator never drops traffic and records +// nothing before the connection has decided whether the segment counts. It +// implements [tcp.Policy]. +func (r *Timer) PreRx(h *tcp.Handler, incoming tcp.Frame) bool { + return true +} + +// PostRx samples the RTT and manages the retransmission timer from a segment the +// connection accepted (RFC 6298 §5.2/§5.3). It implements [tcp.Policy]. +// +// Only accepted segments reach here. Acting on a refused one would let an +// acknowledgement the state machine rejected, for data never sent, collapse the +// backoff and take a bogus RTT sample. +func (r *Timer) PostRx(h *tcp.Handler, prevState tcp.State, accepted tcp.Frame) { + r.postRx(accepted.Segment(len(accepted.Payload())), r.nanotime()) +} + +func (r *Timer) postRx(incoming tcp.Segment, now int64) { + if !r.haveSeq || !incoming.Flags.HasAny(tcp.FlagACK) { + return } ack := incoming.ACK if r.timing && !ack.LessThan(r.timedSeq) { @@ -132,18 +179,24 @@ func (r *RTO) PreRx(incoming Segment, now int64) RxDirective { r.running = true r.deadline = now + int64(r.CurrentRTO()) } - return RxDirective{Keep: true} } // PreTx reports whether the retransmission timer has expired and, if so, applies // the RFC 6298 §5.4–§5.6 timeout response — discard the outstanding RTT sample -// (Karn), back the RTO off exponentially and restart the timer — returning a -// directive that asks the connection to retransmit from snd.UNA (go-back-N). It -// implements [LossRecovery]. -func (r *RTO) PreTx(now int64) TxDirective { +// (Karn), back the RTO off exponentially and restart the timer — and asks the +// connection to retransmit from snd.UNA (go-back-N). It writes no TCP options +// and imposes no transmit limit: retransmission timing needs neither, and +// congestion control belongs to a Policy composing this timer. It implements +// [tcp.Policy]. +func (r *Timer) PreTx(h *tcp.Handler, outgoingOpts tcp.Frame) (newTransmitLimit tcp.Size, rtxFrom tcp.Value, retransmit bool) { + return r.preTx(r.nanotime(), h.ControlBlock().SendUNA()) +} + +func (r *Timer) preTx(now int64, una tcp.Value) (newTransmitLimit tcp.Size, rtxFrom tcp.Value, retransmit bool) { if !r.running || now < r.deadline || r.sndUNA == r.sndNXT { - return TxDirective{} + return tcp.TransmitUnlimited, 0, false } + r.expirations++ r.timing = false // §5.4: do not sample a retransmitted segment. if r.backoff < backoffMax { r.backoff++ @@ -151,20 +204,24 @@ func (r *RTO) PreTx(now int64) TxDirective { } r.running = true r.deadline = now + int64(r.CurrentRTO()) - return TxDirective{RetransmitAll: true} + return tcp.TransmitUnlimited, una, true } // PostTx records an emitted segment: it advances the shadow send sequence, // begins timing newly transmitted data (RFC 6298 §3) and arms the timer (§5.1). // Segments that do not extend the send sequence are retransmissions and are // never RTT-sampled (Karn's algorithm). Control-only segments (no data) are -// ignored. It implements [LossRecovery]. -func (r *RTO) PostTx(outgoing Segment, now int64) { +// ignored. It implements [tcp.Policy]. +func (r *Timer) PostTx(h *tcp.Handler, outgoing tcp.Frame) { + r.postTx(outgoing.Segment(len(outgoing.Payload())), r.nanotime()) +} + +func (r *Timer) postTx(outgoing tcp.Segment, now int64) { if outgoing.DATALEN == 0 { return // only data segments are timed / arm the RTO. } segStart := outgoing.SEQ - segEnd := segStart + Value(outgoing.LEN()) + segEnd := segStart + tcp.Value(outgoing.LEN()) if !r.haveSeq { r.haveSeq = true r.sndUNA = segStart @@ -189,9 +246,20 @@ func (r *RTO) PostTx(outgoing Segment, now int64) { } } +// ObserveRTT folds a round-trip measurement taken by other means into the +// estimator, for a policy that composes this timer and can measure the round trip +// more accurately than acknowledgement timing allows. The RFC 7323 timestamp echo +// is the case this exists for. +// +// Unlike the timer's own sampling this does not apply Karn's algorithm, because a +// sample derived from an echoed timestamp is unambiguous even when the segment +// carrying it was a retransmission (RFC 7323 §4.1). Non-positive samples are +// ignored. +func (r *Timer) ObserveRTT(rtt time.Duration) { r.updateRTT(rtt) } + // updateRTT folds a round-trip measurement into SRTT/RTTVAR/RTO using the // integer-shift form of RFC 6298 §2.2/§2.3. -func (r *RTO) updateRTT(sample time.Duration) { +func (r *Timer) updateRTT(sample time.Duration) { if sample <= 0 { return } diff --git a/tcp/rto/timer_test.go b/tcp/rto/timer_test.go new file mode 100644 index 0000000..14355a3 --- /dev/null +++ b/tcp/rto/timer_test.go @@ -0,0 +1,311 @@ +package rto + +import ( + "testing" + "time" + + "github.com/soypat/lneto/tcp" +) + +const rtoMs = int64(time.Millisecond) + +// dataSeg builds a data segment of datalen octets starting at seq. +func dataSeg(seq uint32, datalen int) tcp.Segment { + return tcp.Segment{SEQ: tcp.Value(seq), DATALEN: tcp.Size(datalen), Flags: tcp.FlagPSH | tcp.FlagACK} +} + +// ackSeg builds a bare ACK acknowledging up to ack. +func ackSeg(ack uint32) tcp.Segment { + return tcp.Segment{ACK: tcp.Value(ack), Flags: tcp.FlagACK} +} + +func newRTO() *Timer { + var r Timer + if err := r.Configure(func() int64 { return 0 }); err != nil { + panic(err) + } + return &r +} + +// frameOf renders a segment as the wire frame the [tcp.Policy] hooks receive. +func frameOf(t *testing.T, s tcp.Segment) tcp.Frame { + t.Helper() + frm, err := tcp.NewFrame(make([]byte, 20+int(s.DATALEN))) + if err != nil { + t.Fatal(err) + } + frm.SetSegment(s, 5) + return frm +} + +func TestRTO_Configure(t *testing.T) { + var r Timer + if err := r.Configure(nil); err == nil { + t.Error("Configure must reject a nil clock") + } + if err := r.Configure(func() int64 { return 0 }); err != nil { + t.Fatal(err) + } + if r.nanotime == nil { + t.Fatal("clock not stored") + } + r.Reset() + if r.nanotime == nil { + t.Error("Reset must preserve the configured clock") + } +} + +func TestRTO_Reset(t *testing.T) { + r := newRTO() + r.Reset() + if r.rto != rtoInitial { + t.Errorf("initial rto=%v, want %v", r.rto, rtoInitial) + } + if r.CurrentRTO() != rtoInitial { + t.Errorf("CurrentRTO=%v, want %v", r.CurrentRTO(), rtoInitial) + } + if r.haveRTT { + t.Error("haveRTT should be false before first sample") + } + if r.Running() || r.NextDeadline() != 0 { + t.Error("timer must be disarmed after Reset") + } +} + +// TestRTO_ArmOnSendSampleOnAck sends data, verifies the timer arms, then acks it +// and verifies an RTT sample is taken and the timer stops once all data is acked. +func TestRTO_ArmOnSendSampleOnAck(t *testing.T) { + r := newRTO() + const iss = uint32(1000) + + r.postTx(dataSeg(iss, 100), 0) + if !r.Running() { + t.Fatal("timer must arm after sending data") + } + if r.NextDeadline() != int64(rtoInitial) { + t.Errorf("deadline=%d, want %d", r.NextDeadline(), int64(rtoInitial)) + } + + // ACK arrives one RTT (40ms) later covering all sent data. + if !r.PreRx(nil, frameOf(t, ackSeg(iss+100))) { + t.Error("PreRx must keep the segment") + } + r.postRx(ackSeg(iss+100), 40*rtoMs) + if r.Running() { + t.Error("timer must stop once all data is acknowledged") + } + if r.SmoothedRTT() != 40*time.Millisecond { + t.Errorf("srtt=%v, want 40ms", r.SmoothedRTT()) + } +} + +// TestRTO_RetransmitOnTimeout verifies PreTx directs a go-back-N retransmit once +// the deadline passes with data outstanding, and backs the RTO off. +func TestRTO_RetransmitOnTimeout(t *testing.T) { + r := newRTO() + const iss = uint32(1000) + r.postTx(dataSeg(iss, 100), 0) + + if _, _, rtx := r.preTx(int64(rtoInitial)-1, tcp.Value(iss)); rtx { + t.Fatal("must not retransmit before the deadline") + } + limit, from, rtx := r.preTx(int64(rtoInitial), tcp.Value(iss)) + if !rtx { + t.Fatal("RTO must fire at the deadline with data outstanding") + } + if limit != tcp.TransmitUnlimited { + t.Error("the estimator never limits new data") + } + if from != tcp.Value(iss) { + t.Errorf("retransmit from %d, want snd.UNA=%d", from, iss) + } + if r.CurrentRTO() != 2*rtoInitial { + t.Errorf("rto=%v after one backoff, want %v", r.CurrentRTO(), 2*rtoInitial) + } + // The connection resends from snd.UNA; postTx sees a retransmission. + r.postTx(dataSeg(iss, 100), int64(rtoInitial)) + if r.timing { + t.Error("retransmitted segment must not be RTT-sampled (Karn)") + } +} + +// TestRTO_KarnNoSampleOnRetransmittedAck verifies that after a retransmission the +// ACK does not produce an RTT sample (Karn's algorithm). +func TestRTO_KarnNoSampleOnRetransmittedAck(t *testing.T) { + r := newRTO() + const iss = uint32(1000) + r.postTx(dataSeg(iss, 100), 0) + // Timeout and retransmit. + r.preTx(int64(rtoInitial), tcp.Value(iss)) + r.postTx(dataSeg(iss, 100), int64(rtoInitial)) + // ACK now arrives; no sample should be taken since timing was discarded. + r.postRx(ackSeg(iss+100), int64(rtoInitial)+10*rtoMs) + if r.haveRTT { + t.Error("no RTT sample should exist after a retransmission (Karn)") + } +} + +// TestRTO_TimerRestartsWhilePartiallyAcked verifies the timer restarts (not +// stops) when an ACK advances UNA but data remains in flight (RFC 6298 §5.3). +func TestRTO_TimerRestartsWhilePartiallyAcked(t *testing.T) { + r := newRTO() + const iss = uint32(1000) + r.postTx(dataSeg(iss, 100), 0) + r.postTx(dataSeg(iss+100, 100), 0) // 200 octets outstanding, iss..iss+200. + + r.postRx(ackSeg(iss+100), 40*rtoMs) // acks first 100 only. + if !r.Running() { + t.Fatal("timer must remain armed while data is still in flight") + } + if r.NextDeadline() != 40*rtoMs+int64(r.CurrentRTO()) { + t.Errorf("deadline=%d, want %d", r.NextDeadline(), 40*rtoMs+int64(r.CurrentRTO())) + } +} + +// TestRTO_NoArmWithoutData verifies control-only segments neither arm the timer +// nor start an RTT sample. +func TestRTO_NoArmWithoutData(t *testing.T) { + r := newRTO() + r.postTx(tcp.Segment{SEQ: 1000, Flags: tcp.FlagACK}, 0) // pure ACK, DATALEN==0. + if r.Running() || r.timing { + t.Error("pure control segment must not arm the timer or start a sample") + } +} + +// TestRTO_BackoffCollapsesOnValidSample verifies a valid RTT measurement +// collapses the exponential backoff counter (RFC 6298 §5.7). +func TestRTO_BackoffCollapsesOnValidSample(t *testing.T) { + r := newRTO() + const iss = uint32(1000) + r.postTx(dataSeg(iss, 100), 0) + r.preTx(int64(rtoInitial), tcp.Value(iss)) // one timeout: backoff=1. + r.postTx(dataSeg(iss, 100), int64(rtoInitial)) // retransmit (no sample). + if r.backoff != 1 { + t.Fatalf("backoff=%d, want 1 after a timeout", r.backoff) + } + // New data sent and freshly sampled, then acked. + r.postTx(dataSeg(iss+100, 100), int64(rtoInitial)+rtoMs) + r.postRx(ackSeg(iss+200), int64(rtoInitial)+30*rtoMs) + if r.backoff != 0 { + t.Errorf("backoff=%d, want 0 after a valid RTT sample", r.backoff) + } +} + +// TestRTO_Clamped verifies CurrentRTO is clamped to [rtoMin, rtoMax]. +func TestRTO_Clamped(t *testing.T) { + r := newRTO() + r.rto = time.Nanosecond + if got := r.CurrentRTO(); got != rtoMin { + t.Errorf("CurrentRTO=%v, want floor %v", got, rtoMin) + } + r.rto = time.Hour + if got := r.CurrentRTO(); got != rtoMax { + t.Errorf("CurrentRTO=%v, want ceiling %v", got, rtoMax) + } +} + +// TestRTO_UpdateRTTFirstSample verifies the first-measurement initialization of +// SRTT/RTTVAR (RFC 6298 §2.2). +func TestRTO_UpdateRTTFirstSample(t *testing.T) { + r := newRTO() + r.updateRTT(100 * time.Millisecond) + if r.srtt != 100*time.Millisecond { + t.Errorf("srtt=%v, want 100ms", r.srtt) + } + if r.rttvar != 50*time.Millisecond { + t.Errorf("rttvar=%v, want 50ms", r.rttvar) + } + // RTO = SRTT + K*RTTVAR = 100 + 4*50 = 300ms. + if r.rto != 300*time.Millisecond { + t.Errorf("rto=%v, want 300ms", r.rto) + } +} + +// TestRTO_PolicyHooksDeriveFromFrame exercises Timer through the [tcp.Policy] +// hooks, verifying it reads the segment out of the frame it is handed: sending +// data arms a deadline and a full ACK disarms it and yields the RTT sample. +func TestRTO_PolicyHooksDeriveFromFrame(t *testing.T) { + var clock int64 + var r Timer + if err := r.Configure(func() int64 { return clock }); err != nil { + t.Fatal(err) + } + var pol tcp.Policy = &r + pol.Reset() + + pol.PostTx(nil, frameOf(t, dataSeg(1000, 100))) + if r.NextDeadline() == 0 { + t.Fatal("expected an armed deadline after sending data") + } + clock = 10 * rtoMs + if !pol.PreRx(nil, frameOf(t, ackSeg(1100))) { + t.Error("PreRx must keep") + } + pol.PostRx(nil, tcp.StateEstablished, frameOf(t, ackSeg(1100))) + if r.NextDeadline() != 0 { + t.Error("expected disarmed timer after full ack") + } + if r.SmoothedRTT() != 10*time.Millisecond { + t.Errorf("srtt=%v, want 10ms sampled through the hooks", r.SmoothedRTT()) + } +} + +// TestRTO_PreRxNeverDrops verifies the estimator keeps every segment and records +// nothing at PreRx time. Dropping is not its business, and the connection has not +// yet judged the segment: an acknowledgement for data never sent would otherwise +// collapse the backoff and take a bogus round-trip sample. Only accepted segments +// reach PostRx, which the Handler guarantees. +func TestRTO_PreRxNeverDrops(t *testing.T) { + r := newRTO() + const iss = uint32(1000) + r.postTx(dataSeg(iss, 100), 0) + armed := r.NextDeadline() + if armed == 0 { + t.Fatal("timer must be armed after sending data") + } + + // An acknowledgement far beyond anything sent, which the connection refuses. + if !r.PreRx(nil, frameOf(t, ackSeg(iss+100000))) { + t.Error("PreRx must keep: dropping is not the estimator's business") + } + if r.NextDeadline() != armed { + t.Errorf("deadline moved to %d at PreRx, want it left at %d", r.NextDeadline(), armed) + } + if r.SmoothedRTT() != 0 { + t.Errorf("took an RTT sample of %v at PreRx", r.SmoothedRTT()) + } + if !r.Running() { + t.Error("timer disarmed at PreRx") + } +} + +// TestRTO_RetransmitsZeroWindowProbe verifies the timer takes over the periodic +// probing of a closed send window. A zero-window probe is a single octet the peer +// cannot accept, so it goes unacknowledged; the timer must keep resending it, with +// exponential backoff, which is the persist-timer behaviour of RFC 9293 §3.8.6.1. +// The tcp package relies on this and refuses to probe without a policy installed. +func TestRTO_RetransmitsZeroWindowProbe(t *testing.T) { + r := newRTO() + const iss = uint32(5000) + probe := dataSeg(iss, 1) // The one-octet probe. + r.postTx(probe, 0) + + now := int64(rtoInitial) + prevRTO := r.CurrentRTO() + for attempt := 1; attempt <= 4; attempt++ { + _, from, rtx := r.preTx(now, tcp.Value(iss)) + if !rtx { + t.Fatalf("attempt %d: timer did not fire; the probe would never be resent", attempt) + } + if from != tcp.Value(iss) { + t.Errorf("attempt %d: retransmit from %d, want the probe octet at %d", attempt, from, iss) + } + if got := r.CurrentRTO(); got <= prevRTO { + t.Errorf("attempt %d: rto %v did not back off past %v", attempt, got, prevRTO) + } + prevRTO = r.CurrentRTO() + // The peer still cannot accept the octet, so it stays unacknowledged. + r.postTx(probe, now) + now += int64(prevRTO) + } +} diff --git a/tcp/rto_test.go b/tcp/rto_test.go deleted file mode 100644 index 44686a2..0000000 --- a/tcp/rto_test.go +++ /dev/null @@ -1,206 +0,0 @@ -package tcp - -import ( - "testing" - "time" -) - -const rtoMs = int64(time.Millisecond) - -// rtoDataSeg builds a data segment of datalen octets starting at seq. -func rtoDataSeg(seq uint32, datalen int) Segment { - return Segment{SEQ: Value(seq), DATALEN: Size(datalen), Flags: FlagPSH | FlagACK} -} - -// rtoAckSeg builds a bare ACK acknowledging up to ack. -func rtoAckSeg(ack uint32) Segment { - return Segment{ACK: Value(ack), Flags: FlagACK} -} - -func newRTO() *RTO { - var r RTO - r.Reset() - return &r -} - -func TestRTO_Reset(t *testing.T) { - var r RTO - r.Reset() - if r.rto != rtoInitial { - t.Errorf("initial rto=%v, want %v", r.rto, rtoInitial) - } - if r.CurrentRTO() != rtoInitial { - t.Errorf("CurrentRTO=%v, want %v", r.CurrentRTO(), rtoInitial) - } - if r.haveRTT { - t.Error("haveRTT should be false before first sample") - } - if r.Running() || r.NextDeadline() != 0 { - t.Error("timer must be disarmed after Reset") - } -} - -// TestRTO_ArmOnSendSampleOnAck sends data, verifies the timer arms, then acks it -// and verifies an RTT sample is taken and the timer stops once all data is acked. -func TestRTO_ArmOnSendSampleOnAck(t *testing.T) { - r := newRTO() - const iss = uint32(1000) - - r.PostTx(rtoDataSeg(iss, 100), 0) - if !r.Running() { - t.Fatal("timer must arm after sending data") - } - if r.NextDeadline() != int64(rtoInitial) { - t.Errorf("deadline=%d, want %d", r.NextDeadline(), int64(rtoInitial)) - } - - // ACK arrives one RTT (40ms) later covering all sent data. - dir := r.PreRx(rtoAckSeg(iss+100), 40*rtoMs) - if !dir.Keep { - t.Error("PreRx must keep the segment") - } - if r.Running() { - t.Error("timer must stop once all data is acknowledged") - } - if r.SmoothedRTT() != 40*time.Millisecond { - t.Errorf("srtt=%v, want 40ms", r.SmoothedRTT()) - } -} - -// TestRTO_RetransmitOnTimeout verifies PreTx directs a go-back-N retransmit once -// the deadline passes with data outstanding, and backs the RTO off. -func TestRTO_RetransmitOnTimeout(t *testing.T) { - r := newRTO() - const iss = uint32(1000) - r.PostTx(rtoDataSeg(iss, 100), 0) - - if r.PreTx(int64(rtoInitial) - 1).RetransmitAll { - t.Fatal("must not retransmit before the deadline") - } - dir := r.PreTx(int64(rtoInitial)) - if !dir.RetransmitAll { - t.Fatal("RTO must fire at the deadline with data outstanding") - } - if r.CurrentRTO() != 2*rtoInitial { - t.Errorf("rto=%v after one backoff, want %v", r.CurrentRTO(), 2*rtoInitial) - } - // The connection resends from snd.UNA; PostTx sees a retransmission. - r.PostTx(rtoDataSeg(iss, 100), int64(rtoInitial)) - if r.timing { - t.Error("retransmitted segment must not be RTT-sampled (Karn)") - } -} - -// TestRTO_KarnNoSampleOnRetransmittedAck verifies that after a retransmission the -// ACK does not produce an RTT sample (Karn's algorithm). -func TestRTO_KarnNoSampleOnRetransmittedAck(t *testing.T) { - r := newRTO() - const iss = uint32(1000) - r.PostTx(rtoDataSeg(iss, 100), 0) - // Timeout and retransmit. - r.PreTx(int64(rtoInitial)) - r.PostTx(rtoDataSeg(iss, 100), int64(rtoInitial)) - // ACK now arrives; no sample should be taken since timing was discarded. - r.PreRx(rtoAckSeg(iss+100), int64(rtoInitial)+10*rtoMs) - if r.haveRTT { - t.Error("no RTT sample should exist after a retransmission (Karn)") - } -} - -// TestRTO_TimerRestartsWhilePartiallyAcked verifies the timer restarts (not -// stops) when an ACK advances UNA but data remains in flight (RFC 6298 §5.3). -func TestRTO_TimerRestartsWhilePartiallyAcked(t *testing.T) { - r := newRTO() - const iss = uint32(1000) - r.PostTx(rtoDataSeg(iss, 100), 0) - r.PostTx(rtoDataSeg(iss+100, 100), 0) // 200 octets outstanding, iss..iss+200. - - dir := r.PreRx(rtoAckSeg(iss+100), 40*rtoMs) // acks first 100 only. - if !r.Running() { - t.Fatal("timer must remain armed while data is still in flight") - } - if r.NextDeadline() != 40*rtoMs+int64(r.CurrentRTO()) { - t.Errorf("deadline=%d, want %d", r.NextDeadline(), 40*rtoMs+int64(r.CurrentRTO())) - } - if !dir.Keep { - t.Error("PreRx must keep the segment") - } -} - -// TestRTO_NoArmWithoutData verifies control-only segments neither arm the timer -// nor start an RTT sample. -func TestRTO_NoArmWithoutData(t *testing.T) { - r := newRTO() - r.PostTx(Segment{SEQ: 1000, Flags: FlagACK}, 0) // pure ACK, DATALEN==0. - if r.Running() || r.timing { - t.Error("pure control segment must not arm the timer or start a sample") - } -} - -// TestRTO_BackoffCollapsesOnValidSample verifies a valid RTT measurement -// collapses the exponential backoff counter (RFC 6298 §5.7). -func TestRTO_BackoffCollapsesOnValidSample(t *testing.T) { - r := newRTO() - const iss = uint32(1000) - r.PostTx(rtoDataSeg(iss, 100), 0) - r.PreTx(int64(rtoInitial)) // one timeout: backoff=1. - r.PostTx(rtoDataSeg(iss, 100), int64(rtoInitial)) // retransmit (no sample). - if r.backoff != 1 { - t.Fatalf("backoff=%d, want 1 after a timeout", r.backoff) - } - // New data sent and freshly sampled, then acked. - r.PostTx(rtoDataSeg(iss+100, 100), int64(rtoInitial)+rtoMs) - r.PreRx(rtoAckSeg(iss+200), int64(rtoInitial)+30*rtoMs) - if r.backoff != 0 { - t.Errorf("backoff=%d, want 0 after a valid RTT sample", r.backoff) - } -} - -// TestRTO_Clamped verifies CurrentRTO is clamped to [rtoMin, rtoMax]. -func TestRTO_Clamped(t *testing.T) { - var r RTO - r.Reset() - r.rto = time.Nanosecond - if got := r.CurrentRTO(); got != rtoMin { - t.Errorf("CurrentRTO=%v, want floor %v", got, rtoMin) - } - r.rto = time.Hour - if got := r.CurrentRTO(); got != rtoMax { - t.Errorf("CurrentRTO=%v, want ceiling %v", got, rtoMax) - } -} - -// TestRTO_UpdateRTTFirstSample verifies the first-measurement initialization of -// SRTT/RTTVAR (RFC 6298 §2.2). -func TestRTO_UpdateRTTFirstSample(t *testing.T) { - var r RTO - r.Reset() - r.updateRTT(100 * time.Millisecond) - if r.srtt != 100*time.Millisecond { - t.Errorf("srtt=%v, want 100ms", r.srtt) - } - if r.rttvar != 50*time.Millisecond { - t.Errorf("rttvar=%v, want 50ms", r.rttvar) - } - // RTO = SRTT + K*RTTVAR = 100 + 4*50 = 300ms. - if r.rto != 300*time.Millisecond { - t.Errorf("rto=%v, want 300ms", r.rto) - } -} - -// TestRTO_ImplementsLossRecovery exercises RTO through the [LossRecovery] -// interface: sending data arms a deadline and a full ACK disarms it. -func TestRTO_ImplementsLossRecovery(t *testing.T) { - var lr LossRecovery = newRTO() - lr.Reset() - lr.PostTx(rtoDataSeg(1000, 100), 0) - if lr.NextDeadline() == 0 { - t.Error("expected an armed deadline after sending data") - } - if !lr.PreRx(rtoAckSeg(1100), 10*rtoMs).Keep { - t.Error("PreRx must keep") - } - if lr.NextDeadline() != 0 { - t.Error("expected disarmed timer after full ack") - } -} diff --git a/tcp/rtointegration_test.go b/tcp/rtointegration_test.go deleted file mode 100644 index 928fd3d..0000000 --- a/tcp/rtointegration_test.go +++ /dev/null @@ -1,120 +0,0 @@ -package tcp - -import ( - "math/rand" - "testing" - "time" - - "github.com/soypat/lneto/ethernet" -) - -// TestHandlerRetransmitsAfterRTO covers the seam between a Handler and its -// LossRecovery, which the RTO unit tests do not: a lost data segment must be -// resent once the timer expires, with nothing arriving to prompt it. -func TestHandlerRetransmitsAfterRTO(t *testing.T) { - const mtu = ethernet.MaxMTU - const maxpackets = 4 - rng := rand.New(rand.NewSource(5)) - client, server := newHandler(t, mtu, maxpackets), newHandler(t, mtu, maxpackets) - - var now int64 // injected monotonic clock, in nanoseconds - client.SetLossRecovery(new(RTO), func() int64 { return now }) - - setupClientServer(t, rng, client, server) - var rawbuf [mtu]byte - establish(t, client, server, rawbuf[:]) - - data := []byte("hello") - if n, err := client.Write(data); err != nil || n != len(data) { - t.Fatal("client write:", n, err) - } - clear(rawbuf[:]) - n, err := client.Send(rawbuf[:]) - if err != nil || n == 0 { - t.Fatal("client send:", n, err) - } - // That frame is lost: it is never handed to the server. - - // Nothing may come back before the timer expires. - var probe [mtu]byte - if n, err := client.Send(probe[:]); err != nil || n != 0 { - t.Fatalf("client sent %d bytes before the RTO expired (err %v)", n, err) - } - - now += int64(3 * time.Second) // past the initial RTO and one backoff - - clear(probe[:]) - n, err = client.Send(probe[:]) - if err != nil { - t.Fatal("client send after RTO:", err) - } - if n == 0 { - t.Fatal("no retransmission after the RTO expired: the loss-recovery directive is never applied") - } - if err := server.Recv(probe[:n]); err != nil { - t.Fatal("server refused the retransmission:", err) - } - got := make([]byte, 16) - nr, err := server.Read(got) - if err != nil || string(got[:nr]) != string(data) { - t.Fatalf("server read %q (%v), want %q", got[:nr], err, data) - } -} - -// TestHandlerRetransmitsAfterCloseWithUnackedData is the write-then-close case -// every server performs. With the last data segment lost, the FIN behind it sits -// above a gap the peer cannot cross, so FIN-WAIT-1 must still retransmit that -// data or both sides wait forever. -func TestHandlerRetransmitsAfterCloseWithUnackedData(t *testing.T) { - const mtu = ethernet.MaxMTU - const maxpackets = 4 - rng := rand.New(rand.NewSource(9)) - client, server := newHandler(t, mtu, maxpackets), newHandler(t, mtu, maxpackets) - - var now int64 - client.SetLossRecovery(new(RTO), func() int64 { return now }) - - setupClientServer(t, rng, client, server) - var rawbuf [mtu]byte - establish(t, client, server, rawbuf[:]) - - data := []byte("last response bytes") - if n, err := client.Write(data); err != nil || n != len(data) { - t.Fatal("client write:", n, err) - } - clear(rawbuf[:]) - n, err := client.Send(rawbuf[:]) // this frame is lost in transit - if err != nil || n == 0 { - t.Fatal("client send:", n, err) - } - - // The application closes right after writing. - if err := client.Close(); err != nil { - t.Fatal("client close:", err) - } - var finbuf [mtu]byte - nfin, err := client.Send(finbuf[:]) // FIN (also lost, or simply unacked) - if err != nil { - t.Fatal("client send FIN:", err) - } - t.Logf("state after close: %s (FIN frame %d bytes)", client.State(), nfin) - - now += int64(3 * time.Second) // past the RTO - - var probe [mtu]byte - n, err = client.Send(probe[:]) - if err != nil { - t.Fatal("client send after RTO:", err) - } - if n == 0 { - t.Fatalf("no retransmission in %s: unacknowledged data is stranded by the close", client.State()) - } - if err := server.Recv(probe[:n]); err != nil { - t.Fatal("server refused the retransmission:", err) - } - got := make([]byte, 32) - nr, err := server.Read(got) - if err != nil || string(got[:nr]) != string(data) { - t.Fatalf("server read %q (%v), want %q", got[:nr], err, data) - } -} diff --git a/tcp/txqueue.go b/tcp/txqueue.go index 386b911..561a100 100644 --- a/tcp/txqueue.go +++ b/tcp/txqueue.go @@ -227,18 +227,36 @@ func (rtx *ringTx) RetransmitFromUNA() { if oldest == nil { return // Nothing in the retransmission queue. } - unaSeq := oldest.seq - if rtx.sentend != 0 { - // Merge sent region [sentoff, sentend) back into unsent. - rtx.unsentoff = rtx.sentoff - if rtx.unsentend == 0 { - rtx.unsentend = rtx.sentend - } - rtx.sentoff = 0 - rtx.sentend = 0 + rtx.RetransmitFrom(oldest.seq) +} + +// RetransmitFrom rewinds the transmit queue so sent-but-unacked data from seq onward +// becomes unsent again causing next [ringTx.MakePacket] to resend them. +// +// Must be called when [ControlBlock.RetransmitFrom] returns true so the +// ring and control block state are coherent. +func (rtx *ringTx) RetransmitFrom(seq Value) { + pkt := rtx.slist.packetContaining(seq) + if pkt == nil { + return // seq not in the retransmission queue. } - // Clear packet metadata; sequence tracking restarts from UNA. - rtx.slist.Reset(cap(rtx.slist.pkts), unaSeq) + rewindOff, rewindSeq := pkt.off, pkt.seq + // The write position is unsentend, except when the unsent region is empty + // (unsentend==0) in which case data ends where the sent region ends. Capture + // it before reopening the unsent region over the rewound packets. + writeEnd := rtx.unsentend + if writeEnd == 0 { + writeEnd = rtx.sentend + } + if rewindOff == rtx.sentoff { + rtx.sentoff = 0 // Whole queue rewound: sent region becomes empty. + rtx.sentend = 0 + } else { + rtx.sentend = rewindOff + } + rtx.unsentoff = rewindOff + rtx.unsentend = writeEnd + rtx.slist.truncateFrom(rewindSeq) } func (rtx *ringTx) consolidateBufs() { @@ -331,6 +349,35 @@ func (sl *sentlist) Free() int { return cap(sl.pkts) - len(sl.pkts) } +// packetContaining returns the queued packet whose sequence range covers seq, or +// nil when no packet does. +func (sl *sentlist) packetContaining(seq Value) *ringidx { + for i := range sl.pkts { + pkt := &sl.pkts[i] + if pkt.seq.LessThanEq(seq) && seq.LessThan(pkt.endSeq()) { + return pkt + } + } + return nil +} + +// truncateFrom drops the packet starting at seq and every packet sent after it, +// so their data can be re-queued as unsent. seq must be a packet start sequence +// (see [sentlist.packetContaining]). When no packet survives, the auxiliary +// sequence counter is rewound to seq so [sentlist.EndSeq] keeps reporting where +// the next packet begins. +func (sl *sentlist) truncateFrom(seq Value) { + for i := range sl.pkts { + if sl.pkts[i].seq == seq { + sl.pkts = sl.pkts[:i] + if i == 0 { + sl.ssn = seq + } + return + } + } +} + func (sl *sentlist) AddPacket(datalen, off, bufsize int, seq Value) *ringidx { free := sl.Free() if free == 0 { diff --git a/tcp/txqueue_retransmit_test.go b/tcp/txqueue_retransmit_test.go new file mode 100644 index 0000000..f444df5 --- /dev/null +++ b/tcp/txqueue_retransmit_test.go @@ -0,0 +1,215 @@ +package tcp + +import ( + "bytes" + "testing" +) + +// newRetransmitQueue builds a queue holding npkt sent packets of pktlen octets +// each, starting at iss, plus any leftover unsent data. It returns the queue and +// the full byte stream that was written. +func newRetransmitQueue(t *testing.T, bufsize, maxPkts, npkt, pktlen, unsent int, iss Value) (*ringTx, []byte) { + t.Helper() + var rtx ringTx + if err := rtx.Reset(make([]byte, bufsize), maxPkts, iss); err != nil { + t.Fatal(err) + } + stream := make([]byte, npkt*pktlen+unsent) + for i := range stream { + stream[i] = byte(i + 1) // Non-zero so a stale ring shows up as a mismatch. + } + if n, err := rtx.Write(stream); err != nil || n != len(stream) { + t.Fatalf("write n=%d err=%v", n, err) + } + seq := iss + scratch := make([]byte, pktlen) + for i := range npkt { + n, err := rtx.MakePacket(scratch, seq) + if err != nil { + t.Fatalf("packet %d: %v", i, err) + } + if n != pktlen { + t.Fatalf("packet %d: n=%d, want %d", i, n, pktlen) + } + seq += Value(n) + } + testQueueSanity(t, &rtx) + return &rtx, stream +} + +// mustRemake asserts the queue re-emits datalen octets at seq matching want. +func mustRemake(t *testing.T, rtx *ringTx, seq Value, want []byte) { + t.Helper() + got := make([]byte, len(want)) + n, err := rtx.MakePacket(got, seq) + if err != nil { + t.Fatalf("MakePacket at seq %d: %v", seq, err) + } + if n != len(want) { + t.Fatalf("MakePacket at seq %d: n=%d, want %d", seq, n, len(want)) + } + if !bytes.Equal(got, want) { + t.Fatalf("MakePacket at seq %d: got %v, want %v", seq, got, want) + } +} + +// TestRingTx_RetransmitFromBoundary rewinds to the start of the second of three +// sent packets: the first stays sent, the rest become unsent and re-emit their +// original bytes. +func TestRingTx_RetransmitFromBoundary(t *testing.T) { + const iss, pktlen = Value(100), 4 + rtx, stream := newRetransmitQueue(t, 64, 4, 3, pktlen, 0, iss) + + sentBefore := rtx.BufferedSent() + rtx.RetransmitFrom(iss + pktlen) // Start of packet 2. + testQueueSanity(t, rtx) + + if got := rtx.BufferedSent(); got != pktlen { + t.Fatalf("sent=%d, want %d (only packet 1 remains sent)", got, pktlen) + } + if got := rtx.BufferedUnsent(); got != sentBefore-pktlen { + t.Fatalf("unsent=%d, want %d", got, sentBefore-pktlen) + } + mustRemake(t, rtx, iss+pktlen, stream[pktlen:2*pktlen]) + testQueueSanity(t, rtx) + mustRemake(t, rtx, iss+2*pktlen, stream[2*pktlen:3*pktlen]) + testQueueSanity(t, rtx) +} + +// TestRingTx_RetransmitFromMidPacket verifies a sequence inside a packet is +// snapped down to that packet's start: the queue tracks whole packets. +func TestRingTx_RetransmitFromMidPacket(t *testing.T) { + const iss, pktlen = Value(100), 4 + rtx, stream := newRetransmitQueue(t, 64, 4, 3, pktlen, 0, iss) + + rtx.RetransmitFrom(iss + pktlen + 2) // Two octets into packet 2. + testQueueSanity(t, rtx) + + if got := rtx.BufferedSent(); got != pktlen { + t.Fatalf("sent=%d, want %d: rewind must floor to the packet start", got, pktlen) + } + mustRemake(t, rtx, iss+pktlen, stream[pktlen:2*pktlen]) +} + +// TestRingTx_RetransmitFromOldest rewinds the whole queue, which must match +// RetransmitFromUNA. +func TestRingTx_RetransmitFromOldest(t *testing.T) { + const iss, pktlen, npkt = Value(100), 4, 3 + rtx, stream := newRetransmitQueue(t, 64, 4, npkt, pktlen, 0, iss) + rtx.RetransmitFrom(iss) + testQueueSanity(t, rtx) + + viaUNA, _ := newRetransmitQueue(t, 64, 4, npkt, pktlen, 0, iss) + viaUNA.RetransmitFromUNA() + testQueueSanity(t, viaUNA) + + if rtx.BufferedSent() != 0 { + t.Fatalf("sent=%d, want 0 after a full rewind", rtx.BufferedSent()) + } + if rtx.BufferedUnsent() != npkt*pktlen { + t.Fatalf("unsent=%d, want %d", rtx.BufferedUnsent(), npkt*pktlen) + } + if rtx.BufferedSent() != viaUNA.BufferedSent() || rtx.BufferedUnsent() != viaUNA.BufferedUnsent() { + t.Fatal("RetransmitFrom(oldest) must match RetransmitFromUNA") + } + mustRemake(t, rtx, iss, stream[:pktlen]) +} + +// TestRingTx_RetransmitFromUnknownSeq verifies a sequence covered by no queued +// packet leaves the queue untouched. +func TestRingTx_RetransmitFromUnknownSeq(t *testing.T) { + const iss, pktlen, npkt = Value(100), 4, 3 + rtx, _ := newRetransmitQueue(t, 64, 4, npkt, pktlen, 0, iss) + sent, unsent := rtx.BufferedSent(), rtx.BufferedUnsent() + + rtx.RetransmitFrom(iss - 1) // Before the queue. + rtx.RetransmitFrom(iss + npkt*pktlen) // One past the last octet sent. + rtx.RetransmitFrom(iss + 1000) // Far beyond. + testQueueSanity(t, rtx) + + if rtx.BufferedSent() != sent || rtx.BufferedUnsent() != unsent { + t.Fatalf("queue moved: sent %d→%d, unsent %d→%d", sent, rtx.BufferedSent(), unsent, rtx.BufferedUnsent()) + } +} + +// TestRingTx_RetransmitWithUnsentTail verifies a rewind reopens the unsent region +// over the rewound packets without losing the unsent tail behind them. +func TestRingTx_RetransmitWithUnsentTail(t *testing.T) { + const iss, pktlen, npkt, tail = Value(100), 4, 2, 5 + rtx, stream := newRetransmitQueue(t, 64, 4, npkt, pktlen, tail, iss) + if got := rtx.BufferedUnsent(); got != tail { + t.Fatalf("unsent tail=%d, want %d", got, tail) + } + + rtx.RetransmitFrom(iss + pktlen) // Rewind the second packet only. + testQueueSanity(t, rtx) + + if got := rtx.BufferedUnsent(); got != pktlen+tail { + t.Fatalf("unsent=%d, want %d (rewound packet plus the tail)", got, pktlen+tail) + } + // The rewound packet re-emits first, then the tail follows in order. + mustRemake(t, rtx, iss+pktlen, stream[pktlen:2*pktlen]) + testQueueSanity(t, rtx) + mustRemake(t, rtx, iss+2*pktlen, stream[2*pktlen:]) +} + +// TestRingTx_RetransmitAfterDrainedUnsent pins the write-position recovery: when +// every octet written has been packetized the unsent region is empty, so the +// rewind must reconstruct where data ends from the sent region. +func TestRingTx_RetransmitAfterDrainedUnsent(t *testing.T) { + const iss, pktlen, npkt = Value(100), 4, 3 + rtx, stream := newRetransmitQueue(t, 64, 4, npkt, pktlen, 0, iss) + if got := rtx.BufferedUnsent(); got != 0 { + t.Fatalf("unsent=%d, want 0: all written data was packetized", got) + } + + rtx.RetransmitFrom(iss + pktlen) + testQueueSanity(t, rtx) + + if got := rtx.BufferedUnsent(); got != 2*pktlen { + t.Fatalf("unsent=%d, want %d: rewind lost the end of the data", got, 2*pktlen) + } + mustRemake(t, rtx, iss+pktlen, stream[pktlen:2*pktlen]) + testQueueSanity(t, rtx) + mustRemake(t, rtx, iss+2*pktlen, stream[2*pktlen:3*pktlen]) +} + +// TestRingTx_RetransmitWrapped exercises a rewind on a queue whose regions wrap +// the end of the ring buffer. +func TestRingTx_RetransmitWrapped(t *testing.T) { + const bufsize, pktlen = 16, 4 + const iss = Value(100) + var rtx ringTx + if err := rtx.Reset(make([]byte, bufsize), 4, iss); err != nil { + t.Fatal(err) + } + // Push the queue most of the way around the ring, acking as we go. + seq := iss + scratch := make([]byte, pktlen) + for round := range 3 { + chunk := make([]byte, pktlen) + for i := range chunk { + chunk[i] = byte(round*pktlen + i + 1) + } + if _, err := rtx.Write(chunk); err != nil { + t.Fatal(err) + } + if _, err := rtx.MakePacket(scratch, seq); err != nil { + t.Fatal(err) + } + seq += Value(pktlen) + if round < 2 { + if err := rtx.RecvACK(seq); err != nil { + t.Fatal(err) + } + } + testQueueSanity(t, &rtx) + } + // Two packets outstanding, straddling the wrap. Rewind the newest. + rewindSeq := seq - Value(pktlen) + want := append([]byte(nil), scratch...) + rtx.RetransmitFrom(rewindSeq) + testQueueSanity(t, &rtx) + mustRemake(t, &rtx, rewindSeq, want) + testQueueSanity(t, &rtx) +}