rpc: rename wsMessageSizeLimit (it's now a default), minor test-changes

This commit is contained in:
Martin Holst Swende 2023-09-05 09:25:35 +02:00
parent e28e681632
commit d6a104c562
No known key found for this signature in database
GPG key ID: 683B438C05A5DDF0
3 changed files with 38 additions and 26 deletions

View file

@ -35,7 +35,7 @@ type clientConfig struct {
// WebSocket options // WebSocket options
wsDialer *websocket.Dialer wsDialer *websocket.Dialer
wsMessageSizeLimit *int64 wsMessageSizeLimit *int64 // wsMessageSizeLimit nil = default, 0 = no limit
// RPC handler options // RPC handler options
idgen func() ID idgen func() ID

View file

@ -38,7 +38,7 @@ const (
wsPingInterval = 30 * time.Second wsPingInterval = 30 * time.Second
wsPingWriteTimeout = 5 * time.Second wsPingWriteTimeout = 5 * time.Second
wsPongTimeout = 30 * time.Second wsPongTimeout = 30 * time.Second
wsMessageSizeLimit = 32 * 1024 * 1024 wsDefaultReadLimit = 32 * 1024 * 1024
) )
var wsBufferPool = new(sync.Pool) var wsBufferPool = new(sync.Pool)
@ -60,7 +60,7 @@ func (s *Server) WebsocketHandler(allowedOrigins []string) http.Handler {
log.Debug("WebSocket upgrade failed", "err", err) log.Debug("WebSocket upgrade failed", "err", err)
return return
} }
codec := newWebsocketCodec(conn, r.Host, r.Header, wsMessageSizeLimit) codec := newWebsocketCodec(conn, r.Host, r.Header, wsDefaultReadLimit)
s.ServeCodec(codec, 0) s.ServeCodec(codec, 0)
}) })
} }
@ -251,11 +251,9 @@ func newClientTransportWS(endpoint string, cfg *clientConfig) (reconnectFunc, er
} }
return nil, hErr return nil, hErr
} }
var messageSizeLimit int64 messageSizeLimit := int64(wsDefaultReadLimit)
if cfg.wsMessageSizeLimit != nil { if cfg.wsMessageSizeLimit != nil && *cfg.wsMessageSizeLimit >= 0 {
messageSizeLimit = *cfg.wsMessageSizeLimit messageSizeLimit = *cfg.wsMessageSizeLimit
} else {
messageSizeLimit = wsMessageSizeLimit
} }
return newWebsocketCodec(conn, dialURL, header, messageSizeLimit), nil return newWebsocketCodec(conn, dialURL, header, messageSizeLimit), nil
} }
@ -289,8 +287,8 @@ type websocketCodec struct {
pongReceived chan struct{} pongReceived chan struct{}
} }
func newWebsocketCodec(conn *websocket.Conn, host string, req http.Header, messageSizeLimit int64) ServerCodec { func newWebsocketCodec(conn *websocket.Conn, host string, req http.Header, readLimit int64) ServerCodec {
conn.SetReadLimit(messageSizeLimit) conn.SetReadLimit(readLimit)
encode := func(v interface{}, isErrorResponse bool) error { encode := func(v interface{}, isErrorResponse bool) error {
return conn.WriteJSON(v) return conn.WriteJSON(v)
} }

View file

@ -125,38 +125,52 @@ func TestWebsocketLargeRead(t *testing.T) {
defer srv.Stop() defer srv.Stop()
defer httpsrv.Close() defer httpsrv.Close()
testLimit := func(limit int64) { testLimit := func(limit *int64) {
opts := []ClientOption{} opts := []ClientOption{}
if limit >= 0 { expLimit := int64(wsDefaultReadLimit)
opts = append(opts, WithWebsocketMessageSizeLimit(limit)) if limit != nil && *limit >= 0 {
} else { opts = append(opts, WithWebsocketMessageSizeLimit(*limit))
limit = wsMessageSizeLimit if *limit > 0 {
expLimit = *limit // 0 means infinite
}
} }
client, err := DialOptions(context.Background(), wsURL, opts...) client, err := DialOptions(context.Background(), wsURL, opts...)
if err != nil { if err != nil {
t.Fatalf("can't dial: %v", err) t.Fatalf("can't dial: %v", err)
} }
defer client.Close() defer client.Close()
// Remove some bytes for json encoding overhead. // Remove some bytes for json encoding overhead.
underLimit := int(limit - 128) underLimit := int(expLimit - 128)
overLimit := expLimit + 1
if expLimit == wsDefaultReadLimit {
// No point trying the full 32MB in tests. Just sanity-check that
// it's not obviously limited.
underLimit = 1024
overLimit = -1
}
var res string var res string
err = client.Call(&res, "test_repeat", "A", underLimit) // Check under limit
if err != nil { if err = client.Call(&res, "test_repeat", "A", underLimit); err != nil {
t.Fatalf("unexpected error with limit %d: %v", limit, err) t.Fatalf("unexpected error with limit %d: %v", expLimit, err)
} }
if len(res) != underLimit || strings.Count(res, "A") != underLimit { if len(res) != underLimit || strings.Count(res, "A") != underLimit {
t.Fatal("incorrect data") t.Fatal("incorrect data")
} }
// Check over limit
err = client.Call(&res, "test_repeat", "A", limit+1) if overLimit > 0 {
err = client.Call(&res, "test_repeat", "A", expLimit+1)
if err == nil || err != websocket.ErrReadLimit { if err == nil || err != websocket.ErrReadLimit {
t.Fatalf("wrong error with limit %d: %v expecting %v", limit, err, websocket.ErrReadLimit) t.Fatalf("wrong error with limit %d: %v expecting %v", expLimit, err, websocket.ErrReadLimit)
} }
} }
}
ptr := func(v int64) *int64 { return &v }
testLimit(-1) testLimit(ptr(-1)) // Should be ignored (use default)
testLimit(wsMessageSizeLimit * 2) testLimit(ptr(0)) // Should be ignored (use default)
testLimit(nil) // Should be ignored (use default)
testLimit(ptr(200))
testLimit(ptr(wsDefaultReadLimit * 2))
} }
func TestWebsocketPeerInfo(t *testing.T) { func TestWebsocketPeerInfo(t *testing.T) {
@ -252,7 +266,7 @@ func TestClientWebsocketLargeMessage(t *testing.T) {
defer srv.Stop() defer srv.Stop()
defer httpsrv.Close() defer httpsrv.Close()
respLength := wsMessageSizeLimit - 50 respLength := wsDefaultReadLimit - 50
srv.RegisterName("test", largeRespService{respLength}) srv.RegisterName("test", largeRespService{respLength})
c, err := DialWebsocket(context.Background(), wsURL, "") c, err := DialWebsocket(context.Background(), wsURL, "")