mirror of
https://git.vectorsigma.ru/public/AdGuardHome.git
synced 2026-08-03 22:09:39 +00:00
all: sync with master
This commit is contained in:
22
internal/agh/agh.go
Normal file
22
internal/agh/agh.go
Normal 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) {}
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
},
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -1,11 +0,0 @@
|
||||
package aghos_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/AdguardTeam/golibs/testutil"
|
||||
)
|
||||
|
||||
func TestMain(m *testing.M) {
|
||||
testutil.DiscardLogOutput(m)
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
)
|
||||
|
||||
@@ -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 },
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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]
|
||||
},
|
||||
|
||||
@@ -2,4 +2,4 @@
|
||||
package configmigrate
|
||||
|
||||
// LastSchemaVersion is the most recent schema version.
|
||||
const LastSchemaVersion uint = 29
|
||||
const LastSchemaVersion uint = 30
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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] {
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
119
internal/configmigrate/testdata/TestMigrateConfig_Migrate/v30/input.yml
vendored
Normal file
119
internal/configmigrate/testdata/TestMigrateConfig_Migrate/v30/input.yml
vendored
Normal 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
|
||||
120
internal/configmigrate/testdata/TestMigrateConfig_Migrate/v30/output.yml
vendored
Normal file
120
internal/configmigrate/testdata/TestMigrateConfig_Migrate/v30/output.yml
vendored
Normal 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
|
||||
33
internal/configmigrate/v30.go
Normal file
33
internal/configmigrate/v30.go
Normal 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
|
||||
}
|
||||
@@ -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:"-"`
|
||||
|
||||
@@ -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,
|
||||
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -201,6 +201,7 @@ func TestServer_clientIDFromDNSContext(t *testing.T) {
|
||||
srv := &Server{
|
||||
conf: ServerConfig{TLSConf: tlsConf},
|
||||
baseLogger: testLogger,
|
||||
logger: testLogger,
|
||||
}
|
||||
|
||||
var (
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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"))
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
})
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 filtering‑rule
|
||||
// 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)
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
},
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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{},
|
||||
},
|
||||
}}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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. "+
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
})
|
||||
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user