diff --git a/tests/extensions/testextflow.nim b/tests/extensions/testextflow.nim index 15282d92b8..46afec7011 100644 --- a/tests/extensions/testextflow.nim +++ b/tests/extensions/testextflow.nim @@ -116,11 +116,11 @@ suite "Encode frame extensions flow": frame.opcode == Opcode.Binary suite "Decode frame extensions flow": - let rng = HmacDrbgContext.new() + let rng = bearSslRng(HmacDrbgContext.new()) var address: TransportAddress server: StreamServer - maskKey = MaskKey.random(rng[]) + maskKey = MaskKey.random(rng) transport: StreamTransport reader: AsyncStreamReader frame: Frame diff --git a/tests/helpers.nim b/tests/helpers.nim index 1308483c59..e6ed46c511 100644 --- a/tests/helpers.nim +++ b/tests/helpers.nim @@ -91,7 +91,7 @@ proc connectClient*( onPing: ControlCb = nil, onPong: ControlCb = nil, onClose: CloseCb = nil, - rng = HmacDrbgContext.new()): Future[WSSession] {.async.} = + rng: RandomBytesRng = bearSslRng(HmacDrbgContext.new())): Future[WSSession] {.async.} = let secure = when defined secure: true else: false return await WebSocket.connect( host = address, diff --git a/tests/testwebsockets.nim b/tests/testwebsockets.nim index 31660cfccc..d5c5ea015b 100644 --- a/tests/testwebsockets.nim +++ b/tests/testwebsockets.nim @@ -19,6 +19,24 @@ import ./helpers let address = initTAddress("127.0.0.1:8888") +suite "Test RNG": + test "custom random-bytes callback": + let rng: RandomBytesRng = + proc(dst: var openArray[byte]): bool {.closure, gcsafe, raises: [].} = + for i in 0 ..< dst.len: + dst[i] = byte(i + 1) + true + + check MaskKey.random(rng) == [byte 1, 2, 3, 4] + + test "custom random-bytes callback failure raises WSRngError": + let rng: RandomBytesRng = + proc(_: var openArray[byte]): bool {.closure, gcsafe, raises: [].} = + false + + expect WSRngError: + discard MaskKey.random(rng) + suite "Test handshake": setup: var @@ -730,7 +748,7 @@ suite "Test Closing": suite "Test Payload": setup: - let rng {.used.} = HmacDrbgContext.new() + let rng {.used.} = bearSslRng(HmacDrbgContext.new()) var server: HttpServer @@ -834,7 +852,7 @@ suite "Test Payload": address = initTAddress("127.0.0.1:8888"), frameSize = maxFrameSize) - let maskKey = MaskKey.random(rng[]) + let maskKey = MaskKey.random(rng) await session.stream.writer.write( (await Frame( fin: false, @@ -894,7 +912,7 @@ suite "Test Payload": pong = true ) - let maskKey = MaskKey.random(rng[]) + let maskKey = MaskKey.random(rng) await session.stream.writer.write( (await Frame( fin: false, diff --git a/websock/session.nim b/websock/session.nim index 73897ea54b..dd8f0128e5 100644 --- a/websock/session.nim +++ b/websock/session.nim @@ -111,7 +111,7 @@ proc doSend( let maskKey = if ws.masked: - MaskKey.random(ws.rng[]) + MaskKey.random(ws.rng) else: default(MaskKey) diff --git a/websock/types.nim b/websock/types.nim index 46248357f5..d740fdf26f 100644 --- a/websock/types.nim +++ b/websock/types.nim @@ -54,6 +54,9 @@ type MaskKey* = array[4, byte] WebSecKey* = array[16, byte] + RandomBytesRng* = proc(dst: var openArray[byte]): bool {. + closure, gcsafe, raises: [] + .} Frame* = ref object fin*: bool ## Indicates that this is the final fragment in a message. @@ -88,7 +91,7 @@ type masked*: bool # send masked packets binary*: bool # is payload binary? flags*: set[TLSFlags] - rng*: ref HmacDrbgContext + rng*: RandomBytesRng frameSize*: int # max frame buffer size onPing*: ControlCb onPong*: ControlCb @@ -167,6 +170,7 @@ type WSInvalidUTF8* = object of WebSocketError WSExtError* = object of WebSocketError WSHookError* = object of WebSocketError + WSRngError* = object of WebSocketError const StatusNotUsed* = (StatusCodes(0)..StatusCodes(999)) @@ -220,5 +224,24 @@ method encode*( method toHttpOptions*(self: Ext): string {.base, gcsafe.} = raiseAssert "Not implemented!" +proc generate*(rng: RandomBytesRng, dst: var openArray[byte]): bool = + if rng.isNil: + return false + rng(dst) + +proc bearSslRng*(rng: ref HmacDrbgContext): RandomBytesRng = + ## Wrap an existing BearSSL HMAC-DRBG context. + doAssert not rng.isNil, "rng cannot be null" + proc(dst: var openArray[byte]): bool {.closure, gcsafe, raises: [].} = + if dst.len > 0: + hmacDrbgGenerate(rng[], addr dst[0], uint dst.len) + true + +proc random*( + T: typedesc[MaskKey | WebSecKey], rng: RandomBytesRng +): T {.raises: [WebSocketError].} = + if not rng.generate(result): + raise newException(WSRngError, "Failed to generate WebSocket random bytes") + func random*(T: typedesc[MaskKey | WebSecKey], rng: var HmacDrbgContext): T = rng.generate(result) diff --git a/websock/websock.nim b/websock/websock.nim index 270840b8f5..8158bbddae 100644 --- a/websock/websock.nim +++ b/websock/websock.nim @@ -103,7 +103,7 @@ proc connect*( onPing: ControlCb = nil, onPong: ControlCb = nil, onClose: CloseCb = nil, - rng = HmacDrbgContext.new(), + rng: RandomBytesRng = bearSslRng(HmacDrbgContext.new()), ): Future[WSSession] {. async: ( raises: @@ -111,7 +111,7 @@ proc connect*( ) .} = let - key = Base64Pad.encode(WebSecKey.random(rng[])) + key = Base64Pad.encode(WebSecKey.random(rng)) hostname = if hostName.len > 0: hostName else: $host var @@ -195,6 +195,47 @@ proc connect*( if not connected: await client.closeWait() +proc connect*( + _: type WebSocket, + host: string | TransportAddress, + path: string, + hostName: string = "", + # override used when the hostname has been externally resolved + protocols: seq[string] = @[], + factories: seq[ExtFactory] = @[], + hooks: seq[Hook] = @[], + secure = false, + flags: set[TLSFlags] = {}, + version = WSDefaultVersion, + frameSize = WSDefaultFrameSize, + onPing: ControlCb = nil, + onPong: ControlCb = nil, + onClose: CloseCb = nil, + rng: ref HmacDrbgContext, +): Future[WSSession] {. + deprecated: "use the RandomBytesRng overload", + async: ( + raises: + [CancelledError, AsyncStreamError, HttpError, TransportError, WebSocketError] + ) +.} = + await WebSocket.connect( + host = host, + path = path, + hostName = hostName, + protocols = protocols, + factories = factories, + hooks = hooks, + secure = secure, + flags = flags, + version = version, + frameSize = frameSize, + onPing = onPing, + onPong = onPong, + onClose = onClose, + rng = bearSslRng(rng), + ) + proc connect*( _: type WebSocket, uri: Uri, @@ -207,7 +248,7 @@ proc connect*( onPing: ControlCb = nil, onPong: ControlCb = nil, onClose: CloseCb = nil, - rng = HmacDrbgContext.new(), + rng: RandomBytesRng = bearSslRng(HmacDrbgContext.new()), ): Future[WSSession] {. async: ( raises: @@ -246,6 +287,39 @@ proc connect*( rng = rng, ) +proc connect*( + _: type WebSocket, + uri: Uri, + protocols: seq[string] = @[], + factories: seq[ExtFactory] = @[], + hooks: seq[Hook] = @[], + flags: set[TLSFlags] = {}, + version = WSDefaultVersion, + frameSize = WSDefaultFrameSize, + onPing: ControlCb = nil, + onPong: ControlCb = nil, + onClose: CloseCb = nil, + rng: ref HmacDrbgContext, +): Future[WSSession] {. + async: ( + raises: + [CancelledError, AsyncStreamError, HttpError, TransportError, WebSocketError] + ) +.} = + await WebSocket.connect( + uri = uri, + protocols = protocols, + factories = factories, + hooks = hooks, + flags = flags, + version = version, + frameSize = frameSize, + onPing = onPing, + onPong = onPong, + onClose = onClose, + rng = bearSslRng(rng), + ) + proc handleRequest*( ws: WSServer, request: HttpRequest, @@ -350,7 +424,7 @@ proc new*( onPing: ControlCb = nil, onPong: ControlCb = nil, onClose: CloseCb = nil, - rng = HmacDrbgContext.new()): WSServer = + rng: RandomBytesRng = bearSslRng(HmacDrbgContext.new())): WSServer = return WSServer( protocols: @protos, @@ -361,3 +435,23 @@ proc new*( onPing: onPing, onPong: onPong, onClose: onClose) + +proc new*( + _: typedesc[WSServer], + protos: openArray[string] = [""], + factories: openArray[ExtFactory] = [], + frameSize = WSDefaultFrameSize, + onPing: ControlCb = nil, + onPong: ControlCb = nil, + onClose: CloseCb = nil, + rng: ref HmacDrbgContext): WSServer {. + deprecated: "use the RandomBytesRng overload".} = + + WSServer.new( + protos = protos, + factories = factories, + frameSize = frameSize, + onPing = onPing, + onPong = onPong, + onClose = onClose, + rng = bearSslRng(rng))