Files
tomswall/internal/nftables/compiler.go
T
benvin e0f54ef320 Add tomswall agent (control-plane pull mode)
Add `tomswall agent`: it pulls this device's compiled config from tomswallapi,
differentially applies it, and reports the applied generation. It caches the
last known-good config and, when the control plane is unreachable, keeps
applying that cache — it never fails closed.

- internal/agent: rendered-config types, HTTP client (fetch + status report),
  on-disk cache, on-device DNS resolver for dns sets (honors the device's
  configured resolver, fail-safe on lookup failure), and the pull-apply-report
  loop behind a mockable Applier.
- Translate the interface-agnostic, address-matched rendered model into native
  tomswall config using the "all:<cidr>" any-interface source/dest form, reusing
  the existing differential engine. Named-set members are inlined as concrete
  addresses (native nft set references are a tracked follow-up).
- cmd/tomswall: wire the `agent` subcommand (flags + TOMSWALL_* env, --once).
- Unit tests: translation, cache, and the don't-fail-closed fallback loop.
- Add DESIGN.md documenting the control-plane architecture.
2026-07-20 20:05:49 +10:00

1688 lines
42 KiB
Go

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<<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 string, 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 string, family config.AddressFamily) []expr.Any {
if strings.ToLower(proto) == "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("", 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
}