all: sync with master

This commit is contained in:
Eugene Burkov
2025-08-14 20:45:22 +03:00
parent 79db602f50
commit c2c1a7d586
149 changed files with 3133 additions and 2320 deletions

22
internal/agh/agh.go Normal file
View File

@@ -0,0 +1,22 @@
// Package agh contains common entities and interfaces of AdGuard Home.
package agh
import (
"context"
)
// ConfigModifier defines an interface for updating the global configuration.
type ConfigModifier interface {
// Apply applies changes to the global configuration.
Apply(ctx context.Context)
}
// EmptyConfigModifier is an empty [ConfigModifier] implementation that does
// nothing.
type EmptyConfigModifier struct{}
// type check
var _ ConfigModifier = EmptyConfigModifier{}
// Apply implements the [ConfigModifier] for EmptyConfigModifier.
func (em EmptyConfigModifier) Apply(ctx context.Context) {}

View File

@@ -1,6 +1,7 @@
package aghnet
import (
"context"
"fmt"
"io"
"io/fs"
@@ -102,7 +103,9 @@ func NewHostsContainer(
func (hc *HostsContainer) Close() (err error) {
log.Debug("%s: closing", hostsContainerPrefix)
err = errors.Annotate(hc.watcher.Close(), "closing fs watcher: %w")
// TODO(s.chzhen): Pass context.
ctx := context.TODO()
err = errors.Annotate(hc.watcher.Shutdown(ctx), "closing fs watcher: %w")
// Go on and close the container either way.
close(hc.done)

View File

@@ -7,7 +7,8 @@ import (
"testing/fstest"
"github.com/AdguardTeam/golibs/errors"
"github.com/AdguardTeam/golibs/testutil/fakefs"
"github.com/AdguardTeam/golibs/testutil"
"github.com/AdguardTeam/golibs/testutil/fakeio/fakefs"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
@@ -64,7 +65,7 @@ func TestHostsContainer_PathsToPatterns(t *testing.T) {
const errStat errors.Error = "bad file"
badFS := &fakefs.StatFS{
OnOpen: func(_ string) (f fs.File, err error) { panic("not implemented") },
OnOpen: func(s string) (f fs.File, err error) { panic(testutil.UnexpectedCall(s)) },
OnStat: func(name string) (fi fs.FileInfo, err error) {
return nil, errStat
},

View File

@@ -1,6 +1,7 @@
package aghnet_test
import (
"context"
"net/netip"
"path"
"sync/atomic"
@@ -67,10 +68,12 @@ func TestNewHostsContainer(t *testing.T) {
}
hc, err := aghnet.NewHostsContainer(testFS, &aghtest.FSWatcher{
OnStart: func() (_ error) { panic("not implemented") },
OnEvents: onEvents,
OnAdd: onAdd,
OnClose: func() (err error) { return nil },
OnStart: func(ctx context.Context) (_ error) {
panic(testutil.UnexpectedCall(ctx))
},
OnEvents: onEvents,
OnAdd: onAdd,
OnShutdown: func(_ context.Context) (err error) { return nil },
}, tc.paths...)
if tc.wantErr != nil {
require.ErrorIs(t, err, tc.wantErr)
@@ -94,11 +97,13 @@ func TestNewHostsContainer(t *testing.T) {
t.Run("nil_fs", func(t *testing.T) {
require.Panics(t, func() {
_, _ = aghnet.NewHostsContainer(nil, &aghtest.FSWatcher{
OnStart: func() (_ error) { panic("not implemented") },
OnStart: func(ctx context.Context) (_ error) {
panic(testutil.UnexpectedCall(ctx))
},
// Those shouldn't panic.
OnEvents: func() (e <-chan struct{}) { return nil },
OnAdd: func(name string) (err error) { return nil },
OnClose: func() (err error) { return nil },
OnEvents: func() (e <-chan struct{}) { return nil },
OnAdd: func(_ string) (err error) { return nil },
OnShutdown: func(_ context.Context) (err error) { return nil },
}, p)
})
})
@@ -113,10 +118,10 @@ func TestNewHostsContainer(t *testing.T) {
const errOnAdd errors.Error = "error"
errWatcher := &aghtest.FSWatcher{
OnStart: func() (_ error) { panic("not implemented") },
OnEvents: func() (e <-chan struct{}) { panic("not implemented") },
OnAdd: func(name string) (err error) { return errOnAdd },
OnClose: func() (err error) { return nil },
OnStart: func(ctx context.Context) (_ error) { panic(testutil.UnexpectedCall(ctx)) },
OnEvents: func() (_ <-chan struct{}) { panic(testutil.UnexpectedCall()) },
OnAdd: func(_ string) (err error) { return errOnAdd },
OnShutdown: func(_ context.Context) (err error) { return nil },
}
hc, err := aghnet.NewHostsContainer(testFS, errWatcher, p)
@@ -158,14 +163,14 @@ func TestHostsContainer_refresh(t *testing.T) {
t.Cleanup(func() { close(eventsCh) })
w := &aghtest.FSWatcher{
OnStart: func() (_ error) { panic("not implemented") },
OnStart: func(ctx context.Context) (_ error) { panic(testutil.UnexpectedCall(ctx)) },
OnEvents: func() (e <-chan event) { return eventsCh },
OnAdd: func(name string) (err error) {
assert.Equal(t, "dir", name)
return nil
},
OnClose: func() (err error) { return nil },
OnShutdown: func(_ context.Context) (err error) { return nil },
}
hc, err := aghnet.NewHostsContainer(testFS, w, "dir")

View File

@@ -9,7 +9,7 @@ import (
"github.com/AdguardTeam/golibs/errors"
"github.com/AdguardTeam/golibs/testutil"
"github.com/AdguardTeam/golibs/testutil/fakefs"
"github.com/AdguardTeam/golibs/testutil/fakeio/fakefs"
"github.com/stretchr/testify/assert"
)
@@ -121,7 +121,7 @@ func TestIfaceSetStaticIP(t *testing.T) {
},
}
panicFsys := &fakefs.FS{
OnOpen: func(name string) (fs.File, error) { panic("not implemented") },
OnOpen: func(name string) (_ fs.File, _ error) { panic(testutil.UnexpectedCall(name)) },
}
testCases := []struct {

View File

@@ -19,11 +19,11 @@ import (
// substRootDirFS replaces the aghos.RootDirFS function used throughout the
// package with fsys for tests ran under t.
func substRootDirFS(t testing.TB, fsys fs.FS) {
t.Helper()
func substRootDirFS(tb testing.TB, fsys fs.FS) {
tb.Helper()
prev := rootDirFS
t.Cleanup(func() { rootDirFS = prev })
tb.Cleanup(func() { rootDirFS = prev })
rootDirFS = fsys
}
@@ -32,11 +32,11 @@ type RunCmdFunc func(cmd string, args ...string) (code int, out []byte, err erro
// substShell replaces the the aghos.RunCommand function used throughout the
// package with rc for tests ran under t.
func substShell(t testing.TB, rc RunCmdFunc) {
t.Helper()
func substShell(tb testing.TB, rc RunCmdFunc) {
tb.Helper()
prev := aghosRunCommand
t.Cleanup(func() { aghosRunCommand = prev })
tb.Cleanup(func() { aghosRunCommand = prev })
aghosRunCommand = rc
}
@@ -72,11 +72,11 @@ type ifaceAddrsFunc func() (ifaces []net.Addr, err error)
// substNetInterfaceAddrs replaces the the net.InterfaceAddrs function used
// throughout the package with f for tests ran under t.
func substNetInterfaceAddrs(t *testing.T, f ifaceAddrsFunc) {
t.Helper()
func substNetInterfaceAddrs(tb testing.TB, f ifaceAddrsFunc) {
tb.Helper()
prev := netInterfaceAddrs
t.Cleanup(func() { netInterfaceAddrs = prev })
tb.Cleanup(func() { netInterfaceAddrs = prev })
netInterfaceAddrs = f
}

View File

@@ -1,11 +0,0 @@
package aghos_test
import (
"testing"
"github.com/AdguardTeam/golibs/testutil"
)
func TestMain(m *testing.M) {
testutil.DiscardLogOutput(m)
}

View File

@@ -13,29 +13,33 @@ import (
"github.com/stretchr/testify/require"
)
func TestFileWalker_Walk(t *testing.T) {
const attribute = `000`
// Common file-walker constants.
const (
attribute = "000"
nl = "\n"
)
makeFileWalker := func(_ string) (fw aghos.FileWalker) {
return func(r io.Reader) (patterns []string, cont bool, err error) {
s := bufio.NewScanner(r)
for s.Scan() {
line := s.Text()
if line == attribute {
return nil, false, nil
}
if len(line) != 0 {
patterns = append(patterns, path.Join(".", line))
}
// newFileWalker returns a new file-walker function that reads patterns from an
// [io.Reader].
func newFileWalker() (fw aghos.FileWalker) {
return func(r io.Reader) (patterns []string, cont bool, err error) {
s := bufio.NewScanner(r)
for s.Scan() {
line := s.Text()
if line == attribute {
return nil, false, nil
}
return patterns, true, s.Err()
if len(line) != 0 {
patterns = append(patterns, path.Join(".", line))
}
}
return patterns, true, s.Err()
}
}
const nl = "\n"
func TestFileWalker_Walk(t *testing.T) {
testCases := []struct {
testFS fstest.MapFS
want assert.BoolAssertionFunc
@@ -88,7 +92,7 @@ func TestFileWalker_Walk(t *testing.T) {
}}
for _, tc := range testCases {
fw := makeFileWalker("")
fw := newFileWalker()
t.Run(tc.name, func(t *testing.T) {
ok, err := fw.Walk(tc.testFS, tc.initPattern)
@@ -100,7 +104,7 @@ func TestFileWalker_Walk(t *testing.T) {
t.Run("pattern_malformed", func(t *testing.T) {
f := fstest.MapFS{}
ok, err := makeFileWalker("").Walk(f, "[]")
ok, err := newFileWalker().Walk(f, "[]")
require.Error(t, err)
assert.False(t, ok)

View File

@@ -1,15 +1,17 @@
package aghos
import (
"context"
"fmt"
"io"
"io/fs"
"log/slog"
"path/filepath"
"github.com/AdguardTeam/golibs/container"
"github.com/AdguardTeam/golibs/errors"
"github.com/AdguardTeam/golibs/log"
"github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/AdguardTeam/golibs/osutil"
"github.com/AdguardTeam/golibs/service"
"github.com/fsnotify/fsnotify"
)
@@ -23,11 +25,7 @@ type event = struct{}
//
// TODO(e.burkov): Add tests.
type FSWatcher interface {
// Start starts watching the added files.
Start() (err error)
// Close stops watching the files and closes an update channel.
io.Closer
service.Interface
// Events returns the channel to notify about the file system events.
Events() (e <-chan event)
@@ -39,6 +37,9 @@ type FSWatcher interface {
// osWatcher tracks the file system provided by the OS.
type osWatcher struct {
// logger is used for logging the operations of the osWatcher.
logger *slog.Logger
// watcher is the actual notifier that is handled by osWatcher.
watcher *fsnotify.Watcher
@@ -54,8 +55,8 @@ type osWatcher struct {
const osWatcherPref = "os watcher"
// NewOSWritesWatcher creates FSWatcher that tracks the real file system of the
// OS and notifies only about writing events.
func NewOSWritesWatcher() (w FSWatcher, err error) {
// OS and notifies only about writing events. l must not be nil.
func NewOSWritesWatcher(l *slog.Logger) (w FSWatcher, err error) {
defer func() { err = errors.Annotate(err, "%s: %w", osWatcherPref) }()
var watcher *fsnotify.Watcher
@@ -65,6 +66,7 @@ func NewOSWritesWatcher() (w FSWatcher, err error) {
}
return &osWatcher{
logger: l,
watcher: watcher,
events: make(chan event, 1),
files: container.NewMapSet[string](),
@@ -74,16 +76,16 @@ func NewOSWritesWatcher() (w FSWatcher, err error) {
// type check
var _ FSWatcher = (*osWatcher)(nil)
// Start implements the FSWatcher interface for *osWatcher.
func (w *osWatcher) Start() (err error) {
go w.handleErrors()
go w.handleEvents()
// Start implements the [FSWatcher] interface for *osWatcher.
func (w *osWatcher) Start(ctx context.Context) (err error) {
go w.handleErrors(ctx)
go w.handleEvents(ctx)
return nil
}
// Close implements the FSWatcher interface for *osWatcher.
func (w *osWatcher) Close() (err error) {
// Shutdown implements the [FSWatcher] interface for *osWatcher.
func (w *osWatcher) Shutdown(_ context.Context) (err error) {
return w.watcher.Close()
}
@@ -120,8 +122,8 @@ func (w *osWatcher) Add(name string) (err error) {
// handleEvents notifies about the received file system's event if needed. It
// is intended to be used as a goroutine.
func (w *osWatcher) handleEvents() {
defer log.OnPanic(fmt.Sprintf("%s: handling events", osWatcherPref))
func (w *osWatcher) handleEvents(ctx context.Context) {
defer slogutil.RecoverAndLog(ctx, w.logger)
defer close(w.events)
@@ -131,33 +133,37 @@ func (w *osWatcher) handleEvents() {
continue
}
// Skip the following events assuming that sometimes the same event
// occurs several times.
for ok := true; ok; {
select {
case _, ok = <-ch:
// Go on.
default:
ok = false
}
}
skipDuplicates(ch)
select {
case w.events <- event{}:
// Go on.
default:
log.Debug("%s: events buffer is full", osWatcherPref)
w.logger.DebugContext(ctx, "events buffer is full")
}
}
}
// skipDuplicates drains the given channel of events, assuming that some events
// might occur multiple times.
func skipDuplicates(ch <-chan fsnotify.Event) {
for {
select {
case <-ch:
// Go on.
default:
return
}
}
}
// handleErrors handles accompanying errors. It used to be called in a separate
// goroutine.
func (w *osWatcher) handleErrors() {
defer log.OnPanic(fmt.Sprintf("%s: handling errors", osWatcherPref))
func (w *osWatcher) handleErrors(ctx context.Context) {
defer slogutil.RecoverAndLog(ctx, w.logger)
for err := range w.watcher.Errors {
log.Error("%s: %s", osWatcherPref, err)
w.logger.ErrorContext(ctx, "handling error", slogutil.KeyError, err)
}
}
@@ -170,13 +176,13 @@ var _ FSWatcher = EmptyFSWatcher{}
// Start implements the [FSWatcher] interface for EmptyFSWatcher. It always
// returns nil error.
func (EmptyFSWatcher) Start() (err error) {
func (EmptyFSWatcher) Start(_ context.Context) (err error) {
return nil
}
// Close implements the [FSWatcher] interface for EmptyFSWatcher. It always
// Shutdown implements the [FSWatcher] interface for EmptyFSWatcher. It always
// returns nil error.
func (EmptyFSWatcher) Close() (err error) {
func (EmptyFSWatcher) Shutdown(_ context.Context) (err error) {
return nil
}

View File

@@ -5,9 +5,11 @@ package aghos
import (
"bufio"
"context"
"fmt"
"io"
"io/fs"
"log/slog"
"os"
"os/exec"
"path"
@@ -17,7 +19,6 @@ import (
"strings"
"github.com/AdguardTeam/golibs/errors"
"github.com/AdguardTeam/golibs/log"
)
// Default file, binary, and directory permissions.
@@ -67,8 +68,14 @@ func RunCommand(command string, arguments ...string) (code int, output []byte, e
}
// PIDByCommand searches for process named command and returns its PID ignoring
// the PIDs from except. If no processes found, the error returned.
func PIDByCommand(command string, except ...int) (pid int, err error) {
// the PIDs from except. If no processes found, the error returned. l must not
// be nil.
func PIDByCommand(
ctx context.Context,
l *slog.Logger,
command string,
except ...int,
) (pid int, err error) {
// Don't use -C flag here since it's a feature of linux's ps
// implementation. Use POSIX-compatible flags instead.
//
@@ -101,7 +108,7 @@ func PIDByCommand(command string, except ...int) (pid int, err error) {
case 1:
// Go on.
default:
log.Info("warning: %d %s instances found", instNum, command)
l.WarnContext(ctx, "instances found", "num", instNum, "command", command)
}
if code := cmd.ProcessState.ExitCode(); code != 0 {

View File

@@ -43,14 +43,14 @@ func TestPendingFile(t *testing.T) {
// newInitialFile is a test helper that returns the path to the file containing
// [initialData].
func newInitialFile(t *testing.T) (targetPath string) {
t.Helper()
func newInitialFile(tb testing.TB) (targetPath string) {
tb.Helper()
dir := t.TempDir()
dir := tb.TempDir()
targetPath = filepath.Join(dir, "target")
err := os.WriteFile(targetPath, initialData, 0o644)
require.NoError(t, err)
require.NoError(tb, err)
return targetPath
}

View File

@@ -34,16 +34,16 @@ func HostToIPs(host string) (ipv4, ipv6 netip.Addr) {
// StartHTTPServer is a helper that starts the HTTP server, which is configured
// to return data on every request, and returns the client and server URL.
func StartHTTPServer(t testing.TB, data []byte) (c *http.Client, u *url.URL) {
t.Helper()
func StartHTTPServer(tb testing.TB, data []byte) (c *http.Client, u *url.URL) {
tb.Helper()
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
_, _ = w.Write(data)
}))
t.Cleanup(srv.Close)
tb.Cleanup(srv.Close)
u, err := url.Parse(srv.URL)
require.NoError(t, err)
require.NoError(tb, err)
return srv.Client(), u
}
@@ -55,8 +55,8 @@ const testTimeout = 1 * time.Second
// StartLocalhostUpstream is a test helper that starts a DNS server on
// localhost.
func StartLocalhostUpstream(t *testing.T, h dns.Handler) (addr *url.URL) {
t.Helper()
func StartLocalhostUpstream(tb *testing.T, h dns.Handler) (addr *url.URL) {
tb.Helper()
startCh := make(chan netip.AddrPort)
defer close(startCh)
@@ -83,12 +83,12 @@ func StartLocalhostUpstream(t *testing.T, h dns.Handler) (addr *url.URL) {
Host: addrPort.String(),
}
testutil.CleanupAndRequireSuccess(t, func() (err error) { return <-errCh })
testutil.CleanupAndRequireSuccess(t, srv.Shutdown)
testutil.CleanupAndRequireSuccess(tb, func() (err error) { return <-errCh })
testutil.CleanupAndRequireSuccess(tb, srv.Shutdown)
case err := <-errCh:
require.NoError(t, err)
require.NoError(tb, err)
case <-time.After(testTimeout):
require.FailNow(t, "timeout exceeded")
require.FailNow(tb, "timeout exceeded")
}
return addr

View File

@@ -2,45 +2,37 @@ package aghtest
import (
"context"
"net"
"net/netip"
"time"
"github.com/AdguardTeam/AdGuardHome/internal/agh"
"github.com/AdguardTeam/AdGuardHome/internal/aghos"
"github.com/AdguardTeam/AdGuardHome/internal/next/agh"
nextagh "github.com/AdguardTeam/AdGuardHome/internal/next/agh"
"github.com/AdguardTeam/AdGuardHome/internal/rdns"
"github.com/AdguardTeam/AdGuardHome/internal/whois"
"github.com/AdguardTeam/dnsproxy/upstream"
"github.com/miekg/dns"
)
// Interface Mocks
//
// Keep entities in this file in alphabetic order.
// Module adguard-home
// Package aghos
// FSWatcher is a fake [aghos.FSWatcher] implementation for tests.
type FSWatcher struct {
OnStart func() (err error)
OnClose func() (err error)
OnEvents func() (e <-chan struct{})
OnAdd func(name string) (err error)
OnStart func(ctx context.Context) (err error)
OnShutdown func(ctx context.Context) (err error)
OnEvents func() (e <-chan struct{})
OnAdd func(name string) (err error)
}
// type check
var _ aghos.FSWatcher = (*FSWatcher)(nil)
// Start implements the [aghos.FSWatcher] interface for *FSWatcher.
func (w *FSWatcher) Start() (err error) {
return w.OnStart()
func (w *FSWatcher) Start(ctx context.Context) (err error) {
return w.OnStart(ctx)
}
// Close implements the [aghos.FSWatcher] interface for *FSWatcher.
func (w *FSWatcher) Close() (err error) {
return w.OnClose()
// Shutdown implements the [aghos.FSWatcher] interface for *FSWatcher.
func (w *FSWatcher) Shutdown(ctx context.Context) (err error) {
return w.OnShutdown(ctx)
}
// Events implements the [aghos.FSWatcher] interface for *FSWatcher.
@@ -53,9 +45,8 @@ func (w *FSWatcher) Add(name string) (err error) {
return w.OnAdd(name)
}
// Package agh
// ServiceWithConfig is a fake [agh.ServiceWithConfig] implementation for tests.
// ServiceWithConfig is a fake [nextagh.ServiceWithConfig] implementation for
// tests.
type ServiceWithConfig[ConfigType any] struct {
OnStart func(ctx context.Context) (err error)
OnShutdown func(ctx context.Context) (err error)
@@ -63,28 +54,26 @@ type ServiceWithConfig[ConfigType any] struct {
}
// type check
var _ agh.ServiceWithConfig[struct{}] = (*ServiceWithConfig[struct{}])(nil)
var _ nextagh.ServiceWithConfig[struct{}] = (*ServiceWithConfig[struct{}])(nil)
// Start implements the [agh.ServiceWithConfig] interface for
// Start implements the [nextagh.ServiceWithConfig] interface for
// *ServiceWithConfig.
func (s *ServiceWithConfig[_]) Start(ctx context.Context) (err error) {
return s.OnStart(ctx)
}
// Shutdown implements the [agh.ServiceWithConfig] interface for
// Shutdown implements the [nextagh.ServiceWithConfig] interface for
// *ServiceWithConfig.
func (s *ServiceWithConfig[_]) Shutdown(ctx context.Context) (err error) {
return s.OnShutdown(ctx)
}
// Config implements the [agh.ServiceWithConfig] interface for
// Config implements the [nextagh.ServiceWithConfig] interface for
// *ServiceWithConfig.
func (s *ServiceWithConfig[ConfigType]) Config() (c ConfigType) {
return s.OnConfig()
}
// Package client
// AddressProcessor is a fake [client.AddressProcessor] implementation for
// tests.
type AddressProcessor struct {
@@ -120,20 +109,6 @@ func (p *AddressUpdater) UpdateAddress(
p.OnUpdateAddress(ctx, ip, host, info)
}
// Package filtering
// Resolver is a fake [filtering.Resolver] implementation for tests.
type Resolver struct {
OnLookupIP func(ctx context.Context, network, host string) (ips []net.IP, err error)
}
// LookupIP implements the [filtering.Resolver] interface for *Resolver.
func (r *Resolver) LookupIP(ctx context.Context, network, host string) (ips []net.IP, err error) {
return r.OnLookupIP(ctx, network, host)
}
// Package rdns
// Exchanger is a fake [rdns.Exchanger] implementation for tests.
type Exchanger struct {
OnExchange func(ip netip.Addr) (host string, ttl time.Duration, err error)
@@ -147,10 +122,6 @@ func (e *Exchanger) Exchange(ip netip.Addr) (host string, ttl time.Duration, err
return e.OnExchange(ip)
}
// Module dnsproxy
// Package upstream
// UpstreamMock is a fake [upstream.Upstream] implementation for tests.
//
// TODO(a.garipov): Replace with all uses of Upstream with UpstreamMock and
@@ -178,3 +149,16 @@ func (u *UpstreamMock) Exchange(req *dns.Msg) (resp *dns.Msg, err error) {
func (u *UpstreamMock) Close() (err error) {
return u.OnClose()
}
// ConfigModifier is a fake [agh.ConfigModifier] implementation for tests.
type ConfigModifier struct {
OnApply func(ctx context.Context)
}
// type check
var _ agh.ConfigModifier = (*ConfigModifier)(nil)
// Apply implements the [ConfigModifier] interface for *ConfigModifier.
func (m *ConfigModifier) Apply(ctx context.Context) {
m.OnApply(ctx)
}

View File

@@ -3,20 +3,12 @@ package aghtest_test
import (
"github.com/AdguardTeam/AdGuardHome/internal/aghtest"
"github.com/AdguardTeam/AdGuardHome/internal/client"
"github.com/AdguardTeam/AdGuardHome/internal/filtering"
)
// Put interface checks that cause import cycles here.
// type check
var _ filtering.Resolver = (*aghtest.Resolver)(nil)
// type check
//
// TODO(s.chzhen): It's here to avoid the import cycle. Remove it.
var _ client.AddressProcessor = (*aghtest.AddressProcessor)(nil)
// type check
//
// TODO(s.chzhen): It's here to avoid the import cycle. Remove it.
var _ client.AddressUpdater = (*aghtest.AddressUpdater)(nil)
// TODO(s.chzhen): Resolve the import cycles and move it to aghtest.
var (
_ client.AddressProcessor = (*aghtest.AddressProcessor)(nil)
_ client.AddressUpdater = (*aghtest.AddressUpdater)(nil)
)

View File

@@ -220,7 +220,7 @@ func NewErrorUpstream() (u *UpstreamMock) {
return &UpstreamMock{
OnAddress: func() (addr string) { return "error.upstream.example" },
OnExchange: func(_ *dns.Msg) (resp *dns.Msg, err error) {
return nil, errors.Error("test upstream error")
return nil, ErrUpstream
},
OnClose: func() (err error) { return nil },
}

View File

@@ -2,26 +2,29 @@
package aghtls
import (
"context"
"crypto/tls"
"crypto/x509"
"fmt"
"log/slog"
"slices"
"github.com/AdguardTeam/golibs/log"
"github.com/AdguardTeam/golibs/netutil"
)
// init makes sure that the cipher name map is filled.
// Init populates the cipherSuites map with the name-to-ID mapping of cipher
// suites from crypto/tls. It must be called only once, and it must be called
// before any function that calls [ParseCiphers].
//
// TODO(a.garipov): Propose a similar API to crypto/tls.
func init() {
func Init(ctx context.Context, l *slog.Logger) {
suites := tls.CipherSuites()
cipherSuites = make(map[string]uint16, len(suites))
for _, s := range suites {
cipherSuites[s.Name] = s.ID
}
log.Debug("tls: known ciphers: %q", cipherSuites)
l.DebugContext(ctx, "known ciphers", "ciphers", cipherSuites)
}
// cipherSuites are a name-to-ID mapping of cipher suites from crypto/tls. It

View File

@@ -3,17 +3,20 @@ package aghtls_test
import (
"crypto/tls"
"testing"
"time"
"github.com/AdguardTeam/AdGuardHome/internal/aghtls"
"github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/AdguardTeam/golibs/testutil"
"github.com/stretchr/testify/assert"
)
func TestMain(m *testing.M) {
testutil.DiscardLogOutput(m)
}
// testTimeout is a common timeout for tests and contexts.
const testTimeout time.Duration = 1 * time.Second
func TestParseCiphers(t *testing.T) {
aghtls.Init(testutil.ContextWithTimeout(t, testTimeout), slogutil.NewDiscardLogger())
testCases := []struct {
name string
wantErrMsg string

View File

@@ -1,7 +1,9 @@
package aghtls
import (
"context"
"crypto/x509"
"log/slog"
)
// SystemRootCAs tries to load root certificates from the operating system. It
@@ -9,6 +11,6 @@ import (
// default algorithm to find system root CA list.
//
// See https://github.com/AdguardTeam/AdGuardHome/issues/1311.
func SystemRootCAs() (roots *x509.CertPool) {
return rootCAs()
func SystemRootCAs(ctx context.Context, l *slog.Logger) (roots *x509.CertPool) {
return rootCAs(ctx, l)
}

View File

@@ -3,15 +3,17 @@
package aghtls
import (
"context"
"crypto/x509"
"log/slog"
"os"
"path/filepath"
"github.com/AdguardTeam/golibs/errors"
"github.com/AdguardTeam/golibs/log"
"github.com/AdguardTeam/golibs/logutil/slogutil"
)
func rootCAs() (roots *x509.CertPool) {
func rootCAs(ctx context.Context, l *slog.Logger) (roots *x509.CertPool) {
// Directories with the system root certificates, which aren't supported by
// Go's crypto/x509.
dirs := []string{
@@ -21,36 +23,51 @@ func rootCAs() (roots *x509.CertPool) {
roots = x509.NewCertPool()
for _, dir := range dirs {
dirEnts, err := os.ReadDir(dir)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
continue
}
// TODO(a.garipov): Improve error handling here and in other places.
log.Error("aghtls: opening directory %q: %s", dir, err)
}
var rootsAdded bool
for _, de := range dirEnts {
var certData []byte
rootFile := filepath.Join(dir, de.Name())
certData, err = os.ReadFile(rootFile)
if err != nil {
log.Error("aghtls: reading root cert: %s", err)
} else {
if roots.AppendCertsFromPEM(certData) {
rootsAdded = true
} else {
log.Error("aghtls: could not add root from %q", rootFile)
}
}
}
if rootsAdded {
if addCertsFromDir(ctx, l, roots, dir) {
return roots
}
}
return nil
}
// addCertsFromDir appends all readable PEM files from dir to pool. It returns
// true if at least one certificate was accepted.
func addCertsFromDir(
ctx context.Context,
l *slog.Logger,
pool *x509.CertPool,
dir string,
) (ok bool) {
dirEnts, err := os.ReadDir(dir)
if err != nil {
if !errors.Is(err, os.ErrNotExist) {
// TODO(a.garipov): Improve error handling here and in other places.
l.ErrorContext(ctx, "opening directory", slogutil.KeyError, err)
}
return false
}
var rootsAdded bool
for _, de := range dirEnts {
var certData []byte
rootFile := filepath.Join(dir, de.Name())
certData, err = os.ReadFile(rootFile)
if err != nil {
l.ErrorContext(ctx, "reading root cert", slogutil.KeyError, err)
continue
}
if !pool.AppendCertsFromPEM(certData) {
l.ErrorContext(ctx, "adding root cert", "file", rootFile, slogutil.KeyError, err)
continue
}
rootsAdded = true
}
return rootsAdded
}

View File

@@ -2,8 +2,12 @@
package aghtls
import "crypto/x509"
import (
"context"
"crypto/x509"
"log/slog"
)
func rootCAs() (roots *x509.CertPool) {
func rootCAs(_ context.Context, _ *slog.Logger) (roots *x509.CertPool) {
return nil
}

View File

@@ -24,12 +24,12 @@ var testdata fs.FS = os.DirFS("./testdata")
type RunCmdFunc func(cmd string, args ...string) (code int, out []byte, err error)
// substShell replaces the the aghos.RunCommand function used throughout the
// package with rc for tests ran under t.
func substShell(t testing.TB, rc RunCmdFunc) {
t.Helper()
// package with rc for tests ran under tb.
func substShell(tb testing.TB, rc RunCmdFunc) {
tb.Helper()
prev := aghosRunCommand
t.Cleanup(func() { aghosRunCommand = prev })
tb.Cleanup(func() { aghosRunCommand = prev })
aghosRunCommand = rc
}

View File

@@ -1208,12 +1208,8 @@ func TestStorage_CustomUpstreamConfig(t *testing.T) {
}
dhcp := &testDHCP{
OnLeases: func() (ls []*dhcpsvc.Lease) {
panic("not implemented")
},
OnHostBy: func(ip netip.Addr) (host string) {
panic("not implemented")
},
OnLeases: func() (_ []*dhcpsvc.Lease) { panic(testutil.UnexpectedCall()) },
OnHostBy: func(ip netip.Addr) (_ string) { panic(testutil.UnexpectedCall(ip)) },
OnMACBy: func(ip netip.Addr) (mac net.HardwareAddr) {
return ipToMAC[ip]
},

View File

@@ -2,4 +2,4 @@
package configmigrate
// LastSchemaVersion is the most recent schema version.
const LastSchemaVersion uint = 29
const LastSchemaVersion uint = 30

View File

@@ -242,8 +242,8 @@ func TestUpgradeSchema8to9(t *testing.T) {
}
// assertEqualExcept removes entries from configs and compares them.
func assertEqualExcept(t *testing.T, oldConf, newConf yobj, oldKeys, newKeys []string) {
t.Helper()
func assertEqualExcept(tb testing.TB, oldConf, newConf yobj, oldKeys, newKeys []string) {
tb.Helper()
for _, k := range oldKeys {
delete(oldConf, k)
@@ -252,7 +252,7 @@ func assertEqualExcept(t *testing.T, oldConf, newConf yobj, oldKeys, newKeys []s
delete(newConf, k)
}
assert.Equal(t, oldConf, newConf)
assert.Equal(tb, oldConf, newConf)
}
func testDiskConf(schemaVersion int) (diskConf yobj) {

View File

@@ -125,6 +125,7 @@ func (m *Migrator) upgradeConfigSchema(current, target uint, diskConf yobj) (err
26: migrateTo27,
27: migrateTo28,
28: m.migrateTo29,
29: m.migrateTo30,
}
for i, migrate := range upgrades[current:target] {

View File

@@ -54,6 +54,8 @@ func getField[T any](t require.TestingT, obj any, indexes ...any) (val T) {
}
func TestMigrateConfig_Migrate(t *testing.T) {
t.Parallel()
const (
inputFileName = "input.yml"
outputFileName = "output.yml"
@@ -193,10 +195,16 @@ func TestMigrateConfig_Migrate(t *testing.T) {
yamlEqFunc: require.YAMLEq,
name: "v27",
targetVersion: 27,
}, {
yamlEqFunc: require.YAMLEq,
name: "v30",
targetVersion: 30,
}}
for _, tc := range testCases {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
body, err := fs.ReadFile(testdata, path.Join(t.Name(), inputFileName))
require.NoError(t, err)

View File

@@ -0,0 +1,119 @@
http:
address: 127.0.0.1:3000
session_ttl: 3h
pprof:
enabled: true
port: 6060
users:
- name: testuser
password: testpassword
dns:
bind_hosts:
- 127.0.0.1
port: 53
parental_sensitivity: 0
upstream_dns:
- tls://1.1.1.1
- tls://1.0.0.1
- quic://8.8.8.8:784
bootstrap_dns:
- 8.8.8.8:53
cache_size: 4194304
edns_client_subnet:
enabled: true
use_custom: false
custom_ip: ""
filtering:
filtering_enabled: true
parental_enabled: false
safebrowsing_enabled: false
safe_fs_patterns: []
safe_search:
enabled: false
bing: true
duckduckgo: true
google: true
pixabay: true
yandex: true
youtube: true
protection_enabled: true
blocked_services:
schedule:
time_zone: Local
ids:
- 500px
blocked_response_ttl: 10
filters:
- url: https://adaway.org/hosts.txt
name: AdAway
enabled: false
- url: FILEPATH
name: Local Filter
enabled: false
clients:
persistent:
- name: localhost
ids:
- 127.0.0.1
- aa:aa:aa:aa:aa:aa
use_global_settings: true
use_global_blocked_services: true
filtering_enabled: false
parental_enabled: false
safebrowsing_enabled: false
safe_search:
enabled: true
bing: true
duckduckgo: true
google: true
pixabay: true
yandex: true
youtube: true
blocked_services:
schedule:
time_zone: Local
ids:
- 500px
runtime_sources:
whois: true
arp: true
rdns: true
dhcp: true
hosts: true
dhcp:
enabled: false
interface_name: vboxnet0
local_domain_name: local
dhcpv4:
gateway_ip: 192.168.0.1
subnet_mask: 255.255.255.0
range_start: 192.168.0.10
range_end: 192.168.0.250
lease_duration: 1234
icmp_timeout_msec: 10
schema_version: 29
user_rules: []
querylog:
enabled: true
file_enabled: true
interval: 720h
size_memory: 1000
ignored:
- '|.^'
statistics:
enabled: true
interval: 240h
ignored:
- '|.^'
os:
group: ''
rlimit_nofile: 123
user: ''
log:
file: ""
max_backups: 0
max_size: 100
max_age: 3
compress: true
local_time: false
verbose: true

View File

@@ -0,0 +1,120 @@
http:
address: 127.0.0.1:3000
session_ttl: 3h
pprof:
enabled: true
port: 6060
users:
- name: testuser
password: testpassword
dns:
bind_hosts:
- 127.0.0.1
port: 53
parental_sensitivity: 0
upstream_dns:
- tls://1.1.1.1
- tls://1.0.0.1
- quic://8.8.8.8:784
bootstrap_dns:
- 8.8.8.8:53
cache_enabled: true
cache_size: 4194304
edns_client_subnet:
enabled: true
use_custom: false
custom_ip: ""
filtering:
filtering_enabled: true
parental_enabled: false
safebrowsing_enabled: false
safe_fs_patterns: []
safe_search:
enabled: false
bing: true
duckduckgo: true
google: true
pixabay: true
yandex: true
youtube: true
protection_enabled: true
blocked_services:
schedule:
time_zone: Local
ids:
- 500px
blocked_response_ttl: 10
filters:
- url: https://adaway.org/hosts.txt
name: AdAway
enabled: false
- url: FILEPATH
name: Local Filter
enabled: false
clients:
persistent:
- name: localhost
ids:
- 127.0.0.1
- aa:aa:aa:aa:aa:aa
use_global_settings: true
use_global_blocked_services: true
filtering_enabled: false
parental_enabled: false
safebrowsing_enabled: false
safe_search:
enabled: true
bing: true
duckduckgo: true
google: true
pixabay: true
yandex: true
youtube: true
blocked_services:
schedule:
time_zone: Local
ids:
- 500px
runtime_sources:
whois: true
arp: true
rdns: true
dhcp: true
hosts: true
dhcp:
enabled: false
interface_name: vboxnet0
local_domain_name: local
dhcpv4:
gateway_ip: 192.168.0.1
subnet_mask: 255.255.255.0
range_start: 192.168.0.10
range_end: 192.168.0.250
lease_duration: 1234
icmp_timeout_msec: 10
schema_version: 30
user_rules: []
querylog:
enabled: true
file_enabled: true
interval: 720h
size_memory: 1000
ignored:
- '|.^'
statistics:
enabled: true
interval: 240h
ignored:
- '|.^'
os:
group: ''
rlimit_nofile: 123
user: ''
log:
file: ""
max_backups: 0
max_size: 100
max_age: 3
compress: true
local_time: false
verbose: true

View File

@@ -0,0 +1,33 @@
package configmigrate
// migrateTo30 performs the following changes:
//
// # BEFORE:
// 'dns':
// 'cache_size': 123456
// # …
//
// # AFTER:
// 'dns':
// 'cache_size': 123456
// 'cache_enabled': true
// # …
//
// If cache_size is zero, then cache_enabled should be false.
func (m Migrator) migrateTo30(diskConf yobj) (err error) {
diskConf["schema_version"] = 30
dnsConf, ok, err := fieldVal[yobj](diskConf, "dns")
if !ok {
return err
}
cacheSize, ok, err := fieldVal[int](dnsConf, "cache_size")
if !ok {
return err
}
dnsConf["cache_enabled"] = cacheSize > 0
return nil
}

View File

@@ -6,6 +6,7 @@ import (
"net/netip"
"time"
"github.com/AdguardTeam/AdGuardHome/internal/agh"
"github.com/AdguardTeam/AdGuardHome/internal/aghhttp"
"github.com/AdguardTeam/AdGuardHome/internal/aghnet"
"github.com/AdguardTeam/AdGuardHome/internal/dhcpsvc"
@@ -15,8 +16,9 @@ import (
// ServerConfig is the configuration for the DHCP server. The order of YAML
// fields is important, since the YAML configuration file follows it.
type ServerConfig struct {
// Called when the configuration is changed by HTTP request
ConfigModified func() `yaml:"-"`
// ConfModifier is used to update the global configuration. It must not be
// nil.
ConfModifier agh.ConfigModifier `yaml:"-"`
// Register an HTTP handler
HTTPRegister aghhttp.RegisterFunc `yaml:"-"`

View File

@@ -107,7 +107,7 @@ var _ Interface = (*server)(nil)
func Create(conf *ServerConfig) (s *server, err error) {
s = &server{
conf: &ServerConfig{
ConfigModified: conf.ConfigModified,
ConfModifier: conf.ConfModifier,
HTTPRegister: conf.HTTPRegister,

View File

@@ -335,7 +335,7 @@ func (s *server) handleDHCPSetConfig(w http.ResponseWriter, r *http.Request) {
}
s.setConfFromJSON(conf, srv4, srv6)
s.conf.ConfigModified()
s.conf.ConfModifier.Apply(r.Context())
err = s.dbLoad()
if err != nil {
@@ -679,7 +679,7 @@ func (s *server) handleReset(w http.ResponseWriter, r *http.Request) {
}
s.conf = &ServerConfig{
ConfigModified: s.conf.ConfigModified,
ConfModifier: s.conf.ConfModifier,
HTTPRegister: s.conf.HTTPRegister,
@@ -702,7 +702,7 @@ func (s *server) handleReset(w http.ResponseWriter, r *http.Request) {
}
s.srv6, _ = v6Create(v6conf)
s.conf.ConfigModified()
s.conf.ConfModifier.Apply(r.Context())
}
func (s *server) handleResetLeases(w http.ResponseWriter, r *http.Request) {

View File

@@ -10,6 +10,7 @@ import (
"net/netip"
"testing"
"github.com/AdguardTeam/AdGuardHome/internal/agh"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
@@ -33,18 +34,18 @@ func defaultResponse() *dhcpStatusResponse {
// handleLease is the helper function that calls handler with provided static
// lease as body and returns modified response recorder.
func handleLease(t *testing.T, lease *leaseStatic, handler http.HandlerFunc) (w *httptest.ResponseRecorder) {
t.Helper()
func handleLease(tb testing.TB, lease *leaseStatic, handler http.HandlerFunc) (w *httptest.ResponseRecorder) {
tb.Helper()
w = httptest.NewRecorder()
b := &bytes.Buffer{}
err := json.NewEncoder(b).Encode(lease)
require.NoError(t, err)
require.NoError(tb, err)
var r *http.Request
r, err = http.NewRequest(http.MethodPost, "", b)
require.NoError(t, err)
require.NoError(tb, err)
handler(w, r)
@@ -84,10 +85,10 @@ func TestServer_handleDHCPStatus(t *testing.T) {
}
s, err := Create(&ServerConfig{
Enabled: true,
Conf4: *defaultV4ServerConf(),
DataDir: t.TempDir(),
ConfigModified: func() {},
Enabled: true,
Conf4: *defaultV4ServerConf(),
DataDir: t.TempDir(),
ConfModifier: agh.EmptyConfigModifier{},
})
require.NoError(t, err)
@@ -178,11 +179,11 @@ func TestServer_HandleUpdateStaticLease(t *testing.T) {
}
s, err := Create(&ServerConfig{
Enabled: true,
Conf4: *defaultV4ServerConf(),
Conf6: V6ServerConf{},
DataDir: t.TempDir(),
ConfigModified: func() {},
Enabled: true,
Conf4: *defaultV4ServerConf(),
Conf6: V6ServerConf{},
DataDir: t.TempDir(),
ConfModifier: agh.EmptyConfigModifier{},
})
require.NoError(t, err)
@@ -266,11 +267,11 @@ func TestServer_HandleUpdateStaticLease_validation(t *testing.T) {
}}
s, err := Create(&ServerConfig{
Enabled: true,
Conf4: *defaultV4ServerConf(),
Conf6: V6ServerConf{},
DataDir: t.TempDir(),
ConfigModified: func() {},
Enabled: true,
Conf4: *defaultV4ServerConf(),
Conf6: V6ServerConf{},
DataDir: t.TempDir(),
ConfModifier: agh.EmptyConfigModifier{},
})
require.NoError(t, err)

View File

@@ -44,12 +44,12 @@ func defaultV4ServerConf() (conf *V4ServerConf) {
// defaultSrv prepares the default DHCPServer to use in tests. The underlying
// type of s is *v4Server.
func defaultSrv(t *testing.T) (s DHCPServer) {
t.Helper()
func defaultSrv(tb testing.TB) (s DHCPServer) {
tb.Helper()
var err error
s, err = v4Create(defaultV4ServerConf())
require.NoError(t, err)
require.NoError(tb, err)
return s
}

View File

@@ -12,7 +12,6 @@ import (
"github.com/AdguardTeam/AdGuardHome/internal/aghhttp"
"github.com/AdguardTeam/AdGuardHome/internal/client"
"github.com/AdguardTeam/golibs/container"
"github.com/AdguardTeam/golibs/log"
"github.com/AdguardTeam/golibs/stringutil"
"github.com/AdguardTeam/urlfilter"
"github.com/AdguardTeam/urlfilter/filterlist"
@@ -230,6 +229,8 @@ func validateStrUniq(clients []string) (uc aghalg.UniqChecker[string], err error
// handleAccessSet handles requests to the POST /control/access/set endpoint.
func (s *Server) handleAccessSet(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
list := &accessListJSON{}
err := json.NewDecoder(r.Body).Decode(&list)
if err != nil {
@@ -253,14 +254,15 @@ func (s *Server) handleAccessSet(w http.ResponseWriter, r *http.Request) {
return
}
defer log.Debug(
"access: updated lists: %d, %d, %d",
len(list.AllowedClients),
len(list.DisallowedClients),
len(list.BlockedHosts),
defer s.logger.DebugContext(
ctx,
"updated access lists",
"allowed", len(list.AllowedClients),
"disallowed", len(list.DisallowedClients),
"blocked_hosts", len(list.BlockedHosts),
)
defer s.conf.ConfigModified()
defer s.conf.ConfModifier.Apply(ctx)
s.serverLock.Lock()
defer s.serverLock.Unlock()

View File

@@ -1,13 +1,13 @@
package dnsforward
import (
"context"
"encoding/binary"
"fmt"
"github.com/AdguardTeam/AdGuardHome/internal/aghnet"
"github.com/AdguardTeam/dnsproxy/proxy"
"github.com/AdguardTeam/golibs/errors"
"github.com/AdguardTeam/golibs/log"
"github.com/miekg/dns"
)
@@ -41,7 +41,13 @@ func (s *Server) HandleBefore(
qt := q.Qtype
host := aghnet.NormalizeDomain(q.Name)
if s.access.isBlockedHost(host, qt) {
log.Debug("access: request %s %s is in access blocklist", dns.Type(qt), host)
// TODO(s.chzhen): Pass context.
s.logger.DebugContext(
context.TODO(),
"request is in access blocklist",
"dns_type", dns.Type(qt),
"host", host,
)
return s.preBlockedResponse(pctx)
}

View File

@@ -9,6 +9,7 @@ import (
"github.com/AdguardTeam/AdGuardHome/internal/aghtest"
"github.com/AdguardTeam/AdGuardHome/internal/filtering"
"github.com/AdguardTeam/dnsproxy/proxy"
"github.com/AdguardTeam/golibs/testutil"
"github.com/miekg/dns"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -132,7 +133,7 @@ func TestServer_HandleBefore_tls(t *testing.T) {
s.conf.DisallowedClients = tc.disallowedClients
s.conf.BlockedHosts = tc.blockedHosts
err := s.Prepare(&s.conf)
err := s.Prepare(testutil.ContextWithTimeout(t, testTimeout), &s.conf)
require.NoError(t, err)
startDeferStop(t, s)

View File

@@ -201,6 +201,7 @@ func TestServer_clientIDFromDNSContext(t *testing.T) {
srv := &Server{
conf: ServerConfig{TLSConf: tlsConf},
baseLogger: testLogger,
logger: testLogger,
}
var (

View File

@@ -1,6 +1,7 @@
package dnsforward
import (
"context"
"crypto/tls"
"crypto/x509"
"fmt"
@@ -11,6 +12,7 @@ import (
"strings"
"time"
"github.com/AdguardTeam/AdGuardHome/internal/agh"
"github.com/AdguardTeam/AdGuardHome/internal/aghhttp"
"github.com/AdguardTeam/AdGuardHome/internal/aghnet"
"github.com/AdguardTeam/AdGuardHome/internal/aghslog"
@@ -25,6 +27,7 @@ import (
"github.com/AdguardTeam/golibs/netutil"
"github.com/AdguardTeam/golibs/stringutil"
"github.com/AdguardTeam/golibs/timeutil"
"github.com/AdguardTeam/golibs/validate"
"github.com/ameshkov/dnscrypt/v2"
)
@@ -102,6 +105,9 @@ type Config struct {
// DNS cache settings
// CacheEnabled defines if the DNS cache should be used.
CacheEnabled bool `yaml:"cache_enabled"`
// CacheSize is the DNS cache size (in bytes).
CacheSize uint32 `yaml:"cache_size"`
@@ -259,8 +265,9 @@ type ServerConfig struct {
// TLSCiphers are the IDs of TLS cipher suites to use.
TLSCiphers []uint16
// Called when the configuration is changed by HTTP request
ConfigModified func()
// ConfModifier is used to update the global configuration. It must not be
// nil.
ConfModifier agh.ConfigModifier
// Register an HTTP handler
HTTPRegister aghhttp.RegisterFunc
@@ -307,7 +314,7 @@ const (
)
// newProxyConfig creates and validates configuration for the main proxy.
func (s *Server) newProxyConfig() (conf *proxy.Config, err error) {
func (s *Server) newProxyConfig(ctx context.Context) (conf *proxy.Config, err error) {
srvConf := s.conf
trustedPrefixes := netutil.UnembedPrefixes(srvConf.TrustedProxies)
@@ -355,17 +362,18 @@ func (s *Server) newProxyConfig() (conf *proxy.Config, err error) {
return nil, fmt.Errorf("bogus_nxdomain: %w", err)
}
err = s.prepareTLS(conf)
err = s.prepareTLS(ctx, conf)
if err != nil {
return nil, fmt.Errorf("validating tls: %w", err)
}
err = s.preparePlain(conf)
err = s.preparePlain(ctx, conf)
if err != nil {
return nil, fmt.Errorf("validating plain: %w", err)
}
conf, err = prepareCacheConfig(conf,
srvConf.CacheEnabled,
srvConf.CacheSize,
srvConf.CacheMinTTL,
srvConf.CacheMaxTTL,
@@ -382,13 +390,20 @@ func (s *Server) newProxyConfig() (conf *proxy.Config, err error) {
// there is one.
func prepareCacheConfig(
conf *proxy.Config,
isEnabled bool,
size uint32,
minTTL uint32,
maxTTL uint32,
) (prepared *proxy.Config, err error) {
if size != 0 {
if isEnabled {
cacheSize := int(size)
err = validate.Positive("cache_size", cacheSize)
if err != nil {
return nil, fmt.Errorf("cache_enabled is true: %w", err)
}
conf.CacheEnabled = true
conf.CacheSizeBytes = int(size)
conf.CacheSizeBytes = cacheSize
}
err = validateCacheTTL(minTTL, maxTTL)
@@ -444,7 +459,7 @@ func (s *Server) initDefaultSettings() {
// prepareIpsetListSettings reads and prepares the ipset configuration either
// from a file or from the data in the configuration file.
func (s *Server) prepareIpsetListSettings() (ipsets []string, err error) {
func (s *Server) prepareIpsetListSettings(ctx context.Context) (ipsets []string, err error) {
fn := s.conf.IpsetListFileName
if fn == "" {
return s.conf.IpsetList, nil
@@ -459,7 +474,7 @@ func (s *Server) prepareIpsetListSettings() (ipsets []string, err error) {
ipsets = stringutil.SplitTrimmed(string(data), "\n")
ipsets = slices.DeleteFunc(ipsets, aghnet.IsCommentOrEmpty)
log.Debug("dns: using %d ipset rules from file %q", len(ipsets), fn)
s.logger.DebugContext(ctx, "using ipset rules from file", "num", len(ipsets), "file", fn)
return ipsets, nil
}
@@ -629,7 +644,7 @@ func (s *Server) prepareDNSCrypt(proxyConf *proxy.Config) {
}
// prepareTLS sets up the TLS configuration for the DNS proxy.
func (s *Server) prepareTLS(proxyConf *proxy.Config) (err error) {
func (s *Server) prepareTLS(ctx context.Context, proxyConf *proxy.Config) (err error) {
s.prepareDNSCrypt(proxyConf)
if s.conf.TLSConf.Cert == nil {
@@ -653,11 +668,20 @@ func (s *Server) prepareTLS(proxyConf *proxy.Config) (err error) {
if s.conf.TLSConf.StrictSNICheck {
if len(cert.DNSNames) != 0 {
s.dnsNames = cert.DNSNames
log.Debug("dns: using certificate's SAN as DNS names: %v", cert.DNSNames)
s.logger.DebugContext(
ctx,
"using certificate's SAN as DNS names",
"dns_names", cert.DNSNames,
)
slices.Sort(s.dnsNames)
} else {
s.dnsNames = []string{cert.Subject.CommonName}
log.Debug("dns: using certificate's CN as DNS name: %s", cert.Subject.CommonName)
s.logger.DebugContext(
ctx,
"using certificate's CN as DNS name",
"common_name",
cert.Subject.CommonName,
)
}
}
@@ -706,15 +730,22 @@ func anyNameMatches(dnsNames []string, sni string) (ok bool) {
// If the server name (from SNI) supplied by client is incorrect - we terminate the ongoing TLS handshake.
func (s *Server) onGetCertificate(ch *tls.ClientHelloInfo) (*tls.Certificate, error) {
if s.conf.TLSConf.StrictSNICheck && !anyNameMatches(s.dnsNames, ch.ServerName) {
log.Info("dns: tls: unknown SNI in Client Hello: %s", ch.ServerName)
// TODO(s.chzhen): Pass context.
s.logger.WarnContext(
context.TODO(),
"unknown SNI in Client Hello",
"server_name", ch.ServerName,
)
return nil, fmt.Errorf("invalid SNI")
}
return s.conf.TLSConf.Cert, nil
}
// preparePlain prepares the plain-DNS configuration for the DNS proxy.
// preparePlain assumes that prepareTLS has already been called.
func (s *Server) preparePlain(proxyConf *proxy.Config) (err error) {
func (s *Server) preparePlain(ctx context.Context, proxyConf *proxy.Config) (err error) {
if s.conf.ServePlainDNS {
proxyConf.UDPListenAddr = s.conf.UDPListenAddrs
proxyConf.TCPListenAddr = s.conf.TCPListenAddrs
@@ -732,14 +763,16 @@ func (s *Server) preparePlain(proxyConf *proxy.Config) (err error) {
return errors.Error("disabling plain dns requires at least one encrypted protocol")
}
log.Info("dnsforward: warning: plain dns is disabled")
s.logger.WarnContext(ctx, "plain dns is disabled")
return nil
}
// UpdatedProtectionStatus updates protection state, if the protection was
// disabled temporarily. Returns the updated state of protection.
func (s *Server) UpdatedProtectionStatus() (enabled bool, disabledUntil *time.Time) {
func (s *Server) UpdatedProtectionStatus(
ctx context.Context,
) (enabled bool, disabledUntil *time.Time) {
s.serverLock.RLock()
defer s.serverLock.RUnlock()
@@ -759,7 +792,7 @@ func (s *Server) UpdatedProtectionStatus() (enabled bool, disabledUntil *time.Ti
//
// See https://github.com/AdguardTeam/AdGuardHome/issues/5661.
if s.protectionUpdateInProgress.CompareAndSwap(false, true) {
go s.enableProtectionAfterPause()
go s.enableProtectionAfterPause(ctx)
}
return true, nil
@@ -767,19 +800,19 @@ func (s *Server) UpdatedProtectionStatus() (enabled bool, disabledUntil *time.Ti
// enableProtectionAfterPause sets the protection configuration to enabled
// values. It is intended to be used as a goroutine.
func (s *Server) enableProtectionAfterPause() {
defer log.OnPanic("dns: enabling protection after pause")
func (s *Server) enableProtectionAfterPause(ctx context.Context) {
defer slogutil.RecoverAndLog(ctx, s.logger)
defer s.protectionUpdateInProgress.Store(false)
defer s.conf.ConfigModified()
defer s.conf.ConfModifier.Apply(ctx)
s.serverLock.Lock()
defer s.serverLock.Unlock()
s.dnsFilter.SetProtectionStatus(true, nil)
log.Info("dns: protection is restarted after pause")
s.logger.InfoContext(ctx, "protection is restarted after pause")
}
// validateCacheTTL returns an error if the configuration of the cache TTL

View File

@@ -9,7 +9,6 @@ import (
"time"
"github.com/AdguardTeam/golibs/errors"
"github.com/AdguardTeam/golibs/log"
"github.com/AdguardTeam/golibs/netutil"
)
@@ -17,7 +16,7 @@ import (
// addr should be a valid host:port address, where host could be a domain name
// or an IP address.
func (s *Server) DialContext(ctx context.Context, network, addr string) (conn net.Conn, err error) {
log.Debug("dnsforward: dialing %q for network %q", addr, network)
s.logger.DebugContext(ctx, "dialing", "addr", addr, "network", network)
host, portStr, err := net.SplitHostPort(addr)
if err != nil {
@@ -45,7 +44,7 @@ func (s *Server) DialContext(ctx context.Context, network, addr string) (conn ne
return nil, fmt.Errorf("no addresses for host %q", host)
}
log.Debug("dnsforward: resolved %q: %v", host, ips)
s.logger.DebugContext(ctx, "resolved", "host", host, "ips", ips)
var dialErrs []error
for _, ip := range ips {

View File

@@ -27,16 +27,16 @@ const maxDNS64SynTTL uint32 = 600
// newRR is a helper that creates a new dns.RR with the given name, qtype, ttl
// and value. It fails the test if the qtype is not supported or the type of
// value doesn't match the qtype.
func newRR(t *testing.T, name string, qtype uint16, ttl uint32, val any) (rr dns.RR) {
t.Helper()
func newRR(tb testing.TB, name string, qtype uint16, ttl uint32, val any) (rr dns.RR) {
tb.Helper()
switch qtype {
case dns.TypeA:
rr = &dns.A{A: testutil.RequireTypeAssert[net.IP](t, val)}
rr = &dns.A{A: testutil.RequireTypeAssert[net.IP](tb, val)}
case dns.TypeAAAA:
rr = &dns.AAAA{AAAA: testutil.RequireTypeAssert[net.IP](t, val)}
rr = &dns.AAAA{AAAA: testutil.RequireTypeAssert[net.IP](tb, val)}
case dns.TypeCNAME:
rr = &dns.CNAME{Target: testutil.RequireTypeAssert[string](t, val)}
rr = &dns.CNAME{Target: testutil.RequireTypeAssert[string](tb, val)}
case dns.TypeSOA:
rr = &dns.SOA{
Ns: "ns." + name,
@@ -48,9 +48,9 @@ func newRR(t *testing.T, name string, qtype uint16, ttl uint32, val any) (rr dns
Minttl: 1,
}
case dns.TypePTR:
rr = &dns.PTR{Ptr: testutil.RequireTypeAssert[string](t, val)}
rr = &dns.PTR{Ptr: testutil.RequireTypeAssert[string](tb, val)}
default:
t.Fatalf("unsupported qtype: %d", qtype)
tb.Fatalf("unsupported qtype: %d", qtype)
}
*rr.Header() = dns.RR_Header{
@@ -325,7 +325,7 @@ func TestServer_dns64WithDisabledRDNS(t *testing.T) {
// Shouldn't go to upstream at all.
panicHdlr := dns.HandlerFunc(func(w dns.ResponseWriter, m *dns.Msg) {
panic("not implemented")
panic(testutil.UnexpectedCall(w, m))
})
upsAddr := aghtest.StartLocalhostUpstream(t, panicHdlr).String()
localUpsAddr := aghtest.StartLocalhostUpstream(t, panicHdlr).String()

View File

@@ -145,6 +145,10 @@ type Server struct {
// have a prefix and must not be nil.
baseLogger *slog.Logger
// logger is used to log the operation of the DNS server. It is created
// during initialization in [NewServer].
logger *slog.Logger
// dnsFilter is the DNS filter for filtering client's DNS requests and
// responses.
dnsFilter *filtering.DNSFilter
@@ -254,6 +258,7 @@ func NewServer(p DNSCreateParams) (s *Server, err error) {
queryLog: p.QueryLog,
privateNets: p.PrivateNets,
baseLogger: p.Logger,
logger: p.Logger.With(slogutil.KeyPrefix, "dnsforward"),
// TODO(e.burkov): Use some case-insensitive string comparison.
localDomainSuffix: strings.ToLower(localDomainSuffix),
etcHosts: etcHosts,
@@ -286,7 +291,7 @@ func NewServer(p DNSCreateParams) (s *Server, err error) {
// its workers finished. But it would require the upstream.Upstream to have the
// Close method to prevent from hanging while waiting for unresponsive server to
// respond.
func (s *Server) Close() {
func (s *Server) Close(ctx context.Context) {
s.serverLock.Lock()
defer s.serverLock.Unlock()
@@ -296,7 +301,7 @@ func (s *Server) Close() {
s.dnsProxy = nil
if err := s.ipset.close(); err != nil {
log.Error("dnsforward: closing ipset: %s", err)
s.logger.ErrorContext(ctx, "closing ipset", slogutil.KeyError, err)
}
}
@@ -461,18 +466,17 @@ func hostFromPTR(resp *dns.Msg) (host string, ttl time.Duration, err error) {
}
// Start starts the DNS server. It must only be called after [Server.Prepare].
func (s *Server) Start() error {
func (s *Server) Start(ctx context.Context) error {
s.serverLock.Lock()
defer s.serverLock.Unlock()
return s.startLocked()
return s.startLocked(ctx)
}
// startLocked starts the DNS server without locking. s.serverLock is expected
// to be locked.
func (s *Server) startLocked() error {
// TODO(e.burkov): Use context properly.
err := s.dnsProxy.Start(context.Background())
func (s *Server) startLocked(ctx context.Context) error {
err := s.dnsProxy.Start(ctx)
if err == nil {
s.isRunning = true
}
@@ -482,7 +486,7 @@ func (s *Server) startLocked() error {
// Prepare initializes parameters of s using data from conf. conf must not be
// nil.
func (s *Server) Prepare(conf *ServerConfig) (err error) {
func (s *Server) Prepare(ctx context.Context, conf *ServerConfig) (err error) {
s.conf = *conf
// dnsFilter can be nil during application update.
@@ -496,13 +500,13 @@ func (s *Server) Prepare(conf *ServerConfig) (err error) {
s.initDefaultSettings()
err = s.prepareInternalDNS()
err = s.prepareInternalDNS(ctx)
if err != nil {
// Don't wrap the error, because it's informative enough as is.
return err
}
proxyConfig, err := s.newProxyConfig()
proxyConfig, err := s.newProxyConfig(ctx)
if err != nil {
return fmt.Errorf("preparing proxy: %w", err)
}
@@ -633,8 +637,8 @@ func (s *Server) prepareLocalResolvers() (uc *proxy.UpstreamConfig, err error) {
// prepareInternalDNS initializes the internal state of s before initializing
// the primary DNS proxy instance. It assumes s.serverLock is locked or the
// Server not running.
func (s *Server) prepareInternalDNS() (err error) {
ipsetList, err := s.prepareIpsetListSettings()
func (s *Server) prepareInternalDNS(ctx context.Context) (err error) {
ipsetList, err := s.prepareIpsetListSettings(ctx)
if err != nil {
return fmt.Errorf("preparing ipset settings: %w", err)
}
@@ -779,27 +783,26 @@ func (s *Server) prepareInternalProxy() (err error) {
}
// Stop stops the DNS server.
func (s *Server) Stop() error {
func (s *Server) Stop(ctx context.Context) error {
s.serverLock.Lock()
defer s.serverLock.Unlock()
s.stopLocked()
s.stopLocked(ctx)
return nil
}
// stopLocked stops the DNS server without locking. s.serverLock is expected to
// be locked.
func (s *Server) stopLocked() {
func (s *Server) stopLocked(ctx context.Context) {
// TODO(e.burkov, a.garipov): Return critical errors, not just log them.
// This will require filtering all the non-critical errors in
// [upstream.Upstream] implementations.
if s.dnsProxy != nil {
// TODO(e.burkov): Use context properly.
err := s.dnsProxy.Shutdown(context.Background())
err := s.dnsProxy.Shutdown(ctx)
if err != nil {
log.Error("dnsforward: closing primary resolvers: %s", err)
s.logger.ErrorContext(ctx, "closing primary resolvers", slogutil.KeyError, err)
}
}
@@ -848,14 +851,14 @@ func (s *Server) proxy() (p *proxy.Proxy) {
// Reconfigure applies the new configuration to the DNS server.
//
// TODO(a.garipov): This whole piece of API is weird and needs to be remade.
func (s *Server) Reconfigure(conf *ServerConfig) error {
func (s *Server) Reconfigure(ctx context.Context, conf *ServerConfig) error {
s.serverLock.Lock()
defer s.serverLock.Unlock()
log.Info("dnsforward: starting reconfiguring server")
defer log.Info("dnsforward: finished reconfiguring server")
s.logger.InfoContext(ctx, "starting reconfiguring server")
defer s.logger.InfoContext(ctx, "finished reconfiguring server")
s.stopLocked()
s.stopLocked(ctx)
// It seems that net.Listener.Close() doesn't close file descriptors right away.
// We wait for some time and hope that this fd will be closed.
@@ -864,7 +867,7 @@ func (s *Server) Reconfigure(conf *ServerConfig) error {
if s.addrProc != nil {
err := s.addrProc.Close()
if err != nil {
log.Error("dnsforward: closing address processor: %s", err)
s.logger.ErrorContext(ctx, "closing address processor", slogutil.KeyError, err)
}
}
@@ -874,12 +877,12 @@ func (s *Server) Reconfigure(conf *ServerConfig) error {
// TODO(e.burkov): It seems an error here brings the server down, which is
// not reliable enough.
err := s.Prepare(conf)
err := s.Prepare(ctx, conf)
if err != nil {
return fmt.Errorf("could not reconfigure the server: %w", err)
}
err = s.startLocked()
err = s.startLocked(ctx)
if err != nil {
return fmt.Errorf("could not reconfigure the server: %w", err)
}
@@ -908,16 +911,29 @@ func (s *Server) IsBlockedClient(ip netip.Addr, clientID string) (blocked bool,
allowlistMode := s.access.allowlistMode()
blockedByClientID := s.access.isBlockedClientID(clientID)
// TODO(s.chzhen): Pass context.
ctx := context.TODO()
// Allow if at least one of the checks allows in allowlist mode, but block
// if at least one of the checks blocks in blocklist mode.
if allowlistMode && blockedByIP && blockedByClientID {
log.Debug("dnsforward: client %v (id %q) is not in access allowlist", ip, clientID)
s.logger.DebugContext(
ctx,
"client is not in access allowlist",
"ip", ip,
"client_id", clientID,
)
// Return now without substituting the empty rule for the
// clientID because the rule can't be empty here.
return true, rule
} else if !allowlistMode && (blockedByIP || blockedByClientID) {
log.Debug("dnsforward: client %v (id %q) is in access blocklist", ip, clientID)
s.logger.DebugContext(
ctx,
"client is in access blocklist",
"ip", ip,
"client_id", clientID,
)
blocked = true
}

View File

@@ -2,6 +2,7 @@ package dnsforward
import (
"cmp"
"context"
"crypto/ecdsa"
"crypto/rand"
"crypto/rsa"
@@ -21,6 +22,7 @@ import (
"testing/fstest"
"time"
"github.com/AdguardTeam/AdGuardHome/internal/agh"
"github.com/AdguardTeam/AdGuardHome/internal/aghnet"
"github.com/AdguardTeam/AdGuardHome/internal/aghtest"
"github.com/AdguardTeam/AdGuardHome/internal/client"
@@ -102,12 +104,17 @@ func (c *clientsContainer) ClearUpstreamCache() {
c.OnClearUpstreamCache()
}
func startDeferStop(t *testing.T, s *Server) {
t.Helper()
// startDeferStop starts the server and stops it when the test ends.
//
// TODO(e.burkov): Replace with [servicetest.RequireRun].
func startDeferStop(tb testing.TB, s *Server) {
tb.Helper()
err := s.Start()
require.NoError(t, err)
testutil.CleanupAndRequireSuccess(t, s.Stop)
err := s.Start(testutil.ContextWithTimeout(tb, testTimeout))
require.NoError(tb, err)
testutil.CleanupAndRequireSuccess(tb, func() (err error) {
return s.Stop(testutil.ContextWithTimeout(tb, testTimeout))
})
}
// applyEmptyClientFiltering is a helper function for tests with
@@ -126,11 +133,11 @@ func emptyFilteringBlockedServices() (bsvc *filtering.BlockedServices) {
// *Server for use in tests, given the provided parameters. It also populates
// the filtering configuration with default parameters.
func createTestServer(
t *testing.T,
tb testing.TB,
filterConf *filtering.Config,
forwardConf ServerConfig,
) (s *Server) {
t.Helper()
tb.Helper()
filterConf.Logger = cmp.Or(filterConf.Logger, testLogger)
@@ -151,14 +158,14 @@ func createTestServer(
}
f, err := filtering.New(filterConf, filters)
require.NoError(t, err)
require.NoError(tb, err)
f.SetEnabled(true)
dhcp := &testDHCP{
OnEnabled: func() (ok bool) { return false },
OnHostByIP: func(ip netip.Addr) (host string) { return "" },
OnIPByHost: func(host string) (ip netip.Addr) { panic("not implemented") },
OnIPByHost: func(host string) (_ netip.Addr) { panic(testutil.UnexpectedCall(host)) },
}
s, err = NewServer(DNSCreateParams{
DHCPServer: dhcp,
@@ -166,23 +173,23 @@ func createTestServer(
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
Logger: testLogger,
})
require.NoError(t, err)
require.NoError(tb, err)
err = s.Prepare(&forwardConf)
require.NoError(t, err)
err = s.Prepare(testutil.ContextWithTimeout(tb, testTimeout), &forwardConf)
require.NoError(tb, err)
return s
}
func createServerTLSConfig(t *testing.T) (*tls.Config, []byte, []byte) {
t.Helper()
func createServerTLSConfig(tb testing.TB) (*tls.Config, []byte, []byte) {
tb.Helper()
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
require.NoErrorf(t, err, "cannot generate RSA key: %s", err)
require.NoErrorf(tb, err, "cannot generate RSA key: %s", err)
serialNumberLimit := new(big.Int).Lsh(big.NewInt(1), 128)
serialNumber, err := rand.Int(rand.Reader, serialNumberLimit)
require.NoErrorf(t, err, "failed to generate serial number: %s", err)
require.NoErrorf(tb, err, "failed to generate serial number: %s", err)
notBefore := time.Now()
notAfter := notBefore.Add(5 * 365 * timeutil.Day)
@@ -203,13 +210,13 @@ func createServerTLSConfig(t *testing.T) (*tls.Config, []byte, []byte) {
template.DNSNames = append(template.DNSNames, tlsServerName)
derBytes, err := x509.CreateCertificate(rand.Reader, &template, &template, publicKey(privateKey), privateKey)
require.NoErrorf(t, err, "failed to create certificate: %s", err)
require.NoErrorf(tb, err, "failed to create certificate: %s", err)
certPem := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: derBytes})
keyPem := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(privateKey)})
cert, err := tls.X509KeyPair(certPem, keyPem)
require.NoErrorf(t, err, "failed to create certificate: %s", err)
require.NoErrorf(tb, err, "failed to create certificate: %s", err)
return &tls.Config{
Certificates: []tls.Certificate{cert},
@@ -218,18 +225,18 @@ func createServerTLSConfig(t *testing.T) (*tls.Config, []byte, []byte) {
}, certPem, keyPem
}
func createTestTLS(t *testing.T, tlsConf *TLSConfig) (s *Server, certPem []byte) {
t.Helper()
func createTestTLS(tb testing.TB, tlsConf *TLSConfig) (s *Server, certPem []byte) {
tb.Helper()
var keyPem []byte
_, certPem, keyPem = createServerTLSConfig(t)
_, certPem, keyPem = createServerTLSConfig(tb)
cert, err := tls.X509KeyPair(certPem, keyPem)
require.NoError(t, err)
require.NoError(tb, err)
tlsConf.Cert = &cert
s = createTestServer(t, &filtering.Config{
s = createTestServer(tb, &filtering.Config{
BlockingMode: filtering.BlockingModeDefault,
}, ServerConfig{
UDPListenAddrs: []*net.UDPAddr{{}},
@@ -243,8 +250,8 @@ func createTestTLS(t *testing.T, tlsConf *TLSConfig) (s *Server, certPem []byte)
ServePlainDNS: true,
})
err = s.Prepare(&s.conf)
require.NoErrorf(t, err, "failed to prepare server: %s", err)
err = s.Prepare(testutil.ContextWithTimeout(tb, testTimeout), &s.conf)
require.NoErrorf(tb, err, "failed to prepare server: %s", err)
return s, certPem
}
@@ -300,25 +307,27 @@ func newResp(rcode int, req *dns.Msg, ans []dns.RR) (resp *dns.Msg) {
return resp
}
func assertGoogleAResponse(t *testing.T, reply *dns.Msg) {
assertResponse(t, reply, netip.AddrFrom4([4]byte{8, 8, 8, 8}))
func assertGoogleAResponse(tb testing.TB, reply *dns.Msg) {
tb.Helper()
assertResponse(tb, reply, netip.AddrFrom4([4]byte{8, 8, 8, 8}))
}
func assertResponse(t *testing.T, reply *dns.Msg, ip netip.Addr) {
t.Helper()
func assertResponse(tb testing.TB, reply *dns.Msg, ip netip.Addr) {
tb.Helper()
require.Lenf(t, reply.Answer, 1, "dns server returned reply with wrong number of answers - %d", len(reply.Answer))
require.Lenf(tb, reply.Answer, 1, "dns server returned reply with wrong number of answers - %d", len(reply.Answer))
a, ok := reply.Answer[0].(*dns.A)
require.Truef(t, ok, "dns server returned wrong answer type instead of A: %v", reply.Answer[0])
assert.Equal(t, net.IP(ip.AsSlice()), a.A)
require.Truef(tb, ok, "dns server returned wrong answer type instead of A: %v", reply.Answer[0])
assert.Equal(tb, net.IP(ip.AsSlice()), a.A)
}
// sendTestMessagesAsync sends messages in parallel to check for race issues.
//
//lint:ignore U1000 it's called from the function which is skipped for now.
func sendTestMessagesAsync(t *testing.T, conn *dns.Conn) {
t.Helper()
func sendTestMessagesAsync(tb testing.TB, conn *dns.Conn) {
tb.Helper()
wg := &sync.WaitGroup{}
@@ -330,29 +339,29 @@ func sendTestMessagesAsync(t *testing.T, conn *dns.Conn) {
defer wg.Done()
err := conn.WriteMsg(msg)
require.NoErrorf(t, err, "cannot write message: %s", err)
require.NoErrorf(tb, err, "cannot write message: %s", err)
res, err := conn.ReadMsg()
require.NoErrorf(t, err, "cannot read response to message: %s", err)
require.NoErrorf(tb, err, "cannot read response to message: %s", err)
assertGoogleAResponse(t, res)
assertGoogleAResponse(tb, res)
}()
}
wg.Wait()
}
func sendTestMessages(t *testing.T, conn *dns.Conn) {
t.Helper()
func sendTestMessages(tb testing.TB, conn *dns.Conn) {
tb.Helper()
for i := range testMessagesCount {
req := createGoogleATestMessage()
err := conn.WriteMsg(req)
assert.NoErrorf(t, err, "cannot write message #%d: %s", i, err)
assert.NoErrorf(tb, err, "cannot write message #%d: %s", i, err)
res, err := conn.ReadMsg()
assert.NoErrorf(t, err, "cannot read response to message #%d: %s", i, err)
assertGoogleAResponse(t, res)
assert.NoErrorf(tb, err, "cannot read response to message #%d: %s", i, err)
assertGoogleAResponse(tb, res)
}
}
@@ -419,7 +428,7 @@ func TestServer_timeout(t *testing.T) {
})
require.NoError(t, err)
err = s.Prepare(srvConf)
err = s.Prepare(testutil.ContextWithTimeout(t, testTimeout), srvConf)
require.NoError(t, err)
assert.Equal(t, testTimeout, s.conf.UpstreamTimeout)
@@ -438,7 +447,7 @@ func TestServer_timeout(t *testing.T) {
Enabled: false,
}
s.conf.Config.ClientsContainer = EmptyClientsContainer{}
err = s.Prepare(&s.conf)
err = s.Prepare(testutil.ContextWithTimeout(t, testTimeout), &s.conf)
require.NoError(t, err)
assert.Equal(t, DefaultTimeout, s.conf.UpstreamTimeout)
@@ -465,7 +474,7 @@ func TestServer_Prepare_fallbacks(t *testing.T) {
})
require.NoError(t, err)
err = s.Prepare(srvConf)
err = s.Prepare(testutil.ContextWithTimeout(t, testTimeout), srvConf)
require.NoError(t, err)
require.NotNil(t, s.dnsProxy.Fallbacks)
@@ -565,8 +574,8 @@ func TestServerRace(t *testing.T) {
UpstreamMode: UpstreamModeLoadBalance,
UpstreamDNS: []string{"8.8.8.8:53", "8.8.4.4:53"},
},
ConfigModified: func() {},
ServePlainDNS: true,
ConfModifier: agh.EmptyConfigModifier{},
ServePlainDNS: true,
}
s := createTestServer(t, filterConf, forwardConf)
s.conf.UpstreamConfig.Upstreams = []upstream.Upstream{newGoogleUpstream()}
@@ -1076,8 +1085,8 @@ func TestBlockedCustomIP(t *testing.T) {
dhcp := &testDHCP{
OnEnabled: func() (ok bool) { return false },
OnHostByIP: func(_ netip.Addr) (host string) { panic("not implemented") },
OnIPByHost: func(_ string) (ip netip.Addr) { panic("not implemented") },
OnHostByIP: func(ip netip.Addr) (_ string) { panic(testutil.UnexpectedCall(ip)) },
OnIPByHost: func(host string) (_ netip.Addr) { panic(testutil.UnexpectedCall(host)) },
}
s, err := NewServer(DNSCreateParams{
DHCPServer: dhcp,
@@ -1103,7 +1112,7 @@ func TestBlockedCustomIP(t *testing.T) {
}
// Invalid BlockingIPv4.
err = s.Prepare(conf)
err = s.Prepare(testutil.ContextWithTimeout(t, testTimeout), conf)
assert.Error(t, err)
s.dnsFilter.SetBlockingMode(
@@ -1111,7 +1120,7 @@ func TestBlockedCustomIP(t *testing.T) {
netip.AddrFrom4([4]byte{0, 0, 0, 1}),
netip.MustParseAddr("::1"))
err = s.Prepare(conf)
err = s.Prepare(testutil.ContextWithTimeout(t, testTimeout), conf)
require.NoError(t, err)
f.SetEnabled(true)
@@ -1252,8 +1261,8 @@ func TestRewrite(t *testing.T) {
dhcp := &testDHCP{
OnEnabled: func() (ok bool) { return false },
OnHostByIP: func(ip netip.Addr) (host string) { panic("not implemented") },
OnIPByHost: func(host string) (ip netip.Addr) { panic("not implemented") },
OnHostByIP: func(ip netip.Addr) (_ string) { panic(testutil.UnexpectedCall(ip)) },
OnIPByHost: func(host string) (_ netip.Addr) { panic(testutil.UnexpectedCall(host)) },
}
s, err := NewServer(DNSCreateParams{
DHCPServer: dhcp,
@@ -1263,7 +1272,7 @@ func TestRewrite(t *testing.T) {
})
require.NoError(t, err)
assert.NoError(t, s.Prepare(&ServerConfig{
assert.NoError(t, s.Prepare(testutil.ContextWithTimeout(t, testTimeout), &ServerConfig{
UDPListenAddrs: []*net.UDPAddr{{}},
TCPListenAddrs: []*net.TCPAddr{{}},
TLSConf: &TLSConfig{},
@@ -1333,7 +1342,7 @@ func TestRewrite(t *testing.T) {
for _, protect := range []bool{true, false} {
val := protect
conf := s.getDNSConfig()
conf := s.getDNSConfig(testutil.ContextWithTimeout(t, testTimeout))
conf.ProtectionEnabled = &val
s.setConfig(conf)
@@ -1388,7 +1397,7 @@ func TestPTRResponseFromDHCPLeases(t *testing.T) {
DNSFilter: flt,
DHCPServer: &testDHCP{
OnEnabled: func() (ok bool) { return true },
OnIPByHost: func(host string) (ip netip.Addr) { panic("not implemented") },
OnIPByHost: func(host string) (_ netip.Addr) { panic(testutil.UnexpectedCall(host)) },
OnHostByIP: func(ip netip.Addr) (host string) {
return "myhost"
},
@@ -1407,12 +1416,12 @@ func TestPTRResponseFromDHCPLeases(t *testing.T) {
s.conf.Config.ClientsContainer = EmptyClientsContainer{}
s.conf.Config.UpstreamMode = UpstreamModeLoadBalance
err = s.Prepare(&s.conf)
err = s.Prepare(testutil.ContextWithTimeout(t, testTimeout), &s.conf)
require.NoError(t, err)
err = s.Start()
err = s.Start(testutil.ContextWithTimeout(t, testTimeout))
require.NoError(t, err)
t.Cleanup(s.Close)
t.Cleanup(func() { s.Close(testutil.ContextWithTimeout(t, testTimeout)) })
addr := s.dnsProxy.Addr(proxy.ProtoUDP)
req := createTestMessageWithType("34.12.168.192.in-addr.arpa.", dns.TypePTR)
@@ -1445,13 +1454,13 @@ func TestPTRResponseFromHosts(t *testing.T) {
dhcp := &testDHCP{
OnEnabled: func() (ok bool) { return false },
OnIPByHost: func(host string) (ip netip.Addr) { panic("not implemented") },
OnIPByHost: func(host string) (_ netip.Addr) { panic(testutil.UnexpectedCall(host)) },
OnHostByIP: func(ip netip.Addr) (host string) { return "" },
}
var eventsCalledCounter uint32
hc, err := aghnet.NewHostsContainer(testFS, &aghtest.FSWatcher{
OnStart: func() (_ error) { panic("not implemented") },
OnStart: func(ctx context.Context) (_ error) { panic(testutil.UnexpectedCall(ctx)) },
OnEvents: func() (e <-chan struct{}) {
assert.Equal(t, uint32(1), atomic.AddUint32(&eventsCalledCounter, 1))
@@ -1462,7 +1471,7 @@ func TestPTRResponseFromHosts(t *testing.T) {
return nil
},
OnClose: func() (err error) { panic("not implemented") },
OnShutdown: func(ctx context.Context) (err error) { panic(testutil.UnexpectedCall(ctx)) },
}, hostsFilename)
require.NoError(t, err)
t.Cleanup(func() {
@@ -1497,12 +1506,12 @@ func TestPTRResponseFromHosts(t *testing.T) {
s.conf.Config.ClientsContainer = EmptyClientsContainer{}
s.conf.Config.UpstreamMode = UpstreamModeLoadBalance
err = s.Prepare(&s.conf)
err = s.Prepare(testutil.ContextWithTimeout(t, testTimeout), &s.conf)
require.NoError(t, err)
err = s.Start()
err = s.Start(testutil.ContextWithTimeout(t, testTimeout))
require.NoError(t, err)
t.Cleanup(s.Close)
t.Cleanup(func() { s.Close(testutil.ContextWithTimeout(t, testTimeout)) })
subTestFunc := func(t *testing.T) {
addr := s.dnsProxy.Addr(proxy.ProtoUDP)
@@ -1523,7 +1532,7 @@ func TestPTRResponseFromHosts(t *testing.T) {
for _, protect := range []bool{true, false} {
val := protect
conf := s.getDNSConfig()
conf := s.getDNSConfig(testutil.ContextWithTimeout(t, testTimeout))
conf.ProtectionEnabled = &val
s.setConfig(conf)

View File

@@ -1,13 +1,13 @@
package dnsforward
import (
"context"
"fmt"
"net/netip"
"github.com/AdguardTeam/AdGuardHome/internal/filtering"
"github.com/AdguardTeam/dnsproxy/proxy"
"github.com/AdguardTeam/golibs/errors"
"github.com/AdguardTeam/golibs/log"
"github.com/AdguardTeam/urlfilter/rules"
"github.com/miekg/dns"
)
@@ -15,6 +15,7 @@ import (
// filterDNSRewriteResponse handles a single DNS rewrite response entry. It
// returns the properly constructed answer resource record.
func (s *Server) filterDNSRewriteResponse(
ctx context.Context,
req *dns.Msg,
rr rules.RRType,
v rules.RRValue,
@@ -27,11 +28,11 @@ func (s *Server) filterDNSRewriteResponse(
case dns.TypeMX:
return s.ansFromDNSRewriteMX(v, rr, req)
case dns.TypeHTTPS, dns.TypeSVCB:
return s.ansFromDNSRewriteSVCB(v, rr, req)
return s.ansFromDNSRewriteSVCB(ctx, v, rr, req)
case dns.TypeSRV:
return s.ansFromDNSRewriteSRV(v, rr, req)
default:
log.Debug("don't know how to handle dns rr type %d, skipping", rr)
s.logger.DebugContext(ctx, "unsupported dns rr type, skipping", "res_record", rr)
return nil, nil
}
@@ -97,6 +98,7 @@ func (s *Server) ansFromDNSRewriteMX(
// ansFromDNSRewriteSVCB creates a new answer resource record from the
// SVCB/HTTPS dnsrewrite rule data.
func (s *Server) ansFromDNSRewriteSVCB(
ctx context.Context,
v rules.RRValue,
rr rules.RRType,
req *dns.Msg,
@@ -111,10 +113,10 @@ func (s *Server) ansFromDNSRewriteSVCB(
}
if rr == dns.TypeHTTPS {
return s.genAnswerHTTPS(req, svcb), nil
return s.genAnswerHTTPS(ctx, req, svcb), nil
}
return s.genAnswerSVCB(req, svcb), nil
return s.genAnswerSVCB(ctx, req, svcb), nil
}
// ansFromDNSRewriteSRV creates a new answer resource record from the SRV
@@ -139,6 +141,7 @@ func (s *Server) ansFromDNSRewriteSRV(
// filterDNSRewrite handles dnsrewrite filters. It constructs a DNS response
// and sets it into pctx.Res. All parameters must not be nil.
func (s *Server) filterDNSRewrite(
ctx context.Context,
req *dns.Msg,
res *filtering.Result,
pctx *proxy.DNSContext,
@@ -164,7 +167,7 @@ func (s *Server) filterDNSRewrite(
values := dnsrr.Response[qtype]
for i, v := range values {
var ans dns.RR
ans, err = s.filterDNSRewriteResponse(req, qtype, v)
ans, err = s.filterDNSRewriteResponse(ctx, req, qtype, v)
if err != nil {
return fmt.Errorf("dns rewrite response for %s[%d]: %w", dns.Type(qtype), i, err)
}

View File

@@ -7,6 +7,7 @@ import (
"github.com/AdguardTeam/AdGuardHome/internal/filtering"
"github.com/AdguardTeam/dnsproxy/proxy"
"github.com/AdguardTeam/golibs/netutil"
"github.com/AdguardTeam/golibs/testutil"
"github.com/AdguardTeam/urlfilter/rules"
"github.com/miekg/dns"
"github.com/stretchr/testify/assert"
@@ -71,7 +72,7 @@ func TestServer_FilterDNSRewrite(t *testing.T) {
res := makeRes(dns.RcodeNameError, 0, nil)
d := &proxy.DNSContext{}
err := srv.filterDNSRewrite(req, res, d)
err := srv.filterDNSRewrite(testutil.ContextWithTimeout(t, testTimeout), req, res, d)
require.NoError(t, err)
assert.Equal(t, dns.RcodeNameError, d.Res.Rcode)
@@ -82,7 +83,7 @@ func TestServer_FilterDNSRewrite(t *testing.T) {
res := makeRes(dns.RcodeSuccess, 0, nil)
d := &proxy.DNSContext{}
err := srv.filterDNSRewrite(req, res, d)
err := srv.filterDNSRewrite(testutil.ContextWithTimeout(t, testTimeout), req, res, d)
require.NoError(t, err)
assert.Equal(t, dns.RcodeSuccess, d.Res.Rcode)
@@ -94,7 +95,7 @@ func TestServer_FilterDNSRewrite(t *testing.T) {
res := makeRes(dns.RcodeSuccess, dns.TypeA, ip4)
d := &proxy.DNSContext{}
err := srv.filterDNSRewrite(req, res, d)
err := srv.filterDNSRewrite(testutil.ContextWithTimeout(t, testTimeout), req, res, d)
require.NoError(t, err)
assert.Equal(t, dns.RcodeSuccess, d.Res.Rcode)
@@ -108,7 +109,7 @@ func TestServer_FilterDNSRewrite(t *testing.T) {
res := makeRes(dns.RcodeSuccess, dns.TypeAAAA, ip6)
d := &proxy.DNSContext{}
err := srv.filterDNSRewrite(req, res, d)
err := srv.filterDNSRewrite(testutil.ContextWithTimeout(t, testTimeout), req, res, d)
require.NoError(t, err)
assert.Equal(t, dns.RcodeSuccess, d.Res.Rcode)
@@ -122,7 +123,7 @@ func TestServer_FilterDNSRewrite(t *testing.T) {
res := makeRes(dns.RcodeSuccess, dns.TypePTR, domain)
d := &proxy.DNSContext{}
err := srv.filterDNSRewrite(req, res, d)
err := srv.filterDNSRewrite(testutil.ContextWithTimeout(t, testTimeout), req, res, d)
require.NoError(t, err)
assert.Equal(t, dns.RcodeSuccess, d.Res.Rcode)
@@ -136,7 +137,7 @@ func TestServer_FilterDNSRewrite(t *testing.T) {
res := makeRes(dns.RcodeSuccess, dns.TypeTXT, domain)
d := &proxy.DNSContext{}
err := srv.filterDNSRewrite(req, res, d)
err := srv.filterDNSRewrite(testutil.ContextWithTimeout(t, testTimeout), req, res, d)
require.NoError(t, err)
assert.Equal(t, dns.RcodeSuccess, d.Res.Rcode)
@@ -150,7 +151,7 @@ func TestServer_FilterDNSRewrite(t *testing.T) {
res := makeRes(dns.RcodeSuccess, dns.TypeMX, mxVal)
d := &proxy.DNSContext{}
err := srv.filterDNSRewrite(req, res, d)
err := srv.filterDNSRewrite(testutil.ContextWithTimeout(t, testTimeout), req, res, d)
require.NoError(t, err)
assert.Equal(t, dns.RcodeSuccess, d.Res.Rcode)
@@ -168,7 +169,7 @@ func TestServer_FilterDNSRewrite(t *testing.T) {
res := makeRes(dns.RcodeSuccess, dns.TypeSVCB, svcbVal)
d := &proxy.DNSContext{}
err := srv.filterDNSRewrite(req, res, d)
err := srv.filterDNSRewrite(testutil.ContextWithTimeout(t, testTimeout), req, res, d)
require.NoError(t, err)
assert.Equal(t, dns.RcodeSuccess, d.Res.Rcode)
@@ -198,7 +199,7 @@ func TestServer_FilterDNSRewrite(t *testing.T) {
res := makeRes(dns.RcodeSuccess, dns.TypeHTTPS, svcbVal)
d := &proxy.DNSContext{}
err := srv.filterDNSRewrite(req, res, d)
err := srv.filterDNSRewrite(testutil.ContextWithTimeout(t, testTimeout), req, res, d)
require.NoError(t, err)
assert.Equal(t, dns.RcodeSuccess, d.Res.Rcode)
@@ -228,7 +229,7 @@ func TestServer_FilterDNSRewrite(t *testing.T) {
res := makeRes(dns.RcodeSuccess, dns.TypeSRV, srvVal)
d := &proxy.DNSContext{}
err := srv.filterDNSRewrite(req, res, d)
err := srv.filterDNSRewrite(testutil.ContextWithTimeout(t, testTimeout), req, res, d)
require.NoError(t, err)
assert.Equal(t, dns.RcodeSuccess, d.Res.Rcode)

View File

@@ -1,13 +1,13 @@
package dnsforward
import (
"context"
"fmt"
"net"
"slices"
"strings"
"github.com/AdguardTeam/AdGuardHome/internal/filtering"
"github.com/AdguardTeam/golibs/log"
"github.com/AdguardTeam/urlfilter/rules"
"github.com/miekg/dns"
)
@@ -24,7 +24,10 @@ func (s *Server) clientRequestFilteringSettings(dctx *dnsContext) (setts *filter
// filterDNSRequest applies the dnsFilter and sets dctx.proxyCtx.Res if the
// request was filtered.
func (s *Server) filterDNSRequest(dctx *dnsContext) (res *filtering.Result, err error) {
func (s *Server) filterDNSRequest(
ctx context.Context,
dctx *dnsContext,
) (res *filtering.Result, err error) {
pctx := dctx.proxyCtx
req := pctx.Req
q := req.Question[0]
@@ -44,12 +47,12 @@ func (s *Server) filterDNSRequest(dctx *dnsContext) (res *filtering.Result, err
dctx.origQuestion = q
req.Question[0].Name = dns.Fqdn(res.CanonName)
case res.IsFiltered:
log.Debug("dnsforward: host %q is filtered, reason: %q", host, res.Reason)
pctx.Res = s.genDNSFilterMessage(pctx, res)
s.logger.DebugContext(ctx, "host is filtered", "host", host, "reason", res.Reason)
pctx.Res = s.genDNSFilterMessage(ctx, pctx, res)
case res.Reason.In(filtering.Rewritten, filtering.FilteredSafeSearch):
pctx.Res = s.getCNAMEWithIPs(req, res.IPList, res.CanonName)
pctx.Res = s.getCNAMEWithIPs(ctx, req, res.IPList, res.CanonName)
case res.Reason.In(filtering.RewrittenRule, filtering.RewrittenAutoHosts):
if err = s.filterDNSRewrite(req, res, pctx); err != nil {
if err = s.filterDNSRewrite(ctx, req, res, pctx); err != nil {
return nil, err
}
}
@@ -90,7 +93,7 @@ func (s *Server) checkHostRules(
// dctx.proxyCtx.Res. It sets dctx.result and dctx.origResp if at least one of
// canonical names, IP addresses, or HTTPS RR hints in it matches the filtering
// rules, as well as sets dctx.proxyCtx.Res to the filtered response.
func (s *Server) filterDNSResponse(dctx *dnsContext) (err error) {
func (s *Server) filterDNSResponse(ctx context.Context, dctx *dnsContext) (err error) {
setts := dctx.setts
if !setts.FilteringEnabled {
return nil
@@ -123,16 +126,27 @@ func (s *Server) filterDNSResponse(dctx *dnsContext) (err error) {
continue
}
log.Debug("dnsforward: checked %s %s for %s", dns.Type(rrtype), host, a.Header().Name)
s.logger.DebugContext(
ctx,
"checked",
"dns_type", dns.Type(rrtype),
"host", host,
"name", a.Header().Name,
)
if err != nil {
return fmt.Errorf("filtering answer at index %d: %w", i, err)
} else if res != nil && res.IsFiltered {
dctx.result = res
dctx.origResp = pctx.Res
pctx.Res = s.genDNSFilterMessage(pctx, res)
pctx.Res = s.genDNSFilterMessage(ctx, pctx, res)
log.Debug("dnsforward: matched %q by response: %q", pctx.Req.Question[0].Name, host)
s.logger.DebugContext(
ctx,
"matched by response",
"name", pctx.Req.Question[0].Name,
"host", host,
)
break
}

View File

@@ -10,6 +10,7 @@ import (
"github.com/AdguardTeam/dnsproxy/proxy"
"github.com/AdguardTeam/dnsproxy/upstream"
"github.com/AdguardTeam/golibs/netutil"
"github.com/AdguardTeam/golibs/testutil"
"github.com/miekg/dns"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -57,8 +58,8 @@ func TestHandleDNSRequest_handleDNSRequest(t *testing.T) {
s, err := NewServer(DNSCreateParams{
DHCPServer: &testDHCP{
OnEnabled: func() (ok bool) { return false },
OnHostByIP: func(ip netip.Addr) (host string) { panic("not implemented") },
OnIPByHost: func(host string) (ip netip.Addr) { panic("not implemented") },
OnHostByIP: func(ip netip.Addr) (_ string) { panic(testutil.UnexpectedCall(ip)) },
OnIPByHost: func(host string) (_ netip.Addr) { panic(testutil.UnexpectedCall(host)) },
},
DNSFilter: f,
PrivateNets: netutil.SubnetSetFunc(netutil.IsLocallyServed),
@@ -66,7 +67,7 @@ func TestHandleDNSRequest_handleDNSRequest(t *testing.T) {
})
require.NoError(t, err)
err = s.Prepare(&forwardConf)
err = s.Prepare(testutil.ContextWithTimeout(t, testTimeout), &forwardConf)
require.NoError(t, err)
s.conf.UpstreamConfig.Upstreams = []upstream.Upstream{
@@ -347,7 +348,7 @@ func TestHandleDNSRequest_filterDNSResponse(t *testing.T) {
},
}
fltErr := s.filterDNSResponse(dctx)
fltErr := s.filterDNSResponse(testutil.ContextWithTimeout(t, testTimeout), dctx)
require.NoError(t, fltErr)
res := dctx.result

View File

@@ -2,6 +2,7 @@ package dnsforward
import (
"cmp"
"context"
"encoding/json"
"fmt"
"io"
@@ -17,7 +18,6 @@ import (
"github.com/AdguardTeam/dnsproxy/proxy"
"github.com/AdguardTeam/dnsproxy/upstream"
"github.com/AdguardTeam/golibs/errors"
"github.com/AdguardTeam/golibs/log"
"github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/AdguardTeam/golibs/netutil"
"github.com/AdguardTeam/golibs/stringutil"
@@ -93,6 +93,9 @@ type jsonDNSConfig struct {
// CacheMaxTTL is custom maximum TTL for cached DNS responses.
CacheMaxTTL *uint32 `json:"cache_ttl_max"`
// CacheEnabled defines if the DNS cache should be used.
CacheEnabled *bool `json:"cache_enabled"`
// CacheOptimistic defines if expired entries should be served.
CacheOptimistic *bool `json:"cache_optimistic"`
@@ -138,8 +141,8 @@ const (
jsonUpstreamModeFastestAddr jsonUpstreamMode = "fastest_addr"
)
func (s *Server) getDNSConfig() (c *jsonDNSConfig) {
protectionEnabled, protectionDisabledUntil := s.UpdatedProtectionStatus()
func (s *Server) getDNSConfig(ctx context.Context) (c *jsonDNSConfig) {
protectionEnabled, protectionDisabledUntil := s.UpdatedProtectionStatus(ctx)
s.serverLock.RLock()
defer s.serverLock.RUnlock()
@@ -162,6 +165,7 @@ func (s *Server) getDNSConfig() (c *jsonDNSConfig) {
enableDNSSEC := s.conf.EnableDNSSEC
aaaaDisabled := s.conf.AAAADisabled
cacheEnabled := s.conf.CacheEnabled
cacheSize := s.conf.CacheSize
cacheMinTTL := s.conf.CacheMinTTL
cacheMaxTTL := s.conf.CacheMaxTTL
@@ -184,7 +188,7 @@ func (s *Server) getDNSConfig() (c *jsonDNSConfig) {
defPTRUps, err := s.defaultLocalPTRUpstreams()
if err != nil {
log.Error("dnsforward: %s", err)
s.logger.ErrorContext(ctx, "getting local ptr upstreams", slogutil.KeyError, err)
}
return &jsonDNSConfig{
@@ -207,6 +211,7 @@ func (s *Server) getDNSConfig() (c *jsonDNSConfig) {
DNSSECEnabled: &enableDNSSEC,
DisableIPv6: &aaaaDisabled,
BlockedResponseTTL: &blockedResponseTTL,
CacheEnabled: &cacheEnabled,
CacheSize: &cacheSize,
CacheMinTTL: &cacheMinTTL,
CacheMaxTTL: &cacheMaxTTL,
@@ -240,7 +245,7 @@ func (s *Server) defaultLocalPTRUpstreams() (ups []string, err error) {
// handleGetConfig handles requests to the GET /control/dns_info endpoint.
func (s *Server) handleGetConfig(w http.ResponseWriter, r *http.Request) {
resp := s.getDNSConfig()
resp := s.getDNSConfig(r.Context())
aghhttp.WriteJSONResponseOK(w, r, resp)
}
@@ -278,6 +283,7 @@ func (req *jsonDNSConfig) validate(
ownAddrs addrPortSet,
sysResolvers SystemResolvers,
privateNets netutil.SubnetSet,
curCacheSize uint32,
) (err error) {
defer func() { err = errors.Annotate(err, "validating dns config: %w") }()
@@ -305,7 +311,7 @@ func (req *jsonDNSConfig) validate(
return err
}
err = req.checkCacheTTL()
err = req.validateCacheSettings(curCacheSize)
if err != nil {
// Don't wrap the error since it's informative enough as is.
return err
@@ -421,9 +427,14 @@ func (req *jsonDNSConfig) validateUpstreamDNSServers(
return nil
}
// checkCacheTTL returns an error if the configuration of the cache TTL is
// invalid.
func (req *jsonDNSConfig) checkCacheTTL() (err error) {
// validateCacheSettings returns an error if the cache configuration is invalid.
func (req *jsonDNSConfig) validateCacheSettings(curCacheSize uint32) (err error) {
err = req.validateCacheSize(curCacheSize)
if err != nil {
// Don't wrap the error because it's informative enough as is.
return err
}
if req.CacheMinTTL == nil && req.CacheMaxTTL == nil {
return nil
}
@@ -440,6 +451,28 @@ func (req *jsonDNSConfig) checkCacheTTL() (err error) {
return validateCacheTTL(minTTL, maxTTL)
}
// validateCacheSize returns an error if the cache size configuration is
// invalid. It also explicitly sets CacheEnabled to support legacy behavior.
func (req *jsonDNSConfig) validateCacheSize(curCacheSize uint32) (err error) {
if req.CacheEnabled != nil && *req.CacheEnabled {
size := curCacheSize
if req.CacheSize != nil {
size = *req.CacheSize
}
if size == 0 {
return errors.Error("cache_size must be greater than zero when cache_enabled is true")
}
}
if req.CacheEnabled == nil && req.CacheSize != nil {
isEnabled := *req.CacheSize > 0
req.CacheEnabled = &isEnabled
}
return nil
}
// checkRatelimitSubnetMaskLen returns an error if the length of the subnet mask
// for IPv4 or IPv6 addresses is invalid.
func (req *jsonDNSConfig) checkRatelimitSubnetMaskLen() (err error) {
@@ -486,6 +519,8 @@ func checkInclusion(ptr *int, minN, maxN int) (err error) {
// handleSetConfig handles requests to the POST /control/dns_config endpoint.
func (s *Server) handleSetConfig(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
req := &jsonDNSConfig{}
err := json.NewDecoder(r.Body).Decode(req)
if err != nil {
@@ -503,7 +538,7 @@ func (s *Server) handleSetConfig(w http.ResponseWriter, r *http.Request) {
return
}
err = req.validate(ourAddrs, s.sysResolvers, s.privateNets)
err = req.validate(ourAddrs, s.sysResolvers, s.privateNets, s.conf.CacheSize)
if err != nil {
aghhttp.Error(r, w, http.StatusBadRequest, "%s", err)
@@ -511,10 +546,10 @@ func (s *Server) handleSetConfig(w http.ResponseWriter, r *http.Request) {
}
restart := s.setConfig(req)
s.conf.ConfigModified()
s.conf.ConfModifier.Apply(ctx)
if restart {
err = s.Reconfigure(nil)
err = s.Reconfigure(ctx, nil)
if err != nil {
aghhttp.Error(r, w, http.StatusInternalServerError, "%s", err)
}
@@ -596,6 +631,7 @@ func (s *Server) setConfigRestartable(dc *jsonDNSConfig) (shouldRestart bool) {
setIfNotNil(&s.conf.FallbackDNS, dc.Fallbacks),
setIfNotNil(&s.conf.EDNSClientSubnet.Enabled, dc.EDNSCSEnabled),
setIfNotNil(&s.conf.EDNSClientSubnet.UseCustom, dc.EDNSCSUseCustom),
setIfNotNil(&s.conf.CacheEnabled, dc.CacheEnabled),
setIfNotNil(&s.conf.CacheSize, dc.CacheSize),
setIfNotNil(&s.conf.CacheMinTTL, dc.CacheMinTTL),
setIfNotNil(&s.conf.CacheMaxTTL, dc.CacheMaxTTL),
@@ -726,7 +762,7 @@ func (s *Server) handleSetProtection(w http.ResponseWriter, r *http.Request) {
s.dnsFilter.SetProtectionStatus(protectionReq.Enabled, disabledUntil)
}()
s.conf.ConfigModified()
s.conf.ConfModifier.Apply(r.Context())
aghhttp.OK(w)
}

View File

@@ -2,6 +2,7 @@ package dnsforward
import (
"bytes"
"context"
"encoding/json"
"io"
"net"
@@ -16,6 +17,7 @@ import (
"testing/fstest"
"time"
"github.com/AdguardTeam/AdGuardHome/internal/agh"
"github.com/AdguardTeam/AdGuardHome/internal/aghhttp"
"github.com/AdguardTeam/AdGuardHome/internal/aghnet"
"github.com/AdguardTeam/AdGuardHome/internal/aghtest"
@@ -42,16 +44,18 @@ func (emptySysResolvers) Addrs() (addrs []netip.AddrPort) {
return nil
}
func loadTestData(t *testing.T, casesFileName string, cases any) {
t.Helper()
// loadTestData loads the test data from the file with the given name into
// cases.
func loadTestData(tb testing.TB, casesFileName string, cases any) {
tb.Helper()
var f *os.File
f, err := os.Open(filepath.Join("testdata", casesFileName))
require.NoError(t, err)
testutil.CleanupAndRequireSuccess(t, f.Close)
require.NoError(tb, err)
testutil.CleanupAndRequireSuccess(tb, f.Close)
err = json.NewDecoder(f).Decode(cases)
require.NoError(t, err)
require.NoError(tb, err)
}
const (
@@ -86,14 +90,16 @@ func TestDNSForwardHTTP_handleGetConfig(t *testing.T) {
EDNSClientSubnet: &EDNSClientSubnet{Enabled: false},
ClientsContainer: EmptyClientsContainer{},
},
ConfigModified: func() {},
ServePlainDNS: true,
ConfModifier: agh.EmptyConfigModifier{},
ServePlainDNS: true,
}
s := createTestServer(t, filterConf, forwardConf)
s.sysResolvers = &emptySysResolvers{}
require.NoError(t, s.Start())
testutil.CleanupAndRequireSuccess(t, s.Stop)
require.NoError(t, s.Start(testutil.ContextWithTimeout(t, testTimeout)))
testutil.CleanupAndRequireSuccess(t, func() (err error) {
return s.Stop(testutil.ContextWithTimeout(t, testTimeout))
})
defaultConf := s.conf
@@ -136,7 +142,7 @@ func TestDNSForwardHTTP_handleGetConfig(t *testing.T) {
t.Cleanup(w.Body.Reset)
s.conf = tc.conf()
s.handleGetConfig(w, nil)
s.handleGetConfig(w, httptest.NewRequest(http.MethodGet, "/", nil))
cType := w.Header().Get(httphdr.ContentType)
assert.Equal(t, aghhttp.HdrValApplicationJSON, cType)
@@ -169,17 +175,19 @@ func TestDNSForwardHTTP_handleSetConfig(t *testing.T) {
EDNSClientSubnet: &EDNSClientSubnet{Enabled: false},
ClientsContainer: EmptyClientsContainer{},
},
ConfigModified: func() {},
ServePlainDNS: true,
ConfModifier: agh.EmptyConfigModifier{},
ServePlainDNS: true,
}
s := createTestServer(t, filterConf, forwardConf)
s.sysResolvers = &emptySysResolvers{}
defaultConf := s.conf
err := s.Start()
err := s.Start(testutil.ContextWithTimeout(t, testTimeout))
assert.NoError(t, err)
testutil.CleanupAndRequireSuccess(t, s.Stop)
testutil.CleanupAndRequireSuccess(t, func() (err error) {
return s.Stop(testutil.ContextWithTimeout(t, testTimeout))
})
w := httptest.NewRecorder()
@@ -223,6 +231,9 @@ func TestDNSForwardHTTP_handleSetConfig(t *testing.T) {
}, {
name: "cache_size",
wantSet: "",
}, {
name: "cache_enabled",
wantSet: "",
}, {
name: "upstream_mode_parallel",
wantSet: "",
@@ -296,21 +307,24 @@ func TestDNSForwardHTTP_handleSetConfig(t *testing.T) {
assert.Equal(t, tc.wantSet, strings.TrimSuffix(w.Body.String(), "\n"))
w.Body.Reset()
s.handleGetConfig(w, nil)
s.handleGetConfig(w, httptest.NewRequest(http.MethodGet, "/", nil))
assert.JSONEq(t, string(caseData.Want), w.Body.String())
w.Body.Reset()
})
}
}
func newLocalUpstreamListener(t *testing.T, port uint16, handler dns.Handler) (real netip.AddrPort) {
t.Helper()
// newLocalUpstreamListener creates a local upstream listener and returns its
// address. The listener is started in a separate goroutine and stopped when
// the tb's test is finished.
func newLocalUpstreamListener(tb testing.TB, port uint16, h dns.Handler) (real netip.AddrPort) {
tb.Helper()
startCh := make(chan struct{})
upsSrv := &dns.Server{
Addr: netip.AddrPortFrom(netutil.IPv4Localhost(), port).String(),
Net: "tcp",
Handler: handler,
Handler: h,
NotifyStartedFunc: func() { close(startCh) },
}
go func() {
@@ -319,9 +333,9 @@ func newLocalUpstreamListener(t *testing.T, port uint16, handler dns.Handler) (r
}()
<-startCh
testutil.CleanupAndRequireSuccess(t, upsSrv.Shutdown)
testutil.CleanupAndRequireSuccess(tb, upsSrv.Shutdown)
return testutil.RequireTypeAssert[*net.TCPAddr](t, upsSrv.Listener.Addr()).AddrPort()
return testutil.RequireTypeAssert[*net.TCPAddr](tb, upsSrv.Listener.Addr()).AddrPort()
}
func TestServer_HandleTestUpstreamDNS(t *testing.T) {
@@ -355,10 +369,10 @@ func TestServer_HandleTestUpstreamDNS(t *testing.T) {
},
},
&aghtest.FSWatcher{
OnStart: func() (_ error) { panic("not implemented") },
OnEvents: func() (e <-chan struct{}) { return nil },
OnAdd: func(_ string) (err error) { return nil },
OnClose: func() (err error) { return nil },
OnStart: func(ctx context.Context) (_ error) { panic(testutil.UnexpectedCall(ctx)) },
OnEvents: func() (e <-chan struct{}) { return nil },
OnAdd: func(_ string) (err error) { return nil },
OnShutdown: func(_ context.Context) (err error) { return nil },
},
hostsFileName,
)
@@ -460,10 +474,8 @@ func TestServer_HandleTestUpstreamDNS(t *testing.T) {
require.NoError(t, err)
require.Contains(t, resp, sleepyUps)
require.IsType(t, "", resp[sleepyUps])
sleepyRes, _ := resp[sleepyUps].(string)
sleepyRes := testutil.RequireTypeAssert[string](t, resp[sleepyUps])
// TODO(e.burkov): Improve the format of an error in dnsproxy.
assert.True(t, strings.HasSuffix(sleepyRes, "i/o timeout"))
})
}

View File

@@ -121,9 +121,7 @@ func ipsFromAnswer(ans []dns.RR) (ip4s, ip6s []net.IP) {
}
// process adds the resolved IP addresses to the domain's ipsets, if any.
func (h *ipsetHandler) process(dctx *dnsContext) (rc resultCode) {
// TODO(s.chzhen): Use passed context.
ctx := context.TODO()
func (h *ipsetHandler) process(ctx context.Context, dctx *dnsContext) (rc resultCode) {
h.logger.DebugContext(ctx, "started processing")
defer h.logger.DebugContext(ctx, "finished processing")

View File

@@ -6,6 +6,7 @@ import (
"testing"
"github.com/AdguardTeam/dnsproxy/proxy"
"github.com/AdguardTeam/golibs/testutil"
"github.com/miekg/dns"
"github.com/stretchr/testify/assert"
)
@@ -62,7 +63,7 @@ func TestIpsetCtx_process(t *testing.T) {
ictx := &ipsetHandler{
logger: testLogger,
}
rc := ictx.process(dctx)
rc := ictx.process(testutil.ContextWithTimeout(t, testTimeout), dctx)
assert.Equal(t, resultCodeSuccess, rc)
err := ictx.close()
@@ -85,7 +86,7 @@ func TestIpsetCtx_process(t *testing.T) {
logger: testLogger,
}
rc := ictx.process(dctx)
rc := ictx.process(testutil.ContextWithTimeout(t, testTimeout), dctx)
assert.Equal(t, resultCodeSuccess, rc)
assert.Equal(t, []net.IP{ip4}, m.ip4s)
assert.Empty(t, m.ip6s)
@@ -110,7 +111,7 @@ func TestIpsetCtx_process(t *testing.T) {
logger: testLogger,
}
rc := ictx.process(dctx)
rc := ictx.process(testutil.ContextWithTimeout(t, testTimeout), dctx)
assert.Equal(t, resultCodeSuccess, rc)
assert.Empty(t, m.ip4s)
assert.Equal(t, []net.IP{ip6}, m.ip6s)

View File

@@ -1,12 +1,13 @@
package dnsforward
import (
"context"
"net/netip"
"slices"
"github.com/AdguardTeam/AdGuardHome/internal/filtering"
"github.com/AdguardTeam/dnsproxy/proxy"
"github.com/AdguardTeam/golibs/log"
"github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/AdguardTeam/urlfilter/rules"
"github.com/miekg/dns"
)
@@ -47,6 +48,7 @@ func ipsFromRules(resRules []*filtering.ResultRule) (ips []netip.Addr) {
// genDNSFilterMessage generates a filtered response to req for the filtering
// result res.
func (s *Server) genDNSFilterMessage(
ctx context.Context,
dctx *proxy.DNSContext,
res *filtering.Result,
) (resp *dns.Msg) {
@@ -63,22 +65,27 @@ func (s *Server) genDNSFilterMessage(
switch res.Reason {
case filtering.FilteredSafeBrowsing:
return s.genBlockedHost(req, s.dnsFilter.SafeBrowsingBlockHost(), dctx)
return s.genBlockedHost(ctx, req, s.dnsFilter.SafeBrowsingBlockHost(), dctx)
case filtering.FilteredParental:
return s.genBlockedHost(req, s.dnsFilter.ParentalBlockHost(), dctx)
return s.genBlockedHost(ctx, req, s.dnsFilter.ParentalBlockHost(), dctx)
case filtering.FilteredSafeSearch:
// If Safe Search generated the necessary IP addresses, use them.
// Otherwise, if there were no errors, there are no addresses for the
// requested IP version, so produce a NODATA response.
return s.getCNAMEWithIPs(req, ipsFromRules(res.Rules), res.CanonName)
return s.getCNAMEWithIPs(ctx, req, ipsFromRules(res.Rules), res.CanonName)
default:
return s.genForBlockingMode(req, ipsFromRules(res.Rules))
return s.genForBlockingMode(ctx, req, ipsFromRules(res.Rules))
}
}
// getCNAMEWithIPs generates a filtered response to req for with CNAME record
// and provided ips.
func (s *Server) getCNAMEWithIPs(req *dns.Msg, ips []netip.Addr, cname string) (resp *dns.Msg) {
func (s *Server) getCNAMEWithIPs(
ctx context.Context,
req *dns.Msg,
ips []netip.Addr,
cname string,
) (resp *dns.Msg) {
resp = s.replyCompressed(req)
originalName := req.Question[0].Name
@@ -94,7 +101,7 @@ func (s *Server) getCNAMEWithIPs(req *dns.Msg, ips []netip.Addr, cname string) (
switch req.Question[0].Qtype {
case dns.TypeA:
ans = append(ans, s.genAnswersWithIPv4s(req, ips)...)
ans = append(ans, s.genAnswersWithIPv4s(ctx, req, ips)...)
case dns.TypeAAAA:
for _, ip := range ips {
if ip.Is6() {
@@ -112,24 +119,28 @@ func (s *Server) getCNAMEWithIPs(req *dns.Msg, ips []netip.Addr, cname string) (
// genForBlockingMode generates a filtered response to req based on the server's
// blocking mode.
func (s *Server) genForBlockingMode(req *dns.Msg, ips []netip.Addr) (resp *dns.Msg) {
func (s *Server) genForBlockingMode(
ctx context.Context,
req *dns.Msg,
ips []netip.Addr,
) (resp *dns.Msg) {
switch mode, bIPv4, bIPv6 := s.dnsFilter.BlockingMode(); mode {
case filtering.BlockingModeCustomIP:
return s.makeResponseCustomIP(req, bIPv4, bIPv6)
return s.makeResponseCustomIP(ctx, req, bIPv4, bIPv6)
case filtering.BlockingModeDefault:
if len(ips) > 0 {
return s.genResponseWithIPs(req, ips)
return s.genResponseWithIPs(ctx, req, ips)
}
return s.makeResponseNullIP(req)
return s.makeResponseNullIP(ctx, req)
case filtering.BlockingModeNullIP:
return s.makeResponseNullIP(req)
return s.makeResponseNullIP(ctx, req)
case filtering.BlockingModeNXDOMAIN:
return s.NewMsgNXDOMAIN(req)
case filtering.BlockingModeREFUSED:
return s.makeResponseREFUSED(req)
default:
log.Error("dnsforward: invalid blocking mode %q", mode)
s.logger.ErrorContext(ctx, "invalid blocking mode", "mode", mode)
return s.replyCompressed(req)
}
@@ -138,6 +149,7 @@ func (s *Server) genForBlockingMode(req *dns.Msg, ips []netip.Addr) (resp *dns.M
// makeResponseCustomIP generates a DNS response message for Custom IP blocking
// mode with the provided IP addresses and an appropriate resource record type.
func (s *Server) makeResponseCustomIP(
ctx context.Context,
req *dns.Msg,
bIPv4 netip.Addr,
bIPv6 netip.Addr,
@@ -150,7 +162,11 @@ func (s *Server) makeResponseCustomIP(
default:
// Generally shouldn't happen, since the types are checked in
// genDNSFilterMessage.
log.Error("dnsforward: invalid msg type %s for custom IP blocking mode", dns.Type(qt))
s.logger.ErrorContext(
ctx,
"invalid message type for custom IP blocking mode",
"dns_type", dns.Type(qt),
)
return s.replyCompressed(req)
}
@@ -234,11 +250,15 @@ func (s *Server) genAnswerTXT(req *dns.Msg, strs []string) (ans *dns.TXT) {
// addresses and an appropriate resource record type. If any of the IPs cannot
// be converted to the correct protocol, genResponseWithIPs returns an empty
// response.
func (s *Server) genResponseWithIPs(req *dns.Msg, ips []netip.Addr) (resp *dns.Msg) {
func (s *Server) genResponseWithIPs(
ctx context.Context,
req *dns.Msg,
ips []netip.Addr,
) (resp *dns.Msg) {
var ans []dns.RR
switch req.Question[0].Qtype {
case dns.TypeA:
ans = s.genAnswersWithIPv4s(req, ips)
ans = s.genAnswersWithIPv4s(ctx, req, ips)
case dns.TypeAAAA:
for _, ip := range ips {
if ip.Is6() {
@@ -258,10 +278,14 @@ func (s *Server) genResponseWithIPs(req *dns.Msg, ips []netip.Addr) (resp *dns.M
// genAnswersWithIPv4s generates DNS A answers provided IPv4 addresses. If any
// of the IPs isn't an IPv4 address, genAnswersWithIPv4s logs a warning and
// returns nil,
func (s *Server) genAnswersWithIPv4s(req *dns.Msg, ips []netip.Addr) (ans []dns.RR) {
func (s *Server) genAnswersWithIPv4s(
ctx context.Context,
req *dns.Msg,
ips []netip.Addr,
) (ans []dns.RR) {
for _, ip := range ips {
if !ip.Is4() {
log.Info("dnsforward: warning: ip %s is not ipv4 address", ip)
s.logger.WarnContext(ctx, "ip is not an ipv4 address", "ip", ip)
return nil
}
@@ -274,16 +298,16 @@ func (s *Server) genAnswersWithIPv4s(req *dns.Msg, ips []netip.Addr) (ans []dns.
// makeResponseNullIP creates a response with 0.0.0.0 for A requests, :: for
// AAAA requests, and an empty response for other types.
func (s *Server) makeResponseNullIP(req *dns.Msg) (resp *dns.Msg) {
func (s *Server) makeResponseNullIP(ctx context.Context, req *dns.Msg) (resp *dns.Msg) {
// Respond with the corresponding zero IP type as opposed to simply
// using one or the other in both cases, because the IPv4 zero IP is
// converted to a IPV6-mapped IPv4 address, while the IPv6 zero IP is
// converted into an empty slice instead of the zero IPv4.
switch req.Question[0].Qtype {
case dns.TypeA:
resp = s.genResponseWithIPs(req, []netip.Addr{netip.IPv4Unspecified()})
resp = s.genResponseWithIPs(ctx, req, []netip.Addr{netip.IPv4Unspecified()})
case dns.TypeAAAA:
resp = s.genResponseWithIPs(req, []netip.Addr{netip.IPv6Unspecified()})
resp = s.genResponseWithIPs(ctx, req, []netip.Addr{netip.IPv6Unspecified()})
default:
resp = s.replyCompressed(req)
}
@@ -291,16 +315,21 @@ func (s *Server) makeResponseNullIP(req *dns.Msg) (resp *dns.Msg) {
return resp
}
func (s *Server) genBlockedHost(request *dns.Msg, newAddr string, d *proxy.DNSContext) *dns.Msg {
func (s *Server) genBlockedHost(
ctx context.Context,
request *dns.Msg,
newAddr string,
d *proxy.DNSContext,
) (msg *dns.Msg) {
if newAddr == "" {
log.Info("dnsforward: block host is not specified")
s.logger.InfoContext(ctx, "block host not specified")
return s.NewMsgSERVFAIL(request)
}
ip, err := netip.ParseAddr(newAddr)
if err == nil {
return s.genResponseWithIPs(request, []netip.Addr{ip})
return s.genResponseWithIPs(ctx, request, []netip.Addr{ip})
}
// look up the hostname, TODO: cache
@@ -316,14 +345,19 @@ func (s *Server) genBlockedHost(request *dns.Msg, newAddr string, d *proxy.DNSCo
prx := s.proxy()
if prx == nil {
log.Debug("dnsforward: %s", srvClosedErr)
s.logger.DebugContext(ctx, "getting current proxy", slogutil.KeyError, srvClosedErr)
return s.NewMsgSERVFAIL(request)
}
err = prx.Resolve(newContext)
if err != nil {
log.Info("dnsforward: looking up replacement host %q: %s", newAddr, err)
s.logger.ErrorContext(
ctx,
"looking up replacement host",
"host", newAddr,
slogutil.KeyError, err,
)
return s.NewMsgSERVFAIL(request)
}

View File

@@ -10,7 +10,6 @@ import (
"github.com/AdguardTeam/AdGuardHome/internal/filtering"
"github.com/AdguardTeam/dnsproxy/proxy"
"github.com/AdguardTeam/golibs/log"
"github.com/AdguardTeam/golibs/netutil"
"github.com/miekg/dns"
)
@@ -83,13 +82,16 @@ const ddrHostFQDN = "_dns.resolver.arpa."
// handleDNSRequest filters the incoming DNS requests and writes them to the query log
func (s *Server) handleDNSRequest(_ *proxy.Proxy, pctx *proxy.DNSContext) error {
// TODO(s.chzhen): Pass context.
ctx := context.TODO()
dctx := &dnsContext{
proxyCtx: pctx,
result: &filtering.Result{},
startTime: time.Now(),
}
type modProcessFunc func(ctx *dnsContext) (rc resultCode)
type modProcessFunc func(ctx context.Context, dctx *dnsContext) (rc resultCode)
// Since (*dnsforward.Server).handleDNSRequest(...) is used as
// proxy.(Config).RequestHandler, there is no need for additional index
@@ -108,7 +110,7 @@ func (s *Server) handleDNSRequest(_ *proxy.Proxy, pctx *proxy.DNSContext) error
s.processQueryLogsAndStats,
}
for _, process := range mods {
r := process(dctx)
r := process(ctx, dctx)
switch r {
case resultCodeSuccess:
// continue: call the next filter
@@ -149,12 +151,12 @@ const healthcheckFQDN = "healthcheck.adguardhome.test."
// needed and enriches dctx with some client-specific information.
//
// TODO(e.burkov): Decompose into less general processors.
func (s *Server) processInitial(dctx *dnsContext) (rc resultCode) {
log.Debug("dnsforward: started processing initial")
defer log.Debug("dnsforward: finished processing initial")
func (s *Server) processInitial(ctx context.Context, dctx *dnsContext) (rc resultCode) {
s.logger.DebugContext(ctx, "started processing initial")
defer s.logger.DebugContext(ctx, "finished processing initial")
pctx := dctx.proxyCtx
s.processClientIP(pctx.Addr.Addr())
s.processClientIP(ctx, pctx.Addr.Addr())
q := pctx.Req.Question[0]
qt := q.Qtype
@@ -184,16 +186,16 @@ func (s *Server) processInitial(dctx *dnsContext) (rc resultCode) {
dctx.clientID = string(s.clientIDCache.Get(key[:]))
// Get the client-specific filtering settings.
dctx.protectionEnabled, _ = s.UpdatedProtectionStatus()
dctx.protectionEnabled, _ = s.UpdatedProtectionStatus(ctx)
dctx.setts = s.clientRequestFilteringSettings(dctx)
return resultCodeSuccess
}
// processClientIP sends the client IP address to s.addrProc, if needed.
func (s *Server) processClientIP(addr netip.Addr) {
func (s *Server) processClientIP(ctx context.Context, addr netip.Addr) {
if !addr.IsValid() {
log.Info("dnsforward: warning: bad client addr %q", addr)
s.logger.WarnContext(ctx, "bad client address", "addr", addr)
return
}
@@ -203,8 +205,7 @@ func (s *Server) processClientIP(addr netip.Addr) {
s.serverLock.RLock()
defer s.serverLock.RUnlock()
// TODO(s.chzhen): Pass context.
s.addrProc.Process(context.TODO(), addr)
s.addrProc.Process(ctx, addr)
}
// processDDRQuery responds to Discovery of Designated Resolvers (DDR) SVCB
@@ -212,9 +213,9 @@ func (s *Server) processClientIP(addr netip.Addr) {
// current user configuration.
//
// See https://www.ietf.org/archive/id/draft-ietf-add-ddr-10.html.
func (s *Server) processDDRQuery(dctx *dnsContext) (rc resultCode) {
log.Debug("dnsforward: started processing ddr")
defer log.Debug("dnsforward: finished processing ddr")
func (s *Server) processDDRQuery(ctx context.Context, dctx *dnsContext) (rc resultCode) {
s.logger.DebugContext(ctx, "started processing ddr")
defer s.logger.DebugContext(ctx, "finished processing ddr")
if !s.conf.HandleDDR {
return resultCodeSuccess
@@ -311,9 +312,9 @@ func (s *Server) makeDDRResponse(req *dns.Msg) (resp *dns.Msg) {
// the request is for AAAA.
//
// TODO(a.garipov): Adapt to AAAA as well.
func (s *Server) processDHCPHosts(dctx *dnsContext) (rc resultCode) {
log.Debug("dnsforward: started processing dhcp hosts")
defer log.Debug("dnsforward: finished processing dhcp hosts")
func (s *Server) processDHCPHosts(ctx context.Context, dctx *dnsContext) (rc resultCode) {
s.logger.DebugContext(ctx, "started processing dhcp hosts")
defer s.logger.DebugContext(ctx, "finished processing dhcp hosts")
pctx := dctx.proxyCtx
req := pctx.Req
@@ -325,7 +326,12 @@ func (s *Server) processDHCPHosts(dctx *dnsContext) (rc resultCode) {
}
if !pctx.IsPrivateClient {
log.Debug("dnsforward: %q requests for dhcp host %q", pctx.Addr, dhcpHost)
s.logger.DebugContext(
ctx,
"requests for dhcp host",
"addr", pctx.Addr,
"dhcp_host", dhcpHost,
)
pctx.Res = s.NewMsgNXDOMAIN(req)
// Do not even put into query log.
@@ -336,12 +342,12 @@ func (s *Server) processDHCPHosts(dctx *dnsContext) (rc resultCode) {
if ip == (netip.Addr{}) {
// Go on and process them with filters, including dnsrewrite ones, and
// possibly route them to a domain-specific upstream.
log.Debug("dnsforward: no dhcp record for %q", dhcpHost)
s.logger.DebugContext(ctx, "no dhcp record", "dhcp_host", dhcpHost)
return resultCodeSuccess
}
log.Debug("dnsforward: dhcp record for %q is %s", dhcpHost, ip)
s.logger.DebugContext(ctx, "dhcp record for", "dhcp_host", dhcpHost, "ip", ip)
resp := s.replyCompressed(req)
switch q.Qtype {
@@ -372,9 +378,9 @@ func (s *Server) processDHCPHosts(dctx *dnsContext) (rc resultCode) {
// processDHCPAddrs responds to PTR requests if the target IP is leased by the
// DHCP server.
func (s *Server) processDHCPAddrs(dctx *dnsContext) (rc resultCode) {
log.Debug("dnsforward: started processing dhcp addrs")
defer log.Debug("dnsforward: finished processing dhcp addrs")
func (s *Server) processDHCPAddrs(ctx context.Context, dctx *dnsContext) (rc resultCode) {
s.logger.DebugContext(ctx, "started processing dhcp addrs")
defer s.logger.DebugContext(ctx, "finished processing dhcp addrs")
pctx := dctx.proxyCtx
if pctx.Res != nil {
@@ -396,7 +402,7 @@ func (s *Server) processDHCPAddrs(dctx *dnsContext) (rc resultCode) {
return resultCodeSuccess
}
log.Debug("dnsforward: dhcp client %s is %q", addr, host)
s.logger.DebugContext(ctx, "dhcp client", "addr", addr, "host", host)
resp := s.replyCompressed(req)
ptr := &dns.PTR{
@@ -417,9 +423,12 @@ func (s *Server) processDHCPAddrs(dctx *dnsContext) (rc resultCode) {
}
// Apply filtering logic
func (s *Server) processFilteringBeforeRequest(dctx *dnsContext) (rc resultCode) {
log.Debug("dnsforward: started processing filtering before req")
defer log.Debug("dnsforward: finished processing filtering before req")
func (s *Server) processFilteringBeforeRequest(
ctx context.Context,
dctx *dnsContext,
) (rc resultCode) {
s.logger.DebugContext(ctx, "started processing filtering before request")
defer s.logger.DebugContext(ctx, "finished processing filtering before request")
if dctx.proxyCtx.RequestedPrivateRDNS != (netip.Prefix{}) {
// There is no need to filter request for locally served ARPA hostname
@@ -439,7 +448,7 @@ func (s *Server) processFilteringBeforeRequest(dctx *dnsContext) (rc resultCode)
defer s.serverLock.RUnlock()
var err error
if dctx.result, err = s.filterDNSRequest(dctx); err != nil {
if dctx.result, err = s.filterDNSRequest(ctx, dctx); err != nil {
dctx.err = err
return resultCodeError
@@ -458,9 +467,9 @@ func ipStringFromAddr(addr net.Addr) (ipStr string) {
}
// processUpstream passes request to upstream servers and handles the response.
func (s *Server) processUpstream(dctx *dnsContext) (rc resultCode) {
log.Debug("dnsforward: started processing upstream")
defer log.Debug("dnsforward: finished processing upstream")
func (s *Server) processUpstream(ctx context.Context, dctx *dnsContext) (rc resultCode) {
s.logger.DebugContext(ctx, "started processing upstream")
defer s.logger.DebugContext(ctx, "finished processing upstream")
pctx := dctx.proxyCtx
req := pctx.Req
@@ -475,13 +484,17 @@ func (s *Server) processUpstream(dctx *dnsContext) (rc resultCode) {
// TODO(a.garipov): Route such queries to a custom upstream for the
// local domain name if there is one.
name := req.Question[0].Name
log.Debug("dnsforward: dhcp client hostname %q was not filtered", name[:len(name)-1])
s.logger.DebugContext(
ctx,
"dhcp client hostname was not filtered",
"hostname", name[:len(name)-1],
)
pctx.Res = s.NewMsgNXDOMAIN(req)
return resultCodeFinish
}
s.setCustomUpstream(pctx, dctx.clientID)
s.setCustomUpstream(ctx, pctx, dctx.clientID)
reqWantsDNSSEC := s.setReqAD(req)
@@ -571,7 +584,7 @@ func (s *Server) dhcpHostFromRequest(q *dns.Question) (reqHost string) {
}
// setCustomUpstream sets custom upstream settings in pctx, if necessary.
func (s *Server) setCustomUpstream(pctx *proxy.DNSContext, clientID string) {
func (s *Server) setCustomUpstream(ctx context.Context, pctx *proxy.DNSContext, clientID string) {
if !pctx.Addr.IsValid() || s.conf.ClientsContainer == nil {
return
}
@@ -579,10 +592,11 @@ func (s *Server) setCustomUpstream(pctx *proxy.DNSContext, clientID string) {
cliAddr := pctx.Addr.Addr()
upsConf := s.conf.ClientsContainer.CustomUpstreamConfig(clientID, cliAddr)
if upsConf != nil {
log.Debug(
"dnsforward: using custom upstreams for client with ip %s and clientid %q",
cliAddr,
clientID,
s.logger.DebugContext(
ctx,
"using custom upstreams for client with",
"ip", cliAddr,
"client_id", clientID,
)
pctx.CustomUpstreamConfig = upsConf
@@ -590,9 +604,9 @@ func (s *Server) setCustomUpstream(pctx *proxy.DNSContext, clientID string) {
}
// Apply filtering logic after we have received response from upstream servers
func (s *Server) processFilteringAfterResponse(dctx *dnsContext) (rc resultCode) {
log.Debug("dnsforward: started processing filtering after resp")
defer log.Debug("dnsforward: finished processing filtering after resp")
func (s *Server) processFilteringAfterResponse(ctx context.Context, dctx *dnsContext) (rc resultCode) {
s.logger.DebugContext(ctx, "started processing filtering after response")
defer s.logger.DebugContext(ctx, "finished processing filtering after response")
switch res := dctx.result; res.Reason {
case filtering.NotFilteredAllowList:
@@ -617,13 +631,13 @@ func (s *Server) processFilteringAfterResponse(dctx *dnsContext) (rc resultCode)
return resultCodeSuccess
default:
return s.filterAfterResponse(dctx)
return s.filterAfterResponse(ctx, dctx)
}
}
// filterAfterResponse returns the result of filtering the response that wasn't
// explicitly allowed or rewritten.
func (s *Server) filterAfterResponse(dctx *dnsContext) (res resultCode) {
func (s *Server) filterAfterResponse(ctx context.Context, dctx *dnsContext) (res resultCode) {
// Check the response only if it's from an upstream. Don't check the
// response if the protection is disabled since dnsrewrite rules aren't
// applied to it anyway.
@@ -631,7 +645,7 @@ func (s *Server) filterAfterResponse(dctx *dnsContext) (res resultCode) {
return resultCodeSuccess
}
err := s.filterDNSResponse(dctx)
err := s.filterDNSResponse(ctx, dctx)
if err != nil {
dctx.err = err

View File

@@ -94,7 +94,7 @@ func TestServer_ProcessInitial(t *testing.T) {
var gotAddr netip.Addr
s.addrProc = &aghtest.AddressProcessor{
OnProcess: func(ctx context.Context, ip netip.Addr) { gotAddr = ip },
OnClose: func() (err error) { panic("not implemented") },
OnClose: func() (_ error) { panic(testutil.UnexpectedCall()) },
}
dctx := &dnsContext{
@@ -105,7 +105,7 @@ func TestServer_ProcessInitial(t *testing.T) {
},
}
gotRC := s.processInitial(dctx)
gotRC := s.processInitial(testutil.ContextWithTimeout(t, testTimeout), dctx)
assert.Equal(t, tc.wantRC, gotRC)
assert.Equal(t, testClientAddrPort.Addr(), gotAddr)
@@ -208,8 +208,8 @@ func TestServer_ProcessFilteringAfterResponse(t *testing.T) {
Addr: testClientAddrPort,
},
}
gotRC := s.processFilteringAfterResponse(dctx)
ctx := testutil.ContextWithTimeout(t, testTimeout)
gotRC := s.processFilteringAfterResponse(ctx, dctx)
assert.Equal(t, tc.wantRC, gotRC)
assert.Equal(t, newResp(dns.RcodeSuccess, tc.req, tc.wantRespAns), dctx.proxyCtx.Res)
})
@@ -353,7 +353,7 @@ func TestServer_ProcessDDRQuery(t *testing.T) {
},
}
res := s.processDDRQuery(dctx)
res := s.processDDRQuery(testutil.ContextWithTimeout(t, testTimeout), dctx)
require.Equal(t, tc.wantRes, res)
if tc.wantRes != resultCodeFinish {
@@ -373,14 +373,14 @@ func TestServer_ProcessDDRQuery(t *testing.T) {
}
// createTestDNSFilter returns the minimum valid DNSFilter.
func createTestDNSFilter(t *testing.T) (f *filtering.DNSFilter) {
t.Helper()
func createTestDNSFilter(tb testing.TB) (f *filtering.DNSFilter) {
tb.Helper()
f, err := filtering.New(&filtering.Config{
Logger: testLogger,
BlockingMode: filtering.BlockingModeDefault,
}, []filtering.Filter{})
require.NoError(t, err)
require.NoError(tb, err)
return f
}
@@ -440,6 +440,7 @@ func TestServer_ProcessDHCPHosts_localRestriction(t *testing.T) {
dhcpServer: dhcp,
localDomainSuffix: localDomainSuffix,
baseLogger: testLogger,
logger: testLogger,
}
req := &dns.Msg{
@@ -460,7 +461,7 @@ func TestServer_ProcessDHCPHosts_localRestriction(t *testing.T) {
},
}
res := s.processDHCPHosts(dctx)
res := s.processDHCPHosts(testutil.ContextWithTimeout(t, testTimeout), dctx)
pctx := dctx.proxyCtx
if !tc.isLocalCli {
@@ -518,7 +519,7 @@ func TestServer_ProcessDHCPHosts(t *testing.T) {
OnIPByHost: func(host string) (ip netip.Addr) {
return knownClients[host]
},
OnHostByIP: func(ip netip.Addr) (host string) { panic("not implemented") },
OnHostByIP: func(ip netip.Addr) (_ string) { panic(testutil.UnexpectedCall(ip)) },
}
testCases := []struct {
@@ -592,6 +593,7 @@ func TestServer_ProcessDHCPHosts(t *testing.T) {
dhcpServer: testDHCP,
localDomainSuffix: tc.suffix,
baseLogger: testLogger,
logger: testLogger,
}
req := (&dns.Msg{}).SetQuestion(dns.Fqdn(tc.host), tc.qtyp)
@@ -604,7 +606,7 @@ func TestServer_ProcessDHCPHosts(t *testing.T) {
}
t.Run(tc.name, func(t *testing.T) {
res := s.processDHCPHosts(dctx)
res := s.processDHCPHosts(testutil.ContextWithTimeout(t, testTimeout), dctx)
pctx := dctx.proxyCtx
assert.Equal(t, tc.wantRes, res)
require.NoError(t, dctx.err)
@@ -812,9 +814,9 @@ func TestServer_ProcessUpstream_localPTR(t *testing.T) {
ServePlainDNS: true,
},
)
ctx := testutil.ContextWithTimeout(t, testTimeout)
pctx := newPrxCtx()
rc := s.processUpstream(&dnsContext{proxyCtx: pctx})
rc := s.processUpstream(ctx, &dnsContext{proxyCtx: pctx})
require.Equal(t, resultCodeSuccess, rc)
require.NotEmpty(t, pctx.Res.Answer)
ptr := testutil.RequireTypeAssert[*dns.PTR](t, pctx.Res.Answer[0])
@@ -844,7 +846,8 @@ func TestServer_ProcessUpstream_localPTR(t *testing.T) {
)
pctx := newPrxCtx()
rc := s.processUpstream(&dnsContext{proxyCtx: pctx})
ctx := testutil.ContextWithTimeout(t, testTimeout)
rc := s.processUpstream(ctx, &dnsContext{proxyCtx: pctx})
require.Equal(t, resultCodeError, rc)
require.Empty(t, pctx.Res.Answer)
})

View File

@@ -1,6 +1,7 @@
package dnsforward
import (
"context"
"net"
"time"
@@ -9,14 +10,13 @@ import (
"github.com/AdguardTeam/AdGuardHome/internal/querylog"
"github.com/AdguardTeam/AdGuardHome/internal/stats"
"github.com/AdguardTeam/dnsproxy/proxy"
"github.com/AdguardTeam/golibs/log"
"github.com/miekg/dns"
)
// Write Stats data and logs
func (s *Server) processQueryLogsAndStats(dctx *dnsContext) (rc resultCode) {
log.Debug("dnsforward: started processing querylog and stats")
defer log.Debug("dnsforward: finished processing querylog and stats")
func (s *Server) processQueryLogsAndStats(ctx context.Context, dctx *dnsContext) (rc resultCode) {
s.logger.DebugContext(ctx, "started processing querylog and stats")
defer s.logger.DebugContext(ctx, "finished processing querylog and stats")
pctx := dctx.proxyCtx
q := pctx.Req.Question[0]
@@ -27,7 +27,7 @@ func (s *Server) processQueryLogsAndStats(dctx *dnsContext) (rc resultCode) {
s.anonymizer.Load()(ip)
ipStr := net.IP(ip).String()
log.Debug("dnsforward: client ip for stats and querylog: %s", ipStr)
s.logger.DebugContext(ctx, "client ip for stats and querylog", "ip", ipStr)
ids := []string{ipStr}
if dctx.clientID != "" {
@@ -47,24 +47,26 @@ func (s *Server) processQueryLogsAndStats(dctx *dnsContext) (rc resultCode) {
if s.shouldLog(host, qt, cl, ids) {
s.logQuery(dctx, ip, processingTime)
} else {
log.Debug(
"dnsforward: request %s %s %q from %s ignored; not adding to querylog",
dns.Class(cl),
dns.Type(qt),
host,
ipStr,
s.logger.DebugContext(
ctx,
"not adding to querylog",
"dns_class", dns.Class(cl),
"dns_type", dns.Type(qt),
"host", host,
"ip", ipStr,
)
}
if s.shouldCountStat(host, qt, cl, ids) {
s.updateStats(dctx, ipStr, processingTime)
} else {
log.Debug(
"dnsforward: request %s %s %q from %s ignored; not counting in stats",
dns.Class(cl),
dns.Type(qt),
host,
ipStr,
s.logger.DebugContext(
ctx,
"not counting in stats",
"dns_class", dns.Class(cl),
"dns_type", dns.Type(qt),
"host", host,
"ip", ipStr,
)
}

View File

@@ -11,6 +11,7 @@ import (
"github.com/AdguardTeam/AdGuardHome/internal/stats"
"github.com/AdguardTeam/dnsproxy/proxy"
"github.com/AdguardTeam/dnsproxy/upstream"
"github.com/AdguardTeam/golibs/testutil"
"github.com/miekg/dns"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -203,6 +204,7 @@ func TestServer_ProcessQueryLogsAndStats(t *testing.T) {
st := &testStats{}
srv := &Server{
baseLogger: testLogger,
logger: testLogger,
queryLog: ql,
stats: st,
anonymizer: aghnet.NewIPMut(nil),
@@ -229,7 +231,7 @@ func TestServer_ProcessQueryLogsAndStats(t *testing.T) {
clientID: tc.clientID,
}
code := srv.processQueryLogsAndStats(dctx)
code := srv.processQueryLogsAndStats(testutil.ContextWithTimeout(t, testTimeout), dctx)
assert.Equal(t, tc.wantCode, code)
assert.Equal(t, tc.wantLogProto, ql.lastParams.ClientProto)
assert.Equal(t, tc.wantStatClient, st.lastEntry.Client)

View File

@@ -1,6 +1,7 @@
package dnsforward
import (
"context"
"encoding/base64"
"net"
"strconv"
@@ -14,9 +15,9 @@ import (
//
// See the comment on genAnswerSVCB for a list of current restrictions on
// parameter values.
func (s *Server) genAnswerHTTPS(req *dns.Msg, svcb *rules.DNSSVCB) (ans *dns.HTTPS) {
func (s *Server) genAnswerHTTPS(ctx context.Context, req *dns.Msg, svcb *rules.DNSSVCB) (ans *dns.HTTPS) {
ans = &dns.HTTPS{
SVCB: *s.genAnswerSVCB(req, svcb),
SVCB: *s.genAnswerSVCB(ctx, req, svcb),
}
ans.Hdr.Rrtype = dns.TypeHTTPS
@@ -163,7 +164,11 @@ var svcbKeyHandlers = map[string]svcbKeyHandler{
// ipv4hint="127.0.0.1,127.0.0.2" // Unsupported.
//
// TODO(a.garipov): Support all of these.
func (s *Server) genAnswerSVCB(req *dns.Msg, svcb *rules.DNSSVCB) (ans *dns.SVCB) {
func (s *Server) genAnswerSVCB(
ctx context.Context,
req *dns.Msg,
svcb *rules.DNSSVCB,
) (ans *dns.SVCB) {
ans = &dns.SVCB{
Hdr: s.hdr(req, dns.TypeSVCB),
Priority: svcb.Priority,
@@ -177,7 +182,7 @@ func (s *Server) genAnswerSVCB(req *dns.Msg, svcb *rules.DNSSVCB) (ans *dns.SVCB
for k, valStr := range svcb.Params {
handler, ok := svcbKeyHandlers[k]
if !ok {
log.Debug("unknown svcb/https key %q, ignoring", k)
s.logger.DebugContext(ctx, "unknown svcb/https key, ignoring", "key", k)
continue
}

View File

@@ -5,6 +5,7 @@ import (
"testing"
"github.com/AdguardTeam/AdGuardHome/internal/filtering"
"github.com/AdguardTeam/golibs/testutil"
"github.com/AdguardTeam/urlfilter/rules"
"github.com/miekg/dns"
"github.com/stretchr/testify/assert"
@@ -152,14 +153,14 @@ func TestGenAnswerHTTPS_andSVCB(t *testing.T) {
want := &dns.HTTPS{SVCB: *tc.want}
want.Hdr.Rrtype = dns.TypeHTTPS
got := s.genAnswerHTTPS(req, tc.svcb)
got := s.genAnswerHTTPS(testutil.ContextWithTimeout(t, testTimeout), req, tc.svcb)
assert.Equal(t, want, got)
})
})
t.Run("svcb", func(t *testing.T) {
t.Run(tc.name, func(t *testing.T) {
got := s.genAnswerSVCB(req, tc.svcb)
got := s.genAnswerSVCB(testutil.ContextWithTimeout(t, testTimeout), req, tc.svcb)
assert.Equal(t, tc.want, got)
})
})

View File

@@ -32,6 +32,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -72,6 +73,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -112,6 +114,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,

View File

@@ -37,6 +37,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -79,6 +80,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -122,6 +124,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -165,6 +168,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -208,6 +212,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -253,6 +258,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -299,6 +305,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -342,6 +349,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -387,6 +395,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -432,6 +441,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -475,6 +485,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -518,6 +529,52 @@
"cache_size": 1024,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": true,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
"local_ptr_upstreams": [],
"edns_cs_use_custom": false,
"edns_cs_custom_ip": ""
}
},
"cache_enabled": {
"req": {
"cache_enabled": true,
"cache_size": 1024
},
"want": {
"upstream_dns": [
"8.8.8.8:53",
"8.8.4.4:53"
],
"upstream_dns_file": "",
"bootstrap_dns": [
"9.9.9.10",
"149.112.112.10",
"2620:fe::10",
"2620:fe::fe:10"
],
"fallback_dns": [],
"protection_enabled": true,
"protection_disabled_until": null,
"ratelimit": 0,
"ratelimit_subnet_len_ipv4": 24,
"ratelimit_subnet_len_ipv6": 56,
"ratelimit_whitelist": [],
"blocking_mode": "default",
"blocking_ipv4": "",
"blocking_ipv6": "",
"blocked_response_ttl": 10,
"upstream_timeout": 10,
"edns_cs_enabled": false,
"dnssec_enabled": false,
"disable_ipv6": false,
"upstream_mode": "",
"cache_size": 1024,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": true,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -561,6 +618,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -604,6 +662,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -649,6 +708,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -694,6 +754,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -738,6 +799,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -781,6 +843,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -826,6 +889,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -874,6 +938,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -917,6 +982,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -964,6 +1030,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -1007,6 +1074,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -1053,6 +1121,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,
@@ -1096,6 +1165,7 @@
"cache_size": 0,
"cache_ttl_min": 0,
"cache_ttl_max": 0,
"cache_enabled": false,
"cache_optimistic": false,
"resolve_clients": false,
"use_private_ptr_resolvers": false,

View File

@@ -159,6 +159,8 @@ func (d *DNSFilter) handleBlockedServicesList(w http.ResponseWriter, r *http.Req
//
// Deprecated: Use handleBlockedServicesUpdate.
func (d *DNSFilter) handleBlockedServicesSet(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
list := []string{}
err := json.NewDecoder(r.Body).Decode(&list)
if err != nil {
@@ -172,10 +174,10 @@ func (d *DNSFilter) handleBlockedServicesSet(w http.ResponseWriter, r *http.Requ
defer d.confMu.Unlock()
d.conf.BlockedServices.IDs = list
d.logger.DebugContext(r.Context(), "updated blocked services list", "len", len(list))
d.logger.DebugContext(ctx, "updated blocked services list", "len", len(list))
}()
d.conf.ConfigModified()
d.conf.ConfModifier.Apply(ctx)
}
// handleBlockedServicesGet is the handler for the GET
@@ -195,6 +197,8 @@ func (d *DNSFilter) handleBlockedServicesGet(w http.ResponseWriter, r *http.Requ
// handleBlockedServicesUpdate is the handler for the PUT
// /control/blocked_services/update HTTP API.
func (d *DNSFilter) handleBlockedServicesUpdate(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
bsvc := &BlockedServices{}
err := json.NewDecoder(r.Body).Decode(bsvc)
if err != nil {
@@ -221,7 +225,7 @@ func (d *DNSFilter) handleBlockedServicesUpdate(w http.ResponseWriter, r *http.R
d.conf.BlockedServices = bsvc
}()
d.logger.DebugContext(r.Context(), "updated blocked services schedule", "len", len(bsvc.IDs))
d.logger.DebugContext(ctx, "updated blocked services schedule", "len", len(bsvc.IDs))
d.conf.ConfigModified()
d.conf.ConfModifier.Apply(ctx)
}

View File

@@ -144,21 +144,22 @@ func (d *DNSFilter) filterSetProperties(
shouldRestart = true
}
if flt.Enabled {
if shouldRestart {
// Download the filter contents.
shouldRestart, err = d.update(flt)
}
} else {
if !flt.Enabled {
// TODO(e.burkov): The validation of the contents of the new URL is
// currently skipped if the rule list is disabled. This makes it
// possible to set a bad rules source, but the validation should still
// kick in when the filter is enabled. Consider changing this behavior
// to be stricter.
flt.unload()
return shouldRestart, err
}
return shouldRestart, err
if !shouldRestart {
return false, nil
}
return d.update(flt)
}
// filterExists returns true if a filter with the same url exists in d. It's
@@ -315,19 +316,7 @@ func (d *DNSFilter) refreshFiltersArray(
return 0, nil, nil, false
}
failNum := 0
for i := range updateFilters {
uf := &updateFilters[i]
updated, err := d.update(uf)
updateFlags = append(updateFlags, updated)
if err != nil {
failNum++
d.logger.ErrorContext(ctx, "updating filter", "url", uf.URL, slogutil.KeyError, err)
continue
}
}
failNum, updateFlags := d.updateFilterList(ctx, updateFilters)
if failNum == len(updateFilters) {
return 0, nil, nil, true
}
@@ -335,6 +324,40 @@ func (d *DNSFilter) refreshFiltersArray(
d.conf.filtersMu.Lock()
defer d.conf.filtersMu.Unlock()
updateCount = d.syncUpdatedFilters(ctx, filters, updateFilters, updateFlags)
return updateCount, updateFilters, updateFlags, false
}
// updateFilterList updates each filter in updateFilters and returns the number
// of failures and the updateFlags slice aligned with updateFilters indicating
// whether each filter's data changed.
func (d *DNSFilter) updateFilterList(
ctx context.Context,
updateFilters []FilterYAML,
) (failNum int, updateFlags []bool) {
for i := range updateFilters {
uf := &updateFilters[i]
updated, err := d.update(uf)
updateFlags = append(updateFlags, updated)
if err != nil {
failNum++
d.logger.ErrorContext(ctx, "updating filter", "url", uf.URL, slogutil.KeyError, err)
}
}
return failNum, updateFlags
}
// syncUpdatedFilters syncs updated filters back to the original filters slice
// and returns the updateCount. filters must not be nil. updateFlags must
// align with updateFilters. d.conf.filtersMu must be locked.
func (d *DNSFilter) syncUpdatedFilters(
ctx context.Context,
filters *[]FilterYAML,
updateFilters []FilterYAML,
updateFlags []bool,
) (updateCount int) {
for i := range updateFilters {
uf := &updateFilters[i]
updated := updateFlags[i]
@@ -365,7 +388,7 @@ func (d *DNSFilter) refreshFiltersArray(
}
}
return updateCount, updateFilters, updateFlags, false
return updateCount
}
// refreshFiltersIntl checks filters and updates them if necessary. If force is
@@ -418,21 +441,23 @@ func (d *DNSFilter) refreshFiltersIntl(block, allow, force bool) (int, bool) {
return 0, true
}
if updNum != 0 {
d.EnableFilters(false)
if updNum == 0 {
return 0, false
}
for i := range lists {
uf := &lists[i]
updated := toUpd[i]
if !updated {
continue
}
d.EnableFilters(false)
p := uf.Path(d.conf.DataDir)
err := os.Remove(p + ".old")
if err != nil {
d.logger.ErrorContext(ctx, "removing old filter", "path", p, slogutil.KeyError, err)
}
for i := range lists {
uf := &lists[i]
updated := toUpd[i]
if !updated {
continue
}
p := uf.Path(d.conf.DataDir)
err := os.Remove(p + ".old")
if err != nil {
d.logger.ErrorContext(ctx, "removing old filter", "path", p, slogutil.KeyError, err)
}
}

View File

@@ -22,17 +22,16 @@ const testTimeout = 5 * time.Second
// serveHTTPLocally starts a new HTTP server, that handles its index with h. It
// also gracefully closes the listener when the test under t finishes.
func serveHTTPLocally(t *testing.T, h http.Handler) (urlStr string) {
t.Helper()
func serveHTTPLocally(tb testing.TB, h http.Handler) (urlStr string) {
tb.Helper()
l, err := net.Listen("tcp", ":0")
require.NoError(t, err)
require.NoError(tb, err)
go func() { _ = http.Serve(l, h) }()
testutil.CleanupAndRequireSuccess(t, l.Close)
testutil.CleanupAndRequireSuccess(tb, l.Close)
addr := l.Addr()
require.IsType(t, (*net.TCPAddr)(nil), addr)
addr := testutil.RequireTypeAssert[*net.TCPAddr](tb, l.Addr())
return (&url.URL{
Scheme: urlutil.SchemeHTTP,
@@ -42,10 +41,10 @@ func serveHTTPLocally(t *testing.T, h http.Handler) (urlStr string) {
// serveFiltersLocally is a helper that concurrently listens on a free port to
// respond with fltContent.
func serveFiltersLocally(t *testing.T, fltContent []byte) (urlStr string) {
t.Helper()
func serveFiltersLocally(tb testing.TB, fltContent []byte) (urlStr string) {
tb.Helper()
return serveHTTPLocally(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
return serveHTTPLocally(tb, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
pt := testutil.PanicT{}
n, werr := w.Write(fltContent)
@@ -57,43 +56,43 @@ func serveFiltersLocally(t *testing.T, fltContent []byte) (urlStr string) {
// updateAndAssert loads filter content from its URL and then asserts rules
// count.
func updateAndAssert(
t *testing.T,
tb testing.TB,
ctx context.Context,
dnsFilter *DNSFilter,
f *FilterYAML,
wantUpd require.BoolAssertionFunc,
wantRulesCount int,
) {
t.Helper()
tb.Helper()
ok, err := dnsFilter.update(f)
require.NoError(t, err)
wantUpd(t, ok)
require.NoError(tb, err)
wantUpd(tb, ok)
assert.Equal(t, wantRulesCount, f.RulesCount)
assert.Equal(tb, wantRulesCount, f.RulesCount)
dir, err := os.ReadDir(filepath.Join(dnsFilter.conf.DataDir, filterDir))
require.NoError(t, err)
require.FileExists(t, f.Path(dnsFilter.conf.DataDir))
require.NoError(tb, err)
require.FileExists(tb, f.Path(dnsFilter.conf.DataDir))
assert.Len(t, dir, 1)
assert.Len(tb, dir, 1)
err = dnsFilter.load(ctx, f)
require.NoError(t, err)
require.NoError(tb, err)
}
// newDNSFilter returns a new properly initialized DNS filter instance.
func newDNSFilter(t *testing.T) (d *DNSFilter) {
t.Helper()
func newDNSFilter(tb testing.TB) (d *DNSFilter) {
tb.Helper()
dnsFilter, err := New(&Config{
Logger: slogutil.NewDiscardLogger(),
DataDir: t.TempDir(),
DataDir: tb.TempDir(),
HTTPClient: &http.Client{
Timeout: testTimeout,
},
}, nil)
require.NoError(t, err)
require.NoError(tb, err)
return dnsFilter
}

View File

@@ -19,6 +19,7 @@ import (
"sync/atomic"
"time"
"github.com/AdguardTeam/AdGuardHome/internal/agh"
"github.com/AdguardTeam/AdGuardHome/internal/aghhttp"
"github.com/AdguardTeam/AdGuardHome/internal/aghos"
"github.com/AdguardTeam/AdGuardHome/internal/filtering/rulelist"
@@ -104,8 +105,9 @@ type Config struct {
// TODO(e.burkov): Move it to dnsforward entirely.
EtcHosts hostsfile.Storage `yaml:"-"`
// Called when the configuration is changed by HTTP request
ConfigModified func() `yaml:"-"`
// ConfModifier is used to update the global configuration. It must not be
// nil.
ConfModifier agh.ConfigModifier `yaml:"-"`
// Register an HTTP handler
HTTPRegister aghhttp.RegisterFunc `yaml:"-"`
@@ -766,42 +768,19 @@ func (d *DNSFilter) matchBlockedServicesRules(
func newRuleStorage(filters []Filter) (rs *filterlist.RuleStorage, err error) {
lists := make([]filterlist.RuleList, 0, len(filters))
for _, f := range filters {
switch id := int(f.ID); {
case len(f.Data) != 0:
lists = append(lists, &filterlist.StringRuleList{
ID: id,
RulesText: string(f.Data),
IgnoreCosmetic: true,
})
case f.FilePath == "":
var rl filterlist.RuleList
var skip bool
rl, skip, err = ruleListFromFilter(f)
if skip {
continue
case runtime.GOOS == "windows":
// On Windows we don't pass a file to urlfilter because it's
// difficult to update this file while it's being used.
var data []byte
data, err = os.ReadFile(f.FilePath)
if errors.Is(err, fs.ErrNotExist) {
continue
} else if err != nil {
return nil, fmt.Errorf("reading filter content: %w", err)
}
lists = append(lists, &filterlist.StringRuleList{
ID: id,
RulesText: string(data),
IgnoreCosmetic: true,
})
default:
var list *filterlist.FileRuleList
list, err = filterlist.NewFileRuleList(id, f.FilePath, true)
if errors.Is(err, fs.ErrNotExist) {
continue
} else if err != nil {
return nil, fmt.Errorf("creating file rule list with %q: %w", f.FilePath, err)
}
lists = append(lists, list)
}
if err != nil {
// Don't wrap the error, because it's informative enough as is.
return nil, err
}
lists = append(lists, rl)
}
rs, err = filterlist.NewRuleStorage(lists)
@@ -812,6 +791,51 @@ func newRuleStorage(filters []Filter) (rs *filterlist.RuleStorage, err error) {
return rs, nil
}
// ruleListFromFilter returns a rule list from a Filter.
func ruleListFromFilter(f Filter) (rl filterlist.RuleList, skip bool, err error) {
id := int(f.ID)
if len(f.Data) != 0 {
return &filterlist.StringRuleList{
ID: id,
RulesText: string(f.Data),
IgnoreCosmetic: true,
}, false, nil
}
if f.FilePath == "" {
return nil, true, nil
}
if runtime.GOOS == "windows" {
// On Windows we don't pass a file to urlfilter because it's
// difficult to update this file while it's being used.
var data []byte
data, err = os.ReadFile(f.FilePath)
if errors.Is(err, fs.ErrNotExist) {
return nil, true, nil
} else if err != nil {
return nil, false, fmt.Errorf("reading filter content: %w", err)
}
return &filterlist.StringRuleList{
ID: id,
RulesText: string(data),
IgnoreCosmetic: true,
}, false, nil
}
var list *filterlist.FileRuleList
list, err = filterlist.NewFileRuleList(id, f.FilePath, true)
if errors.Is(err, fs.ErrNotExist) {
return nil, true, nil
} else if err != nil {
return nil, false, fmt.Errorf("creating file rule list with %q: %w", f.FilePath, err)
}
return list, false, nil
}
// Initialize urlfilter objects.
func (d *DNSFilter) initFiltering(ctx context.Context, allowFilters, blockFilters []Filter) (err error) {
rulesStorage, err := newRuleStorage(blockFilters)
@@ -904,32 +928,37 @@ func (d *DNSFilter) matchHostProcessDNSResult(
return makeResult([]rules.Rule{dnsres.NetworkRule}, reason)
}
switch qtype {
case dns.TypeA:
if dnsres.HostRulesV4 != nil {
res = makeResult(hostRulesToRules(dnsres.HostRulesV4), FilteredBlockList)
for i, hr := range dnsres.HostRulesV4 {
res.Rules[i].IP = hr.IP
}
return res
}
case dns.TypeAAAA:
if dnsres.HostRulesV6 != nil {
res = makeResult(hostRulesToRules(dnsres.HostRulesV6), FilteredBlockList)
for i, hr := range dnsres.HostRulesV6 {
res.Rules[i].IP = hr.IP
}
return res
}
default:
// Go on.
if result, ok := resultFromHostRules(qtype, dnsres); ok {
return result
}
return hostResultForOtherQType(dnsres)
}
// resultFromHostRules handles the HostRulesV4/HostRulesV6 case for
// [matchHostProcessDNSResult]. dnsres must not be nil.
func resultFromHostRules(qtype uint16, dnsres *urlfilter.DNSResult) (res Result, ok bool) {
if qtype == dns.TypeA && dnsres.HostRulesV4 != nil {
res = makeResult(hostRulesToRules(dnsres.HostRulesV4), FilteredBlockList)
for i, hr := range dnsres.HostRulesV4 {
res.Rules[i].IP = hr.IP
}
return res, true
}
if qtype == dns.TypeAAAA && dnsres.HostRulesV6 != nil {
res = makeResult(hostRulesToRules(dnsres.HostRulesV6), FilteredBlockList)
for i, hr := range dnsres.HostRulesV6 {
res.Rules[i].IP = hr.IP
}
return res, true
}
return Result{}, false
}
// hostResultForOtherQType returns a result based on the host rules in dnsres,
// if any. dnsres.HostRulesV4 take precedence over dnsres.HostRulesV6.
func hostResultForOtherQType(dnsres *urlfilter.DNSResult) (res Result) {
@@ -1051,14 +1080,10 @@ func New(c *Config, blockFilters []Filter) (d *DNSFilter, err error) {
confMu: &sync.RWMutex{},
}
for i, p := range c.SafeFSPatterns {
// Use Match to validate the patterns here.
_, err = filepath.Match(p, "test")
if err != nil {
return nil, fmt.Errorf("safe_fs_patterns: at index %d: %w", i, err)
}
d.safeFSPatterns = append(d.safeFSPatterns, p)
err = d.validateSafeFSPatterns(c.SafeFSPatterns)
if err != nil {
// Don't wrap the error, because it's informative enough as is.
return nil, err
}
d.hostCheckers = []hostChecker{{
@@ -1126,6 +1151,22 @@ func New(c *Config, blockFilters []Filter) (d *DNSFilter, err error) {
return d, nil
}
// validateSafeFSPatterns validates and stores patterns for local filteringrule
// files.
func (d *DNSFilter) validateSafeFSPatterns(patterns []string) (err error) {
for i, p := range patterns {
// Use Match to validate the patterns here.
_, err = filepath.Match(p, "test")
if err != nil {
return fmt.Errorf("safe_fs_patterns: at index %d: %w", i, err)
}
d.safeFSPatterns = append(d.safeFSPatterns, p)
}
return nil
}
// Start registers web handlers and starts filters updates loop.
func (d *DNSFilter) Start() {
d.filtersInitializerChan = make(chan filtersInitializerParams, 1)

View File

@@ -63,37 +63,37 @@ func newChecker(host string) Checker {
})
}
func (d *DNSFilter) checkMatch(t *testing.T, hostname string, setts *Settings) {
t.Helper()
func (d *DNSFilter) checkMatch(tb testing.TB, hostname string, setts *Settings) {
tb.Helper()
res, err := d.CheckHost(hostname, dns.TypeA, setts)
require.NoErrorf(t, err, "host %q", hostname)
require.NoErrorf(tb, err, "host %q", hostname)
assert.Truef(t, res.IsFiltered, "host %q", hostname)
assert.Truef(tb, res.IsFiltered, "host %q", hostname)
}
func (d *DNSFilter) checkMatchIP(t *testing.T, hostname, ip string, qtype uint16, setts *Settings) {
t.Helper()
func (d *DNSFilter) checkMatchIP(tb testing.TB, hostname, ip string, qtype uint16, setts *Settings) {
tb.Helper()
res, err := d.CheckHost(hostname, qtype, setts)
require.NoErrorf(t, err, "host %q", hostname, err)
require.NotEmpty(t, res.Rules, "host %q", hostname)
require.NoErrorf(tb, err, "host %q", hostname, err)
require.NotEmpty(tb, res.Rules, "host %q", hostname)
assert.Truef(t, res.IsFiltered, "host %q", hostname)
assert.Truef(tb, res.IsFiltered, "host %q", hostname)
r := res.Rules[0]
require.NotNilf(t, r.IP, "Expected ip %s to match, actual: %v", ip, r.IP)
require.NotNilf(tb, r.IP, "Expected ip %s to match, actual: %v", ip, r.IP)
assert.Equalf(t, ip, r.IP.String(), "host %q", hostname)
assert.Equalf(tb, ip, r.IP.String(), "host %q", hostname)
}
func (d *DNSFilter) checkMatchEmpty(t *testing.T, hostname string, setts *Settings) {
t.Helper()
func (d *DNSFilter) checkMatchEmpty(tb testing.TB, hostname string, setts *Settings) {
tb.Helper()
res, err := d.CheckHost(hostname, dns.TypeA, setts)
require.NoErrorf(t, err, "host %q", hostname)
require.NoErrorf(tb, err, "host %q", hostname)
assert.Falsef(t, res.IsFiltered, "host %q", hostname)
assert.Falsef(tb, res.IsFiltered, "host %q", hostname)
}
func TestDNSFilter_CheckHost_hostRules(t *testing.T) {

View File

@@ -1,6 +1,7 @@
package filtering_test
import (
"context"
"fmt"
"net/netip"
"testing"
@@ -43,10 +44,10 @@ func TestDNSFilter_CheckHost_hostsContainer(t *testing.T) {
},
}
watcher := &aghtest.FSWatcher{
OnStart: func() (_ error) { panic("not implemented") },
OnEvents: func() (e <-chan struct{}) { return nil },
OnAdd: func(name string) (err error) { return nil },
OnClose: func() (err error) { return nil },
OnStart: func(ctx context.Context) (_ error) { panic(testutil.UnexpectedCall(ctx)) },
OnEvents: func() (e <-chan struct{}) { return nil },
OnAdd: func(name string) (err error) { return nil },
OnShutdown: func(_ context.Context) (err error) { return nil },
}
hc, err := aghnet.NewHostsContainer(files, watcher, "hosts")
require.NoError(t, err)

View File

@@ -133,7 +133,7 @@ func (d *DNSFilter) handleFilteringAddURL(w http.ResponseWriter, r *http.Request
return
}
d.conf.ConfigModified()
d.conf.ConfModifier.Apply(r.Context())
d.EnableFilters(true)
_, err = fmt.Fprintf(w, "OK %d rules\n", filt.RulesCount)
@@ -202,7 +202,7 @@ func (d *DNSFilter) handleFilteringRemoveURL(w http.ResponseWriter, r *http.Requ
d.logger.InfoContext(ctx, "deleted filter", "id", deleted.ID)
}()
d.conf.ConfigModified()
d.conf.ConfModifier.Apply(ctx)
d.EnableFilters(true)
// NOTE: The old files "filter.txt.old" aren't deleted. It's not really
@@ -264,7 +264,7 @@ func (d *DNSFilter) handleFilteringSetURL(w http.ResponseWriter, r *http.Request
return
}
d.conf.ConfigModified()
d.conf.ConfModifier.Apply(r.Context())
if restart {
d.EnableFilters(true)
}
@@ -289,7 +289,7 @@ func (d *DNSFilter) handleFilteringSetRules(w http.ResponseWriter, r *http.Reque
}
d.conf.UserRules = req.Rules
d.conf.ConfigModified()
d.conf.ConfModifier.Apply(r.Context())
d.EnableFilters(true)
}
@@ -403,7 +403,7 @@ func (d *DNSFilter) handleFilteringConfig(w http.ResponseWriter, r *http.Request
d.conf.FiltersUpdateIntervalHours = req.Interval
}()
d.conf.ConfigModified()
d.conf.ConfModifier.Apply(r.Context())
d.EnableFilters(true)
}
@@ -571,14 +571,14 @@ func protectedBool(mu *sync.RWMutex, ptr *bool) (val bool) {
// /control/safebrowsing/enable HTTP API.
func (d *DNSFilter) handleSafeBrowsingEnable(w http.ResponseWriter, r *http.Request) {
setProtectedBool(d.confMu, &d.conf.SafeBrowsingEnabled, true)
d.conf.ConfigModified()
d.conf.ConfModifier.Apply(r.Context())
}
// handleSafeBrowsingDisable is the handler for the POST
// /control/safebrowsing/disable HTTP API.
func (d *DNSFilter) handleSafeBrowsingDisable(w http.ResponseWriter, r *http.Request) {
setProtectedBool(d.confMu, &d.conf.SafeBrowsingEnabled, false)
d.conf.ConfigModified()
d.conf.ConfModifier.Apply(r.Context())
}
// handleSafeBrowsingStatus is the handler for the GET
@@ -597,14 +597,14 @@ func (d *DNSFilter) handleSafeBrowsingStatus(w http.ResponseWriter, r *http.Requ
// HTTP API.
func (d *DNSFilter) handleParentalEnable(w http.ResponseWriter, r *http.Request) {
setProtectedBool(d.confMu, &d.conf.ParentalEnabled, true)
d.conf.ConfigModified()
d.conf.ConfModifier.Apply(r.Context())
}
// handleParentalDisable is the handler for the POST /control/parental/disable
// HTTP API.
func (d *DNSFilter) handleParentalDisable(w http.ResponseWriter, r *http.Request) {
setProtectedBool(d.confMu, &d.conf.ParentalEnabled, false)
d.conf.ConfigModified()
d.conf.ConfModifier.Apply(r.Context())
}
// handleParentalStatus is the handler for the GET /control/parental/status

View File

@@ -2,6 +2,7 @@ package filtering
import (
"bytes"
"context"
"encoding/json"
"fmt"
"net/http"
@@ -11,6 +12,7 @@ import (
"testing"
"time"
"github.com/AdguardTeam/AdGuardHome/internal/aghtest"
"github.com/AdguardTeam/AdGuardHome/internal/schedule"
"github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/AdguardTeam/golibs/testutil"
@@ -103,6 +105,10 @@ func TestDNSFilter_handleFilteringSetURL(t *testing.T) {
for _, tc := range testCases {
t.Run(tc.name, func(t *testing.T) {
confModifiedCalled := false
confModifier := &aghtest.ConfigModifier{}
confModifier.OnApply = func(_ context.Context) {
confModifiedCalled = true
}
d, err := New(&Config{
Logger: slogutil.NewDiscardLogger(),
FilteringEnabled: true,
@@ -110,8 +116,8 @@ func TestDNSFilter_handleFilteringSetURL(t *testing.T) {
HTTPClient: &http.Client{
Timeout: 5 * time.Second,
},
ConfigModified: func() { confModifiedCalled = true },
DataDir: filtersDir,
ConfModifier: confModifier,
DataDir: filtersDir,
}, nil)
require.NoError(t, err)
t.Cleanup(d.Close)
@@ -183,13 +189,15 @@ func TestDNSFilter_handleSafeBrowsingStatus(t *testing.T) {
for _, tc := range testCases {
t.Run(tc.name, func(t *testing.T) {
handlers := make(map[string]http.Handler)
confModifier := &aghtest.ConfigModifier{}
confModifier.OnApply = func(_ context.Context) {
testutil.RequireSend(testutil.PanicT{}, confModCh, struct{}{}, testTimeout)
}
d, err := New(&Config{
Logger: slogutil.NewDiscardLogger(),
ConfigModified: func() {
testutil.RequireSend(testutil.PanicT{}, confModCh, struct{}{}, testTimeout)
},
DataDir: filtersDir,
Logger: slogutil.NewDiscardLogger(),
ConfModifier: confModifier,
DataDir: filtersDir,
HTTPRegister: func(_, url string, handler http.HandlerFunc) {
handlers[url] = handler
},
@@ -268,13 +276,15 @@ func TestDNSFilter_handleParentalStatus(t *testing.T) {
for _, tc := range testCases {
t.Run(tc.name, func(t *testing.T) {
handlers := make(map[string]http.Handler)
confModifier := &aghtest.ConfigModifier{}
confModifier.OnApply = func(_ context.Context) {
testutil.RequireSend(testutil.PanicT{}, confModCh, struct{}{}, testTimeout)
}
d, err := New(&Config{
Logger: slogutil.NewDiscardLogger(),
ConfigModified: func() {
testutil.RequireSend(testutil.PanicT{}, confModCh, struct{}{}, testTimeout)
},
DataDir: filtersDir,
Logger: slogutil.NewDiscardLogger(),
ConfModifier: confModifier,
DataDir: filtersDir,
HTTPRegister: func(_, url string, handler http.HandlerFunc) {
handlers[url] = handler
},

View File

@@ -75,13 +75,13 @@ func TestIDGenerator_Fix(t *testing.T) {
// assertUniqueIDs is a test helper that asserts that the IDs of filters are
// unique.
func assertUniqueIDs(t testing.TB, flts []FilterYAML) {
t.Helper()
func assertUniqueIDs(tb testing.TB, flts []FilterYAML) {
tb.Helper()
uc := aghalg.UniqChecker[rulelist.URLFilterID]{}
for _, f := range flts {
uc.Add(f.ID)
}
assert.NoError(t, uc.Validate())
assert.NoError(tb, uc.Validate())
}

View File

@@ -97,16 +97,38 @@ func (s *DefaultStorage) MatchRequest(dReq *urlfilter.DNSRequest) (rws []*rules.
ctx := context.TODO()
rrules := s.rewriteRulesForReq(dReq)
if len(rrules) == 0 {
rewriteRules := s.rewriteRulesForReq(dReq)
if len(rewriteRules) == 0 {
return nil
}
resolvedRules, wildcardRewrite := s.resolveCNAMEChain(ctx, dReq, rewriteRules)
if wildcardRewrite != nil {
return []*rules.DNSRewrite{wildcardRewrite}
}
if resolvedRules == nil {
return nil
}
return s.collectDNSRewrites(resolvedRules, dReq.DNSType)
}
// resolveCNAMEChain follows the CNAME chain for a DNS request, handling loops
// and special cases. dReq must not be nil, and rewriteRules must not contain
// nil elements.
func (s *DefaultStorage) resolveCNAMEChain(
ctx context.Context,
dReq *urlfilter.DNSRequest,
rewriteRules []*rules.NetworkRule,
) (resolvedRules []*rules.NetworkRule, wildcardRewrite *rules.DNSRewrite) {
// TODO(a.garipov): Check cnames for cycles on initialization.
cnames := container.NewMapSet[string]()
host := dReq.Hostname
for len(rrules) > 0 && rrules[0].DNSRewrite != nil && rrules[0].DNSRewrite.NewCNAME != "" {
rule := rrules[0]
for len(rewriteRules) > 0 &&
rewriteRules[0].DNSRewrite != nil &&
rewriteRules[0].DNSRewrite.NewCNAME != "" {
rule := rewriteRules[0]
rwAns := rule.DNSRewrite.NewCNAME
s.logger.DebugContext(ctx, "cname found", "host", host, "cname", rwAns)
@@ -115,37 +137,43 @@ func (s *DefaultStorage) MatchRequest(dReq *urlfilter.DNSRequest) (rws []*rules.
// A request for the hostname itself is an exception rule.
// TODO(d.kolyshev): Check rewrite of a pattern onto itself.
return nil
return nil, nil
}
if host == rwAns && isWildcard(rule.RuleText) {
// An "*.example.com → sub.example.com" rewrite matching in a loop.
//
// See https://github.com/AdguardTeam/AdGuardHome/issues/4016.
return []*rules.DNSRewrite{rule.DNSRewrite}
if isSelfMatchingWildcard(host, rwAns, rule.RuleText) {
return nil, rule.DNSRewrite
}
if cnames.Has(rwAns) {
s.logger.InfoContext(ctx, "rewrite cname loop", "host", dReq.Hostname, "rewrite", rwAns)
s.logger.WarnContext(ctx, "rewrite cname loop", "host", dReq.Hostname, "rewrite", rwAns)
return nil
return nil, nil
}
cnames.Add(rwAns)
drules := s.rewriteRulesForReq(&urlfilter.DNSRequest{
rewriteRulesForReq := s.rewriteRulesForReq(&urlfilter.DNSRequest{
Hostname: rwAns,
DNSType: dReq.DNSType,
})
if drules != nil {
rrules = drules
if rewriteRulesForReq != nil {
rewriteRules = rewriteRulesForReq
}
host = rwAns
}
return s.collectDNSRewrites(rrules, dReq.DNSType)
return rewriteRules, nil
}
// isSelfMatchingWildcard returns true when a wildcard rewrite matches its own
// result.
//
// For example, an "*.example.com → sub.example.com" rewrite matching in a loop.
//
// See https://github.com/AdguardTeam/AdGuardHome/issues/4016.
func isSelfMatchingWildcard(host, rwAns, ruleText string) (ok bool) {
return host == rwAns && isWildcard(ruleText)
}
// collectDNSRewrites filters DNSRewrite by question type.

View File

@@ -36,6 +36,8 @@ func (d *DNSFilter) handleRewriteList(w http.ResponseWriter, r *http.Request) {
// handleRewriteAdd is the handler for the POST /control/rewrite/add HTTP API.
func (d *DNSFilter) handleRewriteAdd(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
rwJSON := rewriteEntryJSON{}
err := json.NewDecoder(r.Body).Decode(&rwJSON)
if err != nil {
@@ -49,7 +51,7 @@ func (d *DNSFilter) handleRewriteAdd(w http.ResponseWriter, r *http.Request) {
Answer: rwJSON.Answer,
}
err = rw.normalize(r.Context(), d.logger)
err = rw.normalize(ctx, 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,7 +66,7 @@ func (d *DNSFilter) handleRewriteAdd(w http.ResponseWriter, r *http.Request) {
d.conf.Rewrites = append(d.conf.Rewrites, rw)
d.logger.DebugContext(
r.Context(),
ctx,
"added rewrite element",
"domain", rw.Domain,
"answer", rw.Answer,
@@ -72,12 +74,14 @@ func (d *DNSFilter) handleRewriteAdd(w http.ResponseWriter, r *http.Request) {
)
}()
d.conf.ConfigModified()
d.conf.ConfModifier.Apply(ctx)
}
// handleRewriteDelete is the handler for the POST /control/rewrite/delete HTTP
// API.
func (d *DNSFilter) handleRewriteDelete(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
jsent := rewriteEntryJSON{}
err := json.NewDecoder(r.Body).Decode(&jsent)
if err != nil {
@@ -92,28 +96,27 @@ func (d *DNSFilter) handleRewriteDelete(w http.ResponseWriter, r *http.Request)
}
arr := []*LegacyRewrite{}
func() {
d.confMu.Lock()
defer d.confMu.Unlock()
defer d.conf.ConfModifier.Apply(ctx)
for _, ent := range d.conf.Rewrites {
if ent.equal(entDel) {
d.logger.DebugContext(
r.Context(),
"removed rewrite element",
"domain", ent.Domain,
"answer", ent.Answer,
)
continue
}
d.confMu.Lock()
defer d.confMu.Unlock()
for _, ent := range d.conf.Rewrites {
if !ent.equal(entDel) {
arr = append(arr, ent)
}
d.conf.Rewrites = arr
}()
d.conf.ConfigModified()
continue
}
d.logger.DebugContext(
ctx,
"removed rewrite element",
"domain", ent.Domain,
"answer", ent.Answer,
)
}
d.conf.Rewrites = arr
}
// rewriteUpdateJSON is a struct for JSON object with rewrite rule update info.
@@ -125,6 +128,8 @@ type rewriteUpdateJSON struct {
// handleRewriteUpdate is the handler for the PUT /control/rewrite/update HTTP
// API.
func (d *DNSFilter) handleRewriteUpdate(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
updateJSON := rewriteUpdateJSON{}
err := json.NewDecoder(r.Body).Decode(&updateJSON)
if err != nil {
@@ -143,7 +148,7 @@ func (d *DNSFilter) handleRewriteUpdate(w http.ResponseWriter, r *http.Request)
Answer: updateJSON.Update.Answer,
}
err = rwAdd.normalize(r.Context(), d.logger)
err = rwAdd.normalize(ctx, 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.
@@ -155,7 +160,7 @@ func (d *DNSFilter) handleRewriteUpdate(w http.ResponseWriter, r *http.Request)
index := -1
defer func() {
if index >= 0 {
d.conf.ConfigModified()
d.conf.ConfModifier.Apply(ctx)
}
}()
@@ -171,7 +176,6 @@ func (d *DNSFilter) handleRewriteUpdate(w http.ResponseWriter, r *http.Request)
d.conf.Rewrites = slices.Replace(d.conf.Rewrites, index, index+1, rwAdd)
ctx := r.Context()
d.logger.DebugContext(
ctx,
"removed rewrite element",

View File

@@ -2,6 +2,7 @@ package filtering_test
import (
"bytes"
"context"
"encoding/json"
"io"
"net/http"
@@ -9,6 +10,7 @@ import (
"testing"
"time"
"github.com/AdguardTeam/AdGuardHome/internal/aghtest"
"github.com/AdguardTeam/AdGuardHome/internal/filtering"
"github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/AdguardTeam/golibs/testutil"
@@ -148,20 +150,17 @@ func TestDNSFilter_handleRewriteHTTP(t *testing.T) {
}}
for _, tc := range testCases {
onConfModified := func() {
if !tc.wantConfMod {
panic("config modified has been fired")
}
testutil.RequireSend(testutil.PanicT{}, confModCh, struct{}{}, testTimeout)
}
t.Run(tc.name, func(t *testing.T) {
handlers := make(map[string]http.Handler)
confModifier := &aghtest.ConfigModifier{}
confModifier.OnApply = func(_ context.Context) {
require.Truef(t, tc.wantConfMod, "config modified has been fired")
testutil.RequireSend(testutil.PanicT{}, confModCh, struct{}{}, testTimeout)
}
d, err := filtering.New(&filtering.Config{
Logger: slogutil.NewDiscardLogger(),
ConfigModified: onConfModified,
Logger: slogutil.NewDiscardLogger(),
ConfModifier: confModifier,
HTTPRegister: func(_, url string, handler http.HandlerFunc) {
handlers[url] = handler
},
@@ -210,20 +209,20 @@ func TestDNSFilter_handleRewriteHTTP(t *testing.T) {
// assertRewritesList checks if rewrites list equals the list received from the
// handler by listURL.
func assertRewritesList(t *testing.T, handler http.Handler, wantList []*rewriteJSON) {
t.Helper()
func assertRewritesList(tb testing.TB, handler http.Handler, wantList []*rewriteJSON) {
tb.Helper()
r := httptest.NewRequest(http.MethodGet, listURL, nil)
w := httptest.NewRecorder()
handler.ServeHTTP(w, r)
require.Equal(t, http.StatusOK, w.Code)
require.Equal(tb, http.StatusOK, w.Code)
var actual []*rewriteJSON
err := json.NewDecoder(w.Body).Decode(&actual)
require.NoError(t, err)
require.NoError(tb, err)
assert.Equal(t, wantList, actual)
assert.Equal(tb, wantList, actual)
}
// rewriteEntriesToLegacyRewrites gets legacy rewrites from json entries.

View File

@@ -52,8 +52,8 @@ func newURLFilterID() (id rulelist.URLFilterID) {
// newFilter is a helper for creating new filters in tests. It does not
// register the closing of the filter using t.Cleanup; callers must do that
// either directly or by using the filter in an engine.
func newFilter(t testing.TB, u *url.URL, name string) (f *rulelist.Filter) {
t.Helper()
func newFilter(tb testing.TB, u *url.URL, name string) (f *rulelist.Filter) {
tb.Helper()
f, err := rulelist.NewFilter(&rulelist.FilterConfig{
URL: u,
@@ -62,7 +62,7 @@ func newFilter(t testing.TB, u *url.URL, name string) (f *rulelist.Filter) {
URLFilterID: newURLFilterID(),
Enabled: true,
})
require.NoError(t, err)
require.NoError(tb, err)
return f
}
@@ -71,24 +71,24 @@ func newFilter(t testing.TB, u *url.URL, name string) (f *rulelist.Filter) {
// file and the HTTP-server. It also registers file removal and server stopping
// using t.Cleanup.
func newFilterLocations(
t testing.TB,
tb testing.TB,
cacheDir string,
fileData string,
httpData string,
) (fileURL, srvURL *url.URL) {
t.Helper()
tb.Helper()
f, err := os.CreateTemp(cacheDir, "")
require.NoError(t, err)
require.NoError(tb, err)
err = f.Close()
require.NoError(t, err)
require.NoError(tb, err)
filePath := f.Name()
err = os.WriteFile(filePath, []byte(fileData), 0o644)
require.NoError(t, err)
require.NoError(tb, err)
testutil.CleanupAndRequireSuccess(t, func() (err error) {
testutil.CleanupAndRequireSuccess(tb, func() (err error) {
return os.Remove(filePath)
})
@@ -98,10 +98,10 @@ func newFilterLocations(
}
srv := newStringHTTPServer(httpData)
t.Cleanup(srv.Close)
tb.Cleanup(srv.Close)
srvURL, err = url.Parse(srv.URL)
require.NoError(t, err)
require.NoError(tb, err)
return fileURL, srvURL
}

View File

@@ -13,7 +13,7 @@ import (
// Deprecated: Use handleSafeSearchSettings.
func (d *DNSFilter) handleSafeSearchEnable(w http.ResponseWriter, r *http.Request) {
setProtectedBool(d.confMu, &d.conf.SafeSearchConf.Enabled, true)
d.conf.ConfigModified()
d.conf.ConfModifier.Apply(r.Context())
}
// handleSafeSearchDisable is the handler for POST /control/safesearch/disable
@@ -22,7 +22,7 @@ func (d *DNSFilter) handleSafeSearchEnable(w http.ResponseWriter, r *http.Reques
// Deprecated: Use handleSafeSearchSettings.
func (d *DNSFilter) handleSafeSearchDisable(w http.ResponseWriter, r *http.Request) {
setProtectedBool(d.confMu, &d.conf.SafeSearchConf.Enabled, false)
d.conf.ConfigModified()
d.conf.ConfModifier.Apply(r.Context())
}
// handleSafeSearchStatus is the handler for GET /control/safesearch/status
@@ -42,6 +42,8 @@ func (d *DNSFilter) handleSafeSearchStatus(w http.ResponseWriter, r *http.Reques
// handleSafeSearchSettings is the handler for PUT /control/safesearch/settings
// HTTP API.
func (d *DNSFilter) handleSafeSearchSettings(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
req := &SafeSearchConfig{}
err := json.NewDecoder(r.Body).Decode(req)
if err != nil {
@@ -51,7 +53,7 @@ func (d *DNSFilter) handleSafeSearchSettings(w http.ResponseWriter, r *http.Requ
}
conf := *req
err = d.safeSearch.Update(r.Context(), conf)
err = d.safeSearch.Update(ctx, conf)
if err != nil {
aghhttp.Error(r, w, http.StatusBadRequest, "updating: %s", err)
@@ -65,7 +67,7 @@ func (d *DNSFilter) handleSafeSearchSettings(w http.ResponseWriter, r *http.Requ
d.conf.SafeSearchConf = conf
}()
d.conf.ConfigModified()
d.conf.ConfModifier.Apply(ctx)
aghhttp.OK(w)
}

View File

@@ -502,6 +502,16 @@ var blockedServices = []blockedService{{
"||globosat.globo.com^",
"||gsatmulti.globo.com^",
},
}, {
ID: "chatgpt",
Name: "ChatGPT",
IconSVG: []byte("<svg xmlns=\"http://www.w3.org/2000/svg\" fill=\"currentColor\" viewBox=\"0 0 50 50\"><path d=\"M45.403 25.562c-.506-1.89-1.518-3.553-2.906-4.862 1.134-2.665.963-5.724-.487-8.237-1.391-2.408-3.636-4.131-6.322-4.851-1.891-.506-3.839-.462-5.669.088C28.276 5.382 25.562 4 22.647 4c-4.906 0-9.021 3.416-10.116 7.991-.01.001-.019-.003-.029-.002-2.902.36-5.404 2.019-6.865 4.549-1.391 2.408-1.76 5.214-1.04 7.9.507 1.891 1.519 3.556 2.909 4.865-1.134 2.666-.97 5.714.484 8.234 1.391 2.408 3.636 4.131 6.322 4.851.896.24 1.807.359 2.711.359 1.003 0 1.995-.161 2.957-.45C21.722 44.619 24.425 46 27.353 46c4.911 0 9.028-3.422 10.12-8.003 2.88-.35 5.431-2.006 6.891-4.535 1.39-2.408 1.759-5.214 1.039-7.9zM35.17 9.543c2.171.581 3.984 1.974 5.107 3.919 1.049 1.817 1.243 4 .569 5.967-.099-.062-.193-.131-.294-.19l-9.169-5.294a1.0072 1.0072 0 0 0-1.01.006l-10.198 6.041-.052-4.607 8.663-5.001C30.733 9.26 33 8.963 35.17 9.543zm-5.433 12.652.062 5.504-4.736 2.805-4.799-2.699-.062-5.504 4.736-2.805 4.799 2.699zm-15.502-7.783C14.235 9.773 18.009 6 22.647 6c2.109 0 4.092.916 5.458 2.488-.105.056-.214.103-.318.163l-9.17 5.294c-.312.181-.504.517-.5.877l.133 11.851-4.015-2.258V14.412zm-7.707 9.509c-.581-2.17-.282-4.438.841-6.383 1.06-1.836 2.823-3.074 4.884-3.474-.004.116-.018.23-.018.348V25c0 .361.195.694.51.872l10.329 5.81-3.964 2.348-8.662-5.002c-1.946-1.123-3.338-2.936-3.92-5.107zm8.302 16.536c-2.171-.581-3.984-1.974-5.107-3.919-1.053-1.824-1.249-4.001-.573-5.97.101.063.196.133.299.193l9.169 5.294a.9998.9998 0 0 0 1.01-.006l10.198-6.041.052 4.607-8.663 5.001c-1.946 1.125-4.214 1.424-6.385.841zm20.935-4.869c0 4.639-3.773 8.412-8.412 8.412-2.119 0-4.094-.919-5.459-2.494.105-.056.216-.098.32-.158l9.17-5.294c.312-.181.504-.517.5-.877l-.134-11.85 4.015 2.258v10.003zm6.866-3.126c-1.056 1.83-2.84 3.086-4.884 3.483.004-.12.018-.237.018-.357V25c0-.361-.195-.694-.51-.872l-10.329-5.81 3.964-2.348 8.662 5.002c1.946 1.123 3.338 2.937 3.92 5.107.581 2.17.282 4.438-.841 6.383z\"/></svg>"),
Rules: []string{
"||chatgpt.com^",
"||oaistatic.com^",
"||oaiusercontent.com^",
"||openai.com^",
},
}, {
ID: "claro",
Name: "Claro",
@@ -530,6 +540,14 @@ var blockedServices = []blockedService{{
"||clarovideo.com^",
"||usclaro.com^",
},
}, {
ID: "claude",
Name: "Claude",
IconSVG: []byte("<svg xmlns=\"http://www.w3.org/2000/svg\" fill=\"currentColor\" viewBox=\"0 0 40 40\"><path d=\"m7.8 26.3 7.7-4.4.1-.4v-.2h-1.8L9.4 21H5.5l-3.7-.2-1-.2-.8-1.2v-.6l.9-.5H2l2.5.3 3.8.2 2.7.2 4 .4h.7v-.3l-.2-.1-.1-.2-4-2.6-4.1-2.8L5 11.8 3.9 11l-.6-.8L3 8.6l1.1-1.2h1.4l.4.2 1.5 1.1 3.1 2.4 4.1 3 .6.6.3-.2v-.1l-.3-.5-2.2-4-2.4-4.1-1-1.7-.3-1A4.9 4.9 0 0 1 9 1.9L10.3.2 11 0l1.7.2.7.6 1 2.3L16 6.8l2.6 5 .7 1.5.4 1.4.2.4h.2v-.3l.2-2.8.4-3.4.4-4.5.1-1.2.7-1.5L23 .6l1 .4.8 1.2-.2.7-.4 3-1 4.8-.5 3.2h.3l.4-.4 1.6-2.1L27.8 8 29 6.6l1.4-1.5 1-.7H33l1.3 1.9-.6 1.9-1.7 2.2-1.5 1.9-2 2.8-1.4 2.2.2.2h.3l4.7-1 2.5-.5 3-.5 1.4.6.2.7-.6 1.3-3.2.8-3.8.8L26 21v.2l2.6.2h3.7l5 .4 1.3 1 .8 1-.1.8-2 1-2.7-.7-6.3-1.5-2.2-.5H26v.2l1.8 1.7 3.3 3 4.1 3.9.2 1-.5.7-.5-.1-3.7-2.7-1.4-1.3-3.1-2.7h-.3v.3l.8 1.1 3.8 5.8.2 1.8-.2.6-1 .3-1.1-.2-2.3-3.2-2.3-3.5-2-3.2-.1.1-1.1 12-.6.6-1.2.4-1-.7-.5-1.3.5-2.4.7-3.2.5-2.5.5-3.1.2-1v-.1h-.2L17 28.4l-3.6 4.9-2.8 3-.7.3-1.2-.6.2-1.1.6-1 4-5 2.3-3 1.5-1.9v-.2L6.7 30.6l-1.9.2L4 30l.1-1.2.4-.4 3.2-2.1Z\"/></svg>"),
Rules: []string{
"||anthropic.com^",
"||claude.ai^",
},
}, {
ID: "cloudflare",
Name: "Cloudflare",
@@ -600,6 +618,13 @@ var blockedServices = []blockedService{{
"||dm-event.net^",
"||dmcdn.net^",
},
}, {
ID: "deepseek",
Name: "DeepSeek",
IconSVG: []byte("<svg xmlns=\"http://www.w3.org/2000/svg\" fill=\"currentColor\" viewBox=\"0 0 30 30\"><path d=\"M29.7 5.8c-.3-.1-.5.2-.7.3l-.1.2c-.5.5-1 .8-1.7.8a3 3 0 0 0-2.7 1c-.2-1-.8-1.5-1.6-2-.4-.1-.9-.3-1.2-.7l-.4-1c0-.2-.1-.4-.3-.4-.3 0-.4.1-.5.3-.4.7-.5 1.5-.5 2.3a5 5 0 0 0 2.3 4.3c.1 0 .2.2.1.4l-.3 1c0 .2-.2.3-.4.2-.8-.4-1.5-.9-2.2-1.5-1-1-2-2.2-3.2-3a13.8 13.8 0 0 0-.9-.7c-1.2-1.2.2-2.1.5-2.3.3 0 .1-.5-1-.5-1 0-2 .4-3.3.9a3.8 3.8 0 0 1-.6.1 12 12 0 0 0-3.6 0 7.8 7.8 0 0 0-5.6 3.2 9.6 9.6 0 0 0-1.6 7.6c.5 2.8 2 5.2 4.2 7 2.3 2 5 2.9 8 2.7 1.9 0 4-.3 6.3-2.3a7.3 7.3 0 0 0 4.4.3c.9-.2.8-1 .5-1.2-2.4-.8-2.7-1.1-2.7-1.1 1.4-1.7 3.4-3.3 4.3-8.8v-1c0-.3 0-.3.2-.4a5.2 5.2 0 0 0 2-.6c1.7-1 2.4-2.5 2.6-4.3 0-.3 0-.6-.3-.8zm-15.2 17C11.9 20.6 10.6 20 10 20c-.5 0-.4.6-.3 1l.4.9c.2.2.3.5 0 .7-1 .6-2.4-.1-2.5-.2a12.2 12.2 0 0 1-5.7-9.7c0-.5.1-.7.6-.7a5.9 5.9 0 0 1 1.9 0c2.7.3 5 1.5 6.8 3.4 1.1 1 2 2.3 2.8 3.6a17.3 17.3 0 0 0 4.2 4.5c-1 0-2.7.1-3.8-.8zm1.2-8.1a.4.4 0 0 1 .5-.4.4.4 0 0 1 .3.4.4.4 0 0 1-.4.3.4.4 0 0 1-.4-.3zm4 2-.8.2c-.4 0-.8-.2-1-.4-.4-.2-.6-.4-.7-1V15c.1-.5 0-.7-.3-1l-.8-.2a.7.7 0 0 1-.4-.1c-.1 0-.2-.2-.1-.4l.2-.3c.5-.3 1-.2 1.5 0 .4.2.7.5 1.2 1l.9 1.1.5 1c.1.3 0 .5-.3.7z\"/></svg>"),
Rules: []string{
"||deepseek.com^",
},
}, {
ID: "deezer",
Name: "Deezer",
@@ -2059,6 +2084,16 @@ var blockedServices = []blockedService{{
"||nvidianews.com^",
"||tegrazone.com^",
},
}, {
ID: "odysee",
Name: "Odysee",
IconSVG: []byte("<svg xmlns=\"http://www.w3.org/2000/svg\" fill=\"currentColor\" viewBox=\"-11 2 178 178\"><path d=\"M82 57c-14 5-20-2-21-14-1-13 12-17 12-17 14-5 18 2 21 12s1 14-12 19m65 85-9-28a67 67 0 0 0-21-23 5 5 0 0 1 0-8c7-6 18-18 22-25 3-4 7-13 8-20 0-6-1-12-8-15s-11 1-11 1c-5 3-6 12-9 21-4 10-10 11-13 11s-1-3-9-24c-7-21-26-17-40-9-19 11-11 35-6 50-3 2-12 4-21 9l-15 8c-6 6-9 11-7 19a12 12 0 0 0 6 7c5 2 13-1 24-10 9-6 19-9 19-9l13 24c7 13-7 17-8 17-2 0-23-2-18 16s31 12 44 3c14-9 10-38 10-38 13-2 17 12 19 19 1 7-2 19 11 20a21 21 0 0 0 6-1c7-2 11-5 13-9a9 9 0 0 0 0-6M88 33a9 9 0 0 0-2-3h-2v2l1 2 1 1h1l1-3m0 7-1 2a7 7 0 0 1 1 5l1 2h1l1-1c1-3 0-5-1-7l-2-1\"/></svg>"),
Rules: []string{
"||odycdn.com^",
"||odysee.com^",
"||odysee.live^",
"||odysee.tv^",
},
}, {
ID: "ok",
Name: "OK.ru",

View File

@@ -1,317 +1,162 @@
package home
import (
"crypto/rand"
"encoding/binary"
"encoding/hex"
"context"
"fmt"
"net/http"
"sync"
"log/slog"
"time"
"github.com/AdguardTeam/AdGuardHome/internal/aghos"
"github.com/AdguardTeam/AdGuardHome/internal/aghuser"
"github.com/AdguardTeam/golibs/errors"
"github.com/AdguardTeam/golibs/log"
"github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/AdguardTeam/golibs/netutil"
"go.etcd.io/bbolt"
"github.com/AdguardTeam/golibs/netutil/httputil"
"github.com/AdguardTeam/golibs/timeutil"
"golang.org/x/crypto/bcrypt"
)
// sessionTokenSize is the length of session token in bytes.
const sessionTokenSize = 16
type session struct {
userName string
// expire is the expiration time, in seconds.
expire uint32
}
func (s *session) serialize() []byte {
const (
expireLen = 4
nameLen = 2
)
data := make([]byte, expireLen+nameLen+len(s.userName))
binary.BigEndian.PutUint32(data[0:4], s.expire)
binary.BigEndian.PutUint16(data[4:6], uint16(len(s.userName)))
copy(data[6:], []byte(s.userName))
return data
}
func (s *session) deserialize(data []byte) bool {
if len(data) < 4+2 {
return false
}
s.expire = binary.BigEndian.Uint32(data[0:4])
nameLen := binary.BigEndian.Uint16(data[4:6])
data = data[6:]
if len(data) < int(nameLen) {
return false
}
s.userName = string(data)
return true
}
// Auth is the global authentication object.
type Auth struct {
trustedProxies netutil.SubnetSet
db *bbolt.DB
rateLimiter *authRateLimiter
sessions map[string]*session
users []webUser
lock sync.Mutex
sessionTTL uint32
}
// sessionsDBName is the name of the file where session data is stored.
const sessionsDBName = "sessions.db"
// webUser represents a user of the Web UI.
//
// TODO(s.chzhen): Improve naming.
type webUser struct {
Name string `yaml:"name"`
// Name represents the login name of the web user.
Name string `yaml:"name"`
// PasswordHash is the hashed representation of the web user password.
PasswordHash string `yaml:"password"`
// UserID is the unique identifier of the web user.
UserID aghuser.UserID `yaml:"-"`
}
// InitAuth initializes the global authentication object.
func InitAuth(
dbFilename string,
users []webUser,
sessionTTL uint32,
rateLimiter *authRateLimiter,
trustedProxies netutil.SubnetSet,
) (a *Auth) {
log.Info("Initializing auth module: %s", dbFilename)
a = &Auth{
sessionTTL: sessionTTL,
rateLimiter: rateLimiter,
sessions: make(map[string]*session),
users: users,
trustedProxies: trustedProxies,
// toUser returns the new properly initialized *aghuser.User using stored
// properties. It panics if there is an error generating the user ID.
func (wu *webUser) toUser() (u *aghuser.User) {
uid := wu.UserID
if uid == (aghuser.UserID{}) {
uid = aghuser.MustNewUserID()
}
var err error
a.db, err = bbolt.Open(dbFilename, aghos.DefaultPermFile, nil)
if err != nil {
log.Error("auth: open DB: %s: %s", dbFilename, err)
if err.Error() == "invalid argument" {
log.Error("AdGuard Home cannot be initialized due to an incompatible file system.\nPlease read the explanation here: https://github.com/AdguardTeam/AdGuardHome/wiki/Getting-Started#limitations")
}
return nil
return &aghuser.User{
Password: aghuser.NewDefaultPassword(wu.PasswordHash),
Login: aghuser.Login(wu.Name),
ID: uid,
}
a.loadSessions()
log.Info("auth: initialized. users:%d sessions:%d", len(a.users), len(a.sessions))
return a
}
// Close closes the authentication database.
func (a *Auth) Close() {
_ = a.db.Close()
// authConfig is the configuration structure for [auth].
type authConfig struct {
// baseLogger is used for creating other loggers. It must not be nil.
baseLogger *slog.Logger
// rateLimiter manages the rate limiting for login attempts. It must not be
// nil.
rateLimiter loginRaateLimiter
// trustedProxies is a set of subnets considered as trusted.
trustedProxies netutil.SubnetSet
// dbFilename is the name of the file where session data is stored. It must
// not be empty.
dbFilename string
// users contains web user information from the configuration file.
users []webUser
// sessionTTL is the TTL (Time To Live) for web user sessions.
sessionTTL time.Duration
// isGLiNet indicates whether GLiNet mode is enabled.
isGLiNet bool
}
func bucketName() []byte {
return []byte("sessions-2")
// auth stores web user information and handles authentication.
type auth struct {
logger *slog.Logger
rateLimiter loginRaateLimiter
trustedProxies netutil.SubnetSet
sessions aghuser.SessionStorage
users aghuser.DB
isGLiNet bool
}
// loadSessions loads sessions from the database file and removes expired
// sessions.
func (a *Auth) loadSessions() {
tx, err := a.db.Begin(true)
if err != nil {
log.Error("auth: bbolt.Begin: %s", err)
return
}
defer func() {
_ = tx.Rollback()
}()
bkt := tx.Bucket(bucketName())
if bkt == nil {
return
}
removed := 0
if tx.Bucket([]byte("sessions")) != nil {
_ = tx.DeleteBucket([]byte("sessions"))
removed = 1
}
now := uint32(time.Now().UTC().Unix())
forEach := func(k, v []byte) error {
s := session{}
if !s.deserialize(v) || s.expire <= now {
err = bkt.Delete(k)
if err != nil {
log.Error("auth: bbolt.Delete: %s", err)
} else {
removed++
}
return nil
}
a.sessions[hex.EncodeToString(k)] = &s
return nil
}
_ = bkt.ForEach(forEach)
if removed != 0 {
err = tx.Commit()
// newAuth returns the new properly initialized *auth.
func newAuth(ctx context.Context, conf *authConfig) (a *auth, err error) {
userDB := aghuser.NewDefaultDB()
for i, u := range conf.users {
err = userDB.Create(ctx, u.toUser())
if err != nil {
log.Error("bolt.Commit(): %s", err)
return nil, fmt.Errorf("users: at index %d: %w", i, err)
}
}
log.Debug("auth: loaded %d sessions from DB (removed %d expired)", len(a.sessions), removed)
s, err := aghuser.NewDefaultSessionStorage(ctx, &aghuser.DefaultSessionStorageConfig{
Logger: conf.baseLogger.With(slogutil.KeyPrefix, "session_storage"),
Clock: timeutil.SystemClock{},
UserDB: userDB,
DBPath: conf.dbFilename,
SessionTTL: conf.sessionTTL,
})
if err != nil {
return nil, fmt.Errorf("creating session storage: %w", err)
}
return &auth{
logger: conf.baseLogger.With(slogutil.KeyPrefix, "auth"),
rateLimiter: conf.rateLimiter,
trustedProxies: conf.trustedProxies,
sessions: s,
users: userDB,
isGLiNet: conf.isGLiNet,
}, nil
}
// addSession adds a new session to the list of sessions and saves it in the
// database file.
func (a *Auth) addSession(data []byte, s *session) {
name := hex.EncodeToString(data)
a.lock.Lock()
a.sessions[name] = s
a.lock.Unlock()
if a.storeSession(data, s) {
log.Debug("auth: created session %s: expire=%d", name, s.expire)
// middleware returns authentication middleware.
func (a *auth) middleware() (mw httputil.Middleware) {
if a.isGLiNet {
return newAuthMiddlewareGLiNet(&authMiddlewareGLiNetConfig{
logger: a.logger,
clock: timeutil.SystemClock{},
tokenFilePrefix: glFilePrefix,
ttl: glTokenTimeout,
maxTokenSize: MaxFileSize,
})
}
return newAuthMiddlewareDefault(&authMiddlewareDefaultConfig{
logger: a.logger,
rateLimiter: a.rateLimiter,
trustedProxies: a.trustedProxies,
sessions: a.sessions,
users: a.users,
})
}
// storeSession saves a session in the database file.
func (a *Auth) storeSession(data []byte, s *session) bool {
tx, err := a.db.Begin(true)
// usersList returns a copy of a users list.
func (a *auth) usersList(ctx context.Context) (webUsers []webUser) {
users, err := a.users.All(ctx)
if err != nil {
log.Error("auth: bbolt.Begin: %s", err)
return false
}
defer func() {
_ = tx.Rollback()
}()
bkt, err := tx.CreateBucketIfNotExists(bucketName())
if err != nil {
log.Error("auth: bbolt.CreateBucketIfNotExists: %s", err)
return false
// Should not happen.
panic(err)
}
err = bkt.Put(data, s.serialize())
if err != nil {
log.Error("auth: bbolt.Put: %s", err)
return false
webUsers = make([]webUser, 0, len(users))
for _, u := range users {
webUsers = append(webUsers, webUser{
Name: string(u.Login),
PasswordHash: string(u.Password.Hash()),
UserID: u.ID,
})
}
err = tx.Commit()
if err != nil {
log.Error("auth: bbolt.Commit: %s", err)
return false
}
return true
return webUsers
}
// removeSessionFromFile removes a stored session from the DB file on disk.
func (a *Auth) removeSessionFromFile(sess []byte) {
tx, err := a.db.Begin(true)
if err != nil {
log.Error("auth: bbolt.Begin: %s", err)
return
}
defer func() {
_ = tx.Rollback()
}()
bkt := tx.Bucket(bucketName())
if bkt == nil {
log.Error("auth: bbolt.Bucket")
return
}
err = bkt.Delete(sess)
if err != nil {
log.Error("auth: bbolt.Put: %s", err)
return
}
err = tx.Commit()
if err != nil {
log.Error("auth: bbolt.Commit: %s", err)
return
}
log.Debug("auth: removed session from DB")
}
// checkSessionResult is the result of checking a session.
type checkSessionResult int
// checkSessionResult constants.
const (
checkSessionOK checkSessionResult = 0
checkSessionNotFound checkSessionResult = -1
checkSessionExpired checkSessionResult = 1
)
// checkSession checks if the session is valid.
func (a *Auth) checkSession(sess string) (res checkSessionResult) {
now := uint32(time.Now().UTC().Unix())
update := false
a.lock.Lock()
defer a.lock.Unlock()
s, ok := a.sessions[sess]
if !ok {
return checkSessionNotFound
}
if s.expire <= now {
delete(a.sessions, sess)
key, _ := hex.DecodeString(sess)
a.removeSessionFromFile(key)
return checkSessionExpired
}
newExpire := now + a.sessionTTL
if s.expire/(24*60*60) != newExpire/(24*60*60) {
// update expiration time once a day
update = true
s.expire = newExpire
}
if update {
key, _ := hex.DecodeString(sess)
if a.storeSession(key, s) {
log.Debug("auth: updated session %s: expire=%d", sess, s.expire)
}
}
return checkSessionOK
}
// removeSession removes the session from the active sessions and the disk.
func (a *Auth) removeSession(sess string) {
key, _ := hex.DecodeString(sess)
a.lock.Lock()
delete(a.sessions, sess)
a.lock.Unlock()
a.removeSessionFromFile(key)
}
// addUser adds a new user with the given password.
func (a *Auth) addUser(u *webUser, password string) (err error) {
// addUser adds a new user with the given password. u must not be nil.
func (a *auth) addUser(ctx context.Context, u *webUser, password string) (err error) {
if len(password) == 0 {
return errors.Error("empty password")
}
@@ -323,97 +168,21 @@ func (a *Auth) addUser(u *webUser, password string) (err error) {
u.PasswordHash = string(hash)
a.lock.Lock()
defer a.lock.Unlock()
err = a.users.Create(ctx, u.toUser())
if err != nil {
// Should not happen.
panic(err)
}
a.users = append(a.users, *u)
log.Debug("auth: added user with login %q", u.Name)
a.logger.DebugContext(ctx, "added user", "login", u.Name)
return nil
}
// findUser returns a user if there is one.
func (a *Auth) findUser(login, password string) (u webUser, ok bool) {
a.lock.Lock()
defer a.lock.Unlock()
for _, u = range a.users {
if u.Name == login &&
bcrypt.CompareHashAndPassword([]byte(u.PasswordHash), []byte(password)) == nil {
return u, true
}
}
return webUser{}, false
}
// getCurrentUser returns the current user. It returns an empty User if the
// user is not found.
func (a *Auth) getCurrentUser(r *http.Request) (u webUser) {
cookie, err := r.Cookie(sessionCookieName)
// close closes the authentication database.
func (a *auth) close(ctx context.Context) {
err := a.sessions.Close()
if err != nil {
// There's no Cookie, check Basic authentication.
user, pass, ok := r.BasicAuth()
if ok {
u, _ = globalContext.auth.findUser(user, pass)
return u
}
return webUser{}
a.logger.ErrorContext(ctx, "closing session storage", slogutil.KeyError, err)
}
a.lock.Lock()
defer a.lock.Unlock()
s, ok := a.sessions[cookie.Value]
if !ok {
return webUser{}
}
for _, u = range a.users {
if u.Name == s.userName {
return u
}
}
return webUser{}
}
// usersList returns a copy of a users list.
func (a *Auth) usersList() (users []webUser) {
a.lock.Lock()
defer a.lock.Unlock()
users = make([]webUser, len(a.users))
copy(users, a.users)
return users
}
// authRequired returns true if a authentication is required.
func (a *Auth) authRequired() bool {
if GLMode {
return true
}
a.lock.Lock()
defer a.lock.Unlock()
return len(a.users) != 0
}
// newSessionToken returns cryptographically secure randomly generated slice of
// bytes of sessionTokenSize length.
//
// TODO(e.burkov): Think about using byte array instead of byte slice.
func newSessionToken() (data []byte) {
randData := make([]byte, sessionTokenSize)
// Since Go 1.24, crypto/rand.Read doesn't return an error and crashes
// unrecoverably instead.
_, _ = rand.Read(randData)
return randData
}

View File

@@ -1,69 +1,52 @@
package home
import (
"encoding/hex"
"path/filepath"
"testing"
"time"
"github.com/AdguardTeam/AdGuardHome/internal/aghuser"
"github.com/AdguardTeam/golibs/testutil"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/crypto/bcrypt"
)
func TestAuth(t *testing.T) {
dir := t.TempDir()
fn := filepath.Join(dir, "sessions.db")
func TestAuth_UsersList(t *testing.T) {
const (
userName = "name"
userPassword = "password"
)
users := []webUser{{
Name: "name",
PasswordHash: "$2y$05$..vyzAECIhJPfaQiOK17IukcQnqEgKJHy0iETyYqxn3YXJl8yZuo2",
}}
a := InitAuth(fn, nil, 60, nil, nil)
s := session{}
user := webUser{Name: "name"}
err := a.addUser(&user, "password")
passwordHash, err := bcrypt.GenerateFromPassword([]byte(userPassword), bcrypt.DefaultCost)
require.NoError(t, err)
assert.Equal(t, checkSessionNotFound, a.checkSession("notfound"))
a.removeSession("notfound")
sessionsDB := filepath.Join(t.TempDir(), "sessions.db")
sess := newSessionToken()
sessStr := hex.EncodeToString(sess)
user := webUser{
Name: userName,
PasswordHash: string(passwordHash),
UserID: aghuser.MustNewUserID(),
}
now := time.Now().UTC().Unix()
// check expiration
s.expire = uint32(now)
a.addSession(sess, &s)
assert.Equal(t, checkSessionExpired, a.checkSession(sessStr))
auth, err := newAuth(testutil.ContextWithTimeout(t, testTimeout), &authConfig{
baseLogger: testLogger,
rateLimiter: emptyRateLimiter{},
trustedProxies: nil,
dbFilename: sessionsDB,
users: nil,
sessionTTL: testTimeout,
isGLiNet: false,
})
require.NoError(t, err)
// add session with TTL = 2 sec
s = session{}
s.expire = uint32(time.Now().UTC().Unix() + 2)
a.addSession(sess, &s)
assert.Equal(t, checkSessionOK, a.checkSession(sessStr))
t.Cleanup(func() { auth.close(testutil.ContextWithTimeout(t, testTimeout)) })
a.Close()
ctx := testutil.ContextWithTimeout(t, testTimeout)
// load saved session
a = InitAuth(fn, users, 60, nil, nil)
assert.Empty(t, auth.usersList(ctx))
// the session is still alive
assert.Equal(t, checkSessionOK, a.checkSession(sessStr))
// reset our expiration time because checkSession() has just updated it
s.expire = uint32(time.Now().UTC().Unix() + 2)
a.storeSession(sess, &s)
a.Close()
err = auth.addUser(ctx, &user, userPassword)
require.NoError(t, err)
u, ok := a.findUser("name", "password")
assert.True(t, ok)
assert.NotEmpty(t, u.Name)
time.Sleep(3 * time.Second)
// load and remove expired sessions
a = InitAuth(fn, users, 60, nil, nil)
assert.Equal(t, checkSessionNotFound, a.checkSession(sessStr))
a.Close()
assert.Equal(t, []webUser{user}, auth.usersList(ctx))
}

View File

@@ -1,116 +1,40 @@
package home
import (
"bytes"
"context"
"encoding/binary"
"io"
"log/slog"
"net"
"net/http"
"net/url"
"os"
"time"
"github.com/AdguardTeam/golibs/ioutil"
"github.com/AdguardTeam/golibs/log"
"github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/AdguardTeam/golibs/netutil/httputil"
"github.com/AdguardTeam/golibs/netutil/urlutil"
"github.com/AdguardTeam/golibs/timeutil"
)
// GLMode - enable GL-Inet compatibility mode
var GLMode bool
// glFilePrefix is the prefix of the filepath where the authentication token is
// stored. Note that it is variable so it can be edited in tests.
//
// TODO(s.chzhen): Make it a constant.
var glFilePrefix = "/tmp/gl_token_"
const (
glTokenTimeoutSeconds = 3600
glCookieName = "Admin-Token"
// glTokenTimeout is the TTL (Time To Live) of the authentication token.
glTokenTimeout = 3600 * time.Second
// glCookieName is the name of the cookie that stores the authentication
// token.
glCookieName = "Admin-Token"
)
func glProcessRedirect(w http.ResponseWriter, r *http.Request) bool {
if !GLMode {
return false
}
// redirect to gl-inet login
host, _, _ := net.SplitHostPort(r.Host)
url := "http://" + host
log.Debug("Auth: redirecting to %s", url)
http.Redirect(w, r, url, http.StatusFound)
return true
}
func glProcessCookie(r *http.Request) bool {
if !GLMode {
return false
}
glCookie, glerr := r.Cookie(glCookieName)
if glerr != nil {
return false
}
log.Debug("Auth: GL cookie value: %s", glCookie.Value)
if glCheckToken(glCookie.Value) {
return true
}
log.Info("Auth: invalid GL cookie value: %s", glCookie)
return false
}
func glCheckToken(sess string) bool {
tokenName := glFilePrefix + sess
_, err := os.Stat(tokenName)
if err != nil {
log.Error("os.Stat: %s", err)
return false
}
tokenDate := glGetTokenDate(tokenName)
now := uint32(time.Now().UTC().Unix())
return now <= (tokenDate + glTokenTimeoutSeconds)
}
// MaxFileSize is a maximum file length in bytes.
const MaxFileSize = 1024 * 1024
func glGetTokenDate(file string) uint32 {
f, err := os.Open(file)
if err != nil {
log.Error("os.Open: %s", err)
return 0
}
defer func() {
derr := f.Close()
if derr != nil {
log.Error("glinet: closing file: %s", err)
}
}()
fileReader := ioutil.LimitReader(f, MaxFileSize)
var dateToken uint32
// This use of ReadAll is now safe, because we limited reader.
bs, err := io.ReadAll(fileReader)
if err != nil {
log.Error("reading token: %s", err)
return 0
}
buf := bytes.NewBuffer(bs)
err = binary.Read(buf, binary.NativeEndian, &dateToken)
if err != nil {
log.Error("decoding token: %s", err)
return 0
}
return dateToken
}
// authMiddlewareGLiNetConfig is the configuration structure for the GLiNet
// authentication middleware.
type authMiddlewareGLiNetConfig struct {
@@ -166,12 +90,37 @@ var _ httputil.Middleware = (*authMiddlewareGLiNet)(nil)
func (mw *authMiddlewareGLiNet) Wrap(h http.Handler) (wrapped http.Handler) {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
path := r.URL.Path
if isPublicResource(path) {
h.ServeHTTP(w, r)
return
}
if mw.isAuthenticated(ctx, r) {
h.ServeHTTP(w, r)
return
}
if path == "/" || path == "/index.html" {
host := r.Host
if h, _, err := net.SplitHostPort(r.Host); err == nil {
host = h
}
u := &url.URL{
Scheme: urlutil.SchemeHTTP,
Host: host,
}
http.Redirect(w, r, u.String(), http.StatusFound)
return
}
w.WriteHeader(http.StatusUnauthorized)
})
}

View File

@@ -56,7 +56,7 @@ func TestAuthMiddlewareGLiNet(t *testing.T) {
}{{
req: httptest.NewRequest(http.MethodGet, "/", nil),
name: "no_cookie",
wantCode: http.StatusUnauthorized,
wantCode: http.StatusFound,
}, {
req: reqValidCookie,
name: "valid_cookie",
@@ -64,7 +64,7 @@ func TestAuthMiddlewareGLiNet(t *testing.T) {
}, {
req: reqInvalidCookie,
name: "invalid_cookie",
wantCode: http.StatusUnauthorized,
wantCode: http.StatusFound,
}}
for _, tc := range testCases {
@@ -78,25 +78,3 @@ func TestAuthMiddlewareGLiNet(t *testing.T) {
})
}
}
func TestAuthGL(t *testing.T) {
dir := t.TempDir()
GLMode = true
t.Cleanup(func() { GLMode = false })
glFilePrefix = dir + "/gl_token_"
data := make([]byte, 4)
binary.NativeEndian.PutUint32(data, 1)
require.NoError(t, os.WriteFile(glFilePrefix+"test", data, 0o644))
assert.False(t, glCheckToken("test"))
data = make([]byte, 4)
binary.NativeEndian.PutUint32(data, uint32(time.Now().UTC().Unix()+60))
require.NoError(t, os.WriteFile(glFilePrefix+"test", data, 0o644))
r, _ := http.NewRequest(http.MethodGet, "http://localhost/", nil)
r.AddCookie(&http.Cookie{Name: glCookieName, Value: "test"})
assert.True(t, glProcessCookie(r))
}

View File

@@ -9,6 +9,7 @@ import (
"net/http"
"net/netip"
"path"
"slices"
"strconv"
"strings"
"time"
@@ -37,40 +38,6 @@ type loginJSON struct {
Password string `json:"password"`
}
// newCookie creates a new authentication cookie.
func (a *Auth) newCookie(req loginJSON, addr string) (c *http.Cookie, err error) {
rateLimiter := a.rateLimiter
u, ok := a.findUser(req.Name, req.Password)
if !ok {
if rateLimiter != nil {
rateLimiter.inc(addr)
}
return nil, errors.Error("invalid username or password")
}
if rateLimiter != nil {
rateLimiter.remove(addr)
}
sess := newSessionToken()
now := time.Now().UTC()
a.addSession(sess, &session{
userName: u.Name,
expire: uint32(now.Unix()) + a.sessionTTL,
})
return &http.Cookie{
Name: sessionCookieName,
Value: hex.EncodeToString(sess),
Path: "/",
Expires: now.Add(cookieTTL),
HttpOnly: true,
SameSite: http.SameSiteLaxMode,
}, nil
}
// realIP extracts the real IP address of the client from an HTTP request using
// the known HTTP headers.
//
@@ -130,7 +97,9 @@ func writeErrorWithIP(
}
// handleLogin is the handler for the POST /control/login HTTP API.
func handleLogin(w http.ResponseWriter, r *http.Request) {
func (web *webAPI) handleLogin(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
req := loginJSON{}
err := json.NewDecoder(r.Body).Decode(&req)
if err != nil {
@@ -140,8 +109,8 @@ func handleLogin(w http.ResponseWriter, r *http.Request) {
}
var remoteIP string
// realIP cannot be used here without taking TrustedProxies into account due
// to security issues.
// The real IP address of the client [realIP] cannot be used here without
// taking trusted proxies into account due to security issues:
//
// See https://github.com/AdguardTeam/AdGuardHome/issues/2799.
if remoteIP, err = netutil.SplitHost(r.RemoteAddr); err != nil {
@@ -157,7 +126,7 @@ func handleLogin(w http.ResponseWriter, r *http.Request) {
return
}
if rateLimiter := globalContext.auth.rateLimiter; rateLimiter != nil {
if rateLimiter := web.auth.rateLimiter; rateLimiter != nil {
if left := rateLimiter.check(remoteIP); left > 0 {
w.Header().Set(httphdr.RetryAfter, strconv.Itoa(int(left.Seconds())))
writeErrorWithIP(
@@ -175,13 +144,18 @@ func handleLogin(w http.ResponseWriter, r *http.Request) {
ip, err := realIP(r)
if err != nil {
log.Error("auth: getting real ip from request with remote ip %s: %s", remoteIP, err)
web.logger.ErrorContext(
ctx,
"getting real ip",
"remote_ip", remoteIP,
slogutil.KeyError, err,
)
}
cookie, err := globalContext.auth.newCookie(req, remoteIP)
cookie, err := newCookie(ctx, web.auth, req, remoteIP)
if err != nil {
logIP := remoteIP
if globalContext.auth.trustedProxies.Contains(ip.Unmap()) {
if web.auth.trustedProxies.Contains(ip.Unmap()) {
logIP = ip.String()
}
@@ -190,7 +164,7 @@ func handleLogin(w http.ResponseWriter, r *http.Request) {
return
}
log.Info("auth: user %q successfully logged in from ip %s", req.Name, ip)
web.logger.InfoContext(ctx, "successful login", "user", req.Name, "ip", ip)
http.SetCookie(w, cookie)
@@ -202,8 +176,54 @@ func handleLogin(w http.ResponseWriter, r *http.Request) {
aghhttp.OK(w)
}
// newCookie creates a new authentication cookie. rateLimiter must not be nil.
func newCookie(
ctx context.Context,
auth *auth,
req loginJSON,
addr string,
) (c *http.Cookie, err error) {
user, err := auth.users.ByLogin(ctx, aghuser.Login(req.Name))
if err != nil {
// Should not happen.
panic(err)
}
rateLimiter := auth.rateLimiter
if user == nil {
rateLimiter.inc(addr)
return nil, errInvalidLogin
}
ok := user.Password.Authenticate(ctx, req.Password)
if !ok {
rateLimiter.inc(addr)
return nil, errInvalidLogin
}
rateLimiter.remove(addr)
sess, err := auth.sessions.New(ctx, user)
if err != nil {
return nil, err
}
return &http.Cookie{
Name: sessionCookieName,
Value: hex.EncodeToString(sess.Token[:]),
Path: "/",
Expires: time.Now().Add(cookieTTL),
HttpOnly: true,
SameSite: http.SameSiteLaxMode,
}, nil
}
// handleLogout is the handler for the GET /control/logout HTTP API.
func handleLogout(w http.ResponseWriter, r *http.Request) {
func (web *webAPI) handleLogout(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
respHdr := w.Header()
c, err := r.Cookie(sessionCookieName)
if err != nil {
@@ -215,7 +235,19 @@ func handleLogout(w http.ResponseWriter, r *http.Request) {
return
}
globalContext.auth.removeSession(c.Value)
t, err := sessionTokenFromHex(c.Value)
if err != nil {
web.logger.ErrorContext(ctx, "getting token", slogutil.KeyError, err)
w.WriteHeader(http.StatusUnauthorized)
return
}
err = web.auth.sessions.DeleteByToken(ctx, t)
if err != nil {
web.logger.ErrorContext(ctx, "removing session by token", slogutil.KeyError, err)
}
c = &http.Cookie{
Name: sessionCookieName,
@@ -233,93 +265,12 @@ func handleLogout(w http.ResponseWriter, r *http.Request) {
}
// RegisterAuthHandlers - register handlers
func RegisterAuthHandlers() {
globalContext.mux.Handle("/control/login", postInstallHandler(ensureHandler(http.MethodPost, handleLogin)))
httpRegister(http.MethodGet, "/control/logout", handleLogout)
}
// optionalAuthThird returns true if a user should authenticate first.
func optionalAuthThird(w http.ResponseWriter, r *http.Request) (mustAuth bool) {
pref := fmt.Sprintf("auth: raddr %s", r.RemoteAddr)
if glProcessCookie(r) {
log.Debug("%s: authentication is handled by gl-inet submodule", pref)
return false
}
// redirect to login page if not authenticated
isAuthenticated := false
cookie, err := r.Cookie(sessionCookieName)
if err != nil {
// The only error that is returned from r.Cookie is [http.ErrNoCookie].
// Check Basic authentication.
user, pass, hasBasic := r.BasicAuth()
if hasBasic {
_, isAuthenticated = globalContext.auth.findUser(user, pass)
if !isAuthenticated {
log.Info("%s: invalid basic authorization value", pref)
}
}
} else {
res := globalContext.auth.checkSession(cookie.Value)
isAuthenticated = res == checkSessionOK
if !isAuthenticated {
log.Debug("%s: invalid cookie value: %q", pref, cookie)
}
}
if isAuthenticated {
return false
}
if p := r.URL.Path; p == "/" || p == "/index.html" {
if glProcessRedirect(w, r) {
log.Debug("%s: redirected to login page by gl-inet submodule", pref)
} else {
log.Debug("%s: redirected to login page", pref)
http.Redirect(w, r, "login.html", http.StatusFound)
}
} else {
log.Debug("%s: responded with forbidden to %s %s", pref, r.Method, p)
w.WriteHeader(http.StatusForbidden)
_, _ = w.Write([]byte("Forbidden"))
}
return true
}
// TODO(a.garipov): Use [http.Handler] consistently everywhere throughout the
// project.
func optionalAuth(
h func(http.ResponseWriter, *http.Request),
) (wrapped func(http.ResponseWriter, *http.Request)) {
return func(w http.ResponseWriter, r *http.Request) {
p := r.URL.Path
authRequired := globalContext.auth != nil && globalContext.auth.authRequired()
if p == "/login.html" {
cookie, err := r.Cookie(sessionCookieName)
if authRequired && err == nil {
// Redirect to the dashboard if already authenticated.
res := globalContext.auth.checkSession(cookie.Value)
if res == checkSessionOK {
http.Redirect(w, r, "", http.StatusFound)
return
}
log.Debug("auth: raddr %s: invalid cookie value: %q", r.RemoteAddr, cookie)
}
} else if isPublicResource(p) {
// Process as usual, no additional auth requirements.
} else if authRequired {
if optionalAuthThird(w, r) {
return
}
}
h(w, r)
}
func RegisterAuthHandlers(web *webAPI) {
globalContext.mux.Handle(
"/control/login",
postInstallHandler(ensureHandler(http.MethodPost, web.handleLogin)),
)
httpRegister(http.MethodGet, "/control/logout", web.handleLogout)
}
// isPublicResource returns true if p is a path to a public resource.
@@ -337,22 +288,23 @@ func isPublicResource(p string) (ok bool) {
panic(fmt.Errorf("bad login pattern: %w", err))
}
return isAsset || isLogin
}
// TODO(s.chzhen): Implement a more strict version.
if strings.HasPrefix(p, "/dns-query/") {
return true
}
// authHandler is a helper structure that implements [http.Handler].
type authHandler struct {
handler http.Handler
}
paths := []string{
"/dns-query",
"/control/login",
"/apple/doh.mobileconfig",
"/apple/dot.mobileconfig",
"/control/install/get_addresses",
"/control/install/check_config",
"/control/install/configure",
"/install.html",
}
// ServeHTTP implements the [http.Handler] interface for *authHandler.
func (a *authHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
optionalAuth(a.handler.ServeHTTP)(w, r)
}
// optionalAuthHandler returns a authentication handler.
func optionalAuthHandler(handler http.Handler) http.Handler {
return &authHandler{handler}
return isAsset || isLogin || slices.Contains(paths, p)
}
const (
@@ -367,6 +319,15 @@ type authMiddlewareDefaultConfig struct {
// be nil.
logger *slog.Logger
// rateLimiter manages the rate limiting for login attempts.
rateLimiter loginRaateLimiter
// trustedProxies is a set of subnets considered as trusted.
//
// TODO(s.chzhen): Use it not only to pass it to the middleware but also to
// log the work of the rate limiter.
trustedProxies netutil.SubnetSet
// sessions contains web user sessions. It must not be nil.
sessions aghuser.SessionStorage
@@ -378,18 +339,22 @@ type authMiddlewareDefaultConfig struct {
// for a web client using an authentication cookie or basic auth credentials and
// passes it with the context.
type authMiddlewareDefault struct {
logger *slog.Logger
sessions aghuser.SessionStorage
users aghuser.DB
logger *slog.Logger
rateLimiter loginRaateLimiter
trustedProxies netutil.SubnetSet
sessions aghuser.SessionStorage
users aghuser.DB
}
// newAuthMiddlewareDefault returns the new properly initialized
// *authMiddlewareDefault.
func newAuthMiddlewareDefault(c *authMiddlewareDefaultConfig) (mw *authMiddlewareDefault) {
return &authMiddlewareDefault{
logger: c.logger,
sessions: c.sessions,
users: c.users,
logger: c.logger,
rateLimiter: c.rateLimiter,
trustedProxies: c.trustedProxies,
sessions: c.sessions,
users: c.users,
}
}
@@ -401,49 +366,61 @@ var _ httputil.Middleware = (*authMiddlewareDefault)(nil)
func (mw *authMiddlewareDefault) Wrap(h http.Handler) (wrapped http.Handler) {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
if !mw.needsAuthentication(ctx, r) {
if !mw.needsAuthentication(ctx) {
h.ServeHTTP(w, r)
return
}
path := r.URL.Path
u, err := mw.userFromRequest(ctx, r)
if err != nil {
mw.logger.ErrorContext(ctx, "retrieving user from request", slogutil.KeyError, err)
}
if u != nil {
if path == "/login.html" {
http.Redirect(w, r, "/", http.StatusFound)
return
}
h.ServeHTTP(w, r.WithContext(withWebUser(ctx, u)))
return
}
if err != nil {
mw.logger.ErrorContext(ctx, "retrieving user from request", slogutil.KeyError, err)
if isPublicResource(path) {
h.ServeHTTP(w, r)
return
}
if path == "/" || path == "/index.html" {
http.Redirect(w, r, "login.html", http.StatusFound)
return
}
w.WriteHeader(http.StatusUnauthorized)
})
}
// needsAuthentication returns true if the current request requires
// authentication.
//
// TODO(s.chzhen): Use the request's path.
func (mw *authMiddlewareDefault) needsAuthentication(
ctx context.Context,
_ *http.Request,
) (ok bool) {
// needsAuthentication returns true if there are stored web users and requests
// should be authenticated first.
func (mw *authMiddlewareDefault) needsAuthentication(ctx context.Context) (ok bool) {
users, err := mw.users.All(ctx)
if err != nil {
// Should not happen.
panic(err)
}
if len(users) == 0 {
return false
}
return true
return len(users) != 0
}
// userFromRequest tries to retrieve a user based on the request.
// userFromRequest tries to retrieve a user based on the request. r must not be
// nil.
func (mw *authMiddlewareDefault) userFromRequest(
ctx context.Context,
r *http.Request,
@@ -451,25 +428,24 @@ func (mw *authMiddlewareDefault) userFromRequest(
defer func() { err = errors.Annotate(err, "getting user from request: %w") }()
cookie, err := r.Cookie(sessionCookieName)
if err == http.ErrNoCookie {
return mw.userFromRequestBasicAuth(ctx, r)
if err == nil {
return mw.userFromCookie(ctx, cookie.Value)
}
sess, err := hex.DecodeString(cookie.Value)
if err != nil {
return nil, fmt.Errorf("decoding cookie: %w", err)
}
return mw.userFromRequestBasicAuth(ctx, r)
}
l := aghuser.SessionTokenLength
// TODO(a.garipov): Add validate.Len.
err = validate.InRange("token length", len(sess), l, l)
// userFromCookie tries to retrieve a user based on the provided cookie value.
func (mw *authMiddlewareDefault) userFromCookie(
ctx context.Context,
val string,
) (u *aghuser.User, err error) {
t, err := sessionTokenFromHex(val)
if err != nil {
// Don't wrap the error because it's informative enough as is.
return nil, err
}
t := aghuser.SessionToken(sess)
s, err := mw.sessions.FindByToken(ctx, t)
if err != nil {
return nil, fmt.Errorf("searching session by token: %w", err)
@@ -487,16 +463,58 @@ func (mw *authMiddlewareDefault) userFromRequest(
return u, nil
}
// userFromRequestBasicAuth searches for a user using Basic Auth credentials.
// sessionTokenFromHex converts a hexadecimal string into a session token.
func sessionTokenFromHex(val string) (token aghuser.SessionToken, err error) {
sess, err := hex.DecodeString(val)
if err != nil {
return token, fmt.Errorf("decoding value: %w", err)
}
l := aghuser.SessionTokenLength
err = validate.Equal("token length", l, len(sess))
if err != nil {
// Don't wrap the error because it's informative enough as is.
return token, err
}
return aghuser.SessionToken(sess), nil
}
// userFromRequestBasicAuth searches for a user using Basic Auth credentials. r
// must not be nil.
func (mw *authMiddlewareDefault) userFromRequestBasicAuth(
ctx context.Context,
r *http.Request,
) (user *aghuser.User, err error) {
login, pass, ok := r.BasicAuth()
if !ok {
return nil, fmt.Errorf("credentials: %w", errors.ErrNoValue)
return nil, nil
}
var remoteIP string
// The real IP address of the client [realIP] cannot be used here without
// taking trusted proxies into account due to security issues:
//
// See https://github.com/AdguardTeam/AdGuardHome/issues/2799.
if remoteIP, err = netutil.SplitHost(r.RemoteAddr); err != nil {
return nil, fmt.Errorf("getting remote address: %w", err)
}
rateLimiter := mw.rateLimiter
if left := rateLimiter.check(remoteIP); left > 0 {
return nil, fmt.Errorf("login attempt blocked for %s", left)
}
rateLimiter.inc(remoteIP)
defer func() {
if err != nil {
return
}
rateLimiter.remove(remoteIP)
}()
user, _ = mw.users.ByLogin(ctx, aghuser.Login(login))
if user == nil {
return nil, errInvalidLogin

View File

@@ -7,19 +7,18 @@ import (
"encoding/binary"
"encoding/hex"
"encoding/json"
"fmt"
"maps"
"net/http"
"net/http/httptest"
"net/netip"
"net/textproto"
"net/url"
"os"
"path/filepath"
"slices"
"testing"
"time"
"github.com/AdguardTeam/AdGuardHome/internal/agh"
"github.com/AdguardTeam/AdGuardHome/internal/aghhttp"
"github.com/AdguardTeam/AdGuardHome/internal/aghuser"
"github.com/AdguardTeam/golibs/httphdr"
@@ -48,20 +47,20 @@ var _ aghuser.SessionStorage = (*testSessionStorage)(nil)
// panic.
func newTestSessionStorage() (ts *testSessionStorage) {
return &testSessionStorage{
onNew: func(_ context.Context, u *aghuser.User) (_ *aghuser.Session, _ error) {
panic(fmt.Errorf("unexpected call to testSessionStorage.New(%v)", u))
onNew: func(ctx context.Context, u *aghuser.User) (_ *aghuser.Session, _ error) {
panic(testutil.UnexpectedCall(ctx, u))
},
onFindByToken: func(
_ context.Context,
ctx context.Context,
t aghuser.SessionToken,
) (_ *aghuser.Session, err error) {
panic(fmt.Errorf("unexpected call to testSessionStorage.FindByToken(%v)", t))
panic(testutil.UnexpectedCall(ctx, t))
},
onDeleteByToken: func(_ context.Context, t aghuser.SessionToken) (_ error) {
panic(fmt.Errorf("unexpected call to testSessionStorage.DeleteByToken(%v)", t))
onDeleteByToken: func(ctx context.Context, t aghuser.SessionToken) (_ error) {
panic(testutil.UnexpectedCall(ctx, t))
},
onClose: func() (_ error) {
panic("unexpected call to testSessionStorage.Close")
panic(testutil.UnexpectedCall())
},
}
}
@@ -110,17 +109,17 @@ type testUsersDB struct {
// newTestUsersDB returns a new *testUsersDB all methods of which panic.
func newTestUsersDB() (ts *testUsersDB) {
return &testUsersDB{
onAll: func(_ context.Context) (_ []*aghuser.User, _ error) {
panic("unexpected call to testUsersDB.All")
onAll: func(ctx context.Context) (_ []*aghuser.User, _ error) {
panic(testutil.UnexpectedCall(ctx))
},
onByLogin: func(_ context.Context, l aghuser.Login) (_ *aghuser.User, _ error) {
panic(fmt.Errorf("unexpected call to testUsersDB.ByLogin(%v)", l))
onByLogin: func(ctx context.Context, l aghuser.Login) (_ *aghuser.User, _ error) {
panic(testutil.UnexpectedCall(ctx, l))
},
onByUUID: func(_ context.Context, id aghuser.UserID) (_ *aghuser.User, _ error) {
panic(fmt.Errorf("unexpected call to testUsersDB.ByUUID(%v)", id))
onByUUID: func(ctx context.Context, id aghuser.UserID) (_ *aghuser.User, _ error) {
panic(testutil.UnexpectedCall(ctx, id))
},
onCreate: func(_ context.Context, u *aghuser.User) (_ error) {
panic(fmt.Errorf("unexpected call to testUsersDB.Create(%v)", u))
onCreate: func(ctx context.Context, u *aghuser.User) (_ error) {
panic(testutil.UnexpectedCall(ctx, u))
},
}
}
@@ -166,40 +165,18 @@ func (h *testAuthHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
h.user, _ = webUserFromContext(r.Context())
}
func TestAuthMiddlewareDefault_firstRun(t *testing.T) {
db := newTestUsersDB()
db.onAll = func(_ context.Context) (users []*aghuser.User, err error) {
return nil, nil
}
mw := newAuthMiddlewareDefault(&authMiddlewareDefaultConfig{
logger: testLogger,
sessions: &testSessionStorage{},
users: db,
})
h := &testAuthHandler{}
wrapped := mw.Wrap(h)
w := httptest.NewRecorder()
r := httptest.NewRequest(http.MethodGet, "/", nil)
wrapped.ServeHTTP(w, r)
assert.Equal(t, http.StatusOK, w.Code)
assert.True(t, h.called)
}
func TestAuthMiddlewareDefault(t *testing.T) {
t.Parallel()
const (
login aghuser.Login = "user_login"
loginStr = "user_login"
passwordStr = "user_password"
passwordRaw = "user_password"
login = aghuser.Login(loginStr)
)
passwordHash, err := bcrypt.GenerateFromPassword(
[]byte(passwordRaw),
[]byte(passwordStr),
bcrypt.DefaultCost,
)
require.NoError(t, err)
@@ -238,22 +215,14 @@ func TestAuthMiddlewareDefault(t *testing.T) {
}
mw := newAuthMiddlewareDefault(&authMiddlewareDefaultConfig{
logger: testLogger,
sessions: ts,
users: usersDB,
logger: testLogger,
rateLimiter: emptyRateLimiter{},
sessions: ts,
users: usersDB,
})
reqCookie := httptest.NewRequest(http.MethodGet, "/", nil)
reqCookie.AddCookie(&http.Cookie{Name: sessionCookieName, Value: tokenHex})
reqInvalidCookie := httptest.NewRequest(http.MethodGet, "/", nil)
reqInvalidCookie.AddCookie(&http.Cookie{Name: sessionCookieName, Value: "invalid_cookie"})
reqBasicAuth := httptest.NewRequest(http.MethodGet, "/", nil)
reqBasicAuth.SetBasicAuth(string(login), passwordRaw)
reqInvalidPassBasicAuth := httptest.NewRequest(http.MethodGet, "/", nil)
reqInvalidPassBasicAuth.SetBasicAuth(string(login), "invalid_password")
cookie := &http.Cookie{Name: sessionCookieName, Value: tokenHex}
invalidCookie := &http.Cookie{Name: sessionCookieName, Value: "123"}
testCases := []struct {
req *http.Request
@@ -263,28 +232,58 @@ func TestAuthMiddlewareDefault(t *testing.T) {
}{{
req: httptest.NewRequest(http.MethodGet, "/", nil),
wantUser: nil,
name: "no_auth",
wantCode: http.StatusUnauthorized,
name: "no_auth_root",
wantCode: http.StatusFound,
}, {
req: reqCookie,
req: httptest.NewRequest(http.MethodGet, "/index.html", nil),
wantUser: nil,
name: "no_auth",
wantCode: http.StatusFound,
}, {
req: authRequest("/", invalidCookie, "", ""),
wantUser: nil,
name: "invalid_auth",
wantCode: http.StatusFound,
}, {
req: authRequest("/", cookie, "", ""),
wantUser: user,
name: "cookie",
wantCode: http.StatusOK,
}, {
req: reqBasicAuth,
req: authRequest("/login.html", cookie, "", ""),
wantUser: nil,
name: "redirect",
wantCode: http.StatusFound,
}, {
req: authRequest("/control/profile", cookie, "", ""),
wantUser: user,
name: "protected",
wantCode: http.StatusOK,
}, {
req: authRequest("/control/profile", invalidCookie, "", ""),
wantUser: nil,
name: "no_auth_protected",
wantCode: http.StatusUnauthorized,
}, {
req: httptest.NewRequest(http.MethodGet, "/control/login", nil),
wantUser: nil,
name: "public",
wantCode: http.StatusOK,
}, {
req: authRequest("/", nil, loginStr, passwordStr),
wantUser: user,
name: "basic_auth",
wantCode: http.StatusOK,
}, {
req: reqInvalidCookie,
req: authRequest("/", invalidCookie, "", ""),
wantUser: nil,
name: "invalid_cookie",
wantCode: http.StatusUnauthorized,
wantCode: http.StatusFound,
}, {
req: reqInvalidPassBasicAuth,
req: authRequest("/", nil, "invalid", "creds"),
wantUser: nil,
name: "invalid_basic_auth",
wantCode: http.StatusUnauthorized,
wantCode: http.StatusFound,
}}
for _, tc := range testCases {
@@ -303,6 +302,22 @@ func TestAuthMiddlewareDefault(t *testing.T) {
}
}
// authRequest is a test helper function that returns a GET request configured
// with the provided credentials and path.
func authRequest(path string, c *http.Cookie, user, pass string) (r *http.Request) {
r = httptest.NewRequest(http.MethodGet, path, nil)
if c != nil {
r.AddCookie(c)
}
if user != "" {
r.SetBasicAuth(user, pass)
}
return r
}
func TestAuth_ServeHTTP_firstRun(t *testing.T) {
storeGlobals(t)
@@ -312,7 +327,17 @@ func TestAuth_ServeHTTP_firstRun(t *testing.T) {
globalContext.mux = mux
ctx := testutil.ContextWithTimeout(t, testTimeout)
web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, false)
web, err := initWeb(
ctx,
options{},
nil,
nil,
testLogger,
nil,
nil,
agh.EmptyConfigModifier{},
false,
)
require.NoError(t, err)
globalContext.web = web
@@ -445,25 +470,49 @@ func TestAuth_ServeHTTP_auth(t *testing.T) {
Name: userName,
PasswordHash: string(passwordHash),
}}
auth := InitAuth(sessionsDB, users, testTTL, nil, nil)
t.Cleanup(auth.Close)
globalContext.auth = auth
mux := http.NewServeMux()
globalContext.mux = mux
auth, err := newAuth(testutil.ContextWithTimeout(t, testTimeout), &authConfig{
baseLogger: testLogger,
rateLimiter: emptyRateLimiter{},
trustedProxies: nil,
dbFilename: sessionsDB,
users: users,
sessionTTL: testTTL * time.Second,
isGLiNet: false,
})
require.NoError(t, err)
t.Cleanup(func() { auth.close(testutil.ContextWithTimeout(t, testTimeout)) })
globalContext.mux = http.NewServeMux()
tlsMgr, err := newTLSManager(testutil.ContextWithTimeout(t, testTimeout), &tlsManagerConfig{
logger: testLogger,
configModified: func() {},
logger: testLogger,
confModifier: agh.EmptyConfigModifier{},
})
require.NoError(t, err)
ctx := testutil.ContextWithTimeout(t, testTimeout)
web, err := initWeb(ctx, options{}, nil, nil, testLogger, tlsMgr, false)
web, err := initWeb(
ctx,
options{},
nil,
nil,
testLogger,
tlsMgr,
auth,
agh.EmptyConfigModifier{},
false,
)
require.NoError(t, err)
globalContext.web = web
mux := auth.middleware().Wrap(globalContext.mux)
auth.isGLiNet = true
gliNetMw := auth.middleware().Wrap(globalContext.mux)
loginCookie := generateAuthCookie(t, mux, userName, userPassword)
testCases := []struct {
@@ -506,7 +555,7 @@ func TestAuth_ServeHTTP_auth(t *testing.T) {
for _, tc := range testCases {
t.Run(tc.path, func(t *testing.T) {
r := httptest.NewRequest(tc.method, tc.path, nil)
assertHandlerStatusCode(t, mux, r, http.StatusForbidden)
assertHandlerStatusCode(t, mux, r, http.StatusUnauthorized)
r = httptest.NewRequest(tc.method, tc.path, nil)
r.SetBasicAuth(userName, userPassword)
@@ -516,22 +565,19 @@ func TestAuth_ServeHTTP_auth(t *testing.T) {
r.AddCookie(loginCookie)
assertHandlerStatusCode(t, mux, r, tc.wantCode)
GLMode = true
t.Cleanup(func() { GLMode = false })
r.AddCookie(&http.Cookie{Name: glCookieName, Value: "test"})
assertHandlerStatusCode(t, mux, r, tc.wantCode)
assertHandlerStatusCode(t, gliNetMw, r, tc.wantCode)
})
}
}
// generateAuthCookie is a helper function that logs in with the provided
// credentials and returns the resulting authentication cookie.
func generateAuthCookie(t *testing.T, mux *http.ServeMux, name, password string) (ac *http.Cookie) {
t.Helper()
func generateAuthCookie(tb testing.TB, mux http.Handler, name, password string) (ac *http.Cookie) {
tb.Helper()
creds, err := json.Marshal(&loginJSON{Name: name, Password: password})
require.NoError(t, err)
require.NoError(tb, err)
r := httptest.NewRequest(http.MethodPost, "/control/login", bytes.NewReader(creds))
r.Header.Set(httphdr.ContentType, aghhttp.HdrValApplicationJSON)
@@ -541,22 +587,26 @@ func generateAuthCookie(t *testing.T, mux *http.ServeMux, name, password string)
for _, c := range w.Result().Cookies() {
if c.Name == sessionCookieName {
return c
ac = c
break
}
}
return nil
require.NotNil(tb, ac)
return ac
}
// assertHandlerStatusCode is a helper function that asserts the response status
// code of a HTTP handler.
func assertHandlerStatusCode(t *testing.T, h http.Handler, r *http.Request, wantCode int) {
t.Helper()
func assertHandlerStatusCode(tb testing.TB, h http.Handler, r *http.Request, wantCode int) {
tb.Helper()
w := httptest.NewRecorder()
h.ServeHTTP(w, r)
assert.Equal(t, wantCode, w.Code)
assert.Equal(tb, wantCode, w.Code)
}
func TestAuth_ServeHTTP_logout(t *testing.T) {
@@ -578,21 +628,40 @@ func TestAuth_ServeHTTP_logout(t *testing.T) {
Name: userName,
PasswordHash: string(passwordHash),
}}
auth := InitAuth(sessionsDB, users, testTTL, nil, nil)
t.Cleanup(auth.Close)
globalContext.auth = auth
mux := http.NewServeMux()
globalContext.mux = mux
auth, err := newAuth(testutil.ContextWithTimeout(t, testTimeout), &authConfig{
baseLogger: testLogger,
rateLimiter: emptyRateLimiter{},
trustedProxies: nil,
dbFilename: sessionsDB,
users: users,
sessionTTL: testTTL * time.Second,
isGLiNet: false,
})
require.NoError(t, err)
t.Cleanup(func() { auth.close(testutil.ContextWithTimeout(t, testTimeout)) })
globalContext.mux = http.NewServeMux()
ctx := testutil.ContextWithTimeout(t, testTimeout)
web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, false)
web, err := initWeb(ctx,
options{},
nil,
nil,
testLogger,
nil,
auth,
agh.EmptyConfigModifier{},
false,
)
require.NoError(t, err)
globalContext.web = web
mux := auth.middleware().Wrap(globalContext.mux)
loginCookie := generateAuthCookie(t, mux, userName, userPassword)
require.NotNil(t, loginCookie)
r := httptest.NewRequest(http.MethodGet, "/control/profile", nil)
r.AddCookie(loginCookie)
@@ -604,110 +673,7 @@ func TestAuth_ServeHTTP_logout(t *testing.T) {
r = httptest.NewRequest(http.MethodGet, "/control/profile", nil)
r.AddCookie(loginCookie)
assertHandlerStatusCode(t, mux, r, http.StatusForbidden)
}
// implements http.ResponseWriter
type testResponseWriter struct {
hdr http.Header
statusCode int
}
func (w *testResponseWriter) Header() http.Header {
return w.hdr
}
func (w *testResponseWriter) Write([]byte) (int, error) {
return 0, nil
}
func (w *testResponseWriter) WriteHeader(statusCode int) {
w.statusCode = statusCode
}
func TestAuthHTTP(t *testing.T) {
dir := t.TempDir()
fn := filepath.Join(dir, "sessions.db")
users := []webUser{
{Name: "name", PasswordHash: "$2y$05$..vyzAECIhJPfaQiOK17IukcQnqEgKJHy0iETyYqxn3YXJl8yZuo2"},
}
globalContext.auth = InitAuth(fn, users, 60, nil, nil)
handlerCalled := false
handler := func(_ http.ResponseWriter, _ *http.Request) {
handlerCalled = true
}
handler2 := optionalAuth(handler)
w := testResponseWriter{}
w.hdr = make(http.Header)
r := http.Request{}
r.Header = make(http.Header)
r.Method = http.MethodGet
// get / - we're redirected to login page
r.URL = &url.URL{Path: "/"}
handlerCalled = false
handler2(&w, &r)
assert.Equal(t, http.StatusFound, w.statusCode)
assert.NotEmpty(t, w.hdr.Get(httphdr.Location))
assert.False(t, handlerCalled)
// go to login page
loginURL := w.hdr.Get(httphdr.Location)
r.URL = &url.URL{Path: loginURL}
handlerCalled = false
handler2(&w, &r)
assert.True(t, handlerCalled)
// perform login
cookie, err := globalContext.auth.newCookie(loginJSON{Name: "name", Password: "password"}, "")
require.NoError(t, err)
require.NotNil(t, cookie)
// get /
handler2 = optionalAuth(handler)
w.hdr = make(http.Header)
r.Header.Set(httphdr.Cookie, cookie.String())
r.URL = &url.URL{Path: "/"}
handlerCalled = false
handler2(&w, &r)
assert.True(t, handlerCalled)
r.Header.Del(httphdr.Cookie)
// get / with basic auth
handler2 = optionalAuth(handler)
w.hdr = make(http.Header)
r.URL = &url.URL{Path: "/"}
r.SetBasicAuth("name", "password")
handlerCalled = false
handler2(&w, &r)
assert.True(t, handlerCalled)
r.Header.Del(httphdr.Authorization)
// get login page with a valid cookie - we're redirected to /
handler2 = optionalAuth(handler)
w.hdr = make(http.Header)
r.Header.Set(httphdr.Cookie, cookie.String())
r.URL = &url.URL{Path: loginURL}
handlerCalled = false
handler2(&w, &r)
assert.NotEmpty(t, w.hdr.Get(httphdr.Location))
assert.False(t, handlerCalled)
r.Header.Del(httphdr.Cookie)
// get login page with an invalid cookie
handler2 = optionalAuth(handler)
w.hdr = make(http.Header)
r.Header.Set(httphdr.Cookie, "bad")
r.URL = &url.URL{Path: loginURL}
handlerCalled = false
handler2(&w, &r)
assert.True(t, handlerCalled)
r.Header.Del(httphdr.Cookie)
globalContext.auth.Close()
assertHandlerStatusCode(t, mux, r, http.StatusUnauthorized)
}
func TestRealIP(t *testing.T) {

View File

@@ -9,6 +9,38 @@ import (
// cache.
const failedAuthTTL = 1 * time.Minute
// loginRaateLimiter is an interface for rate limiting login attempts.
type loginRaateLimiter interface {
// check returns the duration of time left until a user is unblocked.
// A non-positive result indicates that the user is not blocked.
check(usrID string) (left time.Duration)
// inc records a failed login attempt for the specified user.
inc(usrID string)
// remove stops tracking and blocking of the specified user.
remove(usrID string)
}
// emptyRateLimiter is the [loginRateLimiter] interface implementation that does
// nothing.
type emptyRateLimiter struct{}
// type check
var _ emptyRateLimiter = emptyRateLimiter{}
// check implements the [loginRateLimiter] interface for emptyRateLimiter. It
// always returns zero.
func (rl emptyRateLimiter) check(_ string) (left time.Duration) {
return 0
}
// inc implements the [loginRateLimiter] interface for emptyRateLimiter.
func (rl emptyRateLimiter) inc(_ string) {}
// remove implements the [loginRateLimiter] interface for emptyRateLimiter.
func (rl emptyRateLimiter) remove(_ string) {}
// failedAuth is an entry of authRateLimiter's cache.
type failedAuth struct {
until time.Time
@@ -33,6 +65,9 @@ func newAuthRateLimiter(blockDur time.Duration, maxAttempts uint) (ab *authRateL
}
}
// type check
var _ loginRaateLimiter = (*authRateLimiter)(nil)
// cleanupLocked checks each blocked users removing ones with expired TTL. For
// internal use only.
func (ab *authRateLimiter) cleanupLocked(now time.Time) {
@@ -57,8 +92,7 @@ func (ab *authRateLimiter) checkLocked(usrID string, now time.Time) (left time.D
return a.until.Sub(now)
}
// check returns the time left until unblocking. The nonpositive result should
// be interpreted as not blocked attempter.
// check implements the [loginRateLimiter] interface for *authRateLimiter.
func (ab *authRateLimiter) check(usrID string) (left time.Duration) {
now := time.Now()
@@ -91,7 +125,7 @@ func (ab *authRateLimiter) incLocked(usrID string, now time.Time) {
}
}
// inc updates the failed attempt in cache.
// inc implements the [loginRateLimiter] interface for *authRateLimiter.
func (ab *authRateLimiter) inc(usrID string) {
now := time.Now()
@@ -101,7 +135,7 @@ func (ab *authRateLimiter) inc(usrID string) {
ab.incLocked(usrID, now)
}
// remove stops any tracking and any blocking of the user.
// remove implements the [loginRateLimiter] interface for *authRateLimiter.
func (ab *authRateLimiter) remove(usrID string) {
ab.failedAuthsLock.Lock()
defer ab.failedAuthsLock.Unlock()

View File

@@ -9,6 +9,7 @@ import (
"sync"
"time"
"github.com/AdguardTeam/AdGuardHome/internal/agh"
"github.com/AdguardTeam/AdGuardHome/internal/aghnet"
"github.com/AdguardTeam/AdGuardHome/internal/arpdb"
"github.com/AdguardTeam/AdGuardHome/internal/client"
@@ -39,6 +40,10 @@ type clientsContainer struct {
// settings.
clientChecker BlockedClientChecker
// confModifier is used to update the global configuration. It must not be
// nil.
confModifier agh.ConfigModifier
// lock protects all fields.
//
// TODO(a.garipov): Use a pointer and describe which fields are protected in
@@ -52,11 +57,6 @@ type clientsContainer struct {
// safeSearchCacheTTL is the TTL of the safe search cache to use for
// persistent clients.
safeSearchCacheTTL time.Duration
// testing is a flag that disables some features for internal tests.
//
// TODO(a.garipov): Awful. Remove.
testing bool
}
// BlockedClientChecker checks if a client is blocked by the current access
@@ -78,6 +78,7 @@ func (clients *clientsContainer) Init(
arpDB arpdb.Interface,
filteringConf *filtering.Config,
sigHdlr *signalHandler,
confModifier agh.ConfigModifier,
) (err error) {
// TODO(s.chzhen): Refactor it.
if clients.storage != nil {
@@ -88,6 +89,7 @@ func (clients *clientsContainer) Init(
clients.logger = baseLogger.With(slogutil.KeyPrefix, "client_container")
clients.safeSearchCacheSize = filteringConf.SafeSearchCacheSize
clients.safeSearchCacheTTL = time.Minute * time.Duration(filteringConf.CacheTime)
clients.confModifier = confModifier
confClients := make([]*client.Persistent, 0, len(objects))
for i, o := range objects {
@@ -141,10 +143,6 @@ var webHandlersRegistered = false
// Start starts the clients container.
func (clients *clientsContainer) Start(ctx context.Context) (err error) {
if clients.testing {
return
}
if !webHandlersRegistered {
webHandlersRegistered = true
clients.registerWebHandlers()

View File

@@ -3,6 +3,7 @@ package home
import (
"testing"
"github.com/AdguardTeam/AdGuardHome/internal/agh"
"github.com/AdguardTeam/AdGuardHome/internal/client"
"github.com/AdguardTeam/AdGuardHome/internal/filtering"
"github.com/AdguardTeam/golibs/testutil"
@@ -11,14 +12,12 @@ import (
// newClientsContainer is a helper that creates a new clients container for
// tests.
func newClientsContainer(t *testing.T) (c *clientsContainer) {
t.Helper()
func newClientsContainer(tb testing.TB) (c *clientsContainer) {
tb.Helper()
c = &clientsContainer{
testing: true,
}
c = &clientsContainer{}
ctx := testutil.ContextWithTimeout(t, testTimeout)
ctx := testutil.ContextWithTimeout(tb, testTimeout)
err := c.Init(
ctx,
testLogger,
@@ -29,10 +28,11 @@ func newClientsContainer(t *testing.T) (c *clientsContainer) {
&filtering.Config{
Logger: testLogger,
},
newSignalHandler(nil, nil),
newSignalHandler(testLogger, nil, nil),
agh.EmptyConfigModifier{},
)
require.NoError(t, err)
require.NoError(tb, err)
return c
}

View File

@@ -326,6 +326,8 @@ func clientToJSON(c *client.Persistent) (cj *clientJSON) {
// handleAddClient is the handler for POST /control/clients/add HTTP API.
func (clients *clientsContainer) handleAddClient(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
cj := clientJSON{}
err := json.NewDecoder(r.Body).Decode(&cj)
if err != nil {
@@ -334,27 +336,27 @@ func (clients *clientsContainer) handleAddClient(w http.ResponseWriter, r *http.
return
}
c, err := clients.jsonToClient(r.Context(), cj, nil)
c, err := clients.jsonToClient(ctx, cj, nil)
if err != nil {
aghhttp.Error(r, w, http.StatusBadRequest, "%s", err)
return
}
err = clients.storage.Add(r.Context(), c)
err = clients.storage.Add(ctx, c)
if err != nil {
aghhttp.Error(r, w, http.StatusBadRequest, "%s", err)
return
}
if !clients.testing {
onConfigModified()
}
clients.confModifier.Apply(ctx)
}
// handleDelClient is the handler for POST /control/clients/delete HTTP API.
func (clients *clientsContainer) handleDelClient(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
cj := clientJSON{}
err := json.NewDecoder(r.Body).Decode(&cj)
if err != nil {
@@ -369,15 +371,13 @@ func (clients *clientsContainer) handleDelClient(w http.ResponseWriter, r *http.
return
}
if !clients.storage.RemoveByName(r.Context(), cj.Name) {
if !clients.storage.RemoveByName(ctx, cj.Name) {
aghhttp.Error(r, w, http.StatusBadRequest, "Client not found")
return
}
if !clients.testing {
onConfigModified()
}
clients.confModifier.Apply(ctx)
}
// updateJSON contains the name and data of the updated persistent client.
@@ -390,6 +390,8 @@ type updateJSON struct {
//
// TODO(s.chzhen): Accept updated parameters instead of whole structure.
func (clients *clientsContainer) handleUpdateClient(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
dj := updateJSON{}
err := json.NewDecoder(r.Body).Decode(&dj)
if err != nil {
@@ -404,23 +406,21 @@ func (clients *clientsContainer) handleUpdateClient(w http.ResponseWriter, r *ht
return
}
c, err := clients.jsonToClient(r.Context(), dj.Data, nil)
c, err := clients.jsonToClient(ctx, dj.Data, nil)
if err != nil {
aghhttp.Error(r, w, http.StatusBadRequest, "%s", err)
return
}
err = clients.storage.Update(r.Context(), dj.Name, c)
err = clients.storage.Update(ctx, dj.Name, c)
if err != nil {
aghhttp.Error(r, w, http.StatusBadRequest, "%s", err)
return
}
if !clients.testing {
onConfigModified()
}
clients.confModifier.Apply(ctx)
}
// handleFindClient is the handler for GET /control/clients/find HTTP API.
@@ -475,7 +475,13 @@ func (clients *clientsContainer) findClient(
params.RemoteIP,
string(params.ClientID),
)
cj.Disallowed, cj.DisallowedRule = &disallowed, &rule
cj.Disallowed = &disallowed
if disallowed && rule != "" {
// Since "disallowed_rule" is omitted from JSON unless present, it
// should only be set when the client is actually blocked.
cj.DisallowedRule = &rule
}
return cj
}
@@ -554,12 +560,19 @@ func (clients *clientsContainer) findRuntime(
// See https://github.com/AdguardTeam/AdGuardHome/issues/2428.
disallowed, rule := clients.clientChecker.IsBlockedClient(ip, string(params.ClientID))
var disallowedRule *string
if disallowed && rule != "" {
// Since "disallowed_rule" is omitted from JSON unless present, it
// should only be set when the client is actually blocked.
disallowedRule = &rule
}
return &clientJSON{
Name: host,
IDs: []string{idStr},
WHOIS: whois,
Disallowed: &disallowed,
DisallowedRule: &rule,
DisallowedRule: disallowedRule,
}
}

View File

@@ -421,7 +421,6 @@ func TestClientsContainer_HandleSearchClient(t *testing.T) {
allowed = false
dissallowed = true
emptyRule = ""
disallowedRule = "disallowed_rule"
)
@@ -432,7 +431,7 @@ func TestClientsContainer_HandleSearchClient(t *testing.T) {
return true, disallowedRule
}
return false, emptyRule
return false, ""
},
}
@@ -481,11 +480,10 @@ func TestClientsContainer_HandleSearchClient(t *testing.T) {
}},
},
wantRuntime: &clientJSON{
Name: runtimeCli,
IDs: []string{runtimeCliIP},
Disallowed: &allowed,
DisallowedRule: &emptyRule,
WHOIS: &whois.Info{},
Name: runtimeCli,
IDs: []string{runtimeCliIP},
Disallowed: &allowed,
WHOIS: &whois.Info{},
},
}, {
name: "blocked_access",
@@ -508,10 +506,9 @@ func TestClientsContainer_HandleSearchClient(t *testing.T) {
}},
},
wantRuntime: &clientJSON{
IDs: []string{nonExistentCliIP},
Disallowed: &allowed,
DisallowedRule: &emptyRule,
WHOIS: &whois.Info{},
IDs: []string{nonExistentCliIP},
Disallowed: &allowed,
WHOIS: &whois.Info{},
},
}}

View File

@@ -2,13 +2,16 @@ package home
import (
"bytes"
"context"
"fmt"
"log/slog"
"net/netip"
"os"
"path/filepath"
"slices"
"sync"
"github.com/AdguardTeam/AdGuardHome/internal/agh"
"github.com/AdguardTeam/AdGuardHome/internal/aghalg"
"github.com/AdguardTeam/AdGuardHome/internal/aghos"
"github.com/AdguardTeam/AdGuardHome/internal/aghtls"
@@ -22,6 +25,7 @@ import (
"github.com/AdguardTeam/dnsproxy/fastip"
"github.com/AdguardTeam/golibs/errors"
"github.com/AdguardTeam/golibs/log"
"github.com/AdguardTeam/golibs/logutil/slogutil"
"github.com/AdguardTeam/golibs/netutil"
"github.com/AdguardTeam/golibs/timeutil"
"github.com/google/go-cmp/cmp"
@@ -462,7 +466,8 @@ var config = &configuration{
}, {
Prefix: netip.MustParsePrefix("::1/128"),
}},
CacheSize: 4 * 1024 * 1024,
CacheEnabled: true,
CacheSize: 4 * 1024 * 1024,
EDNSClientSubnet: &dnsforward.EDNSClientSubnet{
CustomIP: netip.Addr{},
@@ -743,12 +748,13 @@ func readConfigFile() (fileData []byte, err error) {
}
// Saves configuration to the YAML file and also saves the user filter contents to a file
func (c *configuration) write(tlsMgr *tlsManager) (err error) {
func (c *configuration) write(tlsMgr *tlsManager, auth *auth) (err error) {
c.Lock()
defer c.Unlock()
if globalContext.auth != nil {
config.Users = globalContext.auth.usersList()
if auth != nil {
// TODO(s.chzhen): Pass context.
config.Users = auth.usersList(context.TODO())
}
if tlsMgr != nil {
@@ -836,3 +842,47 @@ func validateTLSCipherIDs(cipherIDs []string) (err error) {
return nil
}
// defaultConfigModifier is a default [agh.ConfigModifier] implementation.
type defaultConfigModifier struct {
auth *auth
config *configuration
logger *slog.Logger
tlsMgr *tlsManager
}
// newDefaultConfigModifier returns the new properly initialized
// *defaultConfigModifier. All arguments must not be nil.
//
// TODO(s.chzhen): Consider using configuration struct.
func newDefaultConfigModifier(
conf *configuration,
l *slog.Logger,
) (cm *defaultConfigModifier) {
return &defaultConfigModifier{
config: conf,
logger: l,
}
}
// type check
var _ agh.ConfigModifier = (*defaultConfigModifier)(nil)
// Apply implements the [agh.ConfigModifier] interface for
// *defaultConfigModifier.
func (cm *defaultConfigModifier) Apply(ctx context.Context) {
err := cm.config.write(cm.tlsMgr, cm.auth)
if err != nil {
cm.logger.ErrorContext(ctx, "writing config", slogutil.KeyError, err)
}
}
// setAuth sets the auth parameters used by Apply.
func (cm *defaultConfigModifier) setAuth(a *auth) {
cm.auth = a
}
// setTLSManager sets the TLS manager used by Apply.
func (cm *defaultConfigModifier) setTLSManager(m *tlsManager) {
cm.tlsMgr = m
}

View File

@@ -115,6 +115,8 @@ type statusResponse struct {
}
func (web *webAPI) handleStatus(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
dnsAddrs, err := collectDNSAddresses(web.tlsManager)
if err != nil {
// Don't add a lot of formatting, since the error is already
@@ -125,14 +127,14 @@ func (web *webAPI) handleStatus(w http.ResponseWriter, r *http.Request) {
}
var (
fltConf *dnsforward.Config
protectionDisabledUntil *time.Time
protectionEnabled bool
fltConf *dnsforward.Config
protDisabledUntil *time.Time
protEnabled bool
)
if globalContext.dnsServer != nil {
fltConf = &dnsforward.Config{}
globalContext.dnsServer.WriteDiskConfig(fltConf)
protectionEnabled, protectionDisabledUntil = globalContext.dnsServer.UpdatedProtectionStatus()
protEnabled, protDisabledUntil = globalContext.dnsServer.UpdatedProtectionStatus(ctx)
}
var resp statusResponse
@@ -141,11 +143,11 @@ func (web *webAPI) handleStatus(w http.ResponseWriter, r *http.Request) {
defer config.RUnlock()
var protectionDisabledDuration int64
if protectionDisabledUntil != nil {
if protDisabledUntil != nil {
// Make sure that we don't send negative numbers to the frontend,
// since enough time might have passed to make the difference less
// than zero.
protectionDisabledDuration = max(0, time.Until(*protectionDisabledUntil).Milliseconds())
protectionDisabledDuration = max(0, time.Until(*protDisabledUntil).Milliseconds())
}
resp = statusResponse{
@@ -155,7 +157,7 @@ func (web *webAPI) handleStatus(w http.ResponseWriter, r *http.Request) {
DNSPort: config.DNS.Port,
HTTPPort: config.HTTPConfig.Address.Port(),
ProtectionDisabledDuration: protectionDisabledDuration,
ProtectionEnabled: protectionEnabled,
ProtectionEnabled: protEnabled,
IsRunning: isRunning(),
}
}()
@@ -171,22 +173,19 @@ func (web *webAPI) handleStatus(w http.ResponseWriter, r *http.Request) {
// registerControlHandlers sets up HTTP handlers for various control endpoints.
// web must not be nil.
func registerControlHandlers(web *webAPI) {
globalContext.mux.HandleFunc(
"/control/version.json",
postInstall(optionalAuth(web.handleVersionJSON)),
)
globalContext.mux.HandleFunc("/control/version.json", postInstall(web.handleVersionJSON))
httpRegister(http.MethodPost, "/control/update", web.handleUpdate)
httpRegister(http.MethodGet, "/control/status", web.handleStatus)
httpRegister(http.MethodPost, "/control/i18n/change_language", handleI18nChangeLanguage)
httpRegister(http.MethodPost, "/control/i18n/change_language", web.handleI18nChangeLanguage)
httpRegister(http.MethodGet, "/control/i18n/current_language", handleI18nCurrentLanguage)
httpRegister(http.MethodGet, "/control/profile", handleGetProfile)
httpRegister(http.MethodPut, "/control/profile/update", handlePutProfile)
httpRegister(http.MethodGet, "/control/profile", web.handleGetProfile)
httpRegister(http.MethodPut, "/control/profile/update", web.handlePutProfile)
// No auth is necessary for DoH/DoT configurations
globalContext.mux.HandleFunc("/apple/doh.mobileconfig", postInstall(handleMobileConfigDoH))
globalContext.mux.HandleFunc("/apple/dot.mobileconfig", postInstall(handleMobileConfigDoT))
RegisterAuthHandlers()
RegisterAuthHandlers(web)
}
// httpRegister registers an HTTP handler.
@@ -197,7 +196,10 @@ func httpRegister(method, url string, handler http.HandlerFunc) {
return
}
globalContext.mux.Handle(url, postInstallHandler(optionalAuthHandler(gziphandler.GzipHandler(ensureHandler(method, handler)))))
globalContext.mux.Handle(
url,
postInstallHandler(gziphandler.GzipHandler(ensureHandler(method, handler))),
)
}
// ensure returns a wrapped handler that makes sure that the request has the

View File

@@ -15,6 +15,7 @@ import (
"time"
"unicode/utf8"
"github.com/AdguardTeam/AdGuardHome/internal/agh"
"github.com/AdguardTeam/AdGuardHome/internal/aghalg"
"github.com/AdguardTeam/AdGuardHome/internal/aghhttp"
"github.com/AdguardTeam/AdGuardHome/internal/aghnet"
@@ -392,6 +393,8 @@ const PasswordMinRunes = 8
// Apply new configuration, start DNS server, restart Web server
func (web *webAPI) handleInstallConfigure(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
req, restartHTTP, err := decodeApplyConfigReq(r.Body)
if err != nil {
aghhttp.Error(r, w, http.StatusBadRequest, "%s", err)
@@ -440,7 +443,7 @@ func (web *webAPI) handleInstallConfigure(w http.ResponseWriter, r *http.Request
u := &webUser{
Name: req.Username,
}
err = globalContext.auth.addUser(u, req.Password)
err = web.auth.addUser(ctx, u, req.Password)
if err != nil {
globalContext.firstRun = true
copyInstallSettings(config, curConfig)
@@ -453,7 +456,7 @@ func (web *webAPI) handleInstallConfigure(w http.ResponseWriter, r *http.Request
// moment we'll allow setting up TLS in the initial configuration or the
// configuration itself will use HTTPS protocol, because the underlying
// functions potentially restart the HTTPS server.
err = startMods(r.Context(), web.baseLogger, web.tlsManager)
err = startMods(ctx, web.baseLogger, web.tlsManager, web.confModifier)
if err != nil {
globalContext.firstRun = true
copyInstallSettings(config, curConfig)
@@ -462,7 +465,7 @@ func (web *webAPI) handleInstallConfigure(w http.ResponseWriter, r *http.Request
return
}
err = config.write(web.tlsManager)
err = config.write(web.tlsManager, web.auth)
if err != nil {
globalContext.firstRun = true
copyInstallSettings(config, curConfig)
@@ -489,11 +492,11 @@ func (web *webAPI) handleInstallConfigure(w http.ResponseWriter, r *http.Request
// and with its own context, because it waits until all requests are handled
// and will be blocked by it's own caller.
go func(timeout time.Duration) {
ctx, cancel := context.WithTimeout(context.Background(), timeout)
defer slogutil.RecoverAndLog(ctx, web.logger)
shutdownCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), timeout)
defer slogutil.RecoverAndLog(shutdownCtx, web.logger)
defer cancel()
shutdownSrv(ctx, web.logger, web.httpServer)
shutdownSrv(shutdownCtx, web.logger, web.httpServer)
}(shutdownTimeout)
}
@@ -530,13 +533,18 @@ func decodeApplyConfigReq(r io.Reader) (req *applyConfigReq, restartHTTP bool, e
// startMods initializes and starts the DNS server after installation.
// baseLogger and tlsMgr must not be nil.
func startMods(ctx context.Context, baseLogger *slog.Logger, tlsMgr *tlsManager) (err error) {
func startMods(
ctx context.Context,
baseLogger *slog.Logger,
tlsMgr *tlsManager,
confModifier agh.ConfigModifier,
) (err error) {
statsDir, querylogDir, err := checkStatsAndQuerylogDirs(&globalContext, config)
if err != nil {
return err
}
err = initDNS(baseLogger, tlsMgr, statsDir, querylogDir)
err = initDNS(ctx, baseLogger, tlsMgr, confModifier, statsDir, querylogDir)
if err != nil {
return err
}
@@ -545,7 +553,7 @@ func startMods(ctx context.Context, baseLogger *slog.Logger, tlsMgr *tlsManager)
err = startDNSServer()
if err != nil {
closeDNSServer()
closeDNSServer(ctx)
return err
}

View File

@@ -12,6 +12,7 @@ import (
"path/filepath"
"time"
"github.com/AdguardTeam/AdGuardHome/internal/agh"
"github.com/AdguardTeam/AdGuardHome/internal/aghalg"
"github.com/AdguardTeam/AdGuardHome/internal/aghhttp"
"github.com/AdguardTeam/AdGuardHome/internal/aghnet"
@@ -38,23 +39,15 @@ const (
defaultPortTLS uint16 = 853
)
// Called by other modules when configuration is changed
//
// TODO(s.chzhen): Remove this after refactoring.
func onConfigModified() {
err := config.write(globalContext.tls)
if err != nil {
log.Error("writing config: %s", err)
}
}
// initDNS updates all the fields of the [globalContext] needed to initialize
// the DNS server and initializes it at last. It also must not be called unless
// [config] and [globalContext] are initialized. baseLogger and tlsMgr must not
// be nil.
// [config] and [globalContext] are initialized. baseLogger, tlsMgr and
// confModfier must not be nil.
func initDNS(
ctx context.Context,
baseLogger *slog.Logger,
tlsMgr *tlsManager,
confModifier agh.ConfigModifier,
statsDir string,
querylogDir string,
) (err error) {
@@ -64,7 +57,7 @@ func initDNS(
Logger: baseLogger.With(slogutil.KeyPrefix, "stats"),
Filename: filepath.Join(statsDir, "stats.db"),
Limit: time.Duration(config.Stats.Interval),
ConfigModified: onConfigModified,
ConfigModifier: confModifier,
HTTPRegister: httpRegister,
Enabled: config.Stats.Enabled,
ShouldCountClient: globalContext.clients.shouldCountClient,
@@ -84,7 +77,7 @@ func initDNS(
conf := querylog.Config{
Logger: baseLogger.With(slogutil.KeyPrefix, "querylog"),
Anonymizer: anonymizer,
ConfigModified: onConfigModified,
ConfigModifier: confModifier,
HTTPRegister: httpRegister,
FindClient: globalContext.clients.findMultiple,
BaseDir: querylogDir,
@@ -113,6 +106,7 @@ func initDNS(
}
return initDNSServer(
ctx,
globalContext.filters,
globalContext.stats,
globalContext.queryLog,
@@ -121,6 +115,7 @@ func initDNS(
httpRegister,
tlsMgr,
baseLogger,
confModifier,
)
}
@@ -131,6 +126,7 @@ func initDNS(
//
// TODO(e.burkov): Use [dnsforward.DNSCreateParams] as a parameter.
func initDNSServer(
ctx context.Context,
filters *filtering.DNSFilter,
sts stats.Interface,
qlog querylog.QueryLog,
@@ -139,6 +135,7 @@ func initDNSServer(
httpReg aghhttp.RegisterFunc,
tlsMgr *tlsManager,
l *slog.Logger,
confModifier agh.ConfigModifier,
) (err error) {
globalContext.dnsServer, err = dnsforward.NewServer(dnsforward.DNSCreateParams{
Logger: l,
@@ -153,7 +150,7 @@ func initDNSServer(
})
defer func() {
if err != nil {
closeDNSServer()
closeDNSServer(ctx)
}
}()
if err != nil {
@@ -169,6 +166,7 @@ func initDNSServer(
tlsMgr,
httpReg,
globalContext.clients.storage,
confModifier,
)
if err != nil {
return fmt.Errorf("newServerConfig: %w", err)
@@ -176,12 +174,12 @@ func initDNSServer(
// Try to prepare the server with disabled private RDNS resolution if it
// failed to prepare as is. See TODO on [dnsforward.PrivateRDNSError].
err = globalContext.dnsServer.Prepare(dnsConf)
err = globalContext.dnsServer.Prepare(ctx, dnsConf)
if privRDNSErr := (&dnsforward.PrivateRDNSError{}); errors.As(err, &privRDNSErr) {
log.Info("WARNING: %s; trying to disable private RDNS resolution", err)
dnsConf.UsePrivateRDNS = false
err = globalContext.dnsServer.Prepare(dnsConf)
err = globalContext.dnsServer.Prepare(ctx, dnsConf)
}
if err != nil {
@@ -245,6 +243,7 @@ func newServerConfig(
tlsMgr *tlsManager,
httpReg aghhttp.RegisterFunc,
clientsContainer dnsforward.ClientsContainer,
confModifier agh.ConfigModifier,
) (newConf *dnsforward.ServerConfig, err error) {
hosts := aghalg.CoalesceSlice(dnsConf.BindHosts, []netip.Addr{netutil.IPv4Localhost()})
@@ -264,7 +263,7 @@ func newServerConfig(
TLSAllowUnencryptedDoH: tlsConf.AllowUnencryptedDoH,
UpstreamTimeout: time.Duration(dnsConf.UpstreamTimeout),
TLSv12Roots: tlsMgr.rootCerts,
ConfigModified: onConfigModified,
ConfModifier: confModifier,
HTTPRegister: httpReg,
LocalPTRResolvers: dnsConf.PrivateRDNSResolvers,
UseDNS64: dnsConf.UseDNS64,
@@ -454,7 +453,7 @@ func startDNSServer() error {
return fmt.Errorf("starting clients container: %w", err)
}
err = globalContext.dnsServer.Start()
err = globalContext.dnsServer.Start(ctx)
if err != nil {
return fmt.Errorf("starting dns server: %w", err)
}
@@ -470,30 +469,30 @@ func startDNSServer() error {
return nil
}
func stopDNSServer() (err error) {
func stopDNSServer(ctx context.Context) (err error) {
if !isRunning() {
return nil
}
err = globalContext.dnsServer.Stop()
err = globalContext.dnsServer.Stop(ctx)
if err != nil {
return fmt.Errorf("stopping forwarding dns server: %w", err)
}
err = globalContext.clients.close(context.TODO())
err = globalContext.clients.close(ctx)
if err != nil {
return fmt.Errorf("closing clients container: %w", err)
}
closeDNSServer()
closeDNSServer(ctx)
return nil
}
func closeDNSServer() {
func closeDNSServer(ctx context.Context) {
// DNS forward module must be closed BEFORE stats or queryLog because it depends on them
if globalContext.dnsServer != nil {
globalContext.dnsServer.Close()
globalContext.dnsServer.Close(ctx)
globalContext.dnsServer = nil
}
@@ -509,8 +508,7 @@ func closeDNSServer() {
}
if globalContext.queryLog != nil {
// TODO(s.chzhen): Pass context.
err := globalContext.queryLog.Shutdown(context.TODO())
err := globalContext.queryLog.Shutdown(ctx)
if err != nil {
log.Error("closing query log: %s", err)
}

View File

@@ -18,10 +18,12 @@ import (
"syscall"
"time"
"github.com/AdguardTeam/AdGuardHome/internal/agh"
"github.com/AdguardTeam/AdGuardHome/internal/aghalg"
"github.com/AdguardTeam/AdGuardHome/internal/aghnet"
"github.com/AdguardTeam/AdGuardHome/internal/aghos"
"github.com/AdguardTeam/AdGuardHome/internal/aghslog"
"github.com/AdguardTeam/AdGuardHome/internal/aghtls"
"github.com/AdguardTeam/AdGuardHome/internal/arpdb"
"github.com/AdguardTeam/AdGuardHome/internal/dhcpd"
"github.com/AdguardTeam/AdGuardHome/internal/dnsforward"
@@ -48,14 +50,20 @@ type homeContext struct {
// Modules
// --
clients clientsContainer // per-client-settings module
stats stats.Interface // statistics module
queryLog querylog.QueryLog // query log module
dnsServer *dnsforward.Server // DNS module
dhcpServer dhcpd.Interface // DHCP module
auth *Auth // HTTP authentication module
filters *filtering.DNSFilter // DNS filtering module
web *webAPI // Web (HTTP, HTTPS) module
clients clientsContainer // per-client-settings module
stats stats.Interface // statistics module
queryLog querylog.QueryLog // query log module
dnsServer *dnsforward.Server // DNS module
dhcpServer dhcpd.Interface // DHCP module
// auth stores web user information and handles authentication.
//
// TODO(s.chzhen): Remove once it is no longer called from different
// modules. See [onConfigModified].
auth *auth
filters *filtering.DNSFilter // DNS filtering module
web *webAPI // Web (HTTP, HTTPS) module
// tls contains the current configuration and state of TLS encryption.
//
@@ -108,13 +116,19 @@ func Main(clientBuildFS fs.FS) {
// package flag.
opts := loadCmdLineOpts()
ls := getLogSettings(opts)
// TODO(a.garipov): Use slog everywhere.
baseLogger := newSlogLogger(ls)
done := make(chan struct{})
signals := make(chan os.Signal, 1)
signal.Notify(signals, syscall.SIGINT, syscall.SIGTERM, syscall.SIGHUP, syscall.SIGQUIT)
ctx := context.Background()
sigHdlr := newSignalHandler(signals, func(ctx context.Context) {
sigHdlrLogger := baseLogger.With(slogutil.KeyPrefix, "signalhdlr")
sigHdlr := newSignalHandler(sigHdlrLogger, signals, func(ctx context.Context) {
cleanup(ctx)
cleanupAlways()
close(done)
@@ -123,24 +137,34 @@ func Main(clientBuildFS fs.FS) {
go sigHdlr.handle(ctx)
if opts.serviceControlAction != "" {
handleServiceControlAction(opts, clientBuildFS, signals, done, sigHdlr)
svcLogger := baseLogger.With(slogutil.KeyPrefix, "service")
handleServiceControlAction(
ctx,
baseLogger,
svcLogger,
opts,
clientBuildFS,
signals,
done,
sigHdlr,
)
return
}
// run the protection
run(opts, clientBuildFS, done, sigHdlr)
run(ctx, baseLogger, opts, clientBuildFS, done, sigHdlr)
}
// setupContext initializes [globalContext] fields. It also reads and upgrades
// config file if necessary.
func setupContext(opts options) (err error) {
// config file if necessary. baseLogger must not be nil.
func setupContext(ctx context.Context, baseLogger *slog.Logger, opts options) (err error) {
globalContext.firstRun = detectFirstRun()
globalContext.mux = http.NewServeMux()
if !opts.noEtcHosts {
err = setupHostsContainer()
err = setupHostsContainer(ctx, baseLogger)
if err != nil {
// Don't wrap the error, because it's informative enough as is.
return err
@@ -224,9 +248,9 @@ func configureOS(conf *configuration) (err error) {
}
// setupHostsContainer initializes the structures to keep up-to-date the hosts
// provided by the OS.
func setupHostsContainer() (err error) {
hostsWatcher, err := aghos.NewOSWritesWatcher()
// provided by the OS. baseLogger must not be nil.
func setupHostsContainer(ctx context.Context, baseLogger *slog.Logger) (err error) {
hostsWatcher, err := aghos.NewOSWritesWatcher(baseLogger.With(slogutil.KeyPrefix, "oswatcher"))
if err != nil {
log.Info("WARNING: initializing filesystem watcher: %s; not watching for changes", err)
@@ -240,7 +264,7 @@ func setupHostsContainer() (err error) {
globalContext.etcHosts, err = aghnet.NewHostsContainer(osutil.RootDirFS(), hostsWatcher, paths...)
if err != nil {
closeErr := hostsWatcher.Close()
closeErr := hostsWatcher.Shutdown(ctx)
if errors.Is(err, aghnet.ErrNoHostsPaths) {
log.Info("warning: initing hosts container: %s", err)
@@ -250,7 +274,7 @@ func setupHostsContainer() (err error) {
return errors.Join(fmt.Errorf("initializing hosts container: %w", err), closeErr)
}
return hostsWatcher.Start()
return hostsWatcher.Start(ctx)
}
// setupOpts sets up command-line options.
@@ -274,12 +298,13 @@ func initContextClients(
ctx context.Context,
logger *slog.Logger,
sigHdlr *signalHandler,
confModifier agh.ConfigModifier,
) (err error) {
//lint:ignore SA1019 Migration is not over.
config.DHCP.WorkDir = globalContext.workDir
config.DHCP.DataDir = globalContext.getDataDir()
config.DHCP.HTTPRegister = httpRegister
config.DHCP.ConfigModified = onConfigModified
config.DHCP.ConfModifier = confModifier
globalContext.dhcpServer, err = dhcpd.Create(config.DHCP)
if globalContext.dhcpServer == nil || err != nil {
@@ -304,6 +329,7 @@ func initContextClients(
arpDB,
config.Filtering,
sigHdlr,
confModifier,
)
}
@@ -354,6 +380,7 @@ func setupDNSFilteringConf(
baseLogger *slog.Logger,
conf *filtering.Config,
tlsMgr *tlsManager,
confModifier agh.ConfigModifier,
) (err error) {
const (
dnsTimeout = 3 * time.Second
@@ -375,7 +402,7 @@ func setupDNSFilteringConf(
conf.EtcHosts = nil
}
conf.ConfigModified = onConfigModified
conf.ConfModifier = confModifier
conf.HTTPRegister = httpRegister
conf.DataDir = globalContext.getDataDir()
conf.Filters = slices.Clone(config.Filters)
@@ -531,8 +558,8 @@ func isUpdateEnabled(
}
}
// initWeb initializes the web module. upd, baseLogger, and tlsMgr must not be
// nil.
// initWeb initializes the web module. upd, baseLogger, tlsMgr, and auth must
// not be nil.
func initWeb(
ctx context.Context,
opts options,
@@ -540,6 +567,8 @@ func initWeb(
upd *updater.Updater,
baseLogger *slog.Logger,
tlsMgr *tlsManager,
auth *auth,
confModifier agh.ConfigModifier,
isCustomUpdURL bool,
) (web *webAPI, err error) {
logger := baseLogger.With(slogutil.KeyPrefix, "webapi")
@@ -559,10 +588,12 @@ func initWeb(
disableUpdate := !isUpdateEnabled(ctx, baseLogger, &opts, isCustomUpdURL)
webConf := &webConfig{
updater: upd,
logger: logger,
baseLogger: baseLogger,
tlsManager: tlsMgr,
updater: upd,
logger: logger,
baseLogger: baseLogger,
confModifier: confModifier,
tlsManager: tlsMgr,
auth: auth,
clientFS: clientFS,
@@ -595,7 +626,14 @@ func fatalOnError(err error) {
// run configures and starts AdGuard Home.
//
// TODO(e.burkov): Make opts a pointer.
func run(opts options, clientBuildFS fs.FS, done chan struct{}, sigHdlr *signalHandler) {
func run(
ctx context.Context,
slogLogger *slog.Logger,
opts options,
clientBuildFS fs.FS,
done chan struct{},
sigHdlr *signalHandler,
) {
// Configure working dir.
err := initWorkingDir(opts)
fatalOnError(err)
@@ -609,10 +647,6 @@ func run(opts options, clientBuildFS fs.FS, done chan struct{}, sigHdlr *signalH
err = configureLogger(ls)
fatalOnError(err)
// TODO(a.garipov): Use slog everywhere.
slogLogger := newSlogLogger(ls)
sigHdlr.swapLogger(slogLogger)
// Print the first message after logger is configured.
log.Info("%s", version.Full())
log.Debug("current working directory is %s", globalContext.workDir)
@@ -620,38 +654,44 @@ func run(opts options, clientBuildFS fs.FS, done chan struct{}, sigHdlr *signalH
log.Info("AdGuard Home is running as a service")
}
err = setupContext(opts)
aghtls.Init(ctx, slogLogger.With(slogutil.KeyPrefix, "aghtls"))
err = setupContext(ctx, slogLogger, opts)
fatalOnError(err)
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(ctx, slogLogger)
err = initContextClients(ctx, slogLogger, sigHdlr)
confModifier := newDefaultConfigModifier(
config,
slogLogger.With(slogutil.KeyPrefix, "config_modifier"),
)
err = initContextClients(ctx, slogLogger, sigHdlr, confModifier)
fatalOnError(err)
tlsMgrLogger := slogLogger.With(slogutil.KeyPrefix, "tls_manager")
tlsMgr, err := newTLSManager(ctx, &tlsManagerConfig{
logger: tlsMgrLogger,
configModified: onConfigModified,
tlsSettings: config.TLS,
servePlainDNS: config.DNS.ServePlainDNS,
logger: tlsMgrLogger,
confModifier: confModifier,
tlsSettings: config.TLS,
servePlainDNS: config.DNS.ServePlainDNS,
})
if err != nil {
tlsMgrLogger.ErrorContext(ctx, "initializing", slogutil.KeyError, err)
onConfigModified()
confModifier.Apply(ctx)
}
globalContext.tls = tlsMgr
confModifier.setTLSManager(tlsMgr)
err = setupDNSFilteringConf(ctx, slogLogger, config.Filtering, tlsMgr)
err = setupDNSFilteringConf(ctx, slogLogger, config.Filtering, tlsMgr, confModifier)
fatalOnError(err)
err = setupOpts(opts)
@@ -671,7 +711,7 @@ func run(opts options, clientBuildFS fs.FS, done chan struct{}, sigHdlr *signalH
if !globalContext.firstRun {
// Save the updated config.
err = config.write(nil)
err = config.write(nil, nil)
fatalOnError(err)
if config.HTTPConfig.Pprof.Enabled {
@@ -683,13 +723,23 @@ func run(opts options, clientBuildFS fs.FS, done chan struct{}, sigHdlr *signalH
err = os.MkdirAll(dataDir, aghos.DefaultPermDir)
fatalOnError(errors.Annotate(err, "creating DNS data dir at %s: %w", dataDir))
GLMode = opts.glinetMode
// Init auth module.
globalContext.auth, err = initUsers()
auth, err := initUsers(ctx, slogLogger, opts.glinetMode)
fatalOnError(err)
web, err := initWeb(ctx, opts, clientBuildFS, upd, slogLogger, tlsMgr, isCustomURL)
globalContext.auth = auth
confModifier.setAuth(auth)
web, err := initWeb(
ctx,
opts,
clientBuildFS,
upd,
slogLogger,
tlsMgr,
auth,
confModifier,
isCustomURL,
)
fatalOnError(err)
globalContext.web = web
@@ -701,7 +751,7 @@ func run(opts options, clientBuildFS fs.FS, done chan struct{}, sigHdlr *signalH
fatalOnError(err)
if !globalContext.firstRun {
err = initDNS(slogLogger, tlsMgr, statsDir, querylogDir)
err = initDNS(ctx, slogLogger, tlsMgr, confModifier, statsDir, querylogDir)
fatalOnError(err)
tlsMgr.start(ctx)
@@ -709,7 +759,7 @@ func run(opts options, clientBuildFS fs.FS, done chan struct{}, sigHdlr *signalH
go func() {
startErr := startDNSServer()
if startErr != nil {
closeDNSServer()
closeDNSServer(ctx)
fatalOnError(startErr)
}
}()
@@ -805,24 +855,33 @@ func checkPermissions(
permcheck.Check(ctx, l, workDir, dataDir, statsDir, querylogDir, confPath)
}
// initUsers initializes context auth module. Clears config users field.
func initUsers() (auth *Auth, err error) {
sessFilename := filepath.Join(globalContext.getDataDir(), "sessions.db")
var rateLimiter *authRateLimiter
// initUsers initializes authentication module and clears the [config.Users]
// field.
func initUsers(
ctx context.Context,
baseLogger *slog.Logger,
isGLiNet bool,
) (auth *auth, err error) {
var rateLimiter loginRaateLimiter
if config.AuthAttempts > 0 && config.AuthBlockMin > 0 {
blockDur := time.Duration(config.AuthBlockMin) * time.Minute
rateLimiter = newAuthRateLimiter(blockDur, config.AuthAttempts)
} else {
log.Info("authratelimiter is disabled")
baseLogger.WarnContext(ctx, "authratelimiter is disabled")
rateLimiter = emptyRateLimiter{}
}
trustedProxies := netutil.SliceSubnetSet(netutil.UnembedPrefixes(config.DNS.TrustedProxies))
sessionTTL := time.Duration(config.HTTPConfig.SessionTTL).Seconds()
auth = InitAuth(sessFilename, config.Users, uint32(sessionTTL), rateLimiter, trustedProxies)
if auth == nil {
return nil, errors.Error("initializing auth module failed")
auth, err = newAuth(ctx, &authConfig{
baseLogger: baseLogger,
rateLimiter: rateLimiter,
trustedProxies: netutil.SliceSubnetSet(netutil.UnembedPrefixes(config.DNS.TrustedProxies)),
dbFilename: filepath.Join(globalContext.getDataDir(), sessionsDBName),
users: config.Users,
sessionTTL: time.Duration(config.HTTPConfig.SessionTTL),
isGLiNet: isGLiNet,
})
if err != nil {
return nil, fmt.Errorf("initializing auth module: %w", err)
}
config.Users = nil
@@ -935,12 +994,8 @@ func cleanup(ctx context.Context) {
globalContext.web.close(ctx)
globalContext.web = nil
}
if globalContext.auth != nil {
globalContext.auth.Close()
globalContext.auth = nil
}
err := stopDNSServer()
err := stopDNSServer(ctx)
if err != nil {
log.Error("stopping dns server: %s", err)
}
@@ -1098,7 +1153,7 @@ func cmdlineUpdate(
//
// TODO(e.burkov): We could probably initialize the internal resolver
// separately.
err := initDNSServer(nil, nil, nil, nil, nil, nil, tlsMgr, l)
err := initDNSServer(ctx, nil, nil, nil, nil, nil, nil, tlsMgr, l, agh.EmptyConfigModifier{})
fatalOnError(err)
l.InfoContext(ctx, "performing update via cli")
@@ -1119,7 +1174,7 @@ func cmdlineUpdate(
err = upd.Update(ctx, globalContext.firstRun)
fatalOnError(err)
err = restartService()
err = restartService(ctx, l)
if err != nil {
l.DebugContext(ctx, "restarting service", slogutil.KeyError, err)
l.InfoContext(ctx, "AdGuard Home was not installed as a service. "+

View File

@@ -64,7 +64,9 @@ func handleI18nCurrentLanguage(w http.ResponseWriter, r *http.Request) {
}
// TODO(d.kolyshev): Deprecated, remove it later.
func handleI18nChangeLanguage(w http.ResponseWriter, r *http.Request) {
func (web *webAPI) handleI18nChangeLanguage(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
if aghhttp.WriteTextPlainDeprecated(w, r) {
return
}
@@ -89,9 +91,10 @@ func handleI18nChangeLanguage(w http.ResponseWriter, r *http.Request) {
defer config.Unlock()
config.Language = lang
log.Printf("home: language is set to %s", lang)
web.logger.InfoContext(ctx, "language is updated", "lang", lang)
}()
onConfigModified()
web.confModifier.Apply(ctx)
aghhttp.OK(w)
}

View File

@@ -39,14 +39,14 @@ func TestLimitRequestBody(t *testing.T) {
want: []byte(nil),
}}
makeHandler := func(t *testing.T, err *error) http.HandlerFunc {
t.Helper()
makeHandler := func(tb testing.TB, err *error) http.HandlerFunc {
tb.Helper()
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
var b []byte
b, *err = io.ReadAll(r.Body)
_, werr := w.Write(b)
require.NoError(t, werr)
require.NoError(tb, werr)
})
}

View File

@@ -15,11 +15,11 @@ import (
// setupDNSIPs is a helper that sets up the server IP address configuration for
// tests and also tears it down in a cleanup function.
func setupDNSIPs(t testing.TB) {
t.Helper()
func setupDNSIPs(tb testing.TB) {
tb.Helper()
prevConfig := config
t.Cleanup(func() {
tb.Cleanup(func() {
config = prevConfig
})

View File

@@ -9,26 +9,34 @@ import (
"github.com/stretchr/testify/require"
)
func testParseOK(t *testing.T, ss ...string) options {
t.Helper()
// testParseOK is a helper that parses the command-line options and returns the
// parsed options.
func testParseOK(tb testing.TB, ss ...string) (o options) {
tb.Helper()
o, _, err := parseCmdOpts("", ss)
require.NoError(t, err)
require.NoError(tb, err)
return o
}
func testParseErr(t *testing.T, descr string, ss ...string) {
t.Helper()
// testParseErr is a helper that asserts that parsing the command-line options
// fails with error.
//
// TODO(a.garipov): Search descr within an error.
func testParseErr(tb testing.TB, descr string, ss ...string) {
tb.Helper()
_, _, err := parseCmdOpts("", ss)
require.Error(t, err)
require.Errorf(tb, err, "should have got error: %s", descr)
}
func testParseParamMissing(t *testing.T, param string) {
t.Helper()
// testParseParamMissing is a helper that asserts that parsing the command-line
// options fails with error due to missing parameter.
func testParseParamMissing(tb testing.TB, param string) {
tb.Helper()
testParseErr(t, fmt.Sprintf("%s parameter missing", param), param)
testParseErr(tb, fmt.Sprintf("%s parameter missing", param), param)
}
func TestParseVerbose(t *testing.T) {

View File

@@ -6,7 +6,6 @@ import (
"net/http"
"github.com/AdguardTeam/AdGuardHome/internal/aghhttp"
"github.com/AdguardTeam/golibs/log"
)
// Theme is an enum of all allowed UI themes.
@@ -46,8 +45,18 @@ type profileJSON struct {
}
// handleGetProfile is the handler for GET /control/profile endpoint.
func handleGetProfile(w http.ResponseWriter, r *http.Request) {
u := globalContext.auth.getCurrentUser(r)
func (web *webAPI) handleGetProfile(w http.ResponseWriter, r *http.Request) {
var name string
if !web.auth.isGLiNet {
u, ok := webUserFromContext(r.Context())
if !ok {
w.WriteHeader(http.StatusUnauthorized)
return
}
name = string(u.Login)
}
var resp profileJSON
func() {
@@ -55,7 +64,7 @@ func handleGetProfile(w http.ResponseWriter, r *http.Request) {
defer config.RUnlock()
resp = profileJSON{
Name: u.Name,
Name: name,
Language: config.Language,
Theme: config.Theme,
}
@@ -65,7 +74,9 @@ func handleGetProfile(w http.ResponseWriter, r *http.Request) {
}
// handlePutProfile is the handler for PUT /control/profile/update endpoint.
func handlePutProfile(w http.ResponseWriter, r *http.Request) {
func (web *webAPI) handlePutProfile(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
if aghhttp.WriteTextPlainDeprecated(w, r) {
return
}
@@ -93,10 +104,10 @@ func handlePutProfile(w http.ResponseWriter, r *http.Request) {
config.Language = lang
config.Theme = theme
log.Printf("home: language is set to %s", lang)
log.Printf("home: theme is set to %s", theme)
web.logger.InfoContext(ctx, "profile updated", "lang", lang, "theme", theme)
}()
onConfigModified()
web.confModifier.Apply(ctx)
aghhttp.OK(w)
}

View File

@@ -1,8 +1,10 @@
package home
import (
"context"
"fmt"
"io/fs"
"log/slog"
"os"
"runtime"
"strconv"
@@ -13,8 +15,9 @@ import (
"github.com/AdguardTeam/AdGuardHome/internal/aghos"
"github.com/AdguardTeam/AdGuardHome/internal/version"
"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/osutil"
"github.com/kardianos/service"
)
@@ -32,10 +35,14 @@ const (
// program represents the program that will be launched by as a service or a
// daemon.
type program struct {
// TODO(s.chzhen): Remove this.
ctx context.Context
clientBuildFS fs.FS
signals chan os.Signal
done chan struct{}
opts options
baseLogger *slog.Logger
logger *slog.Logger
sigHdlr *signalHandler
}
@@ -48,14 +55,14 @@ func (p *program) Start(_ service.Service) (err error) {
args := p.opts
args.runningAsService = true
go run(args, p.clientBuildFS, p.done, p.sigHdlr)
go run(p.ctx, p.baseLogger, args, p.clientBuildFS, p.done, p.sigHdlr)
return nil
}
// Stop implements service.Interface interface for *program.
func (p *program) Stop(_ service.Service) (err error) {
log.Info("service: stopping: waiting for cleanup")
p.logger.InfoContext(p.ctx, "stopping: waiting for cleanup")
aghos.SendShutdownSignal(p.signals)
@@ -84,14 +91,14 @@ func svcStatus(s service.Service) (status service.Status, err error) {
return status, err
}
// svcAction performs the action on the service.
// svcAction performs the action on the service. l must not be nil.
//
// On OpenWrt, the service utility may not exist. We use our service script
// directly in this case.
func svcAction(s service.Service, action string) (err error) {
func svcAction(ctx context.Context, l *slog.Logger, s service.Service, action string) (err error) {
if action == "start" {
if err = aghos.PreCheckActionStart(); err != nil {
log.Error("starting service: %s", err)
l.ErrorContext(ctx, "starting service", slogutil.KeyError, err)
}
}
@@ -105,10 +112,10 @@ func svcAction(s service.Service, action string) (err error) {
}
// Send SIGHUP to a process with PID taken from our .pid file. If it doesn't
// exist, find our PID using 'ps' command.
func sendSigReload() {
// exist, find our PID using 'ps' command. baseLogger and l must not be nil.
func sendSigReload(ctx context.Context, baseLogger, l *slog.Logger) {
if runtime.GOOS == "windows" {
log.Error("service: not implemented on windows")
l.ErrorContext(ctx, "not implemented on windows")
return
}
@@ -117,25 +124,26 @@ func sendSigReload() {
var pid int
data, err := os.ReadFile(pidFile)
if errors.Is(err, os.ErrNotExist) {
if pid, err = aghos.PIDByCommand(serviceName, os.Getpid()); err != nil {
log.Error("service: finding AdGuardHome process: %s", err)
aghosLogger := baseLogger.With(slogutil.KeyPrefix, "aghos")
if pid, err = aghos.PIDByCommand(ctx, aghosLogger, serviceName, os.Getpid()); err != nil {
l.ErrorContext(ctx, "finding adguardhome process", slogutil.KeyError, err)
return
}
} else if err != nil {
log.Error("service: reading pid file %s: %s", pidFile, err)
l.ErrorContext(ctx, "reading", "pid_file", pidFile, slogutil.KeyError, err)
return
} else {
parts := strings.SplitN(string(data), "\n", 2)
if len(parts) == 0 {
log.Error("service: parsing pid file %s: bad value", pidFile)
l.ErrorContext(ctx, "splitting", "pid_file", pidFile, slogutil.KeyError, "bad value")
return
}
if pid, err = strconv.Atoi(strings.TrimSpace(parts[0])); err != nil {
log.Error("service: parsing pid from file %s: %s", pidFile, err)
l.ErrorContext(ctx, "parsing", "pid_file", pidFile, slogutil.KeyError, err)
return
}
@@ -143,23 +151,23 @@ func sendSigReload() {
var proc *os.Process
if proc, err = os.FindProcess(pid); err != nil {
log.Error("service: finding process for pid %d: %s", pid, err)
l.ErrorContext(ctx, "finding process for", "pid", pid, slogutil.KeyError, err)
return
}
if err = proc.Signal(syscall.SIGHUP); err != nil {
log.Error("service: sending signal HUP to pid %d: %s", pid, err)
l.ErrorContext(ctx, "sending sighup to", "pid", pid, slogutil.KeyError, err)
return
}
log.Debug("service: sent signal to pid %d", pid)
l.DebugContext(ctx, "sent sighup to", "pid", pid)
}
// restartService restarts the service. It returns error if the service is not
// running.
func restartService() (err error) {
// running. l must not be nil.
func restartService(ctx context.Context, l *slog.Logger) (err error) {
// Call chooseSystem explicitly to introduce OpenBSD support for service
// package. It's a noop for other GOOS values.
chooseSystem()
@@ -182,7 +190,7 @@ func restartService() (err error) {
return fmt.Errorf("initializing service: %w", err)
}
if err = svcAction(s, "restart"); err != nil {
if err = svcAction(ctx, l, s, "restart"); err != nil {
return fmt.Errorf("restarting service: %w", err)
}
@@ -201,6 +209,9 @@ func restartService() (err error) {
// it is specified when we register a service, and it indicates to the app
// that it is being run as a service/daemon.
func handleServiceControlAction(
ctx context.Context,
baseLogger *slog.Logger,
l *slog.Logger,
opts options,
clientBuildFS fs.FS,
signals chan os.Signal,
@@ -212,25 +223,26 @@ func handleServiceControlAction(
chooseSystem()
action := opts.serviceControlAction
log.Info("%s", version.Full())
log.Info("service: control action: %s", action)
l.InfoContext(ctx, version.Full())
l.InfoContext(ctx, "control", "action", action)
if action == "reload" {
sendSigReload()
sendSigReload(ctx, baseLogger, l)
return
}
pwd, err := os.Getwd()
if err != nil {
log.Fatalf("service: getting current directory: %s", err)
l.ErrorContext(ctx, "getting current directory", slogutil.KeyError, err)
os.Exit(osutil.ExitCodeFailure)
}
runOpts := opts
runOpts.serviceControlAction = "run"
args := optsToArgs(runOpts)
log.Debug("service: using args %q", args)
l.DebugContext(ctx, "using", "args", args)
svcConfig := &service.Config{
Name: serviceName,
@@ -242,33 +254,45 @@ func handleServiceControlAction(
configureService(svcConfig)
s, err := service.New(&program{
ctx: ctx,
clientBuildFS: clientBuildFS,
signals: signals,
done: done,
opts: runOpts,
baseLogger: l,
logger: l.With(slogutil.KeyPrefix, "service"),
sigHdlr: sigHdlr,
}, svcConfig)
if err != nil {
log.Fatalf("service: initializing service: %s", err)
l.ErrorContext(ctx, "initializing service", slogutil.KeyError, err)
os.Exit(osutil.ExitCodeFailure)
}
err = handleServiceCommand(s, action, opts)
err = handleServiceCommand(ctx, l, s, action, opts)
if err != nil {
log.Fatalf("service: %s", err)
l.ErrorContext(ctx, "handling command", slogutil.KeyError, err)
os.Exit(osutil.ExitCodeFailure)
}
log.Printf(
"service: action %s has been done successfully on %s",
action,
service.ChosenSystem(),
l.InfoContext(
ctx,
"action has been done successfully",
"action", action,
"system", service.ChosenSystem(),
)
}
// handleServiceCommand handles service command.
func handleServiceCommand(s service.Service, action string, opts options) (err error) {
func handleServiceCommand(
ctx context.Context,
l *slog.Logger,
s service.Service,
action string,
opts options,
) (err error) {
switch action {
case "status":
handleServiceStatusCommand(s)
handleServiceStatusCommand(ctx, l, s)
case "run":
if err = s.Run(); err != nil {
return fmt.Errorf("failed to run service: %w", err)
@@ -280,11 +304,11 @@ func handleServiceCommand(s service.Service, action string, opts options) (err e
initConfigFilename(opts)
handleServiceInstallCommand(s)
handleServiceInstallCommand(ctx, l, s)
case "uninstall":
handleServiceUninstallCommand(s)
handleServiceUninstallCommand(ctx, l, s)
default:
if err = svcAction(s, action); err != nil {
if err = svcAction(ctx, l, s, action); err != nil {
return fmt.Errorf("executing action %q: %w", action, err)
}
}
@@ -297,29 +321,35 @@ func handleServiceCommand(s service.Service, action string, opts options) (err e
const statusRestartOnFail = service.StatusStopped + 1
// handleServiceStatusCommand handles service "status" command.
func handleServiceStatusCommand(s service.Service) {
func handleServiceStatusCommand(
ctx context.Context,
l *slog.Logger,
s service.Service,
) {
status, errSt := svcStatus(s)
if errSt != nil {
log.Fatalf("service: failed to get service status: %s", errSt)
l.ErrorContext(ctx, "failed to get service status", slogutil.KeyError, errSt)
os.Exit(osutil.ExitCodeFailure)
}
switch status {
case service.StatusUnknown:
log.Printf("service: status is unknown")
l.InfoContext(ctx, "status is unknown")
case service.StatusStopped:
log.Printf("service: stopped")
l.InfoContext(ctx, "stopped")
case service.StatusRunning:
log.Printf("service: running")
l.InfoContext(ctx, "running")
case statusRestartOnFail:
log.Printf("service: restarting after failed start")
l.InfoContext(ctx, "restarting after failed start")
}
}
// handleServiceInstallCommand handles service "install" command.
func handleServiceInstallCommand(s service.Service) {
err := svcAction(s, "install")
func handleServiceInstallCommand(ctx context.Context, l *slog.Logger, s service.Service) {
err := svcAction(ctx, l, s, "install")
if err != nil {
log.Fatalf("service: executing action %q: %s", "install", err)
l.ErrorContext(ctx, "executing install", slogutil.KeyError, err)
os.Exit(osutil.ExitCodeFailure)
}
if aghos.IsOpenWrt() {
@@ -328,56 +358,60 @@ func handleServiceInstallCommand(s service.Service) {
// startup.
_, err = runInitdCommand("enable")
if err != nil {
log.Fatalf("service: running init enable: %s", err)
l.ErrorContext(ctx, "running init enable", slogutil.KeyError, err)
os.Exit(osutil.ExitCodeFailure)
}
}
// Start automatically after install.
err = svcAction(s, "start")
err = svcAction(ctx, l, s, "start")
if err != nil {
log.Fatalf("service: starting: %s", err)
l.ErrorContext(ctx, "starting", slogutil.KeyError, err)
os.Exit(osutil.ExitCodeFailure)
}
log.Printf("service: started")
l.InfoContext(ctx, "started")
if detectFirstRun() {
log.Printf(`Almost ready!
AdGuard Home is successfully installed and will automatically start on boot.
There are a few more things that must be configured before you can use it.
Click on the link below and follow the Installation Wizard steps to finish setup.
AdGuard Home is now available at the following addresses:`)
slogutil.PrintLines(ctx, l, slog.LevelInfo, "", "Almost ready!\n"+
"AdGuard Home is successfully installed and will automatically start on boot.\n"+
"There are a few more things that must be configured before you can use it.\n"+
"Click on the link below and follow the Installation Wizard steps to finish setup.\n"+
"AdGuard Home is now available at the following addresses:")
printHTTPAddresses(urlutil.SchemeHTTP, nil)
}
}
// handleServiceUninstallCommand handles service "uninstall" command.
func handleServiceUninstallCommand(s service.Service) {
func handleServiceUninstallCommand(ctx context.Context, l *slog.Logger, s service.Service) {
if aghos.IsOpenWrt() {
// On OpenWrt it is important to run disable command first
// as it will remove the symlink
_, err := runInitdCommand("disable")
if err != nil {
log.Fatalf("service: running init disable: %s", err)
l.ErrorContext(ctx, "running init disable", slogutil.KeyError, err)
os.Exit(osutil.ExitCodeFailure)
}
}
if err := svcAction(s, "stop"); err != nil {
log.Debug("service: executing action %q: %s", "stop", err)
if err := svcAction(ctx, l, s, "stop"); err != nil {
l.DebugContext(ctx, "executing action stop", slogutil.KeyError, err)
}
if err := svcAction(s, "uninstall"); err != nil {
log.Fatalf("service: executing action %q: %s", "uninstall", err)
if err := svcAction(ctx, l, s, "uninstall"); err != nil {
l.ErrorContext(ctx, "executing action uninstall", slogutil.KeyError, err)
os.Exit(osutil.ExitCodeFailure)
}
if runtime.GOOS == "darwin" {
// Remove log files on cleanup and log errors.
err := os.Remove(launchdStdoutPath)
if err != nil && !errors.Is(err, os.ErrNotExist) {
log.Info("service: warning: removing stdout file: %s", err)
l.WarnContext(ctx, "removing stdout file", slogutil.KeyError, err)
}
err = os.Remove(launchdStderrPath)
if err != nil && !errors.Is(err, os.ErrNotExist) {
log.Info("service: warning: removing stderr file: %s", err)
l.WarnContext(ctx, "removing stderr file", slogutil.KeyError, err)
}
}
}

View File

@@ -5,7 +5,6 @@ import (
"log/slog"
"os"
"sync"
"sync/atomic"
"syscall"
"github.com/AdguardTeam/AdGuardHome/internal/client"
@@ -16,10 +15,8 @@ import (
// signalHandler processes incoming signals. It reloads configurations of
// stored entities on SIGHUP and performs cleanup on all other signals.
type signalHandler struct {
// logger is used to log the operation of the signal handler. Initially,
// [slog.Default] is used, but it should be swapped later using
// [signalHandler.swapLogger].
logger *atomic.Pointer[slog.Logger]
// logger is used to log the operation of the signal handler.
logger *slog.Logger
// mu protects clientStorage and tlsManager.
mu *sync.Mutex
@@ -41,24 +38,16 @@ type signalHandler struct {
// newSignalHandler returns a new properly initialized *signalHandler.
func newSignalHandler(
l *slog.Logger,
signals <-chan os.Signal,
cleanup func(ctx context.Context),
) (h *signalHandler) {
h = &signalHandler{
logger: &atomic.Pointer[slog.Logger]{},
return &signalHandler{
logger: l,
mu: &sync.Mutex{},
signals: signals,
cleanup: cleanup,
}
h.logger.Store(slog.Default())
return h
}
// swapLogger replaces the stored logger with the given logger.
func (h *signalHandler) swapLogger(logger *slog.Logger) {
h.logger.Swap(logger)
}
// addClientStorage stores the client storage.
@@ -89,14 +78,14 @@ func (h *signalHandler) handle(ctx context.Context) {
return
}
slogutil.PrintRecovered(ctx, h.logger.Load(), v)
slogutil.PrintRecovered(ctx, h.logger, v)
os.Exit(osutil.ExitCodeFailure)
}()
for {
sig := <-h.signals
h.logger.Load().InfoContext(ctx, "received signal", "signal", sig)
h.logger.InfoContext(ctx, "received signal", "signal", sig)
switch sig {
case syscall.SIGHUP:
h.reloadConfig(ctx)

View File

@@ -20,6 +20,7 @@ import (
"sync"
"time"
"github.com/AdguardTeam/AdGuardHome/internal/agh"
"github.com/AdguardTeam/AdGuardHome/internal/aghalg"
"github.com/AdguardTeam/AdGuardHome/internal/aghhttp"
"github.com/AdguardTeam/AdGuardHome/internal/aghnet"
@@ -56,9 +57,8 @@ type tlsManager struct {
// conf contains the TLS configuration settings. It must not be nil.
conf *tlsConfigSettings
// configModified is called when the TLS configuration is changed via an
// HTTP request.
configModified func()
// confModifier is used to update the global configuration.
confModifier agh.ConfigModifier
// customCipherIDs are the ID of the cipher suites that AdGuard Home must use.
customCipherIDs []uint16
@@ -73,9 +73,9 @@ type tlsManagerConfig struct {
// be nil.
logger *slog.Logger
// configModified is called when the TLS configuration is changed via an
// HTTP request. It must not be nil.
configModified func()
// confModifier is used to update the global configuration. It must not be
// nil.
confModifier agh.ConfigModifier
// tlsSettings contains the TLS configuration settings.
tlsSettings tlsConfigSettings
@@ -91,15 +91,15 @@ type tlsManagerConfig struct {
// [tlsManager.setWebAPI].
func newTLSManager(ctx context.Context, conf *tlsManagerConfig) (m *tlsManager, err error) {
m = &tlsManager{
logger: conf.logger,
mu: &sync.Mutex{},
configModified: conf.configModified,
status: &tlsConfigStatus{},
conf: &conf.tlsSettings,
servePlainDNS: conf.servePlainDNS,
logger: conf.logger,
mu: &sync.Mutex{},
confModifier: conf.confModifier,
status: &tlsConfigStatus{},
conf: &conf.tlsSettings,
servePlainDNS: conf.servePlainDNS,
}
m.rootCerts = aghtls.SystemRootCAs()
m.rootCerts = aghtls.SystemRootCAs(ctx, conf.logger)
if len(conf.tlsSettings.OverrideTLSCiphers) > 0 {
m.customCipherIDs, err = aghtls.ParseCiphers(config.TLS.OverrideTLSCiphers)
@@ -232,7 +232,7 @@ func (m *tlsManager) reload(ctx context.Context) {
m.certLastMod = fi.ModTime().UTC()
err = m.reconfigureDNSServer()
err = m.reconfigureDNSServer(ctx)
if err != nil {
m.logger.ErrorContext(ctx, "reconfiguring dns server", slogutil.KeyError, err)
}
@@ -245,7 +245,7 @@ func (m *tlsManager) reload(ctx context.Context) {
// reconfigureDNSServer updates the DNS server configuration using the stored
// TLS settings. m.mu is expected to be locked.
func (m *tlsManager) reconfigureDNSServer() (err error) {
func (m *tlsManager) reconfigureDNSServer(ctx context.Context) (err error) {
newConf, err := newServerConfig(
&config.DNS,
config.Clients.Sources,
@@ -253,12 +253,13 @@ func (m *tlsManager) reconfigureDNSServer() (err error) {
m,
httpRegister,
globalContext.clients.storage,
m.confModifier,
)
if err != nil {
return fmt.Errorf("generating forwarding dns server config: %w", err)
}
err = globalContext.dnsServer.Reconfigure(newConf)
err = globalContext.dnsServer.Reconfigure(ctx, newConf)
if err != nil {
return fmt.Errorf("starting forwarding dns server: %w", err)
}
@@ -515,7 +516,7 @@ func (m *tlsManager) handleTLSConfigure(w http.ResponseWriter, r *http.Request)
var restartHTTPS bool
defer func() {
if restartHTTPS {
m.configModified()
m.confModifier.Apply(ctx)
}
}()
@@ -557,7 +558,7 @@ func (m *tlsManager) handleTLSConfigure(w http.ResponseWriter, r *http.Request)
}()
}
err = m.reconfigureDNSServer()
err = m.reconfigureDNSServer(ctx)
if err != nil {
m.logger.ErrorContext(ctx, "reconfiguring dns server", slogutil.KeyError, err)

View File

@@ -20,6 +20,7 @@ import (
"testing"
"time"
"github.com/AdguardTeam/AdGuardHome/internal/agh"
"github.com/AdguardTeam/AdGuardHome/internal/aghalg"
"github.com/AdguardTeam/AdGuardHome/internal/client"
"github.com/AdguardTeam/AdGuardHome/internal/dnsforward"
@@ -66,9 +67,9 @@ func TestValidateCertificates(t *testing.T) {
ctx := testutil.ContextWithTimeout(t, testTimeout)
m, err := newTLSManager(ctx, &tlsManagerConfig{
logger: testLogger,
configModified: func() {},
servePlainDNS: false,
logger: testLogger,
confModifier: agh.EmptyConfigModifier{},
servePlainDNS: false,
})
require.NoError(t, err)
@@ -111,7 +112,6 @@ func TestValidateCertificates(t *testing.T) {
// restores them once the test is complete.
//
// The global variables are:
// - [GLMode]
// - [config]
// - [glFilePrefix]
// - [globalContext.auth]
@@ -126,10 +126,8 @@ func TestValidateCertificates(t *testing.T) {
func storeGlobals(tb testing.TB) {
tb.Helper()
prevGLMode := GLMode
prevConfig := config
prefGLFilePrefix := glFilePrefix
auth := globalContext.auth
storage := globalContext.clients.storage
dnsServer := globalContext.dnsServer
firstRun := globalContext.firstRun
@@ -137,10 +135,8 @@ func storeGlobals(tb testing.TB) {
web := globalContext.web
tb.Cleanup(func() {
GLMode = prevGLMode
config = prevConfig
glFilePrefix = prefGLFilePrefix
globalContext.auth = auth
globalContext.clients.storage = storage
globalContext.dnsServer = dnsServer
globalContext.firstRun = firstRun
@@ -251,8 +247,8 @@ func TestTLSManager_Reload(t *testing.T) {
writeCertAndKey(t, certDER, certPath, key, keyPath)
m, err := newTLSManager(ctx, &tlsManagerConfig{
logger: testLogger,
configModified: func() {},
logger: testLogger,
confModifier: agh.EmptyConfigModifier{},
tlsSettings: tlsConfigSettings{
Enabled: true,
CertificatePath: certPath,
@@ -262,7 +258,7 @@ func TestTLSManager_Reload(t *testing.T) {
})
require.NoError(t, err)
web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, false)
web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, nil, agh.EmptyConfigModifier{}, false)
require.NoError(t, err)
m.setWebAPI(web)
@@ -277,7 +273,9 @@ func TestTLSManager_Reload(t *testing.T) {
// The [tlsManager.reload] method will start the DNS server and it should be
// stopped after the test ends.
testutil.CleanupAndRequireSuccess(t, globalContext.dnsServer.Stop)
testutil.CleanupAndRequireSuccess(t, func() (err error) {
return globalContext.dnsServer.Stop(testutil.ContextWithTimeout(t, testTimeout))
})
conf = m.config()
assertCertSerialNumber(t, conf, snAfter)
@@ -290,8 +288,8 @@ func TestTLSManager_HandleTLSStatus(t *testing.T) {
)
m, err := newTLSManager(ctx, &tlsManagerConfig{
logger: testLogger,
configModified: func() {},
logger: testLogger,
confModifier: agh.EmptyConfigModifier{},
tlsSettings: tlsConfigSettings{
Enabled: true,
CertificateChain: string(testCertChainData),
@@ -326,13 +324,13 @@ func TestValidateTLSSettings(t *testing.T) {
)
m, err := newTLSManager(ctx, &tlsManagerConfig{
logger: testLogger,
configModified: func() {},
servePlainDNS: false,
logger: testLogger,
confModifier: agh.EmptyConfigModifier{},
servePlainDNS: false,
})
require.NoError(t, err)
web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, false)
web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, nil, agh.EmptyConfigModifier{}, false)
require.NoError(t, err)
m.setWebAPI(web)
@@ -425,8 +423,8 @@ func TestTLSManager_HandleTLSValidate(t *testing.T) {
)
m, err := newTLSManager(ctx, &tlsManagerConfig{
logger: testLogger,
configModified: func() {},
logger: testLogger,
confModifier: agh.EmptyConfigModifier{},
tlsSettings: tlsConfigSettings{
Enabled: true,
CertificateChain: string(testCertChainData),
@@ -436,7 +434,7 @@ func TestTLSManager_HandleTLSValidate(t *testing.T) {
})
require.NoError(t, err)
web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, false)
web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, nil, agh.EmptyConfigModifier{}, false)
require.NoError(t, err)
m.setWebAPI(web)
@@ -481,15 +479,17 @@ func TestTLSManager_HandleTLSConfigure(t *testing.T) {
})
require.NoError(t, err)
err = globalContext.dnsServer.Prepare(&dnsforward.ServerConfig{
TLSConf: &dnsforward.TLSConfig{},
Config: dnsforward.Config{
UpstreamMode: dnsforward.UpstreamModeLoadBalance,
EDNSClientSubnet: &dnsforward.EDNSClientSubnet{Enabled: false},
ClientsContainer: dnsforward.EmptyClientsContainer{},
},
ServePlainDNS: true,
})
err = globalContext.dnsServer.Prepare(
testutil.ContextWithTimeout(t, testTimeout),
&dnsforward.ServerConfig{
TLSConf: &dnsforward.TLSConfig{},
Config: dnsforward.Config{
UpstreamMode: dnsforward.UpstreamModeLoadBalance,
EDNSClientSubnet: &dnsforward.EDNSClientSubnet{Enabled: false},
ClientsContainer: dnsforward.EmptyClientsContainer{},
},
ServePlainDNS: true,
})
require.NoError(t, err)
globalContext.clients.storage, err = client.NewStorage(ctx, &client.StorageConfig{
@@ -516,8 +516,8 @@ func TestTLSManager_HandleTLSConfigure(t *testing.T) {
// Initialize the TLS manager and assert its configuration.
m, err := newTLSManager(ctx, &tlsManagerConfig{
logger: testLogger,
configModified: func() {},
logger: testLogger,
confModifier: agh.EmptyConfigModifier{},
tlsSettings: tlsConfigSettings{
Enabled: true,
CertificatePath: certPath,
@@ -527,7 +527,7 @@ func TestTLSManager_HandleTLSConfigure(t *testing.T) {
})
require.NoError(t, err)
web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, false)
web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, nil, agh.EmptyConfigModifier{}, false)
require.NoError(t, err)
m.setWebAPI(web)
@@ -556,7 +556,9 @@ func TestTLSManager_HandleTLSConfigure(t *testing.T) {
// The [tlsManager.handleTLSConfigure] method will start the DNS server and
// it should be stopped after the test ends.
testutil.CleanupAndRequireSuccess(t, globalContext.dnsServer.Stop)
testutil.CleanupAndRequireSuccess(t, func() (err error) {
return globalContext.dnsServer.Stop(testutil.ContextWithTimeout(t, testTimeout))
})
res := &tlsConfig{
tlsConfigStatus: &tlsConfigStatus{},

Some files were not shown because too many files have changed in this diff Show More