"""
12_reliability_validity_tests.py
信度效度检验模块：Cronbach's alpha / AVE / HTMT
Reliability & Validity Testing Module: Cronbach's alpha, AVE, HTMT

用途 / Purpose:
    对三大维度（V/C/F）的构念做经典信度（内部一致性 Cronbach's alpha）
    与效度（收敛效度 AVE、区分效度 HTMT 比率）检验，判断指标聚合是否
    在统计上合理。参考阈值：alpha > 0.70 可接受，AVE > 0.50 可接受，
    HTMT < 0.85 表明区分效度良好。

依赖 / Dependencies: pandas, numpy
"""

import json
import logging
from pathlib import Path

import numpy as np
import pandas as pd

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

DATA_DIR = Path("./data")

DIMENSION_COLUMNS = {
    "V": [
        "gov_expenditure_pct_gdp__minmax",
        "social_protection_expenditure__minmax",
        "fiscal_balance_pct_gdp__minmax",
        "public_debt_pct_gdp__minmax",
        "tax_revenue_pct_gdp__minmax",
        "unemployment_benefit_coverage__minmax",
    ],
    "C": [
        "gdp_growth_pct__minmax",
        "gross_capital_formation_pct_gdp__minmax",
        "patents_residents__minmax",
        "rd_expenditure_pct_gdp__minmax",
        "labor_productivity_growth__minmax",
        "high_tech_exports_pct__minmax",
    ],
    "F": [
        "regulatory_quality__minmax",
        "government_effectiveness__minmax",
        "rule_of_law__minmax",
        "corruption_control__minmax",
        "bureaucracy_delay_index__minmax",
    ],
}


def cronbachs_alpha(df: pd.DataFrame, cols: list) -> float:
    sub = df[cols].dropna()
    k = len(cols)
    item_vars = sub.var(axis=0, ddof=1).sum()
    total_var = sub.sum(axis=1).var(ddof=1)
    if total_var == 0:
        return np.nan
    return (k / (k - 1)) * (1 - item_vars / total_var)


def average_variance_extracted(df: pd.DataFrame, cols: list) -> float:
    """近似 AVE：以指标与其维度均值构念得分的相关系数平方均值估计。"""
    sub = df[cols].dropna()
    construct = sub.mean(axis=1)
    loadings_sq = []
    for col in cols:
        corr = np.corrcoef(sub[col], construct)[0, 1]
        loadings_sq.append(corr ** 2)
    return float(np.mean(loadings_sq))


def htmt_ratio(df: pd.DataFrame, cols_a: list, cols_b: list) -> float:
    """Heterotrait-Monotrait 比率：跨维度指标相关均值 / 同维度指标相关均值。"""
    sub = df[cols_a + cols_b].dropna()
    hetero = [
        abs(sub[a].corr(sub[b])) for a in cols_a for b in cols_b
    ]
    mono_a = [abs(sub[cols_a[i]].corr(sub[cols_a[j]]))
              for i in range(len(cols_a)) for j in range(i + 1, len(cols_a))]
    mono_b = [abs(sub[cols_b[i]].corr(sub[cols_b[j]]))
              for i in range(len(cols_b)) for j in range(i + 1, len(cols_b))]
    mono = mono_a + mono_b
    if not mono:
        return np.nan
    return float(np.mean(hetero) / np.mean(mono))


def main():
    panel = pd.read_csv(DATA_DIR / "panel_normalized_minmax.csv")

    report = {"cronbachs_alpha": {}, "ave": {}}
    for dim, cols in DIMENSION_COLUMNS.items():
        available = [c for c in cols if c in panel.columns]
        alpha = cronbachs_alpha(panel, available)
        ave = average_variance_extracted(panel, available)
        report["cronbachs_alpha"][dim] = round(float(alpha), 3)
        report["ave"][dim] = round(float(ave), 3)
        log.info("%s: Cronbach's alpha=%.3f, AVE=%.3f", dim, alpha, ave)

    report["htmt"] = {}
    dims = list(DIMENSION_COLUMNS.keys())
    for i in range(len(dims)):
        for j in range(i + 1, len(dims)):
            d1, d2 = dims[i], dims[j]
            ratio = htmt_ratio(panel, DIMENSION_COLUMNS[d1], DIMENSION_COLUMNS[d2])
            report["htmt"][f"{d1}-{d2}"] = round(float(ratio), 3)
            log.info("HTMT(%s, %s) = %.3f (threshold < 0.85 for discriminant validity)", d1, d2, ratio)

    with open(DATA_DIR / "reliability_validity_report.json", "w", encoding="utf-8") as f:
        json.dump(report, f, ensure_ascii=False, indent=2)
    log.info("Reliability & validity report saved -> data/reliability_validity_report.json")


if __name__ == "__main__":
    main()
