Files
3x-ui/internal/amneziawgnet/dns.go
T
mrchatam b3a5be9da4 fix(amneziawg): stop AAAA fallback on v4-only tunnels and expose I2–I5 (#6611)
* fix(amneziawg): stop AAAA fallback on v4-only tunnels and expose I2–I5

Gate tunnel DNS queries to address families the device can actually dial,
reject undialable literal IPs early, and surface I2–I5 on the outbound form.

Fixes #6570

* ci: retrigger frontend after npm registry maintenance

The frontend job failed solely on `npm audit` while registry.npmjs.org
returned 503 (Service Under Maintenance). Lint, typecheck, vitest, vite
build, and storybook all passed. Local `npm audit --omit=dev
--audit-level=high` now reports 0 vulnerabilities.

* style(amneziawg): keep the tunnel DNS family comments to two lines

CLAUDE.md caps a comment block at two lines.

---------

Co-authored-by: mrchatam <mrchatam@users.noreply.github.com>
Co-authored-by: MHSanaei <ho3ein.sanaei@gmail.com>
2026-09-26 22:01:04 +02:00

239 lines
6.7 KiB
Go

package amneziawgnet
import (
"context"
"fmt"
"math/rand"
"net/netip"
"strings"
"sync"
"time"
"golang.org/x/net/dns/dnsmessage"
"gvisor.dev/gvisor/pkg/tcpip"
"gvisor.dev/gvisor/pkg/tcpip/adapters/gonet"
"github.com/mhsanaei/3x-ui/v3/internal/amneziawg"
"github.com/mhsanaei/3x-ui/v3/internal/logger"
)
// DefaultTunnelDNSServer resolves domain targets through outbound netstack.
const (
DefaultTunnelDNSServer = "1.1.1.1:53"
DefaultTunnelDNSServerV6 = "[2606:4700:4700::1111]:53"
)
func deviceHasV4(addrs []netip.Addr) bool {
for _, a := range addrs {
if a.Is4() {
return true
}
}
return false
}
func deviceHasV6(addrs []netip.Addr) bool {
for _, a := range addrs {
if a.Is6() && !a.Is4In6() {
return true
}
}
return false
}
// defaultDNSFor picks a resolver matching the tunnel address family:
// IPv4 default (or empty), or IPv6 default when IPv6-only.
func defaultDNSFor(addrs []netip.Addr) string {
if deviceHasV4(addrs) || len(addrs) == 0 {
return DefaultTunnelDNSServer
}
return DefaultTunnelDNSServerV6
}
const (
// tunnelResolveTimeout bounds one lookup inside a live connection handler.
tunnelResolveTimeout = 4 * time.Second
tunnelDNSPacketTimeout = 1200 * time.Millisecond
tunnelDNSAttempts = 3
)
type tunnelDNSCacheEntry struct {
addr netip.Addr
exp time.Time
}
var tunnelDNSCache = struct {
mu sync.Mutex
m map[string]tunnelDNSCacheEntry
}{m: map[string]tunnelDNSCacheEntry{}}
const (
tunnelDNSCacheTTL = 60 * time.Second
tunnelDNSCacheMaxSize = 1024
)
// dnsCacheKey computes cache key scoped by outbound tag, server, and host.
func dnsCacheKey(tag, dnsServer, host string) string {
return tag + "|" + dnsServer + "|" + host
}
func resolveTunnelVia(ctx context.Context, dev *Device, tag string, dnsServer string, host string) (netip.Addr, error) {
normDNS := amneziawg.NormalizeDNSServer(dnsServer)
if normDNS == "" {
normDNS = defaultDNSFor(dev.LocalAddresses())
}
key := dnsCacheKey(tag, normDNS, host)
now := time.Now()
tunnelDNSCache.mu.Lock()
if e, ok := tunnelDNSCache.m[key]; ok && now.Before(e.exp) {
tunnelDNSCache.mu.Unlock()
return e.addr, nil
}
tunnelDNSCache.mu.Unlock()
server, err := netip.ParseAddrPort(normDNS)
if err != nil {
return netip.Addr{}, fmt.Errorf("bad tunnel DNS server %q: %w", normDNS, err)
}
raddr := tcpip.FullAddress{
NIC: 1,
Addr: tcpip.AddrFromSlice(server.Addr().AsSlice()),
Port: server.Port(),
}
conn, derr := gonet.DialUDP(dev.Stack, nil, &raddr, tunnelNetwork(server.Addr()))
if derr != nil {
logger.Warningf("amneziawgnet: resolveTunnel tag=%q host=%q server=%s localAddrs=%v err=%v", tag, host, server, dev.LocalAddresses(), derr)
return netip.Addr{}, fmt.Errorf("dns dial %s: %w", server, derr)
}
defer conn.Close()
addr, rerr := exchangeTunnelDNSWithFallback(ctx, conn, dev.LocalAddresses(), host)
if rerr != nil {
return netip.Addr{}, rerr
}
tunnelDNSCache.mu.Lock()
if len(tunnelDNSCache.m) >= tunnelDNSCacheMaxSize {
tunnelDNSCache.m = map[string]tunnelDNSCacheEntry{}
}
tunnelDNSCache.m[key] = tunnelDNSCacheEntry{addr: addr, exp: now.Add(tunnelDNSCacheTTL)}
tunnelDNSCache.mu.Unlock()
logger.Debugf("amneziawgnet: resolved tag=%q %q -> %s via tunnel", tag, host, addr)
return addr, nil
}
// flushTunnelDNSCacheForTag purges all cached DNS entries for an outbound tag.
func flushTunnelDNSCacheForTag(tag string) {
tunnelDNSCache.mu.Lock()
defer tunnelDNSCache.mu.Unlock()
prefix := tag + "|"
for k := range tunnelDNSCache.m {
if strings.HasPrefix(k, prefix) {
delete(tunnelDNSCache.m, k)
}
}
}
// dnsQueryTypesFor asks only for families the tunnel can dial, so a v4-only tunnel
// never caches an unroutable AAAA answer (#6570). No addresses keeps A then AAAA.
func dnsQueryTypesFor(addrs []netip.Addr) []dnsmessage.Type {
hasV4 := deviceHasV4(addrs)
hasV6 := deviceHasV6(addrs)
switch {
case hasV4 && hasV6:
return []dnsmessage.Type{dnsmessage.TypeA, dnsmessage.TypeAAAA}
case hasV6:
return []dnsmessage.Type{dnsmessage.TypeAAAA}
case hasV4:
return []dnsmessage.Type{dnsmessage.TypeA}
default:
return []dnsmessage.Type{dnsmessage.TypeA, dnsmessage.TypeAAAA}
}
}
// tunnelSupportsAddr reports whether the device stack has a local address in
// the same family as ip (IPv4-mapped IPv6 counts as IPv4).
func tunnelSupportsAddr(addrs []netip.Addr, ip netip.Addr) bool {
if !ip.IsValid() {
return false
}
if ip.Is4() || ip.Is4In6() {
return deviceHasV4(addrs)
}
return deviceHasV6(addrs)
}
// exchangeTunnelDNSWithFallback returns the first answer among the families the
// device stack can route.
func exchangeTunnelDNSWithFallback(ctx context.Context, conn *gonet.UDPConn, addrs []netip.Addr, host string) (netip.Addr, error) {
types := dnsQueryTypesFor(addrs)
var firstErr error
for _, qType := range types {
addr, err := exchangeTunnelDNSQuery(ctx, conn, host, qType)
if err == nil {
return addr, nil
}
if firstErr == nil {
firstErr = err
}
}
return netip.Addr{}, firstErr
}
func exchangeTunnelDNSQuery(ctx context.Context, conn *gonet.UDPConn, host string, qType dnsmessage.Type) (netip.Addr, error) {
name, err := dnsmessage.NewName(host + ".")
if err != nil {
return netip.Addr{}, fmt.Errorf("dns name %q: %w", host, err)
}
id := uint16(rand.Intn(1 << 16))
query := dnsmessage.Message{
Header: dnsmessage.Header{ID: id, RecursionDesired: true},
Questions: []dnsmessage.Question{{
Name: name,
Type: qType,
Class: dnsmessage.ClassINET,
}},
}
wire, err := query.Pack()
if err != nil {
return netip.Addr{}, fmt.Errorf("dns pack %q: %w", host, err)
}
buf := make([]byte, 512)
for attempt := 0; attempt < tunnelDNSAttempts; attempt++ {
select {
case <-ctx.Done():
return netip.Addr{}, ctx.Err()
default:
}
if _, werr := conn.Write(wire); werr != nil {
return netip.Addr{}, fmt.Errorf("dns send %q: %w", host, werr)
}
if derr := conn.SetReadDeadline(time.Now().Add(tunnelDNSPacketTimeout)); derr != nil {
return netip.Addr{}, fmt.Errorf("dns deadline %q: %w", host, derr)
}
for {
n, rerr := conn.Read(buf)
if rerr != nil {
break // per-attempt timeout -> next attempt
}
var resp dnsmessage.Message
if uerr := resp.Unpack(buf[:n]); uerr != nil || resp.ID != id {
continue
}
for _, ans := range resp.Answers {
if a, ok := ans.Body.(*dnsmessage.AResource); ok && qType == dnsmessage.TypeA {
return netip.AddrFrom4(a.A), nil
}
if aaaa, ok := ans.Body.(*dnsmessage.AAAAResource); ok && qType == dnsmessage.TypeAAAA {
return netip.AddrFrom16(aaaa.AAAA), nil
}
}
return netip.Addr{}, fmt.Errorf("dns %q (type %v): rcode=%d answers=%d", host, qType, resp.RCode, len(resp.Answers))
}
}
return netip.Addr{}, fmt.Errorf("dns lookup %q (type %v): no answer after %d attempts", host, qType, tunnelDNSAttempts)
}