Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 30 additions & 3 deletions reladiff/hashdiff_tables.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,8 @@
import os
from functools import cmp_to_key
from numbers import Number
import logging
from typing import Iterator
from typing import Iterator, Sequence, Tuple
from operator import attrgetter
from collections import Counter
from itertools import chain
Expand All @@ -27,6 +28,33 @@
logger = logging.getLogger("hashdiff_tables")


def compare_element(a, b):
"""Compare a and b, treat None as the smallest value.

Return -1 if a < b, 0 if a == b and 1 if a > b.
"""
if a == b:
return 0
if b is not None and ((a is None) or (a < b)):
return -1
return 1


def compare(a: Tuple[str, Sequence], b: Tuple[str, Sequence]) -> int:
"""Compare two sequences of the same length.

Compare a and b until the first element a[1][i] differs from b[1][i].
See compare_element() for detailed comparison rules.

Return -1 if a < b, 0 if a == b and 1 if a > b.
"""
for i in range(len(a[1])):
res = compare_element(a[1][i], b[1][i])
if res != 0:
return res
return 0


def diff_sets(a: list, b: list, skip_sort_results: bool, duplicate_rows_support: bool) -> Iterator:
if duplicate_rows_support:
c = Counter(b)
Expand All @@ -36,8 +64,7 @@ def diff_sets(a: list, b: list, skip_sort_results: bool, duplicate_rows_support:
sa = set(a)
sb = set(b)
diff = chain((("-", x) for x in sa - sb), (("+", x) for x in sb - sa))

return diff if skip_sort_results else sorted(diff, key=lambda i: i[1]) # sort by key
return diff if skip_sort_results else sorted(diff, key=cmp_to_key(compare)) # sort by key


@dataclass(frozen=True)
Expand Down
28 changes: 27 additions & 1 deletion tests/test_diff_tables.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
from sqeleton.queries import table, this, commit
from sqeleton.utils import ArithAlphanumeric, numberToAlphanum

from reladiff.hashdiff_tables import HashDiffer
from reladiff.hashdiff_tables import HashDiffer, compare_element, diff_sets
from reladiff.joindiff_tables import JoinDiffer
from reladiff.table_segment import TableSegment, split_space, Vector
from reladiff import databases as db
Expand Down Expand Up @@ -1071,3 +1071,29 @@ def test_compound_key(self):
self.assertEqual(diff, [("-", (uuid, "9", "9")), ("+", (uuid, "9000", "9"))])

self.assertRaises(ValueError, list, differ.diff_tables(aa, a))


class TestDiffSets(unittest.TestCase):
def test_compare_element(self):
self.assertEqual(compare_element(None, 1), -1)
self.assertEqual(compare_element(1, 2), -1)
self.assertEqual(compare_element(None, 2), -1)
self.assertEqual(compare_element(None, None), 0)
self.assertEqual(compare_element(1, 1), 0)
self.assertEqual(compare_element(2, 2), 0)
self.assertEqual(compare_element(1, None), 1)
self.assertEqual(compare_element(2, 1), 1)
self.assertEqual(compare_element(2, None), 1)
self.assertEqual(compare_element(None, ''), -1)
self.assertEqual(compare_element('', ''), 0)
self.assertEqual(compare_element('', None), 1)

def test_diff_sets(self):
res = diff_sets([(1, None)], [(1, '')], skip_sort_results=False, duplicate_rows_support=True)
self.assertSequenceEqual(res, [('-', (1, None)), ('+', (1, ''))])

res = diff_sets([(1, '')], [(1, None)], skip_sort_results=False, duplicate_rows_support=True)
self.assertSequenceEqual(res, [('+', (1, None)), ('-', (1, ''))])

res = diff_sets([(1, None), (1, None)], [(1, None)], skip_sort_results=False, duplicate_rows_support=True)
self.assertSequenceEqual(res, [('-', (1, None))])