fix cclf transpilation: cast types, concat_str, unique order_by, recursion limit

- transpile.py: raise recursion limit to 5000 for deeply nested pivot SQL
- cclf.py: cast to Float64 before arithmetic in _negate_if_cancelled,
  use nw.concat_str for bill_type_code, add order_by to unique(),
  use keep="any" for stg_beneficiary_demographics
- tests: reduce known transpile failures from 18 to 2, add known
  parse failures for pivot expressions
This commit is contained in:
kert
2026-03-11 14:45:11 -04:00
parent 3abb77904e
commit 36bb3eb62a
3 changed files with 43 additions and 35 deletions

View File

@@ -62,10 +62,11 @@ def _icd_indicator() -> nw.Expr:
def _negate_if_cancelled(col: str) -> nw.Expr: def _negate_if_cancelled(col: str) -> nw.Expr:
"""Negate amount when clm_adjsmt_type_cd = '1' (cancelled).""" """Negate amount when clm_adjsmt_type_cd = '1' (cancelled)."""
amt = nw.col(col).cast(nw.Float64)
return ( return (
nw.when(nw.col("clm_adjsmt_type_cd") == "1") nw.when(nw.col("clm_adjsmt_type_cd") == "1")
.then(nw.col(col) * -1) .then(amt * -1)
.otherwise(nw.col(col)) .otherwise(amt)
.alias(col) .alias(col)
) )
@@ -103,13 +104,13 @@ def stg_beneficiary_xref(
CCLF9 fields used: CRNT_NUM (current MBI), PRVS_NUM (previous MBI), CCLF9 fields used: CRNT_NUM (current MBI), PRVS_NUM (previous MBI),
PRVS_ID_EFCTV_DT (effective date of previous ID — used to pick latest). PRVS_ID_EFCTV_DT (effective date of previous ID — used to pick latest).
""" """
return ( return cclf__cclf9.unique(
cclf__cclf9.sort("prvs_id_efctv_dt", descending=True) subset=["prvs_num"],
.unique(subset=["prvs_num"], keep="first") keep="last",
.select( order_by=["prvs_id_efctv_dt"],
nw.col("crnt_num"), ).select(
nw.col("prvs_num"), nw.col("crnt_num"),
) nw.col("prvs_num"),
) )
@@ -505,10 +506,13 @@ def int_institutional_medical_claim(
nw.col("bene_ptnt_stus_cd").alias("discharge_disposition_code"), nw.col("bene_ptnt_stus_cd").alias("discharge_disposition_code"),
nw.lit(None).cast(nw.String).alias("place_of_service_code"), nw.lit(None).cast(nw.String).alias("place_of_service_code"),
# bill_type_code = fac + clsfctn + freq # bill_type_code = fac + clsfctn + freq
( nw.concat_str(
nw.col("clm_bill_fac_type_cd").cast(nw.String).fill_null("") [
+ nw.col("clm_bill_clsfctn_cd").cast(nw.String).fill_null("") nw.col("clm_bill_fac_type_cd").cast(nw.String).fill_null(""),
+ nw.col("clm_bill_freq_cd").cast(nw.String).fill_null("") nw.col("clm_bill_clsfctn_cd").cast(nw.String).fill_null(""),
nw.col("clm_bill_freq_cd").cast(nw.String).fill_null(""),
],
separator="",
).alias("bill_type_code"), ).alias("bill_type_code"),
nw.lit("ms-drg").alias("drg_code_type"), nw.lit("ms-drg").alias("drg_code_type"),
# DRG: rightmost 3 chars # DRG: rightmost 3 chars
@@ -1175,7 +1179,7 @@ def stg_beneficiary_demographics(
.with_columns( .with_columns(
nw.coalesce("crnt_num", "bene_mbi_id").alias("current_bene_mbi_id"), nw.coalesce("crnt_num", "bene_mbi_id").alias("current_bene_mbi_id"),
) )
.unique(subset=["current_bene_mbi_id"], keep="first") .unique(subset=["current_bene_mbi_id"], keep="any")
) )

View File

@@ -154,11 +154,19 @@ def transpile(
dict[str, str] dict[str, str]
Mapping of expression name to transpiled SQL string. Mapping of expression name to transpiled SQL string.
""" """
import sys
own_connection = False own_connection = False
if con is None: if con is None:
con = _schema_only_connection() con = _schema_only_connection()
own_connection = True own_connection = True
# Pivot/unpivot expressions can generate 100K+ char SQL whose
# deeply nested AST exceeds Python's default recursion limit
# during sqlglot parse, tree traversal, and code generation.
old_limit = sys.getrecursionlimit()
sys.setrecursionlimit(max(old_limit, 5000))
cache: dict[str, Any] = {} cache: dict[str, Any] = {}
result: dict[str, str] = {} result: dict[str, str] = {}
view_registry: dict[str, str] = {} view_registry: dict[str, str] = {}
@@ -209,6 +217,7 @@ def transpile(
) )
result[name] = target_sql result[name] = target_sql
finally: finally:
sys.setrecursionlimit(old_limit)
if own_connection: if own_connection:
con.close() con.close()

View File

@@ -102,30 +102,23 @@ def _flat_cases(all_transpiled):
# ── Test: every expression transpiles without error ─────────────── # ── Test: every expression transpiles without error ───────────────
# Expressions with known sqlglot limitations (deeply nested DuckDB SQL # Schema-only transpilation limitation: nw.concat requires identical
# from pivot/unpivot operations) or narwhals schema-only incompatibilities. # schemas, but the three medical_claim intermediate views produce
# The cclf pipeline cascades: once _int_diagnosis_pivot fails, all # different column counts in schema-only mode (DuckDB doesn't fully
# downstream expressions that depend on it (directly or transitively) # preserve dynamically-generated null columns from narwhals expressions).
# also fail. # Works correctly with real data.
_KNOWN_TRANSPILE_FAILURES = { _KNOWN_TRANSPILE_FAILURES = {
"cclf._stg_beneficiary_xref", "cclf.medical_claim",
# DuckDB and Polars produce different column suffixes after join:
# Polars adds _rev to right-side columns, DuckDB schema-only doesn't.
"cclf._int_institutional_medical_claim",
}
# Pivot expressions produce valid SQL but are too deeply nested for
# sqlglot to re-parse in a second pass (RecursionError).
_KNOWN_PARSE_FAILURES = {
"cclf._int_diagnosis_pivot", "cclf._int_diagnosis_pivot",
"cclf._int_procedure_pivot", "cclf._int_procedure_pivot",
"cclf._int_institutional_medical_claim",
"cclf._stg_physician_claim",
"cclf._int_physician_claim_adr",
"cclf._int_physician_medical_claim",
"cclf._stg_dme_claim",
"cclf._int_dme_claim_adr",
"cclf._int_dme_medical_claim",
"cclf.medical_claim",
"cclf._stg_pharmacy_claim",
"cclf._int_pharmacy_claim_adr",
"cclf.pharmacy_claim",
"cclf._stg_beneficiary_demographics",
"cclf.eligibility",
"provider_attribution._int_current_steps",
"provider_attribution._int_yearly_steps",
} }
@@ -177,6 +170,8 @@ class TestValidDatabricksSQL:
for name, sql in sql_map.items(): for name, sql in sql_map.items():
if sql.startswith("-- ERROR"): if sql.startswith("-- ERROR"):
continue continue
if name in _KNOWN_PARSE_FAILURES:
continue
try: try:
tree = sqlglot.parse_one(sql, dialect="databricks") tree = sqlglot.parse_one(sql, dialect="databricks")
assert isinstance(tree, (exp.Select, exp.Union)) assert isinstance(tree, (exp.Select, exp.Union))