#!/usr/bin/env python3
"""Leakage-aware daily rotation study for the three China T+0 ETF mirrors.

This is intentionally a *daily/overnight* study.  It does not claim to validate
the intraday strategy because the supplied archives contain no ETF minute bars.
Signals are formed after the close and can only rebalance at the next open.
"""

from __future__ import annotations

import argparse
import collections
import datetime as dt
import json
import math
import statistics
import tarfile
import tempfile
from dataclasses import dataclass, field
from pathlib import Path

import pyarrow as pa
import pyarrow.compute as pc
import pyarrow.parquet as pq


TARGETS = ("510900.SH", "511010.SH", "518880.SH")
NAMES = {
    "510900.SH": "H股ETF",
    "511010.SH": "国债ETF",
    "518880.SH": "黄金ETF",
}


def finite(value, default=0.0):
    try:
        number = float(value)
        return number if math.isfinite(number) else default
    except (TypeError, ValueError):
        return default


def metrics(returns: list[float], dates: list[str], start: str, end: str):
    values = [value for value, date in zip(returns, dates) if start <= date <= end]
    selected_dates = [date for date in dates if start <= date <= end]
    if not values:
        return {"observations": 0, "annualizedReturn": 0, "sharpe": 0, "maxDrawdown": 0, "dailyWinRate": 0, "positiveYearRate": 0}
    equity = peak = 1.0
    max_drawdown = 0.0
    yearly = collections.defaultdict(lambda: 1.0)
    for value, date in zip(values, selected_dates):
        equity *= max(0.01, 1 + value)
        peak = max(peak, equity)
        max_drawdown = min(max_drawdown, equity / peak - 1)
        yearly[date[:4]] *= 1 + value
    volatility = statistics.stdev(values) * math.sqrt(252) if len(values) > 1 else 0
    annualized = equity ** (252 / len(values)) - 1 if equity > 0 else -1
    return {
        "observations": len(values),
        "totalReturn": round(equity - 1, 6),
        "annualizedReturn": round(annualized, 6),
        "annualizedVolatility": round(volatility, 6),
        "sharpe": round(statistics.mean(values) * 252 / volatility, 6) if volatility else 0,
        "maxDrawdown": round(max_drawdown, 6),
        "dailyWinRate": round(sum(value > 0 for value in values) / len(values), 6),
        "positiveYearRate": round(sum(value > 1 for value in yearly.values()) / len(yearly), 6),
        "yearly": {year: round(value - 1, 6) for year, value in sorted(yearly.items())},
    }


def score(item):
    return item["annualizedReturn"] + 0.12 * item["sharpe"] - 0.70 * abs(item["maxDrawdown"])


def load_daily(archive_path: Path):
    by_date: dict[str, dict[str, dict[str, float]]] = {}
    with tempfile.TemporaryDirectory(prefix="quant-atlas-etf-rotation-") as temporary:
        root = Path(temporary)
        with tarfile.open(archive_path, "r:gz") as archive:
            members = [member for member in archive.getmembers() if member.isfile() and member.name.endswith(".parquet")]
            archive.extractall(root, members=members, filter="data")
        for path in sorted(root.glob("**/*.parquet")):
            table = pq.read_table(path, columns=["ts_code", "trade_date", "pre_close", "open", "close", "amount"])
            table = table.filter(pc.is_in(table["ts_code"], value_set=pa.array(TARGETS)))
            for row in table.to_pylist():
                code = row["ts_code"]
                raw_date = str(row["trade_date"])
                date = f"{raw_date[:4]}-{raw_date[4:6]}-{raw_date[6:8]}"
                pre_close = finite(row["pre_close"])
                open_price = finite(row["open"])
                close = finite(row["close"])
                if code not in TARGETS or min(pre_close, open_price, close) <= 0:
                    continue
                by_date.setdefault(date, {})[code] = {
                    "preClose": pre_close,
                    "open": open_price,
                    "close": close,
                    "amount": finite(row["amount"]),
                }
    return by_date


@dataclass(frozen=True)
class Config:
    fast: int
    slow: int
    rebalance: int
    max_assets: int
    exposure: float
    vol_penalty: float

    @property
    def identifier(self):
        return f"ETF-D-{self.fast}-{self.slow}-R{self.rebalance}-K{self.max_assets}-E{self.exposure:g}-VP{self.vol_penalty:g}"


@dataclass
class Simulation:
    config: Config
    active: dict[str, float] = field(default_factory=dict)
    pending: dict[str, float] | None = None
    returns: list[float] = field(default_factory=list)
    turnovers: list[float] = field(default_factory=list)
    equity: float = 1.0
    peak: float = 1.0
    cooldown: int = 0


def drift(weights: dict[str, float], factors: dict[str, float]):
    gross_return = sum(weight * (factors.get(code, 1.0) - 1) for code, weight in weights.items())
    divisor = 1 + gross_return
    if divisor > 0:
        weights = {code: weight * factors.get(code, 1.0) / divisor for code, weight in weights.items()}
    return weights, gross_return


