mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-07 05:59:06 +02:00
253 lines
6.1 KiB
Go
253 lines
6.1 KiB
Go
//go:build !android
|
|
|
|
package nftables
|
|
|
|
import (
|
|
"fmt"
|
|
"net/netip"
|
|
|
|
"github.com/google/nftables"
|
|
"github.com/google/nftables/binaryutil"
|
|
"github.com/google/nftables/expr"
|
|
log "github.com/sirupsen/logrus"
|
|
|
|
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
|
)
|
|
|
|
func (r *family) AddInboundDNAT(localAddr netip.Addr, protocol firewall.Protocol, originalPort, translatedPort uint16) error {
|
|
ruleID := firewall.RuleID(fmt.Sprintf("inbound-dnat-%s-%s-%d-%d", localAddr.String(), protocol, originalPort, translatedPort))
|
|
|
|
if _, exists := r.rules[ruleID]; exists {
|
|
return nil
|
|
}
|
|
|
|
protoNum, err := r.af.protoNum(protocol)
|
|
if err != nil {
|
|
return fmt.Errorf("convert protocol to number: %w", err)
|
|
}
|
|
|
|
exprs := []expr.Any{
|
|
&expr.Meta{Key: expr.MetaKeyIIFNAME, Register: 1},
|
|
&expr.Cmp{
|
|
Op: expr.CmpOpEq,
|
|
Register: 1,
|
|
Data: ifname(r.wgIface.Name()),
|
|
},
|
|
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 2},
|
|
&expr.Cmp{
|
|
Op: expr.CmpOpEq,
|
|
Register: 2,
|
|
Data: []byte{protoNum},
|
|
},
|
|
&expr.Payload{
|
|
DestRegister: 3,
|
|
Base: expr.PayloadBaseTransportHeader,
|
|
Offset: 2,
|
|
Len: 2,
|
|
},
|
|
&expr.Cmp{
|
|
Op: expr.CmpOpEq,
|
|
Register: 3,
|
|
Data: binaryutil.BigEndian.PutUint16(originalPort),
|
|
},
|
|
}
|
|
|
|
bits := 32
|
|
if localAddr.Is6() {
|
|
bits = 128
|
|
}
|
|
exprs = append(exprs, prefixMatchExprs(r.af, netip.PrefixFrom(localAddr, bits), false)...)
|
|
|
|
exprs = append(exprs,
|
|
&expr.Immediate{
|
|
Register: 1,
|
|
Data: localAddr.AsSlice(),
|
|
},
|
|
&expr.Immediate{
|
|
Register: 2,
|
|
Data: binaryutil.BigEndian.PutUint16(translatedPort),
|
|
},
|
|
&expr.NAT{
|
|
Type: expr.NATTypeDestNAT,
|
|
Family: uint32(r.af.tableFamily),
|
|
RegAddrMin: 1,
|
|
RegProtoMin: 2,
|
|
RegProtoMax: 0,
|
|
},
|
|
)
|
|
|
|
dnatRule := &nftables.Rule{
|
|
Table: r.workTable,
|
|
Chain: r.chains[chainNameRoutingRdr],
|
|
Exprs: exprs,
|
|
UserData: []byte(ruleID),
|
|
}
|
|
r.conn.AddRule(dnatRule)
|
|
|
|
if err := r.conn.Flush(); err != nil {
|
|
return fmt.Errorf("add inbound DNAT rule: %w", err)
|
|
}
|
|
|
|
r.rules[ruleID] = dnatRule
|
|
|
|
return nil
|
|
}
|
|
|
|
// RemoveInboundDNAT removes an inbound DNAT rule.
|
|
func (r *family) RemoveInboundDNAT(localAddr netip.Addr, protocol firewall.Protocol, originalPort, translatedPort uint16) error {
|
|
if err := r.refreshRulesMap(); err != nil {
|
|
return fmt.Errorf(refreshRulesMapError, err)
|
|
}
|
|
|
|
ruleID := firewall.RuleID(fmt.Sprintf("inbound-dnat-%s-%s-%d-%d", localAddr.String(), protocol, originalPort, translatedPort))
|
|
|
|
rule, exists := r.rules[ruleID]
|
|
if !exists {
|
|
return nil
|
|
}
|
|
|
|
if rule.Handle == 0 {
|
|
log.Warnf("inbound DNAT rule %s has no handle, removing stale entry", ruleID)
|
|
delete(r.rules, ruleID)
|
|
return nil
|
|
}
|
|
|
|
if err := r.conn.DelRule(rule); err != nil {
|
|
return fmt.Errorf("delete inbound DNAT rule %s: %w", ruleID, err)
|
|
}
|
|
if err := r.conn.Flush(); err != nil {
|
|
return fmt.Errorf("flush delete inbound DNAT rule: %w", err)
|
|
}
|
|
delete(r.rules, ruleID)
|
|
|
|
return nil
|
|
}
|
|
|
|
// ensureNATOutputChain lazily creates the OUTPUT NAT chain on first use.
|
|
func (r *family) ensureNATOutputChain() error {
|
|
if _, exists := r.chains[chainNameNATOutput]; exists {
|
|
return nil
|
|
}
|
|
|
|
r.chains[chainNameNATOutput] = r.conn.AddChain(&nftables.Chain{
|
|
Name: chainNameNATOutput,
|
|
Table: r.workTable,
|
|
Hooknum: nftables.ChainHookOutput,
|
|
Priority: nftables.ChainPriorityNATDest,
|
|
Type: nftables.ChainTypeNAT,
|
|
})
|
|
|
|
if err := r.conn.Flush(); err != nil {
|
|
delete(r.chains, chainNameNATOutput)
|
|
return fmt.Errorf("create NAT output chain: %w", err)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// AddOutputDNAT adds an OUTPUT chain DNAT rule for locally-generated traffic.
|
|
func (r *family) AddOutputDNAT(localAddr netip.Addr, protocol firewall.Protocol, originalPort, translatedPort uint16) error {
|
|
ruleID := firewall.RuleID(fmt.Sprintf("output-dnat-%s-%s-%d-%d", localAddr.String(), protocol, originalPort, translatedPort))
|
|
|
|
if _, exists := r.rules[ruleID]; exists {
|
|
return nil
|
|
}
|
|
|
|
if err := r.ensureNATOutputChain(); err != nil {
|
|
return err
|
|
}
|
|
|
|
protoNum, err := r.af.protoNum(protocol)
|
|
if err != nil {
|
|
return fmt.Errorf("convert protocol to number: %w", err)
|
|
}
|
|
|
|
exprs := []expr.Any{
|
|
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
|
|
&expr.Cmp{
|
|
Op: expr.CmpOpEq,
|
|
Register: 1,
|
|
Data: []byte{protoNum},
|
|
},
|
|
&expr.Payload{
|
|
DestRegister: 2,
|
|
Base: expr.PayloadBaseTransportHeader,
|
|
Offset: 2,
|
|
Len: 2,
|
|
},
|
|
&expr.Cmp{
|
|
Op: expr.CmpOpEq,
|
|
Register: 2,
|
|
Data: binaryutil.BigEndian.PutUint16(originalPort),
|
|
},
|
|
}
|
|
|
|
bits := 32
|
|
if localAddr.Is6() {
|
|
bits = 128
|
|
}
|
|
exprs = append(exprs, prefixMatchExprs(r.af, netip.PrefixFrom(localAddr, bits), false)...)
|
|
|
|
exprs = append(exprs,
|
|
&expr.Immediate{
|
|
Register: 1,
|
|
Data: localAddr.AsSlice(),
|
|
},
|
|
&expr.Immediate{
|
|
Register: 2,
|
|
Data: binaryutil.BigEndian.PutUint16(translatedPort),
|
|
},
|
|
&expr.NAT{
|
|
Type: expr.NATTypeDestNAT,
|
|
Family: uint32(r.af.tableFamily),
|
|
RegAddrMin: 1,
|
|
RegProtoMin: 2,
|
|
},
|
|
)
|
|
|
|
dnatRule := &nftables.Rule{
|
|
Table: r.workTable,
|
|
Chain: r.chains[chainNameNATOutput],
|
|
Exprs: exprs,
|
|
UserData: []byte(ruleID),
|
|
}
|
|
r.conn.AddRule(dnatRule)
|
|
|
|
if err := r.conn.Flush(); err != nil {
|
|
return fmt.Errorf("add output DNAT rule: %w", err)
|
|
}
|
|
|
|
r.rules[ruleID] = dnatRule
|
|
|
|
return nil
|
|
}
|
|
|
|
// RemoveOutputDNAT removes an OUTPUT chain DNAT rule.
|
|
func (r *family) RemoveOutputDNAT(localAddr netip.Addr, protocol firewall.Protocol, originalPort, translatedPort uint16) error {
|
|
if err := r.refreshRulesMap(); err != nil {
|
|
return fmt.Errorf(refreshRulesMapError, err)
|
|
}
|
|
|
|
ruleID := firewall.RuleID(fmt.Sprintf("output-dnat-%s-%s-%d-%d", localAddr.String(), protocol, originalPort, translatedPort))
|
|
|
|
rule, exists := r.rules[ruleID]
|
|
if !exists {
|
|
return nil
|
|
}
|
|
|
|
if rule.Handle == 0 {
|
|
log.Warnf("output DNAT rule %s has no handle, removing stale entry", ruleID)
|
|
delete(r.rules, ruleID)
|
|
return nil
|
|
}
|
|
|
|
if err := r.conn.DelRule(rule); err != nil {
|
|
return fmt.Errorf("delete output DNAT rule %s: %w", ruleID, err)
|
|
}
|
|
if err := r.conn.Flush(); err != nil {
|
|
return fmt.Errorf("flush delete output DNAT rule: %w", err)
|
|
}
|
|
delete(r.rules, ruleID)
|
|
|
|
return nil
|
|
}
|