"""Shared primitives: union-find, cluster -> representative selection, normalize.""" from __future__ import annotations import re from collections import defaultdict from typing import Any, Callable, Iterable class UnionFind: """Union-find with path compression and union-by-rank.""" def __init__(self, n: int): self.parent = list(range(n)) self.rank = [0] * n def find(self, x: int) -> int: root = x while self.parent[root] != root: root = self.parent[root] while self.parent[x] != root: self.parent[x], x = root, self.parent[x] return root def union(self, x: int, y: int) -> None: rx, ry = self.find(x), self.find(y) if rx == ry: return if self.rank[rx] < self.rank[ry]: rx, ry = ry, rx self.parent[ry] = rx if self.rank[rx] == self.rank[ry]: self.rank[rx] += 1 def cluster_and_pick( n: int, pairs: Iterable[tuple[int, int]], key_fn: Callable[[int], Any] | None = None, ) -> tuple[list[int], dict[int, list[int]]]: """Cluster items in [0, n) via `pairs`, pick one rep per cluster. Returns (keep, clusters): keep: sorted list of representative indices (one per cluster). clusters: dict {rep_idx -> sorted list of ALL member indices in that cluster, including the rep}. Singletons appear as {i: [i]}. """ uf = UnionFind(n) for i, j in pairs: uf.union(i, j) members_by_root: dict[int, list[int]] = defaultdict(list) for i in range(n): members_by_root[uf.find(i)].append(i) clusters: dict[int, list[int]] = {} for members in members_by_root.values(): rep = min(members) if key_fn is None else min(members, key=key_fn) clusters[rep] = sorted(members) keep = sorted(clusters.keys()) return keep, clusters def pick_representatives( n: int, pairs: Iterable[tuple[int, int]], key_fn: Callable[[int], Any] | None = None, ) -> list[int]: """Cluster items in [0, n) via `pairs`, then keep one rep per cluster. Representative is the cluster member with the smallest `key_fn(i)`. With no key_fn, falls back to smallest index ("first seen"). Returned indices are sorted ascending so caller can re-index in original order. Thin wrapper over `cluster_and_pick` that drops the cluster-membership map. """ keep, _ = cluster_and_pick(n, pairs, key_fn) return keep _WHITESPACE_RE = re.compile(r"\s+") def normalize_text(text: str) -> str: """Lowercase + collapse internal whitespace + strip.""" return _WHITESPACE_RE.sub(" ", text.lower()).strip()