diff --git a/probe/endpoint/reporter.go b/probe/endpoint/reporter.go index 049659b13..8685fa932 100644 --- a/probe/endpoint/reporter.go +++ b/probe/endpoint/reporter.go @@ -26,7 +26,7 @@ type Reporter struct { includeNAT bool conntracker *Conntracker natmapper *natmapper - revResolver *reverseResolver + revResolver *ReverseResolver } // SpyDuration is an exported prometheus metric @@ -65,16 +65,13 @@ func NewReporter(hostID, hostName string, includeProcesses bool, useConntrack bo log.Printf("Failed to start natMapper: %v", err) } } - - revRes := newReverseResolver() - return &Reporter{ hostID: hostID, hostName: hostName, includeProcesses: includeProcesses, conntracker: conntracker, natmapper: natmapper, - revResolver: revRes, + revResolver: NewReverseResolver(), } } @@ -152,7 +149,7 @@ func (r *Reporter) addConnection(rpt *report.Report, localAddr, remoteAddr strin ) // in case we have a reverse resolution for the IP, we can use it for the name... - if revRemoteName, err := r.revResolver.Get(remoteAddr, false); err == nil { + if revRemoteName, err := r.revResolver.Get(remoteAddr); err == nil { remoteNode = remoteNode.AddMetadata(map[string]string{ "name": revRemoteName, }) @@ -191,7 +188,7 @@ func (r *Reporter) addConnection(rpt *report.Report, localAddr, remoteAddr strin ) // in case we have a reverse resolution for the IP, we can use it for the name... - if revRemoteName, err := r.revResolver.Get(remoteAddr, false); err == nil { + if revRemoteName, err := r.revResolver.Get(remoteAddr); err == nil { remoteNode = remoteNode.AddMetadata(map[string]string{ "name": revRemoteName, }) diff --git a/probe/endpoint/resolver.go b/probe/endpoint/resolver.go index c27bef9f4..210e93832 100644 --- a/probe/endpoint/resolver.go +++ b/probe/endpoint/resolver.go @@ -16,25 +16,22 @@ const ( type revResFunc func(addr string) (names []string, err error) -type revResRequest struct { - address string - done chan struct{} -} - // ReverseResolver is a caching, reverse resolver -type reverseResolver struct { - addresses chan revResRequest +type ReverseResolver struct { + addresses chan string cache gcache.Cache - resolver revResFunc + Throttle <-chan time.Time // Made public for mocking + Resolver revResFunc } // NewReverseResolver starts a new reverse resolver that // performs reverse resolutions and caches the result. -func newReverseResolver() *reverseResolver { - r := reverseResolver{ - addresses: make(chan revResRequest, rAddrBacklog), +func NewReverseResolver() *ReverseResolver { + r := ReverseResolver{ + addresses: make(chan string, rAddrBacklog), cache: gcache.New(rAddrCacheLen).LRU().Expiration(rAddrCacheExpiration).Build(), - resolver: net.LookupAddr, + Throttle: time.Tick(time.Second / 10), + Resolver: net.LookupAddr, } go r.loop() return &r @@ -43,43 +40,37 @@ func newReverseResolver() *reverseResolver { // Get the reverse resolution for an IP address if already in the cache, // a gcache.NotFoundKeyError error otherwise. // Note: it returns one of the possible names that can be obtained for that IP. -func (r *reverseResolver) Get(address string, wait bool) (string, error) { +func (r *ReverseResolver) Get(address string) (string, error) { val, err := r.cache.Get(address) if err == nil { return val.(string), nil } if err == gcache.NotFoundKeyError { - request := revResRequest{address: address, done: make(chan struct{})} // we trigger a asynchronous reverse resolution when not cached select { - case r.addresses <- request: - if wait { - <-request.done - } + case r.addresses <- address: default: } } return "", err } -func (r *reverseResolver) loop() { - throttle := time.Tick(time.Second / 10) +func (r *ReverseResolver) loop() { for request := range r.addresses { - <-throttle // rate limit our DNS resolutions + <-r.Throttle // rate limit our DNS resolutions // and check if the answer is already in the cache - if _, err := r.cache.Get(request.address); err == nil { + if _, err := r.cache.Get(request); err == nil { continue } - names, err := r.resolver(request.address) + names, err := r.Resolver(request) if err == nil && len(names) > 0 { name := strings.TrimRight(names[0], ".") - r.cache.Set(request.address, name) + r.cache.Set(request, name) } - close(request.done) } } // Stop the async reverse resolver -func (r *reverseResolver) Stop() { +func (r *ReverseResolver) Stop() { close(r.addresses) } diff --git a/probe/endpoint/resolver_internal_test.go b/probe/endpoint/resolver_internal_test.go deleted file mode 100644 index e9a2d264e..000000000 --- a/probe/endpoint/resolver_internal_test.go +++ /dev/null @@ -1,41 +0,0 @@ -package endpoint - -import ( - "errors" - "testing" -) - -func TestReverseResolver(t *testing.T) { - tests := map[string]string{ - "8.8.8.8": "google-public-dns-a.google.com", - "8.8.4.4": "google-public-dns-b.google.com", - } - - revRes := newReverseResolver() - - // use a mocked resolver function - revRes.resolver = func(addr string) (names []string, err error) { - if name, ok := tests[addr]; ok { - return []string{name}, nil - } - return []string{}, errors.New("invalid IP") - } - - // first time: no names are returned for our reverse resolutions - for ip := range tests { - if have, err := revRes.Get(ip, true); have != "" || err == nil { - t.Errorf("we didn't get an error, or the cache was not empty, when trying to resolve '%q'", ip) - } - } - - // so, if we check again these IPs, we should have the names now - for ip, want := range tests { - have, err := revRes.Get(ip, true) - if err != nil { - t.Errorf("%s: %v", ip, err) - } - if want != have { - t.Errorf("%s: want %q, have %q", ip, want, have) - } - } -} diff --git a/probe/endpoint/resolver_test.go b/probe/endpoint/resolver_test.go new file mode 100644 index 000000000..31c061e7f --- /dev/null +++ b/probe/endpoint/resolver_test.go @@ -0,0 +1,38 @@ +package endpoint_test + +import ( + "errors" + "testing" + "time" + + . "github.com/weaveworks/scope/probe/endpoint" + "github.com/weaveworks/scope/test" +) + +func TestReverseResolver(t *testing.T) { + tests := map[string]string{ + "1.2.3.4": "test.domain.name", + "4.3.2.1": "im.a.little.tea.pot", + } + + revRes := NewReverseResolver() + defer revRes.Stop() + + // use a mocked resolver function + revRes.Resolver = func(addr string) (names []string, err error) { + if name, ok := tests[addr]; ok { + return []string{name}, nil + } + return []string{}, errors.New("invalid IP") + } + + // Up the rate limit so the test runs faster + revRes.Throttle = time.Tick(time.Millisecond) + + for ip, hostname := range tests { + test.Poll(t, 100*time.Millisecond, hostname, func() interface{} { + result, _ := revRes.Get(ip) + return result + }) + } +}