diff --git a/go.mod b/go.mod index 55ef3cf..cb0e935 100644 --- a/go.mod +++ b/go.mod @@ -27,6 +27,7 @@ require ( ) require ( + github.com/ncruces/go-dns v1.3.2 github.com/pelletier/go-toml/v2 v2.2.4 github.com/pires/go-proxyproto v0.11.0 github.com/things-go/go-socks5 v0.1.0 diff --git a/go.sum b/go.sum index 7d99c17..b61ae52 100644 --- a/go.sum +++ b/go.sum @@ -51,6 +51,8 @@ github.com/miekg/dns v1.1.51 h1:0+Xg7vObnhrz/4ZCZcZh7zPXlmU0aveS2HDBd0m0qSo= github.com/miekg/dns v1.1.51/go.mod h1:2Z9d3CP1LQWihRZUf29mQ19yDThaI4DAYzte2CaQW5c= github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA= github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ= +github.com/ncruces/go-dns v1.3.2 h1:kBLuUZBgkQ4qF4WDXZRQ4rG0Gk6sLVJQ5tESkWrxUa0= +github.com/ncruces/go-dns v1.3.2/go.mod h1:tuzixNY8PY/M7yUzcvRbUaeLs3ifIdydpi5H2bfRU+s= github.com/panjf2000/ants/v2 v2.11.5 h1:a7LMnMEeux/ebqTux140tRiaqcFTV0q2bEHF03nl6Rg= github.com/panjf2000/ants/v2 v2.11.5/go.mod h1:8u92CYMUc6gyvTIw8Ru7Mt7+/ESnJahz5EVtqfrilek= github.com/patrickmn/go-cache v2.1.0+incompatible h1:HRMgzkcYKYpi3C8ajMPV8OFXaaRUnok+kx1WdO15EQc= diff --git a/network/v2/dns.go b/network/v2/dns.go new file mode 100644 index 0000000..0728609 --- /dev/null +++ b/network/v2/dns.go @@ -0,0 +1,53 @@ +package network + +import ( + "context" + "fmt" + "net" + "net/url" + "time" + + "github.com/ncruces/go-dns" +) + +var dnsCacheOptions = []dns.CacheOption{ + dns.MaxCacheEntries(dns.DefaultMaxCacheEntries), + dns.MaxCacheTTL(time.Hour), + dns.NegativeCache(false), +} + +func GetDNS(u *url.URL) (*net.Resolver, error) { + if u == nil { + return dns.NewCachingResolver(nil, dnsCacheOptions...), nil + } + + switch u.Scheme { + case "tls": + return dns.NewDoTResolver(u.Host, dns.DoTCache(dnsCacheOptions...)) + case "https": + if u.Path == "" { + u.Path = "/dns-query" + } + + return dns.NewDoHResolver(u.String(), dns.DoHCache(dnsCacheOptions...)) + case "udp", "": + default: + return nil, fmt.Errorf("unsupported DNS %v", u) + } + + port := u.Port() + if port == "" { + port = "53" + } + + hostport := net.JoinHostPort(u.Hostname(), port) + dialer := &net.Dialer{} + resolver := &net.Resolver{ + PreferGo: true, + Dial: func(ctx context.Context, network, address string) (net.Conn, error) { + return dialer.DialContext(ctx, "udp", hostport) + }, + } + + return dns.NewCachingResolver(resolver, dnsCacheOptions...), nil +} diff --git a/network/v2/dns_test.go b/network/v2/dns_test.go new file mode 100644 index 0000000..7a05eca --- /dev/null +++ b/network/v2/dns_test.go @@ -0,0 +1,68 @@ +package network_test + +import ( + "context" + "net" + "net/url" + "testing" + "time" + + "github.com/9seconds/mtg/v2/network/v2" + "github.com/stretchr/testify/require" + "github.com/stretchr/testify/suite" +) + +type DNSTestSuite struct { + suite.Suite +} + +func (suite *DNSTestSuite) TestDefault() { + resolver, err := network.GetDNS(nil) + suite.NoError(err) + suite.doTest(resolver) +} + +func (suite *DNSTestSuite) TestDoH() { + for _, addr := range []string{"1.1.1.1", "cloudflare-dns.com"} { + suite.Run(addr, func() { + u, err := url.Parse("https://" + addr) + require.NoError(suite.T(), err) + + resolver, err := network.GetDNS(u) + suite.NoError(err) + suite.doTest(resolver) + }) + } +} + +func (suite *DNSTestSuite) TestDoT() { + u, err := url.Parse("tls://dns.google") + require.NoError(suite.T(), err) + + resolver, err := network.GetDNS(u) + suite.NoError(err) + suite.doTest(resolver) +} + +func (suite *DNSTestSuite) TestUDP() { + u, err := url.Parse("8.8.8.8") + require.NoError(suite.T(), err) + + resolver, err := network.GetDNS(u) + suite.NoError(err) + suite.doTest(resolver) +} + +func (suite *DNSTestSuite) doTest(resolver *net.Resolver) { + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + ips, err := resolver.LookupIP(ctx, "ip4", "dns.google") + suite.NoError(err) + suite.Greater(len(ips), 0) +} + +func TestGetDNS(t *testing.T) { + t.Parallel() + suite.Run(t, &DNSTestSuite{}) +}