Skip to content
File

Blob: spec/rpc/rpc.go

go214 lines
1package rpc
2 
3import (
4 "context"
5 "encoding/binary"
6 "fmt"
7 "io"
8 "net"
9 "net/http"
10 "strconv"
11 "strings"
12 "sync"
13 "time"
14 
15 "go.miragespace.co/specter/spec/protocol"
16 "go.miragespace.co/specter/spec/transport"
17 "go.miragespace.co/specter/timing"
18 "go.miragespace.co/specter/util/ratecounter"
19 
20 pool "github.com/libp2p/go-buffer-pool"
21 "github.com/twitchtv/twirp"
22)
23 
24const (
25 // uint32
26 LengthSize = 4
27)
28 
29var builderPool = sync.Pool{
30 New: func() any {
31 return &strings.Builder{}
32 },
33}
34 
35type ChordClient interface {
36 protocol.VNodeService
37 protocol.KVService
38 RatePer(interval time.Duration) float64
39}
40 
41type TunnelClient interface {
42 protocol.TunnelService
43 protocol.KeylessService
44}
45 
46func getDynamicDialer(baseCtx context.Context, transport transport.Transport) func(reqCtx context.Context, _, _ string) (net.Conn, error) {
47 return func(reqCtx context.Context, _, _ string) (net.Conn, error) {
48 peer := GetNode(reqCtx)
49 if peer == nil {
50 return nil, fmt.Errorf("node not found in context")
51 }
52 return transport.DialStream(baseCtx, peer, protocol.Stream_RPC)
53 }
54}
55 
56// DynamicChordClient returns a rpc client suitable for both VNodeService and KVService, with the destination
57// set per call dynamically according to the destination in context. Use WithNode(ctx, node) at call site to dynamically dispatch.
58func DynamicChordClient(baseContext context.Context, chordTransport transport.Transport) ChordClient {
59 outboundRate := ratecounter.New(time.Second, time.Second*5)
60 
61 injector := &twirp.ClientHooks{
62 RequestPrepared: func(ctx context.Context, r *http.Request) (context.Context, error) {
63 peer := GetNode(ctx)
64 if peer == nil {
65 return nil, fmt.Errorf("node not found in context")
66 }
67 SerializeContextHeader(ctx, r.Header)
68 
69 // needed to override dialer instead of using http://chord as key
70 sb := builderPool.Get().(*strings.Builder)
71 defer builderPool.Put(sb)
72 defer sb.Reset()
73 sb.WriteString(strconv.FormatUint(peer.GetId(), 10))
74 sb.WriteString(".")
75 sb.WriteString(peer.GetAddress())
76 r.URL.Host = sb.String()
77 
78 outboundRate.Increment()
79 return ctx, nil
80 },
81 }
82 
83 // default to http client pooling
84 t := http.DefaultTransport.(*http.Transport).Clone()
85 t.MaxConnsPerHost = 50
86 t.MaxIdleConnsPerHost = 5
87 t.DisableCompression = true
88 t.IdleConnTimeout = timing.RPCIdleTimeout
89 t.DialTLSContext = getDynamicDialer(baseContext, chordTransport)
90 c := &http.Client{
91 Transport: t,
92 }
93 // disable in testing
94 if disable, ok := baseContext.Value(contextDisablePoolKey).(bool); ok && disable {
95 t.DisableKeepAlives = true
96 t.MaxConnsPerHost = -1
97 }
98 
99 return &struct {
100 protocol.VNodeService
101 protocol.KVService
102 *ratecounter.Rate
103 }{
104 VNodeService: protocol.NewVNodeServiceProtobufClient("https://chord", c, twirp.WithClientHooks(injector)),
105 KVService: protocol.NewKVServiceProtobufClient("https://chord", c, twirp.WithClientHooks(injector)),
106 Rate: outboundRate,
107 }
108}
109 
110// DynamicTunnelClient returns a rpc client suitable for TunnelService, with the destination set per call dynamically
111// according to the destination in context. Use WithNode(ctx, node) at call site to dynamically dispatch. Optionally,
112// use WithClientToken(ctx, token) to include client token.
113func DynamicTunnelClient(baseContext context.Context, tunnelTransport transport.Transport) TunnelClient {
114 injector := &twirp.ClientHooks{
115 RequestPrepared: func(ctx context.Context, r *http.Request) (context.Context, error) {
116 peer := GetNode(ctx)
117 if peer == nil {
118 return nil, fmt.Errorf("peer not found in context")
119 }
120 r.URL.Host = peer.GetAddress() // needed to override dialer instead of using http://tunnel as key
121 return ctx, nil
122 },
123 }
124 
125 // default to http client pooling
126 t := http.DefaultTransport.(*http.Transport).Clone()
127 t.MaxConnsPerHost = 10
128 t.MaxIdleConnsPerHost = 1
129 t.DisableCompression = true
130 t.IdleConnTimeout = timing.RPCIdleTimeout
131 t.DialTLSContext = getDynamicDialer(baseContext, tunnelTransport)
132 c := &http.Client{
133 Transport: t,
134 }
135 // disable in testing
136 if disable, ok := baseContext.Value(contextDisablePoolKey).(bool); ok && disable {
137 t.DisableKeepAlives = true
138 t.MaxConnsPerHost = -1
139 }
140 
141 return &struct {
142 protocol.TunnelService
143 protocol.KeylessService
144 }{
145 TunnelService: protocol.NewTunnelServiceProtobufClient("https://tunnel", c, twirp.WithClientHooks(injector)),
146 KeylessService: protocol.NewKeylessServiceProtobufClient("https://tunnel", c, twirp.WithClientHooks(injector)),
147 }
148}
149 
150func receive(stream io.Reader, rr VTMarshaler, checker func(size uint32) bool) error {
151 var sb [LengthSize]byte
152 
153 n, err := io.ReadFull(stream, sb[:])
154 if err != nil {
155 return fmt.Errorf("reading RPC message buffer size: %w", err)
156 }
157 if n != LengthSize {
158 return fmt.Errorf("expected %d bytes to be read but %d bytes was read", LengthSize, n)
159 }
160 
161 ms := binary.BigEndian.Uint32(sb[:])
162 if !checker(ms) {
163 return fmt.Errorf("RPC message is too large")
164 }
165 
166 mb := pool.Get(int(ms))
167 defer pool.Put(mb)
168 
169 n, err = io.ReadFull(stream, mb)
170 if err != nil {
171 return fmt.Errorf("reading RPC message: %w", err)
172 }
173 if ms != uint32(n) {
174 return fmt.Errorf("expected %d bytes to be read but %d bytes was read", ms, n)
175 }
176 
177 return rr.UnmarshalVT(mb)
178}
179 
180func BoundedReceive(stream io.Reader, rr VTMarshaler, max uint32) error {
181 return receive(stream, rr, func(size uint32) bool {
182 return size <= max
183 })
184}
185 
186func Receive(stream io.Reader, rr VTMarshaler) error {
187 return receive(stream, rr, func(size uint32) bool {
188 return true
189 })
190}
191 
192func Send(stream io.Writer, rr VTMarshaler) error {
193 l := rr.SizeVT()
194 mb := pool.Get(LengthSize + l)
195 defer pool.Put(mb)
196 
197 binary.BigEndian.PutUint32(mb[0:LengthSize], uint32(l))
198 
199 _, err := rr.MarshalToSizedBufferVT(mb[LengthSize:])
200 if err != nil {
201 return fmt.Errorf("encoding outbound RPC message: %w", err)
202 }
203 
204 n, err := stream.Write(mb)
205 if err != nil {
206 return fmt.Errorf("sending RPC message: %w", err)
207 }
208 if n != LengthSize+l {
209 return fmt.Errorf("expected %d bytes sent but %d bytes was sent", l, n)
210 }
211 
212 return nil
213}