dns.go 6.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217
  1. package amneziawgnet
  2. import (
  3. "context"
  4. "fmt"
  5. "math/rand"
  6. "net/netip"
  7. "strings"
  8. "sync"
  9. "time"
  10. "golang.org/x/net/dns/dnsmessage"
  11. "gvisor.dev/gvisor/pkg/tcpip"
  12. "gvisor.dev/gvisor/pkg/tcpip/adapters/gonet"
  13. "github.com/mhsanaei/3x-ui/v3/internal/amneziawg"
  14. "github.com/mhsanaei/3x-ui/v3/internal/logger"
  15. )
  16. // DefaultTunnelDNSServer resolves domain targets through outbound netstack.
  17. const (
  18. DefaultTunnelDNSServer = "1.1.1.1:53"
  19. DefaultTunnelDNSServerV6 = "[2606:4700:4700::1111]:53"
  20. )
  21. func deviceHasV4(addrs []netip.Addr) bool {
  22. for _, a := range addrs {
  23. if a.Is4() {
  24. return true
  25. }
  26. }
  27. return false
  28. }
  29. func deviceHasV6(addrs []netip.Addr) bool {
  30. for _, a := range addrs {
  31. if a.Is6() && !a.Is4In6() {
  32. return true
  33. }
  34. }
  35. return false
  36. }
  37. // defaultDNSFor picks a resolver matching the tunnel address family:
  38. // IPv4 default (or empty), or IPv6 default when IPv6-only.
  39. func defaultDNSFor(addrs []netip.Addr) string {
  40. if deviceHasV4(addrs) || len(addrs) == 0 {
  41. return DefaultTunnelDNSServer
  42. }
  43. return DefaultTunnelDNSServerV6
  44. }
  45. const (
  46. // tunnelResolveTimeout bounds one lookup inside a live connection handler.
  47. tunnelResolveTimeout = 4 * time.Second
  48. tunnelDNSPacketTimeout = 1200 * time.Millisecond
  49. tunnelDNSAttempts = 3
  50. )
  51. type tunnelDNSCacheEntry struct {
  52. addr netip.Addr
  53. exp time.Time
  54. }
  55. var tunnelDNSCache = struct {
  56. mu sync.Mutex
  57. m map[string]tunnelDNSCacheEntry
  58. }{m: map[string]tunnelDNSCacheEntry{}}
  59. const (
  60. tunnelDNSCacheTTL = 60 * time.Second
  61. tunnelDNSCacheMaxSize = 1024
  62. )
  63. // dnsCacheKey computes cache key scoped by outbound tag, server, and host.
  64. func dnsCacheKey(tag, dnsServer, host string) string {
  65. return tag + "|" + dnsServer + "|" + host
  66. }
  67. func resolveTunnelVia(ctx context.Context, dev *Device, tag string, dnsServer string, host string) (netip.Addr, error) {
  68. normDNS := amneziawg.NormalizeDNSServer(dnsServer)
  69. if normDNS == "" {
  70. normDNS = defaultDNSFor(dev.LocalAddresses())
  71. }
  72. key := dnsCacheKey(tag, normDNS, host)
  73. now := time.Now()
  74. tunnelDNSCache.mu.Lock()
  75. if e, ok := tunnelDNSCache.m[key]; ok && now.Before(e.exp) {
  76. tunnelDNSCache.mu.Unlock()
  77. return e.addr, nil
  78. }
  79. tunnelDNSCache.mu.Unlock()
  80. server, err := netip.ParseAddrPort(normDNS)
  81. if err != nil {
  82. return netip.Addr{}, fmt.Errorf("bad tunnel DNS server %q: %w", normDNS, err)
  83. }
  84. raddr := tcpip.FullAddress{
  85. NIC: 1,
  86. Addr: tcpip.AddrFromSlice(server.Addr().AsSlice()),
  87. Port: server.Port(),
  88. }
  89. conn, derr := gonet.DialUDP(dev.Stack, nil, &raddr, tunnelNetwork(server.Addr()))
  90. if derr != nil {
  91. logger.Warningf("amneziawgnet: resolveTunnel tag=%q host=%q server=%s localAddrs=%v err=%v", tag, host, server, dev.LocalAddresses(), derr)
  92. return netip.Addr{}, fmt.Errorf("dns dial %s: %w", server, derr)
  93. }
  94. defer conn.Close()
  95. addr, rerr := exchangeTunnelDNSWithFallback(ctx, conn, dev.LocalAddresses(), host)
  96. if rerr != nil {
  97. return netip.Addr{}, rerr
  98. }
  99. tunnelDNSCache.mu.Lock()
  100. if len(tunnelDNSCache.m) >= tunnelDNSCacheMaxSize {
  101. tunnelDNSCache.m = map[string]tunnelDNSCacheEntry{}
  102. }
  103. tunnelDNSCache.m[key] = tunnelDNSCacheEntry{addr: addr, exp: now.Add(tunnelDNSCacheTTL)}
  104. tunnelDNSCache.mu.Unlock()
  105. logger.Debugf("amneziawgnet: resolved tag=%q %q -> %s via tunnel", tag, host, addr)
  106. return addr, nil
  107. }
  108. // flushTunnelDNSCacheForTag purges all cached DNS entries for an outbound tag.
  109. func flushTunnelDNSCacheForTag(tag string) {
  110. tunnelDNSCache.mu.Lock()
  111. defer tunnelDNSCache.mu.Unlock()
  112. prefix := tag + "|"
  113. for k := range tunnelDNSCache.m {
  114. if strings.HasPrefix(k, prefix) {
  115. delete(tunnelDNSCache.m, k)
  116. }
  117. }
  118. }
  119. // exchangeTunnelDNSWithFallback queries A and/or AAAA depending on the local
  120. // address families configured on the device stack.
  121. func exchangeTunnelDNSWithFallback(ctx context.Context, conn *gonet.UDPConn, addrs []netip.Addr, host string) (netip.Addr, error) {
  122. hasV4 := deviceHasV4(addrs)
  123. hasV6 := deviceHasV6(addrs)
  124. // If the tunnel is IPv6-only, query AAAA first; else query A first.
  125. types := []dnsmessage.Type{dnsmessage.TypeA, dnsmessage.TypeAAAA}
  126. if hasV6 && !hasV4 {
  127. types = []dnsmessage.Type{dnsmessage.TypeAAAA, dnsmessage.TypeA}
  128. }
  129. var firstErr error
  130. for _, qType := range types {
  131. // Skip AAAA if device has no IPv6 capability and has IPv4, unless A failed.
  132. addr, err := exchangeTunnelDNSQuery(ctx, conn, host, qType)
  133. if err == nil {
  134. return addr, nil
  135. }
  136. if firstErr == nil {
  137. firstErr = err
  138. }
  139. }
  140. return netip.Addr{}, firstErr
  141. }
  142. func exchangeTunnelDNSQuery(ctx context.Context, conn *gonet.UDPConn, host string, qType dnsmessage.Type) (netip.Addr, error) {
  143. name, err := dnsmessage.NewName(host + ".")
  144. if err != nil {
  145. return netip.Addr{}, fmt.Errorf("dns name %q: %w", host, err)
  146. }
  147. id := uint16(rand.Intn(1 << 16))
  148. query := dnsmessage.Message{
  149. Header: dnsmessage.Header{ID: id, RecursionDesired: true},
  150. Questions: []dnsmessage.Question{{
  151. Name: name,
  152. Type: qType,
  153. Class: dnsmessage.ClassINET,
  154. }},
  155. }
  156. wire, err := query.Pack()
  157. if err != nil {
  158. return netip.Addr{}, fmt.Errorf("dns pack %q: %w", host, err)
  159. }
  160. buf := make([]byte, 512)
  161. for attempt := 0; attempt < tunnelDNSAttempts; attempt++ {
  162. select {
  163. case <-ctx.Done():
  164. return netip.Addr{}, ctx.Err()
  165. default:
  166. }
  167. if _, werr := conn.Write(wire); werr != nil {
  168. return netip.Addr{}, fmt.Errorf("dns send %q: %w", host, werr)
  169. }
  170. if derr := conn.SetReadDeadline(time.Now().Add(tunnelDNSPacketTimeout)); derr != nil {
  171. return netip.Addr{}, fmt.Errorf("dns deadline %q: %w", host, derr)
  172. }
  173. for {
  174. n, rerr := conn.Read(buf)
  175. if rerr != nil {
  176. break // per-attempt timeout -> next attempt
  177. }
  178. var resp dnsmessage.Message
  179. if uerr := resp.Unpack(buf[:n]); uerr != nil || resp.ID != id {
  180. continue
  181. }
  182. for _, ans := range resp.Answers {
  183. if a, ok := ans.Body.(*dnsmessage.AResource); ok && qType == dnsmessage.TypeA {
  184. return netip.AddrFrom4(a.A), nil
  185. }
  186. if aaaa, ok := ans.Body.(*dnsmessage.AAAAResource); ok && qType == dnsmessage.TypeAAAA {
  187. return netip.AddrFrom16(aaaa.AAAA), nil
  188. }
  189. }
  190. return netip.Addr{}, fmt.Errorf("dns %q (type %v): rcode=%d answers=%d", host, qType, resp.RCode, len(resp.Answers))
  191. }
  192. }
  193. return netip.Addr{}, fmt.Errorf("dns lookup %q (type %v): no answer after %d attempts", host, qType, tunnelDNSAttempts)
  194. }