Implement caching dns resolver

This commit is contained in:
9seconds
2021-03-30 14:14:37 +03:00
parent 841a4d2227
commit 10b78322a3
16 changed files with 163 additions and 47 deletions
+88
View File
@@ -0,0 +1,88 @@
package network
import (
"net/http"
"time"
doh "github.com/babolivier/go-doh-client"
"github.com/dgraph-io/ristretto"
)
const (
dnsResolverSize = 1024 * 1024 // 1mb
dnsResolverKeepTime = 10 * time.Minute
)
type dnsResolver struct {
resolver doh.Resolver
cache *ristretto.Cache
}
func (d dnsResolver) LookupA(hostname string) []string {
key := "\x00." + hostname
if value, ok := d.cache.Get(key); ok {
return value.([]string)
}
var ips []string
if recs, _, err := d.resolver.LookupA(hostname); err == nil {
for _, v := range recs {
ips = append(ips, v.IP4)
}
d.cache.SetWithTTL(key, ips, 0, dnsResolverKeepTime)
}
return ips
}
func (d dnsResolver) LookupAAAA(hostname string) []string {
key := "\x01." + hostname
if value, ok := d.cache.Get(key); ok {
return value.([]string)
}
var ips []string
if recs, _, err := d.resolver.LookupAAAA(hostname); err == nil {
for _, v := range recs {
ips = append(ips, v.IP6)
}
d.cache.SetWithTTL(key, ips, 0, dnsResolverKeepTime)
}
return ips
}
func newDNSResolver(hostname string, httpClient *http.Client) dnsResolver {
cache, err := ristretto.NewCache(&ristretto.Config{
NumCounters: 10 * dnsResolverSize, // nolint: gomnd // taken from official doc as a best practice value
MaxCost: dnsResolverSize,
BufferItems: 64, // nolint: gomnd // taken from official doc as a best practice value
Cost: func(value interface{}) int64 {
var cost int64
for _, v := range value.([]string) {
cost += int64(len([]byte(v)))
}
return cost
},
})
if err != nil {
panic(err)
}
return dnsResolver{
resolver: doh.Resolver{
Host: hostname,
Class: doh.IN,
HTTPClient: httpClient,
},
cache: cache,
}
}
+50
View File
@@ -0,0 +1,50 @@
package network
import (
"net"
"net/http"
"testing"
"time"
"github.com/stretchr/testify/suite"
)
type DNSResolverTestSuite struct {
suite.Suite
d dnsResolver
}
func (suite *DNSResolverTestSuite) TestLookupA() {
suite.d.LookupA("google.com")
time.Sleep(10 * time.Millisecond)
addrs := suite.d.LookupA("google.com")
for _, v := range addrs {
suite.NotEmpty(v)
suite.NotNil(net.ParseIP(v).To4())
}
}
func (suite *DNSResolverTestSuite) TestLookupAAAA() {
suite.d.LookupAAAA("google.com")
time.Sleep(10 * time.Millisecond)
addrs := suite.d.LookupAAAA("google.com")
for _, v := range addrs {
suite.NotEmpty(v)
suite.Nil(net.ParseIP(v).To4())
suite.NotNil(net.ParseIP(v).To16())
}
}
func (suite *DNSResolverTestSuite) SetupTest() {
suite.d = newDNSResolver("1.1.1.1", &http.Client{})
}
func TestDNSResolver(t *testing.T) {
t.Parallel()
suite.Run(t, &DNSResolverTestSuite{})
}
+11 -21
View File
@@ -10,7 +10,6 @@ import (
"time"
"github.com/9seconds/mtg/v2/mtglib"
doh "github.com/babolivier/go-doh-client"
)
type networkHTTPTransport struct {
@@ -26,7 +25,7 @@ func (n networkHTTPTransport) RoundTrip(req *http.Request) (*http.Response, erro
type network struct {
dialer Dialer
dns doh.Resolver
dns dnsResolver
httpTimeout time.Duration
userAgent string
}
@@ -84,14 +83,11 @@ func (n *network) dnsResolve(protocol, address string) ([]string, error) {
go func() {
defer wg.Done()
if recs, _, err := n.dns.LookupA(address); err == nil {
mutex.Lock()
defer mutex.Unlock()
resolved := n.dns.LookupA(address)
for _, v := range recs {
ips = append(ips, v.IP4)
}
}
mutex.Lock()
ips = append(ips, resolved...)
mutex.Unlock()
}()
}
@@ -102,14 +98,11 @@ func (n *network) dnsResolve(protocol, address string) ([]string, error) {
go func() {
defer wg.Done()
if recs, _, err := n.dns.LookupAAAA(address); err == nil {
mutex.Lock()
defer mutex.Unlock()
resolved := n.dns.LookupAAAA(address)
for _, v := range recs {
ips = append(ips, v.IP6)
}
}
mutex.Lock()
ips = append(ips, resolved...)
mutex.Unlock()
}()
}
@@ -140,11 +133,8 @@ func NewNetwork(dialer Dialer,
dialer: dialer,
httpTimeout: httpTimeout,
userAgent: userAgent,
dns: doh.Resolver{
Host: dohHostname,
Class: doh.IN,
HTTPClient: makeHTTPClient(userAgent, DNSTimeout, dialer.DialContext),
},
dns: newDNSResolver(dohHostname,
makeHTTPClient(userAgent, DNSTimeout, dialer.DialContext)),
}, nil
}