Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
41 changes: 38 additions & 3 deletions cmd/profilecli/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,9 +3,11 @@ package main
import (
"fmt"
"net/http"
"os"
"strings"

"connectrpc.com/connect"
dskittls "github.com/grafana/dskit/crypto/tls"
"github.com/prometheus/common/version"
"gopkg.in/alecthomas/kingpin.v2"

Expand Down Expand Up @@ -61,6 +63,7 @@ type phlareClient struct {
Username string
Password string
}
TLS dskittls.ClientConfig
defaultTransport http.RoundTripper
client *http.Client
protocol string
Expand Down Expand Up @@ -92,14 +95,39 @@ func (a *authRoundTripper) RoundTrip(req *http.Request) (*http.Response, error)
return a.next.RoundTrip(req)
}

func (c *phlareClient) tlsConfigured() bool {
return c.TLS.CertPath != "" || c.TLS.KeyPath != "" ||
c.TLS.CAPath != "" || c.TLS.ServerName != "" ||
c.TLS.InsecureSkipVerify
}

func (c *phlareClient) buildTransport() (http.RoundTripper, error) {
transport := c.defaultTransport
if transport == nil {
transport = http.DefaultTransport
}
if !c.tlsConfigured() {
return transport, nil
}
tlsCfg, err := c.TLS.GetTLSConfig()
if err != nil {
return nil, fmt.Errorf("configuring TLS: %w", err)
}
t := transport.(*http.Transport).Clone()
t.TLSClientConfig = tlsCfg
return t, nil
}

func (c *phlareClient) httpClient() *http.Client {
if c.client == nil {
if c.defaultTransport == nil {
c.defaultTransport = http.DefaultTransport
transport, err := c.buildTransport()
if err != nil {
fmt.Fprintf(os.Stderr, "error: %v\n", err)
os.Exit(1)
}
c.client = &http.Client{Transport: &authRoundTripper{
client: c,
next: c.defaultTransport,
next: transport,
}}
}
return c.client
Expand Down Expand Up @@ -133,5 +161,12 @@ func addPhlareClient(cmd commander) *phlareClient {
cmd.Flag("password", "The password to be used for basic auth.").Default("").Envar(envPrefix + "PASSWORD").StringVar(&client.BasicAuth.Password)
cmd.Flag("protocol", "The protocol to be used for communicating with the server.").Default(protocolTypeConnect).EnumVar(&client.protocol,
protocolTypeConnect, protocolTypeGRPC, protocolTypeGRPCWeb)

cmd.Flag("tls-cert-path", "Path to the client TLS certificate for mTLS authentication.").Default("").Envar(envPrefix + "TLS_CERT_PATH").StringVar(&client.TLS.CertPath)
cmd.Flag("tls-key-path", "Path to the client TLS private key for mTLS authentication.").Default("").Envar(envPrefix + "TLS_KEY_PATH").StringVar(&client.TLS.KeyPath)
cmd.Flag("tls-ca-path", "Path to the CA certificate to verify the server certificate against.").Default("").Envar(envPrefix + "TLS_CA_PATH").StringVar(&client.TLS.CAPath)
cmd.Flag("tls-server-name", "Override the expected TLS server name.").Default("").Envar(envPrefix + "TLS_SERVER_NAME").StringVar(&client.TLS.ServerName)
cmd.Flag("tls-insecure-skip-verify", "Skip TLS server certificate verification (insecure).").Default("false").Envar(envPrefix + "TLS_INSECURE_SKIP_VERIFY").BoolVar(&client.TLS.InsecureSkipVerify)

return client
}
227 changes: 227 additions & 0 deletions cmd/profilecli/client_test.go
Original file line number Diff line number Diff line change
@@ -1,9 +1,22 @@
package main

import (
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"encoding/pem"
"math/big"
"net/http"
"os"
"path/filepath"
"testing"
"time"

dskittls "github.com/grafana/dskit/crypto/tls"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

Expand Down Expand Up @@ -77,3 +90,217 @@ func Test_AcceptHeader(t *testing.T) {
})
}
}

func TestTLSConfigured(t *testing.T) {
tests := []struct {
name string
client phlareClient
want bool
}{
{
name: "no flags set",
client: phlareClient{},
want: false,
},
{
name: "cert path set",
client: phlareClient{TLS: dskittls.ClientConfig{CertPath: "/some/cert"}},
want: true,
},
{
name: "key path set",
client: phlareClient{TLS: dskittls.ClientConfig{KeyPath: "/some/key"}},
want: true,
},
{
name: "CA path set",
client: phlareClient{TLS: dskittls.ClientConfig{CAPath: "/some/ca"}},
want: true,
},
{
name: "server name set",
client: phlareClient{TLS: dskittls.ClientConfig{ServerName: "example.com"}},
want: true,
},
{
name: "insecure skip verify set",
client: phlareClient{TLS: dskittls.ClientConfig{InsecureSkipVerify: true}},
want: true,
},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
assert.Equal(t, tt.want, tt.client.tlsConfigured())
})
}
}

