From b88dda327342c6752c60e8b9583629f0b26bdf4b Mon Sep 17 00:00:00 2001 From: Mygod Date: Thu, 16 Jul 2026 17:36:26 -0400 Subject: [PATCH 1/2] Support fallback traffic handling --- README.MD | 3 +- README_ES.MD | 3 +- README_FA.MD | 3 +- README_IT.MD | 3 +- README_RU.MD | 3 +- README_ZH.MD | 3 +- internal/config/server.go | 28 + internal/config/server_test.go | 60 ++ internal/dnsparser/parser.go | 109 +++- internal/dnsparser/parser_lite_test.go | 136 ++++ internal/dnsparser/response.go | 37 +- internal/udpserver/server.go | 119 +++- internal/udpserver/server_fallback_test.go | 490 ++++++++++++++ internal/udpserver/server_ingress.go | 14 +- internal/udpserver/server_ingress_test.go | 23 + internal/udpserver/server_runtime.go | 86 ++- internal/udpserver/udp_fallback.go | 603 ++++++++++++++++++ internal/udpserver/udp_fallback_test.go | 703 +++++++++++++++++++++ server_config.toml.simple | 6 + 19 files changed, 2340 insertions(+), 92 deletions(-) create mode 100644 internal/udpserver/server_fallback_test.go create mode 100644 internal/udpserver/udp_fallback.go create mode 100644 internal/udpserver/udp_fallback_test.go diff --git a/README.MD b/README.MD index 9b65229f..abec256a 100644 --- a/README.MD +++ b/README.MD @@ -716,11 +716,12 @@ These settings are critical on the server: | :--- | :--- | :--- | :--- | | `UDP_HOST` | `"0.0.0.0"` | if empty, this value is used | Address where the DNS server binds.
`0.0.0.0` means listen on all interfaces. | | `UDP_PORT` | `53` | `1..65535` | UDP port used by the server.
In most deployments this should remain `53` so resolvers can query it directly. | +| `FALLBACK` | omitted | empty/off, or `HOST:PORT`; bracket IPv6 addresses | Optional raw UDP proxy for non-DNS datagrams received on the DNS listener. Routing is tracked per source, and only a structurally complete DNS datagram establishes or resets DNS classification. An unknown non-DNS source is forwarded immediately and remains sticky to fallback, while a DNS-classified source switches after 16 consecutive non-DNS packets. While the source remains DNS-classified, every packet retained on the DNS path refreshes its 180-second idle timer; DNS resets the non-DNS streak and non-DNS increments it. After a gap longer than 180 seconds, the next non-DNS packet is treated as a new source and forwarded immediately. Active fallback sessions also expire after 180 seconds without traffic in either direction. Fallback mode uses one ingress reader to preserve packet order and a full 65,535-byte receive buffer; startup rejects a target that resolves back to the listener. Sessions have no hard cap, so untrusted or spoofed traffic can consume file descriptors and CPU; use network filtering or rate limiting on public listeners. | | `UDP_READERS` | `4` | auto-default if `<=0` | Number of goroutines reading directly from the UDP socket.
A larger number may help on very busy servers, but beyond a point it only increases context switching. | | `DNS_REQUEST_WORKERS` | `8` | auto-default if `<=0` | Number of workers that take requests from the front-door queue and pass them into the session/decode layer. | | `MAX_CONCURRENT_REQUESTS` | `16384` | fallback if `<=0` | Capacity of the incoming request queue.
If this queue fills up, packets are dropped and the server emits rate-limited overload logs. | | `SOCKET_BUFFER_SIZE` | `4194304` | fallback if `<=0` | Operating-system socket buffer size request for the UDP listener.
This matters for heavy bursts of incoming traffic. | -| `MAX_PACKET_SIZE` | `65535` | fallback if `<=0` | Size of the largest packet buffer that the packet pool allocates. | +| `MAX_PACKET_SIZE` | `65535` | fallback if `<=0`; at least `65535` with `FALLBACK` | Size of the largest packet buffer that the packet pool allocates. Fallback forces a full-size UDP receive buffer so raw datagrams are not truncated. | | `DROP_LOG_INTERVAL_SECONDS` | `2.0` | fallback if `<=0` | Minimum interval between repeated overload/drop logs, to avoid log spam during pressure. | ### 3.5.3) 🧠 Deferred Session Runtime diff --git a/README_ES.MD b/README_ES.MD index 5de4f2a8..70a1c9fd 100644 --- a/README_ES.MD +++ b/README_ES.MD @@ -718,11 +718,12 @@ Estos ajustes son críticos en el servidor: | :--- | :--- | :--- | :--- | | `UDP_HOST` | `"0.0.0.0"` | si está vacío, se usa este valor | Dirección donde se vincula el servidor DNS.
`0.0.0.0` significa escuchar en todas las interfaces. | | `UDP_PORT` | `53` | `1..65535` | Puerto UDP usado por el servidor.
En la mayoría de los despliegues debe permanecer en `53` para que los resolutores puedan consultarlo directamente. | +| `FALLBACK` | omitido | vacío/desactivado, o `HOST:PORT`; IPv6 entre corchetes | Proxy UDP sin procesar opcional para datagramas no DNS recibidos por el listener DNS. El enrutamiento se mantiene por origen y solo un datagrama DNS estructuralmente completo establece o reinicia la clasificación DNS. Un origen no DNS desconocido se reenvía de inmediato y queda asociado al fallback, mientras que un origen clasificado como DNS cambia después de 16 paquetes no DNS consecutivos. Mientras el origen siga clasificado como DNS, cada paquete retenido en la ruta DNS actualiza su temporizador de inactividad de 180 segundos; un paquete DNS reinicia la racha no DNS y uno no DNS la incrementa. Tras una pausa superior a 180 segundos, el siguiente paquete no DNS se trata como procedente de un origen nuevo y se reenvía de inmediato. Las sesiones fallback activas también vencen tras 180 segundos sin tráfico en ninguna dirección. El modo fallback usa un único lector de entrada para conservar el orden y un búfer completo de 65.535 bytes; el inicio rechaza un destino que resuelva al propio listener. Las sesiones no tienen un límite estricto, por lo que el tráfico no confiable o falsificado puede consumir descriptores de archivo y CPU; use filtrado de red o limitación de tasa en listeners públicos. | | `UDP_READERS` | `4` | valor automático por defecto si es `<=0` | Número de goroutines que leen directamente del socket UDP.
Un número mayor puede ayudar en servidores muy concurridos, pero a partir de cierto punto solo aumenta el cambio de contexto. | | `DNS_REQUEST_WORKERS` | `8` | valor automático por defecto si es `<=0` | Número de workers que toman solicitudes de la cola de entrada y las pasan a la capa de sesión/decodificación. | | `MAX_CONCURRENT_REQUESTS` | `16384` | valor de reserva si es `<=0` | Capacidad de la cola de solicitudes entrantes.
Si esta cola se llena, los paquetes se descartan y el servidor emite registros de sobrecarga limitados por tasa. | | `SOCKET_BUFFER_SIZE` | `4194304` | valor de reserva si es `<=0` | Tamaño de búfer de socket del sistema operativo solicitado para el escuchador UDP.
Esto importa para las ráfagas intensas de tráfico entrante. | -| `MAX_PACKET_SIZE` | `65535` | valor de reserva si es `<=0` | Tamaño del búfer de paquete más grande que asigna el pool de paquetes. | +| `MAX_PACKET_SIZE` | `65535` | valor de reserva si es `<=0`; al menos `65535` con `FALLBACK` | Tamaño del búfer de paquete más grande que asigna el pool de paquetes. Fallback fuerza un búfer UDP de tamaño completo para no truncar datagramas sin procesar. | | `DROP_LOG_INTERVAL_SECONDS` | `2.0` | valor de reserva si es `<=0` | Intervalo mínimo entre registros repetidos de sobrecarga/descarte, para evitar el spam de registros bajo presión. | ### 3.5.3) 🧠 Runtime de Sesión Diferida diff --git a/README_FA.MD b/README_FA.MD index a73ae7c8..3895b23b 100644 --- a/README_FA.MD +++ b/README_FA.MD @@ -687,11 +687,12 @@ sudo journalctl -u masterdnsvpn-client -f | :--- | :--- | :--- | :--- | | `UDP_HOST` | `"0.0.0.0"` | اگر خالی باشد همین مقدار استفاده می‌شود | آدرسی که سرور DNS روی آن bind می‌شود.
`0.0.0.0` یعنی روی همه interfaceها گوش بدهد. | | `UDP_PORT` | `53` | `1..65535` | پورت UDP سرور است.
به‌طور معمول باید همان `53` باشد تا resolverها بتوانند مستقیماً به آن query بفرستند. | +| `FALLBACK` | حذف‌شده | خالی/خاموش، یا `HOST:PORT`؛ IPv6 داخل براکت | یک raw UDP proxy اختیاری برای datagramهای غیر DNS روی listener DNS است. routing برای هر source جدا نگه داشته می‌شود و فقط یک DNS datagram کامل از نظر ساختار، DNS classification را ایجاد یا reset می‌کند. source ناشناخته با اولین packet غیر DNS فوراً forward و روی fallback ثابت می‌شود، ولی source شناخته‌شده به‌عنوان DNS بعد از ۱۶ packet غیر DNS پیاپی جابه‌جا می‌شود. تا وقتی source به‌عنوان DNS شناخته می‌شود، هر packet نگه‌داشته‌شده روی مسیر DNS، timer بی‌کاری ۱۸۰ ثانیه‌ای را refresh می‌کند؛ packet DNS شمارنده non-DNS را reset و packet غیر DNS آن را افزایش می‌دهد. بعد از وقفه‌ای بیش از ۱۸۰ ثانیه، packet غیر DNS بعدی مانند packet یک source جدید فوراً forward می‌شود. session فعال fallback نیز پس از ۱۸۰ ثانیه نبود ترافیک در هر دو جهت منقضی می‌شود. حالت fallback برای حفظ ترتیب از یک ingress reader و buffer کامل ۶۵۵۳۵ بایتی استفاده می‌کند؛ startup مقصدی را که دوباره به همین listener resolve شود رد می‌کند. sessionها سقف سخت ندارند، پس ترافیک untrusted یا spoofed می‌تواند file descriptor و CPU مصرف کند؛ روی listener عمومی از filtering یا rate limit استفاده کنید. | | `UDP_READERS` | `4` | اگر `<=0` باشد auto-default | تعداد goroutineهای خواندن مستقیم از socket UDP.
عدد بالاتر در سرورهای پر ترافیک مفید است، ولی از یک حد به بعد فقط context switching را زیاد می‌کند. | | `DNS_REQUEST_WORKERS` | `8` | اگر `<=0` باشد auto-default | تعداد workerهایی که requestهای ورودی را از front-door queue برمی‌دارند و به لایه session/decode می‌دهند. | | `MAX_CONCURRENT_REQUESTS` | `16384` | اگر `<=0` باشد fallback | ظرفیت صف requestهای ورودی است.
اگر این صف پر شود، پکت‌ها drop می‌شوند و سرور rate-limited overload log می‌دهد. | | `SOCKET_BUFFER_SIZE` | `4194304` | اگر `<=0` باشد fallback | اندازه بافر socket UDP در سطح سیستم‌عامل است.
برای burstهای ورودی زیاد مهم است. | -| `MAX_PACKET_SIZE` | `65535` | اگر `<=0` باشد fallback | اندازه بزرگ‌ترین bufferی که packet pool برای هر packet می‌گیرد. | +| `MAX_PACKET_SIZE` | `65535` | اگر `<=0` باشد fallback؛ با `FALLBACK` حداقل `65535` | اندازه بزرگ‌ترین bufferی که packet pool برای هر packet می‌گیرد. fallback برای جلوگیری از truncate شدن raw datagramها buffer کامل UDP را اجباری می‌کند. | | `DROP_LOG_INTERVAL_SECONDS` | `2.0` | اگر `<=0` باشد fallback | حداقل فاصله بین لاگ‌های drop/overload است تا لاگ سرور در زمان فشار spam نشود. | ### ۳.۵.۳) بخش 🧠 Deferred Session Runtime diff --git a/README_IT.MD b/README_IT.MD index 21941e24..7026f08b 100644 --- a/README_IT.MD +++ b/README_IT.MD @@ -718,11 +718,12 @@ Queste impostazioni sono critiche sul server: | :--- | :--- | :--- | :--- | | `UDP_HOST` | `"0.0.0.0"` | se vuoto, viene usato questo valore | Indirizzo su cui il server DNS si lega.
`0.0.0.0` significa ascoltare su tutte le interfacce. | | `UDP_PORT` | `53` | `1..65535` | Porta UDP usata dal server.
Nella maggior parte delle distribuzioni questa dovrebbe rimanere `53` così che i resolver possano interrogarla direttamente. | +| `FALLBACK` | omesso | vuoto/disattivato, oppure `HOST:PORT`; IPv6 tra parentesi quadre | Proxy UDP raw opzionale per i datagrammi non DNS ricevuti dal listener DNS. Il routing è mantenuto per origine e solo un datagramma DNS strutturalmente completo stabilisce o reimposta la classificazione DNS. Una nuova origine non DNS viene inoltrata subito e resta associata al fallback, mentre un'origine classificata come DNS passa al fallback dopo 16 pacchetti non DNS consecutivi. Finché l'origine resta classificata come DNS, ogni pacchetto trattenuto sul percorso DNS aggiorna il timer di inattività di 180 secondi; un pacchetto DNS azzera la sequenza non DNS e un pacchetto non DNS la incrementa. Dopo una pausa superiore a 180 secondi, il pacchetto non DNS successivo viene trattato come proveniente da una nuova origine e inoltrato immediatamente. Anche le sessioni fallback attive scadono dopo 180 secondi senza traffico in entrambe le direzioni. La modalità fallback usa un solo reader di ingresso per conservare l'ordine e un buffer completo da 65.535 byte; l'avvio rifiuta una destinazione che risolve al listener stesso. Le sessioni non hanno un limite rigido, quindi traffico non attendibile o con origine falsificata può consumare file descriptor e CPU; usare filtri di rete o rate limiting sui listener pubblici. | | `UDP_READERS` | `4` | predefinito automatico se `<=0` | Numero di goroutine che leggono direttamente dal socket UDP.
Un numero maggiore può essere d'aiuto su server molto trafficati, ma oltre un certo punto aumenta solo il context switching. | | `DNS_REQUEST_WORKERS` | `8` | predefinito automatico se `<=0` | Numero di worker che prelevano le richieste dalla coda di ingresso e le passano al livello di sessione/decodifica. | | `MAX_CONCURRENT_REQUESTS` | `16384` | fallback se `<=0` | Capacità della coda delle richieste in entrata.
Se questa coda si riempie, i pacchetti vengono scartati e il server emette log di sovraccarico a frequenza limitata. | | `SOCKET_BUFFER_SIZE` | `4194304` | fallback se `<=0` | Dimensione del buffer del socket richiesta al sistema operativo per il listener UDP.
Questo è importante per i picchi pesanti di traffico in entrata. | -| `MAX_PACKET_SIZE` | `65535` | fallback se `<=0` | Dimensione del buffer del pacchetto più grande che il pool di pacchetti alloca. | +| `MAX_PACKET_SIZE` | `65535` | fallback se `<=0`; almeno `65535` con `FALLBACK` | Dimensione del buffer del pacchetto più grande che il pool di pacchetti alloca. Fallback forza un buffer UDP completo per evitare il troncamento dei datagrammi raw. | | `DROP_LOG_INTERVAL_SECONDS` | `2.0` | fallback se `<=0` | Intervallo minimo tra log di sovraccarico/scarto ripetuti, per evitare lo spam di log durante la pressione. | ### 3.5.3) 🧠 Runtime delle Sessioni Differite diff --git a/README_RU.MD b/README_RU.MD index d7d06191..eb647f9e 100644 --- a/README_RU.MD +++ b/README_RU.MD @@ -721,11 +721,12 @@ Copy-Item client_resolvers.simple client_resolvers.txt | :--- | :--- | :--- | :--- | | `UDP_HOST` | `"0.0.0.0"` | если поле пустое, используется значение по умолчанию. | Адрес, на котором привязан DNS-сервер.
`0.0.0.0` означает прослушивание на всех интерфейсах. | | `UDP_PORT` | `53` | `1..65535` | UDP-порт, используемый сервером.
В большинстве случаев должно быть `53`, чтобы резолверы могли напрямую отправлять запросы на этот порт. | +| `FALLBACK` | не задан | пусто/выключено или `HOST:PORT`; IPv6 в квадратных скобках | Необязательный прокси для пересылки необработанных UDP-датаграмм, не являющихся DNS, которые приходят на DNS-listener. Маршрутизация хранится отдельно для каждого источника, и только структурно полная DNS-датаграмма устанавливает или сбрасывает DNS-классификацию. Новый источник с не-DNS пакетом пересылается сразу и закрепляется за fallback, а источник, классифицированный как DNS, переключается после 16 последовательных не-DNS пакетов. Пока источник остаётся классифицированным как DNS, каждый пакет, удержанный на DNS-пути, обновляет 180-секундный таймер бездействия; DNS-пакет сбрасывает последовательность не-DNS, а не-DNS пакет увеличивает её. После паузы более 180 секунд следующий не-DNS пакет считается пакетом от нового источника и пересылается немедленно. Активные fallback-сессии также истекают после 180 секунд без трафика в обоих направлениях. Режим fallback использует один входной reader для сохранения порядка и полный буфер 65 535 байт; запуск отклоняет адрес, который разрешается обратно в этот listener. Жёсткого лимита числа сессий нет, поэтому недоверенный или поддельный трафик может расходовать файловые дескрипторы и CPU; для публичных listener'ов используйте сетевую фильтрацию или ограничение частоты. | | `UDP_READERS` | `4` | По умолчанию если `<=0` | Количество горутин, осуществляющих чтение непосредственно из UDP-сокета.
Большее количество может помочь на очень загруженных серверах, но после определенного предела это только увеличивает количество переключений контекста. | | `DNS_REQUEST_WORKERS` | `8` | По умолчанию если `<=0` | Количество обработчиков, которые принимают запросы из входной очереди и передают их на уровень сеанса/декодирования. | | `MAX_CONCURRENT_REQUESTS` | `16384` | По умолчанию если `<=0` | Емкость очереди входящих запросов.
Если эта очередь заполняется, пакеты отбрасываются, и сервер выдает сообщения о перегрузке с ограничением скорости. | | `SOCKET_BUFFER_SIZE` | `4194304` | По умолчанию если `<=0` | Запрос размера буфера сокета операционной системы для UDP-прослушивателя.
Это важно при интенсивных всплесках входящего трафика. | -| `MAX_PACKET_SIZE` | `65535` | По умолчанию если `<=0` | Размер самого большого буфера пакетов, выделяемого пулом пакетов. | +| `MAX_PACKET_SIZE` | `65535` | По умолчанию если `<=0`; не менее `65535` с `FALLBACK` | Размер самого большого буфера пакетов, выделяемого пулом пакетов. Fallback принудительно использует полный UDP-буфер, чтобы необработанные датаграммы не обрезались. | | `DROP_LOG_INTERVAL_SECONDS` | `2.0` | По умолчанию если `<=0` | Минимальный интервал между повторными логами о перегрузке/падении, чтобы избежать избыточного количества логов. | ### 3.5.3) 🧠 Отложенная среда выполнения сессии diff --git a/README_ZH.MD b/README_ZH.MD index 97b20612..c267d21f 100644 --- a/README_ZH.MD +++ b/README_ZH.MD @@ -718,11 +718,12 @@ Copy-Item client_resolvers.simple client_resolvers.txt | :--- | :--- | :--- | :--- | | `UDP_HOST` | `"0.0.0.0"` | 若为空则使用此值 | DNS 服务器绑定的地址。
`0.0.0.0` 表示监听所有网络接口。 | | `UDP_PORT` | `53` | `1..65535` | 服务器使用的 UDP 端口。
在大多数部署中应保持为 `53`,以便解析器能直接查询它。 | +| `FALLBACK` | 未设置 | 留空/关闭,或 `HOST:PORT`;IPv6 地址使用方括号 | 可选的原始 UDP 代理,用于转发 DNS 监听器收到的非 DNS 数据报。路由按来源分别维护,只有结构完整的 DNS 数据报才会建立或重置 DNS 分类。未知来源的首个非 DNS 数据包会立即转发并固定使用 fallback;已归类为 DNS 的来源会在连续收到 16 个非 DNS 数据包后切换。来源仍归类为 DNS 时,每个保留在 DNS 路径上的数据包都会刷新 180 秒空闲计时器;DNS 包会重置连续非 DNS 计数,非 DNS 包会递增该计数。若空闲超过 180 秒,下一个非 DNS 包会按新来源处理并立即转发。活动 fallback 会话在双向均无流量 180 秒后也会过期。Fallback 模式使用单个入站读取器来保持顺序,并使用完整的 65,535 字节缓冲区;启动时会拒绝解析回该监听器自身的目标地址。会话数量没有硬上限,因此不可信或伪造的流量可能消耗文件描述符和 CPU;在公网监听器上请使用网络过滤或速率限制。 | | `UDP_READERS` | `4` | 若 `<=0` 则使用自动默认值 | 直接从 UDP 套接字读取的 goroutine 数量。
较大的数值在非常繁忙的服务器上可能有帮助,但超过某个点后只会增加上下文切换。 | | `DNS_REQUEST_WORKERS` | `8` | 若 `<=0` 则使用自动默认值 | 从前门队列中取出请求并将其传入会话/解码层的工作线程数量。 | | `MAX_CONCURRENT_REQUESTS` | `16384` | 若 `<=0` 则使用回退值 | 入站请求队列的容量。
如果此队列填满,数据包会被丢弃,服务器会发出受限速的过载日志。 | | `SOCKET_BUFFER_SIZE` | `4194304` | 若 `<=0` 则使用回退值 | 为 UDP 监听器请求的操作系统套接字缓冲区大小。
这对入站流量的大量突发很重要。 | -| `MAX_PACKET_SIZE` | `65535` | 若 `<=0` 则使用回退值 | 数据包池分配的最大数据包缓冲区大小。 | +| `MAX_PACKET_SIZE` | `65535` | 若 `<=0` 则使用回退值;启用 `FALLBACK` 时至少为 `65535` | 数据包池分配的最大数据包缓冲区大小。Fallback 会强制使用完整大小的 UDP 接收缓冲区,避免原始数据报被截断。 | | `DROP_LOG_INTERVAL_SECONDS` | `2.0` | 若 `<=0` 则使用回退值 | 重复的过载/丢包日志之间的最小间隔,以避免在压力下产生日志刷屏。 | ### 3.5.3) 🧠 延迟会话运行时 diff --git a/internal/config/server.go b/internal/config/server.go index 927e9b2a..1eda777c 100644 --- a/internal/config/server.go +++ b/internal/config/server.go @@ -10,6 +10,7 @@ package config import ( "flag" "fmt" + "net" "os" "path/filepath" "reflect" @@ -29,6 +30,7 @@ type ServerConfig struct { ProtocolType string `toml:"PROTOCOL_TYPE"` UDPHost string `toml:"UDP_HOST"` UDPPort int `toml:"UDP_PORT"` + FallbackAddress string `toml:"FALLBACK"` UDPReaders int `toml:"UDP_READERS"` SocketBufferSize int `toml:"SOCKET_BUFFER_SIZE"` MaxConcurrentRequests int `toml:"MAX_CONCURRENT_REQUESTS"` @@ -116,6 +118,7 @@ func defaultServerConfig() ServerConfig { ProtocolType: "SOCKS5", UDPHost: "0.0.0.0", UDPPort: 53, + FallbackAddress: "", UDPReaders: 4, SocketBufferSize: 8 * 1024 * 1024, MaxConcurrentRequests: 16384, @@ -289,6 +292,31 @@ func finalizeServerConfig(cfg ServerConfig) (ServerConfig, error) { return cfg, fmt.Errorf("invalid UDP_PORT: %d", cfg.UDPPort) } + cfg.FallbackAddress = strings.TrimSpace(cfg.FallbackAddress) + if cfg.FallbackAddress != "" { + host, portText, err := net.SplitHostPort(cfg.FallbackAddress) + if err != nil { + return cfg, fmt.Errorf("invalid FALLBACK address %q: expected HOST:PORT: %w", cfg.FallbackAddress, err) + } + if host == "" || host != strings.TrimSpace(host) { + return cfg, fmt.Errorf("invalid FALLBACK address %q: host must be nonempty", cfg.FallbackAddress) + } + if ip := net.ParseIP(host); ip != nil && ip.IsUnspecified() { + return cfg, fmt.Errorf("invalid FALLBACK address %q: target must not be an unspecified address", cfg.FallbackAddress) + } + portIsNumeric := portText != "" + for _, char := range portText { + if char < '0' || char > '9' { + portIsNumeric = false + break + } + } + port, err := strconv.Atoi(portText) + if !portIsNumeric || err != nil || port < 1 || port > 65535 { + return cfg, fmt.Errorf("invalid FALLBACK address %q: port must be between 1 and 65535", cfg.FallbackAddress) + } + } + if cfg.UDPReaders <= 0 { cfg.UDPReaders = defaultServerConfig().UDPReaders } diff --git a/internal/config/server_test.go b/internal/config/server_test.go index f489041e..664b1cd3 100644 --- a/internal/config/server_test.go +++ b/internal/config/server_test.go @@ -69,6 +69,7 @@ func TestServerConfigFlagBinderBuildsOverridesForSetFlagsOnly(t *testing.T) { if err := fs.Parse([]string{ "-udp-port=5300", + "--fallback=[2001:db8::1]:5353", "-domain=a.example.com,b.example.com", "-use-external-socks5", "-supported-upload-compression-types=0,1", @@ -81,6 +82,9 @@ func TestServerConfigFlagBinderBuildsOverridesForSetFlagsOnly(t *testing.T) { if got, ok := overrides.Values["UDPPort"].(int); !ok || got != 5300 { t.Fatalf("unexpected udp port override: %#v", overrides.Values["UDPPort"]) } + if got, ok := overrides.Values["FallbackAddress"].(string); !ok || got != "[2001:db8::1]:5353" { + t.Fatalf("unexpected fallback override: %#v", overrides.Values["FallbackAddress"]) + } if got, ok := overrides.Values["UseExternalSOCKS5"].(bool); !ok || !got { t.Fatalf("unexpected socks5 override: %#v", overrides.Values["UseExternalSOCKS5"]) } @@ -100,6 +104,62 @@ func TestServerConfigFlagBinderBuildsOverridesForSetFlagsOnly(t *testing.T) { } } +func TestLoadServerConfigAcceptsFallbackAddress(t *testing.T) { + dir := t.TempDir() + configPath := filepath.Join(dir, "server_config.toml") + + if err := os.WriteFile(configPath, []byte(` +PROTOCOL_TYPE = "SOCKS5" +UDP_PORT = 53 +FALLBACK = " fallback.example.com:5353 " +`), 0o644); err != nil { + t.Fatalf("WriteFile config failed: %v", err) + } + + cfg, err := LoadServerConfig(configPath) + if err != nil { + t.Fatalf("LoadServerConfig returned error: %v", err) + } + if cfg.FallbackAddress != "fallback.example.com:5353" { + t.Fatalf("unexpected fallback address: got=%q want=%q", cfg.FallbackAddress, "fallback.example.com:5353") + } +} + +func TestServerConfigFallbackAddressValidation(t *testing.T) { + tests := []struct { + name string + address string + wantErr bool + }{ + {name: "disabled", address: ""}, + {name: "hostname", address: "fallback.example.com:5353"}, + {name: "bracketed ipv6", address: "[2001:db8::1]:5353"}, + {name: "missing port", address: "fallback.example.com", wantErr: true}, + {name: "empty port", address: "fallback.example.com:", wantErr: true}, + {name: "signed port", address: "fallback.example.com:+5353", wantErr: true}, + {name: "zero port", address: "fallback.example.com:0", wantErr: true}, + {name: "port too large", address: "fallback.example.com:65536", wantErr: true}, + {name: "empty host", address: ":5353", wantErr: true}, + {name: "unspecified ipv4", address: "0.0.0.0:5353", wantErr: true}, + {name: "unspecified ipv6", address: "[::]:5353", wantErr: true}, + {name: "unbracketed ipv6", address: "2001:db8::1:5353", wantErr: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cfg := defaultServerConfig() + cfg.FallbackAddress = tt.address + _, err := finalizeServerConfig(cfg) + if tt.wantErr && err == nil { + t.Fatalf("finalizeServerConfig(%q) unexpectedly succeeded", tt.address) + } + if !tt.wantErr && err != nil { + t.Fatalf("finalizeServerConfig(%q) returned error: %v", tt.address, err) + } + }) + } +} + func TestServerConfigEffectiveSizingUsesSmartFloorsAndDerivedCapacities(t *testing.T) { cfg := defaultServerConfig() cfg.ProtocolType = "SOCKS5" diff --git a/internal/dnsparser/parser.go b/internal/dnsparser/parser.go index be8f4425..ceb0a5b2 100644 --- a/internal/dnsparser/parser.go +++ b/internal/dnsparser/parser.go @@ -17,12 +17,14 @@ var ( ErrInvalidName = errors.New("invalid dns name") ErrInvalidQuestion = errors.New("invalid dns question section") ErrInvalidAnswer = errors.New("invalid dns resource record section") + ErrNotDNSMessage = errors.New("packet does not look like a complete dns message") ErrNotDNSRequest = errors.New("packet does not look like a supported dns request") ) const ( - dnsHeaderSize = 12 - maxNameJumps = 10 + dnsHeaderSize = 12 + maxNameJumps = 10 + maxNameWireLength = 255 ) type Header struct { @@ -95,6 +97,42 @@ func ParseDNSRequestLite(data []byte) (LitePacket, error) { return parsePacketLiteWithHeader(data, header) } +// ParseDNSDatagramLite parses a complete, structurally plausible DNS request or +// response. Unlike ParseDNSRequestLite, it validates every declared section and +// rejects trailing data so callers can use success as a routing boundary. +func ParseDNSDatagramLite(data []byte) (LitePacket, error) { + if len(data) < dnsHeaderSize { + return LitePacket{}, ErrPacketTooShort + } + + header := parseHeader(data) + if !isLikelyDNSMessageHeader(header) { + return LitePacket{}, ErrNotDNSMessage + } + + offset, err := skipQuestions(data, dnsHeaderSize, int(header.QDCount)) + if err != nil { + return LitePacket{}, err + } + offset, err = skipResourceRecords(data, offset, int(header.ANCount)) + if err != nil { + return LitePacket{}, err + } + offset, err = skipResourceRecords(data, offset, int(header.NSCount)) + if err != nil { + return LitePacket{}, err + } + offset, err = skipResourceRecords(data, offset, int(header.ARCount)) + if err != nil { + return LitePacket{}, err + } + if offset != len(data) { + return LitePacket{}, ErrNotDNSMessage + } + + return parsePacketLiteWithHeader(data, header) +} + func parsePacketLiteWithHeader(data []byte, header Header) (LitePacket, error) { packet := LitePacket{Header: header} if header.QDCount == 0 { @@ -247,43 +285,54 @@ func parseResourceRecords(data []byte, offset int, count int) ([]ResourceRecord, } func parseName(data []byte, offset int) (string, int, error) { + var name strings.Builder + nextOffset, hasLabel, err := walkName(data, offset, &name) + if err != nil { + return "", nextOffset, err + } + if !hasLabel { + return ".", nextOffset, nil + } + return name.String(), nextOffset, nil +} + +func walkName(data []byte, offset int, name *strings.Builder) (int, bool, error) { dataLen := len(data) if offset >= dataLen { - return "", offset, ErrInvalidName + return offset, false, ErrInvalidName } var ( - jumped bool - jumps int - origNext = offset - name strings.Builder - hasLabel bool + jumped bool + jumps int + nextOffset = offset + hasLabel bool + wireLen = 1 ) for { if offset >= dataLen { - return "", origNext, ErrInvalidName + return nextOffset, hasLabel, ErrInvalidName } length := int(data[offset]) if length == 0 { - offset++ if !jumped { - origNext = offset + nextOffset = offset + 1 } - break + return nextOffset, hasLabel, nil } if length >= 192 { // 0xC0 if offset+1 >= dataLen || jumps >= maxNameJumps { - return "", origNext, ErrInvalidName + return nextOffset, hasLabel, ErrInvalidName } ptr := int(binary.BigEndian.Uint16(data[offset:offset+2]) & 0x3FFF) - if ptr >= dataLen { - return "", origNext, ErrInvalidName + if ptr < dnsHeaderSize || ptr >= offset || ptr >= dataLen { + return nextOffset, hasLabel, ErrInvalidName } if !jumped { - origNext = offset + 2 + nextOffset = offset + 2 jumped = true } offset = ptr @@ -292,33 +341,33 @@ func parseName(data []byte, offset int) (string, int, error) { } if length > 63 { - return "", origNext, ErrInvalidName + return nextOffset, hasLabel, ErrInvalidName + } + if wireLen+length+1 > maxNameWireLength { + return nextOffset, hasLabel, ErrInvalidName } + wireLen += length + 1 offset++ end := offset + length if end > dataLen { - return "", origNext, ErrInvalidName + return nextOffset, hasLabel, ErrInvalidName } - if name.Len() == 0 { - name.Grow(64) - } else { - name.WriteByte('.') + if name != nil { + if !hasLabel { + name.Grow(64) + } else { + name.WriteByte('.') + } + writeLowerASCIILabel(name, data[offset:end]) } - - writeLowerASCIILabel(&name, data[offset:end]) hasLabel = true offset = end if !jumped { - origNext = offset + nextOffset = offset } } - - if !hasLabel { - return ".", origNext, nil - } - return name.String(), origNext, nil } func writeLowerASCIILabel(dst *strings.Builder, label []byte) { diff --git a/internal/dnsparser/parser_lite_test.go b/internal/dnsparser/parser_lite_test.go index 0feba508..95ada1c5 100644 --- a/internal/dnsparser/parser_lite_test.go +++ b/internal/dnsparser/parser_lite_test.go @@ -7,6 +7,7 @@ package dnsparser import ( + "encoding/binary" "testing" Enums "masterdnsvpn-go/internal/enums" @@ -44,6 +45,141 @@ func TestParsePacketLiteParsesAllQuestions(t *testing.T) { } } +func TestParseDNSDatagramLiteRequiresCompleteMessage(t *testing.T) { + query := buildMultiQuestionDNSQuery( + 0x5151, + []liteQuestionSpec{{Name: "example.com", Type: Enums.DNS_RECORD_TYPE_A, Class: Enums.DNSQ_CLASS_IN}}, + false, + ) + response := append([]byte(nil), query...) + response[2] |= 0x80 + + compressedResponse := append([]byte(nil), response...) + binary.BigEndian.PutUint16(compressedResponse[6:8], 1) + compressedResponse = append(compressedResponse, + 0xC0, 0x0C, + 0x00, 0x01, + 0x00, 0x01, + 0x00, 0x00, 0x00, 0x3C, + 0x00, 0x04, + 192, 0, 2, 1, + ) + + multipleQuestions := buildMultiQuestionDNSQuery( + 0x5252, + []liteQuestionSpec{ + {Name: "example.com", Type: Enums.DNS_RECORD_TYPE_A, Class: Enums.DNSQ_CLASS_IN}, + {Name: "example.org", Type: Enums.DNS_RECORD_TYPE_AAAA, Class: Enums.DNSQ_CLASS_IN}, + }, + false, + ) + withOPT := buildMultiQuestionDNSQuery( + 0x5353, + []liteQuestionSpec{{Name: "example.com", Type: Enums.DNS_RECORD_TYPE_A, Class: Enums.DNSQ_CLASS_IN}}, + true, + ) + + incompleteQuestions := append([]byte(nil), query...) + binary.BigEndian.PutUint16(incompleteQuestions[4:6], 2) + missingAnswer := append([]byte(nil), query...) + binary.BigEndian.PutUint16(missingAnswer[6:8], 1) + missingAuthority := append([]byte(nil), query...) + binary.BigEndian.PutUint16(missingAuthority[8:10], 1) + missingAdditional := append([]byte(nil), query...) + binary.BigEndian.PutUint16(missingAdditional[10:12], 1) + oversizedRData := append([]byte(nil), compressedResponse...) + binary.BigEndian.PutUint16(oversizedRData[len(query)+10:len(query)+12], 5) + trailingData := append(append([]byte(nil), query...), 0) + + emptyQuestion := append([]byte(nil), query[:dnsHeaderSize]...) + binary.BigEndian.PutUint16(emptyQuestion[4:6], 0) + emptyResponse := append([]byte(nil), emptyQuestion...) + emptyResponse[2] |= 0x80 + invalidOpcode := append([]byte(nil), query...) + binary.BigEndian.PutUint16(invalidOpcode[2:4], 0x3900) + reservedHeaderBit := append([]byte(nil), query...) + reservedHeaderBit[3] |= 0x40 + + badQuestionPointer := make([]byte, dnsHeaderSize+6) + binary.BigEndian.PutUint16(badQuestionPointer[2:4], 0x0100) + binary.BigEndian.PutUint16(badQuestionPointer[4:6], 1) + copy(badQuestionPointer[dnsHeaderSize:], []byte{0xC0, 0xFF, 0x00, 0x01, 0x00, 0x01}) + headerQuestionPointer := append([]byte(nil), badQuestionPointer...) + headerQuestionPointer[dnsHeaderSize+1] = 0 + cyclicQuestionPointer := append([]byte(nil), badQuestionPointer...) + cyclicQuestionPointer[dnsHeaderSize+1] = dnsHeaderSize + + cyclicAnswerPointer := append([]byte(nil), response...) + binary.BigEndian.PutUint16(cyclicAnswerPointer[6:8], 1) + answerOffset := len(cyclicAnswerPointer) + cyclicAnswerPointer = append(cyclicAnswerPointer, + byte(0xC0|answerOffset>>8), byte(answerOffset), + 0x00, 0x01, + 0x00, 0x01, + 0x00, 0x00, 0x00, 0x00, + 0x00, 0x00, + ) + + tooManyQuestionSpecs := make([]liteQuestionSpec, maxLikelyQuestions+1) + for i := range tooManyQuestionSpecs { + tooManyQuestionSpecs[i] = liteQuestionSpec{Name: ".", Type: Enums.DNS_RECORD_TYPE_A, Class: Enums.DNSQ_CLASS_IN} + } + tooManyQuestions := buildMultiQuestionDNSQuery(0x5454, tooManyQuestionSpecs, false) + + tooManyAnswers := append([]byte(nil), response...) + binary.BigEndian.PutUint16(tooManyAnswers[6:8], maxLikelyAnswers+1) + for range maxLikelyAnswers + 1 { + tooManyAnswers = append(tooManyAnswers, + 0x00, + 0x00, 0x01, + 0x00, 0x01, + 0x00, 0x00, 0x00, 0x00, + 0x00, 0x00, + ) + } + + tests := []struct { + name string + packet []byte + want bool + }{ + {name: "query", packet: query, want: true}, + {name: "response", packet: response, want: true}, + {name: "compressed response", packet: compressedResponse, want: true}, + {name: "multiple questions", packet: multipleQuestions, want: true}, + {name: "EDNS OPT", packet: withOPT, want: true}, + {name: "missing declared question", packet: incompleteQuestions}, + {name: "missing declared answer", packet: missingAnswer}, + {name: "missing declared authority", packet: missingAuthority}, + {name: "missing declared additional", packet: missingAdditional}, + {name: "oversized RDATA", packet: oversizedRData}, + {name: "trailing data", packet: trailingData}, + {name: "empty question", packet: emptyQuestion}, + {name: "empty-question response", packet: emptyResponse, want: true}, + {name: "invalid opcode", packet: invalidOpcode}, + {name: "reserved header bit", packet: reservedHeaderBit}, + {name: "out of range question pointer", packet: badQuestionPointer}, + {name: "question pointer into header", packet: headerQuestionPointer}, + {name: "cyclic question pointer", packet: cyclicQuestionPointer}, + {name: "cyclic answer pointer", packet: cyclicAnswerPointer}, + {name: "too many questions", packet: tooManyQuestions}, + {name: "too many answers", packet: tooManyAnswers}, + {name: "short packet", packet: []byte("not DNS")}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + parsed, err := ParseDNSDatagramLite(tt.packet) + if got := err == nil; got != tt.want { + t.Fatalf("ParseDNSDatagramLite() success=%t want=%t err=%v", got, tt.want, err) + } + if tt.want && parsed.Header.QDCount > 0 && (!parsed.HasQuestion || parsed.QuestionEndOffset <= dnsHeaderSize) { + t.Fatalf("successful parse returned incomplete question metadata: %+v", parsed) + } + }) + } +} + type liteQuestionSpec struct { Name string Type uint16 diff --git a/internal/dnsparser/response.go b/internal/dnsparser/response.go index 01a75e45..0f66ffc2 100644 --- a/internal/dnsparser/response.go +++ b/internal/dnsparser/response.go @@ -156,15 +156,22 @@ func getARCount(optLen int) int { } func isLikelyDNSRequestHeader(header Header) bool { - if header.QR != 0 { + if header.QR != 0 || header.QDCount == 0 { return false } - if header.QDCount == 0 || header.QDCount > maxLikelyQuestions { + return isLikelyDNSMessageHeader(header) +} + +func isLikelyDNSMessageHeader(header Header) bool { + if (header.QR == 0 && header.QDCount == 0) || header.QDCount > maxLikelyQuestions { return false } if header.OpCode > 6 { return false } + if header.Flags&0x0040 != 0 { + return false + } if header.ANCount > maxLikelyAnswers { return false } @@ -389,28 +396,6 @@ func extractRawOPTRecords(data []byte, offset int, count int) ([][]byte, int, in } func skipName(data []byte, offset int) (int, error) { - dataLen := len(data) - for { - if offset >= dataLen { - return offset, ErrInvalidName - } - - length := int(data[offset]) - if length == 0 { - return offset + 1, nil - } - - if length >= 192 { // 0xC0 - if offset+1 >= dataLen { - return offset, ErrInvalidName - } - return offset + 2, nil - } - - if length > 63 { - return offset, ErrInvalidName - } - - offset += length + 1 - } + nextOffset, _, err := walkName(data, offset, nil) + return nextOffset, err } diff --git a/internal/udpserver/server.go b/internal/udpserver/server.go index 068226c7..121ef3ec 100644 --- a/internal/udpserver/server.go +++ b/internal/udpserver/server.go @@ -10,6 +10,7 @@ package udpserver import ( "container/heap" "context" + "fmt" "net" "strconv" "sync" @@ -18,6 +19,7 @@ import ( "masterdnsvpn-go/internal/config" dnsCache "masterdnsvpn-go/internal/dnscache" + DnsParser "masterdnsvpn-go/internal/dnsparser" domainMatcher "masterdnsvpn-go/internal/domainmatcher" fragmentStore "masterdnsvpn-go/internal/fragmentstore" "masterdnsvpn-go/internal/logger" @@ -83,13 +85,18 @@ type Server struct { lastDeferredDropLogUnix atomic.Int64 pongNonce atomic.Uint32 invalidDropMode atomic.Uint32 + fallback *udpFallbackManager } type request struct { - buf []byte - size int - addr *net.UDPAddr - conn *net.UDPConn + buf []byte + size int + addr *net.UDPAddr + conn *net.UDPConn + parsed DnsParser.LitePacket + parseErr error + hasParsedDNS bool + fallbackEpoch *udpFallbackRouteEpoch } type postSessionValidation struct { @@ -119,6 +126,10 @@ func New(cfg config.ServerConfig, log *logger.Logger, codec *security.Codec) *Se sessions := newSessionStore(cfg.EffectiveSessionOrphanQueueInitialCap(), cfg.EffectiveStreamQueueInitialCapacity(), cfg.SessionInitReuseTTL(), cfg.RecentlyClosedStreamTTL(), cfg.RecentlyClosedStreamCap) sessions.maxActiveSessions = cfg.MaxAllowedClientActiveSessions sessions.maxActiveStreams = cfg.MaxAllowedClientActiveStreams + packetBufferSize := cfg.MaxPacketSize + if cfg.FallbackAddress != "" && packetBufferSize < udpFallbackMaxUDPPacketSize { + packetBufferSize = udpFallbackMaxUDPPacketSize + } return &Server{ cfg: cfg, log: log, @@ -167,7 +178,7 @@ func New(cfg config.ServerConfig, log *logger.Logger, codec *security.Codec) *Se deferredInflightIndex: make(map[uint8]map[uint16]map[uint64]struct{}, 64), packetPool: sync.Pool{ New: func() any { - return make([]byte, cfg.MaxPacketSize) + return make([]byte, packetBufferSize) }, }, } @@ -296,20 +307,59 @@ func (s *Server) Run(ctx context.Context) error { runCtx, cancel := context.WithCancel(ctx) defer cancel() - conns, err := s.openUDPListeners() + var fallbackAddr *net.UDPAddr + if s.cfg.FallbackAddress != "" { + resolved, err := net.ResolveUDPAddr("udp", s.cfg.FallbackAddress) + if err != nil { + return fmt.Errorf("resolve fallback address %q: %w", s.cfg.FallbackAddress, err) + } + if resolved.IP.IsUnspecified() { + return fmt.Errorf("fallback address %q resolves to an unspecified address", s.cfg.FallbackAddress) + } + fallbackAddr = resolved + } + + readerCount := s.cfg.EffectiveUDPReaders() + var fallback *udpFallbackManager + if fallbackAddr != nil { + // Fallback classification is stateful and order-sensitive per source. + readerCount = 1 + } + conns, err := s.openUDPListeners(readerCount) if err != nil { return err } - defer func() { + var closeListenersOnce sync.Once + closeListeners := func() { + closeListenersOnce.Do(func() { + for _, conn := range conns { + _ = conn.Close() + } + }) + } + defer closeListeners() + + if fallbackAddr != nil { for _, conn := range conns { - _ = conn.Close() + localAddr, ok := conn.LocalAddr().(*net.UDPAddr) + if ok && fallbackTargetsListener(localAddr, fallbackAddr) { + return fmt.Errorf("fallback address %s resolves to the UDP listener and would loop", fallbackAddr) + } } - }() + + fallback = newUDPFallbackManager(fallbackAddr, s.log) + s.fallback = fallback + defer func() { + closeListeners() + fallback.Close() + s.fallback = nil + }() + } s.log.Infof( "\U0001F4E1 UDP Listener Ready, Addr: %s, Readers: %d, Workers: %d, Queue: %d, Sockets: %d", s.cfg.Address(), - s.cfg.EffectiveUDPReaders(), + readerCount, s.cfg.EffectiveDNSRequestWorkers(), s.cfg.EffectiveMaxConcurrentRequests(), len(conns), @@ -330,16 +380,20 @@ func (s *Server) Run(ctx context.Context) error { go func() { <-runCtx.Done() - for _, conn := range conns { - _ = conn.Close() + closeListeners() + if fallback != nil { + fallback.Close() } }() readErrCh := make(chan error, max(1, len(conns))) var readerWG sync.WaitGroup - s.startReaders(runCtx, conns, reqCh, readErrCh, &readerWG) + s.startReaders(runCtx, conns, readerCount, reqCh, readErrCh, &readerWG) readerWG.Wait() + // A UDP reply can be waiting for send-buffer space. Close the listener + // before waiting for workers or fallback reply loops so teardown unblocks it. + closeListeners() close(reqCh) workerWG.Wait() cancel() @@ -356,3 +410,42 @@ func (s *Server) Run(ctx context.Context) error { return nil } } + +func fallbackTargetsListener(listener *net.UDPAddr, target *net.UDPAddr) bool { + if listener == nil || target == nil || listener.Port != target.Port { + return false + } + if target.IP.IsUnspecified() { + return true + } + if listener.IP.Equal(target.IP) && listener.Zone == target.Zone { + return true + } + if !listener.IP.IsUnspecified() { + return false + } + if listener.IP.To4() != nil && target.IP.To4() == nil { + return false + } + if target.IP.IsLoopback() { + return true + } + + interfaceAddrs, err := net.InterfaceAddrs() + if err != nil { + return false + } + for _, addr := range interfaceAddrs { + var ip net.IP + switch value := addr.(type) { + case *net.IPNet: + ip = value.IP + case *net.IPAddr: + ip = value.IP + } + if ip != nil && ip.Equal(target.IP) { + return true + } + } + return false +} diff --git a/internal/udpserver/server_fallback_test.go b/internal/udpserver/server_fallback_test.go new file mode 100644 index 00000000..828ba28d --- /dev/null +++ b/internal/udpserver/server_fallback_test.go @@ -0,0 +1,490 @@ +package udpserver + +import ( + "bytes" + "context" + "encoding/binary" + "net" + "testing" + "time" + + "masterdnsvpn-go/internal/config" + domainMatcher "masterdnsvpn-go/internal/domainmatcher" + Enums "masterdnsvpn-go/internal/enums" + "masterdnsvpn-go/internal/logger" +) + +const fallbackIntegrationTimeout = 2 * time.Second + +func TestFallbackTargetsListener(t *testing.T) { + tests := []struct { + name string + listener *net.UDPAddr + target *net.UDPAddr + want bool + }{ + { + name: "exact listener", + listener: &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 53}, + target: &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 53}, + want: true, + }, + { + name: "wildcard covers loopback", + listener: &net.UDPAddr{IP: net.IPv4zero, Port: 53}, + target: &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 53}, + want: true, + }, + { + name: "unspecified target reaches specific listener", + listener: &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 53}, + target: &net.UDPAddr{IP: net.IPv4zero, Port: 53}, + want: true, + }, + { + name: "different port", + listener: &net.UDPAddr{IP: net.IPv4zero, Port: 53}, + target: &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 5353}, + }, + { + name: "specific listener does not cover other address", + listener: &net.UDPAddr{IP: net.IPv4(192, 0, 2, 1), Port: 53}, + target: &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 53}, + }, + { + name: "IPv4 wildcard does not cover IPv6", + listener: &net.UDPAddr{IP: net.IPv4zero, Port: 53}, + target: &net.UDPAddr{IP: net.IPv6loopback, Port: 53}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := fallbackTargetsListener(tt.listener, tt.target); got != tt.want { + t.Fatalf("fallbackTargetsListener()=%t want=%t", got, tt.want) + } + }) + } +} + +func TestNewFallbackServerUsesFullUDPPacketBuffer(t *testing.T) { + server := New(config.ServerConfig{ + FallbackAddress: "127.0.0.1:5353", + MaxPacketSize: 512, + }, nil, nil) + buffer := server.packetPool.Get().([]byte) + if len(buffer) != udpFallbackMaxUDPPacketSize { + t.Fatalf("unexpected fallback packet buffer: got=%d want=%d", len(buffer), udpFallbackMaxUDPPacketSize) + } +} + +type fallbackIntegrationEcho struct { + conn *net.UDPConn + received chan []byte + done chan struct{} +} + +func startFallbackIntegrationEcho(t *testing.T) *fallbackIntegrationEcho { + t.Helper() + + conn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatalf("ListenUDP fallback echo failed: %v", err) + } + + echo := &fallbackIntegrationEcho{ + conn: conn, + received: make(chan []byte, 32), + done: make(chan struct{}), + } + go func() { + defer close(echo.done) + buffer := make([]byte, 65535) + for { + n, peer, err := conn.ReadFromUDP(buffer) + if err != nil { + return + } + + packet := append([]byte(nil), buffer[:n]...) + echo.received <- packet + if _, err := conn.WriteToUDP(packet, peer); err != nil { + return + } + } + }() + + return echo +} + +func (e *fallbackIntegrationEcho) Close(t *testing.T) { + t.Helper() + + _ = e.conn.Close() + select { + case <-e.done: + case <-time.After(fallbackIntegrationTimeout): + t.Error("fallback echo goroutine did not stop") + } +} + +func startFallbackIntegrationRuntime(t *testing.T, server *Server, workerConn *net.UDPConn, readerConn *net.UDPConn) { + t.Helper() + if server.packetPool.New == nil { + server.packetPool.New = func() any { + return make([]byte, udpFallbackMaxUDPPacketSize) + } + } + + ctx, cancel := context.WithCancel(context.Background()) + requests := make(chan request, 4) + workerDone := make(chan struct{}) + go func() { + defer close(workerDone) + server.dnsWorker(ctx, workerConn, requests, 1) + }() + readerDone := make(chan error, 1) + go func() { + readerDone <- server.readLoop(ctx, readerConn, requests, 1) + }() + + t.Cleanup(func() { + cancel() + _ = readerConn.Close() + select { + case err := <-readerDone: + if err != nil { + t.Errorf("fallback integration reader failed: %v", err) + } + case <-time.After(fallbackIntegrationTimeout): + t.Error("fallback integration reader did not stop") + } + close(requests) + select { + case <-workerDone: + case <-time.After(fallbackIntegrationTimeout): + t.Error("DNS worker did not stop") + } + }) +} + +func enqueueFallbackIntegrationDatagram( + t *testing.T, + client *net.UDPConn, + listener *net.UDPConn, + packet []byte, +) { + t.Helper() + + if _, err := client.WriteToUDP(packet, listener.LocalAddr().(*net.UDPAddr)); err != nil { + t.Fatalf("send test datagram failed: %v", err) + } +} + +func readFallbackIntegrationReply(t *testing.T, client *net.UDPConn) ([]byte, *net.UDPAddr) { + t.Helper() + + if err := client.SetReadDeadline(time.Now().Add(fallbackIntegrationTimeout)); err != nil { + t.Fatalf("SetReadDeadline client failed: %v", err) + } + buffer := make([]byte, 65535) + n, peer, err := client.ReadFromUDP(buffer) + if err != nil { + t.Fatalf("read test reply failed: %v", err) + } + return append([]byte(nil), buffer[:n]...), peer +} + +func fallbackIntegrationSameUDPAddr(left *net.UDPAddr, right *net.UDPAddr) bool { + return left != nil && right != nil && + left.Port == right.Port && left.Zone == right.Zone && left.IP.Equal(right.IP) +} + +func TestDNSWorkerFallbackRoundTripKeepsOtherDNSPeerLocal(t *testing.T) { + echo := startFallbackIntegrationEcho(t) + t.Cleanup(func() { echo.Close(t) }) + + listener, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatalf("ListenUDP DNS listener failed: %v", err) + } + t.Cleanup(func() { _ = listener.Close() }) + primaryListener, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatalf("ListenUDP primary DNS listener failed: %v", err) + } + t.Cleanup(func() { _ = primaryListener.Close() }) + + server := &Server{ + log: logger.New("Fallback Integration Test", "ERROR"), + domainMatcher: domainMatcher.New([]string{"vpn.example.com"}, 3), + } + server.fallback = newUDPFallbackManager(echo.conn.LocalAddr().(*net.UDPAddr), server.log) + t.Cleanup(server.fallback.Close) + // The worker default is deliberately different from the request listener, + // mirroring the reuseport path where each request carries its ingress socket. + startFallbackIntegrationRuntime(t, server, primaryListener, listener) + + fallbackClient, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatalf("ListenUDP fallback client failed: %v", err) + } + t.Cleanup(func() { _ = fallbackClient.Close() }) + + rawPacket := []byte("not a DNS packet") + enqueueFallbackIntegrationDatagram(t, fallbackClient, listener, rawPacket) + reply, replyPeer := readFallbackIntegrationReply(t, fallbackClient) + if !bytes.Equal(reply, rawPacket) { + t.Fatalf("unexpected fallback reply: got=%q want=%q", reply, rawPacket) + } + if !fallbackIntegrationSameUDPAddr(replyPeer, listener.LocalAddr().(*net.UDPAddr)) { + t.Fatalf("fallback reply did not come from DNS listener: got=%v want=%v", replyPeer, listener.LocalAddr()) + } + select { + case forwarded := <-echo.received: + if !bytes.Equal(forwarded, rawPacket) { + t.Fatalf("unexpected packet at fallback endpoint: got=%q want=%q", forwarded, rawPacket) + } + case <-time.After(fallbackIntegrationTimeout): + t.Fatal("fallback endpoint did not receive raw packet") + } + + dnsClient, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatalf("ListenUDP DNS client failed: %v", err) + } + t.Cleanup(func() { _ = dnsClient.Close() }) + if fallbackIntegrationSameUDPAddr(fallbackClient.LocalAddr().(*net.UDPAddr), dnsClient.LocalAddr().(*net.UDPAddr)) { + t.Fatal("fallback and DNS clients unexpectedly share a source address") + } + + dnsQuery := buildTestDNSQuery(0x7171, "outside.example", Enums.DNS_RECORD_TYPE_A) + enqueueFallbackIntegrationDatagram(t, dnsClient, listener, dnsQuery) + dnsReply, dnsReplyPeer := readFallbackIntegrationReply(t, dnsClient) + if !fallbackIntegrationSameUDPAddr(dnsReplyPeer, listener.LocalAddr().(*net.UDPAddr)) { + t.Fatalf("DNS reply did not come from DNS listener: got=%v want=%v", dnsReplyPeer, listener.LocalAddr()) + } + if len(dnsReply) < 12 { + t.Fatalf("DNS reply too short: %d", len(dnsReply)) + } + if got := binary.BigEndian.Uint16(dnsReply[2:4]) & 0x000F; got != Enums.DNSR_CODE_NAME_ERROR { + t.Fatalf("unexpected DNS rcode: got=%d want=%d", got, Enums.DNSR_CODE_NAME_ERROR) + } + select { + case forwarded := <-echo.received: + t.Fatalf("fallback endpoint received DNS packet from separate source: %x", forwarded) + case <-time.After(200 * time.Millisecond): + } +} + +func TestDNSWorkerFallbackCompleteResponseUsesDNSStreak(t *testing.T) { + echo := startFallbackIntegrationEcho(t) + t.Cleanup(func() { echo.Close(t) }) + + listener, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatalf("ListenUDP DNS listener failed: %v", err) + } + t.Cleanup(func() { _ = listener.Close() }) + + server := &Server{log: logger.New("Fallback DNS Response Test", "ERROR")} + server.fallback = newUDPFallbackManager(echo.conn.LocalAddr().(*net.UDPAddr), server.log) + t.Cleanup(server.fallback.Close) + startFallbackIntegrationRuntime(t, server, listener, listener) + + client, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatalf("ListenUDP fallback client failed: %v", err) + } + t.Cleanup(func() { _ = client.Close() }) + + response := buildTestDNSQuery(0x7272, "example.org", Enums.DNS_RECORD_TYPE_A) + response[2] |= 0x80 + enqueueFallbackIntegrationDatagram(t, client, listener, response) + + nonDNS := buildTestDNSQuery(0x7373, "example.org", Enums.DNS_RECORD_TYPE_A) + binary.BigEndian.PutUint16(nonDNS[4:6], 2) + for range udpFallbackNonDNSStreakLimit - 1 { + enqueueFallbackIntegrationDatagram(t, client, listener, nonDNS) + } + if err := client.SetReadDeadline(time.Now().Add(200 * time.Millisecond)); err != nil { + t.Fatalf("SetReadDeadline client failed: %v", err) + } + buffer := make([]byte, 64) + if _, _, err := client.ReadFromUDP(buffer); err == nil { + t.Fatal("DNS response or pre-threshold non-DNS packet unexpectedly reached fallback") + } else if netErr, ok := err.(net.Error); !ok || !netErr.Timeout() { + t.Fatalf("pre-threshold read failed unexpectedly: %v", err) + } + select { + case forwarded := <-echo.received: + t.Fatalf("fallback endpoint received pre-threshold packet: %x", forwarded) + default: + } + + enqueueFallbackIntegrationDatagram(t, client, listener, nonDNS) + if reply, _ := readFallbackIntegrationReply(t, client); !bytes.Equal(reply, nonDNS) { + t.Fatalf("unexpected threshold fallback reply: got=%x want=%x", reply, nonDNS) + } + select { + case forwarded := <-echo.received: + if !bytes.Equal(forwarded, nonDNS) { + t.Fatalf("unexpected threshold packet at fallback endpoint: got=%x want=%x", forwarded, nonDNS) + } + case <-time.After(fallbackIntegrationTimeout): + t.Fatal("fallback endpoint did not receive threshold packet") + } +} + +func TestDNSWorkerDropsQueuedReplyAfterFallbackTransition(t *testing.T) { + upstream := newUDPFallbackTestUpstream(t, false) + listener := newUDPFallbackTestConn(t) + client := newUDPFallbackTestConn(t) + server := &Server{ + log: logger.New("Fallback Stale DNS Test", "ERROR"), + domainMatcher: domainMatcher.New([]string{"vpn.example.com"}, 3), + } + server.packetPool.New = func() any { + return make([]byte, udpFallbackMaxUDPPacketSize) + } + server.fallback = newUDPFallbackManager(udpFallbackTestAddr(t, upstream.conn), server.log) + t.Cleanup(server.fallback.Close) + + ctx, cancel := context.WithCancel(context.Background()) + requests := make(chan request, 1) + readerDone := make(chan error, 1) + go func() { + readerDone <- server.readLoop(ctx, listener, requests, 1) + }() + t.Cleanup(func() { + cancel() + _ = listener.Close() + select { + case err := <-readerDone: + if err != nil { + t.Errorf("fallback reader failed: %v", err) + } + case <-time.After(fallbackIntegrationTimeout): + t.Error("fallback reader did not stop") + } + }) + + query := buildTestDNSQuery(0x7676, "outside.example", Enums.DNS_RECORD_TYPE_A) + enqueueFallbackIntegrationDatagram(t, client, listener, query) + var queued request + select { + case queued = <-requests: + case <-time.After(fallbackIntegrationTimeout): + t.Fatal("DNS request was not queued") + } + if queued.fallbackEpoch == nil { + t.Fatal("queued DNS request has no routing epoch") + } + + nonDNS := buildTestDNSQuery(0x7777, "example.org", Enums.DNS_RECORD_TYPE_A) + binary.BigEndian.PutUint16(nonDNS[4:6], 2) + for range udpFallbackNonDNSStreakLimit { + enqueueFallbackIntegrationDatagram(t, client, listener, nonDNS) + } + if got := receiveUDPFallbackTestPacket(t, upstream).payload; !bytes.Equal(got, nonDNS) { + t.Fatalf("unexpected threshold fallback packet: got=%x want=%x", got, nonDNS) + } + + workerRequests := make(chan request, 1) + workerRequests <- queued + close(workerRequests) + server.dnsWorker(context.Background(), listener, workerRequests, 1) + + if err := client.SetReadDeadline(time.Now().Add(200 * time.Millisecond)); err != nil { + t.Fatalf("SetReadDeadline client failed: %v", err) + } + buffer := make([]byte, 65535) + if _, _, err := client.ReadFromUDP(buffer); err == nil { + t.Fatal("stale DNS reply reached the active fallback flow") + } else if netErr, ok := err.(net.Error); !ok || !netErr.Timeout() { + t.Fatalf("stale-reply read failed unexpectedly: %v", err) + } +} + +func TestDNSWorkerFallbackForwardsNonDNSDatagrams(t *testing.T) { + emptyQuestion := make([]byte, 12) + binary.BigEndian.PutUint16(emptyQuestion[0:2], 0x7272) + binary.BigEndian.PutUint16(emptyQuestion[2:4], 0x0100) + + incompleteQuestions := buildTestDNSQuery(0x7373, "example.org", Enums.DNS_RECORD_TYPE_A) + binary.BigEndian.PutUint16(incompleteQuestions[4:6], 2) + + for _, tt := range []struct { + name string + packet []byte + }{ + {name: "empty question", packet: emptyQuestion}, + {name: "incomplete declared questions", packet: incompleteQuestions}, + } { + t.Run(tt.name, func(t *testing.T) { + echo := startFallbackIntegrationEcho(t) + t.Cleanup(func() { echo.Close(t) }) + + listener, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatalf("ListenUDP DNS listener failed: %v", err) + } + t.Cleanup(func() { _ = listener.Close() }) + + server := &Server{log: logger.New("Fallback Non-DNS Test", "ERROR")} + server.fallback = newUDPFallbackManager(echo.conn.LocalAddr().(*net.UDPAddr), server.log) + t.Cleanup(server.fallback.Close) + startFallbackIntegrationRuntime(t, server, listener, listener) + + client, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatalf("ListenUDP fallback client failed: %v", err) + } + t.Cleanup(func() { _ = client.Close() }) + + enqueueFallbackIntegrationDatagram(t, client, listener, tt.packet) + + if reply, _ := readFallbackIntegrationReply(t, client); !bytes.Equal(reply, tt.packet) { + t.Fatalf("unexpected fallback reply: got=%x want=%x", reply, tt.packet) + } + select { + case forwarded := <-echo.received: + if !bytes.Equal(forwarded, tt.packet) { + t.Fatalf("unexpected packet at fallback endpoint: got=%x want=%x", forwarded, tt.packet) + } + case <-time.After(fallbackIntegrationTimeout): + t.Fatal("fallback endpoint did not receive datagram") + } + }) + } +} + +func TestDNSWorkerWithoutFallbackSilentlyDropsNonDNS(t *testing.T) { + listener, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatalf("ListenUDP DNS listener failed: %v", err) + } + t.Cleanup(func() { _ = listener.Close() }) + + server := &Server{ + log: logger.New("Fallback Disabled Test", "ERROR"), + } + startFallbackIntegrationRuntime(t, server, listener, listener) + + client, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatalf("ListenUDP client failed: %v", err) + } + t.Cleanup(func() { _ = client.Close() }) + + enqueueFallbackIntegrationDatagram(t, client, listener, []byte("not DNS")) + if err := client.SetReadDeadline(time.Now().Add(200 * time.Millisecond)); err != nil { + t.Fatalf("SetReadDeadline client failed: %v", err) + } + buffer := make([]byte, 64) + if _, _, err := client.ReadFromUDP(buffer); err == nil { + t.Fatal("non-DNS packet unexpectedly received a reply with fallback disabled") + } else if netErr, ok := err.(net.Error); !ok || !netErr.Timeout() { + t.Fatalf("read with fallback disabled failed unexpectedly: %v", err) + } +} diff --git a/internal/udpserver/server_ingress.go b/internal/udpserver/server_ingress.go index 80bc0c87..99b7071d 100644 --- a/internal/udpserver/server_ingress.go +++ b/internal/udpserver/server_ingress.go @@ -8,7 +8,6 @@ package udpserver import ( - "errors" "fmt" "time" @@ -19,13 +18,13 @@ import ( ) func (s *Server) handlePacket(packet []byte) []byte { - parsed, err := DnsParser.ParseDNSRequestLite(packet) - if err != nil { - if errors.Is(err, DnsParser.ErrNotDNSRequest) || errors.Is(err, DnsParser.ErrPacketTooShort) { - return nil - } + parsed, err := DnsParser.ParseDNSDatagramLite(packet) + return s.handleParsedPacket(packet, parsed, err) +} - return s.buildNoDataResponseLogged(packet, "request-parse-failed") +func (s *Server) handleParsedPacket(packet []byte, parsed DnsParser.LitePacket, err error) []byte { + if err != nil || parsed.Header.QR != 0 { + return nil } if !parsed.HasQuestion { @@ -52,7 +51,6 @@ func (s *Server) handlePacket(packet []byte) []byte { return s.buildNoDataResponseLiteLogged(packet, parsed, "domain-match-unknown-action") } } - func (s *Server) handleTunnelCandidate(packet []byte, parsed DnsParser.LitePacket, decision domainMatcher.Decision) []byte { vpnPacket, err := VpnProto.ParseInflatedFromLabels(decision.Labels, s.codec) if err != nil { diff --git a/internal/udpserver/server_ingress_test.go b/internal/udpserver/server_ingress_test.go index 50822850..49944dfa 100644 --- a/internal/udpserver/server_ingress_test.go +++ b/internal/udpserver/server_ingress_test.go @@ -77,6 +77,29 @@ func TestHandlePacketKeepsUnsupportedAllowedAQueryAsNoData(t *testing.T) { } } +func TestHandlePacketDropsNonRequestDatagrams(t *testing.T) { + server := &Server{} + response := buildTestDNSQuery(0x6262, "example.org", Enums.DNS_RECORD_TYPE_A) + response[2] |= 0x80 + + incompleteQuestions := buildTestDNSQuery(0x6363, "example.org", Enums.DNS_RECORD_TYPE_A) + binary.BigEndian.PutUint16(incompleteQuestions[4:6], 2) + + emptyQuestion := make([]byte, 12) + + for name, packet := range map[string][]byte{ + "response": response, + "incomplete questions": incompleteQuestions, + "empty question": emptyQuestion, + } { + t.Run(name, func(t *testing.T) { + if reply := server.handlePacket(packet); reply != nil { + t.Fatalf("non-request datagram received a reply: %x", reply) + } + }) + } +} + func buildTestDNSQuery(id uint16, name string, qtype uint16) []byte { qname := encodeTestDNSName(name) packet := make([]byte, 12+len(qname)+4) diff --git a/internal/udpserver/server_runtime.go b/internal/udpserver/server_runtime.go index 6cd98b20..a57de372 100644 --- a/internal/udpserver/server_runtime.go +++ b/internal/udpserver/server_runtime.go @@ -14,6 +14,7 @@ import ( "sync" "time" + DnsParser "masterdnsvpn-go/internal/dnsparser" "masterdnsvpn-go/internal/logger" ) @@ -27,12 +28,12 @@ func (s *Server) configureSocketBuffers(conn *net.UDPConn) { } } -func (s *Server) openUDPListeners() ([]*net.UDPConn, error) { +func (s *Server) openUDPListeners(readerCount int) ([]*net.UDPConn, error) { addr := &net.UDPAddr{ IP: net.ParseIP(s.cfg.UDPHost), Port: s.cfg.UDPPort, } - desired := s.cfg.EffectiveUDPReaders() + desired := readerCount if desired < 1 { desired = 1 } @@ -74,12 +75,11 @@ func (s *Server) startDNSWorkers(ctx context.Context, conn *net.UDPConn, reqCh < } } -func (s *Server) startReaders(ctx context.Context, conns []*net.UDPConn, reqCh chan<- request, readErrCh chan<- error, readerWG *sync.WaitGroup) { +func (s *Server) startReaders(ctx context.Context, conns []*net.UDPConn, readerCount int, reqCh chan<- request, readErrCh chan<- error, readerWG *sync.WaitGroup) { if len(conns) == 0 { return } - readerCount := s.cfg.EffectiveUDPReaders() if readerCount < 1 { readerCount = 1 } @@ -195,8 +195,40 @@ func (s *Server) readLoop(ctx context.Context, conn *net.UDPConn, reqCh chan<- r return err } + req := request{buf: buffer, size: n, addr: addr, conn: conn} + if fallback := s.fallback; fallback != nil { + packet := buffer[:n] + if fallback.ForwardIfActive(packet, addr, conn) { + s.packetPool.Put(buffer) + continue + } + parsed, parseErr, ok := s.safeParseDNSDatagram(packet) + if !ok { + s.packetPool.Put(buffer) + continue + } + if parseErr != nil { + fallback.RouteNonDNS(packet, addr, conn) + s.packetPool.Put(buffer) + continue + } + routeDNS, epoch := fallback.RouteDNS(packet, addr, conn) + if !routeDNS { + s.packetPool.Put(buffer) + continue + } + if parsed.Header.QR != 0 { + s.packetPool.Put(buffer) + continue + } + req.parsed = parsed + req.parseErr = parseErr + req.hasParsedDNS = true + req.fallbackEpoch = epoch + } + select { - case reqCh <- request{buf: buffer, size: n, addr: addr, conn: conn}: + case reqCh <- req: case <-ctx.Done(): s.packetPool.Put(buffer) return nil @@ -217,13 +249,25 @@ func (s *Server) dnsWorker(ctx context.Context, conn *net.UDPConn, reqCh <-chan return } - response := s.safeHandlePacket(req.buf[:req.size]) + var response []byte + packet := req.buf[:req.size] + if req.hasParsedDNS { + response = s.safeHandleParsedPacket(packet, req.parsed, req.parseErr) + } else { + response = s.safeHandlePacket(packet) + } if len(response) != 0 { writeConn := conn if req.conn != nil { writeConn = req.conn } - if _, err := writeConn.WriteToUDP(response, req.addr); err != nil { + var err error + if req.fallbackEpoch != nil { + err = req.fallbackEpoch.writeIfActive(writeConn, response, req.addr) + } else { + _, err = writeConn.WriteToUDP(response, req.addr) + } + if err != nil { s.log.Debugf( "\U0001F4A5 UDP Write Error, Worker: %d, Remote: %v, Error: %v", workerID, @@ -238,7 +282,31 @@ func (s *Server) dnsWorker(ctx context.Context, conn *net.UDPConn, reqCh <-chan } } -func (s *Server) safeHandlePacket(packet []byte) (response []byte) { +func (s *Server) safeHandlePacket(packet []byte) []byte { + parsed, parseErr, ok := s.safeParseDNSDatagram(packet) + if !ok { + return nil + } + return s.safeHandleParsedPacket(packet, parsed, parseErr) +} + +func (s *Server) safeParseDNSDatagram(packet []byte) (parsed DnsParser.LitePacket, parseErr error, ok bool) { + defer func() { + if recovered := recover(); recovered != nil { + if s.log != nil { + s.log.Errorf( + "\U0001F4A5 Packet Parser Panic Recovered, %v", + recovered, + ) + } + } + }() + + parsed, parseErr = DnsParser.ParseDNSDatagramLite(packet) + return parsed, parseErr, true +} + +func (s *Server) safeHandleParsedPacket(packet []byte, parsed DnsParser.LitePacket, parseErr error) (response []byte) { defer func() { if recovered := recover(); recovered != nil { if s.log != nil { @@ -251,7 +319,7 @@ func (s *Server) safeHandlePacket(packet []byte) (response []byte) { } }() - return s.handlePacket(packet) + return s.handleParsedPacket(packet, parsed, parseErr) } func (s *Server) onDrop(addr *net.UDPAddr, queueLen int, queueCap int) { diff --git a/internal/udpserver/udp_fallback.go b/internal/udpserver/udp_fallback.go new file mode 100644 index 00000000..9b4d6e43 --- /dev/null +++ b/internal/udpserver/udp_fallback.go @@ -0,0 +1,603 @@ +// ============================================================================== +// MasterDnsVPN +// Author: MasterkinG32 +// Github: https://github.com/masterking32 +// Year: 2026 +// ============================================================================== + +package udpserver + +import ( + "errors" + "net" + "sync" + "sync/atomic" + "time" + + "masterdnsvpn-go/internal/logger" +) + +const ( + udpFallbackIdleTimeout = 180 * time.Second + udpFallbackCleanupInterval = 30 * time.Second + udpFallbackNonDNSStreakLimit = 16 + udpFallbackMaxUDPPacketSize = 65535 + udpFallbackSendQueueSize = 16 +) + +type udpFallbackPeerKey struct { + ip string + port int + zone string +} + +type udpFallbackRouteMode uint8 + +const ( + udpFallbackRouteDNS udpFallbackRouteMode = iota + udpFallbackRouteFallback +) + +type udpFallbackReplyWriter interface { + WriteToUDP([]byte, *net.UDPAddr) (int, error) +} + +// udpFallbackWriteBarrier serializes client-bound writes for one peer across +// route epochs. A blocked peer never holds the manager lock or ingress reader. +type udpFallbackWriteBarrier struct { + mu sync.Mutex +} + +// udpFallbackRouteEpoch is a validity token for work queued during one routing +// period. Successive epochs for the same peer share a write barrier. +type udpFallbackRouteEpoch struct { + active atomic.Bool + barrier *udpFallbackWriteBarrier +} + +func newUDPFallbackRouteEpoch(barrier *udpFallbackWriteBarrier) *udpFallbackRouteEpoch { + if barrier == nil { + barrier = &udpFallbackWriteBarrier{} + } + epoch := &udpFallbackRouteEpoch{barrier: barrier} + epoch.active.Store(true) + return epoch +} + +func (e *udpFallbackRouteEpoch) invalidate() { + if e != nil { + e.active.Store(false) + } +} + +func (e *udpFallbackRouteEpoch) writeIfActive( + conn udpFallbackReplyWriter, + packet []byte, + peer *net.UDPAddr, +) error { + if e == nil { + return nil + } + e.barrier.mu.Lock() + defer e.barrier.mu.Unlock() + if !e.active.Load() { + return nil + } + _, err := conn.WriteToUDP(packet, peer) + return err +} + +type udpFallbackPeerState struct { + mode udpFallbackRouteMode + epoch *udpFallbackRouteEpoch + lastSeen time.Time + nonDNSStreak int + peer *net.UDPAddr + listener *net.UDPConn + session *udpFallbackSession +} + +type udpFallbackSession struct { + conn net.Conn + peerLabel string + sendCh chan []byte + closed chan struct{} + done chan struct{} + closeOnce sync.Once +} + +func (s *udpFallbackSession) close() { + if s == nil { + return + } + s.closeOnce.Do(func() { + if s.closed != nil { + close(s.closed) + } + if s.conn != nil { + _ = s.conn.Close() + } + }) +} + +type udpFallbackManager struct { + target *net.UDPAddr + log *logger.Logger + + mu sync.Mutex + peers map[udpFallbackPeerKey]*udpFallbackPeerState + barriers map[udpFallbackPeerKey]*udpFallbackWriteBarrier + dialUDP func(string, *net.UDPAddr, *net.UDPAddr) (*net.UDPConn, error) + closed bool + stopCh chan struct{} + cleanupDone chan struct{} + sessionWG sync.WaitGroup + closeOnce sync.Once +} + +func newUDPFallbackManager(target *net.UDPAddr, log *logger.Logger) *udpFallbackManager { + manager := &udpFallbackManager{ + target: cloneUDPAddr(target), + log: log, + peers: make(map[udpFallbackPeerKey]*udpFallbackPeerState), + barriers: make(map[udpFallbackPeerKey]*udpFallbackWriteBarrier), + dialUDP: net.DialUDP, + stopCh: make(chan struct{}), + cleanupDone: make(chan struct{}), + } + + if target != nil { + manager.log.Infof("Non-DNS UDP packets will be forwarded to %s", target) + } + go manager.cleanupLoop() + return manager +} + +func (m *udpFallbackManager) Close() { + if m == nil { + return + } + + m.closeOnce.Do(func() { + m.mu.Lock() + m.closed = true + close(m.stopCh) + sessions := make([]*udpFallbackSession, 0, len(m.peers)) + for _, state := range m.peers { + state.epoch.invalidate() + if state.session != nil { + sessions = append(sessions, state.session) + state.session = nil + } + } + m.peers = nil + m.barriers = nil + m.mu.Unlock() + + for _, session := range sessions { + session.close() + } + + <-m.cleanupDone + m.sessionWG.Wait() + }) +} + +func (m *udpFallbackManager) ForwardIfActive( + packet []byte, + peer *net.UDPAddr, + listener *net.UDPConn, +) bool { + if m == nil || peer == nil { + return false + } + + key := makeUDPFallbackPeerKey(peer) + now := time.Now() + + m.mu.Lock() + if m.closed { + m.mu.Unlock() + return false + } + state := m.peerStateLocked(key, now) + session, handled := m.fallbackSessionLocked(key, state, peer, listener, now) + m.mu.Unlock() + + if !handled { + return false + } + m.forwardPacket(session, packet, peer) + return true +} + +func (m *udpFallbackManager) RouteDNS( + packet []byte, + peer *net.UDPAddr, + listener *net.UDPConn, +) (bool, *udpFallbackRouteEpoch) { + if m == nil || peer == nil { + return true, nil + } + + key := makeUDPFallbackPeerKey(peer) + now := time.Now() + + m.mu.Lock() + if m.closed { + m.mu.Unlock() + return true, nil + } + state := m.peerStateLocked(key, now) + session, handled := m.fallbackSessionLocked(key, state, peer, listener, now) + if handled { + m.mu.Unlock() + m.forwardPacket(session, packet, peer) + return false, nil + } + + if state == nil { + state = &udpFallbackPeerState{ + mode: udpFallbackRouteDNS, + epoch: newUDPFallbackRouteEpoch(m.routeBarrierLocked(key)), + lastSeen: now, + peer: cloneUDPAddr(peer), + listener: listener, + } + m.peers[key] = state + } else { + state.lastSeen = now + state.nonDNSStreak = 0 + state.peer = cloneUDPAddr(peer) + state.listener = listener + } + epoch := state.epoch + m.mu.Unlock() + return true, epoch +} + +func (m *udpFallbackManager) RouteNonDNS( + packet []byte, + peer *net.UDPAddr, + listener *net.UDPConn, +) { + if m == nil || peer == nil { + return + } + + key := makeUDPFallbackPeerKey(peer) + now := time.Now() + + m.mu.Lock() + if m.closed { + m.mu.Unlock() + return + } + + state := m.peerStateLocked(key, now) + session, handled := m.fallbackSessionLocked(key, state, peer, listener, now) + if handled { + m.mu.Unlock() + m.forwardPacket(session, packet, peer) + return + } + + if state != nil { + state.nonDNSStreak++ + state.lastSeen = now + if state.nonDNSStreak < udpFallbackNonDNSStreakLimit { + m.mu.Unlock() + return + } + state.epoch.invalidate() + session = m.activateFallbackLocked(key, peer, listener, now) + m.mu.Unlock() + m.forwardPacket(session, packet, peer) + return + } + + session = m.activateFallbackLocked(key, peer, listener, now) + m.mu.Unlock() + + m.forwardPacket(session, packet, peer) +} + +func (m *udpFallbackManager) peerStateLocked( + key udpFallbackPeerKey, + now time.Time, +) *udpFallbackPeerState { + state := m.peers[key] + if state == nil { + return nil + } + if now.Sub(state.lastSeen) <= udpFallbackIdleTimeout { + return state + } + + state.epoch.invalidate() + if state.session != nil { + state.session.close() + state.session = nil + } + if state.mode == udpFallbackRouteFallback { + m.log.Debugf("Expired UDP fallback session for %s", state.peer) + } + delete(m.peers, key) + return nil +} + +func (m *udpFallbackManager) fallbackSessionLocked( + key udpFallbackPeerKey, + state *udpFallbackPeerState, + peer *net.UDPAddr, + listener *net.UDPConn, + now time.Time, +) (*udpFallbackSession, bool) { + if state == nil { + return nil, false + } + switch state.mode { + case udpFallbackRouteFallback: + return m.prepareFallbackLocked(key, state, peer, listener, now), true + default: + return nil, false + } +} + +func (m *udpFallbackManager) activateFallbackLocked( + key udpFallbackPeerKey, + peer *net.UDPAddr, + listener *net.UDPConn, + now time.Time, +) *udpFallbackSession { + state := &udpFallbackPeerState{ + mode: udpFallbackRouteFallback, + epoch: newUDPFallbackRouteEpoch(m.routeBarrierLocked(key)), + lastSeen: now, + peer: cloneUDPAddr(peer), + listener: listener, + } + m.peers[key] = state + return m.createSessionLocked(key, state) +} + +func (m *udpFallbackManager) routeBarrierLocked( + key udpFallbackPeerKey, +) *udpFallbackWriteBarrier { + barrier := m.barriers[key] + if barrier == nil { + barrier = &udpFallbackWriteBarrier{} + m.barriers[key] = barrier + } + return barrier +} + +func (m *udpFallbackManager) prepareFallbackLocked( + key udpFallbackPeerKey, + state *udpFallbackPeerState, + peer *net.UDPAddr, + listener *net.UDPConn, + now time.Time, +) *udpFallbackSession { + state.lastSeen = now + state.peer = cloneUDPAddr(peer) + state.listener = listener + + if state.session != nil { + select { + case <-state.session.done: + state.session.close() + state.session = nil + m.log.Debugf("UDP fallback reply loop ended for %s; recreating session", peer) + default: + return state.session + } + } + + return m.createSessionLocked(key, state) +} + +func (m *udpFallbackManager) createSessionLocked( + key udpFallbackPeerKey, + state *udpFallbackPeerState, +) *udpFallbackSession { + peer := state.peer + if m.target == nil || len(m.target.IP) == 0 { + m.log.Warnf("Unable to create UDP fallback session for %s: fallback target is not resolved", peer) + return nil + } + + network := "udp6" + localAddr := &net.UDPAddr{IP: net.IPv6unspecified} + if m.target.IP.To4() != nil { + network = "udp4" + localAddr = &net.UDPAddr{IP: net.IPv4zero} + } + + conn, err := m.dialUDP(network, localAddr, m.target) + if err != nil { + m.log.Warnf("Unable to create UDP fallback session for %s: %v", peer, err) + return nil + } + + session := &udpFallbackSession{ + conn: conn, + peerLabel: peer.String(), + sendCh: make(chan []byte, udpFallbackSendQueueSize), + closed: make(chan struct{}), + done: make(chan struct{}), + } + state.session = session + m.sessionWG.Add(2) + go m.forwardPackets(session) + go m.forwardReplies(key, state.epoch, session) + m.log.Debugf("Created UDP fallback session for %s", peer) + return session +} + +func (m *udpFallbackManager) forwardPacket( + session *udpFallbackSession, + packet []byte, + peer *net.UDPAddr, +) { + if session == nil { + return + } + select { + case <-session.closed: + return + default: + } + if len(session.sendCh) == cap(session.sendCh) { + m.log.Debugf("Dropped UDP fallback packet for %s: send queue is full", peer) + return + } + + queued := append([]byte(nil), packet...) + select { + case session.sendCh <- queued: + case <-session.closed: + default: + m.log.Debugf("Dropped UDP fallback packet for %s: send queue is full", peer) + } +} + +func (m *udpFallbackManager) forwardPackets(session *udpFallbackSession) { + defer m.sessionWG.Done() + + for { + select { + case <-session.closed: + return + default: + } + select { + case packet := <-session.sendCh: + if _, err := session.conn.Write(packet); err != nil && !errors.Is(err, net.ErrClosed) { + m.log.Warnf("UDP fallback write failed for %s via %s: %v", session.peerLabel, m.target, err) + } + case <-session.closed: + return + } + } +} + +func (m *udpFallbackManager) forwardReplies( + key udpFallbackPeerKey, + epoch *udpFallbackRouteEpoch, + session *udpFallbackSession, +) { + defer m.sessionWG.Done() + defer close(session.done) + defer session.close() + + buffer := make([]byte, udpFallbackMaxUDPPacketSize) + for { + size, err := session.conn.Read(buffer) + if err != nil { + if !errors.Is(err, net.ErrClosed) { + m.log.Warnf("UDP fallback reply read failed for %s: %v", session.peerLabel, err) + } + return + } + + now := time.Now() + m.mu.Lock() + state := m.peers[key] + if m.closed || + state == nil || + state.mode != udpFallbackRouteFallback || + state.epoch != epoch || + state.session != session { + m.mu.Unlock() + continue + } + state.lastSeen = now + listener := state.listener + peer := cloneUDPAddr(state.peer) + m.mu.Unlock() + + if listener == nil || peer == nil { + continue + } + err = epoch.writeIfActive(listener, buffer[:size], peer) + if err != nil && !errors.Is(err, net.ErrClosed) { + m.log.Warnf("UDP fallback reply write failed for %s: %v", peer, err) + } + } +} + +func (m *udpFallbackManager) cleanupLoop() { + ticker := time.NewTicker(udpFallbackCleanupInterval) + defer ticker.Stop() + defer close(m.cleanupDone) + + for { + select { + case <-m.stopCh: + return + case now := <-ticker.C: + m.cleanup(now) + } + } +} + +func (m *udpFallbackManager) cleanup(now time.Time) { + if m == nil { + return + } + + m.mu.Lock() + if m.closed { + m.mu.Unlock() + return + } + + expired := make([]*udpFallbackSession, 0) + for key, state := range m.peers { + if now.Sub(state.lastSeen) <= udpFallbackIdleTimeout { + continue + } + + state.epoch.invalidate() + if state.session != nil { + expired = append(expired, state.session) + state.session = nil + } + delete(m.peers, key) + } + // A new epoch must reuse the barrier while an admitted write from the old + // epoch still owns it. Idle barriers can be discarded immediately. + for key, barrier := range m.barriers { + if _, active := m.peers[key]; active || !barrier.mu.TryLock() { + continue + } + barrier.mu.Unlock() + delete(m.barriers, key) + } + m.mu.Unlock() + + for _, session := range expired { + session.close() + } +} + +func makeUDPFallbackPeerKey(peer *net.UDPAddr) udpFallbackPeerKey { + if peer == nil { + return udpFallbackPeerKey{} + } + return udpFallbackPeerKey{ + ip: peer.IP.String(), + port: peer.Port, + zone: peer.Zone, + } +} + +func cloneUDPAddr(addr *net.UDPAddr) *net.UDPAddr { + if addr == nil { + return nil + } + clone := *addr + clone.IP = append(net.IP(nil), addr.IP...) + return &clone +} diff --git a/internal/udpserver/udp_fallback_test.go b/internal/udpserver/udp_fallback_test.go new file mode 100644 index 00000000..edb9cf08 --- /dev/null +++ b/internal/udpserver/udp_fallback_test.go @@ -0,0 +1,703 @@ +package udpserver + +import ( + "bytes" + "errors" + "net" + "sync" + "testing" + "time" + + Enums "masterdnsvpn-go/internal/enums" +) + +type udpFallbackTestPacket struct { + payload []byte + peer *net.UDPAddr +} + +type udpFallbackBlockingConn struct { + writeStarted chan struct{} + closed chan struct{} + startOnce sync.Once + closeOnce sync.Once +} + +func newUDPFallbackBlockingConn() *udpFallbackBlockingConn { + return &udpFallbackBlockingConn{ + writeStarted: make(chan struct{}), + closed: make(chan struct{}), + } +} + +func (c *udpFallbackBlockingConn) Read([]byte) (int, error) { + <-c.closed + return 0, net.ErrClosed +} + +func (c *udpFallbackBlockingConn) Write(packet []byte) (int, error) { + c.startOnce.Do(func() { close(c.writeStarted) }) + <-c.closed + return 0, net.ErrClosed +} + +func (c *udpFallbackBlockingConn) Close() error { + c.closeOnce.Do(func() { close(c.closed) }) + return nil +} + +func (c *udpFallbackBlockingConn) LocalAddr() net.Addr { return &net.UDPAddr{} } +func (c *udpFallbackBlockingConn) RemoteAddr() net.Addr { return &net.UDPAddr{} } +func (c *udpFallbackBlockingConn) SetDeadline(time.Time) error { return nil } +func (c *udpFallbackBlockingConn) SetReadDeadline(time.Time) error { return nil } +func (c *udpFallbackBlockingConn) SetWriteDeadline(time.Time) error { return nil } + +type udpFallbackControlledReplyWriter struct { + label string + started chan struct{} + release <-chan struct{} + events chan<- string +} + +func (w *udpFallbackControlledReplyWriter) WriteToUDP( + packet []byte, + peer *net.UDPAddr, +) (int, error) { + close(w.started) + if w.release != nil { + <-w.release + } + if w.events != nil { + w.events <- w.label + } + return len(packet), nil +} + +type udpFallbackTestUpstream struct { + conn *net.UDPConn + packets chan udpFallbackTestPacket + done chan struct{} +} + +func newUDPFallbackTestConn(t *testing.T) *net.UDPConn { + t.Helper() + conn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatalf("listen UDP: %v", err) + } + t.Cleanup(func() { _ = conn.Close() }) + return conn +} + +func newUDPFallbackTestUpstream(t *testing.T, echo bool) *udpFallbackTestUpstream { + t.Helper() + upstream := &udpFallbackTestUpstream{ + conn: newUDPFallbackTestConn(t), + packets: make(chan udpFallbackTestPacket, 32), + done: make(chan struct{}), + } + go func() { + defer close(upstream.done) + buffer := make([]byte, udpFallbackMaxUDPPacketSize) + for { + size, peer, err := upstream.conn.ReadFromUDP(buffer) + if err != nil { + return + } + payload := append([]byte(nil), buffer[:size]...) + upstream.packets <- udpFallbackTestPacket{payload: payload, peer: cloneUDPAddr(peer)} + if echo { + _, _ = upstream.conn.WriteToUDP(payload, peer) + } + } + }() + t.Cleanup(func() { + _ = upstream.conn.Close() + waitUDPFallbackTestDone(t, upstream.done, "upstream") + }) + return upstream +} + +func udpFallbackTestAddr(t *testing.T, conn *net.UDPConn) *net.UDPAddr { + t.Helper() + addr, ok := conn.LocalAddr().(*net.UDPAddr) + if !ok { + t.Fatalf("unexpected UDP address type %T", conn.LocalAddr()) + } + return cloneUDPAddr(addr) +} + +func receiveUDPFallbackTestPacket(t *testing.T, upstream *udpFallbackTestUpstream) udpFallbackTestPacket { + t.Helper() + select { + case packet := <-upstream.packets: + return packet + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for fallback upstream packet") + return udpFallbackTestPacket{} + } +} + +func receiveUDPFallbackClientPacket(t *testing.T, conn *net.UDPConn) []byte { + t.Helper() + if err := conn.SetReadDeadline(time.Now().Add(2 * time.Second)); err != nil { + t.Fatalf("set UDP read deadline: %v", err) + } + buffer := make([]byte, udpFallbackMaxUDPPacketSize) + size, _, err := conn.ReadFromUDP(buffer) + if err != nil { + t.Fatalf("read fallback client packet: %v", err) + } + return append([]byte(nil), buffer[:size]...) +} + +func waitUDPFallbackTestDone(t *testing.T, done <-chan struct{}, name string) { + t.Helper() + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatalf("%s goroutine did not stop", name) + } +} + +func TestUDPFallbackForwardsBidirectionallyAndKeepsPeersDistinct(t *testing.T) { + upstream := newUDPFallbackTestUpstream(t, true) + listener := newUDPFallbackTestConn(t) + clientOne := newUDPFallbackTestConn(t) + clientTwo := newUDPFallbackTestConn(t) + manager := newUDPFallbackManager(udpFallbackTestAddr(t, upstream.conn), nil) + t.Cleanup(manager.Close) + peerOne := udpFallbackTestAddr(t, clientOne) + peerTwo := udpFallbackTestAddr(t, clientTwo) + + manager.RouteNonDNS([]byte("first"), peerOne, listener) + first := receiveUDPFallbackTestPacket(t, upstream) + if got := string(receiveUDPFallbackClientPacket(t, clientOne)); got != "first" { + t.Fatalf("unexpected first reply %q", got) + } + if !manager.ForwardIfActive([]byte("sticky"), peerOne, listener) { + t.Fatal("expected sticky fallback fast path") + } + receiveUDPFallbackTestPacket(t, upstream) + if got := string(receiveUDPFallbackClientPacket(t, clientOne)); got != "sticky" { + t.Fatalf("unexpected sticky reply %q", got) + } + + manager.RouteNonDNS([]byte("second peer"), peerTwo, listener) + second := receiveUDPFallbackTestPacket(t, upstream) + if first.peer.String() == second.peer.String() { + t.Fatalf("distinct peers shared fallback source %s", first.peer) + } + if got := string(receiveUDPFallbackClientPacket(t, clientTwo)); got != "second peer" { + t.Fatalf("unexpected second-peer reply %q", got) + } + if routeDNS, _ := manager.RouteDNS([]byte("parsed DNS race"), peerOne, listener); routeDNS { + t.Fatal("active fallback peer returned to DNS processing") + } + if got := string(receiveUDPFallbackTestPacket(t, upstream).payload); got != "parsed DNS race" { + t.Fatalf("unexpected raced DNS forwarding %q", got) + } +} + +func TestUDPFallbackDNSStreakThresholdAndReset(t *testing.T) { + upstream := newUDPFallbackTestUpstream(t, false) + listener := newUDPFallbackTestConn(t) + client := newUDPFallbackTestConn(t) + manager := newUDPFallbackManager(udpFallbackTestAddr(t, upstream.conn), nil) + t.Cleanup(manager.Close) + peer := udpFallbackTestAddr(t, client) + key := makeUDPFallbackPeerKey(peer) + + if routeDNS, _ := manager.RouteDNS([]byte("DNS"), peer, listener); !routeDNS { + t.Fatal("new DNS peer should remain DNS") + } + for range 8 { + manager.RouteNonDNS([]byte("stray"), peer, listener) + } + if routeDNS, _ := manager.RouteDNS([]byte("reset"), peer, listener); !routeDNS { + t.Fatal("DNS should reset the non-DNS streak") + } + for range udpFallbackNonDNSStreakLimit - 1 { + manager.RouteNonDNS([]byte("stray"), peer, listener) + } + manager.mu.Lock() + state := manager.peers[key] + sessionExists := state != nil && state.session != nil + manager.mu.Unlock() + if state == nil || + state.mode != udpFallbackRouteDNS || + state.nonDNSStreak != udpFallbackNonDNSStreakLimit-1 || + sessionExists { + t.Fatalf("unexpected pre-threshold state: state=%+v session=%t", state, sessionExists) + } + + manager.RouteNonDNS([]byte("switch"), peer, listener) + if got := string(receiveUDPFallbackTestPacket(t, upstream).payload); got != "switch" { + t.Fatalf("unexpected threshold packet %q", got) + } +} + +func TestUDPFallbackAdmittedDNSWritePrecedesFallbackReply(t *testing.T) { + manager := newUDPFallbackManager(nil, nil) + t.Cleanup(manager.Close) + peer := &net.UDPAddr{IP: net.IPv4(192, 0, 2, 1), Port: 5300} + key := makeUDPFallbackPeerKey(peer) + processDNS, dnsEpoch := manager.RouteDNS([]byte("DNS"), peer, nil) + if !processDNS || dnsEpoch == nil { + t.Fatal("new DNS peer did not receive a routing epoch") + } + + releaseDNS := make(chan struct{}) + var releaseOnce sync.Once + t.Cleanup(func() { releaseOnce.Do(func() { close(releaseDNS) }) }) + events := make(chan string, 2) + dnsWriter := &udpFallbackControlledReplyWriter{ + label: "DNS", + started: make(chan struct{}), + release: releaseDNS, + events: events, + } + dnsDone := make(chan error, 1) + go func() { + dnsDone <- dnsEpoch.writeIfActive(dnsWriter, []byte("DNS reply"), peer) + }() + select { + case <-dnsWriter.started: + case <-time.After(time.Second): + t.Fatal("DNS write was not admitted") + } + + for range udpFallbackNonDNSStreakLimit - 1 { + manager.RouteNonDNS([]byte("non-DNS"), peer, nil) + } + transitionDone := make(chan struct{}) + go func() { + manager.RouteNonDNS([]byte("non-DNS"), peer, nil) + close(transitionDone) + }() + select { + case <-transitionDone: + case <-time.After(time.Second): + t.Fatal("fallback transition waited for the admitted DNS write") + } + manager.mu.Lock() + state := manager.peers[key] + manager.mu.Unlock() + if state == nil || state.mode != udpFallbackRouteFallback { + t.Fatalf("peer did not transition to fallback: %+v", state) + } + fallbackEpoch := state.epoch + if fallbackEpoch.barrier != dnsEpoch.barrier { + t.Fatal("route transition replaced the peer write barrier") + } + + otherPeer := &net.UDPAddr{IP: net.IPv4(192, 0, 2, 2), Port: 5300} + processDNS, otherEpoch := manager.RouteDNS([]byte("other DNS"), otherPeer, nil) + if !processDNS || otherEpoch == nil { + t.Fatal("unrelated DNS peer did not receive a routing epoch") + } + otherWriter := &udpFallbackControlledReplyWriter{started: make(chan struct{})} + otherDone := make(chan error, 1) + go func() { + otherDone <- otherEpoch.writeIfActive(otherWriter, []byte("other reply"), otherPeer) + }() + select { + case err := <-otherDone: + if err != nil { + t.Fatalf("unrelated DNS write failed: %v", err) + } + case <-time.After(time.Second): + t.Fatal("admitted DNS write blocked an unrelated peer") + } + + fallbackWriter := &udpFallbackControlledReplyWriter{ + label: "fallback", + started: make(chan struct{}), + events: events, + } + fallbackAttempted := make(chan struct{}) + fallbackDone := make(chan error, 1) + go func() { + close(fallbackAttempted) + fallbackDone <- fallbackEpoch.writeIfActive(fallbackWriter, []byte("fallback reply"), peer) + }() + <-fallbackAttempted + select { + case <-fallbackWriter.started: + t.Fatal("fallback reply overtook the admitted DNS write") + case <-time.After(100 * time.Millisecond): + } + + releaseOnce.Do(func() { close(releaseDNS) }) + for _, want := range []string{"DNS", "fallback"} { + select { + case got := <-events: + if got != want { + t.Fatalf("unexpected completed client write: got=%s want=%s", got, want) + } + case <-time.After(2 * time.Second): + t.Fatalf("timed out waiting for %s client write", want) + } + } + for _, write := range []struct { + name string + done <-chan error + }{ + {"DNS", dnsDone}, + {"fallback", fallbackDone}, + } { + select { + case err := <-write.done: + if err != nil { + t.Fatalf("%s write failed: %v", write.name, err) + } + case <-time.After(2 * time.Second): + t.Fatalf("timed out waiting for %s writer", write.name) + } + } +} + +func TestUDPFallbackCleanupRetainsBusyWriteBarrier(t *testing.T) { + manager := newUDPFallbackManager(nil, nil) + t.Cleanup(manager.Close) + peer := &net.UDPAddr{IP: net.IPv4(192, 0, 2, 1), Port: 5300} + key := makeUDPFallbackPeerKey(peer) + processDNS, epoch := manager.RouteDNS([]byte("DNS"), peer, nil) + if !processDNS || epoch == nil { + t.Fatal("new DNS peer did not receive a routing epoch") + } + + epoch.barrier.mu.Lock() + var unlockOnce sync.Once + t.Cleanup(func() { unlockOnce.Do(epoch.barrier.mu.Unlock) }) + manager.mu.Lock() + manager.peers[key].lastSeen = time.Now().Add(-udpFallbackIdleTimeout - time.Second) + manager.mu.Unlock() + manager.cleanup(time.Now()) + + manager.mu.Lock() + _, peerExists := manager.peers[key] + retainedBarrier := manager.barriers[key] + manager.mu.Unlock() + if peerExists { + t.Fatal("cleanup retained expired peer state") + } + if retainedBarrier != epoch.barrier { + t.Fatal("cleanup discarded a busy peer write barrier") + } + + processDNS, nextEpoch := manager.RouteDNS([]byte("new DNS"), peer, nil) + if !processDNS || nextEpoch == nil { + t.Fatal("expired peer did not receive a new routing epoch") + } + if nextEpoch.barrier != epoch.barrier { + t.Fatal("peer recreation replaced its busy write barrier") + } + + unlockOnce.Do(epoch.barrier.mu.Unlock) + manager.mu.Lock() + manager.peers[key].lastSeen = time.Now().Add(-udpFallbackIdleTimeout - time.Second) + manager.mu.Unlock() + manager.cleanup(time.Now()) + manager.mu.Lock() + _, peerExists = manager.peers[key] + _, barrierExists := manager.barriers[key] + manager.mu.Unlock() + if peerExists || barrierExists { + t.Fatal("cleanup retained an idle expired peer write barrier") + } +} + +func TestUDPFallbackRetainedNonDNSRefreshesIdleTimeout(t *testing.T) { + upstream := newUDPFallbackTestUpstream(t, false) + listener := newUDPFallbackTestConn(t) + client := newUDPFallbackTestConn(t) + manager := newUDPFallbackManager(udpFallbackTestAddr(t, upstream.conn), nil) + t.Cleanup(manager.Close) + peer := udpFallbackTestAddr(t, client) + key := makeUDPFallbackPeerKey(peer) + + if routeDNS, _ := manager.RouteDNS([]byte("DNS"), peer, listener); !routeDNS { + t.Fatal("new DNS peer should remain DNS") + } + staleLastSeen := time.Now().Add(-udpFallbackIdleTimeout / 2) + manager.mu.Lock() + state := manager.peers[key] + state.lastSeen = staleLastSeen + manager.mu.Unlock() + + beforeRoute := time.Now() + manager.RouteNonDNS([]byte("retained"), peer, listener) + manager.mu.Lock() + state = manager.peers[key] + exists := state != nil + sessionExists := state != nil && state.session != nil + manager.mu.Unlock() + if !exists || state.mode != udpFallbackRouteDNS || sessionExists { + t.Fatalf("retained non-DNS packet changed routing: dns=%t fallback=%t", exists, sessionExists) + } + if state.nonDNSStreak != 1 { + t.Fatalf("unexpected non-DNS streak: got=%d want=1", state.nonDNSStreak) + } + if state.lastSeen.Before(beforeRoute) { + t.Fatalf("retained non-DNS packet did not refresh activity: got=%v before=%v", state.lastSeen, beforeRoute) + } + + manager.cleanup(staleLastSeen.Add(udpFallbackIdleTimeout + time.Second)) + manager.mu.Lock() + _, exists = manager.peers[key] + manager.mu.Unlock() + if !exists { + t.Fatal("cleanup expired a DNS peer that remained active via non-DNS traffic") + } +} + +func TestUDPFallbackBackpressuredUpstreamDoesNotBlockIngressOrClose(t *testing.T) { + manager := newUDPFallbackManager(nil, nil) + t.Cleanup(manager.Close) + peer := &net.UDPAddr{IP: net.IPv4(192, 0, 2, 1), Port: 5300} + key := makeUDPFallbackPeerKey(peer) + conn := newUDPFallbackBlockingConn() + session := &udpFallbackSession{ + conn: conn, + peerLabel: peer.String(), + sendCh: make(chan []byte, udpFallbackSendQueueSize), + closed: make(chan struct{}), + done: make(chan struct{}), + } + state := &udpFallbackPeerState{ + mode: udpFallbackRouteFallback, + epoch: newUDPFallbackRouteEpoch(nil), + lastSeen: time.Now(), + peer: cloneUDPAddr(peer), + session: session, + } + manager.mu.Lock() + manager.peers[key] = state + manager.sessionWG.Add(1) + go manager.forwardPackets(session) + manager.mu.Unlock() + + firstForwardDone := make(chan bool, 1) + go func() { + firstForwardDone <- manager.ForwardIfActive([]byte("blocked"), peer, nil) + }() + select { + case <-conn.writeStarted: + case <-time.After(time.Second): + t.Fatal("fallback writer did not reach the upstream write") + } + select { + case handled := <-firstForwardDone: + if !handled { + t.Fatal("fallback peer was not forwarded") + } + case <-time.After(2 * time.Second): + t.Fatal("fallback dispatch waited for the upstream write") + } + + for range cap(session.sendCh) { + if !manager.ForwardIfActive([]byte("queued"), peer, nil) { + t.Fatal("fallback peer lost sticky classification while its writer was blocked") + } + } + forwardDone := make(chan bool, 1) + go func() { + forwardDone <- manager.ForwardIfActive([]byte("drop"), peer, nil) + }() + select { + case handled := <-forwardDone: + if !handled { + t.Fatal("full fallback queue lost sticky classification") + } + case <-time.After(2 * time.Second): + t.Fatal("full fallback queue blocked ingress") + } + + otherPeer := &net.UDPAddr{IP: net.IPv4(192, 0, 2, 2), Port: 5300} + if processDNS, _ := manager.RouteDNS([]byte("DNS"), otherPeer, nil); !processDNS { + t.Fatal("backpressured fallback peer blocked unrelated DNS routing") + } + + closeDone := make(chan struct{}) + go func() { + manager.Close() + close(closeDone) + }() + select { + case <-closeDone: + case <-time.After(time.Second): + t.Fatal("manager Close did not cancel the blocked fallback write") + } +} + +func TestUDPFallbackExpiredDNSPeerForwardsNextNonDNS(t *testing.T) { + upstream := newUDPFallbackTestUpstream(t, false) + listener := newUDPFallbackTestConn(t) + client := newUDPFallbackTestConn(t) + manager := newUDPFallbackManager(udpFallbackTestAddr(t, upstream.conn), nil) + t.Cleanup(manager.Close) + peer := udpFallbackTestAddr(t, client) + key := makeUDPFallbackPeerKey(peer) + + if routeDNS, _ := manager.RouteDNS([]byte("DNS"), peer, listener); !routeDNS { + t.Fatal("new DNS peer should remain DNS") + } + manager.mu.Lock() + state := manager.peers[key] + state.lastSeen = time.Now().Add(-udpFallbackIdleTimeout - time.Second) + manager.mu.Unlock() + + manager.RouteNonDNS([]byte("after idle gap"), peer, listener) + if got := string(receiveUDPFallbackTestPacket(t, upstream).payload); got != "after idle gap" { + t.Fatalf("unexpected packet after DNS peer expiry: got=%q", got) + } +} + +func TestUDPFallbackInitialDialFailureKeepsFallbackRoute(t *testing.T) { + upstream := newUDPFallbackTestUpstream(t, false) + listener := newUDPFallbackTestConn(t) + client := newUDPFallbackTestConn(t) + manager := newUDPFallbackManager(udpFallbackTestAddr(t, upstream.conn), nil) + t.Cleanup(manager.Close) + peer := udpFallbackTestAddr(t, client) + key := makeUDPFallbackPeerKey(peer) + + realDialUDP := manager.dialUDP + dialAttempts := 0 + manager.dialUDP = func(network string, localAddr, remoteAddr *net.UDPAddr) (*net.UDPConn, error) { + dialAttempts++ + if dialAttempts == 1 { + return nil, errors.New("forced initial dial failure") + } + return realDialUDP(network, localAddr, remoteAddr) + } + + manager.RouteNonDNS([]byte("dropped while dialing"), peer, listener) + manager.mu.Lock() + state := manager.peers[key] + manager.mu.Unlock() + if state == nil || state.mode != udpFallbackRouteFallback || state.session != nil { + t.Fatalf("dial failure lost fallback routing state: %+v", state) + } + + dnsPacket := buildTestDNSQuery(0x7474, "example.org", Enums.DNS_RECORD_TYPE_A) + if routeDNS, _ := manager.RouteDNS(dnsPacket, peer, listener); routeDNS { + t.Fatal("DNS-shaped retry escaped fallback after initial dial failure") + } + if got := receiveUDPFallbackTestPacket(t, upstream).payload; !bytes.Equal(got, dnsPacket) { + t.Fatalf("unexpected recovered fallback packet: got=%x want=%x", got, dnsPacket) + } + if dialAttempts != 2 { + t.Fatalf("unexpected dial attempts: got=%d want=2", dialAttempts) + } +} + +func TestUDPFallbackRedialFailureKeepsFallbackRoute(t *testing.T) { + upstream := newUDPFallbackTestUpstream(t, false) + listener := newUDPFallbackTestConn(t) + client := newUDPFallbackTestConn(t) + manager := newUDPFallbackManager(udpFallbackTestAddr(t, upstream.conn), nil) + t.Cleanup(manager.Close) + peer := udpFallbackTestAddr(t, client) + key := makeUDPFallbackPeerKey(peer) + + realDialUDP := manager.dialUDP + dialAttempts := 0 + manager.dialUDP = func(network string, localAddr, remoteAddr *net.UDPAddr) (*net.UDPConn, error) { + dialAttempts++ + if dialAttempts == 2 { + return nil, errors.New("forced redial failure") + } + return realDialUDP(network, localAddr, remoteAddr) + } + + manager.RouteNonDNS([]byte("activate"), peer, listener) + receiveUDPFallbackTestPacket(t, upstream) + manager.mu.Lock() + oldSession := manager.peers[key].session + manager.mu.Unlock() + oldSession.close() + waitUDPFallbackTestDone(t, oldSession.done, "closed reply loop") + + if !manager.ForwardIfActive([]byte("dropped while redialing"), peer, listener) { + t.Fatal("redial failure lost active fallback classification") + } + manager.mu.Lock() + state := manager.peers[key] + manager.mu.Unlock() + if state == nil || state.mode != udpFallbackRouteFallback || state.session != nil { + t.Fatalf("redial failure lost fallback routing state: %+v", state) + } + + dnsPacket := buildTestDNSQuery(0x7575, "example.org", Enums.DNS_RECORD_TYPE_A) + if routeDNS, _ := manager.RouteDNS(dnsPacket, peer, listener); routeDNS { + t.Fatal("DNS-shaped retry escaped fallback after redial failure") + } + if got := receiveUDPFallbackTestPacket(t, upstream).payload; !bytes.Equal(got, dnsPacket) { + t.Fatalf("unexpected packet after successful redial: got=%x want=%x", got, dnsPacket) + } + if dialAttempts != 3 { + t.Fatalf("unexpected dial attempts: got=%d want=3", dialAttempts) + } +} + +func TestUDPFallbackExpiryRecreationAndClose(t *testing.T) { + upstream := newUDPFallbackTestUpstream(t, false) + listener := newUDPFallbackTestConn(t) + client := newUDPFallbackTestConn(t) + manager := newUDPFallbackManager(udpFallbackTestAddr(t, upstream.conn), nil) + peer := udpFallbackTestAddr(t, client) + key := makeUDPFallbackPeerKey(peer) + + manager.RouteNonDNS([]byte("activate"), peer, listener) + receiveUDPFallbackTestPacket(t, upstream) + manager.mu.Lock() + expiredState := manager.peers[key] + expiredState.lastSeen = time.Now().Add(-udpFallbackIdleTimeout - time.Second) + manager.mu.Unlock() + if manager.ForwardIfActive([]byte("expired"), peer, listener) { + t.Fatal("expired session consumed the fast-path packet") + } + if routeDNS, _ := manager.RouteDNS([]byte("DNS after expiry"), peer, listener); !routeDNS { + t.Fatal("DNS should resume after fallback expiry") + } + + manager.mu.Lock() + dnsState := manager.peers[key] + dnsState.lastSeen = time.Now().Add(-udpFallbackIdleTimeout - time.Second) + manager.mu.Unlock() + manager.cleanup(time.Now()) + manager.RouteNonDNS([]byte("reactivate"), peer, listener) + receiveUDPFallbackTestPacket(t, upstream) + manager.mu.Lock() + oldSession := manager.peers[key].session + manager.mu.Unlock() + oldSession.close() + waitUDPFallbackTestDone(t, oldSession.done, "closed reply loop") + + if !manager.ForwardIfActive([]byte("recreated"), peer, listener) { + t.Fatal("dead reply loop lost active fallback classification") + } + receiveUDPFallbackTestPacket(t, upstream) + manager.mu.Lock() + newSession := manager.peers[key].session + manager.mu.Unlock() + if newSession == oldSession { + t.Fatal("dead fallback session was not recreated") + } + + closeDone := make(chan struct{}) + go func() { + manager.Close() + close(closeDone) + }() + waitUDPFallbackTestDone(t, closeDone, "manager Close") + waitUDPFallbackTestDone(t, newSession.done, "active reply loop") + select { + case <-manager.cleanupDone: + default: + t.Fatal("cleanup loop remained active after Close") + } + manager.Close() +} diff --git a/server_config.toml.simple b/server_config.toml.simple index b37d66b0..d3378b8c 100644 --- a/server_config.toml.simple +++ b/server_config.toml.simple @@ -44,6 +44,12 @@ SUPPORTED_DOWNLOAD_COMPRESSION_TYPES = [0, 1, 2, 3] UDP_HOST = "0.0.0.0" UDP_PORT = 53 +# Optional raw UDP fallback target for non-DNS datagrams received on the DNS +# listener. Leave this unset to disable fallback. The value must be HOST:PORT; +# bracket IPv6 addresses, for example "[2001:db8::1]:5353". The target must +# not resolve back to this listener. Fallback uses one ordered ingress reader. +# FALLBACK = "127.0.0.1:5353" + # UDP readers, DNS workers, and front-door request queue are smart-sized # internally. Only override them if you are profiling a specific host. UDP_READERS = 6 From e98e6cac8d21060ad156be2980e1b9c816a53359 Mon Sep 17 00:00:00 2001 From: Mygod Date: Thu, 23 Jul 2026 01:17:26 -0400 Subject: [PATCH 2/2] Address UDP fallback review feedback --- internal/udpserver/server.go | 29 ++++++++++++++-------- internal/udpserver/server_fallback_test.go | 8 +++++- internal/udpserver/server_ingress.go | 13 ++++++++-- internal/udpserver/server_ingress_test.go | 17 +++++++++++++ internal/udpserver/server_runtime.go | 20 ++++++++++----- 5 files changed, 67 insertions(+), 20 deletions(-) diff --git a/internal/udpserver/server.go b/internal/udpserver/server.go index 121ef3ec..897dacd2 100644 --- a/internal/udpserver/server.go +++ b/internal/udpserver/server.go @@ -342,7 +342,14 @@ func (s *Server) Run(ctx context.Context) error { if fallbackAddr != nil { for _, conn := range conns { localAddr, ok := conn.LocalAddr().(*net.UDPAddr) - if ok && fallbackTargetsListener(localAddr, fallbackAddr) { + if !ok { + continue + } + targetsListener, err := fallbackTargetsListener(localAddr, fallbackAddr) + if err != nil { + return fmt.Errorf("validate fallback address %s against UDP listener %s: %w", fallbackAddr, localAddr, err) + } + if targetsListener { return fmt.Errorf("fallback address %s resolves to the UDP listener and would loop", fallbackAddr) } } @@ -411,29 +418,29 @@ func (s *Server) Run(ctx context.Context) error { } } -func fallbackTargetsListener(listener *net.UDPAddr, target *net.UDPAddr) bool { +func fallbackTargetsListener(listener *net.UDPAddr, target *net.UDPAddr) (bool, error) { if listener == nil || target == nil || listener.Port != target.Port { - return false + return false, nil } if target.IP.IsUnspecified() { - return true + return true, nil } if listener.IP.Equal(target.IP) && listener.Zone == target.Zone { - return true + return true, nil } if !listener.IP.IsUnspecified() { - return false + return false, nil } if listener.IP.To4() != nil && target.IP.To4() == nil { - return false + return false, nil } if target.IP.IsLoopback() { - return true + return true, nil } interfaceAddrs, err := net.InterfaceAddrs() if err != nil { - return false + return false, fmt.Errorf("enumerate local interface addresses: %w", err) } for _, addr := range interfaceAddrs { var ip net.IP @@ -444,8 +451,8 @@ func fallbackTargetsListener(listener *net.UDPAddr, target *net.UDPAddr) bool { ip = value.IP } if ip != nil && ip.Equal(target.IP) { - return true + return true, nil } } - return false + return false, nil } diff --git a/internal/udpserver/server_fallback_test.go b/internal/udpserver/server_fallback_test.go index 828ba28d..020d84f4 100644 --- a/internal/udpserver/server_fallback_test.go +++ b/internal/udpserver/server_fallback_test.go @@ -60,7 +60,11 @@ func TestFallbackTargetsListener(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - if got := fallbackTargetsListener(tt.listener, tt.target); got != tt.want { + got, err := fallbackTargetsListener(tt.listener, tt.target) + if err != nil { + t.Fatalf("fallbackTargetsListener() failed: %v", err) + } + if got != tt.want { t.Fatalf("fallbackTargetsListener()=%t want=%t", got, tt.want) } }) @@ -413,6 +417,7 @@ func TestDNSWorkerFallbackForwardsNonDNSDatagrams(t *testing.T) { incompleteQuestions := buildTestDNSQuery(0x7373, "example.org", Enums.DNS_RECORD_TYPE_A) binary.BigEndian.PutUint16(incompleteQuestions[4:6], 2) + trailingData := append(buildTestDNSQuery(0x7474, "example.org", Enums.DNS_RECORD_TYPE_A), 0) for _, tt := range []struct { name string @@ -420,6 +425,7 @@ func TestDNSWorkerFallbackForwardsNonDNSDatagrams(t *testing.T) { }{ {name: "empty question", packet: emptyQuestion}, {name: "incomplete declared questions", packet: incompleteQuestions}, + {name: "trailing data", packet: trailingData}, } { t.Run(tt.name, func(t *testing.T) { echo := startFallbackIntegrationEcho(t) diff --git a/internal/udpserver/server_ingress.go b/internal/udpserver/server_ingress.go index 99b7071d..fec6478c 100644 --- a/internal/udpserver/server_ingress.go +++ b/internal/udpserver/server_ingress.go @@ -8,6 +8,7 @@ package udpserver import ( + "errors" "fmt" "time" @@ -18,8 +19,16 @@ import ( ) func (s *Server) handlePacket(packet []byte) []byte { - parsed, err := DnsParser.ParseDNSDatagramLite(packet) - return s.handleParsedPacket(packet, parsed, err) + parsed, err := DnsParser.ParseDNSRequestLite(packet) + if err != nil { + if errors.Is(err, DnsParser.ErrNotDNSRequest) || errors.Is(err, DnsParser.ErrPacketTooShort) { + return nil + } + + return s.buildNoDataResponseLogged(packet, "request-parse-failed") + } + + return s.handleParsedPacket(packet, parsed, nil) } func (s *Server) handleParsedPacket(packet []byte, parsed DnsParser.LitePacket, err error) []byte { diff --git a/internal/udpserver/server_ingress_test.go b/internal/udpserver/server_ingress_test.go index 49944dfa..285f202c 100644 --- a/internal/udpserver/server_ingress_test.go +++ b/internal/udpserver/server_ingress_test.go @@ -77,6 +77,23 @@ func TestHandlePacketKeepsUnsupportedAllowedAQueryAsNoData(t *testing.T) { } } +func TestSafeHandlePacketPreservesRequestParsingWithoutFallback(t *testing.T) { + server := &Server{ + domainMatcher: domainMatcher.New([]string{"vpn.example.com"}, 3), + } + request := append(buildTestDNSQuery(0x6262, "example.org", Enums.DNS_RECORD_TYPE_A), 0) + + response := server.safeHandlePacket(request) + if response == nil { + t.Fatal("expected DNS response for request with trailing data, got nil") + } + + flags := binary.BigEndian.Uint16(response[2:4]) + if got := flags & 0x000F; got != Enums.DNSR_CODE_NAME_ERROR { + t.Fatalf("unexpected rcode: got=%d want=%d", got, Enums.DNSR_CODE_NAME_ERROR) + } +} + func TestHandlePacketDropsNonRequestDatagrams(t *testing.T) { server := &Server{} response := buildTestDNSQuery(0x6262, "example.org", Enums.DNS_RECORD_TYPE_A) diff --git a/internal/udpserver/server_runtime.go b/internal/udpserver/server_runtime.go index a57de372..33201e4a 100644 --- a/internal/udpserver/server_runtime.go +++ b/internal/udpserver/server_runtime.go @@ -282,12 +282,20 @@ func (s *Server) dnsWorker(ctx context.Context, conn *net.UDPConn, reqCh <-chan } } -func (s *Server) safeHandlePacket(packet []byte) []byte { - parsed, parseErr, ok := s.safeParseDNSDatagram(packet) - if !ok { - return nil - } - return s.safeHandleParsedPacket(packet, parsed, parseErr) +func (s *Server) safeHandlePacket(packet []byte) (response []byte) { + defer func() { + if recovered := recover(); recovered != nil { + if s.log != nil { + s.log.Errorf( + "\U0001F4A5 Packet Handler Panic Recovered, %v", + recovered, + ) + } + response = nil + } + }() + + return s.handlePacket(packet) } func (s *Server) safeParseDNSDatagram(packet []byte) (parsed DnsParser.LitePacket, parseErr error, ok bool) {