net/tstun: implement tun.MultiQueueDevice (#21665)

Wrapper presents one read queue per queue of the underlying device, and
forwards the flow it is given on the write side.

pollVector still multiplexes injected packets onto the device's own
read, which is queue 0. Wrapper.Read is now queues[0].Read and behaves
exactly as it did. Locally-terminated traffic stays on flow 0.

This is a no-op change: the Wrapper presents a single queue and
wireguard-go runs the single reader goroutine it ran before.

Updates tailscale/corp#37878

Change-Id: I844e95cf430075e237ad7e83bdbaa8e43e067ccb
Signed-off-by: Alex Valiushko <alexvaliushko@tailscale.com>
This commit is contained in:
Alex Valiushko authored and GitHub committed 2026-10-07 12:09:16 -07:00
1 parent 2249be1ac8
commit 1eb2bb62f3
3 files changed
+232 -21

No files matched your search

+34 -5
View File
@@ -13,18 +13,42 @@ import (
type fakeTUN struct {
evchan chan tun.Event
closechan chan struct{}
queues []tun.Queue
}
// NewFake returns a tun.Device that does nothing.
func NewFake() tun.Device {
return &fakeTUN{
func NewFake() tun.Device { return newFakeMQ(1) }
// newFakeMQ returns a fake device with n read queues.
func newFakeMQ(n int) *fakeTUN {
if n < 1 {
panic("newFakeMQ() with less than one queue requested")
}
t := &fakeTUN{
evchan: make(chan tun.Event),
closechan: make(chan struct{}),
}
for range n {
t.queues = append(t.queues, &fakeQueue{closechan: t.closechan})
}
return t
}
type fakeQueue struct {
closechan chan struct{}
}
func (q *fakeQueue) File() *os.File {
panic("fakeTUN.File() called, which makes no sense")
}
func (q *fakeQueue) Read(slab []byte, packets []tun.ReadPacket) (int, error) {
<-q.closechan
return 0, io.EOF
}
func (t *fakeTUN) File() *os.File {
panic("fakeTUN.File() called, which makes no sense")
return t.queues[0].File()
}
func (t *fakeTUN) Close() error {
@@ -34,8 +58,13 @@ func (t *fakeTUN) Close() error {
}
func (t *fakeTUN) Read(slab []byte, packets []tun.ReadPacket) (int, error) {
<-t.closechan
return 0, io.EOF
return t.queues[0].Read(slab, packets)
}
func (t *fakeTUN) Queues() []tun.Queue { return t.queues }
func (t *fakeTUN) WriteTo(_ int, b [][]byte, n int) (int, error) {
return t.Write(b, n)
}
func (t *fakeTUN) Write(b [][]byte, n int) (int, error) {
+124 -16
View File
@@ -81,6 +81,8 @@ var parsedPacketPool = sync.Pool{New: func() any { return new(packet.Parsed) }}
// FilterFunc is a packet-filtering function with access to the Wrapper device.
// It must not hold onto the packet struct, as its backing storage will be reused.
//
// Must be safe to call concurrently.
type FilterFunc func(*packet.Parsed, *Wrapper) filter.Response
// GROFilterFunc is a FilterFunc extended with a *gro.GRO, enabling increased
@@ -91,6 +93,8 @@ type FilterFunc func(*packet.Parsed, *Wrapper) filter.Response
// *gro.GRO is non-nil after the last packet for a given vector is passed
// through the GROFilterFunc, the caller must also call Flush() on it to deliver
// any previously Enqueue()'d packets.
//
// Must be safe to call concurrently.
type GROFilterFunc func(p *packet.Parsed, w *Wrapper, g *gro.GRO) (filter.Response, *gro.GRO)
// Wrapper augments a tun.Device with packet filtering and injection.
@@ -122,8 +126,13 @@ type Wrapper struct {
// peerConfig stores the current NAT configuration.
peerConfig atomic.Pointer[peerConfigTable]
// queues are the Wrapper's read queues, see [tun.QueuesOf].
queues []tun.Queue
// writeTo is tdev's queue-aware write.
writeTo func(flow int, bufs [][]byte, offset int) (int, error)
// startPollingOnce is used to start a [Wrapper.pollVector] goroutine at the
// first call to [Wrapper.Read].
// first read of queue 0.
startPollingOnce sync.Once
// bufferConsumedMu protects bufferConsumed from concurrent sends, closes,
// and send-after-close (by way of bufferConsumedClosed).
@@ -132,7 +141,7 @@ type Wrapper struct {
// read by bufferConsumed writers to prevent send-after-close.
bufferConsumedClosed bool
// bufferConsumed synchronizes access to packet bufs and descriptors shared
// by [Wrapper.Read] and [Wrapper.pollVector].
// by queue 0's read and [Wrapper.pollVector].
//
// Close closes bufferConsumed and sets bufferConsumedClosed to true.
bufferConsumed chan struct{}
@@ -284,7 +293,20 @@ type tunVectorReadResult struct {
injected tunInjectedRead
}
// Start unblocks any Wrapper.Read calls that have already started
// wrapperQueue is one read queue of a [Wrapper], wrapping a single queue of
// the underlying [tun.Device]. Distinct wrapperQueues may be read concurrently.
type wrapperQueue struct {
w *Wrapper
q tun.Queue
// marks queue for injection, see [Wrapper.readMultiplexed]
isInjectionQueue bool
}
// File implements [tun.Queue].
func (q *wrapperQueue) File() *os.File { return q.q.File() }
// Start unblocks any queue reads that have already started
// and makes the Wrapper functional.
//
// Start must be called exactly once after the various Tailscale
@@ -309,6 +331,7 @@ func wrap(logf logger.Logf, tdev tun.Device, isTAP bool, m *usermetric.Registry,
limitedLogf: logger.RateLimitedFn(logf, 1*time.Minute, 2, 10),
isTAP: isTAP,
tdev: tdev,
writeTo: tun.WriteToOf(tdev),
// bufferConsumed is conceptually a condition variable:
// a goroutine should not block when setting it, even with no listeners.
bufferConsumed: make(chan struct{}, 1),
@@ -322,6 +345,13 @@ func wrap(logf logger.Logf, tdev tun.Device, isTAP bool, m *usermetric.Registry,
startCh: make(chan struct{}),
metrics: registerMetrics(m),
}
for i, q := range tun.QueuesOf(tdev) {
w.queues = append(w.queues, &wrapperQueue{
w: w,
isInjectionQueue: i == 0, // injected packets are multiplexed onto queue 0
q: q,
})
}
if buildfeatures.HasTUNDevStats {
if f, ok := HookPollTUNDevStats.GetOk(); ok {
@@ -473,14 +503,25 @@ func (t *Wrapper) Name() (string, error) {
return t.tdev.Name()
}
var (
_ tun.MultiQueueDevice = (*Wrapper)(nil)
_ tun.Queue = (*wrapperQueue)(nil)
)
// Queues implements [tun.MultiQueueDevice] with one queue per queue of the
// underlying device. Injected packets are multiplexed onto queue 0.
func (t *Wrapper) Queues() []tun.Queue {
return slices.Clone(t.queues)
}
// pollVector polls [Wrapper.tdev.Read], writing the oldest unconsumed packet
// slab and packet descriptors into the [Wrapper.vectorOutbound] channel.
// slabLen and packetsLen should originate from the first call to [Wrapper.Read],
// slabLen and packetsLen should originate from the first read of queue 0,
// and are used for sizing the equivalent arguments pollVector passes to
// [Wrapper.tdev.Read].
//
// [Wrapper.tdev.Read] can block, so we poll tdev in a goroutine independent of
// wireguard-go's calls to [Wrapper.Read], in order to support native tdev reads
// wireguard-go's reads of queue 0, in order to support native tdev reads
// alongside packets we inject.
//
// pollVector returns when [t.bufferConsumed] is closed, or when [Wrapper.isClosed]
@@ -877,9 +918,28 @@ func (t *Wrapper) awaitStart() {
// Read implements [tun.Device.Read].
func (t *Wrapper) Read(slab []byte, packets []tun.ReadPacket) (int, error) {
return t.queues[0].Read(slab, packets)
}
// Read implements [tun.Reader]. Queue 0 carries injected packets as well as the
// device's own, multiplexed by [Wrapper.pollVector], every other queue reads
// the device directly into slab.
func (q *wrapperQueue) Read(slab []byte, packets []tun.ReadPacket) (int, error) {
t := q.w
if !t.started.Load() {
t.awaitStart()
}
if q.isInjectionQueue {
return t.readMultiplexed(slab, packets)
}
return q.readDevice(slab, packets)
}
// readMultiplexed multiplexes injected reads into the underlying TUN
// queue's data stream.
//
// TODO(illotum): give injected packets a queue of their own and retire this.
func (t *Wrapper) readMultiplexed(slab []byte, packets []tun.ReadPacket) (int, error) {
t.startPollingOnce.Do(func() {
go t.pollVector(len(slab), len(packets))
})
@@ -899,8 +959,39 @@ func (t *Wrapper) Read(slab []byte, packets []tun.ReadPacket) (int, error) {
if res.real.err != nil && len(res.real.packets) == 0 {
return 0, res.real.err
}
return t.filterOutbound(res.real.slab, res.real.packets, slab, packets), res.real.err
}
metricPacketOut.Add(int64(len(res.real.packets)))
// readDevice reads q's queue of the underlying device directly into slab.
func (q *wrapperQueue) readDevice(slab []byte, packets []tun.ReadPacket) (int, error) {
t := q.w
var n int
var err error
// Empty reads are skipped by WireGuard, it is legal to discard an empty read.
for n == 0 && err == nil {
if t.isClosed() {
return 0, io.EOF
}
n, err = q.q.Read(slab, packets)
if t.isTAP && TAPDebug {
s := fmt.Sprintf("% x", slab)
for strings.HasSuffix(s, " 00") {
s = strings.TrimSuffix(s, " 00")
}
t.logf("TAP read: %v, %v: %s", n, err, s)
}
}
if err != nil && n == 0 {
return 0, err
}
return t.filterOutbound(slab, packets[:n], slab, packets), err
}
// filterOutbound runs the outbound filter pipeline over the packets described
// by read, which live in src. Survivors are copied into slab and their
// descriptors compacted into packets. It returns the number of survivors.
func (t *Wrapper) filterOutbound(src []byte, read []tun.ReadPacket, slab []byte, packets []tun.ReadPacket) int {
metricPacketOut.Add(int64(len(read)))
var numPackets int
p := parsedPacketPool.Get().(*packet.Parsed)
@@ -908,8 +999,8 @@ func (t *Wrapper) Read(slab []byte, packets []tun.ReadPacket) (int, error) {
captHook := t.captureHook.Load()
pc := t.peerConfig.Load()
var buffsGRO *gro.GRO
for _, meta := range res.real.packets {
data := res.real.slab[meta.Offset : meta.Offset+meta.Size]
for _, meta := range read {
data := src[meta.Offset : meta.Offset+meta.Size]
p.Decode(data)
if buildfeatures.HasCapture && captHook != nil {
@@ -932,6 +1023,7 @@ func (t *Wrapper) Read(slab []byte, packets []tun.ReadPacket) (int, error) {
// Make sure to do SNAT after filtering, so that any flow tracking in
// the filter sees the original source address. See #12133.
pc.snat(p)
// A no-op when src is slab: p.Buffer() is then the destination too.
n := copy(slab[meta.Offset:meta.Offset+meta.Size], p.Buffer())
if n != len(data) {
panic(fmt.Sprintf("short copy: %d != %d", n, len(data)))
@@ -944,7 +1036,7 @@ func (t *Wrapper) Read(slab []byte, packets []tun.ReadPacket) (int, error) {
}
t.noteActivity()
return numPackets, res.real.err
return numPackets
}
const (
@@ -1254,11 +1346,19 @@ func (t *Wrapper) filterPacketInboundFromWireGuard(p *packet.Parsed, captHook pa
return filter.Accept, gro
}
// Write accepts incoming packets. The packets begin at buffs[:][offset:],
// Write accepts incoming packets. It is equivalent to [Wrapper.WriteTo]
// with a zero flow.
func (t *Wrapper) Write(buffs [][]byte, offset int) (int, error) {
return t.WriteTo(0, buffs, offset)
}
// WriteTo implements [tun.MultiQueueDevice]. The packets begin at buffs[:][offset:],
// like wireguard-go/tun.Device.Write. Write is called per-peer via
// wireguard-go/device.Peer.RoutineSequentialReceiver, so it MUST be
// thread-safe.
func (t *Wrapper) Write(buffs [][]byte, offset int) (int, error) {
//
// Packets are dispatched to the write queue selected by flow.
func (t *Wrapper) WriteTo(flow int, buffs [][]byte, offset int) (int, error) {
metricPacketIn.Add(int64(len(buffs)))
i := 0
p := parsedPacketPool.Get().(*packet.Parsed)
@@ -1293,7 +1393,7 @@ func (t *Wrapper) Write(buffs [][]byte, offset int) (int, error) {
if len(buffs) > 0 {
t.noteActivity()
_, err := t.tdevWrite(buffs, offset)
_, err := t.tdevWrite(flow, buffs, offset)
if err != nil {
t.metrics.inboundDroppedPacketsTotal.Add(usermetric.DropLabels{
Reason: usermetric.ReasonError,
@@ -1304,7 +1404,7 @@ func (t *Wrapper) Write(buffs [][]byte, offset int) (int, error) {
return 0, nil
}
func (t *Wrapper) tdevWrite(buffs [][]byte, offset int) (int, error) {
func (t *Wrapper) tdevWrite(flow int, buffs [][]byte, offset int) (int, error) {
if buildfeatures.HasNetLog {
if update := t.connCounter.Load(); update != nil {
for i := range buffs {
@@ -1312,7 +1412,7 @@ func (t *Wrapper) tdevWrite(buffs [][]byte, offset int) (int, error) {
}
}
}
return t.tdev.Write(buffs, offset)
return t.writeTo(flow, buffs, offset)
}
func (t *Wrapper) GetFilter() *filter.Filter {
@@ -1356,6 +1456,8 @@ func (t *Wrapper) SetJailedFilter(filt *filter.Filter) {
//
// This path is typically used to deliver synthesized packets to the
// host networking stack.
// Injecting packets from a valid peer will lead to TUN queue switching
// and TCP reorders.
func (t *Wrapper) InjectInboundPacketBuffer(pkt *netstack_PacketBuffer, slab []byte, packets []tun.ReadPacket, writeBufs [][]byte) error {
if !buildfeatures.HasNetstack {
panic("unreachable")
@@ -1404,7 +1506,7 @@ func (t *Wrapper) InjectInboundPacketBuffer(pkt *netstack_PacketBuffer, slab []b
for i, meta := range packets[:n] {
writeBufs[i] = slab[meta.Offset-WritePacketStartOffset : meta.Offset+meta.Size]
}
_, err = t.tdevWrite(writeBufs[:n], WritePacketStartOffset)
_, err = t.tdevWrite(0, writeBufs[:n], WritePacketStartOffset)
return err
}
@@ -1416,6 +1518,9 @@ func (t *Wrapper) InjectInboundPacketBuffer(pkt *netstack_PacketBuffer, slab []b
// The packet contents are to start at &buf[offset].
// offset must be greater or equal to WritePacketStartOffset.
// The space before &buf[offset] will be used by WireGuard.
//
// Injecting packets from a valid peer will lead to TUN queue switching,
// and TCP reorders.
func (t *Wrapper) InjectInboundDirect(buf []byte, offset int) error {
if len(buf) > MaxPacketSize {
return errPacketTooBig
@@ -1428,13 +1533,16 @@ func (t *Wrapper) InjectInboundDirect(buf []byte, offset int) error {
}
// Write to the underlying device to skip filters.
_, err := t.tdevWrite([][]byte{buf}, offset) // TODO(jwhited): alloc?
_, err := t.tdevWrite(0, [][]byte{buf}, offset) // TODO(jwhited): alloc?
return err
}
// InjectInboundCopy takes a packet without leading space,
// reallocates it to conform to the InjectInboundDirect interface
// and calls InjectInboundDirect on it. Injecting a nil packet is a no-op.
//
// Injecting packets from a valid peer will lead to TUN queue switching,
// and TCP reorders.
func (t *Wrapper) InjectInboundCopy(packet []byte) error {
// We duplicate this check from InjectInboundDirect here
// to avoid wasting an allocation on an oversized packet.
+74
View File
@@ -13,8 +13,10 @@ import (
"fmt"
"net/netip"
"reflect"
"slices"
"strconv"
"strings"
"sync"
"testing"
"time"
"unicode"
@@ -1375,3 +1377,75 @@ func TestStackGSOToTunGSO(t *testing.T) {
})
}
}
type flowRecorder struct {
*fakeTUN
mu sync.Mutex
flows []int
}
func (d *flowRecorder) WriteTo(flow int, bufs [][]byte, offset int) (int, error) {
d.mu.Lock()
defer d.mu.Unlock()
d.flows = append(d.flows, flow)
return len(bufs), nil
}
func (d *flowRecorder) Write(bufs [][]byte, offset int) (int, error) {
return d.WriteTo(0, bufs, offset)
}
func (d *flowRecorder) recordedFlows() []int {
d.mu.Lock()
defer d.mu.Unlock()
return slices.Clone(d.flows)
}
func TestWrappedQueuesMatch(t *testing.T) {
for _, q := range []int{1, 4} {
t.Run(fmt.Sprintf("q=%d", q), func(t *testing.T) {
tdev := newFakeMQ(q)
bus := eventbustest.NewBus(t)
want := wgtun.QueuesOf(tdev)
if len(want) != q {
t.Fatalf("newFakeMQ(%d) returned %d queues", q, len(want))
}
w := Wrap(t.Logf, tdev, new(usermetric.Registry), bus)
defer w.Close()
got := w.Queues()
if have, want := len(got), len(want); have != want {
t.Fatalf("len(Queues()) = %d, want %d", have, want)
}
for i, q := range got {
if got, ok := q.(*wrapperQueue); !ok || got.q != want[i] {
t.Errorf("Queues()[%d] wraps the wrong underlying queue", i)
}
}
})
}
}
func TestWriteToForwardsFlow(t *testing.T) {
bus := eventbustest.NewBus(t)
tdev := &flowRecorder{fakeTUN: newFakeMQ(4)}
w := Wrap(t.Logf, tdev, new(usermetric.Registry), bus)
w.disableFilter = true
w.Start()
defer w.Close()
pkt := udp4("100.64.1.2", "100.64.1.3", 1234, 5678)
want := []int{3, 1, 2}
for _, flow := range want {
if _, err := w.WriteTo(flow, [][]byte{pkt}, 0); err != nil {
t.Fatalf("WriteTo(%d): %v", flow, err)
}
}
// Write is WriteTo with a zero flow.
want = append(want, 0)
if _, err := w.Write([][]byte{pkt}, 0); err != nil {
t.Fatalf("Write: %v", err)
}
if got := tdev.recordedFlows(); !slices.Equal(got, want) {
t.Errorf("flows reaching the device = %v, want %v", got, want)
}
}