#!/usr/bin/env python3
"""Second intraday ETF challenger with regime, breadth and cash controls.

The final holdout is evaluated only after calibration/validation selection.  If
the selected candidate has no positive validation edge, CASH_CONTROL wins and
the strategy does not trade.  This prevents a forced-trade optimizer from
turning the least-bad loss into a production recommendation.
"""

from __future__ import annotations

import argparse
import collections
import datetime as dt
import json
import math
import statistics
from dataclasses import asdict, dataclass
from pathlib import Path

import etf_full_universe_intraday_challenger as base


@dataclass(frozen=True)
class Config:
    mode: str
    threshold: float
    quantile: float
    maxAssets: int
    exposure: float
    delayMinutes: int
    costBps: float = 5.0
    maxParticipation: float = 0.0075

    @property
    def identifier(self) -> str:
        return f"ETF-R-{self.mode}-T{self.threshold:g}-Q{self.quantile:g}-K{self.maxAssets}-E{self.exposure:g}-D{self.delayMinutes}"


def tradable(row: dict) -> bool:
    return (
        row.get("mom60") is not None
        and row.get("vol20") is not None
        and row.get("adv20") is not None
        and row["adv20"] >= 20_000_000
        and row["morningAmount"] >= max(2_000_000, row["adv20"] * 0.02)
        and row["morningRange"] <= 0.06
        and row["vol20"] <= 0.06
    )


def daily_candidates(rows: list[dict], config: Config) -> list[tuple[float, dict]]:
    eligible = [row for row in rows if tradable(row)]
    if len(eligible) < 5:
        return []
    ordered = sorted(eligible, key=lambda row: row["morningReturn"])
    ranks = {row["code"]: index / max(1, len(ordered) - 1) for index, row in enumerate(ordered)}
    breadth = sum(row["morningReturn"] > 0 for row in eligible) / len(eligible)
    output = []
    for row in eligible:
        high = row["morningHigh"]
        low = row["morningLow"]
        location = (row["morningClose"] - low) / max(high - low, row["morningClose"] * 1e-5)
        rank = ranks[row["code"]]
        momentum = row["mom60"]
        volatility = row["vol20"]
        score = None
        if config.mode == "RELATIVE_BREAKOUT":
            if breadth >= 0.52 and rank >= config.quantile and row["morningReturn"] >= config.threshold and location >= 0.72 and momentum > 0:
                score = 0.45 * row["morningReturn"] + 0.35 * momentum + 0.20 * location - 0.20 * volatility
        elif config.mode == "SELECTIVE_REVERSAL":
            if breadth >= 0.35 and rank <= 1 - config.quantile and row["morningReturn"] <= -config.threshold and location <= 0.28 and momentum > 0:
                score = -0.55 * row["morningReturn"] + 0.20 * momentum - 0.25 * volatility
        elif config.mode == "LOW_VOL_STRENGTH":
            if breadth >= 0.5 and rank >= config.quantile and row["morningReturn"] >= config.threshold and momentum > 0 and volatility <= 0.025:
                score = 0.35 * row["morningReturn"] + 0.40 * momentum - 0.45 * volatility
        output.append((score, row)) if score is not None else None
    return sorted(output, key=lambda item: item[0], reverse=True)[:config.maxAssets]


def simulate(enriched: dict[str, list[dict]], config: Config, cost_multiplier: float = 1.0):
    equity = 1_000_000.0
    returns, dates, regimes, trades = [], [], [], []
    order_count = filled_order_count = partial_order_count = 0
    for date, rows in sorted(enriched.items()):
        selected = [row for _, row in daily_candidates(rows, config)]
        day_return = 0.0
        if selected:
            target_weight = config.exposure / len(selected)
            for row in selected:
                order_count += 1
                entry = row[f"entry{config.delayMinutes}"]
                exit_price = row["exitPrice"]
                if min(entry, exit_price) <= 0:
                    continue
                order_notional = equity * target_weight
                visible_capacity = row["entryWindowAmount"] * config.maxParticipation
                fill_ratio = min(1.0, visible_capacity / max(order_notional, 1.0))
                if fill_ratio <= 0.05:
                    continue
                filled_order_count += 1
                partial_order_count += fill_ratio < 0.999
                gross = exit_price / entry - 1
                participation = order_notional * fill_ratio / max(row["adv20"], 1.0)
                impact = row["vol20"] * math.sqrt(max(0.0, participation))
                explicit = 2 * config.costBps / 10_000
                cost = cost_multiplier * min(0.01, explicit + impact)
                net = gross - cost
                day_return += target_weight * fill_ratio * net
                trades.append({"date": date, "code": row["code"], "gross": gross, "net": net, "fill": fill_ratio, "regime": row["regime"]})
        day_return = max(-0.04, min(0.04, day_return))
        equity *= 1 + day_return
        dates.append(date)
        returns.append(day_return)
        regimes.append(rows[0]["regime"] if rows else "UNKNOWN")
    return {"returns": returns, "dates": dates, "regimes": regimes, "trades": trades, "orderCount": order_count, "filledOrderCount": filled_order_count, "partialOrderCount": partial_order_count}


def score(metrics: dict) -> float:
    if metrics.get("trades", 0) < 50:
        return -999
    return metrics["annualizedReturn"] + 0.10 * metrics["sharpe"] - 0.70 * abs(metrics["maxDrawdown"])


