diff --git a/conn/bind_std.go b/conn/bind_std.go index 39363448a..20ee5c17b 100644 --- a/conn/bind_std.go +++ b/conn/bind_std.go @@ -208,7 +208,8 @@ func (s *StdNetBind) putMessages(msgs *[]ipv6.Message) { // Non coalesced write paths access only batch.msgs[i].Buffers[0], // but we append more during [coalesceMessages]. // Leave index zero accessible: - (*msgs)[i] = ipv6.Message{Buffers: (*msgs)[i].Buffers[:1], OOB: (*msgs)[i].OOB} + (*msgs)[i] = ipv6.Message{Buffers: (*msgs)[i].Buffers[:1], OOB: (*msgs)[i].OOB[:cap((*msgs)[i].OOB)]} + (*msgs)[i].Buffers[0] = nil } s.msgsPool.Put(msgs) } @@ -236,53 +237,79 @@ func (s *StdNetBind) receiveIP( rxOffload bool, bufs []iobuf.View, eps []Endpoint, -) (n int, err error) { +) (int, error) { msgs := s.getMessages() - // TODO: placeholder until bind implements right-sized buffers. - iobuf.EnsureAllocated(bufs) - for i := range bufs { - (*msgs)[i].Buffers[0] = bufs[i].Bytes - (*msgs)[i].OOB = (*msgs)[i].OOB[:cap((*msgs)[i].OOB)] - } defer s.putMessages(msgs) - var numMsgs int - if runtime.GOOS == "linux" { + var readDataArr [IdealBatchSize]*iobuf.Shared // on the stack + readData := readDataArr[:] + var readDataN int // tracks read buffers to release on return + defer func() { + for i := range readData[:readDataN] { + if readData[i] != nil { + readData[i].Release() + } + } + }() + + switch runtime.GOOS { + case "linux": + readBatchSize := min(len(bufs), len(*msgs)) if rxOffload { - readAt := len(*msgs) - 2 - numMsgs, err = br.ReadBatch((*msgs)[readAt:], 0) - if err != nil { - return 0, err + readBatchSize = min(max(1, len(bufs)/udpSegmentMaxDatagrams), len(*msgs)) + } + readDataN = readBatchSize + for i := range readBatchSize { + readData[i] = iobuf.SharedBufPool.Get() + (*msgs)[i].Buffers[0] = readData[i].Bytes[:] + } + msgsN, err := br.ReadBatch((*msgs)[:readBatchSize], 0) + if err != nil { + return 0, err // expect atomic reads + } + var n int + for i, msg := range (*msgs)[:msgsN] { + if msg.N == 0 { + continue } - numMsgs, err = splitCoalescedMessages(*msgs, readAt, getGSOSize) - if err != nil { - return 0, err + gsoSize := msg.N // Non-offload path splits to one read-sized View. + if rxOffload { + gsoSize, err = getGSOSize(msg.OOB[:msg.NN]) + if err != nil { + iobuf.ReleaseAll(bufs[:n]) + return 0, err + } + if gsoSize == 0 { + gsoSize = msg.N + } } - } else { - numMsgs, err = br.ReadBatch(*msgs, 0) + split, err := readData[i].SplitCoalesced(bufs[n:], gsoSize, msg.N) if err != nil { - return 0, err + iobuf.ReleaseAll(bufs[i:n]) + return i - 1, err + } + addrPort := msg.Addr.(*net.UDPAddr).AddrPort() + ep := &StdNetEndpoint{AddrPort: addrPort} // TODO: remove allocation + getSrcFromControl(msg.OOB[:msg.NN], ep) + for j := n; j < n+split; j++ { + eps[j] = ep } + n += split } - } else { + return n, nil + default: msg := &(*msgs)[0] - msg.N, msg.NN, _, msg.Addr, err = conn.ReadMsgUDP(msg.Buffers[0], msg.OOB) + readDataN = 1 + readData[0] = iobuf.SharedBufPool.Get() + n, nn, _, addr, err := conn.ReadMsgUDP(readData[0].Bytes[:], msg.OOB) if err != nil { return 0, err } - numMsgs = 1 - } - for i := 0; i < numMsgs; i++ { - msg := &(*msgs)[i] - bufs[i].Bytes = bufs[i].Bytes[:msg.N] - if len(bufs[i].Bytes) == 0 { - continue - } - addrPort := msg.Addr.(*net.UDPAddr).AddrPort() - ep := &StdNetEndpoint{AddrPort: addrPort} // TODO: remove allocation - getSrcFromControl(msg.OOB[:msg.NN], ep) - eps[i] = ep + readData[0].Refer(&bufs[0], 0, n) + ep := &StdNetEndpoint{AddrPort: addr.AddrPort()} + getSrcFromControl(msg.OOB[:nn], ep) + eps[0] = ep + return 1, nil } - return numMsgs, nil } func (s *StdNetBind) makeReceiveIPv4(pc *ipv4.PacketConn, conn *net.UDPConn, rxOffload bool) ReceiveFunc { @@ -523,49 +550,3 @@ func coalesceMessages(addr *net.UDPAddr, ep *StdNetEndpoint, bufs [][]byte, offs } return base + 1 } - -type getGSOFunc func(control []byte) (int, error) - -func splitCoalescedMessages(msgs []ipv6.Message, firstMsgAt int, getGSO getGSOFunc) (n int, err error) { - for i := firstMsgAt; i < len(msgs); i++ { - msg := &msgs[i] - if msg.N == 0 { - return n, err - } - var ( - gsoSize int - start int - end = msg.N - numToSplit = 1 - ) - gsoSize, err = getGSO(msg.OOB[:msg.NN]) - if err != nil { - return n, err - } - if gsoSize > 0 { - numToSplit = (msg.N + gsoSize - 1) / gsoSize - end = gsoSize - } - for j := 0; j < numToSplit; j++ { - if n > i { - return n, errors.New("splitting coalesced packet resulted in overflow") - } - copied := copy(msgs[n].Buffers[0], msg.Buffers[0][start:end]) - msgs[n].N = copied - msgs[n].Addr = msg.Addr - start = end - end += gsoSize - if end > msg.N { - end = msg.N - } - n++ - } - if i != n-1 { - // It is legal for bytes to move within msg.Buffers[0] as a result - // of splitting, so we only zero the source msg len when it is not - // the destination of the last split operation above. - msg.N = 0 - } - } - return n, nil -} diff --git a/conn/bind_std_test.go b/conn/bind_std_test.go index 88dcc6c2f..194257d5b 100644 --- a/conn/bind_std_test.go +++ b/conn/bind_std_test.go @@ -135,123 +135,3 @@ func mockGetGSOSize(control []byte) (int, error) { } return int(binary.LittleEndian.Uint16(control)), nil } - -func Test_splitCoalescedMessages(t *testing.T) { - newMsg := func(n, gso int) ipv6.Message { - msg := ipv6.Message{ - Buffers: [][]byte{make([]byte, 1<<16-1)}, - N: n, - OOB: make([]byte, 2), - } - binary.LittleEndian.PutUint16(msg.OOB, uint16(gso)) - if gso > 0 { - msg.NN = 2 - } - return msg - } - - cases := []struct { - name string - msgs []ipv6.Message - firstMsgAt int - wantNumEval int - wantMsgLens []int - wantErr bool - }{ - { - name: "second last split last empty", - msgs: []ipv6.Message{ - newMsg(0, 0), - newMsg(0, 0), - newMsg(3, 1), - newMsg(0, 0), - }, - firstMsgAt: 2, - wantNumEval: 3, - wantMsgLens: []int{1, 1, 1, 0}, - wantErr: false, - }, - { - name: "second last no split last empty", - msgs: []ipv6.Message{ - newMsg(0, 0), - newMsg(0, 0), - newMsg(1, 0), - newMsg(0, 0), - }, - firstMsgAt: 2, - wantNumEval: 1, - wantMsgLens: []int{1, 0, 0, 0}, - wantErr: false, - }, - { - name: "second last no split last no split", - msgs: []ipv6.Message{ - newMsg(0, 0), - newMsg(0, 0), - newMsg(1, 0), - newMsg(1, 0), - }, - firstMsgAt: 2, - wantNumEval: 2, - wantMsgLens: []int{1, 1, 0, 0}, - wantErr: false, - }, - { - name: "second last no split last split", - msgs: []ipv6.Message{ - newMsg(0, 0), - newMsg(0, 0), - newMsg(1, 0), - newMsg(3, 1), - }, - firstMsgAt: 2, - wantNumEval: 4, - wantMsgLens: []int{1, 1, 1, 1}, - wantErr: false, - }, - { - name: "second last split last split", - msgs: []ipv6.Message{ - newMsg(0, 0), - newMsg(0, 0), - newMsg(2, 1), - newMsg(2, 1), - }, - firstMsgAt: 2, - wantNumEval: 4, - wantMsgLens: []int{1, 1, 1, 1}, - wantErr: false, - }, - { - name: "second last no split last split overflow", - msgs: []ipv6.Message{ - newMsg(0, 0), - newMsg(0, 0), - newMsg(1, 0), - newMsg(4, 1), - }, - firstMsgAt: 2, - wantNumEval: 4, - wantMsgLens: []int{1, 1, 1, 1}, - wantErr: true, - }, - } - - for _, tt := range cases { - t.Run(tt.name, func(t *testing.T) { - got, err := splitCoalescedMessages(tt.msgs, 2, mockGetGSOSize) - if err != nil && !tt.wantErr { - t.Fatalf("err: %v", err) - } - if got != tt.wantNumEval { - t.Fatalf("got to eval: %d want: %d", got, tt.wantNumEval) - } - for i, msg := range tt.msgs { - if msg.N != tt.wantMsgLens[i] { - t.Fatalf("msg[%d].N: %d want: %d", i, msg.N, tt.wantMsgLens[i]) - } - } - }) - } -} diff --git a/iobuf/raw.go b/iobuf/raw.go index 91a92399c..c61a84d56 100644 --- a/iobuf/raw.go +++ b/iobuf/raw.go @@ -46,7 +46,7 @@ var DefaultRawPool = NewRawPool(MaxPooledBuffers) // and therefore whether queue finalizers are needed to drain blocked // producers when an autodraining queue is GC'd. func HasAccounting() bool { - return DefaultRawPool.WaitPool.HasAccounting() + return DefaultRawPool.HasAccounting() } // EnsureAllocated fills zero-valued Views from the [DefaultRawPool]. diff --git a/iobuf/shared.go b/iobuf/shared.go new file mode 100644 index 000000000..f7f3d96eb --- /dev/null +++ b/iobuf/shared.go @@ -0,0 +1,109 @@ +/* SPDX-License-Identifier: MIT + * + * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + */ + +package iobuf + +import ( + "errors" + "sync/atomic" + "unsafe" + + "github.com/tailscale/wireguard-go/waitpool" +) + +var _ Recycler = (*Shared)(nil) + +// Shared is a backing array that can be shared by multiple Views. +// It implements [Recycler], and will return itself to the [Pool] when all +// Views referring to it are released. It counts self as reference, and must be +// released as well. +type Shared struct { + Bytes Raw + BytesRecycler Recycler + Refs atomic.Int32 +} + +// Refer sets b to refer to a slice of the backing array. +func (s *Shared) Refer(b *View, start, end int) { + b.Recycler = s + b.BackingGo = unsafe.Pointer(s) + s.Refs.Add(1) + b.Bytes = s.Bytes[start:end:end] +} + +// Recycle is called by a View's Release when the View is no longer in use. +func (s *Shared) Recycle(goPtr unsafe.Pointer, _ uintptr) { + if s.BytesRecycler != nil && s.Refs.Add(-1) == 0 { + s.BytesRecycler.Recycle(goPtr, 0) // return to pool + } +} + +// Release this instance for reuse. Should be called exactly once, and will +// keep its Views alive until they are released as well. +func (s *Shared) Release() { + s.Recycle(unsafe.Pointer(s), 0) +} + +var ( + ErrStrideOutOfRange = errors.New("buffer: stride must be > 0 and <= readLen") + ErrReadLenOverflow = errors.New("buffer: readLen exceeds buffer capacity") + ErrInsufficientBuffers = errors.New("buffer: insufficient buffers") +) + +// SplitCoalesced fills vs with non-overlapping slices of the backing array, +// where stride is the desired slice length, and readLen is the total length of data to split. +// The last slice may be shorter than stride if readLen is not a multiple of stride. +func (s *Shared) SplitCoalesced(vs []View, stride, readLen int) (n int, err error) { + if readLen > len(s.Bytes) { + return 0, ErrReadLenOverflow + } + if stride <= 0 || stride > readLen { + return 0, ErrStrideOutOfRange + } + numToSplit := (readLen + stride - 1) / stride + if numToSplit > len(vs) { + return 0, ErrInsufficientBuffers + } + start, end := 0, stride + for i := range numToSplit { + s.Refer(&vs[i], start, end) + start = end + end += stride + if end > readLen { + end = readLen + } + } + return numToSplit, nil +} + +// SharedBufPool is used for package-level [Get] and [EnsureAllocated]. +var SharedBufPool = NewSharedBufferPool(MaxPooledBuffers) + +// SharedBufferPool is a capped pool of backing arrays. +type SharedBufferPool struct { + *waitpool.WaitPool +} + +func NewSharedBufferPool(limit int) *SharedBufferPool { + pool := &SharedBufferPool{} + pool.WaitPool = waitpool.New(limit, func() any { + p := new(Shared) + p.BytesRecycler = pool + return p + }) + return pool +} + +func (p *SharedBufferPool) Get() *Shared { + arr := p.WaitPool.Get().(*Shared) + arr.Refs.Store(1) // *Shared must be Released as well. + return arr +} + +// Recycle returns the buffer to the pool. +func (p *SharedBufferPool) Recycle(goPtr unsafe.Pointer, _ uintptr) { + arr := (*Shared)(goPtr) + p.Put(arr) +} diff --git a/iobuf/shared_test.go b/iobuf/shared_test.go new file mode 100644 index 000000000..dab0c222f --- /dev/null +++ b/iobuf/shared_test.go @@ -0,0 +1,163 @@ +/* SPDX-License-Identifier: MIT + * + * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved. + */ + +package iobuf + +import ( + "bytes" + "testing" +) + +func TestSharedReferRecycle(t *testing.T) { + pool := NewSharedBufferPool(0) + s := pool.Get() + if got := s.Refs.Load(); got != 1 { + t.Fatalf("Get: Refs = %d, want 1 (pool pre-increment)", got) + } + var v1, v2 View + s.Refer(&v1, 0, 4) + s.Refer(&v2, 4, 8) + if got := s.Refs.Load(); got != 3 { + t.Fatalf("after two Refer: Refs = %d, want 3", got) + } + v1.Release() + if got := s.Refs.Load(); got != 2 { + t.Fatalf("after v1.Release: Refs = %d, want 2", got) + } + v2.Release() + if got := s.Refs.Load(); got != 1 { + t.Fatalf("after v2.Release: Refs = %d, want 1", got) + } + // Only the matching Shared.Release returns to the pool. + s.Bytes[0] = 0xAB + s.Release() + got := pool.Get() + if got != s { + t.Fatal("pool did not return the same *Shared after full Release") + } + if got.Bytes[0] != 0xAB { + t.Fatal("pool did not preserve backing array contents") + } + if r := got.Refs.Load(); r != 1 { + t.Fatalf("pool.Get Refs = %d, want 1", r) + } +} + +func TestSharedNilRecyclerRelease(t *testing.T) { + // An unmanaged Shared (nil BytesRecycler) is GC-managed; draining its refs + // to zero must not panic — Recycle no-ops instead of returning to a pool. + s := &Shared{} + s.Refs.Store(1) + var v View + s.Refer(&v, 0, 4) + v.Release() + s.Release() // drives Refs to 0; must not deref the nil recycler + if got := s.Refs.Load(); got != 0 { + t.Fatalf("after full release: Refs = %d, want 0", got) + } +} + +func BenchmarkSplitCoalesced(b *testing.B) { + s := &Shared{} + for i := range s.Bytes { + s.Bytes[i] = byte(i) + } + vs := make([]View, 64) + for b.Loop() { + _, err := s.SplitCoalesced(vs, 1, 64) + if err != nil { + b.Fatal(err) + } + } + +} + +func TestSplitCoalesced(t *testing.T) { + const numViews = 10 + tests := []struct { + name string + stride int + readLen int + wantN int + wantErr error + }{ + { + name: "non-divisible stride", + stride: 3, + readLen: 16, + wantN: 6, + }, + { + name: "exact stride", + stride: 2, + readLen: 16, + wantN: 8, + }, + { + name: "single segment", + stride: 16, + readLen: 16, + wantN: 1, + }, + { + name: "stride zero", + stride: 0, + readLen: 16, + wantErr: ErrStrideOutOfRange, + }, + { + name: "stride negative", + stride: -1, + readLen: 16, + wantErr: ErrStrideOutOfRange, + }, + { + name: "stride exceeds readLen", + stride: 17, + readLen: 16, + wantErr: ErrStrideOutOfRange, + }, + { + name: "insufficient buffers", + stride: 1, + readLen: 16, + wantErr: ErrInsufficientBuffers, + }, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + s := &Shared{} + for i := 0; i < tc.readLen && i < len(s.Bytes); i++ { + s.Bytes[i] = byte(i) + } + vs := make([]View, numViews) + + n, err := s.SplitCoalesced(vs, tc.stride, tc.readLen) + if err != tc.wantErr { + t.Fatalf("expected error %v, got %v", tc.wantErr, err) + } + if n != tc.wantN { + t.Fatalf("expected %d segments, got %d", tc.wantN, n) + } + + srcOff := 0 + for i, v := range vs[:n] { + got := v.Bytes + wantLen := tc.stride + if i == n-1 { + wantLen = tc.readLen - (n-1)*tc.stride + } + if len(got) != wantLen { + t.Fatalf("segment %d: expected len %d, got %d", i, wantLen, len(got)) + } + want := s.Bytes[srcOff : srcOff+wantLen] + if !bytes.Equal(got, want) { + t.Fatalf("segment %d: data mismatch: got %v, want %v", i, got, want) + } + srcOff += wantLen + } + }) + } +}