"""Home Credit all-column, one-level case aggregation pipeline.

Raw train columns with missing rate > 0.8 are removed before cleaning.
The train keep-list and p99 caps are reused unchanged for test.
Repeated tables aggregate directly by case_id once:
numeric/date-days -> max,last,mean; categorical -> last,mode.
"""

import gc
import glob
import os
from pathlib import Path

import duckdb
import numpy as np
import pandas as pd
import polars as pl
from IPython.display import display
from lightgbm import LGBMClassifier


ROOT_CANDIDATES = [
    Path("/kaggle/input/competitions/home-credit-credit-risk-model-stability"),
    Path("/kaggle/input/home-credit-credit-risk-model-stability"),
]
ROOT = next(
    (
        path
        for path in ROOT_CANDIDATES
        if (path / "parquet_files/train/train_base.parquet").exists()
    ),
    ROOT_CANDIDATES[0],
)
TRAIN_DIR = ROOT / "parquet_files" / "train"
TEST_DIR = ROOT / "parquet_files" / "test"

MISSING_THRESHOLD = 0.8
P99_QUANTILE = 0.99
A2_BATCH_WEIGHT = 8
POSTPROCESS_YEAR_COL = "credit_bureau_a_2__pmts_year_1139T__max"
YEAR_ADJUSTMENTS = {2020: 0.07, 2021: 0.06, 2022: 0.02}

DEPTH0_TABLES = ["static_0", "static_cb_0"]
REPEATED_TABLES = [
    "applprev_1",
    "tax_registry_a_1",
    "credit_bureau_a_1",
    "credit_bureau_a_2",
    "person_1",
    "other_1",
]
GROUP_KEYS = {"case_id", "num_group1", "num_group2"}
TRAIN_KEEP_COLUMNS = {}
TRAIN_P99_CAPS = {}


def table_paths(data_dir, split, table):
    exact = Path(data_dir) / f"{split}_{table}.parquet"
    if exact.exists():
        return [str(exact)]
    paths = sorted(
        glob.glob(str(Path(data_dir) / f"{split}_{table}_*.parquet"))
    )
    if not paths:
        raise FileNotFoundError(f"No files found for {split}_{table}")
    return paths


def read_polars_table(data_dir, split, table):
    frames = [pl.scan_parquet(path) for path in table_paths(data_dir, split, table)]
    return pl.concat(frames, how="vertical_relaxed").collect()


def duckdb_source(data_dir, split, table):
    paths = table_paths(data_dir, split, table)
    path = (
        paths[0]
        if len(paths) == 1
        else str(Path(data_dir) / f"{split}_{table}_*.parquet")
    )
    path = path.replace("'", "''")
    return f"read_parquet('{path}', union_by_name=true, filename=true)"


def qi(name):
    return '"' + name.replace('"', '""') + '"'


def normalize_nan(df):
    cols = [
        name
        for name, dtype in df.schema.items()
        if dtype in (pl.Float32, pl.Float64)
    ]
    return df.with_columns([pl.col(c).fill_nan(None) for c in cols]) if cols else df


def filter_polars_raw(df, table, split, protected):
    df = normalize_nan(df)
    protected = [c for c in protected if c in df.columns]
    candidates = [c for c in df.columns if c not in protected]
    if split == "train":
        rates = (
            df.select([pl.col(c).is_null().mean().alias(c) for c in candidates]).row(0)
            if candidates
            else []
        )
        keep = [
            c
            for c, rate in zip(candidates, rates)
            if rate is not None and float(rate) <= MISSING_THRESHOLD
        ]
        TRAIN_KEEP_COLUMNS[table] = keep
        print(
            f"  [{table}] missing filter kept={len(keep)}, "
            f"dropped={len(candidates) - len(keep)}"
        )
    else:
        if table not in TRAIN_KEEP_COLUMNS:
            raise RuntimeError(f"Missing train keep-list for {table}")
        keep = TRAIN_KEEP_COLUMNS[table]
        absent = [c for c in keep if c not in df.columns]
        if absent:
            raise ValueError(f"{table} test is missing train columns: {absent[:30]}")
    return df.select(protected + keep)


def parse_date_expr(name, dtype):
    if dtype == pl.Date:
        return pl.col(name)
    if dtype == pl.Datetime:
        return pl.col(name).cast(pl.Date)
    return (
        pl.col(name)
        .cast(pl.String)
        .str.strptime(pl.Date, strict=False)
    )


