Skip to content
File

Blob: spec/rpc/context.go

go123 lines
1package rpc
2 
3import (
4 "context"
5 "encoding/base64"
6 "net/http"
7 
8 "go.miragespace.co/specter/spec/protocol"
9 "go.miragespace.co/specter/spec/transport"
10 
11 pool "github.com/libp2p/go-buffer-pool"
12)
13 
14const (
15 HeaderRPCContext = "x-rpc-context"
16)
17 
18type rpcContextKey string
19 
20const (
21 contextNodeKey = rpcContextKey("dial-node") // *protocol.Node to connect
22 contextRPCContextKey = rpcContextKey("rpc-context") // *protocol.Context of the rpc request
23 contextAuthorizationKey = rpcContextKey("auth-header") // "Authorization" header from client rpc request
24 contextClientTokenKey = rpcContextKey("client-token") // *protocol.ClientToken from client or parsed from the header
25 contextClientIdentityKey = rpcContextKey("client-identity") // *protocol.Node of client as matched with delegation and token
26 contextDelegationKey = rpcContextKey("stream-delegation") // *transport.StreamDelegation of the rpc request
27 contextDisablePoolKey = rpcContextKey("disable-http-pool") // disable HTTP client pooling. Used in test to avoid lingering connections
28)
29 
30// Disable HTTP client pool for this client
31func DisablePooling(baseCtx context.Context) context.Context {
32 return context.WithValue(baseCtx, contextDisablePoolKey, true)
33}
34 
35// Connect to the provided node in this request
36func WithNode(ctx context.Context, node *protocol.Node) context.Context {
37 return context.WithValue(ctx, contextNodeKey, node)
38}
39 
40// Retrieve the node of this request
41func GetNode(ctx context.Context) *protocol.Node {
42 if node, ok := ctx.Value(contextNodeKey).(*protocol.Node); ok {
43 return node
44 }
45 return nil
46}
47 
48// Send RPC context for this request
49func WithContext(ctx context.Context, rpcCtx *protocol.Context) context.Context {
50 return context.WithValue(ctx, contextRPCContextKey, rpcCtx)
51}
52 
53// Retrieve the RPC context of this request
54func GetContext(ctx context.Context) *protocol.Context {
55 if r, ok := ctx.Value(contextRPCContextKey).(*protocol.Context); ok {
56 return r
57 }
58 return &protocol.Context{}
59}
60 
61// Attach the delegation triggering the request
62func WithDelegation(ctx context.Context, delegate *transport.StreamDelegate) context.Context {
63 return context.WithValue(ctx, contextDelegationKey, delegate)
64}
65 
66// Retrieve the delegation of this request
67func GetDelegation(ctx context.Context) *transport.StreamDelegate {
68 if delegate, ok := ctx.Value(contextDelegationKey).(*transport.StreamDelegate); ok {
69 return delegate
70 }
71 return nil
72}
73 
74// Serialize RPC context as http headers
75func SerializeContextHeader(ctx context.Context, r http.Header) {
76 rCtx, ok := ctx.Value(contextRPCContextKey).(*protocol.Context)
77 if !ok {
78 return
79 }
80 
81 l := rCtx.SizeVT()
82 mb := pool.Get(l)
83 defer pool.Put(mb)
84 
85 _, err := rCtx.MarshalToSizedBufferVT(mb)
86 if err != nil {
87 return
88 }
89 
90 r.Set(HeaderRPCContext, base64.StdEncoding.EncodeToString(mb))
91}
92 
93// Deserialize RPC context from http headers. The RPC context can be retrieved with GetContext()
94func DeserializeContextHeader(ctx context.Context, r http.Header) (context.Context, bool) {
95 encoded := r.Get(HeaderRPCContext)
96 if len(encoded) < 1 {
97 return ctx, false
98 }
99 
100 mb, err := base64.StdEncoding.DecodeString(encoded)
101 if err != nil {
102 return ctx, false
103 }
104 
105 rCtx := &protocol.Context{}
106 if err := rCtx.UnmarshalVT(mb); err != nil {
107 return ctx, false
108 }
109 
110 return WithContext(ctx, rCtx), true
111}
112 
113// Middleware to attach the deserialized RPC context to the current request
114func ExtractContext(base http.Handler) http.Handler {
115 return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
116 ctx, ok := DeserializeContextHeader(r.Context(), r.Header)
117 if ok {
118 r = r.WithContext(ctx)
119 }
120 base.ServeHTTP(w, r)
121 })
122}