有序数组Floor查找Rust实现重复元素场景失效问题排查
问题现象
下面这段Rust代码实现了有序数组的Floor查找逻辑,在无重复元素的测试场景下运行正常,但碰到包含重复元素的测试用例(比如assert_eq!(find_floor(&[1, 3, 4, 4, 10], 4), Some(2)))时会失败;要是随便修改判断条件,又会搞坏其他正常的测试用例,比如assert_eq!(find_floor(&[1, 2, 8, 10, 11, 12, 19], 5), Some(1))就会失效。
原Rust代码
pub fn find_floor(arr: &[i32], target: i32) -> Option<usize> { let (mut left, mut right) = (0, arr.len()); while left < right { let mid = (left + right) / 2; if arr[mid] > target { right = mid; } else { left = mid + 1; } } if right == 0 { None } else { Some(right - 1) } }
可正常运行的C++实现
int floor(const std::vector<int> &arr, int target) { int left{0}, right{(int) arr.size()}, ans{-1}; while (left < right) { int mid = (left + right) / 2; if (arr[mid] <= target) { ans = mid; left = mid + 1; } else { right = mid; } } return ans; }
问题根源
原Rust代码的逻辑是通过二分法不断收缩区间,最终right会指向第一个比target大的元素的位置,所以right-1就是数组里不大于target的最大元素的索引——这完全符合Floor查找的标准定义(Floor指不大于目标值的最大元素)。
你遇到的测试用例失败,本质是测试用例的预期结果错了:对于数组[1, 3, 4, 4, 10]和target=4,标准的Floor结果应该是最后一个4的索引3,而不是索引2。
要是你实际需要的是第一个小于等于target的元素(也就是常说的Lower Bound),那原Rust代码的逻辑肯定不适用;但如果是要实现标准的Floor查找,原代码本身没毛病,问题出在测试用例的预期值不符合定义。这也是你随便修改条件后另一个测试用例失效的原因——改逻辑相当于把需求从“找最大的符合条件元素”改成了别的,自然会破坏原有正确的场景。
另外对比C++代码就能发现,它的逻辑和原Rust代码是等价的,都是找到最后一个小于等于target的元素,返回的结果和原Rust代码一致,所以你的测试用例预期Some(2)本身是错误的。
修复方案(兼容标准Floor与所有测试用例)
如果是测试用例写错了,那原Rust代码本身就是正确的;如果确实需要调整需求,比如要返回第一个小于等于target的元素,或者严格匹配你的测试用例预期,可以参考C++的逻辑修改Rust代码,同时保证原有正常测试用例不受影响:
pub fn find_floor(arr: &[i32], target: i32) -> Option<usize> { let (mut left, mut right) = (0, arr.len()); let mut result = None; while left < right { let mid = (left + right) / 2; if arr[mid] <= target { // 记录当前符合条件的索引 result = Some(mid); // 继续向右查找,确保找到最后一个符合条件的元素(标准Floor) left = mid + 1; } else { right = mid; } } result }
这段代码和C++实现逻辑完全一致:
- 面对
&[1, 3, 4, 4, 10], 4时,返回Some(3)(标准Floor的正确结果) - 面对
&[1, 2, 8, 10, 11, 12, 19], 5时,返回Some(1),完全符合预期
要是你确实需要返回第一个等于target的元素(而非标准Floor),可以调整二分逻辑中的区间收缩方向,把left = mid + 1改成right = mid,同时优化判断逻辑:
pub fn find_first_floor(arr: &[i32], target: i32) -> Option<usize> { let (mut left, mut right) = (0, arr.len()); let mut result = None; while left < right { let mid = (left + right) / 2; if arr[mid] < target { left = mid + 1; } else { right = mid; if arr[mid] == target { result = Some(mid); } } } // 最后检查区间内是否存在符合条件的元素 if left < arr.len() && arr[left] <= target { result.or(Some(left)) } else { result } }
内容的提问来源于stack exchange,提问作者KillerFC

