all: sync with master

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

View File

@@ -1,7 +1,7 @@
'name': 'build'
'env':
'GO_VERSION': '1.24.5'
'GO_VERSION': '1.24.6'
'NODE_VERSION': '20'
'on':

View File

@@ -1,7 +1,7 @@
'name': 'lint'
'env':
'GO_VERSION': '1.24.5'
'GO_VERSION': '1.24.6'
'on':
'push':

2
.gitignore vendored
View File

@@ -37,5 +37,7 @@ AdGuardHome.exe
AdGuardHome.yaml*
coverage.txt
node_modules/
test-reports/
tmp/
!/build/gitkeep

View File

@@ -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.65...HEAD
[v0.107.65]: https://github.com/AdguardTeam/AdGuardHome/compare/v0.107.64...v0.107.65
[Unreleased]: https://github.com/AdguardTeam/AdGuardHome/compare/v0.107.66...HEAD
[v0.107.66]: https://github.com/AdguardTeam/AdGuardHome/compare/v0.107.65...v0.107.66
-->
[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

View File

@@ -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

View File

@@ -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'

View File

@@ -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'

View File

@@ -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, у байтах"

View File

@@ -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",

View File

@@ -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",

View File

@@ -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",

View File

@@ -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",

View File

@@ -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</0>.",
"topline_expired_certificate": "Tu certificado SSL ha expirado. Actualiza la <0>configuración de cifrado</0>.",
"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)",

View File

@@ -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",

View File

@@ -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",

View File

@@ -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の上書き秒単位",

View File

@@ -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 (초) 무시",

View File

@@ -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",

View File

@@ -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",

View File

@@ -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",

View File

@@ -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",

View File

@@ -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",

View File

@@ -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",

View File

@@ -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 值",

View File

@@ -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",

View File

@@ -337,7 +337,8 @@
}
.button-action--arrow-option:disabled {
display: none;
opacity: 0.5;
cursor: default;
}
.tooltip-custom__container .button-action--arrow-option {

View File

@@ -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<FormData>({
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 (
<form onSubmit={handleSubmit(onSubmit)}>
<div className="row">
<div className="col-12 col-md-7">
<div className="form__group form__group--settings">
<Controller
name="cache_enabled"
control={control}
render={({ field }) => (
<Checkbox
{...field}
data-testid="dns_cache_enabled"
title={t('cache_enabled')}
subtitle={t('cache_enabled_desc')}
disabled={processingSetConfig}
/>
)}
/>
</div>
</div>
{INPUTS_FIELDS.map(({ name, title, description, placeholder }) => (
<div className="col-12" key={name}>
<div className="col-12 col-md-7 p-0">
@@ -102,6 +125,12 @@ const Form = ({ initialValues, onSubmit }: CacheFormProps) => {
setValueAs: (value) => replaceZeroWithEmptyString(value),
})}
/>
{name === CACHE_CONFIG_FIELDS.cache_size && cacheSizeZeroWhenEnabled && (
<span className="form__message form__message--error">
{t('cache_size_validation')}
</span>
)}
</div>
</div>
</div>
@@ -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')}
</button>

View File

@@ -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 = () => {
<div className="form">
<Form
initialValues={{
cache_enabled,
cache_size: replaceZeroWithEmptyString(cache_size),
cache_ttl_max: replaceZeroWithEmptyString(cache_ttl_max),
cache_ttl_min: replaceZeroWithEmptyString(cache_ttl_min),

View File

@@ -322,6 +322,7 @@ export type DnsConfigData = {
ratelimit_subnet_len_ipv6?: number;
edns_cs_use_custom?: boolean;
edns_cs_custom_ip?: string;
cache_enabled?: boolean;
cache_size?: number;
cache_ttl_max?: number;
cache_ttl_min?: number;

47
go.mod
View File

@@ -1,10 +1,11 @@
module github.com/AdguardTeam/AdGuardHome
go 1.24.5
go 1.24.6
require (
github.com/AdguardTeam/dnsproxy v0.76.1
github.com/AdguardTeam/golibs v0.32.15
// TODO(s.chzhen): Use osutil/executil and fakeos/fakeexec.
github.com/AdguardTeam/golibs v0.34.0
github.com/AdguardTeam/urlfilter v0.20.0
github.com/NYTimes/gziphandler v1.1.1
github.com/ameshkov/dnscrypt/v2 v2.4.0
@@ -32,21 +33,21 @@ require (
github.com/stretchr/testify v1.10.0
github.com/ti-mo/netfilter v0.5.3
go.etcd.io/bbolt v1.4.1
golang.org/x/crypto v0.39.0
golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b
golang.org/x/net v0.41.0
golang.org/x/sys v0.33.0
golang.org/x/crypto v0.41.0
golang.org/x/exp v0.0.0-20250813145105-42675adae3e6
golang.org/x/net v0.43.0
golang.org/x/sys v0.35.0
gopkg.in/natefinch/lumberjack.v2 v2.2.1
gopkg.in/yaml.v3 v3.0.1
howett.net/plist v1.0.1
)
require (
cloud.google.com/go v0.121.3 // indirect
cloud.google.com/go v0.121.5 // indirect
cloud.google.com/go/ai v0.12.1 // indirect
cloud.google.com/go/auth v0.16.2 // indirect
cloud.google.com/go/auth v0.16.4 // indirect
cloud.google.com/go/auth/oauth2adapt v0.2.8 // indirect
cloud.google.com/go/compute/metadata v0.7.0 // indirect
cloud.google.com/go/compute/metadata v0.8.0 // indirect
cloud.google.com/go/longrunning v0.6.7 // indirect
github.com/BurntSushi/toml v1.5.0 // indirect
github.com/ameshkov/dnsstamps v1.0.3 // indirect
@@ -61,7 +62,7 @@ require (
github.com/google/generative-ai-go v0.20.1 // indirect
github.com/google/s2a-go v0.1.9 // indirect
github.com/googleapis/enterprise-certificate-proxy v0.3.6 // indirect
github.com/googleapis/gax-go/v2 v2.14.2 // indirect
github.com/googleapis/gax-go/v2 v2.15.0 // indirect
github.com/gookit/color v1.5.4 // indirect
github.com/gordonklaus/ineffassign v0.1.0 // indirect
github.com/josharian/native v1.1.0 // indirect
@@ -75,7 +76,7 @@ require (
github.com/quic-go/qpack v0.5.1 // indirect
github.com/robfig/cron/v3 v3.0.1 // indirect
github.com/rogpeppe/go-internal v1.14.1 // indirect
github.com/securego/gosec/v2 v2.22.5 // indirect
github.com/securego/gosec/v2 v2.22.8 // indirect
github.com/u-root/uio v0.0.0-20240224005618-d2acac8f3701 // indirect
github.com/uudashr/gocognit v1.2.0 // indirect
github.com/xo/terminfo v0.0.0-20220910002029-abceb7e1c41e // indirect
@@ -86,22 +87,22 @@ require (
go.opentelemetry.io/otel/metric v1.37.0 // indirect
go.opentelemetry.io/otel/trace v1.37.0 // indirect
go.uber.org/mock v0.5.2 // indirect
golang.org/x/exp/typeparams v0.0.0-20250620022241-b7579e27df2b // indirect
golang.org/x/mod v0.25.0 // indirect
golang.org/x/exp/typeparams v0.0.0-20250813145105-42675adae3e6 // indirect
golang.org/x/mod v0.27.0 // indirect
golang.org/x/oauth2 v0.30.0 // indirect
golang.org/x/sync v0.15.0 // indirect
golang.org/x/telemetry v0.0.0-20250708141652-5a6bbb13955f // indirect
golang.org/x/term v0.32.0 // indirect
golang.org/x/text v0.26.0 // indirect
golang.org/x/sync v0.16.0 // indirect
golang.org/x/telemetry v0.0.0-20250813145757-41cd51e6ab6a // indirect
golang.org/x/term v0.34.0 // indirect
golang.org/x/text v0.28.0 // indirect
golang.org/x/time v0.12.0 // indirect
golang.org/x/tools v0.34.0 // indirect
golang.org/x/tools v0.36.0 // indirect
golang.org/x/vuln v1.1.4 // indirect
gonum.org/v1/gonum v0.16.0 // indirect
google.golang.org/api v0.240.0 // indirect
google.golang.org/genproto/googleapis/api v0.0.0-20250707201910-8d1bb00bc6a7 // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20250707201910-8d1bb00bc6a7 // indirect
google.golang.org/grpc v1.73.0 // indirect
google.golang.org/protobuf v1.36.6 // indirect
google.golang.org/api v0.247.0 // indirect
google.golang.org/genproto/googleapis/api v0.0.0-20250811230008-5f3141c8851a // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20250811230008-5f3141c8851a // indirect
google.golang.org/grpc v1.74.2 // indirect
google.golang.org/protobuf v1.36.7 // indirect
honnef.co/go/tools v0.6.1 // indirect
mvdan.cc/editorconfig v0.3.0 // indirect
mvdan.cc/gofumpt v0.8.0 // indirect

100
go.sum
View File

@@ -1,19 +1,19 @@
cloud.google.com/go v0.121.3 h1:84RD+hQXNdY5Sw/MWVAx5O9Aui/rd5VQ9HEcdN19afo=
cloud.google.com/go v0.121.3/go.mod h1:6vWF3nJWRrEUv26mMB3FEIU/o1MQNVPG1iHdisa2SJc=
cloud.google.com/go v0.121.5 h1:KU9tFP5NeZiVDSWcsgjJ2P/HosAlD4fCGBimBlGiNXA=
cloud.google.com/go v0.121.5/go.mod h1:coChdst4Ea5vUpiALcYKXEpR1S9ZgXbhEzzMcMR66vI=
cloud.google.com/go/ai v0.12.1 h1:m1n/VjUuHS+pEO/2R4/VbuuEIkgk0w67fDQvFaMngM0=
cloud.google.com/go/ai v0.12.1/go.mod h1:5vIPNe1ZQsVZqCliXIPL4QnhObQQY4d9hAGHdVc4iw4=
cloud.google.com/go/auth v0.16.2 h1:QvBAGFPLrDeoiNjyfVunhQ10HKNYuOwZ5noee0M5df4=
cloud.google.com/go/auth v0.16.2/go.mod h1:sRBas2Y1fB1vZTdurouM0AzuYQBMZinrUYL8EufhtEA=
cloud.google.com/go/auth v0.16.4 h1:fXOAIQmkApVvcIn7Pc2+5J8QTMVbUGLscnSVNl11su8=
cloud.google.com/go/auth v0.16.4/go.mod h1:j10ncYwjX/g3cdX7GpEzsdM+d+ZNsXAbb6qXA7p1Y5M=
cloud.google.com/go/auth/oauth2adapt v0.2.8 h1:keo8NaayQZ6wimpNSmW5OPc283g65QNIiLpZnkHRbnc=
cloud.google.com/go/auth/oauth2adapt v0.2.8/go.mod h1:XQ9y31RkqZCcwJWNSx2Xvric3RrU88hAYYbjDWYDL+c=
cloud.google.com/go/compute/metadata v0.7.0 h1:PBWF+iiAerVNe8UCHxdOt6eHLVc3ydFeOCw78U8ytSU=
cloud.google.com/go/compute/metadata v0.7.0/go.mod h1:j5MvL9PprKL39t166CoB1uVHfQMs4tFQZZcKwksXUjo=
cloud.google.com/go/compute/metadata v0.8.0 h1:HxMRIbao8w17ZX6wBnjhcDkW6lTFpgcaobyVfZWqRLA=
cloud.google.com/go/compute/metadata v0.8.0/go.mod h1:sYOGTp851OV9bOFJ9CH7elVvyzopvWQFNNghtDQ/Biw=
cloud.google.com/go/longrunning v0.6.7 h1:IGtfDWHhQCgCjwQjV9iiLnUta9LBCo8R9QmAFsS/PrE=
cloud.google.com/go/longrunning v0.6.7/go.mod h1:EAFV3IZAKmM56TyiE6VAP3VoTzhZzySwI/YI1s/nRsY=
github.com/AdguardTeam/dnsproxy v0.76.1 h1:ms5vgdbYYXrKGPEpMFqUeql2j3aSfK1tGbCKju9rUgM=
github.com/AdguardTeam/dnsproxy v0.76.1/go.mod h1:9Mw3wQMTYwM/HR9FdtatQAd+m0S8mbwq2J+UZiy/gXc=
github.com/AdguardTeam/golibs v0.32.15 h1:arDRDWiZCH3g5Onr8AqMnOHhaOppNoBpgC3DNhmeDeA=
github.com/AdguardTeam/golibs v0.32.15/go.mod h1:G9CzUOzx87J+2u+eClJrrwWD7lMbROvuUnT8uvDUzIA=
github.com/AdguardTeam/golibs v0.34.0 h1:JQK024DkTYxE7vsPVsYsoyDHW/53Nun7OYb9qscniK8=
github.com/AdguardTeam/golibs v0.34.0/go.mod h1:K4C2EbfSEM1zY5YXoti9SfbTAHN/kIX97LpDtCwORrM=
github.com/AdguardTeam/urlfilter v0.20.0 h1:X32qiuVCVd8WDYCEsbdZKfXMzwdVqrdulamtUi4rmzs=
github.com/AdguardTeam/urlfilter v0.20.0/go.mod h1:gjrywLTxfJh6JOkwi9SU+frhP7kVVEZ5exFGkR99qpk=
github.com/BurntSushi/toml v1.5.0 h1:W5quZX/G/csjUnuI8SUYlsHs9M38FC7znL0lIO+DvMg=
@@ -85,8 +85,8 @@ github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
github.com/googleapis/enterprise-certificate-proxy v0.3.6 h1:GW/XbdyBFQ8Qe+YAmFU9uHLo7OnF5tL52HFAgMmyrf4=
github.com/googleapis/enterprise-certificate-proxy v0.3.6/go.mod h1:MkHOF77EYAE7qfSuSS9PU6g4Nt4e11cnsDUowfwewLA=
github.com/googleapis/gax-go/v2 v2.14.2 h1:eBLnkZ9635krYIPD+ag1USrOAI0Nr0QYF3+/3GqO0k0=
github.com/googleapis/gax-go/v2 v2.14.2/go.mod h1:ON64QhlJkhVtSqp4v1uaK92VyZ2gmvDQsweuyLV+8+w=
github.com/googleapis/gax-go/v2 v2.15.0 h1:SyjDc1mGgZU5LncH8gimWo9lW1DtIfPibOG81vgd/bo=
github.com/googleapis/gax-go/v2 v2.15.0/go.mod h1:zVVkkxAQHa1RQpg9z2AUCMnKhi0Qld9rcmyfL1OZhoc=
github.com/gookit/color v1.5.4 h1:FZmqs7XOyGgCAxmWyPslpiok1k05wmY3SJTytgvYFs0=
github.com/gookit/color v1.5.4/go.mod h1:pZJOeOS8DM43rXbp4AZo1n9zCU2qjpcRko0b6/QJi9w=
github.com/gordonklaus/ineffassign v0.1.0 h1:y2Gd/9I7MdY1oEIt+n+rowjBNDcLQq3RsH5hwJd0f9s=
@@ -128,8 +128,8 @@ github.com/miekg/dns v1.1.66 h1:FeZXOS3VCVsKnEAd+wBkjMC3D2K+ww66Cq3VnCINuJE=
github.com/miekg/dns v1.1.66/go.mod h1:jGFzBsSNbJw6z1HYut1RKBKHA9PBdxeHrZG8J+gC2WE=
github.com/onsi/ginkgo/v2 v2.23.4 h1:ktYTpKJAVZnDT4VjxSbiBenUjmlL/5QkBEocaWXiQus=
github.com/onsi/ginkgo/v2 v2.23.4/go.mod h1:Bt66ApGPBFzHyR+JO10Zbt0Gsp4uWxu5mIOTusL46e8=
github.com/onsi/gomega v1.37.0 h1:CdEG8g0S133B4OswTDC/5XPSzE1OeP29QOioj2PID2Y=
github.com/onsi/gomega v1.37.0/go.mod h1:8D9+Txp43QWKhM24yyOBEdpkzN8FvJyAwecBgsU4KU0=
github.com/onsi/gomega v1.38.0 h1:c/WX+w8SLAinvuKKQFh77WEucCnPk4j2OTUr7lt7BeY=
github.com/onsi/gomega v1.38.0/go.mod h1:OcXcwId0b9QsE7Y49u+BTrL4IdKOBOKnD6VQNTJEB6o=
github.com/patrickmn/go-cache v2.1.0+incompatible h1:HRMgzkcYKYpi3C8ajMPV8OFXaaRUnok+kx1WdO15EQc=
github.com/patrickmn/go-cache v2.1.0+incompatible/go.mod h1:3Qf8kWWT7OJRJbdiICTKqZju1ZixQ/KpMGzzAfe6+WQ=
github.com/pierrec/lz4/v4 v4.1.22 h1:cKFw6uJDK+/gfw5BcDL0JL5aBsAFdsIT18eRtLj7VIU=
@@ -149,8 +149,8 @@ github.com/robfig/cron/v3 v3.0.1 h1:WdRxkvbJztn8LMz/QEvLN5sBU+xKpSqwwUO1Pjr4qDs=
github.com/robfig/cron/v3 v3.0.1/go.mod h1:eQICP3HwyT7UooqI/z+Ov+PtYAWygg1TEWWzGIFLtro=
github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ=
github.com/rogpeppe/go-internal v1.14.1/go.mod h1:MaRKkUm5W0goXpeCfT7UZI6fk/L7L7so1lCWt35ZSgc=
github.com/securego/gosec/v2 v2.22.5 h1:ySws9uwOeE42DsG54v2moaJfh7r08Ev7SAYJuoMDfRA=
github.com/securego/gosec/v2 v2.22.5/go.mod h1:AWfgrFsVewk5LKobsPWlygCHt8K91boVPyL6GUZG5NY=
github.com/securego/gosec/v2 v2.22.8 h1:3NMpmfXO8wAVFZPNsd3EscOTa32Jyo6FLLlW53bexMI=
github.com/securego/gosec/v2 v2.22.8/go.mod h1:ZAw8K2ikuH9qDlfdV87JmNghnVfKB1XC7+TVzk6Utto=
github.com/shirou/gopsutil/v3 v3.24.5 h1:i0t8kL+kQTvpAYToeuiVk3TgDeKOFioZO3Ztz/iZ9pI=
github.com/shirou/gopsutil/v3 v3.24.5/go.mod h1:bsoOS1aStSs9ErQ1WWfxllSeS1K5D+U30r2NfcubMVk=
github.com/shoenig/go-m1cpu v0.1.6 h1:nxdKQNcEB6vzgA2E2bvzKIYRuNj7XNJ4S/aRSwKzFtM=
@@ -203,17 +203,17 @@ go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko=
go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o=
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
golang.org/x/crypto v0.39.0 h1:SHs+kF4LP+f+p14esP5jAoDpHU8Gu/v9lFRK6IT5imM=
golang.org/x/crypto v0.39.0/go.mod h1:L+Xg3Wf6HoL4Bn4238Z6ft6KfEpN0tJGo53AAPC632U=
golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b h1:M2rDM6z3Fhozi9O7NWsxAkg/yqS/lQJ6PmkyIV3YP+o=
golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b/go.mod h1:3//PLf8L/X+8b4vuAfHzxeRUl04Adcb341+IGKfnqS8=
golang.org/x/exp/typeparams v0.0.0-20250620022241-b7579e27df2b h1:KdrhdYPDUvJTvrDK9gdjfFd6JTk8vA1WJoldYSi0kHo=
golang.org/x/exp/typeparams v0.0.0-20250620022241-b7579e27df2b/go.mod h1:LKZHyeOpPuZcMgxeHjJp4p5yvxrCX1xDvH10zYHhjjQ=
golang.org/x/crypto v0.41.0 h1:WKYxWedPGCTVVl5+WHSSrOBT0O8lx32+zxmHxijgXp4=
golang.org/x/crypto v0.41.0/go.mod h1:pO5AFd7FA68rFak7rOAGVuygIISepHftHnr8dr6+sUc=
golang.org/x/exp v0.0.0-20250813145105-42675adae3e6 h1:SbTAbRFnd5kjQXbczszQ0hdk3ctwYf3qBNH9jIsGclE=
golang.org/x/exp v0.0.0-20250813145105-42675adae3e6/go.mod h1:4QTo5u+SEIbbKW1RacMZq1YEfOBqeXa19JeshGi+zc4=
golang.org/x/exp/typeparams v0.0.0-20250813145105-42675adae3e6 h1:zAfbUfwhYzU4abt1UJxExk7WVbeTsYEnOPySl6RVucI=
golang.org/x/exp/typeparams v0.0.0-20250813145105-42675adae3e6/go.mod h1:4Mzdyp/6jzw9auFDJ3OMF5qksa7UvPnzKqTVGcb04ms=
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
golang.org/x/mod v0.1.1-0.20191105210325-c90efee705ee/go.mod h1:QqPTAvyqsEbceGzBzNggFXnrqF1CaUcvgkdR5Ot7KZg=
golang.org/x/mod v0.4.2/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
golang.org/x/mod v0.25.0 h1:n7a+ZbQKQA/Ysbyb0/6IbB1H/X41mKgbhfv7AfG/44w=
golang.org/x/mod v0.25.0/go.mod h1:IXM97Txy2VM4PJ3gI61r1YEk/gAj6zAHN3AdZt6S9Ww=
golang.org/x/mod v0.27.0 h1:kb+q2PyFnEADO2IEF935ehFUXlWiNjJWtRNgBLSfbxQ=
golang.org/x/mod v0.27.0/go.mod h1:rWI627Fq0DEoudcK+MBkNkCe0EetEaDSwJJkCcjpazc=
golang.org/x/net v0.0.0-20190311183353-d8887717615a/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
golang.org/x/net v0.0.0-20190503192946-f4e77d36d62c/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
@@ -221,14 +221,14 @@ golang.org/x/net v0.0.0-20190603091049-60506f45cf65/go.mod h1:HSz+uSET+XFnRR8LxR
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
golang.org/x/net v0.0.0-20210316092652-d523dce5a7f4/go.mod h1:RBQZq4jEuRlivfhVLdyRGr576XBO4/greRjx4P4O3yc=
golang.org/x/net v0.0.0-20210405180319-a5a99cb37ef4/go.mod h1:p54w0d4576C0XHj96bSt6lcn1PtDYWL6XObtHCRCNQM=
golang.org/x/net v0.41.0 h1:vBTly1HeNPEn3wtREYfy4GZ/NECgw2Cnl+nK6Nz3uvw=
golang.org/x/net v0.41.0/go.mod h1:B/K4NNqkfmg07DQYrbwvSluqCJOOXwUjeb/5lOisjbA=
golang.org/x/net v0.43.0 h1:lat02VYK2j4aLzMzecihNvTlJNQUq316m2Mr9rnM6YE=
golang.org/x/net v0.43.0/go.mod h1:vhO1fvI4dGsIjh73sWfUVjj3N7CA9WkKJNQm2svM6Jg=
golang.org/x/oauth2 v0.30.0 h1:dnDm7JmhM45NNpd8FDDeLhK6FwqbOf4MLCM9zb1BOHI=
golang.org/x/oauth2 v0.30.0/go.mod h1:B++QgG3ZKulg6sRPGD/mqlHQs5rB3Ml9erfeDY7xKlU=
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.15.0 h1:KWH3jNZsfyT6xfAfKiz6MRNmd46ByHDYaZ7KSkCtdW8=
golang.org/x/sync v0.15.0/go.mod h1:1dzgHSNfp02xaA81J2MS99Qcpr2w7fw1gpm99rleRqA=
golang.org/x/sync v0.16.0 h1:ycBJEhp9p4vXvUZNszeOq0kGTPghopOL8q0fq3vstxw=
golang.org/x/sync v0.16.0/go.mod h1:1dzgHSNfp02xaA81J2MS99Qcpr2w7fw1gpm99rleRqA=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190312061237-fead79001313/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20190322080309-f49334f85ddc/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
@@ -240,25 +240,29 @@ golang.org/x/sys v0.0.0-20210330210617-4fbd30eecc44/go.mod h1:h1NjWce9XRLGQEsW7w
golang.org/x/sys v0.0.0-20210510120138-977fb7262007/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20210927094055-39ccf1dd6fa6/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20220209214540-3681064d5158/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.33.0 h1:q3i8TbbEz+JRD9ywIRlyRAQbM0qF7hu24q3teo2hbuw=
golang.org/x/sys v0.33.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k=
golang.org/x/telemetry v0.0.0-20250708141652-5a6bbb13955f h1:GnwFSf1cKD9qa+VRJWGjDBK0OHWJgTMaj49bSkN3agw=
golang.org/x/telemetry v0.0.0-20250708141652-5a6bbb13955f/go.mod h1:mUcjA5g0luJpMYCLjhH91f4t4RAUNp+zq9ZmUoqPD7M=
golang.org/x/sys v0.35.0 h1:vz1N37gP5bs89s7He8XuIYXpyY0+QlsKmzipCbUtyxI=
golang.org/x/sys v0.35.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k=
golang.org/x/telemetry v0.0.0-20250813145757-41cd51e6ab6a h1:60aHXv7uIB0ngdMlalCWv6pSC/giLMowErJBA/5xbPQ=
golang.org/x/telemetry v0.0.0-20250813145757-41cd51e6ab6a/go.mod h1:JIJwPkb04vX0KeIBbQ7epGtgIjA8ihHbsAtW4A/lIQ4=
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
golang.org/x/term v0.32.0 h1:DR4lr0TjUs3epypdhTOkMmuF5CDFJ/8pOnbzMZPQ7bg=
golang.org/x/term v0.32.0/go.mod h1:uZG1FhGx848Sqfsq4/DlJr3xGGsYMu/L5GW4abiaEPQ=
golang.org/x/term v0.34.0 h1:O/2T7POpk0ZZ7MAzMeWFSg6S5IpWd/RXDlM9hgM3DR4=
golang.org/x/term v0.34.0/go.mod h1:5jC53AEywhIVebHgPVeg0mj8OD3VO9OzclacVrqpaAw=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
golang.org/x/text v0.26.0 h1:P42AVeLghgTYr4+xUnTRKDMqpar+PtX7KWuNQL21L8M=
golang.org/x/text v0.26.0/go.mod h1:QK15LZJUUQVJxhz7wXgxSy/CJaTFjd0G+YLonydOVQA=
golang.org/x/text v0.28.0 h1:rhazDwis8INMIwQ4tpjLDzUhx6RlXqZNPEM0huQojng=
golang.org/x/text v0.28.0/go.mod h1:U8nCwOR8jO/marOQ0QbDiOngZVEBB7MAiitBuMjXiNU=
golang.org/x/time v0.12.0 h1:ScB/8o8olJvc+CQPWrK3fPZNfh7qgwCrY0zJmoEQLSE=
golang.org/x/time v0.12.0/go.mod h1:CDIdPxbZBQxdj6cxyCIdrNogrJKMJ7pr37NYpMcMDSg=
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo=
golang.org/x/tools v0.0.0-20200130002326-2f3ba24bd6e7/go.mod h1:TB2adYChydJhpapKDTa4BR/hXlZSLoq2Wpct/0txZ28=
golang.org/x/tools v0.1.5/go.mod h1:o0xws9oXOQQZyjljx8fwUC0k7L1pTE6eaCbjGeHmOkk=
golang.org/x/tools v0.34.0 h1:qIpSLOxeCYGg9TrcJokLBG4KFA6d795g0xkBkiESGlo=
golang.org/x/tools v0.34.0/go.mod h1:pAP9OwEaY1CAW3HOmg3hLZC5Z0CCmzjAF2UQMSqNARg=
golang.org/x/tools v0.36.0 h1:kWS0uv/zsvHEle1LbV5LE8QujrxB3wfQyxHfhOk0Qkg=
golang.org/x/tools v0.36.0/go.mod h1:WBDiHKJK8YgLHlcQPYQzNCkUxUypCaa5ZegCVutKm+s=
golang.org/x/tools/go/expect v0.1.1-deprecated h1:jpBZDwmgPhXsKZC6WhL20P4b/wmnpsEAGHaNy0n/rJM=
golang.org/x/tools/go/expect v0.1.1-deprecated/go.mod h1:eihoPOH+FgIqa3FpoTwguz/bVUSGBlGQU67vpBeOrBY=
golang.org/x/tools/go/packages/packagestest v0.1.1-deprecated h1:1h2MnaIAIXISqTFKdENegdpAgUXz6NrPEsbIeWaBRvM=
golang.org/x/tools/go/packages/packagestest v0.1.1-deprecated/go.mod h1:RVAQXBGNv1ib0J382/DPCRS/BPnsGebyM1Gj5VSDpG8=
golang.org/x/vuln v1.1.4 h1:Ju8QsuyhX3Hk8ma3CesTbO8vfJD9EvUBgHvkxHBzj0I=
golang.org/x/vuln v1.1.4/go.mod h1:F+45wmU18ym/ca5PLTPLsSzr2KppzswxPP603ldA67s=
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
@@ -267,18 +271,18 @@ golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8T
golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
gonum.org/v1/gonum v0.16.0 h1:5+ul4Swaf3ESvrOnidPp4GZbzf0mxVQpDCYUQE7OJfk=
gonum.org/v1/gonum v0.16.0/go.mod h1:fef3am4MQ93R2HHpKnLk4/Tbh/s0+wqD5nfa6Pnwy4E=
google.golang.org/api v0.240.0 h1:PxG3AA2UIqT1ofIzWV2COM3j3JagKTKSwy7L6RHNXNU=
google.golang.org/api v0.240.0/go.mod h1:cOVEm2TpdAGHL2z+UwyS+kmlGr3bVWQQ6sYEqkKje50=
google.golang.org/genproto v0.0.0-20250505200425-f936aa4a68b2 h1:1tXaIXCracvtsRxSBsYDiSBN0cuJvM7QYW+MrpIRY78=
google.golang.org/genproto v0.0.0-20250505200425-f936aa4a68b2/go.mod h1:49MsLSx0oWMOZqcpB3uL8ZOkAh1+TndpJ8ONoCBWiZk=
google.golang.org/genproto/googleapis/api v0.0.0-20250707201910-8d1bb00bc6a7 h1:FiusG7LWj+4byqhbvmB+Q93B/mOxJLN2DTozDuZm4EU=
google.golang.org/genproto/googleapis/api v0.0.0-20250707201910-8d1bb00bc6a7/go.mod h1:kXqgZtrWaf6qS3jZOCnCH7WYfrvFjkC51bM8fz3RsCA=
google.golang.org/genproto/googleapis/rpc v0.0.0-20250707201910-8d1bb00bc6a7 h1:pFyd6EwwL2TqFf8emdthzeX+gZE1ElRq3iM8pui4KBY=
google.golang.org/genproto/googleapis/rpc v0.0.0-20250707201910-8d1bb00bc6a7/go.mod h1:qQ0YXyHHx3XkvlzUtpXDkS29lDSafHMZBAZDc03LQ3A=
google.golang.org/grpc v1.73.0 h1:VIWSmpI2MegBtTuFt5/JWy2oXxtjJ/e89Z70ImfD2ok=
google.golang.org/grpc v1.73.0/go.mod h1:50sbHOUqWoCQGI8V2HQLJM0B+LMlIUjNSZmow7EVBQc=
google.golang.org/protobuf v1.36.6 h1:z1NpPI8ku2WgiWnf+t9wTPsn6eP1L7ksHUlkfLvd9xY=
google.golang.org/protobuf v1.36.6/go.mod h1:jduwjTPXsFjZGTmRluh+L6NjiWu7pchiJ2/5YcXBHnY=
google.golang.org/api v0.247.0 h1:tSd/e0QrUlLsrwMKmkbQhYVa109qIintOls2Wh6bngc=
google.golang.org/api v0.247.0/go.mod h1:r1qZOPmxXffXg6xS5uhx16Fa/UFY8QU/K4bfKrnvovM=
google.golang.org/genproto v0.0.0-20250603155806-513f23925822 h1:rHWScKit0gvAPuOnu87KpaYtjK5zBMLcULh7gxkCXu4=
google.golang.org/genproto v0.0.0-20250603155806-513f23925822/go.mod h1:HubltRL7rMh0LfnQPkMH4NPDFEWp0jw3vixw7jEM53s=
google.golang.org/genproto/googleapis/api v0.0.0-20250811230008-5f3141c8851a h1:DMCgtIAIQGZqJXMVzJF4MV8BlWoJh2ZuFiRdAleyr58=
google.golang.org/genproto/googleapis/api v0.0.0-20250811230008-5f3141c8851a/go.mod h1:y2yVLIE/CSMCPXaHnSKXxu1spLPnglFLegmgdY23uuE=
google.golang.org/genproto/googleapis/rpc v0.0.0-20250811230008-5f3141c8851a h1:tPE/Kp+x9dMSwUm/uM0JKK0IfdiJkwAbSMSeZBXXJXc=
google.golang.org/genproto/googleapis/rpc v0.0.0-20250811230008-5f3141c8851a/go.mod h1:gw1tLEfykwDz2ET4a12jcXt4couGAm7IwsVaTy0Sflo=
google.golang.org/grpc v1.74.2 h1:WoosgB65DlWVC9FqI82dGsZhWFNBSLjQ84bjROOpMu4=
google.golang.org/grpc v1.74.2/go.mod h1:CtQ+BGjaAIXHs/5YS3i473GqwBBa1zGQNevxdeBEXrM=
google.golang.org/protobuf v1.36.7 h1:IgrO7UwFQGJdRNXH/sQux4R1Dj1WAKcLElzeeRaXV2A=
google.golang.org/protobuf v1.36.7/go.mod h1:jduwjTPXsFjZGTmRluh+L6NjiWu7pchiJ2/5YcXBHnY=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk=
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q=

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

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

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