@@ -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+ }
0 commit comments