diff options
Diffstat (limited to 'py/objint.c')
| -rw-r--r-- | py/objint.c | 56 |
1 files changed, 28 insertions, 28 deletions
diff --git a/py/objint.c b/py/objint.c index 3b3a3d9c0..9e11871f1 100644 --- a/py/objint.c +++ b/py/objint.c @@ -338,34 +338,34 @@ void mp_obj_int_buffer_overflow_check(mp_obj_t self_in, size_t nbytes, bool is_s void mp_small_int_buffer_overflow_check(mp_int_t val, size_t nbytes, bool is_signed) { // Fast path for zero. - if (val == 0) return; - if (!is_signed) { - if (val >= 0) { - // Using signed constants here, not UINT8_MAX, etc. to avoid any unintended conversions. - if (val <= 0xff) return; // Small values fit in any number of nbytes. - if (nbytes == 2 && val <= 0xffff) return; -#if !defined(__LP64__) - // 32-bit ints and pointers - if (nbytes >= 4) return; // Any mp_int_t will fit. -#else - // 64-bit ints and pointers - if (nbytes == 4 && val <= 0xffffffff) return; - if (nbytes >= 8) return; // Any mp_int_t will fit. -#endif - } // Negative, fall through to failure. - } else { - // signed - if (val >= INT8_MIN && val <= INT8_MAX) return; // Small values fit in any number of nbytes. - if (nbytes == 2 && val >= INT16_MIN && val <= INT16_MAX) return; -#if !defined(__LP64__) - // 32-bit ints and pointers - if (nbytes >= 4) return; // Any mp_int_t will fit. -#else - // 64-bit ints and pointers - if (nbytes == 4 && val >= INT32_MIN && val <= INT32_MAX) return; - if (nbytes >= 8) return; // Any mp_int_t will fit. -#endif - } // Fall through to failure. + if (val == 0) { + return; + } + + // Trying to store negative values in unsigned bytes falls through to failure. + if (is_signed || val >= 0) { + + if (nbytes >= sizeof(val)) { + // All non-negative N bit signed integers fit in an unsigned N bit integer. + // This case prevents shifting too far below. + return; + } + + if (is_signed) { + mp_int_t edge = ((mp_int_t)1 << (nbytes * 8 - 1)); + if (-edge <= val && val < edge) { + return; + } + // Out of range, fall through to failure. + } else { + // Unsigned. We already know val >= 0. + mp_int_t edge = ((mp_int_t)1 << (nbytes * 8)); + if (val < edge) { + return; + } + } + // Fall through to failure. + } mp_raise_OverflowError_varg(translate("value must fit in %d byte(s)"), nbytes); } |
