Files

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
}