diff --git a/client/firewall/nftables/legacy_rule_linux_test.go b/client/firewall/nftables/legacy_rule_linux_test.go new file mode 100644 index 000000000..dc2f1c7a0 --- /dev/null +++ b/client/firewall/nftables/legacy_rule_linux_test.go @@ -0,0 +1,60 @@ +package nftables + +import ( + "testing" + + "github.com/google/nftables/expr" + "github.com/stretchr/testify/require" +) + +func TestBuildLegacyRouteRuleExpressions(t *testing.T) { + sourcePayload := &expr.Payload{} + sourceCmp := &expr.Cmp{} + destinationPayload := &expr.Payload{} + destinationCmp := &expr.Cmp{} + nilSourceDestination := &expr.Payload{} + nilDestinationSource := &expr.Cmp{} + + tests := []struct { + name string + source []expr.Any + destination []expr.Any + matches []expr.Any + }{ + { + name: "both non-empty", + source: []expr.Any{sourcePayload, sourceCmp}, + destination: []expr.Any{destinationPayload, destinationCmp}, + matches: []expr.Any{sourcePayload, sourceCmp, destinationPayload, destinationCmp}, + }, + { + name: "nil source", + destination: []expr.Any{nilSourceDestination}, + matches: []expr.Any{nilSourceDestination}, + }, + { + name: "nil destination", + source: []expr.Any{nilDestinationSource}, + matches: []expr.Any{nilDestinationSource}, + }, + { + name: "both nil", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := buildLegacyRouteRuleExpressions(tt.source, tt.destination) + + require.Len(t, got, len(tt.matches)+2) + for i, match := range tt.matches { + require.Same(t, match, got[i]) + } + + require.IsType(t, &expr.Counter{}, got[len(tt.matches)]) + verdict, ok := got[len(tt.matches)+1].(*expr.Verdict) + require.True(t, ok) + require.Equal(t, expr.VerdictAccept, verdict.Kind) + }) + } +} diff --git a/client/firewall/nftables/router_linux.go b/client/firewall/nftables/router_linux.go index 4214455a9..dfb94c514 100644 --- a/client/firewall/nftables/router_linux.go +++ b/client/firewall/nftables/router_linux.go @@ -953,6 +953,17 @@ func (r *router) addMSSClampingRules() error { return r.conn.Flush() } +func buildLegacyRouteRuleExpressions(sourceExp, destExp []expr.Any) []expr.Any { + exprs := make([]expr.Any, 0, len(sourceExp)+len(destExp)+2) + exprs = append(exprs, sourceExp...) + exprs = append(exprs, destExp...) + exprs = append(exprs, + &expr.Counter{}, + &expr.Verdict{Kind: expr.VerdictAccept}, + ) + return exprs +} + // addLegacyRouteRule adds a legacy routing rule for mgmt servers pre route acls func (r *router) addLegacyRouteRule(pair firewall.RouterPair) error { sourceExp, err := r.applyNetwork(pair.Source, nil, true) @@ -965,15 +976,7 @@ func (r *router) addLegacyRouteRule(pair firewall.RouterPair) error { return fmt.Errorf("apply destination: %w", err) } - exprs := []expr.Any{ - &expr.Counter{}, - &expr.Verdict{ - Kind: expr.VerdictAccept, - }, - } - - exprs = append(exprs, sourceExp...) - exprs = append(exprs, destExp...) + exprs := buildLegacyRouteRuleExpressions(sourceExp, destExp) ruleKey := firewall.GenKey(firewall.ForwardingFormat, pair)