Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 4 additions & 1 deletion cmd/clawpatrol/daemon_linux.go
Original file line number Diff line number Diff line change
Expand Up @@ -566,8 +566,11 @@ func (d *daemon) handle(c net.Conn) {
return
}
defer gvStack.Close()
startTunBridge(tunFile, gvEp, d.transport)
sessionCtx, cancelSession := context.WithCancel(context.Background())
defer cancelSession()
startTunBridge(tunFile, gvEp)
enableTransportTCPForwarder(gvStack, d.transport)
enableTransportUDPForwarder(sessionCtx, gvStack, d.transport, runUDPIdleTimeout)

// 4. Tell the client the bridge is up.
if _, err := io.WriteString(c, "ATTACHED\n"); err != nil {
Expand Down
304 changes: 191 additions & 113 deletions cmd/clawpatrol/daemon_session_linux.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,11 +10,14 @@ package main

import (
"context"
"errors"
"fmt"
"net"
"net/netip"
"os"
"strconv"
"sync"
"sync/atomic"
"time"

"gvisor.dev/gvisor/pkg/buffer"
Expand All @@ -26,6 +29,7 @@ import (
"gvisor.dev/gvisor/pkg/tcpip/network/ipv6"
"gvisor.dev/gvisor/pkg/tcpip/stack"
"gvisor.dev/gvisor/pkg/tcpip/transport/tcp"
"gvisor.dev/gvisor/pkg/tcpip/transport/udp"
"gvisor.dev/gvisor/pkg/waiter"
)

Expand All @@ -35,6 +39,27 @@ import (
// child-side TUN doesn't need to cap.
const runStackTunMTU = 65535

const runUDPIdleTimeout = 2 * time.Minute

// runUDPFlowLimit bounds endpoint, connection, and goroutine state per run
// session. 256 permits substantial DNS/QUIC concurrency while limiting a
// session to at most 512 checked-out 64 KiB relay buffers (32 MiB).
const runUDPFlowLimit = 256

var runUDPRelayBufferPool = sync.Pool{New: func() any {
buf := make([]byte, runStackTunMTU)
return &buf
}}

func getRunUDPRelayBuffer() []byte {
return *runUDPRelayBufferPool.Get().(*[]byte)
}

func putRunUDPRelayBuffer(buf []byte) {
buf = buf[:runStackTunMTU]
runUDPRelayBufferPool.Put(&buf)
}

// newRunStack creates a gVisor TCP/IP stack bound to localIP, which
// is the transport's underlay address (tsnet 100.x.x.x or wg /32).
// Promiscuous + spoofing enabled so the stack accepts inbound
Expand All @@ -47,7 +72,7 @@ func newRunStack(localIP netip.Addr) (*stack.Stack, *channel.Endpoint, error) {
ipv4.NewProtocol, ipv6.NewProtocol,
},
TransportProtocols: []stack.TransportProtocolFactory{
tcp.NewProtocol,
tcp.NewProtocol, udp.NewProtocol,
},
HandleLocal: false,
})
Expand Down Expand Up @@ -105,19 +130,26 @@ func (b *runTunBridge) WriteNotify() {
}
}

// injectRunTunPacket takes ownership of pkt. InjectInbound is synchronous and
// does not consume the caller's reference, so every dispatch and drop path
// releases it here.
func injectRunTunPacket(ep *channel.Endpoint, version byte, pkt *stack.PacketBuffer) {
defer pkt.DecRef()
switch version {
case 4:
ep.InjectInbound(header.IPv4ProtocolNumber, pkt)
case 6:
ep.InjectInbound(header.IPv6ProtocolNumber, pkt)
default:
// Drop packets with an unknown IP version.
}
}

// startTunBridge registers the outbound notification and starts the
// inbound read loop (TUN fd → gVisor InjectInbound). IPv4 UDP is
// intercepted before injection and forwarded directly via the
// transport so DNS / quic-style flows work without a UDP forwarder
// inside the per-session gVisor stack.
func startTunBridge(tunFile *os.File, ep *channel.Endpoint, transport daemonTransport) {
// inbound read loop (TUN fd → gVisor InjectInbound).
func startTunBridge(tunFile *os.File, ep *channel.Endpoint) {
br := &runTunBridge{tunFile: tunFile, ep: ep}
ep.AddNotify(br)
uf := &runUDPForwarder{
transport: transport,
tunFile: tunFile,
flows: map[udpFlowKey]net.Conn{},
}

go func() {
buf := make([]byte, runStackTunMTU)
Expand All @@ -131,132 +163,178 @@ func startTunBridge(tunFile *os.File, ep *channel.Endpoint, transport daemonTran
}
pkt := make([]byte, n)
copy(pkt, buf[:n])
// Intercept IPv4 UDP before injecting into gVisor TCP stack.
if pkt[0]>>4 == 4 && n > 20 && pkt[9] == 17 {
uf.handle(pkt)
continue
}
pkb := stack.NewPacketBuffer(stack.PacketBufferOptions{
Payload: buffer.MakeWithData(pkt),
})
switch pkt[0] >> 4 {
case 4:
ep.InjectInbound(header.IPv4ProtocolNumber, pkb)
case 6:
ep.InjectInbound(header.IPv6ProtocolNumber, pkb)
default:
pkb.DecRef()
}
injectRunTunPacket(ep, pkt[0]>>4, pkb)
}
}()
}

