#!/usr/bin/env python3
"""Run formal falsification and quasi-experimental tests for policy events.

The design is a stacked sector difference-in-differences around five policies
with pre-specified exposed CSI sectors. It does not claim that policy treatment
is exogenous.
"""

from __future__ import annotations

import json
import math
from pathlib import Path

import numpy as np
import pandas as pd


ROOT = Path(__file__).resolve().parents[1]
OUT = ROOT / "data" / "research" / "derived"
SEED = 20260731
PERMUTATIONS = 20000
BOOTSTRAPS = 10000

EVENTS = (
    ("战略性新兴产业决定", "2010-10-18", ("工业", "医药卫生", "信息技术", "电信业务", "公用事业")),
    ("中国制造2025", "2015-05-08", ("工业", "原材料", "信息技术")),
    ("新能源汽车产业规划", "2020-10-20", ("工业", "原材料", "可选消费", "公用事业")),
    ("碳达峰行动方案", "2021-10-24", ("能源", "原材料", "工业", "公用事业")),
    ("十四五数字经济规划", "2022-01-12", ("工业", "信息技术", "电信业务")),
)
SECTORS = ("能源", "原材料", "工业", "可选消费", "主要消费", "医药卫生", "金融地产", "信息技术", "电信业务", "公用事业")
HORIZONS = ((1, 6), (1, 12), (13, 24), (25, 36))


def load_returns() -> pd.DataFrame:
    close = pd.read_csv(OUT / "csi_sector_monthly_closes_2009_2025.csv")
    close["date"] = pd.to_datetime(close["date"])
    close = close.set_index("date").sort_index()
    logs = np.log(close.astype(float)).diff()
    excess = logs[list(SECTORS)].sub(logs["沪深300"], axis=0)
    excess.index = excess.index.to_period("M")
    excess = excess.groupby(level=0).last()
    return excess


def stacked_panel(excess: pd.DataFrame) -> pd.DataFrame:
    rows: list[dict] = []
    for event, date, treated in EVENTS:
        event_month = pd.Timestamp(date).to_period("M")
        for relative_month in range(-12, 37):
            month = event_month + relative_month
            if month not in excess.index:
                continue
            for sector in SECTORS:
                value = excess.at[month, sector]
                if pd.isna(value):
                    continue
                rows.append(
                    {
                        "事件": event,
                        "事件日期": date,
                        "月份": month.to_timestamp("M"),
                        "相对月": relative_month,
                        "行业": sector,
                        "政策暴露": int(sector in treated),
                        "月度对数相对收益": float(value),
                    }
                )
    return pd.DataFrame(rows)


def event_contrast(panel: pd.DataFrame, start: int, end: int, assignments: dict[str, set[str]] | None = None) -> dict[str, float]:
    output: dict[str, float] = {}
    for event, group in panel.groupby("事件"):
        assignment = assignments[event] if assignments else set(
            group.loc[group["政策暴露"] == 1, "行业"].unique()
        )
        pre = group[group["相对月"].between(-12, -1)]
        post = group[group["相对月"].between(start, end)]
        pre_t = pre[pre["行业"].isin(assignment)]["月度对数相对收益"].mean()
        pre_c = pre[~pre["行业"].isin(assignment)]["月度对数相对收益"].mean()
        post_t = post[post["行业"].isin(assignment)]["月度对数相对收益"].mean()
        post_c = post[~post["行业"].isin(assignment)]["月度对数相对收益"].mean()
        output[event] = float((post_t - post_c) - (pre_t - pre_c))
    return output


def pretrend_slope(panel: pd.DataFrame, assignments: dict[str, set[str]] | None = None) -> dict[str, float]:
    output: dict[str, float] = {}
    for event, group in panel[panel["相对月"].between(-12, -1)].groupby("事件"):
        assignment = assignments[event] if assignments else set(
            group.loc[group["政策暴露"] == 1, "行业"].unique()
        )
        by_month = group.assign(treated=group["行业"].isin(assignment)).groupby(
            ["相对月", "treated"]
        )["月度对数相对收益"].mean().unstack()
        difference = by_month[True] - by_month[False]
        x = difference.index.to_numpy(dtype=float)
        y = difference.to_numpy(dtype=float)
        output[event] = float(np.polyfit(x, y, 1)[0])
    return output