def polars_transform(df, table, split, base_dates=None):
    features = [c for c in df.columns if c not in GROUP_KEYS]
    dates = [c for c in features if c.endswith("D")]
    exprs = []
    for c in features:
        dtype = df.schema[c]
        if c.endswith("D"):
            exprs.append(parse_date_expr(c, dtype).alias(c))
        elif c.endswith(("A", "P")):
            exprs.append(pl.col(c).cast(pl.Float64, strict=False).alias(c))
        elif c.endswith("M") or dtype == pl.Boolean:
            exprs.append(pl.col(c).cast(pl.String).alias(c))
    if exprs:
        df = df.with_columns(exprs)
    df = normalize_nan(df)

    if dates:
        if base_dates is None:
            raise ValueError(f"{table} has date columns but no decision dates")
        df = df.join(base_dates, on="case_id", how="left")
        days_exprs = []
        for c in dates:
            raw = (pl.col("date_decision") - pl.col(c)).dt.total_days()
            days_exprs.append(
                pl.when(raw < 0)
                .then(None)
                .otherwise(raw)
                .cast(pl.Int64)
                .alias(f"{c}_days")
            )
        df = df.with_columns(days_exprs).drop(dates + ["date_decision"])

    numeric = [
        c
        for c, dtype in df.schema.items()
        if c not in GROUP_KEYS and dtype.is_numeric()
    ]
    clean = [
        pl.when(pl.col(c) < 0).then(None).otherwise(pl.col(c)).alias(c)
        for c in numeric
        if c.endswith(("A", "P"))
    ]
    if clean:
        df = df.with_columns(clean)

    if split == "train":
        values = (
            df.select([pl.col(c).quantile(P99_QUANTILE).alias(c) for c in numeric]).row(0)
            if numeric
            else []
        )
        caps = {
            c: float(cap)
            for c, cap in zip(numeric, values)
            if cap is not None and np.isfinite(cap) and abs(float(cap)) >= 1e-9
        }
        TRAIN_P99_CAPS[table] = caps
    else:
        if table not in TRAIN_P99_CAPS:
            raise RuntimeError(f"Missing train p99 caps for {table}")
        caps = TRAIN_P99_CAPS[table]

    cap_exprs = [
        pl.when(pl.col(c) > cap)
        .then(pl.lit(cap))
        .otherwise(pl.col(c))
        .alias(c)
        for c, cap in caps.items()
        if c in df.columns
    ]
    return df.with_columns(cap_exprs) if cap_exprs else df


def build_base(data_dir, split):
    base = read_polars_table(data_dir, split, "base")
    protected = ["case_id", "WEEK_NUM", "date_decision"]
    if "target" in base.columns:
        protected.append("target")
    base = filter_polars_raw(base, "base", split, protected)
    base = base.with_columns(
        parse_date_expr("date_decision", base.schema["date_decision"])
        .alias("date_decision")
    )
    extras = [c for c in base.columns if c not in protected]
    if extras:
        base = base.rename({c: f"base__{c}" for c in extras})
    if base["case_id"].n_unique() != base.height:
        raise ValueError(f"{split} base contains duplicate case_id")
    return base


def build_depth0(data_dir, split, table, base_dates):
    df = read_polars_table(data_dir, split, table)
    protected = [
        c for c in ["case_id", "num_group1", "num_group2"] if c in df.columns
    ]
    df = filter_polars_raw(df, table, split, protected)
    df = polars_transform(df, table, split, base_dates)
    if df["case_id"].n_unique() != df.height:
        raise ValueError(f"Depth-0 table {table} is not one row per case_id")
    drop = [c for c in ["num_group1", "num_group2"] if c in df.columns]
    if drop:
        df = df.drop(drop)
    return df.rename(
        {c: f"{table}__{c}" for c in df.columns if c != "case_id"}
    )


def duck_schema(con, source):
    rows = con.execute(f"DESCRIBE SELECT * FROM {source}").fetchall()
    return [(row[0], row[1].upper()) for row in rows if row[0] != "filename"]


def sql_float(dtype):
    return dtype in {"FLOAT", "REAL", "DOUBLE"}


def sql_numeric(dtype):
    prefixes = (
        "TINYINT", "SMALLINT", "INTEGER", "BIGINT", "HUGEINT",
        "UTINYINT", "USMALLINT", "UINTEGER", "UBIGINT",
        "FLOAT", "REAL", "DOUBLE", "DECIMAL",
    )
    return dtype.startswith(prefixes)


