#!/usr/bin/env python3
"""Point-in-time constrained intraday ETF challenger research.

This is a research/backtest program, not an order router.  It scans the current
official VERIFIED_T0 pool, applies each record's listing date, forms signals with
bars available by 10:00, delays entry, exits before the close and models costs,
capacity and partial fills.  Calibration and validation choose a challenger;
the 2022+ holdout is never used for parameter selection.

The current-universe history still has survivorship bias because a complete
historical official T+0 membership archive is not available.  The result is
therefore blocked from automatic promotion even when the numerical gates pass.
"""

from __future__ import annotations

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

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


DATE_RE = re.compile(r"(20\d{6})")


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


def ymd(value: str | None) -> str:
    return re.sub(r"\D", "", value or "")[:8]


def normalize_record(record: dict) -> tuple[str, str]:
    suffix = ".SH" if record.get("venue") == "SSE" else ".SZ"
    return f"{record['symbol']}{suffix}", ymd(record.get("listingDate"))


def load_universe(path: Path):
    payload = json.loads(path.read_text(encoding="utf-8"))
    records = [
        record
        for record in payload.get("records", [])
        if record.get("dayTradeEligibility") == "VERIFIED_T0"
        and record.get("listStatus") == "LISTED"
    ]
    mapping = dict(normalize_record(record) for record in records)
    return records, mapping


def aggregate(table, aggregates):
    if table.num_rows == 0:
        return {}
    grouped = table.group_by("code", use_threads=False).aggregate(aggregates)
    return {row["code"]: row for row in grouped.to_pylist()}


def daily_features(table, date: str, eligible_codes: list[str], listing_dates: dict[str, str]):
    columns = ["code", "trade_time", "open", "high", "low", "close", "vol", "amount"]
    table = table.select(columns)
    table = table.filter(pc.is_in(table["code"], value_set=pa.array(eligible_codes)))
    if table.num_rows == 0:
        return []
    valid = pc.and_(pc.greater(table["close"], 0), pc.greater_equal(table["amount"], 0))
    table = table.filter(valid)
    start = f"{date[:4]}-{date[4:6]}-{date[6:8]} "
    morning = table.filter(pc.less_equal(table["trade_time"], start + "10:00:00"))
    entry_window = table.filter(
        pc.and_(
            pc.greater(table["trade_time"], start + "10:00:00"),
            pc.less_equal(table["trade_time"], start + "10:05:00"),
        )
    )
    exit_window = table.filter(pc.less_equal(table["trade_time"], start + "14:50:00"))
    all_day = aggregate(
        table,
        [("open", "first"), ("close", "last"), ("amount", "sum"), ("high", "max"), ("low", "min"), ("vol", "sum")],
    )
    early = aggregate(
        morning,
        [("open", "first"), ("close", "last"), ("amount", "sum"), ("high", "max"), ("low", "min"), ("vol", "sum")],
    )
    entry_amount = aggregate(entry_window, [("amount", "sum")])
    exit_prices = aggregate(exit_window, [("close", "last")])
    entries = {}
    for delay in (1, 3, 5):
        threshold = start + f"10:{delay:02d}:00"
        delayed = table.filter(pc.greater_equal(table["trade_time"], threshold))
        entries[delay] = aggregate(delayed, [("open", "first")])

    rows = []
    for code, full in all_day.items():
        if listing_dates.get(code) and date < listing_dates[code]:
            continue
        am = early.get(code)
        ex = exit_prices.get(code)
        capacity = entry_amount.get(code)
        if not am or not ex or not capacity:
            continue
        volume = finite(am.get("vol_sum"))
        amount = finite(am.get("amount_sum"))
        open_price = finite(am.get("open_first"))
        morning_close = finite(am.get("close_last"))
        if min(volume, amount, open_price, morning_close, finite(ex.get("close_last"))) <= 0:
            continue
        row = {
            "date": date,
            "code": code,
            "morningOpen": open_price,
            "morningClose": morning_close,
            "morningHigh": finite(am.get("high_max")),
            "morningLow": finite(am.get("low_min")),
            "morningAmount": amount,
            "morningVwap": amount / volume,
            "entryWindowAmount": finite(capacity.get("amount_sum")),
            "exitPrice": finite(ex.get("close_last")),
            "fullClose": finite(full.get("close_last")),
            "fullAmount": finite(full.get("amount_sum")),
        }
        for delay in (1, 3, 5):
            row[f"entry{delay}"] = finite(entries[delay].get(code, {}).get("open_first"))
        if min(row["entry1"], row["entry3"], row["entry5"]) <= 0:
            continue
        rows.append(row)
    return rows


