Skip to content
File

Blob: tun/client/certificate_test.go

go254 lines
1package client
2 
3import (
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 
23func 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 
63func 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 
134func 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 
162func 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}