diff --git a/internal/aghtest/aghtest.go b/internal/aghtest/aghtest.go index 1d9067c5..fcf62be4 100644 --- a/internal/aghtest/aghtest.go +++ b/internal/aghtest/aghtest.go @@ -3,7 +3,6 @@ package aghtest import ( "crypto/sha256" - "io" "net/http" "net/http/httptest" "net/netip" @@ -12,7 +11,6 @@ import ( "time" "github.com/AdguardTeam/dnsproxy/proxy" - "github.com/AdguardTeam/golibs/log" "github.com/AdguardTeam/golibs/netutil" "github.com/AdguardTeam/golibs/testutil" "github.com/miekg/dns" @@ -27,33 +25,6 @@ const ( ReqFQDN = ReqHost + "." ) -// ReplaceLogWriter moves logger output to w and uses Cleanup method of t to -// revert changes. -func ReplaceLogWriter(t testing.TB, w io.Writer) { - t.Helper() - - prev := log.Writer() - t.Cleanup(func() { log.SetOutput(prev) }) - log.SetOutput(w) -} - -// ReplaceLogLevel sets logging level to l and uses Cleanup method of t to -// revert changes. -func ReplaceLogLevel(t testing.TB, l log.Level) { - t.Helper() - - switch l { - case log.INFO, log.DEBUG, log.ERROR: - // Go on. - default: - t.Fatalf("wrong l value (must be one of %v, %v, %v)", log.INFO, log.DEBUG, log.ERROR) - } - - prev := log.GetLevel() - t.Cleanup(func() { log.SetLevel(prev) }) - log.SetLevel(l) -} - // HostToIPs is a helper that generates one IPv4 and one IPv6 address from host. func HostToIPs(host string) (ipv4, ipv6 netip.Addr) { hash := sha256.Sum256([]byte(host)) diff --git a/internal/dnsforward/clientid_internal_test.go b/internal/dnsforward/clientid_internal_test.go index 171e23c4..ec110f60 100644 --- a/internal/dnsforward/clientid_internal_test.go +++ b/internal/dnsforward/clientid_internal_test.go @@ -8,7 +8,6 @@ import ( "testing" "github.com/AdguardTeam/dnsproxy/proxy" - "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/AdguardTeam/golibs/testutil" "github.com/stretchr/testify/assert" ) @@ -201,7 +200,7 @@ func TestServer_clientIDFromDNSContext(t *testing.T) { srv := &Server{ conf: ServerConfig{TLSConf: tlsConf}, - baseLogger: slogutil.NewDiscardLogger(), + baseLogger: testLogger, } var ( diff --git a/internal/dnsforward/dnsforward_internal_test.go b/internal/dnsforward/dnsforward_internal_test.go index 3889fb4c..8263b63c 100644 --- a/internal/dnsforward/dnsforward_internal_test.go +++ b/internal/dnsforward/dnsforward_internal_test.go @@ -63,6 +63,9 @@ const ( // TODO(a.garipov): Use more. var testClientAddrPort = netip.MustParseAddrPort("1.2.3.4:12345") +// testLogger is the common logger for tests. +var testLogger = slogutil.NewDiscardLogger() + // type check var _ ClientsContainer = (*clientsContainer)(nil) @@ -129,6 +132,8 @@ func createTestServer( ) (s *Server) { t.Helper() + filterConf.Logger = cmp.Or(filterConf.Logger, testLogger) + rules := `||nxdomain.example.org ||NULL.example.org^ 127.0.0.1 host.example.org @@ -159,7 +164,7 @@ func createTestServer( DHCPServer: dhcp, DNSFilter: f, PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed), - Logger: slogutil.NewDiscardLogger(), + Logger: testLogger, }) require.NoError(t, err) @@ -410,7 +415,7 @@ func TestServer_timeout(t *testing.T) { s, err := NewServer(DNSCreateParams{ DNSFilter: createTestDNSFilter(t), - Logger: slogutil.NewDiscardLogger(), + Logger: testLogger, }) require.NoError(t, err) @@ -423,7 +428,7 @@ func TestServer_timeout(t *testing.T) { t.Run("default", func(t *testing.T) { s, err := NewServer(DNSCreateParams{ DNSFilter: createTestDNSFilter(t), - Logger: slogutil.NewDiscardLogger(), + Logger: testLogger, }) require.NoError(t, err) @@ -456,7 +461,7 @@ func TestServer_Prepare_fallbacks(t *testing.T) { } s, err := NewServer(DNSCreateParams{ - Logger: slogutil.NewDiscardLogger(), + Logger: testLogger, }) require.NoError(t, err) @@ -584,6 +589,7 @@ func TestSafeSearch(t *testing.T) { } filterConf := &filtering.Config{ + Logger: testLogger, BlockingMode: filtering.BlockingModeDefault, ProtectionEnabled: true, SafeSearchConf: safeSearchConf, @@ -593,7 +599,7 @@ func TestSafeSearch(t *testing.T) { ctx := testutil.ContextWithTimeout(t, testTimeout) safeSearch, err := safesearch.NewDefault(ctx, &safesearch.DefaultConfig{ - Logger: slogutil.NewDiscardLogger(), + Logger: testLogger, ServicesConfig: safeSearchConf, CacheSize: filterConf.SafeSearchCacheSize, CacheTTL: time.Minute * time.Duration(filterConf.CacheTime), @@ -1055,6 +1061,7 @@ func TestBlockedCustomIP(t *testing.T) { }} f, err := filtering.New(&filtering.Config{ + Logger: testLogger, ProtectionEnabled: true, ApplyClientFiltering: applyEmptyClientFiltering, BlockedServices: emptyFilteringBlockedServices(), @@ -1073,7 +1080,7 @@ func TestBlockedCustomIP(t *testing.T) { DHCPServer: dhcp, DNSFilter: f, PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed), - Logger: slogutil.NewDiscardLogger(), + Logger: testLogger, }) require.NoError(t, err) @@ -1173,6 +1180,7 @@ func TestBlockedBySafeBrowsing(t *testing.T) { ) sbChecker := hashprefix.New(&hashprefix.Config{ + Logger: testLogger, CacheTime: cacheTime, CacheSize: cacheSize, Upstream: aghtest.NewBlockUpstream(hostname, true), @@ -1216,6 +1224,7 @@ func TestBlockedBySafeBrowsing(t *testing.T) { func TestRewrite(t *testing.T) { c := &filtering.Config{ + Logger: testLogger, ApplyClientFiltering: applyEmptyClientFiltering, BlockedServices: emptyFilteringBlockedServices(), BlockingMode: filtering.BlockingModeDefault, @@ -1247,7 +1256,7 @@ func TestRewrite(t *testing.T) { DHCPServer: dhcp, DNSFilter: f, PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed), - Logger: slogutil.NewDiscardLogger(), + Logger: testLogger, }) require.NoError(t, err) @@ -1365,6 +1374,7 @@ func TestPTRResponseFromDHCPLeases(t *testing.T) { const localDomain = "lan" flt, err := filtering.New(&filtering.Config{ + Logger: testLogger, ApplyClientFiltering: applyEmptyClientFiltering, BlockedServices: emptyFilteringBlockedServices(), BlockingMode: filtering.BlockingModeDefault, @@ -1381,7 +1391,7 @@ func TestPTRResponseFromDHCPLeases(t *testing.T) { }, }, PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed), - Logger: slogutil.NewDiscardLogger(), + Logger: testLogger, LocalDomain: localDomain, }) require.NoError(t, err) @@ -1457,6 +1467,7 @@ func TestPTRResponseFromHosts(t *testing.T) { }) flt, err := filtering.New(&filtering.Config{ + Logger: testLogger, ApplyClientFiltering: applyEmptyClientFiltering, BlockedServices: emptyFilteringBlockedServices(), BlockingMode: filtering.BlockingModeDefault, @@ -1471,7 +1482,7 @@ func TestPTRResponseFromHosts(t *testing.T) { DHCPServer: dhcp, DNSFilter: flt, PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed), - Logger: slogutil.NewDiscardLogger(), + Logger: testLogger, }) require.NoError(t, err) @@ -1527,27 +1538,27 @@ func TestNewServer(t *testing.T) { }{{ name: "success", in: DNSCreateParams{ - Logger: slogutil.NewDiscardLogger(), + Logger: testLogger, }, wantErrMsg: "", }, { name: "success_local_tld", in: DNSCreateParams{ - Logger: slogutil.NewDiscardLogger(), + Logger: testLogger, LocalDomain: "mynet", }, wantErrMsg: "", }, { name: "success_local_domain", in: DNSCreateParams{ - Logger: slogutil.NewDiscardLogger(), + Logger: testLogger, LocalDomain: "my.local.net", }, wantErrMsg: "", }, { name: "bad_local_domain", in: DNSCreateParams{ - Logger: slogutil.NewDiscardLogger(), + Logger: testLogger, LocalDomain: "!!!", }, wantErrMsg: `local domain: bad domain name "!!!": ` + diff --git a/internal/dnsforward/filter_internal_test.go b/internal/dnsforward/filter_internal_test.go index b14df3f2..8a07b4f1 100644 --- a/internal/dnsforward/filter_internal_test.go +++ b/internal/dnsforward/filter_internal_test.go @@ -9,7 +9,6 @@ import ( "github.com/AdguardTeam/AdGuardHome/internal/filtering" "github.com/AdguardTeam/dnsproxy/proxy" "github.com/AdguardTeam/dnsproxy/upstream" - "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/AdguardTeam/golibs/netutil" "github.com/miekg/dns" "github.com/stretchr/testify/assert" @@ -46,6 +45,7 @@ func TestHandleDNSRequest_handleDNSRequest(t *testing.T) { }} f, err := filtering.New(&filtering.Config{ + Logger: testLogger, ProtectionEnabled: true, ApplyClientFiltering: applyEmptyClientFiltering, BlockedServices: emptyFilteringBlockedServices(), @@ -62,7 +62,7 @@ func TestHandleDNSRequest_handleDNSRequest(t *testing.T) { }, DNSFilter: f, PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed), - Logger: slogutil.NewDiscardLogger(), + Logger: testLogger, }) require.NoError(t, err) @@ -226,7 +226,9 @@ func TestHandleDNSRequest_filterDNSResponse(t *testing.T) { ID: 0, Data: []byte(blockRules), }} - f, err := filtering.New(&filtering.Config{}, filters) + f, err := filtering.New(&filtering.Config{ + Logger: testLogger, + }, filters) require.NoError(t, err) f.SetEnabled(true) @@ -235,7 +237,7 @@ func TestHandleDNSRequest_filterDNSResponse(t *testing.T) { DHCPServer: &testDHCP{}, DNSFilter: f, PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed), - Logger: slogutil.NewDiscardLogger(), + Logger: testLogger, }) require.NoError(t, err) diff --git a/internal/dnsforward/ipset_internal_test.go b/internal/dnsforward/ipset_internal_test.go index 09601ac6..90c200d0 100644 --- a/internal/dnsforward/ipset_internal_test.go +++ b/internal/dnsforward/ipset_internal_test.go @@ -6,7 +6,6 @@ import ( "testing" "github.com/AdguardTeam/dnsproxy/proxy" - "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/miekg/dns" "github.com/stretchr/testify/assert" ) @@ -61,7 +60,7 @@ func TestIpsetCtx_process(t *testing.T) { } ictx := &ipsetHandler{ - logger: slogutil.NewDiscardLogger(), + logger: testLogger, } rc := ictx.process(dctx) assert.Equal(t, resultCodeSuccess, rc) @@ -83,7 +82,7 @@ func TestIpsetCtx_process(t *testing.T) { m := &fakeIpsetMgr{} ictx := &ipsetHandler{ ipsetMgr: m, - logger: slogutil.NewDiscardLogger(), + logger: testLogger, } rc := ictx.process(dctx) @@ -108,7 +107,7 @@ func TestIpsetCtx_process(t *testing.T) { m := &fakeIpsetMgr{} ictx := &ipsetHandler{ ipsetMgr: m, - logger: slogutil.NewDiscardLogger(), + logger: testLogger, } rc := ictx.process(dctx) @@ -132,7 +131,7 @@ func TestIpsetCtx_SkipIpsetProcessing(t *testing.T) { m := &fakeIpsetMgr{} ictx := &ipsetHandler{ ipsetMgr: m, - logger: slogutil.NewDiscardLogger(), + logger: testLogger, } testCases := []struct { diff --git a/internal/dnsforward/process_internal_test.go b/internal/dnsforward/process_internal_test.go index 71a91fdd..8b335832 100644 --- a/internal/dnsforward/process_internal_test.go +++ b/internal/dnsforward/process_internal_test.go @@ -12,7 +12,6 @@ import ( "github.com/AdguardTeam/AdGuardHome/internal/filtering" "github.com/AdguardTeam/dnsproxy/proxy" "github.com/AdguardTeam/dnsproxy/upstream" - "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/AdguardTeam/golibs/netutil" "github.com/AdguardTeam/golibs/testutil" "github.com/AdguardTeam/urlfilter/rules" @@ -378,6 +377,7 @@ func createTestDNSFilter(t *testing.T) (f *filtering.DNSFilter) { t.Helper() f, err := filtering.New(&filtering.Config{ + Logger: testLogger, BlockingMode: filtering.BlockingModeDefault, }, []filtering.Filter{}) require.NoError(t, err) @@ -439,7 +439,7 @@ func TestServer_ProcessDHCPHosts_localRestriction(t *testing.T) { dnsFilter: createTestDNSFilter(t), dhcpServer: dhcp, localDomainSuffix: localDomainSuffix, - baseLogger: slogutil.NewDiscardLogger(), + baseLogger: testLogger, } req := &dns.Msg{ @@ -591,7 +591,7 @@ func TestServer_ProcessDHCPHosts(t *testing.T) { dnsFilter: createTestDNSFilter(t), dhcpServer: testDHCP, localDomainSuffix: tc.suffix, - baseLogger: slogutil.NewDiscardLogger(), + baseLogger: testLogger, } req := (&dns.Msg{}).SetQuestion(dns.Fqdn(tc.host), tc.qtyp) diff --git a/internal/dnsforward/stats_internal_test.go b/internal/dnsforward/stats_internal_test.go index 6e4d5d86..301f8c8b 100644 --- a/internal/dnsforward/stats_internal_test.go +++ b/internal/dnsforward/stats_internal_test.go @@ -11,7 +11,6 @@ import ( "github.com/AdguardTeam/AdGuardHome/internal/stats" "github.com/AdguardTeam/dnsproxy/proxy" "github.com/AdguardTeam/dnsproxy/upstream" - "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/miekg/dns" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -203,7 +202,7 @@ func TestServer_ProcessQueryLogsAndStats(t *testing.T) { ql := &testQueryLog{} st := &testStats{} srv := &Server{ - baseLogger: slogutil.NewDiscardLogger(), + baseLogger: testLogger, queryLog: ql, stats: st, anonymizer: aghnet.NewIPMut(nil), diff --git a/internal/filtering/blocked.go b/internal/filtering/blocked.go index ca59a1b8..8150f309 100644 --- a/internal/filtering/blocked.go +++ b/internal/filtering/blocked.go @@ -1,8 +1,10 @@ package filtering import ( + "context" "encoding/json" "fmt" + "log/slog" "net/http" "slices" "time" @@ -10,7 +12,7 @@ import ( "github.com/AdguardTeam/AdGuardHome/internal/aghhttp" "github.com/AdguardTeam/AdGuardHome/internal/filtering/rulelist" "github.com/AdguardTeam/AdGuardHome/internal/schedule" - "github.com/AdguardTeam/golibs/log" + "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/AdguardTeam/urlfilter/rules" ) @@ -20,23 +22,30 @@ var serviceRules map[string][]*rules.NetworkRule // serviceIDs contains service IDs sorted alphabetically. var serviceIDs []string -// initBlockedServices initializes package-level blocked service data. -func initBlockedServices() { - l := len(blockedServices) - serviceIDs = make([]string, l) - serviceRules = make(map[string][]*rules.NetworkRule, l) +// initBlockedServices initializes package-level blocked service data. l must +// not be nil. +func initBlockedServices(ctx context.Context, l *slog.Logger) { + svcLen := len(blockedServices) + serviceIDs = make([]string, svcLen) + serviceRules = make(map[string][]*rules.NetworkRule, svcLen) for i, s := range blockedServices { netRules := make([]*rules.NetworkRule, 0, len(s.Rules)) for _, text := range s.Rules { rule, err := rules.NewNetworkRule(text, rulelist.URLFilterIDBlockedService) - if err != nil { - log.Error("parsing blocked service %q rule %q: %s", s.ID, text, err) + if err == nil { + netRules = append(netRules, rule) continue } - netRules = append(netRules, rule) + l.ErrorContext( + ctx, + "parsing blocked service rule", + "svc", s.ID, + "rule", text, + slogutil.KeyError, err, + ) } serviceIDs[i] = s.ID @@ -45,7 +54,7 @@ func initBlockedServices() { slices.Sort(serviceIDs) - log.Debug("filtering: initialized %d services", l) + l.DebugContext(ctx, "initialized services", "svc_len", svcLen) } // BlockedServices is the configuration of blocked services. @@ -105,7 +114,7 @@ func (d *DNSFilter) ApplyBlockedServicesList(setts *Settings, list []string) { for _, name := range list { rules, ok := serviceRules[name] if !ok { - log.Error("unknown service name: %s", name) + d.logger.ErrorContext(context.TODO(), "unknown service name", "name", name) continue } @@ -163,7 +172,7 @@ func (d *DNSFilter) handleBlockedServicesSet(w http.ResponseWriter, r *http.Requ defer d.confMu.Unlock() d.conf.BlockedServices.IDs = list - log.Debug("Updated blocked services list: %d", len(list)) + d.logger.DebugContext(r.Context(), "updated blocked services list", "len", len(list)) }() d.conf.ConfigModified() @@ -212,7 +221,7 @@ func (d *DNSFilter) handleBlockedServicesUpdate(w http.ResponseWriter, r *http.R d.conf.BlockedServices = bsvc }() - log.Debug("updated blocked services schedule: %d", len(bsvc.IDs)) + d.logger.DebugContext(r.Context(), "updated blocked services schedule", "len", len(bsvc.IDs)) d.conf.ConfigModified() } diff --git a/internal/filtering/dnsrewrite_test.go b/internal/filtering/dnsrewrite_test.go index 89b6b30d..58353a43 100644 --- a/internal/filtering/dnsrewrite_test.go +++ b/internal/filtering/dnsrewrite_test.go @@ -6,6 +6,7 @@ import ( "testing" "github.com/AdguardTeam/AdGuardHome/internal/filtering" + "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/AdguardTeam/golibs/netutil" "github.com/miekg/dns" "github.com/stretchr/testify/assert" @@ -52,6 +53,7 @@ func TestDNSFilter_CheckHostRules_dnsrewrite(t *testing.T) { ` conf := &filtering.Config{ + Logger: slogutil.NewDiscardLogger(), SafeBrowsingCacheSize: 10000, ParentalCacheSize: 10000, SafeSearchCacheSize: 1000, diff --git a/internal/filtering/filter.go b/internal/filtering/filter.go index 14572d01..240b623d 100644 --- a/internal/filtering/filter.go +++ b/internal/filtering/filter.go @@ -1,6 +1,7 @@ package filtering import ( + "context" "fmt" "io" "net/http" @@ -17,7 +18,7 @@ import ( "github.com/AdguardTeam/AdGuardHome/internal/filtering/rulelist" "github.com/AdguardTeam/golibs/container" "github.com/AdguardTeam/golibs/errors" - "github.com/AdguardTeam/golibs/log" + "github.com/AdguardTeam/golibs/logutil/slogutil" ) // filterDir is the subdirectory of a data directory to store downloaded @@ -105,12 +106,13 @@ func (d *DNSFilter) filterSetProperties( } flt := &filters[i] - log.Debug( - "filtering: set name to %q, url to %s, enabled to %t for filter %s", - newList.Name, - newList.URL, - newList.Enabled, - flt.URL, + d.logger.DebugContext( + context.TODO(), + "updating filter", + "name", newList.Name, + "url", newList.URL, + "enabled", newList.Enabled, + "filter_url", flt.URL, ) defer func(oldURL, oldName string, oldEnabled bool, oldUpdated time.Time, oldRulesCount int) { @@ -213,12 +215,12 @@ func (d *DNSFilter) filterAdd(flt FilterYAML) (err error) { // Load filters from the disk // And if any filter has zero ID, assign a new one -func (d *DNSFilter) loadFilters(array []FilterYAML) { +func (d *DNSFilter) loadFilters(ctx context.Context, array []FilterYAML) { for i := range array { filter := &array[i] // otherwise we're operating on a copy if filter.ID == 0 { newID := d.idGen.next() - log.Info("filtering: warning: filter at index %d has no id; assigning to %d", i, newID) + d.logger.WarnContext(ctx, "filter has no id", "idx", i, "new_id", newID) filter.ID = newID } @@ -228,9 +230,9 @@ func (d *DNSFilter) loadFilters(array []FilterYAML) { continue } - err := d.load(filter) + err := d.load(ctx, filter) if err != nil { - log.Error("filtering: loading filter %d: %s", filter.ID, err) + d.logger.ErrorContext(ctx, "loading filter", "id", filter.ID, slogutil.KeyError, err) } } } @@ -300,10 +302,15 @@ func (d *DNSFilter) listsToUpdate(filters *[]FilterYAML, force bool) (toUpd []Fi return toUpd } -func (d *DNSFilter) refreshFiltersArray(filters *[]FilterYAML, force bool) (int, []FilterYAML, []bool, bool) { - var updateFlags []bool // 'true' if filter data has changed - - updateFilters := d.listsToUpdate(filters, force) +// refreshFiltersArray updates the filters array and returns the number of +// filters that have been refreshed. updateFlags is true if filter data has +// changed. +func (d *DNSFilter) refreshFiltersArray( + ctx context.Context, + filters *[]FilterYAML, + force bool, +) (updateCount int, updateFilters []FilterYAML, updateFlags []bool, isNetErr bool) { + updateFilters = d.listsToUpdate(filters, force) if len(updateFilters) == 0 { return 0, nil, nil, false } @@ -315,7 +322,7 @@ func (d *DNSFilter) refreshFiltersArray(filters *[]FilterYAML, force bool) (int, updateFlags = append(updateFlags, updated) if err != nil { failNum++ - log.Error("filtering: updating filter from url %q: %s\n", uf.URL, err) + d.logger.ErrorContext(ctx, "updating filter", "url", uf.URL, slogutil.KeyError, err) continue } @@ -325,8 +332,6 @@ func (d *DNSFilter) refreshFiltersArray(filters *[]FilterYAML, force bool) (int, return 0, nil, nil, true } - updateCount := 0 - d.conf.filtersMu.Lock() defer d.conf.filtersMu.Unlock() @@ -345,11 +350,12 @@ func (d *DNSFilter) refreshFiltersArray(filters *[]FilterYAML, force bool) (int, continue } - log.Info( - "filtering: updated filter %d; rule count: %d (was %d)", - f.ID, - uf.RulesCount, - f.RulesCount, + d.logger.InfoContext( + ctx, + "updated filter", + "id", f.ID, + "rules_count", uf.RulesCount, + "prev_rules_count", f.RulesCount, ) f.Name = uf.Name @@ -381,19 +387,27 @@ func (d *DNSFilter) refreshFiltersArray(filters *[]FilterYAML, force bool) (int, // // TODO(a.garipov, e.burkov): What the hell? func (d *DNSFilter) refreshFiltersIntl(block, allow, force bool) (int, bool) { + ctx := context.TODO() + updNum := 0 - log.Debug("filtering: starting updating") - defer func() { log.Debug("filtering: finished updating, %d updated", updNum) }() + d.logger.DebugContext(ctx, "starting update") + defer func() { + d.logger.DebugContext(ctx, "finished update", "updated", updNum) + }() var lists []FilterYAML var toUpd []bool isNetErr := false if block { - updNum, lists, toUpd, isNetErr = d.refreshFiltersArray(&d.conf.Filters, force) + updNum, lists, toUpd, isNetErr = d.refreshFiltersArray(ctx, &d.conf.Filters, force) } if allow { - updNumAl, listsAl, toUpdAl, isNetErrAl := d.refreshFiltersArray(&d.conf.WhitelistFilters, force) + updNumAl, listsAl, toUpdAl, isNetErrAl := d.refreshFiltersArray( + ctx, + &d.conf.WhitelistFilters, + force, + ) updNum += updNumAl lists = append(lists, listsAl...) @@ -417,7 +431,7 @@ func (d *DNSFilter) refreshFiltersIntl(block, allow, force bool) (int, bool) { p := uf.Path(d.conf.DataDir) err := os.Remove(p + ".old") if err != nil { - log.Debug("filtering: removing old filter file %q: %s", p, err) + d.logger.ErrorContext(ctx, "removing old filter", "path", p, slogutil.KeyError, err) } } } @@ -427,7 +441,9 @@ func (d *DNSFilter) refreshFiltersIntl(block, allow, force bool) (int, bool) { // update refreshes filter's content and a/mtimes of it's file. func (d *DNSFilter) update(filter *FilterYAML) (b bool, err error) { - b, err = d.updateIntl(filter) + ctx := context.TODO() + + b, err = d.updateIntl(ctx, filter) filter.LastUpdated = time.Now() if !b { chErr := os.Chtimes( @@ -436,7 +452,7 @@ func (d *DNSFilter) update(filter *FilterYAML) (b bool, err error) { filter.LastUpdated, ) if chErr != nil { - log.Error("filtering: os.Chtimes(): %s", chErr) + d.logger.ErrorContext(ctx, "changing last modified time", slogutil.KeyError, chErr) } } @@ -445,8 +461,8 @@ func (d *DNSFilter) update(filter *FilterYAML) (b bool, err error) { // updateIntl updates the flt rewriting it's actual file. It returns true if // the actual update has been performed. -func (d *DNSFilter) updateIntl(flt *FilterYAML) (ok bool, err error) { - log.Debug("filtering: downloading update for filter %d from %q", flt.ID, flt.URL) +func (d *DNSFilter) updateIntl(ctx context.Context, flt *FilterYAML) (ok bool, err error) { + d.logger.DebugContext(ctx, "downloading update for filter", "id", flt.ID, "url", flt.URL) var res *rulelist.ParseResult @@ -454,7 +470,7 @@ func (d *DNSFilter) updateIntl(flt *FilterYAML) (ok bool, err error) { if err != nil { return false, err } - defer func() { err = d.finalizeUpdate(tmpFile, flt, res, err, ok) }() + defer func() { err = d.finalizeUpdate(ctx, tmpFile, flt, res, err, ok) }() r, err := d.reader(flt.URL) if err != nil { @@ -476,6 +492,7 @@ func (d *DNSFilter) updateIntl(flt *FilterYAML) (ok bool, err error) { // according to updated. It also saves new values of flt's name, rules number // and checksum if succeeded. func (d *DNSFilter) finalizeUpdate( + ctx context.Context, file aghrenameio.PendingFile, flt *FilterYAML, res *rulelist.ParseResult, @@ -485,13 +502,13 @@ func (d *DNSFilter) finalizeUpdate( id := flt.ID if !updated { if returned == nil { - log.Debug("filtering: filter %d from url %q has no changes, skipping", id, flt.URL) + d.logger.DebugContext(ctx, "skipping filter with no changes", "id", id, "url", flt.URL) } return errors.WithDeferred(returned, file.Cleanup()) } - log.Info("filtering: saving contents of filter %d into %q", id, flt.Path(d.conf.DataDir)) + d.logger.InfoContext(ctx, "saving contents", "id", id, "path", flt.Path(d.conf.DataDir)) err = file.CloseReplace() if err != nil { @@ -499,7 +516,13 @@ func (d *DNSFilter) finalizeUpdate( } rulesCount := res.RulesCount - log.Info("filtering: updated filter %d: %d bytes, %d rules", id, res.BytesWritten, rulesCount) + d.logger.InfoContext( + ctx, + "filter updated", + "id", id, + "bytes_written", res.BytesWritten, + "rules_count", rulesCount, + ) flt.ensureName(res.Title) flt.checksum = res.Checksum @@ -550,10 +573,10 @@ func (d *DNSFilter) readerFromURL(fltURL string) (r io.ReadCloser, err error) { } // loads filter contents from the file in dataDir -func (d *DNSFilter) load(flt *FilterYAML) (err error) { +func (d *DNSFilter) load(ctx context.Context, flt *FilterYAML) (err error) { fileName := flt.Path(d.conf.DataDir) - log.Debug("filtering: loading filter %d from %q", flt.ID, fileName) + d.logger.DebugContext(ctx, "loading filter", "id", flt.ID, "path", fileName) file, err := os.Open(fileName) if errors.Is(err, os.ErrNotExist) { @@ -569,7 +592,7 @@ func (d *DNSFilter) load(flt *FilterYAML) (err error) { return fmt.Errorf("getting filter file stat: %w", err) } - log.Debug("filtering: file %q, id %d, length %d", fileName, flt.ID, st.Size()) + d.logger.DebugContext(ctx, "filter file", "id", flt.ID, "path", fileName, "len", st.Size()) bufPtr := d.bufPool.Get() defer d.bufPool.Put(bufPtr) @@ -586,14 +609,16 @@ func (d *DNSFilter) load(flt *FilterYAML) (err error) { return nil } +// EnableFilters enables filters. func (d *DNSFilter) EnableFilters(async bool) { d.conf.filtersMu.RLock() defer d.conf.filtersMu.RUnlock() - d.enableFiltersLocked(async) + d.enableFiltersLocked(context.TODO(), async) } -func (d *DNSFilter) enableFiltersLocked(async bool) { +// enableFiltersLocked enables filters under the conf.filtersMu lock. +func (d *DNSFilter) enableFiltersLocked(ctx context.Context, async bool) { filters := make([]Filter, 1, len(d.conf.Filters)+len(d.conf.WhitelistFilters)+1) filters[0] = Filter{ ID: rulelist.URLFilterIDCustom, @@ -623,9 +648,9 @@ func (d *DNSFilter) enableFiltersLocked(async bool) { }) } - err := d.setFilters(filters, allowFilters, async) + err := d.setFilters(ctx, filters, allowFilters, async) if err != nil { - log.Error("filtering: enabling filters: %s", err) + d.logger.ErrorContext(ctx, "enabling filters", slogutil.KeyError, err) } d.SetEnabled(d.conf.FilteringEnabled) diff --git a/internal/filtering/filter_internal_test.go b/internal/filtering/filter_internal_test.go index 8cfcdef9..c2d0de71 100644 --- a/internal/filtering/filter_internal_test.go +++ b/internal/filtering/filter_internal_test.go @@ -1,6 +1,7 @@ package filtering import ( + "context" "net" "net/http" "net/url" @@ -9,6 +10,7 @@ import ( "testing" "time" + "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/AdguardTeam/golibs/netutil/urlutil" "github.com/AdguardTeam/golibs/testutil" "github.com/stretchr/testify/assert" @@ -56,6 +58,7 @@ func serveFiltersLocally(t *testing.T, fltContent []byte) (urlStr string) { // count. func updateAndAssert( t *testing.T, + ctx context.Context, dnsFilter *DNSFilter, f *FilterYAML, wantUpd require.BoolAssertionFunc, @@ -75,7 +78,7 @@ func updateAndAssert( assert.Len(t, dir, 1) - err = dnsFilter.load(f) + err = dnsFilter.load(ctx, f) require.NoError(t, err) } @@ -84,6 +87,7 @@ func newDNSFilter(t *testing.T) (d *DNSFilter) { t.Helper() dnsFilter, err := New(&Config{ + Logger: slogutil.NewDiscardLogger(), DataDir: t.TempDir(), HTTPClient: &http.Client{ Timeout: testTimeout, @@ -95,6 +99,8 @@ func newDNSFilter(t *testing.T) (d *DNSFilter) { } func TestDNSFilter_Update(t *testing.T) { + ctx := testutil.ContextWithTimeout(t, testTimeout) + const content = `||example.org^$third-party # Inline comment example ||example.com^$third-party @@ -111,11 +117,11 @@ func TestDNSFilter_Update(t *testing.T) { dnsFilter := newDNSFilter(t) t.Run("download", func(t *testing.T) { - updateAndAssert(t, dnsFilter, f, require.True, 3) + updateAndAssert(t, ctx, dnsFilter, f, require.True, 3) }) t.Run("refresh_idle", func(t *testing.T) { - updateAndAssert(t, dnsFilter, f, require.False, 3) + updateAndAssert(t, ctx, dnsFilter, f, require.False, 3) }) t.Run("refresh_actually", func(t *testing.T) { @@ -125,11 +131,11 @@ func TestDNSFilter_Update(t *testing.T) { f.URL = serveFiltersLocally(t, anotherContent) t.Cleanup(func() { f.URL = oldURL }) - updateAndAssert(t, dnsFilter, f, require.True, 1) + updateAndAssert(t, ctx, dnsFilter, f, require.True, 1) }) t.Run("load_unload", func(t *testing.T) { - err := dnsFilter.load(f) + err := dnsFilter.load(ctx, f) require.NoError(t, err) f.unload() @@ -137,6 +143,8 @@ func TestDNSFilter_Update(t *testing.T) { } func TestFilterYAML_EnsureName(t *testing.T) { + ctx := testutil.ContextWithTimeout(t, testTimeout) + dnsFilter := newDNSFilter(t) t.Run("title_custom", func(t *testing.T) { @@ -147,7 +155,7 @@ func TestFilterYAML_EnsureName(t *testing.T) { Name: "user-custom", } - updateAndAssert(t, dnsFilter, f, require.True, 1) + updateAndAssert(t, ctx, dnsFilter, f, require.True, 1) assert.Equal(t, "user-custom", f.Name) }) @@ -158,7 +166,7 @@ func TestFilterYAML_EnsureName(t *testing.T) { URL: serveFiltersLocally(t, content), } - updateAndAssert(t, dnsFilter, f, require.True, 1) + updateAndAssert(t, ctx, dnsFilter, f, require.True, 1) assert.Equal(t, "src-title", f.Name) }) @@ -169,7 +177,7 @@ func TestFilterYAML_EnsureName(t *testing.T) { URL: serveFiltersLocally(t, content), } - updateAndAssert(t, dnsFilter, f, require.True, 1) + updateAndAssert(t, ctx, dnsFilter, f, require.True, 1) assert.Equal(t, "List 0", f.Name) }) } diff --git a/internal/filtering/filtering.go b/internal/filtering/filtering.go index df65308f..41d437c3 100644 --- a/internal/filtering/filtering.go +++ b/internal/filtering/filtering.go @@ -5,6 +5,7 @@ import ( "context" "fmt" "io/fs" + "log/slog" "net" "net/http" "net/netip" @@ -24,7 +25,7 @@ import ( "github.com/AdguardTeam/golibs/container" "github.com/AdguardTeam/golibs/errors" "github.com/AdguardTeam/golibs/hostsfile" - "github.com/AdguardTeam/golibs/log" + "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/AdguardTeam/golibs/mathutil" "github.com/AdguardTeam/golibs/syncutil" "github.com/AdguardTeam/urlfilter" @@ -70,6 +71,10 @@ type Resolver interface { // Config allows you to configure DNS filtering with New() or just change variables directly. type Config struct { + // logger is used to log the operations of DNS filtering. It must not be + // nil. + Logger *slog.Logger `yaml:"-"` + // BlockingIPv4 is the IP address to be returned for a blocked A request. BlockingIPv4 netip.Addr `yaml:"blocking_ipv4"` @@ -235,6 +240,9 @@ type Checker interface { // DNSFilter matches hostnames and DNS requests against filtering rules. type DNSFilter struct { + // logger is used for logging the filtering process. + logger *slog.Logger + // idGen is used to generate IDs for package urlfilter. idGen *idGenerator @@ -413,7 +421,12 @@ func (d *DNSFilter) WriteDiskConfig(c *Config) { // filters are ready. // // In this case the caller must ensure that the old filter files are intact. -func (d *DNSFilter) setFilters(blockFilters, allowFilters []Filter, async bool) error { +func (d *DNSFilter) setFilters( + ctx context.Context, + blockFilters []Filter, + allowFilters []Filter, + async bool, +) (err error) { if async { params := filtersInitializerParams{ allowFilters: allowFilters, @@ -439,7 +452,7 @@ func (d *DNSFilter) setFilters(blockFilters, allowFilters []Filter, async bool) return nil } - return d.initFiltering(allowFilters, blockFilters) + return d.initFiltering(ctx, allowFilters, blockFilters) } // Close - close the object @@ -451,19 +464,19 @@ func (d *DNSFilter) Close() { d.done <- struct{}{} } - d.reset() + d.reset(context.TODO()) } -func (d *DNSFilter) reset() { +func (d *DNSFilter) reset(ctx context.Context) { if d.rulesStorage != nil { if err := d.rulesStorage.Close(); err != nil { - log.Error("filtering: rulesStorage.Close: %s", err) + d.logger.ErrorContext(ctx, "closing rules storage", slogutil.KeyError, err) } } if d.rulesStorageAllow != nil { if err := d.rulesStorageAllow.Close(); err != nil { - log.Error("filtering: rulesStorageAllow.Close: %s", err) + d.logger.ErrorContext(ctx, "closing allow rules storage", slogutil.KeyError, err) } } } @@ -649,6 +662,8 @@ func (d *DNSFilter) processRewrites(host string, qtype uint16) (res Result) { d.confMu.RLock() defer d.confMu.RUnlock() + ctx := context.TODO() + rewrites, matched := findRewrites(d.conf.Rewrites, host, qtype) if !matched { return Result{} @@ -663,7 +678,7 @@ func (d *DNSFilter) processRewrites(host string, qtype uint16) (res Result) { rwPat := rw.Domain rwAns := rw.Answer - log.Debug("rewrite: cname for %s is %s", host, rwAns) + d.logger.DebugContext(ctx, "found rewrite", "host", host, "cname", rwAns) if origHost == rwAns || rwPat == rwAns { // Either a request for the hostname itself or a rewrite of @@ -682,7 +697,7 @@ func (d *DNSFilter) processRewrites(host string, qtype uint16) (res Result) { host = rwAns if cnames.Has(host) { - log.Info("rewrite: cname loop for %q on %q", origHost, host) + d.logger.InfoContext(ctx, "cname loop", "host", host, "original", origHost) return res } @@ -692,15 +707,15 @@ func (d *DNSFilter) processRewrites(host string, qtype uint16) (res Result) { rewrites, matched = findRewrites(d.conf.Rewrites, host, qtype) } - setRewriteResult(&res, host, rewrites, qtype) + d.setRewriteResult(ctx, &res, host, rewrites, qtype) return res } // matchBlockedServicesRules checks the host against the blocked services rules -// in settings, if any. The err is always nil, it is only there to make this -// a valid hostChecker function. -func matchBlockedServicesRules( +// in settings, if any. err is always nil, it is only there to make this a +// valid hostChecker function. +func (d *DNSFilter) matchBlockedServicesRules( host string, _ uint16, setts *Settings, @@ -728,8 +743,13 @@ func matchBlockedServicesRules( Text: ruleText, }} - log.Debug("blocked services: matched rule: %s host: %s service: %s", - ruleText, host, s.Name) + d.logger.DebugContext( + context.TODO(), + "blocked services matched rule", + "rule", ruleText, + "host", host, + "service", s.Name, + ) return res, nil } @@ -793,7 +813,7 @@ func newRuleStorage(filters []Filter) (rs *filterlist.RuleStorage, err error) { } // Initialize urlfilter objects. -func (d *DNSFilter) initFiltering(allowFilters, blockFilters []Filter) (err error) { +func (d *DNSFilter) initFiltering(ctx context.Context, allowFilters, blockFilters []Filter) (err error) { rulesStorage, err := newRuleStorage(blockFilters) if err != nil { return err @@ -811,7 +831,7 @@ func (d *DNSFilter) initFiltering(allowFilters, blockFilters []Filter) (err erro d.engineLock.Lock() defer d.engineLock.Unlock() - d.reset() + d.reset(ctx) d.rulesStorage = rulesStorage d.filteringEngine = filteringEngine d.rulesStorageAllow = rulesStorageAllow @@ -821,7 +841,7 @@ func (d *DNSFilter) initFiltering(allowFilters, blockFilters []Filter) (err erro // Make sure that the OS reclaims memory as soon as possible. debug.FreeOSMemory() - log.Debug("filtering: initialized filtering engine") + d.logger.DebugContext(ctx, "initialized filtering engine") return nil } @@ -843,6 +863,7 @@ func hostRulesToRules(netRules []*rules.HostRule) (res []rules.Rule) { // matchHostProcessAllowList processes the allowlist logic of host matching. func (d *DNSFilter) matchHostProcessAllowList( + ctx context.Context, host string, dnsres *urlfilter.DNSResult, ) (res Result, err error) { @@ -859,7 +880,12 @@ func (d *DNSFilter) matchHostProcessAllowList( return Result{}, fmt.Errorf("invalid dns result: rules are empty") } - log.Debug("filtering: allowlist rules for host %q: %+v", host, matchedRules) + d.logger.DebugContext( + ctx, + "allowlist rules for host", + "host", host, + "rules", matchedRules, + ) return makeResult(matchedRules, NotFilteredAllowList), nil } @@ -929,6 +955,8 @@ func (d *DNSFilter) matchHost( return Result{}, nil } + ctx := context.TODO() + ufReq := &urlfilter.DNSRequest{ Hostname: host, SortedClientTags: setts.ClientTags, @@ -947,7 +975,7 @@ func (d *DNSFilter) matchHost( if setts.ProtectionEnabled && d.filteringEngineAllow != nil { dnsres, ok := d.filteringEngineAllow.MatchRequest(ufReq) if ok { - return d.matchHostProcessAllowList(host, dnsres) + return d.matchHostProcessAllowList(ctx, host, dnsres) } } @@ -972,11 +1000,12 @@ func (d *DNSFilter) matchHost( res = d.matchHostProcessDNSResult(rrtype, dnsres) for _, r := range res.Rules { - log.Debug( - "filtering: found rule %q for host %q, filter list id: %d", - r.Text, - host, - r.FilterListID, + d.logger.DebugContext( + ctx, + "found rule for host", + "host", host, + "rule", r.Text, + "filter_list_id", r.FilterListID, ) } @@ -1000,16 +1029,19 @@ func makeResult(matchedRules []rules.Rule, reason Reason) (res Result) { } } -// InitModule manually initializes blocked services map. -func InitModule() { - initBlockedServices() +// InitModule manually initializes blocked services map. l must not be nil. +func InitModule(ctx context.Context, l *slog.Logger) { + initBlockedServices(ctx, l) } // New creates properly initialized DNS Filter that is ready to be used. c must // be non-nil. func New(c *Config, blockFilters []Filter) (d *DNSFilter, err error) { + ctx := context.TODO() + d = &DNSFilter{ - idGen: newIDGenerator(int32(time.Now().Unix())), + logger: c.Logger, + idGen: newIDGenerator(int32(time.Now().Unix()), c.Logger), bufPool: syncutil.NewSlicePool[byte](rulelist.DefaultRuleBufSize), safeSearch: c.SafeSearch, refreshLock: &sync.Mutex{}, @@ -1036,7 +1068,7 @@ func New(c *Config, blockFilters []Filter) (d *DNSFilter, err error) { check: d.matchHost, name: "filtering", }, { - check: matchBlockedServicesRules, + check: d.matchBlockedServicesRules, name: "blocked services", }, { check: d.checkSafeBrowsing, @@ -1054,7 +1086,7 @@ func New(c *Config, blockFilters []Filter) (d *DNSFilter, err error) { d.conf = c d.conf.filtersMu = &sync.RWMutex{} - err = d.prepareRewrites() + err = d.prepareRewrites(ctx) if err != nil { return nil, fmt.Errorf("rewrites: preparing: %w", err) } @@ -1067,7 +1099,7 @@ func New(c *Config, blockFilters []Filter) (d *DNSFilter, err error) { } if blockFilters != nil { - err = d.initFiltering(nil, blockFilters) + err = d.initFiltering(ctx, nil, blockFilters) if err != nil { d.Close() @@ -1082,8 +1114,8 @@ func New(c *Config, blockFilters []Filter) (d *DNSFilter, err error) { return nil, fmt.Errorf("making filtering directory: %w", err) } - d.loadFilters(d.conf.Filters) - d.loadFilters(d.conf.WhitelistFilters) + d.loadFilters(ctx, d.conf.Filters) + d.loadFilters(ctx, d.conf.WhitelistFilters) d.conf.Filters = deduplicateFilters(d.conf.Filters) d.conf.WhitelistFilters = deduplicateFilters(d.conf.WhitelistFilters) @@ -1101,12 +1133,12 @@ func (d *DNSFilter) Start() { d.RegisterFilteringHandlers() - go d.updatesLoop() + go d.updatesLoop(context.TODO()) } // updatesLoop initializes new filters and checks for filters updates in a loop. -func (d *DNSFilter) updatesLoop() { - defer log.OnPanic("filtering: updates loop") +func (d *DNSFilter) updatesLoop(ctx context.Context) { + defer slogutil.RecoverAndLog(ctx, d.logger) ivl := time.Second * 5 t := time.NewTimer(ivl) @@ -1114,9 +1146,9 @@ func (d *DNSFilter) updatesLoop() { for { select { case params := <-d.filtersInitializerChan: - err := d.initFiltering(params.allowFilters, params.blockFilters) + err := d.initFiltering(ctx, params.allowFilters, params.blockFilters) if err != nil { - log.Error("filtering: initializing: %s", err) + d.logger.ErrorContext(ctx, "initializing", slogutil.KeyError, err) continue } @@ -1165,9 +1197,13 @@ func (d *DNSFilter) checkSafeBrowsing( return Result{}, nil } - if log.GetLevel() >= log.DEBUG { - timer := log.StartTimer() - defer timer.LogElapsed("filtering: safebrowsing lookup for %q", host) + ctx := context.TODO() + if d.logger.Enabled(ctx, slogutil.LevelDebug) { + startTime := time.Now() + defer func() { + elapsed := time.Since(startTime) + d.logger.DebugContext(ctx, "safebrowsing lookup", "host", host, "elapsed", elapsed) + }() } res = Result{ @@ -1197,9 +1233,13 @@ func (d *DNSFilter) checkParental( return Result{}, nil } - if log.GetLevel() >= log.DEBUG { - timer := log.StartTimer() - defer timer.LogElapsed("filtering: parental lookup for %q", host) + ctx := context.TODO() + if d.logger.Enabled(ctx, slogutil.LevelDebug) { + startTime := time.Now() + defer func() { + elapsed := time.Since(startTime) + d.logger.DebugContext(ctx, "parental lookup", "host", host, "elapsed", elapsed) + }() } res = Result{ diff --git a/internal/filtering/filtering_internal_test.go b/internal/filtering/filtering_internal_test.go index 27260bf0..5097755c 100644 --- a/internal/filtering/filtering_internal_test.go +++ b/internal/filtering/filtering_internal_test.go @@ -2,13 +2,14 @@ package filtering import ( "bytes" + "cmp" "fmt" "net/netip" "testing" "github.com/AdguardTeam/AdGuardHome/internal/aghtest" "github.com/AdguardTeam/AdGuardHome/internal/filtering/hashprefix" - "github.com/AdguardTeam/golibs/log" + "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/AdguardTeam/golibs/netutil" "github.com/AdguardTeam/golibs/testutil" "github.com/AdguardTeam/urlfilter/rules" @@ -17,15 +18,14 @@ import ( "github.com/stretchr/testify/require" ) -func TestMain(m *testing.M) { - testutil.DiscardLogOutput(m) -} - const ( sbBlocked = "wmconvirus.narod.ru" pcBlocked = "pornhub.com" ) +// testLogger is the common logger for tests. +var testLogger = slogutil.NewDiscardLogger() + // Helpers. func newForTest(t testing.TB, c *Config, filters []Filter) (f *DNSFilter, setts *Settings) { @@ -34,6 +34,7 @@ func newForTest(t testing.TB, c *Config, filters []Filter) (f *DNSFilter, setts FilteringEnabled: true, } if c != nil { + c.Logger = cmp.Or(c.Logger, testLogger) c.SafeBrowsingCacheSize = 10000 c.ParentalCacheSize = 10000 c.SafeSearchCacheSize = 1000 @@ -43,7 +44,9 @@ func newForTest(t testing.TB, c *Config, filters []Filter) (f *DNSFilter, setts setts.ParentalEnabled = c.ParentalEnabled } else { // It must not be nil. - c = &Config{} + c = &Config{ + Logger: testLogger, + } } f, err := New(c, filters) require.NoError(t, err) @@ -53,6 +56,7 @@ func newForTest(t testing.TB, c *Config, filters []Filter) (f *DNSFilter, setts func newChecker(host string) Checker { return hashprefix.New(&hashprefix.Config{ + Logger: testLogger, CacheTime: 10, CacheSize: 100000, Upstream: aghtest.NewBlockUpstream(host, true), @@ -168,12 +172,15 @@ func TestDNSFilter_CheckHost_hostRules(t *testing.T) { func TestSafeBrowsing(t *testing.T) { logOutput := &bytes.Buffer{} - aghtest.ReplaceLogWriter(t, logOutput) - aghtest.ReplaceLogLevel(t, log.DEBUG) - sbChecker := newChecker(sbBlocked) d, setts := newForTest(t, &Config{ + Logger: slogutil.New(&slogutil.Config{ + Level: slogutil.LevelDebug, + Output: logOutput, + Format: slogutil.FormatDefault, + AddTimestamp: false, + }), SafeBrowsingEnabled: true, SafeBrowsingChecker: sbChecker, }, nil) @@ -181,7 +188,7 @@ func TestSafeBrowsing(t *testing.T) { d.checkMatch(t, sbBlocked, setts) - require.Contains(t, logOutput.String(), fmt.Sprintf("safebrowsing lookup for %q", sbBlocked)) + require.Contains(t, logOutput.String(), fmt.Sprintf("safebrowsing lookup host=%s", sbBlocked)) d.checkMatch(t, "test."+sbBlocked, setts) d.checkMatchEmpty(t, "yandex.ru", setts) @@ -216,17 +223,21 @@ func TestParallelSB(t *testing.T) { func TestParentalControl(t *testing.T) { logOutput := &bytes.Buffer{} - aghtest.ReplaceLogWriter(t, logOutput) - aghtest.ReplaceLogLevel(t, log.DEBUG) d, setts := newForTest(t, &Config{ + Logger: slogutil.New(&slogutil.Config{ + Level: slogutil.LevelDebug, + Output: logOutput, + Format: slogutil.FormatDefault, + AddTimestamp: false, + }), ParentalEnabled: true, ParentalControlChecker: newChecker(pcBlocked), }, nil) t.Cleanup(d.Close) d.checkMatch(t, pcBlocked, setts) - require.Contains(t, logOutput.String(), fmt.Sprintf("parental lookup for %q", pcBlocked)) + require.Contains(t, logOutput.String(), fmt.Sprintf("parental lookup host=%s", pcBlocked)) d.checkMatch(t, "www."+pcBlocked, setts) d.checkMatchEmpty(t, "www.yandex.ru", setts) @@ -548,7 +559,8 @@ func TestWhitelist(t *testing.T) { }} d, setts := newForTest(t, nil, filters) - err := d.setFilters(filters, whiteFilters, false) + ctx := testutil.ContextWithTimeout(t, testTimeout) + err := d.setFilters(ctx, filters, whiteFilters, false) require.NoError(t, err) t.Cleanup(d.Close) @@ -663,6 +675,7 @@ func TestClientSettings(t *testing.T) { func BenchmarkSafeBrowsing(b *testing.B) { d, setts := newForTest(b, &Config{ + Logger: testLogger, SafeBrowsingEnabled: true, SafeBrowsingChecker: newChecker(sbBlocked), }, nil) @@ -689,6 +702,7 @@ func BenchmarkSafeBrowsing(b *testing.B) { func BenchmarkSafeBrowsing_parallel(b *testing.B) { d, setts := newForTest(b, &Config{ + Logger: testLogger, SafeBrowsingEnabled: true, SafeBrowsingChecker: newChecker(sbBlocked), }, nil) diff --git a/internal/filtering/hashprefix/cache.go b/internal/filtering/hashprefix/cache.go index 7db2ae22..99d6cc73 100644 --- a/internal/filtering/hashprefix/cache.go +++ b/internal/filtering/hashprefix/cache.go @@ -1,10 +1,9 @@ package hashprefix import ( + "context" "encoding/binary" "time" - - "github.com/AdguardTeam/golibs/log" ) // expirySize is the size of expiry in cacheItem. @@ -91,7 +90,7 @@ func (c *Checker) findInCache( } // storeInCache caches hashes. -func (c *Checker) storeInCache(hashesToRequest, respHashes []hostnameHash) { +func (c *Checker) storeInCache(ctx context.Context, hashesToRequest, respHashes []hostnameHash) { hashToStore := make(map[prefix][]hostnameHash) for _, hash := range respHashes { @@ -102,7 +101,7 @@ func (c *Checker) storeInCache(hashesToRequest, respHashes []hostnameHash) { } for pref, hash := range hashToStore { - c.setCache(pref, hash) + c.setCache(ctx, pref, hash) } for _, hash := range hashesToRequest { @@ -111,18 +110,18 @@ func (c *Checker) storeInCache(hashesToRequest, respHashes []hostnameHash) { var pref prefix copy(pref[:], hash[:]) - c.setCache(pref, nil) + c.setCache(ctx, pref, nil) } } } // setCache stores hash in cache. -func (c *Checker) setCache(pref prefix, hashes []hostnameHash) { +func (c *Checker) setCache(ctx context.Context, pref prefix, hashes []hostnameHash) { item := &cacheItem{ expiry: time.Now().Add(c.cacheTime), hashes: hashes, } c.cache.Set(pref[:], fromCacheItem(item)) - log.Debug("%s: stored in cache: %v", c.svc, pref) + c.logger.DebugContext(ctx, "stored in cache", "pref", pref) } diff --git a/internal/filtering/hashprefix/hashprefix.go b/internal/filtering/hashprefix/hashprefix.go index 55795f9a..0a85c417 100644 --- a/internal/filtering/hashprefix/hashprefix.go +++ b/internal/filtering/hashprefix/hashprefix.go @@ -2,16 +2,18 @@ package hashprefix import ( + "context" "crypto/sha256" "encoding/hex" "fmt" + "log/slog" "slices" "strings" "time" "github.com/AdguardTeam/dnsproxy/upstream" "github.com/AdguardTeam/golibs/cache" - "github.com/AdguardTeam/golibs/log" + "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/AdguardTeam/golibs/netutil" "github.com/AdguardTeam/golibs/stringutil" "github.com/miekg/dns" @@ -52,12 +54,12 @@ func findMatch(a, b []hostnameHash) (matched bool) { // Config is the configuration structure for safe browsing and parental // control. type Config struct { + // Logger is used for logging the check process. It must not be nil. + Logger *slog.Logger + // Upstream is the upstream DNS server. Upstream upstream.Upstream - // ServiceName is the name of the service. - ServiceName string - // TXTSuffix is the TXT suffix for DNS request. TXTSuffix string @@ -70,15 +72,15 @@ type Config struct { } type Checker struct { + // logger is used for logging the check process. + logger *slog.Logger + // upstream is the upstream DNS server. upstream upstream.Upstream // cache stores hostname hashes. cache cache.Cache - // svc is the name of the service. - svc string - // txtSuffix is the TXT suffix for DNS request. txtSuffix string @@ -89,12 +91,12 @@ type Checker struct { // New returns Checker. func New(conf *Config) (c *Checker) { return &Checker{ + logger: conf.Logger, upstream: conf.Upstream, cache: cache.New(cache.Config{ EnableLRU: true, MaxSize: conf.CacheSize, }), - svc: conf.ServiceName, txtSuffix: conf.TXTSuffix, cacheTime: conf.CacheTime, } @@ -102,18 +104,22 @@ func New(conf *Config) (c *Checker) { // Check returns true if request for the host should be blocked. func (c *Checker) Check(host string) (ok bool, err error) { + ctx := context.TODO() + hashes := hostnameToHashes(host) + l := c.logger.With("host", host) + found, blocked, hashesToRequest := c.findInCache(hashes) if found { - log.Debug("%s: found %q in cache, blocked: %t", c.svc, host, blocked) + l.DebugContext(ctx, "found in cache", "blocked", blocked) return blocked, nil } question := c.getQuestion(hashesToRequest) - log.Debug("%s: checking %s: %s", c.svc, host, question) + l.DebugContext(ctx, "checking", "question", question) req := (&dns.Msg{}).SetQuestion(question, dns.TypeTXT) resp, err := c.upstream.Exchange(req) @@ -121,9 +127,9 @@ func (c *Checker) Check(host string) (ok bool, err error) { return false, fmt.Errorf("getting hashes: %w", err) } - matched, receivedHashes := c.processAnswer(hashesToRequest, resp, host) + matched, receivedHashes := c.processAnswer(ctx, l, hashesToRequest, resp) - c.storeInCache(hashesToRequest, receivedHashes) + c.storeInCache(ctx, hashesToRequest, receivedHashes) return matched, nil } @@ -182,11 +188,12 @@ func (c *Checker) getQuestion(hashes []hostnameHash) (q string) { } // processAnswer returns true if DNS response matches the hash, and received -// hashed hostnames from the upstream. +// hashed hostnames from the upstream. l must not be nil. func (c *Checker) processAnswer( + ctx context.Context, + l *slog.Logger, hashesToRequest []hostnameHash, resp *dns.Msg, - host string, ) (matched bool, receivedHashes []hostnameHash) { txtCount := 0 @@ -198,14 +205,14 @@ func (c *Checker) processAnswer( txtCount++ - receivedHashes = c.appendHashesFromTXT(receivedHashes, txt, host) + receivedHashes = c.appendHashesFromTXT(ctx, l, receivedHashes, txt) } - log.Debug("%s: received answer for %s with %d TXT count", c.svc, host, txtCount) + l.DebugContext(ctx, "processing answer with TXT", "txt_count", txtCount) matched = findMatch(hashesToRequest, receivedHashes) if matched { - log.Debug("%s: matched %s", c.svc, host) + l.DebugContext(ctx, "matched") return true, receivedHashes } @@ -213,24 +220,25 @@ func (c *Checker) processAnswer( return false, receivedHashes } -// appendHashesFromTXT appends received hashed hostnames. +// appendHashesFromTXT appends received hashed hostnames. l must not be nil. func (c *Checker) appendHashesFromTXT( + ctx context.Context, + l *slog.Logger, hashes []hostnameHash, txt *dns.TXT, - host string, ) (receivedHashes []hostnameHash) { - log.Debug("%s: received hashes for %s: %v", c.svc, host, txt.Txt) + l.DebugContext(ctx, "received hashes", "txt", txt.Txt) for _, t := range txt.Txt { if len(t) != hexSize { - log.Debug("%s: wrong hex size %d for %s %s", c.svc, len(t), host, t) + l.DebugContext(ctx, "wrong hex size", "len", len(t), "txt", t) continue } buf, err := hex.DecodeString(t) if err != nil { - log.Debug("%s: decoding hex string %s: %s", c.svc, t, err) + l.DebugContext(ctx, "decoding hex string", "txt", t, slogutil.KeyError, err) continue } diff --git a/internal/filtering/hashprefix/hashprefix_internal_test.go b/internal/filtering/hashprefix/hashprefix_internal_test.go index a575d0dd..f1ab3b86 100644 --- a/internal/filtering/hashprefix/hashprefix_internal_test.go +++ b/internal/filtering/hashprefix/hashprefix_internal_test.go @@ -10,6 +10,8 @@ import ( "github.com/AdguardTeam/AdGuardHome/internal/aghtest" "github.com/AdguardTeam/golibs/cache" + "github.com/AdguardTeam/golibs/logutil/slogutil" + "github.com/AdguardTeam/golibs/testutil" "github.com/miekg/dns" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -42,10 +44,10 @@ func TestChcker_getQuestion(t *testing.T) { hash = sha256.Sum256([]byte("com")) assert.False(t, slices.Contains(hashes, hash)) - c := &Checker{ - svc: "SafeBrowsing", - txtSuffix: suf, - } + c := New(&Config{ + Logger: slogutil.NewDiscardLogger(), + TXTSuffix: suf, + }) q := c.getQuestion(hashes) @@ -95,10 +97,13 @@ func TestHostnameToHashes(t *testing.T) { } func TestChecker_storeInCache(t *testing.T) { - c := &Checker{ - svc: "SafeBrowsing", - cacheTime: cacheTime, - } + const testTimeout = 1 * time.Second + + c := New(&Config{ + Logger: slogutil.NewDiscardLogger(), + CacheTime: cacheTime, + }) + conf := cache.Config{} c.cache = cache.New(conf) @@ -112,7 +117,7 @@ func TestChecker_storeInCache(t *testing.T) { hashesArray = append(hashesArray, hash4) hash2 := sha256.Sum256([]byte("host.com")) hashesArray = append(hashesArray, hash2) - c.storeInCache(hashes, hashesArray) + c.storeInCache(testutil.ContextWithTimeout(t, testTimeout), hashes, hashesArray) // match "3.sub.host.com" or "host.com" from cache hashes = []hostnameHash{} @@ -152,10 +157,11 @@ func TestChecker_storeInCache(t *testing.T) { ok = slices.Contains(hashesToRequest, hash) assert.True(t, ok) - c = &Checker{ - svc: "SafeBrowsing", - cacheTime: cacheTime, - } + c = New(&Config{ + Logger: slogutil.NewDiscardLogger(), + CacheTime: cacheTime, + }) + c.cache = cache.New(cache.Config{}) hashes = []hostnameHash{} @@ -189,6 +195,7 @@ func TestChecker_Check(t *testing.T) { for _, tc := range testCases { c := New(&Config{ + Logger: slogutil.NewDiscardLogger(), CacheTime: cacheTime, CacheSize: cacheSize, }) diff --git a/internal/filtering/hosts.go b/internal/filtering/hosts.go index 4943b1af..ba5b3899 100644 --- a/internal/filtering/hosts.go +++ b/internal/filtering/hosts.go @@ -1,12 +1,13 @@ package filtering import ( + "context" "fmt" "net/netip" "github.com/AdguardTeam/AdGuardHome/internal/filtering/rulelist" "github.com/AdguardTeam/golibs/hostsfile" - "github.com/AdguardTeam/golibs/log" + "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/AdguardTeam/golibs/netutil" "github.com/AdguardTeam/urlfilter/rules" "github.com/miekg/dns" @@ -24,7 +25,7 @@ func (d *DNSFilter) matchSysHosts( return Result{}, nil } - vals, rs, matched := hostsRewrites(qtype, host, d.conf.EtcHosts) + vals, rs, matched := d.hostsRewrites(qtype, host, d.conf.EtcHosts) if !matched { return Result{}, nil } @@ -42,11 +43,13 @@ func (d *DNSFilter) matchSysHosts( } // hostsRewrites returns values and rules matched by qt and host within hs. -func hostsRewrites( +func (d *DNSFilter) hostsRewrites( qtype uint16, host string, hs hostsfile.Storage, ) (vals []rules.RRValue, rls []*ResultRule, matched bool) { + ctx := context.TODO() + var isValidProto func(netip.Addr) (ok bool) switch qtype { case dns.TypeA: @@ -56,7 +59,12 @@ func hostsRewrites( case dns.TypePTR: addr, err := netutil.IPFromReversedAddr(host) if err != nil { - log.Debug("filtering: failed to parse PTR record %q: %s", host, err) + d.logger.DebugContext( + ctx, + "failed to parse PTR record", + "host", host, + slogutil.KeyError, err, + ) return nil, nil, false } @@ -73,7 +81,11 @@ func hostsRewrites( return vals, rls, len(names) > 0 default: - log.Debug("filtering: unsupported qtype %d", qtype) + d.logger.DebugContext( + ctx, + "unsupported qtype", + "qtype", qtype, + ) return nil, nil, false } diff --git a/internal/filtering/hosts_test.go b/internal/filtering/hosts_test.go index 14e20adc..5c692814 100644 --- a/internal/filtering/hosts_test.go +++ b/internal/filtering/hosts_test.go @@ -10,6 +10,7 @@ import ( "github.com/AdguardTeam/AdGuardHome/internal/aghtest" "github.com/AdguardTeam/AdGuardHome/internal/filtering" "github.com/AdguardTeam/AdGuardHome/internal/filtering/rulelist" + "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/AdguardTeam/golibs/testutil" "github.com/AdguardTeam/urlfilter/rules" "github.com/miekg/dns" @@ -52,6 +53,7 @@ func TestDNSFilter_CheckHost_hostsContainer(t *testing.T) { testutil.CleanupAndRequireSuccess(t, hc.Close) conf := &filtering.Config{ + Logger: slogutil.NewDiscardLogger(), EtcHosts: hc, } f, err := filtering.New(conf, nil) diff --git a/internal/filtering/http.go b/internal/filtering/http.go index 99acdb16..dca039a0 100644 --- a/internal/filtering/http.go +++ b/internal/filtering/http.go @@ -17,7 +17,7 @@ import ( "github.com/AdguardTeam/AdGuardHome/internal/aghhttp" "github.com/AdguardTeam/AdGuardHome/internal/filtering/rulelist" "github.com/AdguardTeam/golibs/errors" - "github.com/AdguardTeam/golibs/log" + "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/AdguardTeam/golibs/netutil/urlutil" "github.com/miekg/dns" ) @@ -148,6 +148,8 @@ func (d *DNSFilter) handleFilteringRemoveURL(w http.ResponseWriter, r *http.Requ Whitelist bool `json:"whitelist"` } + ctx := r.Context() + req := request{} err := json.NewDecoder(r.Body).Decode(&req) if err != nil { @@ -170,7 +172,12 @@ func (d *DNSFilter) handleFilteringRemoveURL(w http.ResponseWriter, r *http.Requ return flt.URL == req.URL }) if delIdx == -1 { - log.Error("deleting filter with url %q: %s", req.URL, errFilterNotExist) + d.logger.ErrorContext( + ctx, + "deleting filter", + "url", req.URL, + slogutil.KeyError, errFilterNotExist, + ) return } @@ -179,14 +186,20 @@ func (d *DNSFilter) handleFilteringRemoveURL(w http.ResponseWriter, r *http.Requ p := deleted.Path(d.conf.DataDir) err = os.Rename(p, p+".old") if err != nil && !errors.Is(err, os.ErrNotExist) { - log.Error("deleting filter %d: renaming file %q: %s", deleted.ID, p, err) + d.logger.ErrorContext( + ctx, + "renaming filter file", + "id", deleted.ID, + "path", p, + slogutil.KeyError, err, + ) return } *filters = slices.Delete(*filters, delIdx, delIdx+1) - log.Info("deleted filter %d", deleted.ID) + d.logger.InfoContext(ctx, "deleted filter", "id", deleted.ID) }() d.conf.ConfigModified() diff --git a/internal/filtering/http_internal_test.go b/internal/filtering/http_internal_test.go index a46d5d7b..4d45e254 100644 --- a/internal/filtering/http_internal_test.go +++ b/internal/filtering/http_internal_test.go @@ -12,6 +12,7 @@ import ( "time" "github.com/AdguardTeam/AdGuardHome/internal/schedule" + "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/AdguardTeam/golibs/testutil" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -103,6 +104,7 @@ func TestDNSFilter_handleFilteringSetURL(t *testing.T) { t.Run(tc.name, func(t *testing.T) { confModifiedCalled := false d, err := New(&Config{ + Logger: slogutil.NewDiscardLogger(), FilteringEnabled: true, Filters: tc.initial, HTTPClient: &http.Client{ @@ -183,6 +185,7 @@ func TestDNSFilter_handleSafeBrowsingStatus(t *testing.T) { handlers := make(map[string]http.Handler) d, err := New(&Config{ + Logger: slogutil.NewDiscardLogger(), ConfigModified: func() { testutil.RequireSend(testutil.PanicT{}, confModCh, struct{}{}, testTimeout) }, @@ -267,6 +270,7 @@ func TestDNSFilter_handleParentalStatus(t *testing.T) { handlers := make(map[string]http.Handler) d, err := New(&Config{ + Logger: slogutil.NewDiscardLogger(), ConfigModified: func() { testutil.RequireSend(testutil.PanicT{}, confModCh, struct{}{}, testTimeout) }, @@ -370,6 +374,7 @@ func TestDNSFilter_HandleCheckHost(t *testing.T) { } dnsFilter, err := New(&Config{ + Logger: slogutil.NewDiscardLogger(), BlockedServices: &BlockedServices{ Schedule: schedule.EmptyWeekly(), }, diff --git a/internal/filtering/idgenerator.go b/internal/filtering/idgenerator.go index e50f86ee..7d07cf23 100644 --- a/internal/filtering/idgenerator.go +++ b/internal/filtering/idgenerator.go @@ -1,12 +1,13 @@ package filtering import ( + "context" "fmt" + "log/slog" "sync/atomic" "github.com/AdguardTeam/AdGuardHome/internal/filtering/rulelist" "github.com/AdguardTeam/golibs/container" - "github.com/AdguardTeam/golibs/log" ) // idGenerator generates filtering-list IDs in a way broadly compatible with the @@ -16,13 +17,15 @@ import ( // rule-list architecture. type idGenerator struct { current *atomic.Int32 + logger *slog.Logger } // newIDGenerator returns a new ID generator initialized with the given seed // value. -func newIDGenerator(seed int32) (g *idGenerator) { +func newIDGenerator(seed int32, l *slog.Logger) (g *idGenerator) { g = &idGenerator{ current: &atomic.Int32{}, + logger: l, } g.current.Store(seed) @@ -61,11 +64,12 @@ func (g *idGenerator) fix(flts []FilterYAML) { newID = g.next() } - log.Info( - "filtering: warning: filter at index %d has duplicate id %d; reassigning to %d", - i, - id, - newID, + g.logger.WarnContext( + context.TODO(), + "filter has duplicate id; reassigning", + "idx", i, + "id", id, + "new_id", newID, ) flts[i].ID = newID diff --git a/internal/filtering/idgenerator_internal_test.go b/internal/filtering/idgenerator_internal_test.go index 57af4ad1..e9c0db2f 100644 --- a/internal/filtering/idgenerator_internal_test.go +++ b/internal/filtering/idgenerator_internal_test.go @@ -5,6 +5,7 @@ import ( "github.com/AdguardTeam/AdGuardHome/internal/aghalg" "github.com/AdguardTeam/AdGuardHome/internal/filtering/rulelist" + "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/stretchr/testify/assert" ) @@ -64,7 +65,7 @@ func TestIDGenerator_Fix(t *testing.T) { for _, tc := range testCases { t.Run(tc.name, func(t *testing.T) { - g := newIDGenerator(1) + g := newIDGenerator(1, slogutil.NewDiscardLogger()) g.fix(tc.in) assertUniqueIDs(t, tc.in) diff --git a/internal/filtering/rewrite/storage.go b/internal/filtering/rewrite/storage.go index 42b36273..bab84089 100644 --- a/internal/filtering/rewrite/storage.go +++ b/internal/filtering/rewrite/storage.go @@ -2,13 +2,14 @@ package rewrite import ( + "context" "fmt" + "log/slog" "slices" "strings" "sync" "github.com/AdguardTeam/golibs/container" - "github.com/AdguardTeam/golibs/log" "github.com/AdguardTeam/urlfilter" "github.com/AdguardTeam/urlfilter/filterlist" "github.com/AdguardTeam/urlfilter/rules" @@ -30,8 +31,23 @@ type Storage interface { List() (items []*Item) } +// Config is the configuration for DefaultStorage. +type Config struct { + // logger is used for logging storage processes. It must not be nil. + Logger *slog.Logger + + // Rewrites stores the rewrite entries. It must not be nil. + Rewrites []*Item + + // ListID is used as an identifier of the underlying rules list. + ListID int +} + // DefaultStorage is the default storage for rewrite rules. type DefaultStorage struct { + // logger is used for logging storage processes. It must not be nil. + logger *slog.Logger + // mu protects items. mu *sync.RWMutex @@ -51,13 +67,13 @@ type DefaultStorage struct { urlFilterID int } -// NewDefaultStorage returns new rewrites storage. listID is used as an -// identifier of the underlying rules list. rewrites must not be nil. -func NewDefaultStorage(listID int, rewrites []*Item) (s *DefaultStorage, err error) { +// NewDefaultStorage returns new rewrites storage. conf must not be nil. +func NewDefaultStorage(conf *Config) (s *DefaultStorage, err error) { s = &DefaultStorage{ + logger: conf.Logger, mu: &sync.RWMutex{}, - urlFilterID: listID, - rewrites: rewrites, + urlFilterID: conf.ListID, + rewrites: conf.Rewrites, } s.mu.Lock() @@ -79,6 +95,8 @@ func (s *DefaultStorage) MatchRequest(dReq *urlfilter.DNSRequest) (rws []*rules. s.mu.RLock() defer s.mu.RUnlock() + ctx := context.TODO() + rrules := s.rewriteRulesForReq(dReq) if len(rrules) == 0 { return nil @@ -91,7 +109,7 @@ func (s *DefaultStorage) MatchRequest(dReq *urlfilter.DNSRequest) (rws []*rules. rule := rrules[0] rwAns := rule.DNSRewrite.NewCNAME - log.Debug("rewrite: cname for %s is %s", host, rwAns) + s.logger.DebugContext(ctx, "cname found", "host", host, "cname", rwAns) if dReq.Hostname == rwAns { // A request for the hostname itself is an exception rule. @@ -109,7 +127,7 @@ func (s *DefaultStorage) MatchRequest(dReq *urlfilter.DNSRequest) (rws []*rules. } if cnames.Has(rwAns) { - log.Info("rewrite: cname loop for %q on %q", dReq.Hostname, rwAns) + s.logger.InfoContext(ctx, "rewrite cname loop", "host", dReq.Hostname, "rewrite", rwAns) return nil } @@ -168,12 +186,14 @@ func (s *DefaultStorage) Remove(item *Item) (err error) { s.mu.Lock() defer s.mu.Unlock() + ctx := context.TODO() + arr := []*Item{} // TODO(d.kolyshev): Use slices.IndexFunc + slices.Delete? for _, ent := range s.rewrites { if ent.equal(item) { - log.Debug("rewrite: removed element: %s -> %s", ent.Domain, ent.Answer) + s.logger.DebugContext(ctx, "removed element", "domain", ent.Domain, "ans", ent.Answer) continue } @@ -215,7 +235,12 @@ func (s *DefaultStorage) resetRules() (err error) { s.ruleList = strList s.engine = urlfilter.NewDNSEngine(rs) - log.Info("rewrite: filter %d: reset %d rules", s.urlFilterID, s.engine.RulesCount) + s.logger.InfoContext( + context.TODO(), + "reset rules", + "filter", s.urlFilterID, + "count", s.engine.RulesCount, + ) return nil } diff --git a/internal/filtering/rewrite/storage_internal_test.go b/internal/filtering/rewrite/storage_internal_test.go index 502c20b9..10df670c 100644 --- a/internal/filtering/rewrite/storage_internal_test.go +++ b/internal/filtering/rewrite/storage_internal_test.go @@ -4,6 +4,7 @@ import ( "net/netip" "testing" + "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/AdguardTeam/golibs/netutil" "github.com/AdguardTeam/urlfilter" "github.com/AdguardTeam/urlfilter/rules" @@ -18,7 +19,11 @@ func TestNewDefaultStorage(t *testing.T) { Answer: "answer.com", }} - s, err := NewDefaultStorage(-1, items) + s, err := NewDefaultStorage(&Config{ + Logger: slogutil.NewDiscardLogger(), + Rewrites: items, + ListID: -1, + }) require.NoError(t, err) require.Len(t, s.List(), 1) @@ -27,7 +32,11 @@ func TestNewDefaultStorage(t *testing.T) { func TestDefaultStorage_CRUD(t *testing.T) { var items []*Item - s, err := NewDefaultStorage(-1, items) + s, err := NewDefaultStorage(&Config{ + Logger: slogutil.NewDiscardLogger(), + Rewrites: items, + ListID: -1, + }) require.NoError(t, err) require.Len(t, s.List(), 0) @@ -112,7 +121,11 @@ func TestDefaultStorage_MatchRequest(t *testing.T) { Answer: "sub.issue4016.com", }} - s, err := NewDefaultStorage(-1, items) + s, err := NewDefaultStorage(&Config{ + Logger: slogutil.NewDiscardLogger(), + Rewrites: items, + ListID: -1, + }) require.NoError(t, err) testCases := []struct { @@ -284,7 +297,11 @@ func TestDefaultStorage_MatchRequest_Levels(t *testing.T) { Answer: addr3.String(), }} - s, err := NewDefaultStorage(-1, items) + s, err := NewDefaultStorage(&Config{ + Logger: slogutil.NewDiscardLogger(), + Rewrites: items, + ListID: -1, + }) require.NoError(t, err) testCases := []struct { @@ -352,7 +369,11 @@ func TestDefaultStorage_MatchRequest_ExceptionCNAME(t *testing.T) { Answer: "*.sub.host.com", }} - s, err := NewDefaultStorage(-1, items) + s, err := NewDefaultStorage(&Config{ + Logger: slogutil.NewDiscardLogger(), + Rewrites: items, + ListID: -1, + }) require.NoError(t, err) testCases := []struct { @@ -416,7 +437,11 @@ func TestDefaultStorage_MatchRequest_ExceptionIP(t *testing.T) { Answer: "A", }} - s, err := NewDefaultStorage(-1, items) + s, err := NewDefaultStorage(&Config{ + Logger: slogutil.NewDiscardLogger(), + Rewrites: items, + ListID: -1, + }) require.NoError(t, err) testCases := []struct { diff --git a/internal/filtering/rewritehttp.go b/internal/filtering/rewritehttp.go index af2ddf1f..d6415a05 100644 --- a/internal/filtering/rewritehttp.go +++ b/internal/filtering/rewritehttp.go @@ -6,7 +6,6 @@ import ( "slices" "github.com/AdguardTeam/AdGuardHome/internal/aghhttp" - "github.com/AdguardTeam/golibs/log" ) // TODO(d.kolyshev): Use [rewrite.Item] instead. @@ -50,7 +49,7 @@ func (d *DNSFilter) handleRewriteAdd(w http.ResponseWriter, r *http.Request) { Answer: rwJSON.Answer, } - err = rw.normalize() + err = rw.normalize(r.Context(), d.logger) if err != nil { // Shouldn't happen currently, since normalize only returns a non-nil // error when a rewrite is nil, but be change-proof. @@ -64,11 +63,12 @@ func (d *DNSFilter) handleRewriteAdd(w http.ResponseWriter, r *http.Request) { defer d.confMu.Unlock() d.conf.Rewrites = append(d.conf.Rewrites, rw) - log.Debug( - "rewrite: added element: %s -> %s [%d]", - rw.Domain, - rw.Answer, - len(d.conf.Rewrites), + d.logger.DebugContext( + r.Context(), + "added rewrite element", + "domain", rw.Domain, + "answer", rw.Answer, + "rewrites_len", len(d.conf.Rewrites), ) }() @@ -98,7 +98,12 @@ func (d *DNSFilter) handleRewriteDelete(w http.ResponseWriter, r *http.Request) for _, ent := range d.conf.Rewrites { if ent.equal(entDel) { - log.Debug("rewrite: removed element: %s -> %s", ent.Domain, ent.Answer) + d.logger.DebugContext( + r.Context(), + "removed rewrite element", + "domain", ent.Domain, + "answer", ent.Answer, + ) continue } @@ -138,7 +143,7 @@ func (d *DNSFilter) handleRewriteUpdate(w http.ResponseWriter, r *http.Request) Answer: updateJSON.Update.Answer, } - err = rwAdd.normalize() + err = rwAdd.normalize(r.Context(), d.logger) if err != nil { // Shouldn't happen currently, since normalize only returns a non-nil // error when a rewrite is nil, but be change-proof. @@ -166,6 +171,17 @@ func (d *DNSFilter) handleRewriteUpdate(w http.ResponseWriter, r *http.Request) d.conf.Rewrites = slices.Replace(d.conf.Rewrites, index, index+1, rwAdd) - log.Debug("rewrite: removed element: %s -> %s", rwDel.Domain, rwDel.Answer) - log.Debug("rewrite: added element: %s -> %s", rwAdd.Domain, rwAdd.Answer) + ctx := r.Context() + d.logger.DebugContext( + ctx, + "removed rewrite element", + "domain", rwDel.Domain, + "answer", rwDel.Answer, + ) + d.logger.DebugContext( + ctx, + "added rewrite element", + "domain", rwAdd.Domain, + "answer", rwAdd.Answer, + ) } diff --git a/internal/filtering/rewritehttp_test.go b/internal/filtering/rewritehttp_test.go index 93eef85a..b95435b8 100644 --- a/internal/filtering/rewritehttp_test.go +++ b/internal/filtering/rewritehttp_test.go @@ -10,6 +10,7 @@ import ( "time" "github.com/AdguardTeam/AdGuardHome/internal/filtering" + "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/AdguardTeam/golibs/testutil" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -159,6 +160,7 @@ func TestDNSFilter_handleRewriteHTTP(t *testing.T) { handlers := make(map[string]http.Handler) d, err := filtering.New(&filtering.Config{ + Logger: slogutil.NewDiscardLogger(), ConfigModified: onConfModified, HTTPRegister: func(_, url string, handler http.HandlerFunc) { handlers[url] = handler diff --git a/internal/filtering/rewrites.go b/internal/filtering/rewrites.go index 5ac2ffcc..809859b8 100644 --- a/internal/filtering/rewrites.go +++ b/internal/filtering/rewrites.go @@ -1,13 +1,15 @@ package filtering import ( + "context" "fmt" + "log/slog" "net/netip" "slices" "strings" "github.com/AdguardTeam/golibs/errors" - "github.com/AdguardTeam/golibs/log" + "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/miekg/dns" ) @@ -58,7 +60,7 @@ func (rw *LegacyRewrite) matchesQType(qt uint16) (ok bool) { // to domain name case, IP length, and so on. // // If rw is nil, it returns an errors. -func (rw *LegacyRewrite) normalize() (err error) { +func (rw *LegacyRewrite) normalize(ctx context.Context, l *slog.Logger) (err error) { if rw == nil { return errors.Error("nil rewrite entry") } @@ -85,7 +87,7 @@ func (rw *LegacyRewrite) normalize() (err error) { ip, err := netip.ParseAddr(rw.Answer) if err != nil { - log.Debug("normalizing legacy rewrite: %s", err) + l.DebugContext(ctx, "normalizing legacy rewrite", slogutil.KeyError, err) rw.Type = dns.TypeCNAME return nil @@ -136,9 +138,9 @@ func (rw *LegacyRewrite) Compare(b *LegacyRewrite) (res int) { } // prepareRewrites normalizes and validates all legacy DNS rewrites. -func (d *DNSFilter) prepareRewrites() (err error) { +func (d *DNSFilter) prepareRewrites(ctx context.Context) (err error) { for i, r := range d.conf.Rewrites { - err = r.normalize() + err = r.normalize(ctx, d.logger) if err != nil { return fmt.Errorf("at index %d: %w", i, err) } @@ -191,7 +193,13 @@ func findRewrites( // setRewriteResult sets the Reason or IPList of res if necessary. res must not // be nil. -func setRewriteResult(res *Result, host string, rewrites []*LegacyRewrite, qtype uint16) { +func (d *DNSFilter) setRewriteResult( + ctx context.Context, + res *Result, + host string, + rewrites []*LegacyRewrite, + qtype uint16, +) { for _, rw := range rewrites { if rw.Type == qtype && (qtype == dns.TypeA || qtype == dns.TypeAAAA) { if rw.IP == (netip.Addr{}) { @@ -203,7 +211,7 @@ func setRewriteResult(res *Result, host string, rewrites []*LegacyRewrite, qtype res.IPList = append(res.IPList, rw.IP) - log.Debug("rewrite: a/aaaa for %s is %s", host, rw.IP) + d.logger.DebugContext(ctx, "set a/aaaa rewrite", "host", host, "ans", rw.IP) } } } diff --git a/internal/filtering/rewrites_internal_test.go b/internal/filtering/rewrites_internal_test.go index cdec8529..baa17a31 100644 --- a/internal/filtering/rewrites_internal_test.go +++ b/internal/filtering/rewrites_internal_test.go @@ -5,6 +5,7 @@ import ( "net/netip" "testing" + "github.com/AdguardTeam/golibs/testutil" "github.com/miekg/dns" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -88,7 +89,8 @@ func TestRewrites(t *testing.T) { Answer: addr1v4.String(), }} - require.NoError(t, d.prepareRewrites()) + ctx := testutil.ContextWithTimeout(t, testTimeout) + require.NoError(t, d.prepareRewrites(ctx)) testCases := []struct { name string @@ -236,7 +238,8 @@ func TestRewritesLevels(t *testing.T) { Type: dns.TypeA, }} - require.NoError(t, d.prepareRewrites()) + ctx := testutil.ContextWithTimeout(t, testTimeout) + require.NoError(t, d.prepareRewrites(ctx)) testCases := []struct { name string @@ -280,7 +283,8 @@ func TestRewritesExceptionCNAME(t *testing.T) { Answer: "*.sub.host.com", }} - require.NoError(t, d.prepareRewrites()) + ctx := testutil.ContextWithTimeout(t, testTimeout) + require.NoError(t, d.prepareRewrites(ctx)) testCases := []struct { name string @@ -342,7 +346,8 @@ func TestRewritesExceptionIP(t *testing.T) { Type: dns.TypeA, }} - require.NoError(t, d.prepareRewrites()) + ctx := testutil.ContextWithTimeout(t, testTimeout) + require.NoError(t, d.prepareRewrites(ctx)) testCases := []struct { name string diff --git a/internal/home/clients_internal_test.go b/internal/home/clients_internal_test.go index 899ead65..8bbfd185 100644 --- a/internal/home/clients_internal_test.go +++ b/internal/home/clients_internal_test.go @@ -26,7 +26,9 @@ func newClientsContainer(t *testing.T) (c *clientsContainer) { client.EmptyDHCP{}, nil, nil, - &filtering.Config{}, + &filtering.Config{ + Logger: testLogger, + }, newSignalHandler(nil, nil), ) diff --git a/internal/home/controlinstall.go b/internal/home/controlinstall.go index 6e52f80a..602b0f64 100644 --- a/internal/home/controlinstall.go +++ b/internal/home/controlinstall.go @@ -431,6 +431,7 @@ func (web *webAPI) handleInstallConfigure(w http.ResponseWriter, r *http.Request globalContext.firstRun = false config.DNS.BindHosts = []netip.Addr{req.DNS.IP} config.DNS.Port = req.DNS.Port + config.Filtering.Logger = web.baseLogger.With(slogutil.KeyPrefix, "filtering") config.Filtering.SafeFSPatterns = []string{ filepath.Join(globalContext.workDir, userFilterDataDir, "*"), } diff --git a/internal/home/home.go b/internal/home/home.go index 052892f8..d4b6322a 100644 --- a/internal/home/home.go +++ b/internal/home/home.go @@ -357,15 +357,17 @@ func setupDNSFilteringConf( const ( dnsTimeout = 3 * time.Second - sbService = "safe browsing" + sbService = "safe_browsing" defaultSafeBrowsingServer = `https://family.adguard-dns.com/dns-query` sbTXTSuffix = `sb.dns.adguard.com.` - pcService = "parental control" + pcService = "parental_control" defaultParentalServer = `https://family.adguard-dns.com/dns-query` pcTXTSuffix = `pc.dns.adguard.com.` ) + conf.Logger = baseLogger.With(slogutil.KeyPrefix, "filtering") + conf.EtcHosts = globalContext.etcHosts // TODO(s.chzhen): Use empty interface. if globalContext.etcHosts == nil || !config.DNS.HostsFileEnabled { @@ -402,11 +404,11 @@ func setupDNSFilteringConf( } conf.SafeBrowsingChecker = hashprefix.New(&hashprefix.Config{ - Upstream: sbUps, - ServiceName: sbService, - TXTSuffix: sbTXTSuffix, - CacheTime: cacheTime, - CacheSize: conf.SafeBrowsingCacheSize, + Logger: baseLogger.With(slogutil.KeyPrefix, sbService), + Upstream: sbUps, + TXTSuffix: sbTXTSuffix, + CacheTime: cacheTime, + CacheSize: conf.SafeBrowsingCacheSize, }) // Protect against invalid configuration, see #6181. @@ -415,7 +417,11 @@ func setupDNSFilteringConf( // default. if conf.SafeBrowsingBlockHost == "" { host := defaultSafeBrowsingBlockHost - log.Info("%s: warning: empty blocking host; using default: %q", sbService, host) + baseLogger.WarnContext(ctx, + "empty blocking host; set default", + "service", sbService, + "host", host, + ) conf.SafeBrowsingBlockHost = host } @@ -426,11 +432,11 @@ func setupDNSFilteringConf( } conf.ParentalControlChecker = hashprefix.New(&hashprefix.Config{ - Upstream: parUps, - ServiceName: pcService, - TXTSuffix: pcTXTSuffix, - CacheTime: cacheTime, - CacheSize: conf.ParentalCacheSize, + Logger: baseLogger.With(slogutil.KeyPrefix, pcService), + Upstream: parUps, + TXTSuffix: pcTXTSuffix, + CacheTime: cacheTime, + CacheSize: conf.ParentalCacheSize, }) // Protect against invalid configuration, see #6181. @@ -439,7 +445,11 @@ func setupDNSFilteringConf( // default. if conf.ParentalBlockHost == "" { host := defaultParentalBlockHost - log.Info("%s: warning: empty blocking host; using default: %q", pcService, host) + baseLogger.WarnContext(ctx, + "empty blocking host; set default", + "service", pcService, + "host", host, + ) conf.ParentalBlockHost = host } @@ -614,13 +624,13 @@ func run(opts options, clientBuildFS fs.FS, done chan struct{}, sigHdlr *signalH err = configureOS(config) fatalOnError(err) + // TODO(s.chzhen): Use it for the entire initialization process. + ctx := context.Background() + // Clients package uses filtering package's static data // (filtering.BlockedSvcKnown()), so we have to initialize filtering static // data first, but also to avoid relying on automatic Go init() function. - filtering.InitModule() - - // TODO(s.chzhen): Use it for the entire initialization process. - ctx := context.Background() + filtering.InitModule(ctx, slogLogger) err = initContextClients(ctx, slogLogger, sigHdlr) fatalOnError(err)