def member_date(name: str):
    match = DATE_RE.search(Path(name).name)
    return match.group(1) if match else None


def build_feature_cache(root: Path, cache: Path, listing_dates: dict[str, str]):
    minute_root = root / "ETF" / "分钟K线" / "1m"
    eligible = sorted(listing_dates)
    rows = []
    dates_seen = set()
    for archive_path in sorted(minute_root.glob("20??.zip")):
        with ZipFile(archive_path) as archive:
            members = sorted(name for name in archive.namelist() if name.endswith(".parquet"))
            for member in members:
                date = member_date(member)
                if not date or date in dates_seen:
                    continue
                with archive.open(member) as stream:
                    rows.extend(daily_features(pq.read_table(stream), date, eligible, listing_dates))
                dates_seen.add(date)
        print(json.dumps({"featureYear": archive_path.stem, "rows": len(rows)}, ensure_ascii=False), flush=True)
    for path in sorted(minute_root.glob("20??????.parquet")):
        date = member_date(path.name)
        if not date or date in dates_seen:
            continue
        rows.extend(daily_features(pq.read_table(path), date, eligible, listing_dates))
        dates_seen.add(date)
    cache.parent.mkdir(parents=True, exist_ok=True)
    table = pa.Table.from_pylist(rows)
    pq.write_table(table, cache, compression="zstd", compression_level=9)
    return rows


def load_features(root: Path, cache: Path, listing_dates: dict[str, str], rebuild: bool):
    if cache.exists() and not rebuild:
        return pq.read_table(cache).to_pylist()
    return build_feature_cache(root, cache, listing_dates)


def enrich(rows):
    by_date = collections.defaultdict(list)
    for row in rows:
        by_date[row["date"]].append(row)
    close_history = collections.defaultdict(lambda: collections.deque(maxlen=80))
    return_history = collections.defaultdict(lambda: collections.deque(maxlen=60))
    amount_history = collections.defaultdict(lambda: collections.deque(maxlen=30))
    market_returns = collections.deque(maxlen=80)
    enriched = {}
    for date in sorted(by_date):
        rolling_market = 1.0
        for value in list(market_returns)[-60:]:
            rolling_market *= 1 + value
        regime = "BULL" if rolling_market - 1 > 0.06 else "BEAR" if rolling_market - 1 < -0.06 else "SIDEWAYS"
        day = []
        cross_returns = []
        for row in by_date[date]:
            code = row["code"]
            closes = close_history[code]
            returns = return_history[code]
            amounts = amount_history[code]
            if closes:
                cross_returns.append(row["fullClose"] / closes[-1] - 1)
            row = dict(row)
            row["mom20"] = row["morningClose"] / closes[-20] - 1 if len(closes) >= 20 else None
            row["mom60"] = row["morningClose"] / closes[-60] - 1 if len(closes) >= 60 else None
            row["vol20"] = statistics.stdev(list(returns)[-20:]) if len(returns) >= 20 else None
            row["adv20"] = statistics.median(list(amounts)[-20:]) if len(amounts) >= 20 else None
            row["morningReturn"] = row["morningClose"] / row["morningOpen"] - 1
            row["morningRange"] = row["morningHigh"] / row["morningLow"] - 1 if row["morningLow"] > 0 else 0
            row["vwapDeviation"] = row["morningClose"] / row["morningVwap"] - 1 if row["morningVwap"] > 0 else 0
            row["regime"] = regime
            day.append(row)
        for row in by_date[date]:
            code = row["code"]
            closes = close_history[code]
            if closes:
                return_history[code].append(row["fullClose"] / closes[-1] - 1)
            closes.append(row["fullClose"])
            amount_history[code].append(row["fullAmount"])
        if cross_returns:
            market_returns.append(statistics.mean(cross_returns))
        enriched[date] = day
    return enriched


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

    @property
    def identifier(self):
        return f"ETF-I-{self.mode}-T{self.threshold:g}-M{self.momentumDays}-K{self.maxAssets}-E{self.exposure:g}-D{self.delayMinutes}"


