"""IRRepairer — canonicalize near-miss ids in a QueryIR before validation. The planner LLM must copy opaque catalog ids (``c_<12hex>`` for columns, and the analogous ``t_``/``s_`` ids for tables/sources) verbatim into the IR. It occasionally drops, inserts, or mutates a single character — e.g. it emits ``c_b7489c7a4b5`` where the catalog holds ``c_b7489c7a4b5f``. The IRValidator then rejects the IR on an exact-match lookup, the planner retries, produces the same typo, and the whole analysis fails after N attempts. This stage sits *before* the validator and rewrites each unresolvable id to the catalog id within edit-distance 1 — but only when that match is **unique**. When zero or 2+ candidates match, it leaves the id untouched so the validator still fails loudly. It never guesses: the worst case is exactly the pre-repair behaviour (a hard validation error), never a silently-wrong column. Because every catalog id is a fixed-length hash (``c_`` + 12 hex), the candidate space is uniform and collisions within a single table are astronomically unlikely — see the ownership discussion in the ticket. The unique-or-refuse rule is the real safety guarantee regardless. The repairer is best-effort and pure: it returns a repaired *copy* plus the list of edits it made (for observability). It resolves source_id first, then table_id, then assembles the column set from the base table plus any joined tables (column_ids are globally unique, so the union is safe), then repairs every column reference. It never raises — anything it cannot resolve is left for the validator to report. """ from __future__ import annotations from dataclasses import dataclass from ...catalog.models import Catalog, Source, Table from .models import QueryIR @dataclass(frozen=True) class Repair: """One id rewrite the repairer applied.""" where: str # e.g. "select[2].column_id" from_id: str to_id: str class IRRepairer: """Rewrite near-miss ids in a QueryIR to their unique catalog match.""" def repair(self, ir: QueryIR, catalog: Catalog) -> tuple[QueryIR, list[Repair]]: """Return a repaired copy of `ir` and the list of edits applied. Safe by construction: an id is only rewritten when exactly one catalog id lies within edit-distance 1. Ambiguous or unresolvable ids are left as-is for the validator to reject. """ repairs: list[Repair] = [] ir = ir.model_copy(deep=True) # --- source_id ------------------------------------------------------- source = _find_source(catalog, ir.source_id) if source is None: fixed = _unique_near(ir.source_id, [s.source_id for s in catalog.sources]) if fixed is not None: repairs.append(Repair("source_id", ir.source_id, fixed)) ir.source_id = fixed source = _find_source(catalog, fixed) if source is None: # Unknown source and no unique fix — nothing else is resolvable. return ir, repairs known_table_ids = [t.table_id for t in source.tables] # --- table_id -------------------------------------------------------- if ir.table_id not in known_table_ids: fixed = _unique_near(ir.table_id, known_table_ids) if fixed is not None: repairs.append(Repair("table_id", ir.table_id, fixed)) ir.table_id = fixed # --- join target_table_id (before assembling columns) ---------------- for k, join in enumerate(ir.joins): if join.target_table_id not in known_table_ids: fixed = _unique_near(join.target_table_id, known_table_ids) if fixed is not None: repairs.append( Repair(f"joins[{k}].target_table_id", join.target_table_id, fixed) ) join.target_table_id = fixed # --- assemble the column set (base + joined tables) ------------------ # column_ids are globally unique (hash of table/col name), so a union of # every table in play is a safe candidate set for column repairs. tables_in_play: list[Table] = [] base = _find_table(source, ir.table_id) if base is not None: tables_in_play.append(base) for join in ir.joins: tgt = _find_table(source, join.target_table_id) if tgt is not None: tables_in_play.append(tgt) known_col_ids = [c.column_id for t in tables_in_play for c in t.columns] if not known_col_ids: return ir, repairs # --- column references ---------------------------------------------- select_aliases = {s.alias for s in ir.select if s.alias} for i, item in enumerate(ir.select): # AggSelect may carry column_id=None (COUNT(*)) — nothing to repair. if item.column_id is not None: fixed = _maybe_fix(item.column_id, known_col_ids) if fixed is not None: repairs.append(Repair(f"select[{i}].column_id", item.column_id, fixed)) item.column_id = fixed for i, f in enumerate(ir.filters): fixed = _maybe_fix(f.column_id, known_col_ids) if fixed is not None: repairs.append(Repair(f"filters[{i}].column_id", f.column_id, fixed)) f.column_id = fixed for i, col_id in enumerate(ir.group_by): fixed = _maybe_fix(col_id, known_col_ids) if fixed is not None: repairs.append(Repair(f"group_by[{i}]", col_id, fixed)) ir.group_by[i] = fixed for i, ob in enumerate(ir.order_by): # order_by may legitimately reference a select alias, not a column. # Never rewrite an alias reference. if ob.column_id in select_aliases: continue fixed = _maybe_fix(ob.column_id, known_col_ids) if fixed is not None: repairs.append(Repair(f"order_by[{i}].column_id", ob.column_id, fixed)) ob.column_id = fixed for k, join in enumerate(ir.joins): fixed = _maybe_fix(join.left_column_id, known_col_ids) if fixed is not None: repairs.append( Repair(f"joins[{k}].left_column_id", join.left_column_id, fixed) ) join.left_column_id = fixed fixed = _maybe_fix(join.right_column_id, known_col_ids) if fixed is not None: repairs.append( Repair(f"joins[{k}].right_column_id", join.right_column_id, fixed) ) join.right_column_id = fixed return ir, repairs # --------------------------------------------------------------------------- # Matching helpers # --------------------------------------------------------------------------- def _maybe_fix(value: str, known: list[str]) -> str | None: """Return the unique near-miss for `value`, or None if it needs no/ambiguous fix.""" if value in known: return None return _unique_near(value, known) def _unique_near(value: str, known: list[str]) -> str | None: """The single catalog id within edit-distance 1 of `value`, else None. Returns None when `value` already matches, when nothing is close, or when 2+ candidates are equally close (ambiguous — refuse to guess). """ if value in known: return None candidates = [k for k in known if _edit_distance_le_1(value, k)] return candidates[0] if len(candidates) == 1 else None def _edit_distance_le_1(a: str, b: str) -> bool: """True iff `a` and `b` are within Levenshtein distance 1. Covers the three single-character typos an LLM makes when copying a hash id: substitution (same length, one char differs), deletion (a is b with one char removed — this is the "dropped the trailing char" case), and insertion (a is b with one extra char). """ if a == b: return True la, lb = len(a), len(b) if abs(la - lb) > 1: return False if la == lb: return sum(1 for x, y in zip(a, b, strict=True) if x != y) == 1 # Lengths differ by exactly 1: check the shorter is the longer minus one char. shorter, longer = (a, b) if la < lb else (b, a) i = j = 0 edited = False while i < len(shorter) and j < len(longer): if shorter[i] == longer[j]: i += 1 j += 1 elif edited: return False else: edited = True j += 1 # consume one char from the longer string return True def _find_source(catalog: Catalog, source_id: str) -> Source | None: return next((s for s in catalog.sources if s.source_id == source_id), None) def _find_table(source: Source, table_id: str) -> Table | None: return next((t for t in source.tables if t.table_id == table_id), None)