class Solution {
public:
long long countFairPairs(vector<int>& nums, int lower, int upper) {
int n = nums.size();
sort(nums.begin(), nums.end());
long long ans = 0;
for (int i = 0; i <= n - 2; i++) {
ans += countPairsInRange(nums, i + 1, nums[i], lower, upper);
}
return ans;
}
int countPairsInRange(std::vector<int>& nums, int indexStart, int currentNumber, int lower, int upper) {
int indexLower = findLowerBound(nums, indexStart, currentNumber, lower);
int indexUpper = findUpperBound(nums, indexStart, currentNumber, upper);
if (indexLower <= indexUpper && indexUpper < nums.size()) {
return indexUpper - indexLower + 1;
}
return 0;
}
int findLowerBound(std::vector<int>& nums, int indexStart, int currentNumber, int lower) {
int left = indexStart;
int right = nums.size() - 1;
while (left <= right) {
int mid = left + (right - left) / 2;
if (currentNumber + nums[mid] < lower) {
left = mid + 1;
} else {
right = mid - 1;
}
}
return left;
}
int findUpperBound(std::vector<int>& nums, int indexStart, int currentNumber, int upper) {
int left = indexStart;
int right = nums.size() - 1;
while (left <= right) {
int mid = left + (right - left) / 2;
if (currentNumber + nums[mid] <= upper) {
left = mid + 1;
} else {
right = mid - 1;
}
}
return right;
}
};