1936 lines
52 KiB
Go
1936 lines
52 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 {
|
|
fwZone := c.cfg.FirewallZone()
|
|
for i, ct := range c.cfg.Conntrack {
|
|
tag := fmt.Sprintf("conntrack:%d", i)
|
|
if config.HasZoneExclusion(ct.Source) || config.HasZoneExclusion(ct.Dest) {
|
|
return fmt.Errorf("conntrack[%d]: zone exclusions are not supported in conntrack entries", i)
|
|
}
|
|
srcs, dsts := c.zoneSpecs(ct.Source), c.zoneSpecs(ct.Dest)
|
|
if len(srcs) == 0 {
|
|
srcs = []config.ZoneSpec{{}}
|
|
}
|
|
if len(dsts) == 0 {
|
|
dsts = []config.ZoneSpec{{}}
|
|
}
|
|
|
|
for _, src := range srcs {
|
|
if src.Zone == "all" || src.Zone == "any" {
|
|
src.Zone = ""
|
|
}
|
|
chains := []string{"raw_prerouting"}
|
|
switch {
|
|
case ct.Chain == config.ConntrackOutput && src.Zone != fwZone && src.Zone != "":
|
|
return fmt.Errorf("conntrack[%d]: chain output needs SOURCE %s, got %q", i, fwZone, src.Zone)
|
|
case ct.Chain == config.ConntrackPrerouting && src.Zone == fwZone:
|
|
return fmt.Errorf("conntrack[%d]: SOURCE %s cannot use chain prerouting", i, fwZone)
|
|
case ct.Chain != config.ConntrackPrerouting && src.Zone == fwZone, ct.Chain == config.ConntrackOutput:
|
|
chains = []string{"raw_output"}
|
|
case ct.Chain == config.ConntrackBoth && src.Zone == "":
|
|
chains = []string{"raw_prerouting", "raw_output"}
|
|
}
|
|
for _, srcAddr := range splitAddrs(src.Addr) {
|
|
for _, dst := range dsts {
|
|
for _, dstAddr := range splitAddrs(dst.Addr) {
|
|
for _, chain := range chains {
|
|
if err := c.compileConntrackPair(state, tag, chain, ct, src.Zone, srcAddr, dst.Zone, dstAddr); err != nil {
|
|
return fmt.Errorf("conntrack[%d]: %w", i, err)
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// compileConntrackPair matches iif of the source zone in raw_prerouting and oif of the dest zone in raw_output.
|
|
func (c *Compiler) compileConntrackPair(state *FirewallState, tag, chain string, ct config.ConntrackRule,
|
|
srcZone, srcAddr, dstZone, dstAddr string) error {
|
|
if ct.Action == config.ConntrackHelper {
|
|
return nil
|
|
}
|
|
if _, ok := c.cfg.Zones[dstZone]; ok && chain == "raw_prerouting" &&
|
|
(dstAddr == "" || strings.HasPrefix(dstAddr, "!")) {
|
|
return fmt.Errorf("conntrack DEST zone %q needs an address in prerouting", dstZone)
|
|
}
|
|
srcIfaces, dstIfaces := c.resolveZoneInterfaces(srcZone, srcAddr), []string{""}
|
|
if chain == "raw_prerouting" && c.resolveZoneInterfaces(dstZone, dstAddr) == nil {
|
|
return nil
|
|
}
|
|
if chain == "raw_output" {
|
|
srcIfaces, dstIfaces = []string{""}, c.resolveZoneInterfaces(dstZone, dstAddr)
|
|
}
|
|
for _, srcIface := range srcIfaces {
|
|
for _, dstIface := range dstIfaces {
|
|
matches, err := c.buildMatchExprs(srcIface, dstIface, chain, ct.Proto, ct.DPort, ct.SPort, srcAddr, dstAddr)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
for _, m := range matches {
|
|
exprs := m.exprs
|
|
switch ct.Action {
|
|
case config.ConntrackNoTrack:
|
|
exprs = append(exprs, &expr.Notrack{})
|
|
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)*c.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 c.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 c.zoneSpecs(dstSpec) {
|
|
// Exclusion expansion never pairs fw with itself, and pairs a zone with itself only for "all+".
|
|
if src.Zone == dst.Zone && (isZoneExclusion(srcSpec) || isZoneExclusion(dstSpec)) &&
|
|
(src.Zone == fwZone || !strings.Contains(srcSpec, "+!") && !strings.Contains(dstSpec, "+!")) {
|
|
continue
|
|
}
|
|
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 (c *Compiler) specCount(srcSpec, dstSpec, origDest string, action config.RuleAction) int {
|
|
count := func(spec string) (n int) {
|
|
for _, z := range c.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" stay global and "all!x,y" becomes every zone but x and y.
|
|
func (c *Compiler) zoneSpecs(spec string) []config.ZoneSpec {
|
|
zone, addr := splitZoneSpec(spec)
|
|
if !isZoneExclusion(zone) {
|
|
if base := strings.TrimSuffix(zone, "+"); base == "all" || base == "any" {
|
|
return []config.ZoneSpec{{Zone: zone, Addr: addr}}
|
|
}
|
|
return config.SplitZoneList(spec)
|
|
}
|
|
var out []config.ZoneSpec
|
|
for _, z := range c.expandZoneRef(zone) {
|
|
out = append(out, config.ZoneSpec{Zone: z, Addr: addr})
|
|
}
|
|
return out
|
|
}
|
|
|
|
func isZoneExclusion(spec string) bool {
|
|
base, _, ok := strings.Cut(spec, "!")
|
|
base = strings.TrimSuffix(base, "+")
|
|
return ok && (base == "all" || base == "any")
|
|
}
|
|
|
|
// 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, srcAddr)
|
|
dstIfaces := c.resolveZoneInterfaces(dstZone, dstAddr)
|
|
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, srcAddr)
|
|
|
|
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
|
|
natExpr.Specified = true
|
|
}
|
|
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
|
|
natExpr.Specified = true
|
|
}
|
|
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 != "input" {
|
|
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"
|
|
}
|
|
|
|
// resolveZoneInterfaces returns nil (fail closed) for an unknown zone, or one with no interfaces unless a non-negated address match narrows the rule.
|
|
func (c *Compiler) resolveZoneInterfaces(zone, addr string) []string {
|
|
switch zone {
|
|
case "", "all", "all+", "any", "any+":
|
|
return []string{""}
|
|
}
|
|
z, ok := c.cfg.Zones[zone]
|
|
if !ok {
|
|
slog.Warn("compiler: unknown zone, skipping its rules", "zone", zone)
|
|
return nil
|
|
}
|
|
if z.Type == config.ZoneFirewall {
|
|
return []string{""}
|
|
}
|
|
if ifaces := c.cfg.ZoneInterfaces(zone); len(ifaces) > 0 {
|
|
return ifaces
|
|
}
|
|
if addr != "" && !strings.HasPrefix(addr, "!") {
|
|
return []string{""}
|
|
}
|
|
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
|
|
}
|
|
|
|
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+" || base == "any" || base == "any+" {
|
|
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 != "input" {
|
|
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
|
|
}
|