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) }