#!/usr/bin/env python3
"""Full, read-only coverage and integrity audit for the ETF one-minute archive.

The audit streams every parquet member instead of extracting the 4+ GiB archive
tree.  It deliberately distinguishes three different sets:

* every code present in the supplied historical ETF package;
* the current official VERIFIED_T0 universe exported by Quant Atlas;
* the small mirror account, which is not promoted by this audit.

No strategy parameter is selected here.  This file only establishes whether the
data are suitable inputs for a later point-in-time, walk-forward study.
"""

from __future__ import annotations

import argparse
import datetime as dt
import json
import math
import re
from collections import Counter
from pathlib import Path
from zipfile import ZipFile

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


DATE_PATTERN = re.compile(r"(20\d{6})")
YEAR_PATTERN = re.compile(r"(20\d{2})")


def normalize_code(value: object) -> str:
    text = str(value or "").strip().upper()
    match = re.search(r"(\d{6})", text)
    return match.group(1) if match else text


def date_from_name(name: str) -> str | None:
    match = DATE_PATTERN.search(Path(name).name)
    return match.group(1) if match else None


def finite_float(value: object) -> float | None:
    try:
        number = float(value)
        return number if math.isfinite(number) else None
    except (TypeError, ValueError):
        return None


def load_verified_t0(path: Path) -> tuple[set[str], list[dict]]:
    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"
    ]
    return {normalize_code(record.get("symbol")) for record in records}, records


def inspect_table(table, expected_date: str | None) -> dict:
    names = set(table.column_names)
    required = {"code", "trade_time", "open", "high", "low", "close", "vol", "amount"}
    missing = sorted(required - names)
    if missing:
        return {
            "rows": table.num_rows,
            "codes": set(),
            "missingColumns": missing,
            "missingColumnFile": 1,
            "nonPositivePriceRows": None,
            "negativeVolumeRows": None,
            "negativeAmountRows": None,
            "timestampBoundaryMismatchFile": None,
            "barsPerCode": Counter(),
        }

    code_counts = pc.value_counts(table["code"])
    code_values = code_counts.field("values").to_pylist()
    count_values = code_counts.field("counts").to_pylist()
    codes = {normalize_code(value) for value in code_values}
    bars_per_code = Counter(
        {
            normalize_code(code): int(count)
            for code, count in zip(code_values, count_values)
        }
    )
    close = table["close"]
    volume = table["vol"]
    amount = table["amount"]
    non_positive = int(pc.sum(pc.cast(pc.less_equal(close, 0), "int64")).as_py() or 0)
    negative_volume = int(pc.sum(pc.cast(pc.less(volume, 0), "int64")).as_py() or 0)
    negative_amount = int(pc.sum(pc.cast(pc.less(amount, 0), "int64")).as_py() or 0)

    mismatch = 0
    if expected_date:
        boundaries = [pc.min(table["trade_time"]).as_py(), pc.max(table["trade_time"]).as_py()]
        for value in boundaries:
            digits = re.sub(r"\D", "", str(value or ""))
            if len(digits) >= 8 and digits[:8] != expected_date:
                mismatch = 1
                break

    return {
        "rows": table.num_rows,
        "codes": codes,
        "missingColumns": [],
        "missingColumnFile": 0,
        "nonPositivePriceRows": non_positive,
        "negativeVolumeRows": negative_volume,
        "negativeAmountRows": negative_amount,
        "timestampBoundaryMismatchFile": mismatch,
        "barsPerCode": bars_per_code,
    }


def read_member(archive: ZipFile, member: str):
    with archive.open(member) as stream:
        return pq.read_table(
            stream,
            columns=["code", "trade_time", "open", "high", "low", "close", "vol", "amount"],
        )