def run_simulation(config: Config, dates: list[str], data: dict[str, dict[str, dict]], cost_bps: float):
    sim = Simulation(config)
    histories = {code: collections.deque(maxlen=201) for code in TARGETS}
    returns = {code: collections.deque(maxlen=60) for code in TARGETS}
    for date_index, date in enumerate(dates):
        rows = data[date]
        overnight = {code: row["open"] / row["preClose"] for code, row in rows.items()}
        sim.active, overnight_return = drift(sim.active, overnight)

        turnover = 0.0
        if sim.pending is not None:
            executed = dict(sim.pending)
            turnover = sum(abs(sim.active.get(code, 0) - executed.get(code, 0)) for code in set(sim.active) | set(executed))
            sim.active = executed
            sim.pending = None

        intraday = {code: row["close"] / row["open"] for code, row in rows.items()}
        sim.active, intraday_return = drift(sim.active, intraday)
        net_return = max(-0.25, overnight_return + intraday_return - turnover * cost_bps / 10_000)
        sim.returns.append(net_return)
        sim.turnovers.append(turnover)
        sim.equity *= 1 + net_return
        sim.peak = max(sim.peak, sim.equity)
        drawdown = sim.equity / sim.peak - 1
        if sim.cooldown > 0:
            sim.cooldown -= 1
        if drawdown <= -0.10:
            sim.cooldown = max(sim.cooldown, 20)

        for code, row in rows.items():
            history = histories[code]
            if history:
                returns[code].append(row["close"] / history[-1] - 1)
            history.append(row["close"])

        if date_index % config.rebalance or date_index < config.slow:
            continue
        if sim.cooldown:
            sim.pending = {}
            continue
        candidates = []
        for code in TARGETS:
            history = histories[code]
            if len(history) <= config.slow or len(returns[code]) < 40:
                continue
            current = history[-1]
            fast_momentum = current / history[-config.fast - 1] - 1
            slow_momentum = current / history[-config.slow - 1] - 1
            volatility = statistics.stdev(returns[code]) * math.sqrt(252)
            average = statistics.mean(list(history)[-config.slow:])
            liquid = rows.get(code, {}).get("amount", 0) >= 20_000
            if not liquid or current <= average or slow_momentum <= 0 or volatility <= 0:
                continue
            ranking = (0.55 * fast_momentum + 0.45 * slow_momentum) / max(0.04, volatility) ** config.vol_penalty
            candidates.append((ranking, code, volatility))
        selected = sorted(candidates, reverse=True)[: config.max_assets]
        exposure = config.exposure * (0.5 if drawdown <= -0.05 else 1.0)
        if not selected:
            sim.pending = {}
            continue
        inverse_volatility = [1 / max(0.04, item[2]) for item in selected]
        divisor = sum(inverse_volatility)
        sim.pending = {item[1]: min(0.45, exposure * inverse_volatility[index] / divisor) for index, item in enumerate(selected)}
    return sim


