device: further add, revise API for on-demand configuration of peers

Updates tailscale/tailscale#17858
Updates tailscale/corp#35603

Signed-off-by: Brad Fitzpatrick <bradfitz@tailscale.com>
This commit is contained in:
Brad Fitzpatrick 2025-11-26 13:10:21 -08:00 committed by 世界
parent e924a91e99
commit f69b24781e
No known key found for this signature in database
GPG key ID: CD109927C34A63C4
5 changed files with 220 additions and 12 deletions

View file

@ -55,6 +55,25 @@ func commonBits(ip1, ip2 []byte) uint8 {
} }
} }
func commonBits4(ip1 []byte, ip2 [4]byte) uint8 {
a := binary.BigEndian.Uint32(ip1)
b := binary.BigEndian.Uint32(ip2[:])
return uint8(bits.LeadingZeros32(a ^ b))
}
func commonBits6(ip1 []byte, ip2 [16]byte) uint8 {
a := binary.BigEndian.Uint64(ip1)
b := binary.BigEndian.Uint64(ip2[:])
x := a ^ b
if x != 0 {
return uint8(bits.LeadingZeros64(x))
}
a = binary.BigEndian.Uint64(ip1[8:])
b = binary.BigEndian.Uint64(ip2[8:])
x = a ^ b
return 64 + uint8(bits.LeadingZeros64(x))
}
func (node *trieEntry) addToPeerEntries() { func (node *trieEntry) addToPeerEntries() {
node.perPeerElem = node.peer.trieEntries.PushBack(node) node.perPeerElem = node.peer.trieEntries.PushBack(node)
} }
@ -188,7 +207,37 @@ func (trie parentIndirection) insert(ip []byte, cidr uint8, peer *Peer) {
} }
} }
func (node *trieEntry) lookup(ip []byte) *Peer { func (node *trieEntry) lookup4(ip [4]byte) *Peer {
var found *Peer
for node != nil && commonBits4(node.bits, ip) >= node.cidr {
if node.peer != nil {
found = node.peer
}
if node.bitAtByte == 4 {
break
}
bit := (ip[node.bitAtByte] >> node.bitAtShift) & 1
node = node.child[bit]
}
return found
}
func (node *trieEntry) lookup6(ip [16]byte) *Peer {
var found *Peer
for node != nil && commonBits6(node.bits, ip) >= node.cidr {
if node.peer != nil {
found = node.peer
}
if node.bitAtByte == 16 {
break
}
bit := (ip[node.bitAtByte] >> node.bitAtShift) & 1
node = node.child[bit]
}
return found
}
func (node *trieEntry) lookup(ip net.IP) *Peer {
var found *Peer var found *Peer
size := uint8(len(ip)) size := uint8(len(ip))
for node != nil && commonBits(node.bits, ip) >= node.cidr { for node != nil && commonBits(node.bits, ip) >= node.cidr {
@ -208,6 +257,9 @@ type AllowedIPs struct {
mu sync.RWMutex mu sync.RWMutex
ipv4 *trieEntry ipv4 *trieEntry
ipv6 *trieEntry ipv6 *trieEntry
peerByIPPacketFunc PeerByIPPacketFunc // if non-nil, called to look up peers by IP
device *Device // back-reference to parent device; non-nil only if peerByIPPacketFunc is set
} }
func (table *AllowedIPs) EntriesForPeer(peer *Peer, cb func(prefix netip.Prefix) bool) { func (table *AllowedIPs) EntriesForPeer(peer *Peer, cb func(prefix netip.Prefix) bool) {
@ -322,9 +374,53 @@ func (table *AllowedIPs) insertLocked(prefix netip.Prefix, peer *Peer) {
} }
} }
// LookupFromPacket looks up the peer to which an outbound IP packet should be
// sent. It lives on [AllowedIPs] for legacy/structural reasons: historically
// WireGuard's only peer-selection mechanism was the AllowedIPs trie, and the
// send path already had a reference to the table. When a [PeerByIPPacketFunc]
// has been registered via [Device.SetPeerByIPPacketFunc], that callback is used
// instead of the trie and the AllowedIPs table is not consulted at all.
//
// When no callback is registered, only dst is used (standard WireGuard
// AllowedIPs trie lookup). When a callback is registered, all three
// parameters are forwarded to it; see [PeerByIPPacketFunc] for details.
func (table *AllowedIPs) LookupFromPacket(src, dst netip.Addr, ipPkt []byte) *Peer {
table.mu.RLock()
if f := table.peerByIPPacketFunc; f != nil {
device := table.device
table.mu.RUnlock()
if pubk, ok := f(src, dst, ipPkt); ok {
return device.LookupPeer(pubk)
}
return nil
}
defer table.mu.RUnlock()
switch {
case dst.Is6():
return table.ipv6.lookup6(dst.As16())
case dst.Is4():
return table.ipv4.lookup4(dst.As4())
default:
panic(errors.New("looking up unknown address type"))
}
}
// Deprecated: Lookup is only used by legacy tests. It does not call
// [PeerByIPPacketFunc]; use [AllowedIPs.LookupFromPacket] for production lookups.
func (table *AllowedIPs) Lookup(ip []byte) *Peer { func (table *AllowedIPs) Lookup(ip []byte) *Peer {
table.mu.RLock() table.mu.RLock()
defer table.mu.RUnlock() defer table.mu.RUnlock()
return table.lookupLocked(ip)
}
// lookupLocked looks up the peer associated with the given IP address.
// It assumes the caller holds the read lock (or doesn't hold it, but also
// doesn't concurrently mutate AllowedIP).
//
// It returns nil if no peer is associated with the given IP address.
func (table *AllowedIPs) lookupLocked(ip []byte) *Peer {
switch len(ip) { switch len(ip) {
case net.IPv6len: case net.IPv6len:
return table.ipv6.lookup(ip) return table.ipv6.lookup(ip)
@ -334,3 +430,62 @@ func (table *AllowedIPs) Lookup(ip []byte) *Peer {
panic(errors.New("looking up unknown address type")) panic(errors.New("looking up unknown address type"))
} }
} }
// AllowedPeerSourceIP reports whether the given source IP address is allowed
// for the given peer.
func (peer *Peer) AllowedPeerSourceIP(src netip.Addr) bool {
if f := peer.state.testAllowedIP.Load(); f != nil {
return (*f)(src)
}
table := &peer.device.allowedips
table.mu.RLock()
defer table.mu.RUnlock()
switch {
case src.Is6():
return table.ipv6.lookup6(src.As16()) == peer
case src.Is4():
return table.ipv4.lookup4(src.As4()) == peer
}
return false
}
// fakePeer is a zero Peer used only as a placeholder in tries used by mkIPInCIDRsTestFunc.
var fakePeer Peer
// mkIPInCIDRsTestFunc returns a function that tests whether an IP address is
// contained in any of the given CIDRs.
func mkIPInCIDRsTestFunc(cidrs []netip.Prefix) func(netip.Addr) bool {
if len(cidrs) == 0 {
return func(netip.Addr) bool { return false }
}
if len(cidrs) == 1 {
return func(addr netip.Addr) bool { return cidrs[0].Contains(addr) }
}
if len(cidrs) <= 4 {
// For small numbers of CIDRs, just do a linear search. The trie construction
// is more expensive than the linear search, and the test function is faster
// than the trie lookup, so this is a net win.
return func(addr netip.Addr) bool {
for _, c := range cidrs {
if c.Contains(addr) {
return true
}
}
return false
}
}
// Make a trie for faster lookups. We use a dummy Peer.
var a AllowedIPs
for _, c := range cidrs {
a.Insert(c, &fakePeer)
}
return func(addr netip.Addr) bool {
switch {
case addr.Is4():
return a.ipv4.lookup4(addr.As4()) == &fakePeer
default:
return a.ipv6.lookup6(addr.As16()) == &fakePeer
}
}
}

View file

@ -368,8 +368,8 @@ func (device *Device) LookupPeer(pk NoisePublicKey) *Peer {
return p return p
} }
allowedIPs := lookupFunc(pk) conf, ok := lookupFunc(pk)
if allowedIPs == nil { if !ok || conf == nil {
return nil return nil
} }
@ -383,8 +383,11 @@ func (device *Device) LookupPeer(pk NoisePublicKey) *Peer {
device.log.Errorf("Failed to create peer: %v", err) device.log.Errorf("Failed to create peer: %v", err)
return nil return nil
} }
p.SetAllowedIPs(allowedIPs) p.SetAllowedIPs(conf.AllowedIPs)
p.deleteOnIdle = true p.deleteOnIdle = true
if conf.Endpoint != nil {
p.SetEndpointFromPacket(conf.Endpoint)
}
p.Start() p.Start()
return p return p
} }
@ -443,6 +446,17 @@ func (device *Device) RemoveMatchingPeers(shouldRemove func(NoisePublicKey) bool
return numRemoved return numRemoved
} }
// NewPeerConfig are the configuration parameters for a new peer created via a
// [PeerLookupFunc] func.
type NewPeerConfig struct {
// AllowedIPs is the initial set of allowed IPs for the new peer.
AllowedIPs []netip.Prefix
// Endpoint, if non-nil, sets the initial endpoint for newly
// created peers.
Endpoint conn.Endpoint
}
// PeerLookupFunc is the type of function used to look up peers by public key // PeerLookupFunc is the type of function used to look up peers by public key
// when receiving packets for unknown peers. // when receiving packets for unknown peers.
// //
@ -452,7 +466,21 @@ func (device *Device) RemoveMatchingPeers(shouldRemove func(NoisePublicKey) bool
// with the provided allowed IPs. // with the provided allowed IPs.
// //
// See [Device.SetPeerLookupFunc] and [Device.LookupPeer]. // See [Device.SetPeerLookupFunc] and [Device.LookupPeer].
type PeerLookupFunc func(NoisePublicKey) (allowedIPs []netip.Prefix) type PeerLookupFunc func(NoisePublicKey) (_ *NewPeerConfig, ok bool)
// PeerByIPPacketFunc is the type of function used to look up a peer to send to
// for a given src/dst IP pair. The ipPkt parameter is the raw IP packet being
// routed; callers needing transport-layer ports or other header fields may parse
// them from ipPkt, but must handle IP fragmentation (ports may be absent on
// non-first fragments) and protocols that do not use ports (e.g. ICMP).
//
// Except for experimental use cases, dst is the only address
// that should be relied upon when looking up a peer.
//
// If it returns ok=false, the peer is not known.
//
// See [Device.SetPeerByIPPacketFunc] and [Device.SetPeerLookupFunc].
type PeerByIPPacketFunc func(src, dst netip.Addr, ipPkt []byte) (_ NoisePublicKey, ok bool)
// SetPeerLookupFunc sets the function used to look up peers by public key // SetPeerLookupFunc sets the function used to look up peers by public key
// when receiving packets for unknown peers. // when receiving packets for unknown peers.
@ -462,6 +490,15 @@ func (device *Device) SetPeerLookupFunc(f PeerLookupFunc) {
device.peers.lookupFunc = f device.peers.lookupFunc = f
} }
// SetPeerByIPPacketFunc sets the function used to look up peers by IP address
// when sending packets to unknown peers.
func (device *Device) SetPeerByIPPacketFunc(f PeerByIPPacketFunc) {
device.allowedips.mu.Lock()
defer device.allowedips.mu.Unlock()
device.allowedips.peerByIPPacketFunc = f
device.allowedips.device = device
}
func (device *Device) Close() { func (device *Device) Close() {
device.state.Lock() device.state.Lock()
defer device.state.Unlock() defer device.state.Unlock()

View file

@ -57,6 +57,11 @@ type Peer struct {
sync.Mutex // protects against concurrent Start/Stop, and fields below sync.Mutex // protects against concurrent Start/Stop, and fields below
allowedIPs []netip.Prefix allowedIPs []netip.Prefix
// testAllowedIP, if non-nil, is used to test whether the peer is
// allowed to send a packet from the given IP address. It can be read
// without locking, but must be set with the state mutex locked.
testAllowedIP atomic.Pointer[func(netip.Addr) bool]
} }
queue struct { queue struct {
@ -138,7 +143,12 @@ func (p *Peer) SetAllowedIPs(allowedIPs []netip.Prefix) {
return return
} }
p.device.allowedips.setPeerPrefixes(p, allowedIPs) p.device.allowedips.setPeerPrefixes(p, allowedIPs)
p.state.allowedIPs = slices.Clone(allowedIPs) // avoid retaining caller's slice
allowedIPs = slices.Clone(allowedIPs) // avoid retaining caller's slice
p.state.allowedIPs = allowedIPs
f := mkIPInCIDRsTestFunc(allowedIPs)
p.state.testAllowedIP.Store(&f)
} }
// SendBuffers sends buffers to peer. WireGuard packet data in each element of // SendBuffers sends buffers to peer. WireGuard packet data in each element of

