control/controlclient,net/tstun,wgengine/magicsock: fix handling of zero keys in TSMP (#20508)

Updates tailscale/corp#45042

Signed-off-by: Claus Lensbøl <claus@tailscale.com>
This commit is contained in:
Claus Lensbøl
2026-07-17 14:02:53 -04:00
committed by GitHub
parent c1edf7f458
commit 82a381e54b
6 changed files with 121 additions and 49 deletions
+5
View File
@@ -200,6 +200,11 @@ func (ms *mapSession) Close() {
var ErrChangeQueueClosed = errors.New("change queue closed")
func (ms *mapSession) updateDiscoForNode(id tailcfg.NodeID, key key.NodePublic, discoKey key.DiscoPublic, lastSeen time.Time, online bool) error {
if discoKey.IsZero() {
ms.logf("[v1] controlclient: received zero disco key update from nodeID %v", id)
return nil
}
ms.cqmu.Lock()
if ms.changeQueueClosed {
+23 -12
View File
@@ -629,7 +629,7 @@ func TestUpdateDiscoForNode(t *testing.T) {
name string
initialOnline bool
initialLastSeen time.Time
updateDiscoKey bool
updateDiscoKey func() key.DiscoPublic
updateOnline bool
updateLastSeen time.Time
wantUpdate bool
@@ -639,7 +639,7 @@ func TestUpdateDiscoForNode(t *testing.T) {
name: "newer_key_not_online",
initialOnline: true,
initialLastSeen: time.Unix(1, 0),
updateDiscoKey: true,
updateDiscoKey: key.NewDisco().Public,
updateOnline: false,
updateLastSeen: time.Now(),
wantUpdate: true,
@@ -649,7 +649,7 @@ func TestUpdateDiscoForNode(t *testing.T) {
name: "newer_key_online",
initialOnline: true,
initialLastSeen: time.Unix(1, 0),
updateDiscoKey: true,
updateDiscoKey: key.NewDisco().Public,
updateOnline: true,
updateLastSeen: time.Now(),
wantUpdate: true,
@@ -659,7 +659,7 @@ func TestUpdateDiscoForNode(t *testing.T) {
name: "older_key_not_online",
initialOnline: false,
initialLastSeen: time.Now(),
updateDiscoKey: true,
updateDiscoKey: key.NewDisco().Public,
updateOnline: false,
updateLastSeen: time.Unix(1, 0),
wantUpdate: false,
@@ -669,7 +669,7 @@ func TestUpdateDiscoForNode(t *testing.T) {
name: "older_key_online",
initialOnline: false,
initialLastSeen: time.Now(),
updateDiscoKey: true,
updateDiscoKey: key.NewDisco().Public,
updateOnline: true,
updateLastSeen: time.Unix(1, 0),
wantUpdate: true,
@@ -679,7 +679,7 @@ func TestUpdateDiscoForNode(t *testing.T) {
name: "same_newer_key_not_online",
initialOnline: true,
initialLastSeen: time.Unix(1, 0),
updateDiscoKey: false,
updateDiscoKey: nil,
updateOnline: false,
updateLastSeen: time.Now(),
wantUpdate: false,
@@ -689,7 +689,7 @@ func TestUpdateDiscoForNode(t *testing.T) {
name: "same_newer_key_online",
initialOnline: true,
initialLastSeen: time.Unix(1, 0),
updateDiscoKey: false,
updateDiscoKey: nil,
updateOnline: true,
updateLastSeen: time.Now(),
wantUpdate: false,
@@ -699,7 +699,7 @@ func TestUpdateDiscoForNode(t *testing.T) {
name: "same_older_key_not_online",
initialOnline: false,
initialLastSeen: time.Now(),
updateDiscoKey: false,
updateDiscoKey: nil,
updateOnline: false,
updateLastSeen: time.Unix(1, 0),
wantUpdate: false,
@@ -709,7 +709,7 @@ func TestUpdateDiscoForNode(t *testing.T) {
name: "same_older_key_online",
initialOnline: false,
initialLastSeen: time.Now(),
updateDiscoKey: false,
updateDiscoKey: nil,
updateOnline: true,
updateLastSeen: time.Unix(1, 0),
wantUpdate: true,
@@ -718,12 +718,23 @@ func TestUpdateDiscoForNode(t *testing.T) {
{
name: "no_initial_last_seen",
initialOnline: false,
updateDiscoKey: true,
updateDiscoKey: key.NewDisco().Public,
updateOnline: false,
updateLastSeen: time.Now(),
wantUpdate: true,
wantKeyChanged: true,
},
{
name: "zero_key",
initialOnline: false,
updateDiscoKey: func() key.DiscoPublic {
return key.DiscoPublic{}
},
updateOnline: false,
updateLastSeen: time.Now(),
wantUpdate: false,
wantKeyChanged: false,
},
}
for _, tt := range tests {
@@ -755,8 +766,8 @@ func TestUpdateDiscoForNode(t *testing.T) {
}
newKey := oldKey.Public()
if tt.updateDiscoKey {
newKey = key.NewDisco().Public()
if tt.updateDiscoKey != nil {
newKey = tt.updateDiscoKey()
}
ms.updateDiscoForNode(node.ID, node.Key, newKey, tt.updateLastSeen, tt.updateOnline)
+2 -1
View File
@@ -1097,7 +1097,8 @@ func (t *Wrapper) filterPacketInboundFromWireGuard(p *packet.Parsed, captHook pa
t.injectOutboundPong(p, pingReq)
return filter.DropSilently, gro
} else if discoKeyAdvert, ok := p.AsTSMPDiscoAdvertisement(); ok {
if buildfeatures.HasCacheNetMap && envknob.BoolDefaultTrue("TS_USE_CACHED_NETMAP") {
if buildfeatures.HasCacheNetMap && envknob.BoolDefaultTrue("TS_USE_CACHED_NETMAP") &&
!discoKeyAdvert.Key.IsZero() {
t.discoKeyAdvertisementPub.Publish(events.DiscoKeyAdvertisement{
Src: discoKeyAdvert.Src,
Key: discoKeyAdvert.Key,
+23
View File
@@ -1213,3 +1213,26 @@ func TestSetPeerRoutesFastPath(t *testing.T) {
t.Fatalf("peerConfig after second uninstall = %v; want nil", got)
}
}
// Drop empty TSMPDiscoAdvert packets inbound via wireguard.
func TestFilterDropEmptyTSMPDiscoAdvertInbound(t *testing.T) {
var memLog tstest.MemLogger
tw := &Wrapper{logf: memLog.Logf, limitedLogf: memLog.Logf}
ipHdr := packet.IP4Header{
IPProto: ipproto.TSMP,
Src: netaddr.IPv4(1, 2, 3, 4),
Dst: netaddr.IPv4(5, 6, 7, 8),
}
tsmpPayload := make([]byte, 33)
tsmpPayload[0] = byte(packet.TSMPTypeDiscoAdvertisement)
pkt := make([]byte, ipHdr.Len()+len(tsmpPayload))
ipHdr.Marshal(pkt)
copy(pkt[ipHdr.Len():], tsmpPayload)
pp := new(packet.Parsed)
pp.Decode(pkt)
got, _ := tw.filterPacketInboundFromWireGuard(pp, nil, nil, nil)
if got != filter.DropSilently {
t.Errorf("got %v; want DropSilently", got)
}
}
+5
View File
@@ -4495,6 +4495,11 @@ func (c *Conn) PeerRelays() set.Set[netip.Addr] {
// node is the Tailscale tailcfg.NodeView of the peer that sent the update.
func (c *Conn) HandleDiscoKeyAdvertisement(node tailcfg.NodeView, update packet.TSMPDiscoKeyAdvertisement) {
discoKey := update.Key
if discoKey.IsZero() {
c.logf("[v1] magicsock: received zero disco key update from %v", node.StableID())
return
}
c.logf("magicsock: received disco key update %v from %v", discoKey.ShortString(), node.StableID())
metricTSMPDiscoKeyAdvertisementReceived.Add(1)
+63 -36
View File
@@ -4618,46 +4618,73 @@ func TestRotateDiscoKeyMultipleTimes(t *testing.T) {
}
func TestReceiveTSMPDiscoKeyAdvertisement(t *testing.T) {
conn := newTestConn(t)
t.Cleanup(func() { conn.Close() })
peerKey := key.NewNode().Public()
ep := &endpoint{
nodeID: 1,
publicKey: peerKey,
nodeAddr: netip.MustParseAddr("100.64.0.1"),
}
discoKey := key.NewDisco().Public()
ep.disco.Store(&endpointDisco{
key: discoKey,
short: discoKey.ShortString(),
})
ep.c = conn
conn.mu.Lock()
nodeView := (&tailcfg.Node{
Key: ep.publicKey,
Addresses: []netip.Prefix{
netip.MustParsePrefix("100.64.0.1/32"),
tests := []struct {
name string
newKeyFunc func() key.DiscoPublic
wantUpdate bool
}{
{
name: "normal_key_change",
newKeyFunc: key.NewDisco().Public,
wantUpdate: true,
},
{
name: "zero_key_change",
newKeyFunc: func() key.DiscoPublic {
return key.DiscoPublic{}
},
wantUpdate: false,
},
}).View()
conn.peersByID = map[tailcfg.NodeID]tailcfg.NodeView{nodeView.ID(): nodeView}
conn.mu.Unlock()
conn.peerMap.upsertEndpoint(ep, key.DiscoPublic{})
if ep.discoShort() != discoKey.ShortString() {
t.Errorf("Original disco key %s, does not match %s", discoKey.ShortString(), ep.discoShort())
}
newDiscoKey := key.NewDisco().Public()
tka := packet.TSMPDiscoKeyAdvertisement{
Src: netip.MustParseAddr("100.64.0.1"),
Key: newDiscoKey,
}
conn.HandleDiscoKeyAdvertisement(nodeView, tka)
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
conn := newTestConn(t)
t.Cleanup(func() { conn.Close() })
if ep.disco.Load().short != newDiscoKey.ShortString() {
t.Errorf("New disco key %s, does not match %s", newDiscoKey.ShortString(), ep.disco.Load().short)
peerKey := key.NewNode().Public()
ep := &endpoint{
nodeID: 1,
publicKey: peerKey,
nodeAddr: netip.MustParseAddr("100.64.0.1"),
}
discoKey := key.NewDisco().Public()
ep.disco.Store(&endpointDisco{
key: discoKey,
short: discoKey.ShortString(),
})
ep.c = conn
conn.mu.Lock()
nodeView := (&tailcfg.Node{
Key: ep.publicKey,
Addresses: []netip.Prefix{
netip.MustParsePrefix("100.64.0.1/32"),
},
}).View()
conn.peersByID = map[tailcfg.NodeID]tailcfg.NodeView{nodeView.ID(): nodeView}
conn.mu.Unlock()
conn.peerMap.upsertEndpoint(ep, key.DiscoPublic{})
if ep.discoShort() != discoKey.ShortString() {
t.Errorf("Original disco key %s, does not match %s", discoKey.ShortString(), ep.discoShort())
}
newDiscoKey := tt.newKeyFunc()
tka := packet.TSMPDiscoKeyAdvertisement{
Src: netip.MustParseAddr("100.64.0.1"),
Key: newDiscoKey,
}
conn.HandleDiscoKeyAdvertisement(nodeView, tka)
wantDiscoKey := discoKey
if tt.wantUpdate {
wantDiscoKey = newDiscoKey
}
if ep.disco.Load().short != wantDiscoKey.ShortString() {
t.Errorf("New disco key %s, does not match %s", newDiscoKey.ShortString(), ep.disco.Load().short)
}
})
}
}