mirror of
https://github.com/netbirdio/netbird-kubeapi-proxy.git
synced 2026-09-23 09:34:58 -07:00
Cache peers and ensure updates are single flight
This commit is contained in:
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user