#!/usr/bin/env python3
"""Compare production and independent-reference Freqtrade Feather data.

The comparison is read-only.  It freezes a half-open UTC window, applies
authoritative listing timestamps, and compares timestamps and values for each
pair/dataset.  For daily candles, a listing-day partial candle is excluded:
comparison starts at the first complete UTC day after listing.

Example (run inside the Freqtrade container when the host lacks pyarrow):

    python3 /tmp/compare_research_data.py \
      --production-data-dir /freqtrade/user_data/data/binance/futures \
      --reference-data-dir /freqtrade/user_data/data/binance_reference/futures \
      --listing-dates /tmp/binance-futures-listing-dates.json \
      --output-dir /tmp/research-data-comparison
"""

from __future__ import annotations

import argparse
import csv
import hashlib
import json
import platform
import socket
import sys
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Iterable

import numpy as np
import pandas as pd

SCRIPT_VERSION = "1.0"
DEFAULT_START = "2025-12-01T00:00:00Z"
DEFAULT_END = "2026-07-10T00:00:00Z"
DEFAULT_PAIRS = (
    "BTC/USDT:USDT",
    "ETH/USDT:USDT",
    "SOL/USDT:USDT",
    "BNB/USDT:USDT",
    "XRP/USDT:USDT",
    "DOGE/USDT:USDT",
    "ADA/USDT:USDT",
    "ZEC/USDT:USDT",
    "EDGE/USDT:USDT",
    "POWER/USDT:USDT",
    "VANRY/USDT:USDT",
    "UNI/USDT:USDT",
)


@dataclass(frozen=True)
class DatasetSpec:
    key: str
    suffix: str
    value_columns: tuple[str, ...]
    daily: bool = False


OHLCV = ("open", "high", "low", "close", "volume")
DATASETS = (
    DatasetSpec("futures_5m", "5m-futures", OHLCV),
    DatasetSpec("futures_1d", "1d-futures", OHLCV, daily=True),
    DatasetSpec("funding_1h", "1h-funding_rate", ("close",)),
    DatasetSpec("mark_1h", "1h-mark", OHLCV),
)

CSV_FIELDS = (
    "pair",
    "dataset",
    "status",
    "effective_start",
    "window_end",
    "listing_time",
    "production_path",
    "reference_path",
    "production_sha256",
    "reference_sha256",
    "production_rows_in_scope",
    "reference_rows_in_scope",
    "production_unique_timestamps",
    "reference_unique_timestamps",
    "overlap_timestamps",
    "production_only_timestamps",
    "reference_only_timestamps",
    "production_duplicate_timestamps",
    "reference_duplicate_timestamps",
    "production_duplicate_extra_rows",
    "reference_duplicate_extra_rows",
    "production_invalid_timestamps",
    "reference_invalid_timestamps",
    "ambiguous_overlap_timestamps",
    "compared_timestamps",
    "compared_cells",
    "mismatched_cells",
    "max_absolute_difference",
    "max_relative_difference",
    "missing_production_columns",
    "missing_reference_columns",
    "error",
)


def utc_now() -> str:
    return datetime.now(timezone.utc).isoformat()


def sha256_file(path: Path, chunk_size: int = 1024 * 1024) -> str:
    digest = hashlib.sha256()
    with path.open("rb") as handle:
        for chunk in iter(lambda: handle.read(chunk_size), b""):
            digest.update(chunk)
    return digest.hexdigest()


def pair_file(data_dir: Path, pair: str, suffix: str) -> Path:
    contract, *settle_parts = pair.split(":", 1)
    base, quote = contract.split("/", 1)
    settle = settle_parts[0] if settle_parts else quote
    return data_dir / f"{base}_{quote}_{settle}-{suffix}.feather"


def parse_pairs(raw_values: Iterable[str]) -> list[str]:
    result: list[str] = []
    for raw in raw_values:
        for pair in raw.split(","):
            pair = pair.strip()
            if pair and pair not in result:
                result.append(pair)
    return result