def sql_date(dtype, name):
    return (
        name.endswith("D")
        or dtype.startswith("DATE")
        or dtype.startswith("TIMESTAMP")
    )


def missing_condition(name, dtype):
    col = qi(name)
    return f"({col} IS NULL OR isnan({col}))" if sql_float(dtype) else f"({col} IS NULL)"


def clean_numeric(name, dtype):
    col = qi(name)
    value = f"CASE WHEN isnan({col}) THEN NULL ELSE {col} END" if sql_float(dtype) else col
    if name.endswith(("A", "P")):
        value = f"CASE WHEN ({value}) < 0 THEN NULL ELSE ({value}) END"
    return f"CAST(({value}) AS DOUBLE)"


def prepare_duck_rules(con, source, table, split, schema):
    schema_map = dict(schema)
    protected = [
        c for c in ["case_id", "num_group1", "num_group2"] if c in schema_map
    ]
    candidates = [c for c, _ in schema if c not in protected]
    if split == "train":
        if candidates:
            query = ", ".join(
                f"AVG(CASE WHEN {missing_condition(c, schema_map[c])} "
                f"THEN 1.0 ELSE 0.0 END) AS {qi(c)}"
                for c in candidates
            )
            rates = con.execute(f"SELECT {query} FROM {source}").fetchone()
            keep = [
                c
                for c, rate in zip(candidates, rates)
                if rate is not None and float(rate) <= MISSING_THRESHOLD
            ]
        else:
            keep = []
        TRAIN_KEEP_COLUMNS[table] = keep
        print(
            f"  [{table}] missing filter kept={len(keep)}, "
            f"dropped={len(candidates) - len(keep)}"
        )
        numeric = [
            c
            for c in keep
            if sql_numeric(schema_map[c]) and not sql_date(schema_map[c], c)
        ]
        caps = {}
        if numeric:
            query = ", ".join(
                f"approx_quantile({clean_numeric(c, schema_map[c])}, "
                f"{P99_QUANTILE}) AS {qi(c)}"
                for c in numeric
            )
            values = con.execute(f"SELECT {query} FROM {source}").fetchone()
            caps = {
                c: float(cap)
                for c, cap in zip(numeric, values)
                if cap is not None and np.isfinite(cap) and abs(float(cap)) >= 1e-9
            }
        TRAIN_P99_CAPS[table] = caps
    else:
        if table not in TRAIN_KEEP_COLUMNS or table not in TRAIN_P99_CAPS:
            raise RuntimeError(f"Missing train rules for {table}")
        keep = TRAIN_KEEP_COLUMNS[table]
        caps = TRAIN_P99_CAPS[table]
        absent = [c for c in keep if c not in schema_map]
        if absent:
            raise ValueError(f"{table} test is missing train columns: {absent[:30]}")
    return keep, caps


def transformed_expression(name, dtype, cap):
    col = qi(name)
    if sql_date(dtype, name):
        output = f"{name}_days"
        days = (
            f"CAST(date_diff('day', TRY_CAST({col} AS DATE), "
            f"b.date_decision) AS BIGINT)"
        )
        expr = f"CASE WHEN ({days}) < 0 THEN NULL ELSE ({days}) END"
        return expr, output, "numeric"
    if sql_numeric(dtype):
        expr = clean_numeric(name, dtype)
        if cap is not None:
            expr = (
                f"CASE WHEN ({expr}) > {repr(float(cap))} "
                f"THEN {repr(float(cap))} ELSE ({expr}) END"
            )
        return expr, name, "numeric"
    return f"CAST({col} AS VARCHAR)", name, "categorical"


