From df6e422e10a09b1aa61f177b3f7abf21183f7a7b Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Tue, 23 Jun 2026 13:38:16 +0200 Subject: [PATCH] [client] Resolve cold mgmt cache domains synchronously before DNS takeover Async server-domain resolution risked blackholing infra DNS at bootstrap: if OS DNS was reconfigured to route through a dead exit node before the background resolve ran, the cache never populated. Resolve cold domains (no cached record) synchronously while NetBird has not yet taken over the system resolver and will serve DNS (dnsWillBeServed), so the cache is primed via the working OS resolver before takeover. Stale and post-takeover resolves stay async to keep the engine sync lock unblocked. --- client/internal/dns/mgmt/mgmt.go | 70 +++++++++++++++++++++------ client/internal/dns/mgmt/mgmt_test.go | 14 +++--- client/internal/dns/server.go | 6 ++- 3 files changed, 66 insertions(+), 24 deletions(-) diff --git a/client/internal/dns/mgmt/mgmt.go b/client/internal/dns/mgmt/mgmt.go index 2042e8f59..1f2d3710b 100644 --- a/client/internal/dns/mgmt/mgmt.go +++ b/client/internal/dns/mgmt/mgmt.go @@ -519,15 +519,12 @@ func (m *Resolver) GetCachedDomains() domain.List { return domains } -// UpdateFromServerDomains updates the cache with server domains from network configuration. -// It merges new domains with existing ones, replacing entire domain types when updated. -// Empty updates are ignored to prevent clearing infrastructure domains during partial updates. -// UpdateFromServerDomains records the requested domains and kicks off their -// resolution in the background, returning without blocking on DNS so it stays -// off the engine sync lock held by the caller. ctx scopes the background -// resolves to the server lifetime: it is not per-sync, so a fast-returning -// sync won't cancel them, but a server Stop will. -func (m *Resolver) UpdateFromServerDomains(ctx context.Context, serverDomains dnsconfig.ServerDomains) (domain.List, error) { +// UpdateFromServerDomains merges server domains into the cache and resolves +// them. New types replace whole types; empty updates are ignored. Resolution is +// async (off the caller's sync lock) except for cold domains when dnsWillBeServed +// and takeover is pending, which kickoffResolve primes synchronously. ctx is the +// server lifetime, so a fast sync won't cancel resolves but Stop will. +func (m *Resolver) UpdateFromServerDomains(ctx context.Context, serverDomains dnsconfig.ServerDomains, dnsWillBeServed bool) (domain.List, error) { newDomains := m.extractDomainsFromServerDomains(serverDomains) var removedDomains domain.List @@ -545,26 +542,40 @@ func (m *Resolver) UpdateFromServerDomains(ctx context.Context, serverDomains dn removedDomains = m.removeStaleDomains(currentDomains, allDomains) } - m.kickoffResolve(ctx, newDomains) + m.kickoffResolve(ctx, newDomains, dnsWillBeServed) return removedDomains, nil } -// kickoffResolve marks each domain pending and starts a background resolve, -// skipping ones already fresh or in flight. Returns immediately. -func (m *Resolver) kickoffResolve(ctx context.Context, domains domain.List) { +// kickoffResolve resolves each unresolved domain, skipping fresh/in-flight ones. +// Cold domains resolve synchronously only before takeover (no upstream root +// handler) and when dnsWillBeServed, to prime the cache via the working OS +// resolver before OS DNS routes through the tunnel; otherwise async. +func (m *Resolver) kickoffResolve(ctx context.Context, domains domain.List, dnsWillBeServed bool) { + m.mutex.RLock() + chain := m.chain + maxPriority := m.chainMaxPriority + m.mutex.RUnlock() + preTakeover := chain == nil || !chain.HasRootHandlerAtOrBelow(maxPriority) + for _, d := range domains { dnsName := strings.ToLower(dns.Fqdn(d.PunycodeString())) m.mutex.Lock() _, hasPending := m.pending[dnsName] - cached := m.hasFreshRecordLocked(dnsName) - if !hasPending && !cached { + fresh := m.hasFreshRecordLocked(dnsName) + cold := !m.hasAnyRecordLocked(dnsName) + if !hasPending && !fresh { m.pending[dnsName] = pendingEntry{} } m.mutex.Unlock() - if hasPending || cached { + if hasPending || fresh { + continue + } + + if cold && preTakeover && dnsWillBeServed { + m.resolveInitial(ctx, d, dnsName) continue } @@ -572,6 +583,21 @@ func (m *Resolver) kickoffResolve(ctx context.Context, domains domain.List) { } } +// resolveInitial resolves a cold domain synchronously, deduped via resolveGroup +// so a concurrent ServeDNS await joins the same flight. Clears pending when done. +func (m *Resolver) resolveInitial(ctx context.Context, d domain.Domain, dnsName string) { + key := "initial|" + dnsName + _, _, _ = m.resolveGroup.Do(key, func() (any, error) { + defer m.clearPending(dnsName) + if err := m.AddDomain(ctx, d); err != nil { + log.Warnf("initial resolve mgmt domain=%s: %v", d.SafeString(), err) + return struct{}{}, err + } + log.Debugf("added/updated management cache domain=%s", d.SafeString()) + return struct{}{}, nil + }) +} + // scheduleInitialResolve runs AddDomain in the background, deduped per domain // by resolveGroup, clearing the pending marker when it finishes. ctx is the // server-lifetime context so a Stop cancels in-flight resolves. @@ -600,6 +626,18 @@ func (m *Resolver) hasFreshRecordLocked(dnsName string) bool { return false } +// hasAnyRecordLocked reports whether any A or AAAA record exists for the name, +// fresh or stale. Caller holds m.mutex. +func (m *Resolver) hasAnyRecordLocked(dnsName string) bool { + for _, qtype := range []uint16{dns.TypeA, dns.TypeAAAA} { + q := dns.Question{Name: dnsName, Qtype: qtype, Qclass: dns.ClassINET} + if _, ok := m.records[q]; ok { + return true + } + } + return false +} + func (m *Resolver) clearPending(dnsName string) { m.mutex.Lock() delete(m.pending, dnsName) diff --git a/client/internal/dns/mgmt/mgmt_test.go b/client/internal/dns/mgmt/mgmt_test.go index d7605a16f..91ea95440 100644 --- a/client/internal/dns/mgmt/mgmt_test.go +++ b/client/internal/dns/mgmt/mgmt_test.go @@ -325,7 +325,7 @@ func TestResolver_ManagementDomainProtection(t *testing.T) { Relay: []domain.Domain{"cloudflare.com"}, } - _, err = resolver.UpdateFromServerDomains(ctx, serverDomains) + _, err = resolver.UpdateFromServerDomains(ctx, serverDomains, true) if err != nil { t.Logf("Server domains update failed: %v", err) } @@ -363,7 +363,7 @@ func TestResolver_EmptyUpdateDoesNotRemoveDomains(t *testing.T) { } // Add initial domains - _, err := resolver.UpdateFromServerDomains(ctx, initialDomains) + _, err := resolver.UpdateFromServerDomains(ctx, initialDomains, true) if err != nil { t.Skipf("Skipping test due to DNS resolution failure: %v", err) } @@ -375,7 +375,7 @@ func TestResolver_EmptyUpdateDoesNotRemoveDomains(t *testing.T) { // Update with empty ServerDomains (simulating partial network map update) emptyDomains := dnsconfig.ServerDomains{} - removedDomains, err := resolver.UpdateFromServerDomains(ctx, emptyDomains) + removedDomains, err := resolver.UpdateFromServerDomains(ctx, emptyDomains, true) assert.NoError(t, err) // Verify no domains were removed @@ -398,7 +398,7 @@ func TestResolver_PartialUpdateReplacesOnlyUpdatedTypes(t *testing.T) { } // Add initial domains - _, err := resolver.UpdateFromServerDomains(ctx, initialDomains) + _, err := resolver.UpdateFromServerDomains(ctx, initialDomains, true) if err != nil { t.Skipf("Skipping test due to DNS resolution failure: %v", err) } @@ -409,7 +409,7 @@ func TestResolver_PartialUpdateReplacesOnlyUpdatedTypes(t *testing.T) { partialDomains := dnsconfig.ServerDomains{ Signal: "github.com", } - removedDomains, err := resolver.UpdateFromServerDomains(ctx, partialDomains) + removedDomains, err := resolver.UpdateFromServerDomains(ctx, partialDomains, true) if err != nil { t.Skipf("Skipping test due to DNS resolution failure: %v", err) } @@ -444,7 +444,7 @@ func TestResolver_PartialUpdateAddsNewTypePreservesExisting(t *testing.T) { } // Add initial domains - _, err := resolver.UpdateFromServerDomains(ctx, initialDomains) + _, err := resolver.UpdateFromServerDomains(ctx, initialDomains, true) if err != nil { t.Skipf("Skipping test due to DNS resolution failure: %v", err) } @@ -456,7 +456,7 @@ func TestResolver_PartialUpdateAddsNewTypePreservesExisting(t *testing.T) { partialDomains := dnsconfig.ServerDomains{ Flow: "github.com", } - removedDomains, err := resolver.UpdateFromServerDomains(ctx, partialDomains) + removedDomains, err := resolver.UpdateFromServerDomains(ctx, partialDomains, true) if err != nil { t.Skipf("Skipping test due to DNS resolution failure: %v", err) } diff --git a/client/internal/dns/server.go b/client/internal/dns/server.go index c4621ced0..7841145f9 100644 --- a/client/internal/dns/server.go +++ b/client/internal/dns/server.go @@ -613,7 +613,11 @@ func (s *DefaultServer) UpdateServerConfig(domains dnsconfig.ServerDomains) erro defer s.mux.Unlock() if s.mgmtCacheResolver != nil { - removedDomains, err := s.mgmtCacheResolver.UpdateFromServerDomains(s.ctx, domains) + // Mirrors the Initialize guard: without it NetBird never becomes the + // system resolver, so the mgmt cache is never queried and need not be + // primed synchronously. + dnsWillBeServed := !s.disableSys && !netstack.IsEnabled() + removedDomains, err := s.mgmtCacheResolver.UpdateFromServerDomains(s.ctx, domains, dnsWillBeServed) if err != nil { return fmt.Errorf("update management cache resolver: %w", err) }