Improve test coverage

This commit is contained in:
Philip Laine
2026-06-01 15:22:46 +02:00
parent 5d5bbc32ad
commit 0caff618b7
6 changed files with 167 additions and 52 deletions
+1
View File
@@ -19,6 +19,7 @@ jobs:
uses: actions/setup-go@4a3601121dd01d1626a1e23e37211e3254c1c06c #v6.4.0 uses: actions/setup-go@4a3601121dd01d1626a1e23e37211e3254c1c06c #v6.4.0
with: with:
go-version-file: go.mod go-version-file: go.mod
cache: false
- name: Set up Docker - name: Set up Docker
uses: docker/setup-docker-action@b2189fbf2a6592b51fee7cdd93ee2bfaeba733db #v5.1.0 uses: docker/setup-docker-action@b2189fbf2a6592b51fee7cdd93ee2bfaeba733db #v5.1.0
with: with:
+1
View File
@@ -5,6 +5,7 @@ go 1.25.5
toolchain go1.26.3 toolchain go1.26.3
require ( require (
github.com/go-openapi/testify/v2 v2.5.1
github.com/netbirdio/netbird v0.71.2 github.com/netbirdio/netbird v0.71.2
golang.org/x/sync v0.20.0 golang.org/x/sync v0.20.0
) )
+2
View File
@@ -149,6 +149,8 @@ github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre
github.com/go-ole/go-ole v1.2.6/go.mod h1:pprOEPIfldk/42T2oK7lQ4v4JSDwmV0As9GaiUsvbm0= github.com/go-ole/go-ole v1.2.6/go.mod h1:pprOEPIfldk/42T2oK7lQ4v4JSDwmV0As9GaiUsvbm0=
github.com/go-ole/go-ole v1.3.0 h1:Dt6ye7+vXGIKZ7Xtk4s6/xVdGDQynvom7xCFEdWr6uE= github.com/go-ole/go-ole v1.3.0 h1:Dt6ye7+vXGIKZ7Xtk4s6/xVdGDQynvom7xCFEdWr6uE=
github.com/go-ole/go-ole v1.3.0/go.mod h1:5LS6F96DhAwUc7C+1HLexzMXY1xGRSryjyPPKW6zv78= github.com/go-ole/go-ole v1.3.0/go.mod h1:5LS6F96DhAwUc7C+1HLexzMXY1xGRSryjyPPKW6zv78=
github.com/go-openapi/testify/v2 v2.5.1 h1:TMdhCaw8fUNraVSf3Omoob1dO/AzBfhtFAPW0an6sBo=
github.com/go-openapi/testify/v2 v2.5.1/go.mod h1:SgsVHtfooshd0tublTtJ50FPKhujf47YRqauXXOUxfw=
github.com/go-quicktest/qt v1.101.0 h1:O1K29Txy5P2OK0dGo59b7b0LR6wKfIhttaAhHUyn7eI= github.com/go-quicktest/qt v1.101.0 h1:O1K29Txy5P2OK0dGo59b7b0LR6wKfIhttaAhHUyn7eI=
github.com/go-quicktest/qt v1.101.0/go.mod h1:14Bz/f7NwaXPtdYEgzsx46kqSxVwTbzVZsDC26tQJow= github.com/go-quicktest/qt v1.101.0/go.mod h1:14Bz/f7NwaXPtdYEgzsx46kqSxVwTbzVZsDC26tQJow=
github.com/go-sql-driver/mysql v1.9.3 h1:U/N249h2WzJ3Ukj8SowVFjdtZKfu9vlLZxjPXV1aweo= github.com/go-sql-driver/mysql v1.9.3 h1:U/N249h2WzJ3Ukj8SowVFjdtZKfu9vlLZxjPXV1aweo=
+74 -48
View File
@@ -10,6 +10,7 @@ import (
"crypto/tls" "crypto/tls"
"crypto/x509" "crypto/x509"
"crypto/x509/pkix" "crypto/x509/pkix"
"errors"
"fmt" "fmt"
"math/big" "math/big"
"net" "net"
@@ -21,26 +22,22 @@ import (
"github.com/netbirdio/netbird/client/embed" "github.com/netbirdio/netbird/client/embed"
netbird "github.com/netbirdio/netbird/shared/management/client/rest" netbird "github.com/netbirdio/netbird/shared/management/client/rest"
"github.com/netbirdio/netbird/shared/management/http/api"
) )
func Server(embedClient *embed.Client, netbirdClient *netbird.Client, kubeAPIServerURL *url.URL) (*http.Server, error) { type PeerLister interface {
saToken, err := os.ReadFile("/var/run/secrets/kubernetes.io/serviceaccount/token") List(ctx context.Context, opts ...netbird.PeersListOption) ([]api.Peer, error)
if err != nil { }
return nil, err
}
bearerToken := string(saToken)
certPool, err := x509.SystemCertPool() func Server(embedClient *embed.Client, peerLister PeerLister, kubeAPIServerURL *url.URL) (*http.Server, error) {
bearerToken, err := getBearerToken()
if err != nil { if err != nil {
return nil, err return nil, err
} }
k8sCA, err := os.ReadFile("/var/run/secrets/kubernetes.io/serviceaccount/ca.crt") certPool, err := getCertPool()
if err != nil { if err != nil {
return nil, err return nil, err
} }
if ok := certPool.AppendCertsFromPEM(k8sCA); !ok {
return nil, fmt.Errorf("failed to append Kubernetes CA certificate")
}
transport := http.DefaultTransport.(*http.Transport).Clone() transport := http.DefaultTransport.(*http.Transport).Clone()
transport.TLSClientConfig = &tls.Config{ transport.TLSClientConfig = &tls.Config{
@@ -48,44 +45,8 @@ func Server(embedClient *embed.Client, netbirdClient *netbird.Client, kubeAPISer
} }
proxy := &httputil.ReverseProxy{ proxy := &httputil.ReverseProxy{
Transport: transport, Transport: transport,
Rewrite: func(pr *httputil.ProxyRequest) { Rewrite: rewriteHandler(peerLister, kubeAPIServerURL, bearerToken),
allowedHeaders := map[string]any{
"Accept": nil,
"Accept-Encoding": nil,
"Content-Length": nil,
"Content-Type": nil,
"User-Agent": nil,
}
for k := range pr.Out.Header {
if _, ok := allowedHeaders[k]; !ok {
pr.Out.Header.Del(k)
}
}
remoteIP, _, err := net.SplitHostPort(pr.In.RemoteAddr)
if err != nil {
return
}
listCtx, listCancel := context.WithTimeout(pr.In.Context(), 10*time.Second)
defer listCancel()
peers, err := netbirdClient.Peers.List(listCtx, netbird.PeerIPFilter(remoteIP))
if err != nil {
return
}
if len(peers) != 1 {
return
}
peer := peers[0]
pr.Out.Header.Set("Impersonate-User", peer.UserId)
for _, group := range peer.Groups {
pr.Out.Header.Add("Impersonate-Group", group.Name)
}
pr.Out.Header.Set("Authorization", "Bearer "+bearerToken)
pr.SetURL(kubeAPIServerURL)
},
} }
stat, err := embedClient.Status() stat, err := embedClient.Status()
if err != nil { if err != nil {
return nil, err return nil, err
@@ -106,6 +67,71 @@ func Server(embedClient *embed.Client, netbirdClient *netbird.Client, kubeAPISer
return &srv, nil return &srv, nil
} }
func rewriteHandler(peerLister PeerLister, kubeAPIServerURL *url.URL, bearerToken string) func(*httputil.ProxyRequest) {
return func(pr *httputil.ProxyRequest) {
allowedHeaders := map[string]any{
"Accept": nil,
"Accept-Encoding": nil,
"Content-Length": nil,
"Content-Type": nil,
"User-Agent": nil,
}
for k := range pr.Out.Header {
if _, ok := allowedHeaders[k]; !ok {
pr.Out.Header.Del(k)
}
}
remoteIP, _, err := net.SplitHostPort(pr.In.RemoteAddr)
if err != nil {
return
}
listCtx, listCancel := context.WithTimeout(pr.In.Context(), 10*time.Second)
defer listCancel()
peers, err := peerLister.List(listCtx, netbird.PeerIPFilter(remoteIP))
if err != nil {
return
}
if len(peers) != 1 {
return
}
peer := peers[0]
pr.Out.Header.Set("Impersonate-User", peer.UserId)
for _, group := range peer.Groups {
pr.Out.Header.Add("Impersonate-Group", group.Name)
}
pr.Out.Header.Set("Authorization", "Bearer "+bearerToken)
pr.SetURL(kubeAPIServerURL)
}
}
func getBearerToken() (string, error) {
b, err := os.ReadFile("/var/run/secrets/kubernetes.io/serviceaccount/token")
if err != nil {
return "", err
}
if len(b) == 0 {
return "", errors.New("token cannot be empty")
}
return string(b), nil
}
func getCertPool() (*x509.CertPool, error) {
certPool, err := x509.SystemCertPool()
if err != nil {
return nil, err
}
k8sCA, err := os.ReadFile("/var/run/secrets/kubernetes.io/serviceaccount/ca.crt")
if err != nil {
return nil, err
}
if ok := certPool.AppendCertsFromPEM(k8sCA); !ok {
return nil, fmt.Errorf("failed to append Kubernetes CA certificate")
}
return certPool, nil
}
func generateSelfSignedCert(fqdn string) (tls.Certificate, error) { func generateSelfSignedCert(fqdn string) (tls.Certificate, error) {
priv, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) priv, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil { if err != nil {
+86
View File
@@ -0,0 +1,86 @@
// SPDX-License-Identifier: BSD-3-Clause
package proxy
import (
"context"
"crypto/x509"
"fmt"
"net/http"
"net/http/httptest"
"net/http/httputil"
"net/url"
"testing"
"github.com/go-openapi/testify/v2/require"
netbird "github.com/netbirdio/netbird/shared/management/client/rest"
"github.com/netbirdio/netbird/shared/management/http/api"
)
type mockPeerLister struct {
peers map[string]api.Peer
}
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, fmt.Errorf("peer with ip %s could not be found", ip)
}
return []api.Peer{peer}, nil
}
return nil, nil
}
func TestRewriteHandler(t *testing.T) {
t.Parallel()
peerLister := &mockPeerLister{
peers: map[string]api.Peer{
"192.0.2.1": {
UserId: "foo",
Groups: []api.GroupMinimum{
{
Name: "bar",
},
},
},
},
}
serverURL, err := url.Parse("https://internal.name")
require.NoError(t, err)
bearerToken := "foobar"
handler := rewriteHandler(peerLister, serverURL, bearerToken)
req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/version", nil)
pr := &httputil.ProxyRequest{
In: req,
Out: req.Clone(t.Context()),
}
handler(pr)
require.EqualT(t, "example.com", pr.In.Host)
require.EqualT(t, "https://internal.name/version", pr.Out.URL.String())
}
func TestGenerateSelfSignedCert(t *testing.T) {
t.Parallel()
fqdn := "example.com"
tlsCert, err := generateSelfSignedCert(fqdn)
require.NoError(t, err)
x509Cert, err := x509.ParseCertificate(tlsCert.Certificate[0])
require.NoError(t, err)
require.Len(t, x509Cert.DNSNames, 1)
require.EqualT(t, fqdn, x509Cert.DNSNames[0])
}
+3 -4
View File
@@ -10,7 +10,6 @@ import (
"log" "log"
"net/http" "net/http"
"net/url" "net/url"
"os"
"os/signal" "os/signal"
"syscall" "syscall"
@@ -32,8 +31,8 @@ func main() {
clusterName string clusterName string
) )
flag.StringVar(&mgmtURL, "management-url", "https://api.netbird.io", "NetBird management URL") flag.StringVar(&mgmtURL, "management-url", "https://api.netbird.io", "NetBird management URL")
flag.StringVar(&apiKey, "api-key", os.Getenv("NB_API_KEY"), "NetBird API key") flag.StringVar(&apiKey, "api-key", "", "NetBird API key")
flag.StringVar(&setupKey, "setup-key", os.Getenv("NB_SETUP_KEY"), "NetBird setup key") flag.StringVar(&setupKey, "setup-key", "", "NetBird setup key")
flag.StringVar(&kubeAPIServer, "kubernetes-api-server", "https://kubernetes.default.svc.cluster.local", "Target Kubernetes API server URL") flag.StringVar(&kubeAPIServer, "kubernetes-api-server", "https://kubernetes.default.svc.cluster.local", "Target Kubernetes API server URL")
flag.StringVar(&instanceName, "instance-name", "", "Name of the instance") flag.StringVar(&instanceName, "instance-name", "", "Name of the instance")
flag.StringVar(&clusterName, "cluster-name", "", "Name of the cluster") flag.StringVar(&clusterName, "cluster-name", "", "Name of the cluster")
@@ -83,7 +82,7 @@ func run(ctx context.Context, kubeAPIServer, mgmtURL, apiKey, setupKey, instance
return embedClient.Stop(context.Background()) return embedClient.Stop(context.Background())
}) })
proxySrv, err := proxy.Server(embedClient, netbirdClient, kubeAPIServerURL) proxySrv, err := proxy.Server(embedClient, netbirdClient.Peers, kubeAPIServerURL)
if err != nil { if err != nil {
return err return err
} }