a73x

internal/transport/tlsconf_test.go

Ref:   Size: 3.3 KiB   History

package transport

import (
	"crypto/tls"
	"crypto/x509"
	"encoding/pem"
	"testing"
	"time"

	"github.com/stretchr/testify/assert"
	"github.com/stretchr/testify/require"
)

func TestGenerateAndPinRoundTrip(t *testing.T) {
	certPEM, keyPEM, err := GenerateServerCert()
	require.NoError(t, err)

	fp, err := CertFingerprint(certPEM)
	require.NoError(t, err)
	assert.Len(t, fp, 64) // hex sha256

	cert, err := tls.X509KeyPair(certPEM, keyPEM)
	require.NoError(t, err)
	leaf, err := x509.ParseCertificate(cert.Certificate[0])
	require.NoError(t, err)

	cc := ClientTLS(fp)
	state := tls.ConnectionState{PeerCertificates: []*x509.Certificate{leaf}}
	require.NoError(t, cc.VerifyConnection(state), "matching fingerprint must verify")
}

func TestPinRejectsWrongCert(t *testing.T) {
	c1, _, _ := GenerateServerCert()
	fp1, _ := CertFingerprint(c1)
	c2, k2, _ := GenerateServerCert()
	cert2, _ := tls.X509KeyPair(c2, k2)
	leaf2, _ := x509.ParseCertificate(cert2.Certificate[0])

	cc := ClientTLS(fp1)
	state := tls.ConnectionState{PeerCertificates: []*x509.Certificate{leaf2}}
	require.Error(t, cc.VerifyConnection(state), "non-matching fingerprint must be rejected")
}

func TestServerTLSHasALPN(t *testing.T) {
	certPEM, keyPEM, err := GenerateServerCert()
	require.NoError(t, err)
	sc, err := ServerTLS(certPEM, keyPEM)
	require.NoError(t, err)
	assert.Contains(t, sc.NextProtos, ALPN)
}

// TestGeneratedCertValidityIsTwoYears pins the cert lifetime at ~2 years. The
// pin (VerifyConnection) ignores expiry, so a longer validity buys nothing —
// it only widens the forgery window if server.key leaks. Two years sets the
// rotation cadence the renewal warning drives.
func TestGeneratedCertValidityIsTwoYears(t *testing.T) {
	certPEM, _, err := GenerateServerCert()
	require.NoError(t, err)
	block, _ := pem.Decode(certPEM)
	require.NotNil(t, block)
	cert, err := x509.ParseCertificate(block.Bytes)
	require.NoError(t, err)

	lifetime := cert.NotAfter.Sub(cert.NotBefore)
	assert.LessOrEqual(t, lifetime, 2*365*24*time.Hour+31*24*time.Hour,
		"validity must be ~2 years, not the old 10")
	assert.Greater(t, lifetime, 365*24*time.Hour, "validity must exceed 1 year")
}

// TestCertRenewalDue pins the renewal-warning helper: due when now is within
// 90 days of NotAfter (or past it), not due before that, and an unparseable
// cert reports due (fail-loud: an operator should look at a broken cert).
func TestCertRenewalDue(t *testing.T) {
	certPEM, _, err := GenerateServerCert()
	require.NoError(t, err)
	block, _ := pem.Decode(certPEM)
	require.NotNil(t, block)
	cert, err := x509.ParseCertificate(block.Bytes)
	require.NoError(t, err)

	notAfter, due := CertRenewalDue(certPEM, cert.NotAfter.Add(-180*24*time.Hour))
	assert.False(t, due, "6 months out must not be due")
	assert.Equal(t, cert.NotAfter, notAfter)

	_, due = CertRenewalDue(certPEM, cert.NotAfter.Add(-30*24*time.Hour))
	assert.True(t, due, "30 days out must be due")

	_, due = CertRenewalDue(certPEM, cert.NotAfter.Add(24*time.Hour))
	assert.True(t, due, "past expiry must be due")

	_, due = CertRenewalDue([]byte("not a cert"), time.Now())
	assert.True(t, due, "non-PEM input must report due (fail loud)")

	junkDER := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: []byte("junk")})
	_, due = CertRenewalDue(junkDER, time.Now())
	assert.True(t, due, "valid PEM with garbage DER must report due (fail loud)")
}