mirror of
https://github.com/status-im/nim-websock.git
synced 2026-08-27 11:51:13 +00:00
feat: make rng agnostic (#194)
This commit is contained in:
@@ -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
|
||||
|
||||
+1
-1
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
+1
-1
@@ -111,7 +111,7 @@ proc doSend(
|
||||
|
||||
let maskKey =
|
||||
if ws.masked:
|
||||
MaskKey.random(ws.rng[])
|
||||
MaskKey.random(ws.rng)
|
||||
else:
|
||||
default(MaskKey)
|
||||
|
||||
|
||||
+24
-1
@@ -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)
|
||||
|
||||
+98
-4
@@ -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))
|
||||
|
||||
Reference in New Issue
Block a user