Pull request 2442: AGDNS-3061-config-modifier

Merge in DNS/adguard-home from AGDNS-3061-config-modifier to master

Squashed commit of the following:

commit a0068547bd0209d12e8dbf98ddd5e4ed2545cdd0
Merge: 97b798af6 451255675
Author: Stanislav Chzhen <s.chzhen@adguard.com>
Date:   Tue Aug 5 17:39:31 2025 +0300

    Merge branch 'master' into AGDNS-3061-config-modifier

commit 97b798af6a50ee27ee5ed2bcf1c4c3670f5afc5d
Merge: 96d21efc9 b8043e4f0
Author: Stanislav Chzhen <s.chzhen@adguard.com>
Date:   Wed Jul 30 20:27:41 2025 +0300

    Merge branch 'master' into AGDNS-3061-config-modifier

commit 96d21efc984073adef5026de8a03b5bf94542648
Author: Stanislav Chzhen <s.chzhen@adguard.com>
Date:   Wed Jul 30 20:18:43 2025 +0300

    all: imp code

commit 67c5608b4be3bd712a0ab5980f25ecea1ed21d65
Author: Stanislav Chzhen <s.chzhen@adguard.com>
Date:   Tue Jul 29 20:31:19 2025 +0300

    all: imp code

commit 52f45a9f70f57d8e7f7fc0e9e8291ff0dde74343
Author: Stanislav Chzhen <s.chzhen@adguard.com>
Date:   Mon Jul 28 20:32:13 2025 +0300

    all: use config modifier

commit d389ffd286460d8ff1964bd2ee8dabdafb832b9b
Author: Stanislav Chzhen <s.chzhen@adguard.com>
Date:   Fri Jul 25 14:20:31 2025 +0300

    bamboo-specs: fix ci

commit 3f303ac9131a07e4af5783696b237436fd93d65c
Author: Stanislav Chzhen <s.chzhen@adguard.com>
Date:   Fri Jul 25 14:18:42 2025 +0300

    home: config modifier
This commit is contained in:
Stanislav Chzhen
2025-08-05 18:16:39 +03:00
parent 451255675e
commit 86de4e75f0
55 changed files with 813 additions and 491 deletions

View File

@@ -19,6 +19,7 @@ import (
"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"
@@ -327,7 +328,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, nil, false)
web, err := initWeb(
ctx,
options{},
nil,
nil,
testLogger,
nil,
nil,
agh.EmptyConfigModifier{},
false,
)
require.NoError(t, err)
globalContext.web = web
@@ -477,13 +488,23 @@ func TestAuth_ServeHTTP_auth(t *testing.T) {
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, auth, false)
web, err := initWeb(
ctx,
options{},
nil,
nil,
testLogger,
tlsMgr,
auth,
agh.EmptyConfigModifier{},
false,
)
require.NoError(t, err)
globalContext.web = web
@@ -625,7 +646,16 @@ func TestAuth_ServeHTTP_logout(t *testing.T) {
globalContext.mux = http.NewServeMux()
ctx := testutil.ContextWithTimeout(t, testTimeout)
web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, auth, false)
web, err := initWeb(ctx,
options{},
nil,
nil,
testLogger,
nil,
auth,
agh.EmptyConfigModifier{},
false,
)
require.NoError(t, err)
globalContext.web = web

View File

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

View File

