From f68a36a867df2e0442d7ee98c05707436b04b88c Mon Sep 17 00:00:00 2001 From: Philip Laine Date: Thu, 11 Jun 2026 13:26:12 +0200 Subject: [PATCH] Cache peers and ensure updates are single flight --- go.mod | 1 + go.sum | 2 + internal/proxy/peer.go | 98 ++++++++++++++++++++++++++++++++++++ internal/proxy/peer_test.go | 37 ++++++++++++++ internal/proxy/proxy.go | 28 ++++------- internal/proxy/proxy_test.go | 53 ++++++++----------- main.go | 3 +- 7 files changed, 171 insertions(+), 51 deletions(-) create mode 100644 internal/proxy/peer.go create mode 100644 internal/proxy/peer_test.go diff --git a/go.mod b/go.mod index 6f44247..deaecaa 100644 --- a/go.mod +++ b/go.mod @@ -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 ( diff --git a/go.sum b/go.sum index 287fc7d..0ed1a0d 100644 --- a/go.sum +++ b/go.sum @@ -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= diff --git a/internal/proxy/peer.go b/internal/proxy/peer.go new file mode 100644 index 0000000..07a29ed --- /dev/null +++ b/internal/proxy/peer.go @@ -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) + } +} diff --git a/internal/proxy/peer_test.go b/internal/proxy/peer_test.go new file mode 100644 index 0000000..af5bf38 --- /dev/null +++ b/internal/proxy/peer_test.go @@ -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) +} diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go index 82bb335..9b2d8e6 100644 --- a/internal/proxy/proxy.go +++ b/internal/proxy/proxy.go @@ -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)) } } diff --git a/internal/proxy/proxy_test.go b/internal/proxy/proxy_test.go index 3ddc749..fbf64c0 100644 --- a/internal/proxy/proxy_test.go +++ b/internal/proxy/proxy_test.go @@ -19,11 +19,23 @@ 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) { + peers := map[string]api.Peer{ + "192.0.2.1": { + UserId: "foo", + Groups: []api.GroupMinimum{ + { + Name: "group1", + }, + { + Name: "group2", + }, + }, + }, + } + ip := "" for _, o := range opts { k, v := o() @@ -33,7 +45,7 @@ func (p *mockPeerLister) List(ctx context.Context, opts ...netbird.PeersListOpti } } if ip != "" { - peer, ok := p.peers[ip] + peer, ok := peers[ip] if !ok { return nil, nil } @@ -45,21 +57,7 @@ func (p *mockPeerLister) List(ctx context.Context, opts ...netbird.PeersListOpti func TestProxyHandler(t *testing.T) { t.Parallel() - peerLister := &mockPeerLister{ - peers: map[string]api.Peer{ - "192.0.2.1": { - UserId: "foo", - Groups: []api.GroupMinimum{ - { - Name: "group1", - }, - { - Name: "group2", - }, - }, - }, - }, - } + peerStore := NewPeerStore(&mockPeerLister{}) bearerToken := "foobar" srv := httptest.NewTLSServer(http.HandlerFunc(func(rw http.ResponseWriter, req *http.Request) { @@ -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) diff --git a/main.go b/main.go index 3c3abac..1911e73 100644 --- a/main.go +++ b/main.go @@ -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 }