dns.go 6.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238
  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. // dnsQueryTypesFor asks only for families the tunnel can dial, so a v4-only tunnel
  120. // never caches an unroutable AAAA answer (#6570). No addresses keeps A then AAAA.
  121. func dnsQueryTypesFor(addrs []netip.Addr) []dnsmessage.Type {
  122. hasV4 := deviceHasV4(addrs)
  123. hasV6 := deviceHasV6(addrs)
  124. switch {
  125. case hasV4 && hasV6:
  126. return []dnsmessage.Type{dnsmessage.TypeA, dnsmessage.TypeAAAA}
  127. case hasV6:
  128. return []dnsmessage.Type{dnsmessage.TypeAAAA}
  129. case hasV4:
  130. return []dnsmessage.Type{dnsmessage.TypeA}
  131. default:
  132. return []dnsmessage.Type{dnsmessage.TypeA, dnsmessage.TypeAAAA}
  133. }
  134. }
  135. // tunnelSupportsAddr reports whether the device stack has a local address in
  136. // the same family as ip (IPv4-mapped IPv6 counts as IPv4).
  137. func tunnelSupportsAddr(addrs []netip.Addr, ip netip.Addr) bool {
  138. if !ip.IsValid() {
  139. return false
  140. }
  141. if ip.Is4() || ip.Is4In6() {
  142. return deviceHasV4(addrs)
  143. }
  144. return deviceHasV6(addrs)
  145. }
  146. // exchangeTunnelDNSWithFallback returns the first answer among the families the
  147. // device stack can route.
  148. func exchangeTunnelDNSWithFallback(ctx context.Context, conn *gonet.UDPConn, addrs []netip.Addr, host string) (netip.Addr, error) {
  149. types := dnsQueryTypesFor(addrs)
  150. var firstErr error
  151. for _, qType := range types {
  152. addr, err := exchangeTunnelDNSQuery(ctx, conn, host, qType)
  153. if err == nil {
  154. return addr, nil
  155. }
  156. if firstErr == nil {
  157. firstErr = err
  158. }
  159. }
  160. return netip.Addr{}, firstErr
  161. }
  162. func exchangeTunnelDNSQuery(ctx context.Context, conn *gonet.UDPConn, host string, qType dnsmessage.Type) (netip.Addr, error) {
  163. name, err := dnsmessage.NewName(host + ".")
  164. if err != nil {
  165. return netip.Addr{}, fmt.Errorf("dns name %q: %w", host, err)
  166. }
  167. id := uint16(rand.Intn(1 << 16))
  168. query := dnsmessage.Message{
  169. Header: dnsmessage.Header{ID: id, RecursionDesired: true},
  170. Questions: []dnsmessage.Question{{
  171. Name: name,
  172. Type: qType,
  173. Class: dnsmessage.ClassINET,
  174. }},
  175. }
  176. wire, err := query.Pack()
  177. if err != nil {
  178. return netip.Addr{}, fmt.Errorf("dns pack %q: %w", host, err)
  179. }
  180. buf := make([]byte, 512)
  181. for attempt := 0; attempt < tunnelDNSAttempts; attempt++ {
  182. select {
  183. case <-ctx.Done():
  184. return netip.Addr{}, ctx.Err()
  185. default:
  186. }
  187. if _, werr := conn.Write(wire); werr != nil {
  188. return netip.Addr{}, fmt.Errorf("dns send %q: %w", host, werr)
  189. }
  190. if derr := conn.SetReadDeadline(time.Now().Add(tunnelDNSPacketTimeout)); derr != nil {
  191. return netip.Addr{}, fmt.Errorf("dns deadline %q: %w", host, derr)
  192. }
  193. for {
  194. n, rerr := conn.Read(buf)
  195. if rerr != nil {
  196. break // per-attempt timeout -> next attempt
  197. }
  198. var resp dnsmessage.Message
  199. if uerr := resp.Unpack(buf[:n]); uerr != nil || resp.ID != id {
  200. continue
  201. }
  202. for _, ans := range resp.Answers {
  203. if a, ok := ans.Body.(*dnsmessage.AResource); ok && qType == dnsmessage.TypeA {
  204. return netip.AddrFrom4(a.A), nil
  205. }
  206. if aaaa, ok := ans.Body.(*dnsmessage.AAAAResource); ok && qType == dnsmessage.TypeAAAA {
  207. return netip.AddrFrom16(aaaa.AAAA), nil
  208. }
  209. }
  210. return netip.Addr{}, fmt.Errorf("dns %q (type %v): rcode=%d answers=%d", host, qType, resp.RCode, len(resp.Answers))
  211. }
  212. }
  213. return netip.Addr{}, fmt.Errorf("dns lookup %q (type %v): no answer after %d attempts", host, qType, tunnelDNSAttempts)
  214. }