@@ -3,6 +3,7 @@ package home
import (
"testing"
"github.com/AdguardTeam/AdGuardHome/internal/agh"
"github.com/AdguardTeam/AdGuardHome/internal/client"
"github.com/AdguardTeam/AdGuardHome/internal/filtering"
"github.com/AdguardTeam/golibs/testutil"
@@ -14,9 +15,7 @@ import (
func newClientsContainer(t *testing.T) (c *clientsContainer) {
t.Helper()
c = &clientsContainer{
testing: true,
}
c = &clientsContainer{}
ctx := testutil.ContextWithTimeout(t, testTimeout)
err := c.Init(
@@ -30,6 +29,7 @@ func newClientsContainer(t *testing.T) (c *clientsContainer) {
Logger: testLogger,
},
newSignalHandler(testLogger, nil, nil),
agh.EmptyConfigModifier{},
)
require.NoError(t, err)

View File

@@ -326,6 +326,8 @@ func clientToJSON(c *client.Persistent) (cj *clientJSON) {
// handleAddClient is the handler for POST /control/clients/add HTTP API.
func (clients *clientsContainer) handleAddClient(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
cj := clientJSON{}
err := json.NewDecoder(r.Body).Decode(&cj)
if err != nil {
@@ -334,27 +336,27 @@ func (clients *clientsContainer) handleAddClient(w http.ResponseWriter, r *http.
return
}
c, err := clients.jsonToClient(r.Context(), cj, nil)
c, err := clients.jsonToClient(ctx, cj, nil)
if err != nil {
aghhttp.Error(r, w, http.StatusBadRequest, "%s", err)
return
}
err = clients.storage.Add(r.Context(), c)
err = clients.storage.Add(ctx, c)
if err != nil {
aghhttp.Error(r, w, http.StatusBadRequest, "%s", err)
return
}
if !clients.testing {
onConfigModified()
}
clients.confModifier.Apply(ctx)
}
// handleDelClient is the handler for POST /control/clients/delete HTTP API.
func (clients *clientsContainer) handleDelClient(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
cj := clientJSON{}
err := json.NewDecoder(r.Body).Decode(&cj)
if err != nil {
@@ -369,15 +371,13 @@ func (clients *clientsContainer) handleDelClient(w http.ResponseWriter, r *http.
return
}
if !clients.storage.RemoveByName(r.Context(), cj.Name) {
if !clients.storage.RemoveByName(ctx, cj.Name) {
aghhttp.Error(r, w, http.StatusBadRequest, "Client not found")
return
}
if !clients.testing {
onConfigModified()
}
clients.confModifier.Apply(ctx)
}
// updateJSON contains the name and data of the updated persistent client.
@@ -390,6 +390,8 @@ type updateJSON struct {
//
// TODO(s.chzhen): Accept updated parameters instead of whole structure.
func (clients *clientsContainer) handleUpdateClient(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
dj := updateJSON{}
err := json.NewDecoder(r.Body).Decode(&dj)
if err != nil {
@@ -404,23 +406,21 @@ func (clients *clientsContainer) handleUpdateClient(w http.ResponseWriter, r *ht
return
}
c, err := clients.jsonToClient(r.Context(), dj.Data, nil)
c, err := clients.jsonToClient(ctx, dj.Data, nil)
if err != nil {
aghhttp.Error(r, w, http.StatusBadRequest, "%s", err)
return
}
err = clients.storage.Update(r.Context(), dj.Name, c)
err = clients.storage.Update(ctx, dj.Name, c)
if err != nil {
aghhttp.Error(r, w, http.StatusBadRequest, "%s", err)
return
}
if !clients.testing {
onConfigModified()
}
clients.confModifier.Apply(ctx)
}
// handleFindClient is the handler for GET /control/clients/find HTTP API.

View File

@@ -4,12 +4,14 @@ 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"
@@ -23,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"
@@ -838,3 +841,47 @@ func validateTLSCipherIDs(cipherIDs []string) (err error) {
return nil
}
// defaultConfigModifier is a default [agh.ConfigModifier] implementation.
type defaultConfigModifier struct {
auth *auth
config *configuration
logger *slog.Logger
tlsMgr *tlsManager
}
// newDefaultConfigModifier returns the new properly initialized
// *defaultConfigModifier. All arguments must not be nil.
//
// TODO(s.chzhen): Consider using configuration struct.
func newDefaultConfigModifier(
conf *configuration,
l *slog.Logger,
) (cm *defaultConfigModifier) {
return &defaultConfigModifier{
config: conf,
logger: l,
}
}
// type check
var _ agh.ConfigModifier = (*defaultConfigModifier)(nil)
// Apply implements the [agh.ConfigModifier] interface for
// *defaultConfigModifier.
func (cm *defaultConfigModifier) Apply(ctx context.Context) {
err := cm.config.write(cm.tlsMgr, cm.auth)
if err != nil {
cm.logger.ErrorContext(ctx, "writing config", slogutil.KeyError, err)
}
}
// setAuth sets the auth parameters used by Apply.
func (cm *defaultConfigModifier) setAuth(a *auth) {
cm.auth = a
}
// setTLSManager sets the TLS manager used by Apply.
func (cm *defaultConfigModifier) setTLSManager(m *tlsManager) {
cm.tlsMgr = m
}

View File

@@ -115,6 +115,8 @@ type statusResponse struct {
}
func (web *webAPI) handleStatus(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
dnsAddrs, err := collectDNSAddresses(web.tlsManager)
if err != nil {
// Don't add a lot of formatting, since the error is already
@@ -125,14 +127,14 @@ func (web *webAPI) handleStatus(w http.ResponseWriter, r *http.Request) {
}
var (
fltConf *dnsforward.Config
protectionDisabledUntil *time.Time
protectionEnabled bool
fltConf *dnsforward.Config
protDisabledUntil *time.Time
protEnabled bool
)
if globalContext.dnsServer != nil {
fltConf = &dnsforward.Config{}
globalContext.dnsServer.WriteDiskConfig(fltConf)
protectionEnabled, protectionDisabledUntil = globalContext.dnsServer.UpdatedProtectionStatus()
protEnabled, protDisabledUntil = globalContext.dnsServer.UpdatedProtectionStatus(ctx)
}
var resp statusResponse
@@ -141,11 +143,11 @@ func (web *webAPI) handleStatus(w http.ResponseWriter, r *http.Request) {
defer config.RUnlock()
var protectionDisabledDuration int64
if protectionDisabledUntil != nil {
if protDisabledUntil != nil {
// Make sure that we don't send negative numbers to the frontend,
// since enough time might have passed to make the difference less
// than zero.
protectionDisabledDuration = max(0, time.Until(*protectionDisabledUntil).Milliseconds())
protectionDisabledDuration = max(0, time.Until(*protDisabledUntil).Milliseconds())
}
resp = statusResponse{
@@ -155,7 +157,7 @@ func (web *webAPI) handleStatus(w http.ResponseWriter, r *http.Request) {
DNSPort: config.DNS.Port,
HTTPPort: config.HTTPConfig.Address.Port(),
ProtectionDisabledDuration: protectionDisabledDuration,
ProtectionEnabled: protectionEnabled,
ProtectionEnabled: protEnabled,
IsRunning: isRunning(),
}
}()
@@ -175,10 +177,10 @@ func registerControlHandlers(web *webAPI) {
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", web.handleGetProfile)
httpRegister(http.MethodPut, "/control/profile/update", handlePutProfile)
httpRegister(http.MethodPut, "/control/profile/update", web.handlePutProfile)
// No auth is necessary for DoH/DoT configurations
globalContext.mux.HandleFunc("/apple/doh.mobileconfig", postInstall(handleMobileConfigDoH))

View File

@@ -15,6 +15,7 @@ import (
"time"
"unicode/utf8"
"github.com/AdguardTeam/AdGuardHome/internal/agh"
"github.com/AdguardTeam/AdGuardHome/internal/aghalg"
"github.com/AdguardTeam/AdGuardHome/internal/aghhttp"
"github.com/AdguardTeam/AdGuardHome/internal/aghnet"
@@ -455,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(ctx, web.baseLogger, web.tlsManager)
err = startMods(ctx, web.baseLogger, web.tlsManager, web.confModifier)
if err != nil {
globalContext.firstRun = true
copyInstallSettings(config, curConfig)
@@ -532,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
}
@@ -547,7 +553,7 @@ func startMods(ctx context.Context, baseLogger *slog.Logger, tlsMgr *tlsManager)
err = startDNSServer()
if err != nil {
closeDNSServer()
closeDNSServer(ctx)
return err
}

View File

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

View File

@@ -18,6 +18,7 @@ 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"
@@ -297,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 {
@@ -327,6 +329,7 @@ func initContextClients(
arpDB,
config.Filtering,
sigHdlr,
confModifier,
)
}
@@ -377,6 +380,7 @@ func setupDNSFilteringConf(
baseLogger *slog.Logger,
conf *filtering.Config,
tlsMgr *tlsManager,
confModifier agh.ConfigModifier,
) (err error) {
const (
dnsTimeout = 3 * time.Second
@@ -398,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)
@@ -564,6 +568,7 @@ func initWeb(
baseLogger *slog.Logger,
tlsMgr *tlsManager,
auth *auth,
confModifier agh.ConfigModifier,
isCustomUpdURL bool,
) (web *webAPI, err error) {
logger := baseLogger.With(slogutil.KeyPrefix, "webapi")
@@ -583,11 +588,12 @@ func initWeb(
disableUpdate := !isUpdateEnabled(ctx, baseLogger, &opts, isCustomUpdURL)
webConf := &webConfig{
updater: upd,
logger: logger,
baseLogger: baseLogger,
tlsManager: tlsMgr,
auth: auth,
updater: upd,
logger: logger,
baseLogger: baseLogger,
confModifier: confModifier,
tlsManager: tlsMgr,
auth: auth,
clientFS: clientFS,
@@ -661,25 +667,31 @@ func run(
// 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)
@@ -715,8 +727,19 @@ func run(
fatalOnError(err)
globalContext.auth = auth
confModifier.setAuth(auth)
web, err := initWeb(ctx, opts, clientBuildFS, upd, slogLogger, tlsMgr, auth, isCustomURL)
web, err := initWeb(
ctx,
opts,
clientBuildFS,
upd,
slogLogger,
tlsMgr,
auth,
confModifier,
isCustomURL,
)
fatalOnError(err)
globalContext.web = web
@@ -728,7 +751,7 @@ func run(
fatalOnError(err)
if !globalContext.firstRun {
err = initDNS(slogLogger, tlsMgr, statsDir, querylogDir)
err = initDNS(ctx, slogLogger, tlsMgr, confModifier, statsDir, querylogDir)
fatalOnError(err)
tlsMgr.start(ctx)
@@ -736,7 +759,7 @@ func run(
go func() {
startErr := startDNSServer()
if startErr != nil {
closeDNSServer()
closeDNSServer(ctx)
fatalOnError(startErr)
}
}()
@@ -972,7 +995,7 @@ func cleanup(ctx context.Context) {
globalContext.web = nil
}
err := stopDNSServer()
err := stopDNSServer(ctx)
if err != nil {
log.Error("stopping dns server: %s", err)
}
@@ -1130,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")

View File

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

View File

@@ -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.
@@ -75,7 +74,9 @@ func (web *webAPI) 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
}
@@ -103,10 +104,10 @@ func handlePutProfile(w http.ResponseWriter, r *http.Request) {
config.Language = lang
config.Theme = theme
log.Printf("home: language is set to %s", lang)
log.Printf("home: theme is set to %s", theme)
web.logger.InfoContext(ctx, "profile updated", "lang", lang, "theme", theme)
}()
onConfigModified()
web.confModifier.Apply(ctx)
aghhttp.OK(w)
}

View File

@@ -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,12 +91,12 @@ 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(ctx, conf.logger)
@@ -232,7 +232,7 @@ func (m *tlsManager) reload(ctx context.Context) {
m.certLastMod = fi.ModTime().UTC()
err = m.reconfigureDNSServer()
err = m.reconfigureDNSServer(ctx)
if err != nil {
m.logger.ErrorContext(ctx, "reconfiguring dns server", slogutil.KeyError, err)
}
@@ -245,7 +245,7 @@ func (m *tlsManager) reload(ctx context.Context) {
// reconfigureDNSServer updates the DNS server configuration using the stored
// TLS settings. m.mu is expected to be locked.
func (m *tlsManager) reconfigureDNSServer() (err error) {
func (m *tlsManager) reconfigureDNSServer(ctx context.Context) (err error) {
newConf, err := newServerConfig(
&config.DNS,
config.Clients.Sources,
@@ -253,12 +253,13 @@ func (m *tlsManager) reconfigureDNSServer() (err error) {
m,
httpRegister,
globalContext.clients.storage,
m.confModifier,
)
if err != nil {
return fmt.Errorf("generating forwarding dns server config: %w", err)
}
err = globalContext.dnsServer.Reconfigure(newConf)
err = globalContext.dnsServer.Reconfigure(ctx, newConf)
if err != nil {
return fmt.Errorf("starting forwarding dns server: %w", err)
}
@@ -515,7 +516,7 @@ func (m *tlsManager) handleTLSConfigure(w http.ResponseWriter, r *http.Request)
var restartHTTPS bool
defer func() {
if restartHTTPS {
m.configModified()
m.confModifier.Apply(ctx)
}
}()
@@ -557,7 +558,7 @@ func (m *tlsManager) handleTLSConfigure(w http.ResponseWriter, r *http.Request)
}()
}
err = m.reconfigureDNSServer()
err = m.reconfigureDNSServer(ctx)
if err != nil {
m.logger.ErrorContext(ctx, "reconfiguring dns server", slogutil.KeyError, err)

View File

@@ -20,6 +20,7 @@ import (
"testing"
"time"
"github.com/AdguardTeam/AdGuardHome/internal/agh"
"github.com/AdguardTeam/AdGuardHome/internal/aghalg"
"github.com/AdguardTeam/AdGuardHome/internal/client"
"github.com/AdguardTeam/AdGuardHome/internal/dnsforward"
@@ -66,9 +67,9 @@ func TestValidateCertificates(t *testing.T) {
ctx := testutil.ContextWithTimeout(t, testTimeout)
m, err := newTLSManager(ctx, &tlsManagerConfig{
logger: testLogger,
configModified: func() {},
servePlainDNS: false,
logger: testLogger,
confModifier: agh.EmptyConfigModifier{},
servePlainDNS: false,
})
require.NoError(t, err)
@@ -246,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,
@@ -257,7 +258,7 @@ func TestTLSManager_Reload(t *testing.T) {
})
require.NoError(t, err)
web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, nil, false)
web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, nil, agh.EmptyConfigModifier{}, false)
require.NoError(t, err)
m.setWebAPI(web)
@@ -272,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)
@@ -285,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),
@@ -321,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, nil, false)
web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, nil, agh.EmptyConfigModifier{}, false)
require.NoError(t, err)
m.setWebAPI(web)
@@ -420,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),
@@ -431,7 +434,7 @@ func TestTLSManager_HandleTLSValidate(t *testing.T) {
})
require.NoError(t, err)
web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, nil, false)
web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, nil, agh.EmptyConfigModifier{}, false)
require.NoError(t, err)
m.setWebAPI(web)
@@ -476,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{
@@ -511,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,
@@ -522,7 +527,7 @@ func TestTLSManager_HandleTLSConfigure(t *testing.T) {
})
require.NoError(t, err)
web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, nil, false)
web, err := initWeb(ctx, options{}, nil, nil, testLogger, nil, nil, agh.EmptyConfigModifier{}, false)
require.NoError(t, err)
m.setWebAPI(web)
@@ -551,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{},

View File

@@ -12,6 +12,7 @@ import (
"sync"
"time"
"github.com/AdguardTeam/AdGuardHome/internal/agh"
"github.com/AdguardTeam/AdGuardHome/internal/updater"
"github.com/AdguardTeam/golibs/errors"
"github.com/AdguardTeam/golibs/logutil/slogutil"
@@ -47,6 +48,9 @@ type webConfig struct {
// nil.
baseLogger *slog.Logger
// confModifier is used to update the global configuration.
confModifier agh.ConfigModifier
// tlsManager contains the current configuration and state of TLS
// encryption. It must not be nil.
tlsManager *tlsManager
@@ -104,6 +108,9 @@ type httpsServer struct {
type webAPI struct {
conf *webConfig
// confModifier is used to update the global configuration.
confModifier agh.ConfigModifier
// TODO(a.garipov): Refactor all these servers.
httpServer *http.Server
@@ -134,11 +141,12 @@ func newWebAPI(ctx context.Context, conf *webConfig) (w *webAPI) {
conf.logger.InfoContext(ctx, "initializing")
w = &webAPI{
conf: conf,
logger: conf.logger,
baseLogger: conf.baseLogger,
tlsManager: conf.tlsManager,
auth: conf.auth,
conf: conf,
confModifier: conf.confModifier,
logger: conf.logger,
baseLogger: conf.baseLogger,
tlsManager: conf.tlsManager,
auth: conf.auth,
}
clientFS := http.FileServer(http.FS(conf.clientFS))