120 lines
3.8 KiB
Go
120 lines
3.8 KiB
Go
package nftables
|
|
|
|
import (
|
|
"encoding/binary"
|
|
"fmt"
|
|
|
|
"github.com/google/nftables"
|
|
"github.com/google/nftables/expr"
|
|
"github.com/mdlayher/netlink"
|
|
"golang.org/x/sys/unix"
|
|
)
|
|
|
|
const inet = byte(nftables.TableFamilyINet)
|
|
|
|
// exprByName mirrors the expression types google/nftables can parse back from the kernel.
|
|
var exprByName = map[string]func() expr.Any{
|
|
"ct": func() expr.Any { return &expr.Ct{} },
|
|
"range": func() expr.Any { return &expr.Range{} },
|
|
"meta": func() expr.Any { return &expr.Meta{} },
|
|
"cmp": func() expr.Any { return &expr.Cmp{} },
|
|
"counter": func() expr.Any { return &expr.Counter{} },
|
|
"objref": func() expr.Any { return &expr.Objref{} },
|
|
"payload": func() expr.Any { return &expr.Payload{} },
|
|
"lookup": func() expr.Any { return &expr.Lookup{} },
|
|
"immediate": func() expr.Any { return &expr.Immediate{} },
|
|
"bitwise": func() expr.Any { return &expr.Bitwise{} },
|
|
"redir": func() expr.Any { return &expr.Redir{} },
|
|
"nat": func() expr.Any { return &expr.NAT{} },
|
|
"limit": func() expr.Any { return &expr.Limit{} },
|
|
"quota": func() expr.Any { return &expr.Quota{} },
|
|
"dynset": func() expr.Any { return &expr.Dynset{} },
|
|
"log": func() expr.Any { return &expr.Log{} },
|
|
"exthdr": func() expr.Any { return &expr.Exthdr{} },
|
|
"connlimit": func() expr.Any { return &expr.Connlimit{} },
|
|
"queue": func() expr.Any { return &expr.Queue{} },
|
|
"flow_offload": func() expr.Any { return &expr.FlowOffload{} },
|
|
"reject": func() expr.Any { return &expr.Reject{} },
|
|
"masq": func() expr.Any { return &expr.Masq{} },
|
|
"hash": func() expr.Any { return &expr.Hash{} },
|
|
"notrack": func() expr.Any { return &expr.Notrack{} },
|
|
}
|
|
|
|
func encodeState(state *FirewallState) (map[string][]SnapshotRule, error) {
|
|
out := make(map[string][]SnapshotRule, len(state.Rules))
|
|
for chain, rules := range state.Rules {
|
|
for _, r := range rules {
|
|
sr := SnapshotRule{Tag: r.Tag}
|
|
for _, e := range r.Exprs {
|
|
b, err := expr.Marshal(inet, e)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("encoding %s rule %q: %w", chain, r.Tag, err)
|
|
}
|
|
sr.Exprs = append(sr.Exprs, b)
|
|
}
|
|
out[chain] = append(out[chain], sr)
|
|
}
|
|
}
|
|
return out, nil
|
|
}
|
|
|
|
func decodeState(rules map[string][]SnapshotRule) (*FirewallState, error) {
|
|
state := &FirewallState{Rules: make(map[string][]ManagedRule, len(rules))}
|
|
for chain, rs := range rules {
|
|
for _, sr := range rs {
|
|
r := ManagedRule{Chain: chain, Tag: sr.Tag}
|
|
for _, b := range sr.Exprs {
|
|
e, err := decodeExpr(b)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("decoding %s rule %q: %w", chain, sr.Tag, err)
|
|
}
|
|
r.Exprs = append(r.Exprs, e)
|
|
}
|
|
state.Rules[chain] = append(state.Rules[chain], r)
|
|
}
|
|
}
|
|
return state, nil
|
|
}
|
|
|
|
// decodeExpr reverses expr.Marshal, as google/nftables does when reading rules.
|
|
func decodeExpr(b []byte) (expr.Any, error) {
|
|
ad, err := netlink.NewAttributeDecoder(b)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
ad.ByteOrder = binary.BigEndian
|
|
var name string
|
|
var data []byte
|
|
for ad.Next() {
|
|
switch ad.Type() {
|
|
case unix.NFTA_EXPR_NAME:
|
|
name = ad.String()
|
|
case unix.NFTA_EXPR_DATA:
|
|
data = ad.Bytes()
|
|
}
|
|
}
|
|
if err := ad.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
newExpr, ok := exprByName[name]
|
|
if !ok {
|
|
return nil, fmt.Errorf("unsupported expression %q", name)
|
|
}
|
|
e := newExpr()
|
|
if name == "notrack" {
|
|
return e, nil
|
|
}
|
|
if err := expr.Unmarshal(inet, data, e); err != nil {
|
|
return nil, err
|
|
}
|
|
// A verdict is an immediate into the verdict register with no data.
|
|
if imm, ok := e.(*expr.Immediate); ok && imm.Register == unix.NFT_REG_VERDICT && len(imm.Data) == 0 {
|
|
v := &expr.Verdict{}
|
|
if err := expr.Unmarshal(inet, data, v); err != nil {
|
|
return nil, err
|
|
}
|
|
return v, nil
|
|
}
|
|
return e, nil
|
|
}
|