package nftables import ( "encoding/binary" "fmt" "net" "strconv" "strings" "github.com/google/nftables/expr" "golang.org/x/sys/unix" "git.unkin.net/unkin/tomswall/internal/config" ) type Compiler struct { cfg *config.Config } func NewCompiler(cfg *config.Config) *Compiler { return &Compiler{cfg: cfg} } func (c *Compiler) Compile() (*FirewallState, error) { state := &FirewallState{ Rules: make(map[string][]ManagedRule), } c.compileLoopback(state) if err := c.compileConntrackFastPath(state); err != nil { return nil, fmt.Errorf("conntrack fast-path: %w", err) } c.compileAntiSpoof(state) c.compileDHCP(state) if err := c.compileIntraZone(state); err != nil { return nil, fmt.Errorf("intra-zone: %w", err) } if err := c.compileBlrules(state); err != nil { return nil, fmt.Errorf("blrules: %w", err) } if err := c.compileConntrack(state); err != nil { return nil, fmt.Errorf("conntrack: %w", err) } if err := c.compileTunnels(state); err != nil { return nil, fmt.Errorf("tunnels: %w", err) } if err := c.compileRules(state); err != nil { return nil, fmt.Errorf("rules: %w", err) } if err := c.compilePolicies(state); err != nil { return nil, fmt.Errorf("policies: %w", err) } if err := c.compileSNAT(state); err != nil { return nil, fmt.Errorf("snat: %w", err) } if err := c.compileDNAT(state); err != nil { return nil, fmt.Errorf("dnat: %w", err) } if err := c.compileStaticNAT(state); err != nil { return nil, fmt.Errorf("static-nat: %w", err) } c.compileMSSClamp(state) return state, nil } func (c *Compiler) compileConntrackFastPath(state *FirewallState) error { for _, chain := range []string{"input", "forward", "output"} { state.Rules[chain] = append(state.Rules[chain], ManagedRule{ Chain: chain, Exprs: append(matchCtState(ctStateEstablished|ctStateRelated), &expr.Verdict{Kind: expr.VerdictAccept}), Tag: "ct:fastpath:" + chain, }, ManagedRule{ Chain: chain, Exprs: append(matchCtState(ctStateInvalid), &expr.Verdict{Kind: expr.VerdictDrop}), Tag: "ct:invalid:" + chain, }, ) } return nil } func (c *Compiler) compileLoopback(state *FirewallState) { for _, chain := range []string{"input", "output"} { state.Rules[chain] = append(state.Rules[chain], ManagedRule{ Chain: chain, Exprs: append(matchIfaceName(chain == "input", "lo"), &expr.Verdict{Kind: expr.VerdictAccept}), Tag: "loopback:" + chain, }) } } func (c *Compiler) compileAntiSpoof(state *FirewallState) { for _, iface := range c.cfg.Interfaces { if iface.Options.NoSmurfs { state.Rules["input"] = append(state.Rules["input"], ManagedRule{ Chain: "input", Exprs: matchSmurfDrop(iface.PhysicalName()), Tag: fmt.Sprintf("antismurf:%s", iface.Interface), }) } if iface.Options.TCPFlags != nil && *iface.Options.TCPFlags { state.Rules["input"] = append(state.Rules["input"], ManagedRule{ Chain: "input", Exprs: matchTCPFlagsDrop(iface.PhysicalName()), Tag: fmt.Sprintf("tcpflags:%s", iface.Interface), }) } } } func (c *Compiler) compileDHCP(state *FirewallState) { for _, iface := range c.cfg.Interfaces { if !iface.Options.DHCP { continue } name := iface.PhysicalName() // Allow DHCPv4 client traffic (bootpc:68 → bootps:67) state.Rules["input"] = append(state.Rules["input"], ManagedRule{ Chain: "input", Exprs: append(append(append( matchIfaceName(true, name), matchProtoNum(unix.IPPROTO_UDP)...), matchSPort(68)...), matchDPort(67)..., ), Tag: fmt.Sprintf("dhcp:in:%s", iface.Interface), }) // Allow DHCPv4 server → client replies state.Rules["input"] = append(state.Rules["input"], ManagedRule{ Chain: "input", Exprs: append(append(append(append( matchIfaceName(true, name), matchProtoNum(unix.IPPROTO_UDP)...), matchSPort(67)...), matchDPort(68)...), &expr.Verdict{Kind: expr.VerdictAccept}, ), Tag: fmt.Sprintf("dhcp:reply:%s", iface.Interface), }) state.Rules["output"] = append(state.Rules["output"], ManagedRule{ Chain: "output", Exprs: append(append(append(append( matchIfaceName(false, name), matchProtoNum(unix.IPPROTO_UDP)...), matchSPort(68)...), matchDPort(67)...), &expr.Verdict{Kind: expr.VerdictAccept}, ), Tag: fmt.Sprintf("dhcp:out:%s", iface.Interface), }) } } func (c *Compiler) compileIntraZone(state *FirewallState) error { fwZone := c.cfg.FirewallZone() for _, iface := range c.cfg.Interfaces { if iface.Options.RouteBack != nil && *iface.Options.RouteBack { chain := "forward" if iface.Zone == fwZone { continue } state.Rules[chain] = append(state.Rules[chain], ManagedRule{ Chain: chain, Exprs: append( append(matchIfaceName(true, iface.PhysicalName()), matchIfaceName(false, iface.PhysicalName())...), &expr.Verdict{Kind: expr.VerdictAccept}, ), Tag: fmt.Sprintf("intra:%s:%s", iface.Zone, iface.Interface), }) } } return nil } func (c *Compiler) compileBlrules(state *FirewallState) error { fwZone := c.cfg.FirewallZone() blruleToRuleAction := map[config.BlruleAction]config.RuleAction{ config.BlruleAccept: config.RuleAccept, config.BlruleWhitelist: config.RuleAccept, config.BlruleDrop: config.RuleDrop, config.BlruleReject: config.RuleReject, config.BlruleLog: config.RuleLog, config.BlruleContinue: config.RuleContinue, } for i, rule := range c.cfg.Blrules { tag := fmt.Sprintf("blrule:%d", i) action, ok := blruleToRuleAction[rule.Action] if !ok { action = config.RuleDrop } if err := c.compileOneRule(state, tag, rule.Source, rule.Dest, rule.Proto, rule.DPort, rule.SPort, action, rule.Log, "", fwZone, ""); err != nil { return fmt.Errorf("blrule[%d]: %w", i, err) } } return nil } func (c *Compiler) compileConntrack(state *FirewallState) error { for i, ct := range c.cfg.Conntrack { tag := fmt.Sprintf("conntrack:%d", i) chains := []string{"prerouting"} switch ct.Chain { case config.ConntrackOutput: chains = []string{"output"} case config.ConntrackBoth: chains = []string{"prerouting", "output"} } for _, chain := range chains { var exprs []expr.Any 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) } 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, }) } } return nil } func (c *Compiler) compileRules(state *FirewallState) error { fwZone := c.cfg.FirewallZone() for i, rule := range c.cfg.Rules { tag := fmt.Sprintf("rule:%d", i) proto := rule.Proto var dports config.PortSpec var sport config.PortSpec if rule.PortGroup != "" { pg, _ := c.cfg.ResolvePortGroup(rule.PortGroup) proto = pg.Proto dports = pg.Ports } else { dports = rule.DPort sport = rule.SPort } if err := c.compileOneRule(state, tag, rule.Source, rule.Dest, proto, dports, sport, rule.Action, rule.Log, rule.Dest, fwZone, rule.Section); err != nil { return fmt.Errorf("rule[%d]: %w", i, err) } if rule.RateLimit != "" || rule.User != "" || rule.Mark != "" || rule.SetMark != "" || rule.ConnLimit != "" || rule.Time != nil || rule.Action == config.RuleMark || rule.Action == config.RuleConnMark || rule.Action == config.RuleNFQueue { c.applyRuleExtras(state, tag, rule) } } return nil } 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) rules := state.Rules[chain] for idx := len(rules) - 1; idx >= 0; idx-- { if rules[idx].Tag != tag { break } var extra []expr.Any var replaceVerdict []expr.Any if rule.User != "" { extra = append(extra, matchUID(rule.User)...) } if rule.Mark != "" { extra = append(extra, matchMark(rule.Mark)...) } if rule.RateLimit != "" { extra = append(extra, parseRateLimit(rule.RateLimit)...) } if rule.ConnLimit != "" { extra = append(extra, matchConnLimit(rule.ConnLimit)...) } if rule.Time != nil { extra = append(extra, matchTime(rule.Time)...) } if rule.SetMark != "" { extra = append(extra, setMarkExprs(rule.SetMark)...) } if rule.Action == config.RuleNFQueue { replaceVerdict = []expr.Any{&expr.Queue{Num: uint16(rule.NFQueue)}} } if len(extra) > 0 || len(replaceVerdict) > 0 { existingExprs := rules[idx].Exprs var verdict []expr.Any var nonVerdict []expr.Any for _, e := range existingExprs { if _, ok := e.(*expr.Verdict); ok { verdict = append(verdict, e) } else { nonVerdict = append(nonVerdict, e) } } if len(replaceVerdict) > 0 { verdict = replaceVerdict } rules[idx].Exprs = append(append(nonVerdict, extra...), verdict...) } } state.Rules[chain] = rules } func (c *Compiler) compileOneRule(state *FirewallState, tag, srcSpec, dstSpec, proto string, 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) } srcIfaces := c.resolveZoneInterfaces(srcZone) dstIfaces := c.resolveZoneInterfaces(dstZone) chain := c.selectChain(srcZone, dstZone, fwZone) for _, srcIface := range srcIfaces { for _, dstIface := range dstIfaces { exprs, 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)...) } if logLevel != "" { exprs = append(exprs, buildLog(logLevel, tag)...) } verdict := actionVerdict(action, 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 } 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) chain := "prerouting" parts := strings.SplitN(dstSpec, ":", 3) if len(parts) < 2 { return fmt.Errorf("DNAT dest must be zone:address or zone:address:port") } dnatAddr := parts[1] var dnatPort uint16 if len(parts) == 3 { p, err := strconv.ParseUint(parts[2], 10, 16) if err != nil { return fmt.Errorf("invalid DNAT port %q: %w", parts[2], err) } dnatPort = uint16(p) } srcIfaces := c.resolveZoneInterfaces(srcZone) for _, srcIface := range srcIfaces { var exprs []expr.Any if srcIface != "" { exprs = append(exprs, matchIfaceName(true, srcIface)...) } if srcAddr != "" { src, err := matchSourceCIDR(srcAddr) if err != nil { return err } 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 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) } } state.Rules[chain] = append(state.Rules[chain], ManagedRule{ Chain: chain, Exprs: exprs, Tag: tag, }) } return nil } func (c *Compiler) compilePolicies(state *FirewallState) error { fwZone := c.cfg.FirewallZone() for i, pol := range c.cfg.Policy { tag := fmt.Sprintf("policy:%d", i) srcZones := c.expandZoneRef(pol.Source) dstZones := c.expandZoneRef(pol.Dest) for _, sz := range srcZones { for _, dz := range dstZones { if sz == dz && !strings.HasSuffix(pol.Source, "+") { continue } chain := c.selectChain(sz, dz, fwZone) srcIfaces := c.resolveZoneInterfaces(sz) dstIfaces := c.resolveZoneInterfaces(dz) for _, si := range srcIfaces { for _, di := range dstIfaces { var exprs []expr.Any if si != "" { exprs = append(exprs, matchIfaceName(true, si)...) } if di != "" && chain == "forward" { exprs = append(exprs, matchIfaceName(false, di)...) } if pol.RateLimit != "" { exprs = append(exprs, parseRateLimit(pol.RateLimit)...) } if pol.ConnLimit != "" { exprs = append(exprs, matchConnLimit(pol.ConnLimit)...) } if pol.Log != "" { exprs = append(exprs, buildLog(pol.Log, tag)...) } exprs = append(exprs, policyVerdict(pol.Action, c.cfg.Settings.AddressFamily)...) state.Rules[chain] = append(state.Rules[chain], ManagedRule{ Chain: chain, Exprs: exprs, Tag: tag, }) } } } } } return nil } func (c *Compiler) compileSNAT(state *FirewallState) error { for i, snat := range c.cfg.SNAT { tag := fmt.Sprintf("snat:%d", i) var exprs []expr.Any destIface, _ := splitZoneSpec(snat.Dest) exprs = append(exprs, matchIfaceName(false, destIface)...) if snat.Source != "" { srcExprs, err := matchSourceCIDR(snat.Source) if err != nil { return fmt.Errorf("snat[%d]: %w", i, err) } 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...) } if snat.Mark != "" { exprs = append(exprs, matchMark(snat.Mark)...) } if snat.Log != "" { exprs = append(exprs, buildLog(snat.Log, tag)...) } switch snat.Action { case config.SNATMasquerade: masq := &expr.Masq{} if snat.Random { masq.Random = true } exprs = append(exprs, masq) case config.SNATAddress: ip := net.ParseIP(snat.Address) if ip == nil { return fmt.Errorf("snat[%d]: invalid address %q", i, snat.Address) } ip4 := ip.To4() if ip4 != nil { exprs = append(exprs, &expr.Immediate{Register: 1, Data: ip4}, &expr.NAT{ Type: expr.NATTypeSourceNAT, Family: unix.NFPROTO_IPV4, RegAddrMin: 1, RegAddrMax: 1, }, ) } else { exprs = append(exprs, &expr.Immediate{Register: 1, Data: ip.To16()}, &expr.NAT{ Type: expr.NATTypeSourceNAT, Family: unix.NFPROTO_IPV6, RegAddrMin: 1, RegAddrMax: 1, }, ) } } state.Rules["postrouting"] = append(state.Rules["postrouting"], ManagedRule{ Chain: "postrouting", Exprs: exprs, Tag: tag, }) } return nil } func (c *Compiler) compileDNAT(state *FirewallState) error { return nil } func (c *Compiler) compileTunnels(state *FirewallState) error { fwZone := c.cfg.FirewallZone() for i, tun := range c.cfg.Tunnels { baseType, extra, _ := config.ParseTunnelType(tun.Type) tag := fmt.Sprintf("tunnel:%d", i) inChain := c.selectChain(tun.Zone, fwZone, fwZone) outChain := c.selectChain(fwZone, tun.Zone, fwZone) for _, gw := range tun.Gateways { var srcMatch, dstMatch []expr.Any if gw != "0.0.0.0/0" && gw != "::/0" { var err error srcMatch, err = matchSourceCIDR(gw) if err != nil { return fmt.Errorf("tunnel[%d]: %w", i, err) } dstMatch, err = matchDestCIDR(gw) if err != nil { return fmt.Errorf("tunnel[%d]: %w", i, err) } } addTunnelRule := func(chain string, proto byte, dport uint16, srcExprs []expr.Any) { var exprs []expr.Any exprs = append(exprs, srcExprs...) exprs = append(exprs, matchProtoNum(proto)...) if dport > 0 { exprs = append(exprs, matchDPort(dport)...) } exprs = append(exprs, &expr.Verdict{Kind: expr.VerdictAccept}) state.Rules[chain] = append(state.Rules[chain], ManagedRule{ Chain: chain, Exprs: exprs, Tag: tag, }) } switch baseType { case config.TunnelIPSec, config.TunnelIPSecNAT: addTunnelRule(inChain, 50, 0, srcMatch) addTunnelRule(outChain, 50, 0, dstMatch) if extra != "ah" { addTunnelRule(inChain, 51, 0, srcMatch) addTunnelRule(outChain, 51, 0, dstMatch) } addTunnelRule(inChain, unix.IPPROTO_UDP, 500, srcMatch) addTunnelRule(outChain, unix.IPPROTO_UDP, 500, dstMatch) if baseType == config.TunnelIPSecNAT { addTunnelRule(inChain, unix.IPPROTO_UDP, 4500, srcMatch) addTunnelRule(outChain, unix.IPPROTO_UDP, 4500, dstMatch) } case config.TunnelIPIP, config.Tunnel6to4: addTunnelRule(inChain, 4, 0, srcMatch) addTunnelRule(outChain, 4, 0, dstMatch) case config.TunnelGRE: addTunnelRule(inChain, 47, 0, srcMatch) addTunnelRule(outChain, 47, 0, dstMatch) case config.TunnelOpenVPN, config.TunnelOpenVPNClient, config.TunnelOpenVPNServer: proto := unix.IPPROTO_UDP if extra == "tcp" { proto = unix.IPPROTO_TCP } port := uint16(1194) if tun.Port > 0 { port = uint16(tun.Port) } addTunnelRule(inChain, byte(proto), port, srcMatch) addTunnelRule(outChain, byte(proto), port, dstMatch) case config.TunnelL2TP: addTunnelRule(inChain, unix.IPPROTO_UDP, 1701, srcMatch) addTunnelRule(outChain, unix.IPPROTO_UDP, 1701, dstMatch) case config.TunnelTinc: addTunnelRule(inChain, unix.IPPROTO_UDP, 655, srcMatch) addTunnelRule(outChain, unix.IPPROTO_UDP, 655, dstMatch) addTunnelRule(inChain, unix.IPPROTO_TCP, 655, srcMatch) addTunnelRule(outChain, unix.IPPROTO_TCP, 655, dstMatch) case config.TunnelPPTPClient: addTunnelRule(inChain, 47, 0, srcMatch) addTunnelRule(outChain, 47, 0, dstMatch) addTunnelRule(outChain, unix.IPPROTO_TCP, 1723, dstMatch) case config.TunnelPPTPServer: addTunnelRule(inChain, 47, 0, srcMatch) addTunnelRule(outChain, 47, 0, dstMatch) addTunnelRule(inChain, unix.IPPROTO_TCP, 1723, srcMatch) case config.TunnelGeneric: proto := unix.IPPROTO_UDP if extra == "tcp" { proto = unix.IPPROTO_TCP } port := uint16(0) if tun.Port > 0 { port = uint16(tun.Port) } addTunnelRule(inChain, byte(proto), port, srcMatch) addTunnelRule(outChain, byte(proto), port, dstMatch) } } } return nil } func (c *Compiler) compileMSSClamp(state *FirewallState) { for _, iface := range c.cfg.Interfaces { if iface.Options.MSS > 0 { mssBytes := make([]byte, 2) binary.BigEndian.PutUint16(mssBytes, uint16(iface.Options.MSS)) var exprs []expr.Any exprs = append(exprs, matchIfaceName(false, iface.PhysicalName())...) exprs = append(exprs, matchProtoNum(unix.IPPROTO_TCP)...) exprs = append(exprs, matchTCPFlags(0x02, 0x02)...) exprs = append(exprs, &expr.Exthdr{ DestRegister: 1, Type: 2, Offset: 2, Len: 2, Op: 0, }, &expr.Cmp{Op: expr.CmpOpGt, Register: 1, Data: mssBytes}, &expr.Immediate{Register: 1, Data: mssBytes}, &expr.Exthdr{ SourceRegister: 1, Type: 2, Offset: 2, Len: 2, Op: 1, }, ) state.Rules["forward"] = append(state.Rules["forward"], ManagedRule{ Chain: "forward", Exprs: exprs, Tag: fmt.Sprintf("mss:%s", iface.Interface), }) } } } func (c *Compiler) compileStaticNAT(state *FirewallState) error { for i, sn := range c.cfg.StaticNAT { extIP := net.ParseIP(sn.External) intIP := net.ParseIP(sn.Internal) if extIP == nil || intIP == nil { return fmt.Errorf("static-nat[%d]: invalid IP", i) } family := unix.NFPROTO_IPV4 ext4 := extIP.To4() int4 := intIP.To4() if ext4 == nil || int4 == nil { family = unix.NFPROTO_IPV6 } dnatTag := fmt.Sprintf("staticnat:dnat:%d", i) var dnatExprs []expr.Any dnatExprs = append(dnatExprs, matchIfaceName(true, sn.Interface)...) if family == unix.NFPROTO_IPV4 { dst, _ := matchDestCIDR(sn.External) dnatExprs = append(dnatExprs, dst...) dnatExprs = append(dnatExprs, &expr.Immediate{Register: 1, Data: int4}, &expr.NAT{Type: expr.NATTypeDestNAT, Family: uint32(family), RegAddrMin: 1, RegAddrMax: 1}, ) } else { dst, _ := matchDestCIDR(sn.External) dnatExprs = append(dnatExprs, dst...) dnatExprs = append(dnatExprs, &expr.Immediate{Register: 1, Data: intIP.To16()}, &expr.NAT{Type: expr.NATTypeDestNAT, Family: uint32(family), RegAddrMin: 1, RegAddrMax: 1}, ) } state.Rules["prerouting"] = append(state.Rules["prerouting"], ManagedRule{ Chain: "prerouting", Exprs: dnatExprs, Tag: dnatTag, }) snatTag := fmt.Sprintf("staticnat:snat:%d", i) var snatExprs []expr.Any snatExprs = append(snatExprs, matchIfaceName(false, sn.Interface)...) if family == unix.NFPROTO_IPV4 { src, _ := matchSourceCIDR(sn.Internal) snatExprs = append(snatExprs, src...) snatExprs = append(snatExprs, &expr.Immediate{Register: 1, Data: ext4}, &expr.NAT{Type: expr.NATTypeSourceNAT, Family: uint32(family), RegAddrMin: 1, RegAddrMax: 1}, ) } else { src, _ := matchSourceCIDR(sn.Internal) snatExprs = append(snatExprs, src...) snatExprs = append(snatExprs, &expr.Immediate{Register: 1, Data: extIP.To16()}, &expr.NAT{Type: expr.NATTypeSourceNAT, Family: uint32(family), RegAddrMin: 1, RegAddrMax: 1}, ) } state.Rules["postrouting"] = append(state.Rules["postrouting"], ManagedRule{ Chain: "postrouting", Exprs: snatExprs, Tag: snatTag, }) } return nil } func (c *Compiler) selectChain(srcZone, dstZone, fwZone string) string { if dstZone == fwZone { return "input" } if srcZone == fwZone { return "output" } return "forward" } func (c *Compiler) resolveZoneInterfaces(zone string) []string { if zone == "all" || zone == "" { return []string{""} } ifaces := c.cfg.ZoneInterfaces(zone) if len(ifaces) == 0 { return []string{""} } return ifaces } func (c *Compiler) expandZoneRef(ref string) []string { base := ref var excluded map[string]bool if idx := strings.IndexByte(ref, '!'); idx >= 0 { base = ref[:idx] excluded = make(map[string]bool) for _, z := range strings.Split(ref[idx+1:], ",") { z = strings.TrimSpace(z) if z != "" { excluded[z] = true } } } if base == "all" || base == "all+" { var zones []string for name := range c.cfg.Zones { if excluded != nil && excluded[name] { continue } zones = append(zones, name) } return zones } return []string{base} } func (c *Compiler) buildMatchExprs(srcIface, dstIface, chain, proto string, dports, sports config.PortSpec, srcAddr, dstAddr string) ([]expr.Any, error) { var exprs []expr.Any if srcIface != "" { exprs = append(exprs, matchIfaceName(true, srcIface)...) } if dstIface != "" && chain == "forward" { exprs = append(exprs, matchIfaceName(false, dstIface)...) } if srcAddr != "" { src, err := matchSourceCIDR(srcAddr) if err != nil { return nil, err } exprs = append(exprs, src...) } if dstAddr != "" { dst, err := matchDestCIDR(dstAddr) if err != nil { return nil, err } exprs = append(exprs, dst...) } if proto != "" { exprs = append(exprs, matchProto(proto)...) } isICMP := strings.EqualFold(proto, "icmp") || strings.EqualFold(proto, "icmpv6") || strings.EqualFold(proto, "ipv6-icmp") for _, portStr := range dports { if isICMP { pe := matchICMPType(portStr) exprs = append(exprs, pe...) } else { pe, err := parsePortOrRange(portStr) if err != nil { return nil, err } exprs = append(exprs, pe...) } } for _, portStr := range sports { pe, err := parseSPortOrRange(portStr) if err != nil { return nil, err } exprs = append(exprs, pe...) } return exprs, nil } // matchIfaceName matches an interface name, supporting wildcard "+" suffix. func matchIfaceName(input bool, name string) []expr.Any { key := expr.MetaKeyOIFNAME if input { key = expr.MetaKeyIIFNAME } if strings.HasSuffix(name, "+") { prefix := strings.TrimSuffix(name, "+") padded := make([]byte, len(prefix)) copy(padded, prefix) return []expr.Any{ &expr.Meta{Key: key, Register: 1}, &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: padded}, } } padded := make([]byte, 16) copy(padded, name+"\x00") return []expr.Any{ &expr.Meta{Key: key, Register: 1}, &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: padded[:len(name)+1]}, } } func matchProto(proto string) []expr.Any { 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, _ := strconv.Atoi(proto) protoNum = byte(n) } return []expr.Any{ &expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1}, &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{protoNum}}, } } func matchDPort(port uint16) []expr.Any { portBytes := make([]byte, 2) binary.BigEndian.PutUint16(portBytes, port) return []expr.Any{ &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2}, &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: portBytes}, } } func matchDPortRange(low, high uint16) []expr.Any { lowBytes := make([]byte, 2) highBytes := make([]byte, 2) binary.BigEndian.PutUint16(lowBytes, low) binary.BigEndian.PutUint16(highBytes, high) return []expr.Any{ &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2}, &expr.Cmp{Op: expr.CmpOpGte, Register: 1, Data: lowBytes}, &expr.Cmp{Op: expr.CmpOpLte, Register: 1, Data: highBytes}, } } func matchSPort(port uint16) []expr.Any { portBytes := make([]byte, 2) binary.BigEndian.PutUint16(portBytes, port) return []expr.Any{ &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 0, Len: 2}, &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: portBytes}, } } func matchSPortRange(low, high uint16) []expr.Any { lowBytes := make([]byte, 2) highBytes := make([]byte, 2) binary.BigEndian.PutUint16(lowBytes, low) binary.BigEndian.PutUint16(highBytes, high) return []expr.Any{ &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 0, Len: 2}, &expr.Cmp{Op: expr.CmpOpGte, Register: 1, Data: lowBytes}, &expr.Cmp{Op: expr.CmpOpLte, Register: 1, Data: highBytes}, } } func matchProtoNum(proto byte) []expr.Any { return []expr.Any{ &expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1}, &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{proto}}, } } func matchTCPFlags(flags, mask byte) []expr.Any { return []expr.Any{ &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 13, Len: 1}, &expr.Bitwise{ SourceRegister: 1, DestRegister: 1, Len: 1, Mask: []byte{mask}, Xor: []byte{0}, }, &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{flags}}, } } func matchSmurfDrop(iface string) []expr.Any { var exprs []expr.Any exprs = append(exprs, matchIfaceName(true, iface)...) exprs = append(exprs, &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4}, &expr.Bitwise{ SourceRegister: 1, DestRegister: 1, Len: 4, Mask: []byte{0xf0, 0, 0, 0}, Xor: []byte{0, 0, 0, 0}, }, &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{0xe0, 0, 0, 0}}, &expr.Verdict{Kind: expr.VerdictDrop}, ) return exprs } func matchTCPFlagsDrop(iface string) []expr.Any { var exprs []expr.Any exprs = append(exprs, matchIfaceName(true, iface)...) exprs = append(exprs, matchProtoNum(unix.IPPROTO_TCP)...) exprs = append(exprs, &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 13, Len: 1}, &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{0}}, &expr.Verdict{Kind: expr.VerdictDrop}, ) return exprs } var icmpTypeNames = map[string]byte{ "echo-reply": 0, "destination-unreachable": 3, "source-quench": 4, "redirect": 5, "echo-request": 8, "router-advertisement": 9, "router-solicitation": 10, "time-exceeded": 11, "parameter-problem": 12, "timestamp-request": 13, "timestamp-reply": 14, "address-mask-request": 17, "address-mask-reply": 18, } func matchICMPType(spec string) []expr.Any { if strings.Contains(spec, "/") { parts := strings.SplitN(spec, "/", 2) typeVal, ok := resolveICMPType(parts[0]) if !ok { return nil } code, err := strconv.ParseUint(parts[1], 10, 8) if err != nil { return nil } 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)}}, } } typeVal, ok := resolveICMPType(spec) if !ok { return nil } return []expr.Any{ &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 0, Len: 1}, &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{typeVal}}, } } func resolveICMPType(s string) (byte, bool) { if v, ok := icmpTypeNames[strings.ToLower(s)]; ok { return v, true } n, err := strconv.ParseUint(s, 10, 8) if err != nil { return 0, false } return byte(n), true } // Time matching requires NFT_META_TIME_* keys not exposed in google/nftables v0.2.0. func matchTime(_ *config.TimeSpec) []expr.Any { return nil } func matchConnLimit(spec string) []expr.Any { s := spec flags := uint32(0) if strings.HasPrefix(s, "d:") { flags = 1 s = s[2:] } var count uint32 if idx := strings.IndexByte(s, ':'); idx >= 0 { c, err := strconv.ParseUint(s[:idx], 10, 32) if err != nil { return nil } count = uint32(c) } else { c, err := strconv.ParseUint(s, 10, 32) if err != nil { return nil } count = uint32(c) } return []expr.Any{ &expr.Connlimit{ Count: count, Flags: flags, }, } } func parsePortOrRange(s string) ([]expr.Any, error) { if strings.Contains(s, "-") { parts := strings.SplitN(s, "-", 2) low, err := strconv.ParseUint(parts[0], 10, 16) if err != nil { return nil, fmt.Errorf("invalid port range low %q: %w", parts[0], err) } high, err := strconv.ParseUint(parts[1], 10, 16) if err != nil { return nil, fmt.Errorf("invalid port range high %q: %w", parts[1], err) } return matchDPortRange(uint16(low), uint16(high)), nil } p, err := parsePort(s) if err != nil { return nil, err } return matchDPort(p), nil } func parseSPortOrRange(s string) ([]expr.Any, error) { if strings.Contains(s, "-") { parts := strings.SplitN(s, "-", 2) low, err := strconv.ParseUint(parts[0], 10, 16) if err != nil { return nil, fmt.Errorf("invalid sport range low %q: %w", parts[0], err) } high, err := strconv.ParseUint(parts[1], 10, 16) if err != nil { return nil, fmt.Errorf("invalid sport range high %q: %w", parts[1], err) } return matchSPortRange(uint16(low), uint16(high)), nil } p, err := parsePort(s) if err != nil { return nil, err } return matchSPort(p), nil } func matchSourceCIDR(cidr string) ([]expr.Any, error) { return matchAddrCIDR(cidr, true) } func matchDestCIDR(cidr string) ([]expr.Any, error) { return matchAddrCIDR(cidr, false) } func matchAddrCIDR(cidr string, isSrc bool) ([]expr.Any, error) { negated := false if strings.HasPrefix(cidr, "!") { negated = true cidr = cidr[1:] } cmpOp := expr.CmpOpEq if negated { cmpOp = expr.CmpOpNeq } var offset4, offset6 uint32 if isSrc { offset4, offset6 = 12, 8 } else { offset4, offset6 = 16, 24 } ip, ipNet, err := net.ParseCIDR(cidr) if err != nil { singleIP := net.ParseIP(cidr) if singleIP == nil { if isSrc { return nil, fmt.Errorf("invalid source address %q", cidr) } return nil, fmt.Errorf("invalid dest address %q", cidr) } if ip4 := singleIP.To4(); ip4 != nil { return []expr.Any{ &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: offset4, Len: 4}, &expr.Cmp{Op: cmpOp, Register: 1, Data: ip4}, }, nil } ip6 := singleIP.To16() return []expr.Any{ &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: offset6, Len: 16}, &expr.Cmp{Op: cmpOp, Register: 1, Data: ip6}, }, nil } if ip4 := ip.To4(); ip4 != nil { return []expr.Any{ &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: offset4, Len: 4}, &expr.Bitwise{ SourceRegister: 1, DestRegister: 1, Len: 4, Mask: ipNet.Mask, Xor: []byte{0, 0, 0, 0}, }, &expr.Cmp{Op: cmpOp, Register: 1, Data: ipNet.IP.To4()}, }, nil } return []expr.Any{ &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: offset6, Len: 16}, &expr.Bitwise{ SourceRegister: 1, DestRegister: 1, Len: 16, Mask: ipNet.Mask, Xor: make([]byte, 16), }, &expr.Cmp{Op: cmpOp, Register: 1, Data: ipNet.IP.To16()}, }, nil } const ( ctStateInvalid = 1 ctStateEstablished = 2 ctStateRelated = 4 ctStateNew = 8 ctStateUntracked = 64 ) func matchCtState(stateMask uint32) []expr.Any { stateBytes := make([]byte, 4) binary.NativeEndian.PutUint32(stateBytes, stateMask) return []expr.Any{ &expr.Ct{Key: expr.CtKeySTATE, Register: 1}, &expr.Bitwise{ SourceRegister: 1, DestRegister: 1, Len: 4, Mask: stateBytes, Xor: make([]byte, 4), }, &expr.Cmp{Op: expr.CmpOpNeq, Register: 1, Data: make([]byte, 4)}, } } func matchSection(section config.RuleSection) []expr.Any { switch section { case config.SectionEstablished: return matchCtState(ctStateEstablished) case config.SectionRelated: return matchCtState(ctStateRelated) case config.SectionInvalid: return matchCtState(ctStateInvalid) case config.SectionUntracked: return matchCtState(ctStateUntracked) case config.SectionNew: return matchCtState(ctStateNew) default: return nil } } func parseRateLimit(spec string) []expr.Any { s := spec if strings.HasPrefix(s, "s:") || strings.HasPrefix(s, "d:") { s = s[2:] } if idx := strings.IndexByte(s, ':'); idx > 0 { if strings.Contains(s[:idx], "/") { // name:rate/unit:burst → skip name } else { s = s[idx+1:] } } var burst uint32 if idx := strings.LastIndexByte(s, ':'); idx > 0 { b, err := strconv.ParseUint(s[idx+1:], 10, 32) if err == nil { burst = uint32(b) s = s[:idx] } } parts := strings.SplitN(s, "/", 2) if len(parts) != 2 { return nil } rate, err := strconv.ParseUint(parts[0], 10, 64) if err != nil || rate == 0 { return nil } var unit expr.LimitTime switch strings.ToLower(parts[1]) { case "sec", "second": unit = expr.LimitTimeSecond case "min", "minute": unit = expr.LimitTimeMinute case "hour": unit = expr.LimitTimeHour case "day": unit = expr.LimitTimeDay default: return nil } if burst == 0 { burst = 5 } return []expr.Any{ &expr.Limit{ Type: expr.LimitTypePkts, Rate: rate, Unit: unit, Burst: burst, }, } } func matchUID(userSpec string) []expr.Any { negated := false s := userSpec if strings.HasPrefix(s, "!") { negated = true s = s[1:] } if idx := strings.IndexByte(s, ':'); idx >= 0 { s = s[:idx] } uid, err := strconv.ParseUint(s, 10, 32) if err != nil { return nil } uidBytes := make([]byte, 4) binary.NativeEndian.PutUint32(uidBytes, uint32(uid)) op := expr.CmpOpEq if negated { op = expr.CmpOpNeq } return []expr.Any{ &expr.Meta{Key: expr.MetaKeySKUID, Register: 1}, &expr.Cmp{Op: op, Register: 1, Data: uidBytes}, } } func matchMark(markSpec string) []expr.Any { negated := false s := markSpec if strings.HasPrefix(s, "!") { negated = true s = s[1:] } connMark := false if strings.HasSuffix(s, ":C") { connMark = true s = strings.TrimSuffix(s, ":C") } var value, mask uint32 if idx := strings.IndexByte(s, '/'); idx >= 0 { v, err := strconv.ParseUint(s[:idx], 0, 32) if err != nil { return nil } m, err := strconv.ParseUint(s[idx+1:], 0, 32) if err != nil { return nil } value = uint32(v) mask = uint32(m) } else { v, err := strconv.ParseUint(s, 0, 32) if err != nil { return nil } value = uint32(v) mask = 0xffffffff } valBytes := make([]byte, 4) binary.NativeEndian.PutUint32(valBytes, value) maskBytes := make([]byte, 4) binary.NativeEndian.PutUint32(maskBytes, mask) op := expr.CmpOpEq if negated { op = expr.CmpOpNeq } var loadExpr expr.Any if connMark { loadExpr = &expr.Ct{Key: expr.CtKeyMARK, Register: 1} } else { loadExpr = &expr.Meta{Key: expr.MetaKeyMARK, Register: 1} } if mask != 0xffffffff { return []expr.Any{ loadExpr, &expr.Bitwise{ SourceRegister: 1, DestRegister: 1, Len: 4, Mask: maskBytes, Xor: make([]byte, 4), }, &expr.Cmp{Op: op, Register: 1, Data: valBytes}, } } return []expr.Any{ loadExpr, &expr.Cmp{Op: op, Register: 1, Data: valBytes}, } } func setMarkExprs(markSpec string) []expr.Any { var value, mask uint32 if idx := strings.IndexByte(markSpec, '/'); idx >= 0 { v, err := strconv.ParseUint(markSpec[:idx], 0, 32) if err != nil { return nil } m, err := strconv.ParseUint(markSpec[idx+1:], 0, 32) if err != nil { return nil } value = uint32(v) mask = uint32(m) } else { v, err := strconv.ParseUint(markSpec, 0, 32) if err != nil { return nil } value = uint32(v) mask = 0xffffffff } valBytes := make([]byte, 4) binary.NativeEndian.PutUint32(valBytes, value) if mask != 0xffffffff { maskBytes := make([]byte, 4) binary.NativeEndian.PutUint32(maskBytes, mask) return []expr.Any{ &expr.Meta{Key: expr.MetaKeyMARK, Register: 1}, &expr.Bitwise{ SourceRegister: 1, DestRegister: 1, Len: 4, Mask: maskBytes, Xor: valBytes, }, &expr.Meta{Key: expr.MetaKeyMARK, SourceRegister: true, Register: 1}, } } return []expr.Any{ &expr.Immediate{Register: 1, Data: valBytes}, &expr.Meta{Key: expr.MetaKeyMARK, SourceRegister: true, Register: 1}, } } func buildLog(level, prefix string) []expr.Any { nfLevel := logLevelToNF(level) logPrefix := prefix if len(logPrefix) > 63 { logPrefix = logPrefix[:63] } return []expr.Any{ &expr.Log{ Key: 1<