From bd89a24410142ecc7603c8975c5f5f5075509bee Mon Sep 17 00:00:00 2001 From: Jayden Nyamiaka Date: Tue, 3 Sep 2024 11:24:58 -0700 Subject: [PATCH] Implement PayloadSet operation (parsing, interpretation, evaluation, tests). Tests payload set evaluation for all basic fields of IP, IPv6, & TCP headers. Similar to PayloadLoad, these headers were prioritized. Setting other headers should still work but wasn't explicitly tested. Tests for other packets should be added later. PiperOrigin-RevId: 670636213 --- pkg/tcpip/nftables/BUILD | 1 + pkg/tcpip/nftables/nftables.go | 211 ++++++- pkg/tcpip/nftables/nftables_test.go | 850 ++++++++++++++++++++++----- pkg/tcpip/nftables/nftinterp.go | 152 ++++- pkg/tcpip/nftables/nftinterp_test.go | 171 +++++- 5 files changed, 1223 insertions(+), 162 deletions(-) diff --git a/pkg/tcpip/nftables/BUILD b/pkg/tcpip/nftables/BUILD index 6ec1ad4d1..3bc198613 100644 --- a/pkg/tcpip/nftables/BUILD +++ b/pkg/tcpip/nftables/BUILD @@ -13,6 +13,7 @@ go_library( ], deps = [ "//pkg/abi/linux", + "//pkg/tcpip/checksum", "//pkg/tcpip/header", "//pkg/tcpip/stack", ], diff --git a/pkg/tcpip/nftables/nftables.go b/pkg/tcpip/nftables/nftables.go index d2d7a1aab..ac3af4048 100644 --- a/pkg/tcpip/nftables/nftables.go +++ b/pkg/tcpip/nftables/nftables.go @@ -41,10 +41,12 @@ package nftables import ( + "encoding/binary" "fmt" "slices" "gvisor.dev/gvisor/pkg/abi/linux" + "gvisor.dev/gvisor/pkg/tcpip/checksum" "gvisor.dev/gvisor/pkg/tcpip/header" "gvisor.dev/gvisor/pkg/tcpip/stack" ) @@ -567,6 +569,7 @@ var ( _ operation = (*immediate)(nil) _ operation = (*comparison)(nil) _ operation = (*payloadLoad)(nil) + _ operation = (*payloadSet)(nil) ) // immediate is an operation that sets the data in a register. @@ -740,6 +743,32 @@ func validatePayloadBase(base payloadBase) error { } } +// getPayloadBuffer gets the data from the packet payload starting from the +// the beginning of the specified base header. +// Returns nil if the payload is not present or invalid. +func getPayloadBuffer(pkt *stack.PacketBuffer, base payloadBase) []byte { + switch base { + case linux.NFT_PAYLOAD_LL_HEADER: + // Note: Assumes Mac Header is present and valid for necessary use cases. + // Also, doesn't check VLAN tag because VLAN isn't supported by gVisor. + return pkt.LinkHeader().Slice() + case linux.NFT_PAYLOAD_NETWORK_HEADER: + // No checks done in linux kernel. + return pkt.NetworkHeader().Slice() + case linux.NFT_PAYLOAD_TRANSPORT_HEADER: + // Note: Assumes L4 protocol is present and valid for necessary use cases. + + // Errors if the packet is fragmented for IPv4 only. + if net := pkt.NetworkHeader().Slice(); len(net) > 0 && pkt.NetworkProtocolNumber == header.IPv4ProtocolNumber { + if h := header.IPv4(net); h.More() || h.FragmentOffset() != 0 { + break // packet is fragmented + } + } + return pkt.TransportHeader().Slice() + } + return nil +} + // newPayloadLoad creates a new PayloadLoad operation. func newPayloadLoad(base payloadBase, offset, blen, dreg uint8) (*payloadLoad, error) { if dreg == linux.NFT_REG_VERDICT { @@ -757,27 +786,8 @@ func newPayloadLoad(base payloadBase, offset, blen, dreg uint8) (*payloadLoad, e // evaluate for PayloadLoad loads data from the packet payload into the // destination register. func (op payloadLoad) evaluate(regs *registerSet, pkt *stack.PacketBuffer) { - // Gets the data from the packet payload. - var payload []byte = nil - switch op.base { - case linux.NFT_PAYLOAD_LL_HEADER: - // Note: Assumes Mac Header is present and valid for necessary use cases. - // Also, doesn't check VLAN tag because VLAN isn't supported by gVisor. - payload = pkt.LinkHeader().Slice() - case linux.NFT_PAYLOAD_NETWORK_HEADER: - // No checks done in linux kernel. - payload = pkt.NetworkHeader().Slice() - case linux.NFT_PAYLOAD_TRANSPORT_HEADER: - // Note: Assumes L4 protocol is present and valid for necessary use cases. - - // Errors if the packet is fragmented for IPv4 only. - if net := pkt.NetworkHeader().Slice(); len(net) > 0 && pkt.NetworkProtocolNumber == header.IPv4ProtocolNumber { - if h := header.IPv4(net); h.More() || h.FragmentOffset() != 0 { - break // packet is fragmented - } - } - payload = pkt.TransportHeader().Slice() - } + // Gets the packet payload. + payload := getPayloadBuffer(pkt, op.base) // Breaks if could not retrieve packet data. if payload == nil || len(payload) < int(op.offset+op.blen) { @@ -790,6 +800,163 @@ func (op payloadLoad) evaluate(regs *registerSet, pkt *stack.PacketBuffer) { data.storeData(regs, op.dreg) } +// payloadSet is an operation that sets data in the packet payload to the value +// in a register. +// Note: payload operations are not supported for the verdict register. +type payloadSet struct { + base payloadBase // Payload base to access data from. + offset uint8 // Number of bytes to skip after the base for data. + blen uint8 // Number of bytes to load. + sreg uint8 // Number of the source register. + csumType uint8 // Type of checksum to use. + csumOffset uint8 // Number of bytes to skip after the base for checksum. + csumFlags uint8 // Flags for checksum. + + // Note: the only flag defined for csumFlags is NFT_PAYLOAD_L4CSUM_PSEUDOHDR. + // This flag is used to update L4 checksums whenever there has been a change + // to a field that is part of the pseudo-header for the L4 checksum, not when + // data within the L4 header is changed (instead setting csumType to + // NFT_PAYLOAD_CSUM_INET suffices for that case). + + // For example, if any part of the L4 header is changed, csumType is set to + // NFT_PAYLOAD_CSUM_INET and no flag is set for csumFlags since we only need + // to update the checksum of the header specified by the payload base. + // On the other hand, if data in the L3 header is changed that is part of + // the pseudo-header for the L4 checksum (like saddr/daddr), csumType is set + // to NFT_PAYLOAD_CSUM_INET and csumFlags to NFT_PAYLOAD_L4CSUM_PSEUDOHDR + // because in addition to updating the checksum for the header specified by + // the payload base, we need to separately locate and update the L4 checksum. +} + +// validateChecksumType ensures the checksum type is valid. +func validateChecksumType(csumType uint8) error { + switch csumType { + case linux.NFT_PAYLOAD_CSUM_NONE: + return nil + case linux.NFT_PAYLOAD_CSUM_INET: + return nil + case linux.NFT_PAYLOAD_CSUM_SCTP: + return fmt.Errorf("SCTP checksum not supported") + default: + return fmt.Errorf("invalid checksum type: %d", csumType) + } +} + +// newPayloadSet creates a new PayloadSet operation. +func newPayloadSet(base payloadBase, offset, blen, sreg, csumType, csumOffset, csumFlags uint8) (*payloadSet, error) { + if sreg == linux.NFT_REG_VERDICT { + return nil, fmt.Errorf("payload set operation cannot use verdict register as destination") + } + if blen > 16 || (blen > 4 && is4ByteRegister(sreg)) { + return nil, fmt.Errorf("payload length %d is too long for destination register %d", blen, sreg) + } + if err := validatePayloadBase(base); err != nil { + return nil, err + } + if err := validateChecksumType(csumType); err != nil { + return nil, err + } + if csumFlags&^linux.NFT_PAYLOAD_L4CSUM_PSEUDOHDR != 0 { + return nil, fmt.Errorf("invalid checksum flags: %d", csumFlags) + } + return &payloadSet{base: base, offset: offset, blen: blen, sreg: sreg, + csumType: csumType, csumOffset: csumOffset, csumFlags: csumFlags}, nil +} + +// evaluate for PayloadSet sets data in the packet payload to the value in the +// source register. +func (op payloadSet) evaluate(regs *registerSet, pkt *stack.PacketBuffer) { + // Gets the packet payload. + payload := getPayloadBuffer(pkt, op.base) + + // Breaks if could not retrieve packet data. + if payload == nil || len(payload) < int(op.offset+op.blen) { + regs.verdict = Verdict{Code: VC(linux.NFT_BREAK)} + return + } + + // Gets the register data assumed to be in Big Endian. + regData := getRegisterBuffer(regs, op.sreg)[:op.blen] + + // Returns early if the source data is the same as the existing payload data. + if slices.Equal(regData, payload[op.offset:op.offset+op.blen]) { + return + } + + // Sets payload data to source register data after checksum updates. + defer copy(payload[op.offset:op.offset+op.blen], regData) + + // Specifies no checksum updates. + if op.csumType != linux.NFT_PAYLOAD_CSUM_INET && op.csumFlags == 0 { + return + } + + // Calculates partial checksums of old and new data. + // Note: Checksums are done on 2-byte boundaries, so we must append the + // surrounding bytes in our checksum calculations if the beginning or end + // of the checksum is not aligned to a 2-byte boundary. + begin := op.offset + end := op.offset + op.blen + if begin%2 != 0 { + begin-- + } + if end%2 != 0 && end != uint8(len(payload)) { + end++ + } + tempOld := make([]byte, end-begin) + copy(tempOld, payload[begin:end]) + tempNew := make([]byte, end-begin) + if begin != op.offset { + tempNew[0] = payload[begin] + } + copy(tempNew[op.offset-begin:], regData) + if end != op.offset+op.blen { + tempNew[len(tempNew)-1] = payload[end-1] + } + oldDataCsum := checksum.Checksum(tempOld, 0) + newDataCsum := checksum.Checksum(tempNew, 0) + + // Updates the checksum of the header specified by the payload base. + if op.csumType == linux.NFT_PAYLOAD_CSUM_INET { + // Reads the old checksum from the packet payload. + oldTotalCsum := binary.BigEndian.Uint16(payload[op.csumOffset:]) + + // New Total = Old Total - Old Data + New Data + // Logic is very similar to checksum.checksumUpdate2ByteAlignedUint16 + // in gvisor/pkg/tcpip/header/checksum.go + newTotalCsum := checksum.Combine(^oldTotalCsum, checksum.Combine(newDataCsum, ^oldDataCsum)) + checksum.Put(payload[op.csumOffset:], ^newTotalCsum) + } + + // Separately updates the L4 checksum if the pseudo-header flag is set. + // Note: it is possible to update the L4 checksum without updating the + // checksum of the header specified by the payload base (ie type is NONE, + // flag is pseudo-header). Specifically, IPv6 headers don't have their + // own checksum calculations, but the L4 checksum is still updated for any + // TCP/UDP headers following the IPv6 header. + if op.csumFlags&linux.NFT_PAYLOAD_L4CSUM_PSEUDOHDR != 0 { + if tBytes := pkt.TransportHeader().Slice(); pkt.TransportProtocolNumber != 0 && len(tBytes) > 0 { + var transport header.Transport + switch pkt.TransportProtocolNumber { + case header.TCPProtocolNumber: + transport = header.TCP(tBytes) + case header.UDPProtocolNumber: + transport = header.UDP(tBytes) + case header.ICMPv4ProtocolNumber: + transport = header.ICMPv4(tBytes) + case header.ICMPv6ProtocolNumber: + transport = header.ICMPv6(tBytes) + case header.IGMPProtocolNumber: + transport = header.IGMP(tBytes) + } + if transport != nil { // only updates if the transport header is present. + // New Total = Old Total - Old Data + New Data (same as above) + transport.SetChecksum(^checksum.Combine(^transport.Checksum(), checksum.Combine(newDataCsum, ^oldDataCsum))) + } + } + } +} + // // Register and Register-Related Implementations. // Note: Registers are represented by type uint8 for the register number. @@ -942,7 +1109,7 @@ func (rd bytesData) storeData(regs *registerSet, reg uint8) { } // registerSet represents the set of registers supported by the kernel. -// Use RegisterData.StoreData to set data in the registers. +// Use RegisterData.storeData to set data in the registers. // Note: Corresponds to nft_regs from include/net/netfilter/nf_tables.h. type registerSet struct { verdict Verdict // 16-byte verdict register diff --git a/pkg/tcpip/nftables/nftables_test.go b/pkg/tcpip/nftables/nftables_test.go index d1ebeef59..df4731b61 100644 --- a/pkg/tcpip/nftables/nftables_test.go +++ b/pkg/tcpip/nftables/nftables_test.go @@ -18,6 +18,7 @@ import ( "encoding/binary" "fmt" "reflect" + "slices" "testing" "gvisor.dev/gvisor/pkg/abi/linux" @@ -53,11 +54,20 @@ var ( // Packet Constants. const ( - transportProtocol = tcpip.TransportProtocolNumber(6) + tcpTransportProtocol = header.TCPProtocolNumber arbitraryHeaderID = 3 arbitraryTimeToLive = 64 + arbitraryPort = 12345 + arbitraryPort2 = 80 + tcpSeqNum = 32 + tcpAckNum = 165 + tcpWinSize = 65535 + tcpUrgentPointer = 0 + + arbitraryNonZeroFragmentOffset = 16 + // TODO(b/345684870): Use constants defined in the pkg/tcpip/header package. // Ethernet Offsets and Lengths. ethDstAddrOffset = 0 @@ -129,32 +139,75 @@ var ( arbitraryIPv6AddrB2 = [16]byte{0x20, 0x01, 0x0d, 0xb8, 0x85, 0xa3, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0xbb} ipv6MinPayloadLength = 0 - arbitraryPort = 12345 - arbitraryPort2 = 80 - tcpSeqNum = 32 - tcpAckNum = 165 - tcpWinSize = 65535 - tcpUrgentPointer = 0 + // Note: these are functions to make sure they are not modified by tests and + // so each tests gets a new value. + arbitraryEthernetFields = func() *header.EthernetFields { + return &header.EthernetFields{ + SrcAddr: arbitraryLinkAddr, + DstAddr: arbitraryLinkAddr2, + Type: arbitraryEthernetType, + } + } - arbitraryNonZeroFragmentOffset = 16 + arbitraryIPv4Fields = func() *header.IPv4Fields { + return &header.IPv4Fields{ + TOS: 0, + TotalLength: uint16(ipv4MinTotalLength), + ID: arbitraryHeaderID, + FragmentOffset: 0, + TTL: arbitraryTimeToLive, + Protocol: uint8(tcpTransportProtocol), + Checksum: 0, + SrcAddr: tcpip.AddrFrom4(arbitraryIPv4AddrB), + DstAddr: tcpip.AddrFrom4(arbitraryIPv4AddrB2), + Options: nil, + } + } + + fragmentedIPv4Fields = func() *header.IPv4Fields { + fields := arbitraryIPv4Fields() + fields.FragmentOffset = arbitraryNonZeroFragmentOffset + return fields + } + + arbitraryIPv6Fields = func() *header.IPv6Fields { + return &header.IPv6Fields{ + TrafficClass: 0, + FlowLabel: 0, + PayloadLength: uint16(ipv6MinPayloadLength), + TransportProtocol: tcpTransportProtocol, + HopLimit: arbitraryTimeToLive, + SrcAddr: tcpip.AddrFrom16(arbitraryIPv6AddrB), + DstAddr: tcpip.AddrFrom16(arbitraryIPv6AddrB2), + } + } + + arbitraryTCPFields = func() *header.TCPFields { + return &header.TCPFields{ + SrcPort: uint16(arbitraryPort), + DstPort: uint16(arbitraryPort2), + SeqNum: uint32(tcpSeqNum), + AckNum: uint32(tcpAckNum), + DataOffset: header.TCPMinimumSize, + WindowSize: uint16(tcpWinSize), + Checksum: 0, + UrgentPointer: uint16(tcpUrgentPointer), + } + } ) -// makeArbitraryGeneralPacket creates an arbitrary packet for testing. -func makeArbitraryGeneralPacket(reserved int) *stack.PacketBuffer { +// makeArbitraryPacket creates an arbitrary packet for testing. +func makeArbitraryPacket(reserved int) *stack.PacketBuffer { return stack.NewPacketBuffer(stack.PacketBufferOptions{ ReserveHeaderBytes: reserved, Payload: buffer.MakeWithData([]byte{0, 2, 4, 8, 16, 32, 64, 128}), }) } -// makeArbitraryEtherPacket creates a packet with an arbitrary ethernet header. -func makeArbitraryEtherPacket(reserved int) *stack.PacketBuffer { +// makeEthernetPacket creates a packet with an Ethernet header. +func makeEthernetPacket(reserved int, ethFields *header.EthernetFields) *stack.PacketBuffer { eth := make([]byte, header.EthernetMinimumSize) - header.Ethernet(eth).Encode(&header.EthernetFields{ - SrcAddr: arbitraryLinkAddr, - DstAddr: arbitraryLinkAddr2, - Type: arbitraryEthernetType, - }) + header.Ethernet(eth).Encode(ethFields) pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ ReserveHeaderBytes: reserved, Payload: buffer.MakeWithData(eth), @@ -163,8 +216,8 @@ func makeArbitraryEtherPacket(reserved int) *stack.PacketBuffer { return pkt } -// makeArbitraryIPv4Packet creates a packet with an arbitrary IPv4 header. -func makeArbitraryIPv4Packet(reserved int) *stack.PacketBuffer { +// makeIPv4Packet creates a packet with an IPv4 header. +func makeIPv4Packet(reserved int, ipv4Fields *header.IPv4Fields) *stack.PacketBuffer { // Creates a new PacketBuffer with enough space for the IPv4 header. pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ ReserveHeaderBytes: reserved, @@ -174,18 +227,7 @@ func makeArbitraryIPv4Packet(reserved int) *stack.PacketBuffer { ipv4Hdr := header.IPv4(pkt.NetworkHeader().Push(header.IPv4MinimumSize)) // Initializes the IPv4 header with fields. - ipv4Hdr.Encode(&header.IPv4Fields{ - TOS: 0, - TotalLength: uint16(ipv4MinTotalLength), - ID: arbitraryHeaderID, - FragmentOffset: 0, - TTL: arbitraryTimeToLive, - Protocol: uint8(transportProtocol), - Checksum: 0, - SrcAddr: tcpip.AddrFrom4(arbitraryIPv4AddrB), - DstAddr: tcpip.AddrFrom4(arbitraryIPv4AddrB2), - Options: nil, - }) + ipv4Hdr.Encode(ipv4Fields) // Calculates and sets the checksum. ipv4Hdr.SetChecksum(^ipv4Hdr.CalculateChecksum()) @@ -196,32 +238,8 @@ func makeArbitraryIPv4Packet(reserved int) *stack.PacketBuffer { return pkt } -// makeFragmentedIPv4Packet creates a packet with an arbitrary IPv4 header that -// is fragmented (FragmentOffset != 0). -func makeFragmentedIPv4Packet(reserved int) *stack.PacketBuffer { - pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ - ReserveHeaderBytes: reserved, - }) - ipv4Hdr := header.IPv4(pkt.NetworkHeader().Push(header.IPv4MinimumSize)) - ipv4Hdr.Encode(&header.IPv4Fields{ - TOS: 0, - TotalLength: uint16(ipv4MinTotalLength), - ID: arbitraryHeaderID, - FragmentOffset: uint16(arbitraryNonZeroFragmentOffset), - TTL: arbitraryTimeToLive, - Protocol: uint8(transportProtocol), - Checksum: 0, - SrcAddr: tcpip.AddrFrom4(arbitraryIPv4AddrB), - DstAddr: tcpip.AddrFrom4(arbitraryIPv4AddrB2), - Options: nil, - }) - ipv4Hdr.SetChecksum(^ipv4Hdr.CalculateChecksum()) - pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber - return pkt -} - -// makeArbitraryIPv6Packet creates a packet with an arbitrary IPv6 header. -func makeArbitraryIPv6Packet(reserved int) *stack.PacketBuffer { +// makeIPv6Packet creates a packet with an IPv6 header. +func makeIPv6Packet(reserved int, ipv6Fields *header.IPv6Fields) *stack.PacketBuffer { // Creates a new PacketBuffer with enough space for the IPv4 header. pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ ReserveHeaderBytes: reserved, @@ -231,15 +249,9 @@ func makeArbitraryIPv6Packet(reserved int) *stack.PacketBuffer { ipv6Hdr := header.IPv6(pkt.NetworkHeader().Push(header.IPv6MinimumSize)) // Initializes the IPv6 header with fields. - ipv6Hdr.Encode(&header.IPv6Fields{ - TrafficClass: 0, - FlowLabel: 0, - PayloadLength: uint16(ipv6MinPayloadLength), - TransportProtocol: transportProtocol, - HopLimit: arbitraryTimeToLive, - SrcAddr: tcpip.AddrFrom16(arbitraryIPv6AddrB), - DstAddr: tcpip.AddrFrom16(arbitraryIPv6AddrB2), - }) + ipv6Hdr.Encode(ipv6Fields) + + // No checksum for IPv6 (relies on L4 checksum if extra security is needed). // Sets the network protocol number. pkt.NetworkProtocolNumber = header.IPv6ProtocolNumber @@ -247,37 +259,59 @@ func makeArbitraryIPv6Packet(reserved int) *stack.PacketBuffer { return pkt } -// makeArbitraryIPv4TCPPacket creates a packet with an arbitrary IPv4 and TCP -// header. -func makeArbitraryIPv4TCPPacket(reserved int) *stack.PacketBuffer { - pkt := makeArbitraryIPv4Packet(reserved) - +// addTCPHeader adds a TCP header to a packet and returns the header. +// Note: this does not compute the checksum. +func addTCPHeader(pkt *stack.PacketBuffer, tcpFields *header.TCPFields) header.TCP { // Prepends the TCP header to the packet buffer. - tcpHdr := header.TCP(pkt.TransportHeader().Push(header.TCPMinimumSize)) + tcpHdr := header.TCP(pkt.TransportHeader().Push(int(tcpFields.DataOffset))) // Initializes the TCP header with fields. - tcpHdr.Encode(&header.TCPFields{ - SrcPort: uint16(arbitraryPort), - DstPort: uint16(arbitraryPort2), - SeqNum: uint32(tcpSeqNum), - AckNum: uint32(tcpAckNum), - DataOffset: header.TCPMinimumSize, - WindowSize: uint16(tcpWinSize), - Checksum: 0, - UrgentPointer: uint16(tcpUrgentPointer), - }) - - // Calculates the TCP checksum using the pseudo-header and set it in the TCP header. - tcpHdr.SetChecksum(tcpHdr.CalculateChecksum(header.PseudoHeaderChecksum( - header.TCPProtocolNumber, - tcpip.AddrFrom4(arbitraryIPv4AddrB), - tcpip.AddrFrom4(arbitraryIPv4AddrB2), - header.TCPMinimumSize, - ))) + tcpHdr.Encode(tcpFields) // Sets the transport protocol number. pkt.TransportProtocolNumber = header.TCPProtocolNumber + return tcpHdr +} + +// makeIPv4TCPPacket creates a packet with an IPv4 and TCP header. +func makeIPv4TCPPacket(reserved int, ipv4Fields *header.IPv4Fields, tcpFields *header.TCPFields) *stack.PacketBuffer { + // Makes a packet with the L3 IPv4 header (this sets the checksum). + pkt := makeIPv4Packet(reserved, ipv4Fields) + + // Adds the L4 TCP header. + tcpHdr := addTCPHeader(pkt, tcpFields) + + // Calculates the TCP checksum using the pseudo-header and sets it in the TCP header. + tcpHdr.SetChecksum(^tcpHdr.CalculateChecksum(header.PseudoHeaderChecksum( + tcpip.TransportProtocolNumber(ipv4Fields.Protocol), + ipv4Fields.SrcAddr, + ipv4Fields.DstAddr, + uint16(tcpFields.DataOffset), + ))) + + return pkt +} + +// makeIPv6TCPPacket creates a packet with an IPv6 and TCP header. +func makeIPv6TCPPacket(reserved int, ipv6Fields *header.IPv6Fields, tcpFields *header.TCPFields) *stack.PacketBuffer { + // Makes a packet with the L3 IPv6 header (this sets the checksum). + pkt := makeIPv6Packet(reserved, ipv6Fields) + + // Adds the L4 TCP header. + tcpHdr := addTCPHeader(pkt, tcpFields) + + // Calculates the TCP checksum using the pseudo-header and sets it in the TCP header. + tcpHdr.SetChecksum(^tcpHdr.CalculateChecksum(header.PseudoHeaderChecksum( + // Next header is supposed to be in pseudo-header calculation for IPv6 + // transport protocol checksum, not the transport protocol number according + // to RFC 2460 (https://www.rfc-editor.org/rfc/rfc2460.html#section-8.1). + header.IPv6(pkt.NetworkHeader().Slice()).TransportProtocol(), + ipv6Fields.SrcAddr, + ipv6Fields.DstAddr, + uint16(ipv6Fields.PayloadLength), + ))) + return pkt } @@ -285,11 +319,11 @@ func makeArbitraryIPv4TCPPacket(reserved int) *stack.PacketBuffer { // error when evaluating a packet for an unsupported address family. func TestUnsupportedAddressFamily(t *testing.T) { // Makes arbitrary packet for comparison (to check for no changes). - cmpPkt := makeArbitraryGeneralPacket(arbitraryReservedHeaderBytes) + cmpPkt := makeArbitraryPacket(arbitraryReservedHeaderBytes) nf := NewNFTables() for _, unsupportedFamily := range []AddressFamily{AddressFamily(NumAFs), AddressFamily(-1)} { // Note: the Prerouting hook is arbitrary (any hook would work). - pkt := makeArbitraryGeneralPacket(arbitraryReservedHeaderBytes) + pkt := makeArbitraryPacket(arbitraryReservedHeaderBytes) v, err := nf.EvaluateHook(unsupportedFamily, arbitraryHook, pkt) if err == nil { t.Fatalf("expecting error for EvaluateHook with unsupported address family %d; got %v verdict, %s packet, and error %v", @@ -304,12 +338,12 @@ func TestUnsupportedAddressFamily(t *testing.T) { // when evaluating packets at the hook-level. func TestAcceptAllForSupportedHooks(t *testing.T) { // Makes arbitrary packet for comparison (to check for no changes). - cmpPkt := makeArbitraryGeneralPacket(arbitraryReservedHeaderBytes) + cmpPkt := makeArbitraryPacket(arbitraryReservedHeaderBytes) for _, family := range []AddressFamily{IP, IP6, Inet, Arp, Bridge, Netdev} { t.Run(family.String()+" address family", func(t *testing.T) { nf := NewNFTables() for _, hook := range []Hook{Prerouting, Input, Forward, Output, Postrouting, Ingress, Egress} { - pkt := makeArbitraryGeneralPacket(arbitraryReservedHeaderBytes) + pkt := makeArbitraryPacket(arbitraryReservedHeaderBytes) v, err := nf.EvaluateHook(family, hook, pkt) supported := false @@ -515,7 +549,7 @@ func TestEvaluateImmediateVerdict(t *testing.T) { } // Runs evaluation and checks verdict. - pkt := makeArbitraryGeneralPacket(arbitraryReservedHeaderBytes) + pkt := makeArbitraryPacket(arbitraryReservedHeaderBytes) v, err := nf.EvaluateHook(arbitraryFamily, arbitraryHook, pkt) if err != nil { @@ -571,7 +605,7 @@ func TestEvaluateImmediateBytesData(t *testing.T) { } } // Runs evaluation and checks for default policy verdict accept - pkt := makeArbitraryGeneralPacket(arbitraryReservedHeaderBytes) + pkt := makeArbitraryPacket(arbitraryReservedHeaderBytes) v, err := nf.EvaluateHook(arbitraryFamily, arbitraryHook, pkt) if err != nil { t.Fatalf("unexpected error for EvaluateHook: %v", err) @@ -1056,7 +1090,7 @@ func TestEvaluateComparison(t *testing.T) { } // Runs evaluation and checks verdict. - pkt := makeArbitraryGeneralPacket(arbitraryReservedHeaderBytes) + pkt := makeArbitraryPacket(arbitraryReservedHeaderBytes) v, err := nf.EvaluateHook(arbitraryFamily, arbitraryHook, pkt) if err != nil { t.Fatalf("unexpected error for EvaluateHook: %v", err) @@ -1087,10 +1121,10 @@ func TestEvaluateComparison(t *testing.T) { // TODO(b/339691111): Add tests for VLAN, ARP, ICMP, ICMPv6, IGMP, UDP headers. func TestEvaluatePayloadLoad(t *testing.T) { // Sets testing packets. - ethernetPacket := makeArbitraryEtherPacket(0) - ipv4Packet := makeArbitraryIPv4Packet(header.IPv4MinimumSize) - ipv6Packet := makeArbitraryIPv6Packet(header.IPv6MinimumSize) - tcpPacket := makeArbitraryIPv4TCPPacket(header.IPv4MinimumSize + header.TCPMinimumSize) + ethernetPacket := makeEthernetPacket(0, arbitraryEthernetFields()) + ipv4Packet := makeIPv4Packet(header.IPv4MinimumSize, arbitraryIPv4Fields()) + ipv6Packet := makeIPv6Packet(header.IPv6MinimumSize, arbitraryIPv6Fields()) + tcpPacket := makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, arbitraryIPv4Fields(), arbitraryTCPFields()) for _, test := range []struct { tname string @@ -1140,13 +1174,13 @@ func TestEvaluatePayloadLoad(t *testing.T) { }, // Though the packet is fragmented, there should be no issue because we are // changing data within the network header. - { // cmd: add rule ip tab ch ip frag-off 1 + { // cmd: add rule ip tab ch ip frag-off 2 (16 bytes) tname: "load ipv4 header fragment offset non zero for fragmented packet", - pkt: makeFragmentedIPv4Packet(header.IPv4MinimumSize), - op1: mustCreatePayloadLoad(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv4FragOffOffset, ipv4FragOffLen, linux.NFT_REG_1), + pkt: makeIPv4Packet(header.IPv4MinimumSize, fragmentedIPv4Fields()), + op1: mustCreatePayloadLoad(t, linux.NFT_PAYLOAD_NETWORK_HEADER, 6, 2, linux.NFT_REG_1), op2: mustCreateComparison(t, linux.NFT_REG_1, linux.NFT_CMP_EQ, newBytesData(numToBE(arbitraryNonZeroFragmentOffset/8, ipv4FragOffLen))), // we divide by 8 because the fragment offset is in units of 8 bytes, - // which is encoded into the packet in IPv4.Encode() + // which is encoded into the packet in IPv4.Encode(). }, { // cmd: add rule ip tab ch ip ttl 64 tname: "load ipv4 header time to live", @@ -1154,11 +1188,11 @@ func TestEvaluatePayloadLoad(t *testing.T) { op1: mustCreatePayloadLoad(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv4TTLOffset, ipv4TTLLen, linux.NFT_REG32_01), op2: mustCreateComparison(t, linux.NFT_REG32_01, linux.NFT_CMP_EQ, newBytesData(numToBE(arbitraryTimeToLive, ipv4TTLLen))), }, - { // cmd: add rule ip tab ch tcp + { // cmd: add rule ip tab ch ip protocol tcp tname: "load ipv4 header protocol", pkt: ipv4Packet, op1: mustCreatePayloadLoad(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv4ProtocolOffset, ipv4ProtocolLen, linux.NFT_REG_1), - op2: mustCreateComparison(t, linux.NFT_REG_1, linux.NFT_CMP_EQ, newBytesData(numToBE(int(transportProtocol), ipv4ProtocolLen))), + op2: mustCreateComparison(t, linux.NFT_REG_1, linux.NFT_CMP_EQ, newBytesData(numToBE(int(tcpTransportProtocol), ipv4ProtocolLen))), }, { // cmd: add rule ip tab ch ip saddr 192.168.1.1 tname: "load ipv4 header source address", @@ -1190,7 +1224,7 @@ func TestEvaluatePayloadLoad(t *testing.T) { tname: "load ipv6 header next header", pkt: ipv6Packet, op1: mustCreatePayloadLoad(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv6NextHdrOffset, ipv6NextHdrLen, linux.NFT_REG_1), - op2: mustCreateComparison(t, linux.NFT_REG_1, linux.NFT_CMP_EQ, newBytesData(numToBE(int(transportProtocol), ipv6NextHdrLen))), + op2: mustCreateComparison(t, linux.NFT_REG_1, linux.NFT_CMP_EQ, newBytesData(numToBE(int(tcpTransportProtocol), ipv6NextHdrLen))), }, { // cmd: add rule ip6 tab ch ip6 hoplimit 64 tname: "load ipv6 header hop limit", @@ -1204,7 +1238,7 @@ func TestEvaluatePayloadLoad(t *testing.T) { op1: mustCreatePayloadLoad(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv6SrcAddrOffset, ipv6SrcAddrLen, linux.NFT_REG_1), op2: mustCreateComparison(t, linux.NFT_REG_1, linux.NFT_CMP_EQ, newBytesData(arbitraryIPv6AddrB[:])), }, - { // cmd: add rule ip6 tab ch ip6 saddr 2001:db8:85a3::bb + { // cmd: add rule ip6 tab ch ip6 daddr 2001:db8:85a3::bb tname: "load ipv6 header destination address", pkt: ipv6Packet, op1: mustCreatePayloadLoad(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv6DstAddrOffset, ipv6DstAddrLen, linux.NFT_REG_1), @@ -1216,30 +1250,9 @@ func TestEvaluatePayloadLoad(t *testing.T) { // IPv4 packet, this can be problematic, so the evaluation should break. { tname: "load for transport header with a fragmented ipv4 packet", - pkt: func() *stack.PacketBuffer { - p := makeFragmentedIPv4Packet(header.IPv4MinimumSize + header.TCPMinimumSize) - tcpHdr := header.TCP(p.TransportHeader().Push(header.TCPMinimumSize)) - tcpHdr.Encode(&header.TCPFields{ - SrcPort: uint16(arbitraryPort), - DstPort: uint16(arbitraryPort2), - SeqNum: uint32(tcpSeqNum), - AckNum: uint32(tcpAckNum), - DataOffset: header.TCPMinimumSize, - WindowSize: uint16(tcpWinSize), - Checksum: 0, - UrgentPointer: uint16(tcpUrgentPointer), - }) - tcpHdr.SetChecksum(tcpHdr.CalculateChecksum(header.PseudoHeaderChecksum( - header.TCPProtocolNumber, - tcpip.AddrFrom4(arbitraryIPv4AddrB), - tcpip.AddrFrom4(arbitraryIPv4AddrB2), - header.TCPMinimumSize, - ))) - p.TransportProtocolNumber = header.TCPProtocolNumber - return p - }(), - op1: mustCreatePayloadLoad(t, linux.NFT_PAYLOAD_TRANSPORT_HEADER, tcpSrcPortOffset, tcpSrcPortLen, linux.NFT_REG_1), - op2: nil, + pkt: makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, fragmentedIPv4Fields(), arbitraryTCPFields()), + op1: mustCreatePayloadLoad(t, linux.NFT_PAYLOAD_TRANSPORT_HEADER, 0, 2, linux.NFT_REG_1), + op2: nil, }, { // cmd: add rule ip tab ch tcp sport 12345 tname: "load tcp header source port", @@ -1272,13 +1285,13 @@ func TestEvaluatePayloadLoad(t *testing.T) { op1: mustCreatePayloadLoad(t, linux.NFT_PAYLOAD_TRANSPORT_HEADER, tcpWindowOffset, tcpWindowLen, linux.NFT_REG_1), op2: mustCreateComparison(t, linux.NFT_REG_1, linux.NFT_CMP_EQ, newBytesData(numToBE(tcpWinSize, tcpWindowLen))), }, - { // cmd: add rule ip tab ch checksum __ + { // cmd: add rule ip tab ch tcp checksum __ tname: "load tcp header checksum", pkt: tcpPacket, op1: mustCreatePayloadLoad(t, linux.NFT_PAYLOAD_TRANSPORT_HEADER, tcpChecksumOffset, tcpChecksumLen, linux.NFT_REG_1), op2: mustCreateComparison(t, linux.NFT_REG_1, linux.NFT_CMP_EQ, newBytesData(numToBE(int(header.TCP(tcpPacket.TransportHeader().Slice()).Checksum()), tcpChecksumLen))), }, - { // cmd: add rule ip tab ch urgptr 0 + { // cmd: add rule ip tab ch tcp urgptr 0 tname: "load tcp header urgent pointer", pkt: tcpPacket, op1: mustCreatePayloadLoad(t, linux.NFT_PAYLOAD_TRANSPORT_HEADER, tcpUrgPtrOffset, tcpUrgPtrLen, linux.NFT_REG_1), @@ -1339,6 +1352,560 @@ func TestEvaluatePayloadLoad(t *testing.T) { } } +// TestEvaluatePayloadSet tests that the Payload Set operation correctly sets +// the payload from the source register and updates the packet checksums. +// The nft binary commands used to generate these are stated above each test. +// All commands should be preceded by nft --debug=netlink. +// TODO(b/339691111): Add tests for VLAN, ARP, ICMP, ICMPv6, IGMP, UDP headers. +func TestEvaluatePayloadSet(t *testing.T) { + for _, test := range []struct { + tname string + pkt *stack.PacketBuffer + outPkt *stack.PacketBuffer // nil if expecting a break during evaluation. + op1 operation // Immediate operation to load source register. + op2 operation // Payload Set operation to test. + }{ + // Ethernet header statement commands. + { // cmd: add rule ip tab ch ether saddr set 02:02:03:04:05:07 + tname: "set ethernet header source address", + pkt: makeEthernetPacket(0, arbitraryEthernetFields()), + outPkt: func() *stack.PacketBuffer { + fields := arbitraryEthernetFields() + fields.SrcAddr = arbitraryLinkAddr2 + return makeEthernetPacket(0, fields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(arbitraryLinkAddrB2[:])), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_LL_HEADER, ethSrcAddrOffset, ethSrcAddrLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_NONE, 0, 0x0), + }, + { // cmd: add rule ip tab ch ether daddr set 02:02:03:04:05:06 + tname: "set ethernet header destination address", + pkt: makeEthernetPacket(0, arbitraryEthernetFields()), + outPkt: func() *stack.PacketBuffer { + fields := arbitraryEthernetFields() + fields.DstAddr = arbitraryLinkAddr + return makeEthernetPacket(0, fields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_2, newBytesData(arbitraryLinkAddrB[:])), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_LL_HEADER, ethDstAddrOffset, ethDstAddrLen, linux.NFT_REG_2, linux.NFT_PAYLOAD_CSUM_NONE, 0, 0x0), + }, + { // cmd: add rule ip tab ch ether type set ip6 + tname: "set ethernet header type", + pkt: makeEthernetPacket(0, arbitraryEthernetFields()), + outPkt: func() *stack.PacketBuffer { + fields := arbitraryEthernetFields() + fields.Type = header.IPv6ProtocolNumber + return makeEthernetPacket(0, fields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(int(header.IPv6ProtocolNumber), ethTypeLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_LL_HEADER, ethTypeOffset, ethTypeLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_NONE, 0, 0x0), + }, + + // IPv4 header statement commands. + { // cmd: add rule ip tab ch ip length set 30 + tname: "set ipv4 header length", + pkt: makeIPv4Packet(header.IPv4MinimumSize, arbitraryIPv4Fields()), + outPkt: func() *stack.PacketBuffer { + fields := arbitraryIPv4Fields() + fields.TotalLength = uint16(30) + return makeIPv4Packet(header.IPv4MinimumSize, fields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(30, ipv4LengthLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv4LengthOffset, ipv4LengthLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 10, linux.NFT_PAYLOAD_L4CSUM_PSEUDOHDR), + }, + { // cmd: add rule ip tab ch ip id set 12345 + tname: "set ipv4 header ip id", + pkt: makeIPv4Packet(header.IPv4MinimumSize, arbitraryIPv4Fields()), + outPkt: func() *stack.PacketBuffer { + fields := arbitraryIPv4Fields() + fields.ID = uint16(12345) + return makeIPv4Packet(header.IPv4MinimumSize, fields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(12345, ipv4IDLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv4IDOffset, ipv4IDLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 10, 0x0), + }, + // Note: Fragment offsets are divided by 8 because they are in units of 8 + // bytes, which is encoded into the packet in IPv4.Encode(). + { // cmd: add rule ip tab ch ip frag-off set 2 (16 bytes) + tname: "set ipv4 header fragment offset, set fragment on", + pkt: makeIPv4Packet(header.IPv4MinimumSize, arbitraryIPv4Fields()), + outPkt: makeIPv4Packet(header.IPv4MinimumSize, fragmentedIPv4Fields()), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(arbitraryNonZeroFragmentOffset/8, ipv4FragOffLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv4FragOffOffset, ipv4FragOffLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 10, 0x0), + }, + { // cmd: add rule ip tab ch ip frag-off set 0 + tname: "set ipv4 header fragment offset, set fragment off for fragmented packet", + pkt: makeIPv4Packet(header.IPv4MinimumSize, fragmentedIPv4Fields()), + outPkt: makeIPv4Packet(header.IPv4MinimumSize, arbitraryIPv4Fields()), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(0, ipv4FragOffLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv4FragOffOffset, ipv4FragOffLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 10, 0x0), + }, + { // cmd: add rule ip tab ch ip frag-off set 10 (80 bytes) + tname: "set ipv4 header fragment offset, change fragment offset for fragmented packet", + pkt: makeIPv4Packet(header.IPv4MinimumSize, fragmentedIPv4Fields()), + outPkt: func() *stack.PacketBuffer { + fields := arbitraryIPv4Fields() + fields.FragmentOffset = uint16(10 * 8) + return makeIPv4Packet(header.IPv4MinimumSize, fields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(10, ipv4FragOffLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv4FragOffOffset, ipv4FragOffLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 10, 0x0), + }, + { // cmd: add rule ip tab ch ip ttl set 128 + tname: "set ipv4 time to live", + pkt: makeIPv4Packet(header.IPv4MinimumSize, arbitraryIPv4Fields()), + outPkt: func() *stack.PacketBuffer { + fields := arbitraryIPv4Fields() + fields.TTL = uint8(128) + return makeIPv4Packet(header.IPv4MinimumSize, fields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG32_01, newBytesData(numToBE(128, ipv4TTLLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv4TTLOffset, ipv4TTLLen, linux.NFT_REG32_01, linux.NFT_PAYLOAD_CSUM_INET, 10, 0x0), + }, + { // cmd: add rule ip tab ch ip saddr set 192.168.1.9 + tname: "set ipv4 header source address", + pkt: makeIPv4Packet(header.IPv4MinimumSize, arbitraryIPv4Fields()), + outPkt: func() *stack.PacketBuffer { + fields := arbitraryIPv4Fields() + fields.SrcAddr = tcpip.AddrFrom4(arbitraryIPv4AddrB2) + return makeIPv4Packet(header.IPv4MinimumSize, fields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(arbitraryIPv4AddrB2[:])), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv4SrcAddrOffset, ipv4SrcAddrLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 10, linux.NFT_PAYLOAD_L4CSUM_PSEUDOHDR), + }, + { // cmd: add rule ip tab ch ip daddr set 192.168.1.1 + tname: "set ipv4 header destination address", + pkt: makeIPv4Packet(header.IPv4MinimumSize, arbitraryIPv4Fields()), + outPkt: func() *stack.PacketBuffer { + fields := arbitraryIPv4Fields() + fields.DstAddr = tcpip.AddrFrom4(arbitraryIPv4AddrB) + return makeIPv4Packet(header.IPv4MinimumSize, fields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_4, newBytesData(arbitraryIPv4AddrB[:])), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv4DstAddrOffset, ipv4DstAddrLen, linux.NFT_REG_4, linux.NFT_PAYLOAD_CSUM_INET, 10, linux.NFT_PAYLOAD_L4CSUM_PSEUDOHDR), + }, + { // cmd: add rule ip tab ch ip checksum set 6060 + tname: "set ipv4 header checksum", + pkt: makeIPv4Packet(header.IPv4MinimumSize, arbitraryIPv4Fields()), + outPkt: func() *stack.PacketBuffer { + pkt := makeIPv4Packet(header.IPv4MinimumSize, arbitraryIPv4Fields()) + pkt.Network().SetChecksum(6060) + return pkt + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(6060, ipv4ChecksumLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv4ChecksumOffset, ipv4ChecksumLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 10, 0x0), + }, + + // IPv6 header statement commands. + { // cmd: add rule ip6 tab ch ip6 length set 232 + tname: "set ipv6 header length", + pkt: makeIPv6Packet(header.IPv6MinimumSize, arbitraryIPv6Fields()), + outPkt: func() *stack.PacketBuffer { + fields := arbitraryIPv6Fields() + fields.PayloadLength = uint16(232) + return makeIPv6Packet(header.IPv6MinimumSize, fields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(232, ipv6LengthLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv6LengthOffset, ipv6LengthLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_NONE, 0, linux.NFT_PAYLOAD_L4CSUM_PSEUDOHDR), + }, + { // cmd: add rule ip6 tab ch ip6 hoplimit set 54 + tname: "set ipv6 header hop limit", + pkt: makeIPv6Packet(header.IPv6MinimumSize, arbitraryIPv6Fields()), + outPkt: func() *stack.PacketBuffer { + fields := arbitraryIPv6Fields() + fields.HopLimit = uint8(54) + return makeIPv6Packet(header.IPv6MinimumSize, fields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(54, ipv6HopLimitLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv6HopLimitOffset, ipv6HopLimitLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_NONE, 0, 0x0), + }, + { // cmd: add rule ip6 tab ch ip6 saddr set 2001:db8:85a3::bb + tname: "set ipv6 header source address", + pkt: makeIPv6Packet(header.IPv6MinimumSize, arbitraryIPv6Fields()), + outPkt: func() *stack.PacketBuffer { + fields := arbitraryIPv6Fields() + fields.SrcAddr = tcpip.AddrFrom16(arbitraryIPv6AddrB2) + return makeIPv6Packet(header.IPv6MinimumSize, fields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(arbitraryIPv6AddrB2[:])), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv6SrcAddrOffset, ipv6SrcAddrLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_NONE, 0, linux.NFT_PAYLOAD_L4CSUM_PSEUDOHDR), + }, + { // cmd: add rule ip6 tab ch ip6 daddr set 2001:db8:85a3::aa + tname: "set ipv6 header destination address", + pkt: makeIPv6Packet(header.IPv6MinimumSize, arbitraryIPv6Fields()), + outPkt: func() *stack.PacketBuffer { + fields := arbitraryIPv6Fields() + fields.DstAddr = tcpip.AddrFrom16(arbitraryIPv6AddrB) + return makeIPv6Packet(header.IPv6MinimumSize, fields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_3, newBytesData(arbitraryIPv6AddrB[:])), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv6DstAddrOffset, ipv6DstAddrLen, linux.NFT_REG_3, linux.NFT_PAYLOAD_CSUM_NONE, 0, linux.NFT_PAYLOAD_L4CSUM_PSEUDOHDR), + }, + + // TCP with IPv4 header statement commands. + // TCP set commands. + { + // Since we change data within the transport header with a fragmented + // IPv4 packet, this can be problematic, so the evaluation should break. + tname: "set for transport header with a fragmented ipv4 packet", + pkt: makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, fragmentedIPv4Fields(), arbitraryTCPFields()), + outPkt: nil, + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(arbitraryPort, tcpSrcPortLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_TRANSPORT_HEADER, tcpSrcPortOffset, tcpSrcPortLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 16, 0x0), + }, + { // cmd: add rule ip tab ch tcp sport set 80 + tname: "set tcp header with ipv4 source port", + pkt: makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, arbitraryIPv4Fields(), arbitraryTCPFields()), + outPkt: func() *stack.PacketBuffer { + tcpFields := arbitraryTCPFields() + tcpFields.SrcPort = arbitraryPort2 + return makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, arbitraryIPv4Fields(), tcpFields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(arbitraryPort2, tcpSrcPortLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_TRANSPORT_HEADER, tcpSrcPortOffset, tcpSrcPortLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 16, 0x0), + }, + { // cmd: add rule ip tab ch tcp dport set 12345 + tname: "set tcp header with ipv4 destination port", + pkt: makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, arbitraryIPv4Fields(), arbitraryTCPFields()), + outPkt: func() *stack.PacketBuffer { + tcpFields := arbitraryTCPFields() + tcpFields.DstPort = arbitraryPort + return makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, arbitraryIPv4Fields(), tcpFields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(arbitraryPort, tcpDstPortLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_TRANSPORT_HEADER, tcpDstPortOffset, tcpDstPortLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 16, 0x0), + }, + { // cmd: add rule ip tab ch tcp sequence set 33 + tname: "set tcp header with ipv4 sequence number", + pkt: makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, arbitraryIPv4Fields(), arbitraryTCPFields()), + outPkt: func() *stack.PacketBuffer { + tcpFields := arbitraryTCPFields() + tcpFields.SeqNum = uint32(33) + return makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, arbitraryIPv4Fields(), tcpFields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(33, tcpSeqNumLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_TRANSPORT_HEADER, tcpSeqNumOffset, tcpSeqNumLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 16, 0x0), + }, + { // cmd: add rule ip tab ch tcp ackseq set 245 + tname: "set tcp header with ipv4 acknowledgement sequence number", + pkt: makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, arbitraryIPv4Fields(), arbitraryTCPFields()), + outPkt: func() *stack.PacketBuffer { + tcpFields := arbitraryTCPFields() + tcpFields.AckNum = uint32(245) + return makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, arbitraryIPv4Fields(), tcpFields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(245, tcpAckNumLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_TRANSPORT_HEADER, tcpAckNumOffset, tcpAckNumLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 16, 0x0), + }, + { // cmd: add rule ip tab ch tcp window set 91 + tname: "set tcp header with ipv4 window", + pkt: makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, arbitraryIPv4Fields(), arbitraryTCPFields()), + outPkt: func() *stack.PacketBuffer { + tcpFields := arbitraryTCPFields() + tcpFields.WindowSize = 91 + return makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, arbitraryIPv4Fields(), tcpFields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(91, tcpWindowLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_TRANSPORT_HEADER, tcpWindowOffset, tcpWindowLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 16, 0x0), + }, + { // cmd: add rule ip tab ch tcp checksum set 7654 + tname: "set tcp header with ipv4 checksum", + pkt: makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, arbitraryIPv4Fields(), arbitraryTCPFields()), + outPkt: func() *stack.PacketBuffer { + pkt := makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, arbitraryIPv4Fields(), arbitraryTCPFields()) + tcpHdr := header.TCP(pkt.TransportHeader().Slice()) + tcpHdr.SetChecksum(7654) + return pkt + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(7654, tcpChecksumLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_TRANSPORT_HEADER, tcpChecksumOffset, tcpChecksumLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 16, 0x0), + }, + { // cmd: add rule ip tab ch tcp urgptr set 40 + tname: "set tcp header with ipv4 urgent pointer", + pkt: makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, arbitraryIPv4Fields(), arbitraryTCPFields()), + outPkt: func() *stack.PacketBuffer { + tcpFields := arbitraryTCPFields() + tcpFields.UrgentPointer = 40 + return makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, arbitraryIPv4Fields(), tcpFields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(40, tcpUrgPtrLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_TRANSPORT_HEADER, tcpUrgPtrOffset, tcpUrgPtrLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 16, 0x0), + }, + // IPv4 set commands. + { // cmd: add rule ip tab ch ip id set 12345 + tname: "set ipv4 header with tcp ip id", + pkt: makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, arbitraryIPv4Fields(), arbitraryTCPFields()), + outPkt: func() *stack.PacketBuffer { + ipFields := arbitraryIPv4Fields() + ipFields.ID = uint16(12345) + return makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, ipFields, arbitraryTCPFields()) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(12345, ipv4IDLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv4IDOffset, ipv4IDLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 10, 0x0), + }, + { // cmd: add rule ip tab ch ip ttl set 128 + tname: "set ipv4 time to live", + pkt: makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, arbitraryIPv4Fields(), arbitraryTCPFields()), + outPkt: func() *stack.PacketBuffer { + ipFields := arbitraryIPv4Fields() + ipFields.TTL = uint8(128) + return makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, ipFields, arbitraryTCPFields()) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG32_01, newBytesData(numToBE(128, ipv4TTLLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv4TTLOffset, ipv4TTLLen, linux.NFT_REG32_01, linux.NFT_PAYLOAD_CSUM_INET, 10, 0x0), + }, + { // cmd: add rule ip tab ch ip saddr set 192.168.1.9 + tname: "set ipv4 header with tcp source address", + pkt: makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, arbitraryIPv4Fields(), arbitraryTCPFields()), + outPkt: func() *stack.PacketBuffer { + ipFields := arbitraryIPv4Fields() + ipFields.SrcAddr = tcpip.AddrFrom4(arbitraryIPv4AddrB2) + return makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, ipFields, arbitraryTCPFields()) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(arbitraryIPv4AddrB2[:])), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv4SrcAddrOffset, ipv4SrcAddrLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 10, linux.NFT_PAYLOAD_L4CSUM_PSEUDOHDR), + }, + { // cmd: add rule ip tab ch ip daddr set 192.168.1.1 + tname: "set ipv4 header with tcp destination address", + pkt: makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, arbitraryIPv4Fields(), arbitraryTCPFields()), + outPkt: func() *stack.PacketBuffer { + ipFields := arbitraryIPv4Fields() + ipFields.DstAddr = tcpip.AddrFrom4(arbitraryIPv4AddrB) + return makeIPv4TCPPacket(header.IPv4MinimumSize+header.TCPMinimumSize, ipFields, arbitraryTCPFields()) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_4, newBytesData(arbitraryIPv4AddrB[:])), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv4DstAddrOffset, ipv4DstAddrLen, linux.NFT_REG_4, linux.NFT_PAYLOAD_CSUM_INET, 10, linux.NFT_PAYLOAD_L4CSUM_PSEUDOHDR), + }, + + // TCP on IPv6 header statement commands. + // TCP set commands. + { // cmd: add rule ip tab ch tcp sport set 80 + tname: "set tcp header with ipv6 source port", + pkt: makeIPv6TCPPacket(header.IPv6MinimumSize+header.TCPMinimumSize, arbitraryIPv6Fields(), arbitraryTCPFields()), + outPkt: func() *stack.PacketBuffer { + tcpFields := arbitraryTCPFields() + tcpFields.SrcPort = arbitraryPort2 + return makeIPv6TCPPacket(header.IPv6MinimumSize+header.TCPMinimumSize, arbitraryIPv6Fields(), tcpFields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(arbitraryPort2, tcpSrcPortLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_TRANSPORT_HEADER, tcpSrcPortOffset, tcpSrcPortLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 16, 0x0), + }, + { // cmd: add rule ip tab ch tcp dport set 12345 + tname: "set tcp header with ipv6 destination port", + pkt: makeIPv6TCPPacket(header.IPv6MinimumSize+header.TCPMinimumSize, arbitraryIPv6Fields(), arbitraryTCPFields()), + outPkt: func() *stack.PacketBuffer { + tcpFields := arbitraryTCPFields() + tcpFields.DstPort = arbitraryPort + return makeIPv6TCPPacket(header.IPv6MinimumSize+header.TCPMinimumSize, arbitraryIPv6Fields(), tcpFields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(arbitraryPort, tcpDstPortLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_TRANSPORT_HEADER, tcpDstPortOffset, tcpDstPortLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 16, 0x0), + }, + { // cmd: add rule ip tab ch tcp sequence set 33 + tname: "set tcp header with ipv6 sequence number", + pkt: makeIPv6TCPPacket(header.IPv6MinimumSize+header.TCPMinimumSize, arbitraryIPv6Fields(), arbitraryTCPFields()), + outPkt: func() *stack.PacketBuffer { + tcpFields := arbitraryTCPFields() + tcpFields.SeqNum = uint32(33) + return makeIPv6TCPPacket(header.IPv6MinimumSize+header.TCPMinimumSize, arbitraryIPv6Fields(), tcpFields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(33, tcpSeqNumLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_TRANSPORT_HEADER, tcpSeqNumOffset, tcpSeqNumLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 16, 0x0), + }, + { // cmd: add rule ip tab ch tcp ackseq set 245 + tname: "set tcp header with ipv6 acknowledgement sequence number", + pkt: makeIPv6TCPPacket(header.IPv6MinimumSize+header.TCPMinimumSize, arbitraryIPv6Fields(), arbitraryTCPFields()), + outPkt: func() *stack.PacketBuffer { + tcpFields := arbitraryTCPFields() + tcpFields.AckNum = uint32(245) + return makeIPv6TCPPacket(header.IPv6MinimumSize+header.TCPMinimumSize, arbitraryIPv6Fields(), tcpFields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(245, tcpAckNumLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_TRANSPORT_HEADER, tcpAckNumOffset, tcpAckNumLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 16, 0x0), + }, + { // cmd: add rule ip tab ch tcp window set 91 + tname: "set tcp header with ipv6 window", + pkt: makeIPv6TCPPacket(header.IPv6MinimumSize+header.TCPMinimumSize, arbitraryIPv6Fields(), arbitraryTCPFields()), + outPkt: func() *stack.PacketBuffer { + tcpFields := arbitraryTCPFields() + tcpFields.WindowSize = uint16(91) + return makeIPv6TCPPacket(header.IPv6MinimumSize+header.TCPMinimumSize, arbitraryIPv6Fields(), tcpFields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(91, tcpWindowLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_TRANSPORT_HEADER, tcpWindowOffset, tcpWindowLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 16, 0x0), + }, + { // cmd: add rule ip tab ch tcp checksum set 7654 + tname: "set tcp header with ipv6 checksum", + pkt: makeIPv6TCPPacket(header.IPv6MinimumSize+header.TCPMinimumSize, arbitraryIPv6Fields(), arbitraryTCPFields()), + outPkt: func() *stack.PacketBuffer { + pkt := makeIPv6TCPPacket(header.IPv6MinimumSize+header.TCPMinimumSize, arbitraryIPv6Fields(), arbitraryTCPFields()) + tcpHdr := header.TCP(pkt.TransportHeader().Slice()) + tcpHdr.SetChecksum(7654) + return pkt + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(7654, tcpChecksumLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_TRANSPORT_HEADER, tcpChecksumOffset, tcpChecksumLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 16, 0x0), + }, + { // cmd: add rule ip tab ch tcp urgptr set 40 + tname: "set tcp header with ipv6 urgent pointer", + pkt: makeIPv6TCPPacket(header.IPv6MinimumSize+header.TCPMinimumSize, arbitraryIPv6Fields(), arbitraryTCPFields()), + outPkt: func() *stack.PacketBuffer { + tcpFields := arbitraryTCPFields() + tcpFields.UrgentPointer = uint16(40) + return makeIPv6TCPPacket(header.IPv6MinimumSize+header.TCPMinimumSize, arbitraryIPv6Fields(), tcpFields) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(40, tcpUrgPtrLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_TRANSPORT_HEADER, tcpUrgPtrOffset, tcpUrgPtrLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 16, 0x0), + }, + // IPv6 set commands. + { // cmd: add rule ip6 tab ch ip6 length set 232 + tname: "set ipv6 header length", + pkt: makeIPv6TCPPacket(header.IPv6MinimumSize+header.TCPMinimumSize, arbitraryIPv6Fields(), arbitraryTCPFields()), + outPkt: func() *stack.PacketBuffer { + fields := arbitraryIPv6Fields() + fields.PayloadLength = uint16(232) + return makeIPv6TCPPacket(header.IPv6MinimumSize+header.TCPMinimumSize, fields, arbitraryTCPFields()) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(232, ipv6LengthLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv6LengthOffset, ipv6LengthLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_NONE, 0, linux.NFT_PAYLOAD_L4CSUM_PSEUDOHDR), + }, + { // cmd: add rule ip6 tab ch ip6 hoplimit set 54 + tname: "set ipv6 header hop limit", + pkt: makeIPv6TCPPacket(header.IPv6MinimumSize+header.TCPMinimumSize, arbitraryIPv6Fields(), arbitraryTCPFields()), + outPkt: func() *stack.PacketBuffer { + fields := arbitraryIPv6Fields() + fields.HopLimit = uint8(54) + return makeIPv6TCPPacket(header.IPv6MinimumSize+header.TCPMinimumSize, fields, arbitraryTCPFields()) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(numToBE(54, ipv6HopLimitLen))), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv6HopLimitOffset, ipv6HopLimitLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_NONE, 0, 0x0), + }, + { // cmd: add rule ip6 tab ch ip6 saddr set 2001:db8:85a3::bb + tname: "set ipv6 header source address", + pkt: makeIPv6TCPPacket(header.IPv6MinimumSize+header.TCPMinimumSize, arbitraryIPv6Fields(), arbitraryTCPFields()), + outPkt: func() *stack.PacketBuffer { + fields := arbitraryIPv6Fields() + fields.SrcAddr = tcpip.AddrFrom16(arbitraryIPv6AddrB2) + return makeIPv6TCPPacket(header.IPv6MinimumSize+header.TCPMinimumSize, fields, arbitraryTCPFields()) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_1, newBytesData(arbitraryIPv6AddrB2[:])), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv6SrcAddrOffset, ipv6SrcAddrLen, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_NONE, 0, linux.NFT_PAYLOAD_L4CSUM_PSEUDOHDR), + }, + { // cmd: add rule ip6 tab ch ip6 daddr set 2001:db8:85a3::aa + tname: "set ipv6 header destination address", + pkt: makeIPv6TCPPacket(header.IPv6MinimumSize+header.TCPMinimumSize, arbitraryIPv6Fields(), arbitraryTCPFields()), + outPkt: func() *stack.PacketBuffer { + fields := arbitraryIPv6Fields() + fields.DstAddr = tcpip.AddrFrom16(arbitraryIPv6AddrB) + return makeIPv6TCPPacket(header.IPv6MinimumSize+header.TCPMinimumSize, fields, arbitraryTCPFields()) + }(), + op1: mustCreateImmediate(t, linux.NFT_REG_3, newBytesData(arbitraryIPv6AddrB[:])), + op2: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, ipv6DstAddrOffset, ipv6DstAddrLen, linux.NFT_REG_3, linux.NFT_PAYLOAD_CSUM_NONE, 0, linux.NFT_PAYLOAD_L4CSUM_PSEUDOHDR), + }, + } { + t.Run(test.tname, func(t *testing.T) { + // Sets up an NFTables object with a single table, chain, and rule. + nf := NewNFTables() + tab, err := nf.AddTable(arbitraryFamily, "test", "test table", false) + if err != nil { + t.Fatalf("unexpected error for AddTable: %v", err) + } + bc, err := tab.AddChain("base_chain", nil, "test chain", false) + if err != nil { + t.Fatalf("unexpected error for AddChain: %v", err) + } + bc.SetBaseChainInfo(arbitraryInfoPolicyAccept) + rule := &Rule{} + + // Adds testing operations. + if test.op1 != nil { + rule.addOperation(test.op1) + } + if test.op2 != nil { + rule.addOperation(test.op2) + } + + // Adds drop operation. Will be final verdict if payload set evaluation is + // successful (operation breaks if anything goes wrong). + rule.addOperation(mustCreateImmediate(t, linux.NFT_REG_VERDICT, newVerdictData(Verdict{Code: VC(linux.NF_DROP)}))) + + // Registers the rule to the base chain. + if err := bc.RegisterRule(rule, -1); err != nil { + t.Fatalf("unexpected error for RegisterRule: %v", err) + } + + // Runs evaluation. + v, err := nf.EvaluateHook(arbitraryFamily, arbitraryHook, test.pkt) + if err != nil { + t.Fatalf("unexpected error for EvaluateHook: %v", err) + } + + // Checks for final verdict. + if test.outPkt == nil { + // If no output packet is expected, then evaluation should break, + // resulting in Accept as the default policy verdict. + if v.Code != VC(linux.NF_ACCEPT) { + t.Fatalf("expected verdict Accept for break during evaluation, got %v", v) + } + return + } else { + // If an output packet is expected, the evaluation should go until end + // of rule (no errors/breaks), resulting in Drop as the final verdict. + if v.Code != VC(linux.NF_DROP) { + t.Fatalf("expected verdict Drop for successful evaluation, got %v", v) + } + } + + // Compares checksums first for resulting and expected packet. + if test.outPkt.NetworkProtocolNumber != test.pkt.NetworkProtocolNumber { + t.Fatalf("expected network protocol number %d for resulting packet, got %d", test.outPkt.NetworkProtocolNumber, test.pkt.NetworkProtocolNumber) + } + if test.pkt.NetworkHeader().View() != nil && test.outPkt.Network().Checksum() != test.pkt.Network().Checksum() { + t.Fatalf("expected network checksum %d for resulting packet, got %d", test.outPkt.Network().Checksum(), test.pkt.Network().Checksum()) + } + if test.pkt.TransportProtocolNumber != test.outPkt.TransportProtocolNumber { + t.Fatalf("expected transport protocol number %d for resulting packet, got %d", test.outPkt.TransportProtocolNumber, test.pkt.TransportProtocolNumber) + } + if test.pkt.TransportProtocolNumber != 0 { + var transport header.Transport + var transportOut header.Transport + switch tBytes, tOutBytes := test.pkt.TransportHeader().Slice(), + test.outPkt.TransportHeader().Slice(); test.pkt.TransportProtocolNumber { + case header.TCPProtocolNumber: + transport = header.TCP(tBytes) + transportOut = header.TCP(tOutBytes) + case header.UDPProtocolNumber: + transport = header.UDP(tBytes) + transportOut = header.UDP(tOutBytes) + case header.ICMPv4ProtocolNumber: + transport = header.ICMPv4(tBytes) + transportOut = header.ICMPv4(tOutBytes) + case header.ICMPv6ProtocolNumber: + transport = header.ICMPv6(tBytes) + transportOut = header.ICMPv6(tOutBytes) + case header.IGMPProtocolNumber: + transport = header.IGMP(tBytes) + transportOut = header.IGMP(tOutBytes) + } + if transport != nil && transport.Checksum() != transportOut.Checksum() { + t.Fatalf("expected transport checksum %d for resulting packet, got %d", transport.Checksum(), transportOut.Checksum()) + } + } + + // Compares raw packet data in bytes for resulting and expected packet. + actual := test.pkt.AsSlices() + expected := test.outPkt.AsSlices() + if len(actual) != len(expected) { + t.Fatalf("expected %d slices of data for the resulting packet, got %d", len(expected), len(actual)) + } + for i := range actual { + if !slices.Equal(actual[i], expected[i]) { + t.Fatalf("packet data does not match expected packet data (for slice %d)", i) + } + } + }) + } +} + // TestLoopCheckOnRegisterAndUnregister tests the loop checking and accompanying // logic on registering and unregistering rules. func TestLoopCheckOnRegisterAndUnregister(t *testing.T) { @@ -1790,7 +2357,7 @@ func TestLoopCheckOnRegisterAndUnregister(t *testing.T) { } // Runs evaluation and checks verdict. - pkt := makeArbitraryGeneralPacket(arbitraryReservedHeaderBytes) + pkt := makeArbitraryPacket(arbitraryReservedHeaderBytes) v, err := nf.EvaluateHook(arbitraryFamily, arbitraryHook, pkt) if err != nil { if test.verdict.ChainName != "error" { @@ -1897,7 +2464,7 @@ func TestMaxNestedJumps(t *testing.T) { } // Runs evaluation and checks verdict. - pkt := makeArbitraryGeneralPacket(arbitraryReservedHeaderBytes) + pkt := makeArbitraryPacket(arbitraryReservedHeaderBytes) v, err := nf.EvaluateHook(arbitraryFamily, arbitraryHook, pkt) if err != nil { if test.verdict.ChainName != "error" { @@ -1961,3 +2528,12 @@ func mustCreatePayloadLoad(t *testing.T, base payloadBase, offset, len, dreg uin } return pdload } + +// mustCreatePayloadSet wraps the NewPayloadSet function for brevity. +func mustCreatePayloadSet(t *testing.T, base payloadBase, offset uint8, len uint8, sreg uint8, csumType uint8, csumOff uint8, csumFlags uint8) *payloadSet { + pdset, err := newPayloadSet(base, offset, len, sreg, csumType, csumOff, csumFlags) + if err != nil { + t.Fatalf("failed to create payload set: %v", err) + } + return pdset +} diff --git a/pkg/tcpip/nftables/nftinterp.go b/pkg/tcpip/nftables/nftinterp.go index 61f3f2fe8..9db76b2bd 100644 --- a/pkg/tcpip/nftables/nftinterp.go +++ b/pkg/tcpip/nftables/nftinterp.go @@ -140,7 +140,13 @@ func InterpretOperation(line string, lnIdx int) (operation, error) { case "cmp": return InterpretComparison(line, lnIdx) case "payload": - return InterpretPayloadLoad(line, lnIdx) + switch tokens[2] { + case "load": + return InterpretPayloadLoad(line, lnIdx) + case "write": + return InterpretPayloadSet(line, lnIdx) + } + return nil, &SyntaxError{lnIdx, 2, fmt.Sprintf("unrecognized operation type: payload %s", tokens[2])} default: return nil, &SyntaxError{lnIdx, 1, fmt.Sprintf("unrecognized operation type: %s", tokens[1])} } @@ -361,6 +367,141 @@ func InterpretPayloadLoad(line string, lnIdx int) (operation, error) { return pdload, nil } +// InterpretPayloadSet creates a new PayloadSet operation from the given string. +func InterpretPayloadSet(line string, lnIdx int) (operation, error) { + tokens := strings.Fields(line) + + // Requires at least 19 tokens: + // "[", "payload", "write", "reg", register index, "=>", len+"b", "@", payload base, "header", "+", offset, + // "csum_type", checksum type, "csum_off", checksum offset, "csum_flags", checksum flags as hexadecimal, "]". + if len(tokens) != 19 { + return nil, &SyntaxError{lnIdx, 0, fmt.Sprintf("incorrect number of tokens for payload set operation, should be exactly 19, got %d", len(tokens))} + } + + if err := checkOperationBrackets(tokens, lnIdx); err != nil { + return nil, err + } + + tkIdx := 1 + + // First token should be "payload". + if err := consumeToken("payload", tokens, lnIdx, tkIdx); err != nil { + return nil, err + } + tkIdx++ + + // Second token should be "write". + if err := consumeToken("write", tokens, lnIdx, tkIdx); err != nil { + return nil, err + } + tkIdx++ + + // Third token should be "reg" + if err := consumeToken("reg", tokens, lnIdx, tkIdx); err != nil { + return nil, err + } + tkIdx++ + + // Fourth token should be the uint8 representing the register index. + reg, err := parseRegister(tokens[tkIdx], lnIdx, tkIdx) + if err != nil { + return nil, err + } + tkIdx++ + + // Fifth token should be "=>". + if err := consumeToken("=>", tokens, lnIdx, tkIdx); err != nil { + return nil, err + } + tkIdx++ + + // Sixth token should be the length (in bytes) of the payload followed by 'b'. + len, err := parsePayloadLength(tokens[tkIdx], lnIdx, tkIdx) + if err != nil { + return nil, err + } + tkIdx++ + + // Seventh token should be "@". + if err := consumeToken("@", tokens, lnIdx, tkIdx); err != nil { + return nil, err + } + tkIdx++ + + // Eighth token should be the payload base header. + base, err := parsePayloadBase(tokens[tkIdx], lnIdx, tkIdx) + if err != nil { + return nil, err + } + tkIdx++ + + // Ninth token should be "header". + if err := consumeToken("header", tokens, lnIdx, tkIdx); err != nil { + return nil, err + } + tkIdx++ + + // Tenth token should be "+". + if err := consumeToken("+", tokens, lnIdx, tkIdx); err != nil { + return nil, err + } + tkIdx++ + + // Eleventh token should be the uint8 representing the offset. + offset, err := parseUint8(tokens[tkIdx], "offset", lnIdx, tkIdx) + if err != nil { + return nil, err + } + tkIdx++ + + // Twelfth token should be "csum_type". + if err := consumeToken("csum_type", tokens, lnIdx, tkIdx); err != nil { + return nil, err + } + tkIdx++ + + // Thirteenth token should be the uint8 representing the checksum type. + csumType, err := parseUint8(tokens[tkIdx], "checksum type", lnIdx, tkIdx) + if err != nil { + return nil, err + } + tkIdx++ + + // Fourteenth token should be "csum_off". + if err := consumeToken("csum_off", tokens, lnIdx, tkIdx); err != nil { + return nil, err + } + tkIdx++ + + // Fifteenth token should be the uint8 representing the checksum offset. + csumOff, err := parseUint8(tokens[tkIdx], "checksum offset", lnIdx, tkIdx) + if err != nil { + return nil, err + } + tkIdx++ + + // Sixteenth token should be "csum_flags". + if err := consumeToken("csum_flags", tokens, lnIdx, tkIdx); err != nil { + return nil, err + } + tkIdx++ + + // Seventeenth token should be the uint8 representing checksum flags (in hex). + csumFlags, err := parseUint8(tokens[tkIdx], "checksum flags", lnIdx, tkIdx) + if err != nil { + return nil, err + } + tkIdx++ + + // Create the operation with the specified arguments. + pdset, err := newPayloadSet(base, offset, len, reg, csumType, csumOff, csumFlags) + if err != nil { + return nil, &LogicError{lnIdx, tkIdx, err} + } + + return pdset, nil +} + // // Interpreter Helper Functions. // @@ -378,8 +519,15 @@ func checkOperationBrackets(tokens []string, lnIdx int) error { } // parseUint8 parses the uint8 which should be supposed from the given string. +// Input starting with "0x" are parsed as base 16, otherwise assumes base 10. func parseUint8(regString string, supposed string, lnIdx int, tkIdx int) (uint8, error) { - v64, err := strconv.ParseUint(regString, 10, 8) + var v64 uint64 + var err error + if len(regString) > 2 && regString[:2] == "0x" { + v64, err = strconv.ParseUint(regString[2:], 16, 8) + } else { + v64, err = strconv.ParseUint(regString, 10, 8) + } if err != nil { return 0, &SyntaxError{lnIdx, tkIdx, fmt.Sprintf("could not parse uint8 %s: '%s'", supposed, regString)} } diff --git a/pkg/tcpip/nftables/nftinterp_test.go b/pkg/tcpip/nftables/nftinterp_test.go index 35d35d0e6..240d80f75 100644 --- a/pkg/tcpip/nftables/nftinterp_test.go +++ b/pkg/tcpip/nftables/nftinterp_test.go @@ -385,7 +385,7 @@ func TestInterpretPayloadLoadOps(t *testing.T) { opStr: "[ payload load 2b @ transport header + 0 => reg 0 ]", expected: nil, }, - // cmd: add rule ip6 ip tab ch tcp flags syn counter accept + // cmd: add rule ip tab ch tcp flags syn counter accept { tname: "load 1 byte into 4-byte register", opStr: "[ payload load 1b @ transport header + 13 => reg 9 ]", @@ -495,6 +495,175 @@ func checkPayloadLoadOp(tname string, expected operation, actual operation) erro return nil } +// TestInterpretPayloadSetOps tests interpretation of payload set operations. +// Most operations are direct output of nft binary commands. All stated commands +// should be preceded by nft --debug=netlink to generate matching operations. +func TestInterpretPayloadSetOps(t *testing.T) { + for _, test := range []interpretOperationTestAction{ + // Simple checksum type tests. + { + tname: "set checksum type, none", + opStr: "[ payload write reg 1 => 6b @ link header + 0 csum_type 0 csum_off 0 csum_flags 0x0 ]", + expected: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_LL_HEADER, 0, 6, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_NONE, 0, 0x0), + }, + { + tname: "set checksum type, inet", + opStr: "[ payload write reg 1 => 6b @ link header + 0 csum_type 1 csum_off 0 csum_flags 0x0 ]", + expected: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_LL_HEADER, 0, 6, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_INET, 0, 0x0), + }, + { + tname: "set checksum type, sctp", // not supported + opStr: "[ payload write reg 1 => 6b @ link header + 0 csum_type 2 csum_off 0 csum_flags 0x0 ]", + expected: nil, + }, + { + tname: "set out of range checksum type", + opStr: "[ payload write reg 1 => 6b @ link header + 0 csum_type 3 csum_off 0 csum_flags 0x0 ]", + expected: nil, + }, + // Simple checksum offset tests. + { + tname: "set valid offset", + opStr: "[ payload write reg 1 => 6b @ link header + 0 csum_type 0 csum_off 100 csum_flags 0x0 ]", + expected: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_LL_HEADER, 0, 6, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_NONE, 100, 0x0), + }, + { + tname: "set negative checksum offset", + opStr: "[ payload write reg 1 => 6b @ link header + 100 csum_type 1 csum_off -1 csum_flags 0x0 ]", + expected: nil, + }, + // Simple checksum flags tests. + { + tname: "set checksum flags, L4 with psuedoheader flag", + opStr: "[ payload write reg 1 => 6b @ link header + 0 csum_type 0 csum_off 0 csum_flags 0x1 ]", + expected: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_LL_HEADER, 0, 6, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_NONE, 0, linux.NFT_PAYLOAD_L4CSUM_PSEUDOHDR), + }, + { + tname: "set invalid checksum flags", + opStr: "[ payload write reg 1 => 6b @ link header + 0 csum_type 0 csum_off 0 csum_flags 0x2 ]", + expected: nil, + }, + // Invalid register tests. + { + tname: "set from verdict register", + opStr: "[ payload write reg 0 => 4b @ link header + 0 csum_type 0 csum_off 0 csum_flags 0x0 ]", + expected: nil, + }, + { + tname: "set >4 bytes from 4-byte register", + opStr: "[ payload write reg 9 => 6b @ link header + 0 csum_type 0 csum_off 0 csum_flags 0x0 ]", + expected: nil, + }, + { + tname: "set >16 bytes from 16-byte register", + opStr: "[ payload write reg 2 => 20b @ link header + 0 csum_type 0 csum_off 0 csum_flags 0x0 ]", + expected: nil, + }, + + // Valid tests. + // Note: It doesn't seem like the nft binary ever outputs payload set ops + // that have an odd offset or length and checksumming on. This makes sense + // because the offset and length are specified in bytes, but the checksum is + // calculated in half-words (2-bytes), which means the checksum calculation + // is only valid if the offset and length are even. However, the linux + // kernel does not specifically enforce this, so on linux it's technically + // possible to declare payload set operations that undoubtedly result in + // invalid checksums. Since the nft binary is what generates our input, we + // do not test these edge cases either. + + // cmd: add rule ip tab ch @nh,24,8 set 0xab + { + tname: "set 1 byte from 4-byte register with csum NONE and no flags", + opStr: "[ payload write reg 8 => 1b @ network header + 3 csum_type 0 csum_off 0 csum_flags 0x0 ]", + expected: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, 3, 1, linux.NFT_REG32_00, linux.NFT_PAYLOAD_CSUM_NONE, 0, 0x0), + }, + { + tname: "set 1 byte from 16-byte register with csum NONE and no flags", + opStr: "[ payload write reg 1 => 1b @ network header + 4 csum_type 0 csum_off 0 csum_flags 0x0 ]", + expected: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, 4, 1, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_NONE, 0, 0x0), + }, + // cmd: add rule ip tab ch tcp sport set 80 + { + tname: "set 2 bytes from 4-byte register with csum INET and no flags", + opStr: "[ payload write reg 9 => 2b @ transport header + 0 csum_type 1 csum_off 16 csum_flags 0x0 ]", + expected: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_TRANSPORT_HEADER, 0, 2, linux.NFT_REG32_01, linux.NFT_PAYLOAD_CSUM_INET, 16, 0x0), + }, + { + tname: "set 2 bytes from 16-byte register with csum INET and no flags", + opStr: "[ payload write reg 2 => 2b @ transport header + 0 csum_type 1 csum_off 16 csum_flags 0x0 ]", + expected: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_TRANSPORT_HEADER, 0, 2, linux.NFT_REG_2, linux.NFT_PAYLOAD_CSUM_INET, 16, 0x0), + }, + // cmd: add rule ip tab ch @ll,24,24 set 0xabcdef + { + tname: "set 3 bytes from 4-byte register with csum NONE and no flags", + opStr: "[ payload write reg 10 => 3b @ link header + 3 csum_type 0 csum_off 0 csum_flags 0x0 ]", + expected: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_LL_HEADER, 3, 3, linux.NFT_REG32_02, linux.NFT_PAYLOAD_CSUM_NONE, 0, 0x0), + }, + { + tname: "set 3 bytes from 16-byte register with csum NONE and no flags", + opStr: "[ payload write reg 2 => 3b @ link header + 3 csum_type 0 csum_off 0 csum_flags 0x0 ]", + expected: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_LL_HEADER, 3, 3, linux.NFT_REG_2, linux.NFT_PAYLOAD_CSUM_NONE, 0, 0x0), + }, + // cmd: add rule ip tab ch ip daddr set 192.168.1.1 + { + tname: "set 4 bytes from 4-byte register with csum INET and pseudoheader flag", + opStr: "[ payload write reg 11 => 4b @ network header + 16 csum_type 1 csum_off 10 csum_flags 0x1 ]", + expected: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, 16, 4, linux.NFT_REG32_03, linux.NFT_PAYLOAD_CSUM_INET, 10, linux.NFT_PAYLOAD_L4CSUM_PSEUDOHDR), + }, + { + tname: "set 4 bytes from 4-byte register with csum INET and pseudoheader flag", + opStr: "[ payload write reg 3 => 4b @ network header + 16 csum_type 1 csum_off 10 csum_flags 0x1 ]", + expected: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, 16, 4, linux.NFT_REG_3, linux.NFT_PAYLOAD_CSUM_INET, 10, linux.NFT_PAYLOAD_L4CSUM_PSEUDOHDR), + }, + // cmd: add rule ip tab ch ether saddr set 01:23:45:67:89:ab + { + tname: "set 6 bytes from 16-byte register with csum NONE and no flags", + opStr: "[ payload write reg 4 => 6b @ link header + 6 csum_type 0 csum_off 0 csum_flags 0x0 ]", + expected: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_LL_HEADER, 6, 6, linux.NFT_REG_4, linux.NFT_PAYLOAD_CSUM_NONE, 0, 0x0), + }, + // cmd: add rule ip6 tab ch ip6 saddr set 2001:db8::2 + { + tname: "set 16 bytes from 16-byte register with csum NONE and psuedoheader flag", + opStr: "[ payload write reg 1 => 16b @ network header + 8 csum_type 0 csum_off 0 csum_flags 0x1 ]", + expected: mustCreatePayloadSet(t, linux.NFT_PAYLOAD_NETWORK_HEADER, 8, 16, linux.NFT_REG_1, linux.NFT_PAYLOAD_CSUM_NONE, 0, linux.NFT_PAYLOAD_L4CSUM_PSEUDOHDR), + }, + } { + t.Run(test.tname, func(t *testing.T) { checkOp(t, test, checkPayloadSetOp) }) + } +} + +// checkPayloadSetOp checks that the given operation is a payload set +// operation and that it matches the expected payload set operation. +func checkPayloadSetOp(tname string, expected operation, actual operation) error { + expectedPdSet := expected.(*payloadSet) + pdset, ok := actual.(*payloadSet) + if !ok { + return fmt.Errorf("expected operation type to be PayloadLoad for %s, got %T", tname, actual) + } + if pdset.base != expectedPdSet.base { + return fmt.Errorf("expected payload base to be %v for %s, got %v", expectedPdSet.base, tname, pdset.base) + } + if pdset.offset != expectedPdSet.offset { + return fmt.Errorf("expected offset to be %d for %s, got %d", expectedPdSet.offset, tname, pdset.offset) + } + if pdset.blen != expectedPdSet.blen { + return fmt.Errorf("expected length to be %d for %s, got %d", expectedPdSet.blen, tname, pdset.blen) + } + if pdset.sreg != expectedPdSet.sreg { + return fmt.Errorf("expected destination register to be %d for %s, got %d", expectedPdSet.sreg, tname, pdset.sreg) + } + if pdset.csumType != expectedPdSet.csumType { + return fmt.Errorf("expected checksum type to be %d for %s, got %d", expectedPdSet.csumType, tname, pdset.csumType) + } + if pdset.csumOffset != expectedPdSet.csumOffset { + return fmt.Errorf("expected checksum offset to be %d for %s, got %d", expectedPdSet.csumOffset, tname, pdset.csumOffset) + } + if pdset.csumFlags != expectedPdSet.csumFlags { + return fmt.Errorf("expected checksum flags to be %b for %s, got %b", expectedPdSet.csumFlags, tname, pdset.csumFlags) + } + return nil +} + // TestInterpretRule tests the interpretation of basic and general rules as a // list of operations. func TestInterpretRule(t *testing.T) {