def equal_weight_benchmark(dates, data):
    values = []
    active = {}
    for date in dates:
        rows = data[date]
        available = [code for code in TARGETS if code in rows]
        target = {code: 1 / len(available) for code in available} if available else {}
        factors = {code: rows[code]["close"] / rows[code]["preClose"] for code in available}
        active = {code: target.get(code, 0) for code in available}
        _, daily_return = drift(active, factors)
        values.append(daily_return)
    return values


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--archive", type=Path, required=True)
    parser.add_argument("--result", type=Path, required=True)
    parser.add_argument("--report", type=Path, required=True)
    parser.add_argument("--cost-bps", type=float, default=10.0)
    args = parser.parse_args()

    data = load_daily(args.archive)
    dates = sorted(date for date, rows in data.items() if rows)
    ranges = {
        "calibration": {"start": "2013-07-29", "end": "2017-12-31"},
        "validation": {"start": "2018-01-01", "end": "2021-12-31"},
        "finalHoldout": {"start": "2022-01-01", "end": dates[-1]},
        "fullResearch": {"start": dates[0], "end": dates[-1]},
    }
    configs = [
        Config(fast, slow, rebalance, max_assets, exposure, penalty)
        for fast in (10, 20, 40)
        for slow in (60, 120, 200)
        if fast < slow
        for rebalance in (5, 10, 20)
        for max_assets in (1, 2, 3)
        for exposure in (0.25, 0.50, 0.75)
        for penalty in (0.5, 1.0)
    ]
    simulations = [run_simulation(config, dates, data, args.cost_bps) for config in configs]
    evaluated = []
    for sim in simulations:
        item = {"id": sim.config.identifier, "config": sim.config.__dict__, "turnover": round(sum(sim.turnovers), 5)}
        for name, span in ranges.items():
            item[name] = metrics(sim.returns, dates, span["start"], span["end"])
        evaluated.append(item)
    calibration_top = sorted(evaluated, key=lambda item: score(item["calibration"]), reverse=True)[:30]
    winner = max(calibration_top, key=lambda item: score(item["validation"]))
    winner_sim = simulations[next(index for index, sim in enumerate(simulations) if sim.config.identifier == winner["id"])]
    stress_returns = [value - turnover * args.cost_bps / 10_000 for value, turnover in zip(winner_sim.returns, winner_sim.turnovers)]
    cost_stress = {name: metrics(stress_returns, dates, span["start"], span["end"]) for name, span in ranges.items()}
    benchmark_returns = equal_weight_benchmark(dates, data)
    benchmark = {name: metrics(benchmark_returns, dates, span["start"], span["end"]) for name, span in ranges.items()}

    result = {
        "version": "cn-etf-daily-rotation-v1",
        "generatedAt": dt.datetime.now(dt.timezone.utc).isoformat(),
        "mode": "DAILY_RESEARCH_ONLY_INTRADAY_NOT_VALIDATED_REAL_ORDERS_LOCKED",
        "dataset": {
            "archive": str(args.archive),
            "symbols": list(TARGETS),
            "symbolNames": NAMES,
            "sessions": len(dates),
            "start": dates[0],
            "end": dates[-1],
            "rows": sum(len(rows) for rows in data.values()),
        },
        "method": {
            "signalTiming": "T日收盘后形成信号；最早T+1开盘执行",
            "returnPath": "持仓先计隔夜收益，开盘换仓并扣费，再计开盘至收盘收益",
            "selection": "仅校准期选前30；只用验证期选冠军；2022年后最终留出不参与选参",
            "costBpsPerOneWayTurnover": args.cost_bps,
            "doubleCostStressBps": args.cost_bps * 2,
            "intradayClaim": False,
        },
        "ranges": ranges,
        "configsEvaluated": len(evaluated),
        "winner": winner,
        "winnerDoubleCostStress": cost_stress,
        "equalWeightBenchmark": benchmark,
        "topValidationCandidates": sorted(calibration_top, key=lambda item: score(item["validation"]), reverse=True)[:10],
        "gates": {
            "allowedForDailyShadow": True,
            "allowedForIntradayParameterPromotion": False,
            "allowedForRealCapital": False,
            "blockers": [
                "没有ETF历史分钟正文，日线结果不能验证60分钟EMA/VWAP日内策略",
                "没有历史IOPV与券商Level-2队列",
                "必须先经历新的前向影子样本与人工复核",
            ],
        },
    }
    args.result.parent.mkdir(parents=True, exist_ok=True)
    args.result.write_text(json.dumps(result, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")

    pct = lambda value: f"{100 * value:.2f}%"
    report = [
        "# A股场内 T+0 ETF 日线轮动研究 v1",
        "",
        f"- 数据：{len(dates):,}个含目标ETF的交易日、{result['dataset']['rows']:,}条目标ETF日线，{dates[0]} 至 {dates[-1]}。",
        f"- 评估：{len(evaluated)}组配置；2022年后最终留出未参与选参。",
        "- 边界：这是收盘信号、次日开盘执行的日线研究，不是ETF日内回测。",
        "",
        "## 冻结日线挑战者",
        "",
        f"- 配置：`{winner['id']}`。",
        f"- 验证期：年化{pct(winner['validation']['annualizedReturn'])}，Sharpe {winner['validation']['sharpe']:.2f}，最大回撤{pct(winner['validation']['maxDrawdown'])}。",
        f"- 最终留出：年化{pct(winner['finalHoldout']['annualizedReturn'])}，Sharpe {winner['finalHoldout']['sharpe']:.2f}，最大回撤{pct(winner['finalHoldout']['maxDrawdown'])}。",
        f"- 双倍成本留出：年化{pct(cost_stress['finalHoldout']['annualizedReturn'])}，最大回撤{pct(cost_stress['finalHoldout']['maxDrawdown'])}。",
        f"- 三只ETF等权留出：年化{pct(benchmark['finalHoldout']['annualizedReturn'])}，最大回撤{pct(benchmark['finalHoldout']['maxDrawdown'])}。",
        "",
        "## 实盘边界",
        "",
        "日线挑战者可以加入前向影子比较，但不会替换尚未验证的日内参数，也不会产生真实资金订单。需要ETF历史分钟、IOPV和券商Level-2后才能训练真实日内执行层。",
        "",
    ]
    args.report.parent.mkdir(parents=True, exist_ok=True)
    args.report.write_text("\n".join(report), encoding="utf-8")
    print(json.dumps({"dataset": result["dataset"], "configsEvaluated": len(evaluated), "winner": winner, "doubleCostHoldout": cost_stress["finalHoldout"], "benchmarkHoldout": benchmark["finalHoldout"]}, ensure_ascii=False, indent=2))


if __name__ == "__main__":
    main()
