mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-11 16:09:07 +02:00
Merge branch 'main' into feature/gui-quit-daemon-down
This commit is contained in:
@@ -121,6 +121,7 @@ type Manager struct {
|
|||||||
udpTracker *conntrack.UDPTracker
|
udpTracker *conntrack.UDPTracker
|
||||||
icmpTracker *conntrack.ICMPTracker
|
icmpTracker *conntrack.ICMPTracker
|
||||||
tcpTracker *conntrack.TCPTracker
|
tcpTracker *conntrack.TCPTracker
|
||||||
|
fragments *fragmentTracker
|
||||||
forwarder atomic.Pointer[forwarder.Forwarder]
|
forwarder atomic.Pointer[forwarder.Forwarder]
|
||||||
pendingCapture atomic.Pointer[forwarder.PacketCapture]
|
pendingCapture atomic.Pointer[forwarder.PacketCapture]
|
||||||
logger *nblog.Logger
|
logger *nblog.Logger
|
||||||
@@ -183,6 +184,41 @@ func (d *decoder) decodePacket(data []byte) error {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// decodeTransport decodes the transport header of a first fragment (which
|
||||||
|
// gopacket leaves undecoded) into the decoder and appends its layer type to
|
||||||
|
// decoded, so the ACL pipeline can evaluate it like a normal packet. It returns
|
||||||
|
// false if the protocol is unsupported or the header is truncated.
|
||||||
|
func (d *decoder) decodeTransport(proto layers.IPProtocol, payload []byte) bool {
|
||||||
|
var l4 gopacket.DecodingLayer
|
||||||
|
var layerType gopacket.LayerType
|
||||||
|
var minLen int
|
||||||
|
switch proto {
|
||||||
|
case layers.IPProtocolTCP:
|
||||||
|
l4, layerType, minLen = &d.tcp, layers.LayerTypeTCP, 20
|
||||||
|
case layers.IPProtocolUDP:
|
||||||
|
l4, layerType, minLen = &d.udp, layers.LayerTypeUDP, 8
|
||||||
|
case layers.IPProtocolICMPv4:
|
||||||
|
l4, layerType, minLen = &d.icmp4, layers.LayerTypeICMPv4, 8
|
||||||
|
case layers.IPProtocolICMPv6:
|
||||||
|
l4, layerType, minLen = &d.icmp6, layers.LayerTypeICMPv6, 8
|
||||||
|
default:
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reject a fragment too small to hold the full transport header before
|
||||||
|
// decoding: it can't be ACL-evaluated (tiny-fragment attack), and skipping
|
||||||
|
// the decode avoids gopacket allocating an error on the drop path.
|
||||||
|
if len(payload) < minLen {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := l4.DecodeFromBytes(payload, gopacket.NilDecodeFeedback); err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
d.decoded = append(d.decoded, layerType)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
// Create userspace firewall manager constructor
|
// Create userspace firewall manager constructor
|
||||||
func Create(iface common.IFaceMapper, disableServerRoutes bool, flowLogger nftypes.FlowLogger, mtu uint16) (*Manager, error) {
|
func Create(iface common.IFaceMapper, disableServerRoutes bool, flowLogger nftypes.FlowLogger, mtu uint16) (*Manager, error) {
|
||||||
return create(iface, nil, disableServerRoutes, flowLogger, mtu)
|
return create(iface, nil, disableServerRoutes, flowLogger, mtu)
|
||||||
@@ -286,6 +322,8 @@ func create(iface common.IFaceMapper, nativeFirewall firewall.Manager, disableSe
|
|||||||
if err := m.localipmanager.UpdateLocalIPs(iface); err != nil {
|
if err := m.localipmanager.UpdateLocalIPs(iface); err != nil {
|
||||||
return nil, fmt.Errorf("update local IPs: %w", err)
|
return nil, fmt.Errorf("update local IPs: %w", err)
|
||||||
}
|
}
|
||||||
|
m.fragments = newFragmentTracker(m.logger)
|
||||||
|
|
||||||
if disableConntrack {
|
if disableConntrack {
|
||||||
log.Info("conntrack is disabled")
|
log.Info("conntrack is disabled")
|
||||||
} else {
|
} else {
|
||||||
@@ -299,6 +337,7 @@ func create(iface common.IFaceMapper, nativeFirewall firewall.Manager, disableSe
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
if err := iface.SetFilter(m); err != nil {
|
if err := iface.SetFilter(m); err != nil {
|
||||||
|
m.fragments.Close()
|
||||||
return nil, fmt.Errorf("set filter: %w", err)
|
return nil, fmt.Errorf("set filter: %w", err)
|
||||||
}
|
}
|
||||||
return m, nil
|
return m, nil
|
||||||
@@ -694,6 +733,10 @@ func (m *Manager) resetState() {
|
|||||||
m.tcpTracker.Close()
|
m.tcpTracker.Close()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if m.fragments != nil {
|
||||||
|
m.fragments.Close()
|
||||||
|
}
|
||||||
|
|
||||||
if fwder := m.forwarder.Load(); fwder != nil {
|
if fwder := m.forwarder.Load(); fwder != nil {
|
||||||
fwder.SetCapture(nil)
|
fwder.SetCapture(nil)
|
||||||
fwder.Stop()
|
fwder.Stop()
|
||||||
@@ -1046,19 +1089,20 @@ func (m *Manager) filterInbound(packetData []byte, size int) bool {
|
|||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
// TODO: pass fragments of routed packets to forwarder
|
// gopacket does not decode the transport header of any IP fragment, so
|
||||||
|
// fragments take a dedicated path: the first fragment's header is decoded
|
||||||
|
// and ACL-evaluated here, and the remaining fragments inherit its verdict.
|
||||||
if fragment {
|
if fragment {
|
||||||
if m.logger.Enabled(nblog.LevelTrace) {
|
return m.filterInboundFragment(d, srcIP, dstIP, size)
|
||||||
if d.decoded[0] == layers.LayerTypeIPv4 {
|
|
||||||
m.logger.Trace4("packet is a fragment: src=%v dst=%v id=%v flags=%v",
|
|
||||||
srcIP, dstIP, d.ip4.Id, d.ip4.Flags)
|
|
||||||
} else {
|
|
||||||
m.logger.Trace2("packet is an IPv6 fragment: src=%v dst=%v", srcIP, dstIP)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return false
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
return m.filterInboundDecoded(d, srcIP, dstIP, packetData, size)
|
||||||
|
}
|
||||||
|
|
||||||
|
// filterInboundDecoded runs the ACL, DNAT and conntrack pipeline on a fully
|
||||||
|
// decoded (non-fragment) inbound packet. It returns true if the packet should
|
||||||
|
// be dropped.
|
||||||
|
func (m *Manager) filterInboundDecoded(d *decoder, srcIP, dstIP netip.Addr, packetData []byte, size int) bool {
|
||||||
// TODO: optimize port DNAT by caching matched rules in conntrack
|
// TODO: optimize port DNAT by caching matched rules in conntrack
|
||||||
if translated := m.translateInboundPortDNAT(packetData, d, srcIP, dstIP); translated {
|
if translated := m.translateInboundPortDNAT(packetData, d, srcIP, dstIP); translated {
|
||||||
// Re-decode after port DNAT translation to update port information
|
// Re-decode after port DNAT translation to update port information
|
||||||
@@ -1089,33 +1133,226 @@ func (m *Manager) filterInbound(packetData []byte, size int) bool {
|
|||||||
return m.handleRoutedTraffic(d, srcIP, dstIP, packetData, size)
|
return m.handleRoutedTraffic(d, srcIP, dstIP, packetData, size)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// fragmentMeta holds the reassembly identity and layout of an IP fragment,
|
||||||
|
// extracted uniformly for IPv4 and IPv6.
|
||||||
|
type fragmentMeta struct {
|
||||||
|
key fragmentKey
|
||||||
|
// offset is the fragment offset in 8-byte units (zero for the first
|
||||||
|
// fragment).
|
||||||
|
offset uint16
|
||||||
|
// moreFragments is the More Fragments bit. A first fragment with it unset is
|
||||||
|
// an IPv6 atomic fragment (a complete datagram, RFC 6946): it has no trailing
|
||||||
|
// fragments to inherit a verdict, so it must not be recorded.
|
||||||
|
moreFragments bool
|
||||||
|
proto layers.IPProtocol
|
||||||
|
// l4payload is the fragmentable payload of this fragment. For the first
|
||||||
|
// fragment it starts with the transport header.
|
||||||
|
l4payload []byte
|
||||||
|
// headerEndOctets is the first fragment's payload length in 8-byte units:
|
||||||
|
// the smallest offset a trailing fragment may start at without overlapping
|
||||||
|
// the inspected transport header.
|
||||||
|
headerEndOctets uint16
|
||||||
|
}
|
||||||
|
|
||||||
|
// fragmentMetadata extracts the fragment identity and layout from a decoded IP
|
||||||
|
// fragment. It returns false for fragments it can't interpret (e.g. an IPv6
|
||||||
|
// fragment header shorter than 8 bytes), which are then dropped.
|
||||||
|
func fragmentMetadata(d *decoder, srcIP, dstIP netip.Addr) (fragmentMeta, bool) {
|
||||||
|
switch d.decoded[0] {
|
||||||
|
case layers.LayerTypeIPv4:
|
||||||
|
payload := d.ip4.Payload
|
||||||
|
return fragmentMeta{
|
||||||
|
key: fragmentKey{srcIP: srcIP, dstIP: dstIP, id: uint32(d.ip4.Id), proto: uint8(d.ip4.Protocol)},
|
||||||
|
offset: d.ip4.FragOffset,
|
||||||
|
moreFragments: d.ip4.Flags&layers.IPv4MoreFragments != 0,
|
||||||
|
proto: d.ip4.Protocol,
|
||||||
|
l4payload: payload,
|
||||||
|
headerEndOctets: octets(len(payload)),
|
||||||
|
}, true
|
||||||
|
|
||||||
|
case layers.LayerTypeIPv6:
|
||||||
|
// IPv6 fragment extension header: 8 bytes, followed by the fragmentable
|
||||||
|
// payload. Layout: next header (1), reserved (1), offset+flags (2), id (4).
|
||||||
|
payload := d.ip6.Payload
|
||||||
|
if len(payload) < 8 {
|
||||||
|
return fragmentMeta{}, false
|
||||||
|
}
|
||||||
|
nextHeader := layers.IPProtocol(payload[0])
|
||||||
|
offsetFlags := binary.BigEndian.Uint16(payload[2:4])
|
||||||
|
id := binary.BigEndian.Uint32(payload[4:8])
|
||||||
|
l4 := payload[8:]
|
||||||
|
return fragmentMeta{
|
||||||
|
key: fragmentKey{srcIP: srcIP, dstIP: dstIP, id: id, proto: uint8(nextHeader)},
|
||||||
|
offset: offsetFlags >> 3,
|
||||||
|
moreFragments: offsetFlags&1 != 0,
|
||||||
|
proto: nextHeader,
|
||||||
|
l4payload: l4,
|
||||||
|
headerEndOctets: octets(len(l4)),
|
||||||
|
}, true
|
||||||
|
|
||||||
|
default:
|
||||||
|
return fragmentMeta{}, false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// octets rounds a byte length up to whole 8-byte units, the granularity of the
|
||||||
|
// IP fragment offset field.
|
||||||
|
func octets(nbytes int) uint16 {
|
||||||
|
return uint16((nbytes + 7) / 8)
|
||||||
|
}
|
||||||
|
|
||||||
|
// filterInboundFragment decides the fate of an IP fragment. gopacket stops
|
||||||
|
// decoding at the network layer for every fragment, so the first fragment's
|
||||||
|
// transport header is decoded and ACL-evaluated here and its verdict recorded;
|
||||||
|
// the remaining (headerless) fragments inherit that verdict. Anything that
|
||||||
|
// cannot be tied to an allowed, non-overlapping first fragment is dropped.
|
||||||
|
func (m *Manager) filterInboundFragment(d *decoder, srcIP, dstIP netip.Addr, size int) bool {
|
||||||
|
meta, ok := fragmentMetadata(d, srcIP, dstIP)
|
||||||
|
if !ok {
|
||||||
|
if m.logger.Enabled(nblog.LevelTrace) {
|
||||||
|
m.logger.Trace2("dropping unsupported fragment: src=%v dst=%v", srcIP, dstIP)
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
if meta.offset != 0 {
|
||||||
|
return m.filterTrailingFragment(meta, srcIP, dstIP)
|
||||||
|
}
|
||||||
|
|
||||||
|
// A new first fragment supersedes any recorded verdict for this datagram, so
|
||||||
|
// a re-sent or overlapping offset-zero fragment can't inherit the old one.
|
||||||
|
m.fragments.poison(meta.key)
|
||||||
|
|
||||||
|
// First fragment: decode its transport header so the ACL can evaluate it. A
|
||||||
|
// decode failure means the fragment is too small to hold the full transport
|
||||||
|
// header (RFC 1858 §3 tiny-fragment attack); it can't be evaluated, so drop it.
|
||||||
|
if !d.decodeTransport(meta.proto, meta.l4payload) {
|
||||||
|
if m.logger.Enabled(nblog.LevelTrace) {
|
||||||
|
m.logger.Trace3("dropping first fragment without full L4 header: src=%v dst=%v id=%v",
|
||||||
|
srcIP, dstIP, meta.key.id)
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
return m.filterFirstFragment(d, meta, srcIP, dstIP, size)
|
||||||
|
}
|
||||||
|
|
||||||
|
// filterTrailingFragment applies a recorded first-fragment verdict to a
|
||||||
|
// non-first fragment.
|
||||||
|
func (m *Manager) filterTrailingFragment(meta fragmentMeta, srcIP, dstIP netip.Addr) bool {
|
||||||
|
switch m.fragments.verdict(meta.key, meta.offset) {
|
||||||
|
case fragmentAllow:
|
||||||
|
return false
|
||||||
|
case fragmentOverlap:
|
||||||
|
if m.logger.Enabled(nblog.LevelTrace) {
|
||||||
|
m.logger.Trace3("dropping overlapping fragment rewriting inspected header: src=%v dst=%v id=%v",
|
||||||
|
srcIP, dstIP, meta.key.id)
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
default:
|
||||||
|
if m.logger.Enabled(nblog.LevelTrace) {
|
||||||
|
m.logger.Trace3("dropping fragment with no allowed first fragment: src=%v dst=%v id=%v",
|
||||||
|
srcIP, dstIP, meta.key.id)
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// filterFirstFragment runs the verdict part of the inbound pipeline on a first
|
||||||
|
// fragment with its transport header decoded. It mirrors filterInboundDecoded
|
||||||
|
// but skips DNAT (port rewriting on fragments is unsupported) and forwarder
|
||||||
|
// injection (fragments are left to the stack to reassemble, not forwarded).
|
||||||
|
// Allowed fragments have their verdict recorded so the datagram's trailing
|
||||||
|
// fragments inherit it.
|
||||||
|
func (m *Manager) filterFirstFragment(d *decoder, meta fragmentMeta, srcIP, dstIP netip.Addr, size int) bool {
|
||||||
|
if m.stateful && m.isValidTrackedConnection(d, srcIP, dstIP, size) {
|
||||||
|
m.recordFirstFragment(meta)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
if m.localipmanager.IsLocalIP(dstIP) {
|
||||||
|
ruleID, blocked := m.peerACLsBlock(srcIP, d, nil)
|
||||||
|
if blocked {
|
||||||
|
m.storeDropFlow("Dropping local first fragment (ACL denied): rule_id=%s proto=%v src=%s:%d dst=%s:%d",
|
||||||
|
d, srcIP, dstIP, ruleID, size)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
m.trackInbound(d, srcIP, dstIP, ruleID, size)
|
||||||
|
m.recordFirstFragment(meta)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
if !m.routingEnabled.Load() {
|
||||||
|
if m.logger.Enabled(nblog.LevelTrace) {
|
||||||
|
m.logger.Trace2("Dropping routed fragment (routing disabled): src=%s dst=%s", srcIP, dstIP)
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
if m.nativeRouter.Load() {
|
||||||
|
m.trackInbound(d, srcIP, dstIP, nil, size)
|
||||||
|
m.recordFirstFragment(meta)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// TODO: pass fragments of routed packets to the forwarder; until then
|
||||||
|
// allowed routed fragments go to the native stack.
|
||||||
|
srcPort, dstPort := getPortsFromPacket(d)
|
||||||
|
ruleID, pass := m.routeACLsPass(srcIP, dstIP, d.decoded[1], srcPort, dstPort)
|
||||||
|
if !pass {
|
||||||
|
m.storeDropFlow("Dropping routed first fragment (ACL denied): rule_id=%s proto=%v src=%s:%d dst=%s:%d",
|
||||||
|
d, srcIP, dstIP, ruleID, size)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
m.recordFirstFragment(meta)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// recordFirstFragment caches an allowed first fragment's verdict for its
|
||||||
|
// trailing fragments to inherit. Atomic fragments (no More Fragments bit) are
|
||||||
|
// complete datagrams with no trailing fragments, so they are not cached and
|
||||||
|
// cannot exhaust the verdict table.
|
||||||
|
func (m *Manager) recordFirstFragment(meta fragmentMeta) {
|
||||||
|
if !meta.moreFragments {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
m.fragments.recordAllowed(meta.key, meta.headerEndOctets)
|
||||||
|
}
|
||||||
|
|
||||||
|
// storeDropFlow logs and records a netflow drop event for an inbound packet
|
||||||
|
// denied by the ACLs. msg is the trace format taking rule id, protocol, source
|
||||||
|
// and destination.
|
||||||
|
func (m *Manager) storeDropFlow(msg string, d *decoder, srcIP, dstIP netip.Addr, ruleID []byte, size int) {
|
||||||
|
pnum := getProtocolFromPacket(d)
|
||||||
|
srcPort, dstPort := getPortsFromPacket(d)
|
||||||
|
|
||||||
|
if m.logger.Enabled(nblog.LevelTrace) {
|
||||||
|
m.logger.Trace6(msg, ruleID, pnum, srcIP, srcPort, dstIP, dstPort)
|
||||||
|
}
|
||||||
|
|
||||||
|
m.flowLogger.StoreEvent(nftypes.EventFields{
|
||||||
|
FlowID: uuid.New(),
|
||||||
|
Type: nftypes.TypeDrop,
|
||||||
|
RuleID: ruleID,
|
||||||
|
Direction: nftypes.Ingress,
|
||||||
|
Protocol: pnum,
|
||||||
|
SourceIP: srcIP,
|
||||||
|
DestIP: dstIP,
|
||||||
|
SourcePort: srcPort,
|
||||||
|
DestPort: dstPort,
|
||||||
|
// TODO: icmp type/code
|
||||||
|
RxPackets: 1,
|
||||||
|
RxBytes: uint64(size),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
// handleLocalTraffic handles local traffic.
|
// handleLocalTraffic handles local traffic.
|
||||||
// If it returns true, the packet should be dropped.
|
// If it returns true, the packet should be dropped.
|
||||||
func (m *Manager) handleLocalTraffic(d *decoder, srcIP, dstIP netip.Addr, packetData []byte, size int) bool {
|
func (m *Manager) handleLocalTraffic(d *decoder, srcIP, dstIP netip.Addr, packetData []byte, size int) bool {
|
||||||
ruleID, blocked := m.peerACLsBlock(srcIP, d, packetData)
|
ruleID, blocked := m.peerACLsBlock(srcIP, d, packetData)
|
||||||
if blocked {
|
if blocked {
|
||||||
pnum := getProtocolFromPacket(d)
|
m.storeDropFlow("Dropping local packet (ACL denied): rule_id=%s proto=%v src=%s:%d dst=%s:%d",
|
||||||
srcPort, dstPort := getPortsFromPacket(d)
|
d, srcIP, dstIP, ruleID, size)
|
||||||
|
|
||||||
if m.logger.Enabled(nblog.LevelTrace) {
|
|
||||||
m.logger.Trace6("Dropping local packet (ACL denied): rule_id=%s proto=%v src=%s:%d dst=%s:%d",
|
|
||||||
ruleID, pnum, srcIP, srcPort, dstIP, dstPort)
|
|
||||||
}
|
|
||||||
|
|
||||||
m.flowLogger.StoreEvent(nftypes.EventFields{
|
|
||||||
FlowID: uuid.New(),
|
|
||||||
Type: nftypes.TypeDrop,
|
|
||||||
RuleID: ruleID,
|
|
||||||
Direction: nftypes.Ingress,
|
|
||||||
Protocol: pnum,
|
|
||||||
SourceIP: srcIP,
|
|
||||||
DestIP: dstIP,
|
|
||||||
SourcePort: srcPort,
|
|
||||||
DestPort: dstPort,
|
|
||||||
// TODO: icmp type/code
|
|
||||||
RxPackets: 1,
|
|
||||||
RxBytes: uint64(size),
|
|
||||||
})
|
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1168,27 +1405,8 @@ func (m *Manager) handleRoutedTraffic(d *decoder, srcIP, dstIP netip.Addr, packe
|
|||||||
|
|
||||||
ruleID, pass := m.routeACLsPass(srcIP, dstIP, protoLayer, srcPort, dstPort)
|
ruleID, pass := m.routeACLsPass(srcIP, dstIP, protoLayer, srcPort, dstPort)
|
||||||
if !pass {
|
if !pass {
|
||||||
proto := getProtocolFromPacket(d)
|
m.storeDropFlow("Dropping routed packet (ACL denied): rule_id=%s proto=%v src=%s:%d dst=%s:%d",
|
||||||
|
d, srcIP, dstIP, ruleID, size)
|
||||||
if m.logger.Enabled(nblog.LevelTrace) {
|
|
||||||
m.logger.Trace6("Dropping routed packet (ACL denied): rule_id=%s proto=%v src=%s:%d dst=%s:%d",
|
|
||||||
ruleID, proto, srcIP, srcPort, dstIP, dstPort)
|
|
||||||
}
|
|
||||||
|
|
||||||
m.flowLogger.StoreEvent(nftypes.EventFields{
|
|
||||||
FlowID: uuid.New(),
|
|
||||||
Type: nftypes.TypeDrop,
|
|
||||||
RuleID: ruleID,
|
|
||||||
Direction: nftypes.Ingress,
|
|
||||||
Protocol: proto,
|
|
||||||
SourceIP: srcIP,
|
|
||||||
DestIP: dstIP,
|
|
||||||
SourcePort: srcPort,
|
|
||||||
DestPort: dstPort,
|
|
||||||
// TODO: icmp type/code
|
|
||||||
RxPackets: 1,
|
|
||||||
RxBytes: uint64(size),
|
|
||||||
})
|
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -5,7 +5,9 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"os"
|
||||||
"runtime"
|
"runtime"
|
||||||
|
"strconv"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -31,6 +33,11 @@ const (
|
|||||||
defaultMaxInFlight = 1024
|
defaultMaxInFlight = 1024
|
||||||
iosReceiveWindow = 16384
|
iosReceiveWindow = 16384
|
||||||
iosMaxInFlight = 256
|
iosMaxInFlight = 256
|
||||||
|
|
||||||
|
// envForceTCPRACK overrides the platform default for gVisor's RACK loss
|
||||||
|
// detection. Set to a truthy value to force RACK on, or a falsy value to
|
||||||
|
// force it off, on any platform.
|
||||||
|
envForceTCPRACK = "NB_FORCE_TCP_RACK"
|
||||||
)
|
)
|
||||||
|
|
||||||
type Forwarder struct {
|
type Forwarder struct {
|
||||||
@@ -152,6 +159,8 @@ func New(iface common.IFaceMapper, logger *nblog.Logger, flowLogger nftypes.Flow
|
|||||||
maxInFlight = iosMaxInFlight
|
maxInFlight = iosMaxInFlight
|
||||||
}
|
}
|
||||||
|
|
||||||
|
configureTCPRecovery(s)
|
||||||
|
|
||||||
tcpForwarder := tcp.NewForwarder(s, receiveWindow, maxInFlight, f.handleTCP)
|
tcpForwarder := tcp.NewForwarder(s, receiveWindow, maxInFlight, f.handleTCP)
|
||||||
s.SetTransportProtocolHandler(tcp.ProtocolNumber, tcpForwarder.HandlePacket)
|
s.SetTransportProtocolHandler(tcp.ProtocolNumber, tcpForwarder.HandlePacket)
|
||||||
|
|
||||||
@@ -466,3 +475,31 @@ func probeRawICMP(network, addr string, logger *nblog.Logger) bool {
|
|||||||
logger.Debug1("forwarder: raw %s socket access available", network)
|
logger.Debug1("forwarder: raw %s socket access available", network)
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// configureTCPRecovery disables gVisor's RACK loss detection on Windows, where
|
||||||
|
// it interacts poorly with the host and collapses throughput on routed TCP
|
||||||
|
// connections (gVisor issue #9778). Other platforms keep the default. The
|
||||||
|
// EnvForceTCPRACK environment variable overrides the platform default.
|
||||||
|
func configureTCPRecovery(s *stack.Stack) {
|
||||||
|
disableRACK := runtime.GOOS == "windows"
|
||||||
|
|
||||||
|
if val := os.Getenv(envForceTCPRACK); val != "" {
|
||||||
|
force, err := strconv.ParseBool(val)
|
||||||
|
if err != nil {
|
||||||
|
log.Warnf("parse %s: %v", envForceTCPRACK, err)
|
||||||
|
} else {
|
||||||
|
disableRACK = !force
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !disableRACK {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
opt := tcpip.TCPRecovery(0)
|
||||||
|
if err := s.SetTransportProtocolOption(tcp.ProtocolNumber, &opt); err != nil {
|
||||||
|
log.Warnf("disable TCP RACK loss detection: %v", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
log.Info("forwarder: TCP RACK loss detection disabled")
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,204 @@
|
|||||||
|
package uspfilter
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net/netip"
|
||||||
|
"os"
|
||||||
|
"strconv"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
nblog "github.com/netbirdio/netbird/client/firewall/uspfilter/log"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
// defaultFragmentTimeout bounds how long a first-fragment verdict is kept
|
||||||
|
// while the remaining fragments arrive. It mirrors the Linux IP reassembly
|
||||||
|
// timeout (net.ipv4.ipfrag_time).
|
||||||
|
defaultFragmentTimeout = 30 * time.Second
|
||||||
|
// fragmentCleanupInterval is how often expired verdicts are purged.
|
||||||
|
fragmentCleanupInterval = 10 * time.Second
|
||||||
|
// defaultMaxFragmentEntries caps the number of concurrently tracked
|
||||||
|
// fragmented datagrams. The table stays bounded because each datagram is a
|
||||||
|
// single small entry regardless of how many fragments it is split into, and
|
||||||
|
// the 13-bit IPv4 fragment-offset field limits any datagram to 64 KiB.
|
||||||
|
defaultMaxFragmentEntries = 16384
|
||||||
|
|
||||||
|
// EnvFragmentMaxEntries overrides defaultMaxFragmentEntries.
|
||||||
|
EnvFragmentMaxEntries = "NB_FRAGMENT_MAX_ENTRIES"
|
||||||
|
)
|
||||||
|
|
||||||
|
// fragmentVerdict is the decision for a trailing (headerless) fragment.
|
||||||
|
type fragmentVerdict int
|
||||||
|
|
||||||
|
const (
|
||||||
|
// fragmentDeny drops the fragment: no allowed first fragment is on record.
|
||||||
|
fragmentDeny fragmentVerdict = iota
|
||||||
|
// fragmentAllow passes the fragment: it belongs to an allowed datagram and
|
||||||
|
// does not overlap the already-inspected transport header.
|
||||||
|
fragmentAllow
|
||||||
|
// fragmentOverlap drops the fragment and poisons its datagram: it overlaps
|
||||||
|
// the transport header the ACL inspected (RFC 1858 §4, RFC 3128; RFC 5722
|
||||||
|
// requires discarding the whole datagram on overlap for IPv6).
|
||||||
|
fragmentOverlap
|
||||||
|
)
|
||||||
|
|
||||||
|
// fragmentKey identifies a fragmented datagram. It matches the RFC 791 / RFC
|
||||||
|
// 8200 reassembly key: source, destination, protocol and identification. The id
|
||||||
|
// is 32-bit to hold both the IPv4 (16-bit) and IPv6 (32-bit) identification.
|
||||||
|
type fragmentKey struct {
|
||||||
|
srcIP netip.Addr
|
||||||
|
dstIP netip.Addr
|
||||||
|
id uint32
|
||||||
|
proto uint8
|
||||||
|
}
|
||||||
|
|
||||||
|
// fragmentEntry records the verdict of an allowed first fragment.
|
||||||
|
type fragmentEntry struct {
|
||||||
|
// headerEndOctets is the offset, in 8-byte units, at which the first
|
||||||
|
// fragment's payload ended. A trailing fragment starting before this
|
||||||
|
// overlaps bytes the ACL already inspected and is rejected.
|
||||||
|
headerEndOctets uint16
|
||||||
|
// recordedAt is when the first fragment was accepted. The verdict expires a
|
||||||
|
// fixed timeout later and is not refreshed, mirroring the kernel reassembly
|
||||||
|
// timer so a trailing-fragment flood can't keep a datagram alive.
|
||||||
|
recordedAt time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
// fragmentTracker records the ACL verdict of a datagram's first fragment so the
|
||||||
|
// remaining fragments, which carry no L4 header, can inherit the decision
|
||||||
|
// without reassembling the datagram. Only allowed first fragments are stored;
|
||||||
|
// anything that cannot be tied to an allowed, non-overlapping first fragment is
|
||||||
|
// dropped (fail closed).
|
||||||
|
type fragmentTracker struct {
|
||||||
|
logger *nblog.Logger
|
||||||
|
mutex sync.Mutex
|
||||||
|
entries map[fragmentKey]fragmentEntry
|
||||||
|
timeout time.Duration
|
||||||
|
// maxEntries caps the table; atCapacity dedups the capacity warning until
|
||||||
|
// the table drains below the cap again.
|
||||||
|
maxEntries int
|
||||||
|
atCapacity bool
|
||||||
|
cleanupTicker *time.Ticker
|
||||||
|
cancel context.CancelFunc
|
||||||
|
}
|
||||||
|
|
||||||
|
func newFragmentTracker(logger *nblog.Logger) *fragmentTracker {
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
t := &fragmentTracker{
|
||||||
|
logger: logger,
|
||||||
|
entries: make(map[fragmentKey]fragmentEntry),
|
||||||
|
timeout: defaultFragmentTimeout,
|
||||||
|
maxEntries: fragmentMaxEntries(logger),
|
||||||
|
cleanupTicker: time.NewTicker(fragmentCleanupInterval),
|
||||||
|
cancel: cancel,
|
||||||
|
}
|
||||||
|
go t.cleanupRoutine(ctx)
|
||||||
|
return t
|
||||||
|
}
|
||||||
|
|
||||||
|
func fragmentMaxEntries(logger *nblog.Logger) int {
|
||||||
|
v := os.Getenv(EnvFragmentMaxEntries)
|
||||||
|
if v == "" {
|
||||||
|
return defaultMaxFragmentEntries
|
||||||
|
}
|
||||||
|
n, err := strconv.Atoi(v)
|
||||||
|
if err != nil || n <= 0 {
|
||||||
|
logger.Warn2("invalid %s=%q, using default", EnvFragmentMaxEntries, v)
|
||||||
|
return defaultMaxFragmentEntries
|
||||||
|
}
|
||||||
|
return n
|
||||||
|
}
|
||||||
|
|
||||||
|
// recordAllowed stores the verdict of an allowed first fragment. headerEndOctets
|
||||||
|
// is the first fragment's payload length in 8-byte units. When the table is full
|
||||||
|
// the record is dropped, which fails closed: the datagram's trailing fragments
|
||||||
|
// will be denied.
|
||||||
|
func (t *fragmentTracker) recordAllowed(key fragmentKey, headerEndOctets uint16) {
|
||||||
|
t.mutex.Lock()
|
||||||
|
defer t.mutex.Unlock()
|
||||||
|
|
||||||
|
if t.entries == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if _, ok := t.entries[key]; !ok && len(t.entries) >= t.maxEntries {
|
||||||
|
if !t.atCapacity {
|
||||||
|
t.atCapacity = true
|
||||||
|
t.logger.Warn2("fragment verdict table at capacity (%d/%d): trailing fragments of new datagrams will be dropped",
|
||||||
|
len(t.entries), t.maxEntries)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
t.entries[key] = fragmentEntry{
|
||||||
|
headerEndOctets: headerEndOctets,
|
||||||
|
recordedAt: time.Now(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// poison drops any recorded verdict for a datagram, so its later fragments are
|
||||||
|
// denied until a new allowed first fragment is recorded. Called on every
|
||||||
|
// offset-zero fragment to defeat offset-zero overlap rewrites (RFC 3128).
|
||||||
|
func (t *fragmentTracker) poison(key fragmentKey) {
|
||||||
|
t.mutex.Lock()
|
||||||
|
defer t.mutex.Unlock()
|
||||||
|
delete(t.entries, key)
|
||||||
|
}
|
||||||
|
|
||||||
|
// verdict decides the fate of a trailing fragment at fragOffsetOctets (the IPv4
|
||||||
|
// fragment offset, in 8-byte units). A fragment overlapping the inspected
|
||||||
|
// header poisons the datagram: the entry is removed so all further fragments of
|
||||||
|
// that datagram are denied too.
|
||||||
|
func (t *fragmentTracker) verdict(key fragmentKey, fragOffsetOctets uint16) fragmentVerdict {
|
||||||
|
t.mutex.Lock()
|
||||||
|
defer t.mutex.Unlock()
|
||||||
|
|
||||||
|
entry, ok := t.entries[key]
|
||||||
|
if !ok {
|
||||||
|
return fragmentDeny
|
||||||
|
}
|
||||||
|
if time.Since(entry.recordedAt) > t.timeout {
|
||||||
|
delete(t.entries, key)
|
||||||
|
return fragmentDeny
|
||||||
|
}
|
||||||
|
if fragOffsetOctets < entry.headerEndOctets {
|
||||||
|
delete(t.entries, key)
|
||||||
|
return fragmentOverlap
|
||||||
|
}
|
||||||
|
return fragmentAllow
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *fragmentTracker) cleanupRoutine(ctx context.Context) {
|
||||||
|
defer t.cleanupTicker.Stop()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-t.cleanupTicker.C:
|
||||||
|
t.cleanup()
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *fragmentTracker) cleanup() {
|
||||||
|
t.mutex.Lock()
|
||||||
|
defer t.mutex.Unlock()
|
||||||
|
|
||||||
|
for key, entry := range t.entries {
|
||||||
|
if time.Since(entry.recordedAt) > t.timeout {
|
||||||
|
delete(t.entries, key)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(t.entries) < t.maxEntries {
|
||||||
|
t.atCapacity = false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close stops the cleanup routine and releases resources.
|
||||||
|
func (t *fragmentTracker) Close() {
|
||||||
|
t.cancel()
|
||||||
|
|
||||||
|
t.mutex.Lock()
|
||||||
|
t.entries = nil
|
||||||
|
t.mutex.Unlock()
|
||||||
|
}
|
||||||
@@ -0,0 +1,115 @@
|
|||||||
|
package uspfilter
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// benchFilterInbound drives filterInbound over a fixed packet in a tight loop.
|
||||||
|
// Packets are built once, outside the timed region, so the benchmark measures
|
||||||
|
// only pipeline cost, which is what an attacker can amplify.
|
||||||
|
func benchFilterInbound(b *testing.B, pkt []byte) {
|
||||||
|
b.Helper()
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.SetBytes(int64(len(pkt)))
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
m := benchManager
|
||||||
|
m.filterInbound(pkt, len(pkt))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// benchManager is a package-level manager reused across fragment benchmarks so
|
||||||
|
// setup cost stays out of the timed region.
|
||||||
|
var benchManager *Manager
|
||||||
|
|
||||||
|
func setupBenchManager(b *testing.B) *Manager {
|
||||||
|
b.Helper()
|
||||||
|
m := newFragmentTestManager(b)
|
||||||
|
allowUDP(b, m, 8080)
|
||||||
|
// Disable conntrack so the allowed-first-fragment path measures transport
|
||||||
|
// decode + ACL every iteration instead of matching the connection tracked
|
||||||
|
// on the first iteration.
|
||||||
|
m.stateful = false
|
||||||
|
benchManager = m
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkInbound_NormalPacket is the baseline: a full, non-fragmented UDP
|
||||||
|
// packet that passes the ACL. Fragment paths should stay comparable to this.
|
||||||
|
func BenchmarkInbound_NormalPacket(b *testing.B) {
|
||||||
|
setupBenchManager(b)
|
||||||
|
pkt := normalUDPPacket(b, 8080, 32)
|
||||||
|
benchFilterInbound(b, pkt)
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkInbound_FirstFragmentAllowed measures the first-fragment path:
|
||||||
|
// transport decode + ACL evaluation + verdict record.
|
||||||
|
func BenchmarkInbound_FirstFragmentAllowed(b *testing.B) {
|
||||||
|
setupBenchManager(b)
|
||||||
|
pkt := firstFragmentUDP(b, 0x2000, 8080, 32)
|
||||||
|
benchFilterInbound(b, pkt)
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkInbound_TrailingFragmentAllowed measures the common trailing-fragment
|
||||||
|
// path: a single map lookup after the first fragment is on record.
|
||||||
|
func BenchmarkInbound_TrailingFragmentAllowed(b *testing.B) {
|
||||||
|
m := setupBenchManager(b)
|
||||||
|
first := firstFragmentUDP(b, 0x3000, 8080, 32)
|
||||||
|
m.filterInbound(first, len(first))
|
||||||
|
pkt := trailingFragment(b, 0x3000, 5, false, 24)
|
||||||
|
benchFilterInbound(b, pkt)
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkInbound_TrailingFragmentNoFirst is the primary DoS vector: an
|
||||||
|
// attacker floods trailing fragments with no first fragment on record. Each is
|
||||||
|
// a map miss and must be cheap.
|
||||||
|
func BenchmarkInbound_TrailingFragmentNoFirst(b *testing.B) {
|
||||||
|
setupBenchManager(b)
|
||||||
|
pkt := trailingFragment(b, 0x4000, 185, false, 40)
|
||||||
|
benchFilterInbound(b, pkt)
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkInbound_TinyFirstFragment measures the tiny-fragment drop path: a
|
||||||
|
// first fragment too small to decode a transport header.
|
||||||
|
func BenchmarkInbound_TinyFirstFragment(b *testing.B) {
|
||||||
|
setupBenchManager(b)
|
||||||
|
pkt := trailingFragment(b, 0x5000, 0, true, 4)
|
||||||
|
benchFilterInbound(b, pkt)
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkInbound_TrailingFragmentDistinctIDs is the worst case for the
|
||||||
|
// verdict table: an attacker varies the datagram id on every packet so no first
|
||||||
|
// fragment ever matches. Verdict lookups always miss and nothing is recorded,
|
||||||
|
// so the table cannot grow. Each iteration rewrites the id field in place.
|
||||||
|
func BenchmarkInbound_TrailingFragmentDistinctIDs(b *testing.B) {
|
||||||
|
setupBenchManager(b)
|
||||||
|
pkt := trailingFragment(b, 0x6000, 185, false, 40)
|
||||||
|
m := benchManager
|
||||||
|
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.SetBytes(int64(len(pkt)))
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
// IPv4 identification field is at bytes 4:6.
|
||||||
|
binary.BigEndian.PutUint16(pkt[4:6], uint16(i))
|
||||||
|
m.filterInbound(pkt, len(pkt))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkInbound_FirstFragmentDistinctIDs measures sustained first-fragment
|
||||||
|
// pressure with distinct ids: transport decode + ACL + verdict insert until the
|
||||||
|
// table caps, exercising the map growth and capacity guard.
|
||||||
|
func BenchmarkInbound_FirstFragmentDistinctIDs(b *testing.B) {
|
||||||
|
setupBenchManager(b)
|
||||||
|
pkt := firstFragmentUDP(b, 0x7000, 8080, 32)
|
||||||
|
m := benchManager
|
||||||
|
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.SetBytes(int64(len(pkt)))
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
binary.BigEndian.PutUint16(pkt[4:6], uint16(i))
|
||||||
|
m.filterInbound(pkt, len(pkt))
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,554 @@
|
|||||||
|
package uspfilter
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"net"
|
||||||
|
"net/netip"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/google/gopacket"
|
||||||
|
"github.com/google/gopacket/layers"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
fw "github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
|
nbiface "github.com/netbirdio/netbird/client/iface"
|
||||||
|
"github.com/netbirdio/netbird/client/iface/device"
|
||||||
|
"github.com/netbirdio/netbird/client/iface/wgaddr"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
fragTestSrc = "100.10.0.1"
|
||||||
|
fragTestDst = "100.10.0.100"
|
||||||
|
fragTestSrcV6 = "fd00::1"
|
||||||
|
fragTestDstV6 = "fd00::100"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newFragmentTestManager(tb testing.TB) *Manager {
|
||||||
|
tb.Helper()
|
||||||
|
|
||||||
|
ifaceMock := &IFaceMock{
|
||||||
|
SetFilterFunc: func(device.PacketFilter) error { return nil },
|
||||||
|
AddressFunc: func() wgaddr.Address {
|
||||||
|
return wgaddr.Address{
|
||||||
|
IP: netip.MustParseAddr(fragTestDst),
|
||||||
|
Network: netip.MustParsePrefix("100.10.0.0/16"),
|
||||||
|
IPv6: netip.MustParseAddr(fragTestDstV6),
|
||||||
|
IPv6Net: netip.MustParsePrefix("fd00::/64"),
|
||||||
|
}
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
m, err := Create(ifaceMock, false, flowLogger, nbiface.DefaultMTU)
|
||||||
|
require.NoError(tb, err)
|
||||||
|
require.NoError(tb, m.UpdateLocalIPs())
|
||||||
|
tb.Cleanup(func() { require.NoError(tb, m.Close(nil)) })
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
|
||||||
|
// firstFragmentUDPTo builds the first fragment of a fragmented UDP datagram to
|
||||||
|
// the given destination: it carries the full UDP header plus payloadLen bytes
|
||||||
|
// of data, with the More Fragments flag set and offset zero.
|
||||||
|
func firstFragmentUDPTo(tb testing.TB, dst string, id uint16, dstPort uint16, payloadLen int) []byte {
|
||||||
|
tb.Helper()
|
||||||
|
|
||||||
|
ip := &layers.IPv4{
|
||||||
|
Version: 4,
|
||||||
|
TTL: 64,
|
||||||
|
Id: id,
|
||||||
|
Protocol: layers.IPProtocolUDP,
|
||||||
|
SrcIP: net.ParseIP(fragTestSrc),
|
||||||
|
DstIP: net.ParseIP(dst),
|
||||||
|
Flags: layers.IPv4MoreFragments,
|
||||||
|
}
|
||||||
|
udp := &layers.UDP{SrcPort: 40000, DstPort: layers.UDPPort(dstPort)}
|
||||||
|
require.NoError(tb, udp.SetNetworkLayerForChecksum(ip))
|
||||||
|
|
||||||
|
buf := gopacket.NewSerializeBuffer()
|
||||||
|
opts := gopacket.SerializeOptions{ComputeChecksums: true, FixLengths: true}
|
||||||
|
require.NoError(tb, gopacket.SerializeLayers(buf, opts, ip, udp, gopacket.Payload(make([]byte, payloadLen))))
|
||||||
|
return buf.Bytes()
|
||||||
|
}
|
||||||
|
|
||||||
|
func firstFragmentUDP(tb testing.TB, id uint16, dstPort uint16, payloadLen int) []byte {
|
||||||
|
tb.Helper()
|
||||||
|
return firstFragmentUDPTo(tb, fragTestDst, id, dstPort, payloadLen)
|
||||||
|
}
|
||||||
|
|
||||||
|
// firstFragmentTCP builds the first fragment of a fragmented TCP datagram: the
|
||||||
|
// full 20-byte TCP header plus 12 bytes of data, with the More Fragments flag
|
||||||
|
// set and offset zero.
|
||||||
|
func firstFragmentTCP(tb testing.TB, id uint16, dstPort uint16) []byte {
|
||||||
|
tb.Helper()
|
||||||
|
|
||||||
|
ip := &layers.IPv4{
|
||||||
|
Version: 4,
|
||||||
|
TTL: 64,
|
||||||
|
Id: id,
|
||||||
|
Protocol: layers.IPProtocolTCP,
|
||||||
|
SrcIP: net.ParseIP(fragTestSrc),
|
||||||
|
DstIP: net.ParseIP(fragTestDst),
|
||||||
|
Flags: layers.IPv4MoreFragments,
|
||||||
|
}
|
||||||
|
tcp := &layers.TCP{SrcPort: 40000, DstPort: layers.TCPPort(dstPort), SYN: true, Window: 64240}
|
||||||
|
require.NoError(tb, tcp.SetNetworkLayerForChecksum(ip))
|
||||||
|
|
||||||
|
buf := gopacket.NewSerializeBuffer()
|
||||||
|
opts := gopacket.SerializeOptions{ComputeChecksums: true, FixLengths: true}
|
||||||
|
require.NoError(tb, gopacket.SerializeLayers(buf, opts, ip, tcp, gopacket.Payload(make([]byte, 12))))
|
||||||
|
return buf.Bytes()
|
||||||
|
}
|
||||||
|
|
||||||
|
// trailingFragmentTo builds a non-first fragment to the given destination: an
|
||||||
|
// IPv4 header at the given fragment offset (in 8-byte units) carrying raw
|
||||||
|
// payload and no L4 header.
|
||||||
|
func trailingFragmentTo(tb testing.TB, dst string, proto layers.IPProtocol, id uint16, fragOffsetOctets uint16, moreFragments bool, payloadLen int) []byte {
|
||||||
|
tb.Helper()
|
||||||
|
|
||||||
|
ip := &layers.IPv4{
|
||||||
|
Version: 4,
|
||||||
|
TTL: 64,
|
||||||
|
Id: id,
|
||||||
|
Protocol: proto,
|
||||||
|
SrcIP: net.ParseIP(fragTestSrc),
|
||||||
|
DstIP: net.ParseIP(dst),
|
||||||
|
FragOffset: fragOffsetOctets,
|
||||||
|
}
|
||||||
|
if moreFragments {
|
||||||
|
ip.Flags = layers.IPv4MoreFragments
|
||||||
|
}
|
||||||
|
|
||||||
|
buf := gopacket.NewSerializeBuffer()
|
||||||
|
opts := gopacket.SerializeOptions{FixLengths: true}
|
||||||
|
require.NoError(tb, gopacket.SerializeLayers(buf, opts, ip, gopacket.Payload(make([]byte, payloadLen))))
|
||||||
|
return buf.Bytes()
|
||||||
|
}
|
||||||
|
|
||||||
|
func trailingFragment(tb testing.TB, id uint16, fragOffsetOctets uint16, moreFragments bool, payloadLen int) []byte {
|
||||||
|
tb.Helper()
|
||||||
|
return trailingFragmentTo(tb, fragTestDst, layers.IPProtocolUDP, id, fragOffsetOctets, moreFragments, payloadLen)
|
||||||
|
}
|
||||||
|
|
||||||
|
// outboundUDPPacket builds a complete outbound UDP packet from the local
|
||||||
|
// address, used to establish conntrack state for reply-direction tests.
|
||||||
|
func outboundUDPPacket(tb testing.TB, srcPort, dstPort uint16) []byte {
|
||||||
|
tb.Helper()
|
||||||
|
|
||||||
|
ip := &layers.IPv4{
|
||||||
|
Version: 4,
|
||||||
|
TTL: 64,
|
||||||
|
Id: 1,
|
||||||
|
Protocol: layers.IPProtocolUDP,
|
||||||
|
SrcIP: net.ParseIP(fragTestDst),
|
||||||
|
DstIP: net.ParseIP(fragTestSrc),
|
||||||
|
}
|
||||||
|
udp := &layers.UDP{SrcPort: layers.UDPPort(srcPort), DstPort: layers.UDPPort(dstPort)}
|
||||||
|
require.NoError(tb, udp.SetNetworkLayerForChecksum(ip))
|
||||||
|
|
||||||
|
buf := gopacket.NewSerializeBuffer()
|
||||||
|
opts := gopacket.SerializeOptions{ComputeChecksums: true, FixLengths: true}
|
||||||
|
require.NoError(tb, gopacket.SerializeLayers(buf, opts, ip, udp, gopacket.Payload(make([]byte, 16))))
|
||||||
|
return buf.Bytes()
|
||||||
|
}
|
||||||
|
|
||||||
|
// normalUDPPacket builds a complete, non-fragmented UDP packet for baseline
|
||||||
|
// comparisons against the fragment paths.
|
||||||
|
func normalUDPPacket(tb testing.TB, dstPort uint16, payloadLen int) []byte {
|
||||||
|
tb.Helper()
|
||||||
|
|
||||||
|
ip := &layers.IPv4{
|
||||||
|
Version: 4,
|
||||||
|
TTL: 64,
|
||||||
|
Id: 1,
|
||||||
|
Protocol: layers.IPProtocolUDP,
|
||||||
|
SrcIP: net.ParseIP(fragTestSrc),
|
||||||
|
DstIP: net.ParseIP(fragTestDst),
|
||||||
|
}
|
||||||
|
udp := &layers.UDP{SrcPort: 40000, DstPort: layers.UDPPort(dstPort)}
|
||||||
|
require.NoError(tb, udp.SetNetworkLayerForChecksum(ip))
|
||||||
|
|
||||||
|
buf := gopacket.NewSerializeBuffer()
|
||||||
|
opts := gopacket.SerializeOptions{ComputeChecksums: true, FixLengths: true}
|
||||||
|
require.NoError(tb, gopacket.SerializeLayers(buf, opts, ip, udp, gopacket.Payload(make([]byte, payloadLen))))
|
||||||
|
return buf.Bytes()
|
||||||
|
}
|
||||||
|
|
||||||
|
func allowUDP(tb testing.TB, m *Manager, dstPort uint16) {
|
||||||
|
tb.Helper()
|
||||||
|
_, err := m.AddPeerFiltering(nil, net.ParseIP(fragTestSrc), fw.ProtocolUDP, nil,
|
||||||
|
&fw.Port{Values: []uint16{dstPort}}, fw.ActionAccept, "")
|
||||||
|
require.NoError(tb, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_TrailingWithoutFirstDropped is the core bypass repro: a trailing
|
||||||
|
// fragment with no allowed first fragment on record must be dropped. Before the
|
||||||
|
// fix, filterInbound returned false (allow) for any fragment.
|
||||||
|
func TestFragment_TrailingWithoutFirstDropped(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
|
||||||
|
frag := trailingFragment(t, 0x1234, 185, false, 40)
|
||||||
|
require.True(t, m.filterInbound(frag, len(frag)),
|
||||||
|
"trailing fragment without an allowed first fragment must be dropped")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_AllowedFirstPassesTrailing verifies that once a first fragment
|
||||||
|
// passes the ACL, its trailing fragments inherit the allow verdict.
|
||||||
|
func TestFragment_AllowedFirstPassesTrailing(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
allowUDP(t, m, 8080)
|
||||||
|
|
||||||
|
// First fragment: UDP header (8) + 32 payload = 40 octets -> headerEnd = 5.
|
||||||
|
first := firstFragmentUDP(t, 0x2222, 8080, 32)
|
||||||
|
require.False(t, m.filterInbound(first, len(first)),
|
||||||
|
"allowed first fragment should pass and be recorded")
|
||||||
|
|
||||||
|
trailing := trailingFragment(t, 0x2222, 5, false, 24)
|
||||||
|
require.False(t, m.filterInbound(trailing, len(trailing)),
|
||||||
|
"trailing fragment of an allowed datagram should pass")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_DeniedFirstDropsTrailing verifies that a first fragment blocked
|
||||||
|
// by the ACL leaves no verdict, so its trailing fragments are dropped.
|
||||||
|
func TestFragment_DeniedFirstDropsTrailing(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
// No accept rule: local traffic defaults to deny.
|
||||||
|
|
||||||
|
first := firstFragmentUDP(t, 0x3333, 9999, 32)
|
||||||
|
require.True(t, m.filterInbound(first, len(first)),
|
||||||
|
"first fragment to a blocked port should be dropped by the ACL")
|
||||||
|
|
||||||
|
trailing := trailingFragment(t, 0x3333, 5, false, 24)
|
||||||
|
require.True(t, m.filterInbound(trailing, len(trailing)),
|
||||||
|
"trailing fragment of a denied datagram must be dropped")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_OverlappingHeaderDropped covers the RFC 1858 §4 / RFC 3128
|
||||||
|
// overlapping-fragment rewrite: a trailing fragment starting inside the range
|
||||||
|
// the ACL already inspected is dropped and poisons the datagram. TCP is used so
|
||||||
|
// the overlap lands on real header bytes (the flags at byte 13).
|
||||||
|
func TestFragment_OverlappingHeaderDropped(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
_, err := m.AddPeerFiltering(nil, net.ParseIP(fragTestSrc), fw.ProtocolTCP, nil,
|
||||||
|
&fw.Port{Values: []uint16{8080}}, fw.ActionAccept, "")
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// First fragment: TCP header (20) + 12 data = 32 bytes -> headerEnd = 4 octets.
|
||||||
|
first := firstFragmentTCP(t, 0x4444, 8080)
|
||||||
|
require.False(t, m.filterInbound(first, len(first)))
|
||||||
|
|
||||||
|
// Overlapping fragment at offset 1 (byte 8) falls inside the inspected TCP
|
||||||
|
// header, so it could rewrite the flags or port on reassembly.
|
||||||
|
overlap := trailingFragmentTo(t, fragTestDst, layers.IPProtocolTCP, 0x4444, 1, true, 32)
|
||||||
|
require.True(t, m.filterInbound(overlap, len(overlap)),
|
||||||
|
"fragment overlapping the inspected header must be dropped")
|
||||||
|
|
||||||
|
// The datagram is now poisoned: a later, non-overlapping fragment is also
|
||||||
|
// dropped because the verdict was removed.
|
||||||
|
later := trailingFragmentTo(t, fragTestDst, layers.IPProtocolTCP, 0x4444, 4, false, 24)
|
||||||
|
require.True(t, m.filterInbound(later, len(later)),
|
||||||
|
"fragments after an overlap must be dropped (datagram poisoned)")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_OffsetZeroOverlapPoisons covers the RFC 3128 offset-zero rewrite:
|
||||||
|
// an allowed first fragment followed by a denied offset-zero fragment for the
|
||||||
|
// same datagram must not leave the earlier allow verdict in place.
|
||||||
|
func TestFragment_OffsetZeroOverlapPoisons(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
allowUDP(t, m, 8080)
|
||||||
|
|
||||||
|
allowed := firstFragmentUDP(t, 0x5A5A, 8080, 32)
|
||||||
|
require.False(t, m.filterInbound(allowed, len(allowed)),
|
||||||
|
"allowed first fragment should pass and be recorded")
|
||||||
|
|
||||||
|
// A second offset-zero fragment to a denied port supersedes the datagram's
|
||||||
|
// verdict; it is dropped and must not leave the allow in place.
|
||||||
|
denied := firstFragmentUDP(t, 0x5A5A, 9999, 32)
|
||||||
|
require.True(t, m.filterInbound(denied, len(denied)),
|
||||||
|
"denied offset-zero fragment must be dropped")
|
||||||
|
|
||||||
|
trailing := trailingFragment(t, 0x5A5A, 5, false, 24)
|
||||||
|
require.True(t, m.filterInbound(trailing, len(trailing)),
|
||||||
|
"trailing fragment must be denied after the datagram was poisoned")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_TinyFirstDropped covers the tiny-fragment attack: a first
|
||||||
|
// fragment too small to contain the full transport header can't be
|
||||||
|
// ACL-evaluated and must be dropped.
|
||||||
|
func TestFragment_TinyFirstDropped(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
allowUDP(t, m, 8080)
|
||||||
|
|
||||||
|
// IPv4 header + 4 raw bytes, MF set, offset 0: too small for the 8-byte UDP
|
||||||
|
// header, so it decodes to L3 only.
|
||||||
|
tiny := trailingFragment(t, 0x5555, 0, true, 4)
|
||||||
|
require.True(t, m.filterInbound(tiny, len(tiny)),
|
||||||
|
"tiny first fragment without a full L4 header must be dropped")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_TCPFirstFragment verifies the TCP arm of the transport decode: a
|
||||||
|
// first fragment carrying the full 20-byte TCP header is ACL-evaluated and its
|
||||||
|
// trailing fragments inherit the verdict.
|
||||||
|
func TestFragment_TCPFirstFragment(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
_, err := m.AddPeerFiltering(nil, net.ParseIP(fragTestSrc), fw.ProtocolTCP, nil,
|
||||||
|
&fw.Port{Values: []uint16{8080}}, fw.ActionAccept, "")
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// TCP header (20) + 12 data = 32 bytes -> headerEnd = 4 octets.
|
||||||
|
first := firstFragmentTCP(t, 0x6666, 8080)
|
||||||
|
require.False(t, m.filterInbound(first, len(first)),
|
||||||
|
"allowed TCP first fragment should pass and be recorded")
|
||||||
|
|
||||||
|
trailing := trailingFragmentTo(t, fragTestDst, layers.IPProtocolTCP, 0x6666, 4, false, 24)
|
||||||
|
require.False(t, m.filterInbound(trailing, len(trailing)),
|
||||||
|
"trailing fragment of an allowed TCP datagram should pass")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_TCPTinyFirstDropped verifies the TCP minimum header length: 12
|
||||||
|
// bytes would satisfy a UDP header but falls short of the 20-byte TCP header.
|
||||||
|
func TestFragment_TCPTinyFirstDropped(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
_, err := m.AddPeerFiltering(nil, net.ParseIP(fragTestSrc), fw.ProtocolTCP, nil,
|
||||||
|
&fw.Port{Values: []uint16{8080}}, fw.ActionAccept, "")
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
tiny := trailingFragmentTo(t, fragTestDst, layers.IPProtocolTCP, 0x7777, 0, true, 12)
|
||||||
|
require.True(t, m.filterInbound(tiny, len(tiny)),
|
||||||
|
"first fragment shorter than the TCP header must be dropped")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_ConntrackAllowsFirstFragment verifies the conntrack branch: reply
|
||||||
|
// fragments of an outbound-established UDP flow pass without any inbound rule.
|
||||||
|
func TestFragment_ConntrackAllowsFirstFragment(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
|
||||||
|
out := outboundUDPPacket(t, 12345, 40000)
|
||||||
|
require.False(t, m.filterOutbound(out, len(out)))
|
||||||
|
|
||||||
|
first := firstFragmentUDP(t, 0x8888, 12345, 32)
|
||||||
|
require.False(t, m.filterInbound(first, len(first)),
|
||||||
|
"reply first fragment should pass via conntrack")
|
||||||
|
|
||||||
|
trailing := trailingFragment(t, 0x8888, 5, false, 24)
|
||||||
|
require.False(t, m.filterInbound(trailing, len(trailing)),
|
||||||
|
"trailing fragment of a tracked flow should pass")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_RoutingDisabledDropsFragment verifies routed first fragments are
|
||||||
|
// dropped when routing is disabled.
|
||||||
|
func TestFragment_RoutingDisabledDropsFragment(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
m.routingEnabled.Store(false)
|
||||||
|
|
||||||
|
first := firstFragmentUDPTo(t, "198.51.100.10", 0x9999, 8080, 32)
|
||||||
|
require.True(t, m.filterInbound(first, len(first)),
|
||||||
|
"routed first fragment must be dropped when routing is disabled")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_RouteACL verifies the route-ACL branch: fragments to a non-local
|
||||||
|
// destination follow the route rules, allowed datagrams pass their trailing
|
||||||
|
// fragments and denied ones don't.
|
||||||
|
func TestFragment_RouteACL(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
m.routingEnabled.Store(true)
|
||||||
|
m.nativeRouter.Store(false)
|
||||||
|
|
||||||
|
_, err := m.AddRouteFiltering(
|
||||||
|
[]byte("rt-1"),
|
||||||
|
[]netip.Prefix{netip.MustParsePrefix("100.10.0.0/16")},
|
||||||
|
fw.Network{Prefix: netip.MustParsePrefix("198.51.100.0/24")},
|
||||||
|
fw.ProtocolUDP,
|
||||||
|
nil,
|
||||||
|
&fw.Port{Values: []uint16{8080}},
|
||||||
|
fw.ActionAccept,
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
first := firstFragmentUDPTo(t, "198.51.100.10", 0xAAAA, 8080, 32)
|
||||||
|
require.False(t, m.filterInbound(first, len(first)),
|
||||||
|
"route-ACL-allowed first fragment should pass")
|
||||||
|
trailing := trailingFragmentTo(t, "198.51.100.10", layers.IPProtocolUDP, 0xAAAA, 5, false, 24)
|
||||||
|
require.False(t, m.filterInbound(trailing, len(trailing)),
|
||||||
|
"trailing fragment of an allowed routed datagram should pass")
|
||||||
|
|
||||||
|
denied := firstFragmentUDPTo(t, "198.51.100.10", 0xBBBB, 9999, 32)
|
||||||
|
require.True(t, m.filterInbound(denied, len(denied)),
|
||||||
|
"route-ACL-denied first fragment must be dropped")
|
||||||
|
deniedTrailing := trailingFragmentTo(t, "198.51.100.10", layers.IPProtocolUDP, 0xBBBB, 5, false, 24)
|
||||||
|
require.True(t, m.filterInbound(deniedTrailing, len(deniedTrailing)),
|
||||||
|
"trailing fragment of a denied routed datagram must be dropped")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_ExpiredVerdictDropsTrailing verifies a verdict older than the
|
||||||
|
// tracker timeout no longer admits trailing fragments.
|
||||||
|
func TestFragment_ExpiredVerdictDropsTrailing(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
allowUDP(t, m, 8080)
|
||||||
|
|
||||||
|
first := firstFragmentUDP(t, 0xCCCC, 8080, 32)
|
||||||
|
require.False(t, m.filterInbound(first, len(first)))
|
||||||
|
|
||||||
|
m.fragments.mutex.Lock()
|
||||||
|
for key, entry := range m.fragments.entries {
|
||||||
|
entry.recordedAt = time.Now().Add(-defaultFragmentTimeout - time.Second)
|
||||||
|
m.fragments.entries[key] = entry
|
||||||
|
}
|
||||||
|
m.fragments.mutex.Unlock()
|
||||||
|
|
||||||
|
trailing := trailingFragment(t, 0xCCCC, 5, false, 24)
|
||||||
|
require.True(t, m.filterInbound(trailing, len(trailing)),
|
||||||
|
"trailing fragment after verdict expiry must be dropped")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_CapacityFailsClosed verifies the table cap: at capacity, new
|
||||||
|
// datagram verdicts are not recorded (their trailing fragments are dropped)
|
||||||
|
// while already-recorded datagrams keep working.
|
||||||
|
func TestFragment_CapacityFailsClosed(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
allowUDP(t, m, 8080)
|
||||||
|
|
||||||
|
m.fragments.mutex.Lock()
|
||||||
|
m.fragments.maxEntries = 1
|
||||||
|
m.fragments.mutex.Unlock()
|
||||||
|
|
||||||
|
first1 := firstFragmentUDP(t, 0x0101, 8080, 32)
|
||||||
|
require.False(t, m.filterInbound(first1, len(first1)))
|
||||||
|
|
||||||
|
first2 := firstFragmentUDP(t, 0x0202, 8080, 32)
|
||||||
|
require.False(t, m.filterInbound(first2, len(first2)),
|
||||||
|
"first fragment itself still passes at capacity")
|
||||||
|
|
||||||
|
trailing2 := trailingFragment(t, 0x0202, 5, false, 24)
|
||||||
|
require.True(t, m.filterInbound(trailing2, len(trailing2)),
|
||||||
|
"trailing fragment of an unrecorded datagram must be dropped at capacity")
|
||||||
|
|
||||||
|
trailing1 := trailingFragment(t, 0x0101, 5, false, 24)
|
||||||
|
require.False(t, m.filterInbound(trailing1, len(trailing1)),
|
||||||
|
"already-recorded datagram should keep passing at capacity")
|
||||||
|
}
|
||||||
|
|
||||||
|
// v6FragmentHeader builds the 8-byte IPv6 fragment extension header for the
|
||||||
|
// given inner protocol, offset (8-byte units), More Fragments bit and id.
|
||||||
|
func v6FragmentHeader(proto layers.IPProtocol, offsetOctets uint16, moreFragments bool, id uint32) []byte {
|
||||||
|
offsetFlags := offsetOctets << 3
|
||||||
|
if moreFragments {
|
||||||
|
offsetFlags |= 1
|
||||||
|
}
|
||||||
|
hdr := make([]byte, 8)
|
||||||
|
hdr[0] = uint8(proto)
|
||||||
|
binary.BigEndian.PutUint16(hdr[2:4], offsetFlags)
|
||||||
|
binary.BigEndian.PutUint32(hdr[4:8], id)
|
||||||
|
return hdr
|
||||||
|
}
|
||||||
|
|
||||||
|
func v6UDPHeader(dstPort uint16, dataLen int) []byte {
|
||||||
|
hdr := make([]byte, 8)
|
||||||
|
binary.BigEndian.PutUint16(hdr[0:2], 40000)
|
||||||
|
binary.BigEndian.PutUint16(hdr[2:4], dstPort)
|
||||||
|
binary.BigEndian.PutUint16(hdr[4:6], uint16(8+dataLen))
|
||||||
|
return hdr
|
||||||
|
}
|
||||||
|
|
||||||
|
// firstFragmentUDPv6 builds the first fragment of a fragmented IPv6 UDP
|
||||||
|
// datagram: fragment header (offset 0, More Fragments set) + full UDP header +
|
||||||
|
// data.
|
||||||
|
func firstFragmentUDPv6(tb testing.TB, id uint32, dstPort uint16, dataLen int) []byte {
|
||||||
|
tb.Helper()
|
||||||
|
return fragmentUDPv6(tb, id, dstPort, dataLen, true)
|
||||||
|
}
|
||||||
|
|
||||||
|
// fragmentUDPv6 builds an offset-zero IPv6 UDP fragment. With moreFragments
|
||||||
|
// false it is an atomic fragment (a complete datagram, RFC 6946).
|
||||||
|
func fragmentUDPv6(tb testing.TB, id uint32, dstPort uint16, dataLen int, moreFragments bool) []byte {
|
||||||
|
tb.Helper()
|
||||||
|
|
||||||
|
ip := &layers.IPv6{
|
||||||
|
Version: 6,
|
||||||
|
NextHeader: layers.IPProtocolIPv6Fragment,
|
||||||
|
HopLimit: 64,
|
||||||
|
SrcIP: net.ParseIP(fragTestSrcV6),
|
||||||
|
DstIP: net.ParseIP(fragTestDstV6),
|
||||||
|
}
|
||||||
|
payload := append(v6FragmentHeader(layers.IPProtocolUDP, 0, moreFragments, id), v6UDPHeader(dstPort, dataLen)...)
|
||||||
|
payload = append(payload, make([]byte, dataLen)...)
|
||||||
|
|
||||||
|
buf := gopacket.NewSerializeBuffer()
|
||||||
|
require.NoError(tb, gopacket.SerializeLayers(buf, gopacket.SerializeOptions{FixLengths: true}, ip, gopacket.Payload(payload)))
|
||||||
|
return buf.Bytes()
|
||||||
|
}
|
||||||
|
|
||||||
|
// trailingFragmentV6 builds a non-first IPv6 fragment: fragment header at the
|
||||||
|
// given offset carrying raw data and no transport header.
|
||||||
|
func trailingFragmentV6(tb testing.TB, id uint32, offsetOctets uint16, moreFragments bool, dataLen int) []byte {
|
||||||
|
tb.Helper()
|
||||||
|
|
||||||
|
ip := &layers.IPv6{
|
||||||
|
Version: 6,
|
||||||
|
NextHeader: layers.IPProtocolIPv6Fragment,
|
||||||
|
HopLimit: 64,
|
||||||
|
SrcIP: net.ParseIP(fragTestSrcV6),
|
||||||
|
DstIP: net.ParseIP(fragTestDstV6),
|
||||||
|
}
|
||||||
|
payload := append(v6FragmentHeader(layers.IPProtocolUDP, offsetOctets, moreFragments, id), make([]byte, dataLen)...)
|
||||||
|
|
||||||
|
buf := gopacket.NewSerializeBuffer()
|
||||||
|
require.NoError(tb, gopacket.SerializeLayers(buf, gopacket.SerializeOptions{FixLengths: true}, ip, gopacket.Payload(payload)))
|
||||||
|
return buf.Bytes()
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragmentV6_TrailingWithoutFirstDropped verifies the IPv6 bypass is closed:
|
||||||
|
// a trailing fragment with no allowed first fragment is dropped.
|
||||||
|
func TestFragmentV6_TrailingWithoutFirstDropped(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
|
||||||
|
frag := trailingFragmentV6(t, 0xAABBCCDD, 100, false, 40)
|
||||||
|
require.True(t, m.filterInbound(frag, len(frag)),
|
||||||
|
"IPv6 trailing fragment without an allowed first fragment must be dropped")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragmentV6_AllowedFirstPassesTrailing verifies IPv6 fragments are
|
||||||
|
// evaluated like IPv4: an allowed first fragment lets its trailing fragments
|
||||||
|
// through.
|
||||||
|
func TestFragmentV6_AllowedFirstPassesTrailing(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
_, err := m.AddPeerFiltering(nil, net.ParseIP(fragTestSrcV6), fw.ProtocolUDP, nil,
|
||||||
|
&fw.Port{Values: []uint16{8080}}, fw.ActionAccept, "")
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// First fragment: UDP header (8) + 32 data = 40 octets -> headerEnd = 5.
|
||||||
|
first := firstFragmentUDPv6(t, 0xAABBCCDD, 8080, 32)
|
||||||
|
require.False(t, m.filterInbound(first, len(first)),
|
||||||
|
"allowed IPv6 first fragment should pass and be recorded")
|
||||||
|
|
||||||
|
trailing := trailingFragmentV6(t, 0xAABBCCDD, 5, false, 24)
|
||||||
|
require.False(t, m.filterInbound(trailing, len(trailing)),
|
||||||
|
"trailing fragment of an allowed IPv6 datagram should pass")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragmentV6_AtomicNotCached verifies an IPv6 atomic fragment (fragment
|
||||||
|
// header with offset 0 and no More Fragments, a complete datagram per RFC 6946)
|
||||||
|
// is evaluated but not recorded, so a flood of allowed atomic fragments can't
|
||||||
|
// exhaust the verdict table.
|
||||||
|
func TestFragmentV6_AtomicNotCached(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
_, err := m.AddPeerFiltering(nil, net.ParseIP(fragTestSrcV6), fw.ProtocolUDP, nil,
|
||||||
|
&fw.Port{Values: []uint16{8080}}, fw.ActionAccept, "")
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
atomic := fragmentUDPv6(t, 0xA70301C, 8080, 16, false)
|
||||||
|
require.False(t, m.filterInbound(atomic, len(atomic)),
|
||||||
|
"allowed IPv6 atomic fragment should pass")
|
||||||
|
|
||||||
|
m.fragments.mutex.Lock()
|
||||||
|
n := len(m.fragments.entries)
|
||||||
|
m.fragments.mutex.Unlock()
|
||||||
|
require.Zero(t, n, "atomic fragment must not create a verdict entry")
|
||||||
|
|
||||||
|
// A genuine fragmented datagram (More Fragments set) is still recorded.
|
||||||
|
first := fragmentUDPv6(t, 0xBEEF, 8080, 32, true)
|
||||||
|
require.False(t, m.filterInbound(first, len(first)))
|
||||||
|
m.fragments.mutex.Lock()
|
||||||
|
n = len(m.fragments.entries)
|
||||||
|
m.fragments.mutex.Unlock()
|
||||||
|
require.Equal(t, 1, n, "genuine first fragment must record a verdict")
|
||||||
|
}
|
||||||
@@ -3,14 +3,31 @@
|
|||||||
package netstack
|
package netstack
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
"net"
|
||||||
"os"
|
"os"
|
||||||
"strconv"
|
"strconv"
|
||||||
|
|
||||||
log "github.com/sirupsen/logrus"
|
log "github.com/sirupsen/logrus"
|
||||||
)
|
)
|
||||||
|
|
||||||
const EnvUseNetstackMode = "NB_USE_NETSTACK_MODE"
|
const (
|
||||||
|
EnvUseNetstackMode = "NB_USE_NETSTACK_MODE"
|
||||||
|
|
||||||
|
// EnvSocks5ListenerPort overrides the port the SOCKS5 proxy listens on.
|
||||||
|
EnvSocks5ListenerPort = "NB_SOCKS5_LISTENER_PORT"
|
||||||
|
|
||||||
|
// EnvSocks5ListenerAddress overrides the host/IP the SOCKS5 proxy binds to.
|
||||||
|
// The proxy is a bridge for local host applications into the userspace
|
||||||
|
// WireGuard netstack, so it binds to loopback by default. Override this only
|
||||||
|
// when the proxy must be reachable from other hosts (e.g. a container
|
||||||
|
// gateway); doing so exposes an unauthenticated SOCKS5 proxy on that
|
||||||
|
// address.
|
||||||
|
EnvSocks5ListenerAddress = "NB_SOCKS5_LISTENER_ADDRESS"
|
||||||
|
|
||||||
|
// defaultSocks5Host is the loopback address the SOCKS5 proxy binds to unless
|
||||||
|
// overridden via EnvSocks5ListenerAddress.
|
||||||
|
defaultSocks5Host = "127.0.0.1"
|
||||||
|
)
|
||||||
|
|
||||||
// IsEnabled todo: move these function to cmd layer
|
// IsEnabled todo: move these function to cmd layer
|
||||||
func IsEnabled() bool {
|
func IsEnabled() bool {
|
||||||
@@ -18,24 +35,40 @@ func IsEnabled() bool {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func ListenAddr() string {
|
func ListenAddr() string {
|
||||||
sPort := os.Getenv("NB_SOCKS5_LISTENER_PORT")
|
return net.JoinHostPort(listenHost(), strconv.Itoa(listenPort()))
|
||||||
|
}
|
||||||
|
|
||||||
|
// listenHost returns the host/IP the SOCKS5 proxy binds to. It defaults to
|
||||||
|
// loopback and only honors EnvSocks5ListenerAddress when it holds a valid IP.
|
||||||
|
func listenHost() string {
|
||||||
|
addr := os.Getenv(EnvSocks5ListenerAddress)
|
||||||
|
if addr == "" {
|
||||||
|
return defaultSocks5Host
|
||||||
|
}
|
||||||
|
if net.ParseIP(addr) == nil {
|
||||||
|
log.Warnf("invalid socks5 listener address %q, falling back to default: %s", addr, defaultSocks5Host)
|
||||||
|
return defaultSocks5Host
|
||||||
|
}
|
||||||
|
return addr
|
||||||
|
}
|
||||||
|
|
||||||
|
// listenPort returns the port the SOCKS5 proxy binds to, defaulting to
|
||||||
|
// DefaultSocks5Port when EnvSocks5ListenerPort is unset or invalid.
|
||||||
|
func listenPort() int {
|
||||||
|
sPort := os.Getenv(EnvSocks5ListenerPort)
|
||||||
if sPort == "" {
|
if sPort == "" {
|
||||||
return listenAddr(DefaultSocks5Port)
|
return DefaultSocks5Port
|
||||||
}
|
}
|
||||||
|
|
||||||
port, err := strconv.Atoi(sPort)
|
port, err := strconv.Atoi(sPort)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Warnf("invalid socks5 listener port, unable to convert it to int, falling back to default: %d", DefaultSocks5Port)
|
log.Warnf("invalid socks5 listener port, unable to convert it to int, falling back to default: %d", DefaultSocks5Port)
|
||||||
return listenAddr(DefaultSocks5Port)
|
return DefaultSocks5Port
|
||||||
}
|
}
|
||||||
if port < 1 || port > 65535 {
|
if port < 1 || port > 65535 {
|
||||||
log.Warnf("invalid socks5 listener port, it should be in the range 1-65535, falling back to default: %d", DefaultSocks5Port)
|
log.Warnf("invalid socks5 listener port, it should be in the range 1-65535, falling back to default: %d", DefaultSocks5Port)
|
||||||
return listenAddr(DefaultSocks5Port)
|
return DefaultSocks5Port
|
||||||
}
|
}
|
||||||
|
|
||||||
return listenAddr(port)
|
return port
|
||||||
}
|
|
||||||
|
|
||||||
func listenAddr(port int) string {
|
|
||||||
return fmt.Sprintf("0.0.0.0:%d", port)
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,63 @@
|
|||||||
|
//go:build !js
|
||||||
|
|
||||||
|
package netstack
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
"strconv"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestListenAddr_DefaultsToLoopback(t *testing.T) {
|
||||||
|
// No env overrides: must bind loopback, never all interfaces.
|
||||||
|
got := ListenAddr()
|
||||||
|
want := net.JoinHostPort("127.0.0.1", strconv.Itoa(DefaultSocks5Port))
|
||||||
|
if got != want {
|
||||||
|
t.Fatalf("ListenAddr() = %q, want %q", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestListenAddr_AddressOverride(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
env string
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{name: "valid override honored", env: "0.0.0.0", want: "0.0.0.0"},
|
||||||
|
{name: "valid specific ip honored", env: "10.0.0.5", want: "10.0.0.5"},
|
||||||
|
{name: "ipv6 loopback bracketed", env: "::1", want: "::1"},
|
||||||
|
{name: "invalid falls back to loopback", env: "not-an-ip", want: "127.0.0.1"},
|
||||||
|
{name: "empty falls back to loopback", env: "", want: "127.0.0.1"},
|
||||||
|
}
|
||||||
|
for _, tc := range tests {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
t.Setenv(EnvSocks5ListenerAddress, tc.env)
|
||||||
|
want := net.JoinHostPort(tc.want, strconv.Itoa(DefaultSocks5Port))
|
||||||
|
if got := ListenAddr(); got != want {
|
||||||
|
t.Fatalf("ListenAddr() = %q, want %q", got, want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestListenAddr_PortOverride(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
env string
|
||||||
|
want int
|
||||||
|
}{
|
||||||
|
{name: "valid port honored", env: "1081", want: 1081},
|
||||||
|
{name: "non-numeric falls back", env: "abc", want: DefaultSocks5Port},
|
||||||
|
{name: "out of range falls back", env: "70000", want: DefaultSocks5Port},
|
||||||
|
{name: "zero falls back", env: "0", want: DefaultSocks5Port},
|
||||||
|
}
|
||||||
|
for _, tc := range tests {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
t.Setenv(EnvSocks5ListenerPort, tc.env)
|
||||||
|
want := net.JoinHostPort("127.0.0.1", strconv.Itoa(tc.want))
|
||||||
|
if got := ListenAddr(); got != want {
|
||||||
|
t.Fatalf("ListenAddr() = %q, want %q", got, want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -299,7 +299,7 @@ func (d *DeviceAuthorizationFlow) WaitToken(ctx context.Context, info AuthFlowIn
|
|||||||
UseIDToken: d.providerConfig.UseIDToken,
|
UseIDToken: d.providerConfig.UseIDToken,
|
||||||
}
|
}
|
||||||
|
|
||||||
err = isValidAccessToken(tokenInfo.GetTokenToUse(), d.providerConfig.Audience)
|
err = validateTokenAudience(tokenInfo.GetTokenToUse(), d.providerConfig.Audience)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return TokenInfo{}, fmt.Errorf("validate access token failed with error: %v", err)
|
return TokenInfo{}, fmt.Errorf("validate access token failed with error: %v", err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -306,7 +306,7 @@ func (p *PKCEAuthorizationFlow) parseOAuthToken(token *oauth2.Token) (TokenInfo,
|
|||||||
audience = p.providerConfig.ClientID
|
audience = p.providerConfig.ClientID
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := isValidAccessToken(tokenInfo.GetTokenToUse(), audience); err != nil {
|
if err := validateTokenAudience(tokenInfo.GetTokenToUse(), audience); err != nil {
|
||||||
return TokenInfo{}, fmt.Errorf("authentication failed: invalid access token - %w", err)
|
return TokenInfo{}, fmt.Errorf("authentication failed: invalid access token - %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -320,6 +320,11 @@ func (p *PKCEAuthorizationFlow) parseOAuthToken(token *oauth2.Token) (TokenInfo,
|
|||||||
return tokenInfo, nil
|
return tokenInfo, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// parseEmailFromIDToken extracts the email (or name) claim from an ID token
|
||||||
|
// without verifying its signature. The value is best-effort and used only as a
|
||||||
|
// UX convenience (login hint prefill and display); it never drives an
|
||||||
|
// authorization decision. The authoritative identity is established server-side
|
||||||
|
// from the signature-verified token.
|
||||||
func parseEmailFromIDToken(token string) (string, error) {
|
func parseEmailFromIDToken(token string) (string, error) {
|
||||||
parts := strings.Split(token, ".")
|
parts := strings.Split(token, ".")
|
||||||
if len(parts) < 2 {
|
if len(parts) < 2 {
|
||||||
|
|||||||
@@ -20,14 +20,26 @@ func randomBytesInHex(count int) (string, error) {
|
|||||||
return hex.EncodeToString(buf), nil
|
return hex.EncodeToString(buf), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// isValidAccessToken is a simple validation of the access token
|
// validateTokenAudience checks that the token is a well-formed JWT whose
|
||||||
func isValidAccessToken(token string, audience string) error {
|
// audience claim matches the expected audience.
|
||||||
|
//
|
||||||
|
// It does NOT verify the token's cryptographic signature and therefore must not
|
||||||
|
// be treated as an authenticity check. The token is obtained by the client
|
||||||
|
// directly from the IdP token endpoint over TLS, and its signature is verified
|
||||||
|
// server-side by the management server against the IdP's JWKS
|
||||||
|
// (see shared/auth/jwt/validator.go). This function is only a client-side
|
||||||
|
// sanity check that the returned token targets the expected audience.
|
||||||
|
func validateTokenAudience(token string, audience string) error {
|
||||||
if token == "" {
|
if token == "" {
|
||||||
return fmt.Errorf("token received is empty")
|
return fmt.Errorf("token received is empty")
|
||||||
}
|
}
|
||||||
|
|
||||||
encodedClaims := strings.Split(token, ".")[1]
|
parts := strings.Split(token, ".")
|
||||||
claimsString, err := base64.RawURLEncoding.DecodeString(encodedClaims)
|
if len(parts) != 3 {
|
||||||
|
return fmt.Errorf("token is not a well-formed JWT")
|
||||||
|
}
|
||||||
|
|
||||||
|
claimsString, err := base64.RawURLEncoding.DecodeString(parts[1])
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,108 @@
|
|||||||
|
package auth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/base64"
|
||||||
|
"encoding/json"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// makeJWT builds an unsigned JWT-shaped string (header.payload.signature) with
|
||||||
|
// the given claims payload. The signature part is arbitrary because
|
||||||
|
// validateTokenAudience intentionally does not verify it.
|
||||||
|
func makeJWT(t *testing.T, claims map[string]interface{}) string {
|
||||||
|
t.Helper()
|
||||||
|
header := base64.RawURLEncoding.EncodeToString([]byte(`{"alg":"RS256","typ":"JWT"}`))
|
||||||
|
payloadBytes, err := json.Marshal(claims)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("marshal claims: %v", err)
|
||||||
|
}
|
||||||
|
payload := base64.RawURLEncoding.EncodeToString(payloadBytes)
|
||||||
|
return header + "." + payload + ".unverified-signature"
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestValidateTokenAudience(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
token string
|
||||||
|
audience string
|
||||||
|
wantErr bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "empty token",
|
||||||
|
token: "",
|
||||||
|
audience: "netbird",
|
||||||
|
wantErr: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "not a JWT - no dots",
|
||||||
|
token: "notajwt",
|
||||||
|
audience: "netbird",
|
||||||
|
wantErr: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "not a JWT - two parts only",
|
||||||
|
token: "header.payload",
|
||||||
|
audience: "netbird",
|
||||||
|
wantErr: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "matching string audience",
|
||||||
|
token: makeJWT(t, map[string]interface{}{"aud": "netbird"}),
|
||||||
|
audience: "netbird",
|
||||||
|
wantErr: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "mismatching string audience",
|
||||||
|
token: makeJWT(t, map[string]interface{}{"aud": "other"}),
|
||||||
|
audience: "netbird",
|
||||||
|
wantErr: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "matching audience in array",
|
||||||
|
token: makeJWT(t, map[string]interface{}{"aud": []interface{}{"other", "netbird"}}),
|
||||||
|
audience: "netbird",
|
||||||
|
wantErr: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "mismatching audience array",
|
||||||
|
token: makeJWT(t, map[string]interface{}{"aud": []interface{}{"a", "b"}}),
|
||||||
|
audience: "netbird",
|
||||||
|
wantErr: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "missing audience claim",
|
||||||
|
token: makeJWT(t, map[string]interface{}{"sub": "user"}),
|
||||||
|
audience: "netbird",
|
||||||
|
wantErr: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "invalid base64 payload",
|
||||||
|
token: "header.!!!not-base64!!!.sig",
|
||||||
|
audience: "netbird",
|
||||||
|
wantErr: true,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range tests {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
err := validateTokenAudience(tc.token, tc.audience)
|
||||||
|
if tc.wantErr && err == nil {
|
||||||
|
t.Fatalf("expected error, got nil")
|
||||||
|
}
|
||||||
|
if !tc.wantErr && err != nil {
|
||||||
|
t.Fatalf("expected no error, got %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestValidateTokenAudienceNoPanic guards the regression where a non-empty
|
||||||
|
// token without the JWT dot structure caused an index-out-of-range panic.
|
||||||
|
func TestValidateTokenAudienceNoPanic(t *testing.T) {
|
||||||
|
inputs := []string{"a", ".", "a.", "aaaa", "no-dots-here"}
|
||||||
|
for _, in := range inputs {
|
||||||
|
if err := validateTokenAudience(in, "netbird"); err == nil {
|
||||||
|
t.Fatalf("expected error for malformed token %q, got nil", in)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,6 +1,7 @@
|
|||||||
package statemanager
|
package statemanager
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
@@ -305,6 +306,11 @@ func (m *Manager) loadStateFile(deleteCorrupt bool) (map[string]json.RawMessage,
|
|||||||
|
|
||||||
var rawStates map[string]json.RawMessage
|
var rawStates map[string]json.RawMessage
|
||||||
if err := json.Unmarshal(data, &rawStates); err != nil {
|
if err := json.Unmarshal(data, &rawStates); err != nil {
|
||||||
|
if len(bytes.TrimSpace(data)) == 0 {
|
||||||
|
log.Warnf("state file %s is empty (%d bytes)", m.filePath, len(data))
|
||||||
|
} else {
|
||||||
|
log.Warnf("state file %s has malformed content (%d bytes)", m.filePath, len(data))
|
||||||
|
}
|
||||||
m.handleCorruptedState(deleteCorrupt)
|
m.handleCorruptedState(deleteCorrupt)
|
||||||
return nil, fmt.Errorf("unmarshal states: %w", err)
|
return nil, fmt.Errorf("unmarshal states: %w", err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ import (
|
|||||||
log "github.com/sirupsen/logrus"
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
nbssh "github.com/netbirdio/netbird/client/ssh"
|
nbssh "github.com/netbirdio/netbird/client/ssh"
|
||||||
|
"github.com/netbirdio/netbird/shared/management/domain"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -218,11 +219,20 @@ func (m *Manager) buildHostPatterns(peer PeerSSHInfo) []string {
|
|||||||
if peer.IPv6.IsValid() {
|
if peer.IPv6.IsValid() {
|
||||||
hostPatterns = append(hostPatterns, peer.IPv6.String())
|
hostPatterns = append(hostPatterns, peer.IPv6.String())
|
||||||
}
|
}
|
||||||
if peer.FQDN != "" {
|
// Peer FQDNs and hostnames originate from remote peers, so they must be
|
||||||
|
// validated as plain DNS names before being embedded in the ssh_config
|
||||||
|
// "Match host" pattern list. This prevents injection of arbitrary
|
||||||
|
// ssh_config directives via embedded quotes, whitespace, newlines, the
|
||||||
|
// comma pattern separator, or the "*"/"?" pattern metacharacters.
|
||||||
|
if domain.IsValidDomainNoWildcard(peer.FQDN) {
|
||||||
hostPatterns = append(hostPatterns, peer.FQDN)
|
hostPatterns = append(hostPatterns, peer.FQDN)
|
||||||
|
} else if peer.FQDN != "" {
|
||||||
|
log.Warnf("skipping peer FQDN with invalid characters in SSH config: %q", peer.FQDN)
|
||||||
}
|
}
|
||||||
if peer.Hostname != "" && peer.Hostname != peer.FQDN {
|
if peer.Hostname != peer.FQDN && domain.IsValidDomainNoWildcard(peer.Hostname) {
|
||||||
hostPatterns = append(hostPatterns, peer.Hostname)
|
hostPatterns = append(hostPatterns, peer.Hostname)
|
||||||
|
} else if peer.Hostname != "" && peer.Hostname != peer.FQDN {
|
||||||
|
log.Warnf("skipping peer hostname with invalid characters in SSH config: %q", peer.Hostname)
|
||||||
}
|
}
|
||||||
return hostPatterns
|
return hostPatterns
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -148,6 +148,45 @@ func TestManager_MatchHostFormat(t *testing.T) {
|
|||||||
"should use Match host with comma-separated patterns")
|
"should use Match host with comma-separated patterns")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestManager_HostPatternInjection(t *testing.T) {
|
||||||
|
tempDir, err := os.MkdirTemp("", "netbird-ssh-config-test")
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer func() { assert.NoError(t, os.RemoveAll(tempDir)) }()
|
||||||
|
|
||||||
|
manager := &Manager{
|
||||||
|
sshConfigDir: filepath.Join(tempDir, "ssh_config.d"),
|
||||||
|
sshConfigFile: "99-netbird.conf",
|
||||||
|
}
|
||||||
|
|
||||||
|
// A malicious peer FQDN/hostname attempts to break out of the Match host
|
||||||
|
// directive and inject arbitrary ssh_config (a ProxyCommand executing a
|
||||||
|
// command). It must be rejected, not written to the config.
|
||||||
|
peers := []PeerSSHInfo{
|
||||||
|
{
|
||||||
|
Hostname: "evil\"\n ProxyCommand touch /tmp/pwned\nHost x",
|
||||||
|
IP: netip.MustParseAddr("100.125.1.1"),
|
||||||
|
FQDN: "evil\"\n ProxyCommand touch /tmp/pwned\nHost x.nb.internal",
|
||||||
|
},
|
||||||
|
{Hostname: "peer2", IP: netip.MustParseAddr("100.125.1.2"), FQDN: "peer2.nb.internal"},
|
||||||
|
}
|
||||||
|
|
||||||
|
err = manager.SetupSSHClientConfig(peers)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
configPath := filepath.Join(manager.sshConfigDir, manager.sshConfigFile)
|
||||||
|
content, err := os.ReadFile(configPath)
|
||||||
|
require.NoError(t, err)
|
||||||
|
configStr := string(content)
|
||||||
|
|
||||||
|
assert.NotContains(t, configStr, "ProxyCommand touch /tmp/pwned",
|
||||||
|
"injected directive must not appear in generated config")
|
||||||
|
assert.NotContains(t, configStr, "evil",
|
||||||
|
"malicious pattern must be dropped entirely")
|
||||||
|
// The valid peer must still be present, on a single Match host line.
|
||||||
|
assert.Contains(t, configStr, "Match host \"100.125.1.1,100.125.1.2,peer2.nb.internal,peer2\"",
|
||||||
|
"valid peers must survive, injected patterns dropped")
|
||||||
|
}
|
||||||
|
|
||||||
func TestManager_ForcedSSHConfig(t *testing.T) {
|
func TestManager_ForcedSSHConfig(t *testing.T) {
|
||||||
// Set force environment variable
|
// Set force environment variable
|
||||||
t.Setenv(EnvForceSSHConfig, "true")
|
t.Setenv(EnvForceSSHConfig, "true")
|
||||||
|
|||||||
@@ -69,7 +69,8 @@ func parseGetentPasswd(output string) (*user.User, string, error) {
|
|||||||
|
|
||||||
// validateGetentInput checks that the input is safe to pass to getent or id.
|
// validateGetentInput checks that the input is safe to pass to getent or id.
|
||||||
// Allows POSIX usernames, numeric UIDs, and common NSS extensions
|
// Allows POSIX usernames, numeric UIDs, and common NSS extensions
|
||||||
// (@ for Kerberos, $ for Samba, + for NIS compat).
|
// (@ for Kerberos, $ for Samba, + for NIS compat). A leading hyphen is
|
||||||
|
// rejected so the input can never be parsed as a command-line flag.
|
||||||
func validateGetentInput(input string) bool {
|
func validateGetentInput(input string) bool {
|
||||||
maxLen := 32
|
maxLen := 32
|
||||||
if runtime.GOOS == "linux" {
|
if runtime.GOOS == "linux" {
|
||||||
@@ -80,6 +81,10 @@ func validateGetentInput(input string) bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if input[0] == '-' {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
for _, r := range input {
|
for _, r := range input {
|
||||||
if isAllowedGetentChar(r) {
|
if isAllowedGetentChar(r) {
|
||||||
continue
|
continue
|
||||||
|
|||||||
@@ -157,6 +157,9 @@ func TestValidateGetentInput(t *testing.T) {
|
|||||||
{"numeric UID", "1001", true},
|
{"numeric UID", "1001", true},
|
||||||
{"dots and underscores", "alice.bob_test", true},
|
{"dots and underscores", "alice.bob_test", true},
|
||||||
{"hyphen", "alice-bob", true},
|
{"hyphen", "alice-bob", true},
|
||||||
|
{"leading hyphen rejected", "-i", false},
|
||||||
|
{"leading double hyphen rejected", "--no-idn", false},
|
||||||
|
{"lone hyphen rejected", "-", false},
|
||||||
{"kerberos principal", "user@REALM", true},
|
{"kerberos principal", "user@REALM", true},
|
||||||
{"samba machine account", "MACHINE$", true},
|
{"samba machine account", "MACHINE$", true},
|
||||||
{"NIS compat", "+user", true},
|
{"NIS compat", "+user", true},
|
||||||
|
|||||||
@@ -145,6 +145,7 @@ type AuthConfig struct {
|
|||||||
CLIRedirectURIs []string `yaml:"cliRedirectURIs"`
|
CLIRedirectURIs []string `yaml:"cliRedirectURIs"`
|
||||||
Owner *AuthOwnerConfig `yaml:"owner,omitempty"`
|
Owner *AuthOwnerConfig `yaml:"owner,omitempty"`
|
||||||
DashboardPostLogoutRedirectURIs []string `yaml:"dashboardPostLogoutRedirectURIs"`
|
DashboardPostLogoutRedirectURIs []string `yaml:"dashboardPostLogoutRedirectURIs"`
|
||||||
|
GrantTypes []string `yaml:"grantTypes"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// AuthStorageConfig contains auth storage settings
|
// AuthStorageConfig contains auth storage settings
|
||||||
@@ -604,6 +605,7 @@ func (c *CombinedConfig) buildEmbeddedIdPConfig(mgmt ManagementConfig) (*idp.Emb
|
|||||||
DashboardRedirectURIs: mgmt.Auth.DashboardRedirectURIs,
|
DashboardRedirectURIs: mgmt.Auth.DashboardRedirectURIs,
|
||||||
CLIRedirectURIs: mgmt.Auth.CLIRedirectURIs,
|
CLIRedirectURIs: mgmt.Auth.CLIRedirectURIs,
|
||||||
DashboardPostLogoutRedirectURIs: mgmt.Auth.DashboardPostLogoutRedirectURIs,
|
DashboardPostLogoutRedirectURIs: mgmt.Auth.DashboardPostLogoutRedirectURIs,
|
||||||
|
GrantTypes: mgmt.Auth.GrantTypes,
|
||||||
}
|
}
|
||||||
|
|
||||||
if mgmt.Auth.Owner != nil && mgmt.Auth.Owner.Email != "" {
|
if mgmt.Auth.Owner != nil && mgmt.Auth.Owner.Email != "" {
|
||||||
|
|||||||
@@ -335,7 +335,7 @@ replace github.com/cloudflare/circl => codeberg.org/cunicu/circl v0.0.0-20230801
|
|||||||
|
|
||||||
replace github.com/pion/ice/v4 => github.com/netbirdio/ice/v4 v4.0.0-20250908184934-6202be846b51
|
replace github.com/pion/ice/v4 => github.com/netbirdio/ice/v4 v4.0.0-20250908184934-6202be846b51
|
||||||
|
|
||||||
replace github.com/dexidp/dex => github.com/netbirdio/dex v0.244.1-0.20260512110716-8d70ad8647c1
|
replace github.com/dexidp/dex => github.com/netbirdio/dex v0.244.1-0.20260716205454-a163de3129e5
|
||||||
|
|
||||||
replace github.com/dexidp/dex/api/v2 => github.com/netbirdio/dex/api/v2 v2.0.0-20260512110716-8d70ad8647c1
|
replace github.com/dexidp/dex/api/v2 => github.com/netbirdio/dex/api/v2 v2.0.0-20260512110716-8d70ad8647c1
|
||||||
|
|
||||||
|
|||||||
@@ -476,8 +476,8 @@ github.com/morikuni/aec v1.0.0 h1:nP9CBfwrvYnBRgY6qfDQkygYDmYwOilePFkwzv4dU8A=
|
|||||||
github.com/morikuni/aec v1.0.0/go.mod h1:BbKIizmSmc5MMPqRYbxO4ZU0S0+P200+tUnFx7PXmsc=
|
github.com/morikuni/aec v1.0.0/go.mod h1:BbKIizmSmc5MMPqRYbxO4ZU0S0+P200+tUnFx7PXmsc=
|
||||||
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA=
|
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA=
|
||||||
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ=
|
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ=
|
||||||
github.com/netbirdio/dex v0.244.1-0.20260512110716-8d70ad8647c1 h1:4TaYr9O4xX0D2kszeOLclTiCbA3eHq3xWV+9ILJbIYs=
|
github.com/netbirdio/dex v0.244.1-0.20260716205454-a163de3129e5 h1:3PwQv8aR46qN2u16+Dv6udnH3sbVKX5KrGwF35CKSI0=
|
||||||
github.com/netbirdio/dex v0.244.1-0.20260512110716-8d70ad8647c1/go.mod h1:IHH+H8vK2GfqtIt5u/5OdPh18yk0oDHuj2vz5+Goetg=
|
github.com/netbirdio/dex v0.244.1-0.20260716205454-a163de3129e5/go.mod h1:IHH+H8vK2GfqtIt5u/5OdPh18yk0oDHuj2vz5+Goetg=
|
||||||
github.com/netbirdio/dex/api/v2 v2.0.0-20260512110716-8d70ad8647c1 h1:neE7z+FPUkldl3faK/Jt+hJK2L+1XfQ1W33TQhU9m88=
|
github.com/netbirdio/dex/api/v2 v2.0.0-20260512110716-8d70ad8647c1 h1:neE7z+FPUkldl3faK/Jt+hJK2L+1XfQ1W33TQhU9m88=
|
||||||
github.com/netbirdio/dex/api/v2 v2.0.0-20260512110716-8d70ad8647c1/go.mod h1:awuTyT29CYALpEyET0S307EgNlPWrc7fFKRAyhsO45M=
|
github.com/netbirdio/dex/api/v2 v2.0.0-20260512110716-8d70ad8647c1/go.mod h1:awuTyT29CYALpEyET0S307EgNlPWrc7fFKRAyhsO45M=
|
||||||
github.com/netbirdio/easyjson v0.9.0 h1:6Nw2lghSVuy8RSkAYDhDv1thBVEmfVbKZnV7T7Z6Aus=
|
github.com/netbirdio/easyjson v0.9.0 h1:6Nw2lghSVuy8RSkAYDhDv1thBVEmfVbKZnV7T7Z6Aus=
|
||||||
|
|||||||
@@ -613,6 +613,10 @@ func (c *YAMLConfig) ToServerConfig(stor storage.Storage, logger *slog.Logger) s
|
|||||||
cfg.SupportedResponseTypes = c.OAuth2.ResponseTypes
|
cfg.SupportedResponseTypes = c.OAuth2.ResponseTypes
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if len(c.OAuth2.GrantTypes) > 0 {
|
||||||
|
cfg.AllowedGrantTypes = c.OAuth2.GrantTypes
|
||||||
|
}
|
||||||
|
|
||||||
// Apply expiry settings
|
// Apply expiry settings
|
||||||
if c.Expiry.IDTokens != "" {
|
if c.Expiry.IDTokens != "" {
|
||||||
if d, err := parseDuration(c.Expiry.IDTokens); err == nil {
|
if d, err := parseDuration(c.Expiry.IDTokens); err == nil {
|
||||||
|
|||||||
+1
-1
@@ -21,7 +21,7 @@ import (
|
|||||||
"github.com/dexidp/dex/server/signer"
|
"github.com/dexidp/dex/server/signer"
|
||||||
"github.com/dexidp/dex/storage"
|
"github.com/dexidp/dex/storage"
|
||||||
"github.com/dexidp/dex/storage/sql"
|
"github.com/dexidp/dex/storage/sql"
|
||||||
jose "github.com/go-jose/go-jose/v4"
|
"github.com/go-jose/go-jose/v4"
|
||||||
"github.com/google/uuid"
|
"github.com/google/uuid"
|
||||||
"github.com/prometheus/client_golang/prometheus"
|
"github.com/prometheus/client_golang/prometheus"
|
||||||
"golang.org/x/crypto/bcrypt"
|
"golang.org/x/crypto/bcrypt"
|
||||||
|
|||||||
@@ -595,3 +595,90 @@ enablePasswordDB: true
|
|||||||
assert.True(t, cfg.ContinueOnConnectorFailure,
|
assert.True(t, cfg.ContinueOnConnectorFailure,
|
||||||
"buildDexConfig must set ContinueOnConnectorFailure to true so management starts even if an external IdP is down")
|
"buildDexConfig must set ContinueOnConnectorFailure to true so management starts even if an external IdP is down")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestToServerConfig_WiresGrantTypes(t *testing.T) {
|
||||||
|
tmpDir, err := os.MkdirTemp("", "dex-grants-*")
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer os.RemoveAll(tmpDir)
|
||||||
|
|
||||||
|
stor := openTestStorage(t, tmpDir)
|
||||||
|
defer stor.Close()
|
||||||
|
|
||||||
|
logger := slog.New(slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{Level: slog.LevelError}))
|
||||||
|
|
||||||
|
grants := []string{"authorization_code", "refresh_token"}
|
||||||
|
cfg := &YAMLConfig{Issuer: "http://localhost:5599/oauth2", OAuth2: OAuth2{GrantTypes: grants}}
|
||||||
|
assert.Equal(t, grants, cfg.ToServerConfig(stor, logger).AllowedGrantTypes)
|
||||||
|
|
||||||
|
empty := &YAMLConfig{Issuer: "http://localhost:5599/oauth2"}
|
||||||
|
assert.Empty(t, empty.ToServerConfig(stor, logger).AllowedGrantTypes)
|
||||||
|
}
|
||||||
|
|
||||||
|
func newDeviceGuardProvider(t *testing.T, grantTypesYAML string) *Provider {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
tmpDir, err := os.MkdirTemp("", "dex-devguard-*")
|
||||||
|
require.NoError(t, err)
|
||||||
|
t.Cleanup(func() { _ = os.RemoveAll(tmpDir) })
|
||||||
|
|
||||||
|
yamlContent := `
|
||||||
|
issuer: http://localhost:5599/oauth2
|
||||||
|
storage:
|
||||||
|
type: sqlite3
|
||||||
|
config:
|
||||||
|
file: ` + filepath.Join(tmpDir, "dex.db") + `
|
||||||
|
web:
|
||||||
|
http: 127.0.0.1:5599
|
||||||
|
enablePasswordDB: true
|
||||||
|
` + grantTypesYAML
|
||||||
|
|
||||||
|
configPath := filepath.Join(tmpDir, "config.yaml")
|
||||||
|
require.NoError(t, os.WriteFile(configPath, []byte(yamlContent), 0644))
|
||||||
|
|
||||||
|
yamlConfig, err := LoadConfig(configPath)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
provider, err := NewProviderFromYAML(context.Background(), yamlConfig)
|
||||||
|
require.NoError(t, err)
|
||||||
|
t.Cleanup(func() { _ = provider.Stop(context.Background()) })
|
||||||
|
return provider
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandler_BlocksDeviceEndpointsWhenDeviceGrantDisabled(t *testing.T) {
|
||||||
|
provider := newDeviceGuardProvider(t, `
|
||||||
|
oauth2:
|
||||||
|
grantTypes:
|
||||||
|
- authorization_code
|
||||||
|
- refresh_token
|
||||||
|
`)
|
||||||
|
|
||||||
|
devicePaths := []string{
|
||||||
|
"/oauth2/device",
|
||||||
|
"/oauth2/device/code",
|
||||||
|
"/oauth2/device/token",
|
||||||
|
"/oauth2/device/auth/verify_code",
|
||||||
|
"/oauth2/device/callback",
|
||||||
|
}
|
||||||
|
for _, path := range devicePaths {
|
||||||
|
for _, method := range []string{http.MethodGet, http.MethodPost} {
|
||||||
|
req := httptest.NewRequest(method, path, nil)
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
provider.Handler().ServeHTTP(rec, req)
|
||||||
|
assert.Equal(t, http.StatusNotFound, rec.Code, "%s %s must be blocked", method, path)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/oauth2/.well-known/openid-configuration", nil)
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
provider.Handler().ServeHTTP(rec, req)
|
||||||
|
assert.Equal(t, http.StatusOK, rec.Code)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandler_AllowsDeviceEndpointsWhenGrantsDefault(t *testing.T) {
|
||||||
|
provider := newDeviceGuardProvider(t, "")
|
||||||
|
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/oauth2/device/code", nil)
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
provider.Handler().ServeHTTP(rec, req)
|
||||||
|
assert.NotEqual(t, http.StatusNotFound, rec.Code)
|
||||||
|
}
|
||||||
|
|||||||
@@ -4269,7 +4269,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_NetworkRangePreserved(t *te
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Sanity: an actually different range still triggers reallocation.
|
// Sanity: an actually different range still triggers reallocation.
|
||||||
newRange := netip.MustParsePrefix("100.99.0.0/16")
|
newRange := netip.MustParsePrefix("100.60.0.0/16")
|
||||||
_, err = manager.UpdateAccountSettings(ctx, account.Id, userID, &types.Settings{
|
_, err = manager.UpdateAccountSettings(ctx, account.Id, userID, &types.Settings{
|
||||||
PeerLoginExpirationEnabled: true,
|
PeerLoginExpirationEnabled: true,
|
||||||
PeerLoginExpiration: types.DefaultPeerLoginExpiration,
|
PeerLoginExpiration: types.DefaultPeerLoginExpiration,
|
||||||
|
|||||||
@@ -76,6 +76,9 @@ type EmbeddedIdPConfig struct {
|
|||||||
DashboardPostLogoutRedirectURIs []string
|
DashboardPostLogoutRedirectURIs []string
|
||||||
// StaticConnectors are additional connectors to seed during initialization
|
// StaticConnectors are additional connectors to seed during initialization
|
||||||
StaticConnectors []dex.Connector
|
StaticConnectors []dex.Connector
|
||||||
|
// GrantTypes restricts allowed OAuth2 grants; empty means all (Dex default). Omit the
|
||||||
|
// device_code grant to disable the device flow; keep authorization_code and refresh_token.
|
||||||
|
GrantTypes []string
|
||||||
}
|
}
|
||||||
|
|
||||||
// EmbeddedStorageConfig holds storage configuration for the embedded IdP.
|
// EmbeddedStorageConfig holds storage configuration for the embedded IdP.
|
||||||
@@ -175,6 +178,7 @@ func (c *EmbeddedIdPConfig) ToYAMLConfig() (*dex.YAMLConfig, error) {
|
|||||||
},
|
},
|
||||||
OAuth2: dex.OAuth2{
|
OAuth2: dex.OAuth2{
|
||||||
SkipApprovalScreen: true,
|
SkipApprovalScreen: true,
|
||||||
|
GrantTypes: c.GrantTypes,
|
||||||
},
|
},
|
||||||
Frontend: dex.Frontend{
|
Frontend: dex.Frontend{
|
||||||
Issuer: "NetBird",
|
Issuer: "NetBird",
|
||||||
|
|||||||
@@ -1606,7 +1606,8 @@ func (s *SqlStore) getAccount(ctx context.Context, accountID string) (*types.Acc
|
|||||||
settings_routing_peer_dns_resolution_enabled, settings_dns_domain, settings_network_range,
|
settings_routing_peer_dns_resolution_enabled, settings_dns_domain, settings_network_range,
|
||||||
settings_network_range_v6, settings_ipv6_enabled_groups, settings_lazy_connection_enabled,
|
settings_network_range_v6, settings_ipv6_enabled_groups, settings_lazy_connection_enabled,
|
||||||
settings_local_mfa_enabled, settings_metrics_push_enabled, settings_agent_network_only,
|
settings_local_mfa_enabled, settings_metrics_push_enabled, settings_agent_network_only,
|
||||||
settings_dashboard_features,
|
settings_dashboard_features, settings_auto_update_version, settings_auto_update_always,
|
||||||
|
settings_peer_expose_enabled, settings_peer_expose_groups,
|
||||||
-- Embedded ExtraSettings
|
-- Embedded ExtraSettings
|
||||||
settings_extra_peer_approval_enabled, settings_extra_user_approval_required,
|
settings_extra_peer_approval_enabled, settings_extra_user_approval_required,
|
||||||
settings_extra_integrated_validator, settings_extra_integrated_validator_groups
|
settings_extra_integrated_validator, settings_extra_integrated_validator_groups
|
||||||
@@ -1632,6 +1633,10 @@ func (s *SqlStore) getAccount(ctx context.Context, accountID string) (*types.Acc
|
|||||||
sMetricsPushEnabled sql.NullBool
|
sMetricsPushEnabled sql.NullBool
|
||||||
sAgentNetworkOnly sql.NullBool
|
sAgentNetworkOnly sql.NullBool
|
||||||
sDashboardFeatures sql.NullString
|
sDashboardFeatures sql.NullString
|
||||||
|
autoUpdateVersion sql.NullString
|
||||||
|
autoUpdateAlways sql.NullBool
|
||||||
|
peerExposeEnabled sql.NullBool
|
||||||
|
peerExposeGroups sql.NullString
|
||||||
sExtraPeerApprovalEnabled sql.NullBool
|
sExtraPeerApprovalEnabled sql.NullBool
|
||||||
sExtraUserApprovalRequired sql.NullBool
|
sExtraUserApprovalRequired sql.NullBool
|
||||||
sExtraIntegratedValidator sql.NullString
|
sExtraIntegratedValidator sql.NullString
|
||||||
@@ -1655,7 +1660,8 @@ func (s *SqlStore) getAccount(ctx context.Context, accountID string) (*types.Acc
|
|||||||
&sRoutingPeerDNSResolutionEnabled, &sDNSDomain, &sNetworkRange,
|
&sRoutingPeerDNSResolutionEnabled, &sDNSDomain, &sNetworkRange,
|
||||||
&sNetworkRangeV6, &sIPv6EnabledGroups, &sLazyConnectionEnabled,
|
&sNetworkRangeV6, &sIPv6EnabledGroups, &sLazyConnectionEnabled,
|
||||||
&sLocalMFAEnabled, &sMetricsPushEnabled, &sAgentNetworkOnly,
|
&sLocalMFAEnabled, &sMetricsPushEnabled, &sAgentNetworkOnly,
|
||||||
&sDashboardFeatures,
|
&sDashboardFeatures, &autoUpdateVersion, &autoUpdateAlways,
|
||||||
|
&peerExposeEnabled, &peerExposeGroups,
|
||||||
&sExtraPeerApprovalEnabled, &sExtraUserApprovalRequired,
|
&sExtraPeerApprovalEnabled, &sExtraUserApprovalRequired,
|
||||||
&sExtraIntegratedValidator, &sExtraIntegratedValidatorGroups,
|
&sExtraIntegratedValidator, &sExtraIntegratedValidatorGroups,
|
||||||
)
|
)
|
||||||
@@ -1747,6 +1753,18 @@ func (s *SqlStore) getAccount(ctx context.Context, accountID string) (*types.Acc
|
|||||||
if sIPv6EnabledGroups.Valid {
|
if sIPv6EnabledGroups.Valid {
|
||||||
_ = json.Unmarshal([]byte(sIPv6EnabledGroups.String), &account.Settings.IPv6EnabledGroups)
|
_ = json.Unmarshal([]byte(sIPv6EnabledGroups.String), &account.Settings.IPv6EnabledGroups)
|
||||||
}
|
}
|
||||||
|
if autoUpdateAlways.Valid {
|
||||||
|
account.Settings.AutoUpdateAlways = autoUpdateAlways.Bool
|
||||||
|
}
|
||||||
|
if autoUpdateVersion.Valid {
|
||||||
|
account.Settings.AutoUpdateVersion = autoUpdateVersion.String
|
||||||
|
}
|
||||||
|
if peerExposeEnabled.Valid {
|
||||||
|
account.Settings.PeerExposeEnabled = peerExposeEnabled.Bool
|
||||||
|
}
|
||||||
|
if peerExposeGroups.Valid {
|
||||||
|
_ = json.Unmarshal([]byte(peerExposeGroups.String), &account.Settings.PeerExposeGroups)
|
||||||
|
}
|
||||||
|
|
||||||
if sExtraPeerApprovalEnabled.Valid {
|
if sExtraPeerApprovalEnabled.Valid {
|
||||||
account.Settings.Extra.PeerApprovalEnabled = sExtraPeerApprovalEnabled.Bool
|
account.Settings.Extra.PeerApprovalEnabled = sExtraPeerApprovalEnabled.Bool
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ import (
|
|||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
|
"reflect"
|
||||||
"runtime"
|
"runtime"
|
||||||
"sort"
|
"sort"
|
||||||
"sync"
|
"sync"
|
||||||
@@ -34,6 +35,7 @@ import (
|
|||||||
"github.com/netbirdio/netbird/management/server/util"
|
"github.com/netbirdio/netbird/management/server/util"
|
||||||
nbroute "github.com/netbirdio/netbird/route"
|
nbroute "github.com/netbirdio/netbird/route"
|
||||||
"github.com/netbirdio/netbird/shared/management/status"
|
"github.com/netbirdio/netbird/shared/management/status"
|
||||||
|
"github.com/netbirdio/netbird/shared/testing_helpers"
|
||||||
"github.com/netbirdio/netbird/util/crypt"
|
"github.com/netbirdio/netbird/util/crypt"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -296,6 +298,53 @@ func Test_SaveAccount(t *testing.T) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func Test_AccountSettings_SaveAndRetrieve(t *testing.T) {
|
||||||
|
if runtime.GOOS == "windows" {
|
||||||
|
t.Skip("The SQLite store is not properly supported by Windows yet")
|
||||||
|
}
|
||||||
|
|
||||||
|
populateFields := testing_helpers.NewPopulateFields().WithCustomFieldSetter(
|
||||||
|
reflect.PointerTo(reflect.TypeOf(types.ExtraSettings{})), func(this *testing_helpers.PopulateFields, field reflect.Value) (int, error) {
|
||||||
|
es := types.ExtraSettings{}
|
||||||
|
reflectedEs := reflect.ValueOf(&es).Elem()
|
||||||
|
n, err := this.PopulateAll(reflectedEs)
|
||||||
|
if err != nil {
|
||||||
|
return n, err
|
||||||
|
}
|
||||||
|
field.Set(reflectedEs.Addr())
|
||||||
|
return n, nil
|
||||||
|
}).WithCustomFieldSetter(
|
||||||
|
reflect.PointerTo(reflect.TypeOf(types.DashboardFeatures{})), func(this *testing_helpers.PopulateFields, field reflect.Value) (int, error) {
|
||||||
|
t := true
|
||||||
|
df := types.DashboardFeatures{AgentNetwork: &t}
|
||||||
|
reflectedDf := reflect.ValueOf(&df).Elem()
|
||||||
|
field.Set(reflectedDf.Addr())
|
||||||
|
return 1, nil
|
||||||
|
}).WithSkippedTag("gorm", "-")
|
||||||
|
|
||||||
|
runTestForAllEngines(t, "", func(t *testing.T, store Store) {
|
||||||
|
account := newAccountWithId(context.Background(), "account_id", "testuser", "")
|
||||||
|
setupKey, _ := types.GenerateDefaultSetupKey()
|
||||||
|
account.SetupKeys[setupKey.Key] = setupKey
|
||||||
|
|
||||||
|
settings := types.Settings{}
|
||||||
|
numOfExportedFields, err := populateFields.PopulateAll(reflect.ValueOf(&settings).Elem())
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, 27, numOfExportedFields)
|
||||||
|
account.Settings = &settings
|
||||||
|
|
||||||
|
err = store.SaveAccount(context.Background(), account)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
accountFromDb, err := store.GetAccount(context.Background(), account.Id)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, accountFromDb)
|
||||||
|
assert.NotNil(t, accountFromDb.Settings)
|
||||||
|
|
||||||
|
assert.True(t, reflect.DeepEqual(&settings, accountFromDb.Settings), "created settings and settings retrieved from the db should match")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
func TestSqlite_DeleteAccount(t *testing.T) {
|
func TestSqlite_DeleteAccount(t *testing.T) {
|
||||||
if runtime.GOOS == "windows" {
|
if runtime.GOOS == "windows" {
|
||||||
t.Skip("The SQLite store is not properly supported by Windows yet")
|
t.Skip("The SQLite store is not properly supported by Windows yet")
|
||||||
|
|||||||
@@ -51,7 +51,10 @@ func (l *Listener) Listen(acceptFn func(conn relaylistener.Conn)) error {
|
|||||||
|
|
||||||
log.Infof("QUIC client connected from: %s", session.RemoteAddr())
|
log.Infof("QUIC client connected from: %s", session.RemoteAddr())
|
||||||
conn := NewConn(session)
|
conn := NewConn(session)
|
||||||
acceptFn(conn)
|
// Run the accept handler (which performs the pre-auth handshake) in its
|
||||||
|
// own goroutine so a slow or stalled handshake cannot block accepting
|
||||||
|
// further connections.
|
||||||
|
go acceptFn(conn)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,101 @@
|
|||||||
|
package testing_helpers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net/netip"
|
||||||
|
"reflect"
|
||||||
|
)
|
||||||
|
|
||||||
|
type PopulateFields struct {
|
||||||
|
CustomFieldSetters map[reflect.Type]func(this *PopulateFields, field reflect.Value) (int, error)
|
||||||
|
TagsToSkip map[string]string
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewPopulateFields() *PopulateFields {
|
||||||
|
return &PopulateFields{CustomFieldSetters: defaultCustomFieldSetters(), TagsToSkip: make(map[string]string)}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *PopulateFields) WithCustomFieldSetter(t reflect.Type, f func(this *PopulateFields, field reflect.Value) (int, error)) *PopulateFields {
|
||||||
|
p.CustomFieldSetters[t] = f
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *PopulateFields) WithSkippedTag(tag, value string) *PopulateFields {
|
||||||
|
p.TagsToSkip[tag] = value
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *PopulateFields) PopulateAll(v reflect.Value) (int, error) {
|
||||||
|
typ := v.Type()
|
||||||
|
totalExportedFields := 0
|
||||||
|
for i := 0; i < typ.NumField(); i++ {
|
||||||
|
f := typ.Field(i)
|
||||||
|
if f.PkgPath != "" { // unexported
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if p.skippedTagPresent(f.Tag) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
numOfExportedFields, err := p.setNonZero(v.Field(i))
|
||||||
|
totalExportedFields += numOfExportedFields
|
||||||
|
if err != nil {
|
||||||
|
return totalExportedFields, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return totalExportedFields, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// setNonZero assigns a deterministic non-zero value to a field based on its kind,
|
||||||
|
// recursing into nested structs and populating one element of slice fields.
|
||||||
|
func (p *PopulateFields) setNonZero(field reflect.Value) (int, error) {
|
||||||
|
if f, ok := p.CustomFieldSetters[field.Type()]; ok {
|
||||||
|
return f(p, field)
|
||||||
|
}
|
||||||
|
|
||||||
|
switch field.Kind() {
|
||||||
|
case reflect.String:
|
||||||
|
field.SetString("non-zero")
|
||||||
|
case reflect.Bool:
|
||||||
|
field.SetBool(true)
|
||||||
|
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
|
||||||
|
field.SetInt(7)
|
||||||
|
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
|
||||||
|
field.SetUint(7)
|
||||||
|
case reflect.Float32, reflect.Float64:
|
||||||
|
field.SetFloat(7)
|
||||||
|
case reflect.Struct:
|
||||||
|
n, err := p.PopulateAll(field)
|
||||||
|
return n + 1, err
|
||||||
|
case reflect.Slice:
|
||||||
|
s := reflect.MakeSlice(field.Type(), 1, 1)
|
||||||
|
_, err := p.setNonZero(s.Index(0))
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
field.Set(s)
|
||||||
|
default:
|
||||||
|
return 0, fmt.Errorf("unhandled field kind %s; extend setNonZero", field.Kind())
|
||||||
|
}
|
||||||
|
|
||||||
|
return 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func defaultCustomFieldSetters() map[reflect.Type]func(this *PopulateFields, field reflect.Value) (int, error) {
|
||||||
|
return map[reflect.Type]func(this *PopulateFields, field reflect.Value) (int, error){
|
||||||
|
reflect.TypeOf(netip.Prefix{}): func(_ *PopulateFields, field reflect.Value) (int, error) {
|
||||||
|
field.Set(reflect.ValueOf(netip.MustParsePrefix("10.0.0.0/24")))
|
||||||
|
return 1, nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *PopulateFields) skippedTagPresent(t reflect.StructTag) bool {
|
||||||
|
for tag, value := range p.TagsToSkip {
|
||||||
|
if v := t.Get(tag); v == value {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user