diff --git a/src/parent_child/pipeline.py b/src/parent_child/pipeline.py index 00fa25b..65a5883 100644 --- a/src/parent_child/pipeline.py +++ b/src/parent_child/pipeline.py @@ -1,5 +1,6 @@ from __future__ import annotations +import ast import json import pandas as pd import re @@ -28,38 +29,122 @@ _MODULE_DIR = Path(__file__).resolve().parent _CONSTANTS_DIR = _MODULE_DIR.parent / "constants" / "mappings" -def _build_provider_group_label(df: pd.DataFrame) -> pd.Series: - """Human-readable grouping label for reviewers. +_NAME_KEY_STOPWORDS = frozenset({"and", "of", "the", "a", "an"}) - Uses the cleaned provider name, title-cased. When the same name appears - with more than one TIN in the dataset, disambiguate by appending the TIN - (e.g. "Storr Medical Center - 123456789"). Falls back to grouping_key - for rows with no usable provider name. + +def _normalize_provider_name_key(s: str) -> str: + """Lowercase, strip punctuation, drop short connector words. Lets + "CHILDREN'S HOSPITALS AND CLINICS" and "CHILDREN S HOSPITALS CLINICS" + match as the same provider.""" + if not s: + return "" + s = re.sub(r"[^a-z0-9]+", " ", s.lower()) + return " ".join(t for t in s.split() if t and t not in _NAME_KEY_STOPWORDS) + + +def _build_provider_group_label(df: pd.DataFrame) -> pd.Series: + """Title-cased provider name for the Group column. + + Clusters rows by shared TIN, then folds no-TIN rows into a TIN cluster + when their normalized name matches uniquely. Picks the longest name + variant per cluster as canonical. Source columns are never written. """ name_col = "PROV_GROUP_NAME_FULL_cleaned" tin_col = "PROV_GROUP_TIN" + other_tin_col = "PROV_OTHER_TIN" if name_col not in df.columns: return df.get("grouping_key", pd.Series([""] * len(df), index=df.index)).astype( "string" ) - name = df[name_col].fillna("").astype(str).str.strip() - tin = ( - df[tin_col].fillna("").astype(str).str.strip() - if tin_col in df.columns - else pd.Series([""] * len(df), index=df.index) + raw_name = df[name_col].fillna("").astype(str).str.strip() + + n = len(df) + parent_arr = list(range(n)) + + def _find(i: int) -> int: + while parent_arr[i] != i: + parent_arr[i] = parent_arr[parent_arr[i]] + i = parent_arr[i] + return i + + def _union(i: int, j: int) -> None: + ri, rj = _find(i), _find(j) + if ri != rj: + parent_arr[ri] = rj + + group_tin_series = ( + df[tin_col] if tin_col in df.columns else pd.Series([None] * n, index=df.index) + ) + other_tin_series = ( + df[other_tin_col] + if other_tin_col in df.columns + else pd.Series([None] * n, index=df.index) ) - # Count distinct non-empty TINs per name; names with >1 TIN need - # disambiguation. - mask = (name != "") & (tin != "") - tins_per_name: Dict[str, int] = ( - pd.DataFrame({"_n": name[mask], "_t": tin[mask]}) - .groupby("_n")["_t"] - .nunique() - .to_dict() - ) + # Stage 1: union rows sharing any TIN. + row_tin_sets: List[set] = [] + tin_to_anchor_pos: Dict[str, int] = {} + for pos in range(n): + idx = df.index[pos] + tins = _parse_identifier_set(group_tin_series.loc[idx]) | _parse_identifier_set( + other_tin_series.loc[idx] + ) + row_tin_sets.append(tins) + for t in tins: + anchor = tin_to_anchor_pos.get(t) + if anchor is None: + tin_to_anchor_pos[t] = pos + else: + _union(anchor, pos) + + # Stage 2: fold no-TIN rows into a TIN cluster by normalized-name match. + # Skip if the name maps to multiple TIN clusters — those are likely + # different providers and shouldn't be merged. + norm_keys: List[str] = [ + _normalize_provider_name_key(raw_name.iloc[p]) for p in range(n) + ] + key_to_rows: Dict[str, List[int]] = {} + for pos, k in enumerate(norm_keys): + if k: + key_to_rows.setdefault(k, []).append(pos) + + for positions in key_to_rows.values(): + tin_roots = set() + no_tin_rows: List[int] = [] + for pos in positions: + if row_tin_sets[pos]: + tin_roots.add(_find(pos)) + else: + no_tin_rows.append(pos) + if len(tin_roots) == 1: + target = next(iter(tin_roots)) + for pos in no_tin_rows: + _union(pos, target) + elif not tin_roots and len(no_tin_rows) > 1: + anchor = no_tin_rows[0] + for pos in no_tin_rows[1:]: + _union(anchor, pos) + + # Pick canonical name per cluster: longest wins, frequency then alpha tiebreak. + cluster_name_counts: Dict[int, Dict[str, int]] = {} + for pos in range(n): + nm = raw_name.iloc[pos] + if not nm: + continue + root = _find(pos) + bucket = cluster_name_counts.setdefault(root, {}) + bucket[nm] = bucket.get(nm, 0) + 1 + + cluster_canonical: Dict[int, str] = { + root: min(counts.items(), key=lambda kv: (-len(kv[0]), -kv[1], kv[0]))[0] + for root, counts in cluster_name_counts.items() + } + + canonical_name = [ + cluster_canonical.get(_find(pos), raw_name.iloc[pos]) for pos in range(n) + ] fallback = ( df["grouping_key"].fillna("").astype(str) @@ -67,16 +152,11 @@ def _build_provider_group_label(df: pd.DataFrame) -> pd.Series: else pd.Series([""] * len(df), index=df.index) ) - def _label(n: str, t: str, fb: str) -> str: - if not n: - return fb - display = n.title() - if tins_per_name.get(n, 0) > 1 and t: - return f"{display} - {t}" - return display - return pd.Series( - [_label(n, t, fb) for n, t, fb in zip(name, tin, fallback)], + [ + canonical_name[pos].title() if canonical_name[pos] else fallback.iloc[pos] + for pos in range(n) + ], index=df.index, dtype="string", ) @@ -378,17 +458,10 @@ def build_grouping_string(child_row, grouping_cols): def _parse_identifier_set(value) -> set: - """Normalize a TIN/NPI field into a set of individual ID strings. + """Parse a TIN/NPI cell into a set of IDs. - Source data often stores identifiers as JSON-encoded lists - (e.g. `'["133757370", "840611484"]'`) because a single contract can - span multiple legal entities. Plain scalars (`'133757370'`) are also - supported. Returns an empty set for NaN / empty / "nan" values. - - Normalizing to a set lets downstream matchers use intersection - instead of string equality — two rows that share a TIN compare as - compatible even if one encodes it as `["A", "B"]` and the other as - `["B", "C", "D"]`. + Accepts JSON ['["a","b"]'], Python-repr ['['a','b']'], pipe-separated + ('a | b'), plain scalar, or empty/NaN. """ if value is None: return set() @@ -400,15 +473,18 @@ def _parse_identifier_set(value) -> set: s = str(value).strip() if not s or s.lower() == "nan": return set() - # JSON list (the common case from upstream extraction). if s.startswith("[") and s.endswith("]"): - try: - parsed = json.loads(s) + for loader in (json.loads, ast.literal_eval): + try: + parsed = loader(s) + except (json.JSONDecodeError, ValueError, SyntaxError): + continue if isinstance(parsed, list): return {str(v).strip() for v in parsed if str(v).strip()} - except (json.JSONDecodeError, ValueError): - pass - # Plain scalar — wrap as a singleton set. + if "|" in s: + parts = {p.strip() for p in s.split("|") if p.strip()} + if parts: + return parts return {s}