def audit_archives(root: Path, verified_t0: set[str]) -> dict:
    years: list[dict] = []
    all_codes: set[str] = set()
    all_verified_seen: set[str] = set()
    all_dates: set[str] = set()
    latest_date_by_code: dict[str, str] = {}
    seen_date_sources: dict[str, str] = {}
    totals = Counter()
    bar_histogram: Counter[int] = Counter()
    duplicate_daily_files: list[dict] = []

    zip_paths = sorted(
        (path for path in root.glob("*.zip") if YEAR_PATTERN.fullmatch(path.stem)),
        key=lambda path: int(path.stem),
    )
    for zip_path in zip_paths:
        year_codes: set[str] = set()
        year_verified: set[str] = set()
        year_dates: set[str] = set()
        year_totals = Counter()
        with ZipFile(zip_path) as archive:
            members = sorted(
                info.filename
                for info in archive.infolist()
                if not info.is_dir() and info.filename.endswith(".parquet")
            )
            for member in members:
                date = date_from_name(member)
                if date:
                    if date in seen_date_sources:
                        duplicate_daily_files.append({"date": date, "first": seen_date_sources[date], "duplicate": f"{zip_path.name}:{member}"})
                    else:
                        seen_date_sources[date] = f"{zip_path.name}:{member}"
                    year_dates.add(date)
                    all_dates.add(date)
                inspected = inspect_table(read_member(archive, member), date)
                codes = inspected.pop("codes")
                bars = inspected.pop("barsPerCode")
                year_codes.update(codes)
                current_verified = codes & verified_t0
                year_verified.update(current_verified)
                all_codes.update(codes)
                all_verified_seen.update(current_verified)
                for code in codes:
                    if date:
                        latest_date_by_code[code] = max(latest_date_by_code.get(code, date), date)
                for count, frequency in Counter(bars.values()).items():
                    bar_histogram[count] += frequency
                for key, value in inspected.items():
                    if isinstance(value, int):
                        year_totals[key] += value
                        totals[key] += value

        years.append(
            {
                "year": int(zip_path.stem),
                "archive": str(zip_path),
                "dailyFiles": len(members),
                "firstDate": min(year_dates) if year_dates else None,
                "lastDate": max(year_dates) if year_dates else None,
                "rows": year_totals["rows"],
                "uniqueCodes": len(year_codes),
                "currentVerifiedT0Seen": len(year_verified),
                "missingColumnFiles": year_totals["missingColumnFile"],
                "nonPositivePriceRows": year_totals["nonPositivePriceRows"],
                "negativeVolumeRows": year_totals["negativeVolumeRows"],
                "negativeAmountRows": year_totals["negativeAmountRows"],
                "timestampBoundaryMismatchFiles": year_totals["timestampBoundaryMismatchFile"],
            }
        )
        print(
            json.dumps(
                {
                    "year": int(zip_path.stem),
                    "dailyFiles": len(members),
                    "rows": year_totals["rows"],
                    "uniqueCodes": len(year_codes),
                    "currentVerifiedT0Seen": len(year_verified),
                },
                ensure_ascii=False,
            ),
            flush=True,
        )

    # Loose daily files are authoritative additions after the latest yearly zip.
    loose_files = sorted(path for path in root.glob("*.parquet") if date_from_name(path.name))
    loose = []
    for path in loose_files:
        date = date_from_name(path.name)
        inspected = inspect_table(pq.read_table(path), date)
        codes = inspected.pop("codes")
        bars = inspected.pop("barsPerCode")
        duplicate_of = seen_date_sources.get(date or "")
        if date and not duplicate_of:
            seen_date_sources[date] = str(path)
            all_dates.add(date)
        elif date:
            duplicate_daily_files.append({"date": date, "first": duplicate_of, "duplicate": str(path)})
        all_codes.update(codes)
        all_verified_seen.update(codes & verified_t0)
        for code in codes:
            if date:
                latest_date_by_code[code] = max(latest_date_by_code.get(code, date), date)
        for count, frequency in Counter(bars.values()).items():
            bar_histogram[count] += frequency
        loose.append(
            {
                "file": str(path),
                "date": date,
                "rows": inspected["rows"],
                "uniqueCodes": len(codes),
                "currentVerifiedT0Seen": len(codes & verified_t0),
                "duplicateOfYearZip": duplicate_of,
            }
        )

    missing_verified = sorted(verified_t0 - all_verified_seen)
    stale_verified = sorted(
        code for code in all_verified_seen if latest_date_by_code.get(code, "") < (max(all_dates) if all_dates else "")
    )
    common_bars = [
        {"bars": bars, "codeDays": count}
        for bars, count in bar_histogram.most_common(20)
    ]
    return {
        "years": years,
        "looseDailyFiles": loose,
        "totals": {
            "yearArchives": len(zip_paths),
            "uniqueDates": len(all_dates),
            "firstDate": min(all_dates) if all_dates else None,
            "lastDate": max(all_dates) if all_dates else None,
            "rows": totals["rows"] + sum(item["rows"] for item in loose if not item["duplicateOfYearZip"]),
            "historicalCodes": len(all_codes),
            "currentVerifiedT0Requested": len(verified_t0),
            "currentVerifiedT0Seen": len(all_verified_seen),
            "currentVerifiedT0Missing": len(missing_verified),
            "currentVerifiedT0StaleOnLatestDate": len(stale_verified),
            "nonPositivePriceRows": totals["nonPositivePriceRows"],
            "negativeVolumeRows": totals["negativeVolumeRows"],
            "negativeAmountRows": totals["negativeAmountRows"],
            "timestampBoundaryMismatchFiles": totals["timestampBoundaryMismatchFile"],
            "duplicateDailyFiles": len(duplicate_daily_files),
        },
        "currentVerifiedT0MissingSymbols": missing_verified,
        "currentVerifiedT0NotPresentOnLatestDate": stale_verified,
        "duplicateDailyFiles": duplicate_daily_files,
        "barsPerCodeDayTop": common_bars,
    }


