446 lines
12 KiB
Go
446 lines
12 KiB
Go
package nftables
|
|
|
|
import (
|
|
"fmt"
|
|
|
|
"github.com/google/nftables"
|
|
"github.com/google/nftables/expr"
|
|
"github.com/mdlayher/netlink"
|
|
"golang.org/x/sys/unix"
|
|
|
|
"git.unkin.net/unkin/tomswall/internal/config"
|
|
)
|
|
|
|
type Engine struct {
|
|
cfg *config.Config
|
|
conn *nftables.Conn
|
|
}
|
|
|
|
func NewEngine(cfg *config.Config) (*Engine, error) {
|
|
conn, err := nftables.New(nftables.WithSockOptions(largeBuffers))
|
|
if err != nil {
|
|
return nil, fmt.Errorf("connecting to nftables: %w", err)
|
|
}
|
|
return &Engine{cfg: cfg, conn: conn}, nil
|
|
}
|
|
|
|
var tableFamilies = map[config.AddressFamily]nftables.TableFamily{
|
|
config.FamilyINET: nftables.TableFamilyINet,
|
|
config.FamilyIP: nftables.TableFamilyIPv4,
|
|
config.FamilyIP6: nftables.TableFamilyIPv6,
|
|
}
|
|
|
|
func (e *Engine) family() nftables.TableFamily {
|
|
if f, ok := tableFamilies[e.cfg.Settings.AddressFamily]; ok {
|
|
return f
|
|
}
|
|
return nftables.TableFamilyINet
|
|
}
|
|
|
|
func addressFamily(tf nftables.TableFamily) config.AddressFamily {
|
|
for f, t := range tableFamilies {
|
|
if t == tf {
|
|
return f
|
|
}
|
|
}
|
|
return config.FamilyINET
|
|
}
|
|
|
|
// withFamily is the engine for the same table name in another address family.
|
|
func (e *Engine) withFamily(f config.AddressFamily) *Engine {
|
|
cfg := *e.cfg
|
|
cfg.Settings.AddressFamily = f
|
|
return &Engine{cfg: &cfg, conn: e.conn}
|
|
}
|
|
|
|
// overlaps reports whether tables of families a and b filter the same traffic:
|
|
// inet covers both ip and ip6, which do not overlap each other.
|
|
func overlaps(a, b nftables.TableFamily) bool {
|
|
if a == b {
|
|
return false
|
|
}
|
|
return (a == nftables.TableFamilyINet && (b == nftables.TableFamilyIPv4 || b == nftables.TableFamilyIPv6)) ||
|
|
(b == nftables.TableFamilyINet && (a == nftables.TableFamilyIPv4 || a == nftables.TableFamilyIPv6))
|
|
}
|
|
|
|
// staleTables are our-named tables left by a different address_family that
|
|
// would still filter the traffic this family now owns.
|
|
func (e *Engine) staleTables() ([]*nftables.Table, error) {
|
|
tables, err := e.conn.ListTables()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("listing tables: %w", err)
|
|
}
|
|
var stale []*nftables.Table
|
|
for _, t := range tables {
|
|
if t.Name == e.cfg.Settings.TableName && overlaps(e.family(), t.Family) {
|
|
stale = append(stale, t)
|
|
}
|
|
}
|
|
return stale, nil
|
|
}
|
|
|
|
// batchBufSize bounds one batch: the kernel rejects a batch larger than the
|
|
// send buffer (EMSGSIZE) and drops ACKs beyond the receive buffer (ENOBUFS)
|
|
// after committing it.
|
|
// ponytail: fixed cap of tens of thousands of rules; size per batch if exceeded.
|
|
const batchBufSize = 64 << 20
|
|
|
|
// largeBuffers raises both socket buffers, ignoring rmem_max/wmem_max when
|
|
// CAP_NET_ADMIN allows it and falling back to the capped sizes otherwise.
|
|
func largeBuffers(c *netlink.Conn) error {
|
|
rc, err := c.SyscallConn()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
var serr error
|
|
err = rc.Control(func(fd uintptr) {
|
|
for _, o := range [][2]int{{unix.SO_SNDBUFFORCE, unix.SO_SNDBUF}, {unix.SO_RCVBUFFORCE, unix.SO_RCVBUF}} {
|
|
if unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, o[0], batchBufSize) == nil {
|
|
continue
|
|
}
|
|
if serr = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, o[1], batchBufSize); serr != nil {
|
|
return
|
|
}
|
|
}
|
|
})
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return serr
|
|
}
|
|
|
|
func (e *Engine) ensureTable() *nftables.Table {
|
|
return e.conn.AddTable(&nftables.Table{
|
|
Family: e.family(),
|
|
Name: e.cfg.Settings.TableName,
|
|
})
|
|
}
|
|
|
|
// ensureChains declares the base chains; policies overrides their default policy.
|
|
func (e *Engine) ensureChains(table *nftables.Table, policies map[string]nftables.ChainPolicy) map[string]*nftables.Chain {
|
|
chains := map[string]*nftables.Chain{
|
|
"input": {
|
|
Name: "input",
|
|
Table: table,
|
|
Type: nftables.ChainTypeFilter,
|
|
Hooknum: nftables.ChainHookInput,
|
|
Priority: nftables.ChainPriorityFilter,
|
|
Policy: policyPtr(nftables.ChainPolicyDrop),
|
|
},
|
|
"forward": {
|
|
Name: "forward",
|
|
Table: table,
|
|
Type: nftables.ChainTypeFilter,
|
|
Hooknum: nftables.ChainHookForward,
|
|
Priority: nftables.ChainPriorityFilter,
|
|
Policy: policyPtr(nftables.ChainPolicyDrop),
|
|
},
|
|
"output": {
|
|
Name: "output",
|
|
Table: table,
|
|
Type: nftables.ChainTypeFilter,
|
|
Hooknum: nftables.ChainHookOutput,
|
|
Priority: nftables.ChainPriorityFilter,
|
|
Policy: policyPtr(nftables.ChainPolicyAccept),
|
|
},
|
|
"postrouting": {
|
|
Name: "postrouting",
|
|
Table: table,
|
|
Type: nftables.ChainTypeNAT,
|
|
Hooknum: nftables.ChainHookPostrouting,
|
|
Priority: nftables.ChainPriorityNATSource,
|
|
},
|
|
"prerouting": {
|
|
Name: "prerouting",
|
|
Table: table,
|
|
Type: nftables.ChainTypeNAT,
|
|
Hooknum: nftables.ChainHookPrerouting,
|
|
Priority: nftables.ChainPriorityNATDest,
|
|
},
|
|
"helper_prerouting": {
|
|
Name: "helper_prerouting",
|
|
Table: table,
|
|
Type: nftables.ChainTypeFilter,
|
|
Hooknum: nftables.ChainHookPrerouting,
|
|
Priority: nftables.ChainPriorityMangle,
|
|
},
|
|
"helper_output": {
|
|
Name: "helper_output",
|
|
Table: table,
|
|
Type: nftables.ChainTypeFilter,
|
|
Hooknum: nftables.ChainHookOutput,
|
|
Priority: nftables.ChainPriorityMangle,
|
|
},
|
|
"raw_prerouting": {
|
|
Name: "raw_prerouting",
|
|
Table: table,
|
|
Type: nftables.ChainTypeFilter,
|
|
Hooknum: nftables.ChainHookPrerouting,
|
|
Priority: nftables.ChainPriorityRaw,
|
|
Policy: policyPtr(nftables.ChainPolicyAccept),
|
|
},
|
|
"raw_output": {
|
|
Name: "raw_output",
|
|
Table: table,
|
|
Type: nftables.ChainTypeFilter,
|
|
Hooknum: nftables.ChainHookOutput,
|
|
Priority: nftables.ChainPriorityRaw,
|
|
Policy: policyPtr(nftables.ChainPolicyAccept),
|
|
},
|
|
}
|
|
|
|
for name, chain := range chains {
|
|
if p, ok := policies[name]; ok {
|
|
chain.Policy = policyPtr(p)
|
|
}
|
|
chains[name] = e.conn.AddChain(chain)
|
|
}
|
|
return chains
|
|
}
|
|
|
|
func (e *Engine) Plan() (*ChangeSet, error) {
|
|
compiler := NewCompiler(e.cfg)
|
|
|
|
desired, err := compiler.Compile()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("compiling config: %w", err)
|
|
}
|
|
|
|
current, err := e.readCurrentState()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("reading current state: %w", err)
|
|
}
|
|
|
|
return computeDiff(current, desired), nil
|
|
}
|
|
|
|
func (e *Engine) Apply(changes *ChangeSet) error {
|
|
return e.apply(changes, nil)
|
|
}
|
|
|
|
func (e *Engine) apply(changes *ChangeSet, policies map[string]nftables.ChainPolicy) error {
|
|
stale, err := e.staleTables()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
for _, t := range stale {
|
|
e.conn.DelTable(t)
|
|
}
|
|
table := e.ensureTable()
|
|
chains := e.ensureChains(table, policies)
|
|
|
|
for _, r := range changes.Remove {
|
|
e.conn.DelRule(&nftables.Rule{
|
|
Table: table,
|
|
Chain: chains[r.Chain],
|
|
Handle: r.Handle,
|
|
})
|
|
}
|
|
|
|
for _, n := range changes.RemoveHelpers {
|
|
e.conn.DeleteObject(helperObj(table, n, expr.CtHelper{}))
|
|
}
|
|
for _, h := range changes.AddHelpers {
|
|
e.conn.AddObj(helperObj(table, h.Name, h.Helper))
|
|
}
|
|
|
|
for _, r := range changes.Add {
|
|
chain, ok := chains[r.Chain]
|
|
if !ok {
|
|
return fmt.Errorf("unknown chain %q", r.Chain)
|
|
}
|
|
rule := &nftables.Rule{
|
|
Table: table,
|
|
Chain: chain,
|
|
Exprs: r.Exprs,
|
|
UserData: []byte(r.Tag),
|
|
}
|
|
if r.Before != 0 {
|
|
rule.Position = r.Before
|
|
e.conn.InsertRule(rule)
|
|
} else {
|
|
e.conn.AddRule(rule)
|
|
}
|
|
}
|
|
|
|
return e.conn.Flush()
|
|
}
|
|
|
|
func (e *Engine) Flush() error {
|
|
tables, err := e.staleTables()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
own, err := e.findTable()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if own != nil {
|
|
tables = append(tables, own)
|
|
}
|
|
if len(tables) == 0 {
|
|
return nil
|
|
}
|
|
for _, t := range tables {
|
|
e.conn.DelTable(t)
|
|
}
|
|
return e.conn.Flush()
|
|
}
|
|
|
|
func (e *Engine) findTable() (*nftables.Table, error) {
|
|
tables, err := e.conn.ListTables()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("listing tables: %w", err)
|
|
}
|
|
for _, t := range tables {
|
|
if t.Name == e.cfg.Settings.TableName && t.Family == e.family() {
|
|
return t, nil
|
|
}
|
|
}
|
|
return nil, nil
|
|
}
|
|
|
|
func (e *Engine) readCurrentState() (*FirewallState, error) {
|
|
state := &FirewallState{
|
|
Rules: make(map[string][]ManagedRule),
|
|
}
|
|
|
|
ourTable, err := e.findTable()
|
|
if err != nil || ourTable == nil {
|
|
return state, err
|
|
}
|
|
|
|
objs, err := e.conn.GetNamedObjects(ourTable)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("listing objects: %w", err)
|
|
}
|
|
for _, o := range objs {
|
|
if no, ok := o.(*nftables.NamedObj); ok && no.Type == nftables.ObjTypeCtHelper {
|
|
if h, ok := no.Obj.(*expr.CtHelper); ok {
|
|
state.Helpers = append(state.Helpers, Helper{Name: no.Name, Helper: *h})
|
|
}
|
|
}
|
|
}
|
|
|
|
chains, err := e.conn.ListChainsOfTableFamily(e.family())
|
|
if err != nil {
|
|
return nil, fmt.Errorf("listing chains: %w", err)
|
|
}
|
|
|
|
for _, chain := range chains {
|
|
if chain.Table.Name != e.cfg.Settings.TableName {
|
|
continue
|
|
}
|
|
rules, err := e.conn.GetRules(ourTable, chain)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("listing rules of %s: %w", chain.Name, err)
|
|
}
|
|
for _, rule := range rules {
|
|
state.Rules[chain.Name] = append(state.Rules[chain.Name], ManagedRule{
|
|
Chain: chain.Name,
|
|
Handle: rule.Handle,
|
|
Exprs: rule.Exprs,
|
|
Tag: string(rule.UserData),
|
|
})
|
|
}
|
|
}
|
|
|
|
return state, nil
|
|
}
|
|
|
|
// Snapshot is the tomswall table as captured live, serialisable so a revert
|
|
// survives the process that took it.
|
|
type Snapshot struct {
|
|
Table string `json:"table"`
|
|
Family config.AddressFamily `json:"family,omitempty"`
|
|
Present bool `json:"present"`
|
|
Policies map[string]nftables.ChainPolicy `json:"policies,omitempty"`
|
|
Rules map[string][]SnapshotRule `json:"rules,omitempty"`
|
|
Helpers []Helper `json:"helpers,omitempty"`
|
|
}
|
|
|
|
// SnapshotRule is a managed rule with its expressions in netlink wire format.
|
|
type SnapshotRule struct {
|
|
Tag string `json:"tag"`
|
|
Exprs [][]byte `json:"exprs"`
|
|
}
|
|
|
|
// Snapshot captures the live tomswall table so Restore can roll back to it,
|
|
// falling back to the overlapping table of another family that apply replaces.
|
|
func (e *Engine) Snapshot() (*Snapshot, error) {
|
|
t, err := e.findTable()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if t == nil {
|
|
stale, err := e.staleTables()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
// ponytail: captures one stale table; an inet config replacing both ip and ip6 restores only the first.
|
|
if len(stale) > 0 {
|
|
return e.withFamily(addressFamily(stale[0].Family)).Snapshot()
|
|
}
|
|
}
|
|
snap := &Snapshot{Table: e.cfg.Settings.TableName, Family: addressFamily(e.family())}
|
|
if t == nil {
|
|
return snap, nil
|
|
}
|
|
snap.Present = true
|
|
|
|
chains, err := e.conn.ListChainsOfTableFamily(e.family())
|
|
if err != nil {
|
|
return nil, fmt.Errorf("listing chains: %w", err)
|
|
}
|
|
snap.Policies = make(map[string]nftables.ChainPolicy)
|
|
for _, c := range chains {
|
|
if c.Table.Name == snap.Table && c.Policy != nil {
|
|
snap.Policies[c.Name] = *c.Policy
|
|
}
|
|
}
|
|
|
|
state, err := e.readCurrentState()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
snap.Rules, err = encodeState(state, byte(e.family()))
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
snap.Helpers = state.Helpers
|
|
return snap, nil
|
|
}
|
|
|
|
// Restore atomically returns the tomswall table to the snapshot: rule order
|
|
// and chain policies included, or removed if it was absent.
|
|
func (e *Engine) Restore(s *Snapshot) error {
|
|
if s.Table != e.cfg.Settings.TableName {
|
|
return fmt.Errorf("snapshot is of table %q, engine manages %q", s.Table, e.cfg.Settings.TableName)
|
|
}
|
|
// Snapshots predating the family field are of the inet table.
|
|
if f := addressFamily(tableFamilies[s.Family]); f != addressFamily(e.family()) {
|
|
return e.withFamily(f).Restore(s)
|
|
}
|
|
if !s.Present {
|
|
return e.Flush()
|
|
}
|
|
want, err := decodeState(s.Rules, byte(e.family()))
|
|
if err != nil {
|
|
return err
|
|
}
|
|
want.Helpers = s.Helpers
|
|
current, err := e.readCurrentState()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return e.apply(restoreChangeSet(current, want), s.Policies)
|
|
}
|
|
|
|
func helperObj(table *nftables.Table, name string, h expr.CtHelper) *nftables.NamedObj {
|
|
return &nftables.NamedObj{Table: table, Name: name, Type: nftables.ObjTypeCtHelper, Obj: &h}
|
|
}
|
|
|
|
func policyPtr(p nftables.ChainPolicy) *nftables.ChainPolicy {
|
|
return &p
|
|
}
|