e0f54ef320
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.
1688 lines
42 KiB
Go
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
|
|
}
|