diff --git a/CHANGELOG.md b/CHANGELOG.md index 8dcc1b1..47384aa 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -14,6 +14,12 @@ ### Tests + Certinfo and requests: share CA/leaf certificate generation and custom TLS httptest servers via internal/tlstest. + + Tlstest: close the httptest listener when certificate loading fails. + + Tlstest: leave tls.Config.CipherSuites nil by default and accept only TLS 1.0-1.2 suite overrides. + Certinfo: GetRemoteCerts tests apply SetTLSInsecure before SetTLSEndpoint so endpoint certificate retrieval uses the intended TLS verification mode. Requests: httptest TLS servers bind an ephemeral port so parallel cases do not collide on fixed listeners. diff --git a/internal/certinfo/certinfo_handlers_test.go b/internal/certinfo/certinfo_handlers_test.go index 4eb9ed9..e145a4d 100644 --- a/internal/certinfo/certinfo_handlers_test.go +++ b/internal/certinfo/certinfo_handlers_test.go @@ -12,13 +12,15 @@ import ( "time" "github.com/stretchr/testify/require" + "github.com/xenos76/https-wrench/internal/tlstest" ) //nolint:revive func TestCertinfo_GetRemoteCerts(t *testing.T) { tests := []struct { desc string - srvCfg demoHTTPServerConfig + srvCfg tlstest.ServerConfig + serverName string caCertFile string insecure bool expectError bool @@ -27,20 +29,20 @@ func TestCertinfo_GetRemoteCerts(t *testing.T) { }{ { desc: "RSA Cert Success", - srvCfg: demoHTTPServerConfig{ - serverName: "example.com", - serverCertFile: RSASampleCertBundleFile, - serverKeyFile: RSASampleCertKeyFile, + srvCfg: tlstest.ServerConfig{ + ServerCertFile: RSASampleCertBundleFile, + ServerKeyFile: RSASampleCertKeyFile, }, + serverName: "example.com", caCertFile: RSACaCertFile, }, { desc: "Error Secure and No CA Cert", - srvCfg: demoHTTPServerConfig{ - serverName: "example.com", - serverCertFile: RSASampleCertFile, - serverKeyFile: RSASampleCertKeyFile, + srvCfg: tlstest.ServerConfig{ + ServerCertFile: RSASampleCertFile, + ServerKeyFile: RSASampleCertKeyFile, }, + serverName: "example.com", caCertFile: emptyString, //nolint:revive expectError: true, @@ -49,11 +51,11 @@ func TestCertinfo_GetRemoteCerts(t *testing.T) { { desc: "Malformed Server Certificate", - srvCfg: demoHTTPServerConfig{ - serverName: "example.com", - serverCertFile: RSASamplePKCS8Certificate, - serverKeyFile: RSASamplePKCS8PlaintextPrivateKey, + srvCfg: tlstest.ServerConfig{ + ServerCertFile: RSASamplePKCS8Certificate, + ServerKeyFile: RSASamplePKCS8PlaintextPrivateKey, }, + serverName: "example.com", caCertFile: RSACaCertFile, //nolint:revive expectError: true, @@ -61,21 +63,21 @@ func TestCertinfo_GetRemoteCerts(t *testing.T) { }, { desc: "No CA Cert and Insecure", - srvCfg: demoHTTPServerConfig{ - serverName: "example.com", - serverCertFile: RSASampleCertFile, - serverKeyFile: RSASampleCertKeyFile, + srvCfg: tlstest.ServerConfig{ + ServerCertFile: RSASampleCertFile, + ServerKeyFile: RSASampleCertKeyFile, }, + serverName: "example.com", insecure: true, caCertFile: emptyString, }, { desc: "Wrong CA Cert and Secure", - srvCfg: demoHTTPServerConfig{ - serverName: "example.com", - serverCertFile: RSASampleCertFile, - serverKeyFile: RSASampleCertKeyFile, + srvCfg: tlstest.ServerConfig{ + ServerCertFile: RSASampleCertFile, + ServerKeyFile: RSASampleCertKeyFile, }, + serverName: "example.com", caCertFile: RSASamplePKCS8Certificate, //nolint:revive expectError: true, @@ -83,66 +85,66 @@ func TestCertinfo_GetRemoteCerts(t *testing.T) { }, { desc: "Wrong CA Cert and Insecure", - srvCfg: demoHTTPServerConfig{ - serverName: "example.com", - serverCertFile: RSASampleCertFile, - serverKeyFile: RSASampleCertKeyFile, + srvCfg: tlstest.ServerConfig{ + ServerCertFile: RSASampleCertFile, + ServerKeyFile: RSASampleCertKeyFile, }, + serverName: "example.com", caCertFile: RSASamplePKCS8Certificate, insecure: true, }, { desc: "IPV6 Endpoint RSA Cert Success", - srvCfg: demoHTTPServerConfig{ - listenHost: "::1", - serverName: "example.com", - serverCertFile: RSASampleCertFile, - serverKeyFile: RSASampleCertKeyFile, + srvCfg: tlstest.ServerConfig{ + ListenHost: "::1", + ServerCertFile: RSASampleCertFile, + ServerKeyFile: RSASampleCertKeyFile, }, + serverName: "example.com", caCertFile: RSACaCertFile, }, { desc: "Error wrong ServerName", - srvCfg: demoHTTPServerConfig{ - serverName: "example.co.uk", - serverCertFile: RSASampleCertFile, + srvCfg: tlstest.ServerConfig{ + ServerCertFile: RSASampleCertFile, //nolint:revive - serverKeyFile: RSASampleCertKeyFile, + ServerKeyFile: RSASampleCertKeyFile, }, + serverName: "example.co.uk", caCertFile: RSACaCertFile, expectError: true, expectMsg: "TLS handshake failed: tls: failed to verify certificate: x509: certificate is valid for example.com, example.net, example.de, not example.co.uk", }, { desc: "X25519MLKEM768 key exchange", - srvCfg: demoHTTPServerConfig{ - serverName: "example.com", - serverCertFile: RSASampleCertBundleFile, - serverKeyFile: RSASampleCertKeyFile, - tlsCurvePreferences: []tls.CurveID{tls.X25519MLKEM768}, + srvCfg: tlstest.ServerConfig{ + ServerCertFile: RSASampleCertBundleFile, + ServerKeyFile: RSASampleCertKeyFile, + TLSCurvePreferences: []tls.CurveID{tls.X25519MLKEM768}, }, + serverName: "example.com", caCertFile: RSACaCertFile, wantCurveID: "X25519MLKEM768", }, { desc: "SecP256r1MLKEM768 key exchange", - srvCfg: demoHTTPServerConfig{ - serverName: "example.com", - serverCertFile: RSASampleCertBundleFile, - serverKeyFile: RSASampleCertKeyFile, - tlsCurvePreferences: []tls.CurveID{tls.SecP256r1MLKEM768}, + srvCfg: tlstest.ServerConfig{ + ServerCertFile: RSASampleCertBundleFile, + ServerKeyFile: RSASampleCertKeyFile, + TLSCurvePreferences: []tls.CurveID{tls.SecP256r1MLKEM768}, }, + serverName: "example.com", caCertFile: RSACaCertFile, wantCurveID: "SecP256r1MLKEM768", }, { desc: "SecP384r1MLKEM1024 key exchange", - srvCfg: demoHTTPServerConfig{ - serverName: "example.com", - serverCertFile: RSASampleCertBundleFile, - serverKeyFile: RSASampleCertKeyFile, - tlsCurvePreferences: []tls.CurveID{tls.SecP384r1MLKEM1024}, + srvCfg: tlstest.ServerConfig{ + ServerCertFile: RSASampleCertBundleFile, + ServerKeyFile: RSASampleCertKeyFile, + TLSCurvePreferences: []tls.CurveID{tls.SecP384r1MLKEM1024}, }, + serverName: "example.com", caCertFile: RSACaCertFile, wantCurveID: "SecP384r1MLKEM1024", }, @@ -153,18 +155,18 @@ func TestCertinfo_GetRemoteCerts(t *testing.T) { t.Run(tt.desc, func(t *testing.T) { t.Parallel() - ts, err := NewHTTPSTestServer(tt.srvCfg) + ts, err := tlstest.NewServer(tt.srvCfg) require.NoError(t, err) t.Cleanup(ts.Close) - endpoint := testServerHostPort(ts) + endpoint := ts.Listener.Addr().String() host, port, err := net.SplitHostPort(endpoint) require.NoError(t, err) cc, err := New() require.NoError(t, err) - cc.SetTLSServerName(tt.srvCfg.serverName) + cc.SetTLSServerName(tt.serverName) cc.SetCaPoolFromFile(tt.caCertFile, inputReader) cc.SetTLSInsecure(tt.insecure) cc.SetTLSEndpoint(t.Context(), endpoint) @@ -172,7 +174,7 @@ func TestCertinfo_GetRemoteCerts(t *testing.T) { err = cc.GetRemoteCerts(t.Context()) if !tt.expectError { require.NoError(t, err, "check error not expected") - require.Equal(t, tt.srvCfg.serverName, cc.TLSServerName, "check TLSServerName") + require.Equal(t, tt.serverName, cc.TLSServerName, "check TLSServerName") require.Equal(t, host, cc.TLSEndpointHost, "check TLSEndpointHost") require.Equal(t, port, cc.TLSEndpointPort, "check TLSEndpointPort") require.Equal(t, tt.insecure, cc.TLSInsecure, "check TLSInsecure") @@ -398,7 +400,7 @@ func TestCertinfo_PrintData(t *testing.T) { tlsEndpoint string tlsInsecure bool tlsServerName string - srvCfg demoHTTPServerConfig + srvCfg tlstest.ServerConfig expectCertsFetchErr bool expectCertsFetcMsg string }{ @@ -421,10 +423,9 @@ func TestCertinfo_PrintData(t *testing.T) { caCertFile: RSACaCertFile, keyCertMatch: true, tlsServerName: "example.com", - srvCfg: demoHTTPServerConfig{ - serverName: "example.com", - serverCertFile: RSASampleCertFile, - serverKeyFile: RSASampleCertKeyFile, + srvCfg: tlstest.ServerConfig{ + ServerCertFile: RSASampleCertFile, + ServerKeyFile: RSASampleCertKeyFile, }, }, { @@ -433,10 +434,9 @@ func TestCertinfo_PrintData(t *testing.T) { caCertFile: emptyString, tlsServerName: "example.com", //nolint:revive - srvCfg: demoHTTPServerConfig{ - serverName: "example.com", - serverCertFile: RSASampleCertFile, - serverKeyFile: RSASampleCertKeyFile, + srvCfg: tlstest.ServerConfig{ + ServerCertFile: RSASampleCertFile, + ServerKeyFile: RSASampleCertKeyFile, //nolint:revive }, //nolint:revive @@ -451,10 +451,9 @@ func TestCertinfo_PrintData(t *testing.T) { keyCertMatch: true, tlsInsecure: true, tlsServerName: "example.com", - srvCfg: demoHTTPServerConfig{ - serverName: "example.com", - serverCertFile: RSASampleCertFile, - serverKeyFile: RSASampleCertKeyFile, + srvCfg: tlstest.ServerConfig{ + ServerCertFile: RSASampleCertFile, + ServerKeyFile: RSASampleCertKeyFile, }, }, { @@ -463,11 +462,10 @@ func TestCertinfo_PrintData(t *testing.T) { caCertFile: RSACaCertFile, //nolint:revive tlsServerName: emptyString, - srvCfg: demoHTTPServerConfig{ - serverName: "example.com", - serverCertFile: RSASampleCertFile, + srvCfg: tlstest.ServerConfig{ + ServerCertFile: RSASampleCertFile, //nolint:revive - serverKeyFile: RSASampleCertKeyFile, + ServerKeyFile: RSASampleCertKeyFile, //nolint:revive }, expectCertsFetchErr: true, @@ -479,10 +477,9 @@ func TestCertinfo_PrintData(t *testing.T) { caCertFile: RSACaCertFile, keyCertMatch: false, tlsServerName: "example.com", - srvCfg: demoHTTPServerConfig{ - serverName: "example.com", - serverCertFile: RSASampleCertFile, - serverKeyFile: RSASampleCertKeyFile, + srvCfg: tlstest.ServerConfig{ + ServerCertFile: RSASampleCertFile, + ServerKeyFile: RSASampleCertKeyFile, }, }, } @@ -551,7 +548,7 @@ type printDataTestCase struct { tlsEndpoint string tlsInsecure bool tlsServerName string - srvCfg demoHTTPServerConfig + srvCfg tlstest.ServerConfig expectCertsFetchErr bool expectCertsFetcMsg string } @@ -566,12 +563,12 @@ func runPrintDataSubtest(t *testing.T, tt printDataTestCase) { require.NoError(t, cc.SetCertsFromFile(tt.certFile, inputReader)) require.NoError(t, cc.SetCaPoolFromFile(tt.caCertFile, inputReader)) - if tt.srvCfg.serverCertFile != emptyString { - ts, errSrv := NewHTTPSTestServer(tt.srvCfg) + if tt.srvCfg.ServerCertFile != emptyString { + ts, errSrv := tlstest.NewServer(tt.srvCfg) require.NoError(t, errSrv) t.Cleanup(ts.Close) - tt.tlsEndpoint = testServerHostPort(ts) + tt.tlsEndpoint = ts.Listener.Addr().String() if tt.tlsServerName == emptyString { // Cert SANs include 127.0.0.1; dial by hostname so empty SNI still mismatches. _, port, splitErr := net.SplitHostPort(tt.tlsEndpoint) diff --git a/internal/certinfo/main_test.go b/internal/certinfo/main_test.go index 630e74a..8932267 100644 --- a/internal/certinfo/main_test.go +++ b/internal/certinfo/main_test.go @@ -3,45 +3,18 @@ package certinfo import ( "crypto/rand" "crypto/rsa" - "crypto/tls" "crypto/x509" - "crypto/x509/pkix" "encoding/pem" "errors" "fmt" - "math/big" "net" - "net/http" - "net/http/httptest" "os" "testing" - "time" - "github.com/pires/go-proxyproto" + "github.com/xenos76/https-wrench/internal/tlstest" ) type ( - certificateTemplate struct { - cn string - isCA bool - dnsNames []string - ipAddresses []net.IP - key *rsa.PrivateKey - caKey *rsa.PrivateKey - parent *x509.Certificate - } - - demoHTTPServerConfig struct { - listenHost string - proxyprotoEnabled bool - serverName string - tlsCipherSuites []uint16 - tlsCurvePreferences []tls.CurveID - tlsMaxVersion uint16 - serverCertFile string - serverKeyFile string - } - MockErrReader struct{} MockInputReader struct{} mockReader struct{} @@ -180,17 +153,16 @@ func generateRSACertificateData() { fmt.Println(err) } - rsaSampleCertTpl := certificateTemplate{ - cn: "RSA Testing Sample Certificate", - isCA: false, - key: RSASampleCertKey, - caKey: RSACaCertKey, - parent: RSACaCertParent, - dnsNames: []string{"example.com", "example.net", "example.de"}, - ipAddresses: []net.IP{net.ParseIP("::1"), net.ParseIP("127.0.0.1")}, + rsaSampleCertTpl := tlstest.Template{ + CN: "RSA Testing Sample Certificate", + Key: RSASampleCertKey, + CAKey: RSACaCertKey, + Parent: RSACaCertParent, + DNSNames: []string{"example.com", "example.net", "example.de"}, + IPAddresses: []net.IP{net.ParseIP("::1"), net.ParseIP("127.0.0.1")}, } - RSASampleCertPEM, RSASampleCertParent, _ = GenerateCertificate( + RSASampleCertPEM, RSASampleCertParent, _ = tlstest.GenerateCert( rsaSampleCertTpl, ) RSASampleCertPEMString = string(RSASampleCertPEM) @@ -230,12 +202,12 @@ func generateRSACaData() { fmt.Println(err) } - rsaCaCertTpl := certificateTemplate{ - cn: "RSA Testing CA", - isCA: true, - key: RSACaCertKey, + rsaCaCertTpl := tlstest.Template{ + CN: "RSA Testing CA", + IsCA: true, + Key: RSACaCertKey, } - RSACaCertPEM, RSACaCertParent, _ = GenerateCertificate( + RSACaCertPEM, RSACaCertParent, _ = tlstest.GenerateCert( rsaCaCertTpl, ) RSACaCertPEMString = string(RSACaCertPEM) @@ -253,84 +225,6 @@ func generateRSACaData() { } } -// GenerateDemoCert takes as input a demoCertTemplate struct, creates a x509 Certificate -// and returns a PEM encoded version of the certificate, a pointer to the certificate and -// an error. -// The pointer can be used as parent in the creation of a new certificate linked to a self signed -// CA. -// -// Reference: https://shaneutt.com/blog/golang-ca-and-signed-cert-go/ -func GenerateCertificate(tpl certificateTemplate) ([]byte, *x509.Certificate, error) { - // Create a random serial number - serialNumberLimit := new(big.Int).Lsh(big.NewInt(1), 128) - - serialNumber, err := rand.Int(rand.Reader, serialNumberLimit) - if err != nil { - return nil, nil, fmt.Errorf("failed to generate serial number: %w", err) - } - - // Define template for non CA certificate - template := x509.Certificate{ - SerialNumber: serialNumber, - Subject: pkix.Name{ - CommonName: tpl.cn, - }, - NotBefore: time.Now(), - NotAfter: time.Now().Add(1 * 24 * time.Hour), - IsCA: false, - DNSNames: tpl.dnsNames, - IPAddresses: tpl.ipAddresses, - ExtKeyUsage: []x509.ExtKeyUsage{ - x509.ExtKeyUsageClientAuth, - x509.ExtKeyUsageServerAuth, - }, - KeyUsage: x509.KeyUsageDigitalSignature, - } - - certParent := tpl.parent - signingKey := tpl.caKey - - // in case of CA cert we update the template with the proper fields - // use the CA cert key for signing - // and do not reference any previous parent Certificate - if tpl.isCA { - signingKey = tpl.key - template = x509.Certificate{ - SerialNumber: serialNumber, - Subject: pkix.Name{ - CommonName: tpl.cn, - }, - NotBefore: time.Now(), - NotAfter: time.Now().Add(1 * 24 * time.Hour), - IsCA: true, - ExtKeyUsage: []x509.ExtKeyUsage{ - x509.ExtKeyUsageClientAuth, - x509.ExtKeyUsageServerAuth, - }, - KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageCertSign, - BasicConstraintsValid: true, - } - certParent = &template - } - - derBytes, err := x509.CreateCertificate(rand.Reader, - &template, certParent, &tpl.key.PublicKey, signingKey) - if err != nil { - return nil, nil, fmt.Errorf("failed to create certificate: %w", err) - } - - // Encode the certificate to PEM - certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: derBytes}) - - // parse the DER encoded x509.Certificate - certificate, err := x509.ParseCertificate(derBytes) - if err != nil { - return nil, nil, err - } - - return certPEM, certificate, nil -} - func createTmpFileWithContent( tempDir string, filePattern string, @@ -374,86 +268,3 @@ func RSAPrivateKeyToPEM(key *rsa.PrivateKey) []byte { return keyPEM } - -func testServerHostPort(ts *httptest.Server) string { - return ts.Listener.Addr().String() -} - -// NewHTTPSTestServer starts an httptest TLS server configured by cfg. -// Cipher suites default to the TLS 1.3 AEADs, CurvePreferences to the Go 1.27 -// PQ hybrids plus classical fallbacks, and MaxVersion to TLS 1.3. Non-empty -// cfg.tlsCipherSuites, cfg.tlsCurvePreferences, or a non-zero cfg.tlsMaxVersion -// override those defaults. Optional cfg.listenHost and cfg.proxyprotoEnabled -// replace the listener. The caller must Close the returned server. -func NewHTTPSTestServer(cfg demoHTTPServerConfig) (*httptest.Server, error) { - handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - fmt.Fprint(w, "DemoHTTPSServer Handler - client output\n") - fmt.Fprint(w, "Host requested: ", r.Host, "\n") - - fmt.Println("DemoHTTPSServer Handler - shell output") - }) - - ts := httptest.NewUnstartedServer(handler) - ts.EnableHTTP2 = true - - if cfg.listenHost != emptyString { - ln, err := net.Listen("tcp", net.JoinHostPort(cfg.listenHost, "0")) - if err != nil { - return nil, fmt.Errorf("error creating listener: %w", err) - } - - _ = ts.Listener.Close() - ts.Listener = ln - } - - if cfg.proxyprotoEnabled { - ts.Listener = &proxyproto.Listener{ - Listener: ts.Listener, - ReadHeaderTimeout: 10 * time.Second, - } - } - - cert, err := tls.LoadX509KeyPair( - cfg.serverCertFile, - cfg.serverKeyFile, - ) - if err != nil { - return nil, err - } - - // Set default TLS CipherSuites to TLS 1.3 cipher suites - // https://pkg.go.dev/crypto/tls#pkg-constants - tlsCipherSuites := []uint16{ - tls.TLS_AES_128_GCM_SHA256, - tls.TLS_AES_256_GCM_SHA384, - tls.TLS_CHACHA20_POLY1305_SHA256, - } - - if len(cfg.tlsCipherSuites) > 0 { - tlsCipherSuites = cfg.tlsCipherSuites - } - - tlsCurvePreferences := defaultCurvePreferences - - if len(cfg.tlsCurvePreferences) > 0 { - tlsCurvePreferences = cfg.tlsCurvePreferences - } - - // Set default TLS MaxVersion to 1.3 - var tlsMaxVersion uint16 = tls.VersionTLS13 - - if cfg.tlsMaxVersion > 0 { - tlsMaxVersion = cfg.tlsMaxVersion - } - - ts.TLS = &tls.Config{ - Certificates: []tls.Certificate{cert}, - CipherSuites: tlsCipherSuites, - CurvePreferences: tlsCurvePreferences, - MaxVersion: tlsMaxVersion, - } - - ts.StartTLS() - - return ts, nil -} diff --git a/internal/requests/main_test.go b/internal/requests/main_test.go index 5b4d3a3..ba72cd8 100644 --- a/internal/requests/main_test.go +++ b/internal/requests/main_test.go @@ -5,43 +5,19 @@ import ( "crypto/rsa" "crypto/tls" "crypto/x509" - "crypto/x509/pkix" "encoding/pem" "errors" "fmt" "io" - "math/big" "net" "net/http" - "net/http/httptest" "os" "testing" - "time" "github.com/alecthomas/assert/v2" - "github.com/pires/go-proxyproto" + "github.com/xenos76/https-wrench/internal/tlstest" ) -type demoCertTemplate struct { - cn string - isCA bool - dnsNames []string - ipAddresses []net.IP - key *rsa.PrivateKey - caKey *rsa.PrivateKey - parent *x509.Certificate -} - -//nolint:revive -type demoHttpServerData struct { - listenHost string - proxyprotoEnabled bool - serverName string - tlsCipherSuites []uint16 - tlsCurvePreferences []tls.CurveID - tlsMaxVersion uint16 -} - var ( testdataDir = "testdata" systemCertPool *x509.CertPool @@ -60,78 +36,6 @@ var ( tempDir string ) -// GenerateDemoCert takes as input a demoCertTemplate struct, creates a x509 Certificate -// and returns a PEM encoded version of the certificate, a pointer to the certificate and -// an error. -// The pointer can be used as parent in the creation of a new certificate linked to a self signed -// CA. -// -// Reference: https://shaneutt.com/blog/golang-ca-and-signed-cert-go/ -func GenerateDemoCert(tpl demoCertTemplate) ([]byte, *x509.Certificate, error) { - // Create a random serial number - serialNumberLimit := new(big.Int).Lsh(big.NewInt(1), 128) - - serialNumber, err := rand.Int(rand.Reader, serialNumberLimit) - if err != nil { - return nil, nil, fmt.Errorf("failed to generate serial number: %w", err) - } - - // Define template for non CA certificate - template := x509.Certificate{ - SerialNumber: serialNumber, - Subject: pkix.Name{ - CommonName: tpl.cn, - }, - NotBefore: time.Now(), - NotAfter: time.Now().Add(1 * 24 * time.Hour), - IsCA: false, - DNSNames: tpl.dnsNames, - IPAddresses: tpl.ipAddresses, - ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth, x509.ExtKeyUsageServerAuth}, - KeyUsage: x509.KeyUsageDigitalSignature, - } - - certParent := tpl.parent - signingKey := tpl.caKey - - // in case of CA cert we udpate the template with the proper fields - // use the CA cert key for signing - // and do not reference any previuous parent Certificate - if tpl.isCA { - certParent = &template - signingKey = tpl.key - template = x509.Certificate{ - SerialNumber: serialNumber, - Subject: pkix.Name{ - CommonName: tpl.cn, - }, - NotBefore: time.Now(), - NotAfter: time.Now().Add(1 * 24 * time.Hour), - IsCA: true, - ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth, x509.ExtKeyUsageServerAuth}, - KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageCertSign, - BasicConstraintsValid: true, - } - } - - derBytes, err := x509.CreateCertificate(rand.Reader, - &template, certParent, &tpl.key.PublicKey, signingKey) - if err != nil { - return nil, nil, fmt.Errorf("failed to create certificate: %w", err) - } - - // Encode the certificate to PEM - certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: derBytes}) - - // parse the DER encoded x509.Certificate - certificate, err := x509.ParseCertificate(derBytes) - if err != nil { - return nil, nil, err - } - - return certPEM, certificate, nil -} - func GenerateRSAKey(bits int) (*rsa.PrivateKey, error) { priv, err := rsa.GenerateKey(rand.Reader, bits) if err != nil { @@ -171,91 +75,6 @@ func printResponseBody(res *http.Response) { fmt.Println(string(body)) } -func testServerHostPort(ts *httptest.Server) string { - return ts.Listener.Addr().String() -} - -// NewHTTPSTestServer starts an httptest TLS server configured by data. -// Cipher suites default to the TLS 1.3 AEADs, CurvePreferences to the Go 1.27 -// PQ hybrids plus classical fallbacks, and MaxVersion to TLS 1.3. Non-empty -// data.tlsCipherSuites, data.tlsCurvePreferences, or a non-zero data.tlsMaxVersion -// override those defaults. Optional data.listenHost and data.proxyprotoEnabled -// replace the listener. The caller must Close the returned server. -func NewHTTPSTestServer(data demoHttpServerData) (*httptest.Server, error) { - handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - fmt.Fprint(w, "DemoHTTPSServer Handler - client output\n") - fmt.Fprint(w, "Host requested: ", r.Host, "\n") - - fmt.Println("DemoHTTPSServer Handler - shell output") - }) - - ts := httptest.NewUnstartedServer(handler) - ts.EnableHTTP2 = true - - if data.listenHost != emptyString { - ln, err := net.Listen("tcp", net.JoinHostPort(data.listenHost, "0")) - if err != nil { - return nil, fmt.Errorf("error creating listener: %w", err) - } - - _ = ts.Listener.Close() - ts.Listener = ln - } - - if data.proxyprotoEnabled { - ts.Listener = &proxyproto.Listener{ - Listener: ts.Listener, - ReadHeaderTimeout: 10 * time.Second, - } - } - - cert, err := tls.LoadX509KeyPair( - exampleCertFile, - exampleCertKeyFile, - ) - if err != nil { - return nil, err - } - - // Set default TLS CipherSuites to TLS 1.3 cipher suites - // https://pkg.go.dev/crypto/tls#pkg-constants - tlsCipherSuites := []uint16{ - tls.TLS_AES_128_GCM_SHA256, - tls.TLS_AES_256_GCM_SHA384, - tls.TLS_CHACHA20_POLY1305_SHA256, - } - - if len(data.tlsCipherSuites) > 0 { - tlsCipherSuites = data.tlsCipherSuites - } - - tlsCurvePreferences := defaultCurvePreferences - - if len(data.tlsCurvePreferences) > 0 { - tlsCurvePreferences = data.tlsCurvePreferences - } - - // Set default TLS MaxVersion to 1.3 - var tlsMaxVersion uint16 = tls.VersionTLS13 - - if data.tlsMaxVersion > 0 { - tlsMaxVersion = data.tlsMaxVersion - } - - ts.TLS = &tls.Config{ - Certificates: []tls.Certificate{cert}, - CipherSuites: tlsCipherSuites, - CurvePreferences: tlsCurvePreferences, - MaxVersion: tlsMaxVersion, - } - - ts.StartTLS() - - return ts, nil -} - -//nolint:revive - //nolint:revive func TestMain(m *testing.M) { fmt.Printf("Check test data dir: %s\n", testdataDir) @@ -288,8 +107,8 @@ func TestMain(m *testing.M) { fmt.Printf("caCertKeyFile created at %s\n", caCertKeyFile) - caCertTpl := demoCertTemplate{cn: "Demo CA", isCA: true, key: caCertKey} - caCertPEM, caCertParent, _ = GenerateDemoCert(caCertTpl) + caCertTpl := tlstest.Template{CN: "Demo CA", IsCA: true, Key: caCertKey} + caCertPEM, caCertParent, _ = tlstest.GenerateCert(caCertTpl) caCertPEMString = string(caCertPEM) caCertPool = x509.NewCertPool() caCertPool.AppendCertsFromPEM(caCertPEM) @@ -326,17 +145,16 @@ func TestMain(m *testing.M) { fmt.Printf("exampleCertKey file create at %s\n", exampleCertKeyFile) - exampleCertTpl := demoCertTemplate{ - cn: "example.com", - isCA: false, - dnsNames: []string{"example.com", "example.net", "example.de"}, - key: exampleCertKey, - caKey: caCertKey, - ipAddresses: []net.IP{net.IPv4(127, 0, 0, 1), net.ParseIP("::1")}, - parent: caCertParent, + exampleCertTpl := tlstest.Template{ + CN: "example.com", + DNSNames: []string{"example.com", "example.net", "example.de"}, + Key: exampleCertKey, + CAKey: caCertKey, + IPAddresses: []net.IP{net.IPv4(127, 0, 0, 1), net.ParseIP("::1")}, + Parent: caCertParent, } - exampleCertPEM, _, err = GenerateDemoCert(exampleCertTpl) + exampleCertPEM, _, err = tlstest.GenerateCert(exampleCertTpl) if err != nil { fmt.Printf("error while creating exampleCert: %s\n", err) } @@ -394,9 +212,11 @@ func TestHTTPSTestServer(t *testing.T) { t.Run(testname, func(t *testing.T) { t.Parallel() - httpSrvData := demoHttpServerData{listenHost: tt.listenHost} - - ts, err := NewHTTPSTestServer(httpSrvData) + ts, err := tlstest.NewServer(tlstest.ServerConfig{ + ListenHost: tt.listenHost, + ServerCertFile: exampleCertFile, + ServerKeyFile: exampleCertKeyFile, + }) if err != nil { t.Fatal(err) } diff --git a/internal/requests/requests_handlers_test.go b/internal/requests/requests_handlers_test.go index 5e95e70..475ac96 100644 --- a/internal/requests/requests_handlers_test.go +++ b/internal/requests/requests_handlers_test.go @@ -12,6 +12,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "github.com/xenos76/https-wrench/internal/tlstest" ) func TestResponseHeader_Print(t *testing.T) { @@ -432,20 +433,22 @@ func TestRenderTLSData(t *testing.T) { t.Run(tt.reqConf.Name, func(t *testing.T) { t.Parallel() - httpSrvData := demoHttpServerData{ - tlsCipherSuites: []uint16{tt.srvTLSCipherSuite}, - tlsCurvePreferences: tt.tlsCurvePreferences, - tlsMaxVersion: tt.srvTLSMaxVersion, - proxyprotoEnabled: false, - serverName: "localhost", + cfg := tlstest.ServerConfig{ + TLSCurvePreferences: tt.tlsCurvePreferences, + TLSMaxVersion: tt.srvTLSMaxVersion, + ServerCertFile: exampleCertFile, + ServerKeyFile: exampleCertKeyFile, + } + if tt.srvTLSMaxVersion <= tls.VersionTLS12 { + cfg.TLSCipherSuites = []uint16{tt.srvTLSCipherSuite} } - ts, err := NewHTTPSTestServer(httpSrvData) + ts, err := tlstest.NewServer(cfg) require.NoError(t, err) t.Cleanup(ts.Close) - tt.reqConf.TransportOverrideURL = "https://" + testServerHostPort(ts) + tt.reqConf.TransportOverrideURL = "https://" + ts.Listener.Addr().String() respList, err := processHTTPRequestsByHost( context.Background(), @@ -559,16 +562,15 @@ type handleRequestsTestCase struct { func runHandleRequestsSubtest(t *testing.T, tt handleRequestsTestCase) { t.Parallel() - httpSrvData := demoHttpServerData{ - serverName: "localhost", - } - - ts, err := NewHTTPSTestServer(httpSrvData) + ts, err := tlstest.NewServer(tlstest.ServerConfig{ + ServerCertFile: exampleCertFile, + ServerKeyFile: exampleCertKeyFile, + }) require.NoError(t, err) t.Cleanup(ts.Close) - tt.reqMeta.Requests[0].TransportOverrideURL = testServerHostPort(ts) + tt.reqMeta.Requests[0].TransportOverrideURL = ts.Listener.Addr().String() buffer := bytes.Buffer{} respMap, err := HandleRequests(context.Background(), &buffer, &tt.reqMeta) diff --git a/internal/requests/requests_test.go b/internal/requests/requests_test.go index 3d2e1d4..9c99518 100644 --- a/internal/requests/requests_test.go +++ b/internal/requests/requests_test.go @@ -19,6 +19,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "github.com/xenos76/https-wrench/internal/certinfo" + "github.com/xenos76/https-wrench/internal/tlstest" ) func TestNewRequestsMetaConfig(t *testing.T) { @@ -1387,12 +1388,15 @@ func runSetTransportOverrideSubtest(t *testing.T, tt setTransportOverrideTestCas c := NewRequestHTTPClient() - ts, err := NewHTTPSTestServer(demoHttpServerData{}) + ts, err := tlstest.NewServer(tlstest.ServerConfig{ + ServerCertFile: exampleCertFile, + ServerKeyFile: exampleCertKeyFile, + }) require.NoError(t, err) t.Cleanup(ts.Close) - hostPort := testServerHostPort(ts) + hostPort := ts.Listener.Addr().String() transportURL := "https://" + hostPort _, err = c.SetTransportOverride(transportURL) @@ -1444,17 +1448,17 @@ type setProxyProtocolV2TestCase struct { func runSetProxyProtocolV2Subtest(t *testing.T, tt setProxyProtocolV2TestCase) { t.Parallel() - httpSrvData := demoHttpServerData{ - listenHost: tt.listenHost, - proxyprotoEnabled: true, - } - - ts, err := NewHTTPSTestServer(httpSrvData) + ts, err := tlstest.NewServer(tlstest.ServerConfig{ + ListenHost: tt.listenHost, + ProxyprotoEnabled: true, + ServerCertFile: exampleCertFile, + ServerKeyFile: exampleCertKeyFile, + }) require.NoError(t, err) t.Cleanup(ts.Close) - hostPort := testServerHostPort(ts) + hostPort := ts.Listener.Addr().String() transportURL := "https://" + hostPort reqURL := "https://" + tt.serverName @@ -1510,12 +1514,10 @@ type printResponseDebugTestCase struct { func runPrintResponseDebugSubtest(t *testing.T, tt printResponseDebugTestCase) { t.Parallel() - httpSrvData := demoHttpServerData{ - proxyprotoEnabled: false, - serverName: "localhost", - } - - ts, err := NewHTTPSTestServer(httpSrvData) + ts, err := tlstest.NewServer(tlstest.ServerConfig{ + ServerCertFile: exampleCertFile, + ServerKeyFile: exampleCertKeyFile, + }) require.NoError(t, err) t.Cleanup(ts.Close) @@ -1559,17 +1561,15 @@ type processHTTPRequestsByHostTestCase struct { func runProcessHTTPRequestsByHostSubtest(t *testing.T, tt processHTTPRequestsByHostTestCase) { t.Parallel() - httpSrvData := demoHttpServerData{ - proxyprotoEnabled: false, - serverName: "localhost", - } - - ts, err := NewHTTPSTestServer(httpSrvData) + ts, err := tlstest.NewServer(tlstest.ServerConfig{ + ServerCertFile: exampleCertFile, + ServerKeyFile: exampleCertKeyFile, + }) require.NoError(t, err) t.Cleanup(ts.Close) - hostPort := testServerHostPort(ts) + hostPort := ts.Listener.Addr().String() tt.reqConf.TransportOverrideURL = "https://" + hostPort respList, err := processHTTPRequestsByHost( diff --git a/internal/tlstest/cert.go b/internal/tlstest/cert.go new file mode 100644 index 0000000..3e56706 --- /dev/null +++ b/internal/tlstest/cert.go @@ -0,0 +1,133 @@ +// Package tlstest provides TLS certificates and HTTPS test servers for tests. +package tlstest + +import ( + "crypto/rand" + "crypto/rsa" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "errors" + "fmt" + "math/big" + "net" + "time" +) + +const certValidity = 24 * time.Hour + +// Template holds the fields needed to issue a test CA or leaf certificate. +type Template struct { + // CN is the certificate subject Common Name. + CN string + // IsCA issues a self-signed CA when true; otherwise a leaf signed by Parent/CAKey. + IsCA bool + // DNSNames are SAN DNS names for a leaf certificate. + DNSNames []string + // IPAddresses are SAN IP addresses for a leaf certificate. + IPAddresses []net.IP + // Key is the RSA private key whose public key is embedded in the certificate. + Key *rsa.PrivateKey + // CAKey signs a leaf certificate. Ignored when IsCA is true. + CAKey *rsa.PrivateKey + // Parent is the issuing CA certificate for a leaf. Ignored when IsCA is true. + Parent *x509.Certificate +} + +// GenerateCert creates a PEM-encoded x509 certificate from tpl. +// The returned *x509.Certificate can be used as Parent when issuing a leaf. +func GenerateCert(tpl Template) ([]byte, *x509.Certificate, error) { + if err := validateTemplate(tpl); err != nil { + return nil, nil, err + } + + serialNumberLimit := new(big.Int).Lsh(big.NewInt(1), 128) + + serialNumber, err := rand.Int(rand.Reader, serialNumberLimit) + if err != nil { + return nil, nil, fmt.Errorf("failed to generate serial number: %w", err) + } + + notBefore := time.Now() + notAfter := notBefore.Add(certValidity) + + template := x509.Certificate{ + SerialNumber: serialNumber, + Subject: pkix.Name{ + CommonName: tpl.CN, + }, + NotBefore: notBefore, + NotAfter: notAfter, + IsCA: false, + DNSNames: tpl.DNSNames, + IPAddresses: tpl.IPAddresses, + ExtKeyUsage: []x509.ExtKeyUsage{ + x509.ExtKeyUsageClientAuth, + x509.ExtKeyUsageServerAuth, + }, + KeyUsage: x509.KeyUsageDigitalSignature, + } + + certParent := tpl.Parent + signingKey := tpl.CAKey + + if tpl.IsCA { + signingKey = tpl.Key + template = x509.Certificate{ + SerialNumber: serialNumber, + Subject: pkix.Name{ + CommonName: tpl.CN, + }, + NotBefore: notBefore, + NotAfter: notAfter, + IsCA: true, + ExtKeyUsage: []x509.ExtKeyUsage{ + x509.ExtKeyUsageClientAuth, + x509.ExtKeyUsageServerAuth, + }, + KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageCertSign, + BasicConstraintsValid: true, + } + certParent = &template + } + + derBytes, err := x509.CreateCertificate( + rand.Reader, + &template, + certParent, + &tpl.Key.PublicKey, + signingKey, + ) + if err != nil { + return nil, nil, fmt.Errorf("failed to create certificate: %w", err) + } + + certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: derBytes}) + + certificate, err := x509.ParseCertificate(derBytes) + if err != nil { + return nil, nil, err + } + + return certPEM, certificate, nil +} + +func validateTemplate(tpl Template) error { + if tpl.Key == nil { + return errors.New("missing certificate private key") + } + + if tpl.IsCA { + return nil + } + + if tpl.Parent == nil { + return errors.New("missing parent certificate") + } + + if tpl.CAKey == nil { + return errors.New("missing CA private key") + } + + return nil +} diff --git a/internal/tlstest/cert_test.go b/internal/tlstest/cert_test.go new file mode 100644 index 0000000..9ae3132 --- /dev/null +++ b/internal/tlstest/cert_test.go @@ -0,0 +1,106 @@ +package tlstest + +import ( + "crypto/rand" + "crypto/rsa" + "encoding/pem" + "net" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestGenerateCert(t *testing.T) { + t.Parallel() + + caKey, err := rsa.GenerateKey(rand.Reader, 2048) + require.NoError(t, err) + + leafKey, err := rsa.GenerateKey(rand.Reader, 2048) + require.NoError(t, err) + + caPEM, caCert, err := GenerateCert(Template{ + CN: "tlstest CA", + IsCA: true, + Key: caKey, + }) + require.NoError(t, err) + require.True(t, caCert.IsCA) + require.True(t, caCert.BasicConstraintsValid) + require.Equal(t, "tlstest CA", caCert.Subject.CommonName) + require.NotEmpty(t, caPEM) + + block, _ := pem.Decode(caPEM) + require.NotNil(t, block) + require.Equal(t, "CERTIFICATE", block.Type) + + wantDNS := []string{"example.com", "example.net"} + wantIPs := []net.IP{net.IPv4(127, 0, 0, 1), net.ParseIP("::1")} + + leafPEM, leafCert, err := GenerateCert(Template{ + CN: "example.com", + Key: leafKey, + CAKey: caKey, + Parent: caCert, + DNSNames: wantDNS, + IPAddresses: wantIPs, + }) + require.NoError(t, err) + require.False(t, leafCert.IsCA) + require.Equal(t, "example.com", leafCert.Subject.CommonName) + require.Equal(t, wantDNS, leafCert.DNSNames) + require.Len(t, leafCert.IPAddresses, 2) + require.True(t, leafCert.IPAddresses[0].Equal(wantIPs[0])) + require.True(t, leafCert.IPAddresses[1].Equal(wantIPs[1])) + require.NotEmpty(t, leafPEM) + require.NoError(t, leafCert.CheckSignatureFrom(caCert)) +} + +func TestGenerateCert_missingKey(t *testing.T) { + t.Parallel() + + _, _, err := GenerateCert(Template{ + CN: "broken", + IsCA: true, + }) + require.Error(t, err) +} + +func TestGenerateCert_missingIssuer(t *testing.T) { + t.Parallel() + + leafKey, err := rsa.GenerateKey(rand.Reader, 2048) + require.NoError(t, err) + + caKey, err := rsa.GenerateKey(rand.Reader, 2048) + require.NoError(t, err) + + _, caCert, err := GenerateCert(Template{ + CN: "tlstest CA", + IsCA: true, + Key: caKey, + }) + require.NoError(t, err) + + t.Run("missing parent", func(t *testing.T) { + t.Parallel() + + _, _, err := GenerateCert(Template{ + CN: "example.com", + Key: leafKey, + CAKey: caKey, + }) + require.ErrorContains(t, err, "parent certificate") + }) + + t.Run("missing CA key", func(t *testing.T) { + t.Parallel() + + _, _, err := GenerateCert(Template{ + CN: "example.com", + Key: leafKey, + Parent: caCert, + }) + require.ErrorContains(t, err, "CA private key") + }) +} diff --git a/internal/tlstest/server.go b/internal/tlstest/server.go new file mode 100644 index 0000000..2f51dcc --- /dev/null +++ b/internal/tlstest/server.go @@ -0,0 +1,143 @@ +package tlstest + +import ( + "crypto/tls" + "fmt" + "net" + "net/http" + "net/http/httptest" + "slices" + "time" + + "github.com/pires/go-proxyproto" +) + +const proxyProtoReadHeaderTimeout = 10 * time.Second + +// defaultCurvePreferences matches the Go 1.27 TLS hybrids plus classical fallbacks +// used by certinfo and requests production TLS configs. +var defaultCurvePreferences = []tls.CurveID{ + tls.X25519MLKEM768, + tls.SecP256r1MLKEM768, + tls.SecP384r1MLKEM1024, + tls.X25519, + tls.CurveP256, + tls.CurveP384, + tls.CurveP521, +} + +// ServerConfig configures NewServer. +type ServerConfig struct { + // ListenHost binds the listener to this host (ephemeral port). Empty uses httptest's default. + ListenHost string + // ProxyprotoEnabled wraps the listener with a PROXY protocol v2 reader. + ProxyprotoEnabled bool + // TLSCipherSuites, when non-empty, must be TLS 1.0–1.2 cipher suite IDs. + // tls.Config.CipherSuites stays nil by default. Callers that pin suites + // should also set TLSMaxVersion to tls.VersionTLS12. + TLSCipherSuites []uint16 + // TLSCurvePreferences override the Go 1.27 hybrid defaults when non-empty. + TLSCurvePreferences []tls.CurveID + // TLSMaxVersion overrides TLS 1.3 when non-zero. + TLSMaxVersion uint16 + // ServerCertFile is the PEM certificate path loaded for the TLS listener. + ServerCertFile string + // ServerKeyFile is the PEM private key path loaded for the TLS listener. + ServerKeyFile string +} + +// NewServer starts an httptest TLS server configured by cfg. +// CipherSuites stays nil by default (Go's TLS 1.3 AEADs plus default TLS 1.2 +// suites). CurvePreferences default to the Go 1.27 PQ hybrids plus classical +// fallbacks, and MaxVersion to TLS 1.3. Non-empty cfg.TLSCipherSuites must be +// TLS 1.0–1.2 suite IDs; cfg.TLSCurvePreferences and a non-zero +// cfg.TLSMaxVersion override those defaults. Optional cfg.ListenHost and +// cfg.ProxyprotoEnabled replace the listener. The caller must Close the +// returned server. +func NewServer(cfg ServerConfig) (*httptest.Server, error) { + handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + fmt.Fprint(w, "DemoHTTPSServer Handler - client output\n") + fmt.Fprint(w, "Host requested: ", r.Host, "\n") + + fmt.Println("DemoHTTPSServer Handler - shell output") + }) + + ts := httptest.NewUnstartedServer(handler) + ts.EnableHTTP2 = true + + if cfg.ListenHost != "" { + ln, err := net.Listen("tcp", net.JoinHostPort(cfg.ListenHost, "0")) + if err != nil { + return nil, fmt.Errorf("error creating listener: %w", err) + } + + _ = ts.Listener.Close() + ts.Listener = ln + } + + if cfg.ProxyprotoEnabled { + ts.Listener = &proxyproto.Listener{ + Listener: ts.Listener, + ReadHeaderTimeout: proxyProtoReadHeaderTimeout, + } + } + + cert, err := tls.LoadX509KeyPair(cfg.ServerCertFile, cfg.ServerKeyFile) + if err != nil { + _ = ts.Listener.Close() + return nil, err + } + + var tlsCipherSuites []uint16 + + if len(cfg.TLSCipherSuites) > 0 { + if err := validateTLS12CipherSuites(cfg.TLSCipherSuites); err != nil { + _ = ts.Listener.Close() + return nil, err + } + + tlsCipherSuites = cfg.TLSCipherSuites + } + + tlsCurvePreferences := defaultCurvePreferences + if len(cfg.TLSCurvePreferences) > 0 { + tlsCurvePreferences = cfg.TLSCurvePreferences + } + + tlsMaxVersion := uint16(tls.VersionTLS13) + if cfg.TLSMaxVersion > 0 { + tlsMaxVersion = cfg.TLSMaxVersion + } + + ts.TLS = &tls.Config{ + Certificates: []tls.Certificate{cert}, + CipherSuites: tlsCipherSuites, + CurvePreferences: tlsCurvePreferences, + MaxVersion: tlsMaxVersion, + } + + ts.StartTLS() + + return ts, nil +} + +func validateTLS12CipherSuites(ids []uint16) error { + allowed := make(map[uint16]struct{}) + + for _, suite := range slices.Concat(tls.CipherSuites(), tls.InsecureCipherSuites()) { + for _, version := range suite.SupportedVersions { + if version <= tls.VersionTLS12 { + allowed[suite.ID] = struct{}{} + break + } + } + } + + for _, id := range ids { + if _, ok := allowed[id]; !ok { + return fmt.Errorf("cipher suite 0x%04x is not a TLS 1.0-1.2 suite", id) + } + } + + return nil +} diff --git a/internal/tlstest/server_test.go b/internal/tlstest/server_test.go new file mode 100644 index 0000000..97ef623 --- /dev/null +++ b/internal/tlstest/server_test.go @@ -0,0 +1,161 @@ +package tlstest + +import ( + "crypto/rand" + "crypto/rsa" + "crypto/tls" + "crypto/x509" + "encoding/pem" + "net" + "net/http" + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/require" +) + +type leafFiles struct { + certFile string + keyFile string + caCert *x509.Certificate +} + +func writeLeafFiles(t *testing.T) leafFiles { + t.Helper() + + caKey, err := rsa.GenerateKey(rand.Reader, 2048) + require.NoError(t, err) + + leafKey, err := rsa.GenerateKey(rand.Reader, 2048) + require.NoError(t, err) + + _, caCert, err := GenerateCert(Template{ + CN: "tlstest CA", + IsCA: true, + Key: caKey, + }) + require.NoError(t, err) + + leafPEM, _, err := GenerateCert(Template{ + CN: "localhost", + Key: leafKey, + CAKey: caKey, + Parent: caCert, + DNSNames: []string{"localhost"}, + IPAddresses: []net.IP{net.IPv4(127, 0, 0, 1), net.ParseIP("::1")}, + }) + require.NoError(t, err) + + dir := t.TempDir() + certFile := filepath.Join(dir, "leaf.pem") + keyFile := filepath.Join(dir, "leaf.key") + + require.NoError(t, os.WriteFile(certFile, leafPEM, 0o600)) + + keyPEM := pem.EncodeToMemory(&pem.Block{ + Type: "RSA PRIVATE KEY", + Bytes: x509.MarshalPKCS1PrivateKey(leafKey), + }) + require.NoError(t, os.WriteFile(keyFile, keyPEM, 0o600)) + + return leafFiles{ + certFile: certFile, + keyFile: keyFile, + caCert: caCert, + } +} + +func TestNewServer(t *testing.T) { + t.Parallel() + + leaf := writeLeafFiles(t) + + ts, err := NewServer(ServerConfig{ + ListenHost: "127.0.0.1", + ServerCertFile: leaf.certFile, + ServerKeyFile: leaf.keyFile, + }) + require.NoError(t, err) + t.Cleanup(ts.Close) + require.Nil(t, ts.TLS.CipherSuites) + + pool := x509.NewCertPool() + pool.AddCert(leaf.caCert) + + client := &http.Client{ + Transport: &http.Transport{ + TLSClientConfig: &tls.Config{ + RootCAs: pool, + }, + }, + } + + res, err := client.Get(ts.URL) + require.NoError(t, err) + t.Cleanup(func() { _ = res.Body.Close() }) + require.Equal(t, http.StatusOK, res.StatusCode) +} + +func TestNewServerTLS12ConfiguredCipher(t *testing.T) { + t.Parallel() + + leaf := writeLeafFiles(t) + suite := tls.TLS_ECDHE_RSA_WITH_CHACHA20_POLY1305_SHA256 + + ts, err := NewServer(ServerConfig{ + ListenHost: "127.0.0.1", + ServerCertFile: leaf.certFile, + ServerKeyFile: leaf.keyFile, + TLSMaxVersion: tls.VersionTLS12, + TLSCipherSuites: []uint16{suite}, + }) + require.NoError(t, err) + t.Cleanup(ts.Close) + + pool := x509.NewCertPool() + pool.AddCert(leaf.caCert) + + client := &http.Client{ + Transport: &http.Transport{ + TLSClientConfig: &tls.Config{ + RootCAs: pool, + MaxVersion: tls.VersionTLS12, + CipherSuites: []uint16{suite}, + }, + }, + } + + res, err := client.Get(ts.URL) + require.NoError(t, err) + t.Cleanup(func() { _ = res.Body.Close() }) + require.Equal(t, http.StatusOK, res.StatusCode) + require.NotNil(t, res.TLS) + require.Equal(t, uint16(tls.VersionTLS12), res.TLS.Version) + require.Equal(t, suite, res.TLS.CipherSuite) +} + +func TestNewServerRejectsTLS13CipherSuite(t *testing.T) { + t.Parallel() + + leaf := writeLeafFiles(t) + + _, err := NewServer(ServerConfig{ + ListenHost: "127.0.0.1", + ServerCertFile: leaf.certFile, + ServerKeyFile: leaf.keyFile, + TLSCipherSuites: []uint16{tls.TLS_AES_128_GCM_SHA256}, + }) + require.Error(t, err) +} + +func TestNewServerMissingCertClosesListener(t *testing.T) { + t.Parallel() + + _, err := NewServer(ServerConfig{ + ListenHost: "127.0.0.1", + ServerCertFile: filepath.Join(t.TempDir(), "missing.pem"), + ServerKeyFile: filepath.Join(t.TempDir(), "missing.key"), + }) + require.Error(t, err) +}