def build_repeated(data_dir, split, table, base_dates):
    con = duckdb.connect()
    con.execute("SET memory_limit='8GB'")
    con.execute("SET preserve_insertion_order=true")
    con.execute("SET threads TO 2")
    con.register("base_dates", base_dates.to_arrow())

    source = duckdb_source(data_dir, split, table)
    schema = duck_schema(con, source)
    schema_map = dict(schema)
    if "case_id" not in schema_map:
        con.close()
        raise ValueError(f"{table} has no case_id")

    keep, caps = prepare_duck_rules(con, source, table, split, schema)
    selected = ["t.case_id"]
    aggregates = []
    aggregate_groups = []

    for c in keep:
        expr, output, kind = transformed_expression(c, schema_map[c], caps.get(c))
        selected.append(f"{expr} AS {qi(output)}")
        value = qi(output)
        prefix = f"{table}__{output}"
        if kind == "numeric":
            group_aggregates = [
                f"MAX({value}) AS {qi(prefix + '__max')}",
                f"LAST({value}) AS {qi(prefix + '__last')}",
                f"AVG({value}) AS {qi(prefix + '__mean')}",
            ]
        else:
            group_aggregates = [
                f"LAST({value}) AS {qi(prefix + '__last')}",
                (
                    f"MODE({value}) "
                    f"FILTER (WHERE {value} IS NOT NULL) "
                    f"AS {qi(prefix + '__mode')}"
                ),
            ]
        aggregates.extend(group_aggregates)
        aggregate_groups.append((kind, group_aggregates))

    if not aggregates:
        con.close()
        raise ValueError(f"No features survived for {table}")

    selected_sql = ",\n                ".join(selected)
    transformed = f"""
        SELECT
                {selected_sql}
        FROM {source} t
        LEFT JOIN base_dates b USING (case_id)
    """

    checkpoint = None
    if table == "credit_bureau_a_2":
        os.makedirs(".tmp", exist_ok=True)
        checkpoint = f".tmp/{table}_{split}_filtered.parquet"
        if os.path.exists(checkpoint):
            os.remove(checkpoint)
        escaped = checkpoint.replace("'", "''")
        con.execute(
            f"COPY ({transformed}) TO '{escaped}' "
            f"(FORMAT PARQUET, COMPRESSION ZSTD)"
        )
        relation = f"read_parquet('{escaped}')"
    else:
        relation = f"({transformed})"

    batch_paths = []
    if table == "credit_bureau_a_2":
        batches = []
        current_batch = []
        current_weight = 0
        for kind, expressions in aggregate_groups:
            # MODE keeps a much larger per-case state than numeric aggregates.
            weight = 4 if kind == "categorical" else 1
            if current_batch and current_weight + weight > A2_BATCH_WEIGHT:
                batches.append(current_batch)
                current_batch = []
                current_weight = 0
            current_batch.extend(expressions)
            current_weight += weight
        if current_batch:
            batches.append(current_batch)

        for batch_index, batch_aggregates in enumerate(batches):
            batch_path = f".tmp/{table}_{split}_batch_{batch_index}.parquet"
            if os.path.exists(batch_path):
                os.remove(batch_path)
            batch_paths.append(batch_path)
            escaped_batch = batch_path.replace("'", "''")
            batch_sql = ",\n                    ".join(batch_aggregates)
            con.execute(
                f"""
                COPY (
                    SELECT
                        case_id,
                        {batch_sql}
                    FROM {relation}
                    GROUP BY case_id
                ) TO '{escaped_batch}'
                (FORMAT PARQUET, COMPRESSION ZSTD)
                """
            )
            print(
                f"    {table} batch {batch_index + 1}/{len(batches)} "
                f"completed"
            )

        joined_batches = f"read_parquet('{batch_paths[0]}') b0"
        for batch_index, batch_path in enumerate(batch_paths[1:], start=1):
            joined_batches += (
                f" LEFT JOIN read_parquet('{batch_path}') b{batch_index} "
                f"USING (case_id)"
            )
        result = con.execute(f"SELECT * FROM {joined_batches}").pl()
    else:
        aggregate_sql = ",\n            ".join(aggregates)
        result = con.execute(
            f"""
            SELECT
                case_id,
                {aggregate_sql}
            FROM {relation}
            GROUP BY case_id
            """
        ).pl()
    con.close()
    if checkpoint and os.path.exists(checkpoint):
        os.remove(checkpoint)
    for batch_path in batch_paths:
        if os.path.exists(batch_path):
            os.remove(batch_path)

    result = normalize_nan(result)
    float64_cols = [
        name
        for name, dtype in result.schema.items()
        if dtype == pl.Float64
    ]
    if float64_cols:
        result = result.with_columns(
            [pl.col(name).cast(pl.Float32) for name in float64_cols]
        )
    if result["case_id"].n_unique() != result.height:
        raise ValueError(f"{table} aggregation produced duplicate case_id")
    return result


def build_dataset(data_dir, split):
    print(f"=== building {split} ===")
    base = build_base(data_dir, split)
    base_dates = base.select("case_id", "date_decision")
    output = base

    for table in DEPTH0_TABLES:
        print(f"  depth-0 {table}...")
        part = build_depth0(data_dir, split, table, base_dates)
        output = output.join(part, on="case_id", how="left")
        del part
        gc.collect()

    for table in REPEATED_TABLES:
        print(f"  one-level case aggregation {table}...")
        part = build_repeated(data_dir, split, table, base_dates)
        output = output.join(part, on="case_id", how="left")
        del part
        gc.collect()

    if output["case_id"].n_unique() != output.height:
        raise ValueError(f"{split} assembled data contains duplicate case_id")
    print(f"  {split} shape: {output.shape}")
    return output


