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 ( import (
"crypto/sha256" "crypto/sha256"
"io"
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"net/netip" "net/netip"
@@ -12,7 +11,6 @@ import (
"time" "time"
"github.com/AdguardTeam/dnsproxy/proxy" "github.com/AdguardTeam/dnsproxy/proxy"
"github.com/AdguardTeam/golibs/log"
"github.com/AdguardTeam/golibs/netutil" "github.com/AdguardTeam/golibs/netutil"
"github.com/AdguardTeam/golibs/testutil" "github.com/AdguardTeam/golibs/testutil"
"github.com/miekg/dns" "github.com/miekg/dns"
@@ -27,33 +25,6 @@ const (
ReqFQDN = ReqHost + "." 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. // HostToIPs is a helper that generates one IPv4 and one IPv6 address from host.
func HostToIPs(host string) (ipv4, ipv6 netip.Addr) { func HostToIPs(host string) (ipv4, ipv6 netip.Addr) {
hash := sha256.Sum256([]byte(host)) hash := sha256.Sum256([]byte(host))

View File

@@ -8,7 +8,6 @@ import (
"testing" "testing"
"github.com/AdguardTeam/dnsproxy/proxy" "github.com/AdguardTeam/dnsproxy/proxy"
"github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/AdguardTeam/golibs/testutil" "github.com/AdguardTeam/golibs/testutil"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
) )
@@ -201,7 +200,7 @@ func TestServer_clientIDFromDNSContext(t *testing.T) {
srv := &Server{ srv := &Server{
conf: ServerConfig{TLSConf: tlsConf}, conf: ServerConfig{TLSConf: tlsConf},
baseLogger: slogutil.NewDiscardLogger(), baseLogger: testLogger,
} }
var ( var (

View File

@@ -63,6 +63,9 @@ const (
// TODO(a.garipov): Use more. // TODO(a.garipov): Use more.
var testClientAddrPort = netip.MustParseAddrPort("1.2.3.4:12345") var testClientAddrPort = netip.MustParseAddrPort("1.2.3.4:12345")
// testLogger is the common logger for tests.
var testLogger = slogutil.NewDiscardLogger()
// type check // type check
var _ ClientsContainer = (*clientsContainer)(nil) var _ ClientsContainer = (*clientsContainer)(nil)
@@ -129,6 +132,8 @@ func createTestServer(
) (s *Server) { ) (s *Server) {
t.Helper() t.Helper()
filterConf.Logger = cmp.Or(filterConf.Logger, testLogger)
rules := `||nxdomain.example.org rules := `||nxdomain.example.org
||NULL.example.org^ ||NULL.example.org^
127.0.0.1 host.example.org 127.0.0.1 host.example.org
@@ -159,7 +164,7 @@ func createTestServer(
DHCPServer: dhcp, DHCPServer: dhcp,
DNSFilter: f, DNSFilter: f,
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed), PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
Logger: slogutil.NewDiscardLogger(), Logger: testLogger,
}) })
require.NoError(t, err) require.NoError(t, err)
@@ -410,7 +415,7 @@ func TestServer_timeout(t *testing.T) {
s, err := NewServer(DNSCreateParams{ s, err := NewServer(DNSCreateParams{
DNSFilter: createTestDNSFilter(t), DNSFilter: createTestDNSFilter(t),
Logger: slogutil.NewDiscardLogger(), Logger: testLogger,
}) })
require.NoError(t, err) require.NoError(t, err)
@@ -423,7 +428,7 @@ func TestServer_timeout(t *testing.T) {
t.Run("default", func(t *testing.T) { t.Run("default", func(t *testing.T) {
s, err := NewServer(DNSCreateParams{ s, err := NewServer(DNSCreateParams{
DNSFilter: createTestDNSFilter(t), DNSFilter: createTestDNSFilter(t),
Logger: slogutil.NewDiscardLogger(), Logger: testLogger,
}) })
require.NoError(t, err) require.NoError(t, err)
@@ -456,7 +461,7 @@ func TestServer_Prepare_fallbacks(t *testing.T) {
} }
s, err := NewServer(DNSCreateParams{ s, err := NewServer(DNSCreateParams{
Logger: slogutil.NewDiscardLogger(), Logger: testLogger,
}) })
require.NoError(t, err) require.NoError(t, err)
@@ -584,6 +589,7 @@ func TestSafeSearch(t *testing.T) {
} }
filterConf := &filtering.Config{ filterConf := &filtering.Config{
Logger: testLogger,
BlockingMode: filtering.BlockingModeDefault, BlockingMode: filtering.BlockingModeDefault,
ProtectionEnabled: true, ProtectionEnabled: true,
SafeSearchConf: safeSearchConf, SafeSearchConf: safeSearchConf,
@@ -593,7 +599,7 @@ func TestSafeSearch(t *testing.T) {
ctx := testutil.ContextWithTimeout(t, testTimeout) ctx := testutil.ContextWithTimeout(t, testTimeout)
safeSearch, err := safesearch.NewDefault(ctx, &safesearch.DefaultConfig{ safeSearch, err := safesearch.NewDefault(ctx, &safesearch.DefaultConfig{
Logger: slogutil.NewDiscardLogger(), Logger: testLogger,
ServicesConfig: safeSearchConf, ServicesConfig: safeSearchConf,
CacheSize: filterConf.SafeSearchCacheSize, CacheSize: filterConf.SafeSearchCacheSize,
CacheTTL: time.Minute * time.Duration(filterConf.CacheTime), CacheTTL: time.Minute * time.Duration(filterConf.CacheTime),
@@ -1055,6 +1061,7 @@ func TestBlockedCustomIP(t *testing.T) {
}} }}
f, err := filtering.New(&filtering.Config{ f, err := filtering.New(&filtering.Config{
Logger: testLogger,
ProtectionEnabled: true, ProtectionEnabled: true,
ApplyClientFiltering: applyEmptyClientFiltering, ApplyClientFiltering: applyEmptyClientFiltering,
BlockedServices: emptyFilteringBlockedServices(), BlockedServices: emptyFilteringBlockedServices(),
@@ -1073,7 +1080,7 @@ func TestBlockedCustomIP(t *testing.T) {
DHCPServer: dhcp, DHCPServer: dhcp,
DNSFilter: f, DNSFilter: f,
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed), PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
Logger: slogutil.NewDiscardLogger(), Logger: testLogger,
}) })
require.NoError(t, err) require.NoError(t, err)
@@ -1173,6 +1180,7 @@ func TestBlockedBySafeBrowsing(t *testing.T) {
) )
sbChecker := hashprefix.New(&hashprefix.Config{ sbChecker := hashprefix.New(&hashprefix.Config{
Logger: testLogger,
CacheTime: cacheTime, CacheTime: cacheTime,
CacheSize: cacheSize, CacheSize: cacheSize,
Upstream: aghtest.NewBlockUpstream(hostname, true), Upstream: aghtest.NewBlockUpstream(hostname, true),
@@ -1216,6 +1224,7 @@ func TestBlockedBySafeBrowsing(t *testing.T) {
func TestRewrite(t *testing.T) { func TestRewrite(t *testing.T) {
c := &filtering.Config{ c := &filtering.Config{
Logger: testLogger,
ApplyClientFiltering: applyEmptyClientFiltering, ApplyClientFiltering: applyEmptyClientFiltering,
BlockedServices: emptyFilteringBlockedServices(), BlockedServices: emptyFilteringBlockedServices(),
BlockingMode: filtering.BlockingModeDefault, BlockingMode: filtering.BlockingModeDefault,
@@ -1247,7 +1256,7 @@ func TestRewrite(t *testing.T) {
DHCPServer: dhcp, DHCPServer: dhcp,
DNSFilter: f, DNSFilter: f,
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed), PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
Logger: slogutil.NewDiscardLogger(), Logger: testLogger,
}) })
require.NoError(t, err) require.NoError(t, err)
@@ -1365,6 +1374,7 @@ func TestPTRResponseFromDHCPLeases(t *testing.T) {
const localDomain = "lan" const localDomain = "lan"
flt, err := filtering.New(&filtering.Config{ flt, err := filtering.New(&filtering.Config{
Logger: testLogger,
ApplyClientFiltering: applyEmptyClientFiltering, ApplyClientFiltering: applyEmptyClientFiltering,
BlockedServices: emptyFilteringBlockedServices(), BlockedServices: emptyFilteringBlockedServices(),
BlockingMode: filtering.BlockingModeDefault, BlockingMode: filtering.BlockingModeDefault,
@@ -1381,7 +1391,7 @@ func TestPTRResponseFromDHCPLeases(t *testing.T) {
}, },
}, },
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed), PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
Logger: slogutil.NewDiscardLogger(), Logger: testLogger,
LocalDomain: localDomain, LocalDomain: localDomain,
}) })
require.NoError(t, err) require.NoError(t, err)
@@ -1457,6 +1467,7 @@ func TestPTRResponseFromHosts(t *testing.T) {
}) })
flt, err := filtering.New(&filtering.Config{ flt, err := filtering.New(&filtering.Config{
Logger: testLogger,
ApplyClientFiltering: applyEmptyClientFiltering, ApplyClientFiltering: applyEmptyClientFiltering,
BlockedServices: emptyFilteringBlockedServices(), BlockedServices: emptyFilteringBlockedServices(),
BlockingMode: filtering.BlockingModeDefault, BlockingMode: filtering.BlockingModeDefault,
@@ -1471,7 +1482,7 @@ func TestPTRResponseFromHosts(t *testing.T) {
DHCPServer: dhcp, DHCPServer: dhcp,
DNSFilter: flt, DNSFilter: flt,
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed), PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
Logger: slogutil.NewDiscardLogger(), Logger: testLogger,
}) })
require.NoError(t, err) require.NoError(t, err)
@@ -1527,27 +1538,27 @@ func TestNewServer(t *testing.T) {
}{{ }{{
name: "success", name: "success",
in: DNSCreateParams{ in: DNSCreateParams{
Logger: slogutil.NewDiscardLogger(), Logger: testLogger,
}, },
wantErrMsg: "", wantErrMsg: "",
}, { }, {
name: "success_local_tld", name: "success_local_tld",
in: DNSCreateParams{ in: DNSCreateParams{
Logger: slogutil.NewDiscardLogger(), Logger: testLogger,
LocalDomain: "mynet", LocalDomain: "mynet",
}, },
wantErrMsg: "", wantErrMsg: "",
}, { }, {
name: "success_local_domain", name: "success_local_domain",
in: DNSCreateParams{ in: DNSCreateParams{
Logger: slogutil.NewDiscardLogger(), Logger: testLogger,
LocalDomain: "my.local.net", LocalDomain: "my.local.net",
}, },
wantErrMsg: "", wantErrMsg: "",
}, { }, {
name: "bad_local_domain", name: "bad_local_domain",
in: DNSCreateParams{ in: DNSCreateParams{
Logger: slogutil.NewDiscardLogger(), Logger: testLogger,
LocalDomain: "!!!", LocalDomain: "!!!",
}, },
wantErrMsg: `local domain: bad domain name "!!!": ` + wantErrMsg: `local domain: bad domain name "!!!": ` +

View File

@@ -9,7 +9,6 @@ import (
"github.com/AdguardTeam/AdGuardHome/internal/filtering" "github.com/AdguardTeam/AdGuardHome/internal/filtering"
"github.com/AdguardTeam/dnsproxy/proxy" "github.com/AdguardTeam/dnsproxy/proxy"
"github.com/AdguardTeam/dnsproxy/upstream" "github.com/AdguardTeam/dnsproxy/upstream"
"github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/AdguardTeam/golibs/netutil" "github.com/AdguardTeam/golibs/netutil"
"github.com/miekg/dns" "github.com/miekg/dns"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
@@ -46,6 +45,7 @@ func TestHandleDNSRequest_handleDNSRequest(t *testing.T) {
}} }}
f, err := filtering.New(&filtering.Config{ f, err := filtering.New(&filtering.Config{
Logger: testLogger,
ProtectionEnabled: true, ProtectionEnabled: true,
ApplyClientFiltering: applyEmptyClientFiltering, ApplyClientFiltering: applyEmptyClientFiltering,
BlockedServices: emptyFilteringBlockedServices(), BlockedServices: emptyFilteringBlockedServices(),
@@ -62,7 +62,7 @@ func TestHandleDNSRequest_handleDNSRequest(t *testing.T) {
}, },
DNSFilter: f, DNSFilter: f,
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed), PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
Logger: slogutil.NewDiscardLogger(), Logger: testLogger,
}) })
require.NoError(t, err) require.NoError(t, err)
@@ -226,7 +226,9 @@ func TestHandleDNSRequest_filterDNSResponse(t *testing.T) {
ID: 0, Data: []byte(blockRules), 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) require.NoError(t, err)
f.SetEnabled(true) f.SetEnabled(true)
@@ -235,7 +237,7 @@ func TestHandleDNSRequest_filterDNSResponse(t *testing.T) {
DHCPServer: &testDHCP{}, DHCPServer: &testDHCP{},
DNSFilter: f, DNSFilter: f,
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed), PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
Logger: slogutil.NewDiscardLogger(), Logger: testLogger,
}) })
require.NoError(t, err) require.NoError(t, err)

View File

@@ -6,7 +6,6 @@ import (
"testing" "testing"
"github.com/AdguardTeam/dnsproxy/proxy" "github.com/AdguardTeam/dnsproxy/proxy"
"github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/miekg/dns" "github.com/miekg/dns"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
) )
@@ -61,7 +60,7 @@ func TestIpsetCtx_process(t *testing.T) {
} }
ictx := &ipsetHandler{ ictx := &ipsetHandler{
logger: slogutil.NewDiscardLogger(), logger: testLogger,
} }
rc := ictx.process(dctx) rc := ictx.process(dctx)
assert.Equal(t, resultCodeSuccess, rc) assert.Equal(t, resultCodeSuccess, rc)
@@ -83,7 +82,7 @@ func TestIpsetCtx_process(t *testing.T) {
m := &fakeIpsetMgr{} m := &fakeIpsetMgr{}
ictx := &ipsetHandler{ ictx := &ipsetHandler{
ipsetMgr: m, ipsetMgr: m,
logger: slogutil.NewDiscardLogger(), logger: testLogger,
} }
rc := ictx.process(dctx) rc := ictx.process(dctx)
@@ -108,7 +107,7 @@ func TestIpsetCtx_process(t *testing.T) {
m := &fakeIpsetMgr{} m := &fakeIpsetMgr{}
ictx := &ipsetHandler{ ictx := &ipsetHandler{
ipsetMgr: m, ipsetMgr: m,
logger: slogutil.NewDiscardLogger(), logger: testLogger,
} }
rc := ictx.process(dctx) rc := ictx.process(dctx)
@@ -132,7 +131,7 @@ func TestIpsetCtx_SkipIpsetProcessing(t *testing.T) {
m := &fakeIpsetMgr{} m := &fakeIpsetMgr{}
ictx := &ipsetHandler{ ictx := &ipsetHandler{
ipsetMgr: m, ipsetMgr: m,
logger: slogutil.NewDiscardLogger(), logger: testLogger,
} }
testCases := []struct { testCases := []struct {

View File

@@ -12,7 +12,6 @@ import (
"github.com/AdguardTeam/AdGuardHome/internal/filtering" "github.com/AdguardTeam/AdGuardHome/internal/filtering"
"github.com/AdguardTeam/dnsproxy/proxy" "github.com/AdguardTeam/dnsproxy/proxy"
"github.com/AdguardTeam/dnsproxy/upstream" "github.com/AdguardTeam/dnsproxy/upstream"
"github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/AdguardTeam/golibs/netutil" "github.com/AdguardTeam/golibs/netutil"
"github.com/AdguardTeam/golibs/testutil" "github.com/AdguardTeam/golibs/testutil"
"github.com/AdguardTeam/urlfilter/rules" "github.com/AdguardTeam/urlfilter/rules"
@@ -378,6 +377,7 @@ func createTestDNSFilter(t *testing.T) (f *filtering.DNSFilter) {
t.Helper() t.Helper()
f, err := filtering.New(&filtering.Config{ f, err := filtering.New(&filtering.Config{
Logger: testLogger,
BlockingMode: filtering.BlockingModeDefault, BlockingMode: filtering.BlockingModeDefault,
}, []filtering.Filter{}) }, []filtering.Filter{})
require.NoError(t, err) require.NoError(t, err)
@@ -439,7 +439,7 @@ func TestServer_ProcessDHCPHosts_localRestriction(t *testing.T) {
dnsFilter: createTestDNSFilter(t), dnsFilter: createTestDNSFilter(t),
dhcpServer: dhcp, dhcpServer: dhcp,
localDomainSuffix: localDomainSuffix, localDomainSuffix: localDomainSuffix,
baseLogger: slogutil.NewDiscardLogger(), baseLogger: testLogger,
} }
req := &dns.Msg{ req := &dns.Msg{
@@ -591,7 +591,7 @@ func TestServer_ProcessDHCPHosts(t *testing.T) {
dnsFilter: createTestDNSFilter(t), dnsFilter: createTestDNSFilter(t),
dhcpServer: testDHCP, dhcpServer: testDHCP,
localDomainSuffix: tc.suffix, localDomainSuffix: tc.suffix,
baseLogger: slogutil.NewDiscardLogger(), baseLogger: testLogger,
} }
req := (&dns.Msg{}).SetQuestion(dns.Fqdn(tc.host), tc.qtyp) 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/AdGuardHome/internal/stats"
"github.com/AdguardTeam/dnsproxy/proxy" "github.com/AdguardTeam/dnsproxy/proxy"
"github.com/AdguardTeam/dnsproxy/upstream" "github.com/AdguardTeam/dnsproxy/upstream"
"github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/miekg/dns" "github.com/miekg/dns"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
@@ -203,7 +202,7 @@ func TestServer_ProcessQueryLogsAndStats(t *testing.T) {
ql := &testQueryLog{} ql := &testQueryLog{}
st := &testStats{} st := &testStats{}
srv := &Server{ srv := &Server{
baseLogger: slogutil.NewDiscardLogger(), baseLogger: testLogger,
queryLog: ql, queryLog: ql,
stats: st, stats: st,
anonymizer: aghnet.NewIPMut(nil), anonymizer: aghnet.NewIPMut(nil),

View File

@@ -1,8 +1,10 @@
package filtering package filtering
import ( import (
"context"
"encoding/json" "encoding/json"
"fmt" "fmt"
"log/slog"
"net/http" "net/http"
"slices" "slices"
"time" "time"
@@ -10,7 +12,7 @@ import (
"github.com/AdguardTeam/AdGuardHome/internal/aghhttp" "github.com/AdguardTeam/AdGuardHome/internal/aghhttp"
"github.com/AdguardTeam/AdGuardHome/internal/filtering/rulelist" "github.com/AdguardTeam/AdGuardHome/internal/filtering/rulelist"
"github.com/AdguardTeam/AdGuardHome/internal/schedule" "github.com/AdguardTeam/AdGuardHome/internal/schedule"
"github.com/AdguardTeam/golibs/log" "github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/AdguardTeam/urlfilter/rules" "github.com/AdguardTeam/urlfilter/rules"
) )
@@ -20,23 +22,30 @@ var serviceRules map[string][]*rules.NetworkRule
// serviceIDs contains service IDs sorted alphabetically. // serviceIDs contains service IDs sorted alphabetically.
var serviceIDs []string var serviceIDs []string
// initBlockedServices initializes package-level blocked service data. // initBlockedServices initializes package-level blocked service data. l must
func initBlockedServices() { // not be nil.
l := len(blockedServices) func initBlockedServices(ctx context.Context, l *slog.Logger) {
serviceIDs = make([]string, l) svcLen := len(blockedServices)
serviceRules = make(map[string][]*rules.NetworkRule, l) serviceIDs = make([]string, svcLen)
serviceRules = make(map[string][]*rules.NetworkRule, svcLen)
for i, s := range blockedServices { for i, s := range blockedServices {
netRules := make([]*rules.NetworkRule, 0, len(s.Rules)) netRules := make([]*rules.NetworkRule, 0, len(s.Rules))
for _, text := range s.Rules { for _, text := range s.Rules {
rule, err := rules.NewNetworkRule(text, rulelist.URLFilterIDBlockedService) rule, err := rules.NewNetworkRule(text, rulelist.URLFilterIDBlockedService)
if err != nil { if err == nil {
log.Error("parsing blocked service %q rule %q: %s", s.ID, text, err) netRules = append(netRules, rule)
continue continue
} }
netRules = append(netRules, rule) l.ErrorContext(
ctx,
"parsing blocked service rule",
"svc", s.ID,
"rule", text,
slogutil.KeyError, err,
)
} }
serviceIDs[i] = s.ID serviceIDs[i] = s.ID
@@ -45,7 +54,7 @@ func initBlockedServices() {
slices.Sort(serviceIDs) 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. // BlockedServices is the configuration of blocked services.
@@ -105,7 +114,7 @@ func (d *DNSFilter) ApplyBlockedServicesList(setts *Settings, list []string) {
for _, name := range list { for _, name := range list {
rules, ok := serviceRules[name] rules, ok := serviceRules[name]
if !ok { if !ok {
log.Error("unknown service name: %s", name) d.logger.ErrorContext(context.TODO(), "unknown service name", "name", name)
continue continue
} }
@@ -163,7 +172,7 @@ func (d *DNSFilter) handleBlockedServicesSet(w http.ResponseWriter, r *http.Requ
defer d.confMu.Unlock() defer d.confMu.Unlock()
d.conf.BlockedServices.IDs = list 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() d.conf.ConfigModified()
@@ -212,7 +221,7 @@ func (d *DNSFilter) handleBlockedServicesUpdate(w http.ResponseWriter, r *http.R
d.conf.BlockedServices = bsvc 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() d.conf.ConfigModified()
} }

View File

@@ -6,6 +6,7 @@ import (
"testing" "testing"
"github.com/AdguardTeam/AdGuardHome/internal/filtering" "github.com/AdguardTeam/AdGuardHome/internal/filtering"
"github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/AdguardTeam/golibs/netutil" "github.com/AdguardTeam/golibs/netutil"
"github.com/miekg/dns" "github.com/miekg/dns"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
@@ -52,6 +53,7 @@ func TestDNSFilter_CheckHostRules_dnsrewrite(t *testing.T) {
` `
conf := &filtering.Config{ conf := &filtering.Config{
Logger: slogutil.NewDiscardLogger(),
SafeBrowsingCacheSize: 10000, SafeBrowsingCacheSize: 10000,
ParentalCacheSize: 10000, ParentalCacheSize: 10000,
SafeSearchCacheSize: 1000, SafeSearchCacheSize: 1000,

View File

@@ -1,6 +1,7 @@
package filtering package filtering
import ( import (
"context"
"fmt" "fmt"
"io" "io"
"net/http" "net/http"
@@ -17,7 +18,7 @@ import (
"github.com/AdguardTeam/AdGuardHome/internal/filtering/rulelist" "github.com/AdguardTeam/AdGuardHome/internal/filtering/rulelist"
"github.com/AdguardTeam/golibs/container" "github.com/AdguardTeam/golibs/container"
"github.com/AdguardTeam/golibs/errors" "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 // filterDir is the subdirectory of a data directory to store downloaded
@@ -105,12 +106,13 @@ func (d *DNSFilter) filterSetProperties(
} }
flt := &filters[i] flt := &filters[i]
log.Debug( d.logger.DebugContext(
"filtering: set name to %q, url to %s, enabled to %t for filter %s", context.TODO(),
newList.Name, "updating filter",
newList.URL, "name", newList.Name,
newList.Enabled, "url", newList.URL,
flt.URL, "enabled", newList.Enabled,
"filter_url", flt.URL,
) )
defer func(oldURL, oldName string, oldEnabled bool, oldUpdated time.Time, oldRulesCount int) { 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 // Load filters from the disk
// And if any filter has zero ID, assign a new one // 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 { for i := range array {
filter := &array[i] // otherwise we're operating on a copy filter := &array[i] // otherwise we're operating on a copy
if filter.ID == 0 { if filter.ID == 0 {
newID := d.idGen.next() 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 filter.ID = newID
} }
@@ -228,9 +230,9 @@ func (d *DNSFilter) loadFilters(array []FilterYAML) {
continue continue
} }
err := d.load(filter) err := d.load(ctx, filter)
if err != nil { 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 return toUpd
} }
func (d *DNSFilter) refreshFiltersArray(filters *[]FilterYAML, force bool) (int, []FilterYAML, []bool, bool) { // refreshFiltersArray updates the filters array and returns the number of
var updateFlags []bool // 'true' if filter data has changed // filters that have been refreshed. updateFlags is true if filter data has
// changed.
updateFilters := d.listsToUpdate(filters, force) 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 { if len(updateFilters) == 0 {
return 0, nil, nil, false return 0, nil, nil, false
} }
@@ -315,7 +322,7 @@ func (d *DNSFilter) refreshFiltersArray(filters *[]FilterYAML, force bool) (int,
updateFlags = append(updateFlags, updated) updateFlags = append(updateFlags, updated)
if err != nil { if err != nil {
failNum++ 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 continue
} }
@@ -325,8 +332,6 @@ func (d *DNSFilter) refreshFiltersArray(filters *[]FilterYAML, force bool) (int,
return 0, nil, nil, true return 0, nil, nil, true
} }
updateCount := 0
d.conf.filtersMu.Lock() d.conf.filtersMu.Lock()
defer d.conf.filtersMu.Unlock() defer d.conf.filtersMu.Unlock()
@@ -345,11 +350,12 @@ func (d *DNSFilter) refreshFiltersArray(filters *[]FilterYAML, force bool) (int,
continue continue
} }
log.Info( d.logger.InfoContext(
"filtering: updated filter %d; rule count: %d (was %d)", ctx,
f.ID, "updated filter",
uf.RulesCount, "id", f.ID,
f.RulesCount, "rules_count", uf.RulesCount,
"prev_rules_count", f.RulesCount,
) )
f.Name = uf.Name 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? // TODO(a.garipov, e.burkov): What the hell?
func (d *DNSFilter) refreshFiltersIntl(block, allow, force bool) (int, bool) { func (d *DNSFilter) refreshFiltersIntl(block, allow, force bool) (int, bool) {
ctx := context.TODO()
updNum := 0 updNum := 0
log.Debug("filtering: starting updating") d.logger.DebugContext(ctx, "starting update")
defer func() { log.Debug("filtering: finished updating, %d updated", updNum) }() defer func() {
d.logger.DebugContext(ctx, "finished update", "updated", updNum)
}()
var lists []FilterYAML var lists []FilterYAML
var toUpd []bool var toUpd []bool
isNetErr := false isNetErr := false
if block { 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 { 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 updNum += updNumAl
lists = append(lists, listsAl...) lists = append(lists, listsAl...)
@@ -417,7 +431,7 @@ func (d *DNSFilter) refreshFiltersIntl(block, allow, force bool) (int, bool) {
p := uf.Path(d.conf.DataDir) p := uf.Path(d.conf.DataDir)
err := os.Remove(p + ".old") err := os.Remove(p + ".old")
if err != nil { 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. // update refreshes filter's content and a/mtimes of it's file.
func (d *DNSFilter) update(filter *FilterYAML) (b bool, err error) { 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() filter.LastUpdated = time.Now()
if !b { if !b {
chErr := os.Chtimes( chErr := os.Chtimes(
@@ -436,7 +452,7 @@ func (d *DNSFilter) update(filter *FilterYAML) (b bool, err error) {
filter.LastUpdated, filter.LastUpdated,
) )
if chErr != nil { 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 // updateIntl updates the flt rewriting it's actual file. It returns true if
// the actual update has been performed. // the actual update has been performed.
func (d *DNSFilter) updateIntl(flt *FilterYAML) (ok bool, err error) { func (d *DNSFilter) updateIntl(ctx context.Context, flt *FilterYAML) (ok bool, err error) {
log.Debug("filtering: downloading update for filter %d from %q", flt.ID, flt.URL) d.logger.DebugContext(ctx, "downloading update for filter", "id", flt.ID, "url", flt.URL)
var res *rulelist.ParseResult var res *rulelist.ParseResult
@@ -454,7 +470,7 @@ func (d *DNSFilter) updateIntl(flt *FilterYAML) (ok bool, err error) {
if err != nil { if err != nil {
return false, err 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) r, err := d.reader(flt.URL)
if err != nil { 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 // according to updated. It also saves new values of flt's name, rules number
// and checksum if succeeded. // and checksum if succeeded.
func (d *DNSFilter) finalizeUpdate( func (d *DNSFilter) finalizeUpdate(
ctx context.Context,
file aghrenameio.PendingFile, file aghrenameio.PendingFile,
flt *FilterYAML, flt *FilterYAML,
res *rulelist.ParseResult, res *rulelist.ParseResult,
@@ -485,13 +502,13 @@ func (d *DNSFilter) finalizeUpdate(
id := flt.ID id := flt.ID
if !updated { if !updated {
if returned == nil { 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()) 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() err = file.CloseReplace()
if err != nil { if err != nil {
@@ -499,7 +516,13 @@ func (d *DNSFilter) finalizeUpdate(
} }
rulesCount := res.RulesCount 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.ensureName(res.Title)
flt.checksum = res.Checksum 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 // 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) 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) file, err := os.Open(fileName)
if errors.Is(err, os.ErrNotExist) { 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) 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() bufPtr := d.bufPool.Get()
defer d.bufPool.Put(bufPtr) defer d.bufPool.Put(bufPtr)
@@ -586,14 +609,16 @@ func (d *DNSFilter) load(flt *FilterYAML) (err error) {
return nil return nil
} }
// EnableFilters enables filters.
func (d *DNSFilter) EnableFilters(async bool) { func (d *DNSFilter) EnableFilters(async bool) {
d.conf.filtersMu.RLock() d.conf.filtersMu.RLock()
defer d.conf.filtersMu.RUnlock() 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 := make([]Filter, 1, len(d.conf.Filters)+len(d.conf.WhitelistFilters)+1)
filters[0] = Filter{ filters[0] = Filter{
ID: rulelist.URLFilterIDCustom, 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 { if err != nil {
log.Error("filtering: enabling filters: %s", err) d.logger.ErrorContext(ctx, "enabling filters", slogutil.KeyError, err)
} }
d.SetEnabled(d.conf.FilteringEnabled) d.SetEnabled(d.conf.FilteringEnabled)

View File

@@ -1,6 +1,7 @@
package filtering package filtering
import ( import (
"context"
"net" "net"
"net/http" "net/http"
"net/url" "net/url"
@@ -9,6 +10,7 @@ import (
"testing" "testing"
"time" "time"
"github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/AdguardTeam/golibs/netutil/urlutil" "github.com/AdguardTeam/golibs/netutil/urlutil"
"github.com/AdguardTeam/golibs/testutil" "github.com/AdguardTeam/golibs/testutil"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
@@ -56,6 +58,7 @@ func serveFiltersLocally(t *testing.T, fltContent []byte) (urlStr string) {
// count. // count.
func updateAndAssert( func updateAndAssert(
t *testing.T, t *testing.T,
ctx context.Context,
dnsFilter *DNSFilter, dnsFilter *DNSFilter,
f *FilterYAML, f *FilterYAML,
wantUpd require.BoolAssertionFunc, wantUpd require.BoolAssertionFunc,
@@ -75,7 +78,7 @@ func updateAndAssert(
assert.Len(t, dir, 1) assert.Len(t, dir, 1)
err = dnsFilter.load(f) err = dnsFilter.load(ctx, f)
require.NoError(t, err) require.NoError(t, err)
} }
@@ -84,6 +87,7 @@ func newDNSFilter(t *testing.T) (d *DNSFilter) {
t.Helper() t.Helper()
dnsFilter, err := New(&Config{ dnsFilter, err := New(&Config{
Logger: slogutil.NewDiscardLogger(),
DataDir: t.TempDir(), DataDir: t.TempDir(),
HTTPClient: &http.Client{ HTTPClient: &http.Client{
Timeout: testTimeout, Timeout: testTimeout,
@@ -95,6 +99,8 @@ func newDNSFilter(t *testing.T) (d *DNSFilter) {
} }
func TestDNSFilter_Update(t *testing.T) { func TestDNSFilter_Update(t *testing.T) {
ctx := testutil.ContextWithTimeout(t, testTimeout)
const content = `||example.org^$third-party const content = `||example.org^$third-party
# Inline comment example # Inline comment example
||example.com^$third-party ||example.com^$third-party
@@ -111,11 +117,11 @@ func TestDNSFilter_Update(t *testing.T) {
dnsFilter := newDNSFilter(t) dnsFilter := newDNSFilter(t)
t.Run("download", func(t *testing.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) { 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) { t.Run("refresh_actually", func(t *testing.T) {
@@ -125,11 +131,11 @@ func TestDNSFilter_Update(t *testing.T) {
f.URL = serveFiltersLocally(t, anotherContent) f.URL = serveFiltersLocally(t, anotherContent)
t.Cleanup(func() { f.URL = oldURL }) 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) { t.Run("load_unload", func(t *testing.T) {
err := dnsFilter.load(f) err := dnsFilter.load(ctx, f)
require.NoError(t, err) require.NoError(t, err)
f.unload() f.unload()
@@ -137,6 +143,8 @@ func TestDNSFilter_Update(t *testing.T) {
} }
func TestFilterYAML_EnsureName(t *testing.T) { func TestFilterYAML_EnsureName(t *testing.T) {
ctx := testutil.ContextWithTimeout(t, testTimeout)
dnsFilter := newDNSFilter(t) dnsFilter := newDNSFilter(t)
t.Run("title_custom", func(t *testing.T) { t.Run("title_custom", func(t *testing.T) {
@@ -147,7 +155,7 @@ func TestFilterYAML_EnsureName(t *testing.T) {
Name: "user-custom", 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) assert.Equal(t, "user-custom", f.Name)
}) })
@@ -158,7 +166,7 @@ func TestFilterYAML_EnsureName(t *testing.T) {
URL: serveFiltersLocally(t, content), 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) assert.Equal(t, "src-title", f.Name)
}) })
@@ -169,7 +177,7 @@ func TestFilterYAML_EnsureName(t *testing.T) {
URL: serveFiltersLocally(t, content), 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) assert.Equal(t, "List 0", f.Name)
}) })
} }

View File

@@ -5,6 +5,7 @@ import (
"context" "context"
"fmt" "fmt"
"io/fs" "io/fs"
"log/slog"
"net" "net"
"net/http" "net/http"
"net/netip" "net/netip"
@@ -24,7 +25,7 @@ import (
"github.com/AdguardTeam/golibs/container" "github.com/AdguardTeam/golibs/container"
"github.com/AdguardTeam/golibs/errors" "github.com/AdguardTeam/golibs/errors"
"github.com/AdguardTeam/golibs/hostsfile" "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/mathutil"
"github.com/AdguardTeam/golibs/syncutil" "github.com/AdguardTeam/golibs/syncutil"
"github.com/AdguardTeam/urlfilter" "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. // Config allows you to configure DNS filtering with New() or just change variables directly.
type Config struct { 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 is the IP address to be returned for a blocked A request.
BlockingIPv4 netip.Addr `yaml:"blocking_ipv4"` BlockingIPv4 netip.Addr `yaml:"blocking_ipv4"`
@@ -235,6 +240,9 @@ type Checker interface {
// DNSFilter matches hostnames and DNS requests against filtering rules. // DNSFilter matches hostnames and DNS requests against filtering rules.
type DNSFilter struct { type DNSFilter struct {
// logger is used for logging the filtering process.
logger *slog.Logger
// idGen is used to generate IDs for package urlfilter. // idGen is used to generate IDs for package urlfilter.
idGen *idGenerator idGen *idGenerator
@@ -413,7 +421,12 @@ func (d *DNSFilter) WriteDiskConfig(c *Config) {
// filters are ready. // filters are ready.
// //
// In this case the caller must ensure that the old filter files are intact. // 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 { if async {
params := filtersInitializerParams{ params := filtersInitializerParams{
allowFilters: allowFilters, allowFilters: allowFilters,
@@ -439,7 +452,7 @@ func (d *DNSFilter) setFilters(blockFilters, allowFilters []Filter, async bool)
return nil return nil
} }
return d.initFiltering(allowFilters, blockFilters) return d.initFiltering(ctx, allowFilters, blockFilters)
} }
// Close - close the object // Close - close the object
@@ -451,19 +464,19 @@ func (d *DNSFilter) Close() {
d.done <- struct{}{} 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 d.rulesStorage != nil {
if err := d.rulesStorage.Close(); err != 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 d.rulesStorageAllow != nil {
if err := d.rulesStorageAllow.Close(); err != 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() d.confMu.RLock()
defer d.confMu.RUnlock() defer d.confMu.RUnlock()
ctx := context.TODO()
rewrites, matched := findRewrites(d.conf.Rewrites, host, qtype) rewrites, matched := findRewrites(d.conf.Rewrites, host, qtype)
if !matched { if !matched {
return Result{} return Result{}
@@ -663,7 +678,7 @@ func (d *DNSFilter) processRewrites(host string, qtype uint16) (res Result) {
rwPat := rw.Domain rwPat := rw.Domain
rwAns := rw.Answer 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 { if origHost == rwAns || rwPat == rwAns {
// Either a request for the hostname itself or a rewrite of // 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 host = rwAns
if cnames.Has(host) { 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 return res
} }
@@ -692,15 +707,15 @@ func (d *DNSFilter) processRewrites(host string, qtype uint16) (res Result) {
rewrites, matched = findRewrites(d.conf.Rewrites, host, qtype) rewrites, matched = findRewrites(d.conf.Rewrites, host, qtype)
} }
setRewriteResult(&res, host, rewrites, qtype) d.setRewriteResult(ctx, &res, host, rewrites, qtype)
return res return res
} }
// matchBlockedServicesRules checks the host against the blocked services rules // 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 // in settings, if any. err is always nil, it is only there to make this a
// a valid hostChecker function. // valid hostChecker function.
func matchBlockedServicesRules( func (d *DNSFilter) matchBlockedServicesRules(
host string, host string,
_ uint16, _ uint16,
setts *Settings, setts *Settings,
@@ -728,8 +743,13 @@ func matchBlockedServicesRules(
Text: ruleText, Text: ruleText,
}} }}
log.Debug("blocked services: matched rule: %s host: %s service: %s", d.logger.DebugContext(
ruleText, host, s.Name) context.TODO(),
"blocked services matched rule",
"rule", ruleText,
"host", host,
"service", s.Name,
)
return res, nil return res, nil
} }
@@ -793,7 +813,7 @@ func newRuleStorage(filters []Filter) (rs *filterlist.RuleStorage, err error) {
} }
// Initialize urlfilter objects. // 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) rulesStorage, err := newRuleStorage(blockFilters)
if err != nil { if err != nil {
return err return err
@@ -811,7 +831,7 @@ func (d *DNSFilter) initFiltering(allowFilters, blockFilters []Filter) (err erro
d.engineLock.Lock() d.engineLock.Lock()
defer d.engineLock.Unlock() defer d.engineLock.Unlock()
d.reset() d.reset(ctx)
d.rulesStorage = rulesStorage d.rulesStorage = rulesStorage
d.filteringEngine = filteringEngine d.filteringEngine = filteringEngine
d.rulesStorageAllow = rulesStorageAllow 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. // Make sure that the OS reclaims memory as soon as possible.
debug.FreeOSMemory() debug.FreeOSMemory()
log.Debug("filtering: initialized filtering engine") d.logger.DebugContext(ctx, "initialized filtering engine")
return nil return nil
} }
@@ -843,6 +863,7 @@ func hostRulesToRules(netRules []*rules.HostRule) (res []rules.Rule) {
// matchHostProcessAllowList processes the allowlist logic of host matching. // matchHostProcessAllowList processes the allowlist logic of host matching.
func (d *DNSFilter) matchHostProcessAllowList( func (d *DNSFilter) matchHostProcessAllowList(
ctx context.Context,
host string, host string,
dnsres *urlfilter.DNSResult, dnsres *urlfilter.DNSResult,
) (res Result, err error) { ) (res Result, err error) {
@@ -859,7 +880,12 @@ func (d *DNSFilter) matchHostProcessAllowList(
return Result{}, fmt.Errorf("invalid dns result: rules are empty") 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 return makeResult(matchedRules, NotFilteredAllowList), nil
} }
@@ -929,6 +955,8 @@ func (d *DNSFilter) matchHost(
return Result{}, nil return Result{}, nil
} }
ctx := context.TODO()
ufReq := &urlfilter.DNSRequest{ ufReq := &urlfilter.DNSRequest{
Hostname: host, Hostname: host,
SortedClientTags: setts.ClientTags, SortedClientTags: setts.ClientTags,
@@ -947,7 +975,7 @@ func (d *DNSFilter) matchHost(
if setts.ProtectionEnabled && d.filteringEngineAllow != nil { if setts.ProtectionEnabled && d.filteringEngineAllow != nil {
dnsres, ok := d.filteringEngineAllow.MatchRequest(ufReq) dnsres, ok := d.filteringEngineAllow.MatchRequest(ufReq)
if ok { 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) res = d.matchHostProcessDNSResult(rrtype, dnsres)
for _, r := range res.Rules { for _, r := range res.Rules {
log.Debug( d.logger.DebugContext(
"filtering: found rule %q for host %q, filter list id: %d", ctx,
r.Text, "found rule for host",
host, "host", host,
r.FilterListID, "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. // InitModule manually initializes blocked services map. l must not be nil.
func InitModule() { func InitModule(ctx context.Context, l *slog.Logger) {
initBlockedServices() initBlockedServices(ctx, l)
} }
// New creates properly initialized DNS Filter that is ready to be used. c must // New creates properly initialized DNS Filter that is ready to be used. c must
// be non-nil. // be non-nil.
func New(c *Config, blockFilters []Filter) (d *DNSFilter, err error) { func New(c *Config, blockFilters []Filter) (d *DNSFilter, err error) {
ctx := context.TODO()
d = &DNSFilter{ 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), bufPool: syncutil.NewSlicePool[byte](rulelist.DefaultRuleBufSize),
safeSearch: c.SafeSearch, safeSearch: c.SafeSearch,
refreshLock: &sync.Mutex{}, refreshLock: &sync.Mutex{},
@@ -1036,7 +1068,7 @@ func New(c *Config, blockFilters []Filter) (d *DNSFilter, err error) {
check: d.matchHost, check: d.matchHost,
name: "filtering", name: "filtering",
}, { }, {
check: matchBlockedServicesRules, check: d.matchBlockedServicesRules,
name: "blocked services", name: "blocked services",
}, { }, {
check: d.checkSafeBrowsing, check: d.checkSafeBrowsing,
@@ -1054,7 +1086,7 @@ func New(c *Config, blockFilters []Filter) (d *DNSFilter, err error) {
d.conf = c d.conf = c
d.conf.filtersMu = &sync.RWMutex{} d.conf.filtersMu = &sync.RWMutex{}
err = d.prepareRewrites() err = d.prepareRewrites(ctx)
if err != nil { if err != nil {
return nil, fmt.Errorf("rewrites: preparing: %w", err) 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 { if blockFilters != nil {
err = d.initFiltering(nil, blockFilters) err = d.initFiltering(ctx, nil, blockFilters)
if err != nil { if err != nil {
d.Close() 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) return nil, fmt.Errorf("making filtering directory: %w", err)
} }
d.loadFilters(d.conf.Filters) d.loadFilters(ctx, d.conf.Filters)
d.loadFilters(d.conf.WhitelistFilters) d.loadFilters(ctx, d.conf.WhitelistFilters)
d.conf.Filters = deduplicateFilters(d.conf.Filters) d.conf.Filters = deduplicateFilters(d.conf.Filters)
d.conf.WhitelistFilters = deduplicateFilters(d.conf.WhitelistFilters) d.conf.WhitelistFilters = deduplicateFilters(d.conf.WhitelistFilters)
@@ -1101,12 +1133,12 @@ func (d *DNSFilter) Start() {
d.RegisterFilteringHandlers() d.RegisterFilteringHandlers()
go d.updatesLoop() go d.updatesLoop(context.TODO())
} }
// updatesLoop initializes new filters and checks for filters updates in a loop. // updatesLoop initializes new filters and checks for filters updates in a loop.
func (d *DNSFilter) updatesLoop() { func (d *DNSFilter) updatesLoop(ctx context.Context) {
defer log.OnPanic("filtering: updates loop") defer slogutil.RecoverAndLog(ctx, d.logger)
ivl := time.Second * 5 ivl := time.Second * 5
t := time.NewTimer(ivl) t := time.NewTimer(ivl)
@@ -1114,9 +1146,9 @@ func (d *DNSFilter) updatesLoop() {
for { for {
select { select {
case params := <-d.filtersInitializerChan: case params := <-d.filtersInitializerChan:
err := d.initFiltering(params.allowFilters, params.blockFilters) err := d.initFiltering(ctx, params.allowFilters, params.blockFilters)
if err != nil { if err != nil {
log.Error("filtering: initializing: %s", err) d.logger.ErrorContext(ctx, "initializing", slogutil.KeyError, err)
continue continue
} }
@@ -1165,9 +1197,13 @@ func (d *DNSFilter) checkSafeBrowsing(
return Result{}, nil return Result{}, nil
} }
if log.GetLevel() >= log.DEBUG { ctx := context.TODO()
timer := log.StartTimer() if d.logger.Enabled(ctx, slogutil.LevelDebug) {
defer timer.LogElapsed("filtering: safebrowsing lookup for %q", host) startTime := time.Now()
defer func() {
elapsed := time.Since(startTime)
d.logger.DebugContext(ctx, "safebrowsing lookup", "host", host, "elapsed", elapsed)
}()
} }
res = Result{ res = Result{
@@ -1197,9 +1233,13 @@ func (d *DNSFilter) checkParental(
return Result{}, nil return Result{}, nil
} }
if log.GetLevel() >= log.DEBUG { ctx := context.TODO()
timer := log.StartTimer() if d.logger.Enabled(ctx, slogutil.LevelDebug) {
defer timer.LogElapsed("filtering: parental lookup for %q", host) startTime := time.Now()
defer func() {
elapsed := time.Since(startTime)
d.logger.DebugContext(ctx, "parental lookup", "host", host, "elapsed", elapsed)
}()
} }
res = Result{ res = Result{

View File

@@ -2,13 +2,14 @@ package filtering
import ( import (
"bytes" "bytes"
"cmp"
"fmt" "fmt"
"net/netip" "net/netip"
"testing" "testing"
"github.com/AdguardTeam/AdGuardHome/internal/aghtest" "github.com/AdguardTeam/AdGuardHome/internal/aghtest"
"github.com/AdguardTeam/AdGuardHome/internal/filtering/hashprefix" "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/netutil"
"github.com/AdguardTeam/golibs/testutil" "github.com/AdguardTeam/golibs/testutil"
"github.com/AdguardTeam/urlfilter/rules" "github.com/AdguardTeam/urlfilter/rules"
@@ -17,15 +18,14 @@ import (
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
func TestMain(m *testing.M) {
testutil.DiscardLogOutput(m)
}
const ( const (
sbBlocked = "wmconvirus.narod.ru" sbBlocked = "wmconvirus.narod.ru"
pcBlocked = "pornhub.com" pcBlocked = "pornhub.com"
) )
// testLogger is the common logger for tests.
var testLogger = slogutil.NewDiscardLogger()
// Helpers. // Helpers.
func newForTest(t testing.TB, c *Config, filters []Filter) (f *DNSFilter, setts *Settings) { 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, FilteringEnabled: true,
} }
if c != nil { if c != nil {
c.Logger = cmp.Or(c.Logger, testLogger)
c.SafeBrowsingCacheSize = 10000 c.SafeBrowsingCacheSize = 10000
c.ParentalCacheSize = 10000 c.ParentalCacheSize = 10000
c.SafeSearchCacheSize = 1000 c.SafeSearchCacheSize = 1000
@@ -43,7 +44,9 @@ func newForTest(t testing.TB, c *Config, filters []Filter) (f *DNSFilter, setts
setts.ParentalEnabled = c.ParentalEnabled setts.ParentalEnabled = c.ParentalEnabled
} else { } else {
// It must not be nil. // It must not be nil.
c = &Config{} c = &Config{
Logger: testLogger,
}
} }
f, err := New(c, filters) f, err := New(c, filters)
require.NoError(t, err) 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 { func newChecker(host string) Checker {
return hashprefix.New(&hashprefix.Config{ return hashprefix.New(&hashprefix.Config{
Logger: testLogger,
CacheTime: 10, CacheTime: 10,
CacheSize: 100000, CacheSize: 100000,
Upstream: aghtest.NewBlockUpstream(host, true), Upstream: aghtest.NewBlockUpstream(host, true),
@@ -168,12 +172,15 @@ func TestDNSFilter_CheckHost_hostRules(t *testing.T) {
func TestSafeBrowsing(t *testing.T) { func TestSafeBrowsing(t *testing.T) {
logOutput := &bytes.Buffer{} logOutput := &bytes.Buffer{}
aghtest.ReplaceLogWriter(t, logOutput)
aghtest.ReplaceLogLevel(t, log.DEBUG)
sbChecker := newChecker(sbBlocked) sbChecker := newChecker(sbBlocked)
d, setts := newForTest(t, &Config{ d, setts := newForTest(t, &Config{
Logger: slogutil.New(&slogutil.Config{
Level: slogutil.LevelDebug,
Output: logOutput,
Format: slogutil.FormatDefault,
AddTimestamp: false,
}),
SafeBrowsingEnabled: true, SafeBrowsingEnabled: true,
SafeBrowsingChecker: sbChecker, SafeBrowsingChecker: sbChecker,
}, nil) }, nil)
@@ -181,7 +188,7 @@ func TestSafeBrowsing(t *testing.T) {
d.checkMatch(t, sbBlocked, setts) 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.checkMatch(t, "test."+sbBlocked, setts)
d.checkMatchEmpty(t, "yandex.ru", setts) d.checkMatchEmpty(t, "yandex.ru", setts)
@@ -216,17 +223,21 @@ func TestParallelSB(t *testing.T) {
func TestParentalControl(t *testing.T) { func TestParentalControl(t *testing.T) {
logOutput := &bytes.Buffer{} logOutput := &bytes.Buffer{}
aghtest.ReplaceLogWriter(t, logOutput)
aghtest.ReplaceLogLevel(t, log.DEBUG)
d, setts := newForTest(t, &Config{ d, setts := newForTest(t, &Config{
Logger: slogutil.New(&slogutil.Config{
Level: slogutil.LevelDebug,
Output: logOutput,
Format: slogutil.FormatDefault,
AddTimestamp: false,
}),
ParentalEnabled: true, ParentalEnabled: true,
ParentalControlChecker: newChecker(pcBlocked), ParentalControlChecker: newChecker(pcBlocked),
}, nil) }, nil)
t.Cleanup(d.Close) t.Cleanup(d.Close)
d.checkMatch(t, pcBlocked, setts) 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.checkMatch(t, "www."+pcBlocked, setts)
d.checkMatchEmpty(t, "www.yandex.ru", setts) d.checkMatchEmpty(t, "www.yandex.ru", setts)
@@ -548,7 +559,8 @@ func TestWhitelist(t *testing.T) {
}} }}
d, setts := newForTest(t, nil, filters) 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) require.NoError(t, err)
t.Cleanup(d.Close) t.Cleanup(d.Close)
@@ -663,6 +675,7 @@ func TestClientSettings(t *testing.T) {
func BenchmarkSafeBrowsing(b *testing.B) { func BenchmarkSafeBrowsing(b *testing.B) {
d, setts := newForTest(b, &Config{ d, setts := newForTest(b, &Config{
Logger: testLogger,
SafeBrowsingEnabled: true, SafeBrowsingEnabled: true,
SafeBrowsingChecker: newChecker(sbBlocked), SafeBrowsingChecker: newChecker(sbBlocked),
}, nil) }, nil)
@@ -689,6 +702,7 @@ func BenchmarkSafeBrowsing(b *testing.B) {
func BenchmarkSafeBrowsing_parallel(b *testing.B) { func BenchmarkSafeBrowsing_parallel(b *testing.B) {
d, setts := newForTest(b, &Config{ d, setts := newForTest(b, &Config{
Logger: testLogger,
SafeBrowsingEnabled: true, SafeBrowsingEnabled: true,
SafeBrowsingChecker: newChecker(sbBlocked), SafeBrowsingChecker: newChecker(sbBlocked),
}, nil) }, nil)

View File

@@ -1,10 +1,9 @@
package hashprefix package hashprefix
import ( import (
"context"
"encoding/binary" "encoding/binary"
"time" "time"
"github.com/AdguardTeam/golibs/log"
) )
// expirySize is the size of expiry in cacheItem. // expirySize is the size of expiry in cacheItem.
@@ -91,7 +90,7 @@ func (c *Checker) findInCache(
} }
// storeInCache caches hashes. // 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) hashToStore := make(map[prefix][]hostnameHash)
for _, hash := range respHashes { for _, hash := range respHashes {
@@ -102,7 +101,7 @@ func (c *Checker) storeInCache(hashesToRequest, respHashes []hostnameHash) {
} }
for pref, hash := range hashToStore { for pref, hash := range hashToStore {
c.setCache(pref, hash) c.setCache(ctx, pref, hash)
} }
for _, hash := range hashesToRequest { for _, hash := range hashesToRequest {
@@ -111,18 +110,18 @@ func (c *Checker) storeInCache(hashesToRequest, respHashes []hostnameHash) {
var pref prefix var pref prefix
copy(pref[:], hash[:]) copy(pref[:], hash[:])
c.setCache(pref, nil) c.setCache(ctx, pref, nil)
} }
} }
} }
// setCache stores hash in cache. // 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{ item := &cacheItem{
expiry: time.Now().Add(c.cacheTime), expiry: time.Now().Add(c.cacheTime),
hashes: hashes, hashes: hashes,
} }
c.cache.Set(pref[:], fromCacheItem(item)) 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 package hashprefix
import ( import (
"context"
"crypto/sha256" "crypto/sha256"
"encoding/hex" "encoding/hex"
"fmt" "fmt"
"log/slog"
"slices" "slices"
"strings" "strings"
"time" "time"
"github.com/AdguardTeam/dnsproxy/upstream" "github.com/AdguardTeam/dnsproxy/upstream"
"github.com/AdguardTeam/golibs/cache" "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/netutil"
"github.com/AdguardTeam/golibs/stringutil" "github.com/AdguardTeam/golibs/stringutil"
"github.com/miekg/dns" "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 // Config is the configuration structure for safe browsing and parental
// control. // control.
type Config struct { 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 is the upstream DNS server.
Upstream upstream.Upstream Upstream upstream.Upstream
// ServiceName is the name of the service.
ServiceName string
// TXTSuffix is the TXT suffix for DNS request. // TXTSuffix is the TXT suffix for DNS request.
TXTSuffix string TXTSuffix string
@@ -70,15 +72,15 @@ type Config struct {
} }
type Checker struct { type Checker struct {
// logger is used for logging the check process.
logger *slog.Logger
// upstream is the upstream DNS server. // upstream is the upstream DNS server.
upstream upstream.Upstream upstream upstream.Upstream
// cache stores hostname hashes. // cache stores hostname hashes.
cache cache.Cache cache cache.Cache
// svc is the name of the service.
svc string
// txtSuffix is the TXT suffix for DNS request. // txtSuffix is the TXT suffix for DNS request.
txtSuffix string txtSuffix string
@@ -89,12 +91,12 @@ type Checker struct {
// New returns Checker. // New returns Checker.
func New(conf *Config) (c *Checker) { func New(conf *Config) (c *Checker) {
return &Checker{ return &Checker{
logger: conf.Logger,
upstream: conf.Upstream, upstream: conf.Upstream,
cache: cache.New(cache.Config{ cache: cache.New(cache.Config{
EnableLRU: true, EnableLRU: true,
MaxSize: conf.CacheSize, MaxSize: conf.CacheSize,
}), }),
svc: conf.ServiceName,
txtSuffix: conf.TXTSuffix, txtSuffix: conf.TXTSuffix,
cacheTime: conf.CacheTime, cacheTime: conf.CacheTime,
} }
@@ -102,18 +104,22 @@ func New(conf *Config) (c *Checker) {
// Check returns true if request for the host should be blocked. // Check returns true if request for the host should be blocked.
func (c *Checker) Check(host string) (ok bool, err error) { func (c *Checker) Check(host string) (ok bool, err error) {
ctx := context.TODO()
hashes := hostnameToHashes(host) hashes := hostnameToHashes(host)
l := c.logger.With("host", host)
found, blocked, hashesToRequest := c.findInCache(hashes) found, blocked, hashesToRequest := c.findInCache(hashes)
if found { 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 return blocked, nil
} }
question := c.getQuestion(hashesToRequest) 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) req := (&dns.Msg{}).SetQuestion(question, dns.TypeTXT)
resp, err := c.upstream.Exchange(req) 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) 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 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 // 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( func (c *Checker) processAnswer(
ctx context.Context,
l *slog.Logger,
hashesToRequest []hostnameHash, hashesToRequest []hostnameHash,
resp *dns.Msg, resp *dns.Msg,
host string,
) (matched bool, receivedHashes []hostnameHash) { ) (matched bool, receivedHashes []hostnameHash) {
txtCount := 0 txtCount := 0
@@ -198,14 +205,14 @@ func (c *Checker) processAnswer(
txtCount++ 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) matched = findMatch(hashesToRequest, receivedHashes)
if matched { if matched {
log.Debug("%s: matched %s", c.svc, host) l.DebugContext(ctx, "matched")
return true, receivedHashes return true, receivedHashes
} }
@@ -213,24 +220,25 @@ func (c *Checker) processAnswer(
return false, receivedHashes return false, receivedHashes
} }
// appendHashesFromTXT appends received hashed hostnames. // appendHashesFromTXT appends received hashed hostnames. l must not be nil.
func (c *Checker) appendHashesFromTXT( func (c *Checker) appendHashesFromTXT(
ctx context.Context,
l *slog.Logger,
hashes []hostnameHash, hashes []hostnameHash,
txt *dns.TXT, txt *dns.TXT,
host string,
) (receivedHashes []hostnameHash) { ) (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 { for _, t := range txt.Txt {
if len(t) != hexSize { 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 continue
} }
buf, err := hex.DecodeString(t) buf, err := hex.DecodeString(t)
if err != nil { 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 continue
} }

View File

@@ -10,6 +10,8 @@ import (
"github.com/AdguardTeam/AdGuardHome/internal/aghtest" "github.com/AdguardTeam/AdGuardHome/internal/aghtest"
"github.com/AdguardTeam/golibs/cache" "github.com/AdguardTeam/golibs/cache"
"github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/AdguardTeam/golibs/testutil"
"github.com/miekg/dns" "github.com/miekg/dns"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
@@ -42,10 +44,10 @@ func TestChcker_getQuestion(t *testing.T) {
hash = sha256.Sum256([]byte("com")) hash = sha256.Sum256([]byte("com"))
assert.False(t, slices.Contains(hashes, hash)) assert.False(t, slices.Contains(hashes, hash))
c := &Checker{ c := New(&Config{
svc: "SafeBrowsing", Logger: slogutil.NewDiscardLogger(),
txtSuffix: suf, TXTSuffix: suf,
} })
q := c.getQuestion(hashes) q := c.getQuestion(hashes)
@@ -95,10 +97,13 @@ func TestHostnameToHashes(t *testing.T) {
} }
func TestChecker_storeInCache(t *testing.T) { func TestChecker_storeInCache(t *testing.T) {
c := &Checker{ const testTimeout = 1 * time.Second
svc: "SafeBrowsing",
cacheTime: cacheTime, c := New(&Config{
} Logger: slogutil.NewDiscardLogger(),
CacheTime: cacheTime,
})
conf := cache.Config{} conf := cache.Config{}
c.cache = cache.New(conf) c.cache = cache.New(conf)
@@ -112,7 +117,7 @@ func TestChecker_storeInCache(t *testing.T) {
hashesArray = append(hashesArray, hash4) hashesArray = append(hashesArray, hash4)
hash2 := sha256.Sum256([]byte("host.com")) hash2 := sha256.Sum256([]byte("host.com"))
hashesArray = append(hashesArray, hash2) 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 // match "3.sub.host.com" or "host.com" from cache
hashes = []hostnameHash{} hashes = []hostnameHash{}
@@ -152,10 +157,11 @@ func TestChecker_storeInCache(t *testing.T) {
ok = slices.Contains(hashesToRequest, hash) ok = slices.Contains(hashesToRequest, hash)
assert.True(t, ok) assert.True(t, ok)
c = &Checker{ c = New(&Config{
svc: "SafeBrowsing", Logger: slogutil.NewDiscardLogger(),
cacheTime: cacheTime, CacheTime: cacheTime,
} })
c.cache = cache.New(cache.Config{}) c.cache = cache.New(cache.Config{})
hashes = []hostnameHash{} hashes = []hostnameHash{}
@@ -189,6 +195,7 @@ func TestChecker_Check(t *testing.T) {
for _, tc := range testCases { for _, tc := range testCases {
c := New(&Config{ c := New(&Config{
Logger: slogutil.NewDiscardLogger(),
CacheTime: cacheTime, CacheTime: cacheTime,
CacheSize: cacheSize, CacheSize: cacheSize,
}) })

View File

@@ -1,12 +1,13 @@
package filtering package filtering
import ( import (
"context"
"fmt" "fmt"
"net/netip" "net/netip"
"github.com/AdguardTeam/AdGuardHome/internal/filtering/rulelist" "github.com/AdguardTeam/AdGuardHome/internal/filtering/rulelist"
"github.com/AdguardTeam/golibs/hostsfile" "github.com/AdguardTeam/golibs/hostsfile"
"github.com/AdguardTeam/golibs/log" "github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/AdguardTeam/golibs/netutil" "github.com/AdguardTeam/golibs/netutil"
"github.com/AdguardTeam/urlfilter/rules" "github.com/AdguardTeam/urlfilter/rules"
"github.com/miekg/dns" "github.com/miekg/dns"
@@ -24,7 +25,7 @@ func (d *DNSFilter) matchSysHosts(
return Result{}, nil return Result{}, nil
} }
vals, rs, matched := hostsRewrites(qtype, host, d.conf.EtcHosts) vals, rs, matched := d.hostsRewrites(qtype, host, d.conf.EtcHosts)
if !matched { if !matched {
return Result{}, nil return Result{}, nil
} }
@@ -42,11 +43,13 @@ func (d *DNSFilter) matchSysHosts(
} }
// hostsRewrites returns values and rules matched by qt and host within hs. // hostsRewrites returns values and rules matched by qt and host within hs.
func hostsRewrites( func (d *DNSFilter) hostsRewrites(
qtype uint16, qtype uint16,
host string, host string,
hs hostsfile.Storage, hs hostsfile.Storage,
) (vals []rules.RRValue, rls []*ResultRule, matched bool) { ) (vals []rules.RRValue, rls []*ResultRule, matched bool) {
ctx := context.TODO()
var isValidProto func(netip.Addr) (ok bool) var isValidProto func(netip.Addr) (ok bool)
switch qtype { switch qtype {
case dns.TypeA: case dns.TypeA:
@@ -56,7 +59,12 @@ func hostsRewrites(
case dns.TypePTR: case dns.TypePTR:
addr, err := netutil.IPFromReversedAddr(host) addr, err := netutil.IPFromReversedAddr(host)
if err != nil { 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 return nil, nil, false
} }
@@ -73,7 +81,11 @@ func hostsRewrites(
return vals, rls, len(names) > 0 return vals, rls, len(names) > 0
default: default:
log.Debug("filtering: unsupported qtype %d", qtype) d.logger.DebugContext(
ctx,
"unsupported qtype",
"qtype", qtype,
)
return nil, nil, false return nil, nil, false
} }

View File

@@ -10,6 +10,7 @@ import (
"github.com/AdguardTeam/AdGuardHome/internal/aghtest" "github.com/AdguardTeam/AdGuardHome/internal/aghtest"
"github.com/AdguardTeam/AdGuardHome/internal/filtering" "github.com/AdguardTeam/AdGuardHome/internal/filtering"
"github.com/AdguardTeam/AdGuardHome/internal/filtering/rulelist" "github.com/AdguardTeam/AdGuardHome/internal/filtering/rulelist"
"github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/AdguardTeam/golibs/testutil" "github.com/AdguardTeam/golibs/testutil"
"github.com/AdguardTeam/urlfilter/rules" "github.com/AdguardTeam/urlfilter/rules"
"github.com/miekg/dns" "github.com/miekg/dns"
@@ -52,6 +53,7 @@ func TestDNSFilter_CheckHost_hostsContainer(t *testing.T) {
testutil.CleanupAndRequireSuccess(t, hc.Close) testutil.CleanupAndRequireSuccess(t, hc.Close)
conf := &filtering.Config{ conf := &filtering.Config{
Logger: slogutil.NewDiscardLogger(),
EtcHosts: hc, EtcHosts: hc,
} }
f, err := filtering.New(conf, nil) f, err := filtering.New(conf, nil)

View File

@@ -17,7 +17,7 @@ import (
"github.com/AdguardTeam/AdGuardHome/internal/aghhttp" "github.com/AdguardTeam/AdGuardHome/internal/aghhttp"
"github.com/AdguardTeam/AdGuardHome/internal/filtering/rulelist" "github.com/AdguardTeam/AdGuardHome/internal/filtering/rulelist"
"github.com/AdguardTeam/golibs/errors" "github.com/AdguardTeam/golibs/errors"
"github.com/AdguardTeam/golibs/log" "github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/AdguardTeam/golibs/netutil/urlutil" "github.com/AdguardTeam/golibs/netutil/urlutil"
"github.com/miekg/dns" "github.com/miekg/dns"
) )
@@ -148,6 +148,8 @@ func (d *DNSFilter) handleFilteringRemoveURL(w http.ResponseWriter, r *http.Requ
Whitelist bool `json:"whitelist"` Whitelist bool `json:"whitelist"`
} }
ctx := r.Context()
req := request{} req := request{}
err := json.NewDecoder(r.Body).Decode(&req) err := json.NewDecoder(r.Body).Decode(&req)
if err != nil { if err != nil {
@@ -170,7 +172,12 @@ func (d *DNSFilter) handleFilteringRemoveURL(w http.ResponseWriter, r *http.Requ
return flt.URL == req.URL return flt.URL == req.URL
}) })
if delIdx == -1 { 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 return
} }
@@ -179,14 +186,20 @@ func (d *DNSFilter) handleFilteringRemoveURL(w http.ResponseWriter, r *http.Requ
p := deleted.Path(d.conf.DataDir) p := deleted.Path(d.conf.DataDir)
err = os.Rename(p, p+".old") err = os.Rename(p, p+".old")
if err != nil && !errors.Is(err, os.ErrNotExist) { 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 return
} }
*filters = slices.Delete(*filters, delIdx, delIdx+1) *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() d.conf.ConfigModified()

View File

@@ -12,6 +12,7 @@ import (
"time" "time"
"github.com/AdguardTeam/AdGuardHome/internal/schedule" "github.com/AdguardTeam/AdGuardHome/internal/schedule"
"github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/AdguardTeam/golibs/testutil" "github.com/AdguardTeam/golibs/testutil"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
@@ -103,6 +104,7 @@ func TestDNSFilter_handleFilteringSetURL(t *testing.T) {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
confModifiedCalled := false confModifiedCalled := false
d, err := New(&Config{ d, err := New(&Config{
Logger: slogutil.NewDiscardLogger(),
FilteringEnabled: true, FilteringEnabled: true,
Filters: tc.initial, Filters: tc.initial,
HTTPClient: &http.Client{ HTTPClient: &http.Client{
@@ -183,6 +185,7 @@ func TestDNSFilter_handleSafeBrowsingStatus(t *testing.T) {
handlers := make(map[string]http.Handler) handlers := make(map[string]http.Handler)
d, err := New(&Config{ d, err := New(&Config{
Logger: slogutil.NewDiscardLogger(),
ConfigModified: func() { ConfigModified: func() {
testutil.RequireSend(testutil.PanicT{}, confModCh, struct{}{}, testTimeout) testutil.RequireSend(testutil.PanicT{}, confModCh, struct{}{}, testTimeout)
}, },
@@ -267,6 +270,7 @@ func TestDNSFilter_handleParentalStatus(t *testing.T) {
handlers := make(map[string]http.Handler) handlers := make(map[string]http.Handler)
d, err := New(&Config{ d, err := New(&Config{
Logger: slogutil.NewDiscardLogger(),
ConfigModified: func() { ConfigModified: func() {
testutil.RequireSend(testutil.PanicT{}, confModCh, struct{}{}, testTimeout) testutil.RequireSend(testutil.PanicT{}, confModCh, struct{}{}, testTimeout)
}, },
@@ -370,6 +374,7 @@ func TestDNSFilter_HandleCheckHost(t *testing.T) {
} }
dnsFilter, err := New(&Config{ dnsFilter, err := New(&Config{
Logger: slogutil.NewDiscardLogger(),
BlockedServices: &BlockedServices{ BlockedServices: &BlockedServices{
Schedule: schedule.EmptyWeekly(), Schedule: schedule.EmptyWeekly(),
}, },

View File

@@ -1,12 +1,13 @@
package filtering package filtering
import ( import (
"context"
"fmt" "fmt"
"log/slog"
"sync/atomic" "sync/atomic"
"github.com/AdguardTeam/AdGuardHome/internal/filtering/rulelist" "github.com/AdguardTeam/AdGuardHome/internal/filtering/rulelist"
"github.com/AdguardTeam/golibs/container" "github.com/AdguardTeam/golibs/container"
"github.com/AdguardTeam/golibs/log"
) )
// idGenerator generates filtering-list IDs in a way broadly compatible with the // idGenerator generates filtering-list IDs in a way broadly compatible with the
@@ -16,13 +17,15 @@ import (
// rule-list architecture. // rule-list architecture.
type idGenerator struct { type idGenerator struct {
current *atomic.Int32 current *atomic.Int32
logger *slog.Logger
} }
// newIDGenerator returns a new ID generator initialized with the given seed // newIDGenerator returns a new ID generator initialized with the given seed
// value. // value.
func newIDGenerator(seed int32) (g *idGenerator) { func newIDGenerator(seed int32, l *slog.Logger) (g *idGenerator) {
g = &idGenerator{ g = &idGenerator{
current: &atomic.Int32{}, current: &atomic.Int32{},
logger: l,
} }
g.current.Store(seed) g.current.Store(seed)
@@ -61,11 +64,12 @@ func (g *idGenerator) fix(flts []FilterYAML) {
newID = g.next() newID = g.next()
} }
log.Info( g.logger.WarnContext(
"filtering: warning: filter at index %d has duplicate id %d; reassigning to %d", context.TODO(),
i, "filter has duplicate id; reassigning",
id, "idx", i,
newID, "id", id,
"new_id", newID,
) )
flts[i].ID = newID flts[i].ID = newID

View File

@@ -5,6 +5,7 @@ import (
"github.com/AdguardTeam/AdGuardHome/internal/aghalg" "github.com/AdguardTeam/AdGuardHome/internal/aghalg"
"github.com/AdguardTeam/AdGuardHome/internal/filtering/rulelist" "github.com/AdguardTeam/AdGuardHome/internal/filtering/rulelist"
"github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
) )
@@ -64,7 +65,7 @@ func TestIDGenerator_Fix(t *testing.T) {
for _, tc := range testCases { for _, tc := range testCases {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
g := newIDGenerator(1) g := newIDGenerator(1, slogutil.NewDiscardLogger())
g.fix(tc.in) g.fix(tc.in)
assertUniqueIDs(t, tc.in) assertUniqueIDs(t, tc.in)

View File

@@ -2,13 +2,14 @@
package rewrite package rewrite
import ( import (
"context"
"fmt" "fmt"
"log/slog"
"slices" "slices"
"strings" "strings"
"sync" "sync"
"github.com/AdguardTeam/golibs/container" "github.com/AdguardTeam/golibs/container"
"github.com/AdguardTeam/golibs/log"
"github.com/AdguardTeam/urlfilter" "github.com/AdguardTeam/urlfilter"
"github.com/AdguardTeam/urlfilter/filterlist" "github.com/AdguardTeam/urlfilter/filterlist"
"github.com/AdguardTeam/urlfilter/rules" "github.com/AdguardTeam/urlfilter/rules"
@@ -30,8 +31,23 @@ type Storage interface {
List() (items []*Item) 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. // DefaultStorage is the default storage for rewrite rules.
type DefaultStorage struct { type DefaultStorage struct {
// logger is used for logging storage processes. It must not be nil.
logger *slog.Logger
// mu protects items. // mu protects items.
mu *sync.RWMutex mu *sync.RWMutex
@@ -51,13 +67,13 @@ type DefaultStorage struct {
urlFilterID int urlFilterID int
} }
// NewDefaultStorage returns new rewrites storage. listID is used as an // NewDefaultStorage returns new rewrites storage. conf must not be nil.
// identifier of the underlying rules list. rewrites must not be nil. func NewDefaultStorage(conf *Config) (s *DefaultStorage, err error) {
func NewDefaultStorage(listID int, rewrites []*Item) (s *DefaultStorage, err error) {
s = &DefaultStorage{ s = &DefaultStorage{
logger: conf.Logger,
mu: &sync.RWMutex{}, mu: &sync.RWMutex{},
urlFilterID: listID, urlFilterID: conf.ListID,
rewrites: rewrites, rewrites: conf.Rewrites,
} }
s.mu.Lock() s.mu.Lock()
@@ -79,6 +95,8 @@ func (s *DefaultStorage) MatchRequest(dReq *urlfilter.DNSRequest) (rws []*rules.
s.mu.RLock() s.mu.RLock()
defer s.mu.RUnlock() defer s.mu.RUnlock()
ctx := context.TODO()
rrules := s.rewriteRulesForReq(dReq) rrules := s.rewriteRulesForReq(dReq)
if len(rrules) == 0 { if len(rrules) == 0 {
return nil return nil
@@ -91,7 +109,7 @@ func (s *DefaultStorage) MatchRequest(dReq *urlfilter.DNSRequest) (rws []*rules.
rule := rrules[0] rule := rrules[0]
rwAns := rule.DNSRewrite.NewCNAME 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 { if dReq.Hostname == rwAns {
// A request for the hostname itself is an exception rule. // 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) { 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 return nil
} }
@@ -168,12 +186,14 @@ func (s *DefaultStorage) Remove(item *Item) (err error) {
s.mu.Lock() s.mu.Lock()
defer s.mu.Unlock() defer s.mu.Unlock()
ctx := context.TODO()
arr := []*Item{} arr := []*Item{}
// TODO(d.kolyshev): Use slices.IndexFunc + slices.Delete? // TODO(d.kolyshev): Use slices.IndexFunc + slices.Delete?
for _, ent := range s.rewrites { for _, ent := range s.rewrites {
if ent.equal(item) { 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 continue
} }
@@ -215,7 +235,12 @@ func (s *DefaultStorage) resetRules() (err error) {
s.ruleList = strList s.ruleList = strList
s.engine = urlfilter.NewDNSEngine(rs) 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 return nil
} }

View File

@@ -4,6 +4,7 @@ import (
"net/netip" "net/netip"
"testing" "testing"
"github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/AdguardTeam/golibs/netutil" "github.com/AdguardTeam/golibs/netutil"
"github.com/AdguardTeam/urlfilter" "github.com/AdguardTeam/urlfilter"
"github.com/AdguardTeam/urlfilter/rules" "github.com/AdguardTeam/urlfilter/rules"
@@ -18,7 +19,11 @@ func TestNewDefaultStorage(t *testing.T) {
Answer: "answer.com", Answer: "answer.com",
}} }}
s, err := NewDefaultStorage(-1, items) s, err := NewDefaultStorage(&Config{
Logger: slogutil.NewDiscardLogger(),
Rewrites: items,
ListID: -1,
})
require.NoError(t, err) require.NoError(t, err)
require.Len(t, s.List(), 1) require.Len(t, s.List(), 1)
@@ -27,7 +32,11 @@ func TestNewDefaultStorage(t *testing.T) {
func TestDefaultStorage_CRUD(t *testing.T) { func TestDefaultStorage_CRUD(t *testing.T) {
var items []*Item var items []*Item
s, err := NewDefaultStorage(-1, items) s, err := NewDefaultStorage(&Config{
Logger: slogutil.NewDiscardLogger(),
Rewrites: items,
ListID: -1,
})
require.NoError(t, err) require.NoError(t, err)
require.Len(t, s.List(), 0) require.Len(t, s.List(), 0)
@@ -112,7 +121,11 @@ func TestDefaultStorage_MatchRequest(t *testing.T) {
Answer: "sub.issue4016.com", Answer: "sub.issue4016.com",
}} }}
s, err := NewDefaultStorage(-1, items) s, err := NewDefaultStorage(&Config{
Logger: slogutil.NewDiscardLogger(),
Rewrites: items,
ListID: -1,
})
require.NoError(t, err) require.NoError(t, err)
testCases := []struct { testCases := []struct {
@@ -284,7 +297,11 @@ func TestDefaultStorage_MatchRequest_Levels(t *testing.T) {
Answer: addr3.String(), Answer: addr3.String(),
}} }}
s, err := NewDefaultStorage(-1, items) s, err := NewDefaultStorage(&Config{
Logger: slogutil.NewDiscardLogger(),
Rewrites: items,
ListID: -1,
})
require.NoError(t, err) require.NoError(t, err)
testCases := []struct { testCases := []struct {
@@ -352,7 +369,11 @@ func TestDefaultStorage_MatchRequest_ExceptionCNAME(t *testing.T) {
Answer: "*.sub.host.com", Answer: "*.sub.host.com",
}} }}
s, err := NewDefaultStorage(-1, items) s, err := NewDefaultStorage(&Config{
Logger: slogutil.NewDiscardLogger(),
Rewrites: items,
ListID: -1,
})
require.NoError(t, err) require.NoError(t, err)
testCases := []struct { testCases := []struct {
@@ -416,7 +437,11 @@ func TestDefaultStorage_MatchRequest_ExceptionIP(t *testing.T) {
Answer: "A", Answer: "A",
}} }}
s, err := NewDefaultStorage(-1, items) s, err := NewDefaultStorage(&Config{
Logger: slogutil.NewDiscardLogger(),
Rewrites: items,
ListID: -1,
})
require.NoError(t, err) require.NoError(t, err)
testCases := []struct { testCases := []struct {

View File

@@ -6,7 +6,6 @@ import (
"slices" "slices"
"github.com/AdguardTeam/AdGuardHome/internal/aghhttp" "github.com/AdguardTeam/AdGuardHome/internal/aghhttp"
"github.com/AdguardTeam/golibs/log"
) )
// TODO(d.kolyshev): Use [rewrite.Item] instead. // TODO(d.kolyshev): Use [rewrite.Item] instead.
@@ -50,7 +49,7 @@ func (d *DNSFilter) handleRewriteAdd(w http.ResponseWriter, r *http.Request) {
Answer: rwJSON.Answer, Answer: rwJSON.Answer,
} }
err = rw.normalize() err = rw.normalize(r.Context(), d.logger)
if err != nil { if err != nil {
// Shouldn't happen currently, since normalize only returns a non-nil // Shouldn't happen currently, since normalize only returns a non-nil
// error when a rewrite is nil, but be change-proof. // 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() defer d.confMu.Unlock()
d.conf.Rewrites = append(d.conf.Rewrites, rw) d.conf.Rewrites = append(d.conf.Rewrites, rw)
log.Debug( d.logger.DebugContext(
"rewrite: added element: %s -> %s [%d]", r.Context(),
rw.Domain, "added rewrite element",
rw.Answer, "domain", rw.Domain,
len(d.conf.Rewrites), "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 { for _, ent := range d.conf.Rewrites {
if ent.equal(entDel) { 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 continue
} }
@@ -138,7 +143,7 @@ func (d *DNSFilter) handleRewriteUpdate(w http.ResponseWriter, r *http.Request)
Answer: updateJSON.Update.Answer, Answer: updateJSON.Update.Answer,
} }
err = rwAdd.normalize() err = rwAdd.normalize(r.Context(), d.logger)
if err != nil { if err != nil {
// Shouldn't happen currently, since normalize only returns a non-nil // Shouldn't happen currently, since normalize only returns a non-nil
// error when a rewrite is nil, but be change-proof. // 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) d.conf.Rewrites = slices.Replace(d.conf.Rewrites, index, index+1, rwAdd)
log.Debug("rewrite: removed element: %s -> %s", rwDel.Domain, rwDel.Answer) ctx := r.Context()
log.Debug("rewrite: added element: %s -> %s", rwAdd.Domain, rwAdd.Answer) 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" "time"
"github.com/AdguardTeam/AdGuardHome/internal/filtering" "github.com/AdguardTeam/AdGuardHome/internal/filtering"
"github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/AdguardTeam/golibs/testutil" "github.com/AdguardTeam/golibs/testutil"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
@@ -159,6 +160,7 @@ func TestDNSFilter_handleRewriteHTTP(t *testing.T) {
handlers := make(map[string]http.Handler) handlers := make(map[string]http.Handler)
d, err := filtering.New(&filtering.Config{ d, err := filtering.New(&filtering.Config{
Logger: slogutil.NewDiscardLogger(),
ConfigModified: onConfModified, ConfigModified: onConfModified,
HTTPRegister: func(_, url string, handler http.HandlerFunc) { HTTPRegister: func(_, url string, handler http.HandlerFunc) {
handlers[url] = handler handlers[url] = handler

View File

@@ -1,13 +1,15 @@
package filtering package filtering
import ( import (
"context"
"fmt" "fmt"
"log/slog"
"net/netip" "net/netip"
"slices" "slices"
"strings" "strings"
"github.com/AdguardTeam/golibs/errors" "github.com/AdguardTeam/golibs/errors"
"github.com/AdguardTeam/golibs/log" "github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/miekg/dns" "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. // to domain name case, IP length, and so on.
// //
// If rw is nil, it returns an errors. // 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 { if rw == nil {
return errors.Error("nil rewrite entry") return errors.Error("nil rewrite entry")
} }
@@ -85,7 +87,7 @@ func (rw *LegacyRewrite) normalize() (err error) {
ip, err := netip.ParseAddr(rw.Answer) ip, err := netip.ParseAddr(rw.Answer)
if err != nil { if err != nil {
log.Debug("normalizing legacy rewrite: %s", err) l.DebugContext(ctx, "normalizing legacy rewrite", slogutil.KeyError, err)
rw.Type = dns.TypeCNAME rw.Type = dns.TypeCNAME
return nil return nil
@@ -136,9 +138,9 @@ func (rw *LegacyRewrite) Compare(b *LegacyRewrite) (res int) {
} }
// prepareRewrites normalizes and validates all legacy DNS rewrites. // 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 { for i, r := range d.conf.Rewrites {
err = r.normalize() err = r.normalize(ctx, d.logger)
if err != nil { if err != nil {
return fmt.Errorf("at index %d: %w", i, err) 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 // setRewriteResult sets the Reason or IPList of res if necessary. res must not
// be nil. // 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 { for _, rw := range rewrites {
if rw.Type == qtype && (qtype == dns.TypeA || qtype == dns.TypeAAAA) { if rw.Type == qtype && (qtype == dns.TypeA || qtype == dns.TypeAAAA) {
if rw.IP == (netip.Addr{}) { 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) 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" "net/netip"
"testing" "testing"
"github.com/AdguardTeam/golibs/testutil"
"github.com/miekg/dns" "github.com/miekg/dns"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
@@ -88,7 +89,8 @@ func TestRewrites(t *testing.T) {
Answer: addr1v4.String(), Answer: addr1v4.String(),
}} }}
require.NoError(t, d.prepareRewrites()) ctx := testutil.ContextWithTimeout(t, testTimeout)
require.NoError(t, d.prepareRewrites(ctx))
testCases := []struct { testCases := []struct {
name string name string
@@ -236,7 +238,8 @@ func TestRewritesLevels(t *testing.T) {
Type: dns.TypeA, Type: dns.TypeA,
}} }}
require.NoError(t, d.prepareRewrites()) ctx := testutil.ContextWithTimeout(t, testTimeout)
require.NoError(t, d.prepareRewrites(ctx))
testCases := []struct { testCases := []struct {
name string name string
@@ -280,7 +283,8 @@ func TestRewritesExceptionCNAME(t *testing.T) {
Answer: "*.sub.host.com", Answer: "*.sub.host.com",
}} }}
require.NoError(t, d.prepareRewrites()) ctx := testutil.ContextWithTimeout(t, testTimeout)
require.NoError(t, d.prepareRewrites(ctx))
testCases := []struct { testCases := []struct {
name string name string
@@ -342,7 +346,8 @@ func TestRewritesExceptionIP(t *testing.T) {
Type: dns.TypeA, Type: dns.TypeA,
}} }}
require.NoError(t, d.prepareRewrites()) ctx := testutil.ContextWithTimeout(t, testTimeout)
require.NoError(t, d.prepareRewrites(ctx))
testCases := []struct { testCases := []struct {
name string name string

View File

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

View File

@@ -431,6 +431,7 @@ func (web *webAPI) handleInstallConfigure(w http.ResponseWriter, r *http.Request
globalContext.firstRun = false globalContext.firstRun = false
config.DNS.BindHosts = []netip.Addr{req.DNS.IP} config.DNS.BindHosts = []netip.Addr{req.DNS.IP}
config.DNS.Port = req.DNS.Port config.DNS.Port = req.DNS.Port
config.Filtering.Logger = web.baseLogger.With(slogutil.KeyPrefix, "filtering")
config.Filtering.SafeFSPatterns = []string{ config.Filtering.SafeFSPatterns = []string{
filepath.Join(globalContext.workDir, userFilterDataDir, "*"), filepath.Join(globalContext.workDir, userFilterDataDir, "*"),
} }

View File

@@ -357,15 +357,17 @@ func setupDNSFilteringConf(
const ( const (
dnsTimeout = 3 * time.Second dnsTimeout = 3 * time.Second
sbService = "safe browsing" sbService = "safe_browsing"
defaultSafeBrowsingServer = `https://family.adguard-dns.com/dns-query` defaultSafeBrowsingServer = `https://family.adguard-dns.com/dns-query`
sbTXTSuffix = `sb.dns.adguard.com.` sbTXTSuffix = `sb.dns.adguard.com.`
pcService = "parental control" pcService = "parental_control"
defaultParentalServer = `https://family.adguard-dns.com/dns-query` defaultParentalServer = `https://family.adguard-dns.com/dns-query`
pcTXTSuffix = `pc.dns.adguard.com.` pcTXTSuffix = `pc.dns.adguard.com.`
) )
conf.Logger = baseLogger.With(slogutil.KeyPrefix, "filtering")
conf.EtcHosts = globalContext.etcHosts conf.EtcHosts = globalContext.etcHosts
// TODO(s.chzhen): Use empty interface. // TODO(s.chzhen): Use empty interface.
if globalContext.etcHosts == nil || !config.DNS.HostsFileEnabled { if globalContext.etcHosts == nil || !config.DNS.HostsFileEnabled {
@@ -402,11 +404,11 @@ func setupDNSFilteringConf(
} }
conf.SafeBrowsingChecker = hashprefix.New(&hashprefix.Config{ conf.SafeBrowsingChecker = hashprefix.New(&hashprefix.Config{
Upstream: sbUps, Logger: baseLogger.With(slogutil.KeyPrefix, sbService),
ServiceName: sbService, Upstream: sbUps,
TXTSuffix: sbTXTSuffix, TXTSuffix: sbTXTSuffix,
CacheTime: cacheTime, CacheTime: cacheTime,
CacheSize: conf.SafeBrowsingCacheSize, CacheSize: conf.SafeBrowsingCacheSize,
}) })
// Protect against invalid configuration, see #6181. // Protect against invalid configuration, see #6181.
@@ -415,7 +417,11 @@ func setupDNSFilteringConf(
// default. // default.
if conf.SafeBrowsingBlockHost == "" { if conf.SafeBrowsingBlockHost == "" {
host := defaultSafeBrowsingBlockHost 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 conf.SafeBrowsingBlockHost = host
} }
@@ -426,11 +432,11 @@ func setupDNSFilteringConf(
} }
conf.ParentalControlChecker = hashprefix.New(&hashprefix.Config{ conf.ParentalControlChecker = hashprefix.New(&hashprefix.Config{
Upstream: parUps, Logger: baseLogger.With(slogutil.KeyPrefix, pcService),
ServiceName: pcService, Upstream: parUps,
TXTSuffix: pcTXTSuffix, TXTSuffix: pcTXTSuffix,
CacheTime: cacheTime, CacheTime: cacheTime,
CacheSize: conf.ParentalCacheSize, CacheSize: conf.ParentalCacheSize,
}) })
// Protect against invalid configuration, see #6181. // Protect against invalid configuration, see #6181.
@@ -439,7 +445,11 @@ func setupDNSFilteringConf(
// default. // default.
if conf.ParentalBlockHost == "" { if conf.ParentalBlockHost == "" {
host := defaultParentalBlockHost 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 conf.ParentalBlockHost = host
} }
@@ -614,13 +624,13 @@ func run(opts options, clientBuildFS fs.FS, done chan struct{}, sigHdlr *signalH
err = configureOS(config) err = configureOS(config)
fatalOnError(err) fatalOnError(err)
// TODO(s.chzhen): Use it for the entire initialization process.
ctx := context.Background()
// Clients package uses filtering package's static data // Clients package uses filtering package's static data
// (filtering.BlockedSvcKnown()), so we have to initialize filtering static // (filtering.BlockedSvcKnown()), so we have to initialize filtering static
// data first, but also to avoid relying on automatic Go init() function. // data first, but also to avoid relying on automatic Go init() function.
filtering.InitModule() filtering.InitModule(ctx, slogLogger)
// TODO(s.chzhen): Use it for the entire initialization process.
ctx := context.Background()
err = initContextClients(ctx, slogLogger, sigHdlr) err = initContextClients(ctx, slogLogger, sigHdlr)
fatalOnError(err) fatalOnError(err)