Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
48 changes: 34 additions & 14 deletions nimcrypto/sysrand.nim
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,29 @@

{.push raises: [].}

type
RandomFillProc* = proc(pbytes: pointer, nbytes: int): int {.gcsafe, raises: [].}
## Reader used by `fillRandomBytes`. It should write up to `nbytes` bytes at
## `pbytes` and return the number of bytes produced (> 0), `0` to stop, or a
## negative value to indicate a retryable condition (e.g. `EINTR`).

proc fillRandomBytes*(pbytes: pointer, nbytes: int,
reader: RandomFillProc): int =
## Fill `nbytes` of memory at `pbytes` by repeatedly invoking `reader`,
## resuming at the correct offset after short reads. Returns the number of
## bytes actually written.
var res = 0
while res < nbytes:
let p = cast[pointer](cast[uint](pbytes) + uint(res))
let bytesRead = reader(p, nbytes - res)
if bytesRead > 0:
res += bytesRead
elif bytesRead == 0:
break
else:
discard
res

when defined(posix):
import os, posix

Expand Down Expand Up @@ -128,27 +151,24 @@ elif defined(linux):
gSystemRng = newSystemRng()
gSystemRng

proc getrandomReader(p: pointer, n: int): int {.gcsafe, raises: [].} =
let r = int(syscall(SYS_getrandom, p, n, 0))
if r >= 0:
r
elif osLastError().int32 == EINTR:
-1
else:
0

proc randomBytes*(pbytes: pointer, nbytes: int): int =
var p: pointer
let srng = getSystemRng()

if srng.getRandomPresent:
var res = 0
while res < nbytes:
p = cast[pointer](cast[uint](pbytes) + uint(res))
let bytesRead = syscall(SYS_getrandom, pbytes, nbytes - res, 0)
if bytesRead > 0:
res += bytesRead
elif bytesRead == 0:
break
else:
if osLastError().int32 != EINTR:
break

var res = fillRandomBytes(pbytes, nbytes, getrandomReader)
if res <= 0:
res = urandomRead(pbytes, nbytes)
elif res < nbytes:
p = cast[pointer](cast[uint](pbytes) + uint(res))
let p = cast[pointer](cast[uint](pbytes) + uint(res))
let bytesRead = urandomRead(p, nbytes - res)
if bytesRead != -1:
res += bytesRead
Expand Down
28 changes: 28 additions & 0 deletions tests/testsysrand.nim
Original file line number Diff line number Diff line change
Expand Up @@ -41,3 +41,31 @@ suite "OS random source Tests":
result = true

check test() == true

test "getrandom partial-read offset test":
# A source such as getrandom(2) may return fewer bytes than requested for
# large buffers. `fillRandomBytes` must resume at the correct offset after
# each short read; otherwise later reads overwrite the beginning of the
# buffer and leave the tail untouched while still reporting full success.
const Total = 64
var buffer: array[Total, byte]

proc shortReader(pbytes: pointer, nbytes: int): int {.gcsafe, raises: [].} =
# Emulate a source that only ever produces up to 8 bytes per call, writing
# a non-zero marker so an untouched tail is detectable.
let chunk = min(nbytes, 8)
let arr = cast[ptr UncheckedArray[byte]](pbytes)
for i in 0 ..< chunk:
arr[i] = 0xAA'u8
chunk

let filled = fillRandomBytes(addr buffer[0], Total, shortReader)

var allWritten = true
for b in buffer:
if b != 0xAA'u8:
allWritten = false

check:
filled == Total
allWritten == true