def candidate_score(row, config: Config):
    momentum = row[f"mom{config.momentumDays}"]
    volatility = row["vol20"]
    adv = row["adv20"]
    if momentum is None or volatility is None or adv is None or adv < 10_000_000:
        return None
    if row["morningAmount"] < max(1_000_000, adv * 0.015):
        return None
    if row["morningRange"] > 0.08 or volatility > 0.08:
        return None
    if config.mode == "TREND":
        if row["morningReturn"] <= config.threshold or momentum <= 0 or row["vwapDeviation"] <= 0:
            return None
        return 0.55 * row["morningReturn"] + 0.35 * momentum - 0.10 * volatility
    if row["morningReturn"] >= -config.threshold or momentum <= 0 or row["vwapDeviation"] >= 0:
        return None
    return -0.55 * row["morningReturn"] + 0.25 * momentum - 0.20 * volatility


def simulate(enriched, config: Config, cost_multiplier=1.0, initial_equity=1_000_000.0):
    equity = peak = initial_equity
    returns = []
    dates = []
    regimes = []
    trade_results = []
    order_count = filled_order_count = partial_order_count = 0
    gross_turnover = 0.0
    for date, rows in sorted(enriched.items()):
        candidates = []
        for row in rows:
            score = candidate_score(row, config)
            if score is not None:
                candidates.append((score, row))
        selected = [item[1] for item in sorted(candidates, key=lambda item: item[0], reverse=True)[: config.maxAssets]]
        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
                if fill_ratio < 0.999:
                    partial_order_count += 1
                gross = exit_price / entry - 1
                participation = order_notional * fill_ratio / max(row["adv20"], 1.0)
                impact = 2 * 0.5 * row["vol20"] * math.sqrt(max(0.0, participation))
                explicit = 2 * config.costBps / 10_000
                cost = cost_multiplier * min(0.01, explicit + impact)
                realized = gross - cost
                contribution = target_weight * fill_ratio * realized
                day_return += contribution
                trade_results.append({"date": date, "code": row["code"], "gross": gross, "net": realized, "fill": fill_ratio, "regime": row["regime"]})
                gross_turnover += 2 * target_weight * fill_ratio
        day_return = max(-0.08, min(0.08, day_return))
        equity *= 1 + day_return
        peak = max(peak, equity)
        returns.append(day_return)
        dates.append(date)
        regimes.append(rows[0]["regime"] if rows else "UNKNOWN")
    return {
        "returns": returns,
        "dates": dates,
        "regimes": regimes,
        "trades": trade_results,
        "orderCount": order_count,
        "filledOrderCount": filled_order_count,
        "partialOrderCount": partial_order_count,
        "grossTurnover": gross_turnover,
    }