def cash_metrics(observations: int) -> dict:
    return {"observations": observations, "totalReturn": 0.0, "annualizedReturn": 0.0, "annualizedVolatility": 0.0, "sharpe": 0.0, "sortino": 0.0, "maxDrawdown": 0.0, "tradeDays": 0, "trades": 0, "tradeWinRate": None, "positiveYearRate": 0.0, "yearly": {}, "regimes": {}}


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--root", type=Path, required=True)
    parser.add_argument("--universe", type=Path, required=True)
    parser.add_argument("--cache", type=Path, required=True)
    parser.add_argument("--result", type=Path, required=True)
    parser.add_argument("--report", type=Path, required=True)
    args = parser.parse_args()
    records, listing_dates = base.load_universe(args.universe)
    rows = base.load_features(args.root, args.cache, listing_dates, False)
    enriched = base.enrich(rows)
    dates = sorted(enriched)
    ranges = {
        "calibration": {"start": "2013-01-01", "end": "2017-12-31"},
        "validation": {"start": "2018-01-01", "end": "2021-12-31"},
        "finalHoldout": {"start": "2022-01-01", "end": dates[-1]},
    }
    configs = [
        Config(mode, threshold, quantile, maximum, exposure, delay)
        for mode in ("RELATIVE_BREAKOUT", "SELECTIVE_REVERSAL", "LOW_VOL_STRENGTH")
        for threshold in (0.003, 0.006, 0.01)
        for quantile in (0.85, 0.93)
        for maximum in (1, 3)
        for exposure in (0.10, 0.20)
        for delay in (3, 5)
    ]
    evaluated = []
    simulations = {}
    for config in configs:
        sim = simulate(enriched, config)
        simulations[config.identifier] = sim
        evaluated.append({
            "id": config.identifier,
            "config": asdict(config),
            "calibration": base.metrics(sim, **ranges["calibration"]),
            "validation": base.metrics(sim, **ranges["validation"]),
        })
    finalists = sorted(evaluated, key=lambda item: score(item["calibration"]), reverse=True)[:18]
    candidate = max(finalists, key=lambda item: score(item["validation"]))
    candidate_config = Config(**candidate["config"])
    candidate["finalHoldout"] = base.metrics(simulations[candidate["id"]], **ranges["finalHoldout"])
    validation = candidate["validation"]
    candidate_passed = (
        validation.get("annualizedReturn", -1) > 0
        and validation.get("sharpe", -1) > 0
        and validation.get("maxDrawdown", -1) >= -0.10
        and validation.get("trades", 0) >= 50
    )
    selected = candidate["id"] if candidate_passed else "CASH_CONTROL"
    selected_holdout = candidate["finalHoldout"] if candidate_passed else cash_metrics(candidate["finalHoldout"].get("observations", 0))
    stress = {
        f"{multiplier}x": base.metrics(simulate(enriched, candidate_config, multiplier), **ranges["finalHoldout"])
        for multiplier in (1, 2, 3)
    }
    result = {
        "version": "etf-intraday-regime-challenger-v2",
        "generatedAt": dt.datetime.now(dt.timezone.utc).isoformat(),
        "featureRows": len(rows),
        "dateRange": {"start": dates[0], "end": dates[-1], "tradingDays": len(dates)},
        "universe": {"currentOfficialVerifiedT0": len(records), "historicalMembershipComplete": False},
        "selection": {"configurations": len(configs), "finalHoldoutUsedForSelection": False, "cashControlIncluded": True},
        "bestRiskyCandidate": candidate,
        "candidateCostStressFinalHoldout": stress,
        "selectedPolicy": selected,
        "selectedPolicyFinalHoldout": selected_holdout,
        "promotion": "BLOCKED",
        "reasons": [
            "若校准与验证没有正净优势，现金控制优先，禁止强迫交易。",
            "历史逐日官方T+0成员仍不完整；当前成分历史存在幸存者偏差。",
            "缺少历史IOPV、券商Level-2真实队列和券商委托回报。",
            "任何风险策略必须再经过60个新交易日前向影子验证和人工批准。",
        ],
        "realOrders": 0,
    }
    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")
    holdout = candidate["finalHoldout"]
    report = f"""# 全量日内ETF状态与相对强弱挑战者 v2

生成时间：{result['generatedAt']}  
最终留出未参与选参；真实订单 0。

## 研究变化

- 从单标的上午追涨/反转，改为全截面相对强弱、突破位置、反转位置、市场宽度及波动率过滤。
- 延迟 3/5 分钟，考虑参与率、部分成交、冲击与双边成本。
- 新增 `CASH_CONTROL`：验证期没有正净优势时必须空仓。

## 最佳风险候选（不等于获准策略）

`{candidate['id']}`：验证年化 {validation.get('annualizedReturn', 0):.2%}，Sharpe {validation.get('sharpe', 0):.2f}，最大回撤 {validation.get('maxDrawdown', 0):.2%}，交易 {validation.get('trades', 0)} 笔。

最终留出：年化 {holdout.get('annualizedReturn', 0):.2%}，最大回撤 {holdout.get('maxDrawdown', 0):.2%}，交易 {holdout.get('trades', 0)} 笔。该结果只作审计，没有用于选参。

## 决策

当前选择：**{selected}**。即便新的研究假设比第一轮更复杂，也不能用“最不亏”的策略替代现金。继续保留全量227只研究池，但不扩大真实或镜像激活仓位；待补齐点时成员、IOPV/Level-2并获得新前向样本后再挑战。
"""
    args.report.parent.mkdir(parents=True, exist_ok=True)
    args.report.write_text(report, encoding="utf-8")
    print(json.dumps({"selectedPolicy": selected, "candidate": candidate["id"], "validation": validation, "holdout": holdout}, ensure_ascii=False))


if __name__ == "__main__":
    main()