View file

@ -9,6 +9,7 @@ import (
"encoding/binary" "encoding/binary"
"errors" "errors"
"net" "net"
"net/netip"
"sync" "sync"
"time" "time"
@ -482,7 +483,8 @@ func (peer *Peer) RoutineSequentialReceiver(maxBatchSize int) {
} }
elem.packet = elem.packet[:length] elem.packet = elem.packet[:length]
src := elem.packet[IPv4offsetSrc : IPv4offsetSrc+net.IPv4len] src := elem.packet[IPv4offsetSrc : IPv4offsetSrc+net.IPv4len]
if device.allowedips.Lookup(src) != peer { srcAddr, _ := netip.AddrFromSlice(src)
if !peer.AllowedPeerSourceIP(srcAddr) {
device.log.Verbosef("IPv4 packet with disallowed source address from %v", peer) device.log.Verbosef("IPv4 packet with disallowed source address from %v", peer)
continue continue
} }
@ -499,7 +501,8 @@ func (peer *Peer) RoutineSequentialReceiver(maxBatchSize int) {
} }
elem.packet = elem.packet[:length] elem.packet = elem.packet[:length]
src := elem.packet[IPv6offsetSrc : IPv6offsetSrc+net.IPv6len] src := elem.packet[IPv6offsetSrc : IPv6offsetSrc+net.IPv6len]
if device.allowedips.Lookup(src) != peer { srcAddr, _ := netip.AddrFromSlice(src)
if !peer.AllowedPeerSourceIP(srcAddr) {
device.log.Verbosef("IPv6 packet with disallowed source address from %v", peer) device.log.Verbosef("IPv6 packet with disallowed source address from %v", peer)
continue continue
} }