def metrics(sim, start: str, end: str):
    selected = [(date, value, regime) for date, value, regime in zip(sim["dates"], sim["returns"], sim["regimes"]) if start <= date <= end]
    trades = [trade for trade in sim["trades"] if start <= trade["date"] <= end]
    if not selected:
        return {"observations": 0}
    values = [item[1] for item in selected]
    equity = peak = 1.0
    max_drawdown = 0.0
    yearly = collections.defaultdict(lambda: 1.0)
    regime_values = collections.defaultdict(list)
    for date, value, regime in selected:
        equity *= 1 + value
        peak = max(peak, equity)
        max_drawdown = min(max_drawdown, equity / peak - 1)
        yearly[date[:4]] *= 1 + value
        regime_values[regime].append(value)
    mean = statistics.mean(values)
    std = statistics.stdev(values) if len(values) > 1 else 0
    downside = [min(0.0, value) for value in values]
    downside_std = statistics.stdev(downside) if len(downside) > 1 else 0
    years = len(values) / 252
    annualized = equity ** (1 / years) - 1 if equity > 0 and years > 0 else -1
    trade_days = [value for value in values if abs(value) > 1e-12]
    return {
        "observations": len(values),
        "totalReturn": round(equity - 1, 6),
        "annualizedReturn": round(annualized, 6),
        "annualizedVolatility": round(std * math.sqrt(252), 6),
        "sharpe": round(mean / std * math.sqrt(252), 6) if std else 0,
        "sortino": round(mean / downside_std * math.sqrt(252), 6) if downside_std else 0,
        "maxDrawdown": round(max_drawdown, 6),
        "tradeDays": len(trade_days),
        "trades": len(trades),
        "tradeWinRate": round(sum(trade["net"] > 0 for trade in trades) / len(trades), 6) if trades else None,
        "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())},
        "regimes": {
            regime: {
                "observations": len(items),
                "meanDailyReturn": round(statistics.mean(items), 8),
                "positiveDayRate": round(sum(item > 0 for item in items) / len(items), 6),
            }
            for regime, items in sorted(regime_values.items())
        },
    }


def selection_score(item):
    result = item
    if result.get("observations", 0) == 0:
        return -999
    return result["annualizedReturn"] + 0.12 * result["sharpe"] - 0.85 * abs(result["maxDrawdown"])


