File
Blob: tun/client/certificate_test.go
| 1 | package client |
| 2 | |
| 3 | import ( |
| 4 | "crypto/tls" |
| 5 | "crypto/x509" |
| 6 | "net" |
| 7 | "os" |
| 8 | "testing" |
| 9 | "time" |
| 10 | |
| 11 | "go.miragespace.co/specter/spec/chord" |
| 12 | "go.miragespace.co/specter/spec/mocks" |
| 13 | "go.miragespace.co/specter/spec/pki" |
| 14 | "go.miragespace.co/specter/spec/protocol" |
| 15 | |
| 16 | "github.com/stretchr/testify/mock" |
| 17 | "github.com/stretchr/testify/require" |
| 18 | "github.com/zhangyunhao116/skipmap" |
| 19 | "go.uber.org/zap" |
| 20 | "go.uber.org/zap/zaptest" |
| 21 | ) |
| 22 | |
| 23 | func TestCertificateUpdateDuringForwarding(t *testing.T) { |
| 24 | logger := zap.NewNop().With(zap.Uint64("id", 1)) |
| 25 | _, cert, key := makeCertificate(require.New(t), logger, &protocol.Node{Id: 1}, &protocol.ClientToken{Token: []byte("owner")}, nil) |
| 26 | tp := mocks.SelfTransport() |
| 27 | tp.Identify = &protocol.Node{Id: 1} |
| 28 | tp.On("WithClientCertificate", mock.Anything).Return(nil) |
| 29 | c := &Client{ |
| 30 | ClientConfig: ClientConfig{ |
| 31 | Logger: logger, |
| 32 | Configuration: &Config{ |
| 33 | Certificate: cert, |
| 34 | PrivKey: key, |
| 35 | }, |
| 36 | ServerTransport: tp, |
| 37 | }, |
| 38 | forwarder: newForwarder(logger), |
| 39 | } |
| 40 | stop := make(chan struct{}) |
| 41 | done := make(chan struct{}) |
| 42 | go func() { |
| 43 | defer close(done) |
| 44 | for { |
| 45 | select { |
| 46 | case <-stop: |
| 47 | return |
| 48 | default: |
| 49 | } |
| 50 | left, right := net.Pipe() |
| 51 | c.forwarder.handleLink(t.Context(), &protocol.Link{Alpn: protocol.Link_UNKNOWN}, left, route{}) |
| 52 | right.Close() |
| 53 | } |
| 54 | }() |
| 55 | defer func() { close(stop); <-done }() |
| 56 | for i := 0; i < 100; i++ { |
| 57 | require.NoError(t, c.updateTransportCert()) |
| 58 | } |
| 59 | require.Same(t, logger, c.Logger) |
| 60 | require.Same(t, logger, c.forwarder.logger) |
| 61 | } |
| 62 | |
| 63 | func TestRegister_PerformsRenewalWhenNearExpiry(t *testing.T) { |
| 64 | as := require.New(t) |
| 65 | logger := zaptest.NewLogger(t) |
| 66 | |
| 67 | file, err := os.CreateTemp("", "client") |
| 68 | as.NoError(err) |
| 69 | defer os.Remove(file.Name()) |
| 70 | |
| 71 | ctx := t.Context() |
| 72 | |
| 73 | token := &protocol.ClientToken{ |
| 74 | Token: []byte("test"), |
| 75 | } |
| 76 | cl := &protocol.Node{ |
| 77 | Id: chord.Random(), |
| 78 | } |
| 79 | |
| 80 | // Create a certificate that expires in 1 hour (within renewal window) |
| 81 | oldDer, oldCert, key := makeCertificateWithExpiry(as, logger, cl, token, nil, time.Hour) |
| 82 | |
| 83 | // Parse the key so we can create the renewed cert with the same key |
| 84 | privKey, err := pki.UnmarshalPrivateKey([]byte(key)) |
| 85 | as.NoError(err) |
| 86 | |
| 87 | // Create a renewed certificate (fresh, 180 days) |
| 88 | newDer, newCert, _ := makeCertificate(as, logger, cl, token, privKey) |
| 89 | |
| 90 | cfg := &Config{ |
| 91 | path: file.Name(), |
| 92 | router: skipmap.NewString[route](), |
| 93 | Apex: testApex, |
| 94 | Certificate: oldCert, |
| 95 | PrivKey: key, |
| 96 | Tunnels: []Tunnel{ |
| 97 | { |
| 98 | Target: "tcp://127.0.0.1:2345", |
| 99 | }, |
| 100 | }, |
| 101 | } |
| 102 | as.NoError(cfg.validate()) |
| 103 | |
| 104 | pkiClient := new(mocks.PKIClient) |
| 105 | pkiClient.On("RenewCertificate", mock.Anything, mock.Anything).Return(&protocol.CertificateResponse{ |
| 106 | CertDer: newDer, |
| 107 | CertPem: []byte(newCert), |
| 108 | }, nil).Once() |
| 109 | |
| 110 | m := func(s *mocks.TunnelService, t1 *mocks.MemoryTransport, publishCall *mock.Call) { |
| 111 | defaultNoHostnames(s) |
| 112 | // First call: OLD certificate (during NewClient initialization) |
| 113 | // Second call: NEW certificate (after renewal in Register) |
| 114 | t1.On("WithClientCertificate", mock.Anything).Run(func(args mock.Arguments) { |
| 115 | cert := args.Get(0).(tls.Certificate) |
| 116 | parsed, err := x509.ParseCertificate(cert.Certificate[0]) |
| 117 | if err != nil { |
| 118 | panic(err) |
| 119 | } |
| 120 | t1.WithCertificate(parsed) |
| 121 | }).Return(nil).Twice() |
| 122 | _ = oldDer // used implicitly via oldCert |
| 123 | } |
| 124 | |
| 125 | client, _, assertion := setupClient(t, as, ctx, logger, pkiClient, cfg, nil, m, true, 1) |
| 126 | defer assertion() |
| 127 | defer client.Close() |
| 128 | |
| 129 | // Verify that the certificate was renewed (config should have the new cert) |
| 130 | currCfg := client.GetCurrentConfig() |
| 131 | as.Equal(newCert, currCfg.Certificate) |
| 132 | } |
| 133 | |
| 134 | func TestShouldRenewCertificate(t *testing.T) { |
| 135 | as := require.New(t) |
| 136 | |
| 137 | // Test certificate within renewal window (expires in 1 hour) |
| 138 | nearExpiry := &x509.Certificate{ |
| 139 | NotAfter: time.Now().Add(time.Hour), |
| 140 | } |
| 141 | as.True(shouldRenewCertificate(nearExpiry)) |
| 142 | |
| 143 | // Test certificate outside renewal window (expires in 60 days) |
| 144 | farFromExpiry := &x509.Certificate{ |
| 145 | NotAfter: time.Now().Add(60 * 24 * time.Hour), |
| 146 | } |
| 147 | as.False(shouldRenewCertificate(farFromExpiry)) |
| 148 | |
| 149 | // Test certificate exactly at renewal window boundary (30 days) |
| 150 | atBoundary := &x509.Certificate{ |
| 151 | NotAfter: time.Now().Add(30 * 24 * time.Hour), |
| 152 | } |
| 153 | as.True(shouldRenewCertificate(atBoundary)) |
| 154 | |
| 155 | // Test expired certificate (should still return true - needs renewal) |
| 156 | expired := &x509.Certificate{ |
| 157 | NotAfter: time.Now().Add(-time.Hour), |
| 158 | } |
| 159 | as.True(shouldRenewCertificate(expired)) |
| 160 | } |
| 161 | |
| 162 | func TestCertificateMaintainer_RenewsInBackground(t *testing.T) { |
| 163 | as := require.New(t) |
| 164 | logger := zaptest.NewLogger(t) |
| 165 | |
| 166 | file, err := os.CreateTemp("", "client") |
| 167 | as.NoError(err) |
| 168 | defer os.Remove(file.Name()) |
| 169 | |
| 170 | ctx := t.Context() |
| 171 | |
| 172 | token := &protocol.ClientToken{ |
| 173 | Token: []byte("test"), |
| 174 | } |
| 175 | cl := &protocol.Node{ |
| 176 | Id: chord.Random(), |
| 177 | } |
| 178 | |
| 179 | // Create a certificate that expires in 1 hour (within renewal window) |
| 180 | oldDer, oldCert, key := makeCertificateWithExpiry(as, logger, cl, token, nil, time.Hour) |
| 181 | |
| 182 | // Parse the key so we can create the renewed cert with the same key |
| 183 | privKey, err := pki.UnmarshalPrivateKey([]byte(key)) |
| 184 | as.NoError(err) |
| 185 | |
| 186 | // Create a renewed certificate (fresh, 180 days) |
| 187 | newDer, newCert, _ := makeCertificate(as, logger, cl, token, privKey) |
| 188 | |
| 189 | cfg := &Config{ |
| 190 | path: file.Name(), |
| 191 | router: skipmap.NewString[route](), |
| 192 | Apex: testApex, |
| 193 | Certificate: oldCert, |
| 194 | PrivKey: key, |
| 195 | Tunnels: []Tunnel{ |
| 196 | { |
| 197 | Target: "tcp://127.0.0.1:2345", |
| 198 | }, |
| 199 | }, |
| 200 | } |
| 201 | as.NoError(cfg.validate()) |
| 202 | |
| 203 | // Use a channel to signal when RenewCertificate is called |
| 204 | renewCalled := make(chan struct{}, 1) |
| 205 | |
| 206 | pkiClient := new(mocks.PKIClient) |
| 207 | pkiClient.On("RenewCertificate", mock.Anything, mock.Anything).Run(func(args mock.Arguments) { |
| 208 | select { |
| 209 | case renewCalled <- struct{}{}: |
| 210 | default: |
| 211 | } |
| 212 | }).Return(&protocol.CertificateResponse{ |
| 213 | CertDer: newDer, |
| 214 | CertPem: []byte(newCert), |
| 215 | }, nil).Maybe() // May be called multiple times due to both Register and background maintainer |
| 216 | |
| 217 | m := func(s *mocks.TunnelService, t1 *mocks.MemoryTransport, publishCall *mock.Call) { |
| 218 | defaultNoHostnames(s) |
| 219 | // Certificate may be updated multiple times |
| 220 | t1.On("WithClientCertificate", mock.Anything).Run(func(args mock.Arguments) { |
| 221 | cert := args.Get(0).(tls.Certificate) |
| 222 | parsed, err := x509.ParseCertificate(cert.Certificate[0]) |
| 223 | if err != nil { |
| 224 | panic(err) |
| 225 | } |
| 226 | t1.WithCertificate(parsed) |
| 227 | }).Return(nil).Maybe() |
| 228 | _ = oldDer // used implicitly via oldCert |
| 229 | } |
| 230 | |
| 231 | client, _, assertion := setupClient(t, as, ctx, logger, pkiClient, cfg, nil, m, true, 1) |
| 232 | defer assertion() |
| 233 | defer client.Close() |
| 234 | |
| 235 | // Start the client (this starts certificateMaintainer goroutine) |
| 236 | client.Start(ctx) |
| 237 | |
| 238 | // Wait for background renewal to be triggered (certCheckInterval is 100ms in tests) |
| 239 | // The renewal might have already happened in Register, so we just verify it was called |
| 240 | select { |
| 241 | case <-renewCalled: |
| 242 | // RenewCertificate was called (either in Register or by background maintainer) |
| 243 | case <-time.After(time.Second): |
| 244 | t.Fatal("Expected RenewCertificate to be called within timeout") |
| 245 | } |
| 246 | |
| 247 | // Verify that the certificate was renewed (config should have the new cert) |
| 248 | currCfg := client.GetCurrentConfig() |
| 249 | as.Equal(newCert, currCfg.Certificate) |
| 250 | |
| 251 | // Verify PKIClient.RenewCertificate was called at least once |
| 252 | pkiClient.AssertCalled(t, "RenewCertificate", mock.Anything, mock.Anything) |
| 253 | } |