Pull request: AGDNS-2374-imp-filtering-slog

Merge in DNS/adguard-home from AGDNS-2374-imp-filtering-slog to master

Squashed commit of the following:

commit 6040411aa2eeea62305acd3d99a6ea88a3998f38
Merge: 880c024ce 3c05f7799
Author: Dimitry Kolyshev <dkolyshev@adguard.com>
Date:   Wed Jul 9 09:32:45 2025 +0400

    Merge remote-tracking branch 'origin/master' into AGDNS-2374-imp-filtering-slog

commit 880c024ce2ea2190fae4794c262eef3d4ec90687
Author: Dimitry Kolyshev <dkolyshev@adguard.com>
Date:   Tue Jul 1 16:54:08 2025 +0400

    all: imp code

commit de62c55be126c2214e6576a39b07405a6e451aa6
Author: Dimitry Kolyshev <dkolyshev@adguard.com>
Date:   Mon Jun 30 15:16:58 2025 +0400

    filtering: slog

commit 9bc1abfac1b9d183fe44e56d2b57b73d6507441a
Author: Dimitry Kolyshev <dkolyshev@adguard.com>
Date:   Mon Jun 30 15:13:24 2025 +0400

    filtering: slog

commit 88b251e974f78832b1fb4b149c2314471a931490
Author: Dimitry Kolyshev <dkolyshev@adguard.com>
Date:   Mon Jun 30 14:13:37 2025 +0400

    filtering: slog

commit 2a6939865d8398c075d1eae2f9315bb81238b9ff
Author: Dimitry Kolyshev <dkolyshev@adguard.com>
Date:   Mon Jun 30 12:29:48 2025 +0400

    all: init filtering logger

commit 94609444fb54b3a996806606ec7014aec8dcc65e
Author: Dimitry Kolyshev <dkolyshev@adguard.com>
Date:   Mon Jun 30 11:43:51 2025 +0400

    filtering: imp rewrite

commit 84544cce8a65ad4cf6af793b0aaf55a902f0109f
Author: Dimitry Kolyshev <dkolyshev@adguard.com>
Date:   Mon Jun 30 11:39:58 2025 +0400

    filtering: imp hashprefix

commit b88d8799c04c2a45ba43aa1ebd614347686a4148
Author: Dimitry Kolyshev <dkolyshev@adguard.com>
Date:   Fri Jun 27 13:07:09 2025 +0400

    filtering: imp hashprefix

commit 904a847684d6e6d3aa55dfc478d5b6c4dc8767db
Author: Dimitry Kolyshev <dkolyshev@adguard.com>
Date:   Fri Jun 27 12:57:14 2025 +0400

    filtering: rewrite slog

commit 2669388fea2ba4b2e0d470d780714429b43d246b
Author: Dimitry Kolyshev <dkolyshev@adguard.com>
Date:   Fri Jun 27 12:36:34 2025 +0400

    filtering: hashprefix slog
This commit is contained in:
Dimitry Kolyshev
2025-07-09 08:42:44 +03:00
parent 3c05f77991
commit 63c64b10e9
31 changed files with 521 additions and 297 deletions

View File

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

View File

@@ -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 (

View File

@@ -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 "!!!": ` +

View File

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

View File

@@ -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 {

View File

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

View File

@@ -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),

View File

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

View File

@@ -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,

View File

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

View File

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

View File

@@ -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{

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@@ -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()

View File

@@ -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(),
},

View File

@@ -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

View File

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

View File

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

View File

@@ -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 {

View File

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

View File

@@ -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

View File

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

View File

@@ -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

View File

@@ -26,7 +26,9 @@ func newClientsContainer(t *testing.T) (c *clientsContainer) {
client.EmptyDHCP{},
nil,
nil,
&filtering.Config{},
&filtering.Config{
Logger: testLogger,
},
newSignalHandler(nil, nil),
)

View File

@@ -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, "*"),
}

View File

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