Skip to content

Commit 57e16e0

Browse files
infrmtcsNazariiDenha
authored andcommitted
perf(rpc): rpc versions use same pool
1 parent 6d9f798 commit 57e16e0

15 files changed

Lines changed: 360 additions & 26 deletions

jsonrpc/websocket.go

Lines changed: 40 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@ package jsonrpc
22

33
import (
44
"context"
5+
"encoding/json"
56
"errors"
67
"io"
78
"net/http"
@@ -20,13 +21,25 @@ const (
2021
maxConns = 2048 // TODO: an arbitrary default number, should be revisited after monitoring
2122
)
2223

24+
var serverBusyResponse = func() []byte {
25+
b, err := json.Marshal(&response{
26+
Version: "2.0",
27+
Error: &Error{Code: InternalError, Message: ErrServerBusy.Error()},
28+
})
29+
if err != nil {
30+
panic(err)
31+
}
32+
return b
33+
}()
34+
2335
type Websocket struct {
2436
rpc *Server
2537
logger log.StructuredLogger
2638
connParams *WebsocketConnParams
2739
listener NewRequestListener
2840
shutdown <-chan struct{}
2941
requestTimeout time.Duration
42+
gate *Gate
3043

3144
// Add connection tracking
3245
connSem *semaphore.Weighted
@@ -62,6 +75,11 @@ func (ws *Websocket) WithRequestTimeout(d time.Duration) *Websocket {
6275
return ws
6376
}
6477

78+
func (ws *Websocket) WithGate(g *Gate) *Websocket {
79+
ws.gate = g
80+
return ws
81+
}
82+
6583
// WithListener registers a NewRequestListener
6684
func (ws *Websocket) WithListener(listener NewRequestListener) *Websocket {
6785
ws.listener = listener
@@ -116,7 +134,7 @@ func (ws *Websocket) ServeHTTP(w http.ResponseWriter, r *http.Request) {
116134
break
117135
}
118136
ws.listener.OnNewRequest("any")
119-
if err = ws.rpc.HandleReadWriter(wsc.ctx, ws.requestTimeout, wsc); err != nil {
137+
if err = ws.handleMessage(wsc); err != nil {
120138
break
121139
}
122140
// From websocket docs: "Read to EOF otherwise connection will hang."
@@ -146,6 +164,27 @@ func (ws *Websocket) ServeHTTP(w http.ResponseWriter, r *http.Request) {
146164
}
147165
}
148166

167+
func (ws *Websocket) handleMessage(wsc *websocketConn) error {
168+
if ws.gate != nil {
169+
acquireCtx := wsc.ctx
170+
if ws.requestTimeout > 0 {
171+
var cancel context.CancelFunc
172+
acquireCtx, cancel = context.WithTimeout(acquireCtx, ws.requestTimeout)
173+
defer cancel()
174+
}
175+
if err := ws.gate.Acquire(acquireCtx); err != nil {
176+
if errors.Is(err, context.Canceled) {
177+
return err
178+
}
179+
_, writeErr := wsc.Write(serverBusyResponse)
180+
return writeErr
181+
}
182+
defer ws.gate.Release()
183+
}
184+
185+
return ws.rpc.HandleReadWriter(wsc.ctx, ws.requestTimeout, wsc)
186+
}
187+
149188
type WebsocketConnParams struct {
150189
// Maximum message size allowed.
151190
ReadLimit int64

jsonrpc/websocket_test.go

Lines changed: 56 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -296,3 +296,59 @@ func TestWebsocketConnectionLimit(t *testing.T) {
296296
require.Equal(t, http.StatusSwitchingProtocols, resp4.StatusCode)
297297
require.NoError(t, conn4.Close(websocket.StatusNormalClosure, ""))
298298
}
299+
300+
func TestWebsocketGateRejectsWhenBusy(t *testing.T) {
301+
started := make(chan struct{})
302+
release := make(chan struct{})
303+
block := jsonrpc.Method{
304+
Name: "test_block",
305+
Handler: func(ctx context.Context) (int, *jsonrpc.Error) {
306+
close(started)
307+
<-release
308+
return 0, nil
309+
},
310+
}
311+
echo := jsonrpc.Method{
312+
Name: "test_echo",
313+
Params: []jsonrpc.Parameter{{Name: "msg"}},
314+
Handler: func(msg string) (string, *jsonrpc.Error) { return msg, nil },
315+
}
316+
317+
rpc := jsonrpc.NewServer(1, log.NewNopZapLogger())
318+
require.NoError(t, rpc.RegisterMethods(block, echo))
319+
gate := jsonrpc.NewGate(1, 0)
320+
ws := jsonrpc.NewWebsocket(rpc, nil, log.NewNopZapLogger()).WithGate(gate)
321+
srv := httptest.NewServer(ws)
322+
t.Cleanup(srv.Close)
323+
324+
connA, respA, err := websocket.Dial(t.Context(), srv.URL, nil) //nolint:bodyclose // lib closes it
325+
require.NoError(t, err)
326+
require.Equal(t, http.StatusSwitchingProtocols, respA.StatusCode)
327+
defer connA.Close(websocket.StatusNormalClosure, "")
328+
require.NoError(t, connA.Write(t.Context(), websocket.MessageText,
329+
[]byte(`{"jsonrpc":"2.0","method":"test_block","params":[],"id":1}`)))
330+
<-started
331+
332+
connB, respB, err := websocket.Dial(t.Context(), srv.URL, nil) //nolint:bodyclose // lib closes it
333+
require.NoError(t, err)
334+
require.Equal(t, http.StatusSwitchingProtocols, respB.StatusCode)
335+
defer connB.Close(websocket.StatusNormalClosure, "")
336+
require.NoError(t, connB.Write(t.Context(), websocket.MessageText,
337+
[]byte(`{"jsonrpc":"2.0","method":"test_echo","params":["hi"],"id":2}`)))
338+
_, got, err := connB.Read(t.Context())
339+
require.NoError(t, err)
340+
assert.Equal(t,
341+
`{"jsonrpc":"2.0","error":{"code":-32603,"message":"server busy"},"id":null}`,
342+
string(got))
343+
344+
close(release)
345+
_, _, err = connA.Read(t.Context())
346+
require.NoError(t, err)
347+
require.Eventually(t, func() bool { return gate.Running() == 0 }, time.Second, 5*time.Millisecond)
348+
349+
require.NoError(t, connB.Write(t.Context(), websocket.MessageText,
350+
[]byte(`{"jsonrpc":"2.0","method":"test_echo","params":["hi"],"id":3}`)))
351+
_, got, err = connB.Read(t.Context())
352+
require.NoError(t, err)
353+
assert.Equal(t, `{"jsonrpc":"2.0","result":"hi","id":3}`, string(got))
354+
}

node/http.go

Lines changed: 4 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -103,25 +103,13 @@ func makeRPCOverHTTP(
103103
metricsEnabled bool,
104104
corsEnabled bool,
105105
rpcRequestTimeout time.Duration,
106-
maxConcurrentRequests uint,
107-
maxRequestQueue uint,
106+
gate *jsonrpc.Gate,
108107
) *httpService {
109108
var listener jsonrpc.NewRequestListener
110109
if metricsEnabled {
111110
listener = makeHTTPMetrics()
112111
}
113112

114-
// A single gate shared across all RPC servers (v8/v9/v10) so the limit
115-
// protects the whole process, not each version independently. Disabled when
116-
// maxConcurrentRequests is 0.
117-
var gate *jsonrpc.Gate
118-
if maxConcurrentRequests > 0 {
119-
gate = jsonrpc.NewGate(maxConcurrentRequests, uint64(maxRequestQueue))
120-
if metricsEnabled {
121-
makeHTTPGateMetrics(gate)
122-
}
123-
}
124-
125113
mux := http.NewServeMux()
126114
for path, server := range servers {
127115
httpHandler := jsonrpc.NewHTTP(server, logger).
@@ -156,6 +144,7 @@ func makeRPCOverWebsocket(
156144
metricsEnabled bool,
157145
corsEnabled bool,
158146
rpcRequestTimeout time.Duration,
147+
gate *jsonrpc.Gate,
159148
) *httpService {
160149
var listener jsonrpc.NewRequestListener
161150
if metricsEnabled {
@@ -167,7 +156,8 @@ func makeRPCOverWebsocket(
167156
mux := http.NewServeMux()
168157
for path, server := range servers {
169158
wsHandler := jsonrpc.NewWebsocket(server, shutdown, logger).
170-
WithRequestTimeout(rpcRequestTimeout)
159+
WithRequestTimeout(rpcRequestTimeout).
160+
WithGate(gate)
171161
if listener != nil {
172162
wsHandler = wsHandler.WithListener(listener)
173163
}

node/node.go

Lines changed: 9 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -590,6 +590,13 @@ func New(cfg *Config, version string, logLevel *log.Level) (*Node, error) {
590590
"/rpc" + pathV09: jsonrpcServerV09,
591591
"/rpc" + pathV08: jsonrpcServerV08,
592592
}
593+
var rpcGate *jsonrpc.Gate
594+
if cfg.RPCMaxConcurrentRequests > 0 {
595+
rpcGate = jsonrpc.NewGate(cfg.RPCMaxConcurrentRequests, uint64(cfg.RPCMaxRequestQueue))
596+
if cfg.Metrics {
597+
makeHTTPGateMetrics(rpcGate)
598+
}
599+
}
593600
if cfg.HTTP {
594601
readinessHandlers := NewReadinessHandlers(chain, syncReader, cfg.ReadinessBlockTolerance)
595602
httpHandlers := map[string]http.HandlerFunc{
@@ -609,8 +616,7 @@ func New(cfg *Config, version string, logLevel *log.Level) (*Node, error) {
609616
cfg.Metrics,
610617
cfg.RPCCorsEnable,
611618
cfg.RPCRequestTimeout,
612-
cfg.RPCMaxConcurrentRequests,
613-
cfg.RPCMaxRequestQueue,
619+
rpcGate,
614620
),
615621
)
616622
}
@@ -625,6 +631,7 @@ func New(cfg *Config, version string, logLevel *log.Level) (*Node, error) {
625631
cfg.Metrics,
626632
cfg.RPCCorsEnable,
627633
cfg.RPCRequestTimeout,
634+
rpcGate,
628635
),
629636
)
630637
}

rpc/handlers.go

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,7 @@ import (
1919
"github.com/NethermindEth/juno/utils/log"
2020
"github.com/NethermindEth/juno/vm"
2121
"golang.org/x/sync/errgroup"
22+
"golang.org/x/sync/semaphore"
2223
)
2324

2425
const (
@@ -44,9 +45,13 @@ type Handler struct {
4445
func New(bcReader blockchain.Reader, syncReader sync.Reader, virtualMachine vm.VM, version string,
4546
logger log.Logger, network *networks.Network,
4647
) *Handler {
47-
handlerv8 := rpcv8.New(bcReader, syncReader, virtualMachine, logger)
48-
handlerv9 := rpcv9.New(bcReader, syncReader, virtualMachine, logger)
49-
handlerv10 := rpcv10.New(bcReader, syncReader, virtualMachine, logger)
48+
subscriptionLimiter := semaphore.NewWeighted(rpccore.DefaultMaxSubscriptions)
49+
handlerv8 := rpcv8.New(bcReader, syncReader, virtualMachine, logger).
50+
WithSubscriptionLimiter(subscriptionLimiter)
51+
handlerv9 := rpcv9.New(bcReader, syncReader, virtualMachine, logger).
52+
WithSubscriptionLimiter(subscriptionLimiter)
53+
handlerv10 := rpcv10.New(bcReader, syncReader, virtualMachine, logger).
54+
WithSubscriptionLimiter(subscriptionLimiter)
5055

5156
return &Handler{
5257
rpcv8Handler: handlerv8,

rpc/rpccore/rpccore.go

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,8 @@ const (
1818
MaxBlocksBack = 1024
1919
EntrypointNotFoundFelt string = "0x454e545259504f494e545f4e4f545f464f554e44"
2020
ErrEPSNotFound = "Entry point EntryPointSelector(%s) not found in contract."
21+
22+
DefaultMaxSubscriptions int64 = 2048
2123
)
2224

2325
//go:generate mockgen -destination=../mocks/mock_gateway_handler.go -package=mocks github.com/NethermindEth/juno/rpc/rpccore Gateway
@@ -90,4 +92,5 @@ var (
9092

9193
// These errors can be only be returned by Juno-specific methods.
9294
ErrSubscriptionNotFound = &jsonrpc.Error{Code: 100, Message: "Subscription not found"}
95+
ErrTooManySubscriptions = &jsonrpc.Error{Code: 101, Message: "Too many subscriptions"}
9396
)

rpc/v10/handlers.go

Lines changed: 9 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@ import (
2323
"github.com/NethermindEth/juno/utils/lru"
2424
"github.com/NethermindEth/juno/vm"
2525
"github.com/sourcegraph/conc"
26+
"golang.org/x/sync/semaphore"
2627
)
2728

2829
type Handler struct {
@@ -40,8 +41,9 @@ type Handler struct {
4041
l1Heads *feed.Feed[*core.L1Head]
4142
receivedTransactionFeed *feed.Feed[core.Transaction]
4243

43-
idgen func() string
44-
subscriptions stdsync.Map // map[string]*subscription
44+
idgen func() string
45+
subscriptions stdsync.Map // map[string]*subscription
46+
subscriptionLimiter *semaphore.Weighted
4547

4648
blockTraceCache *lru.Cache[felt.Felt, TraceBlockTransactionsResponse]
4749
// todo(rdr): Can this cache be genericified and can it be applied to the `blockTraceCache`
@@ -92,6 +94,11 @@ func New(
9294
}
9395
}
9496

97+
func (h *Handler) WithSubscriptionLimiter(limiter *semaphore.Weighted) *Handler {
98+
h.subscriptionLimiter = limiter
99+
return h
100+
}
101+
95102
func (h *Handler) WithCompiler(compiler compiler.Compiler) *Handler {
96103
h.compiler = compiler
97104
return h

rpc/v10/subscriptions.go

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -54,6 +54,9 @@ func (h *Handler) subscribe(
5454
wsConn jsonrpc.Conn,
5555
subscriber subscriber,
5656
) (SubscriptionID, *jsonrpc.Error) {
57+
if h.subscriptionLimiter != nil && !h.subscriptionLimiter.TryAcquire(1) {
58+
return "", rpccore.ErrTooManySubscriptions
59+
}
5760
id := h.idgen()
5861
//nolint:gosec // G118: cancel called in unsubscribe()
5962
subscriptionCtx, subscriptionCtxCancel := context.WithCancel(wsConn.Context())
@@ -75,6 +78,9 @@ func (h *Handler) subscribe(
7578
)
7679

7780
sub.wg.Go(func() {
81+
if h.subscriptionLimiter != nil {
82+
defer h.subscriptionLimiter.Release(1)
83+
}
7884
defer func() {
7985
h.unsubscribe(sub, id)
8086
unsubscribeFeedSubscription(reorgSub)

rpc/v10/subscriptions_test.go

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,7 @@ import (
3232
"github.com/stretchr/testify/assert"
3333
"github.com/stretchr/testify/require"
3434
"go.uber.org/mock/gomock"
35+
"golang.org/x/sync/semaphore"
3536
)
3637

3738
// mustNewChain builds a pre_confirmed ChainReader from statically valid test
@@ -1214,6 +1215,49 @@ func TestSubscribeNewHeads(t *testing.T) {
12141215
assertNextHead(t, conn, subID, &adaptedHeader3)
12151216
}
12161217

1218+
func TestSubscribeNewHeadsRespectsLimit(t *testing.T) {
1219+
logger := log.NewNopZapLogger()
1220+
client := feeder.NewTestClient(t, &networks.Sepolia)
1221+
block1, commitments1, stateUpdate1 := GetTestBlockWithCommitments(t, client, 56377)
1222+
adaptedHeader := AdaptBlockHeader(block1.Header, commitments1, stateUpdate1.StateDiff)
1223+
1224+
mockCtrl := gomock.NewController(t)
1225+
t.Cleanup(mockCtrl.Finish)
1226+
mockChain := mocks.NewMockReader(mockCtrl)
1227+
1228+
handler := New(mockChain, nil, nil, logger).
1229+
WithSubscriptionLimiter(semaphore.NewWeighted(1))
1230+
1231+
mockChain.EXPECT().Height().Return(block1.Number, nil).Times(3)
1232+
mockChain.EXPECT().BlockHeaderByNumber(block1.Number).Return(block1.Header, nil).Times(2)
1233+
mockChain.EXPECT().BlockCommitmentsByNumber(block1.Number).Return(commitments1, nil).Times(2)
1234+
mockChain.EXPECT().StateUpdateByNumber(block1.Number).Return(stateUpdate1, nil).Times(2)
1235+
1236+
blockIDLatest := BlockIDLatest()
1237+
1238+
subID1, conn1 := createTestNewHeadsWebsocket(t, handler, (*SubscriptionBlockID)(&blockIDLatest))
1239+
assertNextHead(t, conn1, subID1, &adaptedHeader)
1240+
1241+
serverConn, clientConn := net.Pipe()
1242+
t.Cleanup(func() {
1243+
require.NoError(t, serverConn.Close())
1244+
require.NoError(t, clientConn.Close())
1245+
})
1246+
rejConn := &fakeConn{Conn: clientConn, w: serverConn}
1247+
rejCtx := context.WithValue(t.Context(), jsonrpc.ConnKey{}, rejConn)
1248+
id, rpcErr := handler.SubscribeNewHeads(rejCtx, nil)
1249+
assert.Zero(t, id)
1250+
assert.Equal(t, rpccore.ErrTooManySubscriptions, rpcErr)
1251+
1252+
unsubCtx := context.WithValue(t.Context(), jsonrpc.ConnKey{}, conn1)
1253+
ok, rpcErr := handler.Unsubscribe(unsubCtx, string(subID1))
1254+
require.Nil(t, rpcErr)
1255+
require.True(t, ok)
1256+
1257+
subID3, conn3 := createTestNewHeadsWebsocket(t, handler, (*SubscriptionBlockID)(&blockIDLatest))
1258+
assertNextHead(t, conn3, subID3, &adaptedHeader)
1259+
}
1260+
12171261
func TestSubscribeNewHeadsHistorical(t *testing.T) {
12181262
logger := log.NewNopZapLogger()
12191263
client := feeder.NewTestClient(t, &networks.Sepolia)

0 commit comments

Comments
 (0)