def main():
    train = build_dataset(TRAIN_DIR, "train")
    test = build_dataset(TEST_DIR, "test")

    if POSTPROCESS_YEAR_COL not in train.columns:
        raise ValueError(
            f"Year-tip feature was filtered or not generated: "
            f"{POSTPROCESS_YEAR_COL}"
        )
    if POSTPROCESS_YEAR_COL not in test.columns:
        raise ValueError(f"Test is missing {POSTPROCESS_YEAR_COL}")

    excluded = {"case_id", "target", "WEEK_NUM", "date_decision"}
    feature_cols = [c for c in train.columns if c not in excluded]
    if not feature_cols:
        raise ValueError("No model features were generated")

    required_test = {"case_id", "WEEK_NUM", "date_decision", *feature_cols}
    missing = sorted(required_test - set(test.columns))
    extra = sorted(set(test.columns) - required_test)
    if missing or extra:
        raise ValueError(
            f"Train/test feature mismatch; missing={missing[:30]}, "
            f"extra={extra[:30]}"
        )

    x_train = train.select(feature_cols).to_pandas()
    y_train = train["target"].to_numpy()
    x_test = test.select(["case_id"] + feature_cols).to_pandas()
    if not x_test["case_id"].is_unique:
        raise ValueError("Test contains duplicate case_id")
    x_test = x_test.set_index("case_id")
    if list(x_train.columns) != list(x_test.columns):
        raise ValueError("Train/test feature names or order differ")

    postprocess_year = x_test[POSTPROCESS_YEAR_COL].copy()
    categorical = list(
        x_train.select_dtypes(include=["object", "string", "category"]).columns
    )
    for c in categorical:
        train_values = x_train[c].astype("string")
        test_values = x_test[c].astype("string")
        categories = pd.Index(
            pd.concat([train_values, test_values], ignore_index=True)
            .dropna()
            .unique()
        )
        dtype = pd.CategoricalDtype(categories=categories)
        x_train[c] = train_values.astype(dtype)
        x_test[c] = test_values.astype(dtype)

    model = LGBMClassifier(
        objective="binary",
        boosting_type="gbdt",
        n_estimators=500,
        learning_rate=0.05,
        reg_alpha=0.5,
        reg_lambda=4,
        max_depth=5,
        num_leaves=31,
        random_state=42,
        n_jobs=-1,
    )
    model.fit(x_train, y_train)

    probability = model.predict_proba(x_test)[:, 1]
    if not np.isfinite(probability).all():
        raise ValueError("Predictions contain NaN or infinity")
    if ((probability < 0) | (probability > 1)).any():
        raise ValueError("Predictions are outside [0, 1]")

    scores = pd.Series(probability, index=x_test.index, name="score")
    adjusted_counts = {}
    for year, deduction in YEAR_ADJUSTMENTS.items():
        mask = postprocess_year.eq(year)
        adjusted_counts[year] = int(mask.sum())
        scores.loc[mask] = (scores.loc[mask] - deduction).clip(lower=0)

    if not np.isfinite(scores.to_numpy()).all():
        raise ValueError("Adjusted scores contain NaN or infinity")
    if ((scores < 0) | (scores > 1)).any():
        raise ValueError("Adjusted scores are outside [0, 1]")

    sample = pd.read_csv(ROOT / "sample_submission.csv")
    if not {"case_id", "score"}.issubset(sample.columns):
        raise ValueError("sample_submission must contain case_id and score")
    if not sample["case_id"].is_unique:
        raise ValueError("sample_submission has duplicate case_id")
    if len(sample) != len(scores) or set(sample["case_id"]) != set(scores.index):
        raise ValueError("sample_submission and prediction case_id differ")

    submission = sample[["case_id"]].copy()
    submission["score"] = submission["case_id"].map(scores)
    if submission["score"].isna().any():
        raise ValueError("Submission has null or unmatched scores")
    if list(submission.columns) != ["case_id", "score"]:
        raise ValueError("Submission must contain only case_id and score")

    print(f"Using {len(feature_cols):,} model features")
    print(f"Year adjustments applied to: {adjusted_counts}")
    display(submission.head())
    submission.to_csv("submission.csv", index=False)


if __name__ == "__main__":
    main()
