diff --git a/.github/workflows/release.yaml b/.github/workflows/release.yaml index e5e9962..140ae5d 100644 --- a/.github/workflows/release.yaml +++ b/.github/workflows/release.yaml @@ -19,6 +19,7 @@ jobs: uses: actions/setup-go@4a3601121dd01d1626a1e23e37211e3254c1c06c #v6.4.0 with: go-version-file: go.mod + cache: false - name: Set up Docker uses: docker/setup-docker-action@b2189fbf2a6592b51fee7cdd93ee2bfaeba733db #v5.1.0 with: diff --git a/go.mod b/go.mod index 6978347..69dec81 100644 --- a/go.mod +++ b/go.mod @@ -5,6 +5,7 @@ go 1.25.5 toolchain go1.26.3 require ( + github.com/go-openapi/testify/v2 v2.5.1 github.com/netbirdio/netbird v0.71.2 golang.org/x/sync v0.20.0 ) diff --git a/go.sum b/go.sum index 15b6623..ceb0202 100644 --- a/go.sum +++ b/go.sum @@ -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.3.0 h1:Dt6ye7+vXGIKZ7Xtk4s6/xVdGDQynvom7xCFEdWr6uE= 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/go.mod h1:14Bz/f7NwaXPtdYEgzsx46kqSxVwTbzVZsDC26tQJow= github.com/go-sql-driver/mysql v1.9.3 h1:U/N249h2WzJ3Ukj8SowVFjdtZKfu9vlLZxjPXV1aweo= diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go index 66118be..fcb5da3 100644 --- a/internal/proxy/proxy.go +++ b/internal/proxy/proxy.go @@ -10,6 +10,7 @@ import ( "crypto/tls" "crypto/x509" "crypto/x509/pkix" + "errors" "fmt" "math/big" "net" @@ -21,26 +22,22 @@ import ( "github.com/netbirdio/netbird/client/embed" 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) { - saToken, err := os.ReadFile("/var/run/secrets/kubernetes.io/serviceaccount/token") - if err != nil { - return nil, err - } - bearerToken := string(saToken) +type PeerLister interface { + List(ctx context.Context, opts ...netbird.PeersListOption) ([]api.Peer, error) +} - certPool, err := x509.SystemCertPool() +func Server(embedClient *embed.Client, peerLister PeerLister, kubeAPIServerURL *url.URL) (*http.Server, error) { + bearerToken, err := getBearerToken() if err != nil { return nil, err } - k8sCA, err := os.ReadFile("/var/run/secrets/kubernetes.io/serviceaccount/ca.crt") + certPool, err := getCertPool() if err != nil { 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.TLSClientConfig = &tls.Config{ @@ -48,44 +45,8 @@ func Server(embedClient *embed.Client, netbirdClient *netbird.Client, kubeAPISer } proxy := &httputil.ReverseProxy{ Transport: transport, - Rewrite: 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 := 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) - }, + Rewrite: rewriteHandler(peerLister, kubeAPIServerURL, bearerToken), } - stat, err := embedClient.Status() if err != nil { return nil, err @@ -106,6 +67,71 @@ func Server(embedClient *embed.Client, netbirdClient *netbird.Client, kubeAPISer 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) { priv, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) if err != nil { diff --git a/internal/proxy/proxy_test.go b/internal/proxy/proxy_test.go new file mode 100644 index 0000000..dc12eec --- /dev/null +++ b/internal/proxy/proxy_test.go @@ -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]) +} diff --git a/main.go b/main.go index ab5db6e..3aedf84 100644 --- a/main.go +++ b/main.go @@ -10,7 +10,6 @@ import ( "log" "net/http" "net/url" - "os" "os/signal" "syscall" @@ -32,8 +31,8 @@ func main() { clusterName string ) 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(&setupKey, "setup-key", os.Getenv("NB_SETUP_KEY"), "NetBird setup key") + flag.StringVar(&apiKey, "api-key", "", "NetBird API 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(&instanceName, "instance-name", "", "Name of the instance") 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()) }) - proxySrv, err := proxy.Server(embedClient, netbirdClient, kubeAPIServerURL) + proxySrv, err := proxy.Server(embedClient, netbirdClient.Peers, kubeAPIServerURL) if err != nil { return err }