Skip to content
File

Blob: tun/client/fixtures_test.go

go321 lines
1package client
2 
3import (
4 "bytes"
5 "context"
6 "crypto/ed25519"
7 "crypto/rand"
8 "crypto/tls"
9 "crypto/x509"
10 "crypto/x509/pkix"
11 "encoding/pem"
12 "fmt"
13 "math/big"
14 "net"
15 "net/http"
16 "os"
17 "testing"
18 "time"
19 
20 "go.miragespace.co/specter/spec/mocks"
21 "go.miragespace.co/specter/spec/pki"
22 "go.miragespace.co/specter/spec/protocol"
23 "go.miragespace.co/specter/spec/rpc"
24 "go.miragespace.co/specter/spec/rtt"
25 "go.miragespace.co/specter/spec/transport"
26 "go.miragespace.co/specter/util"
27 "go.miragespace.co/specter/util/acceptor"
28 
29 "github.com/go-chi/chi/v5"
30 "github.com/go-chi/chi/v5/middleware"
31 "github.com/stretchr/testify/mock"
32 "github.com/stretchr/testify/require"
33 "github.com/twitchtv/twirp"
34 "go.uber.org/zap"
35)
36 
37func setupRPC(ctx context.Context,
38 logger *zap.Logger,
39 s protocol.TunnelService,
40 k protocol.KeylessService,
41 router *transport.StreamRouter,
42 acc *acceptor.HTTP2Acceptor,
43) {
44 hooks := &twirp.ServerHooks{
45 RequestRouted: func(ctx context.Context) (context.Context, error) {
46 delegation := rpc.GetDelegation(ctx)
47 if delegation.Certificate == nil {
48 return ctx, fmt.Errorf("missing client certificate")
49 }
50 return ctx, nil
51 },
52 Error: func(ctx context.Context, err twirp.Error) context.Context {
53 logger.Error("error handling request", zap.Error(err))
54 return ctx
55 },
56 }
57 tunTwirp := protocol.NewTunnelServiceServer(s, twirp.WithServerHooks(hooks))
58 keylessTwirp := protocol.NewKeylessServiceServer(k, twirp.WithServerHooks(hooks))
59 
60 rpcHandler := chi.NewRouter()
61 rpcHandler.Use(middleware.Recoverer)
62 rpcHandler.Use(util.LimitBody(1 << 10)) // 1KB
63 rpcHandler.Mount(tunTwirp.PathPrefix(), tunTwirp)
64 rpcHandler.Mount(keylessTwirp.PathPrefix(), keylessTwirp)
65 
66 srv := &http.Server{
67 BaseContext: func(l net.Listener) context.Context {
68 return ctx
69 },
70 ConnContext: func(ctx context.Context, c net.Conn) context.Context {
71 return rpc.WithDelegation(ctx, c.(*transport.StreamDelegate))
72 },
73 MaxHeaderBytes: 1 << 10, // 1KB
74 ReadTimeout: time.Second * 3,
75 Handler: rpcHandler,
76 ErrorLog: util.GetStdLogger(logger, "rpc_server"),
77 }
78 
79 go srv.Serve(acc)
80 
81 router.HandleTunnel(protocol.Stream_RPC, func(delegate *transport.StreamDelegate) {
82 acc.Handle(delegate)
83 })
84}
85 
86func setupFakeNodes(rr *mocks.Measurement, s *mocks.TunnelService, expectGenerate bool) []*protocol.Node {
87 fakeNodes := []*protocol.Node{
88 {
89 Id: 1,
90 Address: "192.168.0.1",
91 },
92 {
93 Id: 2,
94 Address: "192.168.0.2",
95 },
96 {
97 Id: 3,
98 Address: "192.168.0.3",
99 },
100 }
101 
102 latencies := []*rtt.Statistics{
103 {
104 Average: time.Hour,
105 },
106 {
107 Average: time.Second,
108 },
109 {
110 Average: time.Minute,
111 },
112 }
113 
114 for i, n := range fakeNodes {
115 // inject fake nodes. Dial target is available in delegation.Identity because PipeTransport()
116 rr.On("Snapshot", mock.MatchedBy(func(key string) bool {
117 return rtt.MakeMeasurementKey(n) == key
118 }), mock.Anything).Return(latencies[i])
119 s.On("Ping", mock.MatchedBy(func(ctx context.Context) bool {
120 delegation := rpc.GetDelegation(ctx)
121 return delegation.Identity.String() == n.String()
122 }), mock.Anything).Return(&protocol.ClientPingResponse{
123 Node: fakeNodes[i],
124 Apex: testApex,
125 }, nil).Maybe()
126 }
127 
128 s.On("Ping", mock.Anything, mock.Anything).Return(&protocol.ClientPingResponse{
129 Node: fakeNodes[0],
130 Apex: testApex,
131 }, nil).Maybe()
132 s.On("GetNodes", mock.Anything, mock.Anything).Return(&protocol.GetNodesResponse{
133 Nodes: fakeNodes,
134 }, nil)
135 if expectGenerate {
136 s.On("GenerateHostname", mock.Anything, mock.Anything).Return(&protocol.GenerateHostnameResponse{
137 Hostname: testHostname,
138 }, nil)
139 }
140 
141 return fakeNodes
142}
143 
144type mockTunnelClient struct {
145 mock.Mock
146 mocks.TunnelService
147 mocks.KeylessService
148}
149 
150func setupClient(
151 t *testing.T,
152 as *require.Assertions,
153 ctx context.Context,
154 logger *zap.Logger,
155 pkiClient *mocks.PKIClient,
156 cfg *Config,
157 reload <-chan os.Signal,
158 m func(s *mocks.TunnelService, t1 *mocks.MemoryTransport, publishCall *mock.Call),
159 preference bool,
160 times int,
161) (
162 *Client,
163 *mocks.MemoryTransport,
164 func(),
165) {
166 t1, t2 := mocks.PipeTransport()
167 rr := new(mocks.Measurement)
168 
169 k := new(mocks.KeylessService)
170 s := new(mocks.TunnelService)
171 s.Keyless = k
172 
173 router := transport.NewStreamRouter(logger, nil, t2)
174 acc := acceptor.NewH2Acceptor(nil)
175 
176 expectGenerate := true
177 if cfg.Tunnels[0].Hostname != "" {
178 expectGenerate = false
179 }
180 fakeNodes := setupFakeNodes(rr, s, expectGenerate)
181 
182 var publishCall *mock.Call
183 if preference {
184 publishCall = s.On("PublishTunnel", mock.Anything, mock.MatchedBy(func(req *protocol.PublishTunnelRequest) bool {
185 // ensure that the node with the lowest latency is the first preference
186 preferenceMatched := len(req.GetServers()) > 0 && req.GetServers()[0].GetAddress() == fakeNodes[1].GetAddress()
187 hostnameMatched := req.GetHostname() == testHostname
188 return preferenceMatched && hostnameMatched
189 })).Return(&protocol.PublishTunnelResponse{
190 Published: fakeNodes,
191 }, nil).Times(times)
192 } else {
193 publishCall = s.On("PublishTunnel", mock.Anything, mock.MatchedBy(func(req *protocol.PublishTunnelRequest) bool {
194 return req.GetHostname() == testHostname
195 })).Return(&protocol.PublishTunnelResponse{
196 Published: fakeNodes,
197 }, nil).Times(times)
198 }
199 
200 if m != nil {
201 m(s, t1, publishCall)
202 }
203 
204 go router.Accept(ctx)
205 
206 setupRPC(ctx, logger, s, k, router, acc)
207 
208 client, err := NewClient(rpc.DisablePooling(ctx), ClientConfig{
209 Logger: logger,
210 Configuration: cfg,
211 PKIClient: pkiClient,
212 ServerTransport: t1,
213 Recorder: rr,
214 ReloadSignal: reload,
215 })
216 as.NoError(err)
217 
218 as.NoError(client.Register(ctx))
219 
220 as.NoError(client.Initialize(ctx, true))
221 
222 return client, t2, func() {
223 acc.Close()
224 rr.AssertExpectations(t)
225 s.AssertExpectations(t)
226 k.AssertExpectations(t)
227 if pkiClient != nil {
228 pkiClient.AssertExpectations(t)
229 }
230 }
231}
232 
233func makeCertificate(as *require.Assertions, logger *zap.Logger, client *protocol.Node, token *protocol.ClientToken, privKey ed25519.PrivateKey) (certDer []byte, certPem, keyPem string) {
234 return makeCertificateWithExpiry(as, logger, client, token, privKey, time.Hour*24*180)
235}
236 
237// makeCertificateWithExpiry creates a certificate that expires in the given duration.
238// Use a short duration (e.g., time.Hour) to create a certificate within the renewal window for testing.
239func makeCertificateWithExpiry(as *require.Assertions, logger *zap.Logger, client *protocol.Node, token *protocol.ClientToken, privKey ed25519.PrivateKey, expiresIn time.Duration) (certDer []byte, certPem, keyPem string) {
240 // generate a CA
241 caPubKey, caPrivKey, err := ed25519.GenerateKey(rand.Reader)
242 as.NoError(err)
243 
244 caTemplate := x509.Certificate{
245 SerialNumber: big.NewInt(1),
246 Subject: pkix.Name{
247 Organization: []string{"dev"},
248 },
249 NotBefore: time.Now(),
250 NotAfter: time.Now().Add(time.Hour * 24 * 365), // CA valid for 1 year
251 KeyUsage: x509.KeyUsageKeyEncipherment | x509.KeyUsageDigitalSignature | x509.KeyUsageCertSign,
252 ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth, x509.ExtKeyUsageClientAuth},
253 BasicConstraintsValid: true,
254 IsCA: true,
255 }
256 
257 caBytes, err := x509.CreateCertificate(rand.Reader, &caTemplate, &caTemplate, caPubKey, caPrivKey)
258 as.NoError(err)
259 
260 var (
261 certPubKey ed25519.PublicKey
262 certPrivKey ed25519.PrivateKey
263 )
264 
265 if privKey != nil {
266 certPubKey = privKey.Public().(ed25519.PublicKey)
267 certPrivKey = privKey
268 } else {
269 certPubKey, certPrivKey, err = ed25519.GenerateKey(rand.Reader)
270 as.NoError(err)
271 }
272 
273 // Create client certificate with custom expiry using pki.GenerateCertificate
274 caTLS := tls.Certificate{
275 Certificate: [][]byte{caBytes},
276 PrivateKey: caPrivKey,
277 }
278 
279 der, err := pki.GenerateCertificate(logger, caTLS, pki.IdentityRequest{
280 PublicKey: certPubKey,
281 Subject: pki.MakeSubjectV2(client.GetId(), token.GetToken()),
282 ValidFor: expiresIn,
283 })
284 as.NoError(err)
285 
286 x509PrivKey, err := x509.MarshalPKCS8PrivateKey(certPrivKey)
287 if err != nil {
288 panic(err)
289 }
290 certPrivKeyPEM := new(bytes.Buffer)
291 pem.Encode(certPrivKeyPEM, &pem.Block{
292 Type: "PRIVATE KEY",
293 Bytes: x509PrivKey,
294 })
295 
296 certDer = der
297 certPem = string(pki.MarshalCertificate(der))
298 keyPem = certPrivKeyPEM.String()
299 return
300}
301 
302func transportHelper(t *mocks.MemoryTransport, der []byte) {
303 t.On("WithClientCertificate", mock.MatchedBy(func(cert tls.Certificate) bool {
304 return bytes.Equal(der, cert.Certificate[0])
305 })).Run(func(args mock.Arguments) {
306 cert := args.Get(0).(tls.Certificate)
307 parsed, err := x509.ParseCertificate(cert.Certificate[0])
308 if err != nil {
309 panic(err)
310 }
311 t.WithCertificate(parsed)
312 }).Return(nil).Once()
313}
314 
315func defaultNoHostnames(s *mocks.TunnelService) {
316 resp := &protocol.RegisteredHostnamesResponse{
317 Hostnames: make([]string, 0),
318 }
319 s.On("RegisteredHostnames", mock.Anything, mock.Anything).Return(resp, nil)
320}