rewrite txqueue sequencing logic to use per-packet sequence numbers

This commit is contained in:
soypat
2025-02-15 17:02:10 -03:00
parent 2e823c0fd5
commit 00cdf471f4
4 changed files with 181 additions and 71 deletions
+27 -4
View File
@@ -42,7 +42,7 @@ func (h *Handler) SetBuffers(txbuf, rxbuf []byte, packets int) error {
if rxbuf != nil {
h.bufRx.Buf = rxbuf
}
if len(h.bufRx.Buf) < 1 {
if len(h.bufRx.Buf) < minBufferSize {
return errors.New("short rx buffer")
}
h.scb.SetRecvWindow(Size(h.bufRx.Size()))
@@ -180,11 +180,9 @@ func (h *Handler) Send(b []byte) (int, error) {
return 0, nil
}
if available > 0 {
n, seq, err := h.bufTx.MakePacket(b[sizeHeaderTCP : sizeHeaderTCP+segment.DATALEN])
n, err := h.bufTx.MakePacket(b[sizeHeaderTCP:sizeHeaderTCP+segment.DATALEN], segment.SEQ)
if err != nil {
return 0, err
} else if seq != segment.SEQ {
panic("mismatching sequence numbers")
} else if n != int(segment.DATALEN) {
panic("expected n == available")
}
@@ -204,6 +202,31 @@ func (h *Handler) Send(b []byte) (int, error) {
return sizeHeaderTCP + int(segment.DATALEN), nil
}
func (h *Handler) Free() int {
return h.bufTx.Free()
}
func (h *Handler) Write(b []byte) (int, error) {
if h.State().IsClosed() { // Reject write call if data cannot be sent.
return 0, net.ErrClosed
}
return h.bufTx.Write(b)
}
func (h *Handler) Read(b []byte) (int, error) {
if h.State().IsClosed() { // Reject read call if state is at StateClosed. Note this is less strict than Write call condition.
return 0, net.ErrClosed
}
return h.bufRx.Read(b)
}
func (h *Handler) Buffered() int {
if h.State().IsClosed() {
return 0
}
return h.bufRx.Buffered()
}
// AwaitingSynResponse checks if the Handler is waiting for a Syn to arrive.
func (h *Handler) AwaitingSynResponse() bool {
return h.remotePort != 0 && h.scb.State() == StateSynSent
+57 -17
View File
@@ -1,6 +1,7 @@
package tcp
import (
"bytes"
"math/rand"
"testing"
)
@@ -9,19 +10,59 @@ func TestHandler(t *testing.T) {
const mtu = 1500
const maxpackets = 3
rng := rand.New(rand.NewSource(0))
client, server := setupClientServer(t, rng, mtu, mtu, maxpackets, mtu, mtu, maxpackets)
client, server := newHandler(t, mtu, maxpackets), newHandler(t, mtu, maxpackets)
setupClientServer(t, rng, client, server)
var rawbuf [mtu]byte
establish(t, client, server, rawbuf[:])
sendDataFull(t, client, server, []byte("hello"), rawbuf[:])
}
func setupClientServer(t *testing.T, rng *rand.Rand, clientTxSize, clientRxSize, clientPackets, serverTxSize, serverRxSize, serverPackets int) (client, server *Handler) {
client = new(Handler)
server = new(Handler)
err := client.SetBuffers(make([]byte, clientTxSize), make([]byte, clientRxSize), clientPackets)
func sendDataFull(t *testing.T, client, server *Handler, data, packetBuf []byte) {
n, err := client.Write(data)
if err != nil {
t.Fatal("client write:", err)
} else if n != len(data) {
t.Fatal("expected client to write full data packet")
}
n, err = client.Send(packetBuf)
if err != nil {
t.Fatal("client sending:", err)
} else if n < len(data)+sizeHeaderTCP {
t.Fatal("expected client to send full data packet", n, len(data)+sizeHeaderTCP)
}
err = server.Recv(packetBuf[:n])
if err != nil {
t.Fatal("server receiving:", err)
} else if server.Buffered() != len(data) {
t.Fatal("server did not receive full data packet", server.Buffered(), len(data))
}
clear(packetBuf)
n, err = server.Read(packetBuf)
if err != nil {
t.Fatal("server read:", err)
} else if n != len(data) {
t.Fatal("expected server to read full data packet")
} else if !bytes.Equal(packetBuf[:n], data) {
t.Fatal("server received unexpected data")
}
}
func newHandler(t *testing.T, mtu, mintaxpackets int) *Handler {
h := new(Handler)
err := h.SetBuffers(make([]byte, mtu), make([]byte, mtu), mintaxpackets)
if err != nil {
t.Fatal(err)
}
err = server.SetBuffers(make([]byte, serverTxSize), make([]byte, serverRxSize), serverPackets)
return h
}
func setupClientServer(t *testing.T, rng *rand.Rand, client, server *Handler) {
// Ensure buffer sizes are OK with reused buffers.
err := client.SetBuffers(nil, nil, 0)
if err != nil {
t.Fatal(err)
}
err = server.SetBuffers(nil, nil, 0)
if err != nil {
t.Fatal(err)
}
@@ -39,21 +80,20 @@ func setupClientServer(t *testing.T, rng *rand.Rand, clientTxSize, clientRxSize,
if !server.AwaitingSynAck() {
t.Fatal("server in wrong state")
}
return client, server
}
func establish(t *testing.T, client, server *Handler, buf []byte) {
func establish(t *testing.T, client, server *Handler, packetBuf []byte) {
if client.State() != StateClosed {
t.Fatal("client in wrong state")
} else if server.State() != StateListen {
t.Fatal("server in wrong state")
}
clear(buf)
clear(packetBuf)
// Commence 3-way handshake: client sends SYN, server sends SYN-ACK, client sends ACK.
// Client sends SYN.
n, err := client.Send(buf)
n, err := client.Send(packetBuf)
if err != nil {
t.Fatal("client sending:", err)
} else if n < sizeHeaderTCP {
@@ -61,15 +101,15 @@ func establish(t *testing.T, client, server *Handler, buf []byte) {
} else if client.State() != StateSynSent {
t.Fatal("client did not transition to SynSent state:", client.State().String())
}
err = server.Recv(buf[:n]) // Server receives SYN.
err = server.Recv(packetBuf[:n]) // Server receives SYN.
if err != nil {
t.Fatal(err)
} else if server.State() != StateSynRcvd {
t.Fatal("server did not transition to SynReceived state:", server.State().String())
}
clear(buf)
clear(packetBuf)
// Server sends SYNACK response to client's SYN.
n, err = server.Send(buf)
n, err = server.Send(packetBuf)
if err != nil {
t.Fatal("server sending:", err)
} else if n < sizeHeaderTCP {
@@ -77,15 +117,15 @@ func establish(t *testing.T, client, server *Handler, buf []byte) {
} else if server.State() != StateSynRcvd {
t.Fatal("server should remain in SynReceived state:", server.State().String())
}
err = client.Recv(buf[:n]) // Client receives SYNACK, is established but must send ACK.
err = client.Recv(packetBuf[:n]) // Client receives SYNACK, is established but must send ACK.
if err != nil {
t.Fatal(err)
} else if client.State() != StateEstablished {
t.Fatal("client did not transition to Established state:", client.State().String())
}
clear(buf)
n, err = client.Send(buf) // Client sends ACK.
clear(packetBuf)
n, err = client.Send(packetBuf) // Client sends ACK.
if err != nil {
t.Fatal("client sending ACK:", err)
} else if n < sizeHeaderTCP {
@@ -93,7 +133,7 @@ func establish(t *testing.T, client, server *Handler, buf []byte) {
} else if client.State() != StateEstablished {
t.Fatal("client should remain in Established state:", client.State().String())
}
err = server.Recv(buf[:n]) // Server receives ACK.
err = server.Recv(packetBuf[:n]) // Server receives ACK.
if err != nil {
t.Fatal(err)
} else if server.State() != StateEstablished {
+47 -24
View File
@@ -28,7 +28,7 @@ type ringTx struct {
sentoff int
// sentend is the offset of end of sent data in rawbuf. If zero then sent buffer is empty.
sentend int
seq Value
// seq Value
// always empty ring.
emptyRing ringidx
}
@@ -39,8 +39,10 @@ type ringidx struct {
off int
// end is the ringed data end offset, non-inclusive. Follows [internal.Ring] semantics.
end int
// seq is the sequence number of the packet.
// seq is the sequence number of the first byte in the packet.
seq Value
// size is the size of the packet in bytes.
size Size
// time is a measure of the instant of time message was sent at.
}
@@ -58,7 +60,6 @@ func (rtx *ringTx) Reset(buf []byte, maxqueuedPackets int, seq Value) error {
*rtx = ringTx{
rawbuf: buf,
packets: rtx.packets[:maxqueuedPackets],
seq: seq,
}
for i := range rtx.packets {
rtx.packets[i].markRcvd()
@@ -112,16 +113,16 @@ func (rtx *ringTx) Write(b []byte) (n int, err error) {
// MakePacket reads from the unsent data ring buffer and generates a new packet segment.
// It fails if the sent packet queue is full.
func (rtx *ringTx) MakePacket(b []byte) (int, Value, error) {
func (rtx *ringTx) MakePacket(b []byte, currentSeq Value) (int, error) {
nxtpkt := rtx.nextPkt()
if rtx.nextPkt() < 0 {
return 0, 0, errors.New("queue full")
if nxtpkt < 0 {
return 0, errors.New("queue full")
}
r, _ := rtx.unsentRing()
start := r.Off
n, err := r.Read(b)
if err != nil {
return n, 0, err
return n, err
}
pkt := &rtx.packets[nxtpkt]
@@ -131,33 +132,34 @@ func (rtx *ringTx) MakePacket(b []byte) (int, Value, error) {
if off == rtx.unsentend {
rtx.unsentend = 0 // Mark unsent as being empty.
}
pkt.off = start
pkt.end = off
// Sequence number updates.
oldseq := rtx.seq
newseq := Add(oldseq, Size(n))
rtx.seq = newseq
pkt.seq = newseq
return n, oldseq, nil
*pkt = ringidx{
off: start,
end: off,
seq: currentSeq,
size: Size(n),
}
return n, nil
}
// RecvSegment processes an incoming segment and updates the sent packet queue
func (rtx *ringTx) RecvACK(ack Value) error {
if ack.LessThan(rtx.seq) {
return errors.New("old packet")
}
first := rtx.firstPkt()
if first < 0 {
return errors.New("no packets to ack")
}
hiSeq := rtx.pkt(first).seq
pkt0 := rtx.pkt(first)
if ack.LessThanEq(pkt0.seq) {
return errors.New("old packet")
}
hiSeq := Add(pkt0.seq, pkt0.size)
for i := 0; i < len(rtx.packets); i++ {
pkt := &rtx.packets[i]
if pkt.sent() && pkt.seq.LessThanEq(ack) {
if hiSeq.LessThanEq(pkt.seq) {
pktendSeq := Add(pkt.seq, pkt.size)
if pkt.sent() && pktendSeq.LessThanEq(ack) {
if hiSeq.LessThanEq(pktendSeq) {
rtx.sentoff = pkt.end
hiSeq = pkt.seq
hiSeq = pktendSeq
}
pkt.markRcvd()
}
@@ -266,7 +268,28 @@ func (rtx *ringTx) consolidateBufs() {
}
}
func (rtx *ringTx) currentSeq() Value { return rtx.seq }
func (rtx *ringTx) endSeq() (Value, bool) {
pkt := rtx.lastPkt()
if pkt < 0 {
return 0, false
}
last := rtx.pkt(pkt)
return Add(last.seq, last.size), true
}
func (rtx *ringTx) lastSeq() (Value, bool) {
pkt := rtx.lastPkt()
if pkt < 0 {
return 0, false
}
return rtx.pkt(pkt).seq, true
}
func (rtx *ringTx) firstSeq() (Value, bool) {
pkt := rtx.firstPkt()
if pkt < 0 {
return 0, false
}
return rtx.pkt(pkt).seq, true
}
// lims returns the limits of free|sent|unsent buffers.
// Example:
+50 -26
View File
@@ -31,20 +31,24 @@ func TestTxQueue(t *testing.T) {
rtx.Reset(ringBuf[:], rng.Intn(4)+1, startAck)
for imsg, msg := range msgs {
// Write and create packet from single messages.
seq := currentAck
currentAck = Add(currentAck, Size(len(msg)))
operateOnRing(t, &rtx, msg, readBuf[:], aux[:], &currentAck)
operateOnRing(t, &rtx, msg, readBuf[:], aux[:], seq, &currentAck)
buffered := rtx.Buffered()
if buffered != 0 {
t.Fatalf("msg%d: want no buffered data after transaction, got %d", imsg, buffered)
}
newSeq := rtx.currentSeq()
wantSeq := currentAck
if newSeq != wantSeq {
t.Fatalf("msg%d: want seq %d, got %d", imsg, wantSeq, newSeq)
}
if t.Failed() {
t.Fatalf("failed on msg %d", imsg)
}
// newSeq, ok := rtx.firstSeq()
// if !ok {
// t.Fatal("no first packet found")
// }
// wantSeq := currentAck
// if newSeq != wantSeq {
// t.Fatalf("msg%d: want seq %d, got %d", imsg, wantSeq, newSeq)
// }
// if t.Failed() {
// t.Fatalf("failed on msg %d", imsg)
// }
}
}
},
@@ -58,14 +62,17 @@ func TestTxQueue(t *testing.T) {
msgs := removeEmptyMsgs(bytes.SplitAfter(msgBuf[:], []byte{0}))
currentAck := Value(startAck)
rtx.Reset(ringBuf[:], rng.Intn(4)+1, startAck)
expectBuffered := 0
for _, msg := range msgs {
// Send all messages.
operateOnRing(t, &rtx, msg, nil, aux[:], nil)
seq := currentAck
operateOnRing(t, &rtx, msg, nil, aux[:], seq, nil)
if t.Failed() {
return
}
gotSeq := rtx.currentSeq()
if gotSeq != startAck {
expectBuffered += len(msg)
buffered := rtx.Buffered()
if buffered != expectBuffered {
t.Fatalf("expected seq to not change during writes")
}
currentAck = Add(currentAck, Size(len(msg)))
@@ -78,7 +85,7 @@ func TestTxQueue(t *testing.T) {
} else if sent != 0 {
t.Fatalf("want no data sent, got %d", sent)
}
operateOnRing(t, &rtx, nil, readBuf[:], aux[:], &currentAck)
operateOnRing(t, &rtx, nil, readBuf[:], aux[:], 0, &currentAck)
unsent = rtx.Buffered()
if unsent != 0 {
t.Fatalf("expected all data to be sent after ack of most recent packet, %d", unsent)
@@ -231,7 +238,7 @@ func removeEmptyMsgs(msgs [][]byte) [][]byte {
return slices.DeleteFunc(msgs, func(b []byte) bool { return len(b) == 0 })
}
func operateOnRing(t *testing.T, rtx *ringTx, write, readPacket, aux []byte, argRecvAck *Value) {
func operateOnRing(t *testing.T, rtx *ringTx, write, readPacket, aux []byte, newPacketSeq Value, argRecvAck *Value) {
if len(aux) < rtx.Size() {
panic("too small auxiliary buffer")
}
@@ -239,7 +246,7 @@ func operateOnRing(t *testing.T, rtx *ringTx, write, readPacket, aux []byte, arg
// Prepare aux with data expected from read after write.
runsent, _ := rtx.unsentRing()
unsent := runsent.Buffered()
startSeq, startSeqOK := rtx.firstSeq()
wantWritten := min(free, len(write))
wantBufRead := aux[:min(unsent+wantWritten, len(readPacket))]
@@ -259,7 +266,6 @@ func operateOnRing(t *testing.T, rtx *ringTx, write, readPacket, aux []byte, arg
copy(wantBufRead[n:], write)
}
prevSeq := rtx.currentSeq()
if len(write) != 0 {
testQueueSanity(t, rtx)
preBuffered := rtx.Buffered()
@@ -284,15 +290,18 @@ func operateOnRing(t *testing.T, rtx *ringTx, write, readPacket, aux []byte, arg
if wantRead != len(wantBufRead) {
t.Fatalf("miscalculated expect read %d != %d", wantRead, len(wantBufRead))
}
n, seq, err := rtx.MakePacket(readPacket)
n, err := rtx.MakePacket(readPacket, newPacketSeq)
if err != nil && wantRead != 0 {
t.Errorf("error reading: %s", err)
} else if n != wantRead {
t.Errorf("want read %d, got %d", wantRead, n)
}
wantSeq := prevSeq
if seq != wantSeq {
t.Errorf("want new seq %d, got %d", wantSeq, seq)
lastSeq, lastSeqOK := rtx.lastSeq()
endSeq, endSeqOK := rtx.endSeq()
if !lastSeqOK || lastSeq != newPacketSeq {
t.Fatalf("expected last seq to be %d, got %d (or lastSeqOK=%v)", newPacketSeq, lastSeq, lastSeqOK)
} else if !endSeqOK || endSeq != Add(newPacketSeq, Size(n)) {
t.Fatalf("expected end seq to be %d, got %d (or endSeqOK=%v)", Add(newPacketSeq, Size(n)), endSeq, endSeqOK)
}
if !bytes.Equal(readPacket[:n], wantBufRead) {
t.Error("data content packet read not match wanted packet")
@@ -303,22 +312,37 @@ func operateOnRing(t *testing.T, rtx *ringTx, write, readPacket, aux []byte, arg
}
}
startSeq2, sseqOK := rtx.firstSeq()
if sseqOK == startSeqOK && startSeq2 != startSeq {
t.Fatalf("expected FIRST seq to not change during writes")
}
if !t.Failed() && argRecvAck != nil {
testQueueSanity(t, rtx)
preAcked := rtx.BufferedSent()
// preAcked := rtx.BufferedSent()
rcvAck := *argRecvAck
seq := rtx.currentSeq()
seq, ok := rtx.firstSeq()
if !ok {
t.Fatal("no first packet found")
}
startSeq := Add(seq, Size(-rtx.BufferedSent()))
acklInSentRange := startSeq.LessThan(rcvAck) && rcvAck.LessThanEq(seq)
err := rtx.RecvACK(rcvAck)
if err != nil && acklInSentRange {
t.Errorf("expected correct acking %d < %d <= %d: %s", startSeq, rcvAck, seq, err)
}
gotCalcAcked := preAcked - rtx.BufferedSent()
wantAcked := int(Sizeof(prevSeq, rtx.currentSeq()))
if gotCalcAcked != wantAcked {
t.Errorf("want acked %d, got %d", wantAcked, gotCalcAcked)
bufSent := rtx.BufferedSent()
gotFirstSeq, ok := rtx.firstSeq()
if !ok && bufSent != 0 {
t.Fatalf("no first packet found after acking")
}
if ok && gotFirstSeq.LessThanEq(rcvAck) {
t.Fatalf("expected first seq %d to be greater than ack %d", gotFirstSeq, rcvAck)
}
// wantAcked := int(Sizeof(prevSeq, gotFirstSeq))
// if gotCalcAcked != wantAcked {
// t.Errorf("want acked %d, got %d", wantAcked, gotCalcAcked)
// }
}
testQueueSanity(t, rtx)
}