diff --git a/shared-bindings/random/__init__.c b/shared-bindings/random/__init__.c index 89ff3572ea9..d3bcd29c9d9 100644 --- a/shared-bindings/random/__init__.c +++ b/shared-bindings/random/__init__.c @@ -116,10 +116,7 @@ static mp_obj_t random_randint(mp_obj_t a_in, mp_obj_t b_in) { if (a > b) { mp_raise_ValueError(NULL); } - if (b == (mp_int_t)(~(mp_uint_t)0 >> 1)) { - mp_raise_ValueError(NULL); - } - return mp_obj_new_int(shared_modules_random_randrange(a, b + 1, 1)); + return mp_obj_new_int(shared_modules_random_randint(a, b)); } static MP_DEFINE_CONST_FUN_OBJ_2(random_randint_obj, random_randint); diff --git a/shared-bindings/random/__init__.h b/shared-bindings/random/__init__.h index 27f255d5d8b..288805b3706 100644 --- a/shared-bindings/random/__init__.h +++ b/shared-bindings/random/__init__.h @@ -13,5 +13,7 @@ void shared_modules_random_seed(mp_uint_t seed); mp_uint_t shared_modules_random_getrandbits(uint8_t n); mp_int_t shared_modules_random_randrange(mp_int_t start, mp_int_t stop, mp_int_t step); +// Returns a random integer in [a, b] inclusive. The caller must ensure a <= b. +mp_int_t shared_modules_random_randint(mp_int_t a, mp_int_t b); mp_float_t shared_modules_random_random(void); mp_float_t shared_modules_random_uniform(mp_float_t a, mp_float_t b); diff --git a/shared-module/random/__init__.c b/shared-module/random/__init__.c index 65876ab5788..a45da46574a 100644 --- a/shared-module/random/__init__.c +++ b/shared-module/random/__init__.c @@ -82,6 +82,16 @@ mp_int_t shared_modules_random_randrange(mp_int_t start, mp_int_t stop, mp_int_t return start + step * yasmarang_randbelow(n); } +mp_int_t shared_modules_random_randint(mp_int_t a, mp_int_t b) { + // Compute the span unsigned so that b - a + 1 cannot overflow for any a <= b. + mp_uint_t n = (mp_uint_t)b - (mp_uint_t)a + 1; + if (n == 0) { + // a and b cover the full mp_int_t range: any value is in range. + return a + (mp_int_t)yasmarang(); + } + return a + yasmarang_randbelow(n); +} + // returns a number in the range [0..1) using Yasmarang to fill in the fraction bits static mp_float_t yasmarang_float(void) { #if MICROPY_FLOAT_IMPL == MICROPY_FLOAT_IMPL_DOUBLE diff --git a/tests/extmod/random_extra.py b/tests/extmod/random_extra.py index aa05053377b..958fb8f5f3f 100644 --- a/tests/extmod/random_extra.py +++ b/tests/extmod/random_extra.py @@ -48,6 +48,11 @@ assert 2 <= random.randint(2, 6) <= 6 assert -2 <= random.randint(-2, 2) <= 2 +# CIRCUITPY-CHANGE: upper bound at the largest 32-bit mp_int_t must not overflow +for i in range(50): + assert 0 <= random.randint(0, 0x7FFFFFFF) <= 0x7FFFFFFF + assert 1 <= random.randint(1, 0x7FFFFFFF) <= 0x7FFFFFFF + # empty range try: random.randint(2, 1)