// generateTestCert creates a self-signed CA cert/key pair and writes them to
// temporary files under t.TempDir(). Returns paths to cert, key, and CA files.
func generateTestCert(t *testing.T) (certPath, keyPath, caPath string) {
t.Helper()
dir := t.TempDir()

// Generate CA key
caKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
require.NoError(t, err)

caTemplate := &x509.Certificate{
SerialNumber: big.NewInt(1),
Subject: pkix.Name{Organization: []string{"Test CA"}},
NotBefore: time.Now(),
NotAfter: time.Now().Add(time.Hour),
KeyUsage: x509.KeyUsageCertSign | x509.KeyUsageCRLSign,
BasicConstraintsValid: true,
IsCA: true,
}
caCertDER, err := x509.CreateCertificate(rand.Reader, caTemplate, caTemplate, &caKey.PublicKey, caKey)
require.NoError(t, err)

// Generate client cert signed by CA
clientKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
require.NoError(t, err)

clientTemplate := &x509.Certificate{
SerialNumber: big.NewInt(2),
Subject: pkix.Name{Organization: []string{"Test Client"}},
NotBefore: time.Now(),
NotAfter: time.Now().Add(time.Hour),
KeyUsage: x509.KeyUsageDigitalSignature,
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth},
}
clientCertDER, err := x509.CreateCertificate(rand.Reader, clientTemplate, caTemplate, &clientKey.PublicKey, caKey)
require.NoError(t, err)

// Write CA cert
caPath = filepath.Join(dir, "ca.pem")
require.NoError(t, os.WriteFile(caPath, pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: caCertDER}), 0o600))

// Write client cert
certPath = filepath.Join(dir, "client.pem")
require.NoError(t, os.WriteFile(certPath, pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: clientCertDER}), 0o600))

// Write client key
keyBytes, err := x509.MarshalECPrivateKey(clientKey)
require.NoError(t, err)
keyPath = filepath.Join(dir, "client-key.pem")
require.NoError(t, os.WriteFile(keyPath, pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyBytes}), 0o600))

return certPath, keyPath, caPath
}

func TestBuildTransport(t *testing.T) {
certPath, keyPath, caPath := generateTestCert(t)

tests := []struct {
name string
client phlareClient
assert func(t *testing.T, rt http.RoundTripper)
wantErr string
}{
{
name: "no TLS returns default transport",
client: phlareClient{},
assert: func(t *testing.T, rt http.RoundTripper) {
assert.Equal(t, http.DefaultTransport, rt)
},
},
{
name: "CA path populates RootCAs",
client: phlareClient{
TLS: dskittls.ClientConfig{CAPath: caPath},
},
assert: func(t *testing.T, rt http.RoundTripper) {
tr := rt.(*http.Transport)
require.NotNil(t, tr.TLSClientConfig)
assert.NotNil(t, tr.TLSClientConfig.RootCAs)
},
},
{
name: "client cert sets GetClientCertificate",
client: phlareClient{
TLS: dskittls.ClientConfig{CertPath: certPath, KeyPath: keyPath},
},
assert: func(t *testing.T, rt http.RoundTripper) {
tr := rt.(*http.Transport)
require.NotNil(t, tr.TLSClientConfig)
assert.NotNil(t, tr.TLSClientConfig.GetClientCertificate)
},
},
{
name: "insecure skip verify",
client: phlareClient{
TLS: dskittls.ClientConfig{InsecureSkipVerify: true},
},
assert: func(t *testing.T, rt http.RoundTripper) {
tr := rt.(*http.Transport)
require.NotNil(t, tr.TLSClientConfig)
assert.True(t, tr.TLSClientConfig.InsecureSkipVerify)
},
},
{
name: "server name override",
client: phlareClient{
TLS: dskittls.ClientConfig{ServerName: "custom.example.com"},
},
assert: func(t *testing.T, rt http.RoundTripper) {
tr := rt.(*http.Transport)
require.NotNil(t, tr.TLSClientConfig)
assert.Equal(t, "custom.example.com", tr.TLSClientConfig.ServerName)
},
},
{
name: "cert without key returns error",
client: phlareClient{
TLS: dskittls.ClientConfig{CertPath: certPath},
},
wantErr: "configuring TLS: certificate given but no key configured",
},
{
name: "key without cert returns error",
client: phlareClient{
TLS: dskittls.ClientConfig{KeyPath: keyPath},
},
wantErr: "configuring TLS: key given but no certificate configured",
},
{
name: "all options combined",
client: phlareClient{
TLS: dskittls.ClientConfig{
CertPath: certPath,
KeyPath: keyPath,
CAPath: caPath,
ServerName: "pyroscope.local",
InsecureSkipVerify: true,
},
},
assert: func(t *testing.T, rt http.RoundTripper) {
tr := rt.(*http.Transport)
require.NotNil(t, tr.TLSClientConfig)
cfg := tr.TLSClientConfig
assert.NotNil(t, cfg.RootCAs)
assert.NotNil(t, cfg.GetClientCertificate)
assert.True(t, cfg.InsecureSkipVerify)
assert.Equal(t, "pyroscope.local", cfg.ServerName)

// Verify GetClientCertificate returns a valid cert
cert, err := cfg.GetClientCertificate(&tls.CertificateRequestInfo{})
require.NoError(t, err)
require.NotNil(t, cert)
},
},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
rt, err := tt.client.buildTransport()
if tt.wantErr != "" {
require.EqualError(t, err, tt.wantErr)
return
}
require.NoError(t, err)
tt.assert(t, rt)
})
}
}