diff --git a/pkg/tcpip/nftables/BUILD b/pkg/tcpip/nftables/BUILD index 4bb1a167b..46dcf7a6b 100644 --- a/pkg/tcpip/nftables/BUILD +++ b/pkg/tcpip/nftables/BUILD @@ -7,7 +7,10 @@ package( go_library( name = "nftables", - srcs = ["nftables.go"], + srcs = [ + "nftables.go", + "nftinterp.go", + ], deps = [ "//pkg/abi/linux", "//pkg/tcpip/stack", @@ -16,7 +19,10 @@ go_library( go_test( name = "nftables_test", - srcs = ["nftables_test.go"], + srcs = [ + "nftables_test.go", + "nftinterp_test.go", + ], library = ":nftables", deps = [ "//pkg/abi/linux", diff --git a/pkg/tcpip/nftables/nftinterp.go b/pkg/tcpip/nftables/nftinterp.go new file mode 100644 index 000000000..b222c9524 --- /dev/null +++ b/pkg/tcpip/nftables/nftinterp.go @@ -0,0 +1,334 @@ +// Copyright 2024 The gVisor Authors. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package nftables + +import ( + "encoding/hex" + "fmt" + "math" + "regexp" + "slices" + "strconv" + "strings" + + "gvisor.dev/gvisor/pkg/abi/linux" +) + +// SyntaxError is an interpretation error due to incorrect syntax. +type SyntaxError struct { + // lnIdx is the index of the line where the error occurred. + lnIdx int + // tkIdx is the index of the token where the error occurred. + tkIdx int + // msg is the error message, defined locally. + msg string +} + +// Error implements error interface for SyntaxError to return an error message. +func (e *SyntaxError) Error() string { + // Adds 1 to line index and token index to account for 0-indexing. + return fmt.Sprintf("syntax error at line %d, token %d: %s", e.lnIdx+1, e.tkIdx+1, e.msg) +} + +// LogicError is an interpretation error from modifying the NFTables state. +type LogicError struct { + // lnIdx is the index of the line where the error occurred. + lnIdx int + // tkIdx is the index of the token where the error occurred. + tkIdx int + // error is the error returned from modifying the NFTables state. + err error +} + +// Error implements error interface for LogicError to return an error message. +func (e *LogicError) Error() string { + // Adds 1 to line index and token index to account for 0-indexing. + return fmt.Sprintf("logic error at line %d, token %d: %v", e.lnIdx+1, e.tkIdx+1, e.err) +} + +// Note: this is a limited set of keywords. +var reservedKeywords []string = []string{ + "include", // include keyword + "define", "undefine", "redefine", // symbolic variables keywords + "ip", "ip6", "inet", "arp", "bridge", "netdev", // address families + "list", "flush", "ruleset", // ruleset operations + "add", "create", "delete", "destroy", "table", "comment", "flags", "handle", // table operations + "rename", "chain", "type", "hook", "device", "priority", "policy", // chain operations + "insert", "reset", "replace", "rule", "index", // rule operations +} + +// Set of reserved specifiers for quick lookup. +var reservedKeywordSet map[string]struct{} = initReservedKeywordSet() + +func initReservedKeywordSet() map[string]struct{} { + set := make(map[string]struct{}) + for _, k := range reservedKeywords { + set[k] = struct{}{} + } + return set +} + +var identifierRegexp = regexp.MustCompile("^[a-zA-Z_][a-zA-Z0-9_/.]*$") + +// validateIdentifier checks if the identifier is valid. +// An identifier is valid if it is not a reserved keyword and begins with an +// alphabetic character or underscore followed by zero or more alphanumeric +// characters, underscores, forward slashes, or periods. +func validateIdentifier(id string, lnIdx int, tkIdx int) error { + if _, ok := reservedKeywordSet[id]; ok { + return &SyntaxError{lnIdx, tkIdx, fmt.Sprintf("cannot use reserved keyword %s as an identifier", id)} + } + + if !identifierRegexp.MatchString(id) { + return &SyntaxError{lnIdx, tkIdx, fmt.Sprintf("invalid identifier %s", id)} + } + + return nil +} + +// InterpretRule creates a new Rule from the given rule string, assumed to be +// represented as a block of text with a single operation per line. +// Note: the rule string should be generated as output from the official nft +// binary (can be accomplished by using flag --debug=netlink). +func InterpretRule(ruleString string) (*Rule, error) { + ruleString = strings.TrimSpace(ruleString) + lines := slices.DeleteFunc(strings.Split(ruleString, "\n"), func(s string) bool { + return s == "" + }) + + r := &Rule{ops: make([]Operation, 0, len(lines))} + + // Interprets all operations in the rule. + for lnIdx, line := range lines { + op, err := InterpretOperation(line, lnIdx) + if err != nil { + return nil, err + } + r.AddOperation(op) + } + + return r, nil +} + +// InterpretOperation creates a new Operation from the given operation string, +// assumed to be a single line of text surrounded in square brackets. +// Note: the operation string should be generated as output from the official nft +// binary (can be accomplished by using flag --debug=netlink). +func InterpretOperation(line string, lnIdx int) (Operation, error) { + tokens := strings.Fields(line) + if len(tokens) < 2 { + return nil, &SyntaxError{lnIdx, 0, fmt.Sprintf("incorrect number of tokens for operation, should be at least 2, got %d", len(tokens))} + } + + // Second token decides the operation type. + switch tokens[1] { + case "immediate": + return InterpretImmediate(line, lnIdx) + default: + return nil, &SyntaxError{lnIdx, 1, fmt.Sprintf("unrecognized operation type: %s", tokens[1])} + } +} + +// InterpretImmediate creates a new Immediate operation from the given string. +func InterpretImmediate(line string, lnIdx int) (Operation, error) { + tokens := strings.Fields(line) + + // Requires at least 6 tokens: + // "[", "immediate", "reg", register index, register value, "]". + if len(tokens) < 6 { + return nil, &SyntaxError{lnIdx, 0, fmt.Sprintf("incorrect number of tokens for immediate operation, should be at least 6, got %d", len(tokens))} + } + + if err := checkOperationBrackets(tokens, lnIdx); err != nil { + return nil, err + } + + tkIdx := 1 + + // First token should be "immediate". + if err := consumeToken("immediate", tokens, lnIdx, tkIdx); err != nil { + return nil, err + } + tkIdx++ + + // Second token should be "reg". + if err := consumeToken("reg", tokens, lnIdx, tkIdx); err != nil { + return nil, err + } + tkIdx++ + + // Third token should be the uint8 representing the register index. + reg, err := parseRegister(tokens[tkIdx], lnIdx, tkIdx) + if err != nil { + return nil, err + } + tkIdx++ + + // Fourth token should be the value. + nextIdx, data, err := parseRegisterData(reg, tokens, lnIdx, tkIdx) + if err != nil { + return nil, err + } + tkIdx = nextIdx + + // Done parsing tokens. + if tkIdx != len(tokens)-1 { + return nil, &SyntaxError{lnIdx, tkIdx, "unexpected token after immediate operation"} + } + + // Create the operation with the specified arguments. + imm, err := NewImmediate(reg, data) + if err != nil { + return nil, &LogicError{lnIdx, tkIdx, err} + } + + return imm, nil +} + +// +// Interpreter Helper Functions. +// + +// checkOperationBrackets checks that the operation string is surrounded by +// square brackets. +func checkOperationBrackets(tokens []string, lnIdx int) error { + if tokens[0] != "[" { + return &SyntaxError{lnIdx, 0, "operation missing opening square bracket"} + } + if tokens[len(tokens)-1] != "]" { + return &SyntaxError{lnIdx, len(tokens) - 1, "operation missing closing square bracket"} + } + return nil +} + +// parseRegister parses the register index from the given string. +func parseRegister(regString string, lnIdx int, tkIdx int) (uint8, error) { + reg64, err := strconv.ParseUint(regString, 10, 8) + if err != nil { + return 0, &SyntaxError{lnIdx, tkIdx, fmt.Sprintf("could not parse uint8 register index: '%s'", regString)} + } + + if reg64 > math.MaxUint8 || !isRegister(uint8(reg64)) { + return 0, &SyntaxError{lnIdx, tkIdx, fmt.Sprintf("invalid register index: %d", reg64)} + } + + return uint8(reg64), nil +} + +// parseRegisterData parses the register data from the given token and returns +// the index of the next token to process (can consume multiple tokens). +// Note: assumes the register index is valid (was checked in parseRegister). +func parseRegisterData(reg uint8, tokens []string, lnIdx int, tkIdx int) (int, RegisterData, error) { + // Handles verdict data. + if isVerdictRegister(reg) { + nextIdx, verdict, err := parseVerdict(tokens, lnIdx, tkIdx) + if err != nil { + return 0, nil, err + } + return nextIdx, NewVerdictData(verdict), nil + } + // Handles hex data (4- or 16-byte). + if len(tokens[tkIdx]) > 1 && tokens[tkIdx][:2] == "0x" { + nextIdx, data, err := parseHexData(tokens, lnIdx, tkIdx) + if err != nil { + return 0, nil, err + } + // Validates the register data type. 4-byte data is valid for both 4- and + // 16-byte registers, but 16-byte data is only valid for 16-byte registers. + if err := data.ValidateRegister(reg); err != nil { + return 0, nil, &LogicError{lnIdx, tkIdx, err} + } + return nextIdx, data, nil + } + // TODO(b/345684870): cases will be added here as more types are supported. + return 0, nil, &SyntaxError{lnIdx, tkIdx, fmt.Sprintf("invalid register data: '%s'", tokens[tkIdx])} +} + +// parseVerdict parses the verdict from the given token and returns +// the index of the next token to process (can consume multiple tokens). +func parseVerdict(tokens []string, lnIdx int, tkIdx int) (int, Verdict, error) { + v := Verdict{} + + switch tokens[tkIdx] { + case "accept": + v.Code = VC(linux.NF_ACCEPT) + case "drop": + v.Code = VC(linux.NF_DROP) + case "continue": + v.Code = VC(linux.NFT_CONTINUE) + case "return": + v.Code = VC(linux.NFT_RETURN) + case "jump": + v.Code = VC(linux.NFT_JUMP) + case "goto": + v.Code = VC(linux.NFT_GOTO) + default: + return 0, v, &SyntaxError{lnIdx, tkIdx, fmt.Sprintf("invalid verdict: '%s'", tokens[tkIdx])} + } + tkIdx++ + + // jump and chain verdicts require 2 more tokens to specify the target chain. + // "jump"/"goto", "->", chain name. + switch v.Code { + case VC(linux.NFT_JUMP), VC(linux.NFT_GOTO): + if err := consumeToken("->", tokens, lnIdx, tkIdx); err != nil { + return 0, v, err + } + tkIdx++ + + if err := validateIdentifier(tokens[tkIdx], lnIdx, tkIdx); err != nil { + return 0, v, err + } + v.ChainName = tokens[tkIdx] + tkIdx++ + } + + return tkIdx, v, nil +} + +// parseHexData parses little endian hexadecimal data from the given token and +// returns the index of the next token to process (can consume multiple tokens). +func parseHexData(tokens []string, lnIdx int, tkIdx int) (int, RegisterData, error) { + var bytes []byte + for ; tkIdx < len(tokens); tkIdx++ { + if len(tokens[tkIdx]) < 2 || tokens[tkIdx][:2] != "0x" { + break + } + + if len(tokens[tkIdx]) != 10 { + return 0, nil, &SyntaxError{lnIdx, tkIdx, fmt.Sprintf("hexadecimal data must be exactly 8 digits long (excluding 0x): '%s'", tokens[tkIdx])} + } + + // Decodes the little endian hex string into bytes + bytes4, err := hex.DecodeString(tokens[tkIdx][2:]) + if err != nil { + return 0, nil, &SyntaxError{lnIdx, tkIdx, fmt.Sprintf("could not decode hexadecimal data: '%s'", tokens[tkIdx])} + } + bytes = append(bytes, bytes4...) + } + if len(bytes) == 4 || len(bytes) == 16 { + return tkIdx, NewBytesData(bytes), nil + } + return 0, nil, &SyntaxError{lnIdx, tkIdx, fmt.Sprintf("incorrect number of bytes for hexadecimal data, should be 4 or 16, got %d", len(bytes))} +} + +// consumeToken is a helper function that checks if the token at the given index +// matches the expected string, returning a SyntaxError if not. +func consumeToken(expected string, tokens []string, lnIdx int, tkIdx int) error { + if tokens[tkIdx] != expected { + return &SyntaxError{lnIdx, tkIdx, fmt.Sprintf("unexpected string: %s", tokens[tkIdx])} + } + return nil +} diff --git a/pkg/tcpip/nftables/nftinterp_test.go b/pkg/tcpip/nftables/nftinterp_test.go new file mode 100644 index 000000000..4fce65597 --- /dev/null +++ b/pkg/tcpip/nftables/nftinterp_test.go @@ -0,0 +1,201 @@ +// Copyright 2024 The gVisor Authors. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package nftables + +import ( + "fmt" + "testing" + + "gvisor.dev/gvisor/pkg/abi/linux" +) + +func TestInterpretImmediateOps(t *testing.T) { + for _, test := range []struct { + tname string + opStr string + op *Immediate // will be nil if an error is expected + }{ + { + tname: "verdict register with accept verdict", + opStr: "[ immediate reg 0 accept ]", + op: mustCreateImmediate(t, linux.NFT_REG_VERDICT, NewVerdictData(Verdict{Code: VC(linux.NF_ACCEPT)})), + }, + { + tname: "verdict register with drop verdict", + opStr: "[ immediate reg 0 drop ]", + op: mustCreateImmediate(t, linux.NFT_REG_VERDICT, NewVerdictData(Verdict{Code: VC(linux.NF_DROP)})), + }, + { + tname: "verdict register with continue verdict", + opStr: "[ immediate reg 0 continue ]", + op: mustCreateImmediate(t, linux.NFT_REG_VERDICT, NewVerdictData(Verdict{Code: VC(linux.NFT_CONTINUE)})), + }, + { + tname: "verdict register with return verdict", + opStr: "[ immediate reg 0 return ]", + op: mustCreateImmediate(t, linux.NFT_REG_VERDICT, NewVerdictData(Verdict{Code: VC(linux.NFT_RETURN)})), + }, + { + tname: "verdict register with jump verdict", + opStr: "[ immediate reg 0 jump -> next_chain ]", + op: mustCreateImmediate(t, linux.NFT_REG_VERDICT, NewVerdictData(Verdict{Code: VC(linux.NFT_JUMP), ChainName: "next_chain"})), + }, + { + tname: "verdict register with goto verdict", + opStr: "[ immediate reg 0 goto -> next_chain ]", + op: mustCreateImmediate(t, linux.NFT_REG_VERDICT, NewVerdictData(Verdict{Code: VC(linux.NFT_GOTO), ChainName: "next_chain"})), + }, + { + tname: "verdict register with 4-byte data", + opStr: "[ immediate reg 0 0x0201a8c0 ]", + op: nil, + }, + { + tname: "verdict register with 16-byte data", + opStr: "[ immediate reg 0 0xb80d0120 0x00000000 0x00000000 0x02000000 ]", + op: nil, + }, + { + tname: "16-byte register with verdict data", + opStr: "[ immediate reg 1 accept ]", + op: nil, + }, + { + tname: "16-byte register with verdict data with target", + opStr: "[ immediate reg 2 jump -> next_chain ]", + op: nil, + }, + { + tname: "16-byte register with 4-byte data", + opStr: "[ immediate reg 3 0x0201a8c0 ]", + op: mustCreateImmediate(t, linux.NFT_REG_3, NewBytesData([]byte{0x02, 0x01, 0xa8, 0xc0})), + }, + { + tname: "16-byte register with 16-byte data", + opStr: "[ immediate reg 4 0xb80d0120 0x00000000 0x00000000 0x02000000 ]", + op: mustCreateImmediate(t, linux.NFT_REG_4, NewBytesData([]byte{0xb8, 0x0d, 0x01, 0x20, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x02, 0x00, 0x00, 0x00})), + }, + { + tname: "16-byte register with 8-byte data", + opStr: "[ immediate reg 4 0xb80d0120 0x00000050 ]", + op: nil, + }, + { + tname: "4-byte register with verdict data", + opStr: "[ immediate reg 8 return ]", + op: nil, + }, + { + tname: "4-byte register with verdict data with target", + opStr: "[ immediate reg 9 goto -> next_chain ]", + op: nil, + }, + { + tname: "4-byte register with 4-byte data", + opStr: "[ immediate reg 10 0x0201a8c0 ]", + op: mustCreateImmediate(t, linux.NFT_REG32_02, NewBytesData([]byte{0x02, 0x01, 0xa8, 0xc0})), + }, + { + tname: "4-byte register with 16-byte data", + opStr: "[ immediate reg 9 0xb80d0120 0x00000000 0x00000000 0x02000000 ]", + op: nil, + }, + } { + t.Run(test.tname, func(t *testing.T) { + rule, err := InterpretRule(test.opStr) + if err != nil { + if test.op == nil { + return + } + t.Fatalf("unexpected interpretation error for %s: %v", test.tname, err) + } + + if len(rule.ops) != 1 { + t.Fatalf("expected single operation for %s, got %d", test.tname, len(rule.ops)) + } + op := rule.ops[0] + if err := checkImmediateOp(test.tname, test.op, op); err != nil { + t.Fatalf(err.Error()) + } + }) + } +} + +func checkImmediateOp(tname string, expected *Immediate, actual Operation) error { + if actual == nil { + return fmt.Errorf("expected non-nil operation for %s, got nil", tname) + } + imm, ok := actual.(*Immediate) + if !ok { + return fmt.Errorf("expected operation type to be Immediate for %s, got %s", tname, actual.TypeString()) + } + if imm.dreg != expected.dreg { + return fmt.Errorf("expected register to be %d for %s, got %d", expected.dreg, tname, imm.dreg) + } + if !imm.data.Equal(expected.data) { + return fmt.Errorf("expected data to be %s for %s, got %s", expected.data.String(), tname, imm.data.String()) + } + return nil +} + +func TestInterpretRule(t *testing.T) { + for _, test := range []struct { + tname string + ruleStr string + rule *Rule // will be nil if an error is expected + }{ + { + tname: "empty ruleset", + ruleStr: ``, + rule: &Rule{}, + }, + { + tname: "empty ruleset with excess whitespace", + ruleStr: ` + + + `, + rule: &Rule{}, + }, + } { + t.Run(test.tname, func(t *testing.T) { + rule, err := InterpretRule(test.ruleStr) + if err != nil { + if test.rule == nil { + return + } + t.Fatalf("unexpected interpretation error for %s: %v", test.tname, err) + } + + if len(rule.ops) != len(test.rule.ops) { + t.Fatalf("expected %d operations for %s, got %d", len(test.rule.ops), test.tname, len(rule.ops)) + } + + // Checks each operation in the rule with the appropriate check function. + for i, op := range rule.ops { + testOp := test.rule.ops[i] + switch testOp.(type) { + case *Immediate: + if err := checkImmediateOp(test.tname, testOp.(*Immediate), op); err != nil { + t.Fatalf(err.Error()) + } + // TODO(b/345684870): cases will be added here as more types are supported. + default: + t.Fatalf("unexpected operation type for %s: %s", test.tname, testOp.TypeString()) + } + } + }) + } +}