diff --git a/net/tstun/fake.go b/net/tstun/fake.go index d6a981576..d0ffcae44 100644 --- a/net/tstun/fake.go +++ b/net/tstun/fake.go @@ -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) { diff --git a/net/tstun/wrap.go b/net/tstun/wrap.go index a29b02c07..511245ef9 100644 --- a/net/tstun/wrap.go +++ b/net/tstun/wrap.go @@ -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. diff --git a/net/tstun/wrap_test.go b/net/tstun/wrap_test.go index 2b92b5a37..091d7d1cc 100644 --- a/net/tstun/wrap_test.go +++ b/net/tstun/wrap_test.go @@ -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) + } +}