diff --git a/nimcrypto/sysrand.nim b/nimcrypto/sysrand.nim index a90e280..43b4d07 100644 --- a/nimcrypto/sysrand.nim +++ b/nimcrypto/sysrand.nim @@ -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 @@ -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 diff --git a/tests/testsysrand.nim b/tests/testsysrand.nim index e7a0bde..cda8575 100644 --- a/tests/testsysrand.nim +++ b/tests/testsysrand.nim @@ -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