Files
tomswall/internal/nftables/compiler.go
T
unkin-agent ff5a52b9e2
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
Reject ORIGDEST on forwarded rules and mixed-family lists
2026-10-03 22:18:48 +10:00

1869 lines
49 KiB
Go

package nftables
import (
"encoding/binary"
"fmt"
"log/slog"
"net"
"sort"
"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
warned map[string]bool
}
func NewCompiler(cfg *config.Config) *Compiler {
return &Compiler{cfg: cfg}
}
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 {
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"}
}
matches, err := l4Matches(ct.Proto, ct.DPort, nil)
if err != nil {
return fmt.Errorf("conntrack[%d]: %w", i, err)
}
for _, chain := range chains {
for _, m := range matches {
exprs := append([]expr.Any{}, m.exprs...)
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 rule.RateLimit != "" || rule.ConnLimit != "" {
matches, err := l4Matches(proto, dports, sport)
if err != nil {
return fmt.Errorf("rule[%d]: %w", i, err)
}
if len(matches)*specCount(rule.Source, rule.Dest, rule.OrigDest, rule.Action) > 1 {
return fmt.Errorf("rule[%d]: ratelimit/connlimit cannot be combined with proto, port, zone or address 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, rule.OrigDest, 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) {
for chain := range state.Rules {
if chain != "prerouting" {
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 {
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), Total: 1}}
}
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, origDest string, fwZone string, section config.RuleSection) error {
for _, src := range zoneSpecs(srcSpec) {
for _, srcAddr := range splitAddrs(src.Addr) {
for _, od := range splitAddrs(origDest) {
if action == config.RuleDNAT || action == config.RuleRedirect {
if err := c.compileDNATRule(state, tag, src.Zone, srcAddr, od, 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, od, proto,
dports, sports, action, logLevel, fwZone, section); err != nil {
return err
}
}
}
}
}
}
return nil
}
// specCount is how many zone/address combinations compileOneRule expands src and dst into.
func specCount(srcSpec, dstSpec, origDest string, action config.RuleAction) int {
count := func(spec string) (n int) {
for _, z := range zoneSpecs(spec) {
n += len(splitAddrs(z.Addr))
}
return n
}
n := count(srcSpec) * len(splitAddrs(origDest))
if action == config.RuleDNAT || action == config.RuleRedirect {
return n
}
return n * count(dstSpec)
}
// zoneSpecs expands a comma zone list; "all"/"any" forms keep their own comma (exclusion) syntax.
func zoneSpecs(spec string) []config.ZoneSpec {
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)
}
// 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, origDest, 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)
// ponytail: forward daddr is post-DNAT; lift with `ct original daddr` (expr.Ct Direction, google/nftables v0.3.0).
if origDest != "" && chain == "forward" {
return fmt.Errorf("origdest: ORIGDEST on forwarded rules is not supported yet")
}
for _, srcIface := range srcIfaces {
for _, dstIface := range dstIfaces {
matches, err := c.buildMatchExprs(srcIface, dstIface, chain, proto, dports, sports, srcAddr, dstAddr)
if err != nil {
return err
}
if origDest != "" {
od, err := matchOrigDest(origDest)
if err != nil {
return fmt.Errorf("origdest: %w", err)
}
for i := range matches {
matches[i].exprs = append(matches[i].exprs, od...)
}
}
for _, m := range matches {
exprs := m.exprs
if section != "" && section != config.SectionAll {
exprs = append(exprs, matchSection(section)...)
}
if logLevel != "" {
exprs = append(exprs, buildLog(logLevel, 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
}
func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcZone, srcAddr, origDest, dstSpec, proto string,
dports config.PortSpec, action config.RuleAction, logLevel string) error {
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)
var odExprs []expr.Any
if origDest != "" {
var err error
if odExprs, err = matchOrigDest(origDest); err != nil {
return fmt.Errorf("origdest: %w", err)
}
}
matches, err := l4Matches(proto, dports, nil)
if err != nil {
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
if srcIface != "" {
exprs = append(exprs, matchIfaceName(true, srcIface)...)
}
if srcAddr != "" {
src, err := matchSourceCIDR(srcAddr)
if err != nil {
return err
}
exprs = append(exprs, src...)
}
exprs = append(exprs, odExprs...)
exprs = append(exprs, m.exprs...)
if logLevel != "" {
exprs = append(exprs, buildLog(logLevel, tag)...)
}
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, RegisterProtoMax: 1, Flags: unix.NF_NAT_RANGE_PROTO_SPECIFIED},
)
} 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...)
}
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)...)
}
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,
},
)
}
}
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
}
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: expr.ExthdrOpTcpopt,
},
&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: expr.ExthdrOpTcpopt,
},
)
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 ifaces
}
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{""}
}
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 {
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)
}
sort.Strings(zones)
return zones
}
return []string{base}
}
func (c *Compiler) buildMatchExprs(srcIface, dstIface, chain, proto string, dports, sports config.PortSpec, srcAddr, dstAddr string) ([]l4Match, 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...)
}
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 byte
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 != "" {
protos = strings.Split(proto, ",")
}
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)
}
var pm []expr.Any
var n byte
isICMP := false
if p != "" {
var err error
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 {
return nil, fmt.Errorf("dport %v is ambiguous with %s in proto list %q", dports, p, proto)
}
parseD = matchICMPType
}
dalts, err := portAlternatives(dports, parseD)
if err != nil {
return nil, err
}
salts, err := portAlternatives(sports, parseSPortOrRange)
if err != nil {
return nil, err
}
for _, d := range dalts {
for _, sp := range salts {
e := append(append(append([]expr.Any{}, pm...), d...), sp...)
out = append(out, l4Match{proto: n, exprs: e})
}
}
}
return out, 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 == "" {
return nil, fmt.Errorf("empty element in port list %q", item)
}
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.
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]},
}
}
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
}
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 {
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, error) {
if strings.Contains(spec, "/") {
parts := strings.SplitN(spec, "/", 2)
typeVal, ok := resolveICMPType(parts[0])
if !ok {
return nil, fmt.Errorf("invalid icmp type %q", parts[0])
}
code, err := strconv.ParseUint(parts[1], 10, 8)
if err != 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, 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) {
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 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)
}
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 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)
}
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)
}
// matchOrigDest guards the daddr match with the address's nfproto so it is family-correct in the inet table.
func matchOrigDest(addr string) ([]expr.Any, error) {
var proto byte
for i, a := range strings.Split(strings.TrimPrefix(addr, "!"), ",") {
a, _, _ = strings.Cut(a, "/")
p := byte(unix.NFPROTO_IPV6)
if ip := net.ParseIP(a); ip != nil && ip.To4() != nil {
p = unix.NFPROTO_IPV4
}
if i > 0 && p != proto {
return nil, fmt.Errorf("%q mixes IPv4 and IPv6 addresses", addr)
}
proto = p
}
dst, err := matchDestCIDR(addr)
if err != nil {
return nil, err
}
return append([]expr.Any{
&expr.Meta{Key: expr.MetaKeyNFPROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{proto}},
}, dst...), nil
}
func matchAddrCIDR(cidr string, isSrc bool) ([]expr.Any, error) {
negated := false
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
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<<unix.NFTA_LOG_PREFIX | 1<<unix.NFTA_LOG_LEVEL,
Level: nfLevel,
Data: []byte(logPrefix),
},
}
}
func logLevelToNF(level string) expr.LogLevel {
switch strings.ToLower(level) {
case "emerg", "panic":
return expr.LogLevelEmerg
case "alert":
return expr.LogLevelAlert
case "crit":
return expr.LogLevelCrit
case "err", "error":
return expr.LogLevelErr
case "warn", "warning":
return expr.LogLevelWarning
case "notice":
return expr.LogLevelNotice
case "info":
return expr.LogLevelInfo
case "debug":
return expr.LogLevelDebug
default:
return expr.LogLevelWarning
}
}
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}}
case config.RuleDrop:
return []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}
case config.RuleReject:
return rejectExprs(proto, family)
case config.RuleLog:
return nil
case config.RuleContinue:
return nil
case config.RuleMark:
return nil
case config.RuleConnMark:
return nil
case config.RuleCount:
return nil
case config.RuleNoNAT:
return []expr.Any{&expr.Verdict{Kind: expr.VerdictReturn}}
default:
return []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}
}
}
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,
}}
}
return []expr.Any{&expr.Reject{
Type: unix.NFT_REJECT_ICMPX_UNREACH,
Code: unix.NFT_REJECT_ICMPX_PORT_UNREACH,
}}
}
func policyVerdict(action config.PolicyAction, family config.AddressFamily) []expr.Any {
switch action {
case config.PolicyAccept:
return []expr.Any{&expr.Verdict{Kind: expr.VerdictAccept}}
case config.PolicyDrop:
return []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}
case config.PolicyReject:
return rejectExprs(0, family)
default:
return []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}
}
}
func splitZoneSpec(spec string) (zone, addr string) {
idx := strings.IndexByte(spec, ':')
if idx < 0 {
return spec, ""
}
return spec[:idx], spec[idx+1:]
}
func parsePort(s string) (uint16, error) {
n, err := strconv.ParseUint(s, 10, 16)
if err != nil {
return 0, fmt.Errorf("invalid port %q: %w", s, err)
}
return uint16(n), nil
}