mirror of
https://github.com/ScuroNeko/mtg.git
synced 2026-08-31 11:44:02 +03:00
Implement caching dns resolver
This commit is contained in:
@@ -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,
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user