srdusr
aboutsummaryrefslogtreecommitdiffstats
path: root/internal/clientcert/clientcert_test.go
blob: 712562b997eba51b53ef4285f4f1257fab0ffc83 (plain) (blame)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
package clientcert

import (
	"bytes"
	"crypto/ecdsa"
	"crypto/elliptic"
	"crypto/rand"
	"crypto/x509"
	"crypto/x509/pkix"
	"encoding/pem"
	"math/big"
	"testing"
	"time"
)

// selfSignedPEM generates a throwaway self-signed cert/key pair entirely
// in memory (crypto/tls + crypto/x509 stdlib only) rather than shelling
// out to openssl or checking in a fixture - keeps the test hermetic and
// portable.
func selfSignedPEM(t *testing.T) (certPEM, keyPEM []byte) {
	t.Helper()
	key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
	if err != nil {
		t.Fatalf("generate key: %v", err)
	}
	tmpl := &x509.Certificate{
		SerialNumber: big.NewInt(1),
		Subject:      pkix.Name{CommonName: "test"},
		NotBefore:    time.Now(),
		NotAfter:     time.Now().Add(time.Hour),
	}
	der, err := x509.CreateCertificate(rand.Reader, tmpl, tmpl, &key.PublicKey, key)
	if err != nil {
		t.Fatalf("create certificate: %v", err)
	}
	keyDER, err := x509.MarshalECPrivateKey(key)
	if err != nil {
		t.Fatalf("marshal key: %v", err)
	}
	var certBuf, keyBuf bytes.Buffer
	pem.Encode(&certBuf, &pem.Block{Type: "CERTIFICATE", Bytes: der})
	pem.Encode(&keyBuf, &pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER})
	return certBuf.Bytes(), keyBuf.Bytes()
}

func TestFindForExactMatch(t *testing.T) {
	certs := []Cert{{ID: 1, Enabled: true, Pattern: "internal.example.com"}}
	got := FindFor(certs, "internal.example.com")
	if got == nil || got.ID != 1 {
		t.Errorf("got %+v, want match on id 1", got)
	}
}

func TestFindForSubstringMatch(t *testing.T) {
	certs := []Cert{{ID: 1, Enabled: true, Pattern: "example.com"}}
	if FindFor(certs, "api.example.com") == nil {
		t.Error("expected subdomain to match by substring, same convention as scope.Rule")
	}
}

func TestFindForRegex(t *testing.T) {
	certs := []Cert{{ID: 1, Enabled: true, Pattern: `^api-\d+\.example\.com$`, IsRegex: true}}
	if FindFor(certs, "api-42.example.com") == nil {
		t.Error("expected regex pattern to match")
	}
	if FindFor(certs, "api-x.example.com") != nil {
		t.Error("expected regex pattern not to match a non-numeric suffix")
	}
}

func TestFindForSkipsDisabled(t *testing.T) {
	certs := []Cert{{ID: 1, Enabled: false, Pattern: "example.com"}}
	if FindFor(certs, "example.com") != nil {
		t.Error("expected a disabled cert not to match")
	}
}

func TestFindForNoMatch(t *testing.T) {
	certs := []Cert{{ID: 1, Enabled: true, Pattern: "example.com"}}
	if FindFor(certs, "other.org") != nil {
		t.Error("expected no match for an unrelated host")
	}
}

func TestFindForFirstMatchWins(t *testing.T) {
	certs := []Cert{
		{ID: 1, Enabled: true, Pattern: "example.com"},
		{ID: 2, Enabled: true, Pattern: "example.com"},
	}
	got := FindFor(certs, "example.com")
	if got == nil || got.ID != 1 {
		t.Errorf("got %+v, want the first matching cert (id 1)", got)
	}
}

func TestTLSCertificateValidPair(t *testing.T) {
	certPEM, keyPEM := selfSignedPEM(t)
	c := Cert{Name: "test", CertPEM: certPEM, KeyPEM: keyPEM}
	if _, err := c.TLSCertificate(); err != nil {
		t.Errorf("unexpected error: %v", err)
	}
}

func TestTLSCertificateMismatchedPairErrors(t *testing.T) {
	certPEM, _ := selfSignedPEM(t)
	_, otherKeyPEM := selfSignedPEM(t)
	c := Cert{Name: "test", CertPEM: certPEM, KeyPEM: otherKeyPEM}
	if _, err := c.TLSCertificate(); err == nil {
		t.Error("expected an error pairing a cert with a key that doesn't match it")
	}
}