From c2c1a7d5864a02cd6f57074c2034ff8fbc1b90b6 Mon Sep 17 00:00:00 2001 From: Eugene Burkov Date: Thu, 14 Aug 2025 20:45:22 +0300 Subject: [PATCH] all: sync with master --- .github/workflows/build.yml | 2 +- .github/workflows/lint.yml | 2 +- .gitignore | 2 + CHANGELOG.md | 55 +- Makefile | 2 +- bamboo-specs/release.yaml | 6 +- bamboo-specs/test.yaml | 37 +- client/src/__locales/be.json | 18 +- client/src/__locales/cs.json | 3 + client/src/__locales/da.json | 3 + client/src/__locales/de.json | 3 + client/src/__locales/en.json | 3 + client/src/__locales/es.json | 11 +- client/src/__locales/fr.json | 3 + client/src/__locales/it.json | 3 + client/src/__locales/ja.json | 3 + client/src/__locales/ko.json | 3 + client/src/__locales/nl.json | 3 + client/src/__locales/pt-br.json | 3 + client/src/__locales/pt-pt.json | 3 + client/src/__locales/ru.json | 3 + client/src/__locales/sk.json | 3 + client/src/__locales/tr.json | 5 +- client/src/__locales/zh-cn.json | 3 + client/src/__locales/zh-tw.json | 3 + client/src/components/Logs/Logs.css | 3 +- .../components/Settings/Dns/Cache/Form.tsx | 33 +- .../components/Settings/Dns/Cache/index.tsx | 3 +- client/src/initialState.ts | 1 + go.mod | 47 +- go.sum | 100 ++-- internal/agh/agh.go | 22 + internal/aghnet/hostscontainer.go | 5 +- .../aghnet/hostscontainer_internal_test.go | 5 +- internal/aghnet/hostscontainer_test.go | 33 +- internal/aghnet/net_darwin_internal_test.go | 4 +- internal/aghnet/net_internal_test.go | 18 +- internal/aghos/aghos_test.go | 11 - internal/aghos/filewalker_test.go | 42 +- internal/aghos/fswatcher.go | 74 +-- internal/aghos/os.go | 15 +- internal/aghrenameio/renameio_test.go | 8 +- internal/aghtest/aghtest.go | 20 +- internal/aghtest/interface.go | 76 ++- internal/aghtest/interface_test.go | 18 +- internal/aghtest/upstream.go | 2 +- internal/aghtls/aghtls.go | 11 +- internal/aghtls/aghtls_test.go | 9 +- internal/aghtls/root.go | 6 +- internal/aghtls/root_linux.go | 75 +-- internal/aghtls/root_others.go | 8 +- internal/arpdb/arpdb_internal_test.go | 8 +- internal/client/storage_test.go | 8 +- internal/configmigrate/configmigrate.go | 2 +- .../configmigrate/migrations_internal_test.go | 6 +- internal/configmigrate/migrator.go | 1 + internal/configmigrate/migrator_test.go | 8 + .../TestMigrateConfig_Migrate/v30/input.yml | 119 +++++ .../TestMigrateConfig_Migrate/v30/output.yml | 120 +++++ internal/configmigrate/v30.go | 33 ++ internal/dhcpd/config.go | 6 +- internal/dhcpd/dhcpd.go | 2 +- internal/dhcpd/http_unix.go | 6 +- internal/dhcpd/http_unix_internal_test.go | 37 +- internal/dhcpd/v4_unix_internal_test.go | 6 +- internal/dnsforward/access.go | 16 +- internal/dnsforward/beforerequest.go | 10 +- .../dnsforward/beforerequest_internal_test.go | 3 +- internal/dnsforward/clientid_internal_test.go | 1 + internal/dnsforward/config.go | 75 ++- internal/dnsforward/dialcontext.go | 5 +- internal/dnsforward/dns64_internal_test.go | 16 +- internal/dnsforward/dnsforward.go | 70 ++- .../dnsforward/dnsforward_internal_test.go | 141 ++--- internal/dnsforward/dnsrewrite.go | 15 +- .../dnsforward/dnsrewrite_internal_test.go | 21 +- internal/dnsforward/filter.go | 34 +- internal/dnsforward/filter_internal_test.go | 9 +- internal/dnsforward/http.go | 62 ++- internal/dnsforward/http_internal_test.go | 66 ++- internal/dnsforward/ipset.go | 4 +- internal/dnsforward/ipset_internal_test.go | 7 +- internal/dnsforward/msg.go | 86 +++- internal/dnsforward/process.go | 104 ++-- internal/dnsforward/process_internal_test.go | 31 +- internal/dnsforward/stats.go | 36 +- internal/dnsforward/stats_internal_test.go | 4 +- internal/dnsforward/svcbmsg.go | 13 +- internal/dnsforward/svcbmsg_internal_test.go | 5 +- .../TestDNSForwardHTTP_handleGetConfig.json | 3 + .../TestDNSForwardHTTP_handleSetConfig.json | 70 +++ internal/filtering/blocked.go | 12 +- internal/filtering/filter.go | 93 ++-- internal/filtering/filter_internal_test.go | 43 +- internal/filtering/filtering.go | 171 ++++--- internal/filtering/filtering_internal_test.go | 30 +- internal/filtering/hosts_test.go | 9 +- internal/filtering/http.go | 18 +- internal/filtering/http_internal_test.go | 34 +- .../filtering/idgenerator_internal_test.go | 6 +- internal/filtering/rewrite/storage.go | 62 ++- internal/filtering/rewritehttp.go | 52 +- internal/filtering/rewritehttp_test.go | 29 +- internal/filtering/rulelist/rulelist_test.go | 22 +- internal/filtering/safesearchhttp.go | 10 +- internal/filtering/servicelist.go | 35 ++ internal/home/auth.go | 481 +++++------------- internal/home/auth_internal_test.go | 79 ++- internal/home/authglinet.go | 125 ++--- internal/home/authglinet_internal_test.go | 26 +- internal/home/authhttp.go | 380 +++++++------- internal/home/authhttp_internal_test.go | 362 ++++++------- internal/home/authratelimiter.go | 42 +- internal/home/clients.go | 16 +- internal/home/clients_internal_test.go | 16 +- internal/home/clientshttp.go | 45 +- internal/home/clientshttp_internal_test.go | 19 +- internal/home/config.go | 58 ++- internal/home/control.go | 34 +- internal/home/controlinstall.go | 26 +- internal/home/dns.go | 52 +- internal/home/home.go | 195 ++++--- internal/home/i18n.go | 9 +- internal/home/middlewares_internal_test.go | 6 +- internal/home/mobileconfig_internal_test.go | 6 +- internal/home/options_internal_test.go | 26 +- internal/home/profilehttp.go | 27 +- internal/home/service.go | 160 +++--- internal/home/signal.go | 25 +- internal/home/tls.go | 37 +- internal/home/tls_internal_test.go | 70 +-- internal/home/web.go | 44 +- internal/next/websvc/dns_test.go | 3 +- internal/next/websvc/websvc_test.go | 66 +-- internal/querylog/http.go | 4 +- internal/querylog/qlog_internal_test.go | 22 +- internal/querylog/qlogfile_internal_test.go | 30 +- internal/querylog/qlogreader_internal_test.go | 14 +- internal/querylog/querylog.go | 7 +- internal/stats/http.go | 4 +- internal/stats/http_internal_test.go | 3 +- internal/stats/stats.go | 14 +- internal/stats/stats_test.go | 12 +- internal/updater/updater.go | 21 +- openapi/CHANGELOG.md | 4 + openapi/openapi.yaml | 8 + scripts/make/go-lint.sh | 75 ++- scripts/make/helper.sh | 42 +- scripts/make/txt-lint.sh | 29 +- 149 files changed, 3133 insertions(+), 2320 deletions(-) create mode 100644 internal/agh/agh.go delete mode 100644 internal/aghos/aghos_test.go create mode 100644 internal/configmigrate/testdata/TestMigrateConfig_Migrate/v30/input.yml create mode 100644 internal/configmigrate/testdata/TestMigrateConfig_Migrate/v30/output.yml create mode 100644 internal/configmigrate/v30.go diff --git a/.github/workflows/build.yml b/.github/workflows/build.yml index 617f3119..26352904 100644 --- a/.github/workflows/build.yml +++ b/.github/workflows/build.yml @@ -1,7 +1,7 @@ 'name': 'build' 'env': - 'GO_VERSION': '1.24.5' + 'GO_VERSION': '1.24.6' 'NODE_VERSION': '20' 'on': diff --git a/.github/workflows/lint.yml b/.github/workflows/lint.yml index 2c0b13e8..7853f146 100644 --- a/.github/workflows/lint.yml +++ b/.github/workflows/lint.yml @@ -1,7 +1,7 @@ 'name': 'lint' 'env': - 'GO_VERSION': '1.24.5' + 'GO_VERSION': '1.24.6' 'on': 'push': diff --git a/.gitignore b/.gitignore index a9598645..87040946 100644 --- a/.gitignore +++ b/.gitignore @@ -37,5 +37,7 @@ AdGuardHome.exe AdGuardHome.yaml* coverage.txt node_modules/ +test-reports/ +tmp/ !/build/gitkeep diff --git a/CHANGELOG.md b/CHANGELOG.md index fa04f0a8..1c8cffca 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -13,7 +13,7 @@ The format is based on [*Keep a Changelog*](https://keepachangelog.com/en/1.0.0/ See also the [v0.107.65 GitHub milestone][ms-v0.107.65]. -[ms-v0.107.65]: https://github.com/AdguardTeam/AdGuardHome/milestone/100?closed=1 +[ms-v0.107.65]: https://github.com/AdguardTeam/AdGuardHome/milestone/101?closed=1 NOTE: Add new changes BELOW THIS COMMENT. --> @@ -21,6 +21,52 @@ NOTE: Add new changes BELOW THIS COMMENT. NOTE: Add new changes ABOVE THIS COMMENT. --> +## [v0.107.65] - 2025-08-18 + +See also the [v0.107.65 GitHub milestone][ms-v0.107.65]. + +### Security + +- Go version has been updated to prevent the possibility of exploiting the Go vulnerabilities fixed in [1.24.6][go-1.24.6]. + +### Added + +- A separate checkbox in the Web UI to enable or disable the global DNS response cache without losing the configured cache size. + +- A new `"cache_enabled"` field to the HTTP API (`GET /control/dns_info` and `POST /control/dns_config`). See `openapi/openapi.yaml` for the full description. + +### Changed + +#### Configuration changes + +In this release, the schema version has changed from 29 to 30. + +- Added a new boolean field `dns.cache_enabled` to the configuration. This field explicitly controls whether DNS caching is enabled, replacing the previous implicit logic based on `dns.cache_size`. + + ```yaml + # BEFORE: + 'dns': + # … + 'cache_size': 123456 + + # AFTER: + 'dns': + # … + 'cache_enabled': true + 'cache_size': 123456 + ``` + + To roll back this change, set the schema_version back to `29`. + +### Fixed + +- Disabled state of *Top clients* action button in web UI ([#7923]). + +[#7923]: https://github.com/AdguardTeam/AdGuardHome/issues/7923 + +[go-1.24.6]: https://groups.google.com/g/golang-announce/c/x5MKroML2yM +[ms-v0.107.65]: https://github.com/AdguardTeam/AdGuardHome/milestone/100?closed=1 + ## [v0.107.64] - 2025-07-28 See also the [v0.107.64 GitHub milestone][ms-v0.107.64]. @@ -3177,11 +3223,12 @@ See also the [v0.104.2 GitHub milestone][ms-v0.104.2]. [ms-v0.104.2]: https://github.com/AdguardTeam/AdGuardHome/milestone/28?closed=1 -[Unreleased]: https://github.com/AdguardTeam/AdGuardHome/compare/v0.107.64...HEAD +[Unreleased]: https://github.com/AdguardTeam/AdGuardHome/compare/v0.107.65...HEAD +[v0.107.65]: https://github.com/AdguardTeam/AdGuardHome/compare/v0.107.64...v0.107.65 [v0.107.64]: https://github.com/AdguardTeam/AdGuardHome/compare/v0.107.63...v0.107.64 [v0.107.63]: https://github.com/AdguardTeam/AdGuardHome/compare/v0.107.62...v0.107.63 [v0.107.62]: https://github.com/AdguardTeam/AdGuardHome/compare/v0.107.61...v0.107.62 diff --git a/Makefile b/Makefile index a3bd3cfe..06b3e420 100644 --- a/Makefile +++ b/Makefile @@ -27,7 +27,7 @@ DIST_DIR = dist GOAMD64 = v1 GOPROXY = https://proxy.golang.org|direct GOTELEMETRY = off -GOTOOLCHAIN = go1.24.5 +GOTOOLCHAIN = go1.24.6 GPG_KEY = devteam@adguard.com GPG_KEY_PASSPHRASE = not-a-real-password NPM = npm diff --git a/bamboo-specs/release.yaml b/bamboo-specs/release.yaml index 116ade6f..838c94ed 100644 --- a/bamboo-specs/release.yaml +++ b/bamboo-specs/release.yaml @@ -8,7 +8,7 @@ 'variables': 'channel': 'edge' 'dockerFrontend': 'adguard/home-js-builder:3.1' - 'dockerGo': 'adguard/go-builder:1.24.5--2' + 'dockerGo': 'adguard/go-builder:1.24.6--1' 'stages': - 'Build frontend': @@ -279,7 +279,7 @@ 'variables': 'channel': 'beta' 'dockerFrontend': 'adguard/home-js-builder:3.1' - 'dockerGo': 'adguard/go-builder:1.24.5--2' + 'dockerGo': 'adguard/go-builder:1.24.6--1' # release-vX.Y.Z branches are the branches from which the actual final # release is built. - '^release-v[0-9]+\.[0-9]+\.[0-9]+': @@ -295,4 +295,4 @@ 'variables': 'channel': 'release' 'dockerFrontend': 'adguard/home-js-builder:3.1' - 'dockerGo': 'adguard/go-builder:1.24.5--2' + 'dockerGo': 'adguard/go-builder:1.24.6--1' diff --git a/bamboo-specs/test.yaml b/bamboo-specs/test.yaml index 3c94be2f..9bb4cf83 100644 --- a/bamboo-specs/test.yaml +++ b/bamboo-specs/test.yaml @@ -6,7 +6,7 @@ 'name': 'AdGuard Home - Build and run tests' 'variables': 'dockerFrontend': 'adguard/home-js-builder:3.1' - 'dockerGo': 'adguard/go-builder:1.24.5--2' + 'dockerGo': 'adguard/go-builder:1.24.6--1' 'channel': 'development' 'stages': @@ -67,6 +67,13 @@ 'volumes': '${system.GO_CACHE_DIR}': '${bamboo.cacheGo}' '${system.GO_PKG_CACHE_DIR}': '${bamboo.cacheGoPkg}' + 'final-tasks': + - 'test-parser': + # The default pattern, '**/test-reports/*.xml', works, so don't set + # the test-results property. + 'type': 'junit' + 'ignore-time': true + - 'clean' 'key': 'GOTEST' 'other': 'clean-working-dir': true @@ -81,16 +88,26 @@ set -e -f -u -x - make\ - GOMAXPROCS=1\ - VERBOSE=1\ + make \ + GOMAXPROCS=1 \ + VERBOSE=1 \ go-deps go-tools go-lint - make\ - VERBOSE=1\ - go-test - 'final-tasks': - - 'clean' + make \ + TEST_REPORTS_DIR="./test-reports/" \ + VERBOSE=1 \ + go-test \ + ; + + exit_code="$(cat ./test-reports/test-exit-code.txt)" + readonly exit_code + + make VERBOSE=1 \ + go-fuzz \ + go-bench \ + ; + + exit "$exit_code" 'requirements': - 'adg-docker': 'true' @@ -234,5 +251,5 @@ # may need to build a few of these. 'variables': 'dockerFrontend': 'adguard/home-js-builder:3.1' - 'dockerGo': 'adguard/go-builder:1.24.5--2' + 'dockerGo': 'adguard/go-builder:1.24.6--1' 'channel': 'candidate' diff --git a/client/src/__locales/be.json b/client/src/__locales/be.json index db963b2e..4484f9ad 100644 --- a/client/src/__locales/be.json +++ b/client/src/__locales/be.json @@ -89,7 +89,7 @@ "form_enter_hostname": "Увядзіце імя хаста", "error_details": "Дэталізацыя памылкі", "response_details": "Дэталі адказу", - "request_details": "Інфармацыя пра запыт", + "request_details": "Падрабязнасці запыту", "client_details": "Дэталі кліента", "details": "Дэталі", "back": "Назад", @@ -108,7 +108,7 @@ "off": "Выкл", "copyright": "Усе правы захаваныя", "homepage": "Хатняя старонка", - "report_an_issue": "Паведаміць пра праблему", + "report_an_issue": "Паведаміць аб праблеме", "privacy_policy": "Палітыка прыватнасці", "enable_protection": "Уключыць абарону", "enabled_protection": "Абарона ўкл.", @@ -736,13 +736,13 @@ "thursday": "Чацвер", "friday": "Пятніца", "saturday": "Субота", - "sunday_short": "Нд.", - "monday_short": "Пн.", - "tuesday_short": "Аў.", - "wednesday_short": "Ср.", - "thursday_short": "Чц.", - "friday_short": "Пт.", - "saturday_short": "Сб.", + "sunday_short": "Ндз", + "monday_short": "Пан", + "tuesday_short": "Аўт", + "wednesday_short": "Срд", + "thursday_short": "Чцв", + "friday_short": "Птн", + "saturday_short": "Суб", "upstream_dns_cache_configuration": "Канфігурацыя кэша upstream сервер DNSаў", "enable_upstream_dns_cache": "Ўключыць кэшаванне для карыстацкай канфігурацыі upstream-сервераў гэтага кліента", "dns_cache_size": "Памер кэша DNS, у байтах" diff --git a/client/src/__locales/cs.json b/client/src/__locales/cs.json index 4851d0ba..9a6106c5 100644 --- a/client/src/__locales/cs.json +++ b/client/src/__locales/cs.json @@ -655,7 +655,10 @@ "safe_search": "Bezpečné vyhledávání", "blocklist": "Zakázaný", "milliseconds_abbreviation": "ms", + "cache_enabled": "Povolit mezipaměť", + "cache_enabled_desc": "Ukládejte odezvy DNS lokálně.", "cache_size": "Velikost mezipaměti", + "cache_size_validation": "Velikost mezipaměti musí být větší než nula, pokud je tato funkce povolena.", "cache_size_desc": "Velikost mezipaměti DNS (v bajtech). Chcete-li ukládání do mezipaměti zakázat, nastavte 0.", "cache_ttl_min_override": "Přepsat minimální hodnotu TTL", "cache_ttl_max_override": "Přepsat maximální hodnotu TTL", diff --git a/client/src/__locales/da.json b/client/src/__locales/da.json index 9acaaece..f74f92f3 100644 --- a/client/src/__locales/da.json +++ b/client/src/__locales/da.json @@ -655,7 +655,10 @@ "safe_search": "Sikker søgning", "blocklist": "Sortliste", "milliseconds_abbreviation": "ms", + "cache_enabled": "Aktivér cache", + "cache_enabled_desc": "Opbevar DNS-svar lokalt.", "cache_size": "Cache-størrelse", + "cache_size_validation": "Cache-størrelsen skal være større end nul, når den er aktiveret.", "cache_size_desc": "DNS cache-størrelse (i bytes). Sæt til 0 for at deaktivere cache.", "cache_ttl_min_override": "Tilsidesæt minimum TTL", "cache_ttl_max_override": "Tilsidesæt maksimal TTL", diff --git a/client/src/__locales/de.json b/client/src/__locales/de.json index 4c855793..e52d9a3b 100644 --- a/client/src/__locales/de.json +++ b/client/src/__locales/de.json @@ -655,7 +655,10 @@ "safe_search": "Sichere Suche", "blocklist": "Sperrliste", "milliseconds_abbreviation": "ms", + "cache_enabled": "Cache aktivieren", + "cache_enabled_desc": "DNS-Antworten lokal speichern.", "cache_size": "Größe des Cache", + "cache_size_validation": "Die Cachegröße muss größer als Null sein, wenn diese Option aktiviert ist.", "cache_size_desc": "Größe des DNS-Cache (in Bytes). Um das Caching zu deaktivieren, setzen Sie den Wert auf 0.", "cache_ttl_min_override": "TTL-Minimalwert überschreiben", "cache_ttl_max_override": "TTL-Höchstwert überschreiben", diff --git a/client/src/__locales/en.json b/client/src/__locales/en.json index e8e17250..cbf089f8 100644 --- a/client/src/__locales/en.json +++ b/client/src/__locales/en.json @@ -655,7 +655,10 @@ "safe_search": "Safe Search", "blocklist": "Blocklist", "milliseconds_abbreviation": "ms", + "cache_enabled": "Enable cache", + "cache_enabled_desc": "Store DNS responses locally.", "cache_size": "Cache size", + "cache_size_validation": "The cache size must be greater than zero when enabled.", "cache_size_desc": "DNS cache size (in bytes). To disable caching, set to 0.", "cache_ttl_min_override": "Override minimum TTL", "cache_ttl_max_override": "Override maximum TTL", diff --git a/client/src/__locales/es.json b/client/src/__locales/es.json index 4b1fec83..11468a15 100644 --- a/client/src/__locales/es.json +++ b/client/src/__locales/es.json @@ -428,9 +428,9 @@ "encryption_hostnames": "Nombres de hosts", "encryption_reset": "¿Estás seguro de que deseas restablecer la configuración de cifrado?", "encryption_warning": "Advertencia", - "encryption_plain_dns_enable": "Activar DNS simple (sin cifrado)", - "encryption_plain_dns_desc": "El DNS simple (sin cifrado) está activado de forma predeterminada. Puedes desactivarlo para obligar a todos los dispositivos a utilizar DNS cifrado. Para ello, debes habilitar al menos un protocolo DNS cifrado", - "encryption_plain_dns_error": "Para desactivar el DNS simple, activa al menos un protocolo DNS cifrado", + "encryption_plain_dns_enable": "Habilitar DNS simple", + "encryption_plain_dns_desc": "El DNS simple está habilitado de manera predeterminada. Puedes deshabilitarlo para obligar a todos los dispositivos a utilizar DNS cifrado. Para ello, debe habilitar al menos un protocolo DNS cifrado", + "encryption_plain_dns_error": "Para deshabilitar el DNS simple, habilita al menos un protocolo DNS cifrado", "topline_expiring_certificate": "Tu certificado SSL está a punto de expirar. Actualiza la <0>configuración de cifrado.", "topline_expired_certificate": "Tu certificado SSL ha expirado. Actualiza la <0>configuración de cifrado.", "form_error_port_range": "Ingresa el número del puerto en el rango de 80 a 65535", @@ -655,8 +655,11 @@ "safe_search": "Búsqueda segura", "blocklist": "Lista de bloqueo", "milliseconds_abbreviation": "ms", + "cache_enabled": "Activar caché", + "cache_enabled_desc": "Almacene las respuestas de DNS localmente.", "cache_size": "Tamaño de la caché", - "cache_size_desc": "Tamaño de la caché DNS (en bytes). Para desactivar el almacenamiento en caché, configúralo en 0.", + "cache_size_validation": "El tamaño de la cache debe ser mayor que cero cuando está habilitado.", + "cache_size_desc": "Tamaño de la caché DNS (en bytes). Para deshabilitar el almacenamiento en caché, establécelo en 0.", "cache_ttl_min_override": "Anular TTL mínimo", "cache_ttl_max_override": "Anular TTL máximo", "enter_cache_size": "Ingresa el tamaño de la caché (bytes)", diff --git a/client/src/__locales/fr.json b/client/src/__locales/fr.json index 92f7afba..01475bad 100644 --- a/client/src/__locales/fr.json +++ b/client/src/__locales/fr.json @@ -655,7 +655,10 @@ "safe_search": "Recherche Sécurisée", "blocklist": "Liste de blocage", "milliseconds_abbreviation": "ms", + "cache_enabled": "Activer le cache", + "cache_enabled_desc": "Stockez les réponses DNS localement.", "cache_size": "Taille du cache", + "cache_size_validation": "La taille du cache doit être supérieure à zéro lorsqu'elle est activée.", "cache_size_desc": "Taille du cache DNS (en octets). Pour désactiver la mise en cache, mettez la valeur sur 0.", "cache_ttl_min_override": "Remplacer le TTL minimum", "cache_ttl_max_override": "Remplacer le TTL maximum", diff --git a/client/src/__locales/it.json b/client/src/__locales/it.json index 1f591b70..7d9fb128 100644 --- a/client/src/__locales/it.json +++ b/client/src/__locales/it.json @@ -655,7 +655,10 @@ "safe_search": "Ricerca Sicura", "blocklist": "Lista nera", "milliseconds_abbreviation": "ms", + "cache_enabled": "Abilita la cache", + "cache_enabled_desc": "Memorizza localmente le risposte DNS.", "cache_size": "Dimensioni cache", + "cache_size_validation": "La dimensione della cache deve essere maggiore di zero quando abilitata.", "cache_size_desc": "Dimensione della memoria temporanea DNS (in byte). Per disabilitare la memoria temporanea, impostare a 0.", "cache_ttl_min_override": "Sovrascrivi TTL minimo", "cache_ttl_max_override": "Sovrascrivi TTL massimo", diff --git a/client/src/__locales/ja.json b/client/src/__locales/ja.json index 2c679a8d..56174cfa 100644 --- a/client/src/__locales/ja.json +++ b/client/src/__locales/ja.json @@ -655,7 +655,10 @@ "safe_search": "セーフサーチ", "blocklist": "ブロックリスト", "milliseconds_abbreviation": "ms", + "cache_enabled": "キャッシュを有効にする", + "cache_enabled_desc": "DNSレスポンスをローカルに保存します。", "cache_size": "キャッシュサイズ", + "cache_size_validation": "キャッシュが有効の場合、キャッシュサイズはゼロより大きい値でなければなりません", "cache_size_desc": "DNSキャッシュサイズ(バイト単位)※キャッシュを無効化するには、「0」(ゼロ)にしてください。", "cache_ttl_min_override": "最小TTLの上書き(秒単位)", "cache_ttl_max_override": "最大TTLの上書き(秒単位)", diff --git a/client/src/__locales/ko.json b/client/src/__locales/ko.json index d4f1cb10..159eda48 100644 --- a/client/src/__locales/ko.json +++ b/client/src/__locales/ko.json @@ -655,7 +655,10 @@ "safe_search": "세이프서치", "blocklist": "차단 목록", "milliseconds_abbreviation": "ms", + "cache_enabled": "캐시 활성화", + "cache_enabled_desc": "DNS 응답을 로컬에 저장합니다.", "cache_size": "캐시 크기", + "cache_size_validation": "활성화된 경우 캐시 크기는 0보다 커야 합니다.", "cache_size_desc": "DNS 캐시 크기(바이트). 캐싱을 사용하지 않으려면 0으로 설정합니다.", "cache_ttl_min_override": "최소 TTL (초) 무시", "cache_ttl_max_override": "최대 TTL (초) 무시", diff --git a/client/src/__locales/nl.json b/client/src/__locales/nl.json index 558e9dd4..a49baf0d 100644 --- a/client/src/__locales/nl.json +++ b/client/src/__locales/nl.json @@ -655,7 +655,10 @@ "safe_search": "Veilig zoeken", "blocklist": "Blokkeerlijst", "milliseconds_abbreviation": "ms", + "cache_enabled": "Cache inschakelen", + "cache_enabled_desc": "DNS-antwoorden lokaal opslaan.", "cache_size": "Cache grootte", + "cache_size_validation": "De cachegrootte moet groter zijn dan nul wanneer deze is ingeschakeld.", "cache_size_desc": "DNS-cachegrootte (in bytes). Om caching uit te schakelen, stel deze in op 0.", "cache_ttl_min_override": "Minimale TTL overschrijven", "cache_ttl_max_override": "Maximale TTL overschrijven", diff --git a/client/src/__locales/pt-br.json b/client/src/__locales/pt-br.json index 78203501..43e1f482 100644 --- a/client/src/__locales/pt-br.json +++ b/client/src/__locales/pt-br.json @@ -655,7 +655,10 @@ "safe_search": "Pesquisa segura", "blocklist": "Lista de bloqueio", "milliseconds_abbreviation": "ms", + "cache_enabled": "Ativar cache", + "cache_enabled_desc": "Armazenar as respostas DNS localmente.", "cache_size": "Tamanho do cache", + "cache_size_validation": "O tamanho do cache deve ser maior que zero quando ativado.", "cache_size_desc": "Tamanho do cache do DNS (em bytes). Para desativar o cache, defina como 0.", "cache_ttl_min_override": "Sobrepor o TTL mínimo", "cache_ttl_max_override": "Sobrepor o TTL máximo", diff --git a/client/src/__locales/pt-pt.json b/client/src/__locales/pt-pt.json index a4e9bcc4..9aa47f44 100644 --- a/client/src/__locales/pt-pt.json +++ b/client/src/__locales/pt-pt.json @@ -655,7 +655,10 @@ "safe_search": "Pesquisa segura", "blocklist": "Lista de bloqueio", "milliseconds_abbreviation": "ms", + "cache_enabled": "Ativar cache", + "cache_enabled_desc": "Armazene as respostas DNS localmente.", "cache_size": "Tamanho do cache", + "cache_size_validation": "O tamanho do cache deve ser maior que zero quando ativado.", "cache_size_desc": "Tamanho do cache DNS (em bytes). Para desativar o cache, defina como 0.", "cache_ttl_min_override": "Sobrepor o TTL mínimo", "cache_ttl_max_override": "Sobrepor o TTL máximo", diff --git a/client/src/__locales/ru.json b/client/src/__locales/ru.json index 81568339..85715714 100644 --- a/client/src/__locales/ru.json +++ b/client/src/__locales/ru.json @@ -655,7 +655,10 @@ "safe_search": "Безопасный поиск", "blocklist": "Чёрный список", "milliseconds_abbreviation": "мс", + "cache_enabled": "Включить кеш", + "cache_enabled_desc": "Сохранять локально ответы DNS.", "cache_size": "Размер кеша", + "cache_size_validation": "Если кеш включен, его размер должен быть больше нуля.", "cache_size_desc": "Размер кеша DNS (в байтах). Чтобы отключить кеширование, установите значение 0.", "cache_ttl_min_override": "Переопределить минимальный TTL", "cache_ttl_max_override": "Переопределить максимальный TTL", diff --git a/client/src/__locales/sk.json b/client/src/__locales/sk.json index 1ba33318..0484565e 100644 --- a/client/src/__locales/sk.json +++ b/client/src/__locales/sk.json @@ -655,7 +655,10 @@ "safe_search": "Bezpečné vyhľadávanie", "blocklist": "Zoznam blokovaní", "milliseconds_abbreviation": "ms", + "cache_enabled": "Povoliť vyrovnávaciu pamäť", + "cache_enabled_desc": "Ukladať DNS odpovede lokálne.", "cache_size": "Veľkosť cache", + "cache_size_validation": "Veľkosť vyrovnávacej pamäte musí byť po povolení väčšia ako nula.", "cache_size_desc": "Veľkosť vyrovnávacej pamäte DNS (v bajtoch). Ak chcete vypnúť ukladanie do vyrovnávacej pamäte, nastavte hodnotu 0.", "cache_ttl_min_override": "Prepísať minimálne TTL", "cache_ttl_max_override": "Prepísať maximálne TTL", diff --git a/client/src/__locales/tr.json b/client/src/__locales/tr.json index aab240e1..97ae1c9c 100644 --- a/client/src/__locales/tr.json +++ b/client/src/__locales/tr.json @@ -245,7 +245,7 @@ "block_for_this_client_only": "Yalnızca bu istemci için engelle", "unblock_for_this_client_only": "Yalnızca bu istemci için engellemeyi kaldır", "add_persistent_client": "Kalıcı istemci olarak ekle", - "time_table_header": "Saat", + "time_table_header": "Süre", "date": "Tarih", "domain_name_table_header": "Alan adı", "domain_or_client": "Alan adı veya istemci", @@ -655,7 +655,10 @@ "safe_search": "Güvenli Arama", "blocklist": "Engel listesi", "milliseconds_abbreviation": "ms", + "cache_enabled": "Önbelleği etkinleştir", + "cache_enabled_desc": "DNS yanıtlarını yerel olarak depolayın.", "cache_size": "Önbellek boyutu", + "cache_size_validation": "Etkinleştirildiğinde önbellek boyutu sıfırdan büyük olmalıdır.", "cache_size_desc": "DNS önbellek boyutu (bayt cinsinden). Önbelleği devre dışı bırakmak için 0 olarak ayarlayın.", "cache_ttl_min_override": "En az kullanım süresini geçersiz kıl", "cache_ttl_max_override": "En fazla kullanım süresini geçersiz kıl", diff --git a/client/src/__locales/zh-cn.json b/client/src/__locales/zh-cn.json index 68ddf586..c2d50561 100644 --- a/client/src/__locales/zh-cn.json +++ b/client/src/__locales/zh-cn.json @@ -655,7 +655,10 @@ "safe_search": "安全搜索", "blocklist": "黑名单", "milliseconds_abbreviation": "毫秒", + "cache_enabled": "启用缓存", + "cache_enabled_desc": "在本地存储 DNS 响应。", "cache_size": "缓存大小", + "cache_size_validation": "启用时,缓存大小必须大于 0。", "cache_size_desc": "DNS 缓存大小(单位:字节)。若要禁用缓存,请设置为 0。", "cache_ttl_min_override": "覆盖最小 TTL 值", "cache_ttl_max_override": "覆盖最大 TTL 值", diff --git a/client/src/__locales/zh-tw.json b/client/src/__locales/zh-tw.json index 3b3e0fed..11a9cf6b 100644 --- a/client/src/__locales/zh-tw.json +++ b/client/src/__locales/zh-tw.json @@ -655,7 +655,10 @@ "safe_search": "安全搜尋", "blocklist": "封鎖清單", "milliseconds_abbreviation": "ms", + "cache_enabled": "啟用快取", + "cache_enabled_desc": "在本機儲存 DNS 回應。", "cache_size": "快取大小", + "cache_size_validation": "啟用時,快取大小必須大於 0。", "cache_size_desc": "DNS 快取大小(位元組)。若要停用快取,請設為 0。", "cache_ttl_min_override": "覆寫最小的存活時間(TTL)", "cache_ttl_max_override": "覆寫最大的存活時間(TTL)", diff --git a/client/src/components/Logs/Logs.css b/client/src/components/Logs/Logs.css index 8ab03da9..3714a06e 100644 --- a/client/src/components/Logs/Logs.css +++ b/client/src/components/Logs/Logs.css @@ -337,7 +337,8 @@ } .button-action--arrow-option:disabled { - display: none; + opacity: 0.5; + cursor: default; } .tooltip-custom__container .button-action--arrow-option { diff --git a/client/src/components/Settings/Dns/Cache/Form.tsx b/client/src/components/Settings/Dns/Cache/Form.tsx index aea6a299..62323e2d 100644 --- a/client/src/components/Settings/Dns/Cache/Form.tsx +++ b/client/src/components/Settings/Dns/Cache/Form.tsx @@ -32,6 +32,7 @@ const INPUTS_FIELDS = [ ]; type FormData = { + cache_enabled: boolean; cache_size: number; cache_ttl_min: number; cache_ttl_max: number; @@ -54,10 +55,11 @@ const Form = ({ initialValues, onSubmit }: CacheFormProps) => { handleSubmit, watch, control, - formState: { isSubmitting, isDirty }, + formState: { isSubmitting }, } = useForm({ mode: 'onBlur', defaultValues: { + cache_enabled: initialValues?.cache_enabled || false, cache_size: initialValues?.cache_size || 0, cache_ttl_min: initialValues?.cache_ttl_min || 0, cache_ttl_max: initialValues?.cache_ttl_max || 0, @@ -65,10 +67,13 @@ const Form = ({ initialValues, onSubmit }: CacheFormProps) => { }, }); + const cache_enabled = watch('cache_enabled'); + const cache_size = watch('cache_size'); const cache_ttl_min = watch('cache_ttl_min'); const cache_ttl_max = watch('cache_ttl_max'); const minExceedsMax = cache_ttl_min > 0 && cache_ttl_max > 0 && cache_ttl_min > cache_ttl_max; + const cacheSizeZeroWhenEnabled = cache_enabled && cache_size === 0; const handleClearCache = () => { if (window.confirm(t('confirm_dns_cache_clear'))) { @@ -79,6 +84,24 @@ const Form = ({ initialValues, onSubmit }: CacheFormProps) => { return (
+
+
+ ( + + )} + /> +
+
+ {INPUTS_FIELDS.map(({ name, title, description, placeholder }) => (
@@ -102,6 +125,12 @@ const Form = ({ initialValues, onSubmit }: CacheFormProps) => { setValueAs: (value) => replaceZeroWithEmptyString(value), })} /> + + {name === CACHE_CONFIG_FIELDS.cache_size && cacheSizeZeroWhenEnabled && ( + + {t('cache_size_validation')} + + )}
@@ -133,7 +162,7 @@ const Form = ({ initialValues, onSubmit }: CacheFormProps) => { type="submit" data-testid="dns_save" className="btn btn-success btn-standard btn-large" - disabled={isSubmitting || !isDirty || processingSetConfig || minExceedsMax}> + disabled={isSubmitting || processingSetConfig || minExceedsMax || cacheSizeZeroWhenEnabled}> {t('save_btn')} diff --git a/client/src/components/Settings/Dns/Cache/index.tsx b/client/src/components/Settings/Dns/Cache/index.tsx index 1eeff122..51c04c41 100644 --- a/client/src/components/Settings/Dns/Cache/index.tsx +++ b/client/src/components/Settings/Dns/Cache/index.tsx @@ -13,7 +13,7 @@ import { RootState } from '../../../../initialState'; const CacheConfig = () => { const { t } = useTranslation(); const dispatch = useDispatch(); - const { cache_size, cache_ttl_max, cache_ttl_min, cache_optimistic } = useSelector( + const { cache_enabled, cache_size, cache_ttl_max, cache_ttl_min, cache_optimistic } = useSelector( (state: RootState) => state.dnsConfig, shallowEqual, ); @@ -32,6 +32,7 @@ const CacheConfig = () => {
0 + + return nil +} diff --git a/internal/dhcpd/config.go b/internal/dhcpd/config.go index d11d9342..4cef310d 100644 --- a/internal/dhcpd/config.go +++ b/internal/dhcpd/config.go @@ -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:"-"` diff --git a/internal/dhcpd/dhcpd.go b/internal/dhcpd/dhcpd.go index edc3d3a4..5da089d3 100644 --- a/internal/dhcpd/dhcpd.go +++ b/internal/dhcpd/dhcpd.go @@ -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, diff --git a/internal/dhcpd/http_unix.go b/internal/dhcpd/http_unix.go index db81bafc..6d6226bd 100644 --- a/internal/dhcpd/http_unix.go +++ b/internal/dhcpd/http_unix.go @@ -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) { diff --git a/internal/dhcpd/http_unix_internal_test.go b/internal/dhcpd/http_unix_internal_test.go index 80d37050..7f4c4390 100644 --- a/internal/dhcpd/http_unix_internal_test.go +++ b/internal/dhcpd/http_unix_internal_test.go @@ -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) diff --git a/internal/dhcpd/v4_unix_internal_test.go b/internal/dhcpd/v4_unix_internal_test.go index 4e0df75b..ae80b5bd 100644 --- a/internal/dhcpd/v4_unix_internal_test.go +++ b/internal/dhcpd/v4_unix_internal_test.go @@ -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 } diff --git a/internal/dnsforward/access.go b/internal/dnsforward/access.go index c5535d30..ef58dff9 100644 --- a/internal/dnsforward/access.go +++ b/internal/dnsforward/access.go @@ -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() diff --git a/internal/dnsforward/beforerequest.go b/internal/dnsforward/beforerequest.go index 952ea253..1d0a138e 100644 --- a/internal/dnsforward/beforerequest.go +++ b/internal/dnsforward/beforerequest.go @@ -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) } diff --git a/internal/dnsforward/beforerequest_internal_test.go b/internal/dnsforward/beforerequest_internal_test.go index 35a1157b..d9b872b8 100644 --- a/internal/dnsforward/beforerequest_internal_test.go +++ b/internal/dnsforward/beforerequest_internal_test.go @@ -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) diff --git a/internal/dnsforward/clientid_internal_test.go b/internal/dnsforward/clientid_internal_test.go index ec110f60..f095c448 100644 --- a/internal/dnsforward/clientid_internal_test.go +++ b/internal/dnsforward/clientid_internal_test.go @@ -201,6 +201,7 @@ func TestServer_clientIDFromDNSContext(t *testing.T) { srv := &Server{ conf: ServerConfig{TLSConf: tlsConf}, baseLogger: testLogger, + logger: testLogger, } var ( diff --git a/internal/dnsforward/config.go b/internal/dnsforward/config.go index a5d83c6c..a2dfeefe 100644 --- a/internal/dnsforward/config.go +++ b/internal/dnsforward/config.go @@ -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 diff --git a/internal/dnsforward/dialcontext.go b/internal/dnsforward/dialcontext.go index 0ed91fb8..669f6c58 100644 --- a/internal/dnsforward/dialcontext.go +++ b/internal/dnsforward/dialcontext.go @@ -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 { diff --git a/internal/dnsforward/dns64_internal_test.go b/internal/dnsforward/dns64_internal_test.go index 2d2c5471..b730ac3f 100644 --- a/internal/dnsforward/dns64_internal_test.go +++ b/internal/dnsforward/dns64_internal_test.go @@ -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() diff --git a/internal/dnsforward/dnsforward.go b/internal/dnsforward/dnsforward.go index 05814288..9b2e91f6 100644 --- a/internal/dnsforward/dnsforward.go +++ b/internal/dnsforward/dnsforward.go @@ -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 } diff --git a/internal/dnsforward/dnsforward_internal_test.go b/internal/dnsforward/dnsforward_internal_test.go index 53be7bbc..df63a140 100644 --- a/internal/dnsforward/dnsforward_internal_test.go +++ b/internal/dnsforward/dnsforward_internal_test.go @@ -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) diff --git a/internal/dnsforward/dnsrewrite.go b/internal/dnsforward/dnsrewrite.go index 7d9fde72..00e4cd56 100644 --- a/internal/dnsforward/dnsrewrite.go +++ b/internal/dnsforward/dnsrewrite.go @@ -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) } diff --git a/internal/dnsforward/dnsrewrite_internal_test.go b/internal/dnsforward/dnsrewrite_internal_test.go index f30c661b..ed523fb5 100644 --- a/internal/dnsforward/dnsrewrite_internal_test.go +++ b/internal/dnsforward/dnsrewrite_internal_test.go @@ -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) diff --git a/internal/dnsforward/filter.go b/internal/dnsforward/filter.go index 6cfd7bea..ebb08f84 100644 --- a/internal/dnsforward/filter.go +++ b/internal/dnsforward/filter.go @@ -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 } diff --git a/internal/dnsforward/filter_internal_test.go b/internal/dnsforward/filter_internal_test.go index 8a07b4f1..99598dd2 100644 --- a/internal/dnsforward/filter_internal_test.go +++ b/internal/dnsforward/filter_internal_test.go @@ -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 diff --git a/internal/dnsforward/http.go b/internal/dnsforward/http.go index 59f3fde8..572a6247 100644 --- a/internal/dnsforward/http.go +++ b/internal/dnsforward/http.go @@ -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) } diff --git a/internal/dnsforward/http_internal_test.go b/internal/dnsforward/http_internal_test.go index f57783c7..2da75e87 100644 --- a/internal/dnsforward/http_internal_test.go +++ b/internal/dnsforward/http_internal_test.go @@ -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")) }) } diff --git a/internal/dnsforward/ipset.go b/internal/dnsforward/ipset.go index 7347890a..4204fbb3 100644 --- a/internal/dnsforward/ipset.go +++ b/internal/dnsforward/ipset.go @@ -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") diff --git a/internal/dnsforward/ipset_internal_test.go b/internal/dnsforward/ipset_internal_test.go index 90c200d0..d7735b91 100644 --- a/internal/dnsforward/ipset_internal_test.go +++ b/internal/dnsforward/ipset_internal_test.go @@ -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) diff --git a/internal/dnsforward/msg.go b/internal/dnsforward/msg.go index e9f1f2d7..6fd8580b 100644 --- a/internal/dnsforward/msg.go +++ b/internal/dnsforward/msg.go @@ -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) } diff --git a/internal/dnsforward/process.go b/internal/dnsforward/process.go index 623edff0..d2d85295 100644 --- a/internal/dnsforward/process.go +++ b/internal/dnsforward/process.go @@ -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 diff --git a/internal/dnsforward/process_internal_test.go b/internal/dnsforward/process_internal_test.go index 8b335832..d9b78180 100644 --- a/internal/dnsforward/process_internal_test.go +++ b/internal/dnsforward/process_internal_test.go @@ -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) }) diff --git a/internal/dnsforward/stats.go b/internal/dnsforward/stats.go index 50818b40..90590455 100644 --- a/internal/dnsforward/stats.go +++ b/internal/dnsforward/stats.go @@ -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, ) } diff --git a/internal/dnsforward/stats_internal_test.go b/internal/dnsforward/stats_internal_test.go index 301f8c8b..3bce993a 100644 --- a/internal/dnsforward/stats_internal_test.go +++ b/internal/dnsforward/stats_internal_test.go @@ -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) diff --git a/internal/dnsforward/svcbmsg.go b/internal/dnsforward/svcbmsg.go index 96983dee..07c171aa 100644 --- a/internal/dnsforward/svcbmsg.go +++ b/internal/dnsforward/svcbmsg.go @@ -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 } diff --git a/internal/dnsforward/svcbmsg_internal_test.go b/internal/dnsforward/svcbmsg_internal_test.go index 7de96da8..9a2e6d93 100644 --- a/internal/dnsforward/svcbmsg_internal_test.go +++ b/internal/dnsforward/svcbmsg_internal_test.go @@ -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) }) }) diff --git a/internal/dnsforward/testdata/TestDNSForwardHTTP_handleGetConfig.json b/internal/dnsforward/testdata/TestDNSForwardHTTP_handleGetConfig.json index d4e76a20..9a3e4750 100644 --- a/internal/dnsforward/testdata/TestDNSForwardHTTP_handleGetConfig.json +++ b/internal/dnsforward/testdata/TestDNSForwardHTTP_handleGetConfig.json @@ -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, diff --git a/internal/dnsforward/testdata/TestDNSForwardHTTP_handleSetConfig.json b/internal/dnsforward/testdata/TestDNSForwardHTTP_handleSetConfig.json index d0967ece..37e25945 100644 --- a/internal/dnsforward/testdata/TestDNSForwardHTTP_handleSetConfig.json +++ b/internal/dnsforward/testdata/TestDNSForwardHTTP_handleSetConfig.json @@ -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, diff --git a/internal/filtering/blocked.go b/internal/filtering/blocked.go index 8150f309..855db7d6 100644 --- a/internal/filtering/blocked.go +++ b/internal/filtering/blocked.go @@ -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) } diff --git a/internal/filtering/filter.go b/internal/filtering/filter.go index 240b623d..ac5e21f1 100644 --- a/internal/filtering/filter.go +++ b/internal/filtering/filter.go @@ -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) } } diff --git a/internal/filtering/filter_internal_test.go b/internal/filtering/filter_internal_test.go index c2d0de71..144ec061 100644 --- a/internal/filtering/filter_internal_test.go +++ b/internal/filtering/filter_internal_test.go @@ -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 } diff --git a/internal/filtering/filtering.go b/internal/filtering/filtering.go index 41d437c3..ac4e2f94 100644 --- a/internal/filtering/filtering.go +++ b/internal/filtering/filtering.go @@ -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) diff --git a/internal/filtering/filtering_internal_test.go b/internal/filtering/filtering_internal_test.go index 5097755c..d5b14aff 100644 --- a/internal/filtering/filtering_internal_test.go +++ b/internal/filtering/filtering_internal_test.go @@ -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) { diff --git a/internal/filtering/hosts_test.go b/internal/filtering/hosts_test.go index 5c692814..9385314a 100644 --- a/internal/filtering/hosts_test.go +++ b/internal/filtering/hosts_test.go @@ -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) diff --git a/internal/filtering/http.go b/internal/filtering/http.go index dca039a0..55b8d4ac 100644 --- a/internal/filtering/http.go +++ b/internal/filtering/http.go @@ -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 diff --git a/internal/filtering/http_internal_test.go b/internal/filtering/http_internal_test.go index 4d45e254..326d91ff 100644 --- a/internal/filtering/http_internal_test.go +++ b/internal/filtering/http_internal_test.go @@ -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 }, diff --git a/internal/filtering/idgenerator_internal_test.go b/internal/filtering/idgenerator_internal_test.go index e9c0db2f..195e9976 100644 --- a/internal/filtering/idgenerator_internal_test.go +++ b/internal/filtering/idgenerator_internal_test.go @@ -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()) } diff --git a/internal/filtering/rewrite/storage.go b/internal/filtering/rewrite/storage.go index bab84089..9fea5eeb 100644 --- a/internal/filtering/rewrite/storage.go +++ b/internal/filtering/rewrite/storage.go @@ -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. diff --git a/internal/filtering/rewritehttp.go b/internal/filtering/rewritehttp.go index d6415a05..d685b27d 100644 --- a/internal/filtering/rewritehttp.go +++ b/internal/filtering/rewritehttp.go @@ -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", diff --git a/internal/filtering/rewritehttp_test.go b/internal/filtering/rewritehttp_test.go index b95435b8..4b57c9db 100644 --- a/internal/filtering/rewritehttp_test.go +++ b/internal/filtering/rewritehttp_test.go @@ -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. diff --git a/internal/filtering/rulelist/rulelist_test.go b/internal/filtering/rulelist/rulelist_test.go index 85a3d362..0e966da2 100644 --- a/internal/filtering/rulelist/rulelist_test.go +++ b/internal/filtering/rulelist/rulelist_test.go @@ -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 } diff --git a/internal/filtering/safesearchhttp.go b/internal/filtering/safesearchhttp.go index 8790b297..b7a6a4f3 100644 --- a/internal/filtering/safesearchhttp.go +++ b/internal/filtering/safesearchhttp.go @@ -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) } diff --git a/internal/filtering/servicelist.go b/internal/filtering/servicelist.go index 0e52ee48..278edcab 100644 --- a/internal/filtering/servicelist.go +++ b/internal/filtering/servicelist.go @@ -502,6 +502,16 @@ var blockedServices = []blockedService{{ "||globosat.globo.com^", "||gsatmulti.globo.com^", }, +}, { + ID: "chatgpt", + Name: "ChatGPT", + IconSVG: []byte(""), + 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(""), + 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(""), + 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(""), + Rules: []string{ + "||odycdn.com^", + "||odysee.com^", + "||odysee.live^", + "||odysee.tv^", + }, }, { ID: "ok", Name: "OK.ru", diff --git a/internal/home/auth.go b/internal/home/auth.go index c5a6d129..f0e074c0 100644 --- a/internal/home/auth.go +++ b/internal/home/auth.go @@ -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 } diff --git a/internal/home/auth_internal_test.go b/internal/home/auth_internal_test.go index 41aabd87..11af39e8 100644 --- a/internal/home/auth_internal_test.go +++ b/internal/home/auth_internal_test.go @@ -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)) } diff --git a/internal/home/authglinet.go b/internal/home/authglinet.go index 784d14b1..18c6fc4c 100644 --- a/internal/home/authglinet.go +++ b/internal/home/authglinet.go @@ -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) }) } diff --git a/internal/home/authglinet_internal_test.go b/internal/home/authglinet_internal_test.go index 258f1652..e0242609 100644 --- a/internal/home/authglinet_internal_test.go +++ b/internal/home/authglinet_internal_test.go @@ -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)) -} diff --git a/internal/home/authhttp.go b/internal/home/authhttp.go index 7dcf5030..bc8e7068 100644 --- a/internal/home/authhttp.go +++ b/internal/home/authhttp.go @@ -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 diff --git a/internal/home/authhttp_internal_test.go b/internal/home/authhttp_internal_test.go index f59e455c..ceb13c25 100644 --- a/internal/home/authhttp_internal_test.go +++ b/internal/home/authhttp_internal_test.go @@ -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) { diff --git a/internal/home/authratelimiter.go b/internal/home/authratelimiter.go index fcd2c127..95e8d571 100644 --- a/internal/home/authratelimiter.go +++ b/internal/home/authratelimiter.go @@ -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() diff --git a/internal/home/clients.go b/internal/home/clients.go index e6459342..a6f37ba1 100644 --- a/internal/home/clients.go +++ b/internal/home/clients.go @@ -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() diff --git a/internal/home/clients_internal_test.go b/internal/home/clients_internal_test.go index 8bbfd185..69f24315 100644 --- a/internal/home/clients_internal_test.go +++ b/internal/home/clients_internal_test.go @@ -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 } diff --git a/internal/home/clientshttp.go b/internal/home/clientshttp.go index 010df861..f1691166 100644 --- a/internal/home/clientshttp.go +++ b/internal/home/clientshttp.go @@ -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, } } diff --git a/internal/home/clientshttp_internal_test.go b/internal/home/clientshttp_internal_test.go index 01197983..054d897e 100644 --- a/internal/home/clientshttp_internal_test.go +++ b/internal/home/clientshttp_internal_test.go @@ -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{}, }, }} diff --git a/internal/home/config.go b/internal/home/config.go index 2ec549d3..65452f8f 100644 --- a/internal/home/config.go +++ b/internal/home/config.go @@ -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 +} diff --git a/internal/home/control.go b/internal/home/control.go index 4be955f3..75c82e5d 100644 --- a/internal/home/control.go +++ b/internal/home/control.go @@ -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 diff --git a/internal/home/controlinstall.go b/internal/home/controlinstall.go index 602b0f64..8ea1cfa2 100644 --- a/internal/home/controlinstall.go +++ b/internal/home/controlinstall.go @@ -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 } diff --git a/internal/home/dns.go b/internal/home/dns.go index dea65bf7..07372efd 100644 --- a/internal/home/dns.go +++ b/internal/home/dns.go @@ -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) } diff --git a/internal/home/home.go b/internal/home/home.go index f38a350f..9536219d 100644 --- a/internal/home/home.go +++ b/internal/home/home.go @@ -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. "+ diff --git a/internal/home/i18n.go b/internal/home/i18n.go index d49ca2fa..4b616605 100644 --- a/internal/home/i18n.go +++ b/internal/home/i18n.go @@ -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) } diff --git a/internal/home/middlewares_internal_test.go b/internal/home/middlewares_internal_test.go index 0b8d7db3..1d3a6b9c 100644 --- a/internal/home/middlewares_internal_test.go +++ b/internal/home/middlewares_internal_test.go @@ -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) }) } diff --git a/internal/home/mobileconfig_internal_test.go b/internal/home/mobileconfig_internal_test.go index 86710003..225ca53e 100644 --- a/internal/home/mobileconfig_internal_test.go +++ b/internal/home/mobileconfig_internal_test.go @@ -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 }) diff --git a/internal/home/options_internal_test.go b/internal/home/options_internal_test.go index dbcf02cc..9a51aff8 100644 --- a/internal/home/options_internal_test.go +++ b/internal/home/options_internal_test.go @@ -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) { diff --git a/internal/home/profilehttp.go b/internal/home/profilehttp.go index 0b1dcf99..41d2c18e 100644 --- a/internal/home/profilehttp.go +++ b/internal/home/profilehttp.go @@ -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) } diff --git a/internal/home/service.go b/internal/home/service.go index e3e1cab6..af3d4a70 100644 --- a/internal/home/service.go +++ b/internal/home/service.go @@ -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) } } } diff --git a/internal/home/signal.go b/internal/home/signal.go index 638d3632..04822ff5 100644 --- a/internal/home/signal.go +++ b/internal/home/signal.go @@ -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) diff --git a/internal/home/tls.go b/internal/home/tls.go index 058a4ba0..7f1d6316 100644 --- a/internal/home/tls.go +++ b/internal/home/tls.go @@ -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) diff --git a/internal/home/tls_internal_test.go b/internal/home/tls_internal_test.go index b65d0a97..6fc13d06 100644 --- a/internal/home/tls_internal_test.go +++ b/internal/home/tls_internal_test.go @@ -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{}, diff --git a/internal/home/web.go b/internal/home/web.go index 9de27782..40f59230 100644 --- a/internal/home/web.go +++ b/internal/home/web.go @@ -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,10 +48,17 @@ 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 + // auth stores web user information and handles authentication. It must not + // be nil. + auth *auth + clientFS fs.FS // BindAddr is the binding address with port for plain HTTP web interface. @@ -100,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 @@ -114,6 +125,9 @@ type webAPI struct { // encryption. tlsManager *tlsManager + // auth stores web user information and handles authentication. + auth *auth + // httpsServer is the server that handles HTTPS traffic. If it is not nil, // [Web.http3Server] must also not be nil. httpsServer httpsServer @@ -127,16 +141,21 @@ 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, + conf: conf, + confModifier: conf.confModifier, + logger: conf.logger, + baseLogger: conf.baseLogger, + tlsManager: conf.tlsManager, + auth: conf.auth, } clientFS := http.FileServer(http.FS(conf.clientFS)) // if not configured, redirect / to /install.html, otherwise redirect /install.html to / - globalContext.mux.Handle("/", withMiddlewares(clientFS, gziphandler.GzipHandler, optionalAuthHandler, postInstallHandler)) + globalContext.mux.Handle( + "/", + withMiddlewares(clientFS, gziphandler.GzipHandler, postInstallHandler), + ) // add handlers for /install paths, we only need them when we're not configured yet if conf.firstRun { @@ -210,7 +229,10 @@ func (web *webAPI) start(ctx context.Context) { errs := make(chan error, 2) // Use an h2c handler to support unencrypted HTTP/2, e.g. for proxies. - hdlr := h2c.NewHandler(withMiddlewares(globalContext.mux, limitRequestBody), &http2.Server{}) + hdlr := h2c.NewHandler( + withMiddlewares(globalContext.mux, limitRequestBody), + &http2.Server{}, + ) logger := web.baseLogger.With(loggerKeyServer, "plain") @@ -221,7 +243,7 @@ func (web *webAPI) start(ctx context.Context) { // Create a new instance, because the Web is not usable after Shutdown. web.httpServer = &http.Server{ Addr: web.conf.BindAddr.String(), - Handler: hdlr, + Handler: web.auth.middleware().Wrap(hdlr), ReadTimeout: web.conf.ReadTimeout, ReadHeaderTimeout: web.conf.ReadHeaderTimeout, WriteTimeout: web.conf.WriteTimeout, @@ -262,6 +284,10 @@ func (web *webAPI) close(ctx context.Context) { shutdownSrv3(ctx, web.logger, web.httpsServer.server3) shutdownSrv(ctx, web.logger, web.httpServer) + if web.auth != nil { + web.auth.close(ctx) + } + web.logger.InfoContext(ctx, "stopped http server") } @@ -303,7 +329,7 @@ func (web *webAPI) tlsServerLoop(ctx context.Context) { web.httpsServer.server = &http.Server{ Addr: addr, - Handler: hdlr, + Handler: web.auth.middleware().Wrap(hdlr), TLSConfig: &tls.Config{ Certificates: []tls.Certificate{web.httpsServer.cert}, RootCAs: web.tlsManager.rootCerts, @@ -344,7 +370,7 @@ func (web *webAPI) mustStartHTTP3(ctx context.Context, address string) { CipherSuites: web.tlsManager.customCipherIDs, MinVersion: tls.VersionTLS12, }, - Handler: withMiddlewares(globalContext.mux, limitRequestBody), + Handler: web.auth.middleware().Wrap(withMiddlewares(globalContext.mux, limitRequestBody)), } web.logger.DebugContext(ctx, "starting http/3 server") diff --git a/internal/next/websvc/dns_test.go b/internal/next/websvc/dns_test.go index 960bf024..c8362b0e 100644 --- a/internal/next/websvc/dns_test.go +++ b/internal/next/websvc/dns_test.go @@ -17,6 +17,7 @@ import ( "github.com/AdguardTeam/AdGuardHome/internal/next/websvc" "github.com/AdguardTeam/dnsproxy/proxy" "github.com/AdguardTeam/golibs/netutil/urlutil" + "github.com/AdguardTeam/golibs/testutil" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -45,7 +46,7 @@ func TestService_HandlePatchSettingsDNS(t *testing.T) { return nil }, - OnShutdown: func(_ context.Context) (err error) { panic("not implemented") }, + OnShutdown: func(ctx context.Context) (_ error) { panic(testutil.UnexpectedCall(ctx)) }, OnConfig: func() (c *dnssvc.Config) { return &dnssvc.Config{} }, } } diff --git a/internal/next/websvc/websvc_test.go b/internal/next/websvc/websvc_test.go index 79e46ac6..979ef21b 100644 --- a/internal/next/websvc/websvc_test.go +++ b/internal/next/websvc/websvc_test.go @@ -19,7 +19,7 @@ import ( "github.com/AdguardTeam/golibs/netutil/httputil" "github.com/AdguardTeam/golibs/netutil/urlutil" "github.com/AdguardTeam/golibs/testutil" - "github.com/AdguardTeam/golibs/testutil/fakefs" + "github.com/AdguardTeam/golibs/testutil/fakeio/fakefs" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -65,13 +65,17 @@ func (m *configManager) UpdateWeb(ctx context.Context, c *websvc.Config) (err er // newConfigManager returns a *configManager all methods of which panic. func newConfigManager() (m *configManager) { return &configManager{ - onDNS: func() (svc agh.ServiceWithConfig[*dnssvc.Config]) { panic("not implemented") }, - onWeb: func() (svc agh.ServiceWithConfig[*websvc.Config]) { panic("not implemented") }, - onUpdateDNS: func(_ context.Context, _ *dnssvc.Config) (err error) { - panic("not implemented") + onDNS: func() (_ agh.ServiceWithConfig[*dnssvc.Config]) { + panic(testutil.UnexpectedCall()) }, - onUpdateWeb: func(_ context.Context, _ *websvc.Config) (err error) { - panic("not implemented") + onWeb: func() (_ agh.ServiceWithConfig[*websvc.Config]) { + panic(testutil.UnexpectedCall()) + }, + onUpdateDNS: func(ctx context.Context, c *dnssvc.Config) (_ error) { + panic(testutil.UnexpectedCall(ctx, c)) + }, + onUpdateWeb: func(ctx context.Context, c *websvc.Config) (_ error) { + panic(testutil.UnexpectedCall(ctx, c)) }, } } @@ -80,10 +84,10 @@ func newConfigManager() (m *configManager) { // sole address. It also registers a cleanup procedure, which shuts the // instance down. func newTestServer( - t testing.TB, + tb testing.TB, confMgr websvc.ConfigManager, ) (svc *websvc.Service, addr netip.AddrPort) { - t.Helper() + tb.Helper() c := &websvc.Config{ Logger: slogutil.NewDiscardLogger(), @@ -103,17 +107,17 @@ func newTestServer( } svc, err := websvc.New(c) - require.NoError(t, err) + require.NoError(tb, err) - err = svc.Start(testutil.ContextWithTimeout(t, testTimeout)) - require.NoError(t, err) - testutil.CleanupAndRequireSuccess(t, func() (err error) { - return svc.Shutdown(testutil.ContextWithTimeout(t, testTimeout)) + err = svc.Start(testutil.ContextWithTimeout(tb, testTimeout)) + require.NoError(tb, err) + testutil.CleanupAndRequireSuccess(tb, func() (err error) { + return svc.Shutdown(testutil.ContextWithTimeout(tb, testTimeout)) }) c = svc.Config() - require.NotNil(t, c) - require.Len(t, c.Addresses, 1) + require.NotNil(tb, c) + require.Len(tb, c.Addresses, 1) return svc, c.Addresses[0] } @@ -125,23 +129,23 @@ type jobj map[string]any // the response as well as checks that the status code is correct. // // TODO(a.garipov): Add helpers for other methods. -func httpGet(t testing.TB, u *url.URL, wantCode int) (body []byte) { - t.Helper() +func httpGet(tb testing.TB, u *url.URL, wantCode int) (body []byte) { + tb.Helper() req, err := http.NewRequest(http.MethodGet, u.String(), nil) - require.NoErrorf(t, err, "creating req") + require.NoErrorf(tb, err, "creating req") httpCli := &http.Client{ Timeout: testTimeout, } resp, err := httpCli.Do(req) - require.NoErrorf(t, err, "performing req") - require.Equal(t, wantCode, resp.StatusCode) + require.NoErrorf(tb, err, "performing req") + require.Equal(tb, wantCode, resp.StatusCode) - testutil.CleanupAndRequireSuccess(t, resp.Body.Close) + testutil.CleanupAndRequireSuccess(tb, resp.Body.Close) body, err = io.ReadAll(resp.Body) - require.NoErrorf(t, err, "reading body") + require.NoErrorf(tb, err, "reading body") return body } @@ -151,26 +155,26 @@ func httpGet(t testing.TB, u *url.URL, wantCode int) (body []byte) { // checks that the status code is correct. // // TODO(a.garipov): Add helpers for other methods. -func httpPatch(t testing.TB, u *url.URL, reqBody any, wantCode int) (body []byte) { - t.Helper() +func httpPatch(tb testing.TB, u *url.URL, reqBody any, wantCode int) (body []byte) { + tb.Helper() b, err := json.Marshal(reqBody) - require.NoErrorf(t, err, "marshaling reqBody") + require.NoErrorf(tb, err, "marshaling reqBody") req, err := http.NewRequest(http.MethodPatch, u.String(), bytes.NewReader(b)) - require.NoErrorf(t, err, "creating req") + require.NoErrorf(tb, err, "creating req") httpCli := &http.Client{ Timeout: testTimeout, } resp, err := httpCli.Do(req) - require.NoErrorf(t, err, "performing req") - require.Equal(t, wantCode, resp.StatusCode) + require.NoErrorf(tb, err, "performing req") + require.Equal(tb, wantCode, resp.StatusCode) - testutil.CleanupAndRequireSuccess(t, resp.Body.Close) + testutil.CleanupAndRequireSuccess(tb, resp.Body.Close) body, err = io.ReadAll(resp.Body) - require.NoErrorf(t, err, "reading body") + require.NoErrorf(tb, err, "reading body") return body } diff --git a/internal/querylog/http.go b/internal/querylog/http.go index fb878e04..727be326 100644 --- a/internal/querylog/http.go +++ b/internal/querylog/http.go @@ -186,7 +186,7 @@ func (l *queryLog) handleQueryLogConfig(w http.ResponseWriter, r *http.Request) return } - defer l.conf.ConfigModified() + defer l.conf.ConfigModifier.Apply(r.Context()) l.confMu.Lock() defer l.confMu.Unlock() @@ -250,7 +250,7 @@ func (l *queryLog) handlePutQueryLogConfig(w http.ResponseWriter, r *http.Reques return } - defer l.conf.ConfigModified() + defer l.conf.ConfigModifier.Apply(r.Context()) l.confMu.Lock() defer l.confMu.Unlock() diff --git a/internal/querylog/qlog_internal_test.go b/internal/querylog/qlog_internal_test.go index 2a688552..65b81708 100644 --- a/internal/querylog/qlog_internal_test.go +++ b/internal/querylog/qlog_internal_test.go @@ -390,20 +390,20 @@ func addEntry(l *queryLog, host string, answerStr, client net.IP) { l.Add(params) } -func assertLogEntry(t *testing.T, entry *logEntry, host string, answer, client net.IP) { - t.Helper() +func assertLogEntry(tb testing.TB, entry *logEntry, host string, answer, client net.IP) { + tb.Helper() - require.NotNil(t, entry) + require.NotNil(tb, entry) - assert.Equal(t, host, entry.QHost) - assert.Equal(t, client, entry.IP) - assert.Equal(t, "A", entry.QType) - assert.Equal(t, "IN", entry.QClass) + assert.Equal(tb, host, entry.QHost) + assert.Equal(tb, client, entry.IP) + assert.Equal(tb, "A", entry.QType) + assert.Equal(tb, "IN", entry.QClass) msg := &dns.Msg{} - require.NoError(t, msg.Unpack(entry.Answer)) - require.Len(t, msg.Answer, 1) + require.NoError(tb, msg.Unpack(entry.Answer)) + require.Len(tb, msg.Answer, 1) - a := testutil.RequireTypeAssert[*dns.A](t, msg.Answer[0]) - assert.Equal(t, answer, a.A.To16()) + a := testutil.RequireTypeAssert[*dns.A](tb, msg.Answer[0]) + assert.Equal(tb, answer, a.A.To16()) } diff --git a/internal/querylog/qlogfile_internal_test.go b/internal/querylog/qlogfile_internal_test.go index 087d43aa..b80b9ce3 100644 --- a/internal/querylog/qlogfile_internal_test.go +++ b/internal/querylog/qlogfile_internal_test.go @@ -20,17 +20,17 @@ import ( // prepareTestFile prepares one test query log file with the specified lines // count. -func prepareTestFile(t *testing.T, dir string, linesNum int) (name string) { - t.Helper() +func prepareTestFile(tb testing.TB, dir string, linesNum int) (name string) { + tb.Helper() f, err := os.CreateTemp(dir, "*.txt") - require.NoError(t, err) + require.NoError(tb, err) // Use defer and not t.Cleanup to make sure that the file is closed // after this function is done. defer func() { derr := f.Close() - require.NoError(t, derr) + require.NoError(tb, derr) }() const ans = `"AAAAAAABAAEAAAAAB2V4YW1wbGUDb3JnAAABAAEHZXhhbXBsZQNvcmcAAAEAAQAAAAAABAECAwQ="` @@ -49,7 +49,7 @@ func prepareTestFile(t *testing.T, dir string, linesNum int) (name string) { line := fmt.Sprintf(format, ip, lineTime.Format(time.RFC3339Nano)) _, err = f.WriteString(line) - require.NoError(t, err) + require.NoError(tb, err) } return f.Name() @@ -57,18 +57,18 @@ func prepareTestFile(t *testing.T, dir string, linesNum int) (name string) { // prepareTestFiles prepares several test query log files, each with the // specified lines count. -func prepareTestFiles(t *testing.T, filesNum, linesNum int) []string { - t.Helper() +func prepareTestFiles(tb testing.TB, filesNum, linesNum int) []string { + tb.Helper() if filesNum == 0 { return []string{} } - dir := t.TempDir() + dir := tb.TempDir() files := make([]string, filesNum) for i := range files { - files[filesNum-i-1] = prepareTestFile(t, dir, linesNum) + files[filesNum-i-1] = prepareTestFile(tb, dir, linesNum) } return files @@ -76,17 +76,17 @@ func prepareTestFiles(t *testing.T, filesNum, linesNum int) []string { // newTestQLogFile creates new *qLogFile for tests and registers the required // cleanup functions. -func newTestQLogFile(t *testing.T, linesNum int) (file *qLogFile) { - t.Helper() +func newTestQLogFile(tb testing.TB, linesNum int) (file *qLogFile) { + tb.Helper() - testFile := prepareTestFiles(t, 1, linesNum)[0] + testFile := prepareTestFiles(tb, 1, linesNum)[0] // Create the new qLogFile instance. file, err := newQLogFile(testFile) - require.NoError(t, err) + require.NoError(tb, err) - assert.NotNil(t, file) - testutil.CleanupAndRequireSuccess(t, file.Close) + assert.NotNil(tb, file) + testutil.CleanupAndRequireSuccess(tb, file.Close) return file } diff --git a/internal/querylog/qlogreader_internal_test.go b/internal/querylog/qlogreader_internal_test.go index bb3ce164..8a1ed009 100644 --- a/internal/querylog/qlogreader_internal_test.go +++ b/internal/querylog/qlogreader_internal_test.go @@ -13,20 +13,20 @@ import ( // newTestQLogReader creates new *qLogReader for tests and registers the // required cleanup functions. -func newTestQLogReader(t *testing.T, filesNum, linesNum int) (reader *qLogReader) { - t.Helper() +func newTestQLogReader(tb testing.TB, filesNum, linesNum int) (reader *qLogReader) { + tb.Helper() - testFiles := prepareTestFiles(t, filesNum, linesNum) + testFiles := prepareTestFiles(tb, filesNum, linesNum) logger := slogutil.NewDiscardLogger() - ctx := testutil.ContextWithTimeout(t, testTimeout) + ctx := testutil.ContextWithTimeout(tb, testTimeout) // Create the new qLogReader instance. reader, err := newQLogReader(ctx, logger, testFiles) - require.NoError(t, err) + require.NoError(tb, err) - assert.NotNil(t, reader) - testutil.CleanupAndRequireSuccess(t, reader.Close) + assert.NotNil(tb, reader) + testutil.CleanupAndRequireSuccess(tb, reader.Close) return reader } diff --git a/internal/querylog/querylog.go b/internal/querylog/querylog.go index c7350f70..c4ed53cc 100644 --- a/internal/querylog/querylog.go +++ b/internal/querylog/querylog.go @@ -8,6 +8,7 @@ import ( "sync" "time" + "github.com/AdguardTeam/AdGuardHome/internal/agh" "github.com/AdguardTeam/AdGuardHome/internal/aghhttp" "github.com/AdguardTeam/AdGuardHome/internal/aghnet" "github.com/AdguardTeam/AdGuardHome/internal/filtering" @@ -47,9 +48,9 @@ type Config struct { // Anonymizer processes the IP addresses to anonymize those if needed. Anonymizer *aghnet.IPMut - // ConfigModified is called when the configuration is changed, for example - // by HTTP requests. - ConfigModified func() + // ConfigModifier is used to update the global configuration. It must not + // be nil. + ConfigModifier agh.ConfigModifier // HTTPRegister registers an HTTP handler. HTTPRegister aghhttp.RegisterFunc diff --git a/internal/stats/http.go b/internal/stats/http.go index c2ea01d0..67b78a97 100644 --- a/internal/stats/http.go +++ b/internal/stats/http.go @@ -166,7 +166,7 @@ func (s *StatsCtx) handleStatsConfig(w http.ResponseWriter, r *http.Request) { limit := time.Duration(reqData.IntervalDays) * timeutil.Day - defer s.configModified() + defer s.configModifier.Apply(ctx) s.confMu.Lock() defer s.confMu.Unlock() @@ -216,7 +216,7 @@ func (s *StatsCtx) handlePutStatsConfig(w http.ResponseWriter, r *http.Request) return } - defer s.configModified() + defer s.configModifier.Apply(ctx) s.confMu.Lock() defer s.confMu.Unlock() diff --git a/internal/stats/http_internal_test.go b/internal/stats/http_internal_test.go index b53668d6..9a53c01a 100644 --- a/internal/stats/http_internal_test.go +++ b/internal/stats/http_internal_test.go @@ -9,6 +9,7 @@ import ( "testing" "time" + "github.com/AdguardTeam/AdGuardHome/internal/agh" "github.com/AdguardTeam/AdGuardHome/internal/aghalg" "github.com/AdguardTeam/golibs/logutil/slogutil" "github.com/AdguardTeam/golibs/testutil" @@ -27,7 +28,7 @@ func TestHandleStatsConfig(t *testing.T) { conf := Config{ Logger: slogutil.NewDiscardLogger(), UnitID: func() (id uint32) { return 0 }, - ConfigModified: func() {}, + ConfigModifier: agh.EmptyConfigModifier{}, ShouldCountClient: func([]string) bool { return true }, Filename: filepath.Join(t.TempDir(), "stats.db"), Limit: time.Hour * 24, diff --git a/internal/stats/stats.go b/internal/stats/stats.go index 3b63df5b..a1506cf4 100644 --- a/internal/stats/stats.go +++ b/internal/stats/stats.go @@ -13,6 +13,7 @@ import ( "sync/atomic" "time" + "github.com/AdguardTeam/AdGuardHome/internal/agh" "github.com/AdguardTeam/AdGuardHome/internal/aghhttp" "github.com/AdguardTeam/AdGuardHome/internal/aghnet" "github.com/AdguardTeam/AdGuardHome/internal/aghos" @@ -55,9 +56,9 @@ type Config struct { // nil, the default function is used, see newUnitID. UnitID UnitIDGenFunc - // ConfigModified will be called each time the configuration changed via web - // interface. - ConfigModified func() + // ConfigModifier is used to update the global configuration. It must not + // be nil. + ConfigModifier agh.ConfigModifier // ShouldCountClient returns client's ignore setting. ShouldCountClient func([]string) bool @@ -123,9 +124,8 @@ type StatsCtx struct { // httpRegister is used to set HTTP handlers. httpRegister aghhttp.RegisterFunc - // configModified is called whenever the configuration is modified via web - // interface. - configModified func() + // configModifier is used to update the global configuration. + configModifier agh.ConfigModifier // confMu protects ignored, limit, and enabled. confMu *sync.RWMutex @@ -165,7 +165,7 @@ func New(conf Config) (s *StatsCtx, err error) { logger: conf.Logger, currMu: &sync.RWMutex{}, httpRegister: conf.HTTPRegister, - configModified: conf.ConfigModified, + configModifier: conf.ConfigModifier, filename: conf.Filename, confMu: &sync.RWMutex{}, diff --git a/internal/stats/stats_test.go b/internal/stats/stats_test.go index 06aa36f3..662d9be0 100644 --- a/internal/stats/stats_test.go +++ b/internal/stats/stats_test.go @@ -26,25 +26,25 @@ import ( // constUnitID is the UnitIDGenFunc which always return 0. func constUnitID() (id uint32) { return 0 } -func assertSuccessAndUnmarshal(t *testing.T, to any, handler http.Handler, req *http.Request) { - t.Helper() +func assertSuccessAndUnmarshal(tb testing.TB, to any, handler http.Handler, req *http.Request) { + tb.Helper() - require.NotNil(t, handler) + require.NotNil(tb, handler) rw := httptest.NewRecorder() handler.ServeHTTP(rw, req) - require.Equal(t, http.StatusOK, rw.Code) + require.Equal(tb, http.StatusOK, rw.Code) data := rw.Body.Bytes() if to == nil { - assert.Empty(t, data) + assert.Empty(tb, data) return } err := json.Unmarshal(data, to) - require.NoError(t, err) + require.NoError(tb, err) } func TestStats(t *testing.T) { diff --git a/internal/updater/updater.go b/internal/updater/updater.go index 93138230..22c44e31 100644 --- a/internal/updater/updater.go +++ b/internal/updater/updater.go @@ -145,19 +145,15 @@ func NewUpdater(conf *Config) *Updater { } } -// Update performs the auto-update. It returns an error if the update failed. +// Update performs the auto-update. It returns an error if the update fails. // If firstRun is true, it assumes the configuration file doesn't exist. func (u *Updater) Update(ctx context.Context, firstRun bool) (err error) { u.mu.Lock() defer u.mu.Unlock() - u.logger.InfoContext(ctx, "staring update", "first_run", firstRun) + u.logger.InfoContext(ctx, "starting update", "first_run", firstRun) defer func() { - if err != nil { - u.logger.ErrorContext(ctx, "update failed", slogutil.KeyError, err) - } else { - u.logger.InfoContext(ctx, "update finished") - } + u.logUpdateResult(ctx, err) }() err = u.prepare(ctx) @@ -197,6 +193,17 @@ func (u *Updater) Update(ctx context.Context, firstRun bool) (err error) { return nil } +// logUpdateResult logs the result of the update operation. +func (u *Updater) logUpdateResult(ctx context.Context, err error) { + if err != nil { + u.logger.ErrorContext(ctx, "update failed", slogutil.KeyError, err) + + return + } + + u.logger.InfoContext(ctx, "update finished") +} + // NewVersion returns the available new version. func (u *Updater) NewVersion() (nv string) { u.mu.RLock() diff --git a/openapi/CHANGELOG.md b/openapi/CHANGELOG.md index 20132c42..06d28e82 100644 --- a/openapi/CHANGELOG.md +++ b/openapi/CHANGELOG.md @@ -4,6 +4,10 @@ ## v0.108.0: API changes +## v0.107.64: API changes + +- The new field `"cache_enabled"` in `GET /control/dns_info` and `POST /control/dns_config`. Setting this flag to true turns the DNS-response cache on and requires a positive `cache_size` value (or a positive `dns.cache_size` in the configuration file). + ## v0.107.58: API changes ### The ability to check rules for query types and/or clients: GET /control/check_host diff --git a/openapi/openapi.yaml b/openapi/openapi.yaml index 77315d41..342e67bc 100644 --- a/openapi/openapi.yaml +++ b/openapi/openapi.yaml @@ -1559,6 +1559,14 @@ 'type': 'integer' 'cache_ttl_max': 'type': 'integer' + 'cache_enabled': + 'type': 'boolean' + 'description': | + Enables or disables the DNS response cache. + + If `cache_enabled` is `true`, the companion field `cache_size` must + be present and greater than 0, or the `dns.cache_size` setting in + the configuration file must already be greater than 0. 'cache_optimistic': 'type': 'boolean' 'upstream_mode': diff --git a/scripts/make/go-lint.sh b/scripts/make/go-lint.sh index 84a4e5b8..2e2448ac 100644 --- a/scripts/make/go-lint.sh +++ b/scripts/make/go-lint.sh @@ -3,7 +3,7 @@ # This comment is used to simplify checking local copies of the script. Bump # this number every time a significant change is made to this script. # -# AdGuard-Project-Version: 13 +# AdGuard-Project-Version: 14 verbose="${VERBOSE:-0}" readonly verbose @@ -26,11 +26,11 @@ set -f -u # Simple analyzers -# blocklist_imports is a simple check against unwanted packages. The following -# packages are banned: +# blocklist_imports is a simple best-effort check against unwanted packages. +# The following packages are banned: # # * Package errors is replaced by our own package in the -# github.com/AdguardTeam/golibs module. +# github.com/AdguardTeam/golibs module. # # * Packages log and github.com/AdguardTeam/golibs/log are replaced by # stdlib's new package log/slog and AdGuard's new utilities package @@ -67,7 +67,10 @@ set -f -u # # TODO(a.garipov): Add golibs/log. blocklist_imports() { - find . \ + import_or_tab="$(printf '^\\(import \\|\t\\)')" + readonly import_or_tab + + find_with_ignore \ -type 'f' \ -name '*.go' \ '!' '(' \ @@ -77,16 +80,16 @@ blocklist_imports() { -exec \ 'grep' \ '-H' \ - '-e' '[[:space:]]"errors"$' \ - '-e' '[[:space:]]"github.com/prometheus/client_golang/prometheus/promauto"$' \ - '-e' '[[:space:]]"golang.org/x/exp/maps"$' \ - '-e' '[[:space:]]"golang.org/x/exp/slices"$' \ - '-e' '[[:space:]]"golang.org/x/net/context"$' \ - '-e' '[[:space:]]"io/ioutil"$' \ - '-e' '[[:space:]]"log"$' \ - '-e' '[[:space:]]"reflect"$' \ - '-e' '[[:space:]]"sort"$' \ - '-e' '[[:space:]]"unsafe"$' \ + '-e' "$import_or_tab"'"errors"$' \ + '-e' "$import_or_tab"'"github.com/prometheus/client_golang/prometheus/promauto"$' \ + '-e' "$import_or_tab"'"golang.org/x/exp/maps"$' \ + '-e' "$import_or_tab"'"golang.org/x/exp/slices"$' \ + '-e' "$import_or_tab"'"golang.org/x/net/context"$' \ + '-e' "$import_or_tab"'"io/ioutil"$' \ + '-e' "$import_or_tab"'"log"$' \ + '-e' "$import_or_tab"'"reflect"$' \ + '-e' "$import_or_tab"'"sort"$' \ + '-e' "$import_or_tab"'"unsafe"$' \ '-n' \ '{}' \ ';' @@ -94,8 +97,11 @@ blocklist_imports() { # method_const is a simple check against the usage of some raw strings and # numbers where one should use named constants. +# +# NOTE: Flag -H for grep is non-POSIX but all of Busybox, GNU, macOS, and +# OpenBSD support it. method_const() { - find . \ + find_with_ignore \ -type 'f' \ -name '*.go' \ -exec \ @@ -116,10 +122,11 @@ method_const() { # use of filenames like client_manager.go. underscores() { underscore_files="$( - find . \ + find_with_ignore \ -type 'f' \ -name '*_*.go' \ - '!' '(' -name '*_bsd.go' \ + '!' '(' \ + -name '*_bsd.go' \ -o -name '*_darwin.go' \ -o -name '*_freebsd.go' \ -o -name '*_generate.go' \ @@ -169,15 +176,6 @@ run_linter gocognit --over='19' \ ./internal/home/ \ ; -run_linter gocognit --over='18' \ - ./internal/aghtls/ \ - ; - -run_linter gocognit --over='15' \ - ./internal/aghos/ \ - ./internal/filtering/ \ - ; - run_linter gocognit --over='14' \ ./internal/dhcpd \ ; @@ -186,33 +184,26 @@ run_linter gocognit --over='13' \ ./internal/aghnet/ \ ; -run_linter gocognit --over='12' \ - ./internal/filtering/rewrite/ \ - ; - -run_linter gocognit --over='11' \ - ./internal/updater/ \ - ; - run_linter gocognit --over='10' \ ./internal/aghalg/ \ ./internal/aghhttp/ \ + ./internal/aghos/ \ ./internal/aghrenameio/ \ ./internal/aghtest/ \ + ./internal/aghtls/ \ ./internal/aghuser/ \ ./internal/arpdb/ \ ./internal/client/ \ ./internal/configmigrate/ \ ./internal/dhcpsvc \ ./internal/dnsforward/ \ - ./internal/filtering/hashprefix/ \ - ./internal/filtering/rulelist/ \ - ./internal/filtering/safesearch/ \ + ./internal/filtering/ \ ./internal/ipset \ ./internal/next/ \ ./internal/rdns/ \ ./internal/schedule/ \ ./internal/stats/ \ + ./internal/updater/ \ ./internal/version/ \ ./internal/whois/ \ ./scripts/ \ @@ -222,13 +213,7 @@ run_linter ineffassign ./... run_linter unparam ./... -find . \ - '(' \ - -name 'node_modules' \ - -type 'd' \ - -prune \ - ')' \ - -o \ +find_with_ignore \ -type 'f' \ '(' \ -name 'Makefile' \ diff --git a/scripts/make/helper.sh b/scripts/make/helper.sh index 8caa0477..f2ccb57a 100644 --- a/scripts/make/helper.sh +++ b/scripts/make/helper.sh @@ -8,13 +8,13 @@ # This comment is used to simplify checking local copies of the script. Bump # this number every time a remarkable change is made to this script. # -# AdGuard-Project-Version: 4 +# AdGuard-Project-Version: 5 # Deferred helpers not_found_msg=' looks like a binary not found error. -make sure you have installed the linter binaries using: +make sure you have installed the linter binaries, including using: $ make go-tools ' @@ -73,3 +73,41 @@ run_linter() ( return "$exitcode" ) + +# find_with_ignore is a wrapper around find that does not descend into ignored +# directories, such as ./tmp/. +# +# NOTE: The arguments must contain on of -exec, -ok, or -print; see +# https://pubs.opengroup.org/onlinepubs/9799919799/utilities/find.html. +# +# TODO(a.garipov): Find a way to integrate the entire gitignore, including the +# global one, without using git, as .git is not copied into the build container. +# +# Keep in sync with .gitignore. +find_with_ignore() { + find . \ + '(' \ + -type 'd' \ + '(' \ + -name '.git' \ + -o -path '/agh-backup' \ + -o -path './bin' \ + -o -path './build' \ + -o -path './client/blob-report' \ + -o -path './client/playwright-report' \ + -o -path './client/playwright/.cache' \ + -o -path './client/test-results' \ + -o -path './data' \ + -o -path './dist' \ + -o -path './launchpad_credentials' \ + -o -path './snapcraft_login' \ + -o -name 'node_modules' \ + -o -name 'test-reports' \ + -o -name 'tmp' \ + ')' \ + -prune \ + ')' \ + -o \ + "$@" \ + ; +} diff --git a/scripts/make/txt-lint.sh b/scripts/make/txt-lint.sh index ed4bb327..7df3bf1a 100644 --- a/scripts/make/txt-lint.sh +++ b/scripts/make/txt-lint.sh @@ -3,7 +3,7 @@ # This comment is used to simplify checking local copies of the script. Bump # this number every time a remarkable change is made to this script. # -# AdGuard-Project-Version: 8 +# AdGuard-Project-Version: 9 verbose="${VERBOSE:-0}" readonly verbose @@ -33,19 +33,7 @@ trailing_newlines() ( nl="$(printf '\n')" readonly nl - find . \ - '(' \ - -type 'd' \ - '(' \ - -name 'node_modules' \ - -o -path './.git' \ - -o -path './bin' \ - -o -path './build' \ - -o -path './client/playwright-report' \ - ')' \ - -prune \ - ')' \ - -o \ + find_with_ignore \ -type 'f' \ '!' '(' \ -name '*.db' \ @@ -71,7 +59,7 @@ trailing_newlines() ( # trailing_whitespace is a simple check that makes sure that there are no # trailing whitespace in plain-text files. trailing_whitespace() { - find . \ + find_with_ignore \ -type 'f' \ '!' '(' \ -name '*.db' \ @@ -84,11 +72,8 @@ trailing_whitespace() { -o -name '*.zip' \ -o -name 'AdGuardHome' \ -o -name 'adguard-home' \ - -o -path '*/node_modules/*' \ - -o -path './.git/*' \ - -o -path './bin/*' \ - -o -path './build/*' \ ')' \ + -print \ | while read -r f; do grep -e '[[:space:]]$' -n -- "$f" \ | sed -e "s:^:${f}\::" -e 's/ \+$/>>>&<<