View file

@ -9,6 +9,7 @@ import (
"encoding/binary" "encoding/binary"
"errors" "errors"
"net" "net"
"net/netip"
"os" "os"
"sync" "sync"
"time" "time"
@ -263,15 +264,17 @@ func (device *Device) RoutineReadFromTUN() {
if len(elem.packet) < ipv4.HeaderLen { if len(elem.packet) < ipv4.HeaderLen {
continue continue
} }
dst := elem.packet[IPv4offsetDst : IPv4offsetDst+net.IPv4len] src := netip.AddrFrom4([4]byte(elem.packet[IPv4offsetSrc : IPv4offsetSrc+net.IPv4len]))
peer = device.allowedips.Lookup(dst) dst := netip.AddrFrom4([4]byte(elem.packet[IPv4offsetDst : IPv4offsetDst+net.IPv4len]))
peer = device.allowedips.LookupFromPacket(src, dst, elem.packet)
case 6: case 6:
if len(elem.packet) < ipv6.HeaderLen { if len(elem.packet) < ipv6.HeaderLen {
continue continue
} }
dst := elem.packet[IPv6offsetDst : IPv6offsetDst+net.IPv6len] src := netip.AddrFrom16([16]byte(elem.packet[IPv6offsetSrc : IPv6offsetSrc+net.IPv6len]))
peer = device.allowedips.Lookup(dst) dst := netip.AddrFrom16([16]byte(elem.packet[IPv6offsetDst : IPv6offsetDst+net.IPv6len]))
peer = device.allowedips.LookupFromPacket(src, dst, elem.packet)
default: default:
device.log.Verbosef("Received packet with unknown IP version") device.log.Verbosef("Received packet with unknown IP version")