Cache peers and ensure updates are single flight

This commit is contained in:
Philip Laine
2026-06-11 14:56:19 +02:00
parent 22e02bd344
commit f68a36a867
7 changed files with 171 additions and 51 deletions
+1
View File
@@ -8,6 +8,7 @@ require (
github.com/go-openapi/testify/v2 v2.5.1
github.com/netbirdio/netbird v0.72.3
golang.org/x/sync v0.21.0
resenje.org/singleflight v0.4.3
)
require (
+2
View File
@@ -706,5 +706,7 @@ gorm.io/gorm v1.25.12 h1:I0u8i2hWQItBq1WfE0o2+WuL9+8L21K9e2HHSTE/0f8=
gorm.io/gorm v1.25.12/go.mod h1:xh7N7RHfYlNc5EmcI/El95gXusucDrQnHXe0+CgWcLQ=
gvisor.dev/gvisor v0.0.0-20260219192049-0f2374377e89 h1:mGJaeA61P8dEHTqdvAgc70ZIV3QoUoJcXCRyyjO26OA=
gvisor.dev/gvisor v0.0.0-20260219192049-0f2374377e89/go.mod h1:QkHjoMIBaYtpVufgwv3keYAbln78mBoCuShZrPrer1Q=
resenje.org/singleflight v0.4.3 h1:l7foFYg8X/VEHPxWs1K/Pw77807RMVzvXgWGb0J1sdM=
resenje.org/singleflight v0.4.3/go.mod h1:lAgQK7VfjG6/pgredbQfmV0RvG/uVhKo6vSuZ0vCWfk=
rsc.io/qr v0.2.0 h1:6vBLea5/NRMVTz8V66gipeLycZMl/+UlFmk8DvqQ6WY=
rsc.io/qr v0.2.0/go.mod h1:IF+uZjkb9fqyeF/4tlBoynqmQxUoPfWEKh921coOuXs=
+98
View File
@@ -0,0 +1,98 @@
// SPDX-License-Identifier: AGPL-3.0
package proxy
import (
"context"
"errors"
"fmt"
"sync"
"time"
"resenje.org/singleflight"
netbird "github.com/netbirdio/netbird/shared/management/client/rest"
"github.com/netbirdio/netbird/shared/management/http/api"
)
var (
ErrNotFound = errors.New("not found")
)
type PeerLister interface {
List(ctx context.Context, opts ...netbird.PeersListOption) ([]api.Peer, error)
}
type PeerStore struct {
peerLister PeerLister
getGroup *singleflight.Group[string, api.Peer]
cacheMx sync.RWMutex
evictMx sync.Mutex
cache map[string]api.Peer
lastUpdated map[string]time.Time
cacheTTL time.Duration
}
func NewPeerStore(peerLister PeerLister) *PeerStore {
return &PeerStore{
peerLister: peerLister,
getGroup: &singleflight.Group[string, api.Peer]{},
cache: map[string]api.Peer{},
lastUpdated: map[string]time.Time{},
cacheTTL: 15 * time.Second,
}
}
func (p *PeerStore) Get(ctx context.Context, ip string) (api.Peer, error) {
p.evictCache()
peer, _, err := p.getGroup.Do(ctx, ip, func(ctx context.Context) (api.Peer, error) {
p.cacheMx.RLock()
if peer, ok := p.cache[ip]; ok {
if time.Since(p.lastUpdated[ip]) < p.cacheTTL {
p.cacheMx.RUnlock()
return peer, nil
}
}
p.cacheMx.RUnlock()
peers, err := p.peerLister.List(ctx, netbird.PeerIPFilter(ip))
if err != nil {
return api.Peer{}, err
}
if len(peers) == 0 {
return api.Peer{}, ErrNotFound
}
if len(peers) > 1 {
return api.Peer{}, fmt.Errorf("receive more than one peer for ip %s", ip)
}
peer := peers[0]
p.cacheMx.Lock()
p.cache[ip] = peer
p.lastUpdated[ip] = time.Now()
p.cacheMx.Unlock()
return peer, nil
})
if err != nil {
return api.Peer{}, err
}
return peer, nil
}
func (p *PeerStore) evictCache() {
if !p.evictMx.TryLock() {
return
}
defer p.evictMx.Unlock()
p.cacheMx.Lock()
defer p.cacheMx.Unlock()
for k := range p.cache {
if time.Since(p.lastUpdated[k]) < p.cacheTTL {
continue
}
delete(p.cache, k)
delete(p.lastUpdated, k)
}
}
+37
View File
@@ -0,0 +1,37 @@
// SPDX-License-Identifier: AGPL-3.0
package proxy
import (
"fmt"
"testing"
"time"
"github.com/go-openapi/testify/v2/require"
"github.com/netbirdio/netbird/shared/management/http/api"
)
func TestPeerStoreEviction(t *testing.T) {
t.Parallel()
peerStore := NewPeerStore(nil)
peerStore.cache["foo"] = api.Peer{}
peerStore.lastUpdated["foo"] = time.Now()
peerStore.evictCache()
require.Len(t, peerStore.cache, 1)
require.Len(t, peerStore.lastUpdated, 1)
for i := range 100 {
peerStore.cache[fmt.Sprintf("%d", i)] = api.Peer{}
peerStore.lastUpdated[fmt.Sprintf("%d", i)] = time.Now().Add(-time.Minute)
}
require.Len(t, peerStore.cache, 101)
require.Len(t, peerStore.lastUpdated, 101)
_, err := peerStore.Get(t.Context(), "foo")
require.NoError(t, err)
require.Len(t, peerStore.cache, 1)
require.Len(t, peerStore.lastUpdated, 1)
}
+11 -17
View File
@@ -21,7 +21,6 @@ import (
"time"
"github.com/netbirdio/netbird/client/embed"
netbird "github.com/netbirdio/netbird/shared/management/client/rest"
"github.com/netbirdio/netbird/shared/management/http/api"
)
@@ -43,11 +42,7 @@ const (
SecWebsocketExtensionsHeader = "Sec-Websocket-Extensions"
)
type PeerLister interface {
List(ctx context.Context, opts ...netbird.PeersListOption) ([]api.Peer, error)
}
func Server(embedClient *embed.Client, peerLister PeerLister, kubeAPIServerURL *url.URL) (*http.Server, error) {
func Server(embedClient *embed.Client, peerStore *PeerStore, kubeAPIServerURL *url.URL) (*http.Server, error) {
bearerToken, err := getBearerToken()
if err != nil {
return nil, err
@@ -56,7 +51,7 @@ func Server(embedClient *embed.Client, peerLister PeerLister, kubeAPIServerURL *
if err != nil {
return nil, err
}
handler := proxyHandler(peerLister, kubeAPIServerURL, certPool, bearerToken)
handler := proxyHandler(peerStore, kubeAPIServerURL, certPool, bearerToken)
stat, err := embedClient.Status()
if err != nil {
@@ -78,7 +73,7 @@ func Server(embedClient *embed.Client, peerLister PeerLister, kubeAPIServerURL *
return &srv, nil
}
func proxyHandler(peerLister PeerLister, kubeAPIServerURL *url.URL, certPool *x509.CertPool, bearerToken string) http.HandlerFunc {
func proxyHandler(peerStore *PeerStore, kubeAPIServerURL *url.URL, certPool *x509.CertPool, bearerToken string) http.HandlerFunc {
type peerCtxKey struct{}
rewrite := func(pr *httputil.ProxyRequest) {
@@ -146,19 +141,18 @@ func proxyHandler(peerLister PeerLister, kubeAPIServerURL *url.URL, certPool *x5
rw.WriteHeader(http.StatusBadRequest)
return
}
listCtx, listCancel := context.WithTimeout(req.Context(), 10*time.Second)
defer listCancel()
peers, err := peerLister.List(listCtx, netbird.PeerIPFilter(remoteIP))
getCtx, getCancel := context.WithTimeout(req.Context(), 10*time.Second)
defer getCancel()
peer, err := peerStore.Get(getCtx, remoteIP)
if errors.Is(err, ErrNotFound) {
rw.WriteHeader(http.StatusUnauthorized)
return
}
if err != nil {
rw.WriteHeader(http.StatusInternalServerError)
return
}
if len(peers) != 1 {
rw.WriteHeader(http.StatusUnauthorized)
return
}
peerCtx := context.WithValue(req.Context(), peerCtxKey{}, peers[0])
peerCtx := context.WithValue(req.Context(), peerCtxKey{}, peer)
proxy.ServeHTTP(rw, req.WithContext(peerCtx))
}
}
+28 -41
View File
@@ -19,34 +19,10 @@ import (
"github.com/netbirdio/netbird/shared/management/http/api"
)
type mockPeerLister struct {
peers map[string]api.Peer
}
type mockPeerLister struct{}
func (p *mockPeerLister) List(ctx context.Context, opts ...netbird.PeersListOption) ([]api.Peer, error) {
ip := ""
for _, o := range opts {
k, v := o()
if k == "ip" {
ip = v
break
}
}
if ip != "" {
peer, ok := p.peers[ip]
if !ok {
return nil, nil
}
return []api.Peer{peer}, nil
}
return nil, nil
}
func TestProxyHandler(t *testing.T) {
t.Parallel()
peerLister := &mockPeerLister{
peers: map[string]api.Peer{
peers := map[string]api.Peer{
"192.0.2.1": {
UserId: "foo",
Groups: []api.GroupMinimum{
@@ -58,9 +34,31 @@ func TestProxyHandler(t *testing.T) {
},
},
},
},
}
ip := ""
for _, o := range opts {
k, v := o()
if k == "ip" {
ip = v
break
}
}
if ip != "" {
peer, ok := peers[ip]
if !ok {
return nil, nil
}
return []api.Peer{peer}, nil
}
return nil, nil
}
func TestProxyHandler(t *testing.T) {
t.Parallel()
peerStore := NewPeerStore(&mockPeerLister{})
bearerToken := "foobar"
srv := httptest.NewTLSServer(http.HandlerFunc(func(rw http.ResponseWriter, req *http.Request) {
token, _ := strings.CutPrefix(req.Header.Get("Authorization"), "Bearer ")
@@ -124,7 +122,7 @@ func TestProxyHandler(t *testing.T) {
req.Header.Add(k, v)
}
rec := httptest.NewRecorder()
handler := proxyHandler(peerLister, kubeAPIServerURL, certPool, bearerToken)
handler := proxyHandler(peerStore, kubeAPIServerURL, certPool, bearerToken)
handler(rec, req)
b, err := io.ReadAll(rec.Result().Body)
require.NoError(t, err)
@@ -138,18 +136,7 @@ func TestProxyHandler(t *testing.T) {
func TestProxyHandlerPreservesUpgrade(t *testing.T) {
t.Parallel()
peerLister := &mockPeerLister{
peers: map[string]api.Peer{
"192.0.2.1": {
UserId: "foo",
Groups: []api.GroupMinimum{
{
Name: "group1",
},
},
},
},
}
peerStore := NewPeerStore(&mockPeerLister{})
bearerToken := "foobar"
srv := httptest.NewTLSServer(http.HandlerFunc(func(rw http.ResponseWriter, req *http.Request) {
@@ -180,7 +167,7 @@ func TestProxyHandlerPreservesUpgrade(t *testing.T) {
req.Header.Set(ImpersonateUserHeader, "system:admin")
rec := httptest.NewRecorder()
handler := proxyHandler(peerLister, kubeAPIServerURL, certPool, bearerToken)
handler := proxyHandler(peerStore, kubeAPIServerURL, certPool, bearerToken)
handler(rec, req)
b, err := io.ReadAll(rec.Result().Body)
+2 -1
View File
@@ -83,7 +83,8 @@ func run(ctx context.Context, kubeAPIServer, mgmtURL, apiKey, setupKey, instance
return embedClient.Stop(context.Background())
})
proxySrv, err := proxy.Server(embedClient, netbirdClient.Peers, kubeAPIServerURL)
peerStore := proxy.NewPeerStore(netbirdClient.Peers)
proxySrv, err := proxy.Server(embedClient, peerStore, kubeAPIServerURL)
if err != nil {
return err
}