cloudflare/cloudflared
Publicmirrored from https://github.com/cloudflare/cloudflaredAvailable
cmd/cloudflared/tunnel/configuration_test.go
236lines · modecode
| 1 | //go:build ignore |
| 2 | |
| 3 | // TODO: Remove the above build tag and include this test when we start compiling with Golang 1.10.0+ |
| 4 | |
| 5 | package tunnel |
| 6 | |
| 7 | import ( |
| 8 | "crypto/x509" |
| 9 | "crypto/x509/pkix" |
| 10 | "encoding/asn1" |
| 11 | "net" |
| 12 | "os" |
| 13 | "testing" |
| 14 | |
| 15 | "github.com/stretchr/testify/assert" |
| 16 | ) |
| 17 | |
| 18 | // Generated using `openssl req -newkey rsa:512 -nodes -x509 -days 3650` |
| 19 | var samplePEM = []byte(` |
| 20 | -----BEGIN CERTIFICATE----- |
| 21 | MIIB4DCCAYoCCQCb/H0EUrdXEjANBgkqhkiG9w0BAQsFADB3MQswCQYDVQQGEwJV |
| 22 | UzEOMAwGA1UECAwFVGV4YXMxDzANBgNVBAcMBkF1c3RpbjEZMBcGA1UECgwQQ2xv |
| 23 | dWRmbGFyZSwgSW5jLjEZMBcGA1UECwwQUHJvZHVjdCBTdHJhdGVneTERMA8GA1UE |
| 24 | AwwIVGVzdCBPbmUwHhcNMTgwNDI2MTYxMDUxWhcNMjgwNDIzMTYxMDUxWjB3MQsw |
| 25 | CQYDVQQGEwJVUzEOMAwGA1UECAwFVGV4YXMxDzANBgNVBAcMBkF1c3RpbjEZMBcG |
| 26 | A1UECgwQQ2xvdWRmbGFyZSwgSW5jLjEZMBcGA1UECwwQUHJvZHVjdCBTdHJhdGVn |
| 27 | eTERMA8GA1UEAwwIVGVzdCBPbmUwXDANBgkqhkiG9w0BAQEFAANLADBIAkEAwVQD |
| 28 | K0SJ25UFLznm2pU3zhzMEvpDEofHVNnCjk4mlDrtVop7PkKZ8pDEmuQANltUrxC8 |
| 29 | yHBE2wXMv+GlH+bDtwIDAQABMA0GCSqGSIb3DQEBCwUAA0EAjVYQzozIFPkt/HRY |
| 30 | uUoZ8zEHIDICb0syFf5VAjm9AgTwIPzUmD+c5vl6LWDnxq7L45nLCzhhQ6YmiwDz |
| 31 | X7Wcyg== |
| 32 | -----END CERTIFICATE----- |
| 33 | -----BEGIN CERTIFICATE----- |
| 34 | MIIB4DCCAYoCCQDZfCdAJ+mwzDANBgkqhkiG9w0BAQsFADB3MQswCQYDVQQGEwJV |
| 35 | UzEOMAwGA1UECAwFVGV4YXMxDzANBgNVBAcMBkF1c3RpbjEZMBcGA1UECgwQQ2xv |
| 36 | dWRmbGFyZSwgSW5jLjEZMBcGA1UECwwQUHJvZHVjdCBTdHJhdGVneTERMA8GA1UE |
| 37 | AwwIVGVzdCBUd28wHhcNMTgwNDI2MTYxMTIwWhcNMjgwNDIzMTYxMTIwWjB3MQsw |
| 38 | CQYDVQQGEwJVUzEOMAwGA1UECAwFVGV4YXMxDzANBgNVBAcMBkF1c3RpbjEZMBcG |
| 39 | A1UECgwQQ2xvdWRmbGFyZSwgSW5jLjEZMBcGA1UECwwQUHJvZHVjdCBTdHJhdGVn |
| 40 | eTERMA8GA1UEAwwIVGVzdCBUd28wXDANBgkqhkiG9w0BAQEFAANLADBIAkEAoHKp |
| 41 | ROVK3zCSsH7ocYeyRAML4V7SFAbZcb4WIwDnE08oMBVRkQVcW5tqEkvG3RiClfzV |
| 42 | wZIJ3CfqKIeSNSDU9wIDAQABMA0GCSqGSIb3DQEBCwUAA0EAJw2gUbnPiq4C2p5b |
| 43 | iWzlA9Q7aKo+VQ4H7IZS7tTccr59nVjvH/TG3eWujpnocr4TOqW9M3CK1DF9mUGP |
| 44 | 3pQ3Jg== |
| 45 | -----END CERTIFICATE----- |
| 46 | `) |
| 47 | |
| 48 | var systemCertPoolSubjects []*pkix.Name |
| 49 | |
| 50 | type certificateFixture struct { |
| 51 | ou string |
| 52 | cn string |
| 53 | } |
| 54 | |
| 55 | func TestMain(m *testing.M) { |
| 56 | systemCertPool, err := x509.SystemCertPool() |
| 57 | if isUnrecoverableError(err) { |
| 58 | os.Exit(1) |
| 59 | } |
| 60 | |
| 61 | if systemCertPool == nil { |
| 62 | // On Windows, let's just assume the system cert pool was empty |
| 63 | systemCertPool = x509.NewCertPool() |
| 64 | } |
| 65 | |
| 66 | systemCertPoolSubjects, err = getCertPoolSubjects(systemCertPool) |
| 67 | if err != nil { |
| 68 | os.Exit(1) |
| 69 | } |
| 70 | |
| 71 | os.Exit(m.Run()) |
| 72 | } |
| 73 | |
| 74 | func TestLoadOriginCertPoolJustSystemPool(t *testing.T) { |
| 75 | certPoolSubjects := loadCertPoolSubjects(t, nil) |
| 76 | extraSubjects := subjectSubtract(systemCertPoolSubjects, certPoolSubjects) |
| 77 | |
| 78 | // Remove extra subjects from the cert pool |
| 79 | var filteredSystemCertPoolSubjects []*pkix.Name |
| 80 | |
| 81 | t.Log(extraSubjects) |
| 82 | |
| 83 | OUTER: |
| 84 | for _, subject := range certPoolSubjects { |
| 85 | for _, extraSubject := range extraSubjects { |
| 86 | if subject == extraSubject { |
| 87 | t.Log(extraSubject) |
| 88 | continue OUTER |
| 89 | } |
| 90 | } |
| 91 | |
| 92 | filteredSystemCertPoolSubjects = append(filteredSystemCertPoolSubjects, subject) |
| 93 | } |
| 94 | |
| 95 | assert.Equal(t, len(filteredSystemCertPoolSubjects), len(systemCertPoolSubjects)) |
| 96 | |
| 97 | difference := subjectSubtract(systemCertPoolSubjects, filteredSystemCertPoolSubjects) |
| 98 | assert.Equal(t, 0, len(difference)) |
| 99 | } |
| 100 | |
| 101 | func TestLoadOriginCertPoolCFCertificates(t *testing.T) { |
| 102 | certPoolSubjects := loadCertPoolSubjects(t, nil) |
| 103 | |
| 104 | extraSubjects := subjectSubtract(systemCertPoolSubjects, certPoolSubjects) |
| 105 | |
| 106 | expected := []*certificateFixture{ |
| 107 | {ou: "CloudFlare Origin SSL ECC Certificate Authority"}, |
| 108 | {ou: "CloudFlare Origin SSL Certificate Authority"}, |
| 109 | {cn: "origin-pull.cloudflare.net"}, |
| 110 | {cn: "Argo Tunnel Sample Hello Server Certificate"}, |
| 111 | } |
| 112 | |
| 113 | assertFixturesMatchSubjects(t, expected, extraSubjects) |
| 114 | } |
| 115 | |
| 116 | func TestLoadOriginCertPoolWithExtraPEMs(t *testing.T) { |
| 117 | certPoolWithoutPEMSubjects := loadCertPoolSubjects(t, nil) |
| 118 | certPoolWithPEMSubjects := loadCertPoolSubjects(t, samplePEM) |
| 119 | |
| 120 | difference := subjectSubtract(certPoolWithoutPEMSubjects, certPoolWithPEMSubjects) |
| 121 | |
| 122 | assert.Equal(t, 2, len(difference)) |
| 123 | |
| 124 | expected := []*certificateFixture{ |
| 125 | {cn: "Test One"}, |
| 126 | {cn: "Test Two"}, |
| 127 | } |
| 128 | |
| 129 | assertFixturesMatchSubjects(t, expected, difference) |
| 130 | } |
| 131 | |
| 132 | func loadCertPoolSubjects(t *testing.T, originCAPoolPEM []byte) []*pkix.Name { |
| 133 | certPool, err := loadOriginCertPool(originCAPoolPEM) |
| 134 | if isUnrecoverableError(err) { |
| 135 | t.Fatal(err) |
| 136 | } |
| 137 | assert.NotEmpty(t, certPool.Subjects()) |
| 138 | certPoolSubjects, err := getCertPoolSubjects(certPool) |
| 139 | if err != nil { |
| 140 | t.Fatal(err) |
| 141 | } |
| 142 | |
| 143 | return certPoolSubjects |
| 144 | } |
| 145 | |
| 146 | func assertFixturesMatchSubjects(t *testing.T, fixtures []*certificateFixture, subjects []*pkix.Name) { |
| 147 | assert.Equal(t, len(fixtures), len(subjects)) |
| 148 | |
| 149 | for _, fixture := range fixtures { |
| 150 | found := false |
| 151 | for _, subject := range subjects { |
| 152 | found = found || fixtureMatchesSubjectPredicate(fixture, subject) |
| 153 | } |
| 154 | |
| 155 | if !found { |
| 156 | t.Fail() |
| 157 | } |
| 158 | } |
| 159 | } |
| 160 | |
| 161 | func fixtureMatchesSubjectPredicate(fixture *certificateFixture, subject *pkix.Name) bool { |
| 162 | cnMatch := true |
| 163 | if fixture.cn != "" { |
| 164 | cnMatch = fixture.cn == subject.CommonName |
| 165 | } |
| 166 | |
| 167 | ouMatch := true |
| 168 | if fixture.ou != "" { |
| 169 | ouMatch = len(subject.OrganizationalUnit) > 0 && fixture.ou == subject.OrganizationalUnit[0] |
| 170 | } |
| 171 | |
| 172 | return cnMatch && ouMatch |
| 173 | } |
| 174 | |
| 175 | func subjectSubtract(left []*pkix.Name, right []*pkix.Name) []*pkix.Name { |
| 176 | var difference []*pkix.Name |
| 177 | |
| 178 | var found bool |
| 179 | for _, r := range right { |
| 180 | found = false |
| 181 | for _, l := range left { |
| 182 | if (*l).String() == (*r).String() { |
| 183 | found = true |
| 184 | } |
| 185 | } |
| 186 | |
| 187 | if !found { |
| 188 | difference = append(difference, r) |
| 189 | } |
| 190 | } |
| 191 | |
| 192 | return difference |
| 193 | } |
| 194 | |
| 195 | func getCertPoolSubjects(certPool *x509.CertPool) ([]*pkix.Name, error) { |
| 196 | var subjects []*pkix.Name |
| 197 | |
| 198 | for _, subject := range certPool.Subjects() { |
| 199 | var sequence pkix.RDNSequence |
| 200 | _, err := asn1.Unmarshal(subject, &sequence) |
| 201 | if err != nil { |
| 202 | return nil, err |
| 203 | } |
| 204 | |
| 205 | name := pkix.Name{} |
| 206 | name.FillFromRDNSequence(&sequence) |
| 207 | |
| 208 | subjects = append(subjects, &name) |
| 209 | } |
| 210 | |
| 211 | return subjects, nil |
| 212 | } |
| 213 | |
| 214 | func isUnrecoverableError(err error) bool { |
| 215 | return err != nil && err.Error() != "crypto/x509: system root pool is not available on Windows" |
| 216 | } |
| 217 | |
| 218 | func TestTestIPBindable(t *testing.T) { |
| 219 | assert.Nil(t, testIPBindable(nil)) |
| 220 | |
| 221 | // Public services - if one of these IPs is on the machine, the test environment is too weird |
| 222 | assert.NotNil(t, testIPBindable(net.ParseIP("8.8.8.8"))) |
| 223 | assert.NotNil(t, testIPBindable(net.ParseIP("1.1.1.1"))) |
| 224 | |
| 225 | addrs, err := net.InterfaceAddrs() |
| 226 | if err != nil { |
| 227 | t.Fatal(err) |
| 228 | } |
| 229 | for i, addr := range addrs { |
| 230 | if i >= 3 { |
| 231 | break |
| 232 | } |
| 233 | ip := addr.(*net.IPNet).IP |
| 234 | assert.Nil(t, testIPBindable(ip)) |
| 235 | } |
| 236 | } |