diff --git a/client/internal/dns/mgmt/mgmt.go b/client/internal/dns/mgmt/mgmt.go index 2037e526e..2042e8f59 100644 --- a/client/internal/dns/mgmt/mgmt.go +++ b/client/internal/dns/mgmt/mgmt.go @@ -57,6 +57,9 @@ type pendingEntry struct{} // Resolver caches critical NetBird infrastructure domains. // records, refreshing, pending, mgmtDomain and serverDomains are all guarded by mutex. type Resolver struct { + // ctx is the server-lifetime context for background resolves. + ctx context.Context + records map[dns.Question]*cachedRecord mgmtDomain *domain.Domain serverDomains *dnsconfig.ServerDomains @@ -85,8 +88,9 @@ type Resolver struct { } // NewResolver creates a new management domains cache resolver. -func NewResolver() *Resolver { +func NewResolver(ctx context.Context) *Resolver { return &Resolver{ + ctx: ctx, records: make(map[dns.Question]*cachedRecord), refreshing: make(map[dns.Question]*atomic.Bool), pending: make(map[string]pendingEntry), @@ -613,7 +617,7 @@ func (m *Resolver) awaitPendingResolve(dnsName string) bool { ch := m.resolveGroup.DoChan(key, func() (any, error) { defer m.clearPending(dnsName) - if err := m.AddDomain(context.Background(), d); err != nil { + if err := m.AddDomain(m.ctx, d); err != nil { return struct{}{}, err } return struct{}{}, nil diff --git a/client/internal/dns/mgmt/mgmt_refresh_test.go b/client/internal/dns/mgmt/mgmt_refresh_test.go index 9faa5a0b8..8db264ba0 100644 --- a/client/internal/dns/mgmt/mgmt_refresh_test.go +++ b/client/internal/dns/mgmt/mgmt_refresh_test.go @@ -130,7 +130,7 @@ func TestResolver_CacheTTLGatesRefresh(t *testing.T) { q := dns.Question{Name: "mgmt.example.com.", Qtype: dns.TypeA, Qclass: dns.ClassINET} t.Run("short TTL treats entry as stale and refreshes", func(t *testing.T) { - r := NewResolver() + r := NewResolver(context.Background()) r.cacheTTL = 10 * time.Millisecond chain := newFakeChain() chain.setAnswer(q.Name, dns.TypeA, "10.0.0.2") @@ -146,7 +146,7 @@ func TestResolver_CacheTTLGatesRefresh(t *testing.T) { }) t.Run("long TTL keeps entry fresh and skips refresh", func(t *testing.T) { - r := NewResolver() + r := NewResolver(context.Background()) r.cacheTTL = time.Hour chain := newFakeChain() chain.setAnswer(q.Name, dns.TypeA, "10.0.0.2") @@ -162,7 +162,7 @@ func TestResolver_CacheTTLGatesRefresh(t *testing.T) { } func TestResolver_ServeFresh_NoRefresh(t *testing.T) { - r := NewResolver() + r := NewResolver(context.Background()) chain := newFakeChain() chain.setAnswer("mgmt.example.com.", dns.TypeA, "10.0.0.2") r.SetChainResolver(chain, 50) @@ -183,7 +183,7 @@ func TestResolver_ServeFresh_NoRefresh(t *testing.T) { } func TestResolver_StaleTriggersAsyncRefresh(t *testing.T) { - r := NewResolver() + r := NewResolver(context.Background()) chain := newFakeChain() chain.setAnswer("mgmt.example.com.", dns.TypeA, "10.0.0.2") r.SetChainResolver(chain, 50) @@ -213,7 +213,7 @@ func TestResolver_StaleTriggersAsyncRefresh(t *testing.T) { } func TestResolver_ConcurrentStaleHitsCollapseRefresh(t *testing.T) { - r := NewResolver() + r := NewResolver(context.Background()) chain := newFakeChain() chain.setAnswer("mgmt.example.com.", dns.TypeA, "10.0.0.2") @@ -262,7 +262,7 @@ func TestResolver_ConcurrentStaleHitsCollapseRefresh(t *testing.T) { } func TestResolver_RefreshFailureArmsBackoff(t *testing.T) { - r := NewResolver() + r := NewResolver(context.Background()) chain := newFakeChain() chain.err = errors.New("boom") r.SetChainResolver(chain, 50) @@ -299,7 +299,7 @@ func TestResolver_RefreshFailureArmsBackoff(t *testing.T) { } func TestResolver_NoRootHandler_SkipsChain(t *testing.T) { - r := NewResolver() + r := NewResolver(context.Background()) chain := newFakeChain() chain.hasRoot = false chain.setAnswer("mgmt.example.com.", dns.TypeA, "10.0.0.2") @@ -320,7 +320,7 @@ func TestResolver_ServeDuringRefreshSetsLoopFlag(t *testing.T) { // ServeDNS being invoked for a question while a refresh for that question // is inflight indicates a resolver loop (OS resolver sent the recursive // query back to us). The inflightRefresh.loopLoggedOnce flag must be set. - r := NewResolver() + r := NewResolver(context.Background()) q := dns.Question{Name: "mgmt.example.com.", Qtype: dns.TypeA, Qclass: dns.ClassINET} r.records[q] = &cachedRecord{ @@ -346,7 +346,7 @@ func TestResolver_ServeDuringRefreshSetsLoopFlag(t *testing.T) { } func TestResolver_LoopFlagOnlyTrippedOncePerRefresh(t *testing.T) { - r := NewResolver() + r := NewResolver(context.Background()) q := dns.Question{Name: "mgmt.example.com.", Qtype: dns.TypeA, Qclass: dns.ClassINET} r.records[q] = &cachedRecord{ @@ -373,7 +373,7 @@ func TestResolver_LoopFlagOnlyTrippedOncePerRefresh(t *testing.T) { } func TestResolver_NoLoopFlagWhenNotRefreshing(t *testing.T) { - r := NewResolver() + r := NewResolver(context.Background()) q := dns.Question{Name: "mgmt.example.com.", Qtype: dns.TypeA, Qclass: dns.ClassINET} r.records[q] = &cachedRecord{ @@ -393,7 +393,7 @@ func TestResolver_NoLoopFlagWhenNotRefreshing(t *testing.T) { } func TestResolver_AddDomain_UsesChainWhenRootRegistered(t *testing.T) { - r := NewResolver() + r := NewResolver(context.Background()) chain := newFakeChain() chain.setAnswer("mgmt.example.com.", dns.TypeA, "10.0.0.2") chain.setAnswer("mgmt.example.com.", dns.TypeAAAA, "fd00::2") diff --git a/client/internal/dns/mgmt/mgmt_test.go b/client/internal/dns/mgmt/mgmt_test.go index 921666664..d7605a16f 100644 --- a/client/internal/dns/mgmt/mgmt_test.go +++ b/client/internal/dns/mgmt/mgmt_test.go @@ -17,7 +17,7 @@ import ( ) func TestResolver_NewResolver(t *testing.T) { - resolver := NewResolver() + resolver := NewResolver(context.Background()) assert.NotNil(t, resolver) assert.NotNil(t, resolver.records) @@ -49,7 +49,7 @@ func TestResolveCacheTTL(t *testing.T) { func TestNewResolver_CacheTTLFromEnv(t *testing.T) { t.Setenv(envMgmtCacheTTL, "7s") - r := NewResolver() + r := NewResolver(context.Background()) assert.Equal(t, 7*time.Second, r.cacheTTL, "NewResolver should evaluate cacheTTL once from env") } @@ -169,7 +169,7 @@ func TestResolver_PopulateFromConfig(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) defer cancel() - resolver := NewResolver() + resolver := NewResolver(context.Background()) // Test with IP address - should return error since IP addresses are rejected mgmtURL, _ := url.Parse("https://127.0.0.1") @@ -184,7 +184,7 @@ func TestResolver_PopulateFromConfig(t *testing.T) { } func TestResolver_ServeDNS(t *testing.T) { - resolver := NewResolver() + resolver := NewResolver(context.Background()) ctx := context.Background() // Add a test domain to the cache - use example.org which is reserved for testing @@ -284,7 +284,7 @@ func TestResolver_ServeDNS(t *testing.T) { } func TestResolver_GetCachedDomains(t *testing.T) { - resolver := NewResolver() + resolver := NewResolver(context.Background()) ctx := context.Background() testDomain, err := domain.FromString("example.org") @@ -304,7 +304,7 @@ func TestResolver_GetCachedDomains(t *testing.T) { } func TestResolver_ManagementDomainProtection(t *testing.T) { - resolver := NewResolver() + resolver := NewResolver(context.Background()) ctx := context.Background() mgmtURL, _ := url.Parse("https://example.org") @@ -352,7 +352,7 @@ func extractDomainFromURL(u *url.URL) (domain.Domain, error) { } func TestResolver_EmptyUpdateDoesNotRemoveDomains(t *testing.T) { - resolver := NewResolver() + resolver := NewResolver(context.Background()) ctx := context.Background() // Set up initial domains using resolvable domains @@ -387,7 +387,7 @@ func TestResolver_EmptyUpdateDoesNotRemoveDomains(t *testing.T) { } func TestResolver_PartialUpdateReplacesOnlyUpdatedTypes(t *testing.T) { - resolver := NewResolver() + resolver := NewResolver(context.Background()) ctx := context.Background() // Set up initial complete domains using resolvable domains @@ -433,7 +433,7 @@ func TestResolver_PartialUpdateReplacesOnlyUpdatedTypes(t *testing.T) { } func TestResolver_PartialUpdateAddsNewTypePreservesExisting(t *testing.T) { - resolver := NewResolver() + resolver := NewResolver(context.Background()) ctx := context.Background() // Set up initial complete domains using resolvable domains diff --git a/client/internal/dns/server.go b/client/internal/dns/server.go index 994ca076d..c4621ced0 100644 --- a/client/internal/dns/server.go +++ b/client/internal/dns/server.go @@ -282,7 +282,7 @@ func newDefaultServer( handlerChain := NewHandlerChain() ctx, stop := context.WithCancel(ctx) - mgmtCacheResolver := mgmt.NewResolver() + mgmtCacheResolver := mgmt.NewResolver(ctx) mgmtCacheResolver.SetChainResolver(handlerChain, PriorityUpstream) defaultServer := &DefaultServer{