diff --git a/wgengine/netstack/gro/gro.go b/wgengine/netstack/gro/gro.go index 152b25295..40ec9f325 100644 --- a/wgengine/netstack/gro/gro.go +++ b/wgengine/netstack/gro/gro.go @@ -25,11 +25,14 @@ import ( // !stack.PacketBuffer.RXChecksumValidated, i.e. it satisfies // stack.CapabilityRXChecksumOffload. Other protocols with checksum fields, // e.g. ICMP{v6}, are still validated by gVisor regardless of rx checksum -// offloading capabilities. +// offloading capabilities. IPv4 fragments cannot have their L4 checksums +// validated before reassembly, so only their IPv4 header checksum is validated +// here. func RXChecksumOffload(p *packet.Parsed) *stack.PacketBuffer { var ( pn tcpip.NetworkProtocolNumber csumStart int + fragment bool ) buf := p.Buffer() @@ -45,6 +48,11 @@ func RXChecksumOffload(p *packet.Parsed) *stack.PacketBuffer { if ^tun.Checksum(buf[:csumStart], 0) != 0 { return nil } + // Non-first fragments (FragmentOffset != 0) arrive here with + // p.IPProto == ipproto.Fragment (set by packet.Parsed.Decode), so + // they already skip the L4 checksum check below. We only need to + // catch the first fragment, which still has its real IPProto. + fragment = header.IPv4(buf).More() pn = header.IPv4ProtocolNumber case 6: if len(buf) < header.IPv6FixedHeaderSize { @@ -80,7 +88,7 @@ func RXChecksumOffload(p *packet.Parsed) *stack.PacketBuffer { } } - if p.IPProto == ipproto.TCP || p.IPProto == ipproto.UDP { + if !fragment && (p.IPProto == ipproto.TCP || p.IPProto == ipproto.UDP) { lenForPseudo := len(buf) - csumStart csum := tun.PseudoHeaderChecksum( uint8(p.IPProto), diff --git a/wgengine/netstack/gro/gro_test.go b/wgengine/netstack/gro/gro_test.go index 49171b78c..687518365 100644 --- a/wgengine/netstack/gro/gro_test.go +++ b/wgengine/netstack/gro/gro_test.go @@ -62,6 +62,23 @@ func Test_RXChecksumOffload(t *testing.T) { at := 20 + 16 tcp4InvalidCsum[at] = ^tcp4InvalidCsum[at] + tcp4FirstFragment := make([]byte, 20+20+60) + copy(tcp4FirstFragment, tcp4[:len(tcp4FirstFragment)]) + ipv4H = header.IPv4(tcp4FirstFragment) + ipv4H.SetTotalLength(uint16(len(tcp4FirstFragment))) + ipv4H.SetFlagsFragmentOffset(header.IPv4FlagMoreFragments, 0) + ipv4H.SetChecksum(0) + ipv4H.SetChecksum(^ipv4H.CalculateChecksum()) + + tcp4SecondFragment := make([]byte, 20+40) + copy(tcp4SecondFragment, tcp4[:20]) + copy(tcp4SecondFragment[20:], tcp4[20+80:]) + ipv4H = header.IPv4(tcp4SecondFragment) + ipv4H.SetTotalLength(uint16(len(tcp4SecondFragment))) + ipv4H.SetFlagsFragmentOffset(0, 80) + ipv4H.SetChecksum(0) + ipv4H.SetChecksum(^ipv4H.CalculateChecksum()) + tcp6ExtHeaderInvalidCsum := make([]byte, len(tcp6ExtHeader)) copy(tcp6ExtHeaderInvalidCsum, tcp6ExtHeader) at = 40 + 8 + 16 @@ -87,6 +104,16 @@ func Test_RXChecksumOffload(t *testing.T) { tcp4InvalidCsum, false, }, + { + "tcp4 first fragment skips L4 csum", + tcp4FirstFragment, + true, + }, + { + "tcp4 second fragment skips L4 csum", + tcp4SecondFragment, + true, + }, { "tcp6 with ext header invalid csum", tcp6ExtHeaderInvalidCsum, diff --git a/wgengine/netstack/netstack_test.go b/wgengine/netstack/netstack_test.go index 46f16af31..8732ab782 100644 --- a/wgengine/netstack/netstack_test.go +++ b/wgengine/netstack/netstack_test.go @@ -17,7 +17,9 @@ import ( "gvisor.dev/gvisor/pkg/tcpip" "gvisor.dev/gvisor/pkg/tcpip/adapters/gonet" "gvisor.dev/gvisor/pkg/tcpip/header" + "gvisor.dev/gvisor/pkg/tcpip/network/ipv4" "gvisor.dev/gvisor/pkg/tcpip/stack" + "gvisor.dev/gvisor/pkg/tcpip/transport/udp" "tailscale.com/envknob" "tailscale.com/ipn" "tailscale.com/ipn/ipnlocal" @@ -1680,6 +1682,114 @@ func udp4raw(t testing.TB, src, dst netip.Addr, sport, dport uint16, payload []b return buf } +func fragmentIPv4ForTest(t testing.TB, pkt []byte, firstPayloadLen uint16) (first, second []byte) { + t.Helper() + if firstPayloadLen%8 != 0 { + t.Fatalf("firstPayloadLen %d is not 8-byte aligned", firstPayloadLen) + } + if len(pkt) < header.IPv4MinimumSize+int(firstPayloadLen) { + t.Fatalf("packet length %d too short for firstPayloadLen %d", len(pkt), firstPayloadLen) + } + ip := header.IPv4(pkt) + if ip.HeaderLength() != header.IPv4MinimumSize { + t.Fatalf("test helper only supports 20-byte IPv4 headers; got %d", ip.HeaderLength()) + } + ipPayloadLen := len(pkt) - header.IPv4MinimumSize + if int(firstPayloadLen) >= ipPayloadLen { + t.Fatalf("firstPayloadLen %d must be smaller than IP payload length %d", firstPayloadLen, ipPayloadLen) + } + + first = make([]byte, header.IPv4MinimumSize+int(firstPayloadLen)) + copy(first, pkt[:len(first)]) + firstIP := header.IPv4(first) + firstIP.SetTotalLength(uint16(len(first))) + firstIP.SetFlagsFragmentOffset(header.IPv4FlagMoreFragments, 0) + firstIP.SetChecksum(0) + firstIP.SetChecksum(^firstIP.CalculateChecksum()) + + secondPayloadLen := ipPayloadLen - int(firstPayloadLen) + second = make([]byte, header.IPv4MinimumSize+secondPayloadLen) + copy(second[:header.IPv4MinimumSize], pkt[:header.IPv4MinimumSize]) + copy(second[header.IPv4MinimumSize:], pkt[header.IPv4MinimumSize+int(firstPayloadLen):]) + secondIP := header.IPv4(second) + secondIP.SetTotalLength(uint16(len(second))) + secondIP.SetFlagsFragmentOffset(0, firstPayloadLen) + secondIP.SetChecksum(0) + secondIP.SetChecksum(^secondIP.CalculateChecksum()) + + return first, second +} + +// TestLinkEndpointInjectInboundIPv4Fragments verifies that Tailscale's inbound +// link endpoint path lets IPv4 fragments reach gVisor for reassembly. +// Previously (see https://github.com/tailscale/tailscale/issues/20320), +// gro.RXChecksumOffload validated L4 checksums before reassembly, so the first +// fragment was dropped and the UDP datagram never reached the socket. +func TestLinkEndpointInjectInboundIPv4Fragments(t *testing.T) { + const nicID tcpip.NICID = 1 + localIP := netip.MustParseAddr("100.64.1.2") + remoteIP := netip.MustParseAddr("100.64.1.3") + payload := []byte("0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789abcdefghijklmnopqrstuvwxyz") + raw := udp4raw(t, remoteIP, localIP, 12345, 8081, payload) + first, second := fragmentIPv4ForTest(t, raw, 80) + + s := stack.New(stack.Options{ + NetworkProtocols: []stack.NetworkProtocolFactory{ + ipv4.NewProtocol, + }, + TransportProtocols: []stack.TransportProtocolFactory{ + udp.NewProtocol, + }, + }) + defer s.Close() + + ep := newLinkEndpoint(64, 1280, "", groNotSupported) + if err := s.CreateNIC(nicID, ep); err != nil { + t.Fatalf("CreateNIC: %v", err) + } + if err := s.AddProtocolAddress(nicID, tcpip.ProtocolAddress{ + Protocol: header.IPv4ProtocolNumber, + AddressWithPrefix: tcpip.AddrFrom4(localIP.As4()).WithPrefix(), + }, stack.AddressProperties{}); err != nil { + t.Fatalf("AddProtocolAddress: %v", err) + } + + pc, err := gonet.DialUDP(s, &tcpip.FullAddress{ + NIC: nicID, + Addr: tcpip.AddrFrom4(localIP.As4()), + Port: 8081, + }, nil, header.IPv4ProtocolNumber) + if err != nil { + t.Fatalf("DialUDP: %v", err) + } + defer pc.Close() + + var parsed packet.Parsed + parsed.Decode(first) + ep.injectInbound(&parsed) + parsed.Decode(second) + ep.injectInbound(&parsed) + + if err := pc.SetReadDeadline(time.Now().Add(5 * time.Second)); err != nil { + t.Fatalf("SetReadDeadline: %v", err) + } + buf := make([]byte, 512) + n, addr, err := pc.ReadFrom(buf) + if err != nil { + t.Fatalf("ReadFrom: %v (fragmented packet was not reassembled and delivered)", err) + } + if got := string(buf[:n]); got != string(payload) { + t.Fatalf("payload = %q, want %q", got, payload) + } + udpAddr, ok := addr.(*net.UDPAddr) + if !ok { + t.Fatalf("remote addr = %T(%v), want *net.UDPAddr", addr, addr) + } + if got := udpAddr.AddrPort(); got != netip.MustParseAddrPort("100.64.1.3:12345") { + t.Fatalf("remote addr = %v, want 100.64.1.3:12345", got) + } +} + // TestInjectLoopback verifies that the inject goroutine delivers self-addressed // packets back into gVisor (via DeliverLoopback) instead of sending them to // WireGuard outbound. This is a regression test for a bug where self-dial