diff --git a/reimplement/set.py b/reimplement/set.py index 5ac4c5f..befc3ed 100644 --- a/reimplement/set.py +++ b/reimplement/set.py @@ -260,10 +260,10 @@ def decode_set(encoded_str: str, Mshift: int) -> List[int]: def rpmsetcmp(str1: str, str2: str) -> int: - if str1[:3] == "set:": - str1 = str1[3:] - if str2[:3] == "set:": - str2 = str2[3:] + if str1.startswith("set:"): + str1 = str1[4:] + if str2.startswith("set:"): + str2 = str2[4:] bpp1, Mshift1 = decode_set_init(str1) bpp2, Mshift2 = decode_set_init(str2) @@ -288,21 +288,28 @@ def rpmsetcmp(str1: str, str2: str) -> int: while i < len(hash_values1) and j < len(hash_values2): if hash_values1[i] < hash_values2[j]: - ge = False + le = False i += 1 elif hash_values1[i] > hash_values2[j]: - le = False + ge = False j += 1 else: i += 1 j += 1 + if i < len(hash_values1): + le = False + if j < len(hash_values2): + ge = False + if ge and le: return 0 elif ge: return 1 - else: + elif le: return -1 + else: + return -2 def set_new() -> Set: diff --git a/tests/test_reimplement_set.py b/tests/test_reimplement_set.py index 917a241..5f60867 100644 --- a/tests/test_reimplement_set.py +++ b/tests/test_reimplement_set.py @@ -81,6 +81,36 @@ class DownsampleSetTest(unittest.TestCase): self.assertEqual(rpmset.downsample_set([8, 10, 14], 3), [0, 2, 6]) +class RpmSetCmpTest(unittest.TestCase): + def encode_values(self, values, bpp=8): + return rpmset.encode_set(values, bpp) + + def test_returns_zero_for_equal_sets(self): + set1 = self.encode_values([1, 3, 5]) + set2 = self.encode_values([1, 3, 5]) + self.assertEqual(rpmset.rpmsetcmp(set1, set2), 0) + + def test_accepts_set_prefix(self): + set1 = "set:" + self.encode_values([1, 3, 5]) + set2 = self.encode_values([1, 3, 5]) + self.assertEqual(rpmset.rpmsetcmp(set1, set2), 0) + + def test_returns_one_when_first_set_is_superset(self): + set1 = self.encode_values([1, 3, 5]) + set2 = self.encode_values([1, 3]) + self.assertEqual(rpmset.rpmsetcmp(set1, set2), 1) + + def test_returns_minus_one_when_first_set_is_subset(self): + set1 = self.encode_values([1, 3]) + set2 = self.encode_values([1, 3, 5]) + self.assertEqual(rpmset.rpmsetcmp(set1, set2), -1) + + def test_returns_minus_two_for_incomparable_sets(self): + set1 = self.encode_values([1, 3, 7]) + set2 = self.encode_values([1, 3, 5]) + self.assertEqual(rpmset.rpmsetcmp(set1, set2), -2) + + class ModuleEntrypointTest(unittest.TestCase): def test_module_has_no_main_side_effects(self): output = io.StringIO()