diff --git a/cmd/juno/juno.go b/cmd/juno/juno.go index 086ffc8c7a..e707c8c627 100644 --- a/cmd/juno/juno.go +++ b/cmd/juno/juno.go @@ -110,6 +110,8 @@ const ( rpcRequestTimeoutF = "rpc-request-timeout" rpcMaxConcurrentRequestsF = "rpc-max-concurrent-requests" rpcMaxRequestQueueF = "rpc-max-request-queue" + rpcMaxWSConnectionsF = "rpc-max-ws-connections" + rpcMaxSubscriptionsF = "rpc-max-subscriptions" rpcMaxBatchSizeF = "rpc-max-batch-size" rpcMaxBatchResponseSizeF = "rpc-max-batch-response-size" rpcBatchConcurrencyF = "rpc-batch-concurrency" @@ -182,6 +184,8 @@ const ( defaultRPCRequestTimeout = 1 * time.Minute defaultRPCMaxConcurrentRequests = 256000 defaultRPCMaxQueuedRequests = 256000 + defaultRPCMaxWSConnections = 1024 + defaultRPCMaxSubscriptions = 128 defaultRPCMaxBatchSize = 1000 defaultRPCMaxBatchResponseSize = 64 // MB defaultRPCBatchConcurrency = uint(0) @@ -278,6 +282,12 @@ const ( "together; 0 disables the limit." rpcMaxRequestQueueUsage = "Maximum number of HTTP RPC requests to queue after " + "reaching rpc-max-concurrent-requests limit. Websocket requests are never queued." + rpcMaxWSConnectionsUsage = "Maximum concurrent websocket connections, across all " + + "RPC versions. 0 disables the limit." + rpcMaxSubscriptionsUsage = "Maximum subscriptions one websocket connection may hold. " + + "The limit is per connection so that one client cannot take the slots another needs; " + + "multiplied by rpc-max-ws-connections it is the most the process will carry. " + + "0 disables the limit." rpcMaxBatchSizeUsage = "Maximum number of calls in a single batch request. " + "0 disables the limit." rpcMaxBatchResponseSizeUsage = "Size (in MBs) at which a batch stops being processed. " + @@ -528,6 +538,8 @@ func NewCmd(config *node.Config, run func(*cobra.Command, []string) error) *cobr defaultRPCMaxQueuedRequests, rpcMaxRequestQueueUsage, ) + junoCmd.Flags().Uint(rpcMaxWSConnectionsF, defaultRPCMaxWSConnections, rpcMaxWSConnectionsUsage) + junoCmd.Flags().Uint(rpcMaxSubscriptionsF, defaultRPCMaxSubscriptions, rpcMaxSubscriptionsUsage) junoCmd.Flags().Uint(rpcMaxBatchSizeF, defaultRPCMaxBatchSize, rpcMaxBatchSizeUsage) junoCmd.Flags().Uint( rpcMaxBatchResponseSizeF, @@ -553,6 +565,8 @@ func NewCmd(config *node.Config, run func(*cobra.Command, []string) error) *cobr rpcRequestTimeoutF, rpcMaxConcurrentRequestsF, rpcMaxRequestQueueF, + rpcMaxWSConnectionsF, + rpcMaxSubscriptionsF, rpcMaxBatchSizeF, rpcMaxBatchResponseSizeF, rpcBatchConcurrencyF, diff --git a/cmd/juno/juno_test.go b/cmd/juno/juno_test.go index dce0cf3169..664c8132a0 100644 --- a/cmd/juno/juno_test.go +++ b/cmd/juno/juno_test.go @@ -72,6 +72,8 @@ func TestConfigPrecedence(t *testing.T) { defaultMaxVMs := uint(3 * runtime.GOMAXPROCS(0)) defaultRPCMaxConcurrentRequests := uint(256000) defaultRPCMaxRequestQueue := uint(256000) + defaultRPCMaxWSConnections := uint(1024) + defaultRPCMaxSubscriptions := uint(128) defaultRPCMaxBatchSize := uint(1000) defaultRPCMaxBatchResponseSize := uint(64) defaultRPCMaxBlockScan := uint(math.MaxUint) @@ -125,6 +127,8 @@ func TestConfigPrecedence(t *testing.T) { MaxVMQueue: 2 * defaultMaxVMs, RPCMaxConcurrentRequests: defaultRPCMaxConcurrentRequests, RPCMaxRequestQueue: defaultRPCMaxRequestQueue, + RPCMaxWSConnections: defaultRPCMaxWSConnections, + RPCMaxSubscriptions: defaultRPCMaxSubscriptions, RPCMaxBatchSize: defaultRPCMaxBatchSize, RPCMaxBatchResponseSize: defaultRPCMaxBatchResponseSize, RPCMaxBlockScan: defaultRPCMaxBlockScan, @@ -177,6 +181,8 @@ func TestConfigPrecedence(t *testing.T) { MaxVMQueue: 2 * defaultMaxVMs, RPCMaxConcurrentRequests: defaultRPCMaxConcurrentRequests, RPCMaxRequestQueue: defaultRPCMaxRequestQueue, + RPCMaxWSConnections: defaultRPCMaxWSConnections, + RPCMaxSubscriptions: defaultRPCMaxSubscriptions, RPCMaxBatchSize: defaultRPCMaxBatchSize, RPCMaxBatchResponseSize: defaultRPCMaxBatchResponseSize, RPCMaxBlockScan: defaultRPCMaxBlockScan, @@ -338,6 +344,8 @@ pprof: true MaxVMQueue: 2 * defaultMaxVMs, RPCMaxConcurrentRequests: defaultRPCMaxConcurrentRequests, RPCMaxRequestQueue: defaultRPCMaxRequestQueue, + RPCMaxWSConnections: defaultRPCMaxWSConnections, + RPCMaxSubscriptions: defaultRPCMaxSubscriptions, RPCMaxBatchSize: defaultRPCMaxBatchSize, RPCMaxBatchResponseSize: defaultRPCMaxBatchResponseSize, RPCMaxBlockScan: defaultRPCMaxBlockScan, @@ -396,6 +404,8 @@ http-port: 4576 MaxVMQueue: 2 * defaultMaxVMs, RPCMaxConcurrentRequests: defaultRPCMaxConcurrentRequests, RPCMaxRequestQueue: defaultRPCMaxRequestQueue, + RPCMaxWSConnections: defaultRPCMaxWSConnections, + RPCMaxSubscriptions: defaultRPCMaxSubscriptions, RPCMaxBatchSize: defaultRPCMaxBatchSize, RPCMaxBatchResponseSize: defaultRPCMaxBatchResponseSize, RPCMaxBlockScan: defaultRPCMaxBlockScan, @@ -452,6 +462,8 @@ http-port: 4576 MaxVMQueue: 2 * defaultMaxVMs, RPCMaxConcurrentRequests: defaultRPCMaxConcurrentRequests, RPCMaxRequestQueue: defaultRPCMaxRequestQueue, + RPCMaxWSConnections: defaultRPCMaxWSConnections, + RPCMaxSubscriptions: defaultRPCMaxSubscriptions, RPCMaxBatchSize: defaultRPCMaxBatchSize, RPCMaxBatchResponseSize: defaultRPCMaxBatchResponseSize, RPCMaxBlockScan: defaultRPCMaxBlockScan, @@ -510,6 +522,8 @@ http-port: 4576 MaxVMQueue: 2 * defaultMaxVMs, RPCMaxConcurrentRequests: defaultRPCMaxConcurrentRequests, RPCMaxRequestQueue: defaultRPCMaxRequestQueue, + RPCMaxWSConnections: defaultRPCMaxWSConnections, + RPCMaxSubscriptions: defaultRPCMaxSubscriptions, RPCMaxBatchSize: defaultRPCMaxBatchSize, RPCMaxBatchResponseSize: defaultRPCMaxBatchResponseSize, RPCMaxBlockScan: defaultRPCMaxBlockScan, @@ -591,6 +605,8 @@ db-cache-size: 1024 MaxVMQueue: 2 * defaultMaxVMs, RPCMaxConcurrentRequests: defaultRPCMaxConcurrentRequests, RPCMaxRequestQueue: defaultRPCMaxRequestQueue, + RPCMaxWSConnections: defaultRPCMaxWSConnections, + RPCMaxSubscriptions: defaultRPCMaxSubscriptions, RPCMaxBatchSize: defaultRPCMaxBatchSize, RPCMaxBatchResponseSize: defaultRPCMaxBatchResponseSize, RPCMaxBlockScan: defaultRPCMaxBlockScan, @@ -651,6 +667,8 @@ network: sepolia MaxVMQueue: 2 * defaultMaxVMs, RPCMaxConcurrentRequests: defaultRPCMaxConcurrentRequests, RPCMaxRequestQueue: defaultRPCMaxRequestQueue, + RPCMaxWSConnections: defaultRPCMaxWSConnections, + RPCMaxSubscriptions: defaultRPCMaxSubscriptions, RPCMaxBatchSize: defaultRPCMaxBatchSize, RPCMaxBatchResponseSize: defaultRPCMaxBatchResponseSize, RPCMaxBlockScan: defaultRPCMaxBlockScan, @@ -707,6 +725,8 @@ network: sepolia MaxVMQueue: 2 * defaultMaxVMs, RPCMaxConcurrentRequests: defaultRPCMaxConcurrentRequests, RPCMaxRequestQueue: defaultRPCMaxRequestQueue, + RPCMaxWSConnections: defaultRPCMaxWSConnections, + RPCMaxSubscriptions: defaultRPCMaxSubscriptions, RPCMaxBatchSize: defaultRPCMaxBatchSize, RPCMaxBatchResponseSize: defaultRPCMaxBatchResponseSize, RPCMaxBlockScan: defaultRPCMaxBlockScan, @@ -761,6 +781,8 @@ network: sepolia MaxVMQueue: 2 * defaultMaxVMs, RPCMaxConcurrentRequests: defaultRPCMaxConcurrentRequests, RPCMaxRequestQueue: defaultRPCMaxRequestQueue, + RPCMaxWSConnections: defaultRPCMaxWSConnections, + RPCMaxSubscriptions: defaultRPCMaxSubscriptions, RPCMaxBatchSize: defaultRPCMaxBatchSize, RPCMaxBatchResponseSize: defaultRPCMaxBatchResponseSize, RPCMaxBlockScan: defaultRPCMaxBlockScan, @@ -816,6 +838,8 @@ network: sepolia MaxVMQueue: 2 * defaultMaxVMs, RPCMaxConcurrentRequests: defaultRPCMaxConcurrentRequests, RPCMaxRequestQueue: defaultRPCMaxRequestQueue, + RPCMaxWSConnections: defaultRPCMaxWSConnections, + RPCMaxSubscriptions: defaultRPCMaxSubscriptions, RPCMaxBatchSize: defaultRPCMaxBatchSize, RPCMaxBatchResponseSize: defaultRPCMaxBatchResponseSize, RPCMaxBlockScan: defaultRPCMaxBlockScan, @@ -871,6 +895,8 @@ network: sepolia MaxVMQueue: 2 * defaultMaxVMs, RPCMaxConcurrentRequests: defaultRPCMaxConcurrentRequests, RPCMaxRequestQueue: defaultRPCMaxRequestQueue, + RPCMaxWSConnections: defaultRPCMaxWSConnections, + RPCMaxSubscriptions: defaultRPCMaxSubscriptions, RPCMaxBatchSize: defaultRPCMaxBatchSize, RPCMaxBatchResponseSize: defaultRPCMaxBatchResponseSize, RPCMaxBlockScan: defaultRPCMaxBlockScan, @@ -925,6 +951,8 @@ network: sepolia MaxVMQueue: 2 * defaultMaxVMs, RPCMaxConcurrentRequests: defaultRPCMaxConcurrentRequests, RPCMaxRequestQueue: defaultRPCMaxRequestQueue, + RPCMaxWSConnections: defaultRPCMaxWSConnections, + RPCMaxSubscriptions: defaultRPCMaxSubscriptions, RPCMaxBatchSize: defaultRPCMaxBatchSize, RPCMaxBatchResponseSize: defaultRPCMaxBatchResponseSize, RPCMaxBlockScan: defaultRPCMaxBlockScan, diff --git a/jsonrpc/server.go b/jsonrpc/server.go index 6392002622..30729015de 100644 --- a/jsonrpc/server.go +++ b/jsonrpc/server.go @@ -287,10 +287,12 @@ type Conn interface { io.Writer Equal(Conn) bool Context() context.Context + SubscriptionSlots } type connection struct { w io.Writer + slots SubscriptionSlots activated <-chan struct{} ctx context.Context @@ -320,6 +322,28 @@ func (c *connection) Context() context.Context { return c.ctx } +// SubscriptionSlots caps how many subscriptions one connection carries at once. +type SubscriptionSlots interface { + // TryAcquireSubscription reserves a slot without waiting + TryAcquireSubscription() bool + // ReleaseSubscription returns a slot taken by TryAcquireSubscription + ReleaseSubscription() +} + +// Transport is what HandleReadWriter needs of a connection +type Transport interface { + io.ReadWriter + SubscriptionSlots +} + +func (c *connection) TryAcquireSubscription() bool { + return c.slots.TryAcquireSubscription() +} + +func (c *connection) ReleaseSubscription() { + c.slots.ReleaseSubscription() +} + // ConnKey the key used to retrieve the connection from the context passed to a handler. // It is exported to allow transports to set it manually if they decide not to use HandleReadWriter, which sets it automatically. // Manually setting the connection can be especially useful when testing handlers. @@ -343,12 +367,13 @@ func ConnFromContext(ctx context.Context) (Conn, bool) { func (s *Server) HandleReadWriter( connCtx context.Context, requestTimeout time.Duration, - rw io.ReadWriter, + rw Transport, ) error { activated := make(chan struct{}) defer close(activated) conn := &connection{ - w: rw.(io.Writer), + w: rw, + slots: rw, activated: activated, ctx: connCtx, } diff --git a/jsonrpc/server_test.go b/jsonrpc/server_test.go index 4c7167fe2b..a0f15cbc6d 100644 --- a/jsonrpc/server_test.go +++ b/jsonrpc/server_test.go @@ -710,6 +710,17 @@ func TestCannotWriteToConnInHandler(t *testing.T) { require.NotNil(t, header) } +// uncappedTransport hands HandleReadWriter a stream with no subscription cap, +// which these tests have no use for. The cap is part of jsonrpc.Transport so +// that a real transport cannot forget it, and saying so here is the price. +type uncappedTransport struct { + net.Conn +} + +func (uncappedTransport) TryAcquireSubscription() bool { return true } + +func (uncappedTransport) ReleaseSubscription() {} + type fakeConn struct { ctx context.Context } @@ -726,6 +737,10 @@ func (fc *fakeConn) Context() context.Context { return fc.ctx } +func (fc *fakeConn) TryAcquireSubscription() bool { return true } + +func (fc *fakeConn) ReleaseSubscription() {} + func TestWriteToConnInHandler(t *testing.T) { testBytes := "written from handler" server := jsonrpc.NewServer(1, log.NewNopZapLogger()) @@ -754,7 +769,7 @@ func TestWriteToConnInHandler(t *testing.T) { }) wg.Go(func() { - err := server.HandleReadWriter(t.Context(), 0, serverConn) + err := server.HandleReadWriter(t.Context(), 0, uncappedTransport{serverConn}) require.NoError(t, err) }) @@ -791,7 +806,7 @@ func TestWriteToClosedConnInHandler(t *testing.T) { }) wg.Go(func() { - err := server.HandleReadWriter(t.Context(), 0, serverConn) + err := server.HandleReadWriter(t.Context(), 0, uncappedTransport{serverConn}) require.ErrorIs(t, err, io.ErrClosedPipe) }) diff --git a/jsonrpc/websocket.go b/jsonrpc/websocket.go index 5b2c33976a..f2825b055d 100644 --- a/jsonrpc/websocket.go +++ b/jsonrpc/websocket.go @@ -7,6 +7,7 @@ import ( "io" "net/http" "strings" + "sync/atomic" "time" "github.com/NethermindEth/juno/db" @@ -16,10 +17,7 @@ import ( "golang.org/x/sync/semaphore" ) -const ( - closeReasonMaxBytes = 125 - maxConns = 2048 // TODO: an arbitrary default number, should be revisited after monitoring -) +const closeReasonMaxBytes = 125 var serverBusyResponse = func() []byte { b, err := json.Marshal(&response{ @@ -43,8 +41,10 @@ type Websocket struct { requestTimeout time.Duration gate *Gate - // Add connection tracking + // connSem bounds concurrent connections connSem *semaphore.Weighted + // maxSubscriptions caps every connection separately + maxSubscriptions int64 } func NewWebsocket(rpc *Server, shutdown <-chan struct{}, logger log.StructuredLogger) *Websocket { @@ -57,15 +57,22 @@ func NewWebsocket(rpc *Server, shutdown <-chan struct{}, logger log.StructuredLo connParams: DefaultWebsocketConnParams(), listener: &SelectiveListener{}, shutdown: shutdown, - connSem: semaphore.NewWeighted(maxConns), } return ws } -// WithMaxConnections sets the maximum number of concurrent websocket connections -func (ws *Websocket) WithMaxConnections(maxConns int64) *Websocket { - ws.connSem = semaphore.NewWeighted(maxConns) +// WithConnLimiter sets the semaphore that bounds concurrent websocket +// connections. nil leaves them unbounded. +func (ws *Websocket) WithConnLimiter(sem *semaphore.Weighted) *Websocket { + ws.connSem = sem + return ws +} + +// WithMaxSubscriptions sets how many subscriptions one connection may hold at +// once; zero or less leaves them unbounded. +func (ws *Websocket) WithMaxSubscriptions(n int64) *Websocket { + ws.maxSubscriptions = n return ws } @@ -92,15 +99,17 @@ func (ws *Websocket) WithListener(listener NewRequestListener) *Websocket { return ws } -// ServeHTTP processes an HTTP request and upgrades it to a websocket connection. -// The connection's entire "lifetime" is spent in this function. -func (ws *Websocket) ServeHTTP(w http.ResponseWriter, r *http.Request) { +// acquireConnSlot takes a connection slot +func (ws *Websocket) acquireConnSlot(ctx context.Context, w http.ResponseWriter) bool { + if ws.connSem == nil { + return true + } + // Create a timeout context for the acquisition const connTimeout = 5 * time.Second - acquireCtx, cancel := context.WithTimeout(r.Context(), connTimeout) + acquireCtx, cancel := context.WithTimeout(ctx, connTimeout) defer cancel() - // Check connection limit if err := ws.connSem.Acquire(acquireCtx, 1); err != nil { if errors.Is(err, context.DeadlineExceeded) { ws.logger.Warn("Connection request timed out while waiting for slot") @@ -108,9 +117,25 @@ func (ws *Websocket) ServeHTTP(w http.ResponseWriter, r *http.Request) { } else { ws.logger.Warn("Connection request was canceled while waiting for slot") } + return false + } + + return true +} + +func (ws *Websocket) releaseConnSlot() { + if ws.connSem != nil { + ws.connSem.Release(1) + } +} + +// ServeHTTP processes an HTTP request and upgrades it to a websocket connection. +// The connection's entire "lifetime" is spent in this function. +func (ws *Websocket) ServeHTTP(w http.ResponseWriter, r *http.Request) { + if !ws.acquireConnSlot(r.Context(), w) { return } - defer ws.connSem.Release(1) + defer ws.releaseConnSlot() conn, err := websocket.Accept(w, r, nil /* TODO: options */) if err != nil { @@ -132,7 +157,7 @@ func (ws *Websocket) ServeHTTP(w http.ResponseWriter, r *http.Request) { } }() - wsc := newWebsocketConn(ctx, conn, ws.connParams) + wsc := newWebsocketConn(ctx, conn, ws.connParams, ws.maxSubscriptions) for { _, wsc.r, err = wsc.conn.Reader(wsc.ctx) @@ -210,17 +235,45 @@ type websocketConn struct { conn *websocket.Conn ctx context.Context params *WebsocketConnParams + + subscriptions atomic.Int64 + maxSubscriptions int64 } -func newWebsocketConn(ctx context.Context, conn *websocket.Conn, params *WebsocketConnParams) *websocketConn { +var _ SubscriptionSlots = (*websocketConn)(nil) + +func newWebsocketConn( + ctx context.Context, + conn *websocket.Conn, + params *WebsocketConnParams, + maxSubscriptions int64, +) *websocketConn { conn.SetReadLimit(params.ReadLimit) return &websocketConn{ - conn: conn, - ctx: ctx, - params: params, + conn: conn, + ctx: ctx, + params: params, + maxSubscriptions: maxSubscriptions, + } +} + +// TryAcquireSubscription takes a slot if the connection is below its limit +func (wsc *websocketConn) TryAcquireSubscription() bool { + for { + held := wsc.subscriptions.Load() + if wsc.maxSubscriptions > 0 && held >= wsc.maxSubscriptions { + return false + } + if wsc.subscriptions.CompareAndSwap(held, held+1) { + return true + } } } +func (wsc *websocketConn) ReleaseSubscription() { + wsc.subscriptions.Add(-1) +} + func (wsc *websocketConn) Read(p []byte) (int, error) { return wsc.r.Read(p) } diff --git a/jsonrpc/websocket_test.go b/jsonrpc/websocket_test.go index adc3df13d2..0473da169a 100644 --- a/jsonrpc/websocket_test.go +++ b/jsonrpc/websocket_test.go @@ -3,6 +3,7 @@ package jsonrpc_test import ( "context" "encoding/json" + "fmt" "net/http" "net/http/httptest" "testing" @@ -14,6 +15,7 @@ import ( "github.com/sourcegraph/conc" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "golang.org/x/sync/semaphore" ) // The caller is responsible for closing the connection. @@ -266,7 +268,8 @@ func TestWebsocketConnectionLimit(t *testing.T) { t.Parallel() rpc := jsonrpc.NewServer(1, log.NewNopZapLogger()) - ws := jsonrpc.NewWebsocket(rpc, nil, log.NewNopZapLogger()).WithMaxConnections(2) + ws := jsonrpc.NewWebsocket(rpc, nil, log.NewNopZapLogger()). + WithConnLimiter(semaphore.NewWeighted(2)) httpSrv := httptest.NewServer(ws) defer httpSrv.Close() @@ -352,3 +355,85 @@ func TestWebsocketGateRejectsWhenBusy(t *testing.T) { require.NoError(t, err) assert.Equal(t, `{"jsonrpc":"2.0","result":"hi","id":3}`, string(got)) } + +func TestWebsocketSubscriptionSlotsArePerConnection(t *testing.T) { + const maxSubs = 2 + + connOf := func(ctx context.Context) (jsonrpc.Conn, *jsonrpc.Error) { + conn, ok := jsonrpc.ConnFromContext(ctx) + if !ok { + return nil, &jsonrpc.Error{Code: 1, Message: "no connection in context"} + } + return conn, nil + } + + subscribe := jsonrpc.Method{ + Name: "test_subscribe", + Handler: func(ctx context.Context) (string, *jsonrpc.Error) { + conn, rpcErr := connOf(ctx) + if rpcErr != nil { + return "", rpcErr + } + if !conn.TryAcquireSubscription() { + return "", &jsonrpc.Error{Code: 101, Message: "Too many subscriptions"} + } + return "subscribed", nil + }, + } + unsubscribe := jsonrpc.Method{ + Name: "test_unsubscribe", + Handler: func(ctx context.Context) (string, *jsonrpc.Error) { + conn, rpcErr := connOf(ctx) + if rpcErr != nil { + return "", rpcErr + } + conn.ReleaseSubscription() + return "unsubscribed", nil + }, + } + + rpc := jsonrpc.NewServer(1, log.NewNopZapLogger()) + require.NoError(t, rpc.RegisterMethods(subscribe, unsubscribe)) + ws := jsonrpc.NewWebsocket(rpc, nil, log.NewNopZapLogger()). + WithMaxSubscriptions(maxSubs) + srv := httptest.NewServer(ws) + t.Cleanup(srv.Close) + + dial := func() *websocket.Conn { + conn, resp, err := websocket.Dial(t.Context(), srv.URL, nil) //nolint:bodyclose // lib closes it + require.NoError(t, err) + require.Equal(t, http.StatusSwitchingProtocols, resp.StatusCode) + t.Cleanup(func() { conn.Close(websocket.StatusNormalClosure, "") }) + return conn + } + call := func(conn *websocket.Conn, method string, id int) string { + require.NoError(t, conn.Write(t.Context(), websocket.MessageText, + []byte(fmt.Sprintf(`{"jsonrpc":"2.0","method":%q,"id":%d}`, method, id)))) + _, got, err := conn.Read(t.Context()) + require.NoError(t, err) + return string(got) + } + + const ( + subscribed = `{"jsonrpc":"2.0","result":"subscribed","id":%d}` + tooMany = `{"jsonrpc":"2.0","error":{"code":101,"message":"Too many subscriptions"},"id":%d}` + ) + + connA, connB := dial(), dial() + + // Interleaved on purpose: if the budget were shared, connB's second call would + // be the fifth overall and would already be refused. + for i := 1; i <= maxSubs; i++ { + assert.JSONEq(t, fmt.Sprintf(subscribed, i), call(connA, "test_subscribe", i)) + assert.JSONEq(t, fmt.Sprintf(subscribed, i), call(connB, "test_subscribe", i)) + } + + assert.JSONEq(t, fmt.Sprintf(tooMany, 3), call(connA, "test_subscribe", 3)) + assert.JSONEq(t, fmt.Sprintf(tooMany, 3), call(connB, "test_subscribe", 3)) + + // A frees one of its own. B stays full, which is the isolation the test is for. + assert.JSONEq(t, `{"jsonrpc":"2.0","result":"unsubscribed","id":4}`, + call(connA, "test_unsubscribe", 4)) + assert.JSONEq(t, fmt.Sprintf(subscribed, 5), call(connA, "test_subscribe", 5)) + assert.JSONEq(t, fmt.Sprintf(tooMany, 5), call(connB, "test_subscribe", 5)) +} diff --git a/node/http.go b/node/http.go index 705d09439a..08c9ddefd6 100644 --- a/node/http.go +++ b/node/http.go @@ -24,6 +24,7 @@ import ( "github.com/prometheus/client_golang/prometheus/promhttp" "github.com/rs/cors" "github.com/sourcegraph/conc" + "golang.org/x/sync/semaphore" "google.golang.org/grpc" ) @@ -145,6 +146,8 @@ func makeRPCOverWebsocket( corsEnabled bool, rpcRequestTimeout time.Duration, gate *jsonrpc.Gate, + maxConns uint, + maxSubscriptions uint, ) *httpService { var listener jsonrpc.NewRequestListener if metricsEnabled { @@ -153,11 +156,20 @@ func makeRPCOverWebsocket( shutdown := make(chan struct{}) + // One connection limiter shared by the per-version handlers below. Zero means + // no limit. + var connLimiter *semaphore.Weighted + if maxConns > 0 { + connLimiter = semaphore.NewWeighted(int64(maxConns)) + } + mux := http.NewServeMux() for path, server := range servers { wsHandler := jsonrpc.NewWebsocket(server, shutdown, logger). WithRequestTimeout(rpcRequestTimeout). - WithGate(gate) + WithGate(gate). + WithConnLimiter(connLimiter). + WithMaxSubscriptions(int64(maxSubscriptions)) if listener != nil { wsHandler = wsHandler.WithListener(listener) } diff --git a/node/node.go b/node/node.go index f9a9ffe0f1..6c19c11599 100644 --- a/node/node.go +++ b/node/node.go @@ -136,6 +136,8 @@ type Config struct { RPCRequestTimeout time.Duration `mapstructure:"rpc-request-timeout"` RPCMaxConcurrentRequests uint `mapstructure:"rpc-max-concurrent-requests"` RPCMaxRequestQueue uint `mapstructure:"rpc-max-request-queue"` + RPCMaxWSConnections uint `mapstructure:"rpc-max-ws-connections"` + RPCMaxSubscriptions uint `mapstructure:"rpc-max-subscriptions"` RPCMaxBatchSize uint `mapstructure:"rpc-max-batch-size"` RPCMaxBatchResponseSize uint `mapstructure:"rpc-max-batch-response-size"` RPCBatchConcurrency uint `mapstructure:"rpc-batch-concurrency"` @@ -632,6 +634,8 @@ func New(cfg *Config, version string, logLevel *log.Level) (*Node, error) { cfg.RPCCorsEnable, cfg.RPCRequestTimeout, rpcGate, + cfg.RPCMaxWSConnections, + cfg.RPCMaxSubscriptions, ), ) } diff --git a/rpc/rpccore/rpccore.go b/rpc/rpccore/rpccore.go index 68bfd472ea..0a3dfbd09b 100644 --- a/rpc/rpccore/rpccore.go +++ b/rpc/rpccore/rpccore.go @@ -90,4 +90,5 @@ var ( // These errors can be only be returned by Juno-specific methods. ErrSubscriptionNotFound = &jsonrpc.Error{Code: 100, Message: "Subscription not found"} + ErrTooManySubscriptions = &jsonrpc.Error{Code: 101, Message: "Too many subscriptions"} ) diff --git a/rpc/v10/subscriptions.go b/rpc/v10/subscriptions.go index af97614425..7ed44f8368 100644 --- a/rpc/v10/subscriptions.go +++ b/rpc/v10/subscriptions.go @@ -54,6 +54,10 @@ func (h *Handler) subscribe( wsConn jsonrpc.Conn, subscriber subscriber, ) (SubscriptionID, *jsonrpc.Error) { + if !wsConn.TryAcquireSubscription() { + return "", rpccore.ErrTooManySubscriptions + } + id := h.idgen() //nolint:gosec // G118: cancel called in unsubscribe() subscriptionCtx, subscriptionCtxCancel := context.WithCancel(wsConn.Context()) @@ -75,6 +79,7 @@ func (h *Handler) subscribe( ) sub.wg.Go(func() { + defer wsConn.ReleaseSubscription() defer func() { h.unsubscribe(sub, id) unsubscribeFeedSubscription(reorgSub) diff --git a/rpc/v10/subscriptions_test.go b/rpc/v10/subscriptions_test.go index 59a5e0ee21..09c355ccd5 100644 --- a/rpc/v10/subscriptions_test.go +++ b/rpc/v10/subscriptions_test.go @@ -10,6 +10,7 @@ import ( "slices" "strconv" stdsync "sync" + "sync/atomic" "testing" "time" @@ -65,6 +66,10 @@ func (fc *fakeConn) Context() context.Context { return fc.ctx } +func (fc *fakeConn) TryAcquireSubscription() bool { return true } + +func (fc *fakeConn) ReleaseSubscription() {} + type fakeSyncer struct { newHeads *feed.Feed[*core.Block] reorgs *feed.Feed[*sync.ReorgBlockRange] @@ -2884,3 +2889,88 @@ func GetTestBlockWithCommitments( return adaptedBlock, commitments, adaptedState } + +// limitedConn is a fakeConn that also caps its subscriptions, which is what a +// real websocket connection does. Equal is redeclared because fakeConn's asserts +// on *fakeConn and would not recognise this type, and Unsubscribe compares the +// stored connection to the caller's. +type limitedConn struct { + *fakeConn + held atomic.Int64 + max int64 +} + +func (lc *limitedConn) Equal(other jsonrpc.Conn) bool { + other2, ok := other.(*limitedConn) + if !ok { + return false + } + return lc.w == other2.w +} + +func (lc *limitedConn) TryAcquireSubscription() bool { + for { + held := lc.held.Load() + if held >= lc.max { + return false + } + if lc.held.CompareAndSwap(held, held+1) { + return true + } + } +} + +func (lc *limitedConn) ReleaseSubscription() { + lc.held.Add(-1) +} + +func TestSubscribeRespectsConnectionLimit(t *testing.T) { + logger := log.NewNopZapLogger() + client := feeder.NewTestClient(t, &networks.Sepolia) + block1, commitments1, _ := GetTestBlockWithCommitments(t, client, 56377) + adaptedHeader := AdaptBlockHeader(block1.Header, commitments1) + + mockCtrl := gomock.NewController(t) + t.Cleanup(mockCtrl.Finish) + mockChain := mocks.NewMockReader(mockCtrl) + handler := New(mockChain, nil, nil, logger) + + mockChain.EXPECT().Height().Return(block1.Number, nil).Times(3) + mockChain.EXPECT().BlockHeaderByNumber(block1.Number).Return(block1.Header, nil).Times(2) + mockChain.EXPECT().BlockCommitmentsByNumber(block1.Number).Return(commitments1, nil).Times(2) + + serverConn, clientConn := net.Pipe() + t.Cleanup(func() { + require.NoError(t, serverConn.Close()) + require.NoError(t, clientConn.Close()) + }) + + connCtx, cancel := context.WithCancel(t.Context()) + t.Cleanup(cancel) + conn := &limitedConn{ + fakeConn: &fakeConn{Conn: clientConn, w: serverConn, ctx: connCtx}, + max: 1, + } + ctx := context.WithValue(connCtx, jsonrpc.ConnKey{}, conn) + + blockIDLatest := BlockIDLatest() + subBlockID := (*SubscriptionBlockID)(&blockIDLatest) + + subID, rpcErr := handler.SubscribeNewHeads(ctx, subBlockID) + require.Nil(t, rpcErr) + assertNextHead(t, clientConn, subID, &adaptedHeader) + assert.Equal(t, int64(1), conn.held.Load()) + + _, rpcErr = handler.SubscribeNewHeads(ctx, subBlockID) + assert.Equal(t, rpccore.ErrTooManySubscriptions, rpcErr) + assert.Equal(t, int64(1), conn.held.Load(), "a refused call must not take a slot") + + ok, rpcErr := handler.Unsubscribe(ctx, string(subID)) + require.Nil(t, rpcErr) + require.True(t, ok) + assert.Equal(t, int64(0), conn.held.Load(), "the slot is free once Unsubscribe returns") + + subID, rpcErr = handler.SubscribeNewHeads(ctx, subBlockID) + require.Nil(t, rpcErr) + assertNextHead(t, clientConn, subID, &adaptedHeader) +} diff --git a/rpc/v8/subscriptions.go b/rpc/v8/subscriptions.go index 85bffdbb98..6808fc8182 100644 --- a/rpc/v8/subscriptions.go +++ b/rpc/v8/subscriptions.go @@ -118,6 +118,10 @@ func (h *Handler) subscribe( wsConn jsonrpc.Conn, subscriber subscriber, ) (SubscriptionID, *jsonrpc.Error) { + if !wsConn.TryAcquireSubscription() { + return "", rpccore.ErrTooManySubscriptions + } + id := h.idgen() //nolint:gosec // G118: cancel called in unsubscribe() subscriptionCtx, subscriptionCtxCancel := context.WithCancel(wsConn.Context()) @@ -132,6 +136,7 @@ func (h *Handler) subscribe( l1HeadSub, l1HeadRecv := getSubscription(subscriber.onL1Head, h.l1Heads) sub.wg.Go(func() { + defer wsConn.ReleaseSubscription() defer func() { h.unsubscribe(sub, id) unsubscribeFeedSubscription(reorgSub) diff --git a/rpc/v8/subscriptions_test.go b/rpc/v8/subscriptions_test.go index f013878035..e6ba1b39a5 100644 --- a/rpc/v8/subscriptions_test.go +++ b/rpc/v8/subscriptions_test.go @@ -58,6 +58,10 @@ func (fc *fakeConn) Context() context.Context { return fc.ctx } +func (fc *fakeConn) TryAcquireSubscription() bool { return true } + +func (fc *fakeConn) ReleaseSubscription() {} + func TestSubscribeEvents(t *testing.T) { logger := log.NewNopZapLogger() diff --git a/rpc/v9/subscriptions.go b/rpc/v9/subscriptions.go index abdd0597bd..00ce449506 100644 --- a/rpc/v9/subscriptions.go +++ b/rpc/v9/subscriptions.go @@ -111,6 +111,10 @@ func (h *Handler) subscribe( wsConn jsonrpc.Conn, subscriber subscriber, ) (SubscriptionID, *jsonrpc.Error) { + if !wsConn.TryAcquireSubscription() { + return "", rpccore.ErrTooManySubscriptions + } + id := h.idgen() //nolint:gosec // G118: cancel called in unsubscribe() subscriptionCtx, subscriptionCtxCancel := context.WithCancel(wsConn.Context()) @@ -130,6 +134,7 @@ func (h *Handler) subscribe( ) sub.wg.Go(func() { + defer wsConn.ReleaseSubscription() defer func() { h.unsubscribe(sub, id) unsubscribeFeedSubscription(reorgSub) diff --git a/rpc/v9/subscriptions_test.go b/rpc/v9/subscriptions_test.go index 2918b3805f..72a63654b9 100644 --- a/rpc/v9/subscriptions_test.go +++ b/rpc/v9/subscriptions_test.go @@ -66,6 +66,10 @@ func (fc *fakeConn) Context() context.Context { return fc.ctx } +func (fc *fakeConn) TryAcquireSubscription() bool { return true } + +func (fc *fakeConn) ReleaseSubscription() {} + type fakeSyncer struct { newHeads *feed.Feed[*core.Block] reorgs *feed.Feed[*sync.ReorgBlockRange]