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 TestDialRequiresTLSConfig(t *testing.T) { if _, err := Dial("bufnet", nil); err == nil { t.Fatal("expected TLS configuration requirement") } } 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 }