def main():
    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)
    parser.add_argument("--rebuild-cache", action="store_true")
    args = parser.parse_args()

    records, listing_dates = load_universe(args.universe)
    rows = load_features(args.root, args.cache, listing_dates, args.rebuild_cache)
    enriched = 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]},
        "fullResearch": {"start": dates[0], "end": dates[-1]},
    }
    configs = [
        Config(mode, threshold, momentum, maximum, exposure, delay)
        for mode in ("TREND", "REVERSAL")
        for threshold in (0.0015, 0.003)
        for momentum in (20, 60)
        for maximum in (3, 6)
        for exposure in (0.25, 0.5)
        for delay in (1, 3)
    ]
    evaluated = []
    simulations = {}
    for config in configs:
        sim = simulate(enriched, config)
        simulations[config.identifier] = sim
        evaluated.append(
            {
                "id": config.identifier,
                "config": asdict(config),
                "calibration": metrics(sim, **ranges["calibration"]),
                "validation": metrics(sim, **ranges["validation"]),
            }
        )
    finalists = sorted(evaluated, key=lambda item: selection_score(item["calibration"]), reverse=True)[:12]
    winner = max(finalists, key=lambda item: selection_score(item["validation"]))
    winner_config = Config(**winner["config"])
    winner_sim = simulations[winner["id"]]
    winner["finalHoldout"] = metrics(winner_sim, **ranges["finalHoldout"])
    winner["fullResearch"] = metrics(winner_sim, **ranges["fullResearch"])
    stress = {}
    for multiplier in (1, 2, 3):
        stressed = simulate(enriched, winner_config, cost_multiplier=multiplier)
        stress[f"{multiplier}x"] = metrics(stressed, **ranges["finalHoldout"])
    holdout = winner["finalHoldout"]
    gates = {
        "holdoutAnnualizedPositive": holdout.get("annualizedReturn", -1) > 0,
        "holdoutMaxDrawdownWithin15Pct": holdout.get("maxDrawdown", -1) >= -0.15,
        "holdoutTradeSampleAtLeast200": holdout.get("trades", 0) >= 200,
        "holdoutPositiveYearRateAtLeast60Pct": holdout.get("positiveYearRate", 0) >= 0.60,
        "doubleCostAnnualizedPositive": stress["2x"].get("annualizedReturn", -1) > 0,
        "completeHistoricalOfficialT0Membership": False,
        "brokerLevel2AndIopv": False,
    }
    result = {
        "version": "etf-full-universe-intraday-challenger-v1",
        "generatedAt": dt.datetime.now(dt.timezone.utc).isoformat(),
        "dataRoot": str(args.root),
        "featureCache": str(args.cache),
        "featureRows": len(rows),
        "dateRange": {"start": dates[0], "end": dates[-1], "tradingDays": len(dates)},
        "universe": {"currentOfficialVerifiedT0": len(records), "pointInTimeListingDateApplied": True, "historicalMembershipComplete": False},
        "selection": {"parameterConfigurations": len(configs), "calibrationTopKept": len(finalists), "finalHoldoutUsedForSelection": False},
        "winner": winner,
        "costStressFinalHoldout": stress,
        "gates": gates,
        "passedNumericalGates": all(value for key, value in gates.items() if key not in {"completeHistoricalOfficialT0Membership", "brokerLevel2AndIopv"}),
        "promotion": "BLOCKED",
        "promotionReasons": [
            "缺少完整历史逐日官方T+0成员快照，仍存在当前成分幸存者偏差。",
            "缺少券商Level-2真实队列、历史IOPV与真实委托回报。",
            "挑战者必须经过至少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")
    report = f"""# 全量日内基金分钟挑战者研究 v1

生成时间：{result['generatedAt']}  
数据：2010—2026 ETF 1分钟K线；{len(rows):,} 个当前官方T+0候选的标的日特征。  
真实订单：0（研究/影子专用）。

## 样本与偏差控制

- 当前官方核验T+0池：{len(records)}只；逐只应用上市日期。
- 参数组合：{len(configs)}；只用2013—2017校准、2018—2021验证。
- 2022—{dates[-1]}最终留出没有参与选参。
- 仍缺少完整历史官方T+0成员快照，因此存在幸存者偏差，禁止自动晋级。

## 胜出配置

`{winner['id']}`

| 区间 | 年化收益 | Sharpe | 最大回撤 | 交易笔数 | 胜率 | 正收益年份率 |
|---|---:|---:|---:|---:|---:|---:|
| 校准 | {winner['calibration'].get('annualizedReturn', 0):.2%} | {winner['calibration'].get('sharpe', 0):.2f} | {winner['calibration'].get('maxDrawdown', 0):.2%} | {winner['calibration'].get('trades', 0)} | {finite(winner['calibration'].get('tradeWinRate')):.2%} | {winner['calibration'].get('positiveYearRate', 0):.2%} |
| 验证 | {winner['validation'].get('annualizedReturn', 0):.2%} | {winner['validation'].get('sharpe', 0):.2f} | {winner['validation'].get('maxDrawdown', 0):.2%} | {winner['validation'].get('trades', 0)} | {finite(winner['validation'].get('tradeWinRate')):.2%} | {winner['validation'].get('positiveYearRate', 0):.2%} |
| 最终留出 | {holdout.get('annualizedReturn', 0):.2%} | {holdout.get('sharpe', 0):.2f} | {holdout.get('maxDrawdown', 0):.2%} | {holdout.get('trades', 0)} | {finite(holdout.get('tradeWinRate')):.2%} | {holdout.get('positiveYearRate', 0):.2%} |

## 成本压力（最终留出）

| 成本 | 年化收益 | 最大回撤 | Sharpe |
|---|---:|---:|---:|
| 1倍 | {stress['1x'].get('annualizedReturn', 0):.2%} | {stress['1x'].get('maxDrawdown', 0):.2%} | {stress['1x'].get('sharpe', 0):.2f} |
| 2倍 | {stress['2x'].get('annualizedReturn', 0):.2%} | {stress['2x'].get('maxDrawdown', 0):.2%} | {stress['2x'].get('sharpe', 0):.2f} |
| 3倍 | {stress['3x'].get('annualizedReturn', 0):.2%} | {stress['3x'].get('maxDrawdown', 0):.2%} | {stress['3x'].get('sharpe', 0):.2f} |

## 结论

状态：**BLOCKED**。数值结果只用于挑战者研究，不是收益承诺。补齐历史点时T+0成员、券商队列/IOPV，并通过60个新交易日前向影子验证和人工批准前，不覆盖当前策略、不连接真实资金。
"""
    args.report.parent.mkdir(parents=True, exist_ok=True)
    args.report.write_text(report, encoding="utf-8")
    print(json.dumps({"result": str(args.result), "report": str(args.report), "winner": winner["id"], "holdout": holdout, "gates": gates}, ensure_ascii=False))


if __name__ == "__main__":
    main()
