#!/usr/bin/env python3
"""Read-only coverage audit for the downloaded TuShare ETF/fund archives.

The script deliberately separates daily ETF coverage from historical intraday
coverage.  It never treats the downloaded stock/index minute archives as ETF
minute data.
"""

from __future__ import annotations

import argparse
import json
import re
import tarfile
import tempfile
from collections import Counter, defaultdict
from pathlib import Path
from zipfile import ZipFile

import pyarrow.parquet as parquet


TARGETS = ("510900.SH", "511010.SH", "518880.SH")


def read_latest_snapshot(tar_path: Path, prefix: str):
    with tarfile.open(tar_path, "r:gz") as archive:
        members = [
            member
            for member in archive.getmembers()
            if member.isfile()
            and member.name.endswith(".parquet")
            and prefix in Path(member.name).name
        ]
        if not members:
            return [], None
        member = sorted(members, key=lambda item: item.name)[-1]
        stream = archive.extractfile(member)
        if stream is None:
            raise RuntimeError(f"Unable to read {member.name}")
        return parquet.read_table(stream).to_pylist(), member.name


def audit_fund_daily(tar_path: Path):
    codes: set[str] = set()
    dates: set[str] = set()
    total_rows = 0
    first_by_code: dict[str, str] = {}
    last_by_code: dict[str, str] = {}
    rows_by_code: Counter[str] = Counter()
    target_rows: dict[str, list[dict]] = defaultdict(list)

    # Parquet readers seek to the footer. Reading thousands of members through
    # TarFile.extractfile would repeatedly seek in the gzip stream and is
    # prohibitively slow. Stage only this 219 MB compressed sub-archive, audit
    # it locally, and let TemporaryDirectory remove the files afterwards.
    with tempfile.TemporaryDirectory(prefix="quant-atlas-etf-daily-") as temporary:
        root = Path(temporary)
        with tarfile.open(tar_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")):
            match = re.search(r"trade_date=(\d{8})", path.name)
            member_date = match.group(1) if match else None
            table = parquet.read_table(
                path,
                columns=[
                    "ts_code",
                    "trade_date",
                    "open",
                    "high",
                    "low",
                    "close",
                    "vol",
                    "amount",
                ],
            )
            columns = table.to_pydict()
            total_rows += table.num_rows
            if member_date:
                dates.add(member_date)
            for index, code in enumerate(columns["ts_code"]):
                date = str(columns["trade_date"][index] or member_date or "")
                if not code or not date:
                    continue
                codes.add(code)
                dates.add(date)
                rows_by_code[code] += 1
                first_by_code[code] = min(first_by_code.get(code, date), date)
                last_by_code[code] = max(last_by_code.get(code, date), date)
                if code in TARGETS:
                    target_rows[code].append(
                        {name: columns[name][index] for name in columns}
                    )

    return {
        "archive": str(tar_path),
        "rows": total_rows,
        "unique_codes": len(codes),
        "unique_trade_dates": len(dates),
        "first_trade_date": min(dates) if dates else None,
        "last_trade_date": max(dates) if dates else None,
        "codes": codes,
        "first_by_code": first_by_code,
        "last_by_code": last_by_code,
        "rows_by_code": rows_by_code,
        "target_rows": target_rows,
    }


def audit_target_adjustments(tar_path: Path):
    result = {code: {"rows": 0, "first": None, "last": None} for code in TARGETS}
    with tarfile.open(tar_path, "r:gz") as archive:
        for member in archive:
            if not member.isfile() or not member.name.endswith(".parquet"):
                continue
            if not any(code in member.name for code in TARGETS):
                continue
            stream = archive.extractfile(member)
            if stream is None:
                continue
            rows = parquet.read_table(
                stream, columns=["ts_code", "trade_date", "adj_factor"]
            ).to_pylist()
            for row in rows:
                code = row.get("ts_code")
                if code not in result:
                    continue
                date = str(row.get("trade_date") or "")
                result[code]["rows"] += 1
                result[code]["first"] = min(result[code]["first"] or date, date)
                result[code]["last"] = max(result[code]["last"] or date, date)
    return result


def audit_minute_archives(paths: list[Path], codes: set[str]):
    normalized = {
        code: {
            code.split(".")[0],
            code.lower().replace(".", ""),
            code.lower().replace(".", "_"),
            code.lower().replace(".", "."),
        }
        for code in codes
    }
    matches: dict[str, list[str]] = {code: [] for code in codes}
    archive_count = 0
    member_count = 0
    for directory in paths:
        for zip_path in sorted(directory.glob("*.zip")):
            year_match = re.match(r"(\d{4})", zip_path.name)
            if year_match and int(year_match.group(1)) < 2010:
                continue
            archive_count += 1
            with ZipFile(zip_path) as archive:
                names = archive.namelist()
                member_count += len(names)
                for name in names:
                    lower = name.lower()
                    for code, variants in normalized.items():
                        if any(variant.lower() in lower for variant in variants):
                            matches[code].append(f"{zip_path.name}:{name}")
    return {
        "archives_scanned": archive_count,
        "members_scanned": member_count,
        "target_matches": matches,
        "historical_etf_minute_confirmed": any(matches.values()),
    }


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--root", type=Path, required=True)
    parser.add_argument("--stock-minute", type=Path, required=True)
    parser.add_argument("--index-minute", type=Path, required=True)
    parser.add_argument("--output", type=Path, required=True)
    args = parser.parse_args()

    daily = audit_fund_daily(args.root / "fund_daily" / "fund_daily.tar.gz")
    basics, basic_member = read_latest_snapshot(
        args.root / "etf_basic" / "etf_basic.tar.gz", "etf_basic__snapshot="
    )
    listed = {
        row["ts_code"]: row
        for row in basics
        if row.get("ts_code") and row.get("list_status") == "L"
    }
    adjustments = audit_target_adjustments(
        args.root / "fund_adj" / "fund_adj.tar.gz"
    )
    minute = audit_minute_archives(
        [args.stock_minute, args.index_minute], set(TARGETS)
    )

    target_coverage = {}
    for code in TARGETS:
        rows = sorted(
            daily["target_rows"].get(code, []),
            key=lambda row: str(row.get("trade_date") or ""),
        )
        target_coverage[code] = {
            "name": (listed.get(code) or {}).get("csname"),
            "etf_type": (listed.get(code) or {}).get("etf_type"),
            "daily_rows": len(rows),
            "daily_first": str(rows[0].get("trade_date")) if rows else None,
            "daily_last": str(rows[-1].get("trade_date")) if rows else None,
            "last_close": rows[-1].get("close") if rows else None,
            "last_amount": rows[-1].get("amount") if rows else None,
            "adjustment": adjustments[code],
            "minute_matches": minute["target_matches"][code],
        }

    result = {
        "audit_date": "2026-08-05",
        "source_root": str(args.root),
        "daily": {
            key: value
            for key, value in daily.items()
            if key
            not in {
                "codes",
                "first_by_code",
                "last_by_code",
                "rows_by_code",
                "target_rows",
            }
        },
        "etf_basic": {
            "snapshot_member": basic_member,
            "rows": len(basics),
            "listed_rows": len(listed),
            "listed_with_daily_history": len(set(listed) & daily["codes"]),
        },
        "target_coverage": target_coverage,
        "minute_archive_audit": minute,
        "capability": {
            "etf_daily_backtest": len(set(listed) & daily["codes"]) > 0,
            "etf_daily_adjusted_backtest_for_targets": all(
                item["adjustment"]["rows"] > 0 for item in target_coverage.values()
            ),
            "etf_intraday_backtest": minute["historical_etf_minute_confirmed"],
            "historical_iopv": False,
            "broker_level2_queue": False,
        },
        "interpretation": [
            "ETF/fund daily OHLCV is available through 2026-07-15.",
            "Target ETF adjustment factors and reference metadata are available.",
            "The downloaded stock/index minute archives do not prove ETF minute coverage.",
            "ETF intraday optimization remains blocked until ETF minute files are supplied.",
            "Historical IOPV and broker Level-2/order acknowledgements are not present.",
        ],
    }
    args.output.parent.mkdir(parents=True, exist_ok=True)
    args.output.write_text(
        json.dumps(result, ensure_ascii=False, indent=2), encoding="utf-8"
    )
    print(json.dumps(result, ensure_ascii=False, indent=2))


if __name__ == "__main__":
    main()