def load_pairs(config: Path | None, raw_pairs: list[str]) -> list[str]:
    if raw_pairs:
        return parse_pairs(raw_pairs)
    if config is not None:
        payload = json.loads(config.read_text(encoding="utf-8"))
        pairs = payload.get("exchange", {}).get("pair_whitelist")
        if not isinstance(pairs, list) or not all(isinstance(item, str) for item in pairs):
            raise ValueError(f"{config}: exchange.pair_whitelist must be a string list")
        return parse_pairs(pairs)
    return list(DEFAULT_PAIRS)


def as_utc(value: Any, label: str) -> pd.Timestamp:
    try:
        stamp = pd.Timestamp(value)
    except Exception as exc:
        raise ValueError(f"{label} is not a valid timestamp: {value!r}") from exc
    if pd.isna(stamp):
        raise ValueError(f"{label} is not a valid timestamp: {value!r}")
    return stamp.tz_localize("UTC") if stamp.tz is None else stamp.tz_convert("UTC")


def load_listing_dates(path: Path) -> dict[str, pd.Timestamp]:
    payload = json.loads(path.read_text(encoding="utf-8"))
    if not isinstance(payload, dict):
        raise ValueError("--listing-dates must contain a JSON object")
    return {str(pair): as_utc(value, f"listing date for {pair}") for pair, value in payload.items()}


def effective_start(
    window_start: pd.Timestamp,
    listing_time: pd.Timestamp | None,
    spec: DatasetSpec,
) -> pd.Timestamp:
    if listing_time is None:
        return window_start
    if spec.daily:
        # ceil(D) preserves an exact 00:00 listing, otherwise advances to the
        # first complete UTC day.  This excludes listing-day partial candles.
        first_full_day = listing_time.ceil("D")
        return max(window_start, first_full_day)
    return max(window_start, listing_time)


def normalize_frame(
    frame: pd.DataFrame,
    start: pd.Timestamp,
    end: pd.Timestamp,
) -> tuple[pd.DataFrame, int]:
    if "date" not in frame.columns:
        raise ValueError("required date column is missing")
    result = frame.copy()
    dates = pd.to_datetime(result["date"], utc=True, errors="coerce")
    invalid = int(dates.isna().sum())
    result = result.assign(date=dates)
    result = result.loc[dates.notna() & dates.ge(start) & dates.lt(end)].copy()
    return result, invalid


def duplicate_metrics(frame: pd.DataFrame) -> tuple[int, int]:
    counts = frame["date"].value_counts()
    duplicate_counts = counts[counts > 1]
    return int(len(duplicate_counts)), int((duplicate_counts - 1).sum())


def timestamp_samples(index: pd.Index, limit: int) -> list[str]:
    return [pd.Timestamp(value).isoformat() for value in index[:limit]]


def finite_max(values: np.ndarray) -> float | None:
    finite = values[np.isfinite(values)]
    return float(finite.max()) if finite.size else None


