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