feat: make rng agnostic (#194)

This commit is contained in:
richΛrd
2026-05-11 21:08:48 -04:00
committed by GitHub
parent 72d4e57558
commit 02616eccd5
6 changed files with 147 additions and 12 deletions
+2 -2
View File
@@ -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
View File
@@ -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,
+21 -3
View File
@@ -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
View File
@@ -111,7 +111,7 @@ proc doSend(
let maskKey =
if ws.masked:
MaskKey.random(ws.rng[])
MaskKey.random(ws.rng)
else:
default(MaskKey)
+24 -1
View File
@@ -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
View File
@@ -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))