From fd38051d5e09c0da998a1a44ce7ae73cef68682f Mon Sep 17 00:00:00 2001 From: Patricio Whittingslow Date: Mon, 15 Jun 2026 18:51:39 -0300 Subject: [PATCH 01/12] begin adding lneto --- .gitignore | 1 + go.mod | 5 +- tun/netstack/lneto.go | 675 +++++++++++++++++++++++++++++++++++++ tun/netstack/lneto_test.go | 300 +++++++++++++++++ 4 files changed, 980 insertions(+), 1 deletion(-) create mode 100644 tun/netstack/lneto.go create mode 100644 tun/netstack/lneto_test.go diff --git a/.gitignore b/.gitignore index e460293e9..66abd3cc2 100644 --- a/.gitignore +++ b/.gitignore @@ -1 +1,2 @@ wireguard-go +*LNETO_EQUIVALENCE.md \ No newline at end of file diff --git a/go.mod b/go.mod index 2a80e0001..e1fa5e866 100644 --- a/go.mod +++ b/go.mod @@ -1,8 +1,9 @@ module golang.zx2c4.com/wireguard -go 1.23.1 +go 1.24 require ( + github.com/soypat/lneto v0.0.0-00010101000000-000000000000 golang.org/x/crypto v0.37.0 golang.org/x/net v0.39.0 golang.org/x/sys v0.32.0 @@ -14,3 +15,5 @@ require ( github.com/google/btree v1.1.2 // indirect golang.org/x/time v0.7.0 // indirect ) + +replace github.com/soypat/lneto => ../lneto diff --git a/tun/netstack/lneto.go b/tun/netstack/lneto.go new file mode 100644 index 000000000..6e7ff54b0 --- /dev/null +++ b/tun/netstack/lneto.go @@ -0,0 +1,675 @@ +/* SPDX-License-Identifier: MIT + * + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. + */ + +package netstack + +import ( + "context" + crand "crypto/rand" + "encoding/binary" + "errors" + "fmt" + "net" + "net/netip" + "os" + "regexp" + "runtime" + "strconv" + "strings" + "sync" + "syscall" + "time" + + "github.com/soypat/lneto" + "github.com/soypat/lneto/dns" + "github.com/soypat/lneto/x/xnet" + "golang.zx2c4.com/wireguard/tun" +) + +// Net2 is a lneto-backed userspace network stack that implements both +// [tun.Device] (for WireGuard integration) and a networking API (Dial/Listen/DNS). +// +// Packet flow: +// - Ingress (WireGuard → stack): [Net2.Write] calls [xnet.StackAsync.IngressIP], +// then pokes the backoff irq so a blocked [Net2.Read] wakes immediately. +// - Egress (stack → WireGuard): [Net2.Read] polls [xnet.StackAsync.EgressIP] +// directly, sleeping on an interruptible backoff between empty polls. +// +// Unlike gVisor's channel.Endpoint there is no native egress notification hook, so +// egress is driven by Read's poll loop. The backoff is interruptible (woken by +// [Net2.interrupt]) and GOMAXPROCS-aware: on a single-threaded runtime it only +// yields cooperatively (sleeping would starve the poll), otherwise it sleeps with +// exponential backoff. This mirrors the proven go-net design. +// +// Lifecycle: created by [CreateNetTUN2], torn down by [Net2.Close] which closes +// the closed channel (unblocking Read and any blocking socket op) and events. +type Net2 struct { + sa xnet.StackAsync + sgo xnet.StackGo // wraps sa; created once in CreateNetTUN2 + + // events carries TUN state changes (e.g. EventUp) consumed by WireGuard's device loop. + events chan tun.Event + + // backoff is the interruptible stack-protocol backoff shared with blk/sgo and + // used by Read's egress poll. backoffirq wakes a sleeping backoff the moment new + // work arrives (ingress packet, socket call) via interrupt. + backoff lneto.BackoffStrategy + backoffirq chan<- event + + // closed is closed once by Close to unblock Read and signal shutdown. + closed chan struct{} + closeOnce sync.Once + + mtu int + dnsServers []netip.Addr + hasV4, hasV6 bool +} + +type TCPConn interface { + Close() error + CloseRead() error + CloseWrite() error + LocalAddr() net.Addr + Read(b []byte) (int, error) + RemoteAddr() net.Addr + SetDeadline(t time.Time) error + SetReadDeadline(t time.Time) error + SetWriteDeadline(t time.Time) error + Write(b []byte) (int, error) +} + +type UDPConn interface { + Close() error + LocalAddr() net.Addr + Read(b []byte) (int, error) + ReadFrom(b []byte) (int, net.Addr, error) + RemoteAddr() net.Addr + SetDeadline(t time.Time) error + SetReadDeadline(t time.Time) error + SetWriteDeadline(t time.Time) error + Write(b []byte) (int, error) + WriteTo(b []byte, addr net.Addr) (int, error) +} + +type TCPListener interface { + Accept() (net.Conn, error) + Addr() net.Addr + Close() error + Shutdown() +} + +type event struct{} + +// interruptBackoff wraps a [lneto.BackoffStrategy] so its sleep can be cut short by +// a write to the returned irq channel. Each call builds its own timer so concurrent +// callers (Read's poll loop and blocking socket ops) do not race on a shared one; +// the capacity-1 irq is shared, so one interrupt wakes exactly one sleeper, which is +// the intended best-effort behavior. The returned strategy always reports +// [lneto.BackoffFlagNop] because the yield is performed entirely inside the wrapper. +func interruptBackoff(backoff lneto.BackoffStrategy) (interrupt chan<- event, _ lneto.BackoffStrategy) { + irq := make(chan event, 1) + wrapped := func(consecutiveBackoffs uint) time.Duration { + switch d := backoff(consecutiveBackoffs); d { + case lneto.BackoffFlagGosched: + runtime.Gosched() + case lneto.BackoffFlagNop: + // Do nothing. + default: + timer := time.NewTimer(d) // per-call: callers must not share a timer. + select { + case <-irq: + if !timer.Stop() && len(timer.C) > 0 { + <-timer.C + } + case <-timer.C: + } + } + return lneto.BackoffFlagNop // yield handled here; signal caller to do nothing. + } + return irq, wrapped +} + +// defaultStackBackoff is the idle backoff for stack protocol loops (DHCP, DNS, the +// egress poll, ...) on multi-threaded runtimes: exponential from 100µs up to 20ms. +func defaultStackBackoff(consecutiveBackoffs uint) time.Duration { + const ( + minWait = 100 * time.Microsecond + maxWait = 20 * time.Millisecond + maxShift = 15 + + _compileTimeOverflowCheck = minWait << maxShift + ) + sleep := minWait << min(consecutiveBackoffs, maxShift) + if sleep > maxWait { + sleep = maxWait + } + return sleep +} + +// defaultTCPBackoff is the per-connection read/write retry backoff for TCP streams. +// Shorter range than defaultStackBackoff to keep interactive sessions responsive. +func defaultTCPBackoff(consecutiveBackoffs uint) time.Duration { + const ( + minWait = 10 * time.Microsecond + maxWait = 1 * time.Millisecond + maxShift = 10 + + _compileTimeOverflowCheck = minWait << maxShift + ) + sleep := minWait << min(consecutiveBackoffs, maxShift) + if sleep > maxWait { + sleep = maxWait + } + return sleep +} + +// backoffYield never sleeps; it only yields cooperatively. Used on GOMAXPROCS==1 +// where sleeping the poll goroutine would starve egress processing. +func backoffYield(consecutiveBackoffs uint) time.Duration { + return lneto.BackoffFlagGosched +} + +// interrupt wakes one sleeper blocked in the interruptible backoff (Read's poll loop +// or a blocking socket op). Non-blocking and safe to call from any goroutine. +func (n *Net2) interrupt() { + select { + case n.backoffirq <- event{}: + default: + } +} + +// --- tun.Device implementation --- + +func (n *Net2) Name() (string, error) { return "go2", nil } +func (n *Net2) File() *os.File { return nil } +func (n *Net2) Events() <-chan tun.Event { return n.events } +func (n *Net2) MTU() (int, error) { return n.mtu, nil } +func (n *Net2) BatchSize() int { return 1 } + +// Write feeds incoming IP packets (WireGuard → stack) into the lneto stack. +func (n *Net2) Write(bufs [][]byte, offset int) (int, error) { + wrote := false + for _, buf := range bufs { + if pkt := buf[offset:]; len(pkt) > 0 { + n.sa.IngressIP(pkt) // errors dropped; stack silently filters bad packets + wrote = true + } + } + if wrote { + n.interrupt() // wake Read: ingress often produces an immediate egress reply. + } + return len(bufs), nil +} + +// Read blocks until the stack has an outgoing IP packet to send to WireGuard, polling +// EgressIP and sleeping on the interruptible backoff between empty polls. It writes +// directly into the caller's buffer (no intermediate copy). Returns os.ErrClosed once +// Close has been called. +func (n *Net2) Read(bufs [][]byte, sizes []int, offset int) (int, error) { + dst := bufs[0][offset:] + var backoffs uint + for { + select { + case <-n.closed: + return 0, os.ErrClosed + default: + } + cnt, _ := n.sa.EgressIP(dst) + if cnt > 0 { + sizes[0] = cnt + return 1, nil + } + n.backoff.Do(backoffs) // interruptible; returns promptly when interrupt fires. + backoffs++ + } +} + +func (n *Net2) Close() error { + n.closeOnce.Do(func() { + close(n.closed) // unblock Read. + n.interrupt() // wake a backoff sleeper so it observes closed promptly. + close(n.events) + }) + return nil +} + +func CreateNetTUN2(localAddresses, dnsServers []netip.Addr, mtu int) (tun.Device, *Net2, error) { + if mtu <= 0 { + mtu = 1500 + } + dev := &Net2{ + events: make(chan tun.Event, 10), + closed: make(chan struct{}), + mtu: mtu, + dnsServers: dnsServers, + } + + var hwAddr [6]byte + if _, err := crand.Read(hwAddr[:]); err != nil { + return nil, nil, fmt.Errorf("CreateNetTUN2: rand MAC: %w", err) + } + hwAddr[0] &^= 0x01 // unicast + hwAddr[0] |= 0x02 // locally administered + + // Pick the first IPv4 and first IPv6 local address; track which families are + // configured (mirrors gVisor's hasV4/hasV6). + var staticAddr4 [4]byte + var staticAddr6 [16]byte + for _, addr := range localAddresses { + if addr.Is4() && !dev.hasV4 { + staticAddr4 = addr.As4() + dev.hasV4 = true + } else if addr.Is6() && !dev.hasV6 { + staticAddr6 = addr.As16() + dev.hasV6 = true + } + } + + // DNS server: prefer IPv4 (the stack's lookup path is currently IPv4-only), + // fall back to the first IPv6 server. + var dnsServer netip.Addr + for _, d := range dnsServers { + if d.Is4() { + dnsServer = d + break + } + } + if !dnsServer.IsValid() { + for _, d := range dnsServers { + if d.Is6() { + dnsServer = d + break + } + } + } + + var randSeed int64 + if err := binary.Read(crand.Reader, binary.LittleEndian, &randSeed); err != nil { + return nil, nil, fmt.Errorf("CreateNetTUN2: rand seed: %w", err) + } + + cfg := xnet.StackConfig{ + HardwareAddress: hwAddr, + StaticAddress4: staticAddr4, + MTU: uint16(mtu), + Hostname: "wg0", + RandSeed: randSeed, + PassivePeers: 0, // no ARP passive learning needed for TUN + ICMPQueueLimit: 4, + MaxActiveTCPPorts: 256, + MaxActiveUDPPorts: 256, + DNSServer: dnsServer, + } + if dev.hasV6 { + cfg.StaticAddress6 = staticAddr6 + cfg.IPv6Stack = xnet.DefaultStack6() + } + // NOTE: ICMP is intentionally NOT enabled here. On a TUN there is no link layer, + // so MAC resolution must be skipped: IPv4 ARP is gated off by leaving the subnet + // unset, and IPv6 NDP is gated off by leaving ICMPv6 unregistered. Enabling ICMP + // would make DialTCP6/DialUDP6 emit Neighbor Solicitations that are never answered + // on a TUN, breaking IPv6 dialing. Ping support (which needs ICMP) is a separate + // follow-up that must reconcile this. + if err := dev.sa.Reset(cfg); err != nil { + return nil, nil, fmt.Errorf("CreateNetTUN2: stack reset: %w", err) + } + + // GOMAXPROCS-aware backoff: on a single-threaded runtime only yield cooperatively + // (sleeping the poll would starve egress), otherwise sleep with exponential backoff. + baseStack := defaultStackBackoff + newTCPBackoff := func() lneto.BackoffStrategy { return defaultTCPBackoff } + if runtime.GOMAXPROCS(0) == 1 { + baseStack = backoffYield + newTCPBackoff = func() lneto.BackoffStrategy { return backoffYield } + } + irq, backoff := interruptBackoff(baseStack) + dev.backoff = backoff + dev.backoffirq = irq + + dev.sgo = dev.sa.StackGo(backoff, xnet.StackGoConfig{ + ListenerPoolConfig: xnet.TCPPoolConfig{ + PoolSize: 256, + QueueSize: 8, + TxBufSize: 32 << 10, + RxBufSize: 32 << 10, + EstablishedTimeout: 30 * time.Second, + ClosingTimeout: 10 * time.Second, + NewBackoff: newTCPBackoff, // required: StackGo panics if nil. + }, + }) + dev.events <- tun.EventUp + return dev, dev, nil +} + +// --- TCP --- + +// socketResult extracts a typed result from a SocketNetip call. +// SocketNetip's TCP dial branch returns connection errors as the value (not err) +// to distinguish stack-level failures from protocol errors, so we handle both. +func socketResult[T any](v any, err error) (T, error) { + var zero T + if err != nil { + return zero, err + } + if e, ok := v.(error); ok { + return zero, e + } + t, ok := v.(T) + if !ok { + return zero, fmt.Errorf("socket: unexpected type %T", v) + } + return t, nil +} + +// socket wraps sgo.SocketNetip. It selects the IPv4 or IPv6 network/family from the +// family-bearing endpoint (the remote for a dial, otherwise the local bind), and pokes +// the egress poll on entry and exit because connection setup (handshake, NDP/ARP) +// queues egress frames that Read must drain promptly. +func (n *Net2) socket(ctx context.Context, proto string, sotype int, laddr, raddr netip.AddrPort) (any, error) { + fam := raddr + if !fam.IsValid() { + fam = laddr + } + network := proto + "4" + family := syscall.AF_INET + if fam.Addr().Is6() { + network = proto + "6" + family = syscall.AF_INET6 + } + n.interrupt() + defer n.interrupt() + return n.sgo.SocketNetip(ctx, network, family, sotype, laddr, raddr) +} + +func (n *Net2) dialTCPCtx(ctx context.Context, addr netip.AddrPort) (TCPConn, error) { + v, err := n.socket(ctx, "tcp", syscall.SOCK_STREAM, netip.AddrPort{}, addr) + return socketResult[TCPConn](v, err) +} + +func (n *Net2) DialContextTCPAddrPort(ctx context.Context, addr netip.AddrPort) (TCPConn, error) { + return n.dialTCPCtx(ctx, addr) +} + +func (n *Net2) DialContextTCP(ctx context.Context, addr *net.TCPAddr) (TCPConn, error) { + if addr == nil { + return n.dialTCPCtx(ctx, netip.AddrPort{}) + } + ip, _ := netip.AddrFromSlice(addr.IP) + return n.dialTCPCtx(ctx, netip.AddrPortFrom(ip.Unmap(), uint16(addr.Port))) +} + +func (n *Net2) DialTCPAddrPort(addr netip.AddrPort) (TCPConn, error) { + return n.dialTCPCtx(context.Background(), addr) +} + +func (n *Net2) DialTCP(addr *net.TCPAddr) (TCPConn, error) { + return n.DialContextTCP(context.Background(), addr) +} + +// --- TCP listener --- + +func (n *Net2) ListenTCPAddrPort(addr netip.AddrPort) (TCPListener, error) { + v, err := n.socket(context.Background(), "tcp", syscall.SOCK_STREAM, addr, netip.AddrPort{}) + return socketResult[TCPListener](v, err) +} + +func (n *Net2) ListenTCP(addr *net.TCPAddr) (TCPListener, error) { + if addr == nil { + return n.ListenTCPAddrPort(netip.AddrPort{}) + } + ip, _ := netip.AddrFromSlice(addr.IP) + return n.ListenTCPAddrPort(netip.AddrPortFrom(ip.Unmap(), uint16(addr.Port))) +} + +// --- UDP --- + +func (n *Net2) ListenUDPAddrPort(laddr netip.AddrPort) (UDPConn, error) { + v, err := n.socket(context.Background(), "udp", syscall.SOCK_DGRAM, laddr, netip.AddrPort{}) + return socketResult[UDPConn](v, err) +} + +func (n *Net2) ListenUDP(laddr *net.UDPAddr) (UDPConn, error) { + if laddr == nil { + return n.ListenUDPAddrPort(netip.AddrPort{}) + } + ip, _ := netip.AddrFromSlice(laddr.IP) + return n.ListenUDPAddrPort(netip.AddrPortFrom(ip.Unmap(), uint16(laddr.Port))) +} + +func (n *Net2) DialUDPAddrPort(laddr, raddr netip.AddrPort) (UDPConn, error) { + v, err := n.socket(context.Background(), "udp", syscall.SOCK_DGRAM, laddr, raddr) + return socketResult[UDPConn](v, err) +} + +func (n *Net2) DialUDP(laddr, raddr *net.UDPAddr) (UDPConn, error) { + var la, ra netip.AddrPort + if laddr != nil { + ip, _ := netip.AddrFromSlice(laddr.IP) + la = netip.AddrPortFrom(ip.Unmap(), uint16(laddr.Port)) + } + if raddr != nil { + ip, _ := netip.AddrFromSlice(raddr.IP) + ra = netip.AddrPortFrom(ip.Unmap(), uint16(raddr.Port)) + } + return n.DialUDPAddrPort(la, ra) +} + +// --- Ping --- + +func (n *Net2) DialPingAddr(_, _ netip.Addr) (*PingConn, error) { + return nil, errors.New("ping not implemented for Net2: PingConn is gvisor-coupled") +} + +func (n *Net2) ListenPingAddr(_ netip.Addr) (*PingConn, error) { + return nil, errors.New("ping not implemented for Net2: PingConn is gvisor-coupled") +} + +func (n *Net2) DialPing(laddr, raddr *PingAddr) (*PingConn, error) { + var la, ra netip.Addr + if laddr != nil { + la = laddr.addr + } + if raddr != nil { + ra = raddr.addr + } + return n.DialPingAddr(la, ra) +} + +func (n *Net2) ListenPing(laddr *PingAddr) (*PingConn, error) { + var la netip.Addr + if laddr != nil { + la = laddr.addr + } + return n.ListenPingAddr(la) +} + +// --- DNS --- + +// dnsError wraps a lookup failure as a *net.DNSError, flagging timeouts when the +// underlying error reports them. Mirrors the error shape produced by the gvisor Net. +func dnsError(host string, err error) *net.DNSError { + de := &net.DNSError{Err: err.Error(), Name: host} + if nerr, ok := err.(net.Error); ok && nerr.Timeout() { + de.IsTimeout = true + } + return de +} + +// LookupContextHost resolves host to a list of IP strings, matching the behaviour of +// the gvisor Net: literal IPs (with IPv6 zone stripping) pass through; empty or +// non-domain hosts and stacks with no address family return an IsNotFound DNSError; +// A and AAAA are queried for the enabled families and, when IPv6 is enabled, IPv6 +// results are ordered first (no RFC 6724). +func (n *Net2) LookupContextHost(ctx context.Context, host string) ([]string, error) { + if host == "" || (!n.hasV4 && !n.hasV6) { + return nil, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true} + } + // Strip any IPv6 zone before attempting to parse a literal address. + zlen := len(host) + if strings.IndexByte(host, ':') != -1 { + if zidx := strings.LastIndexByte(host, '%'); zidx != -1 { + zlen = zidx + } + } + if ip, err := netip.ParseAddr(host[:zlen]); err == nil { + return []string{ip.String()}, nil + } + if !isDomainName(host) { + return nil, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true} + } + + timeout := 5 * time.Second + if dl, ok := ctx.Deadline(); ok { + if rem := time.Until(dl); rem < timeout { + timeout = rem + } + } + blk := n.sa.StackBlocking(n.backoff) + + var addrsV4, addrsV6 []netip.Addr + var lastErr error + if n.hasV4 { + if a, err := blk.DoLookupIPType(host, timeout, dns.TypeA); err != nil { + lastErr = dnsError(host, err) + } else { + addrsV4 = a + } + } + if n.hasV6 { + if a, err := blk.DoLookupIPType(host, timeout, dns.TypeAAAA); err != nil { + if lastErr == nil { + lastErr = dnsError(host, err) + } + } else { + addrsV6 = a + } + } + + // IPv6 first when enabled, mirroring the gvisor Net's ordering. + var addrs []netip.Addr + if n.hasV6 { + addrs = append(addrsV6, addrsV4...) + } else { + addrs = append(addrsV4, addrsV6...) + } + if len(addrs) == 0 { + if lastErr != nil { + return nil, lastErr + } + return nil, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true} + } + out := make([]string, len(addrs)) + for i, a := range addrs { + out[i] = a.String() + } + return out, nil +} + +func (n *Net2) LookupHost(host string) ([]string, error) { + return n.LookupContextHost(context.Background(), host) +} + +// --- Generic Dial --- + +var protoSplitter2 = regexp.MustCompile(`^(tcp|udp|ping)(4|6)?$`) + +func (n *Net2) DialContext(ctx context.Context, network, address string) (net.Conn, error) { + if ctx == nil { + panic("nil context") + } + matches := protoSplitter2.FindStringSubmatch(network) + if matches == nil { + return nil, &net.OpError{Op: "dial", Err: net.UnknownNetworkError(network)} + } + acceptV4 := len(matches[2]) == 0 || matches[2] == "4" + acceptV6 := len(matches[2]) == 0 || matches[2] == "6" + + var host string + var port int + if matches[1] == "ping" { + host = address + } else { + var sport string + var err error + host, sport, err = net.SplitHostPort(address) + if err != nil { + return nil, &net.OpError{Op: "dial", Err: err} + } + port, err = strconv.Atoi(sport) + if err != nil || port < 0 || port > 65535 { + return nil, &net.OpError{Op: "dial", Err: errNumericPort} + } + } + + allAddr, err := n.LookupContextHost(ctx, host) + if err != nil { + return nil, &net.OpError{Op: "dial", Err: err} + } + + var addrs []netip.AddrPort + for _, a := range allAddr { + ip, err := netip.ParseAddr(a) + if err == nil && ((ip.Is4() && acceptV4) || (ip.Is6() && acceptV6)) { + addrs = append(addrs, netip.AddrPortFrom(ip, uint16(port))) + } + } + if len(addrs) == 0 && len(allAddr) != 0 { + return nil, &net.OpError{Op: "dial", Err: errNoSuitableAddress} + } + + var firstErr error + for i, addr := range addrs { + select { + case <-ctx.Done(): + err := ctx.Err() + if err == context.Canceled { + err = errCanceled + } else if err == context.DeadlineExceeded { + err = errTimeout + } + return nil, &net.OpError{Op: "dial", Err: err} + default: + } + dialCtx := ctx + if deadline, hasDeadline := ctx.Deadline(); hasDeadline { + pd, err := partialDeadline(time.Now(), deadline, len(addrs)-i) + if err != nil { + if firstErr == nil { + firstErr = &net.OpError{Op: "dial", Err: err} + } + break + } + if pd.Before(deadline) { + var cancel context.CancelFunc + dialCtx, cancel = context.WithDeadline(ctx, pd) + defer cancel() + } + } + + var c net.Conn + switch matches[1] { + case "tcp": + c, err = n.DialContextTCPAddrPort(dialCtx, addr) + case "udp": + c, err = n.DialUDPAddrPort(netip.AddrPort{}, addr) + case "ping": + c, err = n.DialPingAddr(netip.Addr{}, addr.Addr()) + } + if err == nil { + return c, nil + } + if firstErr == nil { + firstErr = err + } + } + if firstErr == nil { + firstErr = &net.OpError{Op: "dial", Err: errMissingAddress} + } + return nil, firstErr +} + +func (n *Net2) Dial(network, address string) (net.Conn, error) { + return n.DialContext(context.Background(), network, address) +} diff --git a/tun/netstack/lneto_test.go b/tun/netstack/lneto_test.go new file mode 100644 index 000000000..b3c0af4a8 --- /dev/null +++ b/tun/netstack/lneto_test.go @@ -0,0 +1,300 @@ +/* SPDX-License-Identifier: MIT + * + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. + */ + +package netstack + +import ( + "context" + "errors" + "net" + "net/netip" + "os" + "sync" + "testing" + "time" + + "golang.zx2c4.com/wireguard/tun" +) + +// TestNet2_Construct is a regression test: the previous CreateNetTUN2 omitted +// TCPPoolConfig.NewBackoff, which made StackGo panic at construction. +func TestNet2_Construct(t *testing.T) { + dev, net2, err := CreateNetTUN2( + []netip.Addr{netip.MustParseAddr("10.0.0.1")}, + []netip.Addr{netip.MustParseAddr("8.8.8.8")}, + 1500, + ) + if err != nil { + t.Fatal(err) + } + defer dev.Close() + if net2 == nil { + t.Fatal("nil Net2") + } + select { + case ev := <-dev.Events(): + if ev != tun.EventUp { + t.Fatalf("want EventUp, got %v", ev) + } + case <-time.After(time.Second): + t.Fatal("no EventUp event") + } +} + +// TestNet2_CloseRace stresses the Read/Write/Close interaction that previously +// panicked with "send on closed channel". Run with -race. +func TestNet2_CloseRace(t *testing.T) { + for i := 0; i < 50; i++ { + dev, _, err := CreateNetTUN2( + []netip.Addr{netip.MustParseAddr("10.0.0.1")}, + nil, 1500, + ) + if err != nil { + t.Fatal(err) + } + <-dev.Events() // drain EventUp + + var wg sync.WaitGroup + // Reader: must return os.ErrClosed once Close fires, never panic. + wg.Add(1) + go func() { + defer wg.Done() + bufs := [][]byte{make([]byte, 2048)} + sizes := []int{0} + for { + _, err := dev.Read(bufs, sizes, 0) + if errors.Is(err, os.ErrClosed) { + return + } + } + }() + // Writer: feed junk ingress concurrently with Close. + wg.Add(1) + go func() { + defer wg.Done() + pkt := make([]byte, 40) + pkt[0] = 0x45 // IPv4, IHL 5 + for j := 0; j < 1000; j++ { + if _, err := dev.Write([][]byte{pkt}, 0); err != nil { + return + } + } + }() + + time.Sleep(time.Millisecond) + if err := dev.Close(); err != nil { + t.Fatal(err) + } + // Double close must be safe (no panic). + if err := dev.Close(); err != nil { + t.Fatal(err) + } + wg.Wait() + } +} + +// TestNet2_ListenTCPPort0 covers gap D: listening on port 0 must auto-assign an +// ephemeral port instead of failing (the library previously returned ErrZeroSource). +func TestNet2_ListenTCPPort0(t *testing.T) { + dev, net2, err := CreateNetTUN2([]netip.Addr{netip.MustParseAddr("10.0.0.1")}, nil, 1500) + if err != nil { + t.Fatal(err) + } + defer dev.Close() + <-dev.Events() + + ln, err := net2.ListenTCPAddrPort(netip.AddrPort{}) + if err != nil { + t.Fatal("listen on port 0:", err) + } + defer ln.Close() + if a, ok := ln.Addr().(*net.TCPAddr); !ok || a.Port == 0 { + t.Fatalf("expected an auto-assigned ephemeral port, got %v", ln.Addr()) + } +} + +// TestNet2_UDPEcho wires two Net2 instances back-to-back and performs a connected +// UDP dial against a UDP PacketConn listener, exercising the UDP socket paths. +func TestNet2_UDPEcho(t *testing.T) { + const ( + addrA = "10.0.0.1" + addrB = "10.0.0.2" + port = 9999 + ) + devA, netA, err := CreateNetTUN2([]netip.Addr{netip.MustParseAddr(addrA)}, nil, 1500) + if err != nil { + t.Fatal(err) + } + devB, netB, err := CreateNetTUN2([]netip.Addr{netip.MustParseAddr(addrB)}, nil, 1500) + if err != nil { + t.Fatal(err) + } + <-devA.Events() + <-devB.Events() + + var pumps sync.WaitGroup + pumps.Add(2) + go func() { defer pumps.Done(); pump(devA, devB) }() + go func() { defer pumps.Done(); pump(devB, devA) }() + defer func() { + devA.Close() + devB.Close() + pumps.Wait() + }() + + srv, err := netB.ListenUDPAddrPort(netip.AddrPortFrom(netip.MustParseAddr(addrB), port)) + if err != nil { + t.Fatal("listen udp:", err) + } + defer srv.Close() + + const msg = "hello over lneto udp" + srvDone := make(chan error, 1) + go func() { + buf := make([]byte, 512) + srv.SetDeadline(time.Now().Add(5 * time.Second)) + n, from, err := srv.ReadFrom(buf) + if err != nil { + srvDone <- err + return + } + _, err = srv.WriteTo(buf[:n], from) // echo back to sender + srvDone <- err + }() + + cli, err := netA.DialUDPAddrPort(netip.AddrPort{}, netip.AddrPortFrom(netip.MustParseAddr(addrB), port)) + if err != nil { + t.Fatal("dial udp:", err) + } + defer cli.Close() + cli.SetDeadline(time.Now().Add(5 * time.Second)) + if _, err := cli.Write([]byte(msg)); err != nil { + t.Fatal("udp write:", err) + } + echo := make([]byte, 512) + n, err := cli.Read(echo) + if err != nil { + t.Fatal("udp read echo:", err) + } + if string(echo[:n]) != msg { + t.Fatalf("udp echo mismatch: want %q got %q", msg, string(echo[:n])) + } + if err := <-srvDone; err != nil { + t.Fatal("udp server:", err) + } +} + +// pump forwards IP frames produced by src into dst until src is closed. +func pump(src, dst tun.Device) { + bufs := [][]byte{make([]byte, 2048)} + sizes := []int{0} + for { + n, err := src.Read(bufs, sizes, 0) + if err != nil { + return + } + if n == 0 || sizes[0] == 0 { + continue + } + out := make([]byte, sizes[0]) + copy(out, bufs[0][:sizes[0]]) + if _, err := dst.Write([][]byte{out}, 0); err != nil { + return + } + } +} + +// TestNet2_TCPEcho wires two Net2 instances back-to-back (one's egress is the +// other's ingress) and performs a TCP dial + echo, exercising the egress poll, +// the dial path, and the listener/accept path end-to-end for both IPv4 and IPv6. +func TestNet2_TCPEcho(t *testing.T) { + t.Run("ipv4", func(t *testing.T) { testTCPEcho(t, "10.0.0.1", "10.0.0.2") }) + t.Run("ipv6", func(t *testing.T) { testTCPEcho(t, "fd00::1", "fd00::2") }) +} + +func testTCPEcho(t *testing.T, addrA, addrB string) { + const port = 1234 + devA, netA, err := CreateNetTUN2([]netip.Addr{netip.MustParseAddr(addrA)}, nil, 1500) + if err != nil { + t.Fatal(err) + } + devB, netB, err := CreateNetTUN2([]netip.Addr{netip.MustParseAddr(addrB)}, nil, 1500) + if err != nil { + t.Fatal(err) + } + <-devA.Events() + <-devB.Events() + + var pumps sync.WaitGroup + pumps.Add(2) + go func() { defer pumps.Done(); pump(devA, devB) }() + go func() { defer pumps.Done(); pump(devB, devA) }() + + ln, err := netB.ListenTCPAddrPort(netip.AddrPortFrom(netip.MustParseAddr(addrB), port)) + if err != nil { + t.Fatal("listen:", err) + } + // Teardown order matters: stop ingress (close devices → pumps exit) BEFORE + // closing the listener, since tcp.Listener.Close is not synchronized against + // the stack's ingress demux in the lneto library (see review notes). + defer func() { + devA.Close() + devB.Close() + pumps.Wait() + ln.Close() + }() + + const msg = "hello over lneto tcp" + srvDone := make(chan error, 1) + go func() { + conn, err := ln.Accept() + if err != nil { + srvDone <- err + return + } + defer conn.Close() + buf := make([]byte, len(msg)) + conn.SetDeadline(time.Now().Add(5 * time.Second)) + var got int + for got < len(msg) { + n, err := conn.Read(buf[got:]) + if err != nil { + srvDone <- err + return + } + got += n + } + _, err = conn.Write(buf[:got]) // echo back + srvDone <- err + }() + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + conn, err := netA.DialContextTCPAddrPort(ctx, netip.AddrPortFrom(netip.MustParseAddr(addrB), port)) + if err != nil { + t.Fatal("dial:", err) + } + defer conn.Close() + + conn.SetDeadline(time.Now().Add(5 * time.Second)) + if _, err := conn.Write([]byte(msg)); err != nil { + t.Fatal("write:", err) + } + echo := make([]byte, len(msg)) + got := 0 + for got < len(msg) { + n, err := conn.Read(echo[got:]) + if err != nil { + t.Fatal("read echo:", err) + } + got += n + } + if string(echo) != msg { + t.Fatalf("echo mismatch: want %q got %q", msg, string(echo)) + } + if err := <-srvDone; err != nil { + t.Fatal("server:", err) + } +} From 3c5a84dbf47f66707c0f294360cafa6a7c99b189 Mon Sep 17 00:00:00 2001 From: Patricio Whittingslow Date: Tue, 16 Jun 2026 09:58:42 -0300 Subject: [PATCH 02/12] rename gvisor --- device/pools_test.go | 2 +- tun/netstack/{tun.go => gvisor.go} | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) rename tun/netstack/{tun.go => gvisor.go} (100%) diff --git a/device/pools_test.go b/device/pools_test.go index 8381d5a6f..3755e7b13 100644 --- a/device/pools_test.go +++ b/device/pools_test.go @@ -64,7 +64,7 @@ func TestWaitPool(t *testing.T) { } wg.Wait() if max.Load() != p.max { - t.Errorf("Actual maximum count (%d) != ideal maximum count (%d)", max, p.max) + t.Errorf("Actual maximum count (%d) != ideal maximum count (%d)", max.Load(), p.max) } } diff --git a/tun/netstack/tun.go b/tun/netstack/gvisor.go similarity index 100% rename from tun/netstack/tun.go rename to tun/netstack/gvisor.go index a7aec9e82..7b68dc983 100644 --- a/tun/netstack/tun.go +++ b/tun/netstack/gvisor.go @@ -39,6 +39,8 @@ import ( "gvisor.dev/gvisor/pkg/waiter" ) +type Net netTun + type netTun struct { ep *channel.Endpoint stack *stack.Stack @@ -50,8 +52,6 @@ type netTun struct { hasV4, hasV6 bool } -type Net netTun - func CreateNetTUN(localAddresses, dnsServers []netip.Addr, mtu int) (tun.Device, *Net, error) { opts := stack.Options{ NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol}, From 54fbc16a121288aa5c0bc110b6fe8e8f03e55dcc Mon Sep 17 00:00:00 2001 From: Patricio Whittingslow Date: Tue, 16 Jun 2026 10:45:37 -0300 Subject: [PATCH 03/12] modularize lneto/gvisor stacks --- tun/netstack/gvisor.go | 295 +++-------------------- tun/netstack/lneto.go | 476 +++++++++++-------------------------- tun/netstack/lneto_test.go | 18 +- tun/netstack/net.go | 369 ++++++++++++++++++++++++++++ 4 files changed, 548 insertions(+), 610 deletions(-) create mode 100644 tun/netstack/net.go diff --git a/tun/netstack/gvisor.go b/tun/netstack/gvisor.go index 7b68dc983..9c255a6ce 100644 --- a/tun/netstack/gvisor.go +++ b/tun/netstack/gvisor.go @@ -16,8 +16,6 @@ import ( "net" "net/netip" "os" - "regexp" - "strconv" "strings" "syscall" "time" @@ -39,8 +37,6 @@ import ( "gvisor.dev/gvisor/pkg/waiter" ) -type Net netTun - type netTun struct { ep *channel.Endpoint stack *stack.Stack @@ -105,9 +101,11 @@ func CreateNetTUN(localAddresses, dnsServers []netip.Addr, mtu int) (tun.Device, } dev.events <- tun.EventUp - return dev, (*Net)(dev), nil + return dev, &Net{stack: dev}, nil } +var _ Stack = (*netTun)(nil) + func (tun *netTun) Name() (string, error) { return "go", nil } @@ -205,46 +203,34 @@ func convertToFullAddr(endpoint netip.AddrPort) (tcpip.FullAddress, tcpip.Networ }, protoNumber } -func (net *Net) DialContextTCPAddrPort(ctx context.Context, addr netip.AddrPort) (*gonet.TCPConn, error) { +func (tun *netTun) DialContextTCPAddrPort(ctx context.Context, addr netip.AddrPort) (TCPConn, error) { fa, pn := convertToFullAddr(addr) - return gonet.DialContextTCP(ctx, net.stack, fa, pn) -} - -func (net *Net) DialContextTCP(ctx context.Context, addr *net.TCPAddr) (*gonet.TCPConn, error) { - if addr == nil { - return net.DialContextTCPAddrPort(ctx, netip.AddrPort{}) + c, err := gonet.DialContextTCP(ctx, tun.stack, fa, pn) + if err != nil { + return nil, err } - ip, _ := netip.AddrFromSlice(addr.IP) - return net.DialContextTCPAddrPort(ctx, netip.AddrPortFrom(ip, uint16(addr.Port))) + return c, nil } -func (net *Net) DialTCPAddrPort(addr netip.AddrPort) (*gonet.TCPConn, error) { +func (tun *netTun) DialTCPAddrPort(addr netip.AddrPort) (TCPConn, error) { fa, pn := convertToFullAddr(addr) - return gonet.DialTCP(net.stack, fa, pn) -} - -func (net *Net) DialTCP(addr *net.TCPAddr) (*gonet.TCPConn, error) { - if addr == nil { - return net.DialTCPAddrPort(netip.AddrPort{}) + c, err := gonet.DialTCP(tun.stack, fa, pn) + if err != nil { + return nil, err } - ip, _ := netip.AddrFromSlice(addr.IP) - return net.DialTCPAddrPort(netip.AddrPortFrom(ip, uint16(addr.Port))) + return c, nil } -func (net *Net) ListenTCPAddrPort(addr netip.AddrPort) (*gonet.TCPListener, error) { +func (tun *netTun) ListenTCPAddrPort(addr netip.AddrPort) (TCPListener, error) { fa, pn := convertToFullAddr(addr) - return gonet.ListenTCP(net.stack, fa, pn) -} - -func (net *Net) ListenTCP(addr *net.TCPAddr) (*gonet.TCPListener, error) { - if addr == nil { - return net.ListenTCPAddrPort(netip.AddrPort{}) + l, err := gonet.ListenTCP(tun.stack, fa, pn) + if err != nil { + return nil, err } - ip, _ := netip.AddrFromSlice(addr.IP) - return net.ListenTCPAddrPort(netip.AddrPortFrom(ip, uint16(addr.Port))) + return l, nil } -func (net *Net) DialUDPAddrPort(laddr, raddr netip.AddrPort) (*gonet.UDPConn, error) { +func (tun *netTun) DialUDPAddrPort(laddr, raddr netip.AddrPort) (UDPConn, error) { var lfa, rfa *tcpip.FullAddress var pn tcpip.NetworkProtocolNumber if laddr.IsValid() || laddr.Port() > 0 { @@ -257,28 +243,15 @@ func (net *Net) DialUDPAddrPort(laddr, raddr netip.AddrPort) (*gonet.UDPConn, er addr, pn = convertToFullAddr(raddr) rfa = &addr } - return gonet.DialUDP(net.stack, lfa, rfa, pn) -} - -func (net *Net) ListenUDPAddrPort(laddr netip.AddrPort) (*gonet.UDPConn, error) { - return net.DialUDPAddrPort(laddr, netip.AddrPort{}) -} - -func (net *Net) DialUDP(laddr, raddr *net.UDPAddr) (*gonet.UDPConn, error) { - var la, ra netip.AddrPort - if laddr != nil { - ip, _ := netip.AddrFromSlice(laddr.IP) - la = netip.AddrPortFrom(ip, uint16(laddr.Port)) - } - if raddr != nil { - ip, _ := netip.AddrFromSlice(raddr.IP) - ra = netip.AddrPortFrom(ip, uint16(raddr.Port)) + c, err := gonet.DialUDP(tun.stack, lfa, rfa, pn) + if err != nil { + return nil, err } - return net.DialUDPAddrPort(la, ra) + return c, nil } -func (net *Net) ListenUDP(laddr *net.UDPAddr) (*gonet.UDPConn, error) { - return net.DialUDP(laddr, nil) +func (tun *netTun) ListenUDPAddrPort(laddr netip.AddrPort) (UDPConn, error) { + return tun.DialUDPAddrPort(laddr, netip.AddrPort{}) } type PingConn struct { @@ -312,7 +285,7 @@ func PingAddrFromAddr(addr netip.Addr) *PingAddr { return &PingAddr{addr} } -func (net *Net) DialPingAddr(laddr, raddr netip.Addr) (*PingConn, error) { +func (tun *netTun) DialPingAddr(laddr, raddr netip.Addr) (*PingConn, error) { if !laddr.IsValid() && !raddr.IsValid() { return nil, errors.New("ping dial: invalid address") } @@ -339,7 +312,7 @@ func (net *Net) DialPingAddr(laddr, raddr netip.Addr) (*PingConn, error) { } pc.deadline.Stop() - ep, tcpipErr := net.stack.NewEndpoint(tn, pn, &pc.wq) + ep, tcpipErr := tun.stack.NewEndpoint(tn, pn, &pc.wq) if tcpipErr != nil { return nil, fmt.Errorf("ping socket: endpoint: %s", tcpipErr) } @@ -363,27 +336,8 @@ func (net *Net) DialPingAddr(laddr, raddr netip.Addr) (*PingConn, error) { return pc, nil } -func (net *Net) ListenPingAddr(laddr netip.Addr) (*PingConn, error) { - return net.DialPingAddr(laddr, netip.Addr{}) -} - -func (net *Net) DialPing(laddr, raddr *PingAddr) (*PingConn, error) { - var la, ra netip.Addr - if laddr != nil { - la = laddr.addr - } - if raddr != nil { - ra = raddr.addr - } - return net.DialPingAddr(la, ra) -} - -func (net *Net) ListenPing(laddr *PingAddr) (*PingConn, error) { - var la netip.Addr - if laddr != nil { - la = laddr.addr - } - return net.ListenPingAddr(la) +func (tun *netTun) ListenPingAddr(laddr netip.Addr) (*PingConn, error) { + return tun.DialPingAddr(laddr, netip.Addr{}) } func (pc *PingConn) LocalAddr() net.Addr { @@ -475,67 +429,6 @@ func (pc *PingConn) SetReadDeadline(t time.Time) error { return nil } -var ( - errNoSuchHost = errors.New("no such host") - errLameReferral = errors.New("lame referral") - errCannotUnmarshalDNSMessage = errors.New("cannot unmarshal DNS message") - errCannotMarshalDNSMessage = errors.New("cannot marshal DNS message") - errServerMisbehaving = errors.New("server misbehaving") - errInvalidDNSResponse = errors.New("invalid DNS response") - errNoAnswerFromDNSServer = errors.New("no answer from DNS server") - errServerTemporarilyMisbehaving = errors.New("server misbehaving") - errCanceled = errors.New("operation was canceled") - errTimeout = errors.New("i/o timeout") - errNumericPort = errors.New("port must be numeric") - errNoSuitableAddress = errors.New("no suitable address found") - errMissingAddress = errors.New("missing address") -) - -func (net *Net) LookupHost(host string) (addrs []string, err error) { - return net.LookupContextHost(context.Background(), host) -} - -func isDomainName(s string) bool { - l := len(s) - if l == 0 || l > 254 || l == 254 && s[l-1] != '.' { - return false - } - last := byte('.') - nonNumeric := false - partlen := 0 - for i := 0; i < len(s); i++ { - c := s[i] - switch { - default: - return false - case 'a' <= c && c <= 'z' || 'A' <= c && c <= 'Z' || c == '_': - nonNumeric = true - partlen++ - case '0' <= c && c <= '9': - partlen++ - case c == '-': - if last == '.' { - return false - } - partlen++ - nonNumeric = true - case c == '.': - if last == '.' || last == '-' { - return false - } - if partlen > 63 || partlen == 0 { - return false - } - partlen = 0 - } - last = c - } - if last == '-' || partlen > 63 { - return false - } - return nonNumeric -} - func randU16() uint16 { var b [2]byte _, err := rand.Read(b[:]) @@ -650,7 +543,7 @@ func dnsStreamRoundTrip(c net.Conn, id uint16, query dnsmessage.Question, b []by return p, h, nil } -func (tnet *Net) exchange(ctx context.Context, server netip.Addr, q dnsmessage.Question, timeout time.Duration) (dnsmessage.Parser, dnsmessage.Header, error) { +func (tnet *netTun) exchange(ctx context.Context, server netip.Addr, q dnsmessage.Question, timeout time.Duration) (dnsmessage.Parser, dnsmessage.Header, error) { q.Class = dnsmessage.ClassINET id, udpReq, tcpReq, err := newRequest(q) if err != nil { @@ -743,7 +636,7 @@ func skipToAnswer(p *dnsmessage.Parser, qtype dnsmessage.Type) error { } } -func (tnet *Net) tryOneName(ctx context.Context, name string, qtype dnsmessage.Type) (dnsmessage.Parser, string, error) { +func (tnet *netTun) tryOneName(ctx context.Context, name string, qtype dnsmessage.Type) (dnsmessage.Parser, string, error) { var lastErr error n, err := dnsmessage.NewName(name) @@ -810,7 +703,7 @@ func (tnet *Net) tryOneName(ctx context.Context, name string, qtype dnsmessage.T return dnsmessage.Parser{}, "", lastErr } -func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]string, error) { +func (tnet *netTun) LookupContextHost(ctx context.Context, host string) ([]string, error) { if host == "" || (!tnet.hasV6 && !tnet.hasV4) { return nil, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true} } @@ -931,127 +824,3 @@ func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]string, } return saddrs, nil } - -func partialDeadline(now, deadline time.Time, addrsRemaining int) (time.Time, error) { - if deadline.IsZero() { - return deadline, nil - } - timeRemaining := deadline.Sub(now) - if timeRemaining <= 0 { - return time.Time{}, errTimeout - } - timeout := timeRemaining / time.Duration(addrsRemaining) - const saneMinimum = 2 * time.Second - if timeout < saneMinimum { - if timeRemaining < saneMinimum { - timeout = timeRemaining - } else { - timeout = saneMinimum - } - } - return now.Add(timeout), nil -} - -var protoSplitter = regexp.MustCompile(`^(tcp|udp|ping)(4|6)?$`) - -func (tnet *Net) DialContext(ctx context.Context, network, address string) (net.Conn, error) { - if ctx == nil { - panic("nil context") - } - var acceptV4, acceptV6 bool - matches := protoSplitter.FindStringSubmatch(network) - if matches == nil { - return nil, &net.OpError{Op: "dial", Err: net.UnknownNetworkError(network)} - } else if len(matches[2]) == 0 { - acceptV4 = true - acceptV6 = true - } else { - acceptV4 = matches[2][0] == '4' - acceptV6 = !acceptV4 - } - var host string - var port int - if matches[1] == "ping" { - host = address - } else { - var sport string - var err error - host, sport, err = net.SplitHostPort(address) - if err != nil { - return nil, &net.OpError{Op: "dial", Err: err} - } - port, err = strconv.Atoi(sport) - if err != nil || port < 0 || port > 65535 { - return nil, &net.OpError{Op: "dial", Err: errNumericPort} - } - } - allAddr, err := tnet.LookupContextHost(ctx, host) - if err != nil { - return nil, &net.OpError{Op: "dial", Err: err} - } - var addrs []netip.AddrPort - for _, addr := range allAddr { - ip, err := netip.ParseAddr(addr) - if err == nil && ((ip.Is4() && acceptV4) || (ip.Is6() && acceptV6)) { - addrs = append(addrs, netip.AddrPortFrom(ip, uint16(port))) - } - } - if len(addrs) == 0 && len(allAddr) != 0 { - return nil, &net.OpError{Op: "dial", Err: errNoSuitableAddress} - } - - var firstErr error - for i, addr := range addrs { - select { - case <-ctx.Done(): - err := ctx.Err() - if err == context.Canceled { - err = errCanceled - } else if err == context.DeadlineExceeded { - err = errTimeout - } - return nil, &net.OpError{Op: "dial", Err: err} - default: - } - - dialCtx := ctx - if deadline, hasDeadline := ctx.Deadline(); hasDeadline { - partialDeadline, err := partialDeadline(time.Now(), deadline, len(addrs)-i) - if err != nil { - if firstErr == nil { - firstErr = &net.OpError{Op: "dial", Err: err} - } - break - } - if partialDeadline.Before(deadline) { - var cancel context.CancelFunc - dialCtx, cancel = context.WithDeadline(ctx, partialDeadline) - defer cancel() - } - } - - var c net.Conn - switch matches[1] { - case "tcp": - c, err = tnet.DialContextTCPAddrPort(dialCtx, addr) - case "udp": - c, err = tnet.DialUDPAddrPort(netip.AddrPort{}, addr) - case "ping": - c, err = tnet.DialPingAddr(netip.Addr{}, addr.Addr()) - } - if err == nil { - return c, nil - } - if firstErr == nil { - firstErr = err - } - } - if firstErr == nil { - firstErr = &net.OpError{Op: "dial", Err: errMissingAddress} - } - return nil, firstErr -} - -func (tnet *Net) Dial(network, address string) (net.Conn, error) { - return tnet.DialContext(context.Background(), network, address) -} diff --git a/tun/netstack/lneto.go b/tun/netstack/lneto.go index 6e7ff54b0..3721deaca 100644 --- a/tun/netstack/lneto.go +++ b/tun/netstack/lneto.go @@ -14,9 +14,7 @@ import ( "net" "net/netip" "os" - "regexp" "runtime" - "strconv" "strings" "sync" "syscall" @@ -28,26 +26,134 @@ import ( "golang.zx2c4.com/wireguard/tun" ) -// Net2 is a lneto-backed userspace network stack that implements both +func CreateNetTUNLneto(localAddresses, dnsServers []netip.Addr, mtu int) (tun.Device, *Net, error) { + if mtu <= 0 { + mtu = 1500 + } + dev := &lnetoStack{ + events: make(chan tun.Event, 10), + closed: make(chan struct{}), + mtu: mtu, + dnsServers: dnsServers, + } + + var hwAddr [6]byte + if _, err := crand.Read(hwAddr[:]); err != nil { + return nil, nil, fmt.Errorf("CreateNetTUNLneto: rand MAC: %w", err) + } + hwAddr[0] &^= 0x01 // unicast + hwAddr[0] |= 0x02 // locally administered + + // Pick the first IPv4 and first IPv6 local address; track which families are + // configured (mirrors gVisor's hasV4/hasV6). + var staticAddr4 [4]byte + var staticAddr6 [16]byte + for _, addr := range localAddresses { + if addr.Is4() && !dev.hasV4 { + staticAddr4 = addr.As4() + dev.hasV4 = true + } else if addr.Is6() && !dev.hasV6 { + staticAddr6 = addr.As16() + dev.hasV6 = true + } + } + + // DNS server: prefer IPv4 (the stack's lookup path is currently IPv4-only), + // fall back to the first IPv6 server. + var dnsServer netip.Addr + for _, d := range dnsServers { + if d.Is4() { + dnsServer = d + break + } + } + if !dnsServer.IsValid() { + for _, d := range dnsServers { + if d.Is6() { + dnsServer = d + break + } + } + } + + var randSeed int64 + if err := binary.Read(crand.Reader, binary.LittleEndian, &randSeed); err != nil { + return nil, nil, fmt.Errorf("CreateNetTUNLneto: rand seed: %w", err) + } + + cfg := xnet.StackConfig{ + HardwareAddress: hwAddr, + StaticAddress4: staticAddr4, + MTU: uint16(mtu), + Hostname: "wg0", + RandSeed: randSeed, + PassivePeers: 0, // no ARP passive learning needed for TUN + ICMPQueueLimit: 4, + MaxActiveTCPPorts: 256, + MaxActiveUDPPorts: 256, + DNSServer: dnsServer, + } + if dev.hasV6 { + cfg.StaticAddress6 = staticAddr6 + cfg.IPv6Stack = xnet.DefaultStack6() + } + // NOTE: ICMP is intentionally NOT enabled here. On a TUN there is no link layer, + // so MAC resolution must be skipped: IPv4 ARP is gated off by leaving the subnet + // unset, and IPv6 NDP is gated off by leaving ICMPv6 unregistered. Enabling ICMP + // would make DialTCP6/DialUDP6 emit Neighbor Solicitations that are never answered + // on a TUN, breaking IPv6 dialing. Ping support (which needs ICMP) is a separate + // follow-up that must reconcile this. + if err := dev.sa.Reset(cfg); err != nil { + return nil, nil, fmt.Errorf("CreateNetTUNLneto: stack reset: %w", err) + } + + // GOMAXPROCS-aware backoff: on a single-threaded runtime only yield cooperatively + // (sleeping the poll would starve egress), otherwise sleep with exponential backoff. + baseStack := defaultStackBackoff + newTCPBackoff := func() lneto.BackoffStrategy { return defaultTCPBackoff } + if runtime.GOMAXPROCS(0) == 1 { + baseStack = backoffYield + newTCPBackoff = func() lneto.BackoffStrategy { return backoffYield } + } + irq, backoff := interruptBackoff(baseStack) + dev.backoff = backoff + dev.backoffirq = irq + + dev.sgo = dev.sa.StackGo(backoff, xnet.StackGoConfig{ + ListenerPoolConfig: xnet.TCPPoolConfig{ + PoolSize: 256, + QueueSize: 8, + TxBufSize: 32 << 10, + RxBufSize: 32 << 10, + EstablishedTimeout: 30 * time.Second, + ClosingTimeout: 10 * time.Second, + NewBackoff: newTCPBackoff, // required: StackGo panics if nil. + }, + }) + dev.events <- tun.EventUp + return dev, &Net{stack: dev}, nil +} + +// lnetoStack is a lneto-backed userspace network stack that implements both // [tun.Device] (for WireGuard integration) and a networking API (Dial/Listen/DNS). // // Packet flow: -// - Ingress (WireGuard → stack): [Net2.Write] calls [xnet.StackAsync.IngressIP], -// then pokes the backoff irq so a blocked [Net2.Read] wakes immediately. -// - Egress (stack → WireGuard): [Net2.Read] polls [xnet.StackAsync.EgressIP] +// - Ingress (WireGuard → stack): [lnetoStack.Write] calls [xnet.StackAsync.IngressIP], +// then pokes the backoff irq so a blocked [lnetoStack.Read] wakes immediately. +// - Egress (stack → WireGuard): [lnetoStack.Read] polls [xnet.StackAsync.EgressIP] // directly, sleeping on an interruptible backoff between empty polls. // // Unlike gVisor's channel.Endpoint there is no native egress notification hook, so // egress is driven by Read's poll loop. The backoff is interruptible (woken by -// [Net2.interrupt]) and GOMAXPROCS-aware: on a single-threaded runtime it only +// [lnetoStack.interrupt]) and GOMAXPROCS-aware: on a single-threaded runtime it only // yields cooperatively (sleeping would starve the poll), otherwise it sleeps with // exponential backoff. This mirrors the proven go-net design. // -// Lifecycle: created by [CreateNetTUN2], torn down by [Net2.Close] which closes +// Lifecycle: created by [CreateNetTUNLneto], torn down by [lnetoStack.Close] which closes // the closed channel (unblocking Read and any blocking socket op) and events. -type Net2 struct { +type lnetoStack struct { sa xnet.StackAsync - sgo xnet.StackGo // wraps sa; created once in CreateNetTUN2 + sgo xnet.StackGo // wraps sa; created once in CreateNetTUNLneto // events carries TUN state changes (e.g. EventUp) consumed by WireGuard's device loop. events chan tun.Event @@ -67,39 +173,6 @@ type Net2 struct { hasV4, hasV6 bool } -type TCPConn interface { - Close() error - CloseRead() error - CloseWrite() error - LocalAddr() net.Addr - Read(b []byte) (int, error) - RemoteAddr() net.Addr - SetDeadline(t time.Time) error - SetReadDeadline(t time.Time) error - SetWriteDeadline(t time.Time) error - Write(b []byte) (int, error) -} - -type UDPConn interface { - Close() error - LocalAddr() net.Addr - Read(b []byte) (int, error) - ReadFrom(b []byte) (int, net.Addr, error) - RemoteAddr() net.Addr - SetDeadline(t time.Time) error - SetReadDeadline(t time.Time) error - SetWriteDeadline(t time.Time) error - Write(b []byte) (int, error) - WriteTo(b []byte, addr net.Addr) (int, error) -} - -type TCPListener interface { - Accept() (net.Conn, error) - Addr() net.Addr - Close() error - Shutdown() -} - type event struct{} // interruptBackoff wraps a [lneto.BackoffStrategy] so its sleep can be cut short by @@ -173,7 +246,7 @@ func backoffYield(consecutiveBackoffs uint) time.Duration { // interrupt wakes one sleeper blocked in the interruptible backoff (Read's poll loop // or a blocking socket op). Non-blocking and safe to call from any goroutine. -func (n *Net2) interrupt() { +func (n *lnetoStack) interrupt() { select { case n.backoffirq <- event{}: default: @@ -182,14 +255,14 @@ func (n *Net2) interrupt() { // --- tun.Device implementation --- -func (n *Net2) Name() (string, error) { return "go2", nil } -func (n *Net2) File() *os.File { return nil } -func (n *Net2) Events() <-chan tun.Event { return n.events } -func (n *Net2) MTU() (int, error) { return n.mtu, nil } -func (n *Net2) BatchSize() int { return 1 } +func (n *lnetoStack) Name() (string, error) { return "go2", nil } +func (n *lnetoStack) File() *os.File { return nil } +func (n *lnetoStack) Events() <-chan tun.Event { return n.events } +func (n *lnetoStack) MTU() (int, error) { return n.mtu, nil } +func (n *lnetoStack) BatchSize() int { return 1 } // Write feeds incoming IP packets (WireGuard → stack) into the lneto stack. -func (n *Net2) Write(bufs [][]byte, offset int) (int, error) { +func (n *lnetoStack) Write(bufs [][]byte, offset int) (int, error) { wrote := false for _, buf := range bufs { if pkt := buf[offset:]; len(pkt) > 0 { @@ -207,7 +280,7 @@ func (n *Net2) Write(bufs [][]byte, offset int) (int, error) { // EgressIP and sleeping on the interruptible backoff between empty polls. It writes // directly into the caller's buffer (no intermediate copy). Returns os.ErrClosed once // Close has been called. -func (n *Net2) Read(bufs [][]byte, sizes []int, offset int) (int, error) { +func (n *lnetoStack) Read(bufs [][]byte, sizes []int, offset int) (int, error) { dst := bufs[0][offset:] var backoffs uint for { @@ -226,7 +299,7 @@ func (n *Net2) Read(bufs [][]byte, sizes []int, offset int) (int, error) { } } -func (n *Net2) Close() error { +func (n *lnetoStack) Close() error { n.closeOnce.Do(func() { close(n.closed) // unblock Read. n.interrupt() // wake a backoff sleeper so it observes closed promptly. @@ -235,113 +308,7 @@ func (n *Net2) Close() error { return nil } -func CreateNetTUN2(localAddresses, dnsServers []netip.Addr, mtu int) (tun.Device, *Net2, error) { - if mtu <= 0 { - mtu = 1500 - } - dev := &Net2{ - events: make(chan tun.Event, 10), - closed: make(chan struct{}), - mtu: mtu, - dnsServers: dnsServers, - } - - var hwAddr [6]byte - if _, err := crand.Read(hwAddr[:]); err != nil { - return nil, nil, fmt.Errorf("CreateNetTUN2: rand MAC: %w", err) - } - hwAddr[0] &^= 0x01 // unicast - hwAddr[0] |= 0x02 // locally administered - - // Pick the first IPv4 and first IPv6 local address; track which families are - // configured (mirrors gVisor's hasV4/hasV6). - var staticAddr4 [4]byte - var staticAddr6 [16]byte - for _, addr := range localAddresses { - if addr.Is4() && !dev.hasV4 { - staticAddr4 = addr.As4() - dev.hasV4 = true - } else if addr.Is6() && !dev.hasV6 { - staticAddr6 = addr.As16() - dev.hasV6 = true - } - } - - // DNS server: prefer IPv4 (the stack's lookup path is currently IPv4-only), - // fall back to the first IPv6 server. - var dnsServer netip.Addr - for _, d := range dnsServers { - if d.Is4() { - dnsServer = d - break - } - } - if !dnsServer.IsValid() { - for _, d := range dnsServers { - if d.Is6() { - dnsServer = d - break - } - } - } - - var randSeed int64 - if err := binary.Read(crand.Reader, binary.LittleEndian, &randSeed); err != nil { - return nil, nil, fmt.Errorf("CreateNetTUN2: rand seed: %w", err) - } - - cfg := xnet.StackConfig{ - HardwareAddress: hwAddr, - StaticAddress4: staticAddr4, - MTU: uint16(mtu), - Hostname: "wg0", - RandSeed: randSeed, - PassivePeers: 0, // no ARP passive learning needed for TUN - ICMPQueueLimit: 4, - MaxActiveTCPPorts: 256, - MaxActiveUDPPorts: 256, - DNSServer: dnsServer, - } - if dev.hasV6 { - cfg.StaticAddress6 = staticAddr6 - cfg.IPv6Stack = xnet.DefaultStack6() - } - // NOTE: ICMP is intentionally NOT enabled here. On a TUN there is no link layer, - // so MAC resolution must be skipped: IPv4 ARP is gated off by leaving the subnet - // unset, and IPv6 NDP is gated off by leaving ICMPv6 unregistered. Enabling ICMP - // would make DialTCP6/DialUDP6 emit Neighbor Solicitations that are never answered - // on a TUN, breaking IPv6 dialing. Ping support (which needs ICMP) is a separate - // follow-up that must reconcile this. - if err := dev.sa.Reset(cfg); err != nil { - return nil, nil, fmt.Errorf("CreateNetTUN2: stack reset: %w", err) - } - - // GOMAXPROCS-aware backoff: on a single-threaded runtime only yield cooperatively - // (sleeping the poll would starve egress), otherwise sleep with exponential backoff. - baseStack := defaultStackBackoff - newTCPBackoff := func() lneto.BackoffStrategy { return defaultTCPBackoff } - if runtime.GOMAXPROCS(0) == 1 { - baseStack = backoffYield - newTCPBackoff = func() lneto.BackoffStrategy { return backoffYield } - } - irq, backoff := interruptBackoff(baseStack) - dev.backoff = backoff - dev.backoffirq = irq - - dev.sgo = dev.sa.StackGo(backoff, xnet.StackGoConfig{ - ListenerPoolConfig: xnet.TCPPoolConfig{ - PoolSize: 256, - QueueSize: 8, - TxBufSize: 32 << 10, - RxBufSize: 32 << 10, - EstablishedTimeout: 30 * time.Second, - ClosingTimeout: 10 * time.Second, - NewBackoff: newTCPBackoff, // required: StackGo panics if nil. - }, - }) - dev.events <- tun.EventUp - return dev, dev, nil -} +var _ Stack = (*lnetoStack)(nil) // --- TCP --- @@ -367,7 +334,7 @@ func socketResult[T any](v any, err error) (T, error) { // family-bearing endpoint (the remote for a dial, otherwise the local bind), and pokes // the egress poll on entry and exit because connection setup (handshake, NDP/ARP) // queues egress frames that Read must drain promptly. -func (n *Net2) socket(ctx context.Context, proto string, sotype int, laddr, raddr netip.AddrPort) (any, error) { +func (n *lnetoStack) socket(ctx context.Context, proto string, sotype int, laddr, raddr netip.AddrPort) (any, error) { fam := raddr if !fam.IsValid() { fam = laddr @@ -383,106 +350,46 @@ func (n *Net2) socket(ctx context.Context, proto string, sotype int, laddr, radd return n.sgo.SocketNetip(ctx, network, family, sotype, laddr, raddr) } -func (n *Net2) dialTCPCtx(ctx context.Context, addr netip.AddrPort) (TCPConn, error) { +func (n *lnetoStack) dialTCPCtx(ctx context.Context, addr netip.AddrPort) (TCPConn, error) { v, err := n.socket(ctx, "tcp", syscall.SOCK_STREAM, netip.AddrPort{}, addr) return socketResult[TCPConn](v, err) } -func (n *Net2) DialContextTCPAddrPort(ctx context.Context, addr netip.AddrPort) (TCPConn, error) { +func (n *lnetoStack) DialContextTCPAddrPort(ctx context.Context, addr netip.AddrPort) (TCPConn, error) { return n.dialTCPCtx(ctx, addr) } -func (n *Net2) DialContextTCP(ctx context.Context, addr *net.TCPAddr) (TCPConn, error) { - if addr == nil { - return n.dialTCPCtx(ctx, netip.AddrPort{}) - } - ip, _ := netip.AddrFromSlice(addr.IP) - return n.dialTCPCtx(ctx, netip.AddrPortFrom(ip.Unmap(), uint16(addr.Port))) -} - -func (n *Net2) DialTCPAddrPort(addr netip.AddrPort) (TCPConn, error) { +func (n *lnetoStack) DialTCPAddrPort(addr netip.AddrPort) (TCPConn, error) { return n.dialTCPCtx(context.Background(), addr) } -func (n *Net2) DialTCP(addr *net.TCPAddr) (TCPConn, error) { - return n.DialContextTCP(context.Background(), addr) -} - // --- TCP listener --- -func (n *Net2) ListenTCPAddrPort(addr netip.AddrPort) (TCPListener, error) { +func (n *lnetoStack) ListenTCPAddrPort(addr netip.AddrPort) (TCPListener, error) { v, err := n.socket(context.Background(), "tcp", syscall.SOCK_STREAM, addr, netip.AddrPort{}) return socketResult[TCPListener](v, err) } -func (n *Net2) ListenTCP(addr *net.TCPAddr) (TCPListener, error) { - if addr == nil { - return n.ListenTCPAddrPort(netip.AddrPort{}) - } - ip, _ := netip.AddrFromSlice(addr.IP) - return n.ListenTCPAddrPort(netip.AddrPortFrom(ip.Unmap(), uint16(addr.Port))) -} - // --- UDP --- -func (n *Net2) ListenUDPAddrPort(laddr netip.AddrPort) (UDPConn, error) { +func (n *lnetoStack) ListenUDPAddrPort(laddr netip.AddrPort) (UDPConn, error) { v, err := n.socket(context.Background(), "udp", syscall.SOCK_DGRAM, laddr, netip.AddrPort{}) return socketResult[UDPConn](v, err) } -func (n *Net2) ListenUDP(laddr *net.UDPAddr) (UDPConn, error) { - if laddr == nil { - return n.ListenUDPAddrPort(netip.AddrPort{}) - } - ip, _ := netip.AddrFromSlice(laddr.IP) - return n.ListenUDPAddrPort(netip.AddrPortFrom(ip.Unmap(), uint16(laddr.Port))) -} - -func (n *Net2) DialUDPAddrPort(laddr, raddr netip.AddrPort) (UDPConn, error) { +func (n *lnetoStack) DialUDPAddrPort(laddr, raddr netip.AddrPort) (UDPConn, error) { v, err := n.socket(context.Background(), "udp", syscall.SOCK_DGRAM, laddr, raddr) return socketResult[UDPConn](v, err) } -func (n *Net2) DialUDP(laddr, raddr *net.UDPAddr) (UDPConn, error) { - var la, ra netip.AddrPort - if laddr != nil { - ip, _ := netip.AddrFromSlice(laddr.IP) - la = netip.AddrPortFrom(ip.Unmap(), uint16(laddr.Port)) - } - if raddr != nil { - ip, _ := netip.AddrFromSlice(raddr.IP) - ra = netip.AddrPortFrom(ip.Unmap(), uint16(raddr.Port)) - } - return n.DialUDPAddrPort(la, ra) -} - // --- Ping --- -func (n *Net2) DialPingAddr(_, _ netip.Addr) (*PingConn, error) { - return nil, errors.New("ping not implemented for Net2: PingConn is gvisor-coupled") +func (n *lnetoStack) DialPingAddr(_, _ netip.Addr) (*PingConn, error) { + return nil, errors.New("ping not implemented for lnetoStack: PingConn is gvisor-coupled") } -func (n *Net2) ListenPingAddr(_ netip.Addr) (*PingConn, error) { - return nil, errors.New("ping not implemented for Net2: PingConn is gvisor-coupled") -} - -func (n *Net2) DialPing(laddr, raddr *PingAddr) (*PingConn, error) { - var la, ra netip.Addr - if laddr != nil { - la = laddr.addr - } - if raddr != nil { - ra = raddr.addr - } - return n.DialPingAddr(la, ra) -} - -func (n *Net2) ListenPing(laddr *PingAddr) (*PingConn, error) { - var la netip.Addr - if laddr != nil { - la = laddr.addr - } - return n.ListenPingAddr(la) +func (n *lnetoStack) ListenPingAddr(_ netip.Addr) (*PingConn, error) { + return nil, errors.New("ping not implemented for lnetoStack: PingConn is gvisor-coupled") } // --- DNS --- @@ -502,7 +409,7 @@ func dnsError(host string, err error) *net.DNSError { // non-domain hosts and stacks with no address family return an IsNotFound DNSError; // A and AAAA are queried for the enabled families and, when IPv6 is enabled, IPv6 // results are ordered first (no RFC 6724). -func (n *Net2) LookupContextHost(ctx context.Context, host string) ([]string, error) { +func (n *lnetoStack) LookupContextHost(ctx context.Context, host string) ([]string, error) { if host == "" || (!n.hasV4 && !n.hasV6) { return nil, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true} } @@ -566,110 +473,3 @@ func (n *Net2) LookupContextHost(ctx context.Context, host string) ([]string, er } return out, nil } - -func (n *Net2) LookupHost(host string) ([]string, error) { - return n.LookupContextHost(context.Background(), host) -} - -// --- Generic Dial --- - -var protoSplitter2 = regexp.MustCompile(`^(tcp|udp|ping)(4|6)?$`) - -func (n *Net2) DialContext(ctx context.Context, network, address string) (net.Conn, error) { - if ctx == nil { - panic("nil context") - } - matches := protoSplitter2.FindStringSubmatch(network) - if matches == nil { - return nil, &net.OpError{Op: "dial", Err: net.UnknownNetworkError(network)} - } - acceptV4 := len(matches[2]) == 0 || matches[2] == "4" - acceptV6 := len(matches[2]) == 0 || matches[2] == "6" - - var host string - var port int - if matches[1] == "ping" { - host = address - } else { - var sport string - var err error - host, sport, err = net.SplitHostPort(address) - if err != nil { - return nil, &net.OpError{Op: "dial", Err: err} - } - port, err = strconv.Atoi(sport) - if err != nil || port < 0 || port > 65535 { - return nil, &net.OpError{Op: "dial", Err: errNumericPort} - } - } - - allAddr, err := n.LookupContextHost(ctx, host) - if err != nil { - return nil, &net.OpError{Op: "dial", Err: err} - } - - var addrs []netip.AddrPort - for _, a := range allAddr { - ip, err := netip.ParseAddr(a) - if err == nil && ((ip.Is4() && acceptV4) || (ip.Is6() && acceptV6)) { - addrs = append(addrs, netip.AddrPortFrom(ip, uint16(port))) - } - } - if len(addrs) == 0 && len(allAddr) != 0 { - return nil, &net.OpError{Op: "dial", Err: errNoSuitableAddress} - } - - var firstErr error - for i, addr := range addrs { - select { - case <-ctx.Done(): - err := ctx.Err() - if err == context.Canceled { - err = errCanceled - } else if err == context.DeadlineExceeded { - err = errTimeout - } - return nil, &net.OpError{Op: "dial", Err: err} - default: - } - dialCtx := ctx - if deadline, hasDeadline := ctx.Deadline(); hasDeadline { - pd, err := partialDeadline(time.Now(), deadline, len(addrs)-i) - if err != nil { - if firstErr == nil { - firstErr = &net.OpError{Op: "dial", Err: err} - } - break - } - if pd.Before(deadline) { - var cancel context.CancelFunc - dialCtx, cancel = context.WithDeadline(ctx, pd) - defer cancel() - } - } - - var c net.Conn - switch matches[1] { - case "tcp": - c, err = n.DialContextTCPAddrPort(dialCtx, addr) - case "udp": - c, err = n.DialUDPAddrPort(netip.AddrPort{}, addr) - case "ping": - c, err = n.DialPingAddr(netip.Addr{}, addr.Addr()) - } - if err == nil { - return c, nil - } - if firstErr == nil { - firstErr = err - } - } - if firstErr == nil { - firstErr = &net.OpError{Op: "dial", Err: errMissingAddress} - } - return nil, firstErr -} - -func (n *Net2) Dial(network, address string) (net.Conn, error) { - return n.DialContext(context.Background(), network, address) -} diff --git a/tun/netstack/lneto_test.go b/tun/netstack/lneto_test.go index b3c0af4a8..0fa17aa1b 100644 --- a/tun/netstack/lneto_test.go +++ b/tun/netstack/lneto_test.go @@ -18,10 +18,10 @@ import ( "golang.zx2c4.com/wireguard/tun" ) -// TestNet2_Construct is a regression test: the previous CreateNetTUN2 omitted +// TestNet2_Construct is a regression test: the previous CreateNetTUNLneto omitted // TCPPoolConfig.NewBackoff, which made StackGo panic at construction. func TestNet2_Construct(t *testing.T) { - dev, net2, err := CreateNetTUN2( + dev, net2, err := CreateNetTUNLneto( []netip.Addr{netip.MustParseAddr("10.0.0.1")}, []netip.Addr{netip.MustParseAddr("8.8.8.8")}, 1500, @@ -31,7 +31,7 @@ func TestNet2_Construct(t *testing.T) { } defer dev.Close() if net2 == nil { - t.Fatal("nil Net2") + t.Fatal("nil Net") } select { case ev := <-dev.Events(): @@ -47,7 +47,7 @@ func TestNet2_Construct(t *testing.T) { // panicked with "send on closed channel". Run with -race. func TestNet2_CloseRace(t *testing.T) { for i := 0; i < 50; i++ { - dev, _, err := CreateNetTUN2( + dev, _, err := CreateNetTUNLneto( []netip.Addr{netip.MustParseAddr("10.0.0.1")}, nil, 1500, ) @@ -98,7 +98,7 @@ func TestNet2_CloseRace(t *testing.T) { // TestNet2_ListenTCPPort0 covers gap D: listening on port 0 must auto-assign an // ephemeral port instead of failing (the library previously returned ErrZeroSource). func TestNet2_ListenTCPPort0(t *testing.T) { - dev, net2, err := CreateNetTUN2([]netip.Addr{netip.MustParseAddr("10.0.0.1")}, nil, 1500) + dev, net2, err := CreateNetTUNLneto([]netip.Addr{netip.MustParseAddr("10.0.0.1")}, nil, 1500) if err != nil { t.Fatal(err) } @@ -123,11 +123,11 @@ func TestNet2_UDPEcho(t *testing.T) { addrB = "10.0.0.2" port = 9999 ) - devA, netA, err := CreateNetTUN2([]netip.Addr{netip.MustParseAddr(addrA)}, nil, 1500) + devA, netA, err := CreateNetTUNLneto([]netip.Addr{netip.MustParseAddr(addrA)}, nil, 1500) if err != nil { t.Fatal(err) } - devB, netB, err := CreateNetTUN2([]netip.Addr{netip.MustParseAddr(addrB)}, nil, 1500) + devB, netB, err := CreateNetTUNLneto([]netip.Addr{netip.MustParseAddr(addrB)}, nil, 1500) if err != nil { t.Fatal(err) } @@ -216,11 +216,11 @@ func TestNet2_TCPEcho(t *testing.T) { func testTCPEcho(t *testing.T, addrA, addrB string) { const port = 1234 - devA, netA, err := CreateNetTUN2([]netip.Addr{netip.MustParseAddr(addrA)}, nil, 1500) + devA, netA, err := CreateNetTUNLneto([]netip.Addr{netip.MustParseAddr(addrA)}, nil, 1500) if err != nil { t.Fatal(err) } - devB, netB, err := CreateNetTUN2([]netip.Addr{netip.MustParseAddr(addrB)}, nil, 1500) + devB, netB, err := CreateNetTUNLneto([]netip.Addr{netip.MustParseAddr(addrB)}, nil, 1500) if err != nil { t.Fatal(err) } diff --git a/tun/netstack/net.go b/tun/netstack/net.go new file mode 100644 index 000000000..480f59dc7 --- /dev/null +++ b/tun/netstack/net.go @@ -0,0 +1,369 @@ +/* SPDX-License-Identifier: MIT + * + * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. + */ + +package netstack + +import ( + "context" + "errors" + "net" + "net/netip" + "regexp" + "strconv" + "time" +) + +// Net is the userspace networking API (Dial/Listen/DNS) layered over a backend +// [Stack]. It is backend-agnostic: [CreateNetTUN] wraps the gvisor backend and +// [CreateNetTUNLneto] wraps the lneto backend, but both return a *Net with the +// same surface so callers can swap backends transparently. +// +// The backend object also implements [golang.zx2c4.com/wireguard/tun.Device] +// (returned as the first value from the constructors) and is the data plane that +// moves IP packets to and from WireGuard. +type Net struct { + stack Stack +} + +// Stack is the backend-specific primitive set that [Net] delegates to. Both the +// gvisor netTun and the lneto stack implement it. Higher-level conveniences +// (the *net.TCPAddr/*net.UDPAddr/*PingAddr overloads, generic Dial/DialContext, +// LookupHost) live once on [Net] and are written in terms of these primitives. +type Stack interface { + DialContextTCPAddrPort(ctx context.Context, addr netip.AddrPort) (TCPConn, error) + DialTCPAddrPort(addr netip.AddrPort) (TCPConn, error) + ListenTCPAddrPort(addr netip.AddrPort) (TCPListener, error) + DialUDPAddrPort(laddr, raddr netip.AddrPort) (UDPConn, error) + ListenUDPAddrPort(laddr netip.AddrPort) (UDPConn, error) + DialPingAddr(laddr, raddr netip.Addr) (*PingConn, error) + ListenPingAddr(laddr netip.Addr) (*PingConn, error) + LookupContextHost(ctx context.Context, host string) ([]string, error) +} + +// TCPConn is the connection type returned by the TCP dial methods. gvisor's +// *gonet.TCPConn and lneto's TCP connection both satisfy it. +type TCPConn interface { + Close() error + CloseRead() error + CloseWrite() error + LocalAddr() net.Addr + Read(b []byte) (int, error) + RemoteAddr() net.Addr + SetDeadline(t time.Time) error + SetReadDeadline(t time.Time) error + SetWriteDeadline(t time.Time) error + Write(b []byte) (int, error) +} + +// UDPConn is the connection type returned by the UDP dial/listen methods. +type UDPConn interface { + Close() error + LocalAddr() net.Addr + Read(b []byte) (int, error) + ReadFrom(b []byte) (int, net.Addr, error) + RemoteAddr() net.Addr + SetDeadline(t time.Time) error + SetReadDeadline(t time.Time) error + SetWriteDeadline(t time.Time) error + Write(b []byte) (int, error) + WriteTo(b []byte, addr net.Addr) (int, error) +} + +// TCPListener is the listener type returned by the TCP listen methods. +type TCPListener interface { + Accept() (net.Conn, error) + Addr() net.Addr + Close() error + Shutdown() +} + +// --- TCP --- + +func (n *Net) DialContextTCPAddrPort(ctx context.Context, addr netip.AddrPort) (TCPConn, error) { + return n.stack.DialContextTCPAddrPort(ctx, addr) +} + +func (n *Net) DialContextTCP(ctx context.Context, addr *net.TCPAddr) (TCPConn, error) { + if addr == nil { + return n.stack.DialContextTCPAddrPort(ctx, netip.AddrPort{}) + } + ip, _ := netip.AddrFromSlice(addr.IP) + return n.stack.DialContextTCPAddrPort(ctx, netip.AddrPortFrom(ip.Unmap(), uint16(addr.Port))) +} + +func (n *Net) DialTCPAddrPort(addr netip.AddrPort) (TCPConn, error) { + return n.stack.DialTCPAddrPort(addr) +} + +func (n *Net) DialTCP(addr *net.TCPAddr) (TCPConn, error) { + if addr == nil { + return n.stack.DialTCPAddrPort(netip.AddrPort{}) + } + ip, _ := netip.AddrFromSlice(addr.IP) + return n.stack.DialTCPAddrPort(netip.AddrPortFrom(ip.Unmap(), uint16(addr.Port))) +} + +func (n *Net) ListenTCPAddrPort(addr netip.AddrPort) (TCPListener, error) { + return n.stack.ListenTCPAddrPort(addr) +} + +func (n *Net) ListenTCP(addr *net.TCPAddr) (TCPListener, error) { + if addr == nil { + return n.stack.ListenTCPAddrPort(netip.AddrPort{}) + } + ip, _ := netip.AddrFromSlice(addr.IP) + return n.stack.ListenTCPAddrPort(netip.AddrPortFrom(ip.Unmap(), uint16(addr.Port))) +} + +// --- UDP --- + +func (n *Net) DialUDPAddrPort(laddr, raddr netip.AddrPort) (UDPConn, error) { + return n.stack.DialUDPAddrPort(laddr, raddr) +} + +func (n *Net) ListenUDPAddrPort(laddr netip.AddrPort) (UDPConn, error) { + return n.stack.ListenUDPAddrPort(laddr) +} + +func (n *Net) DialUDP(laddr, raddr *net.UDPAddr) (UDPConn, error) { + var la, ra netip.AddrPort + if laddr != nil { + ip, _ := netip.AddrFromSlice(laddr.IP) + la = netip.AddrPortFrom(ip.Unmap(), uint16(laddr.Port)) + } + if raddr != nil { + ip, _ := netip.AddrFromSlice(raddr.IP) + ra = netip.AddrPortFrom(ip.Unmap(), uint16(raddr.Port)) + } + return n.stack.DialUDPAddrPort(la, ra) +} + +func (n *Net) ListenUDP(laddr *net.UDPAddr) (UDPConn, error) { + return n.DialUDP(laddr, nil) +} + +// --- Ping --- + +func (n *Net) DialPingAddr(laddr, raddr netip.Addr) (*PingConn, error) { + return n.stack.DialPingAddr(laddr, raddr) +} + +func (n *Net) ListenPingAddr(laddr netip.Addr) (*PingConn, error) { + return n.stack.ListenPingAddr(laddr) +} + +func (n *Net) DialPing(laddr, raddr *PingAddr) (*PingConn, error) { + var la, ra netip.Addr + if laddr != nil { + la = laddr.addr + } + if raddr != nil { + ra = raddr.addr + } + return n.stack.DialPingAddr(la, ra) +} + +func (n *Net) ListenPing(laddr *PingAddr) (*PingConn, error) { + var la netip.Addr + if laddr != nil { + la = laddr.addr + } + return n.stack.ListenPingAddr(la) +} + +// --- DNS --- + +func (n *Net) LookupContextHost(ctx context.Context, host string) ([]string, error) { + return n.stack.LookupContextHost(ctx, host) +} + +func (n *Net) LookupHost(host string) ([]string, error) { + return n.stack.LookupContextHost(context.Background(), host) +} + +// --- Generic Dial --- + +var protoSplitter = regexp.MustCompile(`^(tcp|udp|ping)(4|6)?$`) + +func (n *Net) DialContext(ctx context.Context, network, address string) (net.Conn, error) { + if ctx == nil { + panic("nil context") + } + var acceptV4, acceptV6 bool + matches := protoSplitter.FindStringSubmatch(network) + if matches == nil { + return nil, &net.OpError{Op: "dial", Err: net.UnknownNetworkError(network)} + } else if len(matches[2]) == 0 { + acceptV4 = true + acceptV6 = true + } else { + acceptV4 = matches[2][0] == '4' + acceptV6 = !acceptV4 + } + var host string + var port int + if matches[1] == "ping" { + host = address + } else { + var sport string + var err error + host, sport, err = net.SplitHostPort(address) + if err != nil { + return nil, &net.OpError{Op: "dial", Err: err} + } + port, err = strconv.Atoi(sport) + if err != nil || port < 0 || port > 65535 { + return nil, &net.OpError{Op: "dial", Err: errNumericPort} + } + } + allAddr, err := n.LookupContextHost(ctx, host) + if err != nil { + return nil, &net.OpError{Op: "dial", Err: err} + } + var addrs []netip.AddrPort + for _, addr := range allAddr { + ip, err := netip.ParseAddr(addr) + if err == nil && ((ip.Is4() && acceptV4) || (ip.Is6() && acceptV6)) { + addrs = append(addrs, netip.AddrPortFrom(ip, uint16(port))) + } + } + if len(addrs) == 0 && len(allAddr) != 0 { + return nil, &net.OpError{Op: "dial", Err: errNoSuitableAddress} + } + + var firstErr error + for i, addr := range addrs { + select { + case <-ctx.Done(): + err := ctx.Err() + if err == context.Canceled { + err = errCanceled + } else if err == context.DeadlineExceeded { + err = errTimeout + } + return nil, &net.OpError{Op: "dial", Err: err} + default: + } + + dialCtx := ctx + if deadline, hasDeadline := ctx.Deadline(); hasDeadline { + partialDeadline, err := partialDeadline(time.Now(), deadline, len(addrs)-i) + if err != nil { + if firstErr == nil { + firstErr = &net.OpError{Op: "dial", Err: err} + } + break + } + if partialDeadline.Before(deadline) { + var cancel context.CancelFunc + dialCtx, cancel = context.WithDeadline(ctx, partialDeadline) + defer cancel() + } + } + + var c net.Conn + switch matches[1] { + case "tcp": + c, err = n.DialContextTCPAddrPort(dialCtx, addr) + case "udp": + c, err = n.DialUDPAddrPort(netip.AddrPort{}, addr) + case "ping": + c, err = n.DialPingAddr(netip.Addr{}, addr.Addr()) + } + if err == nil { + return c, nil + } + if firstErr == nil { + firstErr = err + } + } + if firstErr == nil { + firstErr = &net.OpError{Op: "dial", Err: errMissingAddress} + } + return nil, firstErr +} + +func (n *Net) Dial(network, address string) (net.Conn, error) { + return n.DialContext(context.Background(), network, address) +} + +// --- shared helpers --- + +var ( + errNoSuchHost = errors.New("no such host") + errLameReferral = errors.New("lame referral") + errCannotUnmarshalDNSMessage = errors.New("cannot unmarshal DNS message") + errCannotMarshalDNSMessage = errors.New("cannot marshal DNS message") + errServerMisbehaving = errors.New("server misbehaving") + errInvalidDNSResponse = errors.New("invalid DNS response") + errNoAnswerFromDNSServer = errors.New("no answer from DNS server") + errServerTemporarilyMisbehaving = errors.New("server misbehaving") + errCanceled = errors.New("operation was canceled") + errTimeout = errors.New("i/o timeout") + errNumericPort = errors.New("port must be numeric") + errNoSuitableAddress = errors.New("no suitable address found") + errMissingAddress = errors.New("missing address") +) + +func isDomainName(s string) bool { + l := len(s) + if l == 0 || l > 254 || l == 254 && s[l-1] != '.' { + return false + } + last := byte('.') + nonNumeric := false + partlen := 0 + for i := 0; i < len(s); i++ { + c := s[i] + switch { + default: + return false + case 'a' <= c && c <= 'z' || 'A' <= c && c <= 'Z' || c == '_': + nonNumeric = true + partlen++ + case '0' <= c && c <= '9': + partlen++ + case c == '-': + if last == '.' { + return false + } + partlen++ + nonNumeric = true + case c == '.': + if last == '.' || last == '-' { + return false + } + if partlen > 63 || partlen == 0 { + return false + } + partlen = 0 + } + last = c + } + if last == '-' || partlen > 63 { + return false + } + return nonNumeric +} + +func partialDeadline(now, deadline time.Time, addrsRemaining int) (time.Time, error) { + if deadline.IsZero() { + return deadline, nil + } + timeRemaining := deadline.Sub(now) + if timeRemaining <= 0 { + return time.Time{}, errTimeout + } + timeout := timeRemaining / time.Duration(addrsRemaining) + const saneMinimum = 2 * time.Second + if timeout < saneMinimum { + if timeRemaining < saneMinimum { + timeout = timeRemaining + } else { + timeout = saneMinimum + } + } + return now.Add(timeout), nil +} From deedbc4c8746606c64a07a72929d54f2af94a719 Mon Sep 17 00:00:00 2001 From: Patricio Whittingslow Date: Tue, 16 Jun 2026 12:16:55 -0300 Subject: [PATCH 04/12] split lneto/gvisor with build tag wglneto --- tun/netstack/gvisor.go | 54 +++++++++++--------------------------- tun/netstack/lneto.go | 8 +++--- tun/netstack/lneto_test.go | 2 ++ tun/netstack/net.go | 47 ++++++++++++++++++++++++++++----- tun/netstack/net_gvisor.go | 15 +++++++++++ tun/netstack/net_lneto.go | 15 +++++++++++ 6 files changed, 93 insertions(+), 48 deletions(-) create mode 100644 tun/netstack/net_gvisor.go create mode 100644 tun/netstack/net_lneto.go diff --git a/tun/netstack/gvisor.go b/tun/netstack/gvisor.go index 9c255a6ce..3432e65a5 100644 --- a/tun/netstack/gvisor.go +++ b/tun/netstack/gvisor.go @@ -48,7 +48,7 @@ type netTun struct { hasV4, hasV6 bool } -func CreateNetTUN(localAddresses, dnsServers []netip.Addr, mtu int) (tun.Device, *Net, error) { +func CreateNetTUNGvisor(localAddresses, dnsServers []netip.Addr, mtu int) (tun.Device, *Net, error) { opts := stack.Options{ NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol}, TransportProtocols: []stack.TransportProtocolFactory{tcp.NewProtocol, udp.NewProtocol, icmp.NewProtocol6, icmp.NewProtocol4}, @@ -254,7 +254,8 @@ func (tun *netTun) ListenUDPAddrPort(laddr netip.AddrPort) (UDPConn, error) { return tun.DialUDPAddrPort(laddr, netip.AddrPort{}) } -type PingConn struct { +// pingConn is the gvisor-backed [PingConn] implementation. +type pingConn struct { laddr PingAddr raddr PingAddr wq waiter.Queue @@ -262,30 +263,7 @@ type PingConn struct { deadline *time.Timer } -type PingAddr struct{ addr netip.Addr } - -func (ia PingAddr) String() string { - return ia.addr.String() -} - -func (ia PingAddr) Network() string { - if ia.addr.Is4() { - return "ping4" - } else if ia.addr.Is6() { - return "ping6" - } - return "ping" -} - -func (ia PingAddr) Addr() netip.Addr { - return ia.addr -} - -func PingAddrFromAddr(addr netip.Addr) *PingAddr { - return &PingAddr{addr} -} - -func (tun *netTun) DialPingAddr(laddr, raddr netip.Addr) (*PingConn, error) { +func (tun *netTun) DialPingAddr(laddr, raddr netip.Addr) (PingConn, error) { if !laddr.IsValid() && !raddr.IsValid() { return nil, errors.New("ping dial: invalid address") } @@ -306,7 +284,7 @@ func (tun *netTun) DialPingAddr(laddr, raddr netip.Addr) (*PingConn, error) { pn = ipv6.ProtocolNumber } - pc := &PingConn{ + pc := &pingConn{ laddr: PingAddr{laddr}, deadline: time.NewTimer(time.Hour << 10), } @@ -336,29 +314,29 @@ func (tun *netTun) DialPingAddr(laddr, raddr netip.Addr) (*PingConn, error) { return pc, nil } -func (tun *netTun) ListenPingAddr(laddr netip.Addr) (*PingConn, error) { +func (tun *netTun) ListenPingAddr(laddr netip.Addr) (PingConn, error) { return tun.DialPingAddr(laddr, netip.Addr{}) } -func (pc *PingConn) LocalAddr() net.Addr { +func (pc *pingConn) LocalAddr() net.Addr { return pc.laddr } -func (pc *PingConn) RemoteAddr() net.Addr { +func (pc *pingConn) RemoteAddr() net.Addr { return pc.raddr } -func (pc *PingConn) Close() error { +func (pc *pingConn) Close() error { pc.deadline.Reset(0) pc.ep.Close() return nil } -func (pc *PingConn) SetWriteDeadline(t time.Time) error { +func (pc *pingConn) SetWriteDeadline(t time.Time) error { return errors.New("not implemented") } -func (pc *PingConn) WriteTo(p []byte, addr net.Addr) (n int, err error) { +func (pc *pingConn) WriteTo(p []byte, addr net.Addr) (n int, err error) { var na netip.Addr switch v := addr.(type) { case *PingAddr: @@ -385,11 +363,11 @@ func (pc *PingConn) WriteTo(p []byte, addr net.Addr) (n int, err error) { return int(n64), nil } -func (pc *PingConn) Write(p []byte) (n int, err error) { +func (pc *pingConn) Write(p []byte) (n int, err error) { return pc.WriteTo(p, &pc.raddr) } -func (pc *PingConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) { +func (pc *pingConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) { e, notifyCh := waiter.NewChannelEntry(waiter.EventIn) pc.wq.EventRegister(&e) defer pc.wq.EventUnregister(&e) @@ -413,18 +391,18 @@ func (pc *PingConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) { return res.Count, &PingAddr{remoteAddr}, nil } -func (pc *PingConn) Read(p []byte) (n int, err error) { +func (pc *pingConn) Read(p []byte) (n int, err error) { n, _, err = pc.ReadFrom(p) return } -func (pc *PingConn) SetDeadline(t time.Time) error { +func (pc *pingConn) SetDeadline(t time.Time) error { // pc.SetWriteDeadline is unimplemented return pc.SetReadDeadline(t) } -func (pc *PingConn) SetReadDeadline(t time.Time) error { +func (pc *pingConn) SetReadDeadline(t time.Time) error { pc.deadline.Reset(time.Until(t)) return nil } diff --git a/tun/netstack/lneto.go b/tun/netstack/lneto.go index 3721deaca..1a2dcb13a 100644 --- a/tun/netstack/lneto.go +++ b/tun/netstack/lneto.go @@ -384,12 +384,12 @@ func (n *lnetoStack) DialUDPAddrPort(laddr, raddr netip.AddrPort) (UDPConn, erro // --- Ping --- -func (n *lnetoStack) DialPingAddr(_, _ netip.Addr) (*PingConn, error) { - return nil, errors.New("ping not implemented for lnetoStack: PingConn is gvisor-coupled") +func (n *lnetoStack) DialPingAddr(_, _ netip.Addr) (PingConn, error) { + return nil, errors.New("ping not implemented for lnetoStack") } -func (n *lnetoStack) ListenPingAddr(_ netip.Addr) (*PingConn, error) { - return nil, errors.New("ping not implemented for lnetoStack: PingConn is gvisor-coupled") +func (n *lnetoStack) ListenPingAddr(_ netip.Addr) (PingConn, error) { + return nil, errors.New("ping not implemented for lnetoStack") } // --- DNS --- diff --git a/tun/netstack/lneto_test.go b/tun/netstack/lneto_test.go index 0fa17aa1b..634c258d9 100644 --- a/tun/netstack/lneto_test.go +++ b/tun/netstack/lneto_test.go @@ -1,3 +1,5 @@ +//go:build wglneto + /* SPDX-License-Identifier: MIT * * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. diff --git a/tun/netstack/net.go b/tun/netstack/net.go index 480f59dc7..bc96cdb3c 100644 --- a/tun/netstack/net.go +++ b/tun/netstack/net.go @@ -37,8 +37,8 @@ type Stack interface { ListenTCPAddrPort(addr netip.AddrPort) (TCPListener, error) DialUDPAddrPort(laddr, raddr netip.AddrPort) (UDPConn, error) ListenUDPAddrPort(laddr netip.AddrPort) (UDPConn, error) - DialPingAddr(laddr, raddr netip.Addr) (*PingConn, error) - ListenPingAddr(laddr netip.Addr) (*PingConn, error) + DialPingAddr(laddr, raddr netip.Addr) (PingConn, error) + ListenPingAddr(laddr netip.Addr) (PingConn, error) LookupContextHost(ctx context.Context, host string) ([]string, error) } @@ -79,6 +79,41 @@ type TCPListener interface { Shutdown() } +// PingConn is the ICMP "ping" connection returned by the ping methods. It is both a +// [net.Conn] (Read/Write against the dialed peer) and supports addressed I/O via +// ReadFrom/WriteTo. The gvisor backend returns a concrete implementation; the lneto +// backend does not implement ping and returns an error. +type PingConn interface { + net.Conn + ReadFrom(p []byte) (int, net.Addr, error) + WriteTo(p []byte, addr net.Addr) (int, error) +} + +// PingAddr is a [net.Addr] for the "ping" pseudo-networks. It wraps a bare +// [netip.Addr] (no port) and is backend-neutral. +type PingAddr struct{ addr netip.Addr } + +func (ia PingAddr) String() string { + return ia.addr.String() +} + +func (ia PingAddr) Network() string { + if ia.addr.Is4() { + return "ping4" + } else if ia.addr.Is6() { + return "ping6" + } + return "ping" +} + +func (ia PingAddr) Addr() netip.Addr { + return ia.addr +} + +func PingAddrFromAddr(addr netip.Addr) *PingAddr { + return &PingAddr{addr} +} + // --- TCP --- func (n *Net) DialContextTCPAddrPort(ctx context.Context, addr netip.AddrPort) (TCPConn, error) { @@ -146,15 +181,15 @@ func (n *Net) ListenUDP(laddr *net.UDPAddr) (UDPConn, error) { // --- Ping --- -func (n *Net) DialPingAddr(laddr, raddr netip.Addr) (*PingConn, error) { +func (n *Net) DialPingAddr(laddr, raddr netip.Addr) (PingConn, error) { return n.stack.DialPingAddr(laddr, raddr) } -func (n *Net) ListenPingAddr(laddr netip.Addr) (*PingConn, error) { +func (n *Net) ListenPingAddr(laddr netip.Addr) (PingConn, error) { return n.stack.ListenPingAddr(laddr) } -func (n *Net) DialPing(laddr, raddr *PingAddr) (*PingConn, error) { +func (n *Net) DialPing(laddr, raddr *PingAddr) (PingConn, error) { var la, ra netip.Addr if laddr != nil { la = laddr.addr @@ -165,7 +200,7 @@ func (n *Net) DialPing(laddr, raddr *PingAddr) (*PingConn, error) { return n.stack.DialPingAddr(la, ra) } -func (n *Net) ListenPing(laddr *PingAddr) (*PingConn, error) { +func (n *Net) ListenPing(laddr *PingAddr) (PingConn, error) { var la netip.Addr if laddr != nil { la = laddr.addr diff --git a/tun/netstack/net_gvisor.go b/tun/netstack/net_gvisor.go new file mode 100644 index 000000000..464af07d8 --- /dev/null +++ b/tun/netstack/net_gvisor.go @@ -0,0 +1,15 @@ +//go:build !wglneto + +package netstack + +import ( + "net/netip" + + "golang.zx2c4.com/wireguard/tun" +) + +// CreateNetTUN builds the default gvisor-backed netstack TUN device. Build with +// -tags wglneto to select the lneto backend instead (see CreateNetTUNLneto). +func CreateNetTUN(localAddresses, dnsServers []netip.Addr, mtu int) (tun.Device, *Net, error) { + return CreateNetTUNGvisor(localAddresses, dnsServers, mtu) +} diff --git a/tun/netstack/net_lneto.go b/tun/netstack/net_lneto.go new file mode 100644 index 000000000..f860870e0 --- /dev/null +++ b/tun/netstack/net_lneto.go @@ -0,0 +1,15 @@ +//go:build wglneto + +package netstack + +import ( + "net/netip" + + "golang.zx2c4.com/wireguard/tun" +) + +// CreateNetTUN builds the lneto-backed netstack TUN device. This is the +// implementation selected when building with -tags wglneto. +func CreateNetTUN(localAddresses, dnsServers []netip.Addr, mtu int) (tun.Device, *Net, error) { + return CreateNetTUNLneto(localAddresses, dnsServers, mtu) +} From 8dd8b26b6f6d3c9b4455f12f2f14b5f7dba4dbb9 Mon Sep 17 00:00:00 2001 From: Patricio Whittingslow Date: Wed, 17 Jun 2026 12:46:04 -0300 Subject: [PATCH 05/12] apply coderabbit suggestion fixes --- tun/netstack/lneto.go | 3 +++ tun/netstack/net.go | 7 +++++-- 2 files changed, 8 insertions(+), 2 deletions(-) diff --git a/tun/netstack/lneto.go b/tun/netstack/lneto.go index 1a2dcb13a..3efa6091a 100644 --- a/tun/netstack/lneto.go +++ b/tun/netstack/lneto.go @@ -30,6 +30,9 @@ func CreateNetTUNLneto(localAddresses, dnsServers []netip.Addr, mtu int) (tun.De if mtu <= 0 { mtu = 1500 } + if mtu > 65535 { + return nil, nil, fmt.Errorf("CreateNetTUNLneto: mtu %d exceeds maximum 65535", mtu) + } dev := &lnetoStack{ events: make(chan tun.Event, 10), closed: make(chan struct{}), diff --git a/tun/netstack/net.go b/tun/netstack/net.go index bc96cdb3c..e40af8d3e 100644 --- a/tun/netstack/net.go +++ b/tun/netstack/net.go @@ -283,6 +283,7 @@ func (n *Net) DialContext(ctx context.Context, network, address string) (net.Con } dialCtx := ctx + var cancel context.CancelFunc if deadline, hasDeadline := ctx.Deadline(); hasDeadline { partialDeadline, err := partialDeadline(time.Now(), deadline, len(addrs)-i) if err != nil { @@ -292,9 +293,7 @@ func (n *Net) DialContext(ctx context.Context, network, address string) (net.Con break } if partialDeadline.Before(deadline) { - var cancel context.CancelFunc dialCtx, cancel = context.WithDeadline(ctx, partialDeadline) - defer cancel() } } @@ -307,6 +306,10 @@ func (n *Net) DialContext(ctx context.Context, network, address string) (net.Con case "ping": c, err = n.DialPingAddr(netip.Addr{}, addr.Addr()) } + if cancel != nil { + // This cancel belongs to a function-local context so cancel required to avoid leaks. + cancel() + } if err == nil { return c, nil } From 9b000ed491f197b985b2a08d05a3220c1e647fb9 Mon Sep 17 00:00:00 2001 From: Patricio Whittingslow Date: Wed, 17 Jun 2026 16:49:19 -0300 Subject: [PATCH 06/12] debugIPPackets and wasm concerrency fixes --- tun/netstack/gvisor.go | 2 ++ tun/netstack/lneto.go | 13 ++++++++++--- tun/netstack/lneto_tuncount.go | 30 +++++++++++++++++++++++++++++ tun/netstack/lneto_tuncount_off.go | 7 +++++++ tun/netstack/net_debug.go | 31 ++++++++++++++++++++++++++++++ tun/netstack/net_debugoff.go | 5 +++++ 6 files changed, 85 insertions(+), 3 deletions(-) create mode 100644 tun/netstack/lneto_tuncount.go create mode 100644 tun/netstack/lneto_tuncount_off.go create mode 100644 tun/netstack/net_debug.go create mode 100644 tun/netstack/net_debugoff.go diff --git a/tun/netstack/gvisor.go b/tun/netstack/gvisor.go index 3432e65a5..50e6ff789 100644 --- a/tun/netstack/gvisor.go +++ b/tun/netstack/gvisor.go @@ -129,6 +129,7 @@ func (tun *netTun) Read(buf [][]byte, sizes []int, offset int) (int, error) { return 0, err } sizes[0] = n + debugIPPacket(true, buf[0][offset:offset+n]) return 1, nil } @@ -138,6 +139,7 @@ func (tun *netTun) Write(buf [][]byte, offset int) (int, error) { if len(packet) == 0 { continue } + debugIPPacket(false, packet) pkb := stack.NewPacketBuffer(stack.PacketBufferOptions{Payload: buffer.MakeWithData(packet)}) switch packet[0] >> 4 { diff --git a/tun/netstack/lneto.go b/tun/netstack/lneto.go index 3efa6091a..53dc52a3a 100644 --- a/tun/netstack/lneto.go +++ b/tun/netstack/lneto.go @@ -110,11 +110,16 @@ func CreateNetTUNLneto(localAddresses, dnsServers []netip.Addr, mtu int) (tun.De return nil, nil, fmt.Errorf("CreateNetTUNLneto: stack reset: %w", err) } - // GOMAXPROCS-aware backoff: on a single-threaded runtime only yield cooperatively - // (sleeping the poll would starve egress), otherwise sleep with exponential backoff. + // Backoff selection. A native single-threaded runtime (GOMAXPROCS==1) only yields + // cooperatively, because sleeping the poll would starve egress — there is no other + // thread to drain it. js/wasm is also single-threaded but is the opposite case: it + // MUST sleep so the runtime hands control back to the browser event loop (a bare + // Gosched there starves all WebSocket/WireGuard I/O and freezes the tab). So js, + // like the multi-threaded case, keeps the exponential sleeping backoffs — applied + // to BOTH the stack poll and every per-connection RWBackoff. baseStack := defaultStackBackoff newTCPBackoff := func() lneto.BackoffStrategy { return defaultTCPBackoff } - if runtime.GOMAXPROCS(0) == 1 { + if runtime.GOMAXPROCS(0) == 1 && runtime.GOOS != "js" { baseStack = backoffYield newTCPBackoff = func() lneto.BackoffStrategy { return backoffYield } } @@ -270,6 +275,7 @@ func (n *lnetoStack) Write(bufs [][]byte, offset int) (int, error) { for _, buf := range bufs { if pkt := buf[offset:]; len(pkt) > 0 { n.sa.IngressIP(pkt) // errors dropped; stack silently filters bad packets + debugIPPacket(false, pkt) wrote = true } } @@ -295,6 +301,7 @@ func (n *lnetoStack) Read(bufs [][]byte, sizes []int, offset int) (int, error) { cnt, _ := n.sa.EgressIP(dst) if cnt > 0 { sizes[0] = cnt + debugIPPacket(true, dst[:cnt]) return 1, nil } n.backoff.Do(backoffs) // interruptible; returns promptly when interrupt fires. diff --git a/tun/netstack/lneto_tuncount.go b/tun/netstack/lneto_tuncount.go new file mode 100644 index 000000000..b499df924 --- /dev/null +++ b/tun/netstack/lneto_tuncount.go @@ -0,0 +1,30 @@ +//go:build debugheaplog + +package netstack + +import "fmt" + +// TUN-boundary packet counters, compiled in only under the debugheaplog tag +// (the start.sh -debug flag). They disambiguate a one-way data path: egress is +// what the stack hands to WireGuard (stack -> WG), ingress is what WireGuard +// feeds back into the stack (WG -> stack). A stuck SYN-SENT with egress>0 and +// ingress==0 means replies never reach the TUN; egress==0 means the stack is +// not draining its tx. Single counters are fine: js/wasm is single-threaded. +var ( + egressPkts, egressBytes int + ingressPkts, ingressBytes int +) + +// countEgress records one non-empty packet pulled from the stack toward WireGuard. +func countEgress(n int) { + egressPkts++ + egressBytes += n + fmt.Printf("[TUNCOUNT] egress pkts=%d bytes=%d last=%d\n", egressPkts, egressBytes, n) +} + +// countIngress records one IP packet handed from WireGuard into the stack. +func countIngress(n int) { + ingressPkts++ + ingressBytes += n + fmt.Printf("[TUNCOUNT] ingress pkts=%d bytes=%d last=%d\n", ingressPkts, ingressBytes, n) +} diff --git a/tun/netstack/lneto_tuncount_off.go b/tun/netstack/lneto_tuncount_off.go new file mode 100644 index 000000000..6bdabf943 --- /dev/null +++ b/tun/netstack/lneto_tuncount_off.go @@ -0,0 +1,7 @@ +//go:build !debugheaplog + +package netstack + +// No-op TUN-boundary counters for non-debug builds; calls are inlined away. +func countEgress(int) {} +func countIngress(int) {} diff --git a/tun/netstack/net_debug.go b/tun/netstack/net_debug.go new file mode 100644 index 000000000..ceac04160 --- /dev/null +++ b/tun/netstack/net_debug.go @@ -0,0 +1,31 @@ +//go:build netstackdebug + +package netstack + +import ( + "os" + "sync" + "time" + + "github.com/soypat/lneto/x/xnet" +) + +var ( + _pcap xnet.CapturePrinter + _pcapOnce sync.Once +) + +func debugIPPacket(egress bool, pkt []byte) { + _pcapOnce.Do(func() { + _pcap.Configure(os.Stdout, xnet.CapturePrinterConfig{ + NamespaceWidth: 3, + TimePrecision: 4, + Now: time.Now, + }) + }) + if egress { + _pcap.PrintIP("OUT", pkt) + } else { + _pcap.PrintIP("IN ", pkt) + } +} diff --git a/tun/netstack/net_debugoff.go b/tun/netstack/net_debugoff.go new file mode 100644 index 000000000..2b42ed8bb --- /dev/null +++ b/tun/netstack/net_debugoff.go @@ -0,0 +1,5 @@ +//go:build !netstackdebug + +package netstack + +func debugIPPacket(egress bool, pkt []byte) {} From 3b1390c8ddd76724050a531d5333ec801fe99e80 Mon Sep 17 00:00:00 2001 From: TuteMthCD Date: Sat, 20 Jun 2026 16:32:54 -0300 Subject: [PATCH 07/12] added: retying options --- tun/netstack/lneto.go | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tun/netstack/lneto.go b/tun/netstack/lneto.go index 53dc52a3a..4024d9070 100644 --- a/tun/netstack/lneto.go +++ b/tun/netstack/lneto.go @@ -137,6 +137,8 @@ func CreateNetTUNLneto(localAddresses, dnsServers []netip.Addr, mtu int) (tun.De ClosingTimeout: 10 * time.Second, NewBackoff: newTCPBackoff, // required: StackGo panics if nil. }, + TCPDialTimeout: time.Second, + TCPDialRetries: 30, }) dev.events <- tun.EventUp return dev, &Net{stack: dev}, nil From 7900fcc5ed2e1d09f541c9eef26e8a3f03001827 Mon Sep 17 00:00:00 2001 From: Patricio Whittingslow Date: Wed, 1 Jul 2026 14:48:02 -0300 Subject: [PATCH 08/12] begin adding dns over tcp --- tun/netstack/lneto.go | 132 +++++++++++++++++++++++++++++++++++-- tun/netstack/lneto_test.go | 115 ++++++++++++++++++++++++++++++++ 2 files changed, 243 insertions(+), 4 deletions(-) diff --git a/tun/netstack/lneto.go b/tun/netstack/lneto.go index 4024d9070..12d87870b 100644 --- a/tun/netstack/lneto.go +++ b/tun/netstack/lneto.go @@ -11,10 +11,12 @@ import ( "encoding/binary" "errors" "fmt" + "io" "net" "net/netip" "os" "runtime" + "slices" "strings" "sync" "syscall" @@ -79,6 +81,12 @@ func CreateNetTUNLneto(localAddresses, dnsServers []netip.Addr, mtu int) (tun.De } } + // DNS transport: default to TCP (DNS over TCP) so large responses are not + // truncated; NB_LNETO_DNS_UDP forces the legacy UDP path. dnsServer is the + // IPv4-preferred server the TCP path dials directly. + dev.dnsServer = dnsServer + dev.dnsUDP = os.Getenv("NB_LNETO_DNS_UDP") != "" + var randSeed int64 if err := binary.Read(crand.Reader, binary.LittleEndian, &randSeed); err != nil { return nil, nil, fmt.Errorf("CreateNetTUNLneto: rand seed: %w", err) @@ -180,7 +188,16 @@ type lnetoStack struct { mtu int dnsServers []netip.Addr + dnsServer netip.Addr // IPv4-preferred server used by the TCP DNS path. + dnsUDP bool // when true resolve over UDP (legacy path) instead of TCP. hasV4, hasV6 bool + + // msgaux is reused across DNS lookups to retain the message's slice backing + // arrays (see dns.Message.Reset). msgauxMu guards it: each user is a single + // self-contained critical section (build or decode) with no network I/O held. + msgauxMu sync.Mutex + msgaux dns.Message + msgbuf []byte } type event struct{} @@ -445,19 +462,17 @@ func (n *lnetoStack) LookupContextHost(ctx context.Context, host string) ([]stri timeout = rem } } - blk := n.sa.StackBlocking(n.backoff) - var addrsV4, addrsV6 []netip.Addr var lastErr error if n.hasV4 { - if a, err := blk.DoLookupIPType(host, timeout, dns.TypeA); err != nil { + if a, err := n.lookupIPType(ctx, host, dns.TypeA, timeout); err != nil { lastErr = dnsError(host, err) } else { addrsV4 = a } } if n.hasV6 { - if a, err := blk.DoLookupIPType(host, timeout, dns.TypeAAAA); err != nil { + if a, err := n.lookupIPType(ctx, host, dns.TypeAAAA, timeout); err != nil { if lastErr == nil { lastErr = dnsError(host, err) } @@ -485,3 +500,112 @@ func (n *lnetoStack) LookupContextHost(ctx context.Context, host string) ([]stri } return out, nil } + +var ( + errNoDNSServer = errors.New("no DNS server configured") + errDNSMessageTooLong = errors.New("DNS message exceeds TCP length prefix maximum") +) + +// lookupIPType resolves host for a single record type, selecting the DNS transport: +// TCP by default (see CreateNetTUNLneto) or the legacy UDP path when dnsUDP is set. +func (n *lnetoStack) lookupIPType(ctx context.Context, host string, qtype dns.Type, timeout time.Duration) ([]netip.Addr, error) { + if n.dnsUDP { + return n.sa.StackBlocking(n.backoff).DoLookupIPType(host, timeout, qtype) + } + return n.lookupTCP(ctx, host, qtype, timeout) +} + +// lookupTCP performs a DNS query over TCP (RFC 1035 §4.2.2: each message is preceded +// by a two-byte length field), dialing the configured DNS server on port 53 using the +// lneto TCP socket. It mirrors the gvisor backend's dnsStreamRoundTrip but builds and +// parses the message with lneto's dns package. IPv4 transport only, matching the UDP path. +func (n *lnetoStack) lookupTCP(ctx context.Context, host string, qtype dns.Type, timeout time.Duration) ([]netip.Addr, error) { + if !n.dnsServer.IsValid() { + return nil, errNoDNSServer + } + name, err := dns.NewName(host) + if err != nil { + return nil, err + } + + var txbuf [2]byte + if _, err := crand.Read(txbuf[:]); err != nil { + return nil, err + } + txid := uint16(n.sa.Prand32()) + + // Build the query message: single question, recursion desired. Reuse msgaux + // (guarded) so its slice memory carries over between lookups. + n.msgauxMu.Lock() + defer n.msgauxMu.Unlock() + n.msgaux.Reset() + n.msgaux.AddQuestions([]dns.Question{{Name: name, Type: qtype, Class: dns.ClassINET}}) + msglen := n.msgaux.Len() + n.msgbuf = slices.Grow(n.msgbuf[:0], int(msglen)+2) + binary.BigEndian.PutUint16(n.msgbuf[:2], msglen) + n.msgbuf, err = n.msgaux.AppendTo(n.msgbuf[2:2], txid, dns.NewClientHeaderFlags(dns.OpCodeQuery, true)) + if err != nil { + + return nil, err + } + + if timeout <= 0 { + timeout = 5 * time.Second + } + dialCtx, cancel := context.WithTimeout(ctx, timeout) + defer cancel() + conn, err := n.dialTCPCtx(dialCtx, netip.AddrPortFrom(n.dnsServer, dns.ServerPort)) + if err != nil { + return nil, err + } + defer conn.Close() + if dl, ok := dialCtx.Deadline(); ok { + if err := conn.SetDeadline(dl); err != nil { + return nil, err + } + } + + // Write the 2-byte length prefix followed by the message in a single write. + if _, err := conn.Write(n.msgbuf); err != nil { + return nil, err + } + + // Read the length-prefixed response. + var lenbuf [2]byte + if _, err := io.ReadFull(conn, lenbuf[:]); err != nil { + return nil, err + } + rlen := binary.BigEndian.Uint16(lenbuf[:]) + resp := make([]byte, rlen) + if _, err := io.ReadFull(conn, resp); err != nil { + return nil, err + } + + // Validate the response header against our query. + f, err := dns.NewFrame(resp) + if err != nil { + return nil, err + } + if f.TxID() != txid || !f.Flags().IsResponse() { + return nil, errInvalidDNSResponse + } + if rcode := f.Flags().ResponseCode(); rcode != 0 { + return nil, rcode + } + + dst := make([]netip.Addr, 16) + n.msgauxMu.Lock() + n.msgaux.Reset() + n.msgaux.LimitResourceDecoding(1, 16, 0, 4) // allow up to 16 answers to decode. + var nAns uint16 + if _, incompleteButOK, derr := n.msgaux.Decode(resp); derr != nil && !incompleteButOK { + err = derr + } else { + nAns, err = n.msgaux.WriteAnswers(dst, host) + } + n.msgauxMu.Unlock() + if err != nil && err != lneto.ErrExhausted { + return nil, err + } + return dst[:nAns], nil +} diff --git a/tun/netstack/lneto_test.go b/tun/netstack/lneto_test.go index 634c258d9..c0e01a92f 100644 --- a/tun/netstack/lneto_test.go +++ b/tun/netstack/lneto_test.go @@ -9,7 +9,9 @@ package netstack import ( "context" + "encoding/binary" "errors" + "io" "net" "net/netip" "os" @@ -17,6 +19,7 @@ import ( "testing" "time" + "github.com/soypat/lneto/dns" "golang.zx2c4.com/wireguard/tun" ) @@ -300,3 +303,115 @@ func testTCPEcho(t *testing.T, addrA, addrB string) { t.Fatal("server:", err) } } + +// TestNet2_DNSOverTCP wires two lneto stacks back-to-back and resolves a name over +// TCP: netA (client) runs LookupContextHost with the default TCP transport; netB +// hosts a minimal DNS-over-TCP responder on port 53 that answers with a fixed A +// record. It exercises the new lookupTCP path end-to-end (dial, 2-byte length +// framing per RFC 1035 §4.2.2, message build/parse via lneto's dns package). +func TestNet2_DNSOverTCP(t *testing.T) { + const ( + addrA = "10.0.0.1" + addrB = "10.0.0.2" // also the DNS server address for netA. + host = "example.com" + ) + wantIP := netip.MustParseAddr("1.2.3.4") + + devA, netA, err := CreateNetTUNLneto( + []netip.Addr{netip.MustParseAddr(addrA)}, + []netip.Addr{netip.MustParseAddr(addrB)}, // dnsServers → dnsServer = addrB. + 1500, + ) + if err != nil { + t.Fatal(err) + } + // Default transport must be TCP (not the legacy UDP path). + if stA := devA.(*lnetoStack); stA.dnsUDP { + t.Fatal("expected TCP DNS transport by default") + } + devB, netB, err := CreateNetTUNLneto([]netip.Addr{netip.MustParseAddr(addrB)}, nil, 1500) + if err != nil { + t.Fatal(err) + } + <-devA.Events() + <-devB.Events() + + var pumps sync.WaitGroup + pumps.Add(2) + go func() { defer pumps.Done(); pump(devA, devB) }() + go func() { defer pumps.Done(); pump(devB, devA) }() + + ln, err := netB.ListenTCPAddrPort(netip.AddrPortFrom(netip.MustParseAddr(addrB), dns.ServerPort)) + if err != nil { + t.Fatal("listen dns:", err) + } + defer func() { + devA.Close() + devB.Close() + pumps.Wait() + ln.Close() + }() + + srvDone := make(chan error, 1) + go func() { srvDone <- serveDNSOverTCP(ln, host, wantIP) }() + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + addrs, err := netA.LookupContextHost(ctx, host) + if err != nil { + t.Fatal("lookup:", err) + } + if len(addrs) != 1 || addrs[0] != wantIP.String() { + t.Fatalf("lookup result mismatch: want [%s] got %v", wantIP, addrs) + } + if err := <-srvDone; err != nil { + t.Fatal("dns server:", err) + } +} + +// serveDNSOverTCP accepts one DNS-over-TCP connection, reads the length-prefixed +// query, and replies with a single A record (host → ip). It mirrors what netbird's +// TCP resolver does, minimally, so the client's lookupTCP path can be tested. +func serveDNSOverTCP(ln TCPListener, host string, ip netip.Addr) error { + conn, err := ln.Accept() + if err != nil { + return err + } + defer conn.Close() + conn.SetDeadline(time.Now().Add(5 * time.Second)) + + var lenbuf [2]byte + if _, err := io.ReadFull(conn, lenbuf[:]); err != nil { + return err + } + query := make([]byte, binary.BigEndian.Uint16(lenbuf[:])) + if _, err := io.ReadFull(conn, query); err != nil { + return err + } + f, err := dns.NewFrame(query) + if err != nil { + return err + } + txid := f.TxID() + + name, err := dns.NewName(host) + if err != nil { + return err + } + var resp dns.Message + resp.Reset() + resp.AddQuestions([]dns.Question{{Name: name, Type: dns.TypeA, Class: dns.ClassINET}}) + ip4 := ip.As4() + resp.Answers = append(resp.Answers, dns.NewResource(name, dns.TypeA, dns.ClassINET, 300, ip4[:])) + // Response header: QR=1 (response), RA=1 (recursion available), RCODE=0. + const respFlags dns.HeaderFlags = 1<<15 | 1<<7 + msg, err := resp.AppendTo(nil, txid, respFlags) + if err != nil { + return err + } + framed := make([]byte, 2+len(msg)) + binary.BigEndian.PutUint16(framed[:2], uint16(len(msg))) + copy(framed[2:], msg) + _, err = conn.Write(framed) + return err +} From db0f25343af3681c781ec8dc32ffbb7bdb794209 Mon Sep 17 00:00:00 2001 From: Patricio Whittingslow Date: Wed, 1 Jul 2026 15:06:28 -0300 Subject: [PATCH 09/12] claude suggests --- tun/netstack/lneto.go | 44 ++++++++++++++++++++----------------------- 1 file changed, 20 insertions(+), 24 deletions(-) diff --git a/tun/netstack/lneto.go b/tun/netstack/lneto.go index 12d87870b..765775b56 100644 --- a/tun/netstack/lneto.go +++ b/tun/netstack/lneto.go @@ -198,6 +198,7 @@ type lnetoStack struct { msgauxMu sync.Mutex msgaux dns.Message msgbuf []byte + addrbuf [16]netip.Addr } type event struct{} @@ -528,24 +529,24 @@ func (n *lnetoStack) lookupTCP(ctx context.Context, host string, qtype dns.Type, return nil, err } - var txbuf [2]byte - if _, err := crand.Read(txbuf[:]); err != nil { - return nil, err - } txid := uint16(n.sa.Prand32()) - // Build the query message: single question, recursion desired. Reuse msgaux - // (guarded) so its slice memory carries over between lookups. + // The shared msgaux/msgbuf/addrbuf scratch is used across the whole round trip + // (build → write → read → decode), so hold the lock for the entire function. + // This serializes DNS lookups, which is fine: A and AAAA already run sequentially. n.msgauxMu.Lock() defer n.msgauxMu.Unlock() + + // Build the query message: single question, recursion desired. Reuse msgaux/msgbuf + // so their slice memory carries over between lookups. Layout: 2-byte length prefix + // followed by the message body. n.msgaux.Reset() n.msgaux.AddQuestions([]dns.Question{{Name: name, Type: qtype, Class: dns.ClassINET}}) msglen := n.msgaux.Len() - n.msgbuf = slices.Grow(n.msgbuf[:0], int(msglen)+2) + n.msgbuf = slices.Grow(n.msgbuf[:0], int(msglen)+2)[:2] binary.BigEndian.PutUint16(n.msgbuf[:2], msglen) - n.msgbuf, err = n.msgaux.AppendTo(n.msgbuf[2:2], txid, dns.NewClientHeaderFlags(dns.OpCodeQuery, true)) + n.msgbuf, err = n.msgaux.AppendTo(n.msgbuf, txid, dns.NewClientHeaderFlags(dns.OpCodeQuery, true)) if err != nil { - return nil, err } @@ -571,18 +572,17 @@ func (n *lnetoStack) lookupTCP(ctx context.Context, host string, qtype dns.Type, } // Read the length-prefixed response. - var lenbuf [2]byte - if _, err := io.ReadFull(conn, lenbuf[:]); err != nil { + if _, err := io.ReadFull(conn, n.msgbuf[:2]); err != nil { return nil, err } - rlen := binary.BigEndian.Uint16(lenbuf[:]) - resp := make([]byte, rlen) - if _, err := io.ReadFull(conn, resp); err != nil { + rlen := binary.BigEndian.Uint16(n.msgbuf[:2]) + n.msgbuf = slices.Grow(n.msgbuf[:0], int(rlen))[:rlen] + if _, err := io.ReadFull(conn, n.msgbuf); err != nil { return nil, err } - + msg := n.msgbuf // Validate the response header against our query. - f, err := dns.NewFrame(resp) + f, err := dns.NewFrame(msg) if err != nil { return nil, err } @@ -592,20 +592,16 @@ func (n *lnetoStack) lookupTCP(ctx context.Context, host string, qtype dns.Type, if rcode := f.Flags().ResponseCode(); rcode != 0 { return nil, rcode } - - dst := make([]netip.Addr, 16) - n.msgauxMu.Lock() n.msgaux.Reset() - n.msgaux.LimitResourceDecoding(1, 16, 0, 4) // allow up to 16 answers to decode. + n.msgaux.LimitResourceDecoding(1, uint16(len(n.addrbuf)), 0, 4) // allow up to 16 answers to decode. var nAns uint16 - if _, incompleteButOK, derr := n.msgaux.Decode(resp); derr != nil && !incompleteButOK { + if _, incompleteButOK, derr := n.msgaux.Decode(msg); derr != nil && !incompleteButOK { err = derr } else { - nAns, err = n.msgaux.WriteAnswers(dst, host) + nAns, err = n.msgaux.WriteAnswers(n.addrbuf[:], host) } - n.msgauxMu.Unlock() if err != nil && err != lneto.ErrExhausted { return nil, err } - return dst[:nAns], nil + return slices.Clone(n.addrbuf[:nAns]), nil } From 846407ae204ce4d03485a8c4c5e38085b594467f Mon Sep 17 00:00:00 2001 From: Patricio Whittingslow Date: Wed, 1 Jul 2026 15:16:30 -0300 Subject: [PATCH 10/12] dnsScratch --- tun/netstack/lneto.go | 111 ++++++++++++++++++++++++++---------------- 1 file changed, 69 insertions(+), 42 deletions(-) diff --git a/tun/netstack/lneto.go b/tun/netstack/lneto.go index 765775b56..76007b503 100644 --- a/tun/netstack/lneto.go +++ b/tun/netstack/lneto.go @@ -192,13 +192,8 @@ type lnetoStack struct { dnsUDP bool // when true resolve over UDP (legacy path) instead of TCP. hasV4, hasV6 bool - // msgaux is reused across DNS lookups to retain the message's slice backing - // arrays (see dns.Message.Reset). msgauxMu guards it: each user is a single - // self-contained critical section (build or decode) with no network I/O held. - msgauxMu sync.Mutex - msgaux dns.Message - msgbuf []byte - addrbuf [16]netip.Addr + // dnsScratch holds the reusable buffers for the DNS-over-TCP lookup path. + dnsScratch dnsScratch } type event struct{} @@ -502,10 +497,7 @@ func (n *lnetoStack) LookupContextHost(ctx context.Context, host string) ([]stri return out, nil } -var ( - errNoDNSServer = errors.New("no DNS server configured") - errDNSMessageTooLong = errors.New("DNS message exceeds TCP length prefix maximum") -) +var errNoDNSServer = errors.New("no DNS server configured") // lookupIPType resolves host for a single record type, selecting the DNS transport: // TCP by default (see CreateNetTUNLneto) or the legacy UDP path when dnsUDP is set. @@ -524,28 +516,17 @@ func (n *lnetoStack) lookupTCP(ctx context.Context, host string, qtype dns.Type, if !n.dnsServer.IsValid() { return nil, errNoDNSServer } - name, err := dns.NewName(host) - if err != nil { - return nil, err - } txid := uint16(n.sa.Prand32()) - // The shared msgaux/msgbuf/addrbuf scratch is used across the whole round trip - // (build → write → read → decode), so hold the lock for the entire function. + // The shared dnsScratch buffers are reused across the whole round trip + // (build → write → read → parse), so hold the lock for the entire function. // This serializes DNS lookups, which is fine: A and AAAA already run sequentially. - n.msgauxMu.Lock() - defer n.msgauxMu.Unlock() - - // Build the query message: single question, recursion desired. Reuse msgaux/msgbuf - // so their slice memory carries over between lookups. Layout: 2-byte length prefix - // followed by the message body. - n.msgaux.Reset() - n.msgaux.AddQuestions([]dns.Question{{Name: name, Type: qtype, Class: dns.ClassINET}}) - msglen := n.msgaux.Len() - n.msgbuf = slices.Grow(n.msgbuf[:0], int(msglen)+2)[:2] - binary.BigEndian.PutUint16(n.msgbuf[:2], msglen) - n.msgbuf, err = n.msgaux.AppendTo(n.msgbuf, txid, dns.NewClientHeaderFlags(dns.OpCodeQuery, true)) + s := &n.dnsScratch + s.mu.Lock() + defer s.mu.Unlock() + + framed, err := s.buildQuery(host, qtype, txid) if err != nil { return nil, err } @@ -567,21 +548,67 @@ func (n *lnetoStack) lookupTCP(ctx context.Context, host string, qtype dns.Type, } // Write the 2-byte length prefix followed by the message in a single write. - if _, err := conn.Write(n.msgbuf); err != nil { + if _, err := conn.Write(framed); err != nil { + return nil, err + } + + msg, err := s.readResponse(conn) + if err != nil { + return nil, err + } + return s.parseAnswers(msg, txid, host) +} + +// dnsScratch holds reusable buffers for building and parsing DNS-over-TCP +// messages, retaining slice backing arrays across lookups. Not safe for +// concurrent use: callers must hold mu across an entire build→read→parse +// sequence because buf is reused for both the query and the response. +type dnsScratch struct { + mu sync.Mutex + msg dns.Message + buf []byte // length-prefixed wire buffer + addrs [16]netip.Addr // decode target +} + +// buildQuery resets the scratch, encodes a single-question recursion-desired +// query for host/qtype with txid, and returns the 2-byte-length-prefixed wire +// bytes (aliases buf; valid until the next scratch use). Caller holds mu. +func (s *dnsScratch) buildQuery(host string, qtype dns.Type, txid uint16) ([]byte, error) { + name, err := dns.NewName(host) + if err != nil { return nil, err } + // Layout: 2-byte length prefix followed by the message body. + s.msg.Reset() + s.msg.AddQuestions([]dns.Question{{Name: name, Type: qtype, Class: dns.ClassINET}}) + msglen := s.msg.Len() + s.buf = slices.Grow(s.buf[:0], int(msglen)+2)[:2] + binary.BigEndian.PutUint16(s.buf[:2], msglen) + s.buf, err = s.msg.AppendTo(s.buf, txid, dns.NewClientHeaderFlags(dns.OpCodeQuery, true)) + if err != nil { + return nil, err + } + return s.buf, nil +} - // Read the length-prefixed response. - if _, err := io.ReadFull(conn, n.msgbuf[:2]); err != nil { +// readResponse reads a length-prefixed DNS response from r into buf and returns +// the message bytes (aliases buf; valid until the next scratch use). Caller holds mu. +func (s *dnsScratch) readResponse(r io.Reader) ([]byte, error) { + if _, err := io.ReadFull(r, s.buf[:2]); err != nil { return nil, err } - rlen := binary.BigEndian.Uint16(n.msgbuf[:2]) - n.msgbuf = slices.Grow(n.msgbuf[:0], int(rlen))[:rlen] - if _, err := io.ReadFull(conn, n.msgbuf); err != nil { + rlen := binary.BigEndian.Uint16(s.buf[:2]) + s.buf = slices.Grow(s.buf[:0], int(rlen))[:rlen] + if _, err := io.ReadFull(r, s.buf); err != nil { return nil, err } - msg := n.msgbuf - // Validate the response header against our query. + return s.buf, nil +} + +// parseAnswers validates msg against txid (TxID, IsResponse, ResponseCode), +// decodes it, and returns a freshly cloned slice of answer addresses for host. +// Caller holds mu. Returns the DNS rcode as error when non-zero. +func (s *dnsScratch) parseAnswers(msg []byte, txid uint16, host string) ([]netip.Addr, error) { f, err := dns.NewFrame(msg) if err != nil { return nil, err @@ -592,16 +619,16 @@ func (n *lnetoStack) lookupTCP(ctx context.Context, host string, qtype dns.Type, if rcode := f.Flags().ResponseCode(); rcode != 0 { return nil, rcode } - n.msgaux.Reset() - n.msgaux.LimitResourceDecoding(1, uint16(len(n.addrbuf)), 0, 4) // allow up to 16 answers to decode. + s.msg.Reset() + s.msg.LimitResourceDecoding(1, uint16(len(s.addrs)), 0, 4) // allow up to 16 answers to decode. var nAns uint16 - if _, incompleteButOK, derr := n.msgaux.Decode(msg); derr != nil && !incompleteButOK { + if _, incompleteButOK, derr := s.msg.Decode(msg); derr != nil && !incompleteButOK { err = derr } else { - nAns, err = n.msgaux.WriteAnswers(n.addrbuf[:], host) + nAns, err = s.msg.WriteAnswers(s.addrs[:], host) } if err != nil && err != lneto.ErrExhausted { return nil, err } - return slices.Clone(n.addrbuf[:nAns]), nil + return slices.Clone(s.addrs[:nAns]), nil } From bc01e7c03c81a128303e0b1dd33f4ba13f75d1ac Mon Sep 17 00:00:00 2001 From: Patricio Whittingslow Date: Sat, 4 Jul 2026 14:24:51 -0300 Subject: [PATCH 11/12] dns error and syscall fixes --- .gitignore | 3 ++- tun/netstack/gvisor.go | 12 ++++++++++++ tun/netstack/lneto.go | 34 +++++++++++----------------------- tun/netstack/tinygo.go | 14 ++++++++++++++ 4 files changed, 39 insertions(+), 24 deletions(-) create mode 100644 tun/netstack/tinygo.go diff --git a/.gitignore b/.gitignore index 66abd3cc2..0e4f040c4 100644 --- a/.gitignore +++ b/.gitignore @@ -1,2 +1,3 @@ wireguard-go -*LNETO_EQUIVALENCE.md \ No newline at end of file +*LNETO_EQUIVALENCE.md +*_local.* \ No newline at end of file diff --git a/tun/netstack/gvisor.go b/tun/netstack/gvisor.go index 50e6ff789..f5198cbbc 100644 --- a/tun/netstack/gvisor.go +++ b/tun/netstack/gvisor.go @@ -1,3 +1,5 @@ +//go:build !tinygo + /* SPDX-License-Identifier: MIT * * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved. @@ -804,3 +806,13 @@ func (tnet *netTun) LookupContextHost(ctx context.Context, host string) ([]strin } return saddrs, nil } + +// dnsError wraps a lookup failure as a *net.DNSError, flagging timeouts when the +// underlying error reports them. Mirrors the error shape produced by the gvisor Net. +func dnsError(host, server string, err error, isNotFound bool) *net.DNSError { + de := &net.DNSError{Err: err.Error(), Name: host, Server: server, IsNotFound: isNotFound} + if nerr, ok := err.(net.Error); ok && nerr.Timeout() { + de.IsTimeout = true + } + return de +} diff --git a/tun/netstack/lneto.go b/tun/netstack/lneto.go index 76007b503..b043fa712 100644 --- a/tun/netstack/lneto.go +++ b/tun/netstack/lneto.go @@ -12,14 +12,12 @@ import ( "errors" "fmt" "io" - "net" "net/netip" "os" "runtime" "slices" "strings" "sync" - "syscall" "time" "github.com/soypat/lneto" @@ -365,10 +363,10 @@ func (n *lnetoStack) socket(ctx context.Context, proto string, sotype int, laddr fam = laddr } network := proto + "4" - family := syscall.AF_INET + family := xnet.AF_INET if fam.Addr().Is6() { network = proto + "6" - family = syscall.AF_INET6 + family = xnet.AF_INET6 } n.interrupt() defer n.interrupt() @@ -376,7 +374,7 @@ func (n *lnetoStack) socket(ctx context.Context, proto string, sotype int, laddr } func (n *lnetoStack) dialTCPCtx(ctx context.Context, addr netip.AddrPort) (TCPConn, error) { - v, err := n.socket(ctx, "tcp", syscall.SOCK_STREAM, netip.AddrPort{}, addr) + v, err := n.socket(ctx, "tcp", xnet.SOCK_STREAM, netip.AddrPort{}, addr) return socketResult[TCPConn](v, err) } @@ -391,19 +389,19 @@ func (n *lnetoStack) DialTCPAddrPort(addr netip.AddrPort) (TCPConn, error) { // --- TCP listener --- func (n *lnetoStack) ListenTCPAddrPort(addr netip.AddrPort) (TCPListener, error) { - v, err := n.socket(context.Background(), "tcp", syscall.SOCK_STREAM, addr, netip.AddrPort{}) + v, err := n.socket(context.Background(), "tcp", xnet.SOCK_STREAM, addr, netip.AddrPort{}) return socketResult[TCPListener](v, err) } // --- UDP --- func (n *lnetoStack) ListenUDPAddrPort(laddr netip.AddrPort) (UDPConn, error) { - v, err := n.socket(context.Background(), "udp", syscall.SOCK_DGRAM, laddr, netip.AddrPort{}) + v, err := n.socket(context.Background(), "udp", xnet.SOCK_DGRAM, laddr, netip.AddrPort{}) return socketResult[UDPConn](v, err) } func (n *lnetoStack) DialUDPAddrPort(laddr, raddr netip.AddrPort) (UDPConn, error) { - v, err := n.socket(context.Background(), "udp", syscall.SOCK_DGRAM, laddr, raddr) + v, err := n.socket(context.Background(), "udp", xnet.SOCK_DGRAM, laddr, raddr) return socketResult[UDPConn](v, err) } @@ -419,16 +417,6 @@ func (n *lnetoStack) ListenPingAddr(_ netip.Addr) (PingConn, error) { // --- DNS --- -// dnsError wraps a lookup failure as a *net.DNSError, flagging timeouts when the -// underlying error reports them. Mirrors the error shape produced by the gvisor Net. -func dnsError(host string, err error) *net.DNSError { - de := &net.DNSError{Err: err.Error(), Name: host} - if nerr, ok := err.(net.Error); ok && nerr.Timeout() { - de.IsTimeout = true - } - return de -} - // LookupContextHost resolves host to a list of IP strings, matching the behaviour of // the gvisor Net: literal IPs (with IPv6 zone stripping) pass through; empty or // non-domain hosts and stacks with no address family return an IsNotFound DNSError; @@ -436,7 +424,7 @@ func dnsError(host string, err error) *net.DNSError { // results are ordered first (no RFC 6724). func (n *lnetoStack) LookupContextHost(ctx context.Context, host string) ([]string, error) { if host == "" || (!n.hasV4 && !n.hasV6) { - return nil, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true} + return nil, dnsError(host, "", errNoSuchHost, true) } // Strip any IPv6 zone before attempting to parse a literal address. zlen := len(host) @@ -449,7 +437,7 @@ func (n *lnetoStack) LookupContextHost(ctx context.Context, host string) ([]stri return []string{ip.String()}, nil } if !isDomainName(host) { - return nil, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true} + return nil, dnsError(host, "", errNoSuchHost, true) } timeout := 5 * time.Second @@ -462,7 +450,7 @@ func (n *lnetoStack) LookupContextHost(ctx context.Context, host string) ([]stri var lastErr error if n.hasV4 { if a, err := n.lookupIPType(ctx, host, dns.TypeA, timeout); err != nil { - lastErr = dnsError(host, err) + lastErr = dnsError(host, "", err, false) } else { addrsV4 = a } @@ -470,7 +458,7 @@ func (n *lnetoStack) LookupContextHost(ctx context.Context, host string) ([]stri if n.hasV6 { if a, err := n.lookupIPType(ctx, host, dns.TypeAAAA, timeout); err != nil { if lastErr == nil { - lastErr = dnsError(host, err) + lastErr = dnsError(host, "", err, false) } } else { addrsV6 = a @@ -488,7 +476,7 @@ func (n *lnetoStack) LookupContextHost(ctx context.Context, host string) ([]stri if lastErr != nil { return nil, lastErr } - return nil, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true} + return nil, dnsError(host, "", errNoSuchHost, true) } out := make([]string, len(addrs)) for i, a := range addrs { diff --git a/tun/netstack/tinygo.go b/tun/netstack/tinygo.go new file mode 100644 index 000000000..db9513bc3 --- /dev/null +++ b/tun/netstack/tinygo.go @@ -0,0 +1,14 @@ +//go:build tinygo + +package netstack + +import ( + "fmt" +) + +// dnsError wraps a lookup failure. TinyGo's net package has no net.DNSError, so +// we produce a plain wrapped error mirroring the shape of the gvisor path. The +// isNotFound flag is accepted for signature parity but not otherwise encoded. +func dnsError(host, server string, err error, isNotFound bool) error { + return fmt.Errorf("%w(%s)@ %s", err, server, host) +} From 33c26a62017843a0b366b37c876f6312359cc207 Mon Sep 17 00:00:00 2001 From: Patricio Whittingslow Date: Mon, 6 Jul 2026 12:16:19 -0300 Subject: [PATCH 12/12] local gitignore --- .gitignore | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/.gitignore b/.gitignore index 0e4f040c4..80d3953a7 100644 --- a/.gitignore +++ b/.gitignore @@ -1,3 +1,4 @@ wireguard-go *LNETO_EQUIVALENCE.md -*_local.* \ No newline at end of file +*_local.* +local \ No newline at end of file