def random_assignments(rng: np.random.Generator) -> dict[str, set[str]]:
    output: dict[str, set[str]] = {}
    for event, _, treated in EVENTS:
        output[event] = set(rng.choice(SECTORS, size=len(treated), replace=False).tolist())
    return output


def contrast_from_vector(values: np.ndarray, treated_indices: np.ndarray) -> float:
    mask = np.zeros(len(values), dtype=bool)
    mask[treated_indices] = True
    return float(values[mask].mean() - values[~mask].mean())


def precomputed_event_vectors(
    panel: pd.DataFrame,
) -> tuple[dict[str, dict[tuple[int, int], np.ndarray]], dict[str, np.ndarray]]:
    horizon_vectors: dict[str, dict[tuple[int, int], np.ndarray]] = {}
    slope_vectors: dict[str, np.ndarray] = {}
    for event, group in panel.groupby("事件"):
        matrix = group.pivot(index="相对月", columns="行业", values="月度对数相对收益")
        matrix = matrix.reindex(columns=SECTORS)
        pre = matrix.loc[-12:-1].mean(axis=0).to_numpy(dtype=float)
        horizon_vectors[event] = {}
        for start, end in HORIZONS:
            post = matrix.loc[start:end].mean(axis=0).to_numpy(dtype=float)
            horizon_vectors[event][(start, end)] = post - pre
        x = matrix.loc[-12:-1].index.to_numpy(dtype=float)
        slopes = []
        for sector in SECTORS:
            slopes.append(float(np.polyfit(x, matrix.loc[-12:-1, sector].to_numpy(dtype=float), 1)[0]))
        slope_vectors[event] = np.asarray(slopes)
    return horizon_vectors, slope_vectors


def permutation_distribution(panel: pd.DataFrame) -> tuple[dict[tuple[int, int], np.ndarray], np.ndarray]:
    rng = np.random.default_rng(SEED)
    horizons = {window: np.empty(PERMUTATIONS) for window in HORIZONS}
    slopes = np.empty(PERMUTATIONS)
    horizon_vectors, slope_vectors = precomputed_event_vectors(panel)
    treated_count = {event: len(treated) for event, _, treated in EVENTS}
    for index in range(PERMUTATIONS):
        assignment_indices = {
            event: rng.choice(len(SECTORS), size=treated_count[event], replace=False)
            for event, _, _ in EVENTS
        }
        for window in HORIZONS:
            horizons[window][index] = np.mean(
                [
                    contrast_from_vector(horizon_vectors[event][window], assignment_indices[event])
                    for event, _, _ in EVENTS
                ]
            )
        slopes[index] = np.mean(
            [
                contrast_from_vector(slope_vectors[event], assignment_indices[event])
                for event, _, _ in EVENTS
            ]
        )
    return horizons, slopes


def bootstrap_events(values: dict[str, float]) -> tuple[float, float]:
    rng = np.random.default_rng(SEED + 1)
    array = np.array(list(values.values()), dtype=float)
    draws = rng.choice(array, size=(BOOTSTRAPS, len(array)), replace=True).mean(axis=1)
    return float(np.quantile(draws, 0.025)), float(np.quantile(draws, 0.975))


def bh_adjust(pvalues: list[float]) -> list[float]:
    values = np.asarray(pvalues, dtype=float)
    order = np.argsort(values)
    adjusted = np.empty_like(values)
    running = 1.0
    for rank_index in range(len(values) - 1, -1, -1):
        original_index = order[rank_index]
        rank = rank_index + 1
        candidate = values[original_index] * len(values) / rank
        running = min(running, candidate)
        adjusted[original_index] = min(1.0, running)
    return adjusted.tolist()