def compare_frames(
    production: pd.DataFrame,
    reference: pd.DataFrame,
    *,
    pair: str,
    spec: DatasetSpec,
    window_start: pd.Timestamp,
    window_end: pd.Timestamp,
    listing_time: pd.Timestamp | None,
    absolute_tolerance: float = 0.0,
    relative_tolerance: float = 0.0,
    sample_limit: int = 20,
) -> dict[str, Any]:
    start = effective_start(window_start, listing_time, spec)
    production, production_invalid = normalize_frame(production, start, window_end)
    reference, reference_invalid = normalize_frame(reference, start, window_end)

    production_dup_ts, production_dup_rows = duplicate_metrics(production)
    reference_dup_ts, reference_dup_rows = duplicate_metrics(reference)
    production_counts = production["date"].value_counts()
    reference_counts = reference["date"].value_counts()
    production_times = pd.DatetimeIndex(production_counts.index).sort_values()
    reference_times = pd.DatetimeIndex(reference_counts.index).sort_values()
    production_only = production_times.difference(reference_times)
    reference_only = reference_times.difference(production_times)
    overlap = production_times.intersection(reference_times)
    ambiguous = overlap[
        [
            production_counts.at[stamp] != 1 or reference_counts.at[stamp] != 1
            for stamp in overlap
        ]
    ]
    comparable = overlap.difference(ambiguous)

    missing_production_columns = [
        column for column in spec.value_columns if column not in production.columns
    ]
    missing_reference_columns = [
        column for column in spec.value_columns if column not in reference.columns
    ]
    available_columns = [
        column
        for column in spec.value_columns
        if column in production.columns and column in reference.columns
    ]

    column_metrics: dict[str, dict[str, Any]] = {}
    compared_cells = 0
    mismatched_cells = 0
    all_absolute: list[np.ndarray] = []
    all_relative: list[np.ndarray] = []
    if len(comparable) and available_columns:
        production_unique = production.set_index("date").loc[comparable]
        reference_unique = reference.set_index("date").loc[comparable]
        for column in available_columns:
            prod_values = pd.to_numeric(production_unique[column], errors="coerce").to_numpy(float)
            ref_values = pd.to_numeric(reference_unique[column], errors="coerce").to_numpy(float)
            both_nan = np.isnan(prod_values) & np.isnan(ref_values)
            comparable_values = np.isfinite(prod_values) & np.isfinite(ref_values)
            invalid_mismatch = ~(both_nan | comparable_values)
            absolute = np.abs(prod_values[comparable_values] - ref_values[comparable_values])
            denominator = np.maximum(
                np.maximum(np.abs(prod_values[comparable_values]), np.abs(ref_values[comparable_values])),
                np.finfo(float).tiny,
            )
            relative = absolute / denominator
            tolerance = absolute_tolerance + relative_tolerance * np.abs(
                ref_values[comparable_values]
            )
            numeric_mismatch = absolute > tolerance
            mismatch_count = int(numeric_mismatch.sum() + invalid_mismatch.sum())
            cell_count = int(len(prod_values))
            compared_cells += cell_count
            mismatched_cells += mismatch_count
            all_absolute.append(absolute)
            all_relative.append(relative)
            column_metrics[column] = {
                "compared_cells": cell_count,
                "mismatched_cells": mismatch_count,
                "nonfinite_mismatches": int(invalid_mismatch.sum()),
                "max_absolute_difference": finite_max(absolute),
                "max_relative_difference": finite_max(relative),
            }

    max_absolute = finite_max(np.concatenate(all_absolute)) if all_absolute else None
    max_relative = finite_max(np.concatenate(all_relative)) if all_relative else None
    structural_problem = bool(
        production_invalid
        or reference_invalid
        or production_dup_rows
        or reference_dup_rows
        or len(production_only)
        or len(reference_only)
        or missing_production_columns
        or missing_reference_columns
        or not len(overlap)
    )
    if structural_problem:
        status = "structural_mismatch"
    elif mismatched_cells:
        status = "value_mismatch"
    else:
        status = "match"

    return {
        "pair": pair,
        "dataset": spec.key,
        "status": status,
        "effective_start": start.isoformat(),
        "window_end": window_end.isoformat(),
        "listing_time": listing_time.isoformat() if listing_time is not None else None,
        "listing_filter": (
            "first_complete_utc_day" if spec.daily else "exact_listing_timestamp"
        ),
        "production_rows_in_scope": int(len(production)),
        "reference_rows_in_scope": int(len(reference)),
        "production_unique_timestamps": int(len(production_times)),
        "reference_unique_timestamps": int(len(reference_times)),
        "overlap_timestamps": int(len(overlap)),
        "production_only_timestamps": int(len(production_only)),
        "reference_only_timestamps": int(len(reference_only)),
        "production_only_timestamp_sample": timestamp_samples(production_only, sample_limit),
        "reference_only_timestamp_sample": timestamp_samples(reference_only, sample_limit),
        "production_duplicate_timestamps": production_dup_ts,
        "reference_duplicate_timestamps": reference_dup_ts,
        "production_duplicate_extra_rows": production_dup_rows,
        "reference_duplicate_extra_rows": reference_dup_rows,
        "production_invalid_timestamps": production_invalid,
        "reference_invalid_timestamps": reference_invalid,
        "ambiguous_overlap_timestamps": int(len(ambiguous)),
        "ambiguous_overlap_timestamp_sample": timestamp_samples(ambiguous, sample_limit),
        "compared_timestamps": int(len(comparable)),
        "compared_cells": compared_cells,
        "mismatched_cells": mismatched_cells,
        "max_absolute_difference": max_absolute,
        "max_relative_difference": max_relative,
        "relative_difference_definition": "abs(prod-ref)/max(abs(prod),abs(ref),float_tiny)",
        "absolute_tolerance": absolute_tolerance,
        "relative_tolerance": relative_tolerance,
        "expected_value_columns": list(spec.value_columns),
        "missing_production_columns": missing_production_columns,
        "missing_reference_columns": missing_reference_columns,
        "columns": column_metrics,
    }


