File
Blob: spec/rpc/context.go
| 1 | package rpc |
| 2 | |
| 3 | import ( |
| 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 | |
| 14 | const ( |
| 15 | HeaderRPCContext = "x-rpc-context" |
| 16 | ) |
| 17 | |
| 18 | type rpcContextKey string |
| 19 | |
| 20 | const ( |
| 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 |
| 31 | func DisablePooling(baseCtx context.Context) context.Context { |
| 32 | return context.WithValue(baseCtx, contextDisablePoolKey, true) |
| 33 | } |
| 34 | |
| 35 | // Connect to the provided node in this request |
| 36 | func 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 |
| 41 | func 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 |
| 49 | func 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 |
| 54 | func 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 |
| 62 | func 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 |
| 67 | func 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 |
| 75 | func 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() |
| 94 | func 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 |
| 114 | func 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 | } |