From 9976bc9190b5a0c5cfb111eb11f3a34907e73b92 Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Sat, 3 Oct 2026 20:41:42 +1000 Subject: [PATCH 01/11] Match any listed port or protocol instead of AND-ing them --- internal/nftables/compiler.go | 356 ++++++++++++++++------------- internal/nftables/compiler_test.go | 94 ++++++++ 2 files changed, 286 insertions(+), 164 deletions(-) diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 45e6549..2af98ad 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -220,34 +220,30 @@ func (c *Compiler) compileConntrack(state *FirewallState) error { chains = []string{"prerouting", "output"} } + matches, err := l4Matches(ct.Proto, ct.DPort, nil) + if err != nil { + return fmt.Errorf("conntrack[%d]: %w", i, err) + } + for _, chain := range chains { - var exprs []expr.Any + for _, m := range matches { + exprs := append([]expr.Any{}, m.exprs...) - if ct.Proto != "" { - exprs = append(exprs, matchProto(ct.Proto)...) - } - for _, p := range ct.DPort { - pe, err := parsePortOrRange(p) - if err != nil { - return fmt.Errorf("conntrack[%d]: %w", i, err) + switch ct.Action { + case config.ConntrackNoTrack: + exprs = append(exprs, &expr.Notrack{}) + case config.ConntrackHelper: + continue + case config.ConntrackDrop: + exprs = append(exprs, &expr.Verdict{Kind: expr.VerdictDrop}) } - exprs = append(exprs, pe...) - } - switch ct.Action { - case config.ConntrackNoTrack: - exprs = append(exprs, &expr.Notrack{}) - case config.ConntrackHelper: - continue - case config.ConntrackDrop: - exprs = append(exprs, &expr.Verdict{Kind: expr.VerdictDrop}) + state.Rules[chain] = append(state.Rules[chain], ManagedRule{ + Chain: chain, + Exprs: exprs, + Tag: tag + ":" + chain, + }) } - - state.Rules[chain] = append(state.Rules[chain], ManagedRule{ - Chain: chain, - Exprs: exprs, - Tag: tag + ":" + chain, - }) } } return nil @@ -362,29 +358,33 @@ func (c *Compiler) compileOneRule(state *FirewallState, tag, srcSpec, dstSpec, p for _, srcIface := range srcIfaces { for _, dstIface := range dstIfaces { - exprs, err := c.buildMatchExprs(srcIface, dstIface, chain, proto, dports, sports, srcAddr, dstAddr) + matches, err := c.buildMatchExprs(srcIface, dstIface, chain, proto, dports, sports, srcAddr, dstAddr) if err != nil { return err } - if section != "" && section != config.SectionAll { - exprs = append(exprs, matchSection(section)...) - } + for _, m := range matches { + exprs := m.exprs - if logLevel != "" { - exprs = append(exprs, buildLog(logLevel, tag)...) - } + if section != "" && section != config.SectionAll { + exprs = append(exprs, matchSection(section)...) + } - verdict := actionVerdict(action, proto, c.cfg.Settings.AddressFamily) - if verdict != nil { - exprs = append(exprs, verdict...) - } + if logLevel != "" { + exprs = append(exprs, buildLog(logLevel, tag)...) + } - state.Rules[chain] = append(state.Rules[chain], ManagedRule{ - Chain: chain, - Exprs: exprs, - Tag: tag, - }) + verdict := actionVerdict(action, m.proto, c.cfg.Settings.AddressFamily) + if verdict != nil { + exprs = append(exprs, verdict...) + } + + state.Rules[chain] = append(state.Rules[chain], ManagedRule{ + Chain: chain, + Exprs: exprs, + Tag: tag, + }) + } } } return nil @@ -413,102 +413,99 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcSpec, dstSpec, srcIfaces := c.resolveZoneInterfaces(srcZone) + matches, err := l4Matches(proto, dports, nil) + if err != nil { + return err + } + for _, srcIface := range srcIfaces { - var exprs []expr.Any + for _, m := range matches { + var exprs []expr.Any - if srcIface != "" { - exprs = append(exprs, matchIfaceName(true, srcIface)...) - } - - if srcAddr != "" { - src, err := matchSourceCIDR(srcAddr) - if err != nil { - return err + if srcIface != "" { + exprs = append(exprs, matchIfaceName(true, srcIface)...) } - exprs = append(exprs, src...) - } - if proto != "" { - exprs = append(exprs, matchProto(proto)...) - } - - for _, portStr := range dports { - pe, err := parsePortOrRange(portStr) - if err != nil { - return err - } - exprs = append(exprs, pe...) - } - - if logLevel != "" { - exprs = append(exprs, buildLog(logLevel, tag)...) - } - - ip := net.ParseIP(dnatAddr) - if ip == nil { - return fmt.Errorf("invalid DNAT address %q", dnatAddr) - } - - if action == config.RuleRedirect { - if dnatPort > 0 { - portBytes := make([]byte, 2) - binary.BigEndian.PutUint16(portBytes, dnatPort) - exprs = append(exprs, - &expr.Immediate{Register: 1, Data: portBytes}, - &expr.Redir{RegisterProtoMin: 1}, - ) - } else { - exprs = append(exprs, &expr.Redir{}) - } - } else { - if ip4 := ip.To4(); ip4 != nil { - exprs = append(exprs, - &expr.Immediate{Register: 1, Data: ip4}, - ) - natExpr := &expr.NAT{ - Type: expr.NATTypeDestNAT, - Family: unix.NFPROTO_IPV4, - RegAddrMin: 1, - RegAddrMax: 1, + if srcAddr != "" { + src, err := matchSourceCIDR(srcAddr) + if err != nil { + return err } + exprs = append(exprs, src...) + } + + exprs = append(exprs, m.exprs...) + + if logLevel != "" { + exprs = append(exprs, buildLog(logLevel, tag)...) + } + + ip := net.ParseIP(dnatAddr) + if ip == nil { + return fmt.Errorf("invalid DNAT address %q", dnatAddr) + } + + if action == config.RuleRedirect { if dnatPort > 0 { portBytes := make([]byte, 2) binary.BigEndian.PutUint16(portBytes, dnatPort) exprs = append(exprs, - &expr.Immediate{Register: 2, Data: portBytes}, + &expr.Immediate{Register: 1, Data: portBytes}, + &expr.Redir{RegisterProtoMin: 1}, ) - natExpr.RegProtoMin = 2 - natExpr.RegProtoMax = 2 + } else { + exprs = append(exprs, &expr.Redir{}) } - exprs = append(exprs, natExpr) } else { - exprs = append(exprs, - &expr.Immediate{Register: 1, Data: ip.To16()}, - ) - natExpr := &expr.NAT{ - Type: expr.NATTypeDestNAT, - Family: unix.NFPROTO_IPV6, - RegAddrMin: 1, - RegAddrMax: 1, - } - if dnatPort > 0 { - portBytes := make([]byte, 2) - binary.BigEndian.PutUint16(portBytes, dnatPort) + if ip4 := ip.To4(); ip4 != nil { exprs = append(exprs, - &expr.Immediate{Register: 2, Data: portBytes}, + &expr.Immediate{Register: 1, Data: ip4}, ) - natExpr.RegProtoMin = 2 - natExpr.RegProtoMax = 2 + natExpr := &expr.NAT{ + Type: expr.NATTypeDestNAT, + Family: unix.NFPROTO_IPV4, + RegAddrMin: 1, + RegAddrMax: 1, + } + if dnatPort > 0 { + portBytes := make([]byte, 2) + binary.BigEndian.PutUint16(portBytes, dnatPort) + exprs = append(exprs, + &expr.Immediate{Register: 2, Data: portBytes}, + ) + natExpr.RegProtoMin = 2 + natExpr.RegProtoMax = 2 + } + exprs = append(exprs, natExpr) + } else { + exprs = append(exprs, + &expr.Immediate{Register: 1, Data: ip.To16()}, + ) + natExpr := &expr.NAT{ + Type: expr.NATTypeDestNAT, + Family: unix.NFPROTO_IPV6, + RegAddrMin: 1, + RegAddrMax: 1, + } + if dnatPort > 0 { + portBytes := make([]byte, 2) + binary.BigEndian.PutUint16(portBytes, dnatPort) + exprs = append(exprs, + &expr.Immediate{Register: 2, Data: portBytes}, + ) + natExpr.RegProtoMin = 2 + natExpr.RegProtoMax = 2 + } + exprs = append(exprs, natExpr) } - exprs = append(exprs, natExpr) } - } - state.Rules[chain] = append(state.Rules[chain], ManagedRule{ - Chain: chain, - Exprs: exprs, - Tag: tag, - }) + state.Rules[chain] = append(state.Rules[chain], ManagedRule{ + Chain: chain, + Exprs: exprs, + Tag: tag, + }) + } } return nil } @@ -588,25 +585,12 @@ func (c *Compiler) compileSNAT(state *FirewallState) error { exprs = append(exprs, srcExprs...) } - if snat.Proto != "" { - exprs = append(exprs, matchProto(snat.Proto)...) - } - - for _, portStr := range snat.DPort { - pe, err := parsePortOrRange(portStr) - if err != nil { - return fmt.Errorf("snat[%d] dport: %w", i, err) - } - exprs = append(exprs, pe...) - } - - for _, portStr := range snat.SPort { - pe, err := parseSPortOrRange(portStr) - if err != nil { - return fmt.Errorf("snat[%d] sport: %w", i, err) - } - exprs = append(exprs, pe...) + matches, err := l4Matches(snat.Proto, snat.DPort, snat.SPort) + if err != nil { + return fmt.Errorf("snat[%d]: %w", i, err) } + head := exprs + exprs = nil if snat.Mark != "" { exprs = append(exprs, matchMark(snat.Mark)...) @@ -652,11 +636,13 @@ func (c *Compiler) compileSNAT(state *FirewallState) error { } } - state.Rules["postrouting"] = append(state.Rules["postrouting"], ManagedRule{ - Chain: "postrouting", - Exprs: exprs, - Tag: tag, - }) + for _, m := range matches { + state.Rules["postrouting"] = append(state.Rules["postrouting"], ManagedRule{ + Chain: "postrouting", + Exprs: append(append(append([]expr.Any{}, head...), m.exprs...), exprs...), + Tag: tag, + }) + } } return nil @@ -922,7 +908,7 @@ func (c *Compiler) expandZoneRef(ref string) []string { return []string{base} } -func (c *Compiler) buildMatchExprs(srcIface, dstIface, chain, proto string, dports, sports config.PortSpec, srcAddr, dstAddr string) ([]expr.Any, error) { +func (c *Compiler) buildMatchExprs(srcIface, dstIface, chain, proto string, dports, sports config.PortSpec, srcAddr, dstAddr string) ([]l4Match, error) { var exprs []expr.Any if srcIface != "" { @@ -948,34 +934,76 @@ func (c *Compiler) buildMatchExprs(srcIface, dstIface, chain, proto string, dpor exprs = append(exprs, dst...) } + matches, err := l4Matches(proto, dports, sports) + if err != nil { + return nil, err + } + for i := range matches { + matches[i].exprs = append(append([]expr.Any{}, exprs...), matches[i].exprs...) + } + return matches, nil +} + +type l4Match struct { + proto string + exprs []expr.Any +} + +// l4Matches yields one alternative per (proto, dport, sport): nft ANDs a rule's exprs, so lists need one rule each. +func l4Matches(proto string, dports, sports config.PortSpec) ([]l4Match, error) { + protos := []string{""} if proto != "" { - exprs = append(exprs, matchProto(proto)...) + protos = strings.Split(proto, ",") } - isICMP := strings.EqualFold(proto, "icmp") || strings.EqualFold(proto, "icmpv6") || strings.EqualFold(proto, "ipv6-icmp") - - for _, portStr := range dports { + var out []l4Match + for _, p := range protos { + p = strings.TrimSpace(p) + isICMP := strings.EqualFold(p, "icmp") || strings.EqualFold(p, "icmpv6") || strings.EqualFold(p, "ipv6-icmp") + parseD := parsePortOrRange if isICMP { - pe := matchICMPType(portStr) - exprs = append(exprs, pe...) - } else { - pe, err := parsePortOrRange(portStr) - if err != nil { - return nil, err - } - exprs = append(exprs, pe...) + parseD = func(s string) ([]expr.Any, error) { return matchICMPType(s), nil } } - } - - for _, portStr := range sports { - pe, err := parseSPortOrRange(portStr) + dalts, err := portAlternatives(dports, parseD) if err != nil { return nil, err } - exprs = append(exprs, pe...) + salts, err := portAlternatives(sports, parseSPortOrRange) + if err != nil { + return nil, err + } + for _, d := range dalts { + for _, sp := range salts { + var e []expr.Any + if p != "" { + e = append(e, matchProto(p)...) + } + e = append(append(e, d...), sp...) + out = append(out, l4Match{proto: p, exprs: e}) + } + } } + return out, nil +} - return exprs, nil +func portAlternatives(ports config.PortSpec, parse func(string) ([]expr.Any, error)) ([][]expr.Any, error) { + var alts [][]expr.Any + for _, item := range ports { + for _, s := range strings.Split(item, ",") { + if s = strings.TrimSpace(s); s == "" { + continue + } + e, err := parse(s) + if err != nil { + return nil, err + } + alts = append(alts, e) + } + } + if len(alts) == 0 { + return [][]expr.Any{nil}, nil + } + return alts, nil } // matchIfaceName matches an interface name, supporting wildcard "+" suffix. @@ -1218,8 +1246,8 @@ func matchConnLimit(spec string) []expr.Any { } func parsePortOrRange(s string) ([]expr.Any, error) { - if strings.Contains(s, "-") { - parts := strings.SplitN(s, "-", 2) + if i := strings.IndexAny(s, "-:"); i >= 0 { + parts := []string{s[:i], s[i+1:]} low, err := strconv.ParseUint(parts[0], 10, 16) if err != nil { return nil, fmt.Errorf("invalid port range low %q: %w", parts[0], err) @@ -1238,8 +1266,8 @@ func parsePortOrRange(s string) ([]expr.Any, error) { } func parseSPortOrRange(s string) ([]expr.Any, error) { - if strings.Contains(s, "-") { - parts := strings.SplitN(s, "-", 2) + if i := strings.IndexAny(s, "-:"); i >= 0 { + parts := []string{s[:i], s[i+1:]} low, err := strconv.ParseUint(parts[0], 10, 16) if err != nil { return nil, fmt.Errorf("invalid sport range low %q: %w", parts[0], err) diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index 84fa751..bcc1db7 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -1,6 +1,10 @@ package nftables import ( + "encoding/binary" + "fmt" + "reflect" + "strings" "testing" "github.com/google/nftables/expr" @@ -1531,3 +1535,93 @@ func TestCompile_PolicyRateLimit(t *testing.T) { } t.Error("policy:0 not found in input chain") } + +func TestCompile_PortAndProtoLists(t *testing.T) { + type want struct { + proto byte + dport string // "80" for an exact compare, "8000-8100" for a range + } + tests := []struct { + name string + rule config.Rule + chain string + want []want + }{ + { + name: "multi-port", + rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"80", "443"}}, + chain: "input", + want: []want{{6, "80"}, {6, "443"}}, + }, + { + name: "range in list", + rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"22", "8000:8100"}}, + chain: "input", + want: []want{{6, "22"}, {6, "8000-8100"}}, + }, + { + name: "comma string", + rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"80,443"}}, + chain: "input", + want: []want{{6, "80"}, {6, "443"}}, + }, + { + name: "tcp,udp", + rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp,udp", DPort: config.PortSpec{"53"}}, + chain: "input", + want: []want{{6, "53"}, {17, "53"}}, + }, + { + name: "dnat tcp,udp", + rule: config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "fw:192.0.2.10", Proto: "tcp,udp", DPort: config.PortSpec{"53"}}, + chain: "prerouting", + want: []want{{6, "53"}, {17, "53"}}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cfg := &config.Config{ + Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET}, + Zones: map[string]config.Zone{"fw": {Type: config.ZoneFirewall}, "net": {Type: config.ZoneIP}}, + Interfaces: []config.Interface{{Zone: "net", Interface: "eth0"}}, + Policy: []config.Policy{{Source: "all", Dest: "all", Action: config.PolicyDrop}}, + Rules: []config.Rule{tt.rule}, + PortGroups: make(map[string]config.PortGroup), + } + state, err := NewCompiler(cfg).Compile() + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + + var got []want + for _, r := range state.Rules[tt.chain] { + if r.Tag != "rule:0" { + continue + } + var w want + var portCmps []string + for i, e := range r.Exprs { + if m, ok := e.(*expr.Meta); ok && m.Key == expr.MetaKeyL4PROTO { + w.proto = r.Exprs[i+1].(*expr.Cmp).Data[0] + } + if p, ok := e.(*expr.Payload); ok && p.Base == expr.PayloadBaseTransportHeader && p.Offset == 2 { + for _, c := range r.Exprs[i+1:] { + cmp, ok := c.(*expr.Cmp) + if !ok { + break + } + portCmps = append(portCmps, fmt.Sprint(binary.BigEndian.Uint16(cmp.Data))) + } + } + } + w.dport = strings.Join(portCmps, "-") + got = append(got, w) + } + + if !reflect.DeepEqual(got, tt.want) { + t.Errorf("rules = %+v, want %+v", got, tt.want) + } + }) + } +} From 30758855ffaf964035682099a89e8c28f3eae52b Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Sat, 3 Oct 2026 20:45:16 +1000 Subject: [PATCH 02/11] Reject empty list elements, unknown ICMP types and limits on list rules --- internal/nftables/compiler.go | 29 +++++-- internal/nftables/compiler_test.go | 126 ++++++++++++++++++++++++++++- 2 files changed, 146 insertions(+), 9 deletions(-) diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 2af98ad..161f944 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -267,6 +267,16 @@ func (c *Compiler) compileRules(state *FirewallState) error { sport = rule.SPort } + if rule.RateLimit != "" || rule.ConnLimit != "" { + matches, err := l4Matches(proto, dports, sport) + if err != nil { + return fmt.Errorf("rule[%d]: %w", i, err) + } + if len(matches) > 1 { + return fmt.Errorf("rule[%d]: ratelimit/connlimit cannot be combined with proto or port lists (each expanded rule would get its own limiter)", i) + } + } + if err := c.compileOneRule(state, tag, rule.Source, rule.Dest, proto, dports, sport, rule.Action, rule.Log, rule.Dest, fwZone, rule.Section); err != nil { @@ -959,10 +969,13 @@ func l4Matches(proto string, dports, sports config.PortSpec) ([]l4Match, error) var out []l4Match for _, p := range protos { p = strings.TrimSpace(p) + if proto != "" && p == "" { + return nil, fmt.Errorf("empty element in proto list %q", proto) + } isICMP := strings.EqualFold(p, "icmp") || strings.EqualFold(p, "icmpv6") || strings.EqualFold(p, "ipv6-icmp") parseD := parsePortOrRange if isICMP { - parseD = func(s string) ([]expr.Any, error) { return matchICMPType(s), nil } + parseD = matchICMPType } dalts, err := portAlternatives(dports, parseD) if err != nil { @@ -991,7 +1004,7 @@ func portAlternatives(ports config.PortSpec, parse func(string) ([]expr.Any, err for _, item := range ports { for _, s := range strings.Split(item, ",") { if s = strings.TrimSpace(s); s == "" { - continue + return nil, fmt.Errorf("empty element in port list %q", item) } e, err := parse(s) if err != nil { @@ -1169,33 +1182,33 @@ var icmpTypeNames = map[string]byte{ "address-mask-reply": 18, } -func matchICMPType(spec string) []expr.Any { +func matchICMPType(spec string) ([]expr.Any, error) { if strings.Contains(spec, "/") { parts := strings.SplitN(spec, "/", 2) typeVal, ok := resolveICMPType(parts[0]) if !ok { - return nil + return nil, fmt.Errorf("invalid icmp type %q", parts[0]) } code, err := strconv.ParseUint(parts[1], 10, 8) if err != nil { - return nil + return nil, fmt.Errorf("invalid icmp code %q: %w", parts[1], err) } return []expr.Any{ &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 0, Len: 1}, &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{typeVal}}, &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 1, Len: 1}, &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{byte(code)}}, - } + }, nil } typeVal, ok := resolveICMPType(spec) if !ok { - return nil + return nil, fmt.Errorf("invalid icmp type %q", spec) } return []expr.Any{ &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 0, Len: 1}, &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{typeVal}}, - } + }, nil } func resolveICMPType(s string) (byte, bool) { diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index bcc1db7..50b1263 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -8,6 +8,7 @@ import ( "testing" "github.com/google/nftables/expr" + "golang.org/x/sys/unix" "git.unkin.net/unkin/tomswall/internal/config" ) @@ -1269,7 +1270,10 @@ func TestMatchICMPType(t *testing.T) { } for _, tt := range tests { - exprs := matchICMPType(tt.input) + exprs, err := matchICMPType(tt.input) + if err != nil { + t.Fatalf("matchICMPType(%q) error: %v", tt.input, err) + } if len(exprs) != tt.wantLen { t.Errorf("matchICMPType(%q) returned %d exprs, want %d", tt.input, len(exprs), tt.wantLen) } @@ -1625,3 +1629,123 @@ func TestCompile_PortAndProtoLists(t *testing.T) { }) } } + +func listCfg(mod func(*config.Config)) *config.Config { + cfg := &config.Config{ + Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET}, + Zones: map[string]config.Zone{"fw": {Type: config.ZoneFirewall}, "net": {Type: config.ZoneIP}}, + Interfaces: []config.Interface{{Zone: "net", Interface: "eth0"}}, + Policy: []config.Policy{{Source: "all", Dest: "all", Action: config.PolicyDrop}}, + PortGroups: make(map[string]config.PortGroup), + } + mod(cfg) + return cfg +} + +func taggedRules(state *FirewallState, chain, tag string) []ManagedRule { + var out []ManagedRule + for _, r := range state.Rules[chain] { + if r.Tag == tag { + out = append(out, r) + } + } + return out +} + +func TestCompile_ListExpansionCounts(t *testing.T) { + tests := []struct { + name string + mod func(*config.Config) + chain string + tag string + want int + }{ + {"proto x dport cross product", func(c *config.Config) { + c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp,udp", DPort: config.PortSpec{"80", "443"}}} + }, "input", "rule:0", 4}, + {"sport list", func(c *config.Config) { + c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", SPort: config.PortSpec{"1024,2048"}}} + }, "input", "rule:0", 2}, + {"snat proto x dport", func(c *config.Config) { + c.SNAT = []config.SNATRule{{Action: config.SNATMasquerade, Dest: "eth0", Proto: "tcp,udp", DPort: config.PortSpec{"80,443"}}} + }, "postrouting", "snat:0", 4}, + {"conntrack dport list", func(c *config.Config) { + c.Conntrack = []config.ConntrackRule{{Action: config.ConntrackNoTrack, Proto: "udp", DPort: config.PortSpec{"53", "123"}}} + }, "prerouting", "conntrack:0:prerouting", 2}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + state, err := NewCompiler(listCfg(tt.mod)).Compile() + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + if got := len(taggedRules(state, tt.chain, tt.tag)); got != tt.want { + t.Errorf("%s rules in %s = %d, want %d", tt.tag, tt.chain, got, tt.want) + } + }) + } +} + +func TestCompile_RejectPerProto(t *testing.T) { + state, err := NewCompiler(listCfg(func(c *config.Config) { + c.Rules = []config.Rule{{Action: config.RuleReject, Source: "net", Dest: "fw", Proto: "tcp,udp", DPort: config.PortSpec{"53"}}} + })).Compile() + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + rules := taggedRules(state, "input", "rule:0") + if len(rules) != 2 { + t.Fatalf("got %d rules, want 2", len(rules)) + } + want := []uint32{unix.NFT_REJECT_TCP_RST, unix.NFT_REJECT_ICMPX_UNREACH} + for i, r := range rules { + rej, ok := r.Exprs[len(r.Exprs)-1].(*expr.Reject) + if !ok { + t.Fatalf("rule %d: last expr %T, want *expr.Reject", i, r.Exprs[len(r.Exprs)-1]) + } + if rej.Type != want[i] { + t.Errorf("rule %d: reject type %d, want %d", i, rej.Type, want[i]) + } + } +} + +func TestCompile_ListErrors(t *testing.T) { + tests := []struct { + name string + rule config.Rule + }{ + {"invalid port in list", config.Rule{Proto: "tcp", DPort: config.PortSpec{"80,abc"}}}, + {"trailing empty proto", config.Rule{Proto: "tcp,", DPort: config.PortSpec{"80"}}}, + {"leading empty proto", config.Rule{Proto: ",udp", DPort: config.PortSpec{"80"}}}, + {"empty port element", config.Rule{Proto: "tcp", DPort: config.PortSpec{"80,"}}}, + {"unknown icmp type", config.Rule{Proto: "icmp", DPort: config.PortSpec{"bogus"}}}, + {"ratelimit with port list", config.Rule{Proto: "tcp", DPort: config.PortSpec{"80,443"}, RateLimit: "10/sec"}}, + {"connlimit with proto list", config.Rule{Proto: "tcp,udp", DPort: config.PortSpec{"53"}, ConnLimit: "10"}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + r := tt.rule + r.Action, r.Source, r.Dest = config.RuleAccept, "net", "fw" + _, err := NewCompiler(listCfg(func(c *config.Config) { c.Rules = []config.Rule{r} })).Compile() + if err == nil { + t.Fatal("Compile() succeeded, want error") + } + }) + } +} + +func TestCompile_ListExpansionDiffStable(t *testing.T) { + cfg := listCfg(func(c *config.Config) { + c.Rules = []config.Rule{{Action: config.RuleReject, Source: "net", Dest: "fw", Proto: "tcp,udp", DPort: config.PortSpec{"80,443"}}} + }) + state, err := NewCompiler(cfg).Compile() + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + if n := len(taggedRules(state, "input", "rule:0")); n != 4 { + t.Fatalf("got %d rule:0 rules, want 4", n) + } + if cs := computeDiff(state, state); !cs.Empty() { + t.Errorf("diff not empty:\n%s", cs.Summary()) + } +} From 211bacd5078fac1d090c6163c395edb1b5ff431a Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Sat, 3 Oct 2026 20:45:32 +1000 Subject: [PATCH 03/11] Expand comma zone lists in rule source and dest --- internal/config/blrules.go | 14 +-- internal/config/config_test.go | 33 +++++++ internal/config/rules.go | 38 +++++-- internal/nftables/compiler.go | 90 ++++++++++++++--- internal/nftables/compiler_test.go | 153 +++++++++++++++++++++++++++++ 5 files changed, 297 insertions(+), 31 deletions(-) diff --git a/internal/config/blrules.go b/internal/config/blrules.go index b1fc8e3..f361ea1 100644 --- a/internal/config/blrules.go +++ b/internal/config/blrules.go @@ -57,17 +57,19 @@ func (c *Config) validateBlrules() error { if r.Source != "all" && r.Source != "any" && r.Source != "none" && !hasPrefix(r.Source, "all!") && !hasPrefix(r.Source, "any!") { - srcZone := zoneFromSpec(r.Source) - if _, ok := c.Zones[srcZone]; !ok { - return fmt.Errorf("blrules[%d]: source zone %q not defined", i, srcZone) + for _, zs := range SplitZoneList(r.Source) { + if _, ok := c.Zones[zs.Zone]; !ok { + return fmt.Errorf("blrules[%d]: source zone %q not defined", i, zs.Zone) + } } } if r.Dest != "all" && r.Dest != "any" && r.Dest != "none" && !hasPrefix(r.Dest, "all!") && !hasPrefix(r.Dest, "any!") { - dstZone := zoneFromSpec(r.Dest) - if _, ok := c.Zones[dstZone]; !ok { - return fmt.Errorf("blrules[%d]: dest zone %q not defined", i, dstZone) + for _, zs := range SplitZoneList(r.Dest) { + if _, ok := c.Zones[zs.Zone]; !ok { + return fmt.Errorf("blrules[%d]: dest zone %q not defined", i, zs.Zone) + } } } } diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 4041c1b..012c997 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -3,6 +3,7 @@ package config import ( "os" "path/filepath" + "reflect" "strings" "testing" ) @@ -854,6 +855,19 @@ func TestValidateRules(t *testing.T) { }, wantErr: "source zone \"missing\" not defined", }, + { + name: "comma zone lists are valid", + rules: []Rule{ + {Action: RuleAccept, Source: "fw,loc", Dest: "loc,net:192.0.2.1,198.51.100.1"}, + }, + }, + { + name: "undefined zone in dest list", + rules: []Rule{ + {Action: RuleAccept, Source: "loc", Dest: "net,missing"}, + }, + wantErr: "dest zone \"missing\" not defined", + }, { name: "all keyword is valid source", rules: []Rule{ @@ -1006,3 +1020,22 @@ func TestValidateSNAT(t *testing.T) { }) } } + +func TestSplitZoneList(t *testing.T) { + tests := []struct { + in string + want []ZoneSpec + }{ + {"net", []ZoneSpec{{Zone: "net"}}}, + {"fw,lan,svr", []ZoneSpec{{Zone: "fw"}, {Zone: "lan"}, {Zone: "svr"}}}, + {"svr:192.0.2.17", []ZoneSpec{{Zone: "svr", Addr: "192.0.2.17"}}}, + {"net:192.0.2.1,198.51.100.1", []ZoneSpec{{Zone: "net", Addr: "192.0.2.1,198.51.100.1"}}}, + {"lan,svr:192.0.2.17", []ZoneSpec{{Zone: "lan"}, {Zone: "svr", Addr: "192.0.2.17"}}}, + {"net:2001:db8::1", []ZoneSpec{{Zone: "net", Addr: "2001:db8::1"}}}, + } + for _, tt := range tests { + if got := SplitZoneList(tt.in); !reflect.DeepEqual(got, tt.want) { + t.Errorf("SplitZoneList(%q) = %+v, want %+v", tt.in, got, tt.want) + } + } +} diff --git a/internal/config/rules.go b/internal/config/rules.go index 0a52040..1b36111 100644 --- a/internal/config/rules.go +++ b/internal/config/rules.go @@ -1,6 +1,9 @@ package config -import "fmt" +import ( + "fmt" + "strings" +) type RuleAction string @@ -172,10 +175,9 @@ func (c *Config) validateRules() error { if r.Source != "all" && r.Source != "any" && r.Source != "none" && !hasPrefix(r.Source, "all+") && !hasPrefix(r.Source, "all!") && !hasPrefix(r.Source, "any!") { - for _, srcPart := range splitZones(r.Source) { - srcZone := zoneFromSpec(srcPart) - if _, ok := c.Zones[srcZone]; !ok { - return fmt.Errorf("rule[%d]: source zone %q not defined", i, srcZone) + for _, zs := range SplitZoneList(r.Source) { + if _, ok := c.Zones[zs.Zone]; !ok { + return fmt.Errorf("rule[%d]: source zone %q not defined", i, zs.Zone) } } } @@ -183,10 +185,9 @@ func (c *Config) validateRules() error { if r.Action != RuleDNAT && r.Action != RuleRedirect && r.Action != RuleNoNAT { if r.Dest != "all" && r.Dest != "any" && r.Dest != "none" && !hasPrefix(r.Dest, "all+") && !hasPrefix(r.Dest, "all!") && !hasPrefix(r.Dest, "any!") { - for _, dstPart := range splitZones(r.Dest) { - dstZone := zoneFromSpec(dstPart) - if _, ok := c.Zones[dstZone]; !ok { - return fmt.Errorf("rule[%d]: dest zone %q not defined", i, dstZone) + for _, zs := range SplitZoneList(r.Dest) { + if _, ok := c.Zones[zs.Zone]; !ok { + return fmt.Errorf("rule[%d]: dest zone %q not defined", i, zs.Zone) } } } @@ -223,6 +224,25 @@ func (c *Config) validateRules() error { return nil } +// ZoneSpec is one zone of a SOURCE/DEST list; Addr is its comma-separated address list, if any. +type ZoneSpec struct{ Zone, Addr string } + +// SplitZoneList parses "lan,svr:a,b": commas before the first colon separate zones, +// commas after it separate addresses of the last zone (shorewall semantics). +func SplitZoneList(spec string) []ZoneSpec { + zones, addr, _ := strings.Cut(spec, ":") + var out []ZoneSpec + for _, z := range strings.Split(zones, ",") { + if z = strings.TrimSpace(z); z != "" { + out = append(out, ZoneSpec{Zone: z}) + } + } + if len(out) > 0 { + out[len(out)-1].Addr = addr + } + return out +} + // zoneFromSpec extracts the zone name from a zone spec like "net" or "net:192.168.1.0/24". func zoneFromSpec(spec string) string { for i, c := range spec { diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 2af98ad..e5a9737 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -3,6 +3,7 @@ package nftables import ( "encoding/binary" "fmt" + "log/slog" "net" "strconv" "strings" @@ -285,11 +286,12 @@ func (c *Compiler) compileRules(state *FirewallState) error { } func (c *Compiler) applyRuleExtras(state *FirewallState, tag string, rule config.Rule) { - fwZone := c.cfg.FirewallZone() - srcZone, _ := splitZoneSpec(rule.Source) - dstZone, _ := splitZoneSpec(rule.Dest) - chain := c.selectChain(srcZone, dstZone, fwZone) + for chain := range state.Rules { + c.applyChainExtras(state, chain, tag, rule) + } +} +func (c *Compiler) applyChainExtras(state *FirewallState, chain, tag string, rule config.Rule) { rules := state.Rules[chain] for idx := len(rules) - 1; idx >= 0; idx-- { if rules[idx].Tag != tag { @@ -345,13 +347,47 @@ func (c *Compiler) compileOneRule(state *FirewallState, tag, srcSpec, dstSpec, p dports, sports config.PortSpec, action config.RuleAction, logLevel string, dnatDest string, fwZone string, section config.RuleSection) error { - srcZone, srcAddr := splitZoneSpec(srcSpec) - dstZone, dstAddr := splitZoneSpec(dstSpec) - - if action == config.RuleDNAT || action == config.RuleRedirect { - return c.compileDNATRule(state, tag, srcSpec, dstSpec, proto, dports, sports, action, logLevel, fwZone) + for _, src := range zoneSpecs(srcSpec) { + for _, srcAddr := range splitAddrs(src.Addr) { + if action == config.RuleDNAT || action == config.RuleRedirect { + if err := c.compileDNATRule(state, tag, src.Zone, srcAddr, dstSpec, proto, dports, action, logLevel); err != nil { + return err + } + continue + } + for _, dst := range zoneSpecs(dstSpec) { + for _, dstAddr := range splitAddrs(dst.Addr) { + if err := c.compileZonePair(state, tag, src.Zone, srcAddr, dst.Zone, dstAddr, proto, + dports, sports, action, logLevel, fwZone, section); err != nil { + return err + } + } + } + } } + return nil +} +// zoneSpecs expands a comma zone list; "all"/"any" forms keep their own comma (exclusion) syntax. +func zoneSpecs(spec string) []config.ZoneSpec { + if strings.HasPrefix(spec, "all") || strings.HasPrefix(spec, "any") { + zone, addr := splitZoneSpec(spec) + return []config.ZoneSpec{{Zone: zone, Addr: addr}} + } + return config.SplitZoneList(spec) +} + +// splitAddrs yields one alternative per listed address; a negated list stays one AND-ed match. +func splitAddrs(addr string) []string { + if addr == "" || strings.HasPrefix(addr, "!") { + return []string{addr} + } + return strings.Split(addr, ",") +} + +func (c *Compiler) compileZonePair(state *FirewallState, tag, srcZone, srcAddr, dstZone, dstAddr, proto string, + dports, sports config.PortSpec, action config.RuleAction, logLevel string, + fwZone string, section config.RuleSection) error { srcIfaces := c.resolveZoneInterfaces(srcZone) dstIfaces := c.resolveZoneInterfaces(dstZone) chain := c.selectChain(srcZone, dstZone, fwZone) @@ -390,10 +426,8 @@ func (c *Compiler) compileOneRule(state *FirewallState, tag, srcSpec, dstSpec, p return nil } -func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcSpec, dstSpec, proto string, - dports, sports config.PortSpec, action config.RuleAction, logLevel, fwZone string) error { - - srcZone, srcAddr := splitZoneSpec(srcSpec) +func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr, dstSpec, proto string, + dports config.PortSpec, action config.RuleAction, logLevel string) error { chain := "prerouting" parts := strings.SplitN(dstSpec, ":", 3) @@ -874,10 +908,23 @@ func (c *Compiler) resolveZoneInterfaces(zone string) []string { return []string{""} } ifaces := c.cfg.ZoneInterfaces(zone) - if len(ifaces) == 0 { - return []string{""} + if len(ifaces) > 0 { + return ifaces } - return ifaces + if z, ok := c.cfg.Zones[zone]; ok && z.Type != config.ZoneFirewall && !c.zoneHasHosts(zone) { + slog.Warn("compiler: zone has no interfaces, skipping its rules", "zone", zone) + return nil + } + return []string{""} +} + +func (c *Compiler) zoneHasHosts(zone string) bool { + for _, h := range c.cfg.Hosts { + if h.Zone == zone { + return true + } + } + return false } func (c *Compiler) expandZoneRef(ref string) []string { @@ -1298,6 +1345,17 @@ func matchAddrCIDR(cidr string, isSrc bool) ([]expr.Any, error) { if strings.HasPrefix(cidr, "!") { negated = true cidr = cidr[1:] + if strings.Contains(cidr, ",") { + var all []expr.Any + for _, a := range strings.Split(cidr, ",") { + e, err := matchAddrCIDR("!"+a, isSrc) + if err != nil { + return nil, err + } + all = append(all, e...) + } + return all, nil + } } cmpOp := expr.CmpOpEq diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index bcc1db7..6ae3a27 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -1625,3 +1625,156 @@ func TestCompile_PortAndProtoLists(t *testing.T) { }) } } + +// describeRule renders a rule's iif/oif/saddr/daddr matches, e.g. "iif=eth1 oif=eth2 daddr=192.0.2.1". +func describeRule(r ManagedRule) string { + var parts []string + for i, e := range r.Exprs { + cmp, ok := func() (*expr.Cmp, bool) { + if i+1 >= len(r.Exprs) { + return nil, false + } + c, ok := r.Exprs[i+1].(*expr.Cmp) + return c, ok + }() + if !ok { + continue + } + switch m := e.(type) { + case *expr.Meta: + switch m.Key { + case expr.MetaKeyIIFNAME: + parts = append(parts, "iif="+strings.TrimRight(string(cmp.Data), "\x00")) + case expr.MetaKeyOIFNAME: + parts = append(parts, "oif="+strings.TrimRight(string(cmp.Data), "\x00")) + } + case *expr.Payload: + if m.Base == expr.PayloadBaseNetworkHeader && m.Len == 4 { + name := map[uint32]string{12: "saddr", 16: "daddr"}[m.Offset] + parts = append(parts, fmt.Sprintf("%s=%d.%d.%d.%d", name, cmp.Data[0], cmp.Data[1], cmp.Data[2], cmp.Data[3])) + } + } + } + return strings.Join(parts, " ") +} + +func TestCompile_CommaZoneLists(t *testing.T) { + tests := []struct { + name string + rule config.Rule + want map[string][]string + }{ + { + name: "fw in source list goes to output", + rule: config.Rule{Action: config.RuleAccept, Source: "fw,lan", Dest: "svr", Proto: "tcp", DPort: config.PortSpec{"22"}}, + want: map[string][]string{"output": {""}, "forward": {"iif=eth1 oif=eth2"}}, + }, + { + name: "dest list with fw splits input and forward", + rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "fw,svr,net"}, + want: map[string][]string{"input": {"iif=eth1"}, "forward": {"iif=eth1 oif=eth2", "iif=eth1 oif=eth0"}}, + }, + { + name: "zone without interfaces emits nothing", + rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "svr,dmz"}, + want: map[string][]string{"forward": {"iif=eth1 oif=eth2"}}, + }, + { + name: "address list after colon belongs to one zone", + rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "net:192.0.2.1,198.51.100.1"}, + want: map[string][]string{"forward": {"iif=eth1 oif=eth0 daddr=192.0.2.1", "iif=eth1 oif=eth0 daddr=198.51.100.1"}}, + }, + { + name: "zone:address inside a list", + rule: config.Rule{Action: config.RuleAccept, Source: "lan,svr:203.0.113.7", Dest: "fw"}, + want: map[string][]string{"input": {"iif=eth1", "iif=eth2 saddr=203.0.113.7"}}, + }, + { + name: "dnat source list", + rule: config.Rule{Action: config.RuleDNAT, Source: "net,lan", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, + want: map[string][]string{"prerouting": {"iif=eth0", "iif=eth1"}}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cfg := &config.Config{ + Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET}, + Zones: map[string]config.Zone{ + "fw": {Type: config.ZoneFirewall}, "net": {Type: config.ZoneIP}, + "lan": {Type: config.ZoneIP}, "svr": {Type: config.ZoneIP}, "dmz": {Type: config.ZoneIP}, + }, + Interfaces: []config.Interface{ + {Zone: "net", Interface: "eth0"}, {Zone: "lan", Interface: "eth1"}, {Zone: "svr", Interface: "eth2"}, + }, + Rules: []config.Rule{tt.rule}, + PortGroups: make(map[string]config.PortGroup), + } + state, err := NewCompiler(cfg).Compile() + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + got := map[string][]string{} + for chain, rules := range state.Rules { + for _, r := range rules { + if r.Tag == "rule:0" { + got[chain] = append(got[chain], describeRule(r)) + } + } + } + if !reflect.DeepEqual(got, tt.want) { + t.Errorf("rules = %v, want %v", got, tt.want) + } + }) + } +} + +func TestCompile_CommaZoneListExtras(t *testing.T) { + cfg := &config.Config{ + Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET}, + Zones: map[string]config.Zone{"fw": {Type: config.ZoneFirewall}, "net": {Type: config.ZoneIP}, "lan": {Type: config.ZoneIP}}, + Interfaces: []config.Interface{{Zone: "net", Interface: "eth0"}, {Zone: "lan", Interface: "eth1"}}, + Rules: []config.Rule{{Action: config.RuleAccept, Source: "net", Dest: "fw,lan", RateLimit: "10/sec:5"}}, + PortGroups: make(map[string]config.PortGroup), + } + state, err := NewCompiler(cfg).Compile() + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + for _, chain := range []string{"input", "forward"} { + found := false + for _, r := range state.Rules[chain] { + if r.Tag != "rule:0" { + continue + } + found = true + hasLimit := false + for _, e := range r.Exprs { + if _, ok := e.(*expr.Limit); ok { + hasLimit = true + } + } + if !hasLimit { + t.Errorf("%s rule missing Limit expression", chain) + } + } + if !found { + t.Errorf("rule:0 not found in %s chain", chain) + } + } +} + +func TestNegatedAddressList(t *testing.T) { + exprs, err := matchDestCIDR("!192.0.2.1,198.51.100.1") + if err != nil { + t.Fatalf("matchDestCIDR error: %v", err) + } + if len(exprs) != 4 { + t.Fatalf("expected 4 exprs, got %d", len(exprs)) + } + for _, i := range []int{1, 3} { + if exprs[i].(*expr.Cmp).Op != expr.CmpOpNeq { + t.Errorf("expr %d should be CmpOpNeq", i) + } + } +} From cc12c4a43a1d09f01251e8a6be11325f5020add7 Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Sat, 3 Oct 2026 20:45:57 +1000 Subject: [PATCH 04/11] Add try/confirm safe-apply and agent auto-revert --- cmd/tomswall/main.go | 2 + cmd/tomswall/try.go | 110 +++++++++++++++++++++++++++++++++ cmd/tomswall/try_test.go | 51 +++++++++++++++ internal/agent/agent.go | 64 ++++++++++++++----- internal/agent/agent_test.go | 75 +++++++++++++++++++++- internal/nftables/diff.go | 27 ++++++++ internal/nftables/diff_test.go | 65 +++++++++++++++++++ internal/nftables/engine.go | 65 ++++++++++++++----- 8 files changed, 423 insertions(+), 36 deletions(-) create mode 100644 cmd/tomswall/try.go create mode 100644 cmd/tomswall/try_test.go create mode 100644 internal/nftables/diff_test.go diff --git a/cmd/tomswall/main.go b/cmd/tomswall/main.go index c617769..c96516d 100644 --- a/cmd/tomswall/main.go +++ b/cmd/tomswall/main.go @@ -33,6 +33,8 @@ Use 'tomswall migrate' to convert a shorewall config to YAML.`, root.AddCommand( applyCmd(), + tryCmd(), + confirmCmd(), planCmd(), validateCmd(), statusCmd(), diff --git a/cmd/tomswall/try.go b/cmd/tomswall/try.go new file mode 100644 index 0000000..543d8d8 --- /dev/null +++ b/cmd/tomswall/try.go @@ -0,0 +1,110 @@ +package main + +import ( + "context" + "fmt" + "os" + "os/signal" + "path/filepath" + "strconv" + "strings" + "syscall" + "time" + + "github.com/spf13/cobra" + + "git.unkin.net/unkin/tomswall/internal/agent" +) + +var tryPIDFile = "/run/tomswall/try.pid" + +func tryCmd() *cobra.Command { + var timeout time.Duration + cmd := &cobra.Command{ + Use: "try", + Short: "Apply configuration and revert unless confirmed within a timeout", + Long: `Try snapshots the live tomswall table, applies the configuration, then waits +for 'tomswall confirm'. On timeout, interrupt or hangup the snapshot is restored +atomically. Confirm from a new session to prove new connections still work.`, + RunE: func(cmd *cobra.Command, args []string) error { + cfg, err := loadConfig() + if err != nil { + return err + } + + // Registered before apply so a confirm or hangup cannot be missed. + signal.Ignore(syscall.SIGPIPE) + confirm := make(chan os.Signal, 1) + signal.Notify(confirm, syscall.SIGUSR1) + abort := make(chan os.Signal, 1) + signal.Notify(abort, syscall.SIGHUP, syscall.SIGINT, syscall.SIGTERM) + + if err := os.MkdirAll(filepath.Dir(tryPIDFile), 0o755); err != nil { + return err + } + if err := os.WriteFile(tryPIDFile, []byte(strconv.Itoa(os.Getpid())), 0o644); err != nil { + return err + } + defer os.Remove(tryPIDFile) + + revert, err := agent.EngineApplier{}.Apply(context.Background(), cfg) + if err != nil { + return fmt.Errorf("applying changes: %w", err) + } + fmt.Printf("Applied. Run 'tomswall confirm' within %s or the previous ruleset is restored.\n", timeout) + if err := confirmOrRevert(confirm, abort, timeout, revert); err != nil { + return err + } + fmt.Println("Confirmed.") + return nil + }, + } + cmd.Flags().DurationVar(&timeout, "timeout", 60*time.Second, "time to wait for confirmation before reverting") + return cmd +} + +// confirmOrRevert waits for confirm; on abort or timeout it runs revert. +func confirmOrRevert(confirm, abort <-chan os.Signal, timeout time.Duration, revert func() error) error { + reason := "not confirmed within " + timeout.String() + select { + case <-confirm: + return nil + case s := <-abort: + reason = "interrupted by " + s.String() + case <-time.After(timeout): + } + if err := revert(); err != nil { + return fmt.Errorf("%s; revert failed: %w", reason, err) + } + return fmt.Errorf("%s: previous ruleset restored", reason) +} + +func confirmCmd() *cobra.Command { + return &cobra.Command{ + Use: "confirm", + Short: "Keep the configuration applied by a pending 'tomswall try'", + RunE: func(cmd *cobra.Command, args []string) error { + b, err := os.ReadFile(tryPIDFile) + if os.IsNotExist(err) { + return fmt.Errorf("no pending 'tomswall try'") + } + if err != nil { + return err + } + pid, err := strconv.Atoi(strings.TrimSpace(string(b))) + if err != nil { + return fmt.Errorf("parsing %s: %w", tryPIDFile, err) + } + // A stale pidfile must not signal an unrelated process. + comm, err := os.ReadFile(fmt.Sprintf("/proc/%d/comm", pid)) + if err != nil || strings.TrimSpace(string(comm)) != "tomswall" { + return fmt.Errorf("no pending 'tomswall try' (stale %s)", tryPIDFile) + } + if err := syscall.Kill(pid, syscall.SIGUSR1); err != nil { + return err + } + fmt.Println("Confirmed.") + return nil + }, + } +} diff --git a/cmd/tomswall/try_test.go b/cmd/tomswall/try_test.go new file mode 100644 index 0000000..9652eb5 --- /dev/null +++ b/cmd/tomswall/try_test.go @@ -0,0 +1,51 @@ +package main + +import ( + "errors" + "os" + "syscall" + "testing" + "time" +) + +func TestConfirmOrRevert(t *testing.T) { + tests := []struct { + name string + confirm bool + abort bool + revertErr error + wantRevert bool + wantErr bool + }{ + {name: "confirmed keeps ruleset", confirm: true}, + {name: "timeout reverts", wantRevert: true, wantErr: true}, + {name: "hangup reverts", abort: true, wantRevert: true, wantErr: true}, + {name: "revert failure surfaces", revertErr: errors.New("boom"), wantRevert: true, wantErr: true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + confirm := make(chan os.Signal, 1) + abort := make(chan os.Signal, 1) + if tt.confirm { + confirm <- syscall.SIGUSR1 + } + if tt.abort { + abort <- syscall.SIGHUP + } + reverted := false + err := confirmOrRevert(confirm, abort, 20*time.Millisecond, func() error { + reverted = true + return tt.revertErr + }) + if reverted != tt.wantRevert { + t.Errorf("reverted = %v, want %v", reverted, tt.wantRevert) + } + if (err != nil) != tt.wantErr { + t.Errorf("err = %v, wantErr %v", err, tt.wantErr) + } + if tt.revertErr != nil && !errors.Is(err, tt.revertErr) { + t.Errorf("revert error not wrapped: %v", err) + } + }) + } +} diff --git a/internal/agent/agent.go b/internal/agent/agent.go index bf7d68c..03d97d5 100644 --- a/internal/agent/agent.go +++ b/internal/agent/agent.go @@ -2,18 +2,21 @@ package agent import ( "context" + "errors" "fmt" "log/slog" + "net/url" "time" "git.unkin.net/unkin/tomswall/internal/config" "git.unkin.net/unkin/tomswall/internal/nftables" ) -// Applier applies a translated config to the firewall. Abstracted so the run -// loop is testable without touching the kernel. +// Applier applies a translated config to the firewall and returns a func that +// restores the previous ruleset. Abstracted so the run loop is testable without +// touching the kernel. type Applier interface { - Apply(ctx context.Context, cfg *config.Config) error + Apply(ctx context.Context, cfg *config.Config) (revert func() error, err error) } // Agent runs the pull-apply-report loop for one device. @@ -24,6 +27,10 @@ type Agent struct { Applier Applier // Resolver overrides the DNS resolver (tests); nil derives it per-config. Resolver *Resolver + + // reverted is the last generation rolled back for cutting off the API; it + // is not re-applied until a newer generation is published. + reverted int64 } // Run loops until ctx is cancelled, applying one cycle per Interval (and once @@ -61,16 +68,20 @@ func (a *Agent) RunOnce(ctx context.Context) error { return fmt.Errorf("control plane unreachable and no cached config: %w", err) } // Re-apply last known-good; do not report a generation we didn't fetch. - return a.applyConfig(ctx, cached, false) + return a.applyConfig(ctx, cached, nil) } - if err := a.Cache.Write(raw); err != nil { - slog.Warn("agent: caching config failed", "err", err) + if a.reverted != 0 && rc.Generation == a.reverted { + slog.Warn("agent: skipping reverted generation", "generation", rc.Generation) + return nil } - return a.applyConfig(ctx, rc, true) + return a.applyConfig(ctx, rc, raw) } -func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, report bool) error { +// applyConfig applies rc. A freshly fetched config (raw != nil) is verified by +// reporting status over a new connection through the new ruleset; if the API is +// unreachable the previous ruleset is restored. Only verified configs are cached. +func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte) error { resolver := a.Resolver if resolver == nil { resolver = NewResolver(rc.Resolver) @@ -81,15 +92,28 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, report bool if err != nil { return fmt.Errorf("translate: %w", err) } - if err := a.Applier.Apply(ctx, cfg); err != nil { + revert, err := a.Applier.Apply(ctx, cfg) + if err != nil { return fmt.Errorf("apply: %w", err) } slog.Info("agent: applied config", "generation", rc.Generation, "rules", len(cfg.Rules)) - if report { + if raw != nil { + a.Client.HTTP.CloseIdleConnections() if err := a.Client.ReportStatus(ctx, rc.Generation); err != nil { + var uerr *url.Error + if errors.As(err, &uerr) { + if rerr := revert(); rerr != nil { + return fmt.Errorf("API unreachable after applying generation %d (%v); revert failed: %w", rc.Generation, err, rerr) + } + a.reverted = rc.Generation + return fmt.Errorf("reverted generation %d: API unreachable through new ruleset: %w", rc.Generation, err) + } slog.Warn("agent: reporting status failed", "err", err) } + if err := a.Cache.Write(raw); err != nil { + slog.Warn("agent: caching config failed", "err", err) + } // Report the FIB so the control plane can scope router enforcement. if fib := CollectFIB(ctx); len(fib) > 0 { if err := a.Client.ReportRoutes(ctx, fib); err != nil { @@ -103,18 +127,26 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, report bool // EngineApplier applies via the real nftables differential engine. type EngineApplier struct{} -// Apply computes and applies the differential change set for cfg. -func (EngineApplier) Apply(_ context.Context, cfg *config.Config) error { +// Apply snapshots the live table, then computes and applies the differential +// change set for cfg. The returned func atomically restores the snapshot. +func (EngineApplier) Apply(_ context.Context, cfg *config.Config) (func() error, error) { engine, err := nftables.NewEngine(cfg) if err != nil { - return fmt.Errorf("initializing nftables: %w", err) + return nil, fmt.Errorf("initializing nftables: %w", err) } changes, err := engine.Plan() if err != nil { - return fmt.Errorf("computing changes: %w", err) + return nil, fmt.Errorf("computing changes: %w", err) } if changes.Empty() { - return nil + return func() error { return nil }, nil } - return engine.Apply(changes) + snap, err := engine.Snapshot() + if err != nil { + return nil, fmt.Errorf("snapshotting ruleset: %w", err) + } + if err := engine.Apply(changes); err != nil { + return nil, err + } + return func() error { return engine.Restore(snap) }, nil } diff --git a/internal/agent/agent_test.go b/internal/agent/agent_test.go index 3019815..cda54f7 100644 --- a/internal/agent/agent_test.go +++ b/internal/agent/agent_test.go @@ -126,16 +126,17 @@ func TestTranslateRejectsUnknownAction(t *testing.T) { } } -// fakeApplier records applied configs. +// fakeApplier records applied configs and reverts. type fakeApplier struct { count int32 + reverts int32 lastGen int } -func (f *fakeApplier) Apply(_ context.Context, cfg *config.Config) error { +func (f *fakeApplier) Apply(_ context.Context, cfg *config.Config) (func() error, error) { atomic.AddInt32(&f.count, 1) f.lastGen = len(cfg.Rules) - return nil + return func() error { atomic.AddInt32(&f.reverts, 1); return nil }, nil } const renderedYAML = `generation: 7 @@ -238,3 +239,71 @@ func TestRunOnceNoCacheReturnsError(t *testing.T) { t.Fatal("expected error when unreachable and no cache exists") } } + +// statusServer serves renderedYAML and answers status reports with handler. +func statusServer(t *testing.T, status http.HandlerFunc) *httptest.Server { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodGet { + _, _ = w.Write([]byte(renderedYAML)) + return + } + status(w, r) + })) + t.Cleanup(srv.Close) + return srv +} + +func TestRunOnceRevertsWhenAPIUnreachableAfterApply(t *testing.T) { + // Dropping the connection mimics a pushed rule that severs the API path. + srv := statusServer(t, func(w http.ResponseWriter, r *http.Request) { + conn, _, _ := w.(http.Hijacker).Hijack() + conn.Close() + }) + + applier := &fakeApplier{} + a := &Agent{ + Client: NewClient(srv.URL, "fw-a", "tok"), + Cache: Cache{Path: filepath.Join(t.TempDir(), "cache.yaml")}, + Applier: applier, + } + if err := a.RunOnce(context.Background()); err == nil { + t.Fatal("expected revert error") + } + if applier.reverts != 1 { + t.Errorf("expected 1 revert, got %d", applier.reverts) + } + if cached, _ := a.Cache.Read(); cached != nil { + t.Error("reverted config must not be cached as known-good") + } + + // The same generation is not re-applied on the next cycle. + if err := a.RunOnce(context.Background()); err != nil { + t.Fatalf("second RunOnce: %v", err) + } + if applier.count != 1 { + t.Errorf("reverted generation re-applied: %d applies", applier.count) + } +} + +func TestRunOnceKeepsConfigOnAPIErrorStatus(t *testing.T) { + // An HTTP error still proves the API is reachable — no revert. + srv := statusServer(t, func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + }) + + applier := &fakeApplier{} + a := &Agent{ + Client: NewClient(srv.URL, "fw-a", "tok"), + Cache: Cache{Path: filepath.Join(t.TempDir(), "cache.yaml")}, + Applier: applier, + } + if err := a.RunOnce(context.Background()); err != nil { + t.Fatalf("RunOnce: %v", err) + } + if applier.reverts != 0 { + t.Errorf("unexpected revert") + } + if cached, _ := a.Cache.Read(); cached == nil { + t.Error("expected config cached") + } +} diff --git a/internal/nftables/diff.go b/internal/nftables/diff.go index 82cfa30..faa9860 100644 --- a/internal/nftables/diff.go +++ b/internal/nftables/diff.go @@ -2,6 +2,7 @@ package nftables import ( "fmt" + "sort" "strings" "github.com/google/nftables/expr" @@ -84,6 +85,32 @@ func computeDiff(current, desired *FirewallState) *ChangeSet { return cs } +// restoreChangeSet replaces every managed rule in current with the snapshot's, +// in snapshot order, so a restore cannot reorder rules. +func restoreChangeSet(current, snap *FirewallState) *ChangeSet { + cs := &ChangeSet{} + for _, rules := range current.Rules { + for _, r := range rules { + if r.Tag != "" { + cs.Remove = append(cs.Remove, r) + } + } + } + chains := make([]string, 0, len(snap.Rules)) + for c := range snap.Rules { + chains = append(chains, c) + } + sort.Strings(chains) + for _, c := range chains { + for _, r := range snap.Rules[c] { + if r.Tag != "" { + cs.Add = append(cs.Add, r) + } + } + } + return cs +} + func rulesMatch(a, b []ManagedRule) bool { if len(a) != len(b) { return false diff --git a/internal/nftables/diff_test.go b/internal/nftables/diff_test.go new file mode 100644 index 0000000..33c9b2d --- /dev/null +++ b/internal/nftables/diff_test.go @@ -0,0 +1,65 @@ +package nftables + +import ( + "testing" + + "github.com/google/nftables/expr" +) + +func TestRestoreChangeSet(t *testing.T) { + accept := []expr.Any{&expr.Verdict{Kind: expr.VerdictAccept}} + drop := []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}} + + snap := &FirewallState{Rules: map[string][]ManagedRule{ + "input": { + {Chain: "input", Tag: "ssh", Exprs: accept, Handle: 4}, + {Chain: "input", Tag: "web", Exprs: accept, Handle: 5}, + {Chain: "input", Tag: "", Exprs: drop, Handle: 6}, + }, + "forward": {{Chain: "forward", Tag: "fwd", Exprs: accept, Handle: 7}}, + }} + current := &FirewallState{Rules: map[string][]ManagedRule{ + "input": { + {Chain: "input", Tag: "web", Exprs: accept, Handle: 10}, + {Chain: "input", Tag: "ssh", Exprs: drop, Handle: 11}, + {Chain: "input", Tag: "", Exprs: drop, Handle: 12}, + }, + }} + + cs := restoreChangeSet(current, snap) + + var removed []uint64 + for _, r := range cs.Remove { + removed = append(removed, r.Handle) + } + if len(removed) != 2 || removed[0] != 10 || removed[1] != 11 { + t.Errorf("expected managed handles [10 11] removed, untagged kept; got %v", removed) + } + + var added []string + for _, r := range cs.Add { + added = append(added, r.Tag) + } + want := []string{"fwd", "ssh", "web"} + if len(added) != len(want) { + t.Fatalf("added %v, want %v", added, want) + } + for i := range want { + if added[i] != want[i] { + t.Fatalf("added %v, want %v (snapshot order per chain)", added, want) + } + } + if !exprsEqual(cs.Add[1].Exprs, accept) { + t.Error("ssh not restored to its snapshot exprs") + } +} + +func TestRestoreChangeSetEmptySnapshotRemovesAll(t *testing.T) { + current := &FirewallState{Rules: map[string][]ManagedRule{ + "input": {{Chain: "input", Tag: "x", Handle: 1}}, + }} + cs := restoreChangeSet(current, &FirewallState{Rules: map[string][]ManagedRule{}}) + if len(cs.Remove) != 1 || len(cs.Add) != 0 { + t.Errorf("expected 1 remove 0 add, got %d/%d", len(cs.Remove), len(cs.Add)) + } +} diff --git a/internal/nftables/engine.go b/internal/nftables/engine.go index 0618a6f..5aba7c1 100644 --- a/internal/nftables/engine.go +++ b/internal/nftables/engine.go @@ -135,31 +135,32 @@ func (e *Engine) Flush() error { return nil } +func (e *Engine) findTable() (*nftables.Table, error) { + tables, err := e.conn.ListTables() + if err != nil { + return nil, fmt.Errorf("listing tables: %w", err) + } + for _, t := range tables { + if t.Name == e.cfg.Settings.TableName && t.Family == nftables.TableFamilyINet { + return t, nil + } + } + return nil, nil +} + func (e *Engine) readCurrentState() (*FirewallState, error) { state := &FirewallState{ Rules: make(map[string][]ManagedRule), } - tables, err := e.conn.ListTables() - if err != nil { - return state, nil - } - - var ourTable *nftables.Table - for _, t := range tables { - if t.Name == e.cfg.Settings.TableName && t.Family == nftables.TableFamilyINet { - ourTable = t - break - } - } - - if ourTable == nil { - return state, nil + ourTable, err := e.findTable() + if err != nil || ourTable == nil { + return state, err } chains, err := e.conn.ListChainsOfTableFamily(nftables.TableFamilyINet) if err != nil { - return state, nil + return nil, fmt.Errorf("listing chains: %w", err) } for _, chain := range chains { @@ -168,7 +169,7 @@ func (e *Engine) readCurrentState() (*FirewallState, error) { } rules, err := e.conn.GetRules(ourTable, chain) if err != nil { - continue + return nil, fmt.Errorf("listing rules of %s: %w", chain.Name, err) } for _, rule := range rules { state.Rules[chain.Name] = append(state.Rules[chain.Name], ManagedRule{ @@ -183,6 +184,36 @@ func (e *Engine) readCurrentState() (*FirewallState, error) { return state, nil } +// Snapshot is the tomswall table as captured live; a nil state means it was absent. +type Snapshot struct { + state *FirewallState +} + +// Snapshot captures the live tomswall table so Restore can roll back to it. +func (e *Engine) Snapshot() (*Snapshot, error) { + t, err := e.findTable() + if err != nil || t == nil { + return &Snapshot{}, err + } + state, err := e.readCurrentState() + if err != nil { + return nil, err + } + return &Snapshot{state: state}, nil +} + +// Restore atomically returns the tomswall table to the snapshot, rule order included. +func (e *Engine) Restore(s *Snapshot) error { + if s.state == nil { + return e.Flush() + } + current, err := e.readCurrentState() + if err != nil { + return err + } + return e.Apply(restoreChangeSet(current, s.state)) +} + func policyPtr(p nftables.ChainPolicy) *nftables.ChainPolicy { return &p } From 9c30f1fa544ec1edf713e79c21992e7ccabf9571 Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Sat, 3 Oct 2026 20:47:46 +1000 Subject: [PATCH 05/11] Reject unknown protocols and dports on mixed ICMP proto lists --- internal/nftables/compiler.go | 36 +++++++++----- internal/nftables/compiler_test.go | 78 ++++++++++++++++++++++++++---- 2 files changed, 92 insertions(+), 22 deletions(-) diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 161f944..2fc1e4d 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -428,6 +428,11 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcSpec, dstSpec, return err } + ip := net.ParseIP(dnatAddr) + if ip == nil { + return fmt.Errorf("invalid DNAT address %q", dnatAddr) + } + for _, srcIface := range srcIfaces { for _, m := range matches { var exprs []expr.Any @@ -450,11 +455,6 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcSpec, dstSpec, exprs = append(exprs, buildLog(logLevel, tag)...) } - ip := net.ParseIP(dnatAddr) - if ip == nil { - return fmt.Errorf("invalid DNAT address %q", dnatAddr) - } - if action == config.RuleRedirect { if dnatPort > 0 { portBytes := make([]byte, 2) @@ -975,8 +975,19 @@ func l4Matches(proto string, dports, sports config.PortSpec) ([]l4Match, error) isICMP := strings.EqualFold(p, "icmp") || strings.EqualFold(p, "icmpv6") || strings.EqualFold(p, "ipv6-icmp") parseD := parsePortOrRange if isICMP { + if len(protos) > 1 && len(dports) > 0 { + return nil, fmt.Errorf("dport %v is ambiguous with %s in proto list %q", dports, p, proto) + } parseD = matchICMPType } + var pm []expr.Any + if p != "" { + m, err := matchProto(p) + if err != nil { + return nil, err + } + pm = m + } dalts, err := portAlternatives(dports, parseD) if err != nil { return nil, err @@ -987,11 +998,7 @@ func l4Matches(proto string, dports, sports config.PortSpec) ([]l4Match, error) } for _, d := range dalts { for _, sp := range salts { - var e []expr.Any - if p != "" { - e = append(e, matchProto(p)...) - } - e = append(append(e, d...), sp...) + e := append(append(append([]expr.Any{}, pm...), d...), sp...) out = append(out, l4Match{proto: p, exprs: e}) } } @@ -1044,7 +1051,7 @@ func matchIfaceName(input bool, name string) []expr.Any { } } -func matchProto(proto string) []expr.Any { +func matchProto(proto string) ([]expr.Any, error) { var protoNum byte switch strings.ToLower(proto) { case "tcp": @@ -1064,13 +1071,16 @@ func matchProto(proto string) []expr.Any { case "sctp": protoNum = unix.IPPROTO_SCTP default: - n, _ := strconv.Atoi(proto) + n, err := strconv.ParseUint(proto, 10, 8) + if err != nil { + return nil, fmt.Errorf("unknown protocol %q", proto) + } protoNum = byte(n) } return []expr.Any{ &expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1}, &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{protoNum}}, - } + }, nil } func matchDPort(port uint16) []expr.Any { diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index 50b1263..939c4a9 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -1719,6 +1719,12 @@ func TestCompile_ListErrors(t *testing.T) { {"leading empty proto", config.Rule{Proto: ",udp", DPort: config.PortSpec{"80"}}}, {"empty port element", config.Rule{Proto: "tcp", DPort: config.PortSpec{"80,"}}}, {"unknown icmp type", config.Rule{Proto: "icmp", DPort: config.PortSpec{"bogus"}}}, + {"unknown icmp type with code", config.Rule{Proto: "icmp", DPort: config.PortSpec{"bogus/0"}}}, + {"invalid icmp code", config.Rule{Proto: "icmp", DPort: config.PortSpec{"destination-unreachable/x"}}}, + {"icmp code out of range", config.Rule{Proto: "icmp", DPort: config.PortSpec{"3/256"}}}, + {"unknown proto in list", config.Rule{Proto: "tcp,udpp", DPort: config.PortSpec{"53"}}}, + {"proto number out of range", config.Rule{Proto: "256"}}, + {"dport with icmp in proto list", config.Rule{Proto: "icmp,tcp", DPort: config.PortSpec{"80"}}}, {"ratelimit with port list", config.Rule{Proto: "tcp", DPort: config.PortSpec{"80,443"}, RateLimit: "10/sec"}}, {"connlimit with proto list", config.Rule{Proto: "tcp,udp", DPort: config.PortSpec{"53"}, ConnLimit: "10"}}, } @@ -1734,18 +1740,72 @@ func TestCompile_ListErrors(t *testing.T) { } } -func TestCompile_ListExpansionDiffStable(t *testing.T) { - cfg := listCfg(func(c *config.Config) { - c.Rules = []config.Rule{{Action: config.RuleReject, Source: "net", Dest: "fw", Proto: "tcp,udp", DPort: config.PortSpec{"80,443"}}} - }) - state, err := NewCompiler(cfg).Compile() +func TestCompile_ICMPList(t *testing.T) { + state, err := NewCompiler(listCfg(func(c *config.Config) { + c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "icmp", DPort: config.PortSpec{"echo-request,echo-reply", "destination-unreachable/4"}}} + })).Compile() if err != nil { t.Fatalf("Compile() error: %v", err) } - if n := len(taggedRules(state, "input", "rule:0")); n != 4 { - t.Fatalf("got %d rule:0 rules, want 4", n) + var got [][]byte + for _, r := range taggedRules(state, "input", "rule:0") { + var tc []byte + for i, e := range r.Exprs { + if p, ok := e.(*expr.Payload); ok && p.Base == expr.PayloadBaseTransportHeader && p.Len == 1 { + tc = append(tc, r.Exprs[i+1].(*expr.Cmp).Data[0]) + } + } + got = append(got, tc) } - if cs := computeDiff(state, state); !cs.Empty() { - t.Errorf("diff not empty:\n%s", cs.Summary()) + want := [][]byte{{8}, {0}, {3, 4}} + if !reflect.DeepEqual(got, want) { + t.Errorf("icmp type/code per rule = %v, want %v", got, want) + } +} + +func TestCompile_ColonRanges(t *testing.T) { + tests := []struct { + name string + mod func(*config.Config) + chain string + tag string + offset uint32 + }{ + {"rule sport", func(c *config.Config) { + c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", SPort: config.PortSpec{"1024:2048"}}} + }, "input", "rule:0", 0}, + {"snat sport", func(c *config.Config) { + c.SNAT = []config.SNATRule{{Action: config.SNATMasquerade, Dest: "eth0", Proto: "udp", SPort: config.PortSpec{"1024:2048"}}} + }, "postrouting", "snat:0", 0}, + {"snat dport", func(c *config.Config) { + c.SNAT = []config.SNATRule{{Action: config.SNATMasquerade, Dest: "eth0", Proto: "tcp", DPort: config.PortSpec{"1024:2048"}}} + }, "postrouting", "snat:0", 2}, + {"conntrack dport", func(c *config.Config) { + c.Conntrack = []config.ConntrackRule{{Action: config.ConntrackNoTrack, Proto: "udp", DPort: config.PortSpec{"1024:2048"}}} + }, "prerouting", "conntrack:0:prerouting", 2}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + state, err := NewCompiler(listCfg(tt.mod)).Compile() + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + rules := taggedRules(state, tt.chain, tt.tag) + if len(rules) != 1 { + t.Fatalf("got %d %s rules, want 1", len(rules), tt.tag) + } + var got []string + ex := rules[0].Exprs + for i, e := range ex { + if p, ok := e.(*expr.Payload); ok && p.Base == expr.PayloadBaseTransportHeader && p.Offset == tt.offset && p.Len == 2 && i+2 < len(ex) { + lo, hi := ex[i+1].(*expr.Cmp), ex[i+2].(*expr.Cmp) + got = append(got, fmt.Sprintf("%d>=%d,%d<=%d", lo.Op, binary.BigEndian.Uint16(lo.Data), hi.Op, binary.BigEndian.Uint16(hi.Data))) + } + } + want := []string{fmt.Sprintf("%d>=1024,%d<=2048", expr.CmpOpGte, expr.CmpOpLte)} + if !reflect.DeepEqual(got, want) { + t.Errorf("range match = %v, want %v", got, want) + } + }) } } From 852d6bf2ca38eb1e94fd3a42a8549c8478c2fb81 Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Sat, 3 Oct 2026 20:49:54 +1000 Subject: [PATCH 06/11] Resolve common IANA protocol names and reject ports on portless protocols --- internal/nftables/compiler.go | 71 ++++++++++++++---------------- internal/nftables/compiler_test.go | 21 +++++++++ 2 files changed, 54 insertions(+), 38 deletions(-) diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 2fc1e4d..9bd47eb 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -972,7 +972,19 @@ func l4Matches(proto string, dports, sports config.PortSpec) ([]l4Match, error) if proto != "" && p == "" { return nil, fmt.Errorf("empty element in proto list %q", proto) } - isICMP := strings.EqualFold(p, "icmp") || strings.EqualFold(p, "icmpv6") || strings.EqualFold(p, "ipv6-icmp") + var pm []expr.Any + isICMP := false + if p != "" { + n, err := protoNumber(p) + if err != nil { + return nil, err + } + isICMP = n == unix.IPPROTO_ICMP || n == unix.IPPROTO_ICMPV6 + if (len(dports) > 0 && !isICMP || len(sports) > 0) && !hasPorts(n) { + return nil, fmt.Errorf("protocol %q does not support ports", p) + } + pm = matchProtoNum(n) + } parseD := parsePortOrRange if isICMP { if len(protos) > 1 && len(dports) > 0 { @@ -980,14 +992,6 @@ func l4Matches(proto string, dports, sports config.PortSpec) ([]l4Match, error) } parseD = matchICMPType } - var pm []expr.Any - if p != "" { - m, err := matchProto(p) - if err != nil { - return nil, err - } - pm = m - } dalts, err := portAlternatives(dports, parseD) if err != nil { return nil, err @@ -1051,36 +1055,27 @@ func matchIfaceName(input bool, name string) []expr.Any { } } -func matchProto(proto string) ([]expr.Any, error) { - var protoNum byte - switch strings.ToLower(proto) { - case "tcp": - protoNum = unix.IPPROTO_TCP - case "udp": - protoNum = unix.IPPROTO_UDP - case "icmp": - protoNum = unix.IPPROTO_ICMP - case "icmpv6", "ipv6-icmp": - protoNum = unix.IPPROTO_ICMPV6 - case "gre": - protoNum = 47 - case "esp": - protoNum = 50 - case "ah": - protoNum = 51 - case "sctp": - protoNum = unix.IPPROTO_SCTP - default: - n, err := strconv.ParseUint(proto, 10, 8) - if err != nil { - return nil, fmt.Errorf("unknown protocol %q", proto) - } - protoNum = byte(n) +var protoNumbers = map[string]byte{ + "icmp": 1, "igmp": 2, "ipip": 4, "ipencap": 4, "tcp": 6, "udp": 17, + "gre": 47, "esp": 50, "ah": 51, "icmpv6": 58, "ipv6-icmp": 58, + "ospf": 89, "ospfigp": 89, "pim": 103, "vrrp": 112, "l2tp": 115, + "sctp": 132, "udplite": 136, +} + +// protoNumber resolves a common IANA protocol name or a 0-255 number. +func protoNumber(proto string) (byte, error) { + if n, ok := protoNumbers[strings.ToLower(proto)]; ok { + return n, nil } - return []expr.Any{ - &expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1}, - &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{protoNum}}, - }, nil + n, err := strconv.ParseUint(proto, 10, 8) + if err != nil { + return 0, fmt.Errorf("unknown protocol %q", proto) + } + return byte(n), nil +} + +func hasPorts(proto byte) bool { + return proto == unix.IPPROTO_TCP || proto == unix.IPPROTO_UDP || proto == unix.IPPROTO_SCTP || proto == unix.IPPROTO_UDPLITE } func matchDPort(port uint16) []expr.Any { diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index 939c4a9..912ce65 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -1581,6 +1581,24 @@ func TestCompile_PortAndProtoLists(t *testing.T) { chain: "prerouting", want: []want{{6, "53"}, {17, "53"}}, }, + { + name: "protocol names", + rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "ospf,OSPFIGP,igmp,gre,esp,ah,vrrp,pim,ipencap,ipv6-icmp"}, + chain: "input", + want: []want{{89, ""}, {89, ""}, {2, ""}, {47, ""}, {50, ""}, {51, ""}, {112, ""}, {103, ""}, {4, ""}, {58, ""}}, + }, + { + name: "protocol numbers", + rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "0,89,255"}, + chain: "input", + want: []want{{0, ""}, {89, ""}, {255, ""}}, + }, + { + name: "numeric tcp with port", + rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "6,udplite", DPort: config.PortSpec{"22"}}, + chain: "input", + want: []want{{6, "22"}, {136, "22"}}, + }, } for _, tt := range tests { @@ -1724,6 +1742,9 @@ func TestCompile_ListErrors(t *testing.T) { {"icmp code out of range", config.Rule{Proto: "icmp", DPort: config.PortSpec{"3/256"}}}, {"unknown proto in list", config.Rule{Proto: "tcp,udpp", DPort: config.PortSpec{"53"}}}, {"proto number out of range", config.Rule{Proto: "256"}}, + {"unknown proto name", config.Rule{Proto: "bogus"}}, + {"dport with ospf", config.Rule{Proto: "ospf", DPort: config.PortSpec{"80"}}}, + {"sport with gre", config.Rule{Proto: "gre", SPort: config.PortSpec{"80"}}}, {"dport with icmp in proto list", config.Rule{Proto: "icmp,tcp", DPort: config.PortSpec{"80"}}}, {"ratelimit with port list", config.Rule{Proto: "tcp", DPort: config.PortSpec{"80,443"}, RateLimit: "10/sec"}}, {"connlimit with proto list", config.Rule{Proto: "tcp,udp", DPort: config.PortSpec{"53"}, ConnLimit: "10"}}, From 8efed72c965a8668fab0ba111bb899c352786052 Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Sat, 3 Oct 2026 20:50:36 +1000 Subject: [PATCH 07/11] Match exact all/any zone tokens, skip only interface-less ip zones, keep DNAT free of rule extras, reject '!' inside address lists --- internal/config/blrules.go | 6 ++++ internal/config/config_test.go | 13 +++++++ internal/config/rules.go | 11 ++++++ internal/nftables/compiler.go | 22 ++++++++---- internal/nftables/compiler_test.go | 58 +++++++++++++++++++++++++++--- 5 files changed, 100 insertions(+), 10 deletions(-) diff --git a/internal/config/blrules.go b/internal/config/blrules.go index f361ea1..4118d22 100644 --- a/internal/config/blrules.go +++ b/internal/config/blrules.go @@ -61,6 +61,9 @@ func (c *Config) validateBlrules() error { if _, ok := c.Zones[zs.Zone]; !ok { return fmt.Errorf("blrules[%d]: source zone %q not defined", i, zs.Zone) } + if !validAddrList(zs.Addr) { + return fmt.Errorf("blrules[%d]: source %q: '!' may only prefix the whole address list", i, zs.Addr) + } } } @@ -70,6 +73,9 @@ func (c *Config) validateBlrules() error { if _, ok := c.Zones[zs.Zone]; !ok { return fmt.Errorf("blrules[%d]: dest zone %q not defined", i, zs.Zone) } + if !validAddrList(zs.Addr) { + return fmt.Errorf("blrules[%d]: dest %q: '!' may only prefix the whole address list", i, zs.Addr) + } } } } diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 012c997..7485bae 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -868,6 +868,19 @@ func TestValidateRules(t *testing.T) { }, wantErr: "dest zone \"missing\" not defined", }, + { + name: "negation prefixing the whole address list is valid", + rules: []Rule{ + {Action: RuleAccept, Source: "net:!192.0.2.1,198.51.100.1", Dest: "fw"}, + }, + }, + { + name: "negation inside an address list", + rules: []Rule{ + {Action: RuleAccept, Source: "net", Dest: "loc:192.0.2.1,!198.51.100.1"}, + }, + wantErr: "'!' may only prefix the whole address list", + }, { name: "all keyword is valid source", rules: []Rule{ diff --git a/internal/config/rules.go b/internal/config/rules.go index 1b36111..fb94e31 100644 --- a/internal/config/rules.go +++ b/internal/config/rules.go @@ -179,6 +179,9 @@ func (c *Config) validateRules() error { if _, ok := c.Zones[zs.Zone]; !ok { return fmt.Errorf("rule[%d]: source zone %q not defined", i, zs.Zone) } + if !validAddrList(zs.Addr) { + return fmt.Errorf("rule[%d]: source %q: '!' may only prefix the whole address list", i, zs.Addr) + } } } @@ -189,6 +192,9 @@ func (c *Config) validateRules() error { if _, ok := c.Zones[zs.Zone]; !ok { return fmt.Errorf("rule[%d]: dest zone %q not defined", i, zs.Zone) } + if !validAddrList(zs.Addr) { + return fmt.Errorf("rule[%d]: dest %q: '!' may only prefix the whole address list", i, zs.Addr) + } } } } @@ -243,6 +249,11 @@ func SplitZoneList(spec string) []ZoneSpec { return out } +// validAddrList reports whether '!' appears only at the start, negating the whole list. +func validAddrList(addr string) bool { + return !strings.Contains(strings.TrimPrefix(addr, "!"), "!") +} + // zoneFromSpec extracts the zone name from a zone spec like "net" or "net:192.168.1.0/24". func zoneFromSpec(spec string) string { for i, c := range spec { diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 0a20e5a..1d5c409 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -15,7 +15,8 @@ import ( ) type Compiler struct { - cfg *config.Config + cfg *config.Config + warned map[string]bool } func NewCompiler(cfg *config.Config) *Compiler { @@ -26,6 +27,7 @@ func (c *Compiler) Compile() (*FirewallState, error) { state := &FirewallState{ Rules: make(map[string][]ManagedRule), } + c.warned = nil c.compileLoopback(state) if err := c.compileConntrackFastPath(state); err != nil { @@ -297,7 +299,9 @@ func (c *Compiler) compileRules(state *FirewallState) error { func (c *Compiler) applyRuleExtras(state *FirewallState, tag string, rule config.Rule) { for chain := range state.Rules { - c.applyChainExtras(state, chain, tag, rule) + if chain != "prerouting" { + c.applyChainExtras(state, chain, tag, rule) + } } } @@ -394,8 +398,8 @@ func specCount(srcSpec, dstSpec string, action config.RuleAction) int { // zoneSpecs expands a comma zone list; "all"/"any" forms keep their own comma (exclusion) syntax. func zoneSpecs(spec string) []config.ZoneSpec { - if strings.HasPrefix(spec, "all") || strings.HasPrefix(spec, "any") { - zone, addr := splitZoneSpec(spec) + zone, addr := splitZoneSpec(spec) + if base, _, _ := strings.Cut(strings.TrimSuffix(zone, "+"), "!"); base == "all" || base == "any" { return []config.ZoneSpec{{Zone: zone, Addr: addr}} } return config.SplitZoneList(spec) @@ -935,8 +939,14 @@ func (c *Compiler) resolveZoneInterfaces(zone string) []string { if len(ifaces) > 0 { return ifaces } - if z, ok := c.cfg.Zones[zone]; ok && z.Type != config.ZoneFirewall && !c.zoneHasHosts(zone) { - slog.Warn("compiler: zone has no interfaces, skipping its rules", "zone", zone) + if z, ok := c.cfg.Zones[zone]; ok && z.Type == config.ZoneIP && !c.zoneHasHosts(zone) { + if !c.warned[zone] { + if c.warned == nil { + c.warned = map[string]bool{} + } + c.warned[zone] = true + slog.Warn("compiler: zone has no interfaces, skipping its rules", "zone", zone) + } return nil } return []string{""} diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index 8287e49..75c2eac 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -1655,6 +1655,9 @@ func describeRule(r ManagedRule) string { case *expr.Payload: if m.Base == expr.PayloadBaseNetworkHeader && m.Len == 4 { name := map[uint32]string{12: "saddr", 16: "daddr"}[m.Offset] + if cmp.Op == expr.CmpOpNeq { + name = "!" + name + } parts = append(parts, fmt.Sprintf("%s=%d.%d.%d.%d", name, cmp.Data[0], cmp.Data[1], cmp.Data[2], cmp.Data[3])) } } @@ -1664,9 +1667,10 @@ func describeRule(r ManagedRule) string { func TestCompile_CommaZoneLists(t *testing.T) { tests := []struct { - name string - rule config.Rule - want map[string][]string + name string + rule config.Rule + blrule *config.BlruleRule + want map[string][]string }{ { name: "fw in source list goes to output", @@ -1698,6 +1702,31 @@ func TestCompile_CommaZoneLists(t *testing.T) { rule: config.Rule{Action: config.RuleDNAT, Source: "net,lan", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, want: map[string][]string{"prerouting": {"iif=eth0", "iif=eth1"}}, }, + { + name: "dnat source address list", + rule: config.Rule{Action: config.RuleDNAT, Source: "net:192.0.2.5,198.51.100.5", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, + want: map[string][]string{"prerouting": {"iif=eth0 saddr=192.0.2.5", "iif=eth0 saddr=198.51.100.5"}}, + }, + { + name: "negated address list stays one AND-ed rule", + rule: config.Rule{Action: config.RuleAccept, Source: "net:!192.0.2.5,198.51.100.5", Dest: "fw"}, + want: map[string][]string{"input": {"iif=eth0 !saddr=192.0.2.5 !saddr=198.51.100.5"}}, + }, + { + name: "zone named like all/any keyword is a plain zone", + rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "anycast,net"}, + want: map[string][]string{"forward": {"iif=eth1 oif=eth3", "iif=eth1 oif=eth0"}}, + }, + { + name: "interface-less ipsec zone keeps zone-agnostic rule, ip zone skipped", + rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "vpn,dmz"}, + want: map[string][]string{"forward": {"iif=eth1"}}, + }, + { + name: "blrule zone list", + blrule: &config.BlruleRule{Action: config.BlruleDrop, Source: "net,anycast", Dest: "fw"}, + want: map[string][]string{"input": {"iif=eth0", "iif=eth3"}}, + }, } for _, tt := range tests { @@ -1707,13 +1736,19 @@ func TestCompile_CommaZoneLists(t *testing.T) { Zones: map[string]config.Zone{ "fw": {Type: config.ZoneFirewall}, "net": {Type: config.ZoneIP}, "lan": {Type: config.ZoneIP}, "svr": {Type: config.ZoneIP}, "dmz": {Type: config.ZoneIP}, + "anycast": {Type: config.ZoneIP}, "vpn": {Type: config.ZoneIPSec}, }, Interfaces: []config.Interface{ {Zone: "net", Interface: "eth0"}, {Zone: "lan", Interface: "eth1"}, {Zone: "svr", Interface: "eth2"}, + {Zone: "anycast", Interface: "eth3"}, }, Rules: []config.Rule{tt.rule}, PortGroups: make(map[string]config.PortGroup), } + tag := "rule:0" + if tt.blrule != nil { + cfg.Rules, cfg.Blrules, tag = nil, []config.BlruleRule{*tt.blrule}, "blrule:0" + } state, err := NewCompiler(cfg).Compile() if err != nil { t.Fatalf("Compile() error: %v", err) @@ -1721,7 +1756,7 @@ func TestCompile_CommaZoneLists(t *testing.T) { got := map[string][]string{} for chain, rules := range state.Rules { for _, r := range rules { - if r.Tag == "rule:0" { + if r.Tag == tag { got[chain] = append(got[chain], describeRule(r)) } } @@ -1789,6 +1824,21 @@ func TestCompile_ListExpansionCounts(t *testing.T) { } } +func TestCompile_DNATGetsNoRuleExtras(t *testing.T) { + compile := func(mark string) []expr.Any { + state, err := NewCompiler(listCfg(func(c *config.Config) { + c.Rules = []config.Rule{{Action: config.RuleDNAT, Source: "net", Dest: "fw:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}, Mark: mark}} + })).Compile() + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + return state.Rules["prerouting"][len(state.Rules["prerouting"])-1].Exprs + } + if plain, marked := compile(""), compile("0x1"); !reflect.DeepEqual(plain, marked) { + t.Errorf("DNAT prerouting rule changed by mark extra: %d exprs vs %d", len(plain), len(marked)) + } +} + func TestCompile_CommaZoneListLimitErrors(t *testing.T) { for _, r := range []config.Rule{ {Action: config.RuleAccept, Source: "net", Dest: "fw,lan", RateLimit: "10/sec:5"}, From 457056d58a7531fac1e4e93f8e533fde7311bba4 Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Sat, 3 Oct 2026 20:51:25 +1000 Subject: [PATCH 08/11] Pick reject type by resolved protocol number --- internal/nftables/compiler.go | 16 ++++++---- internal/nftables/compiler_test.go | 51 ++++++++++++++++++------------ 2 files changed, 40 insertions(+), 27 deletions(-) diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 9bd47eb..2994a80 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -955,7 +955,7 @@ func (c *Compiler) buildMatchExprs(srcIface, dstIface, chain, proto string, dpor } type l4Match struct { - proto string + proto byte exprs []expr.Any } @@ -973,9 +973,11 @@ func l4Matches(proto string, dports, sports config.PortSpec) ([]l4Match, error) return nil, fmt.Errorf("empty element in proto list %q", proto) } var pm []expr.Any + var n byte isICMP := false if p != "" { - n, err := protoNumber(p) + var err error + n, err = protoNumber(p) if err != nil { return nil, err } @@ -1003,7 +1005,7 @@ func l4Matches(proto string, dports, sports config.PortSpec) ([]l4Match, error) for _, d := range dalts { for _, sp := range salts { e := append(append(append([]expr.Any{}, pm...), d...), sp...) - out = append(out, l4Match{proto: p, exprs: e}) + out = append(out, l4Match{proto: n, exprs: e}) } } } @@ -1665,7 +1667,7 @@ func logLevelToNF(level string) expr.LogLevel { } } -func actionVerdict(action config.RuleAction, proto string, family config.AddressFamily) []expr.Any { +func actionVerdict(action config.RuleAction, proto byte, family config.AddressFamily) []expr.Any { switch action { case config.RuleAccept: return []expr.Any{&expr.Verdict{Kind: expr.VerdictAccept}} @@ -1690,8 +1692,8 @@ func actionVerdict(action config.RuleAction, proto string, family config.Address } } -func rejectExprs(proto string, family config.AddressFamily) []expr.Any { - if strings.ToLower(proto) == "tcp" { +func rejectExprs(proto byte, family config.AddressFamily) []expr.Any { + if proto == unix.IPPROTO_TCP { return []expr.Any{&expr.Reject{ Type: unix.NFT_REJECT_TCP_RST, Code: 0, @@ -1710,7 +1712,7 @@ func policyVerdict(action config.PolicyAction, family config.AddressFamily) []ex case config.PolicyDrop: return []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}} case config.PolicyReject: - return rejectExprs("", family) + return rejectExprs(0, family) default: return []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}} } diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index 912ce65..1a221b5 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -978,7 +978,7 @@ func TestNegatedAddress(t *testing.T) { } func TestRejectTCPRST(t *testing.T) { - exprs := rejectExprs("tcp", config.FamilyINET) + exprs := rejectExprs(unix.IPPROTO_TCP, config.FamilyINET) if len(exprs) != 1 { t.Fatalf("expected 1 expr, got %d", len(exprs)) } @@ -987,7 +987,7 @@ func TestRejectTCPRST(t *testing.T) { t.Errorf("TCP reject should use NFT_REJECT_TCP_RST (1), got %d", rej.Type) } - exprs = rejectExprs("udp", config.FamilyINET) + exprs = rejectExprs(unix.IPPROTO_UDP, config.FamilyINET) rej = exprs[0].(*expr.Reject) if rej.Type != 2 { t.Errorf("non-TCP reject should use NFT_REJECT_ICMPX_UNREACH (2), got %d", rej.Type) @@ -1705,25 +1705,36 @@ func TestCompile_ListExpansionCounts(t *testing.T) { } func TestCompile_RejectPerProto(t *testing.T) { - state, err := NewCompiler(listCfg(func(c *config.Config) { - c.Rules = []config.Rule{{Action: config.RuleReject, Source: "net", Dest: "fw", Proto: "tcp,udp", DPort: config.PortSpec{"53"}}} - })).Compile() - if err != nil { - t.Fatalf("Compile() error: %v", err) + tests := []struct { + proto string + want []uint32 + }{ + {"tcp,udp", []uint32{unix.NFT_REJECT_TCP_RST, unix.NFT_REJECT_ICMPX_UNREACH}}, + {"6", []uint32{unix.NFT_REJECT_TCP_RST}}, + {"6,17", []uint32{unix.NFT_REJECT_TCP_RST, unix.NFT_REJECT_ICMPX_UNREACH}}, } - rules := taggedRules(state, "input", "rule:0") - if len(rules) != 2 { - t.Fatalf("got %d rules, want 2", len(rules)) - } - want := []uint32{unix.NFT_REJECT_TCP_RST, unix.NFT_REJECT_ICMPX_UNREACH} - for i, r := range rules { - rej, ok := r.Exprs[len(r.Exprs)-1].(*expr.Reject) - if !ok { - t.Fatalf("rule %d: last expr %T, want *expr.Reject", i, r.Exprs[len(r.Exprs)-1]) - } - if rej.Type != want[i] { - t.Errorf("rule %d: reject type %d, want %d", i, rej.Type, want[i]) - } + for _, tt := range tests { + t.Run(tt.proto, func(t *testing.T) { + state, err := NewCompiler(listCfg(func(c *config.Config) { + c.Rules = []config.Rule{{Action: config.RuleReject, Source: "net", Dest: "fw", Proto: tt.proto, DPort: config.PortSpec{"53"}}} + })).Compile() + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + rules := taggedRules(state, "input", "rule:0") + if len(rules) != len(tt.want) { + t.Fatalf("got %d rules, want %d", len(rules), len(tt.want)) + } + for i, r := range rules { + rej, ok := r.Exprs[len(r.Exprs)-1].(*expr.Reject) + if !ok { + t.Fatalf("rule %d: last expr %T, want *expr.Reject", i, r.Exprs[len(r.Exprs)-1]) + } + if rej.Type != tt.want[i] { + t.Errorf("rule %d: reject type %d, want %d", i, rej.Type, tt.want[i]) + } + } + }) } } From 7f9c010e1af587268b651c2bf34171e06cdeaa5e Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Sat, 3 Oct 2026 20:52:42 +1000 Subject: [PATCH 09/11] Drop agent auto-revert; skip agent apply while a try is pending --- internal/agent/agent.go | 71 +++++++++++----------------------- internal/agent/agent_test.go | 75 ++---------------------------------- 2 files changed, 26 insertions(+), 120 deletions(-) diff --git a/internal/agent/agent.go b/internal/agent/agent.go index 03d97d5..868c4e0 100644 --- a/internal/agent/agent.go +++ b/internal/agent/agent.go @@ -2,21 +2,19 @@ package agent import ( "context" - "errors" "fmt" "log/slog" - "net/url" "time" "git.unkin.net/unkin/tomswall/internal/config" "git.unkin.net/unkin/tomswall/internal/nftables" + "git.unkin.net/unkin/tomswall/internal/tryapply" ) -// Applier applies a translated config to the firewall and returns a func that -// restores the previous ruleset. Abstracted so the run loop is testable without -// touching the kernel. +// Applier applies a translated config to the firewall. Abstracted so the run +// loop is testable without touching the kernel. type Applier interface { - Apply(ctx context.Context, cfg *config.Config) (revert func() error, err error) + Apply(ctx context.Context, cfg *config.Config) error } // Agent runs the pull-apply-report loop for one device. @@ -27,10 +25,6 @@ type Agent struct { Applier Applier // Resolver overrides the DNS resolver (tests); nil derives it per-config. Resolver *Resolver - - // reverted is the last generation rolled back for cutting off the API; it - // is not re-applied until a newer generation is published. - reverted int64 } // Run loops until ctx is cancelled, applying one cycle per Interval (and once @@ -68,20 +62,16 @@ func (a *Agent) RunOnce(ctx context.Context) error { return fmt.Errorf("control plane unreachable and no cached config: %w", err) } // Re-apply last known-good; do not report a generation we didn't fetch. - return a.applyConfig(ctx, cached, nil) + return a.applyConfig(ctx, cached, false) } - if a.reverted != 0 && rc.Generation == a.reverted { - slog.Warn("agent: skipping reverted generation", "generation", rc.Generation) - return nil + if err := a.Cache.Write(raw); err != nil { + slog.Warn("agent: caching config failed", "err", err) } - return a.applyConfig(ctx, rc, raw) + return a.applyConfig(ctx, rc, true) } -// applyConfig applies rc. A freshly fetched config (raw != nil) is verified by -// reporting status over a new connection through the new ruleset; if the API is -// unreachable the previous ruleset is restored. Only verified configs are cached. -func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte) error { +func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, report bool) error { resolver := a.Resolver if resolver == nil { resolver = NewResolver(rc.Resolver) @@ -92,28 +82,15 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte) if err != nil { return fmt.Errorf("translate: %w", err) } - revert, err := a.Applier.Apply(ctx, cfg) - if err != nil { + if err := a.Applier.Apply(ctx, cfg); err != nil { return fmt.Errorf("apply: %w", err) } slog.Info("agent: applied config", "generation", rc.Generation, "rules", len(cfg.Rules)) - if raw != nil { - a.Client.HTTP.CloseIdleConnections() + if report { if err := a.Client.ReportStatus(ctx, rc.Generation); err != nil { - var uerr *url.Error - if errors.As(err, &uerr) { - if rerr := revert(); rerr != nil { - return fmt.Errorf("API unreachable after applying generation %d (%v); revert failed: %w", rc.Generation, err, rerr) - } - a.reverted = rc.Generation - return fmt.Errorf("reverted generation %d: API unreachable through new ruleset: %w", rc.Generation, err) - } slog.Warn("agent: reporting status failed", "err", err) } - if err := a.Cache.Write(raw); err != nil { - slog.Warn("agent: caching config failed", "err", err) - } // Report the FIB so the control plane can scope router enforcement. if fib := CollectFIB(ctx); len(fib) > 0 { if err := a.Client.ReportRoutes(ctx, fib); err != nil { @@ -127,26 +104,24 @@ func (a *Agent) applyConfig(ctx context.Context, rc *RenderedConfig, raw []byte) // EngineApplier applies via the real nftables differential engine. type EngineApplier struct{} -// Apply snapshots the live table, then computes and applies the differential -// change set for cfg. The returned func atomically restores the snapshot. -func (EngineApplier) Apply(_ context.Context, cfg *config.Config) (func() error, error) { +// Apply computes and applies the differential change set for cfg. It refuses +// while a 'tomswall try' awaits confirmation. +func (EngineApplier) Apply(_ context.Context, cfg *config.Config) error { + unlock, err := tryapply.Acquire() + if err != nil { + return err + } + defer unlock() engine, err := nftables.NewEngine(cfg) if err != nil { - return nil, fmt.Errorf("initializing nftables: %w", err) + return fmt.Errorf("initializing nftables: %w", err) } changes, err := engine.Plan() if err != nil { - return nil, fmt.Errorf("computing changes: %w", err) + return fmt.Errorf("computing changes: %w", err) } if changes.Empty() { - return func() error { return nil }, nil + return nil } - snap, err := engine.Snapshot() - if err != nil { - return nil, fmt.Errorf("snapshotting ruleset: %w", err) - } - if err := engine.Apply(changes); err != nil { - return nil, err - } - return func() error { return engine.Restore(snap) }, nil + return engine.Apply(changes) } diff --git a/internal/agent/agent_test.go b/internal/agent/agent_test.go index cda54f7..3019815 100644 --- a/internal/agent/agent_test.go +++ b/internal/agent/agent_test.go @@ -126,17 +126,16 @@ func TestTranslateRejectsUnknownAction(t *testing.T) { } } -// fakeApplier records applied configs and reverts. +// fakeApplier records applied configs. type fakeApplier struct { count int32 - reverts int32 lastGen int } -func (f *fakeApplier) Apply(_ context.Context, cfg *config.Config) (func() error, error) { +func (f *fakeApplier) Apply(_ context.Context, cfg *config.Config) error { atomic.AddInt32(&f.count, 1) f.lastGen = len(cfg.Rules) - return func() error { atomic.AddInt32(&f.reverts, 1); return nil }, nil + return nil } const renderedYAML = `generation: 7 @@ -239,71 +238,3 @@ func TestRunOnceNoCacheReturnsError(t *testing.T) { t.Fatal("expected error when unreachable and no cache exists") } } - -// statusServer serves renderedYAML and answers status reports with handler. -func statusServer(t *testing.T, status http.HandlerFunc) *httptest.Server { - srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.Method == http.MethodGet { - _, _ = w.Write([]byte(renderedYAML)) - return - } - status(w, r) - })) - t.Cleanup(srv.Close) - return srv -} - -func TestRunOnceRevertsWhenAPIUnreachableAfterApply(t *testing.T) { - // Dropping the connection mimics a pushed rule that severs the API path. - srv := statusServer(t, func(w http.ResponseWriter, r *http.Request) { - conn, _, _ := w.(http.Hijacker).Hijack() - conn.Close() - }) - - applier := &fakeApplier{} - a := &Agent{ - Client: NewClient(srv.URL, "fw-a", "tok"), - Cache: Cache{Path: filepath.Join(t.TempDir(), "cache.yaml")}, - Applier: applier, - } - if err := a.RunOnce(context.Background()); err == nil { - t.Fatal("expected revert error") - } - if applier.reverts != 1 { - t.Errorf("expected 1 revert, got %d", applier.reverts) - } - if cached, _ := a.Cache.Read(); cached != nil { - t.Error("reverted config must not be cached as known-good") - } - - // The same generation is not re-applied on the next cycle. - if err := a.RunOnce(context.Background()); err != nil { - t.Fatalf("second RunOnce: %v", err) - } - if applier.count != 1 { - t.Errorf("reverted generation re-applied: %d applies", applier.count) - } -} - -func TestRunOnceKeepsConfigOnAPIErrorStatus(t *testing.T) { - // An HTTP error still proves the API is reachable — no revert. - srv := statusServer(t, func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusInternalServerError) - }) - - applier := &fakeApplier{} - a := &Agent{ - Client: NewClient(srv.URL, "fw-a", "tok"), - Cache: Cache{Path: filepath.Join(t.TempDir(), "cache.yaml")}, - Applier: applier, - } - if err := a.RunOnce(context.Background()); err != nil { - t.Fatalf("RunOnce: %v", err) - } - if applier.reverts != 0 { - t.Errorf("unexpected revert") - } - if cached, _ := a.Cache.Read(); cached == nil { - t.Error("expected config cached") - } -} From 6ac03e1012b3ce7fd0627e8c0e543cbcd5ea5b64 Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Sat, 3 Oct 2026 20:52:42 +1000 Subject: [PATCH 10/11] Persist try snapshot and arm a systemd revert timer --- cmd/tomswall/main.go | 1 + cmd/tomswall/try.go | 130 ++++++++++++++------ cmd/tomswall/try_test.go | 66 +++++----- go.mod | 2 +- internal/nftables/engine.go | 62 ++++++++-- internal/nftables/snapshot.go | 119 +++++++++++++++++++ internal/nftables/snapshot_test.go | 131 ++++++++++++++++++++ internal/tryapply/tryapply.go | 185 +++++++++++++++++++++++++++++ internal/tryapply/tryapply_test.go | 137 +++++++++++++++++++++ 9 files changed, 757 insertions(+), 76 deletions(-) create mode 100644 internal/nftables/snapshot.go create mode 100644 internal/nftables/snapshot_test.go create mode 100644 internal/tryapply/tryapply.go create mode 100644 internal/tryapply/tryapply_test.go diff --git a/cmd/tomswall/main.go b/cmd/tomswall/main.go index c96516d..cc4563a 100644 --- a/cmd/tomswall/main.go +++ b/cmd/tomswall/main.go @@ -35,6 +35,7 @@ Use 'tomswall migrate' to convert a shorewall config to YAML.`, applyCmd(), tryCmd(), confirmCmd(), + revertCmd(), planCmd(), validateCmd(), statusCmd(), diff --git a/cmd/tomswall/try.go b/cmd/tomswall/try.go index 543d8d8..0af4a92 100644 --- a/cmd/tomswall/try.go +++ b/cmd/tomswall/try.go @@ -1,31 +1,32 @@ package main import ( - "context" "fmt" "os" "os/signal" - "path/filepath" - "strconv" "strings" "syscall" "time" "github.com/spf13/cobra" - "git.unkin.net/unkin/tomswall/internal/agent" + "git.unkin.net/unkin/tomswall/internal/config" + "git.unkin.net/unkin/tomswall/internal/nftables" + "git.unkin.net/unkin/tomswall/internal/tryapply" ) -var tryPIDFile = "/run/tomswall/try.pid" +// revertGrace lets the in-process revert win before the systemd fallback fires. +const revertGrace = 30 * time.Second func tryCmd() *cobra.Command { var timeout time.Duration cmd := &cobra.Command{ Use: "try", Short: "Apply configuration and revert unless confirmed within a timeout", - Long: `Try snapshots the live tomswall table, applies the configuration, then waits -for 'tomswall confirm'. On timeout, interrupt or hangup the snapshot is restored -atomically. Confirm from a new session to prove new connections still work.`, + Long: `Try snapshots the live tomswall table to disk, applies the configuration, then +waits for 'tomswall confirm'. On timeout, interrupt or hangup the snapshot is +restored atomically. A transient systemd timer restores it too if this process +dies. Confirm from a new session to prove new connections still work.`, RunE: func(cmd *cobra.Command, args []string) error { cfg, err := loadConfig() if err != nil { @@ -39,23 +40,16 @@ atomically. Confirm from a new session to prove new connections still work.`, abort := make(chan os.Signal, 1) signal.Notify(abort, syscall.SIGHUP, syscall.SIGINT, syscall.SIGTERM) - if err := os.MkdirAll(filepath.Dir(tryPIDFile), 0o755); err != nil { + applied, err := tryApply(cfg, timeout+revertGrace) + if err != nil || !applied { return err } - if err := os.WriteFile(tryPIDFile, []byte(strconv.Itoa(os.Getpid())), 0o644); err != nil { - return err - } - defer os.Remove(tryPIDFile) - - revert, err := agent.EngineApplier{}.Apply(context.Background(), cfg) - if err != nil { - return fmt.Errorf("applying changes: %w", err) - } fmt.Printf("Applied. Run 'tomswall confirm' within %s or the previous ruleset is restored.\n", timeout) - if err := confirmOrRevert(confirm, abort, timeout, revert); err != nil { + msg, err := confirmOrRevert(confirm, abort, timeout, tryapply.Revert) + if err != nil { return err } - fmt.Println("Confirmed.") + fmt.Println(msg) return nil }, } @@ -63,20 +57,67 @@ atomically. Confirm from a new session to prove new connections still work.`, return cmd } +func tryApply(cfg *config.Config, fallback time.Duration) (bool, error) { + unlock, err := tryapply.Acquire() + if err != nil { + return false, err + } + defer unlock() + + engine, err := nftables.NewEngine(cfg) + if err != nil { + return false, fmt.Errorf("initializing nftables: %w", err) + } + changes, err := engine.Plan() + if err != nil { + return false, fmt.Errorf("computing changes: %w", err) + } + if changes.Empty() { + fmt.Println("No changes needed — firewall is up to date.") + return false, nil + } + fmt.Println(changes.Summary()) + + snap, err := engine.Snapshot() + if err != nil { + return false, fmt.Errorf("snapshotting ruleset: %w", err) + } + if err := tryapply.Arm(snap, os.Getpid(), fallback); err != nil { + return false, err + } + if err := engine.Apply(changes); err != nil { + if derr := tryapply.Discard(); derr != nil { + err = fmt.Errorf("%w (discarding snapshot: %v)", err, derr) + } + return false, fmt.Errorf("applying changes: %w", err) + } + return true, nil +} + // confirmOrRevert waits for confirm; on abort or timeout it runs revert. -func confirmOrRevert(confirm, abort <-chan os.Signal, timeout time.Duration, revert func() error) error { +// A confirm already delivered wins over a simultaneous abort or timeout. +func confirmOrRevert(confirm, abort <-chan os.Signal, timeout time.Duration, revert func() (bool, error)) (string, error) { reason := "not confirmed within " + timeout.String() select { case <-confirm: - return nil + return "Confirmed.", nil case s := <-abort: reason = "interrupted by " + s.String() case <-time.After(timeout): } - if err := revert(); err != nil { - return fmt.Errorf("%s; revert failed: %w", reason, err) + select { + case <-confirm: + return "Confirmed.", nil + default: } - return fmt.Errorf("%s: previous ruleset restored", reason) + reverted, err := revert() + if err != nil { + return "", fmt.Errorf("%s; revert failed: %w", reason, err) + } + if !reverted { + return "Already resolved by 'tomswall confirm' or 'tomswall revert'.", nil + } + return "", fmt.Errorf("%s: previous ruleset restored", reason) } func confirmCmd() *cobra.Command { @@ -84,27 +125,38 @@ func confirmCmd() *cobra.Command { Use: "confirm", Short: "Keep the configuration applied by a pending 'tomswall try'", RunE: func(cmd *cobra.Command, args []string) error { - b, err := os.ReadFile(tryPIDFile) - if os.IsNotExist(err) { - return fmt.Errorf("no pending 'tomswall try'") - } + pid, ok, err := tryapply.Confirm() if err != nil { return err } - pid, err := strconv.Atoi(strings.TrimSpace(string(b))) - if err != nil { - return fmt.Errorf("parsing %s: %w", tryPIDFile, err) + if !ok { + return fmt.Errorf("no pending 'tomswall try': it was already reverted or confirmed") } - // A stale pidfile must not signal an unrelated process. - comm, err := os.ReadFile(fmt.Sprintf("/proc/%d/comm", pid)) - if err != nil || strings.TrimSpace(string(comm)) != "tomswall" { - return fmt.Errorf("no pending 'tomswall try' (stale %s)", tryPIDFile) - } - if err := syscall.Kill(pid, syscall.SIGUSR1); err != nil { - return err + // A recycled PID must not be signalled. + if comm, err := os.ReadFile(fmt.Sprintf("/proc/%d/comm", pid)); err == nil && strings.TrimSpace(string(comm)) == "tomswall" { + _ = syscall.Kill(pid, syscall.SIGUSR1) } fmt.Println("Confirmed.") return nil }, } } + +func revertCmd() *cobra.Command { + return &cobra.Command{ + Use: "revert", + Short: "Restore the ruleset saved by a pending 'tomswall try'", + RunE: func(cmd *cobra.Command, args []string) error { + reverted, err := tryapply.Revert() + if err != nil { + return err + } + if !reverted { + fmt.Println("No pending 'tomswall try'.") + return nil + } + fmt.Println("Previous ruleset restored.") + return nil + }, + } +} diff --git a/cmd/tomswall/try_test.go b/cmd/tomswall/try_test.go index 9652eb5..a8f99f8 100644 --- a/cmd/tomswall/try_test.go +++ b/cmd/tomswall/try_test.go @@ -10,41 +10,53 @@ import ( func TestConfirmOrRevert(t *testing.T) { tests := []struct { - name string - confirm bool - abort bool - revertErr error - wantRevert bool - wantErr bool + name string + confirm bool + abort bool + resolved bool + revertErr error + wantRevert bool + wantErr bool + wantResolved bool }{ {name: "confirmed keeps ruleset", confirm: true}, + {name: "confirm wins over simultaneous abort", confirm: true, abort: true}, {name: "timeout reverts", wantRevert: true, wantErr: true}, {name: "hangup reverts", abort: true, wantRevert: true, wantErr: true}, {name: "revert failure surfaces", revertErr: errors.New("boom"), wantRevert: true, wantErr: true}, + {name: "resolved elsewhere", resolved: true, wantRevert: true, wantResolved: true}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - confirm := make(chan os.Signal, 1) - abort := make(chan os.Signal, 1) - if tt.confirm { - confirm <- syscall.SIGUSR1 - } - if tt.abort { - abort <- syscall.SIGHUP - } - reverted := false - err := confirmOrRevert(confirm, abort, 20*time.Millisecond, func() error { - reverted = true - return tt.revertErr - }) - if reverted != tt.wantRevert { - t.Errorf("reverted = %v, want %v", reverted, tt.wantRevert) - } - if (err != nil) != tt.wantErr { - t.Errorf("err = %v, wantErr %v", err, tt.wantErr) - } - if tt.revertErr != nil && !errors.Is(err, tt.revertErr) { - t.Errorf("revert error not wrapped: %v", err) + for i := 0; i < 20; i++ { // select between ready channels is random + confirm := make(chan os.Signal, 1) + abort := make(chan os.Signal, 1) + if tt.confirm { + confirm <- syscall.SIGUSR1 + } + if tt.abort { + abort <- syscall.SIGHUP + } + called := false + msg, err := confirmOrRevert(confirm, abort, 10*time.Millisecond, func() (bool, error) { + called = true + return !tt.resolved, tt.revertErr + }) + if called != tt.wantRevert { + t.Fatalf("revert called = %v, want %v", called, tt.wantRevert) + } + if (err != nil) != tt.wantErr { + t.Fatalf("err = %v, wantErr %v", err, tt.wantErr) + } + if tt.revertErr != nil && !errors.Is(err, tt.revertErr) { + t.Fatalf("revert error not wrapped: %v", err) + } + if tt.wantResolved && msg == "" { + t.Fatal("expected already-resolved message") + } + if !(tt.confirm && tt.abort) { + break + } } }) } diff --git a/go.mod b/go.mod index 5d8664e..0f68d47 100644 --- a/go.mod +++ b/go.mod @@ -4,6 +4,7 @@ go 1.23 require ( github.com/google/nftables v0.2.0 + github.com/mdlayher/netlink v1.7.2 github.com/spf13/cobra v1.8.1 golang.org/x/sys v0.18.0 gopkg.in/yaml.v3 v3.0.1 @@ -13,7 +14,6 @@ require ( github.com/google/go-cmp v0.6.0 // indirect github.com/inconshreveable/mousetrap v1.1.0 // indirect github.com/josharian/native v1.1.0 // indirect - github.com/mdlayher/netlink v1.7.2 // indirect github.com/mdlayher/socket v0.5.1 // indirect github.com/spf13/pflag v1.0.5 // indirect golang.org/x/net v0.23.0 // indirect diff --git a/internal/nftables/engine.go b/internal/nftables/engine.go index 5aba7c1..5311490 100644 --- a/internal/nftables/engine.go +++ b/internal/nftables/engine.go @@ -28,7 +28,8 @@ func (e *Engine) ensureTable() *nftables.Table { }) } -func (e *Engine) ensureChains(table *nftables.Table) map[string]*nftables.Chain { +// ensureChains declares the base chains; policies overrides their default policy. +func (e *Engine) ensureChains(table *nftables.Table, policies map[string]nftables.ChainPolicy) map[string]*nftables.Chain { chains := map[string]*nftables.Chain{ "input": { Name: "input", @@ -71,6 +72,9 @@ func (e *Engine) ensureChains(table *nftables.Table) map[string]*nftables.Chain } for name, chain := range chains { + if p, ok := policies[name]; ok { + chain.Policy = policyPtr(p) + } chains[name] = e.conn.AddChain(chain) } return chains @@ -93,8 +97,12 @@ func (e *Engine) Plan() (*ChangeSet, error) { } func (e *Engine) Apply(changes *ChangeSet) error { + return e.apply(changes, nil) +} + +func (e *Engine) apply(changes *ChangeSet, policies map[string]nftables.ChainPolicy) error { table := e.ensureTable() - chains := e.ensureChains(table) + chains := e.ensureChains(table, policies) for _, r := range changes.Remove { e.conn.DelRule(&nftables.Rule{ @@ -184,34 +192,70 @@ func (e *Engine) readCurrentState() (*FirewallState, error) { return state, nil } -// Snapshot is the tomswall table as captured live; a nil state means it was absent. +// Snapshot is the tomswall table as captured live, serialisable so a revert +// survives the process that took it. type Snapshot struct { - state *FirewallState + Table string `json:"table"` + Present bool `json:"present"` + Policies map[string]nftables.ChainPolicy `json:"policies,omitempty"` + Rules map[string][]SnapshotRule `json:"rules,omitempty"` +} + +// SnapshotRule is a managed rule with its expressions in netlink wire format. +type SnapshotRule struct { + Tag string `json:"tag"` + Exprs [][]byte `json:"exprs"` } // Snapshot captures the live tomswall table so Restore can roll back to it. func (e *Engine) Snapshot() (*Snapshot, error) { + snap := &Snapshot{Table: e.cfg.Settings.TableName} t, err := e.findTable() if err != nil || t == nil { - return &Snapshot{}, err + return snap, err } + snap.Present = true + + chains, err := e.conn.ListChainsOfTableFamily(nftables.TableFamilyINet) + if err != nil { + return nil, fmt.Errorf("listing chains: %w", err) + } + snap.Policies = make(map[string]nftables.ChainPolicy) + for _, c := range chains { + if c.Table.Name == snap.Table && c.Policy != nil { + snap.Policies[c.Name] = *c.Policy + } + } + state, err := e.readCurrentState() if err != nil { return nil, err } - return &Snapshot{state: state}, nil + snap.Rules, err = encodeState(state) + if err != nil { + return nil, err + } + return snap, nil } -// Restore atomically returns the tomswall table to the snapshot, rule order included. +// Restore atomically returns the tomswall table to the snapshot: rule order +// and chain policies included, or removed if it was absent. func (e *Engine) Restore(s *Snapshot) error { - if s.state == nil { + if s.Table != e.cfg.Settings.TableName { + return fmt.Errorf("snapshot is of table %q, engine manages %q", s.Table, e.cfg.Settings.TableName) + } + if !s.Present { return e.Flush() } + want, err := decodeState(s.Rules) + if err != nil { + return err + } current, err := e.readCurrentState() if err != nil { return err } - return e.Apply(restoreChangeSet(current, s.state)) + return e.apply(restoreChangeSet(current, want), s.Policies) } func policyPtr(p nftables.ChainPolicy) *nftables.ChainPolicy { diff --git a/internal/nftables/snapshot.go b/internal/nftables/snapshot.go new file mode 100644 index 0000000..bcb9b41 --- /dev/null +++ b/internal/nftables/snapshot.go @@ -0,0 +1,119 @@ +package nftables + +import ( + "encoding/binary" + "fmt" + + "github.com/google/nftables" + "github.com/google/nftables/expr" + "github.com/mdlayher/netlink" + "golang.org/x/sys/unix" +) + +const inet = byte(nftables.TableFamilyINet) + +// exprByName mirrors the expression types google/nftables can parse back from the kernel. +var exprByName = map[string]func() expr.Any{ + "ct": func() expr.Any { return &expr.Ct{} }, + "range": func() expr.Any { return &expr.Range{} }, + "meta": func() expr.Any { return &expr.Meta{} }, + "cmp": func() expr.Any { return &expr.Cmp{} }, + "counter": func() expr.Any { return &expr.Counter{} }, + "objref": func() expr.Any { return &expr.Objref{} }, + "payload": func() expr.Any { return &expr.Payload{} }, + "lookup": func() expr.Any { return &expr.Lookup{} }, + "immediate": func() expr.Any { return &expr.Immediate{} }, + "bitwise": func() expr.Any { return &expr.Bitwise{} }, + "redir": func() expr.Any { return &expr.Redir{} }, + "nat": func() expr.Any { return &expr.NAT{} }, + "limit": func() expr.Any { return &expr.Limit{} }, + "quota": func() expr.Any { return &expr.Quota{} }, + "dynset": func() expr.Any { return &expr.Dynset{} }, + "log": func() expr.Any { return &expr.Log{} }, + "exthdr": func() expr.Any { return &expr.Exthdr{} }, + "connlimit": func() expr.Any { return &expr.Connlimit{} }, + "queue": func() expr.Any { return &expr.Queue{} }, + "flow_offload": func() expr.Any { return &expr.FlowOffload{} }, + "reject": func() expr.Any { return &expr.Reject{} }, + "masq": func() expr.Any { return &expr.Masq{} }, + "hash": func() expr.Any { return &expr.Hash{} }, + "notrack": func() expr.Any { return &expr.Notrack{} }, +} + +func encodeState(state *FirewallState) (map[string][]SnapshotRule, error) { + out := make(map[string][]SnapshotRule, len(state.Rules)) + for chain, rules := range state.Rules { + for _, r := range rules { + sr := SnapshotRule{Tag: r.Tag} + for _, e := range r.Exprs { + b, err := expr.Marshal(inet, e) + if err != nil { + return nil, fmt.Errorf("encoding %s rule %q: %w", chain, r.Tag, err) + } + sr.Exprs = append(sr.Exprs, b) + } + out[chain] = append(out[chain], sr) + } + } + return out, nil +} + +func decodeState(rules map[string][]SnapshotRule) (*FirewallState, error) { + state := &FirewallState{Rules: make(map[string][]ManagedRule, len(rules))} + for chain, rs := range rules { + for _, sr := range rs { + r := ManagedRule{Chain: chain, Tag: sr.Tag} + for _, b := range sr.Exprs { + e, err := decodeExpr(b) + if err != nil { + return nil, fmt.Errorf("decoding %s rule %q: %w", chain, sr.Tag, err) + } + r.Exprs = append(r.Exprs, e) + } + state.Rules[chain] = append(state.Rules[chain], r) + } + } + return state, nil +} + +// decodeExpr reverses expr.Marshal, as google/nftables does when reading rules. +func decodeExpr(b []byte) (expr.Any, error) { + ad, err := netlink.NewAttributeDecoder(b) + if err != nil { + return nil, err + } + ad.ByteOrder = binary.BigEndian + var name string + var data []byte + for ad.Next() { + switch ad.Type() { + case unix.NFTA_EXPR_NAME: + name = ad.String() + case unix.NFTA_EXPR_DATA: + data = ad.Bytes() + } + } + if err := ad.Err(); err != nil { + return nil, err + } + newExpr, ok := exprByName[name] + if !ok { + return nil, fmt.Errorf("unsupported expression %q", name) + } + e := newExpr() + if name == "notrack" { + return e, nil + } + if err := expr.Unmarshal(inet, data, e); err != nil { + return nil, err + } + // A verdict is an immediate into the verdict register with no data. + if imm, ok := e.(*expr.Immediate); ok && imm.Register == unix.NFT_REG_VERDICT && len(imm.Data) == 0 { + v := &expr.Verdict{} + if err := expr.Unmarshal(inet, data, v); err != nil { + return nil, err + } + return v, nil + } + return e, nil +} diff --git a/internal/nftables/snapshot_test.go b/internal/nftables/snapshot_test.go new file mode 100644 index 0000000..8eab0a7 --- /dev/null +++ b/internal/nftables/snapshot_test.go @@ -0,0 +1,131 @@ +package nftables + +import ( + "encoding/json" + "reflect" + "testing" + + "github.com/google/nftables" + "github.com/google/nftables/expr" + "github.com/mdlayher/netlink" + "golang.org/x/sys/unix" + + "git.unkin.net/unkin/tomswall/internal/config" +) + +func TestSnapshotRulesRoundTrip(t *testing.T) { + exprs := []expr.Any{ + &expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1}, + &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_TCP}}, + &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2}, + &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{0, 22}}, + &expr.Ct{Register: 1, Key: expr.CtKeySTATE}, + &expr.Notrack{}, + &expr.Verdict{Kind: expr.VerdictAccept}, + } + state := &FirewallState{Rules: map[string][]ManagedRule{ + "input": {{Chain: "input", Tag: "ssh", Exprs: exprs}, {Chain: "input", Tag: "drop", Exprs: []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}}}, + }} + + rules, err := encodeState(state) + if err != nil { + t.Fatal(err) + } + b, err := json.Marshal(&Snapshot{Table: "tomswall", Present: true, Rules: rules, + Policies: map[string]nftables.ChainPolicy{"input": nftables.ChainPolicyAccept}}) + if err != nil { + t.Fatal(err) + } + var snap Snapshot + if err := json.Unmarshal(b, &snap); err != nil { + t.Fatal(err) + } + if snap.Policies["input"] != nftables.ChainPolicyAccept { + t.Errorf("policy lost: %v", snap.Policies) + } + got, err := decodeState(snap.Rules) + if err != nil { + t.Fatal(err) + } + in := got.Rules["input"] + if len(in) != 2 || in[0].Tag != "ssh" || in[1].Tag != "drop" { + t.Fatalf("rules/order lost: %+v", in) + } + if !reflect.DeepEqual(in[0].Exprs, exprs) { + t.Errorf("exprs changed:\n got %#v\nwant %#v", in[0].Exprs, exprs) + } + if _, ok := in[1].Exprs[0].(*expr.Verdict); !ok { + t.Errorf("verdict decoded as %T", in[1].Exprs[0]) + } +} + +func TestEnsureChainsPolicyOverride(t *testing.T) { + e := testEngine(t, nil) + chains := e.ensureChains(e.ensureTable(), map[string]nftables.ChainPolicy{"input": nftables.ChainPolicyAccept}) + if *chains["input"].Policy != nftables.ChainPolicyAccept { + t.Error("input policy not overridden") + } + if *chains["forward"].Policy != nftables.ChainPolicyDrop { + t.Error("forward policy should keep its default") + } +} + +func TestSnapshotAndRestoreAbsentTable(t *testing.T) { + tablePresent := false + var sent []netlink.HeaderType + e := testEngine(t, func(req []netlink.Message) ([]netlink.Message, error) { + for _, m := range req { + sent = append(sent, m.Header.Type) + if m.Header.Type == nftType(unix.NFT_MSG_GETTABLE) && tablePresent { + data := []byte{inet, 0, 0, 0} + attrs, _ := netlink.MarshalAttributes([]netlink.Attribute{{Type: unix.NFTA_TABLE_NAME, Data: []byte("tomswall\x00")}}) + return []netlink.Message{{Header: netlink.Header{Type: nftType(unix.NFT_MSG_NEWTABLE), Sequence: m.Header.Sequence}, Data: append(data, attrs...)}}, nil + } + } + return nil, nil + }) + + snap, err := e.Snapshot() + if err != nil { + t.Fatal(err) + } + if snap.Present || snap.Table != "tomswall" { + t.Fatalf("want absent tomswall snapshot, got %+v", snap) + } + + // The try created the table; restoring the absent snapshot deletes it. + tablePresent = true + sent = nil + if err := e.Restore(snap); err != nil { + t.Fatal(err) + } + deleted := false + for _, ht := range sent { + deleted = deleted || ht == nftType(unix.NFT_MSG_DELTABLE) + } + if !deleted { + t.Errorf("table not deleted; sent %v", sent) + } +} + +func TestRestoreRejectsOtherTable(t *testing.T) { + e := testEngine(t, nil) + if err := e.Restore(&Snapshot{Table: "other"}); err == nil { + t.Error("expected table mismatch error") + } +} + +func nftType(msg int) netlink.HeaderType { + return netlink.HeaderType(unix.NFNL_SUBSYS_NFTABLES<<8 | msg) +} + +func testEngine(t *testing.T, dial func([]netlink.Message) ([]netlink.Message, error)) *Engine { + if dial == nil { + dial = func([]netlink.Message) ([]netlink.Message, error) { return nil, nil } + } + conn, err := nftables.New(nftables.WithTestDial(dial)) + if err != nil { + t.Fatal(err) + } + return &Engine{cfg: &config.Config{Settings: config.Settings{TableName: "tomswall"}}, conn: conn} +} diff --git a/internal/tryapply/tryapply.go b/internal/tryapply/tryapply.go new file mode 100644 index 0000000..e8b5f30 --- /dev/null +++ b/internal/tryapply/tryapply.go @@ -0,0 +1,185 @@ +// Package tryapply keeps the state of a pending 'tomswall try' on disk so the +// revert survives the try process, backed by a transient systemd timer. +package tryapply + +import ( + "encoding/json" + "errors" + "fmt" + "os" + "os/exec" + "path/filepath" + "strings" + "syscall" + "time" + + "git.unkin.net/unkin/tomswall/internal/config" + "git.unkin.net/unkin/tomswall/internal/nftables" +) + +// Unit is the transient systemd unit that reverts an unconfirmed try. +const Unit = "tomswall-try-revert" + +var ( + // Dir holds the lock and the pending snapshot. + Dir = "/var/lib/tomswall" + // run executes a systemd command; replaced in tests. + run = func(name string, args ...string) error { + if out, err := exec.Command(name, args...).CombinedOutput(); err != nil { + return fmt.Errorf("%s %s: %w: %s", name, strings.Join(args, " "), err, strings.TrimSpace(string(out))) + } + return nil + } +) + +// ErrPending means a try awaits confirmation; nothing else may apply meanwhile. +var ErrPending = errors.New("a 'tomswall try' is pending; run 'tomswall confirm' or 'tomswall revert'") + +type pending struct { + PID int `json:"pid"` + Snapshot *nftables.Snapshot `json:"snapshot"` +} + +func snapshotPath() string { return filepath.Join(Dir, "try-snapshot.json") } + +// Acquire takes the exclusive try lock, failing with ErrPending while a try is unconfirmed. +func Acquire() (unlock func(), err error) { + unlock, err = lock() + if err != nil { + return nil, err + } + if _, err := os.Stat(snapshotPath()); err == nil { + unlock() + return nil, ErrPending + } + return unlock, nil +} + +func lock() (func(), error) { + if err := os.MkdirAll(Dir, 0o755); err != nil { + return nil, err + } + f, err := os.OpenFile(filepath.Join(Dir, "try.lock"), os.O_CREATE|os.O_RDWR, 0o600) + if err != nil { + return nil, err + } + if err := syscall.Flock(int(f.Fd()), syscall.LOCK_EX); err != nil { + f.Close() + return nil, fmt.Errorf("locking %s: %w", f.Name(), err) + } + return func() { f.Close() }, nil +} + +// Arm persists snap and schedules an out-of-process revert after delay. +// The caller must hold the lock from Acquire. +func Arm(snap *nftables.Snapshot, pid int, delay time.Duration) error { + b, err := json.Marshal(pending{PID: pid, Snapshot: snap}) + if err != nil { + return err + } + f, err := os.CreateTemp(Dir, ".try-snapshot-*") + if err != nil { + return err + } + defer os.Remove(f.Name()) + if _, err := f.Write(b); err != nil { + f.Close() + return err + } + if err := f.Sync(); err != nil { + f.Close() + return err + } + if err := f.Close(); err != nil { + return err + } + if err := os.Rename(f.Name(), snapshotPath()); err != nil { + return err + } + + exe, err := os.Executable() + if err != nil { + return discardWith(err) + } + _ = disarm() // a leftover timer from an earlier try would block the unit name + if err := run("systemd-run", "--quiet", "--collect", "--unit", Unit, + fmt.Sprintf("--on-active=%ds", int(delay.Round(time.Second).Seconds())), exe, "revert"); err != nil { + return discardWith(fmt.Errorf("arming revert timer: %w", err)) + } + return nil +} + +// Discard drops the pending snapshot and timer without restoring. The caller must hold the lock. +func Discard() error { + _ = disarm() + if err := os.Remove(snapshotPath()); err != nil && !os.IsNotExist(err) { + return err + } + return nil +} + +func discardWith(err error) error { + if derr := Discard(); derr != nil { + return fmt.Errorf("%w (discarding snapshot: %v)", err, derr) + } + return err +} + +func disarm() error { + return run("systemctl", "stop", Unit+".timer") +} + +// Confirm keeps the tried ruleset. ok is false when no try was pending, i.e. +// it was already reverted; pid is the waiting try process, if any. +func Confirm() (pid int, ok bool, err error) { + unlock, err := lock() + if err != nil { + return 0, false, err + } + defer unlock() + p, err := load() + if err != nil || p == nil { + return 0, false, err + } + return p.PID, true, Discard() +} + +// Revert restores the pending snapshot. reverted is false when nothing was +// pending (already confirmed or reverted). +func Revert() (reverted bool, err error) { + unlock, err := lock() + if err != nil { + return false, err + } + defer unlock() + p, err := load() + if err != nil || p == nil { + return false, err + } + engine, err := nftables.NewEngine(&config.Config{Settings: config.Settings{TableName: p.Snapshot.Table}}) + if err != nil { + return false, err + } + if err := engine.Restore(p.Snapshot); err != nil { + return false, fmt.Errorf("restoring snapshot: %w", err) + } + return true, Discard() +} + +func load() (*pending, error) { + b, err := os.ReadFile(snapshotPath()) + if os.IsNotExist(err) { + return nil, nil + } + if err != nil { + return nil, err + } + var p pending + if err := json.Unmarshal(b, &p); err != nil { + return nil, fmt.Errorf("parsing %s: %w", snapshotPath(), err) + } + if p.Snapshot == nil { + return nil, fmt.Errorf("%s has no snapshot", snapshotPath()) + } + return &p, nil +} diff --git a/internal/tryapply/tryapply_test.go b/internal/tryapply/tryapply_test.go new file mode 100644 index 0000000..ab72aa1 --- /dev/null +++ b/internal/tryapply/tryapply_test.go @@ -0,0 +1,137 @@ +package tryapply + +import ( + "errors" + "os" + "reflect" + "strings" + "testing" + "time" + + "git.unkin.net/unkin/tomswall/internal/nftables" +) + +func setup(t *testing.T) *[]string { + t.Helper() + Dir = t.TempDir() + var cmds []string + orig := run + run = func(name string, args ...string) error { + cmds = append(cmds, name+" "+strings.Join(args, " ")) + return nil + } + t.Cleanup(func() { run = orig }) + return &cmds +} + +func arm(t *testing.T, snap *nftables.Snapshot) { + t.Helper() + unlock, err := Acquire() + if err != nil { + t.Fatal(err) + } + defer unlock() + if err := Arm(snap, 4242, 90*time.Second); err != nil { + t.Fatal(err) + } +} + +func TestArmPersistsSnapshotAndTimer(t *testing.T) { + cmds := setup(t) + snap := &nftables.Snapshot{Table: "tomswall", Present: true, + Rules: map[string][]nftables.SnapshotRule{"input": {{Tag: "ssh", Exprs: [][]byte{{1, 2, 3}}}}}} + arm(t, snap) + + info, err := os.Stat(snapshotPath()) + if err != nil { + t.Fatal(err) + } + if info.Mode().Perm() != 0o600 { + t.Errorf("snapshot mode %v, want 0600", info.Mode().Perm()) + } + p, err := load() + if err != nil { + t.Fatal(err) + } + if p.PID != 4242 || !reflect.DeepEqual(p.Snapshot, snap) { + t.Errorf("round trip mismatch: %+v", p) + } + last := (*cmds)[len(*cmds)-1] + if !strings.HasPrefix(last, "systemd-run ") || !strings.Contains(last, "--unit "+Unit) || + !strings.Contains(last, "--on-active=90s") || !strings.HasSuffix(last, " revert") { + t.Errorf("unexpected arm command %q", last) + } +} + +func TestAbsentTableSnapshotRoundTrip(t *testing.T) { + setup(t) + arm(t, &nftables.Snapshot{Table: "tomswall"}) + p, err := load() + if err != nil { + t.Fatal(err) + } + if p.Snapshot.Present || p.Snapshot.Table != "tomswall" { + t.Errorf("absent table not preserved: %+v", p.Snapshot) + } +} + +func TestAcquireRefusesWhilePending(t *testing.T) { + setup(t) + arm(t, &nftables.Snapshot{Table: "tomswall"}) + if _, err := Acquire(); !errors.Is(err, ErrPending) { + t.Fatalf("second try: got %v, want ErrPending", err) + } +} + +func TestArmFailureDiscardsSnapshot(t *testing.T) { + setup(t) + run = func(name string, args ...string) error { + if name == "systemd-run" { + return errors.New("no systemd") + } + return nil + } + unlock, err := Acquire() + if err != nil { + t.Fatal(err) + } + defer unlock() + if err := Arm(&nftables.Snapshot{Table: "tomswall"}, 1, time.Minute); err == nil { + t.Fatal("expected arm error") + } + if _, err := os.Stat(snapshotPath()); !os.IsNotExist(err) { + t.Error("snapshot left behind without a revert timer") + } +} + +func TestConfirmPendingDisarms(t *testing.T) { + cmds := setup(t) + arm(t, &nftables.Snapshot{Table: "tomswall"}) + pid, ok, err := Confirm() + if err != nil || !ok || pid != 4242 { + t.Fatalf("Confirm = %d, %v, %v", pid, ok, err) + } + if _, err := os.Stat(snapshotPath()); !os.IsNotExist(err) { + t.Error("snapshot not removed") + } + if last := (*cmds)[len(*cmds)-1]; last != "systemctl stop "+Unit+".timer" { + t.Errorf("timer not stopped, last command %q", last) + } + unlock, err := Acquire() + if err != nil { + t.Fatalf("new try refused after confirm: %v", err) + } + unlock() +} + +func TestConfirmAfterRevertFails(t *testing.T) { + setup(t) + _, ok, err := Confirm() + if err != nil || ok { + t.Fatalf("Confirm with nothing pending = %v, %v; want not ok", ok, err) + } + reverted, err := Revert() + if err != nil || reverted { + t.Fatalf("Revert with nothing pending = %v, %v", reverted, err) + } +} From 0b92f1c2f3d9360f8a12c0b4fab0ae696d0c9353 Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Sat, 3 Oct 2026 20:56:52 +1000 Subject: [PATCH 11/11] Refuse mutating commands during a try and scope reverts to the try ID apply, flush and purge take the try lock and fail with ErrPending, which now names 'tomswall revert' as the recovery after a failed automatic revert. Each try gets an ID passed to the timer's 'revert --id', so a stale timer cannot revert a newer try. Adds tests for Revert success/failure/stale ID and for restoring a present table at the engine level. --- cmd/tomswall/guard_test.go | 30 ++++++++ cmd/tomswall/main.go | 21 ++++++ cmd/tomswall/try.go | 36 +++++---- internal/nftables/snapshot_test.go | 114 +++++++++++++++++++++++++++++ internal/tryapply/tryapply.go | 61 +++++++++------ internal/tryapply/tryapply_test.go | 90 +++++++++++++++++++++-- 6 files changed, 306 insertions(+), 46 deletions(-) create mode 100644 cmd/tomswall/guard_test.go diff --git a/cmd/tomswall/guard_test.go b/cmd/tomswall/guard_test.go new file mode 100644 index 0000000..12459a3 --- /dev/null +++ b/cmd/tomswall/guard_test.go @@ -0,0 +1,30 @@ +package main + +import ( + "errors" + "os" + "path/filepath" + "testing" + + "github.com/spf13/cobra" + + "git.unkin.net/unkin/tomswall/internal/tryapply" +) + +func TestMutatingCommandsRefuseWhileTryPending(t *testing.T) { + orig := tryapply.Dir + tryapply.Dir = t.TempDir() + t.Cleanup(func() { tryapply.Dir = orig }) + if err := os.WriteFile(filepath.Join(tryapply.Dir, "try-snapshot.json"), []byte("{}"), 0o600); err != nil { + t.Fatal(err) + } + configPath = "../../tomswall.example.yaml" + + for _, cmd := range []*cobra.Command{applyCmd(), flushCmd(), purgeCmd()} { + t.Run(cmd.Use, func(t *testing.T) { + if err := cmd.RunE(cmd, nil); !errors.Is(err, tryapply.ErrPending) { + t.Errorf("got %v, want ErrPending", err) + } + }) + } +} diff --git a/cmd/tomswall/main.go b/cmd/tomswall/main.go index cc4563a..bd886e1 100644 --- a/cmd/tomswall/main.go +++ b/cmd/tomswall/main.go @@ -12,6 +12,7 @@ import ( "git.unkin.net/unkin/tomswall/internal/config" "git.unkin.net/unkin/tomswall/internal/nftables" "git.unkin.net/unkin/tomswall/internal/shorewall" + "git.unkin.net/unkin/tomswall/internal/tryapply" ) var configPath string @@ -89,6 +90,12 @@ The firewall is never torn down — existing connections are preserved.`, return err } + unlock, err := tryapply.Acquire() + if err != nil { + return err + } + defer unlock() + engine, err := nftables.NewEngine(cfg) if err != nil { return fmt.Errorf("initializing nftables: %w", err) @@ -230,6 +237,14 @@ func purgeCmd() *cobra.Command { return err } + if !dryRun { + unlock, err := tryapply.Acquire() + if err != nil { + return err + } + defer unlock() + } + engine, err := nftables.NewEngine(cfg) if err != nil { return fmt.Errorf("initializing nftables: %w", err) @@ -275,6 +290,12 @@ func flushCmd() *cobra.Command { return err } + unlock, err := tryapply.Acquire() + if err != nil { + return err + } + defer unlock() + engine, err := nftables.NewEngine(cfg) if err != nil { return fmt.Errorf("initializing nftables: %w", err) diff --git a/cmd/tomswall/try.go b/cmd/tomswall/try.go index 0af4a92..ee9e21e 100644 --- a/cmd/tomswall/try.go +++ b/cmd/tomswall/try.go @@ -40,12 +40,12 @@ dies. Confirm from a new session to prove new connections still work.`, abort := make(chan os.Signal, 1) signal.Notify(abort, syscall.SIGHUP, syscall.SIGINT, syscall.SIGTERM) - applied, err := tryApply(cfg, timeout+revertGrace) - if err != nil || !applied { + id, err := tryApply(cfg, timeout+revertGrace) + if err != nil || id == "" { return err } fmt.Printf("Applied. Run 'tomswall confirm' within %s or the previous ruleset is restored.\n", timeout) - msg, err := confirmOrRevert(confirm, abort, timeout, tryapply.Revert) + msg, err := confirmOrRevert(confirm, abort, timeout, func() (bool, error) { return tryapply.Revert(id) }) if err != nil { return err } @@ -57,41 +57,44 @@ dies. Confirm from a new session to prove new connections still work.`, return cmd } -func tryApply(cfg *config.Config, fallback time.Duration) (bool, error) { +// tryApply applies cfg under a pending try and returns its ID, or "" when +// there was nothing to change. +func tryApply(cfg *config.Config, fallback time.Duration) (string, error) { unlock, err := tryapply.Acquire() if err != nil { - return false, err + return "", err } defer unlock() engine, err := nftables.NewEngine(cfg) if err != nil { - return false, fmt.Errorf("initializing nftables: %w", err) + return "", fmt.Errorf("initializing nftables: %w", err) } changes, err := engine.Plan() if err != nil { - return false, fmt.Errorf("computing changes: %w", err) + return "", fmt.Errorf("computing changes: %w", err) } if changes.Empty() { fmt.Println("No changes needed — firewall is up to date.") - return false, nil + return "", nil } fmt.Println(changes.Summary()) snap, err := engine.Snapshot() if err != nil { - return false, fmt.Errorf("snapshotting ruleset: %w", err) + return "", fmt.Errorf("snapshotting ruleset: %w", err) } - if err := tryapply.Arm(snap, os.Getpid(), fallback); err != nil { - return false, err + id, err := tryapply.Arm(snap, os.Getpid(), fallback) + if err != nil { + return "", err } if err := engine.Apply(changes); err != nil { if derr := tryapply.Discard(); derr != nil { err = fmt.Errorf("%w (discarding snapshot: %v)", err, derr) } - return false, fmt.Errorf("applying changes: %w", err) + return "", fmt.Errorf("applying changes: %w", err) } - return true, nil + return id, nil } // confirmOrRevert waits for confirm; on abort or timeout it runs revert. @@ -143,11 +146,12 @@ func confirmCmd() *cobra.Command { } func revertCmd() *cobra.Command { - return &cobra.Command{ + var id string + cmd := &cobra.Command{ Use: "revert", Short: "Restore the ruleset saved by a pending 'tomswall try'", RunE: func(cmd *cobra.Command, args []string) error { - reverted, err := tryapply.Revert() + reverted, err := tryapply.Revert(id) if err != nil { return err } @@ -159,4 +163,6 @@ func revertCmd() *cobra.Command { return nil }, } + cmd.Flags().StringVar(&id, "id", "", "only revert the try with this ID (used by the revert timer)") + return cmd } diff --git a/internal/nftables/snapshot_test.go b/internal/nftables/snapshot_test.go index 8eab0a7..3098119 100644 --- a/internal/nftables/snapshot_test.go +++ b/internal/nftables/snapshot_test.go @@ -1,6 +1,8 @@ package nftables import ( + "bytes" + "encoding/binary" "encoding/json" "reflect" "testing" @@ -108,6 +110,118 @@ func TestSnapshotAndRestoreAbsentTable(t *testing.T) { } } +func TestRestorePresentTable(t *testing.T) { + want := []SnapshotRule{} + for _, r := range []ManagedRule{ + {Tag: "ssh", Exprs: []expr.Any{&expr.Ct{Register: 1, Key: expr.CtKeySTATE}, &expr.Verdict{Kind: expr.VerdictAccept}}}, + {Tag: "drop", Exprs: []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}}, + } { + enc, err := encodeState(&FirewallState{Rules: map[string][]ManagedRule{"input": {r}}}) + if err != nil { + t.Fatal(err) + } + want = append(want, enc["input"]...) + } + snap := &Snapshot{Table: "tomswall", Present: true, + Policies: map[string]nftables.ChainPolicy{"input": nftables.ChainPolicyAccept}, + Rules: map[string][]SnapshotRule{"input": want}} + + // Live state: the tried ruleset left one managed rule (handle 7) in input. + attrs := func(a ...netlink.Attribute) []byte { + b, err := netlink.MarshalAttributes(a) + if err != nil { + t.Fatal(err) + } + return append([]byte{inet, 0, 0, 0}, b...) + } + handle := make([]byte, 8) + binary.BigEndian.PutUint64(handle, 7) + var batch []netlink.Message + e := testEngine(t, func(req []netlink.Message) ([]netlink.Message, error) { + if len(req) == 0 { + return nil, nil + } + reply := func(msg int, data []byte) ([]netlink.Message, error) { + return []netlink.Message{{Header: netlink.Header{Type: nftType(msg), Sequence: req[0].Header.Sequence}, Data: data}}, nil + } + switch req[0].Header.Type { + case nftType(unix.NFT_MSG_GETTABLE): + return reply(unix.NFT_MSG_NEWTABLE, attrs(netlink.Attribute{Type: unix.NFTA_TABLE_NAME, Data: []byte("tomswall\x00")})) + case nftType(unix.NFT_MSG_GETCHAIN): + return reply(unix.NFT_MSG_NEWCHAIN, attrs( + netlink.Attribute{Type: unix.NFTA_CHAIN_TABLE, Data: []byte("tomswall\x00")}, + netlink.Attribute{Type: unix.NFTA_CHAIN_NAME, Data: []byte("input\x00")})) + case nftType(unix.NFT_MSG_GETRULE): + return reply(unix.NFT_MSG_NEWRULE, attrs( + netlink.Attribute{Type: unix.NFTA_RULE_TABLE, Data: []byte("tomswall\x00")}, + netlink.Attribute{Type: unix.NFTA_RULE_CHAIN, Data: []byte("input\x00")}, + netlink.Attribute{Type: unix.NFTA_RULE_HANDLE, Data: handle}, + netlink.Attribute{Type: unix.NFTA_RULE_USERDATA, Data: []byte("tried")})) + } + batch = append(batch, req...) + return nil, nil + }) + + if err := e.Restore(snap); err != nil { + t.Fatal(err) + } + + var deleted []uint64 + var added []SnapshotRule + policy := map[string]uint32{} + for _, m := range batch { + ad, err := netlink.NewAttributeDecoder(m.Data[4:]) + if err != nil { + t.Fatal(err) + } + ad.ByteOrder = binary.BigEndian + var name string + var r SnapshotRule + var h uint64 + var pol *uint32 + for ad.Next() { + switch { + case m.Header.Type == nftType(unix.NFT_MSG_NEWCHAIN) && ad.Type() == unix.NFTA_CHAIN_NAME: + name = ad.String() + case m.Header.Type == nftType(unix.NFT_MSG_NEWCHAIN) && ad.Type() == unix.NFTA_CHAIN_POLICY: + v := ad.Uint32() + pol = &v + case m.Header.Type == nftType(unix.NFT_MSG_DELRULE) && ad.Type() == unix.NFTA_RULE_HANDLE: + h = ad.Uint64() + case m.Header.Type == nftType(unix.NFT_MSG_NEWRULE) && ad.Type() == unix.NFTA_RULE_USERDATA: + r.Tag = string(ad.Bytes()) + case m.Header.Type == nftType(unix.NFT_MSG_NEWRULE) && ad.Type() == unix.NFTA_RULE_EXPRESSIONS: + ad.Nested(func(nad *netlink.AttributeDecoder) error { + for nad.Next() { + r.Exprs = append(r.Exprs, bytes.Clone(nad.Bytes())) + } + return nil + }) + } + } + switch m.Header.Type { + case nftType(unix.NFT_MSG_NEWCHAIN): + if pol != nil { + policy[name] = *pol + } + case nftType(unix.NFT_MSG_DELRULE): + deleted = append(deleted, h) + case nftType(unix.NFT_MSG_NEWRULE): + added = append(added, r) + } + } + + if !reflect.DeepEqual(deleted, []uint64{7}) { + t.Errorf("deleted handles %v, want [7]", deleted) + } + if !reflect.DeepEqual(added, want) { + t.Errorf("restored rules differ from snapshot:\n got %+v\nwant %+v", added, want) + } + if policy["input"] != uint32(nftables.ChainPolicyAccept) || policy["forward"] != uint32(nftables.ChainPolicyDrop) { + t.Errorf("chain policies %v: input must be restored to accept, forward keep drop", policy) + } +} + func TestRestoreRejectsOtherTable(t *testing.T) { e := testEngine(t, nil) if err := e.Restore(&Snapshot{Table: "other"}); err == nil { diff --git a/internal/tryapply/tryapply.go b/internal/tryapply/tryapply.go index e8b5f30..c0e8fb4 100644 --- a/internal/tryapply/tryapply.go +++ b/internal/tryapply/tryapply.go @@ -3,6 +3,8 @@ package tryapply import ( + "crypto/rand" + "encoding/hex" "encoding/json" "errors" "fmt" @@ -30,12 +32,21 @@ var ( } return nil } + // restore rolls the live table back to a snapshot; replaced in tests. + restore = func(s *nftables.Snapshot) error { + engine, err := nftables.NewEngine(&config.Config{Settings: config.Settings{TableName: s.Table}}) + if err != nil { + return err + } + return engine.Restore(s) + } ) // ErrPending means a try awaits confirmation; nothing else may apply meanwhile. -var ErrPending = errors.New("a 'tomswall try' is pending; run 'tomswall confirm' or 'tomswall revert'") +var ErrPending = errors.New("a 'tomswall try' is pending; run 'tomswall confirm' to keep it or 'tomswall revert' to restore the previous ruleset (also the recovery if an automatic revert failed)") type pending struct { + ID string `json:"id"` PID int `json:"pid"` Snapshot *nftables.Snapshot `json:"snapshot"` } @@ -70,43 +81,49 @@ func lock() (func(), error) { return func() { f.Close() }, nil } -// Arm persists snap and schedules an out-of-process revert after delay. +// Arm persists snap and schedules an out-of-process revert after delay, +// returning the try ID that scopes later reverts to this try. // The caller must hold the lock from Acquire. -func Arm(snap *nftables.Snapshot, pid int, delay time.Duration) error { - b, err := json.Marshal(pending{PID: pid, Snapshot: snap}) +func Arm(snap *nftables.Snapshot, pid int, delay time.Duration) (string, error) { + raw := make([]byte, 8) + if _, err := rand.Read(raw); err != nil { + return "", err + } + id := hex.EncodeToString(raw) + b, err := json.Marshal(pending{ID: id, PID: pid, Snapshot: snap}) if err != nil { - return err + return "", err } f, err := os.CreateTemp(Dir, ".try-snapshot-*") if err != nil { - return err + return "", err } defer os.Remove(f.Name()) if _, err := f.Write(b); err != nil { f.Close() - return err + return "", err } if err := f.Sync(); err != nil { f.Close() - return err + return "", err } if err := f.Close(); err != nil { - return err + return "", err } if err := os.Rename(f.Name(), snapshotPath()); err != nil { - return err + return "", err } exe, err := os.Executable() if err != nil { - return discardWith(err) + return "", discardWith(err) } _ = disarm() // a leftover timer from an earlier try would block the unit name if err := run("systemd-run", "--quiet", "--collect", "--unit", Unit, - fmt.Sprintf("--on-active=%ds", int(delay.Round(time.Second).Seconds())), exe, "revert"); err != nil { - return discardWith(fmt.Errorf("arming revert timer: %w", err)) + fmt.Sprintf("--on-active=%ds", int(delay.Round(time.Second).Seconds())), exe, "revert", "--id", id); err != nil { + return "", discardWith(fmt.Errorf("arming revert timer: %w", err)) } - return nil + return id, nil } // Discard drops the pending snapshot and timer without restoring. The caller must hold the lock. @@ -144,23 +161,21 @@ func Confirm() (pid int, ok bool, err error) { return p.PID, true, Discard() } -// Revert restores the pending snapshot. reverted is false when nothing was -// pending (already confirmed or reverted). -func Revert() (reverted bool, err error) { +// Revert restores the pending snapshot. A non-empty id only reverts that try, +// so a stale timer cannot undo a newer one. reverted is false when nothing +// matching was pending (already confirmed or reverted). A failed restore keeps +// the snapshot so 'tomswall revert' can retry. +func Revert(id string) (reverted bool, err error) { unlock, err := lock() if err != nil { return false, err } defer unlock() p, err := load() - if err != nil || p == nil { + if err != nil || p == nil || (id != "" && p.ID != id) { return false, err } - engine, err := nftables.NewEngine(&config.Config{Settings: config.Settings{TableName: p.Snapshot.Table}}) - if err != nil { - return false, err - } - if err := engine.Restore(p.Snapshot); err != nil { + if err := restore(p.Snapshot); err != nil { return false, fmt.Errorf("restoring snapshot: %w", err) } return true, Discard() diff --git a/internal/tryapply/tryapply_test.go b/internal/tryapply/tryapply_test.go index ab72aa1..5b00f2c 100644 --- a/internal/tryapply/tryapply_test.go +++ b/internal/tryapply/tryapply_test.go @@ -24,23 +24,25 @@ func setup(t *testing.T) *[]string { return &cmds } -func arm(t *testing.T, snap *nftables.Snapshot) { +func arm(t *testing.T, snap *nftables.Snapshot) string { t.Helper() unlock, err := Acquire() if err != nil { t.Fatal(err) } defer unlock() - if err := Arm(snap, 4242, 90*time.Second); err != nil { + id, err := Arm(snap, 4242, 90*time.Second) + if err != nil { t.Fatal(err) } + return id } func TestArmPersistsSnapshotAndTimer(t *testing.T) { cmds := setup(t) snap := &nftables.Snapshot{Table: "tomswall", Present: true, Rules: map[string][]nftables.SnapshotRule{"input": {{Tag: "ssh", Exprs: [][]byte{{1, 2, 3}}}}}} - arm(t, snap) + id := arm(t, snap) info, err := os.Stat(snapshotPath()) if err != nil { @@ -53,12 +55,12 @@ func TestArmPersistsSnapshotAndTimer(t *testing.T) { if err != nil { t.Fatal(err) } - if p.PID != 4242 || !reflect.DeepEqual(p.Snapshot, snap) { + if p.ID != id || id == "" || p.PID != 4242 || !reflect.DeepEqual(p.Snapshot, snap) { t.Errorf("round trip mismatch: %+v", p) } last := (*cmds)[len(*cmds)-1] if !strings.HasPrefix(last, "systemd-run ") || !strings.Contains(last, "--unit "+Unit) || - !strings.Contains(last, "--on-active=90s") || !strings.HasSuffix(last, " revert") { + !strings.Contains(last, "--on-active=90s") || !strings.HasSuffix(last, " revert --id "+id) { t.Errorf("unexpected arm command %q", last) } } @@ -78,9 +80,13 @@ func TestAbsentTableSnapshotRoundTrip(t *testing.T) { func TestAcquireRefusesWhilePending(t *testing.T) { setup(t) arm(t, &nftables.Snapshot{Table: "tomswall"}) - if _, err := Acquire(); !errors.Is(err, ErrPending) { + _, err := Acquire() + if !errors.Is(err, ErrPending) { t.Fatalf("second try: got %v, want ErrPending", err) } + if !strings.Contains(err.Error(), "tomswall revert") { + t.Errorf("ErrPending does not name the recovery: %v", err) + } } func TestArmFailureDiscardsSnapshot(t *testing.T) { @@ -96,7 +102,7 @@ func TestArmFailureDiscardsSnapshot(t *testing.T) { t.Fatal(err) } defer unlock() - if err := Arm(&nftables.Snapshot{Table: "tomswall"}, 1, time.Minute); err == nil { + if _, err := Arm(&nftables.Snapshot{Table: "tomswall"}, 1, time.Minute); err == nil { t.Fatal("expected arm error") } if _, err := os.Stat(snapshotPath()); !os.IsNotExist(err) { @@ -130,8 +136,76 @@ func TestConfirmAfterRevertFails(t *testing.T) { if err != nil || ok { t.Fatalf("Confirm with nothing pending = %v, %v; want not ok", ok, err) } - reverted, err := Revert() + reverted, err := Revert("") if err != nil || reverted { t.Fatalf("Revert with nothing pending = %v, %v", reverted, err) } } + +func stubRestore(t *testing.T, err error) *[]*nftables.Snapshot { + t.Helper() + var got []*nftables.Snapshot + orig := restore + restore = func(s *nftables.Snapshot) error { + got = append(got, s) + return err + } + t.Cleanup(func() { restore = orig }) + return &got +} + +func TestRevertPendingRestoresAndDisarms(t *testing.T) { + cmds := setup(t) + restored := stubRestore(t, nil) + snap := &nftables.Snapshot{Table: "tomswall", Present: true, + Rules: map[string][]nftables.SnapshotRule{"input": {{Tag: "ssh", Exprs: [][]byte{{1}}}}}} + id := arm(t, snap) + + reverted, err := Revert(id) + if err != nil || !reverted { + t.Fatalf("Revert = %v, %v", reverted, err) + } + if len(*restored) != 1 || !reflect.DeepEqual((*restored)[0], snap) { + t.Errorf("restored %+v, want the armed snapshot", *restored) + } + if _, err := os.Stat(snapshotPath()); !os.IsNotExist(err) { + t.Error("snapshot not removed") + } + if last := (*cmds)[len(*cmds)-1]; last != "systemctl stop "+Unit+".timer" { + t.Errorf("timer not stopped, last command %q", last) + } +} + +func TestRevertFailureKeepsSnapshot(t *testing.T) { + setup(t) + boom := errors.New("netlink down") + stubRestore(t, boom) + id := arm(t, &nftables.Snapshot{Table: "tomswall"}) + + if _, err := Revert(id); !errors.Is(err, boom) { + t.Fatalf("Revert error = %v, want %v", err, boom) + } + if _, err := os.Stat(snapshotPath()); err != nil { + t.Fatalf("snapshot gone after failed revert: %v", err) + } + if _, err := Acquire(); !errors.Is(err, ErrPending) { + t.Errorf("failed revert must keep the try pending, got %v", err) + } +} + +func TestRevertStaleIDIgnored(t *testing.T) { + setup(t) + restored := stubRestore(t, nil) + arm(t, &nftables.Snapshot{Table: "tomswall"}) + + reverted, err := Revert("stale-try") + if err != nil || reverted { + t.Fatalf("stale Revert = %v, %v; want no-op", reverted, err) + } + if len(*restored) != 0 { + t.Error("stale timer restored a newer try's snapshot") + } + if _, err := os.Stat(snapshotPath()); err != nil { + t.Errorf("newer try's snapshot removed: %v", err) + } +}