dns.go 6.1 KB

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