def main() -> None:
    panel = stacked_panel(load_returns())
    panel.to_csv(OUT / "policy_stacked_did_panel.csv", index=False, encoding="utf-8-sig")
    permutation, permutation_slopes = permutation_distribution(panel)

    rows: list[dict] = []
    pvalues: list[float] = []
    for start, end in HORIZONS:
        by_event = event_contrast(panel, start, end)
        estimate = float(np.mean(list(by_event.values())))
        low, high = bootstrap_events(by_event)
        distribution = permutation[(start, end)]
        pvalue = float((1 + np.sum(np.abs(distribution) >= abs(estimate))) / (1 + len(distribution)))
        pvalues.append(pvalue)
        rows.append(
            {
                "检验": f"发布后{start}至{end}月",
                "起始相对月": start,
                "结束相对月": end,
                "事件数": len(by_event),
                "处理行业事件数": int(
                    sum(len(treated) for _, _, treated in EVENTS)
                ),
                "月均差分估计": estimate,
                "年化近似": math.exp(estimate * 12) - 1,
                "事件级bootstrap_95%下限": low,
                "事件级bootstrap_95%上限": high,
                "置换检验双侧p值": pvalue,
                "结论边界": "行业暴露非随机；只作准实验关联检验，不识别政策因果效应",
            }
        )

    slopes = pretrend_slope(panel)
    slope_estimate = float(np.mean(list(slopes.values())))
    slope_p = float(
        (1 + np.sum(np.abs(permutation_slopes) >= abs(slope_estimate)))
        / (1 + len(permutation_slopes))
    )
    pvalues.append(slope_p)
    rows.append(
        {
            "检验": "发布前12个月处理-对照趋势斜率",
            "起始相对月": -12,
            "结束相对月": -1,
            "事件数": len(slopes),
            "处理行业事件数": int(sum(len(treated) for _, _, treated in EVENTS)),
            "月均差分估计": slope_estimate,
            "年化近似": np.nan,
            "事件级bootstrap_95%下限": bootstrap_events(slopes)[0],
            "事件级bootstrap_95%上限": bootstrap_events(slopes)[1],
            "置换检验双侧p值": slope_p,
            "结论边界": "若前趋势显著，平行趋势假设不成立",
        }
    )

    adjusted = bh_adjust(pvalues)
    for row, value in zip(rows, adjusted):
        row["BH校正p值"] = value
        row["5%水平"] = "拒绝零假设" if value < 0.05 else "不拒绝零假设"
    results = pd.DataFrame(rows)
    results.to_csv(OUT / "policy_stacked_did_results.csv", index=False, encoding="utf-8-sig")

    event_rows: list[dict] = []
    for start, end in HORIZONS:
        for event, estimate in event_contrast(panel, start, end).items():
            event_rows.append(
                {
                    "事件": event,
                    "窗口": f"{start}-{end}月",
                    "差分估计": estimate,
                    "样本角色": "事件异质性与bootstrap单位",
                }
            )
    pd.DataFrame(event_rows).to_csv(
        OUT / "policy_stacked_did_event_estimates.csv", index=False, encoding="utf-8-sig"
    )

    metadata = pd.DataFrame(
        [
            {
                "设计": "五项专项政策堆叠式行业差分",
                "处理组": "政策发布前按文件内容预先映射的中证一级行业",
                "对照组": "同一事件内其他中证一级行业",
                "结果变量": "行业月度对数收益减沪深300月度对数收益",
                "前期": "发布前12至前1月",
                "后期": "1-6、1-12、13-24、25-36月",
                "推断": f"{PERMUTATIONS}次行业标签置换；{BOOTSTRAPS}次事件级bootstrap；BH多重检验校正",
                "固定效应等价处理": "同事件同月份处理-对照差，再减事件发布前均值",
                "宏观共同冲击": "沪深300基准与同事件月份差分吸收；另提供月度宏观背景表",
                "关键限制": "行业暴露非随机、仅5个事件、10个宽行业、事件窗口重叠、价格指数不含股息",
                "证据等级": "准实验关联检验；不是已验证政策因果",
            }
        ]
    )
    metadata.to_csv(
        OUT / "policy_stacked_did_metadata.csv", index=False, encoding="utf-8-sig"
    )

    assert panel["事件"].nunique() == len(EVENTS)
    assert panel["行业"].nunique() == len(SECTORS)
    assert set(results["检验"]).issuperset({"发布后1至6月", "发布前12个月处理-对照趋势斜率"})
    assert results["置换检验双侧p值"].between(0, 1).all()
    assert results["BH校正p值"].between(0, 1).all()
    print("policy-causal-test assertions: OK")


if __name__ == "__main__":
    main()
