Postgres: Switch the datasource plugin from lib/pq to pgx (#83768)
postgres: switch from lib/pq to pgx
This commit is contained in:
@@ -0,0 +1,135 @@
|
||||
package tls
|
||||
|
||||
import (
|
||||
"crypto/tls"
|
||||
"crypto/x509"
|
||||
"errors"
|
||||
|
||||
"github.com/grafana/grafana/pkg/tsdb/sqleng"
|
||||
)
|
||||
|
||||
// we support 4 postgres tls modes:
|
||||
// disable - no tls
|
||||
// require - use tls
|
||||
// verify-ca - use tls, verify root cert but not the hostname
|
||||
// verify-full - use tls, verify root cert
|
||||
// (for all the options except `disable`, you can optionally use client certificates)
|
||||
|
||||
func getTLSConfigRequire(certs *Certs, serverName string) (*tls.Config, error) {
|
||||
// see https://www.postgresql.org/docs/12/libpq-ssl.html ,
|
||||
// mode=require + provided root-cert should behave as mode=verify-ca
|
||||
if certs.rootCerts != nil {
|
||||
return getTLSConfigVerifyCA(certs, serverName)
|
||||
}
|
||||
|
||||
return &tls.Config{
|
||||
InsecureSkipVerify: true, // we do not verify the root cert
|
||||
Certificates: certs.clientCerts,
|
||||
ServerName: serverName,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// to implement the verify-ca mode, we need to do this:
|
||||
// - for the root certificate
|
||||
// - verify that the certificate we receive from the server is trusted,
|
||||
// meaning it relates to our root certificate
|
||||
// - we DO NOT verify that the hostname of the database matches
|
||||
// the hostname in the certificate
|
||||
//
|
||||
// the problem is, `go“ does not offer such an option.
|
||||
// by default, it will verify both things.
|
||||
//
|
||||
// so what we do is:
|
||||
// - we turn off the default-verification with `InsecureSkipVerify`
|
||||
// - we implement our own verification using `VerifyConnection`
|
||||
//
|
||||
// extra info about this:
|
||||
// - there is a rejected feature-request about this at https://github.com/golang/go/issues/21971
|
||||
// - the recommended workaround is based on VerifyPeerCertificate
|
||||
// - there is even example code at https://github.com/golang/go/commit/29cfb4d3c3a97b6f426d1b899234da905be699aa
|
||||
// - but later the example code was changed to use VerifyConnection instead:
|
||||
// https://github.com/golang/go/commit/7eb5941b95a588a23f18fa4c22fe42ff0119c311
|
||||
//
|
||||
// a verifyConnection example is at https://pkg.go.dev/crypto/tls#example-Config-VerifyConnection .
|
||||
//
|
||||
// this is how the `pgx` library handles verify-ca:
|
||||
//
|
||||
// https://github.com/jackc/pgx/blob/5c63f646f820ca9696fc3515c1caf2a557d562e5/pgconn/config.go#L657-L690
|
||||
// (unfortunately pgx only handles this for certificate-provided-as-path, so we cannot rely on it)
|
||||
func getTLSConfigVerifyCA(certs *Certs, serverName string) (*tls.Config, error) {
|
||||
conf := tls.Config{
|
||||
ServerName: serverName,
|
||||
Certificates: certs.clientCerts,
|
||||
InsecureSkipVerify: true, // we turn off the default-verification, we'll do VerifyConnection instead
|
||||
VerifyConnection: func(state tls.ConnectionState) error {
|
||||
// we add all the certificates to the pool, we skip the first cert.
|
||||
intermediates := x509.NewCertPool()
|
||||
for _, c := range state.PeerCertificates[1:] {
|
||||
intermediates.AddCert(c)
|
||||
}
|
||||
|
||||
opts := x509.VerifyOptions{
|
||||
Roots: certs.rootCerts,
|
||||
Intermediates: intermediates,
|
||||
}
|
||||
|
||||
// we call `Verify()` on the first cert (that we skipped previously)
|
||||
_, err := state.PeerCertificates[0].Verify(opts)
|
||||
return err
|
||||
},
|
||||
RootCAs: certs.rootCerts,
|
||||
}
|
||||
|
||||
return &conf, nil
|
||||
}
|
||||
|
||||
func getTLSConfigVerifyFull(certs *Certs, serverName string) (*tls.Config, error) {
|
||||
conf := tls.Config{
|
||||
Certificates: certs.clientCerts,
|
||||
ServerName: serverName,
|
||||
RootCAs: certs.rootCerts,
|
||||
}
|
||||
|
||||
return &conf, nil
|
||||
}
|
||||
|
||||
func IsTLSEnabled(dsInfo sqleng.DataSourceInfo) bool {
|
||||
mode := dsInfo.JsonData.Mode
|
||||
return mode != "disable"
|
||||
}
|
||||
|
||||
// returns `nil` if tls is disabled
|
||||
func GetTLSConfig(dsInfo sqleng.DataSourceInfo, readFile ReadFileFunc, serverName string) (*tls.Config, error) {
|
||||
mode := dsInfo.JsonData.Mode
|
||||
// we need to special-case the no-tls-mode
|
||||
if mode == "disable" {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// for all the remaining cases we need to load
|
||||
// both the root-cert if exists, and the client-cert if exists.
|
||||
certBytes, err := loadCertificateBytes(dsInfo, readFile)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
certs, err := createCertificates(certBytes)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
switch mode {
|
||||
// `disable` already handled
|
||||
case "":
|
||||
// for backward-compatibility reasons this is the same as `require`
|
||||
return getTLSConfigRequire(certs, serverName)
|
||||
case "require":
|
||||
return getTLSConfigRequire(certs, serverName)
|
||||
case "verify-ca":
|
||||
return getTLSConfigVerifyCA(certs, serverName)
|
||||
case "verify-full":
|
||||
return getTLSConfigVerifyFull(certs, serverName)
|
||||
default:
|
||||
return nil, errors.New("tls: invalid mode " + mode)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
package tls
|
||||
|
||||
import (
|
||||
"crypto/tls"
|
||||
"crypto/x509"
|
||||
"errors"
|
||||
|
||||
"github.com/grafana/grafana/pkg/tsdb/sqleng"
|
||||
)
|
||||
|
||||
// this file deals with locating and loading the certificates,
|
||||
// from json-data or from disk.
|
||||
|
||||
type CertBytes struct {
|
||||
rootCert []byte
|
||||
clientKey []byte
|
||||
clientCert []byte
|
||||
}
|
||||
|
||||
type ReadFileFunc = func(name string) ([]byte, error)
|
||||
|
||||
var errPartialClientCertNoKey = errors.New("tls: client cert provided but client key missing")
|
||||
var errPartialClientCertNoCert = errors.New("tls: client key provided but client cert missing")
|
||||
|
||||
// certificates can be stored either as encrypted-json-data, or as file-path
|
||||
func loadCertificateBytes(dsInfo sqleng.DataSourceInfo, readFile ReadFileFunc) (*CertBytes, error) {
|
||||
if dsInfo.JsonData.ConfigurationMethod == "file-content" {
|
||||
return &CertBytes{
|
||||
rootCert: []byte(dsInfo.DecryptedSecureJSONData["tlsCACert"]),
|
||||
clientKey: []byte(dsInfo.DecryptedSecureJSONData["tlsClientKey"]),
|
||||
clientCert: []byte(dsInfo.DecryptedSecureJSONData["tlsClientCert"]),
|
||||
}, nil
|
||||
} else {
|
||||
c := CertBytes{}
|
||||
|
||||
if dsInfo.JsonData.RootCertFile != "" {
|
||||
rootCert, err := readFile(dsInfo.JsonData.RootCertFile)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
c.rootCert = rootCert
|
||||
}
|
||||
|
||||
if dsInfo.JsonData.CertKeyFile != "" {
|
||||
clientKey, err := readFile(dsInfo.JsonData.CertKeyFile)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
c.clientKey = clientKey
|
||||
}
|
||||
|
||||
if dsInfo.JsonData.CertFile != "" {
|
||||
clientCert, err := readFile(dsInfo.JsonData.CertFile)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
c.clientCert = clientCert
|
||||
}
|
||||
|
||||
return &c, nil
|
||||
}
|
||||
}
|
||||
|
||||
type Certs struct {
|
||||
clientCerts []tls.Certificate
|
||||
rootCerts *x509.CertPool
|
||||
}
|
||||
|
||||
func createCertificates(certBytes *CertBytes) (*Certs, error) {
|
||||
certs := Certs{}
|
||||
|
||||
if len(certBytes.rootCert) > 0 {
|
||||
pool := x509.NewCertPool()
|
||||
ok := pool.AppendCertsFromPEM(certBytes.rootCert)
|
||||
if !ok {
|
||||
return nil, errors.New("tls: failed to add root certificate")
|
||||
}
|
||||
certs.rootCerts = pool
|
||||
}
|
||||
|
||||
hasClientKey := len(certBytes.clientKey) > 0
|
||||
hasClientCert := len(certBytes.clientCert) > 0
|
||||
|
||||
if hasClientKey && hasClientCert {
|
||||
cert, err := tls.X509KeyPair(certBytes.clientCert, certBytes.clientKey)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
certs.clientCerts = []tls.Certificate{cert}
|
||||
}
|
||||
|
||||
if hasClientKey && (!hasClientCert) {
|
||||
return nil, errPartialClientCertNoCert
|
||||
}
|
||||
|
||||
if hasClientCert && (!hasClientKey) {
|
||||
return nil, errPartialClientCertNoKey
|
||||
}
|
||||
|
||||
return &certs, nil
|
||||
}
|
||||
@@ -0,0 +1,402 @@
|
||||
package tls
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os"
|
||||
"testing"
|
||||
|
||||
"github.com/grafana/grafana/pkg/tsdb/sqleng"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func noReadFile(path string) ([]byte, error) {
|
||||
return nil, errors.New("not implemented")
|
||||
}
|
||||
|
||||
func TestTLSNoMode(t *testing.T) {
|
||||
// for backward-compatibility reason,
|
||||
// when mode is unset, it defaults to `require`
|
||||
dsInfo := sqleng.DataSourceInfo{
|
||||
JsonData: sqleng.JsonData{
|
||||
ConfigurationMethod: "",
|
||||
},
|
||||
}
|
||||
c, err := GetTLSConfig(dsInfo, noReadFile, "localhost")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, c)
|
||||
require.True(t, c.InsecureSkipVerify)
|
||||
}
|
||||
|
||||
func TestTLSDisable(t *testing.T) {
|
||||
dsInfo := sqleng.DataSourceInfo{
|
||||
JsonData: sqleng.JsonData{
|
||||
Mode: "disable",
|
||||
ConfigurationMethod: "",
|
||||
},
|
||||
}
|
||||
c, err := GetTLSConfig(dsInfo, noReadFile, "localhost")
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, c)
|
||||
}
|
||||
|
||||
func TestTLSRequire(t *testing.T) {
|
||||
dsInfo := sqleng.DataSourceInfo{
|
||||
JsonData: sqleng.JsonData{
|
||||
Mode: "require",
|
||||
ConfigurationMethod: "",
|
||||
},
|
||||
}
|
||||
c, err := GetTLSConfig(dsInfo, noReadFile, "localhost")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, c)
|
||||
require.True(t, c.InsecureSkipVerify)
|
||||
require.Nil(t, c.RootCAs)
|
||||
}
|
||||
|
||||
func TestTLSRequireWithRootCert(t *testing.T) {
|
||||
rootCertBytes, err := CreateRandomRootCertBytes()
|
||||
require.NoError(t, err)
|
||||
|
||||
dsInfo := sqleng.DataSourceInfo{
|
||||
JsonData: sqleng.JsonData{
|
||||
Mode: "require",
|
||||
ConfigurationMethod: "file-content",
|
||||
},
|
||||
DecryptedSecureJSONData: map[string]string{
|
||||
"tlsCACert": string(rootCertBytes),
|
||||
},
|
||||
}
|
||||
c, err := GetTLSConfig(dsInfo, noReadFile, "localhost")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, c)
|
||||
require.True(t, c.InsecureSkipVerify)
|
||||
require.NotNil(t, c.VerifyConnection)
|
||||
require.NotNil(t, c.RootCAs) // TODO: not the best, but nothing better available
|
||||
}
|
||||
|
||||
func TestTLSVerifyCA(t *testing.T) {
|
||||
rootCertBytes, err := CreateRandomRootCertBytes()
|
||||
require.NoError(t, err)
|
||||
|
||||
dsInfo := sqleng.DataSourceInfo{
|
||||
JsonData: sqleng.JsonData{
|
||||
Mode: "verify-ca",
|
||||
ConfigurationMethod: "file-content",
|
||||
},
|
||||
DecryptedSecureJSONData: map[string]string{
|
||||
"tlsCACert": string(rootCertBytes),
|
||||
},
|
||||
}
|
||||
c, err := GetTLSConfig(dsInfo, noReadFile, "localhost")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, c)
|
||||
require.True(t, c.InsecureSkipVerify)
|
||||
require.NotNil(t, c.VerifyConnection)
|
||||
require.NotNil(t, c.RootCAs) // TODO: not the best, but nothing better available
|
||||
}
|
||||
|
||||
func TestTLSVerifyCANoRootCertProvided(t *testing.T) {
|
||||
// this is ok. go will use the default system certs
|
||||
dsInfo := sqleng.DataSourceInfo{
|
||||
JsonData: sqleng.JsonData{
|
||||
Mode: "verify-ca",
|
||||
ConfigurationMethod: "file-content",
|
||||
},
|
||||
DecryptedSecureJSONData: map[string]string{},
|
||||
}
|
||||
_, err := GetTLSConfig(dsInfo, noReadFile, "localhost")
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestTLSClientCert(t *testing.T) {
|
||||
clientKey, clientCert, err := CreateRandomClientCert()
|
||||
require.NoError(t, err)
|
||||
|
||||
dsInfo := sqleng.DataSourceInfo{
|
||||
JsonData: sqleng.JsonData{
|
||||
Mode: "require",
|
||||
ConfigurationMethod: "file-content",
|
||||
},
|
||||
DecryptedSecureJSONData: map[string]string{
|
||||
"tlsClientCert": string(clientCert),
|
||||
"tlsClientKey": string(clientKey),
|
||||
},
|
||||
}
|
||||
c, err := GetTLSConfig(dsInfo, noReadFile, "localhost")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, c)
|
||||
require.Len(t, c.Certificates, 1)
|
||||
}
|
||||
|
||||
func TestTLSMethodFileContentClientCertMissingKey(t *testing.T) {
|
||||
_, clientCert, err := CreateRandomClientCert()
|
||||
require.NoError(t, err)
|
||||
|
||||
dsInfo := sqleng.DataSourceInfo{
|
||||
JsonData: sqleng.JsonData{
|
||||
Mode: "require",
|
||||
ConfigurationMethod: "file-content",
|
||||
},
|
||||
DecryptedSecureJSONData: map[string]string{
|
||||
"tlsClientCert": string(clientCert),
|
||||
},
|
||||
}
|
||||
_, err = GetTLSConfig(dsInfo, noReadFile, "localhost")
|
||||
require.ErrorIs(t, err, errPartialClientCertNoKey)
|
||||
}
|
||||
|
||||
func TestTLSMethodFileContentClientCertMissingCert(t *testing.T) {
|
||||
clientKey, _, err := CreateRandomClientCert()
|
||||
require.NoError(t, err)
|
||||
|
||||
dsInfo := sqleng.DataSourceInfo{
|
||||
JsonData: sqleng.JsonData{
|
||||
Mode: "require",
|
||||
ConfigurationMethod: "file-content",
|
||||
},
|
||||
DecryptedSecureJSONData: map[string]string{
|
||||
"tlsClientKey": string(clientKey),
|
||||
},
|
||||
}
|
||||
_, err = GetTLSConfig(dsInfo, noReadFile, "localhost")
|
||||
require.ErrorIs(t, err, errPartialClientCertNoCert)
|
||||
}
|
||||
|
||||
func TestTLSMethodFilePathClientCertMissingKey(t *testing.T) {
|
||||
clientKey, _, err := CreateRandomClientCert()
|
||||
require.NoError(t, err)
|
||||
|
||||
readFile := newMockReadFile(map[string]([]byte){
|
||||
"path1": clientKey,
|
||||
})
|
||||
|
||||
dsInfo := sqleng.DataSourceInfo{
|
||||
JsonData: sqleng.JsonData{
|
||||
Mode: "require",
|
||||
ConfigurationMethod: "file-path",
|
||||
CertKeyFile: "path1",
|
||||
},
|
||||
}
|
||||
_, err = GetTLSConfig(dsInfo, readFile, "localhost")
|
||||
require.ErrorIs(t, err, errPartialClientCertNoCert)
|
||||
}
|
||||
|
||||
func TestTLSMethodFilePathClientCertMissingCert(t *testing.T) {
|
||||
_, clientCert, err := CreateRandomClientCert()
|
||||
require.NoError(t, err)
|
||||
|
||||
readFile := newMockReadFile(map[string]([]byte){
|
||||
"path1": clientCert,
|
||||
})
|
||||
|
||||
dsInfo := sqleng.DataSourceInfo{
|
||||
JsonData: sqleng.JsonData{
|
||||
Mode: "require",
|
||||
ConfigurationMethod: "file-path",
|
||||
CertFile: "path1",
|
||||
},
|
||||
}
|
||||
_, err = GetTLSConfig(dsInfo, readFile, "localhost")
|
||||
require.ErrorIs(t, err, errPartialClientCertNoKey)
|
||||
}
|
||||
|
||||
func TestTLSVerifyFull(t *testing.T) {
|
||||
rootCertBytes, err := CreateRandomRootCertBytes()
|
||||
require.NoError(t, err)
|
||||
|
||||
dsInfo := sqleng.DataSourceInfo{
|
||||
JsonData: sqleng.JsonData{
|
||||
Mode: "verify-full",
|
||||
ConfigurationMethod: "file-content",
|
||||
},
|
||||
DecryptedSecureJSONData: map[string]string{
|
||||
"tlsCACert": string(rootCertBytes),
|
||||
},
|
||||
}
|
||||
c, err := GetTLSConfig(dsInfo, noReadFile, "localhost")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, c)
|
||||
require.False(t, c.InsecureSkipVerify)
|
||||
require.Nil(t, c.VerifyConnection)
|
||||
require.NotNil(t, c.RootCAs) // TODO: not the best, but nothing better available
|
||||
}
|
||||
|
||||
func TestTLSMethodFileContent(t *testing.T) {
|
||||
rootCertBytes, err := CreateRandomRootCertBytes()
|
||||
require.NoError(t, err)
|
||||
|
||||
clientKey, clientCert, err := CreateRandomClientCert()
|
||||
require.NoError(t, err)
|
||||
|
||||
dsInfo := sqleng.DataSourceInfo{
|
||||
JsonData: sqleng.JsonData{
|
||||
Mode: "verify-full",
|
||||
ConfigurationMethod: "file-content",
|
||||
},
|
||||
DecryptedSecureJSONData: map[string]string{
|
||||
"tlsCACert": string(rootCertBytes),
|
||||
"tlsClientCert": string(clientCert),
|
||||
"tlsClientKey": string(clientKey),
|
||||
},
|
||||
}
|
||||
c, err := GetTLSConfig(dsInfo, noReadFile, "localhost")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, c)
|
||||
require.Len(t, c.Certificates, 1)
|
||||
require.NotNil(t, c.RootCAs) // TODO: not the best, but nothing better available
|
||||
}
|
||||
|
||||
func TestTLSMethodFilePath(t *testing.T) {
|
||||
rootCertBytes, err := CreateRandomRootCertBytes()
|
||||
require.NoError(t, err)
|
||||
|
||||
clientKey, clientCert, err := CreateRandomClientCert()
|
||||
require.NoError(t, err)
|
||||
|
||||
readFile := newMockReadFile(map[string]([]byte){
|
||||
"root-cert-path": rootCertBytes,
|
||||
"client-key-path": clientKey,
|
||||
"client-cert-path": clientCert,
|
||||
})
|
||||
|
||||
dsInfo := sqleng.DataSourceInfo{
|
||||
JsonData: sqleng.JsonData{
|
||||
Mode: "verify-full",
|
||||
ConfigurationMethod: "file-path",
|
||||
RootCertFile: "root-cert-path",
|
||||
CertKeyFile: "client-key-path",
|
||||
CertFile: "client-cert-path",
|
||||
},
|
||||
}
|
||||
c, err := GetTLSConfig(dsInfo, readFile, "localhost")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, c)
|
||||
require.Len(t, c.Certificates, 1)
|
||||
require.NotNil(t, c.RootCAs) // TODO: not the best, but nothing better available
|
||||
}
|
||||
|
||||
func TestTLSMethodFilePathRootCertDoesNotExist(t *testing.T) {
|
||||
readFile := newMockReadFile(map[string]([]byte){})
|
||||
|
||||
dsInfo := sqleng.DataSourceInfo{
|
||||
JsonData: sqleng.JsonData{
|
||||
Mode: "verify-full",
|
||||
ConfigurationMethod: "file-path",
|
||||
RootCertFile: "path1",
|
||||
},
|
||||
}
|
||||
_, err := GetTLSConfig(dsInfo, readFile, "localhost")
|
||||
require.ErrorIs(t, err, os.ErrNotExist)
|
||||
}
|
||||
|
||||
func TestTLSMethodFilePathClientCertKeyDoesNotExist(t *testing.T) {
|
||||
_, clientCert, err := CreateRandomClientCert()
|
||||
require.NoError(t, err)
|
||||
|
||||
readFile := newMockReadFile(map[string]([]byte){
|
||||
"cert-path": clientCert,
|
||||
})
|
||||
|
||||
dsInfo := sqleng.DataSourceInfo{
|
||||
JsonData: sqleng.JsonData{
|
||||
Mode: "require",
|
||||
ConfigurationMethod: "file-path",
|
||||
CertKeyFile: "key-path",
|
||||
CertFile: "cert-path",
|
||||
},
|
||||
}
|
||||
_, err = GetTLSConfig(dsInfo, readFile, "localhost")
|
||||
require.ErrorIs(t, err, os.ErrNotExist)
|
||||
}
|
||||
|
||||
func TestTLSMethodFilePathClientCertCertDoesNotExist(t *testing.T) {
|
||||
clientKey, _, err := CreateRandomClientCert()
|
||||
require.NoError(t, err)
|
||||
|
||||
readFile := newMockReadFile(map[string]([]byte){
|
||||
"key-path": clientKey,
|
||||
})
|
||||
|
||||
dsInfo := sqleng.DataSourceInfo{
|
||||
JsonData: sqleng.JsonData{
|
||||
Mode: "require",
|
||||
ConfigurationMethod: "file-path",
|
||||
CertKeyFile: "key-path",
|
||||
CertFile: "cert-path",
|
||||
},
|
||||
}
|
||||
_, err = GetTLSConfig(dsInfo, readFile, "localhost")
|
||||
require.ErrorIs(t, err, os.ErrNotExist)
|
||||
}
|
||||
|
||||
// method="" equals to method="file-path"
|
||||
func TestTLSMethodEmpty(t *testing.T) {
|
||||
rootCertBytes, err := CreateRandomRootCertBytes()
|
||||
require.NoError(t, err)
|
||||
|
||||
clientKey, clientCert, err := CreateRandomClientCert()
|
||||
require.NoError(t, err)
|
||||
|
||||
readFile := newMockReadFile(map[string]([]byte){
|
||||
"root-cert-path": rootCertBytes,
|
||||
"client-key-path": clientKey,
|
||||
"client-cert-path": clientCert,
|
||||
})
|
||||
|
||||
dsInfo := sqleng.DataSourceInfo{
|
||||
JsonData: sqleng.JsonData{
|
||||
Mode: "verify-full",
|
||||
ConfigurationMethod: "",
|
||||
RootCertFile: "root-cert-path",
|
||||
CertKeyFile: "client-key-path",
|
||||
CertFile: "client-cert-path",
|
||||
},
|
||||
}
|
||||
c, err := GetTLSConfig(dsInfo, readFile, "localhost")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, c)
|
||||
require.Len(t, c.Certificates, 1)
|
||||
require.NotNil(t, c.RootCAs) // TODO: not the best, but nothing better available
|
||||
}
|
||||
|
||||
func TestTLSVerifyFullNoRootCertProvided(t *testing.T) {
|
||||
// this is ok. go will use the default system certs
|
||||
dsInfo := sqleng.DataSourceInfo{
|
||||
JsonData: sqleng.JsonData{
|
||||
Mode: "verify-full",
|
||||
ConfigurationMethod: "file-content",
|
||||
},
|
||||
DecryptedSecureJSONData: map[string]string{},
|
||||
}
|
||||
_, err := GetTLSConfig(dsInfo, noReadFile, "localhost")
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestTLSInvalidMode(t *testing.T) {
|
||||
dsInfo := sqleng.DataSourceInfo{
|
||||
JsonData: sqleng.JsonData{
|
||||
Mode: "not-a-valid-mode",
|
||||
},
|
||||
}
|
||||
|
||||
_, err := GetTLSConfig(dsInfo, noReadFile, "localhost")
|
||||
require.Error(t, err)
|
||||
}
|
||||
|
||||
func TestTLSServerNameSetInEveryMode(t *testing.T) {
|
||||
modes := []string{"require", "verify-ca", "verify-full"}
|
||||
|
||||
for _, mode := range modes {
|
||||
t.Run(mode, func(t *testing.T) {
|
||||
dsInfo := sqleng.DataSourceInfo{
|
||||
JsonData: sqleng.JsonData{
|
||||
Mode: mode,
|
||||
},
|
||||
DecryptedSecureJSONData: map[string]string{},
|
||||
}
|
||||
c, err := GetTLSConfig(dsInfo, noReadFile, "example.com")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "example.com", c.ServerName)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,105 @@
|
||||
package tls
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"crypto/x509"
|
||||
"crypto/x509/pkix"
|
||||
"encoding/pem"
|
||||
"math/big"
|
||||
"os"
|
||||
"time"
|
||||
)
|
||||
|
||||
func CreateRandomRootCertBytes() ([]byte, error) {
|
||||
cert := x509.Certificate{
|
||||
SerialNumber: big.NewInt(42),
|
||||
Subject: pkix.Name{
|
||||
CommonName: "test1",
|
||||
},
|
||||
NotBefore: time.Now(),
|
||||
NotAfter: time.Now().AddDate(10, 0, 0),
|
||||
IsCA: true,
|
||||
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
|
||||
KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageCertSign,
|
||||
BasicConstraintsValid: true,
|
||||
}
|
||||
|
||||
key, err := rsa.GenerateKey(rand.Reader, 2048)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
bytes, err := x509.CreateCertificate(rand.Reader, &cert, &cert, &key.PublicKey, key)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return pem.EncodeToMemory(&pem.Block{
|
||||
Type: "CERTIFICATE",
|
||||
Bytes: bytes,
|
||||
}), nil
|
||||
}
|
||||
|
||||
func CreateRandomClientCert() ([]byte, []byte, error) {
|
||||
caKey, err := rsa.GenerateKey(rand.Reader, 2048)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
key, err := rsa.GenerateKey(rand.Reader, 2048)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
keyBytes := pem.EncodeToMemory(&pem.Block{
|
||||
Type: "RSA PRIVATE KEY",
|
||||
Bytes: x509.MarshalPKCS1PrivateKey(key),
|
||||
})
|
||||
|
||||
caCert := x509.Certificate{
|
||||
SerialNumber: big.NewInt(42),
|
||||
Subject: pkix.Name{
|
||||
CommonName: "test1",
|
||||
},
|
||||
NotBefore: time.Now(),
|
||||
NotAfter: time.Now().AddDate(10, 0, 0),
|
||||
IsCA: true,
|
||||
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth},
|
||||
KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageCertSign,
|
||||
BasicConstraintsValid: true,
|
||||
}
|
||||
|
||||
cert := x509.Certificate{
|
||||
SerialNumber: big.NewInt(2019),
|
||||
Subject: pkix.Name{
|
||||
CommonName: "test1",
|
||||
},
|
||||
NotBefore: time.Now(),
|
||||
NotAfter: time.Now().AddDate(10, 0, 0),
|
||||
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth},
|
||||
KeyUsage: x509.KeyUsageDigitalSignature,
|
||||
}
|
||||
|
||||
certData, err := x509.CreateCertificate(rand.Reader, &cert, &caCert, &key.PublicKey, caKey)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
certBytes := pem.EncodeToMemory(&pem.Block{
|
||||
Type: "CERTIFICATE",
|
||||
Bytes: certData,
|
||||
})
|
||||
|
||||
return keyBytes, certBytes, nil
|
||||
}
|
||||
|
||||
func newMockReadFile(data map[string]([]byte)) ReadFileFunc {
|
||||
return func(path string) ([]byte, error) {
|
||||
bytes, ok := data[path]
|
||||
if !ok {
|
||||
return nil, os.ErrNotExist
|
||||
}
|
||||
return bytes, nil
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user