"""Tests for the union-find / representative-picker / normalization helpers.""" from __future__ import annotations from sdg.preprocessing.dedupe.base import UnionFind, normalize_text, pick_representatives # ── UnionFind ──────────────────────────────────────────────────────────────── def test_unionfind_each_singleton(): uf = UnionFind(5) assert {uf.find(i) for i in range(5)} == {0, 1, 2, 3, 4} def test_unionfind_simple_merge(): uf = UnionFind(5) uf.union(0, 1) assert uf.find(0) == uf.find(1) assert uf.find(2) != uf.find(0) def test_unionfind_transitive_chain(): uf = UnionFind(5) uf.union(0, 1) uf.union(1, 2) uf.union(2, 3) assert uf.find(0) == uf.find(3) assert uf.find(0) != uf.find(4) def test_unionfind_idempotent_self_union(): uf = UnionFind(3) uf.union(1, 1) uf.union(0, 1) uf.union(0, 1) # repeat assert uf.find(0) == uf.find(1) # ── pick_representatives ───────────────────────────────────────────────────── def test_pick_reps_no_pairs_keeps_all(): keep = pick_representatives(4, [], key_fn=None) assert keep == [0, 1, 2, 3] def test_pick_reps_collapses_pair(): keep = pick_representatives(3, [(0, 1)], key_fn=None) assert keep == [0, 2] # (0,1) collapsed to smallest index 0; 2 alone def test_pick_reps_collapses_chain(): keep = pick_representatives(5, [(0, 1), (1, 2), (3, 4)], key_fn=None) assert keep == [0, 3] def test_pick_reps_key_fn_selects_max_by_negative_length(): lengths = [10, 50, 5, 100] key_fn = lambda i: -lengths[i] # longest wins keep = pick_representatives(4, [(0, 1), (2, 3)], key_fn=key_fn) assert keep == [1, 3] # idx 1 (len 50) > idx 0 (10); idx 3 (100) > idx 2 (5) def test_pick_reps_returns_sorted(): keep = pick_representatives(6, [(2, 5), (0, 3)], key_fn=None) assert keep == sorted(keep) def test_pick_reps_empty(): assert pick_representatives(0, [], key_fn=None) == [] # ── normalize_text ─────────────────────────────────────────────────────────── def test_normalize_lowercases(): assert normalize_text("Hello WORLD") == "hello world" def test_normalize_collapses_whitespace(): assert normalize_text("a b\t\tc\n\nd") == "a b c d" def test_normalize_strips_edges(): assert normalize_text(" hi ") == "hi" def test_normalize_idempotent(): text = "Hello\t\tWorld\n" assert normalize_text(normalize_text(text)) == normalize_text(text)