def audit_basic(path: Path) -> dict:
    table = pq.read_table(path)
    rows = table.to_pylist()
    status = Counter(str(row.get("list_status") or row.get("status") or "UNKNOWN") for row in rows)
    categories = Counter(str(row.get("etf_type") or row.get("fund_type") or "UNKNOWN") for row in rows)
    return {
        "file": str(path),
        "rows": len(rows),
        "columns": table.column_names,
        "statusCounts": dict(sorted(status.items())),
        "categoryCounts": dict(categories.most_common()),
    }


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("--output", type=Path, required=True)
    args = parser.parse_args()

    minute_root = args.root / "ETF" / "分钟K线" / "1m"
    basic_path = args.root / "ETF" / "etf_basic.parquet"
    if not minute_root.is_dir() or not basic_path.is_file():
        raise SystemExit("ETF minute/basic data paths are incomplete")

    verified, records = load_verified_t0(args.universe)
    minute = audit_archives(minute_root, verified)
    totals = minute["totals"]
    result = {
        "version": "etf-full-minute-coverage-audit-v1",
        "generatedAt": dt.datetime.now(dt.timezone.utc).isoformat(),
        "sourceRoot": str(args.root),
        "minuteRoot": str(minute_root),
        "universeSource": str(args.universe),
        "officialVerifiedT0": {
            "records": len(records),
            "symbols": sorted(verified),
        },
        "basic": audit_basic(basic_path),
        "minute": minute,
        "capability": {
            "etfHistoricalMinuteAcquired": totals["yearArchives"] >= 17 and totals["lastDate"] >= "20260806",
            "allCurrentVerifiedT0HaveSomeHistory": totals["currentVerifiedT0Missing"] == 0,
            "fullCurrentVerifiedT0LatestDateCoverage": totals["currentVerifiedT0StaleOnLatestDate"] == 0,
            "readyForPointInTimeWalkForwardResearch": totals["yearArchives"] >= 17 and totals["currentVerifiedT0Seen"] > 0,
            "readyForUnfilteredLivePromotion": False,
        },
        "limitations": [
            "Current VERIFIED_T0 membership is not a substitute for historical point-in-time eligibility.",
            "Historical IOPV, broker Level-2 queue position and broker acknowledgements are not supplied by these parquet bars.",
            "Coverage audit does not select parameters or promote a challenger into the mirror account.",
            "Future loose files must be uploaded and re-audited after each vendor update.",
        ],
    }
    args.output.parent.mkdir(parents=True, exist_ok=True)
    args.output.write_text(json.dumps(result, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
    print(json.dumps({"output": str(args.output), "totals": totals, "capability": result["capability"]}, ensure_ascii=False))


if __name__ == "__main__":
    main()
