diff --git a/hscontrol/servertest/via_compat_test.go b/hscontrol/servertest/via_compat_test.go index 755e8f768..e00852c6d 100644 --- a/hscontrol/servertest/via_compat_test.go +++ b/hscontrol/servertest/via_compat_test.go @@ -250,16 +250,66 @@ func runViaMapCompat(t *testing.T, c *testcapture.Capture) { } t.Run(viewerName, func(t *testing.T) { - require.EventuallyWithT(t, func(collect *assert.CollectT) { + requireNetmapHolds(t, func(tt require.TestingT) { nm := cl.Netmap() - require.NotNil(collect, nm, "netmap is nil") + require.NotNil(tt, nm, "netmap is nil") - compareNetmap(collect, nm, capture, clients, saasAddrs) - }, 30*time.Second, 100*time.Millisecond) + compareNetmap(tt, nm, capture, clients, saasAddrs) + }) }) } } +// requireNetmapHolds retries check until it passes on several consecutive +// ticks. One passing tick is not enough: a netmap that matches and then +// drifts, such as a primary flipping away, is the bug these oracles guard. +func requireNetmapHolds(t *testing.T, check func(require.TestingT)) { + t.Helper() + + // Long enough for a batch still in flight to land. + const holdTicks = 5 + + held := 0 + + require.EventuallyWithT(t, func(collect *assert.CollectT) { + tick := &failRecorder{CollectT: collect} + + // Deferred because FailNow exits the goroutine. + defer func() { + if tick.failed { + held = 0 + + return + } + + held++ + if held < holdTicks { + collect.Errorf("netmap held for %d of %d ticks", held, holdTicks) + } + }() + + check(tick) + }, 30*time.Second, 100*time.Millisecond) +} + +// failRecorder notes whether a tick failed, which [assert.CollectT] does +// not expose. +type failRecorder struct { + *assert.CollectT + + failed bool +} + +func (r *failRecorder) Errorf(format string, args ...any) { + r.failed = true + r.CollectT.Errorf(format, args...) +} + +func (r *failRecorder) FailNow() { + r.failed = true + r.CollectT.FailNow() +} + // compareNetmap compares the headscale [tailcfg.MapResponse] against the // captured [netmap.NetworkMap] data in an IP-independent way. It validates: // - Peer visibility (which peers are present, by hostname) @@ -268,7 +318,7 @@ func runViaMapCompat(t *testing.T, c *testcapture.Capture) { // - PrimaryRoutes per peer // - PacketFilter (source, destination, ports) triples // -// It takes [require.TestingT] so it can run inside [require.EventuallyWithT]. +// It takes [require.TestingT] so it can run inside [requireNetmapHolds]. func compareNetmap( t require.TestingT, got *netmap.NetworkMap, diff --git a/hscontrol/servertest/via_ha_compat_test.go b/hscontrol/servertest/via_ha_compat_test.go index fff619f3d..357335564 100644 --- a/hscontrol/servertest/via_ha_compat_test.go +++ b/hscontrol/servertest/via_ha_compat_test.go @@ -196,9 +196,9 @@ func runViaHACompat(t *testing.T, c *testcapture.Capture) { } t.Run(viewerName, func(t *testing.T) { - require.EventuallyWithT(t, func(collect *assert.CollectT) { - compareCaptureNetmap(collect, cl, capture, clients) - }, 30*time.Second, 100*time.Millisecond) + requireNetmapHolds(t, func(tt require.TestingT) { + compareCaptureNetmap(tt, cl, capture, clients) + }) }) } } @@ -206,7 +206,7 @@ func runViaHACompat(t *testing.T, c *testcapture.Capture) { // compareCaptureNetmap compares headscale's [tailcfg.MapResponse] against a // [testcapture.Node]'s [netmap.NetworkMap] data. Same logic as [compareNetmap] but // reads from typed [testcapture] fields instead of goldenFile strings. -// It takes [require.TestingT] so it can run inside [require.EventuallyWithT]. +// It takes [require.TestingT] so it can run inside [requireNetmapHolds]. func compareCaptureNetmap( t require.TestingT, viewer *servertest.TestClient,