diff --git a/tun/offload_linux.go b/tun/offload_linux.go index 54461cc..9a9d38e 100644 --- a/tun/offload_linux.go +++ b/tun/offload_linux.go @@ -748,7 +748,7 @@ const ( udp6GROCandidate ) -func packetIsGROCandidate(b []byte) groCandidateType { +func packetIsGROCandidate(b []byte, canUDPGRO bool) groCandidateType { if len(b) < 28 { return notGROCandidate } @@ -760,14 +760,14 @@ func packetIsGROCandidate(b []byte) groCandidateType { if b[9] == unix.IPPROTO_TCP && len(b) >= 40 { return tcp4GROCandidate } - if b[9] == unix.IPPROTO_UDP { + if b[9] == unix.IPPROTO_UDP && canUDPGRO { return udp4GROCandidate } } else if b[0]>>4 == 6 { if b[6] == unix.IPPROTO_TCP && len(b) >= 60 { return tcp6GROCandidate } - if b[6] == unix.IPPROTO_UDP && len(b) >= 48 { + if b[6] == unix.IPPROTO_UDP && len(b) >= 48 && canUDPGRO { return udp6GROCandidate } } @@ -860,14 +860,15 @@ func udpGRO(bufs [][]byte, offset int, pktI int, table *udpGROTable, isV6 bool) // handleGRO evaluates bufs for GRO, and writes the indices of the resulting // packets into toWrite. toWrite, tcpTable, and udpTable should initially be // empty (but non-nil), and are passed in to save allocs as the caller may reset -// and recycle them across vectors of packets. -func handleGRO(bufs [][]byte, offset int, tcpTable *tcpGROTable, udpTable *udpGROTable, toWrite *[]int) error { +// and recycle them across vectors of packets. canUDPGRO indicates if UDP GRO is +// supported. +func handleGRO(bufs [][]byte, offset int, tcpTable *tcpGROTable, udpTable *udpGROTable, canUDPGRO bool, toWrite *[]int) error { for i := range bufs { if offset < virtioNetHdrLen || offset > len(bufs[i])-1 { return errors.New("invalid offset") } var result groResult - switch packetIsGROCandidate(bufs[i][offset:]) { + switch packetIsGROCandidate(bufs[i][offset:], canUDPGRO) { case tcp4GROCandidate: result = tcpGRO(bufs, offset, i, tcpTable, false) case tcp6GROCandidate: diff --git a/tun/offload_linux_test.go b/tun/offload_linux_test.go index 192232c..91f3941 100644 --- a/tun/offload_linux_test.go +++ b/tun/offload_linux_test.go @@ -286,11 +286,11 @@ func Fuzz_handleGRO(f *testing.F) { pkt9 := udp6Packet(ip6PortA, ip6PortB, 100) pkt10 := udp6Packet(ip6PortA, ip6PortB, 100) pkt11 := udp6Packet(ip6PortA, ip6PortC, 100) - f.Add(pkt0, pkt1, pkt2, pkt3, pkt4, pkt5, pkt6, pkt7, pkt8, pkt9, pkt10, pkt11, offset) - f.Fuzz(func(t *testing.T, pkt0, pkt1, pkt2, pkt3, pkt4, pkt5, pkt6, pkt7, pkt8, pkt9, pkt10, pkt11 []byte, offset int) { + f.Add(pkt0, pkt1, pkt2, pkt3, pkt4, pkt5, pkt6, pkt7, pkt8, pkt9, pkt10, pkt11, true, offset) + f.Fuzz(func(t *testing.T, pkt0, pkt1, pkt2, pkt3, pkt4, pkt5, pkt6, pkt7, pkt8, pkt9, pkt10, pkt11 []byte, canUDPGRO bool, offset int) { pkts := [][]byte{pkt0, pkt1, pkt2, pkt3, pkt4, pkt5, pkt6, pkt7, pkt8, pkt9, pkt10, pkt11} toWrite := make([]int, 0, len(pkts)) - handleGRO(pkts, offset, newTCPGROTable(), newUDPGROTable(), &toWrite) + handleGRO(pkts, offset, newTCPGROTable(), newUDPGROTable(), canUDPGRO, &toWrite) if len(toWrite) > len(pkts) { t.Errorf("len(toWrite): %d > len(pkts): %d", len(toWrite), len(pkts)) } @@ -311,6 +311,7 @@ func Test_handleGRO(t *testing.T) { tests := []struct { name string pktsIn [][]byte + canUDPGRO bool wantToWrite []int wantLens []int wantErr bool @@ -330,10 +331,31 @@ func Test_handleGRO(t *testing.T) { udp6Packet(ip6PortA, ip6PortB, 100), // udp6 flow 1 udp6Packet(ip6PortA, ip6PortB, 100), // udp6 flow 1 }, + true, []int{0, 1, 2, 4, 5, 7, 9}, []int{240, 228, 128, 140, 260, 160, 248}, false, }, + { + "multiple protocols and flows no UDP GRO", + [][]byte{ + tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1), // tcp4 flow 1 + udp4Packet(ip4PortA, ip4PortB, 100), // udp4 flow 1 + udp4Packet(ip4PortA, ip4PortC, 100), // udp4 flow 2 + tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 101), // tcp4 flow 1 + tcp4Packet(ip4PortA, ip4PortC, header.TCPFlagAck, 100, 201), // tcp4 flow 2 + tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 1), // tcp6 flow 1 + tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 101), // tcp6 flow 1 + tcp6Packet(ip6PortA, ip6PortC, header.TCPFlagAck, 100, 201), // tcp6 flow 2 + udp4Packet(ip4PortA, ip4PortB, 100), // udp4 flow 1 + udp6Packet(ip6PortA, ip6PortB, 100), // udp6 flow 1 + udp6Packet(ip6PortA, ip6PortB, 100), // udp6 flow 1 + }, + false, + []int{0, 1, 2, 4, 5, 7, 8, 9, 10}, + []int{240, 128, 128, 140, 260, 160, 128, 148, 148}, + false, + }, { "PSH interleaved", [][]byte{ @@ -346,6 +368,7 @@ func Test_handleGRO(t *testing.T) { tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 201), // v6 flow 1 tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 301), // v6 flow 1 }, + true, []int{0, 2, 4, 6}, []int{240, 240, 260, 260}, false, @@ -360,6 +383,7 @@ func Test_handleGRO(t *testing.T) { udp4Packet(ip4PortA, ip4PortB, 100), udp4Packet(ip4PortA, ip4PortB, 100), }, + true, []int{0, 1, 3, 4}, []int{140, 240, 128, 228}, false, @@ -371,6 +395,7 @@ func Test_handleGRO(t *testing.T) { tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1), // v4 flow 1 seq 1 len 100 tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 201), // v4 flow 1 seq 201 len 100 }, + true, []int{0}, []int{340}, false, @@ -387,6 +412,7 @@ func Test_handleGRO(t *testing.T) { fields.TTL++ }), }, + true, []int{0, 1, 2, 3}, []int{140, 140, 128, 128}, false, @@ -403,6 +429,7 @@ func Test_handleGRO(t *testing.T) { fields.TOS++ }), }, + true, []int{0, 1, 2, 3}, []int{140, 140, 128, 128}, false, @@ -419,6 +446,7 @@ func Test_handleGRO(t *testing.T) { fields.Flags = 1 }), }, + true, []int{0, 1, 2, 3}, []int{140, 140, 128, 128}, false, @@ -435,6 +463,7 @@ func Test_handleGRO(t *testing.T) { fields.Flags = 2 }), }, + true, []int{0, 1, 2, 3}, []int{140, 140, 128, 128}, false, @@ -451,6 +480,7 @@ func Test_handleGRO(t *testing.T) { fields.HopLimit++ }), }, + true, []int{0, 1, 2, 3}, []int{160, 160, 148, 148}, false, @@ -467,6 +497,7 @@ func Test_handleGRO(t *testing.T) { fields.TrafficClass++ }), }, + true, []int{0, 1, 2, 3}, []int{160, 160, 148, 148}, false, @@ -476,7 +507,7 @@ func Test_handleGRO(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { toWrite := make([]int, 0, len(tt.pktsIn)) - err := handleGRO(tt.pktsIn, offset, newTCPGROTable(), newUDPGROTable(), &toWrite) + err := handleGRO(tt.pktsIn, offset, newTCPGROTable(), newUDPGROTable(), tt.canUDPGRO, &toWrite) if err != nil { if tt.wantErr { return @@ -521,74 +552,99 @@ func Test_packetIsGROCandidate(t *testing.T) { udp6TooShort := udp6[:47] tests := []struct { - name string - b []byte - want groCandidateType + name string + b []byte + canUDPGRO bool + want groCandidateType }{ { "tcp4", tcp4, + true, tcp4GROCandidate, }, { "tcp6", tcp6, + true, tcp6GROCandidate, }, { "udp4", udp4, + true, udp4GROCandidate, }, + { + "udp4 no support", + udp4, + false, + notGROCandidate, + }, { "udp6", udp6, + true, udp6GROCandidate, }, + { + "udp6 no support", + udp6, + false, + notGROCandidate, + }, { "udp4 too short", udp4TooShort, + true, notGROCandidate, }, { "udp6 too short", udp6TooShort, + true, notGROCandidate, }, { "tcp4 too short", tcp4TooShort, + true, notGROCandidate, }, { "tcp6 too short", tcp6TooShort, + true, notGROCandidate, }, { "invalid IP version", []byte{0x00}, + true, notGROCandidate, }, { "invalid IP header len", ip4InvalidHeaderLen, + true, notGROCandidate, }, { "ip4 invalid protocol", ip4InvalidProtocol, + true, notGROCandidate, }, { "ip6 invalid protocol", ip6InvalidProtocol, + true, notGROCandidate, }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - if got := packetIsGROCandidate(tt.b); got != tt.want { + if got := packetIsGROCandidate(tt.b, tt.canUDPGRO); got != tt.want { t.Errorf("packetIsGROCandidate() = %v, want %v", got, tt.want) } }) diff --git a/tun/tun_linux.go b/tun/tun_linux.go index 94bffcc..9313ebf 100644 --- a/tun/tun_linux.go +++ b/tun/tun_linux.go @@ -345,7 +345,7 @@ func (tun *NativeTun) Write(bufs [][]byte, offset int) (int, error) { ) tun.toWrite = tun.toWrite[:0] if tun.vnetHdr { - err := handleGRO(bufs, offset, tun.tcpGROTable, tun.udpGROTable, &tun.toWrite) + err := handleGRO(bufs, offset, tun.tcpGROTable, tun.udpGROTable, tun.udpGSO, &tun.toWrite) if err != nil { return 0, err }