From e03a631bc0b582bfb75923722509805d721eb006 Mon Sep 17 00:00:00 2001 From: atarpara Date: Mon, 31 Aug 2026 15:40:02 +0530 Subject: [PATCH] =?UTF-8?q?=F0=9F=90=9E=20Fix=20groupSum=20sorting?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/utils/LibSort.sol | 2 ++ test/LibSort.t.sol | 63 +++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 65 insertions(+) diff --git a/src/utils/LibSort.sol b/src/utils/LibSort.sol index ea7fe145ca..2d811e342e 100644 --- a/src/utils/LibSort.sol +++ b/src/utils/LibSort.sol @@ -650,7 +650,9 @@ library LibSort { /// @dev Sorts and uniquifies `keys`. Updates `values` with the grouped sums by key. function groupSum(int256[] memory keys, uint256[] memory values) internal pure { + _flipSign(keys); groupSum(_toUints(keys), values); + _flipSign(keys); } /// @dev Returns if `a` has any duplicate. Does NOT mutate `a`. `O(n)`. diff --git a/test/LibSort.t.sol b/test/LibSort.t.sol index 14ea5c31be..fc644a05db 100644 --- a/test/LibSort.t.sol +++ b/test/LibSort.t.sol @@ -1350,6 +1350,69 @@ contract LibSortTest is SoladyTest { assertEq(_sum(sums), oriSum); } + function testGroupSumSigned() public { + int256[] memory keys = new int256[](5); + uint256[] memory values = new uint256[](5); + keys[0] = 7; + keys[1] = -1; + keys[2] = 3; + keys[3] = -10; + keys[4] = 0; + unchecked { + for (uint256 i; i < 5; ++i) { + values[i] = i + 1; + } + } + LibSort.groupSum(keys, values); + int256[] memory expectedKeys = new int256[](5); + expectedKeys[0] = -10; + expectedKeys[1] = -1; + expectedKeys[2] = 0; + expectedKeys[3] = 3; + expectedKeys[4] = 7; + uint256[] memory expectedValues = new uint256[](5); + expectedValues[0] = 4; + expectedValues[1] = 2; + expectedValues[2] = 5; + expectedValues[3] = 3; + expectedValues[4] = 1; + assertEq(keys, expectedKeys); + assertEq(values, expectedValues); + } + + function testGroupSumSigned(bytes32) public { + if (_randomChance(2)) { + _misalignFreeMemoryPointer(); + _brutalizeMemory(); + } + uint256 n = _random() & 0x1f; + int256[] memory keys = new int256[](n); + uint256[] memory values = new uint256[](n); + unchecked { + for (uint256 i; i < n; ++i) { + keys[i] = int256(_randomUniform() & 0xf) - 8; // Straddles zero. + values[i] = _randomUniform() & 0xff; + } + } + uint256 oriSum = _sum(values); + int256[] memory uniqueKeys = LibSort.copy(keys); + LibSort.insertionSort(uniqueKeys); + LibSort.uniquifySorted(uniqueKeys); + uint256[] memory sums = new uint256[](uniqueKeys.length); + unchecked { + for (uint256 i; i < n; ++i) { + (, uint256 j) = LibSort.searchSorted(uniqueKeys, keys[i]); + sums[j] += values[i]; + } + } + LibSort.groupSum(keys, values); + _checkMemory(sums); + assertEq(keys, uniqueKeys); + assertEq(values, sums); + assertEq(_sum(sums), oriSum); + assertTrue(LibSort.isSortedAndUniquified(keys)); + } + function _sum(uint256[] memory a) internal pure returns (uint256 result) { unchecked { for (uint256 i; i < a.length; ++i) {