Compare commits

..
Author SHA1 Message Date
JackDoan 71c5297134 improve lighthouse addrMap correctness for multi-address certs 2026-09-02 14:37:27 -05:00
80 changed files with 342 additions and 102 deletions
+11 -2
View File
@@ -79,7 +79,13 @@ func (b *Bits) clearRange(startPos, count uint64) uint64 {
// handle the potential partial word before pos becomes u64 aligned // handle the potential partial word before pos becomes u64 aligned
word := pos >> 6 word := pos >> 6
bit := pos & 63 bit := pos & 63
take := min(min(uint64(64)-bit, remaining), b.length-pos) take := uint64(64) - bit
if take > remaining {
take = remaining
}
if take > b.length-pos {
take = b.length - pos
}
var mask uint64 var mask uint64
if take == 64 { if take == 64 {
mask = math.MaxUint64 mask = math.MaxUint64
@@ -183,7 +189,10 @@ func (b *Bits) Update(l *slog.Logger, i uint64) bool {
func (b *Bits) updateSlow(l *slog.Logger, i uint64) bool { func (b *Bits) updateSlow(l *slog.Logger, i uint64) bool {
// If i is a jump, adjust the window, record lost, update current, and return true // If i is a jump, adjust the window, record lost, update current, and return true
if i > b.current { if i > b.current {
end := min(i, b.current+b.length) end := i
if end > b.current+b.length {
end = b.current + b.length
}
count := end - b.current count := end - b.current
startPos := (b.current + 1) & b.lengthMask startPos := (b.current + 1) & b.lengthMask
+1 -1
View File
@@ -120,7 +120,7 @@ func TestCertificate_SignP256_AlwaysNormalized(t *testing.T) {
pub := elliptic.Marshal(elliptic.P256(), priv.PublicKey.X, priv.PublicKey.Y) pub := elliptic.Marshal(elliptic.P256(), priv.PublicKey.X, priv.PublicKey.Y)
rawPriv := priv.D.FillBytes(make([]byte, 32)) rawPriv := priv.D.FillBytes(make([]byte, 32))
for i := range 1000 { for i := 0; i < 1000; i++ {
if i&1 == 1 { if i&1 == 1 {
tbs.Version = Version1 tbs.Version = Version1
} else { } else {
+4 -4
View File
@@ -151,7 +151,7 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
var groups []string var groups []string
if *cf.groups != "" { if *cf.groups != "" {
for rg := range strings.SplitSeq(*cf.groups, ",") { for _, rg := range strings.Split(*cf.groups, ",") {
g := strings.TrimSpace(rg) g := strings.TrimSpace(rg)
if g != "" { if g != "" {
groups = append(groups, g) groups = append(groups, g)
@@ -171,7 +171,7 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
} }
if *cf.networks != "" { if *cf.networks != "" {
for rs := range strings.SplitSeq(*cf.networks, ",") { for _, rs := range strings.Split(*cf.networks, ",") {
rs := strings.Trim(rs, " ") rs := strings.Trim(rs, " ")
if rs != "" { if rs != "" {
n, err := netip.ParsePrefix(rs) n, err := netip.ParsePrefix(rs)
@@ -193,7 +193,7 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
} }
if *cf.unsafeNetworks != "" { if *cf.unsafeNetworks != "" {
for rs := range strings.SplitSeq(*cf.unsafeNetworks, ",") { for _, rs := range strings.Split(*cf.unsafeNetworks, ",") {
rs := strings.Trim(rs, " ") rs := strings.Trim(rs, " ")
if rs != "" { if rs != "" {
n, err := netip.ParsePrefix(rs) n, err := netip.ParsePrefix(rs)
@@ -221,7 +221,7 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
if !isP11 && *cf.encryption { if !isP11 && *cf.encryption {
passphrase = []byte(os.Getenv("NEBULA_CA_PASSPHRASE")) passphrase = []byte(os.Getenv("NEBULA_CA_PASSPHRASE"))
if len(passphrase) == 0 { if len(passphrase) == 0 {
for range 5 { for i := 0; i < 5; i++ {
errOut.Write([]byte("Enter passphrase: ")) errOut.Write([]byte("Enter passphrase: "))
passphrase, err = pr.ReadPassword() passphrase, err = pr.ReadPassword()
+1
View File
@@ -1,4 +1,5 @@
//go:build !windows //go:build !windows
// +build !windows
package main package main
+4 -4
View File
@@ -146,7 +146,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
passphrase = []byte(os.Getenv("NEBULA_CA_PASSPHRASE")) passphrase = []byte(os.Getenv("NEBULA_CA_PASSPHRASE"))
if len(passphrase) == 0 { if len(passphrase) == 0 {
// ask for a passphrase until we get one // ask for a passphrase until we get one
for range 5 { for i := 0; i < 5; i++ {
errOut.Write([]byte("Enter passphrase: ")) errOut.Write([]byte("Enter passphrase: "))
passphrase, err = pr.ReadPassword() passphrase, err = pr.ReadPassword()
@@ -203,7 +203,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
} }
if *sf.networks != "" { if *sf.networks != "" {
for rs := range strings.SplitSeq(*sf.networks, ",") { for _, rs := range strings.Split(*sf.networks, ",") {
rs := strings.Trim(rs, " ") rs := strings.Trim(rs, " ")
if rs != "" { if rs != "" {
n, err := netip.ParsePrefix(rs) n, err := netip.ParsePrefix(rs)
@@ -228,7 +228,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
} }
if *sf.unsafeNetworks != "" { if *sf.unsafeNetworks != "" {
for rs := range strings.SplitSeq(*sf.unsafeNetworks, ",") { for _, rs := range strings.Split(*sf.unsafeNetworks, ",") {
rs := strings.Trim(rs, " ") rs := strings.Trim(rs, " ")
if rs != "" { if rs != "" {
n, err := netip.ParsePrefix(rs) n, err := netip.ParsePrefix(rs)
@@ -247,7 +247,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
var groups []string var groups []string
if *sf.groups != "" { if *sf.groups != "" {
for rg := range strings.SplitSeq(*sf.groups, ",") { for _, rg := range strings.Split(*sf.groups, ",") {
g := strings.TrimSpace(rg) g := strings.TrimSpace(rg)
if g != "" { if g != "" {
groups = append(groups, g) groups = append(groups, g)
+1
View File
@@ -1,4 +1,5 @@
//go:build !windows //go:build !windows
// +build !windows
package main package main
+1
View File
@@ -1,4 +1,5 @@
//go:build !windows //go:build !windows
// +build !windows
package main package main
+1
View File
@@ -1,4 +1,5 @@
//go:build !linux //go:build !linux
// +build !linux
package main package main
+9 -6
View File
@@ -5,7 +5,6 @@ import (
"errors" "errors"
"fmt" "fmt"
"log/slog" "log/slog"
"maps"
"math" "math"
"os" "os"
"os/signal" "os/signal"
@@ -155,7 +154,9 @@ func (c *C) ReloadConfig() {
defer c.reloadLock.Unlock() defer c.reloadLock.Unlock()
c.oldSettings = make(map[string]any) c.oldSettings = make(map[string]any)
maps.Copy(c.oldSettings, c.Settings) for k, v := range c.Settings {
c.oldSettings[k] = v
}
err := c.Load(c.path) err := c.Load(c.path)
if err != nil { if err != nil {
@@ -176,7 +177,9 @@ func (c *C) ReloadConfigString(raw string) error {
defer c.reloadLock.Unlock() defer c.reloadLock.Unlock()
c.oldSettings = make(map[string]any) c.oldSettings = make(map[string]any)
maps.Copy(c.oldSettings, c.Settings) for k, v := range c.Settings {
c.oldSettings[k] = v
}
err := c.LoadString(raw) err := c.LoadString(raw)
if err != nil { if err != nil {
@@ -213,7 +216,7 @@ func (c *C) GetStringSlice(k string, d []string) []string {
} }
v := make([]string, len(rv)) v := make([]string, len(rv))
for i := range v { for i := 0; i < len(v); i++ {
v[i] = fmt.Sprintf("%v", rv[i]) v[i] = fmt.Sprintf("%v", rv[i])
} }
@@ -307,8 +310,8 @@ func (c *C) IsSet(k string) bool {
} }
func (c *C) get(k string, v any) any { func (c *C) get(k string, v any) any {
parts := strings.SplitSeq(k, ".") parts := strings.Split(k, ".")
for p := range parts { for _, p := range parts {
m, ok := v.(map[string]any) m, ok := v.(map[string]any)
if !ok { if !ok {
return nil return nil
+1 -1
View File
@@ -97,7 +97,7 @@ func TestConnectionState_NextMessageCounter(t *testing.T) {
assert.Equal(t, RejectAfterMessages, cs.messageCounter.Load()) assert.Equal(t, RejectAfterMessages, cs.messageCounter.Load())
// Continued send attempts stay refused and the counter never wraps // Continued send attempts stay refused and the counter never wraps
for range 10 { for i := 0; i < 10; i++ {
_, ok = cs.NextMessageCounter() _, ok = cs.NextMessageCounter()
assert.False(t, ok) assert.False(t, ok)
} }
+1 -1
View File
@@ -247,7 +247,7 @@ func TestControl_ConcurrentStopAndStart(t *testing.T) {
c, _, _ := newReadyControl(t) c, _, _ := newReadyControl(t)
var wg sync.WaitGroup var wg sync.WaitGroup
for range 2 { for i := 0; i < 2; i++ {
wg.Go(func() { c.Stop() }) wg.Go(func() { c.Stop() })
} }
wg.Go(func() { _ = c.Start() }) wg.Go(func() { _ = c.Start() })
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing //go:build e2e_testing
// +build e2e_testing
package e2e package e2e
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing //go:build e2e_testing
// +build e2e_testing
package e2e package e2e
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing //go:build e2e_testing
// +build e2e_testing
package e2e package e2e
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing //go:build e2e_testing
// +build e2e_testing
package e2e package e2e
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing //go:build e2e_testing
// +build e2e_testing
package e2e package e2e
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing //go:build e2e_testing
// +build e2e_testing
package e2e package e2e
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing //go:build e2e_testing
// +build e2e_testing
package e2e package e2e
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing //go:build e2e_testing
// +build e2e_testing
package router package router
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing //go:build e2e_testing
// +build e2e_testing
package router package router
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing //go:build e2e_testing
// +build e2e_testing
package e2e package e2e
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing //go:build e2e_testing
// +build e2e_testing
package e2e package e2e
+1 -1
View File
@@ -22,7 +22,7 @@ func newFixedTicker(t *testing.T, l *slog.Logger, cacheLen int) *ConntrackCacheT
l: l, l: l,
cache: make(ConntrackCache, cacheLen), cache: make(ConntrackCache, cacheLen),
} }
for i := range cacheLen { for i := 0; i < cacheLen; i++ {
c.cache[Packet{LocalPort: uint16(i) + 1}] = struct{}{} c.cache[Packet{LocalPort: uint16(i) + 1}] = struct{}{}
} }
c.cacheTick.Store(1) // cacheV starts at 0, so Get() takes the reset path c.cacheTick.Store(1) // cacheV starts at 0, so Get() takes the reset path
+2 -2
View File
@@ -807,7 +807,7 @@ func (hm *HandshakeManager) beginHandshake(via ViaSender, packet []byte, h *head
} }
hm.sendHandshakeResponse(via, response, hostinfo, false) hm.sendHandshakeResponse(via, response, hostinfo, false)
hostinfo.remotes.RefreshFromHandshake(vpnAddrs) hostinfo.remotes.RefreshFromHandshake(vpnAddrs, remoteCert.Certificate.Version())
// Don't wait for UpdateWorker // Don't wait for UpdateWorker
if f.lightHouse.IsAnyLighthouseAddr(vpnAddrs) { if f.lightHouse.IsAnyLighthouseAddr(vpnAddrs) {
@@ -995,7 +995,7 @@ func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostIn
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore))) f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
} }
hostinfo.remotes.RefreshFromHandshake(vpnAddrs) hostinfo.remotes.RefreshFromHandshake(vpnAddrs, remoteCert.Certificate.Version())
f.metricHandshakes.Update(duration) f.metricHandshakes.Update(duration)
// Don't wait for UpdateWorker // Don't wait for UpdateWorker
+49 -16
View File
@@ -513,17 +513,15 @@ func (lh *LightHouse) QueryServer(vpnAddr netip.Addr) {
} }
func (lh *LightHouse) QueryCache(vpnAddrs []netip.Addr) *RemoteList { func (lh *LightHouse) QueryCache(vpnAddrs []netip.Addr) *RemoteList {
lh.RLock() rl, ok := lh.findRemoteList(vpnAddrs)
if v, ok := lh.addrMap[vpnAddrs[0]]; ok { if ok {
lh.RUnlock() return rl
return v
} }
lh.RUnlock()
lh.Lock() lh.Lock()
defer lh.Unlock() defer lh.Unlock()
// Add an entry if we don't already have one // Add an entry if we don't already have one
return lh.unlockedGetRemoteList(vpnAddrs) //todo CERT-V2 this contains addrmap lookups we could potentially skip return lh.unlockedGetRemoteList(vpnAddrs) //todo this re-calls unlockedFindRemoteList
} }
// queryAndPrepMessage is a lock helper on RemoteList, assisting the caller to build a lighthouse message containing // queryAndPrepMessage is a lock helper on RemoteList, assisting the caller to build a lighthouse message containing
@@ -669,26 +667,61 @@ func (lh *LightHouse) addCalculatedRemotes(vpnAddr netip.Addr) bool {
return len(calculatedV4) > 0 || len(calculatedV6) > 0 return len(calculatedV4) > 0 || len(calculatedV6) > 0
} }
func (lh *LightHouse) findRemoteList(vpnAddrs []netip.Addr) (*RemoteList, bool) {
lh.RLock()
defer lh.RUnlock()
return lh.unlockedFindRemoteList(vpnAddrs)
}
// unlockedFindRemoteList checks addrMap for each of vpnAddrs. It returns the first RemoteList found,
// and true if that RemoteList is present for all vpnAddrs.
// If false, it means the addrMap, and possibly the RemoteList, need to be corrected.
func (lh *LightHouse) unlockedFindRemoteList(vpnAddrs []netip.Addr) (*RemoteList, bool) {
var am *RemoteList
//todo: if a host with addresses A and B is "split", so it only has address A and a new host has only address B
//todo: we don't handly that correctly, I'm pretty sure.
missingOrDifferent := false
for _, addr := range vpnAddrs {
found, ok := lh.addrMap[addr]
if !ok {
missingOrDifferent = true
} else if am == nil {
am = found //the first list we find wins
} else if am != found {
missingOrDifferent = true
}
}
return am, !missingOrDifferent
}
// unlockedGetRemoteList assumes you have the lh lock // unlockedGetRemoteList assumes you have the lh lock
func (lh *LightHouse) unlockedGetRemoteList(allAddrs []netip.Addr) *RemoteList { func (lh *LightHouse) unlockedGetRemoteList(allAddrs []netip.Addr) *RemoteList {
// before we go and make a new remotelist, we need to make sure we don't have one for any of this set of vpnaddrs yet // before we go and make a new remotelist, we need to make sure we don't have one for any of this set of vpnaddrs yet
for i, addr := range allAddrs { am, ok := lh.unlockedFindRemoteList(allAddrs)
am, ok := lh.addrMap[addr]
if ok {
if i != 0 {
lh.addrMap[allAddrs[0]] = am
}
return am
}
}
am := NewRemoteList(allAddrs, lh.shouldAdd) // we failed to find any RemoteLists: make a new one, fill out the addrMap
if am == nil {
am = NewRemoteList(allAddrs, lh.shouldAdd)
for _, addr := range allAddrs { for _, addr := range allAddrs {
lh.addrMap[addr] = am lh.addrMap[addr] = am
} }
return am return am
} }
// we found one! Do we need to fix it?
if !ok {
am.Lock()
am.vpnAddrs = make([]netip.Addr, len(allAddrs))
copy(am.vpnAddrs, allAddrs)
am.Unlock()
for _, addr := range allAddrs {
lh.addrMap[addr] = am
}
}
return am
}
func (lh *LightHouse) shouldAdd(vpnAddrs []netip.Addr, to netip.Addr) bool { func (lh *LightHouse) shouldAdd(vpnAddrs []netip.Addr, to netip.Addr) bool {
allow := lh.GetRemoteAllowList().AllowAll(vpnAddrs, to) allow := lh.GetRemoteAllowList().AllowAll(vpnAddrs, to)
if lh.l.Enabled(context.Background(), logging.LevelTrace) { if lh.l.Enabled(context.Background(), logging.LevelTrace) {
+120
View File
@@ -738,3 +738,123 @@ func TestLighthouse_DeletesWork(t *testing.T) {
out = lh.Query(testHost) out = lh.Query(testHost)
assert.Nil(t, out) assert.Nil(t, out)
} }
// newLHHostUpdateV2 sends a v2-style HostUpdateNotification where the sending tunnel carries
// multiple vpn addrs (a dual-stack v2 cert). Details.VpnAddr is left blank like SendUpdate does.
func newLHHostUpdateV2(fromAddr netip.AddrPort, vpnAddrs []netip.Addr, addrs []netip.AddrPort, lhh *LightHouseHandler) {
req := &NebulaMeta{
Type: NebulaMeta_HostUpdateNotification,
Details: &NebulaMetaDetails{},
}
for _, v := range addrs {
if v.Addr().Is4() {
req.Details.V4AddrPorts = append(req.Details.V4AddrPorts, netAddrToProtoV4AddrPort(v.Addr(), v.Port()))
} else {
req.Details.V6AddrPorts = append(req.Details.V6AddrPorts, netAddrToProtoV6AddrPort(v.Addr(), v.Port()))
}
}
b, err := req.Marshal()
if err != nil {
panic(err)
}
lhh.HandleRequest(fromAddr, vpnAddrs, b, &testEncWriter{})
}
func newIssue1868Lighthouse(t *testing.T) (*LightHouse, *LightHouseHandler) {
l := test.NewLogger()
c := config.NewC(l)
c.Settings["lighthouse"] = map[string]any{"am_lighthouse": true}
c.Settings["listen"] = map[string]any{"port": 4242}
myVpnNet4 := netip.MustParsePrefix("10.128.0.1/24")
myVpnNet6 := netip.MustParsePrefix("fd00::1/64")
nt := new(bart.Lite)
nt.Insert(myVpnNet4)
nt.Insert(myVpnNet6)
cs := &CertState{
myVpnNetworks: []netip.Prefix{myVpnNet4, myVpnNet6},
myVpnNetworksTable: nt,
}
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
require.NoError(t, err)
lh.ifce = &mockEncWriter{}
return lh, lh.NewRequestHandler()
}
// Scenario A: host registers via v2 (both addrs), then a v1 handshake with the same host completes on
// the lighthouse (rehandshake after cert renewal, relay-initiated handshake, traffic to the LH's v4
// addr...). handshake_manager does QueryCache(vpnAddrs) + RefreshFromHandshake(vpnAddrs) with the v1
// cert's single address, which truncates RemoteList.vpnAddrs.
func TestLighthouse_Issue1868_V1HandshakeTruncatesVpnAddrs(t *testing.T) {
lh, lhh := newIssue1868Lighthouse(t)
hostV4 := netip.MustParseAddr("10.128.0.3")
hostV6 := netip.MustParseAddr("fd00::3")
hostUdp := netip.MustParseAddrPort("192.0.2.3:4242")
hostLan := netip.MustParseAddrPort("10.0.0.3:4242")
askerV4 := netip.MustParseAddr("10.128.0.2")
askerUdp := netip.MustParseAddrPort("192.0.2.2:4242")
// Boot: host handshakes with the LH using its v2 cert and sends an update
newLHHostUpdateV2(hostUdp, []netip.Addr{hostV4, hostV6}, []netip.AddrPort{hostLan}, lhh)
// Both addresses resolve
r := newLHHostRequest(askerUdp, askerV4, hostV4, lhh)
require.NotNil(t, r.msg, "v4 query should be answered")
assertIp4InArray(t, r.msg.Details.V4AddrPorts, hostLan)
r = newLHHostRequest(askerUdp, askerV4, hostV6, lhh)
require.NotNil(t, r.msg, "v6 query should be answered before the v1 handshake")
assertIp4InArray(t, r.msg.Details.V4AddrPorts, hostLan)
// Later: a v1 handshake with the same host completes on the LH. This is exactly what
// handshake_manager.go does on completion, with the v1 cert's single vpn addr.
rl := lh.QueryCache([]netip.Addr{hostV4})
rl.RefreshFromHandshake([]netip.Addr{hostV4}, cert.Version1)
// The host keeps sending v2 updates over its v2 tunnel too
newLHHostUpdateV2(hostUdp, []netip.Addr{hostV4, hostV6}, []netip.AddrPort{hostLan}, lhh)
r = newLHHostRequest(askerUdp, askerV4, hostV4, lhh)
require.NotNil(t, r.msg, "v4 query should still be answered")
assertIp4InArray(t, r.msg.Details.V4AddrPorts, hostLan)
r = newLHHostRequest(askerUdp, askerV4, hostV6, lhh)
if assert.NotNil(t, r.msg, "BUG: v6 query is silently dropped after a v1 handshake truncated RemoteList.vpnAddrs") {
assertIp4InArray(t, r.msg.Details.V4AddrPorts, hostLan)
}
}
// Scenario B: the LH first creates a RemoteList for the host keyed only by its v4 addr (a pending
// LH-initiated v1 handshake does QueryCache([v4]) in handleOutbound), then the host arrives with v2.
// unlockedGetRemoteList/QueryCache hit on allAddrs[0] and never add the v6 key to addrMap.
func TestLighthouse_Issue1868_V4OnlyListNeverGainsV6Key(t *testing.T) {
lh, lhh := newIssue1868Lighthouse(t)
hostV4 := netip.MustParseAddr("10.128.0.3")
hostV6 := netip.MustParseAddr("fd00::3")
hostUdp := netip.MustParseAddrPort("192.0.2.3:4242")
hostLan := netip.MustParseAddrPort("10.0.0.3:4242")
askerV4 := netip.MustParseAddr("10.128.0.2")
askerUdp := netip.MustParseAddrPort("192.0.2.2:4242")
// LH is a relay and someone asked it to relay to hostV4 while the host was offline:
// StartHandshake(hostV4) -> handleOutbound -> QueryCache([hostV4]) creates a v4-only list.
_ = lh.QueryCache([]netip.Addr{hostV4})
// Host boots and handshakes v2 with the LH (responder path does QueryCache + RefreshFromHandshake)
rl := lh.QueryCache([]netip.Addr{hostV4, hostV6})
rl.RefreshFromHandshake([]netip.Addr{hostV4, hostV6}, cert.Version2)
newLHHostUpdateV2(hostUdp, []netip.Addr{hostV4, hostV6}, []netip.AddrPort{hostLan}, lhh)
r := newLHHostRequest(askerUdp, askerV4, hostV4, lhh)
require.NotNil(t, r.msg, "v4 query should be answered")
assertIp4InArray(t, r.msg.Details.V4AddrPorts, hostLan)
r = newLHHostRequest(askerUdp, askerV4, hostV6, lhh)
if assert.NotNil(t, r.msg, "BUG: v6 query is silently dropped, addrMap never got the v6 key") {
assertIp4InArray(t, r.msg.Details.V4AddrPorts, hostLan)
}
}
+1
View File
@@ -1,4 +1,5 @@
//go:build boringcrypto //go:build boringcrypto
// +build boringcrypto
package noiseutil package noiseutil
+1
View File
@@ -1,4 +1,5 @@
//go:build boringcrypto //go:build boringcrypto
// +build boringcrypto
package noiseutil package noiseutil
+2 -2
View File
@@ -127,7 +127,7 @@ func TestPseudoSumIPv6MatchesReference(t *testing.T) {
func TestIPv4HdrChecksumMatchesReference(t *testing.T) { func TestIPv4HdrChecksumMatchesReference(t *testing.T) {
rng := rand.New(rand.NewSource(0x1791)) rng := rand.New(rand.NewSource(0x1791))
for _, hdrLen := range []int{20, 24, 40, 60} { for _, hdrLen := range []int{20, 24, 40, 60} {
for trial := range 200 { for trial := 0; trial < 200; trial++ {
hdr := make([]byte, hdrLen) hdr := make([]byte, hdrLen)
rng.Read(hdr) rng.Read(hdr)
hdr[0] = 0x40 | byte(hdrLen/4) hdr[0] = 0x40 | byte(hdrLen/4)
@@ -156,7 +156,7 @@ func TestIPv4HdrChecksumMatchesReference(t *testing.T) {
// the way a receiver does (pseudo-header + L4 must sum to all-ones). // the way a receiver does (pseudo-header + L4 must sum to all-ones).
func TestChecksumSeedReceiverAcceptance(t *testing.T) { func TestChecksumSeedReceiverAcceptance(t *testing.T) {
rng := rand.New(rand.NewSource(0x1826)) rng := rand.New(rand.NewSource(0x1826))
for trial := range 200 { for trial := 0; trial < 200; trial++ {
src := [4]byte{byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256))} src := [4]byte{byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256))}
dst := [4]byte{byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256))} dst := [4]byte{byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256))}
payLen := rng.Intn(1500) payLen := rng.Intn(1500)
+1 -1
View File
@@ -540,7 +540,7 @@ func TestCoalescerCapBySegments(t *testing.T) {
c := newTestTCPCoalescer(t, w) c := newTestTCPCoalescer(t, w)
pay := make([]byte, 512) pay := make([]byte, 512)
seq := uint32(1000) seq := uint32(1000)
for range tcpCoalesceMaxSegs + 5 { for i := 0; i < tcpCoalesceMaxSegs+5; i++ {
if err := c.Commit(buildTCPv4(seq, tcpAck, pay)); err != nil { if err := c.Commit(buildTCPv4(seq, tcpAck, pay)); err != nil {
t.Fatal(err) t.Fatal(err)
} }
+2 -2
View File
@@ -28,7 +28,7 @@ func TestSendBatchReserveCommitFlush(t *testing.T) {
b := NewSendBatch(fw, 4, 32) b := NewSendBatch(fw, 4, 32)
ap := netip.MustParseAddrPort("10.0.0.1:4242") ap := netip.MustParseAddrPort("10.0.0.1:4242")
for i := range 4 { for i := 0; i < 4; i++ {
slot := b.Reserve(32) slot := b.Reserve(32)
if cap(slot) != 32 { if cap(slot) != 32 {
t.Fatalf("slot %d: cap=%d want 32", i, cap(slot)) t.Fatalf("slot %d: cap=%d want 32", i, cap(slot))
@@ -72,7 +72,7 @@ func TestSendBatchSlotsDoNotOverlap(t *testing.T) {
b := NewSendBatch(fw, 3, 8) b := NewSendBatch(fw, 3, 8)
ap := netip.MustParseAddrPort("10.0.0.1:80") ap := netip.MustParseAddrPort("10.0.0.1:80")
for i := range 3 { for i := 0; i < 3; i++ {
s := b.Reserve(8) s := b.Reserve(8)
pkt := append(s[:0], byte(0xA0+i), byte(0xB0+i)) pkt := append(s[:0], byte(0xA0+i), byte(0xB0+i))
b.Commit(pkt, ap) b.Commit(pkt, ap)
+3 -3
View File
@@ -129,7 +129,7 @@ func TestUDPCoalescerCoalescesEqualSized(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w) c := newTestUDPCoalescer(t, w)
pay := make([]byte, 1200) pay := make([]byte, 1200)
for range 3 { for i := 0; i < 3; i++ {
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil { if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -259,7 +259,7 @@ func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w) c := newTestUDPCoalescer(t, w)
pay := make([]byte, 100) pay := make([]byte, 100)
for range udpCoalesceMaxSegs + 5 { for i := 0; i < udpCoalesceMaxSegs+5; i++ {
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil { if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -317,7 +317,7 @@ func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w) c := newTestUDPCoalescer(t, w)
pay := make([]byte, 1200) pay := make([]byte, 1200)
for range 3 { for i := 0; i < 3; i++ {
if err := c.Commit(buildUDPv6(1000, 53, pay)); err != nil { if err := c.Commit(buildUDPv6(1000, 53, pay)); err != nil {
t.Fatal(err) t.Fatal(err)
} }
+1 -1
View File
@@ -151,7 +151,7 @@ func TestChecksumTailPaths(t *testing.T) {
offsets := []int{0, 1, 3, 7, 15} // mix of aligned and odd starts offsets := []int{0, 1, 3, 7, 15} // mix of aligned and odd starts
for k := 0; k <= maxK; k++ { for k := 0; k <= maxK; k++ {
for tail := range 64 { for tail := 0; tail < 64; tail++ {
length := 64*k + tail length := 64*k + tail
for _, seed := range seeds { for _, seed := range seeds {
for _, off := range offsets { for _, off := range offsets {
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing //go:build !e2e_testing
// +build !e2e_testing
package overlay package overlay
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing //go:build !e2e_testing
// +build !e2e_testing
package overlay package overlay
+1
View File
@@ -1,4 +1,5 @@
//go:build linux && !android //go:build linux && !android
// +build linux,!android
package tio package tio
+1
View File
@@ -1,4 +1,5 @@
//go:build linux && !android //go:build linux && !android
// +build linux,!android
package tio package tio
+1
View File
@@ -1,4 +1,5 @@
//go:build linux && !android //go:build linux && !android
// +build linux,!android
package tio package tio
+1
View File
@@ -1,4 +1,5 @@
//go:build linux && !android //go:build linux && !android
// +build linux,!android
package tio package tio
+1
View File
@@ -1,4 +1,5 @@
//go:build linux && !android //go:build linux && !android
// +build linux,!android
package tio package tio
+7 -4
View File
@@ -1,4 +1,5 @@
//go:build linux && !android && !e2e_testing //go:build linux && !android && !e2e_testing
// +build linux,!android,!e2e_testing
package tio package tio
@@ -116,15 +117,17 @@ func TestPoll_ConcurrentWrite_NoRace(t *testing.T) {
}() }()
var wg sync.WaitGroup var wg sync.WaitGroup
for range writers { for w := 0; w < writers; w++ {
wg.Go(func() { wg.Add(1)
for range perWriter { go func() {
defer wg.Done()
for i := 0; i < perWriter; i++ {
if _, werr := p.Write(payload); werr != nil { if _, werr := p.Write(payload); werr != nil {
t.Errorf("write: %v", werr) t.Errorf("write: %v", werr)
return return
} }
} }
}) }()
} }
wg.Wait() wg.Wait()
+1
View File
@@ -1,4 +1,5 @@
//go:build linux && !android //go:build linux && !android
// +build linux,!android
package tio package tio
+6 -5
View File
@@ -1,4 +1,5 @@
//go:build linux && !android && !e2e_testing //go:build linux && !android && !e2e_testing
// +build linux,!android,!e2e_testing
package tio package tio
@@ -136,7 +137,7 @@ func buildTSOv4(t *testing.T, payLen, mss int) ([]byte, virtio.Hdr) {
binary.BigEndian.PutUint16(pkt[34:36], 65535) // window binary.BigEndian.PutUint16(pkt[34:36], 65535) // window
// payload // payload
for i := range payLen { for i := 0; i < payLen; i++ {
pkt[ipLen+tcpLen+i] = byte(i & 0xff) pkt[ipLen+tcpLen+i] = byte(i & 0xff)
} }
return pkt, virtio.NewHeader( return pkt, virtio.NewHeader(
@@ -256,7 +257,7 @@ func TestSegmentTCPv6(t *testing.T) {
pkt[53] = 0x19 // FIN | ACK | PSH — exercise FIN clearing too pkt[53] = 0x19 // FIN | ACK | PSH — exercise FIN clearing too
binary.BigEndian.PutUint16(pkt[54:56], 65535) binary.BigEndian.PutUint16(pkt[54:56], 65535)
for i := range payLen { for i := 0; i < payLen; i++ {
pkt[ipLen+tcpLen+i] = byte(i) pkt[ipLen+tcpLen+i] = byte(i)
} }
@@ -360,7 +361,7 @@ func buildUSOv4(t *testing.T, payLen, gsoSize int) ([]byte, virtio.Hdr) {
binary.BigEndian.PutUint16(pkt[22:24], 53) // dport binary.BigEndian.PutUint16(pkt[22:24], 53) // dport
binary.BigEndian.PutUint16(pkt[24:26], uint16(udpLen+payLen)) // superpacket length binary.BigEndian.PutUint16(pkt[24:26], uint16(udpLen+payLen)) // superpacket length
for i := range payLen { for i := 0; i < payLen; i++ {
pkt[ipLen+udpLen+i] = byte(i & 0xff) pkt[ipLen+udpLen+i] = byte(i & 0xff)
} }
@@ -472,7 +473,7 @@ func TestSegmentUDPv6(t *testing.T) {
// Superpacket-wide length, as the kernel supplies it; see buildUSOv4. // Superpacket-wide length, as the kernel supplies it; see buildUSOv4.
binary.BigEndian.PutUint16(pkt[44:46], uint16(udpLen+payLen)) binary.BigEndian.PutUint16(pkt[44:46], uint16(udpLen+payLen))
for i := range payLen { for i := 0; i < payLen; i++ {
pkt[ipLen+udpLen+i] = byte(i) pkt[ipLen+udpLen+i] = byte(i)
} }
@@ -771,7 +772,7 @@ func buildTSOv6(payLen, gso int) []byte {
pkt[53] = 0x10 // ACK only pkt[53] = 0x10 // ACK only
binary.BigEndian.PutUint16(pkt[54:56], 65535) binary.BigEndian.PutUint16(pkt[54:56], 65535)
for i := range payLen { for i := 0; i < payLen; i++ {
pkt[ipLen+tcpLen+i] = byte(i) pkt[ipLen+tcpLen+i] = byte(i)
} }
return pkt return pkt
+1
View File
@@ -1,4 +1,5 @@
//go:build linux && !android //go:build linux && !android
// +build linux,!android
package virtio package virtio
+11 -4
View File
@@ -1,4 +1,5 @@
//go:build linux && !android //go:build linux && !android
// +build linux,!android
// Package virtio implements the pure validation, header-correction, and // Package virtio implements the pure validation, header-correction, and
// per-segment slicing logic for kernel-supplied TSO/USO superpackets on // per-segment slicing logic for kernel-supplied TSO/USO superpackets on
@@ -253,9 +254,12 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
var savedHdr [maxSegHdrLen]byte var savedHdr [maxSegHdrLen]byte
copy(savedHdr[:headerLen], pkt[:headerLen]) copy(savedHdr[:headerLen], pkt[:headerLen])
for i := range numSeg { for i := 0; i < numSeg; i++ {
segStart := i * gsoSize segStart := i * gsoSize
segEnd := min(segStart+gsoSize, payLen) segEnd := segStart + gsoSize
if segEnd > payLen {
segEnd = payLen
}
segPayLen := segEnd - segStart segPayLen := segEnd - segStart
segLen := headerLen + segPayLen segLen := headerLen + segPayLen
headerOff := i * gsoSize headerOff := i * gsoSize
@@ -355,9 +359,12 @@ func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
var savedHdr [maxSegHdrLen]byte var savedHdr [maxSegHdrLen]byte
copy(savedHdr[:headerLen], pkt[:headerLen]) copy(savedHdr[:headerLen], pkt[:headerLen])
for i := range numSeg { for i := 0; i < numSeg; i++ {
segStart := i * gsoSize segStart := i * gsoSize
segEnd := min(segStart+gsoSize, payLen) segEnd := segStart + gsoSize
if segEnd > payLen {
segEnd = payLen
}
segPayLen := segEnd - segStart segPayLen := segEnd - segStart
segLen := headerLen + segPayLen segLen := headerLen + segPayLen
headerOff := i * gsoSize headerOff := i * gsoSize
+8 -7
View File
@@ -1,4 +1,5 @@
//go:build linux && !android //go:build linux && !android
// +build linux,!android
package virtio package virtio
@@ -57,7 +58,7 @@ func buildTCPv4Super(payLen int) (pkt []byte, hdrLen, csumStart uint16) {
pkt[33] = 0x18 // ACK | PSH pkt[33] = 0x18 // ACK | PSH
binary.BigEndian.PutUint16(pkt[34:36], 65535) // window binary.BigEndian.PutUint16(pkt[34:36], 65535) // window
for i := range payLen { for i := 0; i < payLen; i++ {
pkt[ipLen+tcpLen+i] = byte(i & 0xff) pkt[ipLen+tcpLen+i] = byte(i & 0xff)
} }
return pkt, ipLen + tcpLen, ipLen return pkt, ipLen + tcpLen, ipLen
@@ -81,7 +82,7 @@ func buildUDPv4Super(payLen int) (pkt []byte, hdrLen, csumStart uint16) {
binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport
binary.BigEndian.PutUint16(pkt[22:24], 53) // dport binary.BigEndian.PutUint16(pkt[22:24], 53) // dport
for i := range payLen { for i := 0; i < payLen; i++ {
pkt[ipLen+udpLen+i] = byte(i & 0xff) pkt[ipLen+udpLen+i] = byte(i & 0xff)
} }
return pkt, ipLen + udpLen, ipLen return pkt, ipLen + udpLen, ipLen
@@ -187,7 +188,7 @@ func TestSegmentTCPHeaderNotCorrupted(t *testing.T) {
// Payload bytes must be the original contiguous slice. // Payload bytes must be the original contiguous slice.
segPayLen := len(seg) - int(hdrLen) segPayLen := len(seg) - int(hdrLen)
wantPay := make([]byte, segPayLen) wantPay := make([]byte, segPayLen)
for k := range segPayLen { for k := 0; k < segPayLen; k++ {
wantPay[k] = byte((off + k) & 0xff) wantPay[k] = byte((off + k) & 0xff)
} }
if !bytes.Equal(seg[hdrLen:], wantPay) { if !bytes.Equal(seg[hdrLen:], wantPay) {
@@ -316,7 +317,7 @@ func TestSegmentUDPHeaderNotCorrupted(t *testing.T) {
} }
wantPay := make([]byte, segPayLen) wantPay := make([]byte, segPayLen)
for k := range segPayLen { for k := 0; k < segPayLen; k++ {
wantPay[k] = byte((off + k) & 0xff) wantPay[k] = byte((off + k) & 0xff)
} }
if !bytes.Equal(seg[hdrLen:], wantPay) { if !bytes.Equal(seg[hdrLen:], wantPay) {
@@ -364,7 +365,7 @@ func buildUDPv4Single(payload []byte) (pkt []byte, hdr Hdr) {
// 0xffff, because all-zero is the reserved "no checksum" encoding that IPv6 rejects outright. // 0xffff, because all-zero is the reserved "no checksum" encoding that IPv6 rejects outright.
func TestFinishChecksumUDPZeroStoresAllOnes(t *testing.T) { func TestFinishChecksumUDPZeroStoresAllOnes(t *testing.T) {
var payload []byte var payload []byte
for i := range 0x10000 { for i := 0; i < 0x10000; i++ {
p := []byte{byte(i >> 8), byte(i)} p := []byte{byte(i >> 8), byte(i)}
pkt, hdr := buildUDPv4Single(p) pkt, hdr := buildUDPv4Single(p)
cs, co := int(hdr.CsumStart), int(hdr.CsumOffset) cs, co := int(hdr.CsumStart), int(hdr.CsumOffset)
@@ -543,7 +544,7 @@ func TestBaseSumsMatchZeroingReference(t *testing.T) {
t.Run("ipv4", func(t *testing.T) { t.Run("ipv4", func(t *testing.T) {
for ihl := ipv4HeaderMinLen; ihl <= ipv4HeaderMaxLen; ihl += 4 { for ihl := ipv4HeaderMinLen; ihl <= ipv4HeaderMaxLen; ihl += 4 {
for range 5000 { for iter := 0; iter < 5000; iter++ {
pkt := make([]byte, ihl) pkt := make([]byte, ihl)
for i := range pkt { for i := range pkt {
pkt[i] = randByte(&state) pkt[i] = randByte(&state)
@@ -573,7 +574,7 @@ func TestBaseSumsMatchZeroingReference(t *testing.T) {
for dataOff := 5; dataOff <= 15; dataOff++ { for dataOff := 5; dataOff <= 15; dataOff++ {
tcpLen := dataOff * 4 tcpLen := dataOff * 4
headerLen := csumStart + tcpLen headerLen := csumStart + tcpLen
for range 5000 { for iter := 0; iter < 5000; iter++ {
pkt := make([]byte, headerLen+64) pkt := make([]byte, headerLen+64)
for i := range pkt { for i := range pkt {
pkt[i] = randByte(&state) pkt[i] = randByte(&state)
+2 -2
View File
@@ -93,14 +93,14 @@ func prefixToMask(prefix netip.Prefix) netip.Addr {
} }
func flipBytes(b []byte) []byte { func flipBytes(b []byte) []byte {
for i := range b { for i := 0; i < len(b); i++ {
b[i] ^= 0xFF b[i] ^= 0xFF
} }
return b return b
} }
func orBytes(a []byte, b []byte) []byte { func orBytes(a []byte, b []byte) []byte {
ret := make([]byte, len(a)) ret := make([]byte, len(a))
for i := range a { for i := 0; i < len(a); i++ {
ret[i] = a[i] | b[i] ret[i] = a[i] | b[i]
} }
return ret return ret
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing //go:build !e2e_testing
// +build !e2e_testing
package overlay package overlay
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing //go:build !e2e_testing
// +build !e2e_testing
package overlay package overlay
+1
View File
@@ -1,4 +1,5 @@
//go:build !ios && !e2e_testing //go:build !ios && !e2e_testing
// +build !ios,!e2e_testing
package overlay package overlay
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing //go:build !e2e_testing
// +build !e2e_testing
package overlay package overlay
+1
View File
@@ -1,4 +1,5 @@
//go:build ios && !e2e_testing //go:build ios && !e2e_testing
// +build ios,!e2e_testing
package overlay package overlay
+2 -1
View File
@@ -1,4 +1,5 @@
//go:build !android && !e2e_testing //go:build !android && !e2e_testing
// +build !android,!e2e_testing
package overlay package overlay
@@ -547,7 +548,7 @@ func (t *tun) setDefaultRoute(cidr netip.Prefix) error {
if err != nil { if err != nil {
t.l.Warn("Failed to set default route MTU, retrying", "error", err, "cidr", cidr) t.l.Warn("Failed to set default route MTU, retrying", "error", err, "cidr", cidr)
//retry twice more -- on some systems there appears to be a race condition where if we set routes too soon, netlink says `invalid argument` //retry twice more -- on some systems there appears to be a race condition where if we set routes too soon, netlink says `invalid argument`
for range 2 { for i := 0; i < 2; i++ {
time.Sleep(100 * time.Millisecond) time.Sleep(100 * time.Millisecond)
err = netlink.RouteReplace(&nr) err = netlink.RouteReplace(&nr)
if err == nil { if err == nil {
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing //go:build !e2e_testing
// +build !e2e_testing
package overlay package overlay
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing //go:build !e2e_testing
// +build !e2e_testing
package overlay package overlay
+1
View File
@@ -1,4 +1,5 @@
//go:build !windows //go:build !windows
// +build !windows
package overlay package overlay
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing //go:build !e2e_testing
// +build !e2e_testing
package overlay package overlay
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing //go:build e2e_testing
// +build e2e_testing
package overlay package overlay
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing //go:build !e2e_testing
// +build !e2e_testing
package overlay package overlay
+2 -2
View File
@@ -108,7 +108,7 @@ func TestUserDeviceReadersConcurrentRace(t *testing.T) {
var wg sync.WaitGroup var wg sync.WaitGroup
run := func(idx int) { run := func(idx int) {
defer wg.Done() defer wg.Done()
for range iterations { for i := 0; i < iterations; i++ {
pkts, err := readers[idx].Read() pkts, err := readers[idx].Read()
if err != nil { if err != nil {
errs <- err errs <- err
@@ -136,7 +136,7 @@ func TestUserDeviceReadersConcurrentRace(t *testing.T) {
// waiting reader's private buffer, so reusing buf between writes is safe. // waiting reader's private buffer, so reusing buf between writes is safe.
go func() { go func() {
buf := make([]byte, 32) buf := make([]byte, 32)
for i := range 2 * iterations { for i := 0; i < 2*iterations; i++ {
for j := range buf { for j := range buf {
buf[j] = byte(i + j) buf[j] = byte(i + j)
} }
+7 -1
View File
@@ -11,6 +11,8 @@ import (
"sync" "sync"
"sync/atomic" "sync/atomic"
"time" "time"
"github.com/slackhq/nebula/cert"
) )
// forEachFunc is used to benefit folks that want to do work inside the lock // forEachFunc is used to benefit folks that want to do work inside the lock
@@ -408,11 +410,15 @@ func (r *RemoteList) CopyBlockedRemotes() []netip.AddrPort {
} }
// RefreshFromHandshake locks and updates the RemoteList to account for data learned upon a completed handshake // RefreshFromHandshake locks and updates the RemoteList to account for data learned upon a completed handshake
func (r *RemoteList) RefreshFromHandshake(vpnAddrs []netip.Addr) { func (r *RemoteList) RefreshFromHandshake(vpnAddrs []netip.Addr, v cert.Version) {
r.Lock() r.Lock()
r.badRemotes = nil r.badRemotes = nil
if v != cert.Version1 {
// a handshake from a v1 cert can never expand our knowledge of the number of addresses a host has,
// and, because v2 certs exist, it can also never contract it. So, only update this for non-v1 certs.
r.vpnAddrs = make([]netip.Addr, len(vpnAddrs)) r.vpnAddrs = make([]netip.Addr, len(vpnAddrs))
copy(r.vpnAddrs, vpnAddrs) copy(r.vpnAddrs, vpnAddrs)
}
r.Unlock() r.Unlock()
} }
+3 -3
View File
@@ -27,7 +27,7 @@ func TestPacketsAreBalancedEqually(t *testing.T) {
gw3count := 0 gw3count := 0
iterationCount := uint16(65535) iterationCount := uint16(65535)
for i := range iterationCount { for i := uint16(0); i < iterationCount; i++ {
packet := firewall.Packet{ packet := firewall.Packet{
LocalAddr: netip.MustParseAddr("192.168.1.1"), LocalAddr: netip.MustParseAddr("192.168.1.1"),
RemoteAddr: netip.MustParseAddr("10.0.0.1"), RemoteAddr: netip.MustParseAddr("10.0.0.1"),
@@ -74,7 +74,7 @@ func TestPacketsAreBalancedByPriority(t *testing.T) {
gw2count := 0 gw2count := 0
iterationCount := uint16(65535) iterationCount := uint16(65535)
for i := range iterationCount { for i := uint16(0); i < iterationCount; i++ {
packet := firewall.Packet{ packet := firewall.Packet{
LocalAddr: netip.MustParseAddr("192.168.1.1"), LocalAddr: netip.MustParseAddr("192.168.1.1"),
RemoteAddr: netip.MustParseAddr("10.0.0.1"), RemoteAddr: netip.MustParseAddr("10.0.0.1"),
@@ -115,7 +115,7 @@ func TestBalancePacketDistributsRandomlyAndReturnsFalseIfBucketsNotCalculated(t
gw1count := 0 gw1count := 0
gw2count := 0 gw2count := 0
for i := range iterationCount { for i := uint16(0); i < iterationCount; i++ {
packet := firewall.Packet{ packet := firewall.Packet{
LocalAddr: netip.MustParseAddr("192.168.1.1"), LocalAddr: netip.MustParseAddr("192.168.1.1"),
RemoteAddr: netip.MustParseAddr("10.0.0.1"), RemoteAddr: netip.MustParseAddr("10.0.0.1"),
+4 -5
View File
@@ -3,7 +3,6 @@ package routing
import ( import (
"fmt" "fmt"
"net/netip" "net/netip"
"strings"
) )
const ( const (
@@ -14,14 +13,14 @@ const (
type Gateways []Gateway type Gateways []Gateway
func (g Gateways) String() string { func (g Gateways) String() string {
var str strings.Builder str := ""
for i, gw := range g { for i, gw := range g {
str.WriteString(gw.String()) str += gw.String()
if i < len(g)-1 { if i < len(g)-1 {
str.WriteString(", ") str += ", "
} }
} }
return str.String() return str
} }
type Gateway struct { type Gateway struct {
+8 -4
View File
@@ -1,19 +1,21 @@
package nebula package nebula
import ( import (
"context"
"testing" "testing"
"time" "time"
) )
func TestScheduler_PooledReuse(t *testing.T) { func TestScheduler_PooledReuse(t *testing.T) {
ctx := t.Context() ctx, cancel := context.WithCancel(context.Background())
defer cancel()
s := NewScheduler[int](16) s := NewScheduler[int](16)
delivered := make(chan int, 256) delivered := make(chan int, 256)
go s.Run(ctx, func(item int) { delivered <- item }) go s.Run(ctx, func(item int) { delivered <- item })
const N = 100 const N = 100
for i := range N { for i := 0; i < N; i++ {
s.Schedule(ctx, i, time.Millisecond) s.Schedule(ctx, i, time.Millisecond)
} }
@@ -32,7 +34,8 @@ func TestScheduler_PooledReuse(t *testing.T) {
// BenchmarkScheduler_Schedule reports allocations per Schedule call. // BenchmarkScheduler_Schedule reports allocations per Schedule call.
// In steady state the Scheduler's sync.Pool means we should see zero allocs per op once the pool warms up. // In steady state the Scheduler's sync.Pool means we should see zero allocs per op once the pool warms up.
func BenchmarkScheduler_Schedule(b *testing.B) { func BenchmarkScheduler_Schedule(b *testing.B) {
ctx := b.Context() ctx, cancel := context.WithCancel(context.Background())
defer cancel()
s := NewScheduler[int](b.N) s := NewScheduler[int](b.N)
go s.Run(ctx, func(int) {}) go s.Run(ctx, func(int) {})
@@ -48,7 +51,8 @@ func BenchmarkScheduler_Schedule(b *testing.B) {
// What we'd pay per Schedule if Punchy called time.AfterFunc directly without the pooled Scheduler. // What we'd pay per Schedule if Punchy called time.AfterFunc directly without the pooled Scheduler.
// Allocates a *time.Timer plus a closure each call. // Allocates a *time.Timer plus a closure each call.
func BenchmarkBareAfterFunc(b *testing.B) { func BenchmarkBareAfterFunc(b *testing.B) {
ctx := b.Context() ctx, cancel := context.WithCancel(context.Background())
defer cancel()
queue := make(chan int, b.N) queue := make(chan int, b.N)
go func() { go func() {
+2 -2
View File
@@ -25,7 +25,7 @@ func AssertDeepCopyEqual(t *testing.T, a any, b any) {
} }
func traverseDeepCopy(t *testing.T, v1 reflect.Value, v2 reflect.Value, name string) bool { func traverseDeepCopy(t *testing.T, v1 reflect.Value, v2 reflect.Value, name string) bool {
if v1.Type() == v2.Type() && v1.Type() == reflect.TypeFor[netip.Addr]() { if v1.Type() == v2.Type() && v1.Type() == reflect.TypeOf(netip.Addr{}) {
// Ignore netip.Addr types since they reuse an interned global value // Ignore netip.Addr types since they reuse an interned global value
return false return false
} }
@@ -72,7 +72,7 @@ func traverseDeepCopy(t *testing.T, v1 reflect.Value, v2 reflect.Value, name str
} }
return traverseDeepCopy(t, v1.Elem(), v2.Elem(), name) return traverseDeepCopy(t, v1.Elem(), v2.Elem(), name)
case reflect.Pointer: case reflect.Ptr:
local := reflect.ValueOf(time.Local).Pointer() local := reflect.ValueOf(time.Local).Pointer()
if local == v1.Pointer() && local == v2.Pointer() { if local == v1.Pointer() && local == v2.Pointer() {
return true return true
+1
View File
@@ -1,4 +1,5 @@
//go:build darwin && !ios && !e2e_testing //go:build darwin && !ios && !e2e_testing
// +build darwin,!ios,!e2e_testing
package udp package udp
+1
View File
@@ -1,4 +1,5 @@
//go:build darwin && !ios && !e2e_testing //go:build darwin && !ios && !e2e_testing
// +build darwin,!ios,!e2e_testing
package udp package udp
+1
View File
@@ -1,4 +1,5 @@
//go:build !darwin || ios || e2e_testing //go:build !darwin || ios || e2e_testing
// +build !darwin ios e2e_testing
package udp package udp
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing //go:build !e2e_testing
// +build !e2e_testing
package udp package udp
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing //go:build !e2e_testing
// +build !e2e_testing
package udp package udp
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing //go:build !e2e_testing
// +build !e2e_testing
package udp package udp
+5 -2
View File
@@ -294,7 +294,10 @@ func deliverSegments(r EncReader, from netip.AddrPort, payload []byte, segSize i
return return
} }
for off := 0; off < len(payload); off += segSize { for off := 0; off < len(payload); off += segSize {
end := min(off+segSize, len(payload)) end := off + segSize
if end > len(payload) {
end = len(payload)
}
r(from, payload[off:end:end]) r(from, payload[off:end:end])
} }
} }
@@ -476,7 +479,7 @@ func NewUDPStatsEmitter(udpConns []Conn) func() {
return func() { return func() {
for i, gauges := range udpGauges { for i, gauges := range udpGauges {
if err := udpConns[i].(*StdConn).getMemInfo(&meminfo); err == nil { if err := udpConns[i].(*StdConn).getMemInfo(&meminfo); err == nil {
for j := range unix.SK_MEMINFO_VARS { for j := 0; j < unix.SK_MEMINFO_VARS; j++ {
gauges[j].Update(int64(meminfo[j])) gauges[j].Update(int64(meminfo[j]))
} }
} }
+4 -4
View File
@@ -107,7 +107,7 @@ func TestWriteBatchBadFamilyDeliversOthers(t *testing.T) {
got := map[string]bool{} got := map[string]bool{}
rx.SetReadDeadline(time.Now().Add(2 * time.Second)) rx.SetReadDeadline(time.Now().Add(2 * time.Second))
buf := make([]byte, 64) buf := make([]byte, 64)
for i := range 2 { for i := 0; i < 2; i++ {
n, _, rerr := rx.ReadFromUDPAddrPort(buf) n, _, rerr := rx.ReadFromUDPAddrPort(buf)
if rerr != nil { if rerr != nil {
t.Fatalf("expected 2 delivered packets, read #%d failed: %v", i+1, rerr) t.Fatalf("expected 2 delivered packets, read #%d failed: %v", i+1, rerr)
@@ -161,7 +161,7 @@ func TestWriteBatchUnreachableDestDeliversOthers(t *testing.T) {
got := map[string]bool{} got := map[string]bool{}
rx.SetReadDeadline(time.Now().Add(2 * time.Second)) rx.SetReadDeadline(time.Now().Add(2 * time.Second))
buf := make([]byte, 64) buf := make([]byte, 64)
for i := range 4 { for i := 0; i < 4; i++ {
n, _, rerr := rx.ReadFromUDPAddrPort(buf) n, _, rerr := rx.ReadFromUDPAddrPort(buf)
if rerr != nil { if rerr != nil {
t.Fatalf("expected 4 delivered packets, read #%d failed: %v (got so far: %v)", i+1, rerr, got) t.Fatalf("expected 4 delivered packets, read #%d failed: %v (got so far: %v)", i+1, rerr, got)
@@ -691,7 +691,7 @@ func TestGSOEngagesOnLoopback(t *testing.T) {
// The kernel must deliver the original datagram boundaries and bytes. // The kernel must deliver the original datagram boundaries and bytes.
_ = rx.SetReadDeadline(time.Now().Add(5 * time.Second)) _ = rx.SetReadDeadline(time.Now().Add(5 * time.Second))
got := make([]byte, pktLen+1) got := make([]byte, pktLen+1)
for i := range numPkts { for i := 0; i < numPkts; i++ {
n, _, err := rx.ReadFromUDP(got) n, _, err := rx.ReadFromUDP(got)
if err != nil { if err != nil {
t.Fatalf("rx read %d: %v", i, err) t.Fatalf("rx read %d: %v", i, err)
@@ -699,7 +699,7 @@ func TestGSOEngagesOnLoopback(t *testing.T) {
if n != pktLen { if n != pktLen {
t.Fatalf("rx read %d: len=%d want %d (kernel segmented at wrong boundary)", i, n, pktLen) t.Fatalf("rx read %d: len=%d want %d (kernel segmented at wrong boundary)", i, n, pktLen)
} }
for j := range n { for j := 0; j < n; j++ {
if got[j] != byte(i) { if got[j] != byte(i) {
t.Fatalf("rx read %d: byte %d = %#x, want %#x", i, j, got[j], byte(i)) t.Fatalf("rx read %d: byte %d = %#x, want %#x", i, j, got[j], byte(i))
} }
+6 -3
View File
@@ -112,7 +112,7 @@ func (w *batchWriter) prepareWriteMessages(n int, offloadsEnabled bool) {
w.cmsg = make([]byte, n*w.cmsgSpace) w.cmsg = make([]byte, n*w.cmsgSpace)
for k := range n { for k := 0; k < n; k++ {
base := k * w.cmsgSpace base := k * w.cmsgSpace
seg := (*unix.Cmsghdr)(unsafe.Pointer(&w.cmsg[base])) seg := (*unix.Cmsghdr)(unsafe.Pointer(&w.cmsg[base]))
seg.Level = unix.SOL_UDP seg.Level = unix.SOL_UDP
@@ -214,7 +214,7 @@ func (w *batchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, er
break break
} }
for k := range runLen { for k := 0; k < runLen; k++ {
b := bufs[i+k] b := bufs[i+k]
if len(b) == 0 { if len(b) == 0 {
w.iovs[iovIdx+k].Base = nil w.iovs[iovIdx+k].Base = nil
@@ -318,7 +318,10 @@ func (w *batchWriter) planRun(bufs [][]byte, addrs []netip.AddrPort, start, iovB
return 1, segSize return 1, segSize
} }
dst := addrs[start] dst := addrs[start]
maxLen := min(iovBudget, w.maxGSOSegments) maxLen := w.maxGSOSegments
if iovBudget < maxLen {
maxLen = iovBudget
}
runLen := 1 runLen := 1
total := segSize total := segSize
for runLen < maxLen && start+runLen < len(bufs) { for runLen < maxLen && start+runLen < len(bufs) {
+2 -2
View File
@@ -64,13 +64,13 @@ func TestWriteBatchNoAllocs(t *testing.T) {
addrs = append(addrs, dst) addrs = append(addrs, dst)
} }
// GSO-eligible run with a short tail. // GSO-eligible run with a short tail.
for range 8 { for k := 0; k < 8; k++ {
add(payload, dstA) add(payload, dstA)
} }
add(short, dstA) add(short, dstA)
add(payload, dstA) add(payload, dstA)
// Alternating destinations defeat coalescing entirely. // Alternating destinations defeat coalescing entirely.
for k := range 4 { for k := 0; k < 4; k++ {
dst := dstA dst := dstA
if k%2 == 0 { if k%2 == 0 {
dst = dstB dst = dstB
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing //go:build !e2e_testing
// +build !e2e_testing
package udp package udp
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing //go:build !e2e_testing
// +build !e2e_testing
// Inspired by https://git.zx2c4.com/wireguard-go/tree/conn/bind_windows.go // Inspired by https://git.zx2c4.com/wireguard-go/tree/conn/bind_windows.go
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing //go:build e2e_testing
// +build e2e_testing
package udp package udp
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing //go:build !e2e_testing
// +build !e2e_testing
package udp package udp
+1 -1
View File
@@ -40,7 +40,7 @@ func AllowedCPUs() ([]int, error) {
return nil, err return nil, err
} }
cpus := make([]int, 0, set.Count()) cpus := make([]int, 0, set.Count())
for cpu := range len(set) * 64 { for cpu := 0; cpu < len(set)*64; cpu++ {
if set.IsSet(cpu) { if set.IsSet(cpu) {
cpus = append(cpus, cpu) cpus = append(cpus, cpu)
} }