// runUDPForwarder maintains per-flow transport UDP connections for
// the child netns. Each unique (srcIP:srcPort → dstIP:dstPort)
// 4-tuple gets one transport.Dial("udp", ...) conn.
type runUDPForwarder struct {
transport daemonTransport
tunFile *os.File
mu sync.Mutex
flows map[udpFlowKey]net.Conn
// enableTransportUDPForwarder installs gVisor's dual-stack UDP
// forwarder. gVisor owns packet parsing, checksums and full-tuple flow
// demultiplexing; each accepted endpoint is relayed as datagrams over
// one transport connection.
func enableTransportUDPForwarder(ctx context.Context, s *stack.Stack, transport daemonTransport, idleTimeout time.Duration) {
enableTransportUDPForwarderWithLimit(ctx, s, transport, idleTimeout, runUDPFlowLimit)
}

type udpFlowKey struct {
srcIP, dstIP [4]byte
srcPort, dstPort uint16
func enableTransportUDPForwarderWithLimit(ctx context.Context, s *stack.Stack, transport daemonTransport, idleTimeout time.Duration, flowLimit int) {
slots := make(chan struct{}, flowLimit)
s.SetTransportProtocolHandler(udp.ProtocolNumber, newTransportUDPProtocolHandler(ctx, s, transport, idleTimeout, slots))
}

func (f *runUDPForwarder) handle(pkt []byte) {
ihl := int(pkt[0]&0xf) * 4
if len(pkt) < ihl+8 {
return
// newRunUDPProtocolHandler borrows the stack-owned pkt only for the synchronous
// callback. CreateEndpoint must happen before the callback returns so its queue
// receives the endpoint's own packet clone.
func newRunUDPProtocolHandler(s *stack.Stack, handler udp.ForwarderHandler) func(stack.TransportEndpointID, *stack.PacketBuffer) bool {
return func(id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool {
return handler(udp.NewForwarderRequest(s, id, pkt))
}
var srcIP, dstIP [4]byte
copy(srcIP[:], pkt[12:16])
copy(dstIP[:], pkt[16:20])
srcPort := uint16(pkt[ihl])<<8 | uint16(pkt[ihl+1])
dstPort := uint16(pkt[ihl+2])<<8 | uint16(pkt[ihl+3])
udpLen := int(pkt[ihl+4])<<8 | int(pkt[ihl+5])
if udpLen < 8 || ihl+udpLen > len(pkt) {
return
}
payload := pkt[ihl+8 : ihl+udpLen]
}

key := udpFlowKey{srcIP, dstIP, srcPort, dstPort}
func newTransportUDPProtocolHandler(ctx context.Context, s *stack.Stack, transport daemonTransport, idleTimeout time.Duration, slots chan struct{}) func(stack.TransportEndpointID, *stack.PacketBuffer) bool {
return newRunUDPProtocolHandler(s, func(req *udp.ForwarderRequest) bool {
select {
case slots <- struct{}{}:
default:
// Returning false lets gVisor generate the appropriate unreachable.
return false
}
release := func() { <-slots }

f.mu.Lock()
conn, ok := f.flows[key]
if !ok {
dstAddr := fmt.Sprintf("%d.%d.%d.%d:%d",
dstIP[0], dstIP[1], dstIP[2], dstIP[3], dstPort)
var err error
conn, err = f.transport.Dial(context.Background(), "udp", dstAddr)
if err != nil {
f.mu.Unlock()
return
// CreateEndpoint must remain in the callback, but transport.Dial may
// block for its full timeout and must not stall the sole TUN ingress
// loop. Capture the request ID by value before starting the goroutine.
id := req.ID()
var wq waiter.Queue
ep, terr := req.CreateEndpoint(&wq)
if terr != nil {
release()
return true
}
f.flows[key] = conn
local := gonet.NewUDPConn(&wq, ep)
dstAddr := net.JoinHostPort(id.LocalAddress.String(), strconv.Itoa(int(id.LocalPort)))
go func() {
f.readResponses(conn, dstIP, srcIP, dstPort, srcPort)
f.mu.Lock()
delete(f.flows, key)
f.mu.Unlock()
_ = conn.Close()
dialCtx, cancel := context.WithTimeout(ctx, transportDialTimeout)
defer cancel()
remote, err := transport.Dial(dialCtx, "udp", dstAddr)
if err != nil {
// The callback has already returned true, so no ICMP unreachable
// can be requested here without fabricating a packet. Closing the
// endpoint avoids retaining a black hole and permits tuple redial.
_ = local.Close()
release()
return
}
relayUDPDatagrams(ctx, local, remote, idleTimeout, release)
}()
}
f.mu.Unlock()

_, _ = conn.Write(payload)
return true
})
}

func (f *runUDPForwarder) readResponses(conn net.Conn, srcIP, dstIP [4]byte, srcPort, dstPort uint16) {
buf := make([]byte, 65535)
for {
_ = conn.SetReadDeadline(time.Now().Add(30 * time.Second))
n, err := conn.Read(buf)
if err != nil {
return
}
_, _ = f.tunFile.Write(buildUDPPacket(srcIP, dstIP, srcPort, dstPort, buf[:n]))
func udpIdleRemaining(lastActivity int64, now time.Time, idleTimeout time.Duration) time.Duration {
elapsed := now.Sub(time.Unix(0, lastActivity))
if elapsed < 0 {
return idleTimeout
}
return idleTimeout - elapsed
}

// buildUDPPacket constructs a raw IPv4+UDP packet. UDP checksum is zero
// (optional for IPv4; Linux accepts these from TUN devices).
func buildUDPPacket(srcIP, dstIP [4]byte, srcPort, dstPort uint16, payload []byte) []byte {
udpLen := 8 + len(payload)
ipLen := 20 + udpLen
pkt := make([]byte, ipLen)
pkt[0] = 0x45 // IPv4, IHL=5
pkt[2] = byte(ipLen >> 8)
pkt[3] = byte(ipLen)
pkt[8] = 64 // TTL
pkt[9] = 17 // UDP
copy(pkt[12:16], srcIP[:])
copy(pkt[16:20], dstIP[:])
cs := ipv4Checksum(pkt[:20])
pkt[10] = byte(cs >> 8)
pkt[11] = byte(cs)
pkt[20] = byte(srcPort >> 8)
pkt[21] = byte(srcPort)
pkt[22] = byte(dstPort >> 8)
pkt[23] = byte(dstPort)
pkt[24] = byte(udpLen >> 8)
pkt[25] = byte(udpLen)
// pkt[26:28] = 0 (checksum omitted)
copy(pkt[28:], payload)
return pkt
}

func ipv4Checksum(b []byte) uint16 {
var sum uint32
for i := 0; i+1 < len(b); i += 2 {
sum += uint32(b[i])<<8 | uint32(b[i+1])
func relayUDPDatagrams(ctx context.Context, local, remote net.Conn, idleTimeout time.Duration, release func()) {
defer release()
closeBoth := func() {
_ = local.Close()
_ = remote.Close()
}
activity := make(chan struct{}, 1)
done := make(chan struct{}, 1)
var lastActivity atomic.Int64
lastActivity.Store(time.Now().UnixNano())
signalDone := func() {
select {
case done <- struct{}{}:
default:
}
}
for sum>>16 != 0 {
sum = (sum & 0xffff) + (sum >> 16)
copyDatagrams := func(dst, src net.Conn, boundRead bool) {
buf := getRunUDPRelayBuffer()
defer putRunUDPRelayBuffer(buf)
for {
if boundRead {
now := time.Now()
remaining := udpIdleRemaining(lastActivity.Load(), now, idleTimeout)
if remaining <= 0 {
signalDone()
return
}
if err := src.SetReadDeadline(now.Add(remaining)); err != nil {
signalDone()
return
}
}
n, err := src.Read(buf)
if err != nil {
if boundRead {
var netErr net.Error
if errors.As(err, &netErr) && netErr.Timeout() &&
udpIdleRemaining(lastActivity.Load(), time.Now(), idleTimeout) > 0 {
// Activity in the opposite direction raced this deadline.
// Recalculate from the shared timestamp and keep reading.
continue
}
}
signalDone()
return
}
if _, err := dst.Write(buf[:n]); err != nil {
signalDone()
return
}
lastActivity.Store(time.Now().UnixNano())
select {
case activity <- struct{}{}:
default:
}
}
}
go copyDatagrams(remote, local, false)
go copyDatagrams(local, remote, true)
timer := time.NewTimer(idleTimeout)
defer timer.Stop()
defer closeBoth()
resetFromLastActivity := func() bool {
remaining := udpIdleRemaining(lastActivity.Load(), time.Now(), idleTimeout)
if remaining <= 0 {
return false
}
if !timer.Stop() {
select {
case <-timer.C:
default:
}
}
timer.Reset(remaining)
return true
}
for {
select {
case <-ctx.Done():
return
case <-done:
return
case <-timer.C:
// A successful write can race delivery of timer.C. Re-read the
// race-safe timestamp before deciding the flow is idle.
if !resetFromLastActivity() {
return
}
case <-activity:
if !resetFromLastActivity() {
return
}
}
}
return ^uint16(sum)
}

// transportDialTimeout bounds the upstream transport.Dial while the
Expand Down
Loading
Loading