"""
05_missing_data_imputation.py
缺失数据插补模块：国别内线性插值 + 迭代式多重插补（MICE 思路）
Missing Data Imputation Module: within-country linear interpolation +
iterative multiple imputation (MICE-style)

用途 / Purpose:
    对合并后的原始面板（04_merge_panel.py 输出）执行两阶段插补：
    阶段一：对每个国家的时间序列做线性插值，填补短期缺口；
    阶段二：对阶段一后仍缺失的值，使用 scikit-learn 的
            IterativeImputer（近似 MICE 思路）跨变量迭代插补。

依赖 / Dependencies: pandas, numpy, scikit-learn
随机种子 / Random seed: 42
"""

import logging
from pathlib import Path

import numpy as np
import pandas as pd

try:
    from sklearn.experimental import enable_iterative_imputer  # noqa: F401
    from sklearn.impute import IterativeImputer
except ImportError:  # pragma: no cover
    IterativeImputer = None

logging.basicConfig(level=logging.INFO, format="%(asctime)s  %(message)s")
log = logging.getLogger("05_imputation")

DATA_DIR = Path("./data")
RANDOM_SEED = 42


def linear_interpolate_within_country(panel: pd.DataFrame, value_cols: list) -> pd.DataFrame:
    """Stage 1: within-country linear interpolation over the year axis."""
    panel = panel.sort_values(["iso3", "year"]).copy()
    for col in value_cols:
        panel[col] = panel.groupby("iso3")[col].apply(
            lambda s: s.interpolate(method="linear", limit_direction="both")
        ).reset_index(drop=True)
    return panel


def iterative_impute(panel: pd.DataFrame, value_cols: list, max_iter: int = 20) -> pd.DataFrame:
    """Stage 2: cross-variable iterative imputation (MICE-style)."""
    if IterativeImputer is None:
        log.warning("scikit-learn IterativeImputer unavailable; skipping stage 2.")
        return panel
    imputer = IterativeImputer(max_iter=max_iter, random_state=RANDOM_SEED)
    matrix = panel[value_cols].to_numpy(dtype=float)
    imputed = imputer.fit_transform(matrix)
    panel[value_cols] = imputed
    return panel


def main():
    panel = pd.read_csv(DATA_DIR / "panel_merged_raw.csv")
    id_cols = ["iso3", "year"]
    value_cols = [c for c in panel.columns if c not in id_cols]

    log.info("Stage 1: within-country linear interpolation over %d variables", len(value_cols))
    panel = linear_interpolate_within_country(panel, value_cols)

    missing_after_stage1 = panel[value_cols].isna().mean().mean()
    log.info("Mean missing rate after stage 1: %.4f", missing_after_stage1)

    log.info("Stage 2: iterative multiple imputation (MICE-style, max_iter=20, seed=%d)", RANDOM_SEED)
    panel = iterative_impute(panel, value_cols, max_iter=20)

    missing_after_stage2 = panel[value_cols].isna().mean().mean()
    log.info("Mean missing rate after stage 2: %.4f", missing_after_stage2)

    panel.to_csv(DATA_DIR / "panel_imputed.csv", index=False)
    log.info("Imputed panel saved -> data/panel_imputed.csv (%d x %d)", *panel.shape)


if __name__ == "__main__":
    main()
