1
0
Fork 0

ip/cache.go: skip storing a NetInterfaceIPResolver in struct, Define a CacheDefaultCallback type and have it passed to the new GetWithDefault function instead.

This commit is contained in:
Henrik Hautakoski 2023-12-07 20:40:49 +01:00
parent 0f21902ffd
commit a6c98a3209
3 changed files with 74 additions and 16 deletions

View file

@ -1,18 +1,19 @@
package ip
import (
"errors"
"net"
)
type CacheDefaultCallback func(name string) (net.IP, error)
type Cache struct {
resolver NetInterfaceIPResolver
items map[string]net.IP
items map[string]net.IP
}
func NewCache(resolver NetInterfaceIPResolver) *Cache {
func NewCache() *Cache {
return &Cache{
resolver: resolver,
items: make(map[string]net.IP),
items: make(map[string]net.IP),
}
}
@ -21,8 +22,16 @@ func (c Cache) Get(name string) (net.IP, error) {
if cached, ok := c.items[name]; ok {
return cached, nil
}
return nil, errors.New("key did not exist")
}
ip, err := c.resolver(name)
func (c Cache) GetWithDefault(name string, callback CacheDefaultCallback) (net.IP, error) {
// Return cached entry.
if cached, ok := c.items[name]; ok {
return cached, nil
}
ip, err := callback(name)
if err == nil {
c.Set(name, ip)
}

View file

@ -9,13 +9,20 @@ import (
"github.com/stretchr/testify/assert"
)
func mockResolver(t *testing.T, expected_name string, ip net.IP, err error) NetInterfaceIPResolver {
func defaultCallback(t *testing.T, expected_name string, ip net.IP, err error) CacheDefaultCallback {
return func(name string) (net.IP, error) {
assert.Equal(t, expected_name, name)
return ip, err
}
}
func dontCallDefaultCallback(t *testing.T) CacheDefaultCallback {
return func(name string) (net.IP, error) {
t.Error("Should not have been called")
return nil, nil
}
}
func TestCache_Get(t *testing.T) {
tests := []struct {
name string
@ -24,9 +31,8 @@ func TestCache_Get(t *testing.T) {
want net.IP
wantErr bool
}{
{"FromCache", &Cache{resolver: nil, items: map[string]net.IP{"eth0": net.IPv4(10, 4, 0, 1)}}, "eth0", net.IPv4(10, 4, 0, 1), false},
{"FromResolver", NewCache(mockResolver(t, "eth1", net.IPv4(192, 172, 44, 25), nil)), "eth1", net.IPv4(192, 172, 44, 25), false},
{"NoInterface", NewCache(mockResolver(t, "eth2", nil, errors.New("Invalid interface"))), "eth2", nil, true},
{"Exists in cache", &Cache{items: map[string]net.IP{"eth0": net.IPv4(10, 4, 0, 1)}}, "eth0", net.IPv4(10, 4, 0, 1), false},
{"Did not exist in cache", &Cache{items: map[string]net.IP{}}, "eth0", nil, true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
@ -41,3 +47,30 @@ func TestCache_Get(t *testing.T) {
})
}
}
func TestCache_GetWithDefault(t *testing.T) {
tests := []struct {
name string
c *Cache
def CacheDefaultCallback
iface string
want net.IP
wantErr bool
}{
{"Exists in cache", &Cache{items: map[string]net.IP{"eth0": net.IPv4(10, 4, 0, 1)}}, dontCallDefaultCallback(t), "eth0", net.IPv4(10, 4, 0, 1), false},
{"Did not exists in cache", NewCache(), defaultCallback(t, "eth1", net.IPv4(192, 172, 44, 25), nil), "eth1", net.IPv4(192, 172, 44, 25), false},
{"Callback returns error", NewCache(), defaultCallback(t, "eth1", nil, errors.New("some error")), "eth1", nil, true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := tt.c.GetWithDefault(tt.iface, tt.def)
if (err != nil) != tt.wantErr {
t.Errorf("Cache.Get() error = %v, wantErr %v", err, tt.wantErr)
return
}
if !reflect.DeepEqual(got, tt.want) {
t.Errorf("Cache.Get() = %v, want %v", got, tt.want)
}
})
}
}