diff --git a/pkg/tcpip/network/ipv6/BUILD b/pkg/tcpip/network/ipv6/BUILD index b25e283c4..a48f51d44 100644 --- a/pkg/tcpip/network/ipv6/BUILD +++ b/pkg/tcpip/network/ipv6/BUILD @@ -42,6 +42,7 @@ go_test( "//pkg/bufferv2", "//pkg/refs", "//pkg/refsvfs2", + "//pkg/sync", "//pkg/tcpip", "//pkg/tcpip/checker", "//pkg/tcpip/checksum", diff --git a/pkg/tcpip/network/ipv6/ipv6.go b/pkg/tcpip/network/ipv6/ipv6.go index 943d1e1c1..ec6ccdf86 100644 --- a/pkg/tcpip/network/ipv6/ipv6.go +++ b/pkg/tcpip/network/ipv6/ipv6.go @@ -1046,11 +1046,13 @@ func (e *endpoint) HandlePacket(pkt stack.PacketBufferPtr) { return } - h, ok := e.protocol.parseAndValidate(pkt) + hView, ok := e.protocol.parseAndValidate(pkt) if !ok { stats.MalformedPacketsReceived.Increment() return } + defer hView.Release() + h := header.IPv6(hView.AsSlice()) if !e.nic.IsLoopback() { if !e.protocol.options.AllowExternalLoopbackTraffic { @@ -1101,11 +1103,13 @@ func (e *endpoint) handleLocalPacket(pkt stack.PacketBufferPtr, canSkipRXChecksu defer pkt.DecRef() pkt.RXTransportChecksumValidated = canSkipRXChecksum - h, ok := e.protocol.parseAndValidate(pkt) + hView, ok := e.protocol.parseAndValidate(pkt) if !ok { stats.MalformedPacketsReceived.Increment() return } + defer hView.Release() + h := header.IPv6(hView.AsSlice()) e.handleValidatedPacket(h, pkt, e.nic.Name() /* inNICName */) } @@ -2537,19 +2541,22 @@ func (p *protocol) forwardPendingMulticastPacket(pkt stack.PacketBufferPtr, inst func (*protocol) Wait() {} // parseAndValidate parses the packet (including its transport layer header) and -// returns the parsed IP header. +// returns a view containing the parsed IP header. The caller is responsible +// for releasing the returned View. // // Returns true if the IP header was successfully parsed. -func (p *protocol) parseAndValidate(pkt stack.PacketBufferPtr) (header.IPv6, bool) { +func (p *protocol) parseAndValidate(pkt stack.PacketBufferPtr) (*bufferv2.View, bool) { transProtoNum, hasTransportHdr, ok := p.Parse(pkt) if !ok { return nil, false } - h := header.IPv6(pkt.NetworkHeader().Slice()) + hView := pkt.NetworkHeader().View() + h := header.IPv6(hView.AsSlice()) // Do not include the link header's size when calculating the size of the IP // packet. if !h.IsValid(pkt.Size() - len(pkt.LinkHeader().Slice())) { + hView.Release() return nil, false } @@ -2557,7 +2564,7 @@ func (p *protocol) parseAndValidate(pkt stack.PacketBufferPtr) (header.IPv6, boo p.parseTransport(pkt, transProtoNum) } - return h, true + return hView, true } func (p *protocol) parseTransport(pkt stack.PacketBufferPtr, transProtoNum tcpip.TransportProtocolNumber) { diff --git a/pkg/tcpip/network/ipv6/ipv6_test.go b/pkg/tcpip/network/ipv6/ipv6_test.go index 38de76caa..6b4f315e5 100644 --- a/pkg/tcpip/network/ipv6/ipv6_test.go +++ b/pkg/tcpip/network/ipv6/ipv6_test.go @@ -26,6 +26,7 @@ import ( "github.com/google/go-cmp/cmp" "gvisor.dev/gvisor/pkg/bufferv2" + "gvisor.dev/gvisor/pkg/sync" "gvisor.dev/gvisor/pkg/tcpip" "gvisor.dev/gvisor/pkg/tcpip/checker" "gvisor.dev/gvisor/pkg/tcpip/checksum" @@ -1109,6 +1110,28 @@ type fragmentData struct { data []byte } +func udpGen(payload []byte, multiplier uint8, src, dst tcpip.Address) []byte { + payloadLen := len(payload) + for i := 0; i < payloadLen; i++ { + payload[i] = uint8(i) * multiplier + } + + udpLength := header.UDPMinimumSize + payloadLen + + hdr := prependable.New(udpLength) + u := header.UDP(hdr.Prepend(udpLength)) + u.Encode(&header.UDPFields{ + SrcPort: 5555, + DstPort: 80, + Length: uint16(udpLength), + }) + copy(u.Payload(), payload) + sum := header.PseudoHeaderChecksum(udp.ProtocolNumber, src, dst, uint16(udpLength)) + sum = checksum.Checksum(payload, sum) + u.SetChecksum(^u.CalculateChecksum(sum)) + return hdr.View() +} + func TestReceiveIPv6Fragments(t *testing.T) { const ( udpPayload1Length = 256 @@ -1124,28 +1147,6 @@ func TestReceiveIPv6Fragments(t *testing.T) { routingExtHdrLen = 8 ) - udpGen := func(payload []byte, multiplier uint8, src, dst tcpip.Address) []byte { - payloadLen := len(payload) - for i := 0; i < payloadLen; i++ { - payload[i] = uint8(i) * multiplier - } - - udpLength := header.UDPMinimumSize + payloadLen - - hdr := prependable.New(udpLength) - u := header.UDP(hdr.Prepend(udpLength)) - u.Encode(&header.UDPFields{ - SrcPort: 5555, - DstPort: 80, - Length: uint16(udpLength), - }) - copy(u.Payload(), payload) - sum := header.PseudoHeaderChecksum(udp.ProtocolNumber, src, dst, uint16(udpLength)) - sum = checksum.Checksum(payload, sum) - u.SetChecksum(^u.CalculateChecksum(sum)) - return hdr.View() - } - var udpPayload1Addr1ToAddr2Buf [udpPayload1Length]byte udpPayload1Addr1ToAddr2 := udpPayload1Addr1ToAddr2Buf[:] ipv6Payload1Addr1ToAddr2 := udpGen(udpPayload1Addr1ToAddr2, 1, addr1, addr2) @@ -1936,6 +1937,124 @@ func TestReceiveIPv6Fragments(t *testing.T) { } } +func TestConcurrentFragmentWrites(t *testing.T) { + const udpPayload1Length = 256 + const udpPayload2Length = 128 + var udpPayload1Addr1ToAddr2Buf [udpPayload1Length]byte + udpPayload1Addr1ToAddr2 := udpPayload1Addr1ToAddr2Buf[:] + ipv6Payload1Addr1ToAddr2 := udpGen(udpPayload1Addr1ToAddr2, 1, addr1, addr2) + + var udpPayload2Addr1ToAddr2Buf [udpPayload2Length]byte + udpPayload2Addr1ToAddr2 := udpPayload2Addr1ToAddr2Buf[:] + ipv6Payload2Addr1ToAddr2 := udpGen(udpPayload2Addr1ToAddr2, 2, addr1, addr2) + + fragments := []fragmentData{ + { + srcAddr: addr1, + dstAddr: addr2, + nextHdr: fragmentExtHdrID, + data: append( + // Fragment extension header. + // + // Fragment offset = 0, More = true, ID = 1 + []byte{uint8(header.UDPProtocolNumber), 0, 0, 1, 0, 0, 0, 1}, + ipv6Payload1Addr1ToAddr2[:64]..., + ), + }, + { + srcAddr: addr1, + dstAddr: addr2, + nextHdr: fragmentExtHdrID, + data: append( + // Fragment extension header. + // + // Fragment offset = 0, More = true, ID = 2 + []byte{uint8(header.UDPProtocolNumber), 0, 0, 1, 0, 0, 0, 2}, + ipv6Payload2Addr1ToAddr2[:32]..., + ), + }, + { + srcAddr: addr1, + dstAddr: addr2, + nextHdr: fragmentExtHdrID, + data: append( + // Fragment extension header. + // + // Fragment offset = 8, More = false, ID = 1 + []byte{uint8(header.UDPProtocolNumber), 0, 0, 64, 0, 0, 0, 1}, + ipv6Payload1Addr1ToAddr2[64:]..., + ), + }, + } + + c := newTestContext() + defer c.cleanup() + s := c.s + + e := channel.New(0, header.IPv6MinimumMTU, linkAddr1) + defer e.Close() + if err := s.CreateNIC(nicID, e); err != nil { + t.Fatalf("CreateNIC(%d, _) = %s", nicID, err) + } + protocolAddr := tcpip.ProtocolAddress{ + Protocol: ProtocolNumber, + AddressWithPrefix: addr2.WithPrefix(), + } + if err := s.AddProtocolAddress(nicID, protocolAddr, stack.AddressProperties{}); err != nil { + t.Fatalf("AddProtocolAddress(%d, %+v, {}): %s", nicID, protocolAddr, err) + } + + wq := waiter.Queue{} + we, ch := waiter.NewChannelEntry(waiter.ReadableEvents) + wq.EventRegister(&we) + defer wq.EventUnregister(&we) + defer close(ch) + ep, err := s.NewEndpoint(udp.ProtocolNumber, ProtocolNumber, &wq) + if err != nil { + t.Fatalf("NewEndpoint(%d, %d, _): %s", udp.ProtocolNumber, ProtocolNumber, err) + } + defer ep.Close() + + bindAddr := tcpip.FullAddress{Addr: addr2, Port: 80} + if err := ep.Bind(bindAddr); err != nil { + t.Fatalf("Bind(%+v): %s", bindAddr, err) + } + + var wg sync.WaitGroup + defer wg.Wait() + for i := 0; i < 3; i++ { + wg.Add(1) + go func() { + defer wg.Done() + for i := 0; i < 10; i++ { + for _, f := range fragments { + hdr := prependable.New(header.IPv6MinimumSize) + + // Serialize IPv6 fixed header. + ip := header.IPv6(hdr.Prepend(header.IPv6MinimumSize)) + ip.Encode(&header.IPv6Fields{ + PayloadLength: uint16(len(f.data)), + // We're lying about transport protocol here so that we can generate + // raw extension headers for the tests. + TransportProtocol: tcpip.TransportProtocolNumber(f.nextHdr), + HopLimit: 255, + SrcAddr: f.srcAddr, + DstAddr: f.dstAddr, + }) + + buf := bufferv2.MakeWithData(hdr.View()) + buf.Append(bufferv2.NewViewWithData(f.data)) + pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ + Payload: buf, + }) + e.InjectInbound(ProtocolNumber, pkt) + pkt.DecRef() + } + } + }() + } +} + func TestInvalidIPv6Fragments(t *testing.T) { const ( addr1 = tcpip.Address("\x0a\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x01")