def compare_file(
    production_dir: Path,
    reference_dir: Path,
    *,
    pair: str,
    spec: DatasetSpec,
    window_start: pd.Timestamp,
    window_end: pd.Timestamp,
    listing_time: pd.Timestamp | None,
    absolute_tolerance: float,
    relative_tolerance: float,
    sample_limit: int,
) -> dict[str, Any]:
    production_path = pair_file(production_dir, pair, spec.suffix)
    reference_path = pair_file(reference_dir, pair, spec.suffix)
    base = {
        "pair": pair,
        "dataset": spec.key,
        "production_path": str(production_path),
        "reference_path": str(reference_path),
        "production_sha256": sha256_file(production_path) if production_path.is_file() else None,
        "reference_sha256": sha256_file(reference_path) if reference_path.is_file() else None,
    }
    missing = [
        role
        for role, path in (("production", production_path), ("reference", reference_path))
        if not path.is_file()
    ]
    if missing:
        return {
            **base,
            "status": "missing_file",
            "effective_start": effective_start(window_start, listing_time, spec).isoformat(),
            "window_end": window_end.isoformat(),
            "listing_time": listing_time.isoformat() if listing_time is not None else None,
            "error": f"missing {', '.join(missing)} file(s)",
        }
    try:
        result = compare_frames(
            pd.read_feather(production_path),
            pd.read_feather(reference_path),
            pair=pair,
            spec=spec,
            window_start=window_start,
            window_end=window_end,
            listing_time=listing_time,
            absolute_tolerance=absolute_tolerance,
            relative_tolerance=relative_tolerance,
            sample_limit=sample_limit,
        )
        return {**base, **result}
    except Exception as exc:
        return {
            **base,
            "status": "read_error",
            "effective_start": effective_start(window_start, listing_time, spec).isoformat(),
            "window_end": window_end.isoformat(),
            "listing_time": listing_time.isoformat() if listing_time is not None else None,
            "error": f"{type(exc).__name__}: {exc}",
        }


def csv_value(value: Any) -> Any:
    if isinstance(value, (dict, list)):
        return json.dumps(value, ensure_ascii=False, sort_keys=True)
    return value


def write_csv(path: Path, rows: list[dict[str, Any]]) -> None:
    with path.open("w", encoding="utf-8", newline="") as handle:
        writer = csv.DictWriter(handle, fieldnames=CSV_FIELDS, extrasaction="ignore")
        writer.writeheader()
        for row in rows:
            writer.writerow({field: csv_value(row.get(field)) for field in CSV_FIELDS})


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--production-data-dir", type=Path, required=True)
    parser.add_argument("--reference-data-dir", type=Path, required=True)
    parser.add_argument("--listing-dates", type=Path, required=True)
    parser.add_argument("--config", type=Path)
    parser.add_argument("--pairs", action="append", default=[])
    parser.add_argument("--start", default=DEFAULT_START)
    parser.add_argument("--end", default=DEFAULT_END)
    parser.add_argument("--absolute-tolerance", type=float, default=0.0)
    parser.add_argument("--relative-tolerance", type=float, default=0.0)
    parser.add_argument("--sample-limit", type=int, default=20)
    parser.add_argument("--output-dir", type=Path, required=True)
    return parser


