diff --git a/tcp/txqueue.go b/tcp/txqueue.go index f4ffdd8..3f47ddb 100644 --- a/tcp/txqueue.go +++ b/tcp/txqueue.go @@ -361,35 +361,35 @@ type sentlist struct { pkts []ringidx } -func (sl sentlist) latestPkt() *ringidx { +func (sl sentlist) Newest() *ringidx { if len(sl.pkts) == 0 { return nil } return &sl.pkts[len(sl.pkts)-1] } -func (sl sentlist) oldestPkt() *ringidx { +func (sl sentlist) Oldest() *ringidx { if len(sl.pkts) == 0 { return nil } return &sl.pkts[0] } -func (sl *sentlist) endSeq() Value { +func (sl *sentlist) EndSeq() Value { seq := sl.iss - lastPkt := sl.latestPkt() + lastPkt := sl.Newest() if lastPkt != nil { seq = lastPkt.endSeq() } return seq } -func (sl *sentlist) addPkt(datalen int, bufsize int) { +func (sl *sentlist) AddPacket(datalen int, bufsize int) { free := cap(sl.pkts) - len(sl.pkts) if free == 0 { panic("pkt buffer full") } - lastPkt := sl.latestPkt() + lastPkt := sl.Newest() lastEnd := 0 if lastPkt != nil { lastEnd = lastPkt.end @@ -397,26 +397,27 @@ func (sl *sentlist) addPkt(datalen int, bufsize int) { pkt := ringidx{ off: lastEnd, end: addEnd(lastEnd, datalen, bufsize), - seq: sl.endSeq(), + seq: sl.EndSeq(), size: Size(datalen), } sl.pkts = append(sl.pkts, pkt) } -func (sl *sentlist) recvAck(ack Value, bufsize int) { +func (sl *sentlist) RecvAck(ack Value, bufsize int) { // Mark fully acked. for i := 0; i < len(sl.pkts); i++ { pkt := &sl.pkts[i] endseq := pkt.endSeq() isFullyAcked := endseq.LessThanEq(ack) if isFullyAcked { + sl.iss = endseq pkt.markRcvd() } else { break } } sl.removeRecvd() - maybePartial := sl.oldestPkt() + maybePartial := sl.Oldest() if maybePartial == nil { return // No more packets, all acked. } diff --git a/tcp/txqueue_test.go b/tcp/txqueue_test.go index 1e517bb..38993ff 100644 --- a/tcp/txqueue_test.go +++ b/tcp/txqueue_test.go @@ -8,6 +8,44 @@ import ( "testing" ) +func TestSentlist(t *testing.T) { + sl := sentlist{ + pkts: make([]ringidx, 0, 3), + } + // Test full ack. + const bufsize = 16 + const pkt = 10 + sl.AddPacket(pkt, bufsize) + if sl.Oldest() == nil || sl.Newest() != sl.Oldest() { + t.Error("expected same oldest/newest non-nil packet") + } + ack := Value(pkt) + sl.RecvAck(ack, bufsize) + oldest := sl.Oldest() + if oldest != nil { + t.Fatal("expected packet to be fully read") + } + + // Test partial ack. + sl.AddPacket(pkt, bufsize) + for i := Value(0); i < pkt-1; i++ { + ack++ + sl.RecvAck(ack, bufsize) + oldest = sl.Oldest() + if oldest == nil { + t.Fatal("partially acked packet removed") + } else if oldest.seq != ack { + t.Errorf("want pkt.seq=%d got %d", ack, oldest.seq) + } + } + ack++ + sl.RecvAck(ack, bufsize) + oldest = sl.Oldest() + if oldest != nil { + t.Fatal("expected packet to be fully read") + } +} + func TestTxQueue_multipacket(t *testing.T) { const mtu = 256 const iss = 1