256 lines
8.0 KiB
Go
256 lines
8.0 KiB
Go
package rpc
|
|
|
|
import (
|
|
"crypto/rand"
|
|
"crypto/rsa"
|
|
"crypto/tls"
|
|
"crypto/x509"
|
|
"crypto/x509/pkix"
|
|
"encoding/pem"
|
|
"math/big"
|
|
"net"
|
|
"testing"
|
|
"time"
|
|
)
|
|
|
|
func TestTLSConfigsRequireVerifiedSANPeer(t *testing.T) {
|
|
caPEM, caCert, caKey := testCertificate(t, nil, nil, true, nil, nil)
|
|
serverPEM, _, _ := testCertificate(t, caCert, caKey, false, []string{"agent.local"}, nil)
|
|
clientPEM, _, _ := testCertificate(t, caCert, caKey, false, []string{"dispatcher.local"}, nil)
|
|
|
|
serverConfig, err := NewServerTLSConfig(caPEM.certPEM, serverPEM.certPEM, serverPEM.keyPEM)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if serverConfig.MinVersion != 0x0304 {
|
|
t.Fatalf("MinVersion = %v, want TLS 1.3", serverConfig.MinVersion)
|
|
}
|
|
if serverConfig.ClientAuth != 4 {
|
|
t.Fatalf("ClientAuth = %v, want RequireAndVerifyClientCert", serverConfig.ClientAuth)
|
|
}
|
|
clientConfig, err := NewClientTLSConfig(caPEM.certPEM, clientPEM.certPEM, clientPEM.keyPEM, "agent.local")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if clientConfig.ServerName != "agent.local" || clientConfig.RootCAs == nil {
|
|
t.Fatalf("client config does not verify the configured server name")
|
|
}
|
|
if _, err := NewClientTLSConfig(caPEM.certPEM, clientPEM.certPEM, clientPEM.keyPEM, ""); err == nil {
|
|
t.Fatal("expected empty server name to be rejected")
|
|
}
|
|
if CertificateFingerprint(caCert) == "" {
|
|
t.Fatal("expected certificate fingerprint")
|
|
}
|
|
}
|
|
|
|
func TestTLSConfigsCompleteMutualHandshake(t *testing.T) {
|
|
caPEM, caCert, caKey := testCertificate(t, nil, nil, true, nil, nil)
|
|
serverPEM, _, _ := testCertificate(t, caCert, caKey, false, []string{"agent.local"}, nil)
|
|
clientPEM, _, _ := testCertificate(t, caCert, caKey, false, []string{"dispatcher.local"}, nil)
|
|
serverConfig, err := NewServerTLSConfig(caPEM.certPEM, serverPEM.certPEM, serverPEM.keyPEM)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
clientConfig, err := NewClientTLSConfig(caPEM.certPEM, clientPEM.certPEM, clientPEM.keyPEM, "agent.local")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
listener, err := tls.Listen("tcp", "127.0.0.1:0", serverConfig)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer listener.Close()
|
|
serverDone := make(chan error, 1)
|
|
go func() {
|
|
conn, acceptErr := listener.Accept()
|
|
if acceptErr != nil {
|
|
serverDone <- acceptErr
|
|
return
|
|
}
|
|
defer conn.Close()
|
|
serverDone <- conn.(*tls.Conn).Handshake()
|
|
}()
|
|
client, err := tls.Dial("tcp", listener.Addr().String(), clientConfig)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := client.Handshake(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
_ = client.Close()
|
|
if err := <-serverDone; err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
|
|
func TestTLSConfigsRejectMissingOrUntrustedClient(t *testing.T) {
|
|
caPEM, caCert, caKey := testCertificate(t, nil, nil, true, nil, nil)
|
|
serverPEM, _, _ := testCertificate(t, caCert, caKey, false, []string{"agent.local"}, nil)
|
|
serverConfig, err := NewServerTLSConfig(caPEM.certPEM, serverPEM.certPEM, serverPEM.keyPEM)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
_, rogueCACert, rogueCAKey := testCertificate(t, nil, nil, true, nil, nil)
|
|
rogueClientPEM, _, _ := testCertificate(t, rogueCACert, rogueCAKey, false, []string{"dispatcher.local"}, nil)
|
|
trustedServerPool, err := certPool(caPEM.certPEM)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
clients := map[string]*tls.Config{
|
|
"missing client certificate": {
|
|
MinVersion: tls.VersionTLS13,
|
|
RootCAs: trustedServerPool,
|
|
ServerName: "agent.local",
|
|
NextProtos: []string{"h2"},
|
|
},
|
|
"untrusted client certificate": func() *tls.Config {
|
|
clientConfig, configErr := NewClientTLSConfig(caPEM.certPEM, rogueClientPEM.certPEM, rogueClientPEM.keyPEM, "agent.local")
|
|
if configErr != nil {
|
|
t.Fatal(configErr)
|
|
}
|
|
return clientConfig
|
|
}(),
|
|
}
|
|
for name, clientConfig := range clients {
|
|
t.Run(name, func(t *testing.T) {
|
|
listener, listenErr := tls.Listen("tcp", "127.0.0.1:0", serverConfig)
|
|
if listenErr != nil {
|
|
t.Fatal(listenErr)
|
|
}
|
|
defer listener.Close()
|
|
serverDone := make(chan error, 1)
|
|
go func() {
|
|
conn, acceptErr := listener.Accept()
|
|
if acceptErr != nil {
|
|
serverDone <- acceptErr
|
|
return
|
|
}
|
|
defer conn.Close()
|
|
serverDone <- conn.(*tls.Conn).Handshake()
|
|
}()
|
|
client, dialErr := tls.Dial("tcp", listener.Addr().String(), clientConfig)
|
|
if client != nil {
|
|
_ = client.Close()
|
|
}
|
|
_ = dialErr
|
|
if serverErr := <-serverDone; serverErr == nil {
|
|
t.Fatal("expected server handshake to reject the client")
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestTLSRotationRejectsPreviousClientCA(t *testing.T) {
|
|
caOnePEM, caOneCert, caOneKey := testCertificate(t, nil, nil, true, nil, nil)
|
|
serverOnePEM, _, _ := testCertificate(t, caOneCert, caOneKey, false, []string{"agent.local"}, nil)
|
|
clientOnePEM, _, _ := testCertificate(t, caOneCert, caOneKey, false, []string{"dispatcher.local"}, nil)
|
|
serverOneConfig, err := NewServerTLSConfig(caOnePEM.certPEM, serverOnePEM.certPEM, serverOnePEM.keyPEM)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
clientOneConfig, err := NewClientTLSConfig(caOnePEM.certPEM, clientOnePEM.certPEM, clientOnePEM.keyPEM, "agent.local")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := runTLSHandshake(t, serverOneConfig, clientOneConfig); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
|
|
caTwoPEM, caTwoCert, caTwoKey := testCertificate(t, nil, nil, true, nil, nil)
|
|
serverTwoPEM, _, _ := testCertificate(t, caTwoCert, caTwoKey, false, []string{"agent.local"}, nil)
|
|
clientTwoPEM, _, _ := testCertificate(t, caTwoCert, caTwoKey, false, []string{"dispatcher.local"}, nil)
|
|
serverTwoConfig, err := NewServerTLSConfig(caTwoPEM.certPEM, serverTwoPEM.certPEM, serverTwoPEM.keyPEM)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
clientTwoConfig, err := NewClientTLSConfig(caTwoPEM.certPEM, clientTwoPEM.certPEM, clientTwoPEM.keyPEM, "agent.local")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := runTLSHandshake(t, serverTwoConfig, clientOneConfig); err == nil {
|
|
t.Fatal("expected previous client CA to be rejected after rotation")
|
|
}
|
|
if err := runTLSHandshake(t, serverTwoConfig, clientTwoConfig); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
|
|
func runTLSHandshake(t *testing.T, serverConfig, clientConfig *tls.Config) error {
|
|
t.Helper()
|
|
listener, err := tls.Listen("tcp", "127.0.0.1:0", serverConfig)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer listener.Close()
|
|
serverDone := make(chan error, 1)
|
|
go func() {
|
|
conn, acceptErr := listener.Accept()
|
|
if acceptErr != nil {
|
|
serverDone <- acceptErr
|
|
return
|
|
}
|
|
defer conn.Close()
|
|
serverDone <- conn.(*tls.Conn).Handshake()
|
|
}()
|
|
client, clientErr := tls.Dial("tcp", listener.Addr().String(), clientConfig)
|
|
if client != nil {
|
|
_ = client.Close()
|
|
}
|
|
serverErr := <-serverDone
|
|
if clientErr != nil {
|
|
return clientErr
|
|
}
|
|
return serverErr
|
|
}
|
|
|
|
func testCertificate(t *testing.T, parent *x509.Certificate, parentKey *rsa.PrivateKey, isCA bool, dnsNames []string, ips []net.IP) (pemBundle, *x509.Certificate, *rsa.PrivateKey) {
|
|
t.Helper()
|
|
key, err := rsa.GenerateKey(rand.Reader, 2048)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
serial, err := rand.Int(rand.Reader, new(big.Int).Lsh(big.NewInt(1), 120))
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
now := time.Now()
|
|
template := &x509.Certificate{
|
|
SerialNumber: serial,
|
|
Subject: pkix.Name{CommonName: "sip-go-agent-test"},
|
|
NotBefore: now.Add(-time.Minute),
|
|
NotAfter: now.Add(time.Hour),
|
|
BasicConstraintsValid: true,
|
|
IsCA: isCA,
|
|
DNSNames: dnsNames,
|
|
IPAddresses: ips,
|
|
KeyUsage: x509.KeyUsageDigitalSignature,
|
|
}
|
|
if isCA {
|
|
template.KeyUsage |= x509.KeyUsageCertSign
|
|
}
|
|
if !isCA {
|
|
template.ExtKeyUsage = []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth, x509.ExtKeyUsageClientAuth}
|
|
}
|
|
if parent == nil {
|
|
parent = template
|
|
parentKey = key
|
|
}
|
|
der, err := x509.CreateCertificate(rand.Reader, template, parent, &key.PublicKey, parentKey)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
cert, err := x509.ParseCertificate(der)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
return pemBundle{
|
|
certPEM: pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}),
|
|
keyPEM: pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(key)}),
|
|
}, cert, key
|
|
}
|
|
|
|
type pemBundle struct {
|
|
certPEM []byte
|
|
keyPEM []byte
|
|
}
|