def main(argv: list[str] | None = None) -> int:
    args = build_parser().parse_args(argv)
    for path, label in (
        (args.production_data_dir, "production data directory"),
        (args.reference_data_dir, "reference data directory"),
    ):
        if not path.is_dir():
            raise SystemExit(f"{label} does not exist: {path}")
    if not args.listing_dates.is_file():
        raise SystemExit(f"listing dates file does not exist: {args.listing_dates}")
    if args.config is not None and not args.config.is_file():
        raise SystemExit(f"config does not exist: {args.config}")
    if args.absolute_tolerance < 0 or args.relative_tolerance < 0:
        raise SystemExit("tolerances must be non-negative")
    if args.sample_limit < 0:
        raise SystemExit("--sample-limit must be non-negative")

    window_start = as_utc(args.start, "--start")
    window_end = as_utc(args.end, "--end")
    if window_start >= window_end:
        raise SystemExit("--start must be earlier than --end")
    pairs = load_pairs(args.config, args.pairs)
    if not pairs:
        raise SystemExit("pair list is empty")
    listing_dates = load_listing_dates(args.listing_dates)
    missing_listing_dates = [pair for pair in pairs if pair not in listing_dates]
    if missing_listing_dates:
        raise SystemExit(
            "listing dates missing for pair(s): " + ", ".join(missing_listing_dates)
        )

    records = [
        compare_file(
            args.production_data_dir,
            args.reference_data_dir,
            pair=pair,
            spec=spec,
            window_start=window_start,
            window_end=window_end,
            listing_time=listing_dates[pair],
            absolute_tolerance=args.absolute_tolerance,
            relative_tolerance=args.relative_tolerance,
            sample_limit=args.sample_limit,
        )
        for pair in pairs
        for spec in DATASETS
    ]
    status_counts = {
        status: sum(record["status"] == status for record in records)
        for status in ("match", "value_mismatch", "structural_mismatch", "missing_file", "read_error")
    }
    manifest = {
        "script_version": SCRIPT_VERSION,
        "script_path": str(Path(__file__).resolve()),
        "script_sha256": sha256_file(Path(__file__).resolve()),
        "generated_at": utc_now(),
        "generated_host": socket.gethostname(),
        "production_data_dir": str(args.production_data_dir.resolve()),
        "reference_data_dir": str(args.reference_data_dir.resolve()),
        "listing_dates_path": str(args.listing_dates.resolve()),
        "listing_dates_sha256": sha256_file(args.listing_dates),
        "window": {
            "start": window_start.isoformat(),
            "end": window_end.isoformat(),
            "interval": "half_open",
        },
        "listing_policy": {
            "intraday": "exclude timestamps before exact listing time",
            "daily": "exclude a partial listing-day candle; begin at first complete UTC day",
        },
        "pairs": pairs,
        "datasets": [spec.key for spec in DATASETS],
        "tolerances": {
            "absolute": args.absolute_tolerance,
            "relative": args.relative_tolerance,
        },
        "environment": {
            "python": sys.version,
            "platform": platform.platform(),
            "pandas": pd.__version__,
            "numpy": np.__version__,
        },
        "summary": {
            "comparison_count": len(records),
            **status_counts,
            "all_match": all(record["status"] == "match" for record in records),
        },
        "comparisons": records,
    }
    args.output_dir.mkdir(parents=True, exist_ok=True)
    json_path = args.output_dir / "research_data_comparison.json"
    csv_path = args.output_dir / "research_data_comparison.csv"
    json_path.write_text(
        json.dumps(manifest, ensure_ascii=False, indent=2, sort_keys=True, allow_nan=False) + "\n",
        encoding="utf-8",
    )
    write_csv(csv_path, records)
    print(json.dumps(manifest["summary"], ensure_ascii=False, sort_keys=True))
    print(f"wrote {json_path}")
    print(f"wrote {csv_path}")
    return 0 if manifest["summary"]["all_match"] else 1


if __name__ == "__main__":
    raise SystemExit(main())
