File
Blob: tun/client/fixtures_test.go
| 1 | package client |
| 2 | |
| 3 | import ( |
| 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 | |
| 37 | func 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 | |
| 86 | func 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 | |
| 144 | type mockTunnelClient struct { |
| 145 | mock.Mock |
| 146 | mocks.TunnelService |
| 147 | mocks.KeylessService |
| 148 | } |
| 149 | |
| 150 | func 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 | |
| 233 | func 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. |
| 239 | func 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 | |
| 302 | func 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 | |
| 315 | func 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 | } |