Skip to content
File

Blob: tun/server/client_rpc.go

go483 lines
1package server
2 
3import (
4 "context"
5 "fmt"
6 "net"
7 "net/http"
8 "strings"
9 "time"
10 
11 "go.miragespace.co/specter/spec/chord"
12 "go.miragespace.co/specter/spec/pki"
13 "go.miragespace.co/specter/spec/protocol"
14 "go.miragespace.co/specter/spec/rpc"
15 "go.miragespace.co/specter/spec/transport"
16 "go.miragespace.co/specter/spec/tun"
17 "go.miragespace.co/specter/util"
18 "go.miragespace.co/specter/util/promise"
19 
20 "github.com/go-chi/chi/v5"
21 "github.com/go-chi/chi/v5/middleware"
22 "github.com/go-chi/httprate"
23 "github.com/sethvargo/go-diceware/diceware"
24 "github.com/twitchtv/twirp"
25 "go.uber.org/zap"
26)
27 
28var generator, _ = diceware.NewGenerator(nil)
29 
30const (
31 testDatagramData = "test"
32 lookupTimeout = time.Second * 3
33 publishTimeout = time.Second * 3
34)
35 
36func (s *Server) logError(ctx context.Context, err twirp.Error) context.Context {
37 switch err.Code() {
38 case twirp.FailedPrecondition:
39 fallthrough
40 case twirp.InvalidArgument:
41 return ctx
42 default:
43 }
44 delegation := rpc.GetDelegation(ctx)
45 if delegation != nil {
46 service, _ := twirp.ServiceName(ctx)
47 method, _ := twirp.MethodName(ctx)
48 l := s.Logger.With(
49 zap.String("remote", delegation.RemoteAddr().String()),
50 zap.Object("peer", delegation.Identity),
51 zap.String("service", service),
52 zap.String("method", method),
53 )
54 cause, key := rpc.GetErrorMeta(err)
55 if cause != "" {
56 l = l.With(zap.String("cause", cause))
57 }
58 if key != "" {
59 l = l.With(zap.String("kv-key", key))
60 }
61 l.Error("Error handling RPC request", zap.Error(err))
62 }
63 return ctx
64}
65 
66func (s *Server) attachRPC(ctx context.Context, router *transport.StreamRouter) {
67 tunTwirp := protocol.NewTunnelServiceServer(s, twirp.WithServerHooks(&twirp.ServerHooks{
68 RequestRouted: s.verifyClientIdentity,
69 Error: s.logError,
70 }))
71 keylessTwirp := protocol.NewKeylessServiceServer(s, twirp.WithServerHooks(&twirp.ServerHooks{
72 RequestRouted: s.verifyClientIdentity,
73 Error: s.logError,
74 }))
75 
76 rpcHandler := chi.NewRouter()
77 rpcHandler.Use(middleware.Recoverer)
78 rpcHandler.Use(httprate.LimitByIP(10, time.Second))
79 rpcHandler.Use(util.LimitBody(1 << 10)) // 1KB
80 rpcHandler.Mount(tunTwirp.PathPrefix(), tunTwirp)
81 rpcHandler.Mount(keylessTwirp.PathPrefix(), keylessTwirp)
82 
83 srv := &http.Server{
84 BaseContext: func(l net.Listener) context.Context {
85 return ctx
86 },
87 ConnContext: func(ctx context.Context, c net.Conn) context.Context {
88 return rpc.WithDelegation(ctx, c.(*transport.StreamDelegate))
89 },
90 MaxHeaderBytes: 1 << 10, // 1KB
91 ReadHeaderTimeout: time.Second * 3,
92 Handler: rpcHandler,
93 ErrorLog: util.GetStdLogger(s.Logger, "rpcServer"),
94 }
95 
96 go srv.Serve(s.rpcAcceptor)
97 
98 router.HandleTunnel(protocol.Stream_RPC, func(delegate *transport.StreamDelegate) {
99 s.rpcAcceptor.Handle(delegate)
100 })
101}
102 
103func (s *Server) verifyClientIdentity(ctx context.Context) (context.Context, error) {
104 method, _ := twirp.MethodName(ctx)
105 delegation := rpc.GetDelegation(ctx)
106 if delegation == nil {
107 return nil, twirp.Internal.Error("delegation missing in context")
108 }
109 
110 switch method {
111 case "Ping":
112 return ctx, nil
113 case "RegisterIdentity", "GetNodes", "OpenEphemeralSession", "OpenDelegatedSession":
114 // These handlers validate their certificate or attachment without a
115 // registered owner. All other methods retain owner authentication.
116 return ctx, nil
117 default:
118 token, verifiedClient, err := extractAuthenticated(ctx)
119 if err != nil {
120 return ctx, err
121 }
122 cli, err := s.getClientByToken(ctx, token)
123 if err != nil {
124 return ctx, twirp.Unauthenticated.Errorf("failed to verify client token: %w", err)
125 }
126 // old format: before PKI
127 if cli.GetAddress() == "" || !cli.GetRendezvous() {
128 s.saveClientToken(ctx, token, verifiedClient)
129 }
130 return ctx, nil
131 }
132}
133 
134func (s *Server) Ping(_ context.Context, _ *protocol.ClientPingRequest) (*protocol.ClientPingResponse, error) {
135 return &protocol.ClientPingResponse{
136 Node: s.TunnelTransport.Identity(),
137 Apex: s.Apex,
138 }, nil
139}
140 
141func (s *Server) RegisterIdentity(ctx context.Context, req *protocol.RegisterIdentityRequest) (*protocol.RegisterIdentityResponse, error) {
142 token, verifiedClient, err := extractAuthenticated(ctx)
143 if err != nil {
144 return nil, err
145 }
146 
147 err = s.TunnelTransport.SendDatagram(verifiedClient, []byte(testDatagramData))
148 if err != nil {
149 return nil, twirp.Aborted.Error("client is not connected")
150 }
151 
152 if err := s.saveClientToken(ctx, token, verifiedClient); err != nil {
153 return nil, rpc.WrapErrorKV(tun.ClientTokenKey(token), err)
154 }
155 
156 return &protocol.RegisterIdentityResponse{
157 Apex: s.Apex,
158 }, nil
159}
160 
161func (s *Server) GetNodes(ctx context.Context, _ *protocol.GetNodesRequest) (*protocol.GetNodesResponse, error) {
162 _, _, err := extractAuthenticated(ctx)
163 if err != nil {
164 return nil, err
165 }
166 
167 successors, err := s.Chord.GetSuccessors()
168 if err != nil {
169 return nil, twirp.Internal.Error(err.Error())
170 }
171 
172 vnodes := chord.MakeSuccListByAddress(s.Chord, successors, tun.NumRedundantLinks)
173 lookupJobs := make([]func(context.Context) (*protocol.Node, error), 0)
174 for _, chord := range vnodes {
175 if chord == nil {
176 continue
177 }
178 chord := chord
179 lookupJobs = append(lookupJobs, func(fnCtx context.Context) (*protocol.Node, error) {
180 key := tun.DestinationByChordKey(chord.Identity())
181 destination, err := s.lookupDestination(fnCtx, key)
182 if err != nil {
183 return nil, rpc.WrapErrorKV(key, err)
184 }
185 return destination.GetTunnel(), nil
186 })
187 }
188 
189 lookupCtx, lookupCancel := context.WithTimeout(ctx, lookupTimeout)
190 defer lookupCancel()
191 
192 servers, errors := promise.All(lookupCtx, lookupJobs...)
193 for _, err := range errors {
194 if err != nil {
195 return nil, err
196 }
197 }
198 
199 return &protocol.GetNodesResponse{
200 Nodes: servers,
201 }, nil
202}
203 
204func (s *Server) GenerateHostname(ctx context.Context, req *protocol.GenerateHostnameRequest) (*protocol.GenerateHostnameResponse, error) {
205 token, _, err := extractAuthenticated(ctx)
206 if err != nil {
207 return nil, err
208 }
209 
210 hostname := strings.Join(generator.MustGenerate(5), "-")
211 prefix := tun.ClientHostnamesPrefix(token)
212 if err := s.Chord.PrefixAppend(ctx, []byte(prefix), []byte(hostname)); err != nil {
213 return nil, rpc.WrapErrorKV(prefix, err)
214 }
215 
216 return &protocol.GenerateHostnameResponse{
217 Hostname: hostname,
218 }, nil
219}
220 
221func (s *Server) RegisteredHostnames(ctx context.Context, req *protocol.RegisteredHostnamesRequest) (*protocol.RegisteredHostnamesResponse, error) {
222 token, _, err := extractAuthenticated(ctx)
223 if err != nil {
224 return nil, err
225 }
226 
227 prefix := tun.ClientHostnamesPrefix(token)
228 
229 children, err := s.Chord.PrefixList(ctx, []byte(prefix))
230 if err != nil {
231 return nil, rpc.WrapErrorKV(prefix, err)
232 }
233 
234 hostnames := make([]string, len(children))
235 for i, child := range children {
236 hostnames[i] = string(child)
237 }
238 
239 return &protocol.RegisteredHostnamesResponse{
240 Hostnames: hostnames,
241 }, nil
242}
243 
244func (s *Server) PublishTunnel(ctx context.Context, req *protocol.PublishTunnelRequest) (*protocol.PublishTunnelResponse, error) {
245 token, verifiedClient, err := extractAuthenticated(ctx)
246 if err != nil {
247 return nil, err
248 }
249 
250 requested := uniqueNodes(req.GetServers())
251 if len(requested) > tun.NumRedundantLinks {
252 return nil, twirp.InvalidArgument.Error("too many requested endpoints")
253 }
254 if len(requested) < 1 {
255 return nil, twirp.InvalidArgument.Error("no servers specified in request")
256 }
257 
258 leaseKey := tun.ClientLeaseKey(token)
259 lease, err := s.Chord.Acquire(ctx, []byte(leaseKey), time.Second*30)
260 if err != nil {
261 return nil, rpc.WrapErrorKV(leaseKey, err)
262 }
263 defer s.Chord.Release(ctx, []byte(leaseKey), lease)
264 
265 hostname := req.GetHostname()
266 prefix := tun.ClientHostnamesPrefix(token)
267 b, err := s.Chord.PrefixContains(ctx, []byte(prefix), []byte(hostname))
268 if err != nil {
269 return nil, rpc.WrapErrorKV(prefix, err)
270 }
271 if !b {
272 return nil, twirp.PermissionDenied.Errorf("hostname %s is not registered", hostname)
273 }
274 
275 lookupJobs := make([]func(context.Context) (*protocol.TunnelDestination, error), len(requested))
276 for i, server := range requested {
277 key := tun.DestinationByTunnelKey(server)
278 lookupJobs[i] = func(fnCtx context.Context) (*protocol.TunnelDestination, error) {
279 destination, err := s.lookupDestination(fnCtx, key)
280 if err != nil {
281 return nil, rpc.WrapErrorKV(key, err)
282 }
283 return destination, nil
284 }
285 }
286 
287 lookupCtx, lookupCancel := context.WithTimeout(ctx, lookupTimeout)
288 defer lookupCancel()
289 
290 destinations, errors := promise.All(lookupCtx, lookupJobs...)
291 for _, err := range errors {
292 if err != nil {
293 return nil, err
294 }
295 }
296 
297 publishJobs := make([]func(context.Context) (*protocol.Node, error), len(destinations))
298 for i, destination := range destinations {
299 dst := destination
300 publishJobs[i] = func(fnCtx context.Context) (*protocol.Node, error) {
301 bundle := &protocol.TunnelRoute{
302 ClientDestination: verifiedClient,
303 ChordDestination: dst.GetChord(),
304 TunnelDestination: dst.GetTunnel(),
305 Hostname: hostname,
306 }
307 val, err := bundle.MarshalVT()
308 if err != nil {
309 return nil, err
310 }
311 key := tun.RoutingKey(hostname, i+1)
312 if err := s.Chord.Put(fnCtx, []byte(key), val); err != nil {
313 s.Logger.Error("Error publishing route", zap.String("key", key), zap.Error(err))
314 return nil, nil
315 }
316 return dst.GetTunnel(), nil
317 }
318 }
319 
320 publishCtx, publishCancel := context.WithTimeout(ctx, publishTimeout)
321 defer publishCancel()
322 
323 maybePublished, errors := promise.All(publishCtx, publishJobs...)
324 for _, err := range errors {
325 if err != nil {
326 return nil, twirp.InternalError(err.Error())
327 }
328 }
329 // need to filter possible nil
330 published := make([]*protocol.Node, 0)
331 for _, node := range maybePublished {
332 if node == nil {
333 continue
334 }
335 published = append(published, node)
336 }
337 
338 if len(published) == 0 {
339 return nil, twirp.Unavailable.Error("unable to publish any routes")
340 }
341 
342 return &protocol.PublishTunnelResponse{
343 Published: published,
344 }, nil
345}
346 
347func (s *Server) UnpublishTunnel(ctx context.Context, req *protocol.UnpublishTunnelRequest) (*protocol.UnpublishTunnelResponse, error) {
348 token, _, err := extractAuthenticated(ctx)
349 if err != nil {
350 return nil, err
351 }
352 
353 lease, err := s.Chord.Acquire(ctx, []byte(tun.ClientLeaseKey(token)), time.Second*30)
354 if err != nil {
355 return nil, twirp.Internal.Errorf("error acquiring lease for unpublishing tunnel: %w", err)
356 }
357 defer s.Chord.Release(ctx, []byte(tun.ClientLeaseKey(token)), lease)
358 
359 hostname := req.GetHostname()
360 if err := s.unadvertiseTunnel(ctx, token, hostname); err != nil {
361 return nil, err
362 }
363 
364 return &protocol.UnpublishTunnelResponse{}, nil
365}
366 
367func (s *Server) ReleaseTunnel(ctx context.Context, req *protocol.ReleaseTunnelRequest) (*protocol.ReleaseTunnelResponse, error) {
368 token, client, err := extractAuthenticated(ctx)
369 if err != nil {
370 return nil, err
371 }
372 
373 lease, err := s.Chord.Acquire(ctx, []byte(tun.ClientLeaseKey(token)), time.Second*30)
374 if err != nil {
375 return nil, twirp.Internal.Errorf("error acquiring lease for releasing tunnel: %w", err)
376 }
377 defer s.Chord.Release(ctx, []byte(tun.ClientLeaseKey(token)), lease)
378 
379 hostname := req.GetHostname()
380 if err := s.unadvertiseTunnel(ctx, token, hostname); err != nil {
381 return nil, err
382 }
383 prefix := tun.ClientHostnamesPrefix(token)
384 if err := s.Chord.PrefixRemove(ctx, []byte(prefix), []byte(hostname)); err != nil {
385 return nil, rpc.WrapErrorKV(prefix, err)
386 }
387 
388 if err := tun.RemoveCustomHostname(ctx, s.Chord, hostname); err != nil {
389 s.Logger.Warn("Failed to remove custom hostname when releasing tunnel", zap.String("hostname", hostname), zap.Object("client", client), zap.Error(err))
390 }
391 
392 return &protocol.ReleaseTunnelResponse{}, nil
393}
394 
395func (s *Server) unadvertiseTunnel(ctx context.Context, token *protocol.ClientToken, hostname string) error {
396 prefix := tun.ClientHostnamesPrefix(token)
397 b, err := s.Chord.PrefixContains(ctx, []byte(prefix), []byte(hostname))
398 if err != nil {
399 return rpc.WrapErrorKV(prefix, err)
400 }
401 if !b {
402 return twirp.PermissionDenied.Errorf("hostname %s is not registered", hostname)
403 }
404 
405 unpublishJobs := make([]func(context.Context) (int, error), tun.NumRedundantLinks)
406 for i := range tun.NumRedundantLinks {
407 unpublishJobs[i] = func(fnCtx context.Context) (int, error) {
408 key := tun.RoutingKey(hostname, i+1)
409 err := s.Chord.Delete(fnCtx, []byte(key))
410 return 0, err
411 }
412 }
413 
414 unpublishCtx, unpublishCancel := context.WithTimeout(ctx, publishTimeout)
415 defer unpublishCancel()
416 
417 _, errors := promise.All(unpublishCtx, unpublishJobs...)
418 for _, err := range errors {
419 if err != nil {
420 return twirp.InternalError(err.Error())
421 }
422 }
423 
424 return nil
425}
426 
427func (s *Server) saveClientToken(ctx context.Context, token *protocol.ClientToken, client *protocol.Node) error {
428 val, err := client.MarshalVT()
429 if err != nil {
430 return err
431 }
432 
433 if err := s.Chord.Put(ctx, []byte(tun.ClientTokenKey(token)), val); err != nil {
434 return err
435 }
436 return nil
437}
438 
439func (s *Server) getClientByToken(ctx context.Context, token *protocol.ClientToken) (*protocol.Node, error) {
440 val, err := s.Chord.Get(ctx, []byte(tun.ClientTokenKey(token)))
441 if err != nil {
442 return nil, err
443 }
444 if len(val) == 0 {
445 return nil, fmt.Errorf("no existing client found with given token")
446 }
447 client := &protocol.Node{}
448 if err := client.UnmarshalVT(val); err != nil {
449 return nil, err
450 }
451 return client, nil
452}
453 
454func extractAuthenticated(ctx context.Context) (*protocol.ClientToken, *protocol.Node, error) {
455 delegation := rpc.GetDelegation(ctx)
456 if delegation == nil {
457 return nil, nil, twirp.Internal.Error("delegation missing in context")
458 }
459 if delegation.Certificate == nil {
460 return nil, nil, twirp.Unauthenticated.Error("missing client certificate")
461 }
462 identity, err := pki.ExtractCertificateIdentity(delegation.Certificate)
463 if err != nil {
464 return nil, nil, twirp.Unauthenticated.Error(err.Error())
465 }
466 return &protocol.ClientToken{
467 Token: identity.Token,
468 }, identity.NodeIdentity(), nil
469}
470 
471func uniqueNodes(nodes []*protocol.Node) []*protocol.Node {
472 list := make([]*protocol.Node, 0)
473 seen := make(map[string]bool)
474 for _, node := range nodes {
475 if node == nil || seen[node.GetAddress()] {
476 continue
477 }
478 seen[node.GetAddress()] = true
479 list = append(list, node)
480 }
481 return list
482}