/// @title Math
/// @notice SPDX-License-Identifier: MIT
/// @author manas <https://github.com/manasbir>
/// @author kadenzipfel <https://github.com/kadenzipfel>
/// @notice Math module over Solidity's arithmetic operations
/// @notice Adapted from OpenZeppelin (https://github.com/OpenZeppelin/openzeppelin-contracts/blob/master/contracts/utils/math/Math.sol)
#include "../utils/Errors.huff"
// Interface
#define function sqrt(uint256) pure returns(uint256)
#define function max(uint256, uint256) pure returns(uint256)
#define function min(uint256, uint256) pure returns(uint256)
#define function average(uint256, uint256) pure returns(uint256)
#define function ceilDiv(uint256, uint256) pure returns(uint256)
// Negative 1 in two's compliment
#define constant NEG1 = 0xffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffff
/// @notice Calculates the square root of a stack input
#define macro SQRT() = takes (1) returns (1) {
// input stack // [num]
// if num == 0 return 0
dup1 // [number, number]
iszero // [is_zero, number]
is_zero jumpi
// assign stack vars
dup1 // [x, num]
0x01 // [result, x, num]
// if (x >> 128 > 0) {
// x >>= 128;
// result <<= 64;
// }
dup2 // [x, result, x, num]
0x80 shr // [x >> 128, result, x, num]
dup1 // [x >> 128, x >> 128, result, x, num]
iszero // [x >> 128 == 0, x >> 128, result, x, num]
sh_128_0 jumpi
swap1 0x40 shl // [result, x >> 128, x, num]
swap1 swap2 // [x, result, x >> 128, num]
sh_128_0:
pop
// if (x >> 64 > 0) {
// x >>= 64;
// result <<= 32;
// }
dup2 // [x, result, x, num]
0x40 shr // [x >> 64, result, x, num]
dup1 // [x >> 64, x >> 64, result, x, num]
iszero // [x >> 64 == 0, x >> 64, result, x, num]
sh_64_0 jumpi
swap1 0x20 shl // [result, x >> 64, x, num]
swap1 swap2 // [x, result, x >> 64, num]
sh_64_0:
pop
// if (x >> 32 > 0) {
// x >>= 32;
// result <<= 16;
// }
dup2 // [x, result, x, num]
0x20 shr // [x >> 32, result, x, num]
dup1 // [x >> 32, x >> 32, result, x, num]
iszero // [x >> 32 == 0, x >> 32, result, x, num]
sh_32_0 jumpi
swap1 0x10 shl // [result, x >> 32, x, num]
swap1 swap2 // [x, result, x >> 32, num]
sh_32_0:
pop
// if (x >> 16 > 0) {
// x >>= 16;
// result <<= 8;
// }
dup2 // [x, result, x, num]
0x10 shr // [x >> 16, result, x, num]
dup1 // [x >> 16, x >> 16, result, x, num]
iszero // [x >> 16 == 0, x >> 16, result, x, num]
sh_16_0 jumpi
swap1 0x08 shl // [result, x >> 16, x, num]
swap1 swap2 // [x, result, x >> 16, num]
sh_16_0:
pop
// if (x >> 8 > 0) {
// x >>= 8;
// result <<= 4;
// }
dup2 // [x, result, x, num]
0x08 shr // [x >> 8, result, x, num]
dup1 // [x >> 8, x >> 8, result, x, num]
iszero // [x >> 8 == 0, x >> 8, result, x, num]
sh_8_0 jumpi
swap1 0x04 shl // [result, x >> 8, x, num]
swap1 swap2 // [x, result, x >> 8, num]
sh_8_0:
pop
// if (x >> 4 > 0) {
// x >>= 4;
// result <<= 2;
// }
dup2 // [x, result, x, num]
0x04 shr // [x >> 4, result, x, num]
dup1 // [x >> 4, x >> 4, result, x, num]
iszero // [x >> 4 == 0, x >> 4, result, x, num]
sh_4_0 jumpi
swap1 0x02 shl // [result, x >> 4, x, num]
swap1 swap2 // [x, result, x >> 4, num]
sh_4_0:
pop
// if (x >> 2 > 0) {
// x >>= 2;
// result <<= 1;
// }
dup2 // [x, result, x, num]
0x02 shr // [x >> 2, result, x, num]
dup1 // [x >> 2, x >> 2, result, x, num]
iszero // [x >> 2 == 0, x >> 2, result, x, num]
sh_2_0 jumpi
swap1 0x01 shl // [result, x >> 2, x, num]
swap1 swap2 // [x, result, x >> 2, num]
sh_2_0:
pop
// result = (result + num / result) >> 1;
dup1 // [result, result, x, num]
dup4 // [num, result, result, x, num]
div // [num / result, result x, num]
add // [result + num / result, x, num]
0x01 shr // [result, x, num]
// result = (result + num / result) >> 1;
dup1 // [result, result, x, num]
dup4 // [num, result, result, x, num]
div // [num / result, result, x, num]
add // [result + num / result, x, num]
0x01 shr // [result, x, num]
// result = (result + num / result) >> 1;
dup1 // [result, result, x, num]
dup4 // [num, result, result, x, num]
div // [num / result, result, x, num]
add // [result + num / result, x, num]
0x01 shr // [result, x, num]
// result = (result + num / result) >> 1;
dup1 // [result, result, x, num]
dup4 // [num, result, result, x, num]
div // [num / result, result, x, num]
add // [result + num / result, x, num]
0x01 shr // [result, x, num]
// result = (result + num / result) >> 1;
dup1 // [result, result, x, num]
dup4 // [num, result, result, x, num]
div // [num / result, result, x, num]
add // [result + num / result, x, num]
0x01 shr // [result, x, num]
// result = (result + num / result) >> 1;
dup1 // [result, result, x, num]
dup4 // [num, result, result, x, num]
div // [num / result, result, x, num]
add // [result + num / result, x, num]
0x01 shr // [result, x, num]
// result = (result + num / result) >> 1;
dup1 // [result, result, x, num]
dup4 // [num, result, result, x, num]
div // [num / result, result, x, num]
add // [result + num / result, x, num]
0x01 shr // [result, x, num]
// result = min(result, num / result)
dup1 // [result, result, x, num]
swap3 // [num, result, x, result]
div // [num / result, x, result]
swap2 swap1 pop // [result, num / result]
MIN()
is_zero:
}
/// @notice Returns the maximum value of two values on the stack
#define macro MAX() = takes (2) returns (1) {
// input stack: // [num1, num2]
dup2 // [num2, num1, num2]
dup2 // [num1, num2, num1, num2]
lt // [is_less_than, num1, num2]
less_than jumpi
swap1 // [num1, num2]
less_than:
pop // [max(num2, num1)]
}
/// @notice Returns the ceiling of the division of two values on the stack
#define macro CEIL_DIV() = takes (2) returns (1) {
// input stack: // [num1, num2]
dup1 iszero // [is_zero, num1, num2]
if_zero jumpi
dup2 iszero // [is_zero, num1, num2]
divide_by_zero jumpi
0x01 // [1, num1, num2]
swap1 // [num1, 1, num1]
sub // [num1 - 1, num2]
div 0x01 add // [((num1 - 1) / num2) + 1]
dest jump
divide_by_zero:
[DIVIDE_BY_ZERO] PANIC()
if_zero:
swap1 // [num2, num1]
pop // [num1]
dest:
}
/// @notice Returns the minimum value of two values on the stack
#define macro MIN() = takes (2) returns (1) {
// input stack: // [num1, num2]
dup2 dup2 gt // [is_greater_than, num1, num2]
greater_than jumpi
swap1 // [num2, num1]
greater_than:
pop // [min(num1, num2)]
}
/// @notice Returns the average of two values on the stack
#define macro AVG() = takes (2) returns (1) {
// input stack: // [num1, num2]
dup2 dup2 and // [num1 & num2, num1, num2]
swap2 // [num2, num1, num1 & num2]
xor // [num2 ^ num1, num1 & num2]
0x02 swap1 div // [num2 ^ num1 / 2, num1 & num2]
add // [sum]
}
/// @notice Unsafely subtracts 1 from a uint256 using the 2's complement representation
#define macro UNSAFE_SUB() = takes (1) returns (1) {
// Input Stack: [x]
// Output Stack: [x - 1]
